mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-08-30 12:36:55 +08:00
fix(database): harden worker shutdown and overload propagation
This commit is contained in:
@@ -6,6 +6,8 @@ from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Optional
|
||||
|
||||
from app.application.database import DatabaseWorkerOverloadedError
|
||||
|
||||
|
||||
InstalledPluginsReader = Callable[[], list[str]]
|
||||
InstalledPluginsWriter = Callable[[list[str]], Awaitable[object]]
|
||||
@@ -127,7 +129,7 @@ class PluginInstallCommand:
|
||||
force,
|
||||
)
|
||||
except Exception as err:
|
||||
return await self._failure(
|
||||
result = await self._failure(
|
||||
plugin_id=plugin_id,
|
||||
original_plugins=installed_plugins,
|
||||
checkpoint=checkpoint,
|
||||
@@ -135,6 +137,9 @@ class PluginInstallCommand:
|
||||
message=str(err),
|
||||
package_installed=False,
|
||||
)
|
||||
if isinstance(err, DatabaseWorkerOverloadedError):
|
||||
raise
|
||||
return result
|
||||
if not state:
|
||||
return await self._failure(
|
||||
plugin_id=plugin_id,
|
||||
@@ -152,7 +157,7 @@ class PluginInstallCommand:
|
||||
await self._installed_plugins_writer(updated_plugins)
|
||||
installed_list_persisted = True
|
||||
except Exception as err:
|
||||
return await self._failure(
|
||||
result = await self._failure(
|
||||
plugin_id=plugin_id,
|
||||
original_plugins=installed_plugins,
|
||||
checkpoint=checkpoint,
|
||||
@@ -160,11 +165,14 @@ class PluginInstallCommand:
|
||||
message=str(err),
|
||||
package_installed=True,
|
||||
)
|
||||
if isinstance(err, DatabaseWorkerOverloadedError):
|
||||
raise
|
||||
return result
|
||||
|
||||
try:
|
||||
await self._plugin_reloader(plugin_id)
|
||||
except Exception as err:
|
||||
return await self._failure(
|
||||
result = await self._failure(
|
||||
plugin_id=plugin_id,
|
||||
original_plugins=installed_plugins,
|
||||
checkpoint=checkpoint,
|
||||
@@ -174,11 +182,14 @@ class PluginInstallCommand:
|
||||
installed_list_persisted=installed_list_persisted,
|
||||
runtime_touched=True,
|
||||
)
|
||||
if isinstance(err, DatabaseWorkerOverloadedError):
|
||||
raise
|
||||
return result
|
||||
|
||||
try:
|
||||
await self._registration_refresher(plugin_id)
|
||||
except Exception as err:
|
||||
return await self._failure(
|
||||
result = await self._failure(
|
||||
plugin_id=plugin_id,
|
||||
original_plugins=installed_plugins,
|
||||
checkpoint=checkpoint,
|
||||
@@ -189,6 +200,9 @@ class PluginInstallCommand:
|
||||
runtime_touched=True,
|
||||
registrations_touched=True,
|
||||
)
|
||||
if isinstance(err, DatabaseWorkerOverloadedError):
|
||||
raise
|
||||
return result
|
||||
|
||||
checkpoint_cleanup_error = ""
|
||||
try:
|
||||
@@ -262,7 +276,7 @@ class PluginInstallCommand:
|
||||
registrations_restored = True
|
||||
except Exception as rollback_err:
|
||||
rollback_errors.append(f"路由和服务注册恢复失败:{rollback_err}")
|
||||
return PluginInstallResult(
|
||||
result = PluginInstallResult(
|
||||
success=False,
|
||||
message=f"刷新插件运行态失败:{err}",
|
||||
refreshed_only=True,
|
||||
@@ -275,6 +289,9 @@ class PluginInstallCommand:
|
||||
errors=tuple(rollback_errors),
|
||||
),
|
||||
)
|
||||
if isinstance(err, DatabaseWorkerOverloadedError):
|
||||
raise
|
||||
return result
|
||||
|
||||
reported = False
|
||||
report_error = ""
|
||||
|
||||
+22
-6
@@ -67,6 +67,8 @@ class DatabaseWorker:
|
||||
] = {}
|
||||
self._queued = 0
|
||||
self._running = 0
|
||||
self._reported_queued = 0
|
||||
self._reported_running = 0
|
||||
self._rejected = 0
|
||||
self._closing = False
|
||||
|
||||
@@ -200,12 +202,19 @@ class DatabaseWorker:
|
||||
)
|
||||
self._record_depth()
|
||||
|
||||
async def _wait_until_done(self, future: asyncio.Future[object]) -> None:
|
||||
"""忽略后续取消请求,直到线程内事务结束。"""
|
||||
async def _wait_until_done(
|
||||
self,
|
||||
future: asyncio.Future[object],
|
||||
*,
|
||||
interruptible: bool = False,
|
||||
) -> None:
|
||||
"""等待线程内事务结束,并按调用场景决定是否响应外层取消。"""
|
||||
while not future.done():
|
||||
try:
|
||||
await asyncio.shield(future)
|
||||
except asyncio.CancelledError:
|
||||
if interruptible:
|
||||
raise
|
||||
continue
|
||||
except BaseException:
|
||||
break
|
||||
@@ -224,7 +233,8 @@ class DatabaseWorker:
|
||||
future.cancel()
|
||||
for future, (wrapped, _item) in futures:
|
||||
if not future.cancelled():
|
||||
await self._wait_until_done(wrapped)
|
||||
# 关停超时必须能返回并保留 owner;已开始的数据库事务继续由线程完成。
|
||||
await self._wait_until_done(wrapped, interruptible=True)
|
||||
executor.shutdown(wait=True, cancel_futures=True)
|
||||
while self.snapshot().queued or self.snapshot().running:
|
||||
await asyncio.sleep(0)
|
||||
@@ -232,7 +242,13 @@ class DatabaseWorker:
|
||||
self._record_depth()
|
||||
|
||||
def _record_depth(self) -> None:
|
||||
"""记录当前排队量与运行量。"""
|
||||
"""以状态变化量记录队列和运行中的任务数量。"""
|
||||
stats = self.snapshot()
|
||||
record_metric("db.worker.queue.depth", stats.queued)
|
||||
record_metric("db.worker.active", stats.running)
|
||||
queued_delta = stats.queued - self._reported_queued
|
||||
running_delta = stats.running - self._reported_running
|
||||
if queued_delta:
|
||||
record_metric("db.worker.queue.depth", queued_delta)
|
||||
if running_delta:
|
||||
record_metric("db.worker.active", running_delta)
|
||||
self._reported_queued = stats.queued
|
||||
self._reported_running = stats.running
|
||||
|
||||
@@ -79,6 +79,7 @@ def _native_ai_error_response(
|
||||
protocol: str,
|
||||
status_code: int,
|
||||
message: str,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> JSONResponse:
|
||||
"""按 OpenAI 或 Anthropic 兼容协议构造原生错误响应。"""
|
||||
if protocol == "openai":
|
||||
@@ -98,6 +99,7 @@ def _native_ai_error_response(
|
||||
code=error_type,
|
||||
)
|
||||
).model_dump(mode="json"),
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
error_type = (
|
||||
@@ -112,6 +114,7 @@ def _native_ai_error_response(
|
||||
content=AnthropicErrorResponse(
|
||||
error=AnthropicErrorDetail(type=error_type, message=message)
|
||||
).model_dump(mode="json"),
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
||||
@@ -119,6 +122,7 @@ def _mcp_jsonrpc_error_response(
|
||||
status_code: int,
|
||||
code: int,
|
||||
message: str,
|
||||
headers: dict[str, str] | None = None,
|
||||
) -> JSONResponse:
|
||||
"""构造带 HTTP 状态码的 MCP JSON-RPC 原生错误响应。"""
|
||||
return JSONResponse(
|
||||
@@ -128,6 +132,7 @@ def _mcp_jsonrpc_error_response(
|
||||
id=None,
|
||||
error=McpJsonRpcErrorDetail(code=code, message=message),
|
||||
).model_dump(mode="json"),
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
||||
@@ -202,6 +207,7 @@ async def localized_http_exception_handler(
|
||||
protocol=native_ai_protocol,
|
||||
status_code=exc.status_code,
|
||||
message=message,
|
||||
headers=exc.headers,
|
||||
)
|
||||
if _is_mcp_jsonrpc_request(request):
|
||||
error_codes = {
|
||||
@@ -215,6 +221,7 @@ async def localized_http_exception_handler(
|
||||
status_code=exc.status_code,
|
||||
code=error_codes.get(exc.status_code, -32000),
|
||||
message=message,
|
||||
headers=exc.headers,
|
||||
)
|
||||
return JSONResponse(
|
||||
status_code=exc.status_code,
|
||||
|
||||
@@ -2640,6 +2640,8 @@
|
||||
"app.application.plugin.folders -> app.runtime.log",
|
||||
"app.application.plugin.folders -> app.schemas",
|
||||
"app.application.plugin.folders -> app.schemas.types",
|
||||
"app.application.plugin.install -> app.application",
|
||||
"app.application.plugin.install -> app.application.database",
|
||||
"app.application.recognition -> app.application",
|
||||
"app.application.recognition -> app.application.configuration",
|
||||
"app.application.recognition -> app.schemas",
|
||||
|
||||
@@ -7,6 +7,7 @@ from fastapi import APIRouter, Depends, FastAPI, HTTPException
|
||||
from fastapi.responses import JSONResponse
|
||||
from fastapi.routing import APIRoute
|
||||
from pydantic import BaseModel, ValidationError
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response as StarletteResponse
|
||||
from starlette.responses import StreamingResponse
|
||||
|
||||
@@ -24,6 +25,7 @@ from app.factory import (
|
||||
)
|
||||
from app.application.database import DatabaseWorkerOverloadedError
|
||||
from app.runtime.localization import LocaleHelper
|
||||
from app.runtime.config import settings
|
||||
from app.schemas.common import JsonData
|
||||
from app.schemas.response import Response
|
||||
|
||||
@@ -217,6 +219,40 @@ async def test_database_worker_overload_is_retryable_service_unavailable(
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"path",
|
||||
[
|
||||
f"{settings.API_V1_STR}/openai/v1/chat/completions",
|
||||
f"{settings.API_V1_STR}/anthropic/v1/messages",
|
||||
f"{settings.API_V1_STR}/mcp",
|
||||
],
|
||||
)
|
||||
async def test_database_worker_overload_preserves_retry_after_for_native_protocols(
|
||||
path: str,
|
||||
):
|
||||
"""OpenAI、Anthropic 和 MCP 的原生 503 也必须保留重试提示。"""
|
||||
scope = {
|
||||
"type": "http",
|
||||
"http_version": "1.1",
|
||||
"method": "GET",
|
||||
"scheme": "http",
|
||||
"path": path,
|
||||
"raw_path": path.encode(),
|
||||
"query_string": b"",
|
||||
"headers": [],
|
||||
"server": ("testserver", 80),
|
||||
"client": ("testclient", 123),
|
||||
"root_path": "",
|
||||
}
|
||||
response = await database_worker_overloaded_handler(
|
||||
Request(scope),
|
||||
DatabaseWorkerOverloadedError("worker full"),
|
||||
)
|
||||
|
||||
assert response.status_code == 503
|
||||
assert response.headers["retry-after"] == "1"
|
||||
|
||||
|
||||
async def test_validation_error_uses_unified_model(api_app: FastAPI):
|
||||
"""请求参数校验失败应返回统一协议和明确的错误项结构。"""
|
||||
async with make_client(api_app) as client:
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import asyncio
|
||||
import threading
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -138,3 +139,64 @@ async def test_shutdown_rejects_new_work_and_waits_for_running_work() -> None:
|
||||
assert worker.snapshot().closing is True
|
||||
assert worker.snapshot().queued == 0
|
||||
assert worker.snapshot().running == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_timeout_keeps_running_owner_until_transaction_finishes() -> None:
|
||||
"""关闭超时应返回给生命周期编排,并保留执行器等待事务收敛。"""
|
||||
worker = DatabaseWorker(max_workers=1, capacity=1)
|
||||
await worker.start()
|
||||
started = threading.Event()
|
||||
release = threading.Event()
|
||||
|
||||
def operation() -> None:
|
||||
started.set()
|
||||
release.wait(1)
|
||||
|
||||
running = asyncio.create_task(worker.run(operation))
|
||||
await asyncio.to_thread(started.wait)
|
||||
|
||||
with pytest.raises(asyncio.TimeoutError):
|
||||
await asyncio.wait_for(worker.shutdown(), timeout=0.01)
|
||||
|
||||
assert worker.snapshot().closing is True
|
||||
assert worker._executor is not None
|
||||
|
||||
release.set()
|
||||
await running
|
||||
await asyncio.sleep(0)
|
||||
await worker.shutdown()
|
||||
assert worker._executor is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_worker_depth_metrics_emit_deltas_and_return_to_zero() -> None:
|
||||
"""队列和运行量指标按增减量上报,不能把绝对值累加成漂移。"""
|
||||
worker = DatabaseWorker(max_workers=1, capacity=1)
|
||||
started = threading.Event()
|
||||
release = threading.Event()
|
||||
|
||||
def operation() -> None:
|
||||
started.set()
|
||||
release.wait(1)
|
||||
|
||||
with patch("app.db.worker.record_metric") as record_metric:
|
||||
await worker.start()
|
||||
running = asyncio.create_task(worker.run(operation))
|
||||
await asyncio.to_thread(started.wait)
|
||||
release.set()
|
||||
await running
|
||||
await worker.shutdown()
|
||||
|
||||
queue_values = [
|
||||
call.args[1]
|
||||
for call in record_metric.call_args_list
|
||||
if call.args[0] == "db.worker.queue.depth"
|
||||
]
|
||||
active_values = [
|
||||
call.args[1]
|
||||
for call in record_metric.call_args_list
|
||||
if call.args[0] == "db.worker.active"
|
||||
]
|
||||
assert queue_values == [1.0, -1.0]
|
||||
assert active_values == [1.0, -1.0]
|
||||
|
||||
@@ -2,6 +2,7 @@ from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from app.application.database import DatabaseWorkerOverloadedError
|
||||
from app.application.plugin.install import PluginInstallCommand
|
||||
|
||||
|
||||
@@ -181,6 +182,26 @@ async def test_persistence_failure_restores_package_without_touching_runtime():
|
||||
reloader.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_database_worker_overload_rolls_back_and_reaches_api_boundary():
|
||||
"""配置 worker 背压完成补偿后继续抛出,交由 API 映射为 503。"""
|
||||
checkpoint = object()
|
||||
rollback = AsyncMock()
|
||||
command = _command(
|
||||
checkpointer=AsyncMock(return_value=checkpoint),
|
||||
writer=AsyncMock(side_effect=DatabaseWorkerOverloadedError("worker full")),
|
||||
rollback=rollback,
|
||||
)
|
||||
|
||||
with pytest.raises(DatabaseWorkerOverloadedError):
|
||||
await command.execute(
|
||||
plugin_id="DemoPlugin",
|
||||
repo_url="https://github.com/demo/plugins",
|
||||
)
|
||||
|
||||
rollback.assert_awaited_once_with(checkpoint)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reload_failure_restores_list_files_and_previous_runtime():
|
||||
"""重载失败时依次恢复已安装列表、包文件和旧运行态。"""
|
||||
|
||||
Reference in New Issue
Block a user