Compare commits

...
27 Commits
Author SHA1 Message Date
Alex 05115f7b41 Fix HTTP request behavior (#626) 2026-02-20 10:58:20 +00:00
Alex 554f5fcbe7 Patch: Various feature additions (#625)
- Add admin config for self-settings options visibility. Remove delivery
preferences or notifications from the view.
- Add option to use Booklore's Bookdrop API destination instead of a
specific library
- Add download path options for all torrent clients
2026-02-20 09:53:47 +00:00
Alex 8ff2d776ae Fix ABB magnet parsing (#623) 2026-02-16 16:58:56 +00:00
Alex 6c351f4bf3 Fix TypeScript error (#622) 2026-02-16 14:57:49 +00:00
Alex 7fdf55f5fd Enhancements to ABB handling (#621)
- Migrate download client handling from /prowlarr to /download. Moves
all torrent/usenet handling to app-level and gives ABB this
functionality.
- ABB Scraper now uses shared HTTP infrastructure instead of raw
requests, adding retry and proxy support
- Added author, age and bitrate info to ABB search results
- Added "best match" sorting option for releases
- Added size and bitrate sorting options for ABB
- Removed bundled default ABB hostname, must be configured by the user
- Added URL normalisation for ABB hostname
- Rearranged settings UI, moved download clients to its own section. 
- More tests
2026-02-16 14:52:46 +00:00
bonsai-dreams dd6fd1e199 Feature: Add AudiobookBay release source (#619)
## Summary

Adds AudiobookBay as a web-scraping release source for audiobook
torrents. Once enabled, a new tab shows up in the Find Releases modal.

## What's New

- **AudiobookBay source** – Search AudiobookBay for audiobook torrents
from the Shelfmark UI
- **Torrent downloads** – Extract magnet links from detail pages and add
them to the configured torrent client
- **Audiobook-only** – Source is limited to audiobooks
- **Download clients:** Currently uses the torrent client configured
under **Prowlarr > Download Clients**.
- Audiobook-specific categories (e.g. `QBITTORRENT_CATEGORY_AUDIOBOOK`)
are applied when set.
- **Settings → AudiobookBay**:
  - Enable toggle
  - Hostname
  - Max pages to search (default 5)
  - Rate limit delay in seconds (default 1)


## How It Works

1. User searches for an audiobook; AudiobookBay is queried if enabled.
2. Results show title, language, format, and size
3. User selects a release; the handler fetches the detail page and
extracts the magnet link.
4. Magnet link is sent to the configured torrent client.

## Testing

- Unit tests for source, handler, scraper, and utils
- Mocked HTTP requests and torrent client calls
- Coverage for search, relevance filtering, language mapping, size
parsing, and download flow

### Screenshots
<img width="600" alt="image"
src="https://github.com/user-attachments/assets/2e10a259-5c35-4065-980d-b59a1c961c9f"
/>
2026-02-16 14:44:47 +00:00
Alex ccb39e674e Patch: Request retry, admin-level requests, and various fixes (#620)
- Add request retry in the case of a download failure, admins will be prompted to attach a new file to the request
- Add admin-level "add to requests" button in the release modal
2026-02-16 12:27:56 +00:00
Alex 1931eb96a5 Feature: Notification support + Enhanced request management (#618)
- Added notification support via Apprise dependency
- Notifications can be configured globally or per user, with full
customization of events and notification type.
- Added expanded ActivityCard for increased detail of each request, file
info, and managing the attached file.
- Enhanced tests
2026-02-15 17:59:53 +00:00
Alex b7bee132a1 Requests: Various fixes and improvements (#617)
- Refactored activity backend for full user-level management, using the
db file
- Revamped the activity sidebar UX and categorisation
- Added download history and user filtering
- Added User Preferences modal, giving limited configuration for
non-admins - replaces the "restrict settings" config option.
- Many many bug fixes
- Many many new tests
2026-02-14 18:24:28 +00:00
Alex 68608b6162 Feature: Multi-user request system (#615)
- Adds a comprehensive multi-user request system to the existing
download flow
- Request configuration is policy based. Configure global settings for
content type, or narrow down policy for specific sources (E.g. allow
direct downloads, set prowlarr to request only, block IRC completely,
etc).
- Global policy configuration and per-user overrides for tailored
configs
- Replaced downloads sidebar with ActivitySidebar, combining active
downloads with requests. Admin management of user requests is done here,
and admins have view of downloads from all users. Sidebar can now be
pinned.
- Request either a standard book or a specific release. Release-requests
are used if you permit one source differently than the other. On
book-level requests, admins pick the specific file to be attached to the
fulfilled request.
- Users can request books with a note

This is WIP so some features are still not complete (notifications, more
automatic release selection, among others).
2026-02-14 11:08:20 +00:00
Alex af9d9ec8db Patch: Further multi-user fixes (#613) 2026-02-12 17:47:10 +00:00
arjunsrinivasan1997 a7064939ce feat: Add tag support to qBittorrent (#610)
Added support for adding tag(s) to torrents sent to qBittorrent via
shelfmark.
![Screenshot 2026-02-11 at 3 58
20 AM](https://github.com/user-attachments/assets/aa9b440a-27fd-4166-953b-31f5179688a3)
![Screenshot 2026-02-11 at 3 53
14 AM](https://github.com/user-attachments/assets/15084b44-9a68-493c-85e8-328c92206c85)
2026-02-12 14:52:34 +00:00
Alex 5bed0b20f4 Patch: Multi-user and OIDC polish (#612)
- Moved backend OIDC functionality to external library Authlib to help
maintainability
- Separated User settings UI into individual components, allowing for
standard settings UI decorator components to be used.
- Added full support for reverse proxy and CWA users alongside local and
OIDC
- Added mapping and syncing functionality for OIDC, CWA and reverse
proxy users
- Added per-user settings into the app-wide config system. Each config
can be declared as user-overrideable, and app-wide functionality can now
receive user-specific options via standard config calls.
- Added per-user audiobook destination config
- Updated login modal UI for simplified login, plus custom labels for
OIDC login
- Added user visibility in header dropdown
- Unified "restrict settings to admin" to use app-wide user roles.
2026-02-12 14:38:28 +00:00
Michael Joshua SaulandClaude Opus 4.6 2d2f54729f Add OIDC authentication and multi-user support (#606)
Closes #552

## Summary

Adds OIDC authentication and multi-user support to Shelfmark. Users can
now be managed individually with per-user download settings, while
maintaining full backwards compatibility with existing auth modes
(no-auth, builtin, proxy, CWA).

### Authentication
- **OIDC login** with PKCE, auto-discovery, group-based admin mapping
- **Password fallback** when OIDC is enabled (prevents admin lockout)
- **Auto-provisioning** of OIDC users (configurable on/off)
- **Email-based linking** of pre-created users to OIDC accounts
- **Lockout prevention** — requires a local admin before OIDC can be
enabled

### User Management
- **SQLite user database** (`users.db`) with admin CRUD API
- **Users management tab** in settings UI (admin-only)
- **Settings restricted to admins** in multi-user modes (builtin/OIDC) —
non-admin users cannot access settings
- Create, edit, and delete users with role assignment (admin/user)
- Password management for builtin auth users
- OIDC users shown with provider badge (password fields hidden)
- Per-user configurable settings:
  - **Download destination** — custom folder path per user
- **BookLore library & path** — dropdown select, each user's books go to
their own library
  - **Email recipients** — per-user email delivery targets
- **`{User}` template variable** — use in destination paths (e.g.,
`/books/{User}/`)
- Settings override model: per-user values override globals, empty/unset
falls back to global defaults

### Download Scoping
- **Per-user download visibility** — non-admins only see their own
downloads
- **Username display** in downloads sidebar (shows who requested each
download)
- **WebSocket room-based filtering** — admins see all, users see only
their own
- **Download progress scoping** — progress events routed to correct user
rooms

### BookLore Integration
- **Dynamic dropdown selects** for library/path (replaces text inputs)
- **Per-user library/path overrides** via user settings
- **Options cache refresh** after Test Connection

### Security
- SQL injection prevention (column whitelist on user updates)
- Generic OIDC error messages (no internal detail leakage)
- Admin self-deletion and last-local-admin deletion guards
- OIDC role overwrite fix (only updates role when admin_group is
configured)

## Migration

**No migration script needed.** The `users.db` is created automatically
on first startup. Existing builtin auth users are auto-migrated to the
database on their first login. All other auth modes (no-auth, proxy,
CWA) continue working unchanged.

## Test Plan

- [x] All 519 tests passing, 0 failures
- [ ] Test no-auth mode: settings accessible, downloads work without
login
- [ ] Test builtin auth: legacy credentials auto-migrate on login, new
users can be created
- [ ] Test OIDC auth: login flow, callback, auto-provisioning,
group-based admin
- [ ] Test CWA auth: unchanged behavior
- [ ] Test proxy auth: unchanged behavior
- [ ] Test per-user downloads: non-admin sees only own downloads
- [ ] Test BookLore dropdowns: library/path selection, per-user
overrides
- [ ] Test Docker build: no Dockerfile changes needed

---------

Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-11 17:44:27 +00:00
Alex b5923635a6 Fix recipient modal layout (#604) 2026-02-09 20:13:52 +00:00
Alex e09f5f7757 Feature: Email output mode (#603) 2026-02-09 19:33:12 +00:00
Alex 022e50a0ba Add threading to file system operations (#602) 2026-02-09 18:04:10 +00:00
Alex a560089ce3 Patch: Script improvements + bug fixes (#591)
- Add new booklore API file formats
- Renamed cookie for better login persistence with reverse proxy
- Updated fs.py to try hardlink before atomic move from tmp dir
- Fix transmission URL parsing 
- Fix scenario where file processing of huge files starves the
healthcheck
- Large enhancements to custom scripting, including passing JSON
download info, more consistent activation across output types,
decoupling from staging behavior, and added full documentation.
2026-02-06 13:51:23 +00:00
Alex f84fb082ad Fix: AA mirror behavior (#589)
- Refreshed available AA URLs
- Fixed potential redirect from AA itself causing mirror cache errors
- Added fully customizable mirror list in UI
- Segmented rotation behavior to Auto mode only

Fixes #588
2026-02-06 10:04:31 +00:00
Alex b10458a48b Patch: Migrate bypasser to pure CDP + Misc fixes (#575)
Bypasser:
- Refactored internal bypasser logic to use SeleniumBase Pure CDP mode,
removed chromedriver dependencies and UC code.
- Added dedicated threading for internal bypasser functions, fixes any
potential asyncio CPU spike behavior
- Fixed WebGL issue with Chromium 144. Reverted 1.0.3 hotfix and updated
to latest Chromium

Misc: 
- Added M4A color mapping
- Fix frontend language filtering with multi-language releases
- Added "days" age for usenet/torrent releases
- Improved entrypoint chown efficiency
- Added `ONBOARDING` env variable, default true
2026-02-02 20:32:19 +00:00
Andy Kelk f6dba959c9 Fix: Base path resolution timing issue for subpath deployments (#572)
When deployed under a URL prefix (e.g., /shelfmark), images loaded by
React were not respecting the base path, causing 404 errors. The logo
would incorrectly load from /logo.png instead of /shelfmark/logo.png.

The root cause was that the BASE_PATH constant was being initialized at
module load time, before the DOM was fully parsed. This meant
document.querySelector('base') returned null, causing BASE_PATH to
default to '/' regardless of the actual base tag value.

Changed to lazy initialization pattern where the base path is resolved
on first access, ensuring the DOM and base tag are ready.

Fixes [#571](https://github.com/calibrain/shelfmark/issues/571)
2026-02-02 17:26:06 +00:00
Alex e5ccabe1ef Patch: Various additions (#564)
- Added rich Prowlarr search results for whitelisted indexers
- Added torznab query for whitelisted indexers
- Added flags for all Prowlarr indexers
- Added completed external download retry mechanism and "locating" state
- Added client side preference storage of Book/Audiobook search
preference
- Fixed reverse proxy base URL in edge cases
- Added gevent locking for I/O operations, keeps healthcheck alive on
intensive processing operations
- Added M4A supported audiobook option
- Improved file transfer counting and logging with hardlink fallback
warnings
- Fixed proxy auth header for REMOTE_USER scenario
- Dependency tweak for internal bypasser
2026-01-31 12:53:11 +00:00
Marcel Meier 86082c999c Enhance naming templates with arbitrary prefix/suffix support (#560)
As described in https://github.com/calibrain/shelfmark/issues/559
I would like the option to use prefix/suffix text as part of my file
handling.

I appreciate every feedack
2026-01-30 13:53:05 +00:00
Webysther Sperandio 301b2e5456 Update readme (#557)
Confirmed working with automatic importing in calibre and calibre-web.
2026-01-29 19:39:05 +00:00
Ryan Dawes 4fde128fc7 Feature: Add indexer flags to results (#539)
Display Prowlarr indexer flags by rendering them as distinct,
color-coded badges.

- [New] TAGS Render Type: Added support for a TAGS column type that
renders a list of strings as distinct badges.
- Updated  `ReleaseCell` to handle the TAGS type:
  - Desktop: Renders distinct badges side-by-side.
- Mobile: Renders as a comma-separated text list (e.g., "FREELEECH,
DOUBLE UPLOAD").
- Styling: Added dynamic colors for common flags:
  - Freeleech → Green
  - Double Upload → Blue
  - VIP → Amber
  - Sticky → Yellow
- Prowlarr Source: Updated the "Flags" column to use the new TAGS render
type, enable uppercase styling, and show on mobile devices.

Desktop screenshot:
<img width="2042" height="110" alt="CleanShot 2026-01-25 at 21 32 24@2x"
src="https://github.com/user-attachments/assets/d135b1d6-176c-4cb9-afa8-fbbab0bcbc06"
/>

Mobile screenshot:
<img width="567" height="51" alt="image"
src="https://github.com/user-attachments/assets/b5b38a05-2466-4b6f-b4c4-ab6cce745408"
/>
2026-01-29 19:38:28 +00:00
Patrick VeverkaandCopilot d050417e01 fix newer versions of rtorrent (#549)
Closes https://github.com/calibrain/shelfmark/issues/534

This pull request enhances the rTorrent client testing and
implementation by adding more robust checks for directory paths and
improving how the base path is retrieved. The main focus is on verifying
and obtaining the correct download and base directories for torrents.

---------

Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
2026-01-29 18:41:19 +00:00
KhakisandClaude Opus 4.5 0a7785a333 Docs: Subpath reverse proxy configuration (#542)
## Summary

This PR updates the reverse proxy documentation with comprehensive
configuration examples for subpath deployments, addressing several
issues discovered when running Shelfmark behind nginx at a subpath like
`/shelfmark/`.

## Changes

- **Root path setup**: Added complete nginx server block example
- **Subpath setup without auth**: Complete nginx configuration with all
necessary workarounds
- **Subpath setup with Authelia**: Full example including Authelia
snippets and Shelfmark proxy auth settings
- **Known issues section**: Documents the frontend bugs that require
workarounds

## Issues Addressed

The current documentation's simple example doesn't work for subpath
deployments because:

1. **Socket.IO connects to root**: Frontend connects to `/socket.io/`
instead of `/shelfmark/socket.io/`
2. **API calls use root path**: Cover images request `/api/` instead of
`/shelfmark/api/`
3. **Logo uses root path**: Requested from `/logo.png` instead of
`/shelfmark/logo.png`
4. **Socket.IO backend path**: Always at `/socket.io/` regardless of
`URL_BASE` setting

## Testing

Tested with:
- Nginx reverse proxy
- Authelia authentication proxy
- `URL_BASE=/shelfmark/` configuration
- WebSocket connections working
- Cover images loading
- Proxy authentication with admin group restrictions

Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
2026-01-26 20:03:35 +00:00
268 changed files with 44099 additions and 4410 deletions
+1
View File
@@ -232,3 +232,4 @@ pyrightconfig.json
AGENTS.md
.claude/
.playwright-mcp/
frontend-dist/
+5 -6
View File
@@ -130,11 +130,11 @@ RUN apt-get update && \
xvfb \
# For screen recording
ffmpeg \
# --- Chromium ---
chromium=143.0.7499.169-1~deb13u1 \
chromium-common=143.0.7499.169-1~deb13u1 \
# --- ChromeDriver ---
chromium-driver=143.0.7499.169-1~deb13u1 \
# --- Chromium (unpinned - uses latest from Debian repos) ---
# Chrome 144+ requires --enable-unsafe-swiftshader for WebGL in Docker.
# This flag is set in internal_bypasser.py _get_browser_args()
chromium \
chromium-common \
# For tkinter (pyautogui)
python3-tk \
# For RAR extraction
@@ -152,7 +152,6 @@ RUN --mount=type=cache,target=/root/.cache/pip \
# Grant read/execute permissions to others
RUN chmod -R o+rx /usr/bin/chromium && \
chmod -R o+rx /usr/bin/chromedriver && \
chmod -R o+rwx /usr/local/lib/python3.10/site-packages/seleniumbase/drivers/
# Default command to run the application entrypoint script
+39
View File
@@ -0,0 +1,39 @@
# Bypass testing - switch between dev build and v1.0.1
# Usage:
# Test dev build: docker compose -f docker-compose.bypass-test.yml up shelfmark-dev
# Test v1.0.1: docker compose -f docker-compose.bypass-test.yml up shelfmark-stable
# Pull latest dev: docker compose -f docker-compose.bypass-test.yml build shelfmark-dev
# Pull v1.0.1: docker compose -f docker-compose.bypass-test.yml pull shelfmark-stable
services:
# Dev image from registry
shelfmark-dev:
image: ghcr.io/calibrain/shelfmark:dev
container_name: shelfmark-bypass-dev
environment:
PUID: 1000
PGID: 1000
DEBUG: true
ports:
- 8084:8084
volumes:
- ./.local/bypass-test/config-dev:/config
- ./.local/bypass-test/books:/books
- ./.local/bypass-test/log-dev:/var/log/shelfmark
- ./.local/bypass-test/tmp:/tmp/shelfmark
# Stable v1.0.1 for comparison
shelfmark-stable:
image: ghcr.io/calibrain/shelfmark:1.0.1
container_name: shelfmark-bypass-stable
environment:
PUID: 1000
PGID: 1000
DEBUG: true
ports:
- 8085:8084
volumes:
- ./.local/bypass-test/config-stable:/config
- ./.local/bypass-test/books:/books
- ./.local/bypass-test/log-stable:/var/log/shelfmark
- ./.local/bypass-test/tmp:/tmp/shelfmark
+124 -2
View File
@@ -1,3 +1,125 @@
# Configuration
# Directory and Volume Setup
TODO
This guide explains how to configure directories and Docker volumes for Shelfmark. It focuses on the difference between the destination folder and your download client paths, and how to make those paths line up inside containers.
## Conceptual Overview
```
DIRECT DOWNLOADS
Shelfmark downloads directly -> destination
TORRENT / USENET
Prowlarr -> Download client saves to <client path>
-> Shelfmark reads from <client path>
-> Shelfmark processes to destination
```
Key point: For torrent and usenet downloads, Shelfmark must see the same file path that your download client reports. The container path must match in both containers.
## Direct Download Setup
Direct downloads do not use an external download client. A simple two-folder setup is enough.
Required volumes:
| Container path | Purpose | Notes |
| --- | --- | --- |
| `/config` | Settings, database, cover cache | Configurable via `CONFIG_DIR` |
| `/books` | Destination folder for completed files | Configurable via `INGEST_DIR` and Settings -> Downloads -> Destination |
Example `docker-compose`:
```yaml
services:
shelfmark:
image: ghcr.io/calibrain/shelfmark:latest
volumes:
- /path/to/config:/config
- /path/to/books:/books
```
Notes:
- Point `/books` to your library ingest folder (Calibre-Web, Booklore, Audiobookshelf, etc) for automatic import.
- If you set Books Output Mode to Booklore (API), books are uploaded via API instead of written to `/books`. Audiobooks still use a destination folder.
- Ensure `PUID`/`PGID` (or legacy `UID`/`GID`) match the owner of the host directories to avoid permission errors.
## Torrent / Usenet Setup
For torrents and usenet, your download client reports a path (for example `/data/torrents/books/MyBook.epub`). Shelfmark must be able to read that exact path inside its own container.
Required volumes:
| Container path | Purpose | Notes |
| --- | --- | --- |
| `/config` | Settings, database, cover cache | Configurable via `CONFIG_DIR` |
| `/books` | Destination folder for processed files | Configurable via `INGEST_DIR` |
| `<client path>` | Download client path | Must match the download client container path exactly |
Side-by-side example with qBittorrent:
```yaml
services:
shelfmark:
volumes:
- /path/to/config:/config
- /path/to/books:/books
- /path/to/downloads:/data/torrents # Must match client
qbittorrent:
volumes:
- /path/to/downloads:/data/torrents # Same container path
```
Host paths can be anything. The container path (for example `/data/torrents`) must be identical in both containers.
### Remote Path Mappings
If paths cannot match (different machines or a fixed setup), use Remote Path Mappings.
Where to configure:
- Settings -> Advanced -> Remote Path Mappings
Example:
- Client reports `/data/torrents/books/...`
- Shelfmark can see the same files at `/downloads/books/...`
- Add a mapping from Remote Path `/data/torrents` to Local Path `/downloads`
## File Processing Options
### Transfer Method (Torrent / Usenet Only)
Available methods:
- Copy (default). Works everywhere.
- Hardlink. Preserves seeding without duplicating files.
Hardlink requirements and behavior:
- Source and destination must be on the same filesystem.
- If hardlinking is enabled but not possible, Shelfmark falls back to copying.
- Archive extraction is disabled while hardlinking is enabled.
- Do not use hardlinking if your destination is a library ingest folder.
### File Organization
Shelfmark supports three organization modes for the destination:
- None. Keep original filenames from the source.
- Rename Only. Rename files using a template.
- Rename and Organize. Create folders and rename using templates. Do not use with ingest folders.
Configure templates in Settings -> Downloads. Template syntax details are documented separately.
## Common Mistakes
- "Download failed - file not found": Path mismatch between Shelfmark and the download client. Ensure container paths match or use Remote Path Mappings.
- "Permission denied": `PUID`/`PGID` do not match the host directories. Ensure Shelfmark can read the client path and write to the destination.
- "Hardlinks not working" or "Files being copied instead": Source and destination are on different filesystems. Move the destination or accept copy fallback.
- "Downloads work but library does not see them": Destination does not point to the library ingest folder. Check Settings -> Downloads -> Destination.
- CIFS/SMB shares: Use the `nobrl` mount option to avoid database lock errors. Example: `//server/share /mnt/share cifs nobrl,... 0 0`
## Related Documentation
- Environment Variables Reference: `docs/environment-variables.md`
- Custom Scripts: `docs/custom-scripts.md`
- Installation: `docs/installation.md`
- Troubleshooting: `docs/troubleshooting.md`
+184
View File
@@ -0,0 +1,184 @@
# Custom Scripts
Shelfmark can run an executable you provide after a download task completes successfully. The script runs after the selected output has finished (for example: transfer to the folder destination, or upload to Booklore).
## Quick Start (Recommended)
1. Put your script on the machine that runs Shelfmark.
1. Make it executable.
1. Set it in Shelfmark (Settings -> Advanced -> Custom Script Path).
Example:
```bash
chmod +x /path/to/your/scripts/post_process.sh
```
### Docker Users
If you run Shelfmark in Docker, the script must exist inside the container. The easiest way is to mount a folder of scripts, then point Shelfmark at the container path in the UI.
```yaml
services:
shelfmark:
image: ghcr.io/calibrain/shelfmark:latest
volumes:
- /path/to/your/scripts:/scripts:ro
```
Then set:
- Settings -> Advanced -> Custom Script Path: `/scripts/post_process.sh`
<details>
<summary>Docker Compose: Configure Via Environment Variables (Optional)</summary>
```yaml
services:
shelfmark:
environment:
- CUSTOM_SCRIPT=/scripts/post_process.sh
- CUSTOM_SCRIPT_PATH_MODE=absolute
- CUSTOM_SCRIPT_JSON_PAYLOAD=true
```
</details>
## Script Behaviour
When enabled, Shelfmark runs your script once per successful task:
```bash
<custom_script_path> "<target_path>"
```
- `$1` is always set to the target path.
- If **Custom Script JSON Payload** is enabled, Shelfmark writes a JSON document to stdin (UTF-8).
- If JSON payload is disabled, stdin is empty (EOF).
- Timeout: 300 seconds (5 minutes)
- Exit code: `0` = success; anything else = the task is marked as **Error**
- Concurrency: downloads can run in parallel, so your script may be invoked concurrently for different tasks.
- Runtime: the script runs inside the Shelfmark container (if you use Docker) under the same user as Shelfmark.
## The Target Path (`$1`)
Shelfmark chooses a "best single path" for the task:
- If the output produced exactly one local file: that file path.
- If the output produced multiple local files: a directory path (the common parent directory of those files).
What the target path refers to depends on the output mode:
- Folder output (`output.mode=folder`, `phase=post_transfer`): the final imported file or folder inside your destination.
- Booklore output (`output.mode=booklore`, `phase=post_upload`): the local file or folder that was uploaded (the destination is remote).
By default, `$1` is an absolute path inside the Shelfmark container (or on your host, if you are not using Docker).
## JSON Payload (stdin)
Configure in: Settings -> Advanced -> Custom Script JSON Payload
When enabled, Shelfmark sends a versioned JSON payload to your script via stdin (and still passes `$1`). This is the recommended way to write robust scripts, especially for multi-file imports (audiobooks) and output-specific context (like Booklore).
- The JSON payload always includes absolute paths in `paths.*`, even if you set Custom Script Path Mode to `relative` for `$1`.
- `output.mode` tells you which output ran.
- `output.details` is output-specific. For Booklore output, `output.details.booklore` includes connection details such as `base_url`, `library_id`, and `path_id`.
- `phase` indicates when the script is running. Current values: `post_transfer` (folder output), `post_upload` (Booklore output).
- `transfer` is only included for outputs that do a local transfer (for example the folder output).
If JSON payload is disabled, stdin is empty (EOF). Don't `cat` stdin unless you've enabled the payload.
Example payload shape:
```json
{
"version": 1,
"phase": "post_transfer",
"task": {
"task_id": "abc123",
"source": "direct",
"title": "Foundation",
"author": "Isaac Asimov"
},
"output": {
"mode": "folder",
"organization_mode": "organize"
},
"paths": {
"destination": "/data/library/books",
"target": "/data/library/books/Isaac Asimov/Foundation/Foundation.epub",
"final_paths": [
"/data/library/books/Isaac Asimov/Foundation/Foundation.epub"
]
},
"transfer": {
"op_counts": {"copy": 1, "move": 0, "hardlink": 0},
"use_hardlink": false,
"is_torrent": false,
"preserve_source": false
}
}
```
Example (bash + jq) (JSON payload must be enabled):
```bash
payload="$(cat)"
mode="$(echo "$payload" | jq -r '.output.mode')"
title="$(echo "$payload" | jq -r '.task.title')"
final_paths="$(echo "$payload" | jq -r '.paths.final_paths[]')"
echo "mode=$mode title=$title" >&2
echo "$final_paths" >&2
```
Example (Python) (works whether JSON payload is enabled or not):
```python
#!/usr/bin/env python3
import json
import sys
target = sys.argv[1]
raw = sys.stdin.read()
payload = json.loads(raw) if raw.strip() else None
print(f"target={target}", file=sys.stderr)
if payload:
print(f"mode={payload['output']['mode']} phase={payload['phase']}", file=sys.stderr)
```
<details>
<summary>Advanced Options</summary>
### Absolute vs Relative Target Paths
Configure in: Settings -> Advanced -> Custom Script Path Mode
This setting controls what gets passed as `$1`:
- `absolute` (default): pass an absolute path.
- `relative`: pass a path relative to the output's "destination root", and run the script with `$PWD` set to that root.
For folder output, the destination root is your configured destination folder. For Booklore output, it's the local upload folder.
Example (folder destination is `/data/library/books`, and the imported file ended up in `Isaac Asimov/Foundation/Foundation.epub`):
```bash
# Absolute mode:
$PWD is unchanged
$1 = /data/library/books/Isaac Asimov/Foundation/Foundation.epub
# Relative mode:
$PWD = /data/library/books
$1 = Isaac Asimov/Foundation/Foundation.epub
```
Note: if the target is the destination folder itself, `relative` mode may pass `.`.
</details>
## Notes And Caveats
- **Hardlinks and torrents:** if you use hardlinking to keep seeding, avoid scripts that modify file contents, since hardlinked files share data with the seeding copy.
- **Booklore output mode:** scripts run after upload. `$1` will point at the local uploaded file (or staging folder).
+26
View File
@@ -190,6 +190,31 @@ HeadingField(
)
```
### CustomComponentField
Render a frontend-registered custom settings component while still using the
decorator-based schema.
```python
from shelfmark.core.settings_registry import CustomComponentField
CustomComponentField(
key="request_policy_editor",
component="request_policy_grid", # frontend registry key
label="Request Policy Rules",
description="Custom editor for policy defaults and matrix rules.",
value_fields=[
SelectField(key="REQUEST_POLICY_DEFAULT_EBOOK", label="Default Ebook Mode", default="download"),
SelectField(key="REQUEST_POLICY_DEFAULT_AUDIOBOOK", label="Default Audiobook Mode", default="download"),
TableField(key="REQUEST_POLICY_RULES", label="Rules", columns=_rule_columns, default=[]),
],
wrap_in_field_wrapper=True, # use standard FieldWrapper label/description layout
)
```
When `value_fields` is provided, those backing fields are included in
serialization/save/validation automatically and are hidden from the default renderer.
## Common Field Properties
All field types support these common properties:
@@ -206,6 +231,7 @@ All field types support these common properties:
| `requires_restart` | `bool` | `False` | Whether changes require container restart |
| `show_when` | `dict` | `None` | Conditional visibility (see below) |
| `disabled_when` | `dict` | `None` | Conditional disable (see below) |
| `hidden_in_ui` | `bool` | `False` | Hide from default renderer but keep in schema/save path |
## Conditional Visibility
File diff suppressed because it is too large Load Diff
+136 -24
View File
@@ -1,35 +1,147 @@
# Reverse Proxy & Subpath Hosting
Shelfmark can run behind a reverse proxy at the root path (recommended) or
under a subpath like `/shelfmark`.
Shelfmark can run behind a reverse proxy at the root path (recommended) or under a subpath like `/shelfmark`.
## Subpath setup
## Root path setup (Recommended)
1) Set the base path in Shelfmark:
- UI: Settings → Advanced → Base Path
- Env var: `URL_BASE=/shelfmark`
If you can serve Shelfmark at the root path (`https://shelfmark.example.com/`), leave `URL_BASE` empty. This is the simplest option and avoids extra subpath configuration.
2) Configure your reverse proxy to forward the subpath to Shelfmark and
**strip the prefix** before sending to the backend. The proxy must also allow
WebSocket upgrades for Socket.IO.
```nginx
server {
listen 443 ssl;
server_name shelfmark.example.com;
Example (Nginx-style):
```
location /shelfmark/ {
proxy_pass http://shelfmark:8084/;
proxy_http_version 1.1;
proxy_set_header Upgrade $http_upgrade;
proxy_set_header Connection "upgrade";
location / {
proxy_pass http://shelfmark:8084;
proxy_http_version 1.1;
proxy_set_header Host $host;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
proxy_set_header X-Forwarded-Proto $scheme;
proxy_set_header Upgrade $http_upgrade;
proxy_set_header Connection "upgrade";
}
}
```
Notes:
- Use a trailing slash on the `location` and `proxy_pass` to ensure the
`/shelfmark` prefix is removed.
- Health checks still work at `/api/health` without the subpath.
## Subpath setup
## Root path setup
Running Shelfmark under a subpath like `/shelfmark` is supported without extra rewrite rules.
If you can serve Shelfmark at the root path (`https://shelfmark.example.com/`),
leave `URL_BASE` empty. This is the simplest option.
### 1. Set the base path in Shelfmark
- **UI**: Settings → Advanced → Base Path → `/shelfmark/`
- **Environment variable**: `URL_BASE=/shelfmark/`
### 2. Configure your reverse proxy
All Shelfmark paths (UI, API, assets, Socket.IO) are served under the base path. A single location block is enough.
---
### Without Authentication Proxy
**Complete Nginx configuration for subpath deployment:**
```nginx
location /shelfmark/ {
proxy_pass http://shelfmark:8084/shelfmark/;
proxy_http_version 1.1;
proxy_set_header Host $host;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
proxy_set_header X-Forwarded-Proto $scheme;
proxy_set_header X-Forwarded-Host $host;
proxy_set_header Upgrade $http_upgrade;
proxy_set_header Connection "upgrade";
proxy_read_timeout 86400;
proxy_send_timeout 86400;
proxy_buffering off;
}
```
---
### With Authentication Proxy (Authelia, Authentik, etc.)
Shelfmark supports Proxy Authentication. When enabled, Shelfmark trusts the authenticated user from headers set by your auth proxy.
#### Shelfmark Settings
Configure in Settings → Security:
| Setting | Value |
|---------|-------|
| Authentication Method | Proxy Authentication |
| Proxy Auth User Header | `Remote-User` |
| Proxy Auth Logout URL | `https://auth.example.com/logout` |
| Proxy Auth Admin Group Header | `Remote-Groups` |
| Proxy Auth Admin Group Name | `admins` (or your admin group) |
#### Nginx Configuration with Authelia
This example uses Authelia snippets. Adapt for your auth proxy.
**Authelia auth request snippet** (`/etc/nginx/snippets/authelia-authrequest.conf`):
```nginx
location /authelia {
internal;
proxy_pass http://authelia:9091/api/authz/auth-request;
proxy_pass_request_body off;
proxy_set_header Content-Length "";
proxy_set_header Host $host;
proxy_set_header X-Original-URL $scheme://$http_host$request_uri;
proxy_set_header X-Original-Method $request_method;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
proxy_set_header X-Forwarded-Proto $scheme;
}
```
**Authelia location snippet** (`/etc/nginx/snippets/authelia-location.conf`):
```nginx
auth_request /authelia;
auth_request_set $target_url $scheme://$http_host$request_uri;
auth_request_set $user $upstream_http_remote_user;
auth_request_set $groups $upstream_http_remote_groups;
auth_request_set $name $upstream_http_remote_name;
auth_request_set $email $upstream_http_remote_email;
proxy_set_header Remote-User $user;
proxy_set_header Remote-Groups $groups;
proxy_set_header Remote-Name $name;
proxy_set_header Remote-Email $email;
error_page 401 =302 https://auth.example.com/?rd=$target_url;
```
**Complete Nginx configuration with Authelia:**
```nginx
# Include Authelia auth endpoint in your server block
include /etc/nginx/snippets/authelia-authrequest.conf;
# Main shelfmark location
location /shelfmark/ {
include /etc/nginx/snippets/authelia-location.conf;
proxy_pass http://shelfmark:8084/shelfmark/;
proxy_http_version 1.1;
proxy_set_header Host $host;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
proxy_set_header X-Forwarded-Proto $scheme;
proxy_set_header X-Forwarded-Host $host;
proxy_set_header Upgrade $http_upgrade;
proxy_set_header Connection "upgrade";
proxy_read_timeout 86400;
proxy_send_timeout 86400;
proxy_buffering off;
}
```
---
## Health checks
Health checks work at `/shelfmark/api/health` when using a subpath configuration.
+22 -12
View File
@@ -148,6 +148,7 @@ test_write() {
make_writable() {
folder=$1
did_full_chown=0
set +e
test_write $folder
is_writable=$?
@@ -158,31 +159,40 @@ make_writable() {
echo "Folder $folder is not writable, changing ownership"
change_ownership $folder
chmod -R g+r,g+w $folder || echo "Failed to change group permissions for ${folder}, continuing..."
did_full_chown=1
fi
# Fix any misowned subdirectories/files (e.g., from previous runs as root)
if [ -d "$folder" ]; then
misowned_count=$(find "$folder" -mindepth 1 \( ! -user "$RUN_UID" -o ! -group "$RUN_GID" \) 2>/dev/null | wc -l)
if [ "$misowned_count" -gt 0 ]; then
echo "Fixing ownership of $misowned_count files/directories in $folder"
find "$folder" -mindepth 1 \( ! -user "$RUN_UID" -o ! -group "$RUN_GID" \) \
-exec chown "$RUN_UID:$RUN_GID" {} \; 2>/dev/null || true
fi
if [ "$did_full_chown" -eq 0 ] && [ -d "$folder" ]; then
echo "Checking for misowned files/directories in $folder"
# Stay on the same filesystem to avoid traversing mounted subpaths
# (for example read-only bind mounts under /app in dev setups).
find "$folder" -xdev -mindepth 1 \( ! -user "$RUN_UID" -o ! -group "$RUN_GID" \) \
-exec chown "$RUN_UID:$RUN_GID" {} + 2>/dev/null || true
fi
test_write $folder || echo "Failed to test write to ${folder}, continuing..."
}
fix_misowned() {
folder=$1
mkdir -p $folder
echo "Checking for misowned files/directories in $folder"
# Stay on the same filesystem to avoid traversing mounted subpaths
# (for example read-only bind mounts under /app in dev setups).
find "$folder" -xdev \( ! -user "$RUN_UID" -o ! -group "$RUN_GID" \) \
-exec chown "$RUN_UID:$RUN_GID" {} + 2>/dev/null || true
}
# Ensure proper ownership of application directories
change_ownership() {
folder=$1
mkdir -p $folder
echo "Changing ownership of $folder to $USERNAME:$RUN_GID"
chown -R "${RUN_UID}" "${folder}" || echo "Failed to change user ownership for ${folder}, continuing..."
chown -R ":${RUN_GID}" "${folder}" || echo "Failed to change group ownership for ${folder}, continuing..."
chown -R "${RUN_UID}:${RUN_GID}" "${folder}" || echo "Failed to change ownership for ${folder}, continuing..."
}
change_ownership /app
change_ownership /var/log/shelfmark
change_ownership /tmp/shelfmark
fix_misowned /app
fix_misowned /var/log/shelfmark
fix_misowned /tmp/shelfmark
# SeleniumBase (internal bypasser) writes a patched chromedriver binary (uc_driver)
# into its own drivers directory. Some NAS/docker setups can apply restrictive ACLs
+1
View File
@@ -0,0 +1 @@
../baseline-browser-mapping/dist/cli.js
+17
View File
@@ -0,0 +1,17 @@
{
"name": "shelfmark",
"lockfileVersion": 3,
"requires": true,
"packages": {
"node_modules/baseline-browser-mapping": {
"version": "2.9.19",
"resolved": "https://registry.npmjs.org/baseline-browser-mapping/-/baseline-browser-mapping-2.9.19.tgz",
"integrity": "sha512-ipDqC8FrAl/76p2SSWKSI+H9tFwm7vYqXQrItCuiVPt26Km0jS+NzSsBWAaBusvSbQcfJG+JitdMm+wZAgTYqg==",
"dev": true,
"license": "Apache-2.0",
"bin": {
"baseline-browser-mapping": "dist/cli.js"
}
}
}
}
+201
View File
@@ -0,0 +1,201 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright [yyyy] [name of copyright owner]
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
+463
View File
@@ -0,0 +1,463 @@
# [`baseline-browser-mapping`](https://github.com/web-platform-dx/web-features/packages/baseline-browser-mapping)
By the [W3C WebDX Community Group](https://www.w3.org/community/webdx/) and contributors.
`baseline-browser-mapping` provides:
- An `Array` of browsers compatible with Baseline Widely available and Baseline year feature sets via the [`getCompatibleVersions()` function](#get-baseline-widely-available-browser-versions-or-baseline-year-browser-versions).
- An `Array`, `Object` or `CSV` as a string describing the Baseline feature set support of all browser versions included in the module's data set via the [`getAllVersions()` function](#get-data-for-all-browser-versions).
You can use `baseline-browser-mapping` to help you determine minimum browser version support for your chosen Baseline feature set; or to analyse the level of support for different Baseline feature sets in your site's traffic by joining the data with your analytics data.
## Install for local development
To install the package, run:
`npm install --save-dev baseline-browser-mapping`
`baseline-browser-mapping` depends on `web-features` and `@mdn/browser-compat-data` for core browser version selection, but the data is pre-packaged and minified. This package checks for updates to those modules and the supported [downstream browsers](#downstream-browsers) on a daily basis and is updated frequently. Consider adding a script to your `package.json` to update `baseline-browser-mapping` and using it as part of your build process to ensure your data is as up to date as possible:
```javascript
"scripts": [
"refresh-baseline-browser-mapping": "npm i --save-dev baseline-browser-mapping@latest"
]
```
The minimum supported NodeJS version for `baseline-browser-mapping` is v8 in alignment with `browserslist`. For NodeJS versions earlier than v13.2, the [`require('baseline-browser-mapping')`](https://nodejs.org/api/modules.html#requireid) syntax should be used to import the module.
## Keeping `baseline-browser-mapping` up to date
If you are only using this module to generate minimum browser versions for Baseline Widely available or Baseline year feature sets, you don't need to update this module frequently, as the backward looking data is reasonably stable.
However, if you are targeting Newly available, using the [`getAllVersions()`](#get-data-for-all-browser-versions) function or heavily relying on the data for downstream browsers, you should update this module more frequently. If you target a feature cut off date within the last two months and your installed version of `baseline-browser-mapping` has data that is more than 2 months old, you will receive a console warning advising you to update to the latest version when you call `getCompatibleVersions()` or `getAllVersions()`.
If you want to suppress these warnings you can use the `suppressWarnings: true` option in the configuration object passed to `getCompatibleVersions()` or `getAllVersions()`. Alternatively, you can use the `BASELINE_BROWSER_MAPPING_IGNORE_OLD_DATA=true` environment variable when running your build process. This module also respects the `BROWSERSLIST_IGNORE_OLD_DATA=true` environment variable. Environment variables can also be provided in a `.env` file from Node 20 onwards; however, this module does not load .env files automatically to avoid conflicts with other libraries with different requirements. You will need to use `process.loadEnvFile()` or a library like `dotenv` to load .env files before `baseline-browser-mapping` is called.
If you want to ensure [reproducible builds](https://www.wikiwand.com/en/articles/Reproducible_builds), we strongly recommend using the `widelyAvailableOnDate` option to fix the Widely available date on a per build basis to ensure dependent tools provide the same output and you do not produce data staleness warnings. If you are using [`browserslist`](https://github.com/browserslist/browserslist) to target Baseline Widely available, consider automatically updating your `browserslist` configuration in `package.json` or `.browserslistrc` to `baseline widely available on {YYYY-MM-DD}` as part of your build process to ensure the same or sufficiently similar list of minimum browsers is reproduced for historical builds.
## Importing `baseline-browser-mapping`
This module exposes two functions: `getCompatibleVersions()` and `getAllVersions()`, both which can be imported directly from `baseline-browser-mapping`:
```javascript
import {
getCompatibleVersions,
getAllVersions,
} from "baseline-browser-mapping";
```
If you want to load the script and data directly in a web page without hosting it yourself, consider using a CDN:
```html
<script type="module">
import {
getCompatibleVersions,
getAllVersions,
} from "https://cdn.jsdelivr.net/npm/baseline-browser-mapping";
</script>
```
## Get Baseline Widely available browser versions or Baseline year browser versions
To get the current list of minimum browser versions compatible with Baseline Widely available features from the core browser set, call the `getCompatibleVersions()` function:
```javascript
getCompatibleVersions();
```
Executed on 7th March 2025, the above code returns the following browser versions:
```javascript
[
{ browser: "chrome", version: "105", release_date: "2022-09-02" },
{
browser: "chrome_android",
version: "105",
release_date: "2022-09-02",
},
{ browser: "edge", version: "105", release_date: "2022-09-02" },
{ browser: "firefox", version: "104", release_date: "2022-08-23" },
{
browser: "firefox_android",
version: "104",
release_date: "2022-08-23",
},
{ browser: "safari", version: "15.6", release_date: "2022-09-02" },
{
browser: "safari_ios",
version: "15.6",
release_date: "2022-09-02",
},
];
```
> [!NOTE]
> The minimum versions of each browser are not necessarily the final release before the Widely available cutoff date of `TODAY - 30 MONTHS`. Some earlier versions will have supported the full Widely available feature set.
### `getCompatibleVersions()` configuration options
`getCompatibleVersions()` accepts an `Object` as an argument with configuration options. The defaults are as follows:
```javascript
{
targetYear: undefined,
widelyAvailableOnDate: undefined,
includeDownstreamBrowsers: false,
listAllCompatibleVersions: false,
suppressWarnings: false
}
```
#### `targetYear`
The `targetYear` option returns the minimum browser versions compatible with all **Baseline Newly available** features at the end of the specified calendar year. For example, calling:
```javascript
getCompatibleVersions({
targetYear: 2020,
});
```
Returns the following versions:
```javascript
[
{ browser: "chrome", version: "87", release_date: "2020-11-19" },
{
browser: "chrome_android",
version: "87",
release_date: "2020-11-19",
},
{ browser: "edge", version: "87", release_date: "2020-11-19" },
{ browser: "firefox", version: "83", release_date: "2020-11-17" },
{
browser: "firefox_android",
version: "83",
release_date: "2020-11-17",
},
{ browser: "safari", version: "14", release_date: "2020-09-16" },
{ browser: "safari_ios", version: "14", release_date: "2020-09-16" },
];
```
> [!NOTE]
> The minimum version of each browser is not necessarily the final version released in that calendar year. In the above example, Firefox 84 was the final version released in 2020; however Firefox 83 supported all of the features that were interoperable at the end of 2020.
> [!WARNING]
> You cannot use `targetYear` and `widelyAavailableDate` together. Please only use one of these options at a time.
#### `widelyAvailableOnDate`
The `widelyAvailableOnDate` option returns the minimum versions compatible with Baseline Widely available on a specified date in the format `YYYY-MM-DD`:
```javascript
getCompatibleVersions({
widelyAvailableOnDate: `2023-04-05`,
});
```
> [!TIP]
> This option is useful if you provide a versioned library that targets Baseline Widely available on each version's release date and you need to provide a statement on minimum supported browser versions in your documentation.
#### `includeDownstreamBrowsers`
Setting `includeDownstreamBrowsers` to `true` will include browsers outside of the Baseline core browser set where it is possible to map those browsers to an upstream Chromium or Gecko version:
```javascript
getCompatibleVersions({
includeDownstreamBrowsers: true,
});
```
For more information on downstream browsers, see [the section on downstream browsers](#downstream-browsers) below.
#### `includeKaiOS`
KaiOS is an operating system and app framework based on the Gecko engine from Firefox. KaiOS is based on the Gecko engine and feature support can be derived from the upstream Gecko version that each KaiOS version implements. However KaiOS requires other considerations beyond feature compatibility to ensure a good user experience as it runs on device types that do not have either mouse and keyboard or touch screen input in the way that all the other browsers supported by this module do.
```javascript
getCompatibleVersions({
includeDownstreamBrowsers: true,
includeKaiOS: true,
});
```
> [!NOTE]
> Including KaiOS requires you to include all downstream browsers using the `includeDownstreamBrowsers` option.
#### `listAllCompatibleVersions`
Setting `listAllCompatibleVersions` to true will include the minimum versions of each compatible browser, and all the subsequent versions:
```javascript
getCompatibleVersions({
listAllCompatibleVersions: true,
});
```
#### `suppressWarnings`
Setting `suppressWarnings` to `true` will suppress the console warning about old data:
```javascript
getCompatibleVersions({
suppressWarnings: true,
});
```
## Get data for all browser versions
You may want to obtain data on all the browser versions available in this module for use in an analytics solution or dashboard. To get details of each browser version's level of Baseline support, call the `getAllVersions()` function:
```javascript
import { getAllVersions } from "baseline-browser-mapping";
getAllVersions();
```
By default, this function returns an `Array` of `Objects` and excludes downstream browsers:
```javascript
[
...
{
browser: "firefox_android", // Browser name
version: "125", // Browser version
release_date: "2024-04-16", // Release date
year: 2023, // Baseline year feature set the version supports
wa_compatible: true // Whether the browser version supports Widely available
},
...
]
```
For browser versions in `@mdn/browser-compat-data` that were released before Baseline can be defined, i.e. Baseline 2015, the `year` property is always the string: `"pre_baseline"`.
### Understanding which browsers support Newly available features
You may want to understand which recent browser versions support all Newly available features. You can replace the `wa_compatible` property with a `supports` property using the `useSupport` option:
```javascript
getAllVersions({
useSupports: true,
});
```
The `supports` property is optional and has two possible values:
- `widely` for browser versions that support all Widely available features.
- `newly` for browser versions that support all Newly available features.
Browser versions that do not support Widely or Newly available will not include the `support` property in the `array` or `object` outputs, and in the CSV output, the `support` column will contain an empty string. Browser versions that support all Newly available features also support all Widely available features.
### `getAllVersions()` Configuration options
`getAllVersions()` accepts an `Object` as an argument with configuration options. The defaults are as follows:
```javascript
{
includeDownstreamBrowsers: false,
outputFormat: "array",
suppressWarnings: false
}
```
#### `includeDownstreamBrowsers` (in `getAllVersions()` output)
As with `getCompatibleVersions()`, you can set `includeDownstreamBrowsers` to `true` to include the Chromium and Gecko downstream browsers [listed below](#list-of-downstream-browsers).
```javascript
getAllVersions({
includeDownstreamBrowsers: true,
});
```
Downstream browsers include the same properties as core browsers, as well as the `engine`they use and `engine_version`, for example:
```javascript
[
...
{
browser: "samsunginternet_android",
version: "27.0",
release_date: "2024-11-06",
engine: "Blink",
engine_version: "125",
year: 2023,
supports: "widely"
},
...
]
```
#### `includeKaiOS` (in `getAllVersions()` output)
As with `getCompatibleVersions()` you can include KaiOS in your output. The same requirement to have `includeDownstreamBrowsers: true` applies.
```javascript
getAllVersions({
includeDownstreamBrowsers: true,
includeKaiOS: true,
});
```
#### `suppressWarnings` (in `getAllVersions()` output)
As with `getCompatibleVersions()`, you can set `suppressWarnings` to `true` to suppress the console warning about old data:
```javascript
getAllVersions({
suppressWarnings: true,
});
```
#### `outputFormat`
By default, this function returns an `Array` of `Objects` which can be manipulated in Javascript or output to JSON.
To return an `Object` that nests keys , set `outputFormat` to `object`:
```javascript
getAllVersions({
outputFormat: "object",
});
```
In thise case, `getAllVersions()` returns a nested object with the browser [IDs listed below](#list-of-downstream-browsers) as keys, and versions as keys within them:
```javascript
{
"chrome": {
"53": {
"year": 2016,
"release_date": "2016-09-07"
},
...
}
```
Downstream browsers will include extra fields for `engine` and `engine_versions`
```javascript
{
...
"webview_android": {
"53": {
"year": 2016,
"release_date": "2016-09-07",
"engine": "Blink",
"engine_version": "53"
},
...
}
```
To return a `String` in CSV format, set `outputFormat` to `csv`:
```javascript
getAllVersions({
outputFormat: "csv",
});
```
`getAllVersions` returns a `String` with a header row and comma-separated values for each browser version that you can write to a file or pass to another service. Core browsers will have "NULL" as the value for their `engine` and `engine_version`:
```csv
"browser","version","year","supports","release_date","engine","engine_version"
...
"chrome","24","pre_baseline","","2013-01-10","NULL","NULL"
...
"chrome","53","2016","","2016-09-07","NULL","NULL"
...
"firefox","135","2024","widely","2025-02-04","NULL","NULL"
"firefox","136","2024","newly","2025-03-04","NULL","NULL"
...
"ya_android","20.12","2020","year_only","2020-12-20","Blink","87"
...
```
> [!NOTE]
> The above example uses `"includeDownstreamBrowsers": true`
### Static resources
The outputs of `getAllVersions()` are available as JSON or CSV files generated on a daily basis and hosted on GitHub pages:
- Core browsers only
- [Array](https://web-platform-dx.github.io/baseline-browser-mapping/all_versions_array.json)
- [Object](https://web-platform-dx.github.io/baseline-browser-mapping/all_versions_object.json)
- [CSV](https://web-platform-dx.github.io/baseline-browser-mapping/all_versions.csv)
- Core browsers only, with `supports` property
- [Array](https://web-platform-dx.github.io/baseline-browser-mapping/all_versions_array_with_supports.json)
- [Object](https://web-platform-dx.github.io/baseline-browser-mapping/all_versions_object_with_supports.json)
- [CSV](https://web-platform-dx.github.io/baseline-browser-mapping/all_versions_with_supports.csv)
- Including downstream browsers
- [Array](https://web-platform-dx.github.io/baseline-browser-mapping/with_downstream/all_versions_array.json)
- [Object](https://web-platform-dx.github.io/baseline-browser-mapping/with_downstream/all_versions_object.json)
- [CSV](https://web-platform-dx.github.io/baseline-browser-mapping/with_downstream/all_versions.csv)
- Including downstream browsers with `supports` property
- [Array](https://web-platform-dx.github.io/baseline-browser-mapping/with_downstream/all_versions_array_with_supports.json)
- [Object](https://web-platform-dx.github.io/baseline-browser-mapping/with_downstream/all_versions_object_with_supports.json)
- [CSV](https://web-platform-dx.github.io/baseline-browser-mapping/with_downstream/all_versions_with_supports.csv)
These files are updated on a daily basis.
## CLI
`baseline-browser-mapping` includes a command line interface that exposes the same data and options as the `getCompatibleVersions()` function. To learn more about using the CLI, run:
```sh
npx baseline-browser-mapping --help
```
## Downstream browsers
### Limitations
The browser versions in this module come from two different sources:
- MDN's `browser-compat-data` module.
- Parsed user agent strings provided by [useragents.io](https://useragents.io/)
MDN `browser-compat-data` is an authoritative source of information for the browsers it contains. The release dates for the Baseline core browser set and the mapping of downstream browsers to Chromium versions should be considered accurate.
Browser mappings from useragents.io are provided on a best effort basis. They assume that browser vendors are accurately stating the Chromium version they have implemented. The initial set of version mappings was derived from a bulk export in November 2024. This version was iterated over with a Regex match looking for a major Chrome version and a corresponding version of the browser in question, e.g.:
`Mozilla/5.0 (Linux; U; Android 10; en-US; STK-L21 Build/HUAWEISTK-L21) AppleWebKit/537.36 (KHTML, like Gecko) Version/4.0 Chrome/100.0.4896.58 UCBrowser/13.8.2.1324 Mobile Safari/537.36`
Shows UC Browser Mobile 13.8 implementing Chromium 100, and:
`Mozilla/5.0 (Linux; arm_64; Android 11; Redmi Note 8 Pro) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/128.0.6613.123 YaBrowser/24.10.2.123.00 SA/3 Mobile Safari/537.36`
Shows Yandex Browser Mobile 24.10 implementing Chromium 128. The Chromium version from this string is mapped to the corresponding Chrome version from MDN `browser-compat-data`.
> [!NOTE]
> Where possible, approximate release dates have been included based on useragents.io "first seen" data. useragents.io does not have "first seen" dates prior to June 2020. However, these browsers' Baseline compatibility is determined by their Chromium or Gecko version, so their release dates are more informative than critical.
This data is updated on a daily basis using a [script](https://github.com/web-platform-dx/web-features/tree/main/scripts/refresh-downstream.ts) triggered by a GitHub [action](https://github.com/web-platform-dx/web-features/tree/main/.github/workflows/refresh_downstream.yml). Useragents.io provides a private API for this module which exposes the last 7 days of newly seen user agents for the currently tracked browsers. If a new major version of one of the tracked browsers is encountered with a Chromium version that meets or exceeds the previous latest version of that browser, it is added to the [src/data/downstream-browsers.json](src/data/downstream-browsers.json) file with the date it was first seen by useragents.io as its release date.
KaiOS is an exception - its upstream version mappings are handled separately from the other browsers because they happen very infrequently.
### List of downstream browsers
| Browser | ID | Core | Source |
| --------------------- | ------------------------- | ------- | ------------------------- |
| Chrome | `chrome` | `true` | MDN `browser-compat-data` |
| Chrome for Android | `chrome_android` | `true` | MDN `browser-compat-data` |
| Edge | `edge` | `true` | MDN `browser-compat-data` |
| Firefox | `firefox` | `true` | MDN `browser-compat-data` |
| Firefox for Android | `firefox_android` | `true` | MDN `browser-compat-data` |
| Safari | `safari` | `true` | MDN `browser-compat-data` |
| Safari on iOS | `safari_ios` | `true` | MDN `browser-compat-data` |
| Opera | `opera` | `false` | MDN `browser-compat-data` |
| Opera Android | `opera_android` | `false` | MDN `browser-compat-data` |
| Samsung Internet | `samsunginternet_android` | `false` | MDN `browser-compat-data` |
| WebView Android | `webview_android` | `false` | MDN `browser-compat-data` |
| QQ Browser Mobile | `qq_android` | `false` | useragents.io |
| UC Browser Mobile | `uc_android` | `false` | useragents.io |
| Yandex Browser Mobile | `ya_android` | `false` | useragents.io |
| KaiOS | `kai_os` | `false` | Manual |
| Facebook for Android | `facebook_android` | `false` | useragents.io |
| Instagram for Android | `instagram_android` | `false` | useragents.io |
> [!NOTE]
> All the non-core browsers currently included implement Chromium or Gecko. Their inclusion in any of the above methods is based on the Baseline feature set supported by the Chromium or Gecko version they implement, not their release date.
+64
View File
@@ -0,0 +1,64 @@
{
"name": "baseline-browser-mapping",
"main": "./dist/index.cjs",
"version": "2.9.19",
"description": "A library for obtaining browser versions with their maximum supported Baseline feature set and Widely Available status.",
"exports": {
".": {
"require": "./dist/index.cjs",
"types": "./dist/index.d.ts",
"default": "./dist/index.js"
},
"./legacy": {
"require": "./dist/index.cjs",
"types": "./dist/index.d.ts"
}
},
"jsdelivr": "./dist/index.js",
"files": [
"dist/*",
"!dist/scripts/*",
"LICENSE.txt",
"README.md"
],
"types": "./dist/index.d.ts",
"type": "module",
"bin": {
"baseline-browser-mapping": "dist/cli.js"
},
"scripts": {
"fix-cli-permissions": "output=$(npx baseline-browser-mapping 2>&1); path=$(printf '%s\n' \"$output\" | sed -n 's/^.*: \\(.*\\): Permission denied$/\\1/p; t; s/^\\(.*\\): Permission denied$/\\1/p'); if [ -n \"$path\" ]; then echo \"Permission denied for: $path\"; echo \"Removing $path ...\"; rm -rf \"$path\"; else echo \"$output\"; fi",
"test:format": "npx prettier --check .",
"test:lint": "npx eslint .",
"test:jasmine": "npx jasmine",
"test:jasmine-browser": "npx jasmine-browser-runner runSpecs --config ./spec/support/jasmine-browser.js",
"test": "npm run build && npm run fix-cli-permissions && npm run test:format && npm run test:lint && npm run test:jasmine && npm run test:jasmine-browser",
"build": "rm -rf dist; npx prettier . --write; rollup -c; rm -rf ./dist/scripts/expose-data.d.ts ./dist/cli.d.ts",
"refresh-downstream": "npx tsx scripts/refresh-downstream.ts",
"refresh-static": "npx tsx scripts/refresh-static.ts",
"update-data-file": "npx tsx scripts/update-data-file.ts; npx prettier ./src/data/data.js --write",
"update-data-dependencies": "npm i @mdn/browser-compat-data@latest web-features@latest -D",
"check-data-changes": "git diff --name-only | grep -q '^src/data/data.js$' && echo 'changes-available=TRUE' || echo 'changes-available=FALSE'"
},
"license": "Apache-2.0",
"devDependencies": {
"@mdn/browser-compat-data": "^7.2.5",
"@rollup/plugin-terser": "^0.4.4",
"@rollup/plugin-typescript": "^12.1.3",
"@types/node": "^22.15.17",
"eslint-plugin-new-with-error": "^5.0.0",
"jasmine": "^5.8.0",
"jasmine-browser-runner": "^3.0.0",
"jasmine-spec-reporter": "^7.0.0",
"prettier": "^3.5.3",
"rollup": "^4.44.0",
"tslib": "^2.8.1",
"typescript": "^5.7.2",
"typescript-eslint": "^8.35.0",
"web-features": "^3.14.0"
},
"repository": {
"type": "git",
"url": "git+https://github.com/web-platform-dx/baseline-browser-mapping.git"
}
}
+22
View File
@@ -0,0 +1,22 @@
{
"name": "shelfmark",
"lockfileVersion": 3,
"requires": true,
"packages": {
"": {
"devDependencies": {
"baseline-browser-mapping": "^2.9.19"
}
},
"node_modules/baseline-browser-mapping": {
"version": "2.9.19",
"resolved": "https://registry.npmjs.org/baseline-browser-mapping/-/baseline-browser-mapping-2.9.19.tgz",
"integrity": "sha512-ipDqC8FrAl/76p2SSWKSI+H9tFwm7vYqXQrItCuiVPt26Km0jS+NzSsBWAaBusvSbQcfJG+JitdMm+wZAgTYqg==",
"dev": true,
"license": "Apache-2.0",
"bin": {
"baseline-browser-mapping": "dist/cli.js"
}
}
}
}
+5
View File
@@ -0,0 +1,5 @@
{
"devDependencies": {
"baseline-browser-mapping": "^2.9.19"
}
}
+8 -2
View File
@@ -6,7 +6,12 @@ Formerly *Calibre Web Automated Book Downloader (CWABD)*
Shelfmark is a unified web interface for searching and aggregating books and audiobook downloads from multiple sources - all in one place. Works out of the box with popular web sources, no configuration required. Add metadata providers, additional release sources, and download clients to create a single hub for building your digital library.
**Fully standalone** - no external dependencies required. Works great alongside library tools like [Calibre-Web-Automated](https://github.com/crocodilestick/Calibre-Web-Automated), [Booklore](https://github.com/booklore-app/booklore) or [Audiobookshelf](https://github.com/advplyr/audiobookshelf) for automatic import.
**Fully standalone** - no external dependencies required. Works great alongside the following library tools, with support for automatic imports:
- [Calibre](https://calibre-ebook.com/)
- [Calibre-Web](https://github.com/janeczku/calibre-web)
- [Calibre-Web-Automated](https://github.com/crocodilestick/Calibre-Web-Automated)
- [Booklore](https://github.com/booklore-app/booklore)
- [Audiobookshelf](https://github.com/advplyr/audiobookshelf)
## ✨ Features
@@ -101,6 +106,7 @@ See the full [Environment Variables Reference](docs/environment-variables.md) fo
Some of the additional options available in Settings:
- **Fast Download Key** - Use your paid account to skip Cloudflare challenges entirely and use faster, direct downloads
- **Prowlarr** - Configure indexers and download clients to download books and audiobooks
- **AudiobookBay** - Web scraping source for audiobook torrents (audiobooks only)
- **IRC** - Add details for IRC book sources and download directly from the UI
- **Library Link** - Add a link to your Calibre-Web or Booklore instance in the UI header
- **File processing** - Customiseable download paths, file renaming and directory creation with template-based renaming
@@ -134,7 +140,7 @@ A smaller image without the built-in Cloudflare bypasser. Ideal for:
- **External bypassers** - Already running FlareSolverr or ByParr for other services
- **Fast downloads** - Using fast download sources
- **Alternative sources only** - Exclusively using Prowlarr, IRC, or other sources
- **Alternative sources only** - Exclusively using Prowlarr, AudiobookBay, IRC, or other sources
- **Audiobooks** - Using Shelfmark exclusively for audiobooks
```bash
+47
View File
@@ -0,0 +1,47 @@
## New Features
### OIDC Authentication (#606, #612)
- **OIDC login** with PKCE flow, auto-discovery, and group-based admin mapping
- **Auto-provisioning** of OIDC users (configurable) and email-based account linking
- **Password fallback** when OIDC is enabled to prevent admin lockout
- Backwards compatible with all existing auth modes (no-auth, builtin, proxy, CWA)
### Multi-User Support (#606, #612, #613)
- **User management** -create, edit, and delete users with admin/user roles
- **Per-user settings** -custom download destinations, BookLore library/path, email recipients, and `{User}` template variable
- **Per-user download visibility** -non-admins only see their own downloads
### Multi-User Request System (#615, #617, #620)
- **Book request workflow** -users can request books with notes; admins review, approve, and fulfil requests
- **Policy-based configuration** -set download/request/block policies per content type or per source (e.g. allow direct downloads, set Prowlarr to request-only)
- **Per-user policy overrides** for tailored access control
- **New Activity Sidebar** -replaces downloads sidebar, combining active downloads with requests; sidebar can now be pinned
- Request retry support and admin-level request management
### Notification Support (#618)
- **Apprise-based notifications** for request events and download completions
- Configurable globally or per user, with full customization of events and notification services
- Expanded activity cards with detailed request info and file management
### AudiobookBay Release Source (#619, #621, #623)
- **New release source** -search AudiobookBay for audiobook torrents directly from the UI
- Results include title, language, format, and size
- Downloads via configured torrent client with audiobook-specific category support
- Configurable hostname, max search pages, and rate limit delay
### Email Output Mode (#603, #604)
- **Email delivery** as an alternative output mode for downloaded books
- Per-user email recipient configuration
## Improvements
- Admin-configurable visibility for self-settings options (delivery preferences, notifications) (#625)
- BookLore Bookdrop API destination support as an alternative to specific library selection (#625)
- Download path options for all torrent clients (#625)
- Add tag support to qBittorrent downloads (#610 by @dawescc)
- Add threading to file system operations for improved performance (#602)
- Enhanced custom scripting -JSON download info, more consistent activation, decoupled from staging (#591)
- Hardlink-before-move optimization for file transfers (#591)
- New BookLore API file formats (#591)
- Improved login cookie naming for reverse proxy compatibility (#591)
- Fix Transmission URL parsing (#591)
- Fix healthcheck starvation during large file processing (#591)
+2
View File
@@ -14,3 +14,5 @@ emoji
rarfile
qbittorrent-api
transmission-rpc
authlib>=1.6.6,<1.7
apprise>=1.9.0
-1
View File
@@ -1,5 +1,4 @@
pyvirtualdisplay
pyautogui
selenium==4.39.0
seleniumbase==4.45.10
python-xlib
+6
View File
@@ -178,6 +178,12 @@ def _generate_bootstrap_env_docs() -> List[str]:
"type": "boolean",
"default": "false",
},
{
"name": "ONBOARDING",
"description": "Show the onboarding wizard on first run. Set to false to skip (useful for ephemeral storage).",
"type": "boolean",
"default": "true",
},
]
lines = [
+20 -2
View File
@@ -434,6 +434,10 @@ def test_rtorrent():
version = client.system.library_version()
print(f" Connected to rTorrent {version}")
# default download directory test
default_dir = client.directory.default()
print(f" Default download directory: {default_dir}")
# Get torrent list
torrents = client.download_list()
print(f" Active torrents: {len(torrents)}")
@@ -458,10 +462,13 @@ def test_rtorrent():
torrent_id = "3B245504CF5F11BBDBE1201CEA6A6BF45AEE1BC0" # rtorrent uses uppercase hashes
print(f" Added test torrent: {torrent_id}")
torrents = client.download_list()
print(f" Active torrents: {len(torrents)}")
torrent_list = client.d.multicall.filtered(
"",
"default",
f"equal=d.hash=,cat={torrent_id}",
f"equal={{d.hash=,cat={torrent_id}}}"
"d.hash=",
"d.state=",
"d.completed_bytes=",
@@ -472,11 +479,22 @@ def test_rtorrent():
"d.complete=",
)
torrent = torrent_list[0]
if not torrent:
print(" ERROR: Could not find added torrent in list")
return False
# let's test the base path call
details = client.d.multicall.filtered(
"",
"default",
f"equal=d.hash=,cat={torrent_id}",
"d.base_path=",
)
base_path = details[0][0] if details else None
print(f" Base path: {base_path}")
client.d.erase(torrent_id)
print(" Removed test torrent")
+85 -8
View File
@@ -4,7 +4,7 @@ import logging
import threading
from typing import Optional, Dict, Any, Callable, List
from flask_socketio import SocketIO
from flask_socketio import SocketIO, join_room, leave_room
logger = logging.getLogger(__name__)
@@ -20,6 +20,10 @@ class WebSocketManager:
self._on_first_connect_callbacks: List[Callable[[], None]] = []
self._on_all_disconnect_callbacks: List[Callable[[], None]] = []
self._needs_rewarm = False # Flag to trigger warmup callbacks on next connect
self._user_rooms: Dict[str, int] = {} # room_name -> ref count
self._sid_rooms: Dict[str, str] = {} # sid -> room_name
self._rooms_lock = threading.Lock()
self._queue_status_fn: Optional[Callable] = None # Reference to queue_status()
def init_app(self, app, socketio: SocketIO):
"""Initialize the WebSocket manager with Flask-SocketIO instance."""
@@ -100,19 +104,86 @@ class WebSocketManager:
"""Check if WebSocket is enabled and ready."""
return self._enabled and self.socketio is not None
def set_queue_status_fn(self, fn: Callable):
"""Set the queue_status function reference for per-room filtering."""
self._queue_status_fn = fn
def _increment_user_room_locked(self, room: str):
self._user_rooms[room] = self._user_rooms.get(room, 0) + 1
def _decrement_user_room_locked(self, room: str):
count = self._user_rooms.get(room, 1) - 1
if count <= 0:
self._user_rooms.pop(room, None)
else:
self._user_rooms[room] = count
def _set_sid_room_locked(self, sid: str, room: Optional[str]):
current_room = self._sid_rooms.get(sid)
if current_room == room:
return
if current_room is not None:
leave_room(current_room, sid=sid)
if current_room.startswith("user_"):
self._decrement_user_room_locked(current_room)
self._sid_rooms.pop(sid, None)
if room is not None:
join_room(room, sid=sid)
self._sid_rooms[sid] = room
if room.startswith("user_"):
self._increment_user_room_locked(room)
def sync_user_room(self, sid: str, is_admin: bool, db_user_id: Optional[int] = None):
"""Ensure a SID is in exactly one room matching the current session scope."""
room: Optional[str] = None
if is_admin:
room = "admins"
elif db_user_id is not None:
room = f"user_{db_user_id}"
with self._rooms_lock:
self._set_sid_room_locked(sid, room)
def join_user_room(self, sid: str, is_admin: bool, db_user_id: Optional[int] = None):
"""Join the appropriate room based on user role."""
self.sync_user_room(sid, is_admin, db_user_id)
def leave_user_room(self, sid: str, is_admin: bool = False, db_user_id: Optional[int] = None):
"""Leave whichever room the SID currently belongs to."""
del is_admin, db_user_id # Backward-compatible signature; routing is SID-based.
with self._rooms_lock:
self._set_sid_room_locked(sid, None)
def broadcast_status_update(self, status_data: Dict[str, Any]):
"""Broadcast status update to all connected clients."""
"""Broadcast status update to all connected clients, filtered by user room."""
if not self.is_enabled():
return
try:
# When calling socketio.emit() outside event handlers, it broadcasts by default
self.socketio.emit('status_update', status_data)
logger.debug(f"Broadcasted status update to all clients")
# Admins (and no-auth users) get full status
self.socketio.emit('status_update', status_data, to="admins")
# Each user room gets filtered status
with self._rooms_lock:
active_rooms = list(self._user_rooms.keys())
if active_rooms and self._queue_status_fn:
for room in active_rooms:
try:
# Extract user_id from room name "user_123"
uid = int(room.split("_", 1)[1])
filtered = self._queue_status_fn(user_id=uid)
self.socketio.emit('status_update', filtered, to=room)
except Exception as e:
logger.error(f"Failed to send status update for room {room}: {e}")
logger.debug("Broadcasted status update to all rooms")
except Exception as e:
logger.error(f"Error broadcasting status update: {e}")
def broadcast_download_progress(self, book_id: str, progress: float, status: str):
def broadcast_download_progress(self, book_id: str, progress: float, status: str, user_id: Optional[int] = None):
"""Broadcast download progress update for a specific book."""
if not self.is_enabled():
return
@@ -123,8 +194,14 @@ class WebSocketManager:
'progress': progress,
'status': status
}
# When calling socketio.emit() outside event handlers, it broadcasts by default
self.socketio.emit('download_progress', data)
# Admins always see all progress
self.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)
logger.debug(f"Broadcasted progress for book {book_id}: {progress}%")
except Exception as e:
logger.error(f"Error broadcasting download progress: {e}")
File diff suppressed because it is too large Load Diff
+42
View File
@@ -0,0 +1,42 @@
from __future__ import annotations
from typing import Any
from shelfmark.core.config import config
from shelfmark.download.outputs.email import EmailOutputError, build_email_smtp_config, test_smtp_connection
def test_email_connection(current_values: dict[str, Any] | None = None) -> dict[str, Any]:
"""Test SMTP connectivity using current form values (including unsaved changes)."""
current_values = current_values or {}
def _get_value(key: str, default: Any = None) -> Any:
value = current_values.get(key)
if value not in (None, ""):
return value
if default is None:
return config.get(key)
return config.get(key, default)
settings = {
"EMAIL_SMTP_HOST": _get_value("EMAIL_SMTP_HOST", ""),
"EMAIL_SMTP_PORT": _get_value("EMAIL_SMTP_PORT", 587),
"EMAIL_SMTP_SECURITY": _get_value("EMAIL_SMTP_SECURITY", "starttls"),
"EMAIL_SMTP_USERNAME": _get_value("EMAIL_SMTP_USERNAME", ""),
"EMAIL_SMTP_PASSWORD": _get_value("EMAIL_SMTP_PASSWORD", ""),
"EMAIL_FROM": _get_value("EMAIL_FROM", ""),
"EMAIL_SUBJECT_TEMPLATE": _get_value("EMAIL_SUBJECT_TEMPLATE", "{Title}"),
"EMAIL_SMTP_TIMEOUT_SECONDS": _get_value("EMAIL_SMTP_TIMEOUT_SECONDS", 60),
"EMAIL_ALLOW_UNVERIFIED_TLS": _get_value("EMAIL_ALLOW_UNVERIFIED_TLS", False),
}
try:
smtp_config = build_email_smtp_config(settings)
test_smtp_connection(smtp_config)
return {"success": True, "message": "Connected to SMTP server"}
except EmailOutputError as exc:
return {"success": False, "message": str(exc)}
except Exception as exc:
return {"success": False, "message": f"SMTP test failed: {exc}"}
+9
View File
@@ -113,6 +113,7 @@ FLASK_PORT = int(os.getenv("FLASK_PORT", "8084"))
# =============================================================================
SESSION_COOKIE_SECURE_ENV = os.getenv("SESSION_COOKIE_SECURE", "false")
SESSION_COOKIE_NAME = "shelfmark_session"
CWA_DB_PATH = _resolve_cwa_db_path()
@@ -133,6 +134,14 @@ TOR_VARIANT_AVAILABLE = shutil.which("tor") is not None
USING_TOR = string_to_bool(os.getenv("USING_TOR", "false"))
# =============================================================================
# Onboarding
# =============================================================================
# Set to false to skip the onboarding wizard entirely (useful for ephemeral storage)
ONBOARDING = string_to_bool(os.getenv("ONBOARDING", "true"))
# =============================================================================
# Debug/development settings
# =============================================================================
+123
View File
@@ -0,0 +1,123 @@
"""Configuration migration helpers."""
import json
from typing import Any, Callable
_DEPRECATED_SETTINGS_RESTRICTION_KEYS = (
"PROXY_AUTH_RESTRICT_SETTINGS_TO_ADMIN",
"CWA_RESTRICT_SETTINGS_TO_ADMIN",
"RESTRICT_SETTINGS_TO_ADMIN",
)
def _as_bool(value: Any) -> bool:
if isinstance(value, bool):
return value
if isinstance(value, str):
return value.strip().lower() in {"1", "true", "yes", "on"}
return bool(value)
def _pick_legacy_settings_restriction(config: dict[str, Any]) -> bool | None:
"""Pick the best legacy admin-restriction value to migrate."""
auth_method = str(config.get("AUTH_METHOD", "")).strip().lower()
if (
auth_method == "proxy"
and "PROXY_AUTH_RESTRICT_SETTINGS_TO_ADMIN" in config
):
return _as_bool(config.get("PROXY_AUTH_RESTRICT_SETTINGS_TO_ADMIN"))
if auth_method == "cwa" and "CWA_RESTRICT_SETTINGS_TO_ADMIN" in config:
return _as_bool(config.get("CWA_RESTRICT_SETTINGS_TO_ADMIN"))
if "RESTRICT_SETTINGS_TO_ADMIN" in config:
return _as_bool(config.get("RESTRICT_SETTINGS_TO_ADMIN"))
if "PROXY_AUTH_RESTRICT_SETTINGS_TO_ADMIN" in config:
return _as_bool(config.get("PROXY_AUTH_RESTRICT_SETTINGS_TO_ADMIN"))
if "CWA_RESTRICT_SETTINGS_TO_ADMIN" in config:
return _as_bool(config.get("CWA_RESTRICT_SETTINGS_TO_ADMIN"))
return None
def migrate_security_settings(
*,
load_security_config: Callable[[], dict[str, Any]],
load_users_config: Callable[[], dict[str, Any]],
save_users_config: Callable[[dict[str, Any]], None],
ensure_config_dir: Callable[[], None],
get_config_path: Callable[[], Any],
sync_builtin_admin_user: Callable[[str, str], None],
logger: Any,
) -> None:
"""Migrate legacy security keys and sync builtin admin credentials."""
try:
config = load_security_config()
users_config = load_users_config()
migrated_security = False
migrated_users = False
if "USE_CWA_AUTH" in config:
old_value = config.pop("USE_CWA_AUTH")
if "AUTH_METHOD" not in config:
if old_value:
config["AUTH_METHOD"] = "cwa"
logger.info("Migrated USE_CWA_AUTH=True to AUTH_METHOD='cwa'")
else:
if config.get("BUILTIN_USERNAME") and config.get("BUILTIN_PASSWORD_HASH"):
config["AUTH_METHOD"] = "builtin"
logger.info("Migrated USE_CWA_AUTH=False to AUTH_METHOD='builtin'")
else:
config["AUTH_METHOD"] = "none"
logger.info("Migrated USE_CWA_AUTH=False to AUTH_METHOD='none'")
migrated_security = True
else:
logger.info("Removed deprecated USE_CWA_AUTH setting (AUTH_METHOD already exists)")
migrated_security = True
if "RESTRICT_SETTINGS_TO_ADMIN" not in users_config:
legacy_restrict = _pick_legacy_settings_restriction(config)
if legacy_restrict is not None:
save_users_config({"RESTRICT_SETTINGS_TO_ADMIN": legacy_restrict})
migrated_users = True
logger.info(
"Migrated legacy settings-admin restriction to users.RESTRICT_SETTINGS_TO_ADMIN="
f"{legacy_restrict}"
)
for deprecated_key in _DEPRECATED_SETTINGS_RESTRICTION_KEYS:
if deprecated_key in config:
config.pop(deprecated_key, None)
migrated_security = True
logger.info(f"Removed deprecated security setting: {deprecated_key}")
try:
sync_builtin_admin_user(
config.get("BUILTIN_USERNAME", ""),
config.get("BUILTIN_PASSWORD_HASH", ""),
)
except Exception as exc:
logger.error(
"Failed to sync builtin credentials to users database during migration: "
f"{exc}"
)
if migrated_security:
ensure_config_dir()
config_path = get_config_path()
with open(config_path, "w") as f:
json.dump(config, f, indent=2)
logger.info("Security settings migration completed successfully")
elif migrated_users:
logger.info("Users settings migration completed successfully")
else:
logger.debug("No security settings migration needed")
except FileNotFoundError:
logger.debug("No existing security config file found - nothing to migrate")
except Exception as exc:
logger.error(f"Failed to migrate security settings: {exc}")
+337
View File
@@ -0,0 +1,337 @@
"""Notifications settings tab registration."""
from __future__ import annotations
import re
from typing import Any
from urllib.parse import urlsplit
from shelfmark.core.notifications import NotificationEvent, send_test_notification
from shelfmark.core.settings_registry import (
ActionButton,
HeadingField,
TableField,
load_config_file,
register_on_save,
register_settings,
)
_URL_SCHEME_RE = re.compile(r"^[a-zA-Z][a-zA-Z0-9+.-]*$")
_ROUTE_EVENT_ALL = "all"
_ADMIN_EVENT_OPTIONS = [
{"value": NotificationEvent.REQUEST_CREATED.value, "label": "New request submitted"},
{"value": NotificationEvent.REQUEST_FULFILLED.value, "label": "Request approved"},
{"value": NotificationEvent.REQUEST_REJECTED.value, "label": "Request rejected"},
{"value": NotificationEvent.DOWNLOAD_COMPLETE.value, "label": "Download complete"},
{"value": NotificationEvent.DOWNLOAD_FAILED.value, "label": "Download failed"},
]
_ROUTE_EVENT_OPTIONS = [
{"value": _ROUTE_EVENT_ALL, "label": "All"},
*_ADMIN_EVENT_OPTIONS,
]
_ROUTE_EVENT_ORDER = [option["value"] for option in _ROUTE_EVENT_OPTIONS]
_ROUTE_EVENT_INDEX = {event: index for index, event in enumerate(_ROUTE_EVENT_ORDER)}
_ALLOWED_ROUTE_EVENTS = set(_ROUTE_EVENT_ORDER)
_DEFAULT_ROUTE_ROWS = [{"event": [_ROUTE_EVENT_ALL], "url": ""}]
def _looks_like_apprise_url(url: str) -> bool:
split = urlsplit(url)
if not split.scheme:
return False
if not _URL_SCHEME_RE.match(split.scheme):
return False
return " " not in url
def _coerce_route_rows(value: Any) -> list[dict[str, Any]]:
if value is None:
return []
if isinstance(value, list):
return [row for row in value if isinstance(row, dict)]
if isinstance(value, dict):
return [value]
return []
def _coerce_route_event_values(value: Any) -> list[Any]:
if isinstance(value, list):
return value
if isinstance(value, (tuple, set)):
return list(value)
return [value]
def _normalize_route_events(value: Any) -> list[str]:
normalized: list[str] = []
seen: set[str] = set()
for raw_event in _coerce_route_event_values(value):
event = str(raw_event or "").strip().lower()
if not event or event not in _ALLOWED_ROUTE_EVENTS:
continue
if event in seen:
continue
seen.add(event)
normalized.append(event)
if _ROUTE_EVENT_ALL in seen:
return [_ROUTE_EVENT_ALL]
return sorted(normalized, key=lambda event: _ROUTE_EVENT_INDEX[event])
def _normalize_routes(value: Any) -> list[dict[str, Any]]:
normalized: list[dict[str, Any]] = []
seen: set[tuple[tuple[str, ...], str]] = set()
for row in _coerce_route_rows(value):
events = _normalize_route_events(row.get("event"))
if not events:
continue
url = str(row.get("url") or "").strip()
key = (tuple(events), url)
if key in seen:
continue
seen.add(key)
normalized.append({"event": events, "url": url})
return normalized
def _count_invalid_route_events(value: Any) -> int:
invalid = 0
for row in _coerce_route_rows(value):
raw_events = _coerce_route_event_values(row.get("event"))
if not raw_events:
invalid += 1
continue
for raw_event in raw_events:
event = str(raw_event or "").strip().lower()
if not event or event not in _ALLOWED_ROUTE_EVENTS:
invalid += 1
return invalid
def _count_invalid_route_urls(routes: list[dict[str, Any]]) -> int:
return sum(1 for row in routes if row["url"] and not _looks_like_apprise_url(row["url"]))
def _ensure_default_route_row(routes: list[dict[str, Any]]) -> list[dict[str, Any]]:
return routes if routes else [dict(row) for row in _DEFAULT_ROUTE_ROWS]
def _extract_unique_route_urls(routes: list[dict[str, Any]]) -> list[str]:
urls: list[str] = []
seen: set[str] = set()
for row in routes:
url = row.get("url", "")
if not url:
continue
if url in seen:
continue
seen.add(url)
urls.append(url)
return urls
def build_notification_test_result(routes_input: Any, *, scope_label: str) -> dict[str, Any]:
invalid_event_count = _count_invalid_route_events(routes_input)
if invalid_event_count:
return {
"success": False,
"message": (
f"Found {invalid_event_count} invalid {scope_label} notification route event value(s). "
"Fix route events before running a test."
),
}
normalized_routes = _normalize_routes(routes_input)
invalid_url_count = _count_invalid_route_urls(normalized_routes)
if invalid_url_count:
return {
"success": False,
"message": (
f"Found {invalid_url_count} invalid {scope_label} notification URL(s). "
"Fix route URLs before running a test."
),
}
urls = _extract_unique_route_urls(normalized_routes)
if not urls:
return {
"success": False,
"message": f"Add at least one {scope_label} notification URL route first.",
}
return send_test_notification(urls)
def normalize_notification_routes(value: Any) -> list[dict[str, Any]]:
"""Normalize route table rows for notification preferences."""
return _normalize_routes(value)
def is_valid_notification_url(url: str) -> bool:
"""Shared URL validation for notifications preferences."""
return _looks_like_apprise_url(url)
def _on_save_notifications(values: dict[str, Any]) -> dict[str, Any]:
existing = load_config_file("notifications")
effective: dict[str, Any] = dict(existing)
effective.update(values)
admin_routes_input = effective.get("ADMIN_NOTIFICATION_ROUTES", [])
invalid_admin_event_count = _count_invalid_route_events(admin_routes_input)
if invalid_admin_event_count:
return {
"error": True,
"message": (
f"Found {invalid_admin_event_count} invalid global notification route event value(s)."
),
"values": values,
}
normalized_admin_routes = _normalize_routes(admin_routes_input)
invalid_admin_url_count = _count_invalid_route_urls(normalized_admin_routes)
if invalid_admin_url_count:
return {
"error": True,
"message": (
f"Found {invalid_admin_url_count} invalid global notification URL(s). "
"Use URL values with a valid scheme, e.g. discord://... or ntfys://..."
),
"values": values,
}
user_routes_input = effective.get("USER_NOTIFICATION_ROUTES", [])
invalid_user_event_count = _count_invalid_route_events(user_routes_input)
if invalid_user_event_count:
return {
"error": True,
"message": (
f"Found {invalid_user_event_count} invalid personal notification route event value(s)."
),
"values": values,
}
normalized_user_routes = _normalize_routes(user_routes_input)
invalid_user_url_count = _count_invalid_route_urls(normalized_user_routes)
if invalid_user_url_count:
return {
"error": True,
"message": (
f"Found {invalid_user_url_count} invalid personal notification URL(s). "
"Use URL values with a valid scheme, e.g. discord://... or ntfys://..."
),
"values": values,
}
admin_routes_touched = "ADMIN_NOTIFICATION_ROUTES" in values
if admin_routes_touched:
values["ADMIN_NOTIFICATION_ROUTES"] = _ensure_default_route_row(normalized_admin_routes)
user_routes_touched = "USER_NOTIFICATION_ROUTES" in values
if user_routes_touched:
values["USER_NOTIFICATION_ROUTES"] = _ensure_default_route_row(normalized_user_routes)
return {"error": False, "values": values}
def _test_admin_notification_action(current_values: dict[str, Any]) -> dict[str, Any]:
persisted = load_config_file("notifications")
effective: dict[str, Any] = dict(persisted)
if isinstance(current_values, dict):
effective.update(current_values)
routes_input = effective.get("ADMIN_NOTIFICATION_ROUTES", [])
return build_notification_test_result(routes_input, scope_label="global")
register_on_save("notifications", _on_save_notifications)
@register_settings("notifications", "Notifications", icon="bell", order=7)
def notifications_settings():
"""Global notifications settings."""
return [
HeadingField(
key="notifications_heading",
title="Global Notifications",
description=(
"Global notifications send selected events for all users to configured routes. "
"Users can manage personal notifications in User Preferences."
),
),
TableField(
key="ADMIN_NOTIFICATION_ROUTES",
label="",
description=(
"Create one route per URL. Start with All, then add event-specific routes "
"for targeted delivery. Need format examples? "
"[View Apprise URL formats](https://appriseit.com/services/)."
),
columns=[
{
"key": "event",
"label": "Event",
"type": "multiselect",
"options": _ROUTE_EVENT_OPTIONS,
"defaultValue": [_ROUTE_EVENT_ALL],
"placeholder": "Select events...",
},
{
"key": "url",
"label": "Notification URL",
"type": "text",
"placeholder": "e.g. ntfys://ntfy.sh/shelfmark",
},
],
default=[dict(row) for row in _DEFAULT_ROUTE_ROWS],
add_label="Add Route",
empty_message="No routes configured.",
),
ActionButton(
key="test_admin_notification",
label="Test Notification",
description="Send a test notification to all configured global route URLs.",
style="primary",
callback=_test_admin_notification_action,
),
TableField(
key="USER_NOTIFICATION_ROUTES",
label="",
description=(
"Create one route per URL. Start with All, then add event-specific routes "
"for targeted delivery. Need format examples? "
"[View Apprise URL formats](https://appriseit.com/services/)."
),
columns=[
{
"key": "event",
"label": "Event",
"type": "multiselect",
"options": _ROUTE_EVENT_OPTIONS,
"defaultValue": [_ROUTE_EVENT_ALL],
"placeholder": "Select events...",
},
{
"key": "url",
"label": "Notification URL",
"type": "text",
"placeholder": "e.g. ntfys://ntfy.sh/username-topic",
},
],
default=[dict(row) for row in _DEFAULT_ROUTE_ROWS],
add_label="Add Route",
empty_message="No routes configured.",
user_overridable=True,
hidden_in_ui=True,
),
]
+165 -210
View File
@@ -1,9 +1,12 @@
"""Authentication settings registration."""
from typing import Any, Dict
from werkzeug.security import generate_password_hash
from typing import Any, Dict, Callable
from shelfmark.config.migrations import migrate_security_settings
from shelfmark.config.security_handlers import (
on_save_security,
test_oidc_connection,
)
from shelfmark.core.logger import setup_logger
from shelfmark.core.settings_registry import (
register_settings,
@@ -14,151 +17,57 @@ from shelfmark.core.settings_registry import (
PasswordField,
CheckboxField,
ActionButton,
TagListField,
)
from shelfmark.core.user_db import sync_builtin_admin_user
logger = setup_logger(__name__)
def _auth_condition(auth_method: str) -> dict[str, str]:
return {"field": "AUTH_METHOD", "value": auth_method}
def _ui_field(factory: Callable[..., Any], **kwargs: Any) -> Any:
return factory(env_supported=False, **kwargs)
def _auth_ui_field(factory: Callable[..., Any], auth_method: str, **kwargs: Any) -> Any:
return _ui_field(factory, show_when=_auth_condition(auth_method), **kwargs)
def _migrate_security_settings() -> None:
import json
from shelfmark.core.settings_registry import _get_config_file_path, _ensure_config_dir
from shelfmark.core.settings_registry import (
_get_config_file_path,
_ensure_config_dir,
save_config_file,
)
try:
config = load_config_file("security")
migrated = False
# Migrate USE_CWA_AUTH to AUTH_METHOD
if "USE_CWA_AUTH" in config:
old_value = config.pop("USE_CWA_AUTH")
# Only set AUTH_METHOD if it doesn't already exist
if "AUTH_METHOD" not in config:
if old_value:
config["AUTH_METHOD"] = "cwa"
logger.info("Migrated USE_CWA_AUTH=True to AUTH_METHOD='cwa'")
else:
# If USE_CWA_AUTH was False, determine auth method from credentials
if config.get("BUILTIN_USERNAME") and config.get("BUILTIN_PASSWORD_HASH"):
config["AUTH_METHOD"] = "builtin"
logger.info("Migrated USE_CWA_AUTH=False to AUTH_METHOD='builtin'")
else:
config["AUTH_METHOD"] = "none"
logger.info("Migrated USE_CWA_AUTH=False to AUTH_METHOD='none'")
migrated = True
else:
logger.info("Removed deprecated USE_CWA_AUTH setting (AUTH_METHOD already exists)")
migrated = True
# Migrate RESTRICT_SETTINGS_TO_ADMIN to CWA_RESTRICT_SETTINGS_TO_ADMIN
if "RESTRICT_SETTINGS_TO_ADMIN" in config:
old_value = config.pop("RESTRICT_SETTINGS_TO_ADMIN")
# Only migrate if new key doesn't exist
if "CWA_RESTRICT_SETTINGS_TO_ADMIN" not in config:
config["CWA_RESTRICT_SETTINGS_TO_ADMIN"] = old_value
logger.info(f"Migrated RESTRICT_SETTINGS_TO_ADMIN={old_value} to CWA_RESTRICT_SETTINGS_TO_ADMIN={old_value}")
migrated = True
else:
logger.info("Removed deprecated RESTRICT_SETTINGS_TO_ADMIN setting (CWA_RESTRICT_SETTINGS_TO_ADMIN already exists)")
migrated = True
# Save config if any migrations occurred
if migrated:
_ensure_config_dir("security")
config_path = _get_config_file_path("security")
with open(config_path, 'w') as f:
json.dump(config, f, indent=2)
logger.info("Security settings migration completed successfully")
else:
logger.debug("No security settings migration needed")
except FileNotFoundError:
logger.debug("No existing security config file found - nothing to migrate")
except Exception as e:
logger.error(f"Failed to migrate security settings: {e}")
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),
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,
logger=logger,
)
def _clear_builtin_credentials() -> Dict[str, Any]:
"""Clear built-in credentials to allow public access."""
import json
from shelfmark.core.settings_registry import _get_config_file_path, _ensure_config_dir
try:
config = load_config_file("security")
config.pop("BUILTIN_USERNAME", None)
config.pop("BUILTIN_PASSWORD_HASH", None)
_ensure_config_dir("security")
config_path = _get_config_file_path("security")
with open(config_path, 'w') as f:
json.dump(config, f, indent=2)
logger.info("Cleared credentials")
return {"success": True, "message": "Credentials cleared. The app is now publicly accessible."}
except Exception as e:
logger.error(f"Failed to clear credentials: {e}")
return {"success": False, "message": f"Failed to clear credentials: {str(e)}"}
def _on_save_security(values: Dict[str, Any]) -> Dict[str, Any]:
"""
Custom save handler for security settings.
return on_save_security(values)
Handles password validation and hashing:
- If new password is provided, validate confirmation and hash it
- If password fields are empty, preserve existing hash
- Never store raw passwords
- Ensure username is present if password is set
Returns:
Dict with processed values to save and any validation errors.
"""
password = values.get("BUILTIN_PASSWORD", "")
password_confirm = values.get("BUILTIN_PASSWORD_CONFIRM", "")
# Remove raw password fields - they should never be persisted
values.pop("BUILTIN_PASSWORD", None)
values.pop("BUILTIN_PASSWORD_CONFIRM", None)
# If password is provided, validate and hash it
if password:
if not values.get("BUILTIN_USERNAME"):
return {
"error": True,
"message": "Username cannot be empty",
"values": values
}
if password != password_confirm:
return {
"error": True,
"message": "Passwords do not match",
"values": values
}
if len(password) < 4:
return {
"error": True,
"message": "Password must be at least 4 characters",
"values": values
}
# Hash the password
values["BUILTIN_PASSWORD_HASH"] = generate_password_hash(password)
logger.info("Password hash updated")
# If no password provided but username is being set, preserve existing hash
elif "BUILTIN_USERNAME" in values:
existing = load_config_file("security")
if "BUILTIN_PASSWORD_HASH" in existing:
values["BUILTIN_PASSWORD_HASH"] = existing["BUILTIN_PASSWORD_HASH"]
return {"error": False, "values": values}
def _test_oidc_connection() -> Dict[str, Any]:
return test_oidc_connection(
load_security_config=lambda: load_config_file("security"),
logger=logger,
)
@register_settings("security", "Security", icon="shield", order=5)
def security_settings():
def security_settings():
"""Security and authentication settings."""
from shelfmark.config.env import CWA_DB_PATH
@@ -166,8 +75,9 @@ def security_settings():
auth_method_options = [
{"label": "No Authentication", "value": "none"},
{"label": "Username/Password", "value": "builtin"},
{"label": "Local", "value": "builtin"},
{"label": "Proxy Authentication", "value": "proxy"},
{"label": "OIDC (OpenID Connect)", "value": "oidc"},
]
if cwa_db_available:
auth_method_options.append({"label": "Calibre-Web Database", "value": "cwa"})
@@ -185,106 +95,151 @@ def security_settings():
default="none",
env_supported=False,
),
TextField(
key="BUILTIN_USERNAME",
label="Username",
description="Set a username and password to require login. Leave both empty for public access.",
placeholder="Enter username",
env_supported=False,
show_when={"field": "AUTH_METHOD", "value": "builtin"},
),
PasswordField(
key="BUILTIN_PASSWORD",
label="Set Password",
description="Fill in to set or change the password.",
placeholder="Enter new password",
env_supported=False,
show_when={"field": "AUTH_METHOD", "value": "builtin"},
),
PasswordField(
key="BUILTIN_PASSWORD_CONFIRM",
label="Confirm Password",
placeholder="Confirm new password",
env_supported=False,
show_when={"field": "AUTH_METHOD", "value": "builtin"},
),
ActionButton(
key="clear_credentials",
label="Clear Credentials",
description="Remove login requirement and make the app publicly accessible.",
style="danger",
callback=_clear_builtin_credentials,
show_when={"field": "AUTH_METHOD", "value": "builtin"},
key="open_users_tab",
label="Go to Users",
description="Configure local users and admin access in the Users tab.",
style="primary",
show_when=_auth_condition("builtin"),
),
TextField(
_auth_ui_field(
TextField,
"proxy",
key="PROXY_AUTH_USER_HEADER",
label="Proxy Auth User Header",
description=(
"The HTTP header your proxy uses to pass the authenticated username."
),
description="The HTTP header your proxy uses to pass the authenticated username.",
placeholder="e.g. X-Auth-User",
default="X-Auth-User",
env_supported=False,
show_when={"field": "AUTH_METHOD", "value": "proxy"},
),
TextField(
_auth_ui_field(
TextField,
"proxy",
key="PROXY_AUTH_LOGOUT_URL",
label="Proxy Auth Logout URL",
description=(
"The URL to redirect users to for logging out."
" Leave empty to disable logout functionality."
),
description="The URL to redirect users to for logging out. Leave empty to disable logout functionality.",
placeholder="https://myauth.example.com/logout",
default="",
env_supported=False,
show_when={"field": "AUTH_METHOD", "value": "proxy"},
),
CheckboxField(
key="PROXY_AUTH_RESTRICT_SETTINGS_TO_ADMIN",
label="Restrict Settings to Admins authenticated via Proxy",
description=(
"Only users in the admin group can access settings."
),
default=False,
env_supported=False,
show_when={"field": "AUTH_METHOD", "value": "proxy"},
),
TextField(
_auth_ui_field(
TextField,
"proxy",
key="PROXY_AUTH_ADMIN_GROUP_HEADER",
label="Proxy Auth Admin Group Header",
description=(
"The HTTP header your proxy uses to pass the user's groups/roles."
),
description="Optional: header your proxy uses to pass user groups/roles.",
placeholder="e.g. X-Auth-Groups",
default="X-Auth-Groups",
env_supported=False,
show_when={"field": "PROXY_AUTH_RESTRICT_SETTINGS_TO_ADMIN", "value": True},
),
TextField(
_auth_ui_field(
TextField,
"proxy",
key="PROXY_AUTH_ADMIN_GROUP_NAME",
label="Proxy Auth Admin Group Name",
description=(
"The name of the group/role that should have admin access."
),
label="Proxy Auth Admin Group",
description="Optional: users in this group are treated as admins. Leave blank to skip group-based admin detection.",
placeholder="e.g. admins",
default="admins",
env_supported=False,
show_when={"field": "PROXY_AUTH_RESTRICT_SETTINGS_TO_ADMIN", "value": True},
),
CheckboxField(
key="CWA_RESTRICT_SETTINGS_TO_ADMIN",
label="Restrict Settings to Admins authenticated via Calibre-Web",
description=(
"Only users with admin role in Calibre-Web can access settings."
),
default=False,
env_supported=False,
show_when={"field": "AUTH_METHOD", "value": "cwa"},
default="",
),
]
oidc_specs = [
(
TextField,
{
"key": "OIDC_DISCOVERY_URL",
"label": "Discovery URL",
"description": "OpenID Connect discovery endpoint URL. Usually ends with /.well-known/openid-configuration.",
"placeholder": "https://auth.example.com/.well-known/openid-configuration",
"required": True,
},
),
(
TextField,
{
"key": "OIDC_CLIENT_ID",
"label": "Client ID",
"description": "OAuth2 client ID from your identity provider.",
"placeholder": "shelfmark",
"required": True,
},
),
(
PasswordField,
{
"key": "OIDC_CLIENT_SECRET",
"label": "Client Secret",
"description": "OAuth2 client secret from your identity provider.",
"required": True,
},
),
(
TagListField,
{
"key": "OIDC_SCOPES",
"label": "Scopes",
"description": "OAuth2 scopes to request from the identity provider. Managed automatically: includes essential scopes and the group claim when using admin group authorization.",
"default": ["openid", "email", "profile"],
},
),
(
TextField,
{
"key": "OIDC_GROUP_CLAIM",
"label": "Group Claim Name",
"description": "The name of the claim in the ID token that contains user groups.",
"placeholder": "groups",
"default": "groups",
},
),
(
TextField,
{
"key": "OIDC_ADMIN_GROUP",
"label": "Admin Group Name",
"description": "Users in this group will be given admin access (if enabled below). Leave empty to use database roles only.",
"placeholder": "shelfmark-admins",
"default": "",
},
),
(
CheckboxField,
{
"key": "OIDC_USE_ADMIN_GROUP",
"label": "Use Admin Group for Authorization",
"description": "When enabled, users in the Admin Group are granted admin access. When disabled, admin access is determined solely by database roles.",
"default": True,
},
),
(
CheckboxField,
{
"key": "OIDC_AUTO_PROVISION",
"label": "Auto-Provision Users",
"description": "Automatically create a user account on first OIDC login. When disabled, users must be pre-created by an admin.",
"default": True,
},
),
(
TextField,
{
"key": "OIDC_BUTTON_LABEL",
"label": "Login Button Label",
"description": "Custom label for the OIDC sign-in button on the login page.",
"placeholder": "Sign in with OIDC",
"default": "",
},
),
]
fields.extend(_auth_ui_field(factory, "oidc", **spec) for factory, spec in oidc_specs)
fields.append(
ActionButton(
key="test_oidc",
label="Test Connection",
description="Fetch the OIDC discovery document and validate configuration.",
style="primary",
callback=_test_oidc_connection,
show_when=_auth_condition("oidc"),
)
)
return fields
# Register the on_save handler for this tab
register_on_save("security", _on_save_security)
+72
View File
@@ -0,0 +1,72 @@
"""Operational handlers for security settings (save/actions)."""
import os
from typing import Any, Callable
from shelfmark.core.utils import normalize_http_url
from shelfmark.core.user_db import UserDB
_OIDC_LOCKOUT_MESSAGE = "Create a local admin account first (Users tab) before enabling OIDC. This ensures you can still log in with a password if SSO is unavailable."
def _has_local_password_admin() -> bool:
root = os.environ.get("CONFIG_DIR", "/config")
user_db = UserDB(os.path.join(root, "users.db"))
user_db.initialize()
return any(user.get("password_hash") and user.get("role") == "admin" for user in user_db.list_users())
def on_save_security(
values: dict[str, Any],
) -> dict[str, Any]:
"""Validate security values before persistence."""
normalized_values = values.copy()
discovery_url = normalized_values.get("OIDC_DISCOVERY_URL")
if discovery_url is not None:
normalized_values["OIDC_DISCOVERY_URL"] = normalize_http_url(
str(discovery_url),
default_scheme="https",
)
proxy_logout_url = normalized_values.get("PROXY_AUTH_LOGOUT_URL")
if proxy_logout_url is not None:
normalized_values["PROXY_AUTH_LOGOUT_URL"] = normalize_http_url(
str(proxy_logout_url),
default_scheme="https",
strip_trailing_slash=False,
)
if normalized_values.get("AUTH_METHOD") == "oidc" and not _has_local_password_admin():
return {"error": True, "message": _OIDC_LOCKOUT_MESSAGE, "values": normalized_values}
return {"error": False, "values": normalized_values}
def test_oidc_connection(
*,
load_security_config: Callable[[], dict[str, Any]],
logger: Any,
) -> dict[str, Any]:
"""Fetch and validate the configured OIDC discovery document."""
import requests
try:
discovery_url = load_security_config().get("OIDC_DISCOVERY_URL", "")
if not discovery_url:
return {"success": False, "message": "Discovery URL is not configured."}
response = requests.get(discovery_url, timeout=10)
response.raise_for_status()
document = response.json()
required_fields = ["issuer", "authorization_endpoint", "token_endpoint"]
missing_fields = [field for field in required_fields if field not in document]
if missing_fields:
return {"success": False, "message": f"Discovery document missing fields: {', '.join(missing_fields)}"}
return {"success": True, "message": f"Connected to {document['issuer']}"}
except Exception as exc:
logger.error(f"OIDC connection test failed: {exc}")
return {"success": False, "message": f"Connection failed: {str(exc)}"}
+324 -15
View File
@@ -69,6 +69,7 @@ from shelfmark.config.booklore_settings import (
get_booklore_path_options,
test_booklore_connection,
)
from shelfmark.config.email_settings import test_email_connection
from shelfmark.core.logger import setup_logger
logger = setup_logger(__name__)
@@ -125,6 +126,7 @@ from shelfmark.core.settings_registry import (
CheckboxField,
SelectField,
MultiSelectField,
TagListField,
OrderableListField,
TableField,
HeadingField,
@@ -178,6 +180,7 @@ _FORMAT_OPTIONS = [
_AUDIOBOOK_FORMAT_OPTIONS = [
{"value": "m4b", "label": "M4B"},
{"value": "mp3", "label": "MP3"},
{"value": "m4a", "label": "M4A"},
{"value": "zip", "label": "ZIP"},
{"value": "rar", "label": "RAR"},
]
@@ -224,16 +227,30 @@ _LANGUAGE_OPTIONS = [{"value": lang["code"], "label": lang["language"]} for lang
def _get_aa_base_url_options():
"""Build AA URL options dynamically, including additional mirrors from config."""
from shelfmark.core.mirrors import DEFAULT_AA_MIRRORS, get_aa_mirrors
from shelfmark.core.config import config
from shelfmark.core.utils import normalize_http_url
options = [{"value": "auto", "label": "Auto (Recommended)"}]
# Get all mirrors (defaults + custom)
all_mirrors = get_aa_mirrors()
# 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"),
default_scheme="https",
allow_special=("auto",),
)
if configured_url and configured_url != "auto" and configured_url not in all_mirrors:
all_mirrors = [configured_url] + all_mirrors
for url in all_mirrors:
domain = url.replace("https://", "").replace("http://", "")
is_custom = url not in DEFAULT_AA_MIRRORS
label = f"{domain} (custom)" if is_custom else domain
if configured_url and url == configured_url and is_custom:
label = f"{domain} (configured)"
options.append({"value": url, "label": label})
return options
@@ -608,6 +625,108 @@ def _on_save_downloads(values: dict[str, Any]) -> dict[str, Any]:
"values": values,
}
# Email output (SMTP) validation.
if books_output_mode == "email":
from email.utils import parseaddr
def _is_plain_email_address(addr: str) -> bool:
parsed = parseaddr(addr or "")[1]
return bool(parsed) and "@" in parsed and parsed == addr
# Preferred model: single recipient for global default and per-user override.
raw_recipient = str(effective.get("EMAIL_RECIPIENT", "") or "").strip()
# Optional global fallback: validate only when a default recipient is provided.
if raw_recipient and not _is_plain_email_address(raw_recipient):
return {
"error": True,
"message": "Email recipient must be a valid plain email address.",
"values": values,
}
smtp_host = str(effective.get("EMAIL_SMTP_HOST", "") or "").strip()
if not smtp_host:
return {"error": True, "message": "SMTP host is required", "values": values}
security = str(effective.get("EMAIL_SMTP_SECURITY", "starttls") or "").strip().lower()
if security not in {"none", "starttls", "ssl"}:
return {
"error": True,
"message": "SMTP security must be one of: none, starttls, ssl",
"values": values,
}
try:
port = int(effective.get("EMAIL_SMTP_PORT", 587))
except (TypeError, ValueError):
return {"error": True, "message": "SMTP port must be a number", "values": values}
if port < 1 or port > 65535:
return {"error": True, "message": "SMTP port must be between 1 and 65535", "values": values}
try:
timeout_seconds = int(effective.get("EMAIL_SMTP_TIMEOUT_SECONDS", 60))
except (TypeError, ValueError):
return {"error": True, "message": "SMTP timeout (seconds) must be a number", "values": values}
if timeout_seconds < 1:
return {"error": True, "message": "SMTP timeout (seconds) must be >= 1", "values": values}
username = str(effective.get("EMAIL_SMTP_USERNAME", "") or "").strip()
password = effective.get("EMAIL_SMTP_PASSWORD", "") or ""
if username and not password:
return {"error": True, "message": "SMTP password is required when username is set", "values": values}
try:
attachment_limit_mb = int(effective.get("EMAIL_ATTACHMENT_SIZE_LIMIT_MB", 25))
except (TypeError, ValueError):
return {
"error": True,
"message": "Attachment size limit (MB) must be a number",
"values": values,
}
if attachment_limit_mb < 1 or attachment_limit_mb > 600:
return {
"error": True,
"message": "Attachment size limit (MB) must be between 1 and 600",
"values": values,
}
from_addr = str(effective.get("EMAIL_FROM", "") or "").strip()
if not from_addr:
# If From is empty, default to the SMTP username when it looks like an email address.
username_email = parseaddr(username)[1]
if username_email and "@" in username_email:
from_addr = f"Shelfmark <{username_email}>"
values["EMAIL_FROM"] = from_addr
else:
return {
"error": True,
"message": "From address is required (or set SMTP username to an email address).",
"values": values,
}
else:
from_email = parseaddr(from_addr)[1]
if not from_email or "@" not in from_email:
return {
"error": True,
"message": "From address must be a valid email address",
"values": values,
}
# Persist any normalization/coercion for fields that may have been edited this save.
if "EMAIL_RECIPIENT" in values:
values["EMAIL_RECIPIENT"] = raw_recipient
if "EMAIL_SMTP_SECURITY" in values:
values["EMAIL_SMTP_SECURITY"] = security
if "EMAIL_SMTP_PORT" in values:
values["EMAIL_SMTP_PORT"] = port
if "EMAIL_SMTP_TIMEOUT_SECONDS" in values:
values["EMAIL_SMTP_TIMEOUT_SECONDS"] = timeout_seconds
if "EMAIL_ATTACHMENT_SIZE_LIMIT_MB" in values:
values["EMAIL_ATTACHMENT_SIZE_LIMIT_MB"] = attachment_limit_mb
return {"error": False, "values": values}
@@ -632,6 +751,11 @@ def download_settings():
"label": "Folder",
"description": "Save files to the destination folder",
},
{
"value": "email",
"label": "Email (SMTP)",
"description": "Send files as an email attachment",
},
{
"value": "booklore",
"label": "Booklore (API)",
@@ -639,14 +763,16 @@ def download_settings():
},
],
default="folder",
user_overridable=True,
),
TextField(
key="DESTINATION",
label="Destination",
description="Directory where downloaded files are saved.",
description="Directory where downloaded files are saved. Use {User} for per-user folders (e.g. /books/{User}).",
default="/books",
required=True,
env_var="INGEST_DIR", # Legacy env var name for backwards compatibility
user_overridable=True,
show_when={
"field": "BOOKS_OUTPUT_MODE",
"value": "folder",
@@ -683,7 +809,7 @@ def download_settings():
TextField(
key="TEMPLATE_RENAME",
label="Naming Template",
description="Variables: {Author}, {Title}, {Year}. Universal adds: {Series}, {SeriesPosition}, {Subtitle}. Rename templates are filename-only (no '/' or '\\'); use Organize for folders.",
description="Variables: {Author}, {Title}, {Year}, {User}. Universal adds: {Series}, {SeriesPosition}, {Subtitle}. Use arbitrary prefix/suffix: {Vol. SeriesPosition - } outputs 'Vol. 2 - ' when set, nothing when empty. Rename templates are filename-only (no '/' or '\\'); use Organize for folders.",
default="{Author} - {Title} ({Year})",
placeholder="{Author} - {Title} ({Year})",
show_when=[
@@ -695,7 +821,7 @@ def download_settings():
TextField(
key="TEMPLATE_ORGANIZE",
label="Path Template",
description="Use / to create folders. Variables: {Author}, {Title}, {Year}. Universal adds: {Series}, {SeriesPosition}, {Subtitle}",
description="Use / to create folders. Variables: {Author}, {Title}, {Year}, {User}. Universal adds: {Series}, {SeriesPosition}, {Subtitle}. Use arbitrary prefix/suffix: {Vol. SeriesPosition - } outputs 'Vol. 2 - ' when set, nothing when empty.",
default="{Author}/{Title} ({Year})",
placeholder="{Author}/{Series/}{Title} ({Year})",
show_when=[
@@ -742,13 +868,36 @@ def download_settings():
required=True,
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "booklore"},
),
SelectField(
key="BOOKLORE_DESTINATION",
label="Upload Destination",
description="Choose whether uploads go directly to a specific library path or to Bookdrop for review.",
options=[
{
"value": "library",
"label": "Specific Library",
"description": "Upload directly into the selected library path.",
},
{
"value": "bookdrop",
"label": "Bookdrop",
"description": "Upload into Bookdrop and review metadata before importing to a library.",
},
],
default="library",
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "booklore"},
),
SelectField(
key="BOOKLORE_LIBRARY_ID",
label="Library",
description="Booklore library to upload into.",
options=get_booklore_library_options,
required=True,
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "booklore"},
user_overridable=True,
show_when=[
{"field": "BOOKS_OUTPUT_MODE", "value": "booklore"},
{"field": "BOOKLORE_DESTINATION", "value": "library"},
],
),
SelectField(
key="BOOKLORE_PATH_ID",
@@ -757,7 +906,11 @@ def download_settings():
options=get_booklore_path_options,
required=True,
filter_by_field="BOOKLORE_LIBRARY_ID",
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "booklore"},
user_overridable=True,
show_when=[
{"field": "BOOKS_OUTPUT_MODE", "value": "booklore"},
{"field": "BOOKLORE_DESTINATION", "value": "library"},
],
),
ActionButton(
key="test_booklore",
@@ -767,6 +920,110 @@ def download_settings():
callback=test_booklore_connection,
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "booklore"},
),
HeadingField(
key="email_heading",
title="Email",
description="Send books as email attachments via SMTP. Audiobooks always use folder mode.",
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
TextField(
key="EMAIL_RECIPIENT",
label="Default Email Recipient",
description="Optional fallback email address when no per-user email recipient override is configured.",
placeholder="reader@example.com",
user_overridable=True,
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
NumberField(
key="EMAIL_ATTACHMENT_SIZE_LIMIT_MB",
label="Attachment Size Limit (MB)",
description="Maximum total attachment size per email. Email encoding adds overhead; keep this below your provider's limit.",
default=25,
min_value=1,
max_value=600,
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
TextField(
key="EMAIL_SMTP_HOST",
label="SMTP Host",
description="SMTP server hostname or IP (e.g., smtp.gmail.com).",
placeholder="smtp.example.com",
required=True,
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
NumberField(
key="EMAIL_SMTP_PORT",
label="SMTP Port",
description="SMTP server port (587 is typical for STARTTLS, 465 for SSL).",
default=587,
min_value=1,
max_value=65535,
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
SelectField(
key="EMAIL_SMTP_SECURITY",
label="SMTP Security",
description="Transport security mode for SMTP.",
options=[
{"value": "none", "label": "None", "description": "No TLS (not recommended)."},
{"value": "starttls", "label": "STARTTLS", "description": "Upgrade to TLS after connecting (recommended)."},
{"value": "ssl", "label": "SSL/TLS", "description": "Connect using TLS (SMTPS)."},
],
default="starttls",
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
TextField(
key="EMAIL_SMTP_USERNAME",
label="Username",
description="SMTP username (leave empty for no authentication).",
placeholder="user@example.com",
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
PasswordField(
key="EMAIL_SMTP_PASSWORD",
label="Password",
description="SMTP password (required if Username is set).",
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
TextField(
key="EMAIL_FROM",
label="From Address",
description="From address used for the email. You can include a display name (e.g., Shelfmark <mail@example.com>). Leave blank to default to the SMTP username (when it is an email address).",
placeholder="Shelfmark <mail@example.com>",
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
TextField(
key="EMAIL_SUBJECT_TEMPLATE",
label="Subject Template",
description="Email subject. Variables: {Author}, {Title}, {Year}, {Series}, {SeriesPosition}, {Subtitle}, {Format}.",
default="{Title}",
placeholder="{Title}",
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
NumberField(
key="EMAIL_SMTP_TIMEOUT_SECONDS",
label="SMTP Timeout (seconds)",
description="How long to wait for SMTP operations before failing.",
default=60,
min_value=1,
max_value=600,
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
CheckboxField(
key="EMAIL_ALLOW_UNVERIFIED_TLS",
label="Allow Unverified TLS",
description="Disable TLS certificate verification (not recommended).",
default=False,
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
ActionButton(
key="test_email",
label="Test SMTP Connection",
description="Verify your SMTP configuration (connect + optional login).",
style="primary",
callback=test_email_connection,
show_when={"field": "BOOKS_OUTPUT_MODE", "value": "email"},
),
# === AUDIOBOOKS SECTION ===
# Universal mode only
@@ -779,8 +1036,8 @@ def download_settings():
TextField(
key="DESTINATION_AUDIOBOOK",
label="Destination",
description="Leave empty to use Books destination.",
placeholder="/audiobooks",
description="Directory where downloaded audiobook files are saved. Leave empty to use the Books destination.",
user_overridable=True,
universal_only=True,
),
SelectField(
@@ -799,7 +1056,7 @@ def download_settings():
TextField(
key="TEMPLATE_AUDIOBOOK_RENAME",
label="Naming Template",
description="Variables: {Author}, {Title}, {Year}, {Series}, {SeriesPosition}, {Subtitle}, {PartNumber}. Rename templates are filename-only (no '/' or '\\'); use Organize for folders.",
description="Variables: {Author}, {Title}, {Year}, {User}, {Series}, {SeriesPosition}, {Subtitle}, {PartNumber}. Use arbitrary prefix/suffix: {Vol. SeriesPosition - } outputs 'Vol. 2 - ' when set, nothing when empty. Rename templates are filename-only (no '/' or '\\'); use Organize for folders.",
default="{Author} - {Title}",
placeholder="{Author} - {Title}{ - Part }{PartNumber}",
show_when={"field": "FILE_ORGANIZATION_AUDIOBOOK", "value": "rename"},
@@ -809,7 +1066,7 @@ def download_settings():
TextField(
key="TEMPLATE_AUDIOBOOK_ORGANIZE",
label="Path Template",
description="Use / to create folders. Variables: {Author}, {Title}, {Year}, {Series}, {SeriesPosition}, {Subtitle}, {PartNumber}",
description="Use / to create folders. Variables: {Author}, {Title}, {Year}, {User}, {Series}, {SeriesPosition}, {Subtitle}, {PartNumber}. Use arbitrary prefix/suffix: {Vol. SeriesPosition - } outputs 'Vol. 2 - ' when set, nothing when empty.",
default="{Author}/{Title}",
placeholder="{Author}/{Series/}{Title}{ - Part }{PartNumber}",
show_when={"field": "FILE_ORGANIZATION_AUDIOBOOK", "value": "organize"},
@@ -1102,30 +1359,76 @@ def cloudflare_bypass_settings():
),
]
def _on_save_mirrors(values: Dict[str, Any]) -> Dict[str, Any]:
"""Normalize mirror list settings before persisting."""
from shelfmark.core.logger import setup_logger
from shelfmark.core.mirrors import DEFAULT_AA_MIRRORS
from shelfmark.core.utils import normalize_http_url
logger = setup_logger(__name__)
raw_urls = values.get("AA_MIRROR_URLS")
if raw_urls is None:
return {"error": False, "values": values}
if isinstance(raw_urls, str):
parts = [p.strip() for p in raw_urls.split(",") if p.strip()]
elif isinstance(raw_urls, list):
parts = [str(p).strip() for p in raw_urls if str(p).strip()]
else:
parts = []
normalized: list[str] = []
for url in parts:
if url.lower() == "auto":
continue
norm = normalize_http_url(url, default_scheme="https")
if norm and norm not in normalized:
normalized.append(norm)
if not normalized:
logger.warning("AA_MIRROR_URLS saved empty/invalid; falling back to defaults")
normalized = [normalize_http_url(url, default_scheme="https") for url in DEFAULT_AA_MIRRORS]
normalized = [url for url in normalized if url]
values["AA_MIRROR_URLS"] = normalized
return {"error": False, "values": values}
# Register the on_save handler for this tab
register_on_save("mirrors", _on_save_mirrors)
@register_settings("mirrors", "Mirrors", icon="globe", order=23, group="direct_download")
def mirror_settings():
"""Configure download source mirrors."""
from shelfmark.core.mirrors import DEFAULT_ZLIB_MIRRORS, DEFAULT_WELIB_MIRRORS
from shelfmark.core.mirrors import DEFAULT_AA_MIRRORS, DEFAULT_ZLIB_MIRRORS, DEFAULT_WELIB_MIRRORS
return [
# === PRIMARY SOURCE ===
HeadingField(
key="aa_mirrors_heading",
title="Primary Source",
description="Primary mirror with auto-probe on startup. Additional mirrors used as fallback.",
title="Anna's Archive",
description="Choose a primary mirror, or use Auto to try mirrors from your list below. The mirror list controls which options appear in the dropdown and the order used in Auto mode.",
),
SelectField(
key="AA_BASE_URL",
label="Primary Mirror",
description="Select 'Auto' to probe mirrors on startup, or choose a specific mirror.",
description="Select 'Auto' to try mirrors from your list on startup and fall back on failures. Choosing a specific mirror locks Shelfmark to that mirror (no fallback).",
options=_get_aa_base_url_options,
default="auto",
),
TagListField(
key="AA_MIRROR_URLS",
label="Mirrors",
description="Editable list of AA mirrors. Used to populate the Primary Mirror dropdown and the order used when Auto is selected. Type a URL and press Enter to add. Order matters for auto-rotation",
placeholder="https://annas-archive.gl",
default=DEFAULT_AA_MIRRORS,
),
TextField(
key="AA_ADDITIONAL_URLS",
label="Additional Mirrors",
description="Comma-separated list of custom mirror URLs.",
label="Additional Mirrors (Legacy)",
description="Deprecated. Use Mirrors instead. This is kept for backwards compatibility with existing installs and environment variables.",
show_when={"field": "AA_ADDITIONAL_URLS", "notEmpty": True},
),
# === LIBGEN ===
@@ -1232,6 +1535,12 @@ def advanced_settings():
],
default="absolute",
),
CheckboxField(
key="CUSTOM_SCRIPT_JSON_PAYLOAD",
label="Custom Script JSON Payload",
description="Send a JSON payload to the script via stdin. Useful for multi-file imports (audiobooks) or richer metadata without relying on path parsing.",
default=False,
),
HeadingField(
key="remote_path_mappings_heading",
title="Remote Path Mappings",
+320
View File
@@ -0,0 +1,320 @@
"""Users settings tab registration.
This registers a 'users' tab in the settings sidebar.
The actual user management is handled by a custom frontend component
that talks to /api/admin/users endpoints.
"""
from shelfmark.core.settings_registry import (
CheckboxField,
CustomComponentField,
HeadingField,
MultiSelectField,
NumberField,
SelectField,
TableField,
register_on_save,
register_settings,
)
from shelfmark.core.request_policy import (
get_source_content_type_capabilities,
parse_policy_mode,
validate_policy_rules,
)
_REQUEST_DEFAULT_MODE_OPTIONS = [
{
"value": "download",
"label": "Download",
"description": "Everything can be downloaded directly.",
},
{
"value": "request_release",
"label": "Request Release",
"description": "Users must request a specific release.",
},
{
"value": "request_book",
"label": "Request Book",
"description": "Users request a book, admin picks the release.",
},
{
"value": "blocked",
"label": "Blocked",
"description": "No downloads or requests allowed.",
},
]
_REQUEST_MATRIX_MODE_OPTIONS = [
option for option in _REQUEST_DEFAULT_MODE_OPTIONS if option["value"] != "request_book"
]
_SELF_SETTINGS_SECTION_OPTIONS = [
{
"value": "delivery",
"label": "Delivery Preferences",
"description": "Show personal delivery output and destination settings.",
},
{
"value": "notifications",
"label": "Notifications",
"description": "Show personal notification route settings.",
},
]
_SELF_SETTINGS_SECTION_VALUES = {option["value"] for option in _SELF_SETTINGS_SECTION_OPTIONS}
_SELF_SETTINGS_SECTION_DEFAULTS = [option["value"] for option in _SELF_SETTINGS_SECTION_OPTIONS]
_USERS_HEADING_DESCRIPTION_BY_AUTH_MODE = {
"builtin": (
"Create and manage user accounts directly. Passwords are stored locally and users sign in "
"with their username and password."
),
"oidc": (
"Users sign in through your identity provider. New accounts can be created automatically on "
"first login when auto-provisioning is enabled, or you can pre-create users here and they\u2019ll "
"be linked by email on first sign-in."
),
"proxy": (
"Users are authenticated by your reverse proxy. Accounts are automatically created on first "
"sign-in. If a local user with a matching username already exists, it will be linked instead."
),
"cwa": (
"User accounts are synced from your Calibre-Web database. Users are matched by email, and new "
"accounts are created here when new CWA users are found."
),
"none": "Authentication is disabled. Anyone can access Shelfmark without signing in.",
"default": "Authentication is disabled. Anyone can access Shelfmark without signing in.",
}
def _get_request_source_options():
"""Build request-policy source options from registered release sources."""
from shelfmark.release_sources import list_available_sources
options = []
for source in list_available_sources():
options.append(
{
"value": source["name"],
"label": source["display_name"],
}
)
return options
def _get_request_policy_rule_columns():
source_capabilities = get_source_content_type_capabilities()
content_type_options = []
for source_name, supported_types in source_capabilities.items():
normalized_types = [t for t in ("ebook", "audiobook") if t in supported_types]
for content_type in normalized_types:
content_type_options.append(
{
"value": content_type,
"label": "Ebook" if content_type == "ebook" else "Audiobook",
"childOf": source_name,
}
)
return [
{
"key": "source",
"label": "Source",
"type": "select",
"options": _get_request_source_options(),
"defaultValue": "",
"placeholder": "Select source...",
},
{
"key": "content_type",
"label": "Content Type",
"type": "select",
"options": content_type_options,
"defaultValue": "",
"placeholder": "Select content type...",
"filterByField": "source",
},
{
"key": "mode",
"label": "Mode",
"type": "select",
"options": _REQUEST_MATRIX_MODE_OPTIONS,
"defaultValue": "",
"placeholder": "Select mode...",
},
]
def _on_save_users(values):
"""Validate users/request-policy settings before persistence."""
if "VISIBLE_SELF_SETTINGS_SECTIONS" in values:
raw_sections = values["VISIBLE_SELF_SETTINGS_SECTIONS"]
if raw_sections is None:
candidate_sections: list[str] = []
elif isinstance(raw_sections, str):
candidate_sections = [s.strip() for s in raw_sections.split(",") if s.strip()]
elif isinstance(raw_sections, (list, tuple, set)):
candidate_sections = [str(section).strip() for section in raw_sections if str(section).strip()]
else:
return {
"error": True,
"message": "VISIBLE_SELF_SETTINGS_SECTIONS must be a list of section identifiers",
"values": values,
}
normalized_sections: list[str] = []
for section in candidate_sections:
if section not in _SELF_SETTINGS_SECTION_VALUES:
allowed = ", ".join(sorted(_SELF_SETTINGS_SECTION_VALUES))
return {
"error": True,
"message": (
"VISIBLE_SELF_SETTINGS_SECTIONS contains an unsupported section "
f"'{section}'. Supported values: {allowed}"
),
"values": values,
}
if section not in normalized_sections:
normalized_sections.append(section)
values["VISIBLE_SELF_SETTINGS_SECTIONS"] = normalized_sections
if "REQUEST_POLICY_DEFAULT_EBOOK" in values:
if parse_policy_mode(values["REQUEST_POLICY_DEFAULT_EBOOK"]) is None:
return {
"error": True,
"message": "REQUEST_POLICY_DEFAULT_EBOOK must be a valid policy mode",
"values": values,
}
if "REQUEST_POLICY_DEFAULT_AUDIOBOOK" in values:
if parse_policy_mode(values["REQUEST_POLICY_DEFAULT_AUDIOBOOK"]) is None:
return {
"error": True,
"message": "REQUEST_POLICY_DEFAULT_AUDIOBOOK must be a valid policy mode",
"values": values,
}
if "REQUEST_POLICY_RULES" in values:
normalized_rules, errors = validate_policy_rules(values["REQUEST_POLICY_RULES"])
if errors:
return {
"error": True,
"message": "; ".join(errors),
"values": values,
}
values["REQUEST_POLICY_RULES"] = normalized_rules
return {"error": False, "values": values}
register_on_save("users", _on_save_users)
@register_settings("users", "Users & Requests", icon="users", order=6)
def users_settings():
"""User management tab - rendered as a custom component on the frontend."""
return [
HeadingField(
key="users_heading",
title="Users",
description=_USERS_HEADING_DESCRIPTION_BY_AUTH_MODE["default"],
description_by_auth_mode=_USERS_HEADING_DESCRIPTION_BY_AUTH_MODE,
),
CustomComponentField(
key="users_management",
component="users_management",
),
MultiSelectField(
key="VISIBLE_SELF_SETTINGS_SECTIONS",
label="Visible Self-Settings Sections",
description=(
"Choose which personal settings sections are shown in My Account for non-admin users."
),
options=_SELF_SETTINGS_SECTION_OPTIONS,
default=_SELF_SETTINGS_SECTION_DEFAULTS,
variant="dropdown",
env_supported=False,
),
HeadingField(
key="requests_heading",
title="Requests",
description=(
"Choose what users can download directly and what needs approval first."
),
),
CheckboxField(
key="REQUESTS_ENABLED",
label="Enable Requests",
description=(
"Turn this off to let everyone download directly without needing approval."
),
default=False,
user_overridable=True,
),
CustomComponentField(
key="request_policy_editor",
component="request_policy_grid",
label="Request Rules",
description=(
"Fine-tune access per source. Source rules can only be the same or more restrictive than the default above."
),
show_when={"field": "REQUESTS_ENABLED", "value": True},
wrap_in_field_wrapper=True,
value_fields=[
SelectField(
key="REQUEST_POLICY_DEFAULT_EBOOK",
label="Default Ebook Mode",
description=(
"Sets the baseline for all ebook sources."
),
options=_REQUEST_DEFAULT_MODE_OPTIONS,
default="download",
user_overridable=True,
),
SelectField(
key="REQUEST_POLICY_DEFAULT_AUDIOBOOK",
label="Default Audiobook Mode",
description=(
"Sets the baseline for all audiobook sources."
),
options=_REQUEST_DEFAULT_MODE_OPTIONS,
default="download",
user_overridable=True,
),
TableField(
key="REQUEST_POLICY_RULES",
label="Request Rules",
description=(
"Fine-tune access per source. Source rules can only be the same or more restrictive than the default above."
),
columns=_get_request_policy_rule_columns,
default=[],
add_label="Add Rule",
empty_message="No request policy rules configured.",
env_supported=False,
user_overridable=True,
),
],
),
NumberField(
key="MAX_PENDING_REQUESTS_PER_USER",
label="Max pending requests per user",
description="How many open requests a user can have at a time.",
default=20,
min_value=1,
max_value=1000,
user_overridable=True,
show_when={"field": "REQUESTS_ENABLED", "value": True},
),
CheckboxField(
key="REQUESTS_ALLOW_NOTES",
label="Allow notes on requests",
description="Let users add a note when they submit a request.",
default=True,
user_overridable=True,
show_when={"field": "REQUESTS_ENABLED", "value": True},
),
]
+486
View File
@@ -0,0 +1,486 @@
"""Activity API routes (snapshot, dismiss, history)."""
from __future__ import annotations
from typing import Any, Callable
from flask import Flask, jsonify, request, session
from shelfmark.core.activity_service import ActivityService
from shelfmark.core.logger import setup_logger
from shelfmark.core.user_db import UserDB
logger = setup_logger(__name__)
def _require_authenticated(resolve_auth_mode: Callable[[], str]):
auth_mode = resolve_auth_mode()
if auth_mode == "none":
return None
if "user_id" not in session:
return jsonify({"error": "Unauthorized"}), 401
return None
def _resolve_db_user_id(require_in_auth_mode: bool = True):
raw_db_user_id = session.get("db_user_id")
if raw_db_user_id is None:
if not require_in_auth_mode:
return None, None
return None, (
jsonify(
{
"error": "User identity unavailable for activity workflow",
"code": "user_identity_unavailable",
}
),
403,
)
try:
return int(raw_db_user_id), None
except (TypeError, ValueError):
return None, (
jsonify(
{
"error": "User identity unavailable for activity workflow",
"code": "user_identity_unavailable",
}
),
403,
)
def _emit_activity_event(ws_manager: Any | None, *, room: str, payload: dict[str, Any]) -> None:
if ws_manager is None:
return
try:
socketio = getattr(ws_manager, "socketio", None)
is_enabled = getattr(ws_manager, "is_enabled", None)
if socketio is None or not callable(is_enabled) or not is_enabled():
return
socketio.emit("activity_update", payload, to=room)
except Exception as exc:
logger.warning("Failed to emit activity_update event: %s", exc)
def _list_visible_requests(user_db: UserDB, *, is_admin: bool, db_user_id: int | None) -> list[dict[str, Any]]:
if is_admin:
request_rows = user_db.list_requests()
user_cache: dict[int, str] = {}
for row in request_rows:
requester_id = row["user_id"]
if requester_id not in user_cache:
requester = user_db.get_user(user_id=requester_id)
user_cache[requester_id] = requester.get("username", "") if requester else ""
row["username"] = user_cache[requester_id]
return request_rows
if db_user_id is None:
return []
return user_db.list_requests(user_id=db_user_id)
def _parse_download_item_key(item_key: str) -> str | None:
if not isinstance(item_key, str) or not item_key.startswith("download:"):
return None
task_id = item_key.split(":", 1)[1].strip()
return task_id or None
def _parse_request_item_key(item_key: str) -> int | None:
if not isinstance(item_key, str) or not item_key.startswith("request:"):
return None
raw_id = item_key.split(":", 1)[1].strip()
try:
parsed = int(raw_id)
except (TypeError, ValueError):
return None
return parsed if parsed > 0 else None
def _task_id_from_download_item_key(item_key: str) -> str | None:
task_id = _parse_download_item_key(item_key)
if task_id is None:
return None
return task_id
def _merge_terminal_snapshot_backfill(
*,
status: dict[str, dict[str, Any]],
terminal_rows: list[dict[str, Any]],
) -> None:
existing_task_ids: set[str] = set()
for bucket_key in ("queued", "resolving", "locating", "downloading", "complete", "error", "cancelled"):
bucket = status.get(bucket_key)
if not isinstance(bucket, dict):
continue
existing_task_ids.update(str(task_id) for task_id in bucket.keys())
for row in terminal_rows:
item_key = row.get("item_key")
if not isinstance(item_key, str):
continue
task_id = _task_id_from_download_item_key(item_key)
if not task_id or task_id in existing_task_ids:
continue
final_status = row.get("final_status")
if final_status not in {"complete", "error", "cancelled"}:
continue
snapshot = row.get("snapshot")
if not isinstance(snapshot, dict):
continue
raw_download = snapshot.get("download")
if not isinstance(raw_download, dict):
continue
download_payload = dict(raw_download)
if not isinstance(download_payload.get("id"), str):
download_payload["id"] = task_id
if final_status not in status or not isinstance(status.get(final_status), dict):
status[final_status] = {}
status[final_status][task_id] = download_payload
existing_task_ids.add(task_id)
def _collect_active_download_item_keys(status: dict[str, dict[str, Any]]) -> set[str]:
active_keys: set[str] = set()
for bucket_key in ("queued", "resolving", "locating", "downloading"):
bucket = status.get(bucket_key)
if not isinstance(bucket, dict):
continue
for task_id in bucket.keys():
normalized_task_id = str(task_id).strip()
if not normalized_task_id:
continue
active_keys.add(f"download:{normalized_task_id}")
return active_keys
def _extract_request_source_id(row: dict[str, Any]) -> str | None:
release_data = row.get("release_data")
if not isinstance(release_data, dict):
return None
source_id = release_data.get("source_id")
if not isinstance(source_id, str):
return None
normalized = source_id.strip()
return normalized or None
def _request_terminal_status(row: dict[str, Any]) -> str | None:
request_status = row.get("status")
if request_status == "pending":
return None
if request_status == "rejected":
return "rejected"
if request_status == "cancelled":
return "cancelled"
if request_status != "fulfilled":
return None
delivery_state = str(row.get("delivery_state") or "").strip().lower()
if delivery_state in {"error", "cancelled"}:
return delivery_state
return "complete"
def _minimal_request_snapshot(request_row: dict[str, Any], request_id: int) -> dict[str, Any]:
book_data = request_row.get("book_data")
release_data = request_row.get("release_data")
if not isinstance(book_data, dict):
book_data = {}
if not isinstance(release_data, dict):
release_data = {}
minimal_request = {
"id": request_id,
"user_id": request_row.get("user_id"),
"status": request_row.get("status"),
"request_level": request_row.get("request_level"),
"delivery_state": request_row.get("delivery_state"),
"book_data": book_data,
"release_data": release_data,
"note": request_row.get("note"),
"admin_note": request_row.get("admin_note"),
"created_at": request_row.get("created_at"),
"updated_at": request_row.get("updated_at"),
}
username = request_row.get("username")
if isinstance(username, str):
minimal_request["username"] = username
return {"kind": "request", "request": minimal_request}
def _get_existing_activity_log_id_for_item(
*,
activity_service: ActivityService,
user_db: UserDB,
item_type: str,
item_key: str,
) -> int | None:
if item_type not in {"request", "download"}:
return None
if not isinstance(item_key, str) or not item_key.strip():
return None
existing_log_id = activity_service.get_latest_activity_log_id(
item_type=item_type,
item_key=item_key,
)
if existing_log_id is not None or item_type != "request":
return existing_log_id
request_id = _parse_request_item_key(item_key)
if request_id is None:
return None
row = user_db.get_request(request_id)
if row is None:
return None
final_status = _request_terminal_status(row)
if final_status is None:
return None
source_id = _extract_request_source_id(row)
payload = activity_service.record_terminal_snapshot(
user_id=row.get("user_id"),
item_type="request",
item_key=item_key,
origin="request",
final_status=final_status,
snapshot=_minimal_request_snapshot(row, request_id),
request_id=request_id,
source_id=source_id,
)
return int(payload["id"])
def register_activity_routes(
app: Flask,
user_db: UserDB,
*,
activity_service: ActivityService,
resolve_auth_mode: Callable[[], str],
resolve_status_scope: Callable[[], tuple[bool, int | None, bool]],
queue_status: Callable[..., dict[str, dict[str, Any]]],
sync_request_delivery_states: Callable[..., list[dict[str, Any]]],
emit_request_updates: Callable[[list[dict[str, Any]]], None],
ws_manager: Any | None = None,
) -> None:
"""Register activity routes."""
@app.route("/api/activity/snapshot", methods=["GET"])
def api_activity_snapshot():
auth_gate = _require_authenticated(resolve_auth_mode)
if auth_gate is not None:
return auth_gate
is_admin, db_user_id, can_access_status = resolve_status_scope()
if not can_access_status:
return (
jsonify(
{
"error": "User identity unavailable for activity workflow",
"code": "user_identity_unavailable",
}
),
403,
)
viewer_db_user_id, _ = _resolve_db_user_id(require_in_auth_mode=False)
scoped_user_id = None if is_admin else db_user_id
status = queue_status(user_id=scoped_user_id)
updated_requests = sync_request_delivery_states(
user_db,
queue_status=status,
user_id=scoped_user_id,
)
emit_request_updates(updated_requests)
request_rows = _list_visible_requests(user_db, is_admin=is_admin, db_user_id=db_user_id)
if not is_admin and db_user_id is not None:
try:
terminal_rows = activity_service.get_undismissed_terminal_downloads(db_user_id, limit=200)
_merge_terminal_snapshot_backfill(status=status, terminal_rows=terminal_rows)
except Exception as exc:
logger.warning("Failed to merge terminal snapshot backfill rows: %s", exc)
if viewer_db_user_id is not None:
active_download_keys = _collect_active_download_item_keys(status)
if active_download_keys:
try:
activity_service.clear_dismissals_for_item_keys(
user_id=viewer_db_user_id,
item_type="download",
item_keys=active_download_keys,
)
except Exception as exc:
logger.warning("Failed to clear stale download dismissals for active tasks: %s", exc)
dismissed: list[dict[str, str]] = []
# Admins can view unscoped queue status, but dismissals remain per-viewer.
if viewer_db_user_id is not None:
dismissed = activity_service.get_dismissal_set(viewer_db_user_id)
return jsonify(
{
"status": status,
"requests": request_rows,
"dismissed": dismissed,
}
)
@app.route("/api/activity/dismiss", methods=["POST"])
def api_activity_dismiss():
auth_gate = _require_authenticated(resolve_auth_mode)
if auth_gate is not None:
return auth_gate
db_user_id, db_gate = _resolve_db_user_id()
if db_gate is not None or db_user_id is None:
return db_gate
data = request.get_json(silent=True)
if not isinstance(data, dict):
return jsonify({"error": "Invalid payload"}), 400
activity_log_id = data.get("activity_log_id")
if activity_log_id is None:
try:
activity_log_id = _get_existing_activity_log_id_for_item(
activity_service=activity_service,
user_db=user_db,
item_type=data.get("item_type"),
item_key=data.get("item_key"),
)
except Exception as exc:
logger.warning("Failed to resolve activity snapshot id for dismiss payload: %s", exc)
activity_log_id = None
try:
dismissal = activity_service.dismiss_item(
user_id=db_user_id,
item_type=data.get("item_type"),
item_key=data.get("item_key"),
activity_log_id=activity_log_id,
)
except ValueError as exc:
return jsonify({"error": str(exc)}), 400
_emit_activity_event(
ws_manager,
room=f"user_{db_user_id}",
payload={
"kind": "dismiss",
"user_id": db_user_id,
"item_type": dismissal["item_type"],
"item_key": dismissal["item_key"],
},
)
return jsonify({"status": "dismissed", "item": dismissal})
@app.route("/api/activity/dismiss-many", methods=["POST"])
def api_activity_dismiss_many():
auth_gate = _require_authenticated(resolve_auth_mode)
if auth_gate is not None:
return auth_gate
db_user_id, db_gate = _resolve_db_user_id()
if db_gate is not None or db_user_id is None:
return db_gate
data = request.get_json(silent=True)
if not isinstance(data, dict):
return jsonify({"error": "Invalid payload"}), 400
items = data.get("items")
if not isinstance(items, list):
return jsonify({"error": "items must be an array"}), 400
normalized_items: list[dict[str, Any]] = []
for item in items:
if not isinstance(item, dict):
return jsonify({"error": "items must contain objects"}), 400
activity_log_id = item.get("activity_log_id")
if activity_log_id is None:
try:
activity_log_id = _get_existing_activity_log_id_for_item(
activity_service=activity_service,
user_db=user_db,
item_type=item.get("item_type"),
item_key=item.get("item_key"),
)
except Exception as exc:
logger.warning("Failed to resolve activity snapshot id for dismiss-many item: %s", exc)
activity_log_id = None
normalized_payload = {
"item_type": item.get("item_type"),
"item_key": item.get("item_key"),
}
if activity_log_id is not None:
normalized_payload["activity_log_id"] = activity_log_id
normalized_items.append(normalized_payload)
try:
dismissed_count = activity_service.dismiss_many(user_id=db_user_id, items=normalized_items)
except ValueError as exc:
return jsonify({"error": str(exc)}), 400
_emit_activity_event(
ws_manager,
room=f"user_{db_user_id}",
payload={
"kind": "dismiss_many",
"user_id": db_user_id,
"count": dismissed_count,
},
)
return jsonify({"status": "dismissed", "count": dismissed_count})
@app.route("/api/activity/history", methods=["GET"])
def api_activity_history():
auth_gate = _require_authenticated(resolve_auth_mode)
if auth_gate is not None:
return auth_gate
db_user_id, db_gate = _resolve_db_user_id()
if db_gate is not None or db_user_id is None:
return db_gate
limit = request.args.get("limit", type=int, default=50) or 50
offset = request.args.get("offset", type=int, default=0) or 0
try:
history = activity_service.get_history(db_user_id, limit=limit, offset=offset)
except ValueError as exc:
return jsonify({"error": str(exc)}), 400
return jsonify(history)
@app.route("/api/activity/history", methods=["DELETE"])
def api_activity_history_clear():
auth_gate = _require_authenticated(resolve_auth_mode)
if auth_gate is not None:
return auth_gate
db_user_id, db_gate = _resolve_db_user_id()
if db_gate is not None or db_user_id is None:
return db_gate
deleted_count = activity_service.clear_history(db_user_id)
_emit_activity_event(
ws_manager,
room=f"user_{db_user_id}",
payload={
"kind": "history_cleared",
"user_id": db_user_id,
"count": deleted_count,
},
)
return jsonify({"status": "cleared", "deleted_count": deleted_count})
+618
View File
@@ -0,0 +1,618 @@
"""Persistence helpers for Activity dismissals and terminal snapshots."""
from __future__ import annotations
from datetime import datetime, timezone
import json
import sqlite3
from typing import Any, Iterable
VALID_ITEM_TYPES = frozenset({"download", "request"})
VALID_ORIGINS = frozenset({"direct", "request", "requested"})
VALID_FINAL_STATUSES = frozenset({"complete", "error", "cancelled", "rejected"})
def _now_timestamp() -> str:
return datetime.now(timezone.utc).isoformat(timespec="seconds")
def _normalize_item_type(item_type: Any) -> str:
if not isinstance(item_type, str):
raise ValueError("item_type must be a string")
normalized = item_type.strip().lower()
if normalized not in VALID_ITEM_TYPES:
raise ValueError("item_type must be one of: download, request")
return normalized
def _normalize_item_key(item_key: Any) -> str:
if not isinstance(item_key, str):
raise ValueError("item_key must be a string")
normalized = item_key.strip()
if not normalized:
raise ValueError("item_key must not be empty")
return normalized
def _normalize_origin(origin: Any) -> str:
if not isinstance(origin, str):
raise ValueError("origin must be a string")
normalized = origin.strip().lower()
if normalized not in VALID_ORIGINS:
raise ValueError("origin must be one of: direct, request, requested")
return normalized
def _normalize_final_status(final_status: Any) -> str:
if not isinstance(final_status, str):
raise ValueError("final_status must be a string")
normalized = final_status.strip().lower()
if normalized not in VALID_FINAL_STATUSES:
raise ValueError("final_status must be one of: complete, error, cancelled, rejected")
return normalized
def build_item_key(item_type: str, raw_id: Any) -> str:
"""Build a stable item key used by dismiss/history APIs."""
normalized_type = _normalize_item_type(item_type)
if normalized_type == "request":
try:
request_id = int(raw_id)
except (TypeError, ValueError) as exc:
raise ValueError("request item IDs must be integers") from exc
if request_id < 1:
raise ValueError("request item IDs must be positive integers")
return f"request:{request_id}"
if not isinstance(raw_id, str):
raise ValueError("download item IDs must be strings")
task_id = raw_id.strip()
if not task_id:
raise ValueError("download item IDs must not be empty")
return f"download:{task_id}"
def build_request_item_key(request_id: int) -> str:
"""Build a request item key."""
return build_item_key("request", request_id)
def build_download_item_key(task_id: str) -> str:
"""Build a download item key."""
return build_item_key("download", task_id)
def _parse_request_id_from_item_key(item_key: Any) -> int | None:
if not isinstance(item_key, str) or not item_key.startswith("request:"):
return None
raw_value = item_key.split(":", 1)[1].strip()
try:
parsed = int(raw_value)
except (TypeError, ValueError):
return None
return parsed if parsed > 0 else None
def _request_final_status(request_status: Any, delivery_state: Any) -> str | None:
status = str(request_status or "").strip().lower()
if status == "pending":
return None
if status == "rejected":
return "rejected"
if status == "cancelled":
return "cancelled"
if status != "fulfilled":
return None
delivery = str(delivery_state or "").strip().lower()
if delivery in {"error", "cancelled"}:
return delivery
return "complete"
class ActivityService:
"""Service for per-user activity dismissals and terminal history snapshots."""
def __init__(self, db_path: str):
self._db_path = db_path
def _connect(self) -> sqlite3.Connection:
conn = sqlite3.connect(self._db_path)
conn.row_factory = sqlite3.Row
conn.execute("PRAGMA foreign_keys = ON")
return conn
@staticmethod
def _coerce_positive_int(value: Any, field: str) -> int:
try:
parsed = int(value)
except (TypeError, ValueError) as exc:
raise ValueError(f"{field} must be an integer") from exc
if parsed < 1:
raise ValueError(f"{field} must be a positive integer")
return parsed
@staticmethod
def _row_to_dict(row: sqlite3.Row | None) -> dict[str, Any] | None:
return dict(row) if row is not None else None
@staticmethod
def _parse_json_column(value: Any) -> Any:
if not isinstance(value, str):
return None
try:
return json.loads(value)
except (ValueError, TypeError):
return None
def _build_legacy_request_snapshot(
self,
conn: sqlite3.Connection,
request_id: int,
) -> tuple[dict[str, Any] | None, str | None]:
request_row = conn.execute(
"""
SELECT
id,
user_id,
status,
delivery_state,
request_level,
book_data,
release_data,
note,
admin_note,
created_at,
reviewed_at
FROM download_requests
WHERE id = ?
""",
(request_id,),
).fetchone()
if request_row is None:
return None, None
row_dict = dict(request_row)
book_data = self._parse_json_column(row_dict.get("book_data"))
release_data = self._parse_json_column(row_dict.get("release_data"))
if not isinstance(book_data, dict):
book_data = {}
if not isinstance(release_data, dict):
release_data = {}
snapshot = {
"kind": "request",
"request": {
"id": int(row_dict["id"]),
"user_id": row_dict.get("user_id"),
"status": row_dict.get("status"),
"delivery_state": row_dict.get("delivery_state"),
"request_level": row_dict.get("request_level"),
"book_data": book_data,
"release_data": release_data,
"note": row_dict.get("note"),
"admin_note": row_dict.get("admin_note"),
"created_at": row_dict.get("created_at"),
"updated_at": row_dict.get("reviewed_at") or row_dict.get("created_at"),
},
}
final_status = _request_final_status(row_dict.get("status"), row_dict.get("delivery_state"))
return snapshot, final_status
def record_terminal_snapshot(
self,
*,
user_id: int | None,
item_type: str,
item_key: str,
origin: str,
final_status: str,
snapshot: dict[str, Any],
request_id: int | None = None,
source_id: str | None = None,
terminal_at: str | None = None,
) -> dict[str, Any]:
"""Record a durable terminal-state snapshot for an activity item."""
normalized_item_type = _normalize_item_type(item_type)
normalized_item_key = _normalize_item_key(item_key)
normalized_origin = _normalize_origin(origin)
normalized_final_status = _normalize_final_status(final_status)
if not isinstance(snapshot, dict):
raise ValueError("snapshot must be an object")
if user_id is not None:
user_id = self._coerce_positive_int(user_id, "user_id")
if request_id is not None:
request_id = self._coerce_positive_int(request_id, "request_id")
if source_id is not None and not isinstance(source_id, str):
raise ValueError("source_id must be a string when provided")
if source_id is not None:
source_id = source_id.strip() or None
effective_terminal_at = terminal_at if isinstance(terminal_at, str) and terminal_at.strip() else _now_timestamp()
serialized_snapshot = json.dumps(snapshot, separators=(",", ":"), ensure_ascii=False)
conn = self._connect()
try:
cursor = conn.execute(
"""
INSERT INTO activity_log (
user_id,
item_type,
item_key,
request_id,
source_id,
origin,
final_status,
snapshot_json,
terminal_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
user_id,
normalized_item_type,
normalized_item_key,
request_id,
source_id,
normalized_origin,
normalized_final_status,
serialized_snapshot,
effective_terminal_at,
),
)
snapshot_id = int(cursor.lastrowid)
conn.commit()
row = conn.execute(
"SELECT * FROM activity_log WHERE id = ?",
(snapshot_id,),
).fetchone()
payload = self._row_to_dict(row)
if payload is None:
raise ValueError("Failed to read back recorded activity snapshot")
return payload
finally:
conn.close()
def get_latest_activity_log_id(self, *, item_type: str, item_key: str) -> int | None:
"""Get the newest snapshot ID for an item key."""
normalized_item_type = _normalize_item_type(item_type)
normalized_item_key = _normalize_item_key(item_key)
conn = self._connect()
try:
row = conn.execute(
"""
SELECT id
FROM activity_log
WHERE item_type = ? AND item_key = ?
ORDER BY terminal_at DESC, id DESC
LIMIT 1
""",
(normalized_item_type, normalized_item_key),
).fetchone()
if row is None:
return None
return int(row["id"])
finally:
conn.close()
def dismiss_item(
self,
*,
user_id: int,
item_type: str,
item_key: str,
activity_log_id: int | None = None,
) -> dict[str, Any]:
"""Dismiss an item for a specific user (upsert)."""
normalized_user_id = self._coerce_positive_int(user_id, "user_id")
normalized_item_type = _normalize_item_type(item_type)
normalized_item_key = _normalize_item_key(item_key)
normalized_log_id = (
self._coerce_positive_int(activity_log_id, "activity_log_id")
if activity_log_id is not None
else self.get_latest_activity_log_id(
item_type=normalized_item_type,
item_key=normalized_item_key,
)
)
conn = self._connect()
try:
conn.execute(
"""
INSERT INTO activity_dismissals (
user_id,
item_type,
item_key,
activity_log_id,
dismissed_at
)
VALUES (?, ?, ?, ?, ?)
ON CONFLICT(user_id, item_type, item_key)
DO UPDATE SET
activity_log_id = excluded.activity_log_id,
dismissed_at = excluded.dismissed_at
""",
(
normalized_user_id,
normalized_item_type,
normalized_item_key,
normalized_log_id,
_now_timestamp(),
),
)
conn.commit()
row = conn.execute(
"""
SELECT *
FROM activity_dismissals
WHERE user_id = ? AND item_type = ? AND item_key = ?
""",
(normalized_user_id, normalized_item_type, normalized_item_key),
).fetchone()
payload = self._row_to_dict(row)
if payload is None:
raise ValueError("Failed to read back dismissal row")
return payload
finally:
conn.close()
def dismiss_many(self, *, user_id: int, items: Iterable[dict[str, Any]]) -> int:
"""Dismiss many items for one user."""
normalized_user_id = self._coerce_positive_int(user_id, "user_id")
normalized_items: list[tuple[str, str, int | None]] = []
for item in items:
if not isinstance(item, dict):
raise ValueError("items must contain objects")
normalized_item_type = _normalize_item_type(item.get("item_type"))
normalized_item_key = _normalize_item_key(item.get("item_key"))
raw_log_id = item.get("activity_log_id")
normalized_log_id = (
self._coerce_positive_int(raw_log_id, "activity_log_id")
if raw_log_id is not None
else self.get_latest_activity_log_id(
item_type=normalized_item_type,
item_key=normalized_item_key,
)
)
normalized_items.append((normalized_item_type, normalized_item_key, normalized_log_id))
if not normalized_items:
return 0
conn = self._connect()
try:
timestamp = _now_timestamp()
for item_type, item_key, activity_log_id in normalized_items:
conn.execute(
"""
INSERT INTO activity_dismissals (
user_id,
item_type,
item_key,
activity_log_id,
dismissed_at
)
VALUES (?, ?, ?, ?, ?)
ON CONFLICT(user_id, item_type, item_key)
DO UPDATE SET
activity_log_id = excluded.activity_log_id,
dismissed_at = excluded.dismissed_at
""",
(
normalized_user_id,
item_type,
item_key,
activity_log_id,
timestamp,
),
)
conn.commit()
return len(normalized_items)
finally:
conn.close()
def get_dismissal_set(self, user_id: int) -> list[dict[str, str]]:
"""Return dismissed item keys for one user."""
normalized_user_id = self._coerce_positive_int(user_id, "user_id")
conn = self._connect()
try:
rows = conn.execute(
"""
SELECT item_type, item_key
FROM activity_dismissals
WHERE user_id = ?
ORDER BY dismissed_at DESC, id DESC
""",
(normalized_user_id,),
).fetchall()
return [
{
"item_type": str(row["item_type"]),
"item_key": str(row["item_key"]),
}
for row in rows
]
finally:
conn.close()
def clear_dismissals_for_item_keys(
self,
*,
user_id: int,
item_type: str,
item_keys: Iterable[str],
) -> int:
"""Clear dismissals for one user + item type + item keys."""
normalized_user_id = self._coerce_positive_int(user_id, "user_id")
normalized_item_type = _normalize_item_type(item_type)
normalized_keys = {
_normalize_item_key(item_key)
for item_key in item_keys
if isinstance(item_key, str) and item_key.strip()
}
if not normalized_keys:
return 0
conn = self._connect()
try:
cursor = conn.executemany(
"""
DELETE FROM activity_dismissals
WHERE user_id = ? AND item_type = ? AND item_key = ?
""",
(
(normalized_user_id, normalized_item_type, item_key)
for item_key in normalized_keys
),
)
conn.commit()
return int(cursor.rowcount or 0)
finally:
conn.close()
def get_history(self, user_id: int, *, limit: int = 50, offset: int = 0) -> list[dict[str, Any]]:
"""Return paged dismissal history for one user."""
normalized_user_id = self._coerce_positive_int(user_id, "user_id")
normalized_limit = max(1, min(int(limit), 200))
normalized_offset = max(0, int(offset))
conn = self._connect()
try:
rows = conn.execute(
"""
SELECT
d.id,
d.user_id,
d.item_type,
d.item_key,
d.activity_log_id,
d.dismissed_at,
l.snapshot_json,
l.origin,
l.final_status,
l.terminal_at,
l.request_id,
l.source_id
FROM activity_dismissals d
LEFT JOIN activity_log l ON l.id = d.activity_log_id
WHERE d.user_id = ?
ORDER BY d.dismissed_at DESC, d.id DESC
LIMIT ? OFFSET ?
""",
(normalized_user_id, normalized_limit, normalized_offset),
).fetchall()
payload: list[dict[str, Any]] = []
for row in rows:
row_dict = dict(row)
raw_snapshot_json = row_dict.pop("snapshot_json", None)
snapshot_payload = None
if isinstance(raw_snapshot_json, str):
try:
snapshot_payload = json.loads(raw_snapshot_json)
except (ValueError, TypeError):
snapshot_payload = None
if snapshot_payload is None and row_dict.get("item_type") == "request":
request_id = row_dict.get("request_id")
if request_id is None:
request_id = _parse_request_id_from_item_key(row_dict.get("item_key"))
try:
normalized_request_id = int(request_id) if request_id is not None else None
except (TypeError, ValueError):
normalized_request_id = None
if normalized_request_id and normalized_request_id > 0:
fallback_snapshot, fallback_final_status = self._build_legacy_request_snapshot(
conn,
normalized_request_id,
)
if fallback_snapshot is not None:
snapshot_payload = fallback_snapshot
if not row_dict.get("origin"):
row_dict["origin"] = "request"
if not row_dict.get("final_status") and fallback_final_status is not None:
row_dict["final_status"] = fallback_final_status
row_dict["snapshot"] = snapshot_payload
payload.append(row_dict)
return payload
finally:
conn.close()
def get_undismissed_terminal_downloads(self, user_id: int, *, limit: int = 200) -> list[dict[str, Any]]:
"""Return latest undismissed terminal download snapshots for one user."""
normalized_user_id = self._coerce_positive_int(user_id, "user_id")
normalized_limit = max(1, min(int(limit), 500))
conn = self._connect()
try:
rows = conn.execute(
"""
SELECT
l.id,
l.user_id,
l.item_type,
l.item_key,
l.request_id,
l.source_id,
l.origin,
l.final_status,
l.snapshot_json,
l.terminal_at
FROM activity_log l
LEFT JOIN activity_dismissals d
ON d.user_id = ?
AND d.item_type = l.item_type
AND d.item_key = l.item_key
WHERE l.user_id = ?
AND l.item_type = 'download'
AND l.final_status IN ('complete', 'error', 'cancelled')
AND d.id IS NULL
ORDER BY l.terminal_at DESC, l.id DESC
LIMIT ?
""",
(normalized_user_id, normalized_user_id, normalized_limit * 2),
).fetchall()
payload: list[dict[str, Any]] = []
seen_item_keys: set[str] = set()
for row in rows:
row_dict = dict(row)
item_key = str(row_dict.get("item_key") or "")
if not item_key or item_key in seen_item_keys:
continue
seen_item_keys.add(item_key)
raw_snapshot_json = row_dict.pop("snapshot_json", None)
snapshot_payload = None
if isinstance(raw_snapshot_json, str):
try:
snapshot_payload = json.loads(raw_snapshot_json)
except (ValueError, TypeError):
snapshot_payload = None
row_dict["snapshot"] = snapshot_payload
payload.append(row_dict)
if len(payload) >= normalized_limit:
break
return payload
finally:
conn.close()
def clear_history(self, user_id: int) -> int:
"""Delete all dismissals for a user and return deleted row count."""
normalized_user_id = self._coerce_positive_int(user_id, "user_id")
conn = self._connect()
try:
cursor = conn.execute(
"DELETE FROM activity_dismissals WHERE user_id = ?",
(normalized_user_id,),
)
conn.commit()
return int(cursor.rowcount or 0)
finally:
conn.close()
+439
View File
@@ -0,0 +1,439 @@
"""Admin user management API routes.
Registers /api/admin/users CRUD endpoints for managing users.
All endpoints require admin session.
"""
from functools import wraps
import os
import sqlite3
from typing import Any
from flask import Flask, jsonify, request, session
from werkzeug.security import generate_password_hash
from shelfmark.config.booklore_settings import (
get_booklore_library_options,
get_booklore_path_options,
)
from shelfmark.config.env import CWA_DB_PATH
from shelfmark.core.admin_settings_routes import (
register_admin_settings_routes,
validate_user_settings,
)
from shelfmark.core.auth_modes import (
AUTH_SOURCE_BUILTIN,
AUTH_SOURCE_CWA,
AUTH_SOURCE_OIDC,
AUTH_SOURCE_PROXY,
determine_auth_mode,
has_local_password_admin,
normalize_auth_source,
)
from shelfmark.core.cwa_user_sync import sync_cwa_users_from_rows
from shelfmark.core.logger import setup_logger
from shelfmark.core.settings_registry import load_config_file
from shelfmark.core.user_db import UserDB
logger = setup_logger(__name__)
def _get_user_edit_capabilities(
user: dict[str, Any],
security_config: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Return backend-authored capability flags for the user edit form."""
auth_source = normalize_auth_source(
user.get("auth_source"),
user.get("oidc_subject"),
)
if security_config is None and auth_source == AUTH_SOURCE_OIDC:
security_config = load_config_file("security")
oidc_use_admin_group = bool((security_config or {}).get("OIDC_USE_ADMIN_GROUP", True))
role_managed_by_oidc_group = auth_source == AUTH_SOURCE_OIDC and oidc_use_admin_group
can_edit_role = auth_source == AUTH_SOURCE_BUILTIN or (
auth_source == AUTH_SOURCE_OIDC and not role_managed_by_oidc_group
)
return {
"authSource": auth_source,
"canSetPassword": auth_source == AUTH_SOURCE_BUILTIN,
"canEditRole": can_edit_role,
"canEditEmail": auth_source in {AUTH_SOURCE_BUILTIN, AUTH_SOURCE_PROXY},
"canEditDisplayName": auth_source != AUTH_SOURCE_OIDC,
}
def _get_auth_mode():
"""Get current auth mode from config."""
try:
config = load_config_file("security")
return determine_auth_mode(
config,
CWA_DB_PATH,
has_local_admin=has_local_password_admin(),
)
except Exception:
return "none"
def _require_admin(f):
"""Decorator to require admin session for admin routes.
In no-auth mode, everyone has access (is_admin defaults True).
In auth-required modes, requires an authenticated session with admin role.
"""
@wraps(f)
def decorated(*args, **kwargs):
auth_mode = _get_auth_mode()
if auth_mode != "none":
if "user_id" not in session:
return jsonify({"error": "Authentication required"}), 401
if not session.get("is_admin", False):
return jsonify({"error": "Admin access required"}), 403
return f(*args, **kwargs)
return decorated
def _sanitize_user(user: dict) -> dict:
"""Remove sensitive fields from user dict before returning to client."""
sanitized = dict(user)
sanitized.pop("password_hash", None)
return sanitized
def _oidc_role_management_message(security_config: dict[str, Any]) -> str:
admin_group = security_config.get("OIDC_ADMIN_GROUP", "")
if admin_group:
return (
"Admin roles for OIDC users are managed by the "
f"'{admin_group}' group in your identity provider"
)
return (
"Disable 'Use Admin Group for Authorization' in security settings "
"to manage roles manually"
)
def _is_user_active(user: dict[str, Any], auth_method: str) -> bool:
"""Determine whether a user can authenticate in the current auth mode."""
source = normalize_auth_source(user.get("auth_source"), user.get("oidc_subject"))
if source == AUTH_SOURCE_BUILTIN:
return auth_method in (AUTH_SOURCE_BUILTIN, AUTH_SOURCE_OIDC)
return source == auth_method
def _serialize_user(
user: dict[str, Any],
auth_method: str,
security_config: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Sanitize and enrich a user payload for API responses."""
payload = _sanitize_user(user)
payload["auth_source"] = normalize_auth_source(
payload.get("auth_source"),
payload.get("oidc_subject"),
)
payload["is_active"] = _is_user_active(payload, auth_method)
payload["edit_capabilities"] = _get_user_edit_capabilities(
payload,
security_config=security_config,
)
return payload
def _sync_all_cwa_users(user_db: UserDB) -> dict[str, int]:
"""Sync all users from the Calibre-Web database into users.db."""
if not CWA_DB_PATH or not CWA_DB_PATH.exists():
raise FileNotFoundError("Calibre-Web database is not available")
db_path = os.fspath(CWA_DB_PATH)
db_uri = f"file:{db_path}?mode=ro&immutable=1"
conn = sqlite3.connect(db_uri, uri=True)
try:
cur = conn.cursor()
cur.execute("SELECT name, role, email FROM user")
rows = cur.fetchall()
finally:
conn.close()
return sync_cwa_users_from_rows(user_db, rows)
def register_admin_routes(app: Flask, user_db: UserDB) -> None:
"""Register admin user management routes on the Flask app."""
@app.route("/api/admin/users", methods=["GET"])
@_require_admin
def admin_list_users():
"""List all users."""
users = user_db.list_users()
auth_mode = _get_auth_mode()
security_config = load_config_file("security")
return jsonify([
_serialize_user(u, auth_mode, security_config=security_config)
for u in users
])
@app.route("/api/admin/users", methods=["POST"])
@_require_admin
def admin_create_user():
"""Create a new user with password authentication."""
data = request.get_json() or {}
auth_mode = _get_auth_mode()
username = (data.get("username") or "").strip()
password = data.get("password", "")
email = (data.get("email") or "").strip() or None
display_name = (data.get("display_name") or "").strip() or None
role = data.get("role", "user")
if auth_mode in {AUTH_SOURCE_PROXY, AUTH_SOURCE_CWA}:
return jsonify({
"error": "Local user creation is disabled in this authentication mode",
"message": (
"Users are provisioned by your external authentication source. "
"Switch to builtin or OIDC mode to create local users."
),
}), 400
if not username:
return jsonify({"error": "Username is required"}), 400
if not password or len(password) < 4:
return jsonify({"error": "Password must be at least 4 characters"}), 400
if role not in ("admin", "user"):
return jsonify({"error": "Role must be 'admin' or 'user'"}), 400
# First user is always admin
if not user_db.list_users():
role = "admin"
# Check if username already exists
if user_db.get_user(username=username):
return jsonify({"error": "Username already exists"}), 409
password_hash = generate_password_hash(password)
try:
user = user_db.create_user(
username=username,
password_hash=password_hash,
email=email,
display_name=display_name,
auth_source=AUTH_SOURCE_BUILTIN,
role=role,
)
except ValueError:
return jsonify({"error": "Username already exists"}), 409
logger.info(
"Shelfmark user created "
f"(source=manual_admin_create, created_by={session.get('user_id', 'unknown')}, "
f"username={username}, role={role}, auth_source={AUTH_SOURCE_BUILTIN})"
)
return jsonify(
_serialize_user(
user,
_get_auth_mode(),
security_config=load_config_file("security"),
)
), 201
@app.route("/api/admin/users/<int:user_id>", methods=["GET"])
@_require_admin
def admin_get_user(user_id):
"""Get a user by ID with their settings."""
user = user_db.get_user(user_id=user_id)
if not user:
return jsonify({"error": "User not found"}), 404
result = _serialize_user(
user,
_get_auth_mode(),
security_config=load_config_file("security"),
)
result["settings"] = user_db.get_user_settings(user_id)
return jsonify(result)
@app.route("/api/admin/users/<int:user_id>", methods=["PUT"])
@_require_admin
def admin_update_user(user_id):
"""Update user fields and/or settings."""
user = user_db.get_user(user_id=user_id)
if not user:
return jsonify({"error": "User not found"}), 404
data = request.get_json() or {}
security_config = load_config_file("security")
auth_source = normalize_auth_source(
user.get("auth_source"),
user.get("oidc_subject"),
)
capabilities = _get_user_edit_capabilities(user, security_config=security_config)
# Handle optional password update
password = data.get("password", "")
if password:
if not capabilities["canSetPassword"]:
return jsonify({
"error": f"Cannot set password for {auth_source.upper()} users",
"message": "Password authentication is only available for local users.",
}), 400
if len(password) < 4:
return jsonify({"error": "Password must be at least 4 characters"}), 400
user_db.update_user(user_id, password_hash=generate_password_hash(password))
# Update user fields
user_fields = {}
for field in ("role", "email", "display_name"):
if field in data:
user_fields[field] = data[field]
if "role" in user_fields and user_fields["role"] not in ("admin", "user"):
return jsonify({"error": "Role must be 'admin' or 'user'"}), 400
role_changed = "role" in user_fields and user_fields["role"] != user.get("role")
email_changed = "email" in user_fields and user_fields["email"] != user.get("email")
display_name_changed = (
"display_name" in user_fields
and user_fields["display_name"] != user.get("display_name")
)
if role_changed and not capabilities["canEditRole"]:
if auth_source == AUTH_SOURCE_OIDC:
return jsonify({
"error": "Cannot change role for OIDC user when group-based authorization is enabled",
"message": _oidc_role_management_message(security_config),
}), 400
return jsonify({
"error": f"Cannot change role for {auth_source.upper()} users",
"message": "Role is managed by the external authentication source.",
}), 400
if email_changed and not capabilities["canEditEmail"]:
if auth_source == AUTH_SOURCE_CWA:
return jsonify({
"error": "Cannot change email for CWA users",
"message": "Email is synced from Calibre-Web.",
}), 400
return jsonify({
"error": "Cannot change email for OIDC users",
"message": "Email is managed by your identity provider.",
}), 400
if display_name_changed and not capabilities["canEditDisplayName"]:
return jsonify({
"error": "Cannot change display name for OIDC users",
"message": "Display name is managed by your identity provider.",
}), 400
# Allow demoting the last admin account.
# Auth mode resolution automatically falls back to "none" when no
# local password admin remains.
# Avoid unnecessary writes for no-op field updates.
for field in ("role", "email", "display_name"):
if field in user_fields and user_fields[field] == user.get(field):
user_fields.pop(field)
if user_fields:
user_db.update_user(user_id, **user_fields)
# Update per-user settings
if "settings" in data:
if not isinstance(data["settings"], dict):
return jsonify({"error": "Settings must be an object"}), 400
validated_settings, validation_errors = validate_user_settings(data["settings"])
if validation_errors:
return jsonify({
"error": "Invalid settings payload",
"details": validation_errors,
}), 400
user_db.set_user_settings(user_id, validated_settings)
# Ensure runtime reads see updated per-user overrides immediately.
try:
from shelfmark.core.config import config as app_config
app_config.refresh()
except Exception:
pass
updated = user_db.get_user(user_id=user_id)
result = _serialize_user(
updated,
_get_auth_mode(),
security_config=security_config,
)
result["settings"] = user_db.get_user_settings(user_id)
logger.info(f"Admin updated user {user_id}")
return jsonify(result)
@app.route("/api/admin/users/sync-cwa", methods=["POST"])
@_require_admin
def admin_sync_cwa_users():
"""Manually sync users from Calibre-Web into users.db."""
auth_mode = _get_auth_mode()
if auth_mode != AUTH_SOURCE_CWA:
return jsonify({
"error": "CWA sync is only available when CWA authentication is enabled",
}), 400
try:
summary = _sync_all_cwa_users(user_db)
except FileNotFoundError:
return jsonify({
"error": "Calibre-Web database is not available",
"message": "Verify app.db is mounted and readable at /auth/app.db.",
}), 503
except Exception as exc:
logger.error(f"Failed to sync CWA users: {exc}")
return jsonify({
"error": "Failed to sync users from Calibre-Web",
}), 500
message = (
f"Synced {summary['total']} CWA users "
f"({summary['created']} created, {summary['updated']} updated, "
f"{summary.get('deleted', 0)} deleted)."
)
logger.info(message)
return jsonify({
"success": True,
"message": message,
**summary,
})
register_admin_settings_routes(app, user_db, _require_admin)
@app.route("/api/admin/users/<int:user_id>", methods=["DELETE"])
@_require_admin
def admin_delete_user(user_id):
"""Delete a user."""
# Prevent self-deletion
if session.get("db_user_id") == user_id:
return jsonify({"error": "Cannot delete your own account"}), 400
user = user_db.get_user(user_id=user_id)
if not user:
return jsonify({"error": "User not found"}), 404
auth_mode = _get_auth_mode()
auth_source = normalize_auth_source(
user.get("auth_source"),
user.get("oidc_subject"),
)
if auth_source == AUTH_SOURCE_CWA and auth_source == auth_mode:
return jsonify({
"error": f"Cannot delete active {auth_source.upper()} users",
"message": f"{auth_source.upper()} users are automatically re-provisioned on login.",
}), 400
# Allow deleting the last local admin account.
# Auth mode resolution automatically falls back to "none" when no
# local password admin remains.
user_db.delete_user(user_id)
logger.info(f"Admin deleted user {user_id}: {user['username']}")
return jsonify({"success": True})
+230
View File
@@ -0,0 +1,230 @@
"""Admin settings-introspection routes and settings validation helpers."""
from typing import Any, Callable
from flask import Flask, jsonify, request
from shelfmark.config.notifications_settings import (
build_notification_test_result,
is_valid_notification_url,
normalize_notification_routes,
)
from shelfmark.core.settings_registry import load_config_file
from shelfmark.core.user_settings_overrides import (
build_user_preferences_payload as _build_user_preferences_payload,
get_ordered_user_overridable_fields as _get_ordered_user_overridable_fields,
get_settings_registry as _get_settings_registry,
)
from shelfmark.core.user_db import UserDB
from shelfmark.core.request_policy import parse_policy_mode, validate_policy_rules
def validate_user_settings(settings: dict[str, Any]) -> tuple[dict[str, Any], list[str]]:
settings_registry = _get_settings_registry()
field_map = settings_registry.get_settings_field_map()
overridable_map = settings_registry.get_user_overridable_fields()
valid: dict[str, Any] = {}
errors: list[str] = []
for key, value in settings.items():
if key not in field_map:
errors.append(f"Unknown setting: {key}")
elif key not in overridable_map:
errors.append(f"Setting not user-overridable: {key}")
else:
# null means "clear the per-user override; use global default"
if value is None:
valid[key] = None
continue
if key in {"REQUEST_POLICY_DEFAULT_EBOOK", "REQUEST_POLICY_DEFAULT_AUDIOBOOK"}:
if parse_policy_mode(value) is None:
errors.append(f"Invalid policy mode for {key}: {value}")
continue
if key == "REQUEST_POLICY_RULES":
normalized_rules, rule_errors = validate_policy_rules(value)
if rule_errors:
errors.extend(rule_errors)
continue
valid[key] = normalized_rules
continue
if key == "USER_NOTIFICATION_ROUTES":
normalized_routes = normalize_notification_routes(value)
invalid_count = sum(
1
for row in normalized_routes
if row.get("url") and not is_valid_notification_url(str(row.get("url")))
)
if invalid_count:
errors.append(
(
f"Invalid value for {key}: found {invalid_count} invalid URL(s). "
"Use URL values with a valid scheme, e.g. discord://... or ntfys://..."
)
)
continue
valid[key] = normalized_routes
continue
valid[key] = value
return valid, errors
def build_user_notification_test_response(
*,
user_id: int,
payload: Any,
) -> tuple[dict[str, Any], int]:
from shelfmark.core.config import config as app_config
routes_input = app_config.get("USER_NOTIFICATION_ROUTES", [], user_id=user_id)
if isinstance(payload, dict):
if "USER_NOTIFICATION_ROUTES" in payload:
routes_input = payload.get("USER_NOTIFICATION_ROUTES")
elif "routes" in payload:
routes_input = payload.get("routes")
result = build_notification_test_result(routes_input, scope_label="personal")
status_code = 200 if result.get("success", False) else 400
return result, status_code
def register_admin_settings_routes(
app: Flask,
user_db: UserDB,
require_admin: Callable[[Callable[..., Any]], Callable[..., Any]],
) -> None:
@app.route("/api/admin/download-defaults", methods=["GET"])
@require_admin
def admin_download_defaults():
config = load_config_file("downloads")
defaults = {
key: ("" if (value := config.get(key, field.default)) is None else value)
for key, field in _get_ordered_user_overridable_fields("downloads")
}
security_config = load_config_file("security")
defaults["OIDC_ADMIN_GROUP"] = security_config.get("OIDC_ADMIN_GROUP", "")
defaults["OIDC_USE_ADMIN_GROUP"] = security_config.get("OIDC_USE_ADMIN_GROUP", True)
defaults["OIDC_AUTO_PROVISION"] = security_config.get("OIDC_AUTO_PROVISION", True)
return jsonify(defaults)
@app.route("/api/admin/booklore-options", methods=["GET"])
@require_admin
def admin_booklore_options():
from shelfmark.core import admin_routes
return jsonify({
"libraries": admin_routes.get_booklore_library_options(),
"paths": admin_routes.get_booklore_path_options(),
})
@app.route("/api/admin/users/<int:user_id>/delivery-preferences", methods=["GET"])
@require_admin
def admin_get_delivery_preferences(user_id):
user = user_db.get_user(user_id=user_id)
if not user:
return jsonify({"error": "User not found"}), 404
try:
payload = _build_user_preferences_payload(user_db, user_id, "downloads")
except ValueError:
return jsonify({"error": "Downloads settings tab not found"}), 500
return jsonify(payload)
@app.route("/api/admin/users/<int:user_id>/notification-preferences", methods=["GET"])
@require_admin
def admin_get_notification_preferences(user_id):
user = user_db.get_user(user_id=user_id)
if not user:
return jsonify({"error": "User not found"}), 404
try:
payload = _build_user_preferences_payload(user_db, user_id, "notifications")
except ValueError:
return jsonify({"error": "Notifications settings tab not found"}), 500
return jsonify(payload)
@app.route("/api/admin/users/<int:user_id>/notification-preferences/test", methods=["POST"])
@require_admin
def admin_test_notification_preferences(user_id):
user = user_db.get_user(user_id=user_id)
if not user:
return jsonify({"error": "User not found"}), 404
payload = request.get_json(silent=True)
result, status_code = build_user_notification_test_response(
user_id=user_id,
payload=payload,
)
return jsonify(result), status_code
@app.route("/api/admin/settings/overrides-summary", methods=["GET"])
@require_admin
def admin_settings_overrides_summary():
settings_registry = _get_settings_registry()
tab_name = (request.args.get("tab") or "downloads").strip()
if not settings_registry.get_settings_tab(tab_name):
return jsonify({"error": f"Unknown settings tab: {tab_name}"}), 404
overridable_keys = list(settings_registry.get_user_overridable_fields(tab_name=tab_name))
keys_payload: dict[str, dict[str, Any]] = {}
for user_record in user_db.list_users():
user_settings = user_db.get_user_settings(user_record["id"])
if not isinstance(user_settings, dict):
continue
for key in overridable_keys:
if key not in user_settings or user_settings[key] is None:
continue
entry = keys_payload.setdefault(key, {"count": 0, "users": []})
entry["users"].append({
"userId": user_record["id"],
"username": user_record["username"],
"value": user_settings[key],
})
for summary in keys_payload.values():
summary["count"] = len(summary["users"])
return jsonify({"tab": tab_name, "keys": keys_payload})
@app.route("/api/admin/users/<int:user_id>/effective-settings", methods=["GET"])
@require_admin
def admin_get_effective_settings(user_id):
user = user_db.get_user(user_id=user_id)
if not user:
return jsonify({"error": "User not found"}), 404
from shelfmark.core.config import config as app_config
from shelfmark.core.settings_registry import is_value_from_env
field_map = _get_settings_registry().get_user_overridable_fields()
user_settings = user_db.get_user_settings(user_id)
tab_config_cache: dict[str, dict[str, Any]] = {}
effective: dict[str, dict[str, Any]] = {}
for key, (field, tab_name) in sorted(field_map.items()):
value = app_config.get(key, field.default, user_id=user_id)
source = "default"
if field.env_supported and is_value_from_env(field):
source = "env_var"
elif key in user_settings and user_settings[key] is not None:
source = "user_override"
value = user_settings[key]
else:
tab_config = tab_config_cache.setdefault(tab_name, load_config_file(tab_name))
if key in tab_config:
source = "global_config"
effective[key] = {"value": value, "source": source}
return jsonify(effective)
+126
View File
@@ -0,0 +1,126 @@
"""Authentication mode, auth-source normalization, and admin access policy helpers."""
import os
from typing import Any, Mapping
AUTH_SOURCE_BUILTIN = "builtin"
AUTH_SOURCE_OIDC = "oidc"
AUTH_SOURCE_PROXY = "proxy"
AUTH_SOURCE_CWA = "cwa"
AUTH_SOURCES = (
AUTH_SOURCE_BUILTIN,
AUTH_SOURCE_OIDC,
AUTH_SOURCE_PROXY,
AUTH_SOURCE_CWA,
)
AUTH_SOURCE_SET = frozenset(AUTH_SOURCES)
_ALWAYS_ADMIN_SETTINGS_TABS = frozenset({"security", "users"})
def has_local_password_admin(user_db: Any | None = None) -> bool:
"""Return True when at least one local admin with a password exists."""
try:
db = user_db
if db is None:
from shelfmark.core.user_db import UserDB
config_root = os.environ.get("CONFIG_DIR", "/config")
db = UserDB(os.path.join(config_root, "users.db"))
db.initialize()
return any(
user.get("password_hash") and user.get("role") == "admin"
for user in db.list_users()
)
except Exception:
return False
def normalize_auth_source(
source: Any,
oidc_subject: Any = None,
) -> str:
"""Resolve a stable auth source value from persisted fields."""
normalized = str(source or "").strip().lower()
if normalized in AUTH_SOURCE_SET:
return normalized
if oidc_subject:
return AUTH_SOURCE_OIDC
return AUTH_SOURCE_BUILTIN
def determine_auth_mode(
security_config: Mapping[str, Any],
cwa_db_path: Any | None,
*,
has_local_admin: bool = True,
) -> str:
"""Determine active auth mode from security config and runtime prerequisites."""
auth_mode = security_config.get("AUTH_METHOD", "none")
if auth_mode == AUTH_SOURCE_CWA and cwa_db_path:
return AUTH_SOURCE_CWA
if auth_mode == AUTH_SOURCE_BUILTIN and has_local_admin:
return AUTH_SOURCE_BUILTIN
if auth_mode == AUTH_SOURCE_PROXY and security_config.get("PROXY_AUTH_USER_HEADER"):
return AUTH_SOURCE_PROXY
if (
auth_mode == AUTH_SOURCE_OIDC
and has_local_admin
and security_config.get("OIDC_DISCOVERY_URL")
and security_config.get("OIDC_CLIENT_ID")
):
return AUTH_SOURCE_OIDC
return "none"
def is_settings_or_onboarding_path(path: str) -> bool:
"""Return True when request path targets protected admin settings routes."""
return path.startswith("/api/settings") or path.startswith("/api/onboarding")
def get_settings_tab_from_path(path: str) -> str | None:
"""Extract tab name from /api/settings/<tab>[...] paths."""
if not path.startswith("/api/settings/"):
return None
suffix = path[len("/api/settings/"):]
if not suffix:
return None
return suffix.split("/", 1)[0] or None
def should_restrict_settings_to_admin(
_users_config: Mapping[str, Any],
) -> bool:
"""Settings/onboarding is always admin-only."""
return True
def requires_admin_for_settings_access(
path: str,
users_config: Mapping[str, Any],
) -> bool:
"""Return whether this settings/onboarding request requires admin privileges."""
tab_name = get_settings_tab_from_path(path)
if tab_name in _ALWAYS_ADMIN_SETTINGS_TABS:
return True
return should_restrict_settings_to_admin(users_config)
def get_auth_check_admin_status(
_auth_mode: str,
_users_config: Mapping[str, Any],
session_data: Mapping[str, Any],
) -> bool:
"""Resolve /api/auth/check `is_admin` as the session's real admin role."""
if "user_id" not in session_data:
return False
return bool(session_data.get("is_admin", False))
+84 -2
View File
@@ -1,11 +1,14 @@
"""Configuration singleton with ENV > config file > default resolution."""
import os
import sqlite3
from threading import Lock
from typing import Any, Dict, Optional
# Import lazily to avoid circular imports
_registry_module = None
_env_module = None
_user_db_module = None
def _get_registry():
@@ -26,6 +29,15 @@ def _get_env():
return _env_module
def _get_user_db_module():
"""Lazy import of user DB module to avoid optional dependency loops."""
global _user_db_module
if _user_db_module is None:
from shelfmark.core.user_db import UserDB
_user_db_module = UserDB
return _user_db_module
class Config:
"""
Dynamic configuration singleton that provides live settings access.
@@ -36,7 +48,6 @@ class Config:
_instance: Optional['Config'] = None
_lock = Lock()
def __new__(cls) -> 'Config':
if cls._instance is None:
with cls._lock:
@@ -51,6 +62,10 @@ class Config:
self._cache: Dict[str, Any] = {}
self._field_map: Dict[str, tuple] = {} # key -> (field, tab_name)
self._cache_lock = Lock()
self._user_settings_cache: Dict[int, Dict[str, Any]] = {}
self._user_settings_cache_lock = Lock()
self._user_db = None
self._user_db_load_attempted = False
self._initialized = True
self._loaded = False
@@ -69,6 +84,7 @@ class Config:
# This handles cases where config is accessed before settings are registered
try:
import shelfmark.config.settings # noqa: F401 - main app settings
import shelfmark.config.notifications_settings # noqa: F401 - notifications settings
import shelfmark.release_sources # noqa: F401 - plugin settings
import shelfmark.metadata_providers # noqa: F401 - plugin settings
except ImportError:
@@ -111,19 +127,85 @@ class Config:
with self._cache_lock:
self._loaded = False
self._load_settings()
with self._user_settings_cache_lock:
self._user_settings_cache.clear()
self._user_db = None
self._user_db_load_attempted = False
def get(self, key: str, default: Any = None) -> Any:
def _get_user_db(self):
"""Get or initialize a UserDB handle if available."""
if self._user_db is not None:
return self._user_db
if self._user_db_load_attempted:
return None
self._user_db_load_attempted = True
try:
user_db_cls = _get_user_db_module()
db_path = os.path.join(os.environ.get("CONFIG_DIR", "/config"), "users.db")
user_db = user_db_cls(db_path)
user_db.initialize()
self._user_db = user_db
return self._user_db
except Exception:
# Multi-user support is optional; fall back to global config when unavailable.
return None
def _get_user_settings(self, user_id: int) -> Dict[str, Any]:
"""Get cached per-user settings from user DB."""
with self._user_settings_cache_lock:
if user_id in self._user_settings_cache:
return self._user_settings_cache[user_id]
user_db = self._get_user_db()
if user_db is None:
return {}
try:
settings = user_db.get_user_settings(user_id)
except (sqlite3.OperationalError, OSError, ValueError, TypeError):
return {}
if not isinstance(settings, dict):
settings = {}
with self._user_settings_cache_lock:
self._user_settings_cache[user_id] = settings
return settings
def _get_user_override(self, user_id: int, key: str) -> Any:
"""Get a user override for a specific key."""
user_settings = self._get_user_settings(user_id)
return user_settings.get(key)
def get(self, key: str, default: Any = None, user_id: Optional[int] = None) -> Any:
"""
Get a setting value by key.
Args:
key: The setting key (e.g., 'MAX_RETRY')
default: Default value if setting not found
user_id: Optional DB user ID for per-user setting overrides
Returns:
The setting value, or default if not found
"""
self._ensure_loaded()
if key in self._field_map:
field, _ = self._field_map[key]
registry = _get_registry()
# Deployment-level ENV values always win.
if field.env_supported and registry.is_value_from_env(field):
return self._cache.get(key, default)
# User overrides are only available for explicitly overridable fields.
if user_id is not None and getattr(field, "user_overridable", False):
user_value = self._get_user_override(user_id, key)
if user_value is not None:
return user_value
return self._cache.get(key, default)
def __getattr__(self, name: str) -> Any:
+94
View File
@@ -0,0 +1,94 @@
"""Helpers for provisioning and syncing Calibre-Web users into users.db."""
from __future__ import annotations
from typing import Any, Iterable
from shelfmark.core.auth_modes import AUTH_SOURCE_CWA, normalize_auth_source
from shelfmark.core.external_user_linking import upsert_external_user
from shelfmark.core.user_db import UserDB
_CWA_ALIAS_SUFFIX = "__cwa"
def _normalize_email(value: Any) -> str | None:
if value is None:
return None
email = str(value).strip()
return email or None
def upsert_cwa_user(
user_db: UserDB,
cwa_username: str,
cwa_email: str | None,
role: str,
context: str | None = None,
) -> tuple[dict[str, Any], str]:
"""Create/update a CWA-backed user with collision-safe matching."""
normalized_email = _normalize_email(cwa_email)
collision_strategy = "alias" if normalized_email else "takeover"
user, action = upsert_external_user(
user_db,
auth_source="cwa",
username=cwa_username,
email=normalized_email,
role=role,
allow_email_link=True,
collision_strategy=collision_strategy,
alias_suffix=_CWA_ALIAS_SUFFIX,
context=context,
)
if user is None:
raise RuntimeError("Unexpected CWA user sync result: no user returned")
return user, action
def sync_cwa_users_from_rows(
user_db: UserDB,
rows: Iterable[tuple[Any, Any, Any]],
) -> dict[str, int]:
"""Sync CWA users from raw `(name, role_flags, email)` rows."""
created = 0
updated = 0
active_cwa_user_ids: set[int] = set()
for username, role_flags, email in rows:
normalized_username = str(username or "").strip()
if not normalized_username:
continue
role = "admin" if (int(role_flags or 0) & 1) == 1 else "user"
user, action = upsert_cwa_user(
user_db,
cwa_username=normalized_username,
cwa_email=_normalize_email(email),
role=role,
context="cwa_manual_sync",
)
active_cwa_user_ids.add(int(user["id"]))
if action == "created":
created += 1
else:
updated += 1
deleted = 0
for existing_user in user_db.list_users():
if normalize_auth_source(
existing_user.get("auth_source"),
existing_user.get("oidc_subject"),
) != AUTH_SOURCE_CWA:
continue
existing_id = int(existing_user.get("id") or 0)
if existing_id in active_cwa_user_ids:
continue
user_db.delete_user(existing_id)
deleted += 1
return {
"created": created,
"updated": updated,
"deleted": deleted,
"total": created + updated,
}
+283
View File
@@ -0,0 +1,283 @@
"""Shared external identity matching and provisioning helpers."""
from __future__ import annotations
import re
from typing import Any, Literal
from shelfmark.core.auth_modes import normalize_auth_source
from shelfmark.core.logger import setup_logger
from shelfmark.core.user_db import UserDB
UNSET = object()
CollisionStrategy = Literal["takeover", "suffix", "alias"]
MatchReason = Literal[
"subject_match",
"existing_source_username_match",
"unique_email_match",
]
logger = setup_logger(__name__)
def _normalize_username(value: Any) -> str:
return str(value or "").strip()
def _normalize_email(value: Any) -> str | None:
if value is None:
return None
email = str(value).strip()
return email or None
def _normalize_display_name(value: Any) -> str | None:
if value is None:
return None
name = str(value).strip()
return name or None
def _email_key(value: str | None) -> str:
return (value or "").strip().lower()
def _normalize_role(value: Any) -> str:
return "admin" if str(value or "").strip().lower() == "admin" else "user"
def _get_by_subject(user_db: UserDB, subject_field: str | None, subject: str | None) -> dict[str, Any] | None:
if not subject_field or not subject:
return None
if subject_field == "oidc_subject":
return user_db.get_user(oidc_subject=subject)
return None
def find_unique_user_by_email(user_db: UserDB, email: str | None) -> dict[str, Any] | None:
key = _email_key(_normalize_email(email))
if not key:
return None
matches = [u for u in user_db.list_users() if _email_key(u.get("email")) == key]
return matches[0] if len(matches) == 1 else None
def find_external_user_match(
user_db: UserDB,
*,
auth_source: str,
username: str,
email: str | None,
subject_field: str | None = None,
subject: str | None = None,
allow_email_link: bool = False,
) -> tuple[dict[str, Any] | None, MatchReason | None]:
"""Find an existing local user that should be linked to an external identity."""
normalized_username = _normalize_username(username)
normalized_email = _normalize_email(email)
by_subject = _get_by_subject(user_db, subject_field, subject)
if by_subject is not None:
return by_subject, "subject_match"
by_username = user_db.get_user(username=normalized_username)
if by_username and normalize_auth_source(
by_username.get("auth_source"),
by_username.get("oidc_subject"),
) == auth_source:
return by_username, "existing_source_username_match"
if allow_email_link:
return find_unique_user_by_email(user_db, normalized_email), "unique_email_match"
return None, None
def _build_updates(
*,
auth_source: str,
role: str,
sync_role: bool,
email: str | None | object,
display_name: str | None | object,
subject_field: str | None,
subject: str | None,
) -> dict[str, Any]:
updates: dict[str, Any] = {"auth_source": auth_source}
if sync_role:
updates["role"] = _normalize_role(role)
if email is not UNSET:
updates["email"] = _normalize_email(email)
if display_name is not UNSET:
updates["display_name"] = _normalize_display_name(display_name)
if subject_field == "oidc_subject" and subject:
updates["oidc_subject"] = subject
return updates
def _next_suffix_username(user_db: UserDB, base_username: str) -> str:
candidate = base_username
suffix = 1
while user_db.get_user(username=candidate):
candidate = f"{base_username}_{suffix}"
suffix += 1
return candidate
def _find_existing_alias_user(
user_db: UserDB,
*,
auth_source: str,
alias_base: str,
) -> dict[str, Any] | None:
pattern = re.compile(rf"^{re.escape(alias_base)}(?:_\d+)?$")
candidates = [
user for user in user_db.list_users()
if pattern.match(str(user.get("username") or ""))
and normalize_auth_source(user.get("auth_source"), user.get("oidc_subject")) == auth_source
]
if not candidates:
return None
return sorted(candidates, key=lambda user: int(user.get("id") or 0))[0]
def _resolve_create_username(
user_db: UserDB,
*,
auth_source: str,
requested_username: str,
strategy: CollisionStrategy,
alias_suffix: str,
) -> tuple[str | None, dict[str, Any] | None, str]:
existing = user_db.get_user(username=requested_username)
if not existing:
return requested_username, None, "new_username_available"
if strategy == "takeover":
return None, existing, "username_collision_takeover"
if strategy == "suffix":
return _next_suffix_username(user_db, requested_username), None, "username_collision_suffix"
alias_base = f"{requested_username}{alias_suffix}"
alias_existing = _find_existing_alias_user(
user_db,
auth_source=auth_source,
alias_base=alias_base,
)
if alias_existing is not None:
return None, alias_existing, "reuse_existing_alias"
return _next_suffix_username(user_db, alias_base), None, "username_collision_alias"
def upsert_external_user(
user_db: UserDB,
*,
auth_source: str,
username: str,
role: str,
email: str | None | object = UNSET,
display_name: str | None | object = UNSET,
subject_field: str | None = None,
subject: str | None = None,
allow_email_link: bool = False,
sync_role: bool = True,
allow_create: bool = True,
collision_strategy: CollisionStrategy = "takeover",
alias_suffix: str | None = None,
context: str | None = None,
) -> tuple[dict[str, Any] | None, str]:
"""Create/update a user from an external auth identity.
Returns `(user, action)` where action is one of:
- `"updated"`
- `"created"`
- `"not_found"` (when `allow_create=False` and no link target exists)
"""
normalized_username = _normalize_username(username)
if not normalized_username:
raise ValueError("External username is required")
normalized_email = _normalize_email(email) if email is not UNSET else None
normalized_display_name = (
_normalize_display_name(display_name) if display_name is not UNSET else None
)
normalized_role = _normalize_role(role)
matched, match_reason = find_external_user_match(
user_db,
auth_source=auth_source,
username=normalized_username,
email=normalized_email,
subject_field=subject_field,
subject=subject,
allow_email_link=allow_email_link,
)
updates = _build_updates(
auth_source=auth_source,
role=normalized_role,
sync_role=sync_role,
email=normalized_email if email is not UNSET else UNSET,
display_name=normalized_display_name if display_name is not UNSET else UNSET,
subject_field=subject_field,
subject=subject,
)
if matched is not None:
user_db.update_user(matched["id"], **updates)
mapped = user_db.get_user(user_id=matched["id"]) or matched
logger.info(
"External user mapped to existing Shelfmark user "
f"(source={auth_source}, context={context or 'unspecified'}, reason={match_reason}, "
f"external_username={normalized_username}, shelfmark_user_id={mapped['id']}, "
f"shelfmark_username={mapped['username']})"
)
return mapped, "updated"
if not allow_create:
logger.info(
"External user could not be mapped and creation is disabled "
f"(source={auth_source}, context={context or 'unspecified'}, "
f"external_username={normalized_username})"
)
return None, "not_found"
resolved_alias_suffix = alias_suffix or f"__{auth_source}"
create_username, takeover_target, create_reason = _resolve_create_username(
user_db,
auth_source=auth_source,
requested_username=normalized_username,
strategy=collision_strategy,
alias_suffix=resolved_alias_suffix,
)
if takeover_target is not None:
user_db.update_user(takeover_target["id"], **updates)
mapped = user_db.get_user(user_id=takeover_target["id"]) or takeover_target
logger.info(
"External user mapped to existing Shelfmark user "
f"(source={auth_source}, context={context or 'unspecified'}, reason={create_reason}, "
f"external_username={normalized_username}, shelfmark_user_id={mapped['id']}, "
f"shelfmark_username={mapped['username']})"
)
return mapped, "updated"
create_kwargs: dict[str, Any] = {
"username": create_username,
"auth_source": auth_source,
"role": normalized_role,
}
if email is not UNSET:
create_kwargs["email"] = normalized_email
if display_name is not UNSET:
create_kwargs["display_name"] = normalized_display_name
if subject_field == "oidc_subject" and subject:
create_kwargs["oidc_subject"] = subject
created = user_db.create_user(**create_kwargs)
logger.info(
"External user created Shelfmark user "
f"(source={auth_source}, context={context or 'unspecified'}, reason={create_reason}, "
f"external_username={normalized_username}, shelfmark_user_id={created['id']}, "
f"shelfmark_username={created['username']})"
)
return created, "created"
+29 -13
View File
@@ -39,22 +39,38 @@ class CustomLogger(logging.Logger):
self.debug(msg, *args, exc_info=has_exception, **kwargs)
def log_resource_usage(self):
import psutil
# Best-effort only; this should never raise during exception logging.
try:
import psutil
# Sum RSS of all processes for actual app memory
app_memory_mb = 0
for proc in psutil.process_iter(['memory_info']):
# Sum RSS of all processes for actual app memory (container-friendly),
# but fall back gracefully on platforms that restrict process enumeration.
app_memory_mb = 0.0
try:
if proc.info['memory_info']:
app_memory_mb += proc.info['memory_info'].rss / (1024 * 1024)
except (psutil.NoSuchProcess, psutil.AccessDenied):
continue
for proc in psutil.process_iter(['memory_info']):
try:
mem = proc.info.get('memory_info')
if mem:
app_memory_mb += mem.rss / (1024 * 1024)
except (psutil.NoSuchProcess, psutil.AccessDenied, KeyError, AttributeError):
continue
except (PermissionError, psutil.AccessDenied, OSError):
try:
app_memory_mb = psutil.Process().memory_info().rss / (1024 * 1024)
except Exception:
app_memory_mb = 0.0
memory = psutil.virtual_memory()
system_used_mb = memory.used / (1024 * 1024)
available_mb = memory.available / (1024 * 1024)
cpu_percent = psutil.cpu_percent()
self.debug(f"Container Memory: App={app_memory_mb:.2f} MB, System={system_used_mb:.2f} MB, Available={available_mb:.2f} MB, CPU: {cpu_percent:.2f}%")
memory = psutil.virtual_memory()
system_used_mb = memory.used / (1024 * 1024)
available_mb = memory.available / (1024 * 1024)
cpu_percent = psutil.cpu_percent()
self.debug(
f"Container Memory: App={app_memory_mb:.2f} MB, System={system_used_mb:.2f} MB, "
f"Available={available_mb:.2f} MB, CPU: {cpu_percent:.2f}%"
)
except Exception:
# Avoid breaking the original log call if psutil is missing or restricted.
return
def setup_logger(name: str, log_file: Path = LOG_FILE) -> CustomLogger:
+33 -10
View File
@@ -19,10 +19,8 @@ def _get_config():
# Default mirror lists (hardcoded fallbacks)
DEFAULT_AA_MIRRORS = [
"https://annas-archive.se",
"https://annas-archive.gl",
"https://annas-archive.li",
"https://annas-archive.pm",
"https://annas-archive.in",
]
DEFAULT_LIBGEN_MIRRORS = [
@@ -52,22 +50,47 @@ def _normalize_mirror_url(url: str) -> str:
def get_aa_mirrors() -> List[str]:
"""
Get Anna's Archive mirrors from config + defaults.
Get Anna's Archive mirrors.
Returns:
List of AA mirror URLs, starting with defaults then custom additions.
Ordered list of AA mirror URLs.
If AA_MIRROR_URLS is configured, it is treated as the full list.
Otherwise, defaults are used and AA_ADDITIONAL_URLS (legacy) is appended.
Notes:
- The list is used to populate the AA mirror dropdown in Settings.
- When AA_BASE_URL is set to 'auto', mirrors are tried in the order listed.
"""
mirrors = [_normalize_mirror_url(url) for url in DEFAULT_AA_MIRRORS]
mirrors = [url for url in mirrors if url]
config = _get_config()
additional = config.get("AA_ADDITIONAL_URLS", "")
if additional:
for url in additional.split(","):
mirrors: list[str] = []
configured_list = config.get("AA_MIRROR_URLS", None)
if isinstance(configured_list, list):
for url in configured_list:
normalized = _normalize_mirror_url(str(url))
if normalized and normalized not in mirrors:
mirrors.append(normalized)
elif isinstance(configured_list, str) and configured_list.strip():
# Allow comma-separated env/manual configs.
for url in configured_list.split(","):
normalized = _normalize_mirror_url(url)
if normalized and normalized not in mirrors:
mirrors.append(normalized)
if not mirrors:
mirrors = [_normalize_mirror_url(url) for url in DEFAULT_AA_MIRRORS]
mirrors = [url for url in mirrors if url]
# Backwards-compatible append-only behavior for legacy configs/env.
additional = config.get("AA_ADDITIONAL_URLS", "")
if additional:
for url in additional.split(","):
normalized = _normalize_mirror_url(url)
if normalized and normalized not in mirrors:
mirrors.append(normalized)
return mirrors
+13 -1
View File
@@ -2,7 +2,7 @@
from dataclasses import dataclass, field
from pathlib import Path
from typing import Dict, List, Optional
from typing import Any, Dict, List, Optional
from enum import Enum
import re
import time
@@ -35,6 +35,7 @@ class QueueStatus(str, Enum):
"""Enum for possible book queue statuses."""
QUEUED = "queued"
RESOLVING = "resolving"
LOCATING = "locating"
DOWNLOADING = "downloading"
COMPLETE = "complete"
AVAILABLE = "available"
@@ -75,6 +76,7 @@ class DownloadTask:
size: Optional[str] = None
preview: Optional[str] = None
content_type: Optional[str] = None # "book (fiction)", "audiobook", "magazine", etc.
source_url: Optional[str] = None # Original release URL used by source-specific handlers
# Series info (for library naming templates)
series_name: Optional[str] = None
@@ -88,6 +90,16 @@ class DownloadTask:
# See SearchMode enum for behavioral differences
search_mode: Optional[SearchMode] = None
# Output selection for post-processing.
# This is captured at queue time so in-flight tasks are not affected if the user changes settings later.
output_mode: Optional[str] = None # e.g. "folder", "booklore", "email"
output_args: Dict[str, Any] = field(default_factory=dict) # Per-output parameters (e.g. email recipient)
# User association (multi-user support)
user_id: Optional[int] = None # DB user ID who queued this download
username: Optional[str] = None # Username for {User} template variable
request_id: Optional[int] = None # Origin request ID when queued from request fulfilment
# Runtime state
priority: int = 0
added_time: float = field(default_factory=time.time)
+71 -24
View File
@@ -10,11 +10,21 @@ from shelfmark.core.logger import setup_logger
logger = setup_logger(__name__)
TOKEN_PATTERN = re.compile(
r'\{([- ._/\[(]*)' # prefix: space, dash, dot, underscore, slash, brackets
r'([A-Za-z]+)' # token name
r'([- ._/\])]*)\}' # suffix: space, dash, dot, underscore, slash, brackets
)
# Known variable tokens, sorted longest-first to avoid partial matches
# e.g., "SeriesPosition" must match before "Series"
KNOWN_TOKENS = [
'seriesposition',
'partnumber',
'subtitle',
'author',
'series',
'title',
'year',
'user',
]
# Match any {...} block for template parsing
BRACE_PATTERN = re.compile(r'\{([^}]+)\}')
# Characters that are invalid in filenames on various filesystems
INVALID_CHARS = re.compile(r'[\\/:*?"<>|]')
@@ -88,37 +98,74 @@ def parse_naming_template(
# Normalize metadata keys to lowercase for case-insensitive matching
normalized = {k.lower(): v for k, v in metadata.items()}
def replace_token(match: re.Match) -> str:
prefix = match.group(1)
token_name = match.group(2).lower()
suffix = match.group(3)
def find_token(content: str) -> tuple[Optional[str], int]:
content_lower = content.lower()
for token in KNOWN_TOKENS:
idx = content_lower.find(token)
if idx != -1:
return token, idx
return None, -1
# Get the value for this token
value = normalized.get(token_name)
# Special handling for series position
if token_name == 'seriesposition':
def token_value(token: str) -> str:
value = normalized.get(token)
if token == 'seriesposition':
value = format_series_position(value)
# Convert to string
if value is None:
value = ""
else:
value = str(value).strip()
return ""
return str(value).strip()
# If value is empty, return empty string (no prefix/suffix)
def render_block(content: str) -> Optional[str]:
token, idx = find_token(content)
if token is None:
return None
prefix = content[:idx]
suffix = content[idx + len(token):]
value = token_value(token)
if not value:
return ""
if not allow_path_separators:
value = value.replace("/", "_")
# Sanitize the value
value = sanitize_filename(value)
return f"{prefix}{value}{suffix}"
# Replace all tokens
result = TOKEN_PATTERN.sub(replace_token, template)
# Process brace blocks in order so we can support conditional literal blocks like:
# { - Part }{PartNumber}
matches = list(BRACE_PATTERN.finditer(template))
if not matches:
result = template
else:
parts: list[str] = []
cursor = 0
for idx, match in enumerate(matches):
parts.append(template[cursor:match.start()])
content = match.group(1)
rendered = render_block(content)
if rendered is not None:
parts.append(rendered)
else:
conditional_literal = False
include_literal = False
if idx + 1 < len(matches) and match.end() == matches[idx + 1].start():
next_content = matches[idx + 1].group(1)
next_token, _next_idx = find_token(next_content)
if next_token is not None:
conditional_literal = True
include_literal = bool(token_value(next_token))
if include_literal:
parts.append(content)
elif not conditional_literal:
# Preserve blocks that look like literal text, but treat bare unknown
# placeholders as missing variables.
if re.search(r"\s", content):
parts.append(match.group(0))
cursor = match.end()
parts.append(template[cursor:])
result = "".join(parts)
# Clean up any double slashes that might result from empty tokens
result = re.sub(r'/+', '/', result)
+380
View File
@@ -0,0 +1,380 @@
"""Apprise notification dispatch for global and per-user events."""
from __future__ import annotations
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass
from enum import Enum
from typing import Any, Iterable
try:
import apprise
except Exception: # pragma: no cover - exercised in tests via monkeypatch
apprise = None # type: ignore[assignment]
from shelfmark.core.config import config as app_config
from shelfmark.core.logger import setup_logger
logger = setup_logger(__name__)
# Small pool for non-blocking dispatch. Notification sends are I/O bound and infrequent.
_executor = ThreadPoolExecutor(max_workers=2, thread_name_prefix="Notify")
_ROUTE_EVENT_ALL = "all"
_APPRISE_APP_ID = "Shelfmark"
_APPRISE_APP_DESC = "Shelfmark notifications"
_APPRISE_LOGO_URL = (
"https://raw.githubusercontent.com/calibrain/shelfmark/main/src/frontend/public/logo.png"
)
class NotificationEvent(str, Enum):
"""Global notification event identifiers."""
REQUEST_CREATED = "request_created"
REQUEST_FULFILLED = "request_fulfilled"
REQUEST_REJECTED = "request_rejected"
DOWNLOAD_COMPLETE = "download_complete"
DOWNLOAD_FAILED = "download_failed"
@dataclass
class NotificationContext:
"""Context used to render notification templates."""
event: NotificationEvent
title: str
author: str
username: str | None = None
content_type: str | None = None
format: str | None = None
source: str | None = None
admin_note: str | None = None
error_message: str | None = None
def _normalize_urls(value: Any) -> list[str]:
if value is None:
return []
raw_values: list[Any]
if isinstance(value, list):
raw_values = value
elif isinstance(value, str):
# Support legacy/manual configs.
raw_values = [segment for part in value.splitlines() for segment in part.split(",")]
else:
raw_values = [value]
normalized: list[str] = []
seen: set[str] = set()
for raw_url in raw_values:
url = str(raw_url or "").strip()
if not url:
continue
if url in seen:
continue
seen.add(url)
normalized.append(url)
return normalized
def _normalize_routes(value: Any) -> list[dict[str, str]]:
if not isinstance(value, list):
return []
allowed_events = {_ROUTE_EVENT_ALL, *(event.value for event in NotificationEvent)}
normalized: list[dict[str, str]] = []
seen: set[tuple[str, str]] = set()
for row in value:
if not isinstance(row, dict):
continue
raw_events = row.get("event")
if isinstance(raw_events, list):
event_values = raw_events
elif isinstance(raw_events, (tuple, set)):
event_values = list(raw_events)
else:
event_values = [raw_events]
url = str(row.get("url") or "").strip()
if not url:
continue
row_events: list[str] = []
for raw_event in event_values:
event = str(raw_event or "").strip().lower()
if event not in allowed_events:
continue
if event in row_events:
continue
row_events.append(event)
if _ROUTE_EVENT_ALL in row_events:
row_events = [_ROUTE_EVENT_ALL]
for event in row_events:
key = (event, url)
if key in seen:
continue
seen.add(key)
normalized.append({"event": event, "url": url})
return normalized
def _resolve_admin_routes() -> list[dict[str, str]]:
return _normalize_routes(app_config.get("ADMIN_NOTIFICATION_ROUTES", []))
def _normalize_user_id(value: Any) -> int | None:
try:
user_id = int(value)
except (TypeError, ValueError):
return None
if user_id < 1:
return None
return user_id
def _resolve_user_routes(user_id: int | None) -> list[dict[str, str]]:
normalized_user_id = _normalize_user_id(user_id)
if normalized_user_id is None:
return []
return _normalize_routes(
app_config.get("USER_NOTIFICATION_ROUTES", [], user_id=normalized_user_id)
)
def _resolve_route_urls_for_event(
routes: list[dict[str, str]],
event: NotificationEvent,
) -> list[str]:
selected: list[str] = []
seen: set[str] = set()
event_value = event.value
for row in routes:
row_event = row.get("event", "")
if row_event not in {_ROUTE_EVENT_ALL, event_value}:
continue
url = row.get("url", "")
if not url or url in seen:
continue
seen.add(url)
selected.append(url)
return selected
def _resolve_notify_type(event: NotificationEvent) -> Any:
if apprise is None:
fallback = {
NotificationEvent.REQUEST_CREATED: "info",
NotificationEvent.REQUEST_FULFILLED: "success",
NotificationEvent.REQUEST_REJECTED: "warning",
NotificationEvent.DOWNLOAD_COMPLETE: "success",
NotificationEvent.DOWNLOAD_FAILED: "failure",
}
return fallback[event]
mapping = {
NotificationEvent.REQUEST_CREATED: apprise.NotifyType.INFO,
NotificationEvent.REQUEST_FULFILLED: apprise.NotifyType.SUCCESS,
NotificationEvent.REQUEST_REJECTED: apprise.NotifyType.WARNING,
NotificationEvent.DOWNLOAD_COMPLETE: apprise.NotifyType.SUCCESS,
NotificationEvent.DOWNLOAD_FAILED: apprise.NotifyType.FAILURE,
}
return mapping[event]
def _clean_text(value: Any, fallback: str) -> str:
text = str(value or "").strip()
return text or fallback
def _render_message(context: NotificationContext) -> tuple[str, str]:
event = context.event
title = _clean_text(context.title, "Unknown title")
author = _clean_text(context.author, "Unknown author")
username = _clean_text(context.username, "A user")
if event == NotificationEvent.REQUEST_CREATED:
return "New Request", f'{username} requested "{title}" by {author}'
if event == NotificationEvent.REQUEST_FULFILLED:
return "Request Approved", f'Request for "{title}" by {author} was approved.'
if event == NotificationEvent.REQUEST_REJECTED:
note = _clean_text(context.admin_note, "")
note_line = f"\nNote: {note}" if note else ""
return "Request Rejected", f'Request for "{title}" by {author} was rejected.{note_line}'
if event == NotificationEvent.DOWNLOAD_COMPLETE:
return "Download Complete", f'"{title}" by {author} downloaded successfully.'
error_message = _clean_text(context.error_message, "")
error_line = f"\nError: {error_message}" if error_message else ""
return "Download Failed", f'Failed to download "{title}" by {author}.{error_line}'
def _dispatch_to_apprise(
urls: Iterable[str],
*,
title: str,
body: str,
notify_type: Any,
) -> dict[str, Any]:
normalized_urls = _normalize_urls(list(urls))
if not normalized_urls:
return {"success": False, "message": "No notification URLs configured"}
if apprise is None:
return {"success": False, "message": "Apprise is not installed"}
apobj = _create_apprise_client()
if apobj is None:
return {"success": False, "message": "Apprise is not installed"}
valid_urls = 0
invalid_urls = 0
for url in normalized_urls:
try:
added = bool(apobj.add(url))
except Exception:
added = False
if added:
valid_urls += 1
else:
invalid_urls += 1
if valid_urls == 0:
return {
"success": False,
"message": "No valid notification URLs configured",
}
try:
delivered = bool(apobj.notify(title=title, body=body, notify_type=notify_type))
except Exception as exc:
return {"success": False, "message": f"Notification send failed: {type(exc).__name__}: {exc}"}
if not delivered:
return {"success": False, "message": "Notification delivery failed"}
message = f"Notification sent to {valid_urls} URL(s)"
if invalid_urls:
message += f" ({invalid_urls} invalid URL(s) skipped)"
return {"success": True, "message": message}
def _create_apprise_client() -> Any:
if apprise is None:
return None
apprise_cls = getattr(apprise, "Apprise", None)
if apprise_cls is None:
return None
apprise_asset_cls = getattr(apprise, "AppriseAsset", None)
if apprise_asset_cls is None:
return apprise_cls()
try:
asset = apprise_asset_cls(
app_id=_APPRISE_APP_ID,
app_desc=_APPRISE_APP_DESC,
image_url_logo=_APPRISE_LOGO_URL,
)
except TypeError:
# Support older Apprise versions that do not expose image_url_logo.
asset = apprise_asset_cls(
app_id=_APPRISE_APP_ID,
app_desc=_APPRISE_APP_DESC,
)
except Exception:
return apprise_cls()
try:
return apprise_cls(asset=asset)
except Exception:
return apprise_cls()
def _send_admin_event(event: NotificationEvent, context: NotificationContext, urls: list[str]) -> dict[str, Any]:
title, body = _render_message(context)
notify_type = _resolve_notify_type(event)
return _dispatch_to_apprise(urls, title=title, body=body, notify_type=notify_type)
def notify_admin(event: NotificationEvent, context: NotificationContext) -> None:
"""Send a global admin notification for an event if subscribed."""
routes = _resolve_admin_routes()
urls = _resolve_route_urls_for_event(routes, event)
if not urls:
return
try:
_executor.submit(_dispatch_admin_async, event, context, urls)
except Exception as exc:
logger.warning("Failed to queue admin notification '%s': %s", event.value, exc)
def notify_user(user_id: int | None, event: NotificationEvent, context: NotificationContext) -> None:
"""Send a per-user notification for an event if subscribed."""
normalized_user_id = _normalize_user_id(user_id)
if normalized_user_id is None:
return
routes = _resolve_user_routes(normalized_user_id)
urls = _resolve_route_urls_for_event(routes, event)
if not urls:
return
try:
_executor.submit(_dispatch_user_async, normalized_user_id, event, context, urls)
except Exception as exc:
logger.warning(
"Failed to queue user notification '%s' for user_id=%s: %s",
event.value,
normalized_user_id,
exc,
)
def _dispatch_admin_async(event: NotificationEvent, context: NotificationContext, urls: list[str]) -> None:
result = _send_admin_event(event, context, urls)
if not result.get("success", False):
logger.warning("Admin notification failed for event '%s': %s", event.value, result.get("message"))
def _dispatch_user_async(
user_id: int,
event: NotificationEvent,
context: NotificationContext,
urls: list[str],
) -> None:
result = _send_admin_event(event, context, urls)
if not result.get("success", False):
logger.warning(
"User notification failed for event '%s' (user_id=%s): %s",
event.value,
user_id,
result.get("message"),
)
def send_test_notification(urls: list[str]) -> dict[str, Any]:
"""Send a synchronous test notification to the provided URLs."""
normalized_urls = _normalize_urls(urls)
if not normalized_urls:
return {"success": False, "message": "No notification URLs configured"}
test_context = NotificationContext(
event=NotificationEvent.REQUEST_CREATED,
title="Shelfmark Test Notification",
author="Shelfmark",
username="Shelfmark",
)
return _send_admin_event(NotificationEvent.REQUEST_CREATED, test_context, normalized_urls)
+80
View File
@@ -0,0 +1,80 @@
"""OIDC authentication helpers.
Handles group claim parsing, user info extraction, and user provisioning.
Flask route handlers are registered separately in main.py.
"""
from typing import Any, Dict, List, Optional
from shelfmark.core.external_user_linking import upsert_external_user
from shelfmark.core.user_db import UserDB
def parse_group_claims(id_token: Dict[str, Any], group_claim: str) -> List[str]:
"""Extract group list from an ID token claim.
Supports list, comma-separated string, or pipe-separated string.
Returns empty list if claim is missing.
"""
raw = id_token.get(group_claim)
if raw is None:
return []
if isinstance(raw, list):
return [str(g).strip() for g in raw if str(g).strip()]
if isinstance(raw, str):
delimiter = "," if "," in raw else "|"
return [g.strip() for g in raw.split(delimiter) if g.strip()]
return []
def extract_user_info(id_token: Dict[str, Any]) -> Dict[str, Any]:
"""Extract user info from OIDC ID token claims.
Returns a dict with keys: oidc_subject, username, email, display_name.
Falls back through preferred_username -> email -> sub for username.
"""
sub = id_token.get("sub", "")
email = id_token.get("email")
display_name = id_token.get("name")
username = id_token.get("preferred_username") or email or sub
return {
"oidc_subject": sub,
"username": username,
"email": email,
"display_name": display_name,
}
def provision_oidc_user(
db: UserDB,
user_info: Dict[str, Any],
is_admin: Optional[bool] = None,
allow_email_link: bool = False,
allow_create: bool = True,
) -> Optional[Dict[str, Any]]:
"""Create or update a user from OIDC claims.
Matching and collision handling use the shared external user linker:
- OIDC subject first
- optionally unique email linking (when `allow_email_link=True`)
- username conflict resolution via numeric suffix.
Returns None when no existing user is matchable and `allow_create=False`.
"""
oidc_subject = user_info["oidc_subject"]
user, _ = upsert_external_user(
db,
auth_source="oidc",
username=user_info["username"] or oidc_subject,
role="admin" if is_admin else "user",
email=user_info.get("email"),
display_name=user_info.get("display_name"),
subject_field="oidc_subject",
subject=oidc_subject,
allow_email_link=allow_email_link,
sync_role=is_admin is not None,
allow_create=allow_create,
collision_strategy="suffix",
context="oidc_login",
)
return user
+210
View File
@@ -0,0 +1,210 @@
"""OIDC Flask route handlers using Authlib.
Registers /api/auth/oidc/login and /api/auth/oidc/callback endpoints.
Business logic remains in oidc_auth.py.
"""
from typing import Any
from authlib.jose.errors import InvalidClaimError
from authlib.integrations.flask_client import OAuth
from flask import Flask, jsonify, redirect, request, session
from shelfmark.core.logger import setup_logger
from shelfmark.core.oidc_auth import (
extract_user_info,
parse_group_claims,
provision_oidc_user,
)
from shelfmark.core.settings_registry import load_config_file
from shelfmark.core.user_db import UserDB
logger = setup_logger(__name__)
oauth = OAuth()
def _normalize_claims(raw_claims: Any) -> 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 Exception:
return {}
def _is_email_verified(claims: dict[str, Any]) -> bool:
"""Normalize provider-specific email_verified values into a strict boolean."""
value = claims.get("email_verified", False)
if isinstance(value, bool):
return value
if isinstance(value, str):
return value.strip().lower() == "true"
return False
def _get_oidc_client() -> tuple[Any, dict[str, Any]]:
"""Register and return an OIDC client from the current security config."""
config = load_config_file("security")
discovery_url = config.get("OIDC_DISCOVERY_URL", "")
client_id = config.get("OIDC_CLIENT_ID", "")
if not discovery_url or not client_id:
raise ValueError("OIDC not configured")
configured_scopes = config.get("OIDC_SCOPES", ["openid", "email", "profile"])
if isinstance(configured_scopes, list):
scope_values = [str(scope).strip() for scope in configured_scopes if str(scope).strip()]
elif isinstance(configured_scopes, str):
delimiter = "," if "," in configured_scopes else " "
scope_values = [scope.strip() for scope in configured_scopes.split(delimiter) if scope.strip()]
else:
scope_values = []
scopes = list(dict.fromkeys(["openid"] + scope_values))
admin_group = config.get("OIDC_ADMIN_GROUP", "")
group_claim = config.get("OIDC_GROUP_CLAIM", "groups")
use_admin_group = 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)
oauth._clients.pop("shelfmark_idp", None)
oauth.register(
name="shelfmark_idp",
client_id=client_id,
client_secret=config.get("OIDC_CLIENT_SECRET", ""),
server_metadata_url=discovery_url,
client_kwargs={
"scope": " ".join(scopes),
"code_challenge_method": "S256",
},
overwrite=True,
)
client = oauth.create_client("shelfmark_idp")
if client is None:
raise RuntimeError("OIDC client initialization failed")
return client, config
def register_oidc_routes(app: Flask, user_db: UserDB) -> None:
"""Register OIDC authentication routes on the Flask app."""
oauth.init_app(app)
@app.route("/api/auth/oidc/login", methods=["GET"])
def oidc_login():
"""Initiate OIDC login flow and redirect to the provider."""
try:
client, _ = _get_oidc_client()
redirect_uri = request.url_root.rstrip("/") + "/api/auth/oidc/callback"
return client.authorize_redirect(redirect_uri)
except ValueError:
return jsonify({"error": "OIDC not configured"}), 500
except Exception as e:
logger.error(f"OIDC login error: {e}")
return jsonify({"error": "OIDC login failed"}), 500
@app.route("/api/auth/oidc/callback", methods=["GET"])
def oidc_callback():
"""Handle OIDC callback from identity provider."""
try:
error = request.args.get("error")
if error:
logger.warning(f"OIDC callback error from IdP: {error}")
return jsonify({"error": "Authentication failed"}), 400
client, config = _get_oidc_client()
try:
token = client.authorize_access_token()
except InvalidClaimError as e:
claim_name = getattr(e, "claim_name", "unknown")
discovery_url = str(config.get("OIDC_DISCOVERY_URL", ""))
provider_issuer = ""
try:
metadata = client.load_server_metadata()
if isinstance(metadata, dict):
provider_issuer = str(metadata.get("issuer", ""))
except Exception as metadata_error:
logger.debug(f"OIDC metadata lookup failed during claim diagnostics: {metadata_error}")
logger.error(
"OIDC callback claim validation failed: claim=%s error=%s discovery_url=%s provider_issuer=%s",
claim_name,
e,
discovery_url or "<unset>",
provider_issuer or "<unknown>",
)
if claim_name == "iss":
return (
jsonify(
{
"error": (
"OIDC issuer validation failed. Verify your discovery URL and IdP issuer/"
"external URL configuration."
)
}
),
400,
)
return jsonify({"error": f"OIDC token claim validation failed: {claim_name}"}), 400
claims = _normalize_claims(token.get("userinfo"))
# If userinfo isn't present in token payload, request it explicitly.
if not claims:
try:
claims = _normalize_claims(client.userinfo(token=token))
except TypeError:
claims = _normalize_claims(client.userinfo())
except Exception as e:
logger.error(f"Failed to fetch OIDC userinfo: {e}")
if not claims:
raise ValueError("OIDC authentication failed: missing user claims")
group_claim = config.get("OIDC_GROUP_CLAIM", "groups")
admin_group = config.get("OIDC_ADMIN_GROUP", "")
use_admin_group = config.get("OIDC_USE_ADMIN_GROUP", True)
auto_provision = config.get("OIDC_AUTO_PROVISION", True)
user_info = extract_user_info(claims)
groups = parse_group_claims(claims, group_claim)
is_admin = None
if admin_group and use_admin_group:
is_admin = admin_group in groups
allow_email_link = bool(user_info.get("email")) and _is_email_verified(claims)
user = provision_oidc_user(
user_db,
user_info,
is_admin=is_admin,
allow_email_link=allow_email_link,
allow_create=bool(auto_provision),
)
if user is None:
logger.warning(
f"OIDC login rejected: auto-provision disabled for {user_info['username']}"
)
return jsonify({"error": "Account not found. Contact your administrator."}), 403
session["user_id"] = user["username"]
session["is_admin"] = user.get("role") == "admin"
session["db_user_id"] = user["id"]
session.permanent = True
logger.info(f"OIDC login successful: {user['username']} (admin={is_admin})")
return redirect(request.script_root or "/")
except ValueError as e:
logger.error(f"OIDC callback error: {e}")
return jsonify({"error": str(e)}), 400
except Exception as e:
logger.error(f"OIDC callback error: {e}")
return jsonify({"error": "Authentication failed"}), 500
+6
View File
@@ -34,6 +34,12 @@ def _get_config_dir() -> Path:
def is_onboarding_complete() -> bool:
"""Check if onboarding has been completed."""
from shelfmark.config.env import ONBOARDING
# If onboarding is disabled via env var, treat as complete
if not ONBOARDING:
return True
config_file = _get_config_dir() / "settings.json"
if not config_file.exists():
return False
+75 -16
View File
@@ -5,7 +5,7 @@ import time
from datetime import datetime, timedelta
from pathlib import Path
from threading import Lock, Event
from typing import Dict, List, Optional, Tuple, Any
from typing import Dict, List, Optional, Tuple, Any, Callable
from shelfmark.core.config import config as app_config
from shelfmark.core.models import QueueStatus, QueueItem, DownloadTask
@@ -22,6 +22,9 @@ class BookQueue:
self._status_timestamps: dict[str, datetime] = {} # Track when each status was last updated
self._cancel_flags: dict[str, Event] = {} # Cancellation flags for active downloads
self._active_downloads: dict[str, bool] = {} # Track currently downloading tasks
self._terminal_status_hook: Optional[
Callable[[str, QueueStatus, DownloadTask], None]
] = None
@property
def _status_timeout(self) -> timedelta:
@@ -79,16 +82,47 @@ class BookQueue:
self._status[book_id] = status
self._status_timestamps[book_id] = datetime.now()
def set_terminal_status_hook(
self,
hook: Optional[Callable[[str, QueueStatus, DownloadTask], None]],
) -> None:
"""Register a callback invoked when a task first enters a terminal status."""
with self._lock:
self._terminal_status_hook = hook
def update_status(self, book_id: str, status: QueueStatus) -> None:
"""Update status of a book in the queue."""
hook: Optional[Callable[[str, QueueStatus, DownloadTask], None]] = None
hook_task: Optional[DownloadTask] = None
with self._lock:
previous_status = self._status.get(book_id)
self._update_status(book_id, status)
terminal_statuses = {
QueueStatus.COMPLETE,
QueueStatus.AVAILABLE,
QueueStatus.ERROR,
QueueStatus.DONE,
QueueStatus.CANCELLED,
}
if (
status in terminal_statuses
and previous_status != status
and self._terminal_status_hook is not None
):
current_task = self._task_data.get(book_id)
if current_task is not None:
hook = self._terminal_status_hook
hook_task = current_task
# Clean up active download tracking when finished
if status in [QueueStatus.COMPLETE, QueueStatus.AVAILABLE, QueueStatus.ERROR, QueueStatus.DONE, QueueStatus.CANCELLED]:
self._active_downloads.pop(book_id, None)
self._cancel_flags.pop(book_id, None)
if hook is not None and hook_task is not None:
hook(book_id, status, hook_task)
def update_download_path(self, task_id: str, download_path: str) -> None:
"""Update the download path of a task in the queue."""
with self._lock:
@@ -107,14 +141,22 @@ class BookQueue:
if task_id in self._task_data:
self._task_data[task_id].status_message = message
def get_status(self) -> Dict[QueueStatus, Dict[str, DownloadTask]]:
"""Get current queue status grouped by status."""
def get_status(self, user_id: Optional[int] = None) -> Dict[QueueStatus, Dict[str, DownloadTask]]:
"""Get current queue status grouped by status.
Args:
user_id: If provided, only return tasks belonging to this user
(plus legacy tasks with no user_id). If None, return all.
"""
self.refresh()
with self._lock:
result: Dict[QueueStatus, Dict[str, DownloadTask]] = {status: {} for status in QueueStatus}
for task_id, status in self._status.items():
if task_id in self._task_data:
result[status][task_id] = self._task_data[task_id]
task = self._task_data[task_id]
if user_id is not None and task.user_id is not None and task.user_id != user_id:
continue
result[status][task_id] = task
return result
def get_queue_order(self) -> List[Dict[str, Any]]:
@@ -154,17 +196,11 @@ class BookQueue:
current_status = self._status.get(task_id)
# Allow cancellation during any active state
if current_status in [QueueStatus.RESOLVING, QueueStatus.DOWNLOADING]:
if current_status in [QueueStatus.RESOLVING, QueueStatus.LOCATING, QueueStatus.DOWNLOADING]:
# Signal active download to stop
if task_id in self._cancel_flags:
self._cancel_flags[task_id].set()
self._update_status(task_id, QueueStatus.CANCELLED)
return True
elif current_status == QueueStatus.QUEUED:
# Remove from queue and mark as cancelled
self._update_status(task_id, QueueStatus.CANCELLED)
return True
elif current_status in [QueueStatus.COMPLETE, QueueStatus.DONE, QueueStatus.AVAILABLE, QueueStatus.ERROR, QueueStatus.CANCELLED]:
if current_status in [QueueStatus.COMPLETE, QueueStatus.DONE, QueueStatus.AVAILABLE, QueueStatus.ERROR, QueueStatus.CANCELLED]:
# Clear completed/errored/cancelled items from tracking
self._status.pop(task_id, None)
self._status_timestamps.pop(task_id, None)
@@ -173,7 +209,11 @@ class BookQueue:
self._active_downloads.pop(task_id, None)
return True
return False
if current_status in [QueueStatus.RESOLVING, QueueStatus.LOCATING, QueueStatus.DOWNLOADING, QueueStatus.QUEUED]:
self.update_status(task_id, QueueStatus.CANCELLED)
return True
return False
def set_priority(self, task_id: str, new_priority: int) -> bool:
"""Change the priority of a queued task (lower = higher priority)."""
@@ -245,11 +285,30 @@ class BookQueue:
return True
return any(status == QueueStatus.QUEUED for status in self._status.values())
def clear_completed(self) -> int:
"""Remove all completed, errored, or cancelled tasks from tracking."""
def clear_completed(self, user_id: Optional[int] = None) -> int:
"""Remove terminal tasks from tracking, optionally scoped to one user.
Args:
user_id: If provided, only clear tasks belonging to this user,
plus legacy tasks with no user_id. If None, clear all.
"""
terminal_statuses = {QueueStatus.COMPLETE, QueueStatus.DONE, QueueStatus.AVAILABLE, QueueStatus.ERROR, QueueStatus.CANCELLED}
with self._lock:
to_remove = [task_id for task_id, status in self._status.items() if status in terminal_statuses]
to_remove: list[str] = []
for task_id, status in self._status.items():
if status not in terminal_statuses:
continue
if user_id is None:
to_remove.append(task_id)
continue
task = self._task_data.get(task_id)
if task is None:
# Without task ownership metadata we cannot safely scope removal.
continue
if task.user_id is None or task.user_id == user_id:
to_remove.append(task_id)
for task_id in to_remove:
self._status.pop(task_id, None)
+351
View File
@@ -0,0 +1,351 @@
"""Request-policy resolution helpers.
This module is intentionally pure and side-effect free so it can be reused by
routes/services and tested independently.
"""
from __future__ import annotations
from enum import Enum
from typing import Any, Iterable, Mapping, Sequence
class PolicyMode(str, Enum):
"""Allowed request-policy modes.
Ordered from most to least permissive. The content-type default acts as a
ceiling — matrix rules can only match or restrict further, never upgrade
beyond the default.
"""
DOWNLOAD = "download"
REQUEST_RELEASE = "request_release"
REQUEST_BOOK = "request_book"
BLOCKED = "blocked"
# Permissiveness ordering: lower index = more permissive.
_MODE_PERMISSIVENESS: dict[PolicyMode, int] = {
PolicyMode.DOWNLOAD: 0,
PolicyMode.REQUEST_RELEASE: 1,
PolicyMode.REQUEST_BOOK: 2,
PolicyMode.BLOCKED: 3,
}
# Modes allowed in REQUEST_POLICY_RULES matrix rows.
MATRIX_ALLOWED_MODES = frozenset({PolicyMode.DOWNLOAD, PolicyMode.REQUEST_RELEASE, PolicyMode.BLOCKED})
def cap_mode(mode: PolicyMode, ceiling: PolicyMode) -> PolicyMode:
"""Cap a resolved mode so it cannot be more permissive than the ceiling."""
if _MODE_PERMISSIVENESS[mode] < _MODE_PERMISSIVENESS[ceiling]:
return ceiling
return mode
REQUEST_POLICY_KEYS = frozenset(
{
"REQUESTS_ENABLED",
"REQUEST_POLICY_DEFAULT_EBOOK",
"REQUEST_POLICY_DEFAULT_AUDIOBOOK",
"REQUEST_POLICY_RULES",
"MAX_PENDING_REQUESTS_PER_USER",
"REQUESTS_ALLOW_NOTES",
}
)
REQUEST_POLICY_DEFAULT_FALLBACK_MODE = PolicyMode.REQUEST_BOOK
DEFAULT_SUPPORTED_CONTENT_TYPES = ("ebook", "audiobook")
def filter_request_policy_settings(settings: Mapping[str, Any] | None) -> dict[str, Any]:
"""Return only uppercase request-policy keys from a settings JSON object."""
if not isinstance(settings, Mapping):
return {}
return {key: settings[key] for key in REQUEST_POLICY_KEYS if key in settings}
def merge_request_policy_settings(
global_settings: Mapping[str, Any] | None,
user_settings: Mapping[str, Any] | None = None,
) -> dict[str, Any]:
"""Merge global settings with per-user request-policy overrides."""
merged = filter_request_policy_settings(global_settings)
user_filtered = filter_request_policy_settings(user_settings)
# Preserve global rules by default and treat user rules as per-cell overlays.
# This allows per-user REQUEST_POLICY_RULES payloads to store only explicit
# differences instead of replacing the full global matrix.
global_rules = list(_iter_rules(merged.get("REQUEST_POLICY_RULES", [])))
user_has_rules = "REQUEST_POLICY_RULES" in user_filtered
for key, value in user_filtered.items():
if key == "REQUEST_POLICY_RULES":
continue
merged[key] = value
if user_has_rules:
merged_rules: dict[tuple[str, str], tuple[str, str, PolicyMode]] = {
(source, content_type): (source, content_type, mode)
for source, content_type, mode in global_rules
}
for source, content_type, mode in _iter_rules(user_filtered.get("REQUEST_POLICY_RULES", [])):
merged_rules[(source, content_type)] = (source, content_type, mode)
merged["REQUEST_POLICY_RULES"] = [
{"source": source, "content_type": content_type, "mode": mode.value}
for source, content_type, mode in merged_rules.values()
]
return merged
def normalize_content_type(content_type: Any) -> str:
"""Normalize arbitrary content type values to `ebook` or `audiobook`."""
if not isinstance(content_type, str):
return "ebook"
value = content_type.strip().lower()
if not value:
return "ebook"
if value in {"audiobook", "audiobooks", "audio", "book (audiobook)"}:
return "audiobook"
return "ebook"
def normalize_source(source: Any) -> str:
"""Normalize source values for policy matching."""
if not isinstance(source, str):
return "*"
value = source.strip().lower()
return value or "*"
def parse_policy_mode(mode: Any) -> PolicyMode | None:
"""Parse an arbitrary mode value into a PolicyMode enum member."""
if isinstance(mode, PolicyMode):
return mode
if not isinstance(mode, str):
return None
try:
return PolicyMode(mode.strip().lower())
except ValueError:
return None
def _normalize_rule_content_type(content_type: Any) -> str | None:
if not isinstance(content_type, str):
return None
value = content_type.strip().lower()
if not value:
return None
if value in {"*", "any"}:
return "*"
if value in {"ebook", "book", "books", "book (fiction)"}:
return "ebook"
if value in {"audiobook", "audiobooks", "audio", "book (audiobook)"}:
return "audiobook"
return None
def _normalize_rule_source(source: Any) -> str | None:
if not isinstance(source, str):
return None
value = source.strip().lower()
if not value:
return None
if value in {"*", "any"}:
return "*"
return value
def get_source_content_type_capabilities() -> dict[str, set[str]]:
"""Return source -> supported content type map from registered sources."""
try:
from shelfmark.release_sources import list_available_sources
except Exception:
return {}
capabilities: dict[str, set[str]] = {}
for source in list_available_sources():
raw_name = source.get("name")
name = normalize_source(raw_name)
if not name or name == "*":
continue
raw_types = source.get("supported_content_types", DEFAULT_SUPPORTED_CONTENT_TYPES)
if isinstance(raw_types, str) or not isinstance(raw_types, Sequence):
raw_types = DEFAULT_SUPPORTED_CONTENT_TYPES
normalized_types: set[str] = set()
for content_type in raw_types:
normalized_type = _normalize_rule_content_type(content_type)
if normalized_type and normalized_type != "*":
normalized_types.add(normalized_type)
if not normalized_types:
normalized_types = set(DEFAULT_SUPPORTED_CONTENT_TYPES)
capabilities[name] = normalized_types
return capabilities
def validate_policy_rules(
rules: Any,
source_capabilities: Mapping[str, set[str]] | None = None,
) -> tuple[list[dict[str, str]], list[str]]:
"""Validate and normalize policy rule rows.
Validation covers:
- row shape and required keys
- valid mode/content_type values
- known source names
- source/content-type compatibility from source declarations
"""
capabilities = source_capabilities if source_capabilities is not None else get_source_content_type_capabilities()
normalized_capabilities = {
normalize_source(source): {normalize_content_type(content_type) for content_type in content_types}
for source, content_types in capabilities.items()
}
normalized_rules: list[dict[str, str]] = []
errors: list[str] = []
if rules is None:
return normalized_rules, errors
if not isinstance(rules, list):
return normalized_rules, ["REQUEST_POLICY_RULES must be a list"]
for index, rule in enumerate(rules):
row_label = f"Rule {index + 1}"
if not isinstance(rule, Mapping):
errors.append(f"{row_label}: must be an object")
continue
source = _normalize_rule_source(rule.get("source"))
raw_content_type = rule.get("content_type")
content_type = _normalize_rule_content_type(rule.get("content_type"))
raw_mode = rule.get("mode")
mode = parse_policy_mode(rule.get("mode"))
if source is None:
errors.append(f"{row_label}: source is required")
continue
if (
raw_content_type is None
or (isinstance(raw_content_type, str) and not raw_content_type.strip())
):
errors.append(f"{row_label}: content_type is required")
continue
if content_type is None:
errors.append(f"{row_label}: invalid content_type '{rule.get('content_type')}'")
continue
if (
raw_mode is None
or (isinstance(raw_mode, str) and not raw_mode.strip())
):
errors.append(f"{row_label}: mode is required")
continue
if mode is None:
errors.append(f"{row_label}: invalid mode '{rule.get('mode')}'")
continue
if mode not in MATRIX_ALLOWED_MODES:
errors.append(f"{row_label}: mode '{mode.value}' is not allowed in matrix rules (use content-type defaults instead)")
continue
if source != "*" and source not in normalized_capabilities:
errors.append(f"{row_label}: unknown source '{source}'")
continue
if (
source != "*"
and content_type != "*"
and source in normalized_capabilities
and content_type not in normalized_capabilities[source]
):
errors.append(
f"{row_label}: source '{source}' does not support content_type '{content_type}'"
)
continue
normalized_rules.append(
{
"source": source,
"content_type": content_type,
"mode": mode.value,
}
)
return normalized_rules, errors
def _iter_rules(rules: Any) -> Iterable[tuple[str, str, PolicyMode]]:
if not isinstance(rules, list):
return []
normalized: list[tuple[str, str, PolicyMode]] = []
for rule in rules:
if not isinstance(rule, Mapping):
continue
source = _normalize_rule_source(rule.get("source"))
content_type = _normalize_rule_content_type(rule.get("content_type"))
mode = parse_policy_mode(rule.get("mode"))
if (
source is None
or content_type is None
or mode is None
or mode not in MATRIX_ALLOWED_MODES
):
continue
normalized.append((source, content_type, mode))
return normalized
def resolve_policy_mode(
*,
source: Any,
content_type: Any,
global_settings: Mapping[str, Any] | None,
user_settings: Mapping[str, Any] | None = None,
) -> PolicyMode:
"""Resolve an effective policy mode for a request context.
Resolution:
1. Resolve the content-type default (ceiling).
2. Match rules in specificity order.
3. Cap the matched rule at the ceiling.
4. If no rule matches, return the ceiling.
The content-type default acts as a ceiling — matrix rules can only
match or restrict further, never upgrade beyond the default.
"""
effective = merge_request_policy_settings(global_settings, user_settings)
normalized_source = normalize_source(source)
normalized_content_type = normalize_content_type(content_type)
# Resolve the content-type default (ceiling)
default_key = (
"REQUEST_POLICY_DEFAULT_AUDIOBOOK"
if normalized_content_type == "audiobook"
else "REQUEST_POLICY_DEFAULT_EBOOK"
)
default_mode = parse_policy_mode(effective.get(default_key))
ceiling = default_mode if default_mode is not None else REQUEST_POLICY_DEFAULT_FALLBACK_MODE
# Match rules in specificity order
rules = tuple(_iter_rules(effective.get("REQUEST_POLICY_RULES", [])))
candidates = (
(normalized_source, normalized_content_type),
(normalized_source, "*"),
("*", normalized_content_type),
("*", "*"),
)
for candidate_source, candidate_content_type in candidates:
for rule_source, rule_content_type, rule_mode in rules:
if rule_source == candidate_source and rule_content_type == candidate_content_type:
return cap_mode(rule_mode, ceiling)
return ceiling
+818
View File
@@ -0,0 +1,818 @@
"""Request API routes and policy snapshot endpoint."""
from __future__ import annotations
from typing import Any, Callable
from flask import Flask, jsonify, request, session
from shelfmark.core.logger import setup_logger
from shelfmark.core.request_policy import (
PolicyMode,
REQUEST_POLICY_DEFAULT_FALLBACK_MODE,
get_source_content_type_capabilities,
merge_request_policy_settings,
normalize_content_type,
normalize_source,
parse_policy_mode,
resolve_policy_mode,
)
from shelfmark.core.requests_service import (
RequestServiceError,
cancel_request,
create_request,
fulfil_request,
reject_request,
)
from shelfmark.core.activity_service import ActivityService, build_request_item_key
from shelfmark.core.notifications import (
NotificationContext,
NotificationEvent,
notify_admin,
notify_user,
)
from shelfmark.core.settings_registry import load_config_file
from shelfmark.core.user_db import UserDB
logger = setup_logger(__name__)
def _load_users_request_policy_settings() -> dict[str, Any]:
"""Load global request-policy settings from users config."""
return load_config_file("users")
def _as_bool(value: Any, default: bool = False) -> bool:
if isinstance(value, bool):
return value
if value is None:
return default
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 bool(value)
def _as_int(value: Any, default: int) -> int:
try:
parsed = int(value)
except (TypeError, ValueError):
return default
return parsed
def _error_response(
message: str,
status_code: int,
*,
code: str | None = None,
required_mode: str | None = None,
):
payload: dict[str, Any] = {"error": message}
if code is not None:
payload["code"] = code
if required_mode is not None:
payload["required_mode"] = required_mode
return jsonify(payload), status_code
def _require_request_endpoints_available(resolve_auth_mode: Callable[[], str]):
auth_mode = resolve_auth_mode()
if auth_mode == "none":
return _error_response(
"Request workflow is unavailable in no-auth mode",
403,
code="requests_unavailable",
)
if "user_id" not in session:
return jsonify({"error": "Unauthorized"}), 401
return None
def _require_db_user_id() -> tuple[int | None, Any | None]:
raw_user_id = session.get("db_user_id")
if raw_user_id is None:
return None, _error_response(
"User identity is unavailable for request workflow",
403,
code="user_identity_unavailable",
)
try:
return int(raw_user_id), None
except (TypeError, ValueError):
return None, _error_response(
"User identity is unavailable for request workflow",
403,
code="user_identity_unavailable",
)
def _resolve_effective_policy(
user_db: UserDB,
*,
db_user_id: int | None,
) -> tuple[dict[str, Any], dict[str, Any], dict[str, Any], bool]:
global_settings = _load_users_request_policy_settings()
user_settings = user_db.get_user_settings(db_user_id) if db_user_id is not None else {}
effective = merge_request_policy_settings(global_settings, user_settings)
requests_enabled = _as_bool(effective.get("REQUESTS_ENABLED"), False)
return global_settings, user_settings, effective, requests_enabled
def _emit_request_event(
ws_manager: Any,
*,
event_name: str,
payload: dict[str, Any],
room: str,
) -> None:
if ws_manager is None:
return
try:
socketio = getattr(ws_manager, "socketio", None)
is_enabled = getattr(ws_manager, "is_enabled", None)
if socketio is None or not callable(is_enabled) or not is_enabled():
return
socketio.emit(event_name, payload, to=room)
except Exception as exc:
logger.warning(f"Failed to emit WebSocket event '{event_name}' to room '{room}': {exc}")
def _extract_release_source_id(release_data: Any) -> str | None:
if not isinstance(release_data, dict):
return None
source_id = release_data.get("source_id")
if not isinstance(source_id, str):
return None
normalized = source_id.strip()
return normalized or None
def _record_terminal_request_snapshot(
activity_service: ActivityService | None,
*,
request_row: dict[str, Any],
) -> None:
if activity_service is None:
return
request_status = request_row.get("status")
if request_status not in {"rejected", "cancelled"}:
return
raw_request_id = request_row.get("id")
try:
request_id = int(raw_request_id)
except (TypeError, ValueError):
return
if request_id < 1:
return
raw_user_id = request_row.get("user_id")
try:
user_id = int(raw_user_id)
except (TypeError, ValueError):
user_id = None
source_id = _extract_release_source_id(request_row.get("release_data"))
try:
activity_service.record_terminal_snapshot(
user_id=user_id,
item_type="request",
item_key=build_request_item_key(request_id),
origin="request",
final_status=request_status,
snapshot={"kind": "request", "request": request_row},
request_id=request_id,
source_id=source_id,
)
except Exception as exc:
logger.warning("Failed to record terminal request snapshot for request %s: %s", request_id, exc)
def _normalize_optional_text(value: Any) -> str | None:
if not isinstance(value, str):
return None
normalized = value.strip()
return normalized or None
def _resolve_title_from_book_data(book_data: Any) -> str:
if isinstance(book_data, dict):
title = _normalize_optional_text(book_data.get("title"))
if title is not None:
return title
return "Unknown title"
def _resolve_request_title(request_row: dict[str, Any]) -> str:
return _resolve_title_from_book_data(request_row.get("book_data"))
def _format_user_label(username: str | None, user_id: int | None = None) -> str:
normalized_username = _normalize_optional_text(username)
if normalized_username is not None:
return normalized_username
if user_id is not None and user_id > 0:
return f"user#{user_id}"
return "unknown user"
def _resolve_request_username(
user_db: UserDB,
*,
request_row: dict[str, Any],
fallback_username: str | None = None,
) -> str | None:
normalized_fallback = _normalize_optional_text(fallback_username)
raw_user_id = request_row.get("user_id")
try:
request_user_id = int(raw_user_id)
except (TypeError, ValueError):
return normalized_fallback
requester = user_db.get_user(user_id=request_user_id)
if not isinstance(requester, dict):
return normalized_fallback
return _normalize_optional_text(requester.get("username")) or normalized_fallback
def _resolve_request_source_and_format(request_row: dict[str, Any]) -> tuple[str, str | None]:
release_data = request_row.get("release_data")
if isinstance(release_data, dict):
source = normalize_source(release_data.get("source") or request_row.get("source_hint"))
release_format = _normalize_optional_text(
release_data.get("format")
or release_data.get("filetype")
or release_data.get("extension")
)
return source, release_format
return normalize_source(request_row.get("source_hint")), None
def _resolve_request_user_id(request_row: dict[str, Any]) -> int | None:
raw_user_id = request_row.get("user_id")
try:
user_id = int(raw_user_id)
except (TypeError, ValueError):
return None
return user_id if user_id > 0 else None
def _notify_admin_for_request_event(
user_db: UserDB,
*,
event: NotificationEvent,
request_row: dict[str, Any],
fallback_username: str | None = None,
) -> None:
book_data = request_row.get("book_data")
if not isinstance(book_data, dict):
book_data = {}
source, release_format = _resolve_request_source_and_format(request_row)
context = NotificationContext(
event=event,
title=str(book_data.get("title") or "Unknown title"),
author=str(book_data.get("author") or "Unknown author"),
username=_resolve_request_username(
user_db,
request_row=request_row,
fallback_username=fallback_username,
),
content_type=normalize_content_type(
request_row.get("content_type") or book_data.get("content_type")
),
format=release_format,
source=source,
admin_note=_normalize_optional_text(request_row.get("admin_note")),
error_message=None,
)
owner_user_id = _resolve_request_user_id(request_row)
try:
notify_admin(event, context)
except Exception as exc:
logger.warning(
"Failed to trigger admin notification for request event '%s': %s",
event.value,
exc,
)
if owner_user_id is None:
return
try:
notify_user(owner_user_id, event, context)
except Exception as exc:
logger.warning(
"Failed to trigger user notification for request event '%s' (user_id=%s): %s",
event.value,
owner_user_id,
exc,
)
def register_request_routes(
app: Flask,
user_db: UserDB,
*,
resolve_auth_mode: Callable[[], str],
queue_release: Callable[..., tuple[bool, str | None]],
activity_service: ActivityService | None = None,
ws_manager: Any | None = None,
) -> None:
"""Register request policy and request lifecycle routes."""
@app.route("/api/request-policy", methods=["GET"])
def api_request_policy():
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
if auth_gate is not None:
return auth_gate
is_admin = bool(session.get("is_admin", False))
db_user_id: int | None = None
if not is_admin:
db_user_id, db_gate = _require_db_user_id()
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
global_settings, user_settings, effective, requests_enabled = _resolve_effective_policy(
user_db,
db_user_id=db_user_id,
)
default_ebook_mode = parse_policy_mode(effective.get("REQUEST_POLICY_DEFAULT_EBOOK"))
default_audio_mode = parse_policy_mode(effective.get("REQUEST_POLICY_DEFAULT_AUDIOBOOK"))
source_capabilities = get_source_content_type_capabilities()
source_modes = []
for source_name in sorted(source_capabilities):
supported_types = sorted(
source_capabilities[source_name],
key=lambda ct: (ct != "ebook", ct),
)
modes = {
content_type: resolve_policy_mode(
source=source_name,
content_type=content_type,
global_settings=global_settings,
user_settings=user_settings,
).value
for content_type in supported_types
}
source_modes.append(
{
"source": source_name,
"supported_content_types": supported_types,
"modes": modes,
}
)
return jsonify(
{
"requests_enabled": requests_enabled,
"is_admin": is_admin,
"allow_notes": _as_bool(effective.get("REQUESTS_ALLOW_NOTES"), default=True),
"defaults": {
"ebook": (
default_ebook_mode.value
if default_ebook_mode is not None
else REQUEST_POLICY_DEFAULT_FALLBACK_MODE.value
),
"audiobook": (
default_audio_mode.value
if default_audio_mode is not None
else REQUEST_POLICY_DEFAULT_FALLBACK_MODE.value
),
},
"rules": effective.get("REQUEST_POLICY_RULES", []),
"source_modes": source_modes,
}
)
@app.route("/api/requests", methods=["POST"])
def api_create_request():
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:
return db_gate
actor_username = _normalize_optional_text(session.get("user_id"))
actor_label = _format_user_label(actor_username, db_user_id)
data = request.get_json(silent=True)
if not isinstance(data, dict):
return jsonify({"error": "No data provided"}), 400
context = data.get("context") or {}
if not isinstance(context, dict):
return jsonify({"error": "context must be an object"}), 400
source = normalize_source(context.get("source"))
release_data = data.get("release_data")
request_level = context.get("request_level")
if request_level is None:
request_level = "book" if release_data is None else "release"
book_data = data.get("book_data")
if not isinstance(book_data, dict):
return jsonify({"error": "book_data must be an object"}), 400
request_title = _resolve_title_from_book_data(book_data)
content_type = normalize_content_type(
context.get("content_type")
or data.get("content_type")
or book_data.get("content_type")
)
global_settings, user_settings, effective, requests_enabled = _resolve_effective_policy(
user_db,
db_user_id=db_user_id,
)
if not requests_enabled:
logger.debug(
"Request not created for '%s' by %s: requests are disabled",
request_title,
actor_label,
)
return _error_response(
"Request workflow is disabled by policy",
403,
code="requests_unavailable",
)
max_pending = _as_int(
effective.get("MAX_PENDING_REQUESTS_PER_USER"),
default=20,
)
if max_pending < 1:
max_pending = 1
if max_pending > 1000:
max_pending = 1000
allow_notes = _as_bool(effective.get("REQUESTS_ALLOW_NOTES"), default=True)
note_value = data.get("note") if allow_notes else None
resolved_mode = resolve_policy_mode(
source=source,
content_type=content_type,
global_settings=global_settings,
user_settings=user_settings,
)
logger.debug(
"request create policy user=%s db_user_id=%s source=%s content_type=%s request_level=%s resolved_mode=%s",
session.get("user_id"),
db_user_id,
source,
content_type,
request_level,
resolved_mode.value,
)
if resolved_mode == PolicyMode.BLOCKED:
logger.debug(
"Request blocked by policy for '%s' by %s",
request_title,
actor_label,
)
return _error_response(
"Requesting is blocked by policy",
403,
code="policy_blocked",
required_mode=PolicyMode.BLOCKED.value,
)
if resolved_mode == PolicyMode.REQUEST_BOOK:
requested_level = str(request_level).strip().lower() if isinstance(request_level, str) else ""
if requested_level != "book":
logger.debug(
"Request not created for '%s' by %s: policy requires book-level requests",
request_title,
actor_label,
)
return _error_response(
"Policy requires book-level requests",
403,
code="policy_requires_request",
required_mode=PolicyMode.REQUEST_BOOK.value,
)
try:
created = create_request(
user_db,
user_id=db_user_id,
source_hint=source,
content_type=content_type,
request_level=request_level,
policy_mode=resolved_mode.value,
book_data=book_data,
release_data=release_data,
note=note_value,
max_pending_per_user=max_pending,
)
except RequestServiceError as exc:
return _error_response(str(exc), exc.status_code, code=exc.code)
event_payload = {
"request_id": created["id"],
"status": created["status"],
"title": _resolve_request_title(created),
}
logger.info(
"Request created #%s for '%s' by %s",
created["id"],
event_payload["title"],
actor_label,
)
_emit_request_event(
ws_manager,
event_name="new_request",
payload=event_payload,
room="admins",
)
_emit_request_event(
ws_manager,
event_name="request_update",
payload=event_payload,
room=f"user_{db_user_id}",
)
_notify_admin_for_request_event(
user_db,
event=NotificationEvent.REQUEST_CREATED,
request_row=created,
fallback_username=actor_username,
)
return jsonify(created), 201
@app.route("/api/requests", methods=["GET"])
def api_list_requests():
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:
return db_gate
status = request.args.get("status")
limit = request.args.get("limit", type=int)
offset = request.args.get("offset", type=int, default=0) or 0
try:
rows = user_db.list_requests(
user_id=db_user_id,
status=status,
limit=limit,
offset=offset,
)
except ValueError as exc:
return jsonify({"error": str(exc)}), 400
return jsonify(rows)
@app.route("/api/requests/<int:request_id>", methods=["DELETE"])
def api_cancel_request(request_id: int):
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:
return db_gate
try:
updated = cancel_request(
user_db,
request_id=request_id,
actor_user_id=db_user_id,
)
except RequestServiceError as exc:
return _error_response(str(exc), exc.status_code, code=exc.code)
_record_terminal_request_snapshot(activity_service, request_row=updated)
event_payload = {
"request_id": updated["id"],
"status": updated["status"],
"title": _resolve_request_title(updated),
}
actor_label = _format_user_label(_normalize_optional_text(session.get("user_id")), db_user_id)
logger.info(
"Request cancelled #%s for '%s' by %s",
updated["id"],
event_payload["title"],
actor_label,
)
_emit_request_event(
ws_manager,
event_name="request_update",
payload=event_payload,
room=f"user_{db_user_id}",
)
_emit_request_event(
ws_manager,
event_name="request_update",
payload=event_payload,
room="admins",
)
return jsonify(updated)
@app.route("/api/admin/requests", methods=["GET"])
def api_admin_list_requests():
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
if auth_gate is not None:
return auth_gate
if not session.get("is_admin", False):
return jsonify({"error": "Admin access required"}), 403
status = request.args.get("status")
limit = request.args.get("limit", type=int)
offset = request.args.get("offset", type=int, default=0) or 0
try:
rows = user_db.list_requests(status=status, limit=limit, offset=offset)
except ValueError as exc:
return jsonify({"error": str(exc)}), 400
user_cache: dict[int, str] = {}
for row in rows:
requester_id = row["user_id"]
if requester_id not in user_cache:
requester = user_db.get_user(user_id=requester_id)
user_cache[requester_id] = requester.get("username", "") if requester else ""
row["username"] = user_cache[requester_id]
return jsonify(rows)
@app.route("/api/admin/requests/count", methods=["GET"])
def api_admin_request_counts():
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
if auth_gate is not None:
return auth_gate
if not session.get("is_admin", False):
return jsonify({"error": "Admin access required"}), 403
by_status = {
status: len(user_db.list_requests(status=status))
for status in ("pending", "fulfilled", "rejected", "cancelled")
}
return jsonify(
{
"pending": by_status["pending"],
"total": sum(by_status.values()),
"by_status": by_status,
}
)
@app.route("/api/admin/requests/<int:request_id>/fulfil", methods=["POST"])
def api_admin_fulfil_request(request_id: int):
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
if auth_gate is not None:
return auth_gate
if not session.get("is_admin", False):
return jsonify({"error": "Admin access required"}), 403
raw_admin_id = session.get("db_user_id")
if raw_admin_id is None:
return jsonify({"error": "Admin user identity unavailable"}), 403
try:
admin_user_id = int(raw_admin_id)
except (TypeError, ValueError):
return jsonify({"error": "Admin user identity unavailable"}), 403
data = request.get_json(silent=True) or {}
if not isinstance(data, dict):
return jsonify({"error": "Invalid payload"}), 400
try:
updated = fulfil_request(
user_db,
request_id=request_id,
admin_user_id=admin_user_id,
queue_release=queue_release,
release_data=data.get("release_data"),
admin_note=data.get("admin_note"),
)
except RequestServiceError as exc:
return _error_response(str(exc), exc.status_code, code=exc.code)
event_payload = {
"request_id": updated["id"],
"status": updated["status"],
"title": _resolve_request_title(updated),
}
admin_label = _format_user_label(_normalize_optional_text(session.get("user_id")), admin_user_id)
requester_label = _format_user_label(
_resolve_request_username(user_db, request_row=updated),
_resolve_request_user_id(updated),
)
logger.info(
"Request fulfilled #%s for '%s' by %s (requested by %s)",
updated["id"],
event_payload["title"],
admin_label,
requester_label,
)
_emit_request_event(
ws_manager,
event_name="request_update",
payload=event_payload,
room=f"user_{updated['user_id']}",
)
_emit_request_event(
ws_manager,
event_name="request_update",
payload=event_payload,
room="admins",
)
_notify_admin_for_request_event(
user_db,
event=NotificationEvent.REQUEST_FULFILLED,
request_row=updated,
)
return jsonify(updated)
@app.route("/api/admin/requests/<int:request_id>/reject", methods=["POST"])
def api_admin_reject_request(request_id: int):
auth_gate = _require_request_endpoints_available(resolve_auth_mode)
if auth_gate is not None:
return auth_gate
if not session.get("is_admin", False):
return jsonify({"error": "Admin access required"}), 403
raw_admin_id = session.get("db_user_id")
if raw_admin_id is None:
return jsonify({"error": "Admin user identity unavailable"}), 403
try:
admin_user_id = int(raw_admin_id)
except (TypeError, ValueError):
return jsonify({"error": "Admin user identity unavailable"}), 403
data = request.get_json(silent=True) or {}
if not isinstance(data, dict):
return jsonify({"error": "Invalid payload"}), 400
try:
updated = reject_request(
user_db,
request_id=request_id,
admin_user_id=admin_user_id,
admin_note=data.get("admin_note"),
)
except RequestServiceError as exc:
return _error_response(str(exc), exc.status_code, code=exc.code)
_record_terminal_request_snapshot(activity_service, request_row=updated)
event_payload = {
"request_id": updated["id"],
"status": updated["status"],
"title": _resolve_request_title(updated),
}
admin_label = _format_user_label(_normalize_optional_text(session.get("user_id")), admin_user_id)
requester_label = _format_user_label(
_resolve_request_username(user_db, request_row=updated),
_resolve_request_user_id(updated),
)
logger.info(
"Request rejected #%s for '%s' by %s (requested by %s)",
updated["id"],
event_payload["title"],
admin_label,
requester_label,
)
_emit_request_event(
ws_manager,
event_name="request_update",
payload=event_payload,
room=f"user_{updated['user_id']}",
)
_emit_request_event(
ws_manager,
event_name="request_update",
payload=event_payload,
room="admins",
)
_notify_admin_for_request_event(
user_db,
event=NotificationEvent.REQUEST_REJECTED,
request_row=updated,
)
return jsonify(updated)
+544
View File
@@ -0,0 +1,544 @@
"""Request lifecycle helpers and service-level validation."""
from __future__ import annotations
from datetime import datetime, timezone
import json
from typing import Any, Callable, TYPE_CHECKING
from shelfmark.core.request_policy import normalize_content_type, parse_policy_mode
VALID_REQUEST_STATUSES = frozenset({"pending", "fulfilled", "rejected", "cancelled"})
TERMINAL_REQUEST_STATUSES = frozenset({"fulfilled", "rejected", "cancelled"})
VALID_REQUEST_LEVELS = frozenset({"book", "release"})
VALID_DELIVERY_STATES = frozenset(
{
"none",
"unknown",
"queued",
"resolving",
"locating",
"downloading",
"complete",
"error",
"cancelled",
}
)
MAX_REQUEST_NOTE_LENGTH = 1000
MAX_REQUEST_JSON_BLOB_BYTES = 10 * 1024
if TYPE_CHECKING:
from shelfmark.core.user_db import UserDB
class RequestServiceError(ValueError):
"""Structured error raised by request lifecycle service methods."""
def __init__(
self,
message: str,
*,
status_code: int = 400,
code: str | None = None,
):
super().__init__(message)
self.status_code = status_code
self.code = code
def normalize_request_status(status: Any) -> str:
"""Validate and normalize request status values."""
if not isinstance(status, str):
raise ValueError(f"Invalid request status: {status}")
normalized = status.strip().lower()
if normalized not in VALID_REQUEST_STATUSES:
raise ValueError(f"Invalid request status: {status}")
return normalized
def normalize_policy_mode(mode: Any) -> str:
"""Validate and normalize policy mode values."""
parsed = parse_policy_mode(mode)
if parsed is None:
raise ValueError(f"Invalid policy_mode: {mode}")
return parsed.value
def normalize_request_level(request_level: Any) -> str:
"""Validate and normalize request level values."""
if not isinstance(request_level, str):
raise ValueError(f"Invalid request_level: {request_level}")
normalized = request_level.strip().lower()
if normalized not in VALID_REQUEST_LEVELS:
raise ValueError(f"Invalid request_level: {request_level}")
return normalized
def normalize_delivery_state(state: Any) -> str:
"""Validate and normalize delivery-state values."""
if not isinstance(state, str):
raise ValueError(f"Invalid delivery_state: {state}")
normalized = state.strip().lower()
if normalized not in VALID_DELIVERY_STATES:
raise ValueError(f"Invalid delivery_state: {state}")
return normalized
def validate_request_level_payload(request_level: Any, release_data: Any) -> str:
"""Validate request_level and release_data shape coupling."""
normalized_level = normalize_request_level(request_level)
if normalized_level == "release" and release_data is None:
raise ValueError("request_level=release requires non-null release_data")
if normalized_level == "book" and release_data is not None:
raise ValueError("request_level=book requires null release_data")
return normalized_level
def validate_status_transition(current_status: Any, new_status: Any) -> tuple[str, str]:
"""Validate request status transitions and terminal immutability."""
current = normalize_request_status(current_status)
new = normalize_request_status(new_status)
if current in TERMINAL_REQUEST_STATUSES and new != current:
raise ValueError("Terminal request statuses are immutable")
return current, new
def _normalize_match_text(value: Any) -> str:
if not isinstance(value, str):
return ""
return value.strip().lower()
def normalize_note(note: Any) -> str | None:
"""Validate request notes and normalize empty strings to None."""
if note is None:
return None
if not isinstance(note, str):
raise RequestServiceError("note must be a string", status_code=400)
normalized = note.strip()
if len(normalized) > MAX_REQUEST_NOTE_LENGTH:
raise RequestServiceError(
f"note must be <= {MAX_REQUEST_NOTE_LENGTH} characters",
status_code=400,
)
return normalized or None
def _validate_book_data(book_data: Any) -> dict[str, Any]:
if not isinstance(book_data, dict):
raise RequestServiceError("book_data must be an object", status_code=400)
required_fields = ("title", "author", "provider", "provider_id")
missing = [field for field in required_fields if not _normalize_match_text(book_data.get(field))]
if missing:
raise RequestServiceError(
f"book_data missing required field(s): {', '.join(missing)}",
status_code=400,
)
return dict(book_data)
def _validate_json_blob_size(field: str, payload: Any) -> None:
if payload is None:
return
try:
serialized = json.dumps(payload, separators=(",", ":"), ensure_ascii=False)
except (TypeError, ValueError) as exc:
raise RequestServiceError(f"{field} must be JSON-serializable", status_code=400) from exc
payload_size = len(serialized.encode("utf-8"))
if payload_size > MAX_REQUEST_JSON_BLOB_BYTES:
raise RequestServiceError(
f"{field} must be <= {MAX_REQUEST_JSON_BLOB_BYTES} bytes",
status_code=400,
code="request_payload_too_large",
)
def _find_duplicate_pending_request(
user_db: "UserDB",
*,
user_id: int,
title: str,
author: str,
content_type: str,
) -> dict[str, Any] | None:
pending_rows = user_db.list_requests(user_id=user_id, status="pending")
for row in pending_rows:
row_book_data = row.get("book_data") or {}
if not isinstance(row_book_data, dict):
continue
row_title = _normalize_match_text(row_book_data.get("title"))
row_author = _normalize_match_text(row_book_data.get("author"))
row_content_type = normalize_content_type(
row.get("content_type") or row_book_data.get("content_type")
)
if row_title == title and row_author == author and row_content_type == content_type:
return row
return None
def _now_timestamp() -> str:
return datetime.now(timezone.utc).isoformat(timespec="seconds")
def _extract_release_source_id(release_data: Any) -> str | None:
if not isinstance(release_data, dict):
return None
source_id = release_data.get("source_id")
if not isinstance(source_id, str):
return None
normalized = source_id.strip()
return normalized or None
def _existing_delivery_state(request_row: dict[str, Any]) -> str:
raw_state = request_row.get("delivery_state")
if not isinstance(raw_state, str):
return "none"
normalized = raw_state.strip().lower()
return normalized if normalized in VALID_DELIVERY_STATES else "none"
def sync_delivery_states_from_queue_status(
user_db: "UserDB",
*,
queue_status: dict[str, dict[str, Any]],
user_id: int | None = None,
) -> list[dict[str, Any]]:
"""Persist delivery-state transitions for fulfilled requests based on queue status."""
source_delivery_states: dict[str, str] = {}
for status_key in ("queued", "resolving", "locating", "downloading", "complete", "error", "cancelled"):
status_bucket = queue_status.get(status_key)
if not isinstance(status_bucket, dict):
continue
for source_id in status_bucket:
source_delivery_states[source_id] = status_key
if not source_delivery_states:
return []
fulfilled_rows = user_db.list_requests(user_id=user_id, status="fulfilled")
updated: list[dict[str, Any]] = []
for row in fulfilled_rows:
source_id = _extract_release_source_id(row.get("release_data"))
if source_id is None:
continue
delivery_state = source_delivery_states.get(source_id)
if delivery_state is None:
continue
if _existing_delivery_state(row) == delivery_state:
continue
updated.append(
user_db.update_request(
row["id"],
delivery_state=delivery_state,
delivery_updated_at=_now_timestamp(),
)
)
return updated
def create_request(
user_db: "UserDB",
*,
user_id: int,
source_hint: str | None,
content_type: Any,
request_level: Any,
policy_mode: Any,
book_data: Any,
release_data: Any = None,
note: Any = None,
max_pending_per_user: int | None = None,
) -> dict[str, Any]:
"""Create a pending request after service-level validation."""
validated_book_data = _validate_book_data(book_data)
normalized_note = normalize_note(note)
normalized_content_type = normalize_content_type(
content_type or validated_book_data.get("content_type")
)
validated_book_data["content_type"] = normalized_content_type
try:
normalized_request_level = validate_request_level_payload(request_level, release_data)
normalized_policy_mode = normalize_policy_mode(policy_mode)
except ValueError as exc:
raise RequestServiceError(str(exc), status_code=400) from exc
_validate_json_blob_size("book_data", validated_book_data)
_validate_json_blob_size("release_data", release_data)
if max_pending_per_user is not None:
pending_count = user_db.count_user_pending_requests(user_id)
if pending_count >= max_pending_per_user:
raise RequestServiceError(
"Maximum pending requests reached for this user",
status_code=409,
code="max_pending_reached",
)
duplicate = _find_duplicate_pending_request(
user_db,
user_id=user_id,
title=_normalize_match_text(validated_book_data.get("title")),
author=_normalize_match_text(validated_book_data.get("author")),
content_type=normalized_content_type,
)
if duplicate is not None:
raise RequestServiceError(
"Duplicate pending request exists for this title/author/content_type",
status_code=409,
code="duplicate_pending_request",
)
try:
return user_db.create_request(
user_id=user_id,
source_hint=source_hint,
content_type=normalized_content_type,
request_level=normalized_request_level,
policy_mode=normalized_policy_mode,
book_data=validated_book_data,
release_data=release_data,
note=normalized_note,
)
except ValueError as exc:
raise RequestServiceError(str(exc), status_code=400) from exc
def ensure_request_access(
user_db: "UserDB",
*,
request_id: int,
actor_user_id: int | None,
is_admin: bool,
) -> dict[str, Any]:
"""Get request by ID and enforce ownership for non-admin actors."""
request_row = user_db.get_request(request_id)
if request_row is None:
raise RequestServiceError("Request not found", status_code=404)
if not is_admin:
if actor_user_id is None or request_row["user_id"] != actor_user_id:
raise RequestServiceError("Forbidden", status_code=403)
return request_row
def cancel_request(
user_db: "UserDB",
*,
request_id: int,
actor_user_id: int,
) -> dict[str, Any]:
"""Cancel a pending request owned by the actor."""
request_row = ensure_request_access(
user_db,
request_id=request_id,
actor_user_id=actor_user_id,
is_admin=False,
)
if request_row["status"] != "pending":
raise RequestServiceError(
"Request is already in a terminal state",
status_code=409,
code="stale_transition",
)
try:
return user_db.update_request(
request_id,
expected_current_status="pending",
status="cancelled",
)
except ValueError as exc:
raise RequestServiceError(str(exc), status_code=409, code="stale_transition") from exc
def reject_request(
user_db: "UserDB",
*,
request_id: int,
admin_user_id: int,
admin_note: Any = None,
) -> dict[str, Any]:
"""Reject a pending request as admin."""
request_row = ensure_request_access(
user_db,
request_id=request_id,
actor_user_id=admin_user_id,
is_admin=True,
)
if request_row["status"] != "pending":
raise RequestServiceError(
"Request is already in a terminal state",
status_code=409,
code="stale_transition",
)
normalized_admin_note = None
if admin_note is not None:
if not isinstance(admin_note, str):
raise RequestServiceError("admin_note must be a string", status_code=400)
normalized_admin_note = admin_note.strip() or None
try:
return user_db.update_request(
request_id,
expected_current_status="pending",
status="rejected",
admin_note=normalized_admin_note,
reviewed_by=admin_user_id,
reviewed_at=_now_timestamp(),
)
except ValueError as exc:
raise RequestServiceError(str(exc), status_code=409, code="stale_transition") from exc
def fulfil_request(
user_db: "UserDB",
*,
request_id: int,
admin_user_id: int,
queue_release: Callable[..., tuple[bool, str | None]],
release_data: Any = None,
admin_note: Any = None,
) -> dict[str, Any]:
"""Fulfil a pending request and queue the release under requesting-user identity."""
request_row = ensure_request_access(
user_db,
request_id=request_id,
actor_user_id=admin_user_id,
is_admin=True,
)
if request_row["status"] != "pending":
raise RequestServiceError(
"Request is already in a terminal state",
status_code=409,
code="stale_transition",
)
normalized_admin_note = None
if admin_note is not None:
if not isinstance(admin_note, str):
raise RequestServiceError("admin_note must be a string", status_code=400)
normalized_admin_note = admin_note.strip() or None
selected_release_data = release_data if release_data is not None else request_row.get("release_data")
if selected_release_data is not None and not isinstance(selected_release_data, dict):
raise RequestServiceError("release_data must be an object", status_code=400)
if request_row["request_level"] == "book" and selected_release_data is None:
raise RequestServiceError(
"release_data is required to fulfil book-level requests",
status_code=400,
)
if request_row["request_level"] == "release" and selected_release_data is None:
raise RequestServiceError(
"release_data is required to fulfil release-level requests",
status_code=400,
)
_validate_json_blob_size("release_data", selected_release_data)
requester = user_db.get_user(user_id=request_row["user_id"])
if requester is None:
raise RequestServiceError("Requesting user not found", status_code=404)
queued_release_data = dict(selected_release_data)
queued_release_data["_request_id"] = request_id
success, error = queue_release(
queued_release_data,
0,
user_id=request_row["user_id"],
username=requester.get("username"),
)
if not success:
raise RequestServiceError(
error or "Failed to queue release",
status_code=409,
code="queue_failed",
)
try:
return user_db.update_request(
request_id,
expected_current_status="pending",
status="fulfilled",
release_data=selected_release_data,
delivery_state="queued",
delivery_updated_at=_now_timestamp(),
last_failure_reason=None,
admin_note=normalized_admin_note,
reviewed_by=admin_user_id,
reviewed_at=_now_timestamp(),
)
except ValueError as exc:
raise RequestServiceError(str(exc), status_code=409, code="stale_transition") from exc
def reopen_failed_request(
user_db: "UserDB",
*,
request_id: int,
failure_reason: str | None = None,
) -> dict[str, Any] | None:
"""Reopen a failed fulfilled request so admins can re-approve with a new release."""
normalized_failure_reason = None
if isinstance(failure_reason, str):
normalized_failure_reason = failure_reason.strip() or None
with user_db._lock:
conn = user_db._connect()
try:
current_row = conn.execute(
"SELECT * FROM download_requests WHERE id = ?",
(request_id,),
).fetchone()
current_request = user_db._parse_request_row(current_row)
if current_request is None:
return None
if current_request.get("status") != "fulfilled":
return None
current_delivery_state = _existing_delivery_state(current_request)
# Terminal hook callbacks can run before delivery-state sync persists "error".
# Allow reopening fulfilled requests unless they are already complete.
if current_delivery_state == "complete":
return None
if current_delivery_state not in {"error", "cancelled"} and normalized_failure_reason is None:
return None
conn.execute(
"""
UPDATE download_requests
SET status = 'pending',
delivery_state = 'none',
delivery_updated_at = NULL,
release_data = NULL,
last_failure_reason = ?,
reviewed_by = NULL,
reviewed_at = NULL
WHERE id = ?
""",
(normalized_failure_reason, request_id),
)
updated_row = conn.execute(
"SELECT * FROM download_requests WHERE id = ?",
(request_id,),
).fetchone()
conn.commit()
return user_db._parse_request_row(updated_row)
finally:
conn.close()
+4
View File
@@ -36,6 +36,7 @@ class ReleaseSearchPlan:
title_variants: List[ReleaseSearchVariant]
grouped_title_variants: List[ReleaseSearchVariant]
manual_query: Optional[str] = None
indexers: Optional[List[str]] = None # Indexer names for Prowlarr (overrides settings)
@property
def primary_query(self) -> str:
@@ -86,6 +87,7 @@ def build_release_search_plan(
book: BookMetadata,
languages: Optional[List[str]] = None,
manual_query: Optional[str] = None,
indexers: Optional[List[str]] = None,
) -> ReleaseSearchPlan:
resolved_languages = _normalize_languages(languages)
@@ -106,6 +108,7 @@ def build_release_search_plan(
title_variants=[variant],
grouped_title_variants=[variant],
manual_query=resolved_manual_query,
indexers=indexers,
)
isbn_candidates: List[str] = []
@@ -161,4 +164,5 @@ def build_release_search_plan(
title_variants=title_variants,
grouped_title_variants=grouped_variants,
manual_query=None,
indexers=indexers,
)
+354
View File
@@ -0,0 +1,354 @@
"""Self-service user account routes."""
from functools import wraps
from typing import Any, Callable, Mapping
from flask import Flask, jsonify, request, session
from werkzeug.security import generate_password_hash
from shelfmark.config.env import CWA_DB_PATH
from shelfmark.core.admin_settings_routes import (
build_user_notification_test_response,
validate_user_settings,
)
from shelfmark.core.auth_modes import (
AUTH_SOURCE_BUILTIN,
AUTH_SOURCE_CWA,
AUTH_SOURCE_OIDC,
AUTH_SOURCE_PROXY,
determine_auth_mode,
has_local_password_admin,
normalize_auth_source,
)
from shelfmark.core.logger import setup_logger
from shelfmark.core.settings_registry import load_config_file
from shelfmark.core.user_settings_overrides import (
build_user_preferences_payload as _build_user_preferences_payload,
get_ordered_user_overridable_fields as _get_ordered_user_overridable_fields,
)
from shelfmark.core.user_db import UserDB
logger = setup_logger(__name__)
MIN_PASSWORD_LENGTH = 4
_VISIBLE_SELF_SETTINGS_SECTIONS_KEY = "VISIBLE_SELF_SETTINGS_SECTIONS"
_SELF_SETTINGS_SECTION_DELIVERY = "delivery"
_SELF_SETTINGS_SECTION_NOTIFICATIONS = "notifications"
_VALID_SELF_SETTINGS_SECTIONS = (
_SELF_SETTINGS_SECTION_DELIVERY,
_SELF_SETTINGS_SECTION_NOTIFICATIONS,
)
_DEFAULT_VISIBLE_SELF_SETTINGS_SECTIONS = list(_VALID_SELF_SETTINGS_SECTIONS)
def _get_auth_mode() -> str:
"""Get current auth mode from config."""
try:
config = load_config_file("security")
return determine_auth_mode(
config,
CWA_DB_PATH,
has_local_admin=has_local_password_admin(),
)
except Exception:
return "none"
def _require_authenticated_user(f: Callable[..., Any]) -> Callable[..., Any]:
"""Decorator requiring an authenticated session linked to a local user row."""
@wraps(f)
def decorated(*args, **kwargs):
auth_mode = _get_auth_mode()
if auth_mode != "none" and "user_id" not in session:
return jsonify({"error": "Authentication required"}), 401
if "db_user_id" not in session:
return jsonify({"error": "Authenticated session is missing local user context"}), 403
return f(*args, **kwargs)
return decorated
def _get_current_user(user_db: UserDB) -> tuple[int | None, dict[str, Any] | None, tuple[Any, int] | None]:
raw_user_id = session.get("db_user_id")
try:
user_id = int(raw_user_id)
except (TypeError, ValueError):
return None, None, (jsonify({"error": "Invalid user context"}), 400)
user = user_db.get_user(user_id=user_id)
if not user:
return None, None, (jsonify({"error": "User not found"}), 404)
return user_id, user, None
def _is_user_active(user: Mapping[str, Any], auth_method: str) -> bool:
source = normalize_auth_source(user.get("auth_source"), user.get("oidc_subject"))
if source == AUTH_SOURCE_BUILTIN:
return auth_method in (AUTH_SOURCE_BUILTIN, AUTH_SOURCE_OIDC)
return source == auth_method
def _get_self_edit_capabilities(user: Mapping[str, Any]) -> dict[str, Any]:
auth_source = normalize_auth_source(
user.get("auth_source"),
user.get("oidc_subject"),
)
return {
"authSource": auth_source,
"canSetPassword": auth_source == AUTH_SOURCE_BUILTIN,
"canEditRole": False,
"canEditEmail": auth_source in {AUTH_SOURCE_BUILTIN, AUTH_SOURCE_PROXY},
"canEditDisplayName": auth_source != AUTH_SOURCE_OIDC,
}
def _serialize_self_user(user: Mapping[str, Any], auth_mode: str) -> dict[str, Any]:
payload = dict(user)
payload.pop("password_hash", None)
payload["auth_source"] = normalize_auth_source(
payload.get("auth_source"),
payload.get("oidc_subject"),
)
payload["is_active"] = _is_user_active(payload, auth_mode)
payload["edit_capabilities"] = _get_self_edit_capabilities(payload)
return payload
def _normalize_visible_self_settings_sections(raw_sections: Any) -> list[str]:
"""Normalize users.VISIBLE_SELF_SETTINGS_SECTIONS to a safe ordered list."""
if raw_sections is None:
return list(_DEFAULT_VISIBLE_SELF_SETTINGS_SECTIONS)
if isinstance(raw_sections, str):
candidate_sections = [s.strip() for s in raw_sections.split(",") if s.strip()]
elif isinstance(raw_sections, (list, tuple, set)):
candidate_sections = [str(section).strip() for section in raw_sections if str(section).strip()]
else:
return list(_DEFAULT_VISIBLE_SELF_SETTINGS_SECTIONS)
normalized_sections: list[str] = []
for section in candidate_sections:
if section in _VALID_SELF_SETTINGS_SECTIONS and section not in normalized_sections:
normalized_sections.append(section)
if not normalized_sections and candidate_sections:
# Invalid non-empty config should fail-safe to showing defaults.
return list(_DEFAULT_VISIBLE_SELF_SETTINGS_SECTIONS)
return normalized_sections
def _get_visible_self_settings_sections() -> list[str]:
users_config = load_config_file("users")
raw_sections = users_config.get(_VISIBLE_SELF_SETTINGS_SECTIONS_KEY)
return _normalize_visible_self_settings_sections(raw_sections)
def _get_allowed_self_settings_keys(visible_sections: list[str]) -> set[str]:
allowed_keys: set[str] = set()
visible_sections_set = set(visible_sections)
if _SELF_SETTINGS_SECTION_DELIVERY in visible_sections_set:
allowed_keys |= {
key for key, _field in _get_ordered_user_overridable_fields("downloads")
}
if _SELF_SETTINGS_SECTION_NOTIFICATIONS in visible_sections_set:
allowed_keys |= {
key for key, _field in _get_ordered_user_overridable_fields("notifications")
}
return allowed_keys
def register_self_user_routes(app: Flask, user_db: UserDB) -> None:
"""Register self-service user endpoints."""
@app.route("/api/users/me/edit-context", methods=["GET"])
@_require_authenticated_user
def users_me_edit_context():
user_id, user, user_error = _get_current_user(user_db)
if user_error:
return user_error
auth_mode = _get_auth_mode()
serialized_user = _serialize_self_user(user, auth_mode)
serialized_user["settings"] = user_db.get_user_settings(user_id)
visible_self_settings_sections = _get_visible_self_settings_sections()
delivery_preferences = None
if _SELF_SETTINGS_SECTION_DELIVERY in visible_self_settings_sections:
try:
delivery_preferences = _build_user_preferences_payload(user_db, user_id, "downloads")
except ValueError:
return jsonify({"error": "Downloads settings tab not found"}), 500
except Exception as exc:
logger.warning(f"Failed to build user delivery preferences for user_id={user_id}: {exc}")
delivery_preferences = None
notification_preferences = None
if _SELF_SETTINGS_SECTION_NOTIFICATIONS in visible_self_settings_sections:
try:
notification_preferences = _build_user_preferences_payload(user_db, user_id, "notifications")
except ValueError:
return jsonify({"error": "Notifications settings tab not found"}), 500
except Exception as exc:
logger.warning(f"Failed to build user notification preferences for user_id={user_id}: {exc}")
notification_preferences = None
user_overridable_keys = sorted(
set(delivery_preferences.get("keys", []) if delivery_preferences else [])
| set(notification_preferences.get("keys", []) if notification_preferences else [])
)
return jsonify(
{
"user": serialized_user,
"deliveryPreferences": delivery_preferences,
"notificationPreferences": notification_preferences,
"userOverridableKeys": user_overridable_keys,
"visibleUserSettingsSections": visible_self_settings_sections,
}
)
@app.route("/api/users/me/notification-preferences/test", methods=["POST"])
@_require_authenticated_user
def users_me_test_notification_preferences():
user_id, _user, user_error = _get_current_user(user_db)
if user_error:
return user_error
if user_id is None:
return jsonify({"error": "User not found"}), 404
payload = request.get_json(silent=True)
result, status_code = build_user_notification_test_response(
user_id=user_id,
payload=payload,
)
return jsonify(result), status_code
@app.route("/api/users/me", methods=["PUT"])
@_require_authenticated_user
def users_me_update():
user_id, user, user_error = _get_current_user(user_db)
if user_error:
return user_error
data = request.get_json() or {}
if not isinstance(data, dict):
return jsonify({"error": "Request body must be a JSON object"}), 400
capabilities = _get_self_edit_capabilities(user)
auth_source = capabilities["authSource"]
password = data.get("password", "")
if password:
if not capabilities["canSetPassword"]:
return jsonify(
{
"error": f"Cannot set password for {auth_source.upper()} users",
"message": "Password authentication is only available for local users.",
}
), 400
if len(password) < MIN_PASSWORD_LENGTH:
return jsonify({"error": f"Password must be at least {MIN_PASSWORD_LENGTH} characters"}), 400
user_db.update_user(user_id, password_hash=generate_password_hash(password))
user_fields: dict[str, Any] = {}
if "email" in data:
incoming_email = data.get("email")
if incoming_email is None:
user_fields["email"] = None
else:
user_fields["email"] = str(incoming_email).strip() or None
if "display_name" in data:
incoming_display_name = data.get("display_name")
user_fields["display_name"] = (
str(incoming_display_name).strip() or None
if incoming_display_name is not None
else None
)
email_changed = "email" in user_fields and user_fields["email"] != user.get("email")
display_name_changed = (
"display_name" in user_fields
and user_fields["display_name"] != user.get("display_name")
)
if email_changed and not capabilities["canEditEmail"]:
if auth_source == AUTH_SOURCE_CWA:
return jsonify(
{
"error": "Cannot change email for CWA users",
"message": "Email is synced from Calibre-Web.",
}
), 400
return jsonify(
{
"error": "Cannot change email for OIDC users",
"message": "Email is managed by your identity provider.",
}
), 400
if display_name_changed and not capabilities["canEditDisplayName"]:
return jsonify(
{
"error": "Cannot change display name for OIDC users",
"message": "Display name is managed by your identity provider.",
}
), 400
for field in ("email", "display_name"):
if field in user_fields and user_fields[field] == user.get(field):
user_fields.pop(field)
if user_fields:
user_db.update_user(user_id, **user_fields)
if "settings" in data:
settings_payload = data["settings"]
if not isinstance(settings_payload, dict):
return jsonify({"error": "Settings must be an object"}), 400
visible_self_settings_sections = _get_visible_self_settings_sections()
allowed_user_settings_keys = _get_allowed_self_settings_keys(visible_self_settings_sections)
disallowed_keys = sorted(
key for key in settings_payload if key not in allowed_user_settings_keys
)
if disallowed_keys:
return jsonify(
{
"error": "Some settings are admin-only",
"details": [
f"Setting not user-overridable: {key}" for key in disallowed_keys
],
}
), 400
validated_settings, validation_errors = validate_user_settings(settings_payload)
if validation_errors:
return jsonify(
{
"error": "Invalid settings payload",
"details": validation_errors,
}
), 400
user_db.set_user_settings(user_id, validated_settings)
try:
from shelfmark.core.config import config as app_config
app_config.refresh()
except Exception:
pass
updated = user_db.get_user(user_id=user_id)
if not updated:
return jsonify({"error": "User not found"}), 404
result = _serialize_self_user(updated, _get_auth_mode())
result["settings"] = user_db.get_user_settings(user_id)
logger.info(f"User {user_id} updated their own account")
return jsonify(result)
+231 -15
View File
@@ -22,12 +22,14 @@ class FieldBase:
required: bool = False # Whether field must have a value
env_var: Optional[str] = None # Override env var name (defaults to key)
env_supported: bool = True # Whether this setting can be set via ENV var (False = UI-only)
user_overridable: bool = False # Whether admins can set per-user overrides for this field
disabled: bool = False # Whether field is disabled/greyed out
disabled_reason: str = "" # Explanation shown when disabled
show_when: Optional[Dict[str, Any] | List[Dict[str, Any]]] = None # Conditional visibility: {"field": "key", "value": "expected"} or list of conditions
disabled_when: Optional[Dict[str, Any]] = None # Conditional disable: {"field": "key", "value": "expected", "reason": "..."}
requires_restart: bool = False # Whether changing this setting requires a container restart
universal_only: bool = False # Only show in Universal search mode (hide in Direct mode)
hidden_in_ui: bool = False # Keep field in schema/save path but hide default renderer
def get_env_var_name(self) -> str:
"""Get the environment variable name for this field."""
@@ -83,6 +85,14 @@ class MultiSelectField(FieldBase):
variant: str = "pills" # "pills" (default) or "dropdown" for checkbox dropdown style
@dataclass
class TagListField(FieldBase):
"""Editable list of free-form string values (tag/chip input)."""
placeholder: str = ""
default: List[str] = field(default_factory=list)
normalize_urls: bool = True
@dataclass
class OrderableListField(FieldBase):
# Options can be a list or a callable that returns a list (for lazy evaluation)
@@ -108,6 +118,31 @@ class TableField(FieldBase):
empty_message: str = ""
@dataclass
class CustomComponentField:
"""Render a custom frontend component inside settings content."""
key: str
component: str # Frontend component registry key
label: str = ""
description: str = ""
bind_keys: List[str] = field(default_factory=list) # Related value keys this component edits
value_fields: List[Any] = field(default_factory=list) # Backing value schema for this component
wrap_in_field_wrapper: bool = False # Whether to render with standard FieldWrapper layout
disabled: bool = False
disabled_reason: str = ""
show_when: Optional[Dict[str, Any] | List[Dict[str, Any]]] = None
universal_only: bool = False
def get_field_type(self) -> str:
return "CustomComponentField"
def get_bind_keys(self) -> List[str]:
if self.bind_keys:
return self.bind_keys
return [getattr(f, "key") for f in self.value_fields if getattr(f, "key", None)]
@dataclass
class ActionButton:
key: str # Action identifier
@@ -135,6 +170,7 @@ class HeadingField:
key: str # Unique identifier
title: str # Heading title
description: str = "" # Description text (supports markdown-style links)
description_by_auth_mode: Optional[Dict[str, str]] = None # Optional auth-mode specific description map
link_url: str = "" # Optional URL for a link
link_text: str = "" # Text for the link (defaults to URL if not provided)
show_when: Optional[Dict[str, Any] | List[Dict[str, Any]]] = None # Conditional visibility: {"field": "key", "value": "expected"} or list of conditions
@@ -145,7 +181,20 @@ class HeadingField:
# Type alias for all field types
SettingsField = Union[TextField, PasswordField, NumberField, CheckboxField, SelectField, MultiSelectField, OrderableListField, ActionButton, HeadingField]
SettingsField = Union[
TextField,
PasswordField,
NumberField,
CheckboxField,
SelectField,
MultiSelectField,
TagListField,
OrderableListField,
TableField,
CustomComponentField,
ActionButton,
HeadingField,
]
@dataclass
@@ -240,6 +289,48 @@ 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):
"""Yield value-bearing fields for a tab."""
for field in tab.fields:
if isinstance(field, CustomComponentField):
for value_field in field.value_fields:
if isinstance(value_field, (ActionButton, HeadingField, CustomComponentField)):
continue
yield value_field
continue
if isinstance(field, (ActionButton, HeadingField)):
continue
yield field
def get_settings_field_map(tab_name: Optional[str] = None) -> Dict[str, tuple[SettingsField, str]]:
"""Return key -> (field, tab_name) map for value-bearing settings fields."""
tabs: List[SettingsTab]
if tab_name:
tab = get_settings_tab(tab_name)
if not tab:
return {}
tabs = [tab]
else:
tabs = get_all_settings_tabs()
field_map: Dict[str, tuple[SettingsField, str]] = {}
for tab in tabs:
for field in _iter_value_fields(tab):
field_map[field.key] = (field, tab.name)
return field_map
def get_user_overridable_fields(tab_name: Optional[str] = None) -> Dict[str, tuple[SettingsField, str]]:
"""Return key -> (field, tab_name) map for fields marked user_overridable."""
field_map = get_settings_field_map(tab_name=tab_name)
return {
key: (field, tab)
for key, (field, tab) in field_map.items()
if getattr(field, "user_overridable", False)
}
def list_registered_settings() -> List[str]:
"""List all registered settings tab names."""
return list(_SETTINGS_REGISTRY.keys())
@@ -338,11 +429,7 @@ def initialize_default_configs() -> bool:
# Collect default values for all fields
defaults = {}
for field in tab.fields:
# Skip non-value fields
if isinstance(field, (ActionButton, HeadingField)):
continue
for field in _iter_value_fields(tab):
# Only include fields that have a non-None default
if field.default is not None:
defaults[field.key] = field.default
@@ -374,11 +461,7 @@ def sync_env_to_config() -> None:
for tab in get_all_settings_tabs():
values_to_sync = {}
for field in tab.fields:
# Skip non-value fields
if isinstance(field, (ActionButton, HeadingField)):
continue
for field in _iter_value_fields(tab):
# Skip fields that don't support ENV vars
if not getattr(field, 'env_supported', True):
continue
@@ -398,6 +481,61 @@ def sync_env_to_config() -> None:
logger.debug(f"Synced {len(values_to_sync)} ENV values to {tab.name} config: {list(values_to_sync.keys())}")
migrate_legacy_settings()
migrate_mirror_settings()
def migrate_mirror_settings() -> None:
"""
Migrate legacy AA mirror config into the new editable mirror list setting.
Legacy:
- AA_ADDITIONAL_URLS: comma-separated extra URLs appended to defaults
New:
- AA_MIRROR_URLS: full ordered list of available mirrors (used for Auto mode and for Settings options)
"""
mirrors_config = load_config_file("mirrors")
from shelfmark.core.mirrors import DEFAULT_AA_MIRRORS
from shelfmark.core.utils import normalize_http_url
raw_list = mirrors_config.get("AA_MIRROR_URLS")
raw_additional = mirrors_config.get("AA_ADDITIONAL_URLS", "")
def _normalize_list(values: list[str]) -> list[str]:
out: list[str] = []
for item in values:
if str(item).strip().lower() == "auto":
continue
norm = normalize_http_url(str(item), default_scheme="https")
if norm and norm not in out:
out.append(norm)
return out
# If already a proper list, just ensure it's non-empty.
if isinstance(raw_list, list):
normalized = _normalize_list([str(v) for v in raw_list])
if normalized:
return
save_config_file("mirrors", {"AA_MIRROR_URLS": _normalize_list(DEFAULT_AA_MIRRORS)})
return
# If saved as a string, convert to list.
if isinstance(raw_list, str) and raw_list.strip():
parts = [p.strip() for p in raw_list.split(",") if p.strip()]
normalized = _normalize_list(parts)
if normalized:
save_config_file("mirrors", {"AA_MIRROR_URLS": normalized})
return
save_config_file("mirrors", {"AA_MIRROR_URLS": _normalize_list(DEFAULT_AA_MIRRORS)})
return
# If there's legacy additional mirrors, seed the full list so the UI reflects reality.
if isinstance(raw_additional, str) and raw_additional.strip():
additional_parts = [p.strip() for p in raw_additional.split(",") if p.strip()]
combined = _normalize_list(DEFAULT_AA_MIRRORS + additional_parts)
if combined:
save_config_file("mirrors", {"AA_MIRROR_URLS": combined})
def migrate_legacy_settings() -> None:
@@ -510,7 +648,7 @@ def migrate_legacy_settings() -> None:
def get_setting_value(field: SettingsField, tab_name: str) -> Any:
if isinstance(field, (ActionButton, HeadingField)):
if isinstance(field, (ActionButton, HeadingField, CustomComponentField)):
return None # Actions and headings don't have values
# 1. Check environment variable (if supported for this field)
@@ -542,6 +680,8 @@ def _parse_env_value(value: str, field: SettingsField) -> Any:
return field.default
elif isinstance(field, MultiSelectField):
return [v.strip() for v in value.split(',') if v.strip()]
elif isinstance(field, TagListField):
return [v.strip() for v in value.split(',') if v.strip()]
elif isinstance(field, OrderableListField):
# Parse JSON array: [{"id": "...", "enabled": true}, ...]
try:
@@ -563,7 +703,7 @@ def _parse_env_value(value: str, field: SettingsField) -> Any:
def is_value_from_env(field: SettingsField) -> bool:
"""Check if a field's value comes from an environment variable."""
if isinstance(field, (ActionButton, HeadingField)):
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):
@@ -583,6 +723,36 @@ def serialize_field(field: SettingsField, tab_name: str, include_value: bool = T
Returns:
Dict representation of the field.
"""
# CustomComponentField has a custom structure - handle separately
if isinstance(field, CustomComponentField):
result: Dict[str, Any] = {
"key": field.key,
"label": field.label,
"type": field.get_field_type(),
"description": field.description,
"component": field.component,
"bindKeys": field.get_bind_keys(),
"wrapInFieldWrapper": field.wrap_in_field_wrapper,
"disabled": field.disabled,
"disabledReason": field.disabled_reason,
}
if field.value_fields:
bound_fields = []
for value_field in field.value_fields:
serialized_bound_field = serialize_field(
value_field,
tab_name,
include_value=include_value,
)
serialized_bound_field["hiddenInUi"] = True
bound_fields.append(serialized_bound_field)
result["boundFields"] = bound_fields
if field.show_when:
result["showWhen"] = field.show_when
if field.universal_only:
result["universalOnly"] = True
return result
# HeadingField has a different structure - handle separately
if isinstance(field, HeadingField):
result: Dict[str, Any] = {
@@ -591,6 +761,8 @@ def serialize_field(field: SettingsField, tab_name: str, include_value: bool = T
"title": field.title,
"description": field.description,
}
if field.description_by_auth_mode:
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
@@ -609,6 +781,8 @@ def serialize_field(field: SettingsField, tab_name: str, include_value: bool = T
"disabled": getattr(field, 'disabled', False),
"disabledReason": getattr(field, 'disabled_reason', ''),
"requiresRestart": getattr(field, 'requires_restart', False),
"userOverridable": getattr(field, 'user_overridable', False),
"hiddenInUi": getattr(field, 'hidden_in_ui', False),
}
# Add optional properties if set
@@ -643,6 +817,9 @@ def serialize_field(field: SettingsField, tab_name: str, include_value: bool = T
options = field.options() if callable(field.options) else field.options
result["options"] = options
result["variant"] = field.variant
elif isinstance(field, TagListField):
result["placeholder"] = field.placeholder
result["normalizeUrls"] = field.normalize_urls
elif isinstance(field, OrderableListField):
# Support callable options for lazy evaluation (avoids circular imports)
options = field.options() if callable(field.options) else field.options
@@ -656,7 +833,7 @@ def serialize_field(field: SettingsField, tab_name: str, include_value: bool = T
result["style"] = field.style
result["description"] = field.description
if include_value and not isinstance(field, (ActionButton, HeadingField)):
if include_value and not isinstance(field, (ActionButton, HeadingField, CustomComponentField)):
value = get_setting_value(field, tab_name)
# Ensure select values are serialized as strings so the frontend can
@@ -674,6 +851,16 @@ def serialize_field(field: SettingsField, tab_name: str, include_value: bool = T
value = [v.strip() for v in value.split(",") if v.strip()]
else:
value = []
elif isinstance(field, TagListField):
if value is None:
value = []
elif isinstance(value, list):
value = [str(v) for v in value]
elif isinstance(value, str):
# Support legacy/manual configs where lists were saved as comma-separated strings.
value = [v.strip() for v in value.split(",") if v.strip()]
else:
value = []
elif isinstance(field, TableField):
if value is None:
value = []
@@ -801,6 +988,23 @@ def _apply_dns_settings(config) -> None:
except Exception as e:
logger.warning(f"Failed to apply DNS settings: {e}")
def _apply_aa_mirror_settings(config) -> None:
"""
Apply AA mirror settings changes to the network module.
This ensures AA_BASE_URL / AA_ADDITIONAL_URLS changes take effect immediately
without requiring a container restart.
"""
try:
from shelfmark.download import network
# Reload AA mirror list and configured base URL from refreshed config.
network.init_aa(force=True)
except ImportError:
pass # Network module not available
except Exception as e:
logger.warning(f"Failed to apply AA mirror settings: {e}")
def update_settings(tab_name: str, values: Dict[str, Any]) -> Dict[str, Any]:
tab = get_settings_tab(tab_name)
@@ -808,7 +1012,10 @@ def update_settings(tab_name: str, values: Dict[str, Any]) -> Dict[str, Any]:
return {"success": False, "message": f"Unknown settings tab: {tab_name}", "updated": [], "requiresRestart": False}
# Build a map of field keys to fields (exclude non-value fields)
field_map = {f.key: f for f in tab.fields if not isinstance(f, (ActionButton, HeadingField))}
field_map = {
key: field
for key, (field, _) in get_settings_field_map(tab_name=tab_name).items()
}
# Filter out values that are set via env vars or unknown
values_to_save = {}
@@ -885,6 +1092,15 @@ def update_settings(tab_name: str, values: Dict[str, Any]) -> Dict[str, Any]:
):
_apply_dns_settings(config_obj)
# Apply AA mirror settings changes live (mirrors tab)
aa_keys = {"AA_BASE_URL", "AA_MIRROR_URLS", "AA_ADDITIONAL_URLS"}
if (
config_obj is not None
and tab_name == "mirrors"
and aa_keys.intersection(values_to_save.keys())
):
_apply_aa_mirror_settings(config_obj)
# Sync metadata provider selection when a provider's enabled state changes
tab = get_settings_tab(tab_name)
if tab and tab.group == "metadata_providers":
+754
View File
@@ -0,0 +1,754 @@
"""SQLite user database for multi-user support."""
import json
import os
import sqlite3
import threading
from typing import Any, Dict, List, Optional
from shelfmark.core.auth_modes import AUTH_SOURCE_BUILTIN, AUTH_SOURCE_SET
from shelfmark.core.logger import setup_logger
from shelfmark.core.requests_service import (
normalize_delivery_state,
normalize_policy_mode,
normalize_request_level,
normalize_request_status,
validate_request_level_payload,
validate_status_transition,
)
logger = setup_logger(__name__)
_CREATE_TABLES_SQL = """
CREATE TABLE IF NOT EXISTS users (
id INTEGER PRIMARY KEY AUTOINCREMENT,
username TEXT UNIQUE NOT NULL,
email TEXT,
display_name TEXT,
password_hash TEXT,
oidc_subject TEXT UNIQUE,
auth_source TEXT NOT NULL DEFAULT 'builtin',
role TEXT NOT NULL DEFAULT 'user',
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
CREATE TABLE IF NOT EXISTS user_settings (
user_id INTEGER PRIMARY KEY REFERENCES users(id) ON DELETE CASCADE,
settings_json TEXT NOT NULL DEFAULT '{}'
);
CREATE TABLE IF NOT EXISTS download_requests (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
status TEXT NOT NULL DEFAULT 'pending',
delivery_state TEXT NOT NULL DEFAULT 'none',
source_hint TEXT,
content_type TEXT NOT NULL,
request_level TEXT NOT NULL,
policy_mode TEXT NOT NULL,
book_data TEXT NOT NULL,
release_data TEXT,
note TEXT,
admin_note TEXT,
reviewed_by INTEGER REFERENCES users(id),
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
reviewed_at TIMESTAMP,
delivery_updated_at TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_download_requests_user_status_created_at
ON download_requests (user_id, status, created_at DESC);
CREATE INDEX IF NOT EXISTS idx_download_requests_status_created_at
ON download_requests (status, created_at DESC);
CREATE TABLE IF NOT EXISTS activity_log (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER REFERENCES users(id) ON DELETE CASCADE,
item_type TEXT NOT NULL,
item_key TEXT NOT NULL,
request_id INTEGER,
source_id TEXT,
origin TEXT NOT NULL,
final_status TEXT NOT NULL,
snapshot_json TEXT NOT NULL,
terminal_at TIMESTAMP NOT NULL,
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_activity_log_user_terminal
ON activity_log (user_id, terminal_at DESC);
CREATE INDEX IF NOT EXISTS idx_activity_log_lookup
ON activity_log (user_id, item_type, item_key, id DESC);
CREATE TABLE IF NOT EXISTS activity_dismissals (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
item_type TEXT NOT NULL,
item_key TEXT NOT NULL,
activity_log_id INTEGER REFERENCES activity_log(id) ON DELETE SET NULL,
dismissed_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
UNIQUE(user_id, item_type, item_key)
);
CREATE INDEX IF NOT EXISTS idx_activity_dismissals_user_dismissed_at
ON activity_dismissals (user_id, dismissed_at DESC);
"""
def get_users_db_path(config_dir: Optional[str] = None) -> str:
"""Return the configured users database path."""
root = config_dir or os.environ.get("CONFIG_DIR", "/config")
return os.path.join(root, "users.db")
def sync_builtin_admin_user(
username: str,
password_hash: str,
db_path: Optional[str] = None,
) -> None:
"""Ensure a local admin user exists for configured builtin credentials."""
normalized_username = (username or "").strip()
normalized_hash = password_hash or ""
if not normalized_username or not normalized_hash:
return
user_db = UserDB(db_path or get_users_db_path())
user_db.initialize()
existing = user_db.get_user(username=normalized_username)
if existing:
updates: dict[str, Any] = {}
if existing.get("password_hash") != normalized_hash:
updates["password_hash"] = normalized_hash
if existing.get("role") != "admin":
updates["role"] = "admin"
if existing.get("auth_source") != AUTH_SOURCE_BUILTIN:
updates["auth_source"] = AUTH_SOURCE_BUILTIN
if updates:
user_db.update_user(existing["id"], **updates)
logger.info(f"Updated local admin user '{normalized_username}' from builtin settings")
return
user_db.create_user(
username=normalized_username,
password_hash=normalized_hash,
auth_source=AUTH_SOURCE_BUILTIN,
role="admin",
)
logger.info(f"Created local admin user '{normalized_username}' from builtin settings")
class UserDB:
"""Thread-safe SQLite user database."""
_VALID_AUTH_SOURCES = set(AUTH_SOURCE_SET)
def __init__(self, db_path: str):
self._db_path = db_path
self._lock = threading.Lock()
def _connect(self) -> sqlite3.Connection:
conn = sqlite3.connect(self._db_path)
conn.row_factory = sqlite3.Row
conn.execute("PRAGMA foreign_keys = ON")
return conn
def initialize(self) -> None:
"""Create database and tables if they don't exist."""
with self._lock:
conn = self._connect()
try:
conn.executescript(_CREATE_TABLES_SQL)
self._migrate_auth_source_column(conn)
self._migrate_request_delivery_columns(conn)
self._migrate_activity_tables(conn)
conn.commit()
# WAL mode must be changed outside an open transaction.
conn.execute("PRAGMA journal_mode=WAL")
finally:
conn.close()
def _migrate_auth_source_column(self, conn: sqlite3.Connection) -> None:
"""Ensure users.auth_source exists and backfill historical rows."""
columns = conn.execute("PRAGMA table_info(users)").fetchall()
column_names = {str(col["name"]) for col in columns}
if "auth_source" not in column_names:
conn.execute(
"ALTER TABLE users ADD COLUMN auth_source TEXT NOT NULL DEFAULT 'builtin'"
)
# Backfill OIDC-origin users created before auth_source existed.
conn.execute(
"UPDATE users SET auth_source = 'oidc' WHERE oidc_subject IS NOT NULL"
)
# Defensive cleanup for any legacy null/blank values.
conn.execute(
"UPDATE users SET auth_source = 'builtin' WHERE auth_source IS NULL OR auth_source = ''"
)
def _migrate_request_delivery_columns(self, conn: sqlite3.Connection) -> None:
"""Ensure request delivery-state columns exist and backfill historical rows."""
columns = conn.execute("PRAGMA table_info(download_requests)").fetchall()
column_names = {str(col["name"]) for col in columns}
if "delivery_state" not in column_names:
conn.execute(
"ALTER TABLE download_requests ADD COLUMN delivery_state TEXT NOT NULL DEFAULT 'none'"
)
if "delivery_updated_at" not in column_names:
conn.execute("ALTER TABLE download_requests ADD COLUMN delivery_updated_at TIMESTAMP")
if "last_failure_reason" not in column_names:
conn.execute("ALTER TABLE download_requests ADD COLUMN last_failure_reason TEXT")
conn.execute(
"""
UPDATE download_requests
SET delivery_state = 'unknown'
WHERE status = 'fulfilled' AND (delivery_state IS NULL OR TRIM(delivery_state) = '' OR delivery_state = 'none')
"""
)
conn.execute(
"""
UPDATE download_requests
SET delivery_state = 'none'
WHERE status != 'fulfilled' AND (delivery_state IS NULL OR TRIM(delivery_state) = '')
"""
)
conn.execute(
"""
UPDATE download_requests
SET delivery_updated_at = COALESCE(delivery_updated_at, reviewed_at, created_at)
WHERE delivery_state != 'none' AND delivery_updated_at IS NULL
"""
)
conn.execute(
"""
UPDATE download_requests
SET delivery_state = 'complete'
WHERE delivery_state = 'cleared'
"""
)
def _migrate_activity_tables(self, conn: sqlite3.Connection) -> None:
"""Ensure activity log and dismissal tables exist with current columns/indexes."""
conn.executescript(
"""
CREATE TABLE IF NOT EXISTS activity_log (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER REFERENCES users(id) ON DELETE CASCADE,
item_type TEXT NOT NULL,
item_key TEXT NOT NULL,
request_id INTEGER,
source_id TEXT,
origin TEXT NOT NULL,
final_status TEXT NOT NULL,
snapshot_json TEXT NOT NULL,
terminal_at TIMESTAMP NOT NULL,
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_activity_log_user_terminal
ON activity_log (user_id, terminal_at DESC);
CREATE INDEX IF NOT EXISTS idx_activity_log_lookup
ON activity_log (user_id, item_type, item_key, id DESC);
CREATE TABLE IF NOT EXISTS activity_dismissals (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
item_type TEXT NOT NULL,
item_key TEXT NOT NULL,
activity_log_id INTEGER REFERENCES activity_log(id) ON DELETE SET NULL,
dismissed_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
UNIQUE(user_id, item_type, item_key)
);
CREATE INDEX IF NOT EXISTS idx_activity_dismissals_user_dismissed_at
ON activity_dismissals (user_id, dismissed_at DESC);
"""
)
dismissal_columns = conn.execute("PRAGMA table_info(activity_dismissals)").fetchall()
dismissal_column_names = {str(col["name"]) for col in dismissal_columns}
if "activity_log_id" not in dismissal_column_names:
conn.execute("ALTER TABLE activity_dismissals ADD COLUMN activity_log_id INTEGER")
def create_user(
self,
username: str,
email: Optional[str] = None,
display_name: Optional[str] = None,
password_hash: Optional[str] = None,
oidc_subject: Optional[str] = None,
auth_source: str = "builtin",
role: str = "user",
) -> Dict[str, Any]:
"""Create a new user. Raises ValueError if username or oidc_subject already exists."""
if auth_source not in self._VALID_AUTH_SOURCES:
raise ValueError(f"Invalid auth_source: {auth_source}")
with self._lock:
conn = self._connect()
try:
cursor = conn.execute(
"""INSERT INTO users (
username, email, display_name, password_hash, oidc_subject, auth_source, role
)
VALUES (?, ?, ?, ?, ?, ?, ?)""",
(
username,
email,
display_name,
password_hash,
oidc_subject,
auth_source,
role,
),
)
conn.commit()
user_id = cursor.lastrowid
return self._get_user_by_id(conn, user_id)
except sqlite3.IntegrityError as e:
raise ValueError(f"User already exists: {e}")
finally:
conn.close()
def get_user(
self,
user_id: Optional[int] = None,
username: Optional[str] = None,
oidc_subject: Optional[str] = None,
) -> Optional[Dict[str, Any]]:
"""Get a user by id, username, or oidc_subject. Returns None if not found."""
conn = self._connect()
try:
if user_id is not None:
return self._get_user_by_id(conn, user_id)
elif username is not None:
row = conn.execute(
"SELECT * FROM users WHERE username = ?", (username,)
).fetchone()
elif oidc_subject is not None:
row = conn.execute(
"SELECT * FROM users WHERE oidc_subject = ?", (oidc_subject,)
).fetchone()
else:
return None
return dict(row) if row else None
finally:
conn.close()
def _get_user_by_id(self, conn: sqlite3.Connection, user_id: int) -> Optional[Dict[str, Any]]:
row = conn.execute("SELECT * FROM users WHERE id = ?", (user_id,)).fetchone()
return dict(row) if row else None
_ALLOWED_UPDATE_COLUMNS = {
"email",
"display_name",
"password_hash",
"oidc_subject",
"auth_source",
"role",
}
def update_user(self, user_id: int, **kwargs) -> None:
"""Update user fields. Raises ValueError if user not found or invalid column."""
if not kwargs:
return
for k in kwargs:
if k not in self._ALLOWED_UPDATE_COLUMNS:
raise ValueError(f"Invalid column: {k}")
if "auth_source" in kwargs and kwargs["auth_source"] not in self._VALID_AUTH_SOURCES:
raise ValueError(f"Invalid auth_source: {kwargs['auth_source']}")
with self._lock:
conn = self._connect()
try:
# Verify user exists
if not self._get_user_by_id(conn, user_id):
raise ValueError(f"User {user_id} not found")
sets = ", ".join(f"{k} = ?" for k in kwargs)
values = list(kwargs.values()) + [user_id]
conn.execute(f"UPDATE users SET {sets} WHERE id = ?", values)
conn.commit()
finally:
conn.close()
def delete_user(self, user_id: int) -> None:
"""Delete a user and their settings."""
with self._lock:
conn = self._connect()
try:
conn.execute("DELETE FROM users WHERE id = ?", (user_id,))
conn.commit()
finally:
conn.close()
def list_users(self) -> List[Dict[str, Any]]:
"""List all users."""
conn = self._connect()
try:
rows = conn.execute("SELECT * FROM users ORDER BY id").fetchall()
return [dict(r) for r in rows]
finally:
conn.close()
def get_user_settings(self, user_id: int) -> Dict[str, Any]:
"""Get per-user settings. Returns empty dict if none set."""
conn = self._connect()
try:
row = conn.execute(
"SELECT settings_json FROM user_settings WHERE user_id = ?", (user_id,)
).fetchone()
if row:
return json.loads(row["settings_json"])
return {}
finally:
conn.close()
def set_user_settings(self, user_id: int, settings: Dict[str, Any]) -> None:
"""Merge settings into user's existing settings."""
with self._lock:
conn = self._connect()
try:
existing = {}
row = conn.execute(
"SELECT settings_json FROM user_settings WHERE user_id = ?", (user_id,)
).fetchone()
if row:
existing = json.loads(row["settings_json"])
existing.update(settings)
# Remove keys set to None (meaning "clear this override")
existing = {k: v for k, v in existing.items() if v is not None}
settings_json = json.dumps(existing)
conn.execute(
"""INSERT INTO user_settings (user_id, settings_json) VALUES (?, ?)
ON CONFLICT(user_id) DO UPDATE SET settings_json = ?""",
(user_id, settings_json, settings_json),
)
conn.commit()
finally:
conn.close()
@staticmethod
def _serialize_json(value: Any, field: str) -> Optional[str]:
if value is None:
return None
try:
return json.dumps(value)
except TypeError as exc:
raise ValueError(f"{field} must be JSON-serializable") from exc
@staticmethod
def _parse_request_row(row: Optional[sqlite3.Row]) -> Optional[Dict[str, Any]]:
if row is None:
return None
payload = dict(row)
for key in ("book_data", "release_data"):
raw_value = payload.get(key)
if raw_value is None:
payload[key] = None
continue
try:
payload[key] = json.loads(raw_value)
except (ValueError, TypeError):
payload[key] = None
return payload
def create_request(
self,
*,
user_id: int,
content_type: str,
request_level: str,
policy_mode: str,
book_data: Dict[str, Any],
release_data: Optional[Dict[str, Any]] = None,
status: str = "pending",
source_hint: Optional[str] = None,
note: Optional[str] = None,
admin_note: Optional[str] = None,
reviewed_by: Optional[int] = None,
reviewed_at: Optional[str] = None,
delivery_state: str = "none",
delivery_updated_at: Optional[str] = None,
) -> Dict[str, Any]:
"""Create a download request row and return the created record."""
if not isinstance(book_data, dict):
raise ValueError("book_data must be an object")
if release_data is not None and not isinstance(release_data, dict):
raise ValueError("release_data must be an object when provided")
if not content_type:
raise ValueError("content_type is required")
normalized_status = normalize_request_status(status)
normalized_delivery_state = normalize_delivery_state(delivery_state)
normalized_policy_mode = normalize_policy_mode(policy_mode)
normalized_request_level = validate_request_level_payload(request_level, release_data)
with self._lock:
conn = self._connect()
try:
cursor = conn.execute(
"""
INSERT INTO download_requests (
user_id,
status,
delivery_state,
source_hint,
content_type,
request_level,
policy_mode,
book_data,
release_data,
note,
admin_note,
reviewed_by,
reviewed_at,
delivery_updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
""",
(
user_id,
normalized_status,
normalized_delivery_state,
source_hint,
content_type,
normalized_request_level,
normalized_policy_mode,
self._serialize_json(book_data, "book_data"),
self._serialize_json(release_data, "release_data"),
note,
admin_note,
reviewed_by,
reviewed_at,
delivery_updated_at,
),
)
conn.commit()
request_id = cursor.lastrowid
row = conn.execute(
"SELECT * FROM download_requests WHERE id = ?",
(request_id,),
).fetchone()
parsed = self._parse_request_row(row)
if parsed is None:
raise ValueError(f"Request {request_id} not found after creation")
return parsed
finally:
conn.close()
def get_request(self, request_id: int) -> Optional[Dict[str, Any]]:
"""Get a request row by ID."""
conn = self._connect()
try:
row = conn.execute(
"SELECT * FROM download_requests WHERE id = ?",
(request_id,),
).fetchone()
return self._parse_request_row(row)
finally:
conn.close()
def list_requests(
self,
*,
user_id: Optional[int] = None,
status: Optional[str] = None,
limit: Optional[int] = None,
offset: int = 0,
) -> List[Dict[str, Any]]:
"""List requests with optional user/status filters."""
where_clauses: List[str] = []
params: List[Any] = []
if user_id is not None:
where_clauses.append("user_id = ?")
params.append(user_id)
if status is not None:
where_clauses.append("status = ?")
params.append(normalize_request_status(status))
query = "SELECT * FROM download_requests"
if where_clauses:
query += " WHERE " + " AND ".join(where_clauses)
query += " ORDER BY created_at DESC, id DESC"
if limit is not None:
query += " LIMIT ?"
params.append(int(limit))
if offset:
query += " OFFSET ?"
params.append(offset)
elif offset:
query += " LIMIT -1 OFFSET ?"
params.append(offset)
conn = self._connect()
try:
rows = conn.execute(query, params).fetchall()
results: List[Dict[str, Any]] = []
for row in rows:
parsed = self._parse_request_row(row)
if parsed is not None:
results.append(parsed)
return results
finally:
conn.close()
_ALLOWED_REQUEST_UPDATE_COLUMNS = {
"status",
"source_hint",
"content_type",
"request_level",
"policy_mode",
"book_data",
"release_data",
"note",
"admin_note",
"reviewed_by",
"reviewed_at",
"delivery_state",
"delivery_updated_at",
"last_failure_reason",
}
def update_request(
self,
request_id: int,
expected_current_status: Optional[str] = None,
**kwargs,
) -> Dict[str, Any]:
"""Update request fields and return the updated record."""
if not kwargs:
request = self.get_request(request_id)
if request is None:
raise ValueError(f"Request {request_id} not found")
if expected_current_status is not None:
normalized_expected_status = normalize_request_status(expected_current_status)
if request["status"] != normalized_expected_status:
raise ValueError("Request state changed before update")
return request
for key in kwargs:
if key not in self._ALLOWED_REQUEST_UPDATE_COLUMNS:
raise ValueError(f"Invalid request column: {key}")
with self._lock:
conn = self._connect()
try:
row = conn.execute(
"SELECT * FROM download_requests WHERE id = ?",
(request_id,),
).fetchone()
current = self._parse_request_row(row)
if current is None:
raise ValueError(f"Request {request_id} not found")
if expected_current_status is not None:
normalized_expected_status = normalize_request_status(expected_current_status)
if current["status"] != normalized_expected_status:
raise ValueError("Request state changed before update")
updates = dict(kwargs)
if "status" in updates:
_, normalized_status = validate_status_transition(
current["status"],
updates["status"],
)
updates["status"] = normalized_status
if "policy_mode" in updates:
updates["policy_mode"] = normalize_policy_mode(updates["policy_mode"])
if "delivery_state" in updates:
updates["delivery_state"] = normalize_delivery_state(updates["delivery_state"])
if "delivery_updated_at" in updates:
delivery_updated_at = updates["delivery_updated_at"]
if delivery_updated_at is not None and not isinstance(delivery_updated_at, str):
raise ValueError("delivery_updated_at must be a string when provided")
if "content_type" in updates and not updates["content_type"]:
raise ValueError("content_type is required")
candidate_request_level = updates.get("request_level", current["request_level"])
candidate_release_data = (
updates["release_data"] if "release_data" in updates else current["release_data"]
)
candidate_status = updates.get("status", current["status"])
normalized_request_level = normalize_request_level(candidate_request_level)
normalized_candidate_status = normalize_request_status(candidate_status)
if normalized_request_level == "release" and candidate_release_data is None:
raise ValueError("request_level=release requires non-null release_data")
if (
normalized_request_level == "book"
and candidate_release_data is not None
and normalized_candidate_status != "fulfilled"
):
raise ValueError("request_level=book requires null release_data")
if "request_level" in updates:
updates["request_level"] = normalized_request_level
if "book_data" in updates:
if not isinstance(updates["book_data"], dict):
raise ValueError("book_data must be an object")
updates["book_data"] = self._serialize_json(updates["book_data"], "book_data")
if "release_data" in updates:
if updates["release_data"] is not None and not isinstance(updates["release_data"], dict):
raise ValueError("release_data must be an object when provided")
updates["release_data"] = self._serialize_json(
updates["release_data"],
"release_data",
)
set_clause = ", ".join(f"{column} = ?" for column in updates)
values = list(updates.values()) + [request_id]
conn.execute(
f"UPDATE download_requests SET {set_clause} WHERE id = ?",
values,
)
conn.commit()
updated_row = conn.execute(
"SELECT * FROM download_requests WHERE id = ?",
(request_id,),
).fetchone()
parsed = self._parse_request_row(updated_row)
if parsed is None:
raise ValueError(f"Request {request_id} not found after update")
return parsed
finally:
conn.close()
def count_pending_requests(self) -> int:
"""Count all pending requests."""
conn = self._connect()
try:
row = conn.execute(
"SELECT COUNT(*) AS count FROM download_requests WHERE status = 'pending'"
).fetchone()
return int(row["count"]) if row else 0
finally:
conn.close()
def count_user_pending_requests(self, user_id: int) -> int:
"""Count pending requests for a specific user."""
conn = self._connect()
try:
row = conn.execute(
"SELECT COUNT(*) AS count FROM download_requests WHERE user_id = ? AND status = 'pending'",
(user_id,),
).fetchone()
return int(row["count"]) if row else 0
finally:
conn.close()
+78
View File
@@ -0,0 +1,78 @@
"""Shared helpers for user-overridable settings metadata and payloads."""
from typing import Any
from shelfmark.core.settings_registry import load_config_file
from shelfmark.core.user_db import UserDB
def get_settings_registry():
# Ensure settings modules are loaded before reading registry metadata.
import shelfmark.config.settings # noqa: F401
import shelfmark.config.security # noqa: F401
import shelfmark.config.notifications_settings # noqa: F401
import shelfmark.config.users_settings # noqa: F401
from shelfmark.core import settings_registry
return settings_registry
def get_ordered_user_overridable_fields(tab_name: str) -> list[tuple[str, Any]]:
settings_registry = get_settings_registry()
tab = settings_registry.get_settings_tab(tab_name)
if not tab:
return []
overridable_map = settings_registry.get_user_overridable_fields(tab_name=tab_name)
return [(field.key, field) for field in tab.fields if field.key in overridable_map]
def build_user_preferences_payload(user_db: UserDB, user_id: int, tab_name: str) -> dict[str, Any]:
from shelfmark.core.config import config as app_config
settings_registry = get_settings_registry()
ordered_fields = get_ordered_user_overridable_fields(tab_name)
if not ordered_fields:
tab_label = tab_name.capitalize()
raise ValueError(f"{tab_label} settings tab not found")
tab_config = load_config_file(tab_name)
user_settings = user_db.get_user_settings(user_id)
ordered_keys = [key for key, _ in ordered_fields]
fields_payload: list[dict[str, Any]] = []
global_values: dict[str, Any] = {}
effective: dict[str, dict[str, Any]] = {}
for key, field in ordered_fields:
serialized = settings_registry.serialize_field(field, tab_name, include_value=False)
serialized["fromEnv"] = bool(field.env_supported and settings_registry.is_value_from_env(field))
fields_payload.append(serialized)
global_values[key] = app_config.get(key, field.default)
source = "default"
value = app_config.get(key, field.default, user_id=user_id)
if field.env_supported and settings_registry.is_value_from_env(field):
source = "env_var"
elif key in user_settings and user_settings[key] is not None:
source = "user_override"
value = user_settings[key]
elif key in tab_config:
source = "global_config"
effective[key] = {"value": value, "source": source}
user_overrides = {
key: user_settings[key]
for key in ordered_keys
if key in user_settings and user_settings[key] is not None
}
return {
"tab": tab_name,
"keys": ordered_keys,
"fields": fields_payload,
"globalValues": global_values,
"userOverrides": user_overrides,
"effective": effective,
}
+79 -5
View File
@@ -1,6 +1,8 @@
"""Shared utility functions for the Shelfmark."""
import base64
import os
import re
from pathlib import Path
from typing import Optional
from urllib.parse import urlparse
@@ -118,21 +120,88 @@ _LEGACY_CONTENT_TYPE_TO_CONFIG_KEY = {
"other": "INGEST_DIR_OTHER",
}
_USER_PLACEHOLDER_PATTERN = re.compile(r"\{user\}", re.IGNORECASE)
_INVALID_USER_PATH_CHARS = re.compile(r'[\\/:*?"<>|]')
def get_destination(is_audiobook: bool = False) -> Path:
def _sanitize_user_for_path(username: str) -> str:
"""Sanitize username for path usage in destination placeholders."""
sanitized = _INVALID_USER_PATH_CHARS.sub("_", username.strip())
return sanitized.strip(" .")
def _resolve_destination_username(
user_id: Optional[int] = None,
username: Optional[str] = None,
) -> str:
explicit = str(username or "").strip()
if explicit:
return explicit
if user_id is None:
return ""
try:
from shelfmark.core.user_db import UserDB
user_db = UserDB(os.path.join(os.environ.get("CONFIG_DIR", "/config"), "users.db"))
user_db.initialize()
user = user_db.get_user(user_id=user_id)
if not user:
return ""
return str(user.get("username") or "").strip()
except Exception:
return ""
def _expand_user_destination_placeholder(
path_value: str,
user_id: Optional[int] = None,
username: Optional[str] = None,
) -> str:
"""Expand `{User}` placeholders in destination paths."""
if not isinstance(path_value, str):
return path_value
if not _USER_PLACEHOLDER_PATTERN.search(path_value):
return path_value
resolved_username = _sanitize_user_for_path(
_resolve_destination_username(user_id=user_id, username=username)
)
return _USER_PLACEHOLDER_PATTERN.sub(resolved_username, path_value)
def get_destination(
is_audiobook: bool = False,
user_id: Optional[int] = None,
username: Optional[str] = None,
) -> Path:
"""Get base destination directory. Audiobooks fall back to main destination."""
from shelfmark.core.config import config
if is_audiobook:
# Audiobook destination with fallback to main destination
audiobook_dest = config.get("DESTINATION_AUDIOBOOK", "")
audiobook_dest = config.get("DESTINATION_AUDIOBOOK", "", user_id=user_id)
if audiobook_dest:
return Path(audiobook_dest)
return Path(
_expand_user_destination_placeholder(
str(audiobook_dest),
user_id=user_id,
username=username,
)
)
# Main destination (also fallback for audiobooks)
# Check new setting first, then legacy INGEST_DIR
destination = config.get("DESTINATION", "") or config.get("INGEST_DIR", "/books")
return Path(destination)
destination = config.get("DESTINATION", "", user_id=user_id) or config.get("INGEST_DIR", "/books")
return Path(
_expand_user_destination_placeholder(
str(destination),
user_id=user_id,
username=username,
)
)
def get_aa_content_type_dir(content_type: Optional[str] = None) -> Optional[Path]:
@@ -191,6 +260,11 @@ def transform_cover_url(cover_url: Optional[str], cache_id: str) -> Optional[str
if not is_covers_cache_enabled():
return cover_url
from shelfmark.core.config import config as app_config
# 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", ""))
if base_path:
return f"{base_path}/api/covers/{cache_id}?url={encoded_url}"
return f"/api/covers/{cache_id}?url={encoded_url}"
@@ -1,5 +1,5 @@
"""
Download client infrastructure for Prowlarr integration.
Shared download client infrastructure for external release sources.
This module provides:
- DownloadState: Enum of valid download states
@@ -423,9 +423,9 @@ def get_all_clients() -> Dict[str, List[Type[DownloadClient]]]:
# Import client implementations to trigger registration
# These imports are at the bottom to avoid circular imports
from shelfmark.release_sources.prowlarr.clients import qbittorrent # noqa: F401, E402
from shelfmark.release_sources.prowlarr.clients import nzbget # noqa: F401, E402
from shelfmark.release_sources.prowlarr.clients import sabnzbd # noqa: F401, E402
from shelfmark.release_sources.prowlarr.clients import transmission # noqa: F401, E402
from shelfmark.release_sources.prowlarr.clients import deluge # noqa: F401, E402
from shelfmark.release_sources.prowlarr.clients import rtorrent # noqa: F401, E402
from shelfmark.download.clients import qbittorrent # noqa: F401, E402
from shelfmark.download.clients import nzbget # noqa: F401, E402
from shelfmark.download.clients import sabnzbd # noqa: F401, E402
from shelfmark.download.clients import transmission # noqa: F401, E402
from shelfmark.download.clients import deluge # noqa: F401, E402
from shelfmark.download.clients import rtorrent # noqa: F401, E402
+753
View File
@@ -0,0 +1,753 @@
"""Shared download handler for external torrent/usenet clients."""
import shutil
import time
from abc import ABC, abstractmethod
from dataclasses import dataclass
from pathlib import Path
from threading import Event
from typing import Callable, Optional
from shelfmark.core.config import config
from shelfmark.core.logger import setup_logger
from shelfmark.core.models import DownloadTask
from shelfmark.core.utils import is_audiobook
from shelfmark.download.clients import (
DownloadClient,
DownloadState,
get_client,
list_configured_clients,
)
from shelfmark.download.fs import run_blocking_io
from shelfmark.release_sources import DownloadHandler
logger = setup_logger(__name__)
# How often to poll the download client for status (seconds)
POLL_INTERVAL = 2
# How long to wait for completed files to appear (seconds)
COMPLETED_PATH_RETRY_INTERVAL = 5
COMPLETED_PATH_MAX_ATTEMPTS = 12 # 12 attempts * 5s = 60s grace period
@dataclass(frozen=True)
class DownloadRequest:
"""Source-specific download parameters resolved before sending to a client."""
url: str
protocol: str
release_name: str
expected_hash: Optional[str]
def _diagnose_path_issue(path: str) -> str:
"""
Analyze a path and return diagnostic hints for common issues.
Args:
path: The path that failed to be accessed
Returns:
A hint string to help users diagnose the issue.
"""
# Detect Windows-style paths (won't work in Linux containers)
if len(path) >= 2 and path[1] == ':':
return (
f"Path '{path}' appears to be a Windows path. "
f"Shelfmark runs in Linux and cannot access Windows paths directly. "
f"Ensure your download client uses Linux-style paths (/path/to/files)."
)
# Detect backslashes (Windows path separators)
if "\\" in path:
return (
f"Path '{path}' contains backslashes. "
f"This may indicate a Windows path or incorrect path escaping. "
f"Linux paths should use forward slashes (/)."
)
# Generic hint for Linux paths
return (
f"Path '{path}' is not accessible from Shelfmark's container. "
f"Ensure both containers have matching volume mounts for this directory, "
f"or configure Remote Path Mappings in Settings > Advanced."
)
class ExternalClientHandler(DownloadHandler, ABC):
"""Shared lifecycle handler for sources that hand off to torrent/usenet clients."""
def __init__(self):
# Track downloads that may need client-side cleanup after Shelfmark completes import.
# task_id -> (client, download_id, protocol)
self._cleanup_refs: dict[str, tuple[DownloadClient, str, str]] = {}
@abstractmethod
def _resolve_download(
self,
task: DownloadTask,
status_callback: Callable[[str, Optional[str]], None],
) -> Optional[DownloadRequest]:
"""Resolve source-specific task metadata into a client download request."""
def _on_download_complete(self, task: DownloadTask) -> None:
"""Hook called after successful completion; override for source cleanup."""
return
def _get_client(self, protocol: str) -> Optional[DownloadClient]:
"""Resolve the active client for a protocol."""
return get_client(protocol)
def _list_configured_clients(self) -> list[str]:
"""List protocols with configured clients."""
return list_configured_clients()
def _poll_interval(self) -> float:
"""Polling interval for status checks (seconds)."""
return POLL_INTERVAL
def _completed_path_retry_interval(self) -> float:
"""Retry interval while waiting for completed files (seconds)."""
return COMPLETED_PATH_RETRY_INTERVAL
def _completed_path_max_attempts(self) -> int:
"""Maximum attempts when waiting for completed files."""
return COMPLETED_PATH_MAX_ATTEMPTS
def _get_category_for_task(self, client: DownloadClient, task: DownloadTask) -> Optional[str]:
"""Get audiobook category if configured and applicable, else None for default."""
if not is_audiobook(task.content_type):
return None
# Client-specific audiobook category config keys
audiobook_keys = {
"qbittorrent": "QBITTORRENT_CATEGORY_AUDIOBOOK",
"transmission": "TRANSMISSION_CATEGORY_AUDIOBOOK",
"deluge": "DELUGE_CATEGORY_AUDIOBOOK",
"nzbget": "NZBGET_CATEGORY_AUDIOBOOK",
"sabnzbd": "SABNZBD_CATEGORY_AUDIOBOOK",
}
audiobook_key = audiobook_keys.get(client.name)
return config.get(audiobook_key, "") or None if audiobook_key else None
def post_process_cleanup(self, task: DownloadTask, success: bool) -> None:
if not success:
self._cleanup_refs.pop(task.task_id, None)
return
client_ref = self._cleanup_refs.pop(task.task_id, None)
if client_ref is None:
return
client, download_id, protocol = client_ref
if protocol != "usenet":
return
# "Move" means copy into ingest then let the usenet client delete its own files.
if config.get("PROWLARR_USENET_ACTION", "move") != "move":
return
try:
self._delete_local_download_data(client, download_id)
self._remove_usenet_download(client, download_id, delete_files=True, archive=True)
except Exception as e:
logger.warning(
f"Failed to cleanup usenet download {download_id} in {getattr(client, 'name', 'client')}: {e}"
)
def _remove_usenet_download(
self,
client: DownloadClient,
download_id: str,
*,
delete_files: bool,
archive: bool = True,
) -> None:
"""Remove a usenet download with SABnzbd-specific archive handling."""
if getattr(client, "name", "") == "sabnzbd":
client.remove(download_id, delete_files=delete_files, archive=archive)
else:
client.remove(download_id, delete_files=delete_files)
def _delete_local_download_data(self, client: DownloadClient, download_id: str) -> None:
"""Best-effort local deletion of client download data."""
try:
raw_path = client.get_download_path(download_id)
except Exception as e:
logger.debug(f"Failed to resolve download path for {client.name} {download_id}: {e}")
return
if not raw_path:
logger.debug(f"No download path available for {client.name} {download_id}")
return
from shelfmark.core.path_mappings import (
get_client_host_identifier,
parse_remote_path_mappings,
remap_remote_to_local_with_match,
)
source_path_obj = Path(raw_path)
host = get_client_host_identifier(client) or ""
mapping_value = config.get("PROWLARR_REMOTE_PATH_MAPPINGS", [])
mappings = parse_remote_path_mappings(mapping_value)
remapped, matched_mapping = remap_remote_to_local_with_match(
mappings=mappings,
host=host,
remote_path=source_path_obj,
)
delete_path = remapped if matched_mapping else source_path_obj
if str(delete_path) in ("", "/"):
logger.warning(f"Refusing to delete unsafe path for {client.name} {download_id}: {delete_path}")
return
if not run_blocking_io(delete_path.exists):
logger.debug(f"Local download path does not exist for cleanup: {delete_path}")
return
try:
if run_blocking_io(delete_path.is_dir):
run_blocking_io(shutil.rmtree, delete_path)
else:
run_blocking_io(delete_path.unlink)
logger.info(f"Deleted local download data for {client.name} {download_id}: {delete_path}")
except Exception as e:
logger.warning(f"Failed to delete local download data for {client.name} {download_id}: {e}")
def _safe_remove_download(self, client, download_id: str, protocol: str, reason: str) -> None:
"""Best-effort removal of a failed/cancelled download from the client.
Safety policy:
- torrents: never remove or delete client data (avoid breaking seeding)
- usenet: keep legacy behavior (delete client files on removal)
"""
if protocol != "usenet":
logger.info(
"Skipping download client cleanup for protocol=%s after %s (client=%s id=%s)",
protocol,
reason,
getattr(client, "name", "client"),
download_id,
)
return
try:
# Permanent delete for failed usenet downloads (SABnzbd archive=0).
self._delete_local_download_data(client, download_id)
self._remove_usenet_download(client, download_id, delete_files=True, archive=False)
except Exception as e:
logger.warning(
f"Failed to remove download {download_id} from {client.name} after {reason}: {e}"
)
def _handle_cancelled_download(
self,
client: DownloadClient,
download_id: str,
protocol: str,
status_callback: Callable[[str, Optional[str]], None],
) -> None:
if protocol == "usenet":
logger.info(f"Download cancelled, removing from {client.name}: {download_id}")
try:
self._delete_local_download_data(client, download_id)
self._remove_usenet_download(client, download_id, delete_files=True, archive=True)
except Exception as e:
logger.warning(
f"Failed to remove download {download_id} from {client.name} after cancellation: {e}"
)
else:
logger.info(
f"Download cancelled for protocol={protocol}; leaving in {client.name}: {download_id}"
)
status_callback("cancelled", "Cancelled")
def _resolve_download_path_once(
self,
client: DownloadClient,
download_id: str,
*,
log_details: bool,
) -> tuple[Optional[Path], Optional[str]]:
"""Resolve and validate the completed download path once."""
try:
raw_path = client.get_download_path(download_id)
except Exception as e:
message = (
f"Could not locate completed download in {client.name} (path not returned). "
f"Check volume mappings and category settings."
)
if log_details:
logger.error(
f"Failed to resolve download path for {client.name} {download_id}: {e}"
)
else:
logger.debug(
f"Failed to resolve download path for {client.name} {download_id}: {e}"
)
return None, message
if not raw_path:
message = (
f"Could not locate completed download in {client.name} (path not returned). "
f"Check volume mappings and category settings."
)
if log_details:
logger.error(f"Download client returned empty path for {client.name} {download_id}")
else:
logger.debug(f"Download client returned empty path for {client.name} {download_id}")
return None, message
from shelfmark.core.path_mappings import (
get_client_host_identifier,
parse_remote_path_mappings,
remap_remote_to_local_with_match,
)
source_path_obj = Path(raw_path)
host = get_client_host_identifier(client) or ""
mapping_value = config.get("PROWLARR_REMOTE_PATH_MAPPINGS", [])
mappings = parse_remote_path_mappings(mapping_value)
if log_details:
logger.debug(
"Attempting path remap: client=%s, host=%s, path=%s, mappings=%s",
client.name,
host,
source_path_obj,
[(m.host, m.remote_path, m.local_path) for m in mappings],
)
remapped, matched_mapping = remap_remote_to_local_with_match(
mappings=mappings,
host=host,
remote_path=source_path_obj,
)
if log_details:
remapped_exists = run_blocking_io(remapped.exists)
logger.debug(
"Remap result: %s -> %s (exists=%s, changed=%s, matched=%s)",
source_path_obj,
remapped,
remapped_exists,
remapped != source_path_obj,
matched_mapping,
)
if matched_mapping:
if run_blocking_io(remapped.exists):
logger.info(
"Remapped download path for %s (%s): %s -> %s",
client.name,
download_id,
source_path_obj,
remapped,
)
return remapped, None
message = (
f"Remapped path '{remapped}' does not exist. "
f"Check your Docker volume mounts match the Local Path in Settings > Advanced > Remote Path Mappings."
)
if log_details:
logger.error(
f"Download path does not exist after remapping: {raw_path} -> {remapped}. "
f"Client: {client.name}, ID: {download_id}."
)
else:
logger.debug(
f"Download path does not exist after remapping: {raw_path} -> {remapped}. "
f"Client: {client.name}, ID: {download_id}."
)
return None, message
if mappings:
if run_blocking_io(source_path_obj.exists):
logger.info(
"No remote path mapping matched for %s (%s); using client path: %s",
client.name,
download_id,
source_path_obj,
)
return source_path_obj, None
hint = _diagnose_path_issue(raw_path)
message = f"{hint} No remote path mapping matched for client '{client.name}'."
if log_details:
logger.error(
f"Download path does not exist and no remote path mapping matched for {client.name} "
f"({download_id}): {raw_path}. {hint}"
)
else:
logger.debug(
f"Download path does not exist and no remote path mapping matched for {client.name} "
f"({download_id}): {raw_path}. {hint}"
)
return None, message
if not run_blocking_io(source_path_obj.exists):
hint = _diagnose_path_issue(raw_path)
message = hint
if log_details:
logger.error(
f"Download path does not exist: {raw_path}. "
f"Client: {client.name}, ID: {download_id}. {hint}"
)
else:
logger.debug(
f"Download path does not exist: {raw_path}. "
f"Client: {client.name}, ID: {download_id}. {hint}"
)
return None, message
return source_path_obj, None
def _wait_for_completed_path(
self,
client: DownloadClient,
download_id: str,
*,
cancel_flag: Optional[Event],
status_callback: Callable[[str, Optional[str]], None],
) -> tuple[Optional[Path], Optional[str]]:
"""Wait briefly for completed files to appear on disk."""
last_error: Optional[str] = None
max_attempts = self._completed_path_max_attempts()
retry_interval = self._completed_path_retry_interval()
for attempt in range(1, max_attempts + 1):
if cancel_flag and cancel_flag.is_set():
return None, last_error
log_details = attempt == max_attempts
resolved_path, error = self._resolve_download_path_once(
client,
download_id,
log_details=log_details,
)
if resolved_path:
return resolved_path, None
last_error = error
if attempt < max_attempts:
status_callback("locating", "Waiting for completed files...")
logger.debug(
"Completed files not available yet for %s (%s) (attempt %d/%d)",
client.name,
download_id,
attempt,
max_attempts,
)
if cancel_flag:
if cancel_flag.wait(timeout=retry_interval):
return None, last_error
else:
time.sleep(retry_interval)
return None, last_error
def _build_progress_message(self, status) -> str:
"""Build a progress message from download status."""
msg = f"{status.progress:.0f}%"
if status.download_speed and status.download_speed > 0:
speed_mb = status.download_speed / 1024 / 1024
msg += f" ({speed_mb:.1f} MB/s)"
if status.eta and status.eta > 0:
if status.eta < 60:
msg += f" - {status.eta}s left"
elif status.eta < 3600:
msg += f" - {status.eta // 60}m left"
else:
msg += f" - {status.eta // 3600}h {(status.eta % 3600) // 60}m left"
return msg
def download(
self,
task: DownloadTask,
cancel_flag: Event,
progress_callback: Callable[[float], None],
status_callback: Callable[[str, Optional[str]], None],
) -> Optional[str]:
"""Execute download via configured torrent/usenet client. Returns file path or None."""
try:
if cancel_flag.is_set():
status_callback("cancelled", "Cancelled")
return None
request = self._resolve_download(task, status_callback)
if not request:
return None
client = self._get_client(request.protocol)
if not client:
configured = self._list_configured_clients()
if not configured:
status_callback(
"error",
"No download clients configured. Configure qBittorrent or NZBGet in settings.",
)
else:
status_callback("error", f"No {request.protocol} client configured")
return None
# Check if this download already exists in the client
status_callback("resolving", f"Checking {client.name}")
category = self._get_category_for_task(client, task)
existing = client.find_existing(request.url, category=category)
if existing:
download_id, existing_status = existing
logger.info(f"Found existing download in {client.name}: {download_id}")
# If already complete, skip straight to file handling
if existing_status.complete:
logger.info("Existing download is complete, copying file directly")
status_callback("resolving", "Found existing download, copying to library")
source_path_obj, path_error = self._wait_for_completed_path(
client=client,
download_id=download_id,
cancel_flag=cancel_flag,
status_callback=status_callback,
)
if not source_path_obj:
if cancel_flag.is_set():
return None
status_callback(
"error",
path_error
or f"Could not locate existing download in {client.name}. Check that the file still exists.",
)
return None
result = self._handle_completed_file(
source_path=source_path_obj,
protocol=request.protocol,
task=task,
status_callback=status_callback,
)
if result:
self._on_download_complete(task)
self._cleanup_refs[task.task_id] = (client, download_id, request.protocol)
return result
# Existing but still downloading - join the progress polling
logger.info("Existing download in progress, joining poll loop")
status_callback("downloading", "Resuming existing download")
else:
# No existing download - add new
status_callback("resolving", f"Sending to {client.name}")
try:
download_id = client.add_download(
url=request.url,
name=request.release_name,
category=category,
expected_hash=request.expected_hash,
)
except Exception as e:
logger.error(f"Failed to add to {client.name}: {e}")
status_callback("error", f"Failed to add to {client.name}: {e}")
return None
logger.info(f"Added to {client.name}: {download_id} for '{request.release_name}'")
# Poll for progress
return self._poll_and_complete(
client=client,
download_id=download_id,
protocol=request.protocol,
task=task,
cancel_flag=cancel_flag,
progress_callback=progress_callback,
status_callback=status_callback,
)
except Exception as e:
logger.error(f"External client download error: {e}")
status_callback("error", str(e))
return None
def _poll_and_complete(
self,
client: DownloadClient,
download_id: str,
protocol: str,
task: DownloadTask,
cancel_flag: Event,
progress_callback: Callable[[float], None],
status_callback: Callable[[str, Optional[str]], None],
) -> Optional[str]:
"""Poll the download client for progress and handle completion."""
poll_interval = self._poll_interval()
# Track consecutive "not found" errors - torrents may take time to appear in client
not_found_count = 0
max_not_found_retries = 15 # 15 retries * poll interval ~= 30s grace period
try:
logger.debug(f"Starting poll for {download_id} (content_type={task.content_type})")
while not cancel_flag.is_set():
status = client.get_status(download_id)
progress_callback(status.progress)
# Check for completion
if status.complete:
if status.state == DownloadState.ERROR:
logger.error(f"Download {download_id} completed with error: {status.message}")
status_callback("error", status.message or "Download failed")
self._safe_remove_download(client, download_id, protocol, "completion error")
return None
# Download complete - break to handle file
logger.debug(f"Download {download_id} complete, file_path={status.file_path}")
break
# Check for error state
if status.state == DownloadState.ERROR:
message = (status.message or "").strip()
message_lower = message.lower()
# Only treat *actual* "not found" as retryable.
# qBittorrent auth/network/API failures should surface immediately (more actionable)
# and must not be confused with "torrent missing".
retryable_not_found = any(
token in message_lower
for token in (
"torrent not found",
"not found in qbittorrent",
"download not found",
)
)
non_retryable = any(
token in message_lower
for token in (
"authentication failed",
"cannot connect",
"timed out",
"api request failed",
)
)
if retryable_not_found and not non_retryable:
not_found_count += 1
if not_found_count < max_not_found_retries:
logger.debug(
f"Download {download_id} not yet visible in client "
f"(attempt {not_found_count}/{max_not_found_retries})"
)
status_callback("resolving", "Waiting for download client...")
if cancel_flag.wait(timeout=poll_interval):
break
continue
logger.error(
f"Download {download_id} not found after {max_not_found_retries} attempts"
)
else:
# Fail fast on actionable errors (auth, connectivity, API issues)
logger.error(f"Download {download_id} error state: {status.message}")
status_callback("error", status.message or "Download failed")
self._safe_remove_download(client, download_id, protocol, "download error")
return None
# Reset not-found counter on successful status check
not_found_count = 0
# Build status message - use client message if provided, else build progress
msg = status.message or self._build_progress_message(status)
if status.state == DownloadState.PROCESSING:
# Post-processing (e.g., SABnzbd verifying/extracting)
status_callback("resolving", msg)
else:
status_callback("downloading", msg)
# Wait for next poll (interruptible by cancel)
if cancel_flag.wait(timeout=poll_interval):
break
# Handle cancellation
if cancel_flag.is_set():
self._handle_cancelled_download(client, download_id, protocol, status_callback)
return None
# Handle completed file (wait briefly for files to appear)
source_path_obj, path_error = self._wait_for_completed_path(
client=client,
download_id=download_id,
cancel_flag=cancel_flag,
status_callback=status_callback,
)
if not source_path_obj:
if cancel_flag.is_set():
self._handle_cancelled_download(client, download_id, protocol, status_callback)
return None
status_callback(
"error",
path_error
or f"Could not locate completed download in {client.name} (path not returned). Check volume mappings and category settings.",
)
return None
result = self._handle_completed_file(
source_path=source_path_obj,
protocol=protocol,
task=task,
status_callback=status_callback,
)
# Clean up on success
if result:
self._on_download_complete(task)
self._cleanup_refs[task.task_id] = (client, download_id, protocol)
return result
except Exception as e:
logger.error(f"Error during download polling: {e}")
status_callback("error", str(e))
self._safe_remove_download(client, download_id, protocol, "polling exception")
return None
def _handle_completed_file(
self,
source_path: Path,
protocol: str,
task: DownloadTask,
status_callback: Callable[[str, Optional[str]], None],
) -> Optional[str]:
"""Handle a completed download and return its path.
For external download clients (torrents/usenet), staging large payloads into TMP_DIR
is expensive (and can duplicate multi-GB files). Instead, return the client's
completed path and let the orchestrator perform any required transfer (copy/move/
hardlink) directly from that source.
Torrents also set ``task.original_download_path`` so the orchestrator can detect
seeding data and enable hardlinking when configured.
"""
try:
if protocol == "torrent":
task.original_download_path = str(source_path)
logger.debug(f"Download complete, returning original path: {source_path}")
return str(source_path)
except Exception as e:
logger.error(f"Failed to finalize completed download at {source_path}: {e}")
status_callback("error", f"Failed to finalize completed download: {e}")
return None
def cancel(self, task_id: str) -> bool:
"""Default cancellation (primary cancellation happens via cancel_flag)."""
logger.debug(f"Cancel requested for external client task: {task_id}")
return True
@@ -20,12 +20,12 @@ 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.release_sources.prowlarr.clients import (
from shelfmark.download.clients import (
DownloadClient,
DownloadStatus,
register_client,
)
from shelfmark.release_sources.prowlarr.clients.torrent_utils import (
from shelfmark.download.clients.torrent_utils import (
extract_torrent_info,
)
@@ -97,6 +97,7 @@ class DelugeClient(DownloadClient):
self._rpc_id = 0
self._category = str(config.get("DELUGE_CATEGORY", "books") or "books")
self._download_dir = str(config.get("DELUGE_DOWNLOAD_DIR", "") or "")
def _next_rpc_id(self) -> int:
self._rpc_id += 1
@@ -232,6 +233,8 @@ class DelugeClient(DownloadClient):
raise Exception("Failed to fetch torrent file")
options: dict[str, Any] = {}
if self._download_dir:
options["download_location"] = self._download_dir
if torrent_info.is_magnet:
magnet_url = torrent_info.magnet_url or url
@@ -12,7 +12,7 @@ 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.release_sources.prowlarr.clients import (
from shelfmark.download.clients import (
DownloadClient,
DownloadStatus,
register_client,
@@ -1,18 +1,19 @@
"""qBittorrent download client for Prowlarr integration."""
import time
from pathlib import Path
from types import SimpleNamespace
from typing import Optional, Tuple
from shelfmark.core.config import config
from shelfmark.core.logger import setup_logger
from shelfmark.core.utils import normalize_http_url
from shelfmark.release_sources.prowlarr.clients import (
from shelfmark.download.clients import (
DownloadClient,
DownloadStatus,
register_client,
)
from shelfmark.release_sources.prowlarr.clients.torrent_utils import (
from shelfmark.download.clients.torrent_utils import (
extract_torrent_info,
)
@@ -31,6 +32,35 @@ def _hashes_match(hash1: str, hash2: str) -> bool:
return False
def _normalize_tags(raw_tags: object) -> list[str]:
"""Normalize tag input to a clean, de-duplicated list of strings."""
if raw_tags is None:
return []
if isinstance(raw_tags, str):
parts = [part.strip() for part in raw_tags.split(",")]
elif isinstance(raw_tags, (list, tuple, set)):
parts = []
for item in raw_tags:
if item is None:
continue
parts.append(str(item).strip())
else:
parts = [str(raw_tags).strip()] if raw_tags else []
tags: list[str] = []
seen = set()
for part in parts:
if not part:
continue
if part in seen:
continue
seen.add(part)
tags.append(part)
return tags
@register_client("torrent")
class QBittorrentClient(DownloadClient):
"""qBittorrent download client."""
@@ -109,6 +139,8 @@ class QBittorrentClient(DownloadClient):
password=config.get("QBITTORRENT_PASSWORD", ""),
)
self._category = config.get("QBITTORRENT_CATEGORY", "books")
self._download_dir = config.get("QBITTORRENT_DOWNLOAD_DIR", "")
self._tags = _normalize_tags(config.get("QBITTORRENT_TAG", []))
def _get_torrents_info(
@@ -255,6 +287,7 @@ class QBittorrentClient(DownloadClient):
try:
# Use configured category if not explicitly provided
category = category or self._category
tags = self._tags
# Ensure category exists (may already exist, which is fine)
try:
@@ -270,19 +303,26 @@ class QBittorrentClient(DownloadClient):
torrent_data = torrent_info.torrent_data
# Add the torrent - use file content if we have it, otherwise URL
add_kwargs = {
"category": category,
"rename": name,
}
if self._download_dir:
add_kwargs["save_path"] = self._download_dir
if tags:
add_kwargs["tags"] = ",".join(tags)
if torrent_data:
result = self._client.torrents_add(
torrent_files=torrent_data,
category=category,
rename=name,
**add_kwargs,
)
else:
# Use magnet URL if available, otherwise original URL
add_url = torrent_info.magnet_url or url
result = self._client.torrents_add(
urls=add_url,
category=category,
rename=name,
**add_kwargs,
)
logger.debug(f"qBittorrent add result: {result}")
@@ -463,14 +503,37 @@ class QBittorrentClient(DownloadClient):
Centralizes the logic shared by `get_status()` and `get_download_path()`:
- accept `content_path` only when it's not equal to `save_path`
- when the torrent is complete and both `content_path` and `save_path` are present,
prefer a path rooted at `save_path` to avoid races where qBittorrent briefly reports
a temp/incomplete `content_path` and then moves the payload
- otherwise derive via properties+files
- finally fall back to `save_path + name`
"""
torrent_progress = getattr(torrent, "progress", 0.0)
try:
progress = float(torrent_progress)
except (TypeError, ValueError):
progress = 0.0
# Prefer content_path, but treat content_path == save_path as invalid.
content_path = getattr(torrent, "content_path", "")
save_path = getattr(torrent, "save_path", "")
if content_path and (not save_path or str(content_path) != str(save_path)):
# When using a temp/incomplete directory, qBittorrent can briefly keep reporting
# `content_path` under that temp path right at completion, then move the files
# into `save_path`. Returning the temp path can race with that move.
if save_path and progress >= 1.0:
# Use the basename of content_path under save_path (works for single-file
# torrents and multi-file torrents where content_path is a top-level dir).
try:
content_basename = str(Path(str(content_path)).name)
except Exception:
content_basename = ""
rooted = self._build_path(str(save_path), content_basename)
if rooted:
return rooted
return str(content_path)
download_id = getattr(torrent, "hash", "")
@@ -10,12 +10,12 @@ from urllib.parse import urlparse
from shelfmark.core.config import config
from shelfmark.core.logger import setup_logger
from shelfmark.core.utils import normalize_http_url
from shelfmark.release_sources.prowlarr.clients import (
from shelfmark.download.clients import (
DownloadClient,
DownloadStatus,
register_client,
)
from shelfmark.release_sources.prowlarr.clients.torrent_utils import (
from shelfmark.download.clients.torrent_utils import (
extract_torrent_info,
)
@@ -310,7 +310,17 @@ class RTorrentClient(DownloadClient):
this corresponds to `d.get_base_path()`.
"""
try:
base_path = self._rpc.d.get_base_path(download_id)
return base_path if base_path else None
# rTorrent is case sensitive for hashes; use uppercase as in get_status()
download_hash = download_id.upper()
details = self._rpc.d.multicall.filtered(
"",
"default",
f"equal={{d.hash=,cat={download_hash}}}",
"d.base_path=",
)
if not details:
return None
path = details[0][0]
return path if path else None
except Exception:
return None
@@ -12,7 +12,7 @@ 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.release_sources.prowlarr.clients import (
from shelfmark.download.clients import (
DownloadClient,
DownloadStatus,
register_client,
@@ -195,6 +195,7 @@ class SABnzbdClient(DownloadClient):
return response.content
def _get_prowlarr_headers(self, url: str) -> dict:
# TODO: Move this source-specific Prowlarr auth handling into a source hook.
api_key = str(config.get("PROWLARR_API_KEY", "") or "").strip()
if not api_key:
return {}
+678
View File
@@ -0,0 +1,678 @@
"""Shared download client settings registration."""
from typing import Any, Dict, Optional
from shelfmark.core.settings_registry import (
register_settings,
HeadingField,
TextField,
PasswordField,
ActionButton,
SelectField,
TagListField,
)
from shelfmark.core.utils import normalize_http_url
# ==================== Test Connection Callbacks ====================
def _test_qbittorrent_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test the qBittorrent connection using current form values."""
from shelfmark.core.config import config
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", "")
if not raw_url:
return {"success": False, "message": "qBittorrent URL is required"}
try:
from qbittorrentapi import Client
url = normalize_http_url(raw_url)
if not url:
return {"success": False, "message": "qBittorrent URL is invalid"}
client = Client(host=url, username=username, password=password)
client.auth_log_in()
api_version = client.app.web_api_version
return {"success": True, "message": f"Connected to qBittorrent (API v{api_version})"}
except ImportError:
return {"success": False, "message": "qbittorrent-api package not installed"}
except Exception as e:
return {"success": False, "message": f"Connection failed: {str(e)}"}
def _test_transmission_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test the Transmission connection using current form values."""
from shelfmark.core.config import config
from shelfmark.download.clients.torrent_utils import (
parse_transmission_url,
)
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", "")
if not raw_url:
return {"success": False, "message": "Transmission URL is required"}
url = normalize_http_url(raw_url)
if not url:
return {"success": False, "message": "Transmission URL is invalid"}
try:
from transmission_rpc import Client
# Parse URL to extract host, port, and path
protocol, host, port, path = parse_transmission_url(url)
client_kwargs = {
"host": host,
"port": port,
"path": path,
"username": username if username else None,
"password": password if password else None,
"protocol": protocol,
}
try:
client = Client(**client_kwargs)
except TypeError as e:
if "protocol" not in str(e):
raise
client_kwargs.pop("protocol", None)
client = Client(**client_kwargs)
if protocol == "https" and hasattr(client, "protocol"):
try:
setattr(client, "protocol", protocol)
except Exception:
pass
session = client.get_session()
version = session.version
return {"success": True, "message": f"Connected to Transmission {version}"}
except ImportError:
return {"success": False, "message": "transmission-rpc package not installed"}
except Exception as e:
return {"success": False, "message": f"Connection failed: {str(e)}"}
def _test_deluge_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test Deluge Web UI JSON-RPC connection using current form values."""
from urllib.parse import urlparse
import requests
from shelfmark.core.config import config
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", "")
if not raw_host:
return {"success": False, "message": "Deluge host is required"}
if not password:
return {"success": False, "message": "Deluge password is required"}
raw_host = str(raw_host)
raw_host = normalize_http_url(raw_host, strip_trailing_slash=False) if raw_host else ""
if not raw_host:
return {"success": False, "message": "Deluge host is invalid"}
raw_port = str(raw_port or "8112")
scheme = "http"
base_path = ""
host = raw_host
port = int(raw_port) if raw_port.isdigit() else 8112
# Allow DELUGE_HOST to be a full URL (e.g. http://deluge:8112)
if raw_host.startswith(("http://", "https://")):
parsed = urlparse(raw_host)
scheme = parsed.scheme or "http"
host = parsed.hostname or "localhost"
if parsed.port is not None:
port = parsed.port
base_path = (parsed.path or "").rstrip("/")
else:
# Allow "host:port" in DELUGE_HOST for convenience.
if ":" in raw_host and raw_host.count(":") == 1:
host_part, port_part = raw_host.split(":", 1)
if host_part and port_part.isdigit():
host = host_part
port = int(port_part)
rpc_url = f"{scheme}://{host}:{port}{base_path}/json"
def rpc_call(session: requests.Session, rpc_id: int, method: str, *params: Any) -> Any:
payload = {"id": rpc_id, "method": method, "params": list(params)}
resp = session.post(rpc_url, json=payload, timeout=15)
resp.raise_for_status()
data = resp.json()
if data.get("error"):
error = data["error"]
if isinstance(error, dict):
raise Exception(error.get("message") or str(error))
raise Exception(str(error))
return data.get("result")
def get_daemon_version(session: requests.Session, rpc_id: int) -> Any:
try:
methods = rpc_call(session, rpc_id, "system.listMethods")
if isinstance(methods, list) and "daemon.get_version" in methods:
return rpc_call(session, rpc_id + 1, "daemon.get_version")
except Exception:
# Fall back to daemon.info to preserve existing behavior.
pass
return rpc_call(session, rpc_id + 1, "daemon.info")
try:
session = requests.Session()
if rpc_call(session, 1, "auth.login", password) is not True:
return {"success": False, "message": "Deluge Web UI authentication failed"}
if rpc_call(session, 2, "web.connected") is not True:
hosts = rpc_call(session, 3, "web.get_hosts") or []
if not hosts:
return {
"success": False,
"message": "Deluge Web UI isn't connected to Deluge core (no hosts configured). Add/connect a daemon in Deluge Web UI → Connection Manager.",
}
host_id = hosts[0][0]
for entry in hosts:
if isinstance(entry, list) and len(entry) >= 2 and entry[1] in {"127.0.0.1", "localhost"}:
host_id = entry[0]
break
rpc_call(session, 4, "web.connect", host_id)
if rpc_call(session, 5, "web.connected") is not True:
return {
"success": False,
"message": "Deluge Web UI couldn't connect to Deluge core. Check Deluge Web UI → Connection Manager.",
}
version = get_daemon_version(session, 6)
return {"success": True, "message": f"Connected to Deluge {version}"}
except requests.exceptions.ConnectionError:
return {"success": False, "message": "Could not connect to Deluge Web UI"}
except requests.exceptions.Timeout:
return {"success": False, "message": "Connection timed out"}
except Exception as e:
return {"success": False, "message": f"Connection failed: {str(e)}"}
def _test_rtorrent_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test the rTorrent connection using current form values."""
from shelfmark.core.config import config
from urllib.parse import urlparse
from xmlrpc.client import ServerProxy
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", "")
if not raw_url:
return {"success": False, "message": "rTorrent URL is required"}
url = normalize_http_url(raw_url)
if not url:
return {"success": False, "message": "rTorrent URL is invalid"}
try:
# Add HTTP auth to URL if credentials provided
if username and password:
parsed = urlparse(url)
url = f"{parsed.scheme}://{username}:{password}@{parsed.netloc}{parsed.path}"
rpc = ServerProxy(url.rstrip("/"))
version = rpc.system.client_version()
return {"success": True, "message": f"Connected to rTorrent {version}"}
except Exception as e:
return {"success": False, "message": f"Connection failed: {str(e)}"}
def _test_nzbget_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test the NZBGet connection using current form values."""
import requests
from shelfmark.core.config import config
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", "")
if not raw_url:
return {"success": False, "message": "NZBGet URL is required"}
url = normalize_http_url(raw_url)
if not url:
return {"success": False, "message": "NZBGet URL is invalid"}
try:
rpc_url = f"{url.rstrip('/')}/jsonrpc"
payload = {"jsonrpc": "2.0", "method": "status", "params": [], "id": 1}
response = requests.post(rpc_url, json=payload, auth=(username, password), timeout=30)
response.raise_for_status()
result = response.json()
if "error" in result and result["error"]:
raise Exception(result["error"].get("message", "RPC error"))
version = result.get("result", {}).get("Version", "unknown")
return {"success": True, "message": f"Connected to NZBGet {version}"}
except requests.exceptions.ConnectionError:
return {"success": False, "message": "Could not connect to NZBGet"}
except requests.exceptions.Timeout:
return {"success": False, "message": "Connection timed out"}
except Exception as e:
return {"success": False, "message": f"Connection failed: {str(e)}"}
def _test_sabnzbd_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test the SABnzbd connection using current form values."""
import requests
from shelfmark.core.config import config
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", "")
if not raw_url:
return {"success": False, "message": "SABnzbd URL is required"}
url = normalize_http_url(raw_url)
if not url:
return {"success": False, "message": "SABnzbd URL is invalid"}
if not api_key:
return {"success": False, "message": "API key is required"}
try:
api_url = f"{url.rstrip('/')}/api"
params = {"apikey": api_key, "mode": "version", "output": "json"}
response = requests.get(api_url, params=params, timeout=30)
response.raise_for_status()
result = response.json()
version = result.get("version", "unknown")
return {"success": True, "message": f"Connected to SABnzbd {version}"}
except requests.exceptions.ConnectionError:
return {"success": False, "message": "Could not connect to SABnzbd"}
except requests.exceptions.Timeout:
return {"success": False, "message": "Connection timed out"}
except Exception as e:
return {"success": False, "message": f"Connection failed: {str(e)}"}
# ==================== Download Clients Tab ====================
@register_settings(
name="prowlarr_clients",
display_name="Download Clients",
icon="cog",
order=110,
)
def prowlarr_clients_settings():
"""Download client settings shared by external release sources."""
return [
# --- Torrent Client Selection ---
HeadingField(
key="torrent_heading",
title="Torrent Client",
description="Select and configure a torrent client for downloading torrent releases.",
),
SelectField(
key="PROWLARR_TORRENT_CLIENT",
label="Torrent Client",
description="Choose which torrent client to use",
options=[
{"value": "", "label": "None"},
{"value": "qbittorrent", "label": "qBittorrent"},
{"value": "transmission", "label": "Transmission"},
{"value": "deluge", "label": "Deluge"},
{"value": "rtorrent", "label": "rTorrent"},
],
default="",
),
# --- qBittorrent Settings ---
TextField(
key="QBITTORRENT_URL",
label="qBittorrent URL",
description="Web UI URL of your qBittorrent instance",
placeholder="http://qbittorrent:8080",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
TextField(
key="QBITTORRENT_USERNAME",
label="Username",
description="qBittorrent Web UI username",
placeholder="admin",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
PasswordField(
key="QBITTORRENT_PASSWORD",
label="Password",
description="qBittorrent Web UI password",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
ActionButton(
key="test_qbittorrent",
label="Test Connection",
description="Verify your qBittorrent configuration",
style="primary",
callback=_test_qbittorrent_connection,
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
TextField(
key="QBITTORRENT_CATEGORY",
label="Book Category",
description="Category to assign to book downloads in qBittorrent",
placeholder="books",
default="books",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
TextField(
key="QBITTORRENT_CATEGORY_AUDIOBOOK",
label="Audiobook Category",
description="Category for audiobook downloads. Leave empty to use the book category.",
placeholder="",
default="",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
TextField(
key="QBITTORRENT_DOWNLOAD_DIR",
label="Download Directory",
description="Server-side directory where torrents are downloaded (optional, uses qBittorrent default if not specified)",
placeholder="/downloads",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
TagListField(
key="QBITTORRENT_TAG",
label="Tags",
description="Tag(s) to assign to qBittorrent downloads. Leave empty for no tags.",
placeholder="",
default=[],
normalize_urls=False,
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
# --- Transmission Settings ---
TextField(
key="TRANSMISSION_URL",
label="Transmission URL",
description="URL of your Transmission instance (use https:// for TLS)",
placeholder="http://transmission:9091",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
TextField(
key="TRANSMISSION_USERNAME",
label="Username",
description="Transmission RPC username (if authentication enabled)",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
PasswordField(
key="TRANSMISSION_PASSWORD",
label="Password",
description="Transmission RPC password",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
ActionButton(
key="test_transmission",
label="Test Connection",
description="Verify your Transmission configuration",
style="primary",
callback=_test_transmission_connection,
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
TextField(
key="TRANSMISSION_CATEGORY",
label="Book Label",
description="Label to assign to book downloads in Transmission",
placeholder="books",
default="books",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
TextField(
key="TRANSMISSION_CATEGORY_AUDIOBOOK",
label="Audiobook Label",
description="Label for audiobook downloads. Leave empty to use the book label.",
placeholder="",
default="",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
TextField(
key="TRANSMISSION_DOWNLOAD_DIR",
label="Download Directory",
description="Server-side directory where torrents are downloaded (optional, uses Transmission default if not specified)",
placeholder="/downloads",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
# --- Deluge Settings ---
TextField(
key="DELUGE_HOST",
label="Deluge Web UI Host/URL",
description="Hostname/IP or full URL of your Deluge Web UI (deluge-web)",
placeholder="http://deluge:8112",
default="localhost",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
TextField(
key="DELUGE_PORT",
label="Deluge Web UI Port",
description="Deluge Web UI port (default: 8112)",
placeholder="8112",
default="8112",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
PasswordField(
key="DELUGE_PASSWORD",
label="Password",
description="Deluge Web UI password (default: deluge)",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
ActionButton(
key="test_deluge",
label="Test Connection",
description="Verify your Deluge configuration",
style="primary",
callback=_test_deluge_connection,
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
TextField(
key="DELUGE_CATEGORY",
label="Book Label",
description="Label to assign to book downloads in Deluge",
placeholder="books",
default="books",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
TextField(
key="DELUGE_CATEGORY_AUDIOBOOK",
label="Audiobook Label",
description="Label for audiobook downloads. Leave empty to use the book label.",
placeholder="",
default="",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
TextField(
key="DELUGE_DOWNLOAD_DIR",
label="Download Directory",
description="Server-side directory where torrents are downloaded (optional, uses Deluge default if not specified)",
placeholder="/downloads",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
# --- rTorrent Settings ---
TextField(
key="RTORRENT_URL",
label="rTorrent URL",
description="XML-RPC URL of your rTorrent instance",
placeholder="http://rtorrent:6881/RPC2",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "rtorrent"},
),
TextField(
key="RTORRENT_USERNAME",
label="Username",
description="HTTP Basic auth username (if authentication enabled)",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "rtorrent"},
),
PasswordField(
key="RTORRENT_PASSWORD",
label="Password",
description="HTTP Basic auth password",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "rtorrent"},
),
ActionButton(
key="test_rtorrent",
label="Test Connection",
description="Verify your rTorrent configuration",
style="primary",
callback=_test_rtorrent_connection,
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "rtorrent"},
),
TextField(
key="RTORRENT_LABEL",
label="Book Label",
description="Label to assign to book downloads in rTorrent",
placeholder="cwabd",
default="cwabd",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "rtorrent"},
),
TextField(
key="RTORRENT_DOWNLOAD_DIR",
label="Download Directory",
description="Server-side directory where torrents are downloaded (optional, uses rTorrent default if not specified)",
placeholder="/downloads",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "rtorrent"},
),
# Note: Torrent client download path must be mounted identically in both containers.
# Torrents are always copied (not moved) to preserve seeding capability.
# --- Usenet Client Selection ---
HeadingField(
key="usenet_heading",
title="Usenet Client",
description="Select and configure a usenet client for downloading NZB releases.",
),
SelectField(
key="PROWLARR_USENET_CLIENT",
label="Usenet Client",
description="Choose which usenet client to use",
options=[
{"value": "", "label": "None"},
{"value": "nzbget", "label": "NZBGet"},
{"value": "sabnzbd", "label": "SABnzbd"},
],
default="",
),
# --- NZBGet Settings ---
TextField(
key="NZBGET_URL",
label="NZBGet URL",
description="URL of your NZBGet instance",
placeholder="http://nzbget:6789",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "nzbget"},
),
TextField(
key="NZBGET_USERNAME",
label="Username",
description="NZBGet control username",
placeholder="nzbget",
default="nzbget",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "nzbget"},
),
PasswordField(
key="NZBGET_PASSWORD",
label="Password",
description="NZBGet control password",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "nzbget"},
),
ActionButton(
key="test_nzbget",
label="Test Connection",
description="Verify your NZBGet configuration",
style="primary",
callback=_test_nzbget_connection,
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "nzbget"},
),
TextField(
key="NZBGET_CATEGORY",
label="Book Category",
description="Category to assign to book downloads in NZBGet",
placeholder="Books",
default="Books",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "nzbget"},
),
TextField(
key="NZBGET_CATEGORY_AUDIOBOOK",
label="Audiobook Category",
description="Category for audiobook downloads. Leave empty to use the book category.",
placeholder="",
default="",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "nzbget"},
),
# --- SABnzbd Settings ---
TextField(
key="SABNZBD_URL",
label="SABnzbd URL",
description="URL of your SABnzbd instance",
placeholder="http://sabnzbd:8080",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "sabnzbd"},
),
PasswordField(
key="SABNZBD_API_KEY",
label="API Key",
description="Found in SABnzbd: Config > General > API Key",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "sabnzbd"},
),
ActionButton(
key="test_sabnzbd",
label="Test Connection",
description="Verify your SABnzbd configuration",
style="primary",
callback=_test_sabnzbd_connection,
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "sabnzbd"},
),
TextField(
key="SABNZBD_CATEGORY",
label="Book Category",
description="Category to assign to book downloads in SABnzbd",
placeholder="books",
default="books",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "sabnzbd"},
),
TextField(
key="SABNZBD_CATEGORY_AUDIOBOOK",
label="Audiobook Category",
description="Category for audiobook downloads. Leave empty to use the book category.",
placeholder="",
default="",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "sabnzbd"},
),
# Note: Usenet client download path must be mounted identically in both containers.
SelectField(
key="PROWLARR_USENET_ACTION",
label="NZB Completion Action",
description="Move deletes the job from your usenet client after import; Copy keeps it in the client",
options=[
{"value": "move", "label": "Move"},
{"value": "copy", "label": "Copy"},
],
default="move",
show_when={"field": "PROWLARR_USENET_CLIENT", "notEmpty": True},
),
]
@@ -72,6 +72,7 @@ def extract_torrent_info(
return TorrentInfo(info_hash=expected_hash, torrent_data=None, is_magnet=False)
headers: dict[str, str] = {"Accept": "application/x-bittorrent"}
# TODO: Move this source-specific Prowlarr auth handling into a source hook.
api_key = str(config.get("PROWLARR_API_KEY", "") or "").strip()
if api_key:
headers["X-Api-Key"] = api_key
@@ -134,9 +135,12 @@ def extract_torrent_info(
return TorrentInfo(info_hash=expected_hash, torrent_data=None, is_magnet=False)
def parse_transmission_url(url: str) -> Tuple[str, int, str]:
"""Parse Transmission URL into (host, port, path)."""
def parse_transmission_url(url: str) -> Tuple[str, str, int, str]:
"""Parse Transmission URL into (protocol, host, port, path)."""
parsed = urlparse(url)
protocol = (parsed.scheme or "http").lower()
if protocol not in ("http", "https"):
protocol = "http"
host = parsed.hostname or "localhost"
port = parsed.port or 9091
path = parsed.path or "/transmission/rpc"
@@ -145,7 +149,7 @@ def parse_transmission_url(url: str) -> Tuple[str, int, str]:
if not path.endswith("/rpc"):
path = path.rstrip("/") + "/transmission/rpc"
return host, port, path
return protocol, host, port, path
def bencode_decode(data: bytes) -> tuple:
@@ -10,12 +10,12 @@ from typing import Optional, Tuple
from shelfmark.core.config import config
from shelfmark.core.logger import setup_logger
from shelfmark.core.utils import normalize_http_url
from shelfmark.release_sources.prowlarr.clients import (
from shelfmark.download.clients import (
DownloadClient,
DownloadStatus,
register_client,
)
from shelfmark.release_sources.prowlarr.clients.torrent_utils import (
from shelfmark.download.clients.torrent_utils import (
extract_torrent_info,
parse_transmission_url,
)
@@ -46,16 +46,32 @@ class TransmissionClient(DownloadClient):
password = config.get("TRANSMISSION_PASSWORD", "")
# Parse URL to extract host, port, and path
host, port, path = parse_transmission_url(url)
protocol, host, port, path = parse_transmission_url(url)
self._client = Client(
host=host,
port=port,
path=path,
username=username if username else None,
password=password if password else None,
)
client_kwargs = {
"host": host,
"port": port,
"path": path,
"username": username if username else None,
"password": password if password else None,
"protocol": protocol,
}
try:
self._client = Client(**client_kwargs)
except TypeError as e:
# Older transmission-rpc versions may not accept protocol as a kwarg.
if "protocol" not in str(e):
raise
client_kwargs.pop("protocol", None)
self._client = Client(**client_kwargs)
# Some versions expose protocol as an attribute rather than kwarg.
if protocol == "https" and hasattr(self._client, "protocol"):
try:
setattr(self._client, "protocol", protocol)
except Exception:
pass
self._category = config.get("TRANSMISSION_CATEGORY", "books")
self._download_dir = config.get("TRANSMISSION_DOWNLOAD_DIR", "")
@staticmethod
def is_configured() -> bool:
@@ -100,18 +116,24 @@ class TransmissionClient(DownloadClient):
resolved_category = category or self._category or ""
torrent_info = extract_torrent_info(url, expected_hash=expected_hash)
add_kwargs = {}
if resolved_category:
add_kwargs["labels"] = [resolved_category]
if self._download_dir:
add_kwargs["download_dir"] = self._download_dir
if torrent_info.torrent_data:
torrent = self._client.add_torrent(
torrent=torrent_info.torrent_data,
labels=[resolved_category] if resolved_category else None,
**add_kwargs,
)
else:
# Use magnet URL if available, otherwise original URL
add_url = torrent_info.magnet_url or url
torrent = self._client.add_torrent(
torrent=add_url,
labels=[resolved_category] if resolved_category else None,
**add_kwargs,
)
torrent_hash = torrent.hashString.lower()
+226 -84
View File
@@ -8,14 +8,66 @@ import errno
import os
import shutil
import subprocess
import tempfile
import time
from pathlib import Path
from typing import Any, Callable, Optional, TypeVar, cast
from shelfmark.core.logger import setup_logger
from shelfmark.download.permissions_debug import log_transfer_permission_context
logger = setup_logger(__name__)
try:
from gevent import monkey as _gevent_monkey
from gevent.threadpool import ThreadPool as _GeventThreadPool
except Exception:
_gevent_monkey = None
_GeventThreadPool = None
T = TypeVar("T")
_IO_THREADPOOL: Optional["_GeventThreadPool"] = None
def _use_gevent_threadpool() -> bool:
return bool(
_gevent_monkey
and _GeventThreadPool
and _gevent_monkey.is_module_patched("threading")
)
def _get_io_threadpool() -> "_GeventThreadPool":
global _IO_THREADPOOL
if _IO_THREADPOOL is None:
pool_size = max(2, min(8, os.cpu_count() or 2))
_IO_THREADPOOL = _GeventThreadPool(pool_size)
return _IO_THREADPOOL
def _call_and_capture(func: Callable[..., T], args: tuple[Any, ...], kwargs: dict[str, Any]) -> tuple[bool, T | Exception]:
try:
return True, func(*args, **kwargs)
except Exception as exc:
return False, exc
def run_blocking_io(func: Callable[..., T], *args: Any, **kwargs: Any) -> T:
"""Run blocking I/O in a native thread when under gevent.
gevent's threadpool will eagerly log exceptions raised inside worker threads,
even when the caller expects and handles those errors (e.g. FileExistsError for
collision retries, EXDEV for cross-device moves). Capture and re-raise in the
caller to avoid noisy, misleading tracebacks.
"""
if _use_gevent_threadpool():
ok, result = _get_io_threadpool().apply(_call_and_capture, (func, args, kwargs))
if ok:
return cast(T, result)
exc = cast(Exception, result)
raise exc
return func(*args, **kwargs)
_VERIFY_IO_WAIT_SECONDS = 3.0
@@ -31,7 +83,8 @@ def _verify_transfer_size(
Some filesystems (especially remote NAS/CIFS/NFS) can report stale sizes briefly
after large writes. Do a second stat after a short delay before declaring failure.
"""
actual_size = dest.stat().st_size
# On network filesystems, `stat()` can block long enough to starve the gevent hub.
actual_size = run_blocking_io(dest.stat).st_size
if actual_size == expected_size:
return
@@ -41,7 +94,7 @@ def _verify_transfer_size(
)
time.sleep(_VERIFY_IO_WAIT_SECONDS)
actual_size = dest.stat().st_size
actual_size = run_blocking_io(dest.stat).st_size
if actual_size != expected_size:
raise IOError(
f"File {action} incomplete, data loss may have occurred. "
@@ -74,11 +127,16 @@ def atomic_write(dest_path: Path, data: bytes, max_attempts: int = 100) -> Path:
try_path = dest_path if attempt == 0 else parent / f"{base}_{attempt}{ext}"
try:
# O_CREAT | O_EXCL fails atomically if file exists
fd = os.open(str(try_path), os.O_CREAT | os.O_EXCL | os.O_WRONLY, 0o666)
fd = run_blocking_io(
os.open,
str(try_path),
os.O_CREAT | os.O_EXCL | os.O_WRONLY,
0o666,
)
try:
os.write(fd, data)
run_blocking_io(os.write, fd, data)
finally:
os.close(fd)
run_blocking_io(os.close, fd)
if attempt > 0:
logger.info(f"File collision resolved: {try_path.name}")
return try_path
@@ -96,30 +154,31 @@ def _is_permission_error(e: Exception) -> bool:
def _system_op(op: str, source: Path, dest: Path) -> None:
"""Execute system command (mv or cp) as final fallback."""
logger.warning("Attempting system %s as final fallback: %s -> %s", op, source, dest)
subprocess.run(
run_blocking_io(
subprocess.run,
[op, "-f", str(source), str(dest)],
check=True,
capture_output=True,
text=True
text=True,
)
def _perform_nfs_fallback(source: Path, dest: Path, is_move: bool) -> None:
"""Handle NFS/SMB permission errors by falling back to copyfile -> system op."""
expected_size = source.stat().st_size
expected_size = run_blocking_io(source.stat).st_size
try:
# Fallback 1: copy content only
shutil.copyfile(str(source), str(dest))
run_blocking_io(shutil.copyfile, str(source), str(dest))
_verify_transfer_size(dest, expected_size, "copy")
if is_move:
source.unlink()
run_blocking_io(source.unlink)
return
except Exception as copy_error:
# Clean up failed copy attempt if it exists
dest.unlink(missing_ok=True)
run_blocking_io(dest.unlink, missing_ok=True)
if _is_permission_error(copy_error):
log_transfer_permission_context("nfs_fallback_copyfile", source=source, dest=dest, error=copy_error)
@@ -130,14 +189,14 @@ def _perform_nfs_fallback(source: Path, dest: Path, is_move: bool) -> None:
try:
_system_op(op, source, dest)
# Best-effort verify after external command.
if dest.exists():
if run_blocking_io(dest.exists):
_verify_transfer_size(dest, expected_size, op)
if is_move:
source.unlink(missing_ok=True)
run_blocking_io(source.unlink, missing_ok=True)
except subprocess.CalledProcessError as sys_error:
log_transfer_permission_context("nfs_fallback_system", source=source, dest=dest, error=sys_error)
logger.error("System %s failed (%s -> %s): %s", op, source, dest, sys_error.stderr)
dest.unlink(missing_ok=True)
run_blocking_io(dest.unlink, missing_ok=True)
raise
@@ -147,19 +206,86 @@ def _claim_destination(path: Path) -> bool:
Returns True if the placeholder was created. Caller must replace or unlink it.
"""
try:
fd = os.open(str(path), os.O_CREAT | os.O_EXCL | os.O_WRONLY, 0o666)
fd = run_blocking_io(
os.open,
str(path),
os.O_CREAT | os.O_EXCL | os.O_WRONLY,
0o666,
)
except FileExistsError:
return False
else:
os.close(fd)
run_blocking_io(os.close, fd)
return True
def _hardlink_not_supported(error: OSError) -> bool:
err = error.errno
return err in {
errno.EXDEV,
errno.EMLINK,
errno.EPERM,
errno.EACCES,
getattr(errno, "ENOTSUP", errno.EPERM),
getattr(errno, "EOPNOTSUPP", errno.EPERM),
errno.EINVAL,
}
def _create_temp_path(dest_path: Path) -> Path:
fd, temp_path = run_blocking_io(
tempfile.mkstemp,
prefix=f".{dest_path.name}.",
suffix=".tmp",
dir=str(dest_path.parent),
)
run_blocking_io(os.close, fd)
return Path(temp_path)
def _publish_temp_file(temp_path: Path, dest_path: Path) -> bool:
"""Publish a temp file to its final path without overwriting existing files.
Returns True on success, False if the destination already exists.
"""
try:
run_blocking_io(os.link, str(temp_path), str(dest_path))
run_blocking_io(temp_path.unlink, missing_ok=True)
return True
except FileExistsError:
return False
except OSError as e:
if _is_permission_error(e):
log_transfer_permission_context(
"publish_hardlink",
source=temp_path,
dest=dest_path,
error=e,
)
if _hardlink_not_supported(e):
logger.debug(
"Hardlink publish unsupported; falling back to claim+replace: %s -> %s (%s)",
temp_path,
dest_path,
e,
)
claimed = _claim_destination(dest_path)
if not claimed:
return False
try:
run_blocking_io(os.replace, str(temp_path), str(dest_path))
except Exception:
run_blocking_io(dest_path.unlink, missing_ok=True)
raise
return True
raise
def atomic_move(source_path: Path, dest_path: Path, max_attempts: int = 100) -> Path:
"""Move a file with collision detection.
Uses os.rename() for same-filesystem moves (atomic, triggers inotify events),
falls back to exclusive create + shutil.move for cross-filesystem moves.
falls back to copy-then-publish for cross-filesystem moves.
Note: We use os.rename() instead of hardlink+unlink because os.rename()
triggers proper inotify IN_MOVED_TO events that file watchers (like Calibre's
@@ -185,7 +311,7 @@ def atomic_move(source_path: Path, dest_path: Path, max_attempts: int = 100) ->
# Check for existing file (os.rename would overwrite on Unix)
claimed = False
if try_path.exists():
if run_blocking_io(try_path.exists):
# Some filesystems can report false positives for exists() with
# special characters. Probe with O_EXCL to confirm.
claimed = _claim_destination(try_path)
@@ -195,37 +321,35 @@ def atomic_move(source_path: Path, dest_path: Path, max_attempts: int = 100) ->
try:
# os.rename is atomic on same filesystem and triggers inotify events
if claimed:
os.replace(str(source_path), str(try_path))
run_blocking_io(os.replace, str(source_path), str(try_path))
else:
os.rename(str(source_path), str(try_path))
run_blocking_io(os.rename, str(source_path), str(try_path))
if attempt > 0:
logger.info(f"File collision resolved: {try_path.name}")
return try_path
except FileExistsError:
# Race condition: file created between exists() check and rename()
if claimed:
try_path.unlink(missing_ok=True)
run_blocking_io(try_path.unlink, missing_ok=True)
continue
except OSError as e:
# Cross-filesystem - fall back to exclusive create + verified copy + delete.
# Cross-filesystem - copy to temp and publish atomically.
if e.errno != errno.EXDEV:
if claimed:
try_path.unlink(missing_ok=True)
run_blocking_io(try_path.unlink, missing_ok=True)
raise
expected_size = source_path.stat().st_size
expected_size = run_blocking_io(source_path.stat).st_size
if claimed:
run_blocking_io(try_path.unlink, missing_ok=True)
claimed = False
temp_path: Optional[Path] = None
try:
if not claimed:
# Claim destination path atomically.
fd = os.open(str(try_path), os.O_CREAT | os.O_EXCL | os.O_WRONLY, 0o666)
os.close(fd)
# Copy to a temp file first, then replace to avoid partial files.
temp_path = try_path.parent / f".{try_path.name}.tmp"
try:
temp_path = _create_temp_path(try_path)
try:
shutil.copy2(str(source_path), str(temp_path))
run_blocking_io(shutil.copy2, str(source_path), str(temp_path))
except (PermissionError, OSError) as copy_error:
if _is_permission_error(copy_error):
logger.debug(
@@ -238,21 +362,33 @@ def atomic_move(source_path: Path, dest_path: Path, max_attempts: int = 100) ->
else:
raise
temp_path.replace(try_path)
_verify_transfer_size(try_path, expected_size, "move")
source_path.unlink()
_verify_transfer_size(temp_path, expected_size, "move")
published = _publish_temp_file(temp_path, try_path)
if not published:
run_blocking_io(temp_path.unlink, missing_ok=True)
continue
try:
_verify_transfer_size(try_path, expected_size, "move")
except Exception:
run_blocking_io(try_path.unlink, missing_ok=True)
raise
run_blocking_io(source_path.unlink)
if attempt > 0:
logger.info(f"File collision resolved: {try_path.name}")
return try_path
except FileExistsError:
if temp_path:
run_blocking_io(temp_path.unlink, missing_ok=True)
continue
except Exception:
try_path.unlink(missing_ok=True)
temp_path.unlink(missing_ok=True)
if temp_path:
run_blocking_io(temp_path.unlink, missing_ok=True)
raise
except FileExistsError:
continue
except (PermissionError, OSError) as e:
if _is_permission_error(e):
log_transfer_permission_context(
@@ -306,7 +442,7 @@ def atomic_hardlink(source_path: Path, dest_path: Path, max_attempts: int = 100)
for attempt in range(max_attempts):
try_path = dest_path if attempt == 0 else parent / f"{base}_{attempt}{ext}"
try:
os.link(str(source_path), str(try_path))
run_blocking_io(os.link, str(source_path), str(try_path))
if attempt > 0:
logger.info(f"File collision resolved: {try_path.name}")
return try_path
@@ -336,8 +472,8 @@ def atomic_hardlink(source_path: Path, dest_path: Path, max_attempts: int = 100)
def atomic_copy(source_path: Path, dest_path: Path, max_attempts: int = 100) -> Path:
"""Copy a file with atomic collision detection.
Uses exclusive create to claim destination, then copies via temp file
to avoid partial files on failure.
Uses a temp file in the destination directory and publishes it atomically,
avoiding partial files on failure.
Args:
source_path: Source file to copy
@@ -353,57 +489,63 @@ def atomic_copy(source_path: Path, dest_path: Path, max_attempts: int = 100) ->
base = dest_path.stem
ext = dest_path.suffix
parent = dest_path.parent
expected_size = run_blocking_io(source_path.stat).st_size
for attempt in range(max_attempts):
try_path = dest_path if attempt == 0 else parent / f"{base}_{attempt}{ext}"
if run_blocking_io(try_path.exists):
continue
temp_path: Optional[Path] = None
try:
# Atomically claim the destination by creating an exclusive file
fd = os.open(str(try_path), os.O_CREAT | os.O_EXCL | os.O_WRONLY, 0o666)
os.close(fd)
# Copy to temp file first, then replace to avoid partial files
temp_path = try_path.parent / f".{try_path.name}.tmp"
temp_path = _create_temp_path(try_path)
try:
try:
shutil.copy2(str(source_path), str(temp_path))
except (PermissionError, OSError) as e:
# Handle NFS permission errors immediately here
if _is_permission_error(e):
log_transfer_permission_context(
"atomic_copy",
source=source_path,
dest=temp_path,
error=e,
)
logger.debug(
"Permission error during copy, falling back to copyfile (%s -> %s): %s",
run_blocking_io(shutil.copy2, str(source_path), str(temp_path))
except (PermissionError, OSError) as e:
# Handle NFS permission errors immediately here
if _is_permission_error(e):
log_transfer_permission_context(
"atomic_copy",
source=source_path,
dest=temp_path,
error=e,
)
logger.debug(
"Permission error during copy, falling back to copyfile (%s -> %s): %s",
source_path,
temp_path,
e,
)
try:
_perform_nfs_fallback(source_path, temp_path, is_move=False)
except Exception as fallback_error:
logger.error(
"NFS fallback also failed (%s -> %s): %s",
source_path,
temp_path,
e,
fallback_error,
)
try:
_perform_nfs_fallback(source_path, temp_path, is_move=False)
except Exception as fallback_error:
logger.error(
"NFS fallback also failed (%s -> %s): %s",
source_path,
temp_path,
fallback_error,
)
raise e from fallback_error
else:
raise
temp_path.replace(try_path)
_verify_transfer_size(try_path, source_path.stat().st_size, "copy")
if attempt > 0:
logger.info(f"File collision resolved: {try_path.name}")
return try_path
raise e from fallback_error
else:
raise
_verify_transfer_size(temp_path, expected_size, "copy")
published = _publish_temp_file(temp_path, try_path)
if not published:
run_blocking_io(temp_path.unlink, missing_ok=True)
continue
try:
_verify_transfer_size(try_path, expected_size, "copy")
except Exception:
try_path.unlink(missing_ok=True)
temp_path.unlink(missing_ok=True)
run_blocking_io(try_path.unlink, missing_ok=True)
raise
except FileExistsError:
continue
if attempt > 0:
logger.info(f"File collision resolved: {try_path.name}")
return try_path
except Exception:
if temp_path:
run_blocking_io(temp_path.unlink, missing_ok=True)
raise
raise RuntimeError(f"Could not copy file after {max_attempts} attempts: {dest_path}")
+94 -16
View File
@@ -5,7 +5,7 @@ import time
from io import BytesIO
from threading import Event, Thread
from typing import Callable, Optional
from urllib.parse import urlparse
from urllib.parse import urlparse, urljoin
import requests
from tqdm import tqdm
@@ -177,13 +177,24 @@ def html_get_page(
cancel_flag: Optional[Event] = None,
status_callback: Optional[Callable[[str, Optional[str]], None]] = None,
allow_bypasser_fallback: bool = True,
) -> str:
include_response_url: bool = False,
success_delay: float = 1.0,
session: Optional[requests.Session] = None,
) -> str | tuple[str, str]:
"""Fetch HTML content from a URL with retry mechanism.
Args:
allow_bypasser_fallback: If False, 403 errors will trigger mirror rotation
instead of switching to the bypasser. Use for search operations.
include_response_url: If True, return `(html, final_url)` to expose the
resolved response URL after redirects.
success_delay: Optional delay (seconds) after successful fetch.
"""
def _result(html: str, response_url: str) -> str | tuple[str, str]:
if include_response_url:
return html, response_url
return html
retry = retry if retry is not None else app_config.MAX_RETRY
selector = selector or network.AAMirrorSelector()
original_url = url
@@ -194,7 +205,7 @@ def html_get_page(
# Check for cancellation before each attempt
if cancel_flag and cancel_flag.is_set():
logger.info(f"html_get_page cancelled before attempt {attempt}")
return ""
return _result("", current_url)
try:
if use_bypasser_now and _is_cf_bypass_enabled():
@@ -218,23 +229,90 @@ def html_get_page(
heartbeat_thread.start()
try:
result = get_bypassed_page(current_url, selector, cancel_flag)
return result or ""
return _result(result or "", current_url)
except Exception as e:
logger.warning(f"Bypasser error: {type(e).__name__}: {e}")
return ""
return _result("", current_url)
finally:
heartbeat_stop.set()
if heartbeat_thread:
heartbeat_thread.join(timeout=1)
logger.debug(f"GET: {current_url}")
# Try with CF cookies/UA if available (from previous bypass)
headers = {}
cookies = _apply_cf_bypass(current_url, headers)
response = requests.get(current_url, proxies=get_proxies(current_url), timeout=REQUEST_TIMEOUT, cookies=cookies, headers=headers)
response.raise_for_status()
time.sleep(1)
return response.text
# Use a browser-like UA by default (AA can behave differently for python-requests UA).
headers = {"User-Agent": DOWNLOAD_HEADERS["User-Agent"]}
# AA mirrors sometimes redirect to other (seized/dead) mirror domains. If we let
# requests follow those redirects, the request fails on DNS and we rotate away
# from an otherwise working mirror. Handle AA redirects manually instead.
is_aa_url = network.should_rotate_dns_for_url(current_url)
allow_redirects = not is_aa_url
redirects_followed = 0
while True:
# Try with CF cookies/UA if available (from previous bypass)
cookies = _apply_cf_bypass(current_url, headers)
request_client = session or requests
response = request_client.get(
current_url,
proxies=get_proxies(current_url),
timeout=REQUEST_TIMEOUT,
cookies=cookies,
headers=headers,
allow_redirects=allow_redirects,
)
if is_aa_url and response.is_redirect:
location = response.headers.get("Location", "")
if not location:
raise requests.exceptions.TooManyRedirects(f"Redirect with no Location header: {current_url}")
redirect_url = urljoin(current_url, location)
current_host = urlparse(current_url).hostname or ""
redirect_host = urlparse(redirect_url).hostname or ""
# If an AA mirror redirects to a different hostname, treat that as a mirror
# failure and rotate rather than following the redirect (auto mode only).
if current_host and redirect_host and current_host != redirect_host:
if not network.is_aa_auto_mode():
logger.warning(
"AA mirror locked to %s but redirected to %s: %s",
current_host,
redirect_host,
current_url,
)
return _result("", current_url)
new_url = _try_rotation(original_url, current_url, selector)
if new_url:
current_url = new_url
# Reset per-request state for the new host.
headers = {"User-Agent": DOWNLOAD_HEADERS["User-Agent"]}
is_aa_url = network.should_rotate_dns_for_url(current_url)
allow_redirects = not is_aa_url
redirects_followed = 0
continue
logger.warning(
"AA redirect from %s to %s but mirrors exhausted: %s",
current_host,
redirect_host,
current_url,
)
return _result("", current_url)
# Same-host redirect (relative or absolute) - follow manually.
redirects_followed += 1
if redirects_followed > 5:
raise requests.exceptions.TooManyRedirects(f"Too many redirects for {current_url}")
current_url = redirect_url
continue
response.raise_for_status()
if success_delay > 0:
time.sleep(success_delay)
return _result(response.text, response.url)
except Exception as e:
status = _get_status_code(e)
@@ -248,7 +326,7 @@ def html_get_page(
current_url = new_url
continue
logger.warning(f"403 error, mirrors exhausted: {current_url}")
return ""
return _result("", current_url)
if _is_cf_bypass_enabled() and not use_bypasser_now:
# Before switching to bypasser, check if cookies have become available
@@ -265,12 +343,12 @@ def html_get_page(
use_bypasser_now = True
continue
logger.warning(f"403 error, giving up: {current_url}")
return ""
return _result("", current_url)
# 404 = Not found
if status == 404:
logger.warning(f"404 error: {current_url}")
return ""
return _result("", current_url)
# Try mirror/DNS rotation on retryable errors
if _is_retryable_error(e):
@@ -286,7 +364,7 @@ def html_get_page(
else:
logger.error(f"Giving up after {retry} attempts: {current_url}")
return ""
return _result("", current_url)
def download_url(
+35 -3
View File
@@ -738,11 +738,22 @@ def rotate_dns_and_reset_aa() -> bool:
if not configured_url:
configured_url = "auto"
if configured_url == "auto" or _aa_base_url in _aa_urls:
if configured_url == "auto":
# Auto mode always resets to the first mirror to restart the cascade
_current_aa_url_index = 0
_aa_base_url = _aa_urls[0] if _aa_urls else "https://annas-archive.se"
_aa_base_url = _aa_urls[0] if _aa_urls else "https://annas-archive.gl"
logger.info(f"After DNS switch, resetting AA URL to: {_aa_base_url}")
_save_state(aa_url=_aa_base_url)
else:
# Keep the user's configured primary mirror (if it exists in the list),
# otherwise keep the configured URL as-is (custom/env).
if configured_url in _aa_urls:
_current_aa_url_index = _aa_urls.index(configured_url)
else:
_current_aa_url_index = 0
_aa_base_url = configured_url
logger.info(f"After DNS switch, keeping configured AA URL: {_aa_base_url}")
_save_state(aa_url=_aa_base_url)
return True
def set_dns_provider(provider: str, manual_servers: list[str] | None = None, use_doh: bool | None = None) -> bool:
@@ -916,6 +927,11 @@ def _initialize_aa_state() -> None:
if not configured_url:
configured_url = "auto"
# 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).
if configured_url != "auto" and configured_url not in _aa_urls:
_aa_urls = [configured_url] + _aa_urls
if configured_url == "auto":
if state.get('aa_base_url') and state['aa_base_url'] in _aa_urls:
_current_aa_url_index = _aa_urls.index(state['aa_base_url'])
@@ -1031,6 +1047,17 @@ def get_aa_base_url():
_ensure_initialized()
return _aa_base_url
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"
def get_available_aa_urls():
"""Get list of configured AA URLs (copy)."""
_ensure_initialized()
@@ -1082,12 +1109,17 @@ class AAMirrorSelector:
Returns (new_base, action) where action is 'mirror', 'dns', or 'exhausted'.
"""
self.attempts_this_dns += 1
if self.attempts_this_dns >= len(self.aa_urls):
max_attempts = len(self.aa_urls) if is_aa_auto_mode() else 1
if self.attempts_this_dns >= max_attempts:
if allow_dns and rotate_dns_and_reset_aa():
self._ensure_fresh_state(reset_attempts=True)
return self.current_base, "dns"
return None, "exhausted"
if not is_aa_auto_mode():
# Mirror is explicitly configured; do not fail over to other mirrors.
return None, "exhausted"
next_index = (self._index + 1) % len(self.aa_urls)
set_aa_url_index(next_index)
self._ensure_fresh_state(reset_attempts=False)
+111 -14
View File
@@ -9,6 +9,7 @@ import random
import threading
import time
from concurrent.futures import Future, ThreadPoolExecutor
from email.utils import parseaddr
from pathlib import Path
from threading import Event, Lock
from typing import Any, Dict, List, Optional, Tuple
@@ -17,7 +18,8 @@ from shelfmark.core.config import config
from shelfmark.core.logger import setup_logger
from shelfmark.core.models import BookInfo, DownloadTask, QueueStatus, SearchFilters, SearchMode
from shelfmark.core.queue import book_queue
from shelfmark.core.utils import transform_cover_url
from shelfmark.core.utils import transform_cover_url, is_audiobook as check_audiobook
from shelfmark.download.fs import run_blocking_io
from shelfmark.download.postprocess.pipeline import is_torrent_source, safe_cleanup_path
from shelfmark.download.postprocess.router import post_process_download
from shelfmark.release_sources import direct_download, get_handler, get_source_display_name
@@ -74,7 +76,35 @@ def get_book_info(book_id: str) -> Optional[Dict[str, Any]]:
logger.error_trace(f"Error getting book info: {e}")
raise
def queue_book(book_id: str, priority: int = 0, source: str = "direct_download") -> Tuple[bool, Optional[str]]:
def _is_plain_email_address(value: str) -> bool:
parsed = parseaddr(value or "")[1]
return bool(parsed) and "@" in parsed and parsed == value
def _resolve_email_destination(
user_id: Optional[int] = None,
) -> Tuple[Optional[str], Optional[str]]:
"""Resolve the destination email address for email output mode.
Returns:
(email_to, error_message)
"""
configured_recipient = str(config.get("EMAIL_RECIPIENT", "", user_id=user_id) or "").strip()
if configured_recipient:
if _is_plain_email_address(configured_recipient):
return configured_recipient, None
return None, "Configured email recipient is invalid"
return None, None
def queue_book(
book_id: str,
priority: int = 0,
source: str = "direct_download",
user_id: Optional[int] = None,
username: Optional[str] = None,
) -> Tuple[bool, Optional[str]]:
"""Add a book to the download queue. Returns (success, error_message)."""
try:
book_info = direct_download.get_book_info(book_id, fetch_download_count=False)
@@ -83,6 +113,22 @@ def queue_book(book_id: str, priority: int = 0, source: str = "direct_download")
logger.warning(error_msg)
return False, error_msg
books_output_mode = str(
config.get("BOOKS_OUTPUT_MODE", "folder", user_id=user_id) or "folder"
).strip().lower()
is_audiobook = check_audiobook(book_info.content)
# Capture output mode at queue time so tasks aren't affected if settings change later.
output_mode = "folder" if is_audiobook else books_output_mode
output_args: Dict[str, Any] = {}
if output_mode == "email" and not is_audiobook:
email_to, email_error = _resolve_email_destination(user_id=user_id)
if email_error:
return False, email_error
if email_to:
output_args = {"to": email_to}
# Create a source-agnostic download task
task = DownloadTask(
task_id=book_id,
@@ -94,7 +140,11 @@ def queue_book(book_id: str, priority: int = 0, source: str = "direct_download")
preview=book_info.preview,
content_type=book_info.content,
search_mode=SearchMode.DIRECT,
output_mode=output_mode,
output_args=output_args,
priority=priority,
user_id=user_id,
username=username,
)
if not book_queue.add(task):
@@ -118,23 +168,57 @@ def queue_book(book_id: str, priority: int = 0, source: str = "direct_download")
return False, error_msg
def queue_release(release_data: dict, priority: int = 0) -> Tuple[bool, Optional[str]]:
def queue_release(
release_data: dict,
priority: int = 0,
user_id: Optional[int] = None,
username: Optional[str] = None,
) -> Tuple[bool, Optional[str]]:
"""Add a release to the download queue. Returns (success, error_message)."""
try:
source = release_data.get('source', 'direct_download')
extra = release_data.get('extra', {})
raw_request_id = release_data.get('_request_id')
request_id: Optional[int] = None
if isinstance(raw_request_id, int) and raw_request_id > 0:
request_id = raw_request_id
# Get author, year, preview, and content_type from top-level (preferred) or extra (fallback)
author = release_data.get('author') or extra.get('author')
year = release_data.get('year') or extra.get('year')
preview = release_data.get('preview') or extra.get('preview')
content_type = release_data.get('content_type') or extra.get('content_type')
source_url_raw = (
release_data.get('download_url')
or release_data.get('source_url')
or release_data.get('info_url')
or extra.get('detail_url')
or extra.get('source_url')
)
source_url = source_url_raw.strip() if isinstance(source_url_raw, str) else None
if source_url == "":
source_url = None
# Get series info for library naming templates
series_name = release_data.get('series_name') or extra.get('series_name')
series_position = release_data.get('series_position') or extra.get('series_position')
subtitle = release_data.get('subtitle') or extra.get('subtitle')
books_output_mode = str(
config.get("BOOKS_OUTPUT_MODE", "folder", user_id=user_id) or "folder"
).strip().lower()
is_audiobook = check_audiobook(content_type)
output_mode = "folder" if is_audiobook else books_output_mode
output_args: Dict[str, Any] = {}
if output_mode == "email" and not is_audiobook:
email_to, email_error = _resolve_email_destination(user_id=user_id)
if email_error:
return False, email_error
if email_to:
output_args = {"to": email_to}
# Create a source-agnostic download task from release data
task = DownloadTask(
task_id=release_data['source_id'],
@@ -146,11 +230,17 @@ def queue_release(release_data: dict, priority: int = 0) -> Tuple[bool, Optional
size=release_data.get('size'),
preview=preview,
content_type=content_type,
source_url=source_url,
series_name=series_name,
series_position=series_position,
subtitle=subtitle,
search_mode=SearchMode.UNIVERSAL,
output_mode=output_mode,
output_args=output_args,
priority=priority,
user_id=user_id,
username=username,
request_id=request_id,
)
if not book_queue.add(task):
@@ -179,12 +269,12 @@ def queue_release(release_data: dict, priority: int = 0) -> Tuple[bool, Optional
logger.error_trace(error_msg)
return False, error_msg
def queue_status() -> Dict[str, Dict[str, Any]]:
def queue_status(user_id: Optional[int] = None) -> Dict[str, Dict[str, Any]]:
"""Get current status of the download queue."""
status = book_queue.get_status()
status = book_queue.get_status(user_id=user_id)
for _, tasks in status.items():
for _, task in tasks.items():
if task.download_path and not os.path.exists(task.download_path):
if task.download_path and not run_blocking_io(os.path.exists, task.download_path):
task.download_path = None
# Convert Enum keys to strings and DownloadTask objects to dicts for JSON serialization
@@ -251,6 +341,9 @@ def _task_to_dict(task: DownloadTask) -> Dict[str, Any]:
'status': task.status,
'status_message': task.status_message,
'download_path': task.download_path,
'user_id': task.user_id,
'username': task.username,
'request_id': task.request_id,
}
@@ -295,7 +388,7 @@ def _download_task(task_id: str, cancel_flag: Event) -> Optional[str]:
return None
temp_file = Path(temp_path)
if not temp_file.exists():
if not run_blocking_io(temp_file.exists):
logger.error(f"Handler returned non-existent path: {temp_path}")
return None
@@ -384,7 +477,9 @@ def update_download_progress(book_id: str, progress: float) -> None:
_progress_last_broadcast[f"{book_id}_progress"] = progress
if should_broadcast:
ws_manager.broadcast_download_progress(book_id, progress, 'downloading')
task = book_queue.get_task(book_id)
task_user_id = task.user_id if task else None
ws_manager.broadcast_download_progress(book_id, progress, 'downloading', user_id=task_user_id)
def update_download_status(book_id: str, status: str, message: Optional[str] = None) -> None:
"""Update download status with optional message for UI display."""
@@ -392,6 +487,7 @@ def update_download_status(book_id: str, status: str, message: Optional[str] = N
status_map = {
'queued': QueueStatus.QUEUED,
'resolving': QueueStatus.RESOLVING,
'locating': QueueStatus.LOCATING,
'downloading': QueueStatus.DOWNLOADING,
'complete': QueueStatus.COMPLETE,
'available': QueueStatus.AVAILABLE,
@@ -414,12 +510,13 @@ def update_download_status(book_id: str, status: str, message: Optional[str] = N
return
_last_status_event[book_id] = status_event
book_queue.update_status(book_id, queue_status_enum)
# Update status message if provided (empty string clears the message)
# Update status message first so terminal snapshots capture the final message
# (for example, "Complete" or "Sent to ...") instead of a stale in-progress one.
if message is not None:
book_queue.update_status_message(book_id, message)
book_queue.update_status(book_id, queue_status_enum)
# Broadcast status update via WebSocket
if ws_manager:
ws_manager.broadcast_status_update(queue_status())
@@ -450,9 +547,9 @@ def get_active_downloads() -> List[str]:
"""Get list of currently active downloads."""
return book_queue.get_active_downloads()
def clear_completed() -> int:
"""Clear all completed downloads from tracking."""
return book_queue.clear_completed()
def clear_completed(user_id: Optional[int] = None) -> int:
"""Clear completed downloads from tracking (optionally user-scoped)."""
return book_queue.clear_completed(user_id=user_id)
def _cleanup_progress_tracking(task_id: str) -> None:
"""Clean up progress tracking data for a completed/cancelled download."""
+41
View File
@@ -49,14 +49,55 @@ def load_output_handlers() -> None:
return
from . import booklore # noqa: F401
from . import email # noqa: F401
from . import folder # noqa: F401
_OUTPUTS_LOADED = True
def _normalize_output_mode(value: object) -> str:
return str(value or "").strip().lower()
def _derive_output_mode(task: DownloadTask) -> str:
"""Return the desired output mode for a task.
Prefer the mode captured at queue time. Fall back to current config for
legacy tasks that do not have `output_mode` populated.
"""
mode = _normalize_output_mode(getattr(task, "output_mode", None))
if mode:
return mode
# Legacy / defensive fallback: derive from current config.
from shelfmark.core.config import config
from shelfmark.core.utils import is_audiobook as check_audiobook
if check_audiobook(getattr(task, "content_type", None)):
return "folder"
return _normalize_output_mode(config.get("BOOKS_OUTPUT_MODE", "folder")) or "folder"
def resolve_output_handler(task: DownloadTask) -> Optional[OutputRegistration]:
load_output_handlers()
desired_mode = _derive_output_mode(task)
# Prefer a direct mode match. `supports_task` becomes a capability check
# (e.g., prevent email/booklore for audiobooks).
for entry in _OUTPUT_REGISTRY:
if entry.mode == desired_mode and entry.supports_task(task):
return entry
# If the requested output isn't supported for this task, fall back to folder.
for entry in _OUTPUT_REGISTRY:
if entry.mode == "folder" and entry.supports_task(task):
return entry
# Last-resort fallback: keep the legacy "first supporting handler" behavior.
for entry in _OUTPUT_REGISTRY:
if entry.supports_task(task):
return entry
return None
+104 -14
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
import os
from dataclasses import dataclass
from pathlib import Path
from threading import Event
@@ -12,12 +13,14 @@ from shelfmark.core.logger import setup_logger
from shelfmark.core.models import DownloadTask
from shelfmark.core.utils import is_audiobook as check_audiobook
from shelfmark.download.outputs import register_output
from shelfmark.download.staging import STAGE_MOVE, STAGE_NONE, build_staging_dir
from shelfmark.download.staging import STAGE_MOVE, STAGE_NONE, build_staging_dir, get_staging_dir
logger = setup_logger(__name__)
BOOKLORE_OUTPUT_MODE = "booklore"
BOOKLORE_SUPPORTED_EXTENSIONS = {".cb7", ".cbr", ".cbz", ".epub", ".fb2", ".pdf"}
BOOKLORE_DESTINATION_LIBRARY = "library"
BOOKLORE_DESTINATION_BOOKDROP = "bookdrop"
BOOKLORE_SUPPORTED_EXTENSIONS = {".azw", ".azw3", ".cb7", ".cbr", ".cbz", ".epub", ".fb2", ".mobi", ".pdf"}
BOOKLORE_SUPPORTED_FORMATS_LABEL = ", ".join(
ext.lstrip(".").upper() for ext in sorted(BOOKLORE_SUPPORTED_EXTENSIONS)
)
@@ -35,6 +38,7 @@ class BookloreConfig:
library_id: int
path_id: int
verify_tls: bool = True
upload_to_bookdrop: bool = False
refresh_after_upload: bool = False
@@ -47,7 +51,17 @@ def _parse_int(value: Any, label: str) -> int:
raise BookloreError(f"{label} must be a number") from exc
def build_booklore_config(values: Mapping[str, Any]) -> BookloreConfig:
def _parse_destination(value: Any) -> str:
normalized = str(value or "").strip().lower()
if normalized == BOOKLORE_DESTINATION_BOOKDROP:
return BOOKLORE_DESTINATION_BOOKDROP
return BOOKLORE_DESTINATION_LIBRARY
def build_booklore_config(
values: Mapping[str, Any],
user_id: Optional[int] = None,
) -> BookloreConfig:
base_url = str(values.get("BOOKLORE_HOST", "")).strip()
username = str(values.get("BOOKLORE_USERNAME", "")).strip()
password = values.get("BOOKLORE_PASSWORD", "") or ""
@@ -59,8 +73,32 @@ def build_booklore_config(values: Mapping[str, Any]) -> BookloreConfig:
if not password:
raise BookloreError("Booklore password is required")
library_id = _parse_int(values.get("BOOKLORE_LIBRARY_ID"), "Booklore library ID")
path_id = _parse_int(values.get("BOOKLORE_PATH_ID"), "Booklore path ID")
destination = _parse_destination(
values.get("BOOKLORE_DESTINATION", BOOKLORE_DESTINATION_LIBRARY)
)
upload_to_bookdrop = destination == BOOKLORE_DESTINATION_BOOKDROP
# Resolve library/path through config so user override precedence is centralized.
library_id = 0
path_id = 0
if not upload_to_bookdrop:
if user_id is not None:
library_id_val = core_config.config.get(
"BOOKLORE_LIBRARY_ID",
values.get("BOOKLORE_LIBRARY_ID"),
user_id=user_id,
)
path_id_val = core_config.config.get(
"BOOKLORE_PATH_ID",
values.get("BOOKLORE_PATH_ID"),
user_id=user_id,
)
else:
library_id_val = values.get("BOOKLORE_LIBRARY_ID")
path_id_val = values.get("BOOKLORE_PATH_ID")
library_id = _parse_int(library_id_val, "Booklore library ID")
path_id = _parse_int(path_id_val, "Booklore path ID")
return BookloreConfig(
base_url=base_url.rstrip("/"),
@@ -69,7 +107,8 @@ def build_booklore_config(values: Mapping[str, Any]) -> BookloreConfig:
library_id=library_id,
path_id=path_id,
verify_tls=True,
refresh_after_upload=True, # Always refresh library after upload
upload_to_bookdrop=upload_to_bookdrop,
refresh_after_upload=not upload_to_bookdrop,
)
@@ -123,9 +162,14 @@ def booklore_list_libraries(booklore_config: BookloreConfig, token: str) -> list
def booklore_upload_file(booklore_config: BookloreConfig, token: str, file_path: Path) -> None:
url = f"{booklore_config.base_url}/api/v1/files/upload"
if booklore_config.upload_to_bookdrop:
url = f"{booklore_config.base_url}/api/v1/files/upload/bookdrop"
params = None
else:
url = f"{booklore_config.base_url}/api/v1/files/upload"
params = {"libraryId": booklore_config.library_id, "pathId": booklore_config.path_id}
headers = {"Authorization": f"Bearer {token}"}
params = {"libraryId": booklore_config.library_id, "pathId": booklore_config.path_id}
response = None
@@ -166,9 +210,7 @@ def booklore_refresh_library(booklore_config: BookloreConfig, token: str) -> Non
def _supports_booklore(task: DownloadTask) -> bool:
if check_audiobook(task.content_type):
return False
return core_config.config.get("BOOKS_OUTPUT_MODE", "folder") == BOOKLORE_OUTPUT_MODE
return not check_audiobook(task.content_type)
def _get_booklore_settings() -> Dict[str, Any]:
@@ -176,6 +218,10 @@ def _get_booklore_settings() -> Dict[str, Any]:
"BOOKLORE_HOST": core_config.config.get("BOOKLORE_HOST", ""),
"BOOKLORE_USERNAME": core_config.config.get("BOOKLORE_USERNAME", ""),
"BOOKLORE_PASSWORD": core_config.config.get("BOOKLORE_PASSWORD", ""),
"BOOKLORE_DESTINATION": core_config.config.get(
"BOOKLORE_DESTINATION",
BOOKLORE_DESTINATION_LIBRARY,
),
"BOOKLORE_LIBRARY_ID": core_config.config.get("BOOKLORE_LIBRARY_ID"),
"BOOKLORE_PATH_ID": core_config.config.get("BOOKLORE_PATH_ID"),
}
@@ -197,9 +243,11 @@ def _post_process_booklore(
status_callback,
) -> Optional[str]:
from shelfmark.download.postprocess.pipeline import (
CustomScriptContext,
OutputPlan,
cleanup_output_staging,
is_managed_workspace_path,
maybe_run_custom_script,
prepare_output_files,
)
@@ -208,7 +256,10 @@ def _post_process_booklore(
return None
try:
booklore_config = build_booklore_config(_get_booklore_settings())
booklore_config = build_booklore_config(
_get_booklore_settings(),
user_id=task.user_id,
)
except BookloreError as e:
logger.warning("Task %s: Booklore configuration error: %s", task.task_id, e)
status_callback("error", str(e))
@@ -216,10 +267,13 @@ def _post_process_booklore(
status_callback("resolving", "Preparing Booklore upload")
stage_action = STAGE_MOVE if is_managed_workspace_path(temp_file) else STAGE_NONE
staging_dir = build_staging_dir("booklore", task.task_id) if stage_action != STAGE_NONE else get_staging_dir()
output_plan = OutputPlan(
mode=BOOKLORE_OUTPUT_MODE,
stage_action=STAGE_MOVE if is_managed_workspace_path(temp_file) else STAGE_NONE,
staging_dir=build_staging_dir("booklore", task.task_id),
stage_action=stage_action,
staging_dir=staging_dir,
allow_archive_extraction=True,
)
@@ -265,6 +319,42 @@ def _post_process_booklore(
logger.info("Task %s: uploaded %d file(s) to Booklore", task.task_id, len(prepared.files))
destination: Optional[Path]
if len(prepared.files) == 1:
destination = prepared.files[0].parent
else:
try:
destination = Path(os.path.commonpath([str(p.parent) for p in prepared.files]))
except ValueError:
destination = prepared.files[0].parent if prepared.files else None
script_context = CustomScriptContext(
task=task,
phase="post_upload",
output_mode=BOOKLORE_OUTPUT_MODE,
destination=destination,
final_paths=prepared.files,
output_details={
"booklore": {
"base_url": booklore_config.base_url,
"destination": (
BOOKLORE_DESTINATION_BOOKDROP
if booklore_config.upload_to_bookdrop
else BOOKLORE_DESTINATION_LIBRARY
),
"library_id": (
None
if booklore_config.upload_to_bookdrop
else booklore_config.library_id
),
"path_id": None if booklore_config.upload_to_bookdrop else booklore_config.path_id,
"refresh_after_upload": bool(booklore_config.refresh_after_upload),
}
},
)
if not maybe_run_custom_script(script_context, status_callback=status_callback):
return None
message = "Uploaded to Booklore"
if len(prepared.files) > 1:
message = f"Uploaded to Booklore ({len(prepared.files)} files)"
+428
View File
@@ -0,0 +1,428 @@
from __future__ import annotations
import mimetypes
import smtplib
import ssl
from dataclasses import dataclass
from email.message import EmailMessage
from email.utils import formatdate, make_msgid, parseaddr
from pathlib import Path
from threading import Event
from typing import Any, Dict, Mapping, Optional
import shelfmark.core.config as core_config
from shelfmark.core.logger import setup_logger
from shelfmark.core.models import DownloadTask
from shelfmark.core.utils import is_audiobook as check_audiobook
from shelfmark.download.outputs import register_output
from shelfmark.download.staging import STAGE_MOVE, STAGE_NONE, build_staging_dir, get_staging_dir
logger = setup_logger(__name__)
EMAIL_OUTPUT_MODE = "email"
SECURITY_NONE = "none"
SECURITY_STARTTLS = "starttls"
SECURITY_SSL = "ssl"
ALLOWED_SECURITY = {SECURITY_NONE, SECURITY_STARTTLS, SECURITY_SSL}
class EmailOutputError(Exception):
"""Raised when the email output integration fails."""
@dataclass(frozen=True)
class EmailSmtpConfig:
host: str
port: int
security: str
username: str = ""
password: str = ""
from_addr: str = ""
timeout_seconds: int = 60
allow_unverified_tls: bool = False
subject_template: str = "{Title}"
def _parse_int(value: Any, label: str, *, minimum: int = 1) -> int:
if value is None or value == "":
raise EmailOutputError(f"{label} is required")
try:
parsed = int(value)
except (TypeError, ValueError) as exc:
raise EmailOutputError(f"{label} must be a number") from exc
if parsed < minimum:
raise EmailOutputError(f"{label} must be >= {minimum}")
return parsed
def build_email_smtp_config(values: Mapping[str, Any]) -> EmailSmtpConfig:
host = str(values.get("EMAIL_SMTP_HOST", "") or "").strip()
port = _parse_int(values.get("EMAIL_SMTP_PORT", 587), "SMTP port", minimum=1)
security = str(values.get("EMAIL_SMTP_SECURITY", SECURITY_STARTTLS) or "").strip().lower()
if security not in ALLOWED_SECURITY:
raise EmailOutputError(f"SMTP security must be one of: {', '.join(sorted(ALLOWED_SECURITY))}")
username = str(values.get("EMAIL_SMTP_USERNAME", "") or "").strip()
password = values.get("EMAIL_SMTP_PASSWORD", "") or ""
from_addr = str(values.get("EMAIL_FROM", "") or "").strip()
subject_template = str(values.get("EMAIL_SUBJECT_TEMPLATE", "{Title}") or "").strip()
timeout_seconds = _parse_int(values.get("EMAIL_SMTP_TIMEOUT_SECONDS", 60), "SMTP timeout (seconds)", minimum=1)
allow_unverified_tls = bool(values.get("EMAIL_ALLOW_UNVERIFIED_TLS", False))
if not host:
raise EmailOutputError("SMTP host is required")
if username and not password:
raise EmailOutputError("SMTP password is required when username is set")
if not from_addr:
# If From is not configured, fall back to the SMTP username if it is an email address.
username_email = parseaddr(username)[1]
if username_email and "@" in username_email:
from_addr = f"Shelfmark <{username_email}>"
else:
raise EmailOutputError("From address is required (or set SMTP username to an email address).")
return EmailSmtpConfig(
host=host,
port=port,
security=security,
username=username,
password=password,
from_addr=from_addr,
timeout_seconds=timeout_seconds,
allow_unverified_tls=allow_unverified_tls,
subject_template=subject_template or "{Title}",
)
def _get_email_settings() -> Dict[str, Any]:
return {
"EMAIL_SMTP_HOST": core_config.config.get("EMAIL_SMTP_HOST", ""),
"EMAIL_SMTP_PORT": core_config.config.get("EMAIL_SMTP_PORT", 587),
"EMAIL_SMTP_SECURITY": core_config.config.get("EMAIL_SMTP_SECURITY", SECURITY_STARTTLS),
"EMAIL_SMTP_USERNAME": core_config.config.get("EMAIL_SMTP_USERNAME", ""),
"EMAIL_SMTP_PASSWORD": core_config.config.get("EMAIL_SMTP_PASSWORD", ""),
"EMAIL_FROM": core_config.config.get("EMAIL_FROM", ""),
"EMAIL_SUBJECT_TEMPLATE": core_config.config.get("EMAIL_SUBJECT_TEMPLATE", "{Title}"),
"EMAIL_SMTP_TIMEOUT_SECONDS": core_config.config.get("EMAIL_SMTP_TIMEOUT_SECONDS", 60),
"EMAIL_ALLOW_UNVERIFIED_TLS": core_config.config.get("EMAIL_ALLOW_UNVERIFIED_TLS", False),
}
def _render_subject(template: str, task: DownloadTask) -> str:
mapping = {
"Author": task.author or "",
"Title": task.title or "",
"Year": task.year or "",
"Series": task.series_name or "",
"SeriesPosition": task.series_position or "",
"Subtitle": task.subtitle or "",
"Format": task.format or "",
}
try:
rendered = template.format(**mapping)
except Exception:
rendered = template
rendered = " ".join(str(rendered).split()).strip()
return rendered or "Shelfmark"
def _msgid_domain(from_addr: str) -> str:
try:
from_email = parseaddr(from_addr)[1]
domain = (from_email.partition("@")[2] or "").strip().rstrip(">")
except Exception:
domain = ""
return domain or "shelfmark.local"
def compose_email_message(
smtp_config: EmailSmtpConfig,
*,
task: DownloadTask,
recipient: str,
files: list[Path],
) -> EmailMessage:
message = EmailMessage()
message["From"] = smtp_config.from_addr
message["To"] = recipient
message["Subject"] = _render_subject(smtp_config.subject_template, task)
message["Date"] = formatdate(localtime=True)
message["Message-ID"] = make_msgid(domain=_msgid_domain(smtp_config.from_addr))
# Keep email body empty; attachments carry the content.
message.set_content("")
for file_path in files:
filename = file_path.name
data = file_path.read_bytes()
content_type, encoding = mimetypes.guess_type(filename)
if content_type is None or encoding is not None:
content_type = "application/octet-stream"
main_type, sub_type = content_type.split("/", 1)
message.add_attachment(data, maintype=main_type, subtype=sub_type, filename=filename)
return message
def _create_tls_context(allow_unverified: bool) -> ssl.SSLContext:
context = ssl.create_default_context()
if allow_unverified:
context.check_hostname = False
context.verify_mode = ssl.CERT_NONE
return context
def test_smtp_connection(smtp_config: EmailSmtpConfig) -> None:
"""Connect and (optionally) authenticate to the SMTP server. Does not send mail."""
smtp: Optional[smtplib.SMTP] = None
try:
if smtp_config.security == SECURITY_SSL:
context = _create_tls_context(smtp_config.allow_unverified_tls)
smtp = smtplib.SMTP_SSL(
smtp_config.host,
smtp_config.port,
timeout=smtp_config.timeout_seconds,
context=context,
)
else:
smtp = smtplib.SMTP(smtp_config.host, smtp_config.port, timeout=smtp_config.timeout_seconds)
smtp.ehlo()
if smtp_config.security == SECURITY_STARTTLS:
context = _create_tls_context(smtp_config.allow_unverified_tls)
smtp.starttls(context=context)
smtp.ehlo()
if smtp_config.username:
smtp.login(smtp_config.username, smtp_config.password)
except smtplib.SMTPAuthenticationError as exc:
raise EmailOutputError("SMTP authentication failed") from exc
except (smtplib.SMTPConnectError, smtplib.SMTPServerDisconnected, TimeoutError, OSError) as exc:
raise EmailOutputError(f"Could not connect to SMTP server: {exc}") from exc
finally:
if smtp is not None:
try:
smtp.quit()
except Exception:
try:
smtp.close()
except Exception:
pass
def send_email_message(smtp_config: EmailSmtpConfig, message: EmailMessage) -> None:
smtp: Optional[smtplib.SMTP] = None
try:
if smtp_config.security == SECURITY_SSL:
context = _create_tls_context(smtp_config.allow_unverified_tls)
smtp = smtplib.SMTP_SSL(
smtp_config.host,
smtp_config.port,
timeout=smtp_config.timeout_seconds,
context=context,
)
else:
smtp = smtplib.SMTP(smtp_config.host, smtp_config.port, timeout=smtp_config.timeout_seconds)
smtp.ehlo()
if smtp_config.security == SECURITY_STARTTLS:
context = _create_tls_context(smtp_config.allow_unverified_tls)
smtp.starttls(context=context)
smtp.ehlo()
if smtp_config.username:
smtp.login(smtp_config.username, smtp_config.password)
smtp.send_message(message)
except smtplib.SMTPAuthenticationError as exc:
raise EmailOutputError("SMTP authentication failed") from exc
except (smtplib.SMTPException, TimeoutError, OSError) as exc:
raise EmailOutputError(f"Failed to send email: {exc}") from exc
finally:
if smtp is not None:
try:
smtp.quit()
except Exception:
try:
smtp.close()
except Exception:
pass
def _supports_email(task: DownloadTask) -> bool:
return not check_audiobook(task.content_type)
def _post_process_email(
temp_file: Path,
task: DownloadTask,
cancel_flag: Event,
status_callback,
) -> Optional[str]:
from shelfmark.download.postprocess.pipeline import (
CustomScriptContext,
OutputPlan,
cleanup_output_staging,
is_managed_workspace_path,
maybe_run_custom_script,
prepare_output_files,
)
if cancel_flag.is_set():
logger.info("Task %s: cancelled before email send", task.task_id)
return None
try:
smtp_config = build_email_smtp_config(_get_email_settings())
except EmailOutputError as exc:
logger.warning("Task %s: email configuration error: %s", task.task_id, exc)
status_callback("error", str(exc))
return None
output_args = task.output_args or {}
if not isinstance(output_args, dict):
output_args = {}
recipient = str(output_args.get("to", "") or "").strip()
label = str(output_args.get("label", "") or "").strip() or recipient
if not recipient:
status_callback(
"error",
"No email recipient configured. Set a per-user email recipient or a default in Downloads -> Books.",
)
return None
status_callback("resolving", "Preparing email")
stage_action = STAGE_MOVE if is_managed_workspace_path(temp_file) else STAGE_NONE
staging_dir = build_staging_dir("email", task.task_id) if stage_action != STAGE_NONE else get_staging_dir()
output_plan = OutputPlan(
mode=EMAIL_OUTPUT_MODE,
stage_action=stage_action,
staging_dir=staging_dir,
allow_archive_extraction=True,
)
prepared = prepare_output_files(
temp_file,
task,
EMAIL_OUTPUT_MODE,
status_callback,
output_plan=output_plan,
)
if not prepared:
return None
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
if attachment_limit_mb > 0:
limit_bytes = attachment_limit_mb * 1024 * 1024
file_sizes: list[tuple[Path, int]] = []
total_bytes = 0
for file_path in prepared.files:
try:
size_bytes = file_path.stat().st_size
except OSError:
continue
file_sizes.append((file_path, size_bytes))
total_bytes += size_bytes
too_large = [(path, size) for path, size in file_sizes if size > limit_bytes]
if too_large:
path, size = max(too_large, key=lambda item: item[1])
status_callback(
"error",
f"Attachment '{path.name}' is {size / (1024 * 1024):.1f} MB (limit {attachment_limit_mb} MB)",
)
return None
# Most providers enforce a message size limit and attachments are base64-encoded (~33% overhead).
estimated_encoded_bytes = int(total_bytes * 4 / 3)
if estimated_encoded_bytes > limit_bytes:
status_callback(
"error",
(
f"Attachments total {total_bytes / (1024 * 1024):.1f} MB "
f"(estimated encoded {estimated_encoded_bytes / (1024 * 1024):.1f} MB) "
f"exceeds limit {attachment_limit_mb} MB"
),
)
return None
if cancel_flag.is_set():
logger.info("Task %s: cancelled before email send", task.task_id)
return None
status_callback("resolving", f"Sending email to {label}")
message = compose_email_message(
smtp_config,
task=task,
recipient=recipient,
files=prepared.files,
)
send_email_message(smtp_config, message)
script_context = CustomScriptContext(
task=task,
phase="post_email",
output_mode=EMAIL_OUTPUT_MODE,
destination=prepared.files[0].parent if prepared.files else None,
final_paths=prepared.files,
output_details={
"email": {
"to": recipient,
"label": label,
"host": smtp_config.host,
"port": smtp_config.port,
"security": smtp_config.security,
}
},
)
if not maybe_run_custom_script(script_context, status_callback=status_callback):
return None
status_callback("complete", f"Sent to {label}")
return f"email://{task.task_id}"
except EmailOutputError as exc:
logger.warning("Task %s: email send failed: %s", task.task_id, exc)
status_callback("error", str(exc))
return None
except Exception as exc:
logger.error_trace("Task %s: unexpected error sending email: %s", task.task_id, exc)
status_callback("error", f"Email send failed: {exc}")
return None
finally:
cleanup_output_staging(
prepared.output_plan,
prepared.working_path,
task,
prepared.cleanup_paths,
)
@register_output(EMAIL_OUTPUT_MODE, supports_task=_supports_email, priority=10)
def process_email_output(
temp_file: Path,
task: DownloadTask,
cancel_flag: Event,
status_callback,
) -> Optional[str]:
return _post_process_email(temp_file, task, cancel_flag, status_callback)
+41 -103
View File
@@ -1,7 +1,6 @@
from __future__ import annotations
import os
import subprocess
from dataclasses import dataclass
from pathlib import Path
from threading import Event
@@ -11,7 +10,6 @@ import shelfmark.core.config as core_config
from shelfmark.core.logger import setup_logger
from shelfmark.core.models import DownloadTask
from shelfmark.core.utils import is_audiobook as check_audiobook
from shelfmark.download.archive import is_archive
from shelfmark.download.outputs import register_output
from shelfmark.download.staging import StageAction, STAGE_NONE
@@ -20,17 +18,9 @@ logger = setup_logger(__name__)
FOLDER_OUTPUT_MODE = "folder"
def _resolve_custom_script_target(target_path: Path, destination: Path, path_mode: str) -> Path:
mode = (path_mode or "absolute").strip().lower()
if mode != "relative":
return target_path
try:
return target_path.relative_to(destination)
except ValueError:
if target_path.is_absolute():
return Path(target_path.name)
return target_path
def _format_op_counts(op_counts: dict[str, int]) -> str:
parts = [f"{op}={count}" for op, count in op_counts.items() if count]
return ", ".join(parts) if parts else "none"
@dataclass(frozen=True)
@@ -46,9 +36,7 @@ class _ProcessingPlan:
def _supports_folder_output(task: DownloadTask) -> bool:
if check_audiobook(task.content_type):
return True
return core_config.config.get("BOOKS_OUTPUT_MODE", FOLDER_OUTPUT_MODE) == FOLDER_OUTPUT_MODE
return True
def _build_processing_plan(
@@ -103,12 +91,14 @@ def process_folder_output(
) -> Optional[str]:
"""Post-process download to the configured folder destination."""
from shelfmark.download.postprocess.pipeline import (
CustomScriptContext,
CustomScriptTransferSummary,
cleanup_output_staging,
is_torrent_source,
log_plan_steps,
prepare_output_files,
maybe_run_custom_script,
record_step,
safe_cleanup_path,
transfer_book_files,
)
@@ -141,73 +131,9 @@ def process_folder_output(
step_name = f"stage_{prepared.output_plan.stage_action}"
record_step(steps, step_name, source=str(temp_file), dest=str(prepared.output_plan.staging_dir))
def run_custom_script(script_path: str, target_path: Path, phase: str) -> bool:
path_mode = core_config.config.get("CUSTOM_SCRIPT_PATH_MODE", "absolute")
script_target = _resolve_custom_script_target(target_path, plan.destination, path_mode)
env = {
**os.environ,
"SHELFMARK_CUSTOM_SCRIPT_TARGET": str(target_path),
"SHELFMARK_CUSTOM_SCRIPT_RELATIVE": str(_resolve_custom_script_target(target_path, plan.destination, "relative")),
"SHELFMARK_CUSTOM_SCRIPT_DESTINATION": str(plan.destination),
"SHELFMARK_CUSTOM_SCRIPT_MODE": str(path_mode),
"SHELFMARK_CUSTOM_SCRIPT_PHASE": phase,
}
record_step(
steps,
"custom_script",
script=str(script_path),
target=str(script_target),
target_abs=str(target_path),
mode=str(path_mode),
phase=phase,
)
log_plan_steps(task.task_id, steps)
logger.info(
"Task %s: running custom script %s on %s (%s)",
task.task_id,
script_path,
script_target,
phase,
)
try:
result = subprocess.run(
[script_path, str(script_target)],
check=True,
timeout=300, # 5 minute timeout
capture_output=True,
text=True,
env=env,
)
if result.stdout:
logger.debug("Task %s: custom script stdout: %s", task.task_id, result.stdout.strip())
return True
except FileNotFoundError:
logger.error("Task %s: custom script not found: %s", task.task_id, script_path)
status_callback("error", f"Custom script not found: {script_path}")
return False
except PermissionError:
logger.error("Task %s: custom script not executable: %s", task.task_id, script_path)
status_callback("error", f"Custom script not executable: {script_path}")
return False
except subprocess.TimeoutExpired:
logger.error("Task %s: custom script timed out after 300s: %s", task.task_id, script_path)
status_callback("error", "Custom script timed out")
return False
except subprocess.CalledProcessError as e:
stderr = e.stderr.strip() if e.stderr else "No error output"
logger.error(
"Task %s: custom script failed (exit code %s): %s",
task.task_id,
e.returncode,
stderr,
)
status_callback("error", f"Custom script failed: {stderr[:100]}")
return False
# Custom script is run post-transfer (see below).
# If we staged a copy into TMP_DIR (e.g. for custom script), transfer from the staged
# path and disable hardlinking for this transfer.
# If we staged into TMP_DIR, transfer from the staged path and disable hardlinking.
use_hardlink = plan.use_hardlink and prepared.output_plan.stage_action == STAGE_NONE
source_path = plan.hardlink_source if use_hardlink and plan.hardlink_source else prepared.working_path
is_torrent = is_torrent_source(source_path, task)
@@ -255,7 +181,7 @@ def process_folder_output(
record_step(steps, "cleanup_staging", path=str(prepared.working_path))
log_plan_steps(task.task_id, steps)
final_paths, error = transfer_book_files(
final_paths, error, op_counts = transfer_book_files(
prepared.files,
destination=plan.destination,
task=task,
@@ -271,31 +197,43 @@ def process_folder_output(
return None
logger.info(
"Task %s: transferred %d file(s) to %s (%s)",
"Task %s: transferred %d file(s) to %s (ops: %s)",
task.task_id,
len(final_paths),
plan.destination,
op_label.lower(),
_format_op_counts(op_counts),
)
if use_hardlink and op_counts.get("copy", 0):
logger.warning(
"Task %s: hardlink requested but %d of %d file(s) copied (fallback)",
task.task_id,
op_counts.get("copy", 0),
len(final_paths),
)
script_context = CustomScriptContext(
task=task,
phase="post_transfer",
output_mode=plan.output_mode,
organization_mode=plan.organization_mode,
destination=plan.destination,
final_paths=final_paths,
transfer=CustomScriptTransferSummary(
op_counts=op_counts,
use_hardlink=use_hardlink,
is_torrent=is_torrent,
preserve_source=preserve_source,
),
)
# Run custom script once per successful task, after transfer.
if core_config.config.CUSTOM_SCRIPT:
if len(final_paths) == 1:
target_path = final_paths[0]
else:
try:
target_path = Path(os.path.commonpath([str(p.parent) for p in final_paths]))
except ValueError:
target_path = plan.destination
if not run_custom_script(core_config.config.CUSTOM_SCRIPT, target_path, phase="post_transfer"):
cleanup_output_staging(
prepared.output_plan,
prepared.working_path,
task,
prepared.cleanup_paths,
)
return None
if not maybe_run_custom_script(script_context, status_callback=status_callback, steps=steps):
cleanup_output_staging(
prepared.output_plan,
prepared.working_path,
task,
prepared.cleanup_paths,
)
return None
cleanup_output_staging(
prepared.output_plan,
+24 -7
View File
@@ -16,6 +16,23 @@ from shelfmark.core.logger import setup_logger
logger = setup_logger(__name__)
def _run_io(func, *args, **kwargs):
"""Best-effort offload for potentially blocking filesystem calls.
Keep this module import-cycle safe: `shelfmark.download.fs` imports this module,
so we only import `run_blocking_io` lazily at call-time.
"""
try:
from shelfmark.download.fs import run_blocking_io as _run_blocking_io
except Exception:
return func(*args, **kwargs)
try:
return _run_blocking_io(func, *args, **kwargs)
except Exception:
# Fall back to direct call if threadpool offload is unavailable.
return func(*args, **kwargs)
def _format_uid(uid: int) -> str:
try:
@@ -59,12 +76,12 @@ def log_path_permission_context(label: str, path: Path) -> None:
for probe in [path, path.parent]:
try:
resolved = probe.resolve()
resolved = _run_io(probe.resolve)
except Exception:
resolved = probe
try:
st = probe.stat()
st = _run_io(probe.stat)
logger.debug(
"Path permissions (%s): path=%s resolved=%s mode=%s owner=%s(%d) group=%s(%d) dir=%s symlink=%s",
label,
@@ -75,8 +92,8 @@ def log_path_permission_context(label: str, path: Path) -> None:
st.st_uid,
_format_gid(st.st_gid),
st.st_gid,
probe.is_dir(),
probe.is_symlink(),
_run_io(probe.is_dir),
_run_io(probe.is_symlink),
)
except Exception as stat_error:
logger.debug("Path permissions (%s): stat failed for %s: %s", label, probe, stat_error)
@@ -106,7 +123,7 @@ def log_transfer_permission_context(label: str, source: Path, dest: Path, error:
for probe in [source, dest, dest.parent]:
try:
st = probe.stat()
st = _run_io(probe.stat)
logger.debug(
"Path permissions (%s): path=%s mode=%s owner=%s(%d) group=%s(%d) exists=%s dir=%s",
label,
@@ -116,8 +133,8 @@ def log_transfer_permission_context(label: str, source: Path, dest: Path, error:
st.st_uid,
_format_gid(st.st_gid),
st.st_gid,
probe.exists(),
probe.is_dir(),
_run_io(probe.exists),
_run_io(probe.is_dir),
)
except Exception as stat_error:
logger.debug("Path permissions (%s): stat failed for %s: %s", label, probe, stat_error)
@@ -0,0 +1,296 @@
from __future__ import annotations
import json
import os
import subprocess
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Optional
import shelfmark.core.config as core_config
from shelfmark.core.logger import setup_logger
from shelfmark.core.models import DownloadTask
from shelfmark.download.fs import run_blocking_io
from .steps import log_plan_steps, record_step
from .types import PlanStep
logger = setup_logger(__name__)
DEFAULT_CUSTOM_SCRIPT_TIMEOUT_SECONDS = 300 # 5 minutes
def resolve_custom_script_target(target_path: Path, destination: Path, path_mode: str) -> Path:
"""Resolve the path that should be passed as the custom script argument.
In absolute mode, we pass the full target path.
In relative mode, we pass a path relative to the destination folder. If the
target is not within the destination, fall back to just the filename to
avoid leaking unrelated absolute paths.
"""
mode = (path_mode or "absolute").strip().lower()
if mode != "relative":
return target_path
try:
return target_path.relative_to(destination)
except ValueError:
if target_path.is_absolute():
return Path(target_path.name)
return target_path
@dataclass(frozen=True)
class CustomScriptExecution:
script_path: str
target_arg: Path
target_abs: Path
destination: Path
mode: str
phase: str
payload_json: Optional[str] = None
@dataclass(frozen=True)
class CustomScriptTransferSummary:
op_counts: dict[str, int]
use_hardlink: bool
is_torrent: bool
preserve_source: bool
@dataclass(frozen=True)
class CustomScriptContext:
task: DownloadTask
phase: str
output_mode: str
destination: Optional[Path] = None
final_paths: list[Path] = field(default_factory=list)
target_path: Optional[Path] = None
organization_mode: Optional[str] = None
transfer: Optional[CustomScriptTransferSummary] = None
output_details: dict[str, Any] = field(default_factory=dict)
def prepare_custom_script_execution(
script_path: str,
*,
target_path: Path,
destination: Path,
path_mode: str,
phase: str,
payload: Optional[dict[str, Any]] = None,
) -> CustomScriptExecution:
mode = (path_mode or "absolute").strip().lower()
if mode != "relative":
mode = "absolute"
target_arg = resolve_custom_script_target(target_path, destination, mode)
return CustomScriptExecution(
script_path=str(script_path),
target_arg=target_arg,
target_abs=target_path,
destination=destination,
mode=mode,
phase=phase,
payload_json=json.dumps(payload, indent=2, sort_keys=True) + "\n" if payload else None,
)
def run_custom_script(
execution: CustomScriptExecution,
*,
task_id: str,
status_callback,
timeout_seconds: int = DEFAULT_CUSTOM_SCRIPT_TIMEOUT_SECONDS,
) -> bool:
cwd: Optional[str] = None
if execution.mode == "relative":
# Make relative paths unambiguous by running the script from the destination folder.
cwd = str(execution.destination)
logger.info(
"Task %s: running custom script %s on %s (%s)",
task_id,
execution.script_path,
execution.target_arg,
execution.phase,
)
try:
# If we are not sending a JSON payload, close stdin so scripts that try
# to read it won't block indefinitely.
stdin = None if execution.payload_json is not None else subprocess.DEVNULL
result = run_blocking_io(
subprocess.run,
[execution.script_path, str(execution.target_arg)],
check=True,
timeout=timeout_seconds,
capture_output=True,
text=True,
cwd=cwd,
stdin=stdin,
input=execution.payload_json,
)
if result.stdout:
logger.debug("Task %s: custom script stdout: %s", task_id, result.stdout.strip())
return True
except FileNotFoundError:
logger.error("Task %s: custom script not found: %s", task_id, execution.script_path)
status_callback("error", f"Custom script not found: {execution.script_path}")
return False
except PermissionError:
logger.error("Task %s: custom script not executable: %s", task_id, execution.script_path)
status_callback("error", f"Custom script not executable: {execution.script_path}")
return False
except subprocess.TimeoutExpired:
logger.error(
"Task %s: custom script timed out after %ss: %s",
task_id,
timeout_seconds,
execution.script_path,
)
status_callback("error", "Custom script timed out")
return False
except subprocess.CalledProcessError as exc:
stderr = exc.stderr.strip() if exc.stderr else "No error output"
logger.error(
"Task %s: custom script failed (exit code %s): %s",
task_id,
exc.returncode,
stderr,
)
status_callback("error", f"Custom script failed: {stderr[:100]}")
return False
def _choose_custom_script_target(
*,
explicit_target: Optional[Path],
destination: Optional[Path],
final_paths: list[Path],
) -> Optional[Path]:
if explicit_target is not None:
return explicit_target
if len(final_paths) == 1:
return final_paths[0]
if len(final_paths) > 1:
try:
return Path(os.path.commonpath([str(p.parent) for p in final_paths]))
except ValueError:
return destination or final_paths[0].parent
return destination
def _build_custom_script_payload(context: CustomScriptContext, *, target_path: Path) -> dict[str, Any]:
payload: dict[str, Any] = {
"version": 1,
"phase": context.phase,
"task": {
"task_id": context.task.task_id,
"source": context.task.source,
"search_mode": context.task.search_mode.value if context.task.search_mode else None,
"title": context.task.title,
"author": context.task.author,
"year": context.task.year,
"format": context.task.format,
"content_type": context.task.content_type,
"series_name": context.task.series_name,
"series_position": context.task.series_position,
"subtitle": context.task.subtitle,
"original_download_path": context.task.original_download_path,
},
"output": {
"mode": context.output_mode,
"organization_mode": context.organization_mode,
},
"paths": {
"destination": str(context.destination) if context.destination else None,
"target": str(target_path),
"final_paths": [str(p) for p in context.final_paths],
},
}
if context.output_details:
payload["output"]["details"] = context.output_details
if context.transfer:
payload["transfer"] = {
"op_counts": context.transfer.op_counts,
"use_hardlink": context.transfer.use_hardlink,
"is_torrent": context.transfer.is_torrent,
"preserve_source": context.transfer.preserve_source,
}
return payload
def maybe_run_custom_script(
context: CustomScriptContext,
*,
status_callback,
steps: Optional[list[PlanStep]] = None,
) -> bool:
"""Run the custom script hook (if configured).
The output handler provides a `CustomScriptContext` describing what it did.
This function is responsible for choosing the script target, building the
optional JSON payload, and executing the script.
"""
script_path = getattr(core_config.config, "CUSTOM_SCRIPT", None)
if not isinstance(script_path, str) or not script_path.strip():
return True
target_path = _choose_custom_script_target(
explicit_target=context.target_path,
destination=context.destination,
final_paths=context.final_paths,
)
if not target_path:
logger.warning(
"Task %s: custom script configured but no target could be determined; skipping",
context.task.task_id,
)
return True
path_mode = core_config.config.get("CUSTOM_SCRIPT_PATH_MODE", "absolute")
payload: Optional[dict[str, Any]] = None
if core_config.config.get("CUSTOM_SCRIPT_JSON_PAYLOAD", False):
payload = _build_custom_script_payload(context, target_path=target_path)
# If no destination is available for this output, fall back to the target's
# parent directory so the script can still run consistently.
execution_destination = context.destination or target_path.parent
execution = prepare_custom_script_execution(
script_path,
target_path=target_path,
destination=execution_destination,
path_mode=path_mode,
phase=context.phase,
payload=payload,
)
if steps is not None:
payload_bytes = len(execution.payload_json.encode("utf-8")) if execution.payload_json else 0
record_step(
steps,
"custom_script",
script=str(execution.script_path),
target=str(execution.target_arg),
target_abs=str(execution.target_abs),
mode=str(execution.mode),
phase=str(execution.phase),
payload_stdin=bool(execution.payload_json),
payload_bytes=payload_bytes,
)
log_plan_steps(context.task.task_id, steps)
return run_custom_script(execution, task_id=context.task.task_id, status_callback=status_callback)
@@ -10,6 +10,7 @@ from shelfmark.core.utils import (
get_destination,
is_audiobook as check_audiobook,
)
from shelfmark.download.fs import run_blocking_io
from shelfmark.download.permissions_debug import log_path_permission_context
logger = setup_logger("shelfmark.download.postprocess.pipeline")
@@ -23,14 +24,15 @@ def validate_destination(destination: Path, status_callback) -> bool:
status_callback("error", f"Destination must be absolute: {destination}")
return False
if destination.exists() and not destination.is_dir():
destination_exists = run_blocking_io(destination.exists)
if destination_exists and not run_blocking_io(destination.is_dir):
logger.warning(f"Destination is not a directory: {destination}")
status_callback("error", f"Destination is not a directory: {destination}")
return False
if not destination.exists():
if not destination_exists:
try:
destination.mkdir(parents=True, exist_ok=True)
run_blocking_io(destination.mkdir, parents=True, exist_ok=True)
except (OSError, PermissionError) as exc:
log_path_permission_context("destination_create", destination)
logger.warning(f"Cannot create destination: {destination} ({exc})")
@@ -44,8 +46,8 @@ def validate_destination(destination: Path, status_callback) -> bool:
f"This file was created to verify if '{destination}' is writable. "
"It should've been automatically deleted. Feel free to delete it.\n"
)
test_path.write_text(test_content)
test_path.unlink(missing_ok=True)
run_blocking_io(test_path.write_text, test_content)
run_blocking_io(test_path.unlink, missing_ok=True)
except Exception as exc:
logger.debug("Destination write probe path: %s", test_path)
log_path_permission_context("destination_write_probe", destination)
@@ -66,4 +68,4 @@ def get_final_destination(task: DownloadTask) -> Path:
if override:
return override
return get_destination(is_audiobook)
return get_destination(is_audiobook, user_id=task.user_id, username=task.username)
@@ -17,6 +17,15 @@ implementation stay modular.
from __future__ import annotations
from .custom_script import (
CustomScriptExecution,
CustomScriptContext,
CustomScriptTransferSummary,
maybe_run_custom_script,
prepare_custom_script_execution,
resolve_custom_script_target,
run_custom_script,
)
from .destination import get_final_destination, validate_destination
from .prepare import build_output_plan, prepare_output_files
from .scan import (
@@ -50,6 +59,9 @@ __all__ = [
"PlanStep",
"PreparedFiles",
"TransferPlan",
"CustomScriptExecution",
"CustomScriptContext",
"CustomScriptTransferSummary",
"build_metadata_dict",
"build_output_plan",
"cleanup_output_staging",
@@ -62,10 +74,13 @@ __all__ = [
"is_torrent_source",
"is_within_tmp_dir",
"log_plan_steps",
"maybe_run_custom_script",
"prepare_output_files",
"prepare_custom_script_execution",
"process_directory",
"record_step",
"resolve_hardlink_source",
"resolve_custom_script_target",
"safe_cleanup_path",
"scan_directory_tree",
"should_hardlink",
@@ -73,4 +88,5 @@ __all__ = [
"transfer_directory_to_library",
"transfer_file_to_library",
"validate_destination",
"run_custom_script",
]
+1 -6
View File
@@ -3,10 +3,8 @@ from __future__ import annotations
from pathlib import Path
from typing import Optional
import shelfmark.core.config as core_config
from shelfmark.core.logger import setup_logger
from shelfmark.core.models import DownloadTask
from shelfmark.download.archive import is_archive
from shelfmark.download.staging import STAGE_COPY, STAGE_NONE, get_staging_dir, stage_path
from .scan import collect_staged_files
@@ -27,14 +25,11 @@ def build_output_plan(
"""Build an output plan that describes staging behavior for file-based outputs."""
transfer_plan = resolve_hardlink_source(temp_file, task, destination, status_callback)
runs_custom_script = bool(core_config.config.CUSTOM_SCRIPT) and temp_file.is_file() and not is_archive(temp_file)
stage_action = STAGE_COPY if runs_custom_script and not is_managed_workspace_path(temp_file) else STAGE_NONE
staging_dir = get_staging_dir()
return OutputPlan(
mode=output_mode,
stage_action=stage_action,
stage_action=STAGE_NONE,
staging_dir=staging_dir,
allow_archive_extraction=transfer_plan.allow_archive_extraction,
transfer_plan=transfer_plan,
+53 -22
View File
@@ -8,6 +8,7 @@ from shelfmark.core.logger import setup_logger
from shelfmark.core.models import DownloadTask
from shelfmark.core.utils import is_audiobook as check_audiobook
from shelfmark.download.archive import ArchiveExtractionError, extract_archive, is_archive
from shelfmark.download.fs import run_blocking_io
from shelfmark.download.permissions_debug import log_path_permission_context
from shelfmark.download.postprocess.policy import (
get_supported_audiobook_formats,
@@ -55,7 +56,12 @@ def extract_archive_files(
content_type = task.content_type
try:
extracted_files, warnings, rejected_files = extract_archive(archive_path, output_dir, content_type)
extracted_files, warnings, rejected_files = run_blocking_io(
extract_archive,
archive_path,
output_dir,
content_type,
)
except ArchiveExtractionError as exc:
logger.warning(
"Task %s: archive extraction failed for %s: %s",
@@ -74,7 +80,7 @@ def extract_archive_files(
)
if cleanup_archive:
archive_path.unlink(missing_ok=True)
run_blocking_io(archive_path.unlink, missing_ok=True)
cleanup_paths = [output_dir]
@@ -101,8 +107,12 @@ def scan_directory_tree(
"""Scan a directory tree for book files, trackable-but-unsupported files, and archives."""
try:
with os.scandir(directory) as it:
next(it, None)
def _probe_dir() -> None:
# Force a fast error if the dir is missing/inaccessible.
with os.scandir(directory) as it:
next(it, None)
run_blocking_io(_probe_dir)
except PermissionError as exc:
log_path_permission_context("scan_directory", directory)
logger.warning(f"Permission denied scanning directory: {directory} ({exc})")
@@ -111,10 +121,6 @@ def scan_directory_tree(
logger.warning(f"Cannot access download folder: {directory} ({exc})")
return [], [], [], f"Cannot access download folder: {directory} ({exc})"
book_files: List[Path] = []
rejected_files: List[Path] = []
archive_files: List[Path] = []
supported_formats = get_supported_formats(content_type)
supported_exts = {f".{fmt}" for fmt in supported_formats}
@@ -146,18 +152,35 @@ def scan_directory_tree(
else:
logger.debug(f"Error scanning directory tree: {error}")
for root, _, files in os.walk(directory, onerror=onerror):
for filename in files:
file_path = Path(root) / filename
suffix = file_path.suffix.lower()
def _walk_tree() -> Tuple[List[Path], List[Path], List[Path]]:
book_files: List[Path] = []
rejected_files: List[Path] = []
archive_files: List[Path] = []
if suffix in supported_exts:
book_files.append(file_path)
elif suffix in trackable_exts:
rejected_files.append(file_path)
for root, _, files in os.walk(directory, onerror=onerror):
for filename in files:
file_path = Path(root) / filename
suffix = file_path.suffix.lower()
if is_archive(file_path):
archive_files.append(file_path)
if suffix in supported_exts:
book_files.append(file_path)
elif suffix in trackable_exts:
rejected_files.append(file_path)
if is_archive(file_path):
archive_files.append(file_path)
return book_files, rejected_files, archive_files
try:
book_files, rejected_files, archive_files = run_blocking_io(_walk_tree)
except PermissionError as exc:
log_path_permission_context("scan_directory_walk", directory)
logger.warning(f"Permission denied scanning directory: {directory} ({exc})")
return [], [], [], f"Permission denied accessing download folder: {directory}"
except (FileNotFoundError, NotADirectoryError, OSError) as exc:
logger.warning(f"Cannot access download folder: {directory} ({exc})")
return [], [], [], f"Cannot access download folder: {directory} ({exc})"
return book_files, rejected_files, archive_files, None
@@ -194,12 +217,15 @@ def collect_directory_files(
if archive_files:
if not allow_archive_extraction:
logger.warning(
"Task %s: archive extraction disabled (torrent hardlinking enabled) for %s",
# When extraction is disabled (typically due to torrent hardlinking),
# treat archives as the final importable "files" rather than failing.
logger.info(
"Task %s: archive extraction disabled; importing %d archive(s) as-is from %s",
task.task_id,
len(archive_files),
directory,
)
return [], rejected_files, [], "Archive extraction is disabled when torrent hardlinking is enabled"
return archive_files, rejected_files, [], None
if status_callback:
status_callback("resolving", "Extracting archives")
@@ -258,7 +284,7 @@ def collect_staged_files(
status_callback,
cleanup_archives: bool,
) -> Tuple[List[Path], List[Path], List[Path], Optional[str]]:
if working_path.is_dir():
if run_blocking_io(working_path.is_dir):
if status_callback:
status_callback("resolving", "Processing download folder")
return collect_directory_files(
@@ -293,6 +319,11 @@ def collect_staged_files(
return extracted_files, rejected_files, cleanup_paths, error
if is_archive(working_path) and not allow_archive_extraction:
# When extraction is disabled (typically due to torrent hardlinking),
# import the archive as-is rather than treating it as an unsupported file.
return [working_path], [], [], None
# Single-file download result (non-archive).
# Ensure we respect the user's supported format settings.
suffix = working_path.suffix.lower()
+69 -28
View File
@@ -2,7 +2,7 @@ from __future__ import annotations
import os
from pathlib import Path
from typing import List, Optional, Tuple
from typing import Dict, List, Optional, Tuple
import shelfmark.core.config as core_config
from shelfmark.core.logger import setup_logger
@@ -15,7 +15,7 @@ from shelfmark.core.naming import (
sanitize_filename,
)
from shelfmark.core.utils import is_audiobook as check_audiobook
from shelfmark.download.fs import atomic_copy, atomic_hardlink, atomic_move
from shelfmark.download.fs import atomic_copy, atomic_hardlink, atomic_move, run_blocking_io
from shelfmark.download.postprocess.policy import get_file_organization, get_template
from .scan import collect_directory_files, scan_directory_tree
@@ -53,6 +53,7 @@ def build_metadata_dict(task: DownloadTask) -> dict:
"Year": task.year,
"Series": task.series_name,
"SeriesPosition": task.series_position,
"User": task.username,
}
@@ -70,10 +71,11 @@ def resolve_hardlink_source(
if hardlink_enabled and task.original_download_path:
hardlink_source = Path(task.original_download_path)
if destination and hardlink_source.exists() and same_filesystem(hardlink_source, destination):
hardlink_source_exists = run_blocking_io(hardlink_source.exists)
if destination and hardlink_source_exists and run_blocking_io(same_filesystem, hardlink_source, destination):
use_hardlink = True
source_path = hardlink_source
elif hardlink_source.exists():
elif hardlink_source_exists:
logger.warning(
f"Cannot hardlink: {hardlink_source} and {destination} are on different filesystems. "
"Falling back to copy. To fix: ensure torrent client downloads to same filesystem as destination."
@@ -97,7 +99,7 @@ def is_torrent_source(source_path: Path, task: DownloadTask) -> bool:
original_path = Path(task.original_download_path)
try:
return source_path.resolve() == original_path.resolve()
return run_blocking_io(source_path.resolve) == run_blocking_io(original_path.resolve)
except (OSError, ValueError):
try:
return os.path.normpath(str(source_path)) == os.path.normpath(str(original_path))
@@ -122,7 +124,7 @@ def _transfer_single_file(
if use_hardlink:
final_path = atomic_hardlink(source_path, dest_path, max_attempts=max_attempts)
try:
if os.stat(source_path).st_ino == os.stat(final_path).st_ino:
if run_blocking_io(source_path.stat).st_ino == run_blocking_io(final_path.stat).st_ino:
return final_path, "hardlink"
except OSError:
return final_path, "hardlink"
@@ -142,15 +144,16 @@ def transfer_book_files(
is_torrent: bool,
preserve_source: bool = False,
organization_mode: Optional[str] = None,
) -> Tuple[List[Path], Optional[str]]:
) -> Tuple[List[Path], Optional[str], Dict[str, int]]:
if not book_files:
return [], "No book files found"
return [], "No book files found", {"hardlink": 0, "copy": 0, "move": 0}
is_audiobook = check_audiobook(task.content_type)
organization_mode = organization_mode or get_file_organization(is_audiobook)
max_attempts = _max_attempts_for_batch(len(book_files))
final_paths: List[Path] = []
op_counts: Dict[str, int] = {"hardlink": 0, "copy": 0, "move": 0}
if organization_mode == "organize":
template = get_template(is_audiobook, "organize")
@@ -159,8 +162,14 @@ def transfer_book_files(
if len(book_files) == 1:
source_file = book_files[0]
ext = source_file.suffix.lstrip(".") or task.format or ""
dest_path = build_library_path(str(destination), template, metadata, extension=ext or None)
dest_path.parent.mkdir(parents=True, exist_ok=True)
dest_path = run_blocking_io(
build_library_path,
str(destination),
template,
metadata,
extension=ext or None,
)
run_blocking_io(dest_path.parent.mkdir, parents=True, exist_ok=True)
final_path, op = _transfer_single_file(
source_file,
@@ -171,6 +180,7 @@ def transfer_book_files(
max_attempts=max_attempts,
)
final_paths.append(final_path)
op_counts[op] = op_counts.get(op, 0) + 1
logger.debug(f"{op.capitalize()} to destination: {final_path.name}")
else:
zero_pad_width = max(len(str(len(book_files))), 2)
@@ -179,8 +189,14 @@ def transfer_book_files(
for source_file, part_number in files_with_parts:
ext = source_file.suffix.lstrip(".") or task.format or ""
file_metadata = {**metadata, "PartNumber": part_number}
dest_path = build_library_path(str(destination), template, file_metadata, extension=ext or None)
dest_path.parent.mkdir(parents=True, exist_ok=True)
dest_path = run_blocking_io(
build_library_path,
str(destination),
template,
file_metadata,
extension=ext or None,
)
run_blocking_io(dest_path.parent.mkdir, parents=True, exist_ok=True)
final_path, op = _transfer_single_file(
source_file,
@@ -191,9 +207,10 @@ def transfer_book_files(
max_attempts=max_attempts,
)
final_paths.append(final_path)
op_counts[op] = op_counts.get(op, 0) + 1
logger.debug(f"{op.capitalize()} to destination: {final_path.name}")
return final_paths, None
return final_paths, None, op_counts
for book_file in book_files:
if len(book_files) == 1 and organization_mode != "none":
@@ -223,9 +240,10 @@ def transfer_book_files(
max_attempts=max_attempts,
)
final_paths.append(final_path)
op_counts[op] = op_counts.get(op, 0) + 1
logger.debug(f"{op.capitalize()} to destination: {final_path.name}")
return final_paths, None
return final_paths, None, op_counts
def process_directory(
@@ -257,7 +275,7 @@ def process_directory(
if use_hardlink is None:
use_hardlink = should_hardlink(task)
final_paths, error = transfer_book_files(
final_paths, error, _op_counts = transfer_book_files(
book_files,
destination=ingest_dir,
task=task,
@@ -293,8 +311,8 @@ def transfer_file_to_library(
use_hardlink: bool,
) -> Optional[str]:
extension = source_path.suffix.lstrip(".") or task.format
dest_path = build_library_path(library_base, template, metadata, extension)
dest_path.parent.mkdir(parents=True, exist_ok=True)
dest_path = run_blocking_io(build_library_path, library_base, template, metadata, extension)
run_blocking_io(dest_path.parent.mkdir, parents=True, exist_ok=True)
is_torrent = is_torrent_source(source_path, task)
final_path, op = _transfer_single_file(
@@ -305,6 +323,12 @@ def transfer_file_to_library(
max_attempts=_max_attempts_for_batch(1),
)
logger.info(f"Library {op}: {final_path}")
if use_hardlink and op != "hardlink":
logger.warning(
"Library hardlink requested but %s used instead for %s",
op,
final_path,
)
if use_hardlink and temp_file and not is_torrent_source(temp_file, task):
safe_cleanup_path(temp_file, task)
@@ -339,11 +363,18 @@ def transfer_directory_to_library(
safe_cleanup_path(temp_file, task)
return None
base_library_path = build_library_path(library_base, template, metadata, extension=None)
base_library_path.parent.mkdir(parents=True, exist_ok=True)
base_library_path = run_blocking_io(
build_library_path,
library_base,
template,
metadata,
extension=None,
)
run_blocking_io(base_library_path.parent.mkdir, parents=True, exist_ok=True)
is_torrent = is_torrent_source(source_dir, task)
transferred_paths: List[Path] = []
op_counts: Dict[str, int] = {"hardlink": 0, "copy": 0, "move": 0}
max_attempts = _max_attempts_for_batch(len(source_files))
if len(source_files) == 1:
@@ -359,6 +390,7 @@ def transfer_directory_to_library(
)
logger.debug(f"Library {op}: {source_file.name} -> {final_path}")
transferred_paths.append(final_path)
op_counts[op] = op_counts.get(op, 0) + 1
else:
zero_pad_width = max(len(str(len(source_files))), 2)
files_with_parts = assign_part_numbers(source_files, zero_pad_width)
@@ -366,8 +398,8 @@ def transfer_directory_to_library(
for source_file, part_number in files_with_parts:
ext = source_file.suffix.lstrip(".")
file_metadata = {**metadata, "PartNumber": part_number}
file_path = build_library_path(library_base, template, file_metadata, extension=ext)
file_path.parent.mkdir(parents=True, exist_ok=True)
file_path = run_blocking_io(build_library_path, library_base, template, file_metadata, extension=ext)
run_blocking_io(file_path.parent.mkdir, parents=True, exist_ok=True)
final_path, op = _transfer_single_file(
source_file,
@@ -378,14 +410,23 @@ def transfer_directory_to_library(
)
logger.debug(f"Library {op}: {source_file.name} -> {final_path}")
transferred_paths.append(final_path)
op_counts[op] = op_counts.get(op, 0) + 1
if use_hardlink:
operation = "hardlinks"
elif is_torrent:
operation = "copies"
else:
operation = "files"
logger.info(f"Created {len(transferred_paths)} library {operation} in {base_library_path.parent}")
op_summary = ", ".join(
f"{op}={count}" for op, count in op_counts.items() if count
) or "none"
logger.info(
"Created %d library file(s) in %s (ops: %s)",
len(transferred_paths),
base_library_path.parent,
op_summary,
)
if use_hardlink and op_counts.get("copy", 0):
logger.warning(
"Library hardlink requested but %d of %d file(s) copied (fallback)",
op_counts.get("copy", 0),
len(transferred_paths),
)
if use_hardlink and temp_file and not is_torrent_source(temp_file, task):
safe_cleanup_path(temp_file, task)
+17 -3
View File
@@ -7,6 +7,7 @@ from typing import List, Optional
from shelfmark.config import env as env_config
from shelfmark.core.logger import setup_logger
from shelfmark.core.models import DownloadTask
from shelfmark.download.fs import run_blocking_io
from shelfmark.download.staging import STAGE_NONE
from .types import OutputPlan
@@ -21,8 +22,20 @@ def _tmp_dir() -> Path:
def is_within_tmp_dir(path: Path) -> bool:
"""Legacy helper: True if path is inside TMP_DIR."""
# Fast path: avoid `resolve()` (can block on NFS) for obviously-non-TMP paths.
# This is a *negative* check only; for potential TMP paths we still resolve to
# prevent symlink escapes from being treated as managed.
tmp_dir = _tmp_dir()
try:
path.resolve().relative_to(_tmp_dir().resolve())
if path.is_absolute() and tmp_dir.is_absolute():
if path != tmp_dir and tmp_dir not in path.parents:
return False
except Exception:
# Fall back to the slower resolve-based check below.
pass
try:
run_blocking_io(path.resolve).relative_to(run_blocking_io(tmp_dir.resolve))
return True
except (OSError, ValueError):
return False
@@ -42,7 +55,8 @@ def _is_original_download(path: Optional[Path], task: DownloadTask) -> bool:
if not path or not task.original_download_path:
return False
try:
return path.resolve() == Path(task.original_download_path).resolve()
original = Path(task.original_download_path)
return run_blocking_io(path.resolve) == run_blocking_io(original.resolve)
except (OSError, ValueError):
return False
@@ -59,7 +73,7 @@ def safe_cleanup_path(path: Optional[Path], task: DownloadTask) -> None:
try:
if path.is_dir():
shutil.rmtree(path, ignore_errors=True)
run_blocking_io(shutil.rmtree, path, ignore_errors=True)
elif path.exists():
path.unlink(missing_ok=True)
except (OSError, PermissionError) as exc:
+13 -11
View File
@@ -7,6 +7,7 @@ from typing import Literal
from shelfmark.config import env as env_config
from shelfmark.core.logger import setup_logger
from shelfmark.download.fs import run_blocking_io
logger = setup_logger(__name__)
@@ -19,7 +20,7 @@ STAGE_MOVE: StageAction = "move"
def get_staging_dir() -> Path:
"""Get the staging directory for downloads."""
tmp_dir = env_config.TMP_DIR
tmp_dir.mkdir(parents=True, exist_ok=True)
run_blocking_io(tmp_dir.mkdir, parents=True, exist_ok=True)
return tmp_dir
@@ -40,11 +41,11 @@ def build_staging_dir(prefix: str | None, task_id: str) -> Path:
staging_dir = base_dir / f"{prefix}_{safe_id}"
counter = 1
while staging_dir.exists():
while run_blocking_io(staging_dir.exists):
staging_dir = base_dir / f"{prefix}_{safe_id}_{counter}"
counter += 1
staging_dir.mkdir(parents=True, exist_ok=True)
run_blocking_io(staging_dir.mkdir, parents=True, exist_ok=True)
return staging_dir
@@ -62,23 +63,24 @@ def stage_path(source: Path, staging_dir: Path, action: StageAction) -> Path:
staged_path = staging_dir / source.name
counter = 1
if source.is_dir():
while staged_path.exists():
source_is_dir = run_blocking_io(source.is_dir)
if source_is_dir:
while run_blocking_io(staged_path.exists):
staged_path = staging_dir / f"{source.name}_{counter}"
counter += 1
if action == STAGE_COPY:
shutil.copytree(str(source), str(staged_path))
run_blocking_io(shutil.copytree, str(source), str(staged_path))
else:
shutil.move(str(source), str(staged_path))
run_blocking_io(shutil.move, str(source), str(staged_path))
else:
while staged_path.exists():
while run_blocking_io(staged_path.exists):
staged_path = staging_dir / f"{source.stem}_{counter}{source.suffix}"
counter += 1
if action == STAGE_COPY:
shutil.copy2(str(source), str(staged_path))
run_blocking_io(shutil.copy2, str(source), str(staged_path))
else:
shutil.move(str(source), str(staged_path))
run_blocking_io(shutil.move, str(source), str(staged_path))
staged_kind = "directory" if source.is_dir() else "file"
staged_kind = "directory" if source_is_dir else "file"
logger.debug("Staged %s via %s: %s -> %s", staged_kind, action, source, staged_path)
return staged_path
+747 -85
View File
File diff suppressed because it is too large Load Diff
+16 -1
View File
@@ -58,9 +58,13 @@ class ColumnRenderType(str, Enum):
"""How the frontend should render the column value."""
TEXT = "text" # Plain text
BADGE = "badge" # Colored badge (format, language)
TAGS = "tags" # List of colored badges
SIZE = "size" # File size formatting
NUMBER = "number" # Numeric value
PEERS = "peers" # Peers display: "S/L" with color based on seeder count
INDEXER_PROTOCOL = "indexer_protocol" # Text + colored dot for torrent/usenet
FLAG_ICON = "flag_icon" # Icon with tooltip (VIP, freeleech, etc.)
FORMAT_CONTENT_TYPE = "format_content_type" # Content type icon + format badge
class ColumnAlign(str, Enum):
@@ -123,8 +127,10 @@ class ReleaseColumnConfig:
grid_template: str = "minmax(0,2fr) 60px 80px 80px" # CSS grid-template-columns
leading_cell: Optional[LeadingCellConfig] = None # Defaults to thumbnail mode if None
online_servers: Optional[List[str]] = None # For IRC: list of currently online server nicks
available_indexers: Optional[List[str]] = None # For Prowlarr: list of all enabled indexer names
default_indexers: Optional[List[str]] = None # For Prowlarr: indexers selected in settings (pre-selected in filter)
cache_ttl_seconds: Optional[int] = None # How long to cache results (default: 5 min)
supported_filters: Optional[List[str]] = None # Which filters this source supports: ["format", "language"]
supported_filters: Optional[List[str]] = None # Which filters this source supports: ["format", "language", "indexer"]
action_button: Optional[SourceActionButton] = None # Custom action button (replaces default expand search)
@@ -169,6 +175,14 @@ def serialize_column_config(config: ReleaseColumnConfig) -> Dict[str, Any]:
if config.online_servers is not None:
result["online_servers"] = config.online_servers
# Include available_indexers if provided (e.g., for Prowlarr source)
if config.available_indexers is not None:
result["available_indexers"] = config.available_indexers
# Include default_indexers if provided (indexers selected in settings, for pre-selection)
if config.default_indexers is not None:
result["default_indexers"] = config.default_indexers
# Include cache TTL if specified (sources can request longer caching)
if config.cache_ttl_seconds is not None:
result["cache_ttl_seconds"] = config.cache_ttl_seconds
@@ -352,3 +366,4 @@ def get_source_display_name(name: str) -> str:
from shelfmark.release_sources import direct_download # noqa: F401, E402
from shelfmark.release_sources import prowlarr # noqa: F401, E402
from shelfmark.release_sources import irc # noqa: F401, E402
from shelfmark.release_sources import audiobookbay # noqa: F401, E402
@@ -0,0 +1,6 @@
"""AudiobookBay release source - web scraping for audiobook torrents."""
# Import to trigger registration
from shelfmark.release_sources.audiobookbay import source # noqa: F401, E402
from shelfmark.release_sources.audiobookbay import handler # noqa: F401, E402
from shelfmark.release_sources.audiobookbay import settings # noqa: F401, E402
@@ -0,0 +1,83 @@
"""AudiobookBay download handler - resolves magnet links and uses shared client lifecycle."""
from typing import Callable, Optional
from urllib.parse import urlparse
from shelfmark.core.config import config
from shelfmark.core.logger import setup_logger
from shelfmark.core.models import DownloadTask
from shelfmark.download.clients import DownloadClient, get_client, list_configured_clients
from shelfmark.download.clients.base_handler import DownloadRequest, ExternalClientHandler
from shelfmark.release_sources import register_handler
from shelfmark.release_sources.audiobookbay import scraper
from shelfmark.release_sources.audiobookbay.utils import normalize_hostname
logger = setup_logger(__name__)
@register_handler("audiobookbay")
class AudiobookBayHandler(ExternalClientHandler):
"""Handler for AudiobookBay downloads via configured torrent client."""
@staticmethod
def _resolve_detail_url(task: DownloadTask) -> Optional[str]:
"""Resolve ABB detail URL from queued task metadata."""
source_url = (task.source_url or "").strip()
if source_url:
return source_url
# Backward-compat: older tests and some legacy flows used task_id as URL.
task_id = (task.task_id or "").strip()
if task_id.startswith(("http://", "https://")):
return task_id
return None
def _get_client(self, protocol: str) -> Optional[DownloadClient]:
"""Compatibility shim so module-level patching still works in tests."""
return get_client(protocol)
def _list_configured_clients(self) -> list[str]:
"""Compatibility shim so module-level patching still works in tests."""
return list_configured_clients()
def _resolve_download(
self,
task: DownloadTask,
status_callback: Callable[[str, Optional[str]], None],
) -> Optional[DownloadRequest]:
"""Resolve ABB detail page into a magnet-link download request."""
detail_url = self._resolve_detail_url(task)
if not detail_url:
status_callback("error", "Missing AudiobookBay details URL")
logger.warning(f"Missing details URL for AudiobookBay task: {task.task_id}")
return None
hostname = normalize_hostname(config.get("ABB_HOSTNAME", ""))
if not hostname:
hostname = normalize_hostname(urlparse(detail_url).hostname)
status_callback("resolving", "Extracting magnet link")
magnet_link = scraper.extract_magnet_link(detail_url, hostname)
if not magnet_link:
status_callback("error", "Failed to extract magnet link from detail page")
return None
logger.info(f"Extracted magnet link for task {task.task_id}")
return DownloadRequest(
url=magnet_link,
protocol="torrent",
release_name=task.title or "Unknown",
expected_hash=None,
)
def cancel(self, task_id: str) -> bool:
"""Cancel an in-progress download.
Shelfmark can stop waiting via the queue cancel flag, but once a magnet has
been sent to the torrent client we do not remove it client-side. Users must
cancel/remove it in their torrent client UI.
"""
logger.debug(f"Cancel requested for AudiobookBay task: {task_id}")
return False
@@ -0,0 +1,398 @@
"""Web scraping functions for AudiobookBay."""
import re
import time
from typing import List, Optional, Dict
from urllib.parse import quote
import requests
from bs4 import BeautifulSoup
from shelfmark.core.config import config
from shelfmark.core.logger import setup_logger
from shelfmark.download import http as downloader
logger = setup_logger(__name__)
# Default trackers if none found on page
DEFAULT_TRACKERS = [
"udp://tracker.openbittorrent.com:80",
"udp://opentor.org:2710",
"udp://tracker.ccc.de:80",
"udp://tracker.blackunicorn.xyz:6969",
"udp://tracker.coppersurfer.tk:6969",
"udp://tracker.leechers-paradise.org:6969",
]
# ABB request behavior tuning
SEARCH_PAGE_RETRY_ATTEMPTS = 2
DETAIL_PAGE_RETRY_ATTEMPTS = 2
FIRST_PAGE_SESSION_REFRESH_ATTEMPTS = 2
# Legacy search parameter used by older ABB flows
LEGACY_CATEGORY_QUERY = "undefined%2Cundefined"
# Precompiled patterns used while parsing result cards
LANGUAGE_PATTERN = re.compile(r"Language:\s*([A-Za-z]+)")
POSTED_PATTERN = re.compile(r"Posted:\s*(\d+\s+[A-Za-z]+\s+\d{4})")
FORMAT_PATTERN = re.compile(r"Format:\s*([A-Za-z0-9]+)")
BITRATE_PATTERN = re.compile(r"Bitrate:\s*([\d]+\s*[A-Za-z/]+)")
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 _build_search_url(
hostname: str,
page: int,
query_encoded: str,
*,
include_legacy_category: bool = False,
) -> str:
"""Build an ABB search URL, optionally including legacy category params."""
# Page 1 uses ABB's root search endpoint; pagination continues via /page/{n}/.
if page <= 1:
url = f"https://{hostname}/?s={query_encoded}"
else:
url = f"https://{hostname}/page/{page}/?s={query_encoded}"
if include_legacy_category:
return f"{url}&cat={LEGACY_CATEGORY_QUERY}"
return url
def _is_homepage_redirect(final_url: str, hostname: str) -> bool:
"""Detect whether ABB redirected a search request to its homepage."""
normalized_final = (final_url or "").rstrip("/")
normalized_home = f"https://{hostname}".rstrip("/")
return normalized_final in {normalized_home, f"{normalized_home}/"}
def _encode_search_query(query: str, exact_phrase: bool) -> str:
"""Encode search query using ABB's space-plus style and optional exact phrase wrapping."""
search_query = query.strip()
if exact_phrase and search_query and not (search_query.startswith('"') and search_query.endswith('"')):
search_query = f"\"{search_query}\""
# Keep ABB-friendly encoding style (spaces as '+') while percent-encoding quotes.
return search_query.replace('"', "%22").replace(" ", "+")
def _normalize_result_url(url: str, hostname: str) -> str:
"""Normalize ABB result URLs to absolute HTTPS URLs."""
normalized_url = (url or "").strip()
if not normalized_url:
return ""
if normalized_url.startswith(("http://", "https://")):
return normalized_url
if normalized_url.startswith("//"):
return f"https:{normalized_url}"
if normalized_url.startswith("/"):
return f"https://{hostname}{normalized_url}"
return f"https://{hostname}/{normalized_url.lstrip('/')}"
def _bootstrap_abb_session(
hostname: str,
session: requests.Session,
retry_attempts: int,
) -> None:
"""Warm up ABB session cookies (best effort)."""
downloader.html_get_page(
f"https://{hostname}/",
retry=retry_attempts,
use_bypasser=False,
allow_bypasser_fallback=False,
include_response_url=True,
success_delay=0,
session=session,
)
def search_audiobookbay(
query: str,
max_pages: int = 1,
hostname: str = "audiobookbay.lu",
exact_phrase: bool = False,
) -> List[Dict[str, str]]:
"""Search AudiobookBay for audiobooks matching the query.
Args:
query: Search query string
max_pages: Maximum number of pages to fetch
hostname: AudiobookBay hostname (e.g., "audiobookbay.lu")
exact_phrase: Wrap query in quotes for exact phrase matching
Returns:
List of dicts with keys: title, link, cover, language, format, bitrate, size, posted_date
"""
results = []
rate_limit_delay = config.get("ABB_RATE_LIMIT_DELAY", 1.0)
session = requests.Session()
# Bootstrap ABB session cookie (PHPSESSID). ABB increasingly serves reliable
# search/detail pages only after session initialization, similar to browsers.
_bootstrap_abb_session(hostname, session, SEARCH_PAGE_RETRY_ATTEMPTS)
# Iterate through pages
for page in range(1, max_pages + 1):
# Construct URL - use + for spaces (matching audiobookbay-automated implementation)
# This avoids aggressive encoding that PHP-based sites may reject.
query_encoded = _encode_search_query(query, exact_phrase)
# ABB search expects the legacy category query parameter.
primary_url = _build_search_url(
hostname,
page,
query_encoded,
include_legacy_category=True,
)
try:
# Reuse shared HTTP fetch logic (without bypasser)
page_html, final_url = downloader.html_get_page(
primary_url,
retry=SEARCH_PAGE_RETRY_ATTEMPTS,
use_bypasser=False,
allow_bypasser_fallback=False,
include_response_url=True,
success_delay=0,
session=session,
)
was_home_redirect = _is_homepage_redirect(final_url, hostname)
# ABB can intermittently fail even with a valid URL.
# If page 1 fails, refresh the session and retry the exact same URL.
if page == 1 and (not page_html or was_home_redirect):
for refresh_attempt in range(1, FIRST_PAGE_SESSION_REFRESH_ATTEMPTS + 1):
session = requests.Session()
_bootstrap_abb_session(hostname, session, SEARCH_PAGE_RETRY_ATTEMPTS)
page_html, final_url = downloader.html_get_page(
primary_url,
retry=SEARCH_PAGE_RETRY_ATTEMPTS,
use_bypasser=False,
allow_bypasser_fallback=False,
include_response_url=True,
success_delay=0,
session=session,
)
was_home_redirect = _is_homepage_redirect(final_url, hostname)
if page_html and not was_home_redirect:
break
logger.debug(
"ABB page 1 session refresh %d/%d failed",
refresh_attempt,
FIRST_PAGE_SESSION_REFRESH_ATTEMPTS,
)
if not page_html:
logger.warning(f"Failed to fetch page {page}")
break
# Check if we were redirected to the homepage (search was rejected/blocked)
if was_home_redirect:
# Search was redirected to homepage - this means the search failed
# This can happen due to geo-blocking, rate limiting, or invalid query format
if page == 1:
logger.warning(f"Search query '{query}' was redirected to homepage - search may be blocked or invalid")
break
# Parse HTML
soup = BeautifulSoup(page_html, 'html.parser')
# Extract book entries
posts = soup.select('.post')
if not posts:
# No more results
break
for post in posts:
try:
# Extract title
title_elem = post.select_one('.postTitle > h2 > a')
if not title_elem:
continue
title = title_elem.text.strip()
# Extract link (relative, needs hostname prefix)
href = title_elem.get('href', '')
if not href:
continue
link = _normalize_result_url(href, hostname)
if not link:
continue
# Extract cover image (try .postContent .center img first, then fallback to any img)
cover = None
cover_elem = post.select_one('.postContent .center img') or post.select_one('img')
if cover_elem:
cover = _normalize_result_url(cover_elem.get('src', ''), hostname) or None
# Extract language from .postInfo
language = None
post_info = post.select_one('.postInfo')
if post_info:
info_text = post_info.get_text(separator=' ', strip=True).replace('\xa0', ' ')
lang_match = LANGUAGE_PATTERN.search(info_text)
if lang_match:
language = lang_match.group(1).strip()
# Extract format, bitrate, size, and posted date from .postContent
posted_date = None
format_type = None
bitrate = None
size_str = None
post_content = post.select_one('.postContent')
if post_content:
content_text = post_content.get_text(separator=' ', strip=True).replace('\xa0', ' ')
# Extract posted date
posted_match = POSTED_PATTERN.search(content_text)
if posted_match:
posted_date = posted_match.group(1).strip()
# Extract format (e.g., "M4B", "MP3")
format_match = FORMAT_PATTERN.search(content_text)
if format_match:
format_type = format_match.group(1).strip()
# Extract bitrate (e.g., "256 Kbps")
bitrate_match = BITRATE_PATTERN.search(content_text)
if bitrate_match:
bitrate = bitrate_match.group(1).strip()
# Extract file size (e.g., "11.68 GBs" -> normalized to "11.68 GB")
size_match = SIZE_PATTERN.search(content_text)
if size_match:
size_value = size_match.group(1)
size_unit = size_match.group(2).strip()
if size_unit.lower().endswith("s"):
size_unit = size_unit[:-1]
size_unit = size_unit.upper()
size_str = f"{size_value} {size_unit}"
results.append({
'title': title,
'link': link,
'cover': cover or None,
'language': language,
'format': format_type,
'bitrate': bitrate,
'size': size_str,
'posted_date': posted_date,
})
except Exception as e:
logger.debug(f"Skipping post due to error: {e}")
continue
# Rate limiting delay between pages
if page < max_pages and rate_limit_delay > 0:
time.sleep(rate_limit_delay)
except Exception as e:
logger.error(f"Unexpected error on page {page}: {e}")
break
logger.info(f"Found {len(results)} results for query '{query}'")
return results
def extract_magnet_link(
details_url: str,
hostname: str = "audiobookbay.lu"
) -> Optional[str]:
"""Extract info hash and trackers from book detail page, then construct magnet link.
Args:
details_url: URL of the book's detail page
hostname: AudiobookBay hostname (for logging)
Returns:
Magnet link, or None if extraction fails
"""
try:
session = requests.Session()
_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(
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")
return None
soup = BeautifulSoup(detail_html, 'html.parser')
# 1. Extract Info Hash
# Look for <td>Info Hash</td> and get next sibling value
info_hash = None
info_hash_rows = soup.find_all('td')
for td in info_hash_rows:
if td.text.strip().lower() == 'info hash':
next_td = td.find_next_sibling('td')
if next_td:
info_hash = next_td.text.strip()
break
# Alternative: search for text containing "Info Hash" and get next element
if not info_hash:
for elem in soup.find_all(string=INFO_HASH_LABEL_PATTERN):
parent = elem.parent
if parent and parent.name == 'td':
next_td = parent.find_next_sibling('td')
if next_td:
info_hash = next_td.text.strip()
break
if not info_hash:
logger.warning("Info Hash not found on the page.")
return None
# Clean up info hash (remove whitespace, ensure uppercase)
info_hash = re.sub(r'\s+', '', info_hash).upper()
# 2. Extract Trackers
# Find all <td> containing udp:// or http://
trackers = []
for td in soup.find_all('td'):
text = td.text.strip()
if text.startswith(('udp://', 'http://', 'https://')):
trackers.append(text)
# 3. Use default trackers if none found
if not trackers:
logger.debug("No trackers found on the page. Using default trackers.")
trackers = DEFAULT_TRACKERS
# 4. Construct Magnet Link
# Format: magnet:?xt=urn:btih:{INFO_HASH}&tr={TRACKER1}&tr={TRACKER2}...
tracker_params = "&".join(
f"tr={quote(tracker)}"
for tracker in trackers
)
magnet_link = f"magnet:?xt=urn:btih:{info_hash}&{tracker_params}"
logger.debug(f"Generated Magnet Link: {magnet_link[:100]}...")
return magnet_link
except Exception as e:
logger.error(f"Failed to extract magnet link: {e}")
return None
@@ -0,0 +1,57 @@
"""AudiobookBay settings registration."""
from shelfmark.core.settings_registry import (
register_settings,
CheckboxField,
TextField,
NumberField,
)
# ==================== Register Settings ====================
@register_settings("audiobookbay_config", "AudiobookBay", icon="download", order=45)
def audiobookbay_config_settings():
"""AudiobookBay configuration settings."""
return [
CheckboxField(
key="ABB_ENABLED",
label="Enable AudiobookBay",
description="Enable AudiobookBay as a release source for audiobooks.",
default=False,
),
TextField(
key="ABB_HOSTNAME",
label="Hostname",
description="AudiobookBay domain (e.g., audiobookbay.lu, audiobookbay.is). Required to enable searches.",
placeholder="",
default="",
required=True,
show_when={"field": "ABB_ENABLED", "value": True},
),
NumberField(
key="ABB_PAGE_LIMIT",
label="Max Pages to Search",
description="Maximum number of search result pages to fetch (1-10).",
default=1,
min_value=1,
max_value=10,
show_when={"field": "ABB_ENABLED", "value": True},
),
CheckboxField(
key="ABB_EXACT_PHRASE",
label="Prefer Exact-Phrase Search",
description="Wrap generated queries in quotes for stricter matching. If no results are found, Shelfmark retries without quotes.",
default=False,
show_when={"field": "ABB_ENABLED", "value": True},
),
NumberField(
key="ABB_RATE_LIMIT_DELAY",
label="Rate Limit Delay (seconds)",
description="Delay between requests in seconds to avoid rate limiting (0-10).",
default=1.0,
min_value=0.0,
max_value=10.0,
show_when={"field": "ABB_ENABLED", "value": True},
),
]
@@ -0,0 +1,366 @@
"""AudiobookBay release source - searches AudiobookBay for audiobook torrents."""
import hashlib
import re
from typing import List, Optional, TYPE_CHECKING
if TYPE_CHECKING:
from shelfmark.core.search_plan import ReleaseSearchPlan
from shelfmark.core.config import config
from shelfmark.core.logger import setup_logger
from shelfmark.metadata_providers import BookMetadata
from shelfmark.release_sources import (
Release,
ReleaseProtocol,
ReleaseSource,
register_source,
ReleaseColumnConfig,
ColumnSchema,
ColumnRenderType,
ColumnAlign,
ColumnColorHint,
)
from shelfmark.release_sources.audiobookbay import scraper
from shelfmark.release_sources.audiobookbay.utils import normalize_hostname, parse_size
logger = setup_logger(__name__)
# Map language names to ISO 639-1 codes (matching frontend color maps)
LANGUAGE_MAP = {
"english": "en",
"spanish": "es",
"french": "fr",
"german": "de",
"italian": "it",
"portuguese": "pt",
"russian": "ru",
"japanese": "ja",
"chinese": "zh",
"dutch": "nl",
"swedish": "sv",
"norwegian": "no",
"danish": "da",
"finnish": "fi",
"polish": "pl",
"czech": "cs",
"hungarian": "hu",
"korean": "ko",
"arabic": "ar",
"hebrew": "he",
"turkish": "tr",
"greek": "el",
"hindi": "hi",
"thai": "th",
"vietnamese": "vi",
"indonesian": "id",
"ukrainian": "uk",
"romanian": "ro",
"bulgarian": "bg",
"catalan": "ca",
"croatian": "hr",
"slovenian": "sl",
"serbian": "sr",
}
def _split_title_and_author(raw_title: str) -> tuple[str, Optional[str]]:
"""Split titles in the form 'Title - Author' into title and author.
Args:
raw_title: The raw title string from the scrape.
Returns:
(title, author) where author is None if split is unavailable.
"""
if not raw_title:
return "", None
cleaned_title = raw_title.strip()
if " - " not in cleaned_title:
return cleaned_title, None
title_part, author_part = cleaned_title.rsplit(" - ", 1)
title_part = title_part.strip()
author_part = author_part.strip()
if not title_part or not author_part:
return cleaned_title, None
return title_part, author_part
def _map_language(language: str) -> Optional[str]:
"""Map language name to ISO 639-1 code.
Args:
language: Language name (e.g., "English")
Returns:
ISO 639-1 code (e.g., "en"), or original string if no mapping found, or None if input is empty
"""
if not language:
return None
lang_lower = language.lower().strip()
return LANGUAGE_MAP.get(lang_lower, lang_lower)
def _parse_bitrate_to_kbps(bitrate: Optional[str]) -> Optional[int]:
"""Parse bitrate string to an integer Kbps value.
Args:
bitrate: Human-readable bitrate (e.g., "128 Kbps")
Returns:
Bitrate value in Kbps as integer, or None if parsing fails.
"""
if not bitrate:
return None
match = re.search(r"(\d+(?:\.\d+)?)\s*kbps", bitrate, re.IGNORECASE)
if not match:
return None
try:
return int(float(match.group(1)))
except ValueError:
return None
def _generate_source_id(detail_url: str) -> str:
"""Generate a unique source ID from detail URL."""
return hashlib.md5(detail_url.encode()).hexdigest()
@register_source("audiobookbay")
class AudiobookBaySource(ReleaseSource):
"""Release source for AudiobookBay audiobook torrents."""
name = "audiobookbay"
display_name = "AudiobookBay"
supported_content_types = ["audiobook"] # ONLY audiobooks
def search(
self,
book: BookMetadata,
plan: "ReleaseSearchPlan",
expand_search: bool = False,
content_type: str = "ebook"
) -> List[Release]:
"""Search AudiobookBay for audiobook releases.
Args:
book: Book metadata
plan: Search plan with query variants
expand_search: Ignored (always searches)
content_type: Must be "audiobook" for this source
Returns:
List of Release objects
"""
# Only search for audiobooks
if content_type != "audiobook":
return []
hostname = normalize_hostname(config.get("ABB_HOSTNAME", ""))
if not hostname:
logger.debug("AudiobookBay hostname is not configured")
return []
max_pages = config.get("ABB_PAGE_LIMIT", 1)
exact_phrase = bool(config.get("ABB_EXACT_PHRASE", False))
# Build search query candidates from plan.
query_candidates: list[str] = []
if plan.manual_query:
query_candidates.append(plan.manual_query.strip())
elif plan.title_variants:
variant = plan.title_variants[0]
combined_query = f"{variant.title} {variant.author}".strip()
title_only_query = (variant.title or "").strip()
if combined_query:
query_candidates.append(combined_query)
if title_only_query and title_only_query.lower() != combined_query.lower():
query_candidates.append(title_only_query)
elif book.title:
query_candidates.append(book.title.strip())
# Remove empty and duplicate queries while preserving order.
deduped_queries: list[str] = []
seen_queries: set[str] = set()
for candidate in query_candidates:
normalized = candidate.strip()
if not normalized:
continue
key = normalized.lower()
if key in seen_queries:
continue
seen_queries.add(key)
deduped_queries.append(normalized)
if not deduped_queries:
logger.debug("No search query available")
return []
results = []
query_lower = deduped_queries[0].lower()
try:
for index, query in enumerate(deduped_queries):
query_lower = query.lower()
logger.info(f"Searching AudiobookBay for: {query_lower}")
# Search AudiobookBay
results = scraper.search_audiobookbay(
query=query_lower,
max_pages=max_pages,
hostname=hostname,
exact_phrase=exact_phrase,
)
# For auto-generated queries, fallback to broad matching if exact phrase returns nothing.
if exact_phrase and not results and not plan.manual_query:
logger.info("No exact phrase results, retrying AudiobookBay search without quotes")
results = scraper.search_audiobookbay(
query=query_lower,
max_pages=max_pages,
hostname=hostname,
exact_phrase=False,
)
if results:
break
if index < len(deduped_queries) - 1:
logger.info(
"No AudiobookBay results for '%s', retrying with '%s'",
query_lower,
deduped_queries[index + 1].lower(),
)
# Extract query words for relevance checking
query_words = set(word.lower() for word in query_lower.split() if len(word) > 2)
releases = []
for result in results:
try:
raw_title = result['title']
title, author = _split_title_and_author(raw_title)
title_for_filter = raw_title.lower()
# Basic relevance check: ensure title contains at least one query word
# This filters out homepage "Latest" feed items that may leak through
if query_words:
if not any(word in title_for_filter for word in query_words):
logger.debug(f"Filtering out irrelevant result: {title}")
continue
# Generate unique source ID
source_id = _generate_source_id(result['link'])
# Extract and parse metadata
format_type = result.get('format')
size_str = result.get('size')
size_bytes = parse_size(size_str) if size_str else None
language_raw = result.get('language')
language_code = _map_language(language_raw) if language_raw else None
bitrate = result.get('bitrate')
bitrate_kbps = _parse_bitrate_to_kbps(bitrate)
# Create Release object
release = Release(
source="audiobookbay",
source_id=source_id,
title=title,
format=format_type.lower() if format_type else None,
language=language_code,
size=size_str,
size_bytes=size_bytes,
download_url=result['link'], # Detail page URL (used by handler)
info_url=result['link'], # Make title clickable
protocol=ReleaseProtocol.TORRENT,
indexer="AudiobookBay",
seeders=None, # Not available on search page
peers=None,
content_type="audiobook",
extra={
"preview": result.get('cover'),
"detail_url": result['link'],
"bitrate": bitrate,
"bitrate_value": bitrate_kbps,
"posted_date": result.get('posted_date'),
"title_raw": raw_title,
"language_raw": language_raw, # Keep original for reference
"author": author, # Parsed author from title pattern
}
)
releases.append(release)
except Exception as e:
logger.warning(f"Failed to create release from result: {e}")
continue
logger.info(f"Found {len(releases)} releases from AudiobookBay")
return releases
except Exception as e:
logger.error(f"AudiobookBay search error: {e}")
return []
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", "")))
def get_column_config(self) -> ReleaseColumnConfig:
"""Get column configuration for AudiobookBay releases.
Shows title, language, format, bitrate, and size columns.
No seeders/peers since ABB doesn't show this on search page.
"""
return ReleaseColumnConfig(
columns=[
ColumnSchema(
key="language",
label="Lang",
render_type=ColumnRenderType.BADGE,
align=ColumnAlign.CENTER,
width="60px",
hide_mobile=True,
color_hint=ColumnColorHint(type="map", value="language"),
uppercase=True,
fallback="",
),
ColumnSchema(
key="format",
label="Format",
render_type=ColumnRenderType.BADGE,
align=ColumnAlign.CENTER,
width="80px",
hide_mobile=False,
color_hint=ColumnColorHint(type="map", value="format"),
uppercase=True,
),
ColumnSchema(
key="extra.bitrate",
label="Bitrate",
render_type=ColumnRenderType.NUMBER,
align=ColumnAlign.CENTER,
width="72px",
hide_mobile=False,
fallback="",
sortable=True,
sort_key="extra.bitrate_value",
),
ColumnSchema(
key="size",
label="Size",
render_type=ColumnRenderType.SIZE,
align=ColumnAlign.CENTER,
width="80px",
hide_mobile=False,
sortable=True,
sort_key="size_bytes",
),
],
grid_template="minmax(0,2fr) 60px 80px 72px 80px",
supported_filters=["format", "language"], # Enable format and language filters
)
@@ -0,0 +1,55 @@
"""Utility functions for AudiobookBay integration."""
import re
from typing import Optional
def normalize_hostname(raw: Optional[str]) -> str:
"""Normalize a user-supplied hostname for URL construction.
Strips whitespace, scheme prefixes, trailing slashes, and paths so that
values like "https://audiobookbay.lu/" or " audiobookbay.lu/ " all
resolve to "audiobookbay.lu".
"""
if not raw or not isinstance(raw, str):
return ""
cleaned = raw.strip()
# Strip scheme
for prefix in ("https://", "http://"):
if cleaned.lower().startswith(prefix):
cleaned = cleaned[len(prefix):]
break
# Strip path and trailing slashes
cleaned = cleaned.split("/")[0].strip()
return cleaned
def parse_size(size_str: Optional[str]) -> Optional[int]:
"""Parse size string to bytes.
Args:
size_str: Size string (e.g., "1.5 GB", "500 MB", "11.68 GBs")
Returns:
Size in bytes, or None if parsing fails
"""
if not size_str:
return None
# Match number and unit, handling "GBs" as well as "GB" (case-insensitive)
match = re.search(r"([\d.]+)\s*([BKMGT]B?)S?", size_str.upper())
if not match:
return None
value = float(match.group(1))
unit = match.group(2)
multipliers = {
"B": 1,
"KB": 1024,
"MB": 1024 ** 2,
"GB": 1024 ** 3,
"TB": 1024 ** 4,
}
return int(value * multipliers.get(unit, 1))
@@ -1056,6 +1056,7 @@ def _book_info_to_release(book_info: BookInfo) -> Release:
source_id=book_info.id,
title=book_info.title,
format=book_info.format,
language=book_info.language, # Top-level language for filtering
size=book_info.size,
download_url=book_info.download_urls[0] if book_info.download_urls else None,
info_url=f"{network.get_aa_base_url()}/md5/{book_info.id}",
@@ -7,7 +7,6 @@ across multiple indexers (torrent and usenet).
Includes:
- ProwlarrSource: Search integration with Prowlarr
- ProwlarrHandler: Download handling via external clients
- Download clients: qBittorrent (torrents), NZBGet (usenet)
"""
# Import submodules to trigger decorator registration
@@ -15,12 +14,12 @@ from shelfmark.release_sources.prowlarr import source # noqa: F401
from shelfmark.release_sources.prowlarr import handler # noqa: F401
from shelfmark.release_sources.prowlarr import settings # noqa: F401
# Import clients to trigger client registration
# This is in a try/except to handle optional dependencies gracefully
# Import shared download clients/settings to trigger registration.
# This is in a try/except to handle optional dependencies gracefully.
try:
from shelfmark.release_sources.prowlarr import clients # noqa: F401
from shelfmark.download import clients # noqa: F401
from shelfmark.download.clients import settings as client_settings # noqa: F401
except ImportError as e:
# Log but don't fail - clients require optional dependencies
import logging
logging.getLogger(__name__).debug(f"Prowlarr clients not loaded: {e}")
logging.getLogger(__name__).debug(f"Download clients not loaded: {e}")
+100
View File
@@ -6,6 +6,7 @@ import requests
from shelfmark.core.logger import setup_logger
from shelfmark.core.utils import normalize_http_url
from shelfmark.release_sources.prowlarr.torznab import parse_torznab_xml
logger = setup_logger(__name__)
@@ -90,6 +91,44 @@ class ProwlarrClient:
logger.error(f"Failed to get indexers: {e}")
return []
def get_enabled_indexers_detailed(self) -> List[Dict[str, Any]]:
"""
Get enabled indexers, including implementation metadata.
Note: Prowlarr indexer "name" is user-configurable; prefer
"implementation"/"implementationName" for stable identification.
"""
indexers = self.get_indexers()
return [idx for idx in indexers if idx.get("enable", False)]
def get_enriched_indexer_ids(self, *, restrict_to: Optional[List[int]] = None) -> List[int]:
"""
Return enabled indexer IDs that should use Torznab for richer metadata.
Args:
restrict_to: Optional list of candidate indexer IDs to consider.
"""
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):
continue
if restrict_to is not None and idx_id_int not in restrict_to:
continue
impl = str(idx.get("implementation") or idx.get("implementationName") or idx.get("definitionName") or "")
# Currently only MyAnonamouse provides consistently rich Torznab metadata.
if impl.strip().lower() == "myanonamouse":
enriched_ids.append(idx_id_int)
return enriched_ids
def get_enabled_indexers(self) -> List[Dict[str, Any]]:
"""Get enabled indexers with book capability info."""
indexers = self.get_indexers()
@@ -112,6 +151,67 @@ class ProwlarrClient:
return result
def torznab_search(
self,
*,
indexer_id: int,
query: str,
categories: Optional[List[int]] = None,
search_type: str = "book",
limit: int = 100,
offset: int = 0,
) -> List[Dict[str, Any]]:
"""
Search a specific indexer via Prowlarr's Torznab/Newznab endpoint.
This returns richer fields (e.g., author/booktitle, torznab tags like
FreeLeech) than the JSON /api/v1/search endpoint.
"""
if not query:
return []
endpoint = f"/api/v1/indexer/{int(indexer_id)}/newznab"
url = self.base_url + endpoint
params: Dict[str, Any] = {
"t": search_type,
"q": query,
"limit": limit,
"offset": offset,
}
if categories:
params["cat"] = ",".join(str(c) for c in categories)
logger.debug(f"Prowlarr API: GET {url} (torznab)")
try:
response = self._session.get(
url=url,
params=params,
timeout=self.timeout,
headers={
# Override the session default JSON accept header.
"Accept": "application/rss+xml, application/xml;q=0.9, */*;q=0.8"
},
)
if not response.ok:
try:
error_body = response.text[:500]
logger.error(f"Prowlarr Torznab error response: {error_body}")
except Exception:
pass
response.raise_for_status()
results = parse_torznab_xml(response.text)
# Ensure indexerId is always set (Prowlarr includes it, but be defensive).
for r in results:
if r.get("indexerId") is None:
r["indexerId"] = int(indexer_id)
return results
except Exception as e:
logger.error(f"Prowlarr Torznab search failed for indexer {indexer_id}: {e}")
return []
def _has_book_categories(self, categories: List[Dict[str, Any]]) -> bool:
"""Check if any category or subcategory is in the book range (7000-7999)."""
for cat in categories:
+64 -640
View File
@@ -1,668 +1,92 @@
"""Prowlarr download handler - executes downloads via torrent/usenet clients."""
"""Prowlarr download handler - resolves releases and delegates lifecycle to shared clients."""
import shutil
from pathlib import Path
from threading import Event
from typing import Callable, Optional
from shelfmark.core.config import config
from shelfmark.core.config import config # noqa: F401 (compat patch target in tests)
from shelfmark.core.logger import setup_logger
from shelfmark.core.models import DownloadTask
from shelfmark.core.utils import is_audiobook
from shelfmark.release_sources import DownloadHandler, register_handler
from shelfmark.release_sources.prowlarr.cache import get_release, remove_release
from shelfmark.release_sources.prowlarr.clients import (
DownloadClient,
DownloadState,
get_client,
list_configured_clients,
from shelfmark.download.clients import DownloadClient, get_client, list_configured_clients
from shelfmark.download.clients.base_handler import (
COMPLETED_PATH_MAX_ATTEMPTS as _DEFAULT_COMPLETED_PATH_MAX_ATTEMPTS,
COMPLETED_PATH_RETRY_INTERVAL as _DEFAULT_COMPLETED_PATH_RETRY_INTERVAL,
POLL_INTERVAL as _DEFAULT_POLL_INTERVAL,
DownloadRequest,
ExternalClientHandler,
)
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 get_preferred_download_url, get_protocol
logger = setup_logger(__name__)
# How often to poll the download client for status (seconds)
POLL_INTERVAL = 2
def _diagnose_path_issue(path: str) -> str:
"""
Analyze a path and return diagnostic hints for common issues.
Args:
path: The path that failed to be accessed
Returns:
A hint string to help users diagnose the issue.
"""
# Detect Windows-style paths (won't work in Linux containers)
if len(path) >= 2 and path[1] == ':':
return (
f"Path '{path}' appears to be a Windows path. "
f"Shelfmark runs in Linux and cannot access Windows paths directly. "
f"Ensure your download client uses Linux-style paths (/path/to/files)."
)
# Detect backslashes (Windows path separators)
if '\\' in path:
return (
f"Path '{path}' contains backslashes. "
f"This may indicate a Windows path or incorrect path escaping. "
f"Linux paths should use forward slashes (/)."
)
# Generic hint for Linux paths
return (
f"Path '{path}' is not accessible from Shelfmark's container. "
f"Ensure both containers have matching volume mounts for this directory, "
f"or configure Remote Path Mappings in Settings > Advanced."
)
# Backwards-compat constants for tests patching this module.
POLL_INTERVAL = _DEFAULT_POLL_INTERVAL
COMPLETED_PATH_RETRY_INTERVAL = _DEFAULT_COMPLETED_PATH_RETRY_INTERVAL
COMPLETED_PATH_MAX_ATTEMPTS = _DEFAULT_COMPLETED_PATH_MAX_ATTEMPTS
@register_handler("prowlarr")
class ProwlarrHandler(DownloadHandler):
class ProwlarrHandler(ExternalClientHandler):
"""Handler for Prowlarr downloads via configured torrent or usenet client."""
def __init__(self):
# Track downloads that may need client-side cleanup after Shelfmark completes import.
# task_id -> (client, download_id, protocol)
self._cleanup_refs: dict[str, tuple[DownloadClient, str, str]] = {}
def _get_client(self, protocol: str) -> Optional[DownloadClient]:
"""Compatibility shim so module-level patching still works in tests."""
return get_client(protocol)
def _get_category_for_task(self, client, task: DownloadTask) -> Optional[str]:
"""Get audiobook category if configured and applicable, else None for default."""
if not is_audiobook(task.content_type):
def _list_configured_clients(self) -> list[str]:
"""Compatibility shim so module-level patching still works in tests."""
return list_configured_clients()
def _poll_interval(self) -> float:
return POLL_INTERVAL
def _completed_path_retry_interval(self) -> float:
return COMPLETED_PATH_RETRY_INTERVAL
def _completed_path_max_attempts(self) -> int:
return COMPLETED_PATH_MAX_ATTEMPTS
def _resolve_download(
self,
task: DownloadTask,
status_callback: Callable[[str, Optional[str]], None],
) -> Optional[DownloadRequest]:
"""Resolve Prowlarr cache entry into download request parameters."""
# Look up the cached release
prowlarr_result = get_release(task.task_id)
if not prowlarr_result:
logger.warning(f"Release cache miss: {task.task_id}")
status_callback("error", "Release not found in cache (may have expired)")
return None
# Client-specific audiobook category config keys
audiobook_keys = {
"qbittorrent": "QBITTORRENT_CATEGORY_AUDIOBOOK",
"transmission": "TRANSMISSION_CATEGORY_AUDIOBOOK",
"deluge": "DELUGE_CATEGORY_AUDIOBOOK",
"nzbget": "NZBGET_CATEGORY_AUDIOBOOK",
"sabnzbd": "SABNZBD_CATEGORY_AUDIOBOOK",
}
audiobook_key = audiobook_keys.get(client.name)
return config.get(audiobook_key, "") or None if audiobook_key else None
# Extract download URL
download_url = get_preferred_download_url(prowlarr_result)
if not download_url:
status_callback("error", "No download URL available")
return None
def post_process_cleanup(self, task: DownloadTask, success: bool) -> None:
if not success:
self._cleanup_refs.pop(task.task_id, None)
return
# Determine protocol
protocol = get_protocol(prowlarr_result)
if protocol == "unknown":
status_callback("error", "Could not determine download protocol")
return None
client_ref = self._cleanup_refs.pop(task.task_id, None)
if client_ref is None:
return
release_name = prowlarr_result.get("title") or task.title or "Unknown"
expected_hash = str(prowlarr_result.get("infoHash") or "").strip() or None
client, download_id, protocol = client_ref
if protocol != "usenet":
return
# "Move" means copy into ingest then let the usenet client delete its own files.
if config.get("PROWLARR_USENET_ACTION", "move") != "move":
return
try:
self._delete_local_download_data(client, download_id)
self._remove_usenet_download(client, download_id, delete_files=True, archive=True)
except Exception as e:
logger.warning(f"Failed to cleanup usenet download {download_id} in {getattr(client, 'name', 'client')}: {e}")
def _remove_usenet_download(
self,
client: DownloadClient,
download_id: str,
*,
delete_files: bool,
archive: bool = True,
) -> None:
"""Remove a usenet download with SABnzbd-specific archive handling."""
if getattr(client, "name", "") == "sabnzbd":
client.remove(download_id, delete_files=delete_files, archive=archive)
else:
client.remove(download_id, delete_files=delete_files)
def _delete_local_download_data(self, client: DownloadClient, download_id: str) -> None:
"""Best-effort local deletion of client download data."""
try:
raw_path = client.get_download_path(download_id)
except Exception as e:
logger.debug(f"Failed to resolve download path for {client.name} {download_id}: {e}")
return
if not raw_path:
logger.debug(f"No download path available for {client.name} {download_id}")
return
from shelfmark.core.path_mappings import (
get_client_host_identifier,
parse_remote_path_mappings,
remap_remote_to_local_with_match,
return DownloadRequest(
url=download_url,
protocol=protocol,
release_name=release_name,
expected_hash=expected_hash,
)
source_path_obj = Path(raw_path)
host = get_client_host_identifier(client) or ""
mapping_value = config.get("PROWLARR_REMOTE_PATH_MAPPINGS", [])
mappings = parse_remote_path_mappings(mapping_value)
remapped, matched_mapping = remap_remote_to_local_with_match(
mappings=mappings,
host=host,
remote_path=source_path_obj,
)
delete_path = remapped if matched_mapping else source_path_obj
if str(delete_path) in ("", "/"):
logger.warning(f"Refusing to delete unsafe path for {client.name} {download_id}: {delete_path}")
return
if not delete_path.exists():
logger.debug(f"Local download path does not exist for cleanup: {delete_path}")
return
try:
if delete_path.is_dir():
shutil.rmtree(delete_path)
else:
delete_path.unlink()
logger.info(f"Deleted local download data for {client.name} {download_id}: {delete_path}")
except Exception as e:
logger.warning(f"Failed to delete local download data for {client.name} {download_id}: {e}")
def _safe_remove_download(self, client, download_id: str, protocol: str, reason: str) -> None:
"""Best-effort removal of a failed/cancelled download from the client.
Safety policy:
- torrents: never remove or delete client data (avoid breaking seeding)
- usenet: keep legacy behavior (delete client files on removal)
"""
if protocol != "usenet":
logger.info(
"Skipping download client cleanup for protocol=%s after %s (client=%s id=%s)",
protocol,
reason,
getattr(client, "name", "client"),
download_id,
)
return
try:
# Permanent delete for failed usenet downloads (SABnzbd archive=0).
self._delete_local_download_data(client, download_id)
self._remove_usenet_download(client, download_id, delete_files=True, archive=False)
except Exception as e:
logger.warning(
f"Failed to remove download {download_id} from {client.name} after {reason}: {e}"
)
def _build_progress_message(self, status) -> str:
"""Build a progress message from download status."""
msg = f"{status.progress:.0f}%"
if status.download_speed and status.download_speed > 0:
speed_mb = status.download_speed / 1024 / 1024
msg += f" ({speed_mb:.1f} MB/s)"
if status.eta and status.eta > 0:
if status.eta < 60:
msg += f" - {status.eta}s left"
elif status.eta < 3600:
msg += f" - {status.eta // 60}m left"
else:
msg += f" - {status.eta // 3600}h {(status.eta % 3600) // 60}m left"
return msg
def download(
self,
task: DownloadTask,
cancel_flag: Event,
progress_callback: Callable[[float], None],
status_callback: Callable[[str, Optional[str]], None],
) -> Optional[str]:
"""Execute download via configured torrent/usenet client. Returns file path or None."""
try:
# Look up the cached release
prowlarr_result = get_release(task.task_id)
if not prowlarr_result:
logger.warning(f"Release cache miss: {task.task_id}")
status_callback("error", "Release not found in cache (may have expired)")
return None
# Extract download URL
download_url = get_preferred_download_url(prowlarr_result)
if not download_url:
status_callback("error", "No download URL available")
return None
# Determine protocol
protocol = get_protocol(prowlarr_result)
if protocol == "unknown":
status_callback("error", "Could not determine download protocol")
return None
# Get the appropriate download client
client = get_client(protocol)
if not client:
configured = list_configured_clients()
if not configured:
status_callback("error", "No download clients configured. Configure qBittorrent or NZBGet in settings.")
else:
status_callback("error", f"No {protocol} client configured")
return None
# Check if this download already exists in the client
status_callback("resolving", f"Checking {client.name}")
category = self._get_category_for_task(client, task)
existing = client.find_existing(download_url, category=category)
if existing:
download_id, existing_status = existing
logger.info(f"Found existing download in {client.name}: {download_id}")
# If already complete, skip straight to file handling
if existing_status.complete:
logger.info("Existing download is complete, copying file directly")
status_callback("resolving", "Found existing download, copying to library")
source_path = client.get_download_path(download_id)
if not source_path:
logger.error(
f"Could not get path for existing download. "
f"Client: {client.name}, ID: {download_id}. "
f"The download may have been moved or deleted."
)
status_callback(
"error",
f"Could not locate existing download in {client.name}. "
f"Check that the file still exists."
)
return None
from shelfmark.core.path_mappings import (
get_client_host_identifier,
parse_remote_path_mappings,
remap_remote_to_local_with_match,
)
source_path_obj = Path(source_path)
host = get_client_host_identifier(client) or ""
mapping_value = config.get("PROWLARR_REMOTE_PATH_MAPPINGS", [])
mappings = parse_remote_path_mappings(mapping_value)
remapped, matched_mapping = remap_remote_to_local_with_match(
mappings=mappings,
host=host,
remote_path=source_path_obj,
)
if matched_mapping:
if remapped.exists():
logger.info(
"Remapped existing download path for %s (%s): %s -> %s",
client.name,
download_id,
source_path_obj,
remapped,
)
source_path_obj = remapped
else:
logger.error(
f"Download path does not exist after remapping: {source_path} -> {remapped}. "
f"Client: {client.name}, ID: {download_id}. "
f"Check that the local path in your mapping is mounted correctly."
)
status_callback(
"error",
f"Remapped path '{remapped}' does not exist. "
f"Check your Docker volume mounts match the Local Path in Settings > Advanced > Remote Path Mappings.",
)
return None
elif mappings:
if source_path_obj.exists():
logger.info(
"No remote path mapping matched for %s (%s); using client path: %s",
client.name,
download_id,
source_path_obj,
)
else:
hint = _diagnose_path_issue(source_path)
logger.error(
f"Download path does not exist and no remote path mapping matched for {client.name} "
f"({download_id}): {source_path}. {hint}"
)
status_callback(
"error",
f"{hint} No remote path mapping matched for client '{client.name}'.",
)
return None
elif not source_path_obj.exists():
hint = _diagnose_path_issue(source_path)
logger.error(
f"Download path does not exist: {source_path}. "
f"Client: {client.name}, ID: {download_id}. {hint}"
)
status_callback("error", hint)
return None
result = self._handle_completed_file(
source_path=source_path_obj,
protocol=protocol,
task=task,
status_callback=status_callback,
)
if result:
remove_release(task.task_id)
self._cleanup_refs[task.task_id] = (client, download_id, protocol)
return result
# Existing but still downloading - join the progress polling
logger.info(f"Existing download in progress, joining poll loop")
status_callback("downloading", "Resuming existing download")
else:
# No existing download - add new
status_callback("resolving", f"Sending to {client.name}")
try:
release_name = prowlarr_result.get("title") or task.title or "Unknown"
category = self._get_category_for_task(client, task)
expected_hash = str(prowlarr_result.get("infoHash") or "").strip() or None
download_id = client.add_download(
url=download_url,
name=release_name,
category=category,
expected_hash=expected_hash,
)
except Exception as e:
logger.error(f"Failed to add to {client.name}: {e}")
status_callback("error", f"Failed to add to {client.name}: {e}")
return None
logger.info(f"Added to {client.name}: {download_id} for '{release_name}'")
# Poll for progress
return self._poll_and_complete(
client=client,
download_id=download_id,
protocol=protocol,
task=task,
cancel_flag=cancel_flag,
progress_callback=progress_callback,
status_callback=status_callback,
)
except Exception as e:
logger.error(f"Prowlarr download error: {e}")
status_callback("error", str(e))
return None
def _poll_and_complete(
self,
client,
download_id: str,
protocol: str,
task: DownloadTask,
cancel_flag: Event,
progress_callback: Callable[[float], None],
status_callback: Callable[[str, Optional[str]], None],
) -> Optional[str]:
"""Poll the download client for progress and handle completion."""
# Track consecutive "not found" errors - torrents may take time to appear in client
not_found_count = 0
max_not_found_retries = 15 # 15 retries * 2s poll = 30s grace period
try:
logger.debug(f"Starting poll for {download_id} (content_type={task.content_type})")
while not cancel_flag.is_set():
status = client.get_status(download_id)
progress_callback(status.progress)
# Check for completion
if status.complete:
if status.state == DownloadState.ERROR:
logger.error(f"Download {download_id} completed with error: {status.message}")
status_callback("error", status.message or "Download failed")
self._safe_remove_download(client, download_id, protocol, "completion error")
return None
# Download complete - break to handle file
logger.debug(f"Download {download_id} complete, file_path={status.file_path}")
break
# Check for error state
if status.state == DownloadState.ERROR:
message = (status.message or "").strip()
message_lower = message.lower()
# Only treat *actual* "not found" as retryable.
# qBittorrent auth/network/API failures should surface immediately (more actionable)
# and must not be confused with "torrent missing".
retryable_not_found = any(
token in message_lower
for token in (
"torrent not found",
"not found in qbittorrent",
"download not found",
)
)
non_retryable = any(
token in message_lower
for token in (
"authentication failed",
"cannot connect",
"timed out",
"api request failed",
)
)
if retryable_not_found and not non_retryable:
not_found_count += 1
if not_found_count < max_not_found_retries:
logger.debug(
f"Download {download_id} not yet visible in client "
f"(attempt {not_found_count}/{max_not_found_retries})"
)
status_callback("resolving", "Waiting for download client...")
if cancel_flag.wait(timeout=POLL_INTERVAL):
break
continue
logger.error(
f"Download {download_id} not found after {max_not_found_retries} attempts"
)
else:
# Fail fast on actionable errors (auth, connectivity, API issues)
logger.error(f"Download {download_id} error state: {status.message}")
status_callback("error", status.message or "Download failed")
self._safe_remove_download(client, download_id, protocol, "download error")
return None
# Reset not-found counter on successful status check
not_found_count = 0
# Build status message - use client message if provided, else build progress
msg = status.message or self._build_progress_message(status)
if status.state == DownloadState.PROCESSING:
# Post-processing (e.g., SABnzbd verifying/extracting)
status_callback("resolving", msg)
else:
status_callback("downloading", msg)
# Wait for next poll (interruptible by cancel)
if cancel_flag.wait(timeout=POLL_INTERVAL):
break
# Handle cancellation
if cancel_flag.is_set():
if protocol == "usenet":
logger.info(f"Download cancelled, removing from {client.name}: {download_id}")
try:
self._delete_local_download_data(client, download_id)
self._remove_usenet_download(client, download_id, delete_files=True, archive=True)
except Exception as e:
logger.warning(
f"Failed to remove download {download_id} from {client.name} after cancellation: {e}"
)
else:
logger.info(
f"Download cancelled for protocol={protocol}; leaving in {client.name}: {download_id}"
)
status_callback("cancelled", "Cancelled")
return None
# Handle completed file
source_path = client.get_download_path(download_id)
if not source_path:
logger.error(
f"Download client returned empty path for completed download. "
f"Client: {client.name}, ID: {download_id}. "
f"Check that the download client's completion folder is accessible to Shelfmark."
)
status_callback(
"error",
f"Could not locate completed download in {client.name} (path not returned). "
f"Check volume mappings and category settings."
)
return None
# Apply remote path mappings (client path -> shelfmark container path)
from shelfmark.core.path_mappings import (
get_client_host_identifier,
parse_remote_path_mappings,
remap_remote_to_local_with_match,
)
source_path_obj = Path(source_path)
host = get_client_host_identifier(client) or ""
mapping_value = config.get("PROWLARR_REMOTE_PATH_MAPPINGS", [])
mappings = parse_remote_path_mappings(mapping_value)
logger.debug(
"Attempting path remap: client=%s, host=%s, path=%s, mappings=%s",
client.name,
host,
source_path_obj,
[(m.host, m.remote_path, m.local_path) for m in mappings],
)
remapped, matched_mapping = remap_remote_to_local_with_match(
mappings=mappings,
host=host,
remote_path=source_path_obj,
)
logger.debug(
"Remap result: %s -> %s (exists=%s, changed=%s, matched=%s)",
source_path_obj,
remapped,
remapped.exists(),
remapped != source_path_obj,
matched_mapping,
)
if matched_mapping:
if remapped.exists():
logger.info(
"Remapped download path for %s (%s): %s -> %s",
client.name,
download_id,
source_path_obj,
remapped,
)
source_path_obj = remapped
else:
logger.error(
f"Download path does not exist after remapping: {source_path} -> {remapped}. "
f"Client: {client.name}, ID: {download_id}. "
f"Check that the local path in your mapping is mounted correctly."
)
status_callback(
"error",
f"Remapped path '{remapped}' does not exist. "
f"Check your Docker volume mounts match the Local Path in Settings > Advanced > Remote Path Mappings.",
)
return None
elif mappings:
if source_path_obj.exists():
logger.info(
"No remote path mapping matched for %s (%s); using client path: %s",
client.name,
download_id,
source_path_obj,
)
else:
hint = _diagnose_path_issue(source_path)
logger.error(
f"Download path does not exist and no remote path mapping matched for {client.name} "
f"({download_id}): {source_path}. {hint}"
)
status_callback(
"error",
f"{hint} No remote path mapping matched for client '{client.name}'.",
)
return None
elif not source_path_obj.exists():
hint = _diagnose_path_issue(source_path)
logger.error(
f"Download path does not exist: {source_path}. "
f"Client: {client.name}, ID: {download_id}. {hint}"
)
status_callback("error", hint)
return None
result = self._handle_completed_file(
source_path=source_path_obj,
protocol=protocol,
task=task,
status_callback=status_callback,
)
# Clean up on success
if result:
remove_release(task.task_id)
self._cleanup_refs[task.task_id] = (client, download_id, protocol)
return result
except Exception as e:
logger.error(f"Error during download polling: {e}")
status_callback("error", str(e))
self._safe_remove_download(client, download_id, protocol, "polling exception")
return None
def _handle_completed_file(
self,
source_path: Path,
protocol: str,
task: DownloadTask,
status_callback: Callable[[str, Optional[str]], None],
) -> Optional[str]:
"""Handle a completed download and return its path.
For external download clients (torrents/usenet), staging large payloads into TMP_DIR
is expensive (and can duplicate multi-GB files). Instead, return the client's
completed path and let the orchestrator perform any required transfer (copy/move/
hardlink) directly from that source.
Torrents also set ``task.original_download_path`` so the orchestrator can detect
seeding data and enable hardlinking when configured.
"""
try:
if protocol == "torrent":
task.original_download_path = str(source_path)
logger.debug(f"Download complete, returning original path: {source_path}")
return str(source_path)
except Exception as e:
logger.error(f"Failed to finalize completed download at {source_path}: {e}")
status_callback("error", f"Failed to finalize completed download: {e}")
return None
def _on_download_complete(self, task: DownloadTask) -> None:
"""Remove completed release from the Prowlarr cache."""
remove_release(task.task_id)
def cancel(self, task_id: str) -> bool:
"""Cancel download and clean up cache. Primary cancellation is via cancel_flag."""
logger.debug(f"Cancel requested for Prowlarr task: {task_id}")
# Remove from cache if present
remove_release(task_id)
return True
return super().cancel(task_id)
+7 -643
View File
@@ -1,22 +1,14 @@
"""
Prowlarr settings registration.
Registers Prowlarr settings as a group with multiple tabs:
- Configuration: Prowlarr connection settings + indexer selection
- Download Clients: Torrent and usenet client settings
"""
"""Prowlarr settings registration."""
from typing import Any, Dict, List, Optional
from shelfmark.core.settings_registry import (
register_group,
register_settings,
CheckboxField,
HeadingField,
TextField,
PasswordField,
ActionButton,
SelectField,
MultiSelectField,
)
from shelfmark.core.utils import normalize_http_url
@@ -24,6 +16,7 @@ from shelfmark.core.utils import normalize_http_url
# ==================== Dynamic Options Loaders ====================
def _get_indexer_options() -> List[Dict[str, str]]:
"""
Fetch available indexers from Prowlarr for the multi-select field.
@@ -75,15 +68,14 @@ def _get_indexer_options() -> List[Dict[str, str]]:
return []
# ==================== Test Connection Callbacks ====================
# ==================== Test Connection Callback ====================
def _test_prowlarr_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test the Prowlarr connection using current form values."""
from shelfmark.core.config import config
from shelfmark.core.logger import setup_logger
from shelfmark.release_sources.prowlarr.api import ProwlarrClient
logger = setup_logger(__name__)
current_values = current_values or {}
raw_url = current_values.get("PROWLARR_URL") or config.get("PROWLARR_URL", "")
@@ -106,308 +98,14 @@ def _test_prowlarr_connection(current_values: Optional[Dict[str, Any]] = None) -
return {"success": False, "message": f"Connection failed: {str(e)}"}
def _test_qbittorrent_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test the qBittorrent connection using current form values."""
from shelfmark.core.config import config
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", "")
if not raw_url:
return {"success": False, "message": "qBittorrent URL is required"}
try:
from qbittorrentapi import Client
url = normalize_http_url(raw_url)
if not url:
return {"success": False, "message": "qBittorrent URL is invalid"}
client = Client(host=url, username=username, password=password)
client.auth_log_in()
api_version = client.app.web_api_version
return {"success": True, "message": f"Connected to qBittorrent (API v{api_version})"}
except ImportError:
return {"success": False, "message": "qbittorrent-api package not installed"}
except Exception as e:
return {"success": False, "message": f"Connection failed: {str(e)}"}
def _test_transmission_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test the Transmission connection using current form values."""
from shelfmark.core.config import config
from shelfmark.release_sources.prowlarr.clients.torrent_utils import (
parse_transmission_url,
)
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", "")
if not raw_url:
return {"success": False, "message": "Transmission URL is required"}
url = normalize_http_url(raw_url)
if not url:
return {"success": False, "message": "Transmission URL is invalid"}
try:
from transmission_rpc import Client
# Parse URL to extract host, port, and path
host, port, path = parse_transmission_url(url)
client = Client(
host=host,
port=port,
path=path,
username=username if username else None,
password=password if password else None,
)
session = client.get_session()
version = session.version
return {"success": True, "message": f"Connected to Transmission {version}"}
except ImportError:
return {"success": False, "message": "transmission-rpc package not installed"}
except Exception as e:
return {"success": False, "message": f"Connection failed: {str(e)}"}
def _test_deluge_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test Deluge Web UI JSON-RPC connection using current form values."""
from urllib.parse import urlparse
import requests
from shelfmark.core.config import config
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", "")
if not raw_host:
return {"success": False, "message": "Deluge host is required"}
if not password:
return {"success": False, "message": "Deluge password is required"}
raw_host = str(raw_host)
raw_host = normalize_http_url(raw_host, strip_trailing_slash=False) if raw_host else ""
if not raw_host:
return {"success": False, "message": "Deluge host is invalid"}
raw_port = str(raw_port or "8112")
scheme = "http"
base_path = ""
host = raw_host
port = int(raw_port) if raw_port.isdigit() else 8112
# Allow DELUGE_HOST to be a full URL (e.g. http://deluge:8112)
if raw_host.startswith(("http://", "https://")):
parsed = urlparse(raw_host)
scheme = parsed.scheme or "http"
host = parsed.hostname or "localhost"
if parsed.port is not None:
port = parsed.port
base_path = (parsed.path or "").rstrip("/")
else:
# Allow "host:port" in DELUGE_HOST for convenience.
if ":" in raw_host and raw_host.count(":") == 1:
host_part, port_part = raw_host.split(":", 1)
if host_part and port_part.isdigit():
host = host_part
port = int(port_part)
rpc_url = f"{scheme}://{host}:{port}{base_path}/json"
def rpc_call(session: requests.Session, rpc_id: int, method: str, *params: Any) -> Any:
payload = {"id": rpc_id, "method": method, "params": list(params)}
resp = session.post(rpc_url, json=payload, timeout=15)
resp.raise_for_status()
data = resp.json()
if data.get("error"):
error = data["error"]
if isinstance(error, dict):
raise Exception(error.get("message") or str(error))
raise Exception(str(error))
return data.get("result")
def get_daemon_version(session: requests.Session, rpc_id: int) -> Any:
try:
methods = rpc_call(session, rpc_id, "system.listMethods")
if isinstance(methods, list) and "daemon.get_version" in methods:
return rpc_call(session, rpc_id + 1, "daemon.get_version")
except Exception:
# Fall back to daemon.info to preserve existing behavior.
pass
return rpc_call(session, rpc_id + 1, "daemon.info")
try:
session = requests.Session()
if rpc_call(session, 1, "auth.login", password) is not True:
return {"success": False, "message": "Deluge Web UI authentication failed"}
if rpc_call(session, 2, "web.connected") is not True:
hosts = rpc_call(session, 3, "web.get_hosts") or []
if not hosts:
return {
"success": False,
"message": "Deluge Web UI isn't connected to Deluge core (no hosts configured). Add/connect a daemon in Deluge Web UI → Connection Manager.",
}
host_id = hosts[0][0]
for entry in hosts:
if isinstance(entry, list) and len(entry) >= 2 and entry[1] in {"127.0.0.1", "localhost"}:
host_id = entry[0]
break
rpc_call(session, 4, "web.connect", host_id)
if rpc_call(session, 5, "web.connected") is not True:
return {
"success": False,
"message": "Deluge Web UI couldn't connect to Deluge core. Check Deluge Web UI → Connection Manager.",
}
version = get_daemon_version(session, 6)
return {"success": True, "message": f"Connected to Deluge {version}"}
except requests.exceptions.ConnectionError:
return {"success": False, "message": "Could not connect to Deluge Web UI"}
except requests.exceptions.Timeout:
return {"success": False, "message": "Connection timed out"}
except Exception as e:
return {"success": False, "message": f"Connection failed: {str(e)}"}
def _test_rtorrent_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test the rTorrent connection using current form values."""
from shelfmark.core.config import config
from urllib.parse import urlparse
from xmlrpc.client import ServerProxy
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", "")
if not raw_url:
return {"success": False, "message": "rTorrent URL is required"}
url = normalize_http_url(raw_url)
if not url:
return {"success": False, "message": "rTorrent URL is invalid"}
try:
# Add HTTP auth to URL if credentials provided
if username and password:
parsed = urlparse(url)
url = f"{parsed.scheme}://{username}:{password}@{parsed.netloc}{parsed.path}"
rpc = ServerProxy(url.rstrip("/"))
version = rpc.system.client_version()
return {"success": True, "message": f"Connected to rTorrent {version}"}
except Exception as e:
return {"success": False, "message": f"Connection failed: {str(e)}"}
def _test_nzbget_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test the NZBGet connection using current form values."""
import requests
from shelfmark.core.config import config
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", "")
if not raw_url:
return {"success": False, "message": "NZBGet URL is required"}
url = normalize_http_url(raw_url)
if not url:
return {"success": False, "message": "NZBGet URL is invalid"}
try:
rpc_url = f"{url.rstrip('/')}/jsonrpc"
payload = {"jsonrpc": "2.0", "method": "status", "params": [], "id": 1}
response = requests.post(rpc_url, json=payload, auth=(username, password), timeout=30)
response.raise_for_status()
result = response.json()
if "error" in result and result["error"]:
raise Exception(result["error"].get("message", "RPC error"))
version = result.get("result", {}).get("Version", "unknown")
return {"success": True, "message": f"Connected to NZBGet {version}"}
except requests.exceptions.ConnectionError:
return {"success": False, "message": "Could not connect to NZBGet"}
except requests.exceptions.Timeout:
return {"success": False, "message": "Connection timed out"}
except Exception as e:
return {"success": False, "message": f"Connection failed: {str(e)}"}
def _test_sabnzbd_connection(current_values: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Test the SABnzbd connection using current form values."""
import requests
from shelfmark.core.config import config
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", "")
if not raw_url:
return {"success": False, "message": "SABnzbd URL is required"}
url = normalize_http_url(raw_url)
if not url:
return {"success": False, "message": "SABnzbd URL is invalid"}
if not api_key:
return {"success": False, "message": "API key is required"}
try:
api_url = f"{url.rstrip('/')}/api"
params = {"apikey": api_key, "mode": "version", "output": "json"}
response = requests.get(api_url, params=params, timeout=30)
response.raise_for_status()
result = response.json()
version = result.get("version", "unknown")
return {"success": True, "message": f"Connected to SABnzbd {version}"}
except requests.exceptions.ConnectionError:
return {"success": False, "message": "Could not connect to SABnzbd"}
except requests.exceptions.Timeout:
return {"success": False, "message": "Connection timed out"}
except Exception as e:
return {"success": False, "message": f"Connection failed: {str(e)}"}
# ==================== Register Group ====================
register_group(
name="prowlarr",
display_name="Prowlarr",
icon="download",
order=40,
)
# ==================== Configuration Tab ====================
@register_settings(
name="prowlarr_config",
display_name="Configuration",
display_name="Prowlarr",
icon="download",
order=41,
group="prowlarr",
)
def prowlarr_config_settings():
"""Prowlarr connection and indexer settings."""
@@ -464,337 +162,3 @@ def prowlarr_config_settings():
show_when={"field": "PROWLARR_ENABLED", "value": True},
),
]
# ==================== Download Clients Tab ====================
@register_settings(
name="prowlarr_clients",
display_name="Download Clients",
order=42,
group="prowlarr",
)
def prowlarr_clients_settings():
"""Download client settings for Prowlarr."""
return [
# --- Torrent Client Selection ---
HeadingField(
key="torrent_heading",
title="Torrent Client",
description="Select and configure a torrent client for downloading torrents from Prowlarr.",
),
SelectField(
key="PROWLARR_TORRENT_CLIENT",
label="Torrent Client",
description="Choose which torrent client to use",
options=[
{"value": "", "label": "None"},
{"value": "qbittorrent", "label": "qBittorrent"},
{"value": "transmission", "label": "Transmission"},
{"value": "deluge", "label": "Deluge"},
{"value": "rtorrent", "label": "rTorrent"},
],
default="",
),
# --- qBittorrent Settings ---
TextField(
key="QBITTORRENT_URL",
label="qBittorrent URL",
description="Web UI URL of your qBittorrent instance",
placeholder="http://qbittorrent:8080",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
TextField(
key="QBITTORRENT_USERNAME",
label="Username",
description="qBittorrent Web UI username",
placeholder="admin",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
PasswordField(
key="QBITTORRENT_PASSWORD",
label="Password",
description="qBittorrent Web UI password",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
ActionButton(
key="test_qbittorrent",
label="Test Connection",
description="Verify your qBittorrent configuration",
style="primary",
callback=_test_qbittorrent_connection,
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
TextField(
key="QBITTORRENT_CATEGORY",
label="Book Category",
description="Category to assign to book downloads in qBittorrent",
placeholder="books",
default="books",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
TextField(
key="QBITTORRENT_CATEGORY_AUDIOBOOK",
label="Audiobook Category",
description="Category for audiobook downloads. Leave empty to use the book category.",
placeholder="",
default="",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "qbittorrent"},
),
# --- Transmission Settings ---
TextField(
key="TRANSMISSION_URL",
label="Transmission URL",
description="URL of your Transmission instance",
placeholder="http://transmission:9091",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
TextField(
key="TRANSMISSION_USERNAME",
label="Username",
description="Transmission RPC username (if authentication enabled)",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
PasswordField(
key="TRANSMISSION_PASSWORD",
label="Password",
description="Transmission RPC password",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
ActionButton(
key="test_transmission",
label="Test Connection",
description="Verify your Transmission configuration",
style="primary",
callback=_test_transmission_connection,
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
TextField(
key="TRANSMISSION_CATEGORY",
label="Book Label",
description="Label to assign to book downloads in Transmission",
placeholder="books",
default="books",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
TextField(
key="TRANSMISSION_CATEGORY_AUDIOBOOK",
label="Audiobook Label",
description="Label for audiobook downloads. Leave empty to use the book label.",
placeholder="",
default="",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "transmission"},
),
# --- Deluge Settings ---
TextField(
key="DELUGE_HOST",
label="Deluge Web UI Host/URL",
description="Hostname/IP or full URL of your Deluge Web UI (deluge-web)",
placeholder="http://deluge:8112",
default="localhost",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
TextField(
key="DELUGE_PORT",
label="Deluge Web UI Port",
description="Deluge Web UI port (default: 8112)",
placeholder="8112",
default="8112",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
PasswordField(
key="DELUGE_PASSWORD",
label="Password",
description="Deluge Web UI password (default: deluge)",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
ActionButton(
key="test_deluge",
label="Test Connection",
description="Verify your Deluge configuration",
style="primary",
callback=_test_deluge_connection,
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
TextField(
key="DELUGE_CATEGORY",
label="Book Label",
description="Label to assign to book downloads in Deluge",
placeholder="books",
default="books",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
TextField(
key="DELUGE_CATEGORY_AUDIOBOOK",
label="Audiobook Label",
description="Label for audiobook downloads. Leave empty to use the book label.",
placeholder="",
default="",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "deluge"},
),
# --- rTorrent Settings ---
TextField(
key="RTORRENT_URL",
label="rTorrent URL",
description="XML-RPC URL of your rTorrent instance",
placeholder="http://rtorrent:6881/RPC2",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "rtorrent"},
),
TextField(
key="RTORRENT_USERNAME",
label="Username",
description="HTTP Basic auth username (if authentication enabled)",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "rtorrent"},
),
PasswordField(
key="RTORRENT_PASSWORD",
label="Password",
description="HTTP Basic auth password",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "rtorrent"},
),
ActionButton(
key="test_rtorrent",
label="Test Connection",
description="Verify your rTorrent configuration",
style="primary",
callback=_test_rtorrent_connection,
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "rtorrent"},
),
TextField(
key="RTORRENT_LABEL",
label="Book Label",
description="Label to assign to book downloads in rTorrent",
placeholder="cwabd",
default="cwabd",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "rtorrent"},
),
TextField(
key="RTORRENT_DOWNLOAD_DIR",
label="Download Directory",
description="Server-side directory where torrents are downloaded (optional, uses rTorrent default if not specified)",
placeholder="/downloads",
show_when={"field": "PROWLARR_TORRENT_CLIENT", "value": "rtorrent"},
),
# Note: Torrent client download path must be mounted identically in both containers.
# Torrents are always copied (not moved) to preserve seeding capability.
# --- Usenet Client Selection ---
HeadingField(
key="usenet_heading",
title="Usenet Client",
description="Select and configure a usenet client for downloading NZBs from Prowlarr.",
),
SelectField(
key="PROWLARR_USENET_CLIENT",
label="Usenet Client",
description="Choose which usenet client to use",
options=[
{"value": "", "label": "None"},
{"value": "nzbget", "label": "NZBGet"},
{"value": "sabnzbd", "label": "SABnzbd"},
],
default="",
),
# --- NZBGet Settings ---
TextField(
key="NZBGET_URL",
label="NZBGet URL",
description="URL of your NZBGet instance",
placeholder="http://nzbget:6789",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "nzbget"},
),
TextField(
key="NZBGET_USERNAME",
label="Username",
description="NZBGet control username",
placeholder="nzbget",
default="nzbget",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "nzbget"},
),
PasswordField(
key="NZBGET_PASSWORD",
label="Password",
description="NZBGet control password",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "nzbget"},
),
ActionButton(
key="test_nzbget",
label="Test Connection",
description="Verify your NZBGet configuration",
style="primary",
callback=_test_nzbget_connection,
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "nzbget"},
),
TextField(
key="NZBGET_CATEGORY",
label="Book Category",
description="Category to assign to book downloads in NZBGet",
placeholder="Books",
default="Books",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "nzbget"},
),
TextField(
key="NZBGET_CATEGORY_AUDIOBOOK",
label="Audiobook Category",
description="Category for audiobook downloads. Leave empty to use the book category.",
placeholder="",
default="",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "nzbget"},
),
# --- SABnzbd Settings ---
TextField(
key="SABNZBD_URL",
label="SABnzbd URL",
description="URL of your SABnzbd instance",
placeholder="http://sabnzbd:8080",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "sabnzbd"},
),
PasswordField(
key="SABNZBD_API_KEY",
label="API Key",
description="Found in SABnzbd: Config > General > API Key",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "sabnzbd"},
),
ActionButton(
key="test_sabnzbd",
label="Test Connection",
description="Verify your SABnzbd configuration",
style="primary",
callback=_test_sabnzbd_connection,
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "sabnzbd"},
),
TextField(
key="SABNZBD_CATEGORY",
label="Book Category",
description="Category to assign to book downloads in SABnzbd",
placeholder="books",
default="books",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "sabnzbd"},
),
TextField(
key="SABNZBD_CATEGORY_AUDIOBOOK",
label="Audiobook Category",
description="Category for audiobook downloads. Leave empty to use the book category.",
placeholder="",
default="",
show_when={"field": "PROWLARR_USENET_CLIENT", "value": "sabnzbd"},
),
# Note: Usenet client download path must be mounted identically in both containers.
SelectField(
key="PROWLARR_USENET_ACTION",
label="NZB Completion Action",
description="Move deletes the job from your usenet client after import; Copy keeps it in the client",
options=[
{"value": "move", "label": "Move"},
{"value": "copy", "label": "Copy"},
],
default="move",
show_when={"field": "PROWLARR_USENET_CLIENT", "notEmpty": True},
),
]
+352 -69
View File
@@ -59,6 +59,51 @@ AUDIOBOOK_FORMATS = ["m4b", "mp3", "m4a", "flac", "ogg", "wma", "aac", "wav", "o
# Combined list for format detection (audiobook formats first for priority)
ALL_BOOK_FORMATS = AUDIOBOOK_FORMATS + EBOOK_FORMATS
# Map 3-char MAM language codes to 2-char ISO codes used by frontend color maps
MAM_LANGUAGE_MAP = {
"eng": "en",
"ita": "it",
"spa": "es",
"fra": "fr",
"fre": "fr",
"ger": "de",
"deu": "de",
"por": "pt",
"rus": "ru",
"jpn": "ja",
"jap": "ja",
"chi": "zh",
"zho": "zh",
"dut": "nl",
"nld": "nl",
"swe": "sv",
"nor": "no",
"dan": "da",
"fin": "fi",
"pol": "pl",
"cze": "cs",
"ces": "cs",
"hun": "hu",
"kor": "ko",
"ara": "ar",
"heb": "he",
"tur": "tr",
"gre": "el",
"ell": "el",
"hin": "hi",
"tha": "th",
"vie": "vi",
"ind": "id",
"ukr": "uk",
"rom": "ro",
"ron": "ro",
"bul": "bg",
"cat": "ca",
"hrv": "hr",
"slv": "sl",
"srp": "sr",
}
# Backend safeguard: cap total Prowlarr search time per request.
PROWLARR_SEARCH_TIMEOUT_SECONDS = 120.0
@@ -83,33 +128,79 @@ def _extract_format(title: str) -> Optional[str]:
return None
def _extract_language(title: str) -> Optional[str]:
"""Extract language code from release title (e.g., [German] -> 'de')."""
title_lower = title.lower()
def _extract_mam_language(raw_title: str) -> Optional[str]:
"""
Extract the language code from MyAnonamouse titles.
# Common language names and their codes
languages = {
"english": "en", "eng": "en", "[en]": "en", "(en)": "en",
"german": "de", "deutsch": "de", "[de]": "de", "(de)": "de", "ger": "de",
"french": "fr", "français": "fr", "[fr]": "fr", "(fr)": "fr", "fra": "fr",
"spanish": "es", "español": "es", "[es]": "es", "(es)": "es", "spa": "es",
"italian": "it", "italiano": "it", "[it]": "it", "(it)": "it", "ita": "it",
"portuguese": "pt", "[pt]": "pt", "(pt)": "pt", "por": "pt",
"dutch": "nl", "nederlands": "nl", "[nl]": "nl", "(nl)": "nl", "nld": "nl",
"russian": "ru", "[ru]": "ru", "(ru)": "ru", "rus": "ru",
"polish": "pl", "polski": "pl", "[pl]": "pl", "(pl)": "pl", "pol": "pl",
"chinese": "zh", "[zh]": "zh", "(zh)": "zh", "chi": "zh",
"japanese": "ja", "[ja]": "ja", "(ja)": "ja", "jpn": "ja",
"korean": "ko", "[ko]": "ko", "(ko)": "ko", "kor": "ko",
}
Prowlarr's MAM parser appends a structured bracket segment like:
[ENG / EPUB MOBI PDF]
for lang_pattern, lang_code in languages.items():
if lang_pattern in title_lower:
return lang_code
The language code appears before the "/" - we extract it and map to
the 2-char ISO code used by the frontend color maps.
"""
if not raw_title:
return None
for bracket in re.findall(r"\[([^\]]+)\]", raw_title):
if "/" not in bracket:
continue
before_slash, _ = bracket.split("/", 1)
# Extract the language token (should be a 3-char code like ENG, ITA, etc.)
tokens = re.findall(r"[A-Za-z]+", before_slash.strip())
for token in tokens:
lang_code = token.lower()
if lang_code in MAM_LANGUAGE_MAP:
return MAM_LANGUAGE_MAP[lang_code]
return None
def _extract_mam_formats(raw_title: str) -> List[str]:
"""
Extract a list of formats from MyAnonamouse titles.
Prowlarr's MAM parser appends a structured bracket segment like:
[ENG / EPUB MOBI PDF]
We only trust this structured segment (and do not attempt generic title
heuristics for other indexers).
"""
if not raw_title:
return []
format_set = set(ALL_BOOK_FORMATS)
for bracket in re.findall(r"\[([^\]]+)\]", raw_title):
if "/" not in bracket:
continue
_, after_slash = bracket.split("/", 1)
tokens = re.findall(r"[A-Za-z0-9]+", after_slash)
formats: List[str] = []
for token in tokens:
fmt = token.lower()
if fmt in format_set and fmt not in formats:
formats.append(fmt)
if formats:
return formats
return []
def _formats_display(formats: List[str]) -> Optional[str]:
if not formats:
return None
if len(formats) == 1:
return formats[0]
if len(formats) == 2:
return f"{formats[0]}, {formats[1]}"
# Show first two formats + count of others to prevent overflow
return f"{formats[0]}, {formats[1]} +{len(formats) - 2}"
# Prowlarr category IDs for content type detection
# See: https://wiki.servarr.com/prowlarr/cardigann-yml-definition#categories
AUDIOBOOK_CATEGORY_IDS = {3000, 3030} # 3000 = Audio, 3030 = Audio/Audiobook
@@ -144,9 +235,15 @@ def _detect_content_type_from_categories(categories: list, fallback: str = "book
return "other"
def _prowlarr_result_to_release(result: dict, search_content_type: str = "ebook") -> Release:
def _prowlarr_result_to_release(
result: dict,
search_content_type: str = "ebook",
*,
enable_format_detection: bool = False,
) -> Release:
"""Convert a Prowlarr API result to a Release object."""
title = result.get("title", "Unknown")
raw_title = result.get("title", "Unknown")
title = raw_title
size_bytes = result.get("size")
indexer = result.get("indexer", "Unknown")
protocol = get_protocol(result)
@@ -154,6 +251,27 @@ def _prowlarr_result_to_release(result: dict, search_content_type: str = "ebook"
leechers = result.get("leechers")
categories = result.get("categories", [])
is_torrent = protocol == ReleaseProtocol.TORRENT
raw_indexer_flags = result.get("indexerFlags") or []
indexer_flags: List[str] = []
seen_flags: set[str] = set()
def add_indexer_flag(flag: object) -> None:
if flag is None:
return
flag_str = str(flag).strip()
if not flag_str:
return
lowered = flag_str.lower()
if lowered in seen_flags:
return
seen_flags.add(lowered)
indexer_flags.append(flag_str)
if isinstance(raw_indexer_flags, list):
for flag in raw_indexer_flags:
add_indexer_flag(flag)
elif isinstance(raw_indexer_flags, str):
add_indexer_flag(raw_indexer_flags)
# Format peers display string: "seeders / leechers"
peers_display = (
@@ -162,22 +280,50 @@ def _prowlarr_result_to_release(result: dict, search_content_type: str = "ebook"
else None
)
# For format detection, prefer fileName over title (often cleaner)
file_name = result.get("fileName", "")
format_detected = _extract_format(file_name) if file_name else _extract_format(title)
format_detected: Optional[str] = None
formats: List[str] = []
formats_display: Optional[str] = None
language_detected: Optional[str] = None
if enable_format_detection:
book_title = str(result.get("bookTitle") or "").strip()
if book_title:
title = book_title
formats = _extract_mam_formats(str(raw_title or ""))
format_detected = formats[0] if formats else None
formats_display = _formats_display(formats)
language_detected = _extract_mam_language(str(raw_title or ""))
# Build the source_id from GUID or generate from indexer + title
source_id = result.get("guid") or f"{indexer}:{hash(title)}"
source_id = result.get("guid") or f"{indexer}:{hash(raw_title)}"
# Cache the raw Prowlarr result so handler can look it up by source_id
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
if any(flag.lower() in {"freeleech", "fl"} for flag in indexer_flags):
is_freeleech = True
is_vip = "[vip]" in str(raw_title).lower()
if is_vip:
add_indexer_flag("VIP")
if is_freeleech:
add_indexer_flag("FreeLeech")
return Release(
source="prowlarr",
source_id=source_id,
title=title,
format=format_detected,
language=_extract_language(title),
language=language_detected,
size=_parse_size(size_bytes),
size_bytes=size_bytes,
download_url=get_preferred_download_url(result),
@@ -199,6 +345,20 @@ def _prowlarr_result_to_release(result: dict, search_content_type: str = "ebook"
"indexer_id": result.get("indexerId"),
"files": result.get("files"),
"grabs": result.get("grabs"),
"author": result.get("author"),
"book_title": result.get("bookTitle"),
"indexer_flags": indexer_flags,
"vip": is_vip,
"freeleech": is_freeleech,
"download_volume_factor": result.get("downloadVolumeFactor"),
"upload_volume_factor": result.get("uploadVolumeFactor"),
"minimum_ratio": result.get("minimumRatio"),
"minimum_seed_time": result.get("minimumSeedTime"),
"info_hash": result.get("infoHash"),
"formats": formats if formats else None,
"formats_display": formats_display,
# Raw torznab attributes for rich tooltips (enriched indexers)
"torznab_attrs": result.get("torznabAttrs"),
},
)
@@ -216,48 +376,85 @@ class ProwlarrSource(ReleaseSource):
def get_column_config(self) -> ReleaseColumnConfig:
"""Column configuration for Prowlarr releases."""
# Fetch available indexers from Prowlarr
available_indexers: Optional[List[str]] = None
default_indexers: Optional[List[str]] = None
client = self._get_client()
if client:
try:
enabled_indexers = client.get_enabled_indexers_detailed()
# Get user-selected indexer IDs if configured
selected_ids = self._get_selected_indexer_ids()
all_indexer_names = []
selected_indexer_names = []
for idx in enabled_indexers:
idx_id = idx.get("id")
idx_name = idx.get("name")
if not idx_name:
continue
# Add to all indexers list
all_indexer_names.append(idx_name)
# 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
available_indexers = sorted(all_indexer_names) if all_indexer_names else None
# Only set default_indexers if user has selected specific ones
default_indexers = sorted(selected_indexer_names) if selected_indexer_names else None
except Exception as e:
logger.warning(f"Failed to fetch indexer list for column config: {e}")
return ReleaseColumnConfig(
columns=[
ColumnSchema(
key="indexer",
label="Indexer",
render_type=ColumnRenderType.TEXT,
render_type=ColumnRenderType.INDEXER_PROTOCOL,
align=ColumnAlign.LEFT,
width="minmax(80px, 1fr)",
hide_mobile=True,
width="minmax(140px, 1fr)",
hide_mobile=False,
sortable=True,
),
ColumnSchema(
key="protocol",
label="Type",
render_type=ColumnRenderType.BADGE,
key="extra.indexer_flags",
label="Flags",
render_type=ColumnRenderType.TAGS,
align=ColumnAlign.CENTER,
width="60px",
width="50px",
hide_mobile=False,
color_hint=ColumnColorHint(type="map", value="download_type"),
color_hint=ColumnColorHint(type="map", value="flags"),
fallback="",
uppercase=True,
),
ColumnSchema(
key="peers",
label="Peers",
render_type=ColumnRenderType.PEERS,
key="language",
label="Lang",
render_type=ColumnRenderType.BADGE,
align=ColumnAlign.CENTER,
width="70px",
width="50px",
hide_mobile=True,
fallback="-",
sortable=True,
sort_key="seeders",
color_hint=ColumnColorHint(type="map", value="language"),
uppercase=True,
fallback="",
),
ColumnSchema(
key="content_type",
label="Type",
render_type=ColumnRenderType.BADGE,
key="extra.formats_display",
label="Format",
render_type=ColumnRenderType.FORMAT_CONTENT_TYPE,
align=ColumnAlign.CENTER,
width="90px",
hide_mobile=False,
color_hint=ColumnColorHint(type="map", value="content_type"),
color_hint=ColumnColorHint(type="map", value="format"),
uppercase=True,
fallback="-",
fallback="",
),
ColumnSchema(
key="size",
@@ -270,9 +467,11 @@ class ProwlarrSource(ReleaseSource):
sort_key="size_bytes",
),
],
grid_template="minmax(0,2fr) minmax(80px,1fr) 60px 70px 90px 80px",
grid_template="minmax(0,2fr) minmax(140px,1fr) 50px 50px 90px 80px",
leading_cell=LeadingCellConfig(type=LeadingCellType.NONE), # No leading cell for Prowlarr
supported_filters=["language"], # Enables multi-language query expansion; Prowlarr language metadata is unreliable
available_indexers=available_indexers,
default_indexers=default_indexers,
supported_filters=["language", "indexer"], # Enables multi-language query expansion and indexer filtering
)
def _get_client(self) -> Optional[ProwlarrClient]:
@@ -313,6 +512,39 @@ class ProwlarrSource(ReleaseSource):
logger.warning(f"Invalid PROWLARR_INDEXERS format: {selected} ({e})")
return None
def _resolve_indexer_ids_from_names(
self, client: ProwlarrClient, names: List[str]
) -> Optional[List[int]]:
"""
Convert indexer names to IDs by looking up enabled indexers.
Returns None if no names could be resolved.
"""
if not names:
return None
try:
enabled_indexers = client.get_enabled_indexers_detailed()
name_to_id = {
idx.get("name"): idx.get("id")
for idx in enabled_indexers
if idx.get("name") and idx.get("id") is not None
}
ids = []
for name in names:
idx_id = name_to_id.get(name)
if idx_id is not None:
try:
ids.append(int(idx_id))
except (TypeError, ValueError):
pass
return ids if ids else None
except Exception as e:
logger.warning(f"Failed to resolve indexer names to IDs: {e}")
return None
def search(
self,
book: BookMetadata,
@@ -336,8 +568,12 @@ class ProwlarrSource(ReleaseSource):
logger.warning("No search query available for book")
return []
# Get selected indexer IDs from config (None means search all)
indexer_ids = self._get_selected_indexer_ids()
# Get indexer IDs: prefer plan.indexers (from filter), else use settings
if plan.indexers:
indexer_ids = self._resolve_indexer_ids_from_names(client, plan.indexers)
logger.debug(f"Using filter-specified indexers: {plan.indexers} -> IDs {indexer_ids}")
else:
indexer_ids = self._get_selected_indexer_ids()
# Get search categories based on content type
# Audiobooks use 3030 (Audio/Audiobook), ebooks use 7000 (Books)
@@ -370,37 +606,61 @@ class ProwlarrSource(ReleaseSource):
f"Searching Prowlarr: {query_type} ({len(queries)} variants), {indexer_desc}, categories={categories}"
)
# Identify indexers that should be enriched via Torznab/Newznab.
enriched_indexer_ids = client.get_enriched_indexer_ids(restrict_to=indexer_ids)
non_enriched_indexer_ids: Optional[List[int]] = None
if indexer_ids:
non_enriched_indexer_ids = [i for i in indexer_ids if i not in enriched_indexer_ids]
def search_indexers(query: str, cats: Optional[List[int]]) -> List[dict]:
"""Search indexers with given categories, collecting results."""
results = []
if indexer_ids:
# Prefer a single request for all selected indexers to reduce latency.
try:
raw = client.search(query=query, indexer_ids=indexer_ids, categories=cats)
if raw:
results.extend(raw)
return results
except Exception as e:
logger.warning(
f"Search failed for selected indexers {indexer_ids}: {e}. Falling back to per-indexer search."
)
# Fallback: search specific indexers one at a time
for indexer_id in indexer_ids:
# Search standard indexers via JSON endpoint.
if indexer_ids:
if non_enriched_indexer_ids:
# Prefer a single request for selected indexers to reduce latency.
try:
raw = client.search(query=query, indexer_ids=[indexer_id], categories=cats)
raw = client.search(query=query, indexer_ids=non_enriched_indexer_ids, categories=cats)
if raw:
results.extend(raw)
except Exception as e:
logger.warning(f"Search failed for indexer {indexer_id}: {e}")
logger.warning(
f"Search failed for selected indexers {non_enriched_indexer_ids}: {e}. Falling back to per-indexer search."
)
for indexer_id in non_enriched_indexer_ids:
try:
raw = client.search(query=query, indexer_ids=[indexer_id], categories=cats)
if raw:
results.extend(raw)
except Exception as e:
logger.warning(f"Search failed for indexer {indexer_id}: {e}")
else:
# Search all enabled indexers at once
# Search all enabled indexers at once, then remove enriched results (re-fetched via Torznab).
try:
raw = client.search(query=query, indexer_ids=None, categories=cats)
if raw:
if enriched_indexer_ids:
raw = [r for r in raw if r.get("indexerId") not in enriched_indexer_ids]
results.extend(raw)
except Exception as e:
logger.warning(f"Search failed for all indexers: {e}")
# Search enriched indexers via Torznab/Newznab for richer metadata.
for indexer_id in enriched_indexer_ids:
raw = client.torznab_search(indexer_id=indexer_id, query=query, categories=cats, search_type="book")
if raw:
results.extend(raw)
else:
# Fallback to JSON search for enriched indexers if Torznab fails.
try:
raw_fallback = client.search(query=query, indexer_ids=[indexer_id], categories=cats)
if raw_fallback:
results.extend(raw_fallback)
except Exception as e:
logger.warning(f"Fallback search failed for enriched indexer {indexer_id}: {e}")
return results
try:
@@ -442,7 +702,30 @@ class ProwlarrSource(ReleaseSource):
seen_keys.add(key)
all_results.append(r)
results = [_prowlarr_result_to_release(r, content_type) for r in all_results]
enriched_indexer_ids_set = set(enriched_indexer_ids)
results: List[Release] = []
enriched_source_ids: set[str] = set()
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
is_enriched = bool(idx_id_int is not None and idx_id_int in enriched_indexer_ids_set)
release = _prowlarr_result_to_release(
r,
content_type,
enable_format_detection=is_enriched,
)
results.append(release)
if is_enriched:
enriched_source_ids.add(release.source_id)
# Sort results: enriched indexers first, then others
results.sort(key=lambda r: (0 if r.source_id in enriched_source_ids else 1))
if results:
torrent_count = sum(1 for r in results if r.protocol == ReleaseProtocol.TORRENT)
@@ -0,0 +1,171 @@
"""
Torznab/Newznab (RSS/XML) helpers for Prowlarr.
Used to fetch richer metadata from specific indexers (e.g., MyAnonamouse) that
isn't available via Prowlarr's JSON search endpoint.
"""
from __future__ import annotations
from typing import Any, Dict, List, Optional
from xml.etree import ElementTree as ET
def _local_name(tag: str) -> str:
"""Return tag name without namespace."""
if tag.startswith("{"):
return tag.split("}", 1)[1]
return tag
def _coerce_int(value: Optional[str]) -> Optional[int]:
if value is None:
return None
value = value.strip()
if not value:
return None
try:
return int(value)
except ValueError:
return None
def _coerce_float(value: Optional[str]) -> Optional[float]:
if value is None:
return None
value = value.strip()
if not value:
return None
try:
return float(value)
except ValueError:
return None
def _strip_author_from_title(title: str, author: Optional[str]) -> str:
"""
Prowlarr's MyAnonamouse parser appends " by {author}" into the title while
also emitting author/booktitle fields. Shelfmark's UI shows author
separately, so strip the duplicated " by author" segment when present.
"""
if not title or not author:
return title
needle = f" by {author}"
if needle in title:
return title.replace(needle, "", 1).strip()
return title
def parse_torznab_xml(xml_text: str) -> List[Dict[str, Any]]:
"""
Parse a Torznab/Newznab XML response into a list of dicts that roughly match
Prowlarr's JSON search results shape.
"""
if not xml_text or not xml_text.strip():
return []
try:
root = ET.fromstring(xml_text)
except ET.ParseError:
return []
items = root.findall(".//item")
results: List[Dict[str, Any]] = []
for item in items:
title = (item.findtext("title") or "").strip()
guid = (item.findtext("guid") or "").strip() or None
download_url = (item.findtext("link") or "").strip() or None
info_url = (item.findtext("comments") or "").strip() or None
pub_date = (item.findtext("pubDate") or "").strip() or None
size = _coerce_int(item.findtext("size"))
enclosure = item.find("enclosure")
enclosure_type = enclosure.get("type") if enclosure is not None else None
enclosure_url = enclosure.get("url") if enclosure is not None else None
protocol: Optional[str] = None
if enclosure_type == "application/x-bittorrent":
protocol = "torrent"
elif enclosure_type == "application/x-nzb":
protocol = "usenet"
if not download_url and enclosure_url:
download_url = enclosure_url.strip() or None
prowlarr_indexer_el = item.find("prowlarrindexer")
indexer_id = _coerce_int(prowlarr_indexer_el.get("id")) if prowlarr_indexer_el is not None else None
indexer_name = (prowlarr_indexer_el.text or "").strip() if prowlarr_indexer_el is not None else ""
categories: List[int] = []
for cat_el in item.findall("category"):
cat_id = _coerce_int(cat_el.text)
if cat_id is not None:
categories.append(cat_id)
# Collect torznab/newznab attr elements (namespaced).
attrs: Dict[str, str] = {}
tags: List[str] = []
for el in item.iter():
if _local_name(el.tag) != "attr":
continue
name = (el.get("name") or "").strip()
value = (el.get("value") or "").strip()
if not name:
continue
if name == "tag" and value:
tags.append(value)
continue
if value:
attrs[name] = value
seeders = _coerce_int(attrs.get("seeders"))
peers = _coerce_int(attrs.get("peers"))
leechers: Optional[int] = None
if peers is not None and seeders is not None and peers >= seeders:
leechers = peers - seeders
author = attrs.get("author") or None
book_title = attrs.get("booktitle") or None
info_hash = attrs.get("infohash") or None
download_volume_factor = _coerce_float(attrs.get("downloadvolumefactor"))
upload_volume_factor = _coerce_float(attrs.get("uploadvolumefactor"))
minimum_ratio = _coerce_float(attrs.get("minimumratio"))
minimum_seed_time = _coerce_int(attrs.get("minimumseedtime"))
cleaned_title = _strip_author_from_title(title, author)
results.append({
"title": cleaned_title or title,
"guid": guid or info_url or download_url or f"{indexer_id}:{title}",
"size": size,
"protocol": protocol or "unknown",
"downloadUrl": download_url,
"infoUrl": info_url,
"publishDate": pub_date,
"indexer": indexer_name or None,
"indexerId": indexer_id,
"categories": categories,
"seeders": seeders,
"leechers": leechers,
"files": _coerce_int(attrs.get("files")),
"grabs": _coerce_int(attrs.get("grabs")),
"infoHash": info_hash,
"indexerFlags": tags,
# Optional richer fields (not available via JSON search)
"author": author,
"bookTitle": book_title,
"downloadVolumeFactor": download_volume_factor,
"uploadVolumeFactor": upload_volume_factor,
"minimumRatio": minimum_ratio,
"minimumSeedTime": minimum_seed_time,
# Pass through all torznab attributes for tooltip display
"torznabAttrs": attrs,
})
return results

Some files were not shown because too many files have changed in this diff Show More