Files
MoviePilot/tests/test_database_index_migration.py
T
InfinityPacerandjxxghp c727d9c9f1 test: 隔离单测中的站点原生资源 (#6452)
* test: isolate native site resources

* refactor(test): retain sites stub helper contract

* fix(test): preserve sites stub probe imports

---------

Co-authored-by: jxxghp <jxxghp@gmail.com>
2026-08-25 16:26:46 +08:00

422 lines
14 KiB
Python

import importlib
import os
from pathlib import Path
import subprocess
import sys
import uuid
import pytest
import sqlalchemy as sa
from alembic.migration import MigrationContext
from alembic.operations import Operations
try:
import psycopg2 as postgres_driver
from psycopg2 import sql
except ModuleNotFoundError:
import psycopg as postgres_driver
from psycopg import sql
MIGRATION_MODULE = "database.versions.93f8cb6a4d1e_2_2_4"
MEDIA_TABLES = (
"subscribe",
"subscribehistory",
"downloadhistory",
"transferhistory",
"downloadfailure",
"mediaserveritem",
)
LEGACY_IDENTITY_COLUMNS = {
"tmdbid",
"imdbid",
"tvdbid",
"doubanid",
"bangumiid",
"anilistid",
"mediaid",
}
IDENTITY_INDEX_SIGNATURES = {
"subscribe": {
"ix_subscribe_media_identity": (("media_source", "media_id"), False),
},
"subscribehistory": {
"ix_subscribehistory_media_identity": (
("media_source", "media_id"), False,
),
},
"downloadhistory": {
"ix_downloadhistory_media_identity": (
("media_source", "media_id"), False,
),
},
"transferhistory": {
"ix_transferhistory_media_identity": (
("media_source", "media_id"), False,
),
},
"downloadfailure": {
"ix_downloadfailure_media_identity_site": (
("type", "media_source", "media_id", "site"), False,
),
},
"mediaserveritem": {
"ix_mediaserveritem_media_identity_type": (
("media_source", "media_id", "item_type"), False,
),
},
}
CURRENT_SCHEMA_CHAIN_SCRIPT = """
from app.testing.bootstrap import ensure_sites_stub
# Alembic 会导入引用业务链的旧 revision,迁移验证不应加载本机站点原生制品。
ensure_sites_stub()
from alembic.config import Config
from alembic.script import ScriptDirectory
from sqlalchemy import inspect, MetaData, Table, text
from sqlalchemy.exc import IntegrityError
from app.runtime.config import settings
from app.db import get_engine
from app.startup.initializers.database import init_db, update_db
media_tables = {media_tables!r}
legacy_identity_columns = {legacy_identity_columns!r}
identity_index_signatures = {identity_index_signatures!r}
config = Config()
config.set_main_option("script_location", str(settings.ROOT_PATH / "database"))
heads = ScriptDirectory.from_config(config).get_heads()
assert len(heads) == 1, heads
init_db()
update_db()
update_db()
with get_engine().connect() as connection:
version = connection.execute(
text("SELECT version_num FROM alembic_version")
).scalar_one()
inspector = inspect(connection)
assert version == heads[0], (version, heads)
for table_name in media_tables:
columns = {{
column["name"]
for column in inspector.get_columns(table_name)
}}
indexes = {{
index["name"]: (
tuple(index.get("column_names") or ()),
bool(index.get("unique")),
)
for index in inspector.get_indexes(table_name)
}}
constraint_names = {{
constraint["name"]
for constraint in inspector.get_check_constraints(table_name)
}}
assert {{"media_source", "media_id"}}.issubset(columns), (
table_name,
columns,
)
assert legacy_identity_columns.isdisjoint(columns), (
table_name,
columns,
)
for index_name, signature in identity_index_signatures[table_name].items():
assert indexes.get(index_name) == signature, (
table_name,
index_name,
indexes,
)
constraint_name = f"ck_{{table_name}}_media_identity"
assert constraint_name in constraint_names, (
table_name,
constraint_names,
)
required_values = {{
"subscribe": {{"name": "constraint-test", "state": "N"}},
"subscribehistory": {{"name": "constraint-test"}},
"downloadhistory": {{
"path": "/constraint-test",
"type": "电影",
"title": "constraint-test",
}},
"transferhistory": {{"src_storage": "local"}},
"downloadfailure": {{"fingerprint": "constraint-test"}},
"mediaserveritem": {{}},
}}
invalid_identities = (
(None, "1"),
("acme.video", None),
("", "1"),
(" acme.video", "1"),
("acme.video ", "1"),
("Acme.Video", "1"),
("a" * 65, "1"),
("invalid:source", "1"),
("invalid source", "1"),
("acme.video", ""),
("acme.video", " "),
("acme.video", "0"),
)
for table_name in media_tables:
table = Table(table_name, MetaData(), autoload_with=connection)
constraint_name = f"ck_{{table_name}}_media_identity"
for media_source, media_id in (
(None, None),
("acme.video", "custom-1"),
):
values = {{
**required_values[table_name],
"media_source": media_source,
"media_id": media_id,
}}
savepoint = connection.begin_nested()
try:
connection.execute(table.insert(), values)
finally:
savepoint.rollback()
for media_source, media_id in invalid_identities:
values = {{
**required_values[table_name],
"media_source": media_source,
"media_id": media_id,
}}
try:
with connection.begin_nested():
connection.execute(table.insert(), values)
except IntegrityError as error:
assert constraint_name in str(error.orig), str(error.orig)
else:
raise AssertionError(
"格式非法的媒体身份未被具名检查约束拒绝: "
f"{{table_name}}, {{media_source!r}}, {{media_id!r}}"
)
""".format(
media_tables=MEDIA_TABLES,
legacy_identity_columns=LEGACY_IDENTITY_COLUMNS,
identity_index_signatures=IDENTITY_INDEX_SIGNATURES,
)
def _index_signatures(
connection,
table_name: str,
) -> dict[str, tuple[tuple[str, ...], bool]]:
"""返回索引名称到字段顺序及唯一性的映射。"""
return {
index["name"]: (
tuple(index.get("column_names") or ()),
bool(index.get("unique")),
)
for index in sa.inspect(connection).get_indexes(table_name)
}
def _bind_migration(monkeypatch, connection):
"""把历史 revision 绑定到当前 disposable connection。"""
migration = importlib.import_module(MIGRATION_MODULE)
context = MigrationContext.configure(connection)
monkeypatch.setattr(migration, "op", Operations(context))
return migration
def _run_current_schema_chain(
repository: Path,
environment: dict[str, str],
) -> None:
"""在隔离数据库中执行当前建表、完整升级及最终结构断言。"""
completed = subprocess.run(
[sys.executable, "-c", CURRENT_SCHEMA_CHAIN_SCRIPT],
cwd=repository,
env=environment,
capture_output=True,
text=True,
timeout=180,
check=False,
)
assert completed.returncode == 0, (
f"stdout:\n{completed.stdout}\n"
f"stderr:\n{completed.stderr}"
)
def test_index_migration_preserves_legacy_media_server_semantics(
monkeypatch,
) -> None:
"""旧字段存在时应保持 2.2.4 的索引替换与回滚语义。"""
engine = sa.create_engine("sqlite://")
metadata = sa.MetaData()
media_server = sa.Table(
"mediaserveritem",
metadata,
sa.Column("id", sa.Integer(), primary_key=True),
sa.Column("tmdbid", sa.Integer()),
sa.Column("item_type", sa.String()),
)
sa.Index("ix_mediaserveritem_id", media_server.c.id)
sa.Index("ix_mediaserveritem_tmdbid", media_server.c.tmdbid)
with engine.begin() as connection:
metadata.create_all(connection)
migration = _bind_migration(monkeypatch, connection)
migration.upgrade()
migration.upgrade()
upgraded = _index_signatures(connection, "mediaserveritem")
assert upgraded.get("ix_mediaserveritem_tmdbid_item_type") == (
("tmdbid", "item_type"), False,
)
assert "ix_mediaserveritem_tmdbid" not in upgraded
assert "ix_mediaserveritem_id" not in upgraded
migration.downgrade()
downgraded = _index_signatures(connection, "mediaserveritem")
assert "ix_mediaserveritem_tmdbid_item_type" not in downgraded
assert downgraded.get("ix_mediaserveritem_tmdbid") == (
("tmdbid",), False,
)
assert downgraded.get("ix_mediaserveritem_id") == (("id",), False)
def test_index_migration_skips_only_indexes_with_missing_columns(
monkeypatch,
) -> None:
"""当前 schema 应跳过旧字段索引,同时继续处理其他适用索引。"""
engine = sa.create_engine("sqlite://")
metadata = sa.MetaData()
media_server = sa.Table(
"mediaserveritem",
metadata,
sa.Column("id", sa.Integer(), primary_key=True),
sa.Column("media_source", sa.String()),
sa.Column("media_id", sa.String()),
sa.Column("item_type", sa.String()),
)
sa.Index(
"ix_mediaserveritem_media_identity_type",
media_server.c.media_source,
media_server.c.media_id,
media_server.c.item_type,
)
message = sa.Table(
"message",
metadata,
sa.Column("id", sa.Integer(), primary_key=True),
sa.Column("reg_time", sa.DateTime()),
)
sa.Index("ix_message_reg_time", message.c.reg_time)
with engine.begin() as connection:
metadata.create_all(connection)
migration = _bind_migration(monkeypatch, connection)
migration.upgrade()
media_indexes = _index_signatures(connection, "mediaserveritem")
message_indexes = _index_signatures(connection, "message")
assert "ix_mediaserveritem_tmdbid_item_type" not in media_indexes
assert media_indexes.get("ix_mediaserveritem_media_identity_type") == (
("media_source", "media_id", "item_type"), False,
)
assert "ix_message_reg_time" not in message_indexes
assert message_indexes.get("ix_message_reg_time_id") == (
("reg_time", "id"), False,
)
migration.downgrade()
media_indexes = _index_signatures(connection, "mediaserveritem")
message_indexes = _index_signatures(connection, "message")
assert "ix_mediaserveritem_tmdbid" not in media_indexes
assert media_indexes.get("ix_mediaserveritem_media_identity_type") == (
("media_source", "media_id", "item_type"), False,
)
assert message_indexes.get("ix_message_reg_time") == (
("reg_time",), False,
)
assert "ix_message_reg_time_id" not in message_indexes
def test_current_schema_reaches_current_alembic_head(tmp_path: Path) -> None:
"""真实 fresh 启动链应到动态解析的唯一 head,且重复升级保持幂等。"""
repository = Path(__file__).resolve().parents[1]
environment = os.environ.copy()
environment.update({
"CONFIG_DIR": str(tmp_path),
"DB_TYPE": "sqlite",
"SUPERUSER": "migration-test-admin",
"SUPERUSER_PASSWORD": "MigrationTestPassword123",
})
_run_current_schema_chain(repository, environment)
def test_current_schema_reaches_current_alembic_head_on_postgresql(
tmp_path: Path,
) -> None:
"""PostgreSQL fresh schema 应到唯一 head,不能被吞异常伪装成成功。"""
prefix = "MOVIEPILOT_TEST_POSTGRESQL_"
host = os.getenv(f"{prefix}HOST")
database = os.getenv(f"{prefix}DATABASE")
username = os.getenv(f"{prefix}USERNAME")
if not host or not database or not username:
pytest.skip("未配置隔离 PostgreSQL migration 测试库")
port = os.getenv(f"{prefix}PORT", "5432")
password = os.getenv(f"{prefix}PASSWORD", "")
schema = f"p1_db1_{uuid.uuid4().hex}"
with postgres_driver.connect(
host=host,
port=port,
dbname=database,
user=username,
password=password,
) as connection:
connection.autocommit = True
with connection.cursor() as cursor:
cursor.execute(
sql.SQL("CREATE SCHEMA {}").format(sql.Identifier(schema))
)
repository = Path(__file__).resolve().parents[1]
environment = os.environ.copy()
environment.update({
"CONFIG_DIR": str(tmp_path),
"DB_TYPE": "postgresql",
"DB_POSTGRESQL_HOST": host,
"DB_POSTGRESQL_PORT": port,
"DB_POSTGRESQL_DATABASE": database,
"DB_POSTGRESQL_USERNAME": username,
"DB_POSTGRESQL_PASSWORD": password,
"PGOPTIONS": f"-c search_path={schema}",
"SUPERUSER": "migration-test-admin",
"SUPERUSER_PASSWORD": "MigrationTestPassword123",
})
try:
_run_current_schema_chain(repository, environment)
finally:
with postgres_driver.connect(
host=host,
port=port,
dbname=database,
user=username,
password=password,
) as connection:
connection.autocommit = True
with connection.cursor() as cursor:
cursor.execute(
sql.SQL("DROP SCHEMA IF EXISTS {} CASCADE").format(
sql.Identifier(schema)
)
)