mirror of
https://github.com/calibrain/shelfmark.git
synced 2026-10-06 07:54:38 +01:00
Backend test hardening + quality enforcement (#872)
- Reworked many tests - Enforcing lint + type checking for test suite - Fixed various issues surfaced by the new tests - CI tweaks
This commit is contained in:
+108
-28
@@ -12,7 +12,7 @@ import time
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
@@ -28,6 +28,8 @@ POLL_INTERVAL = 2
|
||||
DOWNLOAD_TIMEOUT = 300 # 5 minutes max for downloads
|
||||
E2E_USERNAME_ENV = "E2E_USERNAME"
|
||||
E2E_PASSWORD_ENV = "E2E_PASSWORD"
|
||||
TERMINAL_DOWNLOAD_STATES = {"complete", "done", "available", "error", "cancelled"}
|
||||
SUCCESS_DOWNLOAD_STATES = {"complete", "done", "available"}
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -58,6 +60,10 @@ class APIClient:
|
||||
kwargs.setdefault("timeout", self.timeout)
|
||||
return self.session.delete(f"{self.base_url}{path}", **kwargs)
|
||||
|
||||
def close(self) -> None:
|
||||
"""Close the underlying HTTP session."""
|
||||
self.session.close()
|
||||
|
||||
def wait_for_health(self, max_wait: int = 30) -> bool:
|
||||
"""Wait for the server to be healthy."""
|
||||
start = time.time()
|
||||
@@ -66,28 +72,75 @@ class APIClient:
|
||||
resp = self.get("/api/health")
|
||||
if resp.status_code == 200:
|
||||
return True
|
||||
except requests.exceptions.ConnectionError:
|
||||
except requests.exceptions.RequestException:
|
||||
pass
|
||||
time.sleep(1)
|
||||
return False
|
||||
|
||||
|
||||
def _get_auth_state(client: APIClient) -> dict[str, object]:
|
||||
def assert_json_object(response: requests.Response, *, context: str) -> dict[str, Any]:
|
||||
"""Assert that a response is a JSON object."""
|
||||
assert response.status_code == 200, f"{context} failed: {response.status_code} {response.text}"
|
||||
try:
|
||||
data = response.json()
|
||||
except ValueError:
|
||||
pytest.fail(f"{context} did not return valid JSON: {response.text}")
|
||||
assert isinstance(data, dict), f"{context} did not return a JSON object: {data!r}"
|
||||
return data
|
||||
|
||||
|
||||
def assert_json_list(response: requests.Response, *, context: str) -> list[Any]:
|
||||
"""Assert that a response is a JSON list."""
|
||||
assert response.status_code == 200, f"{context} failed: {response.status_code} {response.text}"
|
||||
try:
|
||||
data = response.json()
|
||||
except ValueError:
|
||||
pytest.fail(f"{context} did not return valid JSON: {response.text}")
|
||||
assert isinstance(data, list), f"{context} did not return a JSON list: {data!r}"
|
||||
return data
|
||||
|
||||
|
||||
def assert_queued_download_response(
|
||||
response: requests.Response,
|
||||
*,
|
||||
expected_priority: int = 0,
|
||||
) -> dict[str, Any]:
|
||||
"""Assert the shared queued-download payload shape."""
|
||||
data = assert_json_object(response, context="download queue")
|
||||
assert data == {"status": "queued", "priority": expected_priority}
|
||||
return data
|
||||
|
||||
|
||||
def assert_queue_order_response(response: requests.Response) -> list[dict[str, Any]]:
|
||||
"""Assert the queue order response shape."""
|
||||
data = assert_json_object(response, context="queue order")
|
||||
queue = data.get("queue")
|
||||
assert isinstance(queue, list), f"queue order payload missing queue list: {data!r}"
|
||||
for entry in queue:
|
||||
assert isinstance(entry, dict), f"queue entry is not a JSON object: {entry!r}"
|
||||
assert isinstance(entry.get("id"), str)
|
||||
assert isinstance(entry.get("priority"), int)
|
||||
assert isinstance(entry.get("added_time"), (int, float))
|
||||
assert isinstance(entry.get("status"), str)
|
||||
return queue
|
||||
|
||||
|
||||
def _get_auth_state(client: APIClient) -> dict[str, object] | None:
|
||||
"""Read the live server auth state for auth-sensitive E2E tests."""
|
||||
try:
|
||||
response = client.get("/api/auth/check")
|
||||
except requests.exceptions.RequestException:
|
||||
return {}
|
||||
return None
|
||||
|
||||
if response.status_code != 200:
|
||||
return {}
|
||||
return None
|
||||
|
||||
try:
|
||||
payload = response.json()
|
||||
except ValueError:
|
||||
return {}
|
||||
return None
|
||||
|
||||
return payload if isinstance(payload, dict) else {}
|
||||
return payload if isinstance(payload, dict) else None
|
||||
|
||||
|
||||
def _login_with_env_credentials(client: APIClient) -> bool:
|
||||
@@ -122,19 +175,28 @@ def _is_explicit_e2e_run(markexpr: str, args: list[str]) -> bool:
|
||||
base = arg.split("::", maxsplit=1)[0]
|
||||
parts = Path(base).parts
|
||||
if "tests" not in parts:
|
||||
return False
|
||||
continue
|
||||
|
||||
tests_index = parts.index("tests")
|
||||
if len(parts) <= tests_index + 1 or parts[tests_index + 1] != "e2e":
|
||||
return False
|
||||
if len(parts) > tests_index + 1 and parts[tests_index + 1] == "e2e":
|
||||
return True
|
||||
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _require_authenticated_client(client: APIClient, *, strict: bool) -> APIClient:
|
||||
"""Require an authenticated client for protected-route E2E tests."""
|
||||
auth_state = _get_auth_state(client)
|
||||
if not auth_state or not auth_state.get("auth_required"):
|
||||
if auth_state is None:
|
||||
message = (
|
||||
"Unable to read auth state from the live server. "
|
||||
"Check the stack or run against a reachable instance."
|
||||
)
|
||||
if strict:
|
||||
pytest.fail(message)
|
||||
pytest.skip(message)
|
||||
|
||||
if not auth_state.get("auth_required"):
|
||||
return client
|
||||
|
||||
if auth_state.get("authenticated"):
|
||||
@@ -173,6 +235,21 @@ def _require_authenticated_client(client: APIClient, *, strict: bool) -> APIClie
|
||||
pytest.skip(message)
|
||||
|
||||
|
||||
def _require_healthy_server(base_url: str, *, strict: bool) -> str:
|
||||
"""Require a reachable live server before running live E2E tests."""
|
||||
client = APIClient(base_url=base_url)
|
||||
try:
|
||||
if client.wait_for_health():
|
||||
return base_url
|
||||
finally:
|
||||
client.close()
|
||||
|
||||
message = "Server not available - ensure the app is running"
|
||||
if strict:
|
||||
pytest.fail(message)
|
||||
pytest.skip(message)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DownloadTracker:
|
||||
"""Tracks downloads for cleanup after tests."""
|
||||
@@ -188,7 +265,7 @@ class DownloadTracker:
|
||||
def cleanup(self) -> None:
|
||||
"""Cancel all tracked downloads."""
|
||||
for book_id in self.queued_ids:
|
||||
with suppress(Exception):
|
||||
with suppress(requests.exceptions.RequestException):
|
||||
self.client.delete(f"/api/download/{book_id}/cancel")
|
||||
self.queued_ids.clear()
|
||||
|
||||
@@ -218,23 +295,28 @@ class DownloadTracker:
|
||||
continue
|
||||
|
||||
status_data = resp.json()
|
||||
if not isinstance(status_data, dict):
|
||||
time.sleep(POLL_INTERVAL)
|
||||
continue
|
||||
|
||||
# Check each status category
|
||||
for state in target_states:
|
||||
if state in status_data and book_id in status_data[state]:
|
||||
state_entries = status_data.get(state)
|
||||
if isinstance(state_entries, dict) and book_id in state_entries:
|
||||
return {
|
||||
"state": state,
|
||||
"data": status_data[state][book_id],
|
||||
"data": state_entries[book_id],
|
||||
}
|
||||
|
||||
# Check for error state
|
||||
if "error" in status_data and book_id in status_data["error"]:
|
||||
error_entries = status_data.get("error")
|
||||
if isinstance(error_entries, dict) and book_id in error_entries:
|
||||
return {
|
||||
"state": "error",
|
||||
"data": status_data["error"][book_id],
|
||||
"data": error_entries[book_id],
|
||||
}
|
||||
|
||||
except Exception:
|
||||
except requests.exceptions.RequestException, ValueError, TypeError:
|
||||
pass
|
||||
|
||||
time.sleep(POLL_INTERVAL)
|
||||
@@ -249,15 +331,13 @@ def base_url() -> str:
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def healthy_base_url(base_url: str) -> str:
|
||||
def healthy_base_url(base_url: str, request: pytest.FixtureRequest) -> str:
|
||||
"""Ensure the live server is reachable before creating per-test clients."""
|
||||
client = APIClient(base_url=base_url)
|
||||
try:
|
||||
if not client.wait_for_health():
|
||||
pytest.skip("Server not available - ensure the app is running")
|
||||
return base_url
|
||||
finally:
|
||||
client.session.close()
|
||||
strict = _is_explicit_e2e_run(
|
||||
getattr(request.config.option, "markexpr", ""),
|
||||
list(getattr(request.config, "args", [])),
|
||||
)
|
||||
return _require_healthy_server(base_url, strict=strict)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -265,7 +345,7 @@ def api_client(healthy_base_url: str) -> Iterator[APIClient]:
|
||||
"""Create a fresh API client for each E2E test."""
|
||||
client = APIClient(base_url=healthy_base_url)
|
||||
yield client
|
||||
client.session.close()
|
||||
client.close()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -300,4 +380,4 @@ def server_config(healthy_base_url: str) -> dict:
|
||||
return {}
|
||||
return resp.json()
|
||||
finally:
|
||||
client.session.close()
|
||||
client.close()
|
||||
|
||||
+155
-110
@@ -8,7 +8,19 @@ Run with: uv run pytest tests/e2e/ -v -m e2e
|
||||
|
||||
import pytest
|
||||
|
||||
from .conftest import APIClient, DownloadTracker
|
||||
from .conftest import (
|
||||
APIClient,
|
||||
DownloadTracker,
|
||||
assert_queue_order_response,
|
||||
assert_queued_download_response,
|
||||
)
|
||||
|
||||
|
||||
def _assert_json_object(response, *, status_code: int = 200) -> dict:
|
||||
assert response.status_code == status_code
|
||||
data = response.json()
|
||||
assert isinstance(data, dict)
|
||||
return data
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -37,26 +49,23 @@ class TestConfigEndpoint:
|
||||
"""Tests for the configuration endpoint."""
|
||||
|
||||
def test_config_returns_expected_fields(self, protected_api_client: APIClient):
|
||||
"""Test that config includes expected configuration fields."""
|
||||
resp = protected_api_client.get("/api/config")
|
||||
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
# Config should be a dict with various settings
|
||||
assert isinstance(data, dict)
|
||||
# Should have some standard config fields
|
||||
assert "supported_formats" in data or "book_languages" in data
|
||||
"""Test that config exposes the stable frontend contract."""
|
||||
data = _assert_json_object(protected_api_client.get("/api/config"))
|
||||
assert isinstance(data["supported_formats"], list)
|
||||
assert isinstance(data["supported_audiobook_formats"], list)
|
||||
assert isinstance(data["book_languages"], list)
|
||||
assert isinstance(data["settings_enabled"], bool)
|
||||
assert isinstance(data["onboarding_complete"], bool)
|
||||
assert isinstance(data["search_mode"], str)
|
||||
assert isinstance(data["default_release_source"], str)
|
||||
|
||||
def test_config_returns_supported_formats(self, protected_api_client: APIClient):
|
||||
"""Test that config includes supported formats."""
|
||||
resp = protected_api_client.get("/api/config")
|
||||
|
||||
data = resp.json()
|
||||
assert "supported_formats" in data
|
||||
assert isinstance(data["supported_formats"], list)
|
||||
# Should include common ebook formats
|
||||
data = _assert_json_object(protected_api_client.get("/api/config"))
|
||||
formats = data["supported_formats"]
|
||||
assert "epub" in formats or "EPUB" in [f.upper() for f in formats]
|
||||
assert formats
|
||||
assert all(isinstance(fmt, str) for fmt in formats)
|
||||
assert "epub" in {fmt.lower() for fmt in formats}
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -73,12 +82,22 @@ class TestReleaseSourcesEndpoint:
|
||||
|
||||
def test_release_sources_have_required_fields(self, protected_api_client: APIClient):
|
||||
"""Test that each release source has required fields."""
|
||||
resp = protected_api_client.get("/api/release-sources")
|
||||
|
||||
data = resp.json()
|
||||
data = protected_api_client.get("/api/release-sources").json()
|
||||
for source in data:
|
||||
assert "name" in source
|
||||
assert "display_name" in source or "label" in source
|
||||
assert set(source) == {
|
||||
"name",
|
||||
"display_name",
|
||||
"enabled",
|
||||
"supported_content_types",
|
||||
"browse_results_are_releases",
|
||||
"can_be_default",
|
||||
}
|
||||
assert isinstance(source["name"], str)
|
||||
assert isinstance(source["display_name"], str)
|
||||
assert isinstance(source["enabled"], bool)
|
||||
assert isinstance(source["supported_content_types"], list)
|
||||
assert isinstance(source["browse_results_are_releases"], bool)
|
||||
assert isinstance(source["can_be_default"], bool)
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -86,29 +105,32 @@ class TestMetadataProvidersEndpoint:
|
||||
"""Tests for the metadata providers endpoint."""
|
||||
|
||||
def test_providers_returns_data(self, protected_api_client: APIClient):
|
||||
"""Test that providers endpoint returns provider data."""
|
||||
resp = protected_api_client.get("/api/metadata/providers")
|
||||
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
# May be list or dict depending on implementation
|
||||
assert isinstance(data, (list, dict))
|
||||
"""Test that providers endpoint returns the documented object contract."""
|
||||
data = _assert_json_object(protected_api_client.get("/api/metadata/providers"))
|
||||
assert set(data) == {
|
||||
"providers",
|
||||
"configured_provider",
|
||||
"configured_provider_audiobook",
|
||||
"configured_provider_combined",
|
||||
}
|
||||
assert isinstance(data["providers"], list)
|
||||
|
||||
def test_providers_have_required_fields(self, protected_api_client: APIClient):
|
||||
"""Test that each provider has required fields."""
|
||||
resp = protected_api_client.get("/api/metadata/providers")
|
||||
|
||||
data = resp.json()
|
||||
# Handle both list and dict formats
|
||||
if isinstance(data, dict):
|
||||
providers = list(data.values()) if data else []
|
||||
else:
|
||||
providers = data
|
||||
|
||||
for provider in providers:
|
||||
if isinstance(provider, dict):
|
||||
# Should have name or be identifiable
|
||||
assert "name" in provider or "id" in provider or "label" in provider
|
||||
data = _assert_json_object(protected_api_client.get("/api/metadata/providers"))
|
||||
for provider in data["providers"]:
|
||||
assert set(provider) == {
|
||||
"name",
|
||||
"display_name",
|
||||
"requires_auth",
|
||||
"enabled",
|
||||
"available",
|
||||
}
|
||||
assert isinstance(provider["name"], str)
|
||||
assert isinstance(provider["display_name"], str)
|
||||
assert isinstance(provider["requires_auth"], bool)
|
||||
assert isinstance(provider["enabled"], bool)
|
||||
assert isinstance(provider["available"], bool)
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -119,57 +141,56 @@ class TestMetadataSearch:
|
||||
"""Test that search requires a query parameter."""
|
||||
resp = protected_api_client.get("/api/metadata/search")
|
||||
|
||||
# Should return error for missing query
|
||||
assert resp.status_code in [400, 422]
|
||||
assert resp.status_code == 400
|
||||
assert resp.json() == {"error": "Either 'query' or search field values are required"}
|
||||
|
||||
def test_search_returns_results(self, protected_api_client: APIClient):
|
||||
"""Test that search returns results for a known book."""
|
||||
resp = protected_api_client.get("/api/metadata/search", params={"query": "1984 Orwell"})
|
||||
|
||||
# May return 200 with results or 503 if provider unavailable
|
||||
if resp.status_code == 200:
|
||||
data = _assert_json_object(resp)
|
||||
assert isinstance(data["books"], list)
|
||||
assert isinstance(data["provider"], str)
|
||||
assert data["query"] == "1984 Orwell"
|
||||
assert isinstance(data["page"], int)
|
||||
assert isinstance(data["total_found"], int)
|
||||
assert isinstance(data["has_more"], bool)
|
||||
else:
|
||||
assert resp.status_code == 503
|
||||
data = resp.json()
|
||||
# Response may be list directly, or dict with results key
|
||||
assert "results" in data or isinstance(data, list) or "query" in data
|
||||
assert isinstance(data, dict)
|
||||
assert "error" in data
|
||||
assert "message" in data
|
||||
|
||||
def test_search_with_provider_filter(self, protected_api_client: APIClient):
|
||||
"""Test searching with a specific provider."""
|
||||
# Get available providers first
|
||||
providers_resp = protected_api_client.get("/api/metadata/providers")
|
||||
if providers_resp.status_code != 200:
|
||||
pytest.skip("Could not get providers")
|
||||
|
||||
providers_data = providers_resp.json()
|
||||
if not providers_data:
|
||||
providers = providers_data.get("providers", [])
|
||||
if not providers:
|
||||
pytest.skip("No providers available")
|
||||
|
||||
# Handle both list and dict formats
|
||||
if isinstance(providers_data, dict):
|
||||
# Dict format: get first provider name from keys or values
|
||||
if providers_data:
|
||||
first_key = list(providers_data.keys())[0]
|
||||
provider_info = providers_data[first_key]
|
||||
provider_name = (
|
||||
provider_info.get("name", first_key)
|
||||
if isinstance(provider_info, dict)
|
||||
else first_key
|
||||
)
|
||||
else:
|
||||
pytest.skip("No providers available")
|
||||
else:
|
||||
# List format
|
||||
provider_name = providers_data[0].get("name") if providers_data else None
|
||||
|
||||
if not provider_name:
|
||||
pytest.skip("Could not determine provider name")
|
||||
provider_name = providers[0]["name"]
|
||||
|
||||
resp = protected_api_client.get(
|
||||
"/api/metadata/search",
|
||||
params={"query": "Moby Dick", "provider": provider_name},
|
||||
)
|
||||
|
||||
# Should return 200 or 503 (provider unavailable)
|
||||
assert resp.status_code in [200, 503]
|
||||
if resp.status_code == 200:
|
||||
data = _assert_json_object(resp)
|
||||
assert data["provider"] == provider_name
|
||||
assert data["query"] == "Moby Dick"
|
||||
assert isinstance(data["books"], list)
|
||||
else:
|
||||
assert resp.status_code == 503
|
||||
data = resp.json()
|
||||
assert isinstance(data, dict)
|
||||
assert "error" in data
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -182,8 +203,10 @@ class TestStatusEndpoint:
|
||||
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
# Should have standard status categories
|
||||
assert isinstance(data, dict)
|
||||
for status_name, tasks in data.items():
|
||||
assert isinstance(status_name, str)
|
||||
assert isinstance(tasks, dict)
|
||||
|
||||
def test_active_downloads_endpoint(self, protected_api_client: APIClient):
|
||||
"""Test the active downloads endpoint."""
|
||||
@@ -191,7 +214,8 @@ class TestStatusEndpoint:
|
||||
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert isinstance(data, (list, dict))
|
||||
assert data == {"active_downloads": data["active_downloads"]}
|
||||
assert isinstance(data["active_downloads"], list)
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -202,14 +226,8 @@ class TestQueueEndpoint:
|
||||
"""Test that queue order endpoint returns queue data."""
|
||||
resp = protected_api_client.get("/api/queue/order")
|
||||
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
# May return list directly or dict with queue key
|
||||
if isinstance(data, dict):
|
||||
assert "queue" in data
|
||||
assert isinstance(data["queue"], list)
|
||||
else:
|
||||
assert isinstance(data, list)
|
||||
queue = assert_queue_order_response(resp)
|
||||
assert isinstance(queue, list)
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -226,7 +244,16 @@ class TestSettingsEndpoint:
|
||||
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert isinstance(data, (list, dict))
|
||||
assert data == {"tabs": data["tabs"], "groups": data["groups"]}
|
||||
assert isinstance(data["tabs"], list)
|
||||
assert isinstance(data["groups"], list)
|
||||
for tab in data["tabs"]:
|
||||
assert isinstance(tab, dict)
|
||||
assert "name" in tab
|
||||
assert "fields" in tab
|
||||
for group in data["groups"]:
|
||||
assert isinstance(group, dict)
|
||||
assert "name" in group
|
||||
|
||||
def test_get_specific_settings_tab(self, protected_api_client: APIClient):
|
||||
"""Test getting a specific settings tab."""
|
||||
@@ -236,20 +263,20 @@ class TestSettingsEndpoint:
|
||||
pytest.skip("Settings disabled")
|
||||
|
||||
data = resp.json()
|
||||
if not data:
|
||||
tabs = data.get("tabs", []) if isinstance(data, dict) else []
|
||||
if not tabs:
|
||||
pytest.skip("No settings tabs available")
|
||||
|
||||
# Get the first tab
|
||||
if isinstance(data, list):
|
||||
tab_name = data[0].get("name") or data[0].get("id")
|
||||
else:
|
||||
tab_name = list(data.keys())[0] if data else None
|
||||
|
||||
tab_name = tabs[0].get("name")
|
||||
if not tab_name:
|
||||
pytest.skip("Could not determine tab name")
|
||||
|
||||
resp = protected_api_client.get(f"/api/settings/{tab_name}")
|
||||
assert resp.status_code in [200, 404]
|
||||
assert resp.status_code == 200
|
||||
tab_data = resp.json()
|
||||
assert isinstance(tab_data, dict)
|
||||
assert tab_data.get("name") == tab_name
|
||||
assert isinstance(tab_data.get("fields"), list)
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -260,8 +287,9 @@ class TestDownloadFlow:
|
||||
"""Test cancelling a download that doesn't exist."""
|
||||
resp = protected_api_client.delete("/api/download/nonexistent-id-xyz/cancel")
|
||||
|
||||
# Should handle gracefully (may return 200, 204, or 404)
|
||||
assert resp.status_code in [200, 204, 404]
|
||||
assert resp.status_code == 404
|
||||
data = resp.json()
|
||||
assert data.get("error") == "Failed to cancel download or book not found"
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -270,11 +298,14 @@ class TestReleaseDownloadFlow:
|
||||
|
||||
def test_release_download_requires_source_id(self, protected_api_client: APIClient):
|
||||
"""Test that release download requires source_id."""
|
||||
resp = protected_api_client.post("/api/releases/download", json={})
|
||||
resp = protected_api_client.post(
|
||||
"/api/releases/download",
|
||||
json={"source": "test_source"},
|
||||
)
|
||||
|
||||
assert resp.status_code == 400
|
||||
data = resp.json()
|
||||
assert "error" in data
|
||||
assert data == {"error": "source_id is required"}
|
||||
|
||||
def test_release_download_with_minimal_data(
|
||||
self, protected_api_client: APIClient, download_tracker: DownloadTracker
|
||||
@@ -291,10 +322,8 @@ class TestReleaseDownloadFlow:
|
||||
},
|
||||
)
|
||||
|
||||
if resp.status_code == 200:
|
||||
download_tracker.track(test_id)
|
||||
data = resp.json()
|
||||
assert data.get("status") == "queued"
|
||||
download_tracker.track(test_id)
|
||||
assert_queued_download_response(resp)
|
||||
|
||||
def test_cancel_release_with_slash_id(
|
||||
self, protected_api_client: APIClient, download_tracker: DownloadTracker
|
||||
@@ -315,9 +344,11 @@ class TestReleaseDownloadFlow:
|
||||
pytest.skip("Release download endpoint not available")
|
||||
|
||||
download_tracker.track(test_id)
|
||||
assert resp.json() == {"status": "queued", "priority": 0}
|
||||
|
||||
cancel_resp = protected_api_client.delete(f"/api/download/{test_id}/cancel")
|
||||
assert cancel_resp.status_code in [200, 204]
|
||||
assert cancel_resp.status_code == 200
|
||||
assert cancel_resp.json() == {"status": "cancelled", "book_id": test_id}
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -329,8 +360,7 @@ class TestReleasesSearch:
|
||||
resp = protected_api_client.get("/api/releases")
|
||||
|
||||
assert resp.status_code == 400
|
||||
data = resp.json()
|
||||
assert "error" in data
|
||||
assert resp.json() == {"error": "Parameters 'provider' and 'book_id' are required"}
|
||||
|
||||
def test_releases_with_invalid_provider(self, protected_api_client: APIClient):
|
||||
"""Test releases with invalid provider."""
|
||||
@@ -340,8 +370,7 @@ class TestReleasesSearch:
|
||||
)
|
||||
|
||||
assert resp.status_code == 400
|
||||
data = resp.json()
|
||||
assert "error" in data
|
||||
assert resp.json() == {"error": "Unknown metadata provider: nonexistent_provider"}
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -352,8 +381,11 @@ class TestCoverProxy:
|
||||
"""Test that cover endpoint without URL returns error."""
|
||||
resp = protected_api_client.get("/api/covers/test-id")
|
||||
|
||||
# Should return error for missing URL
|
||||
assert resp.status_code in [400, 404]
|
||||
assert resp.status_code == 404
|
||||
assert resp.json() in [
|
||||
{"error": "Cover caching is disabled"},
|
||||
{"error": "Cover URL not provided"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -364,7 +396,9 @@ class TestDirectSourceQueryEndpoint:
|
||||
"""Source query mode requires a query or browse filters."""
|
||||
resp = protected_api_client.get("/api/releases", params={"source": "direct_download"})
|
||||
|
||||
assert resp.status_code in [400, 422]
|
||||
assert resp.status_code == 400
|
||||
data = resp.json()
|
||||
assert data == {"error": "Parameters 'provider' and 'book_id' are required"}
|
||||
|
||||
def test_direct_source_query_returns_results(self, protected_api_client: APIClient):
|
||||
"""Direct mode uses /api/releases source query mode."""
|
||||
@@ -373,11 +407,20 @@ class TestDirectSourceQueryEndpoint:
|
||||
params={"source": "direct_download", "query": "Pride Prejudice"},
|
||||
)
|
||||
|
||||
# May return results or 503 if source unavailable
|
||||
if resp.status_code == 200:
|
||||
data = resp.json()
|
||||
assert data.get("sources_searched") == ["direct_download"]
|
||||
assert isinstance(data.get("releases"), list)
|
||||
expected_keys = {"releases", "book", "sources_searched", "column_config", "search_info"}
|
||||
assert expected_keys <= set(data)
|
||||
assert data["sources_searched"] == ["direct_download"]
|
||||
assert isinstance(data["releases"], list)
|
||||
assert isinstance(data["book"], dict)
|
||||
assert isinstance(data["search_info"], dict)
|
||||
if "errors" in data:
|
||||
assert isinstance(data["errors"], list)
|
||||
else:
|
||||
assert resp.status_code == 503
|
||||
data = resp.json()
|
||||
assert "error" in data
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -390,5 +433,7 @@ class TestSourceRecordEndpoint:
|
||||
"/api/release-sources/direct_download/records/invalid-id-xyz"
|
||||
)
|
||||
|
||||
# Should return 404 or error
|
||||
assert resp.status_code in [404, 500, 503]
|
||||
if resp.status_code == 503:
|
||||
pytest.skip("Direct source record lookup unavailable")
|
||||
assert resp.status_code == 404
|
||||
assert resp.json() == {"error": "Record not found"}
|
||||
|
||||
+194
-131
@@ -9,11 +9,13 @@ from __future__ import annotations
|
||||
import importlib
|
||||
import sqlite3
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Any, Tuple
|
||||
from typing import Any
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
||||
def _as_response(result: Any):
|
||||
"""Normalize Flask view return values to a Response-like object."""
|
||||
@@ -44,31 +46,51 @@ def main_module():
|
||||
|
||||
class TestGetAuthMode:
|
||||
def test_get_auth_mode_none(self, main_module):
|
||||
with patch.object(main_module.app_config, "get", side_effect=_config_getter({"AUTH_METHOD": "none"})):
|
||||
with patch.object(
|
||||
main_module.app_config, "get", side_effect=_config_getter({"AUTH_METHOD": "none"})
|
||||
):
|
||||
assert main_module.get_auth_mode() == "none"
|
||||
|
||||
def test_get_auth_mode_builtin(self, main_module):
|
||||
with patch.object(main_module.app_config, "get", side_effect=_config_getter({"AUTH_METHOD": "builtin"})):
|
||||
with patch("shelfmark.core.auth_modes.has_local_password_admin", return_value=True):
|
||||
assert main_module.get_auth_mode() == "builtin"
|
||||
with (
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({"AUTH_METHOD": "builtin"}),
|
||||
),
|
||||
patch("shelfmark.core.auth_modes.has_local_password_admin", return_value=True),
|
||||
):
|
||||
assert main_module.get_auth_mode() == "builtin"
|
||||
|
||||
def test_get_auth_mode_builtin_without_local_admin_falls_back_to_none(self, main_module):
|
||||
with patch.object(main_module.app_config, "get", side_effect=_config_getter({"AUTH_METHOD": "builtin"})):
|
||||
with patch("shelfmark.core.auth_modes.has_local_password_admin", return_value=False):
|
||||
assert main_module.get_auth_mode() == "none"
|
||||
with (
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({"AUTH_METHOD": "builtin"}),
|
||||
),
|
||||
patch("shelfmark.core.auth_modes.has_local_password_admin", return_value=False),
|
||||
):
|
||||
assert main_module.get_auth_mode() == "none"
|
||||
|
||||
def test_get_auth_mode_proxy(self, main_module):
|
||||
with patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({"AUTH_METHOD": "proxy", "PROXY_AUTH_USER_HEADER": "X-Auth-User"}),
|
||||
side_effect=_config_getter(
|
||||
{"AUTH_METHOD": "proxy", "PROXY_AUTH_USER_HEADER": "X-Auth-User"}
|
||||
),
|
||||
):
|
||||
assert main_module.get_auth_mode() == "proxy"
|
||||
|
||||
def test_get_auth_mode_cwa(self, main_module):
|
||||
with patch.object(main_module.app_config, "get", side_effect=_config_getter({"AUTH_METHOD": "cwa"})):
|
||||
with patch.object(main_module, "CWA_DB_PATH", object()):
|
||||
assert main_module.get_auth_mode() == "cwa"
|
||||
with (
|
||||
patch.object(
|
||||
main_module.app_config, "get", side_effect=_config_getter({"AUTH_METHOD": "cwa"})
|
||||
),
|
||||
patch.object(main_module, "CWA_DB_PATH", object()),
|
||||
):
|
||||
assert main_module.get_auth_mode() == "cwa"
|
||||
|
||||
def test_get_auth_mode_default_on_error(self, main_module):
|
||||
with patch.object(main_module.app_config, "get", side_effect=RuntimeError("boom")):
|
||||
@@ -77,10 +99,12 @@ class TestGetAuthMode:
|
||||
|
||||
class TestAuthCheckEndpoint:
|
||||
def test_auth_check_no_auth(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="none"):
|
||||
with main_module.app.test_request_context("/api/auth/check"):
|
||||
resp = _as_response(main_module.api_auth_check())
|
||||
data = resp.get_json()
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="none"),
|
||||
main_module.app.test_request_context("/api/auth/check"),
|
||||
):
|
||||
resp = _as_response(main_module.api_auth_check())
|
||||
data = resp.get_json()
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert data == {
|
||||
@@ -91,85 +115,108 @@ class TestAuthCheckEndpoint:
|
||||
}
|
||||
|
||||
def test_auth_check_builtin_not_authenticated(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="builtin"):
|
||||
with main_module.app.test_request_context("/api/auth/check"):
|
||||
resp = _as_response(main_module.api_auth_check())
|
||||
data = resp.get_json()
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="builtin"),
|
||||
main_module.app.test_request_context("/api/auth/check"),
|
||||
):
|
||||
resp = _as_response(main_module.api_auth_check())
|
||||
data = resp.get_json()
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert data["authenticated"] is False
|
||||
assert data["auth_required"] is True
|
||||
assert data["auth_mode"] == "builtin"
|
||||
assert data["is_admin"] is False
|
||||
assert data["username"] is None
|
||||
assert data == {
|
||||
"authenticated": False,
|
||||
"auth_required": True,
|
||||
"auth_mode": "builtin",
|
||||
"is_admin": False,
|
||||
"username": None,
|
||||
"display_name": None,
|
||||
}
|
||||
|
||||
def test_auth_check_builtin_authenticated(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="builtin"):
|
||||
with main_module.app.test_request_context("/api/auth/check"):
|
||||
main_module.session["user_id"] = "admin"
|
||||
main_module.session["is_admin"] = True
|
||||
resp = _as_response(main_module.api_auth_check())
|
||||
data = resp.get_json()
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="builtin"),
|
||||
main_module.app.test_request_context("/api/auth/check"),
|
||||
):
|
||||
main_module.session["user_id"] = "admin"
|
||||
main_module.session["is_admin"] = True
|
||||
resp = _as_response(main_module.api_auth_check())
|
||||
data = resp.get_json()
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert data["authenticated"] is True
|
||||
assert data["auth_required"] is True
|
||||
assert data["auth_mode"] == "builtin"
|
||||
assert data["is_admin"] is True
|
||||
assert data["username"] == "admin"
|
||||
assert data == {
|
||||
"authenticated": True,
|
||||
"auth_required": True,
|
||||
"auth_mode": "builtin",
|
||||
"is_admin": True,
|
||||
"username": "admin",
|
||||
"display_name": None,
|
||||
}
|
||||
|
||||
def test_auth_check_proxy_includes_logout_url(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with patch.object(
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({
|
||||
"PROXY_AUTH_USER_HEADER": "X-Auth-User",
|
||||
"PROXY_AUTH_LOGOUT_URL": "https://auth.example.com/logout",
|
||||
}),
|
||||
):
|
||||
with main_module.app.test_request_context("/api/auth/check"):
|
||||
main_module.session["user_id"] = "proxyuser"
|
||||
main_module.session["is_admin"] = True
|
||||
resp = _as_response(main_module.api_auth_check())
|
||||
data = resp.get_json()
|
||||
side_effect=_config_getter(
|
||||
{
|
||||
"PROXY_AUTH_USER_HEADER": "X-Auth-User",
|
||||
"PROXY_AUTH_LOGOUT_URL": "https://auth.example.com/logout",
|
||||
}
|
||||
),
|
||||
),
|
||||
main_module.app.test_request_context("/api/auth/check"),
|
||||
):
|
||||
main_module.session["user_id"] = "proxyuser"
|
||||
main_module.session["is_admin"] = True
|
||||
resp = _as_response(main_module.api_auth_check())
|
||||
data = resp.get_json()
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert data["authenticated"] is True
|
||||
assert data["auth_mode"] == "proxy"
|
||||
assert data["username"] == "proxyuser"
|
||||
assert data["logout_url"] == "https://auth.example.com/logout"
|
||||
assert data == {
|
||||
"authenticated": True,
|
||||
"auth_required": True,
|
||||
"auth_mode": "proxy",
|
||||
"is_admin": True,
|
||||
"username": "proxyuser",
|
||||
"display_name": None,
|
||||
"logout_url": "https://auth.example.com/logout",
|
||||
}
|
||||
|
||||
|
||||
class TestLoginEndpoint:
|
||||
def test_login_proxy_mode_disabled(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with main_module.app.test_request_context(
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
main_module.app.test_request_context(
|
||||
"/api/auth/login",
|
||||
method="POST",
|
||||
json={"anything": "x"},
|
||||
):
|
||||
resp = _as_response(main_module.api_login())
|
||||
data = resp.get_json()
|
||||
),
|
||||
):
|
||||
resp = _as_response(main_module.api_login())
|
||||
data = resp.get_json()
|
||||
|
||||
assert resp.status_code == 401
|
||||
assert "Proxy authentication" in (data.get("error") or "")
|
||||
assert data == {"error": "Proxy authentication is enabled"}
|
||||
|
||||
def test_login_no_auth_success(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="none"):
|
||||
with patch.object(main_module, "is_account_locked", return_value=False):
|
||||
with main_module.app.test_request_context(
|
||||
"/api/auth/login",
|
||||
method="POST",
|
||||
json={"username": "anyuser", "password": "anypass", "remember_me": True},
|
||||
):
|
||||
resp = _as_response(main_module.api_login())
|
||||
data = resp.get_json()
|
||||
assert main_module.session.get("user_id") == "anyuser"
|
||||
assert main_module.session.permanent is True
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="none"),
|
||||
patch.object(main_module, "is_account_locked", return_value=False),
|
||||
main_module.app.test_request_context(
|
||||
"/api/auth/login",
|
||||
method="POST",
|
||||
json={"username": "anyuser", "password": "anypass", "remember_me": True},
|
||||
),
|
||||
):
|
||||
resp = _as_response(main_module.api_login())
|
||||
data = resp.get_json()
|
||||
assert main_module.session.get("user_id") == "anyuser"
|
||||
assert main_module.session.permanent is True
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert data.get("success") is True
|
||||
assert data == {"success": True}
|
||||
|
||||
def test_login_builtin_success(self, main_module):
|
||||
mock_user_db = Mock()
|
||||
@@ -179,21 +226,23 @@ class TestLoginEndpoint:
|
||||
"password_hash": "hash",
|
||||
"role": "admin",
|
||||
}
|
||||
with patch.object(main_module, "get_auth_mode", return_value="builtin"):
|
||||
with patch.object(main_module, "is_account_locked", return_value=False):
|
||||
with patch.object(main_module, "user_db", mock_user_db):
|
||||
with patch.object(main_module, "check_password_hash", return_value=True):
|
||||
with main_module.app.test_request_context(
|
||||
"/api/auth/login",
|
||||
method="POST",
|
||||
json={"username": "admin", "password": "correct", "remember_me": False},
|
||||
):
|
||||
resp = _as_response(main_module.api_login())
|
||||
data = resp.get_json()
|
||||
assert main_module.session.get("user_id") == "admin"
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="builtin"),
|
||||
patch.object(main_module, "is_account_locked", return_value=False),
|
||||
patch.object(main_module, "user_db", mock_user_db),
|
||||
patch.object(main_module, "check_password_hash", return_value=True),
|
||||
main_module.app.test_request_context(
|
||||
"/api/auth/login",
|
||||
method="POST",
|
||||
json={"username": "admin", "password": "correct", "remember_me": False},
|
||||
),
|
||||
):
|
||||
resp = _as_response(main_module.api_login())
|
||||
data = resp.get_json()
|
||||
assert main_module.session.get("user_id") == "admin"
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert data.get("success") is True
|
||||
assert data == {"success": True}
|
||||
|
||||
def test_login_cwa_provisions_db_user(self, main_module, tmp_path):
|
||||
cwa_db_path = tmp_path / "app.db"
|
||||
@@ -210,23 +259,25 @@ class TestLoginEndpoint:
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
with patch.object(main_module, "get_auth_mode", return_value="cwa"):
|
||||
with patch.object(main_module, "is_account_locked", return_value=False):
|
||||
with patch.object(main_module, "CWA_DB_PATH", cwa_db_path):
|
||||
with patch.object(main_module, "check_password_hash", return_value=True):
|
||||
with main_module.app.test_request_context(
|
||||
"/api/auth/login",
|
||||
method="POST",
|
||||
json={"username": username, "password": "correct", "remember_me": False},
|
||||
):
|
||||
resp = _as_response(main_module.api_login())
|
||||
data = resp.get_json()
|
||||
assert main_module.session.get("user_id") == username
|
||||
assert main_module.session.get("is_admin") is True
|
||||
assert main_module.session.get("db_user_id") is not None
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="cwa"),
|
||||
patch.object(main_module, "is_account_locked", return_value=False),
|
||||
patch.object(main_module, "CWA_DB_PATH", cwa_db_path),
|
||||
patch.object(main_module, "check_password_hash", return_value=True),
|
||||
main_module.app.test_request_context(
|
||||
"/api/auth/login",
|
||||
method="POST",
|
||||
json={"username": username, "password": "correct", "remember_me": False},
|
||||
),
|
||||
):
|
||||
resp = _as_response(main_module.api_login())
|
||||
data = resp.get_json()
|
||||
assert main_module.session.get("user_id") == username
|
||||
assert main_module.session.get("is_admin") is True
|
||||
assert main_module.session.get("db_user_id") is not None
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert data.get("success") is True
|
||||
assert data == {"success": True}
|
||||
db_user = main_module.user_db.get_user(username=username)
|
||||
assert db_user["email"] == "cwa@example.com"
|
||||
assert db_user["role"] == "admin"
|
||||
@@ -255,22 +306,24 @@ class TestLoginEndpoint:
|
||||
conn.commit()
|
||||
conn.close()
|
||||
|
||||
with patch.object(main_module, "get_auth_mode", return_value="cwa"):
|
||||
with patch.object(main_module, "is_account_locked", return_value=False):
|
||||
with patch.object(main_module, "CWA_DB_PATH", cwa_db_path):
|
||||
with patch.object(main_module, "check_password_hash", return_value=True):
|
||||
with main_module.app.test_request_context(
|
||||
"/api/auth/login",
|
||||
method="POST",
|
||||
json={"username": username, "password": "correct", "remember_me": False},
|
||||
):
|
||||
resp = _as_response(main_module.api_login())
|
||||
data = resp.get_json()
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="cwa"),
|
||||
patch.object(main_module, "is_account_locked", return_value=False),
|
||||
patch.object(main_module, "CWA_DB_PATH", cwa_db_path),
|
||||
patch.object(main_module, "check_password_hash", return_value=True),
|
||||
main_module.app.test_request_context(
|
||||
"/api/auth/login",
|
||||
method="POST",
|
||||
json={"username": username, "password": "correct", "remember_me": False},
|
||||
),
|
||||
):
|
||||
resp = _as_response(main_module.api_login())
|
||||
data = resp.get_json()
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert data.get("success") is True
|
||||
assert main_module.session.get("user_id") == username
|
||||
assert main_module.session.get("db_user_id") is not None
|
||||
assert resp.status_code == 200
|
||||
assert data == {"success": True}
|
||||
assert main_module.session.get("user_id") == username
|
||||
assert main_module.session.get("db_user_id") is not None
|
||||
|
||||
local_after = main_module.user_db.get_user(user_id=local_user["id"])
|
||||
assert local_after is not None
|
||||
@@ -278,7 +331,8 @@ class TestLoginEndpoint:
|
||||
assert local_after["email"] == "collision.local@example.com"
|
||||
|
||||
provisioned_cwa_user = next(
|
||||
user for user in main_module.user_db.list_users()
|
||||
user
|
||||
for user in main_module.user_db.list_users()
|
||||
if user.get("auth_source") == "cwa" and user.get("email") == external_email
|
||||
)
|
||||
assert provisioned_cwa_user["username"].startswith(f"{username}__cwa")
|
||||
@@ -286,31 +340,40 @@ class TestLoginEndpoint:
|
||||
|
||||
class TestLogoutEndpoint:
|
||||
def test_logout_proxy_returns_logout_url(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with patch.object(
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({"PROXY_AUTH_LOGOUT_URL": "https://auth.example.com/logout"}),
|
||||
):
|
||||
with main_module.app.test_request_context("/api/auth/logout", method="POST"):
|
||||
main_module.session["user_id"] = "proxyuser"
|
||||
resp = _as_response(main_module.api_logout())
|
||||
data = resp.get_json()
|
||||
side_effect=_config_getter(
|
||||
{"PROXY_AUTH_LOGOUT_URL": "https://auth.example.com/logout"}
|
||||
),
|
||||
),
|
||||
main_module.app.test_request_context("/api/auth/logout", method="POST"),
|
||||
):
|
||||
main_module.session["user_id"] = "proxyuser"
|
||||
resp = _as_response(main_module.api_logout())
|
||||
data = resp.get_json()
|
||||
assert "user_id" not in main_module.session
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert data["success"] is True
|
||||
assert data["logout_url"] == "https://auth.example.com/logout"
|
||||
assert data == {
|
||||
"success": True,
|
||||
"logout_url": "https://auth.example.com/logout",
|
||||
}
|
||||
|
||||
def test_logout_basic(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="builtin"):
|
||||
with main_module.app.test_request_context("/api/auth/logout", method="POST"):
|
||||
main_module.session["user_id"] = "admin"
|
||||
resp = _as_response(main_module.api_logout())
|
||||
data = resp.get_json()
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="builtin"),
|
||||
main_module.app.test_request_context("/api/auth/logout", method="POST"),
|
||||
):
|
||||
main_module.session["user_id"] = "admin"
|
||||
resp = _as_response(main_module.api_logout())
|
||||
data = resp.get_json()
|
||||
assert "user_id" not in main_module.session
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert data["success"] is True
|
||||
assert "logout_url" not in data
|
||||
assert data == {"success": True}
|
||||
|
||||
|
||||
class TestRateLimiting:
|
||||
|
||||
+130
-132
@@ -7,9 +7,26 @@ with various authentication modes.
|
||||
Run with: uv run pytest tests/e2e/ -v -m e2e
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
|
||||
from .conftest import APIClient
|
||||
if TYPE_CHECKING:
|
||||
from .conftest import APIClient
|
||||
|
||||
|
||||
def _auth_check(api_client: APIClient) -> dict:
|
||||
"""Fetch the current auth state and assert the response shape."""
|
||||
resp = api_client.get("/api/auth/check")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert isinstance(data, dict)
|
||||
assert "authenticated" in data
|
||||
assert "auth_required" in data
|
||||
assert "auth_mode" in data
|
||||
assert "is_admin" in data
|
||||
return data
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -17,76 +34,104 @@ class TestAuthenticationFlow:
|
||||
"""Tests for the authentication endpoints in a real environment."""
|
||||
|
||||
def test_auth_check_endpoint_exists(self, api_client: APIClient):
|
||||
"""Test that auth check endpoint is accessible."""
|
||||
resp = api_client.get("/api/auth/check")
|
||||
"""Test that auth check returns the stable contract fields."""
|
||||
data = _auth_check(api_client)
|
||||
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert "authenticated" in data
|
||||
assert "auth_required" in data
|
||||
assert "auth_mode" in data
|
||||
if data["auth_mode"] == "none":
|
||||
assert data == {
|
||||
"authenticated": True,
|
||||
"auth_required": False,
|
||||
"auth_mode": "none",
|
||||
"is_admin": True,
|
||||
}
|
||||
return
|
||||
|
||||
assert isinstance(data["authenticated"], bool)
|
||||
assert data["auth_required"] is True
|
||||
assert data["auth_mode"] in ["builtin", "cwa", "proxy", "oidc"]
|
||||
assert isinstance(data["is_admin"], bool)
|
||||
assert data["username"] is None or isinstance(data["username"], str)
|
||||
assert "display_name" in data
|
||||
assert data["display_name"] is None or isinstance(data["display_name"], str)
|
||||
|
||||
def test_auth_check_returns_auth_mode(self, api_client: APIClient):
|
||||
"""Test that auth check returns the current auth mode."""
|
||||
resp = api_client.get("/api/auth/check")
|
||||
|
||||
data = resp.json()
|
||||
assert "auth_mode" in data
|
||||
# Should be one of the valid auth modes
|
||||
"""Test that auth check reports a known auth mode."""
|
||||
data = _auth_check(api_client)
|
||||
assert data["auth_mode"] in ["none", "builtin", "cwa", "proxy", "oidc"]
|
||||
|
||||
def test_auth_check_includes_admin_status(self, api_client: APIClient):
|
||||
"""Test that auth check includes admin status."""
|
||||
resp = api_client.get("/api/auth/check")
|
||||
|
||||
data = resp.json()
|
||||
assert "is_admin" in data
|
||||
"""Test that auth check exposes a boolean admin flag."""
|
||||
data = _auth_check(api_client)
|
||||
assert isinstance(data["is_admin"], bool)
|
||||
|
||||
def test_logout_endpoint_exists(self, api_client: APIClient):
|
||||
"""Test that logout endpoint is accessible."""
|
||||
"""Test that logout returns the stable success contract."""
|
||||
resp = api_client.post("/api/auth/logout")
|
||||
|
||||
# Should return 200 whether authenticated or not
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert "success" in data
|
||||
assert data.get("success") is True
|
||||
assert set(data).issubset({"success", "logout_url"})
|
||||
if "logout_url" in data:
|
||||
assert isinstance(data["logout_url"], str)
|
||||
assert data["logout_url"].startswith("http")
|
||||
|
||||
def test_logout_may_return_logout_url(self, api_client: APIClient):
|
||||
"""Test that logout may return a logout URL for proxy auth."""
|
||||
resp = api_client.post("/api/auth/logout")
|
||||
|
||||
data = resp.json()
|
||||
# logout_url is optional depending on auth mode
|
||||
if "logout_url" in data:
|
||||
assert isinstance(data["logout_url"], str)
|
||||
assert data["logout_url"].startswith("http")
|
||||
|
||||
def test_login_endpoint_exists(self, api_client: APIClient):
|
||||
"""Test that login endpoint is accessible."""
|
||||
"""Test that login obeys the current authentication contract."""
|
||||
auth_data = _auth_check(api_client)
|
||||
username = f"e2e-auth-{uuid4().hex[:8]}"
|
||||
resp = api_client.post(
|
||||
"/api/auth/login", json={"username": "test", "password": "test", "remember_me": False}
|
||||
"/api/auth/login",
|
||||
json={"username": username, "password": "wrong-password", "remember_me": False},
|
||||
)
|
||||
|
||||
# Should return some response (may be success, auth error, or rate limit)
|
||||
assert resp.status_code in [200, 401, 403, 429]
|
||||
auth_mode = auth_data.get("auth_mode")
|
||||
if auth_mode == "none":
|
||||
assert resp.status_code == 200
|
||||
assert resp.json() == {"success": True}
|
||||
elif auth_mode == "proxy":
|
||||
assert resp.status_code == 401
|
||||
assert resp.json() == {"error": "Proxy authentication is enabled"}
|
||||
elif auth_mode in {"builtin", "cwa"}:
|
||||
assert resp.status_code == 401
|
||||
assert resp.json() == {"error": "Invalid username or password."}
|
||||
elif auth_mode == "oidc":
|
||||
if auth_data.get("hide_local_auth"):
|
||||
assert resp.status_code == 403
|
||||
assert resp.json() == {"error": "Local authentication is disabled"}
|
||||
else:
|
||||
assert resp.status_code == 401
|
||||
assert resp.json() == {"error": "Invalid username or password."}
|
||||
else:
|
||||
pytest.fail(f"Unexpected auth mode: {auth_mode}")
|
||||
|
||||
def test_login_with_no_auth_succeeds(self, api_client: APIClient):
|
||||
"""Test that login succeeds when no authentication is required."""
|
||||
# First check if auth is required
|
||||
auth_check = api_client.get("/api/auth/check")
|
||||
auth_data = auth_check.json()
|
||||
auth_data = _auth_check(api_client)
|
||||
|
||||
if not auth_data.get("auth_required"):
|
||||
# Try logging in
|
||||
resp = api_client.post(
|
||||
"/api/auth/login",
|
||||
json={"username": "anyuser", "password": "anypass", "remember_me": False},
|
||||
)
|
||||
|
||||
# Should succeed
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data.get("success") is True
|
||||
assert resp.json() == {"success": True}
|
||||
assert api_client.get("/api/auth/check").json() == {
|
||||
"authenticated": True,
|
||||
"auth_required": False,
|
||||
"auth_mode": "none",
|
||||
"is_admin": True,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -94,34 +139,35 @@ class TestProxyAuthentication:
|
||||
"""Tests for proxy authentication mode."""
|
||||
|
||||
def test_proxy_auth_with_valid_header(self, api_client: APIClient):
|
||||
"""Test proxy auth when valid user header is present."""
|
||||
# Check current auth mode
|
||||
auth_check = api_client.get("/api/auth/check")
|
||||
auth_data = auth_check.json()
|
||||
"""Test proxy auth creates and preserves a session from the proxy header."""
|
||||
auth_data = _auth_check(api_client)
|
||||
|
||||
if auth_data.get("auth_mode") != "proxy":
|
||||
pytest.skip("Proxy authentication not configured")
|
||||
|
||||
# Make a request with proxy auth header
|
||||
# Note: In real deployment, these headers would be set by the proxy
|
||||
resp = api_client.get("/api/config", headers={"X-Auth-User": "proxyuser"})
|
||||
|
||||
if resp.status_code == 401:
|
||||
pytest.skip("Proxy auth header not accepted (check proxy configuration)")
|
||||
|
||||
# Should be able to access the endpoint
|
||||
resp = api_client.get("/api/auth/check", headers={"X-Auth-User": "proxyuser"})
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["authenticated"] is True
|
||||
assert data["auth_mode"] == "proxy"
|
||||
assert data["username"] == "proxyuser"
|
||||
assert data["auth_required"] is True
|
||||
assert isinstance(data["is_admin"], bool)
|
||||
assert "display_name" in data
|
||||
|
||||
follow_up = api_client.get("/api/auth/check")
|
||||
follow_up_data = follow_up.json()
|
||||
assert follow_up.status_code == 200
|
||||
assert follow_up_data["authenticated"] is True
|
||||
assert follow_up_data["username"] == "proxyuser"
|
||||
|
||||
def test_proxy_auth_logout_url_available(self, api_client: APIClient):
|
||||
"""Test that proxy auth provides logout URL if configured."""
|
||||
# Check current auth mode
|
||||
auth_check = api_client.get("/api/auth/check")
|
||||
auth_data = auth_check.json()
|
||||
auth_data = _auth_check(api_client)
|
||||
|
||||
if auth_data.get("auth_mode") != "proxy":
|
||||
pytest.skip("Proxy authentication not configured")
|
||||
|
||||
# Check for logout URL in auth check response
|
||||
if "logout_url" in auth_data:
|
||||
assert isinstance(auth_data["logout_url"], str)
|
||||
assert len(auth_data["logout_url"]) > 0
|
||||
@@ -133,39 +179,32 @@ class TestBuiltinAuthentication:
|
||||
|
||||
def test_builtin_auth_requires_credentials(self, api_client: APIClient):
|
||||
"""Test that endpoints require authentication when builtin auth is enabled."""
|
||||
# Check current auth mode
|
||||
auth_check = api_client.get("/api/auth/check")
|
||||
auth_data = auth_check.json()
|
||||
auth_data = _auth_check(api_client)
|
||||
|
||||
if auth_data.get("auth_mode") != "builtin":
|
||||
pytest.skip("Built-in authentication not configured")
|
||||
|
||||
if not auth_data.get("authenticated"):
|
||||
# Attempt to access protected endpoint without authentication
|
||||
resp = api_client.get("/api/config")
|
||||
if auth_data.get("authenticated"):
|
||||
pytest.skip("Built-in auth session already authenticated")
|
||||
|
||||
# Should be blocked
|
||||
assert resp.status_code == 401
|
||||
resp = api_client.get("/api/config")
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_builtin_auth_invalid_credentials(self, api_client: APIClient):
|
||||
"""Test login with invalid credentials fails."""
|
||||
# Check current auth mode
|
||||
auth_check = api_client.get("/api/auth/check")
|
||||
auth_data = auth_check.json()
|
||||
auth_data = _auth_check(api_client)
|
||||
|
||||
if auth_data.get("auth_mode") != "builtin":
|
||||
pytest.skip("Built-in authentication not configured")
|
||||
|
||||
# Try logging in with invalid credentials
|
||||
username = f"builtin-e2e-{uuid4().hex[:8]}"
|
||||
resp = api_client.post(
|
||||
"/api/auth/login",
|
||||
json={"username": "invalid_user", "password": "wrong_password", "remember_me": False},
|
||||
json={"username": username, "password": "wrong_password", "remember_me": False},
|
||||
)
|
||||
|
||||
# Should fail, or be rate-limited on a live stack after repeated attempts
|
||||
assert resp.status_code in [401, 403, 429]
|
||||
data = resp.json()
|
||||
assert data.get("success") is not True
|
||||
assert resp.status_code == 401
|
||||
assert resp.json() == {"error": "Invalid username or password."}
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -174,53 +213,29 @@ class TestCalibreWebAuthentication:
|
||||
|
||||
def test_cwa_auth_mode_available(self, api_client: APIClient):
|
||||
"""Test that CWA auth mode is reported if configured."""
|
||||
# Check current auth mode
|
||||
auth_check = api_client.get("/api/auth/check")
|
||||
auth_data = auth_check.json()
|
||||
auth_data = _auth_check(api_client)
|
||||
|
||||
if auth_data.get("auth_mode") == "cwa":
|
||||
# CWA mode is active
|
||||
assert auth_data["auth_mode"] == "cwa"
|
||||
# Should have authenticated or auth_required status
|
||||
assert "authenticated" in auth_data
|
||||
assert "auth_required" in auth_data
|
||||
assert auth_data["auth_required"] is True
|
||||
assert isinstance(auth_data["authenticated"], bool)
|
||||
assert isinstance(auth_data["is_admin"], bool)
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
class TestAdminAccess:
|
||||
"""Tests for admin access restrictions."""
|
||||
|
||||
def test_settings_endpoint_respects_admin_restriction(self, api_client: APIClient):
|
||||
"""Test that settings endpoints respect admin restrictions."""
|
||||
# Check current auth status
|
||||
auth_check = api_client.get("/api/auth/check")
|
||||
auth_data = auth_check.json()
|
||||
def test_admin_only_routes_require_auth(self, api_client: APIClient):
|
||||
"""Test that admin-only routes are blocked before auth is established."""
|
||||
auth_data = _auth_check(api_client)
|
||||
|
||||
# If auth is required and user is not admin
|
||||
if auth_data.get("auth_required") and auth_data.get("authenticated"):
|
||||
if not auth_data.get("is_admin"):
|
||||
# Try accessing settings
|
||||
resp = api_client.get("/api/settings")
|
||||
if not auth_data.get("auth_required"):
|
||||
pytest.skip("Authentication is not required in this environment")
|
||||
|
||||
# May be blocked with 403 if admin-only
|
||||
# Or allowed if settings are not restricted
|
||||
assert resp.status_code in [200, 403]
|
||||
|
||||
def test_onboarding_endpoint_respects_admin_restriction(self, api_client: APIClient):
|
||||
"""Test that onboarding endpoints respect admin restrictions."""
|
||||
# Check current auth status
|
||||
auth_check = api_client.get("/api/auth/check")
|
||||
auth_data = auth_check.json()
|
||||
|
||||
# If auth is required and user is not admin
|
||||
if auth_data.get("auth_required") and auth_data.get("authenticated"):
|
||||
if not auth_data.get("is_admin"):
|
||||
# Try accessing onboarding
|
||||
resp = api_client.get("/api/onboarding")
|
||||
|
||||
# May be blocked with 403 if admin-only
|
||||
# Or allowed if settings are not restricted
|
||||
assert resp.status_code in [200, 403]
|
||||
for path in ("/api/settings/security", "/api/settings/users", "/api/onboarding"):
|
||||
resp = api_client.get(path)
|
||||
assert resp.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -229,45 +244,28 @@ class TestAuthenticationWorkflow:
|
||||
|
||||
def test_login_logout_cycle(self, api_client: APIClient):
|
||||
"""Test complete login and logout cycle."""
|
||||
# Check initial auth status
|
||||
auth_check = api_client.get("/api/auth/check")
|
||||
initial_auth = auth_check.json()
|
||||
initial_auth = _auth_check(api_client)
|
||||
|
||||
# If no auth required, skip this test
|
||||
if not initial_auth.get("auth_required"):
|
||||
pytest.skip("No authentication required")
|
||||
|
||||
# Try logout first to clear any existing session
|
||||
logout_resp = api_client.post("/api/auth/logout")
|
||||
assert logout_resp.status_code == 200
|
||||
assert logout_resp.json().get("success") is True
|
||||
|
||||
# Check we're logged out
|
||||
auth_check = api_client.get("/api/auth/check")
|
||||
post_logout_auth = auth_check.json()
|
||||
post_logout_auth = _auth_check(api_client)
|
||||
|
||||
# For builtin/cwa auth, should not be authenticated
|
||||
# For proxy auth, depends on proxy configuration
|
||||
if initial_auth.get("auth_mode") in ["builtin", "cwa"]:
|
||||
if (
|
||||
initial_auth.get("auth_mode") in ["builtin", "cwa"]
|
||||
or initial_auth.get("auth_mode") == "proxy"
|
||||
):
|
||||
assert post_logout_auth.get("authenticated") is False
|
||||
assert post_logout_auth.get("username") is None
|
||||
|
||||
def test_auth_check_consistency(self, api_client: APIClient):
|
||||
"""Test that auth check returns consistent results."""
|
||||
# Make multiple auth check requests
|
||||
resp1 = api_client.get("/api/auth/check")
|
||||
resp2 = api_client.get("/api/auth/check")
|
||||
resp3 = api_client.get("/api/auth/check")
|
||||
data1 = _auth_check(api_client)
|
||||
data2 = _auth_check(api_client)
|
||||
data3 = _auth_check(api_client)
|
||||
|
||||
data1 = resp1.json()
|
||||
data2 = resp2.json()
|
||||
data3 = resp3.json()
|
||||
|
||||
# All should succeed
|
||||
assert resp1.status_code == 200
|
||||
assert resp2.status_code == 200
|
||||
assert resp3.status_code == 200
|
||||
|
||||
# Auth mode should be consistent
|
||||
assert data1["auth_mode"] == data2["auth_mode"] == data3["auth_mode"]
|
||||
|
||||
# Auth required should be consistent
|
||||
assert data1["auth_required"] == data2["auth_required"] == data3["auth_required"]
|
||||
assert data1 == data2 == data3
|
||||
|
||||
@@ -4,26 +4,195 @@ from _pytest.outcomes import Failed, Skipped
|
||||
from tests.e2e import conftest as e2e_conftest
|
||||
|
||||
|
||||
class DummyResponse:
|
||||
def __init__(self, status_code: int, payload: object, text: str = "") -> None:
|
||||
self.status_code = status_code
|
||||
self._payload = payload
|
||||
self.text = text or repr(payload)
|
||||
|
||||
def json(self) -> object:
|
||||
return self._payload
|
||||
|
||||
|
||||
class RaisingResponse(DummyResponse):
|
||||
def __init__(self, status_code: int, exc: Exception, text: str = "") -> None:
|
||||
super().__init__(status_code=status_code, payload=None, text=text)
|
||||
self._exc = exc
|
||||
|
||||
def json(self) -> object:
|
||||
raise self._exc
|
||||
|
||||
|
||||
def _make_client() -> e2e_conftest.APIClient:
|
||||
return e2e_conftest.APIClient(base_url="http://example.com")
|
||||
|
||||
|
||||
def test_api_client_wait_for_health_retries_through_request_exception(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
client = _make_client()
|
||||
calls = {"count": 0}
|
||||
|
||||
def fake_get(_path: str, **_kwargs: object) -> DummyResponse:
|
||||
calls["count"] += 1
|
||||
if calls["count"] == 1:
|
||||
raise e2e_conftest.requests.exceptions.ReadTimeout("timeout")
|
||||
return DummyResponse(200, {"status": "ok"})
|
||||
|
||||
times = [0.0, 0.0, 0.1]
|
||||
|
||||
monkeypatch.setattr(client, "get", fake_get)
|
||||
monkeypatch.setattr(e2e_conftest.time, "time", lambda: times.pop(0))
|
||||
monkeypatch.setattr(e2e_conftest.time, "sleep", lambda _seconds: None)
|
||||
|
||||
try:
|
||||
assert client.wait_for_health(max_wait=1) is True
|
||||
assert calls["count"] == 2
|
||||
finally:
|
||||
client.close()
|
||||
|
||||
|
||||
def test_download_tracker_wait_for_status_ignores_malformed_payloads_and_returns_error(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
class FakeClient:
|
||||
def __init__(self) -> None:
|
||||
self.responses = [
|
||||
DummyResponse(200, ["not", "a", "mapping"]),
|
||||
DummyResponse(200, {"error": {"task-1": {"message": "boom"}}}),
|
||||
]
|
||||
|
||||
def get(self, _path: str) -> DummyResponse:
|
||||
return self.responses.pop(0)
|
||||
|
||||
tracker = e2e_conftest.DownloadTracker(client=FakeClient())
|
||||
times = [0.0, 0.0, 0.1]
|
||||
|
||||
monkeypatch.setattr(e2e_conftest.time, "time", lambda: times.pop(0))
|
||||
monkeypatch.setattr(e2e_conftest.time, "sleep", lambda _seconds: None)
|
||||
|
||||
result = tracker.wait_for_status("task-1", ["complete"], timeout=1)
|
||||
|
||||
assert result == {"state": "error", "data": {"message": "boom"}}
|
||||
|
||||
|
||||
def test_assert_json_object_returns_dict() -> None:
|
||||
response = DummyResponse(200, {"status": "ok"})
|
||||
|
||||
assert e2e_conftest.assert_json_object(response, context="health") == {"status": "ok"}
|
||||
|
||||
|
||||
def test_assert_json_object_rejects_non_object_payload() -> None:
|
||||
response = DummyResponse(200, ["not", "an", "object"])
|
||||
|
||||
with pytest.raises(AssertionError, match="did not return a JSON object"):
|
||||
e2e_conftest.assert_json_object(response, context="health")
|
||||
|
||||
|
||||
def test_assert_json_object_rejects_invalid_json() -> None:
|
||||
response = RaisingResponse(200, ValueError("bad json"))
|
||||
|
||||
with pytest.raises(Failed, match="did not return valid JSON"):
|
||||
e2e_conftest.assert_json_object(response, context="health")
|
||||
|
||||
|
||||
def test_assert_json_list_returns_list() -> None:
|
||||
response = DummyResponse(200, [1, 2, 3])
|
||||
|
||||
assert e2e_conftest.assert_json_list(response, context="queue") == [1, 2, 3]
|
||||
|
||||
|
||||
def test_assert_queue_order_response_validates_entries() -> None:
|
||||
response = DummyResponse(
|
||||
200,
|
||||
{
|
||||
"queue": [
|
||||
{
|
||||
"id": "task-1",
|
||||
"priority": 0,
|
||||
"added_time": 12.5,
|
||||
"status": "queued",
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
|
||||
queue = e2e_conftest.assert_queue_order_response(response)
|
||||
assert queue[0]["id"] == "task-1"
|
||||
|
||||
|
||||
def test_assert_queue_order_response_rejects_missing_entry_fields() -> None:
|
||||
response = DummyResponse(200, {"queue": [{"id": "task-1"}]})
|
||||
|
||||
with pytest.raises(AssertionError):
|
||||
e2e_conftest.assert_queue_order_response(response)
|
||||
|
||||
|
||||
def test_require_authenticated_client_allows_public_server(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
client = e2e_conftest.APIClient(base_url="http://example.com")
|
||||
monkeypatch.setattr(
|
||||
e2e_conftest,
|
||||
"_get_auth_state",
|
||||
lambda _: {"auth_required": False},
|
||||
)
|
||||
client = _make_client()
|
||||
monkeypatch.setattr(e2e_conftest, "_get_auth_state", lambda _: {"auth_required": False})
|
||||
|
||||
try:
|
||||
assert e2e_conftest._require_authenticated_client(client, strict=True) is client
|
||||
finally:
|
||||
client.session.close()
|
||||
client.close()
|
||||
|
||||
|
||||
def test_require_authenticated_client_fails_when_auth_state_unavailable(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
client = _make_client()
|
||||
monkeypatch.setattr(e2e_conftest, "_get_auth_state", lambda _: None)
|
||||
|
||||
try:
|
||||
with pytest.raises(Failed, match="Unable to read auth state"):
|
||||
e2e_conftest._require_authenticated_client(client, strict=True)
|
||||
finally:
|
||||
client.close()
|
||||
|
||||
|
||||
def test_require_healthy_server_fails_in_strict_runs(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
closed: list[str] = []
|
||||
|
||||
monkeypatch.setattr(e2e_conftest.APIClient, "wait_for_health", lambda self, max_wait=30: False)
|
||||
monkeypatch.setattr(e2e_conftest.APIClient, "close", lambda self: closed.append(self.base_url))
|
||||
|
||||
with pytest.raises(Failed, match="Server not available"):
|
||||
e2e_conftest._require_healthy_server("http://example.com", strict=True)
|
||||
|
||||
assert closed == ["http://example.com"]
|
||||
|
||||
|
||||
def test_require_healthy_server_skips_in_non_strict_runs(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(e2e_conftest.APIClient, "wait_for_health", lambda self, max_wait=30: False)
|
||||
|
||||
with pytest.raises(Skipped, match="Server not available"):
|
||||
e2e_conftest._require_healthy_server("http://example.com", strict=False)
|
||||
|
||||
|
||||
def test_require_authenticated_client_skips_when_auth_state_unavailable_and_not_strict(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
client = _make_client()
|
||||
monkeypatch.setattr(e2e_conftest, "_get_auth_state", lambda _: None)
|
||||
|
||||
try:
|
||||
with pytest.raises(Skipped, match="Unable to read auth state"):
|
||||
e2e_conftest._require_authenticated_client(client, strict=False)
|
||||
finally:
|
||||
client.close()
|
||||
|
||||
|
||||
def test_require_authenticated_client_fails_without_env_credentials(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
client = e2e_conftest.APIClient(base_url="http://example.com")
|
||||
client = _make_client()
|
||||
monkeypatch.setattr(
|
||||
e2e_conftest,
|
||||
"_get_auth_state",
|
||||
@@ -36,13 +205,13 @@ def test_require_authenticated_client_fails_without_env_credentials(
|
||||
with pytest.raises(Failed, match="requires authentication"):
|
||||
e2e_conftest._require_authenticated_client(client, strict=True)
|
||||
finally:
|
||||
client.session.close()
|
||||
client.close()
|
||||
|
||||
|
||||
def test_require_authenticated_client_skips_without_env_credentials_when_not_strict(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
client = e2e_conftest.APIClient(base_url="http://example.com")
|
||||
client = _make_client()
|
||||
monkeypatch.setattr(
|
||||
e2e_conftest,
|
||||
"_get_auth_state",
|
||||
@@ -55,13 +224,13 @@ def test_require_authenticated_client_skips_without_env_credentials_when_not_str
|
||||
with pytest.raises(Skipped, match="requires authentication"):
|
||||
e2e_conftest._require_authenticated_client(client, strict=False)
|
||||
finally:
|
||||
client.session.close()
|
||||
client.close()
|
||||
|
||||
|
||||
def test_require_authenticated_client_logs_in_when_credentials_are_available(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
client = e2e_conftest.APIClient(base_url="http://example.com")
|
||||
client = _make_client()
|
||||
auth_states = iter(
|
||||
[
|
||||
{"auth_required": True, "authenticated": False},
|
||||
@@ -76,7 +245,7 @@ def test_require_authenticated_client_logs_in_when_credentials_are_available(
|
||||
try:
|
||||
assert e2e_conftest._require_authenticated_client(client, strict=True) is client
|
||||
finally:
|
||||
client.session.close()
|
||||
client.close()
|
||||
|
||||
|
||||
def test_is_explicit_e2e_run_detects_e2e_path_selection() -> None:
|
||||
@@ -86,6 +255,16 @@ def test_is_explicit_e2e_run_detects_e2e_path_selection() -> None:
|
||||
)
|
||||
|
||||
|
||||
def test_is_explicit_e2e_run_detects_mixed_path_selection() -> None:
|
||||
assert (
|
||||
e2e_conftest._is_explicit_e2e_run(
|
||||
"",
|
||||
["tests/e2e/test_api.py", "tests/core/test_user_db.py"],
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_is_explicit_e2e_run_detects_markexpr_selection() -> None:
|
||||
assert e2e_conftest._is_explicit_e2e_run("e2e", ["tests/"]) is True
|
||||
assert e2e_conftest._is_explicit_e2e_run("slow and e2e", ["tests/"]) is True
|
||||
|
||||
+271
-113
@@ -7,59 +7,204 @@ They require external services to be available and may take longer to run.
|
||||
Run with: uv run pytest tests/e2e/test_download_flow.py -v -m e2e
|
||||
"""
|
||||
|
||||
import os
|
||||
import hashlib
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from .conftest import APIClient, DownloadTracker, DOWNLOAD_TIMEOUT
|
||||
from .conftest import (
|
||||
DOWNLOAD_TIMEOUT,
|
||||
SUCCESS_DOWNLOAD_STATES,
|
||||
APIClient,
|
||||
DownloadTracker,
|
||||
assert_queue_order_response,
|
||||
assert_queued_download_response,
|
||||
)
|
||||
|
||||
|
||||
def _assert_terminal_download_result(
|
||||
result: dict[str, object],
|
||||
*,
|
||||
source_id: str,
|
||||
expected_title: str,
|
||||
expected_source: str | None = None,
|
||||
) -> None:
|
||||
"""Assert that a finished download produced a structured queue payload."""
|
||||
state = result["state"]
|
||||
entry = result["data"]
|
||||
assert isinstance(entry, dict)
|
||||
if state == "error":
|
||||
error_message = str(
|
||||
entry.get("status_message") or entry.get("last_error_message") or ""
|
||||
).strip()
|
||||
pytest.fail(
|
||||
f"{expected_source or source_id} download failed"
|
||||
f"{f': {error_message}' if error_message else f': {entry!r}'}"
|
||||
)
|
||||
|
||||
assert state in SUCCESS_DOWNLOAD_STATES, (
|
||||
f"{expected_source or source_id} ended in unexpected state {state!r}: {entry!r}"
|
||||
)
|
||||
assert entry.get("id") == source_id
|
||||
assert entry.get("title") == expected_title
|
||||
if expected_source is not None:
|
||||
assert entry.get("source") == expected_source
|
||||
status = entry.get("status")
|
||||
assert status is None or status in SUCCESS_DOWNLOAD_STATES | {"queued"}
|
||||
|
||||
|
||||
def _is_duplicate_queue_error(response) -> bool:
|
||||
"""Whether the API refused to queue a release because it already exists."""
|
||||
if response.status_code != 500:
|
||||
return False
|
||||
try:
|
||||
payload = response.json()
|
||||
except ValueError:
|
||||
return False
|
||||
return payload == {"error": "Release is already in the download queue"}
|
||||
|
||||
|
||||
def _require_json_object(
|
||||
response, *, context: str, skip_statuses: set[int] | frozenset[int] = frozenset({503})
|
||||
) -> dict[str, object]:
|
||||
if response.status_code in skip_statuses:
|
||||
pytest.skip(f"{context} unavailable: {response.status_code}")
|
||||
assert response.status_code == 200, f"{context} failed: {response.status_code} {response.text}"
|
||||
payload = response.json()
|
||||
assert isinstance(payload, dict), f"{context} did not return a JSON object: {payload!r}"
|
||||
return payload
|
||||
|
||||
|
||||
def _extract_result_list(payload: object, *, context: str) -> list[dict[str, object]]:
|
||||
if isinstance(payload, dict):
|
||||
if "books" in payload:
|
||||
results = payload["books"]
|
||||
elif "releases" in payload:
|
||||
results = payload["releases"]
|
||||
else:
|
||||
results = payload.get("results", payload)
|
||||
else:
|
||||
results = payload
|
||||
|
||||
if isinstance(results, dict):
|
||||
list_values = [value for value in results.values() if isinstance(value, list)]
|
||||
if not list_values:
|
||||
pytest.fail(f"{context} returned an unexpected result structure: {payload!r}")
|
||||
if not any(list_values):
|
||||
pytest.skip(f"{context} returned no results")
|
||||
results = next(value for value in list_values if value)
|
||||
|
||||
if not isinstance(results, list):
|
||||
pytest.fail(f"{context} did not return a result list: {payload!r}")
|
||||
if not results:
|
||||
pytest.skip(f"{context} returned no results")
|
||||
return results
|
||||
|
||||
|
||||
def _require_queue_entry(
|
||||
api_client: APIClient,
|
||||
book_id: str,
|
||||
*,
|
||||
context: str,
|
||||
timeout: int = 20,
|
||||
) -> dict[str, object]:
|
||||
deadline = time.time() + timeout
|
||||
while time.time() < deadline:
|
||||
queue_resp = api_client.get("/api/queue/order")
|
||||
if queue_resp.status_code == 200:
|
||||
queue_order = assert_queue_order_response(queue_resp)
|
||||
for entry in queue_order:
|
||||
if entry.get("id") == book_id:
|
||||
return entry
|
||||
elif queue_resp.status_code == 503:
|
||||
pytest.skip(f"{context} queue endpoint unavailable")
|
||||
else:
|
||||
pytest.fail(
|
||||
f"{context} queue lookup failed: {queue_resp.status_code} {queue_resp.text}"
|
||||
)
|
||||
|
||||
status_resp = api_client.get("/api/status")
|
||||
if status_resp.status_code == 200:
|
||||
status_data = status_resp.json()
|
||||
if isinstance(status_data, dict):
|
||||
for state in ("complete", "done", "available", "error", "cancelled"):
|
||||
state_entries = status_data.get(state)
|
||||
if isinstance(state_entries, dict) and book_id in state_entries:
|
||||
pytest.fail(
|
||||
f"{context} reached terminal state {state} before it was observed in the queue: "
|
||||
f"{state_entries[book_id]!r}"
|
||||
)
|
||||
|
||||
time.sleep(1)
|
||||
|
||||
pytest.fail(f"{context} never appeared in the queue")
|
||||
|
||||
|
||||
def _wait_for_queue_absence(
|
||||
api_client: APIClient,
|
||||
book_id: str,
|
||||
*,
|
||||
context: str,
|
||||
timeout: int = 20,
|
||||
) -> None:
|
||||
deadline = time.time() + timeout
|
||||
while time.time() < deadline:
|
||||
queue_resp = api_client.get("/api/queue/order")
|
||||
if queue_resp.status_code == 200:
|
||||
queue_order = assert_queue_order_response(queue_resp)
|
||||
if all(entry.get("id") != book_id for entry in queue_order):
|
||||
return
|
||||
elif queue_resp.status_code == 503:
|
||||
pytest.skip(f"{context} queue endpoint unavailable")
|
||||
else:
|
||||
pytest.fail(
|
||||
f"{context} queue lookup failed: {queue_resp.status_code} {queue_resp.text}"
|
||||
)
|
||||
|
||||
time.sleep(1)
|
||||
|
||||
pytest.fail(f"{context} still appeared in the queue after cancellation")
|
||||
|
||||
|
||||
def _find_available_provider(api_client: APIClient) -> str | None:
|
||||
"""Find a working metadata provider."""
|
||||
resp = api_client.get("/api/metadata/providers")
|
||||
if resp.status_code != 200:
|
||||
return None
|
||||
if resp.status_code == 503:
|
||||
pytest.skip("Metadata providers unavailable")
|
||||
assert resp.status_code == 200, f"Metadata providers failed: {resp.status_code} {resp.text}"
|
||||
|
||||
providers_data = resp.json()
|
||||
if not isinstance(providers_data, (dict, list)):
|
||||
pytest.fail(f"Metadata providers returned an unexpected payload: {providers_data!r}")
|
||||
|
||||
# Handle both dict and list formats
|
||||
if isinstance(providers_data, dict):
|
||||
# Dict format: keys are provider names
|
||||
provider_names = list(providers_data.keys())
|
||||
else:
|
||||
# List format
|
||||
provider_names = [
|
||||
p.get("name") for p in providers_data if isinstance(p, dict) and p.get("name")
|
||||
]
|
||||
providers = (
|
||||
providers_data.get("providers", []) if isinstance(providers_data, dict) else providers_data
|
||||
)
|
||||
provider_names = [
|
||||
provider.get("name")
|
||||
for provider in providers
|
||||
if (
|
||||
isinstance(provider, dict)
|
||||
and provider.get("name")
|
||||
and provider.get("enabled") is True
|
||||
and provider.get("available") is True
|
||||
)
|
||||
]
|
||||
|
||||
for name in provider_names:
|
||||
if name:
|
||||
# Try a simple search to verify it works
|
||||
test_resp = api_client.get(
|
||||
"/api/metadata/search",
|
||||
params={"query": "test", "provider": name},
|
||||
timeout=30,
|
||||
)
|
||||
if test_resp.status_code == 200:
|
||||
return name
|
||||
return None
|
||||
|
||||
|
||||
def _find_available_release_source(api_client: APIClient) -> str | None:
|
||||
"""Find a working release source."""
|
||||
resp = api_client.get("/api/release-sources")
|
||||
if resp.status_code != 200:
|
||||
return None
|
||||
|
||||
sources = resp.json()
|
||||
for source in sources:
|
||||
name = source.get("name")
|
||||
# Skip prowlarr unless configured
|
||||
if name and name != "prowlarr":
|
||||
test_resp = api_client.get(
|
||||
"/api/metadata/search",
|
||||
params={"query": "test", "provider": name},
|
||||
timeout=30,
|
||||
)
|
||||
if test_resp.status_code == 200:
|
||||
return name
|
||||
return None
|
||||
if test_resp.status_code != 503:
|
||||
pytest.fail(
|
||||
f"Metadata provider {name} failed during availability check: "
|
||||
f"{test_resp.status_code} {test_resp.text}"
|
||||
)
|
||||
|
||||
pytest.skip("No working metadata providers available")
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -71,8 +216,6 @@ class TestMetadataToReleaseFlow:
|
||||
"""Test searching metadata then finding releases."""
|
||||
# Find a working provider
|
||||
provider = _find_available_provider(protected_api_client)
|
||||
if not provider:
|
||||
pytest.skip("No metadata providers available")
|
||||
|
||||
# Search for a public domain book
|
||||
search_resp = protected_api_client.get(
|
||||
@@ -81,22 +224,8 @@ class TestMetadataToReleaseFlow:
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
if search_resp.status_code != 200:
|
||||
pytest.skip(f"Search failed: {search_resp.status_code}")
|
||||
|
||||
search_data = search_resp.json()
|
||||
results = search_data.get("results", search_data)
|
||||
|
||||
# Handle dict format where results might be nested
|
||||
if isinstance(results, dict):
|
||||
# Results might be under a key like the query or "results"
|
||||
for key, value in results.items():
|
||||
if isinstance(value, list) and value:
|
||||
results = value
|
||||
break
|
||||
|
||||
if not results or not isinstance(results, list):
|
||||
pytest.skip("No search results returned")
|
||||
search_data = _require_json_object(search_resp, context="metadata search")
|
||||
results = _extract_result_list(search_data, context="metadata search")
|
||||
|
||||
# Get the first result
|
||||
first_result = results[0]
|
||||
@@ -115,11 +244,13 @@ class TestMetadataToReleaseFlow:
|
||||
timeout=60,
|
||||
)
|
||||
|
||||
# Releases may fail if sources are unavailable
|
||||
if releases_resp.status_code == 200:
|
||||
releases_data = releases_resp.json()
|
||||
assert "releases" in releases_data
|
||||
assert "book" in releases_data
|
||||
releases_data = _require_json_object(releases_resp, context="release lookup")
|
||||
releases = releases_data.get("releases")
|
||||
if not isinstance(releases, list):
|
||||
pytest.fail(f"Release lookup returned an invalid releases payload: {releases_data!r}")
|
||||
if not releases:
|
||||
pytest.skip("No releases available")
|
||||
assert "book" in releases_data
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -152,21 +283,8 @@ class TestFullDownloadJourney:
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
if search_resp.status_code != 200:
|
||||
pytest.skip(f"Metadata search unavailable: {search_resp.status_code}")
|
||||
|
||||
search_data = search_resp.json()
|
||||
results = search_data.get("results", search_data)
|
||||
|
||||
# Handle dict format where results might be nested
|
||||
if isinstance(results, dict):
|
||||
for key, value in results.items():
|
||||
if isinstance(value, list) and value:
|
||||
results = value
|
||||
break
|
||||
|
||||
if not results or not isinstance(results, list):
|
||||
pytest.skip("No search results")
|
||||
search_data = _require_json_object(search_resp, context="metadata search")
|
||||
results = _extract_result_list(search_data, context="metadata search")
|
||||
|
||||
first_result = results[0]
|
||||
book_id = first_result.get("id") or first_result.get("provider_id")
|
||||
@@ -182,12 +300,8 @@ class TestFullDownloadJourney:
|
||||
timeout=60,
|
||||
)
|
||||
|
||||
if releases_resp.status_code != 200:
|
||||
pytest.skip(f"Releases unavailable: {releases_resp.status_code}")
|
||||
|
||||
releases_data = releases_resp.json()
|
||||
releases_data = _require_json_object(releases_resp, context="release lookup")
|
||||
releases = releases_data.get("releases", [])
|
||||
|
||||
if not releases:
|
||||
pytest.skip("No releases available")
|
||||
|
||||
@@ -218,9 +332,8 @@ class TestFullDownloadJourney:
|
||||
},
|
||||
)
|
||||
|
||||
assert queue_resp.status_code == 200, f"Failed to queue: {queue_resp.text}"
|
||||
queue_data = queue_resp.json()
|
||||
assert queue_data.get("status") == "queued"
|
||||
if not _is_duplicate_queue_error(queue_resp):
|
||||
assert_queued_download_response(queue_resp)
|
||||
|
||||
# Wait for download to complete (or error)
|
||||
result = download_tracker.wait_for_status(
|
||||
@@ -234,12 +347,17 @@ class TestFullDownloadJourney:
|
||||
status_resp = protected_api_client.get("/api/status")
|
||||
if status_resp.status_code == 200:
|
||||
status_data = status_resp.json()
|
||||
if "error" in status_data and source_id in status_data["error"]:
|
||||
error_info = status_data["error"][source_id]
|
||||
pytest.skip(f"Download failed: {error_info}")
|
||||
error_info = status_data.get("error", {}).get(source_id)
|
||||
if error_info:
|
||||
pytest.fail(f"Download failed: {error_info}")
|
||||
pytest.fail("Download timed out")
|
||||
|
||||
assert result["state"] in ["complete", "done", "available"]
|
||||
_assert_terminal_download_result(
|
||||
result,
|
||||
source_id=source_id,
|
||||
expected_title=target_release.get("title", "Test Book"),
|
||||
expected_source=target_release.get("source", "direct_download"),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -260,11 +378,17 @@ class TestDirectSourceReleaseFlow:
|
||||
if search_resp.status_code == 503:
|
||||
pytest.skip("Direct source query unavailable")
|
||||
|
||||
if search_resp.status_code != 200:
|
||||
pytest.skip(f"Direct source query failed: {search_resp.status_code}")
|
||||
assert search_resp.status_code == 200, (
|
||||
f"Direct source query failed: {search_resp.status_code} {search_resp.text}"
|
||||
)
|
||||
|
||||
payload = search_resp.json()
|
||||
assert isinstance(payload, dict), (
|
||||
f"Direct source query returned an unexpected payload: {payload!r}"
|
||||
)
|
||||
results = payload.get("releases") or []
|
||||
if not isinstance(results, list):
|
||||
pytest.fail(f"Direct source query returned an unexpected payload: {payload!r}")
|
||||
if not results:
|
||||
pytest.skip("No direct source query results")
|
||||
|
||||
@@ -277,7 +401,7 @@ class TestDirectSourceReleaseFlow:
|
||||
info_resp = protected_api_client.get(f"/api/release-sources/{source}/records/{source_id}")
|
||||
|
||||
if info_resp.status_code != 200:
|
||||
pytest.skip(f"Source record endpoint failed: {info_resp.status_code}")
|
||||
pytest.fail(f"Source record endpoint failed: {info_resp.status_code}")
|
||||
|
||||
# Queue download from the shared release payload
|
||||
download_tracker.track(source_id)
|
||||
@@ -286,11 +410,21 @@ class TestDirectSourceReleaseFlow:
|
||||
json={**first_result, "content_type": "ebook", "search_mode": "direct"},
|
||||
)
|
||||
|
||||
if download_resp.status_code != 200:
|
||||
pytest.skip(f"Release download queue failed: {download_resp.status_code}")
|
||||
if not _is_duplicate_queue_error(download_resp):
|
||||
assert_queued_download_response(download_resp)
|
||||
|
||||
download_data = download_resp.json()
|
||||
assert download_data.get("status") == "queued"
|
||||
result = download_tracker.wait_for_status(
|
||||
source_id,
|
||||
target_states=["complete", "done", "available"],
|
||||
timeout=DOWNLOAD_TIMEOUT,
|
||||
)
|
||||
assert result is not None, "Direct source download did not reach a terminal state"
|
||||
_assert_terminal_download_result(
|
||||
result,
|
||||
source_id=source_id,
|
||||
expected_title=first_result.get("title", "Unknown title"),
|
||||
expected_source="direct_download",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -314,16 +448,26 @@ class TestDownloadCancellation:
|
||||
},
|
||||
)
|
||||
|
||||
if queue_resp.status_code != 200:
|
||||
pytest.skip("Could not queue test download")
|
||||
assert_queued_download_response(queue_resp)
|
||||
|
||||
# Give it a moment
|
||||
time.sleep(1)
|
||||
_require_queue_entry(
|
||||
protected_api_client,
|
||||
test_id,
|
||||
context="cancel download precondition",
|
||||
)
|
||||
|
||||
# Cancel it
|
||||
cancel_resp = protected_api_client.delete(f"/api/download/{test_id}/cancel")
|
||||
|
||||
assert cancel_resp.status_code in [200, 204]
|
||||
assert cancel_resp.status_code == 200
|
||||
cancel_data = cancel_resp.json()
|
||||
assert cancel_data == {"status": "cancelled", "book_id": test_id}
|
||||
|
||||
_wait_for_queue_absence(
|
||||
protected_api_client,
|
||||
test_id,
|
||||
context="cancel download",
|
||||
)
|
||||
|
||||
def test_cancel_removes_from_queue(
|
||||
self, protected_api_client: APIClient, download_tracker: DownloadTracker
|
||||
@@ -342,18 +486,23 @@ class TestDownloadCancellation:
|
||||
},
|
||||
)
|
||||
|
||||
time.sleep(0.5)
|
||||
_require_queue_entry(
|
||||
protected_api_client,
|
||||
test_id,
|
||||
context="cancel verification precondition",
|
||||
)
|
||||
|
||||
# Cancel it
|
||||
protected_api_client.delete(f"/api/download/{test_id}/cancel")
|
||||
|
||||
time.sleep(0.5)
|
||||
cancel_resp = protected_api_client.delete(f"/api/download/{test_id}/cancel")
|
||||
assert cancel_resp.status_code == 200
|
||||
assert cancel_resp.json() == {"status": "cancelled", "book_id": test_id}
|
||||
|
||||
# Check it's not in the queue
|
||||
queue_resp = protected_api_client.get("/api/queue/order")
|
||||
if queue_resp.status_code == 200:
|
||||
queue_order = queue_resp.json()
|
||||
assert test_id not in queue_order
|
||||
_wait_for_queue_absence(
|
||||
protected_api_client,
|
||||
test_id,
|
||||
context="cancel verification",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -376,10 +525,13 @@ class TestQueuePriority:
|
||||
},
|
||||
)
|
||||
|
||||
if queue_resp.status_code != 200:
|
||||
pytest.skip("Could not queue download")
|
||||
assert_queued_download_response(queue_resp)
|
||||
|
||||
time.sleep(0.5)
|
||||
_require_queue_entry(
|
||||
protected_api_client,
|
||||
test_id,
|
||||
context="priority update precondition",
|
||||
)
|
||||
|
||||
# Update priority
|
||||
priority_resp = protected_api_client.put(
|
||||
@@ -387,5 +539,11 @@ class TestQueuePriority:
|
||||
json={"priority": 10},
|
||||
)
|
||||
|
||||
# Should succeed or return 404 if already processed
|
||||
assert priority_resp.status_code in [200, 404]
|
||||
assert priority_resp.status_code == 200
|
||||
assert priority_resp.json() == {"status": "updated", "book_id": test_id, "priority": 10}
|
||||
|
||||
queue_resp = protected_api_client.get("/api/queue/order")
|
||||
queue_order = assert_queue_order_response(queue_resp)
|
||||
matching_entries = [entry for entry in queue_order if entry.get("id") == test_id]
|
||||
assert len(matching_entries) == 1
|
||||
assert matching_entries[0]["priority"] == 10
|
||||
|
||||
+199
-113
@@ -7,49 +7,147 @@ Requires Prowlarr and a download client (qBittorrent, Transmission, etc.) to be
|
||||
Run with: uv run pytest tests/e2e/test_prowlarr_flow.py -v -m e2e
|
||||
"""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from .conftest import APIClient, DownloadTracker
|
||||
from .conftest import (
|
||||
DOWNLOAD_TIMEOUT,
|
||||
SUCCESS_DOWNLOAD_STATES,
|
||||
APIClient,
|
||||
DownloadTracker,
|
||||
assert_queued_download_response,
|
||||
)
|
||||
|
||||
|
||||
def _assert_terminal_download_result(
|
||||
result: dict[str, object],
|
||||
*,
|
||||
source_id: str,
|
||||
expected_title: str,
|
||||
expected_source: str,
|
||||
) -> None:
|
||||
"""Assert that a finished Prowlarr download produced a structured payload."""
|
||||
state = result["state"]
|
||||
entry = result["data"]
|
||||
assert isinstance(entry, dict)
|
||||
if state == "error":
|
||||
error_message = str(
|
||||
entry.get("status_message") or entry.get("last_error_message") or ""
|
||||
).strip()
|
||||
normalized_error = error_message.lower()
|
||||
if any(
|
||||
marker in normalized_error
|
||||
for marker in (
|
||||
"failed to connect",
|
||||
"connection error",
|
||||
"connection refused",
|
||||
"name or service not known",
|
||||
"max retries exceeded",
|
||||
"timed out",
|
||||
)
|
||||
):
|
||||
pytest.skip(f"{expected_source} dependency unavailable: {error_message}")
|
||||
pytest.fail(
|
||||
f"{expected_source} download failed"
|
||||
f"{f': {error_message}' if error_message else f': {entry!r}'}"
|
||||
)
|
||||
|
||||
assert state in SUCCESS_DOWNLOAD_STATES, (
|
||||
f"{expected_source} ended in unexpected state {state!r}: {entry!r}"
|
||||
)
|
||||
assert entry.get("id") == source_id
|
||||
assert entry.get("title") == expected_title
|
||||
assert entry.get("source") == expected_source
|
||||
status = entry.get("status")
|
||||
assert status is None or status in SUCCESS_DOWNLOAD_STATES | {"queued"}
|
||||
|
||||
|
||||
def _is_duplicate_queue_error(response) -> bool:
|
||||
"""Whether the API refused to queue a release because it already exists."""
|
||||
if response.status_code != 500:
|
||||
return False
|
||||
try:
|
||||
payload = response.json()
|
||||
except ValueError:
|
||||
return False
|
||||
return payload == {"error": "Release is already in the download queue"}
|
||||
|
||||
|
||||
def _require_json_object(
|
||||
response, *, context: str, skip_statuses: set[int] | frozenset[int] = frozenset({503})
|
||||
) -> dict[str, object]:
|
||||
if response.status_code in skip_statuses:
|
||||
pytest.skip(f"{context} unavailable: {response.status_code}")
|
||||
assert response.status_code == 200, f"{context} failed: {response.status_code} {response.text}"
|
||||
payload = response.json()
|
||||
assert isinstance(payload, dict), f"{context} did not return a JSON object: {payload!r}"
|
||||
return payload
|
||||
|
||||
|
||||
def _extract_result_list(payload: object, *, context: str) -> list[dict[str, object]]:
|
||||
if isinstance(payload, dict):
|
||||
if "books" in payload:
|
||||
results = payload["books"]
|
||||
elif "releases" in payload:
|
||||
results = payload["releases"]
|
||||
else:
|
||||
results = payload.get("results", payload)
|
||||
else:
|
||||
results = payload
|
||||
|
||||
if isinstance(results, dict):
|
||||
list_values = [value for value in results.values() if isinstance(value, list)]
|
||||
if not list_values:
|
||||
pytest.fail(f"{context} returned an unexpected result structure: {payload!r}")
|
||||
if not any(list_values):
|
||||
pytest.skip(f"{context} returned no results")
|
||||
results = next(value for value in list_values if value)
|
||||
|
||||
if not isinstance(results, list):
|
||||
pytest.fail(f"{context} did not return a result list: {payload!r}")
|
||||
if not results:
|
||||
pytest.skip(f"{context} returned no results")
|
||||
return results
|
||||
|
||||
|
||||
def _is_prowlarr_configured(api_client: APIClient) -> bool:
|
||||
"""Check if Prowlarr is configured and available."""
|
||||
resp = api_client.get("/api/release-sources")
|
||||
if resp.status_code != 200:
|
||||
return False
|
||||
if resp.status_code == 503:
|
||||
pytest.skip("Release sources unavailable")
|
||||
assert resp.status_code == 200, f"Release sources failed: {resp.status_code} {resp.text}"
|
||||
|
||||
sources = resp.json()
|
||||
for source in sources:
|
||||
if source.get("name") == "prowlarr":
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _get_prowlarr_settings(api_client: APIClient) -> dict | None:
|
||||
"""Get Prowlarr settings if available."""
|
||||
resp = api_client.get("/api/settings/prowlarr")
|
||||
if resp.status_code == 200:
|
||||
return resp.json()
|
||||
return None
|
||||
if not isinstance(sources, list):
|
||||
pytest.fail(f"Release sources returned an unexpected payload: {sources!r}")
|
||||
if not all(isinstance(source, dict) for source in sources):
|
||||
pytest.fail(f"Release sources returned a malformed payload: {sources!r}")
|
||||
return any(source.get("name") == "prowlarr" for source in sources)
|
||||
|
||||
|
||||
def _get_first_provider_name(api_client: APIClient) -> str | None:
|
||||
"""Get the first available provider name."""
|
||||
providers_resp = api_client.get("/api/metadata/providers")
|
||||
if providers_resp.status_code != 200:
|
||||
return None
|
||||
if providers_resp.status_code == 503:
|
||||
pytest.skip("Metadata providers unavailable")
|
||||
assert providers_resp.status_code == 200, (
|
||||
f"Metadata providers failed: {providers_resp.status_code} {providers_resp.text}"
|
||||
)
|
||||
|
||||
providers_data = providers_resp.json()
|
||||
if not providers_data:
|
||||
return None
|
||||
|
||||
# Handle both dict and list formats
|
||||
if isinstance(providers_data, dict):
|
||||
return list(providers_data.keys())[0] if providers_data else None
|
||||
else:
|
||||
return providers_data[0].get("name") if providers_data else None
|
||||
if not isinstance(providers_data, (dict, list)):
|
||||
pytest.fail(f"Metadata providers returned an unexpected payload: {providers_data!r}")
|
||||
providers = (
|
||||
providers_data.get("providers", []) if isinstance(providers_data, dict) else providers_data
|
||||
)
|
||||
for provider in providers:
|
||||
if (
|
||||
isinstance(provider, dict)
|
||||
and provider.get("name")
|
||||
and provider.get("enabled") is True
|
||||
and provider.get("available") is True
|
||||
):
|
||||
return provider["name"]
|
||||
return None
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -62,6 +160,10 @@ class TestProwlarrConfiguration:
|
||||
|
||||
assert resp.status_code == 200
|
||||
sources = resp.json()
|
||||
assert isinstance(sources, list), f"Unexpected release sources payload: {sources!r}"
|
||||
assert all(isinstance(source, dict) for source in sources), (
|
||||
f"Unexpected release sources payload: {sources!r}"
|
||||
)
|
||||
source_names = [s.get("name") for s in sources]
|
||||
assert "prowlarr" in source_names
|
||||
|
||||
@@ -83,9 +185,11 @@ class TestProwlarrConfiguration:
|
||||
all_tab_names = []
|
||||
for group in data.get("groups", []):
|
||||
if isinstance(group, dict):
|
||||
for tab in group.get("tabs", []):
|
||||
if isinstance(tab, dict):
|
||||
all_tab_names.append(tab.get("name") or tab.get("id", ""))
|
||||
all_tab_names.extend(
|
||||
tab.get("name") or tab.get("id", "")
|
||||
for tab in group.get("tabs", [])
|
||||
if isinstance(tab, dict)
|
||||
)
|
||||
tab_names = all_tab_names
|
||||
else:
|
||||
tab_names = list(data.keys())
|
||||
@@ -124,30 +228,12 @@ class TestProwlarrSearch:
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
if search_resp.status_code != 200:
|
||||
pytest.skip("Metadata search unavailable")
|
||||
|
||||
search_data = search_resp.json()
|
||||
results = search_data.get("results", search_data)
|
||||
|
||||
# Handle dict format where results might be nested
|
||||
if isinstance(results, dict) and "results" not in results:
|
||||
# Results might be the actual result list under a different key
|
||||
for key, value in results.items():
|
||||
if isinstance(value, list) and value:
|
||||
results = value
|
||||
break
|
||||
|
||||
if not results or (isinstance(results, dict) and not results):
|
||||
pytest.skip("No metadata results")
|
||||
|
||||
# Get first result
|
||||
if isinstance(results, list):
|
||||
book = results[0]
|
||||
else:
|
||||
pytest.skip("Unexpected results format")
|
||||
search_data = _require_json_object(search_resp, context="metadata search")
|
||||
results = _extract_result_list(search_data, context="metadata search")
|
||||
book = results[0]
|
||||
|
||||
book_id = book.get("id") or book.get("provider_id")
|
||||
assert book_id, "Metadata search result missing ID"
|
||||
|
||||
# Now search releases specifically from Prowlarr
|
||||
releases_resp = protected_api_client.get(
|
||||
@@ -162,14 +248,13 @@ class TestProwlarrSearch:
|
||||
timeout=60,
|
||||
)
|
||||
|
||||
# Prowlarr may not be reachable
|
||||
if releases_resp.status_code == 503:
|
||||
pytest.skip("Prowlarr not reachable")
|
||||
|
||||
if releases_resp.status_code == 200:
|
||||
data = releases_resp.json()
|
||||
assert "releases" in data
|
||||
# Releases may be empty if Prowlarr has no indexers configured
|
||||
data = _require_json_object(releases_resp, context="Prowlarr release search")
|
||||
assert data.get("sources_searched") == ["prowlarr"]
|
||||
releases = data.get("releases")
|
||||
assert isinstance(releases, list), f"Unexpected Prowlarr release payload: {data!r}"
|
||||
if not releases:
|
||||
pytest.skip("No Prowlarr releases found")
|
||||
assert "book" in data
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -200,22 +285,30 @@ class TestProwlarrClientSettings:
|
||||
pytest.skip("Settings not available")
|
||||
|
||||
current = get_resp.json()
|
||||
assert isinstance(current, dict), f"Unexpected prowlarr_clients payload: {current!r}"
|
||||
fields = current.get("fields")
|
||||
assert isinstance(fields, list) and fields, (
|
||||
f"Prowlarr client settings payload missing fields: {current!r}"
|
||||
)
|
||||
|
||||
# Try to save the same settings back (no-op save)
|
||||
if isinstance(current, dict) and "fields" in current:
|
||||
# Extract just the values
|
||||
values = {}
|
||||
for field in current.get("fields", []):
|
||||
key = field.get("key") or field.get("name")
|
||||
if key:
|
||||
values[key] = field.get("value", "")
|
||||
values = {}
|
||||
for field in fields:
|
||||
if not isinstance(field, dict):
|
||||
continue
|
||||
key = field.get("key") or field.get("name")
|
||||
if key:
|
||||
values[key] = field.get("value", "")
|
||||
|
||||
put_resp = protected_api_client.put(
|
||||
"/api/settings/prowlarr_clients",
|
||||
json=values,
|
||||
)
|
||||
# Should succeed (200) or be a no-op
|
||||
assert put_resp.status_code in [200, 204, 400]
|
||||
assert values, f"No editable values found in prowlarr_clients payload: {current!r}"
|
||||
|
||||
put_resp = protected_api_client.put(
|
||||
"/api/settings/prowlarr_clients",
|
||||
json=values,
|
||||
)
|
||||
assert put_resp.status_code in [200, 204]
|
||||
if put_resp.status_code == 200:
|
||||
payload = put_resp.json()
|
||||
assert isinstance(payload, dict)
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -241,24 +334,11 @@ class TestProwlarrDownload:
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
if search_resp.status_code != 200:
|
||||
pytest.skip("Metadata search failed")
|
||||
|
||||
search_data = search_resp.json()
|
||||
results = search_data.get("results", search_data)
|
||||
|
||||
# Handle different result formats
|
||||
if isinstance(results, dict):
|
||||
for key, value in results.items():
|
||||
if isinstance(value, list) and value:
|
||||
results = value
|
||||
break
|
||||
|
||||
if not results or not isinstance(results, list):
|
||||
pytest.skip("No results")
|
||||
|
||||
search_data = _require_json_object(search_resp, context="metadata search")
|
||||
results = _extract_result_list(search_data, context="metadata search")
|
||||
book = results[0]
|
||||
book_id = book.get("id") or book.get("provider_id")
|
||||
assert book_id, "Metadata search result missing ID"
|
||||
|
||||
# Search Prowlarr releases
|
||||
releases_resp = protected_api_client.get(
|
||||
@@ -272,10 +352,9 @@ class TestProwlarrDownload:
|
||||
timeout=60,
|
||||
)
|
||||
|
||||
if releases_resp.status_code != 200:
|
||||
pytest.skip(f"Releases search failed: {releases_resp.status_code}")
|
||||
|
||||
releases = releases_resp.json().get("releases", [])
|
||||
releases_data = _require_json_object(releases_resp, context="Prowlarr release search")
|
||||
releases = releases_data.get("releases")
|
||||
assert isinstance(releases, list), f"Unexpected Prowlarr release payload: {releases_data!r}"
|
||||
if not releases:
|
||||
pytest.skip("No Prowlarr releases found")
|
||||
|
||||
@@ -297,24 +376,29 @@ class TestProwlarrDownload:
|
||||
},
|
||||
)
|
||||
|
||||
# May fail if no download client configured
|
||||
if queue_resp.status_code == 200:
|
||||
data = queue_resp.json()
|
||||
assert data.get("status") == "queued"
|
||||
|
||||
# Wait briefly and check status
|
||||
time.sleep(3)
|
||||
if not _is_duplicate_queue_error(queue_resp):
|
||||
assert_queued_download_response(queue_resp)
|
||||
|
||||
result = download_tracker.wait_for_status(
|
||||
source_id,
|
||||
target_states=["complete", "done", "available"],
|
||||
timeout=DOWNLOAD_TIMEOUT,
|
||||
)
|
||||
if result is None:
|
||||
status_resp = protected_api_client.get("/api/status")
|
||||
if status_resp.status_code == 200:
|
||||
status_data = status_resp.json()
|
||||
# Should be in one of the status categories
|
||||
found = False
|
||||
for category in status_data.values():
|
||||
if isinstance(category, dict) and source_id in category:
|
||||
found = True
|
||||
break
|
||||
# It's ok if not found (may have already processed/errored)
|
||||
error_info = status_data.get("error", {}).get(source_id)
|
||||
if error_info:
|
||||
pytest.fail(f"Prowlarr download failed: {error_info}")
|
||||
pytest.fail("Prowlarr download timed out")
|
||||
|
||||
_assert_terminal_download_result(
|
||||
result,
|
||||
source_id=source_id,
|
||||
expected_title=release.get("title", book.get("title", "Test")),
|
||||
expected_source="prowlarr",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.e2e
|
||||
@@ -328,11 +412,13 @@ class TestProwlarrClientConnection:
|
||||
"/api/settings/prowlarr_clients/action/test_torrent_connection"
|
||||
)
|
||||
|
||||
# May succeed, fail, or not exist depending on configuration
|
||||
# We just verify it returns a response
|
||||
assert resp.status_code in [200, 400, 404, 500]
|
||||
# May succeed, fail, or not exist depending on configuration.
|
||||
# A 500 is a real server error and should fail the test.
|
||||
assert resp.status_code in [200, 400, 404]
|
||||
|
||||
if resp.status_code == 200:
|
||||
data = resp.json()
|
||||
# Should have success/message structure
|
||||
assert "success" in data or "message" in data or "result" in data
|
||||
assert isinstance(data, dict)
|
||||
assert "success" in data or "message" in data
|
||||
if "message" in data:
|
||||
assert isinstance(data["message"], str)
|
||||
|
||||
@@ -9,6 +9,8 @@ from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
||||
def _as_response(result: Any):
|
||||
if isinstance(result, tuple) and len(result) == 2:
|
||||
@@ -36,52 +38,79 @@ def main_module():
|
||||
|
||||
class TestProxyAuthMiddleware:
|
||||
def test_skips_for_non_proxy_mode(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="builtin"):
|
||||
with main_module.app.test_request_context("/api/releases"):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert "user_id" not in main_module.session
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="builtin"),
|
||||
main_module.app.test_request_context("/api/releases"),
|
||||
):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert "user_id" not in main_module.session
|
||||
|
||||
def test_skips_health_endpoint(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with main_module.app.test_request_context("/api/health"):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
main_module.app.test_request_context("/api/health"),
|
||||
):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
|
||||
def test_allows_auth_check_without_header(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with patch.object(
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({"PROXY_AUTH_USER_HEADER": "X-Auth-User"}),
|
||||
):
|
||||
with main_module.app.test_request_context("/api/auth/check"):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert "user_id" not in main_module.session
|
||||
),
|
||||
main_module.app.test_request_context("/api/auth/check"),
|
||||
):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert "user_id" not in main_module.session
|
||||
|
||||
def test_sets_session_from_header(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with patch.object(
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({"PROXY_AUTH_USER_HEADER": "X-Auth-User"}),
|
||||
):
|
||||
with main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
headers={"X-Auth-User": "proxyuser"},
|
||||
):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert main_module.session.get("user_id") == "proxyuser"
|
||||
assert main_module.session.get("is_admin") is True
|
||||
db_user_id = main_module.session.get("db_user_id")
|
||||
assert db_user_id is not None
|
||||
db_user = main_module.user_db.get_user(user_id=db_user_id)
|
||||
assert db_user is not None
|
||||
assert db_user["username"] == "proxyuser"
|
||||
assert db_user["auth_source"] == "proxy"
|
||||
assert main_module.session.permanent is False
|
||||
),
|
||||
main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
headers={"X-Auth-User": "proxyuser"},
|
||||
),
|
||||
):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert main_module.session.get("user_id") == "proxyuser"
|
||||
assert main_module.session.get("is_admin") is True
|
||||
db_user_id = main_module.session.get("db_user_id")
|
||||
assert db_user_id is not None
|
||||
db_user = main_module.user_db.get_user(user_id=db_user_id)
|
||||
assert db_user is not None
|
||||
assert db_user["username"] == "proxyuser"
|
||||
assert db_user["auth_source"] == "proxy"
|
||||
assert main_module.session.permanent is False
|
||||
|
||||
def test_reads_remote_user_wsgi_fallback(self, main_module):
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({"PROXY_AUTH_USER_HEADER": "Remote-User"}),
|
||||
),
|
||||
main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
environ_base={"REMOTE_USER": "proxyremote"},
|
||||
),
|
||||
):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert main_module.session.get("user_id") == "proxyremote"
|
||||
assert main_module.session.get("is_admin") is True
|
||||
assert main_module.session.permanent is False
|
||||
|
||||
def test_proxy_takes_over_existing_local_username(self, main_module):
|
||||
existing = main_module.user_db.create_user(
|
||||
@@ -90,76 +119,82 @@ class TestProxyAuthMiddleware:
|
||||
auth_source="builtin",
|
||||
)
|
||||
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with patch.object(
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({"PROXY_AUTH_USER_HEADER": "X-Auth-User"}),
|
||||
):
|
||||
with main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
headers={"X-Auth-User": "proxy_takeover_local"},
|
||||
):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
),
|
||||
main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
headers={"X-Auth-User": "proxy_takeover_local"},
|
||||
),
|
||||
):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
|
||||
db_user_id = main_module.session.get("db_user_id")
|
||||
db_user = main_module.user_db.get_user(user_id=db_user_id)
|
||||
assert db_user is not None
|
||||
assert db_user["id"] == existing["id"]
|
||||
assert db_user["username"] == "proxy_takeover_local"
|
||||
assert db_user["auth_source"] == "proxy"
|
||||
db_user_id = main_module.session.get("db_user_id")
|
||||
db_user = main_module.user_db.get_user(user_id=db_user_id)
|
||||
assert db_user is not None
|
||||
assert db_user["id"] == existing["id"]
|
||||
assert db_user["username"] == "proxy_takeover_local"
|
||||
assert db_user["auth_source"] == "proxy"
|
||||
|
||||
def test_reprovisions_when_proxy_identity_changes(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with patch.object(
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({"PROXY_AUTH_USER_HEADER": "X-Auth-User"}),
|
||||
):
|
||||
with main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
headers={"X-Auth-User": "proxyuser2"},
|
||||
):
|
||||
main_module.session["user_id"] = "old-user"
|
||||
main_module.session["db_user_id"] = 999999
|
||||
),
|
||||
main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
headers={"X-Auth-User": "proxyuser2"},
|
||||
),
|
||||
):
|
||||
main_module.session["user_id"] = "old-user"
|
||||
main_module.session["db_user_id"] = 999999
|
||||
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert main_module.session.get("user_id") == "proxyuser2"
|
||||
db_user_id = main_module.session.get("db_user_id")
|
||||
db_user = main_module.user_db.get_user(user_id=db_user_id)
|
||||
assert db_user["username"] == "proxyuser2"
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert main_module.session.get("user_id") == "proxyuser2"
|
||||
db_user_id = main_module.session.get("db_user_id")
|
||||
db_user = main_module.user_db.get_user(user_id=db_user_id)
|
||||
assert db_user["username"] == "proxyuser2"
|
||||
|
||||
def test_reprovisions_when_session_db_user_is_stale(self, main_module):
|
||||
stale_user_id = 99999999
|
||||
username = f"proxy_stale_{uuid4().hex[:8]}"
|
||||
assert main_module.user_db.get_user(user_id=stale_user_id) is None
|
||||
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with patch.object(
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({"PROXY_AUTH_USER_HEADER": "X-Auth-User"}),
|
||||
):
|
||||
with main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
headers={"X-Auth-User": username},
|
||||
):
|
||||
main_module.session["user_id"] = username
|
||||
main_module.session["db_user_id"] = stale_user_id
|
||||
),
|
||||
main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
headers={"X-Auth-User": username},
|
||||
),
|
||||
):
|
||||
main_module.session["user_id"] = username
|
||||
main_module.session["db_user_id"] = stale_user_id
|
||||
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert main_module.session.get("user_id") == username
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert main_module.session.get("user_id") == username
|
||||
|
||||
db_user_id = main_module.session.get("db_user_id")
|
||||
assert db_user_id is not None
|
||||
assert db_user_id != stale_user_id
|
||||
db_user_id = main_module.session.get("db_user_id")
|
||||
assert db_user_id is not None
|
||||
assert db_user_id != stale_user_id
|
||||
|
||||
db_user = main_module.user_db.get_user(user_id=db_user_id)
|
||||
assert db_user is not None
|
||||
assert db_user["username"] == username
|
||||
db_user = main_module.user_db.get_user(user_id=db_user_id)
|
||||
assert db_user is not None
|
||||
assert db_user["username"] == username
|
||||
|
||||
def test_reprovisions_when_session_db_user_points_to_other_username(self, main_module):
|
||||
username = f"proxy_target_{uuid4().hex[:8]}"
|
||||
@@ -169,66 +204,74 @@ class TestProxyAuthMiddleware:
|
||||
auth_source="proxy",
|
||||
)
|
||||
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with patch.object(
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({"PROXY_AUTH_USER_HEADER": "X-Auth-User"}),
|
||||
):
|
||||
with main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
headers={"X-Auth-User": username},
|
||||
):
|
||||
main_module.session["user_id"] = username
|
||||
main_module.session["db_user_id"] = other_user["id"]
|
||||
),
|
||||
main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
headers={"X-Auth-User": username},
|
||||
),
|
||||
):
|
||||
main_module.session["user_id"] = username
|
||||
main_module.session["db_user_id"] = other_user["id"]
|
||||
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert main_module.session.get("user_id") == username
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert main_module.session.get("user_id") == username
|
||||
|
||||
db_user_id = main_module.session.get("db_user_id")
|
||||
assert db_user_id is not None
|
||||
assert db_user_id != other_user["id"]
|
||||
db_user_id = main_module.session.get("db_user_id")
|
||||
assert db_user_id is not None
|
||||
assert db_user_id != other_user["id"]
|
||||
|
||||
db_user = main_module.user_db.get_user(user_id=db_user_id)
|
||||
assert db_user is not None
|
||||
assert db_user["username"] == username
|
||||
db_user = main_module.user_db.get_user(user_id=db_user_id)
|
||||
assert db_user is not None
|
||||
assert db_user["username"] == username
|
||||
|
||||
def test_returns_401_when_header_missing_on_protected_path(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with patch.object(
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({"PROXY_AUTH_USER_HEADER": "X-Auth-User"}),
|
||||
):
|
||||
with main_module.app.test_request_context("/api/releases"):
|
||||
resp = _as_response(main_module.proxy_auth_middleware())
|
||||
data = resp.get_json()
|
||||
),
|
||||
main_module.app.test_request_context("/api/releases"),
|
||||
):
|
||||
resp = _as_response(main_module.proxy_auth_middleware())
|
||||
data = resp.get_json()
|
||||
|
||||
assert resp.status_code == 401
|
||||
assert "Authentication required" in (data.get("error") or "")
|
||||
assert data == {"error": "Authentication required. Proxy header not set."}
|
||||
|
||||
def test_admin_group_membership(self, main_module):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with patch.object(
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
patch.object(
|
||||
main_module.app_config,
|
||||
"get",
|
||||
side_effect=_config_getter({
|
||||
"PROXY_AUTH_USER_HEADER": "X-Auth-User",
|
||||
"PROXY_AUTH_ADMIN_GROUP_HEADER": "X-Auth-Groups",
|
||||
"PROXY_AUTH_ADMIN_GROUP_NAME": "admins",
|
||||
}),
|
||||
):
|
||||
with main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
headers={
|
||||
"X-Auth-User": "adminuser",
|
||||
"X-Auth-Groups": "users,admins,devs",
|
||||
},
|
||||
):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert main_module.session.get("is_admin") is True
|
||||
side_effect=_config_getter(
|
||||
{
|
||||
"PROXY_AUTH_USER_HEADER": "X-Auth-User",
|
||||
"PROXY_AUTH_ADMIN_GROUP_HEADER": "X-Auth-Groups",
|
||||
"PROXY_AUTH_ADMIN_GROUP_NAME": "admins",
|
||||
}
|
||||
),
|
||||
),
|
||||
main_module.app.test_request_context(
|
||||
"/api/releases",
|
||||
headers={
|
||||
"X-Auth-User": "adminuser",
|
||||
"X-Auth-Groups": "users,admins,devs",
|
||||
},
|
||||
),
|
||||
):
|
||||
result = main_module.proxy_auth_middleware()
|
||||
assert result is None
|
||||
assert main_module.session.get("is_admin") is True
|
||||
|
||||
|
||||
class TestLoginRequiredDecorator:
|
||||
@@ -240,84 +283,100 @@ class TestLoginRequiredDecorator:
|
||||
return _view
|
||||
|
||||
def test_allows_no_auth(self, main_module, view):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="none"):
|
||||
with main_module.app.test_request_context("/api/releases"):
|
||||
decorated = main_module.login_required(view)
|
||||
resp = decorated()
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="none"),
|
||||
main_module.app.test_request_context("/api/releases"),
|
||||
):
|
||||
decorated = main_module.login_required(view)
|
||||
resp = decorated()
|
||||
|
||||
assert resp[0]["success"] is True
|
||||
|
||||
def test_blocks_when_not_authenticated(self, main_module, view):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="builtin"):
|
||||
with main_module.app.test_request_context("/api/releases"):
|
||||
decorated = main_module.login_required(view)
|
||||
resp = _as_response(decorated())
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="builtin"),
|
||||
main_module.app.test_request_context("/api/releases"),
|
||||
):
|
||||
decorated = main_module.login_required(view)
|
||||
resp = _as_response(decorated())
|
||||
|
||||
assert resp.status_code == 401
|
||||
|
||||
def test_allows_authenticated(self, main_module, view):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="builtin"):
|
||||
with main_module.app.test_request_context("/api/releases"):
|
||||
main_module.session["user_id"] = "user"
|
||||
decorated = main_module.login_required(view)
|
||||
resp = decorated()
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="builtin"),
|
||||
main_module.app.test_request_context("/api/releases"),
|
||||
):
|
||||
main_module.session["user_id"] = "user"
|
||||
decorated = main_module.login_required(view)
|
||||
resp = decorated()
|
||||
|
||||
assert resp[0]["success"] is True
|
||||
|
||||
def test_settings_access_requires_admin_even_when_legacy_toggle_off(self, main_module, view):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="builtin"):
|
||||
with main_module.app.test_request_context("/api/settings/general"):
|
||||
main_module.session["user_id"] = "user"
|
||||
main_module.session["is_admin"] = False
|
||||
decorated = main_module.login_required(view)
|
||||
resp = _as_response(decorated())
|
||||
data = resp.get_json()
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="builtin"),
|
||||
main_module.app.test_request_context("/api/settings/general"),
|
||||
):
|
||||
main_module.session["user_id"] = "user"
|
||||
main_module.session["is_admin"] = False
|
||||
decorated = main_module.login_required(view)
|
||||
resp = _as_response(decorated())
|
||||
data = resp.get_json()
|
||||
|
||||
assert resp.status_code == 403
|
||||
assert "Admin access required" in (data.get("error") or "")
|
||||
|
||||
def test_security_tab_always_blocks_non_admin_even_when_toggle_off(self, main_module, view):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="builtin"):
|
||||
with main_module.app.test_request_context("/api/settings/security"):
|
||||
main_module.session["user_id"] = "user"
|
||||
main_module.session["is_admin"] = False
|
||||
decorated = main_module.login_required(view)
|
||||
resp = _as_response(decorated())
|
||||
data = resp.get_json()
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="builtin"),
|
||||
main_module.app.test_request_context("/api/settings/security"),
|
||||
):
|
||||
main_module.session["user_id"] = "user"
|
||||
main_module.session["is_admin"] = False
|
||||
decorated = main_module.login_required(view)
|
||||
resp = _as_response(decorated())
|
||||
data = resp.get_json()
|
||||
|
||||
assert resp.status_code == 403
|
||||
assert "Admin access required" in (data.get("error") or "")
|
||||
|
||||
def test_users_tab_always_blocks_non_admin_even_when_toggle_off(self, main_module, view):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="builtin"):
|
||||
with main_module.app.test_request_context("/api/settings/users"):
|
||||
main_module.session["user_id"] = "user"
|
||||
main_module.session["is_admin"] = False
|
||||
decorated = main_module.login_required(view)
|
||||
resp = _as_response(decorated())
|
||||
data = resp.get_json()
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="builtin"),
|
||||
main_module.app.test_request_context("/api/settings/users"),
|
||||
):
|
||||
main_module.session["user_id"] = "user"
|
||||
main_module.session["is_admin"] = False
|
||||
decorated = main_module.login_required(view)
|
||||
resp = _as_response(decorated())
|
||||
data = resp.get_json()
|
||||
|
||||
assert resp.status_code == 403
|
||||
assert "Admin access required" in (data.get("error") or "")
|
||||
|
||||
def test_proxy_admin_restriction_blocks_non_admin(self, main_module, view):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="proxy"):
|
||||
with main_module.app.test_request_context("/api/settings/general"):
|
||||
main_module.session["user_id"] = "user"
|
||||
main_module.session["is_admin"] = False
|
||||
decorated = main_module.login_required(view)
|
||||
resp = _as_response(decorated())
|
||||
data = resp.get_json()
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="proxy"),
|
||||
main_module.app.test_request_context("/api/settings/general"),
|
||||
):
|
||||
main_module.session["user_id"] = "user"
|
||||
main_module.session["is_admin"] = False
|
||||
decorated = main_module.login_required(view)
|
||||
resp = _as_response(decorated())
|
||||
data = resp.get_json()
|
||||
|
||||
assert resp.status_code == 403
|
||||
assert "Admin access required" in (data.get("error") or "")
|
||||
|
||||
def test_cwa_admin_restriction_blocks_non_admin(self, main_module, view):
|
||||
with patch.object(main_module, "get_auth_mode", return_value="cwa"):
|
||||
with main_module.app.test_request_context("/api/settings/general"):
|
||||
main_module.session["user_id"] = "user"
|
||||
main_module.session["is_admin"] = False
|
||||
decorated = main_module.login_required(view)
|
||||
resp = _as_response(decorated())
|
||||
with (
|
||||
patch.object(main_module, "get_auth_mode", return_value="cwa"),
|
||||
main_module.app.test_request_context("/api/settings/general"),
|
||||
):
|
||||
main_module.session["user_id"] = "user"
|
||||
main_module.session["is_admin"] = False
|
||||
decorated = main_module.login_required(view)
|
||||
resp = _as_response(decorated())
|
||||
|
||||
assert resp.status_code == 403
|
||||
|
||||
Reference in New Issue
Block a user