From 8a11214a43fddb75b3a09278b8361f15582a7278 Mon Sep 17 00:00:00 2001 From: Aqr-K <95741669+Aqr-K@users.noreply.github.com> Date: Sat, 15 Aug 2026 06:58:38 +0800 Subject: [PATCH] =?UTF-8?q?refactor(db):=20=E4=BF=AE=E5=A4=8D=E5=BC=82?= =?UTF-8?q?=E6=AD=A5=E8=BF=9E=E6=8E=A5=E6=B1=A0=E6=97=A0=E7=95=8C=E5=A2=9E?= =?UTF-8?q?=E9=95=BF=EF=BC=8C=E5=B9=B6=E5=AE=8C=E6=88=90=20SQLAlchemy=202.?= =?UTF-8?q?0=20=E8=BF=81=E7=A7=BB=E4=B8=8E=E5=88=86=E5=B1=82=E5=BD=92?= =?UTF-8?q?=E4=BD=8D=20(#6320)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 3 + app/adapters/external/market.py | 2 +- app/adapters/external/server.py | 9 +- app/agent/llm/provider.py | 2 +- app/agent/mcp.py | 2 +- app/agent/memory/__init__.py | 2 +- app/agent/orchestrator.py | 6 +- app/agent/tools/impl/_filter_rule_utils.py | 4 +- app/agent/tools/impl/_plugin_tool_utils.py | 2 +- app/agent/tools/impl/add_download_tasks.py | 2 +- app/agent/tools/impl/add_subscribe.py | 2 +- app/agent/tools/impl/create_agent_task.py | 4 +- app/agent/tools/impl/delete_agent_task.py | 2 +- .../tools/impl/delete_download_history.py | 2 +- app/agent/tools/impl/delete_subscribe.py | 2 +- .../tools/impl/delete_transfer_history.py | 2 +- app/agent/tools/impl/query_agent_tasks.py | 2 +- .../tools/impl/query_custom_identifiers.py | 2 +- app/agent/tools/impl/query_download_tasks.py | 2 +- app/agent/tools/impl/query_downloaders.py | 2 +- app/agent/tools/impl/query_plugin_data.py | 2 +- app/agent/tools/impl/query_site_userdata.py | 2 +- app/agent/tools/impl/query_sites.py | 2 +- .../tools/impl/query_subscribe_history.py | 2 +- app/agent/tools/impl/query_subscribes.py | 2 +- app/agent/tools/impl/query_system_settings.py | 2 +- .../tools/impl/query_transfer_history.py | 2 +- app/agent/tools/impl/query_workflows.py | 2 +- app/agent/tools/impl/run_agent_task.py | 2 +- app/agent/tools/impl/run_workflow.py | 2 +- app/agent/tools/impl/scrape_metadata.py | 2 +- app/agent/tools/impl/search_media.py | 2 +- app/agent/tools/impl/search_person_credits.py | 2 +- app/agent/tools/impl/search_subscribe.py | 2 +- app/agent/tools/impl/search_torrents.py | 2 +- app/agent/tools/impl/test_site.py | 2 +- app/agent/tools/impl/update_agent_task.py | 2 +- .../tools/impl/update_custom_identifiers.py | 2 +- app/agent/tools/impl/update_site.py | 2 +- app/agent/tools/impl/update_site_cookie.py | 2 +- app/agent/tools/impl/update_subscribe.py | 2 +- .../tools/impl/update_system_settings.py | 2 +- app/{db/user_oper.py => api/deps.py} | 87 +-- app/api/endpoints/agent.py | 5 +- app/api/endpoints/dashboard.py | 2 +- app/api/endpoints/download.py | 6 +- app/api/endpoints/history.py | 2 +- app/api/endpoints/llm.py | 5 +- app/api/endpoints/login.py | 2 +- app/api/endpoints/media.py | 11 +- app/api/endpoints/mediaserver.py | 6 +- app/api/endpoints/message.py | 6 +- app/api/endpoints/mfa.py | 4 +- app/api/endpoints/music.py | 2 +- app/api/endpoints/notification.py | 2 +- app/api/endpoints/plugin.py | 7 +- app/api/endpoints/search.py | 3 +- app/api/endpoints/site.py | 6 +- app/api/endpoints/storage.py | 2 +- app/api/endpoints/subscribe.py | 6 +- app/api/endpoints/system.py | 8 +- app/api/endpoints/tmdb.py | 4 +- app/api/endpoints/torrent.py | 12 +- app/api/endpoints/transfer.py | 5 +- app/api/endpoints/user.py | 8 +- app/api/endpoints/workflow.py | 9 +- app/application/directory.py | 2 +- app/application/filter.py | 2 +- app/application/history.py | 165 +++- app/application/mediaserver.py | 2 +- app/application/messaging/message.py | 2 +- app/application/recognition.py | 2 +- app/application/security/auth.py | 4 +- app/application/storage.py | 2 +- app/application/subscribe.py | 122 +++ app/application/torrent.py | 6 +- app/application/transfer.py | 91 +++ app/chain/__init__.py | 12 +- app/chain/download.py | 8 +- app/chain/media.py | 16 +- app/chain/mediaserver.py | 2 +- app/chain/message.py | 6 +- app/chain/recommend.py | 2 +- app/chain/scraping.py | 9 +- app/chain/search.py | 8 +- app/chain/site.py | 4 +- app/chain/subscribe.py | 21 +- app/chain/torrents.py | 6 +- app/chain/transfer.py | 89 ++- app/chain/user.py | 2 +- app/chain/workflow.py | 2 +- app/db/__init__.py | 664 +++------------- app/db/base.py | 150 ++++ app/db/decorators.py | 264 +++++++ app/db/diagnostics.py | 58 ++ app/db/engine.py | 334 ++++++++ app/db/models/__init__.py | 7 + .../{media_identity.py => _constraints.py} | 8 + app/db/models/_identity.py | 71 ++ app/db/models/agentchat.py | 65 +- app/db/models/agenttask.py | 66 +- app/db/models/agenttaskrun.py | 187 +++-- app/db/models/downloadfailure.py | 76 +- app/db/models/downloadhistory.py | 353 +++------ app/db/models/mediaserver.py | 95 +-- app/db/models/message.py | 70 +- app/db/models/passkey.py | 44 +- app/db/models/plugindata.py | 27 +- app/db/models/site.py | 63 +- app/db/models/siteicon.py | 15 +- app/db/models/sitestatistic.py | 23 +- app/db/models/siteuserdata.py | 99 ++- app/db/models/subscribe.py | 195 +++-- app/db/models/subscribehistory.py | 126 +-- app/db/models/systemconfig.py | 11 +- app/db/models/transferhistory.py | 330 ++++---- app/db/models/transferpending.py | 37 +- app/db/models/user.py | 29 +- app/db/models/userconfig.py | 18 +- app/db/models/workflow.py | 126 ++- app/db/oper/__init__.py | 100 +++ .../{agentchat_oper.py => oper/agentchat.py} | 20 +- .../{agenttask_oper.py => oper/agenttask.py} | 2 +- .../downloadfailure.py} | 2 - .../downloadhistory.py} | 25 +- .../mediaserver.py} | 6 +- app/db/{message_oper.py => oper/message.py} | 22 +- .../plugindata.py} | 0 app/db/{site_oper.py => oper/site.py} | 20 +- app/db/oper/subscribe.py | 280 +++++++ .../subscribehistory.py} | 4 +- .../systemconfig.py} | 2 +- .../transferhistory.py} | 144 +--- .../transferpending.py} | 0 app/db/oper/user.py | 93 +++ .../userconfig.py} | 2 +- app/db/{workflow_oper.py => oper/workflow.py} | 8 +- app/db/session.py | 292 +++++++ app/db/subscribe_oper.py | 317 -------- app/domain/context.py | 2 +- app/domain/media.py | 163 +--- app/domain/meta/metabase.py | 2 +- app/domain/meta/metamusic.py | 2 +- app/domain/metainfo.py | 2 +- app/main.py | 4 +- app/modules/feishu/feishu.py | 2 +- app/modules/indexer/__init__.py | 4 +- app/modules/indexer/spider/haidan.py | 2 +- app/modules/indexer/spider/hddolby.py | 2 +- app/modules/indexer/spider/mtorrent.py | 2 +- app/modules/indexer/spider/rousi.py | 2 +- app/modules/subtitle/__init__.py | 2 +- app/modules/themoviedb/__init__.py | 7 +- app/modules/ugreen/ugreen.py | 2 +- app/monitor/dispatcher.py | 2 +- app/plugins/__init__.py | 4 +- app/runtime/config.py | 34 + app/runtime/extensions/plugin_manager.py | 4 +- app/runtime/extensions/service_registry.py | 2 +- app/scheduler.py | 4 +- app/schemas/media.py | 163 ++++ app/schemas/transfer.py | 55 +- .../database_initializer.py} | 20 +- app/startup/lifecycle.py | 25 + app/startup/modules_initializer.py | 2 +- app/testing/bootstrap.py | 31 +- app/workflow/__init__.py | 2 +- app/workflow/actions/__init__.py | 2 +- app/workflow/actions/add_subscribe.py | 2 +- app/workflow/actions/transfer_file.py | 2 +- database/versions/262735d025da_2_0_1.py | 2 +- database/versions/294b007932ef_2_0_0.py | 2 +- database/versions/3891a5e722a1_2_1_7.py | 2 +- database/versions/486e56a62dcb_2_1_5.py | 2 +- database/versions/4dadad1d161a_3_0_0.py | 2 +- database/versions/89d24811e894_2_1_4.py | 2 +- database/versions/a295e41830a6_2_0_6.py | 2 +- database/versions/a73f2dbf5c09_2_0_4.py | 2 +- database/versions/e8b1c4d7a2f9_2_2_18.py | 2 +- docs/rules/04-design-patterns.md | 20 +- docs/rules/05-architecture.md | 25 +- docs/rules/06-code-styles.md | 4 +- docs/rules/09-external-response.md | 2 +- docs/rules/10-data-and-persistent.md | 47 +- docs/rules/11-quality-and-security.md | 2 +- docs/testing.md | 2 +- scripts/local_setup.py | 12 +- tests/conftest.py | 127 ++- tests/test_agent_chat_history.py | 2 +- tests/test_agent_message_routing.py | 2 +- tests/test_agent_scheduled_tasks.py | 2 +- tests/test_agent_task_runs.py | 2 +- tests/test_api_authorization.py | 2 +- tests/test_async_db_pooling.py | 188 +++++ tests/test_bluray.py | 2 +- tests/test_database_index_migration.py | 6 +- tests/test_database_migration_startup.py | 2 +- tests/test_db_base_crud.py | 148 ++++ tests/test_db_config_user_queries.py | 263 +++++++ tests/test_db_declarative_2_0.py | 242 ++++++ tests/test_db_decorator_error_paths.py | 736 ++++++++++++++++++ tests/test_db_downloadhistory_queries.py | 396 ++++++++++ tests/test_db_engine_postgresql.py | 349 +++++++++ tests/test_db_error_diagnostics.py | 12 +- tests/test_db_lazy_engine.py | 298 +++++++ tests/test_db_media_identity_normalizer.py | 164 ++++ tests/test_db_mediaserver_queries.py | 173 ++++ tests/test_db_oper_layer.py | 644 +++++++++++++++ tests/test_db_oper_layer_extra.py | 308 ++++++++ tests/test_db_plugin_message_agent_queries.py | 456 +++++++++++ tests/test_db_public_api.py | 65 ++ tests/test_db_session_lifecycle.py | 217 ++++++ tests/test_db_site_queries.py | 280 +++++++ tests/test_db_subscribe_queries.py | 337 ++++++++ tests/test_db_transferhistory_queries.py | 519 ++++++++++++ tests/test_db_transferpending_queries.py | 178 +++++ tests/test_db_workflow_queries.py | 235 ++++++ tests/test_lifecycle_shutdown.py | 104 +++ tests/test_manual_transfer_history.py | 2 +- tests/test_media_source_routing.py | 7 +- tests/test_mediascrape.py | 2 +- tests/test_mediaserver_sync_incremental.py | 8 +- tests/test_message_notifications.py | 4 +- tests/test_music_recognize_cache.py | 2 +- tests/test_music_subscribe.py | 25 +- tests/test_music_transfer.py | 9 +- tests/test_notification_template_render.py | 2 +- tests/test_plugin_helper.py | 2 +- tests/test_rust_accel.py | 2 +- tests/test_search_media_sources.py | 2 +- tests/test_subscribe_chain.py | 19 +- tests/test_subscribe_endpoint.py | 8 +- tests/test_subscribe_oper.py | 161 +++- tests/test_subscribe_write_path.py | 422 ++++++++++ tests/test_sunnypt_indexer.py | 2 +- tests/test_system_llm_test.py | 4 +- tests/test_system_nettest.py | 4 +- tests/test_systemconfig_oper.py | 2 +- tests/test_tmdb_cache_management.py | 2 +- tests/test_transfer_history_write_path.py | 291 +++++++ tests/test_transfer_job_manager.py | 164 +++- tests/test_transfer_mark_torrent_completed.py | 18 +- tests/test_transfer_mounted_disk_cleanup.py | 3 +- tests/test_transfer_movie_collection.py | 48 +- tests/test_transfer_overwrite_declined.py | 35 +- tests/test_transfer_pending_replay.py | 3 +- tests/test_transfer_stale_tasks.py | 18 +- tests/test_transfer_sync_extra_files.py | 10 +- tests/test_transfer_tmdb_category.py | 3 +- ..._transferhistory_media_source_migration.py | 14 +- tests/test_web_agent_stream.py | 2 +- tests/test_workflow_authorization.py | 2 +- 252 files changed, 11405 insertions(+), 2889 deletions(-) rename app/{db/user_oper.py => api/deps.py} (58%) create mode 100644 app/application/subscribe.py create mode 100644 app/application/transfer.py create mode 100644 app/db/base.py create mode 100644 app/db/decorators.py create mode 100644 app/db/diagnostics.py create mode 100644 app/db/engine.py rename app/db/models/{media_identity.py => _constraints.py} (57%) create mode 100644 app/db/models/_identity.py create mode 100644 app/db/oper/__init__.py rename app/db/{agentchat_oper.py => oper/agentchat.py} (96%) rename app/db/{agenttask_oper.py => oper/agenttask.py} (99%) rename app/db/{downloadfailure_oper.py => oper/downloadfailure.py} (92%) rename app/db/{downloadhistory_oper.py => oper/downloadhistory.py} (86%) rename app/db/{mediaserver_oper.py => oper/mediaserver.py} (96%) rename app/db/{message_oper.py => oper/message.py} (87%) rename app/db/{plugindata_oper.py => oper/plugindata.py} (100%) rename app/db/{site_oper.py => oper/site.py} (94%) create mode 100644 app/db/oper/subscribe.py rename app/db/{subscribehistory_oper.py => oper/subscribehistory.py} (88%) rename app/db/{systemconfig_oper.py => oper/systemconfig.py} (98%) rename app/db/{transferhistory_oper.py => oper/transferhistory.py} (56%) rename app/db/{transferpending_oper.py => oper/transferpending.py} (100%) create mode 100644 app/db/oper/user.py rename app/db/{userconfig_oper.py => oper/userconfig.py} (96%) rename app/db/{workflow_oper.py => oper/workflow.py} (92%) create mode 100644 app/db/session.py delete mode 100644 app/db/subscribe_oper.py rename app/{db/init.py => startup/database_initializer.py} (55%) create mode 100644 tests/test_async_db_pooling.py create mode 100644 tests/test_db_base_crud.py create mode 100644 tests/test_db_config_user_queries.py create mode 100644 tests/test_db_declarative_2_0.py create mode 100644 tests/test_db_decorator_error_paths.py create mode 100644 tests/test_db_downloadhistory_queries.py create mode 100644 tests/test_db_engine_postgresql.py create mode 100644 tests/test_db_lazy_engine.py create mode 100644 tests/test_db_media_identity_normalizer.py create mode 100644 tests/test_db_mediaserver_queries.py create mode 100644 tests/test_db_oper_layer.py create mode 100644 tests/test_db_oper_layer_extra.py create mode 100644 tests/test_db_plugin_message_agent_queries.py create mode 100644 tests/test_db_public_api.py create mode 100644 tests/test_db_session_lifecycle.py create mode 100644 tests/test_db_site_queries.py create mode 100644 tests/test_db_subscribe_queries.py create mode 100644 tests/test_db_transferhistory_queries.py create mode 100644 tests/test_db_transferpending_queries.py create mode 100644 tests/test_db_workflow_queries.py create mode 100644 tests/test_subscribe_write_path.py create mode 100644 tests/test_transfer_history_write_path.py diff --git a/.gitignore b/.gitignore index dc1308566..5b11aa738 100644 --- a/.gitignore +++ b/.gitignore @@ -11,6 +11,9 @@ nginx/ test.py safety_report.txt app/application/site/*.bin +# 站点数据的运行期下载产物。上游 v3 架构重构后落点从 app/application/site 移到了 +# app/helper,同目录的 .so/.pyd 由上面的通配兜住,只有 .bin 漏了网 +app/helper/*.bin app/plugins/** !app/plugins/__init__.py config/cookies/ diff --git a/app/adapters/external/market.py b/app/adapters/external/market.py index 5da6e7a5a..7c87eaf14 100644 --- a/app/adapters/external/market.py +++ b/app/adapters/external/market.py @@ -30,7 +30,7 @@ from requests import Response from app.runtime.cache import cached, is_fresh from app.runtime.config import settings -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.adapters.system.package import PackageInstallRequest, build_package_install_strategies from app.runtime.log import logger from app.schemas.types import SystemConfigKey diff --git a/app/adapters/external/server.py b/app/adapters/external/server.py index 07e46f520..ecc67f365 100644 --- a/app/adapters/external/server.py +++ b/app/adapters/external/server.py @@ -9,9 +9,9 @@ from app.runtime.cache import cached from app.runtime.config import settings from app.domain.context import MediaInfo, MusicInfo from app.domain.meta.metabase import MetaBase -from app.db.subscribe_oper import SubscribeOper -from app.db.systemconfig_oper import SystemConfigOper -from app.db.workflow_oper import WorkflowOper +from app.db.oper.subscribe import SubscribeOper +from app.db.oper.systemconfig import SystemConfigOper +from app.db.oper.workflow import WorkflowOper from app.runtime.log import logger from app.schemas.types import ( MUSIC_ENTITY_RECORDING, @@ -20,7 +20,8 @@ from app.schemas.types import ( media_type_to_agent, ) from app.adapters.network.http import AsyncRequestUtils, RequestUtils -from app.domain.media import normalize_music_type, resolve_media_identity +from app.domain.media import normalize_music_type +from app.schemas.media import resolve_media_identity from app.adapters.system.host import SystemUtils from version import APP_VERSION, FRONTEND_VERSION diff --git a/app/agent/llm/provider.py b/app/agent/llm/provider.py index 278fc589b..d837e5082 100644 --- a/app/agent/llm/provider.py +++ b/app/agent/llm/provider.py @@ -21,7 +21,7 @@ import httpx import jwt from app.runtime.config import settings -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas.types import SystemConfigKey from app.foundation.singleton import Singleton diff --git a/app/agent/mcp.py b/app/agent/mcp.py index 73ab105e0..8b274868c 100644 --- a/app/agent/mcp.py +++ b/app/agent/mcp.py @@ -12,7 +12,7 @@ from dataclasses import dataclass from typing import Any, Optional from urllib.parse import urljoin -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas.agent import ( AgentMcpServerConfig, diff --git a/app/agent/memory/__init__.py b/app/agent/memory/__init__.py index 310e67914..313b3da41 100644 --- a/app/agent/memory/__init__.py +++ b/app/agent/memory/__init__.py @@ -7,7 +7,7 @@ from typing import Dict, List, Optional from langchain_core.messages import BaseMessage, messages_from_dict, messages_to_dict from app.runtime.config import settings -from app.db.agentchat_oper import AgentChatOper +from app.db.oper.agentchat import AgentChatOper from app.runtime.log import logger from app.schemas.agent import ConversationMemory diff --git a/app/agent/orchestrator.py b/app/agent/orchestrator.py index 12d02eded..f4d3b6581 100644 --- a/app/agent/orchestrator.py +++ b/app/agent/orchestrator.py @@ -72,9 +72,9 @@ from app.chain import ChainBase from app.runtime.config import settings from app.runtime.events import eventmanager from app.runtime.extensions.plugin_manager import PluginManager -from app.db.agentchat_oper import AgentChatOper -from app.db.agenttask_oper import AgentTaskOper -from app.db.user_oper import UserOper +from app.db.oper.agentchat import AgentChatOper +from app.db.oper.agenttask import AgentTaskOper +from app.db.oper.user import UserOper from app.runtime.log import logger from app.schemas import AgentLLMProviderEventData, AgentTokensUsageEventData, Notification, NotificationType from app.schemas.message import ChannelCapabilityManager, ChannelCapability diff --git a/app/agent/tools/impl/_filter_rule_utils.py b/app/agent/tools/impl/_filter_rule_utils.py index e3f047d5d..1bb484999 100644 --- a/app/agent/tools/impl/_filter_rule_utils.py +++ b/app/agent/tools/impl/_filter_rule_utils.py @@ -5,8 +5,8 @@ import re from typing import Any, Dict, Iterable, Optional from app.runtime.events import eventmanager -from app.db.subscribe_oper import SubscribeOper -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.subscribe import SubscribeOper +from app.db.oper.systemconfig import SystemConfigOper from app.application.filter import RuleHelper from app.modules.filter.RuleParser import RuleParser from app.modules.filter.builtin_rules import BUILTIN_RULE_SET diff --git a/app/agent/tools/impl/_plugin_tool_utils.py b/app/agent/tools/impl/_plugin_tool_utils.py index 5dd9e3c6a..5070f8354 100644 --- a/app/agent/tools/impl/_plugin_tool_utils.py +++ b/app/agent/tools/impl/_plugin_tool_utils.py @@ -6,7 +6,7 @@ from typing import Any, Optional from app.runtime.config import settings from app.runtime.extensions.plugin_manager import PluginManager -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.adapters.external.server import MoviePilotServerHelper from app.adapters.external.market import PluginHelper from app.schemas.types import SystemConfigKey diff --git a/app/agent/tools/impl/add_download_tasks.py b/app/agent/tools/impl/add_download_tasks.py index e8cc43649..f68af1c2c 100644 --- a/app/agent/tools/impl/add_download_tasks.py +++ b/app/agent/tools/impl/add_download_tasks.py @@ -15,7 +15,7 @@ from app.chain.search import SearchChain from app.runtime.config import settings from app.domain.context import Context from app.domain.metainfo import MetaInfo -from app.db.site_oper import SiteOper +from app.db.oper.site import SiteOper from app.application.directory import DirectoryHelper, validate_download_save_path from app.runtime.log import logger from app.schemas import FileURI diff --git a/app/agent/tools/impl/add_subscribe.py b/app/agent/tools/impl/add_subscribe.py index a9355f0b7..40a73f74b 100644 --- a/app/agent/tools/impl/add_subscribe.py +++ b/app/agent/tools/impl/add_subscribe.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.chain.subscribe import SubscribeChain -from app.db.user_oper import UserOper +from app.db.oper.user import UserOper from app.runtime.log import logger from app.schemas.types import MUSIC_ENTITY_ALBUM, MediaSource, MediaType, MessageChannel from ._music_utils import normalize_music_type diff --git a/app/agent/tools/impl/create_agent_task.py b/app/agent/tools/impl/create_agent_task.py index de467dd9d..d4383cc40 100644 --- a/app/agent/tools/impl/create_agent_task.py +++ b/app/agent/tools/impl/create_agent_task.py @@ -8,8 +8,8 @@ from pydantic import BaseModel, Field, model_validator from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.runtime.config import settings -from app.db.agentchat_oper import AgentChatOper -from app.db.agenttask_oper import AgentTaskOper +from app.db.oper.agentchat import AgentChatOper +from app.db.oper.agenttask import AgentTaskOper from app.runtime.scheduling import TimerUtils diff --git a/app/agent/tools/impl/delete_agent_task.py b/app/agent/tools/impl/delete_agent_task.py index 411b0da11..642f0d758 100644 --- a/app/agent/tools/impl/delete_agent_task.py +++ b/app/agent/tools/impl/delete_agent_task.py @@ -4,7 +4,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag -from app.db.agenttask_oper import AgentTaskOper +from app.db.oper.agenttask import AgentTaskOper class DeleteAgentTaskInput(BaseModel): diff --git a/app/agent/tools/impl/delete_download_history.py b/app/agent/tools/impl/delete_download_history.py index 64fd50f5b..0b72a0a4e 100644 --- a/app/agent/tools/impl/delete_download_history.py +++ b/app/agent/tools/impl/delete_download_history.py @@ -6,7 +6,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag -from app.db.downloadhistory_oper import DownloadHistoryOper +from app.db.oper.downloadhistory import DownloadHistoryOper from app.runtime.log import logger diff --git a/app/agent/tools/impl/delete_subscribe.py b/app/agent/tools/impl/delete_subscribe.py index 415bc1ffa..e2a7bb055 100644 --- a/app/agent/tools/impl/delete_subscribe.py +++ b/app/agent/tools/impl/delete_subscribe.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.runtime.events import eventmanager -from app.db.subscribe_oper import SubscribeOper +from app.db.oper.subscribe import SubscribeOper from app.adapters.external.server import MoviePilotServerHelper from app.runtime.log import logger from app.schemas.types import EventType diff --git a/app/agent/tools/impl/delete_transfer_history.py b/app/agent/tools/impl/delete_transfer_history.py index e06fcba9f..363449221 100644 --- a/app/agent/tools/impl/delete_transfer_history.py +++ b/app/agent/tools/impl/delete_transfer_history.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.chain.storage import StorageChain -from app.db.transferhistory_oper import TransferHistoryOper +from app.db.oper.transferhistory import TransferHistoryOper from app.runtime.log import logger from app.schemas import FileItem diff --git a/app/agent/tools/impl/query_agent_tasks.py b/app/agent/tools/impl/query_agent_tasks.py index 0ab7f8d6f..9559e024a 100644 --- a/app/agent/tools/impl/query_agent_tasks.py +++ b/app/agent/tools/impl/query_agent_tasks.py @@ -6,7 +6,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.runtime.config import settings -from app.db.agenttask_oper import AgentTaskOper +from app.db.oper.agenttask import AgentTaskOper class QueryAgentTasksInput(BaseModel): diff --git a/app/agent/tools/impl/query_custom_identifiers.py b/app/agent/tools/impl/query_custom_identifiers.py index 15f1b22df..68af8135c 100644 --- a/app/agent/tools/impl/query_custom_identifiers.py +++ b/app/agent/tools/impl/query_custom_identifiers.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas.types import SystemConfigKey diff --git a/app/agent/tools/impl/query_download_tasks.py b/app/agent/tools/impl/query_download_tasks.py index f1336c0ce..ce62c904b 100644 --- a/app/agent/tools/impl/query_download_tasks.py +++ b/app/agent/tools/impl/query_download_tasks.py @@ -8,7 +8,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.chain.download import DownloadChain -from app.db.downloadhistory_oper import DownloadHistoryOper +from app.db.oper.downloadhistory import DownloadHistoryOper from app.runtime.log import logger from app.schemas import DownloaderTorrent from app.schemas.types import MUSIC_ENTITY_RECORDING, TorrentQueryStatus, media_type_to_agent diff --git a/app/agent/tools/impl/query_downloaders.py b/app/agent/tools/impl/query_downloaders.py index 67330971a..1ae9f11b1 100644 --- a/app/agent/tools/impl/query_downloaders.py +++ b/app/agent/tools/impl/query_downloaders.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas.types import SystemConfigKey diff --git a/app/agent/tools/impl/query_plugin_data.py b/app/agent/tools/impl/query_plugin_data.py index c90400ad9..fff735e37 100644 --- a/app/agent/tools/impl/query_plugin_data.py +++ b/app/agent/tools/impl/query_plugin_data.py @@ -12,7 +12,7 @@ from app.agent.tools.impl._plugin_tool_utils import ( build_preview_payload, get_plugin_snapshot, ) -from app.db.plugindata_oper import PluginDataOper +from app.db.oper.plugindata import PluginDataOper from app.runtime.log import logger diff --git a/app/agent/tools/impl/query_site_userdata.py b/app/agent/tools/impl/query_site_userdata.py index 203151125..05926765a 100644 --- a/app/agent/tools/impl/query_site_userdata.py +++ b/app/agent/tools/impl/query_site_userdata.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag -from app.db.site_oper import SiteOper +from app.db.oper.site import SiteOper from app.runtime.log import logger SITE_USERDATA_DETAIL_PREVIEW_LIMIT = 10 diff --git a/app/agent/tools/impl/query_sites.py b/app/agent/tools/impl/query_sites.py index d608d1844..d7cf43a5b 100644 --- a/app/agent/tools/impl/query_sites.py +++ b/app/agent/tools/impl/query_sites.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag -from app.db.site_oper import SiteOper +from app.db.oper.site import SiteOper from app.runtime.log import logger diff --git a/app/agent/tools/impl/query_subscribe_history.py b/app/agent/tools/impl/query_subscribe_history.py index 7748690be..1926ec2a2 100644 --- a/app/agent/tools/impl/query_subscribe_history.py +++ b/app/agent/tools/impl/query_subscribe_history.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag -from app.db.subscribehistory_oper import SubscribeHistoryOper +from app.db.oper.subscribehistory import SubscribeHistoryOper from app.runtime.log import logger from app.schemas.types import MUSIC_ENTITY_RECORDING, MediaType, media_type_to_agent from ._music_utils import normalize_music_type diff --git a/app/agent/tools/impl/query_subscribes.py b/app/agent/tools/impl/query_subscribes.py index c216b7b1c..f20f61acc 100644 --- a/app/agent/tools/impl/query_subscribes.py +++ b/app/agent/tools/impl/query_subscribes.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag -from app.db.subscribe_oper import SubscribeOper +from app.db.oper.subscribe import SubscribeOper from app.runtime.log import logger from app.schemas.subscribe import Subscribe as SubscribeSchema from app.schemas.types import ( diff --git a/app/agent/tools/impl/query_system_settings.py b/app/agent/tools/impl/query_system_settings.py index 7facf698a..e86c29bd8 100644 --- a/app/agent/tools/impl/query_system_settings.py +++ b/app/agent/tools/impl/query_system_settings.py @@ -16,7 +16,7 @@ from app.agent.tools.impl._system_setting_utils import ( should_redact_setting, ) from app.runtime.config import settings -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger diff --git a/app/agent/tools/impl/query_transfer_history.py b/app/agent/tools/impl/query_transfer_history.py index 119bbd396..a84343a53 100644 --- a/app/agent/tools/impl/query_transfer_history.py +++ b/app/agent/tools/impl/query_transfer_history.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag -from app.db.transferhistory_oper import TransferHistoryOper +from app.db.oper.transferhistory import TransferHistoryOper from app.runtime.log import logger from app.schemas.types import media_type_to_agent from app.foundation.text import cut as jieba_cut diff --git a/app/agent/tools/impl/query_workflows.py b/app/agent/tools/impl/query_workflows.py index 9a5f6b377..8dc6fa118 100644 --- a/app/agent/tools/impl/query_workflows.py +++ b/app/agent/tools/impl/query_workflows.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag -from app.db.workflow_oper import WorkflowOper +from app.db.oper.workflow import WorkflowOper from app.runtime.log import logger diff --git a/app/agent/tools/impl/run_agent_task.py b/app/agent/tools/impl/run_agent_task.py index 9a11f4f86..00009064a 100644 --- a/app/agent/tools/impl/run_agent_task.py +++ b/app/agent/tools/impl/run_agent_task.py @@ -6,7 +6,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag -from app.db.agenttask_oper import AgentTaskOper +from app.db.oper.agenttask import AgentTaskOper class RunAgentTaskInput(BaseModel): diff --git a/app/agent/tools/impl/run_workflow.py b/app/agent/tools/impl/run_workflow.py index 26ef9c484..694d2a452 100644 --- a/app/agent/tools/impl/run_workflow.py +++ b/app/agent/tools/impl/run_workflow.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.chain.workflow import WorkflowChain -from app.db.workflow_oper import WorkflowOper +from app.db.oper.workflow import WorkflowOper from app.runtime.log import logger diff --git a/app/agent/tools/impl/scrape_metadata.py b/app/agent/tools/impl/scrape_metadata.py index 66b83240a..a72f0039b 100644 --- a/app/agent/tools/impl/scrape_metadata.py +++ b/app/agent/tools/impl/scrape_metadata.py @@ -19,7 +19,7 @@ from app.schemas.types import ( MediaType, media_type_to_agent, ) -from app.domain.media import normalize_media_source +from app.schemas.media import normalize_media_source from ._music_utils import normalize_music_type, simplify_music_info diff --git a/app/agent/tools/impl/search_media.py b/app/agent/tools/impl/search_media.py index eded9abdb..d59b8dbea 100644 --- a/app/agent/tools/impl/search_media.py +++ b/app/agent/tools/impl/search_media.py @@ -10,7 +10,7 @@ from app.agent.tools.tags import ToolTag from app.chain.media import MediaChain from app.runtime.log import logger from app.schemas.types import MediaType, media_type_to_agent -from app.domain.media import resolve_media_identity +from app.schemas.media import resolve_media_identity from ._music_utils import normalize_music_type, simplify_music_info diff --git a/app/agent/tools/impl/search_person_credits.py b/app/agent/tools/impl/search_person_credits.py index cdfd04822..c9608f58e 100644 --- a/app/agent/tools/impl/search_person_credits.py +++ b/app/agent/tools/impl/search_person_credits.py @@ -11,7 +11,7 @@ from app.chain.douban import DoubanChain from app.chain.tmdb import TmdbChain from app.chain.bangumi import BangumiChain from app.runtime.log import logger -from app.domain.media import resolve_media_identity +from app.schemas.media import resolve_media_identity class SearchPersonCreditsInput(BaseModel): diff --git a/app/agent/tools/impl/search_subscribe.py b/app/agent/tools/impl/search_subscribe.py index e29922442..fdc08d497 100644 --- a/app/agent/tools/impl/search_subscribe.py +++ b/app/agent/tools/impl/search_subscribe.py @@ -8,7 +8,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.chain.subscribe import SubscribeChain -from app.db.subscribe_oper import SubscribeOper +from app.db.oper.subscribe import SubscribeOper from app.runtime.log import logger from app.schemas.types import media_type_to_agent diff --git a/app/agent/tools/impl/search_torrents.py b/app/agent/tools/impl/search_torrents.py index 9320c3825..5a8a09165 100644 --- a/app/agent/tools/impl/search_torrents.py +++ b/app/agent/tools/impl/search_torrents.py @@ -8,7 +8,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.chain.search import SearchChain -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.application.site.sites import SitesHelper # pylint: disable=no-name-in-module from app.runtime.log import logger from app.schemas.types import MediaSource, MediaType, SystemConfigKey diff --git a/app/agent/tools/impl/test_site.py b/app/agent/tools/impl/test_site.py index 3d5b640df..12707bf07 100644 --- a/app/agent/tools/impl/test_site.py +++ b/app/agent/tools/impl/test_site.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.chain.site import SiteChain -from app.db.site_oper import SiteOper +from app.db.oper.site import SiteOper from app.runtime.log import logger diff --git a/app/agent/tools/impl/update_agent_task.py b/app/agent/tools/impl/update_agent_task.py index b2c2cd285..9f7afa213 100644 --- a/app/agent/tools/impl/update_agent_task.py +++ b/app/agent/tools/impl/update_agent_task.py @@ -8,7 +8,7 @@ from pydantic import BaseModel, Field, model_validator from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.runtime.config import settings -from app.db.agenttask_oper import AgentTaskOper +from app.db.oper.agenttask import AgentTaskOper from app.runtime.scheduling import TimerUtils diff --git a/app/agent/tools/impl/update_custom_identifiers.py b/app/agent/tools/impl/update_custom_identifiers.py index b4b27cf81..52dc898fe 100644 --- a/app/agent/tools/impl/update_custom_identifiers.py +++ b/app/agent/tools/impl/update_custom_identifiers.py @@ -8,7 +8,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.domain.metainfo import clear_rust_parse_options_cache -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas.types import SystemConfigKey diff --git a/app/agent/tools/impl/update_site.py b/app/agent/tools/impl/update_site.py index 6eced36be..df2552151 100644 --- a/app/agent/tools/impl/update_site.py +++ b/app/agent/tools/impl/update_site.py @@ -8,7 +8,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.runtime.events import eventmanager -from app.db.site_oper import SiteOper +from app.db.oper.site import SiteOper from app.runtime.log import logger from app.schemas.types import EventType from app.domain.string import StringUtils diff --git a/app/agent/tools/impl/update_site_cookie.py b/app/agent/tools/impl/update_site_cookie.py index 60542270e..3472234a2 100644 --- a/app/agent/tools/impl/update_site_cookie.py +++ b/app/agent/tools/impl/update_site_cookie.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.chain.site import SiteChain -from app.db.site_oper import SiteOper +from app.db.oper.site import SiteOper from app.runtime.log import logger diff --git a/app/agent/tools/impl/update_subscribe.py b/app/agent/tools/impl/update_subscribe.py index 93fd5d1b9..2cf7459fc 100644 --- a/app/agent/tools/impl/update_subscribe.py +++ b/app/agent/tools/impl/update_subscribe.py @@ -8,7 +8,7 @@ from pydantic import BaseModel, Field from app.agent.tools.base import MoviePilotTool from app.agent.tools.tags import ToolTag from app.runtime.events import eventmanager -from app.db.subscribe_oper import SubscribeOper +from app.db.oper.subscribe import SubscribeOper from app.runtime.log import logger from app.schemas.event import SubscribeModifiedEventData from app.schemas.types import EventType, media_type_to_agent diff --git a/app/agent/tools/impl/update_system_settings.py b/app/agent/tools/impl/update_system_settings.py index 7a1197d30..56718f52d 100644 --- a/app/agent/tools/impl/update_system_settings.py +++ b/app/agent/tools/impl/update_system_settings.py @@ -18,7 +18,7 @@ from app.agent.tools.impl._system_setting_utils import ( ) from app.runtime.config import settings from app.runtime.events import eventmanager -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas.event import ConfigChangeEventData from app.schemas.types import EventType diff --git a/app/db/user_oper.py b/app/api/deps.py similarity index 58% rename from app/db/user_oper.py rename to app/api/deps.py index eb504da9b..a7c3745f3 100644 --- a/app/db/user_oper.py +++ b/app/api/deps.py @@ -1,12 +1,18 @@ -from typing import Optional, List +""" +API 层的公共依赖。 +这些是 FastAPI 的路由依赖:从令牌解出用户、校验激活状态与权限,失败一律以 +HTTPException 表达。它们此前住在 app/db/oper/user.py 里,与数据访问混在一处—— +鉴权是 HTTP 层的关注点,产出的是 403/400 而不是数据。放在 db 包里既让数据层反向 +依赖了 fastapi,也使这部分逻辑无法与数据访问分开度量。 +""" from fastapi import Depends, HTTPException from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Session from app import schemas from app.application.security.access import verify_token -from app.db import DbOper, get_db, get_async_db +from app.db import get_async_db, get_db from app.db.models.user import User @@ -112,80 +118,3 @@ async def get_current_active_superuser_async( status_code=400, detail="用户权限不足" ) return current_user - - -class UserOper(DbOper): - """ - 用户管理 - """ - - def list(self) -> List[User]: - """ - 获取用户列表 - """ - return User.list(self._db) - - def add(self, **kwargs): - """ - 新增用户 - """ - user = User(**kwargs) - user.create(self._db) - - def get_by_name(self, name: str) -> User: - """ - 根据用户名获取用户 - """ - return User.get_by_name(self._db, name) - - async def async_get_by_name(self, name: str) -> User: - """ - 异步根据用户名获取用户。 - """ - return await User.async_get_by_name(self._db, name) - - async def async_get_by_id(self, user_id: int) -> User: - """ - 异步根据用户 ID 获取用户。 - """ - return await User.async_get_by_id(self._db, user_id) - - def get_permissions(self, name: str) -> dict: - """ - 获取用户权限 - """ - user = User.get_by_name(self._db, name) - if user: - return user.permissions or {} - return {} - - def get_settings(self, name: str) -> Optional[dict]: - """ - 获取用户个性化设置,返回None表示用户不存在 - """ - user = User.get_by_name(self._db, name) - if user: - return user.settings or {} - return None - - def get_setting(self, name: str, key: str) -> Optional[str]: - """ - 获取用户个性化设置 - """ - settings = self.get_settings(name) - if settings: - return settings.get(key) - return None - - def get_name(self, **kwargs) -> Optional[str]: - """ - 根据绑定账号获取用户名称 - """ - users = self.list() - for user in users: - user_setting = user.settings - if user_setting: - for k, v in kwargs.items(): - if user_setting.get(k) == str(v): - return user.name - return None diff --git a/app/api/endpoints/agent.py b/app/api/endpoints/agent.py index 83e70ff71..6a8a29832 100644 --- a/app/api/endpoints/agent.py +++ b/app/api/endpoints/agent.py @@ -32,10 +32,11 @@ from app.command import Command from app.runtime.config import global_vars, settings from app.runtime.events import Event, EventManager from app.db import get_async_db -from app.db.agentchat_oper import AgentChatOper +from app.db.oper.agentchat import AgentChatOper from app.db.models import User from app.db.models.agentchat import AgentChat -from app.db.user_oper import UserOper, get_current_active_user +from app.db.oper.user import UserOper +from app.api.deps import get_current_active_user from app.application.messaging.agent import attach_web_agent_edit_queue, detach_web_agent_edit_queue from app.application.messaging.interaction import agent_interaction_manager, media_interaction_manager from app.runtime.localization import LocaleHelper diff --git a/app/api/endpoints/dashboard.py b/app/api/endpoints/dashboard.py index 2e1888351..c650fbb5b 100644 --- a/app/api/endpoints/dashboard.py +++ b/app/api/endpoints/dashboard.py @@ -12,7 +12,7 @@ from app.runtime.config import settings from app.application.security.access import verify_apitoken from app.db import get_db from app.db.models.transferhistory import TransferHistory -from app.db.user_oper import get_current_active_superuser +from app.api.deps import get_current_active_superuser from app.application.directory import DirectoryHelper from app.scheduler import Scheduler from app.adapters.system.host import SystemUtils diff --git a/app/api/endpoints/download.py b/app/api/endpoints/download.py index 09f12f1eb..9962ba3e2 100644 --- a/app/api/endpoints/download.py +++ b/app/api/endpoints/download.py @@ -11,9 +11,9 @@ from app.domain.meta.metamusic import MetaMusic from app.domain.metainfo import MetaInfo from app.application.security.access import verify_token from app.db.models.user import User -from app.db.site_oper import SiteOper -from app.db.systemconfig_oper import SystemConfigOper -from app.db.user_oper import get_current_active_user +from app.db.oper.site import SiteOper +from app.db.oper.systemconfig import SystemConfigOper +from app.api.deps import get_current_active_user from app.application.directory import DirectoryHelper from app.schemas.types import ( MUSIC_ENTITY_RECORDING, diff --git a/app/api/endpoints/history.py b/app/api/endpoints/history.py index 068cb15e2..b6f1f5560 100644 --- a/app/api/endpoints/history.py +++ b/app/api/endpoints/history.py @@ -22,7 +22,7 @@ from app.db import get_async_db, get_db from app.db.models import User from app.db.models.downloadhistory import DownloadHistory, DownloadFiles from app.db.models.transferhistory import TransferHistory -from app.db.user_oper import ( +from app.api.deps import ( get_current_active_manage_user, get_current_active_superuser, get_current_active_superuser_async, diff --git a/app/api/endpoints/llm.py b/app/api/endpoints/llm.py index a4bd4c748..1ed38f720 100644 --- a/app/api/endpoints/llm.py +++ b/app/api/endpoints/llm.py @@ -15,10 +15,7 @@ from app.agent.llm import ( ) from app.runtime.config import settings from app.db.models import User -from app.db.user_oper import ( - get_current_active_superuser_async, - get_current_active_user_async, -) +from app.api.deps import get_current_active_superuser_async, get_current_active_user_async from app.runtime.log import logger router = ResponseAPIRouter() diff --git a/app/api/endpoints/login.py b/app/api/endpoints/login.py index 26c1265a0..574003ad1 100644 --- a/app/api/endpoints/login.py +++ b/app/api/endpoints/login.py @@ -10,7 +10,7 @@ from app.api.response import RAW_RESPONSE_OPENAPI_KEY, ResponseAPIRouter from app.chain.user import MfaRequired, UserChain from app.application.security import access as security from app.runtime.config import settings -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.application.site.sites import SitesHelper # pylint: disable=no-name-in-module from app.application.image import WallpaperHelper from app.schemas.types import SystemConfigKey diff --git a/app/api/endpoints/media.py b/app/api/endpoints/media.py index 217fe5dd6..1ca35b697 100644 --- a/app/api/endpoints/media.py +++ b/app/api/endpoints/media.py @@ -17,17 +17,12 @@ from app.domain.meta.metamusic import MetaMusic from app.domain.metainfo import MetaInfo, MetaInfoPath from app.application.security.access import verify_token, verify_apitoken from app.db.models import User -from app.db.user_oper import get_current_active_user, get_current_active_superuser +from app.api.deps import get_current_active_user, get_current_active_superuser from app.schemas import MediaType from app.schemas.category import CategoryConfig from app.schemas.types import MUSIC_ENTITY_RECORDING, MediaSource -from app.domain.media import ( - is_music_media_source, - normalize_media_source, - normalize_music_type, - parse_media_source_selection, - resolve_media_identity, -) +from app.domain.media import is_music_media_source, normalize_music_type, parse_media_source_selection +from app.schemas.media import normalize_media_source, resolve_media_identity router = ResponseAPIRouter() diff --git a/app/api/endpoints/mediaserver.py b/app/api/endpoints/mediaserver.py index f66db0735..23f9dba3e 100644 --- a/app/api/endpoints/mediaserver.py +++ b/app/api/endpoints/mediaserver.py @@ -11,13 +11,13 @@ from app.domain.context import MediaInfo from app.domain.metainfo import MetaInfo from app.application.security.access import verify_token from app.db import get_async_db -from app.db.mediaserver_oper import MediaServerOper +from app.db.oper.mediaserver import MediaServerOper from app.db.models import MediaServerItem -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.application.mediaserver import MediaServerHelper from app.schemas import MediaType, NotExistMediaInfo from app.schemas.types import MediaSource, SystemConfigKey -from app.domain.media import build_media_key, resolve_media_identity +from app.schemas.media import build_media_key, resolve_media_identity router = ResponseAPIRouter() diff --git a/app/api/endpoints/message.py b/app/api/endpoints/message.py index b1de1b232..4f30ba0b9 100644 --- a/app/api/endpoints/message.py +++ b/app/api/endpoints/message.py @@ -14,9 +14,9 @@ from app.runtime.config import settings, global_vars from app.application.security.access import verify_token, verify_apitoken from app.db import get_async_db from app.db.models import User -from app.db.message_oper import MessageOper -from app.db.systemconfig_oper import SystemConfigOper -from app.db.user_oper import get_current_active_superuser +from app.db.oper.message import MessageOper +from app.db.oper.systemconfig import SystemConfigOper +from app.api.deps import get_current_active_superuser from app.runtime.extensions.service_registry import ServiceConfigHelper from app.runtime.log import logger from app.modules.wechat.WXBizMsgCrypt3 import WXBizMsgCrypt diff --git a/app/api/endpoints/mfa.py b/app/api/endpoints/mfa.py index c2f496f82..2a148ebd2 100644 --- a/app/api/endpoints/mfa.py +++ b/app/api/endpoints/mfa.py @@ -18,8 +18,8 @@ from app.runtime.config import settings from app.db import get_async_db from app.db.models.passkey import PassKey from app.db.models.user import User -from app.db.systemconfig_oper import SystemConfigOper -from app.db.user_oper import get_current_active_user, get_current_active_user_async +from app.db.oper.systemconfig import SystemConfigOper +from app.api.deps import get_current_active_user, get_current_active_user_async from app.application.security.passkey import ( PassKeyHelper, PassKeyRegistrationOriginMismatchError, diff --git a/app/api/endpoints/music.py b/app/api/endpoints/music.py index 0997a3b70..94580a310 100644 --- a/app/api/endpoints/music.py +++ b/app/api/endpoints/music.py @@ -10,7 +10,7 @@ from app.schemas.types import MediaSource, MediaType from app.domain.context import MusicAlbumInfo, MusicArtistInfo, MusicInfo from app.application.security.access import verify_token from app.db.models.user import User -from app.db.user_oper import get_current_active_superuser_async +from app.api.deps import get_current_active_superuser_async from app.modules.listenbrainz import ( LISTENBRAINZ_CHART_RANGES, LISTENBRAINZ_FRESH_MAX_DAYS, diff --git a/app/api/endpoints/notification.py b/app/api/endpoints/notification.py index 27effd613..fd3d2a25e 100644 --- a/app/api/endpoints/notification.py +++ b/app/api/endpoints/notification.py @@ -6,7 +6,7 @@ from app import schemas from app.api.response import ResponseAPIRouter from app.runtime.extensions.module_manager import ModuleManager from app.db.models import User -from app.db.user_oper import get_current_active_superuser +from app.api.deps import get_current_active_superuser from app.modules.wechatclawbot.wechatclawbot import WechatClawBot router = ResponseAPIRouter() diff --git a/app/api/endpoints/plugin.py b/app/api/endpoints/plugin.py index 1c340af41..4ba7c565a 100644 --- a/app/api/endpoints/plugin.py +++ b/app/api/endpoints/plugin.py @@ -24,11 +24,8 @@ from app.application.security.access import ( verify_token, ) from app.db.models import User -from app.db.systemconfig_oper import SystemConfigOper -from app.db.user_oper import ( - get_current_active_superuser, - get_current_active_superuser_async, -) +from app.db.oper.systemconfig import SystemConfigOper +from app.api.deps import get_current_active_superuser, get_current_active_superuser_async from app.factory import app from app.adapters.external.server import MoviePilotServerHelper from app.adapters.external.market import PluginHelper diff --git a/app/api/endpoints/search.py b/app/api/endpoints/search.py index 5b9a2e553..0fa902abe 100644 --- a/app/api/endpoints/search.py +++ b/app/api/endpoints/search.py @@ -14,7 +14,8 @@ from app.application.security.access import verify_resource_token, verify_token from app.runtime.localization import LocaleHelper from app.runtime.log import logger from app.schemas.types import MediaSource, MediaType -from app.domain.media import normalize_music_type, resolve_media_identity +from app.domain.media import normalize_music_type +from app.schemas.media import resolve_media_identity from app.application.security.url import SecurityUtils router = ResponseAPIRouter() diff --git a/app/api/endpoints/site.py b/app/api/endpoints/site.py index 8f70a7524..237dbb0d0 100644 --- a/app/api/endpoints/site.py +++ b/app/api/endpoints/site.py @@ -20,9 +20,9 @@ from app.db.models.site import Site from app.db.models.siteicon import SiteIcon from app.db.models.sitestatistic import SiteStatistic from app.db.models.siteuserdata import SiteUserData -from app.db.site_oper import SiteOper -from app.db.systemconfig_oper import SystemConfigOper -from app.db.user_oper import ( +from app.db.oper.site import SiteOper +from app.db.oper.systemconfig import SystemConfigOper +from app.api.deps import ( get_current_active_manage_user, get_current_active_manage_user_async, get_current_active_superuser, diff --git a/app/api/endpoints/storage.py b/app/api/endpoints/storage.py index 1868be428..9a1eee683 100644 --- a/app/api/endpoints/storage.py +++ b/app/api/endpoints/storage.py @@ -15,7 +15,7 @@ from app.chain.transfer import TransferChain from app.runtime.config import settings from app.application.security.access import verify_token from app.db.models import User -from app.db.user_oper import ( +from app.api.deps import ( get_current_active_manage_user, get_current_active_superuser, get_current_active_superuser_async, diff --git a/app/api/endpoints/subscribe.py b/app/api/endpoints/subscribe.py index 400681e2a..f914ede6b 100644 --- a/app/api/endpoints/subscribe.py +++ b/app/api/endpoints/subscribe.py @@ -17,8 +17,8 @@ from app.db import get_async_db, get_db from app.db.models.subscribe import Subscribe from app.db.models.subscribehistory import SubscribeHistory from app.db.models.user import User -from app.db.systemconfig_oper import SystemConfigOper -from app.db.user_oper import get_current_active_user, get_current_active_user_async +from app.db.oper.systemconfig import SystemConfigOper +from app.api.deps import get_current_active_user, get_current_active_user_async from app.adapters.external.server import MoviePilotServerHelper from app.runtime.log import logger from app.scheduler import Scheduler @@ -31,7 +31,7 @@ from app.schemas.types import ( EventType, SystemConfigKey, ) -from app.domain.media import normalize_media_source, resolve_media_identity +from app.schemas.media import normalize_media_source, resolve_media_identity router = ResponseAPIRouter() diff --git a/app/api/endpoints/system.py b/app/api/endpoints/system.py index 9f2a34442..56645b66d 100644 --- a/app/api/endpoints/system.py +++ b/app/api/endpoints/system.py @@ -29,12 +29,8 @@ from app.domain.metainfo import MetaInfo from app.runtime.extensions.module_manager import ModuleManager from app.application.security.access import verify_apitoken, verify_resource_token, verify_token from app.db.models import User -from app.db.systemconfig_oper import SystemConfigOper -from app.db.user_oper import ( - get_current_active_superuser, - get_current_active_superuser_async, - get_current_active_user_async, -) +from app.db.oper.systemconfig import SystemConfigOper +from app.api.deps import get_current_active_superuser, get_current_active_superuser_async, get_current_active_user_async from app.application.image import ImageHelper from app.runtime.localization import LocaleHelper from app.adapters.external.market import ( diff --git a/app/api/endpoints/tmdb.py b/app/api/endpoints/tmdb.py index bd89cf931..f22d8b935 100644 --- a/app/api/endpoints/tmdb.py +++ b/app/api/endpoints/tmdb.py @@ -8,8 +8,8 @@ from app.chain.tmdb import TmdbChain from app.runtime.config import settings from app.application.security.access import verify_token from app.db.models.user import User -from app.db.systemconfig_oper import SystemConfigOper -from app.db.user_oper import get_current_active_superuser_async +from app.db.oper.systemconfig import SystemConfigOper +from app.api.deps import get_current_active_superuser_async from app.modules.themoviedb.tmdb_cache import TmdbCache from app.schemas.types import MediaType, SystemConfigKey diff --git a/app/api/endpoints/torrent.py b/app/api/endpoints/torrent.py index 970ee21fe..dc47ee870 100644 --- a/app/api/endpoints/torrent.py +++ b/app/api/endpoints/torrent.py @@ -11,10 +11,7 @@ from app.domain.context import MediaInfo, MusicInfo from app.domain.meta.metamusic import MetaMusic from app.domain.metainfo import MetaInfo from app.db.models import User -from app.db.user_oper import ( - get_current_active_superuser, - get_current_active_superuser_async, -) +from app.api.deps import get_current_active_superuser, get_current_active_superuser_async from app.schemas.types import ( MUSIC_ENTITY_RECORDING, MediaSource, @@ -22,11 +19,8 @@ from app.schemas.types import ( MusicTargetEntityType, ) from app.foundation.crypto import HashUtils -from app.domain.media import ( - is_music_media_source, - normalize_music_type, - resolve_media_identity, -) +from app.domain.media import is_music_media_source, normalize_music_type +from app.schemas.media import resolve_media_identity router = ResponseAPIRouter() diff --git a/app/api/endpoints/transfer.py b/app/api/endpoints/transfer.py index c4295fd2e..db8549d78 100644 --- a/app/api/endpoints/transfer.py +++ b/app/api/endpoints/transfer.py @@ -13,10 +13,7 @@ from app.application.security.access import verify_token, verify_apitoken from app.db import get_db from app.db.models import User from app.db.models.transferhistory import TransferHistory -from app.db.user_oper import ( - get_current_active_manage_user, - get_current_active_superuser, -) +from app.api.deps import get_current_active_manage_user, get_current_active_superuser from app.application.directory import DirectoryHelper from app.runtime.log import logger from app.schemas import ( diff --git a/app/api/endpoints/user.py b/app/api/endpoints/user.py index 71b171284..1b8fd466e 100644 --- a/app/api/endpoints/user.py +++ b/app/api/endpoints/user.py @@ -10,12 +10,8 @@ from app.api.response import ResponseAPIRouter from app.application.security.access import get_password_hash from app.db import get_async_db from app.db.models.user import User -from app.db.user_oper import ( - get_current_active_superuser_async, - get_current_active_user_async, - get_current_active_user, -) -from app.db.userconfig_oper import UserConfigOper +from app.api.deps import get_current_active_superuser_async, get_current_active_user_async, get_current_active_user +from app.db.oper.userconfig import UserConfigOper router = ResponseAPIRouter() diff --git a/app/api/endpoints/workflow.py b/app/api/endpoints/workflow.py index 202f46095..51a6d59a2 100644 --- a/app/api/endpoints/workflow.py +++ b/app/api/endpoints/workflow.py @@ -14,12 +14,9 @@ from app.runtime.extensions.plugin_manager import PluginManager from app.workflow import WorkFlowManager from app.db import get_async_db, get_db from app.db.models import Workflow, User -from app.db.systemconfig_oper import SystemConfigOper -from app.db.user_oper import ( - get_current_active_manage_user, - get_current_active_manage_user_async, -) -from app.db.workflow_oper import WorkflowOper +from app.db.oper.systemconfig import SystemConfigOper +from app.api.deps import get_current_active_manage_user, get_current_active_manage_user_async +from app.db.oper.workflow import WorkflowOper from app.adapters.external.server import MoviePilotServerHelper from app.scheduler import Scheduler from app.schemas.types import EventType, EVENT_TYPE_NAMES diff --git a/app/application/directory.py b/app/application/directory.py index 390131659..4167b6abc 100644 --- a/app/application/directory.py +++ b/app/application/directory.py @@ -4,7 +4,7 @@ from typing import List, Optional, Tuple from app import schemas from app.domain.context import MediaInfo -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas.types import MediaType, StorageSchema, SystemConfigKey from app.adapters.system.host import SystemUtils diff --git a/app/application/filter.py b/app/application/filter.py index 6121fd9b5..47ac7698c 100644 --- a/app/application/filter.py +++ b/app/application/filter.py @@ -1,6 +1,6 @@ from typing import List, Optional -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.domain.context import MediaInfo from app.schemas import CustomRule, FilterRuleGroup from app.schemas.types import SystemConfigKey diff --git a/app/application/history.py b/app/application/history.py index c0725d6a0..28e6ba6bb 100644 --- a/app/application/history.py +++ b/app/application/history.py @@ -1,10 +1,16 @@ -from typing import Any, Dict, Optional +from typing import Any, Dict, Optional, Union +from app.domain.context import MediaInfo, MusicInfo +from app.schemas.media import resolve_media_identity +from app.domain.meta.metabase import MetaBase +from app.domain.meta.metamusic import MetaMusic from app.runtime.cache import TTLCache from app.runtime.config import settings from app.db.models.transferhistory import TransferHistory -from app.db.transferhistory_oper import TransferHistoryOper +from app.db.oper.transferhistory import TransferHistoryOper from app.runtime.log import logger +from app.schemas import FileItem, TransferInfo +from app.schemas.types import MUSIC_ENTITY_RECORDING # 失败重试次数的合法区间。下界为 1:一次瞬时故障(网络抖动、TMDB 瞬断、移动失败) # 不该让文件永久漏整理,所以不允许关闭重试;上界为 10:永远识别不出的文件重试再多 @@ -420,3 +426,158 @@ def describe_history_gate(history: Optional[TransferHistory], if recorded_size is None and current_size is None: return f"成功记录 #{history.id},大小不可比对" return f"成功记录 #{history.id},大小 {recorded_size} -> {current_size}" + + +# --------------------------------------------------------------------------- # +# 整理历史的写入路径 +# +# 这两个函数把 FileItem / MetaBase / MediaInfo / TransferInfo 四个领域对象翻译成 +# 一行整理历史,是整理历史表的唯一写入口。它们此前长在 TransferHistoryOper 上,但 +# 拼标题、拆季集、取海报、判音乐字段都是整理链的业务规则而非数据访问——Oper 只该 +# 收敛查询,领域对象不该出现在它的入参里。搬到本模块与查重闸(读侧)作伴:同一张 +# 表的读写规则放在一起,字段含义只有一处需要维护。 +# --------------------------------------------------------------------------- # + +def _history_title(meta: MetaBase, + mediainfo: Optional[Union[MediaInfo, MusicInfo]] = None) -> Optional[str]: + """音乐文件优先记录曲目标题,其它媒体保持识别标题。""" + if isinstance(meta, MetaMusic) and meta.title: + return meta.title + if mediainfo and mediainfo.title: + return mediainfo.title + return meta.name + + +def add_transfer_success(fileitem: FileItem, mode: str, meta: MetaBase, + mediainfo: Union[MediaInfo, MusicInfo], transferinfo: TransferInfo, + downloader: Optional[str] = None, + download_hash: Optional[str] = None, + transfer_history_oper: Optional[TransferHistoryOper] = None + ) -> Optional[TransferHistory]: + """ + 新增转移成功历史记录。 + :param fileitem: 源文件项 + :param mode: 整理方式 + :param meta: 文件名识别结果 + :param mediainfo: 媒体识别结果 + :param transferinfo: 整理结果 + :param downloader: 下载器 + :param download_hash: 种子 hash + :param transfer_history_oper: 复用的历史操作对象,未传时新建 + :return: 落库后的整理记录 + """ + oper = transfer_history_oper or TransferHistoryOper() + media_source, media_id = resolve_media_identity(media=mediainfo) + return oper.add_force( + src=fileitem.path, + src_storage=fileitem.storage, + src_fileitem=fileitem.model_dump(), + dest=transferinfo.target_item.path if transferinfo.target_item else None, + dest_storage=transferinfo.target_item.storage if transferinfo.target_item else None, + dest_fileitem=transferinfo.target_item.model_dump() if transferinfo.target_item else None, + mode=mode, + type=mediainfo.type.value, + category=mediainfo.category, + title=_history_title(meta, mediainfo), + year=mediainfo.year, + media_source=media_source, + media_id=media_id, + music_type=getattr(mediainfo, "music_type", None), + total_tracks=getattr(mediainfo, "total_tracks", None), + audio_format=getattr(meta, "audio_format", None), + audio_lossless=getattr(meta, "audio_lossless", None), + bit_depth=getattr(meta, "bit_depth", None), + sample_rate=getattr(meta, "sample_rate", None), + bitrate=getattr(meta, "bitrate", None), + seasons=meta.season, + episodes=meta.episode, + image=mediainfo.get_poster_image(), + downloader=downloader, + download_hash=download_hash, + status=1, + files=transferinfo.file_list + ) + + +def add_transfer_fail(fileitem: FileItem, mode: str, meta: MetaBase, + mediainfo: Optional[Union[MediaInfo, MusicInfo]] = None, + transferinfo: Optional[TransferInfo] = None, + downloader: Optional[str] = None, + download_hash: Optional[str] = None, + transfer_history_oper: Optional[TransferHistoryOper] = None + ) -> Optional[TransferHistory]: + """ + 新增转移失败历史记录。 + + 识别结果与整理结果齐备时按完整字段落库;缺任一项则走「未识别到媒体信息」分支, + 此时只有文件名解析出的元数据可用,不写目标路径。 + :param fileitem: 源文件项 + :param mode: 整理方式 + :param meta: 文件名识别结果 + :param mediainfo: 媒体识别结果,未识别时为 None + :param transferinfo: 整理结果,未进入整理时为 None + :param downloader: 下载器 + :param download_hash: 种子 hash + :param transfer_history_oper: 复用的历史操作对象,未传时新建 + :return: 落库后的整理记录 + """ + oper = transfer_history_oper or TransferHistoryOper() + if mediainfo and transferinfo: + media_source, media_id = resolve_media_identity(media=mediainfo) + his = oper.add_force( + src=fileitem.path, + src_storage=fileitem.storage, + src_fileitem=fileitem.model_dump(), + dest=transferinfo.target_item.path if transferinfo.target_item else None, + dest_storage=transferinfo.target_item.storage if transferinfo.target_item else None, + dest_fileitem=transferinfo.target_item.model_dump() if transferinfo.target_item else None, + mode=mode, + type=mediainfo.type.value, + category=mediainfo.category, + title=_history_title(meta, mediainfo), + year=mediainfo.year or meta.year, + media_source=media_source, + media_id=media_id, + music_type=getattr(mediainfo, "music_type", None), + total_tracks=getattr(mediainfo, "total_tracks", None), + audio_format=getattr(meta, "audio_format", None), + audio_lossless=getattr(meta, "audio_lossless", None), + bit_depth=getattr(meta, "bit_depth", None), + sample_rate=getattr(meta, "sample_rate", None), + bitrate=getattr(meta, "bitrate", None), + seasons=meta.season, + episodes=meta.episode, + image=mediainfo.get_poster_image(), + downloader=downloader, + download_hash=download_hash, + episode_group=mediainfo.episode_group, + status=0, + errmsg=transferinfo.message or '未知错误', + files=transferinfo.file_list + ) + else: + media_source, media_id = resolve_media_identity(media=meta) + his = oper.add_force( + type=meta.type.value if meta.type else None, + title=_history_title(meta), + year=meta.year, + media_source=media_source, + media_id=media_id, + music_type=MUSIC_ENTITY_RECORDING if isinstance(meta, MetaMusic) else None, + audio_format=getattr(meta, "audio_format", None), + audio_lossless=getattr(meta, "audio_lossless", None), + bit_depth=getattr(meta, "bit_depth", None), + sample_rate=getattr(meta, "sample_rate", None), + bitrate=getattr(meta, "bitrate", None), + src=fileitem.path, + src_storage=fileitem.storage, + src_fileitem=fileitem.model_dump(), + mode=mode, + seasons=meta.season, + episodes=meta.episode, + downloader=downloader, + download_hash=download_hash, + status=0, + errmsg="未识别到媒体信息" + ) + return his diff --git a/app/application/mediaserver.py b/app/application/mediaserver.py index 85a23d34e..634ee20c6 100644 --- a/app/application/mediaserver.py +++ b/app/application/mediaserver.py @@ -4,7 +4,7 @@ from typing import Any, Optional from app import schemas from app.domain.context import MusicInfo -from app.domain.media import normalize_media_source, resolve_media_identity +from app.schemas.media import normalize_media_source, resolve_media_identity from app.runtime.extensions.service_registry import ServiceBaseHelper from app.schemas import MediaServerConf, ServiceInfo from app.schemas.types import ( diff --git a/app/application/messaging/message.py b/app/application/messaging/message.py index 3e8bdde96..cb3667c62 100644 --- a/app/application/messaging/message.py +++ b/app/application/messaging/message.py @@ -18,7 +18,7 @@ from app.runtime.config import global_vars from app.domain.context import MediaInfo, MusicInfo, TorrentInfo from app.domain.meta.metabase import MetaBase from app.domain.meta.metamusic import MetaMusic -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas.message import Notification from app.schemas.tmdb import TmdbEpisode diff --git a/app/application/recognition.py b/app/application/recognition.py index b9138e2a5..8a1848f85 100644 --- a/app/application/recognition.py +++ b/app/application/recognition.py @@ -1,6 +1,6 @@ from typing import Optional -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.schemas.types import SystemConfigKey diff --git a/app/application/security/auth.py b/app/application/security/auth.py index 6320f0ddf..240e3cd8a 100644 --- a/app/application/security/auth.py +++ b/app/application/security/auth.py @@ -10,8 +10,8 @@ from app import schemas from app.application.security import access as security from app.runtime.config import settings from app.db.models.user import User -from app.db.systemconfig_oper import SystemConfigOper -from app.db.user_oper import UserOper +from app.db.oper.systemconfig import SystemConfigOper +from app.db.oper.user import UserOper from app.application.site.sites import SitesHelper # pylint: disable=no-name-in-module from app.schemas.types import SystemConfigKey from app.foundation.singleton import Singleton diff --git a/app/application/storage.py b/app/application/storage.py index e9cff4477..6d868525e 100644 --- a/app/application/storage.py +++ b/app/application/storage.py @@ -1,7 +1,7 @@ from typing import List, Optional from app import schemas -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.schemas.types import SystemConfigKey diff --git a/app/application/subscribe.py b/app/application/subscribe.py new file mode 100644 index 000000000..3c1a82e62 --- /dev/null +++ b/app/application/subscribe.py @@ -0,0 +1,122 @@ +""" +订阅的写入路径。 + +这两个函数把 MediaInfo / MusicInfo 翻译成一行订阅,是订阅表的唯一写入口。翻译此前 +长在 SubscribeOper.add 上,但取标题、选海报尺寸、判音乐实体、决定哪几个字段构成一条 +订阅的身份,都是订阅业务的规则而非数据访问——Oper 只该收敛查询,领域对象不该出现在 +它的入参里。搬上来之后 SubscribeOper 收到的是纯粹的持久化字典,与 +app/application/history.py 里整理历史的写入路径同构。 + +留在 Oper 的是列类型强转与建库时间戳:那几步是为 PostgreSQL 的严格类型检查和订阅表 +自己的列类型而存在的,跟着列走比跟着调用方走更不容易漂。 + +字段映射错了不会报错,只会让订阅静静地记错——而搜索、洗版、完成判定、去重全都读这 +张表。同步与异步是两份逐字复制的实现,改一条漏一条就是真实缺陷,故翻译与身份构造由 +下方 _translate 单点承担,两条链路只在「怎么查、怎么写」上分叉。 +""" +from typing import Optional, Tuple + +from app.db.oper.subscribe import SubscribeOper +from app.domain.context import MediaInfo, MusicInfo +from app.schemas.media import resolve_media_identity +from app.schemas.types import MUSIC_ENTITY_ALBUM, MediaType + +# 身份不完整时的固定返回。身份不全的订阅写进去就是一条永远匹配不上资源的僵尸订阅, +# 而后续按身份去重也会失效,所以必须在查询与建模之前短路 +INCOMPLETE_IDENTITY = (0, "媒体身份不完整") + + +def _music_entity(mediainfo: MediaInfo | MusicInfo) -> Optional[str]: + """ + 取音乐实体类型;非音乐媒体一律为空。 + + 影视订阅带着 music_type 会被音乐去重逻辑当成音乐实体,造成串号。查重身份与写入 + 字段都取这一个值,避免两处各算一遍后悄悄分叉。 + :param mediainfo: 识别结果 + :return: 音乐实体类型,非音乐媒体为 None + """ + if mediainfo.type != MediaType.MUSIC: + return None + return getattr(mediainfo, "music_type", None) + + +def _translate(mediainfo: MediaInfo | MusicInfo, + kwargs: dict) -> Optional[Tuple[dict, dict, Optional[str]]]: + """ + 把识别结果翻译成查重身份与写入字段。 + + :param mediainfo: 识别结果 + :param kwargs: 调用方传入的订阅设置,媒体相关的同名字段会被识别结果覆盖 + :return: (查重身份, 写入字段, 限定用户);媒体身份不完整时返回 None + """ + owner_scope = bool(kwargs.pop("owner_scope", False)) + username = kwargs.get("username") if owner_scope else None + media_source, media_id = resolve_media_identity( + media=mediainfo, + media_source=kwargs.get("media_source"), + media_id=kwargs.get("media_id"), + ) + if not media_source or not media_id: + return None + music_type = _music_entity(mediainfo) + identity = { + "media_source": str(media_source), + "media_id": media_id, + "music_type": music_type, + "season": kwargs.get("season"), + "episode_group": mediainfo.episode_group, + } + payload = dict(kwargs) + payload.update({ + "name": mediainfo.title, + "year": mediainfo.year, + "type": mediainfo.type.value, + "media_source": str(media_source), + "media_id": media_id, + "episode_group": mediainfo.episode_group, + "poster": mediainfo.get_poster_image(), + "backdrop": mediainfo.get_backdrop_image(), + "vote": mediainfo.vote_average, + "description": mediainfo.overview, + "music_type": music_type, + # 整专完成判定拿 total_tracks 当分母,单曲带着专辑的曲目数会永远判不到完成 + "total_tracks": getattr(mediainfo, "total_tracks", None) + if music_type == MUSIC_ENTITY_ALBUM else None, + }) + return identity, payload, username + + +def add_subscribe(mediainfo: MediaInfo | MusicInfo, + subscribe_oper: Optional[SubscribeOper] = None, + **kwargs) -> Tuple[int, str]: + """ + 新增订阅。 + :param mediainfo: 识别结果 + :param subscribe_oper: 复用的订阅操作对象,未传时新建 + :param kwargs: 订阅设置;owner_scope 为真时按用户名限定查重范围 + :return: (订阅 ID, 结果说明);ID 为 0 表示未新增 + """ + translated = _translate(mediainfo, kwargs) + if translated is None: + return INCOMPLETE_IDENTITY + identity, payload, username = translated + oper = subscribe_oper or SubscribeOper() + return oper.add(identity=identity, payload=payload, username=username) + + +async def async_add_subscribe(mediainfo: MediaInfo | MusicInfo, + subscribe_oper: Optional[SubscribeOper] = None, + **kwargs) -> Tuple[int, str]: + """ + 异步新增订阅。 + :param mediainfo: 识别结果 + :param subscribe_oper: 复用的订阅操作对象,未传时新建 + :param kwargs: 订阅设置;owner_scope 为真时按用户名限定查重范围 + :return: (订阅 ID, 结果说明);ID 为 0 表示未新增 + """ + translated = _translate(mediainfo, kwargs) + if translated is None: + return INCOMPLETE_IDENTITY + identity, payload, username = translated + oper = subscribe_oper or SubscribeOper() + return await oper.async_add(identity=identity, payload=payload, username=username) diff --git a/app/application/torrent.py b/app/application/torrent.py index c20e50433..2e211bd20 100644 --- a/app/application/torrent.py +++ b/app/application/torrent.py @@ -13,12 +13,12 @@ from app.domain.context import Context, TorrentInfo, MediaInfo from app.domain.meta.metabase import MetaBase from app.domain.meta.metamusic import audio_quality_tier, normalize_audio_format, parse_audio_quality from app.domain.metainfo import MetaInfo -from app.db.site_oper import SiteOper -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.site import SiteOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas.types import MediaType, SystemConfigKey from app.adapters.network.http import RequestUtils -from app.domain.media import resolve_media_identity +from app.schemas.media import resolve_media_identity from app.domain.string import StringUtils diff --git a/app/application/transfer.py b/app/application/transfer.py new file mode 100644 index 000000000..d7b2ac888 --- /dev/null +++ b/app/application/transfer.py @@ -0,0 +1,91 @@ +""" +整理任务:整理链的进程内工作项。 + +TransferTask 此前住在 app/schemas/transfer.py,但它不是出网的 DTO——meta 装的是领域侧 +的 MetaBase 子类,mediainfo 装的是领域侧的 MediaInfo / MusicInfo,都带行为而非纯数据。 +放在 app.schemas 的代价是它没法命名自己真正装的类型:app.schemas 一旦 import 领域类型, +app.schemas -> app.schemas.transfer -> app.domain.* -> app.schemas.types -> app.schemas +就闭环,仓库自己的 test_migrated_modules_are_not_in_import_cycles 会红(已实测)。于是 +两个字段只能标成 Optional[Any],把「这里到底能放什么」这件事整个交给了口头约定。 + +搬到应用层就没有这个约束:app.application 允许依赖 app.domain 与 app.schemas,两个 +字段因此能标出真实类型。它面向前端的投影仍是 app/schemas/transfer.py 里的 +TransferJob / TransferJobTask,那两个用 app.schemas 的同名 DTO——一个是工作项,一个是 +视图,分开表达之后两边都不必再迁就对方。 +""" +from pathlib import Path +from typing import Callable, List, Optional, Union + +from pydantic import BaseModel, ConfigDict + +from app.domain.context import MediaInfo, MusicInfo +from app.domain.meta.metabase import MetaBase +from app.schemas.file import FileItem +from app.schemas.history import DownloadHistory +from app.schemas.media import OptionalMediaIdentityMixin +from app.schemas.system import TransferDirectoryConf +from app.schemas.tmdb import TmdbEpisode +from app.schemas.transfer import TransferInfo +from app.schemas.types import MediaSource, MediaType + + +class TransferTask(OptionalMediaIdentityMixin, BaseModel): + """ + 文件整理任务。 + """ + + # MetaBase 与 MediaInfo / MusicInfo 都是普通类而非 BaseModel,pydantic 需要显式放行 + model_config = ConfigDict(arbitrary_types_allowed=True) + + fileitem: FileItem + meta: Optional[MetaBase] = None + mediainfo: Optional[Union[MusicInfo, MediaInfo]] = None + media_source: Optional[MediaSource] = None + media_id: Optional[str] = None + mtype: Optional[MediaType] = None + target_directory: Optional[TransferDirectoryConf] = None + target_storage: Optional[str] = None + target_path: Optional[Path] = None + transfer_type: Optional[str] = None + scrape: Optional[bool] = False + library_type_folder: Optional[bool] = False + library_category_folder: Optional[bool] = False + episodes_info: Optional[List[TmdbEpisode]] = None + username: Optional[str] = None + downloader: Optional[str] = None + download_hash: Optional[str] = None + download_history: Optional[DownloadHistory] = None + transfer_batch_id: Optional[str] = None + manual: Optional[bool] = False + background: Optional[bool] = True + preview: Optional[bool] = False + + def to_dict(self): + """ + 返回字典。 + + meta 与 mediainfo 用 to_dict() 而非 model_dump():它们是领域对象,没有 + model_dump。此前这里写的是 model_dump(),仓内无人调用才一直没炸——字段类型 + 标成 Any 时,这种错配静态检查也看不出来。 + """ + dicts = vars(self).copy() + dicts["fileitem"] = self.fileitem.model_dump() if self.fileitem else None + dicts["meta"] = self.meta.to_dict() if self.meta else None + dicts["mediainfo"] = self.mediainfo.to_dict() if self.mediainfo else None + dicts["target_directory"] = self.target_directory.model_dump() if self.target_directory else None + return dicts + + +class TransferQueue(BaseModel): + """ + 异步整理队列信息。 + + 和 TransferTask 一起从 app/schemas 搬来:它装着一个 TransferTask 和一个回调函数, + 回调根本不可序列化,因此从来就不是 DTO,只是恰好和视图模型住在同一个文件里。 + """ + # 任务信息 + task: Optional[TransferTask] = None + # 回调函数 + callback: Optional[Callable] = None + # 整理结果 + result: Optional[TransferInfo] = None diff --git a/app/chain/__init__.py b/app/chain/__init__.py index ebbe329ba..606c87b6b 100644 --- a/app/chain/__init__.py +++ b/app/chain/__init__.py @@ -20,9 +20,9 @@ from app.domain.meta.metabase import MetaBase from app.domain.meta.metamusic import MetaMusic from app.runtime.extensions.module_manager import ModuleManager from app.runtime.extensions.plugin_manager import PluginManager -from app.db.message_oper import MessageOper -from app.db.systemconfig_oper import SystemConfigOper -from app.db.user_oper import UserOper +from app.db.oper.message import MessageOper +from app.db.oper.systemconfig import SystemConfigOper +from app.db.oper.user import UserOper from app.application.messaging.message import MessageHelper, MessageQueueManager, MessageTemplateHelper from app.adapters.external.server import MoviePilotServerHelper from app.runtime.extensions.service_registry import ServiceConfigHelper @@ -42,7 +42,7 @@ from app.schemas import ( MessageResponse, ) from app.foundation.identity import normalize_internal_user_id -from app.domain.media import resolve_media_identity +from app.schemas.media import resolve_media_identity from app.schemas.message import ChannelCapability, ChannelCapabilityManager from app.schemas.category import CategoryConfig from app.schemas.types import ( @@ -1773,7 +1773,7 @@ class ChainBase(metaclass=ABCMeta): self, message: Optional[Notification] = None, meta: Optional[MetaBase] = None, - mediainfo: Optional[MediaInfo] = None, + mediainfo: Optional[Union[MediaInfo, MusicInfo]] = None, torrentinfo: Optional[TorrentInfo] = None, transferinfo: Optional[TransferInfo] = None, **kwargs, @@ -1889,7 +1889,7 @@ class ChainBase(metaclass=ABCMeta): self, message: Optional[Notification] = None, meta: Optional[MetaBase] = None, - mediainfo: Optional[MediaInfo] = None, + mediainfo: Optional[Union[MediaInfo, MusicInfo]] = None, torrentinfo: Optional[TorrentInfo] = None, transferinfo: Optional[TransferInfo] = None, **kwargs, diff --git a/app/chain/download.py b/app/chain/download.py index 5f2977012..394a3eb9d 100644 --- a/app/chain/download.py +++ b/app/chain/download.py @@ -26,9 +26,9 @@ from app.runtime.events import eventmanager, Event from app.domain.meta.metabase import MetaBase from app.domain.meta.metamusic import MetaMusic from app.domain.metainfo import MetaInfo -from app.db.downloadfailure_oper import DownloadFailureOper -from app.db.downloadhistory_oper import DownloadHistoryOper -from app.db.mediaserver_oper import MediaServerOper +from app.db.oper.downloadfailure import DownloadFailureOper +from app.db.oper.downloadhistory import DownloadHistoryOper +from app.db.oper.mediaserver import MediaServerOper from app.application.directory import DirectoryHelper, validate_download_save_path from app.runtime.thread import ThreadHelper from app.application.torrent import TorrentHelper @@ -38,7 +38,7 @@ from app.schemas import ExistMediaInfo, FileURI, NotExistMediaInfo, DownloaderTo from app.schemas.types import MUSIC_ENTITY_ALBUM, MediaSource, MediaType, TorrentStatus, EventType, MessageChannel, NotificationType, ContentType, \ ChainEventType from app.adapters.network.http import RequestUtils -from app.domain.media import build_media_key, resolve_media_identity +from app.schemas.media import build_media_key, resolve_media_identity from app.domain.string import StringUtils from app.adapters.system.host import SystemUtils diff --git a/app/chain/media.py b/app/chain/media.py index db7762a33..fe2addf16 100644 --- a/app/chain/media.py +++ b/app/chain/media.py @@ -37,11 +37,8 @@ from app.schemas.types import ( MediaSourceSelection, MediaType, ) -from app.domain.media import ( - is_music_media_source, - normalize_media_source, - resolve_media_identity, -) +from app.domain.media import is_music_media_source +from app.schemas.media import normalize_media_source, resolve_media_identity from app.foundation.singleton import Singleton from app.foundation.text import convert as zhconv_convert from app.domain.string import StringUtils @@ -643,9 +640,9 @@ class MediaChain(ChainBase, metaclass=Singleton): def supplement_tmdb_info( self, - mediainfo: Optional[MediaInfo], + mediainfo: Optional[Union[MediaInfo, MusicInfo]], metainfo: Optional[MetaBase] = None, - ) -> Optional[MediaInfo]: + ) -> Optional[Union[MediaInfo, MusicInfo]]: """ 为任意主识别源补充 TMDB 辅助信息,同时保留原始媒体身份。 @@ -655,7 +652,10 @@ class MediaChain(ChainBase, metaclass=Singleton): """ if not mediainfo: return None - if mediainfo.type == MediaType.MUSIC: + # 音乐原样返回:下面全是 TMDB 影视字段,MusicInfo 上根本没有。用 isinstance + # 而不只看 type,一来静态检查能据此收窄(.type == 的比较收窄不了类型),二来 + # type 没被正确赋值的 MusicInfo 也挡得住,不至于到下一行才 AttributeError + if isinstance(mediainfo, MusicInfo) or mediainfo.type == MediaType.MUSIC: return mediainfo if mediainfo.tmdb_id and mediainfo.tmdb_info and mediainfo.genre_ids: return mediainfo diff --git a/app/chain/mediaserver.py b/app/chain/mediaserver.py index e396edd9c..e46b3a4a1 100644 --- a/app/chain/mediaserver.py +++ b/app/chain/mediaserver.py @@ -4,7 +4,7 @@ from typing import Callable, Dict, List, Union, Optional, Generator, Any from app.chain import ChainBase from app.runtime.config import global_vars -from app.db.mediaserver_oper import MediaServerOper +from app.db.oper.mediaserver import MediaServerOper from app.runtime.extensions.service_registry import ServiceConfigHelper from app.runtime.log import logger from app.schemas import MediaServerLibrary, MediaServerItem, MediaServerSeasonInfo, MediaServerPlayItem diff --git a/app/chain/message.py b/app/chain/message.py index 0762c816d..0f50f8420 100644 --- a/app/chain/message.py +++ b/app/chain/message.py @@ -26,8 +26,8 @@ from app.runtime.config import settings, global_vars from app.domain.context import MediaInfo, Context from app.domain.meta.metabase import MetaBase from app.db.models import TransferHistory -from app.db.transferhistory_oper import TransferHistoryOper -from app.db.user_oper import UserOper +from app.db.oper.transferhistory import TransferHistoryOper +from app.db.oper.user import UserOper from app.application.directory import DirectoryHelper from app.application.messaging.interaction import ( agent_interaction_manager, @@ -42,7 +42,7 @@ from app.schemas.message import ChannelCapabilityManager, ChannelCapability from app.schemas.system import TransferDirectoryConf from app.schemas.types import EventType, MessageChannel, MediaType from app.adapters.network.http import RequestUtils -from app.domain.media import build_media_key, resolve_media_identity +from app.schemas.media import build_media_key, resolve_media_identity from app.domain.string import StringUtils diff --git a/app/chain/recommend.py b/app/chain/recommend.py index 11cf92f2c..91bcb4dda 100644 --- a/app/chain/recommend.py +++ b/app/chain/recommend.py @@ -19,7 +19,7 @@ from app.schemas.types import ( MediaSource, ) from app.runtime.execution import log_execution_time -from app.domain.media import normalize_media_source +from app.schemas.media import normalize_media_source from app.foundation.singleton import Singleton diff --git a/app/chain/scraping.py b/app/chain/scraping.py index 882129d00..3b43c6f32 100644 --- a/app/chain/scraping.py +++ b/app/chain/scraping.py @@ -26,7 +26,7 @@ from app.runtime.events import eventmanager, Event from app.domain.meta.metabase import MetaBase from app.domain.meta.metamusic import MetaMusic from app.domain.metainfo import MetaInfo, MetaInfoPath -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.application.audio import AudioMetadataHelper from app.runtime.log import logger from app.schemas import FileItem @@ -43,11 +43,8 @@ from app.schemas.types import ( SystemConfigKey, ) from app.adapters.network.http import RequestUtils -from app.domain.media import ( - is_music_media_source, - normalize_media_source, - resolve_media_identity, -) +from app.domain.media import is_music_media_source +from app.schemas.media import normalize_media_source, resolve_media_identity from app.runtime.reload import ConfigReloadMixin from app.foundation.singleton import Singleton from app.domain.string import StringUtils diff --git a/app/chain/search.py b/app/chain/search.py index 92f8f88d4..afcdbb83c 100644 --- a/app/chain/search.py +++ b/app/chain/search.py @@ -21,7 +21,7 @@ from app.runtime.events import eventmanager, Event from app.domain.meta.metamusic import MetaMusic from app.domain.metainfo import MetaInfo from app.domain.context import MusicInfo -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.progress import ProgressHelper from app.application.site.sites import SitesHelper # pylint: disable=no-name-in-module from app.application.torrent import TorrentHelper @@ -35,11 +35,7 @@ from app.schemas.types import ( ProgressKey, SystemConfigKey, ) -from app.domain.media import ( - build_media_key, - parse_media_key, - resolve_media_identity, -) +from app.schemas.media import build_media_key, parse_media_key, resolve_media_identity from app.domain.string import StringUtils from app.foundation.text import convert as zhconv_convert diff --git a/app/chain/site.py b/app/chain/site.py index e88e7ebb8..ead8fa06b 100644 --- a/app/chain/site.py +++ b/app/chain/site.py @@ -11,8 +11,8 @@ from app.chain import ChainBase from app.runtime.config import global_vars, settings from app.runtime.events import Event, eventmanager from app.db.models.site import Site -from app.db.site_oper import SiteOper -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.site import SiteOper +from app.db.oper.systemconfig import SystemConfigOper from app.adapters.network.browser import PlaywrightHelper from app.adapters.network.cloudflare import under_challenge from app.application.security.cookie import CookieHelper diff --git a/app/chain/subscribe.py b/app/chain/subscribe.py index b2ddc5b48..59b9fdcc5 100644 --- a/app/chain/subscribe.py +++ b/app/chain/subscribe.py @@ -27,11 +27,11 @@ from app.domain.meta.metabase import MetaBase from app.domain.meta.metamusic import MetaMusic from app.domain.meta.words import WordsMatcher from app.domain.metainfo import MetaInfo -from app.db.downloadhistory_oper import DownloadHistoryOper +from app.db.oper.downloadhistory import DownloadHistoryOper from app.db.models.subscribe import Subscribe -from app.db.site_oper import SiteOper -from app.db.subscribe_oper import SubscribeOper -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.site import SiteOper +from app.db.oper.subscribe import SubscribeOper +from app.db.oper.systemconfig import SystemConfigOper from app.application.messaging.interaction import ( SlashInteractionManager, build_navigation_buttons, @@ -42,6 +42,7 @@ from app.application.messaging.interaction import ( update_or_post_message, ) from app.application.mediaserver import MediaServerHelper +from app.application.subscribe import add_subscribe, async_add_subscribe from app.adapters.external.server import MoviePilotServerHelper from app.application.torrent import TorrentHelper from app.runtime.log import logger @@ -49,12 +50,8 @@ from app.schemas import (SubscribeEpisodesRefreshEventData, SubscribeCompletionCheckEventData) from app.schemas.types import MUSIC_ENTITY_ALBUM, MUSIC_ENTITY_RECORDING, MediaSource, MediaType, SystemConfigKey, MessageChannel, NotificationType, EventType, ChainEventType, \ ContentType -from app.domain.media import ( - MUSIC_SUBSCRIBABLE_TYPES, - build_media_key, - normalize_media_source, - resolve_media_identity, -) +from app.domain.media import MUSIC_SUBSCRIBABLE_TYPES +from app.schemas.media import build_media_key, normalize_media_source, resolve_media_identity subscribe_interaction_manager = SlashInteractionManager() @@ -989,7 +986,7 @@ class SubscribeChain(ChainBase): kwargs.update(self.__get_default_kwargs(mediainfo.type, **kwargs)) # 操作数据库 - sid, err_msg = SubscribeOper().add(mediainfo=mediainfo, season=season, username=username, **kwargs) + sid, err_msg = add_subscribe(mediainfo=mediainfo, season=season, username=username, **kwargs) if not sid: logger.error(f'{mediainfo.title_year} {err_msg}') if not exist_ok and message: @@ -1193,7 +1190,7 @@ class SubscribeChain(ChainBase): kwargs.update(self.__get_default_kwargs(mediainfo.type, **kwargs)) # 操作数据库 - sid, err_msg = await SubscribeOper().async_add(mediainfo=mediainfo, season=season, username=username, **kwargs) + sid, err_msg = await async_add_subscribe(mediainfo=mediainfo, season=season, username=username, **kwargs) if not sid: logger.error(f'{mediainfo.title_year} {err_msg}') if not exist_ok and message: diff --git a/app/chain/torrents.py b/app/chain/torrents.py index 66de95fcb..e6eb9488b 100644 --- a/app/chain/torrents.py +++ b/app/chain/torrents.py @@ -12,14 +12,14 @@ from app.domain.context import TorrentInfo, Context, MediaInfo from app.domain.context import MusicInfo from app.domain.meta.metamusic import MetaMusic from app.domain.metainfo import MetaInfo -from app.db.site_oper import SiteOper -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.site import SiteOper +from app.db.oper.systemconfig import SystemConfigOper from app.application.rss import RssHelper from app.application.torrent import TorrentHelper from app.runtime.log import logger from app.schemas import Notification from app.schemas.types import SystemConfigKey, MessageChannel, NotificationType, MediaType -from app.domain.media import resolve_media_identity +from app.schemas.media import resolve_media_identity from app.domain.string import StringUtils diff --git a/app/chain/transfer.py b/app/chain/transfer.py index a3467726a..147877810 100755 --- a/app/chain/transfer.py +++ b/app/chain/transfer.py @@ -22,17 +22,18 @@ from app.runtime.events import eventmanager from app.domain.meta.metabase import MetaBase from app.domain.meta.metamusic import MetaMusic from app.domain.metainfo import MetaInfoPath -from app.db.downloadhistory_oper import DownloadHistoryOper +from app.db.oper.downloadhistory import DownloadHistoryOper from app.db.models.downloadhistory import DownloadHistory, DownloadFiles from app.db.models.transferhistory import TransferHistory -from app.db.systemconfig_oper import SystemConfigOper -from app.db.transferpending_oper import TransferPendingOper -from app.db.transferhistory_oper import TransferHistoryOper +from app.db.oper.systemconfig import SystemConfigOper +from app.db.oper.transferpending import TransferPendingOper +from app.db.oper.transferhistory import TransferHistoryOper from app.application.directory import DirectoryHelper from app.application.audio import AudioMetadataHelper from app.application.formatting import EpisodeFormatRuleHelper, FormatParser from app.runtime.progress import ProgressHelper -from app.application.history import (clear_transfer_failures, describe_history_gate, +from app.application.history import (add_transfer_fail, add_transfer_success, + clear_transfer_failures, describe_history_gate, evaluate_history_gate, is_skip_action, record_transfer_failure, resolve_history) from app.runtime.log import logger @@ -43,8 +44,6 @@ from app.schemas import ( EpisodeFormat, FileItem, TransferDirectoryConf, - TransferTask, - TransferQueue, TransferJob, TransferJobTask, TmdbEpisode, @@ -65,11 +64,9 @@ from app.schemas.types import ( MediaSource, ) from app.runtime.reload import ConfigReloadMixin -from app.domain.media import ( - normalize_media_source, - normalize_music_type, - resolve_media_identity, -) +from app.application.transfer import TransferQueue, TransferTask +from app.domain.media import normalize_music_type +from app.schemas.media import normalize_media_source, resolve_media_identity from app.foundation.singleton import Singleton from app.domain.string import StringUtils from app.adapters.system.host import SystemUtils @@ -156,7 +153,8 @@ class JobManager: return meta.name, season @staticmethod - def __get_media_id(media: MediaInfo = None, season: Optional[int] = None) -> Tuple: + def __get_media_id(media: Optional[Union[MediaInfo, MusicInfo]] = None, + season: Optional[int] = None) -> Tuple: """ 获取媒体ID;音乐额外区分实体类型,并为无远端ID的曲目构造稳定身份。 """ @@ -225,7 +223,7 @@ class JobManager: return self.__get_id(task) @staticmethod - def __get_media(task: TransferTask) -> schemas.MediaInfo: + def __get_media(task: TransferTask) -> Union[schemas.MediaInfo, schemas.MusicInfo]: """ 获取媒体信息 """ @@ -762,7 +760,7 @@ class JobManager: ) def success_tasks( - self, media: MediaInfo, season: Optional[int] = None + self, media: Union[MediaInfo, MusicInfo], season: Optional[int] = None ) -> List[TransferJobTask]: """ 获取作业中所有成功的任务 @@ -789,7 +787,7 @@ class JobManager: return [] return self._job_view[__mediaid__].tasks - def count(self, media: MediaInfo, season: Optional[int] = None) -> int: + def count(self, media: Union[MediaInfo, MusicInfo], season: Optional[int] = None) -> int: """ 获取作业中成功总数 """ @@ -805,7 +803,7 @@ class JobManager: ] ) - def size(self, media: MediaInfo, season: Optional[int] = None) -> int: + def size(self, media: Union[MediaInfo, MusicInfo], season: Optional[int] = None) -> int: """ 获取作业中所有成功文件总大小 """ @@ -858,7 +856,7 @@ class JobManager: return list(self._job_view.values()) def season_episodes( - self, media: MediaInfo, season: Optional[int] = None + self, media: Union[MediaInfo, MusicInfo], season: Optional[int] = None ) -> List[int]: """ 获取作业的季集清单 @@ -1276,7 +1274,7 @@ class TransferChain(ChainBase, ConfigReloadMixin, metaclass=Singleton): self, history: TransferHistory, src_path: Path, - ) -> Optional[MusicInfo]: + ) -> Optional[Union[MusicInfo, MediaInfo]]: """ 重新整理重试时恢复音乐信息。 @@ -1466,7 +1464,7 @@ class TransferChain(ChainBase, ConfigReloadMixin, metaclass=Singleton): ) # 新增转移失败历史记录 - history = transferhis.add_fail( + history = add_transfer_fail( fileitem=task.fileitem, mode=transferinfo.transfer_type if transferinfo else "", downloader=task.downloader, @@ -1474,6 +1472,7 @@ class TransferChain(ChainBase, ConfigReloadMixin, metaclass=Singleton): meta=task.meta, mediainfo=task.mediainfo, transferinfo=transferinfo, + transfer_history_oper=transferhis, ) # 整理失败事件 @@ -1586,7 +1585,7 @@ class TransferChain(ChainBase, ConfigReloadMixin, metaclass=Singleton): ) # 新增task转移成功历史记录 - history = transferhis.add_success( + history = add_transfer_success( fileitem=task.fileitem, mode=transferinfo.transfer_type if transferinfo else "", downloader=task.downloader, @@ -1594,6 +1593,7 @@ class TransferChain(ChainBase, ConfigReloadMixin, metaclass=Singleton): meta=task.meta, mediainfo=task.mediainfo, transferinfo=transferinfo, + transfer_history_oper=transferhis, ) # task整理完成事件 @@ -2285,7 +2285,9 @@ class TransferChain(ChainBase, ConfigReloadMixin, metaclass=Singleton): try: # 识别 transferhis = TransferHistoryOper() - mediainfo = task.mediainfo + # 显式标注联合:下面既会赋回音乐识别结果(MusicInfo),也会赋回影视识别 + # 结果(MediaInfo),不标注时会被推断成其中一种,另一种就成了假错误 + mediainfo: Optional[Union[MediaInfo, MusicInfo]] = task.mediainfo mediainfo_changed = False need_obtain_images = False if not mediainfo: @@ -2302,8 +2304,10 @@ class TransferChain(ChainBase, ConfigReloadMixin, metaclass=Singleton): and download_history.media_id and not history_year_conflict ): - # 下载记录中已存在识别信息 - mediainfo: Optional[MediaInfo] = MediaChain().recognize_media( + # 下载记录中已存在识别信息。这里不再重复标注类型:函数开头 + # 已把 mediainfo 声明为 MediaInfo | MusicInfo | None,重复 + # 声明会遮蔽它,把音乐识别结果判成类型错误 + mediainfo = MediaChain().recognize_media( mtype=task.mtype or MediaType(download_history.type), media_source=download_history.media_source, media_id=download_history.media_id, @@ -2367,24 +2371,35 @@ class TransferChain(ChainBase, ConfigReloadMixin, metaclass=Singleton): fileid=task.fileitem.fileid if task.fileitem else None, ) # 新增整理失败历史记录 - his = transferhis.add_fail( + his = add_transfer_fail( fileitem=task.fileitem, mode=task.transfer_type, meta=task.meta, downloader=task.downloader, download_hash=task.download_hash, + transfer_history_oper=transferhis, ) self.post_message( Notification( mtype=NotificationType.Manual, title=f"{task.fileitem.name} 未识别到媒体信息,无法入库!", - text=( - "原因:未识别到媒体信息\n" - "如果按钮不可用,可回复:\n" - f"```\n/redo {his.id}\n" - f"/redo {his.id} [media_source]|[media_id]|[类型]\n```\n" - "自动重试或手动识别整理。" - ), + # 历史落库失败时 his 为 None(add_transfer_fail 末尾的 + # get_by_src 查不到即返回 None),此时 /redo 无 ID 可用, + # 只省去这段指引而不是让整条通知连同后续的作业清理、 + # 种子完成标记一起崩在 NoneType 上 + text="\n".join( + [ + "原因:未识别到媒体信息", + ( + "如果按钮不可用,可回复:\n" + f"```\n/redo {his.id}\n" + f"/redo {his.id} [media_source]|[media_id]|[类型]\n```\n" + "自动重试或手动识别整理。" + if his + else "" + ), + ] + ).strip(), username=task.username, link=settings.MP_DOMAIN("#/history"), buttons=self.build_failed_transfer_buttons( @@ -3166,7 +3181,11 @@ class TransferChain(ChainBase, ConfigReloadMixin, metaclass=Singleton): @staticmethod def _is_movie_year_conflict( - file_meta: MetaBase, media: Union[DownloadHistory, MediaInfo] + file_meta: MetaBase, + # 两种 DownloadHistory 都会进来:库模型(本文件按 ORM 行查历史)与 + # schemas DTO(TransferTask.download_history)。本函数只按 getattr 取 + # year 与 type,对两者一视同仁 + media: Union[DownloadHistory, schemas.DownloadHistory, MediaInfo, MusicInfo] ) -> bool: """ 判断文件名年份是否与已识别电影年份冲突。 @@ -3450,7 +3469,7 @@ class TransferChain(ChainBase, ConfigReloadMixin, metaclass=Singleton): self, fileitem: FileItem, meta: MetaBase = None, - mediainfo: MediaInfo = None, + mediainfo: Optional[Union[MediaInfo, MusicInfo]] = None, mtype: Optional[MediaType] = None, media_source: Optional[MediaSource] = None, media_id: Optional[str] = None, @@ -4686,7 +4705,7 @@ class TransferChain(ChainBase, ConfigReloadMixin, metaclass=Singleton): def send_transfer_message( self, meta: MetaBase, - mediainfo: MediaInfo, + mediainfo: Union[MediaInfo, MusicInfo], transferinfo: TransferInfo, season_episode: Optional[str] = None, episodes_info: Optional[List[TmdbEpisode]] = None, diff --git a/app/chain/user.py b/app/chain/user.py index 41d2b8531..5856f7b12 100644 --- a/app/chain/user.py +++ b/app/chain/user.py @@ -6,7 +6,7 @@ from app.chain import ChainBase from app.runtime.config import settings from app.application.security.access import get_password_hash, verify_password from app.db.models.user import User -from app.db.user_oper import UserOper +from app.db.oper.user import UserOper from app.runtime.log import logger from app.schemas import AuthCredentials, AuthInterceptCredentials from app.schemas.types import ChainEventType diff --git a/app/chain/workflow.py b/app/chain/workflow.py index 0584c0c80..717df3f80 100644 --- a/app/chain/workflow.py +++ b/app/chain/workflow.py @@ -16,7 +16,7 @@ from app.chain import ChainBase from app.runtime.config import global_vars from app.runtime.events import Event, eventmanager from app.db.models import Workflow -from app.db.workflow_oper import WorkflowOper +from app.db.oper.workflow import WorkflowOper from app.runtime.log import logger from app.schemas import ActionContext, ActionFlow, Action, ActionExecution, ActionResult from app.schemas.types import EventType diff --git a/app/db/__init__.py b/app/db/__init__.py index 6d0a4c086..29668df49 100644 --- a/app/db/__init__.py +++ b/app/db/__init__.py @@ -1,564 +1,112 @@ -import asyncio -from typing import Any, Generator, List, Optional, Self, Tuple, AsyncGenerator, Union +""" +数据库包入口。 -from sqlalchemy import NullPool, QueuePool, and_, create_engine, event, inspect, text, select, delete, Column, Integer, \ - Sequence, Identity -from sqlalchemy.engine import Engine as SQLAlchemyEngine, ExceptionContext -from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker -from sqlalchemy.orm import Session, as_declarative, declared_attr, scoped_session, sessionmaker +本模块只做符号再导出,不承载实现——具体职责分布在: -from app.runtime.config import settings -from app.runtime.log import logger +- diagnostics 驱动错误的统一分类与日志 +- engine 引擎构建、连接额度核算 +- session 会话获取、异步连接池与配额 +- decorators 同步/异步事务装饰器 +- base ORM 基类与数据访问基类 +- models 表结构声明,一实体一文件 +- oper 数据访问实现,与 models 同名文件一一对应 + +历史上这些代码全部堆在本文件里(782 行),既让包入口承担了实现职责、 +使依赖图难以理清,也让「import 即建立数据库连接」这一副作用被固化下来。 +""" +from typing import TYPE_CHECKING, Any + +from app.db.base import Base, DbOper, execute_dml, get_id_column +from app.db.decorators import async_db_query, async_db_update, db_query, db_update +from app.db.engine import ( + check_connection_budget, + connection_budget, + get_engine, + get_global_async_engine, +) +from app.db.session import ( + AsyncSessionFactory, + ScopedSession, + SessionFactory, + async_session_scope, + close_database, + get_async_db, + get_async_engine, + get_async_session_factory, + get_db, + get_scoped_session, + get_session_factory, +) + +# ==================== 对外契约的分层 ==================== +# 下方 __all__ 是本包**对外承诺**的那一层,仓库外的插件只应依赖其中的名字: +# +# - 数据访问:继承 DbOper 子类(插件基类已备好 self.plugindata / self.systemconfig), +# 或给自己的函数套 db_query / db_update / async_db_query / async_db_update 装饰器。 +# 会话的获取、提交、回滚、释放全部由装饰器收口。 +# - 引擎:Engine / AsyncEngine 保留在契约内。建表、Alembic 迁移、连接诊断这些用途 +# 确实需要引擎对象本身,装饰器覆盖不到,仓库外拿它是正当的。 +# +# SessionFactory / AsyncSessionFactory / ScopedSession 三个名字**不在**契约内,已从 +# __all__ 移除,降级为内部实现细节。它们建出来的是绕过上述装饰器的裸会话——没有提交、 +# 没有回滚、没有释放,谁建谁自己兜底,本身就是误用的形状。仓库内确有几处直接 +# `from app.db import SessionFactory`(scheduler、postgresql 模块、Alembic 迁移脚本), +# 那是包内部的既有用法,直接导入不受 __all__ 约束,照常可用。 +# 若确实需要真正的工厂对象(而非 `X()` 取一个会话),用 get_session_factory() / +# get_scoped_session() / get_async_session_factory()——转发函数上没有 sessionmaker +# 与 scoped_session 的实例接口(.remove() / .configure() / .begin() 等)。 +# +# 实现上,三个工厂名字本身就是转发函数(见 session 模块),直接再导出即可——导入它们 +# 不会碰引擎。Engine / AsyncEngine 则不同:调用方拿到的必须是引擎**对象**而非函数, +# 所以只能靠模块级 __getattr__ 在取属性时才创建。 +# +# 注意这意味着 `from app.db import Engine` 仍会在 import 期把引擎建出来——那是调用方 +# 自己选的时机。本包自身及仓库内代码一律用 get_engine(),所以 `import app.db` 不连库。 +if TYPE_CHECKING: + # 只为静态检查声明这两个名字:运行期由下方 __getattr__ 解析,模块 __dict__ 里并不存在, + # 类型检查器无从知道它们属于本模块(__all__ 里的它们会被报成 reportUnsupportedDunderAll)。 + # 这里同时把类型钉准,比 __getattr__ 的 Any 更有用:调用方拿到的确实是这两类引擎。 + from sqlalchemy.engine import Engine as _SyncEngine + from sqlalchemy.ext.asyncio import AsyncEngine as _SaAsyncEngine + + Engine: _SyncEngine + AsyncEngine: _SaAsyncEngine -def _database_error_metadata(error: BaseException) -> Optional[dict[str, Any]]: - """提取 SQLite 与 PostgreSQL 驱动提供的稳定错误分类字段。""" - metadata = {"error_type": type(error).__name__} - - # DBAPI 驱动字段并不共享统一类型,动态读取可同时兼容 sqlite3、psycopg2 与 asyncpg。 - sqlite_errorcode = getattr(error, "sqlite_errorcode", None) - sqlite_errorname = getattr(error, "sqlite_errorname", None) - if sqlite_errorcode is not None or sqlite_errorname: - if sqlite_errorcode is not None: - metadata["error_code"] = sqlite_errorcode - if sqlite_errorname: - metadata["error_name"] = sqlite_errorname - return metadata - - sqlstate = getattr(error, "sqlstate", None) or getattr(error, "pgcode", None) - if not sqlstate: - sqlstate = getattr(getattr(error, "diag", None), "sqlstate", None) - if sqlstate: - metadata["sqlstate"] = sqlstate - return metadata - - return None - - -def _log_database_error(exception_context: ExceptionContext) -> None: - """记录非敏感驱动错误码,并保持 SQLAlchemy 原有异常传播。""" - metadata = _database_error_metadata(exception_context.original_exception) - if not metadata: - return - - dialect = exception_context.dialect - fields = { - "database": dialect.name, - "driver": dialect.driver, - **metadata, - } - logger.error( - "数据库驱动异常:" + ", ".join(f"{key}={value}" for key, value in fields.items()) - ) - - -def _register_database_error_logging(engine: SQLAlchemyEngine) -> None: - """为主程序 Engine 注册统一的底层驱动错误诊断。""" - event.listen(engine, "handle_error", _log_database_error) - - -def get_id_column(): +def __getattr__(name: str) -> Any: """ - 根据数据库类型返回合适的ID列定义 + 惰性解析 Engine / AsyncEngine 两个旧名字,保持仓库外插件的导入路径可用。 + :param name: 属性名 + :return: 对应的引擎 """ - if settings.DB_TYPE.lower() == "postgresql": - # PostgreSQL使用SERIAL类型,让数据库自动处理序列 - return Column(Integer, Identity(start=1, cycle=True), primary_key=True) - else: - # SQLite使用Sequence - return Column(Integer, Sequence('id'), primary_key=True) - - -def _get_database_engine(is_async: bool = False): - """ - 获取数据库连接参数并设置WAL模式 - :param is_async: 是否创建异步引擎,True - 异步引擎, False - 同步引擎 - :return: 返回对应的数据库引擎 - """ - # 根据数据库类型选择连接方式 - if settings.DB_TYPE.lower() == "postgresql": - return _get_postgresql_engine(is_async) - else: - return _get_sqlite_engine(is_async) - - -def _get_sqlite_engine(is_async: bool = False): - """ - 获取SQLite数据库引擎 - """ - # 连接参数 - _connect_args = { - "timeout": settings.DB_TIMEOUT, - } - # 启用 WAL 模式时的额外配置 - if settings.DB_WAL_ENABLE: - _connect_args["check_same_thread"] = False - - # 创建同步引擎 - if not is_async: - # 根据池类型设置 poolclass 和相关参数 - _pool_class = NullPool if settings.DB_POOL_TYPE == "NullPool" else QueuePool - - # 数据库参数 - _db_kwargs = { - "url": f"sqlite:///{settings.CONFIG_PATH}/user.db", - "pool_pre_ping": settings.DB_POOL_PRE_PING, - "echo": settings.DB_ECHO, - "poolclass": _pool_class, - "pool_recycle": settings.DB_POOL_RECYCLE, - "connect_args": _connect_args - } - - # 当使用 QueuePool 时,添加 QueuePool 特有的参数 - if _pool_class == QueuePool: - _db_kwargs.update({ - "pool_size": settings.DB_SQLITE_POOL_SIZE, - "pool_timeout": settings.DB_POOL_TIMEOUT, - "max_overflow": settings.DB_SQLITE_MAX_OVERFLOW - }) - - # 创建数据库引擎 - engine = create_engine(**_db_kwargs) - _register_database_error_logging(engine) - - # 设置WAL模式 - _journal_mode = "WAL" if settings.DB_WAL_ENABLE else "DELETE" - with engine.connect() as connection: - current_mode = connection.execute(text(f"PRAGMA journal_mode={_journal_mode};")).scalar() - print(f"SQLite database journal mode set to: {current_mode}") - - return engine - else: - # 数据库参数,只能使用 NullPool - _db_kwargs = { - "url": f"sqlite+aiosqlite:///{settings.CONFIG_PATH}/user.db", - "pool_pre_ping": settings.DB_POOL_PRE_PING, - "echo": settings.DB_ECHO, - "poolclass": NullPool, - "pool_recycle": settings.DB_POOL_RECYCLE, - "connect_args": _connect_args - } - # 创建异步数据库引擎 - async_engine = create_async_engine(**_db_kwargs) - _register_database_error_logging(async_engine.sync_engine) - - # 设置WAL模式 - _journal_mode = "WAL" if settings.DB_WAL_ENABLE else "DELETE" - - async def set_async_wal_mode(): - """ - 设置异步引擎的WAL模式 - """ - async with async_engine.connect() as _connection: - result = await _connection.execute(text(f"PRAGMA journal_mode={_journal_mode};")) - _current_mode = result.scalar() - print(f"Async SQLite database journal mode set to: {_current_mode}") - - try: - asyncio.run(set_async_wal_mode()) - except Exception as e: - print(f"Failed to set async SQLite WAL mode: {e}") - - return async_engine - - -def _get_postgresql_engine(is_async: bool = False): - """ - 获取PostgreSQL数据库引擎 - """ - db_url = settings.DB_POSTGRESQL_URL() - - # PostgreSQL连接参数 - _connect_args = {} - - # 创建同步引擎 - if not is_async: - # 根据池类型设置 poolclass 和相关参数 - _pool_class = NullPool if settings.DB_POOL_TYPE == "NullPool" else QueuePool - - # 数据库参数 - _db_kwargs = { - "url": db_url, - "pool_pre_ping": settings.DB_POOL_PRE_PING, - "echo": settings.DB_ECHO, - "poolclass": _pool_class, - "pool_recycle": settings.DB_POOL_RECYCLE, - "connect_args": _connect_args - } - - # 当使用 QueuePool 时,添加 QueuePool 特有的参数 - if _pool_class == QueuePool: - _db_kwargs.update({ - "pool_size": settings.DB_POSTGRESQL_POOL_SIZE, - "pool_timeout": settings.DB_POOL_TIMEOUT, - "max_overflow": settings.DB_POSTGRESQL_MAX_OVERFLOW - }) - - # 创建数据库引擎 - engine = create_engine(**_db_kwargs) - _register_database_error_logging(engine) - print(f"PostgreSQL database connected to {settings.DB_POSTGRESQL_TARGET}/{settings.DB_POSTGRESQL_DATABASE}") - - return engine - else: - async_db_url = settings.DB_POSTGRESQL_URL("asyncpg") - - # 数据库参数,只能使用 NullPool - _db_kwargs = { - "url": async_db_url, - "pool_pre_ping": settings.DB_POOL_PRE_PING, - "echo": settings.DB_ECHO, - "poolclass": NullPool, - "pool_recycle": settings.DB_POOL_RECYCLE, - "connect_args": _connect_args - } - # 创建异步数据库引擎 - async_engine = create_async_engine(**_db_kwargs) - _register_database_error_logging(async_engine.sync_engine) - print(f"Async PostgreSQL database connected to {settings.DB_POSTGRESQL_TARGET}/{settings.DB_POSTGRESQL_DATABASE}") - - return async_engine - - -# 同步数据库引擎 -Engine = _get_database_engine(is_async=False) - -# 异步数据库引擎 -AsyncEngine = _get_database_engine(is_async=True) - -# 同步会话工厂 -SessionFactory = sessionmaker(bind=Engine) - -# 异步会话工厂 -AsyncSessionFactory = async_sessionmaker(bind=AsyncEngine, class_=AsyncSession) - -# 同步多线程全局使用的数据库会话 -ScopedSession = scoped_session(SessionFactory) - - -def get_db() -> Generator: - """ - 获取数据库会话,用于WEB请求 - :return: Session - """ - db = None - try: - db = SessionFactory() - yield db - finally: - if db: - db.close() - - -async def get_async_db() -> AsyncGenerator[AsyncSession, None]: - """ - 获取异步数据库会话,用于WEB请求 - :return: AsyncSession - """ - async with AsyncSessionFactory() as session: - try: - yield session - finally: - await session.close() - - -async def close_database(): - """ - 关闭所有数据库连接并清理资源 - """ - try: - # 释放同步连接池 - Engine.dispose() # noqa - # 释放异步连接池 - await AsyncEngine.dispose() - except Exception as err: - print(f"Error while disposing database connections: {err}") - - -def _get_args_db(args: tuple, kwargs: dict) -> Optional[Session]: - """ - 从参数中获取数据库Session对象 - """ - db = None - if args: - for arg in args: - if isinstance(arg, Session): - db = arg - break - if kwargs: - for key, value in kwargs.items(): - if isinstance(value, Session): - db = value - break - return db - - -def _get_args_async_db(args: tuple, kwargs: dict) -> Optional[AsyncSession]: - """ - 从参数中获取异步数据库AsyncSession对象 - """ - db = None - if args: - for arg in args: - if isinstance(arg, AsyncSession): - db = arg - break - if kwargs: - for key, value in kwargs.items(): - if isinstance(value, AsyncSession): - db = value - break - return db - - -def _update_args_db(args: tuple, kwargs: dict, db: Session) -> Tuple[tuple, dict]: - """ - 更新参数中的数据库Session对象,关键字传参时更新db的值,否则更新第1或第2个参数 - """ - if kwargs and 'db' in kwargs: - kwargs['db'] = db - elif args: - if args[0] is None: - args = (db, *args[1:]) - else: - args = (args[0], db, *args[2:]) - return args, kwargs - - -def _update_args_async_db(args: tuple, kwargs: dict, db: AsyncSession) -> Tuple[tuple, dict]: - """ - 更新参数中的异步数据库AsyncSession对象,关键字传参时更新db的值,否则更新第1或第2个参数 - """ - if kwargs and 'db' in kwargs: - kwargs['db'] = db - elif args: - if args[0] is None: - args = (db, *args[1:]) - else: - args = (args[0], db, *args[2:]) - return args, kwargs - - -def db_update(func): - """ - 数据库更新类操作装饰器,第一个参数必须是数据库会话或存在db参数 - """ - - def wrapper(*args, **kwargs): - # 是否关闭数据库会话 - _close_db = False - # 从参数中获取数据库会话 - db = _get_args_db(args, kwargs) - if not db: - # 如果没有获取到数据库会话,创建一个 - db = ScopedSession() - # 标记需要关闭数据库会话 - _close_db = True - # 更新参数中的数据库会话 - args, kwargs = _update_args_db(args, kwargs, db) - try: - # 执行函数 - result = func(*args, **kwargs) - # 提交事务 - db.commit() - except Exception as err: - # 回滚事务 - db.rollback() - raise err - finally: - # 关闭数据库会话 - if _close_db: - db.close() - return result - - return wrapper - - -def async_db_update(func): - """ - 异步数据库更新类操作装饰器,第一个参数必须是异步数据库会话或存在db参数 - """ - - async def wrapper(*args, **kwargs): - # 是否关闭数据库会话 - _close_db = False - # 从参数中获取异步数据库会话 - db = _get_args_async_db(args, kwargs) - if not db: - # 如果没有获取到异步数据库会话,创建一个 - db = AsyncSessionFactory() - # 标记需要关闭数据库会话 - _close_db = True - # 更新参数中的异步数据库会话 - args, kwargs = _update_args_async_db(args, kwargs, db) - try: - # 执行函数 - result = await func(*args, **kwargs) - # 提交事务 - await db.commit() - except Exception as err: - # 回滚事务 - await db.rollback() - raise err - finally: - # 关闭数据库会话 - if _close_db: - await db.close() - return result - - return wrapper - - -def db_query(func): - """ - 数据库查询操作装饰器,第一个参数必须是数据库会话或存在db参数 - 注意:db.query列表数据时,需要转换为list返回 - """ - - def wrapper(*args, **kwargs): - # 是否关闭数据库会话 - _close_db = False - # 从参数中获取数据库会话 - db = _get_args_db(args, kwargs) - if not db: - # 如果没有获取到数据库会话,创建一个 - db = ScopedSession() - # 标记需要关闭数据库会话 - _close_db = True - # 更新参数中的数据库会话 - args, kwargs = _update_args_db(args, kwargs, db) - try: - # 执行函数 - result = func(*args, **kwargs) - except Exception as err: - raise err - finally: - # 关闭数据库会话 - if _close_db: - db.close() - return result - - return wrapper - - -def async_db_query(func): - """ - 异步数据库查询操作装饰器,第一个参数必须是异步数据库会话或存在db参数 - 注意:db.query列表数据时,需要转换为list返回 - """ - - async def wrapper(*args, **kwargs): - # 是否关闭数据库会话 - _close_db = False - # 从参数中获取异步数据库会话 - db = _get_args_async_db(args, kwargs) - if not db: - # 如果没有获取到异步数据库会话,创建一个 - db = AsyncSessionFactory() - # 标记需要关闭数据库会话 - _close_db = True - # 更新参数中的异步数据库会话 - args, kwargs = _update_args_async_db(args, kwargs, db) - try: - # 执行函数 - result = await func(*args, **kwargs) - except Exception as err: - raise err - finally: - # 关闭数据库会话 - if _close_db: - await db.close() - return result - - return wrapper - - -@as_declarative() -class Base: - id: Any - __name__: str - - @db_update - def create(self, db: Session): - db.add(self) - - @async_db_update - async def async_create(self, db: AsyncSession): - db.add(self) - await db.flush() - return self - - @classmethod - @db_query - def get(cls, db: Session, rid: int) -> Self: - return db.query(cls).filter(and_(cls.id == rid)).first() - - @classmethod - @async_db_query - async def async_get(cls, db: AsyncSession, rid: int) -> Self: - result = await db.execute(select(cls).where(and_(cls.id == rid))) - return result.scalars().first() - - @db_update - def update(self, db: Session, payload: dict): - for key, value in payload.items(): - setattr(self, key, value) - if inspect(self).detached: - db.add(self) - - @async_db_update - async def async_update(self, db: AsyncSession, payload: dict): - for key, value in payload.items(): - setattr(self, key, value) - if inspect(self).detached: - db.add(self) - - @classmethod - @db_update - def delete(cls, db: Session, rid): - db.query(cls).filter(and_(cls.id == rid)).delete() - - @classmethod - @async_db_update - async def async_delete(cls, db: AsyncSession, rid): - result = await db.execute(select(cls).where(and_(cls.id == rid))) - user = result.scalars().first() - if user: - await db.delete(user) - - @classmethod - @db_update - def truncate(cls, db: Session): - db.query(cls).delete() - - @classmethod - @async_db_update - async def async_truncate(cls, db: AsyncSession): - await db.execute(delete(cls)) - - @classmethod - @db_query - def list(cls, db: Session) -> List[Self]: - return db.query(cls).all() - - @classmethod - @async_db_query - async def async_list(cls, db: AsyncSession) -> Sequence[Self]: - result = await db.execute(select(cls)) - return result.scalars().all() - - def to_dict(self): - return {c.name: getattr(self, c.name, None) for c in self.__table__.columns} # noqa - - @declared_attr - def __tablename__(self) -> str: - return self.__name__.lower() - - -class DbOper: - """ - 数据库操作基类 - """ - - def __init__(self, db: Union[Session, AsyncSession] = None): - self._db = db + if name == "Engine": + return get_engine() + if name == "AsyncEngine": + return get_global_async_engine() + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + +__all__ = [ + "AsyncEngine", + "Base", + "DbOper", + "Engine", + "async_db_query", + "async_db_update", + "async_session_scope", + "check_connection_budget", + "close_database", + "connection_budget", + "db_query", + "db_update", + "execute_dml", + "get_async_db", + "get_async_engine", + "get_async_session_factory", + "get_db", + "get_engine", + "get_global_async_engine", + "get_id_column", + "get_scoped_session", + "get_session_factory", +] diff --git a/app/db/base.py b/app/db/base.py new file mode 100644 index 000000000..1dd79ec3d --- /dev/null +++ b/app/db/base.py @@ -0,0 +1,150 @@ +""" +ORM 基类与数据访问基类。 + +Base 提供声明式基类与通用的行为(字典转换、增删改查便利方法); +DbOper 是各业务 Oper 的基类,持有一个可注入的会话。 +""" +from typing import Any, List, Optional, Self, Union, cast + +from sqlalchemy import (CursorResult, Executable, Identity, Integer, Sequence, + and_, delete, inspect, select) +from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import DeclarativeBase, Mapped, Session, declared_attr, mapped_column + +from app.runtime.config import settings +from app.db.decorators import async_db_query, async_db_update, db_query, db_update + + +def execute_dml(db: Session, statement: Executable, + execution_options: Optional[dict] = None) -> int: + """ + 执行 DML 语句并返回影响行数。 + + ``Session.execute`` 的类型标注一律是 ``Result``,只有运行期真正拿到的 + ``CursorResult`` 才带 ``rowcount``——2.0 只为 ``Connection.execute`` 加了 + ``CursorResult`` 重载。这里把转换收口一次,免得每个模型各写一遍 cast。 + :param db: 数据库会话 + :param statement: delete()/update() 等 DML 语句 + :param execution_options: 执行选项;不传即沿用 SQLAlchemy 默认的会话同步策略 + :return: 影响行数 + """ + if execution_options is None: + result = db.execute(statement) + else: + result = db.execute(statement, execution_options=execution_options) + return cast(CursorResult[Any], result).rowcount + + +def get_id_column() -> Mapped[int]: + """ + 根据数据库类型返回合适的ID列定义 + """ + if settings.DB_TYPE.lower() == "postgresql": + # PostgreSQL使用SERIAL类型,让数据库自动处理序列 + return mapped_column(Integer, Identity(start=1, cycle=True), primary_key=True) + else: + # SQLite使用Sequence + return mapped_column(Integer, Sequence('id'), primary_key=True) + + +class Base(DeclarativeBase): + """ + 声明式基类。 + + 2.0 的声明式系统会解释类级 PEP 484 注解,未包裹在 Mapped[] 中的注解会直接报错。 + 仓内模型已全部迁移到 mapped_column() + Mapped[] 注解,因此不设 __allow_unmapped__: + 该标志此前只为「仓外插件可能继承本 Base 自定义 legacy 注解模型」保留,插件生态 + 确定迭代后这条理由不再成立。留着它反而会让回流的 1.x 写法在 import 期悄悄通过, + 等到运行期才以「列不存在」的形式暴露。 + + 继承本类的模型一律使用 mapped_column() + Mapped[] 注解;确需非映射的类级属性时 + 用 ClassVar 显式声明,而不是把这个标志加回来。 + """ + + # 由 get_id_column() 在各模型中提供实际的列定义,这里只声明类型供 IDE 使用 + id: Mapped[int] + + @db_update + def create(self, db: Session): + db.add(self) + + @async_db_update + async def async_create(self, db: AsyncSession): + db.add(self) + await db.flush() + return self + + @classmethod + @db_query + def get(cls, db: Session, rid: int) -> Optional[Self]: + return db.execute(select(cls).where(and_(cls.id == rid))).scalars().first() + + @classmethod + @async_db_query + async def async_get(cls, db: AsyncSession, rid: int) -> Optional[Self]: + result = await db.execute(select(cls).where(and_(cls.id == rid))) + return result.scalars().first() + + @db_update + def update(self, db: Session, payload: dict): + for key, value in payload.items(): + setattr(self, key, value) + if inspect(self).detached: + db.add(self) + + @async_db_update + async def async_update(self, db: AsyncSession, payload: dict): + for key, value in payload.items(): + setattr(self, key, value) + if inspect(self).detached: + db.add(self) + + @classmethod + @db_update + def delete(cls, db: Session, rid): + db.execute(delete(cls).where(and_(cls.id == rid))) + + @classmethod + @async_db_update + async def async_delete(cls, db: AsyncSession, rid): + result = await db.execute(select(cls).where(and_(cls.id == rid))) + user = result.scalars().first() + if user: + await db.delete(user) + + @classmethod + @db_update + def truncate(cls, db: Session): + db.execute(delete(cls)) + + @classmethod + @async_db_update + async def async_truncate(cls, db: AsyncSession): + await db.execute(delete(cls)) + + @classmethod + @db_query + def list(cls, db: Session) -> List[Self]: + return list(db.execute(select(cls)).scalars().all()) + + @classmethod + @async_db_query + async def async_list(cls, db: AsyncSession) -> List[Self]: + result = await db.execute(select(cls)) + return list(result.scalars().all()) + + def to_dict(self): + return {c.name: getattr(self, c.name, None) for c in self.__table__.columns} # noqa + + @declared_attr.directive + def __tablename__(cls) -> str: # noqa: N805 declared_attr 的第一个参数即类本身 + return cls.__name__.lower() + + +class DbOper: + """ + 数据库操作基类 + """ + + def __init__(self, db: Optional[Union[Session, AsyncSession]] = None): + self._db = db diff --git a/app/db/decorators.py b/app/db/decorators.py new file mode 100644 index 000000000..f4acb5ceb --- /dev/null +++ b/app/db/decorators.py @@ -0,0 +1,264 @@ +""" +数据库事务装饰器。 + +同步/异步各一对:查询装饰器负责会话的获取与释放,更新装饰器额外负责提交与回滚。 +未显式传入会话时自动创建,并在结束时归还——异步路径经 async_session_scope 收口, +连接池与配额都在那里生效。 + +收尾故障(rollback / close / __aexit__ 自身抛异常)一律只记日志、不上抛,四个装饰器 +的处理一致。理由与代价都要写明,别当成漏写的 raise: + +- 连接断开、事务已失效这类故障恰恰最容易发生在「出错之后」的收尾阶段。裸写收尾语句时 + 它一抛错就顶替掉原始异常,调用方看到的只剩「connection reset」,业务异常连类型都被 + 换掉,按类型分流的 except(唯一约束冲突要重试、参数错误要报错)一并失配。 +- 代价是成功路径的行为随之改变:func() 成功、close() 失败时,调用方**静默拿到返回值**, + 故障只进日志。这是有意为之——close() 失败时事务已经提交、业务确实成功了,且 + SQLAlchemy 归还连接时已在池层吞掉异常并 invalidate 坏连接,再把释放故障升级成调用方 + 的异常,只会让一次已经落库的写入看起来像失败,诱发重复提交。 +""" +from typing import Any, Awaitable, Callable, Optional, Tuple, TypeVar + +from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import Session + +from app.db.session import ScopedSession, async_session_scope +from app.runtime.log import logger + +_R = TypeVar("_R") + +# 四个装饰器都会重写实参列表:未传会话时自行创建一个并塞回 db 位置。因此包装后的可调用 +# 对象接受的实参与被包装函数的签名并不一致——用 Callable[..., _R] 如实表达「参数由装饰器 +# 接管、返回值原样透传」。否则调用方传 None 或传异步会话都会被判成类型不符,而这恰恰是 +# 装饰器存在的理由(各 Oper 的 self._db 常态就是 None)。 + +def _get_args_db(args: tuple, kwargs: dict) -> Optional[Session]: + """ + 从参数中获取数据库Session对象 + """ + db = None + if args: + for arg in args: + if isinstance(arg, Session): + db = arg + break + if kwargs: + for key, value in kwargs.items(): + if isinstance(value, Session): + db = value + break + return db + + +def _get_args_async_db(args: tuple, kwargs: dict) -> Optional[AsyncSession]: + """ + 从参数中获取异步数据库AsyncSession对象 + """ + db = None + if args: + for arg in args: + if isinstance(arg, AsyncSession): + db = arg + break + if kwargs: + for key, value in kwargs.items(): + if isinstance(value, AsyncSession): + db = value + break + return db + + +def _update_args_db(args: tuple, kwargs: dict, db: Session) -> Tuple[tuple, dict]: + """ + 更新参数中的数据库Session对象,关键字传参时更新db的值,否则更新第1或第2个参数 + """ + if kwargs and 'db' in kwargs: + kwargs['db'] = db + elif args: + if args[0] is None: + args = (db, *args[1:]) + else: + args = (args[0], db, *args[2:]) + return args, kwargs + + +def _update_args_async_db(args: tuple, kwargs: dict, db: AsyncSession) -> Tuple[tuple, dict]: + """ + 更新参数中的异步数据库AsyncSession对象,关键字传参时更新db的值,否则更新第1或第2个参数 + """ + if kwargs and 'db' in kwargs: + kwargs['db'] = db + elif args: + if args[0] is None: + args = (db, *args[1:]) + else: + args = (args[0], db, *args[2:]) + return args, kwargs + + +def db_update(func: Callable[..., _R]) -> Callable[..., _R]: + """ + 数据库更新类操作装饰器,第一个参数必须是数据库会话或存在db参数 + """ + + def wrapper(*args: Any, **kwargs: Any) -> _R: + # 是否关闭数据库会话 + _close_db = False + # 从参数中获取数据库会话 + db = _get_args_db(args, kwargs) + if not db: + # 如果没有获取到数据库会话,创建一个 + db = ScopedSession() + # 标记需要关闭数据库会话 + _close_db = True + # 更新参数中的数据库会话 + args, kwargs = _update_args_db(args, kwargs, db) + try: + # 执行函数 + result = func(*args, **kwargs) + # 提交事务 + db.commit() + except Exception as err: + # 回滚事务。回滚自身失败不得顶替原始异常:连接断开、事务已失效这类收尾故障 + # 恰恰最容易发生在「出错之后」,裸写 db.rollback() 时它一抛错,调用方看到的 + # 就只剩「connection reset」,真正的业务异常连类型都被换掉、按类型分流的 + # except 一并失配。故障本身另行记录,不静默吞掉 + try: + db.rollback() + except Exception as rollback_err: # noqa: BLE001 回滚失败不能掩盖原始异常 + logger.error(f"事务回滚失败,原始异常将原样上抛:{rollback_err}") + raise err + finally: + # 关闭数据库会话。释放失败只记录:既不顶替上面正在传播的业务异常, + # 成功路径下也不把一次已提交的写入变成调用方眼里的失败(见模块说明) + if _close_db: + try: + db.close() + except Exception as close_err: # noqa: BLE001 释放故障不得改变调用结果 + logger.error(f"释放数据库会话失败:{close_err}") + return result + + return wrapper + + +def async_db_update(func: Callable[..., Awaitable[_R]]) -> Callable[..., Awaitable[_R]]: + """ + 异步数据库更新类操作装饰器,第一个参数必须是异步数据库会话或存在db参数 + """ + + async def wrapper(*args: Any, **kwargs: Any) -> _R: + # 是否关闭数据库会话;作用域与 _scope 同生共死,先置空以便静态检查看清 + _close_db = False + _scope = None + # 从参数中获取异步数据库会话 + db = _get_args_async_db(args, kwargs) + if not db: + # 如果没有获取到异步数据库会话,创建一个。经 async_session_scope + # 统一收口:常驻主循环走连接池,其余循环走 NullPool 并占用全局配额 + _scope = async_session_scope() + db = await _scope.__aenter__() + # 标记需要关闭数据库会话 + _close_db = True + # 更新参数中的异步数据库会话 + args, kwargs = _update_args_async_db(args, kwargs, db) + try: + # 执行函数 + result = await func(*args, **kwargs) + # 提交事务 + await db.commit() + except Exception as err: + # 回滚事务;与同步路径同理,回滚失败只记录,不顶替原始异常 + try: + await db.rollback() + except Exception as rollback_err: # noqa: BLE001 回滚失败不能掩盖原始异常 + logger.error(f"事务回滚失败,原始异常将原样上抛:{rollback_err}") + raise err + finally: + # 关闭数据库会话 + if _close_db and _scope is not None: + # 退出会话上下文而不是只 close:配额的释放绑定在 __aexit__ 上, + # 只关会话会让回退路径的全局配额永不归还,最终把自己饿死。 + # 退出失败同样只记录,不改变调用结果(见模块说明) + try: + await _scope.__aexit__(None, None, None) + except Exception as close_err: # noqa: BLE001 释放故障不得改变调用结果 + logger.error(f"释放数据库会话失败:{close_err}") + return result + + return wrapper + + +def db_query(func: Callable[..., _R]) -> Callable[..., _R]: + """ + 数据库查询操作装饰器,第一个参数必须是数据库会话或存在db参数 + 注意:db.query列表数据时,需要转换为list返回 + """ + + def wrapper(*args: Any, **kwargs: Any) -> _R: + # 是否关闭数据库会话 + _close_db = False + # 从参数中获取数据库会话 + db = _get_args_db(args, kwargs) + if not db: + # 如果没有获取到数据库会话,创建一个 + db = ScopedSession() + # 标记需要关闭数据库会话 + _close_db = True + # 更新参数中的数据库会话 + args, kwargs = _update_args_db(args, kwargs, db) + try: + # 执行函数 + result = func(*args, **kwargs) + except Exception as err: + raise err + finally: + # 关闭数据库会话。释放失败只记录,不顶替业务异常、也不影响成功路径的返回值 + # (见模块说明) + if _close_db: + try: + db.close() + except Exception as close_err: # noqa: BLE001 释放故障不得改变调用结果 + logger.error(f"释放数据库会话失败:{close_err}") + return result + + return wrapper + + +def async_db_query(func: Callable[..., Awaitable[_R]]) -> Callable[..., Awaitable[_R]]: + """ + 异步数据库查询操作装饰器,第一个参数必须是异步数据库会话或存在db参数 + 注意:db.query列表数据时,需要转换为list返回 + """ + + async def wrapper(*args: Any, **kwargs: Any) -> _R: + # 是否关闭数据库会话 + _close_db = False + _scope = None + # 从参数中获取异步数据库会话 + db = _get_args_async_db(args, kwargs) + if not db: + # 如果没有获取到异步数据库会话,创建一个。经 async_session_scope + # 统一收口:常驻主循环走连接池,其余循环走 NullPool 并占用全局配额 + _scope = async_session_scope() + db = await _scope.__aenter__() + # 标记需要关闭数据库会话 + _close_db = True + # 更新参数中的异步数据库会话 + args, kwargs = _update_args_async_db(args, kwargs, db) + try: + # 执行函数 + result = await func(*args, **kwargs) + except Exception as err: + raise err + finally: + # 关闭数据库会话 + if _close_db and _scope is not None: + # 退出会话上下文而不是只 close:配额的释放绑定在 __aexit__ 上, + # 只关会话会让回退路径的全局配额永不归还,最终把自己饿死。 + # 退出失败同样只记录,不改变调用结果(见模块说明) + try: + await _scope.__aexit__(None, None, None) + except Exception as close_err: # noqa: BLE001 释放故障不得改变调用结果 + logger.error(f"释放数据库会话失败:{close_err}") + return result + + return wrapper diff --git a/app/db/diagnostics.py b/app/db/diagnostics.py new file mode 100644 index 000000000..02ed91ebd --- /dev/null +++ b/app/db/diagnostics.py @@ -0,0 +1,58 @@ +""" +数据库错误诊断。 + +把驱动层的错误分类字段(sqlite3 / psycopg2 / asyncpg 各不相同)提取成统一结构, +并挂到引擎的异常事件上,使排障不依赖于阅读原始驱动异常。 +""" +from typing import Any, Optional + +from sqlalchemy import event +from sqlalchemy.engine import Engine as SQLAlchemyEngine, ExceptionContext + +from app.runtime.log import logger + + +def _database_error_metadata(error: BaseException) -> Optional[dict[str, Any]]: + """提取 SQLite 与 PostgreSQL 驱动提供的稳定错误分类字段。""" + metadata = {"error_type": type(error).__name__} + + # DBAPI 驱动字段并不共享统一类型,动态读取可同时兼容 sqlite3、psycopg2 与 asyncpg。 + sqlite_errorcode = getattr(error, "sqlite_errorcode", None) + sqlite_errorname = getattr(error, "sqlite_errorname", None) + if sqlite_errorcode is not None or sqlite_errorname: + if sqlite_errorcode is not None: + metadata["error_code"] = sqlite_errorcode + if sqlite_errorname: + metadata["error_name"] = sqlite_errorname + return metadata + + sqlstate = getattr(error, "sqlstate", None) or getattr(error, "pgcode", None) + if not sqlstate: + sqlstate = getattr(getattr(error, "diag", None), "sqlstate", None) + if sqlstate: + metadata["sqlstate"] = sqlstate + return metadata + + return None + + +def _log_database_error(exception_context: ExceptionContext) -> None: + """记录非敏感驱动错误码,并保持 SQLAlchemy 原有异常传播。""" + metadata = _database_error_metadata(exception_context.original_exception) + if not metadata: + return + + dialect = exception_context.dialect + fields = { + "database": dialect.name, + "driver": dialect.driver, + **metadata, + } + logger.error( + "数据库驱动异常:" + ", ".join(f"{key}={value}" for key, value in fields.items()) + ) + + +def _register_database_error_logging(engine: SQLAlchemyEngine) -> None: + """为主程序 Engine 注册统一的底层驱动错误诊断。""" + event.listen(engine, "handle_error", _log_database_error) diff --git a/app/db/engine.py b/app/db/engine.py new file mode 100644 index 000000000..76013bfb8 --- /dev/null +++ b/app/db/engine.py @@ -0,0 +1,334 @@ +""" +数据库引擎的构建与连接额度核算。 + +同步引擎与未池化的全局异步引擎都在此按需创建(首次访问时,不在 import 期); +按事件循环池化的异步引擎由 session 模块创建。三者的构建参数在这里收口。 +""" +import threading +from typing import Dict, Optional, cast + +from sqlalchemy import NullPool, QueuePool, create_engine, text +from sqlalchemy.engine import Engine as SyncEngine +from sqlalchemy.ext.asyncio import AsyncEngine as SaAsyncEngine, create_async_engine + +from app.runtime.config import settings +from app.db.diagnostics import _register_database_error_logging +from app.runtime.log import logger + + +def _async_pool_kwargs(pooled: bool) -> dict: + """ + 异步引擎的连接池参数。 + + 池化时不指定 poolclass:SQLAlchemy 会自动选用异步适配的 + AsyncAdaptedQueuePool,显式传入同步的 QueuePool 反而会出错。 + :param pooled: 是否启用连接池 + :return: 传给 create_async_engine 的池参数 + """ + if not pooled: + return {"poolclass": NullPool} + return { + "pool_size": settings.DB_ASYNC_POOL_SIZE, + "max_overflow": settings.DB_ASYNC_MAX_OVERFLOW, + "pool_timeout": settings.DB_POOL_TIMEOUT, + } + + +def _get_database_engine(is_async: bool = False, pooled: bool = False): + """ + 获取数据库连接参数并设置WAL模式 + :param is_async: 是否创建异步引擎,True - 异步引擎, False - 同步引擎 + :param pooled: 异步引擎是否启用连接池,仅对常驻事件循环使用 + :return: 返回对应的数据库引擎 + """ + # 根据数据库类型选择连接方式 + if settings.DB_TYPE.lower() == "postgresql": + return _get_postgresql_engine(is_async, pooled=pooled) + else: + return _get_sqlite_engine(is_async, pooled=pooled) + + +def _get_sqlite_engine(is_async: bool = False, pooled: bool = False): + """ + 获取SQLite数据库引擎 + """ + # 连接参数 + _connect_args = { + "timeout": settings.DB_TIMEOUT, + } + # 允许部署侧注入驱动级参数(如 PgBouncer 事务模式下的 statement_cache_size) + _connect_args.update(settings.DB_CONNECT_ARGS or {}) + # 启用 WAL 模式时的额外配置 + if settings.DB_WAL_ENABLE: + _connect_args["check_same_thread"] = False + + # 创建同步引擎 + if not is_async: + # 根据池类型设置 poolclass 和相关参数 + _pool_class = NullPool if settings.DB_POOL_TYPE == "NullPool" else QueuePool + + # 数据库参数 + _db_kwargs = { + "url": settings.DB_SQLITE_URL(), + "pool_pre_ping": settings.DB_POOL_PRE_PING, + "echo": settings.DB_ECHO, + "poolclass": _pool_class, + "pool_recycle": settings.DB_POOL_RECYCLE, + "connect_args": _connect_args + } + + # 当使用 QueuePool 时,添加 QueuePool 特有的参数 + if _pool_class == QueuePool: + _db_kwargs.update({ + "pool_size": settings.DB_SQLITE_POOL_SIZE, + "pool_timeout": settings.DB_POOL_TIMEOUT, + "max_overflow": settings.DB_SQLITE_MAX_OVERFLOW + }) + + # 创建数据库引擎 + engine = create_engine(**_db_kwargs) + _register_database_error_logging(engine) + + # 设置WAL模式。 + # 这是引擎构建里唯一的阻塞 I/O,且发生在 get_engine() 的创建锁内——异步侧因此 + # 移除了对称的那一段(见下方 else 分支)。同步侧保留是因为 journal_mode 必须有人 + # 设置一次,而同步引擎的首次创建由 init_db() 在启动期单线程完成,不存在一群线程 + # 等在锁上的场面;即便退化到运行期首次访问,阻塞的也只是本地 SQLite 的一次 PRAGMA。 + _journal_mode = "WAL" if settings.DB_WAL_ENABLE else "DELETE" + with engine.connect() as connection: + current_mode = connection.execute(text(f"PRAGMA journal_mode={_journal_mode};")).scalar() + print(f"SQLite database journal mode set to: {current_mode}") + + return engine + else: + # 数据库参数,只能使用 NullPool + _db_kwargs = { + "url": settings.DB_SQLITE_URL("aiosqlite"), + "pool_pre_ping": settings.DB_POOL_PRE_PING, + "echo": settings.DB_ECHO, + "pool_recycle": settings.DB_POOL_RECYCLE, + "connect_args": _connect_args, + **_async_pool_kwargs(pooled), + } + # 创建异步数据库引擎 + async_engine = create_async_engine(**_db_kwargs) + _register_database_error_logging(async_engine.sync_engine) + + # 异步侧不再设置 WAL。journal_mode 是数据库文件级的持久属性,同步引擎已经设置过, + # 这里重复设置本就是冗余的;而它原本用 asyncio.run() 完成,是异步引擎构建里唯一的 + # 阻塞 I/O。引擎改为惰性创建之后,构建可能发生在任意线程——包括在运行中的事件循环 + # 内部(async_session_scope 首次取全局引擎时),那里调 asyncio.run() 会直接抛 + # RuntimeError;即便不抛,它也是在持有创建锁的状态下阻塞,会把所有等锁的线程拖死。 + return async_engine + + +def _get_postgresql_engine(is_async: bool = False, pooled: bool = False): + """ + 获取PostgreSQL数据库引擎 + """ + db_url = settings.DB_POSTGRESQL_URL() + + # PostgreSQL连接参数。允许部署侧注入驱动级参数, + # 例如经 PgBouncer 事务模式接入时 asyncpg 需要 statement_cache_size=0 + _connect_args = dict(settings.DB_CONNECT_ARGS or {}) + + # 创建同步引擎 + if not is_async: + # 根据池类型设置 poolclass 和相关参数 + _pool_class = NullPool if settings.DB_POOL_TYPE == "NullPool" else QueuePool + + # 数据库参数 + _db_kwargs = { + "url": db_url, + "pool_pre_ping": settings.DB_POOL_PRE_PING, + "echo": settings.DB_ECHO, + "poolclass": _pool_class, + "pool_recycle": settings.DB_POOL_RECYCLE, + "connect_args": _connect_args + } + + # 当使用 QueuePool 时,添加 QueuePool 特有的参数 + if _pool_class == QueuePool: + _db_kwargs.update({ + "pool_size": settings.DB_POSTGRESQL_POOL_SIZE, + "pool_timeout": settings.DB_POOL_TIMEOUT, + "max_overflow": settings.DB_POSTGRESQL_MAX_OVERFLOW + }) + + # 创建数据库引擎 + engine = create_engine(**_db_kwargs) + _register_database_error_logging(engine) + print(f"PostgreSQL database connected to {settings.DB_POSTGRESQL_TARGET}/{settings.DB_POSTGRESQL_DATABASE}") + + return engine + else: + async_db_url = settings.DB_POSTGRESQL_URL("asyncpg") + + # 数据库参数,只能使用 NullPool + _db_kwargs = { + "url": async_db_url, + "pool_pre_ping": settings.DB_POOL_PRE_PING, + "echo": settings.DB_ECHO, + "pool_recycle": settings.DB_POOL_RECYCLE, + "connect_args": _connect_args, + **_async_pool_kwargs(pooled), + } + # 创建异步数据库引擎 + async_engine = create_async_engine(**_db_kwargs) + _register_database_error_logging(async_engine.sync_engine) + print(f"Async PostgreSQL database connected to {settings.DB_POSTGRESQL_TARGET}/{settings.DB_POSTGRESQL_DATABASE}") + + return async_engine + + +# 引擎按需创建,不在 import 期建立连接。 +# +# 此前这两个是模块级常量,`import app.db` 就会按 settings 连库、建出 user.db、SQLite 还要 +# 去设一次 WAL——仅仅把这个包 import 进来(工具脚本、子进程探测、文档生成)就有了副作用。 +# +# 注意惰性化并没有让「隔离 CONFIG_DIR 必须早于 import」这条约束消失:settings 是在 +# import app.runtime.config 时构造的,那一刻 CONFIG_DIR 就定型了,晚建的引擎连的仍是真实库。 +# 它消掉的是「import 本身即产生副作用」,以及由此带来的「测试想换库就必须重新起进程」。 +# +# 惰性化引入的唯一新风险是首次访问的并发:这个项目有上百个调度线程,创建出多个引擎 +# 意味着各自持一份连接池,实际连接数是额度核算的数倍。因此用双重检查加锁收口。 +_sync_engine_lock = threading.RLock() +_async_engine_lock = threading.RLock() +_sync_engine: Optional[SyncEngine] = None +_async_engine: Optional[SaAsyncEngine] = None + + +def get_engine() -> SyncEngine: + """ + 获取同步数据库引擎,首次调用时创建。 + :return: 同步引擎 + """ + global _sync_engine + if _sync_engine is None: + with _sync_engine_lock: + # 锁内复查:等锁期间可能已被其它线程创建 + if _sync_engine is None: + _sync_engine = cast(SyncEngine, _get_database_engine(is_async=False)) + return _sync_engine + + +def get_global_async_engine() -> SaAsyncEngine: + """ + 获取未池化的全局异步引擎,供非常驻事件循环回退使用,首次调用时创建。 + :return: 异步引擎 + """ + global _async_engine + if _async_engine is None: + with _async_engine_lock: + if _async_engine is None: + _async_engine = cast(SaAsyncEngine, _get_database_engine(is_async=True)) + return _async_engine + + +def peek_sync_engine() -> Optional[SyncEngine]: + """ + 取已创建的同步引擎,未创建时返回 None——不触发创建。 + + 关停路径(close_database、测试引导的 atexit)需要「有就释放、没有就算了」: + 走 get_engine() 会为了 dispose 而先连一次库,在从未用过数据库的进程里尤其荒谬。 + 取锁而不是裸读槽位:另一个线程可能正卡在创建里,裸读会看到 None、把那个引擎漏掉, + 取锁则会等它建完。注意这只是把竞态窗口**缩小**,并没有消除——在 peek 返回之后才 + 开始创建的引擎照样漏。真要杜绝得让关停之后的创建直接失败,那是另一层面的改动。 + :return: 同步引擎或 None + """ + with _sync_engine_lock: + return _sync_engine + + +def peek_async_engine() -> Optional[SaAsyncEngine]: + """ + 取已创建的全局异步引擎,未创建时返回 None——不触发创建。 + :return: 异步引擎或 None + """ + with _async_engine_lock: + return _async_engine + + +def _async_pool_enabled() -> bool: + """ + 是否启用异步连接池。设为 NullPool 可回退到池化前的行为。 + """ + return str(settings.DB_ASYNC_POOL_TYPE or "").strip().lower() != "nullpool" + + +def connection_budget() -> Dict[str, int]: + """ + 核算数据库连接的理论峰值。 + + 各连接池此前是彼此独立配置的,没有任何地方核算总和——异步侧从无界收敛到有界 + 之后,真正决定安全与否的就变成了「同步池 + 异步池 + 回退配额」这个总数是否 + 还在数据库的额度之内。这里把它显式算出来,供启动校验与排障使用。 + 连接池是进程级的:多 worker 部署时每个进程各持一份,因此合计要乘上 worker 数。 + 只报单进程用量会让多 worker 在启动校验里一路绿灯,实际第一个 worker 还没起完 + 就顶穿了 max_connections。 + :return: 单进程各项上限、worker 数与合计 + """ + if settings.DB_TYPE.lower() == "postgresql": + sync_max = settings.DB_POSTGRESQL_POOL_SIZE + settings.DB_POSTGRESQL_MAX_OVERFLOW + else: + sync_max = settings.DB_SQLITE_POOL_SIZE + settings.DB_SQLITE_MAX_OVERFLOW + if settings.DB_POOL_TYPE == "NullPool": + # 同步侧也可能被配成 NullPool,此时同样无界,用线程池规模作为可观测的上限估计 + sync_max = settings.CONF.threadpool + async_max = (settings.DB_ASYNC_POOL_SIZE + settings.DB_ASYNC_MAX_OVERFLOW + if _async_pool_enabled() else 0) + fallback = settings.DB_ASYNC_FALLBACK_LIMIT if _async_pool_enabled() else settings.CONF.scheduler + per_worker = sync_max + async_max + fallback + # worker 数非法时按 1 计:退化成 0 会让合计归零、反而误判「额度充足」 + workers = getattr(settings, "API_WORKERS", 1) or 1 + workers = workers if isinstance(workers, int) and workers > 0 else 1 + return { + "sync": sync_max, + "async_pooled": async_max, + "async_fallback": fallback, + "per_worker": per_worker, + "workers": workers, + "total": per_worker * workers, + } + + +def check_connection_budget() -> bool: + """ + 对照数据库的真实连接上限校验理论峰值,超限时告警。 + + 只对 PostgreSQL 生效:SQLite 没有服务端连接上限,其压力体现为 WAL 写争用而非 + 连接耗尽。用真实的 max_connections 而不是猜测值——部署方可能已经调过它。 + :return: 是否在额度之内 + """ + budget = connection_budget() + if settings.DB_TYPE.lower() != "postgresql": + logger.info(f"数据库连接理论峰值: {budget['total']} " + f"(单进程 {budget['per_worker']} = 同步 {budget['sync']} + 异步池 " + f"{budget['async_pooled']} + 回退 {budget['async_fallback']}" + f",worker {budget['workers']})") + return True + try: + with get_engine().connect() as conn: + max_conn = int(conn.execute(text("SHOW max_connections")).scalar() or 0) + reserved = int( + conn.execute(text("SHOW superuser_reserved_connections")).scalar() or 0 + ) + except Exception as err: + logger.warn(f"无法读取 PostgreSQL 连接上限,跳过额度校验: {err}") + return True + available = max_conn - reserved + total = budget["total"] + detail = (f"理论峰值 {total} = 单进程 {budget['per_worker']} (同步 {budget['sync']} " + f"+ 异步池 {budget['async_pooled']} + 回退 {budget['async_fallback']}) " + f"x worker {budget['workers']},数据库可用 {available} " + f"(max_connections {max_conn} - 保留 {reserved})") + if total > available: + logger.error( + f"数据库连接额度不足:{detail}。" + f"突发并发时可能出现 TooManyConnectionsError。" + f"请调大 max_connections,或调小 API_WORKERS / DB_POSTGRESQL_MAX_OVERFLOW / " + f"DB_ASYNC_MAX_OVERFLOW / DB_ASYNC_FALLBACK_LIMIT" + ) + return False + logger.info(f"数据库连接额度校验通过:{detail}") + return True diff --git a/app/db/models/__init__.py b/app/db/models/__init__.py index c882d27a2..e3e79d765 100644 --- a/app/db/models/__init__.py +++ b/app/db/models/__init__.py @@ -1,3 +1,10 @@ +""" +ORM 模型。 + +_identity 必须在此处导入:它在 import 期把媒体身份归一挂到 mapper 事件上,是六张带 +身份列的表的写入不变量。导入任一模型都会先初始化本包,因此这一行让强制点无处可绕。 +""" +from . import _identity # noqa: F401 仅为注册 mapper 事件,不导出符号 from .agentchat import AgentChat from .agenttask import AgentTask from .agenttaskrun import AgentTaskRun diff --git a/app/db/models/media_identity.py b/app/db/models/_constraints.py similarity index 57% rename from app/db/models/media_identity.py rename to app/db/models/_constraints.py index 40480b351..f819a9c4a 100644 --- a/app/db/models/media_identity.py +++ b/app/db/models/_constraints.py @@ -1,3 +1,11 @@ +"""建表约束的共享片段——本模块不声明任何表,只被同包的模型模块拼进 ``__table_args__``。 + +以下划线开头且不进 ``models/__init__.py`` 的再导出,是为了与同目录「一实体一文件」的 +模块区分开:叫 ``media_identity.py`` 时它看着就像一张 MediaIdentity 表。 + +注意:alembic 迁移脚本必须自带 SQL 常量的副本而不是 import 本模块——迁移是历史快照, +跟着当前代码一起演进会让旧库重放出新约束。 +""" from sqlalchemy import CheckConstraint MEDIA_IDENTITY_CHECK_SQL = ( diff --git a/app/db/models/_identity.py b/app/db/models/_identity.py new file mode 100644 index 000000000..470324c88 --- /dev/null +++ b/app/db/models/_identity.py @@ -0,0 +1,71 @@ +""" +媒体身份的持久化不变量。 + +「media_source 与 media_id 必须成对、非零、去空白」这条规则此前由六张表的各个 Oper +在建模前各调一次 normalize_media_identity_payload 来保证——靠调用点的纪律,新加一条 +写入路径忘了调,就会静静写进半对身份,而按身份去重从此对这行失效。 + +这里把它下沉成 flush 前的 mapper 事件:凡同时具备两列的表,任何 ORM 写入都会经过, +忘不掉也绕不开。app/db 里没有 core insert()/bulk 写法(已核对),因此覆盖是完整的。 + +与 DTO 侧的失败语义有意不同,见下方 _normalize_identity 的说明。 +""" +from typing import Any + +from sqlalchemy import event +from sqlalchemy.orm import Mapper + +from app.runtime.log import logger +from app.schemas.media import resolve_media_identity + +# 构成媒体身份的两列,缺一不可 +IDENTITY_COLUMNS = ("media_source", "media_id") + + +def _normalize_identity(mapper: Mapper, connection: Any, target: Any) -> None: + """ + 写库前归一媒体身份;半对、非法或零值身份清空两列并记一条告警。 + + 为什么是「清空 + 告警」而不是像 DTO 侧那样抛错:这六张表都是记账性写入(整理历史、 + 下载历史、失败冷却、媒体服务器同步、订阅历史)。因身份不成对就让整条记录写不进去, + 等于用一个次要字段的问题换掉整条记账——而丢一条整理历史意味着那个文件可能被重复 + 整理。所以持久化侧选择降级保留,但**不再沉默**:告警让半对身份从「查不出的脏数据」 + 变成「日志里可检索的事件」。DTO 侧仍然抛错,那里是用户输入的边界,该当场拒绝。 + + :param mapper: 触发事件的映射器 + :param connection: 本次 flush 使用的连接,未用到 + :param target: 待写入的模型实例 + """ + columns = mapper.columns.keys() + if not all(name in columns for name in IDENTITY_COLUMNS): + return + raw_source = getattr(target, "media_source", None) + raw_id = getattr(target, "media_id", None) + if raw_source is None and raw_id is None: + return + media_source, media_id = resolve_media_identity( + media_source=raw_source, media_id=raw_id + ) + if not (media_source and media_id): + # 用映射类名而非 local_table.name:后者的静态类型是 FromClause,没有 name + logger.warn( + f"{mapper.class_.__name__} 的媒体身份不成对,已清空:" + f"media_source={raw_source!r}, media_id={raw_id!r}" + ) + target.media_source = media_source.value if media_source else None + target.media_id = media_id + + +def register_identity_normalizer() -> None: + """ + 注册身份归一事件。 + + 监听 Mapper 类本身而非某个基类:Base 自身没有表、不是映射类,挂不上 mapper 事件; + 挂在 Mapper 上则覆盖进程内全部映射,包括仓外插件自建的模型。开销由上面那行列名 + 检查兜住——不具备身份列的表直接返回。 + """ + event.listen(Mapper, "before_insert", _normalize_identity) + event.listen(Mapper, "before_update", _normalize_identity) + + +register_identity_normalizer() diff --git a/app/db/models/agentchat.py b/app/db/models/agentchat.py index 56d0160ce..80d5c3e77 100644 --- a/app/db/models/agentchat.py +++ b/app/db/models/agentchat.py @@ -1,8 +1,8 @@ -from typing import Optional +from typing import Any, Optional -from sqlalchemy import Column, Integer, String, JSON, Index, select +from sqlalchemy import Integer, String, JSON, Index, select from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column from app.db import Base, async_db_query, db_query, get_id_column @@ -14,33 +14,33 @@ class AgentChat(Base): id = get_id_column() # Agent 内部会话 ID,用于恢复 LangGraph 对话上下文 - session_id = Column(String, nullable=False) + session_id: Mapped[str] = mapped_column(String, nullable=False) # 前端或渠道侧传入的原始会话标识 - client_session_id = Column(String) + client_session_id: Mapped[Optional[str]] = mapped_column(String) # 用户 ID - user_id = Column(String) + user_id: Mapped[Optional[str]] = mapped_column(String) # 用户名称 - username = Column(String) + username: Mapped[Optional[str]] = mapped_column(String) # 消息渠道 - channel = Column(String) + channel: Mapped[Optional[str]] = mapped_column(String) # 渠道来源配置名 - source = Column(String) + source: Mapped[Optional[str]] = mapped_column(String) # 原聊天 ID,用于区分群聊、频道或私聊 - original_chat_id = Column(String) + original_chat_id: Mapped[Optional[str]] = mapped_column(String) # 会话标题 - title = Column(String) + title: Mapped[Optional[str]] = mapped_column(String) # 会话预览文本 - preview = Column(String) + preview: Mapped[Optional[str]] = mapped_column(String) # 原始 LangChain messages,用于继续会话 - agent_messages = Column(JSON) + agent_messages: Mapped[Optional[Any]] = mapped_column(JSON) # 展示给用户的消息记录,包含文字、工具提示、附件与选择卡片 - display_messages = Column(JSON) + display_messages: Mapped[Optional[Any]] = mapped_column(JSON) # 展示消息数量 - message_count = Column(Integer, default=0) + message_count: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 创建时间 - created_at = Column(String) + created_at: Mapped[Optional[str]] = mapped_column(String) # 更新时间 - updated_at = Column(String) + updated_at: Mapped[Optional[str]] = mapped_column(String) __table_args__ = ( Index("ix_agentchat_session_user", "session_id", "user_id"), @@ -56,10 +56,10 @@ class AgentChat(Base): """ 根据会话 ID 获取 Agent 会话。 """ - query = db.query(cls).filter(cls.session_id == session_id) + statement = select(cls).where(cls.session_id == session_id) if user_id is not None: - query = query.filter(cls.user_id == user_id) - return query.order_by(cls.id.desc()).first() + statement = statement.where(cls.user_id == user_id) + return db.execute(statement.order_by(cls.id.desc())).scalars().first() @classmethod @async_db_query @@ -80,35 +80,34 @@ class AgentChat(Base): def list_by_page( cls, db: Session, - page: Optional[int] = 1, - count: Optional[int] = 30, + page: int = 1, + count: int = 30, user_id: Optional[str] = None, username: Optional[str] = None, ) -> list["AgentChat"]: """ 分页获取 Agent 会话历史。 """ - query = db.query(cls) + statement = select(cls) if user_id is not None and username is not None: - query = query.filter((cls.user_id == user_id) | (cls.username == username)) + statement = statement.where((cls.user_id == user_id) | (cls.username == username)) elif user_id is not None: - query = query.filter(cls.user_id == user_id) + statement = statement.where(cls.user_id == user_id) elif username is not None: - query = query.filter(cls.username == username) - return ( - query.order_by(cls.updated_at.desc(), cls.id.desc()) + statement = statement.where(cls.username == username) + return list(db.execute( + statement.order_by(cls.updated_at.desc(), cls.id.desc()) .offset((page - 1) * count) .limit(count) - .all() - ) + ).scalars().all()) @classmethod @async_db_query async def async_list_by_page( cls, db: AsyncSession, - page: Optional[int] = 1, - count: Optional[int] = 30, + page: int = 1, + count: int = 30, user_id: Optional[str] = None, username: Optional[str] = None, ) -> list["AgentChat"]: @@ -127,4 +126,4 @@ class AgentChat(Base): .offset((page - 1) * count) .limit(count) ) - return result.scalars().all() + return list(result.scalars().all()) diff --git a/app/db/models/agenttask.py b/app/db/models/agenttask.py index e95ad81b3..06d96f4a6 100644 --- a/app/db/models/agenttask.py +++ b/app/db/models/agenttask.py @@ -1,9 +1,9 @@ from typing import Optional -from sqlalchemy import Boolean, Column, Index, Integer, String, Text -from sqlalchemy.orm import Session +from sqlalchemy import Boolean, Index, Integer, String, Text, select, update +from sqlalchemy.orm import Mapped, Session, mapped_column -from app.db import Base, db_query, db_update, get_id_column +from app.db import Base, db_query, db_update, execute_dml, get_id_column class AgentTask(Base): @@ -13,34 +13,34 @@ class AgentTask(Base): id = get_id_column() # 任务名称 - name = Column(String, nullable=False) + name: Mapped[str] = mapped_column(String, nullable=False) # 交给 Agent 执行的完整任务内容 - content = Column(Text, nullable=False) + content: Mapped[str] = mapped_column(Text, nullable=False) # 触发类型:date-单次触发,cron-周期触发 - trigger_type = Column(String, nullable=False) + trigger_type: Mapped[str] = mapped_column(String, nullable=False) # 标准五段 cron 表达式 - cron_expression = Column(String) + cron_expression: Mapped[Optional[str]] = mapped_column(String) # 单次触发时间,使用带时区的 ISO 8601 格式 - run_at = Column(String) + run_at: Mapped[Optional[str]] = mapped_column(String) # 是否继续接受调度 - enabled = Column(Boolean, nullable=False, default=True) + enabled: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True) # 创建任务的用户与会话上下文 - user_id = Column(String, nullable=False) - username = Column(String) - session_id = Column(String, nullable=False) - channel = Column(String) - source = Column(String) - original_chat_id = Column(String) + user_id: Mapped[str] = mapped_column(String, nullable=False) + username: Mapped[Optional[str]] = mapped_column(String) + session_id: Mapped[str] = mapped_column(String, nullable=False) + channel: Mapped[Optional[str]] = mapped_column(String) + source: Mapped[Optional[str]] = mapped_column(String) + original_chat_id: Mapped[Optional[str]] = mapped_column(String) # 最近一次执行状态与结果 - last_status = Column(String, nullable=False, default="waiting") - last_run_at = Column(String) - last_result = Column(Text) + last_status: Mapped[str] = mapped_column(String, nullable=False, default="waiting") + last_run_at: Mapped[Optional[str]] = mapped_column(String) + last_result: Mapped[Optional[str]] = mapped_column(Text) # 最新一次真实执行的公开 ID,用于保护 last_* 投影不被旧运行覆盖 - last_run_id = Column(String) + last_run_id: Mapped[Optional[str]] = mapped_column(String) # 已收口执行次数;进程中断的未完成尝试不计入 - run_count = Column(Integer, nullable=False, default=0) - created_at = Column(String, nullable=False) - updated_at = Column(String, nullable=False) + run_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + created_at: Mapped[str] = mapped_column(String, nullable=False) + updated_at: Mapped[str] = mapped_column(String, nullable=False) __table_args__ = ( Index("ix_agenttask_enabled", "enabled"), @@ -69,10 +69,10 @@ class AgentTask(Base): """ 按任务 ID 和可选用户 ID 查询 Agent 定时任务。 """ - query = db.query(cls).filter(cls.id == task_id) + statement = select(cls).where(cls.id == task_id) if user_id is not None: - query = query.filter(cls.user_id == user_id) - return query.first() + statement = statement.where(cls.user_id == user_id) + return db.execute(statement).scalars().first() @classmethod @db_query @@ -85,12 +85,14 @@ class AgentTask(Base): """ 按用户和启用状态查询 Agent 定时任务。 """ - query = db.query(cls) + statement = select(cls) if user_id is not None: - query = query.filter(cls.user_id == user_id) + statement = statement.where(cls.user_id == user_id) if enabled is not None: - query = query.filter(cls.enabled.is_(enabled)) - return query.order_by(cls.created_at.desc(), cls.id.desc()).all() + statement = statement.where(cls.enabled.is_(enabled)) + return list(db.execute( + statement.order_by(cls.created_at.desc(), cls.id.desc()) + ).scalars().all()) @classmethod @db_update @@ -107,10 +109,10 @@ class AgentTask(Base): 运行状态与配置必须在同一条条件更新中判定,避免执行认领后被并发配置写入 覆盖回可再次执行的状态。 """ - query = db.query(cls).filter( + statement = update(cls).where( cls.id == task_id, cls.last_status != "running", ) if user_id is not None: - query = query.filter(cls.user_id == user_id) - return bool(query.update(payload)) + statement = statement.where(cls.user_id == user_id) + return bool(execute_dml(db, statement.values(payload))) diff --git a/app/db/models/agenttaskrun.py b/app/db/models/agenttaskrun.py index 4b72fd95d..1f013da6a 100644 --- a/app/db/models/agenttaskrun.py +++ b/app/db/models/agenttaskrun.py @@ -1,9 +1,9 @@ -from typing import Optional +from typing import Any, Dict, List, Optional -from sqlalchemy import Column, Index, Integer, String, Text, update -from sqlalchemy.orm import Session +from sqlalchemy import Index, Integer, String, Text, delete, select, update +from sqlalchemy.orm import Mapped, Session, mapped_column -from app.db import Base, db_query, db_update, get_id_column +from app.db import Base, db_query, db_update, execute_dml, get_id_column from app.db.models.agenttask import AgentTask @@ -12,27 +12,27 @@ class AgentTaskRun(Base): id = get_id_column() # 对外稳定的运行身份;内部自增主键不进入 Agent 合同 - run_id = Column(String, nullable=False) + run_id: Mapped[str] = mapped_column(String, nullable=False) # 所属计划及触发入口 - task_id = Column(Integer, nullable=False) - trigger_source = Column(String, nullable=False) + task_id: Mapped[int] = mapped_column(Integer, nullable=False) + trigger_source: Mapped[str] = mapped_column(String, nullable=False) # 执行开始时的任务与用户上下文快照 - name = Column(String, nullable=False) - content = Column(Text, nullable=False) - trigger_type = Column(String, nullable=False) - cron_expression = Column(String) - run_at = Column(String) - user_id = Column(String, nullable=False) - username = Column(String) - session_id = Column(String, nullable=False) - channel = Column(String) - message_source = Column(String) - original_chat_id = Column(String) + name: Mapped[str] = mapped_column(String, nullable=False) + content: Mapped[str] = mapped_column(Text, nullable=False) + trigger_type: Mapped[str] = mapped_column(String, nullable=False) + cron_expression: Mapped[Optional[str]] = mapped_column(String) + run_at: Mapped[Optional[str]] = mapped_column(String) + user_id: Mapped[str] = mapped_column(String, nullable=False) + username: Mapped[Optional[str]] = mapped_column(String) + session_id: Mapped[str] = mapped_column(String, nullable=False) + channel: Mapped[Optional[str]] = mapped_column(String) + message_source: Mapped[Optional[str]] = mapped_column(String) + original_chat_id: Mapped[Optional[str]] = mapped_column(String) # running-success/failed/interrupted;取消沿用 failed 和明确结果文本 - status = Column(String, nullable=False) - started_at = Column(String, nullable=False) - finished_at = Column(String) - result = Column(Text) + status: Mapped[str] = mapped_column(String, nullable=False) + started_at: Mapped[str] = mapped_column(String, nullable=False) + finished_at: Mapped[Optional[str]] = mapped_column(String) + result: Mapped[Optional[str]] = mapped_column(Text) __table_args__ = ( Index("ix_agenttaskrun_run_id", "run_id", unique=True), @@ -117,32 +117,37 @@ class AgentTaskRun(Base): disable_date_task: bool = False, ) -> bool: """原子收口精确运行,并仅在仍为最新运行时更新任务投影。""" - run = db.query(cls).filter( - cls.run_id == run_id, - ).first() + run = db.execute( + select(cls).where(cls.run_id == run_id) + ).scalars().first() if not run: return False status = "success" if success else "failed" - finalized = db.query(cls).filter( - cls.run_id == run_id, - cls.status == "running", - ).update( - { - "status": status, - "result": result, - "finished_at": finished_at, - }, - synchronize_session=False, + finalized = execute_dml( + db, + update(cls) + .where( + cls.run_id == run_id, + cls.status == "running", + ) + .values( + status=status, + result=result, + finished_at=finished_at, + ), + execution_options={"synchronize_session": False}, ) if not finalized: return False - task = db.query(AgentTask).filter( - AgentTask.id == run.task_id, - AgentTask.last_run_id == run_id, - ).first() + task = db.execute( + select(AgentTask).where( + AgentTask.id == run.task_id, + AgentTask.last_run_id == run_id, + ) + ).scalars().first() if task: - payload = { + payload: Dict[str, Any] = { "last_status": status, "last_result": result, "run_count": AgentTask.run_count + 1, @@ -155,10 +160,16 @@ class AgentTaskRun(Base): and task.run_at == run.run_at ): payload["enabled"] = False - db.query(AgentTask).filter( - AgentTask.id == run.task_id, - AgentTask.last_run_id == run_id, - ).update(payload, synchronize_session=False) + execute_dml( + db, + update(AgentTask) + .where( + AgentTask.id == run.task_id, + AgentTask.last_run_id == run_id, + ) + .values(**payload), + execution_options={"synchronize_session": False}, + ) return True @classmethod @@ -171,38 +182,46 @@ class AgentTaskRun(Base): finished_at: str, ) -> bool: """原子标记冷启动时遗留的最新运行及任务投影为结果未知。""" - task = db.query(AgentTask).filter( - AgentTask.id == task_id, - AgentTask.last_status == "running", - ).first() + task = db.execute( + select(AgentTask).where( + AgentTask.id == task_id, + AgentTask.last_status == "running", + ) + ).scalars().first() if not task: return False if task.last_run_id: - interrupted = db.query(cls).filter( - cls.run_id == task.last_run_id, - cls.task_id == task.id, - cls.status == "running", - ).update( - { - "status": "interrupted", - "result": result, - "finished_at": finished_at, - }, - synchronize_session=False, + interrupted = execute_dml( + db, + update(cls) + .where( + cls.run_id == task.last_run_id, + cls.task_id == task.id, + cls.status == "running", + ) + .values( + status="interrupted", + result=result, + finished_at=finished_at, + ), + execution_options={"synchronize_session": False}, ) if not interrupted: return False - return bool(db.query(AgentTask).filter( - AgentTask.id == task.id, - AgentTask.last_status == "running", - AgentTask.last_run_id == task.last_run_id, - ).update( - { - "last_status": "interrupted", - "last_result": result, - "updated_at": finished_at, - }, - synchronize_session=False, + return bool(execute_dml( + db, + update(AgentTask) + .where( + AgentTask.id == task.id, + AgentTask.last_status == "running", + AgentTask.last_run_id == task.last_run_id, + ) + .values( + last_status="interrupted", + last_result=result, + updated_at=finished_at, + ), + execution_options={"synchronize_session": False}, )) @classmethod @@ -214,16 +233,22 @@ class AgentTaskRun(Base): user_id: Optional[str] = None, ) -> bool: """原子删除非运行中任务及其执行历史。""" - query = db.query(AgentTask).filter( + statement = delete(AgentTask).where( AgentTask.id == task_id, AgentTask.last_status != "running", ) if user_id is not None: - query = query.filter(AgentTask.user_id == user_id) - deleted = query.delete(synchronize_session=False) + statement = statement.where(AgentTask.user_id == user_id) + deleted = execute_dml( + db, statement, execution_options={"synchronize_session": False} + ) if not deleted: return False - db.query(cls).filter(cls.task_id == task_id).delete(synchronize_session=False) + execute_dml( + db, + delete(cls).where(cls.task_id == task_id), + execution_options={"synchronize_session": False}, + ) return True @classmethod @@ -234,7 +259,9 @@ class AgentTaskRun(Base): run_id: str, ) -> Optional["AgentTaskRun"]: """按公开运行 ID 查询一次执行。""" - return db.query(cls).filter(cls.run_id == run_id).first() + return db.execute( + select(cls).where(cls.run_id == run_id) + ).scalars().first() @classmethod @db_query @@ -244,11 +271,13 @@ class AgentTaskRun(Base): task_id: int, user_id: Optional[str] = None, limit: int = 10, - ) -> list["AgentTaskRun"]: + ) -> List["AgentTaskRun"]: """按父任务 owner 校验后返回最近的有界运行历史。""" - query = db.query(cls).join(AgentTask, AgentTask.id == cls.task_id).filter( + statement = select(cls).join(AgentTask, AgentTask.id == cls.task_id).where( cls.task_id == task_id, ) if user_id is not None: - query = query.filter(AgentTask.user_id == user_id) - return query.order_by(cls.started_at.desc(), cls.id.desc()).limit(limit).all() + statement = statement.where(AgentTask.user_id == user_id) + return list(db.execute( + statement.order_by(cls.started_at.desc(), cls.id.desc()).limit(limit) + ).scalars().all()) diff --git a/app/db/models/downloadfailure.py b/app/db/models/downloadfailure.py index 695a15008..438359f33 100644 --- a/app/db/models/downloadfailure.py +++ b/app/db/models/downloadfailure.py @@ -1,10 +1,10 @@ from typing import List, Optional -from sqlalchemy import Column, Float, Index, Integer, String -from sqlalchemy.orm import Session +from sqlalchemy import Float, Index, Integer, String, delete, select +from sqlalchemy.orm import Mapped, Session, mapped_column -from app.db import Base, db_query, db_update, get_id_column -from app.db.models.media_identity import media_identity_constraint +from app.db import Base, db_query, db_update, execute_dml, get_id_column +from app.db.models._constraints import media_identity_constraint class DownloadFailure(Base): @@ -14,44 +14,44 @@ class DownloadFailure(Base): id = get_id_column() # 资源失败指纹 - fingerprint = Column(String, nullable=False) + fingerprint: Mapped[str] = mapped_column(String, nullable=False) # 类型 电影/电视剧 - type = Column(String) + type: Mapped[Optional[str]] = mapped_column(String) # 标题 - title = Column(String) + title: Mapped[Optional[str]] = mapped_column(String) # 年份 - year = Column(String) + year: Mapped[Optional[str]] = mapped_column(String) # 媒体数据源与原生ID - media_source = Column(String) - media_id = Column(String) + media_source: Mapped[Optional[str]] = mapped_column(String) + media_id: Mapped[Optional[str]] = mapped_column(String) # Sxx - seasons = Column(String) + seasons: Mapped[Optional[str]] = mapped_column(String) # Exx - episodes = Column(String) + episodes: Mapped[Optional[str]] = mapped_column(String) # 站点ID - site = Column(Integer) + site: Mapped[Optional[int]] = mapped_column(Integer) # 站点名称 - site_name = Column(String) + site_name: Mapped[Optional[str]] = mapped_column(String) # 种子资源键 - torrent_id = Column(String) + torrent_id: Mapped[Optional[str]] = mapped_column(String) # 种子名称 - torrent_name = Column(String) + torrent_name: Mapped[Optional[str]] = mapped_column(String) # 种子大小 - torrent_size = Column(Float) + torrent_size: Mapped[Optional[float]] = mapped_column(Float) # 下载器 - downloader = Column(String) + downloader: Mapped[Optional[str]] = mapped_column(String) # 下载来源 - source = Column(String) + source: Mapped[Optional[str]] = mapped_column(String) # 失败原因 - error_message = Column(String) + error_message: Mapped[Optional[str]] = mapped_column(String) # 重试次数 - retry_count = Column(Integer, default=0) + retry_count: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 首次失败时间 - first_failed_at = Column(String) + first_failed_at: Mapped[Optional[str]] = mapped_column(String) # 最近失败时间 - last_failed_at = Column(String) + last_failed_at: Mapped[Optional[str]] = mapped_column(String) # 下次允许重试时间 - next_retry_at = Column(String) + next_retry_at: Mapped[Optional[str]] = mapped_column(String) __table_args__ = ( media_identity_constraint("downloadfailure"), @@ -74,11 +74,10 @@ class DownloadFailure(Base): normalized = list(dict.fromkeys([fingerprint for fingerprint in fingerprints if fingerprint])) if not normalized: return [] - return ( - db.query(cls) - .filter(cls.fingerprint.in_(normalized), cls.next_retry_at > now_time) - .all() - ) + return list(db.execute( + select(cls) + .where(cls.fingerprint.in_(normalized), cls.next_retry_at > now_time) + ).scalars().all()) @classmethod @db_update @@ -93,7 +92,9 @@ class DownloadFailure(Base): """ 新增或更新资源失败记录。 """ - failure = db.query(cls).filter(cls.fingerprint == fingerprint).first() + failure = db.execute( + select(cls).where(cls.fingerprint == fingerprint) + ).scalars().first() payload = { **kwargs, "fingerprint": fingerprint, @@ -125,14 +126,15 @@ class DownloadFailure(Base): """ 分批清理已过期较久的失败冷却记录。 """ - ids = [ - row[0] - for row in db.query(cls.id) - .filter(cls.next_retry_at < before_time) + ids = db.execute( + select(cls.id) + .where(cls.next_retry_at < before_time) .order_by(cls.id.asc()) .limit(limit) - .all() - ] + ).scalars().all() if not ids: return 0 - return db.query(cls).filter(cls.id.in_(ids)).delete(synchronize_session=False) + return execute_dml( + db, delete(cls).where(cls.id.in_(ids)), + execution_options={"synchronize_session": False}, + ) diff --git a/app/db/models/downloadhistory.py b/app/db/models/downloadhistory.py index 6e9db79de..6e541cf5b 100644 --- a/app/db/models/downloadhistory.py +++ b/app/db/models/downloadhistory.py @@ -1,12 +1,12 @@ import time -from typing import List, Optional +from typing import Any, List, Optional -from sqlalchemy import Column, Integer, String, JSON, Index, select, func +from sqlalchemy import Integer, String, JSON, Index, delete, select, func, update from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column -from app.db import db_query, db_update, get_id_column, Base, async_db_query -from app.db.models.media_identity import media_identity_constraint +from app.db import Base, async_db_query, db_query, db_update, execute_dml, get_id_column +from app.db.models._constraints import media_identity_constraint from app.schemas.types import MediaSource @@ -22,51 +22,51 @@ class DownloadHistory(Base): id = get_id_column() # 保存路径 - path = Column(String, nullable=False, index=True) + path: Mapped[str] = mapped_column(String, nullable=False, index=True) # 类型 电影/电视剧/音乐 - type = Column(String, nullable=False) + type: Mapped[str] = mapped_column(String, nullable=False) # 标题 - title = Column(String, nullable=False) + title: Mapped[str] = mapped_column(String, nullable=False) # 年份 - year = Column(String) - media_source = Column(String, index=True) - media_id = Column(String, index=True) + year: Mapped[Optional[str]] = mapped_column(String) + media_source: Mapped[Optional[str]] = mapped_column(String, index=True) + media_id: Mapped[Optional[str]] = mapped_column(String, index=True) # 音乐实体类型:recording 单曲、album 专辑 - music_type = Column(String) + music_type: Mapped[Optional[str]] = mapped_column(String) # Sxx - seasons = Column(String) + seasons: Mapped[Optional[str]] = mapped_column(String) # Exx - episodes = Column(String) + episodes: Mapped[Optional[str]] = mapped_column(String) # 背景图 - image = Column(String) + image: Mapped[Optional[str]] = mapped_column(String) # 海报 - poster = Column(String) + poster: Mapped[Optional[str]] = mapped_column(String) # 下载器 - downloader = Column(String) + downloader: Mapped[Optional[str]] = mapped_column(String) # 下载任务Hash - download_hash = Column(String) + download_hash: Mapped[Optional[str]] = mapped_column(String) # 种子名称 - torrent_name = Column(String) + torrent_name: Mapped[Optional[str]] = mapped_column(String) # 种子描述 - torrent_description = Column(String) + torrent_description: Mapped[Optional[str]] = mapped_column(String) # 种子站点 - torrent_site = Column(String) + torrent_site: Mapped[Optional[str]] = mapped_column(String) # 下载用户 - userid = Column(String) + userid: Mapped[Optional[str]] = mapped_column(String) # 下载用户名/插件名 - username = Column(String) + username: Mapped[Optional[str]] = mapped_column(String) # 下载渠道 - channel = Column(String) + channel: Mapped[Optional[str]] = mapped_column(String) # 创建时间 - date = Column(String) + date: Mapped[Optional[str]] = mapped_column(String) # 附加信息 - note = Column(JSON) + note: Mapped[Optional[Any]] = mapped_column(JSON) # 自定义媒体类别 - media_category = Column(String) + media_category: Mapped[Optional[str]] = mapped_column(String) # 剧集组 - episode_group = Column(String) + episode_group: Mapped[Optional[str]] = mapped_column(String) # 自定义识别词(用于整理时应用) - custom_words = Column(String) + custom_words: Mapped[Optional[str]] = mapped_column(String) __table_args__ = ( media_identity_constraint("downloadhistory"), @@ -78,12 +78,11 @@ class DownloadHistory(Base): @classmethod @db_query def get_by_hash(cls, db: Session, download_hash: str): - return ( - db.query(DownloadHistory) - .filter(DownloadHistory.download_hash == download_hash) + return db.execute( + select(DownloadHistory) + .where(DownloadHistory.download_hash == download_hash) .order_by(DownloadHistory.date.desc()) - .first() - ) + ).scalars().first() @classmethod @db_query @@ -102,12 +101,11 @@ class DownloadHistory(Base): if not normalized_hashes: return [] - histories = ( - db.query(DownloadHistory) - .filter(DownloadHistory.download_hash.in_(normalized_hashes)) + histories = db.execute( + select(DownloadHistory) + .where(DownloadHistory.download_hash.in_(normalized_hashes)) .order_by(DownloadHistory.download_hash, DownloadHistory.date.desc()) - .all() - ) + ).scalars().all() latest_histories = {} for history in histories: if history.download_hash and history.download_hash not in latest_histories: @@ -128,32 +126,30 @@ class DownloadHistory(Base): """按规范媒体身份查询下载历史。""" if not media_source or media_id is None or not str(media_id).strip(): return [] - query = db.query(DownloadHistory) - query = query.filter( + statement = select(DownloadHistory).where( DownloadHistory.media_source == str(media_source), DownloadHistory.media_id == str(media_id).strip(), ) if music_type: - query = query.filter(DownloadHistory.music_type == music_type) - return query.all() + statement = statement.where(DownloadHistory.music_type == music_type) + return list(db.execute(statement).scalars().all()) @classmethod @db_query def list_by_page( - cls, db: Session, page: Optional[int] = 1, count: Optional[int] = 30 + cls, db: Session, page: int = 1, count: int = 30 ): - return ( - db.query(DownloadHistory) + return list(db.execute( + select(DownloadHistory) .order_by(DownloadHistory.date.desc(), DownloadHistory.id.desc()) .offset((page - 1) * count) .limit(count) - .all() - ) + ).scalars().all()) @classmethod @async_db_query async def async_list_by_page( - cls, db: AsyncSession, page: Optional[int] = 1, count: Optional[int] = 30 + cls, db: AsyncSession, page: int = 1, count: int = 30 ): result = await db.execute( select(cls) @@ -161,7 +157,7 @@ class DownloadHistory(Base): .offset((page - 1) * count) .limit(count) ) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @async_db_query @@ -169,15 +165,15 @@ class DownloadHistory(Base): cls, db: AsyncSession, title: str, - page: Optional[int] = 1, - count: Optional[int] = 30, + page: int = 1, + count: int = 30, ): query = ( select(cls).filter(_title_like(cls.title, title)).order_by(cls.date.desc()) ) query = query.offset((page - 1) * count).limit(count) result = await db.execute(query) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @async_db_query @@ -196,7 +192,9 @@ class DownloadHistory(Base): @classmethod @db_query def get_by_path(cls, db: Session, path: str): - return db.query(DownloadHistory).filter(DownloadHistory.path == path).first() + return db.execute( + select(DownloadHistory).where(DownloadHistory.path == path) + ).scalars().first() @classmethod @db_query @@ -215,107 +213,46 @@ class DownloadHistory(Base): 按媒体身份、季集或标题年份查询下载记录。 """ if media_source and media_id and mtype: - # 电视剧某季某集 - if season is not None and episode: - return ( - db.query(DownloadHistory) - .filter( - DownloadHistory.media_source == str(media_source), - DownloadHistory.media_id == str(media_id), - DownloadHistory.type == mtype, - DownloadHistory.seasons == season, - DownloadHistory.episodes == episode, - ) - .order_by(DownloadHistory.id.desc()) - .all() - ) - # 电视剧某季 - elif season is not None: - return ( - db.query(DownloadHistory) - .filter( - DownloadHistory.media_source == str(media_source), - DownloadHistory.media_id == str(media_id), - DownloadHistory.type == mtype, - DownloadHistory.seasons == season, - ) - .order_by(DownloadHistory.id.desc()) - .all() - ) - else: - # 电视剧所有季集/电影 - return ( - db.query(DownloadHistory) - .filter( - DownloadHistory.media_source == str(media_source), - DownloadHistory.media_id == str(media_id), - DownloadHistory.type == mtype, - ) - .order_by(DownloadHistory.id.desc()) - .all() - ) - # 标题 + 年份 + statement = select(DownloadHistory).where( + DownloadHistory.media_source == str(media_source), + DownloadHistory.media_id == str(media_id), + DownloadHistory.type == mtype, + ) elif title and year: - # 电视剧某季某集 - if season is not None and episode: - return ( - db.query(DownloadHistory) - .filter( - DownloadHistory.title == title, - DownloadHistory.year == year, - DownloadHistory.seasons == season, - DownloadHistory.episodes == episode, - ) - .order_by(DownloadHistory.id.desc()) - .all() - ) - # 电视剧某季 - elif season is not None: - return ( - db.query(DownloadHistory) - .filter( - DownloadHistory.title == title, - DownloadHistory.year == year, - DownloadHistory.seasons == season, - ) - .order_by(DownloadHistory.id.desc()) - .all() - ) - else: - # 电视剧所有季集/电影 - return ( - db.query(DownloadHistory) - .filter( - DownloadHistory.title == title, DownloadHistory.year == year - ) - .order_by(DownloadHistory.id.desc()) - .all() - ) + statement = select(DownloadHistory).where( + DownloadHistory.title == title, + DownloadHistory.year == year, + ) + else: + return [] + # 季、集逐级收窄:给出季才可能给集,与原六条分支等价 + if season is not None: + statement = statement.where(DownloadHistory.seasons == season) + if episode: + statement = statement.where(DownloadHistory.episodes == episode) + return list(db.execute( + statement.order_by(DownloadHistory.id.desc()) + ).scalars().all()) - return [] @classmethod @db_query def list_by_user_date(cls, db: Session, date: str, username: Optional[str] = None): """ - 查询某用户某时间之后的下载历史 + 查询某用户某时间之前的下载历史。 + + 条件是 date < 传入时刻,等于该时刻的那条不计入;oper 层的同名方法描述一致。 + :param db: 数据库会话 + :param date: 时间水位,取该时刻之前的记录 + :param username: 下载用户,不传则跨用户返回 + :return: 下载历史列表,按主键倒序 """ + statement = select(DownloadHistory).where(DownloadHistory.date < date) if username: - return ( - db.query(DownloadHistory) - .filter( - DownloadHistory.date < date, DownloadHistory.username == username - ) - .order_by(DownloadHistory.id.desc()) - .all() - ) - else: - return ( - db.query(DownloadHistory) - .filter(DownloadHistory.date < date) - .order_by(DownloadHistory.id.desc()) - .all() - ) + statement = statement.where(DownloadHistory.username == username) + return list(db.execute( + statement.order_by(DownloadHistory.id.desc()) + ).scalars().all()) @classmethod @db_query @@ -331,46 +268,30 @@ class DownloadHistory(Base): """ 查询某时间之后的下载历史 """ + statement = select(DownloadHistory).where( + DownloadHistory.date > date, + DownloadHistory.type == type, + DownloadHistory.media_source == str(media_source), + DownloadHistory.media_id == str(media_id), + ) if seasons: - return ( - db.query(DownloadHistory) - .filter( - DownloadHistory.date > date, - DownloadHistory.type == type, - DownloadHistory.media_source == str(media_source), - DownloadHistory.media_id == str(media_id), - DownloadHistory.seasons == seasons, - ) - .order_by(DownloadHistory.id.desc()) - .all() - ) - else: - return ( - db.query(DownloadHistory) - .filter( - DownloadHistory.date > date, - DownloadHistory.type == type, - DownloadHistory.media_source == str(media_source), - DownloadHistory.media_id == str(media_id), - ) - .order_by(DownloadHistory.id.desc()) - .all() - ) + statement = statement.where(DownloadHistory.seasons == seasons) + return list(db.execute( + statement.order_by(DownloadHistory.id.desc()) + ).scalars().all()) @classmethod @db_query def list_by_type(cls, db: Session, mtype: str, days: int): - return ( - db.query(DownloadHistory) - .filter( + return list(db.execute( + select(DownloadHistory).where( DownloadHistory.type == mtype, DownloadHistory.date >= time.strftime( "%Y-%m-%d %H:%M:%S", time.localtime(time.time() - 86400 * int(days)) ), ) - .all() - ) + ).scalars().all()) @classmethod @db_update @@ -383,20 +304,17 @@ class DownloadHistory(Base): """ 分批删除指定时间之前的下载历史。 """ - ids = [ - row[0] - for row in db.query(cls.id) - .filter(cls.date < before_time) + ids = db.execute( + select(cls.id) + .where(cls.date < before_time) .order_by(cls.id.asc()) .limit(limit) - .all() - ] + ).scalars().all() if not ids: return 0 - return ( - db.query(cls) - .filter(cls.id.in_(ids)) - .delete(synchronize_session=False) + return execute_dml( + db, delete(cls).where(cls.id.in_(ids)), + execution_options={"synchronize_session": False}, ) @@ -407,19 +325,19 @@ class DownloadFiles(Base): id = get_id_column() # 下载器 - downloader = Column(String) + downloader: Mapped[Optional[str]] = mapped_column(String) # 下载任务Hash - download_hash = Column(String) + download_hash: Mapped[Optional[str]] = mapped_column(String) # 完整路径 - fullpath = Column(String) + fullpath: Mapped[Optional[str]] = mapped_column(String) # 保存路径 - savepath = Column(String, index=True) + savepath: Mapped[Optional[str]] = mapped_column(String, index=True) # 文件相对路径/名称 - filepath = Column(String) + filepath: Mapped[Optional[str]] = mapped_column(String) # 种子名称 - torrentname = Column(String) + torrentname: Mapped[Optional[str]] = mapped_column(String) # 状态 0-已删除 1-正常 - state = Column(Integer, nullable=False, default=1) + state: Mapped[int] = mapped_column(Integer, nullable=False, default=1) __table_args__ = ( Index('ix_downloadfiles_download_hash_state', 'download_hash', 'state'), @@ -429,43 +347,29 @@ class DownloadFiles(Base): @classmethod @db_query def get_by_hash(cls, db: Session, download_hash: str, state: Optional[int] = None): + statement = select(cls).where(cls.download_hash == download_hash) if state is not None: - return ( - db.query(cls) - .filter(cls.download_hash == download_hash, cls.state == state) - .all() - ) - else: - return db.query(cls).filter(cls.download_hash == download_hash).all() + statement = statement.where(cls.state == state) + return list(db.execute(statement).scalars().all()) @classmethod @db_query def get_by_fullpath(cls, db: Session, fullpath: str, all_files: bool = False): - if not all_files: - return ( - db.query(cls) - .filter(cls.fullpath == fullpath) - .order_by(cls.id.desc()) - .first() - ) - else: - return ( - db.query(cls) - .filter(cls.fullpath == fullpath) - .order_by(cls.id.desc()) - .all() - ) + result = db.execute( + select(cls).where(cls.fullpath == fullpath).order_by(cls.id.desc()) + ).scalars() + return list(result.all()) if all_files else result.first() @classmethod @db_query def get_by_savepath(cls, db: Session, savepath: str): - return db.query(cls).filter(cls.savepath == savepath).all() + return list(db.execute(select(cls).where(cls.savepath == savepath)).scalars().all()) @classmethod @db_update def delete_by_fullpath(cls, db: Session, fullpath: str): - db.query(cls).filter(cls.fullpath == fullpath, cls.state == 1).update( - {"state": 0} + db.execute( + update(cls).where(cls.fullpath == fullpath, cls.state == 1).values(state=0) ) @classmethod @@ -481,22 +385,19 @@ class DownloadFiles(Base): downloadfiles 没有时间字段,无法安全地按时间直接裁剪, 因此只清理明确失去父记录的孤儿数据。 """ - ids = [ - row[0] - for row in db.query(cls.id) + ids = db.execute( + select(cls.id) .outerjoin( DownloadHistory, DownloadHistory.download_hash == cls.download_hash, ) - .filter(DownloadHistory.id.is_(None)) + .where(DownloadHistory.id.is_(None)) .order_by(cls.id.asc()) .limit(limit) - .all() - ] + ).scalars().all() if not ids: return 0 - return ( - db.query(cls) - .filter(cls.id.in_(ids)) - .delete(synchronize_session=False) + return execute_dml( + db, delete(cls).where(cls.id.in_(ids)), + execution_options={"synchronize_session": False}, ) diff --git a/app/db/models/mediaserver.py b/app/db/models/mediaserver.py index 4a9d54be5..9d1bd7d20 100644 --- a/app/db/models/mediaserver.py +++ b/app/db/models/mediaserver.py @@ -1,13 +1,13 @@ from datetime import datetime -from typing import Optional, List +from typing import Any, List, Optional -from sqlalchemy import Column, Integer, String, JSON, Index, or_ +from sqlalchemy import Integer, String, JSON, Index, delete, or_ from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column -from app.db import db_query, db_update, get_id_column, async_db_query, Base -from app.db.models.media_identity import media_identity_constraint +from app.db import Base, async_db_query, db_query, db_update, execute_dml, get_id_column +from app.db.models._constraints import media_identity_constraint from app.schemas.types import MediaSource @@ -17,30 +17,30 @@ class MediaServerItem(Base): """ id = get_id_column() # 服务器类型 - server = Column(String) + server: Mapped[Optional[str]] = mapped_column(String) # 媒体库ID - library = Column(String) + library: Mapped[Optional[str]] = mapped_column(String) # ID - item_id = Column(String, index=True) + item_id: Mapped[Optional[str]] = mapped_column(String, index=True) # 类型 - item_type = Column(String) + item_type: Mapped[Optional[str]] = mapped_column(String) # 标题 - title = Column(String, index=True) + title: Mapped[Optional[str]] = mapped_column(String, index=True) # 原标题 - original_title = Column(String) + original_title: Mapped[Optional[str]] = mapped_column(String) # 年份 - year = Column(String) + year: Mapped[Optional[str]] = mapped_column(String) # 媒体数据源与原生ID - media_source = Column(String, index=True) - media_id = Column(String, index=True) + media_source: Mapped[Optional[str]] = mapped_column(String, index=True) + media_id: Mapped[Optional[str]] = mapped_column(String, index=True) # 路径 - path = Column(String) + path: Mapped[Optional[str]] = mapped_column(String) # 季集 - seasoninfo = Column(JSON, default=dict) + seasoninfo: Mapped[Optional[Any]] = mapped_column(JSON, default=dict) # 备注 - note = Column(JSON) + note: Mapped[Optional[Any]] = mapped_column(JSON) # 同步时间 - lst_mod_date = Column(String, default=datetime.now().strftime("%Y-%m-%d %H:%M:%S")) + lst_mod_date: Mapped[Optional[str]] = mapped_column(String, default=datetime.now().strftime("%Y-%m-%d %H:%M:%S")) __table_args__ = ( media_identity_constraint("mediaserveritem"), @@ -54,36 +54,46 @@ class MediaServerItem(Base): @classmethod @db_query def get_by_itemid(cls, db: Session, item_id: str): - return db.query(cls).filter(cls.item_id == item_id).first() + return db.execute(select(cls).where(cls.item_id == item_id)).scalars().first() @classmethod @db_query def get_by_server_itemid(cls, db: Session, server: str, item_id: str): - return db.query(cls).filter(cls.server == server, - cls.item_id == item_id).first() + return db.execute( + select(cls).where(cls.server == server, cls.item_id == item_id) + ).scalars().first() @classmethod @db_update def empty(cls, db: Session, server: Optional[str] = None): - if server is None: - db.query(cls).delete(synchronize_session=False) - else: - db.query(cls).filter(cls.server == server).delete(synchronize_session=False) + statement = delete(cls) + if server is not None: + statement = statement.where(cls.server == server) + db.execute(statement, execution_options={"synchronize_session": False}) @classmethod @db_update def delete_stale(cls, db: Session, server: str, sync_time: str): - return db.query(cls).filter(cls.server == server, - or_(cls.lst_mod_date.is_(None), - cls.lst_mod_date != sync_time)).delete(synchronize_session=False) + return execute_dml( + db, + delete(cls).where( + cls.server == server, + or_(cls.lst_mod_date.is_(None), cls.lst_mod_date != sync_time), + ), + execution_options={"synchronize_session": False}, + ) @classmethod @db_update def delete_excluded_servers(cls, db: Session, servers: List[str]): - if not servers: - return db.query(cls).delete(synchronize_session=False) - return db.query(cls).filter(or_(cls.server.is_(None), - ~cls.server.in_(servers))).delete(synchronize_session=False) + statement = delete(cls) + if servers: + statement = statement.where( + or_(cls.server.is_(None), ~cls.server.in_(servers)) + ) + return execute_dml( + db, statement, execution_options={"synchronize_session": False} + ) @classmethod @db_query @@ -91,26 +101,21 @@ class MediaServerItem(Base): cls, db: Session, media_source: MediaSource, media_id: str, mtype: str, ): """按规范媒体身份和类型查询媒体服务器条目。""" - return db.query(cls).filter( + return db.execute(select(cls).where( cls.media_source == str(media_source), cls.media_id == str(media_id), cls.item_type == mtype, - ).first() + )).scalars().first() @classmethod @db_query def exists_by_title(cls, db: Session, title: str, mtype: str, year: str): - if not mtype and not year: - return db.query(cls).filter(cls.title == title).first() - elif not year: - return db.query(cls).filter(cls.title == title, - cls.item_type == mtype).first() - elif not mtype: - return db.query(cls).filter(cls.title == title, - cls.year == str(year)).first() - return db.query(cls).filter(cls.title == title, - cls.item_type == mtype, - cls.year == str(year)).first() + statement = select(cls).where(cls.title == title) + if mtype: + statement = statement.where(cls.item_type == mtype) + if year: + statement = statement.where(cls.year == str(year)) + return db.execute(statement).scalars().first() @classmethod @async_db_query diff --git a/app/db/models/message.py b/app/db/models/message.py index ca342557a..1a3b707dd 100644 --- a/app/db/models/message.py +++ b/app/db/models/message.py @@ -1,10 +1,10 @@ -from typing import List, Optional +from typing import Any, List, Optional -from sqlalchemy import Column, Integer, String, JSON, Index, and_, or_, select +from sqlalchemy import Integer, String, JSON, Index, and_, delete, or_, select from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column -from app.db import db_query, db_update, Base, get_id_column, async_db_query +from app.db import Base, async_db_query, db_query, db_update, execute_dml, get_id_column class Message(Base): @@ -13,27 +13,27 @@ class Message(Base): """ id = get_id_column() # 消息渠道 - channel = Column(String) + channel: Mapped[Optional[str]] = mapped_column(String) # 消息来源 - source = Column(String) + source: Mapped[Optional[str]] = mapped_column(String) # 消息类型 - mtype = Column(String) + mtype: Mapped[Optional[str]] = mapped_column(String) # 标题 - title = Column(String) + title: Mapped[Optional[str]] = mapped_column(String) # 文本内容 - text = Column(String) + text: Mapped[Optional[str]] = mapped_column(String) # 图片 - image = Column(String) + image: Mapped[Optional[str]] = mapped_column(String) # 链接 - link = Column(String) + link: Mapped[Optional[str]] = mapped_column(String) # 用户ID - userid = Column(String) + userid: Mapped[Optional[str]] = mapped_column(String) # 登记时间 - reg_time = Column(String) + reg_time: Mapped[Optional[str]] = mapped_column(String) # 消息方向:0-接收息,1-发送消息 - action = Column(Integer) + action: Mapped[Optional[int]] = mapped_column(Integer) # 附件json - note = Column(JSON) + note: Mapped[Optional[Any]] = mapped_column(JSON) __table_args__ = ( Index('ix_message_reg_time_id', 'reg_time', 'id'), @@ -50,17 +50,16 @@ class Message(Base): @classmethod @db_query - def list_by_page(cls, db: Session, page: Optional[int] = 1, count: Optional[int] = 30) -> List["Message"]: + def list_by_page(cls, db: Session, page: int = 1, count: int = 30) -> List["Message"]: """ 分页获取消息记录。 """ - return ( - db.query(cls) + return list(db.execute( + select(cls) .order_by(cls.reg_time.desc(), cls.id.desc()) .offset((page - 1) * count) .limit(count) - .all() - ) + ).scalars().all()) @classmethod @db_query @@ -72,12 +71,14 @@ class Message(Base): :param source: 消息来源唯一标识 :return: 是否存在匹配记录 """ - return db.query(cls.id).filter(cls.source == source).first() is not None + return db.execute( + select(cls.id).where(cls.source == source).limit(1) + ).scalars().first() is not None @classmethod @async_db_query async def async_list_by_page( - cls, db: AsyncSession, page: Optional[int] = 1, count: Optional[int] = 30 + cls, db: AsyncSession, page: int = 1, count: int = 30 ) -> List["Message"]: """ 异步分页获取消息记录。 @@ -88,15 +89,15 @@ class Message(Base): .offset((page - 1) * count) .limit(count) ) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @async_db_query async def async_list_sent_by_page( cls, db: AsyncSession, - page: Optional[int] = 1, - count: Optional[int] = 30, + page: int = 1, + count: int = 30, all_clear_before: Optional[str] = None, system_clear_before: Optional[str] = None, media_clear_before: Optional[str] = None, @@ -129,7 +130,7 @@ class Message(Base): .offset((page - 1) * count) .limit(count) ) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_update @@ -142,18 +143,15 @@ class Message(Base): """ 分批删除指定时间之前的消息记录。 """ - ids = [ - row[0] - for row in db.query(cls.id) - .filter(cls.reg_time < before_time) + ids = db.execute( + select(cls.id) + .where(cls.reg_time < before_time) .order_by(cls.id.asc()) .limit(limit) - .all() - ] + ).scalars().all() if not ids: return 0 - return ( - db.query(cls) - .filter(cls.id.in_(ids)) - .delete(synchronize_session=False) + return execute_dml( + db, delete(cls).where(cls.id.in_(ids)), + execution_options={"synchronize_session": False}, ) diff --git a/app/db/models/passkey.py b/app/db/models/passkey.py index 971a94c37..94dc09709 100644 --- a/app/db/models/passkey.py +++ b/app/db/models/passkey.py @@ -1,6 +1,7 @@ -from sqlalchemy import Column, Integer, String, Boolean, DateTime, Text, select, ForeignKey +from typing import Optional +from sqlalchemy import Integer, String, Boolean, DateTime, Text, select, ForeignKey from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column from datetime import datetime from app.db import Base, db_query, db_update, async_db_query, async_db_update, get_id_column @@ -13,31 +14,33 @@ class PassKey(Base): # ID id = get_id_column() # 用户ID - user_id = Column(Integer, ForeignKey('user.id'), nullable=False, index=True) + user_id: Mapped[int] = mapped_column(Integer, ForeignKey('user.id'), nullable=False, index=True) # 凭证ID (credential_id) - credential_id = Column(String, nullable=False, unique=True, index=True) + credential_id: Mapped[str] = mapped_column(String, nullable=False, unique=True, index=True) # 凭证公钥 - public_key = Column(Text, nullable=False) + public_key: Mapped[str] = mapped_column(Text, nullable=False) # 签名计数器 - sign_count = Column(Integer, default=0) + sign_count: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 凭证名称(用户自定义) - name = Column(String, default="通行密钥") + name: Mapped[Optional[str]] = mapped_column(String, default="通行密钥") # AAGUID (Authenticator Attestation GUID) - aaguid = Column(String, nullable=True) + aaguid: Mapped[Optional[str]] = mapped_column(String, nullable=True) # 创建时间 - created_at = Column(DateTime, default=datetime.now) + created_at: Mapped[Optional[datetime]] = mapped_column(DateTime, default=datetime.now) # 最后使用时间 - last_used_at = Column(DateTime, nullable=True) + last_used_at: Mapped[Optional[datetime]] = mapped_column(DateTime, nullable=True) # 是否启用 - is_active = Column(Boolean, default=True) + is_active: Mapped[Optional[bool]] = mapped_column(Boolean, default=True) # 传输方式 (usb, nfc, ble, internal) - transports = Column(String, nullable=True) + transports: Mapped[Optional[str]] = mapped_column(String, nullable=True) @classmethod @db_query def get_by_user_id(cls, db: Session, user_id: int): """获取用户的所有PassKey""" - return db.query(cls).filter(cls.user_id == user_id, cls.is_active.is_(True)).all() + return list(db.execute( + select(cls).where(cls.user_id == user_id, cls.is_active.is_(True)) + ).scalars().all()) @classmethod @async_db_query @@ -46,13 +49,15 @@ class PassKey(Base): result = await db.execute( select(cls).filter(cls.user_id == user_id, cls.is_active.is_(True)) ) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query def get_by_credential_id(cls, db: Session, credential_id: str): """根据凭证ID获取PassKey""" - return db.query(cls).filter(cls.credential_id == credential_id, cls.is_active.is_(True)).first() + return db.execute( + select(cls).where(cls.credential_id == credential_id, cls.is_active.is_(True)) + ).scalars().first() @classmethod @async_db_query @@ -67,7 +72,7 @@ class PassKey(Base): @db_query def get_by_id(cls, db: Session, passkey_id: int): """根据ID获取PassKey""" - return db.query(cls).filter(cls.id == passkey_id).first() + return db.execute(select(cls).where(cls.id == passkey_id)).scalars().first() @classmethod @async_db_query @@ -82,10 +87,9 @@ class PassKey(Base): @db_update def delete_by_id(cls, db: Session, passkey_id: int, user_id: int): """删除指定用户的PassKey""" - passkey = db.query(cls).filter( - cls.id == passkey_id, - cls.user_id == user_id - ).first() + passkey = db.execute( + select(cls).where(cls.id == passkey_id, cls.user_id == user_id) + ).scalars().first() if passkey: passkey.delete(db, passkey.id) return True diff --git a/app/db/models/plugindata.py b/app/db/models/plugindata.py index bfd882b87..ea70a5e39 100644 --- a/app/db/models/plugindata.py +++ b/app/db/models/plugindata.py @@ -1,6 +1,7 @@ -from sqlalchemy import Column, String, JSON, Index, select +from typing import Any, Optional +from sqlalchemy import String, JSON, Index, delete, select from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column from app.db import ( db_query, @@ -16,9 +17,9 @@ class PluginData(Base): 插件数据表 """ id = get_id_column() - plugin_id = Column(String, nullable=False) - key = Column(String, nullable=False) - value = Column(JSON) + plugin_id: Mapped[str] = mapped_column(String, nullable=False) + key: Mapped[str] = mapped_column(String, nullable=False) + value: Mapped[Optional[Any]] = mapped_column(JSON) __table_args__ = ( Index('ix_plugindata_plugin_id_key', 'plugin_id', 'key'), @@ -27,18 +28,20 @@ class PluginData(Base): @classmethod @db_query def get_plugin_data(cls, db: Session, plugin_id: str): - return db.query(cls).filter(cls.plugin_id == plugin_id).all() + return list(db.execute(select(cls).where(cls.plugin_id == plugin_id)).scalars().all()) @classmethod @async_db_query async def async_get_plugin_data(cls, db: AsyncSession, plugin_id: str): result = await db.execute(select(cls).where(cls.plugin_id == plugin_id)) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query def get_plugin_data_by_key(cls, db: Session, plugin_id: str, key: str): - return db.query(cls).filter(cls.plugin_id == plugin_id, cls.key == key).first() + return db.execute( + select(cls).where(cls.plugin_id == plugin_id, cls.key == key) + ).scalars().first() @classmethod @async_db_query @@ -53,17 +56,17 @@ class PluginData(Base): @classmethod @db_update def del_plugin_data_by_key(cls, db: Session, plugin_id: str, key: str): - db.query(cls).filter(cls.plugin_id == plugin_id, cls.key == key).delete() + db.execute(delete(cls).where(cls.plugin_id == plugin_id, cls.key == key)) @classmethod @db_update def del_plugin_data(cls, db: Session, plugin_id: str): - db.query(cls).filter(cls.plugin_id == plugin_id).delete() + db.execute(delete(cls).where(cls.plugin_id == plugin_id)) @classmethod @db_query def get_plugin_data_by_plugin_id(cls, db: Session, plugin_id: str): - return db.query(cls).filter(cls.plugin_id == plugin_id).all() + return list(db.execute(select(cls).where(cls.plugin_id == plugin_id)).scalars().all()) @classmethod @async_db_query @@ -71,4 +74,4 @@ class PluginData(Base): cls, db: AsyncSession, plugin_id: str ): result = await db.execute(select(cls).where(cls.plugin_id == plugin_id)) - return result.scalars().all() + return list(result.scalars().all()) diff --git a/app/db/models/site.py b/app/db/models/site.py index 3c456e5f0..7bd3f22d6 100644 --- a/app/db/models/site.py +++ b/app/db/models/site.py @@ -1,8 +1,9 @@ +from typing import Any, Optional from datetime import datetime -from sqlalchemy import Boolean, Column, Integer, String, JSON, select, delete +from sqlalchemy import Boolean, Integer, String, JSON, select, delete from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column from app.db import db_query, db_update, Base, async_db_query, async_db_update, get_id_column @@ -13,52 +14,52 @@ class Site(Base): """ id = get_id_column() # 站点名 - name = Column(String, nullable=False) + name: Mapped[str] = mapped_column(String, nullable=False) # 域名Key - domain = Column(String, index=True) + domain: Mapped[Optional[str]] = mapped_column(String, index=True) # 站点地址 - url = Column(String, nullable=False) + url: Mapped[str] = mapped_column(String, nullable=False) # 站点优先级 - pri = Column(Integer, default=1) + pri: Mapped[Optional[int]] = mapped_column(Integer, default=1) # RSS地址,未启用 - rss = Column(String) + rss: Mapped[Optional[str]] = mapped_column(String) # Cookie - cookie = Column(String) + cookie: Mapped[Optional[str]] = mapped_column(String) # User-Agent - ua = Column(String) + ua: Mapped[Optional[str]] = mapped_column(String) # ApiKey - apikey = Column(String) + apikey: Mapped[Optional[str]] = mapped_column(String) # Token - token = Column(String) + token: Mapped[Optional[str]] = mapped_column(String) # 是否使用代理 0-否,1-是 - proxy = Column(Integer) + proxy: Mapped[Optional[int]] = mapped_column(Integer) # 过滤规则 - filter = Column(String) + filter: Mapped[Optional[str]] = mapped_column(String) # 是否渲染 - render = Column(Integer) + render: Mapped[Optional[int]] = mapped_column(Integer) # 是否公开站点 - public = Column(Integer) + public: Mapped[Optional[int]] = mapped_column(Integer) # 附加信息 - note = Column(JSON) + note: Mapped[Optional[Any]] = mapped_column(JSON) # 流控单位周期 - limit_interval = Column(Integer, default=0) + limit_interval: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 流控次数 - limit_count = Column(Integer, default=0) + limit_count: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 流控间隔 - limit_seconds = Column(Integer, default=0) + limit_seconds: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 超时时间 - timeout = Column(Integer, default=15) + timeout: Mapped[Optional[int]] = mapped_column(Integer, default=15) # 是否启用 - is_active = Column(Boolean(), default=True) + is_active: Mapped[Optional[bool]] = mapped_column(Boolean(), default=True) # 创建时间 - lst_mod_date = Column(String, default=datetime.now().strftime("%Y-%m-%d %H:%M:%S")) + lst_mod_date: Mapped[Optional[str]] = mapped_column(String, default=datetime.now().strftime("%Y-%m-%d %H:%M:%S")) # 下载器 - downloader = Column(String) + downloader: Mapped[Optional[str]] = mapped_column(String) @classmethod @db_query def get_by_domain(cls, db: Session, domain: str): - return db.query(cls).filter(cls.domain == domain).first() + return db.execute(select(cls).where(cls.domain == domain)).scalars().first() @classmethod @async_db_query @@ -75,34 +76,34 @@ class Site(Base): @classmethod @db_query def get_actives(cls, db: Session): - return db.query(cls).filter(cls.is_active).all() + return list(db.execute(select(cls).where(cls.is_active.is_(True))).scalars().all()) @classmethod @async_db_query async def async_get_actives(cls, db: AsyncSession): - result = await db.execute(select(cls).where(cls.is_active)) - return result.scalars().all() + result = await db.execute(select(cls).where(cls.is_active.is_(True))) + return list(result.scalars().all()) @classmethod @db_query def list_order_by_pri(cls, db: Session): - return db.query(cls).order_by(cls.pri).all() + return list(db.execute(select(cls).order_by(cls.pri)).scalars().all()) @classmethod @async_db_query async def async_list_order_by_pri(cls, db: AsyncSession): result = await db.execute(select(cls).order_by(cls.pri)) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query def get_domains_by_ids(cls, db: Session, ids: list): - return [r[0] for r in db.query(cls.domain).filter(cls.id.in_(ids)).all()] + return list(db.execute(select(cls.domain).where(cls.id.in_(ids))).scalars().all()) @classmethod @db_update def reset(cls, db: Session): - db.query(cls).delete() + db.execute(delete(cls)) @classmethod @async_db_update diff --git a/app/db/models/siteicon.py b/app/db/models/siteicon.py index 05f7593d4..2237b0350 100644 --- a/app/db/models/siteicon.py +++ b/app/db/models/siteicon.py @@ -1,6 +1,7 @@ -from sqlalchemy import Column, String, select +from typing import Optional +from sqlalchemy import String, select from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column from app.db import db_query, Base, get_id_column, async_db_query @@ -11,18 +12,18 @@ class SiteIcon(Base): """ id = get_id_column() # 站点名称 - name = Column(String, nullable=False) + name: Mapped[str] = mapped_column(String, nullable=False) # 域名Key - domain = Column(String, index=True) + domain: Mapped[Optional[str]] = mapped_column(String, index=True) # 图标地址 - url = Column(String, nullable=False) + url: Mapped[str] = mapped_column(String, nullable=False) # 图标Base64 - base64 = Column(String) + base64: Mapped[Optional[str]] = mapped_column(String) @classmethod @db_query def get_by_domain(cls, db: Session, domain: str): - return db.query(cls).filter(cls.domain == domain).first() + return db.execute(select(cls).where(cls.domain == domain)).scalars().first() @classmethod @async_db_query diff --git a/app/db/models/sitestatistic.py b/app/db/models/sitestatistic.py index f2874cffb..a22f2c2d3 100644 --- a/app/db/models/sitestatistic.py +++ b/app/db/models/sitestatistic.py @@ -1,8 +1,9 @@ +from typing import Any, Optional from datetime import datetime -from sqlalchemy import Column, Integer, String, JSON, select +from sqlalchemy import Integer, String, JSON, delete, select from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column from app.db import db_query, db_update, get_id_column, Base, async_db_query @@ -13,24 +14,24 @@ class SiteStatistic(Base): """ id = get_id_column() # 域名Key - domain = Column(String, index=True) + domain: Mapped[Optional[str]] = mapped_column(String, index=True) # 成功次数 - success = Column(Integer) + success: Mapped[Optional[int]] = mapped_column(Integer) # 失败次数 - fail = Column(Integer) + fail: Mapped[Optional[int]] = mapped_column(Integer) # 平均耗时 秒 - seconds = Column(Integer) + seconds: Mapped[Optional[int]] = mapped_column(Integer) # 最后一次访问状态 0-成功 1-失败 - lst_state = Column(Integer) + lst_state: Mapped[Optional[int]] = mapped_column(Integer) # 最后访问时间 - lst_mod_date = Column(String, default=datetime.now().strftime("%Y-%m-%d %H:%M:%S")) + lst_mod_date: Mapped[Optional[str]] = mapped_column(String, default=datetime.now().strftime("%Y-%m-%d %H:%M:%S")) # 耗时记录 Json - note = Column(JSON) + note: Mapped[Optional[Any]] = mapped_column(JSON) @classmethod @db_query def get_by_domain(cls, db: Session, domain: str): - return db.query(cls).filter(cls.domain == domain).first() + return db.execute(select(cls).where(cls.domain == domain)).scalars().first() @classmethod @async_db_query @@ -41,4 +42,4 @@ class SiteStatistic(Base): @classmethod @db_update def reset(cls, db: Session): - db.query(cls).delete() + db.execute(delete(cls)) diff --git a/app/db/models/siteuserdata.py b/app/db/models/siteuserdata.py index c35d910a5..e70dc2315 100644 --- a/app/db/models/siteuserdata.py +++ b/app/db/models/siteuserdata.py @@ -1,11 +1,11 @@ from datetime import datetime -from typing import Optional +from typing import Any, Optional -from sqlalchemy import Column, Integer, String, Float, JSON, Index, func, or_, select +from sqlalchemy import Integer, String, Float, JSON, Index, delete, func, or_, select from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column -from app.db import db_query, db_update, Base, get_id_column, async_db_query +from app.db import Base, async_db_query, db_query, db_update, execute_dml, get_id_column class SiteUserData(Base): @@ -14,45 +14,45 @@ class SiteUserData(Base): """ id = get_id_column() # 站点域名 - domain = Column(String) + domain: Mapped[Optional[str]] = mapped_column(String) # 站点名称 - name = Column(String) + name: Mapped[Optional[str]] = mapped_column(String) # 用户名 - username = Column(String) + username: Mapped[Optional[str]] = mapped_column(String) # 用户ID - userid = Column(String) + userid: Mapped[Optional[str]] = mapped_column(String) # 用户等级 - user_level = Column(String) + user_level: Mapped[Optional[str]] = mapped_column(String) # 加入时间 - join_at = Column(String) + join_at: Mapped[Optional[str]] = mapped_column(String) # 积分 - bonus = Column(Float, default=0) + bonus: Mapped[Optional[float]] = mapped_column(Float, default=0) # 上传量 - upload = Column(Float, default=0) + upload: Mapped[Optional[float]] = mapped_column(Float, default=0) # 下载量 - download = Column(Float, default=0) + download: Mapped[Optional[float]] = mapped_column(Float, default=0) # 分享率 - ratio = Column(Float, default=0) + ratio: Mapped[Optional[float]] = mapped_column(Float, default=0) # 做种数 - seeding = Column(Float, default=0) + seeding: Mapped[Optional[float]] = mapped_column(Float, default=0) # 下载数 - leeching = Column(Float, default=0) + leeching: Mapped[Optional[float]] = mapped_column(Float, default=0) # 做种体积 - seeding_size = Column(Float, default=0) + seeding_size: Mapped[Optional[float]] = mapped_column(Float, default=0) # 下载体积 - leeching_size = Column(Float, default=0) + leeching_size: Mapped[Optional[float]] = mapped_column(Float, default=0) # 做种人数, 种子大小 JSON - seeding_info = Column(JSON, default=dict) + seeding_info: Mapped[Optional[Any]] = mapped_column(JSON, default=dict) # 未读消息 - message_unread = Column(Integer, default=0) + message_unread: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 未读消息内容 JSON - message_unread_contents = Column(JSON, default=list) + message_unread_contents: Mapped[Optional[Any]] = mapped_column(JSON, default=list) # 错误信息 - err_msg = Column(String) + err_msg: Mapped[Optional[str]] = mapped_column(String) # 更新日期 - updated_day = Column(String, default=datetime.now().strftime('%Y-%m-%d')) + updated_day: Mapped[Optional[str]] = mapped_column(String, default=datetime.now().strftime('%Y-%m-%d')) # 更新时间 - updated_time = Column(String, default=datetime.now().strftime('%H:%M:%S')) + updated_time: Mapped[Optional[str]] = mapped_column(String, default=datetime.now().strftime('%H:%M:%S')) __table_args__ = ( Index('ix_siteuserdata_updated_day_id', 'updated_day', 'id'), @@ -62,14 +62,13 @@ class SiteUserData(Base): @classmethod @db_query def get_by_domain(cls, db: Session, domain: str, workdate: Optional[str] = None, worktime: Optional[str] = None): + statement = select(cls).where(cls.domain == domain) if workdate and worktime: - return db.query(cls).filter(cls.domain == domain, - cls.updated_day == workdate, - cls.updated_time == worktime).all() + statement = statement.where(cls.updated_day == workdate, + cls.updated_time == worktime) elif workdate: - return db.query(cls).filter(cls.domain == domain, - cls.updated_day == workdate).all() - return db.query(cls).filter(cls.domain == domain).all() + statement = statement.where(cls.updated_day == workdate) + return list(db.execute(statement).scalars().all()) @classmethod @async_db_query @@ -80,12 +79,12 @@ class SiteUserData(Base): elif workdate: query = query.filter(cls.updated_day == workdate) result = await db.execute(query) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query def get_by_date(cls, db: Session, date: str): - return db.query(cls).filter(cls.updated_day == date).all() + return list(db.execute(select(cls).where(cls.updated_day == date)).scalars().all()) @classmethod @db_query @@ -94,21 +93,23 @@ class SiteUserData(Base): 获取各站点最新一天的数据 """ subquery = ( - db.query( + select( cls.domain, func.max(cls.updated_day).label('latest_update_day') ) + .where(or_(cls.err_msg.is_(None), cls.err_msg == "")) .group_by(cls.domain) - .filter(or_(cls.err_msg.is_(None), cls.err_msg == "")) .subquery() ) # 主查询:按 domain 和 updated_day 获取最新的记录 - return db.query(cls).join( - subquery, - (cls.domain == subquery.c.domain) & - (cls.updated_day == subquery.c.latest_update_day) - ).order_by(cls.updated_time.desc()).all() + return list(db.execute( + select(cls).join( + subquery, + (cls.domain == subquery.c.domain) & + (cls.updated_day == subquery.c.latest_update_day) + ).order_by(cls.updated_time.desc()) + ).scalars().all()) @classmethod @async_db_query @@ -133,7 +134,7 @@ class SiteUserData(Base): (cls.domain == subquery.c.domain) & (cls.updated_day == subquery.c.latest_update_day) ).order_by(cls.updated_time.desc())) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_update @@ -146,19 +147,15 @@ class SiteUserData(Base): """ 分批删除指定日期之前的站点用户快照。 """ - ids = [ - row[0] - for row in db.query(cls.id) - .filter(cls.updated_day < before_day) + ids = db.execute( + select(cls.id) + .where(cls.updated_day < before_day) .order_by(cls.id.asc()) .limit(limit) - .all() - ] + ).scalars().all() if not ids: return 0 - deleted = ( - db.query(cls) - .filter(cls.id.in_(ids)) - .delete(synchronize_session=False) + return execute_dml( + db, delete(cls).where(cls.id.in_(ids)), + execution_options={"synchronize_session": False}, ) - return deleted diff --git a/app/db/models/subscribe.py b/app/db/models/subscribe.py index 814161475..0b5b0717e 100644 --- a/app/db/models/subscribe.py +++ b/app/db/models/subscribe.py @@ -1,12 +1,12 @@ import time -from typing import Optional +from typing import Any, Optional -from sqlalchemy import Column, Integer, String, Float, JSON, Index, or_, select +from sqlalchemy import Integer, String, Float, JSON, Index, delete, or_, select from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column from app.db import db_query, db_update, get_id_column, Base, async_db_query, async_db_update -from app.db.models.media_identity import media_identity_constraint +from app.db.models._constraints import media_identity_constraint from app.schemas.types import MUSIC_ENTITY_RECORDING, MediaSource @@ -16,101 +16,101 @@ class Subscribe(Base): """ id = get_id_column() # 标题 - name = Column(String, nullable=False, index=True) + name: Mapped[str] = mapped_column(String, nullable=False, index=True) # 年份 - year = Column(String) + year: Mapped[Optional[str]] = mapped_column(String) # 类型 - type = Column(String) + type: Mapped[Optional[str]] = mapped_column(String) # 搜索关键字 - keyword = Column(String) - media_source = Column(String, index=True) - media_id = Column(String, index=True) + keyword: Mapped[Optional[str]] = mapped_column(String) + media_source: Mapped[Optional[str]] = mapped_column(String, index=True) + media_id: Mapped[Optional[str]] = mapped_column(String, index=True) # 音乐实体类型:recording 单曲、album 专辑 - music_type = Column(String) + music_type: Mapped[Optional[str]] = mapped_column(String) # 专辑预期总曲目数,供整专资源完整性判断 - total_tracks = Column(Integer) + total_tracks: Mapped[Optional[int]] = mapped_column(Integer) # 季号 - season = Column(Integer) + season: Mapped[Optional[int]] = mapped_column(Integer) # 海报 - poster = Column(String) + poster: Mapped[Optional[str]] = mapped_column(String) # 背景图 - backdrop = Column(String) + backdrop: Mapped[Optional[str]] = mapped_column(String) # 评分,float - vote = Column(Float) + vote: Mapped[Optional[float]] = mapped_column(Float) # 简介 - description = Column(String) + description: Mapped[Optional[str]] = mapped_column(String) # 过滤规则 - filter = Column(String) + filter: Mapped[Optional[str]] = mapped_column(String) # 包含 - include = Column(String) + include: Mapped[Optional[str]] = mapped_column(String) # 排除 - exclude = Column(String) + exclude: Mapped[Optional[str]] = mapped_column(String) # 质量 - quality = Column(String) + quality: Mapped[Optional[str]] = mapped_column(String) # 分辨率 - resolution = Column(String) + resolution: Mapped[Optional[str]] = mapped_column(String) # 特效 - effect = Column(String) + effect: Mapped[Optional[str]] = mapped_column(String) # 音乐音质等级:hires/lossless/lossy,可用正则组合 - audio_quality = Column(String) + audio_quality: Mapped[Optional[str]] = mapped_column(String) # 音频格式,可用正则组合 - audio_format = Column(String) + audio_format: Mapped[Optional[str]] = mapped_column(String) # 最低码率(bps) - min_bitrate = Column(Integer) + min_bitrate: Mapped[Optional[int]] = mapped_column(Integer) # 最低位深(bit) - min_bit_depth = Column(Integer) + min_bit_depth: Mapped[Optional[int]] = mapped_column(Integer) # 最低采样率(Hz) - min_sample_rate = Column(Integer) + min_sample_rate: Mapped[Optional[int]] = mapped_column(Integer) # 总集数 - total_episode = Column(Integer) + total_episode: Mapped[Optional[int]] = mapped_column(Integer) # 开始集数 - start_episode = Column(Integer) + start_episode: Mapped[Optional[int]] = mapped_column(Integer) # 缺失集数 - lack_episode = Column(Integer) + lack_episode: Mapped[Optional[int]] = mapped_column(Integer) # 附加信息 - note = Column(JSON) + note: Mapped[Optional[Any]] = mapped_column(JSON) # 状态:N-新建 R-订阅中 P-待定 S-暂停 - state = Column(String, nullable=False, index=True, default='N') + state: Mapped[str] = mapped_column(String, nullable=False, index=True, default='N') # 最后更新时间 - last_update = Column(String) + last_update: Mapped[Optional[str]] = mapped_column(String) # 创建时间 - date = Column(String) + date: Mapped[Optional[str]] = mapped_column(String) # 订阅用户 - username = Column(String, index=True) + username: Mapped[Optional[str]] = mapped_column(String, index=True) # 订阅站点 - sites = Column(JSON, default=list) + sites: Mapped[Optional[Any]] = mapped_column(JSON, default=list) # 下载器 - downloader = Column(String) + downloader: Mapped[Optional[str]] = mapped_column(String) # 是否洗版 - best_version = Column(Integer, default=0) + best_version: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 是否只洗全集整包,开启后电视剧洗版不按单集下载 - best_version_full = Column(Integer, default=0) + best_version_full: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 当前优先级 - current_priority = Column(Integer) + current_priority: Mapped[Optional[int]] = mapped_column(Integer) # 当前音乐版本格式 - current_audio_format = Column(String) + current_audio_format: Mapped[Optional[str]] = mapped_column(String) # 当前音乐版本码率(bps) - current_bitrate = Column(Integer) + current_bitrate: Mapped[Optional[int]] = mapped_column(Integer) # 当前音乐版本位深(bit) - current_bit_depth = Column(Integer) + current_bit_depth: Mapped[Optional[int]] = mapped_column(Integer) # 当前音乐版本采样率(Hz) - current_sample_rate = Column(Integer) + current_sample_rate: Mapped[Optional[int]] = mapped_column(Integer) # 洗版时已下载剧集的优先级状态,格式:{"1": 90, "2": 100} - episode_priority = Column(JSON) + episode_priority: Mapped[Optional[Any]] = mapped_column(JSON) # 保存路径 - save_path = Column(String) + save_path: Mapped[Optional[str]] = mapped_column(String) # 是否使用 imdbid 搜索 - search_imdbid = Column(Integer, default=0) + search_imdbid: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 是否手动修改过总集数 0否 1是 - manual_total_episode = Column(Integer, default=0) + manual_total_episode: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 自定义识别词 - custom_words = Column(String) + custom_words: Mapped[Optional[str]] = mapped_column(String) # 自定义媒体类别 - media_category = Column(String) + media_category: Mapped[Optional[str]] = mapped_column(String) # 过滤规则组 - filter_groups = Column(JSON, default=list) + filter_groups: Mapped[Optional[Any]] = mapped_column(JSON, default=list) # 选择的剧集组 - episode_group = Column(String) + episode_group: Mapped[Optional[str]] = mapped_column(String) __table_args__ = ( media_identity_constraint("subscribe"), @@ -152,11 +152,11 @@ class Subscribe(Base): ) if condition is None: return None - query = db.query(cls).filter(condition) + statement = select(cls).where(condition) if season is not None: - query = query.filter(cls.season == season) - query = query.filter(cls.episode_group == episode_group) - return query.first() + statement = statement.where(cls.season == season) + statement = statement.where(cls.episode_group == episode_group) + return db.execute(statement).scalars().first() @classmethod @async_db_query @@ -197,11 +197,11 @@ class Subscribe(Base): ) if condition is None: return None - query = db.query(cls).filter(cls.username == username, condition) + statement = select(cls).where(cls.username == username, condition) if season is not None: - query = query.filter(cls.season == season) - query = query.filter(cls.episode_group == episode_group) - return query.first() + statement = statement.where(cls.season == season) + statement = statement.where(cls.episode_group == episode_group) + return db.execute(statement).scalars().first() @classmethod @async_db_query @@ -232,11 +232,11 @@ class Subscribe(Base): @db_query def get_by_state(cls, db: Session, state: str): # 如果 state 为空或 None,返回所有订阅 - if not state: - return db.query(cls).all() - else: + statement = select(cls) + if state: # 如果传入的状态不为空,拆分成多个状态 - return db.query(cls).filter(cls.state.in_(state.split(','))).all() + statement = statement.where(cls.state.in_(state.split(','))) + return list(db.execute(statement).scalars().all()) @classmethod @async_db_query @@ -249,15 +249,15 @@ class Subscribe(Base): result = await db.execute( select(cls).filter(cls.state.in_(state.split(','))) ) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query def get_by_title(cls, db: Session, title: str, season: Optional[int] = None): + statement = select(cls).where(cls.name == title) if season is not None: - return db.query(cls).filter(cls.name == title, - cls.season == season).first() - return db.query(cls).filter(cls.name == title).first() + statement = statement.where(cls.season == season) + return db.execute(statement).scalars().first() @classmethod @async_db_query @@ -286,7 +286,7 @@ class Subscribe(Base): result = await db.execute( select(cls).filter(cls.name == title) ) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query @@ -302,7 +302,7 @@ class Subscribe(Base): ) if condition is None: return [] - return db.query(cls).filter(condition).all() + return list(db.execute(select(cls).where(condition)).scalars().all()) @classmethod @async_db_query @@ -319,7 +319,7 @@ class Subscribe(Base): if condition is None: return [] result = await db.execute(select(cls).filter(condition)) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query @@ -336,10 +336,10 @@ class Subscribe(Base): ) if condition is None: return None - query = db.query(cls).filter(condition, cls.type == type) + statement = select(cls).where(condition, cls.type == type) if season is not None: - query = query.filter(cls.season == season) - return query.first() + statement = statement.where(cls.season == season) + return db.execute(statement).scalars().first() @classmethod @async_db_query @@ -368,13 +368,14 @@ class Subscribe(Base): season: Optional[int] = None, ) -> bool: """按规范媒体身份删除订阅。""" - query = db.query(type(self)).filter( - type(self).media_source == media_source, - type(self).media_id == str(media_id), + model = type(self) + statement = delete(model).where( + model.media_source == media_source, + model.media_id == str(media_id), ) if season is not None: - query = query.filter(type(self).season == season) - query.delete(synchronize_session=False) + statement = statement.where(model.season == season) + db.execute(statement, execution_options={"synchronize_session": False}) return True @async_db_update @@ -394,20 +395,12 @@ class Subscribe(Base): @classmethod @db_query def list_by_username(cls, db: Session, username: str, state: Optional[str] = None, mtype: Optional[str] = None): + statement = select(cls).where(cls.username == username) + if state: + statement = statement.where(cls.state == state) if mtype: - if state: - return db.query(cls).filter(cls.state == state, - cls.username == username, - cls.type == mtype).all() - else: - return db.query(cls).filter(cls.username == username, - cls.type == mtype).all() - else: - if state: - return db.query(cls).filter(cls.state == state, - cls.username == username).all() - else: - return db.query(cls).filter(cls.username == username).all() + statement = statement.where(cls.type == mtype) + return list(db.execute(statement).scalars().all()) @classmethod @async_db_query @@ -431,16 +424,18 @@ class Subscribe(Base): result = await db.execute( select(cls).filter(cls.username == username) ) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query def list_by_type(cls, db: Session, mtype: str, days: int): - return db.query(cls) \ - .filter(cls.type == mtype, - cls.date >= time.strftime("%Y-%m-%d %H:%M:%S", - time.localtime(time.time() - 86400 * int(days))) - ).all() + return list(db.execute( + select(cls).where( + cls.type == mtype, + cls.date >= time.strftime("%Y-%m-%d %H:%M:%S", + time.localtime(time.time() - 86400 * int(days))) + ) + ).scalars().all()) @classmethod @async_db_query @@ -452,4 +447,4 @@ class Subscribe(Base): time.localtime(time.time() - 86400 * int(days))) ) ) - return result.scalars().all() + return list(result.scalars().all()) diff --git a/app/db/models/subscribehistory.py b/app/db/models/subscribehistory.py index f0ecab1cb..be3bfff23 100644 --- a/app/db/models/subscribehistory.py +++ b/app/db/models/subscribehistory.py @@ -1,11 +1,11 @@ -from typing import Optional +from typing import Any, Optional -from sqlalchemy import Column, Integer, String, Float, JSON, Index, or_, select +from sqlalchemy import Integer, String, Float, JSON, Index, or_, select from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column from app.db import db_query, Base, get_id_column, async_db_query -from app.db.models.media_identity import media_identity_constraint +from app.db.models._constraints import media_identity_constraint from app.schemas.types import MUSIC_ENTITY_RECORDING, MediaSource @@ -15,89 +15,89 @@ class SubscribeHistory(Base): """ id = get_id_column() # 标题 - name = Column(String, nullable=False, index=True) + name: Mapped[str] = mapped_column(String, nullable=False, index=True) # 年份 - year = Column(String) + year: Mapped[Optional[str]] = mapped_column(String) # 类型 - type = Column(String) + type: Mapped[Optional[str]] = mapped_column(String) # 搜索关键字 - keyword = Column(String) - media_source = Column(String, index=True) - media_id = Column(String, index=True) + keyword: Mapped[Optional[str]] = mapped_column(String) + media_source: Mapped[Optional[str]] = mapped_column(String, index=True) + media_id: Mapped[Optional[str]] = mapped_column(String, index=True) # 音乐实体类型:recording 单曲、album 专辑 - music_type = Column(String) + music_type: Mapped[Optional[str]] = mapped_column(String) # 专辑预期总曲目数 - total_tracks = Column(Integer) + total_tracks: Mapped[Optional[int]] = mapped_column(Integer) # 季号 - season = Column(Integer) + season: Mapped[Optional[int]] = mapped_column(Integer) # 海报 - poster = Column(String) + poster: Mapped[Optional[str]] = mapped_column(String) # 背景图 - backdrop = Column(String) + backdrop: Mapped[Optional[str]] = mapped_column(String) # 评分,float - vote = Column(Float) + vote: Mapped[Optional[float]] = mapped_column(Float) # 简介 - description = Column(String) + description: Mapped[Optional[str]] = mapped_column(String) # 过滤规则 - filter = Column(String) + filter: Mapped[Optional[str]] = mapped_column(String) # 包含 - include = Column(String) + include: Mapped[Optional[str]] = mapped_column(String) # 排除 - exclude = Column(String) + exclude: Mapped[Optional[str]] = mapped_column(String) # 质量 - quality = Column(String) + quality: Mapped[Optional[str]] = mapped_column(String) # 分辨率 - resolution = Column(String) + resolution: Mapped[Optional[str]] = mapped_column(String) # 特效 - effect = Column(String) + effect: Mapped[Optional[str]] = mapped_column(String) # 音乐音质等级:hires/lossless/lossy,可用正则组合 - audio_quality = Column(String) + audio_quality: Mapped[Optional[str]] = mapped_column(String) # 音频格式,可用正则组合 - audio_format = Column(String) + audio_format: Mapped[Optional[str]] = mapped_column(String) # 最低码率(bps) - min_bitrate = Column(Integer) + min_bitrate: Mapped[Optional[int]] = mapped_column(Integer) # 最低位深(bit) - min_bit_depth = Column(Integer) + min_bit_depth: Mapped[Optional[int]] = mapped_column(Integer) # 最低采样率(Hz) - min_sample_rate = Column(Integer) + min_sample_rate: Mapped[Optional[int]] = mapped_column(Integer) # 总集数 - total_episode = Column(Integer) + total_episode: Mapped[Optional[int]] = mapped_column(Integer) # 开始集数 - start_episode = Column(Integer) + start_episode: Mapped[Optional[int]] = mapped_column(Integer) # 订阅完成时间 - date = Column(String) + date: Mapped[Optional[str]] = mapped_column(String) # 订阅用户 - username = Column(String) + username: Mapped[Optional[str]] = mapped_column(String) # 订阅站点 - sites = Column(JSON) + sites: Mapped[Optional[Any]] = mapped_column(JSON) # 是否洗版 - best_version = Column(Integer, default=0) + best_version: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 是否只洗全集整包,开启后电视剧洗版不按单集下载 - best_version_full = Column(Integer, default=0) + best_version_full: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 完成时的整体优先级 - current_priority = Column(Integer) + current_priority: Mapped[Optional[int]] = mapped_column(Integer) # 完成时的音乐格式 - current_audio_format = Column(String) + current_audio_format: Mapped[Optional[str]] = mapped_column(String) # 完成时的音乐码率(bps) - current_bitrate = Column(Integer) + current_bitrate: Mapped[Optional[int]] = mapped_column(Integer) # 完成时的音乐位深(bit) - current_bit_depth = Column(Integer) + current_bit_depth: Mapped[Optional[int]] = mapped_column(Integer) # 完成时的音乐采样率(Hz) - current_sample_rate = Column(Integer) + current_sample_rate: Mapped[Optional[int]] = mapped_column(Integer) # 洗版时已下载剧集的优先级状态,格式:{"1": 90, "2": 100} - episode_priority = Column(JSON) + episode_priority: Mapped[Optional[Any]] = mapped_column(JSON) # 保存路径 - save_path = Column(String) + save_path: Mapped[Optional[str]] = mapped_column(String) # 是否使用 imdbid 搜索 - search_imdbid = Column(Integer, default=0) + search_imdbid: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 自定义识别词 - custom_words = Column(String) + custom_words: Mapped[Optional[str]] = mapped_column(String) # 自定义媒体类别 - media_category = Column(String) + media_category: Mapped[Optional[str]] = mapped_column(String) # 过滤规则组 - filter_groups = Column(JSON, default=list) + filter_groups: Mapped[Optional[Any]] = mapped_column(JSON, default=list) # 剧集组 - episode_group = Column(String) + episode_group: Mapped[Optional[str]] = mapped_column(String) __table_args__ = ( media_identity_constraint("subscribehistory"), @@ -107,16 +107,18 @@ class SubscribeHistory(Base): @classmethod @db_query - def list_by_type(cls, db: Session, mtype: str, page: Optional[int] = 1, count: Optional[int] = 30): - return db.query(cls).filter( - cls.type == mtype - ).order_by( - cls.date.desc() - ).offset((page - 1) * count).limit(count).all() + def list_by_type(cls, db: Session, mtype: str, page: int = 1, count: int = 30): + return list(db.execute( + select(cls).where( + cls.type == mtype + ).order_by( + cls.date.desc() + ).offset((page - 1) * count).limit(count) + ).scalars().all()) @classmethod @async_db_query - async def async_list_by_type(cls, db: AsyncSession, mtype: str, page: Optional[int] = 1, count: Optional[int] = 30): + async def async_list_by_type(cls, db: AsyncSession, mtype: str, page: int = 1, count: int = 30): result = await db.execute( select(cls).filter( cls.type == mtype @@ -124,7 +126,7 @@ class SubscribeHistory(Base): cls.date.desc() ).offset((page - 1) * count).limit(count) ) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @async_db_query @@ -133,8 +135,8 @@ class SubscribeHistory(Base): db: AsyncSession, mtype: str, username: str, - page: Optional[int] = 1, - count: Optional[int] = 30 + page: int = 1, + count: int = 30 ): """ 按订阅 owner 查询指定类型的历史分页。 @@ -149,7 +151,7 @@ class SubscribeHistory(Base): cls.date.desc() ).offset((page - 1) * count).limit(count) ) - return result.scalars().all() + return list(result.scalars().all()) @classmethod def _identity_condition( @@ -185,11 +187,11 @@ class SubscribeHistory(Base): ) if condition is None: return None - query = db.query(cls).filter(condition) + statement = select(cls).where(condition) if season is not None: - query = query.filter(cls.season == season) - query = query.filter(cls.episode_group == episode_group) - return query.first() + statement = statement.where(cls.season == season) + statement = statement.where(cls.episode_group == episode_group) + return db.execute(statement).scalars().first() @classmethod @async_db_query diff --git a/app/db/models/systemconfig.py b/app/db/models/systemconfig.py index c7b3dac26..f3ab120d5 100644 --- a/app/db/models/systemconfig.py +++ b/app/db/models/systemconfig.py @@ -1,6 +1,7 @@ -from sqlalchemy import Column, String, JSON, select +from typing import Any, Optional +from sqlalchemy import String, JSON, select from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column from app.db import db_query, db_update, Base, async_db_query, get_id_column @@ -11,14 +12,14 @@ class SystemConfig(Base): """ id = get_id_column() # 主键 - key = Column(String, index=True) + key: Mapped[Optional[str]] = mapped_column(String, index=True) # 值 - value = Column(JSON) + value: Mapped[Optional[Any]] = mapped_column(JSON) @classmethod @db_query def get_by_key(cls, db: Session, key: str): - return db.query(cls).filter(cls.key == key).first() + return db.execute(select(cls).where(cls.key == key)).scalars().first() @classmethod @async_db_query diff --git a/app/db/models/transferhistory.py b/app/db/models/transferhistory.py index 9ba2279b0..3be5d8087 100644 --- a/app/db/models/transferhistory.py +++ b/app/db/models/transferhistory.py @@ -1,14 +1,14 @@ import re import time from pathlib import Path -from typing import List, Optional +from typing import Any, List, Optional -from sqlalchemy import Boolean, Column, Index, Integer, JSON, String, func, or_, select +from sqlalchemy import Boolean, Index, Integer, JSON, String, delete, func, or_, select, update from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column -from app.db import db_query, db_update, get_id_column, Base, async_db_query -from app.db.models.media_identity import media_identity_constraint +from app.db import Base, async_db_query, db_query, db_update, execute_dml, get_id_column +from app.db.models._constraints import media_identity_constraint from app.schemas.types import MUSIC_ENTITY_ALBUM, MUSIC_ENTITY_RECORDING, MediaSource, MediaType @@ -25,64 +25,64 @@ class TransferHistory(Base): """ id = get_id_column() # 源路径 - src = Column(String, index=True) + src: Mapped[Optional[str]] = mapped_column(String, index=True) # 源存储 - src_storage = Column(String, nullable=False, default="local") + src_storage: Mapped[str] = mapped_column(String, nullable=False, default="local") # 源文件项 - src_fileitem = Column(JSON, default=dict) + src_fileitem: Mapped[Optional[Any]] = mapped_column(JSON, default=dict) # 目标路径 - dest = Column(String) + dest: Mapped[Optional[str]] = mapped_column(String) # 目标存储 - dest_storage = Column(String) + dest_storage: Mapped[Optional[str]] = mapped_column(String) # 目标文件项 - dest_fileitem = Column(JSON, default=dict) + dest_fileitem: Mapped[Optional[Any]] = mapped_column(JSON, default=dict) # 转移模式 move/copy/link... - mode = Column(String) + mode: Mapped[Optional[str]] = mapped_column(String) # 类型 电影/电视剧 - type = Column(String) + type: Mapped[Optional[str]] = mapped_column(String) # 二级分类 - category = Column(String) + category: Mapped[Optional[str]] = mapped_column(String) # 标题 - title = Column(String, index=True) + title: Mapped[Optional[str]] = mapped_column(String, index=True) # 年份 - year = Column(String) + year: Mapped[Optional[str]] = mapped_column(String) # 媒体数据源与原生ID - media_source = Column(String, index=True) - media_id = Column(String, index=True) + media_source: Mapped[Optional[str]] = mapped_column(String, index=True) + media_id: Mapped[Optional[str]] = mapped_column(String, index=True) # 音乐实体类型:recording 单曲、album 专辑 - music_type = Column(String) + music_type: Mapped[Optional[str]] = mapped_column(String) # 专辑预期总曲目数 - total_tracks = Column(Integer) + total_tracks: Mapped[Optional[int]] = mapped_column(Integer) # 实际音频格式 - audio_format = Column(String) + audio_format: Mapped[Optional[str]] = mapped_column(String) # 是否无损音频 - audio_lossless = Column(Boolean) + audio_lossless: Mapped[Optional[bool]] = mapped_column(Boolean) # 实际位深(bit) - bit_depth = Column(Integer) + bit_depth: Mapped[Optional[int]] = mapped_column(Integer) # 实际采样率(Hz) - sample_rate = Column(Integer) + sample_rate: Mapped[Optional[int]] = mapped_column(Integer) # 实际码率(bps) - bitrate = Column(Integer) + bitrate: Mapped[Optional[int]] = mapped_column(Integer) # Sxx - seasons = Column(String) + seasons: Mapped[Optional[str]] = mapped_column(String) # Exx - episodes = Column(String) + episodes: Mapped[Optional[str]] = mapped_column(String) # 海报 - image = Column(String) + image: Mapped[Optional[str]] = mapped_column(String) # 下载器 - downloader = Column(String) + downloader: Mapped[Optional[str]] = mapped_column(String) # 下载器hash - download_hash = Column(String, index=True) + download_hash: Mapped[Optional[str]] = mapped_column(String, index=True) # 转移成功状态 - status = Column(Boolean(), default=True) + status: Mapped[Optional[bool]] = mapped_column(Boolean(), default=True) # 转移失败信息 - errmsg = Column(String) + errmsg: Mapped[Optional[str]] = mapped_column(String) # 时间 - date = Column(String) + date: Mapped[Optional[str]] = mapped_column(String) # 文件清单,以JSON存储 - files = Column(JSON, default=list) + files: Mapped[Optional[Any]] = mapped_column(JSON, default=list) # 剧集组 - episode_group = Column(String) + episode_group: Mapped[Optional[str]] = mapped_column(String) __table_args__ = ( media_identity_constraint("transferhistory"), @@ -94,8 +94,8 @@ class TransferHistory(Base): @classmethod @db_query - def list_by_title(cls, db: Session, title: str, page: Optional[int] = 1, count: Optional[int] = 30, - status: bool = None, wildcard: bool = False): + def list_by_title(cls, db: Session, title: str, page: int = 1, count: int = 30, + status: Optional[bool] = None, wildcard: bool = False): if wildcard: text_filter = or_( _text_like(cls.title, title, wildcard=True), @@ -108,21 +108,21 @@ class TransferHistory(Base): _text_like(cls.src, f'%{title}%'), _text_like(cls.dest, f'%{title}%'), ) - query = db.query(cls).filter(text_filter) + statement = select(cls).where(text_filter) if status is not None: - query = query.filter(cls.status == status) - query = query.order_by(cls.date.desc()) + statement = statement.where(cls.status == status) + statement = statement.order_by(cls.date.desc()) # 当count为负数时,不限制页数查询所有 if count >= 0: - query = query.offset((page - 1) * count).limit(count) + statement = statement.offset((page - 1) * count).limit(count) - return query.all() + return list(db.execute(statement).scalars().all()) @classmethod @async_db_query - async def async_list_by_title(cls, db: AsyncSession, title: str, page: Optional[int] = 1, count: Optional[int] = 30, - status: bool = None, wildcard: bool = False): + async def async_list_by_title(cls, db: AsyncSession, title: str, page: int = 1, count: int = 30, + status: Optional[bool] = None, wildcard: bool = False): if wildcard: text_filter = or_( _text_like(cls.title, title, wildcard=True), @@ -145,32 +145,26 @@ class TransferHistory(Base): query = query.offset((page - 1) * count).limit(count) result = await db.execute(query) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query - def list_by_page(cls, db: Session, page: Optional[int] = 1, count: Optional[int] = 30, status: bool = None): + def list_by_page(cls, db: Session, page: int = 1, count: int = 30, status: Optional[bool] = None): + statement = select(cls) if status is not None: - query = db.query(cls).filter( - cls.status == status - ).order_by( - cls.date.desc() - ) - else: - query = db.query(cls).order_by( - cls.date.desc() - ) - + statement = statement.where(cls.status == status) + statement = statement.order_by(cls.date.desc()) + # 当count为负数时,不限制页数查询所有 if count >= 0: - query = query.offset((page - 1) * count).limit(count) - - return query.all() + statement = statement.offset((page - 1) * count).limit(count) + + return list(db.execute(statement).scalars().all()) @classmethod @async_db_query - async def async_list_by_page(cls, db: AsyncSession, page: Optional[int] = 1, count: Optional[int] = 30, - status: bool = None): + async def async_list_by_page(cls, db: AsyncSession, page: int = 1, count: int = 30, + status: Optional[bool] = None): if status is not None: query = select(cls).filter( cls.status == status @@ -187,12 +181,14 @@ class TransferHistory(Base): query = query.offset((page - 1) * count).limit(count) result = await db.execute(query) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query def get_by_hash(cls, db: Session, download_hash: str): - return db.query(cls).filter(cls.download_hash == download_hash).first() + return db.execute( + select(cls).where(cls.download_hash == download_hash) + ).scalars().first() @classmethod @db_query @@ -207,11 +203,12 @@ class TransferHistory(Base): :param storage: 源存储类型 :return: 命中的整理记录,未命中时返回 None """ + statement = select(cls).where(cls.src == src) if storage: - query = db.query(cls).filter(cls.src == src, cls.src_storage == storage) - else: - query = db.query(cls).filter(cls.src == src) - return query.order_by(cls.id.desc()).first() + statement = statement.where(cls.src_storage == storage) + return db.execute( + statement.order_by(cls.id.desc()) + ).scalars().first() @classmethod @db_query @@ -228,10 +225,12 @@ class TransferHistory(Base): :param storage: 源存储类型 :return: 命中的成功整理记录,未命中时返回 None """ - query = db.query(cls).filter(cls.src == src, cls.status.is_(True)) + statement = select(cls).where(cls.src == src, cls.status.is_(True)) if storage: - query = query.filter(cls.src_storage == storage) - return query.order_by(cls.id.desc()).first() + statement = statement.where(cls.src_storage == storage) + return db.execute( + statement.order_by(cls.id.desc()) + ).scalars().first() @classmethod @db_query @@ -246,10 +245,12 @@ class TransferHistory(Base): :param storage: 目标存储类型 :return: 命中的整理记录,未命中时返回 None """ - query = db.query(cls).filter(cls.dest == dest) + statement = select(cls).where(cls.dest == dest) if storage: - query = query.filter(cls.dest_storage == storage) - return query.order_by(cls.id.desc()).first() + statement = statement.where(cls.dest_storage == storage) + return db.execute( + statement.order_by(cls.id.desc()) + ).scalars().first() @classmethod @db_query @@ -272,24 +273,24 @@ class TransferHistory(Base): normalized_src = ( Path(str(src).replace("\\", "/")).as_posix().rstrip("/") or "/" ) - query = db.query(cls).filter(cls.status.is_(True)) + statement = select(cls).where(cls.status.is_(True)) if recursive: escaped_src = ( normalized_src.replace("\\", "\\\\") .replace("%", "\\%") .replace("_", "\\_") ) - query = query.filter( + statement = statement.where( or_( cls.src == normalized_src, cls.src.like(f"{escaped_src.rstrip('/')}/%", escape="\\"), ) ) else: - query = query.filter(cls.src == normalized_src) + statement = statement.where(cls.src == normalized_src) if storage: - query = query.filter(cls.src_storage == storage) - return query.all() + statement = statement.where(cls.src_storage == storage) + return list(db.execute(statement).scalars().all()) @classmethod @db_query @@ -312,7 +313,7 @@ class TransferHistory(Base): normalized_dest = ( Path(str(dest).replace("\\", "/")).as_posix().rstrip("/") or "/" ) - query = db.query(cls).filter( + statement = select(cls).where( cls.status.is_(True), cls.mode.contains("move"), ) @@ -322,34 +323,41 @@ class TransferHistory(Base): .replace("%", "\\%") .replace("_", "\\_") ) - query = query.filter( + statement = statement.where( or_( cls.dest == normalized_dest, cls.dest.like(f"{escaped_dest.rstrip('/')}/%", escape="\\"), ) ) else: - query = query.filter(cls.dest == normalized_dest) + statement = statement.where(cls.dest == normalized_dest) if storage: - query = query.filter(cls.dest_storage == storage) - return query.all() + statement = statement.where(cls.dest_storage == storage) + return list(db.execute(statement).scalars().all()) @classmethod @db_query def list_by_hash(cls, db: Session, download_hash: str): - return db.query(cls).filter(cls.download_hash == download_hash).all() + return list(db.execute( + select(cls).where(cls.download_hash == download_hash) + ).scalars().all()) @classmethod @db_query - def statistic(cls, db: Session, days: Optional[int] = 7): + def statistic(cls, db: Session, days: int = 7): """ 统计最近days天的下载历史数量,按日期分组返回每日数量 """ - sub_query = db.query(func.substr(cls.date, 1, 10).label('date'), - cls.id.label('id')).filter( + sub_query = select( + func.substr(cls.date, 1, 10).label('date'), + cls.id.label('id') + ).where( cls.date >= time.strftime("%Y-%m-%d %H:%M:%S", - time.localtime(time.time() - 86400 * days))).subquery() - return db.query(sub_query.c.date, func.count(sub_query.c.id)).group_by(sub_query.c.date).all() + time.localtime(time.time() - 86400 * days)) + ).subquery() + return list(db.execute( + select(sub_query.c.date, func.count(sub_query.c.id)).group_by(sub_query.c.date) + ).all()) @classmethod @db_query @@ -361,11 +369,11 @@ class TransferHistory(Base): 缺少集数时按单条成功整理记录计数;音乐按曲目身份去重,整专记录不能只按专辑 ID 合并。 """ month_prefix = time.strftime("%Y-%m-", time.localtime()) - histories = db.query(cls).filter( + histories = db.execute(select(cls).where( cls.status.is_(True), cls.date.like(f"{month_prefix}%"), cls.type.in_([MediaType.MOVIE.value, MediaType.TV.value, MediaType.MUSIC.value]), - ).all() + )).scalars().all() movie_identities = set() tv_identities = set() episode_count = 0 @@ -419,7 +427,7 @@ class TransferHistory(Base): @classmethod @async_db_query - async def async_statistic(cls, db: AsyncSession, days: Optional[int] = 7): + async def async_statistic(cls, db: AsyncSession, days: int = 7): """ 统计最近days天的下载历史数量,按日期分组返回每日数量 """ @@ -434,15 +442,15 @@ class TransferHistory(Base): @classmethod @db_query - def count(cls, db: Session, status: bool = None): + def count(cls, db: Session, status: Optional[bool] = None): + statement = select(func.count(cls.id)) if status is not None: - return db.query(func.count(cls.id)).filter(cls.status == status).first()[0] - else: - return db.query(func.count(cls.id)).first()[0] + statement = statement.where(cls.status == status) + return db.execute(statement).scalar() @classmethod @async_db_query - async def async_count(cls, db: AsyncSession, status: bool = None): + async def async_count(cls, db: AsyncSession, status: Optional[bool] = None): if status is not None: result = await db.execute( select(func.count(cls.id)).filter(cls.status == status) @@ -455,7 +463,7 @@ class TransferHistory(Base): @classmethod @db_query - def count_by_title(cls, db: Session, title: str, status: bool = None, wildcard: bool = False): + def count_by_title(cls, db: Session, title: str, status: Optional[bool] = None, wildcard: bool = False): if wildcard: text_filter = or_( _text_like(cls.title, title, wildcard=True), @@ -468,14 +476,14 @@ class TransferHistory(Base): _text_like(cls.src, f'%{title}%'), _text_like(cls.dest, f'%{title}%'), ) - query = db.query(func.count(cls.id)).filter(text_filter) + statement = select(func.count(cls.id)).where(text_filter) if status is not None: - query = query.filter(cls.status == status) - return query.first()[0] + statement = statement.where(cls.status == status) + return db.execute(statement).scalar() @classmethod @async_db_query - async def async_count_by_title(cls, db: AsyncSession, title: str, status: bool = None, wildcard: bool = False): + async def async_count_by_title(cls, db: AsyncSession, title: str, status: Optional[bool] = None, wildcard: bool = False): if wildcard: text_filter = or_( _text_like(cls.title, title, wildcard=True), @@ -506,63 +514,31 @@ class TransferHistory(Base): 按媒体身份、季集或标题年份查询整理记录。 """ if media_source and media_id and mtype: - # 电视剧某季某集 - if season is not None and episode: - return db.query(cls).filter(cls.media_source == str(media_source), - cls.media_id == str(media_id), - cls.type == mtype, - cls.seasons == season, - cls.episodes == episode, - cls.dest == dest).all() - # 电视剧某季 - elif season is not None: - return db.query(cls).filter(cls.media_source == str(media_source), - cls.media_id == str(media_id), - cls.type == mtype, - cls.seasons == season).all() - else: - if dest: - # 电影 - return db.query(cls).filter(cls.media_source == str(media_source), - cls.media_id == str(media_id), - cls.type == mtype, - cls.dest == dest).all() - else: - # 电视剧所有季集 - return db.query(cls).filter(cls.media_source == str(media_source), - cls.media_id == str(media_id), - cls.type == mtype).all() - # 标题 + 年份 + statement = select(cls).where(cls.media_source == str(media_source), + cls.media_id == str(media_id), + cls.type == mtype) elif title and year: - # 电视剧某季某集 - if season is not None and episode: - return db.query(cls).filter(cls.title == title, - cls.year == year, - cls.seasons == season, - cls.episodes == episode, - cls.dest == dest).all() - # 电视剧某季 - elif season is not None: - return db.query(cls).filter(cls.title == title, - cls.year == year, - cls.seasons == season).all() - else: - if dest: - # 电影 - return db.query(cls).filter(cls.title == title, - cls.year == year, - cls.dest == dest).all() - else: - # 电视剧所有季集 - return db.query(cls).filter(cls.title == title, - cls.year == year).all() - # 类型 + 转移路径(媒体服务器 webhook 缺少远端身份场景) + statement = select(cls).where(cls.title == title, + cls.year == year) elif mtype and season is not None and dest: + # 类型 + 转移路径(媒体服务器 webhook 缺少远端身份场景) + return list(db.execute(select(cls).where(cls.type == mtype, + cls.seasons == season, + cls.dest.like(f"{dest}%"))).scalars().all()) + else: + return [] + if season is not None and episode: + # 电视剧某季某集:目标路径同样参与匹配,dest 为空即匹配空目标 + statement = statement.where(cls.seasons == season, + cls.episodes == episode, + cls.dest == dest) + elif season is not None: # 电视剧某季 - return db.query(cls).filter(cls.type == mtype, - cls.seasons == season, - cls.dest.like(f"{dest}%")).all() - return [] + statement = statement.where(cls.seasons == season) + elif dest: + # 电影:没有季集,用目标路径区分不同版本 + statement = statement.where(cls.dest == dest) + return list(db.execute(statement).scalars().all()) @classmethod @db_query @@ -571,19 +547,17 @@ class TransferHistory(Base): mtype: Optional[str] = None, ): """按规范媒体身份和类型查询整理记录。""" - return db.query(cls).filter( + return db.execute(select(cls).where( cls.media_source == str(media_source), cls.media_id == str(media_id), cls.type == mtype, - ).first() + )).scalars().first() @classmethod @db_update def update_download_hash(cls, db: Session, historyid: Optional[int] = None, download_hash: Optional[str] = None): - db.query(cls).filter(cls.id == historyid).update( - { - "download_hash": download_hash - } + db.execute( + update(cls).where(cls.id == historyid).values(download_hash=download_hash) ) @classmethod @@ -602,10 +576,13 @@ class TransferHistory(Base): src_storage = kwargs.get("src_storage") or "local" kwargs["src_storage"] = src_storage if src: - db.query(cls).filter( - cls.src == src, - cls.src_storage == src_storage, - ).delete(synchronize_session=False) + db.execute( + delete(cls).where( + cls.src == src, + cls.src_storage == src_storage, + ), + execution_options={"synchronize_session": False}, + ) history = cls(**kwargs) db.add(history) db.flush() @@ -617,7 +594,9 @@ class TransferHistory(Base): """ 查询某时间之后的转移历史 """ - return db.query(cls).filter(cls.date > date).order_by(cls.id.desc()).all() + return list(db.execute( + select(cls).where(cls.date > date).order_by(cls.id.desc()) + ).scalars().all()) @classmethod @db_update @@ -630,18 +609,15 @@ class TransferHistory(Base): """ 分批删除指定时间之前的整理历史。 """ - ids = [ - row[0] - for row in db.query(cls.id) - .filter(cls.date < before_time) + ids = db.execute( + select(cls.id) + .where(cls.date < before_time) .order_by(cls.id.asc()) .limit(limit) - .all() - ] + ).scalars().all() if not ids: return 0 - return ( - db.query(cls) - .filter(cls.id.in_(ids)) - .delete(synchronize_session=False) + return execute_dml( + db, delete(cls).where(cls.id.in_(ids)), + execution_options={"synchronize_session": False}, ) diff --git a/app/db/models/transferpending.py b/app/db/models/transferpending.py index aeb3de804..59269d661 100644 --- a/app/db/models/transferpending.py +++ b/app/db/models/transferpending.py @@ -1,9 +1,9 @@ from typing import List, Optional -from sqlalchemy import Column, Index, String -from sqlalchemy.orm import Session +from sqlalchemy import Index, String, delete, select +from sqlalchemy.orm import Mapped, Session, mapped_column -from app.db import Base, db_query, db_update, get_id_column +from app.db import Base, db_query, db_update, execute_dml, get_id_column class TransferPending(Base): @@ -22,11 +22,11 @@ class TransferPending(Base): id = get_id_column() # 存储 - storage = Column(String, nullable=False) + storage: Mapped[str] = mapped_column(String, nullable=False) # 源文件路径 - src_path = Column(String, nullable=False) + src_path: Mapped[str] = mapped_column(String, nullable=False) # 登记时间 - created_at = Column(String) + created_at: Mapped[Optional[str]] = mapped_column(String) __table_args__ = ( # 同一个文件重复入队只保留一条,回放时不会重复送入整理链 @@ -47,9 +47,9 @@ class TransferPending(Base): """ if not storage or not src_path: return None - pending = db.query(cls).filter( - cls.storage == storage, cls.src_path == src_path - ).first() + pending = db.execute( + select(cls).where(cls.storage == storage, cls.src_path == src_path) + ).scalars().first() if pending: return pending pending = cls(storage=storage, src_path=src_path, created_at=now_time) @@ -68,9 +68,10 @@ class TransferPending(Base): """ if not storage or not src_path: return 0 - return db.query(cls).filter( - cls.storage == storage, cls.src_path == src_path - ).delete(synchronize_session=False) + return execute_dml( + db, delete(cls).where(cls.storage == storage, cls.src_path == src_path), + execution_options={"synchronize_session": False}, + ) @classmethod @db_query @@ -84,12 +85,11 @@ class TransferPending(Base): :param limit: 单次回放上限 :return: 待整理登记列表 """ - return ( - db.query(cls) + return list(db.execute( + select(cls) .order_by(cls.created_at.asc(), cls.id.asc()) .limit(limit) - .all() - ) + ).scalars().all()) @classmethod @db_update @@ -99,4 +99,7 @@ class TransferPending(Base): :param db: 数据库会话 :return: 删除的记录数 """ - return db.query(cls).delete(synchronize_session=False) + return execute_dml( + db, delete(cls), + execution_options={"synchronize_session": False}, + ) diff --git a/app/db/models/user.py b/app/db/models/user.py index 91f37f9c8..6ad0bc211 100644 --- a/app/db/models/user.py +++ b/app/db/models/user.py @@ -1,6 +1,7 @@ -from sqlalchemy import Boolean, Column, JSON, String, select +from typing import Any, Optional +from sqlalchemy import Boolean, JSON, String, select from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy.orm import Session +from sqlalchemy.orm import Mapped, Session, mapped_column from app.db import Base, db_query, db_update, async_db_query, async_db_update, get_id_column @@ -12,30 +13,30 @@ class User(Base): # ID id = get_id_column() # 用户名,唯一值 - name = Column(String, index=True, nullable=False) + name: Mapped[str] = mapped_column(String, index=True, nullable=False) # 邮箱 - email = Column(String) + email: Mapped[Optional[str]] = mapped_column(String) # 加密后密码 - hashed_password = Column(String) + hashed_password: Mapped[Optional[str]] = mapped_column(String) # 是否启用 - is_active = Column(Boolean(), default=True) + is_active: Mapped[Optional[bool]] = mapped_column(Boolean(), default=True) # 是否管理员 - is_superuser = Column(Boolean(), default=False) + is_superuser: Mapped[Optional[bool]] = mapped_column(Boolean(), default=False) # 头像 - avatar = Column(String) + avatar: Mapped[Optional[str]] = mapped_column(String) # 是否启用otp二次验证 - is_otp = Column(Boolean(), default=False) + is_otp: Mapped[Optional[bool]] = mapped_column(Boolean(), default=False) # otp秘钥 - otp_secret = Column(String, default=None) + otp_secret: Mapped[Optional[str]] = mapped_column(String, default=None) # 用户权限 json - permissions = Column(JSON, default=dict) + permissions: Mapped[Optional[Any]] = mapped_column(JSON, default=dict) # 用户个性化设置 json - settings = Column(JSON, default=dict) + settings: Mapped[Optional[Any]] = mapped_column(JSON, default=dict) @classmethod @db_query def get_by_name(cls, db: Session, name: str): - return db.query(cls).filter(cls.name == name).first() + return db.execute(select(cls).where(cls.name == name)).scalars().first() @classmethod @async_db_query @@ -48,7 +49,7 @@ class User(Base): @classmethod @db_query def get_by_id(cls, db: Session, user_id: int): - return db.query(cls).filter(cls.id == user_id).first() + return db.execute(select(cls).where(cls.id == user_id)).scalars().first() @classmethod @async_db_query diff --git a/app/db/models/userconfig.py b/app/db/models/userconfig.py index c99424015..281621d8d 100644 --- a/app/db/models/userconfig.py +++ b/app/db/models/userconfig.py @@ -1,5 +1,6 @@ -from sqlalchemy import Column, String, UniqueConstraint, JSON -from sqlalchemy.orm import Session +from typing import Any, Optional +from sqlalchemy import String, UniqueConstraint, JSON, select +from sqlalchemy.orm import Mapped, Session, mapped_column from app.db import db_query, db_update, get_id_column, Base @@ -10,11 +11,11 @@ class UserConfig(Base): """ id = get_id_column() # 用户名 - username = Column(String) + username: Mapped[Optional[str]] = mapped_column(String) # 配置键 - key = Column(String) + key: Mapped[Optional[str]] = mapped_column(String) # 值 - value = Column(JSON) + value: Mapped[Optional[Any]] = mapped_column(JSON) __table_args__ = ( # 用户名和配置键联合唯一 @@ -24,10 +25,9 @@ class UserConfig(Base): @classmethod @db_query def get_by_key(cls, db: Session, username: str, key: str): - return db.query(cls) \ - .filter(cls.username == username) \ - .filter(cls.key == key) \ - .first() + return db.execute( + select(cls).where(cls.username == username, cls.key == key) + ).scalars().first() @db_update def delete_by_key(self, db: Session, username: str, key: str): diff --git a/app/db/models/workflow.py b/app/db/models/workflow.py index 4a251ecdf..fd91acbf9 100644 --- a/app/db/models/workflow.py +++ b/app/db/models/workflow.py @@ -1,8 +1,9 @@ from datetime import datetime from builtins import list as builtin_list -from typing import Optional +from typing import Any, Optional -from sqlalchemy import Column, Integer, JSON, String, Index, and_, or_, select +from sqlalchemy import Integer, JSON, String, Index, and_, or_, select, update +from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.ext.asyncio import AsyncSession from app.db import Base, db_query, get_id_column, db_update, async_db_query, async_db_update @@ -15,71 +16,60 @@ class Workflow(Base): # ID id = get_id_column() # 名称 - name = Column(String, index=True, nullable=False) + name: Mapped[str] = mapped_column(String, index=True, nullable=False) # 描述 - description = Column(String) + description: Mapped[Optional[str]] = mapped_column(String) # 定时器 - timer = Column(String) + timer: Mapped[Optional[str]] = mapped_column(String) # 触发类型:timer-定时触发 event-事件触发 manual-手动触发 - trigger_type = Column(String, default='timer') + trigger_type: Mapped[Optional[str]] = mapped_column(String, default='timer') # 事件类型(当trigger_type为event时使用) - event_type = Column(String) + event_type: Mapped[Optional[str]] = mapped_column(String) # 事件条件(JSON格式,用于过滤事件) - event_conditions = Column(JSON, default=dict) + event_conditions: Mapped[Optional[Any]] = mapped_column(JSON, default=dict) # 状态:W-等待 R-运行中 P-暂停 S-成功 F-失败 - state = Column(String, nullable=False, index=True, default='W') + state: Mapped[str] = mapped_column(String, nullable=False, index=True, default='W') # 已执行动作(,分隔) - current_action = Column(String) + current_action: Mapped[Optional[str]] = mapped_column(String) # 任务执行结果 - result = Column(String) + result: Mapped[Optional[str]] = mapped_column(String) # 已执行次数 - run_count = Column(Integer, default=0) + run_count: Mapped[Optional[int]] = mapped_column(Integer, default=0) # 任务列表 - actions = Column(JSON, default=builtin_list) + actions: Mapped[Optional[Any]] = mapped_column(JSON, default=builtin_list) # 任务流 - flows = Column(JSON, default=builtin_list) + flows: Mapped[Optional[Any]] = mapped_column(JSON, default=builtin_list) # 执行上下文 - context = Column(JSON, default=dict) + context: Mapped[Optional[Any]] = mapped_column(JSON, default=dict) # 执行配置 - execution_config = Column(JSON, default=dict) + execution_config: Mapped[Optional[Any]] = mapped_column(JSON, default=dict) # 结构化执行状态 - execution_state = Column(JSON, default=dict) + execution_state: Mapped[Optional[Any]] = mapped_column(JSON, default=dict) # 创建时间 - add_time = Column(String, default=lambda: datetime.now().strftime('%Y-%m-%d %H:%M:%S')) + add_time: Mapped[Optional[str]] = mapped_column(String, default=lambda: datetime.now().strftime('%Y-%m-%d %H:%M:%S')) # 最后执行时间 - last_time = Column(String) + last_time: Mapped[Optional[str]] = mapped_column(String) __table_args__ = ( Index('ix_workflow_trigger_type_state', 'trigger_type', 'state'), ) - @classmethod - @db_query - def list(cls, db): - return db.query(cls).all() - - @classmethod - @async_db_query - async def async_list(cls, db: AsyncSession): - result = await db.execute(select(cls)) - return result.scalars().all() - @classmethod @db_query def get_enabled_workflows(cls, db): - return db.query(cls).filter(cls.state != 'P').all() + return list(db.execute(select(cls).where(cls.state != 'P')).scalars().all()) @classmethod @async_db_query async def async_get_enabled_workflows(cls, db: AsyncSession): result = await db.execute(select(cls).where(cls.state != 'P')) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query def get_timer_triggered_workflows(cls, db): """获取定时触发的工作流""" - return db.query(cls).filter( + return list(db.execute(select(cls).where( and_( or_( cls.trigger_type == 'timer', @@ -87,7 +77,7 @@ class Workflow(Base): ), cls.state != 'P' ) - ).all() + )).scalars().all()) @classmethod @async_db_query @@ -102,18 +92,18 @@ class Workflow(Base): cls.state != 'P' ) )) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query def get_event_triggered_workflows(cls, db): """获取事件触发的工作流""" - return db.query(cls).filter( + return list(db.execute(select(cls).where( and_( cls.trigger_type == 'event', cls.state != 'P' ) - ).all() + )).scalars().all()) @classmethod @async_db_query @@ -125,12 +115,12 @@ class Workflow(Base): cls.state != 'P' ) )) - return result.scalars().all() + return list(result.scalars().all()) @classmethod @db_query def get_by_name(cls, db, name: str): - return db.query(cls).filter(cls.name == name).first() + return db.execute(select(cls).where(cls.name == name)).scalars().first() @classmethod @async_db_query @@ -141,45 +131,42 @@ class Workflow(Base): @classmethod @db_update def update_state(cls, db, wid: int, state: str): - db.query(cls).filter(cls.id == wid).update({"state": state}) + db.execute(update(cls).where(cls.id == wid).values(state=state)) return True @classmethod @async_db_update async def async_update_state(cls, db: AsyncSession, wid: int, state: str): - from sqlalchemy import update await db.execute(update(cls).where(cls.id == wid).values(state=state)) return True @classmethod @db_update def start(cls, db, wid: int): - db.query(cls).filter(cls.id == wid).update({ - "state": 'R' - }) + db.execute(update(cls).where(cls.id == wid).values(state='R')) return True @classmethod @async_db_update async def async_start(cls, db: AsyncSession, wid: int): - from sqlalchemy import update await db.execute(update(cls).where(cls.id == wid).values(state='R')) return True @classmethod @db_update def fail(cls, db, wid: int, result: str): - db.query(cls).filter(and_(cls.id == wid, cls.state != "P")).update({ - "state": 'F', - "result": result, - "last_time": datetime.now().strftime('%Y-%m-%d %H:%M:%S') - }) + db.execute(update(cls).where( + and_(cls.id == wid, cls.state != "P") + ).values( + state='F', + result=result, + last_time=datetime.now().strftime('%Y-%m-%d %H:%M:%S') + )) return True @classmethod @async_db_update async def async_fail(cls, db: AsyncSession, wid: int, result: str): - from sqlalchemy import update await db.execute(update(cls).where( and_(cls.id == wid, cls.state != "P") ).values( @@ -192,18 +179,19 @@ class Workflow(Base): @classmethod @db_update def success(cls, db, wid: int, result: Optional[str] = None): - db.query(cls).filter(and_(cls.id == wid, cls.state != "P")).update({ - "state": 'S', - "result": result, - "run_count": cls.run_count + 1, - "last_time": datetime.now().strftime('%Y-%m-%d %H:%M:%S') - }) + db.execute(update(cls).where( + and_(cls.id == wid, cls.state != "P") + ).values( + state='S', + result=result, + run_count=cls.run_count + 1, + last_time=datetime.now().strftime('%Y-%m-%d %H:%M:%S') + )) return True @classmethod @async_db_update async def async_success(cls, db: AsyncSession, wid: int, result: Optional[str] = None): - from sqlalchemy import update await db.execute(update(cls).where( and_(cls.id == wid, cls.state != "P") ).values( @@ -217,20 +205,19 @@ class Workflow(Base): @classmethod @db_update def reset(cls, db, wid: int, reset_count: Optional[bool] = False): - db.query(cls).filter(cls.id == wid).update({ - "state": 'W', - "result": None, - "current_action": None, - "context": {}, - "execution_state": {}, - "run_count": 0 if reset_count else cls.run_count, - }) + db.execute(update(cls).where(cls.id == wid).values( + state='W', + result=None, + current_action=None, + context={}, + execution_state={}, + run_count=0 if reset_count else cls.run_count, + )) return True @classmethod @async_db_update async def async_reset(cls, db: AsyncSession, wid: int, reset_count: Optional[bool] = False): - from sqlalchemy import update await db.execute(update(cls).where(cls.id == wid).values( state='W', result=None, @@ -245,7 +232,7 @@ class Workflow(Base): @db_update def update_current_action(cls, db, wid: int, action_id: str, context: dict, execution_state: Optional[dict] = None): - workflow = db.query(cls).filter(cls.id == wid).first() + workflow = db.execute(select(cls).where(cls.id == wid)).scalars().first() current_actions = [] if workflow and workflow.current_action: current_actions = [item for item in workflow.current_action.split(",") if item] @@ -257,14 +244,13 @@ class Workflow(Base): } if execution_state is not None: update_values["execution_state"] = execution_state - db.query(cls).filter(cls.id == wid).update(update_values) + db.execute(update(cls).where(cls.id == wid).values(**update_values)) return True @classmethod @async_db_update async def async_update_current_action(cls, db: AsyncSession, wid: int, action_id: str, context: dict, execution_state: Optional[dict] = None): - from sqlalchemy import update # 先获取当前current_action result = await db.execute(select(cls.current_action).where(cls.id == wid)) current_action = result.scalar() diff --git a/app/db/oper/__init__.py b/app/db/oper/__init__.py new file mode 100644 index 000000000..4cd94c720 --- /dev/null +++ b/app/db/oper/__init__.py @@ -0,0 +1,100 @@ +""" +数据访问层(Oper)。 + +与 app/db/models 一一对应:models 声明表结构,oper 承载针对该表的读写。 +两个包同名文件互为镜像(models/subscribe.py ↔ oper/subscribe.py), +文件名只写实体,角色由包名表达,因此这里不再有 `_oper` 后缀。 + +本文件只做符号解析,不在 import 期执行任何动作——没有建引擎、没有连库、 +也不会把十六个 Oper 模块一并拉起。`from app.db.oper import SubscribeOper` +经下方 __getattr__ 惰性解析,只导入被点名的那一个模块。 + +这一点不是洁癖:多处测试靠往 sys.modules 塞桩来隔离单个 Oper(例如 +app.db.oper.systemconfig),若本文件改成 models/__init__.py 那样的即时 +re-export,导入任意一个 Oper 都会连带把其余十五个真正拉起来,桩就被绕过了。 +按模块直连(from app.db.oper.subscribe import SubscribeOper)仍是仓库内的 +首选写法,本入口是给「只想要一个类名」的调用方备的门面。 +""" +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + # 运行期由 __getattr__ 解析,模块 __dict__ 里并不存在这些名字; + # 这里为静态检查补上真实类型,同时消掉 __all__ 的未定义告警。 + from app.db.oper.agentchat import AgentChatOper + from app.db.oper.agenttask import AgentTaskOper + from app.db.oper.downloadfailure import DownloadFailureOper + from app.db.oper.downloadhistory import DownloadHistoryOper + from app.db.oper.mediaserver import MediaServerOper + from app.db.oper.message import MessageOper + from app.db.oper.plugindata import PluginDataOper + from app.db.oper.site import SiteOper + from app.db.oper.subscribe import SubscribeOper + from app.db.oper.subscribehistory import SubscribeHistoryOper + from app.db.oper.systemconfig import SystemConfigOper + from app.db.oper.transferhistory import TransferHistoryOper + from app.db.oper.transferpending import TransferPendingOper + from app.db.oper.user import UserOper + from app.db.oper.userconfig import UserConfigOper + from app.db.oper.workflow import WorkflowOper + +# 类名 -> 所在子模块。子模块名即实体名,与 app/db/models 对齐。 +_OPER_MODULES = { + "AgentChatOper": "agentchat", + "AgentTaskOper": "agenttask", + "DownloadFailureOper": "downloadfailure", + "DownloadHistoryOper": "downloadhistory", + "MediaServerOper": "mediaserver", + "MessageOper": "message", + "PluginDataOper": "plugindata", + "SiteOper": "site", + "SubscribeHistoryOper": "subscribehistory", + "SubscribeOper": "subscribe", + "SystemConfigOper": "systemconfig", + "TransferHistoryOper": "transferhistory", + "TransferPendingOper": "transferpending", + "UserConfigOper": "userconfig", + "UserOper": "user", + "WorkflowOper": "workflow", +} + + +def __getattr__(name: str) -> Any: + """ + 惰性解析 Oper 类,只导入被点名的那个子模块。 + :param name: 属性名 + :return: 对应的 Oper 类 + """ + module_name = _OPER_MODULES.get(name) + if module_name is None: + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + from importlib import import_module + + return getattr(import_module(f"{__name__}.{module_name}"), name) + + +def __dir__() -> list[str]: + """ + 让 dir() 与自动补全看得见惰性名字。 + :return: 属性名列表 + """ + return sorted({*globals(), *_OPER_MODULES}) + + +__all__ = [ + "AgentChatOper", + "AgentTaskOper", + "DownloadFailureOper", + "DownloadHistoryOper", + "MediaServerOper", + "MessageOper", + "PluginDataOper", + "SiteOper", + "SubscribeHistoryOper", + "SubscribeOper", + "SystemConfigOper", + "TransferHistoryOper", + "TransferPendingOper", + "UserConfigOper", + "UserOper", + "WorkflowOper", +] diff --git a/app/db/agentchat_oper.py b/app/db/oper/agentchat.py similarity index 96% rename from app/db/agentchat_oper.py rename to app/db/oper/agentchat.py index 3c9dc51ef..328f36f37 100644 --- a/app/db/agentchat_oper.py +++ b/app/db/oper/agentchat.py @@ -16,7 +16,7 @@ class AgentChatOper(DbOper): Agent 会话历史数据管理。 """ - def __init__(self, db: Union[Session, AsyncSession] = None): + def __init__(self, db: Optional[Union[Session, AsyncSession]] = None): super().__init__(db) @staticmethod @@ -96,7 +96,7 @@ class AgentChatOper(DbOper): source: Optional[str] = None, original_chat_id: Optional[str] = None, client_session_id: Optional[str] = None, - ) -> AgentChat: + ) -> Optional[AgentChat]: """ 确保 Agent 会话记录存在,并刷新基础渠道信息。 """ @@ -151,6 +151,8 @@ class AgentChatOper(DbOper): chat = self.get(session_id=session_id) if not chat: chat = self.ensure_session(session_id=session_id, user_id=user_id) + if not chat: + return chat.update( self._db, { @@ -186,6 +188,8 @@ class AgentChatOper(DbOper): original_chat_id=original_chat_id, client_session_id=client_session_id, ) + if not chat: + return if self.has_custom_title(chat.title): return chat.update( @@ -207,7 +211,7 @@ class AgentChatOper(DbOper): original_chat_id: Optional[str] = None, client_session_id: Optional[str] = None, title: Optional[str] = None, - ) -> AgentChat: + ) -> Optional[AgentChat]: """ 保存用户可见的 Agent 会话消息。 """ @@ -221,6 +225,8 @@ class AgentChatOper(DbOper): original_chat_id=original_chat_id, client_session_id=client_session_id, ) + if not chat: + return None normalized_title = ( chat.title if self.has_custom_title(chat.title) @@ -248,7 +254,7 @@ class AgentChatOper(DbOper): source: Optional[str] = None, original_chat_id: Optional[str] = None, client_session_id: Optional[str] = None, - ) -> AgentChat: + ) -> Optional[AgentChat]: """ 追加一组用户可见的 Agent 会话消息。 """ @@ -261,6 +267,8 @@ class AgentChatOper(DbOper): original_chat_id=original_chat_id, client_session_id=client_session_id, ) + if not chat: + return None display_messages = self._normalize_messages(chat.display_messages) display_messages.extend(self._normalize_messages(messages)) title = chat.title if self.has_custom_title(chat.title) else None @@ -278,8 +286,8 @@ class AgentChatOper(DbOper): async def async_list_by_page( self, - page: Optional[int] = 1, - count: Optional[int] = 30, + page: int = 1, + count: int = 30, user_id: Optional[str] = None, username: Optional[str] = None, ) -> list[AgentChat]: diff --git a/app/db/agenttask_oper.py b/app/db/oper/agenttask.py similarity index 99% rename from app/db/agenttask_oper.py rename to app/db/oper/agenttask.py index 512d36f0d..f86cee638 100644 --- a/app/db/agenttask_oper.py +++ b/app/db/oper/agenttask.py @@ -19,7 +19,7 @@ class AgentTaskOper(DbOper): """生成当前数据库时间字符串。""" return datetime.now().strftime("%Y-%m-%d %H:%M:%S") - def add(self, **kwargs: object) -> AgentTask: + def add(self, **kwargs: object) -> Optional[AgentTask]: """ 新增 Agent 定时任务。 """ diff --git a/app/db/downloadfailure_oper.py b/app/db/oper/downloadfailure.py similarity index 92% rename from app/db/downloadfailure_oper.py rename to app/db/oper/downloadfailure.py index 5f935f22b..42faa3d26 100644 --- a/app/db/downloadfailure_oper.py +++ b/app/db/oper/downloadfailure.py @@ -2,7 +2,6 @@ from typing import Dict, List, Optional from app.db import DbOper from app.db.models.downloadfailure import DownloadFailure -from app.domain.media import normalize_media_identity_payload class DownloadFailureOper(DbOper): @@ -39,7 +38,6 @@ class DownloadFailureOper(DbOper): """ 新增或更新资源失败记录。 """ - kwargs = normalize_media_identity_payload(kwargs) return DownloadFailure.record_failure( self._db, fingerprint=fingerprint, diff --git a/app/db/downloadhistory_oper.py b/app/db/oper/downloadhistory.py similarity index 86% rename from app/db/downloadhistory_oper.py rename to app/db/oper/downloadhistory.py index daf19d7e3..41c1d7791 100644 --- a/app/db/downloadhistory_oper.py +++ b/app/db/oper/downloadhistory.py @@ -1,9 +1,8 @@ -from typing import Dict, List, Optional +from typing import Dict, List, Optional, cast from app.db import DbOper from app.db.models.downloadhistory import DownloadHistory, DownloadFiles from app.schemas.types import MediaSource -from app.domain.media import normalize_media_identity_payload class DownloadHistoryOper(DbOper): @@ -11,14 +10,14 @@ class DownloadHistoryOper(DbOper): 下载历史管理 """ - def get_by_path(self, path: str) -> DownloadHistory: + def get_by_path(self, path: str) -> Optional[DownloadHistory]: """ 按路径查询下载记录 :param path: 数据key """ return DownloadHistory.get_by_path(self._db, path) - def get_by_hash(self, download_hash: str) -> DownloadHistory: + def get_by_hash(self, download_hash: str) -> Optional[DownloadHistory]: """ 按Hash查询下载记录 :param download_hash: 数据key @@ -57,7 +56,6 @@ class DownloadHistoryOper(DbOper): """ 新增下载历史 """ - kwargs = normalize_media_identity_payload(kwargs) DownloadHistory(**kwargs).create(self._db) def add_files(self, file_items: List[dict]): @@ -82,19 +80,21 @@ class DownloadHistoryOper(DbOper): """ return DownloadFiles.get_by_hash(self._db, download_hash, state) - def get_file_by_fullpath(self, fullpath: str) -> DownloadFiles: + def get_file_by_fullpath(self, fullpath: str) -> Optional[DownloadFiles]: """ 按fullpath查询下载文件记录 :param fullpath: 数据key """ - return DownloadFiles.get_by_fullpath(self._db, fullpath=fullpath, all_files=False) + return cast(Optional[DownloadFiles], + DownloadFiles.get_by_fullpath(self._db, fullpath=fullpath, all_files=False)) def get_files_by_fullpath(self, fullpath: str) -> List[DownloadFiles]: """ 按fullpath查询下载文件记录 :param fullpath: 数据key """ - return DownloadFiles.get_by_fullpath(self._db, fullpath=fullpath, all_files=True) + return cast(List[DownloadFiles], + DownloadFiles.get_by_fullpath(self._db, fullpath=fullpath, all_files=True)) def get_files_by_savepath(self, fullpath: str) -> List[DownloadFiles]: """ @@ -110,17 +110,18 @@ class DownloadHistoryOper(DbOper): """ DownloadFiles.delete_by_fullpath(self._db, fullpath) - def get_hash_by_fullpath(self, fullpath: str) -> str: + def get_hash_by_fullpath(self, fullpath: str) -> Optional[str]: """ 按fullpath查询下载文件记录hash :param fullpath: 数据key """ - fileinfo: DownloadFiles = DownloadFiles.get_by_fullpath(self._db, fullpath=fullpath, all_files=False) + fileinfo = cast(Optional[DownloadFiles], + DownloadFiles.get_by_fullpath(self._db, fullpath=fullpath, all_files=False)) if fileinfo: return fileinfo.download_hash return "" - def list_by_page(self, page: Optional[int] = 1, count: Optional[int] = 30) -> List[DownloadHistory]: + def list_by_page(self, page: int = 1, count: int = 30) -> List[DownloadHistory]: """ 分页查询下载历史 """ @@ -177,7 +178,7 @@ class DownloadHistoryOper(DbOper): media_id=media_id, seasons=seasons) - def list_by_type(self, mtype: str, days: Optional[int] = 7) -> List[DownloadHistory]: + def list_by_type(self, mtype: str, days: int = 7) -> List[DownloadHistory]: """ 获取指定类型的下载历史 """ diff --git a/app/db/mediaserver_oper.py b/app/db/oper/mediaserver.py similarity index 96% rename from app/db/mediaserver_oper.py rename to app/db/oper/mediaserver.py index 07fdc400d..a37025c9a 100644 --- a/app/db/mediaserver_oper.py +++ b/app/db/oper/mediaserver.py @@ -4,7 +4,6 @@ from sqlalchemy.orm import Session from app.db import DbOper from app.db.models.mediaserver import MediaServerItem -from app.domain.media import normalize_media_identity_payload class MediaServerOper(DbOper): @@ -12,7 +11,7 @@ class MediaServerOper(DbOper): 媒体服务器数据管理 """ - def __init__(self, db: Session = None): + def __init__(self, db: Optional[Session] = None): super().__init__(db) @staticmethod @@ -20,11 +19,10 @@ class MediaServerOper(DbOper): """ 过滤数据库模型不存在或不应由远端覆盖的字段 """ - payload = { + return { k: v for k, v in kwargs.items() if hasattr(MediaServerItem, k) and k != "id" } - return normalize_media_identity_payload(payload) def add(self, **kwargs) -> bool: """ diff --git a/app/db/message_oper.py b/app/db/oper/message.py similarity index 87% rename from app/db/message_oper.py rename to app/db/oper/message.py index 3e63a698b..1f3b39e6a 100644 --- a/app/db/message_oper.py +++ b/app/db/oper/message.py @@ -14,20 +14,20 @@ class MessageOper(DbOper): 消息数据管理 """ - def __init__(self, db: Union[Session, AsyncSession] = None): + def __init__(self, db: Optional[Union[Session, AsyncSession]] = None): super().__init__(db) def add(self, - channel: MessageChannel = None, + channel: Optional[MessageChannel] = None, source: Optional[str] = None, - mtype: NotificationType = None, + mtype: Optional[NotificationType] = None, title: Optional[str] = None, text: Optional[str] = None, image: Optional[str] = None, link: Optional[str] = None, userid: Optional[str] = None, action: Optional[int] = 1, - note: Union[list, dict] = None, + note: Optional[Union[list, dict]] = None, **kwargs) -> dict: """ 新增消息 @@ -64,16 +64,16 @@ class MessageOper(DbOper): return Message(**kwargs).create_and_to_dict(self._db) async def async_add(self, - channel: MessageChannel = None, + channel: Optional[MessageChannel] = None, source: Optional[str] = None, - mtype: NotificationType = None, + mtype: Optional[NotificationType] = None, title: Optional[str] = None, text: Optional[str] = None, image: Optional[str] = None, link: Optional[str] = None, userid: Optional[str] = None, action: Optional[int] = 1, - note: Union[list, dict] = None, + note: Optional[Union[list, dict]] = None, **kwargs) -> Message: """ 异步新增消息 @@ -99,7 +99,7 @@ class MessageOper(DbOper): return await Message(**kwargs).async_create(self._db) - def list_by_page(self, page: Optional[int] = 1, count: Optional[int] = 30) -> list[Message]: + def list_by_page(self, page: int = 1, count: int = 30) -> list[Message]: """ 分页获取消息记录。 """ @@ -115,7 +115,7 @@ class MessageOper(DbOper): return Message.exists_by_source(self._db, source) async def async_list_by_page( - self, page: Optional[int] = 1, count: Optional[int] = 30 + self, page: int = 1, count: int = 30 ) -> list[Message]: """ 分页获取消息记录。 @@ -124,8 +124,8 @@ class MessageOper(DbOper): async def async_list_sent_by_page( self, - page: Optional[int] = 1, - count: Optional[int] = 30, + page: int = 1, + count: int = 30, all_clear_before: Optional[str] = None, system_clear_before: Optional[str] = None, media_clear_before: Optional[str] = None, diff --git a/app/db/plugindata_oper.py b/app/db/oper/plugindata.py similarity index 100% rename from app/db/plugindata_oper.py rename to app/db/oper/plugindata.py diff --git a/app/db/site_oper.py b/app/db/oper/site.py similarity index 94% rename from app/db/site_oper.py rename to app/db/oper/site.py index f5fdda5b5..40c4c4624 100644 --- a/app/db/site_oper.py +++ b/app/db/oper/site.py @@ -23,13 +23,13 @@ class SiteOper(DbOper): return True, "新增站点成功" return False, "站点已存在" - def get(self, sid: int) -> Site: + def get(self, sid: int) -> Optional[Site]: """ 查询单个站点 """ return Site.get(self._db, sid) - async def async_get(self, sid: int) -> Site: + async def async_get(self, sid: int) -> Optional[Site]: """ 异步查询单个站点 """ @@ -71,15 +71,17 @@ class SiteOper(DbOper): """ Site.delete(self._db, sid) - def update(self, sid: int, payload: dict) -> Site: + def update(self, sid: int, payload: dict) -> Optional[Site]: """ 更新站点 """ site = Site.get(self._db, sid) + if not site: + return None site.update(self._db, payload) return site - async def async_update(self, sid: int, payload: dict) -> Site: + async def async_update(self, sid: int, payload: dict) -> Optional[Site]: """ 异步更新站点。 """ @@ -88,25 +90,25 @@ class SiteOper(DbOper): await site.async_update(self._db, payload) return site - def get_by_domain(self, domain: str) -> Site: + def get_by_domain(self, domain: str) -> Optional[Site]: """ 按域名获取站点 """ return Site.get_by_domain(self._db, domain) - async def async_get_by_domain(self, domain: str) -> Site: + async def async_get_by_domain(self, domain: str) -> Optional[Site]: """ 异步按域名获取站点 """ return await Site.async_get_by_domain(self._db, domain) - async def async_get_by_name(self, name: str) -> Site: + async def async_get_by_name(self, name: str) -> Optional[Site]: """ 异步按名称获取站点 """ return await Site.async_get_by_name(self._db, name) - def get_domains_by_ids(self, ids: List[int]) -> List[str]: + def get_domains_by_ids(self, ids: List[int]) -> List[Optional[str]]: """ 按ID获取站点域名 """ @@ -201,7 +203,7 @@ class SiteOper(DbOper): """ return SiteUserData.get_latest(self._db) - def get_icon_by_domain(self, domain: str) -> SiteIcon: + def get_icon_by_domain(self, domain: str) -> Optional[SiteIcon]: """ 按域名获取站点图标 """ diff --git a/app/db/oper/subscribe.py b/app/db/oper/subscribe.py new file mode 100644 index 000000000..fc50cec1e --- /dev/null +++ b/app/db/oper/subscribe.py @@ -0,0 +1,280 @@ +""" +订阅数据访问。 + +本模块只收敛针对订阅表的读写。把 MediaInfo / MusicInfo 翻译成一行订阅是订阅业务的 +规则,住在 app/application/subscribe.py;这里收到的 payload 已经是纯粹的持久化字段, +因此不 import 任何领域对象。 + +留在这一层的只有列类型强转与建库时间戳——它们跟着订阅表的列走,换谁来调都一样。 +""" +import time +from typing import Any, Tuple, List, Optional + +from app.db import DbOper +from app.db.models.subscribe import Subscribe +from app.db.models.subscribehistory import SubscribeHistory +from app.schemas.types import MediaSource + +INTEGER_FLAG_FIELDS = ("best_version", "best_version_full", "search_imdbid", "manual_total_episode") + + +def _normalize_integer_flags(payload: dict, fields: Tuple[str, ...] = INTEGER_FLAG_FIELDS) -> dict: + """ + 将历史兼容的布尔开关转换为整型值,避免 PostgreSQL 严格类型检查失败。 + """ + normalized_payload = dict(payload) + for field in fields: + if isinstance(normalized_payload.get(field), bool): + normalized_payload[field] = int(normalized_payload[field]) + return normalized_payload + + +def _normalize_year(year: Optional[int | str]) -> Optional[str]: + """ + 订阅表的 year 列为字符串类型,而识别链路的媒体年份可能是数字 + (音乐等来源),写库前统一转换为字符串避免数据库类型错误。 + """ + if year is None: + return None + return str(year) + + +def _persistable(payload: dict) -> dict: + """ + 把应用层给的写入字段落成订阅表能收的一行。 + + 做两件事。一是列类型强转:PostgreSQL 的整型列拒收布尔值、字符串列拒收数字,而 + SQLite 会靠类型亲和悄悄替我们转好——漏了只在生产库上炸,所以放在紧挨建模的地方。 + 二是盖建库时间戳:调用方传进来的 date 不作数,否则订阅列表的默认排序与过期清理 + 都会读到一个假的建库时间。 + :param payload: 应用层翻译好的写入字段 + :return: 可直接建模的字段字典 + """ + persistable = _normalize_integer_flags(payload) + persistable["year"] = _normalize_year(persistable.get("year")) + # search_imdbid 参与搜索分支判定,None 与真值都要归一到 0/1,否则同一列在不同 + # 订阅上会存出三种形态,PG 上还会直接拒写 + persistable["search_imdbid"] = 1 if persistable.get("search_imdbid") else 0 + persistable["date"] = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime()) + return persistable + + +class SubscribeOper(DbOper): + """ + 订阅管理 + """ + + def _exists(self, identity: dict, username: Optional[str]) -> Optional[Any]: + """ + 按身份查重。 + :param identity: 查重身份 + :param username: 非空时只在该用户的订阅内查 + :return: 命中的订阅行,未命中为 None + """ + if username: + return Subscribe.exists_by_username(self._db, username=username, **identity) + return Subscribe.exists(self._db, **identity) + + async def _async_exists(self, identity: dict, username: Optional[str]) -> Optional[Any]: + """ + 按身份查重(异步)。 + :param identity: 查重身份 + :param username: 非空时只在该用户的订阅内查 + :return: 命中的订阅行,未命中为 None + """ + if username: + return await Subscribe.async_exists_by_username(self._db, username=username, **identity) + return await Subscribe.async_exists(self._db, **identity) + + def add(self, identity: dict, payload: dict, + username: Optional[str] = None) -> Tuple[int, str]: + """ + 新增订阅:命中既有订阅则原样返回,否则落库后回读。 + + 回读不是多余的一次查询——写入可能被唯一约束或事务回滚吞掉,此时若报成功, + 调用方会继续按订阅已建立往下走,用户看到「订阅成功」却永远等不到资源。 + :param identity: 查重身份(media_source/media_id/music_type/season/episode_group) + :param payload: 订阅表的写入字段,媒体翻译由 app/application/subscribe.py 完成 + :param username: 非空时把查重限定在该用户的订阅内 + :return: (订阅 ID, 结果说明);ID 为 0 表示未新增 + """ + subscribe = self._exists(identity, username) + if subscribe: + return subscribe.id, "订阅已存在" + Subscribe(**_persistable(payload)).create(self._db) + subscribe = self._exists(identity, username) + if not subscribe: + return 0, "新增订阅失败" + return subscribe.id, "新增订阅成功" + + async def async_add(self, identity: dict, payload: dict, + username: Optional[str] = None) -> Tuple[int, str]: + """ + 异步新增订阅,语义与 add 完全一致。 + :param identity: 查重身份(media_source/media_id/music_type/season/episode_group) + :param payload: 订阅表的写入字段,媒体翻译由 app/application/subscribe.py 完成 + :param username: 非空时把查重限定在该用户的订阅内 + :return: (订阅 ID, 结果说明);ID 为 0 表示未新增 + """ + subscribe = await self._async_exists(identity, username) + if subscribe: + return subscribe.id, "订阅已存在" + await Subscribe(**_persistable(payload)).async_create(self._db) + subscribe = await self._async_exists(identity, username) + if not subscribe: + return 0, "新增订阅失败" + return subscribe.id, "新增订阅成功" + + def exists( + self, media_source: MediaSource, media_id: str, + season: Optional[int] = None, episode_group: Optional[str] = None, + music_type: Optional[str] = None, + ) -> bool: + """ + 按媒体身份、季号及可选剧集组判断订阅是否存在。 + """ + identity_params = { + "media_source": media_source, + "media_id": media_id, + "music_type": music_type, + "season": season, + "episode_group": episode_group, + } + return bool(Subscribe.exists(self._db, **identity_params)) + + def get(self, sid: int) -> Optional[Subscribe]: + """ + 获取订阅 + """ + return Subscribe.get(self._db, rid=sid) + + async def async_get(self, sid: int) -> Optional[Subscribe]: + """ + 获取订阅 + """ + return await Subscribe.async_get(self._db, rid=sid) + + def get_by( + self, type: str, media_source: MediaSource, media_id: str, + season: Optional[str] = None, + music_type: Optional[str] = None, + ) -> Optional[Subscribe]: + """ + 根据条件查询订阅 + """ + return Subscribe.get_by( + self._db, type, media_source, media_id, season, music_type, + ) + + async def async_get_by( + self, type: str, media_source: MediaSource, media_id: str, + season: Optional[str] = None, + music_type: Optional[str] = None, + ) -> Optional[Subscribe]: + """ + 根据条件查询订阅 + """ + return await Subscribe.async_get_by( + self._db, type, media_source, media_id, season, music_type, + ) + + def list(self, state: Optional[str] = None) -> List[Subscribe]: + """ + 获取订阅列表 + """ + if state: + return Subscribe.get_by_state(self._db, state) + return Subscribe.list(self._db) + + async def async_list(self, state: Optional[str] = None) -> List[Subscribe]: + """ + 异步获取订阅列表 + """ + if state: + return await Subscribe.async_get_by_state(self._db, state) + return await Subscribe.async_list(self._db) + + def delete(self, sid: int): + """ + 删除订阅 + """ + Subscribe.delete(self._db, rid=sid) + + async def async_delete(self, sid: int): + """ + 异步删除订阅。 + """ + await Subscribe.async_delete(self._db, rid=sid) + + async def async_update(self, sid: int, payload: dict) -> Optional[Subscribe]: + """ + 异步更新订阅。 + """ + subscribe = await self.async_get(sid) + if subscribe: + payload = _normalize_integer_flags(payload) + await subscribe.async_update(self._db, payload) + return subscribe + + async def async_update_filter_groups( + self, sid: int, filter_groups: List[str] + ) -> Optional[Subscribe]: + """ + 异步更新订阅使用的过滤规则组。 + """ + return await self.async_update(sid, {"filter_groups": filter_groups}) + + def update(self, sid: int, payload: dict) -> Optional[Subscribe]: + """ + 更新订阅 + """ + subscribe = self.get(sid) + if subscribe: + payload = _normalize_integer_flags(payload) + subscribe.update(self._db, payload) + return subscribe + + def list_by_username(self, username: str, state: Optional[str] = None, + mtype: Optional[str] = None) -> List[Subscribe]: + """ + 获取指定用户的订阅 + """ + return Subscribe.list_by_username(self._db, username=username, state=state, mtype=mtype) + + def list_by_type(self, mtype: str, days: int = 7) -> List[Subscribe]: + """ + 获取指定类型的订阅 + """ + return Subscribe.list_by_type(self._db, mtype=mtype, days=days) + + def add_history(self, **kwargs): + """ + 新增订阅 + """ + # 去除kwargs中 SubscribeHistory 没有的字段 + kwargs = {k: v for k, v in kwargs.items() if hasattr(SubscribeHistory, k)} + kwargs = _normalize_integer_flags(kwargs) + # 更新完成订阅时间 + kwargs.update({"date": time.strftime("%Y-%m-%d %H:%M:%S", time.localtime())}) + # 去掉主键 + if "id" in kwargs: + kwargs.pop("id") + subscribe = SubscribeHistory(**kwargs) + subscribe.create(self._db) + + def exist_history( + self, media_source: MediaSource, media_id: str, + season: Optional[int] = None, episode_group: Optional[str] = None, + music_type: Optional[str] = None, + ) -> bool: + """ + 按媒体身份、季号及可选剧集组判断订阅历史是否存在。 + """ + identity_params = { + "media_source": media_source, + "media_id": media_id, + "music_type": music_type, + "season": season, + "episode_group": episode_group, + } + return bool(SubscribeHistory.exists(self._db, **identity_params)) diff --git a/app/db/subscribehistory_oper.py b/app/db/oper/subscribehistory.py similarity index 88% rename from app/db/subscribehistory_oper.py rename to app/db/oper/subscribehistory.py index b8bb57036..aef1a9bf8 100644 --- a/app/db/subscribehistory_oper.py +++ b/app/db/oper/subscribehistory.py @@ -12,8 +12,8 @@ class SubscribeHistoryOper(DbOper): async def async_list_by_type( self, mtype: str, - page: Optional[int] = 1, - count: Optional[int] = 30, + page: int = 1, + count: int = 30, ) -> List[SubscribeHistory]: """ 异步按媒体类型分页查询订阅历史。 diff --git a/app/db/systemconfig_oper.py b/app/db/oper/systemconfig.py similarity index 98% rename from app/db/systemconfig_oper.py rename to app/db/oper/systemconfig.py index 3f1229499..3f632daa6 100644 --- a/app/db/systemconfig_oper.py +++ b/app/db/oper/systemconfig.py @@ -87,7 +87,7 @@ class SystemConfigOper(DbOper, metaclass=Singleton): self.__SYSTEMCONF[key] = copy.deepcopy(value) return True - def get(self, key: Union[str, SystemConfigKey] = None) -> Any: + def get(self, key: Optional[Union[str, SystemConfigKey]] = None) -> Any: """ 获取系统设置 """ diff --git a/app/db/transferhistory_oper.py b/app/db/oper/transferhistory.py similarity index 56% rename from app/db/transferhistory_oper.py rename to app/db/oper/transferhistory.py index 4766a3732..f96f94fb3 100644 --- a/app/db/transferhistory_oper.py +++ b/app/db/oper/transferhistory.py @@ -1,14 +1,9 @@ import time from typing import Any, List, Optional -from app.domain.context import MediaInfo -from app.domain.meta.metabase import MetaBase -from app.domain.meta.metamusic import MetaMusic from app.db import DbOper from app.db.models.transferhistory import TransferHistory -from app.schemas import TransferInfo, FileItem -from app.schemas.types import MUSIC_ENTITY_RECORDING, MediaSource -from app.domain.media import normalize_media_identity_payload, resolve_media_identity +from app.schemas.types import MediaSource class TransferHistoryOper(DbOper): @@ -16,14 +11,14 @@ class TransferHistoryOper(DbOper): 转移历史管理 """ - def get(self, historyid: int) -> TransferHistory: + def get(self, historyid: int) -> Optional[TransferHistory]: """ 获取转移历史 :param historyid: 转移历史id """ return TransferHistory.get(self._db, historyid) - async def async_get(self, historyid: int) -> TransferHistory: + async def async_get(self, historyid: int) -> Optional[TransferHistory]: """ 异步获取转移历史。 """ @@ -32,8 +27,8 @@ class TransferHistoryOper(DbOper): async def async_list_by_title( self, title: str, - page: Optional[int] = 1, - count: Optional[int] = 30, + page: int = 1, + count: int = 30, status: Optional[bool] = None, ) -> List[TransferHistory]: """ @@ -45,8 +40,8 @@ class TransferHistoryOper(DbOper): async def async_list_by_page( self, - page: Optional[int] = 1, - count: Optional[int] = 30, + page: int = 1, + count: int = 30, status: Optional[bool] = None, ) -> List[TransferHistory]: """ @@ -56,7 +51,7 @@ class TransferHistoryOper(DbOper): self._db, page=page, count=count, status=status ) - async def async_count(self, status: Optional[bool] = None) -> int: + async def async_count(self, status: Optional[bool] = None) -> Optional[int]: """ 异步统计转移记录数量。 """ @@ -66,7 +61,7 @@ class TransferHistoryOper(DbOper): self, title: str, status: Optional[bool] = None, - ) -> int: + ) -> Optional[int]: """ 异步按标题统计转移记录数量。 """ @@ -166,13 +161,12 @@ class TransferHistoryOper(DbOper): """ 新增转移历史 """ - kwargs = normalize_media_identity_payload(kwargs) kwargs.update({ "date": time.strftime("%Y-%m-%d %H:%M:%S", time.localtime()) }) TransferHistory(**kwargs).create(self._db) - def statistic(self, days: Optional[int] = 7) -> List[Any]: + def statistic(self, days: int = 7) -> List[Any]: """ 统计最近days天的下载历史数量 """ @@ -198,7 +192,7 @@ class TransferHistoryOper(DbOper): def get_by_media_identity( self, media_source: MediaSource, media_id: str, mtype: Optional[str] = None, - ) -> TransferHistory: + ) -> Optional[TransferHistory]: """按规范媒体身份和类型查询整理记录。""" return TransferHistory.get_by_media_identity( db=self._db, @@ -225,11 +219,10 @@ class TransferHistoryOper(DbOper): """ TransferHistory.truncate(self._db) - def add_force(self, **kwargs) -> TransferHistory: + def add_force(self, **kwargs) -> Optional[TransferHistory]: """ 新增转移历史,并以同源存储的记录为准替换旧记录。 """ - kwargs = normalize_media_identity_payload(kwargs) # 文件项的默认存储是 local;归一化旧调用传入的 None,确保运行时语义与 # (src, src_storage) 唯一索引一致。 kwargs["src_storage"] = kwargs.get("src_storage") or "local" @@ -253,119 +246,6 @@ class TransferHistoryOper(DbOper): """ TransferHistory.update_download_hash(self._db, historyid, download_hash) - @staticmethod - def _history_title( - meta: MetaBase, mediainfo: Optional[MediaInfo] = None - ) -> Optional[str]: - """音乐文件优先记录曲目标题,其它媒体保持识别标题。""" - if isinstance(meta, MetaMusic) and meta.title: - return meta.title - if mediainfo and mediainfo.title: - return mediainfo.title - return meta.name - - def add_success(self, fileitem: FileItem, mode: str, meta: MetaBase, - mediainfo: MediaInfo, transferinfo: TransferInfo, - downloader: Optional[str] = None, download_hash: Optional[str] = None): - """ - 新增转移成功历史记录 - """ - media_source, media_id = resolve_media_identity(media=mediainfo) - return self.add_force( - src=fileitem.path, - src_storage=fileitem.storage, - src_fileitem=fileitem.model_dump(), - dest=transferinfo.target_item.path if transferinfo.target_item else None, - dest_storage=transferinfo.target_item.storage if transferinfo.target_item else None, - dest_fileitem=transferinfo.target_item.model_dump() if transferinfo.target_item else None, - mode=mode, - type=mediainfo.type.value, - category=mediainfo.category, - title=self._history_title(meta, mediainfo), - year=mediainfo.year, - media_source=media_source, - media_id=media_id, - music_type=getattr(mediainfo, "music_type", None), - total_tracks=getattr(mediainfo, "total_tracks", None), - audio_format=getattr(meta, "audio_format", None), - audio_lossless=getattr(meta, "audio_lossless", None), - bit_depth=getattr(meta, "bit_depth", None), - sample_rate=getattr(meta, "sample_rate", None), - bitrate=getattr(meta, "bitrate", None), - seasons=meta.season, - episodes=meta.episode, - image=mediainfo.get_poster_image(), - downloader=downloader, - download_hash=download_hash, - status=1, - files=transferinfo.file_list - ) - - def add_fail(self, fileitem: FileItem, mode: str, meta: MetaBase, mediainfo: MediaInfo = None, - transferinfo: TransferInfo = None, downloader: Optional[str] = None, download_hash: Optional[str] = None): - """ - 新增转移失败历史记录 - """ - if mediainfo and transferinfo: - media_source, media_id = resolve_media_identity(media=mediainfo) - his = self.add_force( - src=fileitem.path, - src_storage=fileitem.storage, - src_fileitem=fileitem.model_dump(), - dest=transferinfo.target_item.path if transferinfo.target_item else None, - dest_storage=transferinfo.target_item.storage if transferinfo.target_item else None, - dest_fileitem=transferinfo.target_item.model_dump() if transferinfo.target_item else None, - mode=mode, - type=mediainfo.type.value, - category=mediainfo.category, - title=self._history_title(meta, mediainfo), - year=mediainfo.year or meta.year, - media_source=media_source, - media_id=media_id, - music_type=getattr(mediainfo, "music_type", None), - total_tracks=getattr(mediainfo, "total_tracks", None), - audio_format=getattr(meta, "audio_format", None), - audio_lossless=getattr(meta, "audio_lossless", None), - bit_depth=getattr(meta, "bit_depth", None), - sample_rate=getattr(meta, "sample_rate", None), - bitrate=getattr(meta, "bitrate", None), - seasons=meta.season, - episodes=meta.episode, - image=mediainfo.get_poster_image(), - downloader=downloader, - download_hash=download_hash, - episode_group=mediainfo.episode_group, - status=0, - errmsg=transferinfo.message or '未知错误', - files=transferinfo.file_list - ) - else: - media_source, media_id = resolve_media_identity(media=meta) - his = self.add_force( - type=meta.type.value if meta.type else None, - title=self._history_title(meta), - year=meta.year, - media_source=media_source, - media_id=media_id, - music_type=MUSIC_ENTITY_RECORDING if isinstance(meta, MetaMusic) else None, - audio_format=getattr(meta, "audio_format", None), - audio_lossless=getattr(meta, "audio_lossless", None), - bit_depth=getattr(meta, "bit_depth", None), - sample_rate=getattr(meta, "sample_rate", None), - bitrate=getattr(meta, "bitrate", None), - src=fileitem.path, - src_storage=fileitem.storage, - src_fileitem=fileitem.model_dump(), - mode=mode, - seasons=meta.season, - episodes=meta.episode, - downloader=downloader, - download_hash=download_hash, - status=0, - errmsg="未识别到媒体信息" - ) - return his - def list_by_date(self, date: str) -> List[TransferHistory]: """ 查询某时间之后的转移历史 diff --git a/app/db/transferpending_oper.py b/app/db/oper/transferpending.py similarity index 100% rename from app/db/transferpending_oper.py rename to app/db/oper/transferpending.py diff --git a/app/db/oper/user.py b/app/db/oper/user.py new file mode 100644 index 000000000..f620af491 --- /dev/null +++ b/app/db/oper/user.py @@ -0,0 +1,93 @@ +""" +用户数据访问。 + +认证依赖(get_current_user 等八个)已迁至 app/api/deps.py——那是 HTTP 层的关注点, +产出 403/400 而非数据。本模块只保留 UserOper。 + +这里不为那八个名字留惰性转发。转发曾是给仓外插件备的软着陆,代价是把 +app.db.oper.user -> app.api.deps -> app.application.security 这条边永久焊进依赖图: +数据访问模块从此在静态分析里牵着整个鉴权栈,而仓内没有任何调用方需要它。插件生态 +既已确定迭代,就让旧名字直接以 AttributeError 报错——指向明确、当场可改,好过一条 +悄悄成立的反向依赖。 +""" +from typing import List, Optional + +from app.db import DbOper +from app.db.models.user import User + + +class UserOper(DbOper): + """ + 用户管理 + """ + + def list(self) -> List[User]: + """ + 获取用户列表 + """ + return User.list(self._db) + + def add(self, **kwargs): + """ + 新增用户 + """ + user = User(**kwargs) + user.create(self._db) + + def get_by_name(self, name: str) -> Optional[User]: + """ + 根据用户名获取用户 + """ + return User.get_by_name(self._db, name) + + async def async_get_by_name(self, name: str) -> Optional[User]: + """ + 异步根据用户名获取用户。 + """ + return await User.async_get_by_name(self._db, name) + + async def async_get_by_id(self, user_id: int) -> Optional[User]: + """ + 异步根据用户 ID 获取用户。 + """ + return await User.async_get_by_id(self._db, user_id) + + def get_permissions(self, name: str) -> dict: + """ + 获取用户权限 + """ + user = User.get_by_name(self._db, name) + if user: + return user.permissions or {} + return {} + + def get_settings(self, name: str) -> Optional[dict]: + """ + 获取用户个性化设置,返回None表示用户不存在 + """ + user = User.get_by_name(self._db, name) + if user: + return user.settings or {} + return None + + def get_setting(self, name: str, key: str) -> Optional[str]: + """ + 获取用户个性化设置 + """ + settings = self.get_settings(name) + if settings: + return settings.get(key) + return None + + def get_name(self, **kwargs) -> Optional[str]: + """ + 根据绑定账号获取用户名称 + """ + users = self.list() + for user in users: + user_setting = user.settings + if user_setting: + for k, v in kwargs.items(): + if user_setting.get(k) == str(v): + return user.name + return None diff --git a/app/db/userconfig_oper.py b/app/db/oper/userconfig.py similarity index 96% rename from app/db/userconfig_oper.py rename to app/db/oper/userconfig.py index bdb8c77a1..fd8eb93f5 100644 --- a/app/db/userconfig_oper.py +++ b/app/db/oper/userconfig.py @@ -38,7 +38,7 @@ class UserConfigOper(DbOper, metaclass=Singleton): conf = UserConfig(username=username, key=key, value=value) conf.create(self._db) - def get(self, username: str, key: Union[str, UserConfigKey] = None) -> Any: + def get(self, username: str, key: Optional[Union[str, UserConfigKey]] = None) -> Any: """ 获取用户配置 """ diff --git a/app/db/workflow_oper.py b/app/db/oper/workflow.py similarity index 92% rename from app/db/workflow_oper.py rename to app/db/oper/workflow.py index 73e2f87d2..05eb991df 100644 --- a/app/db/workflow_oper.py +++ b/app/db/oper/workflow.py @@ -19,13 +19,13 @@ class WorkflowOper(DbOper): return True, "新增工作流成功" return False, "工作流已存在" - def get(self, wid: int) -> Workflow: + def get(self, wid: int) -> Optional[Workflow]: """ 查询单个工作流 """ return Workflow.get(self._db, wid) - async def async_get(self, wid: int) -> Workflow: + async def async_get(self, wid: int) -> Optional[Workflow]: """ 异步查询单个工作流 """ @@ -37,7 +37,7 @@ class WorkflowOper(DbOper): """ return Workflow.list(self._db) - async def async_list(self) -> Coroutine[Any, Any, Sequence[Any]]: + async def async_list(self) -> List[Workflow]: """ 异步获取所有工作流列表 """ @@ -67,7 +67,7 @@ class WorkflowOper(DbOper): """ return Workflow.get_by_name(self._db, name) - async def async_get_by_name(self, name: str) -> Workflow: + async def async_get_by_name(self, name: str) -> Optional[Workflow]: """ 异步按名称获取工作流 """ diff --git a/app/db/session.py b/app/db/session.py new file mode 100644 index 000000000..b8311c95a --- /dev/null +++ b/app/db/session.py @@ -0,0 +1,292 @@ +""" +数据库会话与异步连接池。 + +NullPool 之所以曾被硬编码,是因为它「永不复用」从而「永不跨事件循环」——asyncpg 的 +Connection 与 aiosqlite 的线程都绑定在创建它的循环上。但它同时移除了唯一的背压: +每个异步会话独占一条物理连接且无上限。 + +这里按事件循环缓存带池引擎:常驻主循环走连接池(有上限、可复用),其余循环回退 +全局 NullPool 引擎以保持跨循环安全,并由全局配额为其补上背压。 +""" +import asyncio +import inspect +import threading +import time +from contextlib import asynccontextmanager +from typing import Any, AsyncGenerator, Dict, Generator, Optional, Tuple, cast + +from sqlalchemy.ext.asyncio import AsyncEngine as SaAsyncEngine +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker +from sqlalchemy.orm import Session, scoped_session, sessionmaker + +from app.runtime.config import global_vars, settings +from app.db import engine as engine_module +from app.db.engine import (_async_pool_enabled, _get_database_engine, + get_engine, get_global_async_engine) +from app.runtime.log import logger + +# 会话工厂同样惰性:sessionmaker 在构造时就要绑定引擎,模块级构造等于把引擎的 +# 创建时机重新拉回 import 期,惰性化就白做了。 +_factory_lock = threading.RLock() +_session_factory: Optional[sessionmaker] = None +_async_session_factory: Optional[async_sessionmaker] = None +_scoped_session: Optional[scoped_session] = None + + +def get_session_factory() -> sessionmaker: + """ + 获取同步会话工厂,首次调用时按需绑定引擎。 + :return: 同步会话工厂 + """ + global _session_factory + if _session_factory is None: + with _factory_lock: + if _session_factory is None: + _session_factory = sessionmaker(bind=get_engine()) + return _session_factory + + +def get_async_session_factory() -> async_sessionmaker: + """ + 获取异步会话工厂,首次调用时按需绑定全局异步引擎。 + :return: 异步会话工厂 + """ + global _async_session_factory + if _async_session_factory is None: + with _factory_lock: + if _async_session_factory is None: + _async_session_factory = async_sessionmaker( + bind=get_global_async_engine(), class_=AsyncSession) + return _async_session_factory + + +def get_scoped_session() -> scoped_session: + """ + 获取线程局部的会话注册表,首次调用时创建。 + :return: scoped_session + """ + global _scoped_session + if _scoped_session is None: + with _factory_lock: + if _scoped_session is None: + _scoped_session = scoped_session(get_session_factory()) + return _scoped_session + + +# SessionFactory / AsyncSessionFactory / ScopedSession 这三个旧名字保留为**转发函数**, +# 而不是模块级 __getattr__(PEP 562)。 +# +# __getattr__ 只在「对模块对象取属性」时触发。`from app.db.session import ScopedSession` +# 确实会走到它,但那样每个导入方都会在 import 期把引擎创建出来,惰性化等于白做;而模块 +# 自己的函数体里写裸名字 ScopedSession 则**根本不会**触发它——那是全局名字查找,只查模块 +# __dict__ 和 builtins,直接 NameError。 +# +# 做成 __dict__ 里真实存在的函数就两头都成立:导入它不碰引擎,调用它才创建。全仓库对这三个 +# 名字的用法都是 `X()` 取一个会话,这一形式的语义与原先的 sessionmaker / scoped_session 实例 +# 完全一致;patch("app.scheduler.SessionFactory", ...) 这类既有测试替身也照旧生效。 +# +# 但仅限 `X()` 这一形式:它们不再是 sessionmaker / scoped_session 实例,因此实例上的其余接口 +# (ScopedSession.remove()、SessionFactory.configure()、AsyncSessionFactory.begin()、 +# scoped_session(SessionFactory) 等)不再可用。需要真正的工厂对象时用 +# get_scoped_session() / get_session_factory() / get_async_session_factory()。 +# +# 这三个名字已从 app/db/__init__.py 的 __all__ 中移除,属包内实现细节而非对外契约:它们建出 +# 的是绕过事务装饰器的裸会话,提交/回滚/释放全得调用方自己兜底。包内的既有调用方(scheduler、 +# postgresql 模块、Alembic 迁移脚本)走的是直接导入,不受 __all__ 影响。 +def SessionFactory() -> Session: # noqa: N802 + """ + 创建一个同步会话,引擎在首次调用时才建立。 + :return: Session + """ + return get_session_factory()() + + +def AsyncSessionFactory() -> AsyncSession: # noqa: N802 + """ + 创建一个异步会话,全局异步引擎在首次调用时才建立。 + :return: AsyncSession + """ + return get_async_session_factory()() + + +def ScopedSession() -> Session: # noqa: N802 + """ + 取当前线程的会话,引擎在首次调用时才建立。 + :return: Session + """ + return get_scoped_session()() + + +# ==================== 异步引擎的按事件循环池化 ==================== +# NullPool 之所以被选用,是因为它「永不复用」从而「永不跨事件循环」——asyncpg 的 +# Connection 与 aiosqlite 的线程都绑定在创建它的循环上,跨循环复用会抛出 +# "Task got Future attached to a different loop"。 +# +# 但它同时也移除了唯一的背压:每个异步会话独占一条物理连接且无上限。调度器以上百个 +# 线程向主循环投递协程,突发并发会直接顶穿 PostgreSQL 的 max_connections;SQLite 侧 +# 则表现为 WAL 写争用与反复 checkpoint 造成的长时间卡顿。 +# +# 这里按事件循环缓存带池引擎:常驻主循环走连接池(有上限、可复用),其余循环回退 +# 全局 NullPool 引擎以保持跨循环安全,并由 _fallback_slots 为其补上背压。 +_pooled_async_engines: Dict[int, Any] = {} +_pooled_async_lock = threading.Lock() +# 回退路径(未池化的临时循环)共享的全局连接配额。用 threading 信号量而非 +# asyncio.Semaphore:后者绑定单个事件循环,无法跨循环生效 +_fallback_slots = threading.BoundedSemaphore(max(1, settings.DB_ASYNC_FALLBACK_LIMIT)) + + +def _pooled_loop() -> Optional[Any]: + """ + 取当前可池化的事件循环,不可池化时返回 None。 + + 只认常驻主循环:它承载了绝大多数异步 DB 流量,且生命周期与进程一致, + 池中连接不会因循环销毁而失效。 + 直接读 CURRENT_EVENT_LOOP 而不用 global_vars.loop——后者在未设置时会 + 新建一个事件循环,仅为判断就产生副作用是不可接受的。 + """ + if not _async_pool_enabled(): + return None + try: + loop = asyncio.get_running_loop() + except RuntimeError: + # 没有运行中的循环,行为与池化前一致 + return None + if loop is not getattr(global_vars, "CURRENT_EVENT_LOOP", None): + return None + return loop + + +def _resolve_async_engine() -> Tuple[SaAsyncEngine, bool]: + """ + 按事件循环解析异步引擎,并一并给出它是否为池化引擎。 + + 「是否池化」必须由这里给出,不能让调用方拿 `engine is not get_global_async_engine()` + 反推:那个比较本身就会把全局引擎创建出来——池化路径压根用不到它,却因为一次身份比较 + 多出一个从未使用的活引擎,而且第一个异步请求会在事件循环内部去抢引擎创建锁。 + :return: (异步引擎, 是否池化) + """ + loop = _pooled_loop() + if loop is None: + return get_global_async_engine(), False + key = id(loop) + engine = _pooled_async_engines.get(key) + if engine is not None: + return engine, True + with _pooled_async_lock: + engine = _pooled_async_engines.get(key) + if engine is None: + engine = cast(SaAsyncEngine, _get_database_engine(is_async=True, pooled=True)) + _pooled_async_engines[key] = engine + logger.info(f"异步数据库连接池已启用: pool_size={settings.DB_ASYNC_POOL_SIZE}, " + f"max_overflow={settings.DB_ASYNC_MAX_OVERFLOW}") + return engine, True + + +def get_async_engine() -> SaAsyncEngine: + """ + 按事件循环获取异步引擎:常驻主循环用池化引擎,其余回退全局 NullPool 引擎。 + :return: 异步引擎 + """ + return _resolve_async_engine()[0] + + +async def _acquire_fallback_slot(): + """ + 为回退路径申请一个全局连接配额。 + + 池化路径由连接池自身限流;回退路径若不加约束,临时循环上的突发并发会重新 + 变得无界。信号量是线程安全且与事件循环无关的,但不能在协程里阻塞获取, + 因此用非阻塞获取 + 异步让出。 + """ + deadline = time.monotonic() + settings.DB_POOL_TIMEOUT + while not _fallback_slots.acquire(blocking=False): + if time.monotonic() >= deadline: + raise TimeoutError( + f"异步数据库连接配额已耗尽(上限 {settings.DB_ASYNC_FALLBACK_LIMIT})," + f"等待超过 {settings.DB_POOL_TIMEOUT} 秒" + ) + await asyncio.sleep(0.01) + + +@asynccontextmanager +async def async_session_scope() -> AsyncGenerator[AsyncSession, None]: + """ + 获取异步数据库会话。 + + 这是异步会话的唯一入口:池化与配额都在这里收口,调用方无需感知 + 当前运行在哪个事件循环上。 + :return: AsyncSession + """ + engine, pooled = _resolve_async_engine() + if not pooled: + await _acquire_fallback_slot() + try: + async with AsyncSession(bind=engine, expire_on_commit=False) as session: + yield session + finally: + if not pooled: + _fallback_slots.release() + + +def get_db() -> Generator: + """ + 获取数据库会话,用于WEB请求 + :return: Session + """ + db = None + try: + db = SessionFactory() + yield db + finally: + if db: + db.close() + + +async def get_async_db() -> AsyncGenerator[AsyncSession, None]: + """ + 获取异步数据库会话,用于WEB请求 + :return: AsyncSession + """ + async with async_session_scope() as session: + yield session + + +async def _dispose_engine(engine: Any, label: str) -> None: + """ + 释放单个引擎,异常只打印不上抛,避免拖累其余引擎的释放。 + + 同步引擎的 dispose 是普通函数、异步引擎的返回协程,这里一并处理, + 好让三类引擎共用同一套异常隔离。 + :param engine: 待释放的引擎 + :param label: 出错时用于定位是哪一个引擎 + """ + try: + result = engine.dispose() + if inspect.isawaitable(result): + await result + except Exception as err: # noqa: BLE001 + print(f"Error while disposing {label}: {err}") + + +async def close_database(): + """ + 关闭所有数据库连接并清理资源。 + + 逐个引擎隔离异常,而不是外面套一个大 try:套大 try 时同步引擎 dispose 一抛错, + 异步引擎与全部池化引擎就都跳过了释放——一条坏连接拖着其余连接一起泄漏, + 正是关停路径最不该出现的失败方式。 + """ + # 只释放确实创建过的引擎:惰性之后,为了 dispose 而先把引擎创建出来毫无意义, + # 还会在从未用过数据库的进程里凭空连一次库 + sync_engine = engine_module.peek_sync_engine() + if sync_engine is not None: + await _dispose_engine(sync_engine, "sync engine") + async_engine = engine_module.peek_async_engine() + if async_engine is not None: + await _dispose_engine(async_engine, "global async engine") + # 释放按事件循环缓存的池化引擎 + with _pooled_async_lock: + engines = list(_pooled_async_engines.values()) + _pooled_async_engines.clear() + for engine in engines: + await _dispose_engine(engine, "pooled async engine") diff --git a/app/db/subscribe_oper.py b/app/db/subscribe_oper.py deleted file mode 100644 index 93dcc029f..000000000 --- a/app/db/subscribe_oper.py +++ /dev/null @@ -1,317 +0,0 @@ -import time -from typing import Tuple, List, Optional - -from app.domain.context import MediaInfo, MusicInfo -from app.db import DbOper -from app.db.models.subscribe import Subscribe -from app.db.models.subscribehistory import SubscribeHistory -from app.schemas.types import MUSIC_ENTITY_ALBUM, MediaSource, MediaType -from app.domain.media import normalize_media_identity_payload, resolve_media_identity - -INTEGER_FLAG_FIELDS = ("best_version", "best_version_full", "search_imdbid", "manual_total_episode") - - -def _normalize_integer_flags(payload: dict, fields: Tuple[str, ...] = INTEGER_FLAG_FIELDS) -> dict: - """ - 将历史兼容的布尔开关转换为整型值,避免 PostgreSQL 严格类型检查失败。 - """ - normalized_payload = dict(payload) - for field in fields: - if isinstance(normalized_payload.get(field), bool): - normalized_payload[field] = int(normalized_payload[field]) - return normalized_payload - - -def _normalize_year(year: Optional[int | str]) -> Optional[str]: - """ - 订阅表的 year 列为字符串类型,而识别链路的媒体年份可能是数字 - (音乐等来源),写库前统一转换为字符串避免数据库类型错误。 - """ - if year is None: - return None - return str(year) - - -def _music_subscription_fields(mediainfo: MediaInfo | MusicInfo) -> dict: - """从标准媒体信息提取音乐订阅需要持久化的专辑级字段。""" - if mediainfo.type != MediaType.MUSIC: - return {"music_type": None, "total_tracks": None} - music_type = getattr(mediainfo, "music_type", None) - return { - "music_type": music_type, - "total_tracks": getattr(mediainfo, "total_tracks", None) - if music_type == MUSIC_ENTITY_ALBUM else None, - } - - -class SubscribeOper(DbOper): - """ - 订阅管理 - """ - - def add(self, mediainfo: MediaInfo | MusicInfo, **kwargs) -> Tuple[int, str]: - """ - 新增订阅 - """ - owner_scope = bool(kwargs.pop("owner_scope", False)) - username = kwargs.get("username") if owner_scope else None - media_source, media_id = resolve_media_identity( - media=mediainfo, - media_source=kwargs.get("media_source"), - media_id=kwargs.get("media_id"), - ) - if not media_source or not media_id: - return 0, "媒体身份不完整" - identity_params = { - "media_source": str(media_source), - "media_id": media_id, - "music_type": getattr(mediainfo, "music_type", None) - if mediainfo.type == MediaType.MUSIC else None, - "season": kwargs.get("season"), - "episode_group": mediainfo.episode_group, - } - if username: - subscribe = Subscribe.exists_by_username(self._db, - username=username, - **identity_params) - else: - subscribe = Subscribe.exists(self._db, **identity_params) - kwargs.update({ - "name": mediainfo.title, - "year": _normalize_year(mediainfo.year), - "type": mediainfo.type.value, - "media_source": str(media_source), - "media_id": media_id, - "episode_group": mediainfo.episode_group, - "poster": mediainfo.get_poster_image(), - "backdrop": mediainfo.get_backdrop_image(), - "vote": mediainfo.vote_average, - "description": mediainfo.overview, - "search_imdbid": 1 if kwargs.get('search_imdbid') else 0, - "date": time.strftime("%Y-%m-%d %H:%M:%S", time.localtime()) - }) - kwargs.update(_music_subscription_fields(mediainfo)) - kwargs = _normalize_integer_flags(kwargs) - if not subscribe: - subscribe = Subscribe(**kwargs) - subscribe.create(self._db) - # 查询订阅 - if username: - subscribe = Subscribe.exists_by_username(self._db, - username=username, - **identity_params) - else: - subscribe = Subscribe.exists(self._db, **identity_params) - return subscribe.id, "新增订阅成功" - else: - return subscribe.id, "订阅已存在" - - async def async_add(self, mediainfo: MediaInfo | MusicInfo, **kwargs) -> Tuple[int, str]: - """ - 异步新增订阅 - """ - owner_scope = bool(kwargs.pop("owner_scope", False)) - username = kwargs.get("username") if owner_scope else None - media_source, media_id = resolve_media_identity( - media=mediainfo, - media_source=kwargs.get("media_source"), - media_id=kwargs.get("media_id"), - ) - if not media_source or not media_id: - return 0, "媒体身份不完整" - identity_params = { - "media_source": str(media_source), - "media_id": media_id, - "music_type": getattr(mediainfo, "music_type", None) - if mediainfo.type == MediaType.MUSIC else None, - "season": kwargs.get("season"), - "episode_group": mediainfo.episode_group, - } - if username: - subscribe = await Subscribe.async_exists_by_username(self._db, - username=username, - **identity_params) - else: - subscribe = await Subscribe.async_exists(self._db, **identity_params) - kwargs.update({ - "name": mediainfo.title, - "year": _normalize_year(mediainfo.year), - "type": mediainfo.type.value, - "media_source": str(media_source), - "media_id": media_id, - "episode_group": mediainfo.episode_group, - "poster": mediainfo.get_poster_image(), - "backdrop": mediainfo.get_backdrop_image(), - "vote": mediainfo.vote_average, - "description": mediainfo.overview, - "search_imdbid": 1 if kwargs.get('search_imdbid') else 0, - "date": time.strftime("%Y-%m-%d %H:%M:%S", time.localtime()) - }) - kwargs.update(_music_subscription_fields(mediainfo)) - kwargs = _normalize_integer_flags(kwargs) - if not subscribe: - subscribe = Subscribe(**kwargs) - await subscribe.async_create(self._db) - # 查询订阅 - if username: - subscribe = await Subscribe.async_exists_by_username(self._db, - username=username, - **identity_params) - else: - subscribe = await Subscribe.async_exists(self._db, **identity_params) - return subscribe.id, "新增订阅成功" - else: - return subscribe.id, "订阅已存在" - - def exists( - self, media_source: MediaSource, media_id: str, - season: Optional[int] = None, episode_group: Optional[str] = None, - music_type: Optional[str] = None, - ) -> bool: - """ - 按媒体身份、季号及可选剧集组判断订阅是否存在。 - """ - identity_params = { - "media_source": media_source, - "media_id": media_id, - "music_type": music_type, - "season": season, - "episode_group": episode_group, - } - return bool(Subscribe.exists(self._db, **identity_params)) - - def get(self, sid: int) -> Subscribe: - """ - 获取订阅 - """ - return Subscribe.get(self._db, rid=sid) - - async def async_get(self, sid: int) -> Subscribe: - """ - 获取订阅 - """ - return await Subscribe.async_get(self._db, rid=sid) - - def get_by( - self, type: str, media_source: MediaSource, media_id: str, - season: Optional[str] = None, - music_type: Optional[str] = None, - ) -> Optional[Subscribe]: - """ - 根据条件查询订阅 - """ - return Subscribe.get_by( - self._db, type, media_source, media_id, season, music_type, - ) - - async def async_get_by( - self, type: str, media_source: MediaSource, media_id: str, - season: Optional[str] = None, - music_type: Optional[str] = None, - ) -> Optional[Subscribe]: - """ - 根据条件查询订阅 - """ - return await Subscribe.async_get_by( - self._db, type, media_source, media_id, season, music_type, - ) - - def list(self, state: Optional[str] = None) -> List[Subscribe]: - """ - 获取订阅列表 - """ - if state: - return Subscribe.get_by_state(self._db, state) - return Subscribe.list(self._db) - - async def async_list(self, state: Optional[str] = None) -> List[Subscribe]: - """ - 异步获取订阅列表 - """ - if state: - return await Subscribe.async_get_by_state(self._db, state) - return await Subscribe.async_list(self._db) - - def delete(self, sid: int): - """ - 删除订阅 - """ - Subscribe.delete(self._db, rid=sid) - - async def async_delete(self, sid: int): - """ - 异步删除订阅。 - """ - await Subscribe.async_delete(self._db, rid=sid) - - async def async_update(self, sid: int, payload: dict) -> Subscribe: - """ - 异步更新订阅。 - """ - subscribe = await self.async_get(sid) - if subscribe: - payload = _normalize_integer_flags(payload) - await subscribe.async_update(self._db, payload) - return subscribe - - async def async_update_filter_groups(self, sid: int, filter_groups: list) -> Subscribe: - """ - 异步更新订阅使用的过滤规则组。 - """ - return await self.async_update(sid, {"filter_groups": filter_groups}) - - def update(self, sid: int, payload: dict) -> Subscribe: - """ - 更新订阅 - """ - subscribe = self.get(sid) - if subscribe: - payload = _normalize_integer_flags(payload) - subscribe.update(self._db, payload) - return subscribe - - def list_by_username(self, username: str, state: Optional[str] = None, - mtype: Optional[str] = None) -> List[Subscribe]: - """ - 获取指定用户的订阅 - """ - return Subscribe.list_by_username(self._db, username=username, state=state, mtype=mtype) - - def list_by_type(self, mtype: str, days: Optional[int] = 7) -> Subscribe: - """ - 获取指定类型的订阅 - """ - return Subscribe.list_by_type(self._db, mtype=mtype, days=days) - - def add_history(self, **kwargs): - """ - 新增订阅 - """ - # 去除kwargs中 SubscribeHistory 没有的字段 - kwargs = {k: v for k, v in kwargs.items() if hasattr(SubscribeHistory, k)} - kwargs = normalize_media_identity_payload(kwargs) - kwargs = _normalize_integer_flags(kwargs) - # 更新完成订阅时间 - kwargs.update({"date": time.strftime("%Y-%m-%d %H:%M:%S", time.localtime())}) - # 去掉主键 - if "id" in kwargs: - kwargs.pop("id") - subscribe = SubscribeHistory(**kwargs) - subscribe.create(self._db) - - def exist_history( - self, media_source: MediaSource, media_id: str, - season: Optional[int] = None, episode_group: Optional[str] = None, - music_type: Optional[str] = None, - ) -> bool: - """ - 按媒体身份、季号及可选剧集组判断订阅历史是否存在。 - """ - identity_params = { - "media_source": media_source, - "media_id": media_id, - "music_type": music_type, - "season": season, - "episode_group": episode_group, - } - return bool(SubscribeHistory.exists(self._db, **identity_params)) diff --git a/app/domain/context.py b/app/domain/context.py index 685bd6143..f7b8e2334 100644 --- a/app/domain/context.py +++ b/app/domain/context.py @@ -20,7 +20,7 @@ from app.schemas.types import ( MediaSource, MediaType, ) -from app.domain.media import normalize_media_source, resolve_media_identity +from app.schemas.media import normalize_media_source, resolve_media_identity from app.domain.string import StringUtils BANGUMI_MOVIE_PLATFORMS = frozenset({"movie", "电影", "剧场版"}) diff --git a/app/domain/media.py b/app/domain/media.py index 36b1eb2ee..b6aa620d4 100644 --- a/app/domain/media.py +++ b/app/domain/media.py @@ -1,5 +1,16 @@ -from typing import Any, Callable, Optional, Tuple, Union +""" +媒体来源的领域策略。 +这里只留「按什么规则挑来源、什么算可订阅的音乐实体」这类策略——它们依赖运行期配置 +(注入的 SEARCH_SOURCE 提供者)与业务约定,会随产品决策变化。 + +身份的**表示规则**(别名归一、ID 去空白拒零、成对写入、媒体键前缀)不随策略变化, +已迁至 app/schemas/media.py,与两个身份 Mixin 作伴。持久化层因此不必为一条表示规则 +反向依赖领域层。本模块不做 re-export:同一个符号只应有一条 import 路径。 +""" +from typing import Callable, Optional, Tuple, Union + +from app.schemas.media import normalize_media_source from app.schemas.types import ( MUSIC_ENTITY_TYPES, MUSIC_SUBSCRIBABLE_TYPES, @@ -7,47 +18,6 @@ from app.schemas.types import ( MediaSourceSelection, ) -MEDIA_SOURCE_ALIASES = { - "tmdb": MediaSource.TMDB, - "themoviedb": MediaSource.TMDB, - "douban": MediaSource.Douban, - "bangumi": MediaSource.Bangumi, - "anilist": MediaSource.AniList, - "imdb": MediaSource.IMDb, - "tvdb": MediaSource.TVDB, - "musicbrainz": MediaSource.MusicBrainz, - "theaudiodb": MediaSource.TheAudioDB, - "audio_db": MediaSource.TheAudioDB, - "doubanmusic": MediaSource.DoubanMusic, - "douban_music": MediaSource.DoubanMusic, - "bilibili": MediaSource.Bilibili, - "mangguodiscover": MediaSource.MangoTV, - "mango_tv": MediaSource.MangoTV, - "migu": MediaSource.MiguVideo, - "migu_video": MediaSource.MiguVideo, - "tencentvideodiscover": MediaSource.TencentVideo, - "tencent_video": MediaSource.TencentVideo, - "iqiyi": MediaSource.Iqiyi, - "iqiyidiscover": MediaSource.Iqiyi, -} - -MEDIA_SOURCE_PREFIXES = { - MediaSource.TMDB: "tmdb", - MediaSource.Douban: "douban", - MediaSource.Bangumi: "bangumi", - MediaSource.AniList: "anilist", - MediaSource.IMDb: "imdb", - MediaSource.TVDB: "tvdb", - MediaSource.MusicBrainz: "musicbrainz", - MediaSource.TheAudioDB: "theaudiodb", - MediaSource.DoubanMusic: "doubanmusic", - MediaSource.Bilibili: "bilibili", - MediaSource.MangoTV: "mangguodiscover", - MediaSource.MiguVideo: "migu", - MediaSource.TencentVideo: "tencentvideodiscover", - MediaSource.Iqiyi: "iqiyidiscover", -} - MUSIC_MEDIA_SOURCE_ORDER = ( MediaSource.MusicBrainz, MediaSource.TheAudioDB, @@ -81,24 +51,6 @@ def is_music_media_source( return normalize_media_source(source) in MUSIC_MEDIA_SOURCES -def normalize_media_source( - source: Optional[Union[MediaSource, str]], -) -> Optional[MediaSource]: - """将内置别名或插件扩展标识规范化为 MediaSource。""" - if not source: - return None - if isinstance(source, MediaSource): - return source - normalized = str(source).strip().casefold() - builtin_source = MEDIA_SOURCE_ALIASES.get(normalized) - if builtin_source: - return builtin_source - try: - return MediaSource(normalized) - except ValueError: - return None - - def parse_media_source_selection(value: Optional[str]) -> Tuple[MediaSource, ...]: """ 解析 HTTP 查询参数中的逗号分隔来源,并转换为有序枚举集合。 @@ -168,94 +120,3 @@ def is_media_source_enabled( } return source_key in configured_sources return True - - -def parse_media_key( - media_key: Optional[str], -) -> Tuple[Optional[MediaSource], Optional[str]]: - """解析带来源前缀的媒体键,返回规范化数据源与原生 ID。""" - if not media_key or ":" not in str(media_key): - return None, None - prefix, media_id = str(media_key).split(":", 1) - source = normalize_media_source(prefix) - media_id = media_id.strip() - if not source or not media_id or media_id == "0": - return None, None - return source, media_id - - -def resolve_media_identity( - media: Any = None, - media_source: Optional[Union[MediaSource, str]] = None, - media_id: Optional[Any] = None, -) -> Tuple[Optional[MediaSource], Optional[str]]: - """ - 从统一媒体对象或显式字段解析主媒体身份。 - - :param media: 包含 ``media_source`` 和 ``media_id`` 的媒体对象 - :param media_source: 显式媒体来源 - :param media_id: 显式来源原生 ID - :return: 枚举化来源和字符串 ID;任一字段无效时返回空身份 - """ - normalized_source = normalize_media_source(media_source) - if media_source is not None or media_id is not None: - normalized_id = str(media_id).strip() if media_id is not None else "" - if normalized_source and normalized_id and normalized_id != "0": - return normalized_source, normalized_id - return None, None - - if media is None: - return None, None - normalized_source = normalize_media_source( - getattr(media, "media_source", None) - if not isinstance(media, dict) - else media.get("media_source") - ) - object_media_id = ( - getattr(media, "media_id", None) - if not isinstance(media, dict) - else media.get("media_id") - ) - if normalized_source and object_media_id is not None: - normalized_id = str(object_media_id).strip() - if normalized_id and normalized_id != "0": - return normalized_source, normalized_id - return None, None - - -def normalize_media_identity_payload( - payload: dict[str, Any], - *, - include_empty: bool = False, -) -> dict[str, Any]: - """ - 规范化字典中的媒体身份,保证来源与 ID 始终成对写入。 - - :param payload: 待写入或传输的字段字典 - :param include_empty: 字典未声明身份字段时,是否仍补充空身份 - :return: 复制后的规范字典;非法、半对或零值身份会被清空 - """ - normalized = dict(payload) - has_identity = "media_source" in normalized or "media_id" in normalized - if not has_identity and not include_empty: - return normalized - media_source, media_id = resolve_media_identity( - media_source=normalized.get("media_source"), - media_id=normalized.get("media_id"), - ) - normalized["media_source"] = media_source.value if media_source else None - normalized["media_id"] = media_id - return normalized - - -def build_media_key( - media_source: Optional[Union[MediaSource, str]], - media_id: Optional[Any], -) -> str: - """构造 API 使用的带来源前缀媒体键。""" - normalized_source = normalize_media_source(media_source) - normalized_id = str(media_id).strip() if media_id is not None else "" - if not normalized_source or not normalized_id or normalized_id == "0": - return "" - prefix = MEDIA_SOURCE_PREFIXES.get(normalized_source, normalized_source.value) - return f"{prefix}:{normalized_id}" diff --git a/app/domain/meta/metabase.py b/app/domain/meta/metabase.py index 2b29d9de2..663766808 100644 --- a/app/domain/meta/metabase.py +++ b/app/domain/meta/metabase.py @@ -7,7 +7,7 @@ import cn2an import regex as re from app.schemas.types import MediaSource, MediaType -from app.domain.media import resolve_media_identity +from app.schemas.media import resolve_media_identity from app.domain.string import StringUtils diff --git a/app/domain/meta/metamusic.py b/app/domain/meta/metamusic.py index b3e329b54..3e553b2b1 100644 --- a/app/domain/meta/metamusic.py +++ b/app/domain/meta/metamusic.py @@ -7,7 +7,7 @@ from typing import Any, Callable, Optional from app.domain.meta.metabase import MetaBase from app.schemas.types import MediaSource, MediaType -from app.domain.media import resolve_media_identity +from app.schemas.media import resolve_media_identity from app.domain.meta.runtime import get_metainfo_accelerator diff --git a/app/domain/metainfo.py b/app/domain/metainfo.py index 42b700c7d..2b2db4033 100644 --- a/app/domain/metainfo.py +++ b/app/domain/metainfo.py @@ -23,7 +23,7 @@ from app.domain.meta.runtime import ( get_metainfo_accelerator, ) from app.schemas.types import MediaSource, MediaType -from app.domain.media import normalize_media_source, resolve_media_identity +from app.schemas.media import normalize_media_source, resolve_media_identity _ANIME_BRACKET_RE = re.compile(r'【[+0-9XVPI-]+】\s*【', re.IGNORECASE) diff --git a/app/main.py b/app/main.py index 26dfc43b2..1630da944 100644 --- a/app/main.py +++ b/app/main.py @@ -57,7 +57,7 @@ elif SystemUtils.is_frozen(): from app.factory import app from app.runtime.config import global_vars, settings -from app.db.init import init_db, update_db +from app.startup.database_initializer import init_db, update_db # 设置进程名 setproctitle.setproctitle(settings.PROJECT_NAME) @@ -73,7 +73,7 @@ class MoviePilotServer(uvicorn.Server): # uvicorn服务 Server = MoviePilotServer(Config(app, host=settings.HOST, port=settings.PORT, - reload=settings.DEV, workers=multiprocessing.cpu_count() * 2 + 1, + reload=settings.DEV, workers=settings.API_WORKERS, timeout_graceful_shutdown=60)) diff --git a/app/modules/feishu/feishu.py b/app/modules/feishu/feishu.py index 5e8bd723d..fce1ffd26 100644 --- a/app/modules/feishu/feishu.py +++ b/app/modules/feishu/feishu.py @@ -52,7 +52,7 @@ from lark_oapi.event.callback.model.p2_card_action_trigger import ( from app.runtime.config import settings from app.domain.context import Context, MediaInfo -from app.db.user_oper import UserOper +from app.db.oper.user import UserOper from app.application.messaging.agent import matches_channel_admin from app.runtime.log import logger from app.schemas import CommingMessage, Notification diff --git a/app/modules/indexer/__init__.py b/app/modules/indexer/__init__.py index d07770f54..4f617ccf8 100644 --- a/app/modules/indexer/__init__.py +++ b/app/modules/indexer/__init__.py @@ -2,7 +2,7 @@ from datetime import datetime from typing import List, Optional, Tuple, Union from app.domain.context import SubtitleInfo, TorrentInfo -from app.db.site_oper import SiteOper +from app.db.oper.site import SiteOper from app.foundation.reflection import ModuleHelper from app.application.site.sites import SitesHelper # pylint: disable=no-name-in-module from app.runtime.log import logger @@ -17,7 +17,7 @@ from app.modules.indexer.spider.sunnypt import SunnyPTSpider from app.modules.indexer.spider.tnode import TNodeSpider from app.modules.indexer.spider.torrentleech import TorrentLeech from app.schemas.types import MediaSource -from app.domain.media import resolve_media_identity +from app.schemas.media import resolve_media_identity from app.modules.indexer.spider.yema import YemaSpider from app.schemas import SiteUserData from app.schemas.types import MediaType, ModuleType, OtherModulesType diff --git a/app/modules/indexer/spider/haidan.py b/app/modules/indexer/spider/haidan.py index a451650d3..13b37db51 100644 --- a/app/modules/indexer/spider/haidan.py +++ b/app/modules/indexer/spider/haidan.py @@ -2,7 +2,7 @@ import urllib.parse from typing import Tuple, List from app.runtime.config import settings -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas import MediaType from app.adapters.network.http import RequestUtils, AsyncRequestUtils diff --git a/app/modules/indexer/spider/hddolby.py b/app/modules/indexer/spider/hddolby.py index 5fcbe99c9..1031a8dfe 100644 --- a/app/modules/indexer/spider/hddolby.py +++ b/app/modules/indexer/spider/hddolby.py @@ -1,7 +1,7 @@ from typing import Tuple, List, Optional from app.runtime.config import settings -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas import MediaType from app.adapters.network.http import RequestUtils, AsyncRequestUtils diff --git a/app/modules/indexer/spider/mtorrent.py b/app/modules/indexer/spider/mtorrent.py index e7a21b4be..50b0976b8 100644 --- a/app/modules/indexer/spider/mtorrent.py +++ b/app/modules/indexer/spider/mtorrent.py @@ -5,7 +5,7 @@ from typing import Tuple, List, Optional from urllib.parse import urlparse from app.runtime.config import settings -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas import MediaType from app.adapters.network.http import RequestUtils, AsyncRequestUtils diff --git a/app/modules/indexer/spider/rousi.py b/app/modules/indexer/spider/rousi.py index 70da5ed08..cdbdfab87 100644 --- a/app/modules/indexer/spider/rousi.py +++ b/app/modules/indexer/spider/rousi.py @@ -3,7 +3,7 @@ import json from typing import List, Optional, Tuple from app.runtime.config import settings -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas import MediaType from app.adapters.network.http import RequestUtils, AsyncRequestUtils diff --git a/app/modules/subtitle/__init__.py b/app/modules/subtitle/__init__.py index a2e62924a..c2232c658 100644 --- a/app/modules/subtitle/__init__.py +++ b/app/modules/subtitle/__init__.py @@ -10,7 +10,7 @@ from lxml import etree from app.chain.storage import StorageChain from app.runtime.config import settings from app.domain.context import Context -from app.db.site_oper import SiteOper +from app.db.oper.site import SiteOper from app.application.site.sites import SitesHelper # pylint: disable=no-name-in-module from app.application.torrent import TorrentHelper from app.runtime.log import logger diff --git a/app/modules/themoviedb/__init__.py b/app/modules/themoviedb/__init__.py index 3c2efc452..5ea556ba8 100644 --- a/app/modules/themoviedb/__init__.py +++ b/app/modules/themoviedb/__init__.py @@ -24,11 +24,8 @@ from app.schemas.types import ( ModuleType, ) from app.adapters.network.http import RequestUtils -from app.domain.media import ( - is_media_source_enabled, - is_media_source_selected, - normalize_media_source, -) +from app.domain.media import is_media_source_enabled, is_media_source_selected +from app.schemas.media import normalize_media_source from app.foundation.text import convert as zhconv_convert diff --git a/app/modules/ugreen/ugreen.py b/app/modules/ugreen/ugreen.py index 33018931c..ef92680b9 100644 --- a/app/modules/ugreen/ugreen.py +++ b/app/modules/ugreen/ugreen.py @@ -6,7 +6,7 @@ from typing import Any, Dict, Generator, List, Mapping, Optional, Union from urllib.parse import parse_qs, urlparse from app import schemas -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.application.mediaserver import MediaServerIdentityHelper, MusicMediaServerHelper from app.runtime.log import logger from app.modules.ugreen.api import Api diff --git a/app/monitor/dispatcher.py b/app/monitor/dispatcher.py index a1415fe1e..e36ce503c 100644 --- a/app/monitor/dispatcher.py +++ b/app/monitor/dispatcher.py @@ -7,7 +7,7 @@ from typing import Any, Dict, List, Optional, Tuple from app.chain.transfer import TransferChain from app.runtime.cache import TTLCache from app.runtime.config import settings -from app.db.transferhistory_oper import TransferHistoryOper +from app.db.oper.transferhistory import TransferHistoryOper from app.application.directory import DirectoryHelper from app.application.history import (HistoryGateAction, describe_history_gate, evaluate_history_gate, is_skip_action, diff --git a/app/plugins/__init__.py b/app/plugins/__init__.py index 56761914f..66b0d02e9 100644 --- a/app/plugins/__init__.py +++ b/app/plugins/__init__.py @@ -5,8 +5,8 @@ from typing import Any, List, Dict, Tuple, Optional, Type from app.chain import ChainBase from app.core.config import settings from app.core.event import EventManager -from app.db.plugindata_oper import PluginDataOper -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.plugindata import PluginDataOper +from app.db.oper.systemconfig import SystemConfigOper from app.helper.message import MessageHelper from app.schemas import Notification, NotificationType, MessageChannel diff --git a/app/runtime/config.py b/app/runtime/config.py index 00f807a2f..d1d0d928f 100644 --- a/app/runtime/config.py +++ b/app/runtime/config.py @@ -124,6 +124,11 @@ class ConfigModel(BaseModel): # ==================== 数据库配置 ==================== # 数据库类型,支持 sqlite 和 postgresql,默认使用 sqlite + # API 服务的 worker 进程数。连接池是进程级的,每个 worker 各持一份, + # 数据库连接额度校验按它换算总用量。注意:当前主程序以单进程方式启动 + # (uvicorn.Config 的 workers 仅在多进程 supervisor 路径下生效), + # 调大此项前需先解决调度器会在每个 worker 内重复执行的问题 + API_WORKERS: int = 1 DB_TYPE: str = "sqlite" # 是否在控制台输出 SQL 语句,默认关闭 DB_ECHO: bool = False @@ -157,6 +162,26 @@ class ConfigModel(BaseModel): DB_POSTGRESQL_POOL_SIZE: int = 10 # PostgreSQL 连接池溢出数量 DB_POSTGRESQL_MAX_OVERFLOW: int = 50 + # 异步连接池类型:QueuePool / NullPool。 + # NullPool 下每个异步会话独占一条物理连接、零复用且无上限,突发并发会直接顶穿 + # PostgreSQL 的 max_connections(表现为 TooManyConnectionsError),在 SQLite 上 + # 则表现为 WAL 写争用导致的长时间卡顿。默认 QueuePool:仅对常驻主事件循环池化, + # 其余事件循环自动回退 NullPool,避免跨循环复用连接。遇到兼容问题可设为 NullPool + # 回到旧行为。 + DB_ASYNC_POOL_TYPE: str = "QueuePool" + # 异步连接池大小(每个被池化的事件循环) + DB_ASYNC_POOL_SIZE: int = 5 + # 异步连接池溢出数量(每个被池化的事件循环) + DB_ASYNC_MAX_OVERFLOW: int = 10 + # 未被池化的事件循环(临时循环)共享的全局并发连接配额。 + # 池化路径由连接池自身限流,这里只为 NullPool 兜底路径补上背压, + # 防止临时循环上的突发并发再次无界增长。 + # 实测常驻的 feishu/discord 循环不访问数据库,走此路径的只有插件与调度器 + # 兜底分支的零星调用,因此取值不必大——它直接计入连接总额度 + DB_ASYNC_FALLBACK_LIMIT: int = 10 + # 驱动级连接参数,透传给 create_engine/create_async_engine 的 connect_args。 + # 例如经 PgBouncer 事务模式接入时 asyncpg 需要 {"statement_cache_size": 0} + DB_CONNECT_ARGS: dict = Field(default_factory=dict) # ==================== 数据清理配置 ==================== # 是否启用数据表定时清理 @@ -1130,6 +1155,15 @@ class Settings(BaseSettings, ConfigModel, LogConfigModel): return f"{self.DB_POSTGRESQL_HOST}:{self.DB_POSTGRESQL_PORT}" return self.DB_POSTGRESQL_HOST + def DB_SQLITE_URL(self, driver: Optional[str] = None) -> str: + """ + SQLite 连接串。与 DB_POSTGRESQL_URL 对称,避免各调用点各自拼接后悄悄漂移 + ——迁移与应用连到不同的库文件是不会报错的。 + :param driver: 驱动名,如 aiosqlite;留空为同步驱动 + """ + scheme = "sqlite" if not driver else f"sqlite+{driver}" + return f"{scheme}:///{self.CONFIG_PATH}/user.db" + def DB_POSTGRESQL_URL(self, driver: Optional[str] = None) -> str: """按同步或异步驱动构造 PostgreSQL SQLAlchemy URL。""" scheme = "postgresql" if not driver else f"postgresql+{driver}" diff --git a/app/runtime/extensions/plugin_manager.py b/app/runtime/extensions/plugin_manager.py index ed992d1b8..4c73b1395 100644 --- a/app/runtime/extensions/plugin_manager.py +++ b/app/runtime/extensions/plugin_manager.py @@ -20,8 +20,8 @@ from starlette import status from watchfiles import watch from app import schemas -from app.db.plugindata_oper import PluginDataOper -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.plugindata import PluginDataOper +from app.db.oper.systemconfig import SystemConfigOper from app.foundation.crypto import RSAUtils from app.foundation.reflection import ObjectUtils from app.foundation.singleton import Singleton diff --git a/app/runtime/extensions/service_registry.py b/app/runtime/extensions/service_registry.py index a8e42d274..38f559270 100644 --- a/app/runtime/extensions/service_registry.py +++ b/app/runtime/extensions/service_registry.py @@ -3,7 +3,7 @@ from typing import Dict, List, Optional, Type, TypeVar, Generic, Iterator from pydantic import ValidationError from app.runtime.extensions.module_manager import ModuleManager -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas import DownloaderConf, MediaServerConf, NotificationConf, NotificationSwitchConf, ServiceInfo from app.schemas.types import NotificationType, SystemConfigKey, ModuleType diff --git a/app/scheduler.py b/app/scheduler.py index dd76c4edb..e009910e5 100644 --- a/app/scheduler.py +++ b/app/scheduler.py @@ -29,12 +29,12 @@ from app.runtime.config import settings, global_vars from app.runtime.events import Event, eventmanager from app.runtime.extensions.plugin_manager import PluginManager from app.db import SessionFactory -from app.db.agenttask_oper import AgentTaskOper +from app.db.oper.agenttask import AgentTaskOper from app.db.models.downloadhistory import DownloadHistory, DownloadFiles from app.db.models.message import Message from app.db.models.siteuserdata import SiteUserData from app.db.models.transferhistory import TransferHistory -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.application.image import WallpaperHelper from app.application.messaging.message import MessageHelper from app.runtime.progress import ProgressHelper diff --git a/app/schemas/media.py b/app/schemas/media.py index 6566c29aa..ca5e2a341 100644 --- a/app/schemas/media.py +++ b/app/schemas/media.py @@ -1,7 +1,170 @@ +""" +媒体身份的规范化与校验。 + +「media_source 与 media_id 必须成对、非零、去空白」这条不变量在两处生效:DTO 侧由 +下方两个 Mixin 在 pydantic 校验期表达,持久化侧由 app/db/models/_identity.py 在写库前 +表达。两者共用本模块的原语,形态不同但规则同源。 + +这些原语此前住在 app/domain/media.py。它们只依赖 MediaSource 与字符串,是身份的 +**表示规则**而非领域策略——真正的策略(按配置选来源、判音乐实体类型)仍留在 +app/domain/media.py。放在这里,持久化层就不必为了一条表示规则去依赖领域层。 +""" +from typing import Any, Optional, Tuple, Union + from pydantic import model_validator from app.schemas.types import MediaSource +MEDIA_SOURCE_ALIASES = { + "tmdb": MediaSource.TMDB, + "themoviedb": MediaSource.TMDB, + "douban": MediaSource.Douban, + "bangumi": MediaSource.Bangumi, + "anilist": MediaSource.AniList, + "imdb": MediaSource.IMDb, + "tvdb": MediaSource.TVDB, + "musicbrainz": MediaSource.MusicBrainz, + "theaudiodb": MediaSource.TheAudioDB, + "audio_db": MediaSource.TheAudioDB, + "doubanmusic": MediaSource.DoubanMusic, + "douban_music": MediaSource.DoubanMusic, + "bilibili": MediaSource.Bilibili, + "mangguodiscover": MediaSource.MangoTV, + "mango_tv": MediaSource.MangoTV, + "migu": MediaSource.MiguVideo, + "migu_video": MediaSource.MiguVideo, + "tencentvideodiscover": MediaSource.TencentVideo, + "tencent_video": MediaSource.TencentVideo, + "iqiyi": MediaSource.Iqiyi, + "iqiyidiscover": MediaSource.Iqiyi, +} + +MEDIA_SOURCE_PREFIXES = { + MediaSource.TMDB: "tmdb", + MediaSource.Douban: "douban", + MediaSource.Bangumi: "bangumi", + MediaSource.AniList: "anilist", + MediaSource.IMDb: "imdb", + MediaSource.TVDB: "tvdb", + MediaSource.MusicBrainz: "musicbrainz", + MediaSource.TheAudioDB: "theaudiodb", + MediaSource.DoubanMusic: "doubanmusic", + MediaSource.Bilibili: "bilibili", + MediaSource.MangoTV: "mangguodiscover", + MediaSource.MiguVideo: "migu", + MediaSource.TencentVideo: "tencentvideodiscover", + MediaSource.Iqiyi: "iqiyidiscover", +} + + +def normalize_media_source( + source: Optional[Union[MediaSource, str]], +) -> Optional[MediaSource]: + """将内置别名或插件扩展标识规范化为 MediaSource。""" + if not source: + return None + if isinstance(source, MediaSource): + return source + normalized = str(source).strip().casefold() + builtin_source = MEDIA_SOURCE_ALIASES.get(normalized) + if builtin_source: + return builtin_source + try: + return MediaSource(normalized) + except ValueError: + return None + + +def parse_media_key( + media_key: Optional[str], +) -> Tuple[Optional[MediaSource], Optional[str]]: + """解析带来源前缀的媒体键,返回规范化数据源与原生 ID。""" + if not media_key or ":" not in str(media_key): + return None, None + prefix, media_id = str(media_key).split(":", 1) + source = normalize_media_source(prefix) + media_id = media_id.strip() + if not source or not media_id or media_id == "0": + return None, None + return source, media_id + + +def build_media_key( + media_source: Optional[Union[MediaSource, str]], + media_id: Optional[Any], +) -> str: + """构造 API 使用的带来源前缀媒体键。""" + normalized_source = normalize_media_source(media_source) + normalized_id = str(media_id).strip() if media_id is not None else "" + if not normalized_source or not normalized_id or normalized_id == "0": + return "" + prefix = MEDIA_SOURCE_PREFIXES.get(normalized_source, normalized_source.value) + return f"{prefix}:{normalized_id}" + + +def resolve_media_identity( + media: Any = None, + media_source: Optional[Union[MediaSource, str]] = None, + media_id: Optional[Any] = None, +) -> Tuple[Optional[MediaSource], Optional[str]]: + """ + 从统一媒体对象或显式字段解析主媒体身份。 + + :param media: 包含 ``media_source`` 和 ``media_id`` 的媒体对象 + :param media_source: 显式媒体来源 + :param media_id: 显式来源原生 ID + :return: 枚举化来源和字符串 ID;任一字段无效时返回空身份 + """ + normalized_source = normalize_media_source(media_source) + if media_source is not None or media_id is not None: + normalized_id = str(media_id).strip() if media_id is not None else "" + if normalized_source and normalized_id and normalized_id != "0": + return normalized_source, normalized_id + return None, None + + if media is None: + return None, None + normalized_source = normalize_media_source( + getattr(media, "media_source", None) + if not isinstance(media, dict) + else media.get("media_source") + ) + object_media_id = ( + getattr(media, "media_id", None) + if not isinstance(media, dict) + else media.get("media_id") + ) + if normalized_source and object_media_id is not None: + normalized_id = str(object_media_id).strip() + if normalized_id and normalized_id != "0": + return normalized_source, normalized_id + return None, None + + +def normalize_media_identity_payload( + payload: dict[str, Any], + *, + include_empty: bool = False, +) -> dict[str, Any]: + """ + 规范化字典中的媒体身份,保证来源与 ID 始终成对写入。 + + :param payload: 待写入或传输的字段字典 + :param include_empty: 字典未声明身份字段时,是否仍补充空身份 + :return: 复制后的规范字典;非法、半对或零值身份会被清空 + """ + normalized = dict(payload) + has_identity = "media_source" in normalized or "media_id" in normalized + if not has_identity and not include_empty: + return normalized + media_source, media_id = resolve_media_identity( + media_source=normalized.get("media_source"), + media_id=normalized.get("media_id"), + ) + normalized["media_source"] = media_source.value if media_source else None + normalized["media_id"] = media_id + return normalized + class OptionalMediaIdentityMixin: """为可选媒体身份模型统一校验内置或插件来源与原生 ID 的成对约束。""" diff --git a/app/schemas/transfer.py b/app/schemas/transfer.py index cd3793db7..6615fbfa8 100644 --- a/app/schemas/transfer.py +++ b/app/schemas/transfer.py @@ -1,5 +1,5 @@ from pathlib import Path -from typing import Any, Callable, List, Optional, Union +from typing import List, Optional, Union from pydantic import BaseModel, Field @@ -79,43 +79,10 @@ class DownloadingTorrent(DownloaderTorrent): """ -class TransferTask(OptionalMediaIdentityMixin, BaseModel): - """ - 文件整理任务 - """ - fileitem: FileItem - meta: Optional[Any] = None - mediainfo: Optional[Any] = None - media_source: Optional[MediaSource] = None - media_id: Optional[str] = None - mtype: Optional[MediaType] = None - target_directory: Optional[TransferDirectoryConf] = None - target_storage: Optional[str] = None - target_path: Optional[Path] = None - transfer_type: Optional[str] = None - scrape: Optional[bool] = False - library_type_folder: Optional[bool] = False - library_category_folder: Optional[bool] = False - episodes_info: Optional[List[TmdbEpisode]] = None - username: Optional[str] = None - downloader: Optional[str] = None - download_hash: Optional[str] = None - download_history: Optional[DownloadHistory] = None - transfer_batch_id: Optional[str] = None - manual: Optional[bool] = False - background: Optional[bool] = True - preview: Optional[bool] = False - - def to_dict(self): - """ - 返回字典 - """ - dicts = vars(self).copy() - dicts["fileitem"] = self.fileitem.model_dump() if self.fileitem else None - dicts["meta"] = self.meta.model_dump() if self.meta else None - dicts["mediainfo"] = self.mediainfo.model_dump() if self.mediainfo else None - dicts["target_directory"] = self.target_directory.model_dump() if self.target_directory else None - return dicts +# TransferTask 已迁至 app/application/transfer.py:它是整理链的进程内工作项,装的是 +# 领域对象而非 DTO,留在这里只能把两个字段标成 Any——app.schemas 命名领域类型会让 +# app.schemas -> app.schemas.transfer -> app.domain.* -> app.schemas.types -> app.schemas +# 闭环。下面的 TransferJob / TransferJobTask 才是它面向前端的投影,用本包的同名 DTO。 class TransferJobTask(BaseModel): @@ -182,18 +149,6 @@ class TransferInfo(BaseModel): return dicts -class TransferQueue(BaseModel): - """ - 异步整理队列信息 - """ - # 任务信息 - task: Optional[TransferTask] = None - # 回调函数 - callback: Optional[Callable] = None - # 整理结果 - result: Optional[TransferInfo] = None - - class EpisodeFormat(BaseModel): """ 剧集自定义识别格式 diff --git a/app/db/init.py b/app/startup/database_initializer.py similarity index 55% rename from app/db/init.py rename to app/startup/database_initializer.py index 78ce38179..4e32990ac 100644 --- a/app/db/init.py +++ b/app/startup/database_initializer.py @@ -5,7 +5,7 @@ from alembic.command import upgrade from alembic.config import Config from app.runtime.config import settings -from app.db import Engine, Base +from app.db import Base from app.runtime.log import logger @@ -13,11 +13,16 @@ def init_db(): """ 初始化数据库 """ + # 函数内导入而非模块级:写成模块级会让 import 本模块的一方也被迫拉起引擎模块。 + # 引擎一律用 get_engine() 取——旧名字 `app.db.Engine` 只为仓库外插件保留,且它一经 + # 属性访问就把引擎建出来,模块级写法会使本模块反过来依赖「数据库已在别处初始化完成」。 + from app.db.engine import get_engine + # 确保所有模型都已注册到 Base.metadata 中 import app.db.models # noqa: F401 # 全量建表 - Base.metadata.create_all(bind=Engine) # noqa + Base.metadata.create_all(bind=get_engine()) def update_db(): @@ -30,12 +35,10 @@ def update_db(): alembic_cfg.file_config = _ConfigParser(interpolation=None) alembic_cfg.set_main_option('script_location', str(script_location)) - # 根据数据库类型设置不同的URL - if settings.DB_TYPE.lower() == "postgresql": - db_url = settings.DB_POSTGRESQL_URL() - else: - db_location = settings.CONFIG_PATH / 'user.db' - db_url = f"sqlite:///{db_location}" + # 与引擎构建使用同一套 URL 推导:两处各自拼接会在配置变更时悄悄漂移, + # 导致迁移连到与应用不同的库上 + db_url = settings.DB_SQLITE_URL() if settings.DB_TYPE.lower() != "postgresql" \ + else settings.DB_POSTGRESQL_URL() alembic_cfg.set_main_option('sqlalchemy.url', db_url) upgrade(alembic_cfg, 'head') @@ -44,3 +47,4 @@ def update_db(): f'数据库更新失败:{str(error)} - {traceback.format_exc()}' ) raise + diff --git a/app/startup/lifecycle.py b/app/startup/lifecycle.py index 3dfa35552..daf0db1a1 100644 --- a/app/startup/lifecycle.py +++ b/app/startup/lifecycle.py @@ -37,6 +37,7 @@ from app.startup.scheduler_initializer import ( init_scheduler, init_plugin_scheduler, ) +from app.db import check_connection_budget, get_engine, get_global_async_engine from app.startup.transfer_initializer import replay_pending_transfers from app.startup.workflow_initializer import init_workflow, stop_workflow from app.adapters.network.http import ( @@ -88,6 +89,30 @@ async def lifespan(app: FastAPI): configure_domain_dependencies() # 存储当前循环 global_vars.set_loop(asyncio.get_event_loop()) + # 同步与异步引擎各预热一次。引擎改为惰性创建后,两者的首次创建时机都不再由启动路径 + # 决定,这一步把它们拉回来。必须排在所有 init_* 之前,两个理由: + # + # 其一,fail-fast 的落点。异步驱动缺失、异步 URL 拼错这类问题若不在这里暴露,会一路 + # 推迟到第一个异步查询——表现为用户请求 500 或调度任务静默失败,而不是启动即崩。 + # 故意不 try/except:起不来就该起不来,吞掉它等于把 fail-fast 又还回去了。而既然会抛, + # 就必须抛在 init_routers / init_modules 之前——下面的 try/finally 关停块要到 yield 处 + # 才开始,在它之后抛异常,已经初始化好的模块就拿不到 stop_modules() 了。 + # + # 其二,同步引擎的首次创建要落在单线程期。init_db() 会顺带预热它,但那只对 + # run_application() 入口成立;外部 supervisor 直挂 ASGI app(如 + # `gunicorn -k uvicorn.workers.UvicornWorker app.factory:app`)时 init_db() 根本不执行, + # 首次创建便退到运行期——而那时 init_scheduler() / init_monitor() 已经放出上百个线程, + # 引擎构建里那段 PRAGMA journal_mode 会让它们一起堵在创建锁上。 + # + # 代价:异步侧几乎为零,create_async_engine 只校验 URL 与驱动导入、不建立连接;同步侧 + # 会连一次库、设一遍 journal mode,在事件循环上阻塞一小会儿——但那一次本来就免不了, + # 放在这里至少还独占着单线程,而且此刻 uvicorn 尚未开始接请求。 + get_engine() + get_global_async_engine() + # 核算数据库连接理论峰值。各连接池是彼此独立配置的,没有任何地方核算总和, + # 超额只会在突发并发时以 TooManyConnectionsError 的形式暴露;这里在启动期 + # 就对照数据库的真实上限校验一次,把问题前移到可见的位置 + check_connection_budget() # 初始化路由 init_routers(app) # 初始化模块 diff --git a/app/startup/modules_initializer.py b/app/startup/modules_initializer.py index 364db3349..0fdedcf14 100644 --- a/app/startup/modules_initializer.py +++ b/app/startup/modules_initializer.py @@ -31,7 +31,7 @@ from app.adapters.system.resource import ( from app.application.messaging.message import MessageHelper, stop_message from app.adapters.external.server import MoviePilotServerHelper from app.db import close_database -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.command import CommandChain from app.schemas import Notification, NotificationType from app.schemas.types import SystemConfigKey diff --git a/app/testing/bootstrap.py b/app/testing/bootstrap.py index 077ed3412..2762652eb 100644 --- a/app/testing/bootstrap.py +++ b/app/testing/bootstrap.py @@ -6,8 +6,8 @@ 其中 :func:`isolate_config_dir` 为主程序与插件仓共用,``prepare_v1/v2_backend`` 与 :func:`mark_plugin_generation` 为插件仓专用。 -本模块只依赖标准库,``import`` 期不连库、不触发 ``app.db``:调用方可安全地「先 import 本模块、 -再隔离 CONFIG_DIR」,不破坏「隔离必须早于首个 ``import app.db``」这一硬约束。 +本模块只依赖标准库,``import`` 期不触发 ``app.*``:调用方可安全地「先 import 本模块、 +再隔离 CONFIG_DIR」,不破坏「隔离必须早于首个 ``import app.runtime.config``」这一硬约束。 """ from __future__ import annotations @@ -68,9 +68,11 @@ class _SitesHelperStub: def isolate_config_dir() -> str: """把 ``CONFIG_DIR`` 指向进程私有临时目录,隔离主程序真实库与配置(幂等)。 - ``import app.db`` / ``import app.chain.*`` 在 import 期即按 ``settings.CONFIG_PATH`` 连接 - ``user.db``,故本函数必须在首个 ``import app.db`` 之前调用。调用方已显式设置 ``CONFIG_DIR`` - (如 CI 指定隔离目录)时尊重之、不覆盖。 + 数据库引擎已改为惰性创建,``import app.db`` 本身不再连库;但 ``settings`` 是在 + ``import app.runtime.config`` 时构造的,那一刻就把 ``CONFIG_DIR`` 读进字段并建好配置子目录, + 之后再改环境变量对 ``settings.CONFIG_PATH`` 毫无影响——引擎晚点才建,连的仍是真实 ``user.db``。 + 故本函数必须早于首个牵入 ``app.runtime.config`` 的 import(``app.db`` / ``app.chain.*`` 都会牵入)。 + 调用方已显式设置 ``CONFIG_DIR``(如 CI 指定隔离目录)时尊重之、不覆盖。 :return: 实际生效的 CONFIG_DIR 绝对路径 """ @@ -89,13 +91,19 @@ def isolate_config_dir() -> str: """进程退出时释放 SQLite 连接池再删临时目录。 默认参数绑定 ``rmtree``/``path``/``sys_mod``:解释器关停期标准库模块可能已被回收为 ``None``, - 绑定后仍可安全调用。先 ``Engine.dispose`` 释放 ``user.db`` 连接,规避 Windows 下 + 绑定后仍可安全调用。先释放已建立的 ``user.db`` 连接,规避 Windows 下 文件锁导致 ``rmtree`` 静默失败(``ignore_errors``)、残留临时目录。 + + 读 ``peek_sync_engine``(有则取、无则 ``None``)而不是旧名字 ``app.db.Engine``:后者是 + 惰性解析的属性,取它会**创建**引擎——只 ``import`` 过 ``app.db`` 的进程会在解释器关停时 + 凭空连一次库,仅仅为了随后把它 dispose 掉。 """ try: - db_mod = sys_mod.modules.get("app.db") - if db_mod is not None: - db_mod.Engine.dispose() + engine_mod = sys_mod.modules.get("app.db.engine") + peek = getattr(engine_mod, "peek_sync_engine", None) + engine = peek() if peek is not None else None + if engine is not None: + engine.dispose() except Exception: pass rmtree(path, ignore_errors=True) @@ -117,7 +125,8 @@ def ensure_sites_stub() -> None: ``app.application.site.sites`` 由独立仓库动态拉取,CI / 全新环境无该模块,而众多 ``app.chain.*`` / ``app.modules.*`` 在 import 期依赖它。统一补一个最小垫片,省去各测试文件各自打桩;若真实模块 已存在(本地已拉取)则用真实模块、不覆盖,不影响真实行为。须在隔离 CONFIG_DIR 之后调用, - 以免试探性 ``import app.application.site.sites`` 触发的连库落到真实库。 + 以免试探性 ``import app.application.site.sites`` 牵入 ``app.runtime.config``、 + 把配置路径定型到真实目录。 """ if "app.application.site.sites" in sys.modules: return @@ -168,7 +177,7 @@ def prepare_backend() -> None: """ isolate_config_dir() ensure_sites_stub() - from app.db.init import init_db + from app.startup.database_initializer import init_db init_db() # 缓存装饰器在测试模块导入时即创建后端,先装配隔离配置对应的适配器。 from app.startup.cache_initializer import configure_cache_dependencies diff --git a/app/workflow/__init__.py b/app/workflow/__init__.py index 88033f883..9e558836b 100644 --- a/app/workflow/__init__.py +++ b/app/workflow/__init__.py @@ -7,7 +7,7 @@ from pydantic import BaseModel from app.runtime.config import global_vars from app.runtime.events import eventmanager, Event from app.db.models import Workflow -from app.db.workflow_oper import WorkflowOper +from app.db.oper.workflow import WorkflowOper from app.foundation.reflection import ModuleHelper from app.runtime.log import logger from app.schemas import ActionContext, Action, ActionResult diff --git a/app/workflow/actions/__init__.py b/app/workflow/actions/__init__.py index 8dfc01c63..79a0e37c1 100644 --- a/app/workflow/actions/__init__.py +++ b/app/workflow/actions/__init__.py @@ -2,7 +2,7 @@ from abc import ABC, abstractmethod from typing import Any, Union from app.chain import ChainBase -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.schemas import ActionContext, ActionParams, ActionResult diff --git a/app/workflow/actions/add_subscribe.py b/app/workflow/actions/add_subscribe.py index ba6b969df..f8e8e62d3 100644 --- a/app/workflow/actions/add_subscribe.py +++ b/app/workflow/actions/add_subscribe.py @@ -2,7 +2,7 @@ from app.workflow.actions import BaseAction from app.chain.subscribe import SubscribeChain from app.runtime.config import settings, global_vars from app.domain.context import MediaInfo -from app.db.subscribe_oper import SubscribeOper +from app.db.oper.subscribe import SubscribeOper from app.runtime.log import logger from app.schemas import ActionParams, ActionContext diff --git a/app/workflow/actions/transfer_file.py b/app/workflow/actions/transfer_file.py index 52d5d4daf..2b1dbefd1 100644 --- a/app/workflow/actions/transfer_file.py +++ b/app/workflow/actions/transfer_file.py @@ -6,7 +6,7 @@ from pydantic import Field from app.workflow.actions import BaseAction from app.runtime.config import global_vars -from app.db.transferhistory_oper import TransferHistoryOper +from app.db.oper.transferhistory import TransferHistoryOper from app.schemas import ActionParams, ActionContext from app.chain.storage import StorageChain from app.chain.transfer import TransferChain diff --git a/database/versions/262735d025da_2_0_1.py b/database/versions/262735d025da_2_0_1.py index f79b50617..91c35f67a 100644 --- a/database/versions/262735d025da_2_0_1.py +++ b/database/versions/262735d025da_2_0_1.py @@ -6,7 +6,7 @@ Create Date: 2024-09-11 08:07:02.753307 """ -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.schemas.types import SystemConfigKey # revision identifiers, used by Alembic. diff --git a/database/versions/294b007932ef_2_0_0.py b/database/versions/294b007932ef_2_0_0.py index a0f8cb3c0..c1a75f515 100644 --- a/database/versions/294b007932ef_2_0_0.py +++ b/database/versions/294b007932ef_2_0_0.py @@ -12,7 +12,7 @@ from app.runtime.config import settings from app.application.security.access import get_password_hash from app.db import SessionFactory from app.db.models import * -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas.types import SystemConfigKey diff --git a/database/versions/3891a5e722a1_2_1_7.py b/database/versions/3891a5e722a1_2_1_7.py index c3989a68c..2e1e692db 100644 --- a/database/versions/3891a5e722a1_2_1_7.py +++ b/database/versions/3891a5e722a1_2_1_7.py @@ -9,7 +9,7 @@ from alembic import op import sqlalchemy as sa from sqlalchemy.dialects import sqlite -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.schemas.types import SystemConfigKey # revision identifiers, used by Alembic. diff --git a/database/versions/486e56a62dcb_2_1_5.py b/database/versions/486e56a62dcb_2_1_5.py index dc4a333fa..5d20b605e 100644 --- a/database/versions/486e56a62dcb_2_1_5.py +++ b/database/versions/486e56a62dcb_2_1_5.py @@ -7,7 +7,7 @@ Create Date: 2025-05-13 19:49:51.271319 """ import re -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.schemas.types import SystemConfigKey # revision identifiers, used by Alembic. diff --git a/database/versions/4dadad1d161a_3_0_0.py b/database/versions/4dadad1d161a_3_0_0.py index 4546528b1..899e3806c 100644 --- a/database/versions/4dadad1d161a_3_0_0.py +++ b/database/versions/4dadad1d161a_3_0_0.py @@ -6,7 +6,7 @@ Revises: e8b1c4d7a2f9 Create Date: 2026-08-10 """ -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.runtime.log import logger from app.schemas.types import SystemConfigKey diff --git a/database/versions/89d24811e894_2_1_4.py b/database/versions/89d24811e894_2_1_4.py index 7d20f6b7c..7ae0f0de6 100644 --- a/database/versions/89d24811e894_2_1_4.py +++ b/database/versions/89d24811e894_2_1_4.py @@ -6,7 +6,7 @@ Create Date: 2025-05-03 17:29:07.635618 """ -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.schemas.types import SystemConfigKey # revision identifiers, used by Alembic. diff --git a/database/versions/a295e41830a6_2_0_6.py b/database/versions/a295e41830a6_2_0_6.py index 97a21d862..d293aa319 100644 --- a/database/versions/a295e41830a6_2_0_6.py +++ b/database/versions/a295e41830a6_2_0_6.py @@ -6,7 +6,7 @@ Create Date: 2024-11-14 12:49:13.838120 """ -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.schemas.types import SystemConfigKey # revision identifiers, used by Alembic. diff --git a/database/versions/a73f2dbf5c09_2_0_4.py b/database/versions/a73f2dbf5c09_2_0_4.py index c70d17f9a..de542ce3a 100644 --- a/database/versions/a73f2dbf5c09_2_0_4.py +++ b/database/versions/a73f2dbf5c09_2_0_4.py @@ -6,7 +6,7 @@ Create Date: 2024-10-16 15:05:01.775429 """ -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.schemas.types import SystemConfigKey # revision identifiers, used by Alembic. diff --git a/database/versions/e8b1c4d7a2f9_2_2_18.py b/database/versions/e8b1c4d7a2f9_2_2_18.py index 351d074da..82bd3ff32 100644 --- a/database/versions/e8b1c4d7a2f9_2_2_18.py +++ b/database/versions/e8b1c4d7a2f9_2_2_18.py @@ -63,7 +63,7 @@ def upgrade() -> None: ]) # 只升级系统旧默认模板;用户编辑过的模板保持原样。 - from app.db.systemconfig_oper import SystemConfigOper + from app.db.oper.systemconfig import SystemConfigOper from app.schemas.types import SystemConfigKey legacy_organize = """ diff --git a/docs/rules/04-design-patterns.md b/docs/rules/04-design-patterns.md index 73f60d9f8..6fe11fc1f 100644 --- a/docs/rules/04-design-patterns.md +++ b/docs/rules/04-design-patterns.md @@ -112,18 +112,18 @@ eventmanager.send_event(EventType.TransferComplete, data_dict) **When to use:** All database reads and writes. Never issue SQLAlchemy queries directly from chain, module, or endpoint code. -**Convention:** Each SQLAlchemy model in `app/db/models/` has a corresponding `Oper` class in `app/db/_oper.py`. +**Convention:** Each SQLAlchemy model in `app/db/models/` has a corresponding `Oper` class in `app/db/oper/.py` — the two packages mirror each other file for file, so the module name carries the entity and the package carries the role. ``` -app/db/models/subscribe.py → app/db/subscribe_oper.py (SubscribeOper) -app/db/models/systemconfig.py → app/db/systemconfig_oper.py (SystemConfigOper) -app/db/models/transferhistory.py → app/db/transferhistory_oper.py (TransferHistoryOper) +app/db/models/subscribe.py → app/db/oper/subscribe.py (SubscribeOper) +app/db/models/systemconfig.py → app/db/oper/systemconfig.py (SystemConfigOper) +app/db/models/transferhistory.py → app/db/oper/transferhistory.py (TransferHistoryOper) ``` **Usage:** ```python -from app.db.subscribe_oper import SubscribeOper +from app.db.oper.subscribe import SubscribeOper oper = SubscribeOper() subscribe = oper.get(sid=1) @@ -180,11 +180,11 @@ Do not introduce new singletons unless the class genuinely manages global shared **Enum:** `SystemConfigKey` in `app/schemas/types.py` -**Oper class:** `SystemConfigOper` in `app/db/systemconfig_oper.py` +**Oper class:** `SystemConfigOper` in `app/db/oper/systemconfig.py` ```python from app.schemas.types import SystemConfigKey -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper oper = SystemConfigOper() value = oper.get(SystemConfigKey.RssUrls) @@ -199,7 +199,7 @@ oper.set(SystemConfigKey.RssUrls, ["https://..."]) **When to use:** Per-user settings that must survive across sessions but differ by user. -**Oper class:** `UserConfigOper` in `app/db/userconfig_oper.py` +**Oper class:** `UserConfigOper` in `app/db/oper/userconfig.py` Usage mirrors `SystemConfigOper` but scoped to a `user_id`. @@ -212,8 +212,8 @@ Usage mirrors `SystemConfigOper` but scoped to a `user_id`. | `module -> chain` coupling | Move orchestration into `chain` and shared logic into its owning canonical package | | `module -> module` direct calls | Use `chain` to orchestrate cross-module workflows | | Lower-level module importing a chain or manager | Register a callback/resolver from `app/startup/` or move orchestration to `chain` | -| Raw SQLAlchemy queries in endpoints or chains | Use the corresponding `*_oper.py` class | +| Raw SQLAlchemy queries in endpoints or chains | Use the corresponding Oper class in `app/db/oper/` | | Raw string keys for SystemConfig | Define and use a `SystemConfigKey` enum entry | -| HTTP requests via `requests` or `httpx` directly | Host code uses `RequestUtils` from `app/foundation/http.py`; plugins use `app.sdk.network` | +| HTTP requests via `requests` or `httpx` directly | Host code uses `RequestUtils` from `app/adapters/network/http.py`; plugins use `app.sdk.network` | *Last Updated: 2026-08-14* diff --git a/docs/rules/05-architecture.md b/docs/rules/05-architecture.md index e8c0cf9a5..1a4eb7c1e 100644 --- a/docs/rules/05-architecture.md +++ b/docs/rules/05-architecture.md @@ -204,10 +204,27 @@ depend on this established runtime root. ### DB / Oper layer -SQLAlchemy models stay under `app/db/models/`; `*_oper.py` classes encapsulate -queries. Chains, modules, application services and endpoints use Oper classes -instead of issuing SQLAlchemy queries directly. Every schema change requires an -Alembic migration under `database/versions/`. +SQLAlchemy models stay under `app/db/models/`; the data access classes live in +`app/db/oper/` and mirror them one-for-one (`models/subscribe.py` ↔ +`oper/subscribe.py`), so a filename carries only the entity and the package name +carries the role. Chains, modules, application services and endpoints use Oper +classes instead of issuing SQLAlchemy queries directly. Every schema change +requires an Alembic migration under `database/versions/`. + +Oper classes take and return persistence values, not domain objects. Translating +`MediaInfo` / `MetaBase` into a row is business logic and belongs in +`app/application/` — see `application/subscribe.py` and `application/history.py` +for the two write paths. Column-type coercion (numeric year to string, boolean +switches to integers) stays in the Oper because it follows the column, not the +caller. + +Invariants that must hold for *every* write are enforced at the mapper rather +than at each call site: `app/db/models/_identity.py` normalizes +`media_source` / `media_id` on `before_insert` / `before_update`, so a new write +path cannot forget them. Identity representation rules themselves +(alias folding, trimming, rejecting zero) live in `app/schemas/media.py` +alongside the two identity mixins; `app/domain/media.py` keeps only source +policy. `app/db` therefore has no dependency on `app/domain`. ## Composition and Compatibility Boundaries diff --git a/docs/rules/06-code-styles.md b/docs/rules/06-code-styles.md index b7e51ef37..4504eff3b 100644 --- a/docs/rules/06-code-styles.md +++ b/docs/rules/06-code-styles.md @@ -115,8 +115,8 @@ except: ## What Not To Do - Do not introduce new third-party libraries without placing them in the correct dependency entry: runtime packages in `requirements.in`, test/lint/build tooling in `requirements-dev.in`. -- Do not use `requests` or `httpx` directly for external HTTP calls - host code uses `RequestUtils` from `app/foundation/http.py`; plugins use `app.sdk.network`. -- Do not issue raw SQLAlchemy queries from chains, modules, or endpoints — use the `*_oper.py` classes. +- Do not use `requests` or `httpx` directly for external HTTP calls - host code uses `RequestUtils` from `app/adapters/network/http.py`; plugins use `app.sdk.network`. +- Do not issue raw SQLAlchemy queries from chains, modules, or endpoints — use the Oper classes in `app/db/oper/`. - Do not add TODO or FIXME without context. Only keep one if it is genuinely deferred and cannot be addressed in the current task. - Do not add noisy markers like `# change starts here`, `# important`, or `# this is a fix`. - Do not write comments that restate what the code already clearly says. diff --git a/docs/rules/09-external-response.md b/docs/rules/09-external-response.md index ec1b82ab5..0ffe42a9a 100644 --- a/docs/rules/09-external-response.md +++ b/docs/rules/09-external-response.md @@ -2,7 +2,7 @@ ## HTTP Client Conventions -**Rule:** Host outbound HTTP requests must go through `RequestUtils` from `app/foundation/http.py`. Plugins import it from `app.sdk.network`. Do not use `requests`, `httpx`, or `aiohttp` directly. +**Rule:** Host outbound HTTP requests must go through `RequestUtils` from `app/adapters/network/http.py`. Plugins import it from `app.sdk.network`. Do not use `requests`, `httpx`, or `aiohttp` directly. `RequestUtils` handles: - Proxy configuration (from `settings.PROXY_*`) diff --git a/docs/rules/10-data-and-persistent.md b/docs/rules/10-data-and-persistent.md index ea917a63e..29ab8ee50 100644 --- a/docs/rules/10-data-and-persistent.md +++ b/docs/rules/10-data-and-persistent.md @@ -48,21 +48,38 @@ alembic revision -m "describe the change" **Location:** `app/db/` -Each model has a corresponding `*_oper.py` file containing the data access class. Do not write SQLAlchemy queries directly in chain, module, or endpoint code. +Each model has a corresponding file under `app/db/oper/` containing the data access +class, mirroring `app/db/models/` one-for-one. Do not write SQLAlchemy queries +directly in chain, module, or endpoint code. | Oper Class | File | |---|---| -| `SubscribeOper` | `subscribe_oper.py` | -| `SystemConfigOper` | `systemconfig_oper.py` | -| `TransferHistoryOper` | `transferhistory_oper.py` | -| `DownloadHistoryOper` | `downloadhistory_oper.py` | -| `MediaServerOper` | `mediaserver_oper.py` | -| `UserOper` | `user_oper.py` | -| `UserConfigOper` | `userconfig_oper.py` | -| `MessageOper` | `message_oper.py` | -| `SiteOper` | `site_oper.py` | -| `PluginDataOper` | `plugindata_oper.py` | -| `WorkflowOper` | `workflow_oper.py` | +| `AgentChatOper` | `oper/agentchat.py` | +| `AgentTaskOper` | `oper/agenttask.py` | +| `DownloadFailureOper` | `oper/downloadfailure.py` | +| `DownloadHistoryOper` | `oper/downloadhistory.py` | +| `MediaServerOper` | `oper/mediaserver.py` | +| `MessageOper` | `oper/message.py` | +| `PluginDataOper` | `oper/plugindata.py` | +| `SiteOper` | `oper/site.py` | +| `SubscribeHistoryOper` | `oper/subscribehistory.py` | +| `SubscribeOper` | `oper/subscribe.py` | +| `SystemConfigOper` | `oper/systemconfig.py` | +| `TransferHistoryOper` | `oper/transferhistory.py` | +| `TransferPendingOper` | `oper/transferpending.py` | +| `UserConfigOper` | `oper/userconfig.py` | +| `UserOper` | `oper/user.py` | +| `WorkflowOper` | `oper/workflow.py` | + +Import by module (`from app.db.oper.subscribe import SubscribeOper`) — that is the +preferred form in this repository. `app/db/oper/__init__.py` also resolves class +names lazily for callers that only want a name, but it deliberately does not +eagerly re-export: several tests isolate a single Oper by stubbing it in +`sys.modules`, and an eager re-export would pull in the other fifteen and bypass +the stub. + +Oper classes accept and return persistence values. Turning a `MediaInfo` or +`MetaBase` into a row is business logic and lives in `app/application/`. **Standard Oper method conventions:** @@ -83,11 +100,11 @@ oper.delete(sid=1) # Delete by key **Enum:** `SystemConfigKey` in `app/schemas/types.py` -**Oper:** `SystemConfigOper` in `app/db/systemconfig_oper.py` +**Oper:** `SystemConfigOper` in `app/db/oper/systemconfig.py` ```python from app.schemas.types import SystemConfigKey -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper oper = SystemConfigOper() @@ -107,7 +124,7 @@ oper.set(SystemConfigKey.RssUrls, ["https://example.com/rss"]) **Purpose:** Settings that differ per user account. Uses `UserConfigOper`. ```python -from app.db.userconfig_oper import UserConfigOper +from app.db.oper.userconfig import UserConfigOper oper = UserConfigOper() value = oper.get(user_id=1, key="notification_enabled") diff --git a/docs/rules/11-quality-and-security.md b/docs/rules/11-quality-and-security.md index 5590cfae5..96f04f65b 100644 --- a/docs/rules/11-quality-and-security.md +++ b/docs/rules/11-quality-and-security.md @@ -105,7 +105,7 @@ The `API_TOKEN` value in `settings` is the source of truth. It is set at initial ## SQL Injection Prevention -- All database access goes through SQLAlchemy ORM via the `*_oper.py` classes. No raw SQL string construction. +- All database access goes through SQLAlchemy ORM via the Oper classes in `app/db/oper/`. No raw SQL string construction. - If a raw SQL query is ever genuinely necessary, use SQLAlchemy's `text()` with parameterized binds — never string interpolation. --- diff --git a/docs/testing.md b/docs/testing.md index 8e5f06a1b..a74de0c6b 100644 --- a/docs/testing.md +++ b/docs/testing.md @@ -21,7 +21,7 @@ python tests/run.py # 等价于 pytest 全量(参数透 收集任何测试模块、`import app.*` **之前**,conftest 完成两件事: -1. **临时库**:把 `CONFIG_DIR` 指向临时目录并 `init_db()` 建表。`app.db` 在导入期即按 `CONFIG_PATH` 连接 `user.db`,所以必须早于首个 `import app.*`;空库会让运行期查表报 `no such table`,故必须建表。 +1. **临时库**:把 `CONFIG_DIR` 指向临时目录并 `init_db()` 建表。引擎本身已惰性创建(`import app.db` 不再连库),但 `settings` 在 `import app.runtime.config` 那一刻就把 `CONFIG_DIR` 读进字段并建好配置子目录,之后再改环境变量对 `settings.CONFIG_PATH` 毫无影响——引擎晚点才建,连的仍是真实 `user.db`。所以隔离必须早于首个牵入 `app.runtime.config` 的 import(`app.db` / `app.chain.*` 都会牵入);空库会让运行期查表报 `no such table`,故必须建表。 2. **`app.application.site.sites` 垫片**:该模块由独立仓库动态拉取、CI 无此文件,conftest 统一补最小垫片(本地存在真实模块时优先用真实模块)。兼容层会把旧插件的 `app.helper.sites` 导入路由到同一模块。 由此推出两条**硬规范**: diff --git a/scripts/local_setup.py b/scripts/local_setup.py index 4a879f3a6..b8570921a 100644 --- a/scripts/local_setup.py +++ b/scripts/local_setup.py @@ -2408,8 +2408,8 @@ def _apply_local_system_config_inner(config_payload: dict[str, Any]) -> None: sys.path.insert(0, str(ROOT)) try: - from app.db.init import init_db, update_db - from app.db.systemconfig_oper import SystemConfigOper + from app.startup.database_initializer import init_db, update_db + from app.db.oper.systemconfig import SystemConfigOper from app.schemas.types import SystemConfigKey except ModuleNotFoundError as exc: raise RuntimeError( @@ -2498,7 +2498,7 @@ def _ensure_superuser_account_inner() -> None: from app.runtime.config import settings from app.application.security.access import get_password_hash - from app.db.user_oper import UserOper + from app.db.oper.user import UserOper username = str(settings.SUPERUSER or "").strip() username_error = _validate_superuser_name(username) @@ -2548,7 +2548,7 @@ def _ensure_superuser_account_inner() -> None: def _prepare_superuser_password_for_bootstrap() -> Optional[str]: from app.runtime.config import settings - from app.db.user_oper import UserOper + from app.db.oper.user import UserOper username = str(settings.SUPERUSER or "").strip() username_error = _validate_superuser_name(username) @@ -2571,7 +2571,7 @@ def _sync_superuser_account_inner() -> None: sys.path.insert(0, str(ROOT)) try: - from app.db.init import init_db, update_db + from app.startup.database_initializer import init_db, update_db except ModuleNotFoundError as exc: raise RuntimeError( "当前环境尚未安装 MoviePilot 运行依赖,请先执行 moviepilot install deps 或 moviepilot setup" @@ -3671,7 +3671,7 @@ def run_agent_request( sys.path.insert(0, str(ROOT)) try: - from app.db.init import init_db, update_db + from app.startup.database_initializer import init_db, update_db from app.agent import MoviePilotAgent from app.runtime.config import settings except ModuleNotFoundError as exc: diff --git a/tests/conftest.py b/tests/conftest.py index 65bc45c1a..d06fb9630 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -5,9 +5,12 @@ """ import sys -# 必须早于首个 import app.db(其在 import 期即按 CONFIG_PATH 连库):prepare_backend 内部 -# 先隔离 CONFIG_DIR、补 app.application.site.sites 垫片,再建表。app/testing 仅依赖标准库、import 不连库, -# 故此处先 import 再调用是安全的。 +import pytest + +# 必须早于首个牵入 app.runtime.config 的 import(app.db / app.chain.* 都会牵入):引擎本身已惰性, +# import app.db 不再连库,但 settings 在 import 期就把 CONFIG_DIR 读进字段并建好配置目录,之后 +# 改环境变量已经晚了。prepare_backend 内部先隔离 CONFIG_DIR、补 app.application.site.sites 垫片, +# 再建表。app/testing 仅依赖标准库、import 不触发 app.*,故此处先 import 再调用是安全的。 from app.testing.bootstrap import prepare_backend prepare_backend() @@ -16,6 +19,124 @@ prepare_backend() from app.testing.network_guard import block_real_network # noqa: E402,F401 +class DbHarness: + """真实数据库会话的测试载具。 + + ``prepare_backend`` 已把 CONFIG_DIR 指向临时目录并建好表,操作的是一次性数据库; + 但同一次 pytest 会话内所有用例共用这一个库,因此清理必须精确到行——按主键水位回收 + 用例新增的数据,而不是 truncate 整表,否则会连带删掉其他用例依赖的数据。 + + 水位法同时覆盖「被测代码自己写入的行」:只要在写入前登记过该表,其后新增的行 + 都会被回收,测试不必持有每一个模型实例的句柄。 + """ + + def __init__(self, session): + self.session = session + self._watermarks = {} + + def watermark(self, *models) -> None: + """ + 登记若干表的当前最大主键,用例结束时删除其后新增的全部行。 + :param models: 需要纳入回收的模型类 + """ + from sqlalchemy import func, select + + for model in models: + if model in self._watermarks: + continue + current = self.session.execute(select(func.max(model.id))).scalar() + self._watermarks[model] = current or 0 + + def add(self, *rows): + """ + 写入若干行并提交,返回单行或行列表。 + + 写入前自动登记水位,因此这些行以及被测代码后续新增的同表行都会被回收。 + :param rows: 待写入的模型实例 + """ + self.watermark(*{type(row) for row in rows}) + for row in rows: + self.session.add(row) + self.session.commit() + return rows[0] if len(rows) == 1 else list(rows) + + def cleanup(self) -> None: + """按水位删除本用例新增的全部行。""" + from sqlalchemy import delete + + # 用例可能因约束冲突等原因让事务处于待回滚状态,此时任何语句都会被拒绝; + # 先回滚再清理,否则清理会整体失效、数据泄漏到后续用例 + try: + self.session.rollback() + except Exception: # noqa: BLE001 会话已不可用时也要继续尝试清理 + pass + + for model, mark in self._watermarks.items(): + try: + self.session.execute(delete(model).where(model.id > mark)) + self.session.commit() + except Exception: # noqa: BLE001 清理失败不应掩盖用例本身的断言结果 + self.session.rollback() + + +@pytest.fixture +def db(): + """ + 提供真实数据库会话载具,用例结束按主键水位回收新增数据。 + + 数据库查询方法的行为(过滤、排序、分页、去重)无法用替身验证——替身只能证明 + 「调用了什么」,证明不了「查回了什么」,而 1.x Query 到 2.0 select 的改写恰恰 + 只可能在后者上出偏差。 + """ + from app.db.session import ScopedSession + + session = ScopedSession() + harness = DbHarness(session) + try: + yield harness + finally: + harness.cleanup() + session.close() + + +@pytest.fixture +def frozen_now(monkeypatch): + """ + 冻结指定模块看到的 ``time.time()``,其余时间函数原样透传标准库。 + + 形如 ``date >= now - 86400 * days`` 的时间窗查询,窗口起点要到调用那一刻才算得出来, + 不冻结就没法把数据精确摆在窗口起点上——而边界恰恰是 ``>=`` 与 ``>`` 唯一的分界, + 数据不压在边界上,比较符写错也查不出来。 + + :return: ``freeze(module) -> float``,冻结该模块的时钟并返回冻结时刻的时间戳 + """ + import time as real_time + + class _FrozenClock: + """只冻结 ``time()``,``localtime``/``strftime`` 等仍走标准库。""" + + def __init__(self, now: float): + self.now = now + + def time(self) -> float: + return self.now + + def __getattr__(self, name): + return getattr(real_time, name) + + def freeze(module) -> float: + """ + 把模块内的 ``time`` 名字换成冻结时钟。 + :param module: 被测代码所在模块(其内以 ``time.time()`` 取当前时刻) + :return: 冻结时刻的时间戳 + """ + clock = _FrozenClock(real_time.time()) + monkeypatch.setattr(module, "time", clock) + return clock.now + + return freeze + + def _report_session_cleanup_error(session, name: str, err: Exception) -> None: """记录收尾错误;原测试绿色时将会话标记为失败。""" sys.stderr.write(f"\npytest session cleanup failed: {name}: {err!r}\n") diff --git a/tests/test_agent_chat_history.py b/tests/test_agent_chat_history.py index 57c25f00a..b87ebe460 100644 --- a/tests/test_agent_chat_history.py +++ b/tests/test_agent_chat_history.py @@ -6,7 +6,7 @@ from langchain_core.messages import AIMessage, HumanMessage from app.agent import HEARTBEAT_SESSION_PREFIX, MoviePilotAgent from app.agent.memory import memory_manager -from app.db.agentchat_oper import AgentChatOper +from app.db.oper.agentchat import AgentChatOper from app.foundation.identity import SYSTEM_INTERNAL_USER_ID diff --git a/tests/test_agent_message_routing.py b/tests/test_agent_message_routing.py index c246eeec7..73a77145c 100644 --- a/tests/test_agent_message_routing.py +++ b/tests/test_agent_message_routing.py @@ -10,7 +10,7 @@ from app.agent.tools.impl.send_message import SendMessageTool from app.chain.message import MessageChain from app.runtime.config import settings from app.db import SessionFactory -from app.db.message_oper import MessageOper +from app.db.oper.message import MessageOper from app.db.models.message import Message from app.application.messaging.interaction import AgentInteractionOption, agent_interaction_manager, media_interaction_manager from app.schemas.types import MessageChannel, NotificationType diff --git a/tests/test_agent_scheduled_tasks.py b/tests/test_agent_scheduled_tasks.py index 429cca167..d5c7476ee 100644 --- a/tests/test_agent_scheduled_tasks.py +++ b/tests/test_agent_scheduled_tasks.py @@ -38,7 +38,7 @@ from app.agent.tools.impl.update_agent_task import ( from app.agent.tools.tags import ToolTag from app.runtime.config import settings from app.db import SessionFactory -from app.db.agenttask_oper import AgentTaskOper +from app.db.oper.agenttask import AgentTaskOper from app.db.models.agenttask import AgentTask from app.schemas import ScheduleInfo from app.scheduler import Scheduler diff --git a/tests/test_agent_task_runs.py b/tests/test_agent_task_runs.py index d334cb3c6..52bb63483 100644 --- a/tests/test_agent_task_runs.py +++ b/tests/test_agent_task_runs.py @@ -10,7 +10,7 @@ from sqlalchemy.exc import IntegrityError from app.agent import AgentManager from app.agent.tools.impl.query_agent_tasks import QueryAgentTasksTool from app.db import Engine, SessionFactory -from app.db.agenttask_oper import AgentTaskOper +from app.db.oper.agenttask import AgentTaskOper from app.db.models.agenttask import AgentTask from app.db.models.agenttaskrun import AgentTaskRun diff --git a/tests/test_api_authorization.py b/tests/test_api_authorization.py index 2a0e3cad3..d71dc92d5 100644 --- a/tests/test_api_authorization.py +++ b/tests/test_api_authorization.py @@ -18,7 +18,7 @@ from app.api.endpoints import system as system_endpoint from app.api.endpoints import transfer as transfer_endpoint from app.api.endpoints import user as user_endpoint from app.application.security.access import verify_resource_token -from app.db.user_oper import ( +from app.api.deps import ( get_current_active_manage_user, get_current_active_manage_user_async, get_current_active_superuser, diff --git a/tests/test_async_db_pooling.py b/tests/test_async_db_pooling.py new file mode 100644 index 000000000..4b2c191a0 --- /dev/null +++ b/tests/test_async_db_pooling.py @@ -0,0 +1,188 @@ +""" +异步数据库连接池的按事件循环池化测试。 + +NullPool 下每个异步会话独占一条物理连接、零复用且无上限:调度器以上百个线程向 +主事件循环投递协程,突发并发会直接顶穿 PostgreSQL 的 max_connections;SQLite 侧 +则表现为 WAL 写争用与反复 checkpoint 导致的长时间卡顿。 + +NullPool 被选用的唯一理由是「永不复用」从而「永不跨事件循环」——asyncpg 的 +Connection 与 aiosqlite 的线程都绑定在创建它的循环上。因此池化必须严格按循环 +隔离:常驻主循环用池,其余循环回退 NullPool。 + +这些测试固定三项不变量:只有常驻主循环被池化、池化引擎按循环隔离且可复用、 +回退路径受全局配额约束。 +""" +import asyncio +import threading + +import pytest + +# 池化实现位于 app.db.session;app.db 只做 re-export,私有符号不在其上 +import app.db.session as db_module +from app.runtime.config import global_vars, settings +from app.db.engine import _async_pool_kwargs, get_global_async_engine + +# 用 getter 而不是旧名字 AsyncEngine:后者只为仓库外插件保留,模块级导入它会在 pytest +# 的**收集期**就把全局异步引擎建出来——用例还一个没跑,引擎已经在了。getter 是同一个 +# 单例,下面那几处 `is` 断言的语义分毫不差。 + + +@pytest.fixture(autouse=True) +def _restore_state(): + """ + 每个用例后复原全局状态,避免污染其他测试。 + """ + saved_loop = global_vars.CURRENT_EVENT_LOOP + saved_engines = dict(db_module._pooled_async_engines) + yield + global_vars.CURRENT_EVENT_LOOP = saved_loop + db_module._pooled_async_engines.clear() + db_module._pooled_async_engines.update(saved_engines) + + +def test_pool_disabled_falls_back_to_nullpool(monkeypatch): + """ + 配置为 NullPool 时必须完全回到池化前的行为,作为兼容性逃生舱。 + """ + monkeypatch.setattr(settings, "DB_ASYNC_POOL_TYPE", "NullPool", raising=False) + + async def run(): + global_vars.CURRENT_EVENT_LOOP = asyncio.get_running_loop() + return db_module.get_async_engine() + + assert asyncio.run(run()) is get_global_async_engine() + + +def test_non_resident_loop_is_not_pooled(): + """ + 非常驻循环必须回退 NullPool。 + + 池中连接绑定在创建它的循环上,临时循环销毁后连接即失效;若对其池化, + 下次复用会抛 "attached to a different loop"。 + """ + async def run(): + # 当前运行的循环不是注册的常驻循环 + global_vars.CURRENT_EVENT_LOOP = None + return db_module.get_async_engine() + + assert asyncio.run(run()) is get_global_async_engine() + + +def test_no_running_loop_falls_back(): + """ + 没有运行中的事件循环时不得池化,行为与池化前一致。 + """ + assert db_module.get_async_engine() is get_global_async_engine() + + +def test_pooled_loop_gets_dedicated_engine_and_reuses_it(monkeypatch): + """ + 常驻主循环应获得独立的池化引擎,且同一循环内必须复用同一个引擎实例 + ——每次新建引擎等于每次新建一个池,池化就失去了意义。 + """ + monkeypatch.setattr(settings, "DB_ASYNC_POOL_TYPE", "QueuePool", raising=False) + + async def run(): + global_vars.CURRENT_EVENT_LOOP = asyncio.get_running_loop() + db_module._pooled_async_engines.clear() + first = db_module.get_async_engine() + second = db_module.get_async_engine() + return first, second + + first, second = asyncio.run(run()) + assert first is not get_global_async_engine(), "常驻循环没有拿到池化引擎" + assert first is second, "同一循环重复创建了引擎,池被反复丢弃" + + +def test_engine_is_isolated_per_loop(monkeypatch): + """ + 不同事件循环必须拿到不同的引擎实例,绝不能共用一个池。 + """ + monkeypatch.setattr(settings, "DB_ASYNC_POOL_TYPE", "QueuePool", raising=False) + db_module._pooled_async_engines.clear() + engines = [] + + def run_in_own_loop(): + """ + 在独立线程的独立事件循环中取一次引擎。 + """ + async def inner(): + global_vars.CURRENT_EVENT_LOOP = asyncio.get_running_loop() + engines.append(db_module.get_async_engine()) + + asyncio.run(inner()) + + for _ in range(2): + thread = threading.Thread(target=run_in_own_loop) + thread.start() + thread.join() + + assert len(engines) == 2 + assert engines[0] is not engines[1], "两个事件循环共用了同一个连接池" + + +def test_pool_kwargs_shape(): + """ + 池化时不得指定 poolclass:SQLAlchemy 需要自行选用异步适配的 + AsyncAdaptedQueuePool,显式传入同步 QueuePool 会出错。 + """ + pooled = _async_pool_kwargs(True) + assert "poolclass" not in pooled + assert pooled["pool_size"] == settings.DB_ASYNC_POOL_SIZE + assert pooled["max_overflow"] == settings.DB_ASYNC_MAX_OVERFLOW + + fallback = _async_pool_kwargs(False) + assert fallback["poolclass"].__name__ == "NullPool" + + +def test_fallback_slot_is_released(monkeypatch): + """ + 回退路径的配额必须在会话结束后归还,否则连续调用会把自己饿死。 + """ + monkeypatch.setattr(settings, "DB_ASYNC_POOL_TYPE", "NullPool", raising=False) + + async def run(): + before = db_module._fallback_slots._value + for _ in range(3): + async with db_module.async_session_scope(): + pass + return before, db_module._fallback_slots._value + + before, after = asyncio.run(run()) + assert before == after, "配额未归还,回退路径会逐步耗尽" + + +def test_fallback_slot_times_out_when_exhausted(monkeypatch): + """ + 配额耗尽时必须抛出明确错误,而不是无限等待或静默失败 + ——这正是 NullPool 缺失的背压。 + """ + monkeypatch.setattr(settings, "DB_POOL_TIMEOUT", 0.05, raising=False) + monkeypatch.setattr(settings, "DB_ASYNC_FALLBACK_LIMIT", 1, raising=False) + monkeypatch.setattr(db_module, "_fallback_slots", threading.BoundedSemaphore(1)) + + async def run(): + db_module._fallback_slots.acquire() # 占满唯一名额 + with pytest.raises(TimeoutError): + await db_module._acquire_fallback_slot() + + asyncio.run(run()) + + +def test_pooled_path_does_not_consume_fallback_quota(monkeypatch): + """ + 池化路径由连接池自身限流,不应再占用回退配额 + ——否则主循环流量会把兜底名额吃光,临时循环反而被饿死。 + """ + monkeypatch.setattr(settings, "DB_ASYNC_POOL_TYPE", "QueuePool", raising=False) + + async def run(): + global_vars.CURRENT_EVENT_LOOP = asyncio.get_running_loop() + db_module._pooled_async_engines.clear() + before = db_module._fallback_slots._value + async with db_module.async_session_scope(): + during = db_module._fallback_slots._value + return before, during + + before, during = asyncio.run(run()) + assert before == during, "池化路径不应消耗回退配额" diff --git a/tests/test_bluray.py b/tests/test_bluray.py index b096a7324..bfd5bfb72 100644 --- a/tests/test_bluray.py +++ b/tests/test_bluray.py @@ -176,7 +176,7 @@ class BluRayTest(TestCase): @patch("app.chain.scraping.ScrapingChain.metadata_img", return_value=None) # 避免获取图片 @patch("app.chain.ChainBase.__init__", return_value=None) # 避免不必要的模块初始化 - @patch("app.db.transferhistory_oper.TransferHistoryOper.get_by_src") + @patch("app.db.oper.transferhistory.TransferHistoryOper.get_by_src") @patch("app.chain.storage.StorageChain.list_files") @patch("app.chain.storage.StorageChain.get_parent_item") @patch("app.chain.storage.StorageChain.get_file_item") diff --git a/tests/test_database_index_migration.py b/tests/test_database_index_migration.py index 5c71aaf2f..9c4161a8c 100644 --- a/tests/test_database_index_migration.py +++ b/tests/test_database_index_migration.py @@ -75,8 +75,8 @@ from sqlalchemy import inspect, text from sqlalchemy.exc import IntegrityError from app.runtime.config import settings -from app.db import Engine -from app.db.init import init_db, update_db +from app.db import get_engine +from app.startup.database_initializer import init_db, update_db media_tables = {media_tables!r} legacy_identity_columns = {legacy_identity_columns!r} @@ -91,7 +91,7 @@ init_db() update_db() update_db() -with Engine.connect() as connection: +with get_engine().connect() as connection: version = connection.execute( text("SELECT version_num FROM alembic_version") ).scalar_one() diff --git a/tests/test_database_migration_startup.py b/tests/test_database_migration_startup.py index d1dd58b61..8f8877767 100644 --- a/tests/test_database_migration_startup.py +++ b/tests/test_database_migration_startup.py @@ -5,7 +5,7 @@ import uuid import pytest -from app.db import init as db_init +from app.startup import database_initializer as db_init LOCAL_SETUP_PATH = ( diff --git a/tests/test_db_base_crud.py b/tests/test_db_base_crud.py new file mode 100644 index 000000000..0024197db --- /dev/null +++ b/tests/test_db_base_crud.py @@ -0,0 +1,148 @@ +""" +ORM 基类通用增删改查的行为。 + +这几个方法被全部 22 个模型继承,是覆盖面最广的一段代码:get/list/delete/truncate +任何一个出偏差都会同时影响所有表。同步方法与其异步孪生方法必须给出相同结果, +否则同一张表经 API(异步)与经调度任务(同步)会看到不同的数据。 +""" +import asyncio + +import pytest + +from app.db.models.systemconfig import SystemConfig +from app.db.models.userconfig import UserConfig + + +@pytest.fixture(autouse=True) +def _track(db): + """把用于验证基类行为的两张表纳入用例级回收。""" + db.watermark(SystemConfig, UserConfig) + + +def test_create_persists_and_get_reads_back(db): + """ + 创建后可按主键读回,同步与异步取到同一行。 + """ + row = SystemConfig(key="base-create", value={"n": 1}) + row.create(db.session) + + assert row.id is not None + assert SystemConfig.get(db.session, row.id).key == "base-create" + assert asyncio.run(SystemConfig.async_get(rid=row.id)).key == "base-create" + + +def test_get_returns_none_for_missing_id(db): + """ + 主键不存在时返回 None,而不是抛异常或返回任意一行。 + """ + assert SystemConfig.get(db.session, -1) is None + assert asyncio.run(SystemConfig.async_get(rid=-1)) is None + + +def test_async_create_flushes_and_assigns_primary_key(db): + """ + 异步创建必须在返回前拿到主键。 + + 异步路径的调用方常常紧接着用 id 建立关联,拿到 None 会让关联静默丢失。 + """ + created = asyncio.run(SystemConfig(key="base-async-create", value={"n": 2}).async_create()) + + assert created.id is not None + assert SystemConfig.get(db.session, created.id).value == {"n": 2} + + +def test_update_writes_payload_fields(db): + """ + 更新按字典逐字段赋值并落库,同步与异步行为一致。 + """ + row = SystemConfig(key="base-update", value={"n": 1}) + row.create(db.session) + + row.update(db.session, {"value": {"n": 9}}) + assert SystemConfig.get(db.session, row.id).value == {"n": 9} + + asyncio.run(row.async_update(payload={"value": {"n": 10}})) + assert SystemConfig.get(db.session, row.id).value == {"n": 10} + + +def test_delete_removes_only_the_given_row(db): + """ + 按主键删除只影响那一行——条件失效会退化成清表。 + """ + dropped = db.add(SystemConfig(key="base-del", value={"n": 1})) + kept = db.add(SystemConfig(key="base-keep", value={"n": 2})) + + SystemConfig.delete(db.session, dropped.id) + + assert SystemConfig.get(db.session, dropped.id) is None + assert SystemConfig.get(db.session, kept.id) is not None + + +def test_async_delete_removes_only_the_given_row(db): + """ + 异步删除同样只影响目标行。 + """ + dropped = db.add(SystemConfig(key="base-async-del", value={"n": 1})) + kept = db.add(SystemConfig(key="base-async-keep", value={"n": 2})) + + asyncio.run(SystemConfig.async_delete(rid=dropped.id)) + + assert SystemConfig.get(db.session, dropped.id) is None + assert SystemConfig.get(db.session, kept.id) is not None + + +def test_async_delete_tolerates_missing_row(db): + """ + 删除不存在的行不抛异常,保持调用方的幂等语义。 + """ + asyncio.run(SystemConfig.async_delete(rid=-1)) + + +def test_list_returns_every_row_of_that_model_only(db): + """ + 列举必须限定在本模型对应的表,不能跨表。 + """ + db.add(UserConfig(username="base-user", key="k", value="v")) + + listed = UserConfig.list(db.session) + + assert any(item.username == "base-user" for item in listed) + assert all(isinstance(item, UserConfig) for item in listed) + + +def test_async_list_matches_sync_list(db): + """ + 同步与异步列举必须返回同一批主键。 + """ + db.add(UserConfig(username="base-list", key="k", value="v")) + + sync_ids = sorted(item.id for item in UserConfig.list(db.session)) + async_ids = sorted(item.id for item in asyncio.run(UserConfig.async_list())) + + assert sync_ids == async_ids + + +def test_truncate_empties_the_table(db): + """ + 清表后该模型不再有任何行,同步与异步实现须一致。 + """ + db.add(UserConfig(username="base-truncate", key="k", value="v")) + + UserConfig.truncate(db.session) + assert UserConfig.list(db.session) == [] + + db.add(UserConfig(username="base-truncate-async", key="k", value="v")) + asyncio.run(UserConfig.async_truncate()) + assert UserConfig.list(db.session) == [] + + +def test_to_dict_covers_every_mapped_column(db): + """ + 字典转换必须覆盖全部映射列——API 直接把它作为响应体返回,缺列即为接口缺字段。 + """ + row = db.add(SystemConfig(key="base-dict", value={"n": 1})) + + payload = row.to_dict() + + assert set(payload) == {column.name for column in SystemConfig.__table__.columns} + assert payload["key"] == "base-dict" diff --git a/tests/test_db_config_user_queries.py b/tests/test_db_config_user_queries.py new file mode 100644 index 000000000..b2a087c07 --- /dev/null +++ b/tests/test_db_config_user_queries.py @@ -0,0 +1,263 @@ +""" +配置表、用户表与 PassKey 表的查询行为。 + +这几张表决定「谁能登录、看到什么配置」,查错一行的后果是越权或配置串用,而不是 +一个能被日志发现的异常。同步方法都有一个已是 2.0 写法的异步孪生方法,这里对同一 +批数据同时跑两条路径并要求结果一致——同步侧改写后若有偏差,这个断言会直接暴露。 +""" +import asyncio + +import pytest + +from app.db.models.passkey import PassKey +from app.db.models.systemconfig import SystemConfig +from app.db.models.user import User +from app.db.models.userconfig import UserConfig + + +@pytest.fixture(autouse=True) +def _track(db): + """把本文件涉及的表纳入用例级回收。""" + db.watermark(SystemConfig, UserConfig, User, PassKey) + + +# --------------------------------------------------------------------------- # +# SystemConfig +# --------------------------------------------------------------------------- # + +def test_systemconfig_get_by_key_matches_async_twin(db): + """ + 按键取配置的同步与异步结果必须一致,且只命中同名键。 + """ + db.add(SystemConfig(key="mp-test-a", value={"n": 1}), + SystemConfig(key="mp-test-b", value={"n": 2})) + + found = SystemConfig.get_by_key(db.session, "mp-test-a") + assert found.value == {"n": 1} + + async_found = asyncio.run(SystemConfig.async_get_by_key(key="mp-test-a")) + assert async_found.value == found.value + + +def test_systemconfig_get_by_key_returns_none_when_absent(db): + """ + 键不存在时返回 None——调用方据此决定是否落默认值。 + """ + assert SystemConfig.get_by_key(db.session, "mp-test-missing") is None + + +def test_systemconfig_delete_by_key_removes_only_that_key(db): + """ + 按键删除只能删掉那一个键,误删会静默丢失其他配置。 + """ + db.add(SystemConfig(key="mp-test-del", value={"n": 1}), + SystemConfig(key="mp-test-keep", value={"n": 2})) + + assert SystemConfig().delete_by_key(db.session, "mp-test-del") is True + + assert SystemConfig.get_by_key(db.session, "mp-test-del") is None + assert SystemConfig.get_by_key(db.session, "mp-test-keep").value == {"n": 2} + + +def test_systemconfig_delete_by_key_tolerates_missing_key(db): + """ + 删除不存在的键返回 True 而不抛异常,保持调用方的幂等语义。 + """ + assert SystemConfig().delete_by_key(db.session, "mp-test-missing") is True + + +# --------------------------------------------------------------------------- # +# UserConfig +# --------------------------------------------------------------------------- # + +def test_userconfig_get_by_key_scopes_by_username(db): + """ + 用户配置必须同时按用户名和键命中——只按键会把别人的配置读给当前用户。 + """ + db.add(UserConfig(username="alice", key="theme", value="dark"), + UserConfig(username="bob", key="theme", value="light")) + + assert UserConfig.get_by_key(db.session, username="alice", key="theme").value == "dark" + assert UserConfig.get_by_key(db.session, username="bob", key="theme").value == "light" + assert UserConfig.get_by_key(db.session, username="carol", key="theme") is None + + +def test_userconfig_delete_by_key_removes_only_that_user(db): + """ + 删除某用户的配置不能波及同名键的其他用户。 + """ + db.add(UserConfig(username="alice", key="theme", value="dark"), + UserConfig(username="bob", key="theme", value="light")) + + assert UserConfig().delete_by_key(db.session, username="alice", key="theme") is True + + assert UserConfig.get_by_key(db.session, username="alice", key="theme") is None + assert UserConfig.get_by_key(db.session, username="bob", key="theme").value == "light" + + +def test_userconfig_delete_by_key_tolerates_missing_row(db): + """ + 删除不存在的用户配置返回 True,不抛异常。 + """ + assert UserConfig().delete_by_key(db.session, username="nobody", key="theme") is True + + +# --------------------------------------------------------------------------- # +# User +# --------------------------------------------------------------------------- # + +def test_user_lookup_by_name_and_id_matches_async_twin(db): + """ + 按名与按 ID 取用户的同步、异步结果必须指向同一行。 + + 登录链路走同步、API 依赖注入走异步,两者不一致会表现为「能登录但查不到自己」。 + """ + created = db.add(User(name="mp-test-user", email="u@example.com", + hashed_password="x", is_active=True)) + + by_name = User.get_by_name(db.session, "mp-test-user") + by_id = User.get_by_id(db.session, created.id) + assert by_name.id == by_id.id == created.id + + assert asyncio.run(User.async_get_by_name(name="mp-test-user")).id == created.id + assert asyncio.run(User.async_get_by_id(user_id=created.id)).id == created.id + + +def test_user_lookup_returns_none_when_absent(db): + """ + 查无此人时返回 None,而不是抛异常或返回任意一行。 + """ + assert User.get_by_name(db.session, "mp-test-nobody") is None + assert User.get_by_id(db.session, -1) is None + + +def test_user_delete_by_name_and_by_id_remove_only_the_target(db): + """ + 按名、按 ID 删除都只能删掉目标用户。 + """ + keep = db.add(User(name="mp-test-keep", hashed_password="x")) + drop_by_name = db.add(User(name="mp-test-drop-name", hashed_password="x")) + drop_by_id = db.add(User(name="mp-test-drop-id", hashed_password="x")) + + assert User().delete_by_name(db.session, drop_by_name.name) is True + assert User().delete_by_id(db.session, drop_by_id.id) is True + + assert User.get_by_name(db.session, "mp-test-drop-name") is None + assert User.get_by_id(db.session, drop_by_id.id) is None + assert User.get_by_id(db.session, keep.id) is not None + + +def test_user_update_otp_reports_whether_user_existed(db): + """ + 更新 OTP 必须如实反馈用户是否存在。 + + 对不存在的用户返回 True 会让上层以为二次验证已开启,实际并没有。 + """ + db.add(User(name="mp-test-otp", hashed_password="x", is_otp=False)) + + assert User().update_otp_by_name(db.session, "mp-test-otp", True, "SECRET") is True + assert User().update_otp_by_name(db.session, "mp-test-nobody", True, "SECRET") is False + + updated = User.get_by_name(db.session, "mp-test-otp") + assert (updated.is_otp, updated.otp_secret) == (True, "SECRET") + + +def test_user_async_mutations_match_sync_behaviour(db): + """ + 异步的删除与 OTP 更新必须与同步路径给出相同的存在性判断。 + """ + db.add(User(name="mp-test-async-otp", hashed_password="x", is_otp=False)) + + assert asyncio.run(User().async_update_otp_by_name( + name="mp-test-async-otp", otp=True, secret="S2")) is True + assert asyncio.run(User().async_update_otp_by_name( + name="mp-test-nobody", otp=True, secret="S2")) is False + + assert asyncio.run(User().async_delete_by_name(name="mp-test-async-otp")) is True + assert User.get_by_name(db.session, "mp-test-async-otp") is None + + +# --------------------------------------------------------------------------- # +# PassKey +# --------------------------------------------------------------------------- # + +def _passkey(user_id: int, credential_id: str, is_active: bool = True) -> PassKey: + """构造一条 PassKey 记录。""" + return PassKey(user_id=user_id, credential_id=credential_id, + public_key="pk", sign_count=0, is_active=is_active) + + +def test_passkey_listing_excludes_inactive_credentials(db): + """ + 列出用户凭据时必须排除已停用的。 + + 停用的凭据仍能被列出意味着它还会出现在登录选项里,等于停用没生效。 + """ + db.add(_passkey(9001, "cred-active-1"), + _passkey(9001, "cred-active-2"), + _passkey(9001, "cred-inactive", is_active=False), + _passkey(9002, "cred-other")) + + listed = PassKey.get_by_user_id(db.session, 9001) + + assert {p.credential_id for p in listed} == {"cred-active-1", "cred-active-2"} + assert {p.credential_id for p in asyncio.run(PassKey.async_get_by_user_id(user_id=9001))} == \ + {"cred-active-1", "cred-active-2"} + + +def test_passkey_lookup_by_credential_id_skips_inactive(db): + """ + 按凭据 ID 查找同样必须忽略停用记录,否则停用的密钥仍可完成认证。 + """ + db.add(_passkey(9003, "cred-live"), _passkey(9003, "cred-dead", is_active=False)) + + assert PassKey.get_by_credential_id(db.session, "cred-live").user_id == 9003 + assert PassKey.get_by_credential_id(db.session, "cred-dead") is None + assert asyncio.run(PassKey.async_get_by_credential_id(credential_id="cred-dead")) is None + + +def test_passkey_get_by_id_ignores_active_flag(db): + """ + 按主键取记录是管理用途,不应过滤停用状态——否则管理端看不到自己刚停用的凭据。 + """ + dead = db.add(_passkey(9004, "cred-admin", is_active=False)) + + assert PassKey.get_by_id(db.session, dead.id).credential_id == "cred-admin" + assert asyncio.run(PassKey.async_get_by_id(passkey_id=dead.id)).credential_id == "cred-admin" + + +def test_passkey_delete_requires_matching_owner(db): + """ + 删除必须同时匹配主键与所属用户。 + + 只按主键删除即为越权:任意登录用户都能删掉别人的凭据。 + """ + victim = db.add(_passkey(9005, "cred-victim")) + + assert PassKey.delete_by_id(db.session, passkey_id=victim.id, user_id=9999) is False + assert PassKey.get_by_id(db.session, victim.id) is not None + + assert PassKey.delete_by_id(db.session, passkey_id=victim.id, user_id=9005) is True + assert PassKey.get_by_id(db.session, victim.id) is None + + +def test_passkey_async_delete_enforces_the_same_ownership_rule(db): + """ + 异步删除必须与同步路径使用同一套归属判定。 + """ + victim = db.add(_passkey(9006, "cred-async-victim")) + + assert asyncio.run(PassKey.async_delete_by_id(passkey_id=victim.id, user_id=9999)) is False + assert asyncio.run(PassKey.async_delete_by_id(passkey_id=victim.id, user_id=9006)) is True + assert PassKey.get_by_id(db.session, victim.id) is None + + +def test_passkey_update_last_used_persists_sign_count(db): + """ + 签名计数必须落库——它是防重放的依据,不落库等于校验形同虚设。 + """ + key = db.add(_passkey(9007, "cred-count")) + + assert key.update_last_used(db.session, sign_count=42) is True + + assert PassKey.get_by_id(db.session, key.id).sign_count == 42 diff --git a/tests/test_db_declarative_2_0.py b/tests/test_db_declarative_2_0.py new file mode 100644 index 000000000..1a890f91f --- /dev/null +++ b/tests/test_db_declarative_2_0.py @@ -0,0 +1,242 @@ +""" +ORM 声明式写法的 2.0 迁移不变量。 + +基类由 1.x 的 @as_declarative() 迁移到 2.0 的 DeclarativeBase,列声明由 Column() +迁移到 mapped_column()。两者产出的表定义必须完全等价——迁移只改写法,不改 +schema,否则会与既有数据库和 alembic 迁移链产生偏差。 + +仓内模型已全部迁移到 Mapped[] 注解,__allow_unmapped__ 已随之移除;这些测试守住 +「标志不在、且仓内没有任何需要它的写法」这一组不变量,三条断言互为前提:标志一旦 +被加回来,下面两条 AST 守卫就会失去意义(legacy 写法将不再报错,只会静默通过)。 +""" +import ast +import re +from pathlib import Path + +import pytest +from sqlalchemy.orm import DeclarativeBase + +import app.db.models # noqa: F401 确保全部模型完成注册 +from app.db import Base +from app.db.base import get_id_column + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +DB_PACKAGE = PROJECT_ROOT / "app" / "db" + +# 匹配 Mapped[...]、orm.Mapped[...] 等带限定前缀的写法;不带下标的裸 Mapped 不算 +MAPPED_ANNOTATION = re.compile(r"^(?:[\w.]+\.)?Mapped\[") + +SQLALCHEMY = "sqlalchemy" + + +def _sqlalchemy_column_names(tree: ast.Module): + """ + 解析该模块的 import,产出「在本文件中指向 sqlalchemy.Column 的全部名字」。 + + 照字面量 "Column" 硬匹配会两头出错:一头漏掉 ``from sqlalchemy import Column as Col`` + 这类别名,另一头误伤同名的无关符号(rich.table.Column 就叫这个名字,而 rich 是本仓 + 依赖)。按 import 绑定判定,两个方向都准,代价只是多解析一遍 import。 + + 注意 ``sqlalchemy.Column`` 与 ``sqlalchemy.orm.mapped_column`` 是两个东西,这里 + 只认前者:mapped_column 名字里虽然也有 column,但它是 2.0 的正确写法,不该被拦。 + """ + names = set() + for node in ast.walk(tree): + if isinstance(node, ast.ImportFrom): + # from sqlalchemy import Column [as Col] / from sqlalchemy.sql import Column + module = node.module or "" + if module == SQLALCHEMY or module.startswith(f"{SQLALCHEMY}."): + names.update(alias.asname or alias.name + for alias in node.names if alias.name == "Column") + elif isinstance(node, ast.Import): + for alias in node.names: + if alias.name != SQLALCHEMY and not alias.name.startswith(f"{SQLALCHEMY}."): + continue + # import sqlalchemy as sa 绑定 sa;import sqlalchemy[.orm] 绑定顶层包名 + bound = alias.asname or alias.name.split(".")[0] + names.add(f"{bound}.Column") + return names + + +def _class_level_column_assignments(py_file: Path): + """ + 产出该文件中全部「类级 Column(...) 赋值」,形如 (类名, 属性名, 行号)。 + + 取 ClassDef 直接子语句中的 Assign 与 AnnAssign:前者是 1.x 的典型写法 + ``foo = Column(String)``(完全没有注解,因此上面那条注解守卫对它无感);后者兜住 + ``foo: Mapped[str] = Column(String)`` 这种注解已迁完、构造还留在 1.x 的半吊子状态。 + 赋值右侧用 ast.walk 递归找 Call,包一层(如 ``deferred(Column(...))``)也跑不掉。 + """ + tree = ast.parse(py_file.read_text(encoding="utf-8")) + column_names = _sqlalchemy_column_names(tree) + if not column_names: + return + for node in ast.walk(tree): + if not isinstance(node, ast.ClassDef): + continue + for stmt in node.body: + if not isinstance(stmt, (ast.Assign, ast.AnnAssign)) or stmt.value is None: + continue + if not any(isinstance(sub, ast.Call) and ast.unparse(sub.func) in column_names + for sub in ast.walk(stmt.value)): + continue + targets = stmt.targets if isinstance(stmt, ast.Assign) else [stmt.target] + yield (node.name, ", ".join(ast.unparse(t) for t in targets), stmt.lineno) + + +def _class_level_annotations(py_file: Path): + """ + 产出该文件中全部类级注解,形如 (类名, 属性名, 注解源码, 行号)。 + + 只取 ClassDef 直接子语句中的 AnnAssign:函数体内的局部注解、模块级注解都不算 + 类级注解;``if TYPE_CHECKING:`` 块里的注解运行期根本不存在,声明式系统也看不到, + 同样不在此列。 + """ + tree = ast.parse(py_file.read_text(encoding="utf-8")) + for node in ast.walk(tree): + if not isinstance(node, ast.ClassDef): + continue + for stmt in node.body: + if isinstance(stmt, ast.AnnAssign): + yield (node.name, ast.unparse(stmt.target), + ast.unparse(stmt.annotation), stmt.lineno) + + +def test_base_uses_declarative_base(): + """ + 基类必须是 2.0 的 DeclarativeBase,而不是 1.x 的 as_declarative 产物。 + """ + assert issubclass(Base, DeclarativeBase) + + +def test_allow_unmapped_is_not_set(): + """ + __allow_unmapped__ 必须保持缺席。 + + 它此前唯一的存在理由是「仓外插件可能继承本 Base 自定义 legacy 注解模型」;插件 + 生态确定迭代后这条理由已不成立,标志随之移除。这里断言它不存在而不是删掉用例: + 这个标志的危害在于**静默**——加回来之后,2.0 声明式系统不再拒绝未包裹在 + Mapped[] 中的类级注解,本文件另外两条 AST 守卫拦下的 legacy 写法就会一路通过 + 映射,直到运行期以「列不存在」的形式暴露。没有这条断言,谁把它加回来都没人知道。 + + 不用 getattr(..., False) is False:那样无法区分「没有这个属性」和 + 「显式设成了 False」,而后者同样是把这个开关重新引入了代码。 + """ + assert not hasattr(Base, "__allow_unmapped__") + + +def test_no_unmapped_class_level_annotations_in_db_package(): + """ + app/db 内不存在非 Mapped[] 的类级注解——这是移除 __allow_unmapped__ 的前提。 + + 上一条用例断言标志不在,本条断言仓内确实不需要它。两者缺一不可:只断言标志不在, + 则某天有人补进一条 legacy 注解、发现 import 就炸、顺手把标志加回来,上一条用例 + 会跟着被改绿;只断言注解形状,则标志被悄悄加回来时没有任何用例会响。 + + 这条用例还把一句会腐烂的注释钉成了可执行断言:base.py 的 docstring 上一版写着 + 「现有 22 个模型仍是 legacy Column() 写法」,在 329 列全部迁移完之后仍原样留了 + 很久,主动误导读者。 + + 扫全部类而非只扫 Base 子类:判定 Base 子类要么靠运行期 Base.__subclasses__(), + 要么靠 AST 解析基类名。前者会漏掉「新增了模型文件但还没接进 app/db/models/__init__.py」 + 的情况——恰恰是最可能带进 legacy 注解的场景;后者一遇 mixin 或跨文件继承就不准。 + 纯静态扫全部类没有这个盲区,而且当下不需要任何白名单:DbOper 这类非 ORM 类本身 + 就没有类级注解,天然不受影响。 + + 变红时怎么办(二选一,别直接把用例删了): + 1. 常见情况——新模型忘了用 2.0 写法,把它改成 mapped_column() + Mapped[] 即可; + 2. 若确实需要一条非映射的类级属性,用 ClassVar 显式声明——2.0 的声明式系统会 + 跳过 ClassVar,这条路不需要 __allow_unmapped__。本守卫比 SQLAlchemy 更严, + 当下 app/db 里没有这种属性,所以不预留白名单;真要引入时请显式放宽本守卫 + (在 MAPPED_ANNOTATION 之外放行 ClassVar),而不是把那个标志加回来。 + """ + offenders = [ + f"{py_file.relative_to(PROJECT_ROOT)}:{lineno} {cls_name}.{attr} -> {annotation}" + for py_file in sorted(DB_PACKAGE.rglob("*.py")) + for cls_name, attr, annotation, lineno in _class_level_annotations(py_file) + if not MAPPED_ANNOTATION.match(annotation) + ] + assert not offenders, ( + "app/db 内出现了非 Mapped[] 的类级注解,而 __allow_unmapped__ 已移除," + "声明式系统会直接拒绝它们:\n" + + "\n".join(f" {item}" for item in offenders) + + "\n请改用 mapped_column() + Mapped[];若这条注解确实不该被映射," + "用 ClassVar 声明并显式放宽本守卫,不要把 __allow_unmapped__ 加回来。" + ) + + +def test_no_legacy_column_assignments_in_db_package(): + """ + app/db 内不存在 1.x 的 Column() 列声明——一律 mapped_column() + Mapped[]。 + + 这条补的是上一条注解守卫的缝:它只校验「已有的注解是不是 Mapped[]」,对**完全没有 + 类级注解**的裸 ``foo = Column(String)`` 无感。上游新增的 app/db/models/agenttaskrun.py + 整整 254 行、零注解、全是 Column(),就是这么从守卫底下溜过去的——直到 pyright app/db + 从 0 退回 3 errors 才被发现。这一批工作的卖点是「app/db 全量迁到 SQLAlchemy 2.0」, + 守卫拦不住 1.x 写法回流,卖点就只是一次性的。 + + 与上一条分开报而不是合并:两者失败原因不同(注解形状 vs 列构造 API),修法也不同, + 合成一条只会让 offender 列表里混着两种毛病、报错文案被迫说得含糊。 + + Column 按 import 绑定识别而非按名字字面量,别名(``from sqlalchemy import Column as Col``) + 与限定写法(``sa.Column``)都算数,详见 _sqlalchemy_column_names。 + + 变红时怎么办:把 ``foo = Column(String)`` 改成 + ``foo: Mapped[str] = mapped_column(String)``(Optional 列写 Mapped[Optional[str]]), + 主键用 get_id_column()。注意 __allow_unmapped__ 已经移除,没有任何东西替你兜底: + 仓内已全量 2.0,见上一条用例。 + """ + offenders = [ + f"{py_file.relative_to(PROJECT_ROOT)}:{lineno} {cls_name}.{attr}" + for py_file in sorted(DB_PACKAGE.rglob("*.py")) + for cls_name, attr, lineno in _class_level_column_assignments(py_file) + ] + assert not offenders, ( + "app/db 内出现了 1.x 的 Column() 列声明,2.0 迁移被回流破坏:\n" + + "\n".join(f" {item}" for item in offenders) + + "\n请改用 mapped_column() + Mapped[] 注解(主键用 get_id_column())。" + "仓内已全量迁至 2.0,__allow_unmapped__ 已移除,没有兜底可依赖。" + ) + + +def test_id_column_factory_returns_mapped_column(): + """ + 主键工厂必须产出 mapped_column,否则主键仍是 legacy 构造。 + """ + column = get_id_column() + # mapped_column() 返回 MappedColumn,Column() 返回 Column + assert type(column).__name__ == "MappedColumn" + + +def test_all_models_registered_and_mapped(): + """ + 全部模型都应完成映射并拥有主键——迁移过程中最容易出现的失败是某个模型 + 因导入缺失而未注册,此时它不会报错,只是悄悄从 metadata 里消失。 + """ + tables = Base.metadata.tables + assert len(tables) >= 20, f"注册的表过少({len(tables)}),可能有模型未完成导入" + without_pk = [name for name, table in tables.items() if not table.primary_key.columns] + assert not without_pk, f"以下表缺少主键: {without_pk}" + + +@pytest.mark.parametrize("table_name", ["transferhistory", "downloadhistory", "subscribe"]) +def test_core_tables_keep_column_definitions(table_name): + """ + 核心业务表的列定义必须完整。抽查而非全量比对:全量等价性已在迁移时以 + schema 快照逐列核对过,这里只防止后续改动悄悄丢列。 + """ + table = Base.metadata.tables[table_name] + assert len(table.columns) > 5 + assert table.primary_key.columns, f"{table_name} 缺少主键" + # 主键应为自增整型 id + primary_key = list(table.primary_key.columns)[0] + assert primary_key.name == "id" + assert "INTEGER" in str(primary_key.type).upper() + + +def test_no_legacy_declarative_api_in_base(): + """ + base 模块不应再引用 1.x 的 as_declarative,避免两种风格并存。 + """ + source = Path("app/db/base.py").read_text(encoding="utf-8") + assert "as_declarative" not in source diff --git a/tests/test_db_decorator_error_paths.py b/tests/test_db_decorator_error_paths.py new file mode 100644 index 000000000..baf5dcca4 --- /dev/null +++ b/tests/test_db_decorator_error_paths.py @@ -0,0 +1,736 @@ +""" +事务装饰器的异常与回滚路径测试。 + +四个装饰器的正常路径处处都在被间接使用,出错路径却一次都没被验证过——而 +「事务中途抛异常」恰恰是惰性引擎、连接额度核算、按事件循环池化这几项改动 +共同的失败模式:出错时是否真的回滚、异常是否原样上抛、会话与配额是否仍被 +释放,任何一条失守都不会让别的用例变红,只会在生产上表现为脏事务、被顶替 +的异常,或再也拿不回来的连接配额。 + +这里全部用替身构造会话:验的是装饰器的控制流(谁被调用、谁没被调用、异常 +怎么传),不是 SQL 行为,真实会话反而会把这些信号淹掉。 +""" +from unittest.mock import MagicMock + +import pytest +from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import Session + +from app.db import decorators as decorators_module + + +class _BusinessError(Exception): + """被包装函数抛出的业务异常,用于验证它能原样抵达调用方。""" + + +class _FakeScope: + """ + async_session_scope 的替身,记录进入与退出次数。 + + 退出必须能被单独断言:回退路径的全局配额释放绑定在 __aexit__ 上, + 只 close 会话而不退出上下文会让配额永不归还。 + """ + + def __init__(self, session): + self.session = session + self.entered = 0 + self.exited = 0 + + async def __aenter__(self): + self.entered += 1 + return self.session + + async def __aexit__(self, *_exc): + self.exited += 1 + return False + + +class _FailingScope(_FakeScope): + """ + 退出时抛异常的作用域替身:模拟释放阶段才发现连接已断。 + + 异步侧的释放走 __aexit__ 而不是 close(),所以故障也只能从这里注入。 + """ + + def __init__(self, session, error): + super().__init__(session) + self.error = error + + async def __aexit__(self, *_exc): + self.exited += 1 + raise self.error + + +def _sync_session() -> MagicMock: + """ + 造一个能被 _get_args_db 认作调用方会话的同步会话替身。 + + 必须带 spec:装饰器用 isinstance(arg, Session) 判定「调用方是否已传会话」, + 裸 MagicMock 会被当成没传。 + """ + return MagicMock(spec=Session) + + +def _async_session() -> MagicMock: + """ + 造一个异步会话替身,commit/rollback/close 自动是 AsyncMock(可 await)。 + """ + return MagicMock(spec=AsyncSession) + + +def _install_scope(monkeypatch, session) -> _FakeScope: + """ + 把 async_session_scope 换成返回替身作用域,返回该作用域以便断言。 + :param monkeypatch: pytest 的 monkeypatch + :param session: 作用域内交出的会话替身 + """ + scope = _FakeScope(session) + monkeypatch.setattr(decorators_module, "async_session_scope", lambda: scope) + return scope + + +def _install_failing_scope(monkeypatch, session, error) -> _FailingScope: + """ + 同上,但作用域退出时抛出指定异常,用于验证释放故障的处理。 + :param monkeypatch: pytest 的 monkeypatch + :param session: 作用域内交出的会话替身 + :param error: __aexit__ 抛出的异常 + """ + scope = _FailingScope(session, error) + monkeypatch.setattr(decorators_module, "async_session_scope", lambda: scope) + return scope + + +def _capture_logger_errors(monkeypatch) -> list: + """ + 截获 logger.error 的消息,返回随调用不断追加的列表。 + """ + logged = [] + monkeypatch.setattr(decorators_module.logger, "error", lambda msg, *a, **kw: logged.append(msg)) + return logged + + +# ==================== db_update:同步更新 ==================== + +def test_db_update_rolls_back_and_skips_commit_on_error(): + """ + 被包装函数抛异常时必须回滚,且绝不能提交。 + + 漏掉回滚会把半截事务留在会话里:同一线程的 scoped_session 会被后续操作 + 继续复用,脏数据要么被下一次无关的 commit 顺手带进库,要么让后续语句 + 全部撞在「事务已中止」上。 + """ + db = _sync_session() + + @decorators_module.db_update + def _write(db=None): + """必定失败的更新。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + _write(db=db) + + db.rollback.assert_called_once() + db.commit.assert_not_called() + + +def test_db_update_reraises_the_original_exception_object(): + """ + 原异常必须原样上抛:类型与实例都不变,不被包装也不被吞。 + + 调用方靠异常类型分流(唯一约束冲突要重试、参数错误要报错),一旦被换成 + 别的类型,上层的 except 就再也匹配不上。 + """ + db = _sync_session() + boom = _BusinessError("boom") + + @decorators_module.db_update + def _write(db=None): + """抛出一个可辨认的异常实例。""" + raise boom + + with pytest.raises(_BusinessError) as excinfo: + _write(db=db) + + assert excinfo.value is boom, "上抛的不是原始异常实例" + + +def test_db_update_closes_self_created_session_on_error(monkeypatch): + """ + 装饰器自建的会话,异常路径下同样要关闭。 + + 释放写在 finally 里就是为了这个:把它挪进 try 的尾部,正常路径照常绿灯, + 每一次失败的写入却都会漏掉一条连接。 + """ + db = _sync_session() + monkeypatch.setattr(decorators_module, "ScopedSession", lambda: db) + + @decorators_module.db_update + def _write(db=None): + """必定失败的更新。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + _write(db=None) + + db.close.assert_called_once() + + +def test_db_update_does_not_close_caller_session_on_error(): + """ + 调用方传入的会话,异常路径下不得关闭——不是装饰器创建的,就无权释放。 + + 关掉别人的会话比泄漏更糟:调用方后面还要在同一个会话上做别的事, + 而它已经被这次失败连带关掉了。 + """ + db = _sync_session() + + @decorators_module.db_update + def _write(db=None): + """必定失败的更新。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + _write(db=db) + + db.close.assert_not_called() + + +# ==================== async_db_update:异步更新 ==================== + +@pytest.mark.asyncio +async def test_async_db_update_rolls_back_and_skips_commit_on_error(): + """ + 异步更新出错时同样必须回滚、绝不提交(且回滚是 await 的)。 + """ + db = _async_session() + + @decorators_module.async_db_update + async def _write(db=None): + """必定失败的异步更新。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + await _write(db=db) + + db.rollback.assert_awaited_once() + db.commit.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_async_db_update_reraises_the_original_exception_object(): + """ + 异步路径的原异常同样要原样上抛,类型与实例都不变。 + """ + db = _async_session() + boom = _BusinessError("boom") + + @decorators_module.async_db_update + async def _write(db=None): + """抛出一个可辨认的异常实例。""" + raise boom + + with pytest.raises(_BusinessError) as excinfo: + await _write(db=db) + + assert excinfo.value is boom, "上抛的不是原始异常实例" + + +@pytest.mark.asyncio +async def test_async_db_update_exits_scope_on_error(monkeypatch): + """ + 自建会话时,异常路径下必须退出会话作用域,而不是只关会话。 + + 回退路径(非常驻循环)的全局连接配额是在 async_session_scope 的 finally 里 + 归还的,只有走 __aexit__ 才会触发。写成 `await db.close()` 时正常路径与异常 + 路径的断言都照样绿——会话确实关了——但每一次失败都会永久吃掉一个配额名额, + 攒够上限后整个进程的异步数据库访问一起饿死。 + """ + db = _async_session() + scope = _install_scope(monkeypatch, db) + + @decorators_module.async_db_update + async def _write(db=None): + """必定失败的异步更新。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + await _write(db=None) + + assert scope.exited == 1, "异常路径没有退出会话作用域,配额不会归还" + + +@pytest.mark.asyncio +async def test_async_db_update_does_not_release_caller_session_on_error(monkeypatch): + """ + 调用方传入异步会话时,异常路径下既不建作用域也不关它的会话。 + """ + db = _async_session() + + def _boom(): + """任何建作用域的行为都是失败信号。""" + raise AssertionError("调用方已传入会话,装饰器不应再建作用域") + + monkeypatch.setattr(decorators_module, "async_session_scope", _boom) + + @decorators_module.async_db_update + async def _write(db=None): + """必定失败的异步更新。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + await _write(db=db) + + db.close.assert_not_awaited() + + +# ==================== db_query:同步查询 ==================== + +def test_db_query_does_not_touch_transaction_on_error(): + """ + 查询装饰器不管事务:出错时既不提交也不回滚。 + + 替调用方回滚会把它自己那段尚未提交的事务一并抹掉——查询只是借了会话, + 无权处置会话上正在进行的事务。 + """ + db = _sync_session() + + @decorators_module.db_query + def _read(db=None): + """必定失败的查询。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + _read(db=db) + + db.commit.assert_not_called() + db.rollback.assert_not_called() + + +def test_db_query_reraises_the_original_exception_object(): + """ + 查询出错时原异常原样上抛,类型与实例都不变。 + """ + db = _sync_session() + boom = _BusinessError("boom") + + @decorators_module.db_query + def _read(db=None): + """抛出一个可辨认的异常实例。""" + raise boom + + with pytest.raises(_BusinessError) as excinfo: + _read(db=db) + + assert excinfo.value is boom, "上抛的不是原始异常实例" + + +def test_db_query_closes_self_created_session_on_error(monkeypatch): + """ + 自建会话的查询,异常路径下同样要关闭会话。 + """ + db = _sync_session() + monkeypatch.setattr(decorators_module, "ScopedSession", lambda: db) + + @decorators_module.db_query + def _read(db=None): + """必定失败的查询。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + _read(db=None) + + db.close.assert_called_once() + + +# ==================== async_db_query:异步查询 ==================== + +@pytest.mark.asyncio +async def test_async_db_query_does_not_touch_transaction_on_error(): + """ + 异步查询出错时同样既不提交也不回滚。 + """ + db = _async_session() + + @decorators_module.async_db_query + async def _read(db=None): + """必定失败的异步查询。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + await _read(db=db) + + db.commit.assert_not_awaited() + db.rollback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_async_db_query_reraises_the_original_exception_object(): + """ + 异步查询出错时原异常原样上抛,类型与实例都不变。 + """ + db = _async_session() + boom = _BusinessError("boom") + + @decorators_module.async_db_query + async def _read(db=None): + """抛出一个可辨认的异常实例。""" + raise boom + + with pytest.raises(_BusinessError) as excinfo: + await _read(db=db) + + assert excinfo.value is boom, "上抛的不是原始异常实例" + + +@pytest.mark.asyncio +async def test_async_db_query_exits_scope_on_error(monkeypatch): + """ + 自建会话的异步查询,异常路径下必须退出会话作用域(配额同样绑在 __aexit__ 上)。 + """ + db = _async_session() + scope = _install_scope(monkeypatch, db) + + @decorators_module.async_db_query + async def _read(db=None): + """必定失败的异步查询。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + await _read(db=None) + + assert scope.exited == 1, "异常路径没有退出会话作用域,配额不会归还" + + +# ==================== 回滚本身失败 ==================== + +def test_db_update_rollback_failure_does_not_replace_original_error(): + """ + 回滚自身失败时,调用方仍须收到原始业务异常。 + + `except: db.rollback(); raise err` 的写法里,回滚一抛错就直接顶替了原始异常: + 连接断开、事务已失效这类收尾故障恰恰最容易在「出错之后」发生,于是排障时看到的 + 永远是「connection reset」,真正的业务异常连类型都被换掉,调用方按类型分流的 + except 也一并失配。 + """ + db = _sync_session() + db.rollback.side_effect = RuntimeError("connection reset") + boom = _BusinessError("boom") + + @decorators_module.db_update + def _write(db=None): + """抛出一个可辨认的异常实例。""" + raise boom + + with pytest.raises(_BusinessError) as excinfo: + _write(db=db) + + assert excinfo.value is boom, "回滚失败顶替了原始异常" + + +def test_db_update_logs_rollback_failure(monkeypatch): + """ + 回滚失败不能被静默吞掉:原始异常照常上抛,回滚故障要留下记录。 + + 否则「连接已断」这个同样重要的信号会彻底消失——只保原始异常而不记回滚故障, + 等于用一个盲区换掉另一个盲区。 + """ + logged = _capture_logger_errors(monkeypatch) + db = _sync_session() + db.rollback.side_effect = RuntimeError("connection reset") + + @decorators_module.db_update + def _write(db=None): + """必定失败的更新。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + _write(db=db) + + assert any("connection reset" in str(msg) for msg in logged), \ + f"回滚失败被静默吞掉,未留下任何记录:{logged}" + + +@pytest.mark.asyncio +async def test_async_db_update_rollback_failure_does_not_replace_original_error(): + """ + 异步回滚自身失败时,调用方同样必须收到原始业务异常。 + """ + db = _async_session() + db.rollback.side_effect = RuntimeError("connection reset") + boom = _BusinessError("boom") + + @decorators_module.async_db_update + async def _write(db=None): + """抛出一个可辨认的异常实例。""" + raise boom + + with pytest.raises(_BusinessError) as excinfo: + await _write(db=db) + + assert excinfo.value is boom, "回滚失败顶替了原始异常" + + +# ==================== 释放本身失败 ==================== +# +# 释放(同步的 db.close()、异步的 _scope.__aexit__())写在 finally 里,裸写时它一抛错 +# 就成了整个 finally 的出口:异常路径下顶替掉正在传播的业务异常,成功路径下把一次已经 +# 提交的写入变成调用方眼里的失败。三组断言分别钉住三件事——原始异常不被顶替、释放故障 +# 不被静默吞掉、成功路径的返回值不被释放故障拦截。 +# +# 第三条是本次修改**新引入**的行为(改前:func() 成功而 close() 失败,调用方收到异常; +# 改后:静默拿到返回值,故障只进日志)。它是刻意的取舍而非疏漏,因此必须有用例钉住, +# 否则下一个人会把它当 bug「修」回去。 +# +# 释放只发生在装饰器自建会话时,所以这些用例一律走自建路径(monkeypatch 掉会话来源)。 + +def test_db_update_close_failure_does_not_replace_original_error(monkeypatch): + """ + 关闭会话失败时,调用方仍须收到原始业务异常。 + + 与回滚同理:连接已断这类故障最容易出现在收尾阶段,裸写 db.close() 时它一抛错, + 调用方看到的就只剩「connection reset」,业务异常连类型都被换掉。 + """ + db = _sync_session() + db.close.side_effect = RuntimeError("close failed") + monkeypatch.setattr(decorators_module, "ScopedSession", lambda: db) + boom = _BusinessError("boom") + + @decorators_module.db_update + def _write(db=None): + """抛出一个可辨认的异常实例。""" + raise boom + + with pytest.raises(_BusinessError) as excinfo: + _write(db=None) + + assert excinfo.value is boom, "关闭会话失败顶替了原始异常" + + +def test_db_update_logs_close_failure(monkeypatch): + """ + 关闭会话失败不能被静默吞掉:不上抛,但要留下记录。 + + 释放故障是连接池异常的先兆信号,既不上抛又不记录等于让它彻底消失。 + """ + logged = _capture_logger_errors(monkeypatch) + db = _sync_session() + db.close.side_effect = RuntimeError("close failed") + monkeypatch.setattr(decorators_module, "ScopedSession", lambda: db) + + @decorators_module.db_update + def _write(db=None): + """必定失败的更新。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + _write(db=None) + + assert any("close failed" in str(msg) for msg in logged), \ + f"关闭会话失败被静默吞掉,未留下任何记录:{logged}" + + +def test_db_update_close_failure_does_not_break_success_path(monkeypatch): + """ + 业务成功而释放失败时,调用方必须正常拿到返回值(本次修改新引入的行为)。 + + 此时事务已经 commit、数据确实落库了,把释放故障升级成调用方的异常只会让一次成功的 + 写入看起来像失败,诱使上层重试、重复提交。故障降级为日志是刻意的取舍。 + """ + logged = _capture_logger_errors(monkeypatch) + db = _sync_session() + db.close.side_effect = RuntimeError("close failed") + monkeypatch.setattr(decorators_module, "ScopedSession", lambda: db) + + @decorators_module.db_update + def _write(db=None): + """成功的更新。""" + return "written" + + assert _write(db=None) == "written", "释放失败把成功的写入变成了异常" + db.commit.assert_called_once() + assert any("close failed" in str(msg) for msg in logged), \ + f"释放故障既没上抛也没记录,等于彻底消失:{logged}" + + +@pytest.mark.asyncio +async def test_async_db_update_scope_exit_failure_does_not_replace_original_error(monkeypatch): + """ + 异步更新:作用域退出失败时,调用方仍须收到原始业务异常。 + """ + db = _async_session() + _install_failing_scope(monkeypatch, db, RuntimeError("exit failed")) + boom = _BusinessError("boom") + + @decorators_module.async_db_update + async def _write(db=None): + """抛出一个可辨认的异常实例。""" + raise boom + + with pytest.raises(_BusinessError) as excinfo: + await _write(db=None) + + assert excinfo.value is boom, "作用域退出失败顶替了原始异常" + + +@pytest.mark.asyncio +async def test_async_db_update_logs_scope_exit_failure(monkeypatch): + """ + 异步更新:作用域退出失败要留下记录,不能静默吞掉。 + """ + logged = _capture_logger_errors(monkeypatch) + db = _async_session() + _install_failing_scope(monkeypatch, db, RuntimeError("exit failed")) + + @decorators_module.async_db_update + async def _write(db=None): + """必定失败的异步更新。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + await _write(db=None) + + assert any("exit failed" in str(msg) for msg in logged), \ + f"作用域退出失败被静默吞掉,未留下任何记录:{logged}" + + +@pytest.mark.asyncio +async def test_async_db_update_scope_exit_failure_does_not_break_success_path(monkeypatch): + """ + 异步更新:业务成功而作用域退出失败时,调用方仍须正常拿到返回值。 + """ + logged = _capture_logger_errors(monkeypatch) + db = _async_session() + _install_failing_scope(monkeypatch, db, RuntimeError("exit failed")) + + @decorators_module.async_db_update + async def _write(db=None): + """成功的异步更新。""" + return "written" + + assert await _write(db=None) == "written", "释放失败把成功的写入变成了异常" + db.commit.assert_awaited_once() + assert any("exit failed" in str(msg) for msg in logged), \ + f"释放故障既没上抛也没记录,等于彻底消失:{logged}" + + +def test_db_query_close_failure_does_not_replace_original_error(monkeypatch): + """ + 同步查询:关闭会话失败时,调用方仍须收到原始业务异常。 + """ + db = _sync_session() + db.close.side_effect = RuntimeError("close failed") + monkeypatch.setattr(decorators_module, "ScopedSession", lambda: db) + boom = _BusinessError("boom") + + @decorators_module.db_query + def _read(db=None): + """抛出一个可辨认的异常实例。""" + raise boom + + with pytest.raises(_BusinessError) as excinfo: + _read(db=None) + + assert excinfo.value is boom, "关闭会话失败顶替了原始异常" + + +def test_db_query_logs_close_failure(monkeypatch): + """ + 同步查询:关闭会话失败要留下记录,不能静默吞掉。 + """ + logged = _capture_logger_errors(monkeypatch) + db = _sync_session() + db.close.side_effect = RuntimeError("close failed") + monkeypatch.setattr(decorators_module, "ScopedSession", lambda: db) + + @decorators_module.db_query + def _read(db=None): + """必定失败的查询。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + _read(db=None) + + assert any("close failed" in str(msg) for msg in logged), \ + f"关闭会话失败被静默吞掉,未留下任何记录:{logged}" + + +def test_db_query_close_failure_does_not_break_success_path(monkeypatch): + """ + 同步查询:查询成功而释放失败时,调用方必须正常拿到查询结果。 + + 结果已经取到手了,因为归还会话时出的岔子而把它丢掉,是把一次成功的读变成失败。 + """ + logged = _capture_logger_errors(monkeypatch) + db = _sync_session() + db.close.side_effect = RuntimeError("close failed") + monkeypatch.setattr(decorators_module, "ScopedSession", lambda: db) + + @decorators_module.db_query + def _read(db=None): + """成功的查询。""" + return ["row"] + + assert _read(db=None) == ["row"], "释放失败把成功的查询变成了异常" + assert any("close failed" in str(msg) for msg in logged), \ + f"释放故障既没上抛也没记录,等于彻底消失:{logged}" + + +@pytest.mark.asyncio +async def test_async_db_query_scope_exit_failure_does_not_replace_original_error(monkeypatch): + """ + 异步查询:作用域退出失败时,调用方仍须收到原始业务异常。 + """ + db = _async_session() + _install_failing_scope(monkeypatch, db, RuntimeError("exit failed")) + boom = _BusinessError("boom") + + @decorators_module.async_db_query + async def _read(db=None): + """抛出一个可辨认的异常实例。""" + raise boom + + with pytest.raises(_BusinessError) as excinfo: + await _read(db=None) + + assert excinfo.value is boom, "作用域退出失败顶替了原始异常" + + +@pytest.mark.asyncio +async def test_async_db_query_logs_scope_exit_failure(monkeypatch): + """ + 异步查询:作用域退出失败要留下记录,不能静默吞掉。 + """ + logged = _capture_logger_errors(monkeypatch) + db = _async_session() + _install_failing_scope(monkeypatch, db, RuntimeError("exit failed")) + + @decorators_module.async_db_query + async def _read(db=None): + """必定失败的异步查询。""" + raise _BusinessError("boom") + + with pytest.raises(_BusinessError): + await _read(db=None) + + assert any("exit failed" in str(msg) for msg in logged), \ + f"作用域退出失败被静默吞掉,未留下任何记录:{logged}" + + +@pytest.mark.asyncio +async def test_async_db_query_scope_exit_failure_does_not_break_success_path(monkeypatch): + """ + 异步查询:查询成功而作用域退出失败时,调用方仍须正常拿到查询结果。 + """ + logged = _capture_logger_errors(monkeypatch) + db = _async_session() + _install_failing_scope(monkeypatch, db, RuntimeError("exit failed")) + + @decorators_module.async_db_query + async def _read(db=None): + """成功的异步查询。""" + return ["row"] + + assert await _read(db=None) == ["row"], "释放失败把成功的查询变成了异常" + assert any("exit failed" in str(msg) for msg in logged), \ + f"释放故障既没上抛也没记录,等于彻底消失:{logged}" diff --git a/tests/test_db_downloadhistory_queries.py b/tests/test_db_downloadhistory_queries.py new file mode 100644 index 000000000..5c665c53c --- /dev/null +++ b/tests/test_db_downloadhistory_queries.py @@ -0,0 +1,396 @@ +""" +下载历史与下载文件记录的查询行为。 + +下载历史是「这个种子是为哪部片下的」的唯一来源,整理链靠它把文件落到正确的媒体库 +目录;查错一条就是整理到错误的剧里。下载文件记录则决定「删种时该删哪些文件」, +条件写宽会误删别的任务的文件。 +""" +import asyncio +import time as _time + +import pytest + +from app.db.models import downloadhistory as downloadhistory_module +from app.db.models.downloadhistory import DownloadFiles, DownloadHistory +from app.schemas.types import MediaSource, MediaType + +TMDB = str(MediaSource.TMDB) + + +@pytest.fixture(autouse=True) +def _track(db): + """把下载历史与下载文件表纳入用例级回收。""" + db.watermark(DownloadHistory, DownloadFiles) + + +def _history(title: str, download_hash: str = None, media_id: str = "7001", + mtype: str = None, year: str = "2026", seasons: str = None, + episodes: str = None, date: str = "2026-08-13 10:00:00", + username: str = "alice", music_type: str = None, + path: str = None) -> DownloadHistory: + """构造一条下载历史。""" + return DownloadHistory(path=path or f"/downloads/{title}", type=mtype or MediaType.TV.value, + title=title, year=year, media_source=TMDB, media_id=media_id, + music_type=music_type, seasons=seasons, episodes=episodes, + download_hash=download_hash, date=date, username=username) + + +def _file(download_hash: str, fullpath: str, savepath: str = "/downloads", + state: int = 1) -> DownloadFiles: + """构造一条下载文件记录。""" + return DownloadFiles(downloader="qbittorrent", download_hash=download_hash, + fullpath=fullpath, savepath=savepath, + filepath=fullpath.rsplit("/", 1)[-1], + torrentname="种子", state=state) + + +# --------------------------------------------------------------------------- # +# DownloadHistory:按 hash 查询 +# --------------------------------------------------------------------------- # + +def test_get_by_hash_returns_the_latest_record(db): + """ + 同一 hash 存在多条时取最新的一条。 + + 重复下载会留下多条历史,取到旧的那条会让整理用上过期的识别结果。 + """ + db.add(_history("旧记录", download_hash="h-1", date="2026-08-01 10:00:00"), + _history("新记录", download_hash="h-1", date="2026-08-12 10:00:00")) + + assert DownloadHistory.get_by_hash(db.session, "h-1").title == "新记录" + assert DownloadHistory.get_by_hash(db.session, "h-missing") is None + + +def test_get_by_hashes_keeps_request_order_and_dedupes(db): + """ + 批量查询按请求顺序返回、去重、跳过查不到的 hash,每个 hash 只给最新一条。 + + 这个方法存在的理由就是消除 N+1;返回顺序与入参不一致会让上层错位地把 + A 任务的历史贴到 B 任务上。 + """ + db.add(_history("A 旧", download_hash="h-a", date="2026-08-01 10:00:00"), + _history("A 新", download_hash="h-a", date="2026-08-12 10:00:00"), + _history("B", download_hash="h-b", date="2026-08-05 10:00:00")) + + got = DownloadHistory.get_by_hashes(db.session, ["h-b", "h-a", "h-a", "", "h-missing"]) + + assert [h.title for h in got] == ["B", "A 新"] + assert DownloadHistory.get_by_hashes(db.session, []) == [] + assert DownloadHistory.get_by_hashes(db.session, None) == [] + + +# --------------------------------------------------------------------------- # +# DownloadHistory:身份与列表查询 +# --------------------------------------------------------------------------- # + +def test_get_by_media_identity_optionally_narrows_by_music_type(db): + """ + 按媒体身份查询时音乐实体类型可选,给出即须生效;空身份短路成空列表。 + """ + db.add(_history("单曲", media_id="mb-1", music_type="recording"), + _history("专辑", media_id="mb-1", music_type="album")) + + assert len(DownloadHistory.get_by_media_identity(db.session, MediaSource.TMDB, "mb-1")) == 2 + assert [h.title for h in DownloadHistory.get_by_media_identity( + db.session, MediaSource.TMDB, "mb-1", music_type="album")] == ["专辑"] + assert DownloadHistory.get_by_media_identity(db.session, MediaSource.TMDB, " ") == [] + assert DownloadHistory.get_by_media_identity(db.session, None, "mb-1") == [] + + +def test_list_by_page_is_newest_first_and_paged(db): + """ + 历史列表按时间倒序、同时间按主键倒序分页,并与异步孪生方法一致。 + """ + for index in range(4): + db.add(_history(f"p-{index}", date=f"2026-08-13 10:00:0{index}")) + + page1 = DownloadHistory.list_by_page(db.session, page=1, count=2) + assert [h.title for h in page1] == ["p-3", "p-2"] + assert [h.title for h in DownloadHistory.list_by_page(db.session, page=2, count=2)] == \ + ["p-1", "p-0"] + assert [h.title for h in asyncio.run( + DownloadHistory.async_list_by_page(page=1, count=2))] == ["p-3", "p-2"] + + +def test_get_by_path_finds_the_download_directory(db): + """ + 按保存路径查询用于把落地文件反查回下载任务,查不到时返回 None。 + """ + db.add(_history("有路径", path="/downloads/unique-path")) + + assert DownloadHistory.get_by_path(db.session, "/downloads/unique-path").title == "有路径" + assert DownloadHistory.get_by_path(db.session, "/downloads/nope") is None + + +@pytest.mark.parametrize("season,episode,expected", [ + ("S01", "E01", ["季集精确"]), + ("S01", None, ["季集精确", "整季"]), + (None, None, ["季集精确", "整季", "另一季"]), +]) +def test_get_last_by_media_identity_narrows_by_season_and_episode(db, season, episode, expected): + """ + 按媒体身份查询时,季与集逐级收窄。 + + 收窄失效会让「这一集下过没有」误判成整季都下过,订阅直接跳过后续剧集。 + """ + db.add(_history("季集精确", seasons="S01", episodes="E01"), + _history("整季", seasons="S01", episodes=None), + _history("另一季", seasons="S02", episodes=None)) + + got = DownloadHistory.get_last_by(db.session, mtype=MediaType.TV.value, + media_source=MediaSource.TMDB, media_id="7001", + season=season, episode=episode) + + assert sorted(h.title for h in got) == sorted(expected) + + +@pytest.mark.parametrize("season,episode,expected", [ + ("S01", "E01", ["标题季集"]), + ("S01", None, ["标题季集", "标题整季"]), + (None, None, ["标题季集", "标题整季"]), +]) +def test_get_last_by_falls_back_to_title_and_year(db, season, episode, expected): + """ + 没有媒体身份时退回「标题 + 年份」查询,同样支持季集收窄。 + + 这条回退路径服务于识别失败的历史数据,丢了会让这些记录彻底查不到。 + """ + db.add(_history("标题季集", media_id="7900", seasons="S01", episodes="E01", year="2020"), + _history("标题整季", media_id="7901", seasons="S01", episodes=None, year="2020")) + + got = DownloadHistory.get_last_by(db.session, title="标题季集", year="2020", + season=season, episode=episode) + got += DownloadHistory.get_last_by(db.session, title="标题整季", year="2020", + season=season, episode=episode) + + assert sorted(h.title for h in got) == sorted(expected) + + +def test_get_last_by_without_any_identity_returns_empty(db): + """ + 既无媒体身份也无标题年份时返回空列表,不能退化成返回全表。 + """ + db.add(_history("任意")) + + assert DownloadHistory.get_last_by(db.session) == [] + assert DownloadHistory.get_last_by(db.session, title="只有标题") == [] + + +def test_list_by_user_date_scopes_to_owner(db): + """ + 按用户与时间查询必须限定用户名;不给用户名则跨用户返回。 + """ + db.add(_history("alice 的", username="alice", date="2026-08-01 10:00:00"), + _history("bob 的", username="bob", date="2026-08-01 10:00:00"), + _history("太新", username="alice", date="2026-08-20 10:00:00")) + + mine = DownloadHistory.list_by_user_date(db.session, "2026-08-10", username="alice") + assert [h.title for h in mine] == ["alice 的"] + + everyone = DownloadHistory.list_by_user_date(db.session, "2026-08-10") + assert {h.title for h in everyone} >= {"alice 的", "bob 的"} + + +def test_list_by_user_date_excludes_the_row_exactly_at_the_boundary(db): + """ + 取的是「该时刻之前」的历史(``date < date``),正好等于该时刻的那条不算在内。 + + 上面的用例数据离查询时刻有十天之遥,比较符放宽成 ``<=`` 也照样绿; + 这里把行摆在边界上,让开闭区间之差可观测。 + """ + boundary = "2026-08-10 00:00:00" + db.add(_history("边界上", username="carol", date=boundary), + _history("边界前一秒", username="carol", date="2026-08-09 23:59:59")) + + rows = DownloadHistory.list_by_user_date(db.session, boundary, username="carol") + + assert [h.title for h in rows] == ["边界前一秒"] + + +def test_list_by_date_optionally_narrows_by_season(db): + """ + 按时间与媒体身份查询时季号可选,给出即须生效。 + """ + db.add(_history("第一季", seasons="S01", date="2026-08-12 10:00:00"), + _history("第二季", seasons="S02", date="2026-08-12 10:00:00"), + _history("太旧", seasons="S01", date="2026-01-01 10:00:00")) + + scoped = DownloadHistory.list_by_date(db.session, "2026-08-01", MediaType.TV.value, + MediaSource.TMDB, "7001", seasons="S01") + assert [h.title for h in scoped] == ["第一季"] + + both = DownloadHistory.list_by_date(db.session, "2026-08-01", MediaType.TV.value, + MediaSource.TMDB, "7001") + assert {h.title for h in both} == {"第一季", "第二季"} + + +def test_list_by_date_excludes_the_row_exactly_at_the_boundary(db): + """ + 取的是「该时刻之后」的历史(``date > date``),正好等于该时刻的那条不算在内。 + + 这个查询用于判断某媒体近期是否已下载过,边界放宽成 ``>=`` 会把上一轮刚好压线的 + 记录算成「已下过」,从而误跳过一次下载;两侧数据都离边界很远时看不出来。 + """ + boundary = "2026-08-01 00:00:00" + db.add(_history("边界上", media_id="7011", date=boundary), + _history("边界后一秒", media_id="7011", date="2026-08-01 00:00:01")) + + rows = DownloadHistory.list_by_date(db.session, boundary, MediaType.TV.value, + MediaSource.TMDB, "7011") + + assert [h.title for h in rows] == ["边界后一秒"] + + +def test_list_by_type_only_returns_recent_days(db): + """ + 按类型取最近 N 天,超出窗口的不返回——否则首页统计会把全量历史拉出来。 + """ + db.add(_history("最近", date="2099-01-01 00:00:00"), + _history("很久以前", date="2000-01-01 00:00:00")) + + names = {h.title for h in DownloadHistory.list_by_type(db.session, MediaType.TV.value, days=7)} + + assert "最近" in names and "很久以前" not in names + + +def test_list_by_type_includes_the_window_start_boundary(db, frozen_now): + """ + 时间窗是闭区间起点(``date >= 起点``),正好落在起点的那条必须在结果里。 + + 窗口起点由「调用时刻 - N 天」现算,不冻结时钟就摆不到边界上;上面那条用例用的是 + 2099/2000 两个极端值,比较符改成 ``>`` 也照样绿。 + """ + now = frozen_now(downloadhistory_module) + window_start = _time.strftime("%Y-%m-%d %H:%M:%S", _time.localtime(now - 86400 * 7)) + one_second_earlier = _time.strftime("%Y-%m-%d %H:%M:%S", + _time.localtime(now - 86400 * 7 - 1)) + db.add(_history("窗口起点上", media_id="7012", date=window_start), + _history("窗口起点前一秒", media_id="7012", date=one_second_earlier)) + + names = {h.title for h in DownloadHistory.list_by_type(db.session, MediaType.TV.value, days=7)} + + assert "窗口起点上" in names + assert "窗口起点前一秒" not in names + + +def test_delete_before_is_batched_and_keeps_recent(db): + """ + 历史清理分批执行且不碰保留期内的记录。 + """ + for index in range(4): + db.add(_history(f"old-{index}", date=f"2026-01-01 10:00:0{index}")) + db.add(_history("recent", date="2026-08-13 10:00:00")) + + assert DownloadHistory.delete_before(db.session, before_time="2026-08-01", limit=2) == 2 + assert DownloadHistory.delete_before(db.session, before_time="2026-08-01", limit=100) == 2 + assert DownloadHistory.delete_before(db.session, before_time="2026-08-01", limit=100) == 0 + + assert DownloadHistory.list_by_page(db.session, page=1, count=1)[0].title == "recent" + + +def test_delete_before_keeps_the_row_exactly_at_the_boundary(db): + """ + 保留时间点上的历史属于「保留期内」,不能被清理(``date < before_time``)。 + + 上面那条用例的数据离水位有半年,``<`` 写成 ``<=`` 也不可观测; + 这里把行压在水位上,让开闭区间之差暴露出来。 + """ + boundary = "2026-05-01 00:00:00" + at_boundary = db.add(_history("边界上", media_id="7013", date=boundary)) + db.add(_history("边界前一秒", media_id="7013", date="2026-04-30 23:59:59")) + + assert DownloadHistory.delete_before(db.session, before_time=boundary, limit=100) == 1 + + assert db.session.get(DownloadHistory, at_boundary.id) is not None + + +def test_count_and_title_search_match_async_twins(db): + """ + 异步的总数与标题检索必须与已落库的数据一致。 + + 标题检索走大小写不敏感匹配,退化成精确匹配会让搜索框形同虚设。 + """ + db.add(_history("Unique Title Here", date="2026-08-13 10:00:00")) + + assert asyncio.run(DownloadHistory.async_count()) >= 1 + assert asyncio.run(DownloadHistory.async_count_by_title(title="unique title")) == 1 + assert [h.title for h in asyncio.run(DownloadHistory.async_list_by_title( + title="UNIQUE TITLE"))] == ["Unique Title Here"] + + +# --------------------------------------------------------------------------- # +# DownloadFiles +# --------------------------------------------------------------------------- # + +def test_files_get_by_hash_optionally_filters_state(db): + """ + 按 hash 取文件时状态可选:不给返回全部,给出则只返回该状态。 + + 删种时取的是「状态正常」这一批,条件失效会把已删除的文件重复删一遍。 + """ + db.add(_file("fh-1", "/downloads/a.mkv", state=1), + _file("fh-1", "/downloads/b.mkv", state=0), + _file("fh-2", "/downloads/c.mkv", state=1)) + + assert len(DownloadFiles.get_by_hash(db.session, "fh-1")) == 2 + assert [f.fullpath for f in DownloadFiles.get_by_hash(db.session, "fh-1", state=1)] == \ + ["/downloads/a.mkv"] + + +def test_files_get_by_fullpath_returns_newest_or_all(db): + """ + 按完整路径查询默认取最新一条,要求全部时按主键倒序返回。 + + 同一路径可能被多次下载覆盖,取到旧记录会关联到错误的下载任务。 + """ + first = db.add(_file("fh-3", "/downloads/same.mkv")) + second = db.add(_file("fh-4", "/downloads/same.mkv")) + + assert DownloadFiles.get_by_fullpath(db.session, "/downloads/same.mkv").id == second.id + assert [f.id for f in DownloadFiles.get_by_fullpath( + db.session, "/downloads/same.mkv", all_files=True)] == [second.id, first.id] + assert DownloadFiles.get_by_fullpath(db.session, "/downloads/none.mkv") is None + + +def test_files_get_by_savepath_returns_every_file_of_that_directory(db): + """ + 按保存目录查询返回该目录下的全部文件记录。 + """ + db.add(_file("fh-5", "/downloads/dir/a.mkv", savepath="/downloads/dir"), + _file("fh-5", "/downloads/dir/b.mkv", savepath="/downloads/dir"), + _file("fh-6", "/downloads/other/c.mkv", savepath="/downloads/other")) + + assert len(DownloadFiles.get_by_savepath(db.session, "/downloads/dir")) == 2 + + +def test_files_delete_by_fullpath_marks_state_instead_of_removing(db): + """ + 「删除」是把状态置 0 而不是删行,并且只影响状态正常的那条。 + + 保留行是为了让后续整理仍能追溯文件来源;直接删行会让历史断链。 + """ + db.add(_file("fh-7", "/downloads/del.mkv", state=1), + _file("fh-8", "/downloads/keep.mkv", state=1)) + + DownloadFiles.delete_by_fullpath(db.session, "/downloads/del.mkv") + + assert DownloadFiles.get_by_fullpath(db.session, "/downloads/del.mkv").state == 0 + assert DownloadFiles.get_by_fullpath(db.session, "/downloads/keep.mkv").state == 1 + + +def test_files_delete_orphans_only_removes_records_without_parent(db): + """ + 孤儿清理只删掉找不到父下载历史的文件记录,且分批执行。 + + 条件写反会把仍有父记录的文件删光,删种时便再也找不到要删哪些文件。 + """ + db.add(_history("有父记录", download_hash="fh-parent")) + db.add(_file("fh-parent", "/downloads/child.mkv"), + _file("fh-orphan-1", "/downloads/o1.mkv"), + _file("fh-orphan-2", "/downloads/o2.mkv")) + + assert DownloadFiles.delete_orphans(db.session, limit=1) == 1 + assert DownloadFiles.delete_orphans(db.session, limit=100) == 1 + assert DownloadFiles.delete_orphans(db.session, limit=100) == 0 + + assert DownloadFiles.get_by_fullpath(db.session, "/downloads/child.mkv") is not None diff --git a/tests/test_db_engine_postgresql.py b/tests/test_db_engine_postgresql.py new file mode 100644 index 000000000..02c519659 --- /dev/null +++ b/tests/test_db_engine_postgresql.py @@ -0,0 +1,349 @@ +""" +PostgreSQL 引擎构建与连接额度校验测试。 + +生产故障发生在 PostgreSQL 环境(一次 60 站点搜索触发 74 次 +TooManyConnectionsError),但本地开发与 CI 都跑 SQLite——PG 分支此前零执行, +额度校验这类「只在 PG 下生效」的逻辑完全没有测试兜底。 + +这里用 mock 覆盖 PG 路径,不依赖真实 PostgreSQL 实例:额度核算是纯计算, +校验逻辑只需要伪造 SHOW 查询的返回值。 +""" +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from app.runtime.config import settings +from app.db import engine as engine_module +from app.db.engine import connection_budget + + +def _fake_pg_connection(max_connections: int, reserved: int) -> MagicMock: + """ + 伪造一个 PostgreSQL 连接,按顺序返回两条 SHOW 查询的结果。 + :param max_connections: max_connections 的返回值 + :param reserved: superuser_reserved_connections 的返回值 + """ + conn = MagicMock() + conn.execute.side_effect = [ + MagicMock(scalar=MagicMock(return_value=str(max_connections))), + MagicMock(scalar=MagicMock(return_value=str(reserved))), + ] + ctx = MagicMock() + ctx.__enter__ = MagicMock(return_value=conn) + ctx.__exit__ = MagicMock(return_value=False) + return ctx + + +def _patch_engine(monkeypatch, connect) -> None: + """ + 把额度校验取到的同步引擎换成只带 connect 的替身。 + + 额度校验只用引擎做一件事:connect() 出来跑两条 SHOW。打桩 get_engine() 而不是 + 在真引擎上改 connect——后者会为了一次本可全 mock 的校验真的把引擎建出来、连库、 + 设 WAL,正是引擎惰性化要消掉的那种 import/取值副作用。 + :param monkeypatch: pytest 的 monkeypatch 夹具 + :param connect: 替身引擎的 connect 实现 + """ + monkeypatch.setattr(engine_module, "get_engine", + lambda: SimpleNamespace(connect=connect)) + + +# --------------------------------------------------------------------------- # +# 额度核算 +# --------------------------------------------------------------------------- # + +def test_budget_sums_all_connection_sources(monkeypatch): + """ + 理论峰值必须涵盖全部连接来源。 + + 各连接池此前彼此独立配置、没有任何地方核算总和——异步侧从无界收敛到有界后, + 决定安全与否的就变成了这个总数。漏算任何一项都会让校验失去意义。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "postgresql", raising=False) + monkeypatch.setattr(settings, "DB_POSTGRESQL_POOL_SIZE", 10, raising=False) + monkeypatch.setattr(settings, "DB_POSTGRESQL_MAX_OVERFLOW", 50, raising=False) + monkeypatch.setattr(settings, "DB_POOL_TYPE", "QueuePool", raising=False) + monkeypatch.setattr(settings, "DB_ASYNC_POOL_TYPE", "QueuePool", raising=False) + monkeypatch.setattr(settings, "DB_ASYNC_POOL_SIZE", 5, raising=False) + monkeypatch.setattr(settings, "DB_ASYNC_MAX_OVERFLOW", 10, raising=False) + monkeypatch.setattr(settings, "DB_ASYNC_FALLBACK_LIMIT", 10, raising=False) + + budget = engine_module.connection_budget() + + assert budget["sync"] == 60 + assert budget["async_pooled"] == 15 + assert budget["async_fallback"] == 10 + assert budget["total"] == 85 + + +def test_budget_counts_nullpool_async_as_scheduler_sized(monkeypatch): + """ + 异步侧配成 NullPool 时不存在池上限,此时用调度器线程数作为峰值估计 + ——这正是缺陷未修复前的真实状况,额度核算必须如实反映而不是记为 0。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "postgresql", raising=False) + monkeypatch.setattr(settings, "DB_ASYNC_POOL_TYPE", "NullPool", raising=False) + + budget = engine_module.connection_budget() + + assert budget["async_pooled"] == 0, "NullPool 没有池,不应计入池上限" + assert budget["async_fallback"] == settings.CONF.scheduler + + +def test_budget_uses_sqlite_pool_for_sqlite(monkeypatch): + """ + SQLite 后端应取 SQLite 的池配置,而不是 PostgreSQL 的。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "sqlite", raising=False) + monkeypatch.setattr(settings, "DB_SQLITE_POOL_SIZE", 3, raising=False) + monkeypatch.setattr(settings, "DB_SQLITE_MAX_OVERFLOW", 4, raising=False) + monkeypatch.setattr(settings, "DB_POOL_TYPE", "QueuePool", raising=False) + + assert engine_module.connection_budget()["sync"] == 7 + + +# --------------------------------------------------------------------------- # +# 额度校验(PostgreSQL 路径) +# --------------------------------------------------------------------------- # + +def test_check_passes_when_within_available(monkeypatch): + """ + 峰值在数据库可用额度之内时应通过。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "postgresql", raising=False) + monkeypatch.setattr(engine_module, "connection_budget", + lambda: {"sync": 60, "async_pooled": 15, "async_fallback": 10, + "per_worker": 85, "workers": 1, "total": 85}) + _patch_engine(monkeypatch, lambda *_a, **_kw: _fake_pg_connection(100, 3)) + + assert engine_module.check_connection_budget() is True + + +def test_check_fails_when_exceeding_available(monkeypatch): + """ + 峰值超出可用额度时必须返回 False 并报错。 + + 这是本校验存在的全部意义:把「突发并发时才以 TooManyConnectionsError 暴露」 + 的配置问题,前移到启动期就能看见。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "postgresql", raising=False) + monkeypatch.setattr(engine_module, "connection_budget", + lambda: {"sync": 60, "async_pooled": 40, "async_fallback": 30, + "per_worker": 130, "workers": 1, "total": 130}) + _patch_engine(monkeypatch, lambda *_a, **_kw: _fake_pg_connection(100, 3)) + errors = [] + monkeypatch.setattr(engine_module.logger, "error", errors.append) + + assert engine_module.check_connection_budget() is False + assert errors, "超额时必须留下错误日志" + assert "额度不足" in errors[0] + # 报错必须指出可调的参数,否则用户不知道该改什么 + assert "MAX_OVERFLOW" in errors[0] + + +def test_check_uses_real_max_connections_not_assumption(monkeypatch): + """ + 必须读取数据库的真实 max_connections,而不是假定 100 + ——部署方很可能已经调过它,用猜测值会得出相反的结论。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "postgresql", raising=False) + monkeypatch.setattr(engine_module, "connection_budget", + lambda: {"sync": 200, "async_pooled": 15, "async_fallback": 10, + "per_worker": 225, "workers": 1, "total": 225}) + # 数据库已调大到 500,225 应当通过 + _patch_engine(monkeypatch, lambda *_a, **_kw: _fake_pg_connection(500, 3)) + + assert engine_module.check_connection_budget() is True + + +def test_check_tolerates_unreadable_limits(monkeypatch): + """ + 读取上限失败(权限不足、连接异常)不能阻断启动,只记告警。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "postgresql", raising=False) + + def boom(*_args, **_kwargs): + """ + 模拟无权执行 SHOW。 + """ + raise RuntimeError("permission denied for SHOW") + + _patch_engine(monkeypatch, boom) + warnings = [] + monkeypatch.setattr(engine_module.logger, "warn", warnings.append) + + assert engine_module.check_connection_budget() is True + assert warnings + + +def test_check_skips_query_for_sqlite(monkeypatch): + """ + SQLite 没有服务端连接上限,不应执行任何 SHOW 查询。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "sqlite", raising=False) + called = [] + _patch_engine(monkeypatch, lambda *_a, **_kw: called.append(1)) + + assert engine_module.check_connection_budget() is True + assert not called, "SQLite 不应连接数据库查询上限" + + +# --------------------------------------------------------------------------- # +# PostgreSQL 引擎构建 +# --------------------------------------------------------------------------- # + +def test_pg_sync_engine_applies_pool_settings(monkeypatch): + """ + 同步 PG 引擎应带上 QueuePool 的尺寸参数。 + """ + monkeypatch.setattr(settings, "DB_POOL_TYPE", "QueuePool", raising=False) + monkeypatch.setattr(settings, "DB_POSTGRESQL_POOL_SIZE", 7, raising=False) + monkeypatch.setattr(settings, "DB_POSTGRESQL_MAX_OVERFLOW", 9, raising=False) + captured = {} + monkeypatch.setattr(engine_module, "create_engine", + lambda **kw: captured.update(kw) or MagicMock()) + monkeypatch.setattr(engine_module, "_register_database_error_logging", lambda *_a: None) + + engine_module._get_postgresql_engine(is_async=False) + + assert captured["pool_size"] == 7 + assert captured["max_overflow"] == 9 + assert captured["url"].startswith("postgresql") + + +def test_pg_async_engine_pooled_omits_poolclass(monkeypatch): + """ + 池化的异步引擎不得指定 poolclass:SQLAlchemy 需自行选用异步适配的 + AsyncAdaptedQueuePool,传入同步 QueuePool 会出错。 + """ + monkeypatch.setattr(settings, "DB_ASYNC_POOL_SIZE", 5, raising=False) + monkeypatch.setattr(settings, "DB_ASYNC_MAX_OVERFLOW", 10, raising=False) + captured = {} + monkeypatch.setattr(engine_module, "create_async_engine", + lambda **kw: captured.update(kw) or MagicMock(sync_engine=MagicMock())) + monkeypatch.setattr(engine_module, "_register_database_error_logging", lambda *_a: None) + + engine_module._get_postgresql_engine(is_async=True, pooled=True) + + assert "poolclass" not in captured + assert captured["pool_size"] == 5 + assert "asyncpg" in captured["url"] + + +def test_pg_async_engine_unpooled_uses_nullpool(monkeypatch): + """ + 未池化的异步引擎必须用 NullPool,保持跨事件循环的安全性。 + """ + captured = {} + monkeypatch.setattr(engine_module, "create_async_engine", + lambda **kw: captured.update(kw) or MagicMock(sync_engine=MagicMock())) + monkeypatch.setattr(engine_module, "_register_database_error_logging", lambda *_a: None) + + engine_module._get_postgresql_engine(is_async=True, pooled=False) + + assert captured["poolclass"].__name__ == "NullPool" + + +def test_pg_engine_injects_connect_args(monkeypatch): + """ + 驱动级参数必须能注入——经 PgBouncer 事务模式接入时 asyncpg 需要 + statement_cache_size=0,此前无法配置,导致连纯运维手段都用不了。 + """ + monkeypatch.setattr(settings, "DB_CONNECT_ARGS", {"statement_cache_size": 0}, raising=False) + captured = {} + monkeypatch.setattr(engine_module, "create_async_engine", + lambda **kw: captured.update(kw) or MagicMock(sync_engine=MagicMock())) + monkeypatch.setattr(engine_module, "_register_database_error_logging", lambda *_a: None) + + engine_module._get_postgresql_engine(is_async=True, pooled=False) + + assert captured["connect_args"]["statement_cache_size"] == 0 + + +# --------------------------------------------------------------------------- # +# 多 worker 下的额度核算 +# --------------------------------------------------------------------------- # + +def test_budget_reports_per_worker_and_total(monkeypatch): + """ + 连接池是进程级的,多 worker 下每个进程各持一份。 + + 核算必须同时给出「单进程」与「全部 worker 合计」——只报单进程会让多 worker + 部署在启动校验里一路绿灯,实际第一个 worker 还没起完就顶穿 max_connections。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "postgresql", raising=False) + monkeypatch.setattr(settings, "DB_POSTGRESQL_POOL_SIZE", 10, raising=False) + monkeypatch.setattr(settings, "DB_POSTGRESQL_MAX_OVERFLOW", 50, raising=False) + monkeypatch.setattr(settings, "DB_POOL_TYPE", "QueuePool", raising=False) + monkeypatch.setattr(settings, "DB_ASYNC_POOL_TYPE", "QueuePool", raising=False) + monkeypatch.setattr(settings, "DB_ASYNC_POOL_SIZE", 5, raising=False) + monkeypatch.setattr(settings, "DB_ASYNC_MAX_OVERFLOW", 10, raising=False) + monkeypatch.setattr(settings, "DB_ASYNC_FALLBACK_LIMIT", 10, raising=False) + monkeypatch.setattr(settings, "API_WORKERS", 4, raising=False) + + budget = connection_budget() + + assert budget["per_worker"] == 85 + assert budget["workers"] == 4 + assert budget["total"] == 340 + + +def test_budget_single_worker_keeps_total_equal_to_per_worker(monkeypatch): + """ + 单 worker 时合计等于单进程用量,与引入 worker 概念之前的口径一致。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "postgresql", raising=False) + monkeypatch.setattr(settings, "API_WORKERS", 1, raising=False) + + budget = connection_budget() + + assert budget["total"] == budget["per_worker"] + + +@pytest.mark.parametrize("workers", [0, -3, None]) +def test_budget_treats_invalid_worker_count_as_one(monkeypatch, workers): + """ + worker 数非法时按 1 计,不能让核算退化成 0 而误判「额度充足」。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "postgresql", raising=False) + monkeypatch.setattr(settings, "API_WORKERS", workers, raising=False) + + budget = connection_budget() + + assert budget["workers"] == 1 + assert budget["total"] == budget["per_worker"] + + +def test_check_fails_when_workers_multiply_past_the_limit(monkeypatch): + """ + 单进程用量在额度内、但乘上 worker 数后超限时必须报错。 + + 这正是盲区所在:85 条对 max_connections=100 是安全的,17 个 worker 的 1445 条 + 则毫无胜算,而此前的校验对后者一路放行。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "postgresql", raising=False) + monkeypatch.setattr(engine_module, "connection_budget", + lambda: {"sync": 60, "async_pooled": 15, "async_fallback": 10, + "per_worker": 85, "workers": 4, "total": 340}) + _patch_engine(monkeypatch, lambda *_a, **_kw: _fake_pg_connection(100, 3)) + errors = [] + monkeypatch.setattr(engine_module.logger, "error", errors.append) + + assert engine_module.check_connection_budget() is False + assert errors and "额度不足" in errors[0] + # 报错必须点出 worker 数,否则用户看到 340 会以为是池配置写错了 + assert "worker" in errors[0].lower() + + +def test_check_passes_when_workers_stay_within_the_limit(monkeypatch): + """ + 乘上 worker 数后仍在额度内时通过。 + """ + monkeypatch.setattr(settings, "DB_TYPE", "postgresql", raising=False) + monkeypatch.setattr(engine_module, "connection_budget", + lambda: {"sync": 20, "async_pooled": 5, "async_fallback": 5, + "per_worker": 30, "workers": 3, "total": 90}) + _patch_engine(monkeypatch, lambda *_a, **_kw: _fake_pg_connection(100, 3)) + + assert engine_module.check_connection_budget() is True diff --git a/tests/test_db_error_diagnostics.py b/tests/test_db_error_diagnostics.py index 8614237e7..c61846ea0 100644 --- a/tests/test_db_error_diagnostics.py +++ b/tests/test_db_error_diagnostics.py @@ -4,7 +4,11 @@ import pytest from sqlalchemy import create_engine, text from sqlalchemy.exc import OperationalError -import app.db as db_module +# 诊断实现已迁至 app.db.diagnostics;app.db 只做 re-export,私有符号不在其上 +import app.db.diagnostics as db_module +# 用 getter 而不是旧名字 AsyncEngine:后者只为仓库外插件保留,模块级导入它会在 pytest +# 的收集期就把全局异步引擎建出来 +from app.db.engine import get_global_async_engine class _SqliteError(Exception): @@ -62,7 +66,7 @@ def test_database_error_listener_omits_statement_and_parameters(monkeypatch) -> """数据库错误日志不得包含 SQL、参数或驱动返回的原始消息。""" messages = [] engine = create_engine("sqlite:///:memory:") - monkeypatch.setattr("app.db.logger.error", messages.append) + monkeypatch.setattr("app.db.diagnostics.logger.error", messages.append) db_module._register_database_error_logging(engine) with pytest.raises(OperationalError): @@ -84,10 +88,10 @@ def test_database_error_listener_omits_statement_and_parameters(monkeypatch) -> def test_async_database_engine_logs_driver_error_metadata(monkeypatch) -> None: """异步 Engine 应通过底层 sync engine 记录驱动错误码。""" messages = [] - monkeypatch.setattr("app.db.logger.error", messages.append) + monkeypatch.setattr("app.db.diagnostics.logger.error", messages.append) async def query_missing_table() -> None: - async with db_module.AsyncEngine.connect() as connection: + async with get_global_async_engine().connect() as connection: await connection.execute(text("SELECT * FROM async_missing_table")) with pytest.raises(OperationalError): diff --git a/tests/test_db_lazy_engine.py b/tests/test_db_lazy_engine.py new file mode 100644 index 000000000..adc658752 --- /dev/null +++ b/tests/test_db_lazy_engine.py @@ -0,0 +1,298 @@ +""" +引擎的惰性创建。 + +此前引擎在 import 期创建:`import app.db` 就会按 settings 连库、建出 user.db、SQLite +还会去设一次 WAL——仅仅把这个包 import 进来就有副作用,且违反了「隔离 CONFIG_DIR 必须 +早于它」时不会报错,只会静默写进真实的 user.db。 + +惰性化消掉的是这个副作用(排序约束本身仍在,见 engine 模块注释),代价是引入了新的 +正确性问题:首次访问的并发。这个项目有上百个 +调度线程,双重检查一旦写错,会创建出多个引擎、各自持一份连接池,额度核算随之失真。 +这类 bug 在单线程测试里永远不会暴露,所以这里显式并发压它。 +""" +import asyncio +import subprocess +import sys +import tempfile +import threading +import time +from unittest.mock import MagicMock + +import pytest + +from app.runtime.config import global_vars, settings +from app.db import decorators as decorators_module +from app.db import engine as engine_module +from app.db import session as session_module + + +@pytest.fixture +def _reset_engines(): + """复原引擎缓存,避免用例之间相互影响。 + + 只做「存档—还原」,不负责释放:用例必须自行给 ``_get_database_engine`` 打桩, + 绝不能在这里落下真引擎——还原会把它从槽里丢掉,那条连接便再无人 dispose。 + """ + saved_sync, saved_async = engine_module._sync_engine, engine_module._async_engine + yield + engine_module._sync_engine, engine_module._async_engine = saved_sync, saved_async + + +@pytest.fixture +def _reset_pooled_engines(): + """复原按事件循环缓存的池化引擎。""" + saved = dict(session_module._pooled_async_engines) + yield + session_module._pooled_async_engines.clear() + session_module._pooled_async_engines.update(saved) + + +def test_engine_is_not_created_on_import(): + """ + 仅 import 不得创建引擎——这是「独立可测」的全部意义所在。 + + 必须用子进程验证:当前进程早被其它用例触发过引擎创建了。 + """ + code = ( + "import app.db, app.db.engine as e; " + "print('CREATED' if e._sync_engine is not None or e._async_engine is not None " + "else 'LAZY')" + ) + with tempfile.TemporaryDirectory() as tmp: + out = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True, + env={"PATH": "/usr/bin:/bin", "CONFIG_DIR": tmp, + "PYTHONPATH": "."}, timeout=180) + assert "LAZY" in out.stdout, f"import 期即创建了引擎:{out.stdout}{out.stderr[-800:]}" + + +def test_importing_legacy_factory_names_does_not_create_engine(): + """ + `from app.db import SessionFactory` 这类旧写法也不得连带创建引擎。 + + 这三个名字是转发函数而非 sessionmaker 实例,正是为了让「导入」和「创建」分开: + 若改回靠模块级 __getattr__ 解析,每个导入方都会在 import 期把引擎建出来。 + """ + code = ( + "from app.db import SessionFactory, AsyncSessionFactory, ScopedSession; " + "import app.db.engine as e; " + "print('CREATED' if e._sync_engine is not None or e._async_engine is not None " + "else 'LAZY')" + ) + with tempfile.TemporaryDirectory() as tmp: + out = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True, + env={"PATH": "/usr/bin:/bin", "CONFIG_DIR": tmp, + "PYTHONPATH": "."}, timeout=180) + assert "LAZY" in out.stdout, f"导入会话工厂即创建了引擎:{out.stdout}{out.stderr[-800:]}" + + +def test_bootstrap_atexit_cleanup_does_not_create_an_engine(): + """ + 测试引导的退出清理不得为了 dispose 而把引擎创建出来。 + + `isolate_config_dir` 注册的 atexit 回调原本写作 `app.db.Engine.dispose()`——`Engine` + 是惰性解析的属性,取它本身就会**创建**引擎。于是一个只 import 过 `app.db` 的进程会在 + 解释器关停时凭空连一次库、SQLite 还要再设一遍 journal mode,全部只为随后把它 dispose。 + + 必须用子进程,且断言的是「连库的副作用没有发生」而不是引擎槽位:回调在解释器关停期 + 才执行,那时已经没有任何代码能跑断言了,能留下的证据只有 stdout。 + 也必须让子进程自己调 isolate_config_dir()——它只在真的新建了临时目录时才注册回调, + 预先把 CONFIG_DIR 塞进环境会让它直接返回、根本不注册 atexit,用例便成了空跑。 + """ + code = ( + "from app.testing.bootstrap import isolate_config_dir; " + "isolate_config_dir(); " + "import app.db; " + "print('IMPORTED')" + ) + out = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True, + env={"PATH": "/usr/bin:/bin", "PYTHONPATH": "."}, timeout=180) + assert "IMPORTED" in out.stdout, f"子进程没跑到底:{out.stdout}{out.stderr[-800:]}" + assert "journal mode set to" not in out.stdout, ( + f"退出清理凭空建了个引擎、连了一次库:{out.stdout}{out.stderr[-800:]}" + ) + + +def test_db_query_decorator_resolves_session_at_call_time(monkeypatch): + """ + 装饰器必须能在调用期取到会话。 + + 守的是一个真实踩过的坑:曾试图用模块级 __getattr__(PEP 562)把 ScopedSession + 延迟到运行期解析,但 __getattr__ 只对「对模块对象取属性」生效,装饰器函数体里的 + 裸名字 ScopedSession 是**全局名字查找**,只查模块 __dict__ 和 builtins,永远走不 + 到 __getattr__ —— 结果是运行期 NameError。它只在真正调用到某个 Oper 时才炸, + import 与单测都照常绿灯,因此必须显式钉住。 + """ + fake_session = MagicMock() + # 替身打在 session 模块的工厂上,而不是 decorators.ScopedSession:后者会把 + # 名字直接塞进 decorators 的 __dict__,反而掩盖「这个名字本来就该在」的缺陷。 + monkeypatch.setattr(session_module, "get_scoped_session", lambda: (lambda: fake_session)) + + @decorators_module.db_query + def _fetch(db=None): + """装饰器未拿到会话时应自行创建一个并塞回 db 位置。""" + return db + + # 按各 Oper 的常态调用:db 显式传 None,由装饰器补上会话 + assert _fetch(db=None) is fake_session + fake_session.close.assert_called_once() + + +def test_concurrent_first_access_creates_exactly_one_engine(_reset_engines, monkeypatch): + """ + 多线程同时首次取引擎,只能创建出一个实例。 + + 创建多个意味着每个都带一份连接池:实际连接数是额度核算的数倍,而校验对此 + 一无所知——正是这次修复想避免的那类问题。 + """ + engine_module._sync_engine = None + created = [] + barrier = threading.Barrier(16) + + def slow_factory(**_kwargs): + """放大创建耗时,把竞态窗口撑开到必定命中。""" + time.sleep(0.02) + marker = object() + created.append(marker) + return marker + + monkeypatch.setattr(engine_module, "_get_database_engine", slow_factory) + got = [] + + def worker(): + """所有线程在同一时刻发起首次访问。""" + barrier.wait() + got.append(engine_module.get_engine()) + + threads = [threading.Thread(target=worker) for _ in range(16)] + for t in threads: + t.start() + for t in threads: + t.join() + + assert len(created) == 1, f"并发首次访问创建了 {len(created)} 个引擎" + assert len({id(g) for g in got}) == 1, "不同线程拿到了不同的引擎实例" + + +def test_repeated_access_reuses_the_same_engine(_reset_engines, monkeypatch): + """ + 后续访问必须复用,而不是每次重建。 + """ + engine_module._sync_engine = None + calls = [] + monkeypatch.setattr(engine_module, "_get_database_engine", + lambda **kw: calls.append(kw) or object()) + + first = engine_module.get_engine() + second = engine_module.get_engine() + + assert first is second + assert len(calls) == 1 + + +def test_async_engine_has_its_own_lazy_slot(_reset_engines, monkeypatch): + """ + 同步与异步引擎各自独立惰性化,取其中一个不应连带创建另一个。 + """ + engine_module._sync_engine = engine_module._async_engine = None + monkeypatch.setattr(engine_module, "_get_database_engine", lambda **kw: object()) + + engine_module.get_engine() + + assert engine_module._sync_engine is not None + assert engine_module._async_engine is None, "取同步引擎连带创建了异步引擎" + + +def test_pooled_path_does_not_create_the_global_async_engine( + _reset_engines, _reset_pooled_engines, monkeypatch): + """ + 常驻循环走池化引擎时,全局异步引擎必须原封不动地留在「未创建」状态。 + + 钉的是一处已经踩过的实现:async_session_scope 曾用 + `engine is not get_global_async_engine()` 反推是否池化——这个比较**本身**就把被比较的 + 引擎创建了出来。它倒不至于多占连接——全局异步引擎用的是 NullPool,持有 0 条连接; + 真正的代价是第一个异步请求会在事件循环内部去抢引擎创建锁,把本该无锁的热路径变成 + 有锁的,而这个引擎在常驻循环下从头到尾无人使用。这类问题不会让任何断言变红,只能显式压。 + """ + monkeypatch.setattr(settings, "DB_ASYNC_POOL_TYPE", "QueuePool", raising=False) + engine_module._async_engine = None + + def _boom(): + """任何对全局异步引擎的获取都是失败信号。""" + raise AssertionError("池化路径获取了全局异步引擎") + + # 打在 session 模块的名字上:_resolve_async_engine 用的是它自己 __dict__ 里的这个名字 + monkeypatch.setattr(session_module, "get_global_async_engine", _boom) + + async def run(): + """在「当前循环即常驻循环」的前提下真的走一遍会话作用域。""" + global_vars.CURRENT_EVENT_LOOP = asyncio.get_running_loop() + session_module._pooled_async_engines.clear() + try: + async with session_module.async_session_scope(): + pass + finally: + # 用真引擎(不给工厂打桩)才能验证会话确实能建起来;建了就得自己释放 + for pooled in session_module._pooled_async_engines.values(): + await pooled.dispose() + session_module._pooled_async_engines.clear() + + saved_loop = global_vars.CURRENT_EVENT_LOOP + try: + asyncio.run(run()) + finally: + global_vars.CURRENT_EVENT_LOOP = saved_loop + + assert engine_module._async_engine is None, "池化路径把全局异步引擎创建了出来" + + +def test_engine_module_exposes_no_legacy_names(_reset_engines, monkeypatch): + """ + app.db.engine 不再解析 Engine / AsyncEngine 两个旧名字。 + + 这两个名字的对外契约是 `app.db.Engine`(由 app/db/__init__.py 的 __getattr__ 提供)。 + app.db.engine 这个模块是拆分时才出现的,仓库外不可能有代码依赖它,实现模块上那份 + 同名转发因此是纯冗余——两处独立实现同一个契约,改一处漏一处就会各取到一个引擎。 + 仓库内一律用 get_engine() / get_global_async_engine()。 + + 工厂打桩而不是让它建真引擎:_reset_engines 还原槽位时会把真引擎丢掉, + 那条连接就再没人 dispose 了。 + """ + engine_module._sync_engine = engine_module._async_engine = None + monkeypatch.setattr(engine_module, "_get_database_engine", lambda **kw: object()) + + with pytest.raises(AttributeError): + _ = engine_module.Engine + with pytest.raises(AttributeError): + _ = engine_module.AsyncEngine + # 取过之后引擎槽位仍是空的:属性访问没有绕开 getter 把引擎建出来 + assert engine_module._sync_engine is None + assert engine_module._async_engine is None + + +def test_package_entry_resolves_legacy_names(_reset_engines, monkeypatch): + """ + app.db.Engine / app.db.AsyncEngine 仍解析到与 getter 同一个引擎。 + + 这才是真正的对外契约:仓库外的插件按这两个名字取引擎,建表、Alembic 迁移、连接 + 诊断这类用途确实需要引擎对象本身,装饰器覆盖不到。上一条用例删掉了实现模块上的 + 冗余转发,这条钉住包入口那份**不能**跟着删。 + + 与惰性不冲突:属性访问发生在运行期,而不是 import 期。 + """ + import app.db as db_package + + engine_module._sync_engine = engine_module._async_engine = None + monkeypatch.setattr(engine_module, "_get_database_engine", lambda **kw: object()) + + assert db_package.Engine is engine_module.get_engine() + assert db_package.AsyncEngine is engine_module.get_global_async_engine() + + +def test_package_entry_unknown_attribute_still_raises(): + """ + 包入口的模块级 __getattr__ 不能吞掉拼写错误。 + """ + import app.db as db_package + + with pytest.raises(AttributeError): + _ = db_package.NoSuchThing diff --git a/tests/test_db_media_identity_normalizer.py b/tests/test_db_media_identity_normalizer.py new file mode 100644 index 000000000..83bdef1e3 --- /dev/null +++ b/tests/test_db_media_identity_normalizer.py @@ -0,0 +1,164 @@ +""" +媒体身份的持久化不变量:app/db/models/_identity.py 的 mapper 事件。 + +这条不变量此前靠六个 Oper 各自在建模前调一次归一来保证,是调用点纪律;现在下沉成 +flush 前的事件。因此这里断言的是「绕过 Oper 直接建模写库」时不变量仍然成立——那正是 +下沉之前会漏掉的路径,也是这次改动唯一真正新增的保证。 + +半对身份在持久化侧是「清空 + 告警」而非抛错:这几张表都是记账性写入,因一个次要字段 +就丢掉整条整理历史,代价比留下一条无身份记录更大。告警是这次一并补上的——此前它是 +完全沉默的,脏数据只能靠翻库发现。DTO 侧仍然抛错,两边语义的差异是有意的。 +""" +import pytest + +from app.db.models.transferhistory import TransferHistory +from app.db.models.transferpending import TransferPending +from app.schemas.types import MediaSource + + +@pytest.fixture(autouse=True) +def _track(db): + """把两张表纳入用例级回收。""" + db.watermark(TransferHistory, TransferPending) + + +def _history(**identity) -> TransferHistory: + """构造一条最小可写的整理历史,只有身份字段按用例变化。""" + return TransferHistory(src="/downloads/x.mkv", src_storage="local", + dest="/media/x.mkv", dest_storage="local", + mode="link", title="身份用例", status=1, files=[], + **identity) + + +def _write(db, row): + """落库并重新读回,确保拿到的是写入后的值而非内存里的原值。""" + db.add(row) + db.session.expire_all() + return TransferHistory.get(db.session, row.id) + + +# --------------------------------------------------------------------------- # +# 归一 +# --------------------------------------------------------------------------- # + +def test_alias_source_is_persisted_as_canonical_value(db): + """ + 别名来源必须落成规范值。 + + 同一个来源以 tmdb / themoviedb 两种拼法写进去,按身份查重就会把同一部剧 + 当成两条,洗版与去重全部失效。 + """ + row = _write(db, _history(media_source="tmdb", media_id="550")) + + assert row.media_source == MediaSource.TMDB.value + + +def test_media_id_is_stripped(db): + """ + ID 两端空白必须去掉——带空格的 ID 与不带的是两个不同的字符串,查重对不上。 + """ + row = _write(db, _history(media_source=MediaSource.TMDB, media_id=" 550 ")) + + assert row.media_id == "550" + + +def test_enum_source_is_persisted_as_value_not_repr(db): + """ + 传枚举时落库的是它的值,不是枚举的字符串表示。 + """ + row = _write(db, _history(media_source=MediaSource.TMDB, media_id="550")) + + assert row.media_source == MediaSource.TMDB.value + + +# --------------------------------------------------------------------------- # +# 半对身份 +# --------------------------------------------------------------------------- # + +@pytest.mark.parametrize( + "identity", + [ + pytest.param({"media_source": MediaSource.TMDB, "media_id": None}, id="缺ID"), + pytest.param({"media_source": None, "media_id": "550"}, id="缺来源"), + pytest.param({"media_source": MediaSource.TMDB, "media_id": "0"}, id="零值ID"), + pytest.param({"media_source": MediaSource.TMDB, "media_id": " "}, id="空白ID"), + pytest.param({"media_source": "不存在的源", "media_id": "550"}, id="非法来源"), + ], +) +def test_incomplete_identity_is_cleared_on_both_columns(db, identity): + """ + 身份不成对时两列一起清空,不能只留下半边。 + + 只留半边的行既匹配不上任何查重条件,也无法被后续的身份修复流程识别, + 等于永久躺在表里的脏数据。 + """ + row = _write(db, _history(**identity)) + + assert row.media_source is None + assert row.media_id is None + + +def test_incomplete_identity_is_not_silent(db, monkeypatch): + """ + 清空必须留下告警,且告警里带得上被丢弃的原值。 + + 这是本次下沉一并修掉的东西:此前归一是完全沉默的,写进一条无身份记录后没有任何 + 痕迹,只能靠翻库才发现。告警让它变成日志里可检索的事件——因此消息里必须有原值, + 否则只知道"某处丢了身份",仍然定位不到是谁写的。 + + 直接换掉模块里的 logger 而不用 caplog:项目的 LoggerManager 把 propagate 关了 + (app/runtime/log.py:333),caplog 挂在 root 上根本收不到。 + """ + warnings: list[str] = [] + monkeypatch.setattr("app.db.models._identity.logger", + type("_Spy", (), {"warn": staticmethod(warnings.append)})()) + + _write(db, _history(media_source=MediaSource.TMDB, media_id=None)) + + assert len(warnings) == 1 + assert "媒体身份不成对" in warnings[0] + # 原值要出现在消息里,否则日志定位不到是哪条写入 + assert "themoviedb" in warnings[0] and "None" in warnings[0] + + +def test_complete_identity_does_not_warn(db, monkeypatch): + """ + 身份完整时不得告警——告警一旦对正常写入也响,就会被当成噪声忽略掉。 + """ + warnings: list[str] = [] + monkeypatch.setattr("app.db.models._identity.logger", + type("_Spy", (), {"warn": staticmethod(warnings.append)})()) + + _write(db, _history(media_source="tmdb", media_id=" 550 ")) + + assert warnings == [] + + +# --------------------------------------------------------------------------- # +# 覆盖范围 +# --------------------------------------------------------------------------- # + +def test_normalization_also_applies_on_update(db): + """ + 更新路径同样归一——只管 insert 会让一条合法记录被后续更新改成半对身份。 + """ + row = _write(db, _history(media_source=MediaSource.TMDB, media_id="550")) + + row.update(db.session, {"media_source": "douban", "media_id": " 1291546 "}) + db.session.expire_all() + updated = TransferHistory.get(db.session, row.id) + + assert updated.media_source == MediaSource.Douban.value + assert updated.media_id == "1291546" + + +def test_tables_without_identity_columns_are_untouched(db): + """ + 不带身份列的表不受影响——事件挂在 Mapper 上覆盖全部映射,必须靠列名检查收窄, + 否则会去动一张根本没有这两列的表。 + """ + TransferPending.register(db.session, storage="local", src_path="/mnt/a.mkv", + now_time="2026-08-14 10:00:00") + + rows = [r for r in TransferPending.list_all(db.session) if r.src_path == "/mnt/a.mkv"] + assert len(rows) == 1 diff --git a/tests/test_db_mediaserver_queries.py b/tests/test_db_mediaserver_queries.py new file mode 100644 index 000000000..bb2d84c00 --- /dev/null +++ b/tests/test_db_mediaserver_queries.py @@ -0,0 +1,173 @@ +""" +媒体服务器条目表的查询行为。 + +这张表是「媒体库里已经有了吗」的唯一依据:查不到会重复下载,误命中会漏订阅。 +同步清理(delete_stale / delete_excluded_servers)的条件一旦写反,就会把当次刚同步 +进来的条目全删掉,表现为媒体库突然清空。 +""" +import asyncio + +import pytest + +from app.db.models.mediaserver import MediaServerItem +from app.schemas.types import MediaSource + + +@pytest.fixture(autouse=True) +def _track(db): + """把媒体服务器条目表纳入用例级回收。""" + db.watermark(MediaServerItem) + + +def _item(server: str, item_id: str, title: str = "片名", item_type: str = "电影", + year: str = "2026", media_id: str = "1001", + lst_mod_date: str = "2026-08-13 10:00:00") -> MediaServerItem: + """构造一条媒体服务器条目。""" + return MediaServerItem(server=server, library="lib", item_id=item_id, + item_type=item_type, title=title, year=year, + media_source=str(MediaSource.TMDB), media_id=media_id, + lst_mod_date=lst_mod_date) + + +def test_get_by_itemid_matches_async_twin(db): + """ + 按条目 ID 查找的同步、异步结果必须指向同一行。 + """ + db.add(_item("emby", "it-1"), _item("plex", "it-2")) + + assert MediaServerItem.get_by_itemid(db.session, "it-1").server == "emby" + assert asyncio.run(MediaServerItem.async_get_by_itemid(item_id="it-1")).server == "emby" + assert MediaServerItem.get_by_itemid(db.session, "it-missing") is None + + +def test_get_by_server_itemid_scopes_by_server(db): + """ + 条目 ID 只在单个服务器内唯一,查找必须同时限定服务器。 + + 不限定会在多媒体服务器场景下把 Emby 的条目当成 Plex 的,路径与库信息全错。 + """ + db.add(_item("emby", "same-id", title="Emby 的片"), + _item("plex", "same-id", title="Plex 的片")) + + assert MediaServerItem.get_by_server_itemid(db.session, "emby", "same-id").title == "Emby 的片" + assert MediaServerItem.get_by_server_itemid(db.session, "plex", "same-id").title == "Plex 的片" + assert MediaServerItem.get_by_server_itemid(db.session, "jellyfin", "same-id") is None + + +def test_exist_by_media_identity_requires_source_id_and_type(db): + """ + 按媒体身份判存在必须三项齐同:来源、原生 ID、条目类型。 + + 忽略类型会让同一 ID 的电影与剧集互相命中,订阅据此判定「已入库」而跳过。 + """ + db.add(_item("emby", "mi-1", media_id="555", item_type="电影")) + + assert MediaServerItem.exist_by_media_identity( + db.session, MediaSource.TMDB, "555", "电影") is not None + assert MediaServerItem.exist_by_media_identity( + db.session, MediaSource.TMDB, "555", "电视剧") is None + assert MediaServerItem.exist_by_media_identity( + db.session, MediaSource.TMDB, "556", "电影") is None + + assert asyncio.run(MediaServerItem.async_exist_by_media_identity( + media_source=MediaSource.TMDB, media_id="555", mtype="电影")) is not None + + +@pytest.mark.parametrize("mtype,year,expected", [ + (None, None, "标题匹配"), + ("电影", None, "标题匹配"), + (None, "2026", "标题匹配"), + ("电影", "2026", "标题匹配"), + ("电视剧", "2026", None), + ("电影", "2020", None), +]) +def test_exists_by_title_narrows_by_type_and_year(db, mtype, year, expected): + """ + 按标题判存在时,类型与年份各自可选,给出即须生效。 + + 这四条分支是同名不同年、同名不同类型的唯一区分手段,退化成只按标题匹配会把 + 《XX 2020》当成《XX 2026》,订阅直接被跳过。 + """ + db.add(_item("emby", "t-1", title="标题匹配", item_type="电影", year="2026")) + + found = MediaServerItem.exists_by_title(db.session, "标题匹配", mtype, year) + + assert (found.title if found else None) == expected + + +def test_exists_by_title_matches_async_twin(db): + """ + 四种参数组合下同步与异步必须给出相同的命中结果。 + """ + db.add(_item("emby", "t-par", title="并行标题", item_type="电影", year="2026")) + + for mtype, year in ((None, None), ("电影", None), (None, "2026"), ("电影", "2026")): + sync_found = MediaServerItem.exists_by_title(db.session, "并行标题", mtype, year) + async_found = asyncio.run(MediaServerItem.async_exists_by_title( + title="并行标题", mtype=mtype, year=year)) + assert (sync_found is None) == (async_found is None) + + +def test_empty_clears_only_the_given_server(db): + """ + 指定服务器时只清空该服务器的条目,不给则清空全表。 + + 误清其他服务器的条目会让那台服务器的媒体库在本次同步前一直显示为空。 + """ + db.add(_item("emby", "e-1"), _item("plex", "p-1")) + + MediaServerItem.empty(db.session, server="emby") + + assert MediaServerItem.get_by_itemid(db.session, "e-1") is None + assert MediaServerItem.get_by_itemid(db.session, "p-1") is not None + + MediaServerItem.empty(db.session) + assert MediaServerItem.get_by_itemid(db.session, "p-1") is None + + +def test_delete_stale_keeps_items_from_the_current_sync(db): + """ + 清理陈旧条目时必须保留本次同步时间戳的条目。 + + 条件写反会把刚同步进来的条目全删掉,媒体库表现为同步完反而空了。 + """ + db.add(_item("emby", "fresh", lst_mod_date="2026-08-13 12:00:00"), + _item("emby", "stale", lst_mod_date="2026-08-01 12:00:00"), + _item("emby", "never", lst_mod_date=None), + _item("plex", "other", lst_mod_date="2026-08-01 12:00:00")) + + deleted = MediaServerItem.delete_stale(db.session, server="emby", + sync_time="2026-08-13 12:00:00") + + assert deleted == 2 + assert MediaServerItem.get_by_itemid(db.session, "fresh") is not None + assert MediaServerItem.get_by_itemid(db.session, "stale") is None + assert MediaServerItem.get_by_itemid(db.session, "never") is None + assert MediaServerItem.get_by_itemid(db.session, "other") is not None + + +def test_delete_excluded_servers_keeps_configured_ones(db): + """ + 只保留仍在配置中的服务器条目,未配置与来源为空的条目应被清掉。 + """ + db.add(_item("emby", "keep-1"), _item("plex", "drop-1"), + _item(None, "drop-null")) + + deleted = MediaServerItem.delete_excluded_servers(db.session, ["emby"]) + + assert deleted == 2 + assert MediaServerItem.get_by_itemid(db.session, "keep-1") is not None + assert MediaServerItem.get_by_itemid(db.session, "drop-1") is None + assert MediaServerItem.get_by_itemid(db.session, "drop-null") is None + + +def test_delete_excluded_servers_with_empty_list_clears_everything(db): + """ + 一个服务器都没配置时清空全表——否则会留下再也不会被同步到的孤儿条目。 + """ + db.add(_item("emby", "orphan-1"), _item("plex", "orphan-2")) + + MediaServerItem.delete_excluded_servers(db.session, []) + + assert MediaServerItem.get_by_itemid(db.session, "orphan-1") is None + assert MediaServerItem.get_by_itemid(db.session, "orphan-2") is None diff --git a/tests/test_db_oper_layer.py b/tests/test_db_oper_layer.py new file mode 100644 index 000000000..09596c738 --- /dev/null +++ b/tests/test_db_oper_layer.py @@ -0,0 +1,644 @@ +""" +各业务 Oper 的数据访问行为。 + +Oper 层大多是模型方法的薄封装,但薄封装恰恰是最容易出错的地方:参数改名、 +默认值漏传、聚合逻辑写在这一层——这些都绕过了模型侧的测试。这里对着真实数据库 +验证 Oper 的对外契约,而不是验证它调了哪个模型方法。 +""" +import asyncio + +import pytest + +from app.db.oper.downloadhistory import DownloadHistoryOper +from app.db.oper.mediaserver import MediaServerOper +from app.db.models.downloadhistory import DownloadFiles, DownloadHistory +from app.db.models.mediaserver import MediaServerItem +from app.db.models.plugindata import PluginData +from app.db.models.site import Site +from app.db.models.siteicon import SiteIcon +from app.db.models.sitestatistic import SiteStatistic +from app.db.models.siteuserdata import SiteUserData +from app.db.models.user import User +from app.db.models.userconfig import UserConfig +from app.db.models.workflow import Workflow +from app.db.oper.plugindata import PluginDataOper +from app.db.oper.site import SiteOper +from app.db.oper.user import UserOper +from app.db.oper.userconfig import UserConfigOper +from app.db.oper.workflow import WorkflowOper +from app.schemas.types import MediaSource, MediaType + +TMDB = str(MediaSource.TMDB) + + +@pytest.fixture(autouse=True) +def _track(db): + """把本文件涉及的表纳入用例级回收。""" + db.watermark(Site, SiteIcon, SiteStatistic, SiteUserData, PluginData, Workflow, + User, UserConfig, MediaServerItem, DownloadHistory, DownloadFiles) + + +# --------------------------------------------------------------------------- # +# SiteOper +# --------------------------------------------------------------------------- # + +def _site_kwargs(name: str, domain: str, **extra) -> dict: + """构造新增站点的参数。""" + return dict(name=name, domain=domain, url=f"https://{domain}/", **extra) + + +def test_site_oper_add_rejects_duplicate_domain(db): + """ + 同域名不得重复新增,并如实返回原因。 + + 重复新增会让同一站点出现两条配置,Cookie 更新只命中其中一条。 + """ + oper = SiteOper(db=db.session) + + assert oper.add(**_site_kwargs("站点", "op-a.test")) == (True, "新增站点成功") + assert oper.add(**_site_kwargs("站点重复", "op-a.test")) == (False, "站点已存在") + assert oper.exists("op-a.test") is True + assert oper.exists("op-missing.test") is False + + +def test_site_oper_crud_round_trip(db): + """ + 新增、按 ID 取、更新、按域名取、列举、删除构成完整闭环。 + """ + oper = SiteOper(db=db.session) + oper.add(**_site_kwargs("站点", "op-crud.test", pri=5)) + site = oper.get_by_domain("op-crud.test") + + assert oper.get(site.id).id == site.id + assert oper.update(site.id, {"pri": 9}).pri == 9 + assert {s.domain for s in oper.list()} >= {"op-crud.test"} + assert oper.get_domains_by_ids([site.id]) == ["op-crud.test"] + assert {s.domain for s in oper.list_order_by_pri()} >= {"op-crud.test"} + + oper.delete(site.id) + assert oper.get_by_domain("op-crud.test") is None + + +def test_site_oper_async_accessors_match_sync(db): + """ + 异步访问器必须与同步给出一致的结果。 + + Oper 持有的是同步会话,异步模型方法会自行取一个异步会话——传参位置写错时 + 这里会直接失败。 + """ + oper = SiteOper(db=db.session) + oper.add(**_site_kwargs("异步站点", "op-async.test")) + site = oper.get_by_domain("op-async.test") + + assert asyncio.run(oper.async_get(site.id)).id == site.id + assert asyncio.run(oper.async_get_by_domain("op-async.test")).id == site.id + assert asyncio.run(oper.async_get_by_name("异步站点")).id == site.id + assert {s.id for s in asyncio.run(oper.async_list())} >= {site.id} + assert {s.id for s in asyncio.run(oper.async_list_active())} >= {site.id} + assert asyncio.run(oper.async_update(site.id, {"pri": 3})).pri == 3 + + +def test_site_oper_list_active_excludes_disabled(db): + """ + 启用列表排除停用站点。 + """ + oper = SiteOper(db=db.session) + oper.add(**_site_kwargs("启用", "op-on.test", is_active=True)) + oper.add(**_site_kwargs("停用", "op-off.test", is_active=False)) + + domains = {s.domain for s in oper.list_active()} + + assert "op-on.test" in domains and "op-off.test" not in domains + + +def test_site_oper_cookie_and_rss_updates_report_missing_site(db): + """ + 对不存在的站点更新 Cookie / RSS 必须返回失败,而不是静默成功。 + + 静默成功会让 CookieCloud 同步以为已生效,实际站点仍然登录失效。 + """ + oper = SiteOper(db=db.session) + oper.add(**_site_kwargs("站点", "op-cookie.test")) + + assert oper.update_cookie("op-cookie.test", "k=v") == (True, "更新站点Cookie成功") + assert oper.update_rss("op-cookie.test", "https://rss") == (True, "更新站点RSS地址成功") + assert oper.get_by_domain("op-cookie.test").cookie == "k=v" + + assert oper.update_cookie("op-none.test", "k=v")[0] is False + assert oper.update_rss("op-none.test", "https://rss")[0] is False + + +def test_site_oper_update_userdata_upserts_per_day(db): + """ + 站点用户数据按「站点 + 当天」落一条,同日重复上报走更新。 + + 每次插入会让当天出现多条快照,站点数据页面的日环比随之失真。 + """ + oper = SiteOper(db=db.session) + + oper.update_userdata("op-ud.test", "站点", {"upload": 100}) + oper.update_userdata("op-ud.test", "站点", {"upload": 200}) + + rows = oper.get_userdata_by_domain("op-ud.test") + assert len(rows) == 1 and rows[0].upload == 200 + + +def test_site_oper_update_userdata_keeps_last_good_snapshot_on_error(db): + """ + 上报带错误信息时不得覆盖当天已有的成功数据。 + + 抓取失败时用空数据覆盖,页面会显示成「上传量归零」。 + """ + oper = SiteOper(db=db.session) + oper.update_userdata("op-err.test", "站点", {"upload": 100}) + + oper.update_userdata("op-err.test", "站点", {"upload": 0, "err_msg": "登录失败"}) + + assert oper.get_userdata_by_domain("op-err.test")[0].upload == 100 + + +def test_site_oper_userdata_readers(db): + """ + 用户数据的四个读取入口都应命中同一条快照。 + """ + oper = SiteOper(db=db.session) + oper.update_userdata("op-read.test", "站点", {"upload": 50}) + today = oper.get_userdata_by_domain("op-read.test")[0].updated_day + + assert any(r.domain == "op-read.test" for r in oper.get_userdata()) + assert any(r.domain == "op-read.test" for r in oper.get_userdata_by_date(today)) + assert any(r.domain == "op-read.test" for r in oper.get_userdata_latest()) + assert [r.domain for r in asyncio.run( + oper.async_get_userdata_by_domain("op-read.test"))] == ["op-read.test"] + + +def test_site_oper_update_icon_creates_then_only_overwrites_with_content(db): + """ + 图标首次写入后,只有拿到新的 base64 才覆盖。 + + 抓取失败返回空 base64 时覆盖,会把已有图标清成空白。 + """ + oper = SiteOper(db=db.session) + + oper.update_icon("站点", "op-icon.test", "https://op-icon.test/a.ico", "AAA") + first = oper.get_icon_by_domain("op-icon.test").base64 + + oper.update_icon("站点", "op-icon.test", "https://op-icon.test/b.ico", "") + + assert oper.get_icon_by_domain("op-icon.test").base64 == first + assert first.startswith("data:image/ico;base64,") + + +def test_site_oper_success_accumulates_and_records_state(db): + """ + 访问成功累加计数并把最后状态标记为成功。 + """ + oper = SiteOper(db=db.session) + for seconds in range(1, 5): + oper.success("op-stat.test", seconds=seconds) + + stat = SiteStatistic.get_by_domain(db.session, "op-stat.test") + + assert stat.success == 4 + assert stat.lst_state == 0 + assert stat.seconds + + +def test_site_oper_success_caps_the_timing_note_at_ten_entries(db): + """ + 耗时记录最多保留最近 10 条,超出时丢弃最旧的。 + + 不设上限时这个 JSON 字段会随每次访问无限增长,最终把整行撑大到影响查询。 + 直接预置 10 条历史时间戳再上报一次——同一秒内连续调用会写进同一个键, + 只靠循环调用无法触及上限分支。 + """ + old_note = {f"2026-08-13 10:00:{index:02d}": index + 1 for index in range(10)} + db.add(SiteStatistic(domain="op-cap.test", success=10, fail=0, seconds=5, + lst_state=0, note=old_note)) + + SiteOper(db=db.session).success("op-cap.test", seconds=99) + + note = SiteStatistic.get_by_domain(db.session, "op-cap.test").note + assert len(note) == 10 + assert "2026-08-13 10:00:00" not in note, "超出上限时应丢弃最旧的一条" + + +def test_site_oper_fail_creates_then_accumulates(db): + """ + 访问失败首次建行、其后累加,并把最后状态标记为失败。 + """ + oper = SiteOper(db=db.session) + + oper.fail("op-fail.test") + oper.fail("op-fail.test") + + stat = SiteStatistic.get_by_domain(db.session, "op-fail.test") + assert (stat.fail, stat.lst_state) == (2, 1) + + +def test_site_oper_async_success_and_fail_match_sync(db): + """ + 异步的成功/失败统计与同步同规则:首次建行、其后累加。 + """ + oper = SiteOper(db=db.session) + + asyncio.run(oper.async_success("op-async-stat.test", seconds=2)) + asyncio.run(oper.async_success("op-async-stat.test", seconds=4)) + asyncio.run(oper.async_fail("op-async-fail.test")) + asyncio.run(oper.async_fail("op-async-fail.test")) + + assert SiteStatistic.get_by_domain(db.session, "op-async-stat.test").success == 2 + assert SiteStatistic.get_by_domain(db.session, "op-async-fail.test").fail == 2 + + +# --------------------------------------------------------------------------- # +# PluginDataOper +# --------------------------------------------------------------------------- # + +def test_plugindata_oper_save_is_upsert(db): + """ + 同一键重复保存走更新而不是新增。 + + 新增会让读取拿到旧值(取到第一条),插件配置改了却不生效。 + """ + oper = PluginDataOper(db=db.session) + + oper.save("PluginX", "k", {"v": 1}) + oper.save("PluginX", "k", {"v": 2}) + + assert oper.get_data("PluginX", "k") == {"v": 2} + assert len(oper.get_data_all("PluginX")) == 1 + + +def test_plugindata_oper_get_without_key_returns_all_rows(db): + """ + 不给键时返回该插件的全部数据行;键不存在时返回 None。 + """ + oper = PluginDataOper(db=db.session) + oper.save("PluginY", "a", {"v": 1}) + oper.save("PluginY", "b", {"v": 2}) + + assert {row.key for row in oper.get_data("PluginY")} == {"a", "b"} + assert oper.get_data("PluginY", "missing") is None + + +def test_plugindata_oper_delete_scopes_by_key_then_plugin(db): + """ + 删除可精确到键,也可整插件清空,且都不波及其他插件。 + """ + oper = PluginDataOper(db=db.session) + oper.save("PluginZ", "a", {"v": 1}) + oper.save("PluginZ", "b", {"v": 2}) + oper.save("PluginW", "a", {"v": 3}) + + oper.del_data("PluginZ", "a") + assert {row.key for row in oper.get_data("PluginZ")} == {"b"} + + oper.del_data("PluginZ") + assert oper.get_data("PluginZ") == [] + assert oper.get_data("PluginW", "a") == {"v": 3} + + +def test_plugindata_oper_async_accessors_match_sync(db): + """ + 异步读写与同步一致。 + """ + oper = PluginDataOper(db=db.session) + asyncio.run(oper.async_save("PluginAsync", "k", {"v": 1})) + asyncio.run(oper.async_save("PluginAsync", "k", {"v": 2})) + + assert asyncio.run(oper.async_get_data("PluginAsync", "k")) == {"v": 2} + assert asyncio.run(oper.async_get_data("PluginAsync", "missing")) is None + assert len(asyncio.run(oper.async_get_data_all("PluginAsync"))) == 1 + + +# --------------------------------------------------------------------------- # +# WorkflowOper +# --------------------------------------------------------------------------- # + +def _workflow_kwargs(name: str, **extra) -> dict: + """构造新增工作流的参数。""" + return dict(name=name, description=name, timer="0 * * * *", state="W", + actions=[], flows=[], context={}, execution_state={}, **extra) + + +def test_workflow_oper_add_rejects_duplicate_name(db): + """ + 同名工作流不得重复新增。 + """ + oper = WorkflowOper(db=db.session) + + assert oper.add(**_workflow_kwargs("op-wf")) == (True, "新增工作流成功") + assert oper.add(**_workflow_kwargs("op-wf")) == (False, "工作流已存在") + + +def test_workflow_oper_exposes_lists_and_lifecycle(db): + """ + 列表入口与生命周期方法都应透传到模型并落库。 + """ + oper = WorkflowOper(db=db.session) + oper.add(**_workflow_kwargs("op-wf-life", trigger_type="timer")) + flow = oper.get_by_name("op-wf-life") + + assert oper.get(flow.id).id == flow.id + assert {w.name for w in oper.list()} >= {"op-wf-life"} + assert {w.name for w in oper.list_enabled()} >= {"op-wf-life"} + assert {w.name for w in oper.get_timer_triggered_workflows()} >= {"op-wf-life"} + + oper.start(flow.id) + assert oper.get(flow.id).state == "R" + oper.step(flow.id, "a1", {"n": 1}) + assert oper.get(flow.id).current_action == "a1" + oper.success(flow.id, "完成") + assert oper.get(flow.id).state == "S" + oper.fail(flow.id, "出错") + assert oper.get(flow.id).state == "F" + oper.reset(flow.id, reset_count=True) + assert (oper.get(flow.id).state, oper.get(flow.id).run_count) == ("W", 0) + + +def test_workflow_oper_event_list_and_async_accessors(db): + """ + 事件触发列表与异步访问器同样可用。 + """ + oper = WorkflowOper(db=db.session) + oper.add(**_workflow_kwargs("op-wf-event", trigger_type="event")) + flow = oper.get_by_name("op-wf-event") + + assert {w.name for w in oper.get_event_triggered_workflows()} >= {"op-wf-event"} + assert asyncio.run(oper.async_get(flow.id)).id == flow.id + assert asyncio.run(oper.async_get_by_name("op-wf-event")).id == flow.id + assert {w.id for w in asyncio.run(oper.async_list())} >= {flow.id} + + +# --------------------------------------------------------------------------- # +# UserOper / UserConfigOper +# --------------------------------------------------------------------------- # + +def test_user_oper_reads_permissions_and_settings(db): + """ + 权限与个性化设置的读取在用户不存在时各有约定的空值。 + + 权限返回 {} 而设置返回 None——上层据此区分「没有权限」和「没有这个用户」。 + """ + oper = UserOper(db=db.session) + oper.add(name="op-user", hashed_password="x", + permissions={"discovery": True}, settings={"theme": "dark"}) + + assert oper.get_by_name("op-user").name == "op-user" + assert oper.get_permissions("op-user") == {"discovery": True} + assert oper.get_settings("op-user") == {"theme": "dark"} + assert oper.get_setting("op-user", "theme") == "dark" + assert oper.get_setting("op-user", "missing") is None + + assert oper.get_permissions("op-nobody") == {} + assert oper.get_settings("op-nobody") is None + assert oper.get_setting("op-nobody", "theme") is None + assert {u.name for u in oper.list()} >= {"op-user"} + + +def test_user_oper_async_accessors_match_sync(db): + """ + 异步按名、按 ID 取用户与同步结果一致。 + """ + oper = UserOper(db=db.session) + oper.add(name="op-user-async", hashed_password="x") + user = oper.get_by_name("op-user-async") + + assert asyncio.run(oper.async_get_by_name("op-user-async")).id == user.id + assert asyncio.run(oper.async_get_by_id(user.id)).id == user.id + + +def test_userconfig_oper_set_get_and_delete_on_empty_value(db): + """ + 用户配置写入后可读回;写入空值等同于删除该项。 + + 空值删除是「恢复默认」的实现方式,退化成写入空串会让默认值再也拿不回来。 + """ + oper = UserConfigOper() + oper.set("op-cfg-user", "theme", "dark") + + assert oper.get("op-cfg-user", "theme") == "dark" + assert oper.get("op-cfg-user")["theme"] == "dark" + assert UserConfig.get_by_key(db.session, username="op-cfg-user", key="theme") is not None + + oper.set("op-cfg-user", "theme", None) + assert UserConfig.get_by_key(db.session, username="op-cfg-user", key="theme") is None + + +def test_userconfig_oper_scopes_cache_by_username(db): + """ + 内存缓存必须按用户名隔离,且用户名为空时返回全量缓存。 + """ + oper = UserConfigOper() + oper.set("op-cfg-a", "theme", "dark") + oper.set("op-cfg-b", "theme", "light") + + assert oper.get("op-cfg-a", "theme") == "dark" + assert oper.get("op-cfg-b", "theme") == "light" + assert oper.get("op-cfg-missing", "theme") is None + assert oper.get("op-cfg-missing") is None + assert set(oper.get(None)) >= {"op-cfg-a", "op-cfg-b"} + + +# --------------------------------------------------------------------------- # +# MediaServerOper +# --------------------------------------------------------------------------- # + +def _server_item(item_id: str, **extra) -> dict: + """构造媒体服务器条目的写入参数。""" + payload = dict(server="emby", library="lib", item_id=item_id, item_type="电影", + title="片名", year="2026", media_source=TMDB, media_id="5001") + payload.update(extra) + return payload + + +def test_mediaserver_oper_add_requires_item_id(db): + """ + 缺少条目 ID 的数据不得写入——它是媒体服务器侧的唯一标识,缺了就无法再更新。 + """ + oper = MediaServerOper(db=db.session) + + assert oper.add(**_server_item("ms-1")) is True + assert oper.add(**_server_item(None)) is False + + +def test_mediaserver_oper_upsert_updates_existing_item(db): + """ + 同一服务器同一条目重复同步走更新而不是新增,返回值表示「是否新增」。 + + 调用方据这个布尔值统计本次同步新入库了多少条;把更新也算作新增会让 + 同步报告每次都显示全量新增。 + """ + oper = MediaServerOper(db=db.session) + + assert oper.upsert(**_server_item("ms-up", title="旧标题")) is True + assert oper.upsert(**_server_item("ms-up", title="新标题")) is False + + assert MediaServerItem.get_by_server_itemid(db.session, "emby", "ms-up").title == "新标题" + + +def test_mediaserver_oper_exists_by_identity_and_title(db): + """ + 存在性判断支持媒体身份与标题两条路径,条件不足时返回 None。 + """ + oper = MediaServerOper(db=db.session) + oper.add(**_server_item("ms-ex", media_id="5100", title="存在的片")) + + assert oper.exists(media_source=TMDB, media_id="5100", mtype="电影") is not None + assert oper.exists(title="存在的片", mtype="电影", year="2026") is not None + assert oper.exists(title="不存在的片") is None + assert oper.exists() is None + + +def test_mediaserver_oper_exists_checks_season_presence(db): + """ + 要求某一季时必须在季信息里真正存在,否则视为未入库。 + + 季信息缺失却判为已入库,会让整季订阅被跳过。 + """ + oper = MediaServerOper(db=db.session) + oper.add(**_server_item("ms-season", media_id="5200", item_type="电视剧", + seasoninfo={"1": [1, 2]})) + + assert oper.exists(media_source=TMDB, media_id="5200", mtype="电视剧", + season="1") is not None + assert oper.exists(media_source=TMDB, media_id="5200", mtype="电视剧", + season="2") is None + + oper.add(**_server_item("ms-noseason", media_id="5300", item_type="电视剧")) + assert oper.exists(media_source=TMDB, media_id="5300", mtype="电视剧", + season="1") is None + + +def test_mediaserver_oper_get_item_id_and_async_twins(db): + """ + 取条目 ID 与异步版本必须给出相同结果,未命中时返回 None。 + """ + oper = MediaServerOper(db=db.session) + oper.add(**_server_item("ms-id", media_id="5400")) + + assert oper.get_item_id(media_source=TMDB, media_id="5400", mtype="电影") == "ms-id" + assert oper.get_item_id(media_source=TMDB, media_id="5999", mtype="电影") is None + assert asyncio.run(oper.async_get_item_id( + media_source=TMDB, media_id="5400", mtype="电影")) == "ms-id" + assert asyncio.run(oper.async_exists(title="片名", mtype="电影", year="2026")) is not None + + +def test_mediaserver_oper_cleanup_entry_points(db): + """ + 清理入口按服务器、按同步时间、按配置列表三种口径工作。 + """ + oper = MediaServerOper(db=db.session) + oper.add(**_server_item("ms-c1", lst_mod_date="2026-08-13 12:00:00")) + oper.add(**_server_item("ms-c2", lst_mod_date="2026-01-01 12:00:00")) + + assert oper.delete_stale("emby", "2026-08-13 12:00:00") == 1 + assert oper.delete_excluded_servers(["plex"]) == 1 + + oper.add(**_server_item("ms-c3")) + oper.empty("emby") + assert MediaServerItem.get_by_itemid(db.session, "ms-c3") is None + + +# --------------------------------------------------------------------------- # +# DownloadHistoryOper +# --------------------------------------------------------------------------- # + +def test_downloadhistory_oper_get_by_hashes_returns_a_mapping(db): + """ + 批量查询返回「hash -> 历史」映射,供上层直接按 hash 取用。 + + 上层拿到列表还要自己配对,正是 N+1 的温床;这里的契约是映射。 + """ + oper = DownloadHistoryOper(db=db.session) + oper.add(path="/downloads/a", type=MediaType.TV.value, title="A", + download_hash="oh-a", date="2026-08-13 10:00:00") + oper.add(path="/downloads/b", type=MediaType.TV.value, title="B", + download_hash="oh-b", date="2026-08-13 10:00:00") + + mapping = oper.get_by_hashes(["oh-a", "oh-b", "oh-missing"]) + + assert set(mapping) == {"oh-a", "oh-b"} + assert mapping["oh-a"].title == "A" + assert oper.get_by_hashes([]) == {} + + +def test_downloadhistory_oper_file_entry_points(db): + """ + 文件记录的写入与四个读取入口构成完整闭环,删除只置状态。 + """ + oper = DownloadHistoryOper(db=db.session) + oper.add_files([ + dict(downloader="qb", download_hash="oh-f", fullpath="/downloads/f/a.mkv", + savepath="/downloads/f", filepath="a.mkv", torrentname="种子", state=1), + dict(downloader="qb", download_hash="oh-f", fullpath="/downloads/f/b.mkv", + savepath="/downloads/f", filepath="b.mkv", torrentname="种子", state=1), + ]) + + assert len(oper.get_files_by_hash("oh-f")) == 2 + assert len(oper.get_files_by_hash("oh-f", state=1)) == 2 + assert oper.get_file_by_fullpath("/downloads/f/a.mkv") is not None + assert len(oper.get_files_by_fullpath("/downloads/f/a.mkv")) == 1 + assert len(oper.get_files_by_savepath("/downloads/f")) == 2 + assert oper.get_hash_by_fullpath("/downloads/f/a.mkv") == "oh-f" + # 查不到时返回空串而非 None:调用方直接用它拼下载器请求,None 会变成字面量 "None" + assert oper.get_hash_by_fullpath("/downloads/f/none.mkv") == "" + + oper.delete_file_by_fullpath("/downloads/f/a.mkv") + assert oper.get_file_by_fullpath("/downloads/f/a.mkv").state == 0 + + +def test_downloadhistory_oper_query_entry_points(db): + """ + 路径、hash、媒体身份、分页与时间窗口五个查询入口都应透传生效。 + """ + oper = DownloadHistoryOper(db=db.session) + oper.add(path="/downloads/q", type=MediaType.TV.value, title="Q", year="2026", + media_source=TMDB, media_id="4001", seasons="S01", + download_hash="oh-q", username="alice", date="2026-08-13 10:00:00") + + assert oper.get_by_path("/downloads/q").title == "Q" + assert oper.get_by_hash("oh-q").title == "Q" + assert len(oper.get_by_media_identity(media_source=TMDB, media_id="4001")) == 1 + assert oper.list_by_page(page=1, count=1)[0].title == "Q" + assert [h.title for h in oper.list_by_user_date("2026-08-20", username="alice")] == ["Q"] + assert [h.title for h in oper.list_by_date("2026-08-01", MediaType.TV.value, + TMDB, "4001", "S01")] == ["Q"] + assert [h.title for h in oper.list_by_type(MediaType.TV.value, days=36500)] == ["Q"] + assert [h.title for h in oper.get_last_by(mtype=MediaType.TV.value, + media_source=TMDB, media_id="4001")] == ["Q"] + + +def test_downloadhistory_oper_delete_entry_points(db): + """ + 历史与文件记录的删除入口都应真正落库。 + """ + oper = DownloadHistoryOper(db=db.session) + oper.add(path="/downloads/d", type=MediaType.TV.value, title="D", + download_hash="oh-d", date="2026-08-13 10:00:00") + history = oper.get_by_hash("oh-d") + oper.add_files([dict(downloader="qb", download_hash="oh-d", + fullpath="/downloads/d/a.mkv", savepath="/downloads/d", + filepath="a.mkv", torrentname="种子", state=1)]) + file_row = oper.get_file_by_fullpath("/downloads/d/a.mkv") + + oper.delete_downloadfile(file_row.id) + assert oper.get_file_by_fullpath("/downloads/d/a.mkv") is None + + oper.delete_history(history.id) + assert oper.get_by_hash("oh-d") is None + + +def test_downloadhistory_oper_async_delete(db): + """ + 异步删除历史与同步等效。 + """ + oper = DownloadHistoryOper(db=db.session) + oper.add(path="/downloads/ad", type=MediaType.TV.value, title="AD", + download_hash="oh-ad", date="2026-08-13 10:00:00") + history = oper.get_by_hash("oh-ad") + + asyncio.run(oper.async_delete_history(history.id)) + + assert oper.get_by_hash("oh-ad") is None diff --git a/tests/test_db_oper_layer_extra.py b/tests/test_db_oper_layer_extra.py new file mode 100644 index 000000000..0b133caeb --- /dev/null +++ b/tests/test_db_oper_layer_extra.py @@ -0,0 +1,308 @@ +""" +整理历史、订阅、消息与下载冷却四个 Oper 的数据访问行为。 + +与 test_db_oper_layer 同源,拆开只是为了每个文件保持可读的长度。这一组的共同点是 +方法多、每个都很薄——薄封装最容易在参数改名或默认值上出偏差,而调用方拿到的是 +空列表或 None,看起来像「本来就没有数据」。 +""" +import asyncio + +import pytest + +from app.db.oper.downloadfailure import DownloadFailureOper +from app.db.oper.message import MessageOper +from app.db.models.downloadfailure import DownloadFailure +from app.db.models.message import Message +from app.db.models.subscribe import Subscribe +from app.db.models.subscribehistory import SubscribeHistory +from app.db.models.transferhistory import TransferHistory +from app.db.oper.subscribe import SubscribeOper +from app.db.oper.subscribehistory import SubscribeHistoryOper +from app.db.oper.transferhistory import TransferHistoryOper +from app.schemas.types import MediaSource, MediaType + +TMDB = str(MediaSource.TMDB) + + +@pytest.fixture(autouse=True) +def _track(db): + """把本文件涉及的表纳入用例级回收。""" + db.watermark(TransferHistory, Subscribe, SubscribeHistory, Message, DownloadFailure) + + +# --------------------------------------------------------------------------- # +# TransferHistoryOper +# --------------------------------------------------------------------------- # + +def _transfer_kwargs(title: str, src: str, **extra) -> dict: + """构造整理历史的写入参数。""" + payload = dict(src=src, src_storage="local", dest=f"/media/{title}.mkv", + dest_storage="local", mode="move", type=MediaType.TV.value, + title=title, year="2026", media_source=TMDB, media_id="3001", + status=True, date="2026-08-13 10:00:00", files=[]) + payload.update(extra) + return payload + + +def test_transferhistory_oper_path_lookups(db): + """ + 源路径、目标路径、成功记录三个入口都应透传存储参数。 + """ + oper = TransferHistoryOper(db=db.session) + oper.add(**_transfer_kwargs("路径", "/data/op-th.mkv")) + + assert oper.get_by_src("/data/op-th.mkv").title == "路径" + assert oper.get_by_src("/data/op-th.mkv", storage="alist") is None + assert oper.get_success_by_src("/data/op-th.mkv", storage="local").title == "路径" + assert oper.get_by_dest("/media/路径.mkv").title == "路径" + assert oper.get_by_dest("/media/路径.mkv", storage="alist") is None + + +def test_transferhistory_oper_recursive_listings(db): + """ + 源侧与目标侧的递归列举都应包含目录自身与子项,不含同前缀的兄弟目录。 + """ + oper = TransferHistoryOper(db=db.session) + oper.add(**_transfer_kwargs("自身", "/data/op-dir", dest="/media/op-dir")) + oper.add(**_transfer_kwargs("子项", "/data/op-dir/a.mkv", dest="/media/op-dir/a.mkv")) + oper.add(**_transfer_kwargs("兄弟", "/data/op-dir2/b.mkv", dest="/media/op-dir2/b.mkv")) + + assert {h.title for h in oper.list_success_by_src("/data/op-dir", recursive=True)} == \ + {"自身", "子项"} + assert {h.title for h in oper.list_success_move_by_dest("/media/op-dir", recursive=True)} == \ + {"自身", "子项"} + assert [h.title for h in oper.list_success_by_src("/data/op-dir")] == ["自身"] + assert [h.title for h in oper.list_success_move_by_dest("/media/op-dir")] == ["自身"] + + +def test_transferhistory_oper_identity_and_hash_lookups(db): + """ + 按标题、hash、媒体身份查询与按条件组合查询都应命中同一条记录。 + """ + oper = TransferHistoryOper(db=db.session) + oper.add(**_transfer_kwargs("身份", "/data/op-id.mkv", media_id="3100", + download_hash="op-th-hash", seasons="S01")) + + assert [h.title for h in oper.get_by_title("身份")] == ["身份"] + assert [h.title for h in oper.list_by_hash("op-th-hash")] == ["身份"] + assert oper.get_by_media_identity(media_source=TMDB, media_id="3100", + mtype=MediaType.TV.value).title == "身份" + assert [h.title for h in oper.get_by(mtype=MediaType.TV.value, media_source=TMDB, + media_id="3100", season="S01")] == ["身份"] + assert [h.title for h in oper.list_by_date("2026-08-01")] == ["身份"] + assert oper.statistic(days=36500) + + +def test_transferhistory_oper_add_force_replaces_same_source(db): + """ + 强制新增用同源新记录替换旧的,同一源路径只留一条。 + """ + oper = TransferHistoryOper(db=db.session) + oper.add(**_transfer_kwargs("旧记录", "/data/op-force.mkv")) + + created = oper.add_force(**_transfer_kwargs("新记录", "/data/op-force.mkv")) + + assert created.title == "新记录" + assert [h.title for h in oper.list_success_by_src("/data/op-force.mkv")] == ["新记录"] + + +def test_transferhistory_oper_update_hash_and_delete(db): + """ + 补写下载 hash、按 ID 取、删除三个入口都应真正落库。 + """ + oper = TransferHistoryOper(db=db.session) + oper.add(**_transfer_kwargs("待改", "/data/op-upd.mkv")) + history = oper.get_by_src("/data/op-upd.mkv") + + oper.update_download_hash(history.id, "op-new-hash") + assert oper.get(history.id).download_hash == "op-new-hash" + + oper.delete(history.id) + assert oper.get_by_src("/data/op-upd.mkv") is None + + +def test_transferhistory_oper_async_accessors_match_sync(db): + """ + 异步的取单条、检索、分页与计数必须与同步一致。 + """ + oper = TransferHistoryOper(db=db.session) + oper.add(**_transfer_kwargs("AsyncTitle", "/data/op-async.mkv")) + history = oper.get_by_src("/data/op-async.mkv") + + assert asyncio.run(oper.async_get(history.id)).id == history.id + assert [h.title for h in asyncio.run( + oper.async_list_by_title("AsyncTitle", count=-1))] == ["AsyncTitle"] + assert {h.id for h in asyncio.run(oper.async_list_by_page(page=1, count=100))} >= {history.id} + assert asyncio.run(oper.async_count()) >= 1 + assert asyncio.run(oper.async_count_by_title("AsyncTitle")) == 1 + + asyncio.run(oper.async_delete(history.id)) + assert oper.get(history.id) is None + + +def test_transferhistory_oper_truncate_empties_the_table(db): + """ + 清空整理历史后不再有任何记录,供「重置」入口使用。 + """ + oper = TransferHistoryOper(db=db.session) + oper.add(**_transfer_kwargs("待清", "/data/op-truncate.mkv")) + + oper.truncate() + + assert oper.get_by_src("/data/op-truncate.mkv") is None + + +# --------------------------------------------------------------------------- # +# SubscribeOper / SubscribeHistoryOper +# --------------------------------------------------------------------------- # + +def _subscribe(name: str, media_id: str = "2001", state: str = "N", + username: str = "op-alice", mtype: str = None) -> Subscribe: + """构造一条订阅记录。""" + return Subscribe(name=name, type=mtype or MediaType.TV.value, state=state, + media_source=TMDB, media_id=media_id, season=1, + username=username, date="2026-08-13 10:00:00") + + +def test_subscribe_oper_read_entry_points(db): + """ + 按 ID、按条件、按状态、按 owner、按类型五个读取入口都应透传生效。 + """ + row = db.add(_subscribe("订阅一")) + oper = SubscribeOper(db=db.session) + + assert oper.get(row.id).id == row.id + assert oper.get_by(type=MediaType.TV.value, media_source=TMDB, + media_id="2001").id == row.id + assert {s.id for s in oper.list("N")} >= {row.id} + assert {s.id for s in oper.list()} >= {row.id} + assert [s.name for s in oper.list_by_username("op-alice", state="N", + mtype=MediaType.TV.value)] == ["订阅一"] + assert [s.name for s in oper.list_by_type(MediaType.TV.value, days=36500)] == ["订阅一"] + + +def test_subscribe_oper_update_and_delete(db): + """ + 更新与删除都应落库,删除后按 ID 取不到。 + """ + row = db.add(_subscribe("待改订阅", media_id="2100")) + oper = SubscribeOper(db=db.session) + + assert oper.update(row.id, {"state": "R"}).state == "R" + + oper.delete(row.id) + assert oper.get(row.id) is None + + +def test_subscribe_oper_async_entry_points(db): + """ + 异步的取单条、按条件取、列举、更新、删除必须与同步等效。 + """ + row = db.add(_subscribe("异步订阅", media_id="2200")) + oper = SubscribeOper(db=db.session) + + assert asyncio.run(oper.async_get(row.id)).id == row.id + assert asyncio.run(oper.async_get_by(type=MediaType.TV.value, media_source=TMDB, + media_id="2200")).id == row.id + assert {s.id for s in asyncio.run(oper.async_list("N"))} >= {row.id} + assert asyncio.run(oper.async_update(row.id, {"state": "R"})).state == "R" + assert asyncio.run(oper.async_update_filter_groups(row.id, ["g1"])).filter_groups == ["g1"] + + asyncio.run(oper.async_delete(row.id)) + assert oper.get(row.id) is None + + +def test_subscribe_oper_history_round_trip(db): + """ + 订阅完成后写入历史,随后存在性判断应命中。 + + 历史判定失效会让用户能重复订阅一部已经追完的剧。 + """ + oper = SubscribeOper(db=db.session) + oper.add_history(name="历史剧", type=MediaType.TV.value, media_source=TMDB, + media_id="2300", season=1, date="2026-08-13 10:00:00", + username="op-alice", best_version=False) + + assert oper.exist_history(media_source=MediaSource.TMDB, media_id="2300", + season=1) is True + assert oper.exist_history(media_source=MediaSource.TMDB, media_id="2999", + season=1) is False + assert "历史剧" in {h.name for h in asyncio.run( + SubscribeHistoryOper(db=db.session).async_list_by_type( + mtype=MediaType.TV.value, page=1, count=100))} + + +# --------------------------------------------------------------------------- # +# MessageOper +# --------------------------------------------------------------------------- # + +def test_message_oper_add_returns_persisted_payload(db): + """ + 新增消息返回已落库的字段字典,其中必须带上主键。 + """ + oper = MessageOper(db=db.session) + + created = oper.add(title="标题", text="正文", source="op-msg-1", + reg_time="2026-08-13 10:00:00") + + assert created["id"] is not None + assert created["title"] == "标题" + assert oper.exists_by_source("op-msg-1") is True + assert oper.exists_by_source("op-msg-none") is False + + +def test_message_oper_listing_entry_points(db): + """ + 分页列举的同步与异步入口都应返回刚写入的消息。 + """ + oper = MessageOper(db=db.session) + oper.add(title="分页消息", text="正文", source="op-msg-2", + reg_time="2026-08-13 10:00:00") + + assert [m.title for m in oper.list_by_page(page=1, count=1)] == ["分页消息"] + assert [m.title for m in asyncio.run(oper.async_list_by_page(page=1, count=1))] == \ + ["分页消息"] + assert [m.title for m in asyncio.run(oper.async_list_sent_by_page(page=1, count=1))] == \ + ["分页消息"] + + +def test_message_oper_async_add_returns_the_model_not_a_dict(db): + """ + 异步新增返回的是模型实例,而同步新增返回字段字典——两者返回类型并不一致。 + + 把这条不对称固定下来:调用方按字典下标访问异步结果会直接抛 TypeError, + 这里明确它当前的契约,避免后续改写时无意中翻转。 + """ + oper = MessageOper(db=db.session) + + created = asyncio.run(oper.async_add(title="异步消息", text="正文", + source="op-msg-3", + reg_time="2026-08-13 10:00:00")) + + assert isinstance(created, Message) + assert created.id is not None + assert oper.exists_by_source("op-msg-3") is True + + +# --------------------------------------------------------------------------- # +# DownloadFailureOper +# --------------------------------------------------------------------------- # + +def test_downloadfailure_oper_round_trip(db): + """ + 记录失败、查询冷却中、清理过期三个入口构成完整闭环。 + """ + oper = DownloadFailureOper(db=db.session) + oper.record_failure(fingerprint="op-fp-1", now_time="2026-08-13 10:00:00", + next_retry_at="2026-08-13 20:00:00", title="片名") + oper.record_failure(fingerprint="op-fp-old", now_time="2026-01-01 10:00:00", + next_retry_at="2026-01-01 20:00:00", title="旧的") + + # Oper 侧返回「指纹 -> 记录」映射,供上层直接按指纹判定是否仍在冷却 + active = oper.get_active_by_fingerprints(["op-fp-1", "op-fp-old"], + now_time="2026-08-13 12:00:00") + assert set(active) == {"op-fp-1"} + assert active["op-fp-1"].title == "片名" + + assert oper.delete_expired(before_time="2026-08-01", limit=100) == 1 + assert oper.delete_expired(before_time="2026-08-01", limit=100) == 0 diff --git a/tests/test_db_plugin_message_agent_queries.py b/tests/test_db_plugin_message_agent_queries.py new file mode 100644 index 000000000..2579784e4 --- /dev/null +++ b/tests/test_db_plugin_message_agent_queries.py @@ -0,0 +1,456 @@ +""" +插件数据、消息、Agent 会话、Agent 定时任务与下载失败冷却五张表的查询行为。 + +这一组的共同风险是「按用户/插件归属收窄」和「分页 + 排序」:归属条件丢失就是越权, +分页排序错乱则表现为消息重复或漏掉,两者都不会抛异常。 +""" +import asyncio + +import pytest + +from app.db.models.agentchat import AgentChat +from app.db.models.agenttask import AgentTask +from app.db.models.downloadfailure import DownloadFailure +from app.db.models.message import Message +from app.db.models.plugindata import PluginData + + +@pytest.fixture(autouse=True) +def _track(db): + """把本文件涉及的表纳入用例级回收。""" + db.watermark(PluginData, Message, AgentChat, AgentTask, DownloadFailure) + + +# --------------------------------------------------------------------------- # +# PluginData +# --------------------------------------------------------------------------- # + +def test_plugindata_is_scoped_by_plugin_id(db): + """ + 插件数据必须按插件隔离——串读会让一个插件拿到另一个插件的配置。 + """ + db.add(PluginData(plugin_id="PluginA", key="k1", value={"v": 1}), + PluginData(plugin_id="PluginA", key="k2", value={"v": 2}), + PluginData(plugin_id="PluginB", key="k1", value={"v": 3})) + + rows = PluginData.get_plugin_data(db.session, "PluginA") + + assert {r.key for r in rows} == {"k1", "k2"} + assert {r.key for r in asyncio.run(PluginData.async_get_plugin_data(plugin_id="PluginA"))} \ + == {"k1", "k2"} + + +def test_plugindata_get_by_key_needs_both_plugin_and_key(db): + """ + 按键取值必须同时匹配插件与键,只匹配键会取到同名键的别家数据。 + """ + db.add(PluginData(plugin_id="PluginA", key="shared", value={"v": 1}), + PluginData(plugin_id="PluginB", key="shared", value={"v": 2})) + + assert PluginData.get_plugin_data_by_key(db.session, "PluginA", "shared").value == {"v": 1} + assert PluginData.get_plugin_data_by_key(db.session, "PluginB", "shared").value == {"v": 2} + assert PluginData.get_plugin_data_by_key(db.session, "PluginC", "shared") is None + assert asyncio.run(PluginData.async_get_plugin_data_by_key( + plugin_id="PluginA", key="shared")).value == {"v": 1} + + +def test_plugindata_delete_by_key_removes_only_that_entry(db): + """ + 删除单个键不能波及同插件的其他键,也不能波及别的插件。 + """ + db.add(PluginData(plugin_id="PluginA", key="drop", value={"v": 1}), + PluginData(plugin_id="PluginA", key="keep", value={"v": 2}), + PluginData(plugin_id="PluginB", key="drop", value={"v": 3})) + + PluginData.del_plugin_data_by_key(db.session, "PluginA", "drop") + + assert PluginData.get_plugin_data_by_key(db.session, "PluginA", "drop") is None + assert PluginData.get_plugin_data_by_key(db.session, "PluginA", "keep") is not None + assert PluginData.get_plugin_data_by_key(db.session, "PluginB", "drop") is not None + + +def test_plugindata_delete_all_clears_only_that_plugin(db): + """ + 卸载插件时清空其数据,不能连带清掉其他插件——那等于误删用户配置。 + """ + db.add(PluginData(plugin_id="PluginA", key="k1", value={"v": 1}), + PluginData(plugin_id="PluginB", key="k1", value={"v": 2})) + + PluginData.del_plugin_data(db.session, "PluginA") + + assert PluginData.get_plugin_data(db.session, "PluginA") == [] + assert len(PluginData.get_plugin_data_by_plugin_id(db.session, "PluginB")) == 1 + + +# --------------------------------------------------------------------------- # +# Message +# --------------------------------------------------------------------------- # + +def _message(reg_time: str, title: str, source: str = None, action: int = 1, + image: str = None) -> Message: + """构造一条消息记录。""" + return Message(channel="wechat", source=source, mtype="Manual", title=title, + text=title, reg_time=reg_time, action=action, image=image) + + +def test_message_list_by_page_is_newest_first_and_paged(db): + """ + 消息列表按登记时间倒序、同时间按主键倒序,并遵守分页。 + + 排序不稳定时相邻两页会出现重复或漏掉的消息,用户看到的是「消息丢了」。 + """ + for index in range(5): + db.add(_message(f"2026-08-13 10:00:0{index}", f"msg-{index}")) + + first_page = Message.list_by_page(db.session, page=1, count=2) + second_page = Message.list_by_page(db.session, page=2, count=2) + + assert [m.title for m in first_page] == ["msg-4", "msg-3"] + assert [m.title for m in second_page] == ["msg-2", "msg-1"] + + +def test_message_list_by_page_matches_async_twin(db): + """ + 同步与异步分页必须返回同一批消息,前端两条链路才不会互相矛盾。 + """ + for index in range(3): + db.add(_message(f"2026-08-13 11:00:0{index}", f"par-{index}")) + + sync_titles = [m.title for m in Message.list_by_page(db.session, page=1, count=3)] + async_titles = [m.title for m in asyncio.run(Message.async_list_by_page(page=1, count=3))] + + assert sync_titles == async_titles + + +def test_message_exists_by_source_detects_duplicates(db): + """ + 来源标识存在性判断用于消息去重,判错会导致同一条通知重复推送。 + """ + db.add(_message("2026-08-13 10:00:00", "有来源", source="uniq-source-1")) + + assert Message.exists_by_source(db.session, "uniq-source-1") is True + assert Message.exists_by_source(db.session, "uniq-source-missing") is False + + +def test_message_delete_before_is_batched_and_keeps_recent(db): + """ + 历史消息清理必须分批、遵守上限,且不碰保留期内的消息。 + """ + for index in range(4): + db.add(_message(f"2026-01-01 10:00:0{index}", f"old-{index}")) + recent = db.add(_message("2026-08-13 10:00:00", "recent")) + + assert Message.delete_before(db.session, before_time="2026-08-01", limit=2) == 2 + assert Message.delete_before(db.session, before_time="2026-08-01", limit=100) == 2 + assert Message.delete_before(db.session, before_time="2026-08-01", limit=100) == 0 + + assert Message.list_by_page(db.session, page=1, count=1)[0].id == recent.id + + +def test_message_delete_before_keeps_the_row_exactly_at_the_boundary(db): + """ + 保留时间点上的消息属于「保留期内」,不能被清理掉(``reg_time < before_time``)。 + + 比较符若写成 ``<=``,每次清理都会多吃掉恰好落在保留起点的那一批消息; + 数据从不压在边界上时这一字之差完全不可观测,故此处专门把行摆在边界上。 + """ + boundary = "2026-05-01 00:00:00" + at_boundary = db.add(_message(boundary, "边界上")) + db.add(_message("2026-04-30 23:59:59", "边界前一秒")) + + assert Message.delete_before(db.session, before_time=boundary, limit=100) == 1 + + assert db.session.get(Message, at_boundary.id) is not None + + +def test_message_async_list_sent_excludes_the_clear_boundary(db): + """ + 三个清理水位都取「严格晚于水位」的消息,正好落在水位上的必须被滤掉。 + + 水位是「本次清空动作发生的时刻」,与它同一秒的消息属于已清空的那一批; + 比较符放宽成 ``>=`` 会让用户清空后又看见最后一条旧消息。 + """ + boundary, after = "2026-03-01 10:00:00", "2026-03-01 10:00:01" + db.add(_message(boundary, "bd-系统-边界上"), + _message(after, "bd-系统-边界后"), + _message(boundary, "bd-媒体-边界上", image="http://img/1.jpg"), + _message(after, "bd-媒体-边界后", image="http://img/2.jpg")) + + def _titles(**clears) -> set: + """取本用例写入的消息标题集合,隔离其他用例可能残留的消息。""" + rows = asyncio.run(Message.async_list_sent_by_page(page=1, count=100, **clears)) + return {m.title for m in rows if m.title.startswith("bd-")} + + # 全量清空水位:边界上的两条都属于被清空的那一批 + assert _titles(all_clear_before=boundary) == {"bd-系统-边界后", "bd-媒体-边界后"} + # 系统消息(无图)清空水位:只吃无图消息,带图的媒体消息不受影响 + assert _titles(system_clear_before=boundary) == { + "bd-系统-边界后", "bd-媒体-边界上", "bd-媒体-边界后"} + # 媒体消息(有图)清空水位:只吃带图消息,无图的系统消息不受影响 + assert _titles(media_clear_before=boundary) == { + "bd-系统-边界上", "bd-系统-边界后", "bd-媒体-边界后"} + + +def test_message_create_and_to_dict_returns_persisted_fields(db): + """ + 创建后返回的字典必须已带上数据库生成的主键。 + + 返回未落库的字段会让调用方拿到 id 为 None 的消息,后续更新无从下手。 + """ + created = _message("2026-08-13 10:00:00", "新消息").create_and_to_dict(db.session) + + assert created["id"] is not None + assert created["title"] == "新消息" + + +# --------------------------------------------------------------------------- # +# AgentChat +# --------------------------------------------------------------------------- # + +def _chat(session_id: str, user_id: str = "u1", updated_at: str = "2026-08-13 10:00:00", + username: str = None) -> AgentChat: + """构造一条 Agent 会话记录。""" + return AgentChat(session_id=session_id, user_id=user_id, username=username, + channel="web", title=session_id, updated_at=updated_at, + created_at=updated_at, message_count=0) + + +def test_agentchat_get_by_session_takes_the_newest_row(db): + """ + 同一会话 ID 存在多行时取主键最大的那条——它才是最新的会话状态。 + """ + db.add(_chat("s-dup"), _chat("s-dup")) + newest = db.add(_chat("s-dup")) + + assert AgentChat.get_by_session(db.session, "s-dup").id == newest.id + assert asyncio.run(AgentChat.async_get_by_session(session_id="s-dup")).id == newest.id + + +def test_agentchat_get_by_session_enforces_user_scope(db): + """ + 传入用户 ID 时必须同时匹配,否则一个用户能读到另一个用户的会话内容。 + """ + db.add(_chat("s-owned", user_id="alice")) + + assert AgentChat.get_by_session(db.session, "s-owned", user_id="alice") is not None + assert AgentChat.get_by_session(db.session, "s-owned", user_id="bob") is None + assert asyncio.run(AgentChat.async_get_by_session(session_id="s-owned", user_id="bob")) is None + + +def test_agentchat_list_by_page_matches_either_user_or_username(db): + """ + 同时给出用户 ID 与用户名时按「或」匹配。 + + 渠道侧只有用户名、前端只有用户 ID,改成「与」会让两边各自都查不到自己的会话。 + """ + db.add(_chat("s-by-id", user_id="uid-1", username=None), + _chat("s-by-name", user_id="uid-other", username="alice"), + _chat("s-neither", user_id="uid-x", username="bob")) + + listed = AgentChat.list_by_page(db.session, user_id="uid-1", username="alice") + + assert {c.session_id for c in listed} == {"s-by-id", "s-by-name"} + + +@pytest.mark.parametrize("kwargs,expected", [ + ({"user_id": "uid-1"}, {"s-by-id"}), + ({"username": "alice"}, {"s-by-name"}), +]) +def test_agentchat_list_by_page_single_scope(db, kwargs, expected): + """ + 只给用户 ID 或只给用户名时,各自按单一条件收窄。 + """ + db.add(_chat("s-by-id", user_id="uid-1", username=None), + _chat("s-by-name", user_id="uid-other", username="alice")) + + assert {c.session_id for c in AgentChat.list_by_page(db.session, **kwargs)} == expected + + +def test_agentchat_list_by_page_is_newest_first_and_paged(db): + """ + 会话列表按更新时间倒序分页,顺序错乱会让用户的最近会话沉到后面。 + """ + for index in range(4): + db.add(_chat(f"s-p{index}", user_id="uid-page", + updated_at=f"2026-08-13 10:00:0{index}")) + + page1 = AgentChat.list_by_page(db.session, page=1, count=2, user_id="uid-page") + page2 = AgentChat.list_by_page(db.session, page=2, count=2, user_id="uid-page") + + assert [c.session_id for c in page1] == ["s-p3", "s-p2"] + assert [c.session_id for c in page2] == ["s-p1", "s-p0"] + assert [c.session_id for c in asyncio.run( + AgentChat.async_list_by_page(page=1, count=2, user_id="uid-page"))] == ["s-p3", "s-p2"] + + +# --------------------------------------------------------------------------- # +# AgentTask +# --------------------------------------------------------------------------- # + +def _task(name: str, user_id: str = "u1", enabled: bool = True, + created_at: str = "2026-08-13 10:00:00") -> dict: + """构造 Agent 定时任务的新增参数。""" + return dict(name=name, content="做点什么", trigger_type="cron", + cron_expression="0 * * * *", enabled=enabled, user_id=user_id, + session_id=f"sess-{name}", created_at=created_at, + updated_at=created_at, last_status="waiting", run_count=0) + + +def test_agenttask_get_for_user_enforces_ownership(db): + """ + 带用户 ID 查询时必须匹配归属,否则任意用户都能读到别人的定时任务。 + """ + task_id = AgentTask.add_task(db.session, **_task("t1", user_id="alice")) + + assert AgentTask.get_for_user(db.session, task_id).id == task_id + assert AgentTask.get_for_user(db.session, task_id, user_id="alice").id == task_id + assert AgentTask.get_for_user(db.session, task_id, user_id="bob") is None + + +def test_agenttask_list_for_user_filters_by_owner_and_enabled(db): + """ + 列表按归属与启用状态收窄,并按创建时间倒序。 + + 调度器取的是「已启用」这一批,条件失效会把用户停掉的任务重新跑起来。 + """ + AgentTask.add_task(db.session, **_task("t-on", user_id="alice", + created_at="2026-08-13 10:00:00")) + AgentTask.add_task(db.session, **_task("t-off", user_id="alice", enabled=False, + created_at="2026-08-13 11:00:00")) + AgentTask.add_task(db.session, **_task("t-other", user_id="bob")) + + mine = AgentTask.list_for_user(db.session, user_id="alice") + assert [t.name for t in mine] == ["t-off", "t-on"] + + assert [t.name for t in AgentTask.list_for_user(db.session, user_id="alice", enabled=True)] \ + == ["t-on"] + assert [t.name for t in AgentTask.list_for_user(db.session, user_id="alice", enabled=False)] \ + == ["t-off"] + + +def test_agenttask_update_enforces_ownership(db): + """ + 更新必须校验归属,并如实返回是否命中。 + + 删除、认领执行(mark_running)与收尾计数(finish_task)已随运行记录的引入迁出本模型, + 改由 AgentTaskRun / AgentTaskOper 承担,对应用例见 tests/test_agent_task_runs.py 与 + tests/test_agent_scheduled_tasks.py,此处不再重复覆盖。 + """ + task_id = AgentTask.add_task(db.session, **_task("t-own", user_id="alice")) + + assert AgentTask.update_task(db.session, task_id, {"name": "改名"}, user_id="bob") is False + assert AgentTask.update_task(db.session, task_id, {"name": "改名"}, user_id="alice") is True + assert AgentTask.get_for_user(db.session, task_id).name == "改名" + + +# --------------------------------------------------------------------------- # +# DownloadFailure +# --------------------------------------------------------------------------- # + +def _failure(fingerprint: str, next_retry_at: str) -> dict: + """构造下载失败冷却记录的写入参数。""" + return dict(fingerprint=fingerprint, now_time="2026-08-13 10:00:00", + next_retry_at=next_retry_at, title="片名", type="电影") + + +def test_download_failure_active_lookup_excludes_expired_cooldowns(db): + """ + 只返回仍在冷却期内的记录。 + + 冷却已过却仍被判为「冷却中」,资源会被永久跳过、订阅永远下不下来。 + """ + DownloadFailure.record_failure(db.session, **_failure("fp-cold", "2026-08-13 20:00:00")) + DownloadFailure.record_failure(db.session, **_failure("fp-expired", "2026-08-13 09:00:00")) + + active = DownloadFailure.get_active_by_fingerprints( + db.session, ["fp-cold", "fp-expired"], now_time="2026-08-13 12:00:00") + + assert [f.fingerprint for f in active] == ["fp-cold"] + + +def test_download_failure_active_lookup_excludes_the_expiry_boundary(db): + """ + 冷却到点即结束:``next_retry_at`` 恰好等于当前时刻的记录不再算「冷却中」。 + + 条件是 ``next_retry_at > now_time``;放宽成 ``>=`` 会让资源在到点那一秒仍被跳过, + 而两侧数据都离边界一小时时,这一字之差查不出来。 + """ + now_time = "2026-08-13 12:00:00" + DownloadFailure.record_failure(db.session, **_failure("fp-at-boundary", now_time)) + DownloadFailure.record_failure( + db.session, **_failure("fp-past-boundary", "2026-08-13 12:00:01")) + + active = DownloadFailure.get_active_by_fingerprints( + db.session, ["fp-at-boundary", "fp-past-boundary"], now_time=now_time) + + assert [f.fingerprint for f in active] == ["fp-past-boundary"] + + +def test_download_failure_active_lookup_dedupes_and_ignores_blanks(db): + """ + 指纹列表去重并剔除空值,空列表直接短路返回。 + + 条件为空的 IN 查询在部分方言下会退化成全表匹配,把所有资源判成冷却中。 + """ + DownloadFailure.record_failure(db.session, **_failure("fp-a", "2026-08-13 20:00:00")) + + assert DownloadFailure.get_active_by_fingerprints(db.session, [], "2026-08-13 12:00:00") == [] + assert DownloadFailure.get_active_by_fingerprints( + db.session, ["", None], "2026-08-13 12:00:00") == [] + + found = DownloadFailure.get_active_by_fingerprints( + db.session, ["fp-a", "fp-a", ""], "2026-08-13 12:00:00") + assert [f.fingerprint for f in found] == ["fp-a"] + + +def test_download_failure_record_increments_retry_count(db): + """ + 同一指纹再次失败时累加重试次数,而不是新增一行。 + + 每次新增会让冷却窗口永远停留在第一档,退避策略形同虚设。 + """ + first = DownloadFailure.record_failure(db.session, **_failure("fp-retry", "2026-08-13 20:00:00")) + assert first.retry_count == 1 + + second = DownloadFailure.record_failure( + db.session, **_failure("fp-retry", "2026-08-14 20:00:00")) + + assert second.id == first.id + assert second.retry_count == 2 + assert second.next_retry_at == "2026-08-14 20:00:00" + + +def test_download_failure_delete_expired_is_batched(db): + """ + 过期记录清理分批执行,且不碰仍在冷却期内的记录。 + """ + for index in range(3): + DownloadFailure.record_failure( + db.session, **_failure(f"fp-old-{index}", "2026-01-0%d 10:00:00" % (index + 1))) + DownloadFailure.record_failure(db.session, **_failure("fp-live", "2026-12-01 10:00:00")) + + assert DownloadFailure.delete_expired(db.session, before_time="2026-08-01", limit=2) == 2 + assert DownloadFailure.delete_expired(db.session, before_time="2026-08-01", limit=100) == 1 + assert DownloadFailure.delete_expired(db.session, before_time="2026-08-01", limit=100) == 0 + + assert DownloadFailure.get_active_by_fingerprints( + db.session, ["fp-live"], "2026-08-13 12:00:00") + + +def test_download_failure_delete_expired_keeps_the_row_exactly_at_the_boundary(db): + """ + ``next_retry_at`` 恰好等于清理水位的记录不算过期,必须留下(``next_retry_at < before_time``)。 + + 比较符若写成 ``<=``,正好排到水位那一秒的冷却记录会被提前抹掉,该资源随即被重新 + 下载一遍——退避直接失效。上面那条分批用例的数据离水位有半年之遥,压不到边界。 + """ + boundary = "2026-05-01 00:00:00" + DownloadFailure.record_failure(db.session, **_failure("fp-at-boundary", boundary)) + DownloadFailure.record_failure( + db.session, **_failure("fp-before-boundary", "2026-04-30 23:59:59")) + + assert DownloadFailure.delete_expired(db.session, before_time=boundary, limit=100) == 1 + + remaining = DownloadFailure.get_active_by_fingerprints( + db.session, ["fp-at-boundary", "fp-before-boundary"], now_time="2026-01-01 00:00:00") + assert [f.fingerprint for f in remaining] == ["fp-at-boundary"] diff --git a/tests/test_db_public_api.py b/tests/test_db_public_api.py new file mode 100644 index 000000000..406a8b827 --- /dev/null +++ b/tests/test_db_public_api.py @@ -0,0 +1,65 @@ +"""``app.db`` 对外契约的边界。 + +``__all__`` 是这个包唯一一份机器可读的对外承诺,而当前的写法**看起来像个疏漏**: +``SessionFactory`` / ``AsyncSessionFactory`` / ``ScopedSession`` 三个名字在 +``app/db/__init__.py`` 里明明 import 了,却不出现在 ``__all__`` 中。下一个读到这段代码 +的人很容易顺手把它们补回去——那会在无人察觉的情况下把契约重新放宽。 + +这三个名字建出来的是**绕过事务装饰器**的裸会话:不提交、不回滚、不释放,全靠调用方 +自己兜底。它们之所以还留在模块里,只是因为包内的 scheduler、postgresql 模块与 Alembic +迁移脚本用直接导入的方式在用(直接导入不受 ``__all__`` 约束)。 + +所以这里同时钉两头:契约里没有它们,但包内的既有导入不能被这个决定误伤。 +""" +import app.db as db_package + + +# 降级为内部实现细节的三个会话工厂 +INTERNAL_FACTORY_NAMES = ("SessionFactory", "AsyncSessionFactory", "ScopedSession") + + +def test_session_factories_are_not_part_of_the_public_contract(): + """ + 三个会话工厂不得出现在 ``__all__`` 里。 + + 插件要访问数据库应走 ``DbOper`` 子类或 ``db_query`` / ``async_db_query`` 装饰器, + 由装饰器收口会话的提交、回滚与释放。 + """ + leaked = [name for name in INTERNAL_FACTORY_NAMES if name in db_package.__all__] + assert not leaked, f"会话工厂被重新放进了对外契约:{leaked}" + + +def test_session_factories_remain_importable_for_in_repo_callers(): + """ + 契约收窄不等于删除:包内既有的直接导入必须照常可用。 + + ``app/scheduler.py``、``app/modules/postgresql/__init__.py`` 与 + ``database/versions/*.py`` 都在用 ``from app.db import SessionFactory``, + 这类直接导入本就不受 ``__all__`` 影响,此处显式钉住以免连带删除。 + """ + for name in INTERNAL_FACTORY_NAMES: + assert callable(getattr(db_package, name, None)), f"{name} 不再可从 app.db 导入" + + +def test_engines_stay_in_the_public_contract(): + """ + ``Engine`` / ``AsyncEngine`` 留在契约内。 + + 建表、Alembic 迁移与连接诊断确实需要引擎**对象**本身,事务装饰器覆盖不到这些用途, + 仓库外拿它是正当的。 + + 只断言名字在不在 ``__all__``,不去真的取它:这两个名字由模块级 ``__getattr__`` 解析, + 取属性即创建引擎,而本用例并不想在测试进程里凭空建一个没人释放的异步引擎。 + """ + for name in ("Engine", "AsyncEngine"): + assert name in db_package.__all__, f"{name} 被移出了对外契约" + + +def test_documented_plugin_entrypoints_are_exported(): + """ + 契约里必须留有插件真正该走的那条路:``DbOper`` 基类与四个事务装饰器。 + + 否则「工厂降级为内部细节」就成了一句没有出口的话。 + """ + for name in ("DbOper", "db_query", "db_update", "async_db_query", "async_db_update"): + assert name in db_package.__all__, f"插件数据访问入口 {name} 不在 __all__ 中" diff --git a/tests/test_db_session_lifecycle.py b/tests/test_db_session_lifecycle.py new file mode 100644 index 000000000..d24a00ee5 --- /dev/null +++ b/tests/test_db_session_lifecycle.py @@ -0,0 +1,217 @@ +""" +数据库会话生命周期与资源释放测试。 + +会话生成器(FastAPI 依赖注入入口)必须在请求结束时归还连接,close_database 必须 +释放全部引擎——池化之后引擎不再只有一个:除全局同步/异步引擎外,还有按事件循环 +缓存的池化引擎,漏掉任何一类都是连接泄漏。 +""" +import asyncio +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from app.runtime.config import global_vars, settings +from app.db import engine as engine_module +from app.db import session as session_module + + +@pytest.fixture(autouse=True) +def _restore_pooled_engines(): + """ + 复原按循环缓存的引擎,避免用例间相互污染。 + """ + saved = dict(session_module._pooled_async_engines) + yield + session_module._pooled_async_engines.clear() + session_module._pooled_async_engines.update(saved) + + +def test_get_db_closes_session_on_exit(monkeypatch): + """ + 同步会话生成器必须在迭代结束后关闭会话,否则连接不会归还连接池。 + """ + closed = [] + fake = MagicMock() + fake.close = lambda: closed.append(1) + monkeypatch.setattr(session_module, "SessionFactory", lambda: fake) + + gen = session_module.get_db() + assert next(gen) is fake + with pytest.raises(StopIteration): + next(gen) + + assert closed, "生成器正常结束时未关闭会话" + + +def test_get_db_closes_session_even_on_error(monkeypatch): + """ + 调用方提前中止时同样要归还会话——否则一次请求失败就泄漏一条连接。 + """ + closed = [] + fake = MagicMock() + fake.close = lambda: closed.append(1) + monkeypatch.setattr(session_module, "SessionFactory", lambda: fake) + + gen = session_module.get_db() + next(gen) + gen.close() + + assert closed, "生成器被中止时未关闭会话" + + +def test_get_async_db_yields_session_from_scope(monkeypatch): + """ + 异步会话入口必须经 async_session_scope 获取——池化与配额都在那里收口, + 绕过它会同时失去连接复用和背压。 + """ + used = [] + + class _Scope: + """会话作用域替身,记录进入与退出。""" + + async def __aenter__(self): + used.append("enter") + return "SESSION" + + async def __aexit__(self, *_exc): + used.append("exit") + return False + + monkeypatch.setattr(session_module, "async_session_scope", lambda: _Scope()) + + async def run(): + gen = session_module.get_async_db() + got = await gen.__anext__() + with pytest.raises(StopAsyncIteration): + await gen.__anext__() + return got + + assert asyncio.run(run()) == "SESSION" + assert used == ["enter", "exit"], "会话作用域未正确进入/退出" + + +def test_close_database_disposes_pooled_engines(monkeypatch): + """ + close_database 必须释放按事件循环缓存的池化引擎。 + + 池化之后引擎不再只有全局那一个,漏掉缓存中的引擎意味着进程退出时 + 仍持有未归还的物理连接。 + """ + sync_engine = MagicMock() + async_engine = MagicMock(dispose=AsyncMock()) + pooled_a = MagicMock(dispose=AsyncMock()) + pooled_b = MagicMock(dispose=AsyncMock()) + + monkeypatch.setattr(engine_module, "_sync_engine", sync_engine) + monkeypatch.setattr(engine_module, "_async_engine", async_engine) + session_module._pooled_async_engines.clear() + session_module._pooled_async_engines.update({1: pooled_a, 2: pooled_b}) + + asyncio.run(session_module.close_database()) + + sync_engine.dispose.assert_called_once() + async_engine.dispose.assert_awaited_once() + pooled_a.dispose.assert_awaited_once() + pooled_b.dispose.assert_awaited_once() + assert not session_module._pooled_async_engines, "释放后未清空缓存" + + +def test_close_database_does_not_create_engines_to_dispose_them(monkeypatch): + """ + 两个引擎槽都是空的时候,close_database 不得为了 dispose 而把引擎创建出来。 + + 这是惰性化的直接后果,也是最容易在重构中丢掉的一条:写成 `Engine.dispose()` + 同样能跑通上面那几个用例——它们都把 MagicMock 塞进了引擎槽,`is not None` 恒真, + 于是「先创建再释放」和「有才释放」在测试里完全等价。因此必须单独用空槽压一次: + 否则一个从未用过数据库的进程会在关停时凭空连一次库,只为了随后释放它。 + """ + created = [] + monkeypatch.setattr(engine_module, "_sync_engine", None) + monkeypatch.setattr(engine_module, "_async_engine", None) + monkeypatch.setattr(engine_module, "_get_database_engine", + lambda **kw: created.append(kw) or MagicMock(dispose=AsyncMock())) + session_module._pooled_async_engines.clear() + + asyncio.run(session_module.close_database()) + + assert created == [], f"close_database 为了 dispose 创建了引擎:{created}" + assert engine_module._sync_engine is None, "同步引擎槽被 close_database 填上了" + assert engine_module._async_engine is None, "异步引擎槽被 close_database 填上了" + + +def test_close_database_continues_after_single_engine_failure(monkeypatch): + """ + 单个引擎释放失败不能中断其余引擎的释放,否则一个坏连接会让其他连接全部泄漏。 + """ + failing = MagicMock(dispose=AsyncMock(side_effect=RuntimeError("connection reset"))) + healthy = MagicMock(dispose=AsyncMock()) + + monkeypatch.setattr(engine_module, "_sync_engine", MagicMock()) + monkeypatch.setattr(engine_module, "_async_engine", MagicMock(dispose=AsyncMock())) + session_module._pooled_async_engines.clear() + session_module._pooled_async_engines.update({1: failing, 2: healthy}) + + asyncio.run(session_module.close_database()) + + healthy.dispose.assert_awaited_once() + + +@pytest.mark.parametrize("failing", ["sync", "async"]) +def test_close_database_releases_remaining_engines_after_global_failure(monkeypatch, failing): + """ + 全局引擎释放失败:既不能抛出,也不能连累后面的引擎。 + + 「不抛」是因为 close_database 在关闭流程末尾调用,抛异常会掩盖其他关闭步骤的问题。 + 但只断言「不抛」是不够的——在外面套一个大 try 同样不抛,代价是同步引擎一出错, + 异步引擎和全部池化引擎就都跳过了释放:一条坏连接拖着其余连接一起泄漏, + 而这恰恰是兄弟用例 test_close_database_continues_after_single_engine_failure + 的 docstring 已经声称过的不变量。所以这里把它真正钉住:坏的那个失败,其余照常释放。 + """ + sync_engine = MagicMock() + async_engine = MagicMock(dispose=AsyncMock()) + pooled = MagicMock(dispose=AsyncMock()) + if failing == "sync": + sync_engine.dispose.side_effect = RuntimeError("boom") + else: + async_engine.dispose.side_effect = RuntimeError("boom") + + monkeypatch.setattr(engine_module, "_sync_engine", sync_engine) + monkeypatch.setattr(engine_module, "_async_engine", async_engine) + session_module._pooled_async_engines.clear() + session_module._pooled_async_engines.update({1: pooled}) + + asyncio.run(session_module.close_database()) # 不抛异常 + + # 出错的那个也得真被尝试过,排在它后面的一个都不能少 + sync_engine.dispose.assert_called_once() + async_engine.dispose.assert_awaited_once() + pooled.dispose.assert_awaited_once() + assert not session_module._pooled_async_engines, "释放后未清空缓存" + + +def test_pooled_engine_is_reused_within_same_loop(monkeypatch): + """ + 同一事件循环内必须复用同一个池化引擎实例。 + + 每次新建引擎等于每次新建一个连接池,连接无法复用,池化就退化回了 NullPool + 的行为——只是多了一层包装。 + """ + monkeypatch.setattr(settings, "DB_ASYNC_POOL_TYPE", "QueuePool", raising=False) + created = [] + monkeypatch.setattr(session_module, "_get_database_engine", + lambda **kw: created.append(kw) or MagicMock()) + + async def run(): + global_vars.CURRENT_EVENT_LOOP = asyncio.get_running_loop() + session_module._pooled_async_engines.clear() + first = session_module.get_async_engine() + second = session_module.get_async_engine() + return first is second + + saved = global_vars.CURRENT_EVENT_LOOP + try: + assert asyncio.run(run()) is True + assert len(created) == 1, f"引擎被重复创建 {len(created)} 次" + assert created[0]["pooled"] is True + finally: + global_vars.CURRENT_EVENT_LOOP = saved diff --git a/tests/test_db_site_queries.py b/tests/test_db_site_queries.py new file mode 100644 index 000000000..c5d263fc2 --- /dev/null +++ b/tests/test_db_site_queries.py @@ -0,0 +1,280 @@ +""" +站点相关四张表的查询行为:站点、图标、访问统计、用户数据快照。 + +站点数据快照的 get_latest 是这里唯一带子查询与 JOIN 的查询——「每个站点取最新一天」 +用普通过滤写不出来,改写时最容易退化成「全表按时间倒序取第一条」,那样多站点场景下 +只会剩一个站点的数据,而页面不会报错,只是少了几行。 +""" +import asyncio + +import pytest + +from app.db.models.site import Site +from app.db.models.siteicon import SiteIcon +from app.db.models.sitestatistic import SiteStatistic +from app.db.models.siteuserdata import SiteUserData + + +@pytest.fixture(autouse=True) +def _track(db): + """把站点相关表纳入用例级回收。""" + db.watermark(Site, SiteIcon, SiteStatistic, SiteUserData) + + +def _site(name: str, domain: str, pri: int = 1, is_active: bool = True) -> Site: + """构造一条站点记录。""" + return Site(name=name, domain=domain, url=f"https://{domain}/", pri=pri, is_active=is_active) + + +# --------------------------------------------------------------------------- # +# Site +# --------------------------------------------------------------------------- # + +def test_site_get_by_domain_matches_async_twin(db): + """ + 按域名取站点的同步、异步结果必须指向同一行。 + """ + db.add(_site("站点A", "a.test"), _site("站点B", "b.test")) + + assert Site.get_by_domain(db.session, "a.test").name == "站点A" + assert asyncio.run(Site.async_get_by_domain(domain="a.test")).name == "站点A" + assert asyncio.run(Site.async_get_by_name(name="站点B")).domain == "b.test" + + +def test_site_get_by_domain_returns_none_when_absent(db): + """ + 域名不存在时返回 None,调用方据此判断站点是否已配置。 + """ + assert Site.get_by_domain(db.session, "missing.test") is None + + +def test_site_get_actives_excludes_disabled_sites(db): + """ + 取启用站点必须排除已停用的。 + + 停用站点仍被返回意味着它照样会被搜索和刷流访问,等于停用开关没生效。 + """ + db.add(_site("启用1", "on1.test"), _site("启用2", "on2.test"), + _site("停用", "off.test", is_active=False)) + + assert {s.domain for s in Site.get_actives(db.session)} == {"on1.test", "on2.test"} + assert {s.domain for s in asyncio.run(Site.async_get_actives())} == {"on1.test", "on2.test"} + + +def test_site_list_order_by_pri_is_ascending(db): + """ + 站点列表必须按优先级升序——顺序决定搜索与下载的站点先后。 + """ + db.add(_site("三", "p3.test", pri=3), _site("一", "p1.test", pri=1), + _site("二", "p2.test", pri=2)) + + assert [s.domain for s in Site.list_order_by_pri(db.session)] == \ + ["p1.test", "p2.test", "p3.test"] + assert [s.domain for s in asyncio.run(Site.async_list_order_by_pri())] == \ + ["p1.test", "p2.test", "p3.test"] + + +def test_site_get_domains_by_ids_returns_plain_strings(db): + """ + 按 ID 批量取域名必须返回纯字符串列表,且只含请求的那些 ID。 + """ + first = db.add(_site("一", "d1.test")) + second = db.add(_site("二", "d2.test")) + db.add(_site("三", "d3.test")) + + domains = Site.get_domains_by_ids(db.session, [first.id, second.id]) + + assert sorted(domains) == ["d1.test", "d2.test"] + assert all(isinstance(item, str) for item in domains) + + +def test_site_get_domains_by_ids_with_empty_list(db): + """ + ID 列表为空时返回空列表,不能退化成返回全部域名。 + """ + db.add(_site("一", "e1.test")) + + assert Site.get_domains_by_ids(db.session, []) == [] + + +def test_site_reset_empties_the_table(db): + """ + 重置会清空站点表——CookieCloud 全量同步依赖它先清场再写入。 + """ + db.add(_site("一", "r1.test")) + + Site.reset(db.session) + + assert Site.list_order_by_pri(db.session) == [] + + +# --------------------------------------------------------------------------- # +# SiteIcon / SiteStatistic +# --------------------------------------------------------------------------- # + +def test_siteicon_get_by_domain_matches_async_twin(db): + """ + 图标按域名查找的同步、异步结果必须一致。 + """ + db.add(SiteIcon(name="站点A", domain="icon-a.test", url="https://icon-a.test/f.ico"), + SiteIcon(name="站点B", domain="icon-b.test", url="https://icon-b.test/f.ico")) + + assert SiteIcon.get_by_domain(db.session, "icon-a.test").name == "站点A" + assert asyncio.run(SiteIcon.async_get_by_domain(domain="icon-a.test")).name == "站点A" + assert SiteIcon.get_by_domain(db.session, "icon-missing.test") is None + + +def test_sitestatistic_get_by_domain_matches_async_twin(db): + """ + 访问统计按域名查找的同步、异步结果必须一致。 + """ + db.add(SiteStatistic(domain="stat-a.test", success=3, fail=1, seconds=2, lst_state=0), + SiteStatistic(domain="stat-b.test", success=1, fail=0, seconds=1, lst_state=0)) + + assert SiteStatistic.get_by_domain(db.session, "stat-a.test").success == 3 + assert asyncio.run(SiteStatistic.async_get_by_domain(domain="stat-a.test")).success == 3 + assert SiteStatistic.get_by_domain(db.session, "stat-missing.test") is None + + +def test_sitestatistic_reset_empties_the_table(db): + """ + 重置统计会清空整表,供「重置站点数据」入口使用。 + """ + db.add(SiteStatistic(domain="stat-reset.test", success=1, fail=0, seconds=1, lst_state=0)) + + SiteStatistic.reset(db.session) + + assert SiteStatistic.get_by_domain(db.session, "stat-reset.test") is None + + +# --------------------------------------------------------------------------- # +# SiteUserData +# --------------------------------------------------------------------------- # + +def _userdata(domain: str, day: str, time: str, upload: float = 0, + err_msg: str = None) -> SiteUserData: + """构造一条站点用户数据快照。""" + return SiteUserData(domain=domain, name=domain, username="u", upload=upload, + updated_day=day, updated_time=time, err_msg=err_msg) + + +def test_userdata_get_by_domain_narrows_with_date_and_time(db): + """ + 按域名查询时,日期与时刻参数应逐级收窄结果范围。 + """ + db.add(_userdata("ud.test", "2026-08-11", "10:00:00"), + _userdata("ud.test", "2026-08-12", "10:00:00"), + _userdata("ud.test", "2026-08-12", "20:00:00"), + _userdata("other.test", "2026-08-12", "10:00:00")) + + assert len(SiteUserData.get_by_domain(db.session, "ud.test")) == 3 + assert len(SiteUserData.get_by_domain(db.session, "ud.test", workdate="2026-08-12")) == 2 + assert len(SiteUserData.get_by_domain(db.session, "ud.test", + workdate="2026-08-12", worktime="20:00:00")) == 1 + + +def test_userdata_get_by_domain_matches_async_twin(db): + """ + 三种收窄组合下同步与异步必须给出同样多的行。 + """ + db.add(_userdata("ud2.test", "2026-08-12", "10:00:00"), + _userdata("ud2.test", "2026-08-12", "20:00:00")) + + for kwargs in ({}, {"workdate": "2026-08-12"}, + {"workdate": "2026-08-12", "worktime": "20:00:00"}): + sync_rows = SiteUserData.get_by_domain(db.session, "ud2.test", **kwargs) + async_rows = asyncio.run(SiteUserData.async_get_by_domain(domain="ud2.test", **kwargs)) + assert len(sync_rows) == len(async_rows) + + +def test_userdata_get_by_date_returns_all_domains_of_that_day(db): + """ + 按日期查询应跨站点返回当天全部快照。 + """ + db.add(_userdata("day-a.test", "2026-08-12", "10:00:00"), + _userdata("day-b.test", "2026-08-12", "10:00:00"), + _userdata("day-a.test", "2026-08-11", "10:00:00")) + + rows = SiteUserData.get_by_date(db.session, "2026-08-12") + + assert {r.domain for r in rows} == {"day-a.test", "day-b.test"} + + +def test_userdata_get_latest_returns_one_day_per_domain(db): + """ + 每个站点只返回其最新一天的快照,且跨站点互不影响。 + + 这条正是子查询存在的理由:退化成「全表取最新」时,只会剩下日期最大的那个站点。 + """ + db.add(_userdata("late-a.test", "2026-08-10", "10:00:00", upload=1), + _userdata("late-a.test", "2026-08-12", "10:00:00", upload=2), + _userdata("late-b.test", "2026-08-11", "10:00:00", upload=3)) + + latest = {r.domain: r for r in SiteUserData.get_latest(db.session) + if r.domain in ("late-a.test", "late-b.test")} + + assert set(latest) == {"late-a.test", "late-b.test"} + assert latest["late-a.test"].updated_day == "2026-08-12" + assert latest["late-b.test"].updated_day == "2026-08-11" + + +def test_userdata_get_latest_ignores_failed_snapshots_when_picking_the_day(db): + """ + 带错误信息的快照不参与「最新一天」的判定。 + + 抓取失败当天也会留一条记录,若它决定了最新日期,站点数据会显示成空。 + """ + db.add(_userdata("err.test", "2026-08-10", "10:00:00", upload=5), + _userdata("err.test", "2026-08-12", "10:00:00", err_msg="登录失败")) + + rows = [r for r in SiteUserData.get_latest(db.session) if r.domain == "err.test"] + + assert [r.updated_day for r in rows] == ["2026-08-10"] + + +def test_userdata_get_latest_matches_async_twin(db): + """ + 同步与异步的「最新一天」必须选出同一批行。 + """ + db.add(_userdata("par.test", "2026-08-10", "10:00:00"), + _userdata("par.test", "2026-08-12", "10:00:00")) + + sync_rows = [(r.domain, r.updated_day) for r in SiteUserData.get_latest(db.session)] + async_rows = [(r.domain, r.updated_day) for r in asyncio.run(SiteUserData.async_get_latest())] + + assert sorted(sync_rows) == sorted(async_rows) + + +def test_userdata_delete_before_is_batched_and_bounded(db): + """ + 清理旧快照必须分批并遵守上限,且不碰保留期内的数据。 + + 一次性删除大表会长时间持锁,SQLite 下直接表现为整个应用卡住。 + """ + for index in range(5): + db.add(_userdata("old.test", "2026-01-0%d" % (index + 1), "10:00:00")) + db.add(_userdata("old.test", "2026-08-12", "10:00:00")) + + assert SiteUserData.delete_before(db.session, before_day="2026-08-01", limit=2) == 2 + assert SiteUserData.delete_before(db.session, before_day="2026-08-01", limit=100) == 3 + assert SiteUserData.delete_before(db.session, before_day="2026-08-01", limit=100) == 0 + + remaining = SiteUserData.get_by_domain(db.session, "old.test") + assert [r.updated_day for r in remaining] == ["2026-08-12"] + + +def test_userdata_delete_before_keeps_the_row_exactly_at_the_boundary(db): + """ + 保留日期当天的快照属于「保留期内」,不能被清理(``updated_day < before_day``)。 + + 上面那条用例的数据离水位有半年之遥,``<`` 写成 ``<=`` 也照样绿; + 这里把行压在水位当天,让开闭区间之差可观测——差一天就是少一天的站点数据曲线。 + """ + boundary = "2026-05-01" + db.add(_userdata("boundary.test", boundary, "10:00:00"), + _userdata("boundary.test", "2026-04-30", "10:00:00")) + + assert SiteUserData.delete_before(db.session, before_day=boundary, limit=100) == 1 + + remaining = SiteUserData.get_by_domain(db.session, "boundary.test") + assert [r.updated_day for r in remaining] == [boundary] diff --git a/tests/test_db_subscribe_queries.py b/tests/test_db_subscribe_queries.py new file mode 100644 index 000000000..bc7a68fd2 --- /dev/null +++ b/tests/test_db_subscribe_queries.py @@ -0,0 +1,337 @@ +""" +订阅表与订阅历史表的查询行为。 + +订阅身份由「来源 + 原生 ID + 季 + 剧集组 + 音乐实体」五项共同确定,任意一项在查询里 +丢失都会造成误判:判为已存在则新订阅被拒绝,判为不存在则同一部剧被重复订阅。 +这些都不会抛异常,只能靠对真实数据的断言暴露。 +""" +import asyncio +import time as _time + +import pytest + +from app.db.models import subscribe as subscribe_module +from app.db.models.subscribe import Subscribe +from app.db.models.subscribehistory import SubscribeHistory +from app.schemas.types import MediaSource, MediaType + +TMDB = str(MediaSource.TMDB) + + +@pytest.fixture(autouse=True) +def _track(db): + """把订阅与订阅历史表纳入用例级回收。""" + db.watermark(Subscribe, SubscribeHistory) + + +def _sub(name: str, media_id: str = "9001", season: int = 1, episode_group: str = None, + state: str = "N", username: str = "alice", mtype: str = None, + music_type: str = None, date: str = "2026-08-13 10:00:00") -> Subscribe: + """构造一条订阅记录。""" + return Subscribe(name=name, type=mtype or MediaType.TV.value, state=state, + media_source=TMDB, media_id=media_id, season=season, + episode_group=episode_group, username=username, + music_type=music_type, date=date) + + +# --------------------------------------------------------------------------- # +# Subscribe:身份查询 +# --------------------------------------------------------------------------- # + +def test_exists_distinguishes_season_and_episode_group(db): + """ + 同一媒体的不同季、不同剧集组各自是独立订阅身份。 + + 剧集组条件丢失时,主季订阅会命中自定义剧集组的订阅,用户再也加不上第二个组。 + """ + db.add(_sub("主季", season=1, episode_group=None), + _sub("剧集组", season=1, episode_group="eg-1"), + _sub("第二季", season=2, episode_group=None)) + + assert Subscribe.exists(db.session, MediaSource.TMDB, "9001", season=1, + episode_group=None).name == "主季" + assert Subscribe.exists(db.session, MediaSource.TMDB, "9001", season=1, + episode_group="eg-1").name == "剧集组" + assert Subscribe.exists(db.session, MediaSource.TMDB, "9001", season=2, + episode_group=None).name == "第二季" + assert Subscribe.exists(db.session, MediaSource.TMDB, "9001", season=3, + episode_group=None) is None + + +def test_exists_matches_async_twin(db): + """ + 同步与异步的身份判定必须一致,否则 API 与调度任务对「是否已订阅」意见相左。 + """ + db.add(_sub("并行", season=1)) + + sync_found = Subscribe.exists(db.session, MediaSource.TMDB, "9001", season=1) + async_found = asyncio.run(Subscribe.async_exists( + media_source=MediaSource.TMDB, media_id="9001", season=1)) + + assert sync_found.id == async_found.id + + +@pytest.mark.parametrize("media_id", [None, "", " "]) +def test_exists_rejects_blank_media_id(db, media_id): + """ + 媒体 ID 为空时直接返回 None。 + + 否则条件退化,任意一条订阅都会被当成命中,新订阅全部被拒。 + """ + db.add(_sub("有订阅")) + + assert Subscribe.exists(db.session, MediaSource.TMDB, media_id, season=1) is None + + +def test_exists_treats_recording_as_matching_null_music_type(db): + """ + 单曲订阅要兼容历史上未写 music_type 的行。 + + 老数据的 music_type 为空,若严格相等匹配会被判为不存在,用户会重复订阅同一首歌。 + """ + db.add(_sub("老单曲", media_id="mb-1", season=None, music_type=None, + mtype=MediaType.MUSIC.value)) + + found = Subscribe.exists(db.session, MediaSource.TMDB, "mb-1", music_type="recording") + + assert found is not None and found.name == "老单曲" + + +def test_exists_by_username_scopes_to_owner(db): + """ + 按 owner 查询必须限定用户名,且用户名为空时直接返回 None。 + """ + db.add(_sub("alice 的", username="alice"), _sub("bob 的", username="bob", media_id="9002")) + + assert Subscribe.exists_by_username(db.session, "alice", MediaSource.TMDB, + "9001", season=1).name == "alice 的" + assert Subscribe.exists_by_username(db.session, "bob", MediaSource.TMDB, + "9001", season=1) is None + assert Subscribe.exists_by_username(db.session, "", MediaSource.TMDB, + "9001", season=1) is None + + +def test_get_by_narrows_with_type_and_optional_season(db): + """ + 按类型查询时类型必须参与匹配,季号可选但给出即须生效。 + """ + db.add(_sub("剧集", mtype=MediaType.TV.value, season=1), + _sub("电影", mtype=MediaType.MOVIE.value, season=1, media_id="9003")) + + assert Subscribe.get_by(db.session, MediaType.TV.value, MediaSource.TMDB, + "9001").name == "剧集" + assert Subscribe.get_by(db.session, MediaType.MOVIE.value, MediaSource.TMDB, + "9001") is None + assert Subscribe.get_by(db.session, MediaType.TV.value, MediaSource.TMDB, + "9001", season=2) is None + + +def test_list_by_media_identity_returns_all_seasons(db): + """ + 按媒体身份列举会跨季返回全部订阅,空身份则短路成空列表。 + """ + db.add(_sub("第一季", season=1), _sub("第二季", season=2), + _sub("别的剧", media_id="9009")) + + listed = Subscribe.list_by_media_identity(db.session, MediaSource.TMDB, "9001") + + assert {s.season for s in listed} == {1, 2} + assert Subscribe.list_by_media_identity(db.session, MediaSource.TMDB, "") == [] + + +# --------------------------------------------------------------------------- # +# Subscribe:列表查询 +# --------------------------------------------------------------------------- # + +def test_get_by_state_splits_comma_separated_states(db): + """ + 状态支持逗号分隔的多值,为空时返回全部。 + + 订阅刷新按状态取任务,多值解析失效会让一部分订阅永远不被处理。 + """ + db.add(_sub("待订阅", state="N"), _sub("订阅中", state="R", media_id="9004"), + _sub("已完成", state="P", media_id="9005")) + + states = {s.state for s in Subscribe.get_by_state(db.session, "N,R")} + assert states == {"N", "R"} + + assert len(Subscribe.get_by_state(db.session, "")) >= 3 + assert {s.state for s in asyncio.run(Subscribe.async_get_by_state(state="N,R"))} == {"N", "R"} + + +def test_get_by_title_optionally_narrows_by_season(db): + """ + 按标题查询时季号可选,给出即须生效。 + """ + db.add(_sub("同名剧", season=1), _sub("同名剧", season=2)) + + assert Subscribe.get_by_title(db.session, "同名剧", season=2).season == 2 + assert Subscribe.get_by_title(db.session, "同名剧") is not None + assert Subscribe.get_by_title(db.session, "不存在的剧") is None + + +@pytest.mark.parametrize("state,mtype,expected", [ + (None, None, {"剧-N", "剧-R", "影-N"}), + ("N", None, {"剧-N", "影-N"}), + (None, MediaType.TV.value, {"剧-N", "剧-R"}), + ("N", MediaType.TV.value, {"剧-N"}), +]) +def test_list_by_username_covers_all_filter_combinations(db, state, mtype, expected): + """ + 按 owner 列举的四种「状态 × 类型」组合都必须正确收窄。 + + 这四条分支是「我的订阅」页面的全部筛选路径,任何一条串了都会展示别人的订阅 + 或漏掉自己的。 + """ + db.add(_sub("剧-N", state="N", mtype=MediaType.TV.value, media_id="9101"), + _sub("剧-R", state="R", mtype=MediaType.TV.value, media_id="9102"), + _sub("影-N", state="N", mtype=MediaType.MOVIE.value, media_id="9103"), + _sub("别人的", state="N", mtype=MediaType.TV.value, media_id="9104", + username="bob")) + + listed = Subscribe.list_by_username(db.session, "alice", state=state, mtype=mtype) + + assert {s.name for s in listed} == expected + + +def test_list_by_username_matches_async_twin(db): + """ + 四种筛选组合下同步与异步必须返回同一批订阅。 + """ + db.add(_sub("剧-N", state="N", mtype=MediaType.TV.value, media_id="9201"), + _sub("影-R", state="R", mtype=MediaType.MOVIE.value, media_id="9202")) + + for state, mtype in ((None, None), ("N", None), (None, MediaType.TV.value), + ("N", MediaType.TV.value)): + sync_names = sorted(s.name for s in + Subscribe.list_by_username(db.session, "alice", state, mtype)) + async_names = sorted(s.name for s in asyncio.run( + Subscribe.async_list_by_username(username="alice", state=state, mtype=mtype))) + assert sync_names == async_names + + +def test_list_by_type_only_returns_recent_days(db): + """ + 按类型取最近 N 天的订阅,超出窗口的不返回。 + + 时间窗口失效会让「最近订阅」把历史全量拉出来,首页直接卡死。 + """ + db.add(_sub("最近", mtype=MediaType.TV.value, media_id="9301", + date="2099-01-01 00:00:00"), + _sub("很久以前", mtype=MediaType.TV.value, media_id="9302", + date="2000-01-01 00:00:00")) + + names = {s.name for s in Subscribe.list_by_type(db.session, MediaType.TV.value, days=7)} + + assert "最近" in names + assert "很久以前" not in names + + +def test_list_by_type_includes_the_window_start_boundary(db, frozen_now): + """ + 时间窗是闭区间起点(``date >= 起点``),正好落在起点的订阅必须在结果里,同步异步一致。 + + 起点由「调用时刻 - N 天」现算,不冻结时钟就摆不到边界上;上面那条用例用的是 + 2099/2000 两个极端值,比较符改成 ``>`` 照样绿。 + """ + now = frozen_now(subscribe_module) + window_start = _time.strftime("%Y-%m-%d %H:%M:%S", _time.localtime(now - 86400 * 7)) + one_second_earlier = _time.strftime("%Y-%m-%d %H:%M:%S", + _time.localtime(now - 86400 * 7 - 1)) + db.add(_sub("窗口起点上", mtype=MediaType.TV.value, media_id="9303", date=window_start), + _sub("窗口起点前一秒", mtype=MediaType.TV.value, media_id="9304", + date=one_second_earlier)) + + names = {s.name for s in Subscribe.list_by_type(db.session, MediaType.TV.value, days=7)} + async_names = {s.name for s in asyncio.run( + Subscribe.async_list_by_type(mtype=MediaType.TV.value, days=7))} + + assert "窗口起点上" in names and "窗口起点前一秒" not in names + assert "窗口起点上" in async_names and "窗口起点前一秒" not in async_names + + +def test_delete_by_media_identity_removes_matching_seasons_only(db): + """ + 按媒体身份删除时,给出季号只删该季,不给则删全部季。 + """ + db.add(_sub("第一季", season=1), _sub("第二季", season=2)) + + Subscribe().delete_by_media_identity(db.session, TMDB, "9001", season=1) + + remaining = Subscribe.list_by_media_identity(db.session, MediaSource.TMDB, "9001") + assert [s.season for s in remaining] == [2] + + Subscribe().delete_by_media_identity(db.session, TMDB, "9001") + assert Subscribe.list_by_media_identity(db.session, MediaSource.TMDB, "9001") == [] + + +# --------------------------------------------------------------------------- # +# SubscribeHistory +# --------------------------------------------------------------------------- # + +def _history(name: str, mtype: str = MediaType.TV.value, media_id: str = "8001", + season: int = 1, episode_group: str = None, + date: str = "2026-08-13 10:00:00", username: str = "alice") -> SubscribeHistory: + """构造一条订阅历史记录。""" + return SubscribeHistory(name=name, type=mtype, media_source=TMDB, media_id=media_id, + season=season, episode_group=episode_group, date=date, + username=username) + + +def test_history_list_by_type_is_newest_first_and_paged(db): + """ + 历史按完成时间倒序分页,且只返回指定类型。 + """ + db.add(_history("旧", date="2026-08-01 10:00:00", media_id="8101"), + _history("新", date="2026-08-12 10:00:00", media_id="8102"), + _history("电影", mtype=MediaType.MOVIE.value, media_id="8103")) + + page = SubscribeHistory.list_by_type(db.session, MediaType.TV.value, page=1, count=10) + + assert [h.name for h in page] == ["新", "旧"] + assert [h.name for h in SubscribeHistory.list_by_type( + db.session, MediaType.TV.value, page=1, count=1)] == ["新"] + + +def test_history_list_by_type_matches_async_twin(db): + """ + 同步与异步的历史分页必须返回同一批记录。 + """ + db.add(_history("A", date="2026-08-12 10:00:00", media_id="8201"), + _history("B", date="2026-08-11 10:00:00", media_id="8202")) + + sync_names = [h.name for h in SubscribeHistory.list_by_type( + db.session, MediaType.TV.value, page=1, count=10)] + async_names = [h.name for h in asyncio.run(SubscribeHistory.async_list_by_type( + mtype=MediaType.TV.value, page=1, count=10))] + + assert sync_names == async_names + + +def test_history_exists_distinguishes_episode_group(db): + """ + 历史的存在性判定与订阅同规则:剧集组不同即为不同身份。 + + 判错会让已完成的主季订阅挡住自定义剧集组的新订阅。 + """ + db.add(_history("主季历史", season=1, episode_group=None, media_id="8301"), + _history("剧集组历史", season=1, episode_group="eg-1", media_id="8301")) + + assert SubscribeHistory.exists(db.session, MediaSource.TMDB, "8301", season=1, + episode_group=None).name == "主季历史" + assert SubscribeHistory.exists(db.session, MediaSource.TMDB, "8301", season=1, + episode_group="eg-1").name == "剧集组历史" + assert SubscribeHistory.exists(db.session, MediaSource.TMDB, "", season=1) is None + + +def test_history_exists_matches_async_twin(db): + """ + 历史存在性判定的同步与异步结果必须一致。 + """ + db.add(_history("并行历史", season=1, media_id="8401")) + + sync_found = SubscribeHistory.exists(db.session, MediaSource.TMDB, "8401", season=1) + async_found = asyncio.run(SubscribeHistory.async_exists( + media_source=MediaSource.TMDB, media_id="8401", season=1)) + + assert sync_found.id == async_found.id diff --git a/tests/test_db_transferhistory_queries.py b/tests/test_db_transferhistory_queries.py new file mode 100644 index 000000000..d9c5a3436 --- /dev/null +++ b/tests/test_db_transferhistory_queries.py @@ -0,0 +1,519 @@ +""" +整理历史表的查询行为。 + +整理历史同时承担三个职责:查重(这个文件整理过没有)、溯源(这个媒体库文件是从哪 +来的)、统计。查重误判会重复整理或永久漏件——挂载故障那一类问题最终就落在这张表上; +溯源查错会让「重新整理」把不相干的文件搬走。 +""" +import asyncio +import time as _time + +import pytest + +from app.db.models import transferhistory as transferhistory_module +from app.db.models.transferhistory import TransferHistory +from app.schemas.types import MediaSource, MediaType + +TMDB = str(MediaSource.TMDB) + + +@pytest.fixture(autouse=True) +def _track(db): + """把整理历史表纳入用例级回收。""" + db.watermark(TransferHistory) + + +def _hist(title: str = "片名", src: str = None, dest: str = None, + src_storage: str = "local", dest_storage: str = "local", + mode: str = "move", status: bool = True, mtype: str = None, + year: str = "2026", media_id: str = "6001", seasons: str = None, + episodes: str = None, date: str = "2026-08-13 10:00:00", + download_hash: str = None) -> TransferHistory: + """构造一条整理历史。""" + return TransferHistory(src=src or f"/downloads/{title}.mkv", src_storage=src_storage, + dest=dest or f"/media/{title}.mkv", dest_storage=dest_storage, + mode=mode, type=mtype or MediaType.TV.value, title=title, + year=year, media_source=TMDB, media_id=media_id, + seasons=seasons, episodes=episodes, status=status, + date=date, download_hash=download_hash, files=[]) + + +# --------------------------------------------------------------------------- # +# 按路径溯源 +# --------------------------------------------------------------------------- # + +def test_get_by_src_scopes_by_storage(db): + """ + 按源路径查询可限定存储,不限定时取主键最大的那条。 + + 同一路径在不同存储下是不同文件;不限定存储会把本地盘的记录当成网盘的。 + 表上有 (src, src_storage) 唯一约束,所以「最新一条」只在跨存储时才有分歧。 + """ + db.add(_hist("本地", src="/data/same.mkv", src_storage="local")) + latest = db.add(_hist("网盘", src="/data/same.mkv", src_storage="alist")) + + assert TransferHistory.get_by_src(db.session, "/data/same.mkv", + storage="alist").title == "网盘" + assert TransferHistory.get_by_src(db.session, "/data/same.mkv", + storage="local").title == "本地" + assert TransferHistory.get_by_src(db.session, "/data/same.mkv").id == latest.id + assert TransferHistory.get_by_src(db.session, "/data/none.mkv") is None + + +def test_get_success_by_src_matches_path_verbatim(db): + """ + 成功记录按源路径原样精确匹配,不做归一化。 + + 蓝光原盘目录的记录带尾斜杠,归一化后反而匹配不到自己。 + """ + db.add(_hist("原盘", src="/data/BDMV/", status=True), + _hist("失败的", src="/data/fail.mkv", status=False)) + + assert TransferHistory.get_success_by_src(db.session, "/data/BDMV/").title == "原盘" + assert TransferHistory.get_success_by_src(db.session, "/data/BDMV") is None + assert TransferHistory.get_success_by_src(db.session, "/data/fail.mkv") is None + + +def test_get_success_by_src_can_scope_by_storage(db): + """ + 成功记录查询同样支持按源存储收窄。 + """ + db.add(_hist("本地", src="/data/x.mkv", src_storage="local"), + _hist("网盘", src="/data/x.mkv", src_storage="alist")) + + assert TransferHistory.get_success_by_src(db.session, "/data/x.mkv", + storage="alist").title == "网盘" + + +def test_get_by_dest_scopes_by_storage_and_takes_newest(db): + """ + 按目标路径查询用于从媒体库反查来源,可限定目标存储并取最新一条。 + """ + db.add(_hist("旧", src="/data/d1.mkv", dest="/media/same.mkv", dest_storage="local")) + newest = db.add(_hist("新", src="/data/d2.mkv", dest="/media/same.mkv", + dest_storage="local")) + db.add(_hist("别的存储", src="/data/d3.mkv", dest="/media/same.mkv", + dest_storage="alist")) + + assert TransferHistory.get_by_dest(db.session, "/media/same.mkv", + storage="local").id == newest.id + assert TransferHistory.get_by_dest(db.session, "/media/same.mkv", + storage="alist").title == "别的存储" + assert TransferHistory.get_by_dest(db.session, "/media/none.mkv") is None + + +def test_list_success_by_src_normalizes_the_path(db): + """ + 列举成功记录时对入参路径做归一化(反斜杠、尾斜杠)。 + + Windows 侧传来的路径带反斜杠,不归一化会一条都匹配不到。 + """ + db.add(_hist("目录", src="/data/dir")) + + for probe in ("/data/dir", "/data/dir/", "\\data\\dir"): + assert [h.title for h in TransferHistory.list_success_by_src(db.session, probe)] == ["目录"] + + +def test_list_success_by_src_recursive_matches_children_only(db): + """ + 递归模式匹配目录自身及其子项,但不能误伤同前缀的兄弟目录。 + + 直接用 `like("前缀%")` 会把 /data/dir2 也算成 /data/dir 的子项, + 重新整理时会把邻居目录一起搬走。 + """ + db.add(_hist("自身", src="/data/dir"), + _hist("子项", src="/data/dir/a.mkv"), + _hist("兄弟", src="/data/dir2/b.mkv")) + + titles = {h.title for h in TransferHistory.list_success_by_src( + db.session, "/data/dir", recursive=True)} + + assert titles == {"自身", "子项"} + + +def test_list_success_by_src_escapes_like_wildcards(db): + """ + 路径中的 % 与 _ 必须被转义,否则会被当成通配符匹配到别的目录。 + """ + db.add(_hist("含下划线", src="/data/a_b/x.mkv"), + _hist("被误伤", src="/data/axb/y.mkv")) + + titles = {h.title for h in TransferHistory.list_success_by_src( + db.session, "/data/a_b", recursive=True)} + + assert titles == {"含下划线"} + + +def test_list_success_move_by_dest_only_returns_move_mode(db): + """ + 从媒体库现址发起重新整理时只认「移动」模式的记录。 + + 硬链接/复制的源文件仍在原处,把它们算作可回溯的移动记录会导致重复整理。 + """ + db.add(_hist("移动", src="/data/m.mkv", dest="/media/m.mkv", mode="move"), + _hist("硬链", src="/data/l.mkv", dest="/media/l.mkv", mode="link"), + _hist("移动失败", src="/data/f.mkv", dest="/media/f.mkv", mode="move", + status=False)) + + titles = {h.title for h in TransferHistory.list_success_move_by_dest(db.session, "/media/m.mkv")} + assert titles == {"移动"} + assert TransferHistory.list_success_move_by_dest(db.session, "/media/l.mkv") == [] + assert TransferHistory.list_success_move_by_dest(db.session, "/media/f.mkv") == [] + + +def test_list_success_move_by_dest_recursive_matches_children_only(db): + """ + 目标侧的递归匹配同样不能误伤同前缀的兄弟目录。 + """ + db.add(_hist("自身", src="/data/w1.mkv", dest="/media/show", mode="move"), + _hist("子项", src="/data/w2.mkv", dest="/media/show/s01.mkv", mode="move"), + _hist("兄弟", src="/data/w3.mkv", dest="/media/show2/s01.mkv", mode="move")) + + titles = {h.title for h in TransferHistory.list_success_move_by_dest( + db.session, "/media/show", recursive=True)} + + assert titles == {"自身", "子项"} + + +def test_replace_by_src_keeps_one_record_per_source(db): + """ + 同一源路径在同一存储中只保留最新一条记录,且不波及其他存储。 + + 先删后插是在同一事务内完成的:旧的「查一条再删一条」在遗留重复数据下会留下 + 脏记录,查重随后就会命中旧行。 + """ + db.add(_hist("旧一", src="/data/dup.mkv", src_storage="local"), + _hist("别的存储", src="/data/dup.mkv", src_storage="alist")) + + TransferHistory.replace_by_src(db.session, src="/data/dup.mkv", src_storage="local", + dest="/media/new.mkv", type=MediaType.TV.value, + title="新记录", status=True, date="2026-08-13 12:00:00") + + local_rows = TransferHistory.list_success_by_src(db.session, "/data/dup.mkv", + storage="local") + assert [h.title for h in local_rows] == ["新记录"] + assert TransferHistory.get_by_src(db.session, "/data/dup.mkv", + storage="alist").title == "别的存储" + + +def test_replace_by_src_defaults_storage_to_local(db): + """ + 未指定源存储时按 local 处理——历史数据没有这一列,缺省值必须与写入侧一致。 + """ + created = TransferHistory.replace_by_src(db.session, src="/data/nostorage.mkv", + type=MediaType.TV.value, title="无存储", + status=True, date="2026-08-13 12:00:00") + + assert created.src_storage == "local" + + +# --------------------------------------------------------------------------- # +# 按 hash / 身份查询 +# --------------------------------------------------------------------------- # + +def test_hash_lookups_return_single_and_all(db): + """ + 按下载 hash 既可取单条也可取全部,未命中时分别是 None 与空列表。 + """ + db.add(_hist("文件一", download_hash="th-1", src="/data/1.mkv"), + _hist("文件二", download_hash="th-1", src="/data/2.mkv")) + + assert TransferHistory.get_by_hash(db.session, "th-1") is not None + assert len(TransferHistory.list_by_hash(db.session, "th-1")) == 2 + assert TransferHistory.get_by_hash(db.session, "th-none") is None + assert TransferHistory.list_by_hash(db.session, "th-none") == [] + + +def test_get_by_media_identity_requires_matching_type(db): + """ + 按媒体身份查询时类型必须参与匹配,否则电影会命中同 ID 的剧集记录。 + """ + db.add(_hist("剧集", media_id="6100", mtype=MediaType.TV.value)) + + assert TransferHistory.get_by_media_identity( + db.session, MediaSource.TMDB, "6100", MediaType.TV.value) is not None + assert TransferHistory.get_by_media_identity( + db.session, MediaSource.TMDB, "6100", MediaType.MOVIE.value) is None + + +def test_update_download_hash_writes_only_the_target_row(db): + """ + 补写下载 hash 只影响指定的那一条记录。 + """ + target = db.add(_hist("目标", src="/data/t.mkv")) + db.add(_hist("其他", src="/data/o.mkv")) + + TransferHistory.update_download_hash(db.session, historyid=target.id, + download_hash="new-hash") + + assert TransferHistory.get_by_src(db.session, "/data/t.mkv").download_hash == "new-hash" + assert TransferHistory.get_by_src(db.session, "/data/o.mkv").download_hash is None + + +@pytest.mark.parametrize("season,episode,dest,expected", [ + ("S01", "E01", "/media/s1e1.mkv", {"季集"}), + ("S01", None, None, {"季集", "整季"}), + (None, None, None, {"季集", "整季", "另一季"}), +]) +def test_list_by_media_identity_narrows_by_season_episode_and_dest( + db, season, episode, dest, expected): + """ + 按媒体身份查询时季、集、目标路径逐级收窄。 + + 收窄失效会让「这一集整理过没有」误判成整季都整理过,剩余剧集被永久跳过。 + """ + db.add(_hist("季集", src="/data/e1.mkv", seasons="S01", episodes="E01", + dest="/media/s1e1.mkv"), + _hist("整季", src="/data/s1.mkv", seasons="S01", episodes=None, + dest="/media/s1.mkv"), + _hist("另一季", src="/data/s2.mkv", seasons="S02", episodes=None, + dest="/media/s2.mkv")) + + got = TransferHistory.list_by(db.session, mtype=MediaType.TV.value, + media_source=MediaSource.TMDB, media_id="6001", + season=season, episode=episode, dest=dest) + + assert {h.title for h in got} == expected + + +def test_list_by_media_identity_uses_dest_for_movies(db): + """ + 电影没有季集,靠目标路径区分不同版本;给出目标路径即须生效。 + """ + db.add(_hist("4K 版", src="/data/m4k.mkv", mtype=MediaType.MOVIE.value, + media_id="6200", dest="/media/movie-4k.mkv"), + _hist("1080 版", src="/data/m1080.mkv", mtype=MediaType.MOVIE.value, + media_id="6200", dest="/media/movie-1080.mkv")) + + got = TransferHistory.list_by(db.session, mtype=MediaType.MOVIE.value, + media_source=MediaSource.TMDB, media_id="6200", + dest="/media/movie-4k.mkv") + + assert [h.title for h in got] == ["4K 版"] + + +def test_list_by_falls_back_to_title_and_year(db): + """ + 没有媒体身份时退回「标题 + 年份」,服务于识别失败的历史数据。 + """ + db.add(_hist("回退标题", src="/data/fb.mkv", media_id="6300", year="2020", + seasons="S01")) + + assert [h.title for h in TransferHistory.list_by( + db.session, title="回退标题", year="2020")] == ["回退标题"] + assert [h.title for h in TransferHistory.list_by( + db.session, title="回退标题", year="2020", season="S01")] == ["回退标题"] + assert TransferHistory.list_by(db.session, title="回退标题", year="1999") == [] + + +def test_list_by_supports_type_season_and_dest_prefix(db): + """ + 媒体服务器 webhook 缺少远端身份时,按「类型 + 季 + 目标路径前缀」查询。 + + 这是该场景下唯一能定位记录的路径,丢了会让 webhook 触发的刮削全部落空。 + """ + db.add(_hist("剧集一", src="/data/wh1.mkv", mtype=MediaType.TV.value, seasons="S01", + dest="/media/Show/Season 01/e01.mkv"), + _hist("别的剧", src="/data/wh2.mkv", mtype=MediaType.TV.value, seasons="S01", + dest="/media/Other/Season 01/e01.mkv")) + + got = TransferHistory.list_by(db.session, mtype=MediaType.TV.value, season="S01", + dest="/media/Show/") + + assert [h.title for h in got] == ["剧集一"] + + +def test_list_by_without_usable_criteria_returns_empty(db): + """ + 条件不足以定位时返回空列表,不能退化成返回全表。 + """ + db.add(_hist("任意")) + + assert TransferHistory.list_by(db.session) == [] + assert TransferHistory.list_by(db.session, title="只有标题") == [] + + +# --------------------------------------------------------------------------- # +# 列表、计数与统计 +# --------------------------------------------------------------------------- # + +def test_list_by_page_filters_status_and_supports_unbounded_count(db): + """ + 分页可按成功状态收窄;count 为负表示不分页返回全部。 + + 负数这条约定是「导出全部历史」依赖的,当成普通 limit 处理会返回空。 + """ + db.add(_hist("成功一", src="/data/ok1.mkv", status=True, date="2026-08-13 10:00:01"), + _hist("成功二", src="/data/ok2.mkv", status=True, date="2026-08-13 10:00:02"), + _hist("失败", src="/data/ng.mkv", status=False, date="2026-08-13 10:00:03")) + + assert [h.title for h in TransferHistory.list_by_page(db.session, page=1, count=1, + status=True)] == ["成功二"] + assert {h.title for h in TransferHistory.list_by_page(db.session, count=-1, + status=False)} == {"失败"} + assert len(TransferHistory.list_by_page(db.session, count=-1)) >= 3 + + +def test_list_by_title_searches_title_source_and_destination(db): + """ + 标题检索同时匹配标题、源路径与目标路径,且大小写不敏感。 + + 只匹配标题会让用户按文件名搜不到任何记录。 + """ + db.add(_hist("Alpha", src="/downloads/zzz.mkv", dest="/media/zzz.mkv"), + _hist("Beta", src="/downloads/AlphaFile.mkv", dest="/media/beta.mkv"), + _hist("Gamma", src="/downloads/g.mkv", dest="/media/alpha-dir/g.mkv")) + + titles = {h.title for h in TransferHistory.list_by_title(db.session, "alpha", count=-1)} + + assert titles == {"Alpha", "Beta", "Gamma"} + + +def test_list_by_title_wildcard_mode_takes_the_pattern_verbatim(db): + """ + 通配模式下调用方自带 % 通配符,不再额外包裹。 + """ + db.add(_hist("PrefixMatch", src="/downloads/p.mkv", dest="/media/p.mkv"), + _hist("NoPrefix", src="/downloads/n.mkv", dest="/media/n.mkv")) + + titles = {h.title for h in TransferHistory.list_by_title( + db.session, "Prefix%", count=-1, wildcard=True)} + + assert titles == {"PrefixMatch"} + + +def test_list_by_title_can_filter_status(db): + """ + 检索结果同样支持按成功状态收窄。 + """ + db.add(_hist("SearchOk", status=True, src="/downloads/ok.mkv"), + _hist("SearchFail", status=False, src="/downloads/fail.mkv")) + + assert [h.title for h in TransferHistory.list_by_title( + db.session, "SearchFail", count=-1, status=False)] == ["SearchFail"] + + +def test_list_by_title_matches_async_twin(db): + """ + 同步与异步的标题检索必须返回同一批记录。 + """ + db.add(_hist("ParallelSearch", src="/downloads/par.mkv")) + + sync_titles = [h.title for h in TransferHistory.list_by_title( + db.session, "ParallelSearch", count=-1)] + async_titles = [h.title for h in asyncio.run(TransferHistory.async_list_by_title( + title="ParallelSearch", count=-1))] + + assert sync_titles == async_titles + + +def test_count_and_count_by_title_match_async_twins(db): + """ + 计数与带条件计数的同步、异步结果必须一致——分页总数由它决定, + 对不上就会出现「翻到最后一页是空的」。 + """ + db.add(_hist("CountMe", status=True, src="/downloads/c1.mkv", dest="/media/c1.mkv"), + _hist("CountMe", status=False, src="/downloads/c2.mkv", dest="/media/c2.mkv")) + + assert TransferHistory.count(db.session) == asyncio.run(TransferHistory.async_count()) + assert TransferHistory.count(db.session, status=True) == \ + asyncio.run(TransferHistory.async_count(status=True)) + assert TransferHistory.count_by_title(db.session, "CountMe") == 2 + assert TransferHistory.count_by_title(db.session, "CountMe", status=False) == 1 + assert TransferHistory.count_by_title(db.session, "CountMe") == \ + asyncio.run(TransferHistory.async_count_by_title(title="CountMe")) + + +def test_statistic_groups_by_day_within_the_window(db): + """ + 统计按日期分组,且只统计窗口内的记录。 + """ + today = _time.strftime("%Y-%m-%d", _time.localtime()) + db.add(_hist("今天一", src="/data/t1.mkv", date=f"{today} 10:00:00"), + _hist("今天二", src="/data/t2.mkv", date=f"{today} 11:00:00"), + _hist("很久以前", src="/data/t3.mkv", date="2000-01-01 10:00:00")) + + rows = dict(TransferHistory.statistic(db.session, days=7)) + + assert rows.get(today, 0) >= 2 + assert "2000-01-01" not in rows + + +def test_statistic_includes_the_window_start_boundary(db, frozen_now): + """ + 统计窗口是闭区间起点(``date >= 起点``),正好落在起点的记录必须计入,同步异步一致。 + + 起点由「调用时刻 - N 天」现算,不冻结时钟就摆不到边界上;上面那条用例用的是「今天」 + 与 2000 年两个极端值,起点比较符改成 ``>`` 照样绿。 + """ + now = frozen_now(transferhistory_module) + window_start = _time.strftime("%Y-%m-%d %H:%M:%S", _time.localtime(now - 86400 * 7)) + boundary_day = window_start[:10] + db.add(_hist("窗口起点上", src="/data/bstat.mkv", date=window_start)) + + rows = dict(TransferHistory.statistic(db.session, days=7)) + async_rows = dict(asyncio.run(TransferHistory.async_statistic(days=7))) + + assert rows.get(boundary_day, 0) == 1 + assert async_rows.get(boundary_day, 0) == 1 + + +def test_list_by_date_returns_newest_first(db): + """ + 按时间查询返回该时间之后的记录,按主键倒序。 + """ + db.add(_hist("旧", src="/data/od.mkv", date="2026-08-01 10:00:00"), + _hist("新", src="/data/nd.mkv", date="2026-08-12 10:00:00")) + + titles = [h.title for h in TransferHistory.list_by_date(db.session, "2026-08-05")] + + assert titles == ["新"] + + +def test_list_by_date_excludes_the_row_exactly_at_the_boundary(db): + """ + 取的是「该时刻之后」的记录(``date > date``),正好等于该时刻的那条不算在内。 + + 上面那条用例两侧数据各离边界四天与七天,比较符放宽成 ``>=`` 一样绿; + 这里把行压在边界上,让开闭区间之差可观测。 + """ + boundary = "2026-08-05 00:00:00" + db.add(_hist("边界上", src="/data/bd0.mkv", date=boundary), + _hist("边界后一秒", src="/data/bd1.mkv", date="2026-08-05 00:00:01")) + + titles = [h.title for h in TransferHistory.list_by_date(db.session, boundary)] + + assert titles == ["边界后一秒"] + + +def test_delete_before_is_batched_and_keeps_recent(db): + """ + 历史清理分批执行且不碰保留期内的记录。 + """ + for index in range(4): + db.add(_hist(f"old-{index}", src=f"/data/old{index}.mkv", + date=f"2026-01-01 10:00:0{index}")) + db.add(_hist("recent", src="/data/recent.mkv", date="2026-08-13 10:00:00")) + + assert TransferHistory.delete_before(db.session, before_time="2026-08-01", limit=2) == 2 + assert TransferHistory.delete_before(db.session, before_time="2026-08-01", limit=100) == 2 + assert TransferHistory.delete_before(db.session, before_time="2026-08-01", limit=100) == 0 + + assert TransferHistory.get_by_src(db.session, "/data/recent.mkv") is not None + + +def test_delete_before_keeps_the_row_exactly_at_the_boundary(db): + """ + 保留时间点上的整理历史属于「保留期内」,不能被清理(``date < before_time``)。 + + 整理历史被删掉就等于丢失溯源,同一文件会被重新整理一次; + 上面那条用例的数据离水位半年,``<`` 写成 ``<=`` 完全不可观测。 + """ + boundary = "2026-05-01 00:00:00" + db.add(_hist("边界上", src="/data/bdel0.mkv", date=boundary), + _hist("边界前一秒", src="/data/bdel1.mkv", date="2026-04-30 23:59:59")) + + assert TransferHistory.delete_before(db.session, before_time=boundary, limit=100) == 1 + + assert TransferHistory.get_by_src(db.session, "/data/bdel0.mkv") is not None + assert TransferHistory.get_by_src(db.session, "/data/bdel1.mkv") is None diff --git a/tests/test_db_transferpending_queries.py b/tests/test_db_transferpending_queries.py new file mode 100644 index 000000000..55828ce6e --- /dev/null +++ b/tests/test_db_transferpending_queries.py @@ -0,0 +1,178 @@ +""" +待整理登记表的查询行为。 + +这张表是「挂载挂死后重启不漏件」的唯一依据:登记去重、回放顺序、终态注销三件事 +任何一件出偏差,都直接表现为文件被漏整理或被重复整理,而不是一个可见的报错。 +因此这里对着真实数据库断言查回的内容,而不是断言调用了什么。 +""" +import pytest + +from app.db.models.transferpending import TransferPending +from app.db.oper.transferpending import TransferPendingOper + + +@pytest.fixture(autouse=True) +def _track(db): + """把待整理表纳入用例级回收。""" + db.watermark(TransferPending) + + +def test_register_is_idempotent_and_keeps_first_time(db): + """ + 同一文件重复登记只保留一条,且登记时间保持首次的值。 + + 监控在挂载抖动时会对同一个文件反复触发事件,若每次都新增一条,回放时同一个 + 文件会被送进整理链多次。保留首次时间则保证回放顺序仍是「最早发现」的顺序。 + """ + TransferPending.register(db.session, storage="local", src_path="/mnt/a.mkv", + now_time="2026-08-13 10:00:00") + TransferPending.register(db.session, storage="local", src_path="/mnt/a.mkv", + now_time="2026-08-13 12:00:00") + + rows = TransferPending.list_all(db.session) + same_path = [r for r in rows if r.src_path == "/mnt/a.mkv"] + assert len(same_path) == 1 + assert same_path[0].created_at == "2026-08-13 10:00:00" + + +def test_register_scopes_by_storage(db): + """ + 存储不同即为不同文件——路径相同但分属不同存储时不能互相去重。 + """ + TransferPending.register(db.session, storage="local", src_path="/data/x.mkv", + now_time="2026-08-13 10:00:00") + TransferPending.register(db.session, storage="alist", src_path="/data/x.mkv", + now_time="2026-08-13 10:00:01") + + rows = [r for r in TransferPending.list_all(db.session) if r.src_path == "/data/x.mkv"] + assert {r.storage for r in rows} == {"local", "alist"} + + +@pytest.mark.parametrize("storage,src_path", [("", "/mnt/a.mkv"), ("local", ""), ("", "")]) +def test_register_rejects_incomplete_identity(db, storage, src_path): + """ + 缺少存储或路径的登记必须直接丢弃,不能写入半条记录。 + + 半条记录回放时既定位不到文件、也无法被 discard 匹配,会永久留在表里。 + """ + assert TransferPending.register(db.session, storage=storage, src_path=src_path, + now_time="2026-08-13 10:00:00") is None + + +def test_list_all_replays_in_registration_order(db): + """ + 回放顺序必须是登记时间升序、同时间按主键升序。 + + 乱序回放会让后发现的文件先进整理链,与原入队顺序不一致。 + """ + for path, moment in [("/mnt/c.mkv", "2026-08-13 12:00:00"), + ("/mnt/a.mkv", "2026-08-13 10:00:00"), + ("/mnt/b.mkv", "2026-08-13 11:00:00")]: + TransferPending.register(db.session, storage="local", src_path=path, now_time=moment) + + ordered = [r.src_path for r in TransferPending.list_all(db.session) + if r.src_path.startswith("/mnt/")] + assert ordered == ["/mnt/a.mkv", "/mnt/b.mkv", "/mnt/c.mkv"] + + +def test_list_all_honours_limit(db): + """ + 回放上限必须生效——异常积压时一次性全放会把整理链直接压垮。 + """ + for index in range(5): + TransferPending.register(db.session, storage="local", src_path=f"/mnt/{index}.mkv", + now_time=f"2026-08-13 10:00:0{index}") + + assert len(TransferPending.list_all(db.session, limit=3)) == 3 + + +def test_discard_removes_only_the_matching_row(db): + """ + 注销只应删除匹配的那一条,并返回删除条数。 + + 整理到达终态时按「存储 + 路径」注销,误删其他登记等于把别的文件也判成已完成。 + """ + TransferPending.register(db.session, storage="local", src_path="/mnt/a.mkv", + now_time="2026-08-13 10:00:00") + TransferPending.register(db.session, storage="local", src_path="/mnt/b.mkv", + now_time="2026-08-13 10:00:01") + + assert TransferPending.discard(db.session, storage="local", src_path="/mnt/a.mkv") == 1 + + remaining = [r.src_path for r in TransferPending.list_all(db.session) + if r.src_path.startswith("/mnt/")] + assert remaining == ["/mnt/b.mkv"] + + +@pytest.mark.parametrize("storage,src_path", [("", "/mnt/a.mkv"), ("local", "")]) +def test_discard_rejects_incomplete_identity(db, storage, src_path): + """ + 身份不全时必须直接返回 0,不能退化成「条件为空」的全表删除。 + """ + TransferPending.register(db.session, storage="local", src_path="/mnt/a.mkv", + now_time="2026-08-13 10:00:00") + + assert TransferPending.discard(db.session, storage=storage, src_path=src_path) == 0 + assert TransferPending.list_all(db.session) + + +def test_discard_returns_zero_when_absent(db): + """ + 注销不存在的登记返回 0,不抛异常——整理链的终态回调不应因此中断。 + """ + assert TransferPending.discard(db.session, storage="local", src_path="/nope.mkv") == 0 + + +def test_clear_empties_the_table(db): + """ + 清空返回删除条数且表内不再有登记。 + """ + TransferPending.register(db.session, storage="local", src_path="/mnt/a.mkv", + now_time="2026-08-13 10:00:00") + + assert TransferPending.clear(db.session) >= 1 + assert TransferPending.list_all(db.session) == [] + + +def test_oper_returns_plain_tuples_not_orm_instances(db): + """ + 回放接口必须返回纯元组。 + + 回放发生在会话之外,ORM 实例脱离 session 后访问属性会抛 + DetachedInstanceError——那时启动流程已经在跑,报错等于整批漏件。 + """ + oper = TransferPendingOper(db=db.session) + oper.register(storage="local", src_path="/mnt/a.mkv") + + listed = oper.list_all() + + assert ("local", "/mnt/a.mkv") in listed + assert all(isinstance(item, tuple) for item in listed) + + +def test_oper_drops_rows_with_missing_fields(db): + """ + 回放时必须跳过字段残缺的历史遗留行,不能把空存储送进整理链。 + + 列上有 NOT NULL 约束,残缺只可能表现为空串;直接绕过 register 写入, + 模拟历史数据或外部写库留下的半条记录。 + """ + db.add(TransferPending(storage="", src_path="/mnt/broken.mkv", + created_at="2026-08-13 10:00:00")) + oper = TransferPendingOper(db=db.session) + oper.register(storage="local", src_path="/mnt/ok.mkv") + + assert oper.list_all() == [("local", "/mnt/ok.mkv")] + + +def test_oper_discard_and_clear_report_counts(db): + """ + 注销与清空都要如实返回条数,调用方据此判断是否真的清理掉了。 + """ + oper = TransferPendingOper(db=db.session) + oper.register(storage="local", src_path="/mnt/a.mkv") + oper.register(storage="local", src_path="/mnt/b.mkv") + + assert oper.discard(storage="local", src_path="/mnt/a.mkv") == 1 + assert oper.clear() >= 1 + assert oper.list_all() == [] diff --git a/tests/test_db_workflow_queries.py b/tests/test_db_workflow_queries.py new file mode 100644 index 000000000..6e1877be2 --- /dev/null +++ b/tests/test_db_workflow_queries.py @@ -0,0 +1,235 @@ +""" +工作流表的查询与状态流转行为。 + +调度器按「触发类型 + 状态」取工作流,取多了会把用户暂停的流程重新跑起来, +取少了则定时任务永远不触发。状态流转里的 `state != 'P'` 守卫是暂停语义的唯一实现, +`run_count` 的自增必须留在 SQL 侧,否则并发执行会丢计数。 +""" +import asyncio + +import pytest + +from app.db.models.workflow import Workflow + + +@pytest.fixture(autouse=True) +def _track(db): + """把工作流表纳入用例级回收。""" + db.watermark(Workflow) + + +def _flow(name: str, trigger_type: str = "timer", state: str = "W", + run_count: int = 0) -> Workflow: + """构造一条工作流记录。""" + return Workflow(name=name, description=name, timer="0 * * * *", + trigger_type=trigger_type, state=state, run_count=run_count, + actions=[], flows=[], context={}, execution_state={}) + + +# --------------------------------------------------------------------------- # +# 列表查询 +# --------------------------------------------------------------------------- # + +def test_list_and_get_by_name_match_async_twins(db): + """ + 列举与按名查找的同步、异步结果必须一致。 + """ + created = db.add(_flow("wf-name")) + + assert Workflow.get_by_name(db.session, "wf-name").id == created.id + assert asyncio.run(Workflow.async_get_by_name(name="wf-name")).id == created.id + assert Workflow.get_by_name(db.session, "wf-missing") is None + + sync_ids = sorted(w.id for w in Workflow.list(db.session)) + async_ids = sorted(w.id for w in asyncio.run(Workflow.async_list())) + assert sync_ids == async_ids + + +def test_enabled_workflows_exclude_paused(db): + """ + 启用列表排除暂停状态。 + + 暂停是用户显式的「别再跑了」,被列出即等于暂停开关失效。 + """ + db.add(_flow("wf-waiting", state="W"), _flow("wf-running", state="R"), + _flow("wf-paused", state="P")) + + names = {w.name for w in Workflow.get_enabled_workflows(db.session)} + + assert {"wf-waiting", "wf-running"} <= names + assert "wf-paused" not in names + assert "wf-paused" not in {w.name for w in + asyncio.run(Workflow.async_get_enabled_workflows())} + + +def test_timer_triggered_includes_legacy_null_trigger_type(db): + """ + 定时触发列表要包含 trigger_type 为空的历史数据。 + + 该列是后加的,老工作流为空;严格等于 'timer' 会让它们从此再也不被调度, + 而用户看到的只是「任务不跑了」。 + """ + db.add(_flow("wf-timer", trigger_type="timer"), + _flow("wf-legacy", trigger_type=None), + _flow("wf-event", trigger_type="event"), + _flow("wf-timer-paused", trigger_type="timer", state="P")) + + names = {w.name for w in Workflow.get_timer_triggered_workflows(db.session)} + + assert {"wf-timer", "wf-legacy"} <= names + assert "wf-event" not in names + assert "wf-timer-paused" not in names + + +def test_event_triggered_requires_explicit_type(db): + """ + 事件触发列表只认显式的 'event',且同样排除暂停。 + """ + db.add(_flow("wf-event", trigger_type="event"), + _flow("wf-event-paused", trigger_type="event", state="P"), + _flow("wf-legacy", trigger_type=None)) + + names = {w.name for w in Workflow.get_event_triggered_workflows(db.session)} + + assert names >= {"wf-event"} + assert "wf-event-paused" not in names + assert "wf-legacy" not in names + + +def test_trigger_lists_match_async_twins(db): + """ + 定时与事件两条触发列表的同步、异步结果必须一致。 + """ + db.add(_flow("wf-t", trigger_type="timer"), _flow("wf-e", trigger_type="event")) + + assert sorted(w.id for w in Workflow.get_timer_triggered_workflows(db.session)) == \ + sorted(w.id for w in asyncio.run(Workflow.async_get_timer_triggered_workflows())) + assert sorted(w.id for w in Workflow.get_event_triggered_workflows(db.session)) == \ + sorted(w.id for w in asyncio.run(Workflow.async_get_event_triggered_workflows())) + + +# --------------------------------------------------------------------------- # +# 状态流转 +# --------------------------------------------------------------------------- # + +def test_update_state_and_start_write_the_state(db): + """ + 状态更新与启动直接落库,供调度器读到最新状态。 + """ + flow = db.add(_flow("wf-state")) + + Workflow.update_state(db.session, flow.id, "F") + assert Workflow.get_by_name(db.session, "wf-state").state == "F" + + Workflow.start(db.session, flow.id) + assert Workflow.get_by_name(db.session, "wf-state").state == "R" + + +def test_fail_and_success_respect_the_paused_guard(db): + """ + 暂停中的工作流不接受成功/失败结果写入。 + + 守卫丢失后,一个还在跑的旧任务收尾时会把用户刚设的暂停状态改掉, + 下一轮调度它又被跑起来。 + """ + paused = db.add(_flow("wf-paused", state="P")) + + Workflow.fail(db.session, paused.id, "出错了") + assert Workflow.get_by_name(db.session, "wf-paused").state == "P" + + Workflow.success(db.session, paused.id, "完成") + assert Workflow.get_by_name(db.session, "wf-paused").state == "P" + + +def test_fail_records_result_and_timestamp(db): + """ + 失败要同时写入结果与最后执行时间,供界面展示失败原因。 + """ + flow = db.add(_flow("wf-fail", state="R")) + + Workflow.fail(db.session, flow.id, "网络超时") + + updated = Workflow.get_by_name(db.session, "wf-fail") + assert (updated.state, updated.result) == ("F", "网络超时") + assert updated.last_time + + +def test_success_increments_run_count_in_sql(db): + """ + 执行次数必须在 SQL 侧自增。 + + 先读后写会在并发执行时丢计数;连续两次成功后必须是 2。 + """ + flow = db.add(_flow("wf-count", state="R", run_count=0)) + + Workflow.success(db.session, flow.id, "第一次") + Workflow.success(db.session, flow.id, "第二次") + + updated = Workflow.get_by_name(db.session, "wf-count") + assert updated.run_count == 2 + assert updated.state == "S" + + +def test_reset_clears_progress_and_optionally_the_count(db): + """ + 重置清空执行进度;是否清零执行次数由参数决定。 + + 默认保留计数是为了让「重跑」不丢失历史执行统计。 + """ + flow = db.add(_flow("wf-reset", state="F", run_count=5)) + Workflow.update_current_action(db.session, flow.id, "action-1", + {"k": "v"}, {"step": 1}) + + Workflow.reset(db.session, flow.id) + kept = Workflow.get_by_name(db.session, "wf-reset") + assert (kept.state, kept.result, kept.current_action) == ("W", None, None) + assert kept.context == {} and kept.execution_state == {} + assert kept.run_count == 5 + + Workflow.reset(db.session, flow.id, reset_count=True) + assert Workflow.get_by_name(db.session, "wf-reset").run_count == 0 + + +def test_update_current_action_appends_without_duplicating(db): + """ + 已执行动作按逗号追加且不重复登记。 + + 重复登记会让「已执行」列表无限膨胀,重跑时的跳过判断也随之失准。 + """ + flow = db.add(_flow("wf-action")) + + Workflow.update_current_action(db.session, flow.id, "a1", {"n": 1}) + Workflow.update_current_action(db.session, flow.id, "a2", {"n": 2}) + Workflow.update_current_action(db.session, flow.id, "a1", {"n": 3}) + + updated = Workflow.get_by_name(db.session, "wf-action") + assert updated.current_action == "a1,a2" + assert updated.context == {"n": 3} + + +def test_update_current_action_leaves_execution_state_untouched_when_omitted(db): + """ + 不传执行状态时保持原值,避免一次进度更新把结构化状态清空。 + """ + flow = db.add(_flow("wf-keep-state")) + Workflow.update_current_action(db.session, flow.id, "a1", {}, {"step": 7}) + + Workflow.update_current_action(db.session, flow.id, "a2", {"n": 1}) + + assert Workflow.get_by_name(db.session, "wf-keep-state").execution_state == {"step": 7} + + +def test_update_current_action_matches_async_twin(db): + """ + 同步与异步的动作追加必须给出相同的字符串。 + """ + sync_flow = db.add(_flow("wf-sync-action")) + async_flow = db.add(_flow("wf-async-action")) + + for action in ("a1", "a2", "a1"): + Workflow.update_current_action(db.session, sync_flow.id, action, {}) + asyncio.run(Workflow.async_update_current_action( + wid=async_flow.id, action_id=action, context={})) + + assert Workflow.get_by_name(db.session, "wf-sync-action").current_action == \ + Workflow.get_by_name(db.session, "wf-async-action").current_action diff --git a/tests/test_lifecycle_shutdown.py b/tests/test_lifecycle_shutdown.py index aa6ba6894..309d6d520 100644 --- a/tests/test_lifecycle_shutdown.py +++ b/tests/test_lifecycle_shutdown.py @@ -34,6 +34,12 @@ def _patch_lifespan(monkeypatch, *, failing_step: str | None = None) -> dict: monkeypatch.setattr(lifecycle, name, MagicMock()) monkeypatch.setattr(lifecycle, "init_modules", AsyncMock()) + # 启动期的引擎预热与额度核算也要打桩。不打的话这些用例会走真实的引擎创建,在测试 + # 进程里留下一个从此无人释放的全局异步引擎——NullPool 不持连接、无害,但用例就不再 + # 自洽了,而且额度核算还会去连库。 + for name in ("get_engine", "get_global_async_engine", "check_connection_budget"): + monkeypatch.setattr(lifecycle, name, MagicMock()) + system_chain = MagicMock() monkeypatch.setattr(lifecycle, "SystemChain", MagicMock(return_value=system_chain)) monkeypatch.setattr(lifecycle, "init_extra", AsyncMock()) @@ -106,6 +112,104 @@ def test_lifespan_continues_after_each_shutdown_owner_failure( _assert_completed_once(step) +def test_lifespan_creates_global_async_engine_at_startup(monkeypatch): + """启动期必须把全局异步引擎建出来一次,让异步侧恢复 fail-fast + + 引擎改为惰性创建后,启动路径只碰得到同步引擎(init_db 建表),异步驱动没装、 + 异步 URL 拼错这类问题会一路推迟到第一个异步查询——用户拿到 500、调度任务静默死掉, + 而不是启动就崩。create_async_engine 只校验 URL 与驱动导入、不建立连接,代价可以忽略。 + """ + _patch_lifespan(monkeypatch) + created = [] + monkeypatch.setattr(lifecycle, "get_global_async_engine", + lambda: created.append(1) or MagicMock()) + + async def run_lifespan(): + async with lifecycle.lifespan(FastAPI()): + pass + + asyncio.run(run_lifespan()) + + assert created, "启动期未创建全局异步引擎,异步侧的驱动/URL 错误会推迟到运行期才暴露" + + +def test_lifespan_creates_sync_engine_at_startup(monkeypatch): + """启动期也必须把同步引擎建出来一次,把首次创建钉在单线程期 + + 「init_db() 会在启动期单线程预热同步引擎」这个前提只对 run_application() 入口成立。 + 外部 supervisor 直挂 ASGI app(`gunicorn -k uvicorn.workers.UvicornWorker + app.factory:app`、`uvicorn app.main:app`)时 run_application() 不执行、init_db() 也就 + 不执行,同步引擎的首次创建退到运行期——而那时 init_scheduler() / init_monitor() 已经 + 放出上百个线程,引擎构建里那段 PRAGMA journal_mode 会让它们一起堵在创建锁上。 + """ + _patch_lifespan(monkeypatch) + created = [] + monkeypatch.setattr(lifecycle, "get_engine", + lambda: created.append(1) or MagicMock()) + + async def run_lifespan(): + async with lifecycle.lifespan(FastAPI()): + pass + + asyncio.run(run_lifespan()) + + assert created, "启动期未预热同步引擎,首次创建会退到已经放出上百个线程的运行期" + + +def test_lifespan_warms_engines_before_any_initializer(monkeypatch): + """两个引擎的预热必须排在 init_routers / init_modules 之前 + + 排在后面时,预热失败会把已经初始化好的模块晾在那里:lifespan 的 try/finally 关停块 + 要到 yield 处才开始,在它之前抛异常,stop_modules() 根本没有机会执行。 + """ + _patch_lifespan(monkeypatch) + calls = [] + monkeypatch.setattr(lifecycle, "get_engine", lambda: calls.append("sync_engine")) + monkeypatch.setattr(lifecycle, "get_global_async_engine", + lambda: calls.append("async_engine")) + monkeypatch.setattr(lifecycle, "init_routers", lambda _app: calls.append("init_routers")) + async def _init_modules(): + """init_modules 在 v3 是协程,桩也必须可 await。""" + calls.append("init_modules") + + monkeypatch.setattr(lifecycle, "init_modules", _init_modules) + + async def run_lifespan(): + async with lifecycle.lifespan(FastAPI()): + pass + + asyncio.run(run_lifespan()) + + # 不钉同步/异步两者之间的先后:那一层顺序无所谓,要紧的是它们都在 init_* 之前 + assert set(calls[:2]) == {"sync_engine", "async_engine"}, f"引擎预热没有排在最前面:{calls}" + assert calls[2:] == ["init_routers", "init_modules"], f"初始化顺序被打乱:{calls}" + + +def test_lifespan_fails_fast_when_async_engine_cannot_be_built(monkeypatch): + """异步引擎建不起来必须让启动直接失败,不能吞掉继续跑 + + 吞掉等于把 fail-fast 又还回去了:进程起来了、健康检查是绿的,只有异步请求在报错。 + """ + _patch_lifespan(monkeypatch) + + def _boom(): + """模拟异步驱动缺失。""" + raise RuntimeError("no async driver") + + monkeypatch.setattr(lifecycle, "get_global_async_engine", _boom) + + async def run_lifespan(): + async with lifecycle.lifespan(FastAPI()): + pass + + with pytest.raises(RuntimeError, match="no async driver"): + asyncio.run(run_lifespan()) + + # 失败要发生在任何东西被初始化之前,否则模块起来了却没人关:关停块在 yield 处才开始 + lifecycle.init_routers.assert_not_called() + lifecycle.init_modules.assert_not_called() + + def test_uvicorn_signal_publishes_stop_before_server_exit(monkeypatch): """Uvicorn 接管系统信号时必须先发布协作停止标志""" from app import main diff --git a/tests/test_manual_transfer_history.py b/tests/test_manual_transfer_history.py index da71f180d..487997acc 100644 --- a/tests/test_manual_transfer_history.py +++ b/tests/test_manual_transfer_history.py @@ -6,7 +6,7 @@ from app.api.endpoints.transfer import ( ) from app.chain.transfer import TransferChain from app.runtime.config import settings -from app.db.transferhistory_oper import TransferHistoryOper +from app.db.oper.transferhistory import TransferHistoryOper from app.application.history import ( clear_transfer_failures, failed_retry_count, diff --git a/tests/test_media_source_routing.py b/tests/test_media_source_routing.py index 0fd803fce..81593f46b 100644 --- a/tests/test_media_source_routing.py +++ b/tests/test_media_source_routing.py @@ -4,11 +4,8 @@ from app.chain import ChainBase from app.domain.context import MediaInfo from app.domain.meta.metabase import MetaBase from app.schemas.types import MediaSource, MediaType -from app.domain.media import ( - build_media_key, - parse_media_source_selection, - resolve_media_identity, -) +from app.domain.media import parse_media_source_selection +from app.schemas.media import build_media_key, resolve_media_identity def _chain_without_init() -> ChainBase: diff --git a/tests/test_mediascrape.py b/tests/test_mediascrape.py index 9b339dbbf..0ab617f5c 100644 --- a/tests/test_mediascrape.py +++ b/tests/test_mediascrape.py @@ -10,7 +10,7 @@ _systemconfig_stub = MagicMock() _systemconfig_stub.SystemConfigOper.return_value.get.return_value = None with stub_modules({ 'app.application.site.sites': MagicMock(), - 'app.db.systemconfig_oper': _systemconfig_stub, + 'app.db.oper.systemconfig': _systemconfig_stub, }): from app import schemas from app.chain.media import MediaChain diff --git a/tests/test_mediaserver_sync_incremental.py b/tests/test_mediaserver_sync_incremental.py index 9f7b2d3a7..807721bc7 100644 --- a/tests/test_mediaserver_sync_incremental.py +++ b/tests/test_mediaserver_sync_incremental.py @@ -9,7 +9,7 @@ from app import schemas from app.chain import mediaserver as MEDIA_SERVER_CHAIN_MODULE from app.chain.mediaserver import MediaServerChain from app.db import Base -from app.db.mediaserver_oper import MediaServerOper +from app.db.oper.mediaserver import MediaServerOper from app.db.models.mediaserver import MediaServerItem @@ -99,7 +99,7 @@ def test_sync_persists_music_without_querying_tv_episodes(database): ) chain.episodes = lambda *_args, **_kwargs: pytest.fail("音乐条目不应查询电视剧分集") - with patch("app.db.ScopedSession", database), patch.object( + with patch("app.db.decorators.ScopedSession", database), patch.object( MEDIA_SERVER_CHAIN_MODULE.ServiceConfigHelper, "get_mediaserver_configs", return_value=[SimpleNamespace(name="navidrome", enabled=True, sync_libraries=["all"])], @@ -186,7 +186,7 @@ def test_sync_updates_rows_and_removes_stale_entries(database): ) chain.episodes = lambda *_args, **_kwargs: [] - with patch("app.db.ScopedSession", database), patch.object( + with patch("app.db.decorators.ScopedSession", database), patch.object( MEDIA_SERVER_CHAIN_MODULE.ServiceConfigHelper, "get_mediaserver_configs", return_value=[SimpleNamespace(name="plex", enabled=True, sync_libraries=["movies"])], @@ -264,7 +264,7 @@ def test_sync_queries_counts_before_items_and_reports_media_progress(database): chain.items = items chain.episodes = lambda *_args, **_kwargs: [] - with patch("app.db.ScopedSession", database), patch.object( + with patch("app.db.decorators.ScopedSession", database), patch.object( MEDIA_SERVER_CHAIN_MODULE.ServiceConfigHelper, "get_mediaserver_configs", return_value=[ diff --git a/tests/test_message_notifications.py b/tests/test_message_notifications.py index cb16740b8..90c8b42b5 100644 --- a/tests/test_message_notifications.py +++ b/tests/test_message_notifications.py @@ -7,9 +7,9 @@ from app.chain import ChainBase from app.domain.context import Context, MediaInfo, TorrentInfo from app.domain.meta.metabase import MetaBase from app.db import AsyncSessionFactory, SessionFactory -from app.db.message_oper import MessageOper +from app.db.oper.message import MessageOper from app.db.models.message import Message -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.application.messaging.message import MessageHelper from app.schemas import Notification, NotificationClearScope from app.schemas.types import MediaType, NotificationType, SystemConfigKey diff --git a/tests/test_music_recognize_cache.py b/tests/test_music_recognize_cache.py index 007c7fc19..d98b6a971 100644 --- a/tests/test_music_recognize_cache.py +++ b/tests/test_music_recognize_cache.py @@ -11,7 +11,7 @@ from unittest.mock import Mock from app.api.endpoints import music as music_endpoint from app.domain.context import MusicInfo from app.domain.meta.metamusic import MetaMusic -from app.db.user_oper import get_current_active_superuser_async +from app.api.deps import get_current_active_superuser_async from app.modules.musicbrainz import music_cache as music_cache_module from app.modules.musicbrainz import MusicBrainzModule from app.modules.musicbrainz.music_cache import MusicBrainzCache diff --git a/tests/test_music_subscribe.py b/tests/test_music_subscribe.py index fb7ad9ade..51c830196 100644 --- a/tests/test_music_subscribe.py +++ b/tests/test_music_subscribe.py @@ -701,10 +701,13 @@ def test_subscribe_add_music_uses_explicit_entity_recognize(): media_chain = Mock() media_chain.recognize_media = Mock(return_value=target) subscribe_oper = Mock() - subscribe_oper.add.return_value = (1, "") + # 落库入口已迁到 app/application/subscribe.py,链路层能截到的接缝是 add_subscribe, + # 它收到的正是链路交给写入路径的那份字段 + add_subscribe = Mock(return_value=(1, "")) with patch("app.chain.subscribe.MediaChain", return_value=media_chain), \ patch("app.chain.subscribe.SubscribeOper", return_value=subscribe_oper), \ + patch("app.chain.subscribe.add_subscribe", add_subscribe), \ patch("app.chain.subscribe.MoviePilotServerHelper"), \ patch("app.chain.subscribe.eventmanager"): sid, err_msg = SubscribeChain().add( @@ -752,10 +755,11 @@ def test_subscribe_add_music_routes_new_album_sources( media_chain = Mock() media_chain.recognize_media = Mock(return_value=target) subscribe_oper = Mock() - subscribe_oper.add.return_value = (1, "") + add_subscribe = Mock(return_value=(1, "")) with patch("app.chain.subscribe.MediaChain", return_value=media_chain), \ patch("app.chain.subscribe.SubscribeOper", return_value=subscribe_oper), \ + patch("app.chain.subscribe.add_subscribe", add_subscribe), \ patch("app.chain.subscribe.MoviePilotServerHelper"), \ patch("app.chain.subscribe.eventmanager"): sid, err_msg = SubscribeChain().add( @@ -774,8 +778,8 @@ def test_subscribe_add_music_routes_new_album_sources( assert media_chain.recognize_media.call_args.kwargs["media_source"] == MediaSource(media_source) assert media_chain.recognize_media.call_args.kwargs["media_id"] == media_id assert media_chain.recognize_media.call_args.kwargs["music_type"] == MUSIC_ENTITY_ALBUM - assert subscribe_oper.add.call_args.kwargs["media_source"] == MediaSource(media_source) - assert subscribe_oper.add.call_args.kwargs["media_id"] == media_id + assert add_subscribe.call_args.kwargs["media_source"] == MediaSource(media_source) + assert add_subscribe.call_args.kwargs["media_id"] == media_id media_chain.recognize_by_meta.assert_not_called() @@ -784,9 +788,11 @@ def test_subscribe_add_rejects_music_entity_mismatch_before_database_write(): media_chain = Mock() media_chain.recognize_media.return_value = _music_info() subscribe_oper = Mock() + add_subscribe = Mock() with patch("app.chain.subscribe.MediaChain", return_value=media_chain), \ - patch("app.chain.subscribe.SubscribeOper", return_value=subscribe_oper): + patch("app.chain.subscribe.SubscribeOper", return_value=subscribe_oper), \ + patch("app.chain.subscribe.add_subscribe", add_subscribe): sid, err_msg = SubscribeChain().add( title="叶惠美", year="2003", @@ -800,7 +806,7 @@ def test_subscribe_add_rejects_music_entity_mismatch_before_database_write(): assert sid is None assert "类型不匹配" in err_msg media_chain.recognize_by_meta.assert_not_called() - subscribe_oper.add.assert_not_called() + add_subscribe.assert_not_called() def test_subscribe_add_music_fails_fast_on_offline_fallback(): @@ -809,10 +815,11 @@ def test_subscribe_add_music_fails_fast_on_offline_fallback(): media_chain = Mock() media_chain.recognize_by_meta = Mock(return_value=offline) subscribe_oper = Mock() - subscribe_oper.add.return_value = (1, "") + add_subscribe = Mock(return_value=(1, "")) with patch("app.chain.subscribe.MediaChain", return_value=media_chain), \ - patch("app.chain.subscribe.SubscribeOper", return_value=subscribe_oper): + patch("app.chain.subscribe.SubscribeOper", return_value=subscribe_oper), \ + patch("app.chain.subscribe.add_subscribe", add_subscribe): sid, err_msg = SubscribeChain().add( title="未知曲目", year=None, @@ -822,7 +829,7 @@ def test_subscribe_add_music_fails_fast_on_offline_fallback(): assert sid is None assert err_msg == "未识别到媒体信息" - subscribe_oper.add.assert_not_called() + add_subscribe.assert_not_called() def test_follow_preserves_album_entity_and_track_count(): diff --git a/tests/test_music_transfer.py b/tests/test_music_transfer.py index 17b5ac642..d35aa3d93 100644 --- a/tests/test_music_transfer.py +++ b/tests/test_music_transfer.py @@ -12,7 +12,8 @@ from app.domain.context import MusicInfo from app.application.messaging.message import TemplateHelper from app.schemas.file import FileItem from app.schemas.system import TransferDirectoryConf -from app.schemas.transfer import TransferInfo, TransferTask, TransferTorrent +from app.schemas.transfer import TransferInfo, TransferTorrent +from app.application.transfer import TransferTask from app.schemas.types import EventType, MediaType @@ -534,7 +535,11 @@ def test_success_file_aggregation_is_isolated_between_music_jobs_in_same_directo monkeypatch.setattr( "app.chain.transfer.TransferHistoryOper", - lambda: SimpleNamespace(add_success=lambda **kwargs: SimpleNamespace(id=1)), + lambda: SimpleNamespace(), + ) + monkeypatch.setattr( + "app.chain.transfer.add_transfer_success", + lambda **kwargs: SimpleNamespace(id=1), ) for task in tasks: diff --git a/tests/test_notification_template_render.py b/tests/test_notification_template_render.py index d698b7c61..a517d66f9 100644 --- a/tests/test_notification_template_render.py +++ b/tests/test_notification_template_render.py @@ -19,7 +19,7 @@ import pytest from app.domain.context import MUSIC_ENTITY_ALBUM, MusicInfo from app.domain.meta.metamusic import MetaMusic -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.application.messaging.message import MessageTemplateHelper, TemplateContextBuilder, TemplateHelper from app.schemas.message import Notification from app.schemas.types import ContentType, SystemConfigKey diff --git a/tests/test_plugin_helper.py b/tests/test_plugin_helper.py index ae2a06e34..fd330fb2c 100644 --- a/tests/test_plugin_helper.py +++ b/tests/test_plugin_helper.py @@ -916,7 +916,7 @@ class TestPluginHelper: """单插件版本查询不构建全部本地插件信息。""" try: from app.runtime.extensions.plugin_manager import PluginManager - from app.db.systemconfig_oper import SystemConfigOper + from app.db.oper.systemconfig import SystemConfigOper from app.schemas.types import SystemConfigKey except ModuleNotFoundError as exc: pytest.skip(f"missing dependency: {exc}") diff --git a/tests/test_rust_accel.py b/tests/test_rust_accel.py index 5f733ff2f..5d9e4f0c9 100644 --- a/tests/test_rust_accel.py +++ b/tests/test_rust_accel.py @@ -13,7 +13,7 @@ from app.runtime.config import settings from app.domain.meta.customization import CustomizationMatcher from app.domain.meta.releasegroup import ReleaseGroupsMatcher from app.domain.meta.streamingplatform import StreamingPlatforms -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.modules.indexer.spider import SiteSpider from app.schemas.types import SystemConfigKey from app.schemas.types import MediaType diff --git a/tests/test_search_media_sources.py b/tests/test_search_media_sources.py index e6e88a895..5f1ffbb40 100644 --- a/tests/test_search_media_sources.py +++ b/tests/test_search_media_sources.py @@ -10,7 +10,7 @@ from app.chain import subscribe as subscribe_module from app.chain.subscribe import SubscribeChain from app.domain.context import MediaInfo from app.schemas.types import MediaSource, MediaType -from app.domain.media import normalize_media_source +from app.schemas.media import normalize_media_source def test_media_source_normalization_accepts_plugin_source() -> None: diff --git a/tests/test_subscribe_chain.py b/tests/test_subscribe_chain.py index a3d6ec8b8..f5fafa83f 100644 --- a/tests/test_subscribe_chain.py +++ b/tests/test_subscribe_chain.py @@ -333,7 +333,7 @@ def _load_subscribe_chain_class(): db_model_module.Subscribe = _SubscribeModel - subscribe_oper_module = ensure_module("app.db.subscribe_oper", types.ModuleType("app.db.subscribe_oper")) + subscribe_oper_module = ensure_module("app.db.oper.subscribe", types.ModuleType("app.db.oper.subscribe")) class _SubscribeOper: def update(self, *args, **kwargs): @@ -354,9 +354,9 @@ def _load_subscribe_chain_class(): subscribe_oper_module.SubscribeOper = _SubscribeOper simple_oper_modules = { - "app.db.downloadhistory_oper": "DownloadHistoryOper", - "app.db.site_oper": "SiteOper", - "app.db.systemconfig_oper": "SystemConfigOper", + "app.db.oper.downloadhistory": "DownloadHistoryOper", + "app.db.oper.site": "SiteOper", + "app.db.oper.systemconfig": "SystemConfigOper", } for module_name_key, class_name in simple_oper_modules.items(): module = ensure_module(module_name_key, types.ModuleType(module_name_key)) @@ -3325,12 +3325,13 @@ class SubscribeProgressConsolidationTest(TestCase): chain = SubscribeChain() chain.obtain_images = lambda **_kwargs: None - class _SubscribeOper: - def add(self, **kwargs): - added.append(kwargs) - return 41, None + # 落库入口已迁到 app/application/subscribe.py,链路层的接缝是 add_subscribe; + # 截在这里拿到的就是链路交给写入路径的原始字段,正是本用例要断言的总集数 + def _add_subscribe(**kwargs): + added.append(kwargs) + return 41, None - with patch.object(module, "SubscribeOper", return_value=_SubscribeOper()), patch.object( + with patch.object(module, "add_subscribe", _add_subscribe), patch.object( module, "eventmanager", eventmanager, diff --git a/tests/test_subscribe_endpoint.py b/tests/test_subscribe_endpoint.py index 26d8dd423..53cff0141 100644 --- a/tests/test_subscribe_endpoint.py +++ b/tests/test_subscribe_endpoint.py @@ -654,13 +654,14 @@ class SubscribeEndpointTest(TestCase): """ owner-aware 创建不应把他人已有订阅当作当前用户订阅。 """ - from app.db.subscribe_oper import SubscribeOper + from app.application.subscribe import async_add_subscribe + from app.db.oper.subscribe import SubscribeOper other = _EndpointSubscribe(id=21, username="bob") own = _EndpointSubscribe(id=22, username="alice") created = SimpleNamespace(async_create=AsyncMock()) - with patch("app.db.subscribe_oper.Subscribe") as subscribe_model: + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: subscribe_model.async_exists = AsyncMock(return_value=other) subscribe_model.async_exists_by_username = AsyncMock( side_effect=[None, own] @@ -668,7 +669,8 @@ class SubscribeEndpointTest(TestCase): subscribe_model.return_value = created sid, message = asyncio.run( - SubscribeOper(db=object()).async_add( + async_add_subscribe( + subscribe_oper=SubscribeOper(db=object()), mediainfo=_EndpointMediaInfo(), username="alice", owner_scope=True, diff --git a/tests/test_subscribe_oper.py b/tests/test_subscribe_oper.py index ee7889b0b..90fc0e6e1 100644 --- a/tests/test_subscribe_oper.py +++ b/tests/test_subscribe_oper.py @@ -5,13 +5,30 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from app.application.subscribe import add_subscribe, async_add_subscribe from app.db.models.subscribe import Subscribe from app.db.models.subscribehistory import SubscribeHistory -from app.db.subscribe_oper import SubscribeOper +from app.db.oper.subscribe import SubscribeOper from app.domain.context import MusicInfo from app.schemas.types import MediaSource, MediaType +def _add(**kwargs): + """ + 经应用层写入路径新增订阅。 + + 媒体翻译住在 app/application/subscribe.py,查重与落库仍在 SubscribeOper——本文件 + 钉的是查重语义(谁被查、查几次、带哪些身份字段),所以从翻译入口进、把不带真会话 + 的 Oper 注进去,两层的契约一次跑通。 + """ + return add_subscribe(subscribe_oper=SubscribeOper(db=object()), **kwargs) + + +async def _async_add(**kwargs): + """异步写入路径,与 _add 共用注入方式。""" + return await async_add_subscribe(subscribe_oper=SubscribeOper(db=object()), **kwargs) + + def _media(episode_group): """构造订阅新增路径所需的稳定 MediaInfo 契约替身。""" return SimpleNamespace( @@ -71,11 +88,11 @@ def test_add_scopes_duplicate_lookup_by_episode_group(episode_group): persisted = SimpleNamespace(id=88) created = SimpleNamespace(create=MagicMock()) - with patch("app.db.subscribe_oper.Subscribe") as subscribe_model: + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: subscribe_model.exists.side_effect = [None, persisted] subscribe_model.return_value = created - sid, message = SubscribeOper(db=object()).add( + sid, message = _add( mediainfo=_media(episode_group), season=1, ) @@ -89,6 +106,112 @@ def test_add_scopes_duplicate_lookup_by_episode_group(episode_group): created.create.assert_called_once() +# 媒体身份的三种残缺形态。守卫写的是 ``not media_source or not media_id``——只测「两者都空」 +# 时 ``or`` 与 ``and`` 表现一致,必须把「只缺一半」的两种也测到,否则守卫被改宽也没人知道。 +# 注意:真实的 resolve_media_identity 只会返回「两个都有」或「两个都空」,构造不出半残身份, +# 所以这里必须替换掉它才能把守卫本身的契约钉住。 +_INCOMPLETE_IDENTITIES = [ + pytest.param((None, "987654321"), id="缺来源"), + pytest.param((MediaSource.TMDB, None), id="缺原生ID"), + pytest.param((None, None), id="两者皆缺"), +] + + +@pytest.mark.parametrize("identity", _INCOMPLETE_IDENTITIES) +def test_add_rejects_incomplete_media_identity(identity): + """ + 媒体身份只要缺一半就必须拒绝新增,且不得落库。 + + 身份不全的订阅写进去就是一条永远匹配不上资源的僵尸订阅,后续按身份去重也会失效。 + """ + with patch("app.application.subscribe.resolve_media_identity", return_value=identity), \ + patch("app.db.oper.subscribe.Subscribe") as subscribe_model: + result = _add(mediainfo=_media(None), season=1) + + assert result == (0, "媒体身份不完整") + # 守卫必须在查询与建模之前短路,而不是先写进去再补救 + subscribe_model.exists.assert_not_called() + subscribe_model.assert_not_called() + + +@pytest.mark.parametrize("identity", _INCOMPLETE_IDENTITIES) +def test_async_add_rejects_incomplete_media_identity(identity): + """异步新增与同步路径共用同一道身份守卫,两条链路不能一宽一严。""" + with patch("app.application.subscribe.resolve_media_identity", return_value=identity), \ + patch("app.db.oper.subscribe.Subscribe") as subscribe_model: + subscribe_model.async_exists = AsyncMock() + + result = asyncio.run(_async_add( + mediainfo=_media(None), season=1)) + + assert result == (0, "媒体身份不完整") + subscribe_model.async_exists.assert_not_awaited() + subscribe_model.assert_not_called() + + +def test_add_reports_failure_when_the_new_subscribe_cannot_be_read_back(): + """ + 创建后回查落空必须如实报「新增订阅失败」,不能把落空当成功返回。 + + 回查落空意味着写入实际没生效(唯一约束冲突、事务回滚等);此时若返回成功, + 调用方会继续按订阅已建立往下走,用户看到「订阅成功」却永远等不到资源。 + """ + created = SimpleNamespace(create=MagicMock()) + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: + subscribe_model.exists.side_effect = [None, None] + subscribe_model.return_value = created + + result = _add(mediainfo=_media(None), season=1) + + assert result == (0, "新增订阅失败") + created.create.assert_called_once() + + +def test_async_add_reports_failure_when_the_new_subscribe_cannot_be_read_back(): + """异步新增的回查落空路径与同步一致。""" + created = SimpleNamespace(async_create=AsyncMock()) + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: + subscribe_model.async_exists = AsyncMock(side_effect=[None, None]) + subscribe_model.return_value = created + + result = asyncio.run(_async_add( + mediainfo=_media(None), season=1)) + + assert result == (0, "新增订阅失败") + created.async_create.assert_awaited_once() + + +def test_add_reports_existing_subscription_without_creating(): + """ + 首次查询即命中时返回既有订阅,不再建第二条。 + + 重复建订阅会让同一部剧被两条订阅并行搜索、重复下载。 + """ + existing = SimpleNamespace(id=77) + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: + subscribe_model.exists.return_value = existing + + result = _add(mediainfo=_media(None), season=1) + + assert result == (77, "订阅已存在") + assert subscribe_model.exists.call_count == 1 + subscribe_model.assert_not_called() + + +def test_async_add_reports_existing_subscription_without_creating(): + """异步新增命中既有订阅时同样不建第二条。""" + existing = SimpleNamespace(id=78) + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: + subscribe_model.async_exists = AsyncMock(return_value=existing) + + result = asyncio.run(_async_add( + mediainfo=_media(None), season=1)) + + assert result == (78, "订阅已存在") + assert subscribe_model.async_exists.await_count == 1 + subscribe_model.assert_not_called() + + def test_music_subscribe_persists_release_cover_as_poster_and_backdrop(): """音乐订阅应把 MusicBrainz 发行封面写入订阅海报和背景字段。""" persisted = SimpleNamespace(id=92) @@ -100,11 +223,11 @@ def test_music_subscribe_persists_release_cover_as_poster_and_backdrop(): cover_url="https://coverartarchive.org/release-group/example/front-500", ) - with patch("app.db.subscribe_oper.Subscribe") as subscribe_model: + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: subscribe_model.exists.side_effect = [None, persisted] subscribe_model.return_value = created - sid, _ = SubscribeOper(db=object()).add(mediainfo=media, season=None) + sid, _ = _add(mediainfo=media, season=None) assert sid == 92 payload = subscribe_model.call_args.kwargs @@ -124,11 +247,11 @@ def test_music_subscribe_persists_numeric_year_as_string(): cover_url="https://coverartarchive.org/release/example/front-500", ) - with patch("app.db.subscribe_oper.Subscribe") as subscribe_model: + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: subscribe_model.exists.side_effect = [None, persisted] subscribe_model.return_value = created - sid, _ = SubscribeOper(db=object()).add(mediainfo=media, season=None) + sid, _ = _add(mediainfo=media, season=None) assert sid == 93 payload = subscribe_model.call_args.kwargs @@ -148,11 +271,11 @@ def test_music_album_subscription_persists_entity_and_track_count(): total_tracks=11, ) - with patch("app.db.subscribe_oper.Subscribe") as subscribe_model: + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: subscribe_model.exists.side_effect = [None, persisted] subscribe_model.return_value = created - sid, _ = SubscribeOper(db=object()).add(mediainfo=media, season=None) + sid, _ = _add(mediainfo=media, season=None) assert sid == 94 payload = subscribe_model.call_args.kwargs @@ -173,11 +296,11 @@ def test_music_recording_subscription_drops_album_track_count_and_scopes_identit total_tracks=11, ) - with patch("app.db.subscribe_oper.Subscribe") as subscribe_model: + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: subscribe_model.exists.side_effect = [None, persisted] subscribe_model.return_value = created - sid, _ = SubscribeOper(db=object()).add(mediainfo=media, season=None) + sid, _ = _add(mediainfo=media, season=None) assert sid == 95 payload = subscribe_model.call_args.kwargs @@ -195,11 +318,11 @@ def test_async_add_scopes_duplicate_lookup_by_episode_group(episode_group): persisted = SimpleNamespace(id=89) created = SimpleNamespace(async_create=AsyncMock()) - with patch("app.db.subscribe_oper.Subscribe") as subscribe_model: + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: subscribe_model.async_exists = AsyncMock(side_effect=[None, persisted]) subscribe_model.return_value = created - sid, message = asyncio.run(SubscribeOper(db=object()).async_add( + sid, message = asyncio.run(_async_add( mediainfo=_media(episode_group), season=1, )) @@ -218,11 +341,11 @@ def test_owner_scoped_add_forwards_episode_group_sync_and_async(): media = _media("eg-owner") sync_persisted = SimpleNamespace(id=90) sync_created = SimpleNamespace(create=MagicMock()) - with patch("app.db.subscribe_oper.Subscribe") as subscribe_model: + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: subscribe_model.exists_by_username.side_effect = [None, sync_persisted] subscribe_model.return_value = sync_created - sid, _ = SubscribeOper(db=object()).add( + sid, _ = _add( mediainfo=media, season=1, username="alice", @@ -237,13 +360,13 @@ def test_owner_scoped_add_forwards_episode_group_sync_and_async(): async_persisted = SimpleNamespace(id=91) async_created = SimpleNamespace(async_create=AsyncMock()) - with patch("app.db.subscribe_oper.Subscribe") as subscribe_model: + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: subscribe_model.async_exists_by_username = AsyncMock( side_effect=[None, async_persisted] ) subscribe_model.return_value = async_created - sid, _ = asyncio.run(SubscribeOper(db=object()).async_add( + sid, _ = asyncio.run(_async_add( mediainfo=media, season=1, username="alice", @@ -260,7 +383,7 @@ def test_owner_scoped_add_forwards_episode_group_sync_and_async(): def test_exists_defaults_to_main_season_episode_group(): """省略剧集组时按主季查询,显式剧集组按对应范围查询。""" oper = SubscribeOper(db=object()) - with patch("app.db.subscribe_oper.Subscribe") as subscribe_model: + with patch("app.db.oper.subscribe.Subscribe") as subscribe_model: subscribe_model.exists.return_value = SimpleNamespace(id=1) assert oper.exists( @@ -276,7 +399,7 @@ def test_exists_defaults_to_main_season_episode_group(): ) is True assert subscribe_model.exists.call_args.kwargs["episode_group"] == "eg-1" - with patch("app.db.subscribe_oper.SubscribeHistory") as history_model: + with patch("app.db.oper.subscribe.SubscribeHistory") as history_model: history_model.exists.return_value = SimpleNamespace(id=2) assert oper.exist_history( diff --git a/tests/test_subscribe_write_path.py b/tests/test_subscribe_write_path.py new file mode 100644 index 000000000..96ca90d04 --- /dev/null +++ b/tests/test_subscribe_write_path.py @@ -0,0 +1,422 @@ +""" +订阅的写入路径:app/application/subscribe.py 的 add_subscribe / async_add_subscribe。 + +这两个函数是订阅表的唯一写入口,把 MediaInfo / MusicInfo 翻译成一行订阅: +标题、年份、类型、海报背景、评分简介、剧集组、音乐实体与曲目数,再叠上 +持久化类型强转(布尔开关转整型、年份转字符串)。字段映射错了不会报错, +只会让订阅静静地记错——而搜索、洗版、完成判定、去重全都读这张表。 + +因此这里断言的是「落库后每个字段的实际值」,不是「调用了什么」: +同目录的 test_subscribe_oper.py 用替身钉的是查重语义(谁被查、查几次), +证明不了写进去的到底是什么。两者互补,缺一不可。 + +唯一的例外是**持久化类型强转**(年份转字符串、布尔开关转整型)。这两步是 +为 PostgreSQL 的严格类型检查而存在的,而测试库是 SQLite——SQLite 的类型 +亲和会在写入时把 ``2003`` 悄悄转成 ``'2003'``、把 ``True`` 转成 ``1``, +落库后的值对「写入路径自己有没有转」完全无感(已用变异验证确认:删掉 +``_normalize_year`` 后按落库值断言的用例全部照过)。所以这两类契约必须在 +建模那一刻、即 ``Subscribe(**kwargs)`` 的入参上断言,见 ``payloads`` 夹具。 +强转跟着订阅表的列走,因此仍留在 ``app/db/oper/subscribe.py``,夹具也就 +仍然钉在那个模块的 ``Subscribe`` 上。 + +同步与异步两条链路是两份逐字复制的实现,任何一条改了另一条没跟上都属于 +真实缺陷,故每个字段契约都在两条链路上各断言一次。 +""" +import asyncio + +import pytest + +from app.application.subscribe import add_subscribe, async_add_subscribe +from app.db.models.subscribe import Subscribe +from app.db.oper.subscribe import SubscribeOper +from app.domain.context import MediaInfo, MusicInfo +from app.schemas.types import MediaSource, MediaType + + +@pytest.fixture(autouse=True) +def _track(db): + """把订阅表纳入用例级回收。""" + db.watermark(Subscribe) + + +def _media_id(tag: str) -> str: + """给每个用例分配独立媒体 ID,避免共用测试库时互相去重。""" + return f"wp-{tag}" + + +def _mediainfo(media_id: str, title: str = "识别标题", + mtype: MediaType = MediaType.TV, year: str = "2026", + episode_group: str = None) -> MediaInfo: + """构造带完整展示字段的识别结果。""" + media = MediaInfo() + media.type = mtype + media.title = title + media.year = year + media.media_source = MediaSource.TMDB + media.media_id = media_id + media.episode_group = episode_group + media.vote_average = 8.5 + media.overview = "测试简介" + media.poster_path = "https://image.tmdb.org/t/p/original/poster.jpg" + media.backdrop_path = "https://image.tmdb.org/t/p/original/backdrop.jpg" + return media + + +def _musicinfo(media_id: str, music_type: str, **kwargs) -> MusicInfo: + """构造音乐订阅所需的标准音乐信息。""" + return MusicInfo(media_source=MediaSource.MusicBrainz, media_id=media_id, + music_type=music_type, **kwargs) + + +def _add(oper: SubscribeOper, is_async: bool, **kwargs): + """按链路分派到同步或异步新增,让同一份字段契约跑两遍。""" + if is_async: + return asyncio.run(async_add_subscribe(subscribe_oper=oper, **kwargs)) + return add_subscribe(subscribe_oper=oper, **kwargs) + + +def _row(db, subscribe_id: int) -> Subscribe: + """按主键读回落库的订阅行。""" + db.session.expire_all() + return Subscribe.get(db.session, subscribe_id) + + +class _SubscribeSpy: + """记录建模入参并转交真实模型,保持写入路径仍然真的落库。""" + + def __init__(self, recorded: list): + self._recorded = recorded + + def __call__(self, **kwargs): + """截获 ``Subscribe(**kwargs)`` 的入参后构造真实模型实例。""" + self._recorded.append(dict(kwargs)) + return Subscribe(**kwargs) + + def __getattr__(self, name): + """查重用的类方法(exists / async_exists 等)原样透传给真实模型。""" + return getattr(Subscribe, name) + + +@pytest.fixture +def payloads(monkeypatch): + """ + 捕获写入路径建模时的原始 kwargs,即落库**前**的值与类型。 + + 只用于持久化类型强转这一类契约:SQLite 的类型亲和会在写入时替写入路径 + 把类型「修好」,落库后的值证明不了转换真的发生过;而这两步转换恰恰是为 + PostgreSQL 而写的,漏了只会在生产库上炸。 + """ + recorded: list = [] + monkeypatch.setattr("app.db.oper.subscribe.Subscribe", _SubscribeSpy(recorded)) + return recorded + + +# 每个字段契约都在同步与异步两条链路上跑一遍:两份实现是逐字复制的, +# 只测一条等于放任另一条漂移 +_BOTH_PATHS = pytest.mark.parametrize( + "is_async", [pytest.param(False, id="sync"), pytest.param(True, id="async")] +) + + +# --------------------------------------------------------------------------- # +# 展示字段的翻译 +# --------------------------------------------------------------------------- # + +@_BOTH_PATHS +def test_add_maps_every_display_field_onto_the_row(db, is_async): + """ + 识别结果的展示字段必须完整落库。 + + 标题、年份、类型、评分、简介都是订阅列表和通知的唯一数据来源, + 错一项用户就看到一条张冠李戴的订阅。 + """ + oper = SubscribeOper() + media_id = _media_id(f"display-{is_async}") + + sid, message = _add(oper, is_async, mediainfo=_mediainfo(media_id), season=1) + + assert message == "新增订阅成功" + row = _row(db, sid) + assert row.name == "识别标题" + assert row.year == "2026" + assert row.type == MediaType.TV.value + assert row.media_source == MediaSource.TMDB.value + assert row.media_id == media_id + assert row.season == 1 + assert row.vote == 8.5 + assert row.description == "测试简介" + + +@_BOTH_PATHS +def test_add_persists_poster_and_backdrop_from_media(db, is_async): + """ + 海报与背景取自识别结果的图片接口,而不是原始路径字段。 + + 接口会把 original 尺寸换成 w500,直接存 poster_path 会让列表页 + 每张卡片都去拉原图。 + """ + oper = SubscribeOper() + media_id = _media_id(f"image-{is_async}") + + sid, _ = _add(oper, is_async, mediainfo=_mediainfo(media_id), season=1) + + row = _row(db, sid) + assert row.poster == "https://image.tmdb.org/t/p/w500/poster.jpg" + assert row.backdrop == "https://image.tmdb.org/t/p/w500/backdrop.jpg" + + +@_BOTH_PATHS +def test_add_persists_episode_group(db, is_async): + """ + 剧集组必须来自识别结果并落库。 + + 订阅按剧集组去重、搜索也按剧集组匹配集数,丢了它主季与自定义组会互相顶替。 + """ + oper = SubscribeOper() + media_id = _media_id(f"eg-{is_async}") + + sid, _ = _add(oper, is_async, + mediainfo=_mediainfo(media_id, episode_group="eg-1"), season=1) + + assert _row(db, sid).episode_group == "eg-1" + + +@_BOTH_PATHS +def test_add_stamps_creation_date(db, is_async): + """ + 新增时间由写入路径盖戳,调用方传入的值不作数。 + + 订阅列表默认按 date 排序、过期清理也读它,留空会让这条订阅永远排在最后。 + """ + oper = SubscribeOper() + media_id = _media_id(f"date-{is_async}") + + sid, _ = _add(oper, is_async, mediainfo=_mediainfo(media_id), season=1, + date="1970-01-01 00:00:00") + + date = _row(db, sid).date + assert date is not None + assert date != "1970-01-01 00:00:00" + # 形如 2026-08-14 12:34:56 + assert len(date) == 19 and date[4] == "-" and date[13] == ":" + + +# --------------------------------------------------------------------------- # +# 持久化类型强转 +# --------------------------------------------------------------------------- # + +@_BOTH_PATHS +def test_add_converts_boolean_flags_to_integers(db, payloads, is_async): + """ + 历史兼容的布尔开关建模前必须转成整型。 + + PostgreSQL 的整型列拒收布尔值,不转会让新增订阅在 PG 上直接抛类型错误。 + 断言落在建模入参上而非落库值:SQLite 会替我们把 True 存成 1, + 按落库值断言的话删掉转换也照样通过。 + """ + oper = SubscribeOper() + media_id = _media_id(f"flags-{is_async}") + + _add(oper, is_async, mediainfo=_mediainfo(media_id), season=1, + best_version=True, best_version_full=False, manual_total_episode=True) + + payload = payloads[-1] + for field, expected in (("best_version", 1), ("best_version_full", 0), + ("manual_total_episode", 1)): + assert payload[field] == expected + assert type(payload[field]) is int, f"{field} 仍是 {type(payload[field])}" + + +@_BOTH_PATHS +@pytest.mark.parametrize( + "supplied, expected", + [ + pytest.param(True, 1, id="真"), + pytest.param(False, 0, id="假"), + pytest.param(None, 0, id="缺省"), + ], +) +def test_add_normalizes_search_imdbid_to_zero_or_one(db, payloads, is_async, + supplied, expected): + """ + search_imdbid 无论传什么都归一到整型 0/1。 + + 这一列参与搜索分支判定,存进 None 或 True 会让「是否用 imdbid 搜」 + 在不同订阅上表现不一致,在 PG 上还会直接拒写。 + """ + oper = SubscribeOper() + media_id = _media_id(f"imdb-{is_async}-{supplied}") + + _add(oper, is_async, mediainfo=_mediainfo(media_id), season=1, + search_imdbid=supplied) + + payload = payloads[-1] + assert payload["search_imdbid"] == expected + assert type(payload["search_imdbid"]) is int + + +@_BOTH_PATHS +def test_add_converts_numeric_year_to_string(db, payloads, is_async): + """ + 音乐链路的年份是数字,而 year 列是字符串,建模前必须转换。 + + 不转在 PostgreSQL 上直接抛类型错误。同样只能在建模入参上验证: + SQLite 的 TEXT 亲和会把整数 2003 自动存成 '2003',读回来看不出差别。 + """ + oper = SubscribeOper() + media_id = _media_id(f"year-{is_async}") + + sid, _ = _add(oper, is_async, + mediainfo=_musicinfo(media_id, "album", title="叶惠美", year=2003)) + + assert payloads[-1]["year"] == "2003" + assert type(payloads[-1]["year"]) is str + # 落库值也要对得上,转换不能只发生在建模而在写入时被改回去 + assert _row(db, sid).year == "2003" + + +@_BOTH_PATHS +def test_add_keeps_missing_year_as_null(db, payloads, is_async): + """ + 年份缺失时留空,不能变成字符串 "None"。 + + 转字符串离无脑 str() 只有一步之遥,写成 "None" 后年份筛选会命中一个 + 不存在的年份,而这条订阅从此在按年份筛选的界面里凭空消失。 + """ + oper = SubscribeOper() + media_id = _media_id(f"noyear-{is_async}") + + sid, _ = _add(oper, is_async, + mediainfo=_musicinfo(media_id, "album", title="无年份专辑")) + + assert payloads[-1]["year"] is None + assert _row(db, sid).year is None + + +# --------------------------------------------------------------------------- # +# 音乐字段 +# --------------------------------------------------------------------------- # + +@_BOTH_PATHS +def test_add_persists_album_entity_and_track_count(db, is_async): + """ + 专辑订阅要落实体类型和总曲目数。 + + 整专完成判定拿 total_tracks 当分母,缺了这条订阅永远判不到完成。 + """ + oper = SubscribeOper() + media_id = _media_id(f"album-{is_async}") + + sid, _ = _add(oper, is_async, + mediainfo=_musicinfo(media_id, "album", title="叶惠美", total_tracks=11)) + + row = _row(db, sid) + assert row.type == MediaType.MUSIC.value + assert row.music_type == "album" + assert row.total_tracks == 11 + + +@_BOTH_PATHS +def test_add_drops_track_count_for_single_recording(db, is_async): + """ + 单曲订阅只留实体类型,专辑曲目数必须丢弃。 + + 单曲带着专辑的 total_tracks 会让完成判定把一首歌当整专等,永远不完成。 + """ + oper = SubscribeOper() + media_id = _media_id(f"recording-{is_async}") + + sid, _ = _add(oper, is_async, + mediainfo=_musicinfo(media_id, "recording", title="晴天", total_tracks=11)) + + row = _row(db, sid) + assert row.music_type == "recording" + assert row.total_tracks is None + + +@_BOTH_PATHS +def test_add_clears_music_fields_for_non_music_media(db, is_async): + """ + 非音乐媒体的音乐字段一律置空,调用方传进来的也要被覆盖。 + + 影视订阅带上 music_type 会被音乐去重逻辑当成音乐实体,造成串号。 + """ + oper = SubscribeOper() + media_id = _media_id(f"nonmusic-{is_async}") + + sid, _ = _add(oper, is_async, mediainfo=_mediainfo(media_id), season=1, + music_type="album", total_tracks=99) + + row = _row(db, sid) + assert row.music_type is None + assert row.total_tracks is None + + +@_BOTH_PATHS +def test_add_persists_music_cover_as_poster_and_backdrop(db, is_async): + """ + 音乐订阅的海报与背景都用发行封面。 + + 音乐没有独立背景图,留空会让订阅卡片在列表里显示成一块空白。 + """ + oper = SubscribeOper() + media_id = _media_id(f"cover-{is_async}") + cover = "https://coverartarchive.org/release-group/example/front-500" + + sid, _ = _add(oper, is_async, + mediainfo=_musicinfo(media_id, "album", title="封面专辑", cover_url=cover)) + + row = _row(db, sid) + assert row.poster == cover + assert row.backdrop == cover + + +# --------------------------------------------------------------------------- # +# 调用方字段的保留 +# --------------------------------------------------------------------------- # + +@_BOTH_PATHS +def test_add_keeps_caller_supplied_subscription_settings(db, is_async): + """ + 调用方传入的订阅设置要原样保留,不被媒体翻译覆盖。 + + 写入路径只负责翻译媒体身份与展示字段;把用户填的保存路径、过滤词、 + 总集数一并覆盖掉,等于用户每次新增订阅的设置都白填。 + """ + oper = SubscribeOper() + media_id = _media_id(f"settings-{is_async}") + + sid, _ = _add(oper, is_async, mediainfo=_mediainfo(media_id), season=1, + username="alice", save_path="/media/tv", keyword="关键字", + include="内嵌", exclude="预告", total_episode=12, + start_episode=3, downloader="qbittorrent", state="R") + + row = _row(db, sid) + assert row.username == "alice" + assert row.save_path == "/media/tv" + assert row.keyword == "关键字" + assert (row.include, row.exclude) == ("内嵌", "预告") + assert (row.total_episode, row.start_episode) == (12, 3) + assert row.downloader == "qbittorrent" + assert row.state == "R" + + +@_BOTH_PATHS +def test_add_overrides_caller_supplied_media_fields(db, is_async): + """ + 媒体身份与展示字段以识别结果为准,调用方传的同名值要被覆盖。 + + 否则上游一个陈旧的 name/year 就能让订阅记成另一部剧,而去重按身份走、 + 发现不了这种错位。 + """ + oper = SubscribeOper() + media_id = _media_id(f"override-{is_async}") + + sid, _ = _add(oper, is_async, + mediainfo=_mediainfo(media_id, title="正确标题", year="2026"), + season=1, name="错误标题", year="1999", type="电影") + + row = _row(db, sid) + assert row.name == "正确标题" + assert row.year == "2026" + assert row.type == MediaType.TV.value diff --git a/tests/test_sunnypt_indexer.py b/tests/test_sunnypt_indexer.py index 92b3addf7..76fe1a05c 100644 --- a/tests/test_sunnypt_indexer.py +++ b/tests/test_sunnypt_indexer.py @@ -9,7 +9,7 @@ from app.chain.download import DownloadChain from app.chain.site import SiteChain from app.runtime.config import settings from app.domain.context import TorrentInfo -from app.db.message_oper import MessageOper +from app.db.oper.message import MessageOper from app.modules.indexer import IndexerModule from app.modules.indexer.parser.sunnypt import SunnyPTSiteUserInfo from app.modules.indexer.spider.sunnypt import SunnyPTSpider diff --git a/tests/test_system_llm_test.py b/tests/test_system_llm_test.py index 7b93416a5..9d03775e2 100644 --- a/tests/test_system_llm_test.py +++ b/tests/test_system_llm_test.py @@ -45,8 +45,8 @@ _STUB_MODULES = dict([ _stub("app.runtime.extensions.module_manager", ModuleManager=_Dummy), _stub("app.application.security.access", verify_apitoken=_Dummy, verify_resource_token=_Dummy, verify_token=_Dummy), _stub("app.db.models", User=_Dummy), - _stub("app.db.systemconfig_oper", SystemConfigOper=_Dummy), - _stub("app.db.user_oper", get_current_active_superuser=_Dummy, + _stub("app.db.oper.systemconfig", SystemConfigOper=_Dummy), + _stub("app.api.deps", get_current_active_superuser=_Dummy, get_current_active_superuser_async=_Dummy, get_current_active_user_async=_Dummy), _stub("app.application.mediaserver", MediaServerHelper=_Dummy), _stub("app.application.messaging.message", MessageHelper=_Dummy), diff --git a/tests/test_system_nettest.py b/tests/test_system_nettest.py index e56b4574a..9a0ff15af 100644 --- a/tests/test_system_nettest.py +++ b/tests/test_system_nettest.py @@ -44,8 +44,8 @@ _STUB_MODULES = dict([ _stub("app.runtime.extensions.module_manager", ModuleManager=_Dummy), _stub("app.application.security.access", verify_apitoken=_Dummy, verify_resource_token=_Dummy, verify_token=_Dummy), _stub("app.db.models", User=_Dummy), - _stub("app.db.systemconfig_oper", SystemConfigOper=_Dummy), - _stub("app.db.user_oper", get_current_active_superuser=_Dummy, + _stub("app.db.oper.systemconfig", SystemConfigOper=_Dummy), + _stub("app.api.deps", get_current_active_superuser=_Dummy, get_current_active_superuser_async=_Dummy, get_current_active_user_async=_Dummy), _stub("app.agent.llm", LLMHelper=_Dummy, LLMTestError=_DummyError, LLMTestTimeout=_DummyError), _stub("app.application.mediaserver", MediaServerHelper=_Dummy), diff --git a/tests/test_systemconfig_oper.py b/tests/test_systemconfig_oper.py index be708eecb..8d80cf415 100644 --- a/tests/test_systemconfig_oper.py +++ b/tests/test_systemconfig_oper.py @@ -5,7 +5,7 @@ from concurrent.futures import ThreadPoolExecutor import pytest from app.db.models.systemconfig import SystemConfig -from app.db.systemconfig_oper import SystemConfigOper +from app.db.oper.systemconfig import SystemConfigOper from app.schemas.types import SystemConfigKey from app.foundation.singleton import Singleton diff --git a/tests/test_tmdb_cache_management.py b/tests/test_tmdb_cache_management.py index 8c8f54a59..28f846ff6 100644 --- a/tests/test_tmdb_cache_management.py +++ b/tests/test_tmdb_cache_management.py @@ -4,7 +4,7 @@ import pickle from unittest.mock import Mock from app.api.endpoints import tmdb as tmdb_endpoint -from app.db.user_oper import get_current_active_superuser_async +from app.api.deps import get_current_active_superuser_async from app.modules.themoviedb import tmdb_cache as tmdb_cache_module from app.modules.themoviedb.tmdb_cache import TmdbCache from app.schemas.types import MediaType, SystemConfigKey diff --git a/tests/test_transfer_history_write_path.py b/tests/test_transfer_history_write_path.py new file mode 100644 index 000000000..69a3e9e47 --- /dev/null +++ b/tests/test_transfer_history_write_path.py @@ -0,0 +1,291 @@ +""" +整理历史的写入路径:add_transfer_success / add_transfer_fail。 + +这两个函数是整理历史表的**唯一**写入口,整理链的每一次成败都经由它们落库。 +它们干的不是数据访问,而是把 FileItem / MetaBase / MediaInfo / TransferInfo +四个领域对象翻译成一行历史——字段映射错了不会报错,只会让历史里的季集、标题、 +来源静静地记错,而整理查重、媒体库溯源、失败重试全都读这张表。 + +因此这里断言的是「落库后每个字段的实际值」,不是「调用了什么」。 + +它们与同一张表的读侧规则(查重闸,见 test_transfer_history_gate.py)同住 +app/application/history.py;此前长在 TransferHistoryOper 上,故本文件旧名为 +test_db_transferhistory_write_path.py。 +""" +import pytest + +from app import schemas +from app.application.history import add_transfer_fail, add_transfer_success +from app.domain.context import MediaInfo +from app.domain.metainfo import MetaInfo +from app.db.models.transferhistory import TransferHistory +from app.db.oper.transferhistory import TransferHistoryOper +from app.schemas.types import MediaSource, MediaType + + +@pytest.fixture(autouse=True) +def _track(db): + """把整理历史表纳入用例级回收。""" + db.watermark(TransferHistory) + + +def _fileitem(path: str, storage: str = "local") -> schemas.FileItem: + """构造源文件项。""" + return schemas.FileItem(storage=storage, path=path, + name=path.rsplit("/", 1)[-1], type="file") + + +def _transferinfo(dest: str = "/media/片名/Season 01/片名 - S01E02.mkv", + message: str = None, with_target: bool = True) -> schemas.TransferInfo: + """构造整理结果。""" + target = _fileitem(dest) if with_target else None + return schemas.TransferInfo(success=not message, + fileitem=_fileitem("/downloads/片名.S01E02.2026.mkv"), + target_item=target, + file_list=["/downloads/片名.S01E02.2026.mkv"], + message=message) + + +def _mediainfo(title: str = "识别标题", mtype: MediaType = MediaType.TV, + year: str = "2026", category: str = "国产剧", + media_id: str = "5566", episode_group: str = None) -> MediaInfo: + """构造识别结果。""" + media = MediaInfo() + media.type = mtype + media.title = title + media.year = year + media.category = category + media.media_source = MediaSource.TMDB + media.media_id = media_id + media.episode_group = episode_group + return media + + +# --------------------------------------------------------------------------- # +# add_transfer_success +# --------------------------------------------------------------------------- # + +def test_add_success_maps_every_field_onto_the_row(db): + """ + 成功整理的字段映射必须完整落库。 + + 源/目标各自带存储、季集来自识别结果、状态为成功——这些是整理查重与媒体库 + 溯源的全部依据,错一项就查不回来。 + """ + oper = TransferHistoryOper(db=db.session) + meta = MetaInfo("片名.S01E02.2026.mkv") + + add_transfer_success(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/片名.S01E02.2026.mkv"), + mode="move", meta=meta, mediainfo=_mediainfo(), + transferinfo=_transferinfo(), + downloader="qbittorrent", download_hash="hash-ok") + + row = oper.get_by_src("/downloads/片名.S01E02.2026.mkv") + assert row.status is True + assert (row.src_storage, row.dest_storage) == ("local", "local") + assert row.dest == "/media/片名/Season 01/片名 - S01E02.mkv" + assert (row.mode, row.type, row.category) == ("move", MediaType.TV.value, "国产剧") + assert (row.title, row.year) == ("识别标题", "2026") + assert (row.seasons, row.episodes) == ("S01", "E02") + assert (row.downloader, row.download_hash) == ("qbittorrent", "hash-ok") + assert row.files == ["/downloads/片名.S01E02.2026.mkv"] + assert row.errmsg is None + + +def test_add_success_persists_both_file_items(db): + """ + 源与目标的完整文件项都要落库。 + + 整理链回滚、重新整理都要拿原始 FileItem 复原,只存路径字符串不够。 + """ + oper = TransferHistoryOper(db=db.session) + add_transfer_success(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/item.mkv"), mode="link", + meta=MetaInfo("item.mkv"), mediainfo=_mediainfo(), + transferinfo=_transferinfo(dest="/media/item.mkv")) + + row = oper.get_by_src("/downloads/item.mkv") + assert row.src_fileitem["path"] == "/downloads/item.mkv" + assert row.dest_fileitem["path"] == "/media/item.mkv" + + +def test_add_success_tolerates_missing_target_item(db): + """ + 没有目标文件项时,目标三项一并留空而不是崩。 + + 某些存储的整理只回传结果不回传条目,硬取 target_item.path 会直接抛。 + """ + oper = TransferHistoryOper(db=db.session) + add_transfer_success(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/no-target.mkv"), mode="copy", + meta=MetaInfo("no-target.mkv"), mediainfo=_mediainfo(), + transferinfo=_transferinfo(with_target=False)) + + row = oper.get_by_src("/downloads/no-target.mkv") + assert (row.dest, row.dest_storage, row.dest_fileitem) == (None, None, None) + assert row.status is True + + +def test_add_success_replaces_the_previous_record_of_the_same_source(db): + """ + 同一源路径重复整理只保留最新一条。 + + 经 add_force 的先删后插实现;留下旧行会让查重命中过期的目标路径。 + """ + oper = TransferHistoryOper(db=db.session) + for title in ("第一次", "第二次"): + add_transfer_success(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/dup.mkv"), mode="move", + meta=MetaInfo("dup.mkv"), + mediainfo=_mediainfo(title=title), + transferinfo=_transferinfo()) + + rows = oper.list_success_by_src("/downloads/dup.mkv") + assert [r.title for r in rows] == ["第二次"] + + +def test_add_success_prefers_track_title_for_music(db): + """ + 音乐文件记录曲目标题,而不是识别出的专辑/艺人名。 + + 音乐库按曲目组织,记成专辑名会让单曲在历史里全部重名、无法区分。 + """ + oper = TransferHistoryOper(db=db.session) + music_meta = MetaInfo("周杰伦 - 晴天.flac") + + add_transfer_success(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/晴天.flac"), mode="move", + meta=music_meta, + mediainfo=_mediainfo(title="叶惠美", mtype=MediaType.MUSIC), + transferinfo=_transferinfo(dest="/media/晴天.flac")) + + assert oper.get_by_src("/downloads/晴天.flac").title == "晴天" + + +def test_add_success_falls_back_to_recognized_title(db): + """ + 非音乐文件用识别标题,识别标题为空时退回文件名解析出的名字。 + """ + oper = TransferHistoryOper(db=db.session) + meta = MetaInfo("片名.S01E02.2026.mkv") + + add_transfer_success(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/fallback.mkv"), mode="move", + meta=meta, mediainfo=_mediainfo(title=None), + transferinfo=_transferinfo()) + + assert oper.get_by_src("/downloads/fallback.mkv").title == meta.name + + +# --------------------------------------------------------------------------- # +# add_transfer_fail +# --------------------------------------------------------------------------- # + +def test_add_fail_records_the_transfer_error_message(db): + """ + 整理失败时状态为失败,并保留具体错误信息。 + + 失败重试与人工排障都只能看这条 errmsg。 + """ + oper = TransferHistoryOper(db=db.session) + + add_transfer_fail(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/fail.mkv"), mode="move", + meta=MetaInfo("fail.mkv"), mediainfo=_mediainfo(), + transferinfo=_transferinfo(message="目标路径不可写")) + + row = oper.get_by_src("/downloads/fail.mkv") + assert row.status is False + assert row.errmsg == "目标路径不可写" + assert row.title == "识别标题" + + +def test_add_fail_uses_a_default_message_when_none_given(db): + """ + 整理结果没带错误信息时落一个兜底文案,不能留空。 + + 留空会让失败记录在界面上显示成「无错误」,与成功记录无从区分。 + """ + oper = TransferHistoryOper(db=db.session) + + add_transfer_fail(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/nomsg.mkv"), mode="move", + meta=MetaInfo("nomsg.mkv"), mediainfo=_mediainfo(), + transferinfo=_transferinfo(message="")) + + assert oper.get_by_src("/downloads/nomsg.mkv").errmsg == "未知错误" + + +def test_add_fail_persists_episode_group(db): + """ + 失败记录要带上剧集组——重试时靠它还原到同一个剧集组,否则会整理错季。 + """ + oper = TransferHistoryOper(db=db.session) + + add_transfer_fail(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/eg.mkv"), mode="move", + meta=MetaInfo("eg.mkv"), + mediainfo=_mediainfo(episode_group="eg-1"), + transferinfo=_transferinfo(message="出错")) + + assert oper.get_by_src("/downloads/eg.mkv").episode_group == "eg-1" + + +def test_add_fail_without_recognition_takes_the_unidentified_branch(db): + """ + 未识别到媒体信息时走另一条分支:错误文案固定,且不写目标路径。 + + 这条分支是「文件识别不出来」的唯一记录方式,丢了这些文件就彻底无迹可寻。 + """ + oper = TransferHistoryOper(db=db.session) + meta = MetaInfo("无法识别的文件.S02E05.2020.mkv") + + add_transfer_fail(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/unknown.mkv"), mode="move", + meta=meta) + + row = oper.get_by_src("/downloads/unknown.mkv") + assert row.status is False + assert row.errmsg == "未识别到媒体信息" + assert row.dest is None + assert (row.seasons, row.episodes) == ("S02", "E05") + assert row.year == meta.year + assert row.title == meta.name + + +def test_add_fail_unidentified_music_is_marked_as_a_recording(db): + """ + 未识别的音乐文件按单曲登记实体类型。 + + 音乐订阅按单曲/专辑分别查重,实体类型为空会让这条历史两边都匹配不上。 + """ + oper = TransferHistoryOper(db=db.session) + + add_transfer_fail(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/unknown.flac"), mode="move", + meta=MetaInfo("周杰伦 - 晴天.flac")) + + row = oper.get_by_src("/downloads/unknown.flac") + assert row.music_type == "recording" + assert row.type == MediaType.MUSIC.value + assert row.title == "晴天" + + +def test_add_fail_returns_the_persisted_row(db): + """ + 两条分支都要把落库后的记录返回,调用方据此拿主键做后续关联。 + """ + oper = TransferHistoryOper(db=db.session) + + identified = add_transfer_fail(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/r1.mkv"), mode="move", + meta=MetaInfo("r1.mkv"), mediainfo=_mediainfo(), + transferinfo=_transferinfo(message="错误")) + unidentified = add_transfer_fail(transfer_history_oper=oper, + fileitem=_fileitem("/downloads/r2.mkv"), mode="move", + meta=MetaInfo("r2.mkv")) + + assert identified.id is not None + assert unidentified.id is not None + assert identified.id != unidentified.id diff --git a/tests/test_transfer_job_manager.py b/tests/test_transfer_job_manager.py index 8b4159fef..52abe07fd 100644 --- a/tests/test_transfer_job_manager.py +++ b/tests/test_transfer_job_manager.py @@ -5,6 +5,7 @@ from unittest.mock import patch, MagicMock from app.runtime.config import settings from app.domain.context import MediaInfo +from app.domain.meta.metabase import MetaBase from app.domain.meta.metavideo import MetaVideo from app.chain.transfer import JobManager, TransferChain from app.application.history import ( @@ -13,7 +14,8 @@ from app.application.history import ( record_transfer_failure, ) from app.modules.filemanager.transhandler import TransHandler -from app.schemas import EpisodeFormat, FileItem, TransferInfo, TransferTask +from app.schemas import EpisodeFormat, FileItem, TransferInfo +from app.application.transfer import TransferTask from app.schemas.types import EventType, MediaSource, MediaType @@ -22,8 +24,21 @@ def _reset_failed_retries(src_path, storage=None): clear_transfer_failures(src_path, storage) -class FakeMeta: +class FakeMeta(MetaBase): + """ + 整理任务分组需要的最小剧集元数据。 + + 继承 MetaBase 而非鸭子类型:TransferTask.meta 已标注为 MetaBase,pydantic 会做 + isinstance 校验。下面三个类级属性用来遮蔽 MetaBase 上的同名 property,好让 + __init__ 能直接赋值。 + """ + + name = None + episode_list = None + season_episode = None + def __init__(self, episode: int, season: int = 1): + super().__init__(title=f"Test Show S{season:02d}E{episode:02d}") self.name = "Test Show" self.title = f"Test Show S{season:02d}E{episode:02d}" self.year = "2026" @@ -64,9 +79,17 @@ class FakeMeta: } -class FakeMedia: +class FakeMedia(MediaInfo): + """ + 与正式 MediaInfo 身份字段一致的测试媒体对象。 + + 继承 MediaInfo 而非鸭子类型:TransferTask.mediainfo 已标注为 MusicInfo | MediaInfo, + pydantic 会做 isinstance 校验。 + """ + def __init__(self, tmdb_id: int = 12345): """构造与正式 MediaInfo 身份字段一致的测试媒体对象。""" + super().__init__() self.tmdb_id = tmdb_id self.douban_id = None self.bangumi_id = None @@ -403,8 +426,10 @@ class TransferJobManagerTest(unittest.TestCase): ) with patch( - "app.chain.transfer.TransferHistoryOper", - return_value=SimpleNamespace(add_success=lambda **kwargs: SimpleNamespace(id=1)), + "app.chain.transfer.TransferHistoryOper", return_value=SimpleNamespace() + ), patch( + "app.chain.transfer.add_transfer_success", + lambda **kwargs: SimpleNamespace(id=1), ): state, errmsg = chain._TransferChain__default_callback(task, transferinfo) @@ -823,12 +848,13 @@ class TransferJobManagerTest(unittest.TestCase): transfer_type="copy", need_notify=False, ) - failed_history_oper = SimpleNamespace( - add_fail=lambda **kwargs: SimpleNamespace(id=1), - ) + failed_history_oper = SimpleNamespace() with patch( "app.chain.transfer.TransferHistoryOper", return_value=failed_history_oper, + ), patch( + "app.chain.transfer.add_transfer_fail", + lambda **kwargs: SimpleNamespace(id=1), ), patch( "app.chain.transfer.settings.AI_AGENT_ENABLE", False ), patch( @@ -864,8 +890,10 @@ class TransferJobManagerTest(unittest.TestCase): need_notify=False, ) with patch( - "app.chain.transfer.TransferHistoryOper", - return_value=SimpleNamespace(add_success=lambda **kwargs: SimpleNamespace(id=2)), + "app.chain.transfer.TransferHistoryOper", return_value=SimpleNamespace() + ), patch( + "app.chain.transfer.add_transfer_success", + lambda **kwargs: SimpleNamespace(id=2), ): state, _ = chain._TransferChain__default_callback(task, success_transferinfo) @@ -889,13 +917,14 @@ class TransferJobManagerTest(unittest.TestCase): task.download_hash = "abc123" self.assertTrue(chain.jobview.add_task(task)) - transfer_history_oper = SimpleNamespace( - add_fail=lambda **kwargs: SimpleNamespace(id=1) - ) + transfer_history_oper = SimpleNamespace() with patch( "app.chain.transfer.TransferHistoryOper", return_value=transfer_history_oper, + ), patch( + "app.chain.transfer.add_transfer_fail", + lambda **kwargs: SimpleNamespace(id=1), ), patch( "app.chain.transfer.MediaChain" ) as media_chain_cls, patch( @@ -911,6 +940,103 @@ class TransferJobManagerTest(unittest.TestCase): self.assertEqual([("abc123", "qbittorrent")], completed) self.assertEqual([], chain.jobview.list_jobs()) + def test_unrecognized_task_survives_missing_failure_history(self): + """ + 写整理历史失败(``add_transfer_fail`` 返回 None)时,未识别分支仍须走完 + 通知、作业清理与种子完成标记:历史落库是通知的附属信息,不是前置条件。 + 通知正文只省去 ``/redo`` 指引,不得因读取 ``his.id`` 抛 NoneType。 + """ + chain = make_transfer_chain() + notifications = [] + chain.post_message = lambda message, **_kwargs: notifications.append(message) + completed = [] + + def fake_transfer_completed(hashs, downloader): + completed.append((hashs, downloader)) + + chain.transfer_completed = fake_transfer_completed + chain.list_torrents = lambda **kwargs: [SimpleNamespace(progress=100)] + task = make_task(1) + task.downloader = "qbittorrent" + task.download_hash = "abc123" + self.assertTrue(chain.jobview.add_task(task)) + + with patch( + "app.chain.transfer.TransferHistoryOper", + return_value=SimpleNamespace(), + ), patch( + "app.chain.transfer.add_transfer_fail", + lambda **kwargs: None, + ), patch( + "app.chain.transfer.MediaChain" + ) as media_chain_cls, patch( + "app.chain.transfer.settings.AI_AGENT_ENABLE", False + ), patch( + "app.chain.transfer.settings.AI_AGENT_RETRY_TRANSFER", False + ): + media_chain_cls.return_value.recognize_by_meta.return_value = None + state, errmsg = chain._TransferChain__handle_transfer(task) + + self.assertFalse(state) + self.assertEqual("未识别到媒体信息", errmsg) + # 种子完成标记与作业清理都排在通知之后,通知崩掉会把它们一并跳过 + self.assertEqual([("abc123", "qbittorrent")], completed) + self.assertEqual([], chain.jobview.list_jobs()) + # 通知照发,但不含无法使用的 /redo 指引 + self.assertEqual(1, len(notifications)) + notification = notifications[0] + self.assertIn("未识别到媒体信息", notification.text) + self.assertNotIn("/redo", notification.text) + self.assertIsNone(notification.buttons) + + def test_unrecognized_task_keeps_redo_hint_when_history_written(self): + """ + 整理历史正常落库时,未识别通知须保留两条 ``/redo`` 指引与操作按钮, + 防止上一条用例被「一律删掉 /redo」这种偷懒实现蒙混过关。 + """ + chain = make_transfer_chain() + notifications = [] + chain.post_message = lambda message, **_kwargs: notifications.append(message) + chain.transfer_completed = lambda *args, **kwargs: None + chain.list_torrents = lambda **kwargs: [SimpleNamespace(progress=100)] + task = make_task(1) + task.downloader = "qbittorrent" + task.download_hash = "abc123" + self.assertTrue(chain.jobview.add_task(task)) + + with patch( + "app.chain.transfer.TransferHistoryOper", + return_value=SimpleNamespace(), + ), patch( + "app.chain.transfer.add_transfer_fail", + lambda **kwargs: SimpleNamespace(id=77), + ), patch( + "app.chain.transfer.MediaChain" + ) as media_chain_cls, patch( + "app.chain.transfer.settings.AI_AGENT_ENABLE", False + ), patch( + "app.chain.transfer.settings.AI_AGENT_RETRY_TRANSFER", False + ): + media_chain_cls.return_value.recognize_by_meta.return_value = None + chain._TransferChain__handle_transfer(task) + + self.assertEqual(1, len(notifications)) + notification = notifications[0] + self.assertIn("/redo 77\n", notification.text) + self.assertIn("/redo 77 [media_source]|[media_id]|[类型]", notification.text) + self.assertEqual( + [ + [ + {"text": "重试", "callback_data": "transfer_retry_77"}, + { + "text": "智能助手接管", + "callback_data": "transfer_ai_retry_77", + }, + ] + ], + notification.buttons, + ) + def test_do_transfer_syncs_same_stem_extra_files_by_default(self): chain = make_transfer_chain() planned = [] @@ -1338,8 +1464,10 @@ class TransferJobManagerTest(unittest.TestCase): ] with patch( - "app.chain.transfer.TransferHistoryOper", - return_value=SimpleNamespace(add_success=lambda **kwargs: SimpleNamespace(id=1)), + "app.chain.transfer.TransferHistoryOper", return_value=SimpleNamespace() + ), patch( + "app.chain.transfer.add_transfer_success", + lambda **kwargs: SimpleNamespace(id=1), ), patch( "app.chain.transfer.StorageChain" ) as storage_chain_cls: @@ -1402,8 +1530,10 @@ class TransferJobManagerTest(unittest.TestCase): ) with patch( - "app.chain.transfer.TransferHistoryOper", - return_value=SimpleNamespace(add_success=lambda **kwargs: SimpleNamespace(id=1)), + "app.chain.transfer.TransferHistoryOper", return_value=SimpleNamespace() + ), patch( + "app.chain.transfer.add_transfer_success", + lambda **kwargs: SimpleNamespace(id=1), ), patch( "app.chain.transfer.StorageChain" ) as storage_chain_cls: diff --git a/tests/test_transfer_mark_torrent_completed.py b/tests/test_transfer_mark_torrent_completed.py index 3c3e3bf18..cff7ca63b 100644 --- a/tests/test_transfer_mark_torrent_completed.py +++ b/tests/test_transfer_mark_torrent_completed.py @@ -2,16 +2,28 @@ from types import SimpleNamespace from app.chain.transfer import JobManager, TransferChain +from app.domain.meta.metabase import MetaBase from app.runtime.config import settings -from app.schemas import FileItem, TransferTask +from app.schemas import FileItem +from app.application.transfer import TransferTask from app.schemas.types import MediaType -class _FakeMeta: - """构造最小可用的剧集元数据。""" +class _FakeMeta(MetaBase): + """ + 构造最小可用的剧集元数据。 + + 继承 MetaBase 而非鸭子类型:TransferTask.meta 已标注为 MetaBase,pydantic 会做 + isinstance 校验。三个类级属性用来遮蔽 MetaBase 上的同名 property。 + """ + + name = None + episode_list = None + season_episode = None def __init__(self, episode: int, season: int = 1): """初始化剧集编号相关字段。""" + super().__init__(title=f"Test Show S{season:02d}E{episode:02d}") self.name = "Test Show" self.title = f"Test Show S{season:02d}E{episode:02d}" self.year = "2026" diff --git a/tests/test_transfer_mounted_disk_cleanup.py b/tests/test_transfer_mounted_disk_cleanup.py index d85ea2177..c0e61474e 100644 --- a/tests/test_transfer_mounted_disk_cleanup.py +++ b/tests/test_transfer_mounted_disk_cleanup.py @@ -3,7 +3,8 @@ from types import SimpleNamespace from unittest.mock import patch from app.chain.transfer import TransferChain -from app.schemas import FileItem, TransferDirectoryConf, TransferTask +from app.schemas import FileItem, TransferDirectoryConf +from app.application.transfer import TransferTask from app.adapters.system.host import SystemUtils diff --git a/tests/test_transfer_movie_collection.py b/tests/test_transfer_movie_collection.py index 9c31ed1e3..174858e9b 100644 --- a/tests/test_transfer_movie_collection.py +++ b/tests/test_transfer_movie_collection.py @@ -5,7 +5,9 @@ import pytest from app.chain.transfer import TransferChain from app.runtime.config import settings from app.domain.context import MediaInfo -from app.schemas import DownloadHistory, FileItem, TransferTask +from app.domain.meta.metabase import MetaBase +from app.schemas import DownloadHistory, FileItem +from app.application.transfer import TransferTask from app.schemas.types import MediaType @@ -31,16 +33,29 @@ def _make_chain() -> TransferChain: return chain -def _make_file_meta(year: str = "2013") -> SimpleNamespace: +class _FileMeta(MetaBase): + """ + 电影合集文件的元数据。 + + 继承 MetaBase 而非 SimpleNamespace:TransferTask.meta 已标注为 MetaBase, + pydantic 会做 isinstance 校验。name 是 MetaBase 上的 property,用类级属性遮蔽。 + """ + + name = None + + def __init__(self, year: str = "2013"): + super().__init__(title="The Hunger Games Catching Fire") + self.name = "The Hunger Games Catching Fire" + self.year = year + self.type = MediaType.UNKNOWN + self.begin_season = None + self.begin_episode = None + self.part = None + + +def _make_file_meta(year: str = "2013") -> _FileMeta: """构造电影合集文件的元数据。""" - return SimpleNamespace( - name="The Hunger Games Catching Fire", - year=year, - type=MediaType.UNKNOWN, - begin_season=None, - begin_episode=None, - part=None, - ) + return _FileMeta(year=year) def _make_history() -> SimpleNamespace: @@ -175,13 +190,16 @@ def test_movie_collection_conflict_only_drops_automatic_media( monkeypatch.setattr("app.chain.transfer.StorageChain", lambda: SimpleNamespace()) monkeypatch.setattr("app.chain.transfer.MetaInfoPath", lambda *args, **kwargs: file_meta) + # 用真 MediaInfo 而非 SimpleNamespace:它会被装进 TransferTask.mediainfo, + # 那个字段已标注为 MusicInfo | MediaInfo,pydantic 会做 isinstance 校验 + conflicting_media = MediaInfo() + conflicting_media.tmdb_id = 70160 + conflicting_media.type = MediaType.MOVIE + conflicting_media.year = "2012" + chain.do_transfer( fileitem=source_file, - mediainfo=SimpleNamespace( - tmdb_id=70160, - type=MediaType.MOVIE, - year="2012", - ), + mediainfo=conflicting_media, download_hash=history.download_hash, background=False, manual=manual, diff --git a/tests/test_transfer_overwrite_declined.py b/tests/test_transfer_overwrite_declined.py index 351511516..6d5d07890 100644 --- a/tests/test_transfer_overwrite_declined.py +++ b/tests/test_transfer_overwrite_declined.py @@ -15,9 +15,8 @@ from app.schemas.types import EventType from tests.test_transfer_job_manager import FakeMedia, make_task, make_transfer_chain -def make_history_oper(history=None, success_history=None, raise_on_query: bool = False, - add_fail_calls=None): - """构造 __is_overwrite_declined / __default_callback 查询与写入整理历史使用的替身。""" +def make_history_oper(history=None, success_history=None, raise_on_query: bool = False): + """构造 __is_overwrite_declined 查询整理历史使用的替身。""" def get_by_src(src, storage=None): if raise_on_query: @@ -27,18 +26,22 @@ def make_history_oper(history=None, success_history=None, raise_on_query: bool = def get_success_by_src(src, storage=None): return success_history - def add_fail(**kwargs): - if add_fail_calls is not None: - add_fail_calls.append(kwargs) - return SimpleNamespace(id=1) - return SimpleNamespace( get_by_src=get_by_src, get_success_by_src=get_success_by_src, - add_fail=add_fail, ) +def make_fail_recorder(calls): + """替换整理链的失败历史写入函数:只记录调用,不做字段翻译也不落库。""" + + def add_transfer_fail(**kwargs): + calls.append(kwargs) + return SimpleNamespace(id=1) + + return add_transfer_fail + + # --------------------------------------------------------------------------- # TransferChain.__is_overwrite_declined # --------------------------------------------------------------------------- @@ -138,9 +141,7 @@ def test_default_callback_skips_history_and_notification_when_overwrite_declined task = _make_failed_task() success_history = SimpleNamespace(id=99, status=True) add_fail_calls = [] - transfer_history_oper = make_history_oper( - history=success_history, add_fail_calls=add_fail_calls - ) + transfer_history_oper = make_history_oper(history=success_history) transferinfo = TransferInfo( success=False, @@ -154,6 +155,9 @@ def test_default_callback_skips_history_and_notification_when_overwrite_declined with patch( "app.chain.transfer.TransferHistoryOper", return_value=transfer_history_oper, + ), patch( + "app.chain.transfer.add_transfer_fail", + make_fail_recorder(add_fail_calls), ), patch( "app.chain.transfer.settings.AI_AGENT_ENABLE", False ), patch( @@ -183,9 +187,7 @@ def test_default_callback_keeps_original_failure_semantics_without_success_histo task = _make_failed_task() add_fail_calls = [] - transfer_history_oper = make_history_oper( - history=None, add_fail_calls=add_fail_calls - ) + transfer_history_oper = make_history_oper(history=None) transferinfo = TransferInfo( success=False, @@ -199,6 +201,9 @@ def test_default_callback_keeps_original_failure_semantics_without_success_histo with patch( "app.chain.transfer.TransferHistoryOper", return_value=transfer_history_oper, + ), patch( + "app.chain.transfer.add_transfer_fail", + make_fail_recorder(add_fail_calls), ), patch( "app.chain.transfer.settings.AI_AGENT_ENABLE", False ), patch( diff --git a/tests/test_transfer_pending_replay.py b/tests/test_transfer_pending_replay.py index 441f7b2ac..ba37b5ba9 100644 --- a/tests/test_transfer_pending_replay.py +++ b/tests/test_transfer_pending_replay.py @@ -11,7 +11,8 @@ from pathlib import Path from unittest.mock import MagicMock from app.chain.transfer import TransferChain -from app.schemas import FileItem, TransferTask +from app.schemas import FileItem +from app.application.transfer import TransferTask def _build_chain(pendingoper) -> TransferChain: diff --git a/tests/test_transfer_stale_tasks.py b/tests/test_transfer_stale_tasks.py index 1e52643d7..1b23b488d 100644 --- a/tests/test_transfer_stale_tasks.py +++ b/tests/test_transfer_stale_tasks.py @@ -2,14 +2,26 @@ from app.chain import transfer from app.chain.transfer import JobManager -from app.schemas import FileItem, TransferTask +from app.domain.meta.metabase import MetaBase +from app.schemas import FileItem +from app.application.transfer import TransferTask from app.schemas.types import MediaType -class _FakeMeta: - """提供整理任务分组需要的最小元数据。""" +class _FakeMeta(MetaBase): + """ + 提供整理任务分组需要的最小元数据。 + + 继承 MetaBase 而非鸭子类型:TransferTask.meta 已标注为 MetaBase,pydantic 会做 + isinstance 校验。三个类级属性用来遮蔽 MetaBase 上的同名 property。 + """ + + name = None + episode_list = None + season_episode = None def __init__(self): + super().__init__(title="Test Show S01E01") self.name = "Test Show" self.title = "Test Show S01E01" self.year = "2026" diff --git a/tests/test_transfer_sync_extra_files.py b/tests/test_transfer_sync_extra_files.py index a5858439f..c4549ec8f 100644 --- a/tests/test_transfer_sync_extra_files.py +++ b/tests/test_transfer_sync_extra_files.py @@ -2,17 +2,25 @@ from pathlib import Path from types import SimpleNamespace from app.chain.transfer import JobManager, TransferChain +from app.domain.meta.metabase import MetaBase from app.runtime.config import settings from app.schemas import EpisodeFormat, FileItem from app.schemas.types import MediaType -class FakeMeta: +class FakeMeta(MetaBase): """ 构造整理链路所需的最小剧集元数据。 + + 继承 MetaBase 而非鸭子类型:TransferTask.meta 已标注为 MetaBase,pydantic 会做 + isinstance 校验。name 是 MetaBase 上的 property,用类级属性遮蔽;episode_list + 在下方以 property 覆盖,不必再遮蔽。 """ + name = None + def __init__(self, episode: int): + super().__init__(title=f"Test Show S01E{episode:02d}") self.name = "Test Show" self.title = f"Test Show S01E{episode:02d}" self.year = "2026" diff --git a/tests/test_transfer_tmdb_category.py b/tests/test_transfer_tmdb_category.py index 2c82de6e5..9a3c6be17 100644 --- a/tests/test_transfer_tmdb_category.py +++ b/tests/test_transfer_tmdb_category.py @@ -3,7 +3,8 @@ from types import SimpleNamespace from app.chain.transfer import TransferChain from app.domain.context import MediaInfo from app.domain.metainfo import MetaInfo -from app.schemas import FileItem, TransferDirectoryConf, TransferTask +from app.schemas import FileItem, TransferDirectoryConf +from app.application.transfer import TransferTask from app.schemas.types import MediaSource, MediaType diff --git a/tests/test_transferhistory_media_source_migration.py b/tests/test_transferhistory_media_source_migration.py index 0ae882f5e..d4a0ac338 100644 --- a/tests/test_transferhistory_media_source_migration.py +++ b/tests/test_transferhistory_media_source_migration.py @@ -6,10 +6,11 @@ import sqlalchemy as sa from alembic.migration import MigrationContext from alembic.operations import Operations +from app.application.history import add_transfer_fail, add_transfer_success from app.domain.context import MUSIC_ENTITY_ALBUM, MusicInfo from app.domain.meta.metabase import MetaBase from app.domain.meta.metamusic import MetaMusic -from app.db.transferhistory_oper import TransferHistoryOper +from app.db.oper.transferhistory import TransferHistoryOper from app.schemas import FileItem, TransferInfo @@ -69,7 +70,7 @@ def test_failed_transfer_history_preserves_explicit_media_source() -> None: meta.media_source = "anilist" meta.media_id = "154587" - oper.add_fail( + add_transfer_fail( fileitem=FileItem( storage="local", path="/downloads/Frieren.mkv", @@ -77,6 +78,7 @@ def test_failed_transfer_history_preserves_explicit_media_source() -> None: ), mode="copy", meta=meta, + transfer_history_oper=oper, ) call = oper.add_force.call_args @@ -96,7 +98,7 @@ def test_failed_music_history_preserves_media_type_and_entity_namespace() -> Non media_id="recording-1", ) - oper.add_fail( + add_transfer_fail( fileitem=FileItem( storage="local", path="/downloads/周杰伦 - 晴天.flac", @@ -104,6 +106,7 @@ def test_failed_music_history_preserves_media_type_and_entity_namespace() -> Non ), mode="copy", meta=meta, + transfer_history_oper=oper, ) call = oper.add_force.call_args @@ -159,7 +162,7 @@ def test_music_audio_quality_migration_is_idempotent(monkeypatch) -> None: "downloadAdded": "custom download template", } monkeypatch.setattr( - "app.db.systemconfig_oper.SystemConfigOper", + "app.db.oper.systemconfig.SystemConfigOper", lambda: config_oper, ) @@ -222,7 +225,7 @@ def test_transfer_history_preserves_album_entity_context() -> None: meta = MetaMusic(title="叶惠美", artists=["周杰伦"], total_tracks=11) meta.apply_audio_quality("FLAC Lossless 24bit 96kHz 2304kbps") - oper.add_success( + add_transfer_success( fileitem=FileItem( storage="local", path="/downloads/叶惠美/01.flac", @@ -238,6 +241,7 @@ def test_transfer_history_preserves_album_entity_context() -> None: type="file", ), ), + transfer_history_oper=oper, ) call = oper.add_force.call_args diff --git a/tests/test_web_agent_stream.py b/tests/test_web_agent_stream.py index a92cb8aff..8d7102385 100644 --- a/tests/test_web_agent_stream.py +++ b/tests/test_web_agent_stream.py @@ -31,7 +31,7 @@ from app.api.endpoints.agent import ( _split_web_agent_output, ) from app.runtime.events import Event -from app.db.agentchat_oper import AgentChatOper +from app.db.oper.agentchat import AgentChatOper from app.db.models.agentchat import AgentChat from app.application.messaging.agent import build_web_agent_message_update_event from app.application.messaging.interaction import AgentInteractionOption, agent_interaction_manager, skills_interaction_manager diff --git a/tests/test_workflow_authorization.py b/tests/test_workflow_authorization.py index 312408e74..5b1953e3e 100644 --- a/tests/test_workflow_authorization.py +++ b/tests/test_workflow_authorization.py @@ -8,7 +8,7 @@ from fastapi.routing import APIRoute from app.api.endpoints import workflow as workflow_endpoint from app.application.security.access import verify_token -from app.db.user_oper import ( +from app.api.deps import ( get_current_active_manage_user, get_current_active_manage_user_async, get_current_active_user,