mirror of
https://github.com/calibrain/shelfmark.git
synced 2026-10-03 08:15:46 +01:00
Typecheck enforcement (#870)
All BasedPyright standard rules enforced and corrected
This commit is contained in:
+1
-1
@@ -105,7 +105,7 @@ ignore = ["D", "EM", "FBT", "PLR2004", "UP035", "TRY003", "E501", "TD002", "S104
|
||||
include = ["shelfmark"]
|
||||
exclude = [".local", "tests", "**/__pycache__", "**/node_modules"]
|
||||
pythonVersion = "3.14"
|
||||
typeCheckingMode = "off"
|
||||
typeCheckingMode = "standard"
|
||||
|
||||
[tool.vulture]
|
||||
paths = ["shelfmark"]
|
||||
|
||||
+16
-1
@@ -4,5 +4,20 @@ from shelfmark.config.env import FLASK_HOST, FLASK_PORT
|
||||
from shelfmark.core.config import config
|
||||
from shelfmark.main import app, socketio
|
||||
|
||||
|
||||
def _resolve_debug_flag(value: object) -> bool:
|
||||
"""Normalize DEBUG config values for Flask-SocketIO startup."""
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
return value.strip().lower() in {"1", "true", "yes", "on"}
|
||||
return bool(value)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
socketio.run(app, host=FLASK_HOST, port=FLASK_PORT, debug=config.get("DEBUG", False))
|
||||
socketio.run(
|
||||
app,
|
||||
host=FLASK_HOST,
|
||||
port=FLASK_PORT,
|
||||
debug=_resolve_debug_flag(config.get("DEBUG", False)),
|
||||
)
|
||||
|
||||
@@ -56,6 +56,11 @@ class WebSocketManager:
|
||||
"""Check if WebSocket is enabled and ready."""
|
||||
return self._enabled and self.socketio is not None
|
||||
|
||||
def _get_socketio(self) -> SocketIO | None:
|
||||
if not self._enabled:
|
||||
return None
|
||||
return self.socketio
|
||||
|
||||
def set_queue_status_fn(self, fn: Callable) -> None:
|
||||
"""Set the queue_status function reference for per-room filtering."""
|
||||
self._queue_status_fn = fn
|
||||
@@ -126,12 +131,13 @@ class WebSocketManager:
|
||||
|
||||
def broadcast_status_update(self, status_data: dict[str, Any]) -> None:
|
||||
"""Broadcast status update to all connected clients, filtered by user room."""
|
||||
if not self.is_enabled():
|
||||
socketio = self._get_socketio()
|
||||
if socketio is None:
|
||||
return
|
||||
|
||||
try:
|
||||
# Admins (and no-auth users) get full status
|
||||
self.socketio.emit("status_update", status_data, to="admins")
|
||||
socketio.emit("status_update", status_data, to="admins")
|
||||
|
||||
# Each user room gets filtered status
|
||||
with self._rooms_lock:
|
||||
@@ -147,12 +153,16 @@ class WebSocketManager:
|
||||
|
||||
def _broadcast_status_update_to_room(self, room: str) -> None:
|
||||
"""Broadcast status update to one user room."""
|
||||
socketio = self._get_socketio()
|
||||
if socketio is None:
|
||||
return
|
||||
|
||||
try:
|
||||
# Extract user_id from room name "user_123"
|
||||
uid = int(room.split("_", 1)[1])
|
||||
filtered = self._queue_status_fn(user_id=uid) if self._queue_status_fn else None
|
||||
if filtered is not None:
|
||||
self.socketio.emit("status_update", filtered, to=room)
|
||||
socketio.emit("status_update", filtered, to=room)
|
||||
except Exception:
|
||||
logger.exception("Failed to send status update for room %s", room)
|
||||
|
||||
@@ -160,19 +170,20 @@ class WebSocketManager:
|
||||
self, book_id: str, progress: float, status: str, user_id: int | None = None
|
||||
) -> None:
|
||||
"""Broadcast download progress update for a specific book."""
|
||||
if not self.is_enabled():
|
||||
socketio = self._get_socketio()
|
||||
if socketio is None:
|
||||
return
|
||||
|
||||
try:
|
||||
data = {"book_id": book_id, "progress": progress, "status": status}
|
||||
# Admins always see all progress
|
||||
self.socketio.emit("download_progress", data, to="admins")
|
||||
socketio.emit("download_progress", data, to="admins")
|
||||
# If task belongs to a specific user, send to their room too
|
||||
if user_id is not None:
|
||||
room = f"user_{user_id}"
|
||||
with self._rooms_lock:
|
||||
if room in self._user_rooms:
|
||||
self.socketio.emit("download_progress", data, to=room)
|
||||
socketio.emit("download_progress", data, to=room)
|
||||
logger.debug("Broadcasted progress for book %s: %s%%", book_id, progress)
|
||||
except Exception:
|
||||
logger.exception("Error broadcasting download progress")
|
||||
@@ -186,7 +197,8 @@ class WebSocketManager:
|
||||
phase: str = "searching",
|
||||
) -> None:
|
||||
"""Broadcast search status update for a release source search."""
|
||||
if not self.is_enabled():
|
||||
socketio = self._get_socketio()
|
||||
if socketio is None:
|
||||
return
|
||||
|
||||
try:
|
||||
@@ -197,7 +209,7 @@ class WebSocketManager:
|
||||
"message": message,
|
||||
"phase": phase,
|
||||
}
|
||||
self.socketio.emit("search_status", data)
|
||||
socketio.emit("search_status", data)
|
||||
except Exception:
|
||||
logger.exception("Error broadcasting search status")
|
||||
|
||||
|
||||
@@ -31,11 +31,30 @@ BACKOFF_BASE = 1.0
|
||||
BACKOFF_CAP = 10.0
|
||||
|
||||
|
||||
def _coerce_config_str(value: object, default: str) -> str:
|
||||
"""Return a string config value or a safe default."""
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
return default
|
||||
|
||||
|
||||
def _coerce_timeout_ms(value: object, default: int) -> int:
|
||||
"""Return a positive timeout in milliseconds or the default."""
|
||||
if isinstance(value, bool):
|
||||
return default
|
||||
if isinstance(value, int) and value > 0:
|
||||
return value
|
||||
return default
|
||||
|
||||
|
||||
def _fetch_via_bypasser(target_url: str) -> str | None:
|
||||
"""Make a single request to the external bypasser service. Returns HTML or None."""
|
||||
raw_bypasser_url = config.get("EXT_BYPASSER_URL", "http://flaresolverr:8191")
|
||||
bypasser_path = config.get("EXT_BYPASSER_PATH", "/v1")
|
||||
bypasser_timeout = config.get("EXT_BYPASSER_TIMEOUT", 60000)
|
||||
raw_bypasser_url = _coerce_config_str(
|
||||
config.get("EXT_BYPASSER_URL", "http://flaresolverr:8191"),
|
||||
"http://flaresolverr:8191",
|
||||
)
|
||||
bypasser_path = _coerce_config_str(config.get("EXT_BYPASSER_PATH", "/v1"), "/v1")
|
||||
bypasser_timeout = _coerce_timeout_ms(config.get("EXT_BYPASSER_TIMEOUT", 60000), 60000)
|
||||
|
||||
bypasser_url = normalize_http_url(raw_bypasser_url)
|
||||
if not bypasser_url or not bypasser_path:
|
||||
|
||||
@@ -17,7 +17,7 @@ from datetime import UTC, datetime
|
||||
from http import HTTPStatus
|
||||
from pathlib import Path
|
||||
from threading import Event
|
||||
from typing import Any
|
||||
from typing import Any, Protocol, TypedDict, TypeGuard
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import requests
|
||||
@@ -59,7 +59,21 @@ DDOS_GUARD_INDICATORS = [
|
||||
"could not verify your browser automatically",
|
||||
]
|
||||
|
||||
DISPLAY = {
|
||||
|
||||
class _DisplayState(TypedDict):
|
||||
ffmpeg: subprocess.Popen[bytes] | None
|
||||
ffmpeg_output: Path | None
|
||||
|
||||
|
||||
class _PageWithWindowRect(Protocol):
|
||||
async def set_window_rect(self, x: int, _y: int, width: int, height: int) -> object: ...
|
||||
|
||||
|
||||
class _BrowserWithWindowRectPage(Protocol):
|
||||
page: _PageWithWindowRect
|
||||
|
||||
|
||||
DISPLAY: _DisplayState = {
|
||||
"ffmpeg": None,
|
||||
"ffmpeg_output": None,
|
||||
}
|
||||
@@ -95,6 +109,30 @@ _SUBPROCESS_OPERATION_ERRORS = (
|
||||
)
|
||||
|
||||
|
||||
def _coerce_positive_int(value: object, default: int) -> int:
|
||||
"""Return a positive integer config value or the provided default."""
|
||||
if isinstance(value, bool):
|
||||
return default
|
||||
if isinstance(value, int) and value > 0:
|
||||
return value
|
||||
return default
|
||||
|
||||
|
||||
def _coerce_non_negative_float(value: object, default: float) -> float:
|
||||
"""Return a non-negative float config value or the provided default."""
|
||||
if isinstance(value, bool):
|
||||
return default
|
||||
if isinstance(value, int | float) and value >= 0:
|
||||
return float(value)
|
||||
return default
|
||||
|
||||
|
||||
def _has_window_rect_page(candidate: object) -> TypeGuard[_BrowserWithWindowRectPage]:
|
||||
"""Check whether a browser wrapper exposes page.set_window_rect()."""
|
||||
page = getattr(candidate, "page", None)
|
||||
return callable(getattr(page, "set_window_rect", None))
|
||||
|
||||
|
||||
def _describe_runtime_path(path: str | Path) -> str:
|
||||
"""Return compact ownership/mode info for a runtime path."""
|
||||
try:
|
||||
@@ -613,7 +651,9 @@ async def _bypass(
|
||||
page: Any, max_retries: int | None = None, cancel_flag: Event | None = None
|
||||
) -> bool:
|
||||
"""Attempt to bypass Cloudflare/DDOS-Guard protection using multiple methods."""
|
||||
max_retries = max_retries if max_retries is not None else app_config.MAX_RETRY
|
||||
max_retries = (
|
||||
max_retries if max_retries is not None else _coerce_positive_int(app_config.MAX_RETRY, 10)
|
||||
)
|
||||
|
||||
last_challenge_type = None
|
||||
consecutive_same_challenge = 0
|
||||
@@ -790,7 +830,7 @@ async def _get(url: str, driver: Any, cancel_flag: Event | None = None) -> str:
|
||||
|
||||
def get(url: str, retry: int | None = None, cancel_flag: Event | None = None) -> str:
|
||||
"""Fetch a URL with protection bypass. Creates fresh Chrome instance for each bypass."""
|
||||
retry = retry if retry is not None else app_config.MAX_RETRY
|
||||
retry = retry if retry is not None else _coerce_positive_int(app_config.MAX_RETRY, 10)
|
||||
|
||||
with LOCKED:
|
||||
# Try cookies first - another request may have completed bypass while waiting
|
||||
@@ -879,16 +919,17 @@ async def _create_cdp_browser(url: str) -> Any:
|
||||
)
|
||||
raise
|
||||
|
||||
try:
|
||||
await driver.page.set_window_rect(0, 0, screen_width, screen_height)
|
||||
except _CDP_OPERATION_ERRORS as e:
|
||||
logger.debug("Failed to set window size: %s", e)
|
||||
if _has_window_rect_page(driver):
|
||||
try:
|
||||
await driver.page.set_window_rect(0, 0, screen_width, screen_height)
|
||||
except _CDP_OPERATION_ERRORS as e:
|
||||
logger.debug("Failed to set window size: %s", e)
|
||||
|
||||
# Start FFmpeg recording if debug mode (record each bypass session)
|
||||
if app_config.get("DEBUG", False) and not DISPLAY.get("ffmpeg"):
|
||||
_start_ffmpeg_recording(display=os.environ.get("DISPLAY", ":0"))
|
||||
|
||||
await asyncio.sleep(app_config.DEFAULT_SLEEP)
|
||||
await asyncio.sleep(_coerce_non_negative_float(app_config.DEFAULT_SLEEP, 5.0))
|
||||
logger.info("Chrome browser ready (Pure CDP)")
|
||||
logger.log_resource_usage()
|
||||
return driver
|
||||
|
||||
@@ -123,7 +123,7 @@ def get_booklore_library_options() -> list[dict[str, Any]]:
|
||||
|
||||
base_url = str(config.get("BOOKLORE_HOST", "") or "").strip().rstrip("/")
|
||||
username = str(config.get("BOOKLORE_USERNAME", "") or "").strip()
|
||||
password = config.get("BOOKLORE_PASSWORD", "") or ""
|
||||
password = str(config.get("BOOKLORE_PASSWORD", "") or "")
|
||||
|
||||
if not base_url or not username or not password:
|
||||
return []
|
||||
@@ -148,7 +148,7 @@ def get_booklore_path_options() -> list[dict[str, Any]]:
|
||||
|
||||
base_url = str(config.get("BOOKLORE_HOST", "") or "").strip().rstrip("/")
|
||||
username = str(config.get("BOOKLORE_USERNAME", "") or "").strip()
|
||||
password = config.get("BOOKLORE_PASSWORD", "") or ""
|
||||
password = str(config.get("BOOKLORE_PASSWORD", "") or "")
|
||||
|
||||
if not base_url or not username or not password:
|
||||
return []
|
||||
@@ -182,7 +182,7 @@ def check_booklore_connection(
|
||||
|
||||
base_url = str(_get_value("BOOKLORE_HOST", "") or "").strip().rstrip("/")
|
||||
username = str(_get_value("BOOKLORE_USERNAME", "") or "").strip()
|
||||
password = _get_value("BOOKLORE_PASSWORD", "") or ""
|
||||
password = str(_get_value("BOOKLORE_PASSWORD", "") or "")
|
||||
|
||||
if not base_url:
|
||||
return {"success": False, "message": "Grimmory URL is required"}
|
||||
|
||||
@@ -79,7 +79,9 @@ def is_covers_cache_enabled() -> bool:
|
||||
from shelfmark.core.config import config
|
||||
|
||||
setting_enabled = config.get("COVERS_CACHE_ENABLED", True)
|
||||
return setting_enabled and _is_config_dir_writable()
|
||||
if isinstance(setting_enabled, str):
|
||||
return string_to_bool(setting_enabled) and _is_config_dir_writable()
|
||||
return bool(setting_enabled) and _is_config_dir_writable()
|
||||
|
||||
|
||||
# =============================================================================
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
"""Configuration migration helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Protocol
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
from os import PathLike
|
||||
|
||||
_DEPRECATED_SETTINGS_RESTRICTION_KEYS = (
|
||||
"PROXY_AUTH_RESTRICT_SETTINGS_TO_ADMIN",
|
||||
@@ -14,6 +17,16 @@ _DEPRECATED_SETTINGS_RESTRICTION_KEYS = (
|
||||
)
|
||||
|
||||
|
||||
class MigrationLogger(Protocol):
|
||||
"""Logger surface used by config migration helpers."""
|
||||
|
||||
def info(self, msg: str, *args: object) -> object: ...
|
||||
|
||||
def debug(self, msg: str, *args: object) -> object: ...
|
||||
|
||||
def exception(self, msg: str, *args: object) -> object: ...
|
||||
|
||||
|
||||
def _as_bool(value: object) -> bool:
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
@@ -50,9 +63,9 @@ def migrate_security_settings(
|
||||
load_users_config: Callable[[], dict[str, Any]],
|
||||
save_users_config: Callable[[dict[str, Any]], None],
|
||||
ensure_config_dir: Callable[[], None],
|
||||
get_config_path: Callable[[], object],
|
||||
get_config_path: Callable[[], str | PathLike[str]],
|
||||
sync_builtin_admin_user: Callable[[str, str], None],
|
||||
logger: object,
|
||||
logger: MigrationLogger,
|
||||
) -> None:
|
||||
"""Migrate legacy security keys and sync builtin admin credentials."""
|
||||
try:
|
||||
|
||||
@@ -45,10 +45,13 @@ def _migrate_security_settings() -> None:
|
||||
save_config_file,
|
||||
)
|
||||
|
||||
def _save_users_config(values: dict[str, Any]) -> None:
|
||||
save_config_file("users", values)
|
||||
|
||||
migrate_security_settings(
|
||||
load_security_config=lambda: load_config_file("security"),
|
||||
load_users_config=lambda: load_config_file("users"),
|
||||
save_users_config=lambda values: save_config_file("users", values),
|
||||
save_users_config=_save_users_config,
|
||||
ensure_config_dir=lambda: _ensure_config_dir("security"),
|
||||
get_config_path=lambda: _get_config_file_path("security"),
|
||||
sync_builtin_admin_user=sync_builtin_admin_user,
|
||||
|
||||
@@ -255,6 +255,11 @@ _LANGUAGE_OPTIONS = [
|
||||
]
|
||||
|
||||
|
||||
def _string_setting(value: object) -> str:
|
||||
"""Normalize free-form string settings used by select option builders."""
|
||||
return value if isinstance(value, str) else str(value or "")
|
||||
|
||||
|
||||
def _get_aa_base_url_options() -> list[dict[str, str]]:
|
||||
"""Build AA URL options dynamically, including additional mirrors from config."""
|
||||
from shelfmark.core.config import config
|
||||
@@ -269,7 +274,7 @@ def _get_aa_base_url_options() -> list[dict[str, str]]:
|
||||
# If AA_BASE_URL is configured to a custom mirror that isn't present in the
|
||||
# defaults/additional list, include it so the UI can display the active value.
|
||||
configured_url = normalize_http_url(
|
||||
config.get("AA_BASE_URL", "auto"),
|
||||
_string_setting(config.get("AA_BASE_URL", "auto")),
|
||||
default_scheme="https",
|
||||
allow_special=("auto",),
|
||||
)
|
||||
@@ -300,7 +305,7 @@ def _get_zlib_mirror_options() -> list[dict[str, str]]:
|
||||
options.append({"value": url, "label": domain})
|
||||
|
||||
# Add custom mirrors
|
||||
additional = config.get("ZLIB_ADDITIONAL_URLS", "")
|
||||
additional = _string_setting(config.get("ZLIB_ADDITIONAL_URLS", ""))
|
||||
if additional:
|
||||
for raw_url in additional.split(","):
|
||||
url = raw_url.strip()
|
||||
@@ -324,7 +329,7 @@ def _get_welib_mirror_options() -> list[dict[str, str]]:
|
||||
options.append({"value": url, "label": domain})
|
||||
|
||||
# Add custom mirrors
|
||||
additional = config.get("WELIB_ADDITIONAL_URLS", "")
|
||||
additional = _string_setting(config.get("WELIB_ADDITIONAL_URLS", ""))
|
||||
if additional:
|
||||
for raw_url in additional.split(","):
|
||||
url = raw_url.strip()
|
||||
|
||||
@@ -39,6 +39,7 @@ if TYPE_CHECKING:
|
||||
|
||||
logger = setup_logger(__name__)
|
||||
_USER_DB_IDENTITY_ERRORS = (sqlite3.Error, OSError)
|
||||
type ActivityRouteResponse = tuple[Response, int]
|
||||
|
||||
|
||||
def _normalize_log_field(value: object) -> str:
|
||||
@@ -219,12 +220,23 @@ class _ActorContext(NamedTuple):
|
||||
viewer_scope: str
|
||||
|
||||
|
||||
type ActivityActorResolution = tuple[_ActorContext, None] | tuple[None, ActivityRouteResponse]
|
||||
|
||||
|
||||
def _require_activity_actor(actor: _ActorContext | None, *, action: str) -> _ActorContext:
|
||||
"""Convert a resolved actor into the non-optional form route handlers expect."""
|
||||
if actor is None:
|
||||
msg = f"Activity actor missing after successful resolution for {action}"
|
||||
raise RuntimeError(msg)
|
||||
return actor
|
||||
|
||||
|
||||
def _resolve_activity_actor(
|
||||
*,
|
||||
user_db: UserDB,
|
||||
resolve_auth_mode: Callable[[], str],
|
||||
action: str,
|
||||
) -> tuple[_ActorContext | None, object | None]:
|
||||
) -> ActivityActorResolution:
|
||||
"""Resolve acting user identity for activity mutations.
|
||||
|
||||
Returns (actor, error_response). On success actor is non-None.
|
||||
@@ -245,6 +257,9 @@ def _resolve_activity_actor(
|
||||
auth_mode=auth_mode,
|
||||
)
|
||||
if db_user_id is None:
|
||||
if db_gate is None:
|
||||
msg = f"Activity actor resolution failed without an error response for {action}"
|
||||
raise RuntimeError(msg)
|
||||
return None, db_gate
|
||||
|
||||
is_admin = bool(session.get("is_admin"))
|
||||
@@ -267,7 +282,7 @@ def _activity_ws_room(actor: _ActorContext) -> str:
|
||||
return "admins"
|
||||
|
||||
|
||||
def _check_item_ownership(actor: _ActorContext, row: dict[str, Any]) -> object | None:
|
||||
def _check_item_ownership(actor: _ActorContext, row: dict[str, Any]) -> str | None:
|
||||
"""Return an error string if the actor doesn't own the item, else None."""
|
||||
if actor.is_admin:
|
||||
return None
|
||||
@@ -277,14 +292,14 @@ def _check_item_ownership(actor: _ActorContext, row: dict[str, Any]) -> object |
|
||||
return None
|
||||
|
||||
|
||||
def _check_terminal_download(row: dict[str, Any]) -> object | None:
|
||||
def _check_terminal_download(row: dict[str, Any]) -> str | None:
|
||||
final_status = str(row.get("final_status") or "").strip().lower()
|
||||
if final_status not in VALID_TERMINAL_STATUSES:
|
||||
return "Only terminal downloads can be dismissed"
|
||||
return None
|
||||
|
||||
|
||||
def _check_terminal_request(row: dict[str, Any]) -> object | None:
|
||||
def _check_terminal_request(row: dict[str, Any]) -> str | None:
|
||||
if _request_terminal_status(row) is None:
|
||||
return "Only terminal requests can be dismissed"
|
||||
return None
|
||||
@@ -513,6 +528,7 @@ def register_activity_routes(
|
||||
)
|
||||
if actor_error is not None:
|
||||
return actor_error
|
||||
actor = _require_activity_actor(actor, action="snapshot")
|
||||
|
||||
hidden_rows = activity_view_state_service.list_hidden(viewer_scope=actor.viewer_scope)
|
||||
hidden_item_keys = {str(row.get("item_key") or "").strip() for row in hidden_rows}
|
||||
@@ -582,6 +598,7 @@ def register_activity_routes(
|
||||
)
|
||||
if actor_error is not None:
|
||||
return actor_error
|
||||
actor = _require_activity_actor(actor, action="dismiss")
|
||||
|
||||
data = request.get_json(silent=True)
|
||||
if not isinstance(data, dict):
|
||||
@@ -755,6 +772,7 @@ def register_activity_routes(
|
||||
)
|
||||
if actor_error is not None:
|
||||
return actor_error
|
||||
actor = _require_activity_actor(actor, action="dismiss_many")
|
||||
|
||||
data = request.get_json(silent=True)
|
||||
if not isinstance(data, dict):
|
||||
@@ -946,6 +964,7 @@ def register_activity_routes(
|
||||
)
|
||||
if actor_error is not None:
|
||||
return actor_error
|
||||
actor = _require_activity_actor(actor, action="history")
|
||||
|
||||
limit = request.args.get("limit", type=int, default=50)
|
||||
offset = request.args.get("offset", type=int, default=0)
|
||||
@@ -1056,6 +1075,7 @@ def register_activity_routes(
|
||||
)
|
||||
if actor_error is not None:
|
||||
return actor_error
|
||||
actor = _require_activity_actor(actor, action="history_clear")
|
||||
|
||||
cleared_count = activity_view_state_service.clear_history(
|
||||
viewer_scope=actor.viewer_scope,
|
||||
|
||||
@@ -4,10 +4,12 @@ Registers /api/admin/users CRUD endpoints for managing users.
|
||||
All endpoints require admin session.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sqlite3
|
||||
from functools import wraps
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, ParamSpec
|
||||
|
||||
from flask import Flask, Response, g, jsonify, request, session
|
||||
from werkzeug.security import generate_password_hash
|
||||
@@ -37,8 +39,12 @@ from shelfmark.core.logger import setup_logger
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
from flask.typing import ResponseReturnValue
|
||||
|
||||
from shelfmark.core.user_db import UserDB
|
||||
|
||||
P = ParamSpec("P")
|
||||
|
||||
logger = setup_logger(__name__)
|
||||
MIN_PASSWORD_LENGTH = 4
|
||||
_CONFIG_REFRESH_ERRORS = (ImportError, OSError, RuntimeError, TypeError, ValueError)
|
||||
@@ -144,8 +150,8 @@ def register_admin_routes(app: Flask, user_db: UserDB) -> None:
|
||||
"""Register admin user management routes on the Flask app."""
|
||||
|
||||
def _require_admin(
|
||||
f: Callable[..., Response | tuple[Response, int]],
|
||||
) -> Callable[..., Response | tuple[Response, int]]:
|
||||
f: Callable[P, ResponseReturnValue],
|
||||
) -> Callable[P, ResponseReturnValue]:
|
||||
"""Require an admin session for admin routes.
|
||||
|
||||
In no-auth mode, everyone has access (is_admin defaults True).
|
||||
@@ -154,7 +160,7 @@ def register_admin_routes(app: Flask, user_db: UserDB) -> None:
|
||||
"""
|
||||
|
||||
@wraps(f)
|
||||
def decorated(*args: object, **kwargs: object) -> Response | tuple[Response, int]:
|
||||
def decorated(*args: P.args, **kwargs: P.kwargs) -> ResponseReturnValue:
|
||||
auth_mode = load_active_auth_mode(CWA_DB_PATH, user_db=user_db)
|
||||
g.auth_mode = auth_mode
|
||||
if auth_mode != "none":
|
||||
@@ -381,6 +387,8 @@ def register_admin_routes(app: Flask, user_db: UserDB) -> None:
|
||||
)
|
||||
|
||||
updated = user_db.get_user(user_id=user_id)
|
||||
if not isinstance(updated, dict):
|
||||
return jsonify({"error": "User not found"}), 404
|
||||
result = _serialize_user(
|
||||
updated,
|
||||
g.auth_mode,
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
"""Admin settings-introspection routes and settings validation helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from flask import Flask, Response, jsonify, request
|
||||
from flask import Flask, jsonify, request
|
||||
|
||||
from shelfmark.config.notifications_settings import (
|
||||
build_notification_test_result,
|
||||
@@ -26,6 +28,8 @@ from shelfmark.core.user_settings_overrides import (
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
from flask.typing import ResponseReturnValue
|
||||
|
||||
from shelfmark.core.user_db import UserDB
|
||||
|
||||
|
||||
@@ -151,13 +155,15 @@ def build_user_notification_test_response(
|
||||
def register_admin_settings_routes(
|
||||
app: Flask,
|
||||
user_db: UserDB,
|
||||
require_admin: Callable[[Callable[..., object]], Callable[..., object]],
|
||||
require_admin: Callable[
|
||||
[Callable[..., ResponseReturnValue]], Callable[..., ResponseReturnValue]
|
||||
],
|
||||
) -> None:
|
||||
"""Register admin endpoints for user-specific settings and defaults."""
|
||||
|
||||
@app.route("/api/admin/download-defaults", methods=["GET"])
|
||||
@require_admin
|
||||
def admin_download_defaults() -> Response | tuple[Response, int]:
|
||||
def admin_download_defaults() -> ResponseReturnValue:
|
||||
defaults = {
|
||||
key: ("" if (value := app_config.get(key, field.default)) is None else value)
|
||||
for key, field in _get_ordered_user_overridable_fields("downloads")
|
||||
@@ -170,7 +176,7 @@ def register_admin_settings_routes(
|
||||
|
||||
@app.route("/api/admin/booklore-options", methods=["GET"])
|
||||
@require_admin
|
||||
def admin_booklore_options() -> Response | tuple[Response, int]:
|
||||
def admin_booklore_options() -> ResponseReturnValue:
|
||||
from shelfmark.core import admin_routes
|
||||
|
||||
return jsonify(
|
||||
@@ -182,7 +188,7 @@ def register_admin_settings_routes(
|
||||
|
||||
@app.route("/api/admin/users/<int:user_id>/delivery-preferences", methods=["GET"])
|
||||
@require_admin
|
||||
def admin_get_delivery_preferences(user_id: int) -> Response | tuple[Response, int]:
|
||||
def admin_get_delivery_preferences(user_id: int) -> ResponseReturnValue:
|
||||
user = user_db.get_user(user_id=user_id)
|
||||
if not user:
|
||||
return jsonify({"error": "User not found"}), 404
|
||||
@@ -196,7 +202,7 @@ def register_admin_settings_routes(
|
||||
|
||||
@app.route("/api/admin/users/<int:user_id>/search-preferences", methods=["GET"])
|
||||
@require_admin
|
||||
def admin_get_search_preferences(user_id: int) -> Response | tuple[Response, int]:
|
||||
def admin_get_search_preferences(user_id: int) -> ResponseReturnValue:
|
||||
user = user_db.get_user(user_id=user_id)
|
||||
if not user:
|
||||
return jsonify({"error": "User not found"}), 404
|
||||
@@ -210,7 +216,7 @@ def register_admin_settings_routes(
|
||||
|
||||
@app.route("/api/admin/users/<int:user_id>/notification-preferences", methods=["GET"])
|
||||
@require_admin
|
||||
def admin_get_notification_preferences(user_id: int) -> Response | tuple[Response, int]:
|
||||
def admin_get_notification_preferences(user_id: int) -> ResponseReturnValue:
|
||||
user = user_db.get_user(user_id=user_id)
|
||||
if not user:
|
||||
return jsonify({"error": "User not found"}), 404
|
||||
@@ -224,7 +230,7 @@ def register_admin_settings_routes(
|
||||
|
||||
@app.route("/api/admin/users/<int:user_id>/notification-preferences/test", methods=["POST"])
|
||||
@require_admin
|
||||
def admin_test_notification_preferences(user_id: int) -> Response | tuple[Response, int]:
|
||||
def admin_test_notification_preferences(user_id: int) -> ResponseReturnValue:
|
||||
user = user_db.get_user(user_id=user_id)
|
||||
if not user:
|
||||
return jsonify({"error": "User not found"}), 404
|
||||
@@ -238,7 +244,7 @@ def register_admin_settings_routes(
|
||||
|
||||
@app.route("/api/admin/settings/overrides-summary", methods=["GET"])
|
||||
@require_admin
|
||||
def admin_settings_overrides_summary() -> Response | tuple[Response, int]:
|
||||
def admin_settings_overrides_summary() -> ResponseReturnValue:
|
||||
settings_registry = _get_settings_registry()
|
||||
|
||||
tab_name = (request.args.get("tab") or "downloads").strip()
|
||||
@@ -272,7 +278,7 @@ def register_admin_settings_routes(
|
||||
|
||||
@app.route("/api/admin/users/<int:user_id>/effective-settings", methods=["GET"])
|
||||
@require_admin
|
||||
def admin_get_effective_settings(user_id: int) -> Response | tuple[Response, int]:
|
||||
def admin_get_effective_settings(user_id: int) -> ResponseReturnValue:
|
||||
user = user_db.get_user(user_id=user_id)
|
||||
if not user:
|
||||
return jsonify({"error": "User not found"}), 404
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
"""Authentication mode, auth-source normalization, and admin access policy helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sqlite3
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Protocol, TypeGuard
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Mapping
|
||||
@@ -22,6 +24,17 @@ AUTH_SOURCE_SET = frozenset(AUTH_SOURCES)
|
||||
_ALWAYS_ADMIN_SETTINGS_TABS = frozenset({"security", "users"})
|
||||
|
||||
|
||||
class _UserDBWithAdminPassword(Protocol):
|
||||
"""Minimal user DB surface needed for local-admin checks."""
|
||||
|
||||
def has_admin_with_password(self) -> bool: ...
|
||||
|
||||
|
||||
def _has_admin_password_api(candidate: object) -> TypeGuard[_UserDBWithAdminPassword]:
|
||||
"""Return True when *candidate* exposes the admin-password lookup we need."""
|
||||
return callable(getattr(candidate, "has_admin_with_password", None))
|
||||
|
||||
|
||||
def has_local_password_admin(user_db: object | None = None) -> bool:
|
||||
"""Return True when at least one local admin with a password exists."""
|
||||
try:
|
||||
@@ -33,6 +46,8 @@ def has_local_password_admin(user_db: object | None = None) -> bool:
|
||||
db = UserDB(str(Path(config_root) / "users.db"))
|
||||
db.initialize()
|
||||
|
||||
if not _has_admin_password_api(db):
|
||||
return False
|
||||
return db.has_admin_with_password()
|
||||
except AttributeError, ImportError, OSError, RuntimeError, TypeError, ValueError, sqlite3.Error:
|
||||
return False
|
||||
|
||||
+18
-1
@@ -120,6 +120,20 @@ def cache_key(*args: object, **kwargs: object) -> str:
|
||||
return ":".join(parts)
|
||||
|
||||
|
||||
def _coerce_ttl_seconds(value: object, *, default: int) -> int:
|
||||
"""Normalize cache TTL values read from config or decorator arguments."""
|
||||
if isinstance(value, bool):
|
||||
return default
|
||||
if isinstance(value, int):
|
||||
return value if value > 0 else default
|
||||
if isinstance(value, str):
|
||||
stripped = value.strip()
|
||||
if stripped.isdigit():
|
||||
parsed = int(stripped)
|
||||
return parsed if parsed > 0 else default
|
||||
return default
|
||||
|
||||
|
||||
def cacheable(
|
||||
ttl: int | None = None,
|
||||
ttl_key: str | None = None,
|
||||
@@ -142,7 +156,10 @@ def cacheable(
|
||||
if ttl is not None:
|
||||
effective_ttl = ttl
|
||||
elif ttl_key:
|
||||
effective_ttl = config.get(ttl_key, ttl_default)
|
||||
effective_ttl = _coerce_ttl_seconds(
|
||||
config.get(ttl_key, ttl_default),
|
||||
default=ttl_default,
|
||||
)
|
||||
else:
|
||||
effective_ttl = ttl_default
|
||||
|
||||
|
||||
@@ -58,7 +58,7 @@ class Config:
|
||||
Values are cached for performance and can be refreshed when settings change.
|
||||
"""
|
||||
|
||||
_instance: Config | None = None
|
||||
_instance: Self | None = None
|
||||
_lock = Lock()
|
||||
|
||||
def __new__(cls) -> Self:
|
||||
@@ -68,7 +68,11 @@ class Config:
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
cls._instance._initialized = False
|
||||
return cls._instance
|
||||
instance = cls._instance
|
||||
if instance is None:
|
||||
msg = "Config singleton failed to initialize"
|
||||
raise RuntimeError(msg)
|
||||
return instance
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize caches and backing stores for the singleton."""
|
||||
|
||||
@@ -7,7 +7,7 @@ import sqlite3
|
||||
import threading
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, SupportsIndex, SupportsInt, TypeGuard
|
||||
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.models import TERMINAL_QUEUE_STATUSES
|
||||
@@ -25,6 +25,25 @@ ACTIVE_DOWNLOAD_STATUS = "active"
|
||||
VALID_ORIGINS = frozenset({"direct", "requested"})
|
||||
|
||||
|
||||
def _is_convertible_to_int(
|
||||
value: object,
|
||||
) -> TypeGuard[str | bytes | bytearray | SupportsInt | SupportsIndex]:
|
||||
"""Return True when *value* can be safely passed to ``int``."""
|
||||
return (
|
||||
isinstance(value, (str, bytes, bytearray))
|
||||
or hasattr(value, "__int__")
|
||||
or hasattr(value, "__index__")
|
||||
)
|
||||
|
||||
|
||||
def _coerce_int_value(value: object) -> int:
|
||||
"""Normalize int-like values and raise TypeError for unsupported inputs."""
|
||||
if isinstance(value, bool) or not _is_convertible_to_int(value):
|
||||
msg = "limit must be an integer"
|
||||
raise TypeError(msg)
|
||||
return int(value)
|
||||
|
||||
|
||||
def _normalize_task_id(task_id: object) -> str:
|
||||
normalized = normalize_optional_text(task_id)
|
||||
if normalized is None:
|
||||
@@ -60,7 +79,7 @@ def _normalize_limit(value: object, *, default: int, minimum: int, maximum: int)
|
||||
if value is None:
|
||||
return default
|
||||
try:
|
||||
parsed = int(value)
|
||||
parsed = _coerce_int_value(value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
msg = "limit must be an integer"
|
||||
raise ValueError(msg) from exc
|
||||
|
||||
@@ -13,6 +13,7 @@ from urllib.parse import urlparse
|
||||
import requests
|
||||
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.request_helpers import coerce_int
|
||||
from shelfmark.download.network import get_ssl_verify
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -594,8 +595,8 @@ def get_image_cache() -> ImageCacheService:
|
||||
from shelfmark.core.config import config
|
||||
|
||||
cache_dir = CONFIG_DIR / "covers"
|
||||
max_size_mb = config.get("COVERS_CACHE_MAX_SIZE_MB", 500)
|
||||
ttl_days = config.get("COVERS_CACHE_TTL", 0)
|
||||
max_size_mb = coerce_int(config.get("COVERS_CACHE_MAX_SIZE_MB", 500), 500)
|
||||
ttl_days = coerce_int(config.get("COVERS_CACHE_TTL", 0), 0)
|
||||
ttl_seconds = ttl_days * 86400 if ttl_days > 0 else 0
|
||||
|
||||
_instance = ImageCacheService(
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import logging
|
||||
import sys
|
||||
from collections.abc import Mapping
|
||||
from logging.handlers import RotatingFileHandler
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
@@ -17,15 +18,29 @@ class CustomLogger(logging.Logger):
|
||||
def error_trace(self, msg: object, *args: object, **kwargs: object) -> None:
|
||||
"""Log an error message with full stack trace."""
|
||||
self.log_resource_usage()
|
||||
kwargs.pop("exc_info", None)
|
||||
self.error(msg, *args, exc_info=True, **kwargs)
|
||||
stack_info, stacklevel, extra = _extract_log_kwargs(kwargs)
|
||||
self.error(
|
||||
msg,
|
||||
*args,
|
||||
exc_info=True,
|
||||
stack_info=stack_info,
|
||||
stacklevel=stacklevel,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def debug_trace(self, msg: object, *args: object, **kwargs: object) -> None:
|
||||
"""Log a debug message (stack trace only if exception active)."""
|
||||
kwargs.pop("exc_info", None)
|
||||
stack_info, stacklevel, extra = _extract_log_kwargs(kwargs)
|
||||
# Only include exc_info if there's actually an exception
|
||||
has_exception = sys.exc_info()[0] is not None
|
||||
self.debug(msg, *args, exc_info=has_exception, **kwargs)
|
||||
self.debug(
|
||||
msg,
|
||||
*args,
|
||||
exc_info=has_exception,
|
||||
stack_info=stack_info,
|
||||
stacklevel=stacklevel,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def log_resource_usage(self) -> None:
|
||||
"""Log best-effort CPU and memory usage for the current container."""
|
||||
@@ -39,9 +54,13 @@ class CustomLogger(logging.Logger):
|
||||
|
||||
def _get_process_rss_mb(proc: object) -> float | None:
|
||||
try:
|
||||
mem = proc.info.get("memory_info")
|
||||
if mem:
|
||||
return mem.rss / (1024 * 1024)
|
||||
proc_info = getattr(proc, "info", None)
|
||||
if not isinstance(proc_info, Mapping):
|
||||
return None
|
||||
mem = proc_info.get("memory_info")
|
||||
rss = getattr(mem, "rss", None)
|
||||
if isinstance(rss, int | float):
|
||||
return rss / (1024 * 1024)
|
||||
except (
|
||||
psutil.NoSuchProcess,
|
||||
psutil.AccessDenied,
|
||||
@@ -78,6 +97,31 @@ class CustomLogger(logging.Logger):
|
||||
return
|
||||
|
||||
|
||||
def _extract_log_kwargs(
|
||||
kwargs: Mapping[str, object],
|
||||
) -> tuple[bool, int, Mapping[str, object] | None]:
|
||||
stack_info = kwargs.get("stack_info")
|
||||
normalized_stack_info = stack_info if isinstance(stack_info, bool) else False
|
||||
|
||||
stacklevel = kwargs.get("stacklevel")
|
||||
normalized_stacklevel = stacklevel if isinstance(stacklevel, int) else 1
|
||||
|
||||
extra = kwargs.get("extra")
|
||||
normalized_extra = _normalize_log_extra(extra)
|
||||
|
||||
return normalized_stack_info, normalized_stacklevel, normalized_extra
|
||||
|
||||
|
||||
def _normalize_log_extra(value: object) -> Mapping[str, object] | None:
|
||||
if not isinstance(value, Mapping):
|
||||
return None
|
||||
|
||||
if all(isinstance(key, str) for key in value):
|
||||
return value
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def setup_logger(name: str, log_file: Path = LOG_FILE) -> CustomLogger:
|
||||
"""Set up and configure a logger instance.
|
||||
|
||||
|
||||
+25
-12
@@ -1,18 +1,18 @@
|
||||
"""Centralized mirror configuration for all download sources."""
|
||||
|
||||
# Lazy import to avoid circular imports
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from shelfmark.core.utils import normalize_http_url
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from types import ModuleType
|
||||
from shelfmark.core.config import Config
|
||||
|
||||
_config_module = None
|
||||
|
||||
|
||||
def _get_config() -> ModuleType:
|
||||
def _get_config() -> Config:
|
||||
"""Lazy import of config module to avoid circular imports."""
|
||||
global _config_module
|
||||
if _config_module is None:
|
||||
@@ -55,6 +55,11 @@ def _normalize_mirror_url(url: str) -> str:
|
||||
return normalize_http_url(url, default_scheme="https")
|
||||
|
||||
|
||||
def _string_config_value(value: object) -> str:
|
||||
"""Normalize mirror-related config values to strings."""
|
||||
return value if isinstance(value, str) else str(value or "")
|
||||
|
||||
|
||||
def get_aa_mirrors() -> list[str]:
|
||||
"""Get Anna's Archive mirrors.
|
||||
|
||||
@@ -91,7 +96,7 @@ def get_aa_mirrors() -> list[str]:
|
||||
mirrors = [url for url in mirrors if url]
|
||||
|
||||
# Backwards-compatible append-only behavior for legacy configs/env.
|
||||
additional = config.get("AA_ADDITIONAL_URLS", "")
|
||||
additional = _string_config_value(config.get("AA_ADDITIONAL_URLS", ""))
|
||||
if additional:
|
||||
for url in additional.split(","):
|
||||
normalized = _normalize_mirror_url(url)
|
||||
@@ -112,7 +117,7 @@ def get_libgen_mirrors() -> list[str]:
|
||||
mirrors = [url for url in mirrors if url]
|
||||
config = _get_config()
|
||||
|
||||
additional = config.get("LIBGEN_ADDITIONAL_URLS", "")
|
||||
additional = _string_config_value(config.get("LIBGEN_ADDITIONAL_URLS", ""))
|
||||
if additional:
|
||||
for url in additional.split(","):
|
||||
normalized = _normalize_mirror_url(url)
|
||||
@@ -131,7 +136,9 @@ def get_zlib_mirrors() -> list[str]:
|
||||
"""
|
||||
config = _get_config()
|
||||
|
||||
primary = _normalize_mirror_url(config.get("ZLIB_PRIMARY_URL", DEFAULT_ZLIB_MIRRORS[0]))
|
||||
primary = _normalize_mirror_url(
|
||||
_string_config_value(config.get("ZLIB_PRIMARY_URL", DEFAULT_ZLIB_MIRRORS[0]))
|
||||
)
|
||||
if not primary:
|
||||
primary = _normalize_mirror_url(DEFAULT_ZLIB_MIRRORS[0])
|
||||
mirrors = [primary]
|
||||
@@ -143,7 +150,7 @@ def get_zlib_mirrors() -> list[str]:
|
||||
mirrors.append(normalized)
|
||||
|
||||
# Add custom mirrors
|
||||
additional = config.get("ZLIB_ADDITIONAL_URLS", "")
|
||||
additional = _string_config_value(config.get("ZLIB_ADDITIONAL_URLS", ""))
|
||||
if additional:
|
||||
for url in additional.split(","):
|
||||
normalized = _normalize_mirror_url(url)
|
||||
@@ -161,7 +168,9 @@ def get_zlib_primary_url() -> str:
|
||||
|
||||
"""
|
||||
config = _get_config()
|
||||
primary = _normalize_mirror_url(config.get("ZLIB_PRIMARY_URL", DEFAULT_ZLIB_MIRRORS[0]))
|
||||
primary = _normalize_mirror_url(
|
||||
_string_config_value(config.get("ZLIB_PRIMARY_URL", DEFAULT_ZLIB_MIRRORS[0]))
|
||||
)
|
||||
return primary or _normalize_mirror_url(DEFAULT_ZLIB_MIRRORS[0])
|
||||
|
||||
|
||||
@@ -185,7 +194,9 @@ def get_welib_mirrors() -> list[str]:
|
||||
"""
|
||||
config = _get_config()
|
||||
|
||||
primary = _normalize_mirror_url(config.get("WELIB_PRIMARY_URL", DEFAULT_WELIB_MIRRORS[0]))
|
||||
primary = _normalize_mirror_url(
|
||||
_string_config_value(config.get("WELIB_PRIMARY_URL", DEFAULT_WELIB_MIRRORS[0]))
|
||||
)
|
||||
if not primary:
|
||||
primary = _normalize_mirror_url(DEFAULT_WELIB_MIRRORS[0])
|
||||
mirrors = [primary]
|
||||
@@ -197,7 +208,7 @@ def get_welib_mirrors() -> list[str]:
|
||||
mirrors.append(normalized)
|
||||
|
||||
# Add custom mirrors
|
||||
additional = config.get("WELIB_ADDITIONAL_URLS", "")
|
||||
additional = _string_config_value(config.get("WELIB_ADDITIONAL_URLS", ""))
|
||||
if additional:
|
||||
for url in additional.split(","):
|
||||
normalized = _normalize_mirror_url(url)
|
||||
@@ -215,7 +226,9 @@ def get_welib_primary_url() -> str:
|
||||
|
||||
"""
|
||||
config = _get_config()
|
||||
primary = _normalize_mirror_url(config.get("WELIB_PRIMARY_URL", DEFAULT_WELIB_MIRRORS[0]))
|
||||
primary = _normalize_mirror_url(
|
||||
_string_config_value(config.get("WELIB_PRIMARY_URL", DEFAULT_WELIB_MIRRORS[0]))
|
||||
)
|
||||
return primary or _normalize_mirror_url(DEFAULT_WELIB_MIRRORS[0])
|
||||
|
||||
|
||||
@@ -250,7 +263,7 @@ def get_zlib_cookie_domains() -> set:
|
||||
|
||||
# Add custom domains
|
||||
config = _get_config()
|
||||
additional = config.get("ZLIB_ADDITIONAL_URLS", "")
|
||||
additional = _string_config_value(config.get("ZLIB_ADDITIONAL_URLS", ""))
|
||||
if additional:
|
||||
for url in additional.split(","):
|
||||
normalized = _normalize_mirror_url(url)
|
||||
|
||||
@@ -8,7 +8,7 @@ from concurrent.futures import ThreadPoolExecutor
|
||||
from contextlib import contextmanager, suppress
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, Protocol, TypeGuard
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
try:
|
||||
@@ -18,6 +18,7 @@ except ImportError: # pragma: no cover - exercised in tests via monkeypatch
|
||||
|
||||
from shelfmark.core.config import config as app_config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.request_helpers import normalize_positive_int
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Iterable, Iterator
|
||||
@@ -36,6 +37,32 @@ _APPRISE_LOGGER_NAME = "apprise"
|
||||
_APPRISE_DISPATCH_ERRORS = (RuntimeError, TypeError, ValueError)
|
||||
|
||||
|
||||
class _ApprisePluginWithUrl(Protocol):
|
||||
app_id: object
|
||||
|
||||
def url(self, *, privacy: bool = False) -> str:
|
||||
_ = privacy
|
||||
return ""
|
||||
|
||||
|
||||
class _AppriseClient(Protocol):
|
||||
asset: object
|
||||
|
||||
def add(self, plugin: object) -> object: ...
|
||||
|
||||
def notify(self, *, title: str, body: str, notify_type: object) -> object: ...
|
||||
|
||||
|
||||
def _is_apprise_client(candidate: object) -> TypeGuard[_AppriseClient]:
|
||||
return callable(getattr(candidate, "add", None)) and callable(
|
||||
getattr(candidate, "notify", None)
|
||||
)
|
||||
|
||||
|
||||
def _has_plugin_url(candidate: object) -> TypeGuard[_ApprisePluginWithUrl]:
|
||||
return callable(getattr(candidate, "url", None))
|
||||
|
||||
|
||||
class NotificationEvent(StrEnum):
|
||||
"""Global notification event identifiers."""
|
||||
|
||||
@@ -255,13 +282,7 @@ def _resolve_admin_routes() -> list[dict[str, str]]:
|
||||
|
||||
|
||||
def _normalize_user_id(value: object) -> int | None:
|
||||
try:
|
||||
user_id = int(value)
|
||||
except TypeError, ValueError:
|
||||
return None
|
||||
if user_id < 1:
|
||||
return None
|
||||
return user_id
|
||||
return normalize_positive_int(value)
|
||||
|
||||
|
||||
def _resolve_user_routes(user_id: int | None) -> list[dict[str, str]]:
|
||||
@@ -360,8 +381,9 @@ def _plugin_label(plugin: object, fallback_scheme: str) -> str:
|
||||
app_id = getattr(plugin, "app_id", None)
|
||||
if app_id and str(app_id) != fallback_scheme:
|
||||
privacy_url: str | None = None
|
||||
with suppress(Exception):
|
||||
privacy_url = plugin.url(privacy=True)
|
||||
if _has_plugin_url(plugin):
|
||||
with suppress(Exception):
|
||||
privacy_url = plugin.url(privacy=True)
|
||||
|
||||
suffix = str(app_id)
|
||||
if privacy_url:
|
||||
@@ -507,7 +529,7 @@ def _dispatch_to_apprise(
|
||||
return result
|
||||
|
||||
|
||||
def _create_apprise_client() -> object:
|
||||
def _create_apprise_client() -> _AppriseClient | None:
|
||||
if apprise is None:
|
||||
return None
|
||||
|
||||
@@ -517,7 +539,8 @@ def _create_apprise_client() -> object:
|
||||
|
||||
apprise_asset_cls = getattr(apprise, "AppriseAsset", None)
|
||||
if apprise_asset_cls is None:
|
||||
return apprise_cls()
|
||||
client = apprise_cls()
|
||||
return client if _is_apprise_client(client) else None
|
||||
|
||||
try:
|
||||
asset = apprise_asset_cls(
|
||||
@@ -533,12 +556,14 @@ def _create_apprise_client() -> object:
|
||||
app_desc=_APPRISE_APP_DESC,
|
||||
)
|
||||
except TypeError:
|
||||
return apprise_cls()
|
||||
client = apprise_cls()
|
||||
return client if _is_apprise_client(client) else None
|
||||
|
||||
try:
|
||||
return apprise_cls(asset=asset)
|
||||
client = apprise_cls(asset=asset)
|
||||
except TypeError:
|
||||
return apprise_cls()
|
||||
client = apprise_cls()
|
||||
return client if _is_apprise_client(client) else None
|
||||
|
||||
|
||||
def _send_admin_event(
|
||||
|
||||
@@ -4,13 +4,16 @@ Registers /api/auth/oidc/login and /api/auth/oidc/callback endpoints.
|
||||
Business logic remains in oidc_auth.py.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any, Protocol, TypeGuard
|
||||
from urllib.parse import urlencode, urlsplit, urlunsplit
|
||||
|
||||
from authlib.integrations.base_client.errors import OAuthError
|
||||
from authlib.integrations.flask_client import OAuth
|
||||
from authlib.jose.errors import InvalidClaimError
|
||||
from flask import Flask, Response, jsonify, redirect, request, session
|
||||
from flask import Flask, jsonify, redirect, request, session
|
||||
|
||||
from shelfmark.core.config import config as app_config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
@@ -22,6 +25,8 @@ from shelfmark.core.oidc_auth import (
|
||||
from shelfmark.download.network import get_ssl_verify
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from flask.typing import ResponseReturnValue
|
||||
|
||||
from shelfmark.core.user_db import UserDB
|
||||
|
||||
logger = setup_logger(__name__)
|
||||
@@ -30,18 +35,33 @@ _RETURN_TO_SESSION_KEY = "oidc_return_to"
|
||||
_OIDC_CLIENT_ERRORS = (OAuthError, OSError, RuntimeError, TypeError, ValueError)
|
||||
|
||||
|
||||
class _ClaimsMappingLike(Protocol):
|
||||
"""Protocol for Authlib claims payloads that expose a to_dict method."""
|
||||
|
||||
def to_dict(self) -> Mapping[object, object]: ...
|
||||
|
||||
|
||||
def _has_claims_to_dict(candidate: object) -> TypeGuard[_ClaimsMappingLike]:
|
||||
"""Return True when a claims object exposes a callable to_dict method."""
|
||||
return callable(getattr(candidate, "to_dict", None))
|
||||
|
||||
|
||||
def _normalize_claim_mapping(raw_claims: Mapping[object, object]) -> dict[str, Any]:
|
||||
"""Return only string-keyed claims for downstream OIDC helpers."""
|
||||
return {key: value for key, value in raw_claims.items() if isinstance(key, str)}
|
||||
|
||||
|
||||
def _normalize_claims(raw_claims: object) -> dict[str, Any]:
|
||||
"""Return a plain dict for claims from Authlib token/userinfo payloads."""
|
||||
if raw_claims is None:
|
||||
return {}
|
||||
if isinstance(raw_claims, dict):
|
||||
return raw_claims
|
||||
if hasattr(raw_claims, "to_dict"):
|
||||
return raw_claims.to_dict() # type: ignore[no-any-return]
|
||||
try:
|
||||
return dict(raw_claims)
|
||||
except TypeError, ValueError:
|
||||
return {}
|
||||
if isinstance(raw_claims, Mapping):
|
||||
return _normalize_claim_mapping(raw_claims)
|
||||
if _has_claims_to_dict(raw_claims):
|
||||
converted_claims = raw_claims.to_dict()
|
||||
if isinstance(converted_claims, Mapping):
|
||||
return _normalize_claim_mapping(converted_claims)
|
||||
return {}
|
||||
|
||||
|
||||
def _has_username_or_email(claims: dict[str, Any]) -> bool:
|
||||
@@ -135,8 +155,14 @@ def _get_oidc_client() -> tuple[Any, dict[str, Any]]:
|
||||
|
||||
scopes = list(dict.fromkeys(["openid", *scope_values]))
|
||||
|
||||
admin_group = app_config.get("OIDC_ADMIN_GROUP", "")
|
||||
group_claim = app_config.get("OIDC_GROUP_CLAIM", "groups")
|
||||
admin_group_value = app_config.get("OIDC_ADMIN_GROUP", "")
|
||||
admin_group = admin_group_value.strip() if isinstance(admin_group_value, str) else ""
|
||||
group_claim_value = app_config.get("OIDC_GROUP_CLAIM", "groups")
|
||||
group_claim = (
|
||||
group_claim_value.strip()
|
||||
if isinstance(group_claim_value, str) and group_claim_value.strip()
|
||||
else "groups"
|
||||
)
|
||||
use_admin_group = app_config.get("OIDC_USE_ADMIN_GROUP", True)
|
||||
if admin_group and use_admin_group and group_claim and group_claim not in scopes:
|
||||
scopes.append(group_claim)
|
||||
@@ -179,7 +205,7 @@ def register_oidc_routes(app: Flask, user_db: UserDB) -> None:
|
||||
oauth.init_app(app)
|
||||
|
||||
@app.route("/api/auth/oidc/login", methods=["GET"])
|
||||
def oidc_login() -> Response | tuple[Response, int]:
|
||||
def oidc_login() -> ResponseReturnValue:
|
||||
"""Initiate OIDC login flow and redirect to the provider."""
|
||||
try:
|
||||
client, _ = _get_oidc_client()
|
||||
@@ -197,7 +223,7 @@ def register_oidc_routes(app: Flask, user_db: UserDB) -> None:
|
||||
return jsonify({"error": "OIDC login failed"}), 500
|
||||
|
||||
@app.route("/api/auth/oidc/callback", methods=["GET"])
|
||||
def oidc_callback() -> Response | tuple[Response, int]:
|
||||
def oidc_callback() -> ResponseReturnValue:
|
||||
"""Handle OIDC callback from identity provider."""
|
||||
try:
|
||||
error = request.args.get("error")
|
||||
|
||||
@@ -24,7 +24,8 @@ class PrefixMiddleware:
|
||||
|
||||
def __call__(self, environ: dict[str, object], start_response: Callable[..., object]) -> object:
|
||||
"""Rewrite prefixed requests before handing them to the wrapped app."""
|
||||
path = environ.get("PATH_INFO", "") or ""
|
||||
raw_path = environ.get("PATH_INFO", "")
|
||||
path = raw_path if isinstance(raw_path, str) else str(raw_path or "")
|
||||
|
||||
if path in self.bypass_paths:
|
||||
return self.app(environ, start_response)
|
||||
|
||||
+20
-1
@@ -23,6 +23,20 @@ logger = setup_logger(__name__)
|
||||
_QUEUE_HOOK_ERRORS = (OSError, RuntimeError, TypeError, ValueError)
|
||||
|
||||
|
||||
def _coerce_status_timeout_seconds(value: object, *, default: int) -> int:
|
||||
"""Normalize STATUS_TIMEOUT into a usable positive integer."""
|
||||
if isinstance(value, bool):
|
||||
return default
|
||||
if isinstance(value, int):
|
||||
return value if value > 0 else default
|
||||
if isinstance(value, str):
|
||||
stripped = value.strip()
|
||||
if stripped.isdigit():
|
||||
parsed = int(stripped)
|
||||
return parsed if parsed > 0 else default
|
||||
return default
|
||||
|
||||
|
||||
class BookQueue:
|
||||
"""Thread-safe download queue manager with priority support and cancellation."""
|
||||
|
||||
@@ -41,7 +55,12 @@ class BookQueue:
|
||||
@property
|
||||
def _status_timeout(self) -> timedelta:
|
||||
"""Get status timeout from config (allows live updates)."""
|
||||
return timedelta(seconds=app_config.get("STATUS_TIMEOUT", 3600))
|
||||
return timedelta(
|
||||
seconds=_coerce_status_timeout_seconds(
|
||||
app_config.get("STATUS_TIMEOUT", 3600),
|
||||
default=3600,
|
||||
)
|
||||
)
|
||||
|
||||
def add(self, task: DownloadTask) -> bool:
|
||||
"""Add a download task to the queue. Returns False if already exists."""
|
||||
|
||||
@@ -3,13 +3,46 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
from typing import Any, Protocol, SupportsIndex, SupportsInt, TypeGuard
|
||||
|
||||
from shelfmark.core.config import config as app_config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
|
||||
_logger = setup_logger(__name__)
|
||||
|
||||
type _ConvertibleToInt = str | bytes | bytearray | SupportsInt | SupportsIndex
|
||||
|
||||
|
||||
class _MappingWithGet(Protocol):
|
||||
"""Minimal mapping protocol for session-like objects."""
|
||||
|
||||
def get(self, key: str, default: object = None, /) -> object: ...
|
||||
|
||||
|
||||
class _UserDBLike(Protocol):
|
||||
"""Minimal user DB protocol for username population helpers."""
|
||||
|
||||
def get_user(self, *, user_id: int) -> dict[str, Any] | None: ...
|
||||
|
||||
|
||||
def _is_mapping_with_get(candidate: object) -> TypeGuard[_MappingWithGet]:
|
||||
"""Return True when *candidate* exposes a mapping-style get method."""
|
||||
return callable(getattr(candidate, "get", None))
|
||||
|
||||
|
||||
def _is_user_db_like(candidate: object) -> TypeGuard[_UserDBLike]:
|
||||
"""Return True when *candidate* exposes the user lookup API we need."""
|
||||
return callable(getattr(candidate, "get_user", None))
|
||||
|
||||
|
||||
def _is_convertible_to_int(value: object) -> TypeGuard[_ConvertibleToInt]:
|
||||
"""Return True when *value* can be passed to ``int`` safely."""
|
||||
return (
|
||||
isinstance(value, (str, bytes, bytearray))
|
||||
or hasattr(value, "__int__")
|
||||
or hasattr(value, "__index__")
|
||||
)
|
||||
|
||||
|
||||
def now_utc_iso() -> str:
|
||||
"""Return the current UTC time as a seconds-precision ISO 8601 string."""
|
||||
@@ -65,15 +98,17 @@ def coerce_bool(value: object, *, default: bool = False) -> bool:
|
||||
|
||||
def get_session_db_user_id(session_obj: object) -> int | None:
|
||||
"""Extract and coerce `db_user_id` from a Flask session to ``int | None``."""
|
||||
raw = session_obj.get("db_user_id") if session_obj is not None else None
|
||||
raw = session_obj.get("db_user_id") if _is_mapping_with_get(session_obj) else None
|
||||
try:
|
||||
return int(raw) if raw is not None else None
|
||||
return int(raw) if raw is not None and _is_convertible_to_int(raw) else None
|
||||
except TypeError, ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def coerce_int(value: object, default: int) -> int:
|
||||
"""Best-effort integer coercion with fallback to default."""
|
||||
if not _is_convertible_to_int(value):
|
||||
return default
|
||||
try:
|
||||
return int(value)
|
||||
except TypeError, ValueError:
|
||||
@@ -90,6 +125,8 @@ def normalize_optional_text(value: object) -> str | None:
|
||||
|
||||
def normalize_positive_int(value: object) -> int | None:
|
||||
"""Parse *value* as a positive integer, returning ``None`` on failure."""
|
||||
if not _is_convertible_to_int(value):
|
||||
return None
|
||||
try:
|
||||
parsed = int(value)
|
||||
except TypeError, ValueError:
|
||||
@@ -105,6 +142,9 @@ def normalize_optional_positive_int(value: object, field_name: str = "value") ->
|
||||
"""
|
||||
if value is None:
|
||||
return None
|
||||
if not _is_convertible_to_int(value):
|
||||
msg = f"{field_name} must be a positive integer when provided"
|
||||
raise ValueError(msg)
|
||||
try:
|
||||
parsed = int(value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
@@ -118,9 +158,15 @@ def normalize_optional_positive_int(value: object, field_name: str = "value") ->
|
||||
|
||||
def populate_request_usernames(rows: list[dict[str, Any]], user_db: object) -> None:
|
||||
"""Add 'username' to each request row by looking up user_id."""
|
||||
if not _is_user_db_like(user_db):
|
||||
return
|
||||
|
||||
cache: dict[int, str] = {}
|
||||
for row in rows:
|
||||
requester_id = row["user_id"]
|
||||
requester_id = normalize_positive_int(row.get("user_id"))
|
||||
if requester_id is None:
|
||||
row["username"] = ""
|
||||
continue
|
||||
if requester_id not in cache:
|
||||
requester = user_db.get_user(user_id=requester_id)
|
||||
cache[requester_id] = requester.get("username", "") if requester else ""
|
||||
|
||||
@@ -45,6 +45,8 @@ from shelfmark.core.requests_service import (
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
|
||||
from flask.typing import ResponseReturnValue
|
||||
|
||||
from shelfmark.core.user_db import UserDB
|
||||
|
||||
logger = setup_logger(__name__)
|
||||
@@ -81,7 +83,8 @@ def _require_request_endpoints_available(
|
||||
return None
|
||||
|
||||
|
||||
def _require_db_user_id() -> tuple[int | None, object | None]:
|
||||
def _require_db_user_id() -> tuple[int | None, ResponseReturnValue | None]:
|
||||
"""Return the logged-in DB user id or a ready-made error response."""
|
||||
raw_user_id = session.get("db_user_id")
|
||||
if raw_user_id is None:
|
||||
return None, _error_response(
|
||||
@@ -89,26 +92,26 @@ def _require_db_user_id() -> tuple[int | None, object | None]:
|
||||
403,
|
||||
code="user_identity_unavailable",
|
||||
)
|
||||
try:
|
||||
return int(raw_user_id), None
|
||||
except TypeError, ValueError:
|
||||
normalized_user_id = normalize_positive_int(raw_user_id)
|
||||
if normalized_user_id is None:
|
||||
return None, _error_response(
|
||||
"User identity is unavailable for request workflow",
|
||||
403,
|
||||
code="user_identity_unavailable",
|
||||
)
|
||||
return normalized_user_id, None
|
||||
|
||||
|
||||
def _require_admin_user_id() -> tuple[int | None, object | None]:
|
||||
def _require_admin_user_id() -> tuple[int | None, ResponseReturnValue | None]:
|
||||
if not session.get("is_admin", False):
|
||||
return None, (jsonify({"error": "Admin access required"}), 403)
|
||||
raw_admin_id = session.get("db_user_id")
|
||||
if raw_admin_id is None:
|
||||
return None, (jsonify({"error": "Admin user identity unavailable"}), 403)
|
||||
try:
|
||||
return int(raw_admin_id), None
|
||||
except TypeError, ValueError:
|
||||
normalized_admin_user_id = normalize_positive_int(raw_admin_id)
|
||||
if normalized_admin_user_id is None:
|
||||
return None, (jsonify({"error": "Admin user identity unavailable"}), 403)
|
||||
return normalized_admin_user_id, None
|
||||
|
||||
|
||||
def _resolve_effective_policy(
|
||||
@@ -257,13 +260,8 @@ def _resolve_request_user_context(
|
||||
msg = "Admin required"
|
||||
raise RequestServiceError(msg, status_code=403)
|
||||
|
||||
try:
|
||||
target_user_id = int(on_behalf_of_user_id)
|
||||
except (TypeError, ValueError) as exc:
|
||||
msg = "Invalid on_behalf_of_user_id"
|
||||
raise RequestServiceError(msg, status_code=400) from exc
|
||||
|
||||
if target_user_id <= 0:
|
||||
target_user_id = normalize_positive_int(on_behalf_of_user_id)
|
||||
if target_user_id is None:
|
||||
msg = "Invalid on_behalf_of_user_id"
|
||||
raise RequestServiceError(msg, status_code=400)
|
||||
|
||||
@@ -539,7 +537,7 @@ def register_request_routes(
|
||||
"""Register request policy and request lifecycle routes."""
|
||||
|
||||
@app.route("/api/request-policy", methods=["GET"])
|
||||
def api_request_policy() -> Response | tuple[Response, int]:
|
||||
def api_request_policy() -> ResponseReturnValue:
|
||||
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
|
||||
if auth_gate is not None:
|
||||
return auth_gate
|
||||
@@ -551,12 +549,7 @@ def register_request_routes(
|
||||
if db_gate is not None:
|
||||
return db_gate
|
||||
else:
|
||||
raw_id = session.get("db_user_id")
|
||||
if raw_id is not None:
|
||||
try:
|
||||
db_user_id = int(raw_id)
|
||||
except TypeError, ValueError:
|
||||
db_user_id = None
|
||||
db_user_id = normalize_positive_int(session.get("db_user_id"))
|
||||
|
||||
global_settings, user_settings, effective, requests_enabled = _resolve_effective_policy(
|
||||
user_db,
|
||||
@@ -616,7 +609,7 @@ def register_request_routes(
|
||||
)
|
||||
|
||||
@app.route("/api/requests", methods=["POST"])
|
||||
def api_create_request() -> Response | tuple[Response, int]:
|
||||
def api_create_request() -> ResponseReturnValue:
|
||||
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
|
||||
if auth_gate is not None:
|
||||
return auth_gate
|
||||
@@ -682,7 +675,7 @@ def register_request_routes(
|
||||
return jsonify(created), 201
|
||||
|
||||
@app.route("/api/requests/batch", methods=["POST"])
|
||||
def api_create_requests_batch() -> Response | tuple[Response, int]:
|
||||
def api_create_requests_batch() -> ResponseReturnValue:
|
||||
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
|
||||
if auth_gate is not None:
|
||||
return auth_gate
|
||||
@@ -796,14 +789,20 @@ def register_request_routes(
|
||||
return jsonify(ordered_results), status_code
|
||||
|
||||
@app.route("/api/requests", methods=["GET"])
|
||||
def api_list_requests() -> Response | tuple[Response, int]:
|
||||
def api_list_requests() -> ResponseReturnValue:
|
||||
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
|
||||
if auth_gate is not None:
|
||||
return auth_gate
|
||||
|
||||
db_user_id, db_gate = _require_db_user_id()
|
||||
if db_gate is not None or db_user_id is None:
|
||||
if db_gate is not None:
|
||||
return db_gate
|
||||
if db_user_id is None:
|
||||
return _error_response(
|
||||
"User identity is unavailable for request workflow",
|
||||
403,
|
||||
code="user_identity_unavailable",
|
||||
)
|
||||
|
||||
status = request.args.get("status")
|
||||
limit = request.args.get("limit", type=int)
|
||||
@@ -821,14 +820,20 @@ def register_request_routes(
|
||||
return jsonify(rows)
|
||||
|
||||
@app.route("/api/requests/<int:request_id>", methods=["DELETE"])
|
||||
def api_cancel_request(request_id: int) -> Response | tuple[Response, int]:
|
||||
def api_cancel_request(request_id: int) -> ResponseReturnValue:
|
||||
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
|
||||
if auth_gate is not None:
|
||||
return auth_gate
|
||||
|
||||
db_user_id, db_gate = _require_db_user_id()
|
||||
if db_gate is not None or db_user_id is None:
|
||||
if db_gate is not None:
|
||||
return db_gate
|
||||
if db_user_id is None:
|
||||
return _error_response(
|
||||
"User identity is unavailable for request workflow",
|
||||
403,
|
||||
code="user_identity_unavailable",
|
||||
)
|
||||
|
||||
try:
|
||||
updated = cancel_request(
|
||||
@@ -869,7 +874,7 @@ def register_request_routes(
|
||||
return jsonify(updated)
|
||||
|
||||
@app.route("/api/admin/requests", methods=["GET"])
|
||||
def api_admin_list_requests() -> Response | tuple[Response, int]:
|
||||
def api_admin_list_requests() -> ResponseReturnValue:
|
||||
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
|
||||
if auth_gate is not None:
|
||||
return auth_gate
|
||||
@@ -890,7 +895,7 @@ def register_request_routes(
|
||||
return jsonify(rows)
|
||||
|
||||
@app.route("/api/admin/requests/count", methods=["GET"])
|
||||
def api_admin_request_counts() -> Response | tuple[Response, int]:
|
||||
def api_admin_request_counts() -> ResponseReturnValue:
|
||||
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
|
||||
if auth_gate is not None:
|
||||
return auth_gate
|
||||
@@ -907,7 +912,7 @@ def register_request_routes(
|
||||
)
|
||||
|
||||
@app.route("/api/admin/requests/<int:request_id>/fulfil", methods=["POST"])
|
||||
def api_admin_fulfil_request(request_id: int) -> Response | tuple[Response, int]:
|
||||
def api_admin_fulfil_request(request_id: int) -> ResponseReturnValue:
|
||||
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
|
||||
if auth_gate is not None:
|
||||
return auth_gate
|
||||
@@ -915,6 +920,8 @@ def register_request_routes(
|
||||
admin_user_id, admin_gate = _require_admin_user_id()
|
||||
if admin_gate is not None:
|
||||
return admin_gate
|
||||
if admin_user_id is None:
|
||||
return jsonify({"error": "Admin user identity unavailable"}), 403
|
||||
|
||||
data = request.get_json(silent=True) or {}
|
||||
if not isinstance(data, dict):
|
||||
@@ -971,7 +978,7 @@ def register_request_routes(
|
||||
return jsonify(updated)
|
||||
|
||||
@app.route("/api/admin/requests/<int:request_id>/reject", methods=["POST"])
|
||||
def api_admin_reject_request(request_id: int) -> Response | tuple[Response, int]:
|
||||
def api_admin_reject_request(request_id: int) -> ResponseReturnValue:
|
||||
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
|
||||
if auth_gate is not None:
|
||||
return auth_gate
|
||||
@@ -979,6 +986,8 @@ def register_request_routes(
|
||||
admin_user_id, admin_gate = _require_admin_user_id()
|
||||
if admin_gate is not None:
|
||||
return admin_gate
|
||||
if admin_user_id is None:
|
||||
return jsonify({"error": "Admin user identity unavailable"}), 403
|
||||
|
||||
data = request.get_json(silent=True) or {}
|
||||
if not isinstance(data, dict):
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
@@ -53,10 +54,14 @@ class ReleaseSearchPlan:
|
||||
|
||||
def _normalize_languages(languages: list[str] | None) -> list[str] | None:
|
||||
if not languages:
|
||||
default = config.BOOK_LANGUAGE
|
||||
if not default:
|
||||
default = getattr(config, "BOOK_LANGUAGE", None)
|
||||
if isinstance(default, str):
|
||||
default_values: list[object] = [default]
|
||||
elif isinstance(default, Iterable) and not isinstance(default, (bytes, bytearray, dict)):
|
||||
default_values = list(default)
|
||||
else:
|
||||
return None
|
||||
return [str(lang).strip() for lang in default if str(lang).strip()]
|
||||
return [str(lang).strip() for lang in default_values if str(lang).strip()]
|
||||
|
||||
normalized: list[str] = []
|
||||
for lang in languages:
|
||||
|
||||
@@ -55,8 +55,10 @@ _CONFIG_REFRESH_ERRORS = (ImportError, OSError, RuntimeError, TypeError, ValueEr
|
||||
|
||||
def _get_current_user(
|
||||
user_db: UserDB,
|
||||
) -> tuple[int | None, dict[str, Any] | None, tuple[Any, int] | None]:
|
||||
) -> tuple[int | None, dict[str, Any] | None, tuple[Response, int] | None]:
|
||||
raw_user_id = session.get("db_user_id")
|
||||
if raw_user_id is None:
|
||||
return None, None, (jsonify({"error": "Invalid user context"}), 400)
|
||||
try:
|
||||
user_id = int(raw_user_id)
|
||||
except TypeError, ValueError:
|
||||
@@ -206,6 +208,8 @@ def register_self_user_routes(app: Flask, user_db: UserDB) -> None:
|
||||
user_id, user, user_error = _get_current_user(user_db)
|
||||
if user_error:
|
||||
return user_error
|
||||
if user_id is None or user is None:
|
||||
return jsonify({"error": "User not found"}), 404
|
||||
|
||||
serialized_user = _serialize_self_user(user, g.auth_mode)
|
||||
serialized_user["settings"] = user_db.get_user_settings(user_id)
|
||||
@@ -286,6 +290,8 @@ def register_self_user_routes(app: Flask, user_db: UserDB) -> None:
|
||||
user_id, user, user_error = _get_current_user(user_db)
|
||||
if user_error:
|
||||
return user_error
|
||||
if user_id is None or user is None:
|
||||
return jsonify({"error": "User not found"}), 404
|
||||
|
||||
data = request.get_json() or {}
|
||||
if not isinstance(data, dict):
|
||||
|
||||
@@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Any
|
||||
from werkzeug.utils import secure_filename
|
||||
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.request_helpers import coerce_bool, normalize_optional_text
|
||||
|
||||
logger = setup_logger(__name__)
|
||||
_SETTINGS_LIVE_APPLY_ERRORS = (OSError, RuntimeError, TypeError, ValueError)
|
||||
@@ -27,7 +28,7 @@ class FieldBase:
|
||||
key: str # Environment variable / config key
|
||||
label: str # Display label in UI
|
||||
description: str = "" # Help text
|
||||
default: object = None # Default value if not set
|
||||
default: Any = None # Default value if not set
|
||||
required: bool = False # Whether field must have a value
|
||||
env_var: str | None = None # Override env var name (defaults to key)
|
||||
env_supported: bool = True # Whether this setting can be set via ENV var (False = UI-only)
|
||||
@@ -219,7 +220,7 @@ class HeadingField:
|
||||
|
||||
|
||||
# Type alias for all field types
|
||||
SettingsField = (
|
||||
ValueField = (
|
||||
TextField
|
||||
| PasswordField
|
||||
| NumberField
|
||||
@@ -229,11 +230,10 @@ SettingsField = (
|
||||
| TagListField
|
||||
| OrderableListField
|
||||
| TableField
|
||||
| CustomComponentField
|
||||
| ActionButton
|
||||
| HeadingField
|
||||
)
|
||||
|
||||
SettingsField = ValueField | CustomComponentField | ActionButton | HeadingField
|
||||
|
||||
|
||||
@dataclass
|
||||
class SettingsTab:
|
||||
@@ -332,23 +332,21 @@ def get_all_settings_tabs() -> list[SettingsTab]:
|
||||
return sorted(_SETTINGS_REGISTRY.values(), key=lambda t: (t.order, t.name))
|
||||
|
||||
|
||||
def _iter_value_fields(tab: SettingsTab) -> Iterator[SettingsField]:
|
||||
def _iter_value_fields(tab: SettingsTab) -> Iterator[FieldBase]:
|
||||
"""Yield value-bearing fields for a tab."""
|
||||
for settings_field in tab.fields:
|
||||
if isinstance(settings_field, CustomComponentField):
|
||||
for value_field in settings_field.value_fields:
|
||||
if isinstance(value_field, (ActionButton, HeadingField, CustomComponentField)):
|
||||
continue
|
||||
yield value_field
|
||||
if isinstance(value_field, FieldBase):
|
||||
yield value_field
|
||||
continue
|
||||
if isinstance(settings_field, (ActionButton, HeadingField)):
|
||||
continue
|
||||
yield settings_field
|
||||
if isinstance(settings_field, FieldBase):
|
||||
yield settings_field
|
||||
|
||||
|
||||
def get_settings_field_map(
|
||||
tab_name: str | None = None,
|
||||
) -> dict[str, tuple[SettingsField, str]]:
|
||||
) -> dict[str, tuple[FieldBase, str]]:
|
||||
"""Return key -> (field, tab_name) map for value-bearing settings fields."""
|
||||
tabs: list[SettingsTab]
|
||||
if tab_name:
|
||||
@@ -359,7 +357,7 @@ def get_settings_field_map(
|
||||
else:
|
||||
tabs = get_all_settings_tabs()
|
||||
|
||||
field_map: dict[str, tuple[SettingsField, str]] = {}
|
||||
field_map: dict[str, tuple[FieldBase, str]] = {}
|
||||
for tab in tabs:
|
||||
for settings_field in _iter_value_fields(tab):
|
||||
field_map[settings_field.key] = (settings_field, tab.name)
|
||||
@@ -368,7 +366,7 @@ def get_settings_field_map(
|
||||
|
||||
def get_user_overridable_fields(
|
||||
tab_name: str | None = None,
|
||||
) -> dict[str, tuple[SettingsField, str]]:
|
||||
) -> dict[str, tuple[FieldBase, str]]:
|
||||
"""Return key -> (field, tab_name) map for fields marked user_overridable."""
|
||||
field_map = get_settings_field_map(tab_name=tab_name)
|
||||
return {
|
||||
@@ -836,11 +834,8 @@ def migrate_download_to_browser_settings() -> None:
|
||||
logger.exception("Failed to migrate download-to-browser settings")
|
||||
|
||||
|
||||
def get_setting_value(field: SettingsField, tab_name: str) -> object:
|
||||
def get_setting_value(field: FieldBase, tab_name: str) -> object:
|
||||
"""Resolve the effective value for a settings field."""
|
||||
if isinstance(field, (ActionButton, HeadingField, CustomComponentField)):
|
||||
return None # Actions and headings don't have values
|
||||
|
||||
# 1. Check environment variable (if supported for this field)
|
||||
if field.env_supported:
|
||||
env_var_name = field.get_env_var_name()
|
||||
@@ -857,7 +852,7 @@ def get_setting_value(field: SettingsField, tab_name: str) -> object:
|
||||
return field.default
|
||||
|
||||
|
||||
def _parse_env_value(value: str, field: SettingsField) -> object:
|
||||
def _parse_env_value(value: str, field: FieldBase) -> object:
|
||||
"""Parse an environment variable value to the appropriate type."""
|
||||
if isinstance(field, CheckboxField):
|
||||
return value.lower() in ("true", "1", "yes", "on")
|
||||
@@ -889,10 +884,8 @@ def _parse_env_value(value: str, field: SettingsField) -> object:
|
||||
return value
|
||||
|
||||
|
||||
def is_value_from_env(field: SettingsField) -> bool:
|
||||
def is_value_from_env(field: FieldBase) -> bool:
|
||||
"""Check if a field's value comes from an environment variable."""
|
||||
if isinstance(field, (ActionButton, HeadingField, CustomComponentField)):
|
||||
return False
|
||||
# UI-only settings never come from ENV (env_supported=False)
|
||||
if not getattr(field, "env_supported", True):
|
||||
return False
|
||||
@@ -918,7 +911,7 @@ def serialize_field(
|
||||
"""
|
||||
# CustomComponentField has a custom structure - handle separately
|
||||
if isinstance(field, CustomComponentField):
|
||||
result: dict[str, Any] = {
|
||||
component_result: dict[str, Any] = {
|
||||
"key": field.key,
|
||||
"label": field.label,
|
||||
"type": field.get_field_type(),
|
||||
@@ -939,31 +932,31 @@ def serialize_field(
|
||||
)
|
||||
serialized_bound_field["hiddenInUi"] = True
|
||||
bound_fields.append(serialized_bound_field)
|
||||
result["boundFields"] = bound_fields
|
||||
component_result["boundFields"] = bound_fields
|
||||
if field.show_when:
|
||||
result["showWhen"] = field.show_when
|
||||
component_result["showWhen"] = field.show_when
|
||||
if field.universal_only:
|
||||
result["universalOnly"] = True
|
||||
return result
|
||||
component_result["universalOnly"] = True
|
||||
return component_result
|
||||
|
||||
# HeadingField has a different structure - handle separately
|
||||
if isinstance(field, HeadingField):
|
||||
result: dict[str, Any] = {
|
||||
heading_result: dict[str, Any] = {
|
||||
"key": field.key,
|
||||
"type": field.get_field_type(),
|
||||
"title": field.title,
|
||||
"description": field.description,
|
||||
}
|
||||
if field.description_by_auth_mode:
|
||||
result["descriptionByAuthMode"] = field.description_by_auth_mode
|
||||
heading_result["descriptionByAuthMode"] = field.description_by_auth_mode
|
||||
if field.link_url:
|
||||
result["linkUrl"] = field.link_url
|
||||
result["linkText"] = field.link_text or field.link_url
|
||||
heading_result["linkUrl"] = field.link_url
|
||||
heading_result["linkText"] = field.link_text or field.link_url
|
||||
if field.show_when:
|
||||
result["showWhen"] = field.show_when
|
||||
heading_result["showWhen"] = field.show_when
|
||||
if field.universal_only:
|
||||
result["universalOnly"] = True
|
||||
return result
|
||||
heading_result["universalOnly"] = True
|
||||
return heading_result
|
||||
|
||||
result: dict[str, Any] = {
|
||||
"key": field.key,
|
||||
@@ -1166,12 +1159,12 @@ def _apply_dns_settings(config: Config) -> None:
|
||||
try:
|
||||
from shelfmark.download import network
|
||||
|
||||
provider = config.get("CUSTOM_DNS", "auto")
|
||||
use_doh = config.get("USE_DOH", False)
|
||||
provider = normalize_optional_text(config.get("CUSTOM_DNS", "auto")) or "auto"
|
||||
use_doh = coerce_bool(config.get("USE_DOH", False), default=False)
|
||||
manual_servers = None
|
||||
|
||||
if provider == "manual":
|
||||
manual_dns = config.get("CUSTOM_DNS_MANUAL", "")
|
||||
manual_dns = normalize_optional_text(config.get("CUSTOM_DNS_MANUAL", ""))
|
||||
if manual_dns:
|
||||
# Parse comma-separated server list
|
||||
manual_servers = [s.strip() for s in manual_dns.split(",") if s.strip()]
|
||||
|
||||
@@ -116,6 +116,14 @@ WHERE dismissed_at IS NOT NULL;
|
||||
"""
|
||||
|
||||
|
||||
def _require_loaded_user(user: dict[str, Any] | None) -> dict[str, Any]:
|
||||
"""Return a loaded user row or raise when the DB insert result is inconsistent."""
|
||||
if user is None:
|
||||
msg = "Failed to load newly created user"
|
||||
raise RuntimeError(msg)
|
||||
return user
|
||||
|
||||
|
||||
def get_users_db_path(config_dir: str | None = None) -> str:
|
||||
"""Return the configured users database path."""
|
||||
root = config_dir or os.environ.get("CONFIG_DIR", "/config")
|
||||
@@ -296,7 +304,11 @@ class UserDB:
|
||||
)
|
||||
conn.commit()
|
||||
user_id = cursor.lastrowid
|
||||
return self._get_user_by_id(conn, user_id)
|
||||
if not isinstance(user_id, int):
|
||||
msg = "Failed to create user"
|
||||
raise TypeError(msg)
|
||||
created_user = self._get_user_by_id(conn, user_id)
|
||||
return _require_loaded_user(created_user)
|
||||
except sqlite3.IntegrityError as e:
|
||||
msg = f"User already exists: {e}"
|
||||
raise ValueError(msg) from e
|
||||
|
||||
+30
-5
@@ -10,6 +10,8 @@ from threading import Lock
|
||||
from typing import TYPE_CHECKING
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from shelfmark.core.request_helpers import normalize_optional_text
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from types import ModuleType
|
||||
|
||||
@@ -251,9 +253,9 @@ def get_aa_content_type_dir(content_type: str | None = None) -> Path | None:
|
||||
for mapping in (_AA_CONTENT_TYPE_TO_CONFIG_KEY, _LEGACY_CONTENT_TYPE_TO_CONFIG_KEY):
|
||||
config_key = mapping.get(content_type_lower)
|
||||
if config_key:
|
||||
custom_dir = config.get(config_key, "")
|
||||
if custom_dir:
|
||||
return Path(custom_dir)
|
||||
custom_dir = _coerce_config_path(config.get(config_key, ""))
|
||||
if custom_dir is not None:
|
||||
return custom_dir
|
||||
|
||||
return None
|
||||
|
||||
@@ -263,7 +265,11 @@ def get_ingest_dir(content_type: str | None = None) -> Path:
|
||||
from shelfmark.core.config import config
|
||||
|
||||
# Check new DESTINATION setting first, then legacy INGEST_DIR
|
||||
default_ingest_dir = Path(config.get("DESTINATION", "") or config.get("INGEST_DIR", "/books"))
|
||||
default_ingest_dir = _coerce_config_path(config.get("DESTINATION", "")) or _coerce_config_path(
|
||||
config.get("INGEST_DIR", "/books")
|
||||
)
|
||||
if default_ingest_dir is None:
|
||||
default_ingest_dir = Path("/books")
|
||||
|
||||
if not content_type:
|
||||
return default_ingest_dir
|
||||
@@ -295,7 +301,26 @@ def transform_cover_url(cover_url: str | None, cache_id: str) -> str | None:
|
||||
|
||||
# Encode the original URL and create a proxy URL
|
||||
encoded_url = base64.urlsafe_b64encode(cover_url.encode()).decode()
|
||||
base_path = normalize_base_path(app_config.get("URL_BASE", ""))
|
||||
base_path = normalize_base_path(normalize_optional_text(app_config.get("URL_BASE", "")))
|
||||
if base_path:
|
||||
return f"{base_path}/api/covers/{cache_id}?url={encoded_url}"
|
||||
return f"/api/covers/{cache_id}?url={encoded_url}"
|
||||
|
||||
|
||||
def _coerce_config_path(value: object) -> Path | None:
|
||||
if isinstance(value, os.PathLike):
|
||||
path_value = os.fspath(value)
|
||||
if isinstance(path_value, str):
|
||||
normalized = path_value.strip()
|
||||
if normalized:
|
||||
return Path(normalized)
|
||||
return None
|
||||
|
||||
if not isinstance(value, str):
|
||||
return None
|
||||
|
||||
normalized = value.strip()
|
||||
if not normalized:
|
||||
return None
|
||||
|
||||
return Path(normalized)
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
"""Shared coercion helpers for download client config and option values."""
|
||||
|
||||
from shelfmark.core.utils import normalize_http_url
|
||||
|
||||
|
||||
def config_text(value: object, default: str = "") -> str:
|
||||
"""Coerce config values to strings without losing explicit empty defaults."""
|
||||
if value is None:
|
||||
return default
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
return str(value)
|
||||
|
||||
|
||||
def normalize_http_config_url(value: object, *, require_string: bool = False) -> str:
|
||||
"""Normalize HTTP(S) config URLs with optional strict string-only input handling."""
|
||||
if require_string and not isinstance(value, str):
|
||||
return ""
|
||||
return normalize_http_url(config_text(value))
|
||||
|
||||
|
||||
def coerce_optional_int(value: object) -> int | None:
|
||||
"""Convert optional numeric inputs to ints."""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, int):
|
||||
return value
|
||||
if isinstance(value, float):
|
||||
return int(value)
|
||||
if isinstance(value, str):
|
||||
return int(value)
|
||||
msg = f"Expected int-compatible value, got {type(value).__name__}"
|
||||
raise TypeError(msg)
|
||||
|
||||
|
||||
def coerce_optional_float(value: object) -> float | None:
|
||||
"""Convert optional numeric inputs to floats."""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, (int, float)):
|
||||
return float(value)
|
||||
if isinstance(value, str):
|
||||
return float(value)
|
||||
msg = f"Expected float-compatible value, got {type(value).__name__}"
|
||||
raise TypeError(msg)
|
||||
@@ -1,15 +1,18 @@
|
||||
"""Shared download handler for external torrent/usenet clients."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import errno
|
||||
import shutil
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Protocol, TypeGuard
|
||||
|
||||
from shelfmark.core.config import config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.request_helpers import normalize_optional_text
|
||||
from shelfmark.core.utils import is_audiobook
|
||||
from shelfmark.download.clients import (
|
||||
DownloadClient,
|
||||
@@ -31,6 +34,19 @@ if TYPE_CHECKING:
|
||||
logger = setup_logger(__name__)
|
||||
_CLIENT_CLEANUP_ERRORS = (AttributeError, KeyError, OSError, RuntimeError, TypeError, ValueError)
|
||||
|
||||
|
||||
class _SabnzbdLikeClient(Protocol):
|
||||
name: str
|
||||
|
||||
def remove(
|
||||
self, download_id: str, *, delete_files: bool = False, archive: bool = True
|
||||
) -> bool: ...
|
||||
|
||||
|
||||
def _is_sabnzbd_like_client(candidate: DownloadClient) -> TypeGuard[_SabnzbdLikeClient]:
|
||||
return getattr(candidate, "name", "") == "sabnzbd"
|
||||
|
||||
|
||||
# How often to poll the download client for status (seconds)
|
||||
POLL_INTERVAL = 2
|
||||
WINDOWS_DRIVE_PREFIX_LENGTH = 2
|
||||
@@ -165,7 +181,16 @@ class ExternalClientHandler(DownloadHandler, ABC):
|
||||
"sabnzbd": "SABNZBD_CATEGORY_AUDIOBOOK",
|
||||
}
|
||||
audiobook_key = audiobook_keys.get(client.name)
|
||||
return config.get(audiobook_key, "") or None if audiobook_key else None
|
||||
if audiobook_key is None:
|
||||
return None
|
||||
configured_category = config.get(audiobook_key, "")
|
||||
normalized_category = normalize_optional_text(configured_category)
|
||||
if normalized_category is not None:
|
||||
return normalized_category
|
||||
if configured_category is None:
|
||||
return None
|
||||
fallback_category = str(configured_category).strip()
|
||||
return fallback_category or None
|
||||
|
||||
def post_process_cleanup(self, task: DownloadTask, *, success: bool) -> None:
|
||||
"""Clean up external-client state after post-processing finishes."""
|
||||
@@ -216,7 +241,7 @@ class ExternalClientHandler(DownloadHandler, ABC):
|
||||
archive: bool = True,
|
||||
) -> None:
|
||||
"""Remove a usenet download with SABnzbd-specific archive handling."""
|
||||
if getattr(client, "name", "") == "sabnzbd":
|
||||
if _is_sabnzbd_like_client(client):
|
||||
client.remove(download_id, delete_files=delete_files, archive=archive)
|
||||
else:
|
||||
client.remove(download_id, delete_files=delete_files)
|
||||
|
||||
@@ -26,6 +26,11 @@ from shelfmark.download.clients import (
|
||||
DownloadStatus,
|
||||
register_client,
|
||||
)
|
||||
from shelfmark.download.clients._coercion import (
|
||||
coerce_optional_float,
|
||||
coerce_optional_int,
|
||||
config_text,
|
||||
)
|
||||
from shelfmark.download.clients.torrent_utils import (
|
||||
extract_torrent_info,
|
||||
)
|
||||
@@ -78,9 +83,9 @@ class DelugeClient(DownloadClient):
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize the client from the configured Deluge connection settings."""
|
||||
raw_host = str(config.get("DELUGE_HOST", "localhost") or "")
|
||||
raw_port = str(config.get("DELUGE_PORT", "8112") or "8112")
|
||||
password = str(config.get("DELUGE_PASSWORD", "") or "")
|
||||
raw_host = config_text(config.get("DELUGE_HOST", "localhost"))
|
||||
raw_port = config_text(config.get("DELUGE_PORT", "8112"), "8112")
|
||||
password = config_text(config.get("DELUGE_PASSWORD", ""))
|
||||
|
||||
if not raw_host:
|
||||
msg = "DELUGE_HOST is required"
|
||||
@@ -124,14 +129,14 @@ class DelugeClient(DownloadClient):
|
||||
self._connected = False
|
||||
self._rpc_id = 0
|
||||
|
||||
self._category = str(config.get("DELUGE_CATEGORY", "books") or "books")
|
||||
self._download_dir = str(config.get("DELUGE_DOWNLOAD_DIR", "") or "")
|
||||
self._category = config_text(config.get("DELUGE_CATEGORY", "books"), "books")
|
||||
self._download_dir = config_text(config.get("DELUGE_DOWNLOAD_DIR", ""))
|
||||
|
||||
def _next_rpc_id(self) -> int:
|
||||
self._rpc_id += 1
|
||||
return self._rpc_id
|
||||
|
||||
def _rpc_call(self, method: str, *params: object, timeout: int = 15) -> object:
|
||||
def _rpc_call(self, method: str, *params: object, timeout: int = 15) -> Any:
|
||||
payload = {
|
||||
"id": self._next_rpc_id(),
|
||||
"method": method,
|
||||
@@ -239,9 +244,9 @@ class DelugeClient(DownloadClient):
|
||||
@staticmethod
|
||||
def is_configured() -> bool:
|
||||
"""Return whether Deluge is the active configured torrent client."""
|
||||
client = config.get("PROWLARR_TORRENT_CLIENT", "")
|
||||
host = config.get("DELUGE_HOST", "")
|
||||
password = config.get("DELUGE_PASSWORD", "")
|
||||
client = config_text(config.get("PROWLARR_TORRENT_CLIENT", ""))
|
||||
host = config_text(config.get("DELUGE_HOST", ""))
|
||||
password = config_text(config.get("DELUGE_PASSWORD", ""))
|
||||
return client == "deluge" and bool(host) and bool(password)
|
||||
|
||||
def test_connection(self) -> tuple[bool, str]:
|
||||
@@ -279,12 +284,12 @@ class DelugeClient(DownloadClient):
|
||||
options["download_location"] = self._download_dir
|
||||
|
||||
# Per-torrent seeding limits from indexer
|
||||
seeding_time_limit = kwargs.get("seeding_time_limit")
|
||||
seeding_time_limit = coerce_optional_int(kwargs.get("seeding_time_limit"))
|
||||
if seeding_time_limit is not None:
|
||||
options["seed_time_limit"] = int(seeding_time_limit)
|
||||
ratio_limit = kwargs.get("ratio_limit")
|
||||
options["seed_time_limit"] = seeding_time_limit
|
||||
ratio_limit = coerce_optional_float(kwargs.get("ratio_limit"))
|
||||
if ratio_limit is not None:
|
||||
options["stop_at_ratio"] = float(ratio_limit)
|
||||
options["stop_at_ratio"] = ratio_limit
|
||||
options["stop_at_ratio_enabled"] = True
|
||||
|
||||
if torrent_info.is_magnet:
|
||||
|
||||
@@ -4,24 +4,29 @@ Uses NZBGet's JSON-RPC API directly via requests (no external dependency).
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, NoReturn
|
||||
|
||||
import requests
|
||||
|
||||
from shelfmark.core.config import config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.utils import normalize_http_url
|
||||
from shelfmark.download.clients import (
|
||||
DownloadClient,
|
||||
DownloadStatus,
|
||||
register_client,
|
||||
with_retry,
|
||||
)
|
||||
from shelfmark.download.clients._coercion import config_text, normalize_http_config_url
|
||||
from shelfmark.download.network import get_ssl_verify
|
||||
|
||||
logger = setup_logger(__name__)
|
||||
_NZBGET_CLIENT_ERRORS = (AttributeError, OSError, RuntimeError, TypeError, ValueError)
|
||||
|
||||
|
||||
def _raise_runtime_error(message: str) -> NoReturn:
|
||||
raise RuntimeError(message)
|
||||
|
||||
|
||||
@register_client("usenet")
|
||||
class NZBGetClient(DownloadClient):
|
||||
"""NZBGet download client using JSON-RPC API."""
|
||||
@@ -31,24 +36,24 @@ class NZBGetClient(DownloadClient):
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize NZBGet client with settings from config."""
|
||||
raw_url = config.get("NZBGET_URL", "")
|
||||
raw_url = config_text(config.get("NZBGET_URL", ""))
|
||||
if not raw_url:
|
||||
msg = "NZBGET_URL is required"
|
||||
raise ValueError(msg)
|
||||
|
||||
self.url = normalize_http_url(raw_url)
|
||||
self.url = normalize_http_config_url(raw_url)
|
||||
if not self.url:
|
||||
msg = "NZBGET_URL is invalid"
|
||||
raise ValueError(msg)
|
||||
self.username = config.get("NZBGET_USERNAME", "nzbget")
|
||||
self.password = config.get("NZBGET_PASSWORD", "")
|
||||
self._category = config.get("NZBGET_CATEGORY", "Books")
|
||||
self.username = config_text(config.get("NZBGET_USERNAME", "nzbget"), "nzbget")
|
||||
self.password = config_text(config.get("NZBGET_PASSWORD", ""))
|
||||
self._category = config_text(config.get("NZBGET_CATEGORY", "Books"), "Books")
|
||||
|
||||
@staticmethod
|
||||
def is_configured() -> bool:
|
||||
"""Check if NZBGet is configured and selected as the usenet client."""
|
||||
client = config.get("PROWLARR_USENET_CLIENT", "")
|
||||
url = normalize_http_url(config.get("NZBGET_URL", ""))
|
||||
client = config_text(config.get("PROWLARR_USENET_CLIENT", ""))
|
||||
url = normalize_http_config_url(config.get("NZBGET_URL", ""))
|
||||
return client == "nzbget" and bool(url)
|
||||
|
||||
def _try_remove_command(
|
||||
@@ -65,7 +70,7 @@ class NZBGetClient(DownloadClient):
|
||||
return False, None
|
||||
|
||||
@with_retry()
|
||||
def _rpc_call(self, method: str, params: list | None = None) -> object:
|
||||
def _rpc_call(self, method: str, params: list[object] | None = None) -> Any:
|
||||
"""Make a JSON-RPC call to NZBGet.
|
||||
|
||||
Args:
|
||||
@@ -151,11 +156,7 @@ class NZBGetClient(DownloadClient):
|
||||
import base64
|
||||
|
||||
# Use configured category if not explicitly provided
|
||||
category = category or self._category
|
||||
|
||||
def _raise_invalid_nzb_id() -> None:
|
||||
msg = "NZBGet returned invalid ID"
|
||||
raise RuntimeError(msg)
|
||||
resolved_category = category or self._category
|
||||
|
||||
try:
|
||||
# Fetch NZB content from the URL (handles Prowlarr proxy redirects)
|
||||
@@ -175,7 +176,7 @@ class NZBGetClient(DownloadClient):
|
||||
[
|
||||
nzb_filename, # NZBFilename
|
||||
nzb_content, # Content (base64-encoded NZB)
|
||||
category, # Category
|
||||
resolved_category, # Category
|
||||
0, # Priority (0 = normal)
|
||||
False, # AddToTop
|
||||
False, # AddPaused
|
||||
@@ -186,11 +187,17 @@ class NZBGetClient(DownloadClient):
|
||||
],
|
||||
)
|
||||
|
||||
if nzb_id and nzb_id > 0:
|
||||
if isinstance(nzb_id, int) and nzb_id > 0:
|
||||
logger.info("Added NZB to NZBGet: %s", nzb_id)
|
||||
return str(nzb_id)
|
||||
|
||||
_raise_invalid_nzb_id()
|
||||
if isinstance(nzb_id, str):
|
||||
stripped_nzb_id = nzb_id.strip()
|
||||
if stripped_nzb_id.isdigit() and int(stripped_nzb_id) > 0:
|
||||
logger.info("Added NZB to NZBGet: %s", stripped_nzb_id)
|
||||
return stripped_nzb_id
|
||||
|
||||
_raise_runtime_error("NZBGet returned invalid ID")
|
||||
except requests.RequestException as e:
|
||||
logger.exception("Failed to fetch NZB from URL")
|
||||
msg = f"Failed to fetch NZB: {e}"
|
||||
|
||||
@@ -6,18 +6,23 @@ import time
|
||||
from http import HTTPStatus
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import NoReturn
|
||||
from typing import NoReturn, TypedDict
|
||||
|
||||
import requests
|
||||
|
||||
from shelfmark.core.config import config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.utils import normalize_http_url
|
||||
from shelfmark.download.clients import (
|
||||
DownloadClient,
|
||||
DownloadStatus,
|
||||
register_client,
|
||||
)
|
||||
from shelfmark.download.clients._coercion import (
|
||||
coerce_optional_float,
|
||||
coerce_optional_int,
|
||||
config_text,
|
||||
normalize_http_config_url,
|
||||
)
|
||||
from shelfmark.download.clients.torrent_utils import (
|
||||
extract_torrent_info,
|
||||
)
|
||||
@@ -41,6 +46,15 @@ _HTTP_STATUS_NOT_FOUND = HTTPStatus.NOT_FOUND
|
||||
_ONE_WEEK_IN_SECONDS = 604800
|
||||
|
||||
|
||||
class _QBittorrentAddKwargs(TypedDict, total=False):
|
||||
rename: str
|
||||
category: str
|
||||
save_path: str
|
||||
tags: str
|
||||
seeding_time_limit: int
|
||||
ratio_limit: float
|
||||
|
||||
|
||||
def _resolve_qbittorrent_exception_type(candidate: object) -> type[Exception]:
|
||||
if isinstance(candidate, type) and issubclass(candidate, Exception):
|
||||
return candidate
|
||||
@@ -197,21 +211,24 @@ class QBittorrentClient(DownloadClient):
|
||||
raise ValueError(msg)
|
||||
|
||||
# We use `_base_url` for direct HTTP calls, so it must be a fully-qualified URL.
|
||||
self._base_url = normalize_http_url(raw_url)
|
||||
self._base_url = normalize_http_config_url(raw_url, require_string=True)
|
||||
if not self._base_url:
|
||||
msg = "QBITTORRENT_URL is invalid"
|
||||
raise ValueError(msg)
|
||||
|
||||
username = config_text(config.get("QBITTORRENT_USERNAME", ""))
|
||||
password = config_text(config.get("QBITTORRENT_PASSWORD", ""))
|
||||
|
||||
# qbittorrent-api accepts either a full URL or host:port; prefer the normalized URL
|
||||
# for consistency.
|
||||
self._client = Client(
|
||||
host=self._base_url,
|
||||
username=config.get("QBITTORRENT_USERNAME", ""),
|
||||
password=config.get("QBITTORRENT_PASSWORD", ""),
|
||||
username=username,
|
||||
password=password,
|
||||
VERIFY_WEBUI_CERTIFICATE=get_ssl_verify(self._base_url),
|
||||
)
|
||||
self._category = config.get("QBITTORRENT_CATEGORY", "books")
|
||||
self._download_dir = config.get("QBITTORRENT_DOWNLOAD_DIR", "")
|
||||
self._category = config_text(config.get("QBITTORRENT_CATEGORY", "books"))
|
||||
self._download_dir = config_text(config.get("QBITTORRENT_DOWNLOAD_DIR", ""))
|
||||
self._tags = _normalize_tags(config.get("QBITTORRENT_TAG", []))
|
||||
|
||||
def _get_torrents_info(
|
||||
@@ -318,8 +335,8 @@ class QBittorrentClient(DownloadClient):
|
||||
@staticmethod
|
||||
def is_configured() -> bool:
|
||||
"""Check if qBittorrent is configured and selected as the torrent client."""
|
||||
client = config.get("PROWLARR_TORRENT_CLIENT", "")
|
||||
url = normalize_http_url(config.get("QBITTORRENT_URL", ""))
|
||||
client = config_text(config.get("PROWLARR_TORRENT_CLIENT", ""))
|
||||
url = normalize_http_config_url(config.get("QBITTORRENT_URL", ""), require_string=True)
|
||||
return client == "qbittorrent" and bool(url)
|
||||
|
||||
def test_connection(self) -> tuple[bool, str]:
|
||||
@@ -360,6 +377,8 @@ class QBittorrentClient(DownloadClient):
|
||||
# Use configured category if not explicitly provided
|
||||
category = category or self._category
|
||||
tags = self._tags
|
||||
seeding_time_limit: int | None = None
|
||||
ratio_limit: float | None = None
|
||||
|
||||
# Ensure category exists (may already exist, which is fine)
|
||||
if category:
|
||||
@@ -380,24 +399,23 @@ class QBittorrentClient(DownloadClient):
|
||||
expected_hash = torrent_info.info_hash
|
||||
torrent_data = torrent_info.torrent_data
|
||||
|
||||
# Add the torrent - use file content if we have it, otherwise URL
|
||||
add_kwargs = {
|
||||
"rename": name,
|
||||
}
|
||||
# Per-torrent seeding limits from indexer
|
||||
seeding_time_limit_value = kwargs.get("seeding_time_limit")
|
||||
seeding_time_limit = coerce_optional_int(seeding_time_limit_value)
|
||||
ratio_limit_value = kwargs.get("ratio_limit")
|
||||
ratio_limit = coerce_optional_float(ratio_limit_value)
|
||||
|
||||
add_kwargs: _QBittorrentAddKwargs = {"rename": name}
|
||||
if category:
|
||||
add_kwargs["category"] = category
|
||||
if self._download_dir:
|
||||
add_kwargs["save_path"] = self._download_dir
|
||||
if tags:
|
||||
add_kwargs["tags"] = ",".join(tags)
|
||||
|
||||
# Per-torrent seeding limits from indexer
|
||||
seeding_time_limit = kwargs.get("seeding_time_limit")
|
||||
if seeding_time_limit is not None:
|
||||
add_kwargs["seeding_time_limit"] = int(seeding_time_limit)
|
||||
ratio_limit = kwargs.get("ratio_limit")
|
||||
add_kwargs["seeding_time_limit"] = seeding_time_limit
|
||||
if ratio_limit is not None:
|
||||
add_kwargs["ratio_limit"] = float(ratio_limit)
|
||||
add_kwargs["ratio_limit"] = ratio_limit
|
||||
|
||||
if torrent_data:
|
||||
result = self._client.torrents_add(
|
||||
|
||||
@@ -5,17 +5,18 @@ Uses xmlrpc to communicate with rTorrent's RPC interface.
|
||||
|
||||
import ssl
|
||||
import xmlrpc.client as stdlib_xmlrpc_client
|
||||
from typing import NoReturn
|
||||
from typing import Any, NoReturn, Protocol, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from shelfmark.core.config import config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.utils import get_hardened_xmlrpc_client, normalize_http_url
|
||||
from shelfmark.core.utils import get_hardened_xmlrpc_client
|
||||
from shelfmark.download.clients import (
|
||||
DownloadClient,
|
||||
DownloadStatus,
|
||||
register_client,
|
||||
)
|
||||
from shelfmark.download.clients._coercion import config_text, normalize_http_config_url
|
||||
from shelfmark.download.clients.torrent_utils import (
|
||||
extract_torrent_info,
|
||||
)
|
||||
@@ -35,7 +36,38 @@ _RTORRENT_CLIENT_ERRORS = (
|
||||
)
|
||||
|
||||
|
||||
def _create_rtorrent_server_proxy(url: str) -> object:
|
||||
class _RTorrentSystemProtocol(Protocol):
|
||||
def client_version(self) -> object: ...
|
||||
|
||||
|
||||
class _RTorrentLoadProtocol(Protocol):
|
||||
def raw_start(self, target: str, torrent_data: bytes, commands: str) -> object: ...
|
||||
|
||||
def start(self, target: str, url: str, commands: str) -> object: ...
|
||||
|
||||
|
||||
class _RTorrentDownloadProtocol(Protocol):
|
||||
def multicall2(self, *args: object) -> list[list[Any]]: ...
|
||||
|
||||
def delete_tied(self, download_id: str) -> object: ...
|
||||
|
||||
def erase(self, download_id: str) -> object: ...
|
||||
|
||||
def stop(self, download_id: str) -> object: ...
|
||||
|
||||
|
||||
class _RTorrentDirectoryProtocol(Protocol):
|
||||
def default(self) -> str: ...
|
||||
|
||||
|
||||
class _RTorrentRpcProtocol(Protocol):
|
||||
system: _RTorrentSystemProtocol
|
||||
load: _RTorrentLoadProtocol
|
||||
d: _RTorrentDownloadProtocol
|
||||
directory: _RTorrentDirectoryProtocol
|
||||
|
||||
|
||||
def _create_rtorrent_server_proxy(url: str) -> _RTorrentRpcProtocol:
|
||||
"""Create an XML-RPC ServerProxy honoring certificate validation mode."""
|
||||
xmlrpc_client = get_hardened_xmlrpc_client()
|
||||
|
||||
@@ -45,9 +77,9 @@ def _create_rtorrent_server_proxy(url: str) -> object:
|
||||
ssl_context.check_hostname = False
|
||||
ssl_context.verify_mode = ssl.CERT_NONE
|
||||
transport = xmlrpc_client.SafeTransport(context=ssl_context)
|
||||
return xmlrpc_client.ServerProxy(url, transport=transport)
|
||||
return cast(_RTorrentRpcProtocol, xmlrpc_client.ServerProxy(url, transport=transport))
|
||||
|
||||
return xmlrpc_client.ServerProxy(url)
|
||||
return cast(_RTorrentRpcProtocol, xmlrpc_client.ServerProxy(url))
|
||||
|
||||
|
||||
def _raise_runtime_error(message: str) -> NoReturn:
|
||||
@@ -63,32 +95,32 @@ class RTorrentClient(DownloadClient):
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize rTorrent client with settings from config."""
|
||||
raw_url = config.get("RTORRENT_URL", "")
|
||||
raw_url = config_text(config.get("RTORRENT_URL", ""))
|
||||
if not raw_url:
|
||||
msg = "RTORRENT_URL is required"
|
||||
raise ValueError(msg)
|
||||
|
||||
self._base_url = normalize_http_url(raw_url)
|
||||
self._base_url = normalize_http_config_url(raw_url)
|
||||
if not self._base_url:
|
||||
msg = "RTORRENT_URL is invalid"
|
||||
raise ValueError(msg)
|
||||
|
||||
username = config.get("RTORRENT_USERNAME", "")
|
||||
password = config.get("RTORRENT_PASSWORD", "")
|
||||
username = config_text(config.get("RTORRENT_USERNAME", ""))
|
||||
password = config_text(config.get("RTORRENT_PASSWORD", ""))
|
||||
|
||||
if username and password:
|
||||
parsed = urlparse(self._base_url)
|
||||
self._base_url = f"{parsed.scheme}://{username}:{password}@{parsed.netloc}{parsed.path}"
|
||||
|
||||
self._rpc = _create_rtorrent_server_proxy(self._base_url)
|
||||
self._download_dir = config.get("RTORRENT_DOWNLOAD_DIR", "")
|
||||
self._label = config.get("RTORRENT_LABEL", "")
|
||||
self._download_dir = config_text(config.get("RTORRENT_DOWNLOAD_DIR", ""))
|
||||
self._label = config_text(config.get("RTORRENT_LABEL", ""))
|
||||
|
||||
@staticmethod
|
||||
def is_configured() -> bool:
|
||||
"""Check if rTorrent is configured and selected as the torrent client."""
|
||||
client = config.get("PROWLARR_TORRENT_CLIENT", "")
|
||||
url = normalize_http_url(config.get("RTORRENT_URL", ""))
|
||||
client = config_text(config.get("PROWLARR_TORRENT_CLIENT", ""))
|
||||
url = normalize_http_config_url(config.get("RTORRENT_URL", ""))
|
||||
return client == "rtorrent" and bool(url)
|
||||
|
||||
def test_connection(self) -> tuple[bool, str]:
|
||||
@@ -372,4 +404,4 @@ class RTorrentClient(DownloadClient):
|
||||
except _RTORRENT_CLIENT_ERRORS:
|
||||
return None
|
||||
else:
|
||||
return path or None
|
||||
return str(path) if path else None
|
||||
|
||||
@@ -3,19 +3,20 @@
|
||||
Uses SABnzbd's REST API directly via requests (no external dependency).
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import requests
|
||||
|
||||
from shelfmark.core.config import config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.utils import normalize_http_url
|
||||
from shelfmark.download.clients import (
|
||||
DownloadClient,
|
||||
DownloadStatus,
|
||||
register_client,
|
||||
with_retry,
|
||||
)
|
||||
from shelfmark.download.clients._coercion import config_text, normalize_http_config_url
|
||||
from shelfmark.download.network import get_ssl_verify
|
||||
|
||||
logger = setup_logger(__name__)
|
||||
@@ -29,6 +30,7 @@ _SABNZBD_CLIENT_ERRORS = (
|
||||
TypeError,
|
||||
ValueError,
|
||||
)
|
||||
_SabnzbdRequestParam = str | int | float | bool
|
||||
|
||||
|
||||
def _parse_eta(eta_str: str) -> int | None:
|
||||
@@ -110,33 +112,33 @@ class SABnzbdClient(DownloadClient):
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize SABnzbd client with settings from config."""
|
||||
raw_url = config.get("SABNZBD_URL", "")
|
||||
raw_url = config_text(config.get("SABNZBD_URL", ""))
|
||||
if not raw_url:
|
||||
msg = "SABNZBD_URL is required"
|
||||
raise ValueError(msg)
|
||||
|
||||
api_key = config.get("SABNZBD_API_KEY", "")
|
||||
api_key = config_text(config.get("SABNZBD_API_KEY", ""))
|
||||
if not api_key:
|
||||
msg = "SABNZBD_API_KEY is required"
|
||||
raise ValueError(msg)
|
||||
|
||||
self.url = normalize_http_url(raw_url)
|
||||
self.url = normalize_http_config_url(raw_url)
|
||||
if not self.url:
|
||||
msg = "SABNZBD_URL is invalid"
|
||||
raise ValueError(msg)
|
||||
self.api_key = api_key
|
||||
self._category = config.get("SABNZBD_CATEGORY", "books")
|
||||
self._category = config_text(config.get("SABNZBD_CATEGORY", "books"))
|
||||
|
||||
@staticmethod
|
||||
def is_configured() -> bool:
|
||||
"""Check if SABnzbd is configured and selected as the usenet client."""
|
||||
client = config.get("PROWLARR_USENET_CLIENT", "")
|
||||
url = normalize_http_url(config.get("SABNZBD_URL", ""))
|
||||
api_key = config.get("SABNZBD_API_KEY", "")
|
||||
client = config_text(config.get("PROWLARR_USENET_CLIENT", ""))
|
||||
url = normalize_http_config_url(config.get("SABNZBD_URL", ""))
|
||||
api_key = config_text(config.get("SABNZBD_API_KEY", ""))
|
||||
return client == "sabnzbd" and bool(url) and bool(api_key)
|
||||
|
||||
@with_retry()
|
||||
def _api_call(self, mode: str, params: dict | None = None) -> object:
|
||||
def _api_call(self, mode: str, params: dict[str, _SabnzbdRequestParam] | None = None) -> Any:
|
||||
"""Make an API call to SABnzbd.
|
||||
|
||||
Args:
|
||||
@@ -152,7 +154,7 @@ class SABnzbdClient(DownloadClient):
|
||||
"""
|
||||
api_url = f"{self.url}/api"
|
||||
|
||||
request_params = {
|
||||
request_params: dict[str, _SabnzbdRequestParam] = {
|
||||
"apikey": self.api_key,
|
||||
"mode": mode,
|
||||
"output": "json",
|
||||
@@ -177,7 +179,7 @@ class SABnzbdClient(DownloadClient):
|
||||
|
||||
def _api_post_file(
|
||||
self, nzb_content: bytes, filename: str, nzb_name: str, category: str
|
||||
) -> object:
|
||||
) -> Any:
|
||||
"""Upload an NZB file to SABnzbd using addfile.
|
||||
|
||||
Returns:
|
||||
@@ -185,7 +187,7 @@ class SABnzbdClient(DownloadClient):
|
||||
|
||||
"""
|
||||
api_url = f"{self.url}/api"
|
||||
request_params = {
|
||||
request_params: dict[str, _SabnzbdRequestParam] = {
|
||||
"apikey": self.api_key,
|
||||
"mode": "addfile",
|
||||
"output": "json",
|
||||
@@ -224,7 +226,7 @@ class SABnzbdClient(DownloadClient):
|
||||
if not api_key:
|
||||
return {}
|
||||
|
||||
prowlarr_url = normalize_http_url(config.get("PROWLARR_URL", ""))
|
||||
prowlarr_url = normalize_http_config_url(config.get("PROWLARR_URL", ""))
|
||||
if not prowlarr_url:
|
||||
return {}
|
||||
|
||||
@@ -320,13 +322,13 @@ class SABnzbdClient(DownloadClient):
|
||||
|
||||
"""
|
||||
# Use configured category if not explicitly provided
|
||||
category = category or self._category
|
||||
resolved_category = category or self._category
|
||||
|
||||
try:
|
||||
logger.debug("Adding NZB to SABnzbd: %s", name)
|
||||
nzb_filename = self._build_nzb_filename(name, url)
|
||||
nzb_content = self._fetch_nzb_content(url)
|
||||
result = self._api_post_file(nzb_content, nzb_filename, name, category)
|
||||
result = self._api_post_file(nzb_content, nzb_filename, name, resolved_category)
|
||||
nzo_id = self._extract_nzo_id(result)
|
||||
logger.info("Added NZB to SABnzbd: %s", nzo_id)
|
||||
except _SABNZBD_CLIENT_ERRORS as e:
|
||||
@@ -340,7 +342,7 @@ class SABnzbdClient(DownloadClient):
|
||||
{
|
||||
"name": url,
|
||||
"nzbname": name,
|
||||
"cat": category,
|
||||
"cat": resolved_category,
|
||||
},
|
||||
)
|
||||
nzo_id = self._extract_nzo_id(result)
|
||||
|
||||
@@ -1,7 +1,10 @@
|
||||
"""Shared download client settings registration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
from contextlib import contextmanager, suppress
|
||||
from typing import TYPE_CHECKING, Any, NoReturn
|
||||
from typing import TYPE_CHECKING, Any, NoReturn, Protocol, TypeGuard
|
||||
|
||||
from shelfmark.core.settings_registry import (
|
||||
ActionButton,
|
||||
@@ -31,12 +34,24 @@ except ImportError:
|
||||
_ImportedTransmissionError = RuntimeError
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Iterator
|
||||
from collections.abc import Callable, Iterator
|
||||
|
||||
# ==================== Test Connection Callbacks ====================
|
||||
_DELUGE_HOST_ENTRY_MIN_LENGTH = 2
|
||||
|
||||
|
||||
class _SessionWithVerify(Protocol):
|
||||
verify: bool
|
||||
|
||||
|
||||
class _RequestsModuleWithSession(Protocol):
|
||||
Session: Callable[..., _SessionWithVerify]
|
||||
|
||||
|
||||
class _TransmissionClientWithProtocol(Protocol):
|
||||
protocol: str
|
||||
|
||||
|
||||
def _resolve_exception_type(candidate: object) -> type[Exception]:
|
||||
if isinstance(candidate, type) and issubclass(candidate, Exception):
|
||||
return candidate
|
||||
@@ -69,6 +84,39 @@ def _raise_runtime_error(message: str) -> NoReturn:
|
||||
raise RuntimeError(message)
|
||||
|
||||
|
||||
def _is_requests_module_with_session(candidate: object) -> TypeGuard[_RequestsModuleWithSession]:
|
||||
return callable(getattr(candidate, "Session", None))
|
||||
|
||||
|
||||
def _has_protocol_attr(candidate: object) -> TypeGuard[_TransmissionClientWithProtocol]:
|
||||
return hasattr(candidate, "protocol")
|
||||
|
||||
|
||||
def _set_transmission_protocol_if_supported(client: object, protocol: str) -> None:
|
||||
if protocol != "https" or not _has_protocol_attr(client):
|
||||
return
|
||||
with suppress(AttributeError, OSError, RuntimeError, TypeError, ValueError):
|
||||
client.protocol = protocol
|
||||
|
||||
|
||||
def _resolve_string_setting(
|
||||
current_values: dict[str, Any],
|
||||
config_get: Callable[[str, str], object],
|
||||
key: str,
|
||||
*,
|
||||
default: str = "",
|
||||
) -> str:
|
||||
current_value = current_values.get(key)
|
||||
if isinstance(current_value, str) and current_value:
|
||||
return current_value
|
||||
|
||||
config_value = config_get(key, default)
|
||||
if isinstance(config_value, str) and config_value:
|
||||
return config_value
|
||||
|
||||
return default
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _transmission_session_verify_override(url: str) -> Iterator[None]:
|
||||
"""Ensure transmission-rpc constructor uses the configured TLS verify mode."""
|
||||
@@ -78,23 +126,28 @@ def _transmission_session_verify_override(url: str) -> Iterator[None]:
|
||||
return
|
||||
|
||||
try:
|
||||
import transmission_rpc.client as transmission_rpc_client
|
||||
transmission_rpc_client = importlib.import_module("transmission_rpc.client")
|
||||
except ImportError:
|
||||
yield
|
||||
return
|
||||
|
||||
original_session_factory = transmission_rpc_client.requests.Session
|
||||
requests_module = getattr(transmission_rpc_client, "requests", None)
|
||||
if not _is_requests_module_with_session(requests_module):
|
||||
yield
|
||||
return
|
||||
|
||||
original_session_factory = requests_module.Session
|
||||
|
||||
def _session_factory(*args: Any, **kwargs: Any) -> Any:
|
||||
session = original_session_factory(*args, **kwargs)
|
||||
session.verify = False
|
||||
return session
|
||||
|
||||
transmission_rpc_client.requests.Session = _session_factory
|
||||
requests_module.Session = _session_factory
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
transmission_rpc_client.requests.Session = original_session_factory
|
||||
requests_module.Session = original_session_factory
|
||||
|
||||
|
||||
def _test_qbittorrent_connection(current_values: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
@@ -103,9 +156,9 @@ def _test_qbittorrent_connection(current_values: dict[str, Any] | None = None) -
|
||||
|
||||
current_values = current_values or {}
|
||||
|
||||
raw_url = current_values.get("QBITTORRENT_URL") or config.get("QBITTORRENT_URL", "")
|
||||
username = current_values.get("QBITTORRENT_USERNAME") or config.get("QBITTORRENT_USERNAME", "")
|
||||
password = current_values.get("QBITTORRENT_PASSWORD") or config.get("QBITTORRENT_PASSWORD", "")
|
||||
raw_url = _resolve_string_setting(current_values, config.get, "QBITTORRENT_URL")
|
||||
username = _resolve_string_setting(current_values, config.get, "QBITTORRENT_USERNAME")
|
||||
password = _resolve_string_setting(current_values, config.get, "QBITTORRENT_PASSWORD")
|
||||
|
||||
if not raw_url:
|
||||
return {"success": False, "message": "qBittorrent URL is required"}
|
||||
@@ -142,13 +195,9 @@ def _test_transmission_connection(current_values: dict[str, Any] | None = None)
|
||||
|
||||
current_values = current_values or {}
|
||||
|
||||
raw_url = current_values.get("TRANSMISSION_URL") or config.get("TRANSMISSION_URL", "")
|
||||
username = current_values.get("TRANSMISSION_USERNAME") or config.get(
|
||||
"TRANSMISSION_USERNAME", ""
|
||||
)
|
||||
password = current_values.get("TRANSMISSION_PASSWORD") or config.get(
|
||||
"TRANSMISSION_PASSWORD", ""
|
||||
)
|
||||
raw_url = _resolve_string_setting(current_values, config.get, "TRANSMISSION_URL")
|
||||
username = _resolve_string_setting(current_values, config.get, "TRANSMISSION_USERNAME")
|
||||
password = _resolve_string_setting(current_values, config.get, "TRANSMISSION_PASSWORD")
|
||||
|
||||
if not raw_url:
|
||||
return {"success": False, "message": "Transmission URL is required"}
|
||||
@@ -180,9 +229,7 @@ def _test_transmission_connection(current_values: dict[str, Any] | None = None)
|
||||
client_kwargs.pop("protocol", None)
|
||||
with _transmission_session_verify_override(url):
|
||||
client = Client(**client_kwargs)
|
||||
if protocol == "https" and hasattr(client, "protocol"):
|
||||
with suppress(Exception):
|
||||
client.protocol = protocol
|
||||
_set_transmission_protocol_if_supported(client, protocol)
|
||||
|
||||
# Keep session verify aligned for subsequent calls beyond constructor bootstrap.
|
||||
http_session = getattr(client, "_http_session", None)
|
||||
@@ -209,9 +256,11 @@ def _test_deluge_connection(current_values: dict[str, Any] | None = None) -> dic
|
||||
|
||||
current_values = current_values or {}
|
||||
|
||||
raw_host = current_values.get("DELUGE_HOST") or config.get("DELUGE_HOST", "localhost")
|
||||
raw_port = current_values.get("DELUGE_PORT") or config.get("DELUGE_PORT", "8112")
|
||||
password = current_values.get("DELUGE_PASSWORD") or config.get("DELUGE_PASSWORD", "")
|
||||
raw_host = _resolve_string_setting(
|
||||
current_values, config.get, "DELUGE_HOST", default="localhost"
|
||||
)
|
||||
raw_port = _resolve_string_setting(current_values, config.get, "DELUGE_PORT", default="8112")
|
||||
password = _resolve_string_setting(current_values, config.get, "DELUGE_PASSWORD")
|
||||
|
||||
if not raw_host:
|
||||
return {"success": False, "message": "Deluge host is required"}
|
||||
@@ -330,9 +379,9 @@ def _test_rtorrent_connection(current_values: dict[str, Any] | None = None) -> d
|
||||
|
||||
current_values = current_values or {}
|
||||
|
||||
raw_url = current_values.get("RTORRENT_URL") or config.get("RTORRENT_URL", "")
|
||||
username = current_values.get("RTORRENT_USERNAME") or config.get("RTORRENT_USERNAME", "")
|
||||
password = current_values.get("RTORRENT_PASSWORD") or config.get("RTORRENT_PASSWORD", "")
|
||||
raw_url = _resolve_string_setting(current_values, config.get, "RTORRENT_URL")
|
||||
username = _resolve_string_setting(current_values, config.get, "RTORRENT_USERNAME")
|
||||
password = _resolve_string_setting(current_values, config.get, "RTORRENT_PASSWORD")
|
||||
|
||||
if not raw_url:
|
||||
return {"success": False, "message": "rTorrent URL is required"}
|
||||
@@ -343,7 +392,10 @@ def _test_rtorrent_connection(current_values: dict[str, Any] | None = None) -> d
|
||||
|
||||
try:
|
||||
xmlrpc_client = get_hardened_xmlrpc_client()
|
||||
except (RuntimeError, OSError, ValueError, TypeError) as e:
|
||||
return {"success": False, "message": f"Connection failed: {e!s}"}
|
||||
|
||||
try:
|
||||
# Add HTTP auth to URL if credentials provided
|
||||
if username and password:
|
||||
parsed = urlparse(url)
|
||||
@@ -378,9 +430,14 @@ def _test_nzbget_connection(current_values: dict[str, Any] | None = None) -> dic
|
||||
|
||||
current_values = current_values or {}
|
||||
|
||||
raw_url = current_values.get("NZBGET_URL") or config.get("NZBGET_URL", "")
|
||||
username = current_values.get("NZBGET_USERNAME") or config.get("NZBGET_USERNAME", "nzbget")
|
||||
password = current_values.get("NZBGET_PASSWORD") or config.get("NZBGET_PASSWORD", "")
|
||||
raw_url = _resolve_string_setting(current_values, config.get, "NZBGET_URL")
|
||||
username = _resolve_string_setting(
|
||||
current_values,
|
||||
config.get,
|
||||
"NZBGET_USERNAME",
|
||||
default="nzbget",
|
||||
)
|
||||
password = _resolve_string_setting(current_values, config.get, "NZBGET_PASSWORD")
|
||||
|
||||
if not raw_url:
|
||||
return {"success": False, "message": "NZBGet URL is required"}
|
||||
@@ -428,8 +485,8 @@ def _test_sabnzbd_connection(current_values: dict[str, Any] | None = None) -> di
|
||||
|
||||
current_values = current_values or {}
|
||||
|
||||
raw_url = current_values.get("SABNZBD_URL") or config.get("SABNZBD_URL", "")
|
||||
api_key = current_values.get("SABNZBD_API_KEY") or config.get("SABNZBD_API_KEY", "")
|
||||
raw_url = _resolve_string_setting(current_values, config.get, "SABNZBD_URL")
|
||||
api_key = _resolve_string_setting(current_values, config.get, "SABNZBD_API_KEY")
|
||||
|
||||
if not raw_url:
|
||||
return {"success": False, "message": "SABnzbd URL is required"}
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Shared utilities for torrent clients."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import re
|
||||
@@ -31,6 +33,8 @@ _TORRENT_FETCH_ERRORS = (
|
||||
)
|
||||
_TORRENT_PARSE_ERRORS = (IndexError, KeyError, TypeError, ValueError)
|
||||
|
||||
type BencodeValue = dict[str | bytes, BencodeValue] | list[BencodeValue] | int | bytes | str
|
||||
|
||||
|
||||
@dataclass
|
||||
class TorrentInfo:
|
||||
@@ -222,7 +226,7 @@ def bencode_decode(data: bytes) -> tuple:
|
||||
raise ValueError(msg)
|
||||
|
||||
|
||||
def bencode_encode(data: dict[str | bytes, object] | list[object] | int | bytes | str) -> bytes:
|
||||
def bencode_encode(data: BencodeValue) -> bytes:
|
||||
"""Encode data to bencode format."""
|
||||
if isinstance(data, dict):
|
||||
# Keys must be sorted (bencode spec requirement)
|
||||
|
||||
@@ -3,17 +3,25 @@
|
||||
Uses the transmission-rpc library to communicate with Transmission's RPC API.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
from contextlib import contextmanager, suppress
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Protocol, TypeGuard
|
||||
|
||||
from shelfmark.core.config import config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.utils import normalize_http_url
|
||||
from shelfmark.download.clients import (
|
||||
DownloadClient,
|
||||
DownloadStatus,
|
||||
register_client,
|
||||
)
|
||||
from shelfmark.download.clients._coercion import (
|
||||
coerce_optional_float,
|
||||
coerce_optional_int,
|
||||
config_text,
|
||||
normalize_http_config_url,
|
||||
)
|
||||
from shelfmark.download.clients.torrent_utils import (
|
||||
extract_torrent_info,
|
||||
parse_transmission_url,
|
||||
@@ -48,6 +56,39 @@ _TRANSMISSION_CLIENT_ERRORS = (
|
||||
)
|
||||
|
||||
|
||||
class _TransmissionSessionProtocol(Protocol):
|
||||
verify: bool
|
||||
|
||||
|
||||
class _TransmissionSessionFactory(Protocol):
|
||||
def __call__(self, *args: object, **kwargs: object) -> _TransmissionSessionProtocol: ...
|
||||
|
||||
|
||||
class _TransmissionRequestsNamespace(Protocol):
|
||||
Session: _TransmissionSessionFactory
|
||||
|
||||
|
||||
class _TransmissionProtocolAttribute(Protocol):
|
||||
protocol: str
|
||||
|
||||
|
||||
def _is_requests_namespace_with_session(
|
||||
candidate: object,
|
||||
) -> TypeGuard[_TransmissionRequestsNamespace]:
|
||||
return hasattr(candidate, "Session") and callable(getattr(candidate, "Session", None))
|
||||
|
||||
|
||||
def _has_protocol_attr(candidate: object) -> TypeGuard[_TransmissionProtocolAttribute]:
|
||||
return hasattr(candidate, "protocol")
|
||||
|
||||
|
||||
def _set_transmission_protocol_if_supported(client: object, protocol: str) -> None:
|
||||
if protocol != "https" or not _has_protocol_attr(client):
|
||||
return
|
||||
with suppress(AttributeError, OSError, RuntimeError, TypeError, ValueError):
|
||||
client.protocol = protocol
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _transmission_session_verify_override(url: str) -> Iterator[None]:
|
||||
"""Temporarily override transmission-rpc's session factory when verify is disabled.
|
||||
@@ -61,24 +102,29 @@ def _transmission_session_verify_override(url: str) -> Iterator[None]:
|
||||
return
|
||||
|
||||
try:
|
||||
import transmission_rpc.client as transmission_rpc_client
|
||||
|
||||
original_session_factory = transmission_rpc_client.requests.Session
|
||||
except AttributeError, ImportError:
|
||||
transmission_rpc_client = importlib.import_module("transmission_rpc.client")
|
||||
requests_namespace = getattr(transmission_rpc_client, "requests", None)
|
||||
except ImportError:
|
||||
# If internals differ, gracefully fall back to default behavior.
|
||||
yield
|
||||
return
|
||||
|
||||
def _session_factory(*args: object, **kwargs: object) -> object:
|
||||
if not _is_requests_namespace_with_session(requests_namespace):
|
||||
yield
|
||||
return
|
||||
|
||||
original_session_factory = requests_namespace.Session
|
||||
|
||||
def _session_factory(*args: object, **kwargs: object) -> _TransmissionSessionProtocol:
|
||||
session = original_session_factory(*args, **kwargs)
|
||||
session.verify = False
|
||||
return session
|
||||
|
||||
transmission_rpc_client.requests.Session = _session_factory
|
||||
requests_namespace.Session = _session_factory
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
transmission_rpc_client.requests.Session = original_session_factory
|
||||
requests_namespace.Session = original_session_factory
|
||||
|
||||
|
||||
def _apply_transmission_ssl_verify(client: object, url: str) -> None:
|
||||
@@ -103,18 +149,18 @@ class TransmissionClient(DownloadClient):
|
||||
"""Initialize Transmission client with settings from config."""
|
||||
from transmission_rpc import Client
|
||||
|
||||
raw_url = config.get("TRANSMISSION_URL", "")
|
||||
raw_url = config_text(config.get("TRANSMISSION_URL", ""))
|
||||
if not raw_url:
|
||||
msg = "TRANSMISSION_URL is required"
|
||||
raise ValueError(msg)
|
||||
|
||||
url = normalize_http_url(raw_url)
|
||||
url = normalize_http_config_url(raw_url)
|
||||
if not url:
|
||||
msg = "TRANSMISSION_URL is invalid"
|
||||
raise ValueError(msg)
|
||||
|
||||
username = config.get("TRANSMISSION_USERNAME", "")
|
||||
password = config.get("TRANSMISSION_PASSWORD", "")
|
||||
username = config_text(config.get("TRANSMISSION_USERNAME", ""))
|
||||
password = config_text(config.get("TRANSMISSION_PASSWORD", ""))
|
||||
|
||||
# Parse URL to extract host, port, and path
|
||||
protocol, host, port, path = parse_transmission_url(url)
|
||||
@@ -138,18 +184,16 @@ class TransmissionClient(DownloadClient):
|
||||
with _transmission_session_verify_override(url):
|
||||
self._client = Client(**client_kwargs)
|
||||
# Some versions expose protocol as an attribute rather than kwarg.
|
||||
if protocol == "https" and hasattr(self._client, "protocol"):
|
||||
with suppress(Exception):
|
||||
self._client.protocol = protocol
|
||||
_set_transmission_protocol_if_supported(self._client, protocol)
|
||||
_apply_transmission_ssl_verify(self._client, url)
|
||||
self._category = config.get("TRANSMISSION_CATEGORY", "books")
|
||||
self._download_dir = config.get("TRANSMISSION_DOWNLOAD_DIR", "")
|
||||
self._category = config_text(config.get("TRANSMISSION_CATEGORY", "books"))
|
||||
self._download_dir = config_text(config.get("TRANSMISSION_DOWNLOAD_DIR", ""))
|
||||
|
||||
@staticmethod
|
||||
def is_configured() -> bool:
|
||||
"""Check if Transmission is configured and selected as the torrent client."""
|
||||
client = config.get("PROWLARR_TORRENT_CLIENT", "")
|
||||
url = normalize_http_url(config.get("TRANSMISSION_URL", ""))
|
||||
client = config_text(config.get("PROWLARR_TORRENT_CLIENT", ""))
|
||||
url = normalize_http_config_url(config.get("TRANSMISSION_URL", ""))
|
||||
return client == "transmission" and bool(url)
|
||||
|
||||
def test_connection(self) -> tuple[bool, str]:
|
||||
@@ -187,7 +231,7 @@ class TransmissionClient(DownloadClient):
|
||||
|
||||
"""
|
||||
try:
|
||||
resolved_category = category or self._category or ""
|
||||
resolved_category = category or self._category
|
||||
|
||||
torrent_info = extract_torrent_info(url, expected_hash=expected_hash)
|
||||
add_kwargs = {}
|
||||
@@ -215,13 +259,13 @@ class TransmissionClient(DownloadClient):
|
||||
|
||||
# Apply per-torrent seeding limits from indexer
|
||||
seed_kwargs = {}
|
||||
seeding_time_limit = kwargs.get("seeding_time_limit")
|
||||
seeding_time_limit = coerce_optional_int(kwargs.get("seeding_time_limit"))
|
||||
if seeding_time_limit is not None:
|
||||
seed_kwargs["seed_idle_limit"] = int(seeding_time_limit)
|
||||
seed_kwargs["seed_idle_limit"] = seeding_time_limit
|
||||
seed_kwargs["seed_idle_mode"] = 1 # per-torrent
|
||||
ratio_limit = kwargs.get("ratio_limit")
|
||||
ratio_limit = coerce_optional_float(kwargs.get("ratio_limit"))
|
||||
if ratio_limit is not None:
|
||||
seed_kwargs["seed_ratio_limit"] = float(ratio_limit)
|
||||
seed_kwargs["seed_ratio_limit"] = ratio_limit
|
||||
seed_kwargs["seed_ratio_mode"] = 1 # per-torrent
|
||||
if seed_kwargs:
|
||||
try:
|
||||
|
||||
@@ -44,7 +44,11 @@ def _get_io_threadpool() -> ThreadPool:
|
||||
global _IO_THREADPOOL
|
||||
if _IO_THREADPOOL is None:
|
||||
pool_size = max(2, min(8, os.cpu_count() or 2))
|
||||
_IO_THREADPOOL = _GeventThreadPool(pool_size)
|
||||
threadpool_cls = _GeventThreadPool
|
||||
if threadpool_cls is None:
|
||||
msg = "gevent threadpool is unavailable"
|
||||
raise RuntimeError(msg)
|
||||
_IO_THREADPOOL = threadpool_cls(pool_size)
|
||||
return _IO_THREADPOOL
|
||||
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ from tqdm import tqdm
|
||||
from shelfmark.bypass import BypassCancelledError
|
||||
from shelfmark.core.config import config as app_config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.request_helpers import coerce_bool, normalize_positive_int
|
||||
from shelfmark.download import network
|
||||
from shelfmark.download.network import get_proxies, get_ssl_verify
|
||||
|
||||
@@ -90,12 +91,12 @@ def _get_external_bypasser() -> ModuleType:
|
||||
|
||||
def _is_using_external_bypasser() -> bool:
|
||||
"""Check if external bypasser is configured (reads from config, not just env)."""
|
||||
return app_config.get("USING_EXTERNAL_BYPASSER", False)
|
||||
return coerce_bool(app_config.get("USING_EXTERNAL_BYPASSER", False))
|
||||
|
||||
|
||||
def _is_cf_bypass_enabled() -> bool:
|
||||
"""Check if Cloudflare bypass is enabled."""
|
||||
return app_config.get("USE_CF_BYPASS", True)
|
||||
return coerce_bool(app_config.get("USE_CF_BYPASS", True))
|
||||
|
||||
|
||||
def get_bypassed_page(
|
||||
@@ -251,18 +252,22 @@ def html_get_page(
|
||||
return html, response_url
|
||||
return html
|
||||
|
||||
retry = retry if retry is not None else app_config.MAX_RETRY
|
||||
configured_retry = normalize_positive_int(app_config.MAX_RETRY)
|
||||
retry_limit = (
|
||||
retry if retry is not None else (configured_retry if configured_retry is not None else 1)
|
||||
)
|
||||
selector = selector or network.AAMirrorSelector()
|
||||
original_url = url
|
||||
current_url = selector.rewrite(original_url)
|
||||
use_bypasser_now = use_bypasser
|
||||
|
||||
for attempt in range(1, retry + 1):
|
||||
for attempt in range(1, retry_limit + 1):
|
||||
# Check for cancellation before each attempt
|
||||
if cancel_flag and cancel_flag.is_set():
|
||||
logger.info("html_get_page cancelled before attempt %s", attempt)
|
||||
return _result("", current_url)
|
||||
|
||||
cookies: dict[str, str] = {}
|
||||
try:
|
||||
if use_bypasser_now and _is_cf_bypass_enabled():
|
||||
if status_callback:
|
||||
@@ -422,18 +427,18 @@ def html_get_page(
|
||||
continue
|
||||
|
||||
# Retry with backoff
|
||||
if attempt < retry:
|
||||
if attempt < retry_limit:
|
||||
logger.warning(
|
||||
"Retry %s/%s for %s: %s: %s",
|
||||
attempt,
|
||||
retry,
|
||||
retry_limit,
|
||||
current_url,
|
||||
type(e).__name__,
|
||||
e,
|
||||
)
|
||||
time.sleep(_backoff_delay(attempt))
|
||||
else:
|
||||
logger.exception("Giving up after %s attempts: %s", retry, current_url)
|
||||
logger.exception("Giving up after %s attempts: %s", retry_limit, current_url)
|
||||
|
||||
return _result("", current_url)
|
||||
|
||||
|
||||
@@ -16,6 +16,7 @@ from dns.exception import DNSException
|
||||
|
||||
from shelfmark.core.config import config as app_config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.request_helpers import coerce_bool, normalize_optional_text
|
||||
from shelfmark.core.utils import normalize_http_url
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -24,7 +25,7 @@ if TYPE_CHECKING:
|
||||
|
||||
def _get_no_proxy_patterns() -> list[str]:
|
||||
"""Get list of NO_PROXY patterns from config."""
|
||||
no_proxy = app_config.get("NO_PROXY", "")
|
||||
no_proxy = normalize_optional_text(app_config.get("NO_PROXY", ""))
|
||||
if not no_proxy:
|
||||
return []
|
||||
return [p.strip().lower() for p in no_proxy.split(",") if p.strip()]
|
||||
@@ -245,6 +246,32 @@ def _save_state(aa_url: str | None = None, dns_provider: str | None = None) -> N
|
||||
state["chosen_at"] = datetime.now(UTC).isoformat()
|
||||
|
||||
|
||||
def _set_runtime_dns_state(servers: list[str], doh_server: str) -> None:
|
||||
"""Update the module DNS state and mirrored config attributes.
|
||||
|
||||
The config singleton's `get()` values still represent persisted/configured
|
||||
settings. These attribute writes are only for runtime consumers that read
|
||||
the live resolver state via attribute access.
|
||||
"""
|
||||
global CUSTOM_DNS, DOH_SERVER
|
||||
|
||||
CUSTOM_DNS = list(servers)
|
||||
DOH_SERVER = doh_server
|
||||
runtime_config = cast(Any, app_config)
|
||||
runtime_config.CUSTOM_DNS = CUSTOM_DNS
|
||||
runtime_config.DOH_SERVER = DOH_SERVER
|
||||
|
||||
|
||||
def _get_configured_aa_url() -> str:
|
||||
"""Return the configured AA base URL normalized for runtime use."""
|
||||
configured_url = normalize_http_url(
|
||||
normalize_optional_text(app_config.get("AA_BASE_URL", "auto")),
|
||||
default_scheme="https",
|
||||
allow_special=("auto",),
|
||||
)
|
||||
return configured_url or "auto"
|
||||
|
||||
|
||||
# AA URL failover state
|
||||
_current_aa_url_index = 0
|
||||
_aa_urls: list[str] = [] # Initialized lazily in _initialize_aa_state()
|
||||
@@ -856,10 +883,7 @@ def switch_dns_provider() -> bool:
|
||||
|
||||
_current_dns_index += 1
|
||||
name, servers, doh = DNS_PROVIDERS[_current_dns_index]
|
||||
CUSTOM_DNS = servers
|
||||
DOH_SERVER = doh
|
||||
app_config.CUSTOM_DNS = servers
|
||||
app_config.DOH_SERVER = doh
|
||||
_set_runtime_dns_state(servers, doh)
|
||||
|
||||
logger.warning("Switched DNS provider to: %s (using DoH)", name)
|
||||
_save_state(dns_provider=name)
|
||||
@@ -895,13 +919,7 @@ def rotate_dns_and_reset_aa() -> bool:
|
||||
return False
|
||||
# Reset AA URL to first available auto option if using auto AA
|
||||
global _aa_base_url, _current_aa_url_index
|
||||
configured_url = normalize_http_url(
|
||||
app_config.get("AA_BASE_URL", "auto"),
|
||||
default_scheme="https",
|
||||
allow_special=("auto",),
|
||||
)
|
||||
if not configured_url:
|
||||
configured_url = "auto"
|
||||
configured_url = _get_configured_aa_url()
|
||||
|
||||
if configured_url == "auto":
|
||||
# Auto mode always resets to the first mirror to restart the cascade
|
||||
@@ -939,17 +957,18 @@ def set_dns_provider(
|
||||
provider = provider.lower().strip()
|
||||
|
||||
# Determine DoH preference - use provided value or fall back to config setting
|
||||
doh_enabled = use_doh if use_doh is not None else app_config.get("USE_DOH", True)
|
||||
doh_enabled = (
|
||||
use_doh
|
||||
if use_doh is not None
|
||||
else coerce_bool(app_config.get("USE_DOH", True), default=True)
|
||||
)
|
||||
|
||||
with _dns_switch_lock:
|
||||
if provider == "system":
|
||||
# Use system DNS only - no custom resolver, no failover rotation
|
||||
_current_dns_index = -1
|
||||
_dns_exhausted_logged = False
|
||||
CUSTOM_DNS = []
|
||||
DOH_SERVER = ""
|
||||
app_config.CUSTOM_DNS = []
|
||||
app_config.DOH_SERVER = ""
|
||||
_set_runtime_dns_state([], "")
|
||||
# Restore original system getaddrinfo
|
||||
socket.getaddrinfo = original_getaddrinfo
|
||||
logger.info("DNS set to system mode (using OS default resolver)")
|
||||
@@ -961,10 +980,7 @@ def set_dns_provider(
|
||||
# Note: Auto mode always uses DoH when rotating for reliability
|
||||
_current_dns_index = -1
|
||||
_dns_exhausted_logged = False
|
||||
CUSTOM_DNS = []
|
||||
DOH_SERVER = ""
|
||||
app_config.CUSTOM_DNS = []
|
||||
app_config.DOH_SERVER = ""
|
||||
_set_runtime_dns_state([], "")
|
||||
logger.info("DNS set to auto mode (system DNS, will rotate on failure with DoH)")
|
||||
init_dns_resolvers()
|
||||
_notify_dns_rotation("auto", [], "")
|
||||
@@ -975,10 +991,7 @@ def set_dns_provider(
|
||||
logger.warning("Manual DNS requested but no servers provided")
|
||||
return False
|
||||
_current_dns_index = -1 # Not using preset providers
|
||||
CUSTOM_DNS = manual_servers
|
||||
DOH_SERVER = "" # No DoH for manual servers
|
||||
app_config.CUSTOM_DNS = manual_servers
|
||||
app_config.DOH_SERVER = ""
|
||||
_set_runtime_dns_state(manual_servers, "")
|
||||
logger.info("DNS set to manual servers: %s", manual_servers)
|
||||
init_dns_resolvers()
|
||||
_notify_dns_rotation("manual", manual_servers, "")
|
||||
@@ -989,11 +1002,9 @@ def set_dns_provider(
|
||||
if name == provider:
|
||||
_current_dns_index = i
|
||||
_dns_exhausted_logged = False
|
||||
CUSTOM_DNS = servers
|
||||
# Only set DoH server if DoH is enabled
|
||||
DOH_SERVER = doh if doh_enabled else ""
|
||||
app_config.CUSTOM_DNS = servers
|
||||
app_config.DOH_SERVER = DOH_SERVER
|
||||
runtime_doh_server = doh if doh_enabled else ""
|
||||
_set_runtime_dns_state(servers, runtime_doh_server)
|
||||
doh_status = "DoH enabled" if doh_enabled else "standard DNS"
|
||||
logger.info("DNS set to: %s (%s)", name, doh_status)
|
||||
_save_state(dns_provider=name)
|
||||
@@ -1007,21 +1018,13 @@ def set_dns_provider(
|
||||
|
||||
def init_dns_resolvers() -> None:
|
||||
"""Initialize DNS resolvers based on configuration."""
|
||||
global CUSTOM_DNS, DOH_SERVER
|
||||
|
||||
if _is_auto_dns_mode():
|
||||
if _current_dns_index >= 0:
|
||||
name, servers, doh = DNS_PROVIDERS[_current_dns_index]
|
||||
CUSTOM_DNS = servers
|
||||
DOH_SERVER = doh
|
||||
app_config.CUSTOM_DNS = servers
|
||||
app_config.DOH_SERVER = doh
|
||||
_set_runtime_dns_state(servers, doh)
|
||||
logger.info("Using DNS provider: %s (DoH enabled)", name)
|
||||
else:
|
||||
CUSTOM_DNS = []
|
||||
DOH_SERVER = ""
|
||||
app_config.CUSTOM_DNS = []
|
||||
app_config.DOH_SERVER = ""
|
||||
_set_runtime_dns_state([], "")
|
||||
logger.debug("Using system DNS (auto mode - will switch on failure)")
|
||||
socket.getaddrinfo = cast("Any", create_system_failover_getaddrinfo())
|
||||
return
|
||||
@@ -1043,7 +1046,7 @@ def _get_initial_dns_config() -> tuple[str, list[str] | None, bool]:
|
||||
|
||||
"""
|
||||
provider = str(app_config.get("CUSTOM_DNS", "auto")).lower().strip()
|
||||
use_doh = app_config.get("USE_DOH", True)
|
||||
use_doh = coerce_bool(app_config.get("USE_DOH", True), default=True)
|
||||
manual_servers = None
|
||||
|
||||
# Check for manual DNS servers in config
|
||||
@@ -1095,13 +1098,7 @@ def _initialize_aa_state() -> None:
|
||||
_aa_urls = _build_aa_urls()
|
||||
|
||||
# Get configured base URL from config
|
||||
configured_url = normalize_http_url(
|
||||
app_config.get("AA_BASE_URL", "auto"),
|
||||
default_scheme="https",
|
||||
allow_special=("auto",),
|
||||
)
|
||||
if not configured_url:
|
||||
configured_url = "auto"
|
||||
configured_url = _get_configured_aa_url()
|
||||
|
||||
# If AA_BASE_URL is pinned to a custom URL that's not in the mirror list, we still
|
||||
# want to treat it as the active base (and rewrite known mirror links to it).
|
||||
@@ -1222,14 +1219,7 @@ def get_aa_base_url() -> str:
|
||||
|
||||
def is_aa_auto_mode() -> bool:
|
||||
"""Return True when AA_BASE_URL is set to 'auto' (mirror failover enabled)."""
|
||||
configured_url = normalize_http_url(
|
||||
app_config.get("AA_BASE_URL", "auto"),
|
||||
default_scheme="https",
|
||||
allow_special=("auto",),
|
||||
)
|
||||
if not configured_url:
|
||||
configured_url = "auto"
|
||||
return configured_url == "auto"
|
||||
return _get_configured_aa_url() == "auto"
|
||||
|
||||
|
||||
def get_available_aa_urls() -> list[str]:
|
||||
|
||||
@@ -109,20 +109,31 @@ def _parse_release_search_mode(value: object) -> SearchMode:
|
||||
|
||||
|
||||
def _optional_number(value: object) -> float | None:
|
||||
if isinstance(value, bool):
|
||||
return float(value)
|
||||
if isinstance(value, (int, float)) and not isinstance(value, bool):
|
||||
return float(value)
|
||||
try:
|
||||
return float(value)
|
||||
except TypeError, ValueError:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return float(value)
|
||||
except ValueError:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _optional_positive_int(value: object) -> int | None:
|
||||
if isinstance(value, bool):
|
||||
return None
|
||||
try:
|
||||
if isinstance(value, int):
|
||||
parsed = value
|
||||
elif isinstance(value, float):
|
||||
parsed = int(value)
|
||||
except TypeError, ValueError:
|
||||
elif isinstance(value, str):
|
||||
try:
|
||||
parsed = int(value)
|
||||
except ValueError:
|
||||
return None
|
||||
else:
|
||||
return None
|
||||
return parsed if parsed > 0 else None
|
||||
|
||||
@@ -134,6 +145,19 @@ def _seed_time_seconds_to_minutes(value: object) -> int | None:
|
||||
return (seed_time_seconds + 59) // 60
|
||||
|
||||
|
||||
def _config_float(value: object, default: float) -> float:
|
||||
if isinstance(value, bool) or value is None:
|
||||
return default
|
||||
if isinstance(value, (int, float)):
|
||||
return float(value)
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
return float(value)
|
||||
except ValueError:
|
||||
return default
|
||||
return default
|
||||
|
||||
|
||||
def _build_retry_resolution_fields(
|
||||
release_data: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
@@ -689,6 +713,7 @@ def update_download_progress(book_id: str, progress: float) -> None:
|
||||
# Broadcast progress via WebSocket with throttling
|
||||
if ws_manager:
|
||||
current_time = time.time()
|
||||
progress_update_interval = _config_float(config.DOWNLOAD_PROGRESS_UPDATE_INTERVAL, 1.0)
|
||||
should_broadcast = False
|
||||
|
||||
with _progress_lock:
|
||||
@@ -700,7 +725,7 @@ def update_download_progress(book_id: str, progress: float) -> None:
|
||||
should_broadcast = (
|
||||
progress <= _PROGRESS_BROADCAST_START_PERCENT
|
||||
or progress >= _PROGRESS_BROADCAST_COMPLETE_PERCENT
|
||||
or time_elapsed >= config.DOWNLOAD_PROGRESS_UPDATE_INTERVAL
|
||||
or time_elapsed >= progress_update_interval
|
||||
or progress - last_progress >= _PROGRESS_BROADCAST_MIN_DELTA
|
||||
)
|
||||
|
||||
@@ -863,7 +888,8 @@ def _process_single_download(task_id: str, cancel_flag: Event) -> None:
|
||||
|
||||
def concurrent_download_loop() -> None:
|
||||
"""Run the main concurrent download coordinator."""
|
||||
max_workers = config.MAX_CONCURRENT_DOWNLOADS
|
||||
max_workers = normalize_positive_int(config.MAX_CONCURRENT_DOWNLOADS) or 1
|
||||
main_loop_sleep_time = _config_float(config.MAIN_LOOP_SLEEP_TIME, 0.5)
|
||||
logger.info("Starting concurrent download loop with %s workers", max_workers)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=max_workers, thread_name_prefix="Download") as executor:
|
||||
@@ -953,7 +979,7 @@ def concurrent_download_loop() -> None:
|
||||
active_futures[future] = (task_id, cancel_flag)
|
||||
|
||||
# Brief sleep to prevent busy waiting
|
||||
time.sleep(config.MAIN_LOOP_SLEEP_TIME)
|
||||
time.sleep(main_loop_sleep_time)
|
||||
except (AttributeError, KeyError, OSError, RuntimeError, TypeError, ValueError) as e:
|
||||
logger.error_trace("Download coordinator loop error: %s", e)
|
||||
time.sleep(COORDINATOR_LOOP_ERROR_RETRY_DELAY)
|
||||
|
||||
@@ -4,13 +4,29 @@ from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from threading import Event
|
||||
from typing import TYPE_CHECKING, Protocol
|
||||
|
||||
from shelfmark.core.models import DownloadTask
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
from threading import Event
|
||||
|
||||
from shelfmark.core.models import DownloadTask
|
||||
|
||||
StatusCallback = Callable[[str, str | None], None]
|
||||
OutputHandler = Callable[[Path, DownloadTask, Event, StatusCallback, bool], str | None]
|
||||
|
||||
|
||||
class OutputHandler(Protocol):
|
||||
"""Callable contract for post-download output handlers."""
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
temp_file: Path,
|
||||
task: DownloadTask,
|
||||
cancel_flag: Event,
|
||||
status_callback: StatusCallback,
|
||||
*,
|
||||
preserve_source_on_failure: bool = False,
|
||||
) -> str | None: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
||||
@@ -71,6 +71,9 @@ def _parse_int(value: object, label: str) -> int:
|
||||
if value is None or value == "":
|
||||
msg = f"{label} is required"
|
||||
raise BookloreError(msg)
|
||||
if not isinstance(value, (int, float, str)):
|
||||
msg = f"{label} must be a number"
|
||||
raise BookloreError(msg)
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
|
||||
@@ -63,6 +63,9 @@ def _parse_int(value: Any, label: str, *, minimum: int = 1) -> int:
|
||||
if value is None or value == "":
|
||||
msg = f"{label} is required"
|
||||
raise EmailOutputError(msg)
|
||||
if not isinstance(value, (int, float, str)):
|
||||
msg = f"{label} must be a number"
|
||||
raise EmailOutputError(msg)
|
||||
try:
|
||||
parsed = int(value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
@@ -137,6 +140,15 @@ def _get_email_settings() -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _parse_attachment_limit_mb(value: object) -> int:
|
||||
if not isinstance(value, (int, float, str)):
|
||||
return 25
|
||||
try:
|
||||
return int(value)
|
||||
except TypeError, ValueError:
|
||||
return 25
|
||||
|
||||
|
||||
def _render_subject(template: str, task: DownloadTask) -> str:
|
||||
mapping = {
|
||||
"Author": task.author or "",
|
||||
@@ -362,10 +374,7 @@ def _post_process_email(
|
||||
success = False
|
||||
try:
|
||||
limit_mb_raw = core_config.config.get("EMAIL_ATTACHMENT_SIZE_LIMIT_MB", 25)
|
||||
try:
|
||||
attachment_limit_mb = int(limit_mb_raw)
|
||||
except TypeError, ValueError:
|
||||
attachment_limit_mb = 25
|
||||
attachment_limit_mb = _parse_attachment_limit_mb(limit_mb_raw)
|
||||
|
||||
if attachment_limit_mb > 0:
|
||||
limit_bytes = attachment_limit_mb * 1024 * 1024
|
||||
|
||||
@@ -45,7 +45,7 @@ def _log_path_permissions(probe: Path, label: str) -> None:
|
||||
logger.debug("Path permissions (%s): stat failed for %s: %s", label, probe, stat_error)
|
||||
|
||||
|
||||
def _run_io[T](func: Callable[..., _T], *args: Any, **kwargs: Any) -> _T:
|
||||
def _run_io(func: Callable[..., _T], *args: Any, **kwargs: Any) -> _T: # noqa: UP047
|
||||
"""Best-effort offload for potentially blocking filesystem calls.
|
||||
|
||||
Keep this module import-cycle safe: `shelfmark.download.fs` imports this module,
|
||||
|
||||
@@ -284,7 +284,8 @@ def maybe_run_custom_script(
|
||||
)
|
||||
return True
|
||||
|
||||
path_mode = core_config.config.get("CUSTOM_SCRIPT_PATH_MODE", "absolute")
|
||||
configured_path_mode = core_config.config.get("CUSTOM_SCRIPT_PATH_MODE", "absolute")
|
||||
path_mode = configured_path_mode if isinstance(configured_path_mode, str) else "absolute"
|
||||
|
||||
payload: dict[str, Any] | None = None
|
||||
if core_config.config.get("CUSTOM_SCRIPT_JSON_PAYLOAD", False):
|
||||
|
||||
@@ -17,45 +17,50 @@ circular imports (`archive` is used by the pipeline).
|
||||
from __future__ import annotations
|
||||
|
||||
import shelfmark.core.config as core_config
|
||||
from shelfmark.core.request_helpers import coerce_bool
|
||||
|
||||
|
||||
def _normalize_format_list(value: object, default: list[str]) -> list[str]:
|
||||
if isinstance(value, str):
|
||||
return [fmt.strip().lower() for fmt in value.split(",") if fmt.strip()]
|
||||
if isinstance(value, (list, tuple, set)):
|
||||
normalized = [str(fmt).strip().lower() for fmt in value if str(fmt).strip()]
|
||||
return normalized or default
|
||||
return default
|
||||
|
||||
|
||||
def _config_text(value: object) -> str:
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
return ""
|
||||
|
||||
|
||||
def get_supported_formats() -> list[str]:
|
||||
"""Get current supported formats from config singleton."""
|
||||
formats = core_config.config.get(
|
||||
"SUPPORTED_FORMATS",
|
||||
["epub", "mobi", "azw3", "fb2", "djvu", "cbz", "cbr"],
|
||||
)
|
||||
|
||||
# Handle both list (from MultiSelectField) and comma-separated string (legacy/env)
|
||||
if isinstance(formats, str):
|
||||
return [fmt.strip().lower() for fmt in formats.split(",") if fmt.strip()]
|
||||
|
||||
return [fmt.lower() for fmt in formats]
|
||||
default_formats = ["epub", "mobi", "azw3", "fb2", "djvu", "cbz", "cbr"]
|
||||
formats = core_config.config.get("SUPPORTED_FORMATS", default_formats)
|
||||
return _normalize_format_list(formats, default_formats)
|
||||
|
||||
|
||||
def get_supported_audiobook_formats() -> list[str]:
|
||||
"""Get current supported audiobook formats from config singleton."""
|
||||
formats = core_config.config.get("SUPPORTED_AUDIOBOOK_FORMATS", ["m4b", "mp3"])
|
||||
|
||||
# Handle both list (from MultiSelectField) and comma-separated string (legacy/env)
|
||||
if isinstance(formats, str):
|
||||
return [fmt.strip().lower() for fmt in formats.split(",") if fmt.strip()]
|
||||
|
||||
return [fmt.lower() for fmt in formats]
|
||||
default_formats = ["m4b", "mp3"]
|
||||
formats = core_config.config.get("SUPPORTED_AUDIOBOOK_FORMATS", default_formats)
|
||||
return _normalize_format_list(formats, default_formats)
|
||||
|
||||
|
||||
def get_file_organization(*, is_audiobook: bool) -> str:
|
||||
"""Get the file organization mode for the content type."""
|
||||
key = "FILE_ORGANIZATION_AUDIOBOOK" if is_audiobook else "FILE_ORGANIZATION"
|
||||
mode = core_config.config.get(key, "rename")
|
||||
mode = _config_text(core_config.config.get(key, "rename")).strip().lower()
|
||||
|
||||
# Handle legacy settings migration
|
||||
if mode not in ("none", "rename", "organize"):
|
||||
legacy_key = "PROCESSING_MODE_AUDIOBOOK" if is_audiobook else "PROCESSING_MODE"
|
||||
legacy_mode = core_config.config.get(legacy_key, "ingest")
|
||||
legacy_mode = _config_text(core_config.config.get(legacy_key, "ingest")).strip().lower()
|
||||
if legacy_mode == "library":
|
||||
return "organize"
|
||||
if core_config.config.get("USE_BOOK_TITLE", True):
|
||||
if coerce_bool(core_config.config.get("USE_BOOK_TITLE", True), default=True):
|
||||
return "rename"
|
||||
return "none"
|
||||
|
||||
@@ -73,16 +78,16 @@ def get_template(*, is_audiobook: bool, organization_mode: str) -> str:
|
||||
else:
|
||||
key = "TEMPLATE_ORGANIZE" if organization_mode == "organize" else "TEMPLATE_RENAME"
|
||||
|
||||
template = core_config.config.get(key, "")
|
||||
template = _config_text(core_config.config.get(key, ""))
|
||||
|
||||
# Fallback to legacy keys if new keys are empty
|
||||
if not template:
|
||||
legacy_key = "TEMPLATE_AUDIOBOOK" if is_audiobook else "TEMPLATE"
|
||||
template = core_config.config.get(legacy_key, "")
|
||||
template = _config_text(core_config.config.get(legacy_key, ""))
|
||||
|
||||
if not template:
|
||||
legacy_key = "LIBRARY_TEMPLATE_AUDIOBOOK" if is_audiobook else "LIBRARY_TEMPLATE"
|
||||
template = core_config.config.get(legacy_key, "")
|
||||
template = _config_text(core_config.config.get(legacy_key, ""))
|
||||
|
||||
if not template:
|
||||
if organization_mode == "organize":
|
||||
|
||||
@@ -71,7 +71,8 @@ def prepare_output_files(
|
||||
step_label = (
|
||||
"Staging torrent files" if output_plan.stage_action == STAGE_COPY else "Staging files"
|
||||
)
|
||||
status_callback("resolving", step_label)
|
||||
if status_callback is not None:
|
||||
status_callback("resolving", step_label)
|
||||
working_path = stage_path(working_path, output_plan.staging_dir, output_plan.stage_action)
|
||||
|
||||
can_delete_source_archives = (
|
||||
@@ -88,7 +89,8 @@ def prepare_output_files(
|
||||
)
|
||||
|
||||
if error:
|
||||
status_callback("error", error)
|
||||
if status_callback is not None:
|
||||
status_callback("error", error)
|
||||
if not preserve_source_on_failure:
|
||||
cleanup_output_staging(output_plan, working_path, task, cleanup_paths)
|
||||
return None
|
||||
|
||||
@@ -10,7 +10,7 @@ Keeping this separate from `pipeline.py` avoids circular imports:
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Protocol, TypeGuard
|
||||
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.models import DownloadTask, SearchMode
|
||||
@@ -24,6 +24,22 @@ if TYPE_CHECKING:
|
||||
logger = setup_logger(__name__)
|
||||
|
||||
|
||||
class _PostProcessHandler(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
temp_file: Path,
|
||||
task: DownloadTask,
|
||||
cancel_flag: Event,
|
||||
status_callback: Callable[[str, str | None], None],
|
||||
*,
|
||||
preserve_source_on_failure: bool = False,
|
||||
) -> str | None: ...
|
||||
|
||||
|
||||
def _is_post_process_handler(candidate: object) -> TypeGuard[_PostProcessHandler]:
|
||||
return callable(candidate)
|
||||
|
||||
|
||||
def post_process_download(
|
||||
temp_file: Path,
|
||||
task: DownloadTask,
|
||||
@@ -48,7 +64,10 @@ def post_process_download(
|
||||
output_handler = resolve_output_handler(task)
|
||||
if output_handler:
|
||||
logger.info("Task %s: using output mode %s", task.task_id, output_handler.mode)
|
||||
return output_handler.handler(
|
||||
registered_handler = output_handler.handler
|
||||
if not _is_post_process_handler(registered_handler):
|
||||
return None
|
||||
return registered_handler(
|
||||
temp_file,
|
||||
task,
|
||||
cancel_flag,
|
||||
@@ -59,6 +78,8 @@ def post_process_download(
|
||||
from shelfmark.download.outputs.folder import process_folder_output
|
||||
|
||||
logger.info("Task %s: using output mode folder", task.task_id)
|
||||
if not _is_post_process_handler(process_folder_output):
|
||||
return None
|
||||
return process_folder_output(
|
||||
temp_file,
|
||||
task,
|
||||
|
||||
+135
-55
@@ -12,7 +12,7 @@ from datetime import UTC, datetime, timedelta
|
||||
from functools import wraps
|
||||
from importlib import import_module
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Callable, NoReturn
|
||||
from typing import TYPE_CHECKING, Any, Callable, NoReturn, cast
|
||||
|
||||
from flask import Flask, jsonify, request, send_file, send_from_directory, session
|
||||
from flask_cors import CORS
|
||||
@@ -106,14 +106,15 @@ def _raise_runtime_error(message: str) -> NoReturn:
|
||||
PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
||||
FRONTEND_DIST = PROJECT_ROOT / "frontend-dist"
|
||||
|
||||
BASE_PATH = normalize_base_path(app_config.get("URL_BASE", ""))
|
||||
BASE_PATH = normalize_base_path(normalize_optional_text(app_config.get("URL_BASE", "")))
|
||||
|
||||
app = Flask(__name__)
|
||||
app.config["SEND_FILE_MAX_AGE_DEFAULT"] = 0 # Disable caching
|
||||
app.config["APPLICATION_ROOT"] = BASE_PATH or "/"
|
||||
app.wsgi_app = ProxyFix(app.wsgi_app) # type: ignore[assignment]
|
||||
wsgi_app = cast(Any, ProxyFix(app.wsgi_app))
|
||||
if BASE_PATH:
|
||||
app.wsgi_app = PrefixMiddleware(app.wsgi_app, BASE_PATH, bypass_paths={"/api/health"})
|
||||
wsgi_app = cast(Any, PrefixMiddleware(wsgi_app, BASE_PATH, bypass_paths={"/api/health"}))
|
||||
app.wsgi_app = wsgi_app
|
||||
|
||||
# Socket.IO async mode.
|
||||
# We run this app under Gunicorn with a gevent websocket worker (even when DEBUG=true),
|
||||
@@ -123,23 +124,23 @@ socketio_cors_allowed_origins = "*"
|
||||
|
||||
# Initialize Flask-SocketIO with reverse proxy support
|
||||
socketio_path = f"{BASE_PATH}/socket.io" if BASE_PATH else "/socket.io"
|
||||
socketio = SocketIO(
|
||||
app,
|
||||
cors_allowed_origins=socketio_cors_allowed_origins,
|
||||
async_mode=async_mode,
|
||||
logger=False,
|
||||
engineio_logger=False,
|
||||
socketio_init_kwargs: dict[str, Any] = {
|
||||
"cors_allowed_origins": socketio_cors_allowed_origins,
|
||||
"async_mode": async_mode,
|
||||
"logger": False,
|
||||
"engineio_logger": False,
|
||||
# Reverse proxy / Traefik compatibility settings
|
||||
path=socketio_path,
|
||||
ping_timeout=60, # Time to wait for pong response
|
||||
ping_interval=25, # Send ping every 25 seconds
|
||||
"path": socketio_path,
|
||||
"ping_timeout": 60,
|
||||
"ping_interval": 25,
|
||||
# Allow both websocket and polling for better compatibility
|
||||
transports=["websocket", "polling"],
|
||||
"transports": ["websocket", "polling"],
|
||||
# Enable CORS for all origins (you can restrict this in production)
|
||||
allow_upgrades=True,
|
||||
"allow_upgrades": True,
|
||||
# Important for proxies that buffer
|
||||
http_compression=True,
|
||||
)
|
||||
"http_compression": True,
|
||||
}
|
||||
socketio = SocketIO(app, **socketio_init_kwargs)
|
||||
|
||||
# Initialize WebSocket manager
|
||||
ws_manager.init_app(app, socketio)
|
||||
@@ -203,16 +204,44 @@ LOGIN_ATTEMPT_WARNING_THRESHOLD = 5
|
||||
def cleanup_old_lockouts() -> None:
|
||||
"""Remove expired lockout entries to prevent memory buildup."""
|
||||
current_time = datetime.now(UTC)
|
||||
expired_users = [
|
||||
username
|
||||
for username, data in failed_login_attempts.items()
|
||||
if "lockout_until" in data and data["lockout_until"] < current_time
|
||||
]
|
||||
expired_users = []
|
||||
for username in list(failed_login_attempts):
|
||||
lockout_until = _get_lockout_until(username, repair_if_locked=True)
|
||||
if lockout_until is not None and lockout_until < current_time:
|
||||
expired_users.append(username)
|
||||
for username in expired_users:
|
||||
logger.info("Lockout expired for user: %s", username)
|
||||
del failed_login_attempts[username]
|
||||
|
||||
|
||||
def _get_lockout_until(username: str, *, repair_if_locked: bool = False) -> datetime | None:
|
||||
"""Return a valid lockout timestamp for the user when one exists.
|
||||
|
||||
When a user has already crossed the lockout threshold but the timestamp is
|
||||
missing or malformed, optionally repair the state to keep the lockout in
|
||||
force rather than silently letting the user through.
|
||||
"""
|
||||
lockout_state = failed_login_attempts.get(username)
|
||||
if lockout_state is None:
|
||||
return None
|
||||
|
||||
lockout_until = lockout_state.get("lockout_until")
|
||||
if isinstance(lockout_until, datetime):
|
||||
return lockout_until
|
||||
|
||||
attempt_count = lockout_state.get("count")
|
||||
if repair_if_locked and isinstance(attempt_count, int) and attempt_count >= MAX_LOGIN_ATTEMPTS:
|
||||
repaired_lockout_until = datetime.now(UTC) + timedelta(minutes=LOCKOUT_DURATION_MINUTES)
|
||||
lockout_state["lockout_until"] = repaired_lockout_until
|
||||
logger.warning("Repaired missing lockout timestamp for locked account '%s'", username)
|
||||
return repaired_lockout_until
|
||||
|
||||
if lockout_until is not None:
|
||||
logger.warning("Ignoring invalid lockout timestamp for user '%s'", username)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def is_account_locked(username: str) -> bool:
|
||||
"""Check if an account is currently locked due to failed login attempts."""
|
||||
cleanup_old_lockouts()
|
||||
@@ -220,7 +249,7 @@ def is_account_locked(username: str) -> bool:
|
||||
if username not in failed_login_attempts:
|
||||
return False
|
||||
|
||||
lockout_until = failed_login_attempts[username].get("lockout_until")
|
||||
lockout_until = _get_lockout_until(username, repair_if_locked=True)
|
||||
return lockout_until is not None and datetime.now(UTC) < lockout_until
|
||||
|
||||
|
||||
@@ -656,7 +685,10 @@ def proxy_auth_middleware() -> Response | tuple[Response, int] | None:
|
||||
return None
|
||||
|
||||
try:
|
||||
user_header = app_config.get("PROXY_AUTH_USER_HEADER", "X-Auth-User")
|
||||
user_header = (
|
||||
normalize_optional_text(app_config.get("PROXY_AUTH_USER_HEADER", "X-Auth-User"))
|
||||
or "X-Auth-User"
|
||||
)
|
||||
|
||||
# Extract username from proxy header
|
||||
username = get_proxy_header(user_header)
|
||||
@@ -672,8 +704,15 @@ def proxy_auth_middleware() -> Response | tuple[Response, int] | None:
|
||||
# If an admin group is configured, derive from groups header.
|
||||
# Otherwise preserve existing DB role for known users and default
|
||||
# first-time users to admin (to avoid lockouts).
|
||||
admin_group_header = app_config.get("PROXY_AUTH_ADMIN_GROUP_HEADER", "X-Auth-Groups")
|
||||
admin_group_name = str(app_config.get("PROXY_AUTH_ADMIN_GROUP_NAME", "") or "").strip()
|
||||
admin_group_header = (
|
||||
normalize_optional_text(
|
||||
app_config.get("PROXY_AUTH_ADMIN_GROUP_HEADER", "X-Auth-Groups")
|
||||
)
|
||||
or "X-Auth-Groups"
|
||||
)
|
||||
admin_group_name = (
|
||||
normalize_optional_text(app_config.get("PROXY_AUTH_ADMIN_GROUP_NAME", "")) or ""
|
||||
)
|
||||
is_admin = True
|
||||
|
||||
if admin_group_name:
|
||||
@@ -860,7 +899,10 @@ if DEBUG:
|
||||
if app_config.get("USING_EXTERNAL_BYPASSER", False):
|
||||
pass
|
||||
else:
|
||||
from shelfmark.bypass.internal_bypasser import _cleanup_orphan_processes as _stop_gui
|
||||
from shelfmark.bypass.internal_bypasser import _cleanup_orphan_processes
|
||||
|
||||
def _stop_gui() -> None:
|
||||
_cleanup_orphan_processes()
|
||||
|
||||
@app.route("/api/debug", methods=["GET"])
|
||||
@login_required
|
||||
@@ -1076,15 +1118,19 @@ def api_config() -> Response | tuple[Response, int]:
|
||||
"",
|
||||
user_id=db_user_id,
|
||||
)
|
||||
configured_metadata_provider = app_config.get(
|
||||
"METADATA_PROVIDER",
|
||||
"",
|
||||
user_id=db_user_id,
|
||||
configured_metadata_provider = normalize_optional_text(
|
||||
app_config.get(
|
||||
"METADATA_PROVIDER",
|
||||
"",
|
||||
user_id=db_user_id,
|
||||
)
|
||||
)
|
||||
_configured_metadata_provider_audiobook = app_config.get(
|
||||
"METADATA_PROVIDER_AUDIOBOOK",
|
||||
"",
|
||||
user_id=db_user_id,
|
||||
_configured_metadata_provider_audiobook = normalize_optional_text(
|
||||
app_config.get(
|
||||
"METADATA_PROVIDER_AUDIOBOOK",
|
||||
"",
|
||||
user_id=db_user_id,
|
||||
)
|
||||
)
|
||||
metadata_ui_provider = (
|
||||
configured_metadata_provider or _configured_metadata_provider_audiobook
|
||||
@@ -1145,7 +1191,7 @@ def api_health() -> Response | tuple[Response, int]:
|
||||
flask.Response: JSON with status "ok" and optional degraded features.
|
||||
|
||||
"""
|
||||
response = {"status": "ok"}
|
||||
response: dict[str, object] = {"status": "ok"}
|
||||
|
||||
# Report degraded features
|
||||
if not backend.WEBSOCKET_AVAILABLE:
|
||||
@@ -1695,12 +1741,17 @@ def api_retry_download(book_id: str) -> Response | tuple[Response, int]:
|
||||
request_id = normalize_positive_int(history_row.get("request_id"))
|
||||
retry_payload = history_row.get("retry_payload")
|
||||
final_status = history_row.get("final_status")
|
||||
if request_id is not None and not download_history_service.is_retry_available(
|
||||
history_row
|
||||
):
|
||||
return jsonify(
|
||||
{"error": "Forbidden", "code": "requested_download_retry_forbidden"}
|
||||
), 403
|
||||
if request_id is not None:
|
||||
history_service = download_history_service
|
||||
if history_service is None:
|
||||
logger.error(
|
||||
"Download history service unavailable while retrying task %s", book_id
|
||||
)
|
||||
return jsonify({"error": "Download history unavailable"}), 500
|
||||
if not history_service.is_retry_available(history_row):
|
||||
return jsonify(
|
||||
{"error": "Forbidden", "code": "requested_download_retry_forbidden"}
|
||||
), 403
|
||||
success, error = backend.retry_persisted_download(
|
||||
retry_payload,
|
||||
final_status=final_status,
|
||||
@@ -1909,7 +1960,14 @@ def api_login() -> Response | tuple[Response, int]:
|
||||
|
||||
# Check if account is locked due to failed login attempts
|
||||
if is_account_locked(username):
|
||||
lockout_until = failed_login_attempts[username].get("lockout_until")
|
||||
lockout_until = _get_lockout_until(username, repair_if_locked=True)
|
||||
if lockout_until is None:
|
||||
logger.error("Locked account '%s' is missing a lockout timestamp", username)
|
||||
return jsonify(
|
||||
{
|
||||
"error": f"Account temporarily locked due to multiple failed login attempts. Try again in {LOCKOUT_DURATION_MINUTES} minutes."
|
||||
}
|
||||
), 429
|
||||
remaining_time = (lockout_until - datetime.now(UTC)).total_seconds() / 60
|
||||
logger.warning(
|
||||
"Login attempt blocked for locked account '%s' from IP %s", username, ip_address
|
||||
@@ -2652,13 +2710,15 @@ def api_releases() -> Response | tuple[Response, int]:
|
||||
source_results_are_releases,
|
||||
)
|
||||
|
||||
def _search_source_releases(source_name: str) -> tuple[Any | None, list[Any], str | None]:
|
||||
def _search_source_releases(
|
||||
source_name: str, search_book: BookMetadata
|
||||
) -> tuple[Any | None, list[Any], str | None]:
|
||||
"""Search one source and return any error message instead of raising."""
|
||||
try:
|
||||
source = get_source(source_name)
|
||||
|
||||
plan = build_release_search_plan(
|
||||
book,
|
||||
search_book,
|
||||
languages=browse_filters.lang
|
||||
if source_query_filters is not None
|
||||
else languages,
|
||||
@@ -2685,14 +2745,14 @@ def api_releases() -> Response | tuple[Response, int]:
|
||||
source_name,
|
||||
planned_query_type,
|
||||
planned_query,
|
||||
book.title,
|
||||
book.authors,
|
||||
search_book.title,
|
||||
search_book.authors,
|
||||
expand_search,
|
||||
content_type,
|
||||
)
|
||||
|
||||
releases = source.search(
|
||||
book, plan, expand_search=expand_search, content_type=content_type
|
||||
search_book, plan, expand_search=expand_search, content_type=content_type
|
||||
)
|
||||
except ValueError:
|
||||
return None, [], f"Unknown source: {source_name}"
|
||||
@@ -2735,6 +2795,8 @@ def api_releases() -> Response | tuple[Response, int]:
|
||||
source_query_filters = None
|
||||
is_source_provider = bool(provider) and source_results_are_releases(provider)
|
||||
|
||||
book: BookMetadata
|
||||
|
||||
if not provider or not book_id:
|
||||
if not source_filter or not has_browse_filters:
|
||||
return jsonify({"error": "Parameters 'provider' and 'book_id' are required"}), 400
|
||||
@@ -2780,10 +2842,11 @@ def api_releases() -> Response | tuple[Response, int]:
|
||||
# Get book metadata from provider
|
||||
kwargs = get_provider_kwargs(provider)
|
||||
prov = get_provider(provider, **kwargs)
|
||||
book = prov.get_book(book_id)
|
||||
resolved_book = prov.get_book(book_id)
|
||||
|
||||
if not book:
|
||||
if not resolved_book:
|
||||
return jsonify({"error": "Book not found in metadata provider"}), 404
|
||||
book = resolved_book
|
||||
|
||||
# Override title from frontend if available (search results may have better data)
|
||||
# Note: We intentionally DON'T override authors here - get_book() now returns
|
||||
@@ -2808,7 +2871,7 @@ def api_releases() -> Response | tuple[Response, int]:
|
||||
source_instances = {} # Keep source instances for column config
|
||||
|
||||
for source_name in sources_to_search:
|
||||
source, releases, error = _search_source_releases(source_name)
|
||||
source, releases, error = _search_source_releases(source_name, book)
|
||||
if source is not None:
|
||||
source_instances[source_name] = source
|
||||
all_releases.extend(releases)
|
||||
@@ -3139,7 +3202,7 @@ def api_onboarding_skip() -> Response | tuple[Response, int]:
|
||||
# Catch-all route for React Router (must be last)
|
||||
# This handles client-side routing by serving index.html for any unmatched routes
|
||||
@app.route("/<path:path>")
|
||||
def catch_all(path: str) -> Response:
|
||||
def catch_all(path: str) -> Response | tuple[Response, int]:
|
||||
"""Serve the React app for any route not matched by API endpoints.
|
||||
|
||||
This allows React Router to handle client-side routing.
|
||||
@@ -3152,6 +3215,12 @@ def catch_all(path: str) -> Response:
|
||||
return _serve_index_html()
|
||||
|
||||
|
||||
def _get_request_sid() -> str | None:
|
||||
"""Return the Socket.IO session id for the active request when available."""
|
||||
sid = getattr(request, "sid", None)
|
||||
return sid if isinstance(sid, str) and sid else None
|
||||
|
||||
|
||||
# WebSocket event handlers
|
||||
@socketio.on("connect")
|
||||
def handle_connect() -> None:
|
||||
@@ -3163,7 +3232,11 @@ def handle_connect() -> None:
|
||||
|
||||
# Join appropriate room based on authenticated user session
|
||||
is_admin, db_user_id, can_access_status = _resolve_status_scope()
|
||||
ws_manager.join_user_room(request.sid, is_admin=is_admin, db_user_id=db_user_id)
|
||||
sid = _get_request_sid()
|
||||
if sid is None:
|
||||
logger.warning("Socket.IO connect event missing sid")
|
||||
return
|
||||
ws_manager.join_user_room(sid, is_admin=is_admin, db_user_id=db_user_id)
|
||||
|
||||
# Send initial status to the newly connected client (filtered)
|
||||
try:
|
||||
@@ -3184,7 +3257,9 @@ def handle_disconnect() -> None:
|
||||
logger.info("WebSocket client disconnected")
|
||||
|
||||
# Leave room
|
||||
ws_manager.leave_user_room(request.sid)
|
||||
sid = _get_request_sid()
|
||||
if sid is not None:
|
||||
ws_manager.leave_user_room(sid)
|
||||
|
||||
# Track the disconnection
|
||||
ws_manager.client_disconnected()
|
||||
@@ -3195,7 +3270,12 @@ def handle_status_request() -> None:
|
||||
"""Handle manual status request from client."""
|
||||
try:
|
||||
is_admin, db_user_id, can_access_status = _resolve_status_scope()
|
||||
ws_manager.sync_user_room(request.sid, is_admin=is_admin, db_user_id=db_user_id)
|
||||
sid = _get_request_sid()
|
||||
if sid is None:
|
||||
logger.warning("Socket.IO request_status event missing sid")
|
||||
emit("status_update", {})
|
||||
return
|
||||
ws_manager.sync_user_room(sid, is_admin=is_admin, db_user_id=db_user_id)
|
||||
|
||||
if not can_access_status:
|
||||
emit("status_update", {})
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
"""Metadata provider plugin system - base classes and registry."""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Callable
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass, field
|
||||
from enum import StrEnum
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
from typing import Any, ClassVar, TypeVar
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
from shelfmark.core.request_helpers import normalize_optional_text
|
||||
|
||||
|
||||
class SearchType(StrEnum):
|
||||
@@ -431,15 +431,20 @@ class MetadataProvider(ABC):
|
||||
|
||||
# Provider registry
|
||||
_PROVIDERS: dict[str, type[MetadataProvider]] = {}
|
||||
_PROVIDER_KWARGS_FACTORIES: dict[str, Any] = {} # Callable[[], Dict]
|
||||
_PROVIDER_KWARGS_FACTORIES: dict[str, Callable[[], dict[str, Any]]] = {}
|
||||
ProviderType = TypeVar("ProviderType", bound=MetadataProvider)
|
||||
ProviderKwargsFactory = TypeVar(
|
||||
"ProviderKwargsFactory",
|
||||
bound=Callable[[], dict[str, Any]],
|
||||
)
|
||||
|
||||
|
||||
def register_provider(
|
||||
name: str,
|
||||
) -> Callable[[type[MetadataProvider]], type[MetadataProvider]]:
|
||||
) -> Callable[[type[ProviderType]], type[ProviderType]]:
|
||||
"""Register a metadata provider."""
|
||||
|
||||
def decorator(cls: type[MetadataProvider]) -> type[MetadataProvider]:
|
||||
def decorator(cls: type[ProviderType]) -> type[ProviderType]:
|
||||
_PROVIDERS[name] = cls
|
||||
return cls
|
||||
|
||||
@@ -448,7 +453,7 @@ def register_provider(
|
||||
|
||||
def register_provider_kwargs(
|
||||
name: str,
|
||||
) -> Callable[[Callable[[], dict[str, Any]]], Callable[[], dict[str, Any]]]:
|
||||
) -> Callable[[ProviderKwargsFactory], ProviderKwargsFactory]:
|
||||
"""Register a provider kwargs factory.
|
||||
|
||||
The decorated function should return a Dict of kwargs to pass to the
|
||||
@@ -463,7 +468,7 @@ def register_provider_kwargs(
|
||||
|
||||
"""
|
||||
|
||||
def decorator(fn: Callable[[], dict[str, Any]]) -> Callable[[], dict[str, Any]]:
|
||||
def decorator(fn: ProviderKwargsFactory) -> ProviderKwargsFactory:
|
||||
_PROVIDER_KWARGS_FACTORIES[name] = fn
|
||||
return fn
|
||||
|
||||
@@ -528,11 +533,17 @@ def get_configured_provider(
|
||||
|
||||
# For audiobooks, try audiobook-specific provider first, then fall back to main provider
|
||||
if content_type == "audiobook":
|
||||
metadata_provider = app_config.get("METADATA_PROVIDER_AUDIOBOOK", "", user_id=user_id)
|
||||
metadata_provider = normalize_optional_text(
|
||||
app_config.get("METADATA_PROVIDER_AUDIOBOOK", "", user_id=user_id)
|
||||
)
|
||||
if not metadata_provider:
|
||||
metadata_provider = app_config.get("METADATA_PROVIDER", "", user_id=user_id)
|
||||
metadata_provider = normalize_optional_text(
|
||||
app_config.get("METADATA_PROVIDER", "", user_id=user_id)
|
||||
)
|
||||
else:
|
||||
metadata_provider = app_config.get("METADATA_PROVIDER", "", user_id=user_id)
|
||||
metadata_provider = normalize_optional_text(
|
||||
app_config.get("METADATA_PROVIDER", "", user_id=user_id)
|
||||
)
|
||||
|
||||
if not metadata_provider:
|
||||
return None
|
||||
@@ -560,24 +571,28 @@ def get_configured_provider_name(
|
||||
app_config.refresh()
|
||||
|
||||
if content_type == "combined":
|
||||
combined_provider = app_config.get(
|
||||
"METADATA_PROVIDER_COMBINED",
|
||||
"",
|
||||
user_id=user_id,
|
||||
combined_provider = normalize_optional_text(
|
||||
app_config.get(
|
||||
"METADATA_PROVIDER_COMBINED",
|
||||
"",
|
||||
user_id=user_id,
|
||||
)
|
||||
)
|
||||
if combined_provider or not fallback_to_main:
|
||||
return combined_provider
|
||||
return combined_provider or ""
|
||||
|
||||
if content_type == "audiobook":
|
||||
audiobook_provider = app_config.get(
|
||||
"METADATA_PROVIDER_AUDIOBOOK",
|
||||
"",
|
||||
user_id=user_id,
|
||||
audiobook_provider = normalize_optional_text(
|
||||
app_config.get(
|
||||
"METADATA_PROVIDER_AUDIOBOOK",
|
||||
"",
|
||||
user_id=user_id,
|
||||
)
|
||||
)
|
||||
if audiobook_provider or not fallback_to_main:
|
||||
return audiobook_provider
|
||||
return audiobook_provider or ""
|
||||
|
||||
return app_config.get("METADATA_PROVIDER", "", user_id=user_id)
|
||||
return normalize_optional_text(app_config.get("METADATA_PROVIDER", "", user_id=user_id)) or ""
|
||||
|
||||
|
||||
def get_provider_sort_options(
|
||||
@@ -649,7 +664,9 @@ def get_provider_default_sort(
|
||||
|
||||
# Look up provider-specific default sort setting
|
||||
setting_key = f"{provider_name.upper()}_DEFAULT_SORT"
|
||||
return app_config.get(setting_key, "relevance", user_id=user_id)
|
||||
return normalize_optional_text(app_config.get(setting_key, "relevance", user_id=user_id)) or (
|
||||
"relevance"
|
||||
)
|
||||
|
||||
|
||||
def sync_metadata_provider_selection() -> None:
|
||||
|
||||
@@ -15,6 +15,7 @@ import requests
|
||||
from shelfmark.core.cache import cacheable
|
||||
from shelfmark.core.config import config as app_config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.request_helpers import normalize_optional_text
|
||||
from shelfmark.core.settings_registry import (
|
||||
ActionButton,
|
||||
CheckboxField,
|
||||
@@ -30,6 +31,7 @@ from shelfmark.metadata_providers import (
|
||||
DisplayField,
|
||||
MetadataProvider,
|
||||
MetadataSearchOptions,
|
||||
SearchField,
|
||||
SearchType,
|
||||
SortOrder,
|
||||
TextSearchField,
|
||||
@@ -53,10 +55,15 @@ SORT_MAPPING: dict[SortOrder, str | None] = {
|
||||
}
|
||||
|
||||
|
||||
def _normalize_googlebooks_api_key(value: object) -> str:
|
||||
"""Normalize Google Books API keys loaded from config or form values."""
|
||||
return normalize_optional_text(value) or ""
|
||||
|
||||
|
||||
@register_provider_kwargs("googlebooks")
|
||||
def _googlebooks_kwargs() -> dict[str, Any]:
|
||||
"""Provide Google Books-specific constructor kwargs."""
|
||||
return {"api_key": app_config.get("GOOGLEBOOKS_API_KEY", "")}
|
||||
return {"api_key": _normalize_googlebooks_api_key(app_config.get("GOOGLEBOOKS_API_KEY", ""))}
|
||||
|
||||
|
||||
@register_provider("googlebooks")
|
||||
@@ -70,7 +77,7 @@ class GoogleBooksProvider(MetadataProvider):
|
||||
SortOrder.RELEVANCE,
|
||||
SortOrder.NEWEST,
|
||||
)
|
||||
search_fields: ClassVar[tuple[TextSearchField, ...]] = (
|
||||
search_fields: ClassVar[tuple[SearchField, ...]] = (
|
||||
TextSearchField(
|
||||
key="author",
|
||||
label="Author",
|
||||
@@ -85,7 +92,8 @@ class GoogleBooksProvider(MetadataProvider):
|
||||
|
||||
def __init__(self, api_key: str | None = None) -> None:
|
||||
"""Initialize provider with optional API key (falls back to config)."""
|
||||
self.api_key = api_key or app_config.get("GOOGLEBOOKS_API_KEY", "")
|
||||
raw_key = api_key or app_config.get("GOOGLEBOOKS_API_KEY", "")
|
||||
self.api_key = _normalize_googlebooks_api_key(raw_key)
|
||||
self.session = requests.Session()
|
||||
|
||||
def is_available(self) -> bool:
|
||||
@@ -364,7 +372,8 @@ def _test_googlebooks_connection(
|
||||
current_values = current_values or {}
|
||||
|
||||
# Use current form values first, fall back to saved config
|
||||
api_key = current_values.get("GOOGLEBOOKS_API_KEY") or app_config.get("GOOGLEBOOKS_API_KEY", "")
|
||||
raw_key = current_values.get("GOOGLEBOOKS_API_KEY") or app_config.get("GOOGLEBOOKS_API_KEY", "")
|
||||
api_key = _normalize_googlebooks_api_key(raw_key)
|
||||
|
||||
if not api_key:
|
||||
return {
|
||||
|
||||
@@ -13,7 +13,7 @@ import requests
|
||||
from shelfmark.core.cache import cache_key, cacheable, get_metadata_cache
|
||||
from shelfmark.core.config import config as app_config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.request_helpers import coerce_int
|
||||
from shelfmark.core.request_helpers import coerce_bool, coerce_int, normalize_optional_text
|
||||
from shelfmark.core.settings_registry import (
|
||||
ActionButton,
|
||||
CheckboxField,
|
||||
@@ -31,6 +31,7 @@ from shelfmark.metadata_providers import (
|
||||
MetadataCapability,
|
||||
MetadataProvider,
|
||||
MetadataSearchOptions,
|
||||
SearchField,
|
||||
SearchResult,
|
||||
SearchType,
|
||||
SortOrder,
|
||||
@@ -546,6 +547,12 @@ def _normalize_series_position(value: Any) -> float | None:
|
||||
return None
|
||||
|
||||
|
||||
def _normalize_hardcover_api_key(value: object) -> str:
|
||||
"""Normalize Hardcover API keys, stripping copied auth-header prefixes."""
|
||||
normalized_value = normalize_optional_text(value) or ""
|
||||
return normalized_value.removeprefix("Bearer ").strip()
|
||||
|
||||
|
||||
def _normalize_search_text(value: str) -> str:
|
||||
"""Normalize free-text search input for matching and caching."""
|
||||
return " ".join(value.split()).strip()
|
||||
@@ -856,7 +863,7 @@ class HardcoverProvider(MetadataProvider):
|
||||
sort=SortOrder.SERIES_ORDER,
|
||||
),
|
||||
)
|
||||
search_fields: ClassVar[tuple[TextSearchField | DynamicSelectSearchField, ...]] = (
|
||||
search_fields: ClassVar[tuple[SearchField, ...]] = (
|
||||
TextSearchField(
|
||||
key="author",
|
||||
label="Author",
|
||||
@@ -888,8 +895,7 @@ class HardcoverProvider(MetadataProvider):
|
||||
def __init__(self, api_key: str | None = None) -> None:
|
||||
"""Initialize provider with optional API key (falls back to config)."""
|
||||
raw_key = api_key or app_config.get("HARDCOVER_API_KEY", "")
|
||||
# Strip "Bearer " prefix if user pasted the full auth header from Hardcover
|
||||
self.api_key = raw_key.removeprefix("Bearer ").strip() if raw_key else ""
|
||||
self.api_key = _normalize_hardcover_api_key(raw_key)
|
||||
self.session = requests.Session()
|
||||
if self.api_key:
|
||||
self.session.headers.update(
|
||||
@@ -1049,10 +1055,8 @@ class HardcoverProvider(MetadataProvider):
|
||||
if not selected:
|
||||
return SearchResult(books=[], page=page, total_found=0, has_more=False)
|
||||
|
||||
list_id_raw = selected.get("id")
|
||||
try:
|
||||
list_id = int(list_id_raw)
|
||||
except TypeError, ValueError:
|
||||
list_id = coerce_int(selected.get("id"), 0)
|
||||
if list_id < 1:
|
||||
return SearchResult(books=[], page=page, total_found=0, has_more=False)
|
||||
|
||||
return self._fetch_list_books_by_id(list_id, page, limit)
|
||||
@@ -1164,9 +1168,8 @@ class HardcoverProvider(MetadataProvider):
|
||||
if not _query_matches_author_name(query, author_name):
|
||||
continue
|
||||
|
||||
try:
|
||||
author_id = int(item.get("id"))
|
||||
except TypeError, ValueError:
|
||||
author_id = coerce_int(item.get("id"), 0)
|
||||
if author_id < 1:
|
||||
continue
|
||||
|
||||
if author_id not in author_ids:
|
||||
@@ -1229,8 +1232,14 @@ class HardcoverProvider(MetadataProvider):
|
||||
weights=TITLE_SUGGESTION_WEIGHTS,
|
||||
)
|
||||
|
||||
exclude_compilations = app_config.get("HARDCOVER_EXCLUDE_COMPILATIONS", False)
|
||||
exclude_unreleased = app_config.get("HARDCOVER_EXCLUDE_UNRELEASED", False)
|
||||
exclude_compilations = coerce_bool(
|
||||
app_config.get("HARDCOVER_EXCLUDE_COMPILATIONS", False),
|
||||
default=False,
|
||||
)
|
||||
exclude_unreleased = coerce_bool(
|
||||
app_config.get("HARDCOVER_EXCLUDE_UNRELEASED", False),
|
||||
default=False,
|
||||
)
|
||||
current_year = datetime.now(UTC).year
|
||||
|
||||
options: list[dict[str, str]] = []
|
||||
@@ -1375,9 +1384,8 @@ class HardcoverProvider(MetadataProvider):
|
||||
item = _unwrap_hit_document(hit)
|
||||
if item is None:
|
||||
continue
|
||||
try:
|
||||
series_id = int(item.get("id"))
|
||||
except TypeError, ValueError:
|
||||
series_id = coerce_int(item.get("id"), 0)
|
||||
if series_id < 1:
|
||||
continue
|
||||
name = str(item.get("name") or "").strip()
|
||||
if not name:
|
||||
@@ -1715,7 +1723,7 @@ class HardcoverProvider(MetadataProvider):
|
||||
raise ValueError(msg)
|
||||
|
||||
state = self._fetch_book_target_state(book_id_int)
|
||||
options = [
|
||||
options: list[dict[str, Any]] = [
|
||||
dict(option)
|
||||
for option in self.get_user_lists()
|
||||
if option.get("group") in HARDCOVER_WRITABLE_TARGET_GROUPS
|
||||
@@ -1919,7 +1927,7 @@ class HardcoverProvider(MetadataProvider):
|
||||
return {bid: [] for bid in book_ids}
|
||||
|
||||
states = self._fetch_book_target_states_batch(int_ids)
|
||||
writable_options = [
|
||||
writable_options: list[dict[str, Any]] = [
|
||||
dict(option)
|
||||
for option in self.get_user_lists()
|
||||
if option.get("group") in HARDCOVER_WRITABLE_TARGET_GROUPS
|
||||
@@ -2136,8 +2144,14 @@ class HardcoverProvider(MetadataProvider):
|
||||
resolved_series = self._resolve_series_search_value(series_value_from_field)
|
||||
if not resolved_series:
|
||||
return SearchResult(books=[], page=options.page, total_found=0, has_more=False)
|
||||
exclude_compilations = app_config.get("HARDCOVER_EXCLUDE_COMPILATIONS", False)
|
||||
exclude_unreleased = app_config.get("HARDCOVER_EXCLUDE_UNRELEASED", False)
|
||||
exclude_compilations = coerce_bool(
|
||||
app_config.get("HARDCOVER_EXCLUDE_COMPILATIONS", False),
|
||||
default=False,
|
||||
)
|
||||
exclude_unreleased = coerce_bool(
|
||||
app_config.get("HARDCOVER_EXCLUDE_UNRELEASED", False),
|
||||
default=False,
|
||||
)
|
||||
return self._fetch_series_books_by_id(
|
||||
int(resolved_series["id"]),
|
||||
options.page,
|
||||
@@ -2154,8 +2168,14 @@ class HardcoverProvider(MetadataProvider):
|
||||
|
||||
# Build cache key from options (include fields and settings for cache differentiation)
|
||||
fields_key = ":".join(f"{k}={v}" for k, v in sorted(options.fields.items()))
|
||||
exclude_compilations = app_config.get("HARDCOVER_EXCLUDE_COMPILATIONS", False)
|
||||
exclude_unreleased = app_config.get("HARDCOVER_EXCLUDE_UNRELEASED", False)
|
||||
exclude_compilations = coerce_bool(
|
||||
app_config.get("HARDCOVER_EXCLUDE_COMPILATIONS", False),
|
||||
default=False,
|
||||
)
|
||||
exclude_unreleased = coerce_bool(
|
||||
app_config.get("HARDCOVER_EXCLUDE_UNRELEASED", False),
|
||||
default=False,
|
||||
)
|
||||
cache_key = f"{options.query}:{options.search_type.value}:{options.sort.value}:{options.limit}:{options.page}:{fields_key}:excl_comp={exclude_compilations}:excl_unrel={exclude_unreleased}"
|
||||
return self._search_cached(cache_key, options)
|
||||
|
||||
@@ -2214,8 +2234,14 @@ class HardcoverProvider(MetadataProvider):
|
||||
hits, found_count = _extract_typesense_hits(result)
|
||||
|
||||
# Parse hits, filtering compilations and unreleased books if enabled
|
||||
exclude_compilations = app_config.get("HARDCOVER_EXCLUDE_COMPILATIONS", False)
|
||||
exclude_unreleased = app_config.get("HARDCOVER_EXCLUDE_UNRELEASED", False)
|
||||
exclude_compilations = coerce_bool(
|
||||
app_config.get("HARDCOVER_EXCLUDE_COMPILATIONS", False),
|
||||
default=False,
|
||||
)
|
||||
exclude_unreleased = coerce_bool(
|
||||
app_config.get("HARDCOVER_EXCLUDE_UNRELEASED", False),
|
||||
default=False,
|
||||
)
|
||||
current_year = datetime.now(UTC).year
|
||||
books = []
|
||||
for hit in hits:
|
||||
@@ -2720,8 +2746,7 @@ def _test_hardcover_connection(current_values: dict[str, Any] | None = None) ->
|
||||
|
||||
# Use current form values first, fall back to saved config
|
||||
raw_key = current_values.get("HARDCOVER_API_KEY") or app_config.get("HARDCOVER_API_KEY", "")
|
||||
# Strip "Bearer " prefix if user pasted the full auth header from Hardcover
|
||||
api_key = raw_key.removeprefix("Bearer ").strip() if raw_key else ""
|
||||
api_key = _normalize_hardcover_api_key(raw_key)
|
||||
|
||||
key_len = len(api_key) if api_key else 0
|
||||
logger.debug("Hardcover test: key length=%s", key_len)
|
||||
|
||||
@@ -25,6 +25,7 @@ from shelfmark.metadata_providers import (
|
||||
DisplayField,
|
||||
MetadataProvider,
|
||||
MetadataSearchOptions,
|
||||
SearchField,
|
||||
SearchType,
|
||||
SortOrder,
|
||||
TextSearchField,
|
||||
@@ -114,7 +115,7 @@ class OpenLibraryProvider(MetadataProvider):
|
||||
SortOrder.NEWEST,
|
||||
SortOrder.OLDEST,
|
||||
)
|
||||
search_fields: ClassVar[tuple[TextSearchField, ...]] = (
|
||||
search_fields: ClassVar[tuple[SearchField, ...]] = (
|
||||
TextSearchField(
|
||||
key="author",
|
||||
label="Author",
|
||||
|
||||
@@ -26,6 +26,12 @@ if TYPE_CHECKING:
|
||||
logger = setup_logger(__name__)
|
||||
|
||||
|
||||
def _resolve_configured_hostname() -> str:
|
||||
"""Return a normalized ABB hostname from config when available."""
|
||||
configured_hostname = config.get("ABB_HOSTNAME", "")
|
||||
return normalize_hostname(configured_hostname if isinstance(configured_hostname, str) else "")
|
||||
|
||||
|
||||
@register_handler("audiobookbay")
|
||||
class AudiobookBayHandler(ExternalClientHandler):
|
||||
"""Handler for AudiobookBay downloads via configured torrent client."""
|
||||
@@ -63,7 +69,7 @@ class AudiobookBayHandler(ExternalClientHandler):
|
||||
logger.warning("Missing details URL for AudiobookBay task: %s", task.task_id)
|
||||
return None
|
||||
|
||||
hostname = normalize_hostname(config.get("ABB_HOSTNAME", ""))
|
||||
hostname = _resolve_configured_hostname()
|
||||
if not hostname:
|
||||
hostname = normalize_hostname(urlparse(detail_url).hostname)
|
||||
|
||||
|
||||
@@ -40,6 +40,30 @@ SIZE_PATTERN = re.compile(r"File Size:\s*([\d.]+)\s*([A-Za-z]+)")
|
||||
INFO_HASH_LABEL_PATTERN = re.compile(r"Info Hash", re.IGNORECASE)
|
||||
|
||||
|
||||
def _coerce_non_negative_float(value: object, default: float) -> float:
|
||||
"""Return a non-negative float config value or the provided default."""
|
||||
if isinstance(value, bool):
|
||||
return default
|
||||
if isinstance(value, int | float) and value >= 0:
|
||||
return float(value)
|
||||
return default
|
||||
|
||||
|
||||
def _coerce_markup_to_html(value: str | tuple[str, str]) -> str:
|
||||
"""Normalize downloader output to the HTML markup string."""
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
html, _response_url = value
|
||||
return html
|
||||
|
||||
|
||||
def _coerce_attribute_to_str(value: object) -> str:
|
||||
"""Return a plain string HTML attribute value, or an empty string."""
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
return ""
|
||||
|
||||
|
||||
def _build_search_url(
|
||||
hostname: str,
|
||||
page: int,
|
||||
@@ -129,7 +153,7 @@ def search_audiobookbay(
|
||||
|
||||
"""
|
||||
results = []
|
||||
rate_limit_delay = config.get("ABB_RATE_LIMIT_DELAY", 1.0)
|
||||
rate_limit_delay = _coerce_non_negative_float(config.get("ABB_RATE_LIMIT_DELAY", 1.0), 1.0)
|
||||
session = requests.Session()
|
||||
|
||||
# Bootstrap ABB session cookie (PHPSESSID). ABB increasingly serves reliable
|
||||
@@ -221,7 +245,7 @@ def search_audiobookbay(
|
||||
title = title_elem.text.strip()
|
||||
|
||||
# Extract link (relative, needs hostname prefix)
|
||||
href = title_elem.get("href", "")
|
||||
href = _coerce_attribute_to_str(title_elem.get("href", ""))
|
||||
if not href:
|
||||
continue
|
||||
|
||||
@@ -235,7 +259,13 @@ def search_audiobookbay(
|
||||
"img"
|
||||
)
|
||||
if cover_elem:
|
||||
cover = _normalize_result_url(cover_elem.get("src", ""), hostname) or None
|
||||
cover = (
|
||||
_normalize_result_url(
|
||||
_coerce_attribute_to_str(cover_elem.get("src", "")),
|
||||
hostname,
|
||||
)
|
||||
or None
|
||||
)
|
||||
|
||||
# Extract language from .postInfo
|
||||
language = None
|
||||
@@ -328,19 +358,8 @@ def extract_magnet_link(details_url: str, hostname: str = "audiobookbay.lu") ->
|
||||
_bootstrap_abb_session(hostname, session, DETAIL_PAGE_RETRY_ATTEMPTS)
|
||||
|
||||
# Fetch detail page
|
||||
detail_html = downloader.html_get_page(
|
||||
details_url,
|
||||
retry=DETAIL_PAGE_RETRY_ATTEMPTS,
|
||||
use_bypasser=False,
|
||||
allow_bypasser_fallback=False,
|
||||
success_delay=0,
|
||||
session=session,
|
||||
)
|
||||
|
||||
if not detail_html:
|
||||
session = requests.Session()
|
||||
_bootstrap_abb_session(hostname, session, DETAIL_PAGE_RETRY_ATTEMPTS)
|
||||
detail_html = downloader.html_get_page(
|
||||
detail_html = _coerce_markup_to_html(
|
||||
downloader.html_get_page(
|
||||
details_url,
|
||||
retry=DETAIL_PAGE_RETRY_ATTEMPTS,
|
||||
use_bypasser=False,
|
||||
@@ -348,6 +367,21 @@ def extract_magnet_link(details_url: str, hostname: str = "audiobookbay.lu") ->
|
||||
success_delay=0,
|
||||
session=session,
|
||||
)
|
||||
)
|
||||
|
||||
if not detail_html:
|
||||
session = requests.Session()
|
||||
_bootstrap_abb_session(hostname, session, DETAIL_PAGE_RETRY_ATTEMPTS)
|
||||
detail_html = _coerce_markup_to_html(
|
||||
downloader.html_get_page(
|
||||
details_url,
|
||||
retry=DETAIL_PAGE_RETRY_ATTEMPTS,
|
||||
use_bypasser=False,
|
||||
allow_bypasser_fallback=False,
|
||||
success_delay=0,
|
||||
session=session,
|
||||
)
|
||||
)
|
||||
|
||||
if not detail_html:
|
||||
logger.warning("Failed to fetch details page")
|
||||
|
||||
@@ -28,6 +28,20 @@ logger = setup_logger(__name__)
|
||||
MIN_RELEVANCE_QUERY_WORD_LENGTH = 2
|
||||
|
||||
|
||||
def _coerce_hostname_config(value: object) -> str:
|
||||
"""Return a normalized ABB hostname from config."""
|
||||
return normalize_hostname(value if isinstance(value, str) else "")
|
||||
|
||||
|
||||
def _coerce_positive_int(value: object, default: int) -> int:
|
||||
"""Return a positive integer config value or the provided default."""
|
||||
if isinstance(value, bool):
|
||||
return default
|
||||
if isinstance(value, int) and value > 0:
|
||||
return value
|
||||
return default
|
||||
|
||||
|
||||
# Map language names to ISO 639-1 codes (matching frontend color maps)
|
||||
LANGUAGE_MAP = {
|
||||
"english": "en",
|
||||
@@ -169,11 +183,11 @@ class AudiobookBaySource(ReleaseSource):
|
||||
if content_type != "audiobook":
|
||||
return []
|
||||
|
||||
hostname = normalize_hostname(config.get("ABB_HOSTNAME", ""))
|
||||
hostname = _coerce_hostname_config(config.get("ABB_HOSTNAME", ""))
|
||||
if not hostname:
|
||||
logger.debug("AudiobookBay hostname is not configured")
|
||||
return []
|
||||
max_pages = config.get("ABB_PAGE_LIMIT", 1)
|
||||
max_pages = _coerce_positive_int(config.get("ABB_PAGE_LIMIT", 1), 1)
|
||||
exact_phrase = bool(config.get("ABB_EXACT_PHRASE", False))
|
||||
|
||||
# Build search query candidates from plan.
|
||||
@@ -321,7 +335,7 @@ class AudiobookBaySource(ReleaseSource):
|
||||
def is_available(self) -> bool:
|
||||
"""Check if AudiobookBay source is enabled and configured."""
|
||||
return config.get("ABB_ENABLED", False) is True and bool(
|
||||
normalize_hostname(config.get("ABB_HOSTNAME", ""))
|
||||
_coerce_hostname_config(config.get("ABB_HOSTNAME", ""))
|
||||
)
|
||||
|
||||
def get_column_config(self) -> ReleaseColumnConfig:
|
||||
|
||||
@@ -6,11 +6,12 @@ import re
|
||||
import time
|
||||
from dataclasses import replace
|
||||
from http import HTTPStatus
|
||||
from typing import TYPE_CHECKING, ClassVar, NoReturn
|
||||
from typing import TYPE_CHECKING, ClassVar, NoReturn, TypedDict
|
||||
from urllib.parse import quote
|
||||
|
||||
import requests
|
||||
from bs4 import BeautifulSoup, NavigableString, Tag
|
||||
from bs4 import BeautifulSoup, Tag
|
||||
from bs4.element import NavigableString
|
||||
|
||||
from shelfmark.config.env import DEBUG_SKIP_SOURCES, TMP_DIR
|
||||
from shelfmark.core.config import config
|
||||
@@ -37,7 +38,7 @@ from shelfmark.release_sources import (
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Iterable
|
||||
from pathlib import Path
|
||||
from threading import Event
|
||||
|
||||
@@ -47,10 +48,135 @@ if TYPE_CHECKING:
|
||||
logger = setup_logger(__name__)
|
||||
|
||||
|
||||
class SourcePriorityEntry(TypedDict):
|
||||
"""Normalized source priority entry from config."""
|
||||
|
||||
id: str
|
||||
enabled: bool
|
||||
|
||||
|
||||
def _raise_runtime_error(message: str) -> NoReturn:
|
||||
raise RuntimeError(message)
|
||||
|
||||
|
||||
def _coerce_str_list(value: object) -> list[str]:
|
||||
"""Return only string items from a config value."""
|
||||
if not isinstance(value, list | tuple):
|
||||
return []
|
||||
return [item for item in value if isinstance(item, str)]
|
||||
|
||||
|
||||
def _get_supported_formats() -> list[str]:
|
||||
"""Return configured supported formats as a clean string list."""
|
||||
return _coerce_str_list(config.SUPPORTED_FORMATS)
|
||||
|
||||
|
||||
def _parse_source_priority_entries(
|
||||
value: object,
|
||||
*,
|
||||
allowed_ids: set[str] | None = None,
|
||||
excluded_ids: set[str] | None = None,
|
||||
) -> list[SourcePriorityEntry]:
|
||||
"""Normalize orderable-list config values into typed source entries."""
|
||||
if not isinstance(value, list):
|
||||
return []
|
||||
|
||||
entries: list[SourcePriorityEntry] = []
|
||||
for item in value:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
|
||||
source_id = item.get("id")
|
||||
if not isinstance(source_id, str):
|
||||
continue
|
||||
|
||||
if allowed_ids is not None and source_id not in allowed_ids:
|
||||
continue
|
||||
if excluded_ids is not None and source_id in excluded_ids:
|
||||
continue
|
||||
|
||||
entries.append({"id": source_id, "enabled": bool(item.get("enabled", True))})
|
||||
|
||||
return entries
|
||||
|
||||
|
||||
def _html_response_text(response: str | tuple[str, str]) -> str:
|
||||
"""Extract the HTML body from downloader responses."""
|
||||
if isinstance(response, tuple):
|
||||
return response[0]
|
||||
return response
|
||||
|
||||
|
||||
def _attr_to_str(value: object) -> str | None:
|
||||
"""Convert a BeautifulSoup attribute value to a plain string."""
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, list):
|
||||
for item in value:
|
||||
if isinstance(item, str):
|
||||
return item
|
||||
return None
|
||||
|
||||
|
||||
def _get_attr(tag: Tag, attr: str) -> str | None:
|
||||
"""Safely fetch a tag attribute as a string."""
|
||||
return _attr_to_str(tag.get(attr))
|
||||
|
||||
|
||||
def _first_stripped_text(tag: Tag | None) -> str | None:
|
||||
"""Return the first non-empty stripped string from a tag."""
|
||||
if tag is None:
|
||||
return None
|
||||
|
||||
for text in tag.stripped_strings:
|
||||
return text
|
||||
return None
|
||||
|
||||
|
||||
def _iter_child_tags(tag: Tag) -> Iterable[Tag]:
|
||||
"""Iterate only over child tags, skipping text nodes."""
|
||||
for child in tag.children:
|
||||
if isinstance(child, Tag):
|
||||
yield child
|
||||
|
||||
|
||||
def _find_first_anchor_with_text(
|
||||
container: BeautifulSoup | Tag,
|
||||
text: str,
|
||||
*,
|
||||
contains: bool = False,
|
||||
) -> Tag | None:
|
||||
"""Find the first anchor whose text matches the requested value."""
|
||||
expected = text.lower()
|
||||
for anchor in container.find_all("a", href=True):
|
||||
anchor_text = anchor.get_text(strip=True)
|
||||
if not anchor_text:
|
||||
continue
|
||||
candidate = anchor_text.lower()
|
||||
if candidate == expected or (contains and expected in candidate):
|
||||
return anchor
|
||||
return None
|
||||
|
||||
|
||||
def _find_text_node(container: BeautifulSoup | Tag, needle: str) -> NavigableString | None:
|
||||
"""Find a text node containing a case-insensitive substring."""
|
||||
expected = needle.lower()
|
||||
for text_node in container.find_all(string=True):
|
||||
if isinstance(text_node, NavigableString) and expected in text_node.strip().lower():
|
||||
return text_node
|
||||
return None
|
||||
|
||||
|
||||
def _tag_has_class_containing(tag: Tag, needle: str) -> bool:
|
||||
"""Check whether a tag has a CSS class containing a substring."""
|
||||
class_values = tag.get("class")
|
||||
if isinstance(class_values, str):
|
||||
return needle in class_values
|
||||
if isinstance(class_values, list):
|
||||
return any(isinstance(value, str) and needle in value for value in class_values)
|
||||
return False
|
||||
|
||||
|
||||
_aa_slow_rotation = itertools.count()
|
||||
_url_source_types: dict[str, str] = {}
|
||||
|
||||
@@ -114,34 +240,26 @@ _LIBGEN_GET_PATTERNS = [
|
||||
]
|
||||
|
||||
|
||||
def _get_source_priority() -> list[dict]:
|
||||
def _get_source_priority() -> list[SourcePriorityEntry]:
|
||||
"""Get the full source priority list.
|
||||
|
||||
Fast sources come from user config (FAST_SOURCES_DISPLAY).
|
||||
Slow sources come from user config.
|
||||
"""
|
||||
# Fast sources - always first, configurable via settings/env
|
||||
fast_sources: list[dict] = []
|
||||
configured_fast = config.get("FAST_SOURCES_DISPLAY") or []
|
||||
fast_sources = _parse_source_priority_entries(
|
||||
config.get("FAST_SOURCES_DISPLAY"),
|
||||
allowed_ids={"aa-fast", "libgen"},
|
||||
)
|
||||
has_donator_key = bool(config.get("AA_DONATOR_KEY"))
|
||||
|
||||
if isinstance(configured_fast, list):
|
||||
for item in configured_fast:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
source_id = item.get("id")
|
||||
if source_id not in ("aa-fast", "libgen"):
|
||||
continue
|
||||
enabled = bool(item.get("enabled", True))
|
||||
if source_id == "aa-fast" and not has_donator_key:
|
||||
enabled = False
|
||||
fast_sources.append({"id": source_id, "enabled": enabled})
|
||||
for source in fast_sources:
|
||||
if source["id"] == "aa-fast" and not has_donator_key:
|
||||
source["enabled"] = False
|
||||
|
||||
# User's configured slow sources (config won't contain fast sources)
|
||||
slow_sources = config.get("SOURCE_PRIORITY") or []
|
||||
|
||||
# Filter out any legacy fast source entries from old configs
|
||||
slow_sources = [s for s in slow_sources if s["id"] not in ("aa-fast", "libgen")]
|
||||
slow_sources = _parse_source_priority_entries(
|
||||
config.get("SOURCE_PRIORITY"),
|
||||
excluded_ids={"aa-fast", "libgen"},
|
||||
)
|
||||
|
||||
return fast_sources + slow_sources
|
||||
|
||||
@@ -203,7 +321,7 @@ def search_books(query: str, filters: SearchFilters) -> list[BrowseRecord]:
|
||||
for value in filters.content:
|
||||
filters_query += f"&content={quote(value)}"
|
||||
|
||||
formats_to_use = filters.format or config.SUPPORTED_FORMATS
|
||||
formats_to_use = filters.format or _get_supported_formats()
|
||||
|
||||
index = 1
|
||||
for filter_type, filter_values in vars(filters).items():
|
||||
@@ -233,26 +351,30 @@ def search_books(query: str, filters: SearchFilters) -> list[BrowseRecord]:
|
||||
logger.info("No books found for query: %s", query)
|
||||
return []
|
||||
|
||||
soup = BeautifulSoup(html, "html.parser")
|
||||
tbody: Tag | NavigableString | None = soup.find("table")
|
||||
soup = BeautifulSoup(_html_response_text(html), "html.parser")
|
||||
tbody = soup.find("table")
|
||||
|
||||
if not tbody:
|
||||
if tbody is None:
|
||||
logger.warning("No results table found for query: %s", query)
|
||||
msg = "No books found. Please try another query."
|
||||
raise RuntimeError(msg)
|
||||
if not isinstance(tbody, Tag):
|
||||
msg = f"Expected results table tag, got {type(tbody).__name__}"
|
||||
raise TypeError(msg)
|
||||
|
||||
books = []
|
||||
if isinstance(tbody, Tag):
|
||||
for line_tr in tbody.find_all("tr"):
|
||||
book = _parse_search_result_row(line_tr)
|
||||
if book:
|
||||
books.append(book)
|
||||
for line_tr in tbody.find_all("tr"):
|
||||
book = _parse_search_result_row(line_tr)
|
||||
if book:
|
||||
books.append(book)
|
||||
|
||||
supported_formats = _get_supported_formats()
|
||||
|
||||
books.sort(
|
||||
key=lambda x: (
|
||||
config.SUPPORTED_FORMATS.index(x.format)
|
||||
if x.format in config.SUPPORTED_FORMATS
|
||||
else len(config.SUPPORTED_FORMATS)
|
||||
supported_formats.index(x.format)
|
||||
if x.format in supported_formats
|
||||
else len(supported_formats)
|
||||
)
|
||||
)
|
||||
|
||||
@@ -279,7 +401,7 @@ def get_book_info(book_id: str, *, fetch_download_count: bool = True) -> BrowseR
|
||||
msg = "Unable to reach download source. Network restricted or mirrors are blocked."
|
||||
raise SearchUnavailableError(msg)
|
||||
|
||||
soup = BeautifulSoup(html, "html.parser")
|
||||
soup = BeautifulSoup(_html_response_text(html), "html.parser")
|
||||
|
||||
return _parse_book_info_page(soup, book_id, fetch_download_count=fetch_download_count)
|
||||
|
||||
@@ -289,22 +411,52 @@ def _parse_search_result_row(row: Tag) -> BrowseRecord | None:
|
||||
try:
|
||||
if row.text.strip().lower().startswith("your ad here"):
|
||||
return None
|
||||
|
||||
cells = row.find_all("td")
|
||||
anchors = row.find_all("a", href=True)
|
||||
if len(cells) < 11 or not anchors:
|
||||
return None
|
||||
|
||||
record_id = (_get_attr(anchors[0], "href") or "").split("/")[-1]
|
||||
if not record_id:
|
||||
return None
|
||||
|
||||
preview_img = cells[0].find("img")
|
||||
preview = preview_img["src"] if preview_img else None
|
||||
preview = _get_attr(preview_img, "src") if isinstance(preview_img, Tag) else None
|
||||
|
||||
title = _first_stripped_text(cells[1].find("span"))
|
||||
author = _first_stripped_text(cells[2].find("span"))
|
||||
publisher = _first_stripped_text(cells[3].find("span"))
|
||||
year = _first_stripped_text(cells[4].find("span"))
|
||||
language = _first_stripped_text(cells[7].find("span"))
|
||||
content = _first_stripped_text(cells[8].find("span"))
|
||||
file_format = _first_stripped_text(cells[9].find("span"))
|
||||
size = _first_stripped_text(cells[10].find("span"))
|
||||
|
||||
if (
|
||||
title is None
|
||||
or author is None
|
||||
or publisher is None
|
||||
or year is None
|
||||
or language is None
|
||||
or content is None
|
||||
or file_format is None
|
||||
or size is None
|
||||
):
|
||||
return None
|
||||
|
||||
return BrowseRecord(
|
||||
id=row.find_all("a")[0]["href"].split("/")[-1],
|
||||
title=cells[1].find("span").next,
|
||||
id=record_id,
|
||||
title=title,
|
||||
source="direct_download",
|
||||
preview=preview,
|
||||
author=cells[2].find("span").next,
|
||||
publisher=cells[3].find("span").next,
|
||||
year=cells[4].find("span").next,
|
||||
language=cells[7].find("span").next,
|
||||
content=cells[8].find("span").next.lower(),
|
||||
format=cells[9].find("span").next.lower(),
|
||||
size=cells[10].find("span").next,
|
||||
author=author,
|
||||
publisher=publisher,
|
||||
year=year,
|
||||
language=language,
|
||||
content=content.lower() if content else None,
|
||||
format=file_format.lower() if file_format else None,
|
||||
size=size,
|
||||
)
|
||||
except (AttributeError, IndexError, KeyError, TypeError) as e:
|
||||
logger.error_trace(f"Error parsing search result row: {e}")
|
||||
@@ -327,12 +479,24 @@ def _parse_book_info_page(
|
||||
preview: str = ""
|
||||
|
||||
node = data.select_one("div:nth-of-type(1) > img")
|
||||
if node:
|
||||
preview_value = node.get("src", "")
|
||||
preview = preview_value[0] if isinstance(preview_value, list) else preview_value
|
||||
if isinstance(node, Tag):
|
||||
preview = _get_attr(node, "src") or ""
|
||||
|
||||
data = soup.find_all("div", {"class": "main-inner"})[0].find_next("div")
|
||||
divs = list(data.children)
|
||||
main_inner = next(
|
||||
(tag for tag in soup.find_all("div", {"class": "main-inner"}) if isinstance(tag, Tag)),
|
||||
None,
|
||||
)
|
||||
if main_inner is None:
|
||||
msg = f"Failed to parse book details for ID: {book_id}"
|
||||
raise RuntimeError(msg)
|
||||
|
||||
details_container = main_inner.find_next("div")
|
||||
if not isinstance(details_container, Tag):
|
||||
msg = f"Expected details container tag for book ID {book_id}, got {type(details_container).__name__}"
|
||||
raise TypeError(msg)
|
||||
|
||||
original_nodes = list(details_container.children)
|
||||
divs = [node for node in original_nodes if isinstance(node, Tag)]
|
||||
|
||||
slow_urls_no_waitlist: set[str] = set()
|
||||
slow_urls_with_waitlist: set[str] = set()
|
||||
@@ -340,13 +504,20 @@ def _parse_book_info_page(
|
||||
for anchor in soup.find_all("a"):
|
||||
try:
|
||||
text = anchor.text.strip().lower()
|
||||
href = anchor.get("href", "")
|
||||
href = _get_attr(anchor, "href")
|
||||
if not href:
|
||||
continue
|
||||
|
||||
next_text = ""
|
||||
if anchor.next and anchor.next.next:
|
||||
next_text = getattr(anchor.next.next, "text", str(anchor.next.next)).strip().lower()
|
||||
next_elements = anchor.next_elements
|
||||
next(next_elements, None)
|
||||
second_next = next(next_elements, None)
|
||||
if second_next is not None:
|
||||
next_text = (
|
||||
second_next.get_text(strip=True).lower()
|
||||
if isinstance(second_next, Tag)
|
||||
else str(second_next).strip().lower()
|
||||
)
|
||||
|
||||
if text.startswith("slow partner server") and "waitlist" in next_text:
|
||||
if "no waitlist" in next_text:
|
||||
@@ -379,18 +550,19 @@ def _parse_book_info_page(
|
||||
urls.append(abs_url)
|
||||
_url_source_types[abs_url] = "aa-slow-wait"
|
||||
|
||||
original_divs = divs
|
||||
divs = [div for div in divs if div.text.strip() != ""]
|
||||
divs = [div for div in divs if div.get_text(strip=True)]
|
||||
|
||||
all_details = _find_in_divs(divs, " · ")
|
||||
file_format = ""
|
||||
size = ""
|
||||
content = ""
|
||||
supported_formats = _get_supported_formats()
|
||||
|
||||
for _details in all_details:
|
||||
_details = _details.split(" · ")
|
||||
for f in _details:
|
||||
if file_format == "" and f.strip().lower() in config.SUPPORTED_FORMATS:
|
||||
stripped_lower = f.strip().lower()
|
||||
if file_format == "" and stripped_lower in supported_formats:
|
||||
file_format = f.strip().lower()
|
||||
if size == "" and any(u in f.strip().lower() for u in ("mb", "kb", "gb")):
|
||||
size = _normalize_size(f)
|
||||
@@ -427,7 +599,11 @@ def _parse_book_info_page(
|
||||
)
|
||||
|
||||
# Extract additional metadata
|
||||
info = _extract_book_metadata(original_divs[-6])
|
||||
metadata_node = original_nodes[-6]
|
||||
if not isinstance(metadata_node, Tag):
|
||||
msg = f"Expected metadata container tag for book ID {book_id}, got {type(metadata_node).__name__}"
|
||||
raise TypeError(msg)
|
||||
info = _extract_book_metadata(metadata_node)
|
||||
|
||||
if fetch_download_count:
|
||||
try:
|
||||
@@ -436,7 +612,7 @@ def _parse_book_info_page(
|
||||
summary_url, selector=network.AAMirrorSelector(), allow_bypasser_fallback=False
|
||||
)
|
||||
if summary_response:
|
||||
summary_data = json.loads(summary_response)
|
||||
summary_data = json.loads(_html_response_text(summary_response))
|
||||
if "downloads_total" in summary_data:
|
||||
info["Downloads"] = [str(summary_data["downloads_total"])]
|
||||
except (
|
||||
@@ -463,9 +639,9 @@ def _parse_book_info_page(
|
||||
return book_info
|
||||
|
||||
|
||||
def _find_in_divs(divs: list, text: str, *, is_class: bool = False) -> list[str]:
|
||||
def _find_in_divs(divs: list[Tag], text: str, *, is_class: bool = False) -> list[str]:
|
||||
"""Find divs containing text or having a specific class."""
|
||||
results = []
|
||||
results: list[str] = []
|
||||
for div in divs:
|
||||
if is_class:
|
||||
if div.find(class_=text):
|
||||
@@ -527,12 +703,12 @@ def _extract_book_metadata(metadata_divs: Tag) -> dict[str, list[str]]:
|
||||
info: dict[str, set[str]] = {}
|
||||
|
||||
sub_datas = metadata_divs.find_all("div")[0]
|
||||
for sub_data in sub_datas.children:
|
||||
if sub_data.text.strip() == "":
|
||||
for sub_data in _iter_child_tags(sub_datas):
|
||||
if sub_data.get_text(strip=True) == "":
|
||||
continue
|
||||
children = list(sub_data.children)
|
||||
key = children[0].text.strip()
|
||||
value = children[1].text.strip()
|
||||
children = list(_iter_child_tags(sub_data))
|
||||
key = children[0].get_text(strip=True)
|
||||
value = children[1].get_text(strip=True)
|
||||
if key not in info:
|
||||
info[key] = set()
|
||||
info[key].add(value)
|
||||
@@ -755,11 +931,11 @@ def _get_download_urls_from_welib(
|
||||
logger.warning("Welib page empty for %s", book_id)
|
||||
return []
|
||||
|
||||
soup = BeautifulSoup(html, "html.parser")
|
||||
soup = BeautifulSoup(_html_response_text(html), "html.parser")
|
||||
links = [
|
||||
downloader.get_absolute_url(url, a["href"])
|
||||
downloader.get_absolute_url(url, href)
|
||||
for a in soup.find_all("a", href=True)
|
||||
if "/slow_download/" in a["href"]
|
||||
if (href := _get_attr(a, "href")) and "/slow_download/" in href
|
||||
]
|
||||
return list(dict.fromkeys(links)) # Dedupe while preserving order
|
||||
|
||||
@@ -949,7 +1125,11 @@ def _get_download_url(
|
||||
page = downloader.html_get_page(
|
||||
link, selector=sel, cancel_flag=cancel_flag, status_callback=status_callback
|
||||
)
|
||||
return downloader.get_absolute_url(link, json.loads(page).get("download_url", ""))
|
||||
page_data = json.loads(_html_response_text(page))
|
||||
download_url = page_data.get("download_url", "")
|
||||
return (
|
||||
downloader.get_absolute_url(link, download_url) if isinstance(download_url, str) else ""
|
||||
)
|
||||
|
||||
if "/ads.php?md5=" in link and any(domain in link for domain in _get_libgen_domains()):
|
||||
return _extract_libgen_download_url(link, cancel_flag)
|
||||
@@ -960,7 +1140,7 @@ def _get_download_url(
|
||||
if not html:
|
||||
return ""
|
||||
|
||||
soup = BeautifulSoup(html, "html.parser")
|
||||
soup = BeautifulSoup(_html_response_text(html), "html.parser")
|
||||
url = ""
|
||||
|
||||
# Z-Library
|
||||
@@ -973,9 +1153,9 @@ def _get_download_url(
|
||||
link, selector=sel, cancel_flag=cancel_flag, status_callback=status_callback
|
||||
)
|
||||
if html:
|
||||
soup = BeautifulSoup(html, "html.parser")
|
||||
soup = BeautifulSoup(_html_response_text(html), "html.parser")
|
||||
dl = soup.find("a", href=True, class_="addDownloadedBook")
|
||||
url = dl["href"] if dl else ""
|
||||
url = (_get_attr(dl, "href") or "") if isinstance(dl, Tag) else ""
|
||||
|
||||
# AA slow download / partner servers
|
||||
elif "/slow_download/" in link:
|
||||
@@ -984,9 +1164,11 @@ def _get_download_url(
|
||||
)
|
||||
|
||||
else:
|
||||
get_btn = soup.find("a", string="GET") or soup.find("a", string="Download")
|
||||
get_btn = _find_first_anchor_with_text(soup, "GET") or _find_first_anchor_with_text(
|
||||
soup, "Download"
|
||||
)
|
||||
if get_btn:
|
||||
url = get_btn.get("href", "")
|
||||
url = _get_attr(get_btn, "href") or ""
|
||||
else:
|
||||
logger.warning("Unknown source type, couldn't find download link: %s", link)
|
||||
url = ""
|
||||
@@ -1012,24 +1194,30 @@ def _extract_slow_download_url(
|
||||
if url.startswith("http") and "/slow_download/" not in url:
|
||||
return url
|
||||
|
||||
dl_link = soup.find("a", href=True, string="📚 Download now")
|
||||
if not dl_link:
|
||||
dl_link = soup.find("a", href=True, string=lambda s: s and "Download now" in s)
|
||||
dl_link = _find_first_anchor_with_text(soup, "📚 Download now") or _find_first_anchor_with_text(
|
||||
soup, "Download now", contains=True
|
||||
)
|
||||
if dl_link:
|
||||
return dl_link["href"]
|
||||
return _get_attr(dl_link, "href") or ""
|
||||
|
||||
for a_tag in soup.find_all("a", href=True):
|
||||
if a_tag.has_attr("download"):
|
||||
href = a_tag["href"]
|
||||
href = _get_attr(a_tag, "href")
|
||||
if not href:
|
||||
continue
|
||||
if href.startswith("http") and "/slow_download/" not in href:
|
||||
return href
|
||||
|
||||
for span in soup.find_all("span", class_=lambda c: c and "whitespace-normal" in c):
|
||||
for span in soup.find_all("span"):
|
||||
if not _tag_has_class_containing(span, "whitespace-normal"):
|
||||
continue
|
||||
text = span.get_text(strip=True)
|
||||
if text.startswith(("http://", "https://")) and "/slow_download/" not in text:
|
||||
return text
|
||||
|
||||
for span in soup.find_all("span", class_=lambda c: c and "bg-gray-200" in c):
|
||||
for span in soup.find_all("span"):
|
||||
if not _tag_has_class_containing(span, "bg-gray-200"):
|
||||
continue
|
||||
text = span.get_text(strip=True)
|
||||
if text.startswith(("http://", "https://")):
|
||||
return text
|
||||
@@ -1040,20 +1228,20 @@ def _extract_slow_download_url(
|
||||
if url.startswith("http") and "/slow_download/" not in url:
|
||||
return url
|
||||
|
||||
copy_text = soup.find(string=lambda s: s and "copy this url" in s.lower())
|
||||
copy_text = _find_text_node(soup, "copy this url")
|
||||
if copy_text and copy_text.parent:
|
||||
parent = copy_text.parent
|
||||
next_link = parent.find_next("a", href=True)
|
||||
if next_link and next_link.get("href"):
|
||||
return next_link["href"]
|
||||
if isinstance(next_link, Tag):
|
||||
next_href = _get_attr(next_link, "href")
|
||||
if next_href:
|
||||
return next_href
|
||||
code_elem = parent.find_next("code")
|
||||
if code_elem:
|
||||
if isinstance(code_elem, Tag):
|
||||
return code_elem.get_text(strip=True)
|
||||
for sibling in parent.find_next_siblings():
|
||||
text = (
|
||||
sibling.get_text(strip=True)
|
||||
if hasattr(sibling, "get_text")
|
||||
else str(sibling).strip()
|
||||
sibling.get_text(strip=True) if isinstance(sibling, Tag) else str(sibling).strip()
|
||||
)
|
||||
if text.startswith("http"):
|
||||
return text
|
||||
@@ -1101,14 +1289,16 @@ def _extract_slow_download_url(
|
||||
def _extract_countdown_seconds(soup: BeautifulSoup, html_str: str) -> int:
|
||||
"""Extract countdown timer seconds from AA slow download page."""
|
||||
countdown_elem = soup.find("span", class_="js-partner-countdown")
|
||||
if countdown_elem:
|
||||
if isinstance(countdown_elem, Tag):
|
||||
seconds = _parse_countdown_seconds_from_element(countdown_elem)
|
||||
if seconds is not None:
|
||||
return seconds
|
||||
|
||||
for elem in soup.find_all(
|
||||
["span", "div"], class_=lambda c: c and ("timer" in c.lower() or "countdown" in c.lower())
|
||||
):
|
||||
for elem in soup.find_all(["span", "div"]):
|
||||
if not (
|
||||
_tag_has_class_containing(elem, "timer") or _tag_has_class_containing(elem, "countdown")
|
||||
):
|
||||
continue
|
||||
seconds = _parse_countdown_seconds_from_element(elem)
|
||||
if seconds is not None:
|
||||
return seconds
|
||||
|
||||
@@ -28,6 +28,34 @@ DEFAULT_CACHE_TTL = 30 * 24 * 60 * 60
|
||||
_cache_lock = Lock()
|
||||
|
||||
|
||||
def _coerce_cache_ttl(value: object, default: int) -> int:
|
||||
"""Coerce a cache TTL value from config into a non-negative integer."""
|
||||
if isinstance(value, int) and not isinstance(value, bool):
|
||||
return max(value, 0)
|
||||
if isinstance(value, str):
|
||||
stripped = value.strip()
|
||||
if stripped:
|
||||
try:
|
||||
return max(int(stripped), 0)
|
||||
except ValueError:
|
||||
return default
|
||||
return default
|
||||
|
||||
|
||||
def _coerce_timestamp(value: object) -> float:
|
||||
"""Coerce cached timestamps into floats for age calculations."""
|
||||
if isinstance(value, int | float) and not isinstance(value, bool):
|
||||
return float(value)
|
||||
if isinstance(value, str):
|
||||
stripped = value.strip()
|
||||
if stripped:
|
||||
try:
|
||||
return float(stripped)
|
||||
except ValueError:
|
||||
return 0.0
|
||||
return 0.0
|
||||
|
||||
|
||||
def _generate_cache_key(provider: str, provider_id: str, content_type: str | None = None) -> str:
|
||||
"""Generate a cache key from provider, provider_id, and content type."""
|
||||
normalized_content_type = "audiobook" if check_audiobook(content_type) else "ebook"
|
||||
@@ -97,12 +125,7 @@ def get_cached_results(
|
||||
|
||||
if ttl_seconds is None:
|
||||
ttl_value = config.get("IRC_CACHE_TTL", DEFAULT_CACHE_TTL)
|
||||
# Config values are stored as strings, convert to int
|
||||
ttl_seconds = int(ttl_value) if ttl_value else DEFAULT_CACHE_TTL
|
||||
|
||||
# TTL of 0 means cache forever
|
||||
if ttl_seconds == 0:
|
||||
ttl_seconds = float("inf")
|
||||
ttl_seconds = _coerce_cache_ttl(ttl_value, DEFAULT_CACHE_TTL)
|
||||
|
||||
cache_key = _generate_cache_key(provider, provider_id, content_type)
|
||||
|
||||
@@ -114,10 +137,10 @@ def get_cached_results(
|
||||
return None
|
||||
|
||||
# Check expiration
|
||||
cached_at = entry.get("cached_at", 0)
|
||||
cached_at = _coerce_timestamp(entry.get("cached_at", 0))
|
||||
age = time.time() - cached_at
|
||||
|
||||
if age > ttl_seconds:
|
||||
if ttl_seconds != 0 and age > ttl_seconds:
|
||||
title = entry.get("title", cache_key)
|
||||
logger.debug(
|
||||
"IRC cache expired for '%s' (age: %.0fs > TTL: %ss)",
|
||||
@@ -243,8 +266,7 @@ def cleanup_expired(ttl_seconds: int | None = None) -> int:
|
||||
|
||||
if ttl_seconds is None:
|
||||
ttl_value = config.get("IRC_CACHE_TTL", DEFAULT_CACHE_TTL)
|
||||
# Config values are stored as strings, convert to int
|
||||
ttl_seconds = int(ttl_value) if ttl_value else DEFAULT_CACHE_TTL
|
||||
ttl_seconds = _coerce_cache_ttl(ttl_value, DEFAULT_CACHE_TTL)
|
||||
|
||||
current_time = time.time()
|
||||
removed = 0
|
||||
@@ -256,7 +278,8 @@ def cleanup_expired(ttl_seconds: int | None = None) -> int:
|
||||
expired_keys = [
|
||||
key
|
||||
for key, entry in entries.items()
|
||||
if current_time - entry.get("cached_at", 0) > ttl_seconds
|
||||
if ttl_seconds != 0
|
||||
and current_time - _coerce_timestamp(entry.get("cached_at", 0)) > ttl_seconds
|
||||
]
|
||||
|
||||
for key in expired_keys:
|
||||
@@ -280,8 +303,7 @@ def get_cache_stats() -> dict[str, Any]:
|
||||
from shelfmark.core.config import config
|
||||
|
||||
ttl_value = config.get("IRC_CACHE_TTL", DEFAULT_CACHE_TTL)
|
||||
# Config values are stored as strings, convert to int
|
||||
ttl_seconds = int(ttl_value) if ttl_value else DEFAULT_CACHE_TTL
|
||||
ttl_seconds = _coerce_cache_ttl(ttl_value, DEFAULT_CACHE_TTL)
|
||||
current_time = time.time()
|
||||
|
||||
with _cache_lock:
|
||||
@@ -292,7 +314,8 @@ def get_cache_stats() -> dict[str, Any]:
|
||||
expired = sum(
|
||||
1
|
||||
for entry in entries.values()
|
||||
if current_time - entry.get("cached_at", 0) > ttl_seconds
|
||||
if ttl_seconds != 0
|
||||
and current_time - _coerce_timestamp(entry.get("cached_at", 0)) > ttl_seconds
|
||||
)
|
||||
|
||||
# Calculate total releases cached
|
||||
|
||||
@@ -101,6 +101,14 @@ class IRCClient:
|
||||
# Track online servers (elevated users in channel)
|
||||
self.online_servers: set[str] = set()
|
||||
|
||||
def _require_socket(self) -> socket.socket:
|
||||
"""Return the active socket or raise when the client is disconnected."""
|
||||
sock = self._socket
|
||||
if sock is None:
|
||||
msg = "Not connected"
|
||||
raise IRCError(msg)
|
||||
return sock
|
||||
|
||||
def connect(self) -> None:
|
||||
"""Connect to IRC server, send USER/NICK, and wait for welcome."""
|
||||
logger.info("Connecting to %s:%s (TLS=%s)", self.server, self.port, self.use_tls)
|
||||
@@ -132,14 +140,14 @@ class IRCClient:
|
||||
# Wait for 001 (RPL_WELCOME) which confirms registration is complete
|
||||
# Server may take time for hostname lookup, ident check, etc.
|
||||
logger.debug("Waiting for server welcome (001)...")
|
||||
self._socket.settimeout(2.0) # Short timeout for polling
|
||||
sock.settimeout(2.0) # Short timeout for polling
|
||||
|
||||
start = time.time()
|
||||
timeout = 30.0 # Max wait for registration
|
||||
|
||||
while time.time() - start < timeout:
|
||||
try:
|
||||
data = self._socket.recv(RECV_BUFFER)
|
||||
data = sock.recv(RECV_BUFFER)
|
||||
if not data:
|
||||
msg = "Connection closed during registration"
|
||||
raise IRCConnectionError(msg)
|
||||
@@ -162,7 +170,7 @@ class IRCClient:
|
||||
|
||||
# 001 = RPL_WELCOME - registration complete
|
||||
if " 001 " in line:
|
||||
self._socket.settimeout(SOCKET_TIMEOUT) # Restore timeout
|
||||
sock.settimeout(SOCKET_TIMEOUT) # Restore timeout
|
||||
self._connected = True
|
||||
logger.info("Connected as %s", self.nick)
|
||||
return
|
||||
@@ -200,9 +208,10 @@ class IRCClient:
|
||||
self.online_servers.clear()
|
||||
|
||||
if wait_for_join:
|
||||
sock = self._require_socket()
|
||||
# Use a short socket timeout during join so we can check elapsed time
|
||||
original_timeout = self._socket.gettimeout()
|
||||
self._socket.settimeout(2.0) # 2 second recv timeout
|
||||
original_timeout = sock.gettimeout()
|
||||
sock.settimeout(2.0) # 2 second recv timeout
|
||||
|
||||
try:
|
||||
start = time.time()
|
||||
@@ -211,7 +220,7 @@ class IRCClient:
|
||||
while time.time() - start < timeout:
|
||||
# Read data with short timeout
|
||||
try:
|
||||
data = self._socket.recv(RECV_BUFFER)
|
||||
data = sock.recv(RECV_BUFFER)
|
||||
if not data:
|
||||
break
|
||||
self._buffer += data.decode("utf-8", errors="replace")
|
||||
@@ -255,7 +264,7 @@ class IRCClient:
|
||||
|
||||
finally:
|
||||
# Restore original socket timeout
|
||||
self._socket.settimeout(original_timeout)
|
||||
sock.settimeout(original_timeout)
|
||||
|
||||
def send_message(self, target: str, message: str) -> None:
|
||||
"""Send a PRIVMSG to a channel or user."""
|
||||
@@ -289,6 +298,7 @@ class IRCClient:
|
||||
|
||||
def _recv_lines(self) -> Iterator[str]:
|
||||
"""Receive and yield complete CRLF-delimited IRC lines."""
|
||||
sock = self._require_socket()
|
||||
while True:
|
||||
# Check if we have a complete line in buffer
|
||||
while "\r\n" in self._buffer:
|
||||
@@ -298,7 +308,7 @@ class IRCClient:
|
||||
|
||||
# Read more data
|
||||
try:
|
||||
data = self._socket.recv(RECV_BUFFER)
|
||||
data = sock.recv(RECV_BUFFER)
|
||||
if not data:
|
||||
return # Connection closed
|
||||
self._buffer += data.decode("utf-8", errors="replace")
|
||||
|
||||
@@ -27,7 +27,7 @@ class IRCConnectionManager:
|
||||
being idle for IDLE_TIMEOUT seconds.
|
||||
"""
|
||||
|
||||
_instance: IRCConnectionManager | None = None
|
||||
_instance: Self | None = None
|
||||
_lock = threading.Lock()
|
||||
|
||||
def __new__(cls) -> Self:
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
Handles downloading IRC releases via DCC protocol.
|
||||
"""
|
||||
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
@@ -22,6 +23,41 @@ if TYPE_CHECKING:
|
||||
logger = setup_logger(__name__)
|
||||
|
||||
|
||||
def _config_text(key: str) -> str:
|
||||
"""Read a string config value with whitespace trimmed."""
|
||||
value = config.get(key, "")
|
||||
if value is None:
|
||||
return ""
|
||||
return str(value).strip()
|
||||
|
||||
|
||||
def _config_port(key: str, default: int) -> int:
|
||||
"""Read an IRC port value from config, accepting ints and numeric strings."""
|
||||
value = config.get(key, default)
|
||||
if isinstance(value, int) and not isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
stripped = value.strip()
|
||||
if stripped:
|
||||
with suppress(ValueError):
|
||||
return int(stripped)
|
||||
return default
|
||||
|
||||
|
||||
def _config_bool(key: str, default: bool) -> bool:
|
||||
"""Read a boolean config value from config."""
|
||||
value = config.get(key, default)
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized in {"1", "true", "yes", "on"}:
|
||||
return True
|
||||
if normalized in {"0", "false", "no", "off"}:
|
||||
return False
|
||||
return default
|
||||
|
||||
|
||||
@register_handler("irc")
|
||||
class IRCDownloadHandler(DownloadHandler):
|
||||
"""Handle IRC DCC downloads."""
|
||||
@@ -38,11 +74,11 @@ class IRCDownloadHandler(DownloadHandler):
|
||||
logger.info("IRC download: %s...", download_request[:60])
|
||||
|
||||
# Get IRC settings
|
||||
server = config.get("IRC_SERVER", "")
|
||||
port = config.get("IRC_PORT", 6697)
|
||||
use_tls = config.get("IRC_USE_TLS", True)
|
||||
channel = config.get("IRC_CHANNEL", "")
|
||||
nick = config.get("IRC_NICK", "")
|
||||
server = _config_text("IRC_SERVER")
|
||||
port = _config_port("IRC_PORT", 6697)
|
||||
use_tls = _config_bool("IRC_USE_TLS", True)
|
||||
channel = _config_text("IRC_CHANNEL")
|
||||
nick = _config_text("IRC_NICK")
|
||||
|
||||
if not server or not channel or not nick:
|
||||
logger.warning("IRC not fully configured")
|
||||
|
||||
@@ -5,6 +5,7 @@ Parses the text files sent via DCC that contain search results.
|
||||
|
||||
import re
|
||||
import zipfile
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
@@ -55,6 +56,20 @@ ALL_RECOGNIZED_FORMATS = {
|
||||
}
|
||||
|
||||
|
||||
def _normalize_config_formats(raw_formats: object) -> set[str]:
|
||||
"""Normalize configured format values into a lowercase set."""
|
||||
if isinstance(raw_formats, str):
|
||||
return {fmt.strip().lower() for fmt in raw_formats.split(",") if fmt.strip()}
|
||||
if isinstance(raw_formats, Iterable):
|
||||
normalized_formats: set[str] = set()
|
||||
for fmt in raw_formats:
|
||||
normalized = str(fmt).strip().lower()
|
||||
if normalized:
|
||||
normalized_formats.add(normalized)
|
||||
return normalized_formats
|
||||
return set()
|
||||
|
||||
|
||||
def _get_supported_formats(content_type: str | None = None) -> set[str]:
|
||||
"""Get the supported formats for the requested content type."""
|
||||
if check_audiobook(content_type):
|
||||
@@ -64,9 +79,7 @@ def _get_supported_formats(content_type: str | None = None) -> set[str]:
|
||||
"SUPPORTED_FORMATS", ["epub", "mobi", "azw3", "fb2", "djvu", "cbz", "cbr"]
|
||||
)
|
||||
|
||||
if isinstance(formats, str):
|
||||
return {fmt.strip().lower() for fmt in formats.split(",") if fmt.strip()}
|
||||
return {fmt.lower() for fmt in formats}
|
||||
return _normalize_config_formats(formats)
|
||||
|
||||
|
||||
# Regex to parse result lines
|
||||
|
||||
@@ -36,6 +36,43 @@ from .parser import SearchResult, extract_results_from_zip, parse_results_file
|
||||
logger = setup_logger(__name__)
|
||||
|
||||
|
||||
def _config_text(key: str) -> str:
|
||||
"""Read a string config value with whitespace trimmed."""
|
||||
value = config.get(key, "")
|
||||
if value is None:
|
||||
return ""
|
||||
return str(value).strip()
|
||||
|
||||
|
||||
def _config_port(key: str, default: int) -> int:
|
||||
"""Read an IRC port value from config, accepting ints and numeric strings."""
|
||||
value = config.get(key, default)
|
||||
if isinstance(value, int) and not isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
stripped = value.strip()
|
||||
if stripped:
|
||||
try:
|
||||
return int(stripped)
|
||||
except ValueError:
|
||||
return default
|
||||
return default
|
||||
|
||||
|
||||
def _config_bool(key: str, default: bool) -> bool:
|
||||
"""Read a boolean config value from config."""
|
||||
value = config.get(key, default)
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized in {"1", "true", "yes", "on"}:
|
||||
return True
|
||||
if normalized in {"0", "false", "no", "off"}:
|
||||
return False
|
||||
return default
|
||||
|
||||
|
||||
def _emit_status(message: str, phase: str = "searching") -> None:
|
||||
"""Emit search status to frontend via WebSocket."""
|
||||
ws_manager.broadcast_search_status(
|
||||
@@ -81,9 +118,9 @@ class IRCReleaseSource(ReleaseSource):
|
||||
|
||||
def is_available(self) -> bool:
|
||||
"""Check if IRC is configured (server, channel, and nick are set)."""
|
||||
server = config.get("IRC_SERVER", "")
|
||||
channel = config.get("IRC_CHANNEL", "")
|
||||
nick = config.get("IRC_NICK", "")
|
||||
server = _config_text("IRC_SERVER")
|
||||
channel = _config_text("IRC_CHANNEL")
|
||||
nick = _config_text("IRC_NICK")
|
||||
return bool(server and channel and nick)
|
||||
|
||||
def get_column_config(self) -> ReleaseColumnConfig:
|
||||
@@ -162,12 +199,12 @@ class IRCReleaseSource(ReleaseSource):
|
||||
_enforce_rate_limit()
|
||||
|
||||
# Get IRC settings
|
||||
server = config.get("IRC_SERVER", "")
|
||||
port = config.get("IRC_PORT", 6697)
|
||||
use_tls = config.get("IRC_USE_TLS", True)
|
||||
channel = config.get("IRC_CHANNEL", "")
|
||||
nick = config.get("IRC_NICK", "")
|
||||
search_bot = config.get("IRC_SEARCH_BOT", "")
|
||||
server = _config_text("IRC_SERVER")
|
||||
port = _config_port("IRC_PORT", 6697)
|
||||
use_tls = _config_bool("IRC_USE_TLS", True)
|
||||
channel = _config_text("IRC_CHANNEL")
|
||||
nick = _config_text("IRC_NICK")
|
||||
search_bot = _config_text("IRC_SEARCH_BOT")
|
||||
|
||||
client = None
|
||||
try:
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
"""Prowlarr API client for connection testing, indexer listing, and search."""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from contextlib import suppress
|
||||
from http import HTTPStatus
|
||||
from typing import Any
|
||||
@@ -10,6 +11,7 @@ from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.utils import normalize_http_url
|
||||
from shelfmark.download.network import get_ssl_verify
|
||||
from shelfmark.release_sources.prowlarr.torznab import parse_torznab_xml
|
||||
from shelfmark.release_sources.prowlarr.utils import coerce_int_like
|
||||
|
||||
logger = setup_logger(__name__)
|
||||
|
||||
@@ -25,6 +27,31 @@ _PROWLARR_CLIENT_ERRORS = (
|
||||
)
|
||||
|
||||
|
||||
def _normalize_json_object(payload: object, *, context: str) -> dict[str, Any]:
|
||||
"""Return a JSON object payload with string keys or raise on unexpected shapes."""
|
||||
if not isinstance(payload, Mapping):
|
||||
msg = f"Unexpected {context} response payload"
|
||||
raise TypeError(msg)
|
||||
|
||||
normalized: dict[str, Any] = {}
|
||||
for key, value in payload.items():
|
||||
if not isinstance(key, str):
|
||||
msg = f"Unexpected {context} response payload"
|
||||
raise TypeError(msg)
|
||||
normalized[key] = value
|
||||
|
||||
return normalized
|
||||
|
||||
|
||||
def _normalize_json_object_list(payload: object, *, context: str) -> list[dict[str, Any]]:
|
||||
"""Return a list of JSON objects or raise on unexpected item shapes."""
|
||||
if not isinstance(payload, list):
|
||||
msg = f"Unexpected {context} response payload"
|
||||
raise TypeError(msg)
|
||||
|
||||
return [_normalize_json_object(item, context=context) for item in payload]
|
||||
|
||||
|
||||
class ProwlarrClient:
|
||||
"""Client for interacting with the Prowlarr API."""
|
||||
|
||||
@@ -89,7 +116,10 @@ class ProwlarrClient:
|
||||
"""Test connection to Prowlarr. Returns (success, message)."""
|
||||
logger.info("Testing Prowlarr connection to: %s", self.base_url)
|
||||
try:
|
||||
data = self._request("GET", "/api/v1/system/status")
|
||||
data = _normalize_json_object(
|
||||
self._request("GET", "/api/v1/system/status"),
|
||||
context="Prowlarr status",
|
||||
)
|
||||
version = data.get("version", "unknown")
|
||||
except requests.exceptions.ConnectionError:
|
||||
return False, "Could not connect to Prowlarr. Check the URL."
|
||||
@@ -107,7 +137,10 @@ class ProwlarrClient:
|
||||
def get_indexers(self) -> list[dict[str, Any]]:
|
||||
"""Get all configured indexers."""
|
||||
try:
|
||||
return self._request("GET", "/api/v1/indexer")
|
||||
return _normalize_json_object_list(
|
||||
self._request("GET", "/api/v1/indexer"),
|
||||
context="Prowlarr indexer list",
|
||||
)
|
||||
except _PROWLARR_CLIENT_ERRORS:
|
||||
logger.exception("Failed to get indexers")
|
||||
return []
|
||||
@@ -131,12 +164,8 @@ class ProwlarrClient:
|
||||
enriched_ids: list[int] = []
|
||||
|
||||
for idx in self.get_enabled_indexers_detailed():
|
||||
idx_id = idx.get("id")
|
||||
if idx_id is None:
|
||||
continue
|
||||
try:
|
||||
idx_id_int = int(idx_id)
|
||||
except TypeError, ValueError:
|
||||
idx_id_int = coerce_int_like(idx.get("id"))
|
||||
if idx_id_int is None:
|
||||
continue
|
||||
|
||||
if restrict_to is not None and idx_id_int not in restrict_to:
|
||||
|
||||
@@ -26,6 +26,7 @@ from shelfmark.download.clients.base_handler import (
|
||||
from shelfmark.release_sources import register_handler
|
||||
from shelfmark.release_sources.prowlarr.cache import get_release, remove_release
|
||||
from shelfmark.release_sources.prowlarr.utils import (
|
||||
coerce_int_like,
|
||||
get_preferred_download_url,
|
||||
get_protocol,
|
||||
)
|
||||
@@ -55,9 +56,8 @@ def _coerce_seed_time_minutes(raw_seed_time: object) -> int | None:
|
||||
if raw_seed_time is None:
|
||||
return None
|
||||
|
||||
try:
|
||||
seed_time_seconds = int(raw_seed_time)
|
||||
except TypeError, ValueError:
|
||||
seed_time_seconds = coerce_int_like(raw_seed_time)
|
||||
if seed_time_seconds is None:
|
||||
logger.warning("Invalid Prowlarr minimumSeedTime value: %r", raw_seed_time)
|
||||
return None
|
||||
|
||||
|
||||
@@ -4,6 +4,7 @@ from typing import Any
|
||||
|
||||
import requests
|
||||
|
||||
from shelfmark.core.request_helpers import normalize_optional_text
|
||||
from shelfmark.core.settings_registry import (
|
||||
ActionButton,
|
||||
CheckboxField,
|
||||
@@ -28,6 +29,21 @@ _PROWLARR_SETTINGS_ERRORS = (
|
||||
)
|
||||
|
||||
|
||||
def _resolve_setting_text(current_values: dict[str, Any], key: str, *, default: str = "") -> str:
|
||||
"""Prefer current form values, then fall back to persisted config text."""
|
||||
from shelfmark.core.config import config
|
||||
|
||||
current_value = normalize_optional_text(current_values.get(key))
|
||||
if current_value is not None:
|
||||
return current_value
|
||||
|
||||
config_value = normalize_optional_text(config.get(key, default))
|
||||
if config_value is not None:
|
||||
return config_value
|
||||
|
||||
return default
|
||||
|
||||
|
||||
def _get_indexer_options() -> list[dict[str, str]]:
|
||||
"""Fetch available indexers from Prowlarr for the multi-select field.
|
||||
|
||||
@@ -38,8 +54,8 @@ def _get_indexer_options() -> list[dict[str, str]]:
|
||||
|
||||
logger = setup_logger(__name__)
|
||||
|
||||
raw_url = config.get("PROWLARR_URL", "")
|
||||
api_key = config.get("PROWLARR_API_KEY", "")
|
||||
raw_url = normalize_optional_text(config.get("PROWLARR_URL", "")) or ""
|
||||
api_key = normalize_optional_text(config.get("PROWLARR_API_KEY", "")) or ""
|
||||
|
||||
if not raw_url or not api_key:
|
||||
return []
|
||||
@@ -86,13 +102,12 @@ def _get_indexer_options() -> list[dict[str, str]]:
|
||||
|
||||
def _test_prowlarr_connection(current_values: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
"""Test the Prowlarr connection using current form values."""
|
||||
from shelfmark.core.config import config
|
||||
from shelfmark.release_sources.prowlarr.api import ProwlarrClient
|
||||
|
||||
current_values = current_values or {}
|
||||
|
||||
raw_url = current_values.get("PROWLARR_URL") or config.get("PROWLARR_URL", "")
|
||||
api_key = current_values.get("PROWLARR_API_KEY") or config.get("PROWLARR_API_KEY", "")
|
||||
raw_url = _resolve_setting_text(current_values, "PROWLARR_URL")
|
||||
api_key = _resolve_setting_text(current_values, "PROWLARR_API_KEY")
|
||||
|
||||
if not raw_url:
|
||||
return {"success": False, "message": "Prowlarr URL is required"}
|
||||
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
import re
|
||||
import time
|
||||
from contextlib import suppress
|
||||
from typing import TYPE_CHECKING, ClassVar, NoReturn
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -11,6 +10,7 @@ if TYPE_CHECKING:
|
||||
|
||||
from shelfmark.core.config import config
|
||||
from shelfmark.core.logger import setup_logger
|
||||
from shelfmark.core.request_helpers import normalize_optional_text
|
||||
from shelfmark.core.search_plan import ReleaseSearchVariant
|
||||
from shelfmark.core.utils import normalize_http_url
|
||||
from shelfmark.release_sources import (
|
||||
@@ -30,6 +30,8 @@ from shelfmark.release_sources import (
|
||||
from shelfmark.release_sources.prowlarr.api import ProwlarrClient
|
||||
from shelfmark.release_sources.prowlarr.cache import cache_release
|
||||
from shelfmark.release_sources.prowlarr.utils import (
|
||||
coerce_float_like,
|
||||
coerce_int_like,
|
||||
get_preferred_download_url,
|
||||
get_protocol,
|
||||
)
|
||||
@@ -45,6 +47,21 @@ def _raise_timeout_error(message: str) -> NoReturn:
|
||||
raise TimeoutError(message)
|
||||
|
||||
|
||||
def _raise_invalid_indexer_id(item: object) -> NoReturn:
|
||||
msg = f"Invalid indexer id: {item!r}"
|
||||
raise ValueError(msg)
|
||||
|
||||
|
||||
def _raise_invalid_indexer_selection_type(selected: object) -> NoReturn:
|
||||
msg = f"Invalid PROWLARR_INDEXERS type: {type(selected).__name__}"
|
||||
raise TypeError(msg)
|
||||
|
||||
|
||||
def _coerce_indexer_id(value: object) -> int | None:
|
||||
"""Best-effort coercion for indexer identifiers from config/API payloads."""
|
||||
return coerce_int_like(value)
|
||||
|
||||
|
||||
def _parse_size(size_bytes: int | None) -> str | None:
|
||||
"""Convert bytes to human-readable size string."""
|
||||
if size_bytes is None or size_bytes <= 0:
|
||||
@@ -370,13 +387,8 @@ def _prowlarr_result_to_release(
|
||||
cache_release(source_id, result)
|
||||
|
||||
# Derive common indicators from torznab/newznab attrs when present.
|
||||
download_volume_factor = result.get("downloadVolumeFactor")
|
||||
is_freeleech = False
|
||||
try:
|
||||
if download_volume_factor is not None and float(download_volume_factor) == 0.0:
|
||||
is_freeleech = True
|
||||
except TypeError, ValueError:
|
||||
pass
|
||||
download_volume_factor = coerce_float_like(result.get("downloadVolumeFactor"))
|
||||
is_freeleech = download_volume_factor == 0.0
|
||||
|
||||
if any(flag.lower() in {"freeleech", "fl"} for flag in indexer_flags):
|
||||
is_freeleech = True
|
||||
@@ -473,11 +485,9 @@ class ProwlarrSource(ReleaseSource):
|
||||
|
||||
# If user has selected specific indexers, track those separately
|
||||
if selected_ids is not None:
|
||||
try:
|
||||
if int(idx_id) in selected_ids:
|
||||
selected_indexer_names.append(idx_name)
|
||||
except TypeError, ValueError:
|
||||
pass
|
||||
idx_id_int = _coerce_indexer_id(idx_id)
|
||||
if idx_id_int is not None and idx_id_int in selected_ids:
|
||||
selected_indexer_names.append(idx_name)
|
||||
|
||||
available_indexers = sorted(all_indexer_names) if all_indexer_names else None
|
||||
# Only set default_indexers if user has selected specific ones
|
||||
@@ -559,8 +569,8 @@ class ProwlarrSource(ReleaseSource):
|
||||
|
||||
def _get_client(self) -> ProwlarrClient | None:
|
||||
"""Get a configured Prowlarr client or None if not configured."""
|
||||
raw_url = config.get("PROWLARR_URL", "")
|
||||
api_key = config.get("PROWLARR_API_KEY", "")
|
||||
raw_url = normalize_optional_text(config.get("PROWLARR_URL", "")) or ""
|
||||
api_key = normalize_optional_text(config.get("PROWLARR_API_KEY", "")) or ""
|
||||
|
||||
if not raw_url or not api_key:
|
||||
return None
|
||||
@@ -585,10 +595,26 @@ class ProwlarrSource(ReleaseSource):
|
||||
try:
|
||||
if isinstance(selected, list):
|
||||
# Already a list from JSON config
|
||||
ids = [int(x) for x in selected if x]
|
||||
else:
|
||||
ids = []
|
||||
for item in selected:
|
||||
if not item:
|
||||
continue
|
||||
parsed_id = _coerce_indexer_id(item)
|
||||
if parsed_id is None:
|
||||
_raise_invalid_indexer_id(item)
|
||||
ids.append(parsed_id)
|
||||
elif isinstance(selected, str):
|
||||
# Comma-separated string from env var
|
||||
ids = [int(x.strip()) for x in selected.split(",") if x.strip()]
|
||||
ids = []
|
||||
for item in selected.split(","):
|
||||
if not item.strip():
|
||||
continue
|
||||
parsed_id = _coerce_indexer_id(item)
|
||||
if parsed_id is None:
|
||||
_raise_invalid_indexer_id(item)
|
||||
ids.append(parsed_id)
|
||||
else:
|
||||
_raise_invalid_indexer_selection_type(selected)
|
||||
except (ValueError, TypeError) as e:
|
||||
logger.warning("Invalid PROWLARR_INDEXERS format: %s (%s)", selected, e)
|
||||
return None
|
||||
@@ -616,9 +642,9 @@ class ProwlarrSource(ReleaseSource):
|
||||
ids = []
|
||||
for name in names:
|
||||
idx_id = name_to_id.get(name)
|
||||
if idx_id is not None:
|
||||
with suppress(TypeError, ValueError):
|
||||
ids.append(int(idx_id))
|
||||
parsed_id = _coerce_indexer_id(idx_id)
|
||||
if parsed_id is not None:
|
||||
ids.append(parsed_id)
|
||||
except _PROWLARR_SOURCE_ERRORS as e:
|
||||
logger.warning("Failed to resolve indexer names to IDs: %s", e)
|
||||
return None
|
||||
@@ -647,10 +673,10 @@ class ProwlarrSource(ReleaseSource):
|
||||
continue
|
||||
|
||||
indexer_id = indexer.get("id")
|
||||
try:
|
||||
indexer_ids.append(int(indexer_id))
|
||||
except TypeError, ValueError:
|
||||
parsed_indexer_id = _coerce_indexer_id(indexer_id)
|
||||
if parsed_indexer_id is None:
|
||||
continue
|
||||
indexer_ids.append(parsed_indexer_id)
|
||||
|
||||
return indexer_ids
|
||||
|
||||
@@ -815,10 +841,7 @@ class ProwlarrSource(ReleaseSource):
|
||||
|
||||
for r in all_results:
|
||||
idx_id = r.get("indexerId")
|
||||
try:
|
||||
idx_id_int = int(idx_id) if idx_id is not None else None
|
||||
except TypeError, ValueError:
|
||||
idx_id_int = None
|
||||
idx_id_int = _coerce_indexer_id(idx_id)
|
||||
|
||||
is_enriched = bool(
|
||||
idx_id_int is not None and idx_id_int in enriched_indexer_ids_set
|
||||
@@ -864,6 +887,6 @@ class ProwlarrSource(ReleaseSource):
|
||||
"""Check if Prowlarr is enabled and configured."""
|
||||
if not config.get("PROWLARR_ENABLED", False):
|
||||
return False
|
||||
url = normalize_http_url(config.get("PROWLARR_URL", ""))
|
||||
api_key = config.get("PROWLARR_API_KEY", "")
|
||||
url = normalize_http_url(normalize_optional_text(config.get("PROWLARR_URL", "")))
|
||||
api_key = normalize_optional_text(config.get("PROWLARR_API_KEY", "")) or ""
|
||||
return bool(url and api_key)
|
||||
|
||||
@@ -3,12 +3,48 @@
|
||||
Provides common helper functions used across the Prowlarr plugin.
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import TYPE_CHECKING
|
||||
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
|
||||
|
||||
from shelfmark.core.request_helpers import normalize_optional_text
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from pathlib import Path
|
||||
|
||||
_INTEGER_LIKE_PATTERN = re.compile(r"^[+-]?\d+$")
|
||||
_FLOAT_LIKE_PATTERN = re.compile(r"^[+-]?(?:\d+(?:\.\d*)?|\.\d+)$")
|
||||
|
||||
|
||||
def coerce_int_like(value: object) -> int | None:
|
||||
"""Return an integer for int-like config/API values, else None."""
|
||||
if isinstance(value, bool):
|
||||
return None
|
||||
if isinstance(value, int):
|
||||
return value
|
||||
if isinstance(value, float):
|
||||
return int(value) if value.is_integer() else None
|
||||
|
||||
normalized = normalize_optional_text(value)
|
||||
if normalized is None or not _INTEGER_LIKE_PATTERN.fullmatch(normalized):
|
||||
return None
|
||||
|
||||
return int(normalized)
|
||||
|
||||
|
||||
def coerce_float_like(value: object) -> float | None:
|
||||
"""Return a float for float-like config/API values, else None."""
|
||||
if isinstance(value, bool):
|
||||
return None
|
||||
if isinstance(value, (int, float)):
|
||||
return float(value)
|
||||
|
||||
normalized = normalize_optional_text(value)
|
||||
if normalized is None or not _FLOAT_LIKE_PATTERN.fullmatch(normalized):
|
||||
return None
|
||||
|
||||
return float(normalized)
|
||||
|
||||
|
||||
def get_protocol(result: dict) -> str:
|
||||
"""Get the download protocol from a Prowlarr result.
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
"""Focused auth API regression tests for lockout handling."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
from datetime import datetime
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def main_module():
|
||||
"""Import `shelfmark.main` with background startup disabled."""
|
||||
with patch("shelfmark.download.orchestrator.start"):
|
||||
import shelfmark.main as main
|
||||
|
||||
importlib.reload(main)
|
||||
return main
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client(main_module):
|
||||
main_module.failed_login_attempts.clear()
|
||||
try:
|
||||
yield main_module.app.test_client()
|
||||
finally:
|
||||
main_module.failed_login_attempts.clear()
|
||||
|
||||
|
||||
class TestLoginLockoutRepair:
|
||||
def test_is_account_locked_repairs_missing_timestamp(self, main_module):
|
||||
main_module.failed_login_attempts.clear()
|
||||
main_module.failed_login_attempts["locked-user"] = {
|
||||
"count": main_module.MAX_LOGIN_ATTEMPTS
|
||||
}
|
||||
|
||||
assert main_module.is_account_locked("locked-user") is True
|
||||
assert isinstance(
|
||||
main_module.failed_login_attempts["locked-user"].get("lockout_until"), datetime
|
||||
)
|
||||
|
||||
def test_login_keeps_account_locked_when_timestamp_is_missing(self, main_module, client):
|
||||
main_module.failed_login_attempts["locked-user"] = {
|
||||
"count": main_module.MAX_LOGIN_ATTEMPTS
|
||||
}
|
||||
|
||||
with patch.object(main_module, "get_auth_mode", return_value="builtin"):
|
||||
response = client.post(
|
||||
"/api/auth/login",
|
||||
json={"username": "locked-user", "password": "secret", "remember_me": False},
|
||||
)
|
||||
|
||||
assert response.status_code == 429
|
||||
assert "Account temporarily locked" in response.get_json()["error"]
|
||||
assert isinstance(
|
||||
main_module.failed_login_attempts["locked-user"].get("lockout_until"), datetime
|
||||
)
|
||||
Reference in New Issue
Block a user