mirror of
https://github.com/calibrain/shelfmark.git
synced 2026-10-05 08:51:11 +01:00
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:
@@ -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()
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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"]
|
||||
@@ -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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user