diff --git a/app/application/rss.py b/app/application/rss.py
index b3a514991..91675fd6d 100644
--- a/app/application/rss.py
+++ b/app/application/rss.py
@@ -304,6 +304,30 @@ class RssHelper:
parts = hostname.split(".")
return ".".join(parts[-2:]) if len(parts) >= 2 else hostname
+ @staticmethod
+ def normalize_url(url: Any) -> Optional[str]:
+ """校验并规范化 RSS HTTP 地址,保留合法地址中的 fragment。"""
+ if not isinstance(url, str):
+ return None
+ normalized_url = url.strip()
+ if not normalized_url:
+ return None
+ if any(character.isspace() for character in normalized_url):
+ return None
+ try:
+ parsed_url = urlparse(normalized_url)
+ hostname = parsed_url.hostname
+ port = parsed_url.port
+ except ValueError:
+ return None
+ if (
+ parsed_url.scheme.lower() not in ("http", "https")
+ or not hostname
+ or (port is not None and not 0 <= port <= 65535)
+ ):
+ return None
+ return normalized_url
+
@staticmethod
def _parse_publish_time(value: str):
"""将 RSS 常见日期表达解析为 datetime,无法解析时返回 None。"""
@@ -348,8 +372,10 @@ class RssHelper:
"""
# 开始处理
ret_array: list = []
- if not url:
+ normalized_url = self.normalize_url(url)
+ if normalized_url is None:
return False
+ url = normalized_url
http_port, _, _ = _require_rss_ports()
diff --git a/app/chain/torrents.py b/app/chain/torrents.py
index f53faf480..69c75e2c5 100644
--- a/app/chain/torrents.py
+++ b/app/chain/torrents.py
@@ -318,11 +318,12 @@ class TorrentsChain(ChainBase):
if not site:
logger.error(f'站点 {domain} 不存在!')
return []
- if not site.get("rss"):
- logger.error(f'站点 {domain} 未配置RSS地址!')
+ rss_url = RssHelper.normalize_url(site.get("rss"))
+ if rss_url is None:
+ logger.warning(f'站点 {domain} RSS地址无效,跳过获取')
return []
# 解析RSS
- rss_items = RssHelper().parse(site.get("rss"), True if site.get("proxy") else False,
+ rss_items = RssHelper().parse(rss_url, True if site.get("proxy") else False,
timeout=int(site.get("timeout") or 30),
ua=site.get("ua") if site.get("ua") else None)
if rss_items is None:
diff --git a/app/db/oper/subscriptionsearch.py b/app/db/oper/subscriptionsearch.py
index b1651927f..641a581ae 100644
--- a/app/db/oper/subscriptionsearch.py
+++ b/app/db/oper/subscriptionsearch.py
@@ -5,6 +5,8 @@ from typing import Mapping, Optional
from uuid import uuid4
from sqlalchemy import and_, case, func, or_, select, update
+from sqlalchemy.dialects.postgresql import insert as pg_insert
+from sqlalchemy.dialects.sqlite import insert as sqlite_insert
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
@@ -36,6 +38,10 @@ class SubscriptionSearchOper(DbOper):
"""创建批次,并以活动键合并同一订阅的重叠搜索入口。"""
if not isinstance(self._db, Session):
raise RuntimeError("订阅搜索入队需要调用方提供同步 Session")
+ dialect = self._db.get_bind().dialect.name
+ if dialect not in {"postgresql", "sqlite"}:
+ raise RuntimeError(f"订阅搜索入队不支持数据库方言:{dialect}")
+ insert = pg_insert if dialect == "postgresql" else sqlite_insert
now = utc_now_text()
batch = SubscriptionSearchBatch(
batch_id=uuid4().hex,
@@ -54,17 +60,12 @@ class SubscriptionSearchOper(DbOper):
for position, subscription_id in enumerate(dict.fromkeys(subscription_ids)):
active_key = f"subscription:{subscription_id}"
available_at = (
- available_at_by_subscription.get(subscription_id, now)
- if available_at_by_subscription
- else now
+ available_at_by_subscription.get(subscription_id, now) if available_at_by_subscription else now
)
- initial_phase = (
- "scheduled"
- if source == "new" and available_at > now
- else "queued"
- )
- task = SubscriptionSearchTask(
- task_id=uuid4().hex,
+ initial_phase = "scheduled" if source == "new" and available_at > now else "queued"
+ task_id = uuid4().hex
+ statement = insert(SubscriptionSearchTask).values(
+ task_id=task_id,
batch_id=batch.batch_id,
subscription_id=subscription_id,
active_key=active_key,
@@ -77,72 +78,64 @@ class SubscriptionSearchOper(DbOper):
created_at=now,
updated_at=now,
)
- try:
- with self._db.begin_nested():
- self._db.add(task)
- self._db.flush()
- created += 1
- except IntegrityError:
- coalesced += 1
- promote_queued_task = and_(
- SubscriptionSearchTask.priority < priority,
- SubscriptionSearchTask.state == "queued",
- )
- refresh_queued_task = and_(
- refresh_pending,
- SubscriptionSearchTask.state == "queued",
- SubscriptionSearchTask.pending_site_ids.is_not(None),
- )
- execute_dml(
- self._db,
- update(SubscriptionSearchTask)
- .where(SubscriptionSearchTask.active_key == active_key)
- .values(
- source=case(
- (SubscriptionSearchTask.priority < priority, source),
- else_=SubscriptionSearchTask.source,
- ),
- priority=case(
- (SubscriptionSearchTask.priority < priority, priority),
- else_=SubscriptionSearchTask.priority,
- ),
- phase=case(
- (or_(promote_queued_task, refresh_queued_task), "queued"),
- else_=SubscriptionSearchTask.phase,
- ),
- last_error=case(
- (or_(promote_queued_task, refresh_queued_task), None),
- else_=SubscriptionSearchTask.last_error,
- ),
- # 用户重搜和已到期的新周期恢复完整范围;普通合并仍保留补查游标。
- pending_site_ids=case(
- (or_(
- and_(SubscriptionSearchTask.state == "queued", source in {"manual", "targeted"}),
- refresh_queued_task,
- ), None),
- else_=SubscriptionSearchTask.pending_site_ids,
- ),
- available_at=case(
- (
- or_(
- SubscriptionSearchTask.available_at.is_(None),
- SubscriptionSearchTask.available_at > available_at,
- ),
- available_at,
- ),
- else_=SubscriptionSearchTask.available_at,
- ),
- updated_at=now,
+ promote_queued_task = and_(
+ SubscriptionSearchTask.priority < priority,
+ SubscriptionSearchTask.state == "queued",
+ )
+ refresh_queued_task = and_(
+ refresh_pending,
+ SubscriptionSearchTask.state == "queued",
+ SubscriptionSearchTask.pending_site_ids.is_not(None),
+ )
+ # 唯一键仲裁与合并在同一条语句内完成,避免旧任务结束时丢失入队请求。
+ # 只处理 active_key 冲突,其余约束错误继续交给调用方事务处理。
+ statement = statement.on_conflict_do_update(
+ index_elements=[SubscriptionSearchTask.active_key],
+ set_=dict(
+ source=case(
+ (SubscriptionSearchTask.priority < priority, source),
+ else_=SubscriptionSearchTask.source,
),
- execution_options={"synchronize_session": False},
- )
- active_task = self._db.execute(
- select(SubscriptionSearchTask).where(
- SubscriptionSearchTask.active_key == active_key
- )
- ).scalars().first()
- if active_task is not None:
- active_batch_ids.append(active_task.batch_id)
+ priority=case(
+ (SubscriptionSearchTask.priority < priority, priority),
+ else_=SubscriptionSearchTask.priority,
+ ),
+ phase=case(
+ (or_(promote_queued_task, refresh_queued_task), "queued"),
+ else_=SubscriptionSearchTask.phase,
+ ),
+ last_error=case(
+ (or_(promote_queued_task, refresh_queued_task), None),
+ else_=SubscriptionSearchTask.last_error,
+ ),
+ # 用户重搜和已到期的新周期恢复完整范围;普通合并仍保留补查游标。
+ pending_site_ids=case(
+ (or_(
+ and_(SubscriptionSearchTask.state == "queued", source in {"manual", "targeted"}),
+ refresh_queued_task,
+ ), None),
+ else_=SubscriptionSearchTask.pending_site_ids,
+ ),
+ available_at=case(
+ (
+ or_(
+ SubscriptionSearchTask.available_at.is_(None),
+ SubscriptionSearchTask.available_at > available_at,
+ ),
+ available_at,
+ ),
+ else_=SubscriptionSearchTask.available_at,
+ ),
+ updated_at=now,
+ ),
+ ).returning(SubscriptionSearchTask.task_id, SubscriptionSearchTask.batch_id)
+ stored_task_id, stored_batch_id = self._db.execute(statement).one()
+ # 合并保留原任务及批次身份,无需依赖数据库专有的系统列判断插入结果。
+ if stored_task_id == task_id:
+ created += 1
+ else:
+ coalesced += 1
+ active_batch_ids.append(stored_batch_id)
batch.total_count = created
if created == 0:
batch.state = "completed"
diff --git a/tests/test_rss_url_validation.py b/tests/test_rss_url_validation.py
new file mode 100644
index 000000000..70cd81bd3
--- /dev/null
+++ b/tests/test_rss_url_validation.py
@@ -0,0 +1,101 @@
+from types import SimpleNamespace
+from unittest.mock import Mock
+
+import pytest
+
+import app.chain.torrents as torrents_module
+from app.application.rss import RssHelper, configure_rss_ports, reset_rss_ports
+from app.chain.torrents import TorrentsChain
+
+
+@pytest.fixture
+def rss_ports():
+ """装配可记录调用的 RSS 端口,确保用例不产生真实网络请求。"""
+ http_port = Mock()
+ parser_port = Mock()
+ parser_port.parse.return_value = None
+ configure_rss_ports(http=http_port, browser=Mock(), parser=parser_port)
+ yield http_port
+ reset_rss_ports()
+
+
+@pytest.mark.parametrize(
+ "url",
+ [None, "", " ", 123, b"https://example.com/rss", "#", "#fragment",
+ "/rss", "rss.xml", "ftp://example.com/rss", "https://bad host/rss",
+ "https://example.com:abc/rss", "https://example.com:65536/rss"],
+)
+def test_rss_helper_rejects_invalid_urls_before_http_request(url, rss_ports):
+ """无效 RSS 地址应返回普通错误,并在端口层之前被拒绝。"""
+ assert RssHelper().parse(url) is False
+ rss_ports.get.assert_not_called()
+
+
+def test_rss_helper_strips_whitespace_and_allows_http_fragment(rss_ports):
+ """合法 HTTP(S) 地址即使带 fragment 也应去除外层空白后请求。"""
+ rss_ports.get.return_value = SimpleNamespace(
+ status_code=200,
+ content=b"",
+ text="",
+ reason="OK",
+ )
+ rss_ports.decode_xml.return_value = ""
+
+ assert RssHelper().parse(" http://example.com:8080/rss#item\n") == []
+ assert rss_ports.get.call_args.kwargs["url"] == "http://example.com:8080/rss#item"
+
+
+@pytest.mark.parametrize("parse_result, renew_expected", [(False, False), (None, True)])
+def test_torrents_rss_preserves_false_and_none_contract(
+ monkeypatch, parse_result, renew_expected, rss_ports
+):
+ """RSS 普通错误与过期结果应继续分别对应不续期和自动续期。"""
+ site = {
+ "id": 1,
+ "name": "测试站点",
+ "rss": "https://example.com/rss",
+ "proxy": False,
+ "timeout": 30,
+ "ua": None,
+ }
+ sites_helper = Mock()
+ sites_helper.get_indexer.return_value = site
+ monkeypatch.setattr(torrents_module, "SitesHelper", lambda: sites_helper)
+ monkeypatch.setattr(RssHelper, "parse", lambda *_args, **_kwargs: parse_result)
+
+ chain = TorrentsChain()
+ renew = Mock()
+ monkeypatch.setattr(chain, "_TorrentsChain__renew_rss_url", renew)
+
+ assert chain.rss("example.com") == []
+ assert renew.called is renew_expected
+
+
+def test_torrents_rss_warns_once_and_skips_invalid_site_url(monkeypatch, rss_ports):
+ """站点 RSS 配置无效时只记录简短告警,不请求或自动续期。"""
+ site = {
+ "id": 1,
+ "name": "测试站点",
+ "rss": "not-a-url?passkey=secret",
+ "proxy": False,
+ "timeout": 30,
+ "ua": None,
+ }
+ sites_helper = Mock()
+ sites_helper.get_indexer.return_value = site
+ warnings = []
+ errors = []
+ monkeypatch.setattr(torrents_module, "SitesHelper", lambda: sites_helper)
+ monkeypatch.setattr(torrents_module.logger, "warning", warnings.append)
+ monkeypatch.setattr(torrents_module.logger, "error", errors.append)
+
+ chain = TorrentsChain()
+ renew = Mock()
+ monkeypatch.setattr(chain, "_TorrentsChain__renew_rss_url", renew)
+
+ assert chain.rss("example.com") == []
+ assert rss_ports.get.call_count == 0
+ assert errors == []
+ assert warnings == ["站点 example.com RSS地址无效,跳过获取"]
+ assert "passkey=secret" not in warnings[0]
+ renew.assert_not_called()
diff --git a/tests/test_subscription_search_upsert.py b/tests/test_subscription_search_upsert.py
new file mode 100644
index 000000000..7a338b1c0
--- /dev/null
+++ b/tests/test_subscription_search_upsert.py
@@ -0,0 +1,87 @@
+"""验证活动任务原子合并及非预期约束错误的事务边界。"""
+
+import os
+from concurrent.futures import ThreadPoolExecutor
+from threading import Barrier
+from types import SimpleNamespace
+
+import pytest
+from sqlalchemy import create_engine, event, select
+from sqlalchemy.exc import IntegrityError
+from sqlalchemy.orm import Session, sessionmaker
+
+from app.db.adapters.subscriptionsearch import TransactionalSubscriptionSearchRepository
+from app.db.models.subscriptionsearch import SubscriptionSearchBatch, SubscriptionSearchTask
+
+
+@pytest.fixture
+def queue(tmp_path):
+ """默认使用隔离 SQLite;显式测试 URL 只允许指向可清空的专用数据库。"""
+ engine = create_engine(os.environ.get("MOVIEPILOT_TEST_UPSERT_URL", f"sqlite:///{tmp_path / 'queue.db'}"))
+ tables = [SubscriptionSearchBatch.__table__, SubscriptionSearchTask.__table__]
+ for table in tables:
+ table.create(engine, checkfirst=True)
+ errors = []
+ event.listen(engine, "handle_error", errors.append)
+ try:
+ yield TransactionalSubscriptionSearchRepository(sessionmaker(bind=engine)), engine, errors
+ finally:
+ for table in reversed(tables):
+ table.drop(engine, checkfirst=True)
+ engine.dispose()
+
+
+def test_concurrent_enqueue_keeps_one_task_without_driver_errors(queue):
+ """独立事务同时入队时只有一个创建者,其余请求关联原批次且无驱动异常。"""
+ repository, engine, errors = queue
+ barrier = Barrier(4)
+
+ def enqueue():
+ barrier.wait(timeout=10)
+ return repository.enqueue(subscription_ids=(1,), source="fallback", priority=10)
+
+ with ThreadPoolExecutor(max_workers=4) as executor:
+ results = list(executor.map(lambda _: enqueue(), range(4)))
+ assert sum(result.created_count for result in results) == 1
+ assert sum(result.coalesced_count for result in results) == 3
+ created = next(result for result in results if result.created_count)
+ assert all(result.active_batch_ids == (created.batch.batch_id,) for result in results)
+ with Session(engine) as session:
+ assert len(session.scalars(select(SubscriptionSearchTask)).all()) == 1
+ assert errors == []
+
+
+def test_mixed_enqueue_preserves_batch_identity_and_terminal_reenqueue(queue):
+ """同批混合创建和合并保留批次关联,终态释放活动键后可再次创建任务。"""
+ repository, engine, errors = queue
+ first = repository.enqueue(subscription_ids=(1,), source="fallback", priority=10)
+ mixed = repository.enqueue(subscription_ids=(1, 2, 2), source="manual", priority=100)
+ assert (mixed.created_count, mixed.coalesced_count) == (1, 1)
+ assert mixed.active_batch_ids == (mixed.batch.batch_id, first.batch.batch_id)
+ task = repository.claim_next(owner="worker")
+ assert task.subscription_id == 1
+ assert repository.finish_task(task_id=task.task_id, lease_token=task.lease_token, state="completed")
+ again = repository.enqueue(subscription_ids=(1,), source="manual", priority=100)
+ assert (again.created_count, again.coalesced_count) == (1, 0)
+ with Session(engine) as session:
+ rows = session.scalars(select(SubscriptionSearchTask).where(SubscriptionSearchTask.subscription_id == 1)).all()
+ assert len(rows) == 2
+ assert sum(row.active_key is not None for row in rows) == 1
+ assert errors == []
+
+
+def test_unrelated_constraint_error_rolls_back_entire_batch(queue, monkeypatch):
+ """task_id 冲突必须传播,不能误计为活动键合并或留下半个批次。"""
+ repository, engine, errors = queue
+ repository.enqueue(subscription_ids=(1,), source="fallback", priority=10)
+ with Session(engine) as session:
+ old_task = session.scalar(select(SubscriptionSearchTask.task_id))
+
+ identifiers = iter(["new-batch", "new-task", old_task])
+ monkeypatch.setattr("app.db.oper.subscriptionsearch.uuid4", lambda: SimpleNamespace(hex=next(identifiers)))
+ with pytest.raises(IntegrityError):
+ repository.enqueue(subscription_ids=(2, 3), source="fallback", priority=10)
+ with Session(engine) as session:
+ assert len(session.scalars(select(SubscriptionSearchBatch)).all()) == 1
+ assert len(session.scalars(select(SubscriptionSearchTask)).all()) == 1
+ assert len(errors) == 1