Compare commits

..
45 Commits
Author SHA1 Message Date
DDSRem a0ee99aacc chore: bump moviepilot-rust to 0.2.3 (#6128) 2026-07-16 06:31:31 +08:00
InfinityPacer 92918ce380 ci(pr-agent): use shared review runner (#6127) 2026-07-16 06:24:52 +08:00
InfinityPacer a4335fe753 fix(lifecycle): harden application shutdown (#6125) 2026-07-16 06:24:31 +08:00
jxxghp 107ba37834 更新 version.py 2026-07-15 20:22:27 +08:00
jxxghp c27678ce06 fix(ugreen): send client id during login 2026-07-15 17:38:31 +08:00
InfinityPacer 7725342a80 fix(scheduler): refresh plugin jobs after reload (#6124) 2026-07-15 17:29:31 +08:00
InfinityPacer 893269f8c1 fix(modules): serialize configuration reload lifecycle (#6122) 2026-07-15 17:28:49 +08:00
InfinityPacer 00d46f3aab docs: clarify docstring punctuation style (#6121) 2026-07-15 16:01:47 +08:00
jxxghp 077241b6ed Merge remote-tracking branch 'origin/v2' into v2 2026-07-15 10:46:28 +08:00
jxxghp b24a07e388 fix: enhance response data structure in filtering rules with media info 2026-07-15 10:46:21 +08:00
InfinityPacer f814c271cc refactor(runtime): tighten resource cleanup and test isolation (#6116) 2026-07-14 16:03:29 +08:00
InfinityPacer e015c67689 chore(db): add driver error diagnostics (#6115) 2026-07-14 12:31:46 +08:00
qqcomeup 98b16bda8d 优化 Docker 启动完成日志 (#6112) 2026-07-14 12:31:07 +08:00
jxxghp b8233e1789 v2.14.3 2026-07-13 18:46:35 +08:00
InfinityPacer 83107bf447 ci(pr-agent): publish native code reviews (#6110) 2026-07-13 18:41:18 +08:00
jxxghp 3a2f90c567 fix(metainfo): improve regex for episode range recognition with end markers 2026-07-13 18:02:44 +08:00
DDSRem 4826e3301c chore: bump moviepilot-rust to 0.2.2 (#6109) 2026-07-13 18:00:50 +08:00
freeman 1855ba81ec fix(meta): 副标题识别 01-26Fin 等数字范围完结标记集数 (#6105) 2026-07-13 16:50:00 +08:00
jxxghp 96ef431efc feat: 支持豆瓣识别缓存管理 2026-07-13 12:33:56 +08:00
freeman 4f2935c85e fix(transfer): 种子未下载完成时不回写已整理标签 (#6106) 2026-07-13 12:05:01 +08:00
freeman 2a49495e27 fix(jellyfin): 媒体统计改为按用户视图逐库累计 (#5915) (#6104) 2026-07-13 11:51:16 +08:00
jxxghp b628bc7209 fix: 补齐识别缓存多语言响应 2026-07-13 09:58:11 +08:00
jxxghp 29068a5846 feat: 支持 TheMovieDb 识别缓存管理 2026-07-13 09:48:08 +08:00
jxxghp 51a7120c79 完善 qBittorrent 临时标签清理 (#6093) 2026-07-12 16:51:46 +08:00
jxxghp 476dfef7d9 修复 qBittorrent 重复任务临时标签残留 (#6093) 2026-07-12 16:47:12 +08:00
jxxghp bd5ddd6158 fix: remove standalone site collector download section from README 2026-07-12 16:38:53 +08:00
jxxghp a30a48b8f4 fix: allow publishing collector artifacts manually 2026-07-12 16:37:05 +08:00
jxxghp 8e60e5571b fix: support Windows collector console encoding 2026-07-12 16:30:30 +08:00
jxxghp 18c1ec4b82 feat: add standalone site adapter collector 2026-07-12 13:51:07 +08:00
InfinityPacer 30b932e07e fix(subscribe): preserve confirmed episode floor (#6102) 2026-07-12 07:18:29 +08:00
InfinityPacer 54be1143fc feat(plugin): sync federated assets during local development (#6100) 2026-07-11 21:59:53 +08:00
秋澪Akimio 13f27854fd fix: clear Rust parse options cache after updating custom identifiers (#6097) 2026-07-11 18:15:26 +08:00
InfinityPacer 770201c48c fix(plugin): exclude build dependencies from runtime copies (#6096) 2026-07-11 18:15:00 +08:00
Xuanjie Xia 685f044312 fix: 模拟登录时页面跳转导致 page.content() 竞态失败(未知错误) (#6091) 2026-07-10 12:44:36 +08:00
qqcomeup 8c0afac5d1 feat: support prompt-bound plugin input replies (#6087) 2026-07-09 12:52:22 +08:00
qqcomeup 099ef7d5bf fix: avoid blocking plugin release history refresh (#6084) 2026-07-08 12:51:24 +08:00
InfinityPacer f3ac69669c ci(pr-agent): simplify review workflow (#6082) 2026-07-08 12:49:19 +08:00
jxxghp eb4ecd990a fix: restore full test suite 2026-07-08 08:54:53 +08:00
jxxghp b51971ee7d feat: add agent MCP support 2026-07-08 08:44:33 +08:00
InfinityPacer 6f6ed998bb ci(pr-agent): align inline review workflow (#6079) 2026-07-08 07:04:39 +08:00
drdon1234 844407dc41 修复 qBittorrent 已完成但未做种任务识别 (#6076) 2026-07-08 07:01:03 +08:00
jxxghp c54605f8ce fix: support ugreen token_id login response 2026-07-07 20:15:19 +08:00
qqcomeup 0fbf05d72f fix: handle Telegram urllib3 header formatter compatibility (#6074) 2026-07-07 19:58:49 +08:00
jxxghp 09bb32f681 fix: cool down failed subscription resources 2026-07-07 17:07:13 +08:00
qqcomeup a37f118576 perf(docker): skip image path chown by default (#6071) 2026-07-07 16:25:43 +08:00
120 changed files with 9567 additions and 651 deletions
+15 -1
View File
@@ -7,11 +7,13 @@ body:
attributes:
value: |
请说明你希望添加的功能。
站点适配请求请先按 [站点适配采集说明](https://github.com/jxxghp/MoviePilot/blob/v2/docs/site-adapter-capture.md) 生成脱敏 ZIP,并在下方附加。Issue 及附件是公开内容,提交前必须解压预览四个文件。不要上传 Cookie、Authorization、通行密钥、会话字段或任何原始数据。
- type: input
id: version
attributes:
label: 当前程序版本
description: 目前使用的程序版本
description: 目前使用的程序版本;仅提供站点采集文件且未安装 MoviePilot 时填写“不适用”
validations:
required: true
- type: dropdown
@@ -22,6 +24,9 @@ body:
options:
- Docker
- Windows
- macOS
- Linux
- 仅提供站点采集文件
validations:
required: true
- type: dropdown
@@ -32,6 +37,7 @@ body:
options:
- 主程序
- 插件
- 站点适配
- 其他
validations:
required: true
@@ -43,6 +49,14 @@ body:
placeholder: "功能改进"
validations:
required: true
- type: textarea
id: site-adapter-capture
attributes:
label: 站点适配采集文件
description: 站点适配请求必须把采集器生成并人工预览确认过的脱敏 ZIP 拖到这里;Issue 附件公开,严禁附加 Cookie、原始 HTML、HAR 或浏览器网络归档。其他类型请填写“不适用”。
placeholder: "将 moviepilot-site-capture-*.zip 拖到这里;非站点适配填写:不适用"
validations:
required: true
- type: textarea
id: references
attributes:
+15 -86
View File
@@ -1,9 +1,8 @@
name: PR Agent
name: PR-Agent
on:
pull_request_target:
# PR-Agent 通过 base repo 上下文读取 PR diff 并发布 Review,不 checkout 或执行 PR 分支代码。
# pull_request_target 允许 fork PR 使用仓库 secrets,因此 workflow 只运行固定 digest 的 PR-Agent 容器。
# Fork 审查需要目标仓库凭据;该 job 仅通过 GitHub API 读取 PR 内容,不 checkout 或执行 PR 分支代码。
types:
- opened
- reopened
@@ -11,24 +10,17 @@ on:
- review_requested
- synchronize
issue_comment:
# 手动命令如 "/review"、"/describe"、"/improve" 和 "/ask ..." 只在 PR 评论中有意义。
# issue_comment 同时覆盖普通 issue,因此 job 里还会再判断是否属于 PR。
types:
- created
- edited
permissions:
# 读取仓库内容和 PR diff。
contents: read
# 更新 PR 描述、发布 PR Review 或修改 PR 相关元数据。
pull-requests: write
# PR 评论在 GitHub API 中属于 issue comments,手动命令和总结评论需要该权限。
issues: write
jobs:
pr-agent:
name: PR-Agent review and describe
# PR 事件自动处理;评论命令仅允许指定身份在 PR 下触发,避免任意评论消耗模型配额。
if: >-
github.event.sender.type != 'Bot' &&
(
@@ -36,88 +28,25 @@ jobs:
(
github.event_name == 'issue_comment' &&
github.event.issue.pull_request != null &&
contains(fromJSON('["OWNER", "MEMBER", "COLLABORATOR", "CONTRIBUTOR", "FIRST_TIME_CONTRIBUTOR"]'), github.event.comment.author_association) &&
(
github.event.comment.body == '/review' ||
startsWith(github.event.comment.body, '/review ') ||
github.event.comment.body == '/describe' ||
startsWith(github.event.comment.body, '/describe ') ||
github.event.comment.body == '/improve' ||
startsWith(github.event.comment.body, '/improve ') ||
github.event.comment.body == '/ask' ||
startsWith(github.event.comment.body, '/ask ')
)
contains(fromJSON('["OWNER", "MEMBER", "COLLABORATOR", "CONTRIBUTOR", "FIRST_TIME_CONTRIBUTOR"]'), github.event.comment.author_association)
)
)
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.event.issue.number }}
cancel-in-progress: ${{ github.event_name == 'pull_request_target' }}
runs-on: ubuntu-latest
timeout-minutes: 20
steps:
- name: Run PR-Agent
id: pragent
# 使用版本号加 digest 固定容器构建,避免 tag 被重推后改变运行内容。
uses: docker://pragent/pr-agent:0.37.0-github_action@sha256:4ec7bac814050a1bc8c96ab2fab6b7b0f65df0049a5ec43f3fee1a0b551c28ca
- name: Run PR Review
uses: docker://ghcr.io/infinitypacer/pr-review-runner:latest
env:
# PR-Agent 使用该 token 读取 PR 元数据并发布评论。
GITHUB_TOKEN: ${{ github.token }}
# 仓库设置中添加的 SecretSettings -> Secrets and variables -> Actions。
# 该 key 只传给 PR-Agent 运行时,不写入仓库。
OPENAI_KEY: ${{ secrets.OPENAI_KEY }}
# 仓库设置中添加的 Secret。OpenAI 兼容服务通常需要填写以 "/v1" 结尾的 API 根地址。
OPENAI.API_BASE: ${{ secrets.OPENAI_API_BASE }}
# 模型、输出语言和大 diff 处理策略。
config.model: "gpt-5.5"
config.fallback_models: '["gpt-5.4"]'
config.reasoning_effort: "xhigh"
config.ai_timeout: "900"
config.response_language: "zh-CN"
config.large_patch_policy: "clip"
config.ignore_pr_title: '["^\\[Auto\\]", "^Auto"]'
config.ignore_pr_labels: '["skip pr-agent"]'
# pull_request_target 事件默认自动执行 /review 和 /describe/improve 保持手动触发。
github_action_config.auto_review: "true"
github_action_config.auto_describe: "true"
github_action_config.auto_improve: "false"
# 允许触发自动工具的 PR 动作。包含 synchronize,便于新 commit 推送后刷新结果。
github_action_config.pr_actions: '["opened", "reopened", "ready_for_review", "review_requested", "synchronize"]'
# 保留 action outputs,便于后续 workflow 编排或排查。
github_action_config.enable_output: "true"
# /describe 行为控制;与自动触发配置放在同一层,避免使用默认图表和标签策略。
pr_description.generate_ai_title: "false"
pr_description.publish_labels: "false"
pr_description.enable_pr_diagram: "false"
pr_description.collapsible_file_list: "adaptive"
pr_description.add_original_user_description: "true"
# /review 输出策略,聚焦维护者需要处理的风险和缺口。
pr_reviewer.extra_instructions: |
请用中文输出。
优先指出 P0/P1 风险,避免纠结纯格式问题。
重点检查安全、权限、状态一致性、异步/缓存、副作用和测试缺口。
pr_reviewer.num_max_findings: "5"
pr_reviewer.persistent_comment: "true"
pr_reviewer.publish_output_no_suggestions: "true"
pr_reviewer.require_tests_review: "true"
pr_reviewer.require_security_review: "true"
pr_reviewer.require_estimate_effort_to_review: "true"
pr_reviewer.require_can_be_split_review: "true"
pr_reviewer.require_todo_scan: "false"
pr_reviewer.enable_review_labels_effort: "false"
pr_reviewer.enable_review_labels_security: "true"
# /improve 和 /ask 的手动命令策略。
pr_code_suggestions.focus_only_on_problems: "true"
pr_code_suggestions.suggestions_score_threshold: "7"
pr_code_suggestions.commitable_code_suggestions: "false"
pr_questions.use_conversation_history: "true"
# 可选成本和噪音控制:
# github_action_config.auto_improve: "true"
# config.verbosity_level: "1"
# pr_reviewer.num_max_findings: "3"
PRR_AUTO_REVIEW_SCOPE: all
PRR_ALLOWED_ASSOCIATIONS: '["OWNER", "MEMBER", "COLLABORATOR", "CONTRIBUTOR", "FIRST_TIME_CONTRIBUTOR"]'
PRR_DISABLED_COMMANDS: '["/improve"]'
PRR_SKIP_LABEL: skip pr-agent
PRR_SKIP_TITLE_PATTERN: '^(?:\[Auto\]|Auto)'
config.response_language: zh-CN
@@ -0,0 +1,134 @@
name: Site Adapter Collector
on:
workflow_dispatch:
inputs:
release_tag:
description: Existing release tag to receive collector assets; leave empty for artifacts only
required: false
type: string
release:
types:
- published
permissions:
contents: read
jobs:
build:
name: Build ${{ matrix.platform_name }} collector
runs-on: ${{ matrix.runner }}
timeout-minutes: 30
strategy:
fail-fast: false
matrix:
include:
- platform_name: Windows
platform_id: windows
runner: windows-latest
source_name: moviepilot-site-collector.exe
asset_name: moviepilot-site-collector-windows.exe
artifact_name: site-adapter-collector-windows
- platform_name: macOS
platform_id: macos
runner: macos-latest
source_name: moviepilot-site-collector
asset_name: MoviePilot-Site-Collector-macOS.zip
artifact_name: site-adapter-collector-macos
- platform_name: Linux
platform_id: linux
runner: ubuntu-latest
source_name: moviepilot-site-collector
asset_name: moviepilot-site-collector-linux
artifact_name: site-adapter-collector-linux
steps:
- name: Checkout code
uses: actions/checkout@v7
- name: Set up Python
uses: actions/setup-python@v6
with:
python-version: '3.12'
cache: pip
cache-dependency-path: scripts/site_adapter_collector_requirements.txt
- name: Install build dependencies
run: |
python -m pip install --upgrade pip setuptools wheel
pip install -r scripts/site_adapter_collector_requirements.txt
- name: Build single-file collector
run: |
pyinstaller --clean --noconfirm scripts/site_adapter_collector.spec
- name: Smoke-test collector
env:
SOURCE_NAME: ${{ matrix.source_name }}
run: |
python -c "import os, subprocess; from pathlib import Path; subprocess.run([str((Path('dist') / os.environ['SOURCE_NAME']).resolve()), '--help'], check=True)"
- name: Package macOS double-click archive
if: matrix.platform_id == 'macos'
shell: bash
env:
ASSET_NAME: ${{ matrix.asset_name }}
SOURCE_NAME: ${{ matrix.source_name }}
run: |
package_dir="dist/MoviePilot-Collector"
mkdir -p "$package_dir"
cp "dist/$SOURCE_NAME" "$package_dir/moviepilot-site-collector-macos"
cp scripts/start-site-adapter-collector.command "$package_dir/start-site-adapter-collector.command"
chmod +x "$package_dir/moviepilot-site-collector-macos"
chmod +x "$package_dir/start-site-adapter-collector.command"
cd dist
COPYFILE_DISABLE=1 zip -q -r -X "$ASSET_NAME" MoviePilot-Collector
- name: Rename Windows and Linux collector
if: matrix.platform_id != 'macos'
env:
ASSET_NAME: ${{ matrix.asset_name }}
SOURCE_NAME: ${{ matrix.source_name }}
run: |
python -c "import os; from pathlib import Path; (Path('dist') / os.environ['SOURCE_NAME']).replace(Path('dist') / os.environ['ASSET_NAME'])"
- name: Generate SHA-256 checksum
env:
ASSET_NAME: ${{ matrix.asset_name }}
run: |
python -c "import hashlib, os; from pathlib import Path; path = Path('dist') / os.environ['ASSET_NAME']; path.with_name(path.name + '.sha256').write_text(f'{hashlib.sha256(path.read_bytes()).hexdigest()} {path.name}\n', encoding='utf-8')"
- name: Upload collector artifact
uses: actions/upload-artifact@v7
with:
name: ${{ matrix.artifact_name }}
path: |
dist/${{ matrix.asset_name }}
dist/${{ matrix.asset_name }}.sha256
if-no-files-found: error
retention-days: 3
publish:
name: Upload collectors to release
if: github.event_name == 'release' || inputs.release_tag != ''
needs:
- build
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- name: Download collector artifacts
uses: actions/download-artifact@v8
with:
pattern: site-adapter-collector-*
path: release-assets
merge-multiple: true
- name: Upload assets to published release
env:
GH_TOKEN: ${{ github.token }}
RELEASE_TAG: ${{ github.event.release.tag_name || inputs.release_tag }}
run: |
gh release view "$RELEASE_TAG" --repo "$GITHUB_REPOSITORY" >/dev/null
gh release upload "$RELEASE_TAG" release-assets/* --clobber --repo "$GITHUB_REPOSITORY"
+1
View File
@@ -37,6 +37,7 @@ coverage.json
htmlcov/
.vscode
venv
moviepilot-site-capture-*.zip
# Pylint
pylint-report.json
+1
View File
@@ -59,6 +59,7 @@ curl -fsSL https://raw.githubusercontent.com/jxxghp/MoviePilot/v2/scripts/bootst
- 文档规则入口:[docs/rules/README.md](docs/rules/README.md)
- 开发环境与本地源码运行:[docs/development-setup.md](docs/development-setup.md)
- 测试说明:[docs/testing.md](docs/testing.md)
- 新站点适配采集与 Feature Request 提交:[docs/site-adapter-capture.md](docs/site-adapter-capture.md)
- REST API 文档:https://api.movie-pilot.org
- 插件开发说明:https://wiki.movie-pilot.org/zh/plugindev
+1
View File
@@ -58,6 +58,7 @@ Before contributing, read the repository rules and local environment guide, keep
- Rule index: [docs/rules/README.md](docs/rules/README.md)
- Development setup and local source run: [docs/development-setup.md](docs/development-setup.md)
- Testing guide: [docs/testing.md](docs/testing.md)
- New site adapter capture and Feature Request submission: [docs/site-adapter-capture.md](docs/site-adapter-capture.md)
- REST API documentation: https://api.movie-pilot.org
- Plugin development guide: https://wiki.movie-pilot.org/zh/plugindev
+40 -1
View File
@@ -51,7 +51,9 @@ from app.agent.middleware.tool_selection import ToolSelectorMiddleware
from app.agent.middleware.usage import UsageMiddleware
from app.agent.prompt import prompt_manager
from app.agent.runtime import agent_runtime_manager
from app.agent.mcp import agent_mcp_manager
from app.agent.tools.factory import MoviePilotToolFactory
from app.agent.tools.impl.mcp import create_external_mcp_tools
from app.chain import ChainBase
from app.core.config import settings
from app.core.event import eventmanager
@@ -1041,6 +1043,7 @@ class MoviePilotAgent:
settings.LLM_MAX_ITERATIONS,
self._public_runtime_config_signature(runtime_config),
agent_runtime_manager.current_signature(),
agent_mcp_manager.config_signature(),
)
def _get_cached_agent(
@@ -1097,6 +1100,39 @@ class MoviePilotAgent:
allow_message_tools=False,
)
async def _initialize_mcp_tools(self) -> List:
"""
初始化外部 MCP 工具列表。
"""
return await create_external_mcp_tools(
session_id=self.session_id,
user_id=self.user_id,
channel=self.channel,
source=self.source,
username=self.username,
stream_handler=self.stream_handler,
agent_context=self._tool_context,
)
async def _initialize_subagent_mcp_tools(self) -> List:
"""
初始化子代理可用的外部 MCP 工具列表。
"""
return await create_external_mcp_tools(
session_id=self.session_id,
user_id=self.user_id,
channel=self.channel,
source=self.source,
username=self.username,
stream_handler=None,
agent_context={
"user_reply_sent": False,
"reply_mode": None,
"should_dispatch_reply": False,
"is_admin": bool(self._tool_context.get("is_admin")),
},
)
async def _create_agent(self, streaming: bool = False):
"""
创建 LangGraph Agent(使用 create_agent + SummarizationMiddleware
@@ -1126,6 +1162,7 @@ class MoviePilotAgent:
# 工具列表
tools = self._initialize_tools()
tools.extend(await self._initialize_mcp_tools())
skills_middleware = SkillsMiddleware(
sources=[str(agent_runtime_manager.skills_dir)],
bundled_skills_dir=str(settings.ROOT_PATH / "skills"),
@@ -1142,9 +1179,11 @@ class MoviePilotAgent:
activity_log_tools = list(
getattr(activity_log_middleware, "tools", []) or []
)
subagent_tools = self._initialize_subagent_tools()
subagent_tools.extend(await self._initialize_subagent_mcp_tools())
subagent_middlewares, subagent_task_tools = create_subagent_middlewares(
model=non_streaming_model,
tools=self._initialize_subagent_tools(),
tools=subagent_tools,
stream_handler=self.stream_handler,
)
max_tools = settings.LLM_MAX_TOOLS
+600
View File
@@ -0,0 +1,600 @@
"""Agent 外部 MCP 客户端与配置管理。"""
from __future__ import annotations
import asyncio
import hashlib
import json
import os
import re
import uuid
from dataclasses import dataclass
from typing import Any, Optional
from urllib.parse import urljoin
from app.db.systemconfig_oper import SystemConfigOper
from app.log import logger
from app.schemas.agent import (
AgentMcpServerConfig,
AgentMcpServerTestResult,
AgentMcpServerToolInfo,
)
from app.schemas.types import SystemConfigKey
from app.utils.http import AsyncRequestUtils
MCP_PROTOCOL_VERSION = "2025-11-25"
MCP_CLIENT_NAME = "MoviePilot Agent"
DEFAULT_MCP_TIMEOUT = 30
MCP_TOOL_NAME_PATTERN = re.compile(r"[^a-zA-Z0-9_]+")
@dataclass(frozen=True)
class AgentMcpToolSpec:
"""已发现的外部 MCP 工具定义。"""
server: AgentMcpServerConfig
name: str
agent_tool_name: str
description: str
input_schema: dict[str, Any]
def _normalize_identifier(value: str, fallback: str = "mcp") -> str:
"""把服务器或工具名称转换为 Agent 工具可用的标识片段。"""
normalized = MCP_TOOL_NAME_PATTERN.sub("_", str(value or "").strip())
normalized = re.sub(r"_+", "_", normalized).strip("_").lower()
if not normalized:
normalized = fallback
if normalized[0].isdigit():
normalized = f"{fallback}_{normalized}"
return normalized[:64]
def _normalize_timeout(value: Any) -> int:
"""规范化 MCP 连接和调用超时时间。"""
try:
timeout = int(value or DEFAULT_MCP_TIMEOUT)
except (TypeError, ValueError):
timeout = DEFAULT_MCP_TIMEOUT
return min(max(timeout, 1), 600)
def _normalize_string_dict(value: Any) -> dict[str, str]:
"""规范化请求头和环境变量字典,移除空键。"""
if not isinstance(value, dict):
return {}
normalized: dict[str, str] = {}
for key, item in value.items():
normalized_key = str(key or "").strip()
if not normalized_key:
continue
normalized[normalized_key] = str(item or "")
return normalized
def _normalize_input_schema(value: Any) -> dict[str, Any]:
"""规范化 MCP 工具参数 Schema,保证至少是 object schema。"""
if not isinstance(value, dict):
return {"type": "object", "properties": {}, "required": []}
schema = dict(value)
schema.setdefault("type", "object")
schema.setdefault("properties", {})
schema.setdefault("required", [])
return schema
def _build_agent_tool_name(server: AgentMcpServerConfig, tool_name: str) -> str:
"""构造注入 Agent 的外部 MCP 工具名。"""
prefix = server.tool_prefix or f"mcp_{server.name or server.id}"
normalized_prefix = _normalize_identifier(prefix, fallback="mcp")
normalized_tool_name = _normalize_identifier(tool_name, fallback="tool")
if normalized_tool_name.startswith(f"{normalized_prefix}_"):
return normalized_tool_name
return f"{normalized_prefix}_{normalized_tool_name}"[:128]
def _jsonrpc_message(method: str, params: Optional[dict[str, Any]] = None, *, request_id: Optional[str] = None) -> dict:
"""构造 JSON-RPC 2.0 消息。"""
payload = {"jsonrpc": "2.0", "method": method}
if request_id is not None:
payload["id"] = request_id
if params is not None:
payload["params"] = params
return payload
def _raise_for_jsonrpc_error(payload: Any) -> None:
"""检查 JSON-RPC 响应错误并转换为运行时异常。"""
if isinstance(payload, dict) and payload.get("error"):
error = payload["error"]
if isinstance(error, dict):
message = error.get("message") or error
else:
message = error
raise RuntimeError(f"MCP JSON-RPC 错误: {message}")
def _extract_jsonrpc_result(payload: Any, request_id: str) -> Any:
"""从 JSON-RPC 响应中提取 result 字段。"""
if not isinstance(payload, dict):
raise RuntimeError("MCP 响应不是有效 JSON 对象")
if payload.get("id") != request_id:
raise RuntimeError("MCP 响应 ID 与请求不匹配")
_raise_for_jsonrpc_error(payload)
return payload.get("result")
async def _iter_sse_events(response) -> Any:
"""按 SSE 事件格式迭代响应流。"""
event_name = "message"
data_lines: list[str] = []
async for raw_line in response.aiter_lines():
line = raw_line.rstrip("\r")
if not line:
if data_lines:
yield {"event": event_name, "data": "\n".join(data_lines)}
event_name = "message"
data_lines = []
continue
if line.startswith(":"):
continue
field, _, value = line.partition(":")
if value.startswith(" "):
value = value[1:]
if field == "event":
event_name = value or "message"
elif field == "data":
data_lines.append(value)
if data_lines:
yield {"event": event_name, "data": "\n".join(data_lines)}
def _parse_sse_text_response(text: str, request_id: str) -> Any:
"""从非流式 SSE 文本响应中提取匹配请求的 JSON-RPC 结果。"""
event_name = "message"
data_lines: list[str] = []
for raw_line in str(text or "").splitlines():
line = raw_line.rstrip("\r")
if not line:
if data_lines:
payload = _load_sse_json_payload(event_name, "\n".join(data_lines))
if isinstance(payload, dict) and payload.get("id") == request_id:
return _extract_jsonrpc_result(payload, request_id)
event_name = "message"
data_lines = []
continue
if line.startswith(":"):
continue
field, _, value = line.partition(":")
if value.startswith(" "):
value = value[1:]
if field == "event":
event_name = value or "message"
elif field == "data":
data_lines.append(value)
if data_lines:
payload = _load_sse_json_payload(event_name, "\n".join(data_lines))
if isinstance(payload, dict) and payload.get("id") == request_id:
return _extract_jsonrpc_result(payload, request_id)
raise RuntimeError("MCP SSE 响应中未找到匹配请求")
def _load_sse_json_payload(event_name: str, data: str) -> Optional[dict]:
"""解析 SSE data 中的 JSON-RPC 消息。"""
if event_name not in {"message", "messages"}:
return None
try:
payload = json.loads(data)
except (TypeError, ValueError):
return None
return payload if isinstance(payload, dict) else None
class _StdioMcpSession:
"""stdio MCP 会话,按一次操作生命周期启动外部进程。"""
def __init__(self, server: AgentMcpServerConfig) -> None:
self.server = server
self.process: Optional[asyncio.subprocess.Process] = None
self.stderr_task: Optional[asyncio.Task] = None
async def __aenter__(self) -> "_StdioMcpSession":
"""启动 stdio MCP 子进程。"""
if not self.server.command:
raise RuntimeError("stdio MCP 服务器缺少启动命令")
env = os.environ.copy()
env.update(self.server.env or {})
self.process = await asyncio.create_subprocess_exec(
self.server.command,
*(self.server.args or []),
stdin=asyncio.subprocess.PIPE,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
env=env,
)
self.stderr_task = asyncio.create_task(self._drain_stderr())
return self
async def __aexit__(self, exc_type, exc, tb) -> None:
"""结束 stdio MCP 子进程。"""
if self.stderr_task:
self.stderr_task.cancel()
if not self.process:
return
if self.process.returncode is None:
self.process.terminate()
try:
await asyncio.wait_for(self.process.wait(), timeout=2)
except asyncio.TimeoutError:
self.process.kill()
await self.process.wait()
async def _drain_stderr(self) -> None:
"""持续读取子进程 stderr,避免缓冲区阻塞。"""
if not self.process or not self.process.stderr:
return
try:
while True:
line = await self.process.stderr.readline()
if not line:
break
logger.debug(f"MCP stdio[{self.server.name}] stderr: {line.decode(errors='replace').strip()}")
except asyncio.CancelledError:
return
async def notify(self, method: str, params: Optional[dict[str, Any]] = None) -> None:
"""发送不需要响应的 JSON-RPC 通知。"""
await self._write_json(_jsonrpc_message(method, params))
async def request(self, method: str, params: Optional[dict[str, Any]] = None) -> Any:
"""发送 JSON-RPC 请求并等待响应。"""
request_id = uuid.uuid4().hex
await self._write_json(_jsonrpc_message(method, params, request_id=request_id))
while True:
payload = await self._read_json()
if payload.get("id") == request_id:
return _extract_jsonrpc_result(payload, request_id)
async def _write_json(self, payload: dict) -> None:
"""写入一行 JSON-RPC 消息。"""
if not self.process or not self.process.stdin:
raise RuntimeError("stdio MCP 进程未启动")
data = json.dumps(payload, ensure_ascii=False, separators=(",", ":")) + "\n"
self.process.stdin.write(data.encode("utf-8"))
await self.process.stdin.drain()
async def _read_json(self) -> dict:
"""从 stdout 读取一行 JSON-RPC 消息。"""
if not self.process or not self.process.stdout:
raise RuntimeError("stdio MCP 进程未启动")
timeout = _normalize_timeout(self.server.timeout)
while True:
line = await asyncio.wait_for(self.process.stdout.readline(), timeout=timeout)
if not line:
raise RuntimeError("stdio MCP 进程已退出")
try:
payload = json.loads(line.decode("utf-8"))
except ValueError:
logger.debug(f"忽略非 JSON MCP stdout 行: {line!r}")
continue
if isinstance(payload, dict):
return payload
class _HttpMcpSession:
"""Streamable HTTP MCP 会话。"""
def __init__(self, server: AgentMcpServerConfig) -> None:
self.server = server
self.session_id: Optional[str] = None
async def __aenter__(self) -> "_HttpMcpSession":
"""进入 HTTP MCP 会话。"""
return self
async def __aexit__(self, exc_type, exc, tb) -> None:
"""退出 HTTP MCP 会话。"""
return None
async def notify(self, method: str, params: Optional[dict[str, Any]] = None) -> None:
"""发送不需要响应的 JSON-RPC 通知。"""
await self._post(_jsonrpc_message(method, params), expect_response=False)
async def request(self, method: str, params: Optional[dict[str, Any]] = None) -> Any:
"""发送 JSON-RPC 请求并等待响应。"""
request_id = uuid.uuid4().hex
return await self._post(
_jsonrpc_message(method, params, request_id=request_id),
expect_response=True,
request_id=request_id,
)
async def _post(
self,
payload: dict,
*,
expect_response: bool,
request_id: Optional[str] = None,
) -> Any:
"""向 Streamable HTTP MCP 服务发送一条 JSON-RPC 消息。"""
if not self.server.url:
raise RuntimeError("HTTP MCP 服务器缺少 URL")
headers = {
"Accept": "application/json, text/event-stream",
"Content-Type": "application/json",
**(self.server.headers or {}),
}
if self.session_id:
headers["Mcp-Session-Id"] = self.session_id
response = await AsyncRequestUtils(
headers=headers,
timeout=_normalize_timeout(self.server.timeout),
content_type="application/json",
accept_type="application/json, text/event-stream",
http2=False,
).post_res(self.server.url, json=payload, raise_exception=True)
try:
if not response:
raise RuntimeError("HTTP MCP 请求无响应")
response.raise_for_status()
session_id = response.headers.get("Mcp-Session-Id")
if session_id:
self.session_id = session_id
if not expect_response:
return None
content_type = response.headers.get("content-type", "").lower()
if "text/event-stream" in content_type:
return _parse_sse_text_response(response.text, request_id or "")
data = response.json()
return _extract_jsonrpc_result(data, request_id or "")
finally:
if response is not None:
await response.aclose()
class _SseMcpSession:
"""旧版 HTTP+SSE MCP 会话。"""
def __init__(self, server: AgentMcpServerConfig) -> None:
self.server = server
self.response = None
self.endpoint: Optional[str] = None
self._stream_manager = None
self._event_iterator = None
async def __aenter__(self) -> "_SseMcpSession":
"""打开 SSE 流并读取服务端回传的 POST endpoint。"""
if not self.server.url:
raise RuntimeError("SSE MCP 服务器缺少 URL")
self._stream_manager = AsyncRequestUtils(
headers={"Accept": "text/event-stream", **(self.server.headers or {})},
timeout=_normalize_timeout(self.server.timeout),
accept_type="text/event-stream",
http2=False,
).get_stream(self.server.url, raise_exception=True)
self.response = await self._stream_manager.__aenter__()
if not self.response:
raise RuntimeError("SSE MCP 连接无响应")
self.response.raise_for_status()
self._event_iterator = _iter_sse_events(self.response).__aiter__()
self.endpoint = await self._read_endpoint()
return self
async def __aexit__(self, exc_type, exc, tb) -> None:
"""关闭 SSE 流。"""
if self._stream_manager:
await self._stream_manager.__aexit__(exc_type, exc, tb)
async def notify(self, method: str, params: Optional[dict[str, Any]] = None) -> None:
"""发送不需要响应的 JSON-RPC 通知。"""
await self._post(_jsonrpc_message(method, params))
async def request(self, method: str, params: Optional[dict[str, Any]] = None) -> Any:
"""发送 JSON-RPC 请求并等待 SSE 流上的响应。"""
request_id = uuid.uuid4().hex
await self._post(_jsonrpc_message(method, params, request_id=request_id))
timeout = _normalize_timeout(self.server.timeout)
while True:
event = await asyncio.wait_for(self._event_iterator.__anext__(), timeout=timeout)
payload = _load_sse_json_payload(event.get("event", ""), event.get("data", ""))
if isinstance(payload, dict) and payload.get("id") == request_id:
return _extract_jsonrpc_result(payload, request_id)
async def _read_endpoint(self) -> str:
"""读取 SSE endpoint 事件中的 POST 地址。"""
timeout = _normalize_timeout(self.server.timeout)
while True:
event = await asyncio.wait_for(self._event_iterator.__anext__(), timeout=timeout)
if event.get("event") != "endpoint":
continue
endpoint = str(event.get("data") or "").strip()
if not endpoint:
continue
return urljoin(self.server.url, endpoint)
async def _post(self, payload: dict) -> None:
"""向 SSE 握手返回的 endpoint 发送 JSON-RPC 消息。"""
if not self.endpoint:
raise RuntimeError("SSE MCP endpoint 未初始化")
response = await AsyncRequestUtils(
headers={
"Accept": "application/json",
"Content-Type": "application/json",
**(self.server.headers or {}),
},
timeout=_normalize_timeout(self.server.timeout),
content_type="application/json",
accept_type="application/json",
http2=False,
).post_res(self.endpoint, json=payload, raise_exception=True)
try:
if not response:
raise RuntimeError("SSE MCP POST 请求无响应")
response.raise_for_status()
finally:
if response is not None:
await response.aclose()
async def _open_mcp_session(server: AgentMcpServerConfig):
"""根据配置创建对应的 MCP 传输会话。"""
transport = "http" if server.transport == "streamable_http" else server.transport
if transport == "stdio":
return _StdioMcpSession(server)
if transport == "sse":
return _SseMcpSession(server)
if transport == "http":
return _HttpMcpSession(server)
raise RuntimeError(f"不支持的 MCP 传输协议: {server.transport}")
class AgentMcpManager:
"""管理 Agent 外部 MCP 服务器配置、工具发现和工具调用。"""
def get_servers(self) -> list[AgentMcpServerConfig]:
"""读取已保存的外部 MCP 服务器配置。"""
raw_servers = SystemConfigOper().get(SystemConfigKey.AIAgentMcpServers) or []
if not isinstance(raw_servers, list):
return []
servers: list[AgentMcpServerConfig] = []
for raw_server in raw_servers:
try:
servers.append(self.normalize_server(raw_server))
except Exception as err:
logger.warning(f"忽略无效的 Agent MCP 配置: {err}")
return servers
async def save_servers(self, servers: list[AgentMcpServerConfig]) -> bool:
"""保存外部 MCP 服务器配置。"""
normalized_servers = [self.normalize_server(server).model_dump() for server in servers]
return await SystemConfigOper().async_set(
SystemConfigKey.AIAgentMcpServers,
normalized_servers or None,
)
def normalize_server(self, value: Any) -> AgentMcpServerConfig:
"""规范化单个 MCP 服务器配置。"""
if isinstance(value, AgentMcpServerConfig):
raw_server = value.model_dump()
elif isinstance(value, dict):
raw_server = dict(value)
else:
raise ValueError("MCP 服务器配置必须是对象")
raw_server["id"] = str(raw_server.get("id") or uuid.uuid4().hex[:12]).strip()
raw_server["name"] = str(raw_server.get("name") or raw_server["id"]).strip()
raw_server["transport"] = str(raw_server.get("transport") or "stdio").strip()
raw_server["description"] = str(raw_server.get("description") or "").strip() or None
raw_server["command"] = str(raw_server.get("command") or "").strip() or None
raw_server["args"] = [str(item) for item in raw_server.get("args") or []]
raw_server["env"] = _normalize_string_dict(raw_server.get("env"))
raw_server["url"] = str(raw_server.get("url") or "").strip() or None
raw_server["headers"] = _normalize_string_dict(raw_server.get("headers"))
raw_server["timeout"] = _normalize_timeout(raw_server.get("timeout"))
raw_server["tool_prefix"] = str(raw_server.get("tool_prefix") or "").strip() or None
raw_server["require_admin"] = bool(raw_server.get("require_admin", True))
return AgentMcpServerConfig.model_validate(raw_server)
def config_signature(self) -> str:
"""生成外部 MCP 配置签名,用于 Agent 图缓存失效。"""
payload = [server.model_dump() for server in self.get_servers()]
raw_text = json.dumps(payload, ensure_ascii=False, sort_keys=True, default=str)
return hashlib.sha256(raw_text.encode("utf-8")).hexdigest()
async def initialize_session(self, session) -> None:
"""完成 MCP initialize 和 initialized 通知流程。"""
await session.request(
"initialize",
{
"protocolVersion": MCP_PROTOCOL_VERSION,
"capabilities": {},
"clientInfo": {
"name": MCP_CLIENT_NAME,
"version": "1.0.0",
},
},
)
await session.notify("notifications/initialized")
async def list_server_tools(self, server: AgentMcpServerConfig) -> list[AgentMcpToolSpec]:
"""连接单个 MCP 服务器并读取工具列表。"""
normalized_server = self.normalize_server(server)
session_manager = await _open_mcp_session(normalized_server)
async with session_manager as session:
await self.initialize_session(session)
result = await session.request("tools/list")
tools_payload = result.get("tools", []) if isinstance(result, dict) else []
tool_specs: list[AgentMcpToolSpec] = []
for item in tools_payload:
if not isinstance(item, dict) or not item.get("name"):
continue
tool_name = str(item["name"])
tool_specs.append(
AgentMcpToolSpec(
server=normalized_server,
name=tool_name,
agent_tool_name=_build_agent_tool_name(normalized_server, tool_name),
description=str(item.get("description") or ""),
input_schema=_normalize_input_schema(item.get("inputSchema")),
)
)
return tool_specs
async def list_enabled_tool_specs(self) -> list[AgentMcpToolSpec]:
"""读取所有启用 MCP 服务器暴露的工具定义。"""
tool_specs: list[AgentMcpToolSpec] = []
seen_names: set[str] = set()
for server in self.get_servers():
if not server.enabled:
continue
try:
for spec in await self.list_server_tools(server):
if spec.agent_tool_name in seen_names:
logger.warning(f"跳过重复的 MCP Agent 工具名: {spec.agent_tool_name}")
continue
tool_specs.append(spec)
seen_names.add(spec.agent_tool_name)
except Exception as err:
logger.warning(f"读取 MCP 服务器 {server.name} 工具失败: {err}")
return tool_specs
async def call_server_tool(
self,
server: AgentMcpServerConfig,
tool_name: str,
arguments: Optional[dict[str, Any]] = None,
) -> Any:
"""调用单个 MCP 服务器上的指定工具。"""
normalized_server = self.normalize_server(server)
session_manager = await _open_mcp_session(normalized_server)
async with session_manager as session:
await self.initialize_session(session)
return await session.request(
"tools/call",
{
"name": tool_name,
"arguments": arguments or {},
},
)
async def test_server(self, server: AgentMcpServerConfig) -> AgentMcpServerTestResult:
"""测试 MCP 服务器连接并返回工具列表。"""
tool_specs = await self.list_server_tools(server)
tools = [
AgentMcpServerToolInfo(
name=spec.name,
agent_tool_name=spec.agent_tool_name,
description=spec.description,
input_schema=spec.input_schema,
)
for spec in tool_specs
]
return AgentMcpServerTestResult(
success=True,
message=f"连接成功,发现 {len(tools)} 个工具",
tools=tools,
tool_count=len(tools),
)
agent_mcp_manager = AgentMcpManager()
@@ -59,6 +59,10 @@ SYSTEMCONFIG_SETTING_METADATA = {
"group": "ai_agent",
"label": "AI 智能体配置",
},
SystemConfigKey.AIAgentMcpServers.value: {
"group": "ai_agent",
"label": "AI 智能体外部 MCP 服务器",
},
SystemConfigKey.CustomIdentifiers.value: {
"group": "custom_identifiers",
"label": "自定义识别词",
+98
View File
@@ -0,0 +1,98 @@
"""外部 MCP 工具适配器。"""
import json
from typing import Any, Optional
from pydantic import PrivateAttr
from app.agent.mcp import AgentMcpToolSpec, agent_mcp_manager
from app.agent.tools.base import MoviePilotTool
from app.agent.tools.tags import ToolTag
class McpExternalTool(MoviePilotTool):
"""将外部 MCP 工具包装为 MoviePilot Agent 工具。"""
name: str = "mcp_external_tool"
tags: list[str] = [
ToolTag.Read,
ToolTag.Admin,
]
description: str = "Call an external MCP tool configured for MoviePilot Agent."
args_schema: dict[str, Any] = {"type": "object", "properties": {}, "required": []}
require_admin: bool = True
_spec: AgentMcpToolSpec = PrivateAttr()
def __init__(self, spec: AgentMcpToolSpec, session_id: str, user_id: str) -> None:
super().__init__(
session_id=session_id,
user_id=user_id,
name=spec.agent_tool_name,
description=spec.description
or f"Call external MCP tool {spec.name} on {spec.server.name}.",
args_schema=spec.input_schema,
require_admin=spec.server.require_admin,
)
self._spec = spec
def get_tool_message(self, **kwargs) -> Optional[str]:
"""根据 MCP 工具信息生成友好的提示消息。"""
return f"调用 MCP 工具: {self._spec.server.name}/{self._spec.name}"
async def run(self, **kwargs) -> str:
"""
调用外部 MCP 工具。
:param kwargs: 传递给外部 MCP 工具的参数
:return: MCP 工具返回内容
"""
result = await agent_mcp_manager.call_server_tool(
server=self._spec.server,
tool_name=self._spec.name,
arguments=kwargs,
)
return self._format_mcp_result(result)
@staticmethod
def _format_mcp_result(result: Any) -> str:
"""将 MCP tools/call 返回结构转换为 Agent 可读文本。"""
if isinstance(result, dict):
content = result.get("content")
if isinstance(content, list):
parts = []
for item in content:
if not isinstance(item, dict):
continue
if item.get("type") == "text" and item.get("text") is not None:
parts.append(str(item["text"]))
elif item:
parts.append(json.dumps(item, ensure_ascii=False, default=str))
if parts:
return "\n".join(parts)
if result.get("isError"):
return json.dumps(result, ensure_ascii=False, indent=2, default=str)
if isinstance(result, str):
return result
return json.dumps(result, ensure_ascii=False, indent=2, default=str)
async def create_external_mcp_tools(
*,
session_id: str,
user_id: str,
channel: Optional[str] = None,
source: Optional[str] = None,
username: Optional[str] = None,
stream_handler=None,
agent_context: Optional[dict] = None,
) -> list[McpExternalTool]:
"""创建当前已启用的外部 MCP Agent 工具列表。"""
tools = []
for spec in await agent_mcp_manager.list_enabled_tool_specs():
tool = McpExternalTool(spec=spec, session_id=session_id, user_id=user_id)
tool.set_message_attr(channel=channel, source=source, username=username)
tool.set_stream_handler(stream_handler=stream_handler)
tool.set_agent_context(agent_context=agent_context)
tools.append(tool)
return tools
@@ -7,6 +7,7 @@ from pydantic import BaseModel, Field
from app.agent.tools.base import MoviePilotTool
from app.agent.tools.tags import ToolTag
from app.core.metainfo import clear_rust_parse_options_cache
from app.db.systemconfig_oper import SystemConfigOper
from app.log import logger
from app.schemas.types import SystemConfigKey
@@ -85,6 +86,7 @@ class UpdateCustomIdentifiersTool(MoviePilotTool):
SystemConfigKey.CustomIdentifiers, value
)
if success:
clear_rust_parse_options_cache()
return json.dumps(
{
"success": True,
+73
View File
@@ -20,6 +20,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app import schemas
from app.agent import MoviePilotAgent, ReplyMode, StreamingHandler, agent_manager
from app.agent.llm.capability import AgentCapabilityManager
from app.agent.mcp import agent_mcp_manager
from app.chain.message import MessageChain
from app.chain.site import site_interaction_manager
from app.chain.skills import skills_interaction_manager
@@ -56,6 +57,78 @@ _WEB_AGENT_NOTICE_LISTENER_REGISTERED = False
_WEB_AGENT_BACKGROUND_TASKS: set[asyncio.Task] = set()
def _ensure_superuser(user: User) -> None:
"""校验当前用户是否为超级管理员。"""
if not getattr(user, "is_superuser", False):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Forbidden")
@router.get("/mcp/servers", summary="查询 Agent MCP 服务器配置", response_model=schemas.Response)
async def list_agent_mcp_servers(
current_user: User = Depends(get_current_active_user),
) -> schemas.Response:
"""
查询 Agent 外部 MCP 服务器配置。
"""
_ensure_superuser(current_user)
servers = agent_mcp_manager.get_servers()
enabled_count = len([server for server in servers if server.enabled])
return schemas.Response(
success=True,
data={
"servers": [server.model_dump() for server in servers],
"enabled_count": enabled_count,
"total_count": len(servers),
},
)
@router.post("/mcp/servers", summary="保存 Agent MCP 服务器配置", response_model=schemas.Response)
async def save_agent_mcp_servers(
request: schemas.AgentMcpServersSaveRequest,
current_user: User = Depends(get_current_active_user),
) -> schemas.Response:
"""
保存 Agent 外部 MCP 服务器配置。
"""
_ensure_superuser(current_user)
success = await agent_mcp_manager.save_servers(request.servers)
return schemas.Response(
success=success,
message="保存MCP配置成功" if success else "保存MCP配置失败",
)
@router.post("/mcp/servers/test", summary="测试 Agent MCP 服务器", response_model=schemas.Response)
async def test_agent_mcp_server(
request: schemas.AgentMcpServerTestRequest,
current_user: User = Depends(get_current_active_user),
) -> schemas.Response:
"""
测试 Agent 外部 MCP 服务器连接并读取工具列表。
"""
_ensure_superuser(current_user)
try:
result = await agent_mcp_manager.test_server(request.server)
return schemas.Response(
success=result.success,
message=result.message,
data=result.model_dump(),
)
except Exception as err:
logger.warning(f"测试 Agent MCP 服务器失败: {err}")
return schemas.Response(
success=False,
message=f"测试MCP服务器失败: {str(err)}",
data={
"success": False,
"message": str(err),
"tools": [],
"tool_count": 0,
},
)
class _WebAgentStreamingHandler(StreamingHandler):
"""
Web 前端专用流式处理器,将工具提示和文本统一回调给 SSE。
+50
View File
@@ -6,11 +6,61 @@ from app import schemas
from app.chain.douban import DoubanChain
from app.core.context import MediaInfo
from app.core.security import verify_token
from app.db.models.user import User
from app.db.user_oper import get_current_active_superuser_async
from app.modules.douban.douban_cache import DoubanCache
from app.schemas import MediaType
router = APIRouter()
@router.get(
"/cache", summary="查询豆瓣识别缓存", response_model=schemas.Response
)
async def douban_recognition_cache(
_: User = Depends(get_current_active_superuser_async),
) -> schemas.Response:
"""查询可管理的豆瓣识别缓存。"""
cache_items = DoubanCache().list_items()
recognized_count = sum(1 for item in cache_items if item["douban_id"])
return schemas.Response(
success=True,
data={
"count": len(cache_items),
"recognized": recognized_count,
"unrecognized": len(cache_items) - recognized_count,
"data": cache_items,
},
)
@router.delete(
"/cache/{cache_key:path}",
summary="删除指定豆瓣识别缓存",
response_model=schemas.Response,
)
async def delete_douban_recognition_cache(
cache_key: str,
_: User = Depends(get_current_active_superuser_async),
) -> schemas.Response:
"""按缓存键删除单条豆瓣识别缓存。"""
deleted_item = DoubanCache().delete(cache_key)
if not deleted_item:
return schemas.Response(success=False, message="豆瓣识别缓存不存在")
return schemas.Response(success=True, message="豆瓣识别缓存删除成功")
@router.delete(
"/cache", summary="清空豆瓣识别缓存", response_model=schemas.Response
)
async def clear_douban_recognition_cache(
_: User = Depends(get_current_active_superuser_async),
) -> schemas.Response:
"""清空全部豆瓣识别缓存。"""
DoubanCache().clear()
return schemas.Response(success=True, message="豆瓣识别缓存清理完成")
@router.get(
"/person/{person_id}", summary="人物详情", response_model=schemas.MediaPerson
)
+82 -25
View File
@@ -1,3 +1,4 @@
import asyncio
import mimetypes
import shutil
from typing import Annotated, Any, List, Optional
@@ -39,6 +40,67 @@ PROTECTED_ROUTES = {"/api/v1/openapi.json", "/docs", "/docs/oauth2-redirect", "/
PLUGIN_PREFIX = f"{settings.API_V1_STR}/plugin"
router = APIRouter()
_plugin_release_refresh_tasks: set[asyncio.Task] = set()
async def _get_market_plugin_from_repo(
plugin_manager: PluginManager,
plugin_id: str,
repo_url: str,
force: bool,
) -> Optional[schemas.Plugin]:
"""
只读取指定插件仓库的市场元数据,避免单插件详情触发全部市场刷新。
"""
market_plugins = await plugin_manager.async_get_plugins_from_market(
repo_url, settings.VERSION_FLAG, force
)
market_plugin = next(
(
plugin
for plugin in market_plugins or []
if plugin.id == plugin_id
),
None,
)
if market_plugin or not settings.VERSION_FLAG:
return market_plugin
compatible_plugins = await plugin_manager.async_get_plugins_from_market(
repo_url, None, force
)
return next(
(
plugin
for plugin in compatible_plugins or []
if plugin.id == plugin_id
),
None,
)
async def _refresh_plugin_release_versions(plugin_id: str, repo_url: str) -> None:
"""
后台强制刷新 Release 缓存,接口响应路径优先返回已有缓存。
"""
try:
async with async_fresh(True):
await PluginHelper().async_get_plugin_release_versions(plugin_id, repo_url)
except Exception as e:
logger.warning(f"后台刷新插件 {plugin_id} Release 列表失败:{e}")
def _schedule_plugin_release_refresh(plugin_id: str, repo_url: str) -> None:
"""
保留后台任务引用,避免任务被回收,同时让 helper 负责同仓库强刷合并。
"""
task = asyncio.create_task(_refresh_plugin_release_versions(plugin_id, repo_url))
_plugin_release_refresh_tasks.add(task)
def _discard_task(completed_task: asyncio.Task) -> None:
_plugin_release_refresh_tasks.discard(completed_task)
task.add_done_callback(_discard_task)
def register_plugin_api(plugin_id: Optional[str] = None):
@@ -239,6 +301,15 @@ async def _get_plugin_history_detail(
if local_repo_plugin:
return _merge_plugin_market_metadata(installed_plugin, local_repo_plugin)
if installed_plugin.repo_url:
market_plugin = await _get_market_plugin_from_repo(
plugin_manager, plugin_id, installed_plugin.repo_url, force
)
if not market_plugin:
logger.debug(f"插件 {plugin_id} 未从来源仓库获取到更新说明,返回本地插件信息")
return installed_plugin
return _merge_plugin_market_metadata(installed_plugin, market_plugin)
market_plugin = next(
(
plugin
@@ -359,30 +430,9 @@ async def plugin_releases(
}
plugin_manager = PluginManager()
market_plugins = await plugin_manager.async_get_plugins_from_market(
repo_url, settings.VERSION_FLAG, force
market_plugin = await _get_market_plugin_from_repo(
plugin_manager, plugin_id, repo_url, force
)
market_plugin = next(
(
plugin
for plugin in market_plugins or []
if plugin.id == plugin_id
),
None,
)
if not market_plugin and settings.VERSION_FLAG:
compatible_plugins = await plugin_manager.async_get_plugins_from_market(
repo_url, None, force
)
market_plugin = next(
(
plugin
for plugin in compatible_plugins or []
if plugin.id == plugin_id
),
None,
)
latest_version = market_plugin.plugin_version if market_plugin else None
current_version = plugin_manager.get_local_plugin_version(plugin_id)
if not getattr(market_plugin, "release", False):
@@ -393,8 +443,15 @@ async def plugin_releases(
"items": [],
}
async with async_fresh(force):
release_items = await PluginHelper().async_get_plugin_release_versions(plugin_id, repo_url)
plugin_helper = PluginHelper()
has_release_cache = (
await plugin_helper.async_has_plugin_release_cache(repo_url)
if force
else False
)
release_items = await plugin_helper.async_get_plugin_release_versions(plugin_id, repo_url)
if force and has_release_cache:
_schedule_plugin_release_refresh(plugin_id, repo_url)
items = []
for item in release_items:
version = item.get("version")
+36 -14
View File
@@ -1123,33 +1123,64 @@ def ruletest(
"""
过滤规则测试,规则类型 1-订阅,2-洗版,3-搜索
"""
metainfo = MetaInfo(title=title, subtitle=subtitle)
torrent = schemas.TorrentInfo(
title=title,
description=subtitle,
)
# 查询规则组详情
rulegroup = RuleHelper().get_rule_group(rulegroup_name)
result_data = {
"title": title,
"subtitle": subtitle,
"rulegroup_name": rulegroup_name,
"rulegroup": rulegroup.model_dump() if rulegroup else None,
"meta_info": metainfo.to_dict(),
"media_info": None,
"torrent_info": torrent.model_dump(),
"priority": None,
"matched": False,
}
if not rulegroup:
return schemas.Response(
success=False, message=f"过滤规则组 {rulegroup_name} 不存在!"
success=False,
message=f"过滤规则组 {rulegroup_name} 不存在!",
data=result_data,
)
# 根据标题查询媒体信息
media_info = MediaChain().recognize_by_meta(
MetaInfo(title=title, subtitle=subtitle),
metainfo,
obtain_images=False,
)
result_data["media_info"] = media_info.to_dict() if media_info else None
if not media_info:
return schemas.Response(success=False, message="未识别到媒体信息!")
return schemas.Response(
success=False,
message="未识别到媒体信息!",
data=result_data,
)
# 过滤
result = SearchChain().filter_torrents(
rule_groups=[rulegroup.name], torrent_list=[torrent], mediainfo=media_info
)
if not result:
return schemas.Response(success=False, message="不符合过滤规则!")
return schemas.Response(
success=False,
message="不符合过滤规则!",
data=result_data,
)
result_data.update(
{
"matched": True,
"priority": 100 - result[0].pri_order + 1,
"torrent_info": result[0].model_dump(),
}
)
return schemas.Response(
success=True, data={"priority": 100 - result[0].pri_order + 1}
success=True,
data=result_data,
)
@@ -1308,12 +1339,7 @@ def restart_system(_: User = Depends(get_current_active_superuser)):
"""
if not SystemHelper.can_restart():
return schemas.Response(success=False, message="当前运行环境不支持重启操作!")
# 标识停止事件
global_vars.stop_system()
# 执行重启
ret, msg = SystemHelper.restart()
if not ret:
global_vars.resume_system()
return schemas.Response(success=ret, message=msg)
@@ -1331,11 +1357,7 @@ def upgrade_system(
if not SystemHelper.can_restart():
return schemas.Response(success=False, message="当前运行环境不支持升级操作!")
# 标识停止事件
global_vars.stop_system()
ret, msg = SystemHelper.upgrade(mode=mode or "release")
if not ret:
global_vars.resume_system()
return schemas.Response(success=ret, message=msg)
+50
View File
@@ -5,11 +5,61 @@ from fastapi import APIRouter, Depends
from app import schemas
from app.chain.tmdb import TmdbChain
from app.core.security import verify_token
from app.db.models.user import User
from app.db.user_oper import get_current_active_superuser_async
from app.modules.themoviedb.tmdb_cache import TmdbCache
from app.schemas.types import MediaType
router = APIRouter()
@router.get(
"/cache", summary="查询 TheMovieDb 识别缓存", response_model=schemas.Response
)
async def tmdb_recognition_cache(
_: User = Depends(get_current_active_superuser_async),
) -> schemas.Response:
"""查询可管理的 TheMovieDb 识别缓存。"""
cache_items = TmdbCache().list_items()
recognized_count = sum(1 for item in cache_items if item["tmdb_id"])
return schemas.Response(
success=True,
data={
"count": len(cache_items),
"recognized": recognized_count,
"unrecognized": len(cache_items) - recognized_count,
"data": cache_items,
},
)
@router.delete(
"/cache/{cache_key:path}",
summary="删除指定 TheMovieDb 识别缓存",
response_model=schemas.Response,
)
async def delete_tmdb_recognition_cache(
cache_key: str,
_: User = Depends(get_current_active_superuser_async),
) -> schemas.Response:
"""按缓存键删除单条 TheMovieDb 识别缓存。"""
deleted_item = TmdbCache().delete(cache_key)
if not deleted_item:
return schemas.Response(success=False, message="TheMovieDb 识别缓存不存在")
return schemas.Response(success=True, message="TheMovieDb 识别缓存删除成功")
@router.delete(
"/cache", summary="清空 TheMovieDb 识别缓存", response_model=schemas.Response
)
async def clear_tmdb_recognition_cache(
_: User = Depends(get_current_active_superuser_async),
) -> schemas.Response:
"""清空全部 TheMovieDb 识别缓存。"""
TmdbCache().clear()
return schemas.Response(success=True, message="TheMovieDb 识别缓存清理完成")
@router.get(
"/seasons/{tmdbid}", summary="TMDB所有季", response_model=List[schemas.TmdbSeason]
)
+260
View File
@@ -1,11 +1,13 @@
import base64
import copy
import hashlib
import json
import re
import shutil
import time
from pathlib import Path
from typing import List, Optional, Tuple, Set, Dict, Union
from urllib.parse import parse_qs, urlparse
from app import schemas
from app.chain import ChainBase
@@ -16,6 +18,7 @@ from app.core.context import MediaInfo, SubtitleInfo, TorrentInfo, Context
from app.core.event import eventmanager, Event
from app.core.meta import MetaBase
from app.core.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.helper.directory import DirectoryHelper, validate_download_save_path
@@ -31,6 +34,21 @@ from app.utils.string import StringUtils
from app.utils.system import SystemUtils
DOWNLOAD_FAILURE_RESOURCE_TTL_SECONDS = 24 * 60 * 60
DOWNLOAD_FAILURE_TRANSIENT_TTL_SECONDS = 60 * 60
DOWNLOAD_FAILURE_RESOURCE_ERROR_KEYWORDS = (
"无法读取种子文件",
"下载种子内容为空",
"无法获取下载地址",
"种子下载失败",
"torrent not found",
"not found",
"404",
"deleted",
"invalid torrent",
)
class DownloadChain(ChainBase):
"""
下载处理链
@@ -365,6 +383,183 @@ class DownloadChain(ChainBase):
except Exception as err:
logger.error(f"提交下载成功后处理后台任务失败:{str(err)}")
@staticmethod
def _is_subscribe_source(source: Optional[str]) -> bool:
"""
判断下载来源是否为订阅任务。
"""
return bool(source and str(source).startswith("Subscribe|"))
@staticmethod
def _format_failure_episodes(meta: Optional[MetaBase]) -> Optional[str]:
"""
从识别元数据中格式化用于失败记录的集数。
"""
if not meta:
return None
if getattr(meta, "episode", None):
return meta.episode
episode_list = getattr(meta, "episode_list", None)
if episode_list:
return StringUtils.format_ep(list(episode_list))
return None
@staticmethod
def _torrent_resource_key(torrent: Optional[TorrentInfo]) -> str:
"""
生成不保存敏感下载链接的种子资源键。
"""
if not torrent:
return ""
for attr_name in ("torrent_id", "info_hash"):
value = getattr(torrent, attr_name, None)
if value:
return str(value)
for attr_name in ("page_url", "enclosure"):
url = getattr(torrent, attr_name, None)
if not url:
continue
match = re.search(r"\[(.*?)](.*)", str(url))
if match:
url = match.group(2)
parsed = urlparse(str(url))
params = parse_qs(parsed.query)
for param_name in ("id", "torrentid", "torrent_id", "tid", "hash"):
values = params.get(param_name)
if values:
return f"{parsed.netloc}:{param_name}={values[0]}"
if parsed.netloc and parsed.path:
return f"{parsed.netloc}{parsed.path}"
title = getattr(torrent, "title", "") or ""
size = getattr(torrent, "size", "") or ""
return f"title={title}|size={size}"
@classmethod
def _build_download_failure_fingerprint(cls, context: Context) -> Optional[str]:
"""
根据媒体和种子资源信息生成失败冷却指纹。
"""
media = getattr(context, "media_info", None)
torrent = getattr(context, "torrent_info", None)
if not media or not torrent:
return None
media_type = getattr(getattr(media, "type", None), "value", getattr(media, "type", None))
media_key = (
getattr(media, "tmdb_id", None)
or getattr(media, "douban_id", None)
or getattr(media, "imdb_id", None)
or getattr(media, "tvdb_id", None)
or f"{getattr(media, 'title', '')}:{getattr(media, 'year', '')}"
)
meta = getattr(context, "meta_info", None)
site = getattr(torrent, "site", None) or getattr(torrent, "site_name", None)
payload = {
"media_type": str(media_type or ""),
"media_key": str(media_key or ""),
"season": str(getattr(meta, "season", None) or getattr(media, "season", None) or ""),
"episodes": cls._format_failure_episodes(meta) or "",
"site": str(site or ""),
"resource": cls._torrent_resource_key(torrent),
}
if not payload["media_type"] or not payload["media_key"] or not payload["resource"]:
return None
raw_text = json.dumps(payload, ensure_ascii=False, sort_keys=True)
return hashlib.sha256(raw_text.encode("utf-8")).hexdigest()
@staticmethod
def _download_failure_ttl(error_msg: Optional[str]) -> int:
"""
按失败原因确定资源冷却时间。
"""
error_text = str(error_msg or "").lower()
if any(keyword in error_text for keyword in DOWNLOAD_FAILURE_RESOURCE_ERROR_KEYWORDS):
return DOWNLOAD_FAILURE_RESOURCE_TTL_SECONDS
return DOWNLOAD_FAILURE_TRANSIENT_TTL_SECONDS
def _record_download_failure(
self,
context: Context,
error_msg: Optional[str],
downloader: Optional[str] = None,
source: Optional[str] = None,
episodes: Optional[Set[int]] = None,
) -> Optional[str]:
"""
记录资源级下载失败,并返回本次失败指纹。
"""
fingerprint = self._build_download_failure_fingerprint(context)
if not fingerprint:
return None
now_timestamp = time.time()
now_time = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(now_timestamp))
next_retry_at = time.strftime(
"%Y-%m-%d %H:%M:%S",
time.localtime(now_timestamp + self._download_failure_ttl(error_msg)),
)
media = context.media_info
meta = context.meta_info
torrent = context.torrent_info
site = getattr(torrent, "site", None)
try:
DownloadFailureOper().record_failure(
fingerprint=fingerprint,
now_time=now_time,
next_retry_at=next_retry_at,
type=getattr(getattr(media, "type", None), "value", getattr(media, "type", None)),
title=getattr(media, "title", None),
year=getattr(media, "year", None),
tmdbid=getattr(media, "tmdb_id", None),
doubanid=getattr(media, "douban_id", None),
seasons=getattr(meta, "season", None),
episodes=StringUtils.format_ep(list(episodes)) if episodes else self._format_failure_episodes(meta),
site=site if isinstance(site, int) else None,
site_name=getattr(torrent, "site_name", None),
torrent_id=self._torrent_resource_key(torrent),
torrent_name=getattr(torrent, "title", None),
torrent_size=getattr(torrent, "size", None),
downloader=downloader,
source=str(source)[:1000] if source else None,
error_message=str(error_msg or "")[:1000],
)
except Exception as err:
logger.error(f"记录下载失败冷却失败:{str(err)}")
return fingerprint
def _active_download_failure_fingerprints(
self,
contexts: List[Context],
source: Optional[str],
) -> Set[str]:
"""
查询当前订阅候选中仍处于冷却期的失败指纹。
"""
if not self._is_subscribe_source(source):
return set()
fingerprints = [
fingerprint
for fingerprint in [
self._build_download_failure_fingerprint(context)
for context in contexts or []
]
if fingerprint
]
if not fingerprints:
return set()
now_time = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime())
try:
return set(
DownloadFailureOper()
.get_active_by_fingerprints(fingerprints=fingerprints, now_time=now_time)
.keys()
)
except Exception as err:
logger.error(f"查询下载失败冷却失败:{str(err)}")
return set()
def download_torrent(self, torrent: TorrentInfo,
channel: MessageChannel = None,
source: Optional[str] = None,
@@ -578,6 +773,13 @@ class DownloadChain(ChainBase):
torrent_content = cache_backend.get(torrent_file.as_posix(), region="torrents")
if not torrent_content:
self._record_download_failure(
context=context,
error_msg="下载种子内容为空",
downloader=downloader or _site_downloader,
source=source,
episodes=episodes,
)
return (None, "下载种子内容为空") if return_detail else None
# 获取种子文件的文件夹名和文件清单
@@ -732,6 +934,13 @@ class DownloadChain(ChainBase):
# 下载失败
logger.error(f"{_media.title_year} 添加下载任务失败:"
f"{_torrent.title} - {_torrent.enclosure}{error_msg}")
self._record_download_failure(
context=context,
error_msg=error_msg,
downloader=_downloader or downloader or _site_downloader,
source=source,
episodes=episodes,
)
# 只发送给对应渠道和用户
self.post_message(Notification(
channel=channel,
@@ -897,6 +1106,28 @@ class DownloadChain(ChainBase):
# 仅排序,不提前按媒体控重;下载失败时需要继续尝试同组后续候选。
contexts = TorrentHelper().sort_torrents(contexts)
active_failure_fingerprints = self._active_download_failure_fingerprints(
contexts=contexts,
source=source,
)
def __is_context_in_failure_cooldown(_context: Context) -> bool:
"""
判断候选资源是否仍处于失败冷却期。
"""
fingerprint = self._build_download_failure_fingerprint(_context)
if fingerprint and fingerprint in active_failure_fingerprints:
logger.info(f"{_context.torrent_info.title} 近期添加下载失败,暂时跳过该资源")
return True
return False
def __remember_context_failure(_context: Context) -> None:
"""
将本轮失败候选加入内存冷却集合,避免同一批次重复尝试。
"""
fingerprint = self._build_download_failure_fingerprint(_context)
if fingerprint:
active_failure_fingerprints.add(fingerprint)
# 如果是电影,直接下载
downloaded_movies = set()
@@ -904,6 +1135,8 @@ class DownloadChain(ChainBase):
if global_vars.is_system_stopped:
break
if context.media_info.type == MediaType.MOVIE:
if __is_context_in_failure_cooldown(context):
continue
movie_key = __get_movie_download_key(context)
if movie_key in downloaded_movies:
continue
@@ -915,6 +1148,8 @@ class DownloadChain(ChainBase):
logger.info(f"{context.torrent_info.title} 添加下载成功")
downloaded_list.append(context)
downloaded_movies.add(movie_key)
else:
__remember_context_failure(context)
# 电视剧整季匹配
if no_exists:
@@ -959,6 +1194,8 @@ class DownloadChain(ChainBase):
# 不重复添加
if context in downloaded_list:
continue
if __is_context_in_failure_cooldown(context):
continue
# 种子季是需要季或者子集
if set(torrent_season).issubset(set(need_season)):
complete_coverage_matched = False
@@ -968,6 +1205,13 @@ class DownloadChain(ChainBase):
content, _, torrent_files = self.download_torrent(torrent)
if not content:
logger.warn(f"{torrent.title} 种子下载失败!")
self._record_download_failure(
context=context,
error_msg="下载种子内容为空",
downloader=downloader,
source=source,
)
__remember_context_failure(context)
continue
if isinstance(content, str):
logger.warn(f"{meta.org_string} 下载地址是磁力链,无法确定种子文件集数")
@@ -1039,6 +1283,8 @@ class DownloadChain(ChainBase):
if not need_season:
# 全部下载完成
break
else:
__remember_context_failure(context)
# 电视剧季内的集匹配
if no_exists:
logger.info(f"开始电视剧完整集匹配:{no_exists}")
@@ -1079,6 +1325,8 @@ class DownloadChain(ChainBase):
# 不重复添加
if context in downloaded_list:
continue
if __is_context_in_failure_cooldown(context):
continue
# 种子季
torrent_season = meta.season_list
# 只处理单季含集的种子
@@ -1121,6 +1369,8 @@ class DownloadChain(ChainBase):
_sea=need_season,
_current=torrent_episodes)
logger.info(f"{need_season} 剩余需要集:{need_episodes}")
else:
__remember_context_failure(context)
# 仍然缺失的剧集,从整季中选择需要的集数文件下载,仅支持QB和TR
if no_exists:
@@ -1163,6 +1413,8 @@ class DownloadChain(ChainBase):
# 不重复添加
if context in downloaded_list:
continue
if __is_context_in_failure_cooldown(context):
continue
# 没有需要集后退出
if not need_episodes:
break
@@ -1181,6 +1433,13 @@ class DownloadChain(ChainBase):
content, _, torrent_files = self.download_torrent(torrent)
if not content:
logger.info(f"{torrent.title} 种子下载失败!")
self._record_download_failure(
context=context,
error_msg="下载种子内容为空",
downloader=downloader,
source=source,
)
__remember_context_failure(context)
continue
if isinstance(content, str):
logger.warn(f"{meta.org_string} 下载地址是磁力链,无法解析种子文件集数")
@@ -1209,6 +1468,7 @@ class DownloadChain(ChainBase):
custom_words=custom_words
)
if not download_id:
__remember_context_failure(context)
continue
# 下载成功
logger.info(f"{torrent.title} 添加下载成功")
+21 -3
View File
@@ -141,9 +141,9 @@ class MessageChain(ChainBase):
logger.debug(f"未识别到消息内容::{body}{form}{args}")
return
# 获取原消息ID信息
original_message_id = info.message_id
original_chat_id = info.chat_id
reply_to_message_id = info.reply_to_message_id
# 处理消息
self.handle_message(
@@ -154,6 +154,7 @@ class MessageChain(ChainBase):
text=text,
original_message_id=original_message_id,
original_chat_id=original_chat_id,
reply_to_message_id=reply_to_message_id,
images=images,
audio_refs=audio_refs,
files=files,
@@ -171,6 +172,7 @@ class MessageChain(ChainBase):
images: Optional[List[CommingMessage.MessageImage]] = None,
audio_refs: Optional[List[str]] = None,
files: Optional[List[CommingMessage.MessageAttachment]] = None,
reply_to_message_id: Optional[Union[str, int]] = None,
) -> None:
"""
识别消息内容执行操作
@@ -213,6 +215,7 @@ class MessageChain(ChainBase):
username=username,
text=text,
original_chat_id=original_chat_id,
reply_to_message_id=reply_to_message_id,
images=images,
audio_refs=audio_refs,
files=files,
@@ -255,6 +258,7 @@ class MessageChain(ChainBase):
text=text,
original_message_id=original_message_id,
original_chat_id=original_chat_id,
reply_to_message_id=reply_to_message_id,
images=images,
audio_refs=audio_refs,
files=files,
@@ -286,6 +290,7 @@ class MessageChain(ChainBase):
files: Optional[List[CommingMessage.MessageAttachment]] = None,
has_audio_input: bool = False,
processing_status: Optional[_ProcessingStatus] = None,
reply_to_message_id: Optional[Union[str, int]] = None,
) -> bool:
"""执行实际消息路由,便于统一包裹处理中状态。"""
@@ -316,6 +321,7 @@ class MessageChain(ChainBase):
username=username,
text=text,
original_chat_id=original_chat_id,
reply_to_message_id=reply_to_message_id,
images=images,
audio_refs=audio_refs,
files=files,
@@ -444,6 +450,8 @@ class MessageChain(ChainBase):
"userid": userid,
"channel": channel,
"source": source,
"chat_id": original_chat_id,
"reply_to_message_id": reply_to_message_id,
},
)
return False
@@ -460,6 +468,7 @@ class MessageChain(ChainBase):
audio_refs: Optional[List[str]] = None,
files: Optional[List[CommingMessage.MessageAttachment]] = None,
has_audio_input: bool = False,
reply_to_message_id: Optional[Union[str, int]] = None,
) -> bool:
"""
将插件输入会话中的下一条普通文本派发给指定插件
@@ -469,8 +478,14 @@ class MessageChain(ChainBase):
if text.startswith("CALLBACK:"):
return False
is_cancel_text = text.strip().lower() in {"取消", "退出", "q", "quit", "exit"}
request, status = plugin_input_interaction_manager.consume_by_user(
userid, channel, source, original_chat_id
userid,
channel,
source,
original_chat_id,
reply_to_message_id=reply_to_message_id,
bypass_reply_check=is_cancel_text,
)
if not request:
return False
@@ -487,6 +502,7 @@ class MessageChain(ChainBase):
"source": source,
"username": username,
"chat_id": original_chat_id,
"reply_to_message_id": reply_to_message_id,
"prompt_id": request.prompt_id,
"input_session_id": request.request_id,
"expired": True,
@@ -505,7 +521,7 @@ class MessageChain(ChainBase):
)
return not text.strip().startswith("/")
if text.strip().lower() in {"取消", "退出", "q", "quit", "exit"}:
if is_cancel_text:
self.eventmanager.send_event(
EventType.MessageAction,
{
@@ -517,6 +533,7 @@ class MessageChain(ChainBase):
"source": source,
"username": username,
"chat_id": original_chat_id,
"reply_to_message_id": reply_to_message_id,
"prompt_id": request.prompt_id,
"input_session_id": request.request_id,
"cancelled": True,
@@ -547,6 +564,7 @@ class MessageChain(ChainBase):
"source": source,
"username": username,
"chat_id": original_chat_id,
"reply_to_message_id": reply_to_message_id,
"prompt_id": request.prompt_id,
"input_session_id": request.request_id,
"payload": request.payload,
+89 -2
View File
@@ -1993,6 +1993,14 @@ class SubscribeChain(ChainBase):
tmdbid=subscribe.tmdbid, doubanid=subscribe.doubanid,
subscribe_id=subscribe.id, scene="refresh")
old_total_episode = subscribe.total_episode or 0
if total_episode and total_episode < old_total_episode:
total_episode = self.__resolve_total_episode_decrease(
subscribe=subscribe,
candidate_total=total_episode,
meta=meta,
mediainfo=mediainfo,
mediakey=subscribe.tmdbid or subscribe.doubanid,
)
if total_episode and total_episode != old_total_episode:
progress_update = self.__prepare_total_episode_change_fields(
subscribe=subscribe,
@@ -3655,7 +3663,12 @@ class SubscribeChain(ChainBase):
- exist_flag (bool): 布尔值表示媒体是否已经完全下载或已存在
- no_exists (dict): 缺失的媒体信息包含缺失的集数或其他相关信息
"""
self.__refresh_total_episode_before_completion(subscribe=subscribe, mediainfo=mediainfo)
self.__refresh_total_episode_before_completion(
subscribe=subscribe,
mediainfo=mediainfo,
meta=meta,
mediakey=mediakey,
)
exist_flag, no_exists = self.resolve_subscribe_missing(
subscribe=subscribe,
@@ -3759,6 +3772,66 @@ class SubscribeChain(ChainBase):
return bool(downloaded), no_exists
return False, no_exists
def __resolve_total_episode_decrease(
self,
subscribe: Subscribe,
candidate_total: int,
meta: MetaBase,
mediainfo: MediaInfo,
mediakey: Optional[Union[str, int]] = None,
) -> int:
"""以旧目标范围内已确认存在的最高集号限制总集数回落。"""
old_total = subscribe.total_episode or 0
if candidate_total >= old_total or not old_total:
return candidate_total
if subscribe.type != MediaType.TV.value or self.__is_full_best_version_enabled(subscribe):
return candidate_total
target_key = mediakey or subscribe.tmdbid or subscribe.doubanid
target_season = subscribe.season
target_start = subscribe.start_episode or 1
snapshot = copy.copy(subscribe)
snapshot.total_episode = old_total
try:
satisfied, no_exists = self.resolve_subscribe_missing(
subscribe=snapshot,
meta=meta,
mediainfo=mediainfo,
mediakey=target_key,
best_version_accept_downloaded=bool(subscribe.best_version),
)
except Exception as err:
logger.warning(f"订阅 {subscribe.name} 已存在分集事实查询失败,按元数据总集数继续:{err}")
return candidate_total
if satisfied:
return old_total
if not isinstance(no_exists, dict):
return candidate_total
seasons = no_exists.get(target_key)
if not isinstance(seasons, dict):
return candidate_total
missing_info = seasons.get(target_season)
if not missing_info:
return candidate_total
try:
scope_matches = missing_info.season == target_season \
and missing_info.start_episode == target_start \
and missing_info.total_episode == old_total
episodes = missing_info.episodes
except AttributeError:
return candidate_total
if not scope_matches:
return candidate_total
if not isinstance(episodes, list) or not episodes:
return candidate_total
if any(isinstance(episode, bool) or not isinstance(episode, int)
or episode < target_start or episode > old_total for episode in episodes):
return candidate_total
confirmed = set(range(target_start, old_total + 1)).difference(episodes)
return max(candidate_total, max(confirmed) if confirmed else 0)
@staticmethod
def __resolve_effective_total_episode(subscribe: Subscribe, mediainfo: MediaInfo) -> int:
"""
@@ -3828,7 +3901,13 @@ class SubscribeChain(ChainBase):
return result.total_episode
return current_total
def __refresh_total_episode_before_completion(self, subscribe: Subscribe, mediainfo: MediaInfo):
def __refresh_total_episode_before_completion(
self,
subscribe: Subscribe,
mediainfo: MediaInfo,
meta: Optional[MetaBase] = None,
mediakey: Optional[Union[str, int]] = None,
) -> None:
"""
在完成判断前按最新识别结果兜底修正订阅总集数防止旧总集数导致误完成
"""
@@ -3846,6 +3925,14 @@ class SubscribeChain(ChainBase):
tmdbid=subscribe.tmdbid, doubanid=subscribe.doubanid,
subscribe_id=subscribe.id, scene="precheck")
old_total_episode = subscribe.total_episode or 0
if meta is not None and new_total_episode and new_total_episode < old_total_episode:
new_total_episode = self.__resolve_total_episode_decrease(
subscribe=subscribe,
candidate_total=new_total_episode,
meta=meta,
mediainfo=mediainfo,
mediakey=mediakey,
)
if not new_total_episode or new_total_episode == old_total_episode:
return
-4
View File
@@ -38,8 +38,6 @@ class SystemChain(ChainBase):
"""
重启系统
"""
from app.core.config import global_vars
if channel and userid:
self.post_message(Notification(
channel=channel,
@@ -54,8 +52,6 @@ class SystemChain(ChainBase):
}, self._restart_file)
# 主动备份一次插件
self.backup_plugins()
# 设置停止标志,通知所有模块准备停止
global_vars.stop_system()
# 重启
SystemHelper.restart()
+35 -5
View File
@@ -1240,14 +1240,44 @@ class TransferChain(ChainBase, ConfigReloadMixin, metaclass=Singleton):
history_exists: bool = True,
):
"""
当同一种子的任务都已结束时回写下载器已整理标签
当同一种子的任务都已结束且种子已完成下载回写下载器已整理标签
"""
if (
history_exists
and download_hash
and self.jobview.is_torrent_done(download_hash)
not history_exists
or not download_hash
or not self.jobview.is_torrent_done(download_hash)
):
self.transfer_completed(hashs=download_hash, downloader=downloader)
return
# 作业视图只包含已登记的整理任务;多集种子部分文件先下载完成时,
# 剩余文件尚未产生任务,此时打已整理标签会使下载器轮询永久跳过
# 剩余文件(#6009),因此必须确认种子已整体下载完成。
if not self.__is_torrent_download_completed(download_hash, downloader):
logger.debug(
f"种子 {download_hash} 尚未下载完成或状态未知,暂不设置已整理标签"
)
return
if not self.jobview.is_torrent_done(download_hash):
logger.debug(
f"种子 {download_hash} 存在新登记的整理任务,暂不设置已整理标签"
)
return
self.transfer_completed(hashs=download_hash, downloader=downloader)
def __is_torrent_download_completed(
self, download_hash: str, downloader: Optional[str]
) -> bool:
"""
检查种子在下载器中是否已完成下载查询不到或查询失败时视为未完成
留待下载器定时轮询兜底避免误打已整理标签
"""
try:
torrents = self.list_torrents(hashs=download_hash, downloader=downloader)
if not torrents:
return False
return all((torrent.progress or 0) >= 100 for torrent in torrents)
except Exception as e:
logger.error(f"检查种子 {download_hash} 下载进度失败:{e}")
return False
def __send_metadata_scrape_event(
self, task: TransferTask, transferinfo: TransferInfo
+22
View File
@@ -1214,8 +1214,19 @@ def cached(region: Optional[str] = None, maxsize: Optional[int] = 1024, ttl: Opt
"""
await cache_backend.clear(region=cache_region)
async def cache_exists(*args, **kwargs) -> bool:
"""
判断当前参数对应的有效缓存是否存在
"""
cache_key = __get_cache_key(args, kwargs)
cached_value = await cache_backend.get(cache_key, region=cache_region)
return should_cache(cached_value) and await async_is_valid_cache_value(
cache_key, cached_value, cache_region
)
async_wrapper.cache_region = cache_region
async_wrapper.cache_clear = cache_clear
async_wrapper.cache_exists = cache_exists
return async_wrapper
else:
# 同步函数使用同步缓存后端
@@ -1246,8 +1257,19 @@ def cached(region: Optional[str] = None, maxsize: Optional[int] = 1024, ttl: Opt
"""
cache_backend.clear(region=cache_region)
def cache_exists(*args, **kwargs) -> bool:
"""
判断当前参数对应的有效缓存是否存在
"""
cache_key = __get_cache_key(args, kwargs)
cached_value = cache_backend.get(cache_key, region=cache_region)
return should_cache(cached_value) and is_valid_cache_value(
cache_key, cached_value, cache_region
)
wrapper.cache_region = cache_region
wrapper.cache_clear = cache_clear
wrapper.cache_exists = cache_exists
return wrapper
return decorator
-6
View File
@@ -1211,12 +1211,6 @@ class GlobalVar(object):
"""
self.STOP_EVENT.set()
def resume_system(self):
"""
恢复系统运行标记
"""
self.STOP_EVENT.clear()
@property
def is_system_stopped(self):
"""
+37
View File
@@ -24,6 +24,13 @@ SUBTITLE_EPISODE_ALL_RE = re.compile(
r"([0-9一二三四五六七八九十百零]+)\s*集\s*全|[全共]\s*([0-9一二三四五六七八九十百零]+)\s*[集话話期幕]",
re.IGNORECASE,
)
# 结尾分支显式区分有无右方括号,避免可选括号回溯后绕过数字后缀边界
SUBTITLE_EPISODE_RANGE_FIN_RE = re.compile(
r"(?<!\d)\[?\s*(\d{1,4})\s*-\s*(\d{1,4})\s*"
r"(?:(?:Fin|End)(?![a-z0-9])|完结(?![\u4e00-\u9fff]))"
r"(?:\s*\](?!\d)|(?!\s*(?:\]\d|\d))\s*)",
re.IGNORECASE,
)
VIDEO_BIT_RE = re.compile(
r"(?<![A-Za-z0-9])(?P<bit>8|10|12|16)[\s._-]*bits?(?![A-Za-z0-9])",
re.IGNORECASE,
@@ -292,6 +299,36 @@ class MetaBase(object):
self.type = MediaType.TV
self._subtitle_flag = True
return
# 01-26Fin 等数字范围+完结标记
self.__init_episode_range_fin(title_text)
else:
# 副标题无中文季集标记时,仍识别 01-26Fin 等数字范围+完结标记
self.__init_episode_range_fin(title_text)
def __init_episode_range_fin(self, title_text: str):
"""
识别 01-26Fin / [01-38 END] "数字范围+完结标记"格式的集数信息
"""
episode_range_str = SUBTITLE_EPISODE_RANGE_FIN_RE.search(title_text)
if not episode_range_str:
return
try:
begin_episode = int(episode_range_str.group(1))
end_episode = int(episode_range_str.group(2))
except Exception as err:
logger.debug(f'识别集失败:{str(err)} - {traceback.format_exc()}')
return
if begin_episode < 1 or begin_episode > end_episode or end_episode >= 10000:
return
# 两个数字都落在常见年份区间时视为年份范围而非集数(如 2019-2020完结)
if begin_episode >= 1900 and end_episode <= 2155:
return
if self.begin_episode is None:
self.begin_episode = begin_episode
self.end_episode = end_episode
self.total_episode = end_episode
self.type = MediaType.TV
self._subtitle_flag = True
@property
def season(self) -> str:
+5 -6
View File
@@ -58,12 +58,11 @@ class ModuleManager(metaclass=Singleton):
"""
logger.info("正在停止所有模块...")
for module_id, module in self._running_modules.items():
if hasattr(module, "stop"):
try:
module.stop()
logger.debug(f"Moudle Stoped{module_id}")
except Exception as err:
logger.error(f"Stop Moudle Error{module_id}{str(err)} - {traceback.format_exc()}", exc_info=True)
try:
module.stop()
logger.debug(f"Moudle Stoped{module_id}")
except Exception as err:
logger.error(f"Stop Moudle Error{module_id}{str(err)} - {traceback.format_exc()}", exc_info=True)
logger.info("所有模块停止完成")
def reload(self):
+85 -5
View File
@@ -363,6 +363,20 @@ class PluginManager(ConfigReloadMixin, metaclass=Singleton):
logger.warn(f"检测到本地插件 {candidate.get('id')} 依赖文件变化,请重新安装本地插件以安装依赖")
continue
federated_change = self._get_federated_plugin_change(event_path)
if federated_change:
pid, candidate, remote_entry_ready = federated_change
# 运行目录由构建方直接写入;外部本地仓库只在入口完整时同步运行副本。
if candidate and remote_entry_ready:
if candidate.get("compatible") is False:
logger.info(
f"检测到本地插件 {pid} 联邦构建产物变化,"
f"但跳过同步:{candidate.get('skip_reason')}"
)
elif pid not in local_plugins_to_sync:
local_plugins_to_sync[pid] = (candidate, event_path, False)
continue
# 跳过非 .py 文件
if not event_path.name.endswith(".py"):
continue
@@ -385,13 +399,14 @@ class PluginManager(ConfigReloadMixin, metaclass=Singleton):
f"文件:{event_path},但跳过同步:{local_candidate.get('skip_reason')}"
)
continue
local_plugins_to_sync[local_candidate.get("id")] = (local_candidate, event_path)
local_plugins_to_sync[local_candidate.get("id")] = (local_candidate, event_path, True)
for pid, (candidate, event_path) in local_plugins_to_sync.items():
for pid, (candidate, event_path, should_reload) in local_plugins_to_sync.items():
package_version = candidate.get("package_version")
source_root = f"plugins.{package_version}" if package_version else "plugins"
logger.info(f"检测到本地插件 {pid} 文件变化,来源:{source_root},文件:{event_path}")
if self._sync_local_plugin_if_installed(pid, candidate):
change_name = "Python 文件" if should_reload else "联邦构建产物"
logger.info(f"检测到本地插件 {pid} {change_name}变化,来源:{source_root},文件:{event_path}")
if self._sync_local_plugin_if_installed(pid, candidate) and should_reload:
plugins_to_reload.add(pid)
# 触发重载
@@ -403,6 +418,71 @@ class PluginManager(ConfigReloadMixin, metaclass=Singleton):
except Exception as e:
logger.error(f"插件 {pid} 热重载失败: {e}", exc_info=True)
def _get_federated_plugin_change(
self,
event_path: Path,
) -> Optional[Tuple[str, Optional[dict], bool]]:
"""
识别运行态 Vue 插件声明目录内的构建产物变化
:return: 插件 ID本地仓库候选和联邦入口是否完整非联邦目录变化返回 None
"""
try:
event_path = event_path.resolve()
candidate = self._get_local_plugin_candidate_from_path(event_path)
if candidate:
pid = candidate.get("id")
plugin_dir = Path(candidate.get("path")).resolve()
else:
runtime_root = (settings.ROOT_PATH / "app" / "plugins").resolve()
if not event_path.is_relative_to(runtime_root):
return None
relative_parts = event_path.relative_to(runtime_root).parts
if not relative_parts:
return None
plugin_dir = runtime_root / relative_parts[0]
pid = next(
(
plugin_id
for plugin_id in self._running_plugins
if plugin_id.lower() == relative_parts[0].lower()
),
None,
)
if not pid:
return None
plugin = self._running_plugins.get(pid)
if not plugin:
return None
render_mode, dist_path = plugin.get_render_mode()
if render_mode != "vue" or not isinstance(dist_path, str) or not dist_path:
return None
relative_dist_path = Path(dist_path)
if relative_dist_path.is_absolute() or ".." in relative_dist_path.parts or "\\" in dist_path:
return None
plugin_dir = plugin_dir.resolve()
dist_dir = (plugin_dir / relative_dist_path).resolve()
if (
dist_dir == plugin_dir
or not dist_dir.is_relative_to(plugin_dir)
or not event_path.is_relative_to(dist_dir)
):
return None
remote_entry = dist_dir / "remoteEntry.js"
remote_entry_ready = (
remote_entry.is_file()
and remote_entry.resolve().is_relative_to(plugin_dir)
)
return pid, candidate, remote_entry_ready
except Exception as e:
logger.error(f"识别插件联邦构建产物变化时出错: {e}")
return None
@staticmethod
def _get_plugin_id_from_path(event_path: Path) -> Optional[str]:
"""
@@ -517,7 +597,7 @@ class PluginManager(ConfigReloadMixin, metaclass=Singleton):
source_dir,
dest_dir,
dirs_exist_ok=True,
ignore=shutil.ignore_patterns("__pycache__", "*.pyc", ".DS_Store")
ignore=shutil.ignore_patterns("__pycache__", "*.pyc", ".DS_Store", "node_modules")
)
PluginManager()._recent_local_sync[pid] = time.time()
logger.info(f"已同步本地插件 {pid}{source_dir} -> {dest_dir}")
+53 -1
View File
@@ -1,12 +1,60 @@
import asyncio
from typing import Any, Generator, List, Optional, Self, Tuple, AsyncGenerator, Union
from sqlalchemy import NullPool, QueuePool, and_, create_engine, inspect, text, select, delete, Column, Integer, \
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.core.config import settings
from app.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)
def get_id_column():
@@ -71,6 +119,7 @@ def _get_sqlite_engine(is_async: bool = False):
# 创建数据库引擎
engine = create_engine(**_db_kwargs)
_register_database_error_logging(engine)
# 设置WAL模式
_journal_mode = "WAL" if settings.DB_WAL_ENABLE else "DELETE"
@@ -91,6 +140,7 @@ def _get_sqlite_engine(is_async: bool = False):
}
# 创建异步数据库引擎
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"
@@ -146,6 +196,7 @@ def _get_postgresql_engine(is_async: bool = False):
# 创建数据库引擎
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
@@ -163,6 +214,7 @@ def _get_postgresql_engine(is_async: bool = False):
}
# 创建异步数据库引擎
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
+61
View File
@@ -0,0 +1,61 @@
from typing import Dict, List, Optional
from app.db import DbOper
from app.db.models.downloadfailure import DownloadFailure
class DownloadFailureOper(DbOper):
"""
下载失败冷却记录管理
"""
def get_active_by_fingerprints(
self,
fingerprints: List[str],
now_time: str,
) -> Dict[str, DownloadFailure]:
"""
批量按指纹查询仍在冷却期的失败记录
"""
failures = DownloadFailure.get_active_by_fingerprints(
self._db,
fingerprints=fingerprints,
now_time=now_time,
)
return {
failure.fingerprint: failure
for failure in failures
if failure and failure.fingerprint
}
def record_failure(
self,
fingerprint: str,
now_time: str,
next_retry_at: str,
**kwargs: object,
) -> DownloadFailure:
"""
新增或更新资源失败记录
"""
return DownloadFailure.record_failure(
self._db,
fingerprint=fingerprint,
now_time=now_time,
next_retry_at=next_retry_at,
**kwargs,
)
def delete_expired(
self,
before_time: str,
limit: Optional[int] = 500,
) -> int:
"""
删除已过期较久的失败记录
"""
return DownloadFailure.delete_expired(
self._db,
before_time=before_time,
limit=limit,
)
+1
View File
@@ -1,4 +1,5 @@
from .agentchat import AgentChat
from .downloadfailure import DownloadFailure
from .downloadhistory import DownloadHistory, DownloadFiles
from .mediaserver import MediaServerItem
from .message import Message
+137
View File
@@ -0,0 +1,137 @@
from typing import List, Optional
from sqlalchemy import Column, Float, Index, Integer, String
from sqlalchemy.orm import Session
from app.db import Base, db_query, db_update, get_id_column
class DownloadFailure(Base):
"""
下载失败冷却记录
"""
id = get_id_column()
# 资源失败指纹
fingerprint = Column(String, nullable=False)
# 类型 电影/电视剧
type = Column(String)
# 标题
title = Column(String)
# 年份
year = Column(String)
# TMDBID
tmdbid = Column(Integer)
# 豆瓣ID
doubanid = Column(String)
# Sxx
seasons = Column(String)
# Exx
episodes = Column(String)
# 站点ID
site = Column(Integer)
# 站点名称
site_name = Column(String)
# 种子资源键
torrent_id = Column(String)
# 种子名称
torrent_name = Column(String)
# 种子大小
torrent_size = Column(Float)
# 下载器
downloader = Column(String)
# 下载来源
source = Column(String)
# 失败原因
error_message = Column(String)
# 重试次数
retry_count = Column(Integer, default=0)
# 首次失败时间
first_failed_at = Column(String)
# 最近失败时间
last_failed_at = Column(String)
# 下次允许重试时间
next_retry_at = Column(String)
__table_args__ = (
Index("ux_downloadfailure_fingerprint", "fingerprint", unique=True),
Index("ix_downloadfailure_next_retry_at", "next_retry_at"),
Index("ix_downloadfailure_media_site", "type", "tmdbid", "doubanid", "site"),
)
@classmethod
@db_query
def get_active_by_fingerprints(
cls,
db: Session,
fingerprints: List[str],
now_time: str,
) -> List["DownloadFailure"]:
"""
按指纹批量查询仍处于冷却期的失败记录
"""
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()
)
@classmethod
@db_update
def record_failure(
cls,
db: Session,
fingerprint: str,
now_time: str,
next_retry_at: str,
**kwargs: object,
) -> "DownloadFailure":
"""
新增或更新资源失败记录
"""
failure = db.query(cls).filter(cls.fingerprint == fingerprint).first()
payload = {
**kwargs,
"fingerprint": fingerprint,
"last_failed_at": now_time,
"next_retry_at": next_retry_at,
}
if failure:
payload["retry_count"] = (failure.retry_count or 0) + 1
for key, value in payload.items():
setattr(failure, key, value)
return failure
failure = cls(
**payload,
retry_count=1,
first_failed_at=now_time,
)
db.add(failure)
return failure
@classmethod
@db_update
def delete_expired(
cls,
db: Session,
before_time: str,
limit: Optional[int] = 500,
) -> int:
"""
分批清理已过期较久的失败冷却记录
"""
ids = [
row[0]
for row in db.query(cls.id)
.filter(cls.next_retry_at < before_time)
.order_by(cls.id.asc())
.limit(limit)
.all()
]
if not ids:
return 0
return db.query(cls).filter(cls.id.in_(ids)).delete(synchronize_session=False)
+45 -4
View File
@@ -1,4 +1,5 @@
import base64
import time
from typing import Tuple, Optional
from lxml import etree
@@ -57,6 +58,36 @@ class CookieHelper:
]
}
@staticmethod
def get_page_content(page: BrowserPage, retries: int = 3, interval: float = 1.0) -> Optional[str]:
"""
获取页面源码页面跳转中如登录前后的重定向会导致 page.content() 抛出
"Unable to retrieve content because the page is navigating" 异常等待加载完成后重试
:param page: 浏览器页面
:param retries: 最大重试次数
:param interval: 重试间隔
:return: 页面源码
"""
for i in range(retries):
# 等待加载失败不代表源码不可读取,最后一次等待失败时仍尝试直接获取源码
try:
page.wait_for_load_state("domcontentloaded", timeout=10 * 1000)
except Exception as e:
if i < retries - 1:
logger.warning(f"等待页面加载完成失败:{str(e)}{interval}秒后重试 ({i + 1}/{retries - 1})")
time.sleep(interval)
continue
logger.warning(f"等待页面加载完成失败:{str(e)},尝试直接获取源码")
try:
return page.content()
except Exception as e:
if i >= retries - 1:
logger.error(f"获取页面源码失败:{str(e)}")
return None
logger.warning(f"获取页面源码失败:{str(e)}{interval}秒后重试 ({i + 1}/{retries - 1})")
time.sleep(interval)
return None
@staticmethod
def parse_cookies(cookies: list) -> str:
"""
@@ -93,11 +124,13 @@ class CookieHelper:
:return: Cookie和UA
"""
# 登录页面代码
html_text = page.content()
html_text = self.get_page_content(page)
if not html_text:
return None, None, "获取源码失败"
# 查找用户名输入框
html = etree.HTML(html_text)
if html is None:
return None, None, "解析网页源码失败"
try:
username_xpath = None
for xpath in self._SITE_LOGIN_XPATH.get("username"):
@@ -189,7 +222,12 @@ class CookieHelper:
if "verify" in page.url:
if not otp_code:
return None, None, "需要二次验证码"
html = etree.HTML(page.content())
html_text = self.get_page_content(page)
if not html_text:
return None, None, "获取网页源码失败"
html = etree.HTML(html_text)
if html is None:
return None, None, "解析网页源码失败"
for xpath in self._SITE_LOGIN_XPATH.get("twostep"):
if html.xpath(xpath):
try:
@@ -205,14 +243,17 @@ class CookieHelper:
break
# 登录后的源码
html_text = page.content()
html_text = self.get_page_content(page)
if not html_text:
return None, None, "获取网页源码失败"
if SiteUtils.is_logged_in(html_text):
return self.parse_cookies(page.context.cookies()), \
page.evaluate("() => window.navigator.userAgent"), ""
else:
# 读取错误信息
# 从登录后的页面读取错误信息
html = etree.HTML(html_text)
if html is None:
return None, None, "登录失败"
error_xpath = None
for xpath in self._SITE_LOGIN_XPATH.get("error"):
if html.xpath(xpath):
+44 -11
View File
@@ -18,8 +18,10 @@ from app.log import logger
from app.utils.mixins import ConfigReloadMixin
from app.utils.singleton import Singleton
# 定义一个全局线程池执行器
_executor = concurrent.futures.ThreadPoolExecutor()
# DoH 关闭时需要释放线程池;保持惰性创建可避免未启用 DoH 时占用进程级资源
_executor: Optional[concurrent.futures.ThreadPoolExecutor] = None
_executor_lock = Lock()
_doh_enabled = False
# 定义默认的DoH配置
_doh_timeout = 5
@@ -29,11 +31,21 @@ _doh_lock = Lock()
_orig_getaddrinfo = socket.getaddrinfo
def _get_executor_locked() -> concurrent.futures.ThreadPoolExecutor:
"""在持有执行器锁时按需获取 DoH 查询线程池"""
global _executor
if _executor is None:
_executor = concurrent.futures.ThreadPoolExecutor()
return _executor
def enable_doh(enable: bool) -> None:
"""
socket.getaddrinfo 进行补丁
"""
global _doh_enabled
def _patched_getaddrinfo(host: str, *args, **kwargs):
"""
socket.getaddrinfo的补丁版本
@@ -47,9 +59,15 @@ def enable_doh(enable: bool) -> None:
logger.info(f"已解析 [{host}] 为 [{ip}] (缓存)")
return _orig_getaddrinfo(ip, *args, **kwargs)
# 使用DoH解析主机
futures = []
for resolver in settings.DOH_RESOLVERS.split(","):
futures.append(_executor.submit(_doh_query, resolver, host))
with _executor_lock:
if not _doh_enabled:
return _orig_getaddrinfo(host, *args, **kwargs)
executor = _get_executor_locked()
# 一次解析的任务必须在同一临界区提交完,避免关闭过程中部分任务落入新线程池
futures = [
executor.submit(_doh_query, resolver, host)
for resolver in settings.DOH_RESOLVERS.split(",")
]
for future in concurrent.futures.as_completed(futures):
ip = future.result()
if ip is not None:
@@ -60,11 +78,9 @@ def enable_doh(enable: bool) -> None:
break
return _orig_getaddrinfo(host, *args, **kwargs)
if enable:
# 替换 socket.getaddrinfo 方法
socket.getaddrinfo = _patched_getaddrinfo
else:
socket.getaddrinfo = _orig_getaddrinfo
with _executor_lock:
_doh_enabled = enable
socket.getaddrinfo = _patched_getaddrinfo if enable else _orig_getaddrinfo
class DohHelper(ConfigReloadMixin, metaclass=Singleton):
@@ -77,14 +93,31 @@ class DohHelper(ConfigReloadMixin, metaclass=Singleton):
enable_doh(settings.DOH_ENABLE)
def on_config_changed(self) -> None:
if not settings.DOH_ENABLE:
self.shutdown()
return
with _doh_lock:
# DOH配置有变动的情况下,清空缓存
_doh_cache.clear()
enable_doh(settings.DOH_ENABLE)
enable_doh(True)
def get_reload_name(self) -> str:
return 'DoH'
def shutdown(self) -> None:
"""恢复系统 DNS 并释放 DoH 查询线程池"""
global _executor, _doh_enabled
with _executor_lock:
_doh_enabled = False
socket.getaddrinfo = _orig_getaddrinfo
executor = _executor
_executor = None
with _doh_lock:
_doh_cache.clear()
if executor:
executor.shutdown(wait=True)
def _doh_query(resolver: str, host: str) -> Optional[str]:
"""
使用给定的DoH解析器查询给定主机的IP地址
+69 -8
View File
@@ -415,6 +415,8 @@ class PendingPluginInputInteraction:
payload: Optional[Any] = None
timeout_seconds: int = 120
created_at: datetime = field(default_factory=datetime.now)
# Optional reply binding for channels that can report reply_to_message_id.
prompt_message_id: Optional[str] = None
@property
def expires_at(self) -> datetime:
@@ -504,6 +506,8 @@ class PluginInputInteractionManager:
prompt_id: Optional[str] = None,
timeout_seconds: int = 120,
payload: Optional[Any] = None,
*,
prompt_message_id: Optional[Union[str, int]] = None,
) -> PendingPluginInputInteraction:
with self._lock:
self._cleanup_locked()
@@ -526,6 +530,13 @@ class PluginInputInteractionManager:
if not self._keys_overlap(stored_key, key)
}
normalized_chat_id = str(chat_id) if chat_id not in (None, "") else None
normalized_prompt_message_id = (
str(prompt_message_id)
if channel == MessageChannel.Telegram and normalized_chat_id and prompt_message_id not in (None, "")
else None
)
request = PendingPluginInputInteraction(
request_id=uuid.uuid4().hex[:12],
user_id=str(user_id),
@@ -533,8 +544,9 @@ class PluginInputInteractionManager:
channel=channel,
source=source,
username=username,
chat_id=str(chat_id) if chat_id not in (None, "") else None,
chat_id=normalized_chat_id,
prompt_id=prompt_id,
prompt_message_id=normalized_prompt_message_id,
timeout_seconds=timeout_seconds,
payload=payload,
)
@@ -563,8 +575,16 @@ class PluginInputInteractionManager:
source: Optional[str] = None,
chat_id: Optional[Union[str, int]] = None,
) -> Optional[PendingPluginInputInteraction]:
request, _ = self.consume_by_user(user_id, channel, source, chat_id)
return request
with self._lock:
self._cleanup_locked()
key, request_id = self._find_key_and_request_id_locked(user_id, channel, source, chat_id)
if request_id:
self._by_user_channel.pop(key, None)
return self._by_id.pop(request_id, None)
expired_key, request = self._find_expired_key_and_request_locked(user_id, channel, source, chat_id)
if expired_key:
self._expired_by_user_channel.pop(expired_key, None)
return request
def consume_by_user(
self,
@@ -572,23 +592,64 @@ class PluginInputInteractionManager:
channel: Optional[MessageChannel] = None,
source: Optional[str] = None,
chat_id: Optional[Union[str, int]] = None,
*,
reply_to_message_id: Optional[Union[str, int]] = None,
bypass_reply_check: bool = False,
) -> Tuple[Optional[PendingPluginInputInteraction], Optional[str]]:
with self._lock:
key, request_id = self._find_key_and_request_id_locked(user_id, channel, source, chat_id)
if request_id:
self._by_user_channel.pop(key, None)
request = self._by_id.pop(request_id, None)
if request:
status = "expired" if request.expires_at < datetime.now() else "active"
return request, status
request = self._by_id.get(request_id)
if not request:
self._by_user_channel.pop(key, None)
elif request.expires_at < datetime.now():
self._by_user_channel.pop(key, None)
self._by_id.pop(request_id, None)
if request.prompt_message_id:
return None, None
return request, "expired"
elif not self._reply_matches_prompt(
request,
chat_id,
reply_to_message_id,
ignore_reply_to_message_id=bypass_reply_check,
):
return None, None
else:
self._by_user_channel.pop(key, None)
self._by_id.pop(request_id, None)
return request, "active"
self._cleanup_locked()
key, request = self._find_expired_key_and_request_locked(user_id, channel, source, chat_id)
if request:
self._expired_by_user_channel.pop(key, None)
if request.prompt_message_id:
return None, None
return request, "expired"
self._cleanup_locked()
return None, None
@staticmethod
def _reply_matches_prompt(
request: PendingPluginInputInteraction,
chat_id: Optional[Union[str, int]],
reply_to_message_id: Optional[Union[str, int]],
*,
ignore_reply_to_message_id: bool = False,
) -> bool:
if not request.prompt_message_id:
return True
if not request.chat_id or chat_id in (None, ""):
return False
if str(chat_id) != str(request.chat_id):
return False
if ignore_reply_to_message_id:
return True
if reply_to_message_id in (None, ""):
return False
return str(reply_to_message_id) == str(request.prompt_message_id)
def _find_request_id_locked(
self,
user_id: Union[str, int],
+9 -5
View File
@@ -605,6 +605,7 @@ class MessageQueueManager(metaclass=SingletonClass):
self.check_interval = check_interval
self._running = True
self._stop_event = threading.Event()
self.thread = threading.Thread(target=self._monitor_loop, daemon=True)
self.thread.start()
@@ -752,13 +753,15 @@ class MessageQueueManager(metaclass=SingletonClass):
logger.info(f"队列剩余消息:{self.queue.qsize()}")
except queue.Empty:
break
time.sleep(self.check_interval)
if self._stop_event.wait(self.check_interval):
break
def stop(self) -> None:
"""
停止队列管理器
"""
self._running = False
self._stop_event.set()
logger.info("正在停止消息队列...")
self.thread.join()
logger.info("消息队列已停止")
@@ -841,7 +844,8 @@ def stop_message():
"""
停止消息服务
"""
# 停止消息队列
MessageQueueManager().stop()
# 关闭消息演染器
TemplateHelper().close()
# 只关闭已启动的服务,避免清理路径反向创建后台线程和缓存
if queue_manager := MessageQueueManager.get_existing_instance():
queue_manager.stop()
if template_helper := TemplateHelper.get_existing_instance():
template_helper.close()
+32 -19
View File
@@ -758,7 +758,7 @@ class PluginHelper(metaclass=WeakSingleton):
source_dir,
dest_dir,
dirs_exist_ok=True,
ignore=shutil.ignore_patterns("__pycache__", "*.pyc", ".DS_Store")
ignore=shutil.ignore_patterns("__pycache__", "*.pyc", ".DS_Store", "node_modules")
)
return True, ""
except Exception as e:
@@ -2218,35 +2218,48 @@ class PluginHelper(metaclass=WeakSingleton):
normal_task_key = (loop, normalized_repo_url, False)
force_task_key = (loop, normalized_repo_url, True)
with self._release_task_lock:
force_task = self._release_tasks.get(force_task_key)
if force_task and not force_task.done():
task_key = force_task_key
task = force_task
elif is_fresh():
pending_normal_task = self._release_tasks.get(normal_task_key)
if pending_normal_task and pending_normal_task.done():
pending_normal_task = None
task_key = force_task_key
task = loop.create_task(
self._async_refresh_plugin_repo_releases(normalized_repo_url, pending_normal_task)
)
self._release_tasks[task_key] = task
task.add_done_callback(
lambda completed_task: self._remove_release_task(task_key, completed_task)
)
if is_fresh():
force_task = self._release_tasks.get(force_task_key)
if force_task and not force_task.done():
task_key = force_task_key
task = force_task
else:
pending_normal_task = self._release_tasks.get(normal_task_key)
if pending_normal_task and pending_normal_task.done():
pending_normal_task = None
task_key = force_task_key
task = loop.create_task(
self._async_refresh_plugin_repo_releases(normalized_repo_url, pending_normal_task)
)
self._release_tasks[task_key] = task
task.add_done_callback(
lambda completed_task: self._remove_release_task(task_key, completed_task)
)
else:
task_key = normal_task_key
task = self._release_tasks.get(task_key)
if task is None or task.done():
pending_normal_task = self._release_tasks.get(normal_task_key)
if pending_normal_task is None or pending_normal_task.done():
task = loop.create_task(self._async_get_plugin_repo_releases(normalized_repo_url))
self._release_tasks[task_key] = task
task.add_done_callback(
lambda completed_task: self._remove_release_task(task_key, completed_task)
)
else:
task = pending_normal_task
payload = await asyncio.shield(task)
return self.__parse_plugin_release_response(pid, payload)
async def async_has_plugin_release_cache(self, repo_url: str) -> bool:
"""
判断指定仓库的 Release 列表缓存是否已经存在
"""
if not repo_url:
return False
return await self._async_get_plugin_repo_releases.cache_exists(
self, repo_url.rstrip("/")
)
async def _async_refresh_plugin_repo_releases(
self,
repo_url: str,
+12
View File
@@ -122,6 +122,8 @@
"刮削路径不存在": "Scraping path does not exist",
"保存成功": "Saved successfully",
"保存失败": "Failed to save",
"保存MCP配置成功": "MCP configuration saved successfully",
"保存MCP配置失败": "Failed to save MCP configuration",
"参数错误": "Invalid parameters",
"未配置媒体服务器": "Media server is not configured",
"未找到播放地址": "Playback URL not found",
@@ -174,6 +176,12 @@
"未找到指定的种子": "Specified torrent not found",
"种子删除成功": "Torrent deleted successfully",
"种子缓存清理完成": "Torrent cache cleanup completed",
"TheMovieDb 识别缓存不存在": "TheMovieDb recognition cache does not exist",
"TheMovieDb 识别缓存删除成功": "TheMovieDb recognition cache deleted successfully",
"TheMovieDb 识别缓存清理完成": "TheMovieDb recognition cache cleanup completed",
"豆瓣识别缓存不存在": "Douban recognition cache does not exist",
"豆瓣识别缓存删除成功": "Douban recognition cache deleted successfully",
"豆瓣识别缓存清理完成": "Douban recognition cache cleanup completed",
"重新识别完成": "Re-recognition completed",
"未识别到新名称": "Unable to recognize new name",
"缺少参数": "Missing parameters",
@@ -402,6 +410,10 @@
"source": "获取工具Schema失败: {reason}",
"target": "Failed to get tool schema: {reason}"
},
{
"source": "测试MCP服务器失败: {reason}",
"target": "Failed to test MCP server: {reason}"
},
{
"source": "插件 {plugin} 不存在或未安装",
"target": "Plugin {plugin} does not exist or is not installed"
+13 -1
View File
@@ -102,8 +102,16 @@
"豆瓣网络连接失败": "豆瓣网络连接失败",
"Bangumi网络连接失败": "Bangumi网络连接失败",
"fanart网络连接失败": "fanart网络连接失败",
"保存MCP配置成功": "保存MCP配置成功",
"保存MCP配置失败": "保存MCP配置失败",
"未配置站点或未通过用户认证": "未配置站点或未通过用户认证",
"Redis连接失败,请检查配置": "Redis连接失败,请检查配置"
"Redis连接失败,请检查配置": "Redis连接失败,请检查配置",
"TheMovieDb 识别缓存不存在": "TheMovieDb 识别缓存不存在",
"TheMovieDb 识别缓存删除成功": "TheMovieDb 识别缓存删除成功",
"TheMovieDb 识别缓存清理完成": "TheMovieDb 识别缓存清理完成",
"豆瓣识别缓存不存在": "豆瓣识别缓存不存在",
"豆瓣识别缓存删除成功": "豆瓣识别缓存删除成功",
"豆瓣识别缓存清理完成": "豆瓣识别缓存清理完成"
},
"message_patterns": [
{
@@ -181,6 +189,10 @@
{
"source": "{domain} 网络连接失败",
"target": "{domain} 网络连接失败"
},
{
"source": "测试MCP服务器失败: {reason}",
"target": "测试MCP服务器失败: {reason}"
}
]
}
+12
View File
@@ -122,6 +122,8 @@
"刮削路径不存在": "刮削路徑不存在",
"保存成功": "儲存成功",
"保存失败": "儲存失敗",
"保存MCP配置成功": "MCP 設定儲存成功",
"保存MCP配置失败": "MCP 設定儲存失敗",
"参数错误": "參數錯誤",
"未配置媒体服务器": "未設定媒體伺服器",
"未找到播放地址": "未找到播放位址",
@@ -174,6 +176,12 @@
"未找到指定的种子": "未找到指定的種子",
"种子删除成功": "種子刪除成功",
"种子缓存清理完成": "種子快取清理完成",
"TheMovieDb 识别缓存不存在": "TheMovieDb 識別快取不存在",
"TheMovieDb 识别缓存删除成功": "TheMovieDb 識別快取刪除成功",
"TheMovieDb 识别缓存清理完成": "TheMovieDb 識別快取清理完成",
"豆瓣识别缓存不存在": "豆瓣識別快取不存在",
"豆瓣识别缓存删除成功": "豆瓣識別快取刪除成功",
"豆瓣识别缓存清理完成": "豆瓣識別快取清理完成",
"重新识别完成": "重新識別完成",
"未识别到新名称": "未識別到新名稱",
"缺少参数": "缺少參數",
@@ -402,6 +410,10 @@
"source": "获取工具Schema失败: {reason}",
"target": "取得工具 Schema 失敗: {reason}"
},
{
"source": "测试MCP服务器失败: {reason}",
"target": "測試 MCP 伺服器失敗: {reason}"
},
{
"source": "插件 {plugin} 不存在或未安装",
"target": "插件 {plugin} 不存在或未安裝"
+58 -35
View File
@@ -124,7 +124,7 @@ class NonBlockingFileHandler:
"""
_instance = None
_lock = threading.Lock()
_rotating_handlers = {}
_stop_sentinel = object()
def __new__(cls):
if cls._instance is None:
@@ -138,6 +138,9 @@ class NonBlockingFileHandler:
return
self._initialized = True
self._state_lock = threading.RLock()
self._handlers_lock = threading.Lock()
self._rotating_handlers = {}
self._write_queue = queue.Queue(maxsize=log_settings.ASYNC_FILE_QUEUE_SIZE)
self._executor = ThreadPoolExecutor(max_workers=log_settings.ASYNC_FILE_WORKERS,
thread_name_prefix="LogWriter")
@@ -151,27 +154,28 @@ class NonBlockingFileHandler:
"""
获取或创建RotatingFileHandler实例
"""
if file_path not in self._rotating_handlers:
# 确保目录存在
file_path.parent.mkdir(parents=True, exist_ok=True)
with self._handlers_lock:
if file_path not in self._rotating_handlers:
# 确保目录存在
file_path.parent.mkdir(parents=True, exist_ok=True)
# 创建RotatingFileHandler
handler = RotatingFileHandler(
filename=str(file_path),
maxBytes=log_settings.LOG_MAX_FILE_SIZE_BYTES,
backupCount=log_settings.LOG_BACKUP_COUNT,
encoding='utf-8'
)
# 创建RotatingFileHandler
handler = RotatingFileHandler(
filename=str(file_path),
maxBytes=log_settings.LOG_MAX_FILE_SIZE_BYTES,
backupCount=log_settings.LOG_BACKUP_COUNT,
encoding='utf-8'
)
# 设置格式化器
formatter = logging.Formatter(log_settings.LOG_FILE_FORMAT)
handler.setFormatter(formatter)
# 设置格式化器
formatter = logging.Formatter(log_settings.LOG_FILE_FORMAT)
handler.setFormatter(formatter)
self._rotating_handlers[file_path] = handler
self._rotating_handlers[file_path] = handler
return self._rotating_handlers[file_path]
return self._rotating_handlers[file_path]
def write_log(self, level: str, message: str, file_path: Path):
def write_log(self, level: str, message: str, file_path: Path) -> None:
"""
写入日志 - 自动检测协程环境并使用合适的方式
"""
@@ -181,8 +185,11 @@ class NonBlockingFileHandler:
if self._is_in_event_loop():
# 在协程环境中,使用非阻塞方式
self._write_non_blocking(entry)
else:
# 不在协程环境中,直接同步写入
return
with self._state_lock:
if not self._running:
return
# 不在协程环境中,持锁同步写入,避免关闭文件处理器时仍有写操作进行
self._write_sync(entry)
@staticmethod
@@ -196,15 +203,19 @@ class NonBlockingFileHandler:
except RuntimeError:
return False
def _write_non_blocking(self, entry: LogEntry):
def _write_non_blocking(self, entry: LogEntry) -> bool:
"""
非阻塞写入用于协程环境
"""
try:
self._write_queue.put_nowait(entry)
except queue.Full:
# 队列满时,使用线程池处理
self._executor.submit(self._write_sync, entry)
with self._state_lock:
if not self._running:
return False
try:
self._write_queue.put_nowait(entry)
except queue.Full:
# 队列满时,使用线程池处理
self._executor.submit(self._write_sync, entry)
return True
@staticmethod
def _write_sync(entry: LogEntry):
@@ -215,8 +226,7 @@ class NonBlockingFileHandler:
# 获取RotatingFileHandler实例
handler = NonBlockingFileHandler()._get_rotating_handler(entry.file_path)
# 使用RotatingFileHandler的emit方法,只传递原始消息
handler.emit(logging.LogRecord(
handler.handle(logging.LogRecord(
name='',
level=getattr(logging, entry.level.upper(), logging.INFO),
pathname='',
@@ -235,22 +245,28 @@ class NonBlockingFileHandler:
"""
后台批量写入线程
"""
while self._running:
while True:
try:
# 收集一批日志条目
batch = []
should_stop = False
end_time = time.time() + log_settings.WRITE_TIMEOUT
while len(batch) < log_settings.BATCH_WRITE_SIZE and time.time() < end_time:
try:
remaining_time = max(0, end_time - time.time())
entry = self._write_queue.get(timeout=remaining_time)
if entry is self._stop_sentinel:
should_stop = True
break
batch.append(entry)
except queue.Empty:
break
if batch:
self._write_batch(batch)
if should_stop:
break
except Exception as e:
print(f"批量写入线程错误: {e}")
@@ -275,8 +291,7 @@ class NonBlockingFileHandler:
# 批量写入
for entry in entries:
# 使用RotatingFileHandler的emit方法,只传递原始消息
handler.emit(logging.LogRecord(
handler.handle(logging.LogRecord(
name='',
level=getattr(logging, entry.level.upper(), logging.INFO),
pathname='',
@@ -294,15 +309,23 @@ class NonBlockingFileHandler:
def shutdown(self):
"""
关闭文件处理器
排空异步日志并关闭文件处理器
"""
self._running = False
if hasattr(self, '_write_thread'):
self._write_thread.join(timeout=5)
with self._state_lock:
if not self._running:
return
self._running = False
if hasattr(self, '_write_thread') and self._write_thread.is_alive():
# 状态锁保证停止标记之后不会再有生产者入队
self._write_queue.put(self._stop_sentinel)
if hasattr(self, '_write_thread') and self._write_thread.is_alive():
self._write_thread.join()
if self._executor:
self._executor.shutdown(wait=True)
# 清理缓存
for handler in self._rotating_handlers.values():
handler.flush()
handler.close()
self._rotating_handlers.clear()
+27 -7
View File
@@ -30,16 +30,31 @@ elif SystemUtils.is_frozen():
sys.stderr = open(os.devnull, 'w')
from app.factory import app
from app.core.config import settings
from app.core.config import global_vars, settings
from app.db.init import init_db, update_db
# 设置进程名
setproctitle.setproctitle(settings.PROJECT_NAME)
class MoviePilotServer(uvicorn.Server):
"""在 Uvicorn 开始优雅退出前发布应用协作停止标志"""
def handle_exit(self, sig, frame) -> None:
global_vars.stop_system()
super().handle_exit(sig, frame)
# uvicorn服务
Server = uvicorn.Server(Config(app, host=settings.HOST, port=settings.PORT,
reload=settings.DEV, workers=multiprocessing.cpu_count() * 2 + 1,
timeout_graceful_shutdown=60))
Server = MoviePilotServer(Config(app, host=settings.HOST, port=settings.PORT,
reload=settings.DEV, workers=multiprocessing.cpu_count() * 2 + 1,
timeout_graceful_shutdown=60))
def request_shutdown() -> None:
"""发布协作停止标志并请求 Uvicorn 退出"""
global_vars.stop_system()
Server.should_exit = True
def start_tray():
@@ -64,8 +79,8 @@ def start_tray():
"""
退出程序
"""
request_shutdown()
TrayIcon.stop()
Server.should_exit = True
import pystray
@@ -93,10 +108,11 @@ def signal_handler(signum, frame):
信号处理函数用于优雅停止服务
"""
print(f"收到信号 {signum},开始优雅停止服务...")
Server.should_exit = True
request_shutdown()
if __name__ == '__main__':
def run_application() -> None:
"""初始化进程并启动 API 服务"""
# 注册信号处理器
signal.signal(signal.SIGTERM, signal_handler)
signal.signal(signal.SIGINT, signal_handler)
@@ -109,3 +125,7 @@ if __name__ == '__main__':
update_db()
# 启动API服务
Server.run()
if __name__ == '__main__':
run_application()
+17 -2
View File
@@ -1,8 +1,10 @@
import threading
from abc import abstractmethod, ABCMeta
from typing import Generic, Tuple, Union, TypeVar, Type, Dict, Optional, Callable
from pathlib import Path
from app.helper.service import ServiceConfigHelper
from app.log import logger
from app.schemas import Notification, NotificationConf, MediaServerConf, DownloaderConf
from app.schemas.types import ModuleType, DownloaderType, MediaServerType, MessageChannel, StorageSchema, \
OtherModulesType, SystemConfigKey
@@ -15,8 +17,21 @@ class _ModuleBase(ConfigReloadMixin, metaclass=ABCMeta):
输入参数与输出参数一致的或没有输出的可以被多个模块重复实现
"""
def on_config_changed(self):
self.init_module()
def __init__(self) -> None:
"""初始化模块生命周期锁"""
super().__init__()
self._reload_lock = threading.RLock()
def on_config_changed(self) -> None:
"""串行停止旧资源并按最新配置重新初始化模块"""
with self._reload_lock:
try:
self.stop()
except Exception as err:
logger.error(
f"停止 {self.get_reload_name()} 旧资源失败,继续按最新配置初始化:{err}"
)
self.init_module()
def get_reload_name(self):
return self.get_name()
+6 -6
View File
@@ -58,7 +58,6 @@ class DiscordModule(_ModuleBase, _MessageBase[Discord]):
if not Discord:
logger.error("Discord 依赖未就绪(需要安装 discord.py==2.6.4),模块未启动")
return
self.stop()
super().init_service(
service_name=Discord.__name__.lower(), service_type=Discord
)
@@ -89,12 +88,13 @@ class DiscordModule(_ModuleBase, _MessageBase[Discord]):
"""
return 4
def stop(self):
"""
停止模块
"""
def stop(self) -> None:
"""停止模块"""
for client in self.get_instances().values():
client.stop()
try:
client.stop()
except Exception as err:
logger.error(f"停止Discord模块实例失败:{err}")
def test(self) -> Optional[Tuple[bool, str]]:
"""
+28 -1
View File
@@ -25,10 +25,11 @@ class DoubanCache(metaclass=WeakSingleton):
"type": MediaType
}
"""
# TMDB缓存过期
# 豆瓣缓存过期
_douban_cache_expire: bool = True
def __init__(self):
"""初始化豆瓣识别缓存并恢复本地持久化数据。"""
self.maxsize = settings.CONF.douban
self.ttl = settings.CONF.meta
self.region = "__douban_cache__"
@@ -46,6 +47,30 @@ class DoubanCache(metaclass=WeakSingleton):
"""
with lock:
self._cache.clear()
self.save(force=True)
def list_items(self) -> list[dict]:
"""返回可供管理界面展示的豆瓣识别缓存列表。"""
with lock:
cache_items = []
for key, value in self._cache.items():
if not isinstance(value, dict):
continue
media_type = value.get("type")
if not isinstance(media_type, MediaType):
try:
media_type = MediaType(media_type)
except (TypeError, ValueError):
media_type = None
cache_items.append({
"key": key,
"douban_id": value.get("id") or 0,
"title": value.get("title") or "",
"year": value.get("year") or "",
"media_type": media_type.to_agent() if media_type else "unknown",
"poster_path": value.get("poster_path") or "",
})
return sorted(cache_items, key=lambda item: item["key"])
@staticmethod
def __get_key(meta: MetaBase) -> str:
@@ -73,6 +98,7 @@ class DoubanCache(metaclass=WeakSingleton):
redis_data = self._cache.get(key)
if redis_data:
self._cache.delete(key)
self.save(force=True)
return redis_data
return {}
@@ -169,4 +195,5 @@ class DoubanCache(metaclass=WeakSingleton):
pickle.dump(new_meta_data, f, pickle.HIGHEST_PROTOCOL) # noqa
def __del__(self):
"""实例释放前保存非 Redis 缓存。"""
self.save()
+6 -7
View File
@@ -10,7 +10,6 @@ from app.schemas.types import ModuleType
class FeishuModule(_ModuleBase, _MessageBase[Feishu]):
def init_module(self) -> None:
self.stop()
super().init_service(service_name=Feishu.__name__.lower(), service_type=Feishu)
self._channel = MessageChannel.Feishu
@@ -30,13 +29,13 @@ class FeishuModule(_ModuleBase, _MessageBase[Feishu]):
def get_priority() -> int:
return 2
def stop(self):
def stop(self) -> None:
"""停止模块"""
for client in self.get_instances().values():
if hasattr(client, "stop"):
try:
client.stop()
except Exception as err:
logger.error(f"停止飞书模块实例失败:{err}")
try:
client.stop()
except Exception as err:
logger.error(f"停止飞书模块实例失败:{err}")
def test(self) -> Optional[Tuple[bool, str]]:
if not self.get_instances():
+2 -11
View File
@@ -92,13 +92,6 @@ class FilterModule(_ModuleBase):
self.rule_set = deepcopy(self.builtin_rule_set)
self.__init_custom_rules()
def on_config_changed(self) -> None:
"""
自定义过滤或 Meta 识别配置变更后重建规则集并刷新 Rust Meta 配置缓存
"""
clear_rust_parse_options_cache()
self.init_module()
def __init_custom_rules(self):
"""
加载用户自定义规则如跟内置规则冲突以用户自定义规则为准
@@ -137,10 +130,8 @@ class FilterModule(_ModuleBase):
return 4
def stop(self) -> None:
"""
停止过滤器模块
"""
pass
"""停止模块"""
clear_rust_parse_options_cache()
def test(self) -> None:
"""
+40 -4
View File
@@ -317,10 +317,18 @@ class Jellyfin:
def get_medias_count(self) -> schemas.Statistic:
"""
获得电影电视剧动漫媒体数量
:return: MovieCount SeriesCount SongCount
优先遍历用户媒体库视图逐库统计全局 `Items/Counts` 按数据库原始条目
计数同一影片在库内有多个版本/多个文件夹拷贝时会重复累计#5915),
而用户级 `Users/{user}/Items` 查询会折叠版本 Jellyfin 页面显示一致
仅在用户视图不可用时回退到 `Items/Counts`
:return: MovieCount SeriesCount EpisodeCount
"""
if not self._host or not self._apikey:
return schemas.Statistic()
stat = self.__count_medias_by_librarys()
if stat is not None:
return stat
url = f"{self._host}Items/Counts"
params = {
'api_key': self._apikey
@@ -341,6 +349,32 @@ class Jellyfin:
logger.error(f"连接Items/Counts出错:" + str(e))
return schemas.Statistic()
def __count_medias_by_librarys(self) -> Optional[schemas.Statistic]:
"""
遍历用户媒体库视图逐库统计媒体数量
`Users/{user}/Views` 每个媒体库仅返回一条记录库包含多个文件夹时
也不会重复 `CollectionType` 分桶后用用户级条目查询累计
:return: 统计结果用户或媒体库视图不可用时返回None由调用方回退
"""
if not self.user:
return None
librarys = self.__get_jellyfin_librarys()
if not librarys:
return None
stat = schemas.Statistic()
for library in librarys:
library_id = library.get("Id")
if not library_id:
continue
collection_type = library.get("CollectionType")
if collection_type == "movies":
stat.movie_count += self.get_items_count(library_id, include_item_types="Movie") or 0
elif collection_type == "tvshows":
stat.tv_count += self.get_items_count(library_id, include_item_types="Series") or 0
stat.episode_count += self.get_items_count(library_id, include_item_types="Episode") or 0
return stat
def __get_jellyfin_series_id_by_name(self, name: str, year: str) -> Optional[str]:
"""
根据名称查询Jellyfin中剧集的SeriesId
@@ -809,11 +843,13 @@ class Jellyfin:
logger.error(f"连接Users/{self.user}/Items/{itemid}" + str(e))
return None
def get_items_count(self, parent: Union[str, int]) -> Optional[int]:
def get_items_count(self, parent: Union[str, int],
include_item_types: str = "Movie,Series") -> Optional[int]:
"""
获取指定媒体库可同步的电影和剧集总数
获取指定媒体库可同步的媒体条目总数
:param parent: 媒体库ID
:param include_item_types: 统计的条目类型默认电影和剧集
:return: 媒体条目总数查询失败时返回None
"""
if not parent or not self._host or not self._apikey or not self.user:
@@ -822,7 +858,7 @@ class Jellyfin:
params = {
"ParentId": parent,
"Recursive": "true",
"IncludeItemTypes": "Movie,Series",
"IncludeItemTypes": include_item_types,
"Limit": 0,
"api_key": self._apikey,
}
+7 -6
View File
@@ -44,13 +44,14 @@ class PlexModule(_ModuleBase, _MediaServerBase[Plex]):
"""
return 3
def stop(self):
"""
停止模块服务
"""
def stop(self) -> None:
"""停止模块"""
for server in self.get_instances().values():
if server:
server.close()
try:
if server:
server.close()
except Exception as err:
logger.error(f"停止Plex模块实例失败:{err}")
def test(self) -> Optional[Tuple[bool, str]]:
"""
+6 -4
View File
@@ -195,9 +195,6 @@ class QbittorrentModule(_ModuleBase, _DownloaderBase[Qbittorrent]):
ignore_category_check=False
)
# 获取种子内容布局: `Original: 原始, Subfolder: 创建子文件夹, NoSubfolder: 不创建子文件夹`
torrent_layout = server.get_content_layout()
if not state:
# 查询所有下载器的种子
torrents, error = server.get_torrents()
@@ -210,7 +207,8 @@ class QbittorrentModule(_ModuleBase, _DownloaderBase[Qbittorrent]):
if torrent.get("name") == getattr(torrent_from_file, 'name', '') \
and torrent.get("total_size") == getattr(torrent_from_file, 'total_size', 0):
torrent_hash = torrent.get("hash")
torrent_tags = [str(tag).strip() for tag in torrent.get("tags").split(',')]
server.delete_torrents_tag(torrent_hash, tag)
torrent_tags = [str(tag).strip() for tag in (torrent.get("tags") or "").split(',')]
logger.warn(f"下载器中已存在该种子任务:{torrent_hash} - {torrent.get('name')}")
# 给种子打上标签
if "已整理" in torrent_tags:
@@ -218,6 +216,8 @@ class QbittorrentModule(_ModuleBase, _DownloaderBase[Qbittorrent]):
if settings.TORRENT_TAG and settings.TORRENT_TAG not in torrent_tags:
logger.info(f"给种子 {torrent_hash} 打上标签:{settings.TORRENT_TAG}")
server.set_torrents_tag(ids=torrent_hash, tags=[settings.TORRENT_TAG])
# 获取种子内容布局: `Original: 原始, Subfolder: 创建子文件夹, NoSubfolder: 不创建子文件夹`
torrent_layout = server.get_content_layout()
return downloader or self.get_default_config_name(), torrent_hash, torrent_layout, f"下载任务已存在"
finally:
torrents.clear()
@@ -233,6 +233,8 @@ class QbittorrentModule(_ModuleBase, _DownloaderBase[Qbittorrent]):
if not torrent_hash:
return None, None, None, f"下载任务添加成功,但获取Qbittorrent任务信息失败:{content}"
else:
# 获取种子内容布局: `Original: 原始, Subfolder: 创建子文件夹, NoSubfolder: 不创建子文件夹`
torrent_layout = server.get_content_layout()
if is_paused:
# 种子文件
torrent_files = server.get_files(
+32 -5
View File
@@ -259,9 +259,34 @@ class Qbittorrent:
"""
if not self.qbc:
return None
# completed会包含移动状态 改为获取seeding状态 包含活动上传, 正在做种, 及强制做种
torrents, error = self.get_torrents(status="seeding", ids=ids, tags=tags)
return None if error else torrents or []
torrents, error = self.get_torrents(status="completed", ids=ids, tags=tags)
if error:
return None
ret_torrents = []
for torrent in torrents or []:
state = str(torrent.get("state") or "").strip().lower()
progress = torrent.get("progress") or 0
amount_left = torrent.get("amount_left") or 0
if (
progress >= 1
and amount_left <= 0
and state not in {
"allocating",
"checkingdl",
"checkingup",
"downloading",
"error",
"forceddl",
"missingfiles",
"metadl",
"moving",
"queueddl",
"stalleddl",
"unknown",
}
):
ret_torrents.append(torrent)
return ret_torrents
def get_downloading_torrents(self, ids: Union[str, list] = None,
tags: Union[str, list] = None) -> Optional[List[TorrentDictionary]]:
@@ -278,14 +303,16 @@ class Qbittorrent:
def delete_torrents_tag(self, ids: Union[str, list], tag: Union[str, list]) -> bool:
"""
删除Tag
从指定种子移除标签并删除全局标签定义
:param ids: 种子Hash列表
:param tag: 标签内容
:return: 是否删除成功
"""
if not self.qbc:
return False
try:
self.qbc.torrents_delete_tags(torrent_hashes=ids, tags=tag)
self.qbc.torrents_remove_tags(torrent_hashes=ids, tags=tag)
self.qbc.torrents_delete_tags(tags=tag)
return True
except Exception as err:
logger.error(f"删除种子Tag出错:{str(err)}")
+4 -2
View File
@@ -46,7 +46,6 @@ class QQBotModule(_ModuleBase, _MessageBase[QQBot]):
)
def init_module(self) -> None:
self.stop()
super().init_service(service_name=QQBot.__name__.lower(), service_type=QQBot)
self._channel = MessageChannel.QQ
@@ -67,9 +66,12 @@ class QQBotModule(_ModuleBase, _MessageBase[QQBot]):
return 10
def stop(self) -> None:
"""停止模块"""
for client in self.get_instances().values():
if hasattr(client, "stop"):
try:
client.stop()
except Exception as err:
logger.error(f"停止QQ Bot模块实例失败:{err}")
def test(self) -> Optional[Tuple[bool, str]]:
if not self.get_instances():
+6 -5
View File
@@ -69,12 +69,13 @@ class SlackModule(_ModuleBase, _MessageBase[Slack]):
"""
return 3
def stop(self):
"""
停止模块
"""
def stop(self) -> None:
"""停止模块"""
for client in self.get_instances().values():
client.stop()
try:
client.stop()
except Exception as err:
logger.error(f"停止Slack模块实例失败:{err}")
def test(self) -> Optional[Tuple[bool, str]]:
"""
+24 -6
View File
@@ -62,12 +62,13 @@ class TelegramModule(_ModuleBase, _MessageBase[Telegram]):
"""
return 0
def stop(self):
"""
停止模块
"""
def stop(self) -> None:
"""停止模块"""
for client in self.get_instances().values():
client.stop()
try:
client.stop()
except Exception as err:
logger.error(f"停止Telegram模块实例失败:{err}")
def test(self) -> Optional[Tuple[bool, str]]:
"""
@@ -252,9 +253,11 @@ class TelegramModule(_ModuleBase, _MessageBase[Telegram]):
处理普通文本消息
"""
text = msg.get("text") or msg.get("caption")
message_id = msg.get("message_id")
user_id = msg.get("from", {}).get("id")
user_name = msg.get("from", {}).get("username")
chat_id = msg.get("chat", {}).get("id")
reply_to_message_id = (msg.get("reply_to_message") or {}).get("message_id")
# 将 text_link 实体中的 URL 嵌入到文本中
if text:
@@ -309,7 +312,9 @@ class TelegramModule(_ModuleBase, _MessageBase[Telegram]):
userid=user_id,
username=user_name,
text=cleaned_text,
message_id=message_id,
chat_id=str(chat_id) if chat_id else None,
reply_to_message_id=reply_to_message_id,
images=images if images else None,
audio_refs=audio_refs if audio_refs else None,
files=files if files else None,
@@ -514,6 +519,12 @@ class TelegramModule(_ModuleBase, _MessageBase[Telegram]):
parse_mode=message.parse_mode,
)
else:
# Telegram 的 reply_markup 不能同时承载 InlineKeyboard 和 ForceReply。
# 普通通知只清空可编辑消息 ID,仍保留原会话作为新消息目标。
has_interaction_context = bool(message.buttons or message.force_reply)
original_message_id = (
message.original_message_id if has_interaction_context else None
)
client.send_msg(
title=message.title,
text=message.text,
@@ -522,7 +533,7 @@ class TelegramModule(_ModuleBase, _MessageBase[Telegram]):
link=message.link,
buttons=message.buttons,
force_reply=message.force_reply,
original_message_id=message.original_message_id,
original_message_id=original_message_id,
original_chat_id=message.original_chat_id,
disable_web_page_preview=message.disable_web_page_preview,
parse_mode=message.parse_mode,
@@ -735,12 +746,19 @@ class TelegramModule(_ModuleBase, _MessageBase[Telegram]):
parse_mode=message.parse_mode,
)
else:
# direct message 只禁用编辑旧消息;仅 ForceReply 使用 original_chat_id
# 发回原会话,并保留 original_message_id 让 client reply_to 原消息。
original_chat_id = message.original_chat_id if message.force_reply else None
original_message_id = message.original_message_id if message.force_reply else None
result = client.send_msg(
title=message.title,
text=message.text,
image=message.image,
userid=userid,
link=message.link,
force_reply=message.force_reply,
original_message_id=original_message_id,
original_chat_id=original_chat_id,
disable_web_page_preview=message.disable_web_page_preview,
parse_mode=message.parse_mode,
)
+24
View File
@@ -0,0 +1,24 @@
def ensure_urllib3_header_param_compat() -> None:
"""
pyTelegramBotAPI imports urllib3.fields.format_header_param at import time.
Some urllib3-future builds only expose newer formatter names.
RFC 2231 formatting is kept as the last fallback because it encodes
non-ASCII values differently from urllib3's old default.
"""
try:
from urllib3 import fields
except ImportError:
return
if hasattr(fields, "format_header_param"):
return
for fallback_name in (
"format_header_param_html5",
"format_multipart_header_param",
"format_header_param_rfc2231",
):
fallback = getattr(fields, fallback_name, None)
if fallback is not None:
fields.format_header_param = fallback
return
+30 -24
View File
@@ -8,36 +8,41 @@ from pathlib import Path
from typing import Any, Callable, Dict, List, Optional, Union
from urllib.parse import urljoin, quote
from telebot import TeleBot, apihelper
from telebot.types import (
from app.modules.telegram.compat import ensure_urllib3_header_param_compat
# Must run before importing pyTelegramBotAPI.
ensure_urllib3_header_param_compat()
from telebot import TeleBot, apihelper # noqa: E402
from telebot.types import ( # noqa: E402
BotCommand,
InlineKeyboardMarkup,
InlineKeyboardButton,
InputMediaPhoto,
)
try:
from telebot.types import ForceReply
from telebot.types import ForceReply # noqa: E402
except ImportError:
ForceReply = None
from telegramify_markdown import standardize, telegramify # noqa
from telegramify_markdown import standardize, telegramify # noqa: E402
try:
from telegramify_markdown import entities_to_markdownv2 # noqa
from telegramify_markdown import entities_to_markdownv2 # noqa: E402
except ImportError:
entities_to_markdownv2 = None
try:
from telegramify_markdown.content import ContentTypes, File, Photo, Text
from telegramify_markdown.content import ContentTypes, File, Photo, Text # noqa: E402
except ImportError:
from telegramify_markdown.type import ContentTypes, File, Photo, Text
from telegramify_markdown.type import ContentTypes, File, Photo, Text # noqa: E402
from app.core.config import settings
from app.core.context import MediaInfo, Context
from app.core.metainfo import MetaInfo
from app.helper.image import ImageHelper
from app.helper.thread import ThreadHelper
from app.log import logger
from app.utils.common import retry
from app.utils.http import RequestUtils
from app.utils.string import StringUtils
from app.core.config import settings # noqa: E402
from app.core.context import MediaInfo, Context # noqa: E402
from app.core.metainfo import MetaInfo # noqa: E402
from app.helper.image import ImageHelper # noqa: E402
from app.helper.thread import ThreadHelper # noqa: E402
from app.log import logger # noqa: E402
from app.utils.common import retry # noqa: E402
from app.utils.http import RequestUtils # noqa: E402
from app.utils.string import StringUtils # noqa: E402
TELEGRAM_PARSE_MODE_MARKDOWN = "MarkdownV2"
@@ -274,8 +279,6 @@ class Telegram:
@staticmethod
def _telegramify_item_text(item: Text) -> str:
"""将 telegramify 文本片段转换为 Telegram MarkdownV2 字符串。"""
if hasattr(item, "content"):
return item.content
if entities_to_markdownv2:
return entities_to_markdownv2(item.text, item.entities)
return standardize(item.text)
@@ -285,8 +288,6 @@ class Telegram:
"""将 telegramify 文本或媒体片段转换为 Telegram MarkdownV2 caption。"""
if isinstance(item, Text):
return Telegram._telegramify_item_text(item)
if hasattr(item, "caption"):
return item.caption
if entities_to_markdownv2:
return entities_to_markdownv2(item.caption_text, item.caption_entities)
return standardize(item.caption_text)
@@ -1535,14 +1536,19 @@ class Telegram:
# 清理菜单命令
self._bot.delete_my_commands()
def stop(self):
def stop(self) -> None:
"""
停止Telegram消息接收服务
"""
# 停止所有typing任务
for chat_id in list(self._typing_tasks.keys()):
self._stop_typing_task(chat_id)
if self._bot:
self._bot.stop_polling()
if not self._bot:
return
self._bot.stop_bot()
if self._polling_thread:
self._polling_thread.join()
logger.info("Telegram消息接收服务已停止")
self._polling_thread = None
self._bot = None
logger.info("Telegram消息接收服务已停止")
+7 -9
View File
@@ -43,12 +43,6 @@ class TheMovieDbModule(_ModuleBase):
self.category = CategoryHelper()
self.scraper = TmdbScraper()
def on_config_changed(self):
# 停止模块
self.stop()
# 初始化模块
self.init_module()
@staticmethod
def get_name() -> str:
return "TheMovieDb"
@@ -74,9 +68,13 @@ class TheMovieDbModule(_ModuleBase):
"""
return 1
def stop(self):
self.cache.save()
self.tmdb.close()
def stop(self) -> None:
"""停止模块"""
# 缓存持久化失败不能阻断 HTTP 客户端关闭
try:
self.cache.save()
finally:
self.tmdb.close()
def test(self) -> Tuple[bool, str]:
"""
+30
View File
@@ -27,6 +27,7 @@ class TmdbCache(metaclass=WeakSingleton):
_tmdb_cache_expire: bool = True
def __init__(self):
"""初始化 TMDB 识别缓存并恢复本地持久化数据。"""
self.maxsize = settings.CONF.douban
self.ttl = settings.CONF.meta
self.region = "__tmdb_cache__"
@@ -44,6 +45,33 @@ class TmdbCache(metaclass=WeakSingleton):
"""
with lock:
self._cache.clear()
self.save(force=True)
def list_items(self) -> list[dict]:
"""
返回可供管理界面展示的 TMDB 识别缓存列表
"""
with lock:
cache_items = []
for key, value in self._cache.items():
if not isinstance(value, dict):
continue
media_type = value.get("type")
if not isinstance(media_type, MediaType):
try:
media_type = MediaType(media_type)
except (TypeError, ValueError):
media_type = None
cache_items.append({
"key": key,
"tmdb_id": value.get("id") or 0,
"title": value.get("title") or "",
"year": value.get("year") or "",
"media_type": media_type.to_agent() if media_type else "unknown",
"poster_path": value.get("poster_path") or "",
"backdrop_path": value.get("backdrop_path") or "",
})
return sorted(cache_items, key=lambda item: item["key"])
@staticmethod
def __get_key(meta: MetaBase) -> str:
@@ -71,6 +99,7 @@ class TmdbCache(metaclass=WeakSingleton):
redis_data = self._cache.get(key)
if redis_data:
self._cache.delete(key)
self.save(force=True)
return redis_data
return {}
@@ -156,4 +185,5 @@ class TmdbCache(metaclass=WeakSingleton):
pickle.dump(new_meta_data, f, pickle.HIGHEST_PROTOCOL) # type: ignore
def __del__(self):
"""实例释放前保存非 Redis 缓存。"""
self.save()
-1
View File
@@ -114,7 +114,6 @@ class TheTvDbModule(_ModuleBase):
return 4
def stop(self):
logger.info("TheTvDbModule 停止。正在清除 TVDB 会话。")
with self.__auth_lock:
self.tvdb = None
+7 -3
View File
@@ -61,10 +61,14 @@ class TrimeMediaModule(_ModuleBase, _MediaServerBase[TrimeMedia]):
logger.info(f"飞牛影视 {name} 连接断开,尝试重连 ...")
server.reconnect()
def stop(self):
def stop(self) -> None:
"""停止模块"""
for server in self.get_instances().values():
if server.is_authenticated():
server.disconnect()
try:
if server.is_authenticated():
server.disconnect()
except Exception as err:
logger.error(f"停止飞牛影视模块实例失败:{err}")
def test(self) -> Optional[Tuple[bool, str]]:
"""
+7 -3
View File
@@ -60,10 +60,14 @@ class UgreenModule(_ModuleBase, _MediaServerBase[Ugreen]):
logger.info(f"绿联影视 {name} 连接断开,尝试重连 ...")
server.reconnect()
def stop(self):
def stop(self) -> None:
"""停止模块"""
for server in self.get_instances().values():
if server.is_authenticated():
server.disconnect()
try:
if server.is_authenticated():
server.disconnect()
except Exception as err:
logger.error(f"停止绿联影视模块实例失败:{err}")
def test(self) -> Optional[Tuple[bool, str]]:
"""
+50 -11
View File
@@ -1,7 +1,7 @@
import base64
import uuid
from dataclasses import dataclass
from typing import Any, Dict, Mapping, Optional, Union
from typing import Any, Mapping, Optional, Union
from urllib.parse import urlsplit, urlunsplit
from requests import Session
@@ -13,6 +13,10 @@ from app.utils.url import UrlUtils
@dataclass
class ApiResult:
"""
绿联接口标准响应封装
"""
code: int = -1
msg: str = ""
data: Any = None
@@ -21,6 +25,7 @@ class ApiResult:
@property
def success(self) -> bool:
"""判断绿联接口是否返回成功状态"""
return self.code == 200
@@ -59,7 +64,17 @@ class Api:
ug_agent: str = "PC/WEB",
timeout: int = 20,
verify_ssl: bool = True,
):
) -> None:
"""
初始化绿联影视 API 客户端
:param host: 绿联服务端地址
:param client_version: 绿联 Web 客户端版本号
:param language: 请求语言
:param ug_agent: 绿联客户端标识
:param timeout: HTTP 请求超时时间
:param verify_ssl: 是否校验 HTTPS 证书
"""
self._host = self._normalize_base_url(host)
self._session = Session()
@@ -80,25 +95,30 @@ class Api:
@property
def host(self) -> str:
"""获取规范化后的绿联服务端地址"""
return self._host
@property
def token(self) -> Optional[str]:
"""获取当前登录会话 token"""
return self._token
@property
def static_token(self) -> Optional[str]:
"""获取可用于静态资源访问的 token"""
return self._static_token
@property
def is_ugk(self) -> bool:
"""判断当前会话是否使用 ugk 访问参数"""
return self._is_ugk
@property
def public_key(self) -> Optional[str]:
"""获取当前会话加密公钥"""
return self._public_key
def close(self):
def close(self) -> None:
"""
关闭底层 HTTP 会话
"""
@@ -141,13 +161,14 @@ class Api:
def _common_headers(self) -> dict[str, str]:
"""
获取绿联 Web 端通用请求头
获取绿联 Web 端通用请求头兼容新版登录客户端标识
"""
return {
"Accept": "application/json, text/plain, */*",
"Client-Id": self._client_id,
"Client-Version": self._client_version,
"UG-Agent": self._ug_agent,
"UG-Client-Id": self._client_id,
"X-Specify-Language": self._language,
}
@@ -262,14 +283,32 @@ class Api:
logger.error(f"绿联登录失败:{login_result.msg}")
return None
token = str(login_result.data.get("token") or "").strip()
public_key = self._decode_public_key(str(login_result.data.get("public_key") or ""))
token = str(
login_result.data.get("token")
or login_result.data.get("token_id")
or login_result.data.get("tokenId")
or ""
).strip()
public_key = (
self._decode_public_key(
str(
login_result.data.get("public_key")
or login_result.data.get("publicKey")
or ""
)
)
or login_public_key
)
if not token or not public_key:
logger.error("绿联登录失败:未返回 token/public_key")
logger.error("绿联登录失败:未返回 token/token_id 或可用公钥")
return None
self._token = token
static_token = str(login_result.data.get("static_token") or "").strip()
static_token = str(
login_result.data.get("static_token")
or login_result.data.get("staticToken")
or ""
).strip()
self._static_token = static_token or self._token
self._is_ugk = bool(login_result.data.get("is_ugk"))
self._public_key = public_key
@@ -365,7 +404,7 @@ class Api:
)
return True
def logout(self):
def logout(self) -> None:
"""
登出并清理本地认证状态
"""
@@ -569,7 +608,7 @@ class Api:
"""
获取海报墙文件夹与条目可按目录路径递归展开
"""
params: Dict[str, Any] = {
params: dict[str, Any] = {
"page": page,
"page_size": page_size,
"sort_type": sort_type,
@@ -590,7 +629,7 @@ class Api:
"""
获取电影详情
"""
params: Dict[str, Any] = {
params: dict[str, Any] = {
"id": item_id,
"media_lib_set_id": media_lib_set_id,
"fileVersion": "true",
+6 -6
View File
@@ -24,7 +24,6 @@ class WechatModule(_ModuleBase, _MessageBase[WeChat]):
"""
初始化模块
"""
self.stop()
super().init_service(service_name=WeChat.__name__.lower(),
service_type=self._create_client)
self._channel = MessageChannel.Wechat
@@ -54,13 +53,14 @@ class WechatModule(_ModuleBase, _MessageBase[WeChat]):
"""
return 1
def stop(self):
def stop(self) -> None:
"""停止模块"""
for client in self.get_instances().values():
if hasattr(client, "stop"):
try:
try:
if hasattr(client, "stop"):
client.stop()
except Exception as err:
logger.error(f"停止微信模块实例失败:{err}")
except Exception as err:
logger.error(f"停止微信模块实例失败:{err}")
@staticmethod
def _is_bot_mode(config: dict) -> bool:
+6 -8
View File
@@ -23,7 +23,6 @@ class WechatClawBotModule(_ModuleBase, _MessageBase[WechatClawBot]):
def init_module(self) -> None:
"""初始化模块。"""
self.stop()
super().init_service(
service_name=WechatClawBot.__name__.lower(), service_type=WechatClawBot
)
@@ -49,14 +48,13 @@ class WechatClawBotModule(_ModuleBase, _MessageBase[WechatClawBot]):
"""获取模块优先级。"""
return 2
def stop(self):
"""停止模块"""
def stop(self) -> None:
"""停止模块"""
for client in self.get_instances().values():
if hasattr(client, "stop"):
try:
client.stop()
except Exception as err:
logger.error(f"停止微信 ClawBot 模块实例失败:{err}")
try:
client.stop()
except Exception as err:
logger.error(f"停止微信 ClawBot 模块实例失败:{err}")
def test(self) -> Optional[Tuple[bool, str]]:
"""测试模块连接性。"""
+9 -1
View File
@@ -25,7 +25,7 @@ from app.chain.subscribe import SubscribeChain
from app.chain.transfer import TransferChain
from app.chain.workflow import WorkflowChain
from app.core.config import settings, global_vars
from app.core.event import eventmanager
from app.core.event import Event, eventmanager
from app.core.plugin import PluginManager
from app.db import SessionFactory
from app.db.models.downloadhistory import DownloadHistory, DownloadFiles
@@ -987,6 +987,14 @@ class Scheduler(ConfigReloadMixin, metaclass=SingletonClass):
for pid in PluginManager().get_running_plugin_ids():
self.update_plugin_job(pid)
@eventmanager.register(EventType.PluginReload)
def on_plugin_reload(self, event: Event) -> None:
"""插件重载后按当前实例重新注册全部定时服务"""
plugin_id = event.event_data.get("plugin_id")
if not plugin_id:
return
self.update_plugin_job(plugin_id)
def init_workflow_jobs(self):
"""
初始化工作流定时服务
+53 -1
View File
@@ -1,7 +1,7 @@
"""AI智能体相关数据模型"""
from datetime import datetime
from typing import List, Optional, Union
from typing import Any, List, Literal, Optional, Union
from langchain_core.messages import BaseMessage
from pydantic import BaseModel, Field, ConfigDict, field_serializer
@@ -57,6 +57,58 @@ class ToolResult(BaseModel):
error: Optional[str] = Field(default=None, description="错误信息")
class AgentMcpServerConfig(BaseModel):
"""Agent 外部 MCP 服务器配置。"""
id: str = Field(..., description="服务器唯一 ID")
name: str = Field(..., description="服务器显示名称")
enabled: bool = Field(default=True, description="是否启用")
transport: Literal["stdio", "sse", "http", "streamable_http"] = Field(
default="stdio", description="MCP 传输协议"
)
description: Optional[str] = Field(None, description="服务器说明")
command: Optional[str] = Field(None, description="stdio 启动命令")
args: list[str] = Field(default_factory=list, description="stdio 启动参数")
env: dict[str, str] = Field(default_factory=dict, description="stdio 环境变量")
url: Optional[str] = Field(None, description="HTTP/SSE MCP 入口地址")
headers: dict[str, str] = Field(default_factory=dict, description="HTTP 请求头")
timeout: int = Field(default=30, description="连接和调用超时时间(秒)")
tool_prefix: Optional[str] = Field(None, description="注入 Agent 的工具名前缀")
require_admin: bool = Field(default=True, description="是否仅管理员可调用")
class AgentMcpServersSaveRequest(BaseModel):
"""Agent 外部 MCP 服务器保存请求。"""
servers: list[AgentMcpServerConfig] = Field(
default_factory=list, description="MCP 服务器配置列表"
)
class AgentMcpServerTestRequest(BaseModel):
"""Agent 外部 MCP 服务器测试请求。"""
server: AgentMcpServerConfig = Field(..., description="待测试的 MCP 服务器配置")
class AgentMcpServerToolInfo(BaseModel):
"""Agent 外部 MCP 工具摘要。"""
name: str = Field(..., description="原始 MCP 工具名称")
agent_tool_name: str = Field(..., description="注入 Agent 后的工具名称")
description: str = Field(default="", description="工具说明")
input_schema: dict[str, Any] = Field(default_factory=dict, description="工具参数 Schema")
class AgentMcpServerTestResult(BaseModel):
"""Agent 外部 MCP 服务器测试结果。"""
success: bool = Field(..., description="测试是否成功")
message: str = Field(default="", description="测试消息")
tools: list[AgentMcpServerToolInfo] = Field(default_factory=list, description="工具列表")
tool_count: int = Field(default=0, description="工具数量")
class AgentChatAttachment(BaseModel):
"""
Agent 会话展示附件
+2
View File
@@ -175,6 +175,8 @@ class CommingMessage(BaseModel):
message_id: Optional[Union[str, int]] = None
# 聊天ID(用于回调时定位聊天)
chat_id: Optional[str] = None
# 回复目标消息ID(用于 ForceReply 等回复场景)
reply_to_message_id: Optional[Union[str, int]] = None
# 完整的回调查询信息(原始数据)
callback_query: Optional[Dict] = None
# 图片列表(图片URL或file_id
+2
View File
@@ -265,6 +265,8 @@ class SystemConfigKey(Enum):
NotificationSendTime = "NotificationSendTime"
# AI智能体配置
AIAgentConfig = "AIAgentConfig"
# AI智能体外部MCP服务器配置
AIAgentMcpServers = "AIAgentMcpServers"
# 通知消息格式模板
NotificationTemplates = "NotificationTemplates"
# 通知中心清理时间
+32 -17
View File
@@ -1,5 +1,7 @@
import asyncio
import inspect
from contextlib import asynccontextmanager
from typing import Callable
from fastapi import FastAPI
@@ -20,6 +22,7 @@ from app.chain.system import SystemChain
from app.core.config import global_vars, settings
from app.helper.server import MoviePilotServerHelper
from app.helper.system import SystemHelper
from app.log import logger, LoggerManager
from app.startup.command_initializer import init_command, stop_command, restart_command
from app.startup.modules_initializer import init_modules, stop_modules
from app.startup.monitor_initializer import stop_monitor, init_monitor
@@ -55,6 +58,16 @@ async def init_extra():
await MoviePilotServerHelper.async_report_usage()
async def run_shutdown_step(name: str, callback: Callable[[], object]) -> None:
"""隔离单个关闭阶段的异常,确保后续资源仍有机会释放"""
try:
result = callback()
if inspect.isawaitable(result):
await result
except Exception as err:
logger.error(f"关闭{name}失败:{err}")
@asynccontextmanager
async def lifespan(app: FastAPI):
"""
@@ -89,6 +102,7 @@ async def lifespan(app: FastAPI):
yield
finally:
print("Shutting down...")
global_vars.stop_system()
# 取消同步插件任务
try:
sync_plugins_task.cancel()
@@ -97,20 +111,21 @@ async def lifespan(app: FastAPI):
pass
except Exception as e:
print(str(e))
if not settings.MOVIEPILOT_SAFE_MODE:
# 备份插件
SystemChain().backup_plugins()
# 停止工作流
stop_workflow()
# 停止命令
stop_command()
# 停止监控器
stop_monitor()
# 停止定时器
stop_scheduler()
# 停止插件
stop_plugins()
# 停止模块
await stop_modules()
# 关闭共享的异步 HTTP 连接池,释放底层连接资源
await aclose_shared_async_transports()
try:
if not settings.MOVIEPILOT_SAFE_MODE:
await run_shutdown_step(
"插件备份", lambda: SystemChain().backup_plugins()
)
await run_shutdown_step("工作流", stop_workflow)
await run_shutdown_step("命令服务", stop_command)
await run_shutdown_step("监控器", stop_monitor)
await run_shutdown_step("定时器", stop_scheduler)
await run_shutdown_step("插件", stop_plugins)
await run_shutdown_step("模块服务", stop_modules)
await run_shutdown_step(
"共享异步 HTTP 连接池",
aclose_shared_async_transports,
)
finally:
# 日志最后关闭,确保其他组件的收尾信息已写入文件
LoggerManager.shutdown()
+23 -21
View File
@@ -1,4 +1,6 @@
import inspect
import sys
from typing import Callable
from app.helper.redis import RedisHelper, AsyncRedisHelper
@@ -129,27 +131,27 @@ async def stop_modules():
"""
服务关闭
"""
# 停止AI智能体
await stop_agent()
# 停止模块
ModuleManager().stop()
# 停止事件消费
EventManager().stop()
# 停止虚拟显示
DisplayHelper().stop()
# 停止线程池
ThreadHelper().shutdown()
# 停止消息服务
stop_message()
# 关闭Redis缓存连接
RedisHelper().close()
await AsyncRedisHelper().close()
# 停止数据库连接
await close_database()
# 停止前端服务
stop_frontend()
# 清理临时文件
clear_temp()
async def run_step(name: str, callback: Callable[[], object]) -> None:
"""单个模块资源关闭失败时继续执行后续阶段"""
try:
result = callback()
if inspect.isawaitable(result):
await result
except Exception as err:
logger.error(f"关闭{name}失败:{err}")
await run_step("AI智能体", stop_agent)
await run_step("模块", lambda: ModuleManager().stop())
await run_step("事件消费", lambda: EventManager().stop())
await run_step("虚拟显示", lambda: DisplayHelper().stop())
await run_step("DoH服务", lambda: DohHelper().shutdown())
await run_step("线程池", lambda: ThreadHelper().shutdown())
await run_step("消息服务", stop_message)
await run_step("Redis缓存连接", lambda: RedisHelper().close())
await run_step("异步Redis缓存连接", lambda: AsyncRedisHelper().close())
await run_step("数据库连接", close_database)
await run_step("前端服务", stop_frontend)
await run_step("临时文件", clear_temp)
def init_modules():
+34 -8
View File
@@ -9,6 +9,8 @@ fixture 一并识别,autouse 自动作用于每个用例,无需逐用例改
"""
from __future__ import annotations
import ipaddress
import pytest
# 本地回环/通配地址放行,其余主机一律视为真实出站;getaddrinfo 的 host 可能为 str 或 bytes
@@ -20,21 +22,45 @@ def block_real_network(monkeypatch):
"""防御纵深:拦截对非本地主机的真实出站,强制测试零真实网络。
补在各用例自身 mock 之上某用例万一漏 mock 外部依赖TMDB / LLM 目录 / 下载器 /
媒体服务器 / 任意外链真实 DNS 解析会在此被拦并报错而非静默发请求本地回环放行
sqlite asyncio 默认解析器经线程池调用 ``socket.getaddrinfo``故拦此一处即覆盖
同步与异步出站``monkeypatch`` 在用例结束后自动还原不影响其他用例与进程退出
媒体服务器 / 任意外链 DNS 解析 socket 连接会被拦截本地回环放行sqlite
所有拦截记录会在用例收尾再次断言避免业务代码捕获网络异常后让漏 mock 的用例静默通过
``monkeypatch`` 在用例结束后自动还原不影响其他用例与进程退出
"""
import socket
_real_getaddrinfo = socket.getaddrinfo
_real_connect = socket.socket.connect
attempts = []
def _is_allowed_host(host) -> bool:
normalized = host.decode() if isinstance(host, (bytes, bytearray)) else host
if normalized is None or normalized in _ALLOWED_NETWORK_HOSTS:
return True
try:
address = ipaddress.ip_address(str(normalized).split("%", 1)[0])
return address.is_loopback or address.is_unspecified
except ValueError:
return False
def _blocked(operation: str, host):
attempts.append((operation, host))
raise RuntimeError(
f"测试禁止真实出站网络:尝试通过 {operation} 访问 {host!r};请 mock 对应外部依赖"
)
def _guarded_getaddrinfo(host, *args, **kwargs):
normalized = host.decode() if isinstance(host, (bytes, bytearray)) else host
if normalized is not None and normalized not in _ALLOWED_NETWORK_HOSTS:
raise RuntimeError(
f"测试禁止真实出站网络:尝试解析 {normalized!r};请 mock 对应外部依赖"
)
if not _is_allowed_host(host):
_blocked("DNS", host)
return _real_getaddrinfo(host, *args, **kwargs)
def _guarded_connect(sock, address):
if isinstance(address, tuple) and address and not _is_allowed_host(address[0]):
_blocked("socket", address[0])
return _real_connect(sock, address)
monkeypatch.setattr(socket, "getaddrinfo", _guarded_getaddrinfo)
monkeypatch.setattr(socket.socket, "connect", _guarded_connect)
yield
if attempts:
details = ", ".join(f"{operation}:{host}" for operation, host in attempts)
pytest.fail(f"测试期间发生真实出站网络尝试:{details}")
+26 -7
View File
@@ -76,6 +76,14 @@ _REQUESTS_RETRY_IDEMPOTENT_METHODS = ("GET", "HEAD", "OPTIONS")
_pending_eviction_tasks: set[asyncio.Task] = set()
def _discard_pending_eviction_task(task: asyncio.Task) -> None:
"""从跨线程共享集合移除已完成的 transport 关闭任务"""
with _shared_async_transports_lock:
_pending_eviction_tasks.discard(task)
if not task.cancelled() and (error := task.exception()):
logger.debug(f"LRU 淘汰共享 transport 时关闭失败: {error!r}")
def _get_shared_async_transport(
proxy: Optional[str],
verify: Union[bool, str],
@@ -140,8 +148,9 @@ def _get_shared_async_transport(
try:
task = loop.create_task(evicted_transport.aclose())
# 强引用避免 task 仅被 loop 弱持有而触发 "Task was destroyed but pending"
_pending_eviction_tasks.add(task)
task.add_done_callback(_pending_eviction_tasks.discard)
with _shared_async_transports_lock:
_pending_eviction_tasks.add(task)
task.add_done_callback(_discard_pending_eviction_task)
except Exception as e: # pragma: no cover - 防御性
logger.debug(f"LRU 淘汰共享 transport 时调度关闭失败: {e!r}")
@@ -160,17 +169,27 @@ async def aclose_shared_async_transports() -> None:
# 弹出而非 get+clear,避免外层 dict 残留空 OrderedDict 占位
with _shared_async_transports_lock:
per_loop = _shared_async_transports.pop(loop, None)
if not per_loop:
pending_evictions = [
task
for task in _pending_eviction_tasks
if task.get_loop() is loop
]
transports = list(per_loop.values()) if per_loop else []
if per_loop:
per_loop.clear()
if not transports and not pending_evictions:
return
transports = list(per_loop.values())
per_loop.clear()
# 并行关闭:每个 transport 的 TLS close_notify 各占一个 RTT
# 顺序等待会线性放大 shutdown 耗时;return_exceptions 让单点失败
# 不影响其他 transport 的释放
results = await asyncio.gather(
*(t.aclose() for t in transports), return_exceptions=True
*pending_evictions,
*(t.aclose() for t in transports),
return_exceptions=True,
)
for result in results:
with _shared_async_transports_lock:
_pending_eviction_tasks.difference_update(pending_evictions)
for result in results[len(pending_evictions):]:
if isinstance(result, BaseException):
logger.debug(f"关闭共享 AsyncHTTPTransport 失败: {result!r}")
+9
View File
@@ -10,6 +10,11 @@ class Singleton(abc.ABCMeta, type):
_instances: dict = {}
def get_existing_instance(cls, *args, **kwargs):
"""按相同参数返回已创建实例,不触发初始化"""
key = (cls, args, frozenset(kwargs.items()))
return cls._instances.get(key)
def __call__(cls, *args, **kwargs):
key = (cls, args, frozenset(kwargs.items()))
if key not in cls._instances:
@@ -31,6 +36,10 @@ class SingletonClass(abc.ABCMeta, type):
_instances: dict = {}
def get_existing_instance(cls):
"""返回已创建实例,不触发初始化"""
return cls._instances.get(cls)
def __call__(cls, *args, **kwargs):
if cls not in cls._instances:
cls._instances[cls] = super().__call__(*args, **kwargs)
+80
View File
@@ -0,0 +1,80 @@
"""2.2.11
新增下载失败资源冷却表
Revision ID: b7d4a9c2e6f1
Revises: 8ab72c49d1e3
Create Date: 2026-07-07
"""
from alembic import op
import sqlalchemy as sa
revision = "b7d4a9c2e6f1"
down_revision = "8ab72c49d1e3"
branch_labels = None
depends_on = None
def _has_table(inspector: sa.Inspector, table_name: str) -> bool:
"""检查数据表是否已存在。"""
return table_name in inspector.get_table_names()
def upgrade() -> None:
"""升级数据库结构。"""
inspector = sa.inspect(op.get_bind())
if _has_table(inspector, "downloadfailure"):
return
op.create_table(
"downloadfailure",
sa.Column("id", sa.Integer(), nullable=False),
sa.Column("fingerprint", sa.String(), nullable=False),
sa.Column("type", sa.String(), nullable=True),
sa.Column("title", sa.String(), nullable=True),
sa.Column("year", sa.String(), nullable=True),
sa.Column("tmdbid", sa.Integer(), nullable=True),
sa.Column("doubanid", sa.String(), nullable=True),
sa.Column("seasons", sa.String(), nullable=True),
sa.Column("episodes", sa.String(), nullable=True),
sa.Column("site", sa.Integer(), nullable=True),
sa.Column("site_name", sa.String(), nullable=True),
sa.Column("torrent_id", sa.String(), nullable=True),
sa.Column("torrent_name", sa.String(), nullable=True),
sa.Column("torrent_size", sa.Float(), nullable=True),
sa.Column("downloader", sa.String(), nullable=True),
sa.Column("source", sa.String(), nullable=True),
sa.Column("error_message", sa.String(), nullable=True),
sa.Column("retry_count", sa.Integer(), nullable=True),
sa.Column("first_failed_at", sa.String(), nullable=True),
sa.Column("last_failed_at", sa.String(), nullable=True),
sa.Column("next_retry_at", sa.String(), nullable=True),
sa.PrimaryKeyConstraint("id"),
)
op.create_index(
"ux_downloadfailure_fingerprint",
"downloadfailure",
["fingerprint"],
unique=True,
)
op.create_index(
"ix_downloadfailure_next_retry_at",
"downloadfailure",
["next_retry_at"],
)
op.create_index(
"ix_downloadfailure_media_site",
"downloadfailure",
["type", "tmdbid", "doubanid", "site"],
)
def downgrade() -> None:
"""回滚数据库结构。"""
inspector = sa.inspect(op.get_bind())
if not _has_table(inspector, "downloadfailure"):
return
op.drop_index("ix_downloadfailure_media_site", table_name="downloadfailure")
op.drop_index("ix_downloadfailure_next_retry_at", table_name="downloadfailure")
op.drop_index("ux_downloadfailure_fingerprint", table_name="downloadfailure")
op.drop_table("downloadfailure")
+119 -11
View File
@@ -20,6 +20,18 @@ function WARN() {
echo -e "${WARN} ${1}"
}
ENTRYPOINT_START_TIME="$(date +%s)"
function normalize_env_value() {
printf '%s' "${1:-}" | tr '[:upper:]' '[:lower:]'
}
function is_truthy_value() {
local value
value="$(normalize_env_value "${1:-}")"
[ "${value}" = "true" ] || [ "${value}" = "1" ] || [ "${value}" = "yes" ]
}
# 设置虚拟环境路径(兼容群晖等系统必须这样配置)
VENV_PATH="${VENV_PATH:-/opt/venv}"
export PATH="${VENV_PATH}/bin:$PATH"
@@ -47,6 +59,42 @@ function run_package_command() {
fi
}
function wait_backend_ready() {
local entrypoint_start_time="${1:-$(date +%s)}"
local backend_start_time="${2:-$(date +%s)}"
local python_pid="${3:-}"
local backend_port="${PORT:-3001}"
local web_port="${NGINX_PORT:-3000}"
local timeout="${MOVIEPILOT_BACKEND_READY_TIMEOUT:-300}"
local ready_url="http://127.0.0.1:${backend_port}/api/v1/system/global?token=moviepilot"
local deadline
if ! [[ "${timeout}" =~ ^[0-9]+$ ]] || [ "$((10#${timeout}))" -le 0 ]; then
WARN "→ MOVIEPILOT_BACKEND_READY_TIMEOUT=${timeout} 无效,使用默认 300 秒。"
timeout=300
else
timeout=$((10#${timeout}))
fi
deadline=$(( $(date +%s) + timeout ))
while [ "$(date +%s)" -lt "${deadline}" ]; do
if [ -n "${python_pid}" ] && ! kill -0 "${python_pid}" >/dev/null 2>&1; then
WARN "→ 后端服务启动完成探测已停止:后端进程已退出。"
return 1
fi
if curl -fsS --max-time 2 "${ready_url}" >/dev/null 2>&1; then
local now
now="$(date +%s)"
INFO "→ MoviePilot Web 已可访问,启动总耗时 $(( now - entrypoint_start_time )) 秒,后端就绪耗时 $(( now - backend_start_time )) 秒,后端端口 ${backend_port},前端端口 ${web_port}"
return 0
fi
sleep 1
done
WARN "→ 后端服务启动完成探测超时,已等待 ${timeout} 秒,后端端口 ${backend_port},继续等待进程日志..."
return 1
}
# 环境变量补全
# 优先级: 系统环境变量 -> .env 文件 (即使为空字符串) -> 预设默认值
# 精准适配 Python 端 set_key (quote_mode="always", 单引号包裹, \' 转义)
@@ -65,6 +113,7 @@ function load_config_from_app_env() {
["GITHUB_TOKEN"]=""
["MOVIEPILOT_AUTO_UPDATE"]="release"
["MOVIEPILOT_DOCKER_KEEPALIVE_ON_FAILURE"]="true"
["MOVIEPILOT_FORCE_CHOWN"]="false"
["MOVIEPILOT_SAFE_MODE"]="false"
["BROWSER_EMULATION"]="cloakbrowser"
@@ -261,8 +310,8 @@ function graceful_exit() {
# 后端异常退出时默认保留容器,避免无法 docker exec 进入容器运行 doctor。
function diagnostic_keepalive() {
local exit_code=${1:-1}
local keepalive="${MOVIEPILOT_DOCKER_KEEPALIVE_ON_FAILURE:-true}"
keepalive="${keepalive,,}"
local keepalive
keepalive="$(normalize_env_value "${MOVIEPILOT_DOCKER_KEEPALIVE_ON_FAILURE:-true}")"
if [ "${keepalive}" = "false" ] || [ "${keepalive}" = "0" ] || [ "${keepalive}" = "no" ]; then
graceful_exit "$exit_code" "python_exit"
@@ -315,6 +364,70 @@ function ensure_backend_runtime_dependencies() {
INFO "→ 已自动恢复主程序依赖,继续启动后端。"
}
function path_owner_id() {
local target="${1:-}"
[ -n "${target}" ] || return 0
stat -c '%u:%g' "${target}" 2>/dev/null || stat -f '%u:%g' "${target}" 2>/dev/null || true
}
function force_chown_image_paths_if_requested() {
if ! is_truthy_value "${MOVIEPILOT_FORCE_CHOWN:-false}"; then
return 0
fi
WARN "→ MOVIEPILOT_FORCE_CHOWN 已启用,将递归修复 /app、/public 权限,可能显著增加启动耗时。"
local path
for path in "$@"; do
[ -e "${path}" ] || continue
chown -R moviepilot:moviepilot "${path}"
done
}
function correct_home_permissions() {
[ -e "${HOME}" ] || return 0
chown moviepilot:moviepilot "${HOME}"
[ -e "${HOME}/.cloakbrowser" ] && chown -h moviepilot:moviepilot "${HOME}/.cloakbrowser"
if is_truthy_value "${MOVIEPILOT_FORCE_CHOWN:-false}"; then
[ -e "${HOME}/.cloakbrowser" ] && chown -R moviepilot:moviepilot "${HOME}/.cloakbrowser"
elif [ -e "${HOME}/.cloakbrowser" ]; then
INFO "→ 默认跳过 ${HOME}/.cloakbrowser 递归权限校正,如遇浏览器缓存权限错误可设置 MOVIEPILOT_FORCE_CHOWN=true 后重启一次。"
fi
find "${HOME}" -mindepth 1 -maxdepth 1 ! -name ".cloakbrowser" -exec chown -R moviepilot:moviepilot {} +
}
function chown_plugin_runtime_path() {
local plugin_path="${1:-}"
[ -n "${plugin_path}" ] || return 0
[ -e "${plugin_path}" ] || return 0
local current_owner
current_owner="$(path_owner_id "${plugin_path}")"
[ "${current_owner}" = "${PUID}:${PGID}" ] && return 0
chown -h moviepilot:moviepilot "${plugin_path}"
}
function correct_file_permissions() {
local chown_start
local chown_end
chown_start=$(date +%s)
INFO "→ 正在校正文件权限..."
force_chown_image_paths_if_requested /app /public
chown_plugin_runtime_path /app/app/plugins
correct_home_permissions
chown -R moviepilot:moviepilot \
"${CONFIG_DIR}" \
/var/lib/nginx \
/var/log/nginx
chown moviepilot:moviepilot /etc/hosts /tmp
chown_end=$(date +%s)
INFO "→ 文件权限校正完成,耗时 $(( chown_end - chown_start )) 秒。"
}
# 使用env配置
load_config_from_app_env
apply_package_cache_env
@@ -354,14 +467,7 @@ groupmod -o -g "${PGID}" moviepilot
usermod -o -u "${PUID}" moviepilot
# 更改文件权限
chown -R moviepilot:moviepilot \
"${HOME}" \
/app \
/public \
"${CONFIG_DIR}" \
/var/lib/nginx \
/var/log/nginx
chown moviepilot:moviepilot /etc/hosts /tmp
correct_file_permissions
# 启动前优先确认主运行环境仍然健康,避免插件依赖污染导致服务直接起不来。
ensure_backend_runtime_dependencies
@@ -369,7 +475,7 @@ ensure_backend_runtime_dependencies
# 下载浏览器内核
function install_browser_kernel() {
local emulation="${BROWSER_EMULATION:-cloakbrowser}"
emulation="${emulation,,}"
emulation="$(normalize_env_value "${emulation}")"
local proxy="${HTTPS_PROXY:-${https_proxy:-$PROXY_HOST}}"
if [ "${emulation}" != "cloakbrowser" ] && [ "${emulation}" != "flaresolverr" ] && [ -n "${emulation}" ]; then
@@ -412,12 +518,14 @@ umask "${UMASK}"
# 启动后端服务
INFO "→ 启动后端服务..."
BACKEND_START_TIME="$(date +%s)"
if [ "${START_NOGOSU:-false}" = "true" ]; then
"${VENV_PATH}/bin/python3" app/main.py > /dev/stdout 2> /dev/stderr &
else
gosu moviepilot:moviepilot "${VENV_PATH}/bin/python3" app/main.py > /dev/stdout 2> /dev/stderr &
fi
PYTHON_PID=$!
wait_backend_ready "${ENTRYPOINT_START_TIME}" "${BACKEND_START_TIME}" "${PYTHON_PID}" &
# 等待 Python 进程退出。
# 如果收到信号,trap 会中断 wait,并执行 graceful_exit。
+33
View File
@@ -67,6 +67,26 @@ MCP 使用系统配置中的 `API_TOKEN` 作为认证密钥,文档中的 API K
}
```
## 4.1 Agent 外部 MCP Client 配置
MoviePilot 的内置 Agent 也可以作为 MCP Client 连接外部 MCP 服务器,将外部工具注入到智能助手工具列表中。当前支持:
- `stdio`:按配置的命令和参数启动本地 MCP 进程,通过标准输入输出交换 JSON-RPC 消息。
- `sse`:连接旧版 HTTP+SSE MCP 服务,先读取 `endpoint` 事件,再向返回的 endpoint POST JSON-RPC 消息。
- `http` / `streamable_http`:连接 Streamable HTTP MCP 服务,直接向配置 URL POST JSON-RPC 消息。
这些配置是管理员级 Agent 运行时配置,保存在 `SystemConfigKey.AIAgentMcpServers` 中。外部 MCP 工具默认要求管理员上下文调用,避免普通用户触发高权限外部工具。
### Agent MCP 配置接口
这些接口使用登录态鉴权,并要求当前用户为超级管理员。
| 方法 | 路径 | 说明 |
| :--- | :--- | :--- |
| GET | `/api/v1/message/agent/mcp/servers` | 查询已配置的外部 MCP 服务器列表 |
| POST | `/api/v1/message/agent/mcp/servers` | 保存外部 MCP 服务器列表 |
| POST | `/api/v1/message/agent/mcp/servers/test` | 测试单个外部 MCP 服务器,返回发现的工具列表 |
## 5. 错误码说明
| 错误码 | 消息 | 说明 |
@@ -133,6 +153,19 @@ FastAPI 异常响应保留 `detail` 字段,并在错误详情为文本时返
| POST | `/api/v1/system/setting/PLUGIN_MARKET/sync-wiki` | 管理员从 MoviePilot Wiki 的插件文档同步公开插件仓库清单,和本地 `PLUGIN_MARKET` 合并去重后写入配置 |
| GET | `/api/v1/system/modulelist` | 查询已加载模块,保留 `name` 原始中文字段,并提供 `name_i18n``name_key` 给多语言前端展示 |
| GET | `/api/v1/system/moduletest/{moduleid}` | 测试指定模块可用性,保留原 `message`,并在标准响应顶层返回 `message_i18n` |
| GET | `/api/v1/message/agent/mcp/servers` | 管理员查询 Agent 外部 MCP 服务器配置 |
| POST | `/api/v1/message/agent/mcp/servers` | 管理员保存 Agent 外部 MCP 服务器配置 |
| POST | `/api/v1/message/agent/mcp/servers/test` | 管理员测试单个 Agent 外部 MCP 服务器并读取工具列表 |
#### 缓存管理
以下接口使用登录态鉴权,并要求当前用户为超级管理员。
| 方法 | 路径 | 说明 |
| :--- | :--- | :--- |
| GET | `/api/v1/tmdb/cache` | 查询 TheMovieDb 识别缓存及识别成功、失败条目统计 |
| DELETE | `/api/v1/tmdb/cache/{cache_key}` | 按缓存键删除单条 TheMovieDb 识别缓存,缓存键需要进行 URL 编码 |
| DELETE | `/api/v1/tmdb/cache` | 清空全部 TheMovieDb 识别缓存 |
### 插件补充接口
+23 -106
View File
@@ -1,131 +1,48 @@
# PR-Agent 使用说明
本仓库通过 GitHub Actions 运行开源 PR-Agent用于自动生成 PR 说明和进行 AI Review
本仓库通过 PR Review Runner 运行 PR-Agent帮助贡献者维护 PR 摘要、获取代码审查结果和提出 PR 相关问题
## Secrets
## 自动执行
在仓库的 `Settings -> Secrets and variables -> Actions -> Repository secrets` 中配置:
同仓分支和来自 fork 的 PR 都会自动执行 PR-Agent。
- `OPENAI_KEY`OpenAI 或 OpenAI 兼容服务的 API Key。
- `OPENAI_API_BASE`OpenAI 兼容接口的 API Base,通常需要包含 `/v1`,以服务商文档为准。
PR 在以下场景会自动处理:
`GITHUB_TOKEN` 使用 GitHub Actions 自动注入的 `${{ github.token }}`,不需要手工添加
- 打开或重新打开 PR
- 将草稿 PR 标记为可审查。
- 请求审查。
- 每次推送新的 commit。
## 触发方式
PR 带有 `skip pr-agent` 标签,或标题以 `[Auto]``Auto` 开头时,自动和手工路径都会跳过。
`.github/workflows/pr-agent.yml` 监听:
## 手工命令
- `pull_request_target`:PR 打开、重新打开、标记 ready、请求 review、推送新 commit 时自动运行。
- `issue_comment`:允许身份在 PR 评论里写允许的命令时手动运行。
PR 事件会自动执行受控审查,包含同仓 PR 和 fork PR。允许身份也可以在 PR 评论中使用允许的命令触发受控审查。
允许身份包括 `OWNER``MEMBER``COLLABORATOR``CONTRIBUTOR``FIRST_TIME_CONTRIBUTOR`
## Workflow 权限
workflow 设置了最小可用权限:
- `contents: read`:读取仓库内容和 PR diff。
- `pull-requests: write`:更新 PR 描述、发布 PR Review 或修改 PR 相关元数据。
- `issues: write`PR 评论在 GitHub API 中属于 issue comments,手动命令和总结评论需要该权限。
没有开启 `contents: write`。当前配置不让 PR-Agent 往仓库推代码或提交 changelog,因此不需要内容写权限。
默认自动执行:
- `/review`:检查 PR 风险、潜在 bug、安全问题、测试缺口和可维护性问题。
- `/describe`:生成或更新 PR 描述、变更摘要和文件说明。
默认不自动执行:
- `/improve`:给出代码改进建议。这个工具更容易产生噪音和额外成本,建议先用评论命令手动触发。
## 常用评论命令
以下身份可在 PR 评论中使用:
- `OWNER`:仓库所有者。
- `MEMBER`:组织仓库中的组织成员。
- `COLLABORATOR`:仓库协作者。
- `CONTRIBUTOR`:曾经向仓库提交并合入过代码的贡献者。
- `FIRST_TIME_CONTRIBUTOR`:首次向仓库贡献 PR 的用户。
在 PR 的普通讨论评论中使用以下命令:
```text
/review
/describe
/improve
/review
/ask 这次改动有没有遗漏权限校验?
```
评论触发依赖 `issue_comment` 事件。普通 issue 评论、Bot 评论、非允许身份评论、以及不以允许命令开头的评论都会跳过
- `/describe`:更新 PR Body 内的 `PR-Agent 摘要`,并保留贡献者原有的 PR 描述
- `/review`:发起一次代码审查。
- `/ask ...`:就当前 PR 提问,回复会发布在普通 PR 评论中。
## 配置来源
本仓库禁用 `/improve` 及其等价别名。其他命令是否可用由 runner 所包含的 PR-Agent 能力决定。
PR-Agent 配置集中在 `.github/workflows/pr-agent.yml``env` 中维护
手工命令仅允许以下 GitHub 身份关联的用户使用:`OWNER``MEMBER``COLLABORATOR``CONTRIBUTOR``FIRST_TIME_CONTRIBUTOR`
当前主要设置:
新建的合法命令评论会触发执行;编辑后仍为合法命令的评论也会触发。编辑普通讨论评论不会调用模型。
- `config.model = "gpt-5.5"`:默认使用 GPT-5.5。
- `config.fallback_models = ["gpt-5.4"]`:主模型不可用时降级到 GPT-5.4。
- `config.reasoning_effort = "xhigh"`:使用更高审查推理强度。
- `config.ai_timeout = "900"`:模型调用最长等待 900 秒。
- `config.response_language = "zh-CN"`:让 PR-Agent 默认中文输出。
- `config.large_patch_policy = "clip"`:大 PR 截断分析,不直接跳过。
- `config.ignore_pr_title` / `config.ignore_pr_labels`:跳过自动生成 PR 或带 `skip pr-agent` 标签的 PR。
- `pr_reviewer.extra_instructions`:要求中文输出,优先指出 P0/P1 风险,并关注安全、权限、状态一致性、异步/缓存、副作用和测试缺口。
- `pr_reviewer.require_security_review = true`:要求输出安全审查部分。
- `pr_reviewer.require_tests_review = true`:要求输出测试审查部分。
- `pr_reviewer.enable_review_labels_effort = false`:不添加 `Review effort x/5` 工作量标签。
- `pr_reviewer.enable_review_labels_security = true`:保留明确安全风险标签。
- `pr_description.generate_ai_title = false`:默认不改 PR 标题。
- `pr_description.publish_labels = false`:默认不添加 PR 类型标签。
- `pr_description.enable_pr_diagram = false`:默认不生成图表。
- `pr_code_suggestions.focus_only_on_problems = true`:手动 `/improve` 时优先输出问题型建议。
- `pr_code_suggestions.suggestions_score_threshold = 7`:过滤低置信度建议。
## 审查结果
标签来源:
本仓库固定使用中文生成 PR-Agent 内容。`/describe` 的结果位于 PR Body 的 `PR-Agent 摘要` 区域,用于概览本次变更。
- `/review` 可添加安全标签和工作量标签;当前只保留安全标签,关闭工作量标签
- `/describe` 可按 PR 类型添加 `Bug fix``Tests``Bug fix with tests``Enhancement``Documentation``Other` 等标签;当前 `pr_description.publish_labels = false`,不会添加类型标签。
- 自定义标签默认未启用。
`/review` 和自动审查会通过原生 GitHub Review 发布,结果位于 Review 页签,标题固定为 `PR-Agent Code Review`。可定位到本次变更的问题会在对应代码行以行内评论呈现,并使用 high、medium 或 low 风险标识
可按需再启用的工具配置:
- `[pr_update_changelog]`:配合 `/update_changelog` 生成 changelog 建议。
- `[pr_add_docs]`:配合 `/add_docs` 生成文档建议。
- `[pr_test]`:配合 `/test` 生成测试建议;它不会替代仓库自己的测试命令。
- `[pr_questions]`:配合 `/ask ...` 回答 PR 相关问题。
Review 摘要会自然概括本次变更和整体审查结论,不会复制行内评论。未发现需要处理的问题时,摘要会概括变更并自然说明暂无其他反馈。审查不会额外创建专用的 issue comment 摘要。
## 安全边界
PR-Agent Action 会读取 `OPENAI_KEY`,因此依赖的 Docker 镜像在 workflow 中固定版本号和 digest
不使用浮动的 `latest` 或仅依赖可变 tag。
当前使用 `pull_request_target` 支持 fork PR 自动审查,但 workflow 不 checkout 或执行来自 fork 的代码,
只运行固定 digest 的 PR-Agent 容器并通过 GitHub API 读取 PR diff。`issue_comment` 属于 base repo
事件,因此评论命令只允许指定身份触发。
API Key 建议使用低额度、可轮换的专用 key。`OPENAI_API_BASE` 本身通常不是敏感信息,但继续按 secret 管理可以避免暴露服务商信息。
## 调整自动行为
自动行为在 workflow 的 `env` 中控制:
```yaml
config.model: "gpt-5.5"
config.fallback_models: '["gpt-5.4"]'
config.reasoning_effort: "xhigh"
config.ai_timeout: "900"
config.response_language: "zh-CN"
github_action_config.auto_review: "true"
github_action_config.auto_describe: "true"
github_action_config.auto_improve: "false"
github_action_config.pr_actions: '["opened", "reopened", "ready_for_review", "review_requested", "synchronize"]'
pr_description.generate_ai_title: "false"
pr_description.publish_labels: "false"
pr_description.enable_pr_diagram: "false"
pr_reviewer.enable_review_labels_effort: "false"
pr_reviewer.enable_review_labels_security: "true"
```
如果需要自动运行 `/improve`,把 `github_action_config.auto_improve` 改为 `"true"`。建议先观察手动 `/improve` 的质量和成本,再决定是否开启。
自动审查通过 `pull_request_target` 在目标仓库上下文中读取 PR 信息,并使用共享 runner 的 `latest` 镜像,但不会 checkout 或执行 PR 分支代码。权限保持最小化:只授予读取仓库内容所需的 `contents: read`,以及更新 PR Body、发布 Review 和回复 PR 评论所需的写权限;不会向仓库推送代码或创建提交。
+17 -1
View File
@@ -269,4 +269,20 @@ moviepilot help tool
moviepilot help scheduler
```
*Last Updated: 2026-05-25*
---
## Site Adapter Capture — macOS / Linux
```bash
# Run from a MoviePilot source checkout and reuse its virtual environment
bash scripts/collect-site-adapter.sh
```
**Rules:**
- The default collector asks only for the site HTTPS address, opens an isolated local Chrome/Edge profile, and reads the completed search page after the user confirms.
- Users must not be asked to inspect HTML or copy Cookie/User-Agent values in the default flow. `--manual-cookie` is an advanced fallback only.
- Run only the collector shipped with a trusted local MoviePilot source checkout or installation package. Do not pipe a remote branch script into a shell.
- Never put a Cookie or other credential in command arguments or shell history.
- Feature Request attachments are public. Review all four files in the generated ZIP before attaching it, and never attach raw HTML, HAR, or browser network archives.
*Last Updated: 2026-07-12*
+4 -2
View File
@@ -10,6 +10,8 @@ All **public classes**, **public methods**, and **public functions** in this pro
## Docstring Format
Short, label-style docstrings should follow the surrounding code style and must not gain a period mechanically. Complete sentences that explain non-obvious behavior should use normal Chinese punctuation.
### Single-line (for simple, obvious descriptions)
```python
@@ -28,7 +30,7 @@ def download(
download_dir: Path,
) -> Optional[str]:
"""
添加下载任务到下载器
添加下载任务到下载器
:param context: 当前媒体上下文,包含识别结果和种子选择信息
:param torrent: 要下载的种子信息
@@ -43,7 +45,7 @@ def download(
```python
class DownloadChain(ChainBase):
"""
下载处理链,负责协调搜索结果的种子选择、下载器调度和下载后处理
下载处理链,负责协调搜索结果的种子选择、下载器调度和下载后处理
"""
```
+60
View File
@@ -0,0 +1,60 @@
# 站点适配采集
当开发者没有目标站点账号时,可以由已有账号的用户在本地采集一份经过裁剪和脱敏的搜索页结构,并把采集 ZIP 附加到站点适配 Feature Request。开发者和自动化流程只处理脱敏包,不需要获取用户账号。
## 普通用户一键采集
普通用户只需准备两样东西:目标站点账号,以及已安装的 Chrome、Edge 或 Chromium。无需安装 Python、Git、MoviePilot、Docker,也不需要查看 HTML、复制 Cookie 或填写 User-Agent。
1. 从 MoviePilot 官方 Release 下载与 Windows、macOS 或 Linux 对应的 `moviepilot-site-collector-*` 单文件采集器。
2. 运行采集器,只输入站点首页地址,例如 `https://tracker.example.com`
3. 程序会打开一个临时浏览器窗口。在这个窗口里正常登录站点,搜索一个能返回至少 3 条结果的常见关键词,并保持搜索结果页打开。
4. 回到采集器按回车。程序会自动识别搜索地址、关键词、Cookie 和 User-Agent,完成本地裁剪与脱敏后生成 ZIP。
临时浏览器使用独立的一次性用户目录,不会读取日常浏览器的历史登录状态。采集完成后程序会关闭临时浏览器并清理这次登录数据。原始页面和 Cookie 只在内存中处理,不会写入采集包。
各系统下载文件和首次运行方式见 [站点适配采集器下载说明](site-adapter-collector-release.md)。
## 开发者源码入口
已经有 MoviePilot 源码和 Python 环境的开发者,也可以在项目目录执行:
```bash
bash scripts/collect-site-adapter.sh
```
源码入口和独立程序使用同一套浏览器采集流程。只有排查兼容问题时才使用 `--manual-cookie` 高级模式;普通用户不需要接触 Cookie。
## 第一版限制
当前采集器会读取浏览器渲染后的页面,因此可由用户在临时窗口中完成验证码、Cloudflare 检查和普通登录。但自动适配协议仍要求搜索结果地址能够表示为 HTTPS GET URL。以下场景第一版不做自动适配:
- 必须 POST 表单才能搜索的站点。
- 搜索完成后地址栏完全没有关键词或可复用搜索参数的站点。
- 需要专用 API、复杂签名或无法从一次搜索结果页观察出 FREE/HR 规则的站点。
遇到这些场景时,请在 Feature Request 中说明失败步骤和终端错误文字,等待人工确认采集方案。不要用原始 HAR、原始 HTML 或包含账号信息的截图替代脱敏 ZIP。
## 脱敏范围
采集器只接受 HTTPS 地址。默认模式从本机临时浏览器只读当前搜索结果页;高级手动模式设置 30 秒超时和 5 MiB 响应上限,并禁止携带 Cookie 跨 origin 重定向。读取页面后会在本机完成以下处理:
- 只保留种子列表、表头和最多 25 条结果相关 DOM,丢弃账号导航、页脚、脚本、样式、隐藏表单和其他页面内容。
- 替换种子标题、用户名、邮箱、IP、时间、大小、统计值和长随机标识,只保留适配所需的标签、class、字段与链接结构。
- URL 转为同源相对路径,查询值统一替换为占位符,凭据语义的字段直接移除。
- Cookie 和浏览器 UA 仅用于本次请求,不写入采集包;写入前还会使用 Cookie 原值执行二次泄露检查。
输出 ZIP 根目录固定包含:
- `manifest.json`:包版本、站点标识、采集时间、结果行数、HTML 摘要和隐私声明。
- `request.json`:仅包含 GET、origin、相对路径和脱敏后的查询参数,搜索值固定为 `{keyword}`
- `search.html`:本地裁剪并脱敏后的种子列表结构。
- `redaction-report.json`:脱敏状态和各类处理计数。
## 提交 Feature Request
在 GitHub 创建“功能改进” Issue,类型选择“站点适配”,然后把生成的 `moviepilot-site-capture-*.zip` 直接拖入“站点适配采集文件”输入框。
Feature Request 及其附件是公开内容。提交前请先在本地解压 ZIP,确认根目录只有上述四个文件,并逐一预览确认没有站点账号、搜索隐私或其他不希望公开的信息。
站点适配请求必须附加采集器生成并人工复核过的 ZIP。严禁手工上传 Cookie、Authorization、通行密钥、会话字段、原始 HTML、原始 HAR、浏览器网络归档或截图中的账号信息。如果采集失败,请只提交终端错误文字,不要用任何原始数据代替脱敏包。
+36
View File
@@ -0,0 +1,36 @@
# 站点适配采集器下载说明
普通用户优先使用 MoviePilot 正式 Release 提供的单文件采集器。单文件已经包含 Python 和采集器依赖,不需要安装 Python、pip、Git、MoviePilot 后端,也不需要下载源码。电脑只需已安装 Chrome、Edge 或 Chromium 浏览器。
## 选择下载文件
请只从 MoviePilot 官方 GitHub Release 下载与系统匹配的文件:
| 系统 | 下载文件 | 用户侧运行环境 |
|---|---|---|
| Windows | `moviepilot-site-collector-windows.exe` | Chrome、Edge 或 Chromium |
| macOS | `MoviePilot-Site-Collector-macOS.zip` | Chrome、Edge 或 Chromium |
| Linux | `moviepilot-site-collector-linux` | Chrome、Edge 或 Chromium |
每个程序旁边还有同名的 `.sha256` 文件,可用于核对下载文件是否完整。GitHub Actions 的手动构建产物主要用于维护者测试;普通用户应使用正式 Release 资产。
## 运行采集器
Windows 用户下载后双击 `.exe`,按窗口提示操作即可。macOS 用户解压 ZIP 后双击 `start-site-adapter-collector.command`,不要打开构建目录中的 `.pkg` 文件。Linux 用户在下载目录打开终端,只需首次赋予执行权限后运行:
```bash
chmod +x moviepilot-site-collector-linux
./moviepilot-site-collector-linux
```
运行后只需输入站点首页地址,随后在弹出的临时浏览器中登录并搜索,最后回到采集器按回车。采集器会在当前目录生成 `moviepilot-site-capture-*.zip`,用户只需把这个 ZIP 附加到站点适配 Feature Request,不需要提交任何源码、Cookie 或 HTML。
## 系统安全提示
当前自动构建产物尚未接入 Windows 或 Apple 代码签名。Windows SmartScreen 或 macOS Gatekeeper 可能因此显示安全提示。仅在文件来自 MoviePilot 官方 GitHub Release,且校验摘要一致时运行;不要从聊天、网盘或第三方站点接收采集器。
如果系统阻止运行,可改用随 MoviePilot 源码提供的本地采集脚本;该方式需要 Python 3.11 及完整后端依赖,不适合作为普通用户的首选路径。
## 维护者发布流程
`.github/workflows/site-adapter-collector.yml` 支持手动触发,也会在 Release 发布后自动构建 Windows、macOS 和 Linux 单文件程序。每个平台先执行 `--help` 启动检查,再上传程序及 SHA-256 摘要为 Workflow Artifact。Release 事件会在三个平台全部成功后,把文件附加到触发本次任务的 Release;手动触发时填写已有的 `release_tag` 也会上传到该 Release,留空则只生成 3 天的测试 Artifact。
-4
View File
@@ -7,9 +7,5 @@ timeout_method = thread
# 让本仓自身的新告警更醒目。本仓代码引发的告警一律不在此忽略,应在源码/用例处修复。
filterwarnings =
ignore:datetime.datetime.utcfromtimestamp\(\) is deprecated:DeprecationWarning
ignore:websockets.legacy is deprecated:DeprecationWarning
ignore:websockets.InvalidStatusCode is deprecated:DeprecationWarning
ignore:pkg_resources is deprecated as an API:DeprecationWarning
ignore:Deprecated call to .pkg_resources.declare_namespace:DeprecationWarning
ignore:'crypt' is deprecated:DeprecationWarning
ignore:'audioop' is deprecated:DeprecationWarning
+1 -1
View File
@@ -1,4 +1,4 @@
moviepilot-rust~=0.2.1
moviepilot-rust~=0.2.3
pydantic>=2.13.4,<3.0.0
pydantic-settings>=2.14.1,<3.0.0
SQLAlchemy~=2.0.50
+99
View File
@@ -0,0 +1,99 @@
#!/usr/bin/env bash
set -euo pipefail
ORIGINAL_DIR="$PWD"
SCRIPT_DIR=""
PROJECT_ROOT=""
PYTHON_BIN=""
COLLECTOR_PATH=""
# 判断候选 Python 是否满足 3.11 最低版本。
python_version_ok() {
"$1" - <<'PY' >/dev/null 2>&1
import sys
raise SystemExit(0 if sys.version_info >= (3, 11) else 1)
PY
}
# 从当前源码目录、环境变量或脚本位置查找 MoviePilot 根目录。
find_project_root() {
local candidate=""
local source_path="${BASH_SOURCE[0]:-}"
if [[ -n "$source_path" && -f "$source_path" ]]; then
SCRIPT_DIR="$(cd "$(dirname "$source_path")" && pwd)"
fi
for candidate in "${MOVIEPILOT_ROOT:-}" "$ORIGINAL_DIR" "${SCRIPT_DIR:+$SCRIPT_DIR/..}"; do
if [[ -n "$candidate" && -f "$candidate/scripts/site_adapter_collector.py" ]]; then
PROJECT_ROOT="$(cd "$candidate" && pwd)"
return 0
fi
done
return 1
}
# 查找可用的项目虚拟环境或系统 Python。
find_python() {
local candidate=""
local resolved=""
for candidate in \
"${PROJECT_ROOT:+$PROJECT_ROOT/venv/bin/python}" \
"${PROJECT_ROOT:+$PROJECT_ROOT/.venv/bin/python}" \
"${VIRTUAL_ENV:+$VIRTUAL_ENV/bin/python}" \
python3.13 python3.12 python3.11 python3; do
[[ -n "$candidate" ]] || continue
if [[ "$candidate" == */* ]]; then
resolved="$candidate"
else
resolved="$(command -v "$candidate" 2>/dev/null || true)"
fi
if [[ -n "$resolved" && -x "$resolved" ]] && python_version_ok "$resolved"; then
PYTHON_BIN="$resolved"
return 0
fi
done
return 1
}
# 判断 Python 是否已具备独立采集器所需的最小运行依赖。
project_runtime_ready() {
PYTHONPATH="$PROJECT_ROOT${PYTHONPATH:+:$PYTHONPATH}" "$PYTHON_BIN" - <<'PY' >/dev/null 2>&1
import requests
import websocket
from bs4 import BeautifulSoup
PY
}
# 让脚本始终从终端安全读取交互输入。
restore_terminal_input() {
if [[ -r /dev/tty ]]; then
exec </dev/tty
else
echo "需要可交互终端读取站点地址并等待浏览器采集确认。" >&2
exit 1
fi
}
# 定位本地项目运行环境后启动随发行包提供的采集器。
main() {
if ! find_project_root; then
echo "未找到本地 MoviePilot 源码或安装目录,请在 MoviePilot 目录中运行此脚本。" >&2
exit 1
fi
if ! find_python; then
echo "未找到 Python 3.11 或更高版本,请先安装 Python。" >&2
exit 1
fi
if ! project_runtime_ready; then
echo "本地 Python 环境缺少采集器依赖,请优先下载官方 Release 的独立采集器。" >&2
exit 1
fi
COLLECTOR_PATH="$PROJECT_ROOT/scripts/site_adapter_collector.py"
restore_terminal_input
cd "$ORIGINAL_DIR"
PYTHONPATH="$PROJECT_ROOT${PYTHONPATH:+:$PYTHONPATH}" "$PYTHON_BIN" "$COLLECTOR_PATH" "$@"
}
main "$@"
File diff suppressed because it is too large Load Diff
+41
View File
@@ -0,0 +1,41 @@
# -*- mode: python ; coding: utf-8 -*-
""" Python """
from pathlib import Path
PROJECT_ROOT = Path(SPECPATH).resolve().parent
ENTRYPOINT = PROJECT_ROOT / "scripts" / "site_adapter_collector.py"
analysis = Analysis(
[str(ENTRYPOINT)],
pathex=[str(PROJECT_ROOT)],
binaries=[],
datas=[],
hiddenimports=[],
hookspath=[],
hooksconfig={},
runtime_hooks=[],
excludes=[],
noarchive=False,
optimize=0,
)
python_archive = PYZ(analysis.pure)
executable = EXE(
python_archive,
analysis.scripts,
analysis.binaries,
analysis.datas,
[],
name="moviepilot-site-collector",
debug=False,
bootloader_ignore_signals=False,
strip=False,
upx=False,
console=True,
disable_windowed_traceback=False,
argv_emulation=False,
target_arch=None,
codesign_identity=None,
entitlements_file=None,
)
@@ -0,0 +1,4 @@
beautifulsoup4~=4.15.0
PyInstaller>=6.14,<7.0
requests>=2.32,<3.0
websocket-client~=1.9.0
+29
View File
@@ -0,0 +1,29 @@
#!/usr/bin/env bash
set -u
# 从解压目录启动 macOS 单文件采集器,并在结束后保留终端窗口供用户查看结果。
main() {
local script_dir=""
local collector_path=""
local status=0
script_dir="$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)"
collector_path="$script_dir/moviepilot-site-collector-macos"
if [[ ! -f "$collector_path" ]]; then
echo "未找到 moviepilot-site-collector-macos,请完整解压 ZIP 后再双击启动。" >&2
status=1
else
chmod +x "$collector_path"
cd "$script_dir" || status=1
if [[ "$status" -eq 0 ]]; then
"$collector_path" || status=$?
fi
fi
echo
read -r -p "按回车关闭窗口..." _
exit "$status"
}
main "$@"
+8
View File
@@ -470,6 +470,14 @@ All endpoints are under the base URL `{MP_HOST}`. Path parameters are shown as `
| GET | `/api/v1/mcp/tools/{tool_name}` | Get tool definition |
| GET | `/api/v1/mcp/tools/{tool_name}/schema` | Get tool input schema |
### Agent MCP Client (3 endpoints)
| Method | Path | Description |
|--------|------|-------------|
| GET | `/api/v1/message/agent/mcp/servers` | List external MCP servers configured for the built-in Agent. Superuser login required |
| POST | `/api/v1/message/agent/mcp/servers` | Save external MCP servers for the built-in Agent. Body: `{"servers":[...]}` |
| POST | `/api/v1/message/agent/mcp/servers/test` | Test one external MCP server and return discovered tools. Body: `{"server":{...}}` |
### Webhook (2 endpoints)
| Method | Path | Description |
+15 -7
View File
@@ -16,9 +16,11 @@ prepare_backend()
from app.testing.network_guard import block_real_network # noqa: E402,F401
def _report_session_cleanup_error(name: str, err: Exception) -> None:
"""测试收尾清理失败只记录诊断,不覆盖原始 pytest 退出状态"""
def _report_session_cleanup_error(session, name: str, err: Exception) -> None:
"""记录收尾错误;原测试绿色时将会话标记为失败"""
sys.stderr.write(f"\npytest session cleanup failed: {name}: {err!r}\n")
if session.exitstatus == 0:
session.exitstatus = 1
def pytest_sessionfinish(session, exitstatus):
@@ -28,21 +30,27 @@ def pytest_sessionfinish(session, exitstatus):
shutdown_blocking_executors(cancel_futures=True)
except Exception as err:
_report_session_cleanup_error("agent blocking executors", err)
_report_session_cleanup_error(session, "agent blocking executors", err)
try:
from app.helper.thread import ThreadHelper
from app.utils.singleton import Singleton
helper = Singleton._instances.get((ThreadHelper, (), frozenset()))
helper = ThreadHelper.get_existing_instance()
if helper:
helper.shutdown()
except Exception as err:
_report_session_cleanup_error("thread helper", err)
_report_session_cleanup_error(session, "thread helper", err)
try:
from app.helper.message import stop_message
stop_message()
except Exception as err:
_report_session_cleanup_error(session, "message service", err)
try:
from app.log import LoggerManager
LoggerManager.shutdown()
except Exception as err:
_report_session_cleanup_error("logger manager", err)
_report_session_cleanup_error(session, "logger manager", err)
+143
View File
@@ -0,0 +1,143 @@
import sys
import textwrap
import pytest
from app.agent.mcp import AgentMcpManager, AgentMcpToolSpec
from app.agent.tools.impl.mcp import McpExternalTool
from app.schemas.agent import AgentMcpServerConfig
def _write_stdio_mcp_server(tmp_path):
"""写入一个用于测试的最小 stdio MCP 服务。"""
server_path = tmp_path / "stdio_mcp_server.py"
server_path.write_text(
textwrap.dedent(
"""
import json
import sys
TOOLS = [
{
"name": "echo",
"description": "Echo input text.",
"inputSchema": {
"type": "object",
"properties": {
"text": {"type": "string", "description": "Text to echo"}
},
"required": ["text"],
},
}
]
for line in sys.stdin:
request = json.loads(line)
request_id = request.get("id")
method = request.get("method")
if request_id is None:
continue
if method == "initialize":
result = {
"protocolVersion": request["params"]["protocolVersion"],
"capabilities": {"tools": {}},
"serverInfo": {"name": "Fake MCP", "version": "1.0.0"},
}
elif method == "tools/list":
result = {"tools": TOOLS}
elif method == "tools/call":
args = request.get("params", {}).get("arguments", {})
result = {"content": [{"type": "text", "text": args.get("text", "")}]}
else:
result = {}
print(json.dumps({"jsonrpc": "2.0", "id": request_id, "result": result}), flush=True)
"""
),
encoding="utf-8",
)
return server_path
@pytest.mark.anyio
async def test_stdio_mcp_server_lists_tools(tmp_path):
"""stdio MCP 服务器应能被初始化并读取工具列表。"""
server_path = _write_stdio_mcp_server(tmp_path)
manager = AgentMcpManager()
server = AgentMcpServerConfig(
id="fake",
name="Fake MCP",
transport="stdio",
command=sys.executable,
args=[str(server_path)],
timeout=5,
)
tools = await manager.list_server_tools(server)
assert len(tools) == 1
assert tools[0].name == "echo"
assert tools[0].agent_tool_name == "mcp_fake_mcp_echo"
assert tools[0].input_schema["properties"]["text"]["type"] == "string"
@pytest.mark.anyio
async def test_stdio_mcp_server_calls_tool(tmp_path):
"""stdio MCP 工具应能通过 tools/call 返回内容。"""
server_path = _write_stdio_mcp_server(tmp_path)
manager = AgentMcpManager()
server = AgentMcpServerConfig(
id="fake",
name="Fake MCP",
transport="stdio",
command=sys.executable,
args=[str(server_path)],
timeout=5,
)
result = await manager.call_server_tool(server, "echo", {"text": "hello"})
assert result == {"content": [{"type": "text", "text": "hello"}]}
def test_normalize_server_generates_runtime_defaults():
"""MCP 配置规范化应补齐默认值并清理空字段。"""
manager = AgentMcpManager()
server = manager.normalize_server(
{
"id": "demo",
"name": "Demo",
"transport": "http",
"url": " https://example.com/mcp ",
"headers": {" Authorization ": "Bearer token", "": "ignored"},
"timeout": "bad",
}
)
assert server.id == "demo"
assert server.url == "https://example.com/mcp"
assert server.headers == {"Authorization": "Bearer token"}
assert server.timeout == 30
assert server.require_admin is True
def test_mcp_external_tool_uses_discovered_schema():
"""外部 MCP 工具应保留发现到的 JSON Schema。"""
server = AgentMcpServerConfig(id="fake", name="Fake MCP")
spec = AgentMcpToolSpec(
server=server,
name="echo",
agent_tool_name="mcp_fake_echo",
description="Echo input text.",
input_schema={
"type": "object",
"properties": {"text": {"type": "string"}},
"required": ["text"],
},
)
tool = McpExternalTool(spec=spec, session_id="session-1", user_id="10001")
assert tool.name == "mcp_fake_echo"
assert tool.args_schema["properties"]["text"]["type"] == "string"
assert tool.require_admin is True
+106
View File
@@ -0,0 +1,106 @@
import asyncio
import pytest
from sqlalchemy import create_engine, text
from sqlalchemy.exc import OperationalError
import app.db as db_module
class _SqliteError(Exception):
"""模拟 sqlite3 异常暴露的扩展错误字段。"""
sqlite_errorcode = 266
sqlite_errorname = "SQLITE_IOERR_READ"
class _PsycopgError(Exception):
"""模拟 psycopg2 异常暴露的 SQLSTATE 字段。"""
pgcode = "40001"
class _AsyncpgError(Exception):
"""模拟 asyncpg 适配异常暴露的 SQLSTATE 字段。"""
sqlstate = "23505"
@pytest.mark.parametrize(
("error", "expected"),
[
(
_SqliteError("disk I/O error"),
{
"error_type": "_SqliteError",
"error_code": 266,
"error_name": "SQLITE_IOERR_READ",
},
),
(
_PsycopgError("serialization failure"),
{
"error_type": "_PsycopgError",
"sqlstate": "40001",
},
),
(
_AsyncpgError("duplicate key"),
{
"error_type": "_AsyncpgError",
"sqlstate": "23505",
},
),
],
)
def test_database_error_metadata_extracts_driver_codes(error, expected) -> None:
"""诊断元数据应兼容 SQLite、psycopg2 与 asyncpg 的稳定错误字段。"""
assert db_module._database_error_metadata(error) == expected
def test_database_error_listener_omits_statement_and_parameters(monkeypatch) -> None:
"""数据库错误日志不得包含 SQL、参数或驱动返回的原始消息。"""
messages = []
engine = create_engine("sqlite:///:memory:")
monkeypatch.setattr("app.db.logger.error", messages.append)
db_module._register_database_error_logging(engine)
with pytest.raises(OperationalError):
with engine.connect() as connection:
connection.execute(
text("SELECT * FROM missing_table WHERE token = :token"),
{"token": "private-token"},
)
assert len(messages) == 1
assert "database=sqlite" in messages[0]
assert "driver=pysqlite" in messages[0]
assert "error_code=1" in messages[0]
assert "error_name=SQLITE_ERROR" in messages[0]
assert "missing_table" not in messages[0]
assert "private-token" not in messages[0]
def test_async_database_engine_logs_driver_error_metadata(monkeypatch) -> None:
"""异步 Engine 应通过底层 sync engine 记录驱动错误码。"""
messages = []
monkeypatch.setattr("app.db.logger.error", messages.append)
async def query_missing_table() -> None:
async with db_module.AsyncEngine.connect() as connection:
await connection.execute(text("SELECT * FROM async_missing_table"))
with pytest.raises(OperationalError):
asyncio.run(query_missing_table())
assert len(messages) == 1
assert "database=sqlite" in messages[0]
assert "driver=aiosqlite" in messages[0]
assert "error_code=1" in messages[0]
assert "error_name=SQLITE_ERROR" in messages[0]
assert "async_missing_table" not in messages[0]
def test_database_error_metadata_ignores_unclassified_errors() -> None:
"""没有驱动错误码时不应制造无效诊断日志。"""
assert db_module._database_error_metadata(RuntimeError("plain failure")) is None
+245
View File
@@ -0,0 +1,245 @@
import os
import subprocess
import textwrap
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
def _write_entrypoint_functions(tmp_path: Path) -> Path:
content = (ROOT / "docker" / "entrypoint.sh").read_text(encoding="utf-8")
marker = "# 使用env配置"
assert marker in content
functions = tmp_path / "entrypoint-functions.sh"
functions.write_text(content.split(marker, 1)[0], encoding="utf-8")
return functions
def _write_fake_chown(tmp_path: Path) -> Path:
fake_bin = tmp_path / "bin"
fake_bin.mkdir()
chown = fake_bin / "chown"
chown.write_text(
textwrap.dedent(
"""\
#!/usr/bin/env bash
printf '%s\\n' "$*" >> "${MP_CHOWN_LOG}"
"""
),
encoding="utf-8",
)
chown.chmod(0o755)
return fake_bin
def _run_permission_case(tmp_path: Path, body: str, env: dict[str, str] | None = None) -> str:
functions = _write_entrypoint_functions(tmp_path)
fake_bin = _write_fake_chown(tmp_path)
chown_log = tmp_path / "chown.log"
app_dir = tmp_path / "app"
public_dir = tmp_path / "public"
home_dir = tmp_path / "home"
(app_dir / "app" / "plugins").mkdir(parents=True)
public_dir.mkdir()
(home_dir / ".cloakbrowser").mkdir(parents=True)
(home_dir / "runtime").mkdir()
(app_dir / "app" / "plugins" / "plugin.py").write_text("# plugin\n", encoding="utf-8")
(public_dir / "index.html").write_text("<!doctype html>\n", encoding="utf-8")
(home_dir / ".cloakbrowser" / "chrome").write_text("browser cache\n", encoding="utf-8")
(home_dir / "runtime" / "state").write_text("state\n", encoding="utf-8")
external_target = tmp_path / "external-target"
external_target.write_text("external\n", encoding="utf-8")
(app_dir / "external-link").symlink_to(external_target)
case_env = {
**os.environ,
"PATH": f"{fake_bin}:{os.environ['PATH']}",
"MP_CHOWN_LOG": str(chown_log),
"ENTRYPOINT_FUNCTIONS": str(functions),
"APP_DIR": str(app_dir),
"PUBLIC_DIR": str(public_dir),
"HOME_DIR": str(home_dir),
"CONFIG_DIR": str(tmp_path / "config"),
"PUID": str(os.getuid()),
"PGID": str(os.getgid()),
}
if env:
case_env.update(env)
script = textwrap.dedent(
f"""\
set -euo pipefail
source "${{ENTRYPOINT_FUNCTIONS}}"
{body}
"""
)
subprocess.run(["bash", "-c", script], check=True, env=case_env)
return chown_log.read_text(encoding="utf-8") if chown_log.exists() else ""
def _run_entrypoint_case(tmp_path: Path, body: str, env: dict[str, str] | None = None) -> str:
functions = _write_entrypoint_functions(tmp_path)
case_env = {
**os.environ,
"ENTRYPOINT_FUNCTIONS": str(functions),
}
if env:
case_env.update(env)
script = textwrap.dedent(
f"""\
set -euo pipefail
source "${{ENTRYPOINT_FUNCTIONS}}"
{body}
"""
)
result = subprocess.run(["bash", "-c", script], check=True, env=case_env, text=True, capture_output=True)
return result.stdout
def test_image_paths_are_not_chowned_by_default_regardless_of_owner(tmp_path: Path) -> None:
log = _run_permission_case(
tmp_path,
'force_chown_image_paths_if_requested "${APP_DIR}" "${PUBLIC_DIR}"',
env={"PUID": "999999", "PGID": "999999"},
)
assert log == ""
def test_image_paths_force_chown_uses_recursive_repair(tmp_path: Path) -> None:
log = _run_permission_case(
tmp_path,
'MOVIEPILOT_FORCE_CHOWN=true force_chown_image_paths_if_requested "${APP_DIR}" "${PUBLIC_DIR}"',
)
assert log.startswith("-R moviepilot:moviepilot ")
assert "/app" in log
assert "/public" in log
def test_image_paths_force_chown_accepts_numeric_and_yes_values(tmp_path: Path) -> None:
for force_value in ("1", "YES"):
case_path = tmp_path / force_value.lower()
case_path.mkdir()
log = _run_permission_case(
case_path,
f'MOVIEPILOT_FORCE_CHOWN={force_value} force_chown_image_paths_if_requested "${{APP_DIR}}" "${{PUBLIC_DIR}}"',
)
assert log.startswith("-R moviepilot:moviepilot ")
assert "/app" in log
assert "/public" in log
def test_plugin_directory_skips_chown_when_owner_matches(tmp_path: Path) -> None:
log = _run_permission_case(
tmp_path,
'chown_plugin_runtime_path "${APP_DIR}/app/plugins"',
)
assert log == ""
def test_plugin_directory_chowns_only_root_directory_when_owner_mismatches(tmp_path: Path) -> None:
log = _run_permission_case(
tmp_path,
'chown_plugin_runtime_path "${APP_DIR}/app/plugins"',
env={"PUID": "999999", "PGID": "999999"},
)
assert log == f"-h moviepilot:moviepilot {tmp_path}/app/app/plugins\n"
def test_home_permissions_skip_cloakbrowser_cache_by_default(tmp_path: Path) -> None:
log = _run_permission_case(
tmp_path,
'HOME="${HOME_DIR}" correct_home_permissions',
)
lines = log.splitlines()
assert f"moviepilot:moviepilot {tmp_path}/home" in lines
assert f"-h moviepilot:moviepilot {tmp_path}/home/.cloakbrowser" in lines
assert f"-R moviepilot:moviepilot {tmp_path}/home/runtime" in lines
assert not any(line.startswith("-R ") and ".cloakbrowser" in line for line in lines)
def test_home_permissions_force_chown_repairs_cloakbrowser_cache(tmp_path: Path) -> None:
log = _run_permission_case(
tmp_path,
'MOVIEPILOT_FORCE_CHOWN=yes HOME="${HOME_DIR}" correct_home_permissions',
)
assert f"-R moviepilot:moviepilot {tmp_path}/home/.cloakbrowser" in log
def test_runtime_writable_paths_are_still_corrected(tmp_path: Path) -> None:
log = _run_permission_case(
tmp_path,
'HOME="${HOME_DIR}" correct_file_permissions',
env={"PUID": "999999", "PGID": "999999"},
)
lines = log.splitlines()
assert f"moviepilot:moviepilot {tmp_path}/home" in lines
assert f"-h moviepilot:moviepilot {tmp_path}/home/.cloakbrowser" in lines
assert f"-R moviepilot:moviepilot {tmp_path}/home/runtime" in lines
assert f"-R moviepilot:moviepilot {tmp_path}/config /var/lib/nginx /var/log/nginx" in lines
assert "moviepilot:moviepilot /etc/hosts /tmp" in lines
assert not any(line.startswith("-R ") and ".cloakbrowser" in line for line in lines)
assert not any(f"{tmp_path}/app " in line for line in lines)
assert not any(f"{tmp_path}/public" in line for line in lines)
def test_backend_ready_log_uses_configured_ports(tmp_path: Path) -> None:
curl_log = tmp_path / "curl.log"
output = _run_entrypoint_case(
tmp_path,
"""
INFO() { printf '[INFO] %s\\n' "$1"; }
curl() {
printf '%s\\n' "$*" > "${CURL_LOG}"
return 0
}
PORT=4321 NGINX_PORT=8765 wait_backend_ready 1 2 "$$"
""",
env={"CURL_LOG": str(curl_log)},
)
assert curl_log.read_text(encoding="utf-8") == (
"-fsS --max-time 2 http://127.0.0.1:4321/api/v1/system/global?token=moviepilot\n"
)
assert "MoviePilot Web 已可访问" in output
assert "后端就绪耗时" in output
assert "后端端口 4321" in output
assert "前端端口 8765" in output
def test_backend_ready_timeout_falls_back_to_default_for_invalid_value(tmp_path: Path) -> None:
output = _run_entrypoint_case(
tmp_path,
"""
WARN() { printf '[WARN] %s\\n' "$1"; }
curl() { return 1; }
MOVIEPILOT_BACKEND_READY_TIMEOUT=invalid wait_backend_ready 1 2 999999 || true
""",
)
assert "MOVIEPILOT_BACKEND_READY_TIMEOUT=invalid 无效,使用默认 300 秒" in output
assert "后端服务启动完成探测已停止:后端进程已退出" in output
def test_backend_ready_timeout_accepts_leading_zero_decimal(tmp_path: Path) -> None:
output = _run_entrypoint_case(
tmp_path,
"""
INFO() { printf '[INFO] %s\\n' "$1"; }
WARN() { printf '[WARN] %s\\n' "$1"; }
curl() { return 0; }
MOVIEPILOT_BACKEND_READY_TIMEOUT=08 wait_backend_ready 1 2 "$$"
""",
)
assert "MOVIEPILOT_BACKEND_READY_TIMEOUT=08 无效" not in output
assert "MoviePilot Web 已可访问" in output
+56
View File
@@ -3,6 +3,61 @@ import socket
from app.helper import doh
def test_doh_executor_is_lazy_and_shutdown_restores_socket(monkeypatch):
"""DoH 线程池按需创建,并在模块关闭时恢复系统 DNS"""
original_getaddrinfo = socket.getaddrinfo
helper = object.__new__(doh.DohHelper)
monkeypatch.setattr(doh.settings, "DOH_DOMAINS", "example.com")
monkeypatch.setattr(doh.settings, "DOH_RESOLVERS", "resolver.test")
monkeypatch.setattr(doh, "_doh_query", lambda resolver, host: "203.0.113.7")
monkeypatch.setattr(doh, "_orig_getaddrinfo", lambda host, *args, **kwargs: [])
try:
helper.shutdown()
assert doh._executor is None
doh.enable_doh(True)
socket.getaddrinfo("example.com", None)
executor = doh._executor
assert executor is not None
helper.shutdown()
assert doh._executor is None
assert socket.getaddrinfo is doh._orig_getaddrinfo
assert getattr(executor, "_shutdown", False)
finally:
helper.shutdown()
socket.getaddrinfo = original_getaddrinfo
def test_doh_config_reload_disables_and_closes_executor(monkeypatch):
"""热更新关闭 DoH 时恢复系统 DNS 并释放已创建的线程池"""
original_getaddrinfo = socket.getaddrinfo
helper = object.__new__(doh.DohHelper)
monkeypatch.setattr(doh.settings, "DOH_DOMAINS", "example.com")
monkeypatch.setattr(doh.settings, "DOH_RESOLVERS", "resolver.test")
monkeypatch.setattr(doh, "_doh_query", lambda resolver, host: "203.0.113.7")
monkeypatch.setattr(doh, "_orig_getaddrinfo", lambda host, *args, **kwargs: [])
try:
helper.shutdown()
doh.enable_doh(True)
socket.getaddrinfo("example.com", None)
executor = doh._executor
assert executor is not None
monkeypatch.setattr(doh.settings, "DOH_ENABLE", False)
helper.on_config_changed()
assert doh._executor is None
assert getattr(executor, "_shutdown", False)
assert socket.getaddrinfo is doh._orig_getaddrinfo
finally:
helper.shutdown()
socket.getaddrinfo = original_getaddrinfo
def test_enable_doh_reuses_cached_host_resolution(monkeypatch):
"""
同一 DoH 域名第二次解析应命中缓存避免重复请求远端解析器
@@ -33,6 +88,7 @@ def test_enable_doh_reuses_cached_host_resolution(monkeypatch):
socket.getaddrinfo("example.com", None)
socket.getaddrinfo("example.com", None)
finally:
object.__new__(doh.DohHelper).shutdown()
socket.getaddrinfo = original_getaddrinfo
with doh._doh_lock:
doh._doh_cache.clear()
+136
View File
@@ -0,0 +1,136 @@
import asyncio
import inspect
from app.api.endpoints import douban as douban_endpoint
from app.db.user_oper import get_current_active_superuser_async
from app.modules.douban.douban_cache import DoubanCache
from app.schemas.types import MediaType
class _MemoryCacheStub:
"""提供豆瓣缓存管理测试所需的最小内存后端。"""
def __init__(self, data: dict):
"""使用给定字典初始化测试缓存。"""
self.data = data
def items(self):
"""返回全部缓存条目。"""
return self.data.items()
def get(self, key: str):
"""读取指定缓存条目。"""
return self.data.get(key)
def delete(self, key: str):
"""删除指定缓存条目。"""
self.data.pop(key, None)
def clear(self):
"""清空全部缓存条目。"""
self.data.clear()
def _build_douban_cache(data: dict) -> DoubanCache:
"""构造绕过单例初始化的豆瓣缓存测试实例。"""
cache = object.__new__(DoubanCache)
cache._cache = _MemoryCacheStub(data)
cache.save = lambda force=False: None
return cache
def test_douban_cache_management_endpoints_require_superuser():
"""豆瓣识别缓存管理接口必须仅允许超级管理员访问。"""
endpoints = [
douban_endpoint.douban_recognition_cache,
douban_endpoint.delete_douban_recognition_cache,
douban_endpoint.clear_douban_recognition_cache,
]
for endpoint in endpoints:
dependency = inspect.signature(endpoint).parameters["_"].default.dependency
assert dependency is get_current_active_superuser_async
def test_douban_cache_list_items_normalizes_media_type_and_sorting():
"""豆瓣管理列表应输出稳定顺序和前端可识别的媒体类型。"""
cache = _build_douban_cache({
"[电视剧]Zulu-2024-1": {
"id": "2",
"title": "Zulu",
"type": MediaType.TV,
"year": "2024",
},
"[电影]Alpha-2023-None": {
"id": "1",
"title": "Alpha",
"type": "电影",
"year": "2023",
"poster_path": "https://example.com/alpha.jpg",
},
"[电影]Missing-2022-None": {"id": 0},
})
items = cache.list_items()
assert [item["title"] for item in items] == ["Alpha", "", "Zulu"]
assert [item["media_type"] for item in items] == ["movie", "unknown", "tv"]
assert items[0]["poster_path"] == "https://example.com/alpha.jpg"
assert items[1]["douban_id"] == 0
def test_douban_cache_delete_and_clear_persist_immediately(monkeypatch):
"""豆瓣管理操作应修改运行时缓存并立即触发本地持久化。"""
cache = _build_douban_cache({"first": {"id": "1"}, "second": {"id": "2"}})
saved_forces = []
monkeypatch.setattr(cache, "save", lambda force=False: saved_forces.append(force))
assert cache.delete("first") == {"id": "1"}
assert cache.delete("missing") == {}
cache.clear()
assert cache.list_items() == []
assert saved_forces == [True, True]
def test_douban_cache_endpoint_returns_management_statistics(monkeypatch):
"""豆瓣查询接口应返回识别成功和失败条目的统计。"""
cache = _build_douban_cache({
"recognized": {"id": "1", "title": "Alpha", "type": MediaType.MOVIE},
"unrecognized": {"id": 0},
})
monkeypatch.setattr(douban_endpoint, "DoubanCache", lambda: cache)
response = asyncio.run(douban_endpoint.douban_recognition_cache(None))
assert response.success is True
assert response.data["count"] == 2
assert response.data["recognized"] == 1
assert response.data["unrecognized"] == 1
def test_douban_cache_delete_endpoint_reports_missing_item(monkeypatch):
"""豆瓣删除接口应区分成功删除与缓存不存在。"""
cache = _build_douban_cache({"existing": {"id": "1"}})
monkeypatch.setattr(douban_endpoint, "DoubanCache", lambda: cache)
deleted_response = asyncio.run(
douban_endpoint.delete_douban_recognition_cache("existing", None)
)
missing_response = asyncio.run(
douban_endpoint.delete_douban_recognition_cache("missing", None)
)
assert deleted_response.success is True
assert missing_response.success is False
def test_douban_cache_clear_endpoint_removes_all_items(monkeypatch):
"""豆瓣清空接口应删除全部识别缓存。"""
cache = _build_douban_cache({"existing": {"id": "1"}})
monkeypatch.setattr(douban_endpoint, "DoubanCache", lambda: cache)
response = asyncio.run(douban_endpoint.clear_douban_recognition_cache(None))
assert response.success is True
assert cache.list_items() == []
+141
View File
@@ -568,6 +568,147 @@ def test_batch_download_threads_custom_words_to_download_single(monkeypatch):
assert chain.download_single.call_args.kwargs["custom_words"] == custom_words
def test_download_single_records_failure_cooldown_when_downloader_rejects(monkeypatch):
"""
下载器拒绝种子且没有返回 hash 应记录资源级失败冷却
"""
captured = {}
class _CapturingDownloadFailureOper:
"""
捕获下载失败冷却记录避免测试写入数据库
"""
def record_failure(self, **kwargs: object) -> SimpleNamespace:
"""
保存写入字段供断言使用
"""
captured.update(kwargs)
return SimpleNamespace(id=1)
monkeypatch.setattr(
"app.helper.directory.DirectoryHelper.get_download_dirs",
lambda _self: _download_dirs(),
)
monkeypatch.setattr(download_module, "TorrentHelper", _FakeTorrentHelper)
monkeypatch.setattr(download_module, "DownloadFailureOper", _CapturingDownloadFailureOper)
monkeypatch.setattr(download_module.eventmanager, "send_event", lambda *args, **kwargs: None)
chain = DownloadChain.__new__(DownloadChain)
error_msg = "添加种子任务失败:无法读取种子文件"
chain.download = MagicMock(return_value=("qb", None, "Original", error_msg))
chain.post_message = MagicMock()
context = Context(
meta_info=MetaInfo("Demo Movie 2026"),
media_info=MediaInfo(
type=MediaType.MOVIE,
title="Demo Movie",
year="2026",
tmdb_id=1,
genre_ids=[18],
),
torrent_info=TorrentInfo(
site=12,
site_name="AGSVPT",
title="Demo Movie 2026 1080p",
enclosure="https://example.com/download.php?id=484660",
size=1024,
),
)
download_id, returned_error = chain.download_single(
context=context,
torrent_content=b"torrent-content",
save_path="/downloads",
source="Subscribe|{}",
return_detail=True,
)
assert download_id is None
assert returned_error == error_msg
assert captured["fingerprint"] == DownloadChain._build_download_failure_fingerprint(context)
assert captured["torrent_id"] == "example.com:id=484660"
assert captured["site"] == 12
assert captured["error_message"] == error_msg
assert captured["next_retry_at"] > captured["now_time"]
def test_batch_download_skips_failed_subscription_resource_and_tries_next(monkeypatch):
"""
订阅自动下载应跳过冷却中的失败资源但继续尝试同媒体的后续候选
"""
_FakeBatchTorrentHelper.episodes = []
monkeypatch.setattr(download_module, "TorrentHelper", _FakeBatchTorrentHelper)
monkeypatch.setattr(download_module.eventmanager, "send_event", lambda *args, **kwargs: None)
first_context = SimpleNamespace(
media_info=SimpleNamespace(
type=MediaType.MOVIE,
title="Demo Movie",
year="2026",
title_year="Demo Movie (2026)",
tmdb_id=1,
douban_id=None,
),
meta_info=SimpleNamespace(season=None, episode=None, episode_list=[], season_episode=""),
torrent_info=SimpleNamespace(
site=12,
site_name="AGSVPT",
title="Demo Movie Bad",
torrent_id="484660",
size=1024,
),
)
second_context = SimpleNamespace(
media_info=SimpleNamespace(
type=MediaType.MOVIE,
title="Demo Movie",
year="2026",
title_year="Demo Movie (2026)",
tmdb_id=1,
douban_id=None,
),
meta_info=SimpleNamespace(season=None, episode=None, episode_list=[], season_episode=""),
torrent_info=SimpleNamespace(
site=13,
site_name="OtherSite",
title="Demo Movie Good",
torrent_id="999999",
size=2048,
),
)
failed_fingerprint = DownloadChain._build_download_failure_fingerprint(first_context)
class _ActiveDownloadFailureOper:
"""
返回第一个候选的活跃失败冷却记录
"""
def get_active_by_fingerprints(self, fingerprints: list[str], now_time: str) -> dict:
"""
模拟数据库批量查询活跃失败记录
"""
assert now_time
assert failed_fingerprint in fingerprints
return {failed_fingerprint: SimpleNamespace(fingerprint=failed_fingerprint)}
monkeypatch.setattr(download_module, "DownloadFailureOper", _ActiveDownloadFailureOper)
chain = DownloadChain.__new__(DownloadChain)
chain.download_single = MagicMock(return_value="hash")
downloads, lefts = chain.batch_download(
contexts=[first_context, second_context],
source="Subscribe|{}",
)
assert downloads == [second_context]
assert lefts is None
chain.download_single.assert_called_once()
assert chain.download_single.call_args.args[0] is second_context
def test_batch_download_accepts_complete_coverage_when_files_cover_target_range(monkeypatch):
"""
自定义起始集场景按目标范围覆盖判断100-143 可满足 start=100total=143
+1
View File
@@ -102,6 +102,7 @@ class EmbyDashboardLinksTest(unittest.TestCase):
with (
patch.object(client, "_Emby__get_emby_librarys") as librarys,
patch.object(client, "_Emby__get_local_image_by_id") as image_by_id,
patch.object(client, "get_items_count", return_value=0),
):
librarys.return_value = [
{
+126
View File
@@ -0,0 +1,126 @@
# -*- coding: utf-8 -*-
from unittest.mock import patch
from app.modules.jellyfin.jellyfin import Jellyfin
class _FakeResponse:
"""模拟媒体服务器 HTTP 响应。"""
def __init__(self, payload: dict, status_code: int = 200):
"""保存响应数据和状态码。"""
self._payload = payload
self.status_code = status_code
def json(self) -> dict:
"""返回模拟的 JSON 数据。"""
return self._payload
def _make_client(user: str = "user-id") -> Jellyfin:
"""构造跳过初始化的 Jellyfin 客户端。"""
client = Jellyfin.__new__(Jellyfin)
client._host = "http://media.local/"
client._apikey = "token"
client.user = user
client._sync_libraries = []
return client
def _routed_get_res(views: dict, counts: dict, global_counts: dict = None):
"""按 URL 分发的 get_res 模拟:视图列表、单库统计与全局统计。"""
def _get_res(url, params=None, **_kwargs):
if url.endswith("/Views"):
return _FakeResponse(views)
if url.endswith("/Items") and params:
key = (params.get("ParentId"), params.get("IncludeItemTypes"))
return _FakeResponse({"TotalRecordCount": counts.get(key, 0)})
if url.endswith("Items/Counts"):
return _FakeResponse(global_counts or {})
raise AssertionError(f"意外的请求地址:{url}")
return _get_res
def test_medias_count_deduplicates_multi_folder_library():
"""多文件夹电影库应按用户视图统计,避免版本重复计数(#5915)。"""
client = _make_client()
views = {"Items": [{"Id": "lib-movie", "CollectionType": "movies"}]}
# 用户级查询会折叠同一影片的多个版本,返回 67 而非数据库原始行数 201
counts = {("lib-movie", "Movie"): 67}
with patch("app.modules.jellyfin.jellyfin.RequestUtils") as request_utils_cls:
request_utils_cls.return_value.get_res.side_effect = _routed_get_res(
views, counts, global_counts={"MovieCount": 201}
)
stat = client.get_medias_count()
assert stat.movie_count == 67
assert stat.tv_count == 0
assert stat.episode_count == 0
def test_medias_count_buckets_by_collection_type():
"""电影与剧集库应按视图类型分桶累计,未知类型库不参与统计。"""
client = _make_client()
views = {
"Items": [
{"Id": "lib-movie", "CollectionType": "movies"},
{"Id": "lib-tv", "CollectionType": "tvshows"},
{"Id": "lib-music", "CollectionType": "music"},
{"Id": None, "CollectionType": "movies"},
]
}
counts = {
("lib-movie", "Movie"): 12,
("lib-tv", "Series"): 3,
("lib-tv", "Episode"): 45,
}
with patch("app.modules.jellyfin.jellyfin.RequestUtils") as request_utils_cls:
request_utils_cls.return_value.get_res.side_effect = _routed_get_res(
views, counts
)
stat = client.get_medias_count()
assert stat.movie_count == 12
assert stat.tv_count == 3
assert stat.episode_count == 45
def test_medias_count_falls_back_without_user():
"""无可用用户时应回退到全局 Items/Counts 统计。"""
client = _make_client(user=None)
with patch("app.modules.jellyfin.jellyfin.RequestUtils") as request_utils_cls:
request_utils_cls.return_value.get_res.return_value = _FakeResponse(
{"MovieCount": 5, "SeriesCount": 2, "EpisodeCount": 30}
)
stat = client.get_medias_count()
assert stat.movie_count == 5
assert stat.tv_count == 2
assert stat.episode_count == 30
args = request_utils_cls.return_value.get_res.call_args.args
assert args[0] == "http://media.local/Items/Counts"
def test_medias_count_falls_back_when_views_unavailable():
"""媒体库视图查询失败时应回退到全局 Items/Counts 统计。"""
client = _make_client()
def _get_res(url, params=None, **_kwargs):
if url.endswith("/Views"):
return None
if url.endswith("Items/Counts"):
return _FakeResponse({"MovieCount": 7, "SeriesCount": 1, "EpisodeCount": 9})
raise AssertionError(f"意外的请求地址:{url}")
with patch("app.modules.jellyfin.jellyfin.RequestUtils") as request_utils_cls:
request_utils_cls.return_value.get_res.side_effect = _get_res
stat = client.get_medias_count()
assert stat.movie_count == 7
assert stat.tv_count == 1
assert stat.episode_count == 9
+405
View File
@@ -0,0 +1,405 @@
import asyncio
import signal
import threading
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import FastAPI
from app.startup import lifecycle, modules_initializer
from app.utils import http as http_utils
def _assert_completed_once(mock: MagicMock) -> None:
if isinstance(mock, AsyncMock):
mock.assert_awaited_once_with()
else:
mock.assert_called_once_with()
def _patch_lifespan(monkeypatch, *, failing_step: str | None = None) -> dict:
"""隔离 lifespan 的外部依赖,并按名称注入一个关闭失败"""
monkeypatch.setattr(lifecycle.settings, "MOVIEPILOT_SAFE_MODE", False)
monkeypatch.setattr(lifecycle.global_vars, "set_loop", MagicMock())
monkeypatch.setattr(lifecycle.global_vars, "stop_system", MagicMock())
for name in (
"init_routers",
"init_modules",
"init_plugins",
"init_scheduler",
"init_monitor",
"init_command",
"init_workflow",
):
monkeypatch.setattr(lifecycle, name, MagicMock())
system_chain = MagicMock()
monkeypatch.setattr(lifecycle, "SystemChain", MagicMock(return_value=system_chain))
monkeypatch.setattr(lifecycle, "init_extra", AsyncMock())
shutdown_steps = {
"backup_plugins": system_chain.backup_plugins,
"stop_workflow": MagicMock(),
"stop_command": MagicMock(),
"stop_monitor": MagicMock(),
"stop_scheduler": MagicMock(),
"stop_plugins": MagicMock(),
"stop_modules": AsyncMock(),
"close_http": AsyncMock(),
}
for name in (
"stop_workflow",
"stop_command",
"stop_monitor",
"stop_scheduler",
"stop_plugins",
):
monkeypatch.setattr(lifecycle, name, shutdown_steps[name])
monkeypatch.setattr(lifecycle, "stop_modules", shutdown_steps["stop_modules"])
monkeypatch.setattr(
lifecycle,
"aclose_shared_async_transports",
shutdown_steps["close_http"],
)
if failing_step:
shutdown_steps[failing_step].side_effect = RuntimeError(
f"{failing_step} failed"
)
logger_shutdown = MagicMock()
monkeypatch.setattr(lifecycle.LoggerManager, "shutdown", logger_shutdown)
shutdown_steps["logger"] = logger_shutdown
return shutdown_steps
@pytest.mark.parametrize(
"failing_step",
[
"backup_plugins",
"stop_workflow",
"stop_command",
"stop_monitor",
"stop_scheduler",
"stop_plugins",
"stop_modules",
"close_http",
],
)
def test_lifespan_continues_after_each_shutdown_owner_failure(
monkeypatch,
failing_step,
):
"""任一关闭阶段失败都不能跳过后续资源所有者"""
shutdown_steps = _patch_lifespan(monkeypatch, failing_step=failing_step)
async def run_lifespan():
async with lifecycle.lifespan(FastAPI()):
pass
asyncio.run(run_lifespan())
lifecycle.global_vars.stop_system.assert_called_once_with()
for step in shutdown_steps.values():
_assert_completed_once(step)
def test_uvicorn_signal_publishes_stop_before_server_exit(monkeypatch):
"""Uvicorn 接管系统信号时必须先发布协作停止标志"""
from app import main
calls = []
monkeypatch.setattr(main.global_vars, "stop_system", lambda: calls.append("stop"))
monkeypatch.setattr(
main.uvicorn.Server,
"handle_exit",
lambda _self, _sig, _frame: calls.append("uvicorn"),
)
server = object.__new__(main.MoviePilotServer)
server.handle_exit(signal.SIGTERM, None)
assert calls == ["stop", "uvicorn"]
def test_application_preserves_stop_requested_before_startup(monkeypatch):
"""启动流程不能清除初始化前已经发布的退出请求"""
from app import main
stop_event = threading.Event()
stop_event.set()
monkeypatch.setattr(main.global_vars, "STOP_EVENT", stop_event)
calls = []
monkeypatch.setattr(
main.signal,
"signal",
lambda *_args: calls.append("signal"),
)
monkeypatch.setattr(main, "start_tray", lambda: calls.append("tray"))
monkeypatch.setattr(main, "init_db", lambda: calls.append("init_db"))
monkeypatch.setattr(main, "update_db", lambda: calls.append("update_db"))
monkeypatch.setattr(main.Server, "run", lambda: calls.append("server"))
main.run_application()
assert stop_event.is_set()
assert calls == [
"signal",
"signal",
"tray",
"init_db",
"update_db",
"server",
]
def test_uvicorn_preserves_stop_requested_before_serve(monkeypatch):
"""Uvicorn 启动不能清除数据库初始化阶段已经发布的停止请求"""
from app import main
stop_event = threading.Event()
monkeypatch.setattr(main.global_vars, "STOP_EVENT", stop_event)
main.global_vars.stop_system()
async def serve(_self, sockets=None):
assert main.global_vars.is_system_stopped
monkeypatch.setattr(main.uvicorn.Server, "serve", serve)
server = object.__new__(main.MoviePilotServer)
asyncio.run(server.serve())
@pytest.mark.parametrize("endpoint_name", ["restart_system", "upgrade_system"])
@pytest.mark.parametrize(
"initially_stopped",
[False, True],
ids=["running", "stopping"],
)
def test_restart_endpoint_failure_preserves_stop_state(
monkeypatch,
endpoint_name,
initially_stopped,
):
"""重启或升级失败不能发布或撤销停止请求"""
from app.api.endpoints import system
stop_event = threading.Event()
if initially_stopped:
stop_event.set()
monkeypatch.setattr(system.global_vars, "STOP_EVENT", stop_event)
monkeypatch.setattr(system.SystemHelper, "can_restart", MagicMock(return_value=True))
monkeypatch.setattr(
system.SystemHelper,
"restart" if endpoint_name == "restart_system" else "upgrade",
MagicMock(return_value=(False, "restart failed")),
)
if endpoint_name == "restart_system":
response = system.restart_system(None)
else:
response = system.upgrade_system(None, None)
assert not response.success
assert stop_event.is_set() is initially_stopped
def test_command_restart_failure_does_not_publish_stop_request(monkeypatch):
"""命令重启失败时进程仍在运行,不能提前发布停止请求"""
from app.chain.system import SystemChain
from app.core.config import global_vars
stop_event = threading.Event()
monkeypatch.setattr(global_vars, "STOP_EVENT", stop_event)
monkeypatch.setattr(SystemChain, "backup_plugins", MagicMock())
restart = MagicMock(return_value=(False, "restart failed"))
monkeypatch.setattr("app.chain.system.SystemHelper.restart", restart)
chain = object.__new__(SystemChain)
chain.restart(channel=None, userid=None)
restart.assert_called_once_with()
assert not stop_event.is_set()
def test_stop_modules_continues_after_internal_owner_failures(monkeypatch):
"""模块关闭编排中的多个失败不能阻断其余清理"""
stop_agent = AsyncMock(side_effect=RuntimeError("agent failed"))
monkeypatch.setattr(modules_initializer, "stop_agent", stop_agent)
dependencies = _patch_module_shutdown_dependencies(monkeypatch)
dependencies["module"].side_effect = RuntimeError("module failed")
asyncio.run(modules_initializer.stop_modules())
stop_agent.assert_awaited_once_with()
for dependency in dependencies.values():
_assert_completed_once(dependency)
def _patch_module_shutdown_dependencies(monkeypatch) -> dict:
"""替换 stop_modules 的资源所有者,避免测试启动真实后台服务"""
dependencies = {}
for name, method_name in (
("ModuleManager", "stop"),
("EventManager", "stop"),
("DisplayHelper", "stop"),
("DohHelper", "shutdown"),
("ThreadHelper", "shutdown"),
("RedisHelper", "close"),
):
instance = MagicMock()
setattr(instance, method_name, MagicMock())
monkeypatch.setattr(
modules_initializer,
name,
MagicMock(return_value=instance),
)
key = name.removesuffix("Helper").removesuffix("Manager").lower()
dependencies[key] = getattr(instance, method_name)
for name in ("stop_message", "stop_frontend", "clear_temp"):
dependency = MagicMock()
monkeypatch.setattr(modules_initializer, name, dependency)
dependencies[name] = dependency
async_redis = MagicMock()
async_redis.close = AsyncMock()
monkeypatch.setattr(
modules_initializer,
"AsyncRedisHelper",
MagicMock(return_value=async_redis),
)
dependencies["async_redis"] = async_redis.close
close_database = AsyncMock()
monkeypatch.setattr(modules_initializer, "close_database", close_database)
dependencies["close_database"] = close_database
return dependencies
def test_shared_http_close_waits_for_real_lru_eviction(monkeypatch):
"""最终 HTTP 关闭必须等待真实 LRU 淘汰任务并消费其异常"""
class FakeTransport:
created = []
def __init__(self, **_kwargs):
self.close_started = asyncio.Event()
self.release_close = asyncio.Event()
self.closed = False
self.fail_on_close = not self.created
if not self.fail_on_close:
self.release_close.set()
self.created.append(self)
async def aclose(self):
self.close_started.set()
await self.release_close.wait()
self.closed = True
if self.fail_on_close:
raise RuntimeError("eviction close failed")
monkeypatch.setattr(http_utils, "_MAX_SHARED_TRANSPORTS_PER_LOOP", 1)
monkeypatch.setattr(http_utils.httpx, "AsyncHTTPTransport", FakeTransport)
debug = MagicMock()
monkeypatch.setattr(http_utils.logger, "debug", debug)
async def run_test():
transport_kwargs = {
"proxy": None,
"verify": True,
"http2": False,
"max_keepalive_connections": 1,
"max_connections": 1,
}
evicted_transport = http_utils._get_shared_async_transport(
**transport_kwargs,
keepalive_expiry=1,
)
active_transport = http_utils._get_shared_async_transport(
**transport_kwargs,
keepalive_expiry=2,
)
await asyncio.wait_for(evicted_transport.close_started.wait(), timeout=1)
loop = asyncio.get_running_loop()
with http_utils._shared_async_transports_lock:
eviction_tasks = [
task
for task in http_utils._pending_eviction_tasks
if task.get_loop() is loop
]
assert len(eviction_tasks) == 1
close_task = asyncio.create_task(http_utils.aclose_shared_async_transports())
await asyncio.sleep(0)
try:
assert not close_task.done()
evicted_transport.release_close.set()
await close_task
await asyncio.sleep(0)
assert eviction_tasks[0].done()
assert evicted_transport.closed
assert active_transport.closed
with http_utils._shared_async_transports_lock:
assert not any(
task.get_loop() is loop
for task in http_utils._pending_eviction_tasks
)
finally:
evicted_transport.release_close.set()
active_transport.release_close.set()
await asyncio.gather(close_task, return_exceptions=True)
await http_utils.aclose_shared_async_transports()
asyncio.run(run_test())
debug.assert_any_call(
"LRU 淘汰共享 transport 时关闭失败: "
"RuntimeError('eviction close failed')"
)
def test_shared_http_close_ignores_eviction_from_other_loop():
"""当前事件循环关闭不能等待其他循环持有的淘汰任务"""
ready = threading.Event()
release = threading.Event()
failures = []
state = {}
def run_foreign_loop():
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
async def delayed_close():
while not release.is_set():
await asyncio.sleep(0.01)
task = loop.create_task(delayed_close())
state["task"] = task
with http_utils._shared_async_transports_lock:
http_utils._pending_eviction_tasks.add(task)
task.add_done_callback(http_utils._discard_pending_eviction_task)
ready.set()
try:
loop.run_until_complete(task)
loop.run_until_complete(asyncio.sleep(0))
except BaseException as err:
failures.append(err)
finally:
with http_utils._shared_async_transports_lock:
http_utils._pending_eviction_tasks.discard(task)
loop.close()
thread = threading.Thread(target=run_foreign_loop)
thread.start()
try:
assert ready.wait(timeout=2)
asyncio.run(http_utils.aclose_shared_async_transports())
assert thread.is_alive()
assert not state["task"].done()
finally:
release.set()
thread.join(timeout=2)
assert not thread.is_alive()
assert not failures
+149
View File
@@ -0,0 +1,149 @@
import threading
import time
from unittest.mock import MagicMock
from app.log import LogEntry, NonBlockingFileHandler, log_settings
def test_non_blocking_file_handler_shutdown_wakes_writer_and_closes_handlers(tmp_path):
"""日志关闭应立即唤醒空闲写线程,并关闭所有已打开的文件处理器"""
original_instance = NonBlockingFileHandler._instance
NonBlockingFileHandler._instance = None
handler = NonBlockingFileHandler()
handler._rotating_handlers = {}
log_handler = handler._get_rotating_handler(tmp_path / "shutdown.log")
try:
started_at = time.monotonic()
handler.shutdown()
elapsed = time.monotonic() - started_at
assert elapsed < 1
assert not handler._write_thread.is_alive()
assert log_handler.stream is None
assert handler._write_non_blocking(
LogEntry("info", "late-message", tmp_path / "shutdown.log")
) is False
assert handler._write_queue.empty()
finally:
if handler._write_thread.is_alive():
handler._running = False
handler._write_thread.join(timeout=5)
if log_handler.stream is not None:
log_handler.close()
NonBlockingFileHandler._instance = original_instance
def test_non_blocking_file_handler_shutdown_drains_queued_batches(monkeypatch, tmp_path):
"""停止标记之前已进入队列的日志应跨批次全部写完"""
original_instance = NonBlockingFileHandler._instance
NonBlockingFileHandler._instance = None
monkeypatch.setattr(log_settings, "BATCH_WRITE_SIZE", 2)
handler = NonBlockingFileHandler()
handler._rotating_handlers = {}
written = []
monkeypatch.setattr(
handler,
"_write_batch",
lambda batch: written.extend(entry.message for entry in batch),
)
try:
for index in range(5):
handler._write_non_blocking(
LogEntry("info", f"message-{index}", tmp_path / "drain.log")
)
handler.shutdown()
assert written == [f"message-{index}" for index in range(5)]
assert not handler._write_thread.is_alive()
finally:
if handler._write_thread.is_alive():
handler._running = False
handler._write_queue.put(handler._stop_sentinel)
handler._write_thread.join(timeout=5)
NonBlockingFileHandler._instance = original_instance
def test_non_blocking_file_handler_creates_one_handler_for_concurrent_first_write(monkeypatch, tmp_path):
"""同一路径首次并发写入时只创建并关闭一个文件处理器"""
original_instance = NonBlockingFileHandler._instance
NonBlockingFileHandler._instance = None
handler = NonBlockingFileHandler()
handler._rotating_handlers = {}
first_created = threading.Event()
second_started = threading.Event()
release_first = threading.Event()
created_handlers = []
results = []
class ProbeHandler:
def __init__(self, **kwargs):
self.closed = False
created_handlers.append(self)
if len(created_handlers) == 1:
first_created.set()
release_first.wait(timeout=2)
@staticmethod
def setFormatter(formatter):
pass
@staticmethod
def flush():
pass
def close(self):
self.closed = True
monkeypatch.setattr("app.log.RotatingFileHandler", ProbeHandler)
file_path = tmp_path / "concurrent.log"
def get_handler(started=None):
if started:
started.set()
results.append(handler._get_rotating_handler(file_path))
first = threading.Thread(target=get_handler)
second = threading.Thread(target=get_handler, args=(second_started,))
try:
first.start()
assert first_created.wait(timeout=1)
second.start()
assert second_started.wait(timeout=1)
time.sleep(0.05)
release_first.set()
first.join(timeout=2)
second.join(timeout=2)
assert len(created_handlers) == 1
assert results[0] is results[1]
handler.shutdown()
assert created_handlers[0].closed is True
finally:
release_first.set()
first.join(timeout=2)
second.join(timeout=2)
handler.shutdown()
NonBlockingFileHandler._instance = original_instance
def test_non_blocking_file_handler_uses_handler_lock(monkeypatch, tmp_path):
"""日志写入通过 Handler 入口串行化 emit 与 rollover"""
original_instance = NonBlockingFileHandler._instance
NonBlockingFileHandler._instance = None
handler = NonBlockingFileHandler()
handler._rotating_handlers = {}
log_handler = MagicMock()
monkeypatch.setattr(handler, "_get_rotating_handler", MagicMock(return_value=log_handler))
try:
handler._write_sync(LogEntry("info", "message", tmp_path / "locked.log"))
log_handler.handle.assert_called_once()
log_handler.emit.assert_not_called()
finally:
handler.shutdown()
NonBlockingFileHandler._instance = original_instance
+561
View File
@@ -20,6 +20,16 @@ def clear_media_interactions():
plugin_input_interaction_manager.clear()
@pytest.fixture(autouse=True)
def mock_default_media_search():
"""未显式验证搜索结果的消息路由用例不访问真实媒体元数据服务"""
with patch(
"app.chain.media.MediaChain.search",
side_effect=lambda title: (_build_meta(title), []),
):
yield
def _build_meta(name: str) -> MetaBase:
"""构造媒体识别元数据。"""
meta = MetaBase(name)
@@ -160,6 +170,110 @@ def test_message_routes_text_reply_to_media_interaction_before_ai():
handle_ai.assert_not_called()
def test_message_process_preserves_parser_message_id_context():
"""消息链不按渠道解释 message_id,只透传解析器给出的原消息上下文。"""
chain = MessageChain()
incoming = CommingMessage(
channel=MessageChannel.Telegram,
source="telegram-test",
userid="10001",
username="tester",
text="东张西望",
message_id=101,
chat_id="chat-a",
reply_to_message_id=99,
)
with patch.object(chain, "message_parser", return_value=incoming), patch.object(
chain, "handle_message"
) as handle_message:
chain.process(body=None, form=None, args={"source": "telegram-test"})
handle_message.assert_called_once()
kwargs = handle_message.call_args.kwargs
assert kwargs["original_message_id"] == 101
assert kwargs["original_chat_id"] == "chat-a"
assert kwargs["reply_to_message_id"] == 99
def test_message_process_keeps_callback_message_id_as_edit_context():
"""按钮回调的 message_id 仍应作为机器人原消息 ID 传递,供编辑原消息使用。"""
chain = MessageChain()
incoming = CommingMessage(
channel=MessageChannel.Telegram,
source="telegram-test",
userid="10001",
username="tester",
text="CALLBACK:demo",
is_callback=True,
message_id=101,
chat_id="chat-a",
)
with patch.object(chain, "message_parser", return_value=incoming), patch.object(
chain, "handle_message"
) as handle_message:
chain.process(body=None, form=None, args={"source": "telegram-test"})
handle_message.assert_called_once()
kwargs = handle_message.call_args.kwargs
assert kwargs["original_message_id"] == 101
assert kwargs["original_chat_id"] == "chat-a"
def test_message_process_preserves_non_telegram_plain_message_id():
"""非 Telegram 渠道保持旧行为,普通消息 ID 仍向下传递给渠道实现自行解释。"""
chain = MessageChain()
incoming = CommingMessage(
channel=MessageChannel.Slack,
source="slack-test",
userid="10001",
username="tester",
text="hello",
message_id="slack-message-ts",
chat_id="slack-channel",
)
with patch.object(chain, "message_parser", return_value=incoming), patch.object(
chain, "handle_message"
) as handle_message:
chain.process(body=None, form=None, args={"source": "slack-test"})
handle_message.assert_called_once()
kwargs = handle_message.call_args.kwargs
assert kwargs["original_message_id"] == "slack-message-ts"
assert kwargs["original_chat_id"] == "slack-channel"
def test_handle_message_keeps_legacy_positional_images_argument():
"""新增 reply_to_message_id 不应改变旧位置参数 images/audio/files 的含义。"""
chain = MessageChain()
images = [CommingMessage.MessageImage(ref="tg://file_id/photo-1")]
with patch.object(
chain, "_handle_plugin_input_interaction", return_value=False
), patch.object(
chain, "_mark_message_processing_started", return_value=None
), patch.object(
chain, "_mark_message_processing_finished"
), patch.object(chain, "_handle_message_core", return_value=False) as handle_core:
chain.handle_message(
MessageChannel.Telegram,
"telegram-test",
"10001",
"tester",
"带图消息",
None,
"chat-a",
images,
)
handle_core.assert_called_once()
kwargs = handle_core.call_args.kwargs
assert kwargs["images"] == images
assert kwargs["reply_to_message_id"] is None
def test_plugin_input_session_captures_plain_text_before_media_interaction():
"""插件输入会话存在时,普通文本应派发给插件而不是媒体交互。"""
chain = MessageChain()
@@ -209,6 +323,7 @@ def test_plugin_input_session_captures_plain_text_before_media_interaction():
"source": "wechat-test",
"username": "tester",
"chat_id": None,
"reply_to_message_id": None,
"prompt_id": "prompt-1",
"input_session_id": request.request_id,
"payload": {"step": "name"},
@@ -528,6 +643,326 @@ def test_plugin_input_session_does_not_capture_other_chat_text():
assert payload["chat_id"] == "chat-a"
def test_plugin_input_prompt_message_requires_matching_reply():
"""绑定提示消息 ID 的插件输入只应消费当前 ForceReply 回复。"""
chain = MessageChain()
request = plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Telegram,
source="telegram-test",
username="tester",
chat_id="chat-a",
prompt_message_id="prompt-current",
payload={"step": "keyword"},
)
with patch.object(chain, "_record_user_message") as record_message, patch.object(
chain.eventmanager, "send_event"
) as send_event:
chain.handle_message(
channel=MessageChannel.Telegram,
source="telegram-test",
userid="10001",
username="tester",
text="旧回复框文本",
original_chat_id="chat-a",
reply_to_message_id="prompt-old",
)
record_message.assert_called_once()
assert plugin_input_interaction_manager.get_by_user(
"10001", MessageChannel.Telegram, "telegram-test", "chat-a"
) == request
assert not any(
call.args and call.args[0] == EventType.MessageAction
for call in send_event.call_args_list
)
with patch.object(chain, "_record_user_message") as record_message, patch.object(
chain.eventmanager, "send_event"
) as send_event:
chain.handle_message(
channel=MessageChannel.Telegram,
source="telegram-test",
userid="10001",
username="tester",
text="当前回复框文本",
original_chat_id="chat-a",
reply_to_message_id="prompt-current",
)
record_message.assert_not_called()
send_event.assert_called_once()
event_type, payload = send_event.call_args.args
assert event_type == EventType.MessageAction
assert payload["input_session_id"] == request.request_id
assert payload["input_text"] == "当前回复框文本"
assert payload["reply_to_message_id"] == "prompt-current"
assert plugin_input_interaction_manager.get_by_user(
"10001", MessageChannel.Telegram, "telegram-test", "chat-a"
) is None
def test_plugin_input_prompt_message_matches_integer_reply_ids():
"""真实 Telegram message_id 为 int,应与内部 str 归一化后的 prompt_message_id 匹配。"""
chain = MessageChain()
request = plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Telegram,
source="telegram-test",
username="tester",
chat_id=10001,
prompt_message_id=99,
payload={"step": "keyword"},
)
with patch.object(chain.eventmanager, "send_event") as send_event:
chain.handle_message(
channel=MessageChannel.Telegram,
source="telegram-test",
userid="10001",
username="tester",
text="翡翠台",
original_chat_id=10001,
reply_to_message_id=99,
)
send_event.assert_called_once()
event_type, payload = send_event.call_args.args
assert event_type == EventType.MessageAction
assert payload["input_session_id"] == request.request_id
assert payload["input_text"] == "翡翠台"
def test_plugin_input_prompt_message_ignores_plain_text_without_reply():
"""用户未使用 ForceReply 回复框直接发文本时,绑定会话不应消费该文本。"""
chain = MessageChain()
request = plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Telegram,
source="telegram-test",
username="tester",
chat_id="chat-a",
prompt_message_id="prompt-current",
payload={"step": "keyword"},
)
with patch.object(chain, "_record_user_message") as record_message, patch.object(
chain.eventmanager, "send_event"
) as send_event:
chain.handle_message(
channel=MessageChannel.Telegram,
source="telegram-test",
userid="10001",
username="tester",
text="直接输入文本",
original_chat_id="chat-a",
)
record_message.assert_called_once()
assert plugin_input_interaction_manager.get_by_user(
"10001", MessageChannel.Telegram, "telegram-test", "chat-a"
) == request
assert not any(
call.args and call.args[0] == EventType.MessageAction
for call in send_event.call_args_list
)
def test_plugin_input_prompt_message_allows_direct_cancel_without_reply():
"""绑定 ForceReply 时,取消词应能直接结束会话,避免用户被残留回复框卡住。"""
chain = MessageChain()
request = plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Telegram,
source="telegram-test",
username="tester",
chat_id="chat-a",
prompt_message_id="prompt-current",
payload={"step": "keyword"},
)
with patch.object(chain, "_record_user_message") as record_message, patch.object(
chain.eventmanager, "send_event"
) as send_event, patch.object(chain, "post_message") as post_message:
chain.handle_message(
channel=MessageChannel.Telegram,
source="telegram-test",
userid="10001",
username="tester",
text="取消",
original_chat_id="chat-a",
)
record_message.assert_not_called()
send_event.assert_called_once()
event_type, payload = send_event.call_args.args
assert event_type == EventType.MessageAction
assert payload["input_session_id"] == request.request_id
assert payload["cancelled"] is True
post_message.assert_called_once()
assert plugin_input_interaction_manager.get_by_user(
"10001", MessageChannel.Telegram, "telegram-test", "chat-a"
) is None
def test_expired_prompt_message_cancel_text_falls_back_to_normal_search_without_notice():
"""绑定 ForceReply 过期后,即使输入取消词也应静默放行给普通文本链路。"""
chain = MessageChain()
plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Telegram,
source="telegram-test",
username="tester",
chat_id="chat-a",
prompt_message_id="prompt-expired",
timeout_seconds=60,
).created_at = datetime.now() - timedelta(seconds=61)
with patch.object(chain, "_record_user_message") as record_message, patch.object(
chain.eventmanager, "send_event"
) as send_event, patch.object(
chain, "_handle_message_core", return_value=False
) as handle_core:
chain.handle_message(
channel=MessageChannel.Telegram,
source="telegram-test",
userid="10001",
username="tester",
text="取消",
original_chat_id="chat-a",
reply_to_message_id="prompt-expired",
)
record_message.assert_called_once()
handle_core.assert_called_once()
assert not any(
call.args and call.args[0] == EventType.MessageAction
for call in send_event.call_args_list
)
assert plugin_input_interaction_manager.get_by_user(
"10001", MessageChannel.Telegram, "telegram-test", "chat-a"
) is None
def test_plugin_input_prompt_message_requires_matching_chat_id():
"""绑定提示消息 ID 时还必须匹配 chat_id,避免跨聊天同号消息误消费。"""
chain = MessageChain()
request = plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Telegram,
source="telegram-test",
username="tester",
chat_id="chat-a",
prompt_message_id="prompt-current",
payload={"step": "keyword"},
)
with patch.object(chain, "_record_user_message") as record_message, patch.object(
chain.eventmanager, "send_event"
) as send_event:
chain.handle_message(
channel=MessageChannel.Telegram,
source="telegram-test",
userid="10001",
username="tester",
text="其他聊天同号回复",
original_chat_id="chat-b",
reply_to_message_id="prompt-current",
)
record_message.assert_called_once()
assert plugin_input_interaction_manager.get_by_user(
"10001", MessageChannel.Telegram, "telegram-test", "chat-a"
) == request
assert not any(
call.args and call.args[0] == EventType.MessageAction
for call in send_event.call_args_list
)
def test_expired_prompt_message_input_falls_back_to_normal_search_without_notice():
"""回复过期 ForceReply 时不提示插件输入超时,交回普通文本搜索。"""
chain = MessageChain()
request = plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Telegram,
source="telegram-test",
username="tester",
chat_id="chat-a",
prompt_message_id="prompt-expired",
timeout_seconds=60,
)
request.created_at = datetime.now() - timedelta(seconds=61)
with patch.object(chain, "_record_user_message") as record_message, patch.object(
chain.eventmanager, "send_event"
) as send_event, patch.object(chain, "post_message") as post_message:
chain.handle_message(
channel=MessageChannel.Telegram,
source="telegram-test",
userid="10001",
username="tester",
text="过期回复框文本",
original_chat_id="chat-a",
reply_to_message_id="prompt-expired",
)
record_message.assert_called_once()
post_message.assert_not_called()
assert not any(
call.args and call.args[0] == EventType.MessageAction
for call in send_event.call_args_list
)
assert plugin_input_interaction_manager.get_by_user(
"10001", MessageChannel.Telegram, "telegram-test", "chat-a"
) is None
def test_expired_prompt_message_without_reply_falls_back_and_clears_state():
"""绑定会话过期后,未命中回复框的文本也应放行并清理过期状态。"""
chain = MessageChain()
request = plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Telegram,
source="telegram-test",
username="tester",
chat_id="chat-a",
prompt_message_id="prompt-expired",
timeout_seconds=60,
)
request.created_at = datetime.now() - timedelta(seconds=61)
with patch.object(chain, "_record_user_message") as record_message, patch.object(
chain.eventmanager, "send_event"
) as send_event:
chain.handle_message(
channel=MessageChannel.Telegram,
source="telegram-test",
userid="10001",
username="tester",
text="过期后直接输入",
original_chat_id="chat-a",
)
record_message.assert_called_once()
assert not any(
call.args and call.args[0] == EventType.MessageAction
for call in send_event.call_args_list
)
assert plugin_input_interaction_manager.get_by_user(
"10001", MessageChannel.Telegram, "telegram-test", "chat-a"
) is None
def test_plugin_input_chatless_session_keeps_legacy_chat_fallback():
"""旧插件未绑定 chat_id 时,同 source 消息仍可兼容消费。"""
chain = MessageChain()
@@ -875,6 +1310,92 @@ def test_plugin_input_session_with_no_channel_and_no_source_does_not_match_speci
)
def test_plugin_input_create_or_replace_keeps_legacy_positional_timeout_and_payload():
"""新增 prompt_message_id 不应改变旧位置参数 timeout_seconds/payload 的含义。"""
request = plugin_input_interaction_manager.create_or_replace(
"10001",
"demo_plugin",
MessageChannel.Telegram,
"telegram-test",
"tester",
"chat-a",
"prompt-id",
30,
{"step": "legacy"},
)
assert request.timeout_seconds == 30
assert request.payload == {"step": "legacy"}
assert request.prompt_message_id is None
def test_plugin_input_create_or_replace_ignores_prompt_message_without_chat_id():
"""缺少 chat_id 时不启用 prompt_message_id 绑定,避免创建永远无法消费的会话。"""
request = plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Telegram,
source="telegram-test",
username="tester",
prompt_message_id="prompt-current",
)
assert request.chat_id is None
assert request.prompt_message_id is None
def test_plugin_input_create_or_replace_ignores_prompt_message_for_non_telegram_channel():
"""非 Telegram 渠道不启用 prompt_message_id 绑定,避免渠道无法上报回复 ID 时卡死。"""
request = plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Slack,
source="slack-test",
username="tester",
chat_id="slack-channel",
prompt_message_id="prompt-current",
)
assert request.chat_id == "slack-channel"
assert request.prompt_message_id is None
consumed, status = plugin_input_interaction_manager.consume_by_user(
"10001",
MessageChannel.Slack,
"slack-test",
"slack-channel",
)
assert consumed == request
assert status == "active"
def test_plugin_input_bypass_reply_check_still_requires_matching_chat_id():
"""取消词绕过 reply_id 校验时,仍必须匹配绑定会话的 chat_id。"""
request = plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Telegram,
source="telegram-test",
username="tester",
chat_id="chat-a",
prompt_message_id="prompt-current",
)
consumed, status = plugin_input_interaction_manager.consume_by_user(
"10001",
MessageChannel.Telegram,
"telegram-test",
"chat-b",
bypass_reply_check=True,
)
assert consumed is None
assert status is None
assert plugin_input_interaction_manager.get_by_user(
"10001", MessageChannel.Telegram, "telegram-test", "chat-a"
) == request
def test_plugin_input_specific_session_replaces_overlapping_no_channel_session():
"""同用户创建具体渠道会话时,应替换重叠的无渠道会话,避免下一条消息被连环接管。"""
old_request = plugin_input_interaction_manager.create_or_replace(
@@ -918,6 +1439,46 @@ def test_plugin_input_session_pop_by_user_consumes_once():
) is None
def test_plugin_input_session_pop_by_user_ignores_prompt_message_binding():
"""主动清理会话时不应要求提供 ForceReply 的 reply_to_message_id。"""
request = plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Telegram,
source="telegram-test",
username="tester",
prompt_message_id="prompt-current",
)
assert plugin_input_interaction_manager.pop_by_user(
"10001", MessageChannel.Telegram, "telegram-test"
) == request
assert plugin_input_interaction_manager.get_by_user(
"10001", MessageChannel.Telegram, "telegram-test"
) is None
def test_plugin_input_session_pop_by_user_removes_expired_prompt_session():
"""主动清理已过期会话时,也应移除过期表中的绑定 ForceReply 会话。"""
request = plugin_input_interaction_manager.create_or_replace(
user_id="10001",
plugin_id="demo_plugin",
channel=MessageChannel.Telegram,
source="telegram-test",
username="tester",
prompt_message_id="prompt-current",
timeout_seconds=60,
)
request.created_at = datetime.now() - timedelta(seconds=61)
assert plugin_input_interaction_manager.pop_by_user(
"10001", MessageChannel.Telegram, "telegram-test"
) == request
assert plugin_input_interaction_manager.pop_by_user(
"10001", MessageChannel.Telegram, "telegram-test"
) is None
def test_target_plugin_filter_only_allows_target_plugin_handler():
"""带目标插件的输入事件不应投递给其他插件或模块级处理器。"""
+30
View File
@@ -0,0 +1,30 @@
import time
from app.helper.message import MessageQueueManager, TemplateHelper, stop_message
from app.utils.singleton import SingletonClass
def test_message_queue_stop_wakes_idle_monitor(monkeypatch):
"""消息队列停止时应唤醒空闲监控线程,不等待完整检查周期"""
monkeypatch.setattr(MessageQueueManager, "init_config", lambda self: None)
manager = object.__new__(MessageQueueManager)
manager.__init__(check_interval=10)
started_at = time.monotonic()
manager.stop()
elapsed = time.monotonic() - started_at
assert elapsed < 1
assert not manager.thread.is_alive()
def test_stop_message_does_not_initialize_absent_services(monkeypatch):
"""消息服务未初始化时,关闭入口不应为了清理而创建后台资源"""
monkeypatch.setattr(SingletonClass, "_instances", {})
assert MessageQueueManager.get_existing_instance() is None
assert TemplateHelper.get_existing_instance() is None
stop_message()
assert MessageQueueManager not in SingletonClass._instances
assert TemplateHelper not in SingletonClass._instances

Some files were not shown because too many files have changed in this diff Show More