mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-10 18:06:48 +08:00
feat(subscribe): expose governed execution status
This commit is contained in:
@@ -34,6 +34,7 @@ from app.application.subscription.delete import (
|
||||
from app.application.subscription.delete import (
|
||||
DeleteSubscribeCommand,
|
||||
)
|
||||
from app.application.subscription.execution import SubscriptionSearchRepository
|
||||
from app.application.subscription.identity import (
|
||||
DeleteSubscriptionsByIdentityCommand,
|
||||
)
|
||||
@@ -42,6 +43,7 @@ from app.application.subscription.mutation import (
|
||||
)
|
||||
from app.application.subscription.query import SubscriptionQueryService
|
||||
from app.application.subscription.search import SearchSubscriptionsCommand
|
||||
from app.application.subscription.status import SubscriptionExecutionStatusService
|
||||
from app.application.subscription.write import (
|
||||
SubscriptionBatchWritePort,
|
||||
)
|
||||
@@ -176,6 +178,27 @@ def get_subscription_query_service(
|
||||
)
|
||||
|
||||
|
||||
def get_subscription_execution_status_service(
|
||||
db: AsyncSession = Depends(get_async_session),
|
||||
runtime: HostRuntime = Depends(get_host_runtime),
|
||||
) -> SubscriptionExecutionStatusService:
|
||||
"""组装请求级订阅执行状态投影服务。"""
|
||||
factory = runtime.subscription.execution_status_repository
|
||||
if factory is None:
|
||||
raise RuntimeError("订阅执行状态仓储未注册")
|
||||
repository = factory(db)
|
||||
return SubscriptionExecutionStatusService(repository=repository) # type: ignore[arg-type]
|
||||
|
||||
|
||||
def get_subscription_search_repository(
|
||||
runtime: HostRuntime = Depends(get_host_runtime),
|
||||
) -> SubscriptionSearchRepository:
|
||||
"""返回宿主组合根持有的订阅搜索队列端口。"""
|
||||
if runtime.subscription.search_repository is None:
|
||||
raise RuntimeError("订阅搜索队列未注册")
|
||||
return cast(SubscriptionSearchRepository, runtime.subscription.search_repository)
|
||||
|
||||
|
||||
def get_subscription_mutation_service(
|
||||
repository_port: SessionSubscriptionPort = Depends(get_subscription_repository),
|
||||
history_repository: SubscriptionHistoryStagingPort = Depends(
|
||||
|
||||
@@ -22,8 +22,10 @@ from app.api.dependencies.subscription import (
|
||||
get_delete_subscribe_command,
|
||||
get_delete_subscriptions_by_identity_command,
|
||||
get_search_subscriptions_command,
|
||||
get_subscription_execution_status_service,
|
||||
get_subscription_mutation_service,
|
||||
get_subscription_query_service,
|
||||
get_subscription_search_repository,
|
||||
)
|
||||
from app.api.principal import ApiPrincipal
|
||||
from app.api.response import (
|
||||
@@ -44,6 +46,7 @@ from app.application.subscription.delete import (
|
||||
DeleteSubscribeCommand,
|
||||
SubscribeDeletionActor,
|
||||
)
|
||||
from app.application.subscription.execution import SubscriptionSearchRepository
|
||||
from app.application.subscription.identity import (
|
||||
DeleteSubscriptionsByIdentityCommand,
|
||||
)
|
||||
@@ -56,6 +59,7 @@ from app.application.subscription.search import (
|
||||
SearchSubscriptionsCommand,
|
||||
SubscribeSearchActor,
|
||||
)
|
||||
from app.application.subscription.status import SubscriptionExecutionStatusService
|
||||
from app.chain.subscribe.facade import SubscribeChain
|
||||
from app.domain.context import MediaInfo
|
||||
from app.domain.metainfo import MetaInfo
|
||||
@@ -68,6 +72,8 @@ from app.schemas.subscribe import SubscrbieInfo as _SchemaSubscrbieInfo
|
||||
from app.schemas.subscribe import SubscribeDeletionResult as _SchemaSubscribeDeletionResult
|
||||
from app.schemas.subscribe import SubscribeShare as _SchemaSubscribeShare
|
||||
from app.schemas.subscribe import SubscribeShareStatistics as _SchemaSubscribeShareStatistics
|
||||
from app.schemas.subscribe import SubscriptionBatchStatus as _SchemaSubscriptionBatchStatus
|
||||
from app.schemas.subscribe import SubscriptionExecutionStatus as _SchemaSubscriptionExecutionStatus
|
||||
from app.schemas.token import TokenPayload as _SchemaTokenPayload
|
||||
from app.schemas.types import (
|
||||
MUSIC_ENTITY_ALBUM,
|
||||
@@ -82,6 +88,34 @@ from app.schemas.workflow import Subscribe as _SchemaSubscribe
|
||||
router = ResponseAPIRouter()
|
||||
|
||||
|
||||
async def _attach_execution_status(
|
||||
subscribes: list[_SchemaSubscribe],
|
||||
status_service: object,
|
||||
) -> list[_SchemaSubscribe]:
|
||||
"""批量附加当前执行状态,保持无执行记录时的旧响应形状。"""
|
||||
loader = getattr(status_service, "for_subscriptions", None)
|
||||
if not callable(loader):
|
||||
return subscribes
|
||||
statuses = await loader(
|
||||
tuple(item.id for item in subscribes if item.id is not None)
|
||||
)
|
||||
for subscribe in subscribes:
|
||||
if subscribe.id is not None and (status := statuses.get(subscribe.id)) is not None:
|
||||
subscribe.execution_status = _SchemaSubscriptionExecutionStatus.model_validate(status)
|
||||
return subscribes
|
||||
|
||||
|
||||
async def _accessible_subscription_ids(
|
||||
query: SubscriptionQueryService,
|
||||
current_user: ApiPrincipal,
|
||||
) -> Optional[set[int]]:
|
||||
"""返回普通用户可访问订阅 ID;超级用户以 None 表示不限制。"""
|
||||
if current_user.is_superuser:
|
||||
return None
|
||||
subscribes = await query.list_public(current_user.name)
|
||||
return {item.id for item in subscribes if item.id is not None}
|
||||
|
||||
|
||||
def start_subscribe_add(
|
||||
title: str,
|
||||
year: str,
|
||||
@@ -162,6 +196,9 @@ def matches_subscribe_music_type(
|
||||
async def read_subscribes(
|
||||
response: Response = None,
|
||||
query: SubscriptionQueryService = Depends(get_subscription_query_service),
|
||||
status_service: SubscriptionExecutionStatusService = Depends(
|
||||
get_subscription_execution_status_service
|
||||
),
|
||||
current_user: ApiPrincipal = Depends(get_current_active_user_async),
|
||||
page: CompatiblePageParam = None,
|
||||
count: CompatibleCountParam = None,
|
||||
@@ -175,7 +212,8 @@ async def read_subscribes(
|
||||
response.headers[COLLECTION_TOTAL_HEADER] = str(
|
||||
await query.count_public(username)
|
||||
)
|
||||
return await query.list_public(username, page=page, count=count)
|
||||
subscribes = await query.list_public(username, page=page, count=count)
|
||||
return await _attach_execution_status(subscribes, status_service)
|
||||
|
||||
|
||||
@router.get(
|
||||
@@ -187,6 +225,9 @@ async def read_subscribes(
|
||||
async def list_subscribes(
|
||||
response: Response = None,
|
||||
query: SubscriptionQueryService = Depends(get_subscription_query_service),
|
||||
status_service: SubscriptionExecutionStatusService = Depends(
|
||||
get_subscription_execution_status_service
|
||||
),
|
||||
_: Annotated[str, Depends(verify_apitoken)] = None,
|
||||
page: CompatiblePageParam = None,
|
||||
count: CompatibleCountParam = None,
|
||||
@@ -197,7 +238,8 @@ async def list_subscribes(
|
||||
page, count = resolve_compatible_pagination(page, count)
|
||||
if response is not None:
|
||||
response.headers[COLLECTION_TOTAL_HEADER] = str(await query.count_public())
|
||||
return await query.list_public(page=page, count=count)
|
||||
subscribes = await query.list_public(page=page, count=count)
|
||||
return await _attach_execution_status(subscribes, status_service)
|
||||
|
||||
|
||||
@router.post(
|
||||
@@ -460,6 +502,83 @@ async def search_subscribe(
|
||||
return _SchemaResponse(success=True)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/execution/batches",
|
||||
summary="查询订阅搜索批次状态",
|
||||
response_model=List[_SchemaSubscriptionBatchStatus],
|
||||
)
|
||||
async def list_subscription_execution_batches(
|
||||
limit: int = 10,
|
||||
status_service: SubscriptionExecutionStatusService = Depends(
|
||||
get_subscription_execution_status_service
|
||||
),
|
||||
query: SubscriptionQueryService = Depends(get_subscription_query_service),
|
||||
current_user: ApiPrincipal = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""返回当前用户完整可见的最近搜索批次。"""
|
||||
accessible_ids = await _accessible_subscription_ids(query, current_user)
|
||||
batches = await status_service.list_batches(
|
||||
accessible_subscription_ids=accessible_ids,
|
||||
limit=limit,
|
||||
)
|
||||
return [_SchemaSubscriptionBatchStatus.model_validate(batch) for batch in batches]
|
||||
|
||||
|
||||
@router.get(
|
||||
"/execution/batches/{batch_id}",
|
||||
summary="查询订阅搜索批次",
|
||||
response_model=_SchemaSubscriptionBatchStatus,
|
||||
)
|
||||
async def get_subscription_execution_batch(
|
||||
batch_id: str,
|
||||
status_service: SubscriptionExecutionStatusService = Depends(
|
||||
get_subscription_execution_status_service
|
||||
),
|
||||
query: SubscriptionQueryService = Depends(get_subscription_query_service),
|
||||
current_user: ApiPrincipal = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""按稳定 ID 返回当前用户可访问的搜索批次。"""
|
||||
accessible_ids = await _accessible_subscription_ids(query, current_user)
|
||||
batch = await status_service.get_batch(
|
||||
batch_id,
|
||||
accessible_subscription_ids=accessible_ids,
|
||||
)
|
||||
if batch is None:
|
||||
raise HTTPException(status_code=404, detail="订阅搜索批次不存在")
|
||||
return _SchemaSubscriptionBatchStatus.model_validate(batch)
|
||||
|
||||
|
||||
@router.put(
|
||||
"/execution/batches/{batch_id}/cancel",
|
||||
summary="取消订阅搜索批次",
|
||||
response_model=_SchemaResponse[None],
|
||||
)
|
||||
async def cancel_subscription_execution_batch(
|
||||
batch_id: str,
|
||||
status_service: SubscriptionExecutionStatusService = Depends(
|
||||
get_subscription_execution_status_service
|
||||
),
|
||||
query: SubscriptionQueryService = Depends(get_subscription_query_service),
|
||||
search_repository: SubscriptionSearchRepository = Depends(
|
||||
get_subscription_search_repository
|
||||
),
|
||||
current_user: ApiPrincipal = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""在权限校验后请求取消尚未越过下载副作用边界的任务。"""
|
||||
accessible_ids = await _accessible_subscription_ids(query, current_user)
|
||||
batch = await status_service.get_batch(
|
||||
batch_id,
|
||||
accessible_subscription_ids=accessible_ids,
|
||||
)
|
||||
if batch is None:
|
||||
return _SchemaResponse(success=False, message="订阅搜索批次不存在")
|
||||
cancelled = await run_in_threadpool(search_repository.request_cancel, batch_id)
|
||||
return _SchemaResponse(
|
||||
success=bool(cancelled),
|
||||
message="" if cancelled else "订阅搜索批次已结束或无法取消",
|
||||
)
|
||||
|
||||
|
||||
@router.delete("/media/{media_id}", summary="删除订阅", response_model=_SchemaResponse[None])
|
||||
async def delete_subscribe_by_media_identity(
|
||||
media_id: str,
|
||||
@@ -692,6 +811,9 @@ async def user_subscribes(
|
||||
username: str,
|
||||
response: Response = None,
|
||||
query: SubscriptionQueryService = Depends(get_subscription_query_service),
|
||||
status_service: SubscriptionExecutionStatusService = Depends(
|
||||
get_subscription_execution_status_service
|
||||
),
|
||||
current_user: ApiPrincipal = Depends(get_current_active_user_async),
|
||||
page: CompatiblePageParam = None,
|
||||
count: CompatibleCountParam = None,
|
||||
@@ -706,7 +828,8 @@ async def user_subscribes(
|
||||
response.headers[COLLECTION_TOTAL_HEADER] = str(
|
||||
await query.count_public(username)
|
||||
)
|
||||
return await query.list_public(username, page=page, count=count)
|
||||
subscribes = await query.list_public(username, page=page, count=count)
|
||||
return await _attach_execution_status(subscribes, status_service)
|
||||
|
||||
|
||||
@router.get(
|
||||
|
||||
@@ -37,6 +37,7 @@ class SearchTaskSnapshot:
|
||||
priority: int
|
||||
position: int
|
||||
state: str
|
||||
phase: str
|
||||
attempt_count: int
|
||||
cancel_requested: bool
|
||||
lease_token: Optional[str]
|
||||
@@ -46,6 +47,7 @@ class SearchTaskSnapshot:
|
||||
started_at: Optional[str] = None
|
||||
finished_at: Optional[str] = None
|
||||
last_error: Optional[str] = None
|
||||
current_site_id: Optional[int] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
@@ -86,6 +88,17 @@ class SubscriptionSearchRepository(Protocol):
|
||||
"""以租约令牌收口任务,并推进所属批次聚合状态。"""
|
||||
...
|
||||
|
||||
def update_task_phase(
|
||||
self,
|
||||
*,
|
||||
task_id: str,
|
||||
lease_token: str,
|
||||
phase: str,
|
||||
current_site_id: Optional[int] = None,
|
||||
) -> bool:
|
||||
"""以当前租约令牌更新用户可见阶段和正在处理的站点。"""
|
||||
...
|
||||
|
||||
def release_task(
|
||||
self,
|
||||
*,
|
||||
|
||||
@@ -84,6 +84,7 @@ class SubscriptionSiteBudget:
|
||||
random_uniform: Callable[[float, float], float] = random.uniform,
|
||||
sleeper: Callable[[float], None] = time.sleep,
|
||||
clock: Callable[[], datetime] = _utc_now,
|
||||
phase_changed: Optional[Callable[[str, Optional[int]], None]] = None,
|
||||
) -> None:
|
||||
"""保存持久化端口及可注入的时钟、随机数和等待实现。"""
|
||||
self._repository = repository
|
||||
@@ -96,6 +97,7 @@ class SubscriptionSiteBudget:
|
||||
self._random_uniform = random_uniform
|
||||
self._sleeper = sleeper
|
||||
self._clock = clock
|
||||
self._phase_changed = phase_changed
|
||||
|
||||
def acquire(self, site_id: int) -> SiteBudgetClaim:
|
||||
"""循环认领指定站点,并在每秒边界检查取消与停机。"""
|
||||
@@ -108,7 +110,9 @@ class SubscriptionSiteBudget:
|
||||
lease_seconds=self._lease_seconds,
|
||||
)
|
||||
if claim.acquired:
|
||||
self._report_phase("searching", site_id)
|
||||
return claim
|
||||
self._report_phase("waiting_site_budget", site_id)
|
||||
retry_at = datetime.fromisoformat(claim.retry_at)
|
||||
remaining = max(0.0, (retry_at - self._clock()).total_seconds())
|
||||
if remaining > max(0.0, deadline - time.monotonic()):
|
||||
@@ -118,6 +122,11 @@ class SubscriptionSiteBudget:
|
||||
)
|
||||
self._sleeper(min(max(remaining, 0.05), 1.0))
|
||||
|
||||
def _report_phase(self, phase: str, site_id: Optional[int]) -> None:
|
||||
"""向任务所有者报告不改变预算语义的业务阶段。"""
|
||||
if self._phase_changed:
|
||||
self._phase_changed(phase, site_id)
|
||||
|
||||
def finish(self, claim: SiteBudgetClaim, observation: SiteSearchObservation) -> bool:
|
||||
"""依据调用结果计算随机间隔或错误冷却并释放租约。"""
|
||||
if not claim.lease_token:
|
||||
|
||||
@@ -0,0 +1,269 @@
|
||||
"""订阅执行状态的业务投影与访问范围治理。"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Protocol
|
||||
|
||||
from app.application.download.admission import SubscriptionDownloadSnapshot
|
||||
from app.application.subscription.execution import SearchBatchSnapshot, SearchTaskSnapshot
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SubscriptionExecutionStatus:
|
||||
"""一个订阅跨搜索与下载账本合并后的用户可见状态。"""
|
||||
|
||||
state: str
|
||||
phase: str
|
||||
updated_at: str
|
||||
source: Optional[str] = None
|
||||
batch_id: Optional[str] = None
|
||||
task_id: Optional[str] = None
|
||||
current_site_id: Optional[int] = None
|
||||
error: Optional[str] = None
|
||||
can_cancel: bool = False
|
||||
can_retry: bool = False
|
||||
requires_reconciliation: bool = False
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SubscriptionBatchStatus:
|
||||
"""订阅搜索批次的进度、当前工作和操作能力。"""
|
||||
|
||||
batch_id: str
|
||||
source: str
|
||||
state: str
|
||||
phase: str
|
||||
total_count: int
|
||||
processed_count: int
|
||||
finished_count: int
|
||||
failed_count: int
|
||||
cancelled_count: int
|
||||
created_at: str
|
||||
updated_at: str
|
||||
current_subscription_id: Optional[int] = None
|
||||
current_site_id: Optional[int] = None
|
||||
error: Optional[str] = None
|
||||
can_cancel: bool = False
|
||||
|
||||
|
||||
class SubscriptionExecutionReadRepository(Protocol):
|
||||
"""请求级读取搜索任务、批次与下载提交事实的端口。"""
|
||||
|
||||
async def latest_search_tasks(
|
||||
self,
|
||||
subscription_ids: tuple[int, ...],
|
||||
) -> dict[int, SearchTaskSnapshot]:
|
||||
"""返回每条订阅最近更新的搜索任务。"""
|
||||
...
|
||||
|
||||
async def latest_download_submissions(
|
||||
self,
|
||||
subscription_ids: tuple[int, ...],
|
||||
) -> dict[int, SubscriptionDownloadSnapshot]:
|
||||
"""返回每条订阅最近更新的下载提交。"""
|
||||
...
|
||||
|
||||
async def list_batches(
|
||||
self,
|
||||
*,
|
||||
limit: int,
|
||||
) -> list[SearchBatchSnapshot]:
|
||||
"""返回最近更新的搜索批次。"""
|
||||
...
|
||||
|
||||
async def get_batch(self, batch_id: str) -> Optional[SearchBatchSnapshot]:
|
||||
"""按稳定 ID 返回一个搜索批次。"""
|
||||
...
|
||||
|
||||
async def list_batch_tasks(self, batch_id: str) -> list[SearchTaskSnapshot]:
|
||||
"""按稳定位置返回批次任务。"""
|
||||
...
|
||||
|
||||
|
||||
class SubscriptionExecutionStatusService:
|
||||
"""把搜索队列与下载幂等账本投影为稳定业务状态。"""
|
||||
|
||||
_DOWNLOAD_STATES = {
|
||||
"submitting": "submitting",
|
||||
"accepted": "accepted",
|
||||
"succeeded": "completed",
|
||||
"retryable": "retryable",
|
||||
"reconcile_required": "reconcile_required",
|
||||
"cancelled": "cancelled",
|
||||
}
|
||||
_ACTIVE_STATES = {
|
||||
"queued",
|
||||
"running",
|
||||
"matching",
|
||||
"searching",
|
||||
"waiting_site_budget",
|
||||
"preparing",
|
||||
"submitting",
|
||||
"accepted",
|
||||
"cancelling",
|
||||
}
|
||||
_DOWNLOAD_OVERRIDE_STATES = {
|
||||
"submitting",
|
||||
"accepted",
|
||||
"retryable",
|
||||
"reconcile_required",
|
||||
}
|
||||
|
||||
def __init__(self, repository: SubscriptionExecutionReadRepository) -> None:
|
||||
"""保存请求会话绑定的状态读取端口。"""
|
||||
self._repository = repository
|
||||
|
||||
async def for_subscriptions(
|
||||
self,
|
||||
subscription_ids: tuple[int, ...],
|
||||
) -> dict[int, SubscriptionExecutionStatus]:
|
||||
"""批量投影订阅状态,避免列表接口逐条查询。"""
|
||||
ids = tuple(dict.fromkeys(subscription_ids))
|
||||
if not ids:
|
||||
return {}
|
||||
tasks = await self._repository.latest_search_tasks(ids)
|
||||
downloads = await self._repository.latest_download_submissions(ids)
|
||||
result: dict[int, SubscriptionExecutionStatus] = {}
|
||||
for subscription_id in ids:
|
||||
task = tasks.get(subscription_id)
|
||||
download = downloads.get(subscription_id)
|
||||
if download and self._download_wins(task, download):
|
||||
result[subscription_id] = self._from_download(download, task)
|
||||
elif task:
|
||||
result[subscription_id] = self._from_task(task)
|
||||
return result
|
||||
|
||||
async def list_batches(
|
||||
self,
|
||||
*,
|
||||
accessible_subscription_ids: Optional[set[int]],
|
||||
limit: int = 10,
|
||||
) -> list[SubscriptionBatchStatus]:
|
||||
"""列出访问范围完整覆盖的最近批次。"""
|
||||
batches = await self._repository.list_batches(limit=max(1, min(limit, 50)))
|
||||
result = []
|
||||
for batch in batches:
|
||||
tasks = await self._repository.list_batch_tasks(batch.batch_id)
|
||||
if not self._can_access_tasks(tasks, accessible_subscription_ids):
|
||||
continue
|
||||
result.append(self._from_batch(batch, tasks))
|
||||
return result
|
||||
|
||||
async def get_batch(
|
||||
self,
|
||||
batch_id: str,
|
||||
*,
|
||||
accessible_subscription_ids: Optional[set[int]],
|
||||
) -> Optional[SubscriptionBatchStatus]:
|
||||
"""读取一个访问范围完整覆盖的批次。"""
|
||||
batch = await self._repository.get_batch(batch_id)
|
||||
if batch is None:
|
||||
return None
|
||||
tasks = await self._repository.list_batch_tasks(batch_id)
|
||||
if not self._can_access_tasks(tasks, accessible_subscription_ids):
|
||||
return None
|
||||
return self._from_batch(batch, tasks)
|
||||
|
||||
@classmethod
|
||||
def _download_wins(
|
||||
cls,
|
||||
task: Optional[SearchTaskSnapshot],
|
||||
download: SubscriptionDownloadSnapshot,
|
||||
) -> bool:
|
||||
"""下载风险状态优先,其余事实按更新时间选择。"""
|
||||
if download.state in cls._DOWNLOAD_OVERRIDE_STATES:
|
||||
return True
|
||||
return task is None or download.updated_at >= task.updated_at
|
||||
|
||||
@classmethod
|
||||
def _from_task(cls, task: SearchTaskSnapshot) -> SubscriptionExecutionStatus:
|
||||
"""把搜索任务状态归一为稳定业务词汇。"""
|
||||
if task.cancel_requested and task.state == "running":
|
||||
state = phase = "cancelling"
|
||||
elif task.state == "running":
|
||||
state = phase = task.phase or "running"
|
||||
else:
|
||||
state = phase = task.state
|
||||
return SubscriptionExecutionStatus(
|
||||
state=state,
|
||||
phase=phase,
|
||||
source=task.source,
|
||||
batch_id=task.batch_id,
|
||||
task_id=task.task_id,
|
||||
current_site_id=task.current_site_id,
|
||||
updated_at=task.updated_at,
|
||||
error=cls._safe_error(task.last_error),
|
||||
can_cancel=state in cls._ACTIVE_STATES,
|
||||
can_retry=state == "failed",
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _from_download(
|
||||
cls,
|
||||
download: SubscriptionDownloadSnapshot,
|
||||
task: Optional[SearchTaskSnapshot],
|
||||
) -> SubscriptionExecutionStatus:
|
||||
"""把下载提交账本状态投影为业务状态并保留搜索来源。"""
|
||||
state = cls._DOWNLOAD_STATES.get(download.state, download.state)
|
||||
return SubscriptionExecutionStatus(
|
||||
state=state,
|
||||
phase=state,
|
||||
source=task.source if task else None,
|
||||
batch_id=task.batch_id if task else None,
|
||||
task_id=download.task_id or (task.task_id if task else None),
|
||||
current_site_id=task.current_site_id if task else None,
|
||||
updated_at=download.updated_at,
|
||||
error=cls._safe_error(download.last_error),
|
||||
can_cancel=state == "submitting",
|
||||
can_retry=state == "retryable",
|
||||
requires_reconciliation=state == "reconcile_required",
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _from_batch(
|
||||
cls,
|
||||
batch: SearchBatchSnapshot,
|
||||
tasks: list[SearchTaskSnapshot],
|
||||
) -> SubscriptionBatchStatus:
|
||||
"""组合批次计数与当前运行任务。"""
|
||||
current = next((task for task in tasks if task.state == "running"), None)
|
||||
if current is None:
|
||||
current = next((task for task in tasks if task.state == "queued"), None)
|
||||
processed = batch.finished_count + batch.failed_count + batch.cancelled_count
|
||||
phase = current.phase if current else batch.state
|
||||
return SubscriptionBatchStatus(
|
||||
batch_id=batch.batch_id,
|
||||
source=batch.source,
|
||||
state=batch.state,
|
||||
phase=phase,
|
||||
total_count=batch.total_count,
|
||||
processed_count=processed,
|
||||
finished_count=batch.finished_count,
|
||||
failed_count=batch.failed_count,
|
||||
cancelled_count=batch.cancelled_count,
|
||||
current_subscription_id=current.subscription_id if current else None,
|
||||
current_site_id=current.current_site_id if current else None,
|
||||
created_at=batch.created_at,
|
||||
updated_at=batch.updated_at,
|
||||
error=cls._safe_error(batch.last_error),
|
||||
can_cancel=batch.state in {"queued", "running", "cancelling"}
|
||||
and not batch.cancel_requested,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _can_access_tasks(
|
||||
tasks: list[SearchTaskSnapshot],
|
||||
accessible_subscription_ids: Optional[set[int]],
|
||||
) -> bool:
|
||||
"""超级用户不限制;普通用户必须拥有批次内全部订阅。"""
|
||||
if accessible_subscription_ids is None:
|
||||
return True
|
||||
return bool(tasks) and all(
|
||||
task.subscription_id in accessible_subscription_ids for task in tasks
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _safe_error(error: Optional[str]) -> Optional[str]:
|
||||
"""压平并限制内部错误文本,避免把堆栈或超长响应暴露给界面。"""
|
||||
if not error:
|
||||
return None
|
||||
return " ".join(str(error).split())[:500]
|
||||
@@ -300,18 +300,26 @@ class SubscribeSearchOwner(_SubscribeOwnerBase):
|
||||
)
|
||||
continue
|
||||
current = subscribe
|
||||
phase_changed = partial(
|
||||
self._update_search_task_phase,
|
||||
queue,
|
||||
task_id,
|
||||
task.lease_token,
|
||||
)
|
||||
searchchain.configure_subscription_site_budget(
|
||||
SubscriptionSiteBudget(
|
||||
repository=queue,
|
||||
owner=f"{owner}:{task_id}",
|
||||
cancelled=cancelled,
|
||||
stop_state=getattr(self, "stop_state", runtime_stop_state),
|
||||
phase_changed=phase_changed,
|
||||
)
|
||||
)
|
||||
self._subscription_download_task_id = task_id
|
||||
self._subscription_download_cancelled = cancelled
|
||||
self._subscription_download_crossed_boundary = False
|
||||
self._subscription_download_mark_started = self._mark_subscription_download_started
|
||||
self._subscription_execution_phase = phase_changed
|
||||
try:
|
||||
current = self._process_search_subscription(subscribe, searchchain)
|
||||
if queue.is_cancel_requested(task.task_id):
|
||||
@@ -353,6 +361,7 @@ class SubscribeSearchOwner(_SubscribeOwnerBase):
|
||||
delattr(self, "_subscription_download_task_id")
|
||||
delattr(self, "_subscription_download_cancelled")
|
||||
delattr(self, "_subscription_download_mark_started")
|
||||
delattr(self, "_subscription_execution_phase")
|
||||
self._subscription_download_crossed_boundary = False
|
||||
searchchain.configure_subscription_site_budget(None)
|
||||
if current and current.state == "N":
|
||||
@@ -383,6 +392,25 @@ class SubscribeSearchOwner(_SubscribeOwnerBase):
|
||||
def _mark_subscription_download_started(self) -> None:
|
||||
"""记录当前搜索任务已提交或复用了真实下载结果。"""
|
||||
self._subscription_download_crossed_boundary = True
|
||||
phase_changed = getattr(self, "_subscription_execution_phase", None)
|
||||
if phase_changed:
|
||||
phase_changed("submitting", None)
|
||||
|
||||
@staticmethod
|
||||
def _update_search_task_phase(
|
||||
queue: SubscriptionSearchRepository,
|
||||
task_id: str,
|
||||
lease_token: str,
|
||||
phase: str,
|
||||
current_site_id: Optional[int] = None,
|
||||
) -> None:
|
||||
"""以当前任务租约持久化业务阶段,过期执行者不得覆盖新状态。"""
|
||||
queue.update_task_phase(
|
||||
task_id=task_id,
|
||||
lease_token=lease_token,
|
||||
phase=phase,
|
||||
current_site_id=current_site_id,
|
||||
)
|
||||
|
||||
def resume_search_queue(
|
||||
self,
|
||||
@@ -537,6 +565,9 @@ class SubscribeSearchOwner(_SubscribeOwnerBase):
|
||||
if subscribe.best_version
|
||||
else SystemConfigKey.SubscribeFilterRuleGroups
|
||||
)
|
||||
phase_changed = getattr(self, "_subscription_execution_phase", None)
|
||||
if phase_changed:
|
||||
phase_changed("searching", None)
|
||||
contexts = searchchain.process(
|
||||
mediainfo=mediainfo,
|
||||
keyword=subscribe.keyword,
|
||||
@@ -564,6 +595,8 @@ class SubscribeSearchOwner(_SubscribeOwnerBase):
|
||||
self.finish_subscribe_or_not(subscribe=subscribe, meta=meta, mediainfo=mediainfo, lefts=no_exists)
|
||||
self._raise_site_budget_failures(site_budget_failures)
|
||||
return subscribe
|
||||
if phase_changed:
|
||||
phase_changed("preparing", None)
|
||||
downloads, lefts = self._SubscribeChain__download_best_version_with_full_pack_first(
|
||||
contexts=matched,
|
||||
no_exists=no_exists,
|
||||
|
||||
@@ -51,6 +51,7 @@ def _task(record: SubscriptionSearchTask) -> SearchTaskSnapshot:
|
||||
priority=record.priority,
|
||||
position=record.position,
|
||||
state=record.state,
|
||||
phase=record.phase,
|
||||
attempt_count=record.attempt_count,
|
||||
cancel_requested=bool(record.cancel_requested),
|
||||
lease_token=record.lease_token,
|
||||
@@ -60,6 +61,7 @@ def _task(record: SubscriptionSearchTask) -> SearchTaskSnapshot:
|
||||
started_at=record.started_at,
|
||||
finished_at=record.finished_at,
|
||||
last_error=record.last_error,
|
||||
current_site_id=record.current_site_id,
|
||||
)
|
||||
|
||||
|
||||
@@ -140,6 +142,24 @@ class TransactionalSubscriptionSearchRepository:
|
||||
)
|
||||
)
|
||||
|
||||
def update_task_phase(
|
||||
self,
|
||||
*,
|
||||
task_id: str,
|
||||
lease_token: str,
|
||||
phase: str,
|
||||
current_site_id: Optional[int] = None,
|
||||
) -> bool:
|
||||
"""以当前租约更新任务阶段。"""
|
||||
return self._write(
|
||||
lambda repository: repository.update_task_phase(
|
||||
task_id=task_id,
|
||||
lease_token=lease_token,
|
||||
phase=phase,
|
||||
current_site_id=current_site_id,
|
||||
)
|
||||
)
|
||||
|
||||
def release_task(
|
||||
self,
|
||||
*,
|
||||
|
||||
@@ -0,0 +1,148 @@
|
||||
"""订阅执行状态的请求级异步 SQLAlchemy 适配器。"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.application.download.admission import SubscriptionDownloadSnapshot
|
||||
from app.application.subscription.execution import SearchBatchSnapshot, SearchTaskSnapshot
|
||||
from app.db.models.subscriptiondownload import SubscriptionDownloadSubmission
|
||||
from app.db.models.subscriptionsearch import SubscriptionSearchBatch, SubscriptionSearchTask
|
||||
|
||||
|
||||
def _task(record: SubscriptionSearchTask) -> SearchTaskSnapshot:
|
||||
"""复制可脱离 AsyncSession 使用的搜索任务快照。"""
|
||||
return SearchTaskSnapshot(
|
||||
task_id=record.task_id,
|
||||
batch_id=record.batch_id,
|
||||
subscription_id=record.subscription_id,
|
||||
source=record.source,
|
||||
priority=record.priority,
|
||||
position=record.position,
|
||||
state=record.state,
|
||||
phase=record.phase,
|
||||
attempt_count=record.attempt_count,
|
||||
cancel_requested=bool(record.cancel_requested),
|
||||
lease_token=record.lease_token,
|
||||
created_at=record.created_at,
|
||||
updated_at=record.updated_at,
|
||||
available_at=record.available_at,
|
||||
started_at=record.started_at,
|
||||
finished_at=record.finished_at,
|
||||
last_error=record.last_error,
|
||||
current_site_id=record.current_site_id,
|
||||
)
|
||||
|
||||
|
||||
def _batch(record: SubscriptionSearchBatch) -> SearchBatchSnapshot:
|
||||
"""复制可脱离 AsyncSession 使用的搜索批次快照。"""
|
||||
return SearchBatchSnapshot(
|
||||
batch_id=record.batch_id,
|
||||
source=record.source,
|
||||
state=record.state,
|
||||
priority=record.priority,
|
||||
total_count=record.total_count,
|
||||
finished_count=record.finished_count,
|
||||
failed_count=record.failed_count,
|
||||
cancelled_count=record.cancelled_count,
|
||||
cancel_requested=bool(record.cancel_requested),
|
||||
created_at=record.created_at,
|
||||
updated_at=record.updated_at,
|
||||
started_at=record.started_at,
|
||||
finished_at=record.finished_at,
|
||||
last_error=record.last_error,
|
||||
)
|
||||
|
||||
|
||||
def _download(record: SubscriptionDownloadSubmission) -> SubscriptionDownloadSnapshot:
|
||||
"""复制可脱离 AsyncSession 使用的下载提交快照。"""
|
||||
return SubscriptionDownloadSnapshot(
|
||||
idempotency_key=record.idempotency_key,
|
||||
subscription_id=record.subscription_id,
|
||||
task_id=record.task_id,
|
||||
state=record.state,
|
||||
attempt_count=record.attempt_count,
|
||||
attempt_token=record.attempt_token,
|
||||
downloader=record.downloader,
|
||||
download_hash=record.download_hash,
|
||||
available_at=record.available_at,
|
||||
last_error=record.last_error,
|
||||
created_at=record.created_at,
|
||||
updated_at=record.updated_at,
|
||||
)
|
||||
|
||||
|
||||
class SessionSubscriptionExecutionStatusRepository:
|
||||
"""复用请求 AsyncSession 批量读取搜索和下载执行事实。"""
|
||||
|
||||
def __init__(self, session: AsyncSession) -> None:
|
||||
"""绑定请求持有的异步会话。"""
|
||||
self._session = session
|
||||
|
||||
async def latest_search_tasks(
|
||||
self,
|
||||
subscription_ids: tuple[int, ...],
|
||||
) -> dict[int, SearchTaskSnapshot]:
|
||||
"""按更新时间倒序读取并在内存中保留每条订阅首项。"""
|
||||
result = await self._session.execute(
|
||||
select(SubscriptionSearchTask)
|
||||
.where(SubscriptionSearchTask.subscription_id.in_(subscription_ids))
|
||||
.order_by(
|
||||
SubscriptionSearchTask.updated_at.desc(),
|
||||
SubscriptionSearchTask.id.desc(),
|
||||
)
|
||||
)
|
||||
snapshots: dict[int, SearchTaskSnapshot] = {}
|
||||
for record in result.scalars().all():
|
||||
snapshots.setdefault(record.subscription_id, _task(record))
|
||||
return snapshots
|
||||
|
||||
async def latest_download_submissions(
|
||||
self,
|
||||
subscription_ids: tuple[int, ...],
|
||||
) -> dict[int, SubscriptionDownloadSnapshot]:
|
||||
"""按更新时间倒序读取并在内存中保留每条订阅首项。"""
|
||||
result = await self._session.execute(
|
||||
select(SubscriptionDownloadSubmission)
|
||||
.where(SubscriptionDownloadSubmission.subscription_id.in_(subscription_ids))
|
||||
.order_by(
|
||||
SubscriptionDownloadSubmission.updated_at.desc(),
|
||||
SubscriptionDownloadSubmission.id.desc(),
|
||||
)
|
||||
)
|
||||
snapshots: dict[int, SubscriptionDownloadSnapshot] = {}
|
||||
for record in result.scalars().all():
|
||||
snapshots.setdefault(record.subscription_id, _download(record))
|
||||
return snapshots
|
||||
|
||||
async def list_batches(self, *, limit: int) -> list[SearchBatchSnapshot]:
|
||||
"""返回最近更新的批次,访问范围由应用服务依据任务校验。"""
|
||||
result = await self._session.execute(
|
||||
select(SubscriptionSearchBatch)
|
||||
.order_by(
|
||||
SubscriptionSearchBatch.updated_at.desc(),
|
||||
SubscriptionSearchBatch.id.desc(),
|
||||
)
|
||||
.limit(limit)
|
||||
)
|
||||
return [_batch(record) for record in result.scalars().all()]
|
||||
|
||||
async def get_batch(self, batch_id: str) -> Optional[SearchBatchSnapshot]:
|
||||
"""按稳定批次 ID 返回状态快照。"""
|
||||
result = await self._session.execute(
|
||||
select(SubscriptionSearchBatch).where(
|
||||
SubscriptionSearchBatch.batch_id == batch_id
|
||||
)
|
||||
)
|
||||
record = result.scalars().first()
|
||||
return _batch(record) if record else None
|
||||
|
||||
async def list_batch_tasks(self, batch_id: str) -> list[SearchTaskSnapshot]:
|
||||
"""按持久位置返回批次任务,供访问校验和当前阶段投影。"""
|
||||
result = await self._session.execute(
|
||||
select(SubscriptionSearchTask)
|
||||
.where(SubscriptionSearchTask.batch_id == batch_id)
|
||||
.order_by(SubscriptionSearchTask.position, SubscriptionSearchTask.id)
|
||||
)
|
||||
return [_task(record) for record in result.scalars().all()]
|
||||
@@ -45,6 +45,8 @@ class SubscriptionSearchTask(Base):
|
||||
priority: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
||||
position: Mapped[int] = mapped_column(Integer, nullable=False)
|
||||
state: Mapped[str] = mapped_column(String(32), nullable=False, default="queued")
|
||||
phase: Mapped[str] = mapped_column(String(32), nullable=False, default="queued")
|
||||
current_site_id: Mapped[Optional[int]] = mapped_column(Integer)
|
||||
attempt_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
||||
cancel_requested: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
||||
lease_owner: Mapped[Optional[str]] = mapped_column(String(128))
|
||||
|
||||
@@ -60,6 +60,7 @@ class SubscriptionSearchOper(DbOper):
|
||||
priority=priority,
|
||||
position=position,
|
||||
state="queued",
|
||||
phase="queued",
|
||||
available_at=available_at or now,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
@@ -170,6 +171,8 @@ class SubscriptionSearchOper(DbOper):
|
||||
)
|
||||
.values(
|
||||
state="running",
|
||||
phase="matching",
|
||||
current_site_id=None,
|
||||
lease_owner=owner,
|
||||
lease_token=lease_token,
|
||||
lease_expires_at=lease_expires_at,
|
||||
@@ -207,6 +210,52 @@ class SubscriptionSearchOper(DbOper):
|
||||
return claimed_task
|
||||
return None
|
||||
|
||||
def update_task_phase(
|
||||
self,
|
||||
*,
|
||||
task_id: str,
|
||||
lease_token: str,
|
||||
phase: str,
|
||||
current_site_id: Optional[int],
|
||||
) -> bool:
|
||||
"""只允许当前运行租约推进用户可见阶段。"""
|
||||
if not isinstance(self._db, Session):
|
||||
raise RuntimeError("订阅搜索阶段更新需要调用方提供同步 Session")
|
||||
task = self._db.execute(
|
||||
select(SubscriptionSearchTask).where(
|
||||
SubscriptionSearchTask.task_id == task_id,
|
||||
SubscriptionSearchTask.state == "running",
|
||||
SubscriptionSearchTask.lease_token == lease_token,
|
||||
)
|
||||
).scalars().first()
|
||||
if task is None:
|
||||
return False
|
||||
now = utc_now_text()
|
||||
updated = execute_dml(
|
||||
self._db,
|
||||
update(SubscriptionSearchTask)
|
||||
.where(
|
||||
SubscriptionSearchTask.id == task.id,
|
||||
SubscriptionSearchTask.state == "running",
|
||||
SubscriptionSearchTask.lease_token == lease_token,
|
||||
)
|
||||
.values(
|
||||
phase=phase,
|
||||
current_site_id=current_site_id,
|
||||
updated_at=now,
|
||||
),
|
||||
execution_options={"synchronize_session": False},
|
||||
)
|
||||
if updated:
|
||||
execute_dml(
|
||||
self._db,
|
||||
update(SubscriptionSearchBatch)
|
||||
.where(SubscriptionSearchBatch.batch_id == task.batch_id)
|
||||
.values(updated_at=now),
|
||||
execution_options={"synchronize_session": False},
|
||||
)
|
||||
return bool(updated)
|
||||
|
||||
def finish_task(
|
||||
self,
|
||||
*,
|
||||
@@ -240,6 +289,8 @@ class SubscriptionSearchOper(DbOper):
|
||||
)
|
||||
.values(
|
||||
state=state,
|
||||
phase=state,
|
||||
current_site_id=None,
|
||||
active_key=None,
|
||||
lease_owner=None,
|
||||
lease_token=None,
|
||||
@@ -293,6 +344,8 @@ class SubscriptionSearchOper(DbOper):
|
||||
)
|
||||
.values(
|
||||
state="queued",
|
||||
phase="queued",
|
||||
current_site_id=None,
|
||||
lease_owner=None,
|
||||
lease_token=None,
|
||||
lease_expires_at=None,
|
||||
@@ -336,6 +389,8 @@ class SubscriptionSearchOper(DbOper):
|
||||
)
|
||||
.values(
|
||||
state="cancelled",
|
||||
phase="cancelled",
|
||||
current_site_id=None,
|
||||
active_key=None,
|
||||
cancel_requested=1,
|
||||
finished_at=now,
|
||||
@@ -350,7 +405,7 @@ class SubscriptionSearchOper(DbOper):
|
||||
SubscriptionSearchTask.batch_id == batch_id,
|
||||
SubscriptionSearchTask.state == "running",
|
||||
)
|
||||
.values(cancel_requested=1, updated_at=now),
|
||||
.values(cancel_requested=1, phase="cancelling", updated_at=now),
|
||||
execution_options={"synchronize_session": False},
|
||||
)
|
||||
self._refresh_batch(batch_id, now=now, error=None)
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import json
|
||||
from typing import Optional, List, Dict, Any, ClassVar, Literal
|
||||
from typing import Any, ClassVar, Dict, List, Literal, Optional
|
||||
|
||||
from pydantic import BaseModel, Field, ConfigDict, model_validator, field_validator
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
|
||||
from app.schemas.media import OptionalMediaIdentityMixin
|
||||
from app.schemas.types import MediaSource, MediaType
|
||||
@@ -45,6 +45,46 @@ def compute_subscribe_completed_episode(subscribe: "Subscribe") -> Optional[int]
|
||||
return min(max(start_episode - 1, 0), total_episode) + priority_completed
|
||||
|
||||
|
||||
class SubscriptionExecutionStatus(BaseModel):
|
||||
"""订阅列表可见的当前业务执行状态。"""
|
||||
|
||||
state: str
|
||||
phase: str
|
||||
updated_at: str
|
||||
source: Optional[str] = None
|
||||
batch_id: Optional[str] = None
|
||||
task_id: Optional[str] = None
|
||||
current_site_id: Optional[int] = None
|
||||
error: Optional[str] = None
|
||||
can_cancel: bool = False
|
||||
can_retry: bool = False
|
||||
requires_reconciliation: bool = False
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
class SubscriptionBatchStatus(BaseModel):
|
||||
"""订阅搜索批次的用户可见进度和操作能力。"""
|
||||
|
||||
batch_id: str
|
||||
source: str
|
||||
state: str
|
||||
phase: str
|
||||
total_count: int
|
||||
processed_count: int
|
||||
finished_count: int
|
||||
failed_count: int
|
||||
cancelled_count: int
|
||||
created_at: str
|
||||
updated_at: str
|
||||
current_subscription_id: Optional[int] = None
|
||||
current_site_id: Optional[int] = None
|
||||
error: Optional[str] = None
|
||||
can_cancel: bool = False
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
class Subscribe(OptionalMediaIdentityMixin, BaseModel):
|
||||
"""订阅输入与响应模型,媒体身份必须为空对或完整有效对。"""
|
||||
|
||||
@@ -59,6 +99,7 @@ class Subscribe(OptionalMediaIdentityMixin, BaseModel):
|
||||
"id", "poster", "backdrop", "vote", "description", "lack_episode", "completed_episode",
|
||||
"note", "state", "last_update", "username", "current_priority", "episode_priority", "date",
|
||||
"current_audio_format", "current_bitrate", "current_bit_depth", "current_sample_rate",
|
||||
"execution_status",
|
||||
})
|
||||
|
||||
id: Optional[int] = None
|
||||
@@ -158,6 +199,8 @@ class Subscribe(OptionalMediaIdentityMixin, BaseModel):
|
||||
filter_groups: Optional[List[str]] = Field(default_factory=list)
|
||||
# 剧集组
|
||||
episode_group: Optional[str] = None
|
||||
# 当前搜索或下载执行状态,只用于响应投影
|
||||
execution_status: Optional[SubscriptionExecutionStatus] = None
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
@@ -100,6 +100,14 @@ class SubscriptionBatchWriterFactory(Protocol):
|
||||
...
|
||||
|
||||
|
||||
class SubscriptionExecutionStatusRepositoryFactory(Protocol):
|
||||
"""由请求会话构造订阅执行状态读取仓储的工厂。"""
|
||||
|
||||
def __call__(self, session: object) -> object:
|
||||
"""绑定请求会话并返回执行状态读取端口。"""
|
||||
...
|
||||
|
||||
|
||||
class AsyncSessionProvider(Protocol):
|
||||
"""FastAPI 请求级异步会话提供器。"""
|
||||
|
||||
@@ -230,6 +238,8 @@ class SubscriptionRuntime:
|
||||
rule_group_mutation_scope: Callable[[], AbstractContextManager[SyncRuleGroupMutationService]]
|
||||
async_rule_group_mutation_scope: Callable[[], AbstractAsyncContextManager[AsyncRuleGroupMutationService]]
|
||||
site_reference_mutation_scope: Callable[[], AbstractContextManager[SyncSiteReferenceMutationService]]
|
||||
execution_status_repository: SubscriptionExecutionStatusRepositoryFactory | None = None
|
||||
search_repository: object | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
||||
@@ -153,6 +153,9 @@ def compose_runtime(inputs: RuntimeInputs) -> RuntimeComposition:
|
||||
SessionSubscriptionHistoryRepository,
|
||||
SessionSubscriptionRepository,
|
||||
)
|
||||
from app.db.adapters.subscriptionstatus import (
|
||||
SessionSubscriptionExecutionStatusRepository,
|
||||
)
|
||||
from app.db.oper.mediaserver import MediaServerOper
|
||||
from app.db.oper.message import MessageOper
|
||||
from app.db.oper.workflow import WorkflowOper
|
||||
@@ -214,6 +217,8 @@ def compose_runtime(inputs: RuntimeInputs) -> RuntimeComposition:
|
||||
async_session=get_async_db,
|
||||
repository=SessionSubscriptionRepository,
|
||||
history_repository=SessionSubscriptionHistoryRepository,
|
||||
execution_status_repository=SessionSubscriptionExecutionStatusRepository,
|
||||
search_repository=dependencies.subscription_search,
|
||||
transaction=SqlAlchemyAsyncUnitOfWork,
|
||||
outbox=SqlAlchemyAsyncOutboxStager,
|
||||
dispatch_store=SqlAlchemyAsyncOutboxDispatchStore(async_session_scope),
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
"""3.0.22 增加订阅搜索业务阶段与当前站点。
|
||||
|
||||
Revision ID: f3c8a1d6b2e9
|
||||
Revises: e1b6d4f8a2c7
|
||||
Create Date: 2026-09-01
|
||||
"""
|
||||
|
||||
# Alembic 的 op 是运行期代理,静态分析无法看到实际操作方法。
|
||||
# pylint: disable=no-member
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "f3c8a1d6b2e9"
|
||||
down_revision = "e1b6d4f8a2c7"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
_TABLE = "subscriptionsearchtask"
|
||||
|
||||
|
||||
def _column_names() -> set[str]:
|
||||
"""返回当前订阅搜索任务列名集合。"""
|
||||
inspector = sa.inspect(op.get_bind())
|
||||
if _TABLE not in set(inspector.get_table_names()):
|
||||
return set()
|
||||
return {column["name"] for column in inspector.get_columns(_TABLE)}
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""为存量队列增加带默认值的可观察阶段字段。"""
|
||||
columns = _column_names()
|
||||
if not columns:
|
||||
return
|
||||
if "phase" not in columns:
|
||||
op.add_column(
|
||||
_TABLE,
|
||||
sa.Column(
|
||||
"phase",
|
||||
sa.String(length=32),
|
||||
nullable=False,
|
||||
server_default="queued",
|
||||
),
|
||||
)
|
||||
if "current_site_id" not in columns:
|
||||
op.add_column(
|
||||
_TABLE,
|
||||
sa.Column("current_site_id", sa.Integer(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""移除业务阶段字段,保留原搜索队列事实。"""
|
||||
columns = _column_names()
|
||||
if "current_site_id" in columns:
|
||||
op.drop_column(_TABLE, "current_site_id")
|
||||
if "phase" in columns:
|
||||
op.drop_column(_TABLE, "phase")
|
||||
@@ -0,0 +1,159 @@
|
||||
"""订阅执行状态合并、批次权限和操作能力测试。"""
|
||||
|
||||
import asyncio
|
||||
|
||||
from app.application.download.admission import SubscriptionDownloadSnapshot
|
||||
from app.application.subscription.execution import SearchBatchSnapshot, SearchTaskSnapshot
|
||||
from app.application.subscription.status import SubscriptionExecutionStatusService
|
||||
|
||||
|
||||
def _task(
|
||||
subscription_id: int,
|
||||
*,
|
||||
state: str = "running",
|
||||
phase: str = "searching",
|
||||
updated_at: str = "2026-09-01T01:00:00+00:00",
|
||||
batch_id: str = "batch-1",
|
||||
) -> SearchTaskSnapshot:
|
||||
"""构造最小搜索任务快照。"""
|
||||
return SearchTaskSnapshot(
|
||||
task_id=f"task-{subscription_id}",
|
||||
batch_id=batch_id,
|
||||
subscription_id=subscription_id,
|
||||
source="manual",
|
||||
priority=100,
|
||||
position=subscription_id,
|
||||
state=state,
|
||||
phase=phase,
|
||||
attempt_count=1,
|
||||
cancel_requested=False,
|
||||
lease_token="lease" if state == "running" else None,
|
||||
created_at="2026-09-01T00:00:00+00:00",
|
||||
updated_at=updated_at,
|
||||
current_site_id=9 if phase == "waiting_site_budget" else None,
|
||||
last_error=" provider\n timeout " if state == "failed" else None,
|
||||
)
|
||||
|
||||
|
||||
def _download(subscription_id: int, state: str) -> SubscriptionDownloadSnapshot:
|
||||
"""构造一个比搜索任务更早的下载提交快照。"""
|
||||
return SubscriptionDownloadSnapshot(
|
||||
idempotency_key=f"key-{subscription_id}",
|
||||
subscription_id=subscription_id,
|
||||
task_id=f"task-{subscription_id}",
|
||||
state=state,
|
||||
attempt_count=1,
|
||||
attempt_token="attempt",
|
||||
downloader=None,
|
||||
download_hash=None,
|
||||
available_at=None,
|
||||
last_error="downloader response uncertain" if state == "reconcile_required" else None,
|
||||
created_at="2026-09-01T00:30:00+00:00",
|
||||
updated_at="2026-09-01T00:59:00+00:00",
|
||||
)
|
||||
|
||||
|
||||
class _Repository:
|
||||
"""保存测试快照的异步状态仓储。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""初始化可由测试覆盖的快照集合。"""
|
||||
self.tasks: dict[int, SearchTaskSnapshot] = {}
|
||||
self.downloads: dict[int, SubscriptionDownloadSnapshot] = {}
|
||||
self.batch = SearchBatchSnapshot(
|
||||
batch_id="batch-1",
|
||||
source="manual",
|
||||
state="running",
|
||||
priority=100,
|
||||
total_count=2,
|
||||
finished_count=0,
|
||||
failed_count=0,
|
||||
cancelled_count=0,
|
||||
cancel_requested=False,
|
||||
created_at="2026-09-01T00:00:00+00:00",
|
||||
updated_at="2026-09-01T01:00:00+00:00",
|
||||
)
|
||||
|
||||
async def latest_search_tasks(self, subscription_ids):
|
||||
"""返回请求范围内搜索任务。"""
|
||||
return {key: value for key, value in self.tasks.items() if key in subscription_ids}
|
||||
|
||||
async def latest_download_submissions(self, subscription_ids):
|
||||
"""返回请求范围内下载提交。"""
|
||||
return {key: value for key, value in self.downloads.items() if key in subscription_ids}
|
||||
|
||||
async def list_batches(self, *, limit):
|
||||
"""返回一个测试批次。"""
|
||||
return [self.batch][:limit]
|
||||
|
||||
async def get_batch(self, batch_id):
|
||||
"""按 ID 返回测试批次。"""
|
||||
return self.batch if batch_id == self.batch.batch_id else None
|
||||
|
||||
async def list_batch_tasks(self, batch_id):
|
||||
"""返回属于测试批次的任务。"""
|
||||
return [task for task in self.tasks.values() if task.batch_id == batch_id]
|
||||
|
||||
|
||||
def test_execution_status_exposes_site_wait_and_cancel_capability():
|
||||
"""站点预算等待必须保留当前站点和取消能力。"""
|
||||
repository = _Repository()
|
||||
repository.tasks[1] = _task(1, phase="waiting_site_budget")
|
||||
|
||||
statuses = asyncio.run(
|
||||
SubscriptionExecutionStatusService(repository).for_subscriptions((1,))
|
||||
)
|
||||
|
||||
assert statuses[1].state == "waiting_site_budget"
|
||||
assert statuses[1].current_site_id == 9
|
||||
assert statuses[1].can_cancel is True
|
||||
|
||||
|
||||
def test_reconciliation_state_overrides_newer_search_terminal():
|
||||
"""不确定下载副作用不得被稍晚写入的搜索失败掩盖。"""
|
||||
repository = _Repository()
|
||||
repository.tasks[2] = _task(2, state="failed", phase="failed")
|
||||
repository.downloads[2] = _download(2, "reconcile_required")
|
||||
|
||||
statuses = asyncio.run(
|
||||
SubscriptionExecutionStatusService(repository).for_subscriptions((2,))
|
||||
)
|
||||
|
||||
assert statuses[2].state == "reconcile_required"
|
||||
assert statuses[2].requires_reconciliation is True
|
||||
assert statuses[2].can_retry is False
|
||||
assert statuses[2].error == "downloader response uncertain"
|
||||
|
||||
|
||||
def test_failed_search_exposes_safe_retryable_error():
|
||||
"""搜索失败文本必须压平且仅声明安全重试能力。"""
|
||||
repository = _Repository()
|
||||
repository.tasks[3] = _task(3, state="failed", phase="failed")
|
||||
|
||||
statuses = asyncio.run(
|
||||
SubscriptionExecutionStatusService(repository).for_subscriptions((3,))
|
||||
)
|
||||
|
||||
assert statuses[3].state == "failed"
|
||||
assert statuses[3].can_retry is True
|
||||
assert statuses[3].error == "provider timeout"
|
||||
|
||||
|
||||
def test_batch_requires_complete_subscription_access():
|
||||
"""普通用户不得读取混合其他 owner 订阅的批次聚合。"""
|
||||
repository = _Repository()
|
||||
repository.tasks = {1: _task(1), 2: _task(2)}
|
||||
service = SubscriptionExecutionStatusService(repository)
|
||||
|
||||
hidden = asyncio.run(
|
||||
service.get_batch("batch-1", accessible_subscription_ids={1})
|
||||
)
|
||||
visible = asyncio.run(
|
||||
service.get_batch("batch-1", accessible_subscription_ids={1, 2})
|
||||
)
|
||||
|
||||
assert hidden is None
|
||||
assert visible is not None
|
||||
assert visible.current_subscription_id == 1
|
||||
assert visible.processed_count == 0
|
||||
assert visible.can_cancel is True
|
||||
@@ -68,6 +68,34 @@ def test_search_queue_recovers_expired_lease_with_same_task_identity(tmp_path):
|
||||
assert recovered.attempt_count == 2
|
||||
|
||||
|
||||
def test_search_queue_phase_update_requires_current_lease(tmp_path):
|
||||
"""过期执行者不得覆盖当前任务的用户可见阶段。"""
|
||||
repository, _engine = _repository(tmp_path)
|
||||
repository.enqueue(subscription_ids=(30,), source="manual", priority=100)
|
||||
task = repository.claim_next(owner="worker-a")
|
||||
|
||||
assert repository.update_task_phase(
|
||||
task_id=task.task_id,
|
||||
lease_token="stale-token",
|
||||
phase="searching",
|
||||
current_site_id=7,
|
||||
) is False
|
||||
assert repository.update_task_phase(
|
||||
task_id=task.task_id,
|
||||
lease_token=task.lease_token,
|
||||
phase="waiting_site_budget",
|
||||
current_site_id=7,
|
||||
) is True
|
||||
|
||||
current = repository.claim_next(owner="worker-b")
|
||||
assert current is None
|
||||
assert repository.finish_task(
|
||||
task_id=task.task_id,
|
||||
lease_token=task.lease_token,
|
||||
state="completed",
|
||||
) is True
|
||||
|
||||
|
||||
def test_search_queue_cancel_finishes_queued_and_running_tasks(tmp_path):
|
||||
"""取消立即终止未发请求任务,运行中任务在租约边界收口。"""
|
||||
repository, engine = _repository(tmp_path)
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
"""订阅搜索业务阶段字段的 Alembic 可逆迁移测试。"""
|
||||
|
||||
import importlib
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic.migration import MigrationContext
|
||||
from alembic.operations import Operations
|
||||
|
||||
MIGRATION = "database.versions.f3c8a1d6b2e9_3_0_22"
|
||||
|
||||
|
||||
def _bind_migration(monkeypatch, connection):
|
||||
"""把 3.0.22 迁移绑定到隔离 SQLite 连接。"""
|
||||
migration = importlib.import_module(MIGRATION)
|
||||
monkeypatch.setattr(
|
||||
migration,
|
||||
"op",
|
||||
Operations(MigrationContext.configure(connection)),
|
||||
)
|
||||
return migration
|
||||
|
||||
|
||||
def _create_legacy_search_task(connection) -> None:
|
||||
"""创建迁移前的最小订阅搜索任务表和一条排队记录。"""
|
||||
metadata = sa.MetaData()
|
||||
task = sa.Table(
|
||||
"subscriptionsearchtask",
|
||||
metadata,
|
||||
sa.Column("id", sa.Integer(), primary_key=True),
|
||||
sa.Column("task_id", sa.String(64), nullable=False),
|
||||
sa.Column("state", sa.String(32), nullable=False),
|
||||
)
|
||||
metadata.create_all(connection)
|
||||
connection.execute(task.insert(), {"id": 1, "task_id": "task-1", "state": "queued"})
|
||||
|
||||
|
||||
def test_subscription_status_migration_upgrade_downgrade_reupgrade(monkeypatch) -> None:
|
||||
"""存量任务应获得默认阶段,迁移可重复执行并完整回滚。"""
|
||||
engine = sa.create_engine("sqlite://")
|
||||
with engine.begin() as connection:
|
||||
_create_legacy_search_task(connection)
|
||||
migration = _bind_migration(monkeypatch, connection)
|
||||
|
||||
migration.upgrade()
|
||||
migration.upgrade()
|
||||
|
||||
columns = {
|
||||
column["name"]
|
||||
for column in sa.inspect(connection).get_columns("subscriptionsearchtask")
|
||||
}
|
||||
assert {"phase", "current_site_id"}.issubset(columns)
|
||||
row = connection.execute(
|
||||
sa.text(
|
||||
"SELECT phase, current_site_id FROM subscriptionsearchtask WHERE id = 1"
|
||||
)
|
||||
).mappings().one()
|
||||
assert dict(row) == {"phase": "queued", "current_site_id": None}
|
||||
|
||||
migration.downgrade()
|
||||
downgraded = {
|
||||
column["name"]
|
||||
for column in sa.inspect(connection).get_columns("subscriptionsearchtask")
|
||||
}
|
||||
assert "phase" not in downgraded
|
||||
assert "current_site_id" not in downgraded
|
||||
|
||||
migration.upgrade()
|
||||
assert connection.execute(
|
||||
sa.text("SELECT phase FROM subscriptionsearchtask WHERE id = 1")
|
||||
).scalar_one() == "queued"
|
||||
engine.dispose()
|
||||
Reference in New Issue
Block a user