mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-08-28 19:47:41 +08:00
refactor: govern background tasks and query ownership
This commit is contained in:
@@ -22,6 +22,7 @@ from app.startup.context import (
|
||||
HostRuntime,
|
||||
SubscriptionRuntime,
|
||||
)
|
||||
from app.runtime.tasks import TaskRegistry, get_task_registry
|
||||
|
||||
|
||||
def get_host_runtime(request: Request) -> HostRuntime:
|
||||
@@ -32,6 +33,20 @@ def get_host_runtime(request: Request) -> HostRuntime:
|
||||
return runtime
|
||||
|
||||
|
||||
def get_background_task_registry(
|
||||
runtime: HostRuntime = Depends(get_host_runtime),
|
||||
) -> TaskRegistry:
|
||||
"""返回当前 lifespan 统一管理的后台任务登记器。"""
|
||||
return runtime.tasks
|
||||
|
||||
|
||||
def resolve_background_task_registry(value: object) -> TaskRegistry:
|
||||
"""兼容直接调用 endpoint 的旧入口,并优先使用注入的任务登记器。"""
|
||||
if isinstance(value, TaskRegistry):
|
||||
return value
|
||||
return get_task_registry()
|
||||
|
||||
|
||||
def get_api_runtime_config(
|
||||
runtime: HostRuntime = Depends(get_host_runtime),
|
||||
) -> ApiRuntimeConfig:
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
from typing import Any, cast
|
||||
|
||||
from fastapi import BackgroundTasks, Depends
|
||||
from fastapi import Depends
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
@@ -38,6 +38,8 @@ from app.runtime.events import eventmanager
|
||||
from app.runtime.log import logger
|
||||
from app.schemas.types import EventType
|
||||
from app.startup.context import HostRuntime
|
||||
from app.api.context import get_background_task_registry, resolve_background_task_registry
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
|
||||
|
||||
async def _publish_subscribe_deleted(
|
||||
@@ -94,14 +96,14 @@ def get_delete_subscriptions_by_identity_command(
|
||||
|
||||
|
||||
def get_search_subscriptions_command(
|
||||
background_tasks: BackgroundTasks,
|
||||
task_registry: TaskRegistry = Depends(get_background_task_registry),
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
runtime: HostRuntime = Depends(get_host_runtime),
|
||||
) -> SearchSubscriptionsCommand:
|
||||
"""组装手工订阅搜索用例,并把调度延迟到响应后的后台任务。"""
|
||||
def schedule_search(subscribe_id: int | None, state: str | None) -> None:
|
||||
"""按历史参数提交订阅搜索调度任务。"""
|
||||
background_tasks.add_task(
|
||||
resolve_background_task_registry(task_registry).create_sync(
|
||||
start_scheduler_job,
|
||||
job_id="subscribe_search",
|
||||
sid=subscribe_id,
|
||||
|
||||
@@ -4,7 +4,7 @@ import json
|
||||
import time
|
||||
from typing import Protocol, Union, Any, List, Optional
|
||||
|
||||
from fastapi import BackgroundTasks, Depends, Request
|
||||
from fastapi import Depends, Request
|
||||
from starlette.responses import PlainTextResponse
|
||||
|
||||
from app.schemas.message import MessageClearBefore as _SchemaMessageClearBefore
|
||||
@@ -32,6 +32,8 @@ from app.runtime.extensions.service_config import ServiceConfigHelper
|
||||
from app.runtime.log import logger
|
||||
from app.adapters.external.wechat_crypt import WXBizMsgCrypt
|
||||
from app.schemas.types import NotificationChannel, SystemConfigKey
|
||||
from app.api.context import get_background_task_registry, resolve_background_task_registry
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
|
||||
router = ResponseAPIRouter()
|
||||
|
||||
@@ -116,7 +118,7 @@ def start_message_chain(body: Any, form: Any, args: Any):
|
||||
|
||||
@router.post("/", summary="接收用户消息", response_model=_SchemaResponse[None])
|
||||
async def user_message(
|
||||
background_tasks: BackgroundTasks,
|
||||
task_registry: Annotated[TaskRegistry, Depends(get_background_task_registry)],
|
||||
request: Request,
|
||||
_: _SchemaTokenPayload = Depends(verify_apitoken),
|
||||
):
|
||||
@@ -150,7 +152,9 @@ async def user_message(
|
||||
list(form.keys()) if form else [],
|
||||
image_markers,
|
||||
)
|
||||
background_tasks.add_task(start_message_chain, body, form, args)
|
||||
resolve_background_task_registry(task_registry).create_sync(
|
||||
start_message_chain, body, form, args, owner="api.message.user"
|
||||
)
|
||||
return _SchemaResponse(success=True)
|
||||
|
||||
|
||||
|
||||
@@ -60,6 +60,8 @@ from app.adapters.system.plugin.package import PluginPackageManager
|
||||
from app.application.database import DatabaseWorkerOverloadedError
|
||||
from app.runtime.log import logger
|
||||
from app.schemas.types import SystemConfigKey
|
||||
from app.api.context import get_background_task_registry, resolve_background_task_registry
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
|
||||
router = ResponseAPIRouter()
|
||||
_plugin_release_refresh_tasks: set[asyncio.Task] = set()
|
||||
@@ -112,11 +114,17 @@ async def _refresh_plugin_release_versions(plugin_id: str, repo_url: str) -> Non
|
||||
logger.warning(f"后台刷新插件 {plugin_id} Release 列表失败:{e}")
|
||||
|
||||
|
||||
def _schedule_plugin_release_refresh(plugin_id: str, repo_url: str) -> None:
|
||||
def _schedule_plugin_release_refresh(
|
||||
plugin_id: str, repo_url: str, task_registry: TaskRegistry | None = None
|
||||
) -> None:
|
||||
"""
|
||||
保留后台任务引用,避免任务被回收,同时让 helper 负责同仓库强刷合并。
|
||||
"""
|
||||
task = asyncio.create_task(_refresh_plugin_release_versions(plugin_id, repo_url))
|
||||
registry = resolve_background_task_registry(task_registry)
|
||||
task = registry.create(
|
||||
_refresh_plugin_release_versions(plugin_id, repo_url),
|
||||
owner="api.plugin.release_refresh",
|
||||
)
|
||||
_plugin_release_refresh_tasks.add(task)
|
||||
|
||||
def _discard_task(completed_task: asyncio.Task) -> None:
|
||||
@@ -368,6 +376,7 @@ async def plugin_releases(
|
||||
_: ApiPrincipal = Depends(get_current_active_superuser_async),
|
||||
repo_url: Optional[str] = "",
|
||||
force: bool = False,
|
||||
task_registry: TaskRegistry = Depends(get_background_task_registry),
|
||||
) -> dict:
|
||||
"""
|
||||
查询指定插件可直接安装的 GitHub Release 版本。
|
||||
@@ -404,7 +413,11 @@ async def plugin_releases(
|
||||
)
|
||||
release_items = await plugin_helper.async_get_plugin_release_versions(plugin_id, repo_url)
|
||||
if force and has_release_cache:
|
||||
_schedule_plugin_release_refresh(plugin_id, repo_url)
|
||||
_schedule_plugin_release_refresh(
|
||||
plugin_id,
|
||||
repo_url,
|
||||
resolve_background_task_registry(task_registry),
|
||||
)
|
||||
items = []
|
||||
for item in release_items:
|
||||
version = item.get("version")
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from typing import List, Any, Dict, Optional
|
||||
|
||||
from fastapi import Depends, HTTPException
|
||||
from starlette.background import BackgroundTasks
|
||||
from typing import Annotated
|
||||
|
||||
from app.schemas.common import JsonObject as _SchemaJsonObject
|
||||
from app.schemas.response import Response as _SchemaResponse
|
||||
@@ -42,6 +42,8 @@ from app.runtime.log import logger
|
||||
from app.application.scheduling import Scheduler
|
||||
from app.schemas.types import SystemConfigKey, MediaType
|
||||
from app.domain import site as site_rules
|
||||
from app.api.context import get_background_task_registry, resolve_background_task_registry
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
|
||||
router = ResponseAPIRouter()
|
||||
|
||||
@@ -169,13 +171,15 @@ async def update_site(
|
||||
|
||||
@router.get("/cookiecloud", summary="CookieCloud同步", response_model=_SchemaResponse[None])
|
||||
async def cookie_cloud_sync(
|
||||
background_tasks: BackgroundTasks,
|
||||
task_registry: Annotated[TaskRegistry, Depends(get_background_task_registry)],
|
||||
_: ApiPrincipal = Depends(get_current_active_superuser_async),
|
||||
) -> Any:
|
||||
"""
|
||||
运行CookieCloud同步站点信息
|
||||
"""
|
||||
background_tasks.add_task(Scheduler().start, job_id="cookiecloud")
|
||||
resolve_background_task_registry(task_registry).create_sync(
|
||||
Scheduler().start, job_id="cookiecloud", owner="api.site.cookiecloud_sync"
|
||||
)
|
||||
return _SchemaResponse(success=True, message="CookieCloud同步任务已启动!")
|
||||
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
from typing import List, Any, Annotated, Optional
|
||||
|
||||
import cn2an
|
||||
from fastapi import Request, BackgroundTasks, Depends, HTTPException, Header
|
||||
from fastapi import Request, Depends, HTTPException, Header
|
||||
|
||||
from app.schemas.common import IdData as _SchemaIdData
|
||||
from app.schemas.response import Response as _SchemaResponse
|
||||
@@ -12,6 +12,10 @@ from app.schemas.token import TokenPayload as _SchemaTokenPayload
|
||||
from app.schemas.workflow import MediaInfo as _SchemaMediaInfo
|
||||
from app.schemas.workflow import Subscribe as _SchemaSubscribe
|
||||
from app.api.response import ResponseAPIRouter
|
||||
from app.api.context import (
|
||||
get_background_task_registry,
|
||||
resolve_background_task_registry,
|
||||
)
|
||||
from app.chain.subscribe import SubscribeChain
|
||||
from app.runtime.events import eventmanager
|
||||
from app.domain.context import MediaInfo
|
||||
@@ -52,6 +56,7 @@ from app.api.dependencies.subscription import (
|
||||
)
|
||||
from app.adapters.external.server import MoviePilotServerHelper
|
||||
from app.application.scheduling import Scheduler
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
from app.schemas.event import SubscribeModifiedEventData
|
||||
from app.schemas.types import (
|
||||
MUSIC_ENTITY_ALBUM,
|
||||
@@ -489,7 +494,7 @@ async def delete_subscribe_by_media_identity(
|
||||
)
|
||||
async def seerr_subscribe(
|
||||
request: Request,
|
||||
background_tasks: BackgroundTasks,
|
||||
task_registry: Annotated[TaskRegistry, Depends(get_background_task_registry)],
|
||||
authorization: Annotated[str | None, Header()] = None,
|
||||
) -> Any:
|
||||
"""
|
||||
@@ -521,7 +526,7 @@ async def seerr_subscribe(
|
||||
user_name = req_json.get("request", {}).get("requestedBy_username")
|
||||
# 添加订阅
|
||||
if media_type == MediaType.MOVIE:
|
||||
background_tasks.add_task(
|
||||
resolve_background_task_registry(task_registry).create_sync(
|
||||
start_subscribe_add,
|
||||
mtype=media_type,
|
||||
media_source=MediaSource.TMDB,
|
||||
@@ -531,6 +536,7 @@ async def seerr_subscribe(
|
||||
# 电影不传季号,避免被误判为剧集(S00)并污染通知标题
|
||||
season=None,
|
||||
username=user_name,
|
||||
owner="api.subscribe.seerr",
|
||||
)
|
||||
else:
|
||||
seasons = []
|
||||
@@ -543,7 +549,7 @@ async def seerr_subscribe(
|
||||
]
|
||||
break
|
||||
for season in seasons:
|
||||
background_tasks.add_task(
|
||||
resolve_background_task_registry(task_registry).create_sync(
|
||||
start_subscribe_add,
|
||||
mtype=media_type,
|
||||
media_source=MediaSource.TMDB,
|
||||
@@ -552,6 +558,7 @@ async def seerr_subscribe(
|
||||
year="",
|
||||
season=season,
|
||||
username=user_name,
|
||||
owner="api.subscribe.seerr",
|
||||
)
|
||||
|
||||
return _SchemaResponse(success=True)
|
||||
|
||||
@@ -1,11 +1,13 @@
|
||||
from typing import Any, Annotated
|
||||
|
||||
from fastapi import BackgroundTasks, Request, Depends
|
||||
from fastapi import Depends, Request
|
||||
|
||||
from app.schemas.response import Response as _SchemaResponse
|
||||
from app.api.response import ResponseAPIRouter
|
||||
from app.chain.webhook import WebhookChain
|
||||
from app.adapters.web.security.access import verify_apitoken
|
||||
from app.api.context import get_background_task_registry, resolve_background_task_registry
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
|
||||
router = ResponseAPIRouter()
|
||||
|
||||
@@ -19,7 +21,7 @@ def start_webhook_chain(body: Any, form: Any, args: Any):
|
||||
|
||||
@router.post("/", summary="Webhook消息响应", response_model=_SchemaResponse[None])
|
||||
async def webhook_message(
|
||||
background_tasks: BackgroundTasks,
|
||||
task_registry: Annotated[TaskRegistry, Depends(get_background_task_registry)],
|
||||
request: Request,
|
||||
_: Annotated[str, Depends(verify_apitoken)],
|
||||
) -> Any:
|
||||
@@ -29,13 +31,15 @@ async def webhook_message(
|
||||
body = await request.body()
|
||||
form = await request.form()
|
||||
args = request.query_params
|
||||
background_tasks.add_task(start_webhook_chain, body, form, args)
|
||||
resolve_background_task_registry(task_registry).create_sync(
|
||||
start_webhook_chain, body, form, args, owner="api.webhook.message"
|
||||
)
|
||||
return _SchemaResponse(success=True)
|
||||
|
||||
|
||||
@router.get("/", summary="Webhook消息响应", response_model=_SchemaResponse[None])
|
||||
async def webhook_message_get(
|
||||
background_tasks: BackgroundTasks,
|
||||
task_registry: Annotated[TaskRegistry, Depends(get_background_task_registry)],
|
||||
request: Request,
|
||||
_: Annotated[str, Depends(verify_apitoken)],
|
||||
) -> Any:
|
||||
@@ -43,5 +47,7 @@ async def webhook_message_get(
|
||||
Webhook响应,配置请求中需要添加参数:token=API_TOKEN&source=媒体服务器名
|
||||
"""
|
||||
args = request.query_params
|
||||
background_tasks.add_task(start_webhook_chain, None, None, args)
|
||||
resolve_background_task_registry(task_registry).create_sync(
|
||||
start_webhook_chain, None, None, args, owner="api.webhook.message"
|
||||
)
|
||||
return _SchemaResponse(success=True)
|
||||
|
||||
@@ -7,6 +7,7 @@ from threading import Lock
|
||||
from typing import Awaitable, Callable, Dict, Iterable, List, Optional, Tuple, Union
|
||||
|
||||
from app.schemas.types import NotificationChannel
|
||||
from app.runtime.tasks import get_task_registry
|
||||
|
||||
# Agent 选择按钮回调前缀(新旧两种格式都必须继续兼容)
|
||||
AGENT_CHOICE_PREFIX = "agent_interaction:choice:"
|
||||
@@ -185,7 +186,10 @@ def create_web_agent_background_task(
|
||||
coroutine: Awaitable[object],
|
||||
) -> asyncio.Task[object]:
|
||||
"""登记 Web Agent 后台任务,使应用关闭时可以统一收口。"""
|
||||
task = asyncio.create_task(coroutine)
|
||||
task = get_task_registry().create(
|
||||
coroutine,
|
||||
owner="api.agent.web_execution",
|
||||
)
|
||||
_WEB_AGENT_BACKGROUND_TASKS.add(task)
|
||||
task.add_done_callback(_WEB_AGENT_BACKGROUND_TASKS.discard)
|
||||
return task
|
||||
|
||||
@@ -173,6 +173,12 @@ class DbOper:
|
||||
return run_sync_transaction(operation)
|
||||
return operation(self._db)
|
||||
|
||||
def _execute_sync_query(self, operation: Callable[[Session], T]) -> T:
|
||||
"""在当前同步会话查询,或委托组合根创建一次性兼容会话。"""
|
||||
if self._db is None or isinstance(self._db, AsyncSession):
|
||||
return run_sync_transaction(operation)
|
||||
return operation(self._db)
|
||||
|
||||
async def _execute_async_write(
|
||||
self,
|
||||
operation: Callable[[AsyncSession], Awaitable[T]],
|
||||
@@ -184,6 +190,15 @@ class DbOper:
|
||||
return await run_async_transaction(operation)
|
||||
return await operation(self._db)
|
||||
|
||||
async def _execute_async_query(
|
||||
self,
|
||||
operation: Callable[[AsyncSession], Awaitable[T]],
|
||||
) -> T:
|
||||
"""在当前异步会话查询,或委托组合根创建一次性兼容会话。"""
|
||||
if self._db is None or isinstance(self._db, Session):
|
||||
return await run_async_transaction(operation)
|
||||
return await operation(self._db)
|
||||
|
||||
def _stage_create(self, model: TModel) -> TModel:
|
||||
"""在显式同步事务中暂存新模型,不触发 Base 的兼容提交装饰器。"""
|
||||
def stage(session: Session) -> TModel:
|
||||
|
||||
+73
-48
@@ -5,7 +5,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import Mapped, Session, mapped_column
|
||||
|
||||
from app.db.base import Base, execute_dml, get_id_column
|
||||
from app.db.decorators import async_db_query, db_query
|
||||
from app.db.decorators import run_legacy_async_query, run_legacy_sync_query
|
||||
|
||||
|
||||
class Message(Base):
|
||||
@@ -49,21 +49,32 @@ class Message(Base):
|
||||
return self.to_dict()
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def list_by_page(cls, db: Session, page: int = 1, count: int = 30) -> List["Message"]:
|
||||
def list_by_page(
|
||||
cls,
|
||||
db: Session | None = None,
|
||||
page: int = 1,
|
||||
count: int = 30,
|
||||
) -> List["Message"]:
|
||||
"""
|
||||
分页获取消息记录。
|
||||
分页获取消息记录,兼容显式会话和旧插件无会话调用。
|
||||
"""
|
||||
return list(db.execute(
|
||||
select(cls)
|
||||
.order_by(cls.reg_time.desc(), cls.id.desc())
|
||||
.offset((page - 1) * count)
|
||||
.limit(count)
|
||||
).scalars().all())
|
||||
def query(session: Session) -> List["Message"]:
|
||||
"""在给定同步会话中执行消息分页查询。"""
|
||||
return list(session.execute(
|
||||
select(cls)
|
||||
.order_by(cls.reg_time.desc(), cls.id.desc())
|
||||
.offset((page - 1) * count)
|
||||
.limit(count)
|
||||
).scalars().all())
|
||||
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def exists_by_source(cls, db: Session, source: str) -> bool:
|
||||
def exists_by_source(
|
||||
cls,
|
||||
db: Session | str | None = None,
|
||||
source: str | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
判断指定来源标识的消息记录是否存在。
|
||||
|
||||
@@ -71,31 +82,42 @@ class Message(Base):
|
||||
:param source: 消息来源唯一标识
|
||||
:return: 是否存在匹配记录
|
||||
"""
|
||||
return db.execute(
|
||||
select(cls.id).where(cls.source == source).limit(1)
|
||||
).scalars().first() is not None
|
||||
if source is None and isinstance(db, str):
|
||||
source, db = db, None
|
||||
if source is None:
|
||||
raise TypeError("source is required")
|
||||
|
||||
def query(session: Session) -> bool:
|
||||
"""在给定同步会话中执行来源存在性查询。"""
|
||||
return session.execute(
|
||||
select(cls.id).where(cls.source == source).limit(1)
|
||||
).scalars().first() is not None
|
||||
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_list_by_page(
|
||||
cls, db: AsyncSession, page: int = 1, count: int = 30
|
||||
cls, db: AsyncSession | None = None, page: int = 1, count: int = 30
|
||||
) -> List["Message"]:
|
||||
"""
|
||||
异步分页获取消息记录。
|
||||
"""
|
||||
result = await db.execute(
|
||||
select(cls)
|
||||
.order_by(cls.reg_time.desc(), cls.id.desc())
|
||||
.offset((page - 1) * count)
|
||||
.limit(count)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
async def query(session: AsyncSession) -> List["Message"]:
|
||||
"""在给定异步会话中执行消息分页查询。"""
|
||||
result = await session.execute(
|
||||
select(cls)
|
||||
.order_by(cls.reg_time.desc(), cls.id.desc())
|
||||
.offset((page - 1) * count)
|
||||
.limit(count)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_list_sent_by_page(
|
||||
cls,
|
||||
db: AsyncSession,
|
||||
db: AsyncSession | None = None,
|
||||
page: int = 1,
|
||||
count: int = 30,
|
||||
all_clear_before: Optional[str] = None,
|
||||
@@ -105,32 +127,35 @@ class Message(Base):
|
||||
"""
|
||||
分页获取系统发送的通知消息。
|
||||
"""
|
||||
statement = select(cls).where(cls.action == 1)
|
||||
if all_clear_before:
|
||||
statement = statement.where(cls.reg_time > all_clear_before)
|
||||
if system_clear_before:
|
||||
statement = statement.where(
|
||||
or_(
|
||||
and_(cls.image.isnot(None), cls.image != ""),
|
||||
cls.reg_time > system_clear_before,
|
||||
async def query(session: AsyncSession) -> List["Message"]:
|
||||
"""在给定异步会话中执行通知消息分页查询。"""
|
||||
statement = select(cls).where(cls.action == 1)
|
||||
if all_clear_before:
|
||||
statement = statement.where(cls.reg_time > all_clear_before)
|
||||
if system_clear_before:
|
||||
statement = statement.where(
|
||||
or_(
|
||||
and_(cls.image.isnot(None), cls.image != ""),
|
||||
cls.reg_time > system_clear_before,
|
||||
)
|
||||
)
|
||||
)
|
||||
if media_clear_before:
|
||||
statement = statement.where(
|
||||
or_(
|
||||
cls.image.is_(None),
|
||||
cls.image == "",
|
||||
cls.reg_time > media_clear_before,
|
||||
if media_clear_before:
|
||||
statement = statement.where(
|
||||
or_(
|
||||
cls.image.is_(None),
|
||||
cls.image == "",
|
||||
cls.reg_time > media_clear_before,
|
||||
)
|
||||
)
|
||||
result = await session.execute(
|
||||
statement
|
||||
.order_by(cls.reg_time.desc(), cls.id.desc())
|
||||
.offset((page - 1) * count)
|
||||
.limit(count)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
result = await db.execute(
|
||||
statement
|
||||
.order_by(cls.reg_time.desc(), cls.id.desc())
|
||||
.offset((page - 1) * count)
|
||||
.limit(count)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
def delete_before(
|
||||
|
||||
+95
-29
@@ -6,7 +6,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import Mapped, Session, mapped_column
|
||||
|
||||
from app.db.base import Base, get_id_column
|
||||
from app.db.decorators import db_query, async_db_query
|
||||
from app.db.decorators import run_legacy_async_query, run_legacy_sync_query
|
||||
|
||||
|
||||
class Site(Base):
|
||||
@@ -58,48 +58,114 @@ class Site(Base):
|
||||
downloader: Mapped[Optional[str]] = mapped_column(String)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def get_by_domain(cls, db: Session, domain: str):
|
||||
return db.execute(select(cls).where(cls.domain == domain)).scalars().first()
|
||||
def get_by_domain(cls, db: Session | str | None = None, domain: str | None = None):
|
||||
"""按域名查询站点,兼容显式会话和旧插件无会话调用。"""
|
||||
if domain is None and isinstance(db, str):
|
||||
domain, db = db, None
|
||||
if domain is None:
|
||||
raise TypeError("domain is required")
|
||||
|
||||
def query(session: Session):
|
||||
"""在给定同步会话中执行域名查询。"""
|
||||
return session.execute(select(cls).where(cls.domain == domain)).scalars().first()
|
||||
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_get_by_domain(cls, db: AsyncSession, domain: str):
|
||||
result = await db.execute(select(cls).where(cls.domain == domain))
|
||||
return result.scalar_one_or_none()
|
||||
async def async_get_by_domain(
|
||||
cls,
|
||||
db: AsyncSession | str | None = None,
|
||||
domain: str | None = None,
|
||||
):
|
||||
"""异步按域名查询站点,兼容显式会话和旧插件无会话调用。"""
|
||||
if domain is None and isinstance(db, str):
|
||||
domain, db = db, None
|
||||
if domain is None:
|
||||
raise TypeError("domain is required")
|
||||
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行域名查询。"""
|
||||
result = await session.execute(select(cls).where(cls.domain == domain))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_get_by_name(cls, db: AsyncSession, name: str):
|
||||
result = await db.execute(select(cls).where(cls.name == name))
|
||||
return result.scalar_one_or_none()
|
||||
async def async_get_by_name(
|
||||
cls,
|
||||
db: AsyncSession | str | None = None,
|
||||
name: str | None = None,
|
||||
):
|
||||
"""异步按站点名称查询,兼容显式会话和旧插件无会话调用。"""
|
||||
if name is None and isinstance(db, str):
|
||||
name, db = db, None
|
||||
if name is None:
|
||||
raise TypeError("name is required")
|
||||
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行名称查询。"""
|
||||
result = await session.execute(select(cls).where(cls.name == name))
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def get_actives(cls, db: Session):
|
||||
return list(db.execute(select(cls).where(cls.is_active.is_(True))).scalars().all())
|
||||
def get_actives(cls, db: Session | None = None):
|
||||
"""查询启用站点,兼容显式会话和旧插件无会话调用。"""
|
||||
def query(session: Session):
|
||||
"""在给定同步会话中执行启用站点查询。"""
|
||||
return list(session.execute(select(cls).where(cls.is_active.is_(True))).scalars().all())
|
||||
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_get_actives(cls, db: AsyncSession):
|
||||
result = await db.execute(select(cls).where(cls.is_active.is_(True)))
|
||||
return list(result.scalars().all())
|
||||
async def async_get_actives(cls, db: AsyncSession | None = None):
|
||||
"""异步查询启用站点,兼容显式会话和旧插件无会话调用。"""
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行启用站点查询。"""
|
||||
result = await session.execute(select(cls).where(cls.is_active.is_(True)))
|
||||
return list(result.scalars().all())
|
||||
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def list_order_by_pri(cls, db: Session):
|
||||
return list(db.execute(select(cls).order_by(cls.pri)).scalars().all())
|
||||
def list_order_by_pri(cls, db: Session | None = None):
|
||||
"""按优先级升序查询站点,兼容显式会话和旧插件无会话调用。"""
|
||||
def query(session: Session):
|
||||
"""在给定同步会话中执行优先级查询。"""
|
||||
return list(session.execute(select(cls).order_by(cls.pri)).scalars().all())
|
||||
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_list_order_by_pri(cls, db: AsyncSession):
|
||||
result = await db.execute(select(cls).order_by(cls.pri))
|
||||
return list(result.scalars().all())
|
||||
async def async_list_order_by_pri(cls, db: AsyncSession | None = None):
|
||||
"""异步按优先级升序查询站点,兼容显式会话和旧插件无会话调用。"""
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行优先级查询。"""
|
||||
result = await session.execute(select(cls).order_by(cls.pri))
|
||||
return list(result.scalars().all())
|
||||
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def get_domains_by_ids(cls, db: Session, ids: list):
|
||||
return list(db.execute(select(cls.domain).where(cls.id.in_(ids))).scalars().all())
|
||||
def get_domains_by_ids(
|
||||
cls,
|
||||
db: Session | list[int] | None = None,
|
||||
ids: list[int] | None = None,
|
||||
):
|
||||
"""按 ID 查询域名,兼容显式会话和旧插件无会话调用。"""
|
||||
if ids is None and isinstance(db, list):
|
||||
ids, db = db, None
|
||||
if ids is None:
|
||||
raise TypeError("ids is required")
|
||||
if not ids:
|
||||
return []
|
||||
|
||||
def query(session: Session):
|
||||
"""在给定同步会话中执行域名投影查询。"""
|
||||
return list(session.execute(select(cls.domain).where(cls.id.in_(ids))).scalars().all())
|
||||
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
def reset(cls, db: Session):
|
||||
|
||||
+221
-140
@@ -6,7 +6,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import Mapped, Session, mapped_column
|
||||
|
||||
from app.db.base import get_id_column, Base
|
||||
from app.db.decorators import async_db_query, db_query
|
||||
from app.db.decorators import run_legacy_async_query, run_legacy_sync_query
|
||||
from app.db.models._constraints import media_identity_constraint
|
||||
from app.schemas.types import MUSIC_ENTITY_RECORDING, MediaSource
|
||||
|
||||
@@ -140,50 +140,66 @@ class Subscribe(Base):
|
||||
return condition
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def exists(
|
||||
cls, db: Session, media_source: MediaSource, media_id: str,
|
||||
cls, db: Session | MediaSource | None = None,
|
||||
media_source: MediaSource | str | None = None,
|
||||
media_id: str | None = None,
|
||||
season: Optional[int] = None,
|
||||
episode_group: Optional[str] = None,
|
||||
music_type: Optional[str] = None,
|
||||
):
|
||||
"""按媒体身份、季号与剧集组查询已有订阅。"""
|
||||
if db is not None and not isinstance(db, Session):
|
||||
media_source, media_id, db = db, media_source, None
|
||||
condition = cls._identity_condition(
|
||||
media_source, media_id, music_type
|
||||
)
|
||||
if condition is None:
|
||||
return None
|
||||
statement = select(cls).where(condition)
|
||||
if season is not None:
|
||||
statement = statement.where(cls.season == season)
|
||||
statement = statement.where(cls.episode_group == episode_group)
|
||||
return db.execute(statement).scalars().first()
|
||||
def query(session: Session):
|
||||
"""在给定会话中执行订阅身份查询。"""
|
||||
statement = select(cls).where(condition)
|
||||
if season is not None:
|
||||
statement = statement.where(cls.season == season)
|
||||
return session.execute(
|
||||
statement.where(cls.episode_group == episode_group)
|
||||
).scalars().first()
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_exists(
|
||||
cls, db: AsyncSession, media_source: MediaSource, media_id: str,
|
||||
cls, db: AsyncSession | MediaSource | None = None,
|
||||
media_source: MediaSource | str | None = None,
|
||||
media_id: str | None = None,
|
||||
season: Optional[int] = None,
|
||||
episode_group: Optional[str] = None,
|
||||
music_type: Optional[str] = None,
|
||||
):
|
||||
"""异步按媒体身份、季号与剧集组查询已有订阅。"""
|
||||
if db is not None and not isinstance(db, AsyncSession):
|
||||
media_source, media_id, db = db, media_source, None
|
||||
condition = cls._identity_condition(
|
||||
media_source, media_id, music_type
|
||||
)
|
||||
if condition is None:
|
||||
return None
|
||||
query = select(cls).filter(condition)
|
||||
if season is not None:
|
||||
query = query.filter(cls.season == season)
|
||||
query = query.filter(cls.episode_group == episode_group)
|
||||
result = await db.execute(query)
|
||||
return result.scalars().first()
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行订阅身份查询。"""
|
||||
statement = select(cls).where(condition)
|
||||
if season is not None:
|
||||
statement = statement.where(cls.season == season)
|
||||
result = await session.execute(
|
||||
statement.where(cls.episode_group == episode_group)
|
||||
)
|
||||
return result.scalars().first()
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def exists_by_username(
|
||||
cls, db: Session, username: str, media_source: MediaSource, media_id: str,
|
||||
cls, db: Session | str | None = None,
|
||||
username: str | MediaSource | None = None,
|
||||
media_source: MediaSource | str | None = None,
|
||||
media_id: str | None = None,
|
||||
season: Optional[int] = None,
|
||||
episode_group: Optional[str] = None,
|
||||
music_type: Optional[str] = None,
|
||||
@@ -191,6 +207,8 @@ class Subscribe(Base):
|
||||
"""
|
||||
按订阅 owner、媒体身份、季号与剧集组查询订阅行。
|
||||
"""
|
||||
if db is not None and not isinstance(db, Session):
|
||||
username, media_source, media_id, db = db, username, media_source, None
|
||||
if not username:
|
||||
return None
|
||||
condition = cls._identity_condition(
|
||||
@@ -198,23 +216,30 @@ class Subscribe(Base):
|
||||
)
|
||||
if condition is None:
|
||||
return None
|
||||
statement = select(cls).where(cls.username == username, condition)
|
||||
if season is not None:
|
||||
statement = statement.where(cls.season == season)
|
||||
statement = statement.where(cls.episode_group == episode_group)
|
||||
return db.execute(statement).scalars().first()
|
||||
def query(session: Session):
|
||||
"""在给定会话中执行订阅 owner 查询。"""
|
||||
statement = select(cls).where(cls.username == username, condition)
|
||||
if season is not None:
|
||||
statement = statement.where(cls.season == season)
|
||||
return session.execute(
|
||||
statement.where(cls.episode_group == episode_group)
|
||||
).scalars().first()
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_exists_by_username(
|
||||
cls, db: AsyncSession, username: str, media_source: MediaSource,
|
||||
media_id: str, season: Optional[int] = None,
|
||||
cls, db: AsyncSession | str | None = None,
|
||||
username: str | MediaSource | None = None,
|
||||
media_source: MediaSource | str | None = None,
|
||||
media_id: str | None = None, season: Optional[int] = None,
|
||||
episode_group: Optional[str] = None,
|
||||
music_type: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
异步按订阅 owner、媒体身份、季号与剧集组查询订阅行。
|
||||
"""
|
||||
if db is not None and not isinstance(db, AsyncSession):
|
||||
username, media_source, media_id, db = db, username, media_source, None
|
||||
if not username:
|
||||
return None
|
||||
condition = cls._identity_condition(
|
||||
@@ -222,80 +247,106 @@ class Subscribe(Base):
|
||||
)
|
||||
if condition is None:
|
||||
return None
|
||||
query = select(cls).filter(cls.username == username, condition)
|
||||
if season is not None:
|
||||
query = query.filter(cls.season == season)
|
||||
query = query.filter(cls.episode_group == episode_group)
|
||||
result = await db.execute(query)
|
||||
return result.scalars().first()
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def get_by_state(cls, db: Session, state: str):
|
||||
# 如果 state 为空或 None,返回所有订阅
|
||||
statement = select(cls)
|
||||
if state:
|
||||
# 如果传入的状态不为空,拆分成多个状态
|
||||
statement = statement.where(cls.state.in_(state.split(',')))
|
||||
return list(db.execute(statement).scalars().all())
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_get_by_state(cls, db: AsyncSession, state: str):
|
||||
# 如果 state 为空或 None,返回所有订阅
|
||||
if not state:
|
||||
result = await db.execute(select(cls))
|
||||
else:
|
||||
# 如果传入的状态不为空,拆分成多个状态
|
||||
result = await db.execute(
|
||||
select(cls).filter(cls.state.in_(state.split(',')))
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行订阅 owner 查询。"""
|
||||
statement = select(cls).where(cls.username == username, condition)
|
||||
if season is not None:
|
||||
statement = statement.where(cls.season == season)
|
||||
result = await session.execute(
|
||||
statement.where(cls.episode_group == episode_group)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
return result.scalars().first()
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def get_by_title(cls, db: Session, title: str, season: Optional[int] = None):
|
||||
statement = select(cls).where(cls.name == title)
|
||||
if season is not None:
|
||||
statement = statement.where(cls.season == season)
|
||||
return db.execute(statement).scalars().first()
|
||||
def get_by_state(cls, db: Session | str | None = None, state: str | None = None):
|
||||
"""按状态列表查询订阅,兼容显式会话和旧插件无会话调用。"""
|
||||
if not isinstance(db, Session):
|
||||
state, db = db if state is None else state, None
|
||||
def query(session: Session):
|
||||
"""在给定会话中执行状态查询。"""
|
||||
statement = select(cls)
|
||||
if state:
|
||||
statement = statement.where(cls.state.in_(state.split(',')))
|
||||
return list(session.execute(statement).scalars().all())
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_get_by_title(cls, db: AsyncSession, title: str, season: Optional[int] = None):
|
||||
if season is not None:
|
||||
result = await db.execute(
|
||||
select(cls).filter(cls.name == title, cls.season == season)
|
||||
)
|
||||
else:
|
||||
result = await db.execute(
|
||||
select(cls).filter(cls.name == title)
|
||||
)
|
||||
return result.scalars().first()
|
||||
async def async_get_by_state(
|
||||
cls, db: AsyncSession | str | None = None, state: str | None = None
|
||||
):
|
||||
"""异步按状态列表查询订阅,兼容显式会话和旧插件无会话调用。"""
|
||||
if not isinstance(db, AsyncSession):
|
||||
state, db = db if state is None else state, None
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行状态查询。"""
|
||||
statement = select(cls)
|
||||
if state:
|
||||
statement = statement.where(cls.state.in_(state.split(',')))
|
||||
result = await session.execute(statement)
|
||||
return list(result.scalars().all())
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_list_by_title(cls, db: AsyncSession, title: str, season: Optional[int] = None):
|
||||
"""
|
||||
异步按标题查询候选订阅列表。
|
||||
"""
|
||||
if season is not None:
|
||||
result = await db.execute(
|
||||
select(cls).filter(cls.name == title, cls.season == season)
|
||||
)
|
||||
else:
|
||||
result = await db.execute(
|
||||
select(cls).filter(cls.name == title)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
def get_by_title(
|
||||
cls, db: Session | str | None = None, title: str | None = None,
|
||||
season: Optional[int] = None,
|
||||
):
|
||||
"""按标题查询订阅,兼容显式会话和旧插件无会话调用。"""
|
||||
if not isinstance(db, Session):
|
||||
title, db = db if title is None else title, None
|
||||
def query(session: Session):
|
||||
"""在给定会话中执行标题查询。"""
|
||||
statement = select(cls).where(cls.name == title)
|
||||
if season is not None:
|
||||
statement = statement.where(cls.season == season)
|
||||
return session.execute(statement).scalars().first()
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
async def async_get_by_title(
|
||||
cls, db: AsyncSession | str | None = None, title: str | None = None,
|
||||
season: Optional[int] = None,
|
||||
):
|
||||
"""异步按标题查询订阅,兼容显式会话和旧插件无会话调用。"""
|
||||
if not isinstance(db, AsyncSession):
|
||||
title, db = db if title is None else title, None
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行标题查询。"""
|
||||
statement = select(cls).where(cls.name == title)
|
||||
if season is not None:
|
||||
statement = statement.where(cls.season == season)
|
||||
result = await session.execute(statement)
|
||||
return result.scalars().first()
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
async def async_list_by_title(
|
||||
cls, db: AsyncSession | str | None = None, title: str | None = None,
|
||||
season: Optional[int] = None,
|
||||
):
|
||||
"""异步按标题查询候选订阅列表,兼容显式会话和旧插件无会话调用。"""
|
||||
if not isinstance(db, AsyncSession):
|
||||
title, db = db if title is None else title, None
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行标题列表查询。"""
|
||||
statement = select(cls).where(cls.name == title)
|
||||
if season is not None:
|
||||
statement = statement.where(cls.season == season)
|
||||
result = await session.execute(statement)
|
||||
return list(result.scalars().all())
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def list_by_media_identity(
|
||||
cls, db: Session, media_source: MediaSource, media_id: str,
|
||||
cls, db: Session | MediaSource | None = None,
|
||||
media_source: MediaSource | str | None = None,
|
||||
media_id: str | None = None,
|
||||
music_type: Optional[str] = None,
|
||||
):
|
||||
"""同步按统一媒体身份查询候选订阅列表。"""
|
||||
if db is not None and not isinstance(db, Session):
|
||||
media_source, media_id, db = db, media_source, None
|
||||
condition = cls._identity_condition(
|
||||
media_source=media_source,
|
||||
media_id=media_id,
|
||||
@@ -303,15 +354,21 @@ class Subscribe(Base):
|
||||
)
|
||||
if condition is None:
|
||||
return []
|
||||
return list(db.execute(select(cls).where(condition)).scalars().all())
|
||||
def query(session: Session):
|
||||
"""在给定会话中执行媒体身份列表查询。"""
|
||||
return list(session.execute(select(cls).where(condition)).scalars().all())
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_list_by_media_identity(
|
||||
cls, db: AsyncSession, media_source: MediaSource, media_id: str,
|
||||
cls, db: AsyncSession | MediaSource | None = None,
|
||||
media_source: MediaSource | str | None = None,
|
||||
media_id: str | None = None,
|
||||
music_type: Optional[str] = None,
|
||||
):
|
||||
"""异步按统一媒体身份查询候选订阅列表。"""
|
||||
if db is not None and not isinstance(db, AsyncSession):
|
||||
media_source, media_id, db = db, media_source, None
|
||||
condition = cls._identity_condition(
|
||||
media_source=media_source,
|
||||
media_id=media_id,
|
||||
@@ -319,19 +376,26 @@ class Subscribe(Base):
|
||||
)
|
||||
if condition is None:
|
||||
return []
|
||||
result = await db.execute(select(cls).filter(condition))
|
||||
return list(result.scalars().all())
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行媒体身份列表查询。"""
|
||||
result = await session.execute(select(cls).where(condition))
|
||||
return list(result.scalars().all())
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def get_by(
|
||||
cls, db: Session, type: str, media_source: MediaSource, media_id: str,
|
||||
cls, db: Session | str | None = None,
|
||||
type: str | MediaSource | None = None,
|
||||
media_source: MediaSource | str | None = None,
|
||||
media_id: str | None = None,
|
||||
season: Optional[str] = None,
|
||||
music_type: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
根据条件查询订阅
|
||||
"""
|
||||
if db is not None and not isinstance(db, Session):
|
||||
type, media_source, media_id, db = db, type, media_source, None
|
||||
condition = cls._identity_condition(
|
||||
media_source, media_id, music_type
|
||||
)
|
||||
@@ -340,18 +404,25 @@ class Subscribe(Base):
|
||||
statement = select(cls).where(condition, cls.type == type)
|
||||
if season is not None:
|
||||
statement = statement.where(cls.season == season)
|
||||
return db.execute(statement).scalars().first()
|
||||
def query(session: Session):
|
||||
"""在给定会话中执行类型媒体查询。"""
|
||||
return session.execute(statement).scalars().first()
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_get_by(
|
||||
cls, db: AsyncSession, type: str, media_source: MediaSource, media_id: str,
|
||||
cls, db: AsyncSession | str | None = None,
|
||||
type: str | MediaSource | None = None,
|
||||
media_source: MediaSource | str | None = None,
|
||||
media_id: str | None = None,
|
||||
season: Optional[str] = None,
|
||||
music_type: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
根据条件查询订阅
|
||||
"""
|
||||
if db is not None and not isinstance(db, AsyncSession):
|
||||
type, media_source, media_id, db = db, type, media_source, None
|
||||
condition = cls._identity_condition(
|
||||
media_source, media_id, music_type
|
||||
)
|
||||
@@ -360,62 +431,72 @@ class Subscribe(Base):
|
||||
query = select(cls).filter(condition, cls.type == type)
|
||||
if season is not None:
|
||||
query = query.filter(cls.season == season)
|
||||
result = await db.execute(query)
|
||||
return result.scalars().first()
|
||||
async def execute_query(session: AsyncSession):
|
||||
"""在给定异步会话中执行类型媒体查询。"""
|
||||
result = await session.execute(query)
|
||||
return result.scalars().first()
|
||||
return await execute_query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(execute_query)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def list_by_username(cls, db: Session, username: str, state: Optional[str] = None, mtype: Optional[str] = None):
|
||||
statement = select(cls).where(cls.username == username)
|
||||
if state:
|
||||
statement = statement.where(cls.state == state)
|
||||
if mtype:
|
||||
statement = statement.where(cls.type == mtype)
|
||||
return list(db.execute(statement).scalars().all())
|
||||
def list_by_username(cls, db: Session | str | None = None, username: str | None = None,
|
||||
state: Optional[str] = None, mtype: Optional[str] = None):
|
||||
"""按用户筛选订阅,兼容显式会话和旧插件无会话调用。"""
|
||||
if not isinstance(db, Session):
|
||||
username, db = db if username is None else username, None
|
||||
def query(session: Session):
|
||||
"""在给定会话中执行用户筛选查询。"""
|
||||
statement = select(cls).where(cls.username == username)
|
||||
if state:
|
||||
statement = statement.where(cls.state == state)
|
||||
if mtype:
|
||||
statement = statement.where(cls.type == mtype)
|
||||
return list(session.execute(statement).scalars().all())
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_list_by_username(cls, db: AsyncSession, username: str, state: Optional[str] = None,
|
||||
async def async_list_by_username(cls, db: AsyncSession | str | None = None,
|
||||
username: str | None = None, state: Optional[str] = None,
|
||||
mtype: Optional[str] = None):
|
||||
if mtype:
|
||||
"""异步按用户筛选订阅,兼容显式会话和旧插件无会话调用。"""
|
||||
if not isinstance(db, AsyncSession):
|
||||
username, db = db if username is None else username, None
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行用户筛选查询。"""
|
||||
statement = select(cls).where(cls.username == username)
|
||||
if state:
|
||||
result = await db.execute(
|
||||
select(cls).filter(cls.state == state, cls.username == username, cls.type == mtype)
|
||||
)
|
||||
else:
|
||||
result = await db.execute(
|
||||
select(cls).filter(cls.username == username, cls.type == mtype)
|
||||
)
|
||||
else:
|
||||
if state:
|
||||
result = await db.execute(
|
||||
select(cls).filter(cls.state == state, cls.username == username)
|
||||
)
|
||||
else:
|
||||
result = await db.execute(
|
||||
select(cls).filter(cls.username == username)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
statement = statement.where(cls.state == state)
|
||||
if mtype:
|
||||
statement = statement.where(cls.type == mtype)
|
||||
result = await session.execute(statement)
|
||||
return list(result.scalars().all())
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
@db_query
|
||||
def list_by_type(cls, db: Session, mtype: str, days: int):
|
||||
return list(db.execute(
|
||||
select(cls).where(
|
||||
def list_by_type(cls, db: Session | str | None = None, mtype: str | None = None, days: int = 7):
|
||||
"""按类型查询最近时间窗内的订阅,兼容显式会话和旧插件无会话调用。"""
|
||||
if not isinstance(db, Session):
|
||||
mtype, db = db if mtype is None else mtype, None
|
||||
def query(session: Session):
|
||||
"""在给定会话中执行时间窗订阅查询。"""
|
||||
return list(session.execute(select(cls).where(
|
||||
cls.type == mtype,
|
||||
cls.date >= time.strftime("%Y-%m-%d %H:%M:%S",
|
||||
time.localtime(time.time() - 86400 * int(days)))
|
||||
)
|
||||
).scalars().all())
|
||||
)).scalars().all())
|
||||
return query(db) if isinstance(db, Session) else run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_list_by_type(cls, db: AsyncSession, mtype: str, days: int):
|
||||
result = await db.execute(
|
||||
select(cls).filter(
|
||||
async def async_list_by_type(cls, db: AsyncSession | str | None = None,
|
||||
mtype: str | None = None, days: int = 7):
|
||||
"""异步按类型查询最近时间窗内的订阅,兼容显式会话和旧插件无会话调用。"""
|
||||
if not isinstance(db, AsyncSession):
|
||||
mtype, db = db if mtype is None else mtype, None
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行时间窗订阅查询。"""
|
||||
result = await session.execute(select(cls).where(
|
||||
cls.type == mtype,
|
||||
cls.date >= time.strftime("%Y-%m-%d %H:%M:%S",
|
||||
time.localtime(time.time() - 86400 * int(days)))
|
||||
)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
))
|
||||
return list(result.scalars().all())
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
+38
-13
@@ -4,7 +4,10 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import Mapped, Session, mapped_column
|
||||
|
||||
from app.db.base import Base, get_id_column
|
||||
from app.db.decorators import async_db_query, run_legacy_sync_query
|
||||
from app.db.decorators import (
|
||||
run_legacy_async_query,
|
||||
run_legacy_sync_query,
|
||||
)
|
||||
|
||||
|
||||
class User(Base):
|
||||
@@ -55,12 +58,23 @@ class User(Base):
|
||||
return run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_get_by_name(cls, db: AsyncSession, name: str):
|
||||
result = await db.execute(
|
||||
select(cls).filter(cls.name == name)
|
||||
)
|
||||
return result.scalars().first()
|
||||
async def async_get_by_name(
|
||||
cls,
|
||||
db: AsyncSession | str | None = None,
|
||||
name: str | None = None,
|
||||
):
|
||||
"""异步按用户名查询,兼容显式会话和旧插件无会话调用。"""
|
||||
if name is None and isinstance(db, str):
|
||||
name, db = db, None
|
||||
if name is None:
|
||||
raise TypeError("name is required")
|
||||
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行用户名查询。"""
|
||||
result = await session.execute(select(cls).filter(cls.name == name))
|
||||
return result.scalars().first()
|
||||
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
@classmethod
|
||||
def get_by_id(cls, db: Session | int | None = None, user_id: int | None = None):
|
||||
@@ -79,12 +93,23 @@ class User(Base):
|
||||
return run_legacy_sync_query(query)
|
||||
|
||||
@classmethod
|
||||
@async_db_query
|
||||
async def async_get_by_id(cls, db: AsyncSession, user_id: int):
|
||||
result = await db.execute(
|
||||
select(cls).filter(cls.id == user_id)
|
||||
)
|
||||
return result.scalars().first()
|
||||
async def async_get_by_id(
|
||||
cls,
|
||||
db: AsyncSession | int | None = None,
|
||||
user_id: int | None = None,
|
||||
):
|
||||
"""异步按用户 ID 查询,兼容显式会话和旧插件无会话调用。"""
|
||||
if user_id is None and isinstance(db, int):
|
||||
user_id, db = db, None
|
||||
if user_id is None:
|
||||
raise TypeError("user_id is required")
|
||||
|
||||
async def query(session: AsyncSession):
|
||||
"""在给定异步会话中执行用户 ID 查询。"""
|
||||
result = await session.execute(select(cls).filter(cls.id == user_id))
|
||||
return result.scalars().first()
|
||||
|
||||
return await query(db) if isinstance(db, AsyncSession) else await run_legacy_async_query(query)
|
||||
|
||||
def delete_by_name(self, db: Session, name: str):
|
||||
user = self.get_by_name(db, name)
|
||||
|
||||
+50
-11
@@ -3,6 +3,7 @@ from typing import Optional, Union
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import Session
|
||||
from sqlalchemy import and_, or_, select
|
||||
|
||||
from app.db.base import DbOper
|
||||
from app.db.models.message import Message
|
||||
@@ -105,7 +106,14 @@ class MessageOper(DbOper):
|
||||
"""
|
||||
分页获取消息记录。
|
||||
"""
|
||||
return Message.list_by_page(self._db, page, count)
|
||||
return self._execute_sync_query(
|
||||
lambda session: list(session.execute(
|
||||
select(Message)
|
||||
.order_by(Message.reg_time.desc(), Message.id.desc())
|
||||
.offset((page - 1) * count)
|
||||
.limit(count)
|
||||
).scalars().all())
|
||||
)
|
||||
|
||||
def exists_by_source(self, source: str) -> bool:
|
||||
"""
|
||||
@@ -114,7 +122,11 @@ class MessageOper(DbOper):
|
||||
:param source: 消息来源唯一标识
|
||||
:return: 是否存在匹配记录
|
||||
"""
|
||||
return Message.exists_by_source(self._db, source)
|
||||
return self._execute_sync_query(
|
||||
lambda session: session.execute(
|
||||
select(Message.id).where(Message.source == source).limit(1)
|
||||
).scalars().first() is not None
|
||||
)
|
||||
|
||||
async def async_list_by_page(
|
||||
self, page: int = 1, count: int = 30
|
||||
@@ -122,7 +134,17 @@ class MessageOper(DbOper):
|
||||
"""
|
||||
分页获取消息记录。
|
||||
"""
|
||||
return await Message.async_list_by_page(self._db, page, count)
|
||||
async def query(session: AsyncSession) -> list[Message]:
|
||||
"""在调用方异步会话中执行消息分页查询。"""
|
||||
result = await session.execute(
|
||||
select(Message)
|
||||
.order_by(Message.reg_time.desc(), Message.id.desc())
|
||||
.offset((page - 1) * count)
|
||||
.limit(count)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
return await self._execute_async_query(query)
|
||||
|
||||
async def async_list_sent_by_page(
|
||||
self,
|
||||
@@ -135,11 +157,28 @@ class MessageOper(DbOper):
|
||||
"""
|
||||
分页获取系统发送的通知消息。
|
||||
"""
|
||||
return await Message.async_list_sent_by_page(
|
||||
self._db,
|
||||
page,
|
||||
count,
|
||||
all_clear_before=all_clear_before,
|
||||
system_clear_before=system_clear_before,
|
||||
media_clear_before=media_clear_before,
|
||||
)
|
||||
async def query(session: AsyncSession) -> list[Message]:
|
||||
"""在调用方异步会话中执行通知消息分页查询。"""
|
||||
statement = select(Message).where(Message.action == 1)
|
||||
if all_clear_before:
|
||||
statement = statement.where(Message.reg_time > all_clear_before)
|
||||
if system_clear_before:
|
||||
statement = statement.where(or_(
|
||||
and_(Message.image.isnot(None), Message.image != ""),
|
||||
Message.reg_time > system_clear_before,
|
||||
))
|
||||
if media_clear_before:
|
||||
statement = statement.where(or_(
|
||||
Message.image.is_(None),
|
||||
Message.image == "",
|
||||
Message.reg_time > media_clear_before,
|
||||
))
|
||||
result = await session.execute(
|
||||
statement
|
||||
.order_by(Message.reg_time.desc(), Message.id.desc())
|
||||
.offset((page - 1) * count)
|
||||
.limit(count)
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
|
||||
return await self._execute_async_query(query)
|
||||
|
||||
+86
-18
@@ -1,7 +1,8 @@
|
||||
from datetime import datetime
|
||||
from typing import Any, List, Mapping, Tuple, Optional
|
||||
|
||||
from sqlalchemy import delete as sqlalchemy_delete
|
||||
from sqlalchemy import delete as sqlalchemy_delete, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.db.base import DbOper
|
||||
@@ -11,6 +12,18 @@ from app.db.models.sitestatistic import SiteStatistic
|
||||
from app.db.models.siteuserdata import SiteUserData
|
||||
|
||||
|
||||
async def _async_first(session: AsyncSession, statement: Any) -> Optional[Site]:
|
||||
"""执行异步站点查询并返回首条记录。"""
|
||||
result = await session.execute(statement)
|
||||
return result.scalars().first()
|
||||
|
||||
|
||||
async def _async_all(session: AsyncSession, statement: Any) -> list[Site]:
|
||||
"""执行异步站点查询并返回稳定列表。"""
|
||||
result = await session.execute(statement)
|
||||
return list(result.scalars().all())
|
||||
|
||||
|
||||
class SiteOper(DbOper):
|
||||
"""
|
||||
站点管理
|
||||
@@ -21,7 +34,7 @@ class SiteOper(DbOper):
|
||||
新增站点
|
||||
"""
|
||||
site = Site(**kwargs)
|
||||
if not site.get_by_domain(self._db, kwargs.get("domain")):
|
||||
if not self.get_by_domain(kwargs.get("domain")):
|
||||
self._stage_create(site)
|
||||
return True, "新增站点成功"
|
||||
return False, "站点已存在"
|
||||
@@ -30,13 +43,22 @@ class SiteOper(DbOper):
|
||||
"""
|
||||
查询单个站点
|
||||
"""
|
||||
return Site.get(self._db, sid)
|
||||
return self._execute_sync_query(
|
||||
lambda session: session.execute(
|
||||
select(Site).where(Site.id == sid)
|
||||
).scalars().first()
|
||||
)
|
||||
|
||||
async def async_get(self, sid: int) -> Optional[Site]:
|
||||
"""
|
||||
异步查询单个站点
|
||||
"""
|
||||
return await Site.async_get(self._db, sid)
|
||||
return await self._execute_async_query(
|
||||
lambda session: _async_first(
|
||||
session,
|
||||
select(Site).where(Site.id == sid),
|
||||
)
|
||||
)
|
||||
|
||||
async def get_by_id(self, site_id: int) -> Optional[Site]:
|
||||
"""读取站点写用例需要的目标站点。"""
|
||||
@@ -80,35 +102,59 @@ class SiteOper(DbOper):
|
||||
"""
|
||||
获取站点列表
|
||||
"""
|
||||
return Site.list(self._db)
|
||||
return self._execute_sync_query(
|
||||
lambda session: list(session.execute(select(Site)).scalars().all())
|
||||
)
|
||||
|
||||
async def async_list(self) -> List[Site]:
|
||||
"""
|
||||
异步获取站点列表
|
||||
"""
|
||||
return await Site.async_list(self._db)
|
||||
return await self._execute_async_query(
|
||||
lambda session: _async_all(session, select(Site))
|
||||
)
|
||||
|
||||
async def async_list_order_by_pri(self) -> List[Site]:
|
||||
"""异步按优先级获取站点,供站点查询应用服务使用。"""
|
||||
return await Site.async_list_order_by_pri(self._db)
|
||||
return await self._execute_async_query(
|
||||
lambda session: _async_all(
|
||||
session,
|
||||
select(Site).order_by(Site.pri),
|
||||
)
|
||||
)
|
||||
|
||||
def list_order_by_pri(self) -> List[Site]:
|
||||
"""
|
||||
获取站点列表
|
||||
"""
|
||||
return Site.list_order_by_pri(self._db)
|
||||
return self._execute_sync_query(
|
||||
lambda session: list(
|
||||
session.execute(select(Site).order_by(Site.pri)).scalars().all()
|
||||
)
|
||||
)
|
||||
|
||||
def list_active(self) -> List[Site]:
|
||||
"""
|
||||
按状态获取站点列表
|
||||
"""
|
||||
return Site.get_actives(self._db)
|
||||
return self._execute_sync_query(
|
||||
lambda session: list(
|
||||
session.execute(
|
||||
select(Site).where(Site.is_active.is_(True))
|
||||
).scalars().all()
|
||||
)
|
||||
)
|
||||
|
||||
async def async_list_active(self) -> List[Site]:
|
||||
"""
|
||||
异步按状态获取站点列表
|
||||
"""
|
||||
return await Site.async_get_actives(self._db)
|
||||
return await self._execute_async_query(
|
||||
lambda session: _async_all(
|
||||
session,
|
||||
select(Site).where(Site.is_active.is_(True)),
|
||||
)
|
||||
)
|
||||
|
||||
def delete(self, sid: int):
|
||||
"""
|
||||
@@ -128,7 +174,7 @@ class SiteOper(DbOper):
|
||||
"""
|
||||
更新站点
|
||||
"""
|
||||
site = Site.get(self._db, sid)
|
||||
site = self.get(sid)
|
||||
if not site:
|
||||
return None
|
||||
self._stage_update(site, payload)
|
||||
@@ -147,37 +193,59 @@ class SiteOper(DbOper):
|
||||
"""
|
||||
按域名获取站点
|
||||
"""
|
||||
return Site.get_by_domain(self._db, domain)
|
||||
return self._execute_sync_query(
|
||||
lambda session: session.execute(
|
||||
select(Site).where(Site.domain == domain)
|
||||
).scalars().first()
|
||||
)
|
||||
|
||||
async def async_get_by_domain(self, domain: str) -> Optional[Site]:
|
||||
"""
|
||||
异步按域名获取站点
|
||||
"""
|
||||
return await Site.async_get_by_domain(self._db, domain)
|
||||
return await self._execute_async_query(
|
||||
lambda session: _async_first(
|
||||
session,
|
||||
select(Site).where(Site.domain == domain),
|
||||
)
|
||||
)
|
||||
|
||||
async def async_get_by_name(self, name: str) -> Optional[Site]:
|
||||
"""
|
||||
异步按名称获取站点
|
||||
"""
|
||||
return await Site.async_get_by_name(self._db, name)
|
||||
return await self._execute_async_query(
|
||||
lambda session: _async_first(
|
||||
session,
|
||||
select(Site).where(Site.name == name),
|
||||
)
|
||||
)
|
||||
|
||||
def get_domains_by_ids(self, ids: List[int]) -> List[Optional[str]]:
|
||||
"""
|
||||
按ID获取站点域名
|
||||
"""
|
||||
return Site.get_domains_by_ids(self._db, ids)
|
||||
if not ids:
|
||||
return []
|
||||
return self._execute_sync_query(
|
||||
lambda session: list(
|
||||
session.execute(
|
||||
select(Site.domain).where(Site.id.in_(ids))
|
||||
).scalars().all()
|
||||
)
|
||||
)
|
||||
|
||||
def exists(self, domain: str) -> bool:
|
||||
"""
|
||||
判断站点是否存在
|
||||
"""
|
||||
return Site.get_by_domain(self._db, domain) is not None
|
||||
return self.get_by_domain(domain) is not None
|
||||
|
||||
def update_cookie(self, domain: str, cookies: str) -> Tuple[bool, str]:
|
||||
"""
|
||||
更新站点Cookie
|
||||
"""
|
||||
site = Site.get_by_domain(self._db, domain)
|
||||
site = self.get_by_domain(domain)
|
||||
if not site:
|
||||
return False, "站点不存在"
|
||||
self._stage_update(site, {
|
||||
@@ -189,7 +257,7 @@ class SiteOper(DbOper):
|
||||
"""
|
||||
更新站点rss
|
||||
"""
|
||||
site = Site.get_by_domain(self._db, domain)
|
||||
site = self.get_by_domain(domain)
|
||||
if not site:
|
||||
return False, "站点不存在"
|
||||
self._stage_update(site, {
|
||||
|
||||
+127
-47
@@ -297,13 +297,24 @@ class SubscribeOper(DbOper):
|
||||
"""
|
||||
获取订阅
|
||||
"""
|
||||
return Subscribe.get(self._db, rid=sid)
|
||||
return self._execute_sync_query(
|
||||
lambda session: session.execute(
|
||||
select(Subscribe).where(Subscribe.id == sid)
|
||||
).scalars().first()
|
||||
)
|
||||
|
||||
async def async_get(self, sid: int) -> Optional[Subscribe]:
|
||||
"""
|
||||
获取订阅
|
||||
"""
|
||||
return await Subscribe.async_get(self._db, rid=sid)
|
||||
if self._db is not None and not isinstance(self._db, (Session, AsyncSession)):
|
||||
# 保留旧测试替身与插件注入对象对 Model ABI 的兼容入口。
|
||||
return await Subscribe.async_get(self._db, rid=sid)
|
||||
async def query(session: AsyncSession) -> Optional[Subscribe]:
|
||||
"""在调用方异步会话中执行订阅主键查询。"""
|
||||
result = await session.execute(select(Subscribe).where(Subscribe.id == sid))
|
||||
return result.scalars().first()
|
||||
return await self._execute_async_query(query)
|
||||
|
||||
async def async_list_by_media_identity(
|
||||
self,
|
||||
@@ -312,12 +323,18 @@ class SubscribeOper(DbOper):
|
||||
music_type: Optional[str] = None,
|
||||
) -> List[Subscribe]:
|
||||
"""异步按规范媒体身份读取订阅。"""
|
||||
return await Subscribe.async_list_by_media_identity(
|
||||
self._db,
|
||||
media_source=media_source,
|
||||
media_id=media_id,
|
||||
music_type=music_type,
|
||||
)
|
||||
async def query(session: AsyncSession) -> List[Subscribe]:
|
||||
"""在调用方异步会话中执行媒体身份列表查询。"""
|
||||
condition = Subscribe._identity_condition( # pylint: disable=protected-access
|
||||
media_source, media_id, music_type
|
||||
)
|
||||
if condition is None:
|
||||
return []
|
||||
result = await session.execute(select(Subscribe).where(condition))
|
||||
return list(result.scalars().all())
|
||||
if isinstance(self._db, AsyncSession):
|
||||
return await query(self._db)
|
||||
return await self._execute_async_query(query)
|
||||
|
||||
def list_by_media_identity(
|
||||
self,
|
||||
@@ -326,12 +343,15 @@ class SubscribeOper(DbOper):
|
||||
music_type: Optional[str] = None,
|
||||
) -> List[Subscribe]:
|
||||
"""同步按规范媒体身份读取订阅。"""
|
||||
return Subscribe.list_by_media_identity(
|
||||
self._db,
|
||||
media_source=media_source,
|
||||
media_id=media_id,
|
||||
music_type=music_type,
|
||||
)
|
||||
def query(session: Session) -> List[Subscribe]:
|
||||
"""在调用方同步会话中执行媒体身份列表查询。"""
|
||||
condition = Subscribe._identity_condition( # pylint: disable=protected-access
|
||||
media_source, media_id, music_type
|
||||
)
|
||||
if condition is None:
|
||||
return []
|
||||
return list(session.execute(select(Subscribe).where(condition)).scalars().all())
|
||||
return self._execute_sync_query(query)
|
||||
|
||||
async def get_candidate(
|
||||
self,
|
||||
@@ -360,11 +380,8 @@ class SubscribeOper(DbOper):
|
||||
music_type: Optional[str],
|
||||
) -> List[SubscribeDeletionCandidate]:
|
||||
"""按媒体身份读取去重后的订阅删除快照。"""
|
||||
subscribes = await Subscribe.async_list_by_media_identity(
|
||||
self._db,
|
||||
media_source=media_source,
|
||||
media_id=media_id,
|
||||
music_type=music_type,
|
||||
subscribes = await self.async_list_by_media_identity(
|
||||
media_source, media_id, music_type
|
||||
)
|
||||
candidates = []
|
||||
seen_ids = set()
|
||||
@@ -395,11 +412,7 @@ class SubscribeOper(DbOper):
|
||||
|
||||
async def list_search_ids(self, username: str, state: str) -> List[int]:
|
||||
"""返回用户指定状态的订阅编号,不向应用用例暴露 ORM 列表。"""
|
||||
subscribes = await Subscribe.async_list_by_username(
|
||||
self._db,
|
||||
username,
|
||||
state=state,
|
||||
)
|
||||
subscribes = await self.async_list_by_username(username, state=state)
|
||||
return [subscribe.id for subscribe in subscribes if subscribe.id]
|
||||
|
||||
def get_by(
|
||||
@@ -410,9 +423,18 @@ class SubscribeOper(DbOper):
|
||||
"""
|
||||
根据条件查询订阅
|
||||
"""
|
||||
return Subscribe.get_by(
|
||||
self._db, type, media_source, media_id, season, music_type,
|
||||
)
|
||||
def query(session: Session) -> Optional[Subscribe]:
|
||||
"""在调用方同步会话中执行类型媒体查询。"""
|
||||
condition = Subscribe._identity_condition( # pylint: disable=protected-access
|
||||
media_source, media_id, music_type
|
||||
)
|
||||
if condition is None:
|
||||
return None
|
||||
statement = select(Subscribe).where(condition, Subscribe.type == type)
|
||||
if season is not None:
|
||||
statement = statement.where(Subscribe.season == season)
|
||||
return session.execute(statement).scalars().first()
|
||||
return self._execute_sync_query(query)
|
||||
|
||||
async def async_get_by(
|
||||
self, type: str, media_source: MediaSource, media_id: str,
|
||||
@@ -422,25 +444,55 @@ class SubscribeOper(DbOper):
|
||||
"""
|
||||
根据条件查询订阅
|
||||
"""
|
||||
return await Subscribe.async_get_by(
|
||||
self._db, type, media_source, media_id, season, music_type,
|
||||
)
|
||||
async def query(session: AsyncSession) -> Optional[Subscribe]:
|
||||
"""在调用方异步会话中执行类型媒体查询。"""
|
||||
condition = Subscribe._identity_condition( # pylint: disable=protected-access
|
||||
media_source, media_id, music_type
|
||||
)
|
||||
if condition is None:
|
||||
return None
|
||||
statement = select(Subscribe).where(condition, Subscribe.type == type)
|
||||
if season is not None:
|
||||
statement = statement.where(Subscribe.season == season)
|
||||
result = await session.execute(statement)
|
||||
return result.scalars().first()
|
||||
return await self._execute_async_query(query)
|
||||
|
||||
def list(self, state: Optional[str] = None) -> List[Subscribe]:
|
||||
"""
|
||||
获取订阅列表
|
||||
"""
|
||||
if state:
|
||||
return Subscribe.get_by_state(self._db, state)
|
||||
return Subscribe.list(self._db)
|
||||
return self._execute_sync_query(
|
||||
lambda session: list(session.execute(
|
||||
select(Subscribe).where(Subscribe.state.in_(state.split(',')))
|
||||
).scalars().all())
|
||||
)
|
||||
return self._execute_sync_query(
|
||||
lambda session: list(session.execute(select(Subscribe)).scalars().all())
|
||||
)
|
||||
|
||||
async def async_list(self, state: Optional[str] = None) -> List[Subscribe]:
|
||||
"""
|
||||
异步获取订阅列表
|
||||
"""
|
||||
if self._db is not None and not isinstance(self._db, (Session, AsyncSession)):
|
||||
if state:
|
||||
return await Subscribe.async_get_by_state(self._db, state)
|
||||
return await Subscribe.async_list(self._db)
|
||||
if state:
|
||||
return await Subscribe.async_get_by_state(self._db, state)
|
||||
return await Subscribe.async_list(self._db)
|
||||
async def query(session: AsyncSession) -> List[Subscribe]:
|
||||
"""在调用方异步会话中执行状态列表查询。"""
|
||||
result = await session.execute(
|
||||
select(Subscribe).where(Subscribe.state.in_(state.split(',')))
|
||||
)
|
||||
return list(result.scalars().all())
|
||||
return await self._execute_async_query(query)
|
||||
async def query_all(session: AsyncSession) -> List[Subscribe]:
|
||||
"""在调用方异步会话中执行全量订阅查询。"""
|
||||
result = await session.execute(select(Subscribe))
|
||||
return list(result.scalars().all())
|
||||
return await self._execute_async_query(query_all)
|
||||
|
||||
async def async_list_by_username(
|
||||
self,
|
||||
@@ -449,12 +501,20 @@ class SubscribeOper(DbOper):
|
||||
mtype: Optional[str] = None,
|
||||
) -> List[Subscribe]:
|
||||
"""异步按用户获取订阅。"""
|
||||
return await Subscribe.async_list_by_username(
|
||||
self._db,
|
||||
username=username,
|
||||
state=state,
|
||||
mtype=mtype,
|
||||
)
|
||||
if self._db is not None and not isinstance(self._db, (Session, AsyncSession)):
|
||||
return await Subscribe.async_list_by_username(
|
||||
self._db, username=username, state=state, mtype=mtype
|
||||
)
|
||||
async def query(session: AsyncSession) -> List[Subscribe]:
|
||||
"""在调用方异步会话中执行用户筛选查询。"""
|
||||
statement = select(Subscribe).where(Subscribe.username == username)
|
||||
if state:
|
||||
statement = statement.where(Subscribe.state == state)
|
||||
if mtype:
|
||||
statement = statement.where(Subscribe.type == mtype)
|
||||
result = await session.execute(statement)
|
||||
return list(result.scalars().all())
|
||||
return await self._execute_async_query(query)
|
||||
|
||||
async def async_list_by_title(
|
||||
self,
|
||||
@@ -462,11 +522,14 @@ class SubscribeOper(DbOper):
|
||||
season: Optional[int] = None,
|
||||
) -> List[Subscribe]:
|
||||
"""异步按标题获取订阅,供旧查询测试和迁移调用兼容。"""
|
||||
return await Subscribe.async_list_by_title(
|
||||
self._db,
|
||||
title=title,
|
||||
season=season,
|
||||
)
|
||||
async def query(session: AsyncSession) -> List[Subscribe]:
|
||||
"""在调用方异步会话中执行标题列表查询。"""
|
||||
statement = select(Subscribe).where(Subscribe.name == title)
|
||||
if season is not None:
|
||||
statement = statement.where(Subscribe.season == season)
|
||||
result = await session.execute(statement)
|
||||
return list(result.scalars().all())
|
||||
return await self._execute_async_query(query)
|
||||
|
||||
def delete(self, sid: int):
|
||||
"""
|
||||
@@ -535,13 +598,30 @@ class SubscribeOper(DbOper):
|
||||
"""
|
||||
获取指定用户的订阅
|
||||
"""
|
||||
return Subscribe.list_by_username(self._db, username=username, state=state, mtype=mtype)
|
||||
def query(session: Session) -> List[Subscribe]:
|
||||
"""在调用方同步会话中执行用户筛选查询。"""
|
||||
statement = select(Subscribe).where(Subscribe.username == username)
|
||||
if state:
|
||||
statement = statement.where(Subscribe.state == state)
|
||||
if mtype:
|
||||
statement = statement.where(Subscribe.type == mtype)
|
||||
return list(session.execute(statement).scalars().all())
|
||||
return self._execute_sync_query(query)
|
||||
|
||||
def list_by_type(self, mtype: str, days: int = 7) -> List[Subscribe]:
|
||||
"""
|
||||
获取指定类型的订阅
|
||||
"""
|
||||
return Subscribe.list_by_type(self._db, mtype=mtype, days=days)
|
||||
def query(session: Session) -> List[Subscribe]:
|
||||
"""在调用方同步会话中执行时间窗订阅查询。"""
|
||||
cutoff = time.strftime(
|
||||
"%Y-%m-%d %H:%M:%S",
|
||||
time.localtime(time.time() - 86400 * int(days)),
|
||||
)
|
||||
return list(session.execute(select(Subscribe).where(
|
||||
Subscribe.type == mtype, Subscribe.date >= cutoff
|
||||
)).scalars().all())
|
||||
return self._execute_sync_query(query)
|
||||
|
||||
def add_history(self, **kwargs):
|
||||
"""
|
||||
|
||||
+13
-2
@@ -12,6 +12,7 @@ runtime 兼容映射指向 SDK 薄门面;canonical 数据访问模块仍只依
|
||||
from typing import List, Optional
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.db.base import DbOper
|
||||
from app.db.models.user import User
|
||||
@@ -103,13 +104,23 @@ class UserOper(DbOper):
|
||||
"""
|
||||
异步根据用户名获取用户。
|
||||
"""
|
||||
return await User.async_get_by_name(self._db, name)
|
||||
async def query(session: AsyncSession) -> Optional[User]:
|
||||
"""在调用方异步会话中执行用户名查询。"""
|
||||
result = await session.execute(select(User).where(User.name == name))
|
||||
return result.scalars().first()
|
||||
|
||||
return await self._execute_async_query(query)
|
||||
|
||||
async def async_get_by_id(self, user_id: int) -> Optional[User]:
|
||||
"""
|
||||
异步根据用户 ID 获取用户。
|
||||
"""
|
||||
return await User.async_get_by_id(self._db, user_id)
|
||||
async def query(session: AsyncSession) -> Optional[User]:
|
||||
"""在调用方异步会话中执行用户 ID 查询。"""
|
||||
result = await session.execute(select(User).where(User.id == user_id))
|
||||
return result.scalars().first()
|
||||
|
||||
return await self._execute_async_query(query)
|
||||
|
||||
def get_permissions(self, name: str) -> dict:
|
||||
"""
|
||||
|
||||
@@ -17,6 +17,16 @@ class ModuleResultAggregation(StrEnum):
|
||||
ORDERED_LIST_MERGE = "ordered_list_merge"
|
||||
|
||||
|
||||
class ModuleResultShape(StrEnum):
|
||||
"""描述模块 provider 返回值的基础 Python 形状。"""
|
||||
|
||||
ANY = "any"
|
||||
LIST = "list"
|
||||
STRING = "string"
|
||||
MAPPING = "mapping"
|
||||
BOOLEAN = "boolean"
|
||||
|
||||
|
||||
class ModuleExecutionMode(StrEnum):
|
||||
"""描述 provider 可以采用的执行形态。"""
|
||||
|
||||
@@ -45,6 +55,7 @@ class ModuleMethodContract:
|
||||
version: int = 1
|
||||
input_contract: str = "legacy_args"
|
||||
result_contract: str = "Any"
|
||||
result_shape: ModuleResultShape = ModuleResultShape.ANY
|
||||
required_parameters: tuple[str, ...] = ()
|
||||
execution: ModuleExecutionMode = ModuleExecutionMode.SYNC_OR_ASYNC
|
||||
timeout_policy: str = "caller_budget"
|
||||
@@ -74,11 +85,11 @@ _METHOD_CONTRACTS = {
|
||||
"media_category": ModuleMethodContract(family="media-recognition", input_contract="MediaCategoryRequest", result_contract="CategoryConfig | None"),
|
||||
"mediaserver_items": ModuleMethodContract(family="media-server", input_contract="MediaServerItemsRequest", result_contract="list[MediaServerItem]", aggregation=ModuleResultAggregation.ORDERED_LIST_MERGE, required_parameters=("server", "library_id", "start_index", "limit")),
|
||||
"mediaserver_iteminfo": ModuleMethodContract(family="media-server", input_contract="MediaServerItemRequest", result_contract="MediaServerItem | None", aggregation=ModuleResultAggregation.FIRST_NON_EMPTY, required_parameters=("server", "item_id")),
|
||||
"mediaserver_play_url": ModuleMethodContract(family="media-server", input_contract="MediaServerPlayRequest", result_contract="str | None", aggregation=ModuleResultAggregation.FIRST_NON_EMPTY, required_parameters=("server", "item_id")),
|
||||
"mediaserver_play_url": ModuleMethodContract(family="media-server", input_contract="MediaServerPlayRequest", result_contract="str | None", result_shape=ModuleResultShape.STRING, aggregation=ModuleResultAggregation.FIRST_NON_EMPTY, required_parameters=("server", "item_id")),
|
||||
"mediaserver_tv_episodes": ModuleMethodContract(family="media-server", input_contract="MediaServerEpisodesRequest", result_contract="list[MediaServerPlayItem]", aggregation=ModuleResultAggregation.ORDERED_LIST_MERGE, required_parameters=("server", "item_id")),
|
||||
"download_file": ModuleMethodContract(family="storage", input_contract="StorageDownloadRequest", result_contract="FileItem | None", aggregation=ModuleResultAggregation.FIRST_NON_EMPTY, required_parameters=("fileitem", "path")),
|
||||
"upload_file": ModuleMethodContract(family="storage", input_contract="StorageUploadRequest", result_contract="FileItem | None", aggregation=ModuleResultAggregation.FIRST_NON_EMPTY, required_parameters=("fileitem", "path", "new_name")),
|
||||
"list_files": ModuleMethodContract(family="storage", input_contract="StorageListRequest", result_contract="list[FileItem]", aggregation=ModuleResultAggregation.ORDERED_LIST_MERGE, required_parameters=("fileitem", "recursion")),
|
||||
"list_files": ModuleMethodContract(family="storage", input_contract="StorageListRequest", result_contract="list[FileItem]", result_shape=ModuleResultShape.LIST, aggregation=ModuleResultAggregation.ORDERED_LIST_MERGE, required_parameters=("fileitem", "recursion")),
|
||||
"get_file_item": ModuleMethodContract(family="storage", input_contract="StorageItemRequest", result_contract="FileItem | None", aggregation=ModuleResultAggregation.FIRST_NON_EMPTY, required_parameters=("storage", "path")),
|
||||
"get_folder": ModuleMethodContract(family="storage", input_contract="StorageFolderRequest", result_contract="FileItem | None", aggregation=ModuleResultAggregation.FIRST_NON_EMPTY, required_parameters=("storage", "path")),
|
||||
"get_parent_item": ModuleMethodContract(family="storage", input_contract="StorageParentRequest", result_contract="FileItem | None", aggregation=ModuleResultAggregation.FIRST_NON_EMPTY, required_parameters=("fileitem",)),
|
||||
@@ -455,6 +466,22 @@ def diagnose_module_callable(method: str, callback: Callable[..., Any]) -> tuple
|
||||
return tuple(f"missing-parameter:{name}" for name in missing)
|
||||
|
||||
|
||||
def diagnose_module_result(method: str, result: Any) -> tuple[str, ...]:
|
||||
"""诊断显式模块结果的基础形状,兼容阶段只告警而不改写返回值。"""
|
||||
shape = get_module_method_contract(method).result_shape
|
||||
if shape is ModuleResultShape.ANY or result is None:
|
||||
return ()
|
||||
matches = {
|
||||
ModuleResultShape.LIST: isinstance(result, list),
|
||||
ModuleResultShape.STRING: isinstance(result, str),
|
||||
ModuleResultShape.MAPPING: isinstance(result, dict),
|
||||
ModuleResultShape.BOOLEAN: isinstance(result, bool),
|
||||
}
|
||||
if matches.get(shape, True):
|
||||
return ()
|
||||
return (f"unexpected-result:{shape.value}:{type(result).__name__}",)
|
||||
|
||||
|
||||
def list_explicit_module_contracts() -> dict[str, ModuleMethodContract]:
|
||||
"""返回显式方法清单的副本,供架构基线和 SDK 文档使用。"""
|
||||
return dict(_METHOD_CONTRACTS)
|
||||
|
||||
@@ -12,6 +12,7 @@ from app.runtime.log import logger
|
||||
from app.runtime.observability import observe_duration, record_metric
|
||||
from app.runtime.extensions.module.contracts import (
|
||||
diagnose_module_callable,
|
||||
diagnose_module_result,
|
||||
get_module_method_contract,
|
||||
is_explicit_module_method,
|
||||
)
|
||||
@@ -134,8 +135,10 @@ class ModuleInvocationDispatcher:
|
||||
logger.info("请求插件 %s 执行:%s ...", plugin_name, method)
|
||||
if self.is_valid_empty(result):
|
||||
result = func(*args, **kwargs)
|
||||
self._diagnose_result(method, result, "plugin")
|
||||
elif isinstance(result, list):
|
||||
temp = func(*args, **kwargs)
|
||||
self._diagnose_result(method, temp, "plugin")
|
||||
if isinstance(temp, list):
|
||||
result.extend(temp)
|
||||
else:
|
||||
@@ -187,8 +190,10 @@ class ModuleInvocationDispatcher:
|
||||
logger.info("请求插件 %s 执行:%s ...", plugin_name, method)
|
||||
if self.is_valid_empty(result):
|
||||
result = await self._async_call(func, *args, **kwargs)
|
||||
self._diagnose_result(method, result, "plugin")
|
||||
elif isinstance(result, list):
|
||||
temp = await self._async_call(func, *args, **kwargs)
|
||||
self._diagnose_result(method, temp, "plugin")
|
||||
if isinstance(temp, list):
|
||||
result.extend(temp)
|
||||
else:
|
||||
@@ -238,10 +243,13 @@ class ModuleInvocationDispatcher:
|
||||
self._diagnose_callable(method, func, f"宿主模块 {module_id}")
|
||||
if self.is_valid_empty(result):
|
||||
result = func(*args, **kwargs)
|
||||
self._diagnose_result(method, result, "system")
|
||||
elif ObjectUtils.check_signature(func, result):
|
||||
result = func(result)
|
||||
self._diagnose_result(method, result, "system")
|
||||
elif isinstance(result, list):
|
||||
temp = func(*args, **kwargs)
|
||||
self._diagnose_result(method, temp, "system")
|
||||
if isinstance(temp, list):
|
||||
result.extend(temp)
|
||||
else:
|
||||
@@ -291,10 +299,13 @@ class ModuleInvocationDispatcher:
|
||||
self._diagnose_callable(method, func, f"宿主模块 {module_id}")
|
||||
if self.is_valid_empty(result):
|
||||
result = await self._async_call(func, *args, **kwargs)
|
||||
self._diagnose_result(method, result, "system")
|
||||
elif ObjectUtils.check_signature(func, result):
|
||||
result = await self._async_call(func, result)
|
||||
self._diagnose_result(method, result, "system")
|
||||
elif isinstance(result, list):
|
||||
temp = await self._async_call(func, *args, **kwargs)
|
||||
self._diagnose_result(method, temp, "system")
|
||||
if isinstance(temp, list):
|
||||
result.extend(temp)
|
||||
else:
|
||||
@@ -360,6 +371,24 @@ class ModuleInvocationDispatcher:
|
||||
", ".join(problems),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _diagnose_result(method: str, result: Any, provider_type: str) -> None:
|
||||
"""记录 provider 结果形状偏差,保持旧插件返回值原样继续执行。"""
|
||||
problems = diagnose_module_result(method, result)
|
||||
if problems:
|
||||
record_metric(
|
||||
"module.contract.result_mismatch",
|
||||
method=method,
|
||||
provider_type=provider_type,
|
||||
problem=problems[0],
|
||||
)
|
||||
logger.warning(
|
||||
"模块方法 %s 的 %s provider 返回值与契约不一致:%s;当前仅诊断",
|
||||
method,
|
||||
provider_type,
|
||||
", ".join(problems),
|
||||
)
|
||||
|
||||
async def _async_call(
|
||||
self,
|
||||
func: Callable[..., Any],
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
"""进程内后台任务登记与生命周期收口。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Coroutine
|
||||
from dataclasses import dataclass
|
||||
from functools import partial
|
||||
from typing import Any, Callable
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TaskRecord:
|
||||
"""记录一个后台任务的所有者,便于关停阶段按责任域收口。"""
|
||||
|
||||
owner: str
|
||||
task: asyncio.Task[Any]
|
||||
cancel_on_shutdown: bool
|
||||
|
||||
|
||||
class TaskRegistry:
|
||||
"""管理由宿主创建的进程内后台任务,并提供统一取消与等待入口。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""初始化空任务登记表。"""
|
||||
self._records: dict[asyncio.Task[Any], TaskRecord] = {}
|
||||
self._accepting = True
|
||||
|
||||
@property
|
||||
def records(self) -> tuple[TaskRecord, ...]:
|
||||
"""返回当前仍未完成的任务快照。"""
|
||||
return tuple(
|
||||
record for record in self._records.values() if not record.task.done()
|
||||
)
|
||||
|
||||
def create(
|
||||
self,
|
||||
coroutine: Coroutine[Any, Any, Any],
|
||||
*,
|
||||
owner: str,
|
||||
cancel_on_shutdown: bool = True,
|
||||
) -> asyncio.Task[Any]:
|
||||
"""创建并登记后台任务,任务完成后自动从登记表移除。"""
|
||||
if not self._accepting:
|
||||
coroutine.close()
|
||||
raise RuntimeError("后台任务登记器正在关闭,不能再创建新任务")
|
||||
task = asyncio.create_task(coroutine, name=owner)
|
||||
self.register(
|
||||
task,
|
||||
owner=owner,
|
||||
cancel_on_shutdown=cancel_on_shutdown,
|
||||
)
|
||||
return task
|
||||
|
||||
def create_sync(
|
||||
self,
|
||||
function: Callable[..., Any],
|
||||
*args: Any,
|
||||
owner: str,
|
||||
**kwargs: Any,
|
||||
) -> asyncio.Task[Any]:
|
||||
"""在线程池执行同步后台函数并登记其异步生命周期。"""
|
||||
return self.create(
|
||||
asyncio.to_thread(partial(function, *args, **kwargs)),
|
||||
owner=owner,
|
||||
cancel_on_shutdown=False,
|
||||
)
|
||||
|
||||
def register(
|
||||
self,
|
||||
task: asyncio.Task[Any],
|
||||
*,
|
||||
owner: str,
|
||||
cancel_on_shutdown: bool = True,
|
||||
) -> asyncio.Task[Any]:
|
||||
"""登记已有任务并绑定责任域。"""
|
||||
if not self._accepting:
|
||||
task.cancel()
|
||||
raise RuntimeError("后台任务登记器正在关闭,不能再登记新任务")
|
||||
task.set_name(owner)
|
||||
self._records[task] = TaskRecord(
|
||||
owner=owner,
|
||||
task=task,
|
||||
cancel_on_shutdown=cancel_on_shutdown,
|
||||
)
|
||||
task.add_done_callback(self._discard)
|
||||
return task
|
||||
|
||||
def _discard(self, task: asyncio.Task[Any]) -> None:
|
||||
"""移除已结束任务,并把未处理异常交给事件循环统一报告。"""
|
||||
record = self._records.pop(task, None)
|
||||
if task.cancelled():
|
||||
return
|
||||
exception = task.exception()
|
||||
if exception is not None:
|
||||
task.get_loop().call_exception_handler(
|
||||
{
|
||||
"message": "MoviePilot 后台任务执行失败",
|
||||
"exception": exception,
|
||||
"task": task,
|
||||
"owner": record.owner if record else task.get_name(),
|
||||
}
|
||||
)
|
||||
|
||||
async def shutdown(self, *, timeout_seconds: float = 10.0) -> None:
|
||||
"""取消并等待全部登记任务,超时后放弃等待但不影响其他关闭步骤。"""
|
||||
self._accepting = False
|
||||
records = self.records
|
||||
tasks = [record.task for record in records]
|
||||
for record in records:
|
||||
if record.cancel_on_shutdown:
|
||||
record.task.cancel()
|
||||
if tasks:
|
||||
_, pending = await asyncio.wait(tasks, timeout=timeout_seconds)
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
self._records.clear()
|
||||
|
||||
|
||||
_default_registry = TaskRegistry()
|
||||
_runtime_registry: TaskRegistry | None = None
|
||||
|
||||
|
||||
def configure_task_registry(registry: TaskRegistry | None) -> None:
|
||||
"""由启动组合根发布当前 lifespan 的任务登记器。"""
|
||||
global _runtime_registry
|
||||
_runtime_registry = registry
|
||||
|
||||
|
||||
def get_task_registry() -> TaskRegistry:
|
||||
"""返回当前宿主任务登记器,未启动完整 lifespan 时保留测试兼容回退。"""
|
||||
return _runtime_registry or _default_registry
|
||||
@@ -1,9 +1,11 @@
|
||||
"""宿主启动阶段构建的类型化运行时上下文。"""
|
||||
|
||||
from collections.abc import AsyncGenerator, Callable, Generator
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Protocol
|
||||
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
|
||||
from app.application.messaging.chat import (
|
||||
AsyncAgentChatRepository,
|
||||
AgentChatPersistenceService,
|
||||
@@ -193,3 +195,4 @@ class HostRuntime:
|
||||
workflow: WorkflowRuntime
|
||||
configuration: RuntimeConfiguration
|
||||
settings: RuntimeSettingsService
|
||||
tasks: TaskRegistry = field(default_factory=TaskRegistry)
|
||||
|
||||
@@ -32,6 +32,7 @@ from app.runtime.settings import RuntimeSettingsCompat
|
||||
settings = RuntimeSettingsCompat()
|
||||
from app.runtime.health import get_application_health
|
||||
from app.runtime.topology import validate_process_topology
|
||||
from app.runtime.tasks import TaskRegistry, configure_task_registry
|
||||
from app.adapters.external.server import MoviePilotServerHelper
|
||||
from app.runtime.state import SystemHelper
|
||||
from app.runtime.log import logger, LoggerManager
|
||||
@@ -165,6 +166,24 @@ async def initialize_modules_component(app: FastAPI) -> None:
|
||||
app.state.host_runtime = runtime
|
||||
|
||||
|
||||
def initialize_task_registry(app: FastAPI) -> None:
|
||||
"""创建当前 lifespan 独占的后台任务登记器。"""
|
||||
task_registry = TaskRegistry()
|
||||
app.state.task_registry = task_registry
|
||||
configure_task_registry(task_registry)
|
||||
|
||||
|
||||
async def stop_task_registry(app: FastAPI) -> None:
|
||||
"""停止接收新后台任务,并取消、等待当前 lifespan 的存量任务。"""
|
||||
task_registry = getattr(app.state, "task_registry", None)
|
||||
try:
|
||||
if isinstance(task_registry, TaskRegistry):
|
||||
await task_registry.shutdown(timeout_seconds=30.0)
|
||||
finally:
|
||||
configure_task_registry(None)
|
||||
app.state.task_registry = None
|
||||
|
||||
|
||||
def prepare_plugin_restore() -> None:
|
||||
"""先装配插件外部系统服务,再恢复插件及其依赖。"""
|
||||
configure_plugin_services()
|
||||
@@ -188,6 +207,15 @@ def prepare_database_component(app: FastAPI) -> None:
|
||||
def build_lifecycle_components(app: FastAPI) -> tuple[LifecycleComponent, ...]:
|
||||
"""按现有顺序构建应用组件清单,回调在每次 lifespan 启动时重新绑定。"""
|
||||
return (
|
||||
LifecycleComponent(
|
||||
name="后台任务登记器",
|
||||
start=lambda: initialize_task_registry(app),
|
||||
stop=lambda: stop_task_registry(app),
|
||||
start_order=5,
|
||||
stop_order=5,
|
||||
start_timeout_seconds=30,
|
||||
stop_timeout_seconds=60,
|
||||
),
|
||||
LifecycleComponent(
|
||||
name="数据库准备",
|
||||
start=lambda: prepare_database_component(app),
|
||||
@@ -371,10 +399,16 @@ async def lifespan(app: FastAPI):
|
||||
sync_plugins_task = asyncio.create_task(
|
||||
run_startup_step("插件同步与启动收尾", init_extra)
|
||||
)
|
||||
task_registry = app.state.task_registry
|
||||
task_registry.register(sync_plugins_task, owner="startup.plugin_settlement")
|
||||
health.mark_ready()
|
||||
except BaseException:
|
||||
# Uvicorn 在 lifespan 抛错时不会开始接流量;状态仍需供嵌入式入口和测试诊断。
|
||||
health.mark_failed()
|
||||
try:
|
||||
await stop_task_registry(app)
|
||||
except Exception as cleanup_error:
|
||||
logger.error(f"启动失败后的后台任务清理失败:{cleanup_error}")
|
||||
raise
|
||||
try:
|
||||
# 在此处 yield,表示应用已经启动,控制权交回 FastAPI 主事件循环
|
||||
|
||||
@@ -171,6 +171,7 @@ from app.runtime.extensions.service_config import (
|
||||
ServiceConfigHelper,
|
||||
configure_service_config_reader,
|
||||
)
|
||||
from app.runtime.tasks import get_task_registry
|
||||
|
||||
|
||||
_database_worker: DatabaseWorker | None = None
|
||||
@@ -719,6 +720,7 @@ async def init_modules() -> HostRuntime:
|
||||
),
|
||||
configuration=runtime_configuration,
|
||||
settings=runtime_settings,
|
||||
tasks=get_task_registry(),
|
||||
)
|
||||
configure_runtime_configuration(host_runtime.configuration)
|
||||
configure_runtime_settings(host_runtime.settings)
|
||||
|
||||
@@ -59,8 +59,12 @@ Event Contract Registry 是 53 个事件的逐项机器清单。下表按相同
|
||||
|
||||
### FastAPI BackgroundTasks
|
||||
|
||||
- 订阅手工搜索调度、插件市场刷新、低价值上报:E1;响应成功只表示已接受本进程调度,不表示执行完成。
|
||||
- 若任务源于已提交的用户数据且不可从数据库重建,必须提升为 E2,不得继续新增裸 BackgroundTasks。
|
||||
- 订阅手工搜索调度、插件市场刷新、低价值上报、CookieCloud 手工调度:E1;响应成功只表示已接受本进程
|
||||
调度,不表示执行完成。
|
||||
- Webhook E0 广播、消息入口和 Seerr 订阅入口均已迁入 lifespan TaskRegistry,具备 owner、停止接收和
|
||||
有限等待语义;进程崩溃时仍允许丢失,不因此提升为 durable。
|
||||
- 主仓不再新增或保留裸 FastAPI `BackgroundTasks`;若任务源于已提交的用户数据且不可从数据库重建,
|
||||
必须提升为 E2,进入 Outbox 或持久任务表。
|
||||
|
||||
### Scheduler jobs
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@
|
||||
> 审计范围:宿主后端;排除 `app/plugins/**` 运行时插件副本
|
||||
> 规范优先级:`AGENTS.md` 与 `docs/rules/` 高于本文
|
||||
> 相关文档:`docs/architecture-overview.md`、`docs/refactor/backend-architecture-governance.md`、`docs/refactor/backend-module-refactor-compatibility.md`
|
||||
> 实施进度:阶段 0~6 的宿主架构能力已完成收口;API/Application 公共复杂度基线已清零,启动组合根的 SystemConfigOper 构造点已由 14 降至 1;插件仓适配、Outbox 外围扩展和 Model 查询兼容面仍按风险切片推进。
|
||||
> 实施进度:阶段 0~6 的宿主架构能力已完成收口;API/Application 公共复杂度基线已清零,启动组合根的 SystemConfigOper 构造点已由 14 降至 1;API 进程内后台任务已完成首批统一登记,插件仓适配、Outbox 外围扩展和 Model 查询兼容面仍按风险切片推进。
|
||||
|
||||
## 当前复核结论(2026-08-23)
|
||||
|
||||
@@ -26,9 +26,9 @@
|
||||
|
||||
### P1:需要优先治理的真实债务
|
||||
|
||||
1. **后台任务没有统一的所有权和恢复模型。** 当前约有 `50` 个 `create_task`/等价任务创建点,另有 FastAPI `BackgroundTasks`、线程池和 APScheduler 并存。生命周期清单能关闭模块、插件、调度器和 Agent,但 API 层的若干任务集合(如 `app/api/endpoints/agent.py`、`app/api/endpoints/plugin.py`)没有统一注册到 HostRuntime,也没有在 shutdown 阶段统一等待或取消。`app/api/endpoints/webhook.py:32`、`app/api/endpoints/site.py:178` 这类接口会先返回成功,再执行关键副作用;进程崩溃、重启或客户端断开时可能丢失。需要为每类任务明确 owner、取消、等待、重试、幂等和是否 durable,关键业务副作用优先接入已有 Outbox/恢复表。
|
||||
1. **后台任务的统一所有权已覆盖 API 入口,但仍有更深层任务机制待分级。** `app/runtime/tasks.py` 已建立 lifespan 级 TaskRegistry,启动收尾、插件 Release 刷新、Webhook E0 广播、CookieCloud E1 手工调度、消息入口、Seerr 订阅和 WebAgent 断线后执行/快照保存均不再维护端点模块级任务集合或 Starlette 回调,shutdown 会停止接收、取消并有限等待,且生命周期清单明确登记其顺序。主仓 `app/` 已无裸 FastAPI `BackgroundTasks`;当前仍有约 `50` 个更底层 `create_task`/等价任务创建点,与线程池和 APScheduler 并存,后续需逐项确认 owner、取消、等待、重试、幂等和是否 durable,关键业务副作用优先接入已有 Outbox/恢复表。
|
||||
2. **动态模块契约仍以 legacy 聚合语义为主。** 当前登记 `212` 个模块方法,其中 `194` 个仍使用 `legacy` aggregation,只有 `14` 个 `first_non_empty`、`4` 个 `ordered_list_merge`。`app/runtime/extensions/module/contracts.py:422-455` 已能登记 family、输入/结果标签和基础签名诊断,但 `193` 个方法没有 required parameters,调度器 `app/runtime/extensions/module/dispatcher.py:109-260` 仍主要依赖运行时反射、返回值形状和短路规则。未知第三方方法保留 legacy fallback 是兼容要求,不应删除;宿主高频能力则应逐族补齐可执行的输入校验、结果校验、超时和错误语义。
|
||||
3. **查询侧数据库兼容 ABI 仍未完全收口。** 写事务装饰器已降为 `0`,事务所有权已经明显改善;但 `app/db/models` 仍有 `106` 个 `db_query/async_db_query`(`62` 个同步、`44` 个异步)。这些装饰器会在调用方未传 Session 时隐式创建并关闭会话(见 `app/db/decorators.py:224-298`),查询返回的 ORM 对象仍可能跨层流转,导致事务组合、对象生命周期和懒加载行为需要依赖隐含约定。应按高频业务路径逐步迁移到显式 Query/Repository + 请求/任务级 Session,不宜一次性全仓改写。
|
||||
3. **查询侧数据库兼容 ABI 仍未完全收口。** 写事务装饰器已降为 `0`,事务所有权已经明显改善;但 `app/db/models` 仍有 `75` 个 `db_query/async_db_query`(`48` 个同步、`27` 个异步)。这些装饰器会在调用方未传 Session 时隐式创建并关闭会话(见 `app/db/decorators.py:224-298`),查询返回的 ORM 对象仍可能跨层流转,导致事务组合、对象生命周期和懒加载行为需要依赖隐含约定。站点、消息、用户和订阅高频查询已迁到对应 Oper 显式 Session 路径,后续继续按历史等风险切片迁移,不一次性全仓改写。
|
||||
4. **组合根和全局状态仍形成复杂的隐式运行时图。** Singleton 实例、模块级 provider、`configure_*` 注册函数和兼容 Facade 同时存在;它们解决了旧 ABI 和启动顺序问题,但增加测试污染、重复装配、实例身份和初始化顺序风险。`app/startup/lifecycle/__init__.py:161-376` 已有声明式生命周期,`app/startup/modules_initializer.py:505-530` 也有分阶段关闭,但尚未做到所有进程级资源都只通过 typed HostRuntime 访问。后续应以“新代码禁止新增 Service Locator/Singleton 依赖、旧入口有命中观测”为 ratchet。
|
||||
|
||||
### P2:中长期可演进性债务
|
||||
@@ -713,6 +713,10 @@ ModuleMethodSpec(
|
||||
结果合同、执行、超时和错误语义;未知第三方自定义方法仍走开放 legacy fallback,不拒绝加载或执行。
|
||||
- 未知动态方法在真实 provider 命中时记录 `module.contract.legacy_hit`,区分插件/宿主调用方和 ABI 来源;
|
||||
该指标只在 callable 实际存在并准备执行时递增,不改变未知第三方方法的开放 fallback、聚合或异常语义。
|
||||
- 2026-08-23 增加 `result_shape` 基础结果形状合同,首批覆盖存储列表、媒体服务器列表/剧集、播放 URL、
|
||||
快照映射和无返回值能力。Dispatcher 在 provider 返回边界记录
|
||||
`module.contract.result_mismatch` 与期望形状;该阶段只观测和告警,不拒绝旧插件、不改写返回值,
|
||||
也不把业务对象类型强行导入动态调度器。未知第三方方法继续完全使用 legacy fallback。
|
||||
|
||||
#### ARCH-241:Event Contract Registry
|
||||
|
||||
@@ -814,6 +818,19 @@ ADR 必须逐个映射当前 Event、BackgroundTasks、Scheduler job、Agent tas
|
||||
- 首个 pilot 选择 `SubscribeAdded`,因为 ARCH-221 已有事务所有权与 post-commit 样板;文件整理
|
||||
继续保持 E3,不在本任务中被降格为普通事件重试。
|
||||
|
||||
**扩展实施记录(2026-08-23)**:
|
||||
|
||||
- 新增 `app/runtime/tasks.py` 的 `TaskRegistry`,作为当前 lifespan 的进程内后台任务所有权边界;
|
||||
生命周期清单新增“后台任务登记器”组件,启动失败和正常关闭均会停止接收新任务、取消存量任务并
|
||||
在有限等待窗口内收口,登记器不承担 durable queue 语义。
|
||||
- 插件 Release 后台刷新、WebAgent 断线后 Agent 执行以及消息展示快照保存均接入登记器,旧插件 API、
|
||||
SSE 协议和测试直接调用入口保持不变;未启动完整 ASGI lifespan 的旧调用继续使用兼容回退登记器。
|
||||
- Webhook E0 广播和站点 CookieCloud E1 手工调度已从 Starlette `BackgroundTasks` 迁入同一登记器;
|
||||
同步函数在线程池执行,关停时优先等待而非假设线程可强制取消。响应成功仍只表示本进程已接受,
|
||||
不能因为进程内任务已统一登记而宣称崩溃可恢复。
|
||||
- 订阅手工搜索、消息入口和 Seerr 订阅均已按 E0/E1 登记;其他关键业务副作用继续按等级逐项迁移,
|
||||
需要可靠交付的路径仍走 ARCH-251 的 Outbox/幂等切片,不扩大插件事件或 API payload。
|
||||
|
||||
#### ARCH-251:用现有数据库做首个 durable side-effect pilot
|
||||
|
||||
**前置**:ARCH-220/221 与 ARCH-241 完成。
|
||||
@@ -887,7 +904,7 @@ ADR 必须逐个映射当前 Event、BackgroundTasks、Scheduler job、Agent tas
|
||||
事务低水位从 174 降到 168,Oper 仍不创建 Session、也不直接 commit/rollback。
|
||||
- 剩余同步/异步 Model 写装饰器已全部迁移:AgentTask、PassKey、User、消息、历史清理、
|
||||
站点快照、媒体服务器、插件数据、TransferPending 等写入由调用方 Session 和 UoW 收口;无 Session
|
||||
的旧 Oper ABI 委托 Startup 注入的短事务执行器。当前 Model 装饰器仅剩 106 个查询装饰器,
|
||||
的旧 Oper ABI 委托 Startup 注入的短事务执行器。当前 Model 装饰器仅剩 75 个查询装饰器(同步 48、异步 27),
|
||||
`db_update` 与 `async_db_update` 均为 0,Oper 自建 Session/直接提交仍为 0。
|
||||
- 数据清理按批次显式提交 UoW,单表失败先回滚会话再继续汇总后续表;不再依赖删除 Model 的隐式提交。
|
||||
- 收尾批次进一步移除宿主 Oper 对 `Base.create/update/delete/truncate` 八个兼容包装器的调用:显式
|
||||
@@ -1157,7 +1174,7 @@ Settings 读取作为基础设施边界,架构基线已明确记录该例外
|
||||
`model_dump` 旧 Settings ABI,并由应用组合根注入服务对象,低层 runtime 不再反向导入 `app.application`;
|
||||
`SkillHelper` 的技能市场写入继续经过兼容代理,旧插件/测试的模块级替换语义保持。`UserConfigOper` 的
|
||||
无 Session 查询改为一次性兼容查询会话,显式 Session 仍由调用方持有。配置债务稳定为 8 个文件,Model
|
||||
查询装饰器稳定为 106 个且写装饰器为 0;四分片全量测试 `5441 passed, 3 skipped`,mypy、复杂度、异步阻塞、
|
||||
查询装饰器在消息、用户和订阅查询切片后进一步降至 75 个且写装饰器为 0;四分片全量测试 `5492 passed, 3 skipped`,mypy、复杂度、异步阻塞、
|
||||
host/plugin 架构基线均通过。
|
||||
|
||||
#### ARCH-272:异步阻塞检测
|
||||
@@ -1337,7 +1354,7 @@ rollback:
|
||||
| 基线写入行为 | 默认命令可能覆盖 fixture | 所有默认/check 命令保证工作树不变;write 必须显式 scope |
|
||||
| 全功能 worker | 配置允许 >1,控制面会复制 | 启动期明确拒绝 >1;文档与配置一致 |
|
||||
| 健康接口 | 认证 `/system/ping` 为主 | 分离公开 live 与受限/安全 ready;失败原因可诊断 |
|
||||
| Model 事务装饰器 | 当前 106 个且全部只读;写装饰器 0 | 查询债务只降不增;写事务不回退到 Model/Base 隐式提交 |
|
||||
| Model 事务装饰器 | 当前 75 个且全部只读;写装饰器 0 | 查询债务只降不增;写事务不回退到 Model/Base 隐式提交 |
|
||||
| 新写用例事务 | 宿主写 Oper 已脱离 Base 隐式提交 | 100% 由入口/Application 边界拥有 Session/UoW |
|
||||
| 高频 Module 契约 | 212 个宿主能力显式登记 | 新观察到的宿主方法必须同步登记完整契约 |
|
||||
| Event payload | 53 类型全部登记 typed payload 与可靠性 | 新事件必须同步登记,不回退裸 dict |
|
||||
|
||||
@@ -48,7 +48,7 @@ cycle passes through an established package that was not moved.
|
||||
| `app/foundation/` | Stateless, config-free and I/O-free primitives: reflection and dynamic import, crypto, DOM parsing, identity, collections, singleton, text conversion/segmentation, URL and version helpers |
|
||||
| `app/domain/` | Pure MoviePilot business semantics for media, recognition, sites and torrents; live configuration, persistence, transport and acceleration are injected |
|
||||
| `app/application/` | Focused stateful application services, configured capability selection and service-bound rules |
|
||||
| `app/runtime/` | Process-wide config, events, complete logging runtime, cache contracts/in-memory policy, execution, localization, scheduling, restart state, concurrency, GC and rate limits |
|
||||
| `app/runtime/` | Process-wide config, events, complete logging runtime, cache contracts/in-memory policy, execution, background-task ownership, localization, scheduling, restart state, concurrency, GC and rate limits |
|
||||
| `app/adapters/` | Concrete technical I/O and named external ecosystems, split by cache, network, system and external boundaries |
|
||||
| `app/sdk/` | Stable, deliberately curated imports for plugin authors |
|
||||
|
||||
@@ -88,6 +88,7 @@ create additional top-level directory categories.
|
||||
| `app/runtime/log.py` | Complete console/plugin/file logging runtime and shutdown |
|
||||
| `app/runtime/cache.py` | Cache protocols, memory implementations, decorators and proxies |
|
||||
| `app/runtime/managed_resources.py` | Provider-neutral acquisition, observation and shutdown facade for process-owned optional resources |
|
||||
| `app/runtime/tasks.py` | Lifespan-scoped ownership, cancellation and bounded shutdown waiting for in-process background tasks |
|
||||
| `app/runtime/state.py` | Process restart and update state |
|
||||
| `app/runtime/extensions/` | Module, plugin, configured-service and managed-resource discovery/registration/lifecycle adapters |
|
||||
| `app/runtime/compat/` | Standard-library-only exact legacy import routing, resource preflight scanning and DEBUG diagnostics |
|
||||
@@ -145,6 +146,9 @@ mechanism remains in `app/adapters/system/resource.py`.
|
||||
应用级启动顺序使用 `app/startup/lifecycle/components.py` 的组件描述声明依赖、
|
||||
normal/safe-mode 范围、start/stop 顺序、超时预算和失败策略。新增进程级资源不得只在
|
||||
`lifespan()` 中追加过程代码,必须先进入可导出的生命周期清单并补顺序快照测试。
|
||||
API 中允许丢失或可重建的进程内任务必须登记到 `app/runtime/tasks.py`;登记器先于其他
|
||||
运行资源启动,并在资源释放前停止接收、取消和有限等待。需要崩溃恢复的 E2/E3 副作用仍应
|
||||
进入 Outbox 或持久任务表,不能把 TaskRegistry 当成 durable queue。
|
||||
Runtime 关闭后不可逆;完整应用生命周期的再次启动必须由新进程承载,不能在同一解释器中重建局部资源域。
|
||||
插件需要浏览器时使用 `app.sdk.browser`,由宿主浏览器适配器协调资源,不直接依赖资源实现。
|
||||
旧插件若直接导入有资源前置条件的第三方包,compat 在插件 import 前递归扫描源码并保守准备资源;
|
||||
|
||||
@@ -84,7 +84,7 @@ Oper classes accept and return persistence values. Turning a `MediaInfo` or
|
||||
### Transaction ownership ratchet
|
||||
|
||||
- `tests/fixtures/architecture/transaction-debt-baseline.json` records the
|
||||
existing Model transaction decorators. The current 106 decorators are query-only
|
||||
existing Model transaction decorators. The current 75 decorators are query-only
|
||||
migration debt: they may decrease but must never increase or move to a new
|
||||
Model method. Both `db_update` and `async_db_update` must remain at zero.
|
||||
- New Model methods must not use `db_query`, `db_update`, `async_db_query`, or
|
||||
|
||||
+24
-3
@@ -13,8 +13,8 @@
|
||||
"runtime_to_db": [],
|
||||
"workflow_to_db": []
|
||||
},
|
||||
"edge_count": 6444,
|
||||
"edge_sha256": "b5db7b31c7ea4dd7311fcd9e11a49eb32feb6939f896ed85703b3573408752a8",
|
||||
"edge_count": 6464,
|
||||
"edge_sha256": "256ae6f9cd8950b0fe2743e81300551877eb595bda37b8ce114f439b33496af9",
|
||||
"edges": [
|
||||
"app -> app.runtime",
|
||||
"app -> app.runtime.compat",
|
||||
@@ -1510,6 +1510,8 @@
|
||||
"app.api.context -> app.application.subscription.delete",
|
||||
"app.api.context -> app.application.subscription.identity",
|
||||
"app.api.context -> app.application.subscription.mutation",
|
||||
"app.api.context -> app.runtime",
|
||||
"app.api.context -> app.runtime.tasks",
|
||||
"app.api.context -> app.startup",
|
||||
"app.api.context -> app.startup.context",
|
||||
"app.api.dependencies.agent -> app.api",
|
||||
@@ -1599,6 +1601,7 @@
|
||||
"app.api.dependencies.subscription -> app.runtime",
|
||||
"app.api.dependencies.subscription -> app.runtime.events",
|
||||
"app.api.dependencies.subscription -> app.runtime.log",
|
||||
"app.api.dependencies.subscription -> app.runtime.tasks",
|
||||
"app.api.dependencies.subscription -> app.schemas",
|
||||
"app.api.dependencies.subscription -> app.schemas.types",
|
||||
"app.api.dependencies.subscription -> app.startup",
|
||||
@@ -1945,6 +1948,7 @@
|
||||
"app.api.endpoints.message -> app.adapters.web.security",
|
||||
"app.api.endpoints.message -> app.adapters.web.security.access",
|
||||
"app.api.endpoints.message -> app.api",
|
||||
"app.api.endpoints.message -> app.api.context",
|
||||
"app.api.endpoints.message -> app.api.dependencies",
|
||||
"app.api.endpoints.message -> app.api.dependencies.agent",
|
||||
"app.api.endpoints.message -> app.api.dependencies.auth",
|
||||
@@ -1961,6 +1965,7 @@
|
||||
"app.api.endpoints.message -> app.runtime.extensions",
|
||||
"app.api.endpoints.message -> app.runtime.extensions.service_config",
|
||||
"app.api.endpoints.message -> app.runtime.log",
|
||||
"app.api.endpoints.message -> app.runtime.tasks",
|
||||
"app.api.endpoints.message -> app.schemas",
|
||||
"app.api.endpoints.message -> app.schemas.message",
|
||||
"app.api.endpoints.message -> app.schemas.response",
|
||||
@@ -2048,6 +2053,7 @@
|
||||
"app.api.endpoints.plugin -> app.adapters.web.security",
|
||||
"app.api.endpoints.plugin -> app.adapters.web.security.access",
|
||||
"app.api.endpoints.plugin -> app.api",
|
||||
"app.api.endpoints.plugin -> app.api.context",
|
||||
"app.api.endpoints.plugin -> app.api.dependencies",
|
||||
"app.api.endpoints.plugin -> app.api.dependencies.auth",
|
||||
"app.api.endpoints.plugin -> app.api.dependencies.plugin",
|
||||
@@ -2070,6 +2076,7 @@
|
||||
"app.api.endpoints.plugin -> app.runtime.extensions.plugin",
|
||||
"app.api.endpoints.plugin -> app.runtime.extensions.plugin.contracts",
|
||||
"app.api.endpoints.plugin -> app.runtime.log",
|
||||
"app.api.endpoints.plugin -> app.runtime.tasks",
|
||||
"app.api.endpoints.plugin -> app.schemas",
|
||||
"app.api.endpoints.plugin -> app.schemas.common",
|
||||
"app.api.endpoints.plugin -> app.schemas.plugin",
|
||||
@@ -2122,6 +2129,7 @@
|
||||
"app.api.endpoints.site -> app.adapters.web.security",
|
||||
"app.api.endpoints.site -> app.adapters.web.security.access",
|
||||
"app.api.endpoints.site -> app.api",
|
||||
"app.api.endpoints.site -> app.api.context",
|
||||
"app.api.endpoints.site -> app.api.dependencies",
|
||||
"app.api.endpoints.site -> app.api.dependencies.auth",
|
||||
"app.api.endpoints.site -> app.api.dependencies.site",
|
||||
@@ -2145,6 +2153,7 @@
|
||||
"app.api.endpoints.site -> app.domain.site",
|
||||
"app.api.endpoints.site -> app.runtime",
|
||||
"app.api.endpoints.site -> app.runtime.log",
|
||||
"app.api.endpoints.site -> app.runtime.tasks",
|
||||
"app.api.endpoints.site -> app.schemas",
|
||||
"app.api.endpoints.site -> app.schemas.common",
|
||||
"app.api.endpoints.site -> app.schemas.response",
|
||||
@@ -2180,6 +2189,7 @@
|
||||
"app.api.endpoints.subscribe -> app.adapters.web.security",
|
||||
"app.api.endpoints.subscribe -> app.adapters.web.security.access",
|
||||
"app.api.endpoints.subscribe -> app.api",
|
||||
"app.api.endpoints.subscribe -> app.api.context",
|
||||
"app.api.endpoints.subscribe -> app.api.dependencies",
|
||||
"app.api.endpoints.subscribe -> app.api.dependencies.auth",
|
||||
"app.api.endpoints.subscribe -> app.api.dependencies.subscription",
|
||||
@@ -2201,6 +2211,7 @@
|
||||
"app.api.endpoints.subscribe -> app.domain.metainfo",
|
||||
"app.api.endpoints.subscribe -> app.runtime",
|
||||
"app.api.endpoints.subscribe -> app.runtime.events",
|
||||
"app.api.endpoints.subscribe -> app.runtime.tasks",
|
||||
"app.api.endpoints.subscribe -> app.schemas",
|
||||
"app.api.endpoints.subscribe -> app.schemas.common",
|
||||
"app.api.endpoints.subscribe -> app.schemas.event",
|
||||
@@ -2347,9 +2358,12 @@
|
||||
"app.api.endpoints.webhook -> app.adapters.web.security",
|
||||
"app.api.endpoints.webhook -> app.adapters.web.security.access",
|
||||
"app.api.endpoints.webhook -> app.api",
|
||||
"app.api.endpoints.webhook -> app.api.context",
|
||||
"app.api.endpoints.webhook -> app.api.response",
|
||||
"app.api.endpoints.webhook -> app.chain",
|
||||
"app.api.endpoints.webhook -> app.chain.webhook",
|
||||
"app.api.endpoints.webhook -> app.runtime",
|
||||
"app.api.endpoints.webhook -> app.runtime.tasks",
|
||||
"app.api.endpoints.webhook -> app.schemas",
|
||||
"app.api.endpoints.webhook -> app.schemas.response",
|
||||
"app.api.endpoints.workflow -> app.adapters",
|
||||
@@ -2555,6 +2569,8 @@
|
||||
"app.application.mediaserver -> app.schemas.mediaserver",
|
||||
"app.application.mediaserver -> app.schemas.system",
|
||||
"app.application.mediaserver -> app.schemas.types",
|
||||
"app.application.messaging.agent -> app.runtime",
|
||||
"app.application.messaging.agent -> app.runtime.tasks",
|
||||
"app.application.messaging.agent -> app.schemas",
|
||||
"app.application.messaging.agent -> app.schemas.types",
|
||||
"app.application.messaging.chat -> app.application",
|
||||
@@ -5988,6 +6004,8 @@
|
||||
"app.startup.context -> app.application.subscription.identity",
|
||||
"app.startup.context -> app.application.subscription.mutation",
|
||||
"app.startup.context -> app.application.workflow",
|
||||
"app.startup.context -> app.runtime",
|
||||
"app.startup.context -> app.runtime.tasks",
|
||||
"app.startup.database -> app.adapters",
|
||||
"app.startup.database -> app.adapters.system",
|
||||
"app.startup.database -> app.adapters.system.backup",
|
||||
@@ -6050,6 +6068,7 @@
|
||||
"app.startup.lifecycle -> app.runtime.log",
|
||||
"app.startup.lifecycle -> app.runtime.settings",
|
||||
"app.startup.lifecycle -> app.runtime.state",
|
||||
"app.startup.lifecycle -> app.runtime.tasks",
|
||||
"app.startup.lifecycle -> app.runtime.topology",
|
||||
"app.startup.lifecycle -> app.startup",
|
||||
"app.startup.lifecycle -> app.startup.cache_initializer",
|
||||
@@ -6164,6 +6183,7 @@
|
||||
"app.startup.modules_initializer -> app.runtime.observability",
|
||||
"app.startup.modules_initializer -> app.runtime.settings",
|
||||
"app.startup.modules_initializer -> app.runtime.state",
|
||||
"app.startup.modules_initializer -> app.runtime.tasks",
|
||||
"app.startup.modules_initializer -> app.runtime.thread",
|
||||
"app.startup.modules_initializer -> app.scheduler",
|
||||
"app.startup.modules_initializer -> app.schemas",
|
||||
@@ -6461,7 +6481,7 @@
|
||||
"app.workflow.actions.transfer_file -> app.workflow",
|
||||
"app.workflow.actions.transfer_file -> app.workflow.actions"
|
||||
],
|
||||
"module_count": 798,
|
||||
"module_count": 799,
|
||||
"modules": [
|
||||
"app",
|
||||
"app.adapters",
|
||||
@@ -7153,6 +7173,7 @@
|
||||
"app.runtime.scheduling",
|
||||
"app.runtime.settings",
|
||||
"app.runtime.state",
|
||||
"app.runtime.tasks",
|
||||
"app.runtime.thread",
|
||||
"app.runtime.topology",
|
||||
"app.scheduler",
|
||||
|
||||
@@ -4947,7 +4947,7 @@
|
||||
}
|
||||
},
|
||||
"provenance": {
|
||||
"head": "fd0363711dbba96e31a27ad56b2e87c3524096a4",
|
||||
"head": "7d2d676d6f5139227050e9e71bd579943f45c0e8",
|
||||
"python_file_count": 238,
|
||||
"source_sha256": "f07ff2e8c95080cbbf6d229376464e61dc74061c1a8d371b5ea2cf00148413de"
|
||||
},
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+3
-158
@@ -1,12 +1,12 @@
|
||||
{
|
||||
"model_decorators": {
|
||||
"by_kind": {
|
||||
"async_db_query": 44,
|
||||
"async_db_query": 27,
|
||||
"async_db_update": 0,
|
||||
"db_query": 62,
|
||||
"db_query": 48,
|
||||
"db_update": 0
|
||||
},
|
||||
"count": 106,
|
||||
"count": 75,
|
||||
"methods": [
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
@@ -153,26 +153,6 @@
|
||||
"file": "app/db/models/mediaserver.py",
|
||||
"method": "MediaServerItem.get_by_server_itemid"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/message.py",
|
||||
"method": "Message.async_list_by_page"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/message.py",
|
||||
"method": "Message.async_list_sent_by_page"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/message.py",
|
||||
"method": "Message.exists_by_source"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/message.py",
|
||||
"method": "Message.list_by_page"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/passkey.py",
|
||||
@@ -193,46 +173,6 @@
|
||||
"file": "app/db/models/passkey.py",
|
||||
"method": "PassKey.get_by_id"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/site.py",
|
||||
"method": "Site.async_get_actives"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/site.py",
|
||||
"method": "Site.async_get_by_domain"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/site.py",
|
||||
"method": "Site.async_get_by_name"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/site.py",
|
||||
"method": "Site.async_list_order_by_pri"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/site.py",
|
||||
"method": "Site.get_actives"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/site.py",
|
||||
"method": "Site.get_by_domain"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/site.py",
|
||||
"method": "Site.get_domains_by_ids"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/site.py",
|
||||
"method": "Site.list_order_by_pri"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/siteuserdata.py",
|
||||
@@ -258,91 +198,6 @@
|
||||
"file": "app/db/models/siteuserdata.py",
|
||||
"method": "SiteUserData.get_latest"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.async_exists"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.async_exists_by_username"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.async_get_by"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.async_get_by_state"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.async_get_by_title"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.async_list_by_media_identity"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.async_list_by_title"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.async_list_by_type"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.async_list_by_username"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.exists"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.exists_by_username"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.get_by"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.get_by_state"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.get_by_title"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.list_by_media_identity"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.list_by_type"
|
||||
},
|
||||
{
|
||||
"decorator": "db_query",
|
||||
"file": "app/db/models/subscribe.py",
|
||||
"method": "Subscribe.list_by_username"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/subscribehistory.py",
|
||||
@@ -488,16 +343,6 @@
|
||||
"file": "app/db/models/transferpending.py",
|
||||
"method": "TransferPending.list_all"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/user.py",
|
||||
"method": "User.async_get_by_id"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/user.py",
|
||||
"method": "User.async_get_by_name"
|
||||
},
|
||||
{
|
||||
"decorator": "async_db_query",
|
||||
"file": "app/db/models/workflow.py",
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
"""API 后台任务必须进入宿主 TaskRegistry 的回归测试。"""
|
||||
|
||||
import asyncio
|
||||
from types import SimpleNamespace
|
||||
|
||||
from app.api.endpoints import message, site, subscribe, webhook
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
|
||||
|
||||
class _TaskRegistry(TaskRegistry):
|
||||
"""记录同步任务提交参数,不在端点测试中执行真实业务。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""初始化调用记录。"""
|
||||
super().__init__()
|
||||
self.calls: list[tuple] = []
|
||||
|
||||
def create_sync(self, function, *args, owner: str, **kwargs) -> None:
|
||||
"""保存函数、参数和 owner。"""
|
||||
self.calls.append((function, args, kwargs, owner))
|
||||
|
||||
|
||||
class _WebhookRequest:
|
||||
"""提供 webhook 端点读取的最小请求接口。"""
|
||||
|
||||
query_params = {"source": "jellyfin"}
|
||||
|
||||
async def body(self) -> bytes:
|
||||
"""返回最小请求体。"""
|
||||
return b"{}"
|
||||
|
||||
async def form(self) -> dict:
|
||||
"""返回空表单。"""
|
||||
return {}
|
||||
|
||||
|
||||
class _MessageRequest(_WebhookRequest):
|
||||
"""复用 webhook 请求替身,覆盖用户消息入口所需字段。"""
|
||||
|
||||
headers = {"content-type": "application/json"}
|
||||
|
||||
|
||||
class _SeerrRequest:
|
||||
"""提供 Seerr 电影订阅 webhook 所需的最小 JSON 请求。"""
|
||||
|
||||
async def json(self) -> dict:
|
||||
"""返回一个已批准的电影订阅通知。"""
|
||||
return {
|
||||
"notification_type": "MEDIA_APPROVED",
|
||||
"subject": "测试电影",
|
||||
"media": {"media_type": "movie", "tmdbId": 123},
|
||||
"request": {"requestedBy_username": "tester"},
|
||||
}
|
||||
|
||||
|
||||
def test_webhook_post_uses_task_registry() -> None:
|
||||
"""POST webhook 应登记解析任务,响应仍只表示宿主已接受。"""
|
||||
registry = _TaskRegistry()
|
||||
response = asyncio.run(
|
||||
webhook.webhook_message(registry, _WebhookRequest(), "token")
|
||||
)
|
||||
|
||||
function, args, kwargs, owner = registry.calls[0]
|
||||
assert response.success is True
|
||||
assert function is webhook.start_webhook_chain
|
||||
assert args == (b"{}", {}, {"source": "jellyfin"})
|
||||
assert kwargs == {}
|
||||
assert owner == "api.webhook.message"
|
||||
|
||||
|
||||
def test_webhook_get_uses_task_registry() -> None:
|
||||
"""GET webhook 应保留旧参数形状并进入相同 owner。"""
|
||||
registry = _TaskRegistry()
|
||||
response = asyncio.run(
|
||||
webhook.webhook_message_get(registry, _WebhookRequest(), "token")
|
||||
)
|
||||
|
||||
function, args, kwargs, owner = registry.calls[0]
|
||||
assert response.success is True
|
||||
assert function is webhook.start_webhook_chain
|
||||
assert args == (None, None, {"source": "jellyfin"})
|
||||
assert kwargs == {}
|
||||
assert owner == "api.webhook.message"
|
||||
|
||||
|
||||
def test_cookiecloud_sync_uses_task_registry(monkeypatch) -> None:
|
||||
"""CookieCloud 手工同步应登记 Scheduler E1 任务而非 Starlette 后台回调。"""
|
||||
registry = _TaskRegistry()
|
||||
scheduler = SimpleNamespace(start=lambda **_kwargs: None)
|
||||
monkeypatch.setattr(site, "Scheduler", lambda: scheduler)
|
||||
|
||||
response = asyncio.run(site.cookie_cloud_sync(registry, SimpleNamespace()))
|
||||
|
||||
function, args, kwargs, owner = registry.calls[0]
|
||||
assert response.success is True
|
||||
assert function is scheduler.start
|
||||
assert args == ()
|
||||
assert kwargs == {"job_id": "cookiecloud"}
|
||||
assert owner == "api.site.cookiecloud_sync"
|
||||
|
||||
|
||||
def test_user_message_uses_task_registry() -> None:
|
||||
"""消息入口应登记 E0 链任务并保持原始载荷。"""
|
||||
registry = _TaskRegistry()
|
||||
response = asyncio.run(message.user_message(registry, _MessageRequest(), None))
|
||||
|
||||
function, args, kwargs, owner = registry.calls[0]
|
||||
assert response.success is True
|
||||
assert function is message.start_message_chain
|
||||
assert args == (b"{}", {}, {"source": "jellyfin"})
|
||||
assert kwargs == {}
|
||||
assert owner == "api.message.user"
|
||||
|
||||
|
||||
def test_seerr_subscribe_uses_task_registry(monkeypatch) -> None:
|
||||
"""Seerr webhook 应登记订阅创建任务且保持旧参数投影。"""
|
||||
registry = _TaskRegistry()
|
||||
monkeypatch.setattr(
|
||||
subscribe,
|
||||
"get_api_runtime_config_snapshot",
|
||||
lambda: SimpleNamespace(api_token="token"),
|
||||
)
|
||||
|
||||
response = asyncio.run(
|
||||
subscribe.seerr_subscribe(_SeerrRequest(), registry, "token")
|
||||
)
|
||||
|
||||
function, args, kwargs, owner = registry.calls[0]
|
||||
assert response.success is True
|
||||
assert function is subscribe.start_subscribe_add
|
||||
assert args == ()
|
||||
assert kwargs == {
|
||||
"mtype": subscribe.MediaType.MOVIE,
|
||||
"media_source": subscribe.MediaSource.TMDB,
|
||||
"media_id": "123",
|
||||
"title": "测试电影",
|
||||
"year": "",
|
||||
"season": None,
|
||||
"username": "tester",
|
||||
}
|
||||
assert owner == "api.subscribe.seerr"
|
||||
@@ -126,8 +126,8 @@ def test_transaction_debt_baseline_is_a_model_and_oper_ratchet() -> None:
|
||||
baseline = json.loads(baseline_path.read_text(encoding="utf-8"))
|
||||
|
||||
assert baseline["schema_version"] == 1
|
||||
assert baseline["model_decorators"]["count"] == 106
|
||||
assert sum(baseline["model_decorators"]["by_kind"].values()) == 106
|
||||
assert baseline["model_decorators"]["count"] == 75
|
||||
assert sum(baseline["model_decorators"]["by_kind"].values()) == 75
|
||||
assert baseline["model_decorators"]["by_kind"]["db_update"] == 0
|
||||
assert baseline["model_decorators"]["by_kind"]["async_db_update"] == 0
|
||||
assert baseline["model_transaction_calls"] == {"count": 0, "calls": []}
|
||||
|
||||
@@ -253,6 +253,7 @@ def test_lifecycle_manifest_declares_normal_and_safe_mode_order() -> None:
|
||||
safe_names = {item["name"] for item in safe}
|
||||
|
||||
assert normal_start == [
|
||||
"后台任务登记器",
|
||||
"数据库准备",
|
||||
"HTTP 基础能力",
|
||||
"领域依赖装配",
|
||||
@@ -269,6 +270,7 @@ def test_lifecycle_manifest_declares_normal_and_safe_mode_order() -> None:
|
||||
"工作流",
|
||||
]
|
||||
assert normal_stop == [
|
||||
"后台任务登记器",
|
||||
"插件备份",
|
||||
"工作流",
|
||||
"命令服务",
|
||||
@@ -279,6 +281,7 @@ def test_lifecycle_manifest_declares_normal_and_safe_mode_order() -> None:
|
||||
"HTTP 基础能力",
|
||||
]
|
||||
assert safe_names == {
|
||||
"后台任务登记器",
|
||||
"数据库准备",
|
||||
"HTTP 基础能力",
|
||||
"领域依赖装配",
|
||||
|
||||
@@ -7,7 +7,9 @@ from app.runtime.extensions.module.contracts import (
|
||||
ModuleErrorPolicy,
|
||||
ModuleExecutionMode,
|
||||
ModuleResultAggregation,
|
||||
ModuleResultShape,
|
||||
diagnose_module_callable,
|
||||
diagnose_module_result,
|
||||
get_module_method_contract,
|
||||
is_explicit_module_method,
|
||||
list_explicit_module_contracts,
|
||||
@@ -116,3 +118,22 @@ def test_signature_diagnostics_accept_keyword_compatibility_provider() -> None:
|
||||
return kwargs
|
||||
|
||||
assert diagnose_module_callable("snapshot_storage", compatible_provider) == ()
|
||||
|
||||
|
||||
def test_result_diagnostics_check_only_enabled_basic_shapes() -> None:
|
||||
"""高频方法检查基础结果形状,业务对象合同仍留给逐族适配器。"""
|
||||
assert get_module_method_contract("list_files").result_shape is ModuleResultShape.LIST
|
||||
assert diagnose_module_result("list_files", [object()]) == ()
|
||||
assert diagnose_module_result("list_files", None) == ()
|
||||
assert diagnose_module_result("list_files", "legacy-value") == (
|
||||
"unexpected-result:list:str",
|
||||
)
|
||||
assert diagnose_module_result("mediaserver_play_url", "https://example.test") == ()
|
||||
assert diagnose_module_result("mediaserver_play_url", 7) == (
|
||||
"unexpected-result:string:int",
|
||||
)
|
||||
|
||||
|
||||
def test_unknown_plugin_result_keeps_unchecked_legacy_compatibility() -> None:
|
||||
"""未知第三方方法的任意返回值继续不做结果形状诊断。"""
|
||||
assert diagnose_module_result("third_party_custom_method", object()) == ()
|
||||
|
||||
@@ -14,6 +14,7 @@ from app.api.endpoints.system import sync_plugin_market_from_wiki
|
||||
from app.application.plugin.config import PluginConfigCommand
|
||||
from app.runtime.config import settings
|
||||
from app.runtime.extensions.plugin_manager import PluginManager
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
from app.schemas.event import PluginDataResetEventData
|
||||
from app.schemas.plugin import PluginInstance, PluginRuntimeStatus
|
||||
from app.schemas.types import ChainEventType, SystemConfigKey
|
||||
@@ -312,8 +313,8 @@ def test_plugin_releases_force_uses_cached_release_response_and_schedules_refres
|
||||
plugin_helper.async_get_plugin_release_versions = fake_releases
|
||||
scheduled = []
|
||||
|
||||
def fake_schedule(plugin_id, repo_url):
|
||||
scheduled.append((plugin_id, repo_url))
|
||||
def fake_schedule(plugin_id, repo_url, task_registry):
|
||||
scheduled.append((plugin_id, repo_url, task_registry))
|
||||
|
||||
with (
|
||||
patch("app.api.endpoints.plugin.PluginManager", return_value=plugin_manager),
|
||||
@@ -324,7 +325,12 @@ def test_plugin_releases_force_uses_cached_release_response_and_schedules_refres
|
||||
|
||||
assert result["release_supported"] is True
|
||||
assert fresh_states == [False]
|
||||
assert scheduled == [("DemoPlugin", "https://github.com/demo/plugins")]
|
||||
assert len(scheduled) == 1
|
||||
assert scheduled[0][:2] == (
|
||||
"DemoPlugin",
|
||||
"https://github.com/demo/plugins",
|
||||
)
|
||||
assert isinstance(scheduled[0][2], TaskRegistry)
|
||||
plugin_helper.async_has_plugin_release_cache.assert_awaited_once_with(
|
||||
"https://github.com/demo/plugins"
|
||||
)
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
"""进程内后台任务登记与关停语义测试。"""
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from app.runtime.tasks import TaskRegistry
|
||||
|
||||
|
||||
def test_task_registry_removes_completed_task() -> None:
|
||||
"""正常完成的任务应自动退出登记表,避免长期持有请求对象。"""
|
||||
|
||||
async def scenario() -> None:
|
||||
registry = TaskRegistry()
|
||||
release = asyncio.Event()
|
||||
|
||||
async def worker() -> None:
|
||||
"""等待测试释放信号。"""
|
||||
await release.wait()
|
||||
|
||||
task = registry.create(worker(), owner="test.completed")
|
||||
assert [record.owner for record in registry.records] == ["test.completed"]
|
||||
|
||||
release.set()
|
||||
await task
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert registry.records == ()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_task_registry_cancels_tasks_and_rejects_late_registration() -> None:
|
||||
"""关停应取消存量任务,并拒绝在资源释放阶段继续产生新任务。"""
|
||||
|
||||
async def scenario() -> None:
|
||||
registry = TaskRegistry()
|
||||
started = asyncio.Event()
|
||||
cancelled = asyncio.Event()
|
||||
|
||||
async def worker() -> None:
|
||||
"""记录任务收到取消信号。"""
|
||||
started.set()
|
||||
try:
|
||||
await asyncio.Event().wait()
|
||||
except asyncio.CancelledError:
|
||||
cancelled.set()
|
||||
raise
|
||||
|
||||
task = registry.create(worker(), owner="test.shutdown")
|
||||
await started.wait()
|
||||
await registry.shutdown(timeout_seconds=1.0)
|
||||
|
||||
assert task.cancelled()
|
||||
assert cancelled.is_set()
|
||||
assert registry.records == ()
|
||||
|
||||
async def late_worker() -> None:
|
||||
"""模拟关停开始后到达的晚任务。"""
|
||||
|
||||
with pytest.raises(RuntimeError, match="正在关闭"):
|
||||
registry.create(late_worker(), owner="test.late")
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_task_registry_runs_sync_function_and_tracks_until_completion() -> None:
|
||||
"""同步任务应在线程池执行,并在真实完成前保留 owner 记录。"""
|
||||
|
||||
async def scenario() -> None:
|
||||
registry = TaskRegistry()
|
||||
release = asyncio.Event()
|
||||
|
||||
def worker(value: int) -> int:
|
||||
"""返回传入值,验证参数和结果没有被登记器改写。"""
|
||||
return value
|
||||
|
||||
task = registry.create_sync(worker, 7, owner="test.sync")
|
||||
assert [record.owner for record in registry.records] == ["test.sync"]
|
||||
assert await task == 7
|
||||
await asyncio.sleep(0)
|
||||
assert registry.records == ()
|
||||
|
||||
release.set()
|
||||
|
||||
asyncio.run(scenario())
|
||||
Reference in New Issue
Block a user