Misc fixes (#718)

- Update file movement to prefer copy
- Improved mirror config overwriting on app updates
- Request / user DB hardening
This commit is contained in:
Alex
2026-03-07 10:30:47 +00:00
committed by GitHub
parent edb437e905
commit 80aa289a64
16 changed files with 425 additions and 119 deletions
+11 -4
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import os
import sqlite3
import threading
from datetime import datetime
from datetime import datetime, timezone
from typing import Any
from shelfmark.core.logger import setup_logger
@@ -115,9 +115,12 @@ class DownloadHistoryService:
return None
normalized = value.strip().replace("Z", "+00:00")
try:
return datetime.fromisoformat(normalized).timestamp()
parsed = datetime.fromisoformat(normalized)
except ValueError:
return None
if parsed.tzinfo is None:
parsed = parsed.replace(tzinfo=timezone.utc)
return parsed.timestamp()
@classmethod
def to_history_row(cls, row: dict[str, Any], *, dismissed_at: str) -> dict[str, Any]:
@@ -177,6 +180,7 @@ class DownloadHistoryService:
if normalized_title is None:
raise ValueError("title must be a non-empty string")
normalized_origin = _normalize_origin(origin)
recorded_at = now_utc_iso()
with self._lock:
conn = self._connect()
@@ -191,12 +195,12 @@ class DownloadHistoryService:
status_message, download_path,
queued_at, terminal_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'active', NULL, NULL, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'active', NULL, NULL, ?, ?)
ON CONFLICT(task_id) DO UPDATE SET
final_status = 'active',
status_message = NULL,
download_path = NULL,
terminal_at = CURRENT_TIMESTAMP
terminal_at = ?
""",
(
normalized_task_id,
@@ -212,6 +216,9 @@ class DownloadHistoryService:
normalize_optional_text(preview),
normalize_optional_text(content_type),
normalized_origin,
recorded_at,
recorded_at,
recorded_at,
),
)
conn.commit()
+65 -31
View File
@@ -17,7 +17,7 @@ from shelfmark.core.request_validation import (
validate_request_level_payload,
validate_status_transition,
)
from shelfmark.core.request_helpers import extract_release_source_id
from shelfmark.core.request_helpers import extract_release_source_id, normalize_positive_int
MAX_REQUEST_NOTE_LENGTH = 1000
@@ -139,26 +139,44 @@ def sync_delivery_states_from_queue_status(
user_id: int | None = None,
) -> list[dict[str, Any]]:
"""Persist delivery-state transitions for fulfilled requests based on queue status."""
source_delivery_states: dict[str, str] = {}
for status_key in QueueStatus:
status_bucket = queue_status.get(status_key)
if not isinstance(status_bucket, dict):
continue
for source_id in status_bucket:
source_delivery_states[source_id] = status_key
if not source_delivery_states:
fulfilled_rows = user_db.list_requests(user_id=user_id, status=RequestStatus.FULFILLED)
if not fulfilled_rows:
return []
fulfilled_rows = user_db.list_requests(user_id=user_id, status=RequestStatus.FULFILLED)
updated: list[dict[str, Any]] = []
unique_request_ids_by_source: dict[str, int] = {}
ambiguous_source_ids: set[str] = set()
for row in fulfilled_rows:
source_id = extract_release_source_id(row.get("release_data"))
if source_id is None:
continue
if source_id in unique_request_ids_by_source:
ambiguous_source_ids.add(source_id)
continue
unique_request_ids_by_source[source_id] = int(row["id"])
for source_id in ambiguous_source_ids:
unique_request_ids_by_source.pop(source_id, None)
delivery_state = source_delivery_states.get(source_id)
request_delivery_states: dict[int, str] = {}
for status_key in QueueStatus:
status_bucket = queue_status.get(status_key)
if not isinstance(status_bucket, dict):
continue
for source_id, task_payload in status_bucket.items():
request_id = None
if isinstance(task_payload, dict):
request_id = normalize_positive_int(task_payload.get("request_id"))
if request_id is None:
request_id = unique_request_ids_by_source.get(str(source_id).strip())
if request_id is None:
continue
request_delivery_states[request_id] = status_key
if not request_delivery_states:
return []
updated: list[dict[str, Any]] = []
for row in fulfilled_rows:
delivery_state = request_delivery_states.get(int(row["id"]))
if delivery_state is None:
continue
@@ -385,24 +403,9 @@ def fulfil_request(
if requester is None:
raise RequestServiceError("Requesting user not found", status_code=404)
queued_release_data = dict(selected_release_data)
queued_release_data["_request_id"] = request_id
success, error = queue_release(
queued_release_data,
0,
user_id=request_row["user_id"],
username=requester.get("username"),
)
if not success:
raise RequestServiceError(
error or "Failed to queue release",
status_code=409,
code="queue_failed",
)
original_release_data = request_row.get("release_data")
try:
return user_db.update_request(
claimed_request = user_db.update_request(
request_id,
expected_current_status=RequestStatus.PENDING,
status=RequestStatus.FULFILLED,
@@ -417,6 +420,37 @@ def fulfil_request(
except ValueError as exc:
raise RequestServiceError(str(exc), status_code=409, code="stale_transition") from exc
queued_release_data = dict(selected_release_data)
queued_release_data["_request_id"] = request_id
try:
success, error = queue_release(
queued_release_data,
0,
user_id=request_row["user_id"],
username=requester.get("username"),
)
except Exception:
user_db.rollback_request_fulfilment(
request_id,
release_data=original_release_data,
last_failure_reason="Queue dispatch raised an exception",
)
raise
if not success:
user_db.rollback_request_fulfilment(
request_id,
release_data=original_release_data,
last_failure_reason=error,
)
raise RequestServiceError(
error or "Failed to queue release",
status_code=409,
code="queue_failed",
)
return claimed_request
def reopen_failed_request(
user_db: "UserDB",
+58 -14
View File
@@ -486,22 +486,20 @@ def sync_env_to_config() -> None:
def migrate_mirror_settings() -> None:
"""
Migrate legacy AA mirror config into the new editable mirror list setting.
Sync AA mirror list when code defaults change between versions.
Legacy:
- AA_ADDITIONAL_URLS: comma-separated extra URLs appended to defaults
On startup, compares a hash of DEFAULT_AA_MIRRORS against the hash stored
in the config file. If they differ (i.e., an update shipped new defaults),
the config is overwritten with the new defaults. If they match, the user's
customizations are left untouched.
New:
- AA_MIRROR_URLS: full ordered list of available mirrors (used for Auto mode and for Settings options)
Also handles legacy migration from AA_ADDITIONAL_URLS.
"""
mirrors_config = load_config_file("mirrors")
import hashlib
from shelfmark.core.mirrors import DEFAULT_AA_MIRRORS
from shelfmark.core.utils import normalize_http_url
raw_list = mirrors_config.get("AA_MIRROR_URLS")
raw_additional = mirrors_config.get("AA_ADDITIONAL_URLS", "")
def _normalize_list(values: list[str]) -> list[str]:
out: list[str] = []
for item in values:
@@ -512,12 +510,42 @@ def migrate_mirror_settings() -> None:
out.append(norm)
return out
def _hash_mirrors(mirrors: list[str]) -> str:
return hashlib.sha256(",".join(mirrors).encode()).hexdigest()
normalized_defaults = _normalize_list(DEFAULT_AA_MIRRORS)
current_defaults_hash = _hash_mirrors(normalized_defaults)
mirrors_config = load_config_file("mirrors")
stored_hash = mirrors_config.get("_AA_MIRRORS_DEFAULTS_HASH")
raw_list = mirrors_config.get("AA_MIRROR_URLS")
raw_additional = mirrors_config.get("AA_ADDITIONAL_URLS", "")
def _save_mirrors(values: dict[str, Any]) -> None:
merged = dict(mirrors_config)
merged.update(values)
save_config_file("mirrors", merged)
mirrors_config.update(values)
# Defaults changed since last startup — push new mirrors to config
if stored_hash != current_defaults_hash:
_save_mirrors({
"AA_MIRROR_URLS": normalized_defaults,
"_AA_MIRRORS_DEFAULTS_HASH": current_defaults_hash,
})
return
# --- Legacy migration (only runs if hash already matches / first time) ---
# If already a proper list, just ensure it's non-empty.
if isinstance(raw_list, list):
normalized = _normalize_list([str(v) for v in raw_list])
if normalized:
return
save_config_file("mirrors", {"AA_MIRROR_URLS": _normalize_list(DEFAULT_AA_MIRRORS)})
_save_mirrors({
"AA_MIRROR_URLS": normalized_defaults,
"_AA_MIRRORS_DEFAULTS_HASH": current_defaults_hash,
})
return
# If saved as a string, convert to list.
@@ -525,17 +553,33 @@ def migrate_mirror_settings() -> None:
parts = [p.strip() for p in raw_list.split(",") if p.strip()]
normalized = _normalize_list(parts)
if normalized:
save_config_file("mirrors", {"AA_MIRROR_URLS": normalized})
_save_mirrors({
"AA_MIRROR_URLS": normalized,
"_AA_MIRRORS_DEFAULTS_HASH": current_defaults_hash,
})
return
save_config_file("mirrors", {"AA_MIRROR_URLS": _normalize_list(DEFAULT_AA_MIRRORS)})
_save_mirrors({
"AA_MIRROR_URLS": normalized_defaults,
"_AA_MIRRORS_DEFAULTS_HASH": current_defaults_hash,
})
return
# If there's legacy additional mirrors, seed the full list so the UI reflects reality.
# If there's legacy additional mirrors, seed the full list.
if isinstance(raw_additional, str) and raw_additional.strip():
additional_parts = [p.strip() for p in raw_additional.split(",") if p.strip()]
combined = _normalize_list(DEFAULT_AA_MIRRORS + additional_parts)
if combined:
save_config_file("mirrors", {"AA_MIRROR_URLS": combined})
_save_mirrors({
"AA_MIRROR_URLS": combined,
"_AA_MIRRORS_DEFAULTS_HASH": current_defaults_hash,
})
return
# No config at all yet — write defaults
_save_mirrors({
"AA_MIRROR_URLS": normalized_defaults,
"_AA_MIRRORS_DEFAULTS_HASH": current_defaults_hash,
})
def migrate_legacy_settings() -> None:
+58
View File
@@ -137,6 +137,14 @@ def sync_builtin_admin_user(
existing = user_db.get_user(username=normalized_username)
if existing:
existing_auth_source = str(existing.get("auth_source") or AUTH_SOURCE_BUILTIN).strip().lower()
if existing_auth_source != AUTH_SOURCE_BUILTIN:
logger.warning(
"Skipped builtin admin sync for username '%s' because it belongs to auth_source='%s'",
normalized_username,
existing_auth_source,
)
return
updates: dict[str, Any] = {}
if existing.get("password_hash") != normalized_hash:
updates["password_hash"] = normalized_hash
@@ -777,6 +785,56 @@ class UserDB:
finally:
conn.close()
def rollback_request_fulfilment(
self,
request_id: int,
*,
release_data: Optional[Dict[str, Any]],
last_failure_reason: Optional[str] = None,
) -> Dict[str, Any]:
"""Restore a request to pending after fulfilment claimed it but queueing failed."""
with self._lock:
conn = self._connect()
try:
row = conn.execute(
"SELECT * FROM download_requests WHERE id = ?",
(request_id,),
).fetchone()
current = self._parse_request_row(row)
if current is None:
raise ValueError(f"Request {request_id} not found")
conn.execute(
"""
UPDATE download_requests
SET status = 'pending',
release_data = ?,
admin_note = NULL,
reviewed_by = NULL,
reviewed_at = NULL,
delivery_state = 'none',
delivery_updated_at = NULL,
last_failure_reason = ?
WHERE id = ?
""",
(
self._serialize_json(release_data, "release_data"),
last_failure_reason,
request_id,
),
)
updated_row = conn.execute(
"SELECT * FROM download_requests WHERE id = ?",
(request_id,),
).fetchone()
conn.commit()
parsed = self._parse_request_row(updated_row)
if parsed is None:
raise ValueError(f"Request {request_id} not found after rollback")
return parsed
finally:
conn.close()
def count_pending_requests(self) -> int:
"""Count all pending requests."""
conn = self._connect()
+26 -33
View File
@@ -263,10 +263,12 @@ def _hardlink_not_supported(error: OSError) -> bool:
return err in {
errno.EXDEV,
errno.EMLINK,
errno.EIO,
errno.EPERM,
errno.EACCES,
getattr(errno, "ENOTSUP", errno.EPERM),
getattr(errno, "EOPNOTSUPP", errno.EPERM),
getattr(errno, "ENOSYS", errno.EPERM),
errno.EINVAL,
}
@@ -287,43 +289,33 @@ def _publish_temp_file(temp_path: Path, dest_path: Path) -> bool:
Returns True on success, False if the destination already exists.
"""
claimed = _claim_destination(dest_path)
if not claimed:
return False
try:
run_blocking_io(os.link, str(temp_path), str(dest_path))
# Trigger IN_CLOSE_WRITE so inotify-based file watchers (e.g. CWA) detect the new file.
# os.link() only generates IN_CREATE which many watchers don't monitor.
# Publish by renaming the fully-written temp file into place. This gives
# watchers an IN_MOVED_TO-style event on the final path instead of relying
# on hardlink support in the destination filesystem.
run_blocking_io(os.replace, str(temp_path), str(dest_path))
# Best-effort nudge for watchers that only react to close-write on the
# final filename rather than rename/move events.
try:
fd = run_blocking_io(os.open, str(dest_path), os.O_WRONLY)
run_blocking_io(os.close, fd)
except OSError:
pass
run_blocking_io(temp_path.unlink, missing_ok=True)
return True
except FileExistsError:
return False
except OSError as e:
except Exception as e:
if _is_permission_error(e):
log_transfer_permission_context(
"publish_hardlink",
"publish_replace",
source=temp_path,
dest=dest_path,
error=e,
)
if _hardlink_not_supported(e):
logger.debug(
"Hardlink publish unsupported; falling back to claim+replace: %s -> %s (%s)",
temp_path,
dest_path,
e,
)
claimed = _claim_destination(dest_path)
if not claimed:
return False
try:
run_blocking_io(os.replace, str(temp_path), str(dest_path))
except Exception:
run_blocking_io(dest_path.unlink, missing_ok=True)
raise
return True
run_blocking_io(dest_path.unlink, missing_ok=True)
raise
@@ -505,14 +497,15 @@ def atomic_hardlink(source_path: Path, dest_path: Path, max_attempts: int = 100)
except FileExistsError:
continue
except OSError as e:
if _is_permission_error(e) or e.errno in (errno.EXDEV, errno.EMLINK):
if _is_permission_error(e):
log_transfer_permission_context(
"atomic_hardlink",
source=source_path,
dest=try_path,
error=e,
)
permission_error = _is_permission_error(e)
if permission_error:
log_transfer_permission_context(
"atomic_hardlink",
source=source_path,
dest=try_path,
error=e,
)
if permission_error or _hardlink_not_supported(e):
logger.debug(
"Hardlink failed (%s), falling back to copy: %s -> %s",
e,
@@ -528,7 +521,7 @@ def atomic_hardlink(source_path: Path, dest_path: Path, max_attempts: int = 100)
def atomic_copy(source_path: Path, dest_path: Path, max_attempts: int = 100) -> Path:
"""Copy a file with atomic collision detection.
Uses a temp file in the destination directory and publishes it atomically,
Uses a temp file in the destination directory and publishes it via rename,
avoiding partial files on failure.
Args:
+1 -1
View File
@@ -581,7 +581,7 @@ def _compute_search_title(
if match:
suffix = _strip_parenthetical_suffix(match.group(2).strip())
if normalized_subtitle.lower() == suffix.lower() or normalized_subtitle.lower() in suffix.lower():
return match.group(1).strip()
return None
# Prefer subtitle when it looks like the real title.
if normalized_subtitle and not _is_probably_series_position(normalized_subtitle):
+2 -2
View File
@@ -4,8 +4,8 @@ def test_aa_base_url_options_include_configured_custom_url(monkeypatch):
# Use a custom mirror not present in defaults/additional.
monkeypatch.setenv("AA_BASE_URL", "https://custom-aa.example")
config_obj.refresh()
monkeypatch.delattr(config_obj, "_env_synced", raising=False)
config_obj.refresh(force=True)
options = settings._get_aa_base_url_options()
assert any(opt["value"] == "https://custom-aa.example" for opt in options)
+3 -3
View File
@@ -140,9 +140,9 @@ class TestOIDCFieldShowWhen:
class TestOIDCFieldsEnvSupport:
"""Tests that OIDC fields are UI-only (no env var support)."""
"""Tests that OIDC fields support env configuration."""
def test_oidc_fields_not_env_supported(self):
def test_oidc_fields_env_supported(self):
fields = _reload_security_module()
oidc_keys = [
"OIDC_DISCOVERY_URL",
@@ -157,4 +157,4 @@ class TestOIDCFieldsEnvSupport:
for key in oidc_keys:
field = _get_field(fields, key)
assert field is not None, f"Field {key} not found"
assert field.env_supported is False, f"Field {key} should not support env vars"
assert field.env_supported is True, f"Field {key} should support env vars"
+6 -3
View File
@@ -267,8 +267,8 @@ class TestSecurityMigration:
class TestSecuritySettings:
"""Tests for security settings registration."""
def test_security_settings_without_cwa(self):
"""CWA option should be hidden when DB is unavailable."""
def test_security_settings_without_cwa_shows_warning_but_keeps_option(self):
"""CWA remains selectable but warns when the DB is unavailable."""
with patch("shelfmark.config.env.CWA_DB_PATH", None):
import importlib
import shelfmark.config.security
@@ -284,7 +284,10 @@ class TestSecuritySettings:
assert "none" in option_values
assert "builtin" in option_values
assert "proxy" in option_values
assert "cwa" not in option_values
assert "cwa" in option_values
cwa_warning_field = next((f for f in fields if f.key == "cwa_db_missing"), None)
assert cwa_warning_field is not None
def test_security_settings_with_cwa(self):
"""CWA option should be shown when DB is mounted."""
+29
View File
@@ -0,0 +1,29 @@
"""Tests for syncing builtin-admin credentials into the users DB."""
from __future__ import annotations
import os
import tempfile
from shelfmark.core.user_db import UserDB, sync_builtin_admin_user
def test_sync_builtin_admin_user_does_not_overwrite_external_user_with_same_username():
with tempfile.TemporaryDirectory() as tmpdir:
db_path = os.path.join(tmpdir, "users.db")
user_db = UserDB(db_path)
user_db.initialize()
existing = user_db.create_user(
username="admin",
auth_source="oidc",
oidc_subject="oidc-subject",
role="user",
)
sync_builtin_admin_user("admin", "builtin-hash", db_path=db_path)
refreshed = user_db.get_user(user_id=existing["id"])
assert refreshed is not None
assert refreshed["auth_source"] == "oidc"
assert refreshed["role"] == "user"
assert refreshed["password_hash"] is None
@@ -0,0 +1,51 @@
"""Tests for persisted download-history helpers."""
from __future__ import annotations
import os
import tempfile
from shelfmark.core.download_history_service import DownloadHistoryService
from shelfmark.core.user_db import UserDB
def test_iso_to_epoch_treats_naive_sqlite_timestamp_as_utc():
epoch = DownloadHistoryService._iso_to_epoch("2026-01-02 03:04:05")
assert epoch == 1767323045.0
def test_record_download_stores_utc_iso_timestamps():
with tempfile.TemporaryDirectory() as tmpdir:
db_path = os.path.join(tmpdir, "users.db")
user_db = UserDB(db_path)
user_db.initialize()
service = DownloadHistoryService(db_path)
service.record_download(
task_id="task-1",
user_id=None,
username=None,
request_id=None,
source="direct_download",
source_display_name="Direct Download",
title="Example",
author=None,
format=None,
size=None,
preview=None,
content_type="ebook",
origin="direct",
)
conn = user_db._connect()
try:
row = conn.execute(
"SELECT queued_at, terminal_at FROM download_history WHERE task_id = ?",
("task-1",),
).fetchone()
finally:
conn.close()
assert row is not None
assert "+00:00" in row["queued_at"]
assert "+00:00" in row["terminal_at"]
+25
View File
@@ -224,6 +224,31 @@ class TestAtomicCopy:
assert result.exists()
assert result.read_text() == "content"
def test_publish_does_not_depend_on_hardlinks(self, tmp_path, monkeypatch):
"""Temp-file publish succeeds even when hardlinks are unavailable."""
from shelfmark.download.fs import atomic_copy as _atomic_copy
source = tmp_path / "source.txt"
source.write_text("content")
dest = tmp_path / "dest.txt"
link_calls = {"count": 0}
def _link_should_not_be_used(*_args, **_kwargs):
link_calls["count"] += 1
raise AssertionError("atomic_copy publish should not call os.link")
monkeypatch.setattr(os, "link", _link_should_not_be_used)
result = _atomic_copy(source, dest)
assert result == dest
assert result.exists()
assert result.read_text() == "content"
assert source.exists()
assert link_calls["count"] == 0
def test_max_attempts_exceeded(self, tmp_path):
"""Raises after max collision attempts."""
from shelfmark.download.fs import atomic_copy as _atomic_copy
+26
View File
@@ -246,6 +246,32 @@ class TestAtomicHardlink:
assert source.exists()
assert os.stat(source).st_ino != os.stat(result).st_ino
def test_falls_back_to_copy_on_input_output_error(self, tmp_path, monkeypatch):
"""Falls back to copy when filesystem reports hardlinks are unsupported via EIO."""
import errno
from shelfmark.download.fs import atomic_hardlink as _atomic_hardlink
source = tmp_path / "source.txt"
source.write_text("content")
dest = tmp_path / "dest.txt"
original_link = os.link
def _raise_eio_for_initial_link(src, dst, *_args, **_kwargs):
if Path(src) == source:
raise OSError(errno.EIO, "Input/output error")
return original_link(src, dst)
monkeypatch.setattr(os, "link", _raise_eio_for_initial_link)
result = _atomic_hardlink(source, dest)
assert result == dest
assert result.read_text() == "content"
assert source.exists()
assert os.stat(source).st_ino != os.stat(result).st_ino
class TestAtomicMove:
"""Tests for _atomic_move() function."""
+2 -3
View File
@@ -34,8 +34,7 @@ def test_get_aa_mirrors_falls_back_to_defaults_and_legacy_additional(monkeypatch
monkeypatch.setattr(mirrors, "_get_config", lambda: dummy)
aa = mirrors.get_aa_mirrors()
assert "https://annas-archive.gl" in aa
assert "https://annas-archive.li" in aa
for default_mirror in mirrors.DEFAULT_AA_MIRRORS:
assert default_mirror in aa
assert "https://extra.example" in aa
assert "https://extra2.example" in aa
+60 -23
View File
@@ -431,7 +431,7 @@ def test_fulfil_request_queues_as_requesting_user(user_db):
assert "_request_id" not in (fulfilled["release_data"] or {})
def test_fulfil_request_rejects_when_state_changes_after_queue_dispatch(user_db):
def test_fulfil_request_claims_request_before_queue_dispatch(user_db):
alice = user_db.create_user(username="alice")
admin = user_db.create_user(username="admin", role="admin")
created = create_request(
@@ -445,31 +445,22 @@ def test_fulfil_request_rejects_when_state_changes_after_queue_dispatch(user_db)
release_data=_release_data(),
)
release_data = _release_data()
def fake_queue_release(_release_data_arg, _priority, user_id=None, username=None):
# Simulate another worker fulfilling the same request while this call is in-flight.
user_db.update_request(
created["id"],
status="fulfilled",
release_data=release_data,
delivery_state="queued",
delivery_updated_at="2026-01-01T00:00:00+00:00",
reviewed_by=admin["id"],
reviewed_at="2026-01-01T00:00:00+00:00",
)
in_flight = user_db.get_request(created["id"])
assert in_flight is not None
assert in_flight["status"] == "fulfilled"
assert in_flight["delivery_state"] == "queued"
assert in_flight["reviewed_by"] == admin["id"]
return True, None
with pytest.raises(RequestServiceError) as exc_info:
fulfil_request(
user_db,
request_id=created["id"],
admin_user_id=admin["id"],
queue_release=fake_queue_release,
)
fulfilled = fulfil_request(
user_db,
request_id=created["id"],
admin_user_id=admin["id"],
queue_release=fake_queue_release,
)
assert exc_info.value.status_code == 409
assert exc_info.value.code == "stale_transition"
assert fulfilled["status"] == "fulfilled"
assert user_db.get_request(created["id"])["status"] == "fulfilled"
@@ -674,6 +665,47 @@ def test_sync_delivery_states_from_queue_status_updates_matching_fulfilled_reque
assert refreshed_bob["delivery_state"] == "queued"
def test_sync_delivery_states_from_queue_status_uses_request_id_for_duplicate_source_ids(user_db):
alice = user_db.create_user(username="alice")
older_request = user_db.create_request(
user_id=alice["id"],
source_hint="prowlarr",
content_type="ebook",
request_level="release",
policy_mode="request_release",
book_data=_book_data(),
release_data={"source": "prowlarr", "source_id": "shared-rel", "title": "Shared Release"},
status="fulfilled",
delivery_state="complete",
)
newer_request = user_db.create_request(
user_id=alice["id"],
source_hint="prowlarr",
content_type="ebook",
request_level="release",
policy_mode="request_release",
book_data=_book_data(),
release_data={"source": "prowlarr", "source_id": "shared-rel", "title": "Shared Release"},
status="fulfilled",
delivery_state="queued",
)
updated = sync_delivery_states_from_queue_status(
user_db,
queue_status={
"downloading": {
"shared-rel": {"id": "shared-rel", "request_id": newer_request["id"]},
},
},
user_id=alice["id"],
)
assert [row["id"] for row in updated] == [newer_request["id"]]
assert user_db.get_request(older_request["id"])["delivery_state"] == "complete"
assert user_db.get_request(newer_request["id"])["delivery_state"] == "downloading"
# ---------------------------------------------------------------------------
# book_data validation
# ---------------------------------------------------------------------------
@@ -1185,9 +1217,14 @@ def test_fulfil_queue_failure_returns_error(user_db):
assert exc_info.value.status_code == 409
assert exc_info.value.code == "queue_failed"
# Request should still be pending since queue failed before status update.
# Request should be rolled back to pending when queue dispatch fails.
row = user_db.get_request(created["id"])
assert row["status"] == "pending"
assert row["delivery_state"] == "none"
assert row["reviewed_by"] is None
assert row["reviewed_at"] is None
assert row["last_failure_reason"] == "Torrent client unreachable"
assert row["release_data"]["source_id"] == "release-123"
def test_fulfil_admin_can_override_release_data(user_db):
+2 -2
View File
@@ -45,7 +45,7 @@ class TestGetAuthMode:
"shelfmark.core.settings_registry.load_config_file",
return_value={"AUTH_METHOD": "builtin"},
):
with patch.object(main_module, "has_local_password_admin", return_value=True):
with 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):
@@ -53,7 +53,7 @@ class TestGetAuthMode:
"shelfmark.core.settings_registry.load_config_file",
return_value={"AUTH_METHOD": "builtin"},
):
with patch.object(main_module, "has_local_password_admin", return_value=False):
with 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):