refactor: make model sessions explicit

This commit is contained in:
jxxghp
2026-08-23 23:33:07 +08:00
parent 820582ab12
commit 6e69258e3c
65 changed files with 1299 additions and 2010 deletions
+43 -57
View File
@@ -10,11 +10,11 @@ import time as _time
import pytest
from app.db import decorators
from app.db import base as db_base
from app.db.models import subscribe as subscribe_module
from app.db.models.subscribe import Subscribe
from app.db.models.subscribehistory import SubscribeHistory
from app.db.session import SessionFactory, async_session_scope
from app.db.session import async_session_scope
from app.schemas.types import MediaSource, MediaType
TMDB = str(MediaSource.TMDB)
@@ -67,8 +67,11 @@ def test_exists_matches_async_twin(db):
db.add(_sub("并行", season=1))
sync_found = Subscribe.exists(db.session, MediaSource.TMDB, "9001", season=1)
async_found = asyncio.run(Subscribe.async_exists(
media_source=MediaSource.TMDB, media_id="9001", season=1))
async_found = db.run_async_session(
lambda session: Subscribe.async_exists(
session, media_source=MediaSource.TMDB, media_id="9001", season=1
)
)
assert sync_found.id == async_found.id
@@ -77,9 +80,11 @@ def test_history_queries_reuse_explicit_sessions(db, monkeypatch):
"""订阅历史同步/异步查询必须复用调用方会话。"""
row = db.add(_history("显式历史", media_id="8501"))
monkeypatch.setattr(
decorators,
"ScopedSession",
lambda: (_ for _ in ()).throw(AssertionError("不应创建额外同步会话")),
db_base,
"run_sync_transaction",
lambda _operation: (_ for _ in ()).throw(
AssertionError("不应创建额外同步事务")
),
)
assert SubscribeHistory.list_by_type(
db.session, MediaType.TV.value, page=1, count=10
@@ -92,9 +97,11 @@ def test_history_queries_reuse_explicit_sessions(db, monkeypatch):
"""验证异步订阅历史查询复用显式 AsyncSession。"""
async with async_session_scope() as session:
monkeypatch.setattr(
decorators,
"async_session_scope",
lambda: (_ for _ in ()).throw(AssertionError("不应创建额外异步会话")),
db_base,
"run_async_transaction",
lambda _operation: (_ for _ in ()).throw(
AssertionError("不应创建额外异步事务")
),
)
assert await SubscribeHistory.async_list_by_type(
session, MediaType.TV.value, page=1, count=10
@@ -109,44 +116,6 @@ def test_history_queries_reuse_explicit_sessions(db, monkeypatch):
asyncio.run(check())
def test_history_queries_keep_legacy_keyword_abi(db, monkeypatch):
"""旧插件关键字直调订阅历史查询时仍自动补入兼容会话。"""
row = db.add(_history("关键字历史", media_id="8601"))
opened_sync = []
monkeypatch.setattr(
decorators,
"ScopedSession",
lambda: (opened_sync.append(True) or SessionFactory()),
)
assert SubscribeHistory.list_by_type(
mtype=MediaType.TV.value, page=1, count=10
)
assert SubscribeHistory.exists(
media_source=MediaSource.TMDB, media_id="8601", season=1
).id == row.id
assert opened_sync == [True, True]
opened_async = []
original_scope = async_session_scope
def tracked_scope():
"""记录旧异步 ABI 创建的兼容会话作用域。"""
opened_async.append(True)
return original_scope()
monkeypatch.setattr(decorators, "async_session_scope", tracked_scope)
assert asyncio.run(SubscribeHistory.async_list_by_type(
mtype=MediaType.TV.value, page=1, count=10
))
assert asyncio.run(SubscribeHistory.async_list_by_type_and_username(
mtype=MediaType.TV.value, username="alice", page=1, count=10
))
assert asyncio.run(SubscribeHistory.async_exists(
media_source=MediaSource.TMDB, media_id="8601", season=1
)) is not None
assert opened_async == [True, True, True]
@pytest.mark.parametrize("media_id", [None, "", " "])
def test_exists_rejects_blank_media_id(db, media_id):
"""
@@ -232,7 +201,9 @@ def test_get_by_state_splits_comma_separated_states(db):
assert states == {"N", "R"}
assert len(Subscribe.get_by_state(db.session, "")) >= 3
assert {s.state for s in asyncio.run(Subscribe.async_get_by_state(state="N,R"))} == {"N", "R"}
assert {s.state for s in db.run_async_session(
lambda session: Subscribe.async_get_by_state(session, state="N,R")
)} == {"N", "R"}
def test_get_by_title_optionally_narrows_by_season(db):
@@ -281,8 +252,11 @@ def test_list_by_username_matches_async_twin(db):
("N", MediaType.TV.value)):
sync_names = sorted(s.name for s in
Subscribe.list_by_username(db.session, "alice", state, mtype))
async_names = sorted(s.name for s in asyncio.run(
Subscribe.async_list_by_username(username="alice", state=state, mtype=mtype)))
async_names = sorted(s.name for s in db.run_async_session(
lambda session: Subscribe.async_list_by_username(
session, username="alice", state=state, mtype=mtype
)
))
assert sync_names == async_names
@@ -319,8 +293,11 @@ def test_list_by_type_includes_the_window_start_boundary(db, frozen_now):
date=one_second_earlier))
names = {s.name for s in Subscribe.list_by_type(db.session, MediaType.TV.value, days=7)}
async_names = {s.name for s in asyncio.run(
Subscribe.async_list_by_type(mtype=MediaType.TV.value, days=7))}
async_names = {s.name for s in db.run_async_session(
lambda session: Subscribe.async_list_by_type(
session, mtype=MediaType.TV.value, days=7
)
)}
assert "窗口起点上" in names and "窗口起点前一秒" not in names
assert "窗口起点上" in async_names and "窗口起点前一秒" not in async_names
@@ -363,8 +340,11 @@ def test_history_list_by_type_matches_async_twin(db):
sync_names = [h.name for h in SubscribeHistory.list_by_type(
db.session, MediaType.TV.value, page=1, count=10)]
async_names = [h.name for h in asyncio.run(SubscribeHistory.async_list_by_type(
mtype=MediaType.TV.value, page=1, count=10))]
async_names = [h.name for h in db.run_async_session(
lambda session: SubscribeHistory.async_list_by_type(
session, mtype=MediaType.TV.value, page=1, count=10
)
)]
assert sync_names == async_names
@@ -392,7 +372,13 @@ def test_history_exists_matches_async_twin(db):
db.add(_history("并行历史", season=1, media_id="8401"))
sync_found = SubscribeHistory.exists(db.session, MediaSource.TMDB, "8401", season=1)
async_found = asyncio.run(SubscribeHistory.async_exists(
media_source=MediaSource.TMDB, media_id="8401", season=1))
async_found = db.run_async_session(
lambda session: SubscribeHistory.async_exists(
session,
media_source=MediaSource.TMDB,
media_id="8401",
season=1,
)
)
assert sync_found.id == async_found.id