diff --git a/app/api/endpoints/plugin.py b/app/api/endpoints/plugin.py index 41957db8b..76fd14fd5 100644 --- a/app/api/endpoints/plugin.py +++ b/app/api/endpoints/plugin.py @@ -874,7 +874,10 @@ async def save_plugin_folders( 保存插件文件夹分组配置 """ try: - get_configured_system_config().set(SystemConfigKey.PluginFolders, folders) + await get_configured_system_config().async_set( + SystemConfigKey.PluginFolders, + folders, + ) return _SchemaResponse(success=True) except Exception as e: logger.error(f"[文件夹API] 保存文件夹配置失败: {str(e)}") @@ -893,7 +896,10 @@ async def create_plugin_folder( folders = get_configured_system_config().get(SystemConfigKey.PluginFolders) or {} if folder_name not in folders: folders[folder_name] = [] - get_configured_system_config().set(SystemConfigKey.PluginFolders, folders) + await get_configured_system_config().async_set( + SystemConfigKey.PluginFolders, + folders, + ) return _SchemaResponse( success=True, message=f"文件夹 '{folder_name}' 创建成功" ) diff --git a/app/api/endpoints/site.py b/app/api/endpoints/site.py index ebd257279..57c6abe79 100644 --- a/app/api/endpoints/site.py +++ b/app/api/endpoints/site.py @@ -188,8 +188,8 @@ async def reset( 清空所有站点数据并重新同步CookieCloud站点信息 """ result = await command.reset() - get_configured_system_config().set(SystemConfigKey.IndexerSites, []) - get_configured_system_config().set(SystemConfigKey.RssSites, []) + await get_configured_system_config().async_set(SystemConfigKey.IndexerSites, []) + await get_configured_system_config().async_set(SystemConfigKey.RssSites, []) # 启动定时服务 Scheduler().start("cookiecloud", manual=True) # 插件站点删除 diff --git a/app/api/endpoints/user.py b/app/api/endpoints/user.py index 9f986aba1..94163a1f8 100644 --- a/app/api/endpoints/user.py +++ b/app/api/endpoints/user.py @@ -16,7 +16,6 @@ from app.application.security.user import UserService from app.api.dependencies.auth import ( get_current_active_superuser_async, get_current_active_user_async, - get_current_active_user, get_user_service, ) from app.application.security.userconfig import get_configured_user_configuration @@ -140,7 +139,10 @@ async def upload_avatar( summary="查询用户配置", response_model=_SchemaResponse[_SchemaValueData], ) -def get_config(key: str, current_user: Any = Depends(get_current_active_user)): +async def get_config( + key: str, + current_user: Any = Depends(get_current_active_user_async), +): """ 查询用户配置 """ @@ -149,15 +151,15 @@ def get_config(key: str, current_user: Any = Depends(get_current_active_user)): @router.post("/config/{key}", summary="更新用户配置", response_model=_SchemaResponse[None]) -def set_config( +async def set_config( key: str, value: Annotated[Union[list, dict, bool, int, str] | None, Body()] = None, - current_user: Any = Depends(get_current_active_user), + current_user: Any = Depends(get_current_active_user_async), ): """ 更新用户配置 """ - get_configured_user_configuration().set( + await get_configured_user_configuration().async_set( username=current_user.name, key=key, value=value, diff --git a/app/application/configuration.py b/app/application/configuration.py index 59f62d669..01dd656ff 100644 --- a/app/application/configuration.py +++ b/app/application/configuration.py @@ -4,9 +4,11 @@ from __future__ import annotations from collections.abc import Callable from dataclasses import dataclass +from functools import partial from pathlib import Path from typing import Any, Optional, Protocol +from app.application.database import AsyncDatabaseExecutor from app.schemas.types import MediaType @@ -16,9 +18,6 @@ class SystemConfigReader(Protocol): def get(self, key: Any = None) -> Any: """读取配置。""" - async def async_get(self, key: Any = None) -> Any: - """异步读取配置。""" - class SystemConfigWriter(Protocol): """持久化用户配置的最小写入端口。""" @@ -26,9 +25,6 @@ class SystemConfigWriter(Protocol): def set(self, key: Any, value: Any) -> bool | None: """写入配置。""" - async def async_set(self, key: Any, value: Any) -> bool | None: - """异步写入配置。""" - def delete(self, key: Any) -> Any: """删除配置。""" @@ -296,14 +292,16 @@ class SystemConfigService: *, reader: SystemConfigReader | None = None, writer: SystemConfigWriter | None = None, + async_executor: AsyncDatabaseExecutor | None = None, ) -> None: - """注入可分离的读写端口,并兼容旧的单仓储装配参数。""" + """注入读写端口及可选的异步事务执行能力。""" resolved_reader = reader or repository resolved_writer = writer or repository if resolved_reader is None or resolved_writer is None: raise ValueError("系统配置服务必须同时提供 reader 与 writer") self._reader = resolved_reader self._writer = resolved_writer + self._async_executor = async_executor def get(self, key: Any = None) -> Any: """读取配置。""" @@ -313,13 +311,11 @@ class SystemConfigService: """写入配置。""" return self._writer.set(key, value) - async def async_get(self, key: Any = None) -> Any: - """异步读取配置。""" - return await self._reader.async_get(key) - async def async_set(self, key: Any, value: Any) -> bool | None: - """异步写入配置。""" - return await self._writer.async_set(key, value) + """异步写入配置,并等待数据库提交或回滚完成。""" + if self._async_executor is None: + raise RuntimeError("系统配置异步数据库执行端口尚未配置") + return await self._async_executor.run(partial(self._writer.set, key, value)) def delete(self, key: Any) -> Any: """删除配置。""" diff --git a/app/application/database.py b/app/application/database.py index 511350f1d..d69414d27 100644 --- a/app/application/database.py +++ b/app/application/database.py @@ -3,17 +3,27 @@ from __future__ import annotations from collections.abc import Callable -from typing import Any, Optional +from typing import TYPE_CHECKING, Any, Optional, Protocol, TypeVar -from app.application.backup import ( - BackupArtifact, - BackupVerification, - DatabaseBackupService, -) -from app.application.maintenance import DataCleanupService +if TYPE_CHECKING: + from app.application.backup import ( + BackupArtifact, + BackupVerification, + DatabaseBackupService, + ) + from app.application.maintenance import DataCleanupService DatabaseProbe = Callable[[], Optional[str]] +T = TypeVar("T") + + +class AsyncDatabaseExecutor(Protocol): + """让异步业务调用同步短事务而不阻塞事件循环。""" + + async def run(self, operation: Callable[[], T]) -> T: + """等待数据库操作完成提交或回滚,并返回执行结果。""" + ... class DatabaseHealthService: diff --git a/app/application/security/userconfig.py b/app/application/security/userconfig.py index 09617ee53..739656363 100644 --- a/app/application/security/userconfig.py +++ b/app/application/security/userconfig.py @@ -2,8 +2,11 @@ from __future__ import annotations +from functools import partial from typing import Any, Protocol +from app.application.database import AsyncDatabaseExecutor + class UserConfigurationRepository(Protocol): """用户配置数据端口。""" @@ -18,9 +21,15 @@ class UserConfigurationRepository(Protocol): class UserConfigurationService: """编排用户个性化配置读写。""" - def __init__(self, repository: UserConfigurationRepository) -> None: - """注入用户配置数据端口。""" + def __init__( + self, + repository: UserConfigurationRepository, + *, + async_executor: AsyncDatabaseExecutor | None = None, + ) -> None: + """注入用户配置数据端口及可选的异步事务执行能力。""" self._repository = repository + self._async_executor = async_executor def get(self, username: str, key: str) -> Any: """读取用户配置。""" @@ -30,6 +39,14 @@ class UserConfigurationService: """写入用户配置。""" return self._repository.set(username=username, key=key, value=value) + async def async_set(self, username: str, key: str, value: Any) -> Any: + """异步写入用户配置,并等待数据库提交或回滚完成。""" + if self._async_executor is None: + raise RuntimeError("用户配置异步数据库执行端口尚未配置") + return await self._async_executor.run( + partial(self._repository.set, username=username, key=key, value=value) + ) + _configured_user_configuration: UserConfigurationService | None = None diff --git a/app/db/oper/systemconfig.py b/app/db/oper/systemconfig.py index df0212bac..63be40ce0 100644 --- a/app/db/oper/systemconfig.py +++ b/app/db/oper/systemconfig.py @@ -1,4 +1,3 @@ -import asyncio import copy import threading from typing import Any, Optional, Union @@ -14,15 +13,38 @@ class SystemConfigOper(DbOper, metaclass=Singleton): 系统配置管理 """ def __init__(self): - """ - 加载配置到内存 - """ + """初始化空快照,数据库加载由启动组合根显式执行。""" super().__init__() self.__SYSTEMCONF = {} - self._rlock = threading.RLock() - self._alock = asyncio.Lock() - for item in SystemConfig.list(self._db): - self.__SYSTEMCONF[item.key] = item.value + self._snapshot_lock = threading.RLock() + self._write_lock = threading.RLock() + self._loaded = False + + def load_snapshot(self) -> None: + """从数据库加载完整配置,并一次性发布新的内存快照。""" + with self._write_lock: + snapshot = { + item.key: copy.deepcopy(item.value) + for item in SystemConfig.list(self._db) + } + with self._snapshot_lock: + self.__SYSTEMCONF = snapshot + self._loaded = True + + def _require_loaded(self) -> None: + """阻止消费者读取尚未完成启动加载的半成品快照。""" + if not self._loaded: + raise RuntimeError("系统配置快照尚未加载") + + def _publish_value(self, key: str, value: Any) -> None: + """在事务成功后短暂持锁发布单项配置。""" + with self._snapshot_lock: + self.__SYSTEMCONF[key] = copy.deepcopy(value) + + def _publish_delete(self, key: str) -> None: + """在事务成功后短暂持锁移除单项配置。""" + with self._snapshot_lock: + self.__SYSTEMCONF.pop(key, None) def set(self, key: Union[str, SystemConfigKey], value: Any) -> Optional[bool]: """ @@ -33,59 +55,25 @@ class SystemConfigOper(DbOper, metaclass=Singleton): """ if isinstance(key, SystemConfigKey): key = key.value - with self._rlock: - # 旧值 - old_value = self.__SYSTEMCONF.get(key) - # 更新内存(deepcopy避免内存共享) - self.__SYSTEMCONF[key] = copy.deepcopy(value) - conf = SystemConfig.get_by_key(self._db, key) - if conf: - if old_value != value: - # 假值(False/0/None/空容器)同样落库而不是删除记录: - # 读取端以「无记录」表示未配置并回落默认值,删除会使布尔开关的关闭态无法持久化 - self._stage_update(conf, {"value": value}) - return True - return None - else: - conf = SystemConfig(key=key, value=value) - self._stage_create(conf) + self._require_loaded() + with self._write_lock: + + def write(db): + """在当前事务中创建或更新配置记录。""" + conf = SystemConfig.get_by_key(db, key) + if conf: + if conf.value == value: + return None + # 假值同样是有效配置;删除记录会使读取端错误回落默认值。 + conf.value = copy.deepcopy(value) + else: + db.add(SystemConfig(key=key, value=copy.deepcopy(value))) return True - async def async_set(self, key: Union[str, SystemConfigKey], value: Any) -> Optional[bool]: - """ - 异步设置系统设置 - :param key: 配置键 - :param value: 配置值 - :return: 是否设置成功(True 成功/False 失败/None 无需更新) - """ - if isinstance(key, SystemConfigKey): - key = key.value - async with self._alock: - conf = await SystemConfig.async_get_by_key(self._db, key) - # 确定是否需要更新数据库 - needs_db_update = False - if conf: - if conf.value != value: - needs_db_update = True - else: # 记录不存在,总是需要创建/更新 - needs_db_update = True - if not needs_db_update: - # 即使数据库值相同,也要确保缓存同步 - with self._rlock: - self.__SYSTEMCONF[key] = copy.deepcopy(value) - return None - # 执行数据库更新 - if conf: - # 假值(False/0/None/空容器)同样落库而不是删除记录: - # 读取端以「无记录」表示未配置并回落默认值,删除会使布尔开关的关闭态无法持久化 - await self._stage_async_update(conf, {"value": value}) - else: - conf = SystemConfig(key=key, value=value) - await self._stage_async_create(conf) - # 数据库更新成功后,再更新缓存 - with self._rlock: - self.__SYSTEMCONF[key] = copy.deepcopy(value) - return True + result = self._execute_sync_write(write) + # 数据库操作返回时事务已经提交,读取方不会看到尚未持久化的配置。 + self._publish_value(key, value) + return result def get(self, key: Optional[Union[str, SystemConfigKey]] = None) -> Any: """ @@ -95,7 +83,8 @@ class SystemConfigOper(DbOper, metaclass=Singleton): key = key.value if not key: return self.all() - with self._rlock: + with self._snapshot_lock: + self._require_loaded() # 避免将__SYSTEMCONF内的值引用出去,会导致set时误判没有变动 return copy.deepcopy(self.__SYSTEMCONF.get(key)) @@ -107,7 +96,8 @@ class SystemConfigOper(DbOper, metaclass=Singleton): :param step: 递增步长 :return: 递增后的整数值 """ - with self._rlock: + self._require_loaded() + with self._write_lock: value = int(self.get(key) or 0) + step self.set(key, value) return value @@ -116,7 +106,8 @@ class SystemConfigOper(DbOper, metaclass=Singleton): """ 获取所有系统设置 """ - with self._rlock: + with self._snapshot_lock: + self._require_loaded() # 避免将__SYSTEMCONF内的值引用出去,会导致set时误判没有变动 return copy.deepcopy(self.__SYSTEMCONF) @@ -126,11 +117,15 @@ class SystemConfigOper(DbOper, metaclass=Singleton): """ if isinstance(key, SystemConfigKey): key = key.value - with self._rlock: - # 更新内存 - self.__SYSTEMCONF.pop(key, None) - # 写入数据库 - conf = SystemConfig.get_by_key(self._db, key) - if conf: - self._stage_delete(SystemConfig, conf.id) + self._require_loaded() + with self._write_lock: + + def delete(db): + """在当前事务中删除配置记录。""" + conf = SystemConfig.get_by_key(db, key) + if conf: + db.delete(conf) + + self._execute_sync_write(delete) + self._publish_delete(key) return True diff --git a/app/db/oper/userconfig.py b/app/db/oper/userconfig.py index 5edfc3b97..7b1656bdc 100644 --- a/app/db/oper/userconfig.py +++ b/app/db/oper/userconfig.py @@ -1,3 +1,5 @@ +import copy +import threading from typing import Any, Union, Dict, Optional from app.db.base import DbOper @@ -11,13 +13,30 @@ class UserConfigOper(DbOper, metaclass=Singleton): 用户配置管理 """ def __init__(self): - """ - 加载配置到内存 - """ + """初始化空快照,数据库加载由启动组合根显式执行。""" super().__init__() self.__USERCONF = {} - for item in UserConfig.list(self._db): - self.__set_config_cache(username=item.username, key=item.key, value=item.value) + self._snapshot_lock = threading.RLock() + self._write_lock = threading.RLock() + self._loaded = False + + def load_snapshot(self) -> None: + """从数据库加载完整用户配置,并一次性发布新的内存快照。""" + with self._write_lock: + snapshot: dict[str, dict[str, Any]] = {} + for item in UserConfig.list(self._db): + if item.username and item.key: + snapshot.setdefault(item.username, {})[item.key] = copy.deepcopy( + item.value + ) + with self._snapshot_lock: + self.__USERCONF = snapshot + self._loaded = True + + def _require_loaded(self) -> None: + """阻止消费者读取尚未完成启动加载的半成品快照。""" + if not self._loaded: + raise RuntimeError("用户配置快照尚未加载") def set(self, username: str, key: Union[str, UserConfigKey], value: Any): """ @@ -25,30 +44,43 @@ class UserConfigOper(DbOper, metaclass=Singleton): """ if isinstance(key, UserConfigKey): key = key.value - # 更新内存 - self.__set_config_cache(username=username, key=key, value=value) - # 写入数据库 - conf = UserConfig.get_by_key(db=self._db, username=username, key=key) - if conf: - if value: - self._stage_update(conf, {"value": value}) - else: - self._stage_delete(UserConfig, conf.id) - else: - conf = UserConfig(username=username, key=key, value=value) - self._stage_create(conf) + self._require_loaded() + with self._write_lock: + + def write(db): + """在当前事务中按用户配置的假值规则写入记录。""" + conf = UserConfig.get_by_key(db=db, username=username, key=key) + if conf: + if value: + conf.value = copy.deepcopy(value) + else: + db.delete(conf) + else: + db.add( + UserConfig( + username=username, + key=key, + value=copy.deepcopy(value), + ) + ) + + self._execute_sync_write(write) + # 既有运行时语义会保留刚写入的假值,即使其数据库记录被删除。 + self.__set_config_cache(username=username, key=key, value=value) def get(self, username: str, key: Optional[Union[str, UserConfigKey]] = None) -> Any: """ 获取用户配置 """ - if not username: - return self.__USERCONF - if isinstance(key, UserConfigKey): - key = key.value - if not key: - return self.__get_config_caches(username=username) - return self.__get_config_cache(username=username, key=key) + with self._snapshot_lock: + self._require_loaded() + if not username: + return copy.deepcopy(self.__USERCONF) + if isinstance(key, UserConfigKey): + key = key.value + if not key: + return copy.deepcopy(self.__get_config_caches(username=username)) + return copy.deepcopy(self.__get_config_cache(username=username, key=key)) def __set_config_cache(self, username: str, key: str, value: Any): """ @@ -56,15 +88,9 @@ class UserConfigOper(DbOper, metaclass=Singleton): """ if not username or not key: return - cache = self.__USERCONF - if not cache: - cache = {} - user_cache = cache.get(username) - if not user_cache: - user_cache = {} - cache[username] = user_cache - user_cache[key] = value - self.__USERCONF = cache + with self._snapshot_lock: + user_cache = self.__USERCONF.setdefault(username, {}) + user_cache[key] = copy.deepcopy(value) def __get_config_caches(self, username: str) -> Optional[Dict[str, Any]]: """ diff --git a/app/db/worker.py b/app/db/worker.py new file mode 100644 index 000000000..b99f0883c --- /dev/null +++ b/app/db/worker.py @@ -0,0 +1,234 @@ +"""同步数据库短事务的异步执行器。""" + +from __future__ import annotations + +import asyncio +import threading +import time +from concurrent.futures import Future, ThreadPoolExecutor +from contextvars import copy_context +from dataclasses import dataclass +from typing import Callable, TypeVar + +from app.runtime.observability import record_metric + + +T = TypeVar("T") + + +class DatabaseWorkerClosedError(RuntimeError): + """数据库执行器尚未启动或已经停止。""" + + +class DatabaseWorkerOverloadedError(RuntimeError): + """数据库执行器的运行与排队容量已经用尽。""" + + +@dataclass(frozen=True, slots=True) +class DatabaseWorkerStats: + """数据库执行器当前的容量和任务数量。""" + + max_workers: int + capacity: int + queued: int + running: int + rejected: int + closing: bool + + +@dataclass(slots=True) +class _WorkItem: + """记录一个任务的排队时间和执行状态。""" + + submitted_at: float + started_at: float | None = None + + +class DatabaseWorker: + """以有限线程和队列执行不能原生异步化的数据库短事务。""" + + def __init__(self, *, max_workers: int = 4, capacity: int = 32) -> None: + """保存容量配置,线程只在显式启动后创建。""" + if max_workers < 1: + raise ValueError("数据库 worker 线程数必须大于 0") + if capacity < max_workers: + raise ValueError("数据库 worker 总容量不能小于线程数") + self._max_workers = max_workers + self._capacity = capacity + self._state_lock = threading.Lock() + self._loop: asyncio.AbstractEventLoop | None = None + self._executor: ThreadPoolExecutor | None = None + self._futures: dict[ + Future[object], tuple[asyncio.Future[object], _WorkItem] + ] = {} + self._queued = 0 + self._running = 0 + self._rejected = 0 + self._closing = False + + async def start(self) -> None: + """绑定当前事件循环并准备专属线程池。""" + if self._executor is not None: + if self._loop is not asyncio.get_running_loop(): + raise RuntimeError("数据库 worker 不能跨事件循环复用") + return + self._loop = asyncio.get_running_loop() + self._executor = ThreadPoolExecutor( + max_workers=self._max_workers, + thread_name_prefix="moviepilot-db", + ) + self._closing = False + self._record_depth() + + def snapshot(self) -> DatabaseWorkerStats: + """返回无需访问任务对象的低基数运行快照。""" + with self._state_lock: + return DatabaseWorkerStats( + max_workers=self._max_workers, + capacity=self._capacity, + queued=self._queued, + running=self._running, + rejected=self._rejected, + closing=self._closing, + ) + + async def run(self, operation: Callable[[], T]) -> T: + """执行短事务,取消时仍等待已开始的事务取得最终结果。""" + loop = asyncio.get_running_loop() + executor = self._executor + if executor is None or self._loop is not loop or self._closing: + raise DatabaseWorkerClosedError("数据库 worker 当前不可接收任务") + + with self._state_lock: + if self._queued + self._running >= self._capacity: + self._rejected += 1 + record_metric("db.worker.rejected") + raise DatabaseWorkerOverloadedError( + f"数据库 worker 容量已用尽(上限 {self._capacity})" + ) + self._queued += 1 + + item = _WorkItem(submitted_at=time.perf_counter()) + context = copy_context() + try: + future = executor.submit(self._execute, item, context.run, operation) + except BaseException: + with self._state_lock: + self._queued -= 1 + self._record_depth() + raise + + wrapped = asyncio.wrap_future(future, loop=loop) + with self._state_lock: + self._futures[future] = (wrapped, item) + future.add_done_callback( + lambda completed: self._schedule_completion(completed, item) + ) + self._record_depth() + + try: + return await asyncio.shield(wrapped) + except asyncio.CancelledError: + if not future.cancel(): + await self._wait_until_done(wrapped) + if wrapped.done() and not wrapped.cancelled(): + wrapped.exception() + raise + + def _execute( + self, + item: _WorkItem, + context_run: Callable[..., T], + operation: Callable[[], T], + ) -> T: + """在线程中标记任务开始并保留提交时的上下文。""" + item.started_at = time.perf_counter() + loop = self._loop + if loop is not None: + try: + loop.call_soon_threadsafe(self._mark_running, item) + except RuntimeError: + pass + return context_run(operation) + + def _mark_running(self, item: _WorkItem) -> None: + """把任务从排队状态移入运行状态。""" + with self._state_lock: + self._queued -= 1 + self._running += 1 + started_at = item.started_at or time.perf_counter() + record_metric("db.worker.wait", started_at - item.submitted_at) + self._record_depth() + + def _schedule_completion( + self, + future: Future[object], + item: _WorkItem, + ) -> None: + """把线程完成通知安全地回投到所属事件循环。""" + loop = self._loop + if loop is None: + return + try: + loop.call_soon_threadsafe(self._complete, future, item) + except RuntimeError: + pass + + def _complete(self, future: Future[object], item: _WorkItem) -> None: + """释放 admission,并记录任务的最终结果。""" + with self._state_lock: + self._futures.pop(future, None) + if future.running() or future.done() and not future.cancelled(): + self._running -= 1 + else: + self._queued -= 1 + outcome = "cancelled" if future.cancelled() else "success" + if not future.cancelled(): + try: + future.result() + except BaseException: + outcome = "error" + started_at = item.started_at or item.submitted_at + record_metric( + "db.worker.duration", + time.perf_counter() - started_at, + outcome=outcome, + ) + self._record_depth() + + async def _wait_until_done(self, future: asyncio.Future[object]) -> None: + """忽略后续取消请求,直到线程内事务结束。""" + while not future.done(): + try: + await asyncio.shield(future) + except asyncio.CancelledError: + continue + except BaseException: + break + + async def shutdown(self) -> None: + """拒绝新任务,取消排队任务并等待运行中事务结束。""" + executor = self._executor + if executor is None: + return + if self._loop is not asyncio.get_running_loop(): + raise RuntimeError("数据库 worker 必须在所属事件循环中停止") + self._closing = True + with self._state_lock: + futures = tuple(self._futures.items()) + for future, _state in futures: + future.cancel() + for future, (wrapped, _item) in futures: + if not future.cancelled(): + await self._wait_until_done(wrapped) + executor.shutdown(wait=True, cancel_futures=True) + while self.snapshot().queued or self.snapshot().running: + await asyncio.sleep(0) + self._executor = None + 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) diff --git a/app/runtime/observability/__init__.py b/app/runtime/observability/__init__.py index cf7f198b3..1da6e1d10 100644 --- a/app/runtime/observability/__init__.py +++ b/app/runtime/observability/__init__.py @@ -35,6 +35,11 @@ METRIC_SPECS = { MetricSpec("db.pool.wait", MetricKind.HISTOGRAM, frozenset({"backend", "outcome"})), MetricSpec("db.pool.checked_out", MetricKind.GAUGE, frozenset({"backend"})), MetricSpec("db.pool.timeout", MetricKind.COUNTER, frozenset({"backend"})), + MetricSpec("db.worker.wait", MetricKind.HISTOGRAM, frozenset()), + MetricSpec("db.worker.duration", MetricKind.HISTOGRAM, frozenset({"outcome"})), + MetricSpec("db.worker.queue.depth", MetricKind.GAUGE, frozenset()), + MetricSpec("db.worker.active", MetricKind.GAUGE, frozenset()), + MetricSpec("db.worker.rejected", MetricKind.COUNTER, frozenset()), MetricSpec("event.queue.depth", MetricKind.GAUGE, frozenset({"delivery"})), MetricSpec("event.handler.duration", MetricKind.HISTOGRAM, frozenset({"event_type", "handler_type", "outcome"})), MetricSpec("module.provider.duration", MetricKind.HISTOGRAM, frozenset({"method", "provider_type", "outcome"})), diff --git a/app/startup/lifecycle/__init__.py b/app/startup/lifecycle/__init__.py index 4f3bf3bd6..68705f3b6 100644 --- a/app/startup/lifecycle/__init__.py +++ b/app/startup/lifecycle/__init__.py @@ -130,7 +130,13 @@ async def run_startup_step( async def initialize_modules_component(app: FastAPI) -> None: """启动模块并把其类型化运行时发布到当前 FastAPI AppState。""" - runtime = await init_modules() + try: + runtime = await init_modules() + except BaseException: + from app.startup.modules_initializer import stop_database_worker + + await stop_database_worker() + raise if runtime is not None: app.state.host_runtime = runtime diff --git a/app/startup/modules_initializer.py b/app/startup/modules_initializer.py index 52260d4d7..598da1dc7 100644 --- a/app/startup/modules_initializer.py +++ b/app/startup/modules_initializer.py @@ -89,6 +89,7 @@ from app.db.session import ( get_async_db, get_db, ) +from app.db.worker import DatabaseWorker from app.db.uow import ( SqlAlchemyAsyncUnitOfWork, SqlAlchemyUnitOfWork, @@ -157,6 +158,40 @@ from app.runtime.extensions.service_config import ( ) +_database_worker: DatabaseWorker | None = None + + +async def stop_database_worker() -> None: + """停止当前进程的数据库短事务 worker。""" + global _database_worker + worker = _database_worker + _database_worker = None + if worker is not None: + await worker.shutdown() + + +async def _initialize_configuration_services( + database_worker: DatabaseWorker, +) -> None: + """加载完整配置快照后发布系统与用户配置服务。""" + system_config = SystemConfigOper() + user_config = UserConfigOper() + await database_worker.run(system_config.load_snapshot) + await database_worker.run(user_config.load_snapshot) + configure_system_config( + SystemConfigService( + repository=system_config, + async_executor=database_worker, + ) + ) + configure_user_configuration( + UserConfigurationService( + repository=user_config, + async_executor=database_worker, + ) + ) + + async def _async_get_subscribe(subscribe_id: int): """通过数据库操作器异步读取订阅,供服务端共享用例使用。""" return await SubscribeOper().async_get(subscribe_id) @@ -506,6 +541,7 @@ async def stop_modules(): await run_step("消息服务", stop_message) await run_step("Redis缓存连接", lambda: RedisHelper().close()) await run_step("异步Redis缓存连接", lambda: AsyncRedisHelper().close()) + await run_step("数据库任务", stop_database_worker) await run_step("数据库连接", close_database) await run_step("前端服务", stop_frontend) await run_step("临时文件", clear_temp) @@ -515,6 +551,7 @@ async def init_modules() -> HostRuntime: """ 启动模块并返回本次 lifespan 唯一的类型化 HostRuntime。 """ + global _database_worker # 兼容 Oper 的无 Session 写入口仍由组合根持有事务,避免模型恢复自动提交。 transaction_runner = TransactionalWriteRunner( sync_session=SessionFactory, @@ -524,6 +561,14 @@ async def init_modules() -> HostRuntime: sync=transaction_runner.sync, async_=transaction_runner.async_, ) + database_worker = DatabaseWorker() + await database_worker.start() + _database_worker = database_worker + try: + await _initialize_configuration_services(database_worker) + except BaseException: + await stop_database_worker() + raise # 数据访问能力统一在启动组合根注入,Runtime 和 Adapter 不再直接依赖 Oper。 api_data = ApiDataPorts( sync_session=get_db, @@ -599,8 +644,6 @@ async def init_modules() -> HostRuntime: configure_runtime_settings(host_runtime.settings) configure_runtime_setting_provider(lambda key: getattr(settings, key)) configure_token_runtime_config(lambda: build_token_runtime_config(settings)) - # 先发布系统配置服务,后续启动组合步骤统一复用同一配置端口。 - configure_system_config(SystemConfigService(repository=SystemConfigOper())) # 旧 app.api.data 导入只保留 ABI 转发,正式 API 依赖全部读取 HostRuntime。 configure_api_data_runtime(api_data) configure_runtime_data_providers() @@ -640,7 +683,6 @@ async def init_modules() -> HostRuntime: ) ) configure_passkey_service(PasskeyService(repository=PassKeyOper())) - configure_user_configuration(UserConfigurationService(repository=UserConfigOper())) configure_transfer_history_provider(lambda: TransferHistoryOper()) configure_site_query_service(SiteQueryService(repository=SiteOper())) configure_site_health_service(SiteHealthService(repository=SiteOper())) diff --git a/app/testing/bootstrap.py b/app/testing/bootstrap.py index 67c453a95..5c4e7a030 100644 --- a/app/testing/bootstrap.py +++ b/app/testing/bootstrap.py @@ -171,7 +171,7 @@ def ensure_optional_stub(name: str, **attrs) -> None: def prepare_backend() -> None: - """隔离 CONFIG_DIR、补 sites 垫片并建表(后端须已在 ``sys.path`` 上)。 + """隔离 CONFIG_DIR、补 sites 垫片、建表并加载配置快照。 主程序中后端即当前包;插件仓由其 ``tests/_bootstrap.py`` shim 在 import 本模块前 先把后端目录注入 ``sys.path``。顺序固定:先隔离 CONFIG_DIR,再补 ``app.application.site.sites`` 垫片, @@ -182,6 +182,11 @@ def prepare_backend() -> None: ensure_sites_stub() from app.startup.database_initializer import init_db init_db() + from app.db.oper.systemconfig import SystemConfigOper + from app.db.oper.userconfig import UserConfigOper + + SystemConfigOper().load_snapshot() + UserConfigOper().load_snapshot() # 缓存装饰器在测试模块导入时即创建后端,先装配隔离配置对应的适配器。 from app.startup.cache_initializer import configure_cache_dependencies configure_cache_dependencies() diff --git a/scripts/local_setup.py b/scripts/local_setup.py index 8035e10f8..86d19ced8 100644 --- a/scripts/local_setup.py +++ b/scripts/local_setup.py @@ -2400,6 +2400,7 @@ def _apply_local_system_config_inner(config_payload: dict[str, Any]) -> None: print_step(f"超级管理员初始密码:{generated_password}") system_config = SystemConfigOper() + system_config.load_snapshot() directory_items = config_payload.get("directories") or [] if directory_items: current_directories = system_config.get(SystemConfigKey.Directories) or [] diff --git a/tests/conftest.py b/tests/conftest.py index 8584ed70c..4092f5a7e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -3,6 +3,7 @@ 引导与网络守卫均复用 ``app/testing`` 的共享 harness(与插件仓 conftest 同源), 引导逻辑只在 ``app/testing`` 维护一处。 """ +import asyncio import sys import pytest @@ -20,6 +21,14 @@ prepare_backend() from app.testing.network_guard import block_real_network # noqa: E402,F401 +class _TestDatabaseExecutor: + """让绕过完整 lifespan 的测试仍通过线程执行同步数据库写入。""" + + async def run(self, operation): + """在线程中执行测试事务。""" + return await asyncio.to_thread(operation) + + @pytest.fixture(autouse=True) def configure_plugin_system_services(): """为绕过完整启动流程的单元测试装配真实插件系统适配器。""" @@ -61,6 +70,11 @@ def configure_plugin_system_services(): configure_transaction_runners, ) from app.db.oper.systemconfig import SystemConfigOper + from app.db.oper.userconfig import UserConfigOper + from app.application.security.userconfig import ( + UserConfigurationService, + configure_user_configuration, + ) configure_token_codec(create_access_token, decode_access_token) configure_runtime_configuration( @@ -73,7 +87,23 @@ def configure_plugin_system_services(): configure_runtime_settings(RuntimeSettingsService(settings)) configure_runtime_setting_provider(lambda key: getattr(settings, key)) configure_token_runtime_config(lambda: build_token_runtime_config(settings)) - configure_system_config(SystemConfigService(repository=SystemConfigOper())) + database_executor = _TestDatabaseExecutor() + system_config = SystemConfigOper() + system_config.load_snapshot() + user_config = UserConfigOper() + user_config.load_snapshot() + configure_system_config( + SystemConfigService( + repository=system_config, + async_executor=database_executor, + ) + ) + configure_user_configuration( + UserConfigurationService( + repository=user_config, + async_executor=database_executor, + ) + ) configure_transfer_retry_config( lambda: TransferRetryConfig( max_failed_retries=settings.TRANSFER_MAX_FAILED_RETRIES, diff --git a/tests/test_agent_lifecycle.py b/tests/test_agent_lifecycle.py index 2ec843909..fe4e42b76 100644 --- a/tests/test_agent_lifecycle.py +++ b/tests/test_agent_lifecycle.py @@ -138,7 +138,10 @@ async def test_agent_initialization_failure_does_not_stop_module_startup( monkeypatch.setattr(modules_initializer, "start_frontend", start_frontend) monkeypatch.setattr(modules_initializer, "check_auth", check_auth) - await modules_initializer.init_modules() + try: + await modules_initializer.init_modules() + finally: + await modules_initializer.stop_database_worker() manager.initialize.assert_awaited_once_with() start_frontend.assert_called_once_with() diff --git a/tests/test_cache_system.py b/tests/test_cache_system.py index 95016906a..aa38f68e2 100644 --- a/tests/test_cache_system.py +++ b/tests/test_cache_system.py @@ -164,7 +164,13 @@ def test_init_modules_does_not_clear_package_tool_cache(monkeypatch): monkeypatch.setattr(modules_initializer, "start_frontend", lambda: None) monkeypatch.setattr(modules_initializer, "check_auth", lambda: None) - asyncio.run(modules_initializer.init_modules()) + async def initialize_modules() -> None: + try: + await modules_initializer.init_modules() + finally: + await modules_initializer.stop_database_worker() + + asyncio.run(initialize_modules()) assert called is False init_agent.assert_awaited_once_with() diff --git a/tests/test_configuration_initializer.py b/tests/test_configuration_initializer.py new file mode 100644 index 000000000..3f3be0f5a --- /dev/null +++ b/tests/test_configuration_initializer.py @@ -0,0 +1,104 @@ +"""配置快照启动顺序测试。""" + +from unittest.mock import AsyncMock + +import pytest + +from app.startup import modules_initializer +from app.startup.lifecycle import initialize_modules_component + + +class _InlineWorker: + """按提交顺序执行配置加载操作。""" + + async def run(self, operation): + """执行并返回操作结果。""" + return operation() + + +@pytest.mark.asyncio +async def test_configuration_services_publish_after_both_snapshots_load( + monkeypatch, +) -> None: + """两个完整快照加载成功前不发布任一配置服务。""" + events = [] + + class _SystemConfig: + def load_snapshot(self): + events.append("load-system") + + class _UserConfig: + def load_snapshot(self): + events.append("load-user") + + monkeypatch.setattr(modules_initializer, "SystemConfigOper", _SystemConfig) + monkeypatch.setattr(modules_initializer, "UserConfigOper", _UserConfig) + monkeypatch.setattr( + modules_initializer, + "configure_system_config", + lambda _service: events.append("publish-system"), + ) + monkeypatch.setattr( + modules_initializer, + "configure_user_configuration", + lambda _service: events.append("publish-user"), + ) + + await modules_initializer._initialize_configuration_services(_InlineWorker()) + + assert events == [ + "load-system", + "load-user", + "publish-system", + "publish-user", + ] + + +@pytest.mark.asyncio +async def test_configuration_load_failure_does_not_publish_partial_service( + monkeypatch, +) -> None: + """任一快照加载失败时不发布半套配置服务。""" + published = [] + + class _SystemConfig: + def load_snapshot(self): + return None + + class _UserConfig: + def load_snapshot(self): + raise RuntimeError("load failed") + + monkeypatch.setattr(modules_initializer, "SystemConfigOper", _SystemConfig) + monkeypatch.setattr(modules_initializer, "UserConfigOper", _UserConfig) + monkeypatch.setattr( + modules_initializer, + "configure_system_config", + lambda service: published.append(service), + ) + monkeypatch.setattr( + modules_initializer, + "configure_user_configuration", + lambda service: published.append(service), + ) + + with pytest.raises(RuntimeError, match="load failed"): + await modules_initializer._initialize_configuration_services(_InlineWorker()) + + assert published == [] + + +@pytest.mark.asyncio +async def test_modules_startup_failure_stops_database_worker(monkeypatch) -> None: + """模块启动失败时立即关闭已创建的数据库 worker。""" + monkeypatch.setattr( + "app.startup.lifecycle.init_modules", + AsyncMock(side_effect=RuntimeError("startup failed")), + ) + stop_worker = AsyncMock() + monkeypatch.setattr(modules_initializer, "stop_database_worker", stop_worker) + + with pytest.raises(RuntimeError, match="startup failed"): + await initialize_modules_component(object()) + + stop_worker.assert_awaited_once_with() diff --git a/tests/test_configuration_ports.py b/tests/test_configuration_ports.py index ac0f1b46e..b45a6cb45 100644 --- a/tests/test_configuration_ports.py +++ b/tests/test_configuration_ports.py @@ -2,7 +2,7 @@ import asyncio from dataclasses import FrozenInstanceError -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import MagicMock import pytest @@ -19,6 +19,15 @@ from app.application.configuration import ( get_api_runtime_config_snapshot, get_transfer_retry_config, ) +from app.application.security.userconfig import UserConfigurationService + + +class _InlineDatabaseExecutor: + """同步执行测试操作,并保留异步应用端口的调用形态。""" + + async def run(self, operation): + """执行并返回操作结果。""" + return operation() class _MutableSettings: @@ -66,23 +75,44 @@ def test_system_config_service_supports_separate_reader_and_writer() -> None: """应用服务可以分别注入只读与写入适配器。""" reader = MagicMock() reader.get.return_value = "old" - reader.async_get = AsyncMock(return_value="async-old") writer = MagicMock() writer.set.return_value = True - writer.async_set = AsyncMock(return_value=True) - service = SystemConfigService(reader=reader, writer=writer) + service = SystemConfigService( + reader=reader, + writer=writer, + async_executor=_InlineDatabaseExecutor(), + ) assert service.get("key") == "old" assert service.set("key", "new") is True - assert asyncio.run(service.async_get("key")) == "async-old" assert asyncio.run(service.async_set("key", "new")) is True service.delete("key") reader.get.assert_called_once_with("key") - writer.set.assert_called_once_with("key", "new") + assert writer.set.call_args_list == [ + (("key", "new"), {}), + (("key", "new"), {}), + ] writer.delete.assert_called_once_with("key") +def test_user_configuration_service_supports_sync_and_async_writes() -> None: + """用户配置服务的同步与异步入口执行同一仓储方法。""" + repository = MagicMock() + repository.set.return_value = True + service = UserConfigurationService( + repository, + async_executor=_InlineDatabaseExecutor(), + ) + + assert service.set("alice", "theme", "dark") is True + assert asyncio.run(service.async_set("alice", "theme", "light")) is True + assert repository.set.call_args_list == [ + ((), {"username": "alice", "key": "theme", "value": "dark"}), + ((), {"username": "alice", "key": "theme", "value": "light"}), + ] + + def test_transfer_retry_provider_returns_frozen_snapshot_per_call() -> None: """配置工厂在每次用例入口创建新快照,旧快照不受 reload 后状态影响。""" state = {"value": 2} diff --git a/tests/test_database_worker.py b/tests/test_database_worker.py new file mode 100644 index 000000000..563079d29 --- /dev/null +++ b/tests/test_database_worker.py @@ -0,0 +1,140 @@ +"""数据库短事务 worker 的容量、取消与关闭合同测试。""" + +import asyncio +import threading + +import pytest + +from app.db.worker import ( + DatabaseWorker, + DatabaseWorkerClosedError, + DatabaseWorkerOverloadedError, +) + + +@pytest.mark.asyncio +async def test_worker_requires_explicit_start() -> None: + """构造对象不会隐式创建可执行线程池。""" + worker = DatabaseWorker(max_workers=1, capacity=1) + + with pytest.raises(DatabaseWorkerClosedError): + await worker.run(lambda: None) + + +@pytest.mark.asyncio +async def test_worker_rejects_work_beyond_running_and_queue_capacity() -> None: + """运行与排队任务达到总容量后立即拒绝新任务。""" + worker = DatabaseWorker(max_workers=1, capacity=2) + await worker.start() + started = threading.Event() + release = threading.Event() + + def block() -> None: + started.set() + release.wait(1) + + running = asyncio.create_task(worker.run(block)) + await asyncio.to_thread(started.wait) + queued = asyncio.create_task(worker.run(lambda: None)) + await asyncio.sleep(0) + + with pytest.raises(DatabaseWorkerOverloadedError): + await worker.run(lambda: None) + + assert worker.snapshot().running == 1 + assert worker.snapshot().queued == 1 + assert worker.snapshot().rejected == 1 + + release.set() + await asyncio.gather(running, queued) + await worker.shutdown() + + +@pytest.mark.asyncio +async def test_cancelling_queued_work_prevents_execution() -> None: + """尚未取得线程的任务取消后不得执行数据库操作。""" + worker = DatabaseWorker(max_workers=1, capacity=2) + await worker.start() + started = threading.Event() + release = threading.Event() + queued_executed = threading.Event() + + def block() -> None: + started.set() + release.wait(1) + + running = asyncio.create_task(worker.run(block)) + await asyncio.to_thread(started.wait) + queued = asyncio.create_task(worker.run(queued_executed.set)) + await asyncio.sleep(0) + queued.cancel() + with pytest.raises(asyncio.CancelledError): + await queued + + release.set() + await running + await worker.shutdown() + + assert queued_executed.is_set() is False + assert worker.snapshot().queued == 0 + assert worker.snapshot().running == 0 + + +@pytest.mark.asyncio +async def test_cancelling_running_work_waits_for_transaction_terminal_state() -> None: + """线程内操作开始后,取消结果必须晚于操作的最终状态。""" + worker = DatabaseWorker(max_workers=1, capacity=1) + await worker.start() + started = threading.Event() + release = threading.Event() + completed = threading.Event() + + def operation() -> None: + started.set() + release.wait(1) + completed.set() + + task = asyncio.create_task(worker.run(operation)) + await asyncio.to_thread(started.wait) + task.cancel() + await asyncio.sleep(0.01) + + assert task.done() is False + assert completed.is_set() is False + + release.set() + with pytest.raises(asyncio.CancelledError): + await task + + assert completed.is_set() is True + await worker.shutdown() + + +@pytest.mark.asyncio +async def test_shutdown_rejects_new_work_and_waits_for_running_work() -> None: + """关闭期间不接收新任务,并等待已开始的操作结束。""" + worker = DatabaseWorker(max_workers=1, capacity=2) + 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) + shutdown = asyncio.create_task(worker.shutdown()) + await asyncio.sleep(0) + + with pytest.raises(DatabaseWorkerClosedError): + await worker.run(lambda: None) + assert shutdown.done() is False + + release.set() + await running + await shutdown + + assert worker.snapshot().closing is True + assert worker.snapshot().queued == 0 + assert worker.snapshot().running == 0 diff --git a/tests/test_systemconfig_oper.py b/tests/test_systemconfig_oper.py index 8d80cf415..5e6b77d61 100644 --- a/tests/test_systemconfig_oper.py +++ b/tests/test_systemconfig_oper.py @@ -1,6 +1,7 @@ import threading import uuid from concurrent.futures import ThreadPoolExecutor +from types import SimpleNamespace import pytest @@ -16,15 +17,105 @@ def _unique_key() -> str: def _fresh_oper() -> SystemConfigOper: - """重置单例并从数据库重新加载配置缓存。""" + """重置单例并显式加载数据库配置快照。""" Singleton._instances.pop((SystemConfigOper, (), frozenset()), None) - return SystemConfigOper() + oper = SystemConfigOper() + oper.load_snapshot() + return oper + + +def test_constructor_does_not_query_database(monkeypatch): + """构造配置对象时不打开数据库会话。""" + Singleton._instances.pop((SystemConfigOper, (), frozenset()), None) + monkeypatch.setattr( + SystemConfig, + "list", + lambda _db: pytest.fail("构造阶段不应查询数据库"), + ) + + oper = SystemConfigOper() + + with pytest.raises(RuntimeError, match="快照尚未加载"): + oper.get("key") + + +def test_load_snapshot_publishes_complete_dictionary(monkeypatch): + """重新加载期间读取方只会看到完整旧快照或完整新快照。""" + Singleton._instances.pop((SystemConfigOper, (), frozenset()), None) + oper = SystemConfigOper() + values = [SimpleNamespace(key="key", value="old")] + entered = threading.Event() + release = threading.Event() + + monkeypatch.setattr(SystemConfig, "list", lambda _db: values) + oper.load_snapshot() + + def load_new_snapshot(_db): + entered.set() + release.wait(1) + return [SimpleNamespace(key="key", value="new")] + + monkeypatch.setattr(SystemConfig, "list", load_new_snapshot) + thread = threading.Thread(target=oper.load_snapshot) + thread.start() + assert entered.wait(1) + + assert oper.get("key") == "old" + + release.set() + thread.join(1) + assert thread.is_alive() is False + assert oper.get("key") == "new" + + +def test_read_does_not_wait_for_slow_write_transaction(monkeypatch): + """数据库写入期间,内存读取仍返回最近一次已提交值。""" + oper = _fresh_oper() + key = _unique_key() + oper.set(key, "old") + entered = threading.Event() + release = threading.Event() + + def slow_write(_operation): + entered.set() + release.wait(1) + return True + + monkeypatch.setattr(oper, "_execute_sync_write", slow_write) + thread = threading.Thread(target=lambda: oper.set(key, "new")) + thread.start() + assert entered.wait(1) + + assert oper.get(key) == "old" + + release.set() + thread.join(1) + assert thread.is_alive() is False + assert oper.get(key) == "new" + + +def test_failed_write_keeps_committed_snapshot(monkeypatch): + """事务失败时内存快照保持最近一次已提交值。""" + oper = _fresh_oper() + key = _unique_key() + oper.set(key, "old") + + def fail_write(_operation): + raise RuntimeError("write failed") + + monkeypatch.setattr(oper, "_execute_sync_write", fail_write) + + with pytest.raises(RuntimeError, match="write failed"): + oper.set(key, "new") + assert oper.get(key) == "old" def test_increment_serializes_concurrent_counter_updates(monkeypatch): """并发递增系统计数时不应丢失更新。""" oper = object.__new__(SystemConfigOper) - oper._rlock = threading.RLock() + oper._snapshot_lock = threading.RLock() + oper._write_lock = threading.RLock() + oper._loaded = True stored_value = {"value": 0} monkeypatch.setattr(oper, "get", lambda _key: stored_value["value"]) @@ -51,7 +142,9 @@ def test_increment_serializes_concurrent_counter_updates(monkeypatch): def test_increment_supports_custom_step(monkeypatch): """整数系统计数应支持指定递增步长。""" oper = object.__new__(SystemConfigOper) - oper._rlock = threading.RLock() + oper._snapshot_lock = threading.RLock() + oper._write_lock = threading.RLock() + oper._loaded = True stored_value = {"value": 4} monkeypatch.setattr(oper, "get", lambda _key: stored_value["value"]) @@ -95,7 +188,17 @@ async def test_async_set_persists_falsy_value_on_existing_record(): oper = _fresh_oper() oper.set(key, True) - assert await oper.async_set(key, False) is True + from app.application.configuration import SystemConfigService + + class _InlineDatabaseExecutor: + async def run(self, operation): + return operation() + + service = SystemConfigService( + repository=oper, + async_executor=_InlineDatabaseExecutor(), + ) + assert await service.async_set(key, False) is True assert oper.get(key) is False assert SystemConfig.get_by_key(oper._db, key).value is False @@ -106,7 +209,17 @@ async def test_async_set_creates_record_for_falsy_value(): key = _unique_key() oper = _fresh_oper() - assert await oper.async_set(key, 0) is True + from app.application.configuration import SystemConfigService + + class _InlineDatabaseExecutor: + async def run(self, operation): + return operation() + + service = SystemConfigService( + repository=oper, + async_executor=_InlineDatabaseExecutor(), + ) + assert await service.async_set(key, 0) is True assert oper.get(key) == 0 assert SystemConfig.get_by_key(oper._db, key).value == 0 diff --git a/tests/test_user_config_endpoint.py b/tests/test_user_config_endpoint.py new file mode 100644 index 000000000..112993435 --- /dev/null +++ b/tests/test_user_config_endpoint.py @@ -0,0 +1,54 @@ +"""用户配置 API 的异步应用端口测试。""" + +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from app.api.endpoints import user as user_endpoint + + +@pytest.mark.asyncio +async def test_set_config_waits_for_async_configuration_write(monkeypatch) -> None: + """更新接口等待用户配置事务完成后再返回成功。""" + service = MagicMock() + service.async_set = AsyncMock(return_value=True) + monkeypatch.setattr( + user_endpoint, + "get_configured_user_configuration", + lambda: service, + ) + + response = await user_endpoint.set_config( + "theme", + "dark", + current_user=SimpleNamespace(name="alice"), + ) + + assert response.success is True + service.async_set.assert_awaited_once_with( + username="alice", + key="theme", + value="dark", + ) + + +@pytest.mark.asyncio +async def test_get_config_reads_loaded_snapshot(monkeypatch) -> None: + """查询接口直接读取已加载的用户配置快照。""" + service = MagicMock() + service.get.return_value = "dark" + monkeypatch.setattr( + user_endpoint, + "get_configured_user_configuration", + lambda: service, + ) + + response = await user_endpoint.get_config( + "theme", + current_user=SimpleNamespace(name="alice"), + ) + + assert response.success is True + assert response.data == {"value": "dark"} + service.get.assert_called_once_with(username="alice", key="theme") diff --git a/tests/test_userconfig_oper.py b/tests/test_userconfig_oper.py new file mode 100644 index 000000000..1221307b4 --- /dev/null +++ b/tests/test_userconfig_oper.py @@ -0,0 +1,93 @@ +"""用户配置快照与异步写入合同测试。""" + +import asyncio +import threading +from types import SimpleNamespace + +import pytest + +from app.application.security.userconfig import UserConfigurationService +from app.db.models.userconfig import UserConfig +from app.db.oper.userconfig import UserConfigOper +from app.foundation.singleton import Singleton + + +class _ThreadDatabaseExecutor: + """在线程中执行测试事务。""" + + async def run(self, operation): + """执行并返回操作结果。""" + return await asyncio.to_thread(operation) + + +def _fresh_oper() -> UserConfigOper: + """重置单例并显式加载用户配置快照。""" + Singleton._instances.pop((UserConfigOper, (), frozenset()), None) + oper = UserConfigOper() + oper.load_snapshot() + return oper + + +def test_constructor_does_not_query_database(monkeypatch): + """构造用户配置对象时不打开数据库会话。""" + Singleton._instances.pop((UserConfigOper, (), frozenset()), None) + monkeypatch.setattr( + UserConfig, + "list", + lambda _db: pytest.fail("构造阶段不应查询数据库"), + ) + + oper = UserConfigOper() + + with pytest.raises(RuntimeError, match="快照尚未加载"): + oper.get("alice", "theme") + + +def test_load_snapshot_publishes_complete_dictionary(monkeypatch): + """重新加载期间读取方不会看到逐项构造的用户配置。""" + Singleton._instances.pop((UserConfigOper, (), frozenset()), None) + oper = UserConfigOper() + monkeypatch.setattr( + UserConfig, + "list", + lambda _db: [SimpleNamespace(username="alice", key="theme", value="old")], + ) + oper.load_snapshot() + entered = threading.Event() + release = threading.Event() + + def load_new_snapshot(_db): + entered.set() + release.wait(1) + return [SimpleNamespace(username="alice", key="theme", value="new")] + + monkeypatch.setattr(UserConfig, "list", load_new_snapshot) + thread = threading.Thread(target=oper.load_snapshot) + thread.start() + assert entered.wait(1) + + assert oper.get("alice", "theme") == "old" + + release.set() + thread.join(1) + assert thread.is_alive() is False + assert oper.get("alice", "theme") == "new" + + +@pytest.mark.asyncio +async def test_async_write_uses_same_repository_rule() -> None: + """异步入口提交后同步读取立即看到相同结果。""" + oper = _fresh_oper() + service = UserConfigurationService( + oper, + async_executor=_ThreadDatabaseExecutor(), + ) + + await service.async_set("async-user", "theme", "dark") + + assert service.get("async-user", "theme") == "dark" + assert UserConfig.get_by_key( + oper._db, + username="async-user", + key="theme", + ).value == "dark"