diff --git a/alembic/versions/i6c7d8e9f012_add_indexer_download_client.py b/alembic/versions/i6c7d8e9f012_add_indexer_download_client.py new file mode 100644 index 00000000..31d0be41 --- /dev/null +++ b/alembic/versions/i6c7d8e9f012_add_indexer_download_client.py @@ -0,0 +1,35 @@ +"""Let an indexer send every grab to one chosen download client. + +Revision ID: i6c7d8e9f012 +Revises: h5b6c7d8e901 +Create Date: 2026-10-03 +""" + +import sqlalchemy as sa + +from alembic import op + +revision = "i6c7d8e9f012" +down_revision = "h5b6c7d8e901" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + """Add the nullable client pin; existing indexers stay unpinned.""" + with op.batch_alter_table("indexer_configs") as batch_op: + batch_op.add_column(sa.Column("download_client_id", sa.Integer(), nullable=True)) + batch_op.create_foreign_key( + "fk_indexer_configs_download_client", + "download_client_configs", + ["download_client_id"], + ["id"], + ondelete="SET NULL", + ) + + +def downgrade() -> None: + """Drop the pin; those indexers fall back to priority-based selection.""" + with op.batch_alter_table("indexer_configs") as batch_op: + batch_op.drop_constraint("fk_indexer_configs_download_client", type_="foreignkey") + batch_op.drop_column("download_client_id") diff --git a/src/pullbox/api/v1/clients.py b/src/pullbox/api/v1/clients.py index ba7f7f1a..c20af9e9 100644 --- a/src/pullbox/api/v1/clients.py +++ b/src/pullbox/api/v1/clients.py @@ -243,16 +243,6 @@ async def add_client( session: DbSession, ) -> ClientResponse: """Add a new download client configuration.""" - if body.client_type is not DownloadClientType.AIRDCPP: - existing = await session.execute( - select(DownloadClientConfig).where(DownloadClientConfig.client_type == body.client_type) - ) - if existing.scalar_one_or_none(): - raise ValidationError( - f"A {body.client_type.value} client is already configured. " - "Only one instance per client type is allowed." - ) - if body.client_type is DownloadClientType.AIRDCPP: if not get_settings().airdcpp_enabled: raise ValidationError("AirDC++ integration is disabled by the feature flag.") diff --git a/src/pullbox/api/v1/indexers.py b/src/pullbox/api/v1/indexers.py index 9e2f1dec..c30b3437 100644 --- a/src/pullbox/api/v1/indexers.py +++ b/src/pullbox/api/v1/indexers.py @@ -9,6 +9,7 @@ from sqlalchemy.exc import OperationalError from pullbox.api.deps import DbSession, InteractiveOperatorUser +from pullbox.core.acquisition import AcquisitionProtocol from pullbox.core.encryption import decrypt_secret, encrypt_secret from pullbox.core.exceptions import NotFoundError, ValidationError from pullbox.core.sqlite_lock import ( @@ -16,7 +17,9 @@ is_sqlite_locked_error, sqlite_lock_retry_delay, ) -from pullbox.models.indexer import IndexerConfig +from pullbox.models.client import DownloadClientConfig +from pullbox.models.download import DownloadClientType +from pullbox.models.indexer import IndexerConfig, IndexerType from pullbox.schemas.indexer import ( IndexerCreate, IndexerResponse, @@ -54,6 +57,7 @@ def _redact_indexer(indexer: IndexerConfig) -> dict[str, object]: "enable_automatic_search": indexer.enable_automatic_search, "enable_interactive_search": indexer.enable_interactive_search, "resolver_enabled": indexer.resolver_enabled, + "download_client_id": indexer.download_client_id, "last_success_at": indexer.last_success_at, "last_failure_at": indexer.last_failure_at, "last_error": indexer.last_error, @@ -136,6 +140,7 @@ async def add_indexer( source="manual", resolver_enabled=body.resolver_enabled, ) + await _validate_download_client(session, body.indexer_type, body.download_client_id) indexer = IndexerConfig( name=body.name, indexer_type=body.indexer_type, @@ -149,6 +154,7 @@ async def add_indexer( enable_automatic_search=body.enable_automatic_search, enable_interactive_search=body.enable_interactive_search, resolver_enabled=body.resolver_enabled, + download_client_id=body.download_client_id, ) session.add(indexer) await session.flush() @@ -501,6 +507,11 @@ async def update_indexer( resolver_enabled=True, ) + if update_data.get("download_client_id") is not None: + await _validate_download_client( + session, indexer.indexer_type, update_data["download_client_id"] + ) + # Encrypt api_key if provided; omitted → keep existing encrypted value if update_data.get("api_key"): update_data["api_key"] = encrypt_secret(update_data["api_key"]) @@ -530,6 +541,35 @@ def _validate_resolver_scope( ) +_PROTOCOL_BY_INDEXER_TYPE: dict[str, tuple[AcquisitionProtocol, ...]] = { + IndexerType.TORZNAB: (AcquisitionProtocol.TORRENT,), + IndexerType.NEWZNAB: (AcquisitionProtocol.USENET,), +} + + +async def _validate_download_client( + session: DbSession, + indexer_type: object, + download_client_id: int | None, +) -> None: + """Check a pinned client exists and takes the protocol this indexer returns.""" + if download_client_id is None: + return + client = await session.get(DownloadClientConfig, download_client_id) + if client is None: + raise ValidationError(f"Download client {download_client_id} does not exist.") + wanted = _PROTOCOL_BY_INDEXER_TYPE.get( + str(indexer_type), + (AcquisitionProtocol.TORRENT, AcquisitionProtocol.USENET), + ) + if DownloadClientType(client.client_type).acquisition_protocol not in wanted: + kind = " or ".join(protocol.value for protocol in wanted) + raise ValidationError( + f'"{client.name}" is not a {kind} download client, so it cannot ' + "take releases from this indexer." + ) + + # ── Delete ─────────────────────────────────────────────────────────── diff --git a/src/pullbox/models/indexer.py b/src/pullbox/models/indexer.py index d074a081..f4c51308 100644 --- a/src/pullbox/models/indexer.py +++ b/src/pullbox/models/indexer.py @@ -3,7 +3,7 @@ import enum from datetime import datetime -from sqlalchemy import Boolean, Integer, String, Text +from sqlalchemy import Boolean, ForeignKey, Integer, String, Text from sqlalchemy import Enum as SQLAlchemyEnum from sqlalchemy.orm import Mapped, mapped_column @@ -35,6 +35,13 @@ class IndexerConfig(Base, IdentityMixin, TimestampMixin): priority: Mapped[int] = mapped_column(Integer, default=50) categories: Mapped[str | None] = mapped_column(Text) + # Download client every grab from this indexer is sent to. Null means the + # highest-priority client for the release's protocol. + download_client_id: Mapped[int | None] = mapped_column( + ForeignKey("download_client_configs.id", ondelete="SET NULL"), + nullable=True, + ) + # Manager sync tracking. The legacy Prowlarr integer remains during the # additive migration so existing databases and integrations stay readable. source: Mapped[str] = mapped_column(String(20), default="manual") diff --git a/src/pullbox/schemas/indexer.py b/src/pullbox/schemas/indexer.py index d4bfa712..7a5e80da 100644 --- a/src/pullbox/schemas/indexer.py +++ b/src/pullbox/schemas/indexer.py @@ -25,6 +25,10 @@ class IndexerCreate(BaseModel): False, description="Allow a manual Torznab indexer to use the ranked browser resolver chain", ) + download_client_id: int | None = Field( + None, + description="Send every grab from this indexer to this download client", + ) @field_validator("url") @classmethod @@ -51,6 +55,10 @@ class IndexerUpdate(BaseModel): None, description="Allow a manual Torznab indexer to use the ranked browser resolver chain", ) + download_client_id: int | None = Field( + None, + description="Send every grab to this download client; null for the default client", + ) @field_validator("url") @classmethod @@ -82,6 +90,7 @@ class IndexerResponse(BaseModel): enable_automatic_search: bool = True enable_interactive_search: bool = True resolver_enabled: bool = False + download_client_id: int | None = None last_success_at: datetime | None = None last_failure_at: datetime | None = None last_error: str | None = None diff --git a/src/pullbox/services/download_service.py b/src/pullbox/services/download_service.py index 7363817c..f8876b6e 100644 --- a/src/pullbox/services/download_service.py +++ b/src/pullbox/services/download_service.py @@ -66,7 +66,7 @@ async def send_to_client( indexer_id = release.indexer_id # Select the appropriate client - client = self._select_client(release.protocol) + client = await self._client_for_release(session, release.protocol, indexer_id) if not client: raise ProviderError( "download", @@ -503,6 +503,54 @@ async def on_attempt(event: object) -> None: clear_download_progress(download_id) raise + async def _client_for_release( + self, + session: AsyncSession, + protocol: AcquisitionProtocol, + indexer_id: int | None, + ) -> DownloadClient | None: + """Use the indexer's pinned client when it has one, else the protocol default. + + A pinned client is never swapped for another one: if it is disabled or + cannot take this protocol the grab fails, so releases from that indexer + never land in a client the user kept them out of. + """ + if indexer_id is not None: + from pullbox.models.indexer import IndexerConfig + + indexer = await session.get(IndexerConfig, indexer_id) + if indexer is not None and indexer.download_client_id is not None: + return await self._pinned_client( + session, indexer.download_client_id, protocol, indexer.name + ) + return self._select_client(protocol) + + async def _pinned_client( + self, + session: AsyncSession, + client_config_id: int, + protocol: AcquisitionProtocol, + indexer_name: str, + ) -> DownloadClient: + client = self._registry.get_download_client(client_config_id) + if client is None: + from pullbox.models.client import DownloadClientConfig + + config = await session.get(DownloadClientConfig, client_config_id) + label = config.name if config is not None else f"#{client_config_id}" + raise ProviderError( + "download", + f'Indexer "{indexer_name}" sends its grabs to download client "{label}", ' + "which is disabled or unavailable.", + ) + if DownloadClientType(client.client_type).acquisition_protocol is not protocol: + raise ProviderError( + "download", + f'Indexer "{indexer_name}" is set to download client "{client.name}", ' + f"which cannot take {protocol.value} releases.", + ) + return client + def _select_client(self, protocol: AcquisitionProtocol) -> DownloadClient | None: """Select the highest-priority client for an acquisition protocol.""" if protocol is AcquisitionProtocol.TORRENT: @@ -541,9 +589,19 @@ async def _persisted_config_id_for_client( def get_client_for_download(self, download: DownloadHistory) -> DownloadClient | None: """Resolve the exact persisted client, with fallback for legacy null rows.""" - if download.download_client_config_id is not None: - return self._registry.get_download_client(download.download_client_config_id) - return self._registry.get_client_for_type(str(download.download_client)) + return self.get_client_for_identity( + download.download_client_config_id, download.download_client + ) + + def get_client_for_identity( + self, + client_config_id: int | None, + client_type: object, + ) -> DownloadClient | None: + """Resolve a client by its config ID, or by type for rows recorded without one.""" + if client_config_id is not None: + return self._registry.get_download_client(client_config_id) + return self._registry.get_client_for_type(str(client_type)) def get_client_for_type(self, client_type: object) -> DownloadClient | None: """Get the client for a given DownloadClientType value.""" diff --git a/src/pullbox/tasks/download_monitor_poll.py b/src/pullbox/tasks/download_monitor_poll.py index 3c7683d9..6e8eaf76 100644 --- a/src/pullbox/tasks/download_monitor_poll.py +++ b/src/pullbox/tasks/download_monitor_poll.py @@ -30,7 +30,9 @@ async def poll_download_clients( client_type = item["download_client"] existing_path = item["downloaded_path"] - client = download_svc.get_client_for_type(client_type) + client = download_svc.get_client_for_identity( + item.get("download_client_config_id"), client_type + ) if not client: continue diff --git a/src/pullbox/tasks/download_monitor_read.py b/src/pullbox/tasks/download_monitor_read.py index 02d4ea3e..2400a61c 100644 --- a/src/pullbox/tasks/download_monitor_read.py +++ b/src/pullbox/tasks/download_monitor_read.py @@ -37,6 +37,7 @@ def build_poll_item(download: Any) -> dict[str, object]: "external_id": download.external_id, "title": download.title, "download_client": download.download_client, + "download_client_config_id": download.download_client_config_id, "downloaded_path": download.downloaded_path, "issue_id": download.issue_id, "retry_count": download.retry_count, diff --git a/src/pullbox/tasks/download_post_processing_sources.py b/src/pullbox/tasks/download_post_processing_sources.py index 159d22b1..822e5d84 100644 --- a/src/pullbox/tasks/download_post_processing_sources.py +++ b/src/pullbox/tasks/download_post_processing_sources.py @@ -15,6 +15,7 @@ from sqlalchemy.ext.asyncio import AsyncSession + from pullbox.models.client import DownloadClientConfig from pullbox.models.download import DownloadHistory logger = structlog.get_logger(__name__) @@ -251,13 +252,7 @@ async def _resolve_local_path( ) return str(mapped) - result = await session.execute( - select(DownloadClientConfig).where( - DownloadClientConfig.client_type == download.download_client, - DownloadClientConfig.enabled.is_(True), - ) - ) - client_cfg = result.scalars().first() + client_cfg = await _download_client_config(session, download) if client_cfg and client_cfg.remote_path and client_cfg.download_dir: windows_origin = _is_windows_origin_path(client_cfg.remote_path) @@ -299,20 +294,38 @@ async def _resolve_local_path( return raw_path -async def _resolve_local_download_root( +async def _download_client_config( session: AsyncSession, download: DownloadHistory, -) -> Path | None: - """Return the configured local root that bounds source cleanup.""" +) -> DownloadClientConfig | None: + """Return the config of the client that handled this download. + + Several clients can share a type, so the exact config recorded on the + download wins; rows recorded before that identity existed fall back to the + first enabled client of the same type. The exact config is used even when + it has since been disabled: its paths still describe where this download is. + """ from pullbox.models.client import DownloadClientConfig + if download.download_client_config_id is not None: + exact = await session.get(DownloadClientConfig, download.download_client_config_id) + if exact is not None: + return exact result = await session.execute( select(DownloadClientConfig).where( DownloadClientConfig.client_type == download.download_client, DownloadClientConfig.enabled.is_(True), ) ) - client_cfg = result.scalars().first() + return result.scalars().first() + + +async def _resolve_local_download_root( + session: AsyncSession, + download: DownloadHistory, +) -> Path | None: + """Return the configured local root that bounds source cleanup.""" + client_cfg = await _download_client_config(session, download) if client_cfg is None or not client_cfg.download_dir: return None download_dir = client_cfg.download_dir.strip() diff --git a/src/pullbox/ui/settings_routes.py b/src/pullbox/ui/settings_routes.py index 445ee041..a410918d 100644 --- a/src/pullbox/ui/settings_routes.py +++ b/src/pullbox/ui/settings_routes.py @@ -13,12 +13,14 @@ from pullbox.api.deps import AuthenticatedUser, DbSession from pullbox.config import get_settings +from pullbox.core.acquisition import AcquisitionProtocol from pullbox.core.naming import ( resolve_collection_non_standard_file_template, resolve_single_non_standard_file_template, ) from pullbox.models.client import DownloadClientConfig from pullbox.models.config import SystemConfig +from pullbox.models.download import DownloadClientType from pullbox.models.indexer import IndexerConfig page_router = APIRouter() @@ -181,6 +183,29 @@ def load_client_status_seed( return seed +async def _indexer_download_client_choices(session: DbSession) -> list[dict[str, object]]: + """Clients an indexer can be pinned to: the torrent and Usenet ones.""" + result = await session.execute( + select(DownloadClientConfig).order_by( + DownloadClientConfig.priority, DownloadClientConfig.name + ) + ) + choices: list[dict[str, object]] = [] + for client in result.scalars().all(): + protocol = DownloadClientType(client.client_type).acquisition_protocol + if protocol not in {AcquisitionProtocol.TORRENT, AcquisitionProtocol.USENET}: + continue + choices.append( + { + "id": client.id, + "name": client.name, + "protocol": protocol.value, + "enabled": client.enabled, + } + ) + return choices + + def load_indexer_status_seed( request: Request, indexers: Sequence[IndexerConfig], @@ -337,6 +362,7 @@ async def load_settings_tab(request: Request, session: DbSession, tab: str) -> d ctx["browser_resolver_available"] = ( await session.scalar(select(DirectResolverConfig.id).limit(1)) is not None ) + ctx["indexer_download_clients"] = await _indexer_download_client_choices(session) manager_sources_by_name: dict[str, set[str]] = {} manager_display_names: dict[str, str] = {} for indexer in indexers: diff --git a/src/pullbox/ui/templates/partials/settings_clients.html b/src/pullbox/ui/templates/partials/settings_clients.html index 133f80bd..de07a991 100644 --- a/src/pullbox/ui/templates/partials/settings_clients.html +++ b/src/pullbox/ui/templates/partials/settings_clients.html @@ -336,69 +336,59 @@

What type of download client are you adding?

{% if airdcpp_enabled %}
+
+ + +

Send every grab from this indexer to this client. If it is disabled or unreachable the grab fails rather than going to another client.

+
+
wanted.includes(client.protocol)); + }, form: { indexer_type: '', name: '', url: '', api_key: '', has_api_key: false, enabled: true, priority: 50, categories: '', source: 'manual', enable_rss: true, enable_automatic_search: true, enable_interactive_search: true, - resolver_enabled: false, + resolver_enabled: false, download_client_id: '', }, // Prowlarr connection state @@ -1014,7 +1039,7 @@ indexer_type: '', name: '', url: '', api_key: '', has_api_key: false, enabled: true, priority: 50, categories: '', source: 'manual', enable_rss: true, enable_automatic_search: true, enable_interactive_search: true, - resolver_enabled: false, + resolver_enabled: false, download_client_id: '', }; }, @@ -1170,6 +1195,7 @@ enable_automatic_search: data.enable_automatic_search !== false, enable_interactive_search: data.enable_interactive_search !== false, resolver_enabled: this.resolverChainAvailable && data.resolver_enabled === true, + download_client_id: data.download_client_id == null ? '' : String(data.download_client_id), }; this.showModal = true; }) @@ -1204,6 +1230,7 @@ && this.form.source === 'manual' ? this.form.resolver_enabled : false, + download_client_id: this.form.download_client_id ? Number(this.form.download_client_id) : null, }; // Only send api_key if user entered a new one diff --git a/tests/api/test_clients_api.py b/tests/api/test_clients_api.py index faaaac8f..7e33b936 100644 --- a/tests/api/test_clients_api.py +++ b/tests/api/test_clients_api.py @@ -14,7 +14,7 @@ from pullbox.api.v1 import clients as clients_api from pullbox.core.encryption import decrypt_secret, encrypt_secret -from pullbox.core.exceptions import NotFoundError, ProviderError, ValidationError +from pullbox.core.exceptions import NotFoundError, ProviderError from pullbox.models.client import DownloadClientConfig from pullbox.models.download import DownloadClientType from pullbox.providers.base import ProviderHealthResult @@ -136,14 +136,17 @@ async def test_create_encrypts_secret_and_rejects_duplicate_client_type( assert decrypt_secret(row.api_key) == "raw-sab-key" assert row.category == "comics" - duplicate = await authenticated_client.post( + # A second client of the same type is allowed, e.g. one per tracker. + second = await authenticated_client.post( "/api/v1/clients", json=_client_payload(name="Second SAB", api_key="other-key"), headers=_csrf_header_for(authenticated_client), ) - assert duplicate.status_code == 422 - assert "already configured" in duplicate.text + assert second.status_code == 201 + assert second.json()["id"] != created["id"] + listed = await authenticated_client.get("/api/v1/clients") + assert sorted(item["name"] for item in listed.json()) == ["SAB", "Second SAB"] async def test_update_preserves_blank_password_and_clears_blank_api_key( self, @@ -461,12 +464,12 @@ async def test_crud_route_functions_redact_encrypt_and_validate( assert stored.password != "raw-password" assert decrypt_secret(stored.password) == "raw-password" - with pytest.raises(ValidationError): - await clients_api.add_client( - _create_model(name="Second qBit", client_type="qbittorrent"), - object(), # type: ignore[arg-type] - session, - ) + second = await clients_api.add_client( + _create_model(name="Second qBit", client_type="qbittorrent"), + object(), # type: ignore[arg-type] + session, + ) + assert second.id != created.id updated = await clients_api.update_client( created.id, diff --git a/tests/api/test_indexer_download_client.py b/tests/api/test_indexer_download_client.py new file mode 100644 index 00000000..72d58448 --- /dev/null +++ b/tests/api/test_indexer_download_client.py @@ -0,0 +1,206 @@ +"""API contracts for pinning an indexer's grabs to one download client.""" + +from __future__ import annotations + +import os +import sys +from typing import TYPE_CHECKING + +import pytest +from sqlalchemy import event, text +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine + +from pullbox.api.v1 import clients as clients_api +from pullbox.api.v1 import indexers as indexers_api +from pullbox.core.exceptions import ValidationError +from pullbox.models import Base +from pullbox.models.client import DownloadClientConfig +from pullbox.models.download import DownloadClientType +from pullbox.models.indexer import IndexerConfig +from pullbox.schemas.indexer import IndexerCreate, IndexerUpdate + +if TYPE_CHECKING: + from collections.abc import AsyncGenerator + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) + +pytest_plugins = ["conftest_security"] + + +async def _seed_clients(session: AsyncSession) -> tuple[DownloadClientConfig, DownloadClientConfig]: + torrent = DownloadClientConfig( + name="qBittorrent (tracker)", + client_type=DownloadClientType.QBITTORRENT, + url="http://qbit-tracker.test", + priority=90, + ) + usenet = DownloadClientConfig( + name="SABnzbd", + client_type=DownloadClientType.SABNZBD, + url="http://sab.test", + ) + session.add_all([torrent, usenet]) + await session.flush() + return torrent, usenet + + +def _torznab(**overrides: object) -> IndexerCreate: + payload: dict[str, object] = { + "name": "Private Tracker", + "indexer_type": "torznab", + "url": "http://tracker.example", + "api_key": "tracker-secret", + } + payload.update(overrides) + return IndexerCreate.model_validate(payload) + + +@pytest.mark.asyncio +class TestIndexerDownloadClientRoutes: + async def test_create_and_read_back_pinned_client( + self, + sec_db: async_sessionmaker[AsyncSession], + ) -> None: + async with sec_db() as session: + torrent, _usenet = await _seed_clients(session) + created = await indexers_api.add_indexer( + _torznab(download_client_id=torrent.id), + object(), # type: ignore[arg-type] + session, + ) + fetched = await indexers_api.get_indexer( + created.id, + object(), # type: ignore[arg-type] + session, + ) + + assert created.download_client_id == torrent.id + assert fetched.download_client_id == torrent.id + + async def test_unpinned_by_default( + self, + sec_db: async_sessionmaker[AsyncSession], + ) -> None: + async with sec_db() as session: + created = await indexers_api.add_indexer( + _torznab(), + object(), # type: ignore[arg-type] + session, + ) + + assert created.download_client_id is None + + async def test_rejects_client_of_the_wrong_protocol( + self, + sec_db: async_sessionmaker[AsyncSession], + ) -> None: + async with sec_db() as session: + _torrent, usenet = await _seed_clients(session) + with pytest.raises(ValidationError, match="torrent download client"): + await indexers_api.add_indexer( + _torznab(download_client_id=usenet.id), + object(), # type: ignore[arg-type] + session, + ) + + async def test_rejects_unknown_client( + self, + sec_db: async_sessionmaker[AsyncSession], + ) -> None: + async with sec_db() as session: + with pytest.raises(ValidationError, match="does not exist"): + await indexers_api.add_indexer( + _torznab(download_client_id=4242), + object(), # type: ignore[arg-type] + session, + ) + + async def test_update_sets_validates_and_clears_pin( + self, + sec_db: async_sessionmaker[AsyncSession], + ) -> None: + async with sec_db() as session: + torrent, usenet = await _seed_clients(session) + created = await indexers_api.add_indexer( + _torznab(), + object(), # type: ignore[arg-type] + session, + ) + + pinned = await indexers_api.update_indexer( + created.id, + IndexerUpdate(download_client_id=torrent.id), + object(), # type: ignore[arg-type] + session, + ) + assert pinned.download_client_id == torrent.id + + with pytest.raises(ValidationError, match="torrent download client"): + await indexers_api.update_indexer( + created.id, + IndexerUpdate(download_client_id=usenet.id), + object(), # type: ignore[arg-type] + session, + ) + + # An unrelated edit leaves the pin alone. + renamed = await indexers_api.update_indexer( + created.id, + IndexerUpdate(name="Renamed Tracker"), + object(), # type: ignore[arg-type] + session, + ) + assert renamed.download_client_id == torrent.id + + cleared = await indexers_api.update_indexer( + created.id, + IndexerUpdate.model_validate({"download_client_id": None}), + object(), # type: ignore[arg-type] + session, + ) + + assert cleared.download_client_id is None + + +@pytest.fixture +async def fk_db() -> AsyncGenerator[async_sessionmaker[AsyncSession], None]: + """In-memory database with SQLite foreign keys enforced, as in production.""" + engine = create_async_engine("sqlite+aiosqlite:///:memory:", echo=False) + + @event.listens_for(engine.sync_engine, "connect") + def _enable_foreign_keys(dbapi_connection: object, _record: object) -> None: + cursor = dbapi_connection.cursor() # type: ignore[attr-defined] + cursor.execute("PRAGMA foreign_keys=ON") + cursor.close() + + async with engine.begin() as conn: + await conn.run_sync(Base.metadata.create_all) + yield async_sessionmaker(engine, expire_on_commit=False) + await engine.dispose() + + +@pytest.mark.asyncio +async def test_deleting_the_client_unpins_the_indexer( + fk_db: async_sessionmaker[AsyncSession], +) -> None: + async with fk_db() as session: + assert (await session.execute(text("PRAGMA foreign_keys"))).scalar_one() == 1 + torrent, _usenet = await _seed_clients(session) + created = await indexers_api.add_indexer( + _torznab(download_client_id=torrent.id), + object(), # type: ignore[arg-type] + session, + ) + await session.commit() + + await clients_api.delete_client( + torrent.id, + object(), # type: ignore[arg-type] + session, + ) + await session.commit() + session.expire_all() + indexer = await session.get(IndexerConfig, created.id) + + assert indexer is not None + assert indexer.download_client_id is None diff --git a/tests/integration/test_download_clients_integration.py b/tests/integration/test_download_clients_integration.py index a88679a9..c6e64dba 100644 --- a/tests/integration/test_download_clients_integration.py +++ b/tests/integration/test_download_clients_integration.py @@ -18,6 +18,7 @@ from pullbox.core.acquisition import AcquisitionProtocol from pullbox.core.events import EventBus +from pullbox.core.exceptions import ProviderError from pullbox.models import Base from pullbox.models.client import DownloadClientConfig from pullbox.models.download import DownloadClientType, DownloadHistory, DownloadState @@ -291,6 +292,118 @@ async def test_no_nzb_client_raises( # ── Client type affinity ───────────────────────────────────────── +async def _pin_indexer( + session: AsyncSession, + client_config_id: int, + client_type: DownloadClientType = DownloadClientType.QBITTORRENT, +) -> None: + session.add( + DownloadClientConfig( + id=client_config_id, + name=f"Client {client_config_id}", + client_type=client_type, + url=f"http://client-{client_config_id}.test", + ) + ) + await session.flush() + indexer = await session.get(IndexerConfig, 7) + assert indexer is not None + indexer.download_client_id = client_config_id + await session.flush() + + +@pytest.mark.asyncio +class TestIndexerPinnedClient: + """An indexer pinned to a download client always sends its grabs there.""" + + async def test_pinned_client_wins_over_priority( + self, + session: AsyncSession, + issue: Issue, + ) -> None: + reg = ProviderRegistry() + default = _mock_client("qbittorrent", "default") + tracker = _mock_client("qbittorrent", "tracker") + reg.register_download_client(1, default, priority=10) + reg.register_download_client(2, tracker, priority=90) + await _pin_indexer(session, 2) + + svc = DownloadService(registry=reg, event_bus=EventBus()) + release = _make_release("Pinned", is_torrent=True, registry=reg) + dl = await svc.send_to_client(session, release, issue.id) + + tracker.add_torrent_data.assert_awaited_once_with(b"fixture-torrent", "Pinned") + default.add_torrent_data.assert_not_awaited() + assert dl.download_client_config_id == 2 + assert svc.get_client_for_download(dl) is tracker + + async def test_unpinned_indexer_keeps_priority_selection( + self, + session: AsyncSession, + issue: Issue, + ) -> None: + reg = ProviderRegistry() + default = _mock_client("qbittorrent", "default") + tracker = _mock_client("qbittorrent", "tracker") + reg.register_download_client(1, default, priority=10) + reg.register_download_client(2, tracker, priority=90) + session.add( + DownloadClientConfig( + id=1, + name="Default", + client_type=DownloadClientType.QBITTORRENT, + url="http://default.test", + ) + ) + await session.flush() + + svc = DownloadService(registry=reg, event_bus=EventBus()) + release = _make_release("Unpinned", is_torrent=True, registry=reg) + dl = await svc.send_to_client(session, release, issue.id) + + default.add_torrent_data.assert_awaited_once() + tracker.add_torrent_data.assert_not_awaited() + assert dl.download_client_config_id == 1 + + async def test_unavailable_pinned_client_fails_instead_of_falling_back( + self, + session: AsyncSession, + issue: Issue, + ) -> None: + reg = ProviderRegistry() + default = _mock_client("qbittorrent", "default") + reg.register_download_client(1, default, priority=10) + # Client 2 exists but is disabled, so it is not registered. + await _pin_indexer(session, 2) + + svc = DownloadService(registry=reg, event_bus=EventBus()) + release = _make_release("Pinned", is_torrent=True, registry=reg) + with pytest.raises(ProviderError, match="Client 2"): + await svc.send_to_client(session, release, issue.id) + + default.add_torrent_data.assert_not_awaited() + + async def test_pinned_client_of_another_protocol_is_refused( + self, + session: AsyncSession, + issue: Issue, + ) -> None: + reg = ProviderRegistry() + default = _mock_client("qbittorrent", "default") + sab = _mock_client("sabnzbd", "sab") + reg.register_download_client(1, default, priority=10) + reg.register_download_client(2, sab, priority=10) + await _pin_indexer(session, 2, DownloadClientType.SABNZBD) + + svc = DownloadService(registry=reg, event_bus=EventBus()) + release = _make_release("Pinned", is_torrent=True, registry=reg) + with pytest.raises(ProviderError, match="cannot take torrent"): + await svc.send_to_client(session, release, issue.id) + + default.add_torrent_data.assert_not_awaited() + sab.add_nzb.assert_not_awaited() + + @pytest.mark.asyncio class TestClientTypeAffinity: """Downloads are polled from the correct client based on recorded type.""" diff --git a/tests/integration/test_migrations.py b/tests/integration/test_migrations.py index fe80b861..97a2d22b 100644 --- a/tests/integration/test_migrations.py +++ b/tests/integration/test_migrations.py @@ -1870,10 +1870,37 @@ def test_library_root_management_downgrade_refuses_unrepresentable_default( finally: engine.dispose() + def test_indexer_download_client_pin_round_trips(self, alembic_cfg) -> None: + cfg, sync_url = alembic_cfg + command.upgrade(cfg, "i6c7d8e9f012") + + engine = create_engine(sync_url) + try: + columns = {c["name"] for c in inspect(engine).get_columns("indexer_configs")} + assert "download_client_id" in columns + foreign_keys = inspect(engine).get_foreign_keys("indexer_configs") + assert [ + (fk["referred_table"], fk["constrained_columns"], fk["options"].get("ondelete")) + for fk in foreign_keys + ] == [("download_client_configs", ["download_client_id"], "SET NULL")] + finally: + engine.dispose() + + command.downgrade(cfg, "h5b6c7d8e901") + + engine = create_engine(sync_url) + try: + columns = {c["name"] for c in inspect(engine).get_columns("indexer_configs")} + assert "download_client_id" not in columns + assert inspect(engine).get_foreign_keys("indexer_configs") == [] + finally: + engine.dispose() + def test_root_removal_protects_dependencies_without_losing_other_fk_actions(self, alembic_cfg): cfg, sync_url = alembic_cfg script = ScriptDirectory.from_config(cfg) - assert script.get_heads() == ["h5b6c7d8e901"] + assert script.get_heads() == ["i6c7d8e9f012"] + assert script.get_revision("i6c7d8e9f012").down_revision == "h5b6c7d8e901" assert script.get_revision("h5b6c7d8e901").down_revision == "g4a5b6c7d890" assert script.get_revision("g4a5b6c7d890").down_revision == "f3z4a5b6c789" assert script.get_revision("f3z4a5b6c789").down_revision == "e2y3z4a5b678" diff --git a/tests/tasks/test_download_monitor_poll.py b/tests/tasks/test_download_monitor_poll.py index ae490b37..eae21927 100644 --- a/tests/tasks/test_download_monitor_poll.py +++ b/tests/tasks/test_download_monitor_poll.py @@ -32,9 +32,11 @@ async def get_download_status(self, external_id: str): class _FakeService: def __init__(self, client: _FakeClient) -> None: self.client = client + self.identities: list[tuple[object, object]] = [] - def get_client_for_type(self, client_type: object) -> _FakeClient: + def get_client_for_identity(self, client_config_id: object, client_type: object) -> _FakeClient: assert client_type == "qbittorrent" + self.identities.append((client_config_id, client_type)) return self.client @@ -88,3 +90,33 @@ async def test_poll_download_clients_matches_missing_external_id_by_title() -> N }, ] assert client.status_calls == ["matched-hash"] + + +@pytest.mark.asyncio +async def test_poll_download_clients_asks_for_the_exact_client() -> None: + """With two clients of one type, polling must reach the one that has the download.""" + from pullbox.tasks import download_monitor_poll + from pullbox.tasks.download_monitor_updates import build_status_update + + service = _FakeService(_FakeClient()) + + await download_monitor_poll.poll_download_clients( + [ + { + "id": 8, + "external_id": "hash-8", + "title": "Batman 001.cbz", + "download_client": "qbittorrent", + "download_client_config_id": 5, + "downloaded_path": None, + "issue_id": 99, + } + ], + service, + record_download_progress=lambda download_id, status, event_logger: False, + build_status_update=build_status_update, + build_status_check_error_update=lambda **kwargs: None, + event_logger=_FakeLogger(), + ) + + assert service.identities == [(5, "qbittorrent")] diff --git a/tests/tasks/test_download_monitor_read.py b/tests/tasks/test_download_monitor_read.py index 3521a59c..b8c84673 100644 --- a/tests/tasks/test_download_monitor_read.py +++ b/tests/tasks/test_download_monitor_read.py @@ -16,6 +16,7 @@ def test_build_poll_item_snapshots_detached_download_fields() -> None: external_id="abc", title="Batman 001.cbz", download_client="sabnzbd", + download_client_config_id=4, downloaded_path="/downloads/Batman 001.cbz", issue_id=22, retry_count=1, @@ -27,6 +28,7 @@ def test_build_poll_item_snapshots_detached_download_fields() -> None: "external_id": "abc", "title": "Batman 001.cbz", "download_client": "sabnzbd", + "download_client_config_id": 4, "downloaded_path": "/downloads/Batman 001.cbz", "issue_id": 22, "retry_count": 1, diff --git a/tests/tasks/test_download_post_processing_sources.py b/tests/tasks/test_download_post_processing_sources.py index a3196f4e..84a96646 100644 --- a/tests/tasks/test_download_post_processing_sources.py +++ b/tests/tasks/test_download_post_processing_sources.py @@ -35,7 +35,9 @@ async def test_resolve_local_download_root_uses_enabled_client_directory() -> No result = MagicMock() result.scalars.return_value.first.return_value = SimpleNamespace(download_dir="/downloads/") session = SimpleNamespace(execute=AsyncMock(return_value=result)) - download = SimpleNamespace(download_client=DownloadClientType.SABNZBD) + download = SimpleNamespace( + download_client_config_id=None, download_client=DownloadClientType.SABNZBD + ) root = await _resolve_local_download_root(session, download) @@ -51,7 +53,9 @@ async def test_resolve_local_download_root_requires_configured_directory() -> No result = MagicMock() result.scalars.return_value.first.return_value = SimpleNamespace(download_dir=None) session = SimpleNamespace(execute=AsyncMock(return_value=result)) - download = SimpleNamespace(download_client=DownloadClientType.SABNZBD) + download = SimpleNamespace( + download_client_config_id=None, download_client=DownloadClientType.SABNZBD + ) root = await _resolve_local_download_root(session, download) @@ -97,6 +101,7 @@ async def test_resolve_local_path_normalizes_windows_remote_path() -> None: ) session = SimpleNamespace(execute=AsyncMock(return_value=result)) download = SimpleNamespace( + download_client_config_id=None, download_client=DownloadClientType.SABNZBD, downloaded_path=r"E:\Temp\Pullbox_Downloads\Release\file.cbr", ) @@ -120,6 +125,7 @@ async def test_resolve_local_path_rejects_unmapped_windows_path() -> None: ) session = SimpleNamespace(execute=AsyncMock(return_value=result)) download = SimpleNamespace( + download_client_config_id=None, download_client=DownloadClientType.SABNZBD, downloaded_path=r"E:\Temp\Release\file.cbr", ) @@ -141,6 +147,7 @@ async def test_resolve_local_path_matches_windows_paths_case_insensitively() -> ) session = SimpleNamespace(execute=AsyncMock(return_value=result)) download = SimpleNamespace( + download_client_config_id=None, download_client=DownloadClientType.SABNZBD, downloaded_path=r"e:\temp\pullbox_downloads\Release\file.cbr", ) @@ -163,6 +170,7 @@ async def test_resolve_local_path_requires_windows_component_boundary() -> None: ) session = SimpleNamespace(execute=AsyncMock(return_value=result)) download = SimpleNamespace( + download_client_config_id=None, download_client=DownloadClientType.SABNZBD, downloaded_path=r"E:\Temp\Pullbox_Downloads-old\Release\file.cbr", ) @@ -184,6 +192,7 @@ async def test_resolve_local_path_maps_windows_unc_path() -> None: ) session = SimpleNamespace(execute=AsyncMock(return_value=result)) download = SimpleNamespace( + download_client_config_id=None, download_client=DownloadClientType.SABNZBD, downloaded_path=r"\\SERVER\SHARE\Downloads\Release\file.cbr", ) @@ -206,6 +215,7 @@ async def test_resolve_local_path_rejects_windows_parent_traversal() -> None: ) session = SimpleNamespace(execute=AsyncMock(return_value=result)) download = SimpleNamespace( + download_client_config_id=None, download_client=DownloadClientType.SABNZBD, downloaded_path=r"E:\Temp\Pullbox_Downloads\..\outside\file.cbr", ) @@ -227,6 +237,7 @@ async def test_resolve_local_path_preserves_posix_literal_backslash() -> None: ) session = SimpleNamespace(execute=AsyncMock(return_value=result)) download = SimpleNamespace( + download_client_config_id=None, download_client=DownloadClientType.SABNZBD, downloaded_path="/remote/file.cbz", ) @@ -247,7 +258,9 @@ async def test_resolve_local_download_root_preserves_posix_literal_backslash() - download_dir="/downloads/Batman\\Superman" ) session = SimpleNamespace(execute=AsyncMock(return_value=result)) - download = SimpleNamespace(download_client=DownloadClientType.SABNZBD) + download = SimpleNamespace( + download_client_config_id=None, download_client=DownloadClientType.SABNZBD + ) root = await _resolve_local_download_root(session, download) @@ -284,3 +297,45 @@ def test_post_processing_integrity_exception_distinguishes_missing_source() -> N ) assert isinstance(exc, FileNotFoundError) + + +@pytest.mark.asyncio +async def test_resolve_local_path_uses_the_exact_client_when_types_repeat() -> None: + """Two clients of one type can map different paths; the download's own client wins.""" + from pullbox.models.download import DownloadClientType + from pullbox.tasks.download_post_processing_sources import _resolve_local_path + + exact = SimpleNamespace(remote_path="/tracker/complete", download_dir="/tracker-downloads") + session = SimpleNamespace( + get=AsyncMock(return_value=exact), + execute=AsyncMock(side_effect=AssertionError("type lookup must not run")), + ) + download = SimpleNamespace( + download_client=DownloadClientType.QBITTORRENT, + download_client_config_id=5, + downloaded_path="/tracker/complete/Release/file.cbz", + ) + + resolved = await _resolve_local_path(session, download) + + assert resolved == "/tracker-downloads/Release/file.cbz" + assert session.get.await_args.args[1] == 5 + + +@pytest.mark.asyncio +async def test_resolve_local_download_root_uses_the_exact_client_when_types_repeat() -> None: + from pullbox.models.download import DownloadClientType + from pullbox.tasks.download_post_processing_sources import _resolve_local_download_root + + session = SimpleNamespace( + get=AsyncMock(return_value=SimpleNamespace(download_dir="/tracker-downloads")), + execute=AsyncMock(side_effect=AssertionError("type lookup must not run")), + ) + download = SimpleNamespace( + download_client=DownloadClientType.QBITTORRENT, + download_client_config_id=5, + ) + + root = await _resolve_local_download_root(session, download) + + assert root == Path("/tracker-downloads") diff --git a/tests/tasks/test_download_task_latency.py b/tests/tasks/test_download_task_latency.py index 6f744816..f320724e 100644 --- a/tests/tasks/test_download_task_latency.py +++ b/tests/tasks/test_download_task_latency.py @@ -260,7 +260,7 @@ async def test_completed_download_triggers_immediate_post_processing( ) ) fake_service = MagicMock() - fake_service.get_client_for_type.return_value = fake_client + fake_service.get_client_for_identity.return_value = fake_client trigger = MagicMock() fake_logger = _FakeLogger() @@ -337,7 +337,7 @@ async def test_failed_download_emits_retry_scheduled_summary_once( ) ) fake_service = MagicMock() - fake_service.get_client_for_type.return_value = fake_client + fake_service.get_client_for_identity.return_value = fake_client fake_logger = _FakeLogger() monkeypatch.setattr(download_task, "get_session_factory", lambda: db_factory) @@ -399,7 +399,7 @@ async def test_failed_download_emits_terminal_failure_summary_once( ) ) fake_service = MagicMock() - fake_service.get_client_for_type.return_value = fake_client + fake_service.get_client_for_identity.return_value = fake_client fake_logger = _FakeLogger() monkeypatch.setattr(download_task, "get_session_factory", lambda: db_factory) @@ -455,7 +455,7 @@ async def test_heartbeat_only_poll_does_not_emit_lifecycle_summary( ) ) fake_service = MagicMock() - fake_service.get_client_for_type.return_value = fake_client + fake_service.get_client_for_identity.return_value = fake_client fake_logger = _FakeLogger() monkeypatch.setattr(download_task, "get_session_factory", lambda: db_factory) @@ -491,7 +491,7 @@ async def test_removed_externally_emits_summary_and_restores_issue( fake_client = MagicMock() fake_client.get_download_status = AsyncMock(side_effect=RuntimeError("not found")) fake_service = MagicMock() - fake_service.get_client_for_type.return_value = fake_client + fake_service.get_client_for_identity.return_value = fake_client fake_logger = _FakeLogger() monkeypatch.setattr(download_task, "get_session_factory", lambda: db_factory) diff --git a/tests/ui/test_settings_shell_ui_routes.py b/tests/ui/test_settings_shell_ui_routes.py index 18302e8a..eee22b90 100644 --- a/tests/ui/test_settings_shell_ui_routes.py +++ b/tests/ui/test_settings_shell_ui_routes.py @@ -100,6 +100,32 @@ async def test_settings_clients_prefers_persisted_server_status_message( assert "localhost and None: # type: ignore[no-untyped-def] + async with sec_db() as session: + session.add( + DownloadClientConfig( + name="qBittorrent", + client_type=DownloadClientType.QBITTORRENT, + url="http://qbit.test", + ) + ) + await session.commit() + + response = await authenticated_client.get("/settings?tab=clients") + + assert response.status_code == 200 + picker = _script_block( + response.text, + 'data-testid="settings-clients-picker-qbittorrent"', + "", + ) + assert "disabled" not in picker + assert "Already configured" not in response.text + async def test_settings_clients_describes_process_completed_as_recovery_sweep( self, authenticated_client, @@ -288,6 +314,51 @@ async def test_settings_indexers_show_jackett_source_and_retired_state( assert "Unavailable in Jackett" in response.text assert "Jackett owns tracker challenge resolution" in response.text + async def test_settings_indexers_offer_and_show_a_pinned_download_client( + self, + authenticated_client, + sec_db, + ) -> None: # type: ignore[no-untyped-def] + async with sec_db() as session: + tracker_client = DownloadClientConfig( + name="qBittorrent (tracker)", + client_type=DownloadClientType.QBITTORRENT, + url="http://qbit-tracker.test", + ) + session.add_all( + [ + tracker_client, + DownloadClientConfig( + name="Direct downloads", + client_type=DownloadClientType.DIRECT, + url="http://direct.test", + ), + ] + ) + await session.flush() + indexer = IndexerConfig( + name="Private Tracker", + indexer_type=IndexerType.TORZNAB, + url="http://tracker.example", + api_key=encrypt_secret("tracker-key"), + download_client_id=tracker_client.id, + ) + session.add(indexer) + await session.commit() + indexer_id = indexer.id + + response = await authenticated_client.get("/settings?tab=indexers") + + assert response.status_code == 200 + assert 'data-testid="settings-indexers-download-client"' in response.text + assert ( + f'data-testid="settings-indexers-client-{indexer_id}">Client: qBittorrent (tracker)' + in response.text + ) + choices = _script_block(response.text, "downloadClients: ", "\n") + assert "qBittorrent (tracker)" in choices + assert "Direct downloads" not in choices + async def test_settings_indexers_warn_when_managers_sync_the_same_tracker( self, authenticated_client,