mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-08-27 19:20:19 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b1b6a81cef | ||
|
|
c6b94d4908 | ||
|
|
52ca375f3d | ||
|
|
d8adb4fbfe | ||
|
|
51d2ed1200 | ||
|
|
702801d0dc | ||
|
|
3e32eab98f | ||
|
|
4292678672 | ||
|
|
7c3f9629bf | ||
|
|
93761fe7e4 | ||
|
|
593139faac | ||
|
|
6c89f1eb4b | ||
|
|
2310a3a456 | ||
|
|
48852350a0 | ||
|
|
a23fce1491 | ||
|
|
c976741574 | ||
|
|
04facef64d | ||
|
|
33a97eb2c8 | ||
|
|
cf80b551f9 | ||
|
|
e011b20210 | ||
|
|
bdf395f494 | ||
|
|
68686bc23a | ||
|
|
1a528c7803 | ||
|
|
6b4a255f26 | ||
|
|
1a3c1b8b39 | ||
|
|
8788dae34b | ||
|
|
bb00814d7a | ||
|
|
3d55d44457 | ||
|
|
7a5e565b15 | ||
|
|
cae28d8c03 | ||
|
|
3d020c8ceb | ||
|
|
1b065cc08b | ||
|
|
a1a6376adc | ||
|
|
197a09b2a4 | ||
|
|
8d099b9581 | ||
|
|
ff9ba79b60 | ||
|
|
6354a48405 | ||
|
|
f6df6cc093 | ||
|
|
98e69d1b45 | ||
|
|
4b9af5b8c7 | ||
|
|
059a50f7f8 | ||
|
|
14fed2d70b | ||
|
|
875984ad39 | ||
|
|
297cd04fbc | ||
|
|
3dde94be0f | ||
|
|
98ee939236 | ||
|
|
c6611f6210 | ||
|
|
503ee90c0c | ||
|
|
fb32c59713 | ||
|
|
a3c90c64ca | ||
|
|
de97cb3c0a | ||
|
|
3b709b7f2e | ||
|
|
6f8b6cfbc9 | ||
|
|
e3f80af74f | ||
|
|
1d708870c9 | ||
|
|
053e1b7562 | ||
|
|
4ca3e40507 | ||
|
|
318d2ab7d7 | ||
|
|
7c2390908a | ||
|
|
5c2b503a74 | ||
|
|
44fa202778 | ||
|
|
ed92be08af | ||
|
|
9ed0704c5b | ||
|
|
e46b4e5ba0 | ||
|
|
87ad7988b2 | ||
|
|
1382975b18 | ||
|
|
d9a42c672a | ||
|
|
b042086efa | ||
|
|
36fefa14e0 | ||
|
|
1332576c3f | ||
|
|
4300af0e9c | ||
|
|
405350c774 | ||
|
|
d666134ed2 | ||
|
|
5588e37c6d | ||
|
|
2056aa0b2c | ||
|
|
b0ff3ae3c7 | ||
|
|
31544629b4 | ||
|
|
142393f2d3 | ||
|
|
5cf79e0360 | ||
|
|
f152a0381d | ||
|
|
428c19b6ba | ||
|
|
8b5524a321 | ||
|
|
b972b46747 | ||
|
|
0598fbdd75 | ||
|
|
572299a45e | ||
|
|
229824a417 | ||
|
|
a0ee99aacc | ||
|
|
92918ce380 | ||
|
|
a4335fe753 | ||
|
|
107ba37834 | ||
|
|
c27678ce06 | ||
|
|
7725342a80 | ||
|
|
893269f8c1 | ||
|
|
00d46f3aab | ||
|
|
077241b6ed | ||
|
|
b24a07e388 | ||
|
|
f814c271cc | ||
|
|
e015c67689 | ||
|
|
98b16bda8d | ||
|
|
b8233e1789 | ||
|
|
83107bf447 | ||
|
|
3a2f90c567 | ||
|
|
4826e3301c | ||
|
|
1855ba81ec | ||
|
|
96ef431efc | ||
|
|
4f2935c85e | ||
|
|
2a49495e27 | ||
|
|
b628bc7209 | ||
|
|
29068a5846 | ||
|
|
51a7120c79 | ||
|
|
476dfef7d9 | ||
|
|
bd5ddd6158 | ||
|
|
a30a48b8f4 | ||
|
|
8e60e5571b | ||
|
|
18c1ec4b82 | ||
|
|
30b932e07e | ||
|
|
54be1143fc | ||
|
|
13f27854fd | ||
|
|
770201c48c | ||
|
|
685f044312 | ||
|
|
8c0afac5d1 | ||
|
|
099ef7d5bf | ||
|
|
f3ac69669c | ||
|
|
eb4ecd990a | ||
|
|
b51971ee7d | ||
|
|
6f6ed998bb | ||
|
|
844407dc41 | ||
|
|
c54605f8ce | ||
|
|
0fbf05d72f | ||
|
|
09bb32f681 | ||
|
|
a37f118576 | ||
|
|
e635bc8e04 | ||
|
|
8245124e82 | ||
|
|
827ed8330c | ||
|
|
136c1baed3 | ||
|
|
992031ef95 | ||
|
|
b16c50b03a | ||
|
|
76803ae7a3 | ||
|
|
56bda11947 | ||
|
|
1b12d7664e | ||
|
|
db9960d9b9 | ||
|
|
2f0c1252da | ||
|
|
36d4434596 | ||
|
|
93e907d032 | ||
|
|
132f27c1c6 | ||
|
|
b231ad415f | ||
|
|
0f183ae08e | ||
|
|
a71d3ea03f | ||
|
|
7f82a9ea4d | ||
|
|
d977e4c48a | ||
|
|
95b6adbeee | ||
|
|
964fee1106 | ||
|
|
656473f3aa | ||
|
|
ab5995a609 | ||
|
|
064e6535d5 | ||
|
|
cab2ac400a | ||
|
|
d14d401c86 | ||
|
|
6c3c5e042d | ||
|
|
f3e5be37fd | ||
|
|
d8f7fa70af | ||
|
|
6916ee0988 | ||
|
|
6fef533527 | ||
|
|
c57985d553 | ||
|
|
ec07379a67 | ||
|
|
2764742b86 | ||
|
|
a0f613fa1e | ||
|
|
73d5c95f4e | ||
|
|
4d30dee74c | ||
|
|
302d8bbf5c | ||
|
|
b646cbb4f6 | ||
|
|
dd73b97095 | ||
|
|
0cb0bac0e1 | ||
|
|
9eb71c744b | ||
|
|
8bf826faa0 | ||
|
|
df4e45c644 | ||
|
|
494f809ef0 | ||
|
|
a4f6e13881 | ||
|
|
36fb82b7aa | ||
|
|
9b1bdb0cb2 | ||
|
|
2a89bfd25c | ||
|
|
d3d1c18316 | ||
|
|
027330f714 | ||
|
|
6fc1c672ad | ||
|
|
d383c9ffd1 | ||
|
|
27e1f634cb | ||
|
|
a9197c434e | ||
|
|
2a8708498c | ||
|
|
098db56a36 | ||
|
|
a8ccb08dd3 | ||
|
|
ff7a238309 | ||
|
|
318cfe68e9 | ||
|
|
0f98cda3b4 | ||
|
|
544ed6d84d | ||
|
|
bb9b6ec5d0 | ||
|
|
7d2a730b0c | ||
|
|
1b6a548dee | ||
|
|
52c5f2900f | ||
|
|
43e89ebf77 | ||
|
|
bc52653ec1 | ||
|
|
4233ebfba6 | ||
|
|
af32c5e9bb | ||
|
|
cd02e55879 | ||
|
|
d49c57b11e | ||
|
|
724b4a59d5 | ||
|
|
109be58abe | ||
|
|
c1d8676b25 | ||
|
|
c1cefa3f40 | ||
|
|
0c53fb86fd | ||
|
|
126279c63b | ||
|
|
d7aa66853a | ||
|
|
d18c2d6f72 | ||
|
|
3e3883a57f | ||
|
|
2ebe7c27c2 | ||
|
|
0cd049bfc2 | ||
|
|
dc773337d3 | ||
|
|
5c649ff1d1 | ||
|
|
3407cc8edd | ||
|
|
f9ea0118d9 | ||
|
|
ad73434e2c | ||
|
|
a6afa0fbc0 | ||
|
|
3306d196b7 | ||
|
|
e44a6f41b5 | ||
|
|
7a7b27858e |
@@ -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:
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
name: PR-Agent
|
||||
|
||||
on:
|
||||
pull_request_target:
|
||||
# Fork 审查需要目标仓库凭据;该 job 仅通过 GitHub API 读取 PR 内容,不 checkout 或执行 PR 分支代码。
|
||||
types:
|
||||
- opened
|
||||
- reopened
|
||||
- ready_for_review
|
||||
- review_requested
|
||||
- synchronize
|
||||
issue_comment:
|
||||
types:
|
||||
- created
|
||||
- edited
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: write
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
pr-agent:
|
||||
if: >-
|
||||
github.event.sender.type != 'Bot' &&
|
||||
(
|
||||
github.event_name == 'pull_request_target' ||
|
||||
(
|
||||
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)
|
||||
)
|
||||
)
|
||||
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 Review
|
||||
uses: docker://ghcr.io/infinitypacer/pr-review-runner:latest
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ github.token }}
|
||||
OPENAI_KEY: ${{ secrets.OPENAI_KEY }}
|
||||
OPENAI.API_BASE: ${{ secrets.OPENAI_API_BASE }}
|
||||
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
|
||||
@@ -23,24 +23,15 @@ jobs:
|
||||
uses: actions/cache@v4
|
||||
with:
|
||||
path: ~/.cache/pip
|
||||
key: ${{ runner.os }}-pip-${{ hashFiles('**/requirements.txt', '**/requirements.in') }}
|
||||
key: ${{ runner.os }}-pip-${{ hashFiles('**/requirements.in', '**/requirements-dev.in', '**/requirements.txt') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-pip-
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip setuptools wheel
|
||||
pip install pylint
|
||||
# 安装项目依赖
|
||||
if [ -f requirements.txt ]; then
|
||||
echo "📦 安装 requirements.txt 中的依赖..."
|
||||
pip install -r requirements.txt
|
||||
elif [ -f requirements.in ]; then
|
||||
echo "📦 安装 requirements.in 中的依赖..."
|
||||
pip install -r requirements.in
|
||||
else
|
||||
echo "⚠️ 未找到依赖文件,仅安装 pylint"
|
||||
fi
|
||||
# Pylint 属于开发/静态检查依赖,统一通过 dev 入口安装。
|
||||
pip install -r requirements-dev.in
|
||||
|
||||
- name: Verify pylint config
|
||||
run: |
|
||||
@@ -88,4 +79,4 @@ jobs:
|
||||
run: |
|
||||
echo "🎉 Pylint 检查完成!"
|
||||
echo "✅ 没有发现语法错误或严重问题"
|
||||
echo "📊 详细报告已保存为构建工件"
|
||||
echo "📊 详细报告已保存为构建工件"
|
||||
|
||||
@@ -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"
|
||||
@@ -38,16 +38,15 @@ jobs:
|
||||
uses: actions/cache@v5
|
||||
with:
|
||||
path: ~/.cache/pip
|
||||
key: ${{ runner.os }}-pip-${{ hashFiles('**/requirements.in', '**/requirements.txt') }}
|
||||
key: ${{ runner.os }}-pip-${{ hashFiles('**/requirements.in', '**/requirements-dev.in', '**/requirements.txt') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-pip-
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip setuptools wheel
|
||||
# 用 requirements.in 还原 CI / 全新环境(含 pytest 与 moviepilot-rust 等可选扩展),
|
||||
# 与本地"干净 venv 复现"一致;测试运行器 pytest 已在 requirements.in 中声明。
|
||||
pip install -r requirements.in
|
||||
# 单测需要开发/测试依赖;运行时入口 requirements.in 不携带测试与构建辅助工具。
|
||||
pip install -r requirements-dev.in
|
||||
|
||||
- name: Run tests
|
||||
timeout-minutes: 10
|
||||
|
||||
+3
-1
@@ -16,7 +16,7 @@ app/helper/*.pyd
|
||||
app/helper/*.bin
|
||||
app/plugins/**
|
||||
!app/plugins/__init__.py
|
||||
config/cookies/**
|
||||
config/cookies/
|
||||
config/app.env
|
||||
config/user.db*
|
||||
config/sites/**
|
||||
@@ -25,6 +25,7 @@ config/logs/
|
||||
config/plugins/
|
||||
config/temp/
|
||||
config/cache/
|
||||
config/.cache/
|
||||
.runtime/
|
||||
public/
|
||||
.moviepilot.env
|
||||
@@ -36,6 +37,7 @@ coverage.json
|
||||
htmlcov/
|
||||
.vscode
|
||||
venv
|
||||
moviepilot-site-capture-*.zip
|
||||
|
||||
# Pylint
|
||||
pylint-report.json
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
+336
-37
@@ -1,4 +1,5 @@
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import re
|
||||
import traceback
|
||||
@@ -50,11 +51,14 @@ 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
|
||||
from app.db.agentchat_oper import AgentChatOper
|
||||
from app.db.agenttask_oper import AgentTaskOper
|
||||
from app.db.user_oper import UserOper
|
||||
from app.log import logger
|
||||
from app.schemas import AgentLLMProviderEventData, AgentTokensUsageEventData, Notification, NotificationType
|
||||
@@ -64,6 +68,8 @@ from app.utils.identity import SYSTEM_INTERNAL_USER_ID
|
||||
|
||||
|
||||
class AgentChain(ChainBase):
|
||||
"""Agent 业务处理链。"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
@@ -148,6 +154,16 @@ class _SessionUsageSnapshot:
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class _CompiledAgentBundle:
|
||||
"""会话内可复用的 Agent 图及其构造签名。"""
|
||||
|
||||
signature: tuple[Any, ...]
|
||||
agent: Any
|
||||
streaming: bool
|
||||
created_at: datetime
|
||||
|
||||
|
||||
class _ThinkTagStripper:
|
||||
"""
|
||||
流式剥离 <think>...</think> 标签的辅助类。
|
||||
@@ -280,6 +296,8 @@ class MoviePilotAgent:
|
||||
self._llm_runtime_config: Optional[Dict[str, Any]] = None
|
||||
self._llm_provider_selection: Dict[str, Any] = {}
|
||||
self._agent_started_at: Optional[datetime] = None
|
||||
self._compiled_agent_bundle: Optional[_CompiledAgentBundle] = None
|
||||
self._last_agent_cache_hit = False
|
||||
|
||||
# 流式token管理
|
||||
self.stream_handler = StreamingHandler()
|
||||
@@ -697,7 +715,7 @@ class MoviePilotAgent:
|
||||
"""
|
||||
通过链式事件解析本次 Agent 可用的 LLM 运行时配置。
|
||||
|
||||
若没有插件返回 selected_provider_id,则沿用系统配置,保持既有行为。
|
||||
插件未返回有效配置时沿用系统配置,显式返回的配置优先。
|
||||
"""
|
||||
if self._llm_runtime_config is not None:
|
||||
return self._llm_runtime_config
|
||||
@@ -710,7 +728,8 @@ class MoviePilotAgent:
|
||||
base_url_preset=settings.LLM_BASE_URL_PRESET,
|
||||
user_agent=settings.LLM_USER_AGENT,
|
||||
use_proxy=settings.LLM_USE_PROXY,
|
||||
thinking_level=None,
|
||||
thinking_level=settings.LLM_THINKING_LEVEL,
|
||||
api_protocol=settings.LLM_API_PROTOCOL,
|
||||
)
|
||||
selected_event = await eventmanager.async_send_event(
|
||||
ChainEventType.AgentLLMProvider,
|
||||
@@ -745,9 +764,15 @@ class MoviePilotAgent:
|
||||
use_proxy = self._get_event_value(resolved_data, "use_proxy")
|
||||
if use_proxy is None:
|
||||
use_proxy = settings.LLM_USE_PROXY
|
||||
thinking_level = self._clean_optional_text(
|
||||
self._get_event_value(resolved_data, "thinking_level")
|
||||
thinking_level = (
|
||||
self._clean_optional_text(
|
||||
self._get_event_value(resolved_data, "thinking_level")
|
||||
)
|
||||
or settings.LLM_THINKING_LEVEL
|
||||
)
|
||||
api_protocol = self._clean_optional_text(
|
||||
self._get_event_value(resolved_data, "api_protocol")
|
||||
) or settings.LLM_API_PROTOCOL
|
||||
selected_provider_id = self._clean_optional_text(
|
||||
self._get_event_value(resolved_data, "selected_provider_id")
|
||||
)
|
||||
@@ -773,6 +798,7 @@ class MoviePilotAgent:
|
||||
"user_agent": user_agent,
|
||||
"use_proxy": bool(use_proxy),
|
||||
"thinking_level": thinking_level,
|
||||
"api_protocol": api_protocol,
|
||||
}
|
||||
return self._llm_runtime_config
|
||||
|
||||
@@ -838,6 +864,8 @@ class MoviePilotAgent:
|
||||
detail = cls._exception_detail_text(error).lower()
|
||||
if "no endpoints found that support image input" in detail:
|
||||
return True
|
||||
if "not a vlm" in detail or "text-only prompts" in detail:
|
||||
return True
|
||||
if "unknown variant" in detail and "image_url" in detail:
|
||||
return True
|
||||
if "image input" not in detail and "images" not in detail:
|
||||
@@ -978,6 +1006,93 @@ class MoviePilotAgent:
|
||||
allow_message_tools=self.allow_message_tools,
|
||||
)
|
||||
|
||||
def _refresh_tool_context(self, values: Dict[str, object]) -> None:
|
||||
"""
|
||||
刷新本轮工具共享上下文。
|
||||
|
||||
工具对象可能随会话内 Agent 图缓存被复用,因此这里保留 dict 对象本身,
|
||||
只替换其中内容,确保缓存工具看到的是最新权限与回复状态。
|
||||
"""
|
||||
self._tool_context.clear()
|
||||
self._tool_context.update(values)
|
||||
|
||||
@staticmethod
|
||||
def _public_runtime_config_signature(runtime_config: Dict[str, Any]) -> tuple:
|
||||
"""生成不包含密钥明文的 LLM 运行时签名。"""
|
||||
api_key = runtime_config.get("api_key") or ""
|
||||
api_key_digest = (
|
||||
hashlib.sha256(str(api_key).encode("utf-8")).hexdigest()[:12]
|
||||
if api_key
|
||||
else ""
|
||||
)
|
||||
return (
|
||||
runtime_config.get("provider"),
|
||||
runtime_config.get("model"),
|
||||
api_key_digest,
|
||||
runtime_config.get("base_url"),
|
||||
runtime_config.get("base_url_preset"),
|
||||
runtime_config.get("user_agent"),
|
||||
bool(runtime_config.get("use_proxy")),
|
||||
runtime_config.get("thinking_level"),
|
||||
runtime_config.get("api_protocol"),
|
||||
)
|
||||
|
||||
async def _agent_bundle_signature(self, streaming: bool) -> tuple[Any, ...]:
|
||||
"""构造会话内 Agent 图缓存签名。"""
|
||||
runtime_config = await self._resolve_llm_runtime_config()
|
||||
return (
|
||||
streaming,
|
||||
self.channel,
|
||||
self.source,
|
||||
self.user_id,
|
||||
self.username,
|
||||
self.allow_message_tools,
|
||||
bool(self._tool_context.get("is_admin")),
|
||||
self.has_message_context,
|
||||
self.is_background,
|
||||
settings.AI_AGENT_VERBOSE,
|
||||
settings.LLM_TEMPERATURE,
|
||||
settings.LLM_MAX_TOOLS,
|
||||
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(
|
||||
self, signature: tuple[Any, ...], streaming: bool
|
||||
) -> Optional[Any]:
|
||||
"""按签名读取当前会话已编译的 Agent 图。"""
|
||||
bundle = self._compiled_agent_bundle
|
||||
if (
|
||||
bundle
|
||||
and bundle.streaming == streaming
|
||||
and bundle.signature == signature
|
||||
):
|
||||
return bundle.agent
|
||||
return None
|
||||
|
||||
def _cache_agent(
|
||||
self,
|
||||
*,
|
||||
signature: tuple[Any, ...],
|
||||
agent: Any,
|
||||
streaming: bool,
|
||||
) -> Any:
|
||||
"""保存当前会话可复用的 Agent 图。"""
|
||||
self._compiled_agent_bundle = _CompiledAgentBundle(
|
||||
signature=signature,
|
||||
agent=agent,
|
||||
streaming=streaming,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
return agent
|
||||
|
||||
@staticmethod
|
||||
def _latest_turn_messages(messages: List[BaseMessage]) -> List[BaseMessage]:
|
||||
"""从完整历史中提取本轮新增用户消息。"""
|
||||
return [messages[-1]] if messages else []
|
||||
|
||||
def _initialize_subagent_tools(self) -> List:
|
||||
"""
|
||||
初始化子代理专用静默工具列表。
|
||||
@@ -998,12 +1113,52 @@ 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)
|
||||
:param streaming: 是否启用流式输出
|
||||
"""
|
||||
try:
|
||||
bundle_signature = await self._agent_bundle_signature(streaming)
|
||||
cached_agent = self._get_cached_agent(bundle_signature, streaming)
|
||||
self._last_agent_cache_hit = bool(cached_agent)
|
||||
if cached_agent:
|
||||
logger.debug(f"复用会话内 Agent 图: session_id={self.session_id}")
|
||||
return cached_agent
|
||||
|
||||
# 系统提示词
|
||||
system_prompt = prompt_manager.get_agent_prompt(channel=self.channel)
|
||||
|
||||
@@ -1020,6 +1175,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"),
|
||||
@@ -1036,9 +1192,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
|
||||
@@ -1111,13 +1269,18 @@ class MoviePilotAgent:
|
||||
)
|
||||
)
|
||||
|
||||
return create_agent(
|
||||
agent = create_agent(
|
||||
model=agent_model,
|
||||
tools=[*tools, *skill_tools, *activity_log_tools],
|
||||
system_prompt=system_prompt,
|
||||
middleware=middlewares,
|
||||
checkpointer=InMemorySaver(),
|
||||
)
|
||||
return self._cache_agent(
|
||||
signature=bundle_signature,
|
||||
agent=agent,
|
||||
streaming=streaming,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"创建 Agent 失败: {e}")
|
||||
raise e
|
||||
@@ -1135,12 +1298,15 @@ class MoviePilotAgent:
|
||||
user_display_saved = False
|
||||
try:
|
||||
logger.info(
|
||||
f"Agent推理: session_id={self.session_id}, input={message}, "
|
||||
f"Agent推理: session_id={self.session_id}, "
|
||||
f"input_chars={len(message or '')}, "
|
||||
f"images={len(images) if images else 0}, files={len(files) if files else 0}, "
|
||||
f"audio_input={has_audio_input}"
|
||||
)
|
||||
self._tool_context = await self._build_tool_context(
|
||||
should_dispatch_reply=self.should_dispatch_reply
|
||||
self._refresh_tool_context(
|
||||
await self._build_tool_context(
|
||||
should_dispatch_reply=self.should_dispatch_reply
|
||||
)
|
||||
)
|
||||
self._streamed_output = ""
|
||||
|
||||
@@ -1328,6 +1494,11 @@ class MoviePilotAgent:
|
||||
|
||||
# 创建智能体(根据是否流式传入不同 LLM)
|
||||
agent = await self._create_agent(streaming=use_streaming)
|
||||
input_messages = (
|
||||
self._latest_turn_messages(messages)
|
||||
if self._last_agent_cache_hit
|
||||
else messages
|
||||
)
|
||||
|
||||
if use_streaming:
|
||||
self.stream_handler.set_dispatch_policy(
|
||||
@@ -1346,7 +1517,7 @@ class MoviePilotAgent:
|
||||
# 流式运行智能体,token 直接推送到 stream_handler
|
||||
await self._stream_agent_tokens(
|
||||
agent=agent,
|
||||
messages={"messages": messages},
|
||||
messages={"messages": input_messages},
|
||||
config=agent_config,
|
||||
on_token=self._handle_stream_text,
|
||||
)
|
||||
@@ -1385,7 +1556,7 @@ class MoviePilotAgent:
|
||||
else:
|
||||
# 非流式模式:后台任务或渠道不支持消息编辑
|
||||
await agent.ainvoke(
|
||||
{"messages": messages},
|
||||
{"messages": input_messages},
|
||||
config=agent_config,
|
||||
)
|
||||
|
||||
@@ -1444,9 +1615,11 @@ class MoviePilotAgent:
|
||||
|
||||
except asyncio.CancelledError:
|
||||
logger.info(f"Agent执行被取消: session_id={self.session_id}")
|
||||
self._compiled_agent_bundle = None
|
||||
execution_error = "任务已取消"
|
||||
return "任务已取消", {}
|
||||
except Exception as e:
|
||||
self._compiled_agent_bundle = None
|
||||
execution_error = str(e)
|
||||
if self._messages_have_image_input(messages) and self._is_unsupported_image_input_error(e):
|
||||
logger.warning(
|
||||
@@ -1467,22 +1640,24 @@ class MoviePilotAgent:
|
||||
if not streaming_stopped:
|
||||
await self.stream_handler.stop_streaming()
|
||||
|
||||
async def send_agent_message(self, message: str, title: str = ""):
|
||||
async def send_agent_message(self, message: str, title: str = "") -> None:
|
||||
"""
|
||||
通过原渠道发送消息给用户
|
||||
发送 Agent 消息;后台任务不绑定原渠道,交由通知链广播。
|
||||
"""
|
||||
broadcast = self.is_background
|
||||
self._save_assistant_display_message_once(message)
|
||||
await AgentChain().async_post_message(
|
||||
Notification(
|
||||
channel=self.channel,
|
||||
source=self.source,
|
||||
channel=None if broadcast else self.channel,
|
||||
source=None if broadcast else self.source,
|
||||
mtype=NotificationType.Agent,
|
||||
userid=self.user_id,
|
||||
username=self.username,
|
||||
original_message_id=self.original_message_id,
|
||||
original_chat_id=self.original_chat_id,
|
||||
userid=None if broadcast else self.user_id,
|
||||
username=self.username or (settings.SUPERUSER if broadcast else None),
|
||||
original_message_id=None if broadcast else self.original_message_id,
|
||||
original_chat_id=None if broadcast else self.original_chat_id,
|
||||
title=title,
|
||||
text=message,
|
||||
save_history=False,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -1490,6 +1665,7 @@ class MoviePilotAgent:
|
||||
"""
|
||||
清理智能体资源
|
||||
"""
|
||||
self._compiled_agent_bundle = None
|
||||
logger.info(f"MoviePilot智能体已清理: session_id={self.session_id}")
|
||||
|
||||
|
||||
@@ -1513,6 +1689,10 @@ class _MessageTask:
|
||||
processing_status: Optional[dict] = None
|
||||
reply_mode: ReplyMode = ReplyMode.DISPATCH
|
||||
allow_message_tools: bool = True
|
||||
output_callback: Optional[Callable[[str], None]] = None
|
||||
notification_callback: Optional[Callable[[Any], None]] = None
|
||||
agent_factory: Optional[Callable[..., MoviePilotAgent]] = None
|
||||
completion_future: Optional[asyncio.Future] = None
|
||||
|
||||
|
||||
class AgentManager:
|
||||
@@ -1617,6 +1797,12 @@ class AgentManager:
|
||||
queue = self._session_queues.get(session_id)
|
||||
return bool(queue and not queue.empty())
|
||||
|
||||
def is_session_busy(self, session_id: str) -> bool:
|
||||
"""
|
||||
查询会话是否仍有正在执行或排队的任务。
|
||||
"""
|
||||
return self._is_session_busy(session_id)
|
||||
|
||||
def _expired_idle_sessions(self) -> list[tuple[str, str]]:
|
||||
"""
|
||||
收集已经超过空闲时间且当前不忙的会话。
|
||||
@@ -1658,11 +1844,18 @@ class AgentManager:
|
||||
original_chat_id: Optional[str] = None,
|
||||
reply_mode: ReplyMode = ReplyMode.DISPATCH,
|
||||
allow_message_tools: bool = True,
|
||||
output_callback: Optional[Callable[[str], None]] = None,
|
||||
notification_callback: Optional[Callable[[Any], None]] = None,
|
||||
agent_factory: Optional[Callable[..., MoviePilotAgent]] = None,
|
||||
wait_for_completion: bool = False,
|
||||
) -> str:
|
||||
"""
|
||||
处理用户消息:将消息放入会话队列,按顺序依次处理。
|
||||
同一会话的消息排队等待,不同会话之间互不影响。
|
||||
"""
|
||||
completion_future = (
|
||||
asyncio.get_running_loop().create_future() if wait_for_completion else None
|
||||
)
|
||||
task = _MessageTask(
|
||||
session_id=session_id,
|
||||
user_id=user_id,
|
||||
@@ -1677,6 +1870,10 @@ class AgentManager:
|
||||
original_chat_id=original_chat_id,
|
||||
reply_mode=reply_mode,
|
||||
allow_message_tools=allow_message_tools,
|
||||
output_callback=output_callback,
|
||||
notification_callback=notification_callback,
|
||||
agent_factory=agent_factory,
|
||||
completion_future=completion_future,
|
||||
)
|
||||
self._record_session_activity(session_id, user_id)
|
||||
|
||||
@@ -1709,6 +1906,8 @@ class AgentManager:
|
||||
self._session_worker(session_id)
|
||||
)
|
||||
|
||||
if completion_future:
|
||||
return await completion_future
|
||||
return ""
|
||||
|
||||
async def _session_worker(self, session_id: str):
|
||||
@@ -1732,9 +1931,17 @@ class AgentManager:
|
||||
|
||||
try:
|
||||
await self._start_task_processing_status(task)
|
||||
await self._process_message_internal(task)
|
||||
result = await self._process_message_internal(task)
|
||||
if task.completion_future and not task.completion_future.done():
|
||||
task.completion_future.set_result(result)
|
||||
except asyncio.CancelledError as err:
|
||||
if task.completion_future and not task.completion_future.done():
|
||||
task.completion_future.set_exception(err)
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"处理会话 {session_id} 的消息失败: {e}")
|
||||
if task.completion_future and not task.completion_future.done():
|
||||
task.completion_future.set_exception(e)
|
||||
finally:
|
||||
await self._finish_task_processing_status(task)
|
||||
queue.task_done()
|
||||
@@ -1773,35 +1980,55 @@ class AgentManager:
|
||||
实际处理单条消息
|
||||
"""
|
||||
session_id = task.session_id
|
||||
existing_agent = self.active_agents.get(session_id)
|
||||
if (
|
||||
existing_agent
|
||||
and task.agent_factory
|
||||
and isinstance(task.agent_factory, type)
|
||||
and not isinstance(existing_agent, task.agent_factory)
|
||||
):
|
||||
await existing_agent.cleanup()
|
||||
self.active_agents.pop(session_id, None)
|
||||
|
||||
if session_id not in self.active_agents:
|
||||
logger.info(
|
||||
f"创建新的AI智能体实例,session_id: {session_id}, user_id: {task.user_id}"
|
||||
)
|
||||
agent = MoviePilotAgent(
|
||||
session_id=session_id,
|
||||
user_id=task.user_id,
|
||||
channel=task.channel,
|
||||
source=task.source,
|
||||
username=task.username,
|
||||
original_message_id=task.original_message_id,
|
||||
original_chat_id=task.original_chat_id,
|
||||
replay_mode=task.reply_mode,
|
||||
allow_message_tools=task.allow_message_tools,
|
||||
)
|
||||
agent_factory = task.agent_factory or MoviePilotAgent
|
||||
agent_kwargs = {
|
||||
"session_id": session_id,
|
||||
"user_id": task.user_id,
|
||||
"channel": task.channel,
|
||||
"source": task.source,
|
||||
"username": task.username,
|
||||
"original_message_id": task.original_message_id,
|
||||
"original_chat_id": task.original_chat_id,
|
||||
"replay_mode": task.reply_mode,
|
||||
"allow_message_tools": task.allow_message_tools,
|
||||
"output_callback": task.output_callback,
|
||||
}
|
||||
if task.notification_callback is not None and task.agent_factory:
|
||||
agent_kwargs["notification_callback"] = task.notification_callback
|
||||
agent = agent_factory(**agent_kwargs)
|
||||
self.active_agents[session_id] = agent
|
||||
else:
|
||||
agent = self.active_agents[session_id]
|
||||
agent.user_id = task.user_id
|
||||
if task.channel:
|
||||
agent.channel = task.channel
|
||||
if task.source:
|
||||
agent.source = task.source
|
||||
if task.username:
|
||||
agent.username = task.username
|
||||
# 每条队列任务都携带完整消息上下文,None 也必须覆盖,避免后台任务
|
||||
# 复用会话 Agent 时继续沿用上一条入站消息的渠道。
|
||||
agent.channel = task.channel
|
||||
agent.source = task.source
|
||||
agent.username = task.username
|
||||
agent.original_message_id = task.original_message_id
|
||||
agent.original_chat_id = task.original_chat_id
|
||||
agent.reply_mode = task.reply_mode
|
||||
agent.allow_message_tools = task.allow_message_tools
|
||||
if hasattr(agent, "set_output_callback"):
|
||||
agent.set_output_callback(task.output_callback)
|
||||
else:
|
||||
agent.output_callback = task.output_callback
|
||||
if task.notification_callback is not None and hasattr(agent, "set_notification_callback"):
|
||||
agent.set_notification_callback(task.notification_callback)
|
||||
|
||||
process_kwargs = {
|
||||
"images": task.images,
|
||||
@@ -1909,6 +2136,78 @@ class AgentManager:
|
||||
await agent.cleanup()
|
||||
memory_manager.clear_memory(session_id, user_id)
|
||||
|
||||
async def execute_scheduled_task(self, task_id: int) -> tuple[bool, str]:
|
||||
"""
|
||||
按持久化上下文唤醒 Agent 执行自主定时任务并向用户回传结果。
|
||||
|
||||
:param task_id: Agent 定时任务 ID
|
||||
:return: 执行是否成功及结果摘要
|
||||
"""
|
||||
if not settings.AI_AGENT_ENABLE:
|
||||
return False, "AI Agent 未启用"
|
||||
oper = AgentTaskOper()
|
||||
task = oper.get(task_id)
|
||||
if not task or not task.enabled:
|
||||
return False, "Agent 定时任务不存在或已停用"
|
||||
if not oper.mark_running(task_id):
|
||||
return False, "Agent 定时任务当前不可执行"
|
||||
|
||||
task_message = (
|
||||
f"定时任务已按计划触发。请立即完成下面的任务,不要只确认收到,"
|
||||
f"也不要重复创建同一个定时任务。\n\n"
|
||||
f"任务名称:{task.name}\n"
|
||||
f"任务内容:{task.content}\n\n"
|
||||
"完成后请直接向用户发送消息报告本次执行结果;如果无法完成,也需发送消息说明原因。"
|
||||
)
|
||||
success = True
|
||||
result = ""
|
||||
notification_username = task.username or settings.SUPERUSER
|
||||
try:
|
||||
result = await self.process_message(
|
||||
session_id=task.session_id,
|
||||
user_id=task.user_id,
|
||||
message=task_message,
|
||||
channel=None,
|
||||
source=None,
|
||||
username=notification_username,
|
||||
original_chat_id=None,
|
||||
reply_mode=ReplyMode.DISPATCH,
|
||||
allow_message_tools=True,
|
||||
wait_for_completion=True,
|
||||
)
|
||||
result_text = str(result or "").strip()
|
||||
success = not result_text.startswith(
|
||||
(AGENT_EXECUTION_ERROR_PREFIX, "处理消息时发生错误")
|
||||
)
|
||||
except Exception as err:
|
||||
success = False
|
||||
result = f"Agent 定时任务执行失败:{str(err)}"
|
||||
logger.error(f"Agent 定时任务 {task_id} 执行失败: {str(err)}")
|
||||
await AgentChain().async_post_message(
|
||||
Notification(
|
||||
mtype=NotificationType.Agent,
|
||||
username=notification_username,
|
||||
title=f"定时任务执行失败:{task.name}",
|
||||
text=result,
|
||||
save_history=False,
|
||||
)
|
||||
)
|
||||
finally:
|
||||
current_task = oper.get(task_id)
|
||||
oper.finish(
|
||||
task_id=task_id,
|
||||
success=success,
|
||||
result=str(result or ""),
|
||||
disable=bool(
|
||||
current_task
|
||||
and task.trigger_type == "date"
|
||||
and current_task.trigger_type == task.trigger_type
|
||||
and current_task.run_at == task.run_at
|
||||
),
|
||||
)
|
||||
|
||||
return success, str(result or "任务执行完成")
|
||||
|
||||
@staticmethod
|
||||
def _build_heartbeat_prompt() -> str:
|
||||
"""使用程序内置 System Tasks 定义构建心跳任务提示词。"""
|
||||
|
||||
@@ -536,6 +536,7 @@ class StreamingHandler:
|
||||
original_chat_id=self._original_chat_id,
|
||||
title=self._title,
|
||||
text=current_text,
|
||||
save_history=False,
|
||||
),
|
||||
)
|
||||
if response and response.success and response.message_id:
|
||||
@@ -581,6 +582,7 @@ class StreamingHandler:
|
||||
original_chat_id=self._original_chat_id,
|
||||
title=self._title,
|
||||
text=current_text,
|
||||
save_history=False,
|
||||
),
|
||||
)
|
||||
if response and response.success and response.message_id:
|
||||
|
||||
@@ -691,7 +691,9 @@ class AgentCapabilityManager:
|
||||
@staticmethod
|
||||
def supports_image_input() -> bool:
|
||||
"""当前 Agent 是否启用图片输入能力。"""
|
||||
return bool(settings.LLM_SUPPORT_IMAGE_INPUT)
|
||||
from app.agent.llm.helper import LLMHelper
|
||||
|
||||
return LLMHelper.supports_image_input()
|
||||
|
||||
@staticmethod
|
||||
def supports_audio_input() -> bool:
|
||||
|
||||
+205
-25
@@ -5,7 +5,7 @@ import inspect
|
||||
import json
|
||||
import time
|
||||
from functools import wraps
|
||||
from typing import Any, List
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk
|
||||
|
||||
@@ -700,11 +700,85 @@ class LLMHelper:
|
||||
return {}
|
||||
|
||||
@staticmethod
|
||||
def supports_image_input() -> bool:
|
||||
def _metadata_supports_image_input(metadata: Any) -> Optional[bool]:
|
||||
"""从模型元数据中读取图片输入能力,未知时返回 None。"""
|
||||
if not isinstance(metadata, dict):
|
||||
return None
|
||||
|
||||
modalities = metadata.get("modalities") or {}
|
||||
input_modalities = modalities.get("input")
|
||||
if isinstance(input_modalities, str):
|
||||
input_modalities = [input_modalities]
|
||||
if isinstance(input_modalities, list):
|
||||
normalized_modalities = {
|
||||
str(item or "").strip().lower() for item in input_modalities
|
||||
}
|
||||
return "image" in normalized_modalities
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def _resolve_catalog_image_input_support(
|
||||
cls,
|
||||
provider: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
base_url: Optional[str] = None,
|
||||
base_url_preset: Optional[str] = None,
|
||||
) -> Optional[bool]:
|
||||
"""复用 provider 目录缓存解析当前模型是否支持图片输入。"""
|
||||
provider_name = str(provider if provider is not None else settings.LLM_PROVIDER).strip()
|
||||
model_name = str(model if model is not None else settings.LLM_MODEL).strip()
|
||||
if not provider_name or not model_name:
|
||||
return None
|
||||
|
||||
try:
|
||||
from app.agent.llm.provider import LLMProviderManager
|
||||
|
||||
metadata = LLMProviderManager().resolve_cached_model_metadata(
|
||||
provider_id=provider_name,
|
||||
model_id=model_name,
|
||||
base_url=base_url if base_url is not None else settings.LLM_BASE_URL,
|
||||
base_url_preset_id=(
|
||||
base_url_preset
|
||||
if base_url_preset is not None
|
||||
else settings.LLM_BASE_URL_PRESET
|
||||
),
|
||||
)
|
||||
except Exception as err:
|
||||
logger.debug(f"解析模型图片能力失败: {err}")
|
||||
return None
|
||||
|
||||
return cls._metadata_supports_image_input(metadata)
|
||||
|
||||
@classmethod
|
||||
def supports_image_input(
|
||||
cls,
|
||||
provider: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
base_url: Optional[str] = None,
|
||||
base_url_preset: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
判断当前模型是否启用了图片输入能力。
|
||||
|
||||
用户开关为总开关;当内置模型目录明确标注当前模型不支持 image 输入时,
|
||||
即使总开关开启也降级为纯文本,避免文本模型收到 `image_url` 内容块后
|
||||
被兼容端点以 400 拒绝。无参调用保持旧版“只读总开关”语义,
|
||||
未知自定义模型也保持原有开关语义。
|
||||
"""
|
||||
return bool(settings.LLM_SUPPORT_IMAGE_INPUT)
|
||||
if not settings.LLM_SUPPORT_IMAGE_INPUT:
|
||||
return False
|
||||
if provider is None and model is None:
|
||||
return True
|
||||
|
||||
image_support = cls._resolve_catalog_image_input_support(
|
||||
provider=provider,
|
||||
model=model,
|
||||
base_url=base_url,
|
||||
base_url_preset=base_url_preset,
|
||||
)
|
||||
if image_support is not None:
|
||||
return image_support
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _build_legacy_runtime(
|
||||
@@ -772,19 +846,31 @@ class LLMHelper:
|
||||
provider: str,
|
||||
model: str | None,
|
||||
runtime: dict[str, Any],
|
||||
api_protocol: str | None = None,
|
||||
) -> bool | None:
|
||||
"""
|
||||
判断官方 ChatGPT API Key 模式是否应使用 Responses API。
|
||||
判断本次 OpenAI 兼容调用是否应使用 Responses API。
|
||||
|
||||
GPT-5/o 系推理模型在 Chat Completions 中组合 function tools 与
|
||||
reasoning_effort 时会被官方端点拒绝,因此 ChatGPT 官方 API Key
|
||||
模式需要显式切到 Responses API;通用 OpenAI-compatible 入口保持
|
||||
provider 目录解析出的默认行为,避免误伤第三方兼容服务。
|
||||
优先级:
|
||||
1. 运行时显式要求(ChatGPT Plus/Pro OAuth、Codex 等端点契约),始终保留;
|
||||
2. 用户通过 ``LLM_API_PROTOCOL`` 显式指定 ``responses`` / ``chat_completions``;
|
||||
3. ``auto``(默认)保持原有 ChatGPT 官方 API Key + GPT-5/o 系推理模型
|
||||
自动切换逻辑,通用 OpenAI 兼容入口仍走 Chat Completions,
|
||||
避免误伤第三方兼容服务。
|
||||
|
||||
:param api_protocol: 显式传入的 API 协议,未传入时读取 ``LLM_API_PROTOCOL``
|
||||
:return: True/False 强制指定协议;None 表示交由 LangChain 默认行为
|
||||
"""
|
||||
runtime_use_responses_api = runtime.get("use_responses_api")
|
||||
if runtime_use_responses_api is not None:
|
||||
return bool(runtime_use_responses_api)
|
||||
|
||||
protocol = cls._normalize_api_protocol(api_protocol)
|
||||
if protocol == "responses":
|
||||
return True
|
||||
if protocol == "chat_completions":
|
||||
return False
|
||||
|
||||
provider_name = (provider or "").strip().lower()
|
||||
if provider_name != "chatgpt":
|
||||
return None
|
||||
@@ -798,6 +884,53 @@ class LLMHelper:
|
||||
return True
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _normalize_api_protocol(api_protocol: str | None) -> str:
|
||||
"""
|
||||
规范化 API 协议配置,未知值统一回退为 ``auto`` 以保持兼容。
|
||||
"""
|
||||
normalized = str(api_protocol or settings.LLM_API_PROTOCOL or "").strip().lower()
|
||||
if normalized in {"auto", "chat_completions", "responses"}:
|
||||
return normalized
|
||||
if normalized:
|
||||
logger.warning(f"忽略不支持的 LLM_API_PROTOCOL 配置: {api_protocol}")
|
||||
return "auto"
|
||||
|
||||
@staticmethod
|
||||
def _attach_runtime_metadata(model: Any, runtime: dict[str, Any]) -> None:
|
||||
"""
|
||||
将 MoviePilot 已解析出的 provider 运行时信息挂到模型实例上。
|
||||
|
||||
这些字段只供内部中间件识别协议能力,不参与 LangChain 请求序列化。
|
||||
"""
|
||||
runtime_metadata = {
|
||||
"runtime": runtime.get("runtime"),
|
||||
"provider_id": runtime.get("provider_id"),
|
||||
"base_url": runtime.get("base_url"),
|
||||
}
|
||||
|
||||
def _set_metadata_attr(name: str, value: Any) -> None:
|
||||
try:
|
||||
setattr(model, name, value)
|
||||
except Exception:
|
||||
object.__setattr__(model, name, value)
|
||||
|
||||
try:
|
||||
_set_metadata_attr("_moviepilot_llm_runtime", runtime_metadata["runtime"])
|
||||
_set_metadata_attr(
|
||||
"_moviepilot_llm_provider_id",
|
||||
runtime_metadata["provider_id"],
|
||||
)
|
||||
_set_metadata_attr("_moviepilot_llm_base_url", runtime_metadata["base_url"])
|
||||
except Exception as err:
|
||||
logger.debug(f"LLM运行时元数据附加失败: {str(err)}")
|
||||
|
||||
profile = getattr(model, "profile", None)
|
||||
if isinstance(profile, dict):
|
||||
profile["moviepilot_runtime"] = runtime_metadata["runtime"]
|
||||
profile["moviepilot_provider_id"] = runtime_metadata["provider_id"]
|
||||
profile["moviepilot_base_url"] = runtime_metadata["base_url"]
|
||||
|
||||
@classmethod
|
||||
def _resolve_thinking_level(
|
||||
cls,
|
||||
@@ -843,7 +976,9 @@ class LLMHelper:
|
||||
base_url: str | None = None,
|
||||
base_url_preset: str | None = None,
|
||||
user_agent: str | None = None,
|
||||
temperature: Optional[float] = None,
|
||||
use_proxy: bool | None = None,
|
||||
api_protocol: str | None = None,
|
||||
):
|
||||
"""
|
||||
获取LLM实例
|
||||
@@ -858,7 +993,12 @@ class LLMHelper:
|
||||
:param base_url: API Base URL。未显式传入时使用当前配置项 LLM_BASE_URL。
|
||||
:param base_url_preset: Base URL 预设。未显式传入时使用当前配置项 LLM_BASE_URL_PRESET。
|
||||
:param user_agent: OpenAI兼容接口请求 User-Agent。未显式传入时使用配置项 LLM_USER_AGENT。
|
||||
:param temperature: LLM 温度参数。未显式传入时使用配置项 LLM_TEMPERATURE。
|
||||
:param use_proxy: 是否为本次 LLM 调用使用系统代理。未显式传入时使用配置项 LLM_USE_PROXY。
|
||||
:param api_protocol: OpenAI 兼容接口 API 协议
|
||||
(auto/chat_completions/responses)。未显式传入时使用配置项 LLM_API_PROTOCOL。
|
||||
仅对 OpenAI 兼容运行时生效;``responses`` 强制走 Responses API,
|
||||
``chat_completions`` 强制走 Chat Completions,``auto`` 保持原有自动判断。
|
||||
:return: LLM实例
|
||||
"""
|
||||
provider_name = str(provider if provider is not None else settings.LLM_PROVIDER).lower()
|
||||
@@ -869,6 +1009,7 @@ class LLMHelper:
|
||||
base_url_preset if base_url_preset is not None else settings.LLM_BASE_URL_PRESET
|
||||
)
|
||||
user_agent_value = user_agent if user_agent is not None else settings.LLM_USER_AGENT
|
||||
temperature_value = temperature if temperature is not None else settings.LLM_TEMPERATURE
|
||||
normalized_thinking_level = cls._resolve_thinking_level(
|
||||
thinking_level=thinking_level,
|
||||
)
|
||||
@@ -909,6 +1050,7 @@ class LLMHelper:
|
||||
provider=provider_name,
|
||||
model=model_name,
|
||||
runtime=runtime,
|
||||
api_protocol=api_protocol,
|
||||
)
|
||||
llm_proxy = _resolve_llm_proxy(use_proxy)
|
||||
|
||||
@@ -925,7 +1067,7 @@ class LLMHelper:
|
||||
model=model_name,
|
||||
api_key=runtime["api_key"],
|
||||
retries=3,
|
||||
temperature=settings.LLM_TEMPERATURE,
|
||||
temperature=temperature_value,
|
||||
streaming=streaming,
|
||||
client_args=_build_google_client_args(llm_proxy),
|
||||
**thinking_kwargs,
|
||||
@@ -939,13 +1081,36 @@ class LLMHelper:
|
||||
api_key=runtime["api_key"],
|
||||
api_base=runtime["base_url"],
|
||||
max_retries=3,
|
||||
temperature=settings.LLM_TEMPERATURE,
|
||||
temperature=temperature_value,
|
||||
streaming=streaming,
|
||||
stream_usage=True,
|
||||
http_client=_build_httpx_client(llm_proxy),
|
||||
http_async_client=_build_httpx_client(llm_proxy, async_client=True),
|
||||
**thinking_kwargs,
|
||||
)
|
||||
elif runtime["runtime"] == "bedrock":
|
||||
from langchain_aws import ChatBedrockConverse
|
||||
|
||||
from app.agent.llm.provider import LLMProviderManager
|
||||
|
||||
aws_region = runtime.get("aws_region") or "us-east-1"
|
||||
aws_auth = runtime.get("aws_auth") or {}
|
||||
# Bearer 认证需要跳过 SigV4 签名并注入 Authorization 头,SigV4 认证
|
||||
# 直接以 AK/SK 签名;两种方式统一由 provider 管理器构造 boto3 客户端。
|
||||
bedrock_client = LLMProviderManager().create_bedrock_client(
|
||||
"bedrock-runtime",
|
||||
region=aws_region,
|
||||
credentials=aws_auth,
|
||||
base_url=runtime.get("base_url"),
|
||||
use_proxy=use_proxy,
|
||||
read_timeout=settings.LLM_TOOL_TIMEOUT,
|
||||
)
|
||||
model = ChatBedrockConverse(
|
||||
model_id=model_name,
|
||||
client=bedrock_client,
|
||||
temperature=temperature_value,
|
||||
disable_streaming=not streaming,
|
||||
)
|
||||
elif runtime["runtime"] in {"anthropic_compatible", "copilot_anthropic"}:
|
||||
from langchain_anthropic import ChatAnthropic
|
||||
|
||||
@@ -954,7 +1119,7 @@ class LLMHelper:
|
||||
api_key=runtime["api_key"],
|
||||
base_url=runtime["base_url"],
|
||||
max_retries=3,
|
||||
temperature=settings.LLM_TEMPERATURE,
|
||||
temperature=temperature_value,
|
||||
streaming=streaming,
|
||||
stream_usage=True,
|
||||
anthropic_proxy=llm_proxy,
|
||||
@@ -975,7 +1140,7 @@ class LLMHelper:
|
||||
api_key=runtime["api_key"],
|
||||
max_retries=3,
|
||||
base_url=runtime.get("base_url"),
|
||||
temperature=settings.LLM_TEMPERATURE,
|
||||
temperature=temperature_value,
|
||||
streaming=streaming,
|
||||
stream_usage=True,
|
||||
openai_proxy=llm_proxy,
|
||||
@@ -995,7 +1160,11 @@ class LLMHelper:
|
||||
# 优先使用 provider / models.dev 目录中的上下文上限,减少用户手填成本。
|
||||
model_profile = getattr(model, "profile", None)
|
||||
if model_profile:
|
||||
logger.debug(f"使用LLM模型: {model.model},Profile: {model.profile}")
|
||||
# ChatBedrockConverse 等模型类没有 model 属性,模型名存放在 model_id。
|
||||
logged_model_name = getattr(model, "model", None) or getattr(
|
||||
model, "model_id", model_name
|
||||
)
|
||||
logger.debug(f"使用LLM模型: {logged_model_name},Profile: {model_profile}")
|
||||
else:
|
||||
model_record = runtime.get("model_record") or {}
|
||||
model_metadata = runtime.get("model_metadata") or {}
|
||||
@@ -1011,6 +1180,7 @@ class LLMHelper:
|
||||
"max_input_tokens": int(max_input_tokens),
|
||||
}
|
||||
|
||||
cls._attach_runtime_metadata(model, runtime)
|
||||
return model
|
||||
|
||||
@staticmethod
|
||||
@@ -1068,25 +1238,35 @@ class LLMHelper:
|
||||
base_url: str | None = None,
|
||||
base_url_preset: str | None = None,
|
||||
user_agent: str | None = None,
|
||||
temperature: Optional[float] = None,
|
||||
use_proxy: bool | None = None,
|
||||
api_protocol: str | None = None,
|
||||
) -> dict:
|
||||
"""
|
||||
使用当前已保存配置执行一次最小 LLM 调用。
|
||||
使用当前配置或显式传入的临时配置执行一次最小 LLM 调用。
|
||||
|
||||
:param temperature: LLM 温度参数。未显式传入时沿用已保存配置。
|
||||
:param api_protocol: OpenAI 兼容接口 API 协议,未显式传入时沿用已保存配置。
|
||||
"""
|
||||
provider_name = provider if provider is not None else settings.LLM_PROVIDER
|
||||
model_name = model if model is not None else settings.LLM_MODEL
|
||||
start = time.perf_counter()
|
||||
llm = await LLMHelper.get_llm(
|
||||
streaming=False,
|
||||
provider=provider_name,
|
||||
model=model_name,
|
||||
thinking_level=thinking_level,
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
base_url_preset=base_url_preset,
|
||||
user_agent=user_agent,
|
||||
use_proxy=use_proxy,
|
||||
)
|
||||
llm_kwargs = {
|
||||
"streaming": False,
|
||||
"provider": provider_name,
|
||||
"model": model_name,
|
||||
"thinking_level": thinking_level,
|
||||
"api_key": api_key,
|
||||
"base_url": base_url,
|
||||
"base_url_preset": base_url_preset,
|
||||
"user_agent": user_agent,
|
||||
"use_proxy": use_proxy,
|
||||
"api_protocol": api_protocol,
|
||||
}
|
||||
if temperature is not None:
|
||||
llm_kwargs["temperature"] = temperature
|
||||
|
||||
llm = await LLMHelper.get_llm(**llm_kwargs)
|
||||
try:
|
||||
response = await asyncio.wait_for(llm.ainvoke(prompt), timeout=timeout)
|
||||
except TimeoutError as err:
|
||||
|
||||
+602
-32
@@ -7,13 +7,14 @@ import base64
|
||||
import copy
|
||||
import hashlib
|
||||
import json
|
||||
import re
|
||||
import secrets
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
from urllib.parse import urlencode
|
||||
from urllib.parse import urlencode, urlsplit
|
||||
|
||||
import aiofiles
|
||||
import httpx
|
||||
@@ -106,6 +107,90 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
_MODELS_DEV_BUNDLED_PATH = Path(__file__).with_name("models.json")
|
||||
_MODELS_DEV_CACHE_TTL = 7 * 24 * 60 * 60
|
||||
_AUTH_SESSION_DONE_RETENTION = 300
|
||||
_BEDROCK_DEFAULT_REGION = "us-east-1"
|
||||
_BEDROCK_API_KEY_PREFIX = "bedrock-api-key-"
|
||||
_BEDROCK_GPT_OSS_BASE_REGIONS = (
|
||||
"ap-northeast-1",
|
||||
"ap-south-1",
|
||||
"ap-southeast-2",
|
||||
"eu-central-1",
|
||||
"eu-north-1",
|
||||
"eu-west-1",
|
||||
"eu-west-2",
|
||||
"sa-east-1",
|
||||
"us-east-1",
|
||||
"us-east-2",
|
||||
"us-west-2",
|
||||
)
|
||||
_BEDROCK_GPT_OSS_SAFEGUARD_REGIONS = (
|
||||
"ap-northeast-1",
|
||||
"ap-south-1",
|
||||
"ap-southeast-2",
|
||||
"eu-west-1",
|
||||
"eu-west-2",
|
||||
"sa-east-1",
|
||||
"us-east-1",
|
||||
"us-east-2",
|
||||
"us-west-2",
|
||||
)
|
||||
_BEDROCK_ON_DEMAND_MODEL_REGIONS = {
|
||||
"openai.gpt-oss-120b-1:0": _BEDROCK_GPT_OSS_BASE_REGIONS,
|
||||
"openai.gpt-oss-20b-1:0": _BEDROCK_GPT_OSS_BASE_REGIONS,
|
||||
"openai.gpt-oss-safeguard-120b": _BEDROCK_GPT_OSS_SAFEGUARD_REGIONS,
|
||||
"openai.gpt-oss-safeguard-20b": _BEDROCK_GPT_OSS_SAFEGUARD_REGIONS,
|
||||
"amazon.nova-lite-v1:0": (
|
||||
"ap-northeast-1",
|
||||
"ap-southeast-2",
|
||||
"eu-west-2",
|
||||
"us-east-1",
|
||||
"us-gov-west-1",
|
||||
),
|
||||
"amazon.nova-micro-v1:0": (
|
||||
"ap-southeast-2",
|
||||
"eu-west-2",
|
||||
"us-east-1",
|
||||
"us-gov-west-1",
|
||||
),
|
||||
"amazon.nova-pro-v1:0": (
|
||||
"ap-southeast-2",
|
||||
"eu-west-2",
|
||||
"us-east-1",
|
||||
"us-gov-west-1",
|
||||
),
|
||||
"anthropic.claude-3-5-haiku-20241022-v1:0": (
|
||||
"us-west-2",
|
||||
),
|
||||
"anthropic.claude-3-5-sonnet-20240620-v1:0": (
|
||||
"ap-northeast-1",
|
||||
"ap-northeast-2",
|
||||
"ap-southeast-1",
|
||||
"eu-central-1",
|
||||
"eu-central-2",
|
||||
"us-east-1",
|
||||
"us-gov-west-1",
|
||||
"us-west-2",
|
||||
),
|
||||
"anthropic.claude-3-5-sonnet-20241022-v2:0": (
|
||||
"ap-southeast-2",
|
||||
"us-west-2",
|
||||
),
|
||||
"anthropic.claude-3-7-sonnet-20250219-v1:0": (
|
||||
"eu-west-2",
|
||||
"us-gov-west-1",
|
||||
),
|
||||
"anthropic.claude-3-haiku-20240307-v1:0": (
|
||||
"ap-northeast-1",
|
||||
"ap-northeast-2",
|
||||
"ap-south-1",
|
||||
"ap-southeast-2",
|
||||
"eu-central-1",
|
||||
"eu-west-1",
|
||||
"eu-west-3",
|
||||
"us-east-1",
|
||||
"us-gov-west-1",
|
||||
"us-west-2",
|
||||
),
|
||||
}
|
||||
_CHATGPT_CLIENT_ID = "app_EMoamEEZ73f0CkXaXp7hrann"
|
||||
_CHATGPT_ISSUER = "https://auth.openai.com"
|
||||
_CHATGPT_CODEX_BASE_URL = "https://chatgpt.com/backend-api/codex"
|
||||
@@ -367,6 +452,50 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
api_key_hint="填写 Anthropic API Key。",
|
||||
description="Anthropic Claude 官方端点。",
|
||||
),
|
||||
ProviderSpec(
|
||||
id="amazon-bedrock",
|
||||
name="Amazon Bedrock",
|
||||
runtime="bedrock",
|
||||
models_dev_provider_id="amazon-bedrock",
|
||||
default_base_url="https://bedrock-runtime.us-east-1.amazonaws.com",
|
||||
base_url_presets=(
|
||||
url_preset(
|
||||
id="bedrock-us-east-1",
|
||||
label="美东(弗吉尼亚北部)us-east-1",
|
||||
value="https://bedrock-runtime.us-east-1.amazonaws.com",
|
||||
),
|
||||
url_preset(
|
||||
id="bedrock-us-west-2",
|
||||
label="美西(俄勒冈)us-west-2",
|
||||
value="https://bedrock-runtime.us-west-2.amazonaws.com",
|
||||
),
|
||||
url_preset(
|
||||
id="bedrock-eu-central-1",
|
||||
label="欧洲(法兰克福)eu-central-1",
|
||||
value="https://bedrock-runtime.eu-central-1.amazonaws.com",
|
||||
),
|
||||
url_preset(
|
||||
id="bedrock-ap-northeast-1",
|
||||
label="亚太(东京)ap-northeast-1",
|
||||
value="https://bedrock-runtime.ap-northeast-1.amazonaws.com",
|
||||
),
|
||||
url_preset(
|
||||
id="bedrock-ap-southeast-1",
|
||||
label="亚太(新加坡)ap-southeast-1",
|
||||
value="https://bedrock-runtime.ap-southeast-1.amazonaws.com",
|
||||
),
|
||||
),
|
||||
base_url_editable=True,
|
||||
api_key_label="Bedrock API Key / AK:SK",
|
||||
api_key_hint=(
|
||||
"支持两种认证方式:填写 Amazon Bedrock API Key(bedrock-api-key- 开头,"
|
||||
"Bearer 认证);或填写 Access Key ID:Secret Access Key(可选追加 :Session Token,"
|
||||
"SigV4 认证)。Base URL 决定 AWS Region。"
|
||||
),
|
||||
model_list_strategy="bedrock",
|
||||
description="Amazon Bedrock 托管模型服务,支持 Bedrock API Key 与 AK/SK 双认证。",
|
||||
sort_order=35,
|
||||
),
|
||||
ProviderSpec(
|
||||
id="deepseek",
|
||||
name="DeepSeek",
|
||||
@@ -1564,6 +1693,101 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
return models[candidate]
|
||||
return None
|
||||
|
||||
def _cached_models_dev_model(
|
||||
self,
|
||||
provider_id: str,
|
||||
model_id: str,
|
||||
base_url: Optional[str] = None,
|
||||
base_url_preset_id: Optional[str] = None,
|
||||
) -> dict[str, Any] | None:
|
||||
"""从已缓存或内置的 models.dev 数据中同步读取模型元数据。"""
|
||||
try:
|
||||
spec = self.get_provider(provider_id)
|
||||
except LLMProviderError:
|
||||
return None
|
||||
|
||||
models_dev_provider_id = self._resolve_provider_models_dev_provider_id(
|
||||
spec,
|
||||
base_url,
|
||||
base_url_preset_id=base_url_preset_id,
|
||||
)
|
||||
if not models_dev_provider_id:
|
||||
return None
|
||||
|
||||
payload = self._cached_models_dev_payload().get(models_dev_provider_id, {}) or {}
|
||||
models = payload.get("models") if isinstance(payload, dict) else None
|
||||
if not isinstance(models, dict):
|
||||
return None
|
||||
|
||||
candidates = [model_id]
|
||||
if model_id.startswith("models/"):
|
||||
candidates.append(model_id.removeprefix("models/"))
|
||||
|
||||
for candidate in candidates:
|
||||
if candidate in models:
|
||||
return models[candidate]
|
||||
return None
|
||||
|
||||
def resolve_cached_model_metadata(
|
||||
self,
|
||||
provider_id: str,
|
||||
model_id: Optional[str],
|
||||
base_url: Optional[str] = None,
|
||||
base_url_preset_id: Optional[str] = None,
|
||||
) -> dict[str, Any] | None:
|
||||
"""同步解析缓存中的模型元数据,不触发远端 models.dev 刷新。"""
|
||||
if not model_id:
|
||||
return None
|
||||
metadata = self._cached_models_dev_model(
|
||||
provider_id,
|
||||
model_id,
|
||||
base_url=base_url,
|
||||
base_url_preset_id=base_url_preset_id,
|
||||
)
|
||||
if metadata:
|
||||
return metadata
|
||||
if provider_id == "chatgpt":
|
||||
return self._cached_models_dev_model("openai", model_id)
|
||||
if provider_id == "openai":
|
||||
return (
|
||||
self._cached_models_dev_payload()
|
||||
.get("openai", {})
|
||||
.get("models", {})
|
||||
.get(model_id)
|
||||
)
|
||||
return None
|
||||
|
||||
def _resolve_cached_model_record(
|
||||
self,
|
||||
provider_id: str,
|
||||
model_id: Optional[str],
|
||||
base_url: Optional[str] = None,
|
||||
base_url_preset_id: Optional[str] = None,
|
||||
transport: str = "openai",
|
||||
) -> dict[str, Any] | None:
|
||||
"""从缓存中的模型元数据构造轻量模型记录,不触发远端模型列表刷新。"""
|
||||
if not model_id:
|
||||
return None
|
||||
metadata = self.resolve_cached_model_metadata(
|
||||
provider_id,
|
||||
model_id,
|
||||
base_url=base_url,
|
||||
base_url_preset_id=base_url_preset_id,
|
||||
) or {}
|
||||
if not metadata:
|
||||
return self._normalize_model_record(
|
||||
model_id=model_id,
|
||||
transport=transport,
|
||||
source="configured",
|
||||
)
|
||||
return self._normalize_model_record(
|
||||
model_id=model_id,
|
||||
display_name=metadata.get("name") or model_id,
|
||||
metadata=metadata,
|
||||
transport=transport,
|
||||
source="models.dev-cache",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_model_record(
|
||||
model_id: str,
|
||||
@@ -1648,6 +1872,112 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
return normalized[:-3]
|
||||
return normalized
|
||||
|
||||
@classmethod
|
||||
def _extract_bedrock_region(cls, base_url: Optional[str]) -> str:
|
||||
"""
|
||||
从 Bedrock 运行时端点 URL 中提取 AWS Region
|
||||
|
||||
兼容标准端点、FIPS 端点与 PrivateLink(VPCE)端点等主机名形态,
|
||||
从中识别 Region 段。
|
||||
|
||||
:param base_url: 形如 https://bedrock-runtime.us-east-1.amazonaws.com 的端点地址
|
||||
:return: 提取到的 Region,无法识别时回退 us-east-1
|
||||
"""
|
||||
hostname = urlsplit((base_url or "").strip().lower()).hostname or ""
|
||||
match = re.search(
|
||||
r"(?:^|\.)(?:bedrock(?:-runtime)?(?:-fips)?)"
|
||||
r"\.([a-z0-9-]+-\d+)(?:\.|$)",
|
||||
hostname,
|
||||
)
|
||||
if match:
|
||||
return match.group(1)
|
||||
return cls._BEDROCK_DEFAULT_REGION
|
||||
|
||||
# Inference Profile 的地理前缀与可用 Region 的对应关系,用于降级目录按
|
||||
# 当前 Region 过滤掉不可调用的 Profile 条目。
|
||||
_BEDROCK_GEO_PREFIXES: dict[str, tuple[str, ...]] = {
|
||||
"us": ("us-east-", "us-west-"),
|
||||
"eu": ("eu-",),
|
||||
"apac": ("ap-",),
|
||||
"au": ("ap-southeast-2", "ap-southeast-4"),
|
||||
"jp": ("ap-northeast-1", "ap-northeast-3"),
|
||||
"ca": ("ca-",),
|
||||
}
|
||||
_BEDROCK_NON_COMMERCIAL_REGION_PREFIXES = (
|
||||
"cn-",
|
||||
"eu-isoe-",
|
||||
"us-gov-",
|
||||
"us-iso-",
|
||||
"us-isob-",
|
||||
"us-isof-",
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _bedrock_model_matches_region(cls, model_id: str, region: str) -> bool:
|
||||
"""
|
||||
判断目录中的模型 ID 在指定 Region 是否可调用
|
||||
|
||||
models.dev 目录同时收录裸模型 ID(直连调用)与带地理前缀的
|
||||
Inference Profile ID(us./eu./apac./global. 等)。带前缀的条目只在
|
||||
对应地理分区和 AWS 分区的 Region 可用;global Profile 仅允许商业
|
||||
AWS 分区。裸 ID 仅在明确记录的 ON_DEMAND Region 可用,未知条目
|
||||
按不可直连处理。
|
||||
|
||||
:param model_id: 目录中的模型 ID
|
||||
:param region: 当前 Base URL 对应的 AWS Region
|
||||
:return: 该模型在当前 Region 可调用时返回 True
|
||||
"""
|
||||
prefix = model_id.split(".", 1)[0]
|
||||
if prefix == "global":
|
||||
return not region.startswith(cls._BEDROCK_NON_COMMERCIAL_REGION_PREFIXES)
|
||||
region_prefixes = cls._BEDROCK_GEO_PREFIXES.get(prefix)
|
||||
if region_prefixes is not None:
|
||||
return (
|
||||
not region.startswith(cls._BEDROCK_NON_COMMERCIAL_REGION_PREFIXES)
|
||||
and region.startswith(region_prefixes)
|
||||
)
|
||||
on_demand_regions = cls._BEDROCK_ON_DEMAND_MODEL_REGIONS.get(model_id)
|
||||
return on_demand_regions is not None and region in on_demand_regions
|
||||
|
||||
@classmethod
|
||||
def _parse_bedrock_credentials(cls, api_key: Optional[str]) -> dict[str, Any]:
|
||||
"""
|
||||
解析 Bedrock 凭证字符串,识别 Bearer 与 SigV4 两种认证方式
|
||||
|
||||
- Bedrock API Key(bedrock-api-key- 开头的长期 Key,或控制台生成的短期
|
||||
Token)走 Bearer 认证;
|
||||
- `AccessKeyId:SecretAccessKey` 或 `AccessKeyId:SecretAccessKey:SessionToken`
|
||||
走 SigV4 认证,AWS Access Key ID 均以 "AKIA"/"ASIA" 开头。
|
||||
|
||||
:param api_key: 用户在 API Key 输入框填写的凭证内容
|
||||
:return: 含 auth_scheme 及对应凭证字段的字典
|
||||
"""
|
||||
normalized = str(api_key or "").strip()
|
||||
if not normalized:
|
||||
raise LLMProviderAuthError(
|
||||
"Amazon Bedrock 需要填写 Bedrock API Key 或 Access Key ID:Secret Access Key"
|
||||
)
|
||||
|
||||
if not normalized.startswith(cls._BEDROCK_API_KEY_PREFIX):
|
||||
parts = [part.strip() for part in normalized.split(":")]
|
||||
if len(parts) in {2, 3} and all(parts):
|
||||
credentials = {
|
||||
"auth_scheme": "sigv4",
|
||||
"access_key_id": parts[0],
|
||||
"secret_access_key": parts[1],
|
||||
}
|
||||
if len(parts) == 3:
|
||||
credentials["session_token"] = parts[2]
|
||||
return credentials
|
||||
if ":" in normalized:
|
||||
raise LLMProviderAuthError(
|
||||
"Amazon Bedrock AK/SK 凭证格式不正确,"
|
||||
"请按 AccessKeyId:SecretAccessKey 或 "
|
||||
"AccessKeyId:SecretAccessKey:SessionToken 填写"
|
||||
)
|
||||
|
||||
return {"auth_scheme": "bearer", "bearer_token": normalized}
|
||||
|
||||
async def _list_models_from_google(
|
||||
self,
|
||||
api_key: str,
|
||||
@@ -1762,6 +2092,235 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
)
|
||||
return sorted(results, key=lambda item: item["name"].lower())
|
||||
|
||||
def _build_bedrock_boto3_config(
|
||||
self,
|
||||
use_proxy: Optional[bool] = None,
|
||||
) -> Any:
|
||||
"""
|
||||
构造 Bedrock boto3 客户端配置,统一超时、重试与代理策略
|
||||
|
||||
:param use_proxy: 是否使用系统代理,None 时读取 LLM_USE_PROXY 配置
|
||||
:return: botocore Config 实例
|
||||
"""
|
||||
from botocore.config import Config
|
||||
|
||||
should_use_proxy = settings.LLM_USE_PROXY if use_proxy is None else use_proxy
|
||||
proxies = None
|
||||
if should_use_proxy and settings.PROXY_HOST:
|
||||
proxies = {"http": settings.PROXY_HOST, "https": settings.PROXY_HOST}
|
||||
return Config(
|
||||
connect_timeout=10,
|
||||
read_timeout=60,
|
||||
retries={"max_attempts": 3, "mode": "standard"},
|
||||
proxies=proxies,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _bedrock_endpoint_url(
|
||||
service_name: str, base_url: Optional[str]
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
解析应传给 boto3 客户端的自定义端点 URL
|
||||
|
||||
标准公有端点交由 boto3 按 Region 自行推导;用户填写 PrivateLink、
|
||||
FIPS 等非标准端点时才显式透传,保证所选网络路径实际生效。
|
||||
|
||||
:param service_name: boto3 服务名(bedrock 或 bedrock-runtime)
|
||||
:param base_url: 用户配置的 Base URL
|
||||
:return: 需要显式指定端点时返回 URL,否则返回 None
|
||||
"""
|
||||
normalized = (base_url or "").strip().rstrip("/")
|
||||
if not normalized:
|
||||
return None
|
||||
if re.fullmatch(
|
||||
rf"https://{service_name}\.[a-z0-9-]+\.amazonaws\.com",
|
||||
normalized,
|
||||
):
|
||||
return None
|
||||
return normalized
|
||||
|
||||
def create_bedrock_client(
|
||||
self,
|
||||
service_name: str,
|
||||
region: str,
|
||||
credentials: dict[str, Any],
|
||||
base_url: Optional[str] = None,
|
||||
use_proxy: Optional[bool] = None,
|
||||
read_timeout: Optional[int] = None,
|
||||
) -> Any:
|
||||
"""
|
||||
按解析后的凭证创建 Bedrock boto3 客户端,Bearer 方式注入 Authorization 头
|
||||
|
||||
:param service_name: boto3 服务名(bedrock 或 bedrock-runtime)
|
||||
:param region: AWS Region
|
||||
:param credentials: `_parse_bedrock_credentials` 的解析结果
|
||||
:param base_url: 用户配置的 Base URL,非标准端点(PrivateLink/FIPS 等)时透传给 boto3
|
||||
:param use_proxy: 是否使用系统代理
|
||||
:param read_timeout: 读取超时秒数,None 时使用默认值
|
||||
:return: boto3 客户端实例
|
||||
"""
|
||||
import boto3
|
||||
from botocore import UNSIGNED
|
||||
|
||||
config = self._build_bedrock_boto3_config(use_proxy)
|
||||
if read_timeout:
|
||||
config = config.merge(type(config)(read_timeout=read_timeout))
|
||||
endpoint_kwargs: dict[str, Any] = {}
|
||||
endpoint_url = self._bedrock_endpoint_url(service_name, base_url)
|
||||
if endpoint_url:
|
||||
endpoint_kwargs["endpoint_url"] = endpoint_url
|
||||
|
||||
if credentials["auth_scheme"] == "sigv4":
|
||||
return boto3.client(
|
||||
service_name,
|
||||
region_name=region,
|
||||
aws_access_key_id=credentials["access_key_id"],
|
||||
aws_secret_access_key=credentials["secret_access_key"],
|
||||
aws_session_token=credentials.get("session_token"),
|
||||
config=config,
|
||||
**endpoint_kwargs,
|
||||
)
|
||||
|
||||
# Bearer 认证:以 UNSIGNED 跳过 SigV4 签名,再把 API Key 注入 Authorization 头。
|
||||
bearer_token = credentials["bearer_token"]
|
||||
config = config.merge(type(config)(signature_version=UNSIGNED))
|
||||
client = boto3.client(
|
||||
service_name,
|
||||
region_name=region,
|
||||
aws_access_key_id="unsigned",
|
||||
aws_secret_access_key="unsigned",
|
||||
config=config,
|
||||
**endpoint_kwargs,
|
||||
)
|
||||
|
||||
def _inject_bearer(request: Any, **_kwargs: Any) -> None:
|
||||
request.headers["Authorization"] = f"Bearer {bearer_token}"
|
||||
|
||||
client.meta.events.register(
|
||||
f"request-created.{service_name}",
|
||||
_inject_bearer,
|
||||
)
|
||||
return client
|
||||
|
||||
async def _list_models_from_bedrock_fallback(
|
||||
self,
|
||||
region: str,
|
||||
use_proxy: Optional[bool] = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
从 models.dev 目录筛选当前 Region 可调用的 Bedrock 模型
|
||||
|
||||
:param region: 当前 Base URL 对应的 AWS Region
|
||||
:param use_proxy: 是否使用系统代理
|
||||
:return: 过滤后的标准化模型记录列表
|
||||
"""
|
||||
models = await self._list_models_from_models_dev_only(
|
||||
provider_id="amazon-bedrock",
|
||||
use_proxy=use_proxy,
|
||||
)
|
||||
return [
|
||||
model
|
||||
for model in models
|
||||
if self._bedrock_model_matches_region(model["id"], region)
|
||||
]
|
||||
|
||||
async def _list_models_from_bedrock(
|
||||
self,
|
||||
api_key: str,
|
||||
base_url: Optional[str],
|
||||
use_proxy: Optional[bool] = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""
|
||||
从 Bedrock 控制面拉取模型目录,聚合跨区 Inference Profile 与直连模型
|
||||
|
||||
Bedrock 多数新模型仅允许通过 Inference Profile(us./eu./apac./global. 前缀)
|
||||
调用,因此优先列出 Profile,再补充支持 ON_DEMAND 直连的基础模型。
|
||||
|
||||
:param api_key: 用户填写的凭证内容(Bedrock API Key 或 AK/SK)
|
||||
:param base_url: Bedrock 运行时端点,决定 Region
|
||||
:param use_proxy: 是否使用系统代理
|
||||
:return: 标准化后的模型记录列表
|
||||
"""
|
||||
credentials = self._parse_bedrock_credentials(api_key)
|
||||
region = self._extract_bedrock_region(base_url)
|
||||
# runtime VPCE 无法安全推导对应的控制面 VPCE;FIPS 端点也不能绕回
|
||||
# 公有非 FIPS 控制面,因此直接使用本地目录。
|
||||
if self._bedrock_endpoint_url("bedrock-runtime", base_url):
|
||||
return await self._list_models_from_bedrock_fallback(region, use_proxy)
|
||||
client = self.create_bedrock_client(
|
||||
"bedrock",
|
||||
region=region,
|
||||
credentials=credentials,
|
||||
use_proxy=use_proxy,
|
||||
)
|
||||
|
||||
def _fetch() -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
|
||||
profiles: list[dict[str, Any]] = []
|
||||
paginator = client.get_paginator("list_inference_profiles")
|
||||
for page in paginator.paginate(typeEquals="SYSTEM_DEFINED"):
|
||||
profiles.extend(page.get("inferenceProfileSummaries") or [])
|
||||
foundation = client.list_foundation_models(
|
||||
byOutputModality="TEXT",
|
||||
byInferenceType="ON_DEMAND",
|
||||
).get("modelSummaries") or []
|
||||
return profiles, foundation
|
||||
|
||||
try:
|
||||
profile_summaries, foundation_summaries = await asyncio.to_thread(_fetch)
|
||||
except Exception as err:
|
||||
# 部分 Bedrock API Key 的授权范围仅覆盖 bedrock-runtime 推理接口,
|
||||
# 控制面查询被拒时降级到 models.dev 目录,保证仍能选择模型。
|
||||
logger.warning(
|
||||
f"获取 Amazon Bedrock 控制面模型列表失败,降级 models.dev 目录: {err}"
|
||||
)
|
||||
return await self._list_models_from_bedrock_fallback(region, use_proxy)
|
||||
finally:
|
||||
await asyncio.to_thread(client.close)
|
||||
|
||||
results: list[dict[str, Any]] = []
|
||||
seen_ids: set[str] = set()
|
||||
|
||||
def _append_record(model_id: str, display_name: Optional[str]) -> None:
|
||||
if not model_id or model_id in seen_ids:
|
||||
return
|
||||
seen_ids.add(model_id)
|
||||
# Inference Profile 带区域前缀,models.dev 目录按基础模型 ID 收录,
|
||||
# 去掉首个前缀段再查一次元数据。
|
||||
metadata = self._cached_models_dev_model("amazon-bedrock", model_id)
|
||||
if not metadata and "." in model_id:
|
||||
metadata = self._cached_models_dev_model(
|
||||
"amazon-bedrock",
|
||||
model_id.split(".", 1)[1],
|
||||
)
|
||||
results.append(
|
||||
self._normalize_model_record(
|
||||
model_id=model_id,
|
||||
display_name=display_name or (metadata or {}).get("name") or model_id,
|
||||
metadata=metadata or {},
|
||||
source="provider",
|
||||
)
|
||||
)
|
||||
|
||||
for profile in profile_summaries:
|
||||
if (profile.get("status") or "ACTIVE") != "ACTIVE":
|
||||
continue
|
||||
_append_record(
|
||||
str(profile.get("inferenceProfileId") or "").strip(),
|
||||
profile.get("inferenceProfileName"),
|
||||
)
|
||||
# 控制面已按当前 Region 和 ON_DEMAND 筛选,不能复用仅面向
|
||||
# models.dev 降级目录的静态白名单,否则 AWS 新增模型会被遗漏。
|
||||
for summary in foundation_summaries:
|
||||
lifecycle = (summary.get("modelLifecycle") or {}).get("status") or "ACTIVE"
|
||||
if lifecycle != "ACTIVE":
|
||||
continue
|
||||
_append_record(
|
||||
str(summary.get("modelId") or "").strip(),
|
||||
summary.get("modelName"),
|
||||
)
|
||||
|
||||
return sorted(results, key=lambda item: item["name"].lower())
|
||||
|
||||
@staticmethod
|
||||
def _copilot_headers(
|
||||
token: Optional[str] = None, include_auth: bool = True
|
||||
@@ -1969,6 +2528,13 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
use_proxy=use_proxy,
|
||||
)
|
||||
|
||||
if resolved_model_list_strategy == "bedrock":
|
||||
return await self._list_models_from_bedrock(
|
||||
api_key=runtime["api_key"],
|
||||
base_url=runtime.get("base_url"),
|
||||
use_proxy=use_proxy,
|
||||
)
|
||||
|
||||
if resolved_model_list_strategy == "anthropic_compatible":
|
||||
return await self._list_models_from_models_dev_only(
|
||||
provider_id=provider_id,
|
||||
@@ -2040,7 +2606,7 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
try:
|
||||
return jwt.decode(token, options={"verify_signature": False})
|
||||
except Exception as err:
|
||||
print(err)
|
||||
logger.debug(f"解析 JWT token 内容失败: {err}")
|
||||
return {}
|
||||
|
||||
@staticmethod
|
||||
@@ -2523,40 +3089,29 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
)
|
||||
normalized_api_key = str(api_key or "").strip() or None
|
||||
normalized_base_url = self._sanitize_base_url(base_url)
|
||||
model_record = None
|
||||
if model:
|
||||
try:
|
||||
model_record = next(
|
||||
(
|
||||
item
|
||||
for item in await self.list_models(
|
||||
normalized_provider_id,
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
base_url_preset_id=normalized_base_url_preset_id,
|
||||
user_agent=user_agent,
|
||||
use_proxy=use_proxy,
|
||||
)
|
||||
if item["id"] == model
|
||||
),
|
||||
None,
|
||||
)
|
||||
except Exception as err:
|
||||
print(err)
|
||||
model_record = None
|
||||
default_transport = (
|
||||
"anthropic" if resolved_runtime == "anthropic_compatible" else "openai"
|
||||
)
|
||||
model_record = self._resolve_cached_model_record(
|
||||
normalized_provider_id,
|
||||
model,
|
||||
base_url=base_url,
|
||||
base_url_preset_id=normalized_base_url_preset_id,
|
||||
transport=default_transport,
|
||||
)
|
||||
model_metadata = self.resolve_cached_model_metadata(
|
||||
normalized_provider_id,
|
||||
model,
|
||||
base_url=base_url,
|
||||
base_url_preset_id=normalized_base_url_preset_id,
|
||||
)
|
||||
|
||||
result: dict[str, Any] = {
|
||||
"provider_id": normalized_provider_id,
|
||||
"runtime": resolved_runtime,
|
||||
"model_id": model,
|
||||
"model_record": model_record,
|
||||
"model_metadata": await self.resolve_model_metadata(
|
||||
normalized_provider_id,
|
||||
model,
|
||||
base_url=base_url,
|
||||
base_url_preset_id=normalized_base_url_preset_id,
|
||||
use_proxy=use_proxy,
|
||||
),
|
||||
"model_metadata": model_metadata,
|
||||
"default_headers": None,
|
||||
"use_responses_api": None,
|
||||
"auth_mode": "api_key",
|
||||
@@ -2567,8 +3122,7 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
try:
|
||||
auth = await self._resolve_chatgpt_oauth()
|
||||
except Exception as err:
|
||||
print(err)
|
||||
pass
|
||||
logger.debug(f"解析 ChatGPT OAuth 鉴权失败,回退 API Key 模式: {err}")
|
||||
|
||||
if auth:
|
||||
headers = {"originator": "moviepilot"}
|
||||
@@ -2648,6 +3202,22 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
)
|
||||
return result
|
||||
|
||||
if resolved_runtime == "bedrock":
|
||||
effective_base_url = normalized_base_url or self._default_base_url_for_provider(
|
||||
spec
|
||||
)
|
||||
credentials = self._parse_bedrock_credentials(normalized_api_key)
|
||||
result.update(
|
||||
{
|
||||
"api_key": normalized_api_key,
|
||||
"base_url": effective_base_url,
|
||||
"aws_region": self._extract_bedrock_region(effective_base_url),
|
||||
"aws_auth": credentials,
|
||||
"auth_mode": "api_key",
|
||||
}
|
||||
)
|
||||
return result
|
||||
|
||||
if resolved_runtime == "anthropic_compatible":
|
||||
effective_base_url = normalized_base_url or self._default_base_url_for_provider(
|
||||
spec
|
||||
|
||||
@@ -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()
|
||||
@@ -6,13 +6,16 @@
|
||||
并在每次 Agent 启动时注入轻量索引,完整日志由工具按需查询。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Awaitable, Callable
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
from typing import Annotated, Any, NotRequired, Optional, TypedDict
|
||||
|
||||
import anyio
|
||||
from anyio import Path as AsyncPath
|
||||
from langchain.agents.middleware.types import (
|
||||
AgentMiddleware,
|
||||
@@ -86,10 +89,6 @@ ACTIVITY_ENTRY_PATTERN = re.compile(r"^-\s+\*\*(?P<time>\d{2}:\d{2})\*\*\s+(?P<s
|
||||
class QueryActivityLogInput(BaseModel):
|
||||
"""查询活动日志工具的输入参数模型。"""
|
||||
|
||||
explanation: Optional[str] = Field(
|
||||
None,
|
||||
description="Clear explanation of why this tool is being used in the current context",
|
||||
)
|
||||
keyword: Optional[str] = Field(
|
||||
None,
|
||||
description=(
|
||||
@@ -286,17 +285,15 @@ class _ActivityLogToolProvider:
|
||||
date: Optional[str] = None,
|
||||
days: Optional[int] = DEFAULT_QUERY_DAYS,
|
||||
limit: Optional[int] = DEFAULT_QUERY_LIMIT,
|
||||
explanation: Optional[str] = None,
|
||||
) -> str:
|
||||
"""查询活动日志并返回 JSON 字符串。"""
|
||||
logger.info(
|
||||
"查询活动日志: keyword=%s, use_regex=%s, date=%s, days=%s, limit=%s, explanation=%s",
|
||||
"查询活动日志: keyword=%s, use_regex=%s, date=%s, days=%s, limit=%s",
|
||||
keyword,
|
||||
use_regex,
|
||||
date,
|
||||
days,
|
||||
limit,
|
||||
explanation or "-",
|
||||
)
|
||||
try:
|
||||
payload = query_activity_logs(
|
||||
@@ -503,6 +500,7 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
||||
self.retention_days = retention_days
|
||||
self.prompt_load_days = prompt_load_days
|
||||
self.stream_handler = stream_handler
|
||||
self._background_tasks: set[asyncio.Task[None]] = set()
|
||||
self._tool_provider = _ActivityLogToolProvider(activity_dir=activity_dir)
|
||||
self.tools = [
|
||||
StructuredTool.from_function(
|
||||
@@ -579,14 +577,29 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
||||
entry = f"- **{now_str}** {summary}\n"
|
||||
try:
|
||||
if await log_path.exists():
|
||||
existing = await log_path.read_text(encoding="utf-8", errors="replace")
|
||||
await log_path.write_text(existing + entry, encoding="utf-8")
|
||||
async with await anyio.open_file(
|
||||
log_path,
|
||||
mode="a",
|
||||
encoding="utf-8",
|
||||
) as stream:
|
||||
await stream.write(entry)
|
||||
else:
|
||||
header = f"# {today_str} 活动日志\n\n"
|
||||
await log_path.write_text(header + entry, encoding="utf-8")
|
||||
logger.debug("Activity logged: %s", summary[:80])
|
||||
try:
|
||||
fd = os.open(log_path, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o644)
|
||||
except FileExistsError:
|
||||
async with await anyio.open_file(
|
||||
log_path,
|
||||
mode="a",
|
||||
encoding="utf-8",
|
||||
) as stream:
|
||||
await stream.write(entry)
|
||||
else:
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as stream:
|
||||
stream.write(header + entry)
|
||||
logger.debug(f"Activity logged: {summary[:80]}")
|
||||
except Exception as e:
|
||||
logger.warning("Failed to append activity log: %s", e)
|
||||
logger.warning(f"Failed to append activity log: {e}")
|
||||
|
||||
async def _cleanup_old_logs(self) -> None:
|
||||
"""清理超过保留天数的旧日志文件。"""
|
||||
@@ -608,20 +621,54 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
||||
file_date = datetime.strptime(match.group(1), "%Y-%m-%d").date()
|
||||
if file_date < cutoff_date:
|
||||
await path.unlink()
|
||||
logger.debug("Cleaned up old activity log: %s", path.name)
|
||||
logger.debug(f"Cleaned up old activity log: {path.name}")
|
||||
except ValueError:
|
||||
continue
|
||||
except Exception as e:
|
||||
logger.warning("Failed to cleanup old activity logs: %s", e)
|
||||
logger.warning(f"Failed to cleanup old activity logs: {e}")
|
||||
|
||||
def _schedule_activity_recording(self, messages: list) -> None:
|
||||
"""提交后台活动记录任务,不阻塞当前 Agent 会话结束。"""
|
||||
task = asyncio.create_task(self._record_activity(messages))
|
||||
self._background_tasks.add(task)
|
||||
task.add_done_callback(self._on_activity_recording_done)
|
||||
|
||||
def _on_activity_recording_done(self, task: asyncio.Task[None]) -> None:
|
||||
"""清理已完成的后台任务并记录未捕获异常。"""
|
||||
self._background_tasks.discard(task)
|
||||
try:
|
||||
task.result()
|
||||
except asyncio.CancelledError:
|
||||
logger.debug("活动日志后台记录任务已取消")
|
||||
except Exception as err:
|
||||
logger.warning(f"活动日志后台记录任务失败: {err}")
|
||||
|
||||
async def _record_activity(self, messages: list) -> None:
|
||||
"""在后台生成本轮活动摘要并写入活动日志。"""
|
||||
try:
|
||||
# 提取本轮交互
|
||||
round_messages = _extract_last_round(messages)
|
||||
if not round_messages:
|
||||
return
|
||||
if _should_skip_activity_summary(round_messages):
|
||||
return
|
||||
|
||||
# 格式化对话文本
|
||||
conversation_text = _format_conversation_for_summary(round_messages)
|
||||
if not conversation_text:
|
||||
return
|
||||
|
||||
# 调用 LLM 生成摘要
|
||||
summary = await _summarize_with_llm(conversation_text)
|
||||
if summary:
|
||||
await self._append_activity(summary)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to record activity: {e}")
|
||||
|
||||
async def abefore_agent(
|
||||
self, state: ActivityLogState, runtime: Runtime
|
||||
) -> Optional[ActivityLogStateUpdate]:
|
||||
"""在 Agent 执行前加载近期活动日志。"""
|
||||
# 如果已经加载则跳过
|
||||
if "activity_log_contents" in state:
|
||||
return None
|
||||
|
||||
contents = await self._load_recent_logs()
|
||||
|
||||
# 趁机清理旧日志(低频操作,不影响性能)
|
||||
@@ -686,30 +733,14 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
||||
async def aafter_agent(
|
||||
self, state: ActivityLogState, runtime: Runtime
|
||||
) -> Optional[dict[str, Any]]:
|
||||
"""Agent 执行完毕后,调用 LLM 对本轮对话生成摘要并追加到当日活动日志。"""
|
||||
"""Agent 执行完毕后,异步提交活动日志记录任务。"""
|
||||
try:
|
||||
messages = state.get("messages", [])
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
# 提取本轮交互
|
||||
round_messages = _extract_last_round(messages)
|
||||
if not round_messages:
|
||||
return None
|
||||
if _should_skip_activity_summary(round_messages):
|
||||
return None
|
||||
|
||||
# 格式化对话文本
|
||||
conversation_text = _format_conversation_for_summary(round_messages)
|
||||
if not conversation_text:
|
||||
return None
|
||||
|
||||
# 调用 LLM 生成摘要
|
||||
summary = await _summarize_with_llm(conversation_text)
|
||||
if summary:
|
||||
await self._append_activity(summary)
|
||||
self._schedule_activity_recording(list(messages))
|
||||
except Exception as e:
|
||||
logger.warning("Failed to record activity: %s", e)
|
||||
logger.warning(f"Failed to record activity: {e}")
|
||||
|
||||
return None
|
||||
|
||||
|
||||
@@ -204,9 +204,14 @@ You have a scheduled jobs system for user-requested delayed or recurring work.
|
||||
{jobs_list}
|
||||
|
||||
Rules:
|
||||
- Create jobs only when the user asks for delayed, recurring, reminder, or monitoring behavior.
|
||||
- Do not create jobs for immediate one-time work or work already handled by MoviePilot schedulers.
|
||||
- Each job lives in its own directory with a `JOB.md`; read the listed file before executing or updating an active job.
|
||||
- For new delayed, recurring, reminder, or monitoring work, use the dedicated
|
||||
`create_agent_task`, `query_agent_tasks`, `update_agent_task`, `run_agent_task`,
|
||||
and `delete_agent_task` tools. These tools use integer task IDs. Do not create
|
||||
or edit JOB.md files for new tasks.
|
||||
- Use `query_schedulers` and `run_scheduler` only for MoviePilot system, plugin,
|
||||
or workflow runtime services; never pass their string job IDs to Agent task tools.
|
||||
- Do not create tasks for immediate one-time work or work already handled by MoviePilot schedulers.
|
||||
- Entries listed above are legacy JOB.md tasks. Read their files only when a heartbeat asks you to execute them.
|
||||
- During heartbeat checks, act only on `pending` or `in_progress` jobs, update status/last_run/logs, and leave recurring jobs `pending` after each run.
|
||||
</jobs_system>
|
||||
"""
|
||||
@@ -230,7 +235,7 @@ class JobsMiddleware(AgentMiddleware[JobsState, ContextT, ResponseT]): # noqa
|
||||
def _format_jobs_list(jobs: list[JobMetadata]) -> str:
|
||||
"""格式化任务元数据列表用于系统提示词。"""
|
||||
if not jobs:
|
||||
return "(No active jobs. You can create jobs when users request periodic or scheduled tasks.)"
|
||||
return "(No active legacy JOB.md tasks. Use create_agent_task for new scheduled work.)"
|
||||
|
||||
lines = []
|
||||
for job in jobs:
|
||||
@@ -283,12 +288,7 @@ class JobsMiddleware(AgentMiddleware[JobsState, ContextT, ResponseT]): # noqa
|
||||
) -> JobsStateUpdate | None:
|
||||
"""在 Agent 执行前异步加载任务元数据。
|
||||
|
||||
每个会话仅加载一次。若 state 中已有则跳过。
|
||||
"""
|
||||
# 如果 state 中已存在元数据则跳过
|
||||
if "jobs_metadata" in state:
|
||||
return None
|
||||
|
||||
return JobsStateUpdate(
|
||||
jobs_metadata=await load_jobs_metadata(self.sources)
|
||||
)
|
||||
|
||||
@@ -302,7 +302,6 @@ class MemoryMiddleware(AgentMiddleware[MemoryState, ContextT, ResponseT]): # no
|
||||
"""在代理执行前扫描记忆目录并加载所有 .md 文件的内容。
|
||||
|
||||
自动发现目录下所有 `.md` 文件并加载其内容到状态中。
|
||||
如果状态中尚未存在则进行加载。
|
||||
同时检测记忆文件是否为空,设置 memory_empty 标志位,
|
||||
以便在系统提示词中触发初始化引导流程。
|
||||
|
||||
@@ -314,10 +313,6 @@ class MemoryMiddleware(AgentMiddleware[MemoryState, ContextT, ResponseT]): # no
|
||||
返回:
|
||||
填充了 memory_contents 和 memory_empty 的状态更新。
|
||||
"""
|
||||
# 如果已经加载则跳过
|
||||
if "memory_contents" in state:
|
||||
return None
|
||||
|
||||
# 扫描目录下所有 .md 文件
|
||||
md_files = await self._scan_memory_files()
|
||||
|
||||
|
||||
@@ -92,10 +92,6 @@ class SkillsStateUpdate(TypedDict):
|
||||
class SkillToolInput(BaseModel):
|
||||
"""Skill 加载工具的输入参数模型。"""
|
||||
|
||||
explanation: Optional[str] = Field(
|
||||
None,
|
||||
description="Clear explanation of why this skill is needed in the current context",
|
||||
)
|
||||
name: str = Field(
|
||||
...,
|
||||
description="Skill name or id from the available skills list.",
|
||||
@@ -322,7 +318,7 @@ def _extract_version(skill_md: Path) -> int:
|
||||
try:
|
||||
content = skill_md.read_text(encoding="utf-8", errors="replace")
|
||||
except Exception as err:
|
||||
print(err)
|
||||
logger.debug(f"读取技能版本失败: {err}")
|
||||
return 0
|
||||
match = re.match(r"^---\s*\n(.*?)\n---\s*\n", content, re.DOTALL)
|
||||
if not match:
|
||||
@@ -460,9 +456,9 @@ class _SkillToolProvider:
|
||||
raw_content = await handle.read(MAX_SKILL_FILE_SIZE)
|
||||
return raw_content.decode("utf-8", errors="replace"), truncated
|
||||
|
||||
async def load_skill(self, name: str, explanation: Optional[str] = None) -> str:
|
||||
async def load_skill(self, name: str) -> str:
|
||||
"""加载指定 Skill 的完整说明并返回 JSON 字符串。"""
|
||||
logger.info(f"加载 Skill: name={name}, explanation={explanation or '-'}")
|
||||
logger.info(f"加载 Skill: name={name}")
|
||||
try:
|
||||
skill = await self._find_skill(name)
|
||||
if not skill:
|
||||
@@ -627,13 +623,8 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
|
||||
) -> SkillsStateUpdate | None: # ty: ignore[invalid-method-override]
|
||||
"""在 Agent 执行前异步加载技能元数据。
|
||||
|
||||
每个会话仅加载一次。若 state 中已有则跳过。
|
||||
首次加载时,会先将内置技能同步到用户目录(如不存在)。
|
||||
"""
|
||||
# 如果 state 中已存在元数据则跳过
|
||||
if "skills_metadata" in state:
|
||||
return None
|
||||
|
||||
self._sync_bundled_skills()
|
||||
|
||||
all_skills: dict[str, SkillMetadata] = {}
|
||||
@@ -679,8 +670,7 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
|
||||
if not isinstance(tool_args, dict):
|
||||
tool_args = {}
|
||||
logger.info(
|
||||
f"开始执行 Skill 工具: name={tool_args.get('name') or '-'}, "
|
||||
f"explanation={tool_args.get('explanation') or '-'}"
|
||||
f"开始执行 Skill 工具: name={tool_args.get('name') or '-'}"
|
||||
)
|
||||
if self.stream_handler and getattr(self.stream_handler, "is_streaming", False):
|
||||
self.stream_handler.record_tool_call(
|
||||
|
||||
@@ -197,18 +197,44 @@ def is_subagent_stream_metadata(metadata: Any) -> bool:
|
||||
) == SUBAGENT_STREAM_MARKER_VALUE:
|
||||
return True
|
||||
|
||||
return bool(metadata.get("lc_agent_name") in builtin_subagent_names())
|
||||
return bool(
|
||||
metadata.get("lc_agent_name")
|
||||
in builtin_subagent_names(agent_runtime_manager.current_signature())
|
||||
)
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def builtin_subagent_names() -> frozenset[str]:
|
||||
def builtin_subagent_names(
|
||||
runtime_signature: Optional[tuple[tuple[str, int, int], ...]] = None,
|
||||
) -> frozenset[str]:
|
||||
"""返回内置子代理名称集合。"""
|
||||
return frozenset(profile.name for profile in _builtin_subagent_profiles())
|
||||
runtime_signature = runtime_signature or agent_runtime_manager.current_signature()
|
||||
return _cached_builtin_subagent_names(runtime_signature)
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _builtin_subagent_profiles() -> tuple[_SubAgentProfile, ...]:
|
||||
@lru_cache(maxsize=8)
|
||||
def _cached_builtin_subagent_names(
|
||||
runtime_signature: tuple[tuple[str, int, int], ...],
|
||||
) -> frozenset[str]:
|
||||
"""按运行时签名缓存内置子代理名称集合。"""
|
||||
return frozenset(
|
||||
profile.name
|
||||
for profile in _builtin_subagent_profiles(runtime_signature)
|
||||
)
|
||||
|
||||
|
||||
def _builtin_subagent_profiles(
|
||||
runtime_signature: Optional[tuple[tuple[str, int, int], ...]] = None,
|
||||
) -> tuple[_SubAgentProfile, ...]:
|
||||
"""从运行时配置目录加载 MoviePilot 子代理定义。"""
|
||||
runtime_signature = runtime_signature or agent_runtime_manager.current_signature()
|
||||
return _cached_builtin_subagent_profiles(runtime_signature)
|
||||
|
||||
|
||||
@lru_cache(maxsize=8)
|
||||
def _cached_builtin_subagent_profiles(
|
||||
runtime_signature: tuple[tuple[str, int, int], ...],
|
||||
) -> tuple[_SubAgentProfile, ...]:
|
||||
"""按运行时签名缓存 MoviePilot 子代理定义。"""
|
||||
definitions = agent_runtime_manager.list_subagents()
|
||||
profiles = tuple(
|
||||
_profile_from_runtime_definition(definition)
|
||||
@@ -237,6 +263,10 @@ def _builtin_subagent_profiles() -> tuple[_SubAgentProfile, ...]:
|
||||
)
|
||||
|
||||
|
||||
builtin_subagent_names.cache_clear = _cached_builtin_subagent_names.cache_clear
|
||||
_builtin_subagent_profiles.cache_clear = _cached_builtin_subagent_profiles.cache_clear
|
||||
|
||||
|
||||
def _profile_from_runtime_definition(
|
||||
definition: SubAgentDefinition,
|
||||
) -> _SubAgentProfile:
|
||||
@@ -1044,6 +1074,7 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
||||
if unfinished_records:
|
||||
logger.info(f"Agent 结束,取消未完成子代理任务: tasks={len(unfinished_records)}")
|
||||
await self._cancel_records(unfinished_records)
|
||||
self._tasks.clear()
|
||||
|
||||
async def awrap_tool_call(
|
||||
self,
|
||||
@@ -1083,9 +1114,8 @@ def create_subagent_middlewares(
|
||||
stream_handler: Any = None,
|
||||
) -> tuple[list[AgentMiddleware], list[BaseTool]]:
|
||||
"""创建子代理中间件列表和任务工具列表。"""
|
||||
_builtin_subagent_profiles.cache_clear()
|
||||
builtin_subagent_names.cache_clear()
|
||||
profiles = _builtin_subagent_profiles()
|
||||
runtime_signature = agent_runtime_manager.current_signature()
|
||||
profiles = _builtin_subagent_profiles(runtime_signature)
|
||||
subagent_middleware = MoviePilotSubAgentMiddleware(
|
||||
model=model,
|
||||
profiles=profiles,
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
"""MoviePilot 自定义工具筛选中间件。"""
|
||||
|
||||
from dataclasses import replace
|
||||
from dataclasses import dataclass, replace
|
||||
import json
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Annotated, Any, NotRequired
|
||||
@@ -20,7 +20,7 @@ from langchain.agents.middleware.tool_selection import (
|
||||
LLMToolSelectorMiddleware,
|
||||
)
|
||||
from langchain_core.language_models.chat_models import BaseChatModel
|
||||
from langchain_core.messages import AIMessage, HumanMessage
|
||||
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langchain_core.tools import BaseTool
|
||||
from langgraph.runtime import Runtime
|
||||
@@ -68,19 +68,25 @@ class ToolSelectionStateUpdate(TypedDict):
|
||||
selected_tool_names: list[str] | None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ToolSelectionAttempt:
|
||||
"""工具筛选尝试结果,用于统一记录最终日志。"""
|
||||
|
||||
request: ModelRequest
|
||||
selected_tool_names: list[str]
|
||||
status: str
|
||||
detail: str = ""
|
||||
|
||||
|
||||
class ToolSelectorMiddleware(LLMToolSelectorMiddleware):
|
||||
"""
|
||||
为 DeepSeek 兼容端点提供更稳妥的工具筛选实现。
|
||||
使用 provider-neutral JSON 提示执行工具筛选。
|
||||
|
||||
LangChain 默认会通过 `with_structured_output()` 走 OpenAI 的
|
||||
`response_format=json_schema` 路径,但 DeepSeek 官方 OpenAI 兼容端点公开文档
|
||||
仅保证 `json_object` 模式可用。对于 `deepseek-reasoner`,这会在工具筛选阶段
|
||||
提前触发 400,导致 Agent 还没真正开始执行工具就失败。
|
||||
|
||||
因此这里仅在识别到 DeepSeek 模型/端点时,退回到显式 JSON 输出模式:
|
||||
1. 使用 `response_format={"type": "json_object"}`;
|
||||
2. 在提示词中明确约束返回 JSON 结构;
|
||||
3. 手动解析 `{"tools": [...]}`,其余模型继续沿用 LangChain 默认实现。
|
||||
LangChain 默认会通过 `with_structured_output()` 走 provider-specific 的
|
||||
结构化输出能力,不同 OpenAI/Anthropic 兼容端点对 `response_format`、
|
||||
JSON schema 和工具绑定的支持并不一致。工具筛选只是 Agent 执行前的
|
||||
辅助优化,失败时也会恢复使用全部工具,因此这里统一使用文本提示约束
|
||||
模型返回 `{"tools": [...]}` 并手动解析,避免在筛选阶段引入额外兼容分支。
|
||||
|
||||
另外,LangChain 原生工具筛选挂在 `wrap_model_call` 上,会在同一条用户请求
|
||||
的每次“模型回合”前都重新筛选一次工具。对于会多轮调用工具的复杂任务,
|
||||
@@ -322,21 +328,16 @@ class ToolSelectorMiddleware(LLMToolSelectorMiddleware):
|
||||
request: ModelRequest[ContextT],
|
||||
) -> ModelRequest[ContextT]:
|
||||
"""
|
||||
处理工具筛选响应,并保留空结果回退所有工具的 MoviePilot 策略。
|
||||
处理工具筛选响应,并在正常空结果时禁用可筛选工具。
|
||||
"""
|
||||
if response.get("tools") == []:
|
||||
logger.warning("工具筛选结果为空,将恢复使用所有工具。")
|
||||
|
||||
always_included_tools: list[BaseTool] = [
|
||||
tool
|
||||
for tool in request.tools
|
||||
if not isinstance(tool, dict) and tool.name in self.always_include
|
||||
]
|
||||
provider_tools = [tool for tool in request.tools if isinstance(tool, dict)]
|
||||
|
||||
return request.override(
|
||||
tools=[*available_tools, *always_included_tools, *provider_tools]
|
||||
)
|
||||
return request.override(tools=[*always_included_tools, *provider_tools])
|
||||
|
||||
response["tools"] = self._complete_low_count_selection(
|
||||
selected_tool_names=[
|
||||
@@ -347,47 +348,21 @@ class ToolSelectorMiddleware(LLMToolSelectorMiddleware):
|
||||
valid_tool_names=valid_tool_names,
|
||||
available_tools=available_tools,
|
||||
)
|
||||
return super()._process_selection_response(
|
||||
modified_request = super()._process_selection_response(
|
||||
response,
|
||||
available_tools,
|
||||
valid_tool_names,
|
||||
request,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _is_deepseek_compatible_model(model: BaseChatModel) -> bool:
|
||||
"""
|
||||
判断当前模型是否应当走 DeepSeek JSON 兼容分支。
|
||||
|
||||
除了官方 `langchain_deepseek`,用户也可能通过 OpenAI-compatible
|
||||
配置把 DeepSeek 端点接到 `ChatOpenAI`。因此这里同时检查模块名、模型名
|
||||
和 Base URL,避免只靠单一条件漏判。
|
||||
"""
|
||||
module_name = type(model).__module__.lower()
|
||||
model_name = (
|
||||
str(getattr(model, "model_name", "") or getattr(model, "model", ""))
|
||||
.strip()
|
||||
.lower()
|
||||
)
|
||||
base_url = (
|
||||
str(getattr(model, "openai_api_base", "") or getattr(model, "api_base", ""))
|
||||
.strip()
|
||||
.lower()
|
||||
)
|
||||
|
||||
return (
|
||||
"deepseek" in module_name
|
||||
or model_name.startswith("deepseek-")
|
||||
or "api.deepseek.com" in base_url
|
||||
)
|
||||
return modified_request
|
||||
|
||||
@staticmethod
|
||||
def _parse_json_object(text: str) -> dict[str, Any]:
|
||||
"""
|
||||
解析模型返回的 JSON。
|
||||
|
||||
DeepSeek 在 JSON 模式下通常会返回纯 JSON,但这里仍做一层兜底,
|
||||
兼容模型偶发输出围栏或前后说明文本的情况。
|
||||
不同模型可能偶发输出 Markdown 围栏或前后说明文本,因此这里从
|
||||
响应中提取第一个 JSON 对象作为兜底。
|
||||
"""
|
||||
stripped_text = text.strip()
|
||||
if not stripped_text:
|
||||
@@ -440,16 +415,16 @@ class ToolSelectorMiddleware(LLMToolSelectorMiddleware):
|
||||
)
|
||||
return f"Capability groups from tool tags:\n{rendered_groups}\n\n"
|
||||
|
||||
def _build_deepseek_selection_prompt(self, selection_request: Any) -> str:
|
||||
def _build_json_selection_prompt(self, selection_request: Any) -> str:
|
||||
"""
|
||||
为 DeepSeek 生成显式 JSON 输出提示。
|
||||
生成显式 JSON 输出提示。
|
||||
|
||||
DeepSeek 官方文档要求在 JSON 输出模式下,提示词中必须明确包含 JSON
|
||||
约束,否则兼容端点可能返回空内容或无意义输出。
|
||||
使用纯提示约束可覆盖更多兼容端点,避免在工具筛选阶段依赖某个
|
||||
provider 专属的 `response_format` 或 schema 能力。
|
||||
"""
|
||||
limit_instruction = ""
|
||||
if self.max_tools:
|
||||
limit_instruction = f"- Select up to {self.max_tools} tools. IF NO TOOLS ARE RELEVANT, DO NOT RETURN AN EMPTY ARRAY. SELECT THE MOST APPLICABLE ONES TO ENSURE THE REQUEST IS HANDLED."
|
||||
limit_instruction = f"- Select up to {self.max_tools} tools. Return an empty array if no tools are relevant."
|
||||
|
||||
return (
|
||||
f"{selection_request.system_message}\n\n"
|
||||
@@ -469,7 +444,7 @@ class ToolSelectorMiddleware(LLMToolSelectorMiddleware):
|
||||
|
||||
def _normalize_selection_response(self, response: Any) -> dict[str, list[str]]:
|
||||
"""
|
||||
解析并标准化 DeepSeek JSON 模式的工具筛选结果。
|
||||
解析并标准化显式 JSON 模式的工具筛选结果。
|
||||
"""
|
||||
content = getattr(response, "content", response)
|
||||
text = LLMHelper.extract_text_content(content)
|
||||
@@ -486,22 +461,21 @@ class ToolSelectorMiddleware(LLMToolSelectorMiddleware):
|
||||
logger.debug(f"工具筛选标准化结果: {normalized_tools}")
|
||||
return {"tools": normalized_tools}
|
||||
|
||||
async def _aselect_tools_with_deepseek(
|
||||
async def _aselect_tools_with_json_prompt(
|
||||
self, selection_request: Any
|
||||
) -> dict[str, list[str]]:
|
||||
"""
|
||||
使用 DeepSeek 兼容的 JSON 输出模式执行异步工具筛选。
|
||||
使用 JSON 提示执行异步工具筛选。
|
||||
|
||||
:param selection_request: LangChain 工具筛选请求
|
||||
:return: 标准化后的工具名列表
|
||||
"""
|
||||
logger.debug("工具筛选走 DeepSeek JSON 兼容分支")
|
||||
structured_model = selection_request.model.bind(
|
||||
response_format={"type": "json_object"}
|
||||
)
|
||||
response = await structured_model.ainvoke(
|
||||
logger.debug("工具筛选走 JSON 提示分支")
|
||||
response = await selection_request.model.ainvoke(
|
||||
[
|
||||
{
|
||||
"role": "system",
|
||||
"content": self._build_deepseek_selection_prompt(selection_request),
|
||||
},
|
||||
SystemMessage(
|
||||
content=self._build_json_selection_prompt(selection_request)
|
||||
),
|
||||
selection_request.last_user_message,
|
||||
]
|
||||
)
|
||||
@@ -512,6 +486,31 @@ class ToolSelectorMiddleware(LLMToolSelectorMiddleware):
|
||||
"""从已筛选后的请求中提取最终工具名,保留原有顺序。"""
|
||||
return [tool.name for tool in request.tools if not isinstance(tool, dict)]
|
||||
|
||||
@staticmethod
|
||||
def _count_request_tools(request: ModelRequest) -> int:
|
||||
"""统计当前请求中的 LangChain 工具数量,不包含 provider 原生工具字典。"""
|
||||
return len([tool for tool in request.tools if not isinstance(tool, dict)])
|
||||
|
||||
@classmethod
|
||||
def _log_selection_attempt(cls, attempt: _ToolSelectionAttempt) -> None:
|
||||
"""按工具筛选最终状态记录稳定日志。"""
|
||||
tool_count = cls._count_request_tools(attempt.request)
|
||||
if attempt.status == "selected":
|
||||
selected_text = ", ".join(attempt.selected_tool_names) or "无有效工具"
|
||||
logger.info(f"工具筛选结果: {selected_text}")
|
||||
return
|
||||
if attempt.status == "failed_fallback":
|
||||
logger.warning(
|
||||
f"工具筛选失败,将恢复使用所有工具(共 {tool_count} 个): {attempt.detail}"
|
||||
)
|
||||
return
|
||||
if attempt.status == "skipped":
|
||||
logger.info(f"工具筛选跳过: {attempt.detail}。")
|
||||
return
|
||||
if attempt.status == "reused":
|
||||
selected_text = ", ".join(attempt.selected_tool_names) or "无有效工具"
|
||||
logger.info(f"工具筛选复用已有结果: {selected_text}")
|
||||
|
||||
@staticmethod
|
||||
def _apply_selected_tools(
|
||||
request: ModelRequest[ContextT],
|
||||
@@ -523,9 +522,6 @@ class ToolSelectorMiddleware(LLMToolSelectorMiddleware):
|
||||
这里只复用首次筛选出的客户端工具名;provider-specific 的 dict 工具仍然
|
||||
原样保留,避免破坏 LangChain/provider 自身的工具绑定约定。
|
||||
"""
|
||||
if not selected_tool_names:
|
||||
return request
|
||||
|
||||
current_tools_by_name = {
|
||||
tool.name: tool for tool in request.tools if not isinstance(tool, dict)
|
||||
}
|
||||
@@ -546,30 +542,43 @@ class ToolSelectorMiddleware(LLMToolSelectorMiddleware):
|
||||
这里单独抽成 helper,便于首次筛选后缓存结果,也便于测试覆盖
|
||||
“首轮筛选,后续复用”的行为。
|
||||
"""
|
||||
return (await self._aselect_request_once_with_status(request)).request
|
||||
|
||||
async def _aselect_request_once_with_status(
|
||||
self, request: ModelRequest[ContextT]
|
||||
) -> _ToolSelectionAttempt:
|
||||
"""
|
||||
执行一次真实工具筛选,并携带最终状态供调用方统一记录日志。
|
||||
"""
|
||||
selection_request = self._prepare_selection_request(request)
|
||||
if selection_request is None:
|
||||
return request
|
||||
return _ToolSelectionAttempt(
|
||||
request=request,
|
||||
selected_tool_names=self._extract_selected_tool_names(request),
|
||||
status="skipped",
|
||||
detail="没有需要筛选的工具",
|
||||
)
|
||||
|
||||
if not self._is_deepseek_compatible_model(selection_request.model):
|
||||
captured_request: ModelRequest[ContextT] = request
|
||||
|
||||
async def _capture_handler(
|
||||
updated_request: ModelRequest[ContextT],
|
||||
) -> ModelRequest[ContextT]:
|
||||
nonlocal captured_request
|
||||
captured_request = updated_request
|
||||
return updated_request
|
||||
|
||||
await super().awrap_model_call(request, _capture_handler)
|
||||
return captured_request
|
||||
|
||||
response = await self._aselect_tools_with_deepseek(selection_request)
|
||||
return self._process_selection_response(
|
||||
response,
|
||||
selection_request.available_tools,
|
||||
selection_request.valid_tool_names,
|
||||
request,
|
||||
)
|
||||
try:
|
||||
response = await self._aselect_tools_with_json_prompt(selection_request)
|
||||
modified_request = self._process_selection_response(
|
||||
response,
|
||||
selection_request.available_tools,
|
||||
selection_request.valid_tool_names,
|
||||
request,
|
||||
)
|
||||
return _ToolSelectionAttempt(
|
||||
request=modified_request,
|
||||
selected_tool_names=self._extract_selected_tool_names(modified_request),
|
||||
status="selected",
|
||||
)
|
||||
except Exception as err:
|
||||
return _ToolSelectionAttempt(
|
||||
request=request,
|
||||
selected_tool_names=self._extract_selected_tool_names(request),
|
||||
status="failed_fallback",
|
||||
detail=str(err),
|
||||
)
|
||||
|
||||
async def abefore_agent( # noqa
|
||||
self,
|
||||
@@ -583,10 +592,22 @@ class ToolSelectorMiddleware(LLMToolSelectorMiddleware):
|
||||
这样后续多轮 `model -> tools -> model` 循环都只复用这一次结果,
|
||||
不会为每次模型回合重复追加一笔 selector LLM 开销。
|
||||
"""
|
||||
if "selected_tool_names" in state:
|
||||
return None
|
||||
|
||||
if not self.selection_tools or self.model is None:
|
||||
detail = "没有可筛选工具" if not self.selection_tools else "未配置筛选模型"
|
||||
self._log_selection_attempt(
|
||||
_ToolSelectionAttempt(
|
||||
request=ModelRequest(
|
||||
model=self.model,
|
||||
tools=list(self.selection_tools),
|
||||
messages=state["messages"],
|
||||
state=state,
|
||||
runtime=runtime,
|
||||
),
|
||||
selected_tool_names=[],
|
||||
status="skipped",
|
||||
detail=detail,
|
||||
)
|
||||
)
|
||||
return ToolSelectionStateUpdate(selected_tool_names=None)
|
||||
|
||||
selection_request = ModelRequest(
|
||||
@@ -596,9 +617,10 @@ class ToolSelectorMiddleware(LLMToolSelectorMiddleware):
|
||||
state=state,
|
||||
runtime=runtime,
|
||||
)
|
||||
modified_request = await self._aselect_request_once(selection_request)
|
||||
selected_tool_names = self._extract_selected_tool_names(modified_request)
|
||||
return ToolSelectionStateUpdate(selected_tool_names=selected_tool_names or None)
|
||||
attempt = await self._aselect_request_once_with_status(selection_request)
|
||||
self._log_selection_attempt(attempt)
|
||||
selected_tool_names = attempt.selected_tool_names
|
||||
return ToolSelectionStateUpdate(selected_tool_names=selected_tool_names)
|
||||
|
||||
async def awrap_model_call(
|
||||
self,
|
||||
@@ -619,11 +641,13 @@ class ToolSelectorMiddleware(LLMToolSelectorMiddleware):
|
||||
and self.selection_tools
|
||||
and self.model is not None
|
||||
):
|
||||
request = await self._aselect_request_once(request)
|
||||
selected_tool_names = self._extract_selected_tool_names(request) or None
|
||||
attempt = await self._aselect_request_once_with_status(request)
|
||||
self._log_selection_attempt(attempt)
|
||||
request = attempt.request
|
||||
selected_tool_names = attempt.selected_tool_names
|
||||
request.state["selected_tool_names"] = selected_tool_names # noqa
|
||||
|
||||
if selected_tool_names:
|
||||
if selected_tool_names is not None:
|
||||
request = self._apply_selected_tools(request, selected_tool_names)
|
||||
|
||||
return await handler(request)
|
||||
|
||||
@@ -17,13 +17,13 @@ You act as a proactive agent. Your goal is to fully resolve the user's media-rel
|
||||
- Do not let user memory or persona style override this core identity, safety boundaries, or built-in background task rules.
|
||||
- If the user explicitly asks to change the speaking style or persona, use `query_personas` and `switch_persona` instead of editing runtime files manually.
|
||||
- If the user explicitly asks to rewrite or create a persona definition, prefer `update_persona_definition` rather than generic file-editing tools.
|
||||
- Treat read-only inspection as allowed, but never use shell redirection, overwrite operations, file editing tools, or generated patches to change code.
|
||||
</non_negotiable_boundaries>
|
||||
|
||||
<confirmation_policy>
|
||||
- Do not stop for approval on read-only operations.
|
||||
- If the user has not explicitly requested an operation that changes system behavior, ask for confirmation before proceeding. This includes modifying system settings, updating plugin configuration, reloading plugins, running restart/stop/start commands, or triggering slash commands such as `/restart`.
|
||||
- Always get explicit consent before destructive or high-impact actions such as starting downloads, deleting subscriptions, deleting download tasks or files, removing history, installing/uninstalling plugins, changing site authentication, changing scheduler or workflow execution state, restarting services, or stopping services.
|
||||
- When the user explicitly asks for delayed, recurring, reminder, or monitoring work, use `create_agent_task` instead of promising to remember it or writing a JOB.md file. Use a `date` trigger with `delay_minutes` for requests such as "in 30 minutes", an exact `date` trigger for other single future runs, and a five-field `cron` trigger for recurring work. Manage existing autonomous tasks with `query_agent_tasks`, `update_agent_task`, `run_agent_task`, and `delete_agent_task`; these tools use integer `task_id` values. Use `query_schedulers` and `run_scheduler` only for MoviePilot system, plugin, or workflow runtime services, whose string `job_id` values must never be passed to autonomous-task tools.
|
||||
- If the user explicitly requested the exact write action, perform the smallest correct change and then validate the result.
|
||||
- If a requested action is ambiguous between read-only inspection and state change, inspect first and ask a short confirmation question before the state-changing step.
|
||||
</confirmation_policy>
|
||||
@@ -65,7 +65,11 @@ You act as a proactive agent. Your goal is to fully resolve the user's media-rel
|
||||
- If `search_media` fails, fall back to `search_web` or `recognize_media`. Only ask the user when automated paths are exhausted.
|
||||
- If torrent search yields no useful result, check site scope, site health, and recognition quality before concluding that the resource is unavailable.
|
||||
- Reuse the latest torrent search cache for `get_search_results` and `add_download_tasks` instead of re-running the same search unnecessarily.
|
||||
- Use `execute_command` only for diagnostics, read-only inspection, or commands the user explicitly asked to run. Its default `action=start` starts a managed background session and returns `session_id`, `status`, `last_seq`, and `output_until_seq`; call the same tool again with `action=read`, `action=wait`, `action=write`, or `action=kill` to poll output, wait in short segments, send stdin, or stop the process.
|
||||
- For administrator code discovery across local files, use `execute_command(action="run")` with `rg` and narrow globs or paths. Use `list_directory` to inspect one known directory or a supported remote storage backend, and use `read_file` when the exact local file is known.
|
||||
- Read the relevant file before changing it. Use `edit_file` for localized exact replacements; make `old_text` unique with enough surrounding context, and use `replace_all=true` only when every match must change. Use `write_file` for new files; set `overwrite=true` only for an intentional full rewrite, and use `read_file(include_metadata=true)` plus `expected_sha256` when preserving the previously read version matters.
|
||||
- When implementation depends on a Python or Node.js API, first identify the installed or locked dependency version from environment metadata, requirements, package manifests, lockfiles, local source, and type declarations. Use `rg` against the relevant package directory, `.venv`, or `node_modules` instead of scanning the entire project without bounds. If local evidence is insufficient, use `search_web` and then `browse_webpage` to read the matching version of the official documentation. Do not guess signatures from memory, mix examples from incompatible versions, or install a package only to inspect its API.
|
||||
- Use structured file tools for source edits because they enforce file access boundaries and conflict checks. Never use shell redirection, inline scripts, or another tool to bypass a file-tool permission denial.
|
||||
- Use `execute_command` for administrator-only multi-file diagnostics, tests, Git, service operations, SSH, or an exact command the user requested. Use `action=run` for short bounded commands. Use `action=start` for long-running or interactive commands, including SSH; then continue with `read`, `wait`, `write`, or `kill` using the returned `session_id`. Do not start a background session for a short command that can finish within `action=run`.
|
||||
</tool_strategy>
|
||||
|
||||
<media_rules>
|
||||
|
||||
@@ -79,6 +79,10 @@ task_types:
|
||||
- "- Transfer mode: {transfer_mode}"
|
||||
- "- Current TMDB ID: {tmdbid}"
|
||||
- "- Current Douban ID: {doubanid}"
|
||||
- "- Current Bangumi ID: {bangumiid}"
|
||||
- "- Current AniList ID: {anilistid}"
|
||||
- "- Current media source: {media_source}"
|
||||
- "- Current source-native ID: {media_id}"
|
||||
- "- Error message: {error_message}"
|
||||
steps_title: "Required workflow"
|
||||
steps:
|
||||
@@ -90,7 +94,7 @@ task_types:
|
||||
- "Only continue when you have high confidence in the target media."
|
||||
- "Before re-organizing, delete the old transfer history record with `delete_transfer_history` so the system will not skip the source file."
|
||||
- "Then use `transfer_file` to organize the source path directly."
|
||||
- "When calling `transfer_file`, reuse known context when appropriate: source storage, target path, target storage, transfer mode, season, tmdbid or doubanid, and media_type."
|
||||
- "When calling `transfer_file`, reuse known context when appropriate: source storage, target path, target storage, transfer mode, season, all known media IDs, media_source, media_id, and media_type."
|
||||
- "If this record is already correct and no re-organize is needed, do not perform destructive actions; simply report that no change is necessary."
|
||||
task_rules:
|
||||
- "Do NOT rely on previous chat context. Work only from the record above."
|
||||
@@ -116,7 +120,7 @@ task_types:
|
||||
- "If a source file no longer exists or cannot be safely processed, skip that record and note the reason."
|
||||
- "Before re-organizing a record, delete the old transfer history record with `delete_transfer_history` so the system will not skip the source file."
|
||||
- "Then use `transfer_file` to organize the source path directly."
|
||||
- "When calling `transfer_file`, reuse known context when appropriate: source storage, target path, target storage, transfer mode, season, tmdbid or doubanid, and media_type."
|
||||
- "When calling `transfer_file`, reuse known context when appropriate: source storage, target path, target storage, transfer mode, season, all known media IDs, media_source, media_id, and media_type."
|
||||
- "If a record is already correct and no re-organize is needed, do not perform destructive actions; simply mark it as skipped."
|
||||
- "Report only the aggregate outcome, including how many records succeeded, skipped, and failed."
|
||||
task_rules:
|
||||
|
||||
@@ -315,8 +315,6 @@ class PromptManager:
|
||||
"项目根目录": settings.ROOT_PATH,
|
||||
"配置目录": settings.CONFIG_PATH,
|
||||
"临时目录": settings.TEMP_PATH,
|
||||
"日志目录": settings.LOG_PATH,
|
||||
"主日志文件": settings.LOG_PATH / "moviepilot.log",
|
||||
}
|
||||
return [f" - {label}: `{path}`" for label, path in paths.items()]
|
||||
|
||||
|
||||
@@ -32,6 +32,10 @@ def build_manual_redo_template_context(history: Any) -> dict[str, int | str]:
|
||||
"transfer_mode": history.mode or "unknown",
|
||||
"tmdbid": history.tmdbid or "none",
|
||||
"doubanid": history.doubanid or "none",
|
||||
"bangumiid": history.bangumiid or "none",
|
||||
"anilistid": history.anilistid or "none",
|
||||
"media_source": history.media_source or "none",
|
||||
"media_id": history.media_id or "none",
|
||||
"error_message": history.errmsg or "none",
|
||||
}
|
||||
|
||||
@@ -55,6 +59,10 @@ def format_manual_redo_record_context(history: Any) -> str:
|
||||
f"- Transfer mode: {context['transfer_mode']}",
|
||||
f"- Current TMDB ID: {context['tmdbid']}",
|
||||
f"- Current Douban ID: {context['doubanid']}",
|
||||
f"- Current Bangumi ID: {context['bangumiid']}",
|
||||
f"- Current AniList ID: {context['anilistid']}",
|
||||
f"- Current media source: {context['media_source']}",
|
||||
f"- Current source-native ID: {context['media_id']}",
|
||||
f"- Error message: {context['error_message']}",
|
||||
]
|
||||
)
|
||||
|
||||
+36
-8
@@ -5,6 +5,7 @@ from __future__ import annotations
|
||||
import re
|
||||
import shutil
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable, Optional
|
||||
@@ -243,9 +244,15 @@ class AgentRuntimeManager:
|
||||
self._cache_lock = threading.Lock()
|
||||
self._cached_signature: Optional[tuple[tuple[str, int, int], ...]] = None
|
||||
self._cached_config: Optional[AgentRuntimeConfig] = None
|
||||
self._cached_signature_checked_at = 0.0
|
||||
self._signature_check_interval = 1.0
|
||||
self._layout_ready = False
|
||||
|
||||
def ensure_layout(self) -> None:
|
||||
"""创建目录、同步默认文件,并清理废弃的旧版 runtime 文件。"""
|
||||
with self._cache_lock:
|
||||
if self._layout_ready:
|
||||
return
|
||||
self.agent_root_dir.mkdir(parents=True, exist_ok=True)
|
||||
self.runtime_dir.mkdir(parents=True, exist_ok=True)
|
||||
self.memory_dir.mkdir(parents=True, exist_ok=True)
|
||||
@@ -257,11 +264,13 @@ class AgentRuntimeManager:
|
||||
self._remove_obsolete_runtime_files()
|
||||
self._sync_bundled_defaults()
|
||||
self._migrate_root_memory_files()
|
||||
with self._cache_lock:
|
||||
self._layout_ready = True
|
||||
|
||||
def load_runtime_config(self) -> AgentRuntimeConfig:
|
||||
"""加载配置。用户目录损坏时自动回退到内置默认配置。"""
|
||||
self.ensure_layout()
|
||||
signature = self._build_signature()
|
||||
signature = self.current_signature()
|
||||
with self._cache_lock:
|
||||
if self._cached_signature == signature and self._cached_config:
|
||||
return self._cached_config
|
||||
@@ -269,7 +278,7 @@ class AgentRuntimeManager:
|
||||
try:
|
||||
config = self._load_from_root(self.runtime_dir)
|
||||
except AgentRuntimeConfigError as err:
|
||||
logger.warning("Agent 根层配置无效,回退到内置默认配置: %s", err)
|
||||
logger.warning(f"Agent 根层配置无效,回退到内置默认配置: {err}")
|
||||
config = self._load_from_root(self.bundled_defaults_dir)
|
||||
config.used_fallback = True
|
||||
config.warnings.insert(
|
||||
@@ -285,6 +294,25 @@ class AgentRuntimeManager:
|
||||
with self._cache_lock:
|
||||
self._cached_signature = None
|
||||
self._cached_config = None
|
||||
self._cached_signature_checked_at = 0.0
|
||||
self._layout_ready = False
|
||||
|
||||
def current_signature(self) -> tuple[tuple[str, int, int], ...]:
|
||||
"""返回当前运行时配置文件签名,供调用方判断缓存是否仍可复用。"""
|
||||
now = time.monotonic()
|
||||
with self._cache_lock:
|
||||
if (
|
||||
self._cached_signature is not None
|
||||
and now - self._cached_signature_checked_at
|
||||
< self._signature_check_interval
|
||||
):
|
||||
return self._cached_signature
|
||||
|
||||
signature = self._build_signature()
|
||||
with self._cache_lock:
|
||||
self._cached_signature = signature
|
||||
self._cached_signature_checked_at = now
|
||||
return signature
|
||||
|
||||
def set_active_persona(self, persona_query: str) -> AgentRuntimeConfig:
|
||||
"""切换当前激活人格,并立即刷新缓存。"""
|
||||
@@ -308,7 +336,7 @@ class AgentRuntimeManager:
|
||||
)
|
||||
current_path.write_text(document, encoding="utf-8")
|
||||
self.invalidate_cache()
|
||||
logger.info("已切换 Agent 人格: %s", persona.persona_id)
|
||||
logger.info(f"已切换 Agent 人格: {persona.persona_id}")
|
||||
return self.load_runtime_config()
|
||||
|
||||
def list_personas(self) -> list[PersonaDefinition]:
|
||||
@@ -439,7 +467,7 @@ class AgentRuntimeManager:
|
||||
continue
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
shutil.copy2(path, target)
|
||||
logger.info("已同步默认 Agent 运行时文件: %s", target)
|
||||
logger.info(f"已同步默认 Agent 运行时文件: {target}")
|
||||
|
||||
@classmethod
|
||||
def _should_update_bundled_subagent(
|
||||
@@ -478,7 +506,7 @@ class AgentRuntimeManager:
|
||||
return
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
source.rename(target)
|
||||
logger.info("已迁移旧版 Agent 根配置文件: %s -> %s", source, target)
|
||||
logger.info(f"已迁移旧版 Agent 根配置文件: {source} -> {target}")
|
||||
|
||||
def _remove_obsolete_runtime_files(self) -> None:
|
||||
"""删除不再支持的旧版 Agent 配置文件,避免被误迁移到 memory。"""
|
||||
@@ -487,14 +515,14 @@ class AgentRuntimeManager:
|
||||
if not path.exists() or not path.is_file():
|
||||
continue
|
||||
path.unlink()
|
||||
logger.info("已删除废弃的 Agent 根配置文件: %s", path)
|
||||
logger.info(f"已删除废弃的 Agent 根配置文件: {path}")
|
||||
|
||||
for relative_path in sorted(OBSOLETE_RUNTIME_FILES):
|
||||
path = self.runtime_dir / relative_path
|
||||
if not path.exists() or not path.is_file():
|
||||
continue
|
||||
path.unlink()
|
||||
logger.info("已删除废弃的 Agent 运行时文件: %s", path)
|
||||
logger.info(f"已删除废弃的 Agent 运行时文件: {path}")
|
||||
|
||||
def _migrate_root_memory_files(self) -> None:
|
||||
"""将旧版根目录 memory 文件移入 `config/agent/memory`。"""
|
||||
@@ -505,7 +533,7 @@ class AgentRuntimeManager:
|
||||
if target.exists():
|
||||
continue
|
||||
path.rename(target)
|
||||
logger.info("已迁移旧版 Agent memory 文件: %s -> %s", path, target)
|
||||
logger.info(f"已迁移旧版 Agent memory 文件: {path} -> {target}")
|
||||
|
||||
def _load_from_root(self, root: Path) -> AgentRuntimeConfig:
|
||||
current_persona_path = root / CURRENT_PERSONA_FILE
|
||||
|
||||
+25
-16
@@ -238,10 +238,6 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
||||
|
||||
# 获取工具执行提示消息
|
||||
tool_message = self.get_tool_message(**kwargs)
|
||||
if not tool_message:
|
||||
explanation = kwargs.get("explanation")
|
||||
if explanation:
|
||||
tool_message = explanation
|
||||
|
||||
# 发送工具执行过程消息(流式传输且非最后终结工具时)
|
||||
if self._stream_handler and self._stream_handler.is_streaming and not self.return_direct:
|
||||
@@ -325,16 +321,13 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
||||
获取工具执行时的友好提示消息。
|
||||
|
||||
子类可以重写此方法,根据实际参数生成个性化的提示消息。
|
||||
如果返回 None 或空字符串,将回退使用 explanation 参数。
|
||||
|
||||
Args:
|
||||
**kwargs: 工具的所有参数(包括 explanation)
|
||||
**kwargs: 工具的所有参数
|
||||
|
||||
Returns:
|
||||
str: 友好的提示消息,如果返回 None 或空字符串则使用 explanation
|
||||
str: 友好的提示消息
|
||||
"""
|
||||
explanation = kwargs.get("explanation")
|
||||
return str(explanation) if explanation else None
|
||||
return None
|
||||
|
||||
@abstractmethod
|
||||
async def run(self, **kwargs) -> str:
|
||||
@@ -431,7 +424,7 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
||||
:return: 普通用户允许读写的本地目录列表
|
||||
"""
|
||||
roots = [
|
||||
settings.CONFIG_PATH / "agent"
|
||||
settings.CONFIG_PATH / "agent",
|
||||
]
|
||||
resolved_roots = []
|
||||
for root in roots:
|
||||
@@ -467,7 +460,7 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
||||
allowed_text = "、".join(str(root) for root in allowed_roots)
|
||||
return (
|
||||
resolved_path,
|
||||
f"抱歉,普通用户只能{operation}配置目录、Agent记忆目录和日志目录内的文件或目录:{allowed_text}",
|
||||
f"抱歉,普通用户只能{operation}Agent配置目录内的文件或目录:{allowed_text}",
|
||||
)
|
||||
|
||||
async def _check_local_storage_access(
|
||||
@@ -489,7 +482,7 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
||||
return None, None
|
||||
return (
|
||||
None,
|
||||
f"抱歉,普通用户只能{operation}本地配置目录、Agent记忆目录和日志目录,不能访问远程存储。",
|
||||
f"抱歉,普通用户只能{operation}本地Agent配置目录,不能访问远程存储。",
|
||||
)
|
||||
|
||||
return await self._check_local_file_access(path=path, operation=operation)
|
||||
@@ -515,8 +508,8 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
||||
return (
|
||||
"抱歉,您没有执行此工具的权限。"
|
||||
"只有渠道管理员或系统管理员才能执行工具操作。"
|
||||
"如需执行工具,请联系渠道管理员将您的用户ID添加到渠道管理员列表中,"
|
||||
"或联系系统管理员为您设置权限。"
|
||||
"如需执行工具,请联系管理员将您的用户ID添加到渠道管理员列表中(设定 -> 通知 -> 对应渠道配置 -> 管理员名单),"
|
||||
"或联系系统管理员为您设置管理员权限。"
|
||||
)
|
||||
|
||||
async def _has_channel_admin_permission(self) -> bool:
|
||||
@@ -627,7 +620,8 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
||||
发送工具通知消息。
|
||||
|
||||
WebAgent 渠道没有后端模块实例,前端流式面板通过 Agent 上下文中的
|
||||
回调直接接收通知;其它渠道继续走统一消息链。
|
||||
回调直接接收通知;无渠道的后台任务清空渠道侧定位信息后交由消息链广播,
|
||||
其它渠道继续走统一消息链。
|
||||
"""
|
||||
callback = self._agent_context.get("notification_callback")
|
||||
if (
|
||||
@@ -637,6 +631,20 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
||||
callback(notification)
|
||||
return
|
||||
|
||||
if not self._channel or not self._source:
|
||||
notification = notification.model_copy(
|
||||
update={
|
||||
"channel": None,
|
||||
"source": None,
|
||||
"userid": None,
|
||||
"username": notification.username
|
||||
or self._username
|
||||
or settings.SUPERUSER,
|
||||
"original_message_id": None,
|
||||
"original_chat_id": None,
|
||||
}
|
||||
)
|
||||
|
||||
await ToolChain().async_post_message(notification)
|
||||
|
||||
async def send_tool_message(
|
||||
@@ -655,5 +663,6 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
||||
title=title,
|
||||
text=message,
|
||||
image=image,
|
||||
save_history=False,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -42,8 +42,13 @@ from app.agent.tools.impl.send_message import SendMessageTool
|
||||
from app.agent.tools.impl.ask_user_choice import AskUserChoiceTool
|
||||
from app.agent.tools.impl.send_local_file import SendLocalFileTool
|
||||
from app.agent.tools.impl.send_voice_message import SendVoiceMessageTool
|
||||
from app.agent.tools.impl.create_agent_task import CreateAgentTaskTool
|
||||
from app.agent.tools.impl.delete_agent_task import DeleteAgentTaskTool
|
||||
from app.agent.tools.impl.query_agent_tasks import QueryAgentTasksTool
|
||||
from app.agent.tools.impl.query_schedulers import QuerySchedulersTool
|
||||
from app.agent.tools.impl.run_agent_task import RunAgentTaskTool
|
||||
from app.agent.tools.impl.run_scheduler import RunSchedulerTool
|
||||
from app.agent.tools.impl.update_agent_task import UpdateAgentTaskTool
|
||||
from app.agent.tools.impl.query_workflows import QueryWorkflowsTool
|
||||
from app.agent.tools.impl.run_workflow import RunWorkflowTool
|
||||
from app.agent.tools.impl.query_personas import QueryPersonasTool
|
||||
@@ -141,6 +146,11 @@ class MoviePilotToolFactory:
|
||||
QueryTransferHistoryTool,
|
||||
TransferFileTool,
|
||||
SendMessageTool,
|
||||
CreateAgentTaskTool,
|
||||
QueryAgentTasksTool,
|
||||
UpdateAgentTaskTool,
|
||||
RunAgentTaskTool,
|
||||
DeleteAgentTaskTool,
|
||||
QuerySchedulersTool,
|
||||
RunSchedulerTool,
|
||||
QueryWorkflowsTool,
|
||||
@@ -181,6 +191,8 @@ class MoviePilotToolFactory:
|
||||
"edit_file",
|
||||
"execute_command",
|
||||
"ask_user_choice",
|
||||
"create_agent_task",
|
||||
"query_agent_tasks",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
"""Agent 命令工具的安全校验逻辑。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os.path
|
||||
import re
|
||||
import shlex
|
||||
|
||||
|
||||
COMMAND_FORBIDDEN_KEYWORDS = (
|
||||
":(){ :|:& };:",
|
||||
"dd if=/dev/zero",
|
||||
"mkfs",
|
||||
"reboot",
|
||||
"shutdown",
|
||||
)
|
||||
|
||||
COMMAND_DANGEROUS_PATTERNS = (
|
||||
re.compile(r"\brm\s+[^;&|]*-[^\s;&|]*[rR][fF]?[^\s;&|]*\s+/(?:\s|$|[;&|])"),
|
||||
re.compile(r"\bdd\s+[^;&|]*(?:of=/dev/(?:sd[a-z]\d*|nvme\d+n\d+p?\d*|disk\d+)|if=/dev/zero)"),
|
||||
re.compile(r"\b(?:mkfs|fdisk|parted|diskutil)\b"),
|
||||
re.compile(r"\b(?:chmod|chown)\s+[^;&|]*-R[^;&|]*\s+/(?:\s|$|[;&|])"),
|
||||
re.compile(r"\b(?:reboot|shutdown|halt|poweroff)\b"),
|
||||
)
|
||||
|
||||
|
||||
def _command_tokens(command: str) -> list[str]:
|
||||
"""尽力解析 shell 命令 token,解析失败时退回空白分割。"""
|
||||
try:
|
||||
return shlex.split(command, posix=True)
|
||||
except ValueError:
|
||||
return re.split(r"\s+", command.strip())
|
||||
|
||||
|
||||
def _contains_recursive_root_delete(command: str) -> bool:
|
||||
"""识别递归删除根目录或一级目录的 rm 命令。"""
|
||||
tokens = _command_tokens(command)
|
||||
if not any(token == "rm" or token.endswith("/rm") for token in tokens):
|
||||
return False
|
||||
has_recursive = any(
|
||||
token.startswith("-") and ("r" in token or "R" in token)
|
||||
for token in tokens
|
||||
)
|
||||
if not has_recursive:
|
||||
return False
|
||||
|
||||
for token in tokens:
|
||||
clean_token = re.match(r"^([^;|&><]+)", token)
|
||||
if not clean_token:
|
||||
continue
|
||||
path_value = clean_token.group(1).strip("\"'")
|
||||
if not path_value.startswith("/"):
|
||||
continue
|
||||
norm_path = os.path.normpath(path_value)
|
||||
if norm_path == "/" or re.match(r"^/[^/]+$", norm_path):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def detect_dangerous_command(command: str) -> str:
|
||||
"""返回危险命令原因,安全时返回空字符串。"""
|
||||
normalized = str(command or "").strip()
|
||||
if not normalized:
|
||||
return "命令不能为空"
|
||||
for keyword in COMMAND_FORBIDDEN_KEYWORDS:
|
||||
if keyword in normalized:
|
||||
return f"命令包含禁止使用的关键字 '{keyword}'"
|
||||
if _contains_recursive_root_delete(normalized):
|
||||
return "命令疑似递归删除根目录或一级目录"
|
||||
for pattern in COMMAND_DANGEROUS_PATTERNS:
|
||||
if pattern.search(normalized):
|
||||
return "命令匹配高危系统操作模式"
|
||||
return ""
|
||||
|
||||
|
||||
def validate_command_safety(command: str, *, confirmed: bool = False) -> None:
|
||||
"""
|
||||
校验 shell 命令安全性。
|
||||
|
||||
:param command: 待执行命令
|
||||
:param confirmed: 是否已经通过显式参数确认高危操作
|
||||
"""
|
||||
reason = detect_dangerous_command(command)
|
||||
if not reason:
|
||||
return
|
||||
if confirmed and reason != "命令不能为空":
|
||||
return
|
||||
raise ValueError(f"{reason}。如确认需要执行,请设置 confirm_dangerous=true")
|
||||
@@ -0,0 +1,53 @@
|
||||
"""Agent 文件写入工具的共享辅助函数。"""
|
||||
|
||||
import hashlib
|
||||
import os
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class FileVersionConflictError(RuntimeError):
|
||||
"""目标文件在准备写入期间发生变化。"""
|
||||
|
||||
|
||||
def calculate_file_sha256(path: Path) -> str:
|
||||
"""计算文件原始字节的 SHA-256,用于检测陈旧写入。"""
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as file_handle:
|
||||
for chunk in iter(lambda: file_handle.read(64 * 1024), b""):
|
||||
digest.update(chunk)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def atomic_write_text(
|
||||
path: Path,
|
||||
content: str,
|
||||
expected_sha256: str | None = None,
|
||||
) -> None:
|
||||
"""校验目标版本后,在同目录写入临时文件并原子替换文本。"""
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
descriptor, temp_name = tempfile.mkstemp(
|
||||
dir=path.parent,
|
||||
prefix=f".{path.name}.",
|
||||
suffix=".tmp",
|
||||
)
|
||||
temp_path = Path(temp_name)
|
||||
try:
|
||||
with os.fdopen(descriptor, "w", encoding="utf-8", newline="") as file_handle:
|
||||
file_handle.write(content)
|
||||
file_handle.flush()
|
||||
os.fsync(file_handle.fileno())
|
||||
|
||||
if expected_sha256:
|
||||
if (
|
||||
not path.is_file()
|
||||
or calculate_file_sha256(path).casefold()
|
||||
!= expected_sha256.casefold()
|
||||
):
|
||||
raise FileVersionConflictError(str(path))
|
||||
if path.exists():
|
||||
os.chmod(temp_path, path.stat().st_mode)
|
||||
os.replace(temp_path, path)
|
||||
finally:
|
||||
if temp_path.exists():
|
||||
temp_path.unlink()
|
||||
@@ -5,8 +5,7 @@ import re
|
||||
from typing import Any, Dict, Iterable, Optional
|
||||
|
||||
from app.core.event import eventmanager
|
||||
from app.db import AsyncSessionFactory
|
||||
from app.db.models.subscribe import Subscribe
|
||||
from app.db.subscribe_oper import SubscribeOper
|
||||
from app.db.systemconfig_oper import SystemConfigOper
|
||||
from app.helper.rule import RuleHelper
|
||||
from app.modules.filter.RuleParser import RuleParser
|
||||
@@ -284,23 +283,22 @@ async def collect_rule_group_usages(
|
||||
continue
|
||||
ensure_usage(name)["used_in_global_best_version"] = True
|
||||
|
||||
async with AsyncSessionFactory() as db:
|
||||
subscribes = await Subscribe.async_list(db)
|
||||
for subscribe in subscribes:
|
||||
filter_groups = subscribe.filter_groups or []
|
||||
for name in filter_groups:
|
||||
if target_names and name not in target_names:
|
||||
continue
|
||||
ensure_usage(name)["subscribes"].append(
|
||||
{
|
||||
"subscribe_id": subscribe.id,
|
||||
"name": subscribe.name,
|
||||
"season": subscribe.season,
|
||||
"type": subscribe.type,
|
||||
"username": subscribe.username,
|
||||
"best_version": bool(subscribe.best_version),
|
||||
}
|
||||
)
|
||||
subscribes = await SubscribeOper().async_list()
|
||||
for subscribe in subscribes:
|
||||
filter_groups = subscribe.filter_groups or []
|
||||
for name in filter_groups:
|
||||
if target_names and name not in target_names:
|
||||
continue
|
||||
ensure_usage(name)["subscribes"].append(
|
||||
{
|
||||
"subscribe_id": subscribe.id,
|
||||
"name": subscribe.name,
|
||||
"season": subscribe.season,
|
||||
"type": subscribe.type,
|
||||
"username": subscribe.username,
|
||||
"best_version": bool(subscribe.best_version),
|
||||
}
|
||||
)
|
||||
|
||||
return usage_map
|
||||
|
||||
@@ -482,22 +480,22 @@ async def rename_rule_group_references(old_name: str, new_name: str) -> dict:
|
||||
await save_system_config(config_key, updated)
|
||||
changed["global_settings"][config_key.value] = updated
|
||||
|
||||
async with AsyncSessionFactory() as db:
|
||||
subscribes = await Subscribe.async_list(db)
|
||||
for subscribe in subscribes:
|
||||
original = subscribe.filter_groups or []
|
||||
updated = replace_group_name_in_list(original, old_name, new_name)
|
||||
if updated == original:
|
||||
continue
|
||||
await subscribe.async_update(db, {"filter_groups": updated})
|
||||
changed["subscribes"].append(
|
||||
{
|
||||
"subscribe_id": subscribe.id,
|
||||
"name": subscribe.name,
|
||||
"season": subscribe.season,
|
||||
"filter_groups": updated,
|
||||
}
|
||||
)
|
||||
subscribe_oper = SubscribeOper()
|
||||
subscribes = await subscribe_oper.async_list()
|
||||
for subscribe in subscribes:
|
||||
original = subscribe.filter_groups or []
|
||||
updated = replace_group_name_in_list(original, old_name, new_name)
|
||||
if updated == original:
|
||||
continue
|
||||
await subscribe_oper.async_update_filter_groups(subscribe.id, updated)
|
||||
changed["subscribes"].append(
|
||||
{
|
||||
"subscribe_id": subscribe.id,
|
||||
"name": subscribe.name,
|
||||
"season": subscribe.season,
|
||||
"filter_groups": updated,
|
||||
}
|
||||
)
|
||||
|
||||
return changed
|
||||
|
||||
@@ -520,21 +518,21 @@ async def remove_rule_group_references(group_name: str) -> dict:
|
||||
await save_system_config(config_key, updated)
|
||||
changed["global_settings"][config_key.value] = updated
|
||||
|
||||
async with AsyncSessionFactory() as db:
|
||||
subscribes = await Subscribe.async_list(db)
|
||||
for subscribe in subscribes:
|
||||
original = subscribe.filter_groups or []
|
||||
updated = [value for value in original if value != group_name]
|
||||
if updated == original:
|
||||
continue
|
||||
await subscribe.async_update(db, {"filter_groups": updated})
|
||||
changed["subscribes"].append(
|
||||
{
|
||||
"subscribe_id": subscribe.id,
|
||||
"name": subscribe.name,
|
||||
"season": subscribe.season,
|
||||
"filter_groups": updated,
|
||||
}
|
||||
)
|
||||
subscribe_oper = SubscribeOper()
|
||||
subscribes = await subscribe_oper.async_list()
|
||||
for subscribe in subscribes:
|
||||
original = subscribe.filter_groups or []
|
||||
updated = [value for value in original if value != group_name]
|
||||
if updated == original:
|
||||
continue
|
||||
await subscribe_oper.async_update_filter_groups(subscribe.id, updated)
|
||||
changed["subscribes"].append(
|
||||
{
|
||||
"subscribe_id": subscribe.id,
|
||||
"name": subscribe.name,
|
||||
"season": subscribe.season,
|
||||
"filter_groups": updated,
|
||||
}
|
||||
)
|
||||
|
||||
return changed
|
||||
|
||||
@@ -103,6 +103,79 @@ def summarize_plugin(plugin: Any) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _merge_plugin_source_metadata(plugin: Any, source_plugin: Any) -> Any:
|
||||
"""
|
||||
将插件市场或本地仓库中的来源元数据合并到已安装插件对象。
|
||||
"""
|
||||
repo_url = getattr(source_plugin, "repo_url", None)
|
||||
if repo_url:
|
||||
setattr(plugin, "repo_url", repo_url)
|
||||
|
||||
for attr in (
|
||||
"has_update",
|
||||
"release",
|
||||
"system_version",
|
||||
"system_version_compatible",
|
||||
"system_version_message",
|
||||
):
|
||||
value = getattr(source_plugin, attr, None)
|
||||
if value is not None:
|
||||
setattr(plugin, attr, value)
|
||||
|
||||
return plugin
|
||||
|
||||
|
||||
def _map_plugins_by_id(plugins: list[Any]) -> dict[str, Any]:
|
||||
"""
|
||||
按插件 ID 建立稳定映射,保留同 ID 首个候选来源。
|
||||
"""
|
||||
plugin_map: dict[str, Any] = {}
|
||||
for plugin in plugins:
|
||||
plugin_id = getattr(plugin, "id", None)
|
||||
if plugin_id and plugin_id not in plugin_map:
|
||||
plugin_map[plugin_id] = plugin
|
||||
return plugin_map
|
||||
|
||||
|
||||
async def enrich_installed_plugin_sources(
|
||||
installed_plugins: list[Any],
|
||||
force_refresh: bool = False,
|
||||
) -> list[Any]:
|
||||
"""
|
||||
为已安装插件补齐安装来源仓库地址。
|
||||
|
||||
本地插件对象只包含运行目录中的静态元数据,通常没有 repo_url。这里按需从
|
||||
本地插件仓库和插件市场补齐来源,保证 Agent 后续安装、升级判断可以拿到仓库地址。
|
||||
"""
|
||||
missing_source_plugins = [
|
||||
plugin for plugin in installed_plugins if not getattr(plugin, "repo_url", None)
|
||||
]
|
||||
if not missing_source_plugins:
|
||||
return installed_plugins
|
||||
|
||||
plugin_manager = PluginManager()
|
||||
local_repo_map = _map_plugins_by_id(plugin_manager.get_local_repo_plugins())
|
||||
for plugin in missing_source_plugins:
|
||||
source_plugin = local_repo_map.get(getattr(plugin, "id", None))
|
||||
if source_plugin:
|
||||
_merge_plugin_source_metadata(plugin, source_plugin)
|
||||
|
||||
missing_source_plugins = [
|
||||
plugin for plugin in installed_plugins if not getattr(plugin, "repo_url", None)
|
||||
]
|
||||
if not missing_source_plugins:
|
||||
return installed_plugins
|
||||
|
||||
market_plugins = await plugin_manager.async_get_online_plugins(force=force_refresh)
|
||||
market_map = _map_plugins_by_id(market_plugins or [])
|
||||
for plugin in missing_source_plugins:
|
||||
source_plugin = market_map.get(getattr(plugin, "id", None))
|
||||
if source_plugin:
|
||||
_merge_plugin_source_metadata(plugin, source_plugin)
|
||||
|
||||
return installed_plugins
|
||||
|
||||
|
||||
async def load_market_plugins(force_refresh: bool = False) -> list[Any]:
|
||||
"""
|
||||
聚合插件市场与本地插件仓库中的候选插件。
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""系统设置工具共用的键解析与分组元数据。"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
from typing import Any, Optional
|
||||
|
||||
from app.core.config import Settings
|
||||
from app.schemas.types import SystemConfigKey
|
||||
@@ -15,6 +15,7 @@ class SettingSpec:
|
||||
source: str
|
||||
group: str
|
||||
label: str
|
||||
systemconfig_key: Optional[SystemConfigKey] = None
|
||||
|
||||
|
||||
SYSTEMCONFIG_SETTING_METADATA = {
|
||||
@@ -58,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": "自定义识别词",
|
||||
@@ -234,6 +239,7 @@ def _build_specs() -> tuple[dict[str, SettingSpec], dict[str, SettingSpec]]:
|
||||
source="systemconfig",
|
||||
group=metadata.get("group", "misc"),
|
||||
label=metadata.get("label", item.value),
|
||||
systemconfig_key=item,
|
||||
)
|
||||
return core_specs, system_specs
|
||||
|
||||
@@ -333,3 +339,57 @@ def list_setting_specs(
|
||||
|
||||
def get_default_list_match_field(setting_key: str) -> Optional[str]:
|
||||
return LIST_ITEM_MATCH_FIELD_DEFAULTS.get(setting_key)
|
||||
|
||||
|
||||
SECRET_KEYWORDS = (
|
||||
"api_key",
|
||||
"apikey",
|
||||
"token",
|
||||
"secret",
|
||||
"password",
|
||||
"passwd",
|
||||
"cookie",
|
||||
"authorization",
|
||||
"refresh_token",
|
||||
"access_token",
|
||||
)
|
||||
|
||||
|
||||
def is_secret_setting_key(key: str) -> bool:
|
||||
"""判断设置键名是否疑似敏感字段。"""
|
||||
normalized = _normalize_token(key)
|
||||
return any(keyword in normalized for keyword in SECRET_KEYWORDS)
|
||||
|
||||
|
||||
def redact_secret_value(value: Any, *, redact_scalar: bool = False) -> Any:
|
||||
"""递归脱敏配置值中的密钥、Cookie、Token 等敏感字段。"""
|
||||
if isinstance(value, dict):
|
||||
return {
|
||||
key: "***"
|
||||
if is_secret_setting_key(str(key))
|
||||
else redact_secret_value(item, redact_scalar=redact_scalar)
|
||||
for key, item in value.items()
|
||||
}
|
||||
if isinstance(value, list):
|
||||
return [
|
||||
redact_secret_value(item, redact_scalar=redact_scalar)
|
||||
for item in value
|
||||
]
|
||||
if isinstance(value, str):
|
||||
return "***" if value and redact_scalar else value
|
||||
return value
|
||||
|
||||
|
||||
def should_redact_setting(spec: SettingSpec, value: Any) -> bool:
|
||||
"""判断某项设置在默认查询响应中是否需要脱敏。"""
|
||||
if is_secret_setting_key(spec.key):
|
||||
return True
|
||||
if isinstance(value, dict):
|
||||
return any(is_secret_setting_key(str(key)) for key in value.keys())
|
||||
if isinstance(value, list):
|
||||
return any(
|
||||
should_redact_setting(spec, item)
|
||||
for item in value
|
||||
if isinstance(item, dict)
|
||||
)
|
||||
return False
|
||||
|
||||
@@ -13,6 +13,7 @@ from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
from app.agent.tools.impl._command_safety import validate_command_safety
|
||||
from app.core.config import settings
|
||||
from app.log import logger
|
||||
|
||||
@@ -34,14 +35,6 @@ TERMINAL_PTY_POLL_INTERVAL = 0.05
|
||||
TERMINAL_WAIT_DEFAULT_MS = 1000
|
||||
TERMINAL_WAIT_MAX_MS = 60 * 1000
|
||||
TERMINAL_KILL_GRACE_SECONDS = 3
|
||||
TERMINAL_FORBIDDEN_KEYWORDS = (
|
||||
"rm -rf /",
|
||||
":(){ :|:& };:",
|
||||
"dd if=/dev/zero",
|
||||
"mkfs",
|
||||
"reboot",
|
||||
"shutdown",
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -176,13 +169,9 @@ class _TerminalSessionManager:
|
||||
return merged_env
|
||||
|
||||
@staticmethod
|
||||
def _validate_command(command: str) -> None:
|
||||
def _validate_command(command: str, *, confirmed: bool = False) -> None:
|
||||
"""拒绝明显危险或空白命令。"""
|
||||
if not command or not command.strip():
|
||||
raise ValueError("命令不能为空")
|
||||
for keyword in TERMINAL_FORBIDDEN_KEYWORDS:
|
||||
if keyword in command:
|
||||
raise ValueError(f"命令包含禁止使用的关键字 '{keyword}'")
|
||||
validate_command_safety(command, confirmed=confirmed)
|
||||
|
||||
@staticmethod
|
||||
def _set_nonblocking(fd: int) -> None:
|
||||
@@ -213,9 +202,10 @@ class _TerminalSessionManager:
|
||||
cwd: Optional[str] = None,
|
||||
env: Optional[dict[str, Any]] = None,
|
||||
use_pty: Any = True,
|
||||
confirm_dangerous: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
"""启动后台命令并立即返回会话 ID。"""
|
||||
self._validate_command(command)
|
||||
self._validate_command(command, confirmed=confirm_dangerous)
|
||||
normalized_cwd = self._normalize_cwd(cwd)
|
||||
normalized_env = self._build_env(env)
|
||||
should_use_pty = self._normalize_bool(use_pty, default=True) and os.name == "posix"
|
||||
@@ -313,7 +303,10 @@ class _TerminalSessionManager:
|
||||
continue
|
||||
except OSError as err:
|
||||
if err.errno not in {errno.EIO, errno.EBADF}:
|
||||
logger.debug("PTY 输出读取异常: session_id=%s, error=%s", session.session_id, err)
|
||||
logger.debug(
|
||||
f"PTY 输出读取异常: session_id={session.session_id}, "
|
||||
f"error={err}"
|
||||
)
|
||||
break
|
||||
|
||||
if not data:
|
||||
@@ -343,7 +336,9 @@ class _TerminalSessionManager:
|
||||
session.mark_finished(session.exit_code)
|
||||
except Exception as err:
|
||||
session.mark_error(str(err))
|
||||
logger.warning("等待 PTY 进程失败: session_id=%s, error=%s", session.session_id, err)
|
||||
logger.warning(
|
||||
f"等待 PTY 进程失败: session_id={session.session_id}, error={err}"
|
||||
)
|
||||
finally:
|
||||
await self._finish_reader_tasks(session)
|
||||
session.close_pty()
|
||||
@@ -358,7 +353,9 @@ class _TerminalSessionManager:
|
||||
session.mark_finished(exit_code)
|
||||
except Exception as err:
|
||||
session.mark_error(str(err))
|
||||
logger.warning("等待管道进程失败: session_id=%s, error=%s", session.session_id, err)
|
||||
logger.warning(
|
||||
f"等待管道进程失败: session_id={session.session_id}, error={err}"
|
||||
)
|
||||
finally:
|
||||
await self._finish_reader_tasks(session)
|
||||
|
||||
|
||||
@@ -127,8 +127,19 @@ def filter_contexts(items: List[Context],
|
||||
return filtered_items
|
||||
|
||||
|
||||
def simplify_search_result(context: Context, index: int) -> dict:
|
||||
"""精简单条搜索结果"""
|
||||
def simplify_search_result(
|
||||
context: Context,
|
||||
index: int,
|
||||
include_description: bool = False,
|
||||
) -> dict:
|
||||
"""
|
||||
精简单条搜索结果
|
||||
|
||||
:param context: 搜索结果上下文
|
||||
:param index: 搜索结果在原始缓存中的序号
|
||||
:param include_description: 是否返回种子简介
|
||||
:return: 精简后的搜索结果
|
||||
"""
|
||||
simplified = {}
|
||||
torrent_info = context.torrent_info
|
||||
meta_info = context.meta_info
|
||||
@@ -147,6 +158,8 @@ def simplify_search_result(context: Context, index: int) -> dict:
|
||||
"freedate_diff": torrent_info.freedate_diff,
|
||||
"pubdate": torrent_info.pubdate,
|
||||
}
|
||||
if include_description:
|
||||
simplified["torrent_info"]["description"] = torrent_info.description
|
||||
|
||||
if media_info:
|
||||
simplified["media_info"] = {
|
||||
|
||||
@@ -20,8 +20,6 @@ from app.schemas.types import SystemConfigKey
|
||||
class AddCustomFilterRuleInput(BaseModel):
|
||||
"""新增自定义过滤规则工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
rule_id: str = Field(
|
||||
...,
|
||||
description="Unique custom rule ID. Only letters and numbers are allowed.",
|
||||
|
||||
@@ -15,7 +15,7 @@ from app.core.config import settings
|
||||
from app.core.context import Context
|
||||
from app.core.metainfo import MetaInfo
|
||||
from app.db.site_oper import SiteOper
|
||||
from app.helper.directory import DirectoryHelper
|
||||
from app.helper.directory import DirectoryHelper, validate_download_save_path
|
||||
from app.log import logger
|
||||
from app.schemas import FileURI, TorrentInfo
|
||||
from app.utils.crypto import HashUtils
|
||||
@@ -24,7 +24,6 @@ from app.utils.crypto import HashUtils
|
||||
class AddDownloadTasksInput(BaseModel):
|
||||
"""添加下载任务工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
torrent_url: List[str] = Field(
|
||||
...,
|
||||
description="One or more torrent_url values. Supports refs from get_search_results (`hash:id`) and magnet links."
|
||||
@@ -184,8 +183,8 @@ class AddDownloadTasksTool(MoviePilotTool):
|
||||
@staticmethod
|
||||
def _resolve_direct_download_dir(save_path: Optional[str]) -> Optional[Path]:
|
||||
"""解析直接下载使用的目录,优先使用 save_path,其次使用默认下载目录"""
|
||||
if save_path:
|
||||
return Path(save_path)
|
||||
if save_path is not None:
|
||||
return Path(validate_download_save_path(save_path))
|
||||
|
||||
download_dirs = DirectoryHelper().get_download_dirs()
|
||||
if not download_dirs:
|
||||
@@ -226,6 +225,8 @@ class AddDownloadTasksTool(MoviePilotTool):
|
||||
merged_labels: Optional[str],
|
||||
) -> tuple[Optional[str], Optional[str]]:
|
||||
"""同步提交带上下文的下载任务,避免站点下载与下载器调用阻塞事件循环。"""
|
||||
if save_path is not None:
|
||||
save_path = validate_download_save_path(save_path)
|
||||
return DownloadChain().download_single(
|
||||
context=context,
|
||||
downloader=downloader,
|
||||
@@ -246,6 +247,12 @@ class AddDownloadTasksTool(MoviePilotTool):
|
||||
if not torrent_inputs:
|
||||
return "错误:torrent_url 不能为空。"
|
||||
|
||||
if save_path is not None:
|
||||
try:
|
||||
save_path = validate_download_save_path(save_path)
|
||||
except ValueError as err:
|
||||
return f"参数错误:save_path {str(err)}"
|
||||
|
||||
merged_labels = self._merge_labels_with_system_tag(labels)
|
||||
success_count = 0
|
||||
failed_messages = []
|
||||
|
||||
@@ -24,8 +24,6 @@ from app.schemas.types import SystemConfigKey
|
||||
class AddRuleGroupInput(BaseModel):
|
||||
"""新增过滤规则组工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
name: str = Field(..., description="New rule group name.")
|
||||
rule_string: str = Field(
|
||||
...,
|
||||
|
||||
@@ -15,8 +15,6 @@ from app.schemas.types import MediaType, MessageChannel
|
||||
class AddSubscribeInput(BaseModel):
|
||||
"""添加订阅工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
title: str = Field(
|
||||
...,
|
||||
description="The title of the media to subscribe to (e.g., 'The Matrix', 'Breaking Bad')",
|
||||
@@ -41,6 +39,10 @@ class AddSubscribeInput(BaseModel):
|
||||
None,
|
||||
description="Douban ID for precise media identification (optional, alternative to tmdb_id)",
|
||||
)
|
||||
bangumi_id: Optional[int] = Field(None, description="Bangumi media ID")
|
||||
anilist_id: Optional[int] = Field(None, description="AniList media ID")
|
||||
media_source: Optional[str] = Field(None, description="Media metadata source")
|
||||
media_id: Optional[str] = Field(None, description="Native ID for media_source")
|
||||
start_episode: Optional[int] = Field(
|
||||
None,
|
||||
description="Starting episode number for TV shows (optional, defaults to 1 if not specified)",
|
||||
@@ -99,7 +101,7 @@ class AddSubscribeTool(MoviePilotTool):
|
||||
message += f" ({year})"
|
||||
if media_type:
|
||||
message += f" [{media_type}]"
|
||||
if season:
|
||||
if season is not None:
|
||||
message += f" 第{season}季"
|
||||
elif media_type == "tv":
|
||||
message += " 第1季(默认)"
|
||||
@@ -146,6 +148,10 @@ class AddSubscribeTool(MoviePilotTool):
|
||||
season: Optional[int] = None,
|
||||
tmdb_id: Optional[int] = None,
|
||||
douban_id: Optional[str] = None,
|
||||
bangumi_id: Optional[int] = None,
|
||||
anilist_id: Optional[int] = None,
|
||||
media_source: Optional[str] = None,
|
||||
media_id: Optional[str] = None,
|
||||
start_episode: Optional[int] = None,
|
||||
total_episode: Optional[int] = None,
|
||||
quality: Optional[str] = None,
|
||||
@@ -199,6 +205,10 @@ class AddSubscribeTool(MoviePilotTool):
|
||||
year=year,
|
||||
tmdbid=tmdb_id,
|
||||
doubanid=douban_id,
|
||||
bangumiid=bangumi_id,
|
||||
anilistid=anilist_id,
|
||||
media_source=media_source,
|
||||
media_id=media_id,
|
||||
season=season,
|
||||
username=subscribe_username,
|
||||
**subscribe_kwargs,
|
||||
|
||||
@@ -27,6 +27,7 @@ class UserChoiceOptionInput(BaseModel):
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_option(self):
|
||||
"""校验按钮选项的文案和值不能为空。"""
|
||||
label = str(self.label)
|
||||
value = str(self.value)
|
||||
if not label.strip():
|
||||
@@ -39,8 +40,6 @@ class UserChoiceOptionInput(BaseModel):
|
||||
class AskUserChoiceInput(BaseModel):
|
||||
"""按钮选择工具输入。"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why the agent needs the user to choose from buttons",)
|
||||
message: str = Field(
|
||||
...,
|
||||
description="Question or prompt shown to the user together with the buttons",
|
||||
@@ -56,6 +55,7 @@ class AskUserChoiceInput(BaseModel):
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_payload(self):
|
||||
"""校验按钮选择工具必须提供问题和选项。"""
|
||||
message = str(self.message)
|
||||
if not message.strip():
|
||||
raise ValueError("message 不能为空")
|
||||
@@ -85,6 +85,7 @@ class AskUserChoiceTool(MoviePilotTool):
|
||||
args_schema: Type[BaseModel] = AskUserChoiceInput
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
"""生成工具执行提示文案。"""
|
||||
message = kwargs.get("message", "") or ""
|
||||
if len(message) > 40:
|
||||
message = message[:40] + "..."
|
||||
@@ -92,6 +93,7 @@ class AskUserChoiceTool(MoviePilotTool):
|
||||
|
||||
@staticmethod
|
||||
def _truncate_button_text(text: str, max_length: int) -> str:
|
||||
"""按渠道限制截断按钮文案。"""
|
||||
if max_length <= 0 or len(text) <= max_length:
|
||||
return text
|
||||
if max_length <= 3:
|
||||
@@ -114,6 +116,14 @@ class AskUserChoiceTool(MoviePilotTool):
|
||||
title: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""
|
||||
发送按钮选择消息,并登记待回调的交互上下文。
|
||||
|
||||
:param message: 展示给用户的问题
|
||||
:param options: 可点击的选项列表
|
||||
:param title: 可选标题
|
||||
:return: 工具执行结果描述
|
||||
"""
|
||||
if self._blocked_by_feedback_quality_gate():
|
||||
logger.warning(
|
||||
"ask_user_choice blocked after feedback issue rejected_quality: "
|
||||
@@ -148,7 +158,8 @@ class AskUserChoiceTool(MoviePilotTool):
|
||||
|
||||
choice_options = [
|
||||
AgentInteractionOption(
|
||||
label=option.label.strip(), value=option.value.strip()
|
||||
label=option.label.strip(),
|
||||
value=option.value.strip(),
|
||||
)
|
||||
for option in options
|
||||
]
|
||||
|
||||
@@ -47,8 +47,6 @@ class BrowserAction(str, Enum):
|
||||
class BrowseWebpageInput(BaseModel):
|
||||
"""浏览器操作工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this browser action is being performed",)
|
||||
action: str = Field(
|
||||
...,
|
||||
description=(
|
||||
@@ -228,6 +226,11 @@ class BrowseWebpageTool(MoviePilotTool):
|
||||
return "错误: 'fill_ref' 操作需要提供 value 参数"
|
||||
if browser_action == BrowserAction.EVALUATE and not script:
|
||||
return "错误: 'evaluate' 操作需要提供 script 参数"
|
||||
if (
|
||||
browser_action == BrowserAction.EVALUATE
|
||||
and not await self.is_admin_user()
|
||||
):
|
||||
return "错误: 'evaluate' 操作仅允许管理员使用"
|
||||
if (
|
||||
browser_action in (BrowserAction.FOCUS_TAB, BrowserAction.CLOSE_TAB)
|
||||
and tab_index is None
|
||||
|
||||
@@ -0,0 +1,161 @@
|
||||
import json
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Literal, Optional, Type
|
||||
|
||||
import pytz
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.core.config import settings
|
||||
from app.db.agentchat_oper import AgentChatOper
|
||||
from app.db.agenttask_oper import AgentTaskOper
|
||||
from app.utils.timer import TimerUtils
|
||||
|
||||
|
||||
class CreateAgentTaskInput(BaseModel):
|
||||
"""创建 Agent 自主定时任务的输入参数。"""
|
||||
|
||||
name: str = Field(
|
||||
...,
|
||||
min_length=1,
|
||||
max_length=100,
|
||||
description="Short task name shown in task management and execution reports.",
|
||||
)
|
||||
content: str = Field(
|
||||
...,
|
||||
min_length=1,
|
||||
max_length=10000,
|
||||
description="Complete instructions that the agent must execute when the task fires.",
|
||||
)
|
||||
trigger_type: Literal["date", "cron"] = Field(
|
||||
...,
|
||||
description="Use 'date' for one exact future run or 'cron' for recurring work.",
|
||||
)
|
||||
trigger: Optional[str] = Field(
|
||||
None,
|
||||
min_length=1,
|
||||
max_length=200,
|
||||
description=(
|
||||
"For date, an ISO 8601 local or timezone-aware time such as "
|
||||
"2026-07-19 20:30:00; for cron, a standard five-field expression "
|
||||
"(minute hour day month weekday). The MoviePilot system timezone is used."
|
||||
),
|
||||
)
|
||||
delay_minutes: Optional[int] = Field(
|
||||
None,
|
||||
ge=1,
|
||||
le=525600,
|
||||
description=(
|
||||
"For a one-time date task expressed as 'in N minutes', provide this instead "
|
||||
"of trigger. MoviePilot calculates and persists the exact future run time."
|
||||
),
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_trigger(self) -> "CreateAgentTaskInput":
|
||||
"""校验任务触发配置并统一格式。"""
|
||||
self.name = self.name.strip()
|
||||
self.content = self.content.strip()
|
||||
if not self.name or not self.content:
|
||||
raise ValueError("name 和 content 不能只包含空白字符")
|
||||
if self.trigger_type == "date":
|
||||
if self.delay_minutes is not None:
|
||||
# LangChain 会在 run() 前后各校验一次,延迟时间在持久化前统一计算。
|
||||
self.trigger = None
|
||||
return self
|
||||
if self.trigger is None:
|
||||
raise ValueError("date 任务必须提供 trigger 或 delay_minutes")
|
||||
elif self.trigger is None or self.delay_minutes is not None:
|
||||
raise ValueError("cron 任务必须提供 trigger,且不能提供 delay_minutes")
|
||||
self.trigger_type, self.trigger = TimerUtils.normalize_schedule_trigger(
|
||||
trigger_type=self.trigger_type,
|
||||
trigger_value=self.trigger,
|
||||
timezone_name=settings.TZ,
|
||||
require_future=True,
|
||||
)
|
||||
return self
|
||||
|
||||
|
||||
class CreateAgentTaskTool(MoviePilotTool):
|
||||
"""创建可精确唤醒当前 Agent 会话的自主定时任务。"""
|
||||
|
||||
name: str = "create_agent_task"
|
||||
tags: list[str] = [ToolTag.Write, ToolTag.AgentTask, ToolTag.Admin]
|
||||
description: str = (
|
||||
"Create a persistent autonomous agent task only when the user explicitly asks "
|
||||
"for delayed, scheduled, recurring, reminder, or monitoring work. Use trigger_type "
|
||||
"'date' with delay_minutes for requests such as 'check in 30 minutes', an exact "
|
||||
"trigger time for other one-time work, and 'cron' for recurring schedules. When "
|
||||
"fired, MoviePilot wakes the agent in this conversation, executes content, and "
|
||||
"broadcasts user-facing messages through the configured notification channels."
|
||||
)
|
||||
args_schema: Type[BaseModel] = CreateAgentTaskInput
|
||||
require_admin: bool = True
|
||||
|
||||
def get_tool_message(self, **kwargs: object) -> Optional[str]:
|
||||
"""生成创建定时任务的提示消息。"""
|
||||
return f"创建自主定时任务:{kwargs.get('name', '')}"
|
||||
|
||||
def _create_task(self, payload: CreateAgentTaskInput) -> dict:
|
||||
"""持久化任务并立即注册到运行时调度器。"""
|
||||
from app.scheduler import Scheduler
|
||||
|
||||
trigger_value = payload.trigger
|
||||
if payload.trigger_type == "date" and payload.delay_minutes is not None:
|
||||
timezone = pytz.timezone(settings.TZ)
|
||||
trigger_value = (
|
||||
datetime.now(timezone) + timedelta(minutes=payload.delay_minutes)
|
||||
).isoformat(timespec="seconds")
|
||||
_, trigger_value = TimerUtils.normalize_schedule_trigger(
|
||||
trigger_type=payload.trigger_type,
|
||||
trigger_value=trigger_value,
|
||||
timezone_name=settings.TZ,
|
||||
require_future=True,
|
||||
)
|
||||
chat = AgentChatOper().get(
|
||||
session_id=self._session_id,
|
||||
user_id=self._user_id,
|
||||
)
|
||||
task = AgentTaskOper().add(
|
||||
name=payload.name.strip(),
|
||||
content=payload.content.strip(),
|
||||
trigger_type=payload.trigger_type,
|
||||
cron_expression=trigger_value if payload.trigger_type == "cron" else None,
|
||||
run_at=trigger_value if payload.trigger_type == "date" else None,
|
||||
user_id=str(self._user_id),
|
||||
username=self._username or (chat.username if chat else None),
|
||||
session_id=str(self._session_id),
|
||||
channel=self._channel or (chat.channel if chat else None),
|
||||
source=self._source or (chat.source if chat else None),
|
||||
original_chat_id=chat.original_chat_id if chat else None,
|
||||
)
|
||||
scheduler = Scheduler()
|
||||
next_run_at = scheduler.update_agent_task_job(task.id)
|
||||
return AgentTaskOper.to_dict(
|
||||
task,
|
||||
next_run_at=next_run_at,
|
||||
timezone=settings.TZ,
|
||||
)
|
||||
|
||||
async def run(
|
||||
self,
|
||||
name: str,
|
||||
content: str,
|
||||
trigger_type: str,
|
||||
trigger: Optional[str] = None,
|
||||
delay_minutes: Optional[int] = None,
|
||||
**kwargs: object,
|
||||
) -> str:
|
||||
"""创建 Agent 自主定时任务。"""
|
||||
if not settings.AI_AGENT_ENABLE:
|
||||
return "AI Agent 未启用,无法创建自主定时任务"
|
||||
payload = CreateAgentTaskInput(
|
||||
name=name,
|
||||
content=content,
|
||||
trigger_type=trigger_type,
|
||||
trigger=trigger,
|
||||
delay_minutes=delay_minutes,
|
||||
)
|
||||
task = await self.run_blocking("db", self._create_task, payload)
|
||||
return json.dumps(task, ensure_ascii=False, indent=2)
|
||||
@@ -0,0 +1,50 @@
|
||||
from typing import Optional, Type
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.db.agenttask_oper import AgentTaskOper
|
||||
|
||||
|
||||
class DeleteAgentTaskInput(BaseModel):
|
||||
"""删除 Agent 自主定时任务的输入参数。"""
|
||||
|
||||
task_id: int = Field(..., ge=1, description="ID of the task to permanently delete.")
|
||||
|
||||
|
||||
class DeleteAgentTaskTool(MoviePilotTool):
|
||||
"""永久删除 Agent 自主定时任务。"""
|
||||
|
||||
name: str = "delete_agent_task"
|
||||
tags: list[str] = [ToolTag.Write, ToolTag.AgentTask, ToolTag.Admin]
|
||||
description: str = (
|
||||
"Permanently delete an autonomous agent task and remove its runtime schedule. "
|
||||
"Use update_agent_task with enabled=false when the user only wants to pause it."
|
||||
)
|
||||
args_schema: Type[BaseModel] = DeleteAgentTaskInput
|
||||
require_admin: bool = True
|
||||
|
||||
def get_tool_message(self, **kwargs: object) -> Optional[str]:
|
||||
"""生成删除定时任务的提示消息。"""
|
||||
return f"删除自主定时任务:{kwargs.get('task_id', '')}"
|
||||
|
||||
def _delete_task(self, task_id: int) -> bool:
|
||||
"""删除当前用户的任务并移除运行时调度。"""
|
||||
from app.scheduler import Scheduler
|
||||
|
||||
deleted = AgentTaskOper().delete(
|
||||
task_id=task_id,
|
||||
user_id=str(self._user_id),
|
||||
)
|
||||
if deleted:
|
||||
Scheduler().remove_agent_task_job(task_id)
|
||||
return deleted
|
||||
|
||||
async def run(self, task_id: int, **kwargs: object) -> str:
|
||||
"""删除 Agent 自主定时任务。"""
|
||||
payload = DeleteAgentTaskInput(task_id=task_id)
|
||||
deleted = await self.run_blocking("db", self._delete_task, payload.task_id)
|
||||
if not deleted:
|
||||
return f"Agent 定时任务 {task_id} 不存在或不属于当前用户"
|
||||
return f"Agent 定时任务 {task_id} 已删除"
|
||||
@@ -20,8 +20,6 @@ from app.schemas.types import SystemConfigKey
|
||||
class DeleteCustomFilterRuleInput(BaseModel):
|
||||
"""删除自定义过滤规则工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
rule_id: str = Field(..., description="Custom rule ID to delete.")
|
||||
|
||||
|
||||
|
||||
@@ -6,16 +6,13 @@ from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.db import AsyncSessionFactory
|
||||
from app.db.models.downloadhistory import DownloadHistory
|
||||
from app.db.downloadhistory_oper import DownloadHistoryOper
|
||||
from app.log import logger
|
||||
|
||||
|
||||
class DeleteDownloadHistoryInput(BaseModel):
|
||||
"""删除下载历史记录工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
history_id: int = Field(
|
||||
..., description="The ID of the download history record to delete"
|
||||
)
|
||||
@@ -40,9 +37,8 @@ class DeleteDownloadHistoryTool(MoviePilotTool):
|
||||
logger.info(f"执行工具: {self.name}, 参数: history_id={history_id}")
|
||||
|
||||
try:
|
||||
async with AsyncSessionFactory() as db:
|
||||
await DownloadHistory.async_delete(db, history_id)
|
||||
return f"下载历史记录 ID: {history_id} 已成功删除"
|
||||
await DownloadHistoryOper().async_delete_history(history_id)
|
||||
return f"下载历史记录 ID: {history_id} 已成功删除"
|
||||
except Exception as e:
|
||||
logger.error(f"删除下载历史记录失败: {e}", exc_info=True)
|
||||
return f"删除下载历史记录时发生错误: {str(e)}"
|
||||
|
||||
@@ -13,8 +13,6 @@ from app.log import logger
|
||||
class DeleteDownloadTasksInput(BaseModel):
|
||||
"""删除下载任务工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
hash: str = Field(
|
||||
..., description="Task hash (can be obtained from query_download_tasks tool)"
|
||||
)
|
||||
|
||||
@@ -19,8 +19,6 @@ from app.schemas.types import SystemConfigKey
|
||||
class DeleteRuleGroupInput(BaseModel):
|
||||
"""删除过滤规则组工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
name: str = Field(..., description="Rule group name to delete.")
|
||||
|
||||
|
||||
|
||||
@@ -16,8 +16,6 @@ from app.schemas.types import EventType
|
||||
class DeleteSubscribeInput(BaseModel):
|
||||
"""删除订阅工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
subscribe_id: int = Field(
|
||||
...,
|
||||
description="The ID of the subscription to delete (can be obtained from query_subscribes tool)",
|
||||
@@ -56,7 +54,15 @@ class DeleteSubscribeTool(MoviePilotTool):
|
||||
await subscribe_oper.async_delete(subscribe_id)
|
||||
# 分享订阅统计刷新本身已异步化,这里只需要在删除后触发即可。
|
||||
MoviePilotServerHelper.sub_done_async(
|
||||
{"tmdbid": subscribe.tmdbid, "doubanid": subscribe.doubanid}
|
||||
{
|
||||
"tmdbid": subscribe.tmdbid,
|
||||
"doubanid": subscribe.doubanid,
|
||||
"bangumiid": subscribe.bangumiid,
|
||||
"anilistid": subscribe.anilistid,
|
||||
"media_source": subscribe.media_source,
|
||||
"media_id": subscribe.media_id,
|
||||
"season": subscribe.season,
|
||||
}
|
||||
)
|
||||
|
||||
# 发送事件
|
||||
|
||||
@@ -15,8 +15,6 @@ from app.schemas import FileItem
|
||||
class DeleteTransferHistoryInput(BaseModel):
|
||||
"""删除整理历史记录工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
history_id: int = Field(
|
||||
..., description="The ID of the transfer history record to delete"
|
||||
)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""文件编辑工具"""
|
||||
"""文件精确编辑工具。"""
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Optional, Type
|
||||
@@ -7,28 +7,59 @@ from anyio import Path as AsyncPath
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.impl._file_write_utils import (
|
||||
FileVersionConflictError,
|
||||
atomic_write_text,
|
||||
calculate_file_sha256,
|
||||
)
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.log import logger
|
||||
|
||||
|
||||
class EditFileInput(BaseModel):
|
||||
"""Input parameters for edit file tool"""
|
||||
"""文件编辑工具的输入参数模型。"""
|
||||
|
||||
file_path: str = Field(..., description="The absolute path of the file to edit")
|
||||
old_text: str = Field(..., description="The exact old text to be replaced")
|
||||
old_text: str = Field(
|
||||
...,
|
||||
description=(
|
||||
"The exact old text to replace. It must be non-empty and uniquely "
|
||||
"identify one location unless replace_all is true."
|
||||
),
|
||||
)
|
||||
new_text: str = Field(..., description="The new text to replace with")
|
||||
replace_all: bool = Field(
|
||||
False,
|
||||
description=(
|
||||
"Replace every exact match. Keep false for normal code edits so an "
|
||||
"ambiguous match fails instead of changing multiple locations."
|
||||
),
|
||||
)
|
||||
expected_sha256: Optional[str] = Field(
|
||||
None,
|
||||
pattern=r"^[0-9a-fA-F]{64}$",
|
||||
description=(
|
||||
"Optional SHA-256 returned by read_file(include_metadata=true). The "
|
||||
"edit fails if the file changed after it was read."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class EditFileTool(MoviePilotTool):
|
||||
"""使用精确文本匹配安全编辑本地文件。"""
|
||||
|
||||
name: str = "edit_file"
|
||||
tags: list[str] = [
|
||||
ToolTag.Write,
|
||||
ToolTag.File,
|
||||
]
|
||||
description: str = (
|
||||
"Edit a local text file by replacing specific old text with new text. "
|
||||
"Non-admin users can only edit files inside the MoviePilot config, "
|
||||
"Agent memory/activity, and log directories."
|
||||
"Edit an existing local text file using an exact text match. By default "
|
||||
"the match must occur exactly once; use replace_all only for intentional "
|
||||
"bulk replacement. old_text cannot be empty, and new files must be "
|
||||
"created with write_file. Supports an optional SHA-256 conflict check. "
|
||||
"Non-admin users can only edit files inside the MoviePilot Agent config "
|
||||
"directory."
|
||||
)
|
||||
args_schema: Type[BaseModel] = EditFileInput
|
||||
|
||||
@@ -38,7 +69,16 @@ class EditFileTool(MoviePilotTool):
|
||||
file_name = Path(file_path).name if file_path else "未知文件"
|
||||
return f"编辑文件: {file_name}"
|
||||
|
||||
async def run(self, file_path: str, old_text: str, new_text: str, **kwargs) -> str:
|
||||
async def run(
|
||||
self,
|
||||
file_path: str,
|
||||
old_text: str,
|
||||
new_text: str,
|
||||
replace_all: bool = False,
|
||||
expected_sha256: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""校验精确匹配和可选文件版本后,以原子方式写入编辑结果。"""
|
||||
logger.info(f"执行工具: {self.name}, 参数: file_path={file_path}")
|
||||
|
||||
try:
|
||||
@@ -48,37 +88,74 @@ class EditFileTool(MoviePilotTool):
|
||||
if access_error:
|
||||
return access_error
|
||||
|
||||
path = AsyncPath(resolved_path)
|
||||
# 校验逻辑:如果要替换特定文本,文件必须存在且包含该文本
|
||||
if not await path.exists():
|
||||
# 如果 old_text 为空,可能用户想直接创建文件,但通常 edit_file 需要匹配旧内容
|
||||
if old_text:
|
||||
return f"错误:文件 {resolved_path} 不存在,无法进行内容替换。"
|
||||
if not old_text:
|
||||
return "错误:old_text 不能为空;创建或完整写入文件请使用 write_file。"
|
||||
|
||||
if await path.exists() and not await path.is_file():
|
||||
path = AsyncPath(resolved_path)
|
||||
if not await path.exists():
|
||||
return f"错误:文件 {resolved_path} 不存在;创建文件请使用 write_file。"
|
||||
|
||||
if not await path.is_file():
|
||||
return f"错误:{resolved_path} 不是一个文件"
|
||||
|
||||
if await path.exists():
|
||||
content = await path.read_text(encoding="utf-8", errors="replace")
|
||||
if old_text not in content:
|
||||
logger.warning(f"编辑文件 {resolved_path} 失败:未找到指定的旧文本块")
|
||||
return f"错误:在文件 {resolved_path} 中未找到指定的旧文本。请确保包含所有的空格、缩进 and 换行符。"
|
||||
occurrences = content.count(old_text)
|
||||
new_content = content.replace(old_text, new_text)
|
||||
else:
|
||||
# 文件不存在且 old_text 为空的情形(初始化新文件)
|
||||
new_content = new_text
|
||||
occurrences = 1
|
||||
local_path = Path(resolved_path)
|
||||
current_sha256 = await self.run_blocking(
|
||||
"default", calculate_file_sha256, local_path
|
||||
)
|
||||
if (
|
||||
expected_sha256
|
||||
and current_sha256.casefold() != expected_sha256.casefold()
|
||||
):
|
||||
return (
|
||||
f"错误:文件 {resolved_path} 已在读取后发生变化,拒绝覆盖。"
|
||||
"请重新读取文件并基于最新内容编辑。"
|
||||
)
|
||||
|
||||
# 自动创建父目录
|
||||
await path.parent.mkdir(parents=True, exist_ok=True)
|
||||
content = await path.read_text(encoding="utf-8", errors="strict")
|
||||
occurrences = content.count(old_text)
|
||||
if occurrences == 0:
|
||||
logger.warning(f"编辑文件 {resolved_path} 失败:未找到指定的旧文本块")
|
||||
return (
|
||||
f"错误:在文件 {resolved_path} 中未找到指定的旧文本。"
|
||||
"请重新读取文件并确认空格、缩进和换行。"
|
||||
)
|
||||
if occurrences > 1 and not replace_all:
|
||||
return (
|
||||
f"错误:old_text 在文件 {resolved_path} 中匹配到 {occurrences} 处,"
|
||||
"为避免误改已拒绝编辑。请提供更多上下文使其唯一,或明确设置 "
|
||||
"replace_all=true。"
|
||||
)
|
||||
|
||||
# 写入文件
|
||||
await path.write_text(new_content, encoding="utf-8")
|
||||
replacement_count = occurrences if replace_all else 1
|
||||
new_content = content.replace(
|
||||
old_text,
|
||||
new_text,
|
||||
-1 if replace_all else 1,
|
||||
)
|
||||
await self.run_blocking(
|
||||
"default",
|
||||
atomic_write_text,
|
||||
local_path,
|
||||
new_content,
|
||||
current_sha256,
|
||||
)
|
||||
new_sha256 = await self.run_blocking(
|
||||
"default", calculate_file_sha256, local_path
|
||||
)
|
||||
|
||||
logger.info(f"成功编辑文件 {resolved_path},替换了 {occurrences} 处内容")
|
||||
return f"成功编辑文件 {resolved_path} (替换了 {occurrences} 处匹配内容)"
|
||||
logger.info(
|
||||
f"成功编辑文件 {resolved_path},替换了 {replacement_count} 处内容"
|
||||
)
|
||||
return (
|
||||
f"成功编辑文件 {resolved_path}(替换了 {replacement_count} 处匹配内容,"
|
||||
f"sha256={new_sha256})"
|
||||
)
|
||||
|
||||
except FileVersionConflictError:
|
||||
return (
|
||||
f"错误:文件 {file_path} 在编辑期间发生变化,拒绝覆盖。"
|
||||
"请重新读取文件并再次编辑。"
|
||||
)
|
||||
except PermissionError:
|
||||
return f"错误:没有访问/修改 {file_path} 的权限"
|
||||
except UnicodeDecodeError:
|
||||
|
||||
@@ -13,6 +13,7 @@ from typing import Any, Literal, Optional, TextIO, Type
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.impl._command_safety import validate_command_safety
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.agent.tools.impl._terminal_session import (
|
||||
@@ -30,14 +31,6 @@ MAX_OUTPUT_PREVIEW_BYTES = 10 * 1024
|
||||
READ_CHUNK_SIZE = 4096
|
||||
KILL_GRACE_SECONDS = 3
|
||||
COMMAND_CONCURRENCY_LIMIT = 2
|
||||
COMMAND_FORBIDDEN_KEYWORDS = (
|
||||
":(){ :|:& };:",
|
||||
"dd if=/dev/zero",
|
||||
"mkfs",
|
||||
"reboot",
|
||||
"shutdown",
|
||||
)
|
||||
|
||||
_command_semaphore = asyncio.Semaphore(COMMAND_CONCURRENCY_LIMIT)
|
||||
|
||||
|
||||
@@ -142,7 +135,6 @@ class _CommandOutput:
|
||||
class ExecuteCommandInput(BaseModel):
|
||||
"""执行 Shell 命令工具的输入参数模型。"""
|
||||
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this command action is needed")
|
||||
action: Optional[Literal["start", "read", "wait", "write", "kill", "run"]] = Field(
|
||||
"start",
|
||||
description=(
|
||||
@@ -195,6 +187,13 @@ class ExecuteCommandInput(BaseModel):
|
||||
60,
|
||||
description="For action=run, max execution time in seconds.",
|
||||
)
|
||||
confirm_dangerous: Optional[bool] = Field(
|
||||
False,
|
||||
description=(
|
||||
"Explicit confirmation for high-risk commands such as recursive root deletion, "
|
||||
"disk formatting, shutdown/reboot, or destructive permission changes."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class ExecuteCommandTool(MoviePilotTool):
|
||||
@@ -255,34 +254,9 @@ class ExecuteCommandTool(MoviePilotTool):
|
||||
return command
|
||||
|
||||
@staticmethod
|
||||
def _validate_command(command: str) -> None:
|
||||
def _validate_command(command: str, *, confirmed: bool = False) -> None:
|
||||
"""复用旧工具的基础危险命令过滤,避免明显破坏性命令进入 shell。"""
|
||||
for keyword in COMMAND_FORBIDDEN_KEYWORDS:
|
||||
if keyword in command:
|
||||
raise ValueError(f"命令包含禁止使用的关键字 '{keyword}'")
|
||||
|
||||
# 检查是否使用了 rm -r/R 删除根目录或一级目录,防止误杀多级目录
|
||||
import re
|
||||
import os.path
|
||||
tokens = re.split(r'\s+', command.strip())
|
||||
if any(t == "rm" or t.endswith("/rm") for t in tokens):
|
||||
has_r = False
|
||||
for token in tokens:
|
||||
if token.startswith("-") and ("r" in token or "R" in token):
|
||||
has_r = True
|
||||
break
|
||||
|
||||
if has_r:
|
||||
for token in tokens:
|
||||
# 提取可能包含目标路径的部分(去除重定向、管道、分号等末尾干扰)
|
||||
m = re.match(r'^([^;\|&><]+)', token)
|
||||
if m:
|
||||
clean_token = m.group(1).strip('"\'')
|
||||
# 仅对绝对路径进行一级目录限制
|
||||
if clean_token.startswith('/'):
|
||||
norm_path = os.path.normpath(clean_token)
|
||||
if re.match(r'^/[^/]*$', norm_path) or re.match(r'^/[^/]*/$', norm_path):
|
||||
raise ValueError(f"不允许使用 rm 命令删除根目录或一级目录: {clean_token}")
|
||||
validate_command_safety(command, confirmed=confirmed)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_timeout(timeout: Optional[int]) -> tuple[int, Optional[str]]:
|
||||
@@ -367,7 +341,7 @@ class ExecuteCommandTool(MoviePilotTool):
|
||||
asyncio.shield(wait_task), timeout=KILL_GRACE_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning("命令进程强制清理超时: pid=%s", process.pid)
|
||||
logger.warning(f"命令进程强制清理超时: pid={process.pid}")
|
||||
|
||||
@staticmethod
|
||||
async def _finish_reader_tasks(reader_tasks: list[asyncio.Task]) -> None:
|
||||
@@ -382,7 +356,7 @@ class ExecuteCommandTool(MoviePilotTool):
|
||||
if isinstance(result, Exception) and not isinstance(
|
||||
result, asyncio.CancelledError
|
||||
):
|
||||
logger.debug("命令输出读取任务异常: %s", result)
|
||||
logger.debug(f"命令输出读取任务异常: {result}")
|
||||
|
||||
@staticmethod
|
||||
def _format_run_result(
|
||||
@@ -425,9 +399,10 @@ class ExecuteCommandTool(MoviePilotTool):
|
||||
command: str,
|
||||
timeout: Optional[int],
|
||||
cwd: Optional[str] = None,
|
||||
confirm_dangerous: bool = False,
|
||||
) -> str:
|
||||
"""按旧模式一次性执行命令,等待完成或超时后返回文本结果。"""
|
||||
self._validate_command(command)
|
||||
self._validate_command(command, confirmed=confirm_dangerous)
|
||||
normalized_timeout, timeout_note = self._normalize_timeout(timeout)
|
||||
|
||||
async with _command_semaphore:
|
||||
@@ -482,27 +457,29 @@ class ExecuteCommandTool(MoviePilotTool):
|
||||
max_bytes: Optional[int] = TERMINAL_DEFAULT_READ_BYTES,
|
||||
timeout_ms: Optional[int] = TERMINAL_WAIT_DEFAULT_MS,
|
||||
timeout: Optional[int] = 60,
|
||||
confirm_dangerous: Optional[bool] = False,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""执行命令动作:默认后台启动,也支持读取、等待、写入、终止和一次性执行。"""
|
||||
normalized_action = (action or "start").strip().lower()
|
||||
logger.info(
|
||||
"执行工具: %s, action=%s, command=%s, session_id=%s",
|
||||
self.name,
|
||||
normalized_action,
|
||||
command,
|
||||
session_id,
|
||||
f"执行工具: {self.name}, action={normalized_action}, "
|
||||
f"command={command}, session_id={session_id}"
|
||||
)
|
||||
|
||||
try:
|
||||
if normalized_action == "start":
|
||||
start_command = self._require_command(command)
|
||||
self._validate_command(start_command)
|
||||
self._validate_command(
|
||||
start_command,
|
||||
confirmed=bool(confirm_dangerous),
|
||||
)
|
||||
payload = await terminal_session_manager.start(
|
||||
command=start_command,
|
||||
cwd=cwd,
|
||||
env=env,
|
||||
use_pty=use_pty,
|
||||
confirm_dangerous=bool(confirm_dangerous),
|
||||
)
|
||||
return self._dump(payload)
|
||||
|
||||
@@ -542,9 +519,10 @@ class ExecuteCommandTool(MoviePilotTool):
|
||||
command=self._require_command(command),
|
||||
timeout=timeout,
|
||||
cwd=cwd,
|
||||
confirm_dangerous=bool(confirm_dangerous),
|
||||
)
|
||||
|
||||
raise ValueError(f"不支持的 action: {action}")
|
||||
except Exception as err:
|
||||
logger.error("执行命令 action 失败: %s", err, exc_info=True)
|
||||
logger.error(f"执行命令 action 失败: {err}", exc_info=True)
|
||||
return self._dump({"error": str(err), "status": "error", "action": normalized_action})
|
||||
|
||||
@@ -15,8 +15,6 @@ from app.schemas.types import MediaType, media_type_to_agent
|
||||
class GetRecommendationsInput(BaseModel):
|
||||
"""获取推荐工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
source: Optional[str] = Field(
|
||||
"tmdb_trending",
|
||||
description="Recommendation source: "
|
||||
@@ -212,6 +210,10 @@ class GetRecommendationsTool(MoviePilotTool):
|
||||
"tmdb_id": r.get("tmdb_id"),
|
||||
"imdb_id": r.get("imdb_id"),
|
||||
"douban_id": r.get("douban_id"),
|
||||
"bangumi_id": r.get("bangumi_id"),
|
||||
"anilist_id": r.get("anilist_id"),
|
||||
"media_source": r.get("source"),
|
||||
"media_id": r.get("media_id"),
|
||||
"vote_average": r.get("vote_average"),
|
||||
"poster_path": r.get("poster_path"),
|
||||
"detail_link": r.get("detail_link"),
|
||||
|
||||
@@ -21,8 +21,6 @@ from ._torrent_search_utils import (
|
||||
class GetSearchResultsInput(BaseModel):
|
||||
"""获取搜索结果工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
site: Optional[List[str]] = Field(None, description="Site name filters")
|
||||
season: Optional[List[str]] = Field(None, description="Season or episode filters")
|
||||
free_state: Optional[List[str]] = Field(None, description="Promotion state filters")
|
||||
@@ -36,6 +34,14 @@ class GetSearchResultsInput(BaseModel):
|
||||
None,
|
||||
description="Regular expression pattern to filter torrent titles (e.g., '4K|2160p|UHD', '1080p.*BluRay')",
|
||||
)
|
||||
content_pattern: Optional[str] = Field(
|
||||
None,
|
||||
description="Regular expression pattern to filter torrent titles, descriptions, and labels (e.g., '特效字幕|国语|DIY')",
|
||||
)
|
||||
include_description: Optional[bool] = Field(
|
||||
False,
|
||||
description="Whether to include torrent descriptions in returned results",
|
||||
)
|
||||
show_filter_options: Optional[bool] = Field(
|
||||
False,
|
||||
description="Whether to return only optional filter options for re-checking available conditions",
|
||||
@@ -47,6 +53,8 @@ class GetSearchResultsInput(BaseModel):
|
||||
|
||||
|
||||
class GetSearchResultsTool(MoviePilotTool):
|
||||
"""获取并筛选最近一次种子搜索结果"""
|
||||
|
||||
name: str = "get_search_results"
|
||||
tags: list[str] = [
|
||||
ToolTag.Read,
|
||||
@@ -56,6 +64,7 @@ class GetSearchResultsTool(MoviePilotTool):
|
||||
args_schema: Type[BaseModel] = GetSearchResultsInput
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
"""返回工具执行提示"""
|
||||
return "获取搜索结果"
|
||||
|
||||
async def run(
|
||||
@@ -68,13 +77,33 @@ class GetSearchResultsTool(MoviePilotTool):
|
||||
resolution: Optional[List[str]] = None,
|
||||
release_group: Optional[List[str]] = None,
|
||||
title_pattern: Optional[str] = None,
|
||||
content_pattern: Optional[str] = None,
|
||||
include_description: bool = False,
|
||||
show_filter_options: bool = False,
|
||||
page: Optional[int] = 1,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""
|
||||
获取并筛选最近一次种子搜索结果
|
||||
|
||||
:param site: 站点名称筛选项
|
||||
:param season: 季集筛选项
|
||||
:param free_state: 促销状态筛选项
|
||||
:param video_code: 视频编码筛选项
|
||||
:param edition: 制作版本筛选项
|
||||
:param resolution: 分辨率筛选项
|
||||
:param release_group: 发布组筛选项
|
||||
:param title_pattern: 仅匹配种子标题的正则表达式
|
||||
:param content_pattern: 匹配种子标题、简介和标签的正则表达式
|
||||
:param include_description: 是否在结果中返回种子简介
|
||||
:param show_filter_options: 是否只返回可用筛选项
|
||||
:param page: 分页页码
|
||||
:param kwargs: 工具框架附加参数
|
||||
:return: JSON 格式的搜索结果或错误提示
|
||||
"""
|
||||
page = max(1, page or 1)
|
||||
logger.info(
|
||||
f"执行工具: {self.name}, 参数: site={site}, season={season}, free_state={free_state}, video_code={video_code}, edition={edition}, resolution={resolution}, release_group={release_group}, title_pattern={title_pattern}, show_filter_options={show_filter_options}, page={page}"
|
||||
f"执行工具: {self.name}, 参数: site={site}, season={season}, free_state={free_state}, video_code={video_code}, edition={edition}, resolution={resolution}, release_group={release_group}, title_pattern={title_pattern}, content_pattern={content_pattern}, include_description={include_description}, show_filter_options={show_filter_options}, page={page}"
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -89,14 +118,22 @@ class GetSearchResultsTool(MoviePilotTool):
|
||||
}
|
||||
return json.dumps(payload, ensure_ascii=False, indent=2)
|
||||
|
||||
regex_pattern = None
|
||||
title_regex_pattern = None
|
||||
if title_pattern:
|
||||
try:
|
||||
regex_pattern = re.compile(title_pattern, re.IGNORECASE)
|
||||
title_regex_pattern = re.compile(title_pattern, re.IGNORECASE)
|
||||
except re.error as e:
|
||||
logger.warning(f"正则表达式编译失败: {title_pattern}, 错误: {e}")
|
||||
return f"正则表达式格式错误: {str(e)}"
|
||||
|
||||
content_regex_pattern = None
|
||||
if content_pattern:
|
||||
try:
|
||||
content_regex_pattern = re.compile(content_pattern, re.IGNORECASE)
|
||||
except re.error as e:
|
||||
logger.warning(f"正则表达式编译失败: {content_pattern}, 错误: {e}")
|
||||
return f"正则表达式格式错误: {str(e)}"
|
||||
|
||||
filtered_items = filter_contexts(
|
||||
items=items,
|
||||
site=site,
|
||||
@@ -107,14 +144,29 @@ class GetSearchResultsTool(MoviePilotTool):
|
||||
resolution=resolution,
|
||||
release_group=release_group,
|
||||
)
|
||||
if regex_pattern:
|
||||
if title_regex_pattern:
|
||||
filtered_items = [
|
||||
item
|
||||
for item in filtered_items
|
||||
if item.torrent_info
|
||||
and item.torrent_info.title
|
||||
and regex_pattern.search(item.torrent_info.title)
|
||||
and title_regex_pattern.search(item.torrent_info.title)
|
||||
]
|
||||
if content_regex_pattern:
|
||||
content_filtered_items = []
|
||||
for item in filtered_items:
|
||||
torrent_info = item.torrent_info
|
||||
if not torrent_info:
|
||||
continue
|
||||
content_values = [torrent_info.title, torrent_info.description]
|
||||
content_values.extend(torrent_info.labels or [])
|
||||
if any(
|
||||
content_regex_pattern.search(str(value))
|
||||
for value in content_values
|
||||
if value
|
||||
):
|
||||
content_filtered_items.append(item)
|
||||
filtered_items = content_filtered_items
|
||||
if not filtered_items:
|
||||
return "没有符合筛选条件的搜索结果,请调整筛选条件"
|
||||
|
||||
@@ -137,7 +189,11 @@ class GetSearchResultsTool(MoviePilotTool):
|
||||
return f"第 {page} 页没有数据,共 {total_count} 条结果,共 {(total_count + page_size - 1) // page_size} 页。"
|
||||
|
||||
results = [
|
||||
simplify_search_result(item, index)
|
||||
simplify_search_result(
|
||||
item,
|
||||
index,
|
||||
include_description=include_description,
|
||||
)
|
||||
for item, index in zip(page_items, page_indices)
|
||||
]
|
||||
total_pages = (total_count + page_size - 1) // page_size
|
||||
|
||||
@@ -19,8 +19,6 @@ from app.log import logger
|
||||
class InstallPluginInput(BaseModel):
|
||||
"""安装插件工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
plugin_id: str = Field(
|
||||
...,
|
||||
description="Exact plugin ID to install. Use query_market_plugins first to find the correct plugin_id.",
|
||||
|
||||
@@ -17,7 +17,6 @@ from app.utils.string import StringUtils
|
||||
|
||||
class ListDirectoryInput(BaseModel):
|
||||
"""查询文件系统目录内容工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
path: str = Field(..., description="Directory path to list contents (e.g., '/home/user/downloads' or 'C:/Downloads')")
|
||||
storage: Optional[str] = Field("local", description="Storage type (default: 'local' for local file system, can be 'smb', 'alist', etc.)")
|
||||
sort_by: Optional[str] = Field("name", description="Sort order: 'name' for alphabetical sorting, 'time' for modification time sorting (default: 'name')")
|
||||
|
||||
@@ -13,8 +13,6 @@ from app.log import logger
|
||||
class ListSlashCommandsInput(BaseModel):
|
||||
"""查询所有可用斜杠命令工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
|
||||
|
||||
class ListSlashCommandsTool(MoviePilotTool):
|
||||
|
||||
@@ -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
|
||||
@@ -0,0 +1,88 @@
|
||||
import json
|
||||
from typing import Optional, Type
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.core.config import settings
|
||||
from app.db.agenttask_oper import AgentTaskOper
|
||||
|
||||
|
||||
class QueryAgentTasksInput(BaseModel):
|
||||
"""查询 Agent 自主定时任务的输入参数。"""
|
||||
|
||||
task_id: Optional[int] = Field(
|
||||
None,
|
||||
ge=1,
|
||||
description="Optional task ID. Omit it to list tasks owned by the current user.",
|
||||
)
|
||||
enabled: Optional[bool] = Field(
|
||||
None,
|
||||
description="Optional enabled-state filter used when listing tasks.",
|
||||
)
|
||||
|
||||
|
||||
class QueryAgentTasksTool(MoviePilotTool):
|
||||
"""查询当前用户创建的 Agent 自主定时任务。"""
|
||||
|
||||
name: str = "query_agent_tasks"
|
||||
tags: list[str] = [ToolTag.Read, ToolTag.AgentTask, ToolTag.Admin]
|
||||
description: str = (
|
||||
"Query persistent autonomous agent tasks owned by the current user, including "
|
||||
"reminders, monitoring tasks, and recurring agent work. Returns the integer "
|
||||
"task_id, instructions, trigger, enabled state, next run time, and latest result. "
|
||||
"Do not use this for MoviePilot system, plugin, or workflow scheduler services."
|
||||
)
|
||||
args_schema: Type[BaseModel] = QueryAgentTasksInput
|
||||
require_admin: bool = True
|
||||
|
||||
def get_tool_message(self, **kwargs: object) -> Optional[str]:
|
||||
"""生成查询定时任务的提示消息。"""
|
||||
task_id = kwargs.get("task_id")
|
||||
return f"查询自主定时任务:{task_id}" if task_id else "查询自主定时任务"
|
||||
|
||||
def _query_tasks(
|
||||
self,
|
||||
task_id: Optional[int],
|
||||
enabled: Optional[bool],
|
||||
) -> list[dict]:
|
||||
"""读取当前用户的任务及运行时下一次触发时间。"""
|
||||
from app.scheduler import Scheduler
|
||||
|
||||
oper = AgentTaskOper()
|
||||
if task_id:
|
||||
task = oper.get(task_id=task_id, user_id=str(self._user_id))
|
||||
tasks = [task] if task else []
|
||||
else:
|
||||
tasks = oper.list(user_id=str(self._user_id), enabled=enabled)
|
||||
scheduler = Scheduler()
|
||||
result = []
|
||||
for task in tasks:
|
||||
data = oper.to_dict(
|
||||
task,
|
||||
next_run_at=scheduler.get_agent_task_next_run(task.id),
|
||||
timezone=settings.TZ,
|
||||
)
|
||||
result.append(data)
|
||||
return result
|
||||
|
||||
async def run(
|
||||
self,
|
||||
task_id: Optional[int] = None,
|
||||
enabled: Optional[bool] = None,
|
||||
**kwargs: object,
|
||||
) -> str:
|
||||
"""查询 Agent 自主定时任务。"""
|
||||
payload = QueryAgentTasksInput(task_id=task_id, enabled=enabled)
|
||||
tasks = await self.run_blocking(
|
||||
"db",
|
||||
self._query_tasks,
|
||||
payload.task_id,
|
||||
payload.enabled,
|
||||
)
|
||||
return json.dumps(
|
||||
{"total": len(tasks), "tasks": tasks},
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
)
|
||||
@@ -18,8 +18,6 @@ from app.log import logger
|
||||
class QueryBuiltinFilterRulesInput(BaseModel):
|
||||
"""查询内置过滤规则工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
rule_ids: Optional[List[str]] = Field(
|
||||
None,
|
||||
description="Optional list of built-in rule IDs to query. If omitted, return all built-in rules.",
|
||||
|
||||
@@ -19,8 +19,6 @@ from app.log import logger
|
||||
class QueryCustomFilterRulesInput(BaseModel):
|
||||
"""查询自定义过滤规则工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
rule_ids: Optional[List[str]] = Field(
|
||||
None,
|
||||
description="Optional list of custom rule IDs to query. If omitted, return all custom rules.",
|
||||
|
||||
@@ -15,8 +15,6 @@ from app.schemas.types import SystemConfigKey
|
||||
class QueryCustomIdentifiersInput(BaseModel):
|
||||
"""查询自定义识别词工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
|
||||
|
||||
class QueryCustomIdentifiersTool(MoviePilotTool):
|
||||
|
||||
@@ -13,7 +13,6 @@ from app.log import logger
|
||||
|
||||
class QueryDirectorySettingsInput(BaseModel):
|
||||
"""查询系统目录设置工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
directory_type: Optional[str] = Field("all",
|
||||
description="Filter directories by type: 'download' for download directories, 'library' for media library directories, 'all' for all directories")
|
||||
storage_type: Optional[str] = Field("all",
|
||||
|
||||
@@ -14,10 +14,6 @@ from app.log import logger
|
||||
class QueryDoctorReportInput(BaseModel):
|
||||
"""查询 Doctor 诊断报告工具的输入参数模型。"""
|
||||
|
||||
explanation: Optional[str] = Field(
|
||||
None,
|
||||
description="Clear explanation of why this tool is being used in the current context",
|
||||
)
|
||||
deep: Optional[bool] = Field(
|
||||
False,
|
||||
description=(
|
||||
@@ -48,7 +44,9 @@ class QueryDoctorReportTool(MoviePilotTool):
|
||||
description: str = (
|
||||
"Run MoviePilot Doctor in read-only mode and return a structured diagnostic report for troubleshooting. "
|
||||
"Use this tool when analyzing startup failures, Docker/runtime issues, port conflicts, dependency problems, "
|
||||
"database health, frontend assets, safe mode, or recent log error clues. This tool never applies fixes."
|
||||
"database health, frontend assets, safe mode, or recent log error clues. Plugin-only log findings remain "
|
||||
"visible with affects_report_status=false and do not downgrade the overall status. This tool never applies "
|
||||
"fixes."
|
||||
)
|
||||
require_admin: bool = True
|
||||
args_schema: Type[BaseModel] = QueryDoctorReportInput
|
||||
@@ -77,6 +75,7 @@ class QueryDoctorReportTool(MoviePilotTool):
|
||||
"title": item.get("title"),
|
||||
"fixable": item.get("fixable"),
|
||||
"fixed": item.get("fixed"),
|
||||
"affects_report_status": item.get("affects_report_status", True),
|
||||
}
|
||||
for item in report.get("findings") or []
|
||||
if isinstance(item, dict)
|
||||
|
||||
@@ -16,7 +16,6 @@ from app.schemas.types import TorrentQueryStatus, media_type_to_agent
|
||||
|
||||
class QueryDownloadTasksInput(BaseModel):
|
||||
"""查询下载工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
downloader: Optional[str] = Field(None,
|
||||
description="Name of specific downloader to query (optional, if not provided queries all configured downloaders)")
|
||||
status: Optional[str] = Field("all",
|
||||
|
||||
@@ -14,9 +14,6 @@ from app.schemas.types import SystemConfigKey
|
||||
|
||||
class QueryDownloadersInput(BaseModel):
|
||||
"""查询下载器工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
|
||||
|
||||
class QueryDownloadersTool(MoviePilotTool):
|
||||
name: str = "query_downloaders"
|
||||
tags: list[str] = [
|
||||
|
||||
@@ -13,7 +13,6 @@ from app.log import logger
|
||||
|
||||
class QueryEpisodeScheduleInput(BaseModel):
|
||||
"""查询剧集上映时间工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
tmdb_id: int = Field(..., description="TMDB ID of the TV series (can be obtained from search_media tool)")
|
||||
season: int = Field(..., description="Season number to query")
|
||||
episode_group: Optional[str] = Field(None, description="Episode group ID (optional)")
|
||||
|
||||
@@ -10,6 +10,7 @@ from app.agent.tools.tags import ToolTag
|
||||
from app.agent.tools.impl._plugin_tool_utils import (
|
||||
DEFAULT_PLUGIN_CANDIDATE_LIMIT,
|
||||
MAX_PLUGIN_CANDIDATE_LIMIT,
|
||||
enrich_installed_plugin_sources,
|
||||
list_installed_plugins,
|
||||
search_plugin_candidates,
|
||||
summarize_candidates,
|
||||
@@ -21,8 +22,6 @@ from app.log import logger
|
||||
class QueryInstalledPluginsInput(BaseModel):
|
||||
"""查询已安装插件工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
query: Optional[str] = Field(
|
||||
None,
|
||||
description="Optional keyword to filter installed plugins by plugin ID, name, description, or author.",
|
||||
@@ -31,9 +30,15 @@ class QueryInstalledPluginsInput(BaseModel):
|
||||
DEFAULT_PLUGIN_CANDIDATE_LIMIT,
|
||||
description="Maximum number of plugins to return. Defaults to 50, capped at 200.",
|
||||
)
|
||||
force_refresh_market: bool = Field(
|
||||
False,
|
||||
description="Whether to refresh plugin market caches before completing missing repo_url values.",
|
||||
)
|
||||
|
||||
|
||||
class QueryInstalledPluginsTool(MoviePilotTool):
|
||||
"""查询已安装插件并返回 Agent 可消费的摘要信息。"""
|
||||
|
||||
name: str = "query_installed_plugins"
|
||||
tags: list[str] = [
|
||||
ToolTag.Read,
|
||||
@@ -67,9 +72,15 @@ class QueryInstalledPluginsTool(MoviePilotTool):
|
||||
self,
|
||||
query: Optional[str] = None,
|
||||
max_results: Optional[int] = DEFAULT_PLUGIN_CANDIDATE_LIMIT,
|
||||
force_refresh_market: bool = False,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
logger.info(f"执行工具: {self.name}, 参数: query={query}")
|
||||
"""
|
||||
查询已安装插件列表,并在可能时补齐插件来源仓库地址。
|
||||
"""
|
||||
logger.info(
|
||||
f"执行工具: {self.name}, 参数: query={query}, force_refresh_market={force_refresh_market}"
|
||||
)
|
||||
try:
|
||||
installed_plugins = list_installed_plugins()
|
||||
if not installed_plugins:
|
||||
@@ -77,6 +88,10 @@ class QueryInstalledPluginsTool(MoviePilotTool):
|
||||
{"success": False, "message": "当前没有已安装的插件"},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
installed_plugins = await enrich_installed_plugin_sources(
|
||||
installed_plugins,
|
||||
force_refresh=force_refresh_market,
|
||||
)
|
||||
|
||||
limit = self._clamp_results(max_results)
|
||||
if query:
|
||||
|
||||
@@ -77,9 +77,12 @@ def _build_tv_server_result(existing_seasons: OrderedDict, total_seasons: Ordere
|
||||
|
||||
class QueryLibraryExistsInput(BaseModel):
|
||||
"""查询媒体库工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
tmdb_id: Optional[int] = Field(None, description="TMDB ID (can be obtained from search_media tool). Either tmdb_id or douban_id must be provided.")
|
||||
douban_id: Optional[str] = Field(None, description="Douban ID (can be obtained from search_media tool). Either tmdb_id or douban_id must be provided.")
|
||||
tmdb_id: Optional[int] = Field(None, description="TMDB media ID")
|
||||
douban_id: Optional[str] = Field(None, description="Douban media ID")
|
||||
bangumi_id: Optional[int] = Field(None, description="Bangumi media ID")
|
||||
anilist_id: Optional[int] = Field(None, description="AniList media ID")
|
||||
media_source: Optional[str] = Field(None, description="Media metadata source")
|
||||
media_id: Optional[str] = Field(None, description="Native ID for media_source")
|
||||
media_type: Optional[str] = Field(None, description="Allowed values: movie, tv")
|
||||
|
||||
|
||||
@@ -90,21 +93,24 @@ class QueryLibraryExistsTool(MoviePilotTool):
|
||||
ToolTag.Library,
|
||||
ToolTag.Media,
|
||||
]
|
||||
description: str = "Check whether media already exists in Plex, Emby, or Jellyfin by media ID. Results are grouped by media server; TV results include existing episodes, total episodes, and missing episodes/seasons. Requires tmdb_id or douban_id from search_media."
|
||||
description: str = "Check whether media already exists in Plex, Emby, or Jellyfin by a TMDB, Douban, Bangumi, AniList, or source-native media ID. Results are grouped by media server; TV results include existing episodes, total episodes, and missing episodes/seasons."
|
||||
args_schema: Type[BaseModel] = QueryLibraryExistsInput
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
"""根据查询参数生成友好的提示消息"""
|
||||
tmdb_id = kwargs.get("tmdb_id")
|
||||
douban_id = kwargs.get("douban_id")
|
||||
media_type = kwargs.get("media_type")
|
||||
|
||||
if tmdb_id:
|
||||
message = f"查询媒体库: TMDB={tmdb_id}"
|
||||
elif douban_id:
|
||||
message = f"查询媒体库: 豆瓣={douban_id}"
|
||||
else:
|
||||
message = "查询媒体库"
|
||||
identities = (
|
||||
("TMDB", kwargs.get("tmdb_id")),
|
||||
("豆瓣", kwargs.get("douban_id")),
|
||||
("Bangumi", kwargs.get("bangumi_id")),
|
||||
("AniList", kwargs.get("anilist_id")),
|
||||
(kwargs.get("media_source") or "媒体源", kwargs.get("media_id")),
|
||||
)
|
||||
label, identity = next(
|
||||
((label, identity) for label, identity in identities if identity is not None),
|
||||
(None, None),
|
||||
)
|
||||
message = f"查询媒体库: {label}={identity}" if label else "查询媒体库"
|
||||
if media_type:
|
||||
message += f" [{media_type}]"
|
||||
return message
|
||||
@@ -120,11 +126,13 @@ class QueryLibraryExistsTool(MoviePilotTool):
|
||||
return MediaServerChain().media_exists(mediainfo=mediainfo, server=server)
|
||||
|
||||
async def run(self, tmdb_id: Optional[int] = None, douban_id: Optional[str] = None,
|
||||
bangumi_id: Optional[int] = None, anilist_id: Optional[int] = None,
|
||||
media_source: Optional[str] = None, media_id: Optional[str] = None,
|
||||
media_type: Optional[str] = None, **kwargs) -> str:
|
||||
logger.info(f"执行工具: {self.name}, 参数: tmdb_id={tmdb_id}, douban_id={douban_id}, media_type={media_type}")
|
||||
try:
|
||||
if not tmdb_id and not douban_id:
|
||||
return "参数错误:tmdb_id 和 douban_id 至少需要提供一个,请先使用 search_media 工具获取媒体 ID。"
|
||||
if not any((tmdb_id, douban_id, bangumi_id, anilist_id, media_id)):
|
||||
return "参数错误:至少需要提供一个媒体 ID,请先使用 search_media 工具获取媒体信息。"
|
||||
|
||||
media_type_enum = None
|
||||
if media_type:
|
||||
@@ -136,11 +144,15 @@ class QueryLibraryExistsTool(MoviePilotTool):
|
||||
mediainfo = await media_chain.async_recognize_media(
|
||||
tmdbid=tmdb_id,
|
||||
doubanid=douban_id,
|
||||
bangumiid=bangumi_id,
|
||||
anilistid=anilist_id,
|
||||
source=media_source,
|
||||
mediaid=media_id,
|
||||
mtype=media_type_enum,
|
||||
)
|
||||
if not mediainfo:
|
||||
media_id = f"TMDB={tmdb_id}" if tmdb_id else f"豆瓣={douban_id}"
|
||||
return f"未识别到媒体信息: {media_id}"
|
||||
identity = media_id or tmdb_id or douban_id or bangumi_id or anilist_id
|
||||
return f"未识别到媒体信息: {identity}"
|
||||
|
||||
# 2. 遍历所有媒体服务器,分别查询存在性信息
|
||||
server_results = OrderedDict()
|
||||
|
||||
@@ -18,8 +18,6 @@ PAGE_SIZE = 20
|
||||
class QueryLibraryLatestInput(BaseModel):
|
||||
"""查询媒体服务器最近入库影片工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
server: Optional[str] = Field(
|
||||
None,
|
||||
description="Media server name (optional, if not specified queries all enabled media servers)",
|
||||
|
||||
@@ -21,8 +21,6 @@ from app.log import logger
|
||||
class QueryMarketPluginsInput(BaseModel):
|
||||
"""查询插件市场工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
query: Optional[str] = Field(
|
||||
None,
|
||||
description="Optional keyword to filter plugin market results by plugin ID, name, description, or author.",
|
||||
|
||||
@@ -18,9 +18,12 @@ SEASON_PREVIEW_LIMIT = 100
|
||||
|
||||
class QueryMediaDetailInput(BaseModel):
|
||||
"""查询媒体详情工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
tmdb_id: Optional[int] = Field(None, description="TMDB ID of the media (movie or TV series, can be obtained from search_media tool)")
|
||||
douban_id: Optional[str] = Field(None, description="Douban ID of the media (alternative to tmdb_id)")
|
||||
bangumi_id: Optional[int] = Field(None, description="Bangumi media ID")
|
||||
anilist_id: Optional[int] = Field(None, description="AniList media ID")
|
||||
media_source: Optional[str] = Field(None, description="Media metadata source")
|
||||
media_id: Optional[str] = Field(None, description="Native ID for media_source")
|
||||
media_type: str = Field(..., description="Allowed values: movie, tv")
|
||||
|
||||
|
||||
@@ -30,24 +33,37 @@ class QueryMediaDetailTool(MoviePilotTool):
|
||||
ToolTag.Read,
|
||||
ToolTag.Media,
|
||||
]
|
||||
description: str = "Query supplementary media details from TMDB by ID and media_type. Accepts tmdb_id or douban_id (at least one required). media_type accepts 'movie' or 'tv'. Returns non-duplicated detail fields such as status, genres, directors, actors, and season info for TV series."
|
||||
description: str = "Query supplementary media details from a metadata source by ID and media_type. Accepts a TMDB, Douban, Bangumi, AniList, or source-native media ID. media_type accepts 'movie' or 'tv'. Returns non-duplicated detail fields such as status, genres, directors, actors, and season info for TV series."
|
||||
args_schema: Type[BaseModel] = QueryMediaDetailInput
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
"""根据查询参数生成友好的提示消息"""
|
||||
tmdb_id = kwargs.get("tmdb_id")
|
||||
douban_id = kwargs.get("douban_id")
|
||||
if tmdb_id:
|
||||
return f"查询媒体详情: TMDB ID {tmdb_id}"
|
||||
return f"查询媒体详情: 豆瓣 ID {douban_id}"
|
||||
identities = (
|
||||
("TMDB", kwargs.get("tmdb_id")),
|
||||
("豆瓣", kwargs.get("douban_id")),
|
||||
("Bangumi", kwargs.get("bangumi_id")),
|
||||
("AniList", kwargs.get("anilist_id")),
|
||||
)
|
||||
for label, identity in identities:
|
||||
if identity is not None:
|
||||
return f"查询媒体详情: {label} ID {identity}"
|
||||
return (
|
||||
f"查询媒体详情: {kwargs.get('media_source') or '媒体源'} "
|
||||
f"ID {kwargs.get('media_id')}"
|
||||
)
|
||||
|
||||
async def run(self, media_type: str, tmdb_id: Optional[int] = None, douban_id: Optional[str] = None, **kwargs) -> str:
|
||||
async def run(
|
||||
self, media_type: str, tmdb_id: Optional[int] = None,
|
||||
douban_id: Optional[str] = None, bangumi_id: Optional[int] = None,
|
||||
anilist_id: Optional[int] = None, media_source: Optional[str] = None,
|
||||
media_id: Optional[str] = None, **kwargs,
|
||||
) -> str:
|
||||
logger.info(f"执行工具: {self.name}, 参数: tmdb_id={tmdb_id}, douban_id={douban_id}, media_type={media_type}")
|
||||
|
||||
if tmdb_id is None and douban_id is None:
|
||||
if not any((tmdb_id, douban_id, bangumi_id, anilist_id, media_id)):
|
||||
return json.dumps({
|
||||
"success": False,
|
||||
"message": "必须提供 tmdb_id 或 douban_id 之一"
|
||||
"message": "必须提供至少一个媒体 ID"
|
||||
}, ensure_ascii=False)
|
||||
|
||||
try:
|
||||
@@ -60,10 +76,22 @@ class QueryMediaDetailTool(MoviePilotTool):
|
||||
"message": f"无效的媒体类型 '{media_type}',支持的类型:'movie', 'tv'"
|
||||
}, ensure_ascii=False)
|
||||
|
||||
mediainfo = await media_chain.async_recognize_media(tmdbid=tmdb_id, doubanid=douban_id, mtype=media_type_enum)
|
||||
mediainfo = await media_chain.async_recognize_media(
|
||||
tmdbid=tmdb_id,
|
||||
doubanid=douban_id,
|
||||
bangumiid=bangumi_id,
|
||||
anilistid=anilist_id,
|
||||
source=media_source,
|
||||
mediaid=media_id,
|
||||
mtype=media_type_enum,
|
||||
)
|
||||
|
||||
if not mediainfo:
|
||||
id_info = f"TMDB ID {tmdb_id}" if tmdb_id else f"豆瓣 ID {douban_id}"
|
||||
id_info = (
|
||||
f"{media_source or '媒体源'} ID {media_id}"
|
||||
if media_id else
|
||||
f"媒体 ID {tmdb_id or douban_id or bangumi_id or anilist_id}"
|
||||
)
|
||||
return json.dumps({
|
||||
"success": False,
|
||||
"message": f"未找到 {id_info} 的媒体信息"
|
||||
@@ -140,5 +168,9 @@ class QueryMediaDetailTool(MoviePilotTool):
|
||||
"success": False,
|
||||
"message": error_message,
|
||||
"tmdb_id": tmdb_id,
|
||||
"douban_id": douban_id
|
||||
"douban_id": douban_id,
|
||||
"bangumi_id": bangumi_id,
|
||||
"anilist_id": anilist_id,
|
||||
"media_source": media_source,
|
||||
"media_id": media_id,
|
||||
}, ensure_ascii=False)
|
||||
|
||||
@@ -14,8 +14,6 @@ from app.log import logger
|
||||
class QueryPersonasInput(BaseModel):
|
||||
"""查询人格工具的输入参数模型。"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
query: Optional[str] = Field(
|
||||
None,
|
||||
description=(
|
||||
|
||||
@@ -14,8 +14,6 @@ from app.log import logger
|
||||
class QueryPluginCapabilitiesInput(BaseModel):
|
||||
"""查询插件能力工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
plugin_id: Optional[str] = Field(
|
||||
None,
|
||||
description="Optional plugin ID to query capabilities for a specific plugin. "
|
||||
|
||||
@@ -15,8 +15,6 @@ from app.log import logger
|
||||
class QueryPluginConfigInput(BaseModel):
|
||||
"""查询插件配置工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
plugin_id: str = Field(
|
||||
...,
|
||||
description="The plugin ID to query. Use query_installed_plugins first to discover valid plugin IDs.",
|
||||
|
||||
@@ -19,8 +19,6 @@ from app.log import logger
|
||||
class QueryPluginDataInput(BaseModel):
|
||||
"""查询插件数据工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
plugin_id: str = Field(
|
||||
...,
|
||||
description="The plugin ID to query. Use query_installed_plugins first to discover valid plugin IDs.",
|
||||
|
||||
@@ -18,7 +18,6 @@ MAX_PAGE_SIZE = 50
|
||||
|
||||
class QueryPopularSubscribesInput(BaseModel):
|
||||
"""查询热门订阅工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
media_type: str = Field(..., description="Allowed values: movie, tv")
|
||||
page: Optional[int] = Field(1, description="Page number for pagination (default: 1)")
|
||||
count: Optional[int] = Field(30, description="Number of items per page (default: 30, max: 50)")
|
||||
@@ -119,7 +118,7 @@ class QueryPopularSubscribesTool(MoviePilotTool):
|
||||
# 处理标题
|
||||
title = sub.get("name")
|
||||
season = sub.get("season")
|
||||
if season and int(season) > 1 and media.tmdb_id:
|
||||
if season not in (None, "") and int(season) != 1 and media.tmdb_id:
|
||||
# 小写数据转大写
|
||||
season_str = cn2an.an2cn(season, "low")
|
||||
title = f"{title} 第{season_str}季"
|
||||
@@ -127,6 +126,8 @@ class QueryPopularSubscribesTool(MoviePilotTool):
|
||||
media.year = sub.get("year")
|
||||
media.douban_id = sub.get("doubanid")
|
||||
media.bangumi_id = sub.get("bangumiid")
|
||||
media.anilist_id = sub.get("anilistid")
|
||||
media.source = sub.get("media_source")
|
||||
media.tvdb_id = sub.get("tvdbid")
|
||||
media.imdb_id = sub.get("imdbid")
|
||||
media.season = sub.get("season")
|
||||
@@ -150,6 +151,9 @@ class QueryPopularSubscribesTool(MoviePilotTool):
|
||||
"tmdb_id": media_dict.get("tmdb_id"),
|
||||
"douban_id": media_dict.get("douban_id"),
|
||||
"bangumi_id": media_dict.get("bangumi_id"),
|
||||
"anilist_id": media_dict.get("anilist_id"),
|
||||
"media_source": media_dict.get("source"),
|
||||
"media_id": media_dict.get("media_id"),
|
||||
"tvdb_id": media_dict.get("tvdb_id"),
|
||||
"imdb_id": media_dict.get("imdb_id"),
|
||||
"season": media_dict.get("season"),
|
||||
|
||||
@@ -19,8 +19,6 @@ from app.log import logger
|
||||
class QueryRuleGroupsInput(BaseModel):
|
||||
"""查询规则组工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
group_names: Optional[List[str]] = Field(
|
||||
None,
|
||||
description="Optional list of rule group names to query. If omitted, return all rule groups.",
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
import json
|
||||
from typing import Optional, Type
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
@@ -11,50 +11,69 @@ from app.log import logger
|
||||
|
||||
|
||||
class QuerySchedulersInput(BaseModel):
|
||||
"""查询定时服务工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
"""查询运行时定时服务的输入参数模型。"""
|
||||
|
||||
|
||||
class QuerySchedulersTool(MoviePilotTool):
|
||||
"""查询系统、插件和工作流注册的运行时定时服务。"""
|
||||
|
||||
name: str = "query_schedulers"
|
||||
tags: list[str] = [
|
||||
ToolTag.Read,
|
||||
ToolTag.Scheduler,
|
||||
ToolTag.Admin,
|
||||
]
|
||||
description: str = "Query scheduled tasks and list all available scheduler jobs. Shows job status, next run time, and provider information."
|
||||
description: str = (
|
||||
"Query runtime scheduler services registered by MoviePilot system components, "
|
||||
"plugins, and workflows. It excludes user-created autonomous agent tasks; use "
|
||||
"query_agent_tasks for reminders, monitoring tasks, and other agent schedules."
|
||||
)
|
||||
args_schema: Type[BaseModel] = QuerySchedulersInput
|
||||
require_admin: bool = True
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
"""生成友好的提示消息"""
|
||||
return "查询定时服务"
|
||||
def get_tool_message(self, **kwargs: object) -> Optional[str]:
|
||||
"""生成查询运行时定时服务的提示消息。"""
|
||||
return "查询系统定时服务"
|
||||
|
||||
async def run(self, **kwargs) -> str:
|
||||
async def run(self, **kwargs: object) -> str:
|
||||
"""查询非 Agent 自主任务的运行时定时服务。"""
|
||||
logger.info(f"执行工具: {self.name}")
|
||||
try:
|
||||
from app.scheduler import Scheduler
|
||||
from app.scheduler import AGENT_TASK_JOB_PREFIX, Scheduler
|
||||
|
||||
scheduler = Scheduler()
|
||||
schedulers = scheduler.list()
|
||||
agent_task_prefix = f"{AGENT_TASK_JOB_PREFIX}-"
|
||||
schedulers = [
|
||||
scheduler_item
|
||||
for scheduler_item in scheduler.list()
|
||||
if not str(scheduler_item.id or "").startswith(agent_task_prefix)
|
||||
]
|
||||
if schedulers:
|
||||
# 转换为字典列表以便JSON序列化
|
||||
schedulers_list = []
|
||||
for s in schedulers:
|
||||
schedulers_list.append({
|
||||
"id": s.id,
|
||||
"name": s.name,
|
||||
"provider": s.provider,
|
||||
"status": s.status,
|
||||
"next_run": s.next_run
|
||||
})
|
||||
schedulers_list = [
|
||||
{
|
||||
"id": scheduler_item.id,
|
||||
"name": scheduler_item.name,
|
||||
"provider": scheduler_item.provider,
|
||||
"status": scheduler_item.status,
|
||||
"next_run": scheduler_item.next_run,
|
||||
}
|
||||
for scheduler_item in schedulers
|
||||
]
|
||||
result_json = json.dumps(schedulers_list, ensure_ascii=False, indent=2)
|
||||
# 限制最多30条结果
|
||||
total_count = len(schedulers_list)
|
||||
if total_count > 30:
|
||||
limited_schedulers = schedulers_list[:30]
|
||||
limited_json = json.dumps(limited_schedulers, ensure_ascii=False, indent=2)
|
||||
return f"注意:查询结果共找到 {total_count} 条,为节省上下文空间,仅显示前 30 条结果。\n\n{limited_json}"
|
||||
limited_json = json.dumps(
|
||||
limited_schedulers,
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
)
|
||||
return (
|
||||
f"注意:查询结果共找到 {total_count} 条,为节省上下文空间,"
|
||||
f"仅显示前 30 条结果。\n\n{limited_json}"
|
||||
)
|
||||
return result_json
|
||||
return "未找到定时服务"
|
||||
return "未找到系统、插件或工作流定时服务"
|
||||
except Exception as e:
|
||||
logger.error(f"查询定时服务失败: {e}", exc_info=True)
|
||||
return f"查询定时服务时发生错误: {str(e)}"
|
||||
|
||||
@@ -7,9 +7,7 @@ from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.db import AsyncSessionFactory
|
||||
from app.db.models.site import Site
|
||||
from app.db.models.siteuserdata import SiteUserData
|
||||
from app.db.site_oper import SiteOper
|
||||
from app.log import logger
|
||||
|
||||
SITE_USERDATA_DETAIL_PREVIEW_LIMIT = 10
|
||||
@@ -24,8 +22,6 @@ def _preview_list(value, limit: int = SITE_USERDATA_DETAIL_PREVIEW_LIMIT) -> tup
|
||||
class QuerySiteUserdataInput(BaseModel):
|
||||
"""查询站点用户数据工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
site_id: int = Field(
|
||||
...,
|
||||
description="The ID of the site to query user data for (can be obtained from query_sites tool)",
|
||||
@@ -66,118 +62,115 @@ class QuerySiteUserdataTool(MoviePilotTool):
|
||||
)
|
||||
|
||||
try:
|
||||
# 获取数据库会话
|
||||
async with AsyncSessionFactory() as db:
|
||||
# 获取站点
|
||||
site = await Site.async_get(db, site_id)
|
||||
if not site:
|
||||
return json.dumps(
|
||||
{"success": False, "message": f"站点不存在: {site_id}"},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
|
||||
# 获取站点用户数据
|
||||
user_data_list = await SiteUserData.async_get_by_domain(
|
||||
db, domain=site.domain, workdate=workdate
|
||||
site_oper = SiteOper()
|
||||
site = await site_oper.async_get(site_id)
|
||||
if not site:
|
||||
return json.dumps(
|
||||
{"success": False, "message": f"站点不存在: {site_id}"},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
|
||||
if not user_data_list:
|
||||
return json.dumps(
|
||||
{
|
||||
"success": False,
|
||||
"message": f"站点 {site.name} ({site.domain}) 暂无用户数据",
|
||||
"site_id": site_id,
|
||||
"site_name": site.name,
|
||||
"site_domain": site.domain,
|
||||
"workdate": workdate,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
user_data_list = await site_oper.async_get_userdata_by_domain(
|
||||
domain=site.domain, workdate=workdate
|
||||
)
|
||||
|
||||
# 格式化用户数据
|
||||
result = {
|
||||
"success": True,
|
||||
"site_id": site_id,
|
||||
"site_name": site.name,
|
||||
"site_domain": site.domain,
|
||||
"workdate": workdate,
|
||||
"data_count": len(user_data_list),
|
||||
"user_data": [],
|
||||
if not user_data_list:
|
||||
return json.dumps(
|
||||
{
|
||||
"success": False,
|
||||
"message": f"站点 {site.name} ({site.domain}) 暂无用户数据",
|
||||
"site_id": site_id,
|
||||
"site_name": site.name,
|
||||
"site_domain": site.domain,
|
||||
"workdate": workdate,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
|
||||
# 格式化用户数据
|
||||
result = {
|
||||
"success": True,
|
||||
"site_id": site_id,
|
||||
"site_name": site.name,
|
||||
"site_domain": site.domain,
|
||||
"workdate": workdate,
|
||||
"data_count": len(user_data_list),
|
||||
"user_data": [],
|
||||
}
|
||||
|
||||
for user_data in user_data_list:
|
||||
# 格式化上传/下载量(转换为可读格式)
|
||||
upload_gb = user_data.upload / (1024**3) if user_data.upload else 0
|
||||
download_gb = (
|
||||
user_data.download / (1024**3) if user_data.download else 0
|
||||
)
|
||||
seeding_size_gb = (
|
||||
user_data.seeding_size / (1024**3)
|
||||
if user_data.seeding_size
|
||||
else 0
|
||||
)
|
||||
leeching_size_gb = (
|
||||
user_data.leeching_size / (1024**3)
|
||||
if user_data.leeching_size
|
||||
else 0
|
||||
)
|
||||
|
||||
seeding_preview, seeding_count, seeding_truncated = _preview_list(
|
||||
user_data.seeding_info
|
||||
)
|
||||
unread_preview, unread_count, unread_truncated = _preview_list(
|
||||
user_data.message_unread_contents
|
||||
)
|
||||
|
||||
user_data_dict = {
|
||||
"domain": user_data.domain,
|
||||
"name": user_data.name,
|
||||
"username": user_data.username,
|
||||
"userid": user_data.userid,
|
||||
"user_level": user_data.user_level,
|
||||
"join_at": user_data.join_at,
|
||||
"bonus": user_data.bonus,
|
||||
"upload": user_data.upload,
|
||||
"upload_gb": round(upload_gb, 2),
|
||||
"download": user_data.download,
|
||||
"download_gb": round(download_gb, 2),
|
||||
"ratio": round(user_data.ratio, 2) if user_data.ratio else 0,
|
||||
"seeding": int(user_data.seeding) if user_data.seeding else 0,
|
||||
"leeching": int(user_data.leeching)
|
||||
if user_data.leeching
|
||||
else 0,
|
||||
"seeding_size": user_data.seeding_size,
|
||||
"seeding_size_gb": round(seeding_size_gb, 2),
|
||||
"leeching_size": user_data.leeching_size,
|
||||
"leeching_size_gb": round(leeching_size_gb, 2),
|
||||
"seeding_info_count": seeding_count,
|
||||
"seeding_info": seeding_preview,
|
||||
"seeding_info_truncated": seeding_truncated,
|
||||
"message_unread": user_data.message_unread,
|
||||
"message_unread_contents_count": unread_count,
|
||||
"message_unread_contents": unread_preview,
|
||||
"message_unread_contents_truncated": unread_truncated,
|
||||
"err_msg": user_data.err_msg,
|
||||
"updated_day": user_data.updated_day,
|
||||
"updated_time": user_data.updated_time,
|
||||
}
|
||||
result["user_data"].append(user_data_dict)
|
||||
|
||||
for user_data in user_data_list:
|
||||
# 格式化上传/下载量(转换为可读格式)
|
||||
upload_gb = user_data.upload / (1024**3) if user_data.upload else 0
|
||||
download_gb = (
|
||||
user_data.download / (1024**3) if user_data.download else 0
|
||||
)
|
||||
seeding_size_gb = (
|
||||
user_data.seeding_size / (1024**3)
|
||||
if user_data.seeding_size
|
||||
else 0
|
||||
)
|
||||
leeching_size_gb = (
|
||||
user_data.leeching_size / (1024**3)
|
||||
if user_data.leeching_size
|
||||
else 0
|
||||
)
|
||||
# 如果有多条数据,只返回最新的(按更新时间排序)
|
||||
if len(result["user_data"]) > 1:
|
||||
result["user_data"].sort(
|
||||
key=lambda x: (
|
||||
x.get("updated_day", ""),
|
||||
x.get("updated_time", ""),
|
||||
),
|
||||
reverse=True,
|
||||
)
|
||||
result["message"] = (
|
||||
f"找到 {len(result['user_data'])} 条数据,显示最新的一条"
|
||||
)
|
||||
result["user_data"] = [result["user_data"][0]]
|
||||
|
||||
seeding_preview, seeding_count, seeding_truncated = _preview_list(
|
||||
user_data.seeding_info
|
||||
)
|
||||
unread_preview, unread_count, unread_truncated = _preview_list(
|
||||
user_data.message_unread_contents
|
||||
)
|
||||
|
||||
user_data_dict = {
|
||||
"domain": user_data.domain,
|
||||
"name": user_data.name,
|
||||
"username": user_data.username,
|
||||
"userid": user_data.userid,
|
||||
"user_level": user_data.user_level,
|
||||
"join_at": user_data.join_at,
|
||||
"bonus": user_data.bonus,
|
||||
"upload": user_data.upload,
|
||||
"upload_gb": round(upload_gb, 2),
|
||||
"download": user_data.download,
|
||||
"download_gb": round(download_gb, 2),
|
||||
"ratio": round(user_data.ratio, 2) if user_data.ratio else 0,
|
||||
"seeding": int(user_data.seeding) if user_data.seeding else 0,
|
||||
"leeching": int(user_data.leeching)
|
||||
if user_data.leeching
|
||||
else 0,
|
||||
"seeding_size": user_data.seeding_size,
|
||||
"seeding_size_gb": round(seeding_size_gb, 2),
|
||||
"leeching_size": user_data.leeching_size,
|
||||
"leeching_size_gb": round(leeching_size_gb, 2),
|
||||
"seeding_info_count": seeding_count,
|
||||
"seeding_info": seeding_preview,
|
||||
"seeding_info_truncated": seeding_truncated,
|
||||
"message_unread": user_data.message_unread,
|
||||
"message_unread_contents_count": unread_count,
|
||||
"message_unread_contents": unread_preview,
|
||||
"message_unread_contents_truncated": unread_truncated,
|
||||
"err_msg": user_data.err_msg,
|
||||
"updated_day": user_data.updated_day,
|
||||
"updated_time": user_data.updated_time,
|
||||
}
|
||||
result["user_data"].append(user_data_dict)
|
||||
|
||||
# 如果有多条数据,只返回最新的(按更新时间排序)
|
||||
if len(result["user_data"]) > 1:
|
||||
result["user_data"].sort(
|
||||
key=lambda x: (
|
||||
x.get("updated_day", ""),
|
||||
x.get("updated_time", ""),
|
||||
),
|
||||
reverse=True,
|
||||
)
|
||||
result["message"] = (
|
||||
f"找到 {len(result['user_data'])} 条数据,显示最新的一条"
|
||||
)
|
||||
result["user_data"] = [result["user_data"][0]]
|
||||
|
||||
return json.dumps(result, ensure_ascii=False, indent=2)
|
||||
return json.dumps(result, ensure_ascii=False, indent=2)
|
||||
|
||||
except Exception as e:
|
||||
error_message = f"查询站点用户数据失败: {str(e)}"
|
||||
|
||||
@@ -14,8 +14,6 @@ from app.log import logger
|
||||
class QuerySitesInput(BaseModel):
|
||||
"""查询站点工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
status: Optional[str] = Field(
|
||||
"all",
|
||||
description="Filter sites by status: 'active' for enabled sites, 'inactive' for disabled sites, 'all' for all sites",
|
||||
|
||||
@@ -7,8 +7,7 @@ from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.db import AsyncSessionFactory
|
||||
from app.db.models.subscribehistory import SubscribeHistory
|
||||
from app.db.subscribehistory_oper import SubscribeHistoryOper
|
||||
from app.log import logger
|
||||
from app.schemas.types import media_type_to_agent
|
||||
|
||||
@@ -18,8 +17,6 @@ PAGE_SIZE = 20
|
||||
class QuerySubscribeHistoryInput(BaseModel):
|
||||
"""查询订阅历史工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
media_type: Optional[str] = Field(
|
||||
"all", description="Allowed values: movie, tv, all"
|
||||
)
|
||||
@@ -74,88 +71,87 @@ class QuerySubscribeHistoryTool(MoviePilotTool):
|
||||
if media_type not in ["all", "movie", "tv"]:
|
||||
return f"错误:无效的媒体类型 '{media_type}',支持的类型:'movie', 'tv', 'all'"
|
||||
|
||||
# 获取数据库会话
|
||||
async with AsyncSessionFactory() as db:
|
||||
if name:
|
||||
# 有名称过滤时,获取足够多的记录在内存中过滤,不分页
|
||||
fetch_count = 500
|
||||
if media_type == "all":
|
||||
movie_history = await SubscribeHistory.async_list_by_type(
|
||||
db, mtype="movie", page=1, count=fetch_count
|
||||
)
|
||||
tv_history = await SubscribeHistory.async_list_by_type(
|
||||
db, mtype="tv", page=1, count=fetch_count
|
||||
)
|
||||
all_history = list(movie_history) + list(tv_history)
|
||||
all_history.sort(key=lambda x: x.date or "", reverse=True)
|
||||
else:
|
||||
all_history = list(
|
||||
await SubscribeHistory.async_list_by_type(
|
||||
db, mtype=media_type, page=1, count=fetch_count
|
||||
)
|
||||
)
|
||||
|
||||
# 按名称过滤
|
||||
name_lower = name.lower()
|
||||
filtered_history = [
|
||||
record
|
||||
for record in all_history
|
||||
if record.name and name_lower in record.name.lower()
|
||||
]
|
||||
|
||||
if not filtered_history:
|
||||
return "未找到相关订阅历史记录"
|
||||
|
||||
# 名称过滤时直接返回所有匹配结果,不分页
|
||||
simplified_records = self._simplify_records(filtered_history)
|
||||
result_json = json.dumps(
|
||||
simplified_records, ensure_ascii=False, indent=2
|
||||
subscribe_history_oper = SubscribeHistoryOper()
|
||||
if name:
|
||||
# 有名称过滤时,获取足够多的记录在内存中过滤,不分页
|
||||
fetch_count = 500
|
||||
if media_type == "all":
|
||||
movie_history = await subscribe_history_oper.async_list_by_type(
|
||||
mtype="movie", page=1, count=fetch_count
|
||||
)
|
||||
return result_json
|
||||
tv_history = await subscribe_history_oper.async_list_by_type(
|
||||
mtype="tv", page=1, count=fetch_count
|
||||
)
|
||||
all_history = list(movie_history) + list(tv_history)
|
||||
all_history.sort(key=lambda x: x.date or "", reverse=True)
|
||||
else:
|
||||
# 无名称过滤时,直接利用数据库分页
|
||||
if media_type == "all":
|
||||
movie_history = await SubscribeHistory.async_list_by_type(
|
||||
db, mtype="movie", page=1, count=page * PAGE_SIZE
|
||||
)
|
||||
tv_history = await SubscribeHistory.async_list_by_type(
|
||||
db, mtype="tv", page=1, count=page * PAGE_SIZE
|
||||
)
|
||||
all_history = list(movie_history) + list(tv_history)
|
||||
all_history.sort(key=lambda x: x.date or "", reverse=True)
|
||||
filtered_history = all_history
|
||||
else:
|
||||
filtered_history = list(
|
||||
await SubscribeHistory.async_list_by_type(
|
||||
db, mtype=media_type, page=1, count=page * PAGE_SIZE
|
||||
)
|
||||
all_history = list(
|
||||
await subscribe_history_oper.async_list_by_type(
|
||||
mtype=media_type, page=1, count=fetch_count
|
||||
)
|
||||
)
|
||||
|
||||
# 按名称过滤
|
||||
name_lower = name.lower()
|
||||
filtered_history = [
|
||||
record
|
||||
for record in all_history
|
||||
if record.name and name_lower in record.name.lower()
|
||||
]
|
||||
|
||||
if not filtered_history:
|
||||
return "未找到相关订阅历史记录"
|
||||
|
||||
# 分页切片
|
||||
total_count = len(filtered_history)
|
||||
start = (page - 1) * PAGE_SIZE
|
||||
end = start + PAGE_SIZE
|
||||
page_records = filtered_history[start:end]
|
||||
|
||||
if not page_records:
|
||||
return f"第 {page} 页没有数据。"
|
||||
|
||||
simplified_records = self._simplify_records(page_records)
|
||||
# 名称过滤时直接返回所有匹配结果,不分页
|
||||
simplified_records = self._simplify_records(filtered_history)
|
||||
result_json = json.dumps(
|
||||
simplified_records, ensure_ascii=False, indent=2
|
||||
)
|
||||
|
||||
has_more = total_count > end
|
||||
payload_msg = f"第 {page} 页,当前页 {len(simplified_records)} 条结果。"
|
||||
if has_more:
|
||||
payload_msg += (
|
||||
f" 可能有更多数据,可使用 page={page + 1} 获取下一页。"
|
||||
return result_json
|
||||
else:
|
||||
# 无名称过滤时,直接利用数据库分页
|
||||
if media_type == "all":
|
||||
movie_history = await subscribe_history_oper.async_list_by_type(
|
||||
mtype="movie", page=1, count=page * PAGE_SIZE
|
||||
)
|
||||
tv_history = await subscribe_history_oper.async_list_by_type(
|
||||
mtype="tv", page=1, count=page * PAGE_SIZE
|
||||
)
|
||||
all_history = list(movie_history) + list(tv_history)
|
||||
all_history.sort(key=lambda x: x.date or "", reverse=True)
|
||||
filtered_history = all_history
|
||||
else:
|
||||
filtered_history = list(
|
||||
await subscribe_history_oper.async_list_by_type(
|
||||
mtype=media_type, page=1, count=page * PAGE_SIZE
|
||||
)
|
||||
)
|
||||
|
||||
return f"{payload_msg}\n\n{result_json}"
|
||||
if not filtered_history:
|
||||
return "未找到相关订阅历史记录"
|
||||
|
||||
# 分页切片
|
||||
total_count = len(filtered_history)
|
||||
start = (page - 1) * PAGE_SIZE
|
||||
end = start + PAGE_SIZE
|
||||
page_records = filtered_history[start:end]
|
||||
|
||||
if not page_records:
|
||||
return f"第 {page} 页没有数据。"
|
||||
|
||||
simplified_records = self._simplify_records(page_records)
|
||||
result_json = json.dumps(
|
||||
simplified_records, ensure_ascii=False, indent=2
|
||||
)
|
||||
|
||||
has_more = total_count > end
|
||||
payload_msg = f"第 {page} 页,当前页 {len(simplified_records)} 条结果。"
|
||||
if has_more:
|
||||
payload_msg += (
|
||||
f" 可能有更多数据,可使用 page={page + 1} 获取下一页。"
|
||||
)
|
||||
|
||||
return f"{payload_msg}\n\n{result_json}"
|
||||
except Exception as e:
|
||||
logger.error(f"查询订阅历史失败: {e}", exc_info=True)
|
||||
return f"查询订阅历史时发生错误: {str(e)}"
|
||||
@@ -174,6 +170,9 @@ class QuerySubscribeHistoryTool(MoviePilotTool):
|
||||
"tmdbid": record.tmdbid,
|
||||
"doubanid": record.doubanid,
|
||||
"bangumiid": record.bangumiid,
|
||||
"anilistid": record.anilistid,
|
||||
"media_source": record.media_source,
|
||||
"media_id": record.media_id,
|
||||
"poster": record.poster,
|
||||
"vote": record.vote,
|
||||
"total_episode": record.total_episode,
|
||||
|
||||
@@ -15,7 +15,6 @@ MAX_PAGE_SIZE = 50
|
||||
|
||||
class QuerySubscribeSharesInput(BaseModel):
|
||||
"""查询订阅分享工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
name: Optional[str] = Field(None, description="Filter shares by media name (partial match, optional)")
|
||||
page: Optional[int] = Field(1, description="Page number for pagination (default: 1)")
|
||||
count: Optional[int] = Field(30, description="Number of items per page (default: 30, max: 50)")
|
||||
@@ -98,6 +97,9 @@ class QuerySubscribeSharesTool(MoviePilotTool):
|
||||
"tmdbid": share.get("tmdbid"),
|
||||
"doubanid": share.get("doubanid"),
|
||||
"bangumiid": share.get("bangumiid"),
|
||||
"anilistid": share.get("anilistid"),
|
||||
"media_source": share.get("media_source"),
|
||||
"media_id": share.get("media_id"),
|
||||
"poster": share.get("poster"),
|
||||
"vote": share.get("vote"),
|
||||
"share_title": share.get("share_title"),
|
||||
|
||||
@@ -48,8 +48,6 @@ QUERY_SUBSCRIBE_OUTPUT_FIELDS = [
|
||||
class QuerySubscribesInput(BaseModel):
|
||||
"""查询订阅工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
status: Optional[str] = Field(
|
||||
"all",
|
||||
description="Filter subscriptions by status: 'R' for enabled subscriptions, 'S' for paused ones, 'all' for all subscriptions",
|
||||
@@ -65,6 +63,10 @@ class QuerySubscribesInput(BaseModel):
|
||||
None,
|
||||
description="Filter by Douban ID to check if a specific media is already subscribed",
|
||||
)
|
||||
bangumi_id: Optional[int] = Field(None, description="Filter by Bangumi ID")
|
||||
anilist_id: Optional[int] = Field(None, description="Filter by AniList ID")
|
||||
media_source: Optional[str] = Field(None, description="Filter by media source")
|
||||
media_id: Optional[str] = Field(None, description="Filter by source-native media ID")
|
||||
page: Optional[int] = Field(
|
||||
1, description="Page number for pagination (default: 1, 100 items per page)"
|
||||
)
|
||||
@@ -106,6 +108,10 @@ class QuerySubscribesTool(MoviePilotTool):
|
||||
media_type: Optional[str] = "all",
|
||||
tmdb_id: Optional[int] = None,
|
||||
douban_id: Optional[str] = None,
|
||||
bangumi_id: Optional[int] = None,
|
||||
anilist_id: Optional[int] = None,
|
||||
media_source: Optional[str] = None,
|
||||
media_id: Optional[str] = None,
|
||||
page: Optional[int] = 1,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
@@ -132,6 +138,14 @@ class QuerySubscribesTool(MoviePilotTool):
|
||||
continue
|
||||
if douban_id is not None and sub.doubanid != douban_id:
|
||||
continue
|
||||
if bangumi_id is not None and sub.bangumiid != bangumi_id:
|
||||
continue
|
||||
if anilist_id is not None and sub.anilistid != anilist_id:
|
||||
continue
|
||||
if media_source is not None and sub.media_source != media_source:
|
||||
continue
|
||||
if media_id is not None and sub.media_id != media_id:
|
||||
continue
|
||||
filtered_subscribes.append(sub)
|
||||
if filtered_subscribes:
|
||||
total_count = len(filtered_subscribes)
|
||||
|
||||
@@ -9,8 +9,11 @@ from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.agent.tools.impl._system_setting_utils import (
|
||||
SettingSpec,
|
||||
is_secret_setting_key,
|
||||
list_setting_specs,
|
||||
redact_secret_value,
|
||||
resolve_setting_spec,
|
||||
should_redact_setting,
|
||||
)
|
||||
from app.core.config import settings
|
||||
from app.db.systemconfig_oper import SystemConfigOper
|
||||
@@ -20,8 +23,6 @@ from app.log import logger
|
||||
class QuerySystemSettingsInput(BaseModel):
|
||||
"""查询系统设置工具的输入参数模型。"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
setting_key: Optional[str] = Field(
|
||||
None,
|
||||
description=(
|
||||
@@ -53,6 +54,13 @@ class QuerySystemSettingsInput(BaseModel):
|
||||
"when multiple settings are matched it returns summaries only unless this is explicitly set to true."
|
||||
),
|
||||
)
|
||||
show_secrets: Optional[bool] = Field(
|
||||
False,
|
||||
description=(
|
||||
"Whether to return raw secret values such as API keys, tokens, cookies, and passwords. "
|
||||
"Defaults to false; secret-like fields are redacted in returned values and previews."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class QuerySystemSettingsTool(MoviePilotTool):
|
||||
@@ -85,15 +93,18 @@ class QuerySystemSettingsTool(MoviePilotTool):
|
||||
|
||||
@staticmethod
|
||||
def _load_setting_value(spec: SettingSpec):
|
||||
"""读取指定设置项的当前值。"""
|
||||
if spec.source == "settings":
|
||||
return getattr(settings, spec.key)
|
||||
return SystemConfigOper().get(spec.key)
|
||||
return SystemConfigOper().get(spec.systemconfig_key)
|
||||
|
||||
@staticmethod
|
||||
def _summarize_value(value) -> dict:
|
||||
def _summarize_value(value, *, redacted: bool = False) -> dict:
|
||||
"""生成设置值摘要,避免列表和字典默认输出过长。"""
|
||||
summary = {
|
||||
"has_value": value is not None,
|
||||
"value_type": type(value).__name__,
|
||||
"redacted": redacted,
|
||||
}
|
||||
if isinstance(value, list):
|
||||
summary["item_count"] = len(value)
|
||||
@@ -122,14 +133,12 @@ class QuerySystemSettingsTool(MoviePilotTool):
|
||||
group: Optional[str] = "all",
|
||||
keyword: Optional[str] = None,
|
||||
include_values: Optional[bool] = None,
|
||||
show_secrets: Optional[bool] = False,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
logger.info(
|
||||
"执行工具: %s, setting_key=%s, group=%s, keyword=%s",
|
||||
self.name,
|
||||
setting_key,
|
||||
group,
|
||||
keyword,
|
||||
f"执行工具: {self.name}, setting_key={setting_key}, "
|
||||
f"group={group}, keyword={keyword}"
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -158,18 +167,30 @@ class QuerySystemSettingsTool(MoviePilotTool):
|
||||
should_include_values = (
|
||||
include_values if include_values is not None else len(specs) == 1
|
||||
)
|
||||
allow_secret_values = bool(show_secrets) and await self.is_admin_user()
|
||||
settings_payload = []
|
||||
for spec in specs:
|
||||
value = self._load_setting_value(spec)
|
||||
should_redact = (
|
||||
should_redact_setting(spec, value) and not allow_secret_values
|
||||
)
|
||||
response_value = (
|
||||
redact_secret_value(
|
||||
value,
|
||||
redact_scalar=is_secret_setting_key(spec.key),
|
||||
)
|
||||
if should_redact
|
||||
else value
|
||||
)
|
||||
item = {
|
||||
"setting_key": spec.key,
|
||||
"source": spec.source,
|
||||
"group": spec.group,
|
||||
"label": spec.label,
|
||||
}
|
||||
item.update(self._summarize_value(value))
|
||||
item.update(self._summarize_value(response_value, redacted=should_redact))
|
||||
if should_include_values:
|
||||
item["value"] = value
|
||||
item["value"] = response_value
|
||||
settings_payload.append(item)
|
||||
|
||||
return json.dumps(
|
||||
@@ -177,6 +198,7 @@ class QuerySystemSettingsTool(MoviePilotTool):
|
||||
"success": True,
|
||||
"matched_count": len(settings_payload),
|
||||
"include_values": should_include_values,
|
||||
"show_secrets": allow_secret_values,
|
||||
"settings": settings_payload,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
|
||||
@@ -7,8 +7,7 @@ from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.db import AsyncSessionFactory
|
||||
from app.db.models.transferhistory import TransferHistory
|
||||
from app.db.transferhistory_oper import TransferHistoryOper
|
||||
from app.log import logger
|
||||
from app.schemas.types import media_type_to_agent
|
||||
from app.utils.jieba import cut as jieba_cut
|
||||
@@ -16,7 +15,6 @@ from app.utils.jieba import cut as jieba_cut
|
||||
|
||||
class QueryTransferHistoryInput(BaseModel):
|
||||
"""查询整理历史记录工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
title: Optional[str] = Field(None, description="Search by title (optional, supports partial match)")
|
||||
status: Optional[str] = Field("all",
|
||||
description="Filter by status: 'success' for successful transfers, 'failed' for failed transfers, 'all' for all records (default: 'all')")
|
||||
@@ -70,70 +68,77 @@ class QueryTransferHistoryTool(MoviePilotTool):
|
||||
# 每页固定 30 条,与工具说明保持一致,避免整理路径等字段撑大上下文。
|
||||
count = 30
|
||||
|
||||
# 获取数据库会话
|
||||
async with AsyncSessionFactory() as db:
|
||||
# 处理标题搜索
|
||||
if title:
|
||||
# 使用统一分词封装处理标题,便于替换底层实现。
|
||||
words = jieba_cut(title, HMM=False)
|
||||
title_search = "%".join(words)
|
||||
# 查询记录
|
||||
result = await TransferHistory.async_list_by_title(
|
||||
db, title=title_search, page=page, count=count, status=status_bool
|
||||
)
|
||||
total = await TransferHistory.async_count_by_title(
|
||||
db, title=title_search, status=status_bool
|
||||
)
|
||||
else:
|
||||
# 查询所有记录
|
||||
result = await TransferHistory.async_list_by_page(
|
||||
db, page=page, count=count, status=status_bool
|
||||
)
|
||||
total = await TransferHistory.async_count(db, status=status_bool)
|
||||
transferhis = TransferHistoryOper()
|
||||
# 处理标题搜索
|
||||
if title:
|
||||
# 使用统一分词封装处理标题,便于替换底层实现。
|
||||
words = jieba_cut(title, HMM=False)
|
||||
title_search = "%".join(words)
|
||||
# 查询记录
|
||||
result = await transferhis.async_list_by_title(
|
||||
title=title_search, page=page, count=count, status=status_bool
|
||||
)
|
||||
total = await transferhis.async_count_by_title(
|
||||
title=title_search, status=status_bool
|
||||
)
|
||||
else:
|
||||
# 查询所有记录
|
||||
result = await transferhis.async_list_by_page(
|
||||
page=page, count=count, status=status_bool
|
||||
)
|
||||
total = await transferhis.async_count(status=status_bool)
|
||||
|
||||
if not result:
|
||||
return "未找到相关整理历史记录"
|
||||
if not result:
|
||||
return "未找到相关整理历史记录"
|
||||
|
||||
# 转换为字典格式,只保留关键信息
|
||||
simplified_records = []
|
||||
for record in result:
|
||||
simplified = {
|
||||
"id": record.id,
|
||||
"title": record.title,
|
||||
"year": record.year,
|
||||
"type": media_type_to_agent(record.type),
|
||||
"category": record.category,
|
||||
"seasons": record.seasons,
|
||||
"episodes": record.episodes,
|
||||
"src": record.src,
|
||||
"dest": record.dest,
|
||||
"mode": record.mode,
|
||||
"status": "成功" if record.status else "失败",
|
||||
"date": record.date,
|
||||
"downloader": record.downloader,
|
||||
"download_hash": record.download_hash
|
||||
}
|
||||
# 如果失败,添加错误信息
|
||||
if not record.status and record.errmsg:
|
||||
simplified["errmsg"] = record.errmsg
|
||||
# 添加媒体ID信息(如果有)
|
||||
if record.tmdbid:
|
||||
simplified["tmdbid"] = record.tmdbid
|
||||
if record.imdbid:
|
||||
simplified["imdbid"] = record.imdbid
|
||||
if record.doubanid:
|
||||
simplified["doubanid"] = record.doubanid
|
||||
simplified_records.append(simplified)
|
||||
# 转换为字典格式,只保留关键信息
|
||||
simplified_records = []
|
||||
for record in result:
|
||||
simplified = {
|
||||
"id": record.id,
|
||||
"title": record.title,
|
||||
"year": record.year,
|
||||
"type": media_type_to_agent(record.type),
|
||||
"category": record.category,
|
||||
"seasons": record.seasons,
|
||||
"episodes": record.episodes,
|
||||
"src": record.src,
|
||||
"dest": record.dest,
|
||||
"mode": record.mode,
|
||||
"status": "成功" if record.status else "失败",
|
||||
"date": record.date,
|
||||
"downloader": record.downloader,
|
||||
"download_hash": record.download_hash
|
||||
}
|
||||
# 如果失败,添加错误信息
|
||||
if not record.status and record.errmsg:
|
||||
simplified["errmsg"] = record.errmsg
|
||||
# 添加媒体ID信息(如果有)
|
||||
if record.tmdbid:
|
||||
simplified["tmdbid"] = record.tmdbid
|
||||
if record.imdbid:
|
||||
simplified["imdbid"] = record.imdbid
|
||||
if record.doubanid:
|
||||
simplified["doubanid"] = record.doubanid
|
||||
if record.bangumiid:
|
||||
simplified["bangumiid"] = record.bangumiid
|
||||
if record.anilistid:
|
||||
simplified["anilistid"] = record.anilistid
|
||||
if record.media_source:
|
||||
simplified["media_source"] = record.media_source
|
||||
if record.media_id:
|
||||
simplified["media_id"] = record.media_id
|
||||
simplified_records.append(simplified)
|
||||
|
||||
result_json = json.dumps(simplified_records, ensure_ascii=False, indent=2)
|
||||
result_json = json.dumps(simplified_records, ensure_ascii=False, indent=2)
|
||||
|
||||
# 计算总页数
|
||||
total_pages = (total + count - 1) // count if total > 0 else 1
|
||||
# 计算总页数
|
||||
total_pages = (total + count - 1) // count if total > 0 else 1
|
||||
|
||||
# 构建分页信息
|
||||
pagination_info = f"第 {page}/{total_pages} 页,共 {total} 条记录(每页 {count} 条)"
|
||||
# 构建分页信息
|
||||
pagination_info = f"第 {page}/{total_pages} 页,共 {total} 条记录(每页 {count} 条)"
|
||||
|
||||
return f"{pagination_info}\n\n{result_json}"
|
||||
return f"{pagination_info}\n\n{result_json}"
|
||||
except Exception as e:
|
||||
logger.error(f"查询整理历史记录失败: {e}", exc_info=True)
|
||||
return f"查询整理历史记录时发生错误: {str(e)}"
|
||||
|
||||
@@ -7,14 +7,12 @@ from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.db import AsyncSessionFactory
|
||||
from app.db.workflow_oper import WorkflowOper
|
||||
from app.log import logger
|
||||
|
||||
|
||||
class QueryWorkflowsInput(BaseModel):
|
||||
"""查询工作流工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
state: Optional[str] = Field("all", description="Filter workflows by state: 'W' for waiting, 'R' for running, 'P' for paused, 'S' for success, 'F' for failed, 'all' for all workflows (default: 'all')")
|
||||
name: Optional[str] = Field(None, description="Filter workflows by name (partial match, optional)")
|
||||
trigger_type: Optional[str] = Field("all", description="Filter workflows by trigger type: 'timer' for scheduled, 'event' for event-triggered, 'manual' for manual, 'all' for all types (default: 'all')")
|
||||
@@ -56,75 +54,73 @@ class QueryWorkflowsTool(MoviePilotTool):
|
||||
logger.info(f"执行工具: {self.name}, 参数: state={state}, name={name}, trigger_type={trigger_type}")
|
||||
|
||||
try:
|
||||
# 获取数据库会话
|
||||
async with AsyncSessionFactory() as db:
|
||||
workflow_oper = WorkflowOper(db)
|
||||
workflows = await workflow_oper.async_list()
|
||||
|
||||
# 过滤工作流
|
||||
filtered_workflows = []
|
||||
for wf in workflows:
|
||||
# 按状态过滤
|
||||
if state != "all" and wf.state != state:
|
||||
workflow_oper = WorkflowOper()
|
||||
workflows = await workflow_oper.async_list()
|
||||
|
||||
# 过滤工作流
|
||||
filtered_workflows = []
|
||||
for wf in workflows:
|
||||
# 按状态过滤
|
||||
if state != "all" and wf.state != state:
|
||||
continue
|
||||
|
||||
# 按触发类型过滤
|
||||
if trigger_type != "all":
|
||||
if trigger_type == "timer" and wf.trigger_type not in ["timer", None]:
|
||||
continue
|
||||
|
||||
# 按触发类型过滤
|
||||
if trigger_type != "all":
|
||||
if trigger_type == "timer" and wf.trigger_type not in ["timer", None]:
|
||||
continue
|
||||
elif trigger_type == "event" and wf.trigger_type != "event":
|
||||
continue
|
||||
elif trigger_type == "manual" and wf.trigger_type != "manual":
|
||||
continue
|
||||
|
||||
# 按名称过滤(部分匹配)
|
||||
if name and wf.name and name.lower() not in wf.name.lower():
|
||||
elif trigger_type == "event" and wf.trigger_type != "event":
|
||||
continue
|
||||
|
||||
filtered_workflows.append(wf)
|
||||
|
||||
if not filtered_workflows:
|
||||
return "未找到相关工作流"
|
||||
|
||||
# 转换为字典格式,只保留关键信息
|
||||
simplified_workflows = []
|
||||
for wf in filtered_workflows:
|
||||
# 状态说明
|
||||
state_map = {
|
||||
"W": "等待",
|
||||
"R": "运行中",
|
||||
"P": "暂停",
|
||||
"S": "成功",
|
||||
"F": "失败"
|
||||
}
|
||||
state_desc = state_map.get(wf.state, wf.state)
|
||||
|
||||
# 触发类型说明
|
||||
trigger_type_map = {
|
||||
"timer": "定时触发",
|
||||
"event": "事件触发",
|
||||
"manual": "手动触发"
|
||||
}
|
||||
trigger_type_desc = trigger_type_map.get(wf.trigger_type, wf.trigger_type or "定时触发")
|
||||
|
||||
simplified = {
|
||||
"id": wf.id,
|
||||
"name": wf.name,
|
||||
"description": wf.description,
|
||||
"trigger_type": trigger_type_desc,
|
||||
"state": state_desc,
|
||||
"run_count": wf.run_count,
|
||||
"timer": wf.timer,
|
||||
"event_type": wf.event_type,
|
||||
"add_time": wf.add_time,
|
||||
"last_time": wf.last_time,
|
||||
"current_action": wf.current_action
|
||||
}
|
||||
# wf.result 往往是执行日志或上下文快照,不适合作为列表查询结果返回。
|
||||
simplified_workflows.append(simplified)
|
||||
|
||||
result_json = json.dumps(simplified_workflows, ensure_ascii=False, indent=2)
|
||||
return result_json
|
||||
elif trigger_type == "manual" and wf.trigger_type != "manual":
|
||||
continue
|
||||
|
||||
# 按名称过滤(部分匹配)
|
||||
if name and wf.name and name.lower() not in wf.name.lower():
|
||||
continue
|
||||
|
||||
filtered_workflows.append(wf)
|
||||
|
||||
if not filtered_workflows:
|
||||
return "未找到相关工作流"
|
||||
|
||||
# 转换为字典格式,只保留关键信息
|
||||
simplified_workflows = []
|
||||
for wf in filtered_workflows:
|
||||
# 状态说明
|
||||
state_map = {
|
||||
"W": "等待",
|
||||
"R": "运行中",
|
||||
"P": "暂停",
|
||||
"S": "成功",
|
||||
"F": "失败"
|
||||
}
|
||||
state_desc = state_map.get(wf.state, wf.state)
|
||||
|
||||
# 触发类型说明
|
||||
trigger_type_map = {
|
||||
"timer": "定时触发",
|
||||
"event": "事件触发",
|
||||
"manual": "手动触发"
|
||||
}
|
||||
trigger_type_desc = trigger_type_map.get(wf.trigger_type, wf.trigger_type or "定时触发")
|
||||
|
||||
simplified = {
|
||||
"id": wf.id,
|
||||
"name": wf.name,
|
||||
"description": wf.description,
|
||||
"trigger_type": trigger_type_desc,
|
||||
"state": state_desc,
|
||||
"run_count": wf.run_count,
|
||||
"timer": wf.timer,
|
||||
"event_type": wf.event_type,
|
||||
"add_time": wf.add_time,
|
||||
"last_time": wf.last_time,
|
||||
"current_action": wf.current_action
|
||||
}
|
||||
# wf.result 往往是执行日志或上下文快照,不适合作为列表查询结果返回。
|
||||
simplified_workflows.append(simplified)
|
||||
|
||||
result_json = json.dumps(simplified_workflows, ensure_ascii=False, indent=2)
|
||||
return result_json
|
||||
except Exception as e:
|
||||
logger.error(f"查询工作流失败: {e}", exc_info=True)
|
||||
return f"查询工作流时发生错误: {str(e)}"
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""文件读取工具"""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Optional, Type
|
||||
|
||||
@@ -15,13 +17,23 @@ MAX_READ_SIZE = 50 * 1024
|
||||
|
||||
|
||||
class ReadFileInput(BaseModel):
|
||||
"""Input parameters for read file tool"""
|
||||
"""文件读取工具的输入参数模型。"""
|
||||
|
||||
file_path: str = Field(..., description="The absolute path of the file to read")
|
||||
start_line: Optional[int] = Field(None, description="The starting line number (1-based, inclusive). If not provided, reading starts from the beginning of the file.")
|
||||
end_line: Optional[int] = Field(None, description="The ending line number (1-based, inclusive). If not provided, reading goes until the end of the file.")
|
||||
include_metadata: bool = Field(
|
||||
False,
|
||||
description=(
|
||||
"Return structured JSON containing content, size, truncation state, "
|
||||
"and SHA-256. Use before a guarded full-file overwrite."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class ReadFileTool(MoviePilotTool):
|
||||
"""按行范围读取本地文本文件,并可返回文件版本元数据。"""
|
||||
|
||||
name: str = "read_file"
|
||||
tags: list[str] = [
|
||||
ToolTag.Read,
|
||||
@@ -36,8 +48,15 @@ class ReadFileTool(MoviePilotTool):
|
||||
file_name = Path(file_path).name if file_path else "未知文件"
|
||||
return f"读取文件: {file_name}"
|
||||
|
||||
async def run(self, file_path: str, start_line: Optional[int] = None,
|
||||
end_line: Optional[int] = None, **kwargs) -> str:
|
||||
async def run(
|
||||
self,
|
||||
file_path: str,
|
||||
start_line: Optional[int] = None,
|
||||
end_line: Optional[int] = None,
|
||||
include_metadata: bool = False,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""读取指定文本范围,必要时附带完整文件的 SHA-256 元数据。"""
|
||||
logger.info(f"执行工具: {self.name}, 参数: file_path={file_path}, start_line={start_line}, end_line={end_line}")
|
||||
|
||||
try:
|
||||
@@ -55,7 +74,8 @@ class ReadFileTool(MoviePilotTool):
|
||||
if not await path.is_file():
|
||||
return f"错误:{resolved_path} 不是一个文件"
|
||||
|
||||
content = await path.read_text(encoding="utf-8", errors="replace")
|
||||
raw_content = await path.read_bytes()
|
||||
content = raw_content.decode("utf-8", errors="replace")
|
||||
truncated = False
|
||||
|
||||
if start_line is not None or end_line is not None:
|
||||
@@ -78,6 +98,21 @@ class ReadFileTool(MoviePilotTool):
|
||||
content = content_bytes[:MAX_READ_SIZE].decode("utf-8", errors="replace")
|
||||
truncated = True
|
||||
|
||||
if include_metadata:
|
||||
return json.dumps(
|
||||
{
|
||||
"file_path": str(resolved_path),
|
||||
"sha256": hashlib.sha256(raw_content).hexdigest(),
|
||||
"size_bytes": len(raw_content),
|
||||
"start_line": start_line,
|
||||
"end_line": end_line,
|
||||
"truncated": truncated,
|
||||
"content": content,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
)
|
||||
|
||||
if truncated:
|
||||
return f"{content}\n\n[警告:文件内容已超过50KB限制,以上内容已被截断。请使用 start_line/end_line 参数分段读取。]"
|
||||
|
||||
|
||||
@@ -15,10 +15,6 @@ from app.log import logger
|
||||
class RecognizeCaptchaInput(BaseModel):
|
||||
"""识别图形验证码工具的输入参数模型。"""
|
||||
|
||||
explanation: Optional[str] = Field(
|
||||
None,
|
||||
description="Clear explanation of why this captcha image needs to be recognized",
|
||||
)
|
||||
image_url: str = Field(
|
||||
...,
|
||||
description=(
|
||||
|
||||
@@ -16,7 +16,6 @@ from app.schemas.types import media_type_to_agent
|
||||
|
||||
class RecognizeMediaInput(BaseModel):
|
||||
"""识别媒体信息工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
title: Optional[str] = Field(None, description="The title of the torrent/media to recognize (required for torrent recognition)")
|
||||
subtitle: Optional[str] = Field(None, description="The subtitle or description of the torrent (optional, helps improve recognition accuracy)")
|
||||
path: Optional[str] = Field(None, description="The file path to recognize (required for file recognition, mutually exclusive with title)")
|
||||
@@ -143,6 +142,9 @@ class RecognizeMediaTool(MoviePilotTool):
|
||||
"imdb_id": media_info.get("imdb_id"),
|
||||
"douban_id": media_info.get("douban_id"),
|
||||
"bangumi_id": media_info.get("bangumi_id"),
|
||||
"anilist_id": media_info.get("anilist_id"),
|
||||
"media_source": media_info.get("source"),
|
||||
"media_id": media_info.get("media_id"),
|
||||
"overview": media_info.get("overview"),
|
||||
"vote_average": media_info.get("vote_average"),
|
||||
"poster_path": media_info.get("poster_path"),
|
||||
@@ -168,7 +170,11 @@ class RecognizeMediaTool(MoviePilotTool):
|
||||
"season_episode": meta_info.get("season_episode"),
|
||||
"episode_list": meta_info.get("episode_list"),
|
||||
"tmdbid": meta_info.get("tmdbid"),
|
||||
"doubanid": meta_info.get("doubanid")
|
||||
"doubanid": meta_info.get("doubanid"),
|
||||
"bangumiid": meta_info.get("bangumiid"),
|
||||
"anilistid": meta_info.get("anilistid"),
|
||||
"media_source": meta_info.get("media_source"),
|
||||
"media_id": meta_info.get("media_id"),
|
||||
}
|
||||
|
||||
return json.dumps(result, ensure_ascii=False, indent=2)
|
||||
|
||||
@@ -17,8 +17,6 @@ from app.log import logger
|
||||
class ReloadPluginInput(BaseModel):
|
||||
"""重载插件工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
plugin_id: str = Field(
|
||||
...,
|
||||
description="The plugin ID to reload so the latest saved config takes effect.",
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
"""立即执行 Agent 自主定时任务工具。"""
|
||||
|
||||
from typing import Optional, Type
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.db.agenttask_oper import AgentTaskOper
|
||||
|
||||
|
||||
class RunAgentTaskInput(BaseModel):
|
||||
"""立即执行 Agent 自主定时任务的输入参数。"""
|
||||
|
||||
task_id: int = Field(
|
||||
...,
|
||||
ge=1,
|
||||
description=(
|
||||
"Integer autonomous task ID returned by query_agent_tasks. Do not pass a "
|
||||
"runtime scheduler job_id such as agent-task-12."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class RunAgentTaskTool(MoviePilotTool):
|
||||
"""将当前用户的 Agent 自主定时任务提交为立即执行。"""
|
||||
|
||||
name: str = "run_agent_task"
|
||||
tags: list[str] = [ToolTag.Write, ToolTag.AgentTask, ToolTag.Admin]
|
||||
description: str = (
|
||||
"Queue an enabled autonomous agent task owned by the current user for immediate "
|
||||
"execution. Use the integer task_id returned by query_agent_tasks. The task runs "
|
||||
"after the current agent turn can finish and broadcasts its result through the "
|
||||
"configured notification channels."
|
||||
)
|
||||
args_schema: Type[BaseModel] = RunAgentTaskInput
|
||||
require_admin: bool = True
|
||||
|
||||
def get_tool_message(self, **kwargs: object) -> Optional[str]:
|
||||
"""生成立即执行 Agent 任务的提示消息。"""
|
||||
return f"立即执行自主定时任务:{kwargs.get('task_id', '')}"
|
||||
|
||||
def _get_task_state(self, task_id: int) -> tuple[str, Optional[str]]:
|
||||
"""校验任务归属和状态,返回可执行性及任务名称。"""
|
||||
task = AgentTaskOper().get(
|
||||
task_id=task_id,
|
||||
user_id=str(self._user_id),
|
||||
)
|
||||
if not task:
|
||||
return "not_found", None
|
||||
if not task.enabled:
|
||||
return "disabled", task.name
|
||||
if task.last_status == "running":
|
||||
return "running", task.name
|
||||
return "ready", task.name
|
||||
|
||||
async def run(self, task_id: int, **kwargs: object) -> str:
|
||||
"""立即执行当前用户拥有且已启用的 Agent 自主定时任务。"""
|
||||
from app.scheduler import Scheduler
|
||||
|
||||
payload = RunAgentTaskInput(task_id=task_id)
|
||||
status, task_name = await self.run_blocking(
|
||||
"db",
|
||||
self._get_task_state,
|
||||
payload.task_id,
|
||||
)
|
||||
if status == "not_found":
|
||||
return f"Agent 定时任务 {task_id} 不存在或不属于当前用户"
|
||||
if status == "disabled":
|
||||
return f"Agent 定时任务 {task_id} 已暂停,请先恢复后再执行"
|
||||
if status == "running":
|
||||
return f"Agent 定时任务 {task_id} 正在执行,请勿重复触发"
|
||||
if not Scheduler().start_agent_task(payload.task_id):
|
||||
return f"Agent 定时任务 {task_id} 尚未注册到运行时调度器,无法立即执行"
|
||||
return (
|
||||
f"Agent 定时任务 {task_id} 已提交立即执行:{task_name}。"
|
||||
"执行完成后将通过已配置的通知渠道广播结果"
|
||||
)
|
||||
@@ -12,27 +12,34 @@ from app.log import logger
|
||||
class RunSchedulerInput(BaseModel):
|
||||
"""运行定时服务工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
job_id: str = Field(
|
||||
...,
|
||||
description="The ID of the scheduled job to run (can be obtained from query_schedulers tool)",
|
||||
description=(
|
||||
"Runtime scheduler job ID returned by query_schedulers. Do not pass an "
|
||||
"autonomous agent task ID or an agent-task-* runtime ID."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class RunSchedulerTool(MoviePilotTool):
|
||||
"""立即运行系统、插件或工作流注册的定时服务。"""
|
||||
|
||||
name: str = "run_scheduler"
|
||||
tags: list[str] = [
|
||||
ToolTag.Write,
|
||||
ToolTag.Scheduler,
|
||||
ToolTag.Admin,
|
||||
]
|
||||
description: str = "Manually trigger a scheduled task to run immediately. This will execute the specified scheduler job by its ID."
|
||||
description: str = (
|
||||
"Manually trigger a MoviePilot system, plugin, or workflow scheduler service by "
|
||||
"the runtime job_id returned from query_schedulers. This tool does not run "
|
||||
"user-created autonomous agent tasks; use run_agent_task with an integer task_id."
|
||||
)
|
||||
args_schema: Type[BaseModel] = RunSchedulerInput
|
||||
require_admin: bool = True
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
"""根据运行参数生成友好的提示消息"""
|
||||
def get_tool_message(self, **kwargs: object) -> Optional[str]:
|
||||
"""根据运行参数生成友好的提示消息。"""
|
||||
job_id = kwargs.get("job_id", "")
|
||||
return f"运行定时服务 (ID: {job_id})"
|
||||
|
||||
@@ -48,10 +55,18 @@ class RunSchedulerTool(MoviePilotTool):
|
||||
return True, scheduler_item.name
|
||||
return False, ""
|
||||
|
||||
async def run(self, job_id: str, **kwargs) -> str:
|
||||
async def run(self, job_id: str, **kwargs: object) -> str:
|
||||
"""立即运行非 Agent 自主任务的运行时定时服务。"""
|
||||
logger.info(f"执行工具: {self.name}, 参数: job_id={job_id}")
|
||||
|
||||
try:
|
||||
from app.scheduler import AGENT_TASK_JOB_PREFIX
|
||||
|
||||
if job_id.startswith(f"{AGENT_TASK_JOB_PREFIX}-"):
|
||||
return (
|
||||
"Agent 自主定时任务不能通过 run_scheduler 运行,"
|
||||
"请使用 query_agent_tasks 查询整数 task_id 后调用 run_agent_task"
|
||||
)
|
||||
job_exists, job_name = await self.run_blocking(
|
||||
"workflow", self._run_scheduler_sync, job_id
|
||||
)
|
||||
|
||||
@@ -15,8 +15,6 @@ from app.schemas.types import EventType, MessageChannel
|
||||
class RunSlashCommandInput(BaseModel):
|
||||
"""运行斜杠命令工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
command: str = Field(
|
||||
...,
|
||||
description="The slash command to execute, e.g. '/cookiecloud'. "
|
||||
|
||||
@@ -7,7 +7,6 @@ from pydantic import BaseModel, Field
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.chain.workflow import WorkflowChain
|
||||
from app.db import AsyncSessionFactory
|
||||
from app.db.workflow_oper import WorkflowOper
|
||||
from app.log import logger
|
||||
|
||||
@@ -15,8 +14,6 @@ from app.log import logger
|
||||
class RunWorkflowInput(BaseModel):
|
||||
"""执行工作流工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
workflow_id: int = Field(
|
||||
..., description="Workflow ID (can be obtained from query_workflows tool)"
|
||||
)
|
||||
@@ -65,26 +62,23 @@ class RunWorkflowTool(MoviePilotTool):
|
||||
)
|
||||
|
||||
try:
|
||||
# 获取数据库会话
|
||||
async with AsyncSessionFactory() as db:
|
||||
workflow_oper = WorkflowOper(db)
|
||||
workflow = await workflow_oper.async_get(workflow_id)
|
||||
workflow = await WorkflowOper().async_get(workflow_id)
|
||||
|
||||
if not workflow:
|
||||
return f"未找到工作流:{workflow_id},请使用 query_workflows 工具查询可用的工作流"
|
||||
if not workflow:
|
||||
return f"未找到工作流:{workflow_id},请使用 query_workflows 工具查询可用的工作流"
|
||||
|
||||
# 工作流执行链路包含大量同步步骤,统一放到 workflow 线程池。
|
||||
state, errmsg = await self.run_blocking(
|
||||
"workflow",
|
||||
self._run_workflow_sync,
|
||||
workflow.id,
|
||||
from_begin,
|
||||
)
|
||||
# 工作流执行链路包含大量同步步骤,统一放到 workflow 线程池。
|
||||
state, errmsg = await self.run_blocking(
|
||||
"workflow",
|
||||
self._run_workflow_sync,
|
||||
workflow.id,
|
||||
from_begin,
|
||||
)
|
||||
|
||||
if not state:
|
||||
return f"执行工作流失败:{workflow.name} (ID: {workflow.id})\n错误原因:{errmsg}"
|
||||
else:
|
||||
return f"工作流执行成功:{workflow.name} (ID: {workflow.id})"
|
||||
if not state:
|
||||
return f"执行工作流失败:{workflow.name} (ID: {workflow.id})\n错误原因:{errmsg}"
|
||||
else:
|
||||
return f"工作流执行成功:{workflow.name} (ID: {workflow.id})"
|
||||
except Exception as e:
|
||||
logger.error(f"执行工作流失败: {e}", exc_info=True)
|
||||
return f"执行工作流时发生错误: {str(e)}"
|
||||
|
||||
@@ -16,8 +16,6 @@ from app.schemas import FileItem
|
||||
class ScrapeMetadataInput(BaseModel):
|
||||
"""刮削媒体元数据工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(None,
|
||||
description="Clear explanation of why this tool is being used in the current context",)
|
||||
path: str = Field(
|
||||
...,
|
||||
description="Path to the file or directory to scrape metadata for (e.g., '/path/to/file.mkv' or '/path/to/directory')",
|
||||
|
||||
@@ -10,11 +10,11 @@ from app.agent.tools.tags import ToolTag
|
||||
from app.chain.media import MediaChain
|
||||
from app.log import logger
|
||||
from app.schemas.types import MediaType, media_type_to_agent
|
||||
from app.utils.media import resolve_media_identity
|
||||
|
||||
|
||||
class SearchMediaInput(BaseModel):
|
||||
"""搜索媒体工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
title: str = Field(..., description="The title of the media to search for (e.g., 'The Matrix', 'Breaking Bad')")
|
||||
year: Optional[str] = Field(None, description="Release year of the media (optional, helps narrow down results)")
|
||||
media_type: Optional[str] = Field(None,
|
||||
@@ -44,7 +44,7 @@ class SearchMediaTool(MoviePilotTool):
|
||||
message += f" ({year})"
|
||||
if media_type:
|
||||
message += f" [{media_type}]"
|
||||
if season:
|
||||
if season is not None:
|
||||
message += f" 第{season}季"
|
||||
|
||||
return message
|
||||
@@ -84,6 +84,7 @@ class SearchMediaTool(MoviePilotTool):
|
||||
# 精简字段,只保留关键信息
|
||||
simplified_results = []
|
||||
for r in limited_results:
|
||||
media_source, media_id = resolve_media_identity(media=r)
|
||||
simplified = {
|
||||
"title": r.title,
|
||||
"en_title": r.en_title,
|
||||
@@ -93,6 +94,10 @@ class SearchMediaTool(MoviePilotTool):
|
||||
"tmdb_id": r.tmdb_id,
|
||||
"imdb_id": r.imdb_id,
|
||||
"douban_id": r.douban_id,
|
||||
"bangumi_id": r.bangumi_id,
|
||||
"anilist_id": r.anilist_id,
|
||||
"media_source": media_source,
|
||||
"media_id": media_id,
|
||||
"overview": r.overview[:200] + "..." if r.overview and len(r.overview) > 200 else r.overview,
|
||||
"vote_average": r.vote_average,
|
||||
"poster_path": r.poster_path,
|
||||
|
||||
@@ -13,7 +13,6 @@ from app.log import logger
|
||||
|
||||
class SearchPersonInput(BaseModel):
|
||||
"""搜索人物工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
name: str = Field(..., description="The name of the person to search for (e.g., 'Tom Hanks', '周杰伦')")
|
||||
|
||||
|
||||
|
||||
@@ -11,11 +11,11 @@ from app.chain.douban import DoubanChain
|
||||
from app.chain.tmdb import TmdbChain
|
||||
from app.chain.bangumi import BangumiChain
|
||||
from app.log import logger
|
||||
from app.utils.media import resolve_media_identity
|
||||
|
||||
|
||||
class SearchPersonCreditsInput(BaseModel):
|
||||
"""搜索演员参演作品工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
person_id: int = Field(..., description="The ID of the person/actor to search for credits (e.g., 31 for Tom Hanks in TMDB)")
|
||||
source: str = Field(..., description="The data source: 'tmdb' for TheMovieDB, 'douban' for Douban, 'bangumi' for Bangumi")
|
||||
page: Optional[int] = Field(1, description="Page number for pagination (default: 1)")
|
||||
@@ -60,6 +60,7 @@ class SearchPersonCreditsTool(MoviePilotTool):
|
||||
# 精简字段,只保留关键信息
|
||||
simplified_results = []
|
||||
for media in limited_medias:
|
||||
media_source, media_id = resolve_media_identity(media=media)
|
||||
simplified = {
|
||||
"title": media.title,
|
||||
"en_title": media.en_title,
|
||||
@@ -69,6 +70,10 @@ class SearchPersonCreditsTool(MoviePilotTool):
|
||||
"tmdb_id": media.tmdb_id,
|
||||
"imdb_id": media.imdb_id,
|
||||
"douban_id": media.douban_id,
|
||||
"bangumi_id": media.bangumi_id,
|
||||
"anilist_id": media.anilist_id,
|
||||
"media_source": media_source,
|
||||
"media_id": media_id,
|
||||
"overview": media.overview[:200] + "..." if media.overview and len(media.overview) > 200 else media.overview,
|
||||
"vote_average": media.vote_average,
|
||||
"poster_path": media.poster_path,
|
||||
|
||||
@@ -15,7 +15,6 @@ from app.schemas.types import media_type_to_agent
|
||||
|
||||
class SearchSubscribeInput(BaseModel):
|
||||
"""搜索订阅缺失剧集工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
subscribe_id: int = Field(..., description="The ID of the subscription to search for missing episodes (can be obtained from query_subscribes tool)")
|
||||
manual: Optional[bool] = Field(False, description="Whether this is a manual search (default: False)")
|
||||
filter_groups: Optional[List[str]] = Field(None,
|
||||
@@ -71,7 +70,11 @@ class SearchSubscribeTool(MoviePilotTool):
|
||||
"total_episode": subscribe.total_episode,
|
||||
"lack_episode": subscribe.lack_episode,
|
||||
"tmdbid": subscribe.tmdbid,
|
||||
"doubanid": subscribe.doubanid
|
||||
"doubanid": subscribe.doubanid,
|
||||
"bangumiid": subscribe.bangumiid,
|
||||
"anilistid": subscribe.anilistid,
|
||||
"media_source": subscribe.media_source,
|
||||
"media_id": subscribe.media_id,
|
||||
}
|
||||
|
||||
# 检查订阅状态
|
||||
|
||||
@@ -9,7 +9,7 @@ from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.chain.search import SearchChain
|
||||
from app.db.systemconfig_oper import SystemConfigOper
|
||||
from app.helper.sites import SitesHelper
|
||||
from app.helper.sites import SitesHelper # noqa
|
||||
from app.log import logger
|
||||
from app.schemas.types import MediaType, SystemConfigKey
|
||||
from ._torrent_search_utils import (
|
||||
@@ -20,14 +20,18 @@ from ._torrent_search_utils import (
|
||||
|
||||
class SearchTorrentsInput(BaseModel):
|
||||
"""搜索种子工具的输入参数模型"""
|
||||
explanation: Optional[str] = Field(None, description="Clear explanation of why this tool is being used in the current context")
|
||||
tmdb_id: Optional[int] = Field(None, description="TMDB ID (can be obtained from search_media tool). Either tmdb_id or douban_id must be provided.")
|
||||
douban_id: Optional[str] = Field(None, description="Douban ID (can be obtained from search_media tool). Either tmdb_id or douban_id must be provided.")
|
||||
tmdb_id: Optional[int] = Field(None, description="TMDB media ID")
|
||||
douban_id: Optional[str] = Field(None, description="Douban media ID")
|
||||
bangumi_id: Optional[int] = Field(None, description="Bangumi media ID")
|
||||
anilist_id: Optional[int] = Field(None, description="AniList media ID")
|
||||
media_source: Optional[str] = Field(None, description="Media metadata source")
|
||||
media_id: Optional[str] = Field(None, description="Native ID for media_source")
|
||||
media_type: Optional[str] = Field(None, description="Allowed values: movie, tv")
|
||||
area: Optional[str] = Field(None, description="Search scope: 'title' (default) or 'imdbid'")
|
||||
sites: Optional[List[int]] = Field(None,
|
||||
description="Array of specific site IDs to search on (optional, if not provided searches all configured sites)")
|
||||
|
||||
|
||||
class SearchTorrentsTool(MoviePilotTool):
|
||||
name: str = "search_torrents"
|
||||
tags: list[str] = [
|
||||
@@ -36,23 +40,27 @@ class SearchTorrentsTool(MoviePilotTool):
|
||||
ToolTag.Site,
|
||||
ToolTag.Media,
|
||||
]
|
||||
description: str = ("Search for torrent files by media ID across configured indexer sites, cache the matched results, "
|
||||
"and return available filter options for follow-up selection. "
|
||||
"Requires tmdb_id or douban_id (can be obtained from search_media tool) for accurate matching.")
|
||||
description: str = (
|
||||
"Search for torrent files by media ID across configured indexer sites, cache the matched results, "
|
||||
"and return available filter options for follow-up selection. "
|
||||
"Accepts a TMDB, Douban, Bangumi, AniList, or source-native media ID for accurate matching.")
|
||||
args_schema: Type[BaseModel] = SearchTorrentsInput
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
"""根据搜索参数生成友好的提示消息"""
|
||||
tmdb_id = kwargs.get("tmdb_id")
|
||||
douban_id = kwargs.get("douban_id")
|
||||
media_type = kwargs.get("media_type")
|
||||
|
||||
if tmdb_id:
|
||||
message = f"搜索种子: TMDB={tmdb_id}"
|
||||
elif douban_id:
|
||||
message = f"搜索种子: 豆瓣={douban_id}"
|
||||
else:
|
||||
message = "搜索种子"
|
||||
identities = (
|
||||
("TMDB", kwargs.get("tmdb_id")),
|
||||
("豆瓣", kwargs.get("douban_id")),
|
||||
("Bangumi", kwargs.get("bangumi_id")),
|
||||
("AniList", kwargs.get("anilist_id")),
|
||||
(kwargs.get("media_source") or "媒体源", kwargs.get("media_id")),
|
||||
)
|
||||
label, identity = next(
|
||||
((label, identity) for label, identity in identities if identity is not None),
|
||||
(None, None),
|
||||
)
|
||||
message = f"搜索种子: {label}={identity}" if label else "搜索种子"
|
||||
if media_type:
|
||||
message += f" [{media_type}]"
|
||||
return message
|
||||
@@ -63,13 +71,15 @@ class SearchTorrentsTool(MoviePilotTool):
|
||||
return SystemConfigOper().get(SystemConfigKey.IndexerSites) or []
|
||||
|
||||
async def run(self, tmdb_id: Optional[int] = None, douban_id: Optional[str] = None,
|
||||
bangumi_id: Optional[int] = None, anilist_id: Optional[int] = None,
|
||||
media_source: Optional[str] = None, media_id: Optional[str] = None,
|
||||
media_type: Optional[str] = None, area: Optional[str] = None,
|
||||
sites: Optional[List[int]] = None, **kwargs) -> str:
|
||||
logger.info(
|
||||
f"执行工具: {self.name}, 参数: tmdb_id={tmdb_id}, douban_id={douban_id}, media_type={media_type}, area={area}, sites={sites}")
|
||||
|
||||
if not tmdb_id and not douban_id:
|
||||
return "参数错误:tmdb_id 和 douban_id 至少需要提供一个,请先使用 search_media 工具获取媒体 ID。"
|
||||
if not any((tmdb_id, douban_id, bangumi_id, anilist_id, media_id)):
|
||||
return "参数错误:至少需要提供一个媒体 ID,请先使用 search_media 工具获取媒体信息。"
|
||||
|
||||
try:
|
||||
search_chain = SearchChain()
|
||||
@@ -82,6 +92,10 @@ class SearchTorrentsTool(MoviePilotTool):
|
||||
filtered_torrents = await search_chain.async_search_by_id(
|
||||
tmdbid=tmdb_id,
|
||||
doubanid=douban_id,
|
||||
bangumiid=bangumi_id,
|
||||
anilistid=anilist_id,
|
||||
source=media_source,
|
||||
mediaid=media_id,
|
||||
mtype=media_type_enum,
|
||||
area=area or "title",
|
||||
sites=sites,
|
||||
@@ -108,9 +122,9 @@ class SearchTorrentsTool(MoviePilotTool):
|
||||
}, ensure_ascii=False, indent=2)
|
||||
return result_json
|
||||
else:
|
||||
media_id = f"TMDB={tmdb_id}" if tmdb_id else f"豆瓣={douban_id}"
|
||||
identity = media_id or tmdb_id or douban_id or bangumi_id or anilist_id
|
||||
result_json = json.dumps({
|
||||
"message": f"未找到相关种子资源: {media_id}",
|
||||
"message": f"未找到相关种子资源: {identity}",
|
||||
"all_sites": all_sites,
|
||||
"search_site_ids": search_site_ids,
|
||||
}, ensure_ascii=False, indent=2)
|
||||
|
||||
@@ -48,10 +48,6 @@ class _SearchSiteFilter:
|
||||
class SearchWebInput(BaseModel):
|
||||
"""搜索网络内容工具的输入参数模型"""
|
||||
|
||||
explanation: Optional[str] = Field(
|
||||
None,
|
||||
description="Clear explanation of why this tool is being used in the current context",
|
||||
)
|
||||
query: str = Field(
|
||||
..., description="The search query string to search for on the web"
|
||||
)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user