mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-08-26 10:40:32 +08:00
Compare commits
218 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
684e76e518 | ||
|
|
458c08a137 | ||
|
|
7012e0e305 | ||
|
|
91ce365f78 | ||
|
|
17be4304c1 | ||
|
|
c6bd396794 | ||
|
|
7985268f10 | ||
|
|
865635c59d | ||
|
|
759b9e47eb | ||
|
|
63e492be7c | ||
|
|
b5eca00ba3 | ||
|
|
b80642f56f | ||
|
|
6d3161f3cb | ||
|
|
ea8d1f8d26 | ||
|
|
5654512d41 | ||
|
|
a52e1fdc1c | ||
|
|
44db45ea28 | ||
|
|
83409c1439 | ||
|
|
cf6c73d85b | ||
|
|
987c1722d7 | ||
|
|
57220c93db | ||
|
|
4abae809c2 | ||
|
|
7fe7be6d71 | ||
|
|
4b1df72a4a | ||
|
|
a23ac6c56d | ||
|
|
48f4bd5f18 | ||
|
|
7a2b003d02 | ||
|
|
a700ba9379 | ||
|
|
36dd381260 | ||
|
|
f25ee25bb0 | ||
|
|
d625196b08 | ||
|
|
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 |
@@ -71,6 +71,7 @@ test_*
|
||||
|
||||
# Build artifacts
|
||||
build/
|
||||
.build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
rust/**/target/
|
||||
|
||||
16
.github/ISSUE_TEMPLATE/feature_request.yml
vendored
16
.github/ISSUE_TEMPLATE/feature_request.yml
vendored
@@ -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:
|
||||
|
||||
27
.github/workflows/beta.yml
vendored
27
.github/workflows/beta.yml
vendored
@@ -2,6 +2,10 @@ name: MoviePilot Builder Beta
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
jobs:
|
||||
Docker-build:
|
||||
runs-on: ubuntu-latest
|
||||
@@ -16,6 +20,25 @@ jobs:
|
||||
app_version=$(cat version.py |sed -ne "s/APP_VERSION\s=\s'v\(.*\)'/\1/gp")
|
||||
echo "app_version=$app_version" >> $GITHUB_ENV
|
||||
|
||||
- name: Checkout Wiki Plugin Market
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
repository: jxxghp/MoviePilot-Wiki
|
||||
ref: main
|
||||
path: .build/moviepilot-wiki
|
||||
sparse-checkout: plugin.md
|
||||
sparse-checkout-cone-mode: false
|
||||
persist-credentials: false
|
||||
|
||||
- name: Generate Plugin Market Default
|
||||
id: plugin_market
|
||||
run: |
|
||||
python3 -m scripts.generate_plugin_market_default \
|
||||
--wiki-file .build/moviepilot-wiki/plugin.md \
|
||||
--config-file app/core/config.py
|
||||
wiki_commit=$(git -C .build/moviepilot-wiki rev-parse HEAD)
|
||||
echo "wiki_commit=$wiki_commit" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Docker Meta
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
@@ -55,6 +78,8 @@ jobs:
|
||||
linux/arm64/v8
|
||||
push: true
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
labels: |
|
||||
${{ steps.meta.outputs.labels }}
|
||||
org.moviepilot.plugin-market-wiki-revision=${{ steps.plugin_market.outputs.wiki_commit }}
|
||||
cache-from: type=gha,scope=moviepilot-docker,version=2
|
||||
cache-to: type=gha,scope=moviepilot-docker,mode=max,version=2
|
||||
|
||||
14
.github/workflows/build-v3.yml
vendored
Normal file
14
.github/workflows/build-v3.yml
vendored
Normal file
@@ -0,0 +1,14 @@
|
||||
name: MoviePilot Builder v3
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
select-v3:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
# GitHub 仅从默认分支登记手动工作流;选择 v3 后会加载 v3 分支的完整构建配置。
|
||||
- name: Require v3 branch
|
||||
run: |
|
||||
echo "::error::请在 Run workflow 中选择 v3 分支"
|
||||
exit 1
|
||||
57
.github/workflows/build.yml
vendored
57
.github/workflows/build.yml
vendored
@@ -7,6 +7,10 @@ on:
|
||||
paths:
|
||||
- 'version.py'
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
packages: write
|
||||
|
||||
jobs:
|
||||
Docker-build:
|
||||
runs-on: ubuntu-latest
|
||||
@@ -23,6 +27,39 @@ jobs:
|
||||
run: |
|
||||
app_version=$(cat version.py |sed -ne "s/APP_VERSION\s=\s'v\(.*\)'/\1/gp")
|
||||
echo "app_version=$app_version" >> $GITHUB_ENV
|
||||
echo "SOURCE_COMMIT=$(git rev-parse HEAD)" >> $GITHUB_ENV
|
||||
|
||||
- name: Checkout Wiki Plugin Market
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
repository: jxxghp/MoviePilot-Wiki
|
||||
ref: main
|
||||
path: .build/moviepilot-wiki
|
||||
sparse-checkout: plugin.md
|
||||
sparse-checkout-cone-mode: false
|
||||
persist-credentials: false
|
||||
|
||||
- name: Generate Plugin Market Default
|
||||
id: plugin_market
|
||||
run: |
|
||||
python3 -m scripts.generate_plugin_market_default \
|
||||
--wiki-file .build/moviepilot-wiki/plugin.md \
|
||||
--config-file app/core/config.py
|
||||
wiki_commit=$(git -C .build/moviepilot-wiki rev-parse HEAD)
|
||||
echo "wiki_commit=$wiki_commit" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Create Release Snapshot
|
||||
id: release_snapshot
|
||||
env:
|
||||
WIKI_COMMIT: ${{ steps.plugin_market.outputs.wiki_commit }}
|
||||
run: |
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "41898282+github-actions[bot]@users.noreply.github.com"
|
||||
git add app/core/config.py
|
||||
if ! git diff --cached --quiet; then
|
||||
git commit -m "build(plugin-market): sync default from MoviePilot-Wiki@${WIKI_COMMIT:0:12}"
|
||||
fi
|
||||
echo "release_commit=$(git rev-parse HEAD)" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Docker Meta
|
||||
id: meta
|
||||
@@ -65,7 +102,10 @@ jobs:
|
||||
linux/arm64/v8
|
||||
push: true
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
labels: |
|
||||
${{ steps.meta.outputs.labels }}
|
||||
org.opencontainers.image.revision=${{ steps.release_snapshot.outputs.release_commit }}
|
||||
org.moviepilot.plugin-market-wiki-revision=${{ steps.plugin_market.outputs.wiki_commit }}
|
||||
cache-from: type=gha,scope=moviepilot-docker,version=2
|
||||
cache-to: type=gha,scope=moviepilot-docker,mode=max,version=2
|
||||
|
||||
@@ -78,9 +118,9 @@ jobs:
|
||||
|
||||
# 使用 || 作为分隔符,同时获取 commit 消息和作者 GitHub 用户名
|
||||
if [ -z "$PREVIOUS_TAG" ]; then
|
||||
COMMITS=$(git log --pretty=format:"%s||%an" HEAD)
|
||||
COMMITS=$(git log --pretty=format:"%s||%an" "${SOURCE_COMMIT}")
|
||||
else
|
||||
COMMITS=$(git log --pretty=format:"%s||%an" ${PREVIOUS_TAG}..HEAD)
|
||||
COMMITS=$(git log --pretty=format:"%s||%an" "${PREVIOUS_TAG}..${SOURCE_COMMIT}")
|
||||
fi
|
||||
|
||||
# 分类收集 commit 消息(使用关联数组去重)
|
||||
@@ -188,6 +228,17 @@ jobs:
|
||||
delete_release: true
|
||||
github_token: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Publish Release Tag
|
||||
env:
|
||||
RELEASE_COMMIT: ${{ steps.release_snapshot.outputs.release_commit }}
|
||||
run: |
|
||||
tag_name="v${{ env.app_version }}"
|
||||
if git show-ref --verify --quiet "refs/tags/${tag_name}"; then
|
||||
git tag -d "$tag_name"
|
||||
fi
|
||||
git tag "$tag_name" "$RELEASE_COMMIT"
|
||||
git push origin "refs/tags/${tag_name}"
|
||||
|
||||
- name: Generate Release
|
||||
uses: softprops/action-gh-release@v2
|
||||
with:
|
||||
|
||||
101
.github/workflows/pr-agent.yml
vendored
101
.github/workflows/pr-agent.yml
vendored
@@ -1,9 +1,8 @@
|
||||
name: PR Agent
|
||||
name: PR-Agent
|
||||
|
||||
on:
|
||||
pull_request_target:
|
||||
# PR-Agent 通过 base repo 上下文读取 PR diff 并发布 Review,不 checkout 或执行 PR 分支代码。
|
||||
# pull_request_target 允许 fork PR 使用仓库 secrets,因此 workflow 只运行固定 digest 的 PR-Agent 容器。
|
||||
# Fork 审查需要目标仓库凭据;该 job 仅通过 GitHub API 读取 PR 内容,不 checkout 或执行 PR 分支代码。
|
||||
types:
|
||||
- opened
|
||||
- reopened
|
||||
@@ -11,24 +10,17 @@ on:
|
||||
- review_requested
|
||||
- synchronize
|
||||
issue_comment:
|
||||
# 手动命令如 "/review"、"/describe"、"/improve" 和 "/ask ..." 只在 PR 评论中有意义。
|
||||
# issue_comment 同时覆盖普通 issue,因此 job 里还会再判断是否属于 PR。
|
||||
types:
|
||||
- created
|
||||
- edited
|
||||
|
||||
permissions:
|
||||
# 读取仓库内容和 PR diff。
|
||||
contents: read
|
||||
# 更新 PR 描述、发布 PR Review 或修改 PR 相关元数据。
|
||||
pull-requests: write
|
||||
# PR 评论在 GitHub API 中属于 issue comments,手动命令和总结评论需要该权限。
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
pr-agent:
|
||||
name: PR-Agent review and describe
|
||||
# PR 事件自动处理;评论命令仅允许指定身份在 PR 下触发,避免任意评论消耗模型配额。
|
||||
if: >-
|
||||
github.event.sender.type != 'Bot' &&
|
||||
(
|
||||
@@ -36,88 +28,25 @@ jobs:
|
||||
(
|
||||
github.event_name == 'issue_comment' &&
|
||||
github.event.issue.pull_request != null &&
|
||||
contains(fromJSON('["OWNER", "MEMBER", "COLLABORATOR", "CONTRIBUTOR", "FIRST_TIME_CONTRIBUTOR"]'), github.event.comment.author_association) &&
|
||||
(
|
||||
github.event.comment.body == '/review' ||
|
||||
startsWith(github.event.comment.body, '/review ') ||
|
||||
github.event.comment.body == '/describe' ||
|
||||
startsWith(github.event.comment.body, '/describe ') ||
|
||||
github.event.comment.body == '/improve' ||
|
||||
startsWith(github.event.comment.body, '/improve ') ||
|
||||
github.event.comment.body == '/ask' ||
|
||||
startsWith(github.event.comment.body, '/ask ')
|
||||
)
|
||||
contains(fromJSON('["OWNER", "MEMBER", "COLLABORATOR", "CONTRIBUTOR", "FIRST_TIME_CONTRIBUTOR"]'), github.event.comment.author_association)
|
||||
)
|
||||
)
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.event.issue.number }}
|
||||
cancel-in-progress: ${{ github.event_name == 'pull_request_target' }}
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20
|
||||
|
||||
steps:
|
||||
- name: Run PR-Agent
|
||||
id: pragent
|
||||
# 使用版本号加 digest 固定容器构建,避免 tag 被重推后改变运行内容。
|
||||
uses: docker://pragent/pr-agent:0.37.0-github_action@sha256:4ec7bac814050a1bc8c96ab2fab6b7b0f65df0049a5ec43f3fee1a0b551c28ca
|
||||
- name: Run PR Review
|
||||
uses: docker://ghcr.io/infinitypacer/pr-review-runner:latest
|
||||
env:
|
||||
# PR-Agent 使用该 token 读取 PR 元数据并发布评论。
|
||||
GITHUB_TOKEN: ${{ github.token }}
|
||||
|
||||
# 仓库设置中添加的 Secret:Settings -> Secrets and variables -> Actions。
|
||||
# 该 key 只传给 PR-Agent 运行时,不写入仓库。
|
||||
OPENAI_KEY: ${{ secrets.OPENAI_KEY }}
|
||||
|
||||
# 仓库设置中添加的 Secret。OpenAI 兼容服务通常需要填写以 "/v1" 结尾的 API 根地址。
|
||||
OPENAI.API_BASE: ${{ secrets.OPENAI_API_BASE }}
|
||||
|
||||
# 模型、输出语言和大 diff 处理策略。
|
||||
config.model: "gpt-5.5"
|
||||
config.fallback_models: '["gpt-5.4"]'
|
||||
config.reasoning_effort: "xhigh"
|
||||
config.ai_timeout: "900"
|
||||
config.response_language: "zh-CN"
|
||||
config.large_patch_policy: "clip"
|
||||
config.ignore_pr_title: '["^\\[Auto\\]", "^Auto"]'
|
||||
config.ignore_pr_labels: '["skip pr-agent"]'
|
||||
|
||||
# pull_request_target 事件默认自动执行 /review 和 /describe;/improve 保持手动触发。
|
||||
github_action_config.auto_review: "true"
|
||||
github_action_config.auto_describe: "true"
|
||||
github_action_config.auto_improve: "false"
|
||||
|
||||
# 允许触发自动工具的 PR 动作。包含 synchronize,便于新 commit 推送后刷新结果。
|
||||
github_action_config.pr_actions: '["opened", "reopened", "ready_for_review", "review_requested", "synchronize"]'
|
||||
|
||||
# 保留 action outputs,便于后续 workflow 编排或排查。
|
||||
github_action_config.enable_output: "true"
|
||||
|
||||
# /describe 行为控制;与自动触发配置放在同一层,避免使用默认图表和标签策略。
|
||||
pr_description.generate_ai_title: "false"
|
||||
pr_description.publish_labels: "false"
|
||||
pr_description.enable_pr_diagram: "false"
|
||||
pr_description.collapsible_file_list: "adaptive"
|
||||
pr_description.add_original_user_description: "true"
|
||||
|
||||
# /review 输出策略,聚焦维护者需要处理的风险和缺口。
|
||||
pr_reviewer.extra_instructions: |
|
||||
请用中文输出。
|
||||
优先指出 P0/P1 风险,避免纠结纯格式问题。
|
||||
重点检查安全、权限、状态一致性、异步/缓存、副作用和测试缺口。
|
||||
pr_reviewer.num_max_findings: "5"
|
||||
pr_reviewer.persistent_comment: "true"
|
||||
pr_reviewer.publish_output_no_suggestions: "true"
|
||||
pr_reviewer.require_tests_review: "true"
|
||||
pr_reviewer.require_security_review: "true"
|
||||
pr_reviewer.require_estimate_effort_to_review: "true"
|
||||
pr_reviewer.require_can_be_split_review: "true"
|
||||
pr_reviewer.require_todo_scan: "false"
|
||||
pr_reviewer.enable_review_labels_effort: "false"
|
||||
pr_reviewer.enable_review_labels_security: "true"
|
||||
|
||||
# /improve 和 /ask 的手动命令策略。
|
||||
pr_code_suggestions.focus_only_on_problems: "true"
|
||||
pr_code_suggestions.suggestions_score_threshold: "7"
|
||||
pr_code_suggestions.commitable_code_suggestions: "false"
|
||||
pr_questions.use_conversation_history: "true"
|
||||
|
||||
# 可选成本和噪音控制:
|
||||
# github_action_config.auto_improve: "true"
|
||||
# config.verbosity_level: "1"
|
||||
# pr_reviewer.num_max_findings: "3"
|
||||
PRR_AUTO_REVIEW_SCOPE: all
|
||||
PRR_ALLOWED_ASSOCIATIONS: '["OWNER", "MEMBER", "COLLABORATOR", "CONTRIBUTOR", "FIRST_TIME_CONTRIBUTOR"]'
|
||||
PRR_DISABLED_COMMANDS: '["/improve"]'
|
||||
PRR_SKIP_LABEL: skip pr-agent
|
||||
PRR_SKIP_TITLE_PATTERN: '^(?:\[Auto\]|Auto)'
|
||||
config.response_language: zh-CN
|
||||
|
||||
134
.github/workflows/site-adapter-collector.yml
vendored
Normal file
134
.github/workflows/site-adapter-collector.yml
vendored
Normal file
@@ -0,0 +1,134 @@
|
||||
name: Site Adapter Collector
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
release_tag:
|
||||
description: Existing release tag to receive collector assets; leave empty for artifacts only
|
||||
required: false
|
||||
type: string
|
||||
release:
|
||||
types:
|
||||
- published
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
build:
|
||||
name: Build ${{ matrix.platform_name }} collector
|
||||
runs-on: ${{ matrix.runner }}
|
||||
timeout-minutes: 30
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
include:
|
||||
- platform_name: Windows
|
||||
platform_id: windows
|
||||
runner: windows-latest
|
||||
source_name: moviepilot-site-collector.exe
|
||||
asset_name: moviepilot-site-collector-windows.exe
|
||||
artifact_name: site-adapter-collector-windows
|
||||
- platform_name: macOS
|
||||
platform_id: macos
|
||||
runner: macos-latest
|
||||
source_name: moviepilot-site-collector
|
||||
asset_name: MoviePilot-Site-Collector-macOS.zip
|
||||
artifact_name: site-adapter-collector-macos
|
||||
- platform_name: Linux
|
||||
platform_id: linux
|
||||
runner: ubuntu-latest
|
||||
source_name: moviepilot-site-collector
|
||||
asset_name: moviepilot-site-collector-linux
|
||||
artifact_name: site-adapter-collector-linux
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v7
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: '3.12'
|
||||
cache: pip
|
||||
cache-dependency-path: scripts/site_adapter_collector_requirements.txt
|
||||
|
||||
- name: Install build dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip setuptools wheel
|
||||
pip install -r scripts/site_adapter_collector_requirements.txt
|
||||
|
||||
- name: Build single-file collector
|
||||
run: |
|
||||
pyinstaller --clean --noconfirm scripts/site_adapter_collector.spec
|
||||
|
||||
- name: Smoke-test collector
|
||||
env:
|
||||
SOURCE_NAME: ${{ matrix.source_name }}
|
||||
run: |
|
||||
python -c "import os, subprocess; from pathlib import Path; subprocess.run([str((Path('dist') / os.environ['SOURCE_NAME']).resolve()), '--help'], check=True)"
|
||||
|
||||
- name: Package macOS double-click archive
|
||||
if: matrix.platform_id == 'macos'
|
||||
shell: bash
|
||||
env:
|
||||
ASSET_NAME: ${{ matrix.asset_name }}
|
||||
SOURCE_NAME: ${{ matrix.source_name }}
|
||||
run: |
|
||||
package_dir="dist/MoviePilot-Collector"
|
||||
mkdir -p "$package_dir"
|
||||
cp "dist/$SOURCE_NAME" "$package_dir/moviepilot-site-collector-macos"
|
||||
cp scripts/start-site-adapter-collector.command "$package_dir/start-site-adapter-collector.command"
|
||||
chmod +x "$package_dir/moviepilot-site-collector-macos"
|
||||
chmod +x "$package_dir/start-site-adapter-collector.command"
|
||||
cd dist
|
||||
COPYFILE_DISABLE=1 zip -q -r -X "$ASSET_NAME" MoviePilot-Collector
|
||||
|
||||
- name: Rename Windows and Linux collector
|
||||
if: matrix.platform_id != 'macos'
|
||||
env:
|
||||
ASSET_NAME: ${{ matrix.asset_name }}
|
||||
SOURCE_NAME: ${{ matrix.source_name }}
|
||||
run: |
|
||||
python -c "import os; from pathlib import Path; (Path('dist') / os.environ['SOURCE_NAME']).replace(Path('dist') / os.environ['ASSET_NAME'])"
|
||||
|
||||
- name: Generate SHA-256 checksum
|
||||
env:
|
||||
ASSET_NAME: ${{ matrix.asset_name }}
|
||||
run: |
|
||||
python -c "import hashlib, os; from pathlib import Path; path = Path('dist') / os.environ['ASSET_NAME']; path.with_name(path.name + '.sha256').write_text(f'{hashlib.sha256(path.read_bytes()).hexdigest()} {path.name}\n', encoding='utf-8')"
|
||||
|
||||
- name: Upload collector artifact
|
||||
uses: actions/upload-artifact@v7
|
||||
with:
|
||||
name: ${{ matrix.artifact_name }}
|
||||
path: |
|
||||
dist/${{ matrix.asset_name }}
|
||||
dist/${{ matrix.asset_name }}.sha256
|
||||
if-no-files-found: error
|
||||
retention-days: 3
|
||||
|
||||
publish:
|
||||
name: Upload collectors to release
|
||||
if: github.event_name == 'release' || inputs.release_tag != ''
|
||||
needs:
|
||||
- build
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
|
||||
steps:
|
||||
- name: Download collector artifacts
|
||||
uses: actions/download-artifact@v8
|
||||
with:
|
||||
pattern: site-adapter-collector-*
|
||||
path: release-assets
|
||||
merge-multiple: true
|
||||
|
||||
- name: Upload assets to published release
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
RELEASE_TAG: ${{ github.event.release.tag_name || inputs.release_tag }}
|
||||
run: |
|
||||
gh release view "$RELEASE_TAG" --repo "$GITHUB_REPOSITORY" >/dev/null
|
||||
gh release upload "$RELEASE_TAG" release-assets/* --clobber --repo "$GITHUB_REPOSITORY"
|
||||
1
.gitignore
vendored
1
.gitignore
vendored
@@ -37,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
|
||||
|
||||
|
||||
@@ -51,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
|
||||
@@ -65,6 +68,8 @@ from app.utils.identity import SYSTEM_INTERNAL_USER_ID
|
||||
|
||||
|
||||
class AgentChain(ChainBase):
|
||||
"""Agent 业务处理链。"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
@@ -124,9 +129,18 @@ class _SessionUsageSnapshot:
|
||||
last_output_tokens: int = 0
|
||||
last_total_tokens: int = 0
|
||||
last_context_usage_ratio: Optional[float] = None
|
||||
last_cache_usage_available: bool = False
|
||||
last_cache_read_input_tokens: int = 0
|
||||
last_cache_write_input_tokens: int = 0
|
||||
last_uncached_input_tokens: int = 0
|
||||
last_cache_hit_ratio: Optional[float] = None
|
||||
total_input_tokens: int = 0
|
||||
total_output_tokens: int = 0
|
||||
total_tokens: int = 0
|
||||
total_cache_read_input_tokens: int = 0
|
||||
total_cache_write_input_tokens: int = 0
|
||||
total_uncached_input_tokens: int = 0
|
||||
cache_usage_available: bool = False
|
||||
model_call_count: int = 0
|
||||
last_updated_at: Optional[datetime] = None
|
||||
|
||||
@@ -139,9 +153,23 @@ class _SessionUsageSnapshot:
|
||||
"last_output_tokens": self.last_output_tokens,
|
||||
"last_total_tokens": self.last_total_tokens,
|
||||
"last_context_usage_ratio": self.last_context_usage_ratio,
|
||||
"last_cache_usage_available": self.last_cache_usage_available,
|
||||
"last_cache_read_input_tokens": self.last_cache_read_input_tokens,
|
||||
"last_cache_write_input_tokens": self.last_cache_write_input_tokens,
|
||||
"last_uncached_input_tokens": self.last_uncached_input_tokens,
|
||||
"last_cache_hit_ratio": self.last_cache_hit_ratio,
|
||||
"total_input_tokens": self.total_input_tokens,
|
||||
"total_output_tokens": self.total_output_tokens,
|
||||
"total_tokens": self.total_tokens,
|
||||
"total_cache_read_input_tokens": self.total_cache_read_input_tokens,
|
||||
"total_cache_write_input_tokens": self.total_cache_write_input_tokens,
|
||||
"total_uncached_input_tokens": self.total_uncached_input_tokens,
|
||||
"cache_usage_available": self.cache_usage_available,
|
||||
"total_cache_hit_ratio": (
|
||||
self.total_cache_read_input_tokens / self.total_input_tokens
|
||||
if self.cache_usage_available and self.total_input_tokens
|
||||
else None
|
||||
),
|
||||
"model_call_count": self.model_call_count,
|
||||
"last_updated_at": self.last_updated_at.strftime("%Y-%m-%d %H:%M:%S")
|
||||
if self.last_updated_at
|
||||
@@ -313,13 +341,19 @@ class MoviePilotAgent:
|
||||
"""
|
||||
构造可展示的 Agent 会话消息。
|
||||
"""
|
||||
normalized_content = content or ""
|
||||
return {
|
||||
"id": f"{role}-{uuid.uuid4().hex}",
|
||||
"role": role,
|
||||
"content": content or "",
|
||||
"content": normalized_content,
|
||||
"createdAt": cls._current_timestamp_ms(),
|
||||
"status": status,
|
||||
"tools": [],
|
||||
"segments": (
|
||||
[{"type": "text", "content": normalized_content}]
|
||||
if normalized_content
|
||||
else []
|
||||
),
|
||||
"attachments": attachments or [],
|
||||
"choices": [],
|
||||
}
|
||||
@@ -543,9 +577,33 @@ class MoviePilotAgent:
|
||||
self._session_usage.last_output_tokens = output_tokens
|
||||
self._session_usage.last_total_tokens = total_tokens
|
||||
self._session_usage.last_context_usage_ratio = usage.get("context_usage_ratio")
|
||||
cache_usage_available = bool(usage.get("cache_usage_available"))
|
||||
cache_read_input_tokens = self._coerce_int(
|
||||
usage.get("cache_read_input_tokens")
|
||||
) or 0
|
||||
cache_write_input_tokens = self._coerce_int(
|
||||
usage.get("cache_write_input_tokens")
|
||||
) or 0
|
||||
uncached_input_tokens = self._coerce_int(
|
||||
usage.get("uncached_input_tokens")
|
||||
)
|
||||
if uncached_input_tokens is None:
|
||||
uncached_input_tokens = max(
|
||||
input_tokens - cache_read_input_tokens - cache_write_input_tokens,
|
||||
0,
|
||||
)
|
||||
self._session_usage.last_cache_usage_available = cache_usage_available
|
||||
self._session_usage.last_cache_read_input_tokens = cache_read_input_tokens
|
||||
self._session_usage.last_cache_write_input_tokens = cache_write_input_tokens
|
||||
self._session_usage.last_uncached_input_tokens = uncached_input_tokens
|
||||
self._session_usage.last_cache_hit_ratio = usage.get("cache_hit_ratio")
|
||||
self._session_usage.total_input_tokens += input_tokens
|
||||
self._session_usage.total_output_tokens += output_tokens
|
||||
self._session_usage.total_tokens += total_tokens
|
||||
self._session_usage.total_cache_read_input_tokens += cache_read_input_tokens
|
||||
self._session_usage.total_cache_write_input_tokens += cache_write_input_tokens
|
||||
self._session_usage.total_uncached_input_tokens += uncached_input_tokens
|
||||
self._session_usage.cache_usage_available |= cache_usage_available
|
||||
|
||||
def get_session_status(self) -> dict[str, Any]:
|
||||
if not self._session_usage.model:
|
||||
@@ -579,6 +637,17 @@ class MoviePilotAgent:
|
||||
input_tokens=self._session_usage.total_input_tokens,
|
||||
output_tokens=self._session_usage.total_output_tokens,
|
||||
total_tokens=self._session_usage.total_tokens,
|
||||
cache_read_input_tokens=self._session_usage.total_cache_read_input_tokens,
|
||||
cache_write_input_tokens=self._session_usage.total_cache_write_input_tokens,
|
||||
uncached_input_tokens=self._session_usage.total_uncached_input_tokens,
|
||||
cache_hit_ratio=(
|
||||
self._session_usage.total_cache_read_input_tokens
|
||||
/ self._session_usage.total_input_tokens
|
||||
if self._session_usage.cache_usage_available
|
||||
and self._session_usage.total_input_tokens
|
||||
else None
|
||||
),
|
||||
cache_usage_available=self._session_usage.cache_usage_available,
|
||||
model_call_count=self._session_usage.model_call_count,
|
||||
success=success,
|
||||
error=error,
|
||||
@@ -710,7 +779,7 @@ class MoviePilotAgent:
|
||||
"""
|
||||
通过链式事件解析本次 Agent 可用的 LLM 运行时配置。
|
||||
|
||||
若没有插件返回 selected_provider_id,则沿用系统配置,保持既有行为。
|
||||
插件未返回有效配置时沿用系统配置,显式返回的配置优先。
|
||||
"""
|
||||
if self._llm_runtime_config is not None:
|
||||
return self._llm_runtime_config
|
||||
@@ -723,7 +792,9 @@ 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,
|
||||
web_search_mode=settings.LLM_WEB_SEARCH_MODE,
|
||||
)
|
||||
selected_event = await eventmanager.async_send_event(
|
||||
ChainEventType.AgentLLMProvider,
|
||||
@@ -758,9 +829,18 @@ 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
|
||||
web_search_mode = self._clean_optional_text(
|
||||
self._get_event_value(resolved_data, "web_search_mode")
|
||||
) or settings.LLM_WEB_SEARCH_MODE
|
||||
selected_provider_id = self._clean_optional_text(
|
||||
self._get_event_value(resolved_data, "selected_provider_id")
|
||||
)
|
||||
@@ -786,6 +866,8 @@ class MoviePilotAgent:
|
||||
"user_agent": user_agent,
|
||||
"use_proxy": bool(use_proxy),
|
||||
"thinking_level": thinking_level,
|
||||
"api_protocol": api_protocol,
|
||||
"web_search_mode": web_search_mode,
|
||||
}
|
||||
return self._llm_runtime_config
|
||||
|
||||
@@ -795,7 +877,17 @@ class MoviePilotAgent:
|
||||
:param streaming: 是否启用流式输出
|
||||
"""
|
||||
runtime_config = await self._resolve_llm_runtime_config()
|
||||
return await LLMHelper.get_llm(streaming=streaming, **runtime_config)
|
||||
return await LLMHelper.get_llm(
|
||||
streaming=streaming,
|
||||
prompt_cache_key=self._build_prompt_cache_key(),
|
||||
**runtime_config,
|
||||
)
|
||||
|
||||
def _build_prompt_cache_key(self) -> str:
|
||||
"""生成不暴露用户标识、且在同一会话内稳定的提示词缓存键。"""
|
||||
cache_identity = f"{self.user_id or ''}\x00{self.session_id}"
|
||||
digest = hashlib.sha256(cache_identity.encode("utf-8")).hexdigest()[:32]
|
||||
return f"moviepilot-agent-{digest}"
|
||||
|
||||
@classmethod
|
||||
def _has_image_input_content(cls, content: Any) -> bool:
|
||||
@@ -993,6 +1085,13 @@ class MoviePilotAgent:
|
||||
allow_message_tools=self.allow_message_tools,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _filter_local_web_search_tools(tools: List, enabled: bool) -> List:
|
||||
"""按联网搜索策略保留或移除本地 search_web 工具。"""
|
||||
if enabled:
|
||||
return tools
|
||||
return [tool for tool in tools if getattr(tool, "name", None) != "search_web"]
|
||||
|
||||
def _refresh_tool_context(self, values: Dict[str, object]) -> None:
|
||||
"""
|
||||
刷新本轮工具共享上下文。
|
||||
@@ -1021,6 +1120,8 @@ class MoviePilotAgent:
|
||||
runtime_config.get("user_agent"),
|
||||
bool(runtime_config.get("use_proxy")),
|
||||
runtime_config.get("thinking_level"),
|
||||
runtime_config.get("api_protocol"),
|
||||
runtime_config.get("web_search_mode"),
|
||||
)
|
||||
|
||||
async def _agent_bundle_signature(self, streaming: bool) -> tuple[Any, ...]:
|
||||
@@ -1037,10 +1138,12 @@ class MoviePilotAgent:
|
||||
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(
|
||||
@@ -1097,6 +1200,39 @@ class MoviePilotAgent:
|
||||
allow_message_tools=False,
|
||||
)
|
||||
|
||||
async def _initialize_mcp_tools(self) -> List:
|
||||
"""
|
||||
初始化外部 MCP 工具列表。
|
||||
"""
|
||||
return await create_external_mcp_tools(
|
||||
session_id=self.session_id,
|
||||
user_id=self.user_id,
|
||||
channel=self.channel,
|
||||
source=self.source,
|
||||
username=self.username,
|
||||
stream_handler=self.stream_handler,
|
||||
agent_context=self._tool_context,
|
||||
)
|
||||
|
||||
async def _initialize_subagent_mcp_tools(self) -> List:
|
||||
"""
|
||||
初始化子代理可用的外部 MCP 工具列表。
|
||||
"""
|
||||
return await create_external_mcp_tools(
|
||||
session_id=self.session_id,
|
||||
user_id=self.user_id,
|
||||
channel=self.channel,
|
||||
source=self.source,
|
||||
username=self.username,
|
||||
stream_handler=None,
|
||||
agent_context={
|
||||
"user_reply_sent": False,
|
||||
"reply_mode": None,
|
||||
"should_dispatch_reply": False,
|
||||
"is_admin": bool(self._tool_context.get("is_admin")),
|
||||
},
|
||||
)
|
||||
|
||||
async def _create_agent(self, streaming: bool = False):
|
||||
"""
|
||||
创建 LangGraph Agent(使用 create_agent + SummarizationMiddleware)
|
||||
@@ -1116,6 +1252,8 @@ class MoviePilotAgent:
|
||||
# LLM 模型(用于 agent 执行)
|
||||
agent_model = await self._initialize_llm(streaming=streaming)
|
||||
self._sync_model_profile(agent_model)
|
||||
server_tools = LLMHelper.get_server_tools(agent_model)
|
||||
use_local_web_search = LLMHelper.should_use_local_web_search(agent_model)
|
||||
|
||||
# 为内部模型调用准备非流式 LLM,避免与用户流式回复复用同一实例。
|
||||
non_streaming_model = (
|
||||
@@ -1125,7 +1263,11 @@ class MoviePilotAgent:
|
||||
)
|
||||
|
||||
# 工具列表
|
||||
tools = self._initialize_tools()
|
||||
tools = self._filter_local_web_search_tools(
|
||||
self._initialize_tools(),
|
||||
enabled=use_local_web_search,
|
||||
)
|
||||
tools.extend(await self._initialize_mcp_tools())
|
||||
skills_middleware = SkillsMiddleware(
|
||||
sources=[str(agent_runtime_manager.skills_dir)],
|
||||
bundled_skills_dir=str(settings.ROOT_PATH / "skills"),
|
||||
@@ -1142,9 +1284,15 @@ class MoviePilotAgent:
|
||||
activity_log_tools = list(
|
||||
getattr(activity_log_middleware, "tools", []) or []
|
||||
)
|
||||
subagent_tools = self._filter_local_web_search_tools(
|
||||
self._initialize_subagent_tools(),
|
||||
enabled=use_local_web_search,
|
||||
)
|
||||
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,
|
||||
server_tools=server_tools,
|
||||
stream_handler=self.stream_handler,
|
||||
)
|
||||
max_tools = settings.LLM_MAX_TOOLS
|
||||
@@ -1219,7 +1367,7 @@ class MoviePilotAgent:
|
||||
|
||||
agent = create_agent(
|
||||
model=agent_model,
|
||||
tools=[*tools, *skill_tools, *activity_log_tools],
|
||||
tools=[*tools, *skill_tools, *activity_log_tools, *server_tools],
|
||||
system_prompt=system_prompt,
|
||||
middleware=middlewares,
|
||||
checkpointer=InMemorySaver(),
|
||||
@@ -1588,22 +1736,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,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -1960,12 +2110,11 @@ class AgentManager:
|
||||
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
|
||||
@@ -2083,6 +2232,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:
|
||||
|
||||
@@ -5,13 +5,17 @@ import inspect
|
||||
import json
|
||||
import time
|
||||
from functools import wraps
|
||||
from typing import Any, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
from langchain_core.messages import AIMessage, AIMessageChunk
|
||||
|
||||
from app.core.config import settings
|
||||
from app.log import logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from app.agent.llm.server_tools import ServerToolResolution
|
||||
|
||||
|
||||
class LLMTestError(RuntimeError):
|
||||
"""LLM 测试调用异常,附带请求耗时。"""
|
||||
@@ -224,74 +228,76 @@ def _is_deepseek_thinking_enabled(model_name: str | None, extra_body: Any) -> bo
|
||||
return False
|
||||
|
||||
|
||||
def _patch_deepseek_reasoning_content_support():
|
||||
"""
|
||||
修补 langchain-deepseek 在 tool-call 场景下遗漏 reasoning_content 回传的问题。
|
||||
|
||||
DeepSeek thinking mode 要求:若 assistant 历史消息包含 tool_calls,
|
||||
后续请求中必须带回该条消息的顶层 reasoning_content。
|
||||
某些 langchain-deepseek 版本虽然能从响应中拿到 reasoning_content,
|
||||
但不会在重放消息历史时写回请求载荷,导致 400。
|
||||
"""
|
||||
try:
|
||||
from langchain_deepseek import ChatDeepSeek
|
||||
except Exception as err:
|
||||
logger.debug(f"跳过 langchain-deepseek reasoning_content 修补:{err}")
|
||||
def _patch_interleaved_reasoning_request_support(
|
||||
model_cls: Any,
|
||||
*,
|
||||
patch_marker: str,
|
||||
thinking_filter: Any = None,
|
||||
normalize_deepseek_messages: bool = False,
|
||||
inject_missing_as_empty: bool = False,
|
||||
) -> None:
|
||||
"""为兼容模型统一补回工具调用历史中的 reasoning_content。"""
|
||||
if getattr(model_cls, patch_marker, False):
|
||||
return
|
||||
|
||||
if getattr(ChatDeepSeek, "_moviepilot_reasoning_content_patched", False):
|
||||
return
|
||||
|
||||
original_get_request_payload = getattr(ChatDeepSeek, "_get_request_payload", None)
|
||||
original_get_request_payload = getattr(model_cls, "_get_request_payload", None)
|
||||
if not callable(original_get_request_payload):
|
||||
logger.warning("langchain-deepseek 缺少 _get_request_payload,无法修补 reasoning_content")
|
||||
logger.warning(
|
||||
f"{model_cls.__name__} 缺少 _get_request_payload,无法修补 reasoning_content"
|
||||
)
|
||||
return
|
||||
|
||||
@wraps(original_get_request_payload)
|
||||
def _patched_get_request_payload(self, input_, *, stop=None, **kwargs):
|
||||
payload = original_get_request_payload(self, input_, stop=stop, **kwargs)
|
||||
if "messages" not in payload:
|
||||
return payload
|
||||
|
||||
extra_body = (getattr(self, "model_kwargs", None) or {}).get("extra_body")
|
||||
if not _is_deepseek_thinking_enabled(
|
||||
extra_body = getattr(self, "extra_body", None)
|
||||
if extra_body is None:
|
||||
extra_body = (getattr(self, "model_kwargs", None) or {}).get("extra_body")
|
||||
if thinking_filter is not None and not thinking_filter(
|
||||
getattr(self, "model_name", None) or getattr(self, "model", None),
|
||||
extra_body,
|
||||
):
|
||||
return payload
|
||||
|
||||
# 从原始 LangChain 消息中取回 reasoning_content。上游 payload 构造器
|
||||
# 不会自动透传这个 DeepSeek 扩展字段。
|
||||
messages = self._convert_input(input_).to_messages()
|
||||
|
||||
for i, message in enumerate(payload["messages"]):
|
||||
if message["role"] == "tool" and isinstance(message["content"], list):
|
||||
message["content"] = json.dumps(message["content"])
|
||||
elif message["role"] == "assistant":
|
||||
if isinstance(message["content"], list):
|
||||
# DeepSeek API 要求 assistant content 为字符串;工具场景下
|
||||
# LangChain 可能保留为内容块列表,这里只拼回可见文本块。
|
||||
text_parts = [
|
||||
block.get("text", "")
|
||||
for block in message["content"]
|
||||
if isinstance(block, dict) and block.get("type") == "text"
|
||||
]
|
||||
message["content"] = "".join(text_parts) if text_parts else ""
|
||||
|
||||
# DeepSeek thinking mode 要求历史 assistant 消息携带
|
||||
# reasoning_content,即便本地只保存到了 additional_kwargs。
|
||||
if (
|
||||
"reasoning_content" not in message
|
||||
and i < len(messages)
|
||||
and isinstance(messages[i], AIMessage)
|
||||
for index, payload_message in enumerate(payload["messages"]):
|
||||
if normalize_deepseek_messages:
|
||||
if payload_message.get("role") == "tool" and isinstance(
|
||||
payload_message.get("content"), list
|
||||
):
|
||||
message["reasoning_content"] = messages[i].additional_kwargs.get(
|
||||
"reasoning_content", ""
|
||||
payload_message["content"] = json.dumps(payload_message["content"])
|
||||
elif payload_message.get("role") == "assistant" and isinstance(
|
||||
payload_message.get("content"), list
|
||||
):
|
||||
payload_message["content"] = "".join(
|
||||
block.get("text", "")
|
||||
for block in payload_message["content"]
|
||||
if isinstance(block, dict) and block.get("type") == "text"
|
||||
)
|
||||
|
||||
if (
|
||||
payload_message.get("role") != "assistant"
|
||||
or index >= len(messages)
|
||||
or not isinstance(messages[index], AIMessage)
|
||||
or "reasoning_content" in payload_message
|
||||
):
|
||||
continue
|
||||
|
||||
reasoning_content = messages[index].additional_kwargs.get(
|
||||
"reasoning_content"
|
||||
)
|
||||
if reasoning_content is not None:
|
||||
payload_message["reasoning_content"] = reasoning_content
|
||||
elif inject_missing_as_empty:
|
||||
payload_message["reasoning_content"] = ""
|
||||
|
||||
return payload
|
||||
|
||||
ChatDeepSeek._get_request_payload = _patched_get_request_payload
|
||||
ChatDeepSeek._moviepilot_reasoning_content_patched = True
|
||||
logger.debug("已修补 langchain-deepseek thinking tool-call 的 reasoning_content 回传兼容性")
|
||||
model_cls._get_request_payload = _patched_get_request_payload
|
||||
setattr(model_cls, patch_marker, True)
|
||||
|
||||
|
||||
def _patch_openai_interleaved_reasoning_content_support():
|
||||
@@ -352,42 +358,10 @@ def _patch_openai_interleaved_reasoning_content_support():
|
||||
|
||||
_openai_base._moviepilot_reasoning_response_patched = True
|
||||
|
||||
if getattr(ChatOpenAI, "_moviepilot_interleaved_reasoning_patched", False):
|
||||
return
|
||||
|
||||
original_get_request_payload = getattr(ChatOpenAI, "_get_request_payload", None)
|
||||
if not callable(original_get_request_payload):
|
||||
logger.warning("langchain-openai 缺少 _get_request_payload,无法修补 reasoning_content")
|
||||
return
|
||||
|
||||
@wraps(original_get_request_payload)
|
||||
def _patched_get_request_payload(self, input_, *, stop=None, **kwargs):
|
||||
payload = original_get_request_payload(self, input_, stop=stop, **kwargs)
|
||||
if "messages" not in payload:
|
||||
return payload
|
||||
|
||||
messages = self._convert_input(input_).to_messages()
|
||||
for index, payload_message in enumerate(payload["messages"]):
|
||||
if (
|
||||
payload_message.get("role") != "assistant"
|
||||
or index >= len(messages)
|
||||
or not isinstance(messages[index], AIMessage)
|
||||
or "reasoning_content" in payload_message
|
||||
):
|
||||
continue
|
||||
|
||||
reasoning_content = messages[index].additional_kwargs.get(
|
||||
"reasoning_content"
|
||||
)
|
||||
if reasoning_content is not None:
|
||||
# 只回传模型真实返回过的思考字段。普通模型没有该字段时,
|
||||
# payload 保持原样,不额外塞未知参数。
|
||||
payload_message["reasoning_content"] = reasoning_content
|
||||
|
||||
return payload
|
||||
|
||||
ChatOpenAI._get_request_payload = _patched_get_request_payload
|
||||
ChatOpenAI._moviepilot_interleaved_reasoning_patched = True
|
||||
_patch_interleaved_reasoning_request_support(
|
||||
ChatOpenAI,
|
||||
patch_marker="_moviepilot_interleaved_reasoning_patched",
|
||||
)
|
||||
logger.debug("已修补 langchain-openai interleaved reasoning_content 回传兼容性")
|
||||
|
||||
|
||||
@@ -840,25 +814,122 @@ class LLMHelper:
|
||||
headers["User-Agent"] = normalized_user_agent
|
||||
return headers or None
|
||||
|
||||
@staticmethod
|
||||
def _matches_endpoint_host(base_url: str | None, expected_host: str) -> bool:
|
||||
"""严格匹配官方 API 主机,避免向兼容端点发送供应商专属参数。"""
|
||||
try:
|
||||
return (urlsplit(str(base_url or "")).hostname or "").lower() == expected_host
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
def _build_openai_prompt_cache_options(
|
||||
cls,
|
||||
*,
|
||||
provider: str,
|
||||
base_url: str | None,
|
||||
use_responses_api: bool | None,
|
||||
prompt_cache_key: str | None,
|
||||
default_headers: dict[str, str] | None,
|
||||
model_kwargs: dict[str, Any],
|
||||
) -> tuple[dict[str, str] | None, dict[str, Any]]:
|
||||
"""为 OpenAI 与 xAI 官方端点构造稳定提示词缓存路由参数。"""
|
||||
cache_key = str(prompt_cache_key or "").strip()
|
||||
headers = dict(default_headers or {})
|
||||
kwargs = dict(model_kwargs)
|
||||
provider_name = str(provider or "").strip().lower()
|
||||
if not cache_key:
|
||||
return headers or None, kwargs
|
||||
|
||||
is_openai = provider_name in {"chatgpt", "openai"} and cls._matches_endpoint_host(
|
||||
base_url,
|
||||
"api.openai.com",
|
||||
)
|
||||
is_xai = provider_name == "xai" and cls._matches_endpoint_host(
|
||||
base_url,
|
||||
"api.x.ai",
|
||||
)
|
||||
if not is_openai and not is_xai:
|
||||
return headers or None, kwargs
|
||||
|
||||
if is_xai and use_responses_api is not True:
|
||||
headers["x-grok-conv-id"] = cache_key
|
||||
return headers, kwargs
|
||||
|
||||
extra_body = dict(kwargs.get("extra_body") or {})
|
||||
extra_body["prompt_cache_key"] = cache_key
|
||||
kwargs["extra_body"] = extra_body
|
||||
return headers or None, kwargs
|
||||
|
||||
@staticmethod
|
||||
def _with_prompt_cache_control(
|
||||
model_cls: type,
|
||||
cache_control: dict[str, str],
|
||||
) -> type:
|
||||
"""创建在最终模型绑定阶段保留缓存控制参数的适配类。"""
|
||||
|
||||
class PromptCachingModel(model_cls):
|
||||
"""在 LangChain 工具绑定后仍保留提示词缓存参数的模型适配器。"""
|
||||
|
||||
def bind(self, **kwargs: Any) -> Any:
|
||||
"""绑定调用参数,并补入当前 Provider 的默认缓存控制。"""
|
||||
kwargs.setdefault("cache_control", dict(cache_control))
|
||||
return super().bind(**kwargs)
|
||||
|
||||
PromptCachingModel.__name__ = f"PromptCaching{model_cls.__name__}"
|
||||
return PromptCachingModel
|
||||
|
||||
@classmethod
|
||||
def _use_anthropic_prompt_cache(
|
||||
cls,
|
||||
*,
|
||||
provider: str,
|
||||
runtime: dict[str, Any],
|
||||
prompt_cache_key: str | None,
|
||||
) -> bool:
|
||||
"""判断当前运行时是否为可安全启用缓存的 Anthropic 官方端点。"""
|
||||
return (
|
||||
bool(str(prompt_cache_key or "").strip())
|
||||
and str(provider or "").strip().lower() == "anthropic"
|
||||
and str(runtime.get("runtime") or "").strip().lower()
|
||||
== "anthropic_compatible"
|
||||
and cls._matches_endpoint_host(
|
||||
runtime.get("base_url"),
|
||||
"api.anthropic.com",
|
||||
)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _should_use_openai_responses_api(
|
||||
cls,
|
||||
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
|
||||
@@ -872,6 +943,18 @@ 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:
|
||||
"""
|
||||
@@ -907,6 +990,36 @@ class LLMHelper:
|
||||
profile["moviepilot_provider_id"] = runtime_metadata["provider_id"]
|
||||
profile["moviepilot_base_url"] = runtime_metadata["base_url"]
|
||||
|
||||
@staticmethod
|
||||
def _attach_server_tool_metadata(
|
||||
model: Any,
|
||||
resolution: "ServerToolResolution",
|
||||
) -> None:
|
||||
"""把服务端工具解析结果挂到模型实例,供 Agent 组装工具列表。"""
|
||||
metadata = {
|
||||
"mode": resolution.mode,
|
||||
"use_local_web_search": resolution.use_local_web_search,
|
||||
"server_tools": [dict(tool) for tool in resolution.server_tools],
|
||||
"available": resolution.available,
|
||||
"reason": resolution.reason,
|
||||
}
|
||||
try:
|
||||
setattr(model, "_moviepilot_server_tool_metadata", metadata)
|
||||
except Exception:
|
||||
object.__setattr__(model, "_moviepilot_server_tool_metadata", metadata)
|
||||
|
||||
@staticmethod
|
||||
def get_server_tools(model: Any) -> list[dict[str, Any]]:
|
||||
"""读取模型已解析的服务端工具定义。"""
|
||||
metadata = getattr(model, "_moviepilot_server_tool_metadata", {}) or {}
|
||||
return [dict(tool) for tool in metadata.get("server_tools", [])]
|
||||
|
||||
@staticmethod
|
||||
def should_use_local_web_search(model: Any) -> bool:
|
||||
"""判断当前模型是否应保留 MoviePilot 本地联网搜索工具。"""
|
||||
metadata = getattr(model, "_moviepilot_server_tool_metadata", {}) or {}
|
||||
return bool(metadata.get("use_local_web_search", True))
|
||||
|
||||
@classmethod
|
||||
def _resolve_thinking_level(
|
||||
cls,
|
||||
@@ -952,7 +1065,11 @@ 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,
|
||||
web_search_mode: str | None = None,
|
||||
prompt_cache_key: str | None = None,
|
||||
):
|
||||
"""
|
||||
获取LLM实例
|
||||
@@ -967,7 +1084,16 @@ 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`` 保持原有自动判断。
|
||||
:param web_search_mode: 联网搜索模式
|
||||
(local/builtin/auto/disabled)。未显式传入时使用配置项
|
||||
``LLM_WEB_SEARCH_MODE``。
|
||||
:param prompt_cache_key: 同一 Agent 会话内稳定且脱敏的提示词缓存路由键。
|
||||
:return: LLM实例
|
||||
"""
|
||||
provider_name = str(provider if provider is not None else settings.LLM_PROVIDER).lower()
|
||||
@@ -978,6 +1104,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,
|
||||
)
|
||||
@@ -1005,6 +1132,40 @@ class LLMHelper:
|
||||
user_agent=user_agent_value,
|
||||
)
|
||||
model_name = runtime.get("model_id") or model_name
|
||||
from app.agent.llm.server_tools import (
|
||||
ServerToolRegistry,
|
||||
ServerToolUnavailableError,
|
||||
)
|
||||
|
||||
server_tool_resolution = ServerToolRegistry.resolve_web_search(
|
||||
provider=provider_name,
|
||||
model=model_name,
|
||||
mode=(
|
||||
web_search_mode
|
||||
if web_search_mode is not None
|
||||
else getattr(settings, "LLM_WEB_SEARCH_MODE", "local")
|
||||
),
|
||||
api_protocol=(
|
||||
api_protocol
|
||||
if api_protocol is not None
|
||||
else settings.LLM_API_PROTOCOL
|
||||
),
|
||||
base_url=runtime.get("base_url"),
|
||||
)
|
||||
if (
|
||||
server_tool_resolution.mode == "builtin"
|
||||
and not server_tool_resolution.available
|
||||
):
|
||||
raise ServerToolUnavailableError(
|
||||
provider=provider_name,
|
||||
model=str(model_name or ""),
|
||||
tool_id="web_search",
|
||||
)
|
||||
effective_api_protocol = (
|
||||
server_tool_resolution.required_api_protocol
|
||||
if server_tool_resolution.required_api_protocol == "responses"
|
||||
else api_protocol
|
||||
)
|
||||
default_headers = cls._build_openai_default_headers(
|
||||
runtime.get("default_headers"),
|
||||
user_agent=user_agent_value,
|
||||
@@ -1018,6 +1179,15 @@ class LLMHelper:
|
||||
provider=provider_name,
|
||||
model=model_name,
|
||||
runtime=runtime,
|
||||
api_protocol=effective_api_protocol,
|
||||
)
|
||||
default_headers, openai_model_kwargs = cls._build_openai_prompt_cache_options(
|
||||
provider=provider_name,
|
||||
base_url=runtime.get("base_url"),
|
||||
use_responses_api=use_responses_api,
|
||||
prompt_cache_key=prompt_cache_key,
|
||||
default_headers=default_headers,
|
||||
model_kwargs=thinking_kwargs,
|
||||
)
|
||||
llm_proxy = _resolve_llm_proxy(use_proxy)
|
||||
|
||||
@@ -1034,36 +1204,92 @@ 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,
|
||||
)
|
||||
elif runtime["runtime"] == "deepseek":
|
||||
elif (
|
||||
runtime["runtime"] == "deepseek"
|
||||
and server_tool_resolution.client_adapter != "openai_responses"
|
||||
and use_responses_api is not True
|
||||
):
|
||||
from langchain_deepseek import ChatDeepSeek
|
||||
|
||||
_patch_deepseek_reasoning_content_support()
|
||||
_patch_interleaved_reasoning_request_support(
|
||||
ChatDeepSeek,
|
||||
patch_marker="_moviepilot_reasoning_content_patched",
|
||||
thinking_filter=lambda model_name, extra_body: (
|
||||
_is_deepseek_thinking_enabled(model_name, extra_body)
|
||||
),
|
||||
normalize_deepseek_messages=True,
|
||||
inject_missing_as_empty=True,
|
||||
)
|
||||
model = ChatDeepSeek(
|
||||
model=model_name,
|
||||
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
|
||||
|
||||
bedrock_model_cls = ChatBedrockConverse
|
||||
if (
|
||||
str(prompt_cache_key or "").strip()
|
||||
and runtime.get("supports_prompt_cache")
|
||||
):
|
||||
bedrock_model_cls = cls._with_prompt_cache_control(
|
||||
ChatBedrockConverse,
|
||||
{"type": "default"},
|
||||
)
|
||||
|
||||
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 = bedrock_model_cls(
|
||||
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
|
||||
|
||||
model = ChatAnthropic(
|
||||
anthropic_model_cls = ChatAnthropic
|
||||
if cls._use_anthropic_prompt_cache(
|
||||
provider=provider_name,
|
||||
runtime=runtime,
|
||||
prompt_cache_key=prompt_cache_key,
|
||||
):
|
||||
anthropic_model_cls = cls._with_prompt_cache_control(
|
||||
ChatAnthropic,
|
||||
{"type": "ephemeral"},
|
||||
)
|
||||
|
||||
model = anthropic_model_cls(
|
||||
model=model_name,
|
||||
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,
|
||||
@@ -1084,7 +1310,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,
|
||||
@@ -1098,13 +1324,18 @@ class LLMHelper:
|
||||
),
|
||||
default_headers=default_headers,
|
||||
use_responses_api=use_responses_api,
|
||||
**thinking_kwargs,
|
||||
output_version=("responses/v1" if use_responses_api else None),
|
||||
**openai_model_kwargs,
|
||||
)
|
||||
|
||||
# 优先使用 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 {}
|
||||
@@ -1121,6 +1352,7 @@ class LLMHelper:
|
||||
}
|
||||
|
||||
cls._attach_runtime_metadata(model, runtime)
|
||||
cls._attach_server_tool_metadata(model, server_tool_resolution)
|
||||
return model
|
||||
|
||||
@staticmethod
|
||||
@@ -1178,25 +1410,38 @@ 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,
|
||||
web_search_mode: str | None = None,
|
||||
) -> dict:
|
||||
"""
|
||||
使用当前已保存配置执行一次最小 LLM 调用。
|
||||
使用当前配置或显式传入的临时配置执行一次最小 LLM 调用。
|
||||
|
||||
:param temperature: LLM 温度参数。未显式传入时沿用已保存配置。
|
||||
:param api_protocol: OpenAI 兼容接口 API 协议,未显式传入时沿用已保存配置。
|
||||
:param web_search_mode: 联网搜索模式,未显式传入时沿用已保存配置。
|
||||
"""
|
||||
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,
|
||||
"web_search_mode": web_search_mode,
|
||||
}
|
||||
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:
|
||||
@@ -1240,7 +1485,7 @@ class LLMHelper:
|
||||
try:
|
||||
from app.agent.llm.provider import LLMProviderManager
|
||||
|
||||
return await LLMProviderManager().list_models(
|
||||
models = await LLMProviderManager().list_models(
|
||||
provider_id=provider,
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
@@ -1249,16 +1494,25 @@ class LLMHelper:
|
||||
use_proxy=use_proxy,
|
||||
force_refresh=force_refresh,
|
||||
)
|
||||
return self._attach_server_tool_capabilities(
|
||||
provider,
|
||||
models,
|
||||
base_url=base_url,
|
||||
)
|
||||
except Exception as err:
|
||||
logger.debug(f"LLM provider 目录不可用,回退旧模型列表逻辑: {err}")
|
||||
if provider == "google":
|
||||
return [
|
||||
{"id": model_id, "name": model_id}
|
||||
for model_id in await self._get_google_models(
|
||||
api_key or "",
|
||||
use_proxy=use_proxy,
|
||||
)
|
||||
]
|
||||
return self._attach_server_tool_capabilities(
|
||||
provider,
|
||||
[
|
||||
{"id": model_id, "name": model_id}
|
||||
for model_id in await self._get_google_models(
|
||||
api_key or "",
|
||||
use_proxy=use_proxy,
|
||||
)
|
||||
],
|
||||
base_url=base_url,
|
||||
)
|
||||
try:
|
||||
from app.agent.llm.provider import LLMProviderManager
|
||||
|
||||
@@ -1272,16 +1526,40 @@ class LLMHelper:
|
||||
)
|
||||
except Exception:
|
||||
model_list_base_url = base_url
|
||||
return [
|
||||
{"id": model_id, "name": model_id}
|
||||
for model_id in await self._get_openai_compatible_models(
|
||||
provider,
|
||||
api_key or "",
|
||||
model_list_base_url,
|
||||
user_agent=user_agent,
|
||||
use_proxy=use_proxy,
|
||||
)
|
||||
]
|
||||
return self._attach_server_tool_capabilities(
|
||||
provider,
|
||||
[
|
||||
{"id": model_id, "name": model_id}
|
||||
for model_id in await self._get_openai_compatible_models(
|
||||
provider,
|
||||
api_key or "",
|
||||
model_list_base_url,
|
||||
user_agent=user_agent,
|
||||
use_proxy=use_proxy,
|
||||
)
|
||||
],
|
||||
base_url=base_url,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _attach_server_tool_capabilities(
|
||||
provider: str,
|
||||
models: List[dict[str, Any]],
|
||||
base_url: Optional[str] = None,
|
||||
) -> List[dict[str, Any]]:
|
||||
"""为模型目录附加通用服务端工具能力元数据。"""
|
||||
from app.agent.llm.server_tools import ServerToolRegistry
|
||||
|
||||
result = []
|
||||
for item in models:
|
||||
model_item = dict(item)
|
||||
model_item["server_tools"] = ServerToolRegistry.list_capabilities(
|
||||
provider=provider,
|
||||
model=str(model_item.get("id") or ""),
|
||||
base_url=base_url,
|
||||
)
|
||||
result.append(model_item)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
async def _get_google_models(api_key: str, use_proxy: bool | None = None) -> List[str]:
|
||||
|
||||
@@ -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",
|
||||
@@ -1536,6 +1665,20 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
await self.get_models_dev_data(use_proxy=use_proxy)
|
||||
).get(models_dev_provider_id, {}) or {}
|
||||
|
||||
@staticmethod
|
||||
def _models_dev_model_candidates(
|
||||
provider_id: str,
|
||||
model_id: str,
|
||||
) -> tuple[str, ...]:
|
||||
"""生成模型目录查询候选,兼容 Provider 添加的透明模型前缀。"""
|
||||
candidates = [model_id]
|
||||
if model_id.startswith("models/"):
|
||||
candidates.append(model_id.removeprefix("models/"))
|
||||
if provider_id == "amazon-bedrock" and "." in model_id:
|
||||
# Cross-region Inference Profile 会增加 us./eu./global. 等前缀。
|
||||
candidates.append(model_id.split(".", 1)[1])
|
||||
return tuple(dict.fromkeys(candidates))
|
||||
|
||||
async def _models_dev_model(
|
||||
self,
|
||||
provider_id: str,
|
||||
@@ -1555,15 +1698,32 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
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:
|
||||
for candidate in self._models_dev_model_candidates(provider_id, model_id):
|
||||
if candidate in models:
|
||||
return models[candidate]
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _metadata_supports_prompt_cache(metadata: Any) -> bool:
|
||||
"""从统一模型元数据中判断是否声明了提示词缓存能力。"""
|
||||
if not isinstance(metadata, dict):
|
||||
return False
|
||||
|
||||
explicit_capability = metadata.get("prompt_cache")
|
||||
if isinstance(explicit_capability, bool):
|
||||
return explicit_capability
|
||||
|
||||
capabilities = metadata.get("capabilities")
|
||||
if isinstance(capabilities, dict):
|
||||
explicit_capability = capabilities.get("prompt_cache")
|
||||
if isinstance(explicit_capability, bool):
|
||||
return explicit_capability
|
||||
|
||||
cost = metadata.get("cost")
|
||||
return isinstance(cost, dict) and any(
|
||||
key in cost for key in ("cache_read", "cache_write")
|
||||
)
|
||||
|
||||
def _cached_models_dev_model(
|
||||
self,
|
||||
provider_id: str,
|
||||
@@ -1590,11 +1750,7 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
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:
|
||||
for candidate in self._models_dev_model_candidates(provider_id, model_id):
|
||||
if candidate in models:
|
||||
return models[candidate]
|
||||
return None
|
||||
@@ -1743,6 +1899,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,
|
||||
@@ -1857,6 +2119,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
|
||||
@@ -2064,6 +2555,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,
|
||||
@@ -2641,6 +3139,9 @@ class LLMProviderManager(metaclass=Singleton):
|
||||
"model_id": model,
|
||||
"model_record": model_record,
|
||||
"model_metadata": model_metadata,
|
||||
"supports_prompt_cache": self._metadata_supports_prompt_cache(
|
||||
model_metadata
|
||||
),
|
||||
"default_headers": None,
|
||||
"use_responses_api": None,
|
||||
"auth_mode": "api_key",
|
||||
@@ -2731,6 +3232,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
|
||||
|
||||
249
app/agent/llm/server_tools.py
Normal file
249
app/agent/llm/server_tools.py
Normal file
@@ -0,0 +1,249 @@
|
||||
"""LLM 服务端工具能力注册与解析。"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from fnmatch import fnmatch
|
||||
from typing import Any, Optional
|
||||
|
||||
|
||||
WEB_SEARCH_MODES = frozenset({"local", "builtin", "auto", "disabled"})
|
||||
|
||||
|
||||
class ServerToolUnavailableError(ValueError):
|
||||
"""表示用户强制选择了当前模型不可用的服务端工具。"""
|
||||
|
||||
def __init__(self, *, provider: str, model: str, tool_id: str) -> None:
|
||||
"""初始化服务端工具不可用异常。"""
|
||||
self.provider = provider
|
||||
self.model = model
|
||||
self.tool_id = tool_id
|
||||
super().__init__(
|
||||
f"当前模型 {provider}/{model} 或接口地址不支持服务端联网搜索,"
|
||||
"请改用“自动”或“MoviePilot 本地搜索”"
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ServerToolCapability:
|
||||
"""描述一个模型可用的服务端工具能力。"""
|
||||
|
||||
tool_id: str
|
||||
provider_ids: tuple[str, ...]
|
||||
model_patterns: tuple[str, ...]
|
||||
required_api_protocol: str
|
||||
client_adapter: str
|
||||
tool_definition: dict[str, Any]
|
||||
base_url_patterns: tuple[str, ...] = ()
|
||||
match_without_base_url: bool = True
|
||||
|
||||
def matches(self, provider: str, model: str, base_url: Optional[str] = None) -> bool:
|
||||
"""判断给定 provider/model 是否匹配当前能力。"""
|
||||
normalized_provider = str(provider or "").strip().lower()
|
||||
normalized_model = str(model or "").strip().lower().removeprefix("models/")
|
||||
normalized_base_url = str(base_url or "").strip().lower()
|
||||
return (
|
||||
normalized_provider in self.provider_ids
|
||||
and any(fnmatch(normalized_model, pattern) for pattern in self.model_patterns)
|
||||
and (
|
||||
(not normalized_base_url and self.match_without_base_url)
|
||||
or not self.base_url_patterns
|
||||
or any(
|
||||
pattern in normalized_base_url
|
||||
for pattern in self.base_url_patterns
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
def serialize(self) -> dict[str, Any]:
|
||||
"""返回供 API 与前端使用的能力元数据。"""
|
||||
return {
|
||||
"id": self.tool_id,
|
||||
"required_api_protocol": self.required_api_protocol,
|
||||
"client_adapter": self.client_adapter,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ServerToolResolution:
|
||||
"""记录本次联网搜索模式解析后的执行策略。"""
|
||||
|
||||
mode: str
|
||||
use_local_web_search: bool
|
||||
server_tools: tuple[dict[str, Any], ...] = ()
|
||||
client_adapter: Optional[str] = None
|
||||
required_api_protocol: Optional[str] = None
|
||||
available: bool = False
|
||||
reason: Optional[str] = None
|
||||
|
||||
|
||||
class ServerToolRegistry:
|
||||
"""集中注册模型服务端工具,并解析通用执行策略。"""
|
||||
|
||||
_CAPABILITIES = (
|
||||
ServerToolCapability(
|
||||
tool_id="web_search",
|
||||
provider_ids=("chatgpt",),
|
||||
model_patterns=("gpt-5*", "gpt-4.1*", "o4-mini*"),
|
||||
base_url_patterns=("api.openai.com",),
|
||||
required_api_protocol="responses",
|
||||
client_adapter="openai_responses",
|
||||
tool_definition={"type": "web_search"},
|
||||
),
|
||||
ServerToolCapability(
|
||||
tool_id="web_search",
|
||||
provider_ids=("openai",),
|
||||
model_patterns=("gpt-5*", "gpt-4.1*", "o4-mini*"),
|
||||
base_url_patterns=("api.openai.com",),
|
||||
required_api_protocol="responses",
|
||||
client_adapter="openai_responses",
|
||||
tool_definition={"type": "web_search"},
|
||||
match_without_base_url=False,
|
||||
),
|
||||
ServerToolCapability(
|
||||
tool_id="web_search",
|
||||
provider_ids=("anthropic",),
|
||||
model_patterns=(
|
||||
"claude-opus-4*",
|
||||
"claude-sonnet-4*",
|
||||
"claude-haiku-4*",
|
||||
"claude-opus-5*",
|
||||
"claude-sonnet-5*",
|
||||
"claude-haiku-5*",
|
||||
"claude-fable-5*",
|
||||
"claude-mythos-5*",
|
||||
),
|
||||
base_url_patterns=("api.anthropic.com",),
|
||||
required_api_protocol="native",
|
||||
client_adapter="anthropic_native",
|
||||
tool_definition={
|
||||
"type": "web_search_20250305",
|
||||
"name": "web_search",
|
||||
},
|
||||
),
|
||||
ServerToolCapability(
|
||||
tool_id="web_search",
|
||||
provider_ids=("google",),
|
||||
model_patterns=("gemini-3*", "gemini-2.5*", "gemini-2.0-flash*"),
|
||||
required_api_protocol="native",
|
||||
client_adapter="google_native",
|
||||
tool_definition={"google_search": {}},
|
||||
),
|
||||
ServerToolCapability(
|
||||
tool_id="web_search",
|
||||
provider_ids=("xai",),
|
||||
model_patterns=("grok-4.5*",),
|
||||
base_url_patterns=("api.x.ai",),
|
||||
required_api_protocol="responses",
|
||||
client_adapter="openai_responses",
|
||||
tool_definition={"type": "web_search"},
|
||||
),
|
||||
ServerToolCapability(
|
||||
tool_id="web_search",
|
||||
provider_ids=("deepseek",),
|
||||
model_patterns=("deepseek-v4-flash",),
|
||||
base_url_patterns=("api.deepseek.com",),
|
||||
required_api_protocol="responses",
|
||||
client_adapter="openai_responses",
|
||||
tool_definition={"type": "web_search"},
|
||||
),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def normalize_web_search_mode(cls, mode: Optional[str]) -> str:
|
||||
"""规范化联网搜索模式,未知值回退为本地搜索。"""
|
||||
normalized = str(mode or "local").strip().lower()
|
||||
return normalized if normalized in WEB_SEARCH_MODES else "local"
|
||||
|
||||
@classmethod
|
||||
def get_capability(
|
||||
cls,
|
||||
*,
|
||||
provider: str,
|
||||
model: str,
|
||||
base_url: Optional[str] = None,
|
||||
tool_id: str,
|
||||
) -> Optional[ServerToolCapability]:
|
||||
"""查找指定模型的服务端工具能力。"""
|
||||
return next(
|
||||
(
|
||||
capability
|
||||
for capability in cls._CAPABILITIES
|
||||
if capability.tool_id == tool_id
|
||||
and capability.matches(provider, model, base_url)
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def list_capabilities(
|
||||
cls,
|
||||
*,
|
||||
provider: str,
|
||||
model: str,
|
||||
base_url: Optional[str] = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""列出指定模型可用的服务端工具能力。"""
|
||||
return [
|
||||
capability.serialize()
|
||||
for capability in cls._CAPABILITIES
|
||||
if capability.matches(provider, model, base_url)
|
||||
]
|
||||
|
||||
@classmethod
|
||||
def resolve_web_search(
|
||||
cls,
|
||||
*,
|
||||
provider: str,
|
||||
model: str,
|
||||
mode: Optional[str],
|
||||
api_protocol: Optional[str],
|
||||
base_url: Optional[str] = None,
|
||||
) -> ServerToolResolution:
|
||||
"""解析联网搜索应使用本地工具还是模型服务端工具。"""
|
||||
normalized_mode = cls.normalize_web_search_mode(mode)
|
||||
normalized_protocol = str(api_protocol or "auto").strip().lower()
|
||||
capability = cls.get_capability(
|
||||
provider=provider,
|
||||
model=model,
|
||||
base_url=base_url,
|
||||
tool_id="web_search",
|
||||
)
|
||||
|
||||
if normalized_mode == "disabled":
|
||||
return ServerToolResolution(
|
||||
mode=normalized_mode,
|
||||
use_local_web_search=False,
|
||||
reason="web_search_disabled",
|
||||
)
|
||||
if normalized_mode == "local":
|
||||
return ServerToolResolution(
|
||||
mode=normalized_mode,
|
||||
use_local_web_search=True,
|
||||
reason="local_web_search_selected",
|
||||
)
|
||||
if capability is None:
|
||||
return ServerToolResolution(
|
||||
mode=normalized_mode,
|
||||
use_local_web_search=normalized_mode == "auto",
|
||||
reason="builtin_web_search_unavailable",
|
||||
)
|
||||
if (
|
||||
normalized_mode == "auto"
|
||||
and normalized_protocol == "chat_completions"
|
||||
and capability.required_api_protocol == "responses"
|
||||
):
|
||||
return ServerToolResolution(
|
||||
mode=normalized_mode,
|
||||
use_local_web_search=True,
|
||||
available=True,
|
||||
reason="chat_completions_uses_local_fallback",
|
||||
)
|
||||
|
||||
return ServerToolResolution(
|
||||
mode=normalized_mode,
|
||||
use_local_web_search=False,
|
||||
server_tools=(dict(capability.tool_definition),),
|
||||
client_adapter=capability.client_adapter,
|
||||
required_api_protocol=capability.required_api_protocol,
|
||||
available=True,
|
||||
reason="builtin_web_search_selected",
|
||||
)
|
||||
600
app/agent/mcp.py
Normal file
600
app/agent/mcp.py
Normal file
@@ -0,0 +1,600 @@
|
||||
"""Agent 外部 MCP 客户端与配置管理。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Optional
|
||||
from urllib.parse import urljoin
|
||||
|
||||
from app.db.systemconfig_oper import SystemConfigOper
|
||||
from app.log import logger
|
||||
from app.schemas.agent import (
|
||||
AgentMcpServerConfig,
|
||||
AgentMcpServerTestResult,
|
||||
AgentMcpServerToolInfo,
|
||||
)
|
||||
from app.schemas.types import SystemConfigKey
|
||||
from app.utils.http import AsyncRequestUtils
|
||||
|
||||
MCP_PROTOCOL_VERSION = "2025-11-25"
|
||||
MCP_CLIENT_NAME = "MoviePilot Agent"
|
||||
DEFAULT_MCP_TIMEOUT = 30
|
||||
MCP_TOOL_NAME_PATTERN = re.compile(r"[^a-zA-Z0-9_]+")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AgentMcpToolSpec:
|
||||
"""已发现的外部 MCP 工具定义。"""
|
||||
|
||||
server: AgentMcpServerConfig
|
||||
name: str
|
||||
agent_tool_name: str
|
||||
description: str
|
||||
input_schema: dict[str, Any]
|
||||
|
||||
|
||||
def _normalize_identifier(value: str, fallback: str = "mcp") -> str:
|
||||
"""把服务器或工具名称转换为 Agent 工具可用的标识片段。"""
|
||||
normalized = MCP_TOOL_NAME_PATTERN.sub("_", str(value or "").strip())
|
||||
normalized = re.sub(r"_+", "_", normalized).strip("_").lower()
|
||||
if not normalized:
|
||||
normalized = fallback
|
||||
if normalized[0].isdigit():
|
||||
normalized = f"{fallback}_{normalized}"
|
||||
return normalized[:64]
|
||||
|
||||
|
||||
def _normalize_timeout(value: Any) -> int:
|
||||
"""规范化 MCP 连接和调用超时时间。"""
|
||||
try:
|
||||
timeout = int(value or DEFAULT_MCP_TIMEOUT)
|
||||
except (TypeError, ValueError):
|
||||
timeout = DEFAULT_MCP_TIMEOUT
|
||||
return min(max(timeout, 1), 600)
|
||||
|
||||
|
||||
def _normalize_string_dict(value: Any) -> dict[str, str]:
|
||||
"""规范化请求头和环境变量字典,移除空键。"""
|
||||
if not isinstance(value, dict):
|
||||
return {}
|
||||
normalized: dict[str, str] = {}
|
||||
for key, item in value.items():
|
||||
normalized_key = str(key or "").strip()
|
||||
if not normalized_key:
|
||||
continue
|
||||
normalized[normalized_key] = str(item or "")
|
||||
return normalized
|
||||
|
||||
|
||||
def _normalize_input_schema(value: Any) -> dict[str, Any]:
|
||||
"""规范化 MCP 工具参数 Schema,保证至少是 object schema。"""
|
||||
if not isinstance(value, dict):
|
||||
return {"type": "object", "properties": {}, "required": []}
|
||||
schema = dict(value)
|
||||
schema.setdefault("type", "object")
|
||||
schema.setdefault("properties", {})
|
||||
schema.setdefault("required", [])
|
||||
return schema
|
||||
|
||||
|
||||
def _build_agent_tool_name(server: AgentMcpServerConfig, tool_name: str) -> str:
|
||||
"""构造注入 Agent 的外部 MCP 工具名。"""
|
||||
prefix = server.tool_prefix or f"mcp_{server.name or server.id}"
|
||||
normalized_prefix = _normalize_identifier(prefix, fallback="mcp")
|
||||
normalized_tool_name = _normalize_identifier(tool_name, fallback="tool")
|
||||
if normalized_tool_name.startswith(f"{normalized_prefix}_"):
|
||||
return normalized_tool_name
|
||||
return f"{normalized_prefix}_{normalized_tool_name}"[:128]
|
||||
|
||||
|
||||
def _jsonrpc_message(method: str, params: Optional[dict[str, Any]] = None, *, request_id: Optional[str] = None) -> dict:
|
||||
"""构造 JSON-RPC 2.0 消息。"""
|
||||
payload = {"jsonrpc": "2.0", "method": method}
|
||||
if request_id is not None:
|
||||
payload["id"] = request_id
|
||||
if params is not None:
|
||||
payload["params"] = params
|
||||
return payload
|
||||
|
||||
|
||||
def _raise_for_jsonrpc_error(payload: Any) -> None:
|
||||
"""检查 JSON-RPC 响应错误并转换为运行时异常。"""
|
||||
if isinstance(payload, dict) and payload.get("error"):
|
||||
error = payload["error"]
|
||||
if isinstance(error, dict):
|
||||
message = error.get("message") or error
|
||||
else:
|
||||
message = error
|
||||
raise RuntimeError(f"MCP JSON-RPC 错误: {message}")
|
||||
|
||||
|
||||
def _extract_jsonrpc_result(payload: Any, request_id: str) -> Any:
|
||||
"""从 JSON-RPC 响应中提取 result 字段。"""
|
||||
if not isinstance(payload, dict):
|
||||
raise RuntimeError("MCP 响应不是有效 JSON 对象")
|
||||
if payload.get("id") != request_id:
|
||||
raise RuntimeError("MCP 响应 ID 与请求不匹配")
|
||||
_raise_for_jsonrpc_error(payload)
|
||||
return payload.get("result")
|
||||
|
||||
|
||||
async def _iter_sse_events(response) -> Any:
|
||||
"""按 SSE 事件格式迭代响应流。"""
|
||||
event_name = "message"
|
||||
data_lines: list[str] = []
|
||||
async for raw_line in response.aiter_lines():
|
||||
line = raw_line.rstrip("\r")
|
||||
if not line:
|
||||
if data_lines:
|
||||
yield {"event": event_name, "data": "\n".join(data_lines)}
|
||||
event_name = "message"
|
||||
data_lines = []
|
||||
continue
|
||||
if line.startswith(":"):
|
||||
continue
|
||||
field, _, value = line.partition(":")
|
||||
if value.startswith(" "):
|
||||
value = value[1:]
|
||||
if field == "event":
|
||||
event_name = value or "message"
|
||||
elif field == "data":
|
||||
data_lines.append(value)
|
||||
if data_lines:
|
||||
yield {"event": event_name, "data": "\n".join(data_lines)}
|
||||
|
||||
|
||||
def _parse_sse_text_response(text: str, request_id: str) -> Any:
|
||||
"""从非流式 SSE 文本响应中提取匹配请求的 JSON-RPC 结果。"""
|
||||
event_name = "message"
|
||||
data_lines: list[str] = []
|
||||
for raw_line in str(text or "").splitlines():
|
||||
line = raw_line.rstrip("\r")
|
||||
if not line:
|
||||
if data_lines:
|
||||
payload = _load_sse_json_payload(event_name, "\n".join(data_lines))
|
||||
if isinstance(payload, dict) and payload.get("id") == request_id:
|
||||
return _extract_jsonrpc_result(payload, request_id)
|
||||
event_name = "message"
|
||||
data_lines = []
|
||||
continue
|
||||
if line.startswith(":"):
|
||||
continue
|
||||
field, _, value = line.partition(":")
|
||||
if value.startswith(" "):
|
||||
value = value[1:]
|
||||
if field == "event":
|
||||
event_name = value or "message"
|
||||
elif field == "data":
|
||||
data_lines.append(value)
|
||||
if data_lines:
|
||||
payload = _load_sse_json_payload(event_name, "\n".join(data_lines))
|
||||
if isinstance(payload, dict) and payload.get("id") == request_id:
|
||||
return _extract_jsonrpc_result(payload, request_id)
|
||||
raise RuntimeError("MCP SSE 响应中未找到匹配请求")
|
||||
|
||||
|
||||
def _load_sse_json_payload(event_name: str, data: str) -> Optional[dict]:
|
||||
"""解析 SSE data 中的 JSON-RPC 消息。"""
|
||||
if event_name not in {"message", "messages"}:
|
||||
return None
|
||||
try:
|
||||
payload = json.loads(data)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
return payload if isinstance(payload, dict) else None
|
||||
|
||||
|
||||
class _StdioMcpSession:
|
||||
"""stdio MCP 会话,按一次操作生命周期启动外部进程。"""
|
||||
|
||||
def __init__(self, server: AgentMcpServerConfig) -> None:
|
||||
self.server = server
|
||||
self.process: Optional[asyncio.subprocess.Process] = None
|
||||
self.stderr_task: Optional[asyncio.Task] = None
|
||||
|
||||
async def __aenter__(self) -> "_StdioMcpSession":
|
||||
"""启动 stdio MCP 子进程。"""
|
||||
if not self.server.command:
|
||||
raise RuntimeError("stdio MCP 服务器缺少启动命令")
|
||||
env = os.environ.copy()
|
||||
env.update(self.server.env or {})
|
||||
self.process = await asyncio.create_subprocess_exec(
|
||||
self.server.command,
|
||||
*(self.server.args or []),
|
||||
stdin=asyncio.subprocess.PIPE,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
env=env,
|
||||
)
|
||||
self.stderr_task = asyncio.create_task(self._drain_stderr())
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb) -> None:
|
||||
"""结束 stdio MCP 子进程。"""
|
||||
if self.stderr_task:
|
||||
self.stderr_task.cancel()
|
||||
if not self.process:
|
||||
return
|
||||
if self.process.returncode is None:
|
||||
self.process.terminate()
|
||||
try:
|
||||
await asyncio.wait_for(self.process.wait(), timeout=2)
|
||||
except asyncio.TimeoutError:
|
||||
self.process.kill()
|
||||
await self.process.wait()
|
||||
|
||||
async def _drain_stderr(self) -> None:
|
||||
"""持续读取子进程 stderr,避免缓冲区阻塞。"""
|
||||
if not self.process or not self.process.stderr:
|
||||
return
|
||||
try:
|
||||
while True:
|
||||
line = await self.process.stderr.readline()
|
||||
if not line:
|
||||
break
|
||||
logger.debug(f"MCP stdio[{self.server.name}] stderr: {line.decode(errors='replace').strip()}")
|
||||
except asyncio.CancelledError:
|
||||
return
|
||||
|
||||
async def notify(self, method: str, params: Optional[dict[str, Any]] = None) -> None:
|
||||
"""发送不需要响应的 JSON-RPC 通知。"""
|
||||
await self._write_json(_jsonrpc_message(method, params))
|
||||
|
||||
async def request(self, method: str, params: Optional[dict[str, Any]] = None) -> Any:
|
||||
"""发送 JSON-RPC 请求并等待响应。"""
|
||||
request_id = uuid.uuid4().hex
|
||||
await self._write_json(_jsonrpc_message(method, params, request_id=request_id))
|
||||
while True:
|
||||
payload = await self._read_json()
|
||||
if payload.get("id") == request_id:
|
||||
return _extract_jsonrpc_result(payload, request_id)
|
||||
|
||||
async def _write_json(self, payload: dict) -> None:
|
||||
"""写入一行 JSON-RPC 消息。"""
|
||||
if not self.process or not self.process.stdin:
|
||||
raise RuntimeError("stdio MCP 进程未启动")
|
||||
data = json.dumps(payload, ensure_ascii=False, separators=(",", ":")) + "\n"
|
||||
self.process.stdin.write(data.encode("utf-8"))
|
||||
await self.process.stdin.drain()
|
||||
|
||||
async def _read_json(self) -> dict:
|
||||
"""从 stdout 读取一行 JSON-RPC 消息。"""
|
||||
if not self.process or not self.process.stdout:
|
||||
raise RuntimeError("stdio MCP 进程未启动")
|
||||
timeout = _normalize_timeout(self.server.timeout)
|
||||
while True:
|
||||
line = await asyncio.wait_for(self.process.stdout.readline(), timeout=timeout)
|
||||
if not line:
|
||||
raise RuntimeError("stdio MCP 进程已退出")
|
||||
try:
|
||||
payload = json.loads(line.decode("utf-8"))
|
||||
except ValueError:
|
||||
logger.debug(f"忽略非 JSON MCP stdout 行: {line!r}")
|
||||
continue
|
||||
if isinstance(payload, dict):
|
||||
return payload
|
||||
|
||||
|
||||
class _HttpMcpSession:
|
||||
"""Streamable HTTP MCP 会话。"""
|
||||
|
||||
def __init__(self, server: AgentMcpServerConfig) -> None:
|
||||
self.server = server
|
||||
self.session_id: Optional[str] = None
|
||||
|
||||
async def __aenter__(self) -> "_HttpMcpSession":
|
||||
"""进入 HTTP MCP 会话。"""
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb) -> None:
|
||||
"""退出 HTTP MCP 会话。"""
|
||||
return None
|
||||
|
||||
async def notify(self, method: str, params: Optional[dict[str, Any]] = None) -> None:
|
||||
"""发送不需要响应的 JSON-RPC 通知。"""
|
||||
await self._post(_jsonrpc_message(method, params), expect_response=False)
|
||||
|
||||
async def request(self, method: str, params: Optional[dict[str, Any]] = None) -> Any:
|
||||
"""发送 JSON-RPC 请求并等待响应。"""
|
||||
request_id = uuid.uuid4().hex
|
||||
return await self._post(
|
||||
_jsonrpc_message(method, params, request_id=request_id),
|
||||
expect_response=True,
|
||||
request_id=request_id,
|
||||
)
|
||||
|
||||
async def _post(
|
||||
self,
|
||||
payload: dict,
|
||||
*,
|
||||
expect_response: bool,
|
||||
request_id: Optional[str] = None,
|
||||
) -> Any:
|
||||
"""向 Streamable HTTP MCP 服务发送一条 JSON-RPC 消息。"""
|
||||
if not self.server.url:
|
||||
raise RuntimeError("HTTP MCP 服务器缺少 URL")
|
||||
headers = {
|
||||
"Accept": "application/json, text/event-stream",
|
||||
"Content-Type": "application/json",
|
||||
**(self.server.headers or {}),
|
||||
}
|
||||
if self.session_id:
|
||||
headers["Mcp-Session-Id"] = self.session_id
|
||||
response = await AsyncRequestUtils(
|
||||
headers=headers,
|
||||
timeout=_normalize_timeout(self.server.timeout),
|
||||
content_type="application/json",
|
||||
accept_type="application/json, text/event-stream",
|
||||
http2=False,
|
||||
).post_res(self.server.url, json=payload, raise_exception=True)
|
||||
try:
|
||||
if not response:
|
||||
raise RuntimeError("HTTP MCP 请求无响应")
|
||||
response.raise_for_status()
|
||||
session_id = response.headers.get("Mcp-Session-Id")
|
||||
if session_id:
|
||||
self.session_id = session_id
|
||||
if not expect_response:
|
||||
return None
|
||||
content_type = response.headers.get("content-type", "").lower()
|
||||
if "text/event-stream" in content_type:
|
||||
return _parse_sse_text_response(response.text, request_id or "")
|
||||
data = response.json()
|
||||
return _extract_jsonrpc_result(data, request_id or "")
|
||||
finally:
|
||||
if response is not None:
|
||||
await response.aclose()
|
||||
|
||||
|
||||
class _SseMcpSession:
|
||||
"""旧版 HTTP+SSE MCP 会话。"""
|
||||
|
||||
def __init__(self, server: AgentMcpServerConfig) -> None:
|
||||
self.server = server
|
||||
self.response = None
|
||||
self.endpoint: Optional[str] = None
|
||||
self._stream_manager = None
|
||||
self._event_iterator = None
|
||||
|
||||
async def __aenter__(self) -> "_SseMcpSession":
|
||||
"""打开 SSE 流并读取服务端回传的 POST endpoint。"""
|
||||
if not self.server.url:
|
||||
raise RuntimeError("SSE MCP 服务器缺少 URL")
|
||||
self._stream_manager = AsyncRequestUtils(
|
||||
headers={"Accept": "text/event-stream", **(self.server.headers or {})},
|
||||
timeout=_normalize_timeout(self.server.timeout),
|
||||
accept_type="text/event-stream",
|
||||
http2=False,
|
||||
).get_stream(self.server.url, raise_exception=True)
|
||||
self.response = await self._stream_manager.__aenter__()
|
||||
if not self.response:
|
||||
raise RuntimeError("SSE MCP 连接无响应")
|
||||
self.response.raise_for_status()
|
||||
self._event_iterator = _iter_sse_events(self.response).__aiter__()
|
||||
self.endpoint = await self._read_endpoint()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb) -> None:
|
||||
"""关闭 SSE 流。"""
|
||||
if self._stream_manager:
|
||||
await self._stream_manager.__aexit__(exc_type, exc, tb)
|
||||
|
||||
async def notify(self, method: str, params: Optional[dict[str, Any]] = None) -> None:
|
||||
"""发送不需要响应的 JSON-RPC 通知。"""
|
||||
await self._post(_jsonrpc_message(method, params))
|
||||
|
||||
async def request(self, method: str, params: Optional[dict[str, Any]] = None) -> Any:
|
||||
"""发送 JSON-RPC 请求并等待 SSE 流上的响应。"""
|
||||
request_id = uuid.uuid4().hex
|
||||
await self._post(_jsonrpc_message(method, params, request_id=request_id))
|
||||
timeout = _normalize_timeout(self.server.timeout)
|
||||
while True:
|
||||
event = await asyncio.wait_for(self._event_iterator.__anext__(), timeout=timeout)
|
||||
payload = _load_sse_json_payload(event.get("event", ""), event.get("data", ""))
|
||||
if isinstance(payload, dict) and payload.get("id") == request_id:
|
||||
return _extract_jsonrpc_result(payload, request_id)
|
||||
|
||||
async def _read_endpoint(self) -> str:
|
||||
"""读取 SSE endpoint 事件中的 POST 地址。"""
|
||||
timeout = _normalize_timeout(self.server.timeout)
|
||||
while True:
|
||||
event = await asyncio.wait_for(self._event_iterator.__anext__(), timeout=timeout)
|
||||
if event.get("event") != "endpoint":
|
||||
continue
|
||||
endpoint = str(event.get("data") or "").strip()
|
||||
if not endpoint:
|
||||
continue
|
||||
return urljoin(self.server.url, endpoint)
|
||||
|
||||
async def _post(self, payload: dict) -> None:
|
||||
"""向 SSE 握手返回的 endpoint 发送 JSON-RPC 消息。"""
|
||||
if not self.endpoint:
|
||||
raise RuntimeError("SSE MCP endpoint 未初始化")
|
||||
response = await AsyncRequestUtils(
|
||||
headers={
|
||||
"Accept": "application/json",
|
||||
"Content-Type": "application/json",
|
||||
**(self.server.headers or {}),
|
||||
},
|
||||
timeout=_normalize_timeout(self.server.timeout),
|
||||
content_type="application/json",
|
||||
accept_type="application/json",
|
||||
http2=False,
|
||||
).post_res(self.endpoint, json=payload, raise_exception=True)
|
||||
try:
|
||||
if not response:
|
||||
raise RuntimeError("SSE MCP POST 请求无响应")
|
||||
response.raise_for_status()
|
||||
finally:
|
||||
if response is not None:
|
||||
await response.aclose()
|
||||
|
||||
|
||||
async def _open_mcp_session(server: AgentMcpServerConfig):
|
||||
"""根据配置创建对应的 MCP 传输会话。"""
|
||||
transport = "http" if server.transport == "streamable_http" else server.transport
|
||||
if transport == "stdio":
|
||||
return _StdioMcpSession(server)
|
||||
if transport == "sse":
|
||||
return _SseMcpSession(server)
|
||||
if transport == "http":
|
||||
return _HttpMcpSession(server)
|
||||
raise RuntimeError(f"不支持的 MCP 传输协议: {server.transport}")
|
||||
|
||||
|
||||
class AgentMcpManager:
|
||||
"""管理 Agent 外部 MCP 服务器配置、工具发现和工具调用。"""
|
||||
|
||||
def get_servers(self) -> list[AgentMcpServerConfig]:
|
||||
"""读取已保存的外部 MCP 服务器配置。"""
|
||||
raw_servers = SystemConfigOper().get(SystemConfigKey.AIAgentMcpServers) or []
|
||||
if not isinstance(raw_servers, list):
|
||||
return []
|
||||
servers: list[AgentMcpServerConfig] = []
|
||||
for raw_server in raw_servers:
|
||||
try:
|
||||
servers.append(self.normalize_server(raw_server))
|
||||
except Exception as err:
|
||||
logger.warning(f"忽略无效的 Agent MCP 配置: {err}")
|
||||
return servers
|
||||
|
||||
async def save_servers(self, servers: list[AgentMcpServerConfig]) -> bool:
|
||||
"""保存外部 MCP 服务器配置。"""
|
||||
normalized_servers = [self.normalize_server(server).model_dump() for server in servers]
|
||||
return await SystemConfigOper().async_set(
|
||||
SystemConfigKey.AIAgentMcpServers,
|
||||
normalized_servers or None,
|
||||
)
|
||||
|
||||
def normalize_server(self, value: Any) -> AgentMcpServerConfig:
|
||||
"""规范化单个 MCP 服务器配置。"""
|
||||
if isinstance(value, AgentMcpServerConfig):
|
||||
raw_server = value.model_dump()
|
||||
elif isinstance(value, dict):
|
||||
raw_server = dict(value)
|
||||
else:
|
||||
raise ValueError("MCP 服务器配置必须是对象")
|
||||
|
||||
raw_server["id"] = str(raw_server.get("id") or uuid.uuid4().hex[:12]).strip()
|
||||
raw_server["name"] = str(raw_server.get("name") or raw_server["id"]).strip()
|
||||
raw_server["transport"] = str(raw_server.get("transport") or "stdio").strip()
|
||||
raw_server["description"] = str(raw_server.get("description") or "").strip() or None
|
||||
raw_server["command"] = str(raw_server.get("command") or "").strip() or None
|
||||
raw_server["args"] = [str(item) for item in raw_server.get("args") or []]
|
||||
raw_server["env"] = _normalize_string_dict(raw_server.get("env"))
|
||||
raw_server["url"] = str(raw_server.get("url") or "").strip() or None
|
||||
raw_server["headers"] = _normalize_string_dict(raw_server.get("headers"))
|
||||
raw_server["timeout"] = _normalize_timeout(raw_server.get("timeout"))
|
||||
raw_server["tool_prefix"] = str(raw_server.get("tool_prefix") or "").strip() or None
|
||||
raw_server["require_admin"] = bool(raw_server.get("require_admin", True))
|
||||
return AgentMcpServerConfig.model_validate(raw_server)
|
||||
|
||||
def config_signature(self) -> str:
|
||||
"""生成外部 MCP 配置签名,用于 Agent 图缓存失效。"""
|
||||
payload = [server.model_dump() for server in self.get_servers()]
|
||||
raw_text = json.dumps(payload, ensure_ascii=False, sort_keys=True, default=str)
|
||||
return hashlib.sha256(raw_text.encode("utf-8")).hexdigest()
|
||||
|
||||
async def initialize_session(self, session) -> None:
|
||||
"""完成 MCP initialize 和 initialized 通知流程。"""
|
||||
await session.request(
|
||||
"initialize",
|
||||
{
|
||||
"protocolVersion": MCP_PROTOCOL_VERSION,
|
||||
"capabilities": {},
|
||||
"clientInfo": {
|
||||
"name": MCP_CLIENT_NAME,
|
||||
"version": "1.0.0",
|
||||
},
|
||||
},
|
||||
)
|
||||
await session.notify("notifications/initialized")
|
||||
|
||||
async def list_server_tools(self, server: AgentMcpServerConfig) -> list[AgentMcpToolSpec]:
|
||||
"""连接单个 MCP 服务器并读取工具列表。"""
|
||||
normalized_server = self.normalize_server(server)
|
||||
session_manager = await _open_mcp_session(normalized_server)
|
||||
async with session_manager as session:
|
||||
await self.initialize_session(session)
|
||||
result = await session.request("tools/list")
|
||||
tools_payload = result.get("tools", []) if isinstance(result, dict) else []
|
||||
tool_specs: list[AgentMcpToolSpec] = []
|
||||
for item in tools_payload:
|
||||
if not isinstance(item, dict) or not item.get("name"):
|
||||
continue
|
||||
tool_name = str(item["name"])
|
||||
tool_specs.append(
|
||||
AgentMcpToolSpec(
|
||||
server=normalized_server,
|
||||
name=tool_name,
|
||||
agent_tool_name=_build_agent_tool_name(normalized_server, tool_name),
|
||||
description=str(item.get("description") or ""),
|
||||
input_schema=_normalize_input_schema(item.get("inputSchema")),
|
||||
)
|
||||
)
|
||||
return tool_specs
|
||||
|
||||
async def list_enabled_tool_specs(self) -> list[AgentMcpToolSpec]:
|
||||
"""读取所有启用 MCP 服务器暴露的工具定义。"""
|
||||
tool_specs: list[AgentMcpToolSpec] = []
|
||||
seen_names: set[str] = set()
|
||||
for server in self.get_servers():
|
||||
if not server.enabled:
|
||||
continue
|
||||
try:
|
||||
for spec in await self.list_server_tools(server):
|
||||
if spec.agent_tool_name in seen_names:
|
||||
logger.warning(f"跳过重复的 MCP Agent 工具名: {spec.agent_tool_name}")
|
||||
continue
|
||||
tool_specs.append(spec)
|
||||
seen_names.add(spec.agent_tool_name)
|
||||
except Exception as err:
|
||||
logger.warning(f"读取 MCP 服务器 {server.name} 工具失败: {err}")
|
||||
return tool_specs
|
||||
|
||||
async def call_server_tool(
|
||||
self,
|
||||
server: AgentMcpServerConfig,
|
||||
tool_name: str,
|
||||
arguments: Optional[dict[str, Any]] = None,
|
||||
) -> Any:
|
||||
"""调用单个 MCP 服务器上的指定工具。"""
|
||||
normalized_server = self.normalize_server(server)
|
||||
session_manager = await _open_mcp_session(normalized_server)
|
||||
async with session_manager as session:
|
||||
await self.initialize_session(session)
|
||||
return await session.request(
|
||||
"tools/call",
|
||||
{
|
||||
"name": tool_name,
|
||||
"arguments": arguments or {},
|
||||
},
|
||||
)
|
||||
|
||||
async def test_server(self, server: AgentMcpServerConfig) -> AgentMcpServerTestResult:
|
||||
"""测试 MCP 服务器连接并返回工具列表。"""
|
||||
tool_specs = await self.list_server_tools(server)
|
||||
tools = [
|
||||
AgentMcpServerToolInfo(
|
||||
name=spec.name,
|
||||
agent_tool_name=spec.agent_tool_name,
|
||||
description=spec.description,
|
||||
input_schema=spec.input_schema,
|
||||
)
|
||||
for spec in tool_specs
|
||||
]
|
||||
return AgentMcpServerTestResult(
|
||||
success=True,
|
||||
message=f"连接成功,发现 {len(tools)} 个工具",
|
||||
tools=tools,
|
||||
tool_count=len(tools),
|
||||
)
|
||||
|
||||
|
||||
agent_mcp_manager = AgentMcpManager()
|
||||
@@ -3,7 +3,7 @@
|
||||
|
||||
按日期存储在 CONFIG_PATH/agent/activity/YYYY-MM-DD.md 中,
|
||||
每次 Agent 执行完毕后自动调用 LLM 对本轮对话生成简洁的活动摘要,
|
||||
并在每次 Agent 启动时注入轻量索引,完整日志由工具按需查询。
|
||||
系统提示词只注入稳定的检索规则,完整日志由工具按需查询。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
@@ -447,7 +447,7 @@ async def _summarize_with_llm(conversation_text: str) -> Optional[str]:
|
||||
llm = await LLMHelper.get_llm(streaming=False)
|
||||
prompt = SUMMARY_PROMPT.format(conversation=conversation_text)
|
||||
response = await llm.ainvoke(prompt)
|
||||
summary = response.content.strip()
|
||||
summary = LLMHelper.extract_text_content(response.content).strip()
|
||||
# 清理模型可能输出的前缀(如 "摘要:" "总结:")
|
||||
summary = re.sub(r"^(摘要|总结|活动记录)[::]\s*", "", summary)
|
||||
if summary.strip().upper() == SUMMARY_SKIP_MARKER:
|
||||
@@ -459,12 +459,8 @@ async def _summarize_with_llm(conversation_text: str) -> Optional[str]:
|
||||
|
||||
|
||||
ACTIVITY_LOG_SYSTEM_PROMPT = """<activity_log>
|
||||
<activity_log_index>
|
||||
{activity_log_index}
|
||||
</activity_log_index>
|
||||
|
||||
<activity_log_guidelines>
|
||||
The index only shows recent dates and entry counts, not full log contents.
|
||||
Activity log contents and indexes are not included in the default context.
|
||||
Use `query_activity_log` only when the user references previous work, asks to continue a prior task, or recent activity is clearly relevant.
|
||||
Activity logs are read-only and retained for {retention_days} days; use MEMORY.md for durable preferences.
|
||||
</activity_log_guidelines>
|
||||
@@ -473,10 +469,10 @@ ACTIVITY_LOG_SYSTEM_PROMPT = """<activity_log>
|
||||
|
||||
|
||||
class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, ResponseT]): # noqa
|
||||
"""自动记录 Agent 活动日志并注入轻量索引的中间件。
|
||||
"""自动记录 Agent 活动日志并注入稳定检索规则的中间件。
|
||||
|
||||
- abefore_agent: 加载近几天的活动日志索引
|
||||
- awrap_model_call: 将活动日志索引和检索规则注入系统提示词
|
||||
- awrap_model_call: 将固定的活动日志检索规则注入系统提示词
|
||||
- aafter_agent: 从本次对话中提取摘要并追加到当日日志文件
|
||||
|
||||
参数:
|
||||
@@ -516,31 +512,9 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
||||
"""获取指定日期的日志文件路径。"""
|
||||
return AsyncPath(self.activity_dir) / f"{date_str}.md"
|
||||
|
||||
def _format_activity_log(self, contents: dict[str, str]) -> str:
|
||||
"""格式化活动日志索引用于系统提示词注入。"""
|
||||
if not contents:
|
||||
return ACTIVITY_LOG_SYSTEM_PROMPT.format(
|
||||
activity_log_index="(近期暂无活动日志索引。需要历史上下文时可调用 query_activity_log。)",
|
||||
retention_days=self.retention_days,
|
||||
)
|
||||
|
||||
# 按日期排序(最近的在前)
|
||||
sorted_dates = sorted(contents.keys(), reverse=True)
|
||||
sections = []
|
||||
for date_str in sorted_dates:
|
||||
content = contents[date_str].strip()
|
||||
if content:
|
||||
sections.append(f"### {date_str}\n{content}")
|
||||
|
||||
if not sections:
|
||||
return ACTIVITY_LOG_SYSTEM_PROMPT.format(
|
||||
activity_log_index="(近期暂无活动日志索引。需要历史上下文时可调用 query_activity_log。)",
|
||||
retention_days=self.retention_days,
|
||||
)
|
||||
|
||||
log_body = "\n".join(sections)
|
||||
def _format_activity_log(self, _contents: dict[str, str]) -> str:
|
||||
"""生成不受活动日志内容变化影响的系统提示词。"""
|
||||
return ACTIVITY_LOG_SYSTEM_PROMPT.format(
|
||||
activity_log_index=log_body,
|
||||
retention_days=self.retention_days,
|
||||
)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -27,8 +27,10 @@ from app.agent.middleware.utils import append_to_system_message
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.log import logger
|
||||
|
||||
# 安全提示: SKILL.md 文件最大限制为 10MB,防止 DoS 攻击
|
||||
MAX_SKILL_FILE_SIZE = 10 * 1024 * 1024
|
||||
# 磁盘读取上限与模型返回上限分离,避免异常大的 Skill 文件撑爆内存或上下文。
|
||||
MAX_SKILL_FILE_SIZE = 1 * 1024 * 1024
|
||||
MAX_SKILL_RESULT_CHARS = 64 * 1024
|
||||
SKILL_CONTENT_TRUNCATION_SUFFIX = "\n...(Skill 内容已截断)"
|
||||
|
||||
# Agent Skills 规范约束 (https://agentskills.io/specification)
|
||||
MAX_SKILL_NAME_LENGTH = 64
|
||||
@@ -248,7 +250,17 @@ async def _alist_skills(source_path: AsyncPath) -> list[SkillMetadata]:
|
||||
for skill_path in skill_dirs:
|
||||
skill_md_path = skill_path / "SKILL.md"
|
||||
|
||||
skill_content = await skill_md_path.read_text(encoding="utf-8", errors="replace")
|
||||
stat = await skill_md_path.stat()
|
||||
if stat.st_size > MAX_SKILL_FILE_SIZE:
|
||||
logger.warning(
|
||||
"Skipping %s: file too large (%d bytes)",
|
||||
skill_md_path,
|
||||
stat.st_size,
|
||||
)
|
||||
continue
|
||||
skill_content = (await skill_md_path.read_bytes()).decode(
|
||||
"utf-8", errors="replace"
|
||||
)
|
||||
|
||||
# 解析元数据
|
||||
skill_metadata = _parse_skill_metadata(
|
||||
@@ -280,7 +292,16 @@ def _list_skills(source_path: Path) -> list[SkillMetadata]:
|
||||
skills: list[SkillMetadata] = []
|
||||
for skill_path in skill_dirs:
|
||||
skill_md_path = skill_path / "SKILL.md"
|
||||
skill_content = skill_md_path.read_text(encoding="utf-8", errors="replace")
|
||||
if skill_md_path.stat().st_size > MAX_SKILL_FILE_SIZE:
|
||||
logger.warning(
|
||||
"Skipping %s: file too large (%d bytes)",
|
||||
skill_md_path,
|
||||
skill_md_path.stat().st_size,
|
||||
)
|
||||
continue
|
||||
skill_content = skill_md_path.read_bytes().decode(
|
||||
"utf-8", errors="replace"
|
||||
)
|
||||
skill_metadata = _parse_skill_metadata(
|
||||
content=skill_content,
|
||||
skill_path=str(skill_md_path),
|
||||
@@ -456,6 +477,46 @@ class _SkillToolProvider:
|
||||
raw_content = await handle.read(MAX_SKILL_FILE_SIZE)
|
||||
return raw_content.decode("utf-8", errors="replace"), truncated
|
||||
|
||||
@staticmethod
|
||||
def _serialize_skill_payload(payload: dict[str, Any]) -> str:
|
||||
"""序列化 Skill 返回值,并严格限制最终进入模型的字符数。"""
|
||||
serialized = json.dumps(payload, ensure_ascii=False, indent=2)
|
||||
if len(serialized) <= MAX_SKILL_RESULT_CHARS:
|
||||
return serialized
|
||||
|
||||
original_content = str(payload.get("content") or "")
|
||||
truncated_payload = dict(payload)
|
||||
truncated_payload["truncated"] = True
|
||||
low = 0
|
||||
high = len(original_content)
|
||||
best_result = json.dumps(
|
||||
{
|
||||
**truncated_payload,
|
||||
"content": SKILL_CONTENT_TRUNCATION_SUFFIX.strip(),
|
||||
},
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
)
|
||||
while low <= high:
|
||||
middle = (low + high) // 2
|
||||
candidate = json.dumps(
|
||||
{
|
||||
**truncated_payload,
|
||||
"content": (
|
||||
original_content[:middle]
|
||||
+ SKILL_CONTENT_TRUNCATION_SUFFIX
|
||||
),
|
||||
},
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
)
|
||||
if len(candidate) <= MAX_SKILL_RESULT_CHARS:
|
||||
best_result = candidate
|
||||
low = middle + 1
|
||||
else:
|
||||
high = middle - 1
|
||||
return best_result
|
||||
|
||||
async def load_skill(self, name: str) -> str:
|
||||
"""加载指定 Skill 的完整说明并返回 JSON 字符串。"""
|
||||
logger.info(f"加载 Skill: name={name}")
|
||||
@@ -471,7 +532,7 @@ class _SkillToolProvider:
|
||||
)
|
||||
|
||||
content, truncated = await self._read_skill_content(skill["path"])
|
||||
return json.dumps(
|
||||
return self._serialize_skill_payload(
|
||||
{
|
||||
"success": True,
|
||||
"skill": {
|
||||
@@ -483,9 +544,7 @@ class _SkillToolProvider:
|
||||
},
|
||||
"content": content,
|
||||
"truncated": truncated,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
}
|
||||
)
|
||||
except Exception as err:
|
||||
logger.error(f"加载 Skill 失败: {err}", exc_info=True)
|
||||
|
||||
@@ -377,11 +377,13 @@ class _SubAgentAgentProvider:
|
||||
model: BaseChatModel,
|
||||
profiles: tuple[_SubAgentProfile, ...],
|
||||
tools: list[BaseTool],
|
||||
server_tools: Optional[list[dict[str, Any]]] = None,
|
||||
) -> None:
|
||||
"""初始化子代理执行器。"""
|
||||
self._model = model
|
||||
self._profiles = {profile.name: profile for profile in profiles}
|
||||
self._tools = tools
|
||||
self._server_tools = server_tools or []
|
||||
self._agents = {}
|
||||
self._default_agent_name = "general-purpose"
|
||||
|
||||
@@ -404,7 +406,7 @@ class _SubAgentAgentProvider:
|
||||
)
|
||||
agent = create_agent(
|
||||
model=self._model,
|
||||
tools=subagent_tools,
|
||||
tools=[*subagent_tools, *self._server_tools],
|
||||
system_prompt=profile.prompt,
|
||||
name=profile.name,
|
||||
)
|
||||
@@ -462,16 +464,19 @@ class MoviePilotSubAgentMiddleware(AgentMiddleware):
|
||||
model: BaseChatModel,
|
||||
profiles: tuple[_SubAgentProfile, ...],
|
||||
tools: list[BaseTool],
|
||||
server_tools: Optional[list[dict[str, Any]]] = None,
|
||||
system_prompt: str = SUBAGENT_PARENT_PROMPT,
|
||||
task_description: str = SUBAGENT_TASK_DESCRIPTION,
|
||||
stream_handler: Any = None,
|
||||
) -> None:
|
||||
"""初始化同步子代理中间件。"""
|
||||
self.system_prompt = system_prompt
|
||||
self.stream_handler = stream_handler
|
||||
self._provider = _SubAgentAgentProvider(
|
||||
model=model,
|
||||
profiles=profiles,
|
||||
tools=tools,
|
||||
server_tools=server_tools,
|
||||
)
|
||||
self.tools = [
|
||||
StructuredTool.from_function(
|
||||
@@ -549,6 +554,7 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
||||
model: BaseChatModel,
|
||||
profiles: tuple[_SubAgentProfile, ...],
|
||||
tools: list[BaseTool],
|
||||
server_tools: Optional[list[dict[str, Any]]] = None,
|
||||
task_description: str = SUBAGENT_CONTROL_DESCRIPTION,
|
||||
stream_handler: Any = None,
|
||||
) -> None:
|
||||
@@ -558,6 +564,7 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
||||
model=model,
|
||||
profiles=profiles,
|
||||
tools=tools,
|
||||
server_tools=server_tools,
|
||||
)
|
||||
self._semaphore = asyncio.Semaphore(SUBAGENT_MAX_CONCURRENT_TASKS)
|
||||
self._tasks: dict[str, _SubAgentRuntimeTask] = {}
|
||||
@@ -1111,6 +1118,7 @@ def create_subagent_middlewares(
|
||||
*,
|
||||
model: BaseChatModel,
|
||||
tools: list[BaseTool],
|
||||
server_tools: Optional[list[dict[str, Any]]] = None,
|
||||
stream_handler: Any = None,
|
||||
) -> tuple[list[AgentMiddleware], list[BaseTool]]:
|
||||
"""创建子代理中间件列表和任务工具列表。"""
|
||||
@@ -1120,12 +1128,14 @@ def create_subagent_middlewares(
|
||||
model=model,
|
||||
profiles=profiles,
|
||||
tools=tools,
|
||||
server_tools=server_tools or [],
|
||||
stream_handler=stream_handler,
|
||||
)
|
||||
control_middleware = SubAgentTaskControlMiddleware(
|
||||
model=model,
|
||||
profiles=profiles,
|
||||
tools=tools,
|
||||
server_tools=server_tools or [],
|
||||
stream_handler=stream_handler,
|
||||
)
|
||||
|
||||
|
||||
@@ -51,6 +51,18 @@ class UsageMiddleware(AgentMiddleware):
|
||||
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def _first_int(
|
||||
cls,
|
||||
candidates: tuple[tuple[Any, tuple[str, ...]], ...],
|
||||
) -> int | None:
|
||||
"""按优先级返回首个可用的 usage 整数值。"""
|
||||
for container, keys in candidates:
|
||||
value = cls._lookup_int(container, *keys)
|
||||
if value is not None:
|
||||
return value
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def _extract_model_name(cls, model: Any) -> str | None:
|
||||
return (
|
||||
@@ -82,6 +94,131 @@ class UsageMiddleware(AgentMiddleware):
|
||||
or {}
|
||||
)
|
||||
|
||||
input_token_details = None
|
||||
if usage_metadata:
|
||||
getter = getattr(usage_metadata, "get", None)
|
||||
input_token_details = (
|
||||
getter("input_token_details")
|
||||
if callable(getter)
|
||||
else getattr(usage_metadata, "input_token_details", None)
|
||||
)
|
||||
|
||||
cache_read_tokens = cls._first_int(
|
||||
(
|
||||
(
|
||||
input_token_details,
|
||||
(
|
||||
"cache_read",
|
||||
"cached_tokens",
|
||||
"cache_read_input_tokens",
|
||||
"cacheReadInputTokens",
|
||||
),
|
||||
),
|
||||
(
|
||||
token_usage,
|
||||
(
|
||||
"prompt_cache_hit_tokens",
|
||||
"cache_read_input_tokens",
|
||||
"cacheReadInputTokens",
|
||||
),
|
||||
),
|
||||
(
|
||||
response_metadata,
|
||||
(
|
||||
"prompt_cache_hit_tokens",
|
||||
"cache_read_input_tokens",
|
||||
"cacheReadInputTokens",
|
||||
"cached_tokens",
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
if cache_read_tokens is None:
|
||||
cache_read_tokens = cls._first_int(
|
||||
(
|
||||
(
|
||||
token_usage.get("prompt_tokens_details", {}),
|
||||
("cached_tokens", "cache_read"),
|
||||
),
|
||||
(
|
||||
token_usage.get("input_tokens_details", {}),
|
||||
("cached_tokens", "cache_read"),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
cache_write_tokens = cls._first_int(
|
||||
(
|
||||
(
|
||||
input_token_details,
|
||||
(
|
||||
"cache_creation",
|
||||
"cache_write",
|
||||
"cache_write_tokens",
|
||||
"cache_write_input_tokens",
|
||||
"cacheWriteInputTokens",
|
||||
),
|
||||
),
|
||||
(
|
||||
token_usage,
|
||||
(
|
||||
"cache_creation_input_tokens",
|
||||
"cache_write_tokens",
|
||||
"cache_write_input_tokens",
|
||||
"cacheWriteInputTokens",
|
||||
),
|
||||
),
|
||||
(
|
||||
response_metadata,
|
||||
(
|
||||
"cache_creation_input_tokens",
|
||||
"cache_write_tokens",
|
||||
"cache_write_input_tokens",
|
||||
"cacheWriteInputTokens",
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
if cache_write_tokens is None:
|
||||
cache_write_tokens = cls._first_int(
|
||||
(
|
||||
(
|
||||
token_usage.get("prompt_tokens_details", {}),
|
||||
("cache_write_tokens", "cache_creation"),
|
||||
),
|
||||
(
|
||||
token_usage.get("input_tokens_details", {}),
|
||||
("cache_write_tokens", "cache_creation"),
|
||||
),
|
||||
)
|
||||
)
|
||||
cache_write_ttl_tokens = sum(
|
||||
cls._lookup_int(
|
||||
input_token_details,
|
||||
ttl_key,
|
||||
)
|
||||
or 0
|
||||
for ttl_key in (
|
||||
"ephemeral_5m_input_tokens",
|
||||
"ephemeral_1h_input_tokens",
|
||||
)
|
||||
)
|
||||
if cache_write_ttl_tokens:
|
||||
cache_write_tokens = cache_write_ttl_tokens
|
||||
|
||||
cache_miss_tokens = cls._first_int(
|
||||
(
|
||||
(
|
||||
token_usage,
|
||||
("prompt_cache_miss_tokens", "cache_miss_input_tokens"),
|
||||
),
|
||||
(
|
||||
response_metadata,
|
||||
("prompt_cache_miss_tokens", "cache_miss_input_tokens"),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
if input_tokens is None:
|
||||
input_tokens = cls._lookup_int(
|
||||
token_usage,
|
||||
@@ -94,6 +231,27 @@ class UsageMiddleware(AgentMiddleware):
|
||||
"prompt_token_count",
|
||||
"input_tokens",
|
||||
)
|
||||
if input_tokens is None:
|
||||
bedrock_input_tokens = cls._lookup_int(token_usage, "inputTokens")
|
||||
if bedrock_input_tokens is not None:
|
||||
input_tokens = (
|
||||
bedrock_input_tokens
|
||||
+ (cache_read_tokens or 0)
|
||||
+ (cache_write_tokens or 0)
|
||||
)
|
||||
if input_tokens is None and any(
|
||||
value is not None
|
||||
for value in (
|
||||
cache_read_tokens,
|
||||
cache_write_tokens,
|
||||
cache_miss_tokens,
|
||||
)
|
||||
):
|
||||
input_tokens = (
|
||||
(cache_read_tokens or 0)
|
||||
+ (cache_write_tokens or 0)
|
||||
+ (cache_miss_tokens or 0)
|
||||
)
|
||||
|
||||
if output_tokens is None:
|
||||
output_tokens = cls._lookup_int(
|
||||
@@ -113,8 +271,24 @@ class UsageMiddleware(AgentMiddleware):
|
||||
if total_tokens is None:
|
||||
total_tokens = cls._lookup_int(response_metadata, "total_token_count")
|
||||
|
||||
has_cache_usage = any(
|
||||
value is not None
|
||||
for value in (
|
||||
cache_read_tokens,
|
||||
cache_write_tokens,
|
||||
cache_miss_tokens,
|
||||
)
|
||||
)
|
||||
has_usage = any(
|
||||
value is not None for value in (input_tokens, output_tokens, total_tokens)
|
||||
value is not None
|
||||
for value in (
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
total_tokens,
|
||||
cache_read_tokens,
|
||||
cache_write_tokens,
|
||||
cache_miss_tokens,
|
||||
)
|
||||
)
|
||||
resolved_input = input_tokens or 0
|
||||
resolved_output = output_tokens or 0
|
||||
@@ -123,12 +297,32 @@ class UsageMiddleware(AgentMiddleware):
|
||||
if total_tokens is not None
|
||||
else resolved_input + resolved_output
|
||||
)
|
||||
resolved_cache_read = cache_read_tokens or 0
|
||||
resolved_cache_write = cache_write_tokens or 0
|
||||
uncached_input_tokens = (
|
||||
cache_miss_tokens
|
||||
if cache_miss_tokens is not None
|
||||
else max(
|
||||
resolved_input - resolved_cache_read - resolved_cache_write,
|
||||
0,
|
||||
)
|
||||
)
|
||||
cache_hit_ratio = (
|
||||
resolved_cache_read / resolved_input
|
||||
if has_cache_usage and resolved_input
|
||||
else None
|
||||
)
|
||||
|
||||
return {
|
||||
"has_usage": has_usage,
|
||||
"cache_usage_available": has_cache_usage,
|
||||
"input_tokens": resolved_input,
|
||||
"output_tokens": resolved_output,
|
||||
"total_tokens": resolved_total,
|
||||
"cache_read_input_tokens": resolved_cache_read,
|
||||
"cache_write_input_tokens": resolved_cache_write,
|
||||
"uncached_input_tokens": uncached_input_tokens,
|
||||
"cache_hit_ratio": cache_hit_ratio,
|
||||
}
|
||||
|
||||
async def awrap_model_call(
|
||||
@@ -157,9 +351,14 @@ class UsageMiddleware(AgentMiddleware):
|
||||
if ai_message
|
||||
else {
|
||||
"has_usage": False,
|
||||
"cache_usage_available": False,
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
"total_tokens": 0,
|
||||
"cache_read_input_tokens": 0,
|
||||
"cache_write_input_tokens": 0,
|
||||
"uncached_input_tokens": 0,
|
||||
"cache_hit_ratio": None,
|
||||
}
|
||||
)
|
||||
context_window_tokens = self._extract_context_window_tokens(request.model)
|
||||
|
||||
@@ -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; large searches may be split with narrower globs, paths, or `rg --files` filters. Use `list_directory` to inspect one known directory or a supported remote storage backend; request its `limit`/`offset` page fields when more than the first page is needed, and use `read_file` when the exact local file is known. If `read_file` reports truncation, continue with smaller `start_line` and `end_line` ranges instead of assuming the file ended.
|
||||
- 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>
|
||||
@@ -81,7 +85,14 @@ You act as a proactive agent. Your goal is to fully resolve the user's media-rel
|
||||
</agent_core>
|
||||
|
||||
<communication_runtime>
|
||||
{verbose_spec}
|
||||
<progress_updates>
|
||||
- Base progress updates on meaningful changes in understanding or execution, not on elapsed time or the number of tool calls.
|
||||
- Do not send a progress update merely because work is starting or because one or two tools have finished. Work through a coherent batch of investigation first.
|
||||
- Send an intermediate update when you have a useful preliminary conclusion, complete or validate a meaningful stage, discover evidence that materially changes the working direction, or encounter a sustained blocker the user should know about.
|
||||
- Explain the result or new direction with enough context to be useful, including the key evidence and what you will do next. An update may use several sentences when the finding needs explanation; brevity is not a goal by itself.
|
||||
- Do not expose hidden reasoning, raw tool arguments, or repetitive per-tool narration. Do not repeat an unchanged status.
|
||||
- Continue working after each update. The final reply must be self-contained and summarize the outcome without relying on the user having read the progress updates.
|
||||
</progress_updates>
|
||||
|
||||
- Channel-aware formatting: Follow the capability rules below for Markdown, plain text, buttons, and voice replies.
|
||||
{button_choice_spec}
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -24,6 +24,7 @@ SYSTEM_TASKS_FILE = "System Tasks.yaml"
|
||||
SYSTEM_TASKS_SCHEMA_VERSION = 2
|
||||
COMMON_SHELL_COMMANDS = (
|
||||
"ssh",
|
||||
"sshpass",
|
||||
"scp",
|
||||
"sftp",
|
||||
"git",
|
||||
@@ -139,19 +140,6 @@ class PromptManager:
|
||||
markdown_spec = self._generate_formatting_instructions(caps)
|
||||
button_choice_spec = self._generate_button_choice_instructions(msg_channel)
|
||||
|
||||
# 啰嗦模式
|
||||
verbose_spec = ""
|
||||
if not settings.AI_AGENT_VERBOSE:
|
||||
verbose_spec = (
|
||||
"\n\n[Important Instruction] STRICTLY ENFORCED: "
|
||||
"If tools are needed, DO NOT output any conversational text, explanations, progress updates, "
|
||||
"or acknowledgements before the first tool call or between tool calls. "
|
||||
"Call tools directly without any transitional phrases. "
|
||||
"You MUST remain completely silent until all required tools have finished and you have the final result. "
|
||||
"Only then may you send one final user-facing reply. "
|
||||
"DO NOT output any intermediate content whatsoever."
|
||||
)
|
||||
|
||||
# MoviePilot系统信息
|
||||
moviepilot_info = self._get_moviepilot_info()
|
||||
voice_reply_spec = self._generate_voice_reply_instructions()
|
||||
@@ -159,7 +147,6 @@ class PromptManager:
|
||||
# 始终替换占位符,避免后续 .format() 时因残留花括号报 KeyError
|
||||
base_prompt = base_prompt.format(
|
||||
markdown_spec=markdown_spec,
|
||||
verbose_spec=verbose_spec,
|
||||
moviepilot_info=moviepilot_info,
|
||||
voice_reply_spec=voice_reply_spec,
|
||||
button_choice_spec=button_choice_spec,
|
||||
@@ -315,8 +302,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']}",
|
||||
]
|
||||
)
|
||||
|
||||
@@ -28,7 +28,6 @@ class ToolChain(ChainBase):
|
||||
# 单个工具结果的兜底上限。各工具仍应优先在自身逻辑中分页或摘要化;
|
||||
# 这里用于拦截遗漏路径,避免超大结果直接进入模型上下文。
|
||||
DEFAULT_TOOL_RESULT_MAX_CHARS = 64 * 1024
|
||||
MIN_TOOL_RESULT_PREVIEW_CHARS = 512
|
||||
|
||||
|
||||
def serialize_tool_result_for_agent(result: Any) -> str:
|
||||
@@ -59,20 +58,35 @@ def format_tool_result_for_agent(
|
||||
if not max_chars or max_chars <= 0 or len(formatted_result) <= max_chars:
|
||||
return formatted_result
|
||||
|
||||
preview_limit = max(MIN_TOOL_RESULT_PREVIEW_CHARS, max_chars)
|
||||
preview = formatted_result[:preview_limit]
|
||||
payload = {
|
||||
"tool_result_truncated": True,
|
||||
"tool_name": tool_name,
|
||||
"total_chars": len(formatted_result),
|
||||
"returned_chars": len(preview),
|
||||
"content_preview": preview,
|
||||
"message": (
|
||||
f"工具返回内容超过 {max_chars} 字符,已截断为预览;"
|
||||
"请使用更精确的筛选条件、分页参数或专用查询参数继续获取。"
|
||||
),
|
||||
}
|
||||
return json.dumps(payload, ensure_ascii=False, indent=2)
|
||||
def _dump_preview(preview: str) -> str:
|
||||
"""序列化截断结果,并让 returned_chars 与实际预览保持一致。"""
|
||||
payload = {
|
||||
"tool_result_truncated": True,
|
||||
"tool_name": tool_name,
|
||||
"total_chars": len(formatted_result),
|
||||
"returned_chars": len(preview),
|
||||
"content_preview": preview,
|
||||
"message": (
|
||||
f"工具返回内容超过 {max_chars} 字符,已截断为预览;"
|
||||
"请使用更精确的筛选条件、分页参数或专用查询参数继续获取。"
|
||||
),
|
||||
}
|
||||
return json.dumps(payload, ensure_ascii=False, indent=2)
|
||||
|
||||
# JSON 会转义换行、引号和反斜杠,预览本身等于上限时,最终返回值仍可能
|
||||
# 明显超限。通过二分查找预留包装开销,确保进入模型的最终字符串是硬上限。
|
||||
low = 0
|
||||
high = min(len(formatted_result), max_chars)
|
||||
best_result = _dump_preview("")
|
||||
while low <= high:
|
||||
middle = (low + high) // 2
|
||||
candidate = _dump_preview(formatted_result[:middle])
|
||||
if len(candidate) <= max_chars:
|
||||
best_result = candidate
|
||||
low = middle + 1
|
||||
else:
|
||||
high = middle - 1
|
||||
return best_result
|
||||
|
||||
|
||||
# 将常见的阻塞调用按能力域拆分到独立线程池,避免外部慢 IO 抢占同一批 worker。
|
||||
@@ -425,7 +439,6 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
||||
"""
|
||||
roots = [
|
||||
settings.CONFIG_PATH / "agent",
|
||||
settings.LOG_PATH,
|
||||
]
|
||||
resolved_roots = []
|
||||
for root in roots:
|
||||
@@ -461,7 +474,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(
|
||||
@@ -483,7 +496,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)
|
||||
@@ -509,8 +522,8 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
||||
return (
|
||||
"抱歉,您没有执行此工具的权限。"
|
||||
"只有渠道管理员或系统管理员才能执行工具操作。"
|
||||
"如需执行工具,请联系渠道管理员将您的用户ID添加到渠道管理员列表中,"
|
||||
"或联系系统管理员为您设置权限。"
|
||||
"如需执行工具,请联系管理员将您的用户ID添加到渠道管理员列表中(设定 -> 通知 -> 对应渠道配置 -> 管理员名单),"
|
||||
"或联系系统管理员为您设置管理员权限。"
|
||||
)
|
||||
|
||||
async def _has_channel_admin_permission(self) -> bool:
|
||||
@@ -621,7 +634,8 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
||||
发送工具通知消息。
|
||||
|
||||
WebAgent 渠道没有后端模块实例,前端流式面板通过 Agent 上下文中的
|
||||
回调直接接收通知;其它渠道继续走统一消息链。
|
||||
回调直接接收通知;无渠道的后台任务清空渠道侧定位信息后交由消息链广播,
|
||||
其它渠道继续走统一消息链。
|
||||
"""
|
||||
callback = self._agent_context.get("notification_callback")
|
||||
if (
|
||||
@@ -631,6 +645,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(
|
||||
@@ -649,5 +677,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
|
||||
|
||||
53
app/agent/tools/impl/_file_write_utils.py
Normal file
53
app/agent/tools/impl/_file_write_utils.py
Normal file
@@ -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()
|
||||
@@ -59,6 +59,10 @@ SYSTEMCONFIG_SETTING_METADATA = {
|
||||
"group": "ai_agent",
|
||||
"label": "AI 智能体配置",
|
||||
},
|
||||
SystemConfigKey.AIAgentMcpServers.value: {
|
||||
"group": "ai_agent",
|
||||
"label": "AI 智能体外部 MCP 服务器",
|
||||
},
|
||||
SystemConfigKey.CustomIdentifiers.value: {
|
||||
"group": "custom_identifiers",
|
||||
"label": "自定义识别词",
|
||||
|
||||
@@ -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"] = {
|
||||
|
||||
@@ -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
|
||||
@@ -183,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:
|
||||
@@ -225,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,
|
||||
@@ -245,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 = []
|
||||
|
||||
@@ -39,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)",
|
||||
@@ -97,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季(默认)"
|
||||
@@ -144,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,
|
||||
@@ -197,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,
|
||||
|
||||
@@ -12,8 +12,8 @@ from app.agent.tools.tags import ToolTag
|
||||
from app.helper.browser import BrowserSessionHelper
|
||||
from app.log import logger
|
||||
|
||||
# 页面内容最大长度
|
||||
MAX_CONTENT_LENGTH = 8000
|
||||
# 页面内容最大长度;保留在全局工具结果兜底上限以内。
|
||||
MAX_CONTENT_LENGTH = 12_000
|
||||
# 默认超时时间(秒)
|
||||
DEFAULT_TIMEOUT = 30
|
||||
# 截图最大宽度
|
||||
|
||||
161
app/agent/tools/impl/create_agent_task.py
Normal file
161
app/agent/tools/impl/create_agent_task.py
Normal file
@@ -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)
|
||||
50
app/agent/tools/impl/delete_agent_task.py
Normal file
50
app/agent/tools/impl/delete_agent_task.py
Normal file
@@ -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} 已删除"
|
||||
@@ -54,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,
|
||||
}
|
||||
)
|
||||
|
||||
# 发送事件
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""文件编辑工具"""
|
||||
"""文件精确编辑工具。"""
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Optional, Type
|
||||
@@ -7,6 +7,11 @@ 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
|
||||
|
||||
@@ -15,20 +20,46 @@ class EditFileInput(BaseModel):
|
||||
"""文件编辑工具的输入参数模型。"""
|
||||
|
||||
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. "
|
||||
"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 "
|
||||
"and log directories."
|
||||
"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:
|
||||
|
||||
@@ -7,6 +7,7 @@ import json
|
||||
import os
|
||||
import signal
|
||||
import subprocess
|
||||
from collections import deque
|
||||
from dataclasses import dataclass, field
|
||||
from tempfile import NamedTemporaryFile
|
||||
from typing import Any, Literal, Optional, TextIO, Type
|
||||
@@ -27,7 +28,9 @@ from app.log import logger
|
||||
|
||||
DEFAULT_TIMEOUT_SECONDS = 60
|
||||
MAX_TIMEOUT_SECONDS = 300
|
||||
MAX_OUTPUT_PREVIEW_BYTES = 10 * 1024
|
||||
MAX_OUTPUT_PREVIEW_BYTES = 32 * 1024
|
||||
MAX_OUTPUT_HEAD_BYTES = 16 * 1024
|
||||
MAX_OUTPUT_TAIL_BYTES = 16 * 1024
|
||||
READ_CHUNK_SIZE = 4096
|
||||
KILL_GRACE_SECONDS = 3
|
||||
COMMAND_CONCURRENCY_LIMIT = 2
|
||||
@@ -36,11 +39,13 @@ _command_semaphore = asyncio.Semaphore(COMMAND_CONCURRENCY_LIMIT)
|
||||
|
||||
@dataclass
|
||||
class _CommandOutput:
|
||||
"""保存前 10KB 预览,并在超限时将完整输出写入临时文件。"""
|
||||
"""保存命令头尾预览,并在超限时将完整输出写入临时文件。"""
|
||||
|
||||
preview_limit_bytes: int
|
||||
preview_entries: list[tuple[str, str]] = field(default_factory=list)
|
||||
tail_entries: deque[tuple[str, str]] = field(default_factory=deque)
|
||||
captured_bytes: int = 0
|
||||
tail_bytes: int = 0
|
||||
preview_truncated: bool = False
|
||||
temp_file_path: Optional[str] = None
|
||||
temp_file_handle: Optional[TextIO] = None
|
||||
@@ -93,10 +98,12 @@ class _CommandOutput:
|
||||
self.temp_file_handle = None
|
||||
|
||||
def append(self, stream_name: str, text: str) -> None:
|
||||
"""追加一段输出,超出预览上限后只保留完整日志文件。"""
|
||||
"""追加一段输出,超出预览上限后保留头尾预览和完整日志文件。"""
|
||||
if not text:
|
||||
return
|
||||
|
||||
self._append_tail(stream_name, text)
|
||||
|
||||
if self.temp_file_handle:
|
||||
self._write_chunk(stream_name, text)
|
||||
return
|
||||
@@ -117,6 +124,60 @@ class _CommandOutput:
|
||||
self.preview_entries.append((stream_name, preview))
|
||||
self.captured_bytes += len(preview.encode("utf-8"))
|
||||
|
||||
def _append_tail(self, stream_name: str, text: str) -> None:
|
||||
"""维护固定字节大小的尾部输出,方便定位测试和构建失败信息。"""
|
||||
self.tail_entries.append((stream_name, text))
|
||||
self.tail_bytes += len(text.encode("utf-8"))
|
||||
while self.tail_bytes > MAX_OUTPUT_TAIL_BYTES and self.tail_entries:
|
||||
old_stream, old_text = self.tail_entries.popleft()
|
||||
old_bytes = len(old_text.encode("utf-8"))
|
||||
overflow = self.tail_bytes - MAX_OUTPUT_TAIL_BYTES
|
||||
if old_bytes <= overflow:
|
||||
self.tail_bytes -= old_bytes
|
||||
continue
|
||||
kept_text = old_text.encode("utf-8")[overflow:].decode(
|
||||
"utf-8", errors="ignore"
|
||||
)
|
||||
kept_bytes = len(kept_text.encode("utf-8"))
|
||||
self.tail_bytes -= old_bytes
|
||||
if kept_text:
|
||||
self.tail_entries.appendleft((old_stream, kept_text))
|
||||
self.tail_bytes += kept_bytes
|
||||
|
||||
@staticmethod
|
||||
def _format_entries(entries: list[tuple[str, str]]) -> str:
|
||||
"""按 stdout/stderr 切换插入可读的输出分段标题。"""
|
||||
parts: list[str] = []
|
||||
last_stream: Optional[str] = None
|
||||
for stream_name, text in entries:
|
||||
if stream_name != last_stream:
|
||||
title = "标准输出" if stream_name == "stdout" else "错误输出"
|
||||
parts.append(f"\n[{title}]\n")
|
||||
last_stream = stream_name
|
||||
parts.append(text)
|
||||
return "".join(parts).strip()
|
||||
|
||||
@property
|
||||
def combined_preview(self) -> str:
|
||||
"""返回完整输出或头尾组合预览。"""
|
||||
if not self.preview_truncated:
|
||||
return self._format_entries(self.preview_entries)
|
||||
|
||||
head_entries: list[tuple[str, str]] = []
|
||||
remaining = MAX_OUTPUT_HEAD_BYTES
|
||||
for stream_name, text in self.preview_entries:
|
||||
if remaining <= 0:
|
||||
break
|
||||
clipped = self._clip_text_to_bytes(text, remaining)
|
||||
if clipped:
|
||||
head_entries.append((stream_name, clipped))
|
||||
remaining -= len(clipped.encode("utf-8"))
|
||||
head = self._format_entries(head_entries)
|
||||
tail = self._format_entries(list(self.tail_entries))
|
||||
return (
|
||||
f"{head}\n\n...(中间输出已省略,完整内容在临时文件中)...\n\n{tail}"
|
||||
).strip()
|
||||
|
||||
@property
|
||||
def stdout(self) -> str:
|
||||
"""返回当前保留的 stdout 预览。"""
|
||||
@@ -295,7 +356,7 @@ class ExecuteCommandTool(MoviePilotTool):
|
||||
stream_name: str,
|
||||
output: _CommandOutput,
|
||||
) -> None:
|
||||
"""按块读取一次性命令输出,只把前 10KB 保留在返回结果中。"""
|
||||
"""按块读取一次性命令输出,保留 32KB 头尾预览。"""
|
||||
while True:
|
||||
chunk = await stream.read(READ_CHUNK_SIZE)
|
||||
if not chunk:
|
||||
@@ -379,17 +440,16 @@ class ExecuteCommandTool(MoviePilotTool):
|
||||
file_note = "截至命令终止前的完整输出" if timed_out else "完整输出"
|
||||
result += (
|
||||
"\n\n提示:\n"
|
||||
f"命令输出超过 10KB,仅返回前 {MAX_OUTPUT_PREVIEW_BYTES} 字节内容。\n"
|
||||
f"命令输出超过 {MAX_OUTPUT_PREVIEW_BYTES // 1024}KB,"
|
||||
f"仅返回前后各 {MAX_OUTPUT_HEAD_BYTES // 1024}KB 预览。\n"
|
||||
f"{file_note}已写入临时文件: {output.temp_file_path}\n"
|
||||
"如需完整内容,请继续读取该文件。"
|
||||
)
|
||||
if output.stdout:
|
||||
result += f"\n\n标准输出:\n{output.stdout}"
|
||||
if output.stderr:
|
||||
result += f"\n\n错误输出:\n{output.stderr}"
|
||||
if output.combined_preview:
|
||||
result += f"\n\n命令输出预览:\n{output.combined_preview}"
|
||||
if output.preview_truncated:
|
||||
result += "\n\n...(仅展示前 10KB 内容)"
|
||||
if not output.stdout and not output.stderr:
|
||||
result += "\n\n...(仅展示前后各 16KB 内容)"
|
||||
if not output.combined_preview:
|
||||
result += "\n\n(无输出内容)"
|
||||
return result
|
||||
|
||||
|
||||
@@ -210,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"),
|
||||
|
||||
@@ -34,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",
|
||||
@@ -45,6 +53,8 @@ class GetSearchResultsInput(BaseModel):
|
||||
|
||||
|
||||
class GetSearchResultsTool(MoviePilotTool):
|
||||
"""获取并筛选最近一次种子搜索结果"""
|
||||
|
||||
name: str = "get_search_results"
|
||||
tags: list[str] = [
|
||||
ToolTag.Read,
|
||||
@@ -54,6 +64,7 @@ class GetSearchResultsTool(MoviePilotTool):
|
||||
args_schema: Type[BaseModel] = GetSearchResultsInput
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
"""返回工具执行提示"""
|
||||
return "获取搜索结果"
|
||||
|
||||
async def run(
|
||||
@@ -66,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:
|
||||
@@ -87,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,
|
||||
@@ -105,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 "没有符合筛选条件的搜索结果,请调整筛选条件"
|
||||
|
||||
@@ -135,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
|
||||
|
||||
@@ -15,21 +15,47 @@ from app.schemas.file import FileItem
|
||||
from app.utils.string import StringUtils
|
||||
|
||||
|
||||
DEFAULT_DIRECTORY_PAGE_SIZE = 50
|
||||
MAX_DIRECTORY_PAGE_SIZE = 200
|
||||
|
||||
|
||||
class ListDirectoryInput(BaseModel):
|
||||
"""查询文件系统目录内容工具的输入参数模型"""
|
||||
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')")
|
||||
limit: Optional[int] = Field(
|
||||
DEFAULT_DIRECTORY_PAGE_SIZE,
|
||||
ge=1,
|
||||
le=MAX_DIRECTORY_PAGE_SIZE,
|
||||
description=(
|
||||
f"Maximum items to return in this page (default: {DEFAULT_DIRECTORY_PAGE_SIZE}, "
|
||||
f"maximum: {MAX_DIRECTORY_PAGE_SIZE})"
|
||||
),
|
||||
)
|
||||
offset: Optional[int] = Field(
|
||||
0,
|
||||
ge=0,
|
||||
description="Number of sorted directory items to skip before this page",
|
||||
)
|
||||
|
||||
|
||||
class ListDirectoryTool(MoviePilotTool):
|
||||
"""分页查询本地或远程存储目录中的文件和子目录。"""
|
||||
|
||||
name: str = "list_directory"
|
||||
tags: list[str] = [
|
||||
ToolTag.Read,
|
||||
ToolTag.Directory,
|
||||
ToolTag.File,
|
||||
]
|
||||
description: str = "List actual files and folders in a file system directory (NOT configuration). Shows files and subdirectories with their names, types, sizes, and modification times. Returns up to 20 items and the total count if there are more items. Use 'query_directory_settings' to query directory configuration settings."
|
||||
description: str = (
|
||||
"List actual files and folders in a file system directory (NOT configuration). "
|
||||
"Shows files and subdirectories with their names, types, sizes, and modification "
|
||||
f"times. Returns a page of up to {DEFAULT_DIRECTORY_PAGE_SIZE} items with total "
|
||||
f"count and next offset; limit is capped at {MAX_DIRECTORY_PAGE_SIZE}. "
|
||||
"Use 'query_directory_settings' to query directory configuration settings."
|
||||
)
|
||||
args_schema: Type[BaseModel] = ListDirectoryInput
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
@@ -45,10 +71,14 @@ class ListDirectoryTool(MoviePilotTool):
|
||||
|
||||
@staticmethod
|
||||
def _list_directory_sync(
|
||||
path: str, storage: Optional[str] = "local", sort_by: Optional[str] = "name"
|
||||
path: str,
|
||||
storage: Optional[str] = "local",
|
||||
sort_by: Optional[str] = "name",
|
||||
limit: Optional[int] = DEFAULT_DIRECTORY_PAGE_SIZE,
|
||||
offset: Optional[int] = 0,
|
||||
) -> str:
|
||||
"""
|
||||
目录遍历可能触发本地磁盘或远程存储请求,统一放到线程池中执行。
|
||||
目录遍历可能触发本地磁盘或远程存储请求,统一放到线程池中执行并分页返回。
|
||||
"""
|
||||
if not path:
|
||||
return "错误:路径不能为空"
|
||||
@@ -64,9 +94,6 @@ class ListDirectoryTool(MoviePilotTool):
|
||||
|
||||
if file_list is None:
|
||||
return f"无法访问目录:{path},请检查路径是否正确或存储是否可用"
|
||||
if not file_list:
|
||||
return f"目录 {path} 为空"
|
||||
|
||||
if sort_by == "time":
|
||||
file_list.sort(key=lambda x: x.modify_time or 0, reverse=True)
|
||||
else:
|
||||
@@ -78,7 +105,14 @@ class ListDirectoryTool(MoviePilotTool):
|
||||
)
|
||||
|
||||
total_count = len(file_list)
|
||||
limited_list = file_list[:20]
|
||||
normalized_limit = max(
|
||||
1,
|
||||
min(int(limit or DEFAULT_DIRECTORY_PAGE_SIZE), MAX_DIRECTORY_PAGE_SIZE),
|
||||
)
|
||||
normalized_offset = max(0, int(offset or 0))
|
||||
limited_list = file_list[
|
||||
normalized_offset : normalized_offset + normalized_limit
|
||||
]
|
||||
simplified_items = []
|
||||
for item in limited_list:
|
||||
size_str = StringUtils.str_filesize(item.size) if item.size else None
|
||||
@@ -102,16 +136,39 @@ class ListDirectoryTool(MoviePilotTool):
|
||||
simplified["extension"] = item.extension
|
||||
simplified_items.append(simplified)
|
||||
|
||||
result_json = json.dumps(simplified_items, ensure_ascii=False, indent=2)
|
||||
if total_count > 20:
|
||||
return (
|
||||
f"注意:目录中共有 {total_count} 个项目,为节省上下文空间,仅显示前 20 个项目。\n\n"
|
||||
f"{result_json}"
|
||||
)
|
||||
return result_json
|
||||
returned_count = len(simplified_items)
|
||||
has_more = normalized_offset + returned_count < total_count
|
||||
return json.dumps(
|
||||
{
|
||||
"items": simplified_items,
|
||||
"total_count": total_count,
|
||||
"returned_count": returned_count,
|
||||
"limit": normalized_limit,
|
||||
"offset": normalized_offset,
|
||||
"has_more": has_more,
|
||||
"next_offset": (
|
||||
normalized_offset + returned_count if has_more else None
|
||||
),
|
||||
},
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
)
|
||||
|
||||
async def run(self, path: str, storage: Optional[str] = "local",
|
||||
sort_by: Optional[str] = "name", **kwargs) -> str:
|
||||
sort_by: Optional[str] = "name",
|
||||
limit: Optional[int] = DEFAULT_DIRECTORY_PAGE_SIZE,
|
||||
offset: Optional[int] = 0,
|
||||
**kwargs) -> str:
|
||||
"""
|
||||
分页查询指定目录的文件和子目录。
|
||||
|
||||
:param path: 要查询的目录路径
|
||||
:param storage: 存储类型,默认为本地存储
|
||||
:param sort_by: 排序方式,支持名称或修改时间
|
||||
:param limit: 当前页最大条数,最高不超过工具上限
|
||||
:param offset: 当前页起始偏移量
|
||||
:return: 包含项目列表和分页元数据的 JSON 字符串
|
||||
"""
|
||||
logger.info(f"执行工具: {self.name}, 参数: path={path}, storage={storage}, sort_by={sort_by}")
|
||||
|
||||
try:
|
||||
@@ -123,7 +180,13 @@ class ListDirectoryTool(MoviePilotTool):
|
||||
if resolved_path:
|
||||
path = str(resolved_path)
|
||||
return await self.run_blocking(
|
||||
"storage", self._list_directory_sync, path, storage, sort_by
|
||||
"storage",
|
||||
self._list_directory_sync,
|
||||
path,
|
||||
storage,
|
||||
sort_by,
|
||||
limit,
|
||||
offset,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"查询目录内容失败: {e}", exc_info=True)
|
||||
|
||||
98
app/agent/tools/impl/mcp.py
Normal file
98
app/agent/tools/impl/mcp.py
Normal file
@@ -0,0 +1,98 @@
|
||||
"""外部 MCP 工具适配器。"""
|
||||
|
||||
import json
|
||||
from typing import Any, Optional
|
||||
|
||||
from pydantic import PrivateAttr
|
||||
|
||||
from app.agent.mcp import AgentMcpToolSpec, agent_mcp_manager
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
|
||||
|
||||
class McpExternalTool(MoviePilotTool):
|
||||
"""将外部 MCP 工具包装为 MoviePilot Agent 工具。"""
|
||||
|
||||
name: str = "mcp_external_tool"
|
||||
tags: list[str] = [
|
||||
ToolTag.Read,
|
||||
ToolTag.Admin,
|
||||
]
|
||||
description: str = "Call an external MCP tool configured for MoviePilot Agent."
|
||||
args_schema: dict[str, Any] = {"type": "object", "properties": {}, "required": []}
|
||||
require_admin: bool = True
|
||||
|
||||
_spec: AgentMcpToolSpec = PrivateAttr()
|
||||
|
||||
def __init__(self, spec: AgentMcpToolSpec, session_id: str, user_id: str) -> None:
|
||||
super().__init__(
|
||||
session_id=session_id,
|
||||
user_id=user_id,
|
||||
name=spec.agent_tool_name,
|
||||
description=spec.description
|
||||
or f"Call external MCP tool {spec.name} on {spec.server.name}.",
|
||||
args_schema=spec.input_schema,
|
||||
require_admin=spec.server.require_admin,
|
||||
)
|
||||
self._spec = spec
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
"""根据 MCP 工具信息生成友好的提示消息。"""
|
||||
return f"调用 MCP 工具: {self._spec.server.name}/{self._spec.name}"
|
||||
|
||||
async def run(self, **kwargs) -> str:
|
||||
"""
|
||||
调用外部 MCP 工具。
|
||||
|
||||
:param kwargs: 传递给外部 MCP 工具的参数
|
||||
:return: MCP 工具返回内容
|
||||
"""
|
||||
result = await agent_mcp_manager.call_server_tool(
|
||||
server=self._spec.server,
|
||||
tool_name=self._spec.name,
|
||||
arguments=kwargs,
|
||||
)
|
||||
return self._format_mcp_result(result)
|
||||
|
||||
@staticmethod
|
||||
def _format_mcp_result(result: Any) -> str:
|
||||
"""将 MCP tools/call 返回结构转换为 Agent 可读文本。"""
|
||||
if isinstance(result, dict):
|
||||
content = result.get("content")
|
||||
if isinstance(content, list):
|
||||
parts = []
|
||||
for item in content:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
if item.get("type") == "text" and item.get("text") is not None:
|
||||
parts.append(str(item["text"]))
|
||||
elif item:
|
||||
parts.append(json.dumps(item, ensure_ascii=False, default=str))
|
||||
if parts:
|
||||
return "\n".join(parts)
|
||||
if result.get("isError"):
|
||||
return json.dumps(result, ensure_ascii=False, indent=2, default=str)
|
||||
if isinstance(result, str):
|
||||
return result
|
||||
return json.dumps(result, ensure_ascii=False, indent=2, default=str)
|
||||
|
||||
|
||||
async def create_external_mcp_tools(
|
||||
*,
|
||||
session_id: str,
|
||||
user_id: str,
|
||||
channel: Optional[str] = None,
|
||||
source: Optional[str] = None,
|
||||
username: Optional[str] = None,
|
||||
stream_handler=None,
|
||||
agent_context: Optional[dict] = None,
|
||||
) -> list[McpExternalTool]:
|
||||
"""创建当前已启用的外部 MCP Agent 工具列表。"""
|
||||
tools = []
|
||||
for spec in await agent_mcp_manager.list_enabled_tool_specs():
|
||||
tool = McpExternalTool(spec=spec, session_id=session_id, user_id=user_id)
|
||||
tool.set_message_attr(channel=channel, source=source, username=username)
|
||||
tool.set_stream_handler(stream_handler=stream_handler)
|
||||
tool.set_agent_context(agent_context=agent_context)
|
||||
tools.append(tool)
|
||||
return tools
|
||||
88
app/agent/tools/impl/query_agent_tasks.py
Normal file
88
app/agent/tools/impl/query_agent_tasks.py
Normal file
@@ -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,
|
||||
)
|
||||
@@ -44,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
|
||||
@@ -73,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)
|
||||
|
||||
@@ -77,8 +77,12 @@ def _build_tv_server_result(existing_seasons: OrderedDict, total_seasons: Ordere
|
||||
|
||||
class QueryLibraryExistsInput(BaseModel):
|
||||
"""查询媒体库工具的输入参数模型"""
|
||||
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")
|
||||
|
||||
|
||||
@@ -89,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
|
||||
@@ -119,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:
|
||||
@@ -135,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()
|
||||
|
||||
@@ -20,6 +20,10 @@ class QueryMediaDetailInput(BaseModel):
|
||||
"""查询媒体详情工具的输入参数模型"""
|
||||
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")
|
||||
|
||||
|
||||
@@ -29,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:
|
||||
@@ -59,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} 的媒体信息"
|
||||
@@ -139,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)
|
||||
|
||||
@@ -118,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}季"
|
||||
@@ -126,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")
|
||||
@@ -149,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"),
|
||||
|
||||
@@ -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,47 +11,69 @@ from app.log import logger
|
||||
|
||||
|
||||
class QuerySchedulersInput(BaseModel):
|
||||
"""查询定时服务工具的输入参数模型"""
|
||||
"""查询运行时定时服务的输入参数模型。"""
|
||||
|
||||
|
||||
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)}"
|
||||
|
||||
@@ -170,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,
|
||||
|
||||
@@ -97,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"),
|
||||
|
||||
@@ -63,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)"
|
||||
)
|
||||
@@ -104,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:
|
||||
@@ -130,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)
|
||||
|
||||
@@ -120,6 +120,14 @@ class QueryTransferHistoryTool(MoviePilotTool):
|
||||
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)
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""文件读取工具"""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Optional, Type
|
||||
|
||||
@@ -12,22 +14,40 @@ from app.log import logger
|
||||
|
||||
# 最大读取大小 50KB
|
||||
MAX_READ_SIZE = 50 * 1024
|
||||
READ_FILE_TRUNCATION_MESSAGE = (
|
||||
"文件内容超过50KB,本次结果已截断。"
|
||||
"请使用 start_line 和 end_line 参数指定行号范围分段读取。"
|
||||
)
|
||||
|
||||
|
||||
class ReadFileInput(BaseModel):
|
||||
"""文件读取工具的输入参数模型。"""
|
||||
|
||||
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,
|
||||
ToolTag.File,
|
||||
]
|
||||
description: str = "Read the content of a text file. Supports reading by line range. Each read is limited to 50KB; content exceeding this limit will be truncated."
|
||||
description: str = (
|
||||
"Read the content of a text file. Supports reading by line range. Each "
|
||||
"read is limited to 50KB; when content is truncated, continue with "
|
||||
"smaller start_line and end_line ranges."
|
||||
)
|
||||
args_schema: Type[BaseModel] = ReadFileInput
|
||||
|
||||
def get_tool_message(self, **kwargs) -> Optional[str]:
|
||||
@@ -36,8 +56,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 +82,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,8 +106,22 @@ class ReadFileTool(MoviePilotTool):
|
||||
content = content_bytes[:MAX_READ_SIZE].decode("utf-8", errors="replace")
|
||||
truncated = True
|
||||
|
||||
if include_metadata:
|
||||
payload = {
|
||||
"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,
|
||||
}
|
||||
if truncated:
|
||||
payload["truncation_message"] = READ_FILE_TRUNCATION_MESSAGE
|
||||
payload["content"] = content
|
||||
return json.dumps(payload, ensure_ascii=False, indent=2)
|
||||
|
||||
if truncated:
|
||||
return f"{content}\n\n[警告:文件内容已超过50KB限制,以上内容已被截断。请使用 start_line/end_line 参数分段读取。]"
|
||||
return f"{content}\n\n[警告:{READ_FILE_TRUNCATION_MESSAGE}]"
|
||||
|
||||
return content
|
||||
|
||||
|
||||
@@ -142,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"),
|
||||
@@ -167,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)
|
||||
|
||||
78
app/agent/tools/impl/run_agent_task.py
Normal file
78
app/agent/tools/impl/run_agent_task.py
Normal file
@@ -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}。"
|
||||
"执行完成后将通过已配置的通知渠道广播结果"
|
||||
)
|
||||
@@ -14,23 +14,32 @@ class RunSchedulerInput(BaseModel):
|
||||
|
||||
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})"
|
||||
|
||||
@@ -46,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
|
||||
)
|
||||
|
||||
@@ -10,6 +10,7 @@ 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):
|
||||
@@ -43,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
|
||||
@@ -83,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,
|
||||
@@ -92,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,
|
||||
|
||||
@@ -11,6 +11,7 @@ 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):
|
||||
@@ -59,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,
|
||||
@@ -68,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,
|
||||
|
||||
@@ -70,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,
|
||||
}
|
||||
|
||||
# 检查订阅状态
|
||||
|
||||
@@ -20,13 +20,18 @@ from ._torrent_search_utils import (
|
||||
|
||||
class SearchTorrentsInput(BaseModel):
|
||||
"""搜索种子工具的输入参数模型"""
|
||||
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] = [
|
||||
@@ -35,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
|
||||
@@ -62,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()
|
||||
@@ -81,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,
|
||||
@@ -107,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)
|
||||
|
||||
@@ -105,6 +105,7 @@ class SendLocalFileTool(MoviePilotTool):
|
||||
text=message,
|
||||
file_path=str(resolved_path),
|
||||
file_name=file_name or resolved_path.name,
|
||||
save_history=False,
|
||||
)
|
||||
)
|
||||
return "本地附件已发送"
|
||||
|
||||
@@ -100,6 +100,7 @@ class SendMessageTool(MoviePilotTool):
|
||||
title=title,
|
||||
text=text,
|
||||
image=image_url,
|
||||
save_history=False,
|
||||
)
|
||||
)
|
||||
self._agent_context["user_reply_sent"] = True
|
||||
|
||||
@@ -96,6 +96,7 @@ class SendVoiceMessageTool(MoviePilotTool):
|
||||
if voice_path and settings.AUDIO_OUTPUT_INCLUDE_TEXT
|
||||
else None
|
||||
),
|
||||
save_history=False,
|
||||
)
|
||||
)
|
||||
self._agent_context["user_reply_sent"] = True
|
||||
|
||||
@@ -38,6 +38,10 @@ class TransferFileInput(BaseModel):
|
||||
doubanid: Optional[str] = Field(
|
||||
None, description="Douban ID for media identification (optional)"
|
||||
)
|
||||
bangumiid: Optional[int] = Field(None, description="Bangumi media ID")
|
||||
anilistid: 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")
|
||||
season: Optional[int] = Field(
|
||||
None, description="Season number for TV shows (optional)"
|
||||
)
|
||||
@@ -109,6 +113,10 @@ class TransferFileTool(MoviePilotTool):
|
||||
media_type: Optional[str] = None,
|
||||
tmdbid: Optional[int] = None,
|
||||
doubanid: Optional[str] = None,
|
||||
bangumiid: Optional[int] = None,
|
||||
anilistid: Optional[int] = None,
|
||||
media_source: Optional[str] = None,
|
||||
media_id: Optional[str] = None,
|
||||
season: Optional[int] = None,
|
||||
transfer_type: Optional[str] = None,
|
||||
background: Optional[bool] = False,
|
||||
@@ -148,6 +156,10 @@ class TransferFileTool(MoviePilotTool):
|
||||
target_path=target_path_obj,
|
||||
tmdbid=tmdbid,
|
||||
doubanid=doubanid,
|
||||
bangumiid=bangumiid,
|
||||
anilistid=anilistid,
|
||||
media_source=media_source,
|
||||
media_id=media_id,
|
||||
mtype=media_type_enum,
|
||||
season=season,
|
||||
transfer_type=transfer_type,
|
||||
@@ -178,6 +190,10 @@ class TransferFileTool(MoviePilotTool):
|
||||
media_type: Optional[str] = None,
|
||||
tmdbid: Optional[int] = None,
|
||||
doubanid: Optional[str] = None,
|
||||
bangumiid: Optional[int] = None,
|
||||
anilistid: Optional[int] = None,
|
||||
media_source: Optional[str] = None,
|
||||
media_id: Optional[str] = None,
|
||||
season: Optional[int] = None,
|
||||
transfer_type: Optional[str] = None,
|
||||
background: Optional[bool] = False,
|
||||
@@ -200,6 +216,10 @@ class TransferFileTool(MoviePilotTool):
|
||||
media_type,
|
||||
tmdbid,
|
||||
doubanid,
|
||||
bangumiid,
|
||||
anilistid,
|
||||
media_source,
|
||||
media_id,
|
||||
season,
|
||||
transfer_type,
|
||||
background,
|
||||
|
||||
193
app/agent/tools/impl/update_agent_task.py
Normal file
193
app/agent/tools/impl/update_agent_task.py
Normal file
@@ -0,0 +1,193 @@
|
||||
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.agenttask_oper import AgentTaskOper
|
||||
from app.utils.timer import TimerUtils
|
||||
|
||||
|
||||
class UpdateAgentTaskInput(BaseModel):
|
||||
"""更新 Agent 自主定时任务的输入参数。"""
|
||||
|
||||
task_id: int = Field(..., ge=1, description="ID of the task to update.")
|
||||
name: Optional[str] = Field(None, min_length=1, max_length=100)
|
||||
content: Optional[str] = Field(None, min_length=1, max_length=10000)
|
||||
trigger_type: Optional[Literal["date", "cron"]] = Field(
|
||||
None,
|
||||
description="New trigger type. Must be provided together with trigger.",
|
||||
)
|
||||
trigger: Optional[str] = Field(
|
||||
None,
|
||||
min_length=1,
|
||||
max_length=200,
|
||||
description="New ISO 8601 date or five-field cron expression.",
|
||||
)
|
||||
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 together with trigger_type='date'."
|
||||
),
|
||||
)
|
||||
enabled: Optional[bool] = Field(
|
||||
None,
|
||||
description="Set false to pause the task or true to resume it.",
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_update(self) -> "UpdateAgentTaskInput":
|
||||
"""校验更新内容和触发参数组合。"""
|
||||
if self.name is not None:
|
||||
self.name = self.name.strip()
|
||||
if not self.name:
|
||||
raise ValueError("name 不能只包含空白字符")
|
||||
if self.content is not None:
|
||||
self.content = self.content.strip()
|
||||
if not self.content:
|
||||
raise ValueError("content 不能只包含空白字符")
|
||||
has_schedule_update = any(
|
||||
value is not None
|
||||
for value in (self.trigger_type, self.trigger, self.delay_minutes)
|
||||
)
|
||||
if has_schedule_update:
|
||||
if self.trigger_type is None:
|
||||
raise ValueError("修改触发配置时必须提供 trigger_type")
|
||||
if self.trigger_type == "date":
|
||||
if self.delay_minutes is not None:
|
||||
# 保持校验幂等,具体绝对时间在更新调度前只计算一次。
|
||||
self.trigger = None
|
||||
elif 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")
|
||||
if all(
|
||||
value is None
|
||||
for value in (
|
||||
self.name,
|
||||
self.content,
|
||||
self.trigger_type,
|
||||
self.enabled,
|
||||
)
|
||||
):
|
||||
raise ValueError("至少需要提供一个要更新的字段")
|
||||
return self
|
||||
|
||||
|
||||
class UpdateAgentTaskTool(MoviePilotTool):
|
||||
"""修改、暂停或恢复 Agent 自主定时任务。"""
|
||||
|
||||
name: str = "update_agent_task"
|
||||
tags: list[str] = [ToolTag.Write, ToolTag.AgentTask, ToolTag.Admin]
|
||||
description: str = (
|
||||
"Update an autonomous agent task's name, instructions, exact date or cron "
|
||||
"trigger, relative delay_minutes, or enabled state. Use enabled=false to pause "
|
||||
"and enabled=true to resume."
|
||||
)
|
||||
args_schema: Type[BaseModel] = UpdateAgentTaskInput
|
||||
require_admin: bool = True
|
||||
|
||||
def get_tool_message(self, **kwargs: object) -> Optional[str]:
|
||||
"""生成更新定时任务的提示消息。"""
|
||||
return f"更新自主定时任务:{kwargs.get('task_id', '')}"
|
||||
|
||||
def _update_task(self, payload: UpdateAgentTaskInput) -> Optional[dict]:
|
||||
"""更新当前用户的任务并刷新运行时调度。"""
|
||||
from app.scheduler import Scheduler
|
||||
|
||||
oper = AgentTaskOper()
|
||||
task = oper.get(task_id=payload.task_id, user_id=str(self._user_id))
|
||||
if not task:
|
||||
return None
|
||||
if task.last_status == "running":
|
||||
return {"error": f"Agent 定时任务 {payload.task_id} 正在执行,请稍后再修改"}
|
||||
|
||||
trigger_type = payload.trigger_type or task.trigger_type
|
||||
trigger_value = payload.trigger
|
||||
if 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")
|
||||
if trigger_value is None:
|
||||
trigger_value = (
|
||||
task.cron_expression if trigger_type == "cron" else task.run_at
|
||||
)
|
||||
enabled = task.enabled if payload.enabled is None else payload.enabled
|
||||
normalized_type, normalized_trigger = TimerUtils.normalize_schedule_trigger(
|
||||
trigger_type=trigger_type,
|
||||
trigger_value=trigger_value,
|
||||
timezone_name=settings.TZ,
|
||||
require_future=bool(enabled and trigger_type == "date"),
|
||||
)
|
||||
|
||||
update_payload = {}
|
||||
if payload.name is not None:
|
||||
update_payload["name"] = payload.name.strip()
|
||||
if payload.content is not None:
|
||||
update_payload["content"] = payload.content.strip()
|
||||
if payload.trigger_type is not None:
|
||||
update_payload.update(
|
||||
{
|
||||
"trigger_type": normalized_type,
|
||||
"cron_expression": (
|
||||
normalized_trigger if normalized_type == "cron" else None
|
||||
),
|
||||
"run_at": normalized_trigger if normalized_type == "date" else None,
|
||||
"last_status": "waiting",
|
||||
"last_result": None,
|
||||
}
|
||||
)
|
||||
if payload.enabled is not None:
|
||||
update_payload["enabled"] = payload.enabled
|
||||
if payload.enabled:
|
||||
update_payload["last_status"] = "waiting"
|
||||
|
||||
oper.update(
|
||||
task_id=payload.task_id,
|
||||
payload=update_payload,
|
||||
user_id=str(self._user_id),
|
||||
)
|
||||
scheduler = Scheduler()
|
||||
next_run_at = scheduler.update_agent_task_job(payload.task_id)
|
||||
updated_task = oper.get(task_id=payload.task_id, user_id=str(self._user_id))
|
||||
return oper.to_dict(
|
||||
updated_task,
|
||||
next_run_at=next_run_at,
|
||||
timezone=settings.TZ,
|
||||
)
|
||||
|
||||
async def run(
|
||||
self,
|
||||
task_id: int,
|
||||
name: Optional[str] = None,
|
||||
content: Optional[str] = None,
|
||||
trigger_type: Optional[str] = None,
|
||||
trigger: Optional[str] = None,
|
||||
delay_minutes: Optional[int] = None,
|
||||
enabled: Optional[bool] = None,
|
||||
**kwargs: object,
|
||||
) -> str:
|
||||
"""更新 Agent 自主定时任务。"""
|
||||
payload = UpdateAgentTaskInput(
|
||||
task_id=task_id,
|
||||
name=name,
|
||||
content=content,
|
||||
trigger_type=trigger_type,
|
||||
trigger=trigger,
|
||||
delay_minutes=delay_minutes,
|
||||
enabled=enabled,
|
||||
)
|
||||
task = await self.run_blocking("db", self._update_task, payload)
|
||||
if not task:
|
||||
return f"Agent 定时任务 {task_id} 不存在或不属于当前用户"
|
||||
if task.get("error"):
|
||||
return task["error"]
|
||||
return json.dumps(task, ensure_ascii=False, indent=2)
|
||||
@@ -7,6 +7,7 @@ from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.core.metainfo import clear_rust_parse_options_cache
|
||||
from app.db.systemconfig_oper import SystemConfigOper
|
||||
from app.log import logger
|
||||
from app.schemas.types import SystemConfigKey
|
||||
@@ -85,6 +86,7 @@ class UpdateCustomIdentifiersTool(MoviePilotTool):
|
||||
SystemConfigKey.CustomIdentifiers, value
|
||||
)
|
||||
if success:
|
||||
clear_rust_parse_options_cache()
|
||||
return json.dumps(
|
||||
{
|
||||
"success": True,
|
||||
|
||||
@@ -8,6 +8,7 @@ from pydantic import BaseModel, Field
|
||||
from app.agent.tools.base import MoviePilotTool
|
||||
from app.agent.tools.tags import ToolTag
|
||||
from app.chain.download import DownloadChain
|
||||
from app.helper.directory import validate_download_save_path
|
||||
from app.log import logger
|
||||
|
||||
|
||||
@@ -150,6 +151,18 @@ class UpdateDownloadTasksTool(MoviePilotTool):
|
||||
],
|
||||
}
|
||||
|
||||
if save_path is not None:
|
||||
try:
|
||||
save_path = validate_download_save_path(save_path)
|
||||
except ValueError:
|
||||
return {
|
||||
"hash": hash_value,
|
||||
"downloader": resolved_downloader,
|
||||
"results": [
|
||||
cls._build_result("save_path", False, "保存目录不在允许的下载目录范围内")
|
||||
],
|
||||
}
|
||||
|
||||
results = []
|
||||
if tags:
|
||||
tag_result = download_chain.set_torrents_tag(
|
||||
|
||||
@@ -10,6 +10,7 @@ from app.agent.tools.tags import ToolTag
|
||||
from app.core.event import eventmanager
|
||||
from app.db.subscribe_oper import SubscribeOper
|
||||
from app.log import logger
|
||||
from app.schemas.event import SubscribeModifiedEventData
|
||||
from app.schemas.types import EventType
|
||||
|
||||
|
||||
@@ -261,13 +262,14 @@ class UpdateSubscribeTool(MoviePilotTool):
|
||||
# 发送订阅调整事件
|
||||
await eventmanager.async_send_event(
|
||||
EventType.SubscribeModified,
|
||||
{
|
||||
"subscribe_id": subscribe_id,
|
||||
"old_subscribe_info": old_subscribe_dict,
|
||||
"subscribe_info": updated_subscribe.to_dict()
|
||||
SubscribeModifiedEventData(
|
||||
subscribe_id=subscribe_id,
|
||||
old_subscribe_info=old_subscribe_dict,
|
||||
subscribe_info=updated_subscribe.to_dict()
|
||||
if updated_subscribe
|
||||
else {},
|
||||
},
|
||||
scene="agent_update",
|
||||
).to_dict(),
|
||||
)
|
||||
|
||||
# 构建返回结果
|
||||
|
||||
@@ -7,6 +7,11 @@ 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
|
||||
|
||||
@@ -16,17 +21,36 @@ class WriteFileInput(BaseModel):
|
||||
|
||||
file_path: str = Field(..., description="The absolute path of the file to write")
|
||||
content: str = Field(..., description="The content to write into the file")
|
||||
overwrite: bool = Field(
|
||||
False,
|
||||
description=(
|
||||
"Allow replacing an existing file in full. Keep false when creating a "
|
||||
"new file; prefer edit_file for localized changes."
|
||||
),
|
||||
)
|
||||
expected_sha256: Optional[str] = Field(
|
||||
None,
|
||||
pattern=r"^[0-9a-fA-F]{64}$",
|
||||
description=(
|
||||
"Optional SHA-256 returned by read_file(include_metadata=true). When "
|
||||
"overwriting, fail if the existing file no longer has this hash."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class WriteFileTool(MoviePilotTool):
|
||||
"""创建本地文本文件,或在显式允许后完整覆盖已有文件。"""
|
||||
|
||||
name: str = "write_file"
|
||||
tags: list[str] = [
|
||||
ToolTag.Write,
|
||||
ToolTag.File,
|
||||
]
|
||||
description: str = (
|
||||
"Write full content to a local text file. Non-admin users can only write "
|
||||
"inside the MoviePilot Agent config and log directories."
|
||||
"Create a local text file with complete content. Existing files are "
|
||||
"protected unless overwrite=true; localized changes should use edit_file. "
|
||||
"Supports an optional SHA-256 conflict check and writes atomically. "
|
||||
"Non-admin users can only write inside the MoviePilot Agent config directory."
|
||||
)
|
||||
args_schema: Type[BaseModel] = WriteFileInput
|
||||
|
||||
@@ -36,7 +60,15 @@ class WriteFileTool(MoviePilotTool):
|
||||
file_name = Path(file_path).name if file_path else "未知文件"
|
||||
return f"写入文件: {file_name}"
|
||||
|
||||
async def run(self, file_path: str, content: str, **kwargs) -> str:
|
||||
async def run(
|
||||
self,
|
||||
file_path: str,
|
||||
content: str,
|
||||
overwrite: bool = False,
|
||||
expected_sha256: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
"""创建或显式覆盖文件,并通过可选哈希阻止陈旧写入。"""
|
||||
logger.info(f"执行工具: {self.name}, 参数: file_path={file_path}")
|
||||
|
||||
try:
|
||||
@@ -48,18 +80,52 @@ class WriteFileTool(MoviePilotTool):
|
||||
|
||||
path = AsyncPath(resolved_path)
|
||||
|
||||
if await path.exists() and not await path.is_file():
|
||||
exists = await path.exists()
|
||||
if exists and not await path.is_file():
|
||||
return f"错误:{resolved_path} 路径已存在但不是一个文件"
|
||||
if exists and not overwrite:
|
||||
return (
|
||||
f"错误:文件 {resolved_path} 已存在,拒绝完整覆盖。"
|
||||
"局部修改请使用 edit_file;确需重写时设置 overwrite=true。"
|
||||
)
|
||||
if expected_sha256 and not exists:
|
||||
return (
|
||||
f"错误:文件 {resolved_path} 不存在,无法校验 expected_sha256。"
|
||||
"请确认路径和最新文件状态。"
|
||||
)
|
||||
|
||||
# 自动创建父目录
|
||||
await path.parent.mkdir(parents=True, exist_ok=True)
|
||||
local_path = Path(resolved_path)
|
||||
current_sha256 = None
|
||||
if exists:
|
||||
current_sha256 = await self.run_blocking(
|
||||
"default", calculate_file_sha256, local_path
|
||||
)
|
||||
if expected_sha256:
|
||||
if current_sha256.casefold() != expected_sha256.casefold():
|
||||
return (
|
||||
f"错误:文件 {resolved_path} 已在读取后发生变化,拒绝覆盖。"
|
||||
"请重新读取文件并基于最新内容写入。"
|
||||
)
|
||||
|
||||
# 写入文件
|
||||
await path.write_text(content, encoding="utf-8")
|
||||
await self.run_blocking(
|
||||
"default",
|
||||
atomic_write_text,
|
||||
local_path,
|
||||
content,
|
||||
current_sha256,
|
||||
)
|
||||
new_sha256 = await self.run_blocking(
|
||||
"default", calculate_file_sha256, local_path
|
||||
)
|
||||
|
||||
logger.info(f"成功写入文件 {resolved_path}")
|
||||
return f"成功写入文件 {resolved_path}"
|
||||
return f"成功写入文件 {resolved_path}(sha256={new_sha256})"
|
||||
|
||||
except FileVersionConflictError:
|
||||
return (
|
||||
f"错误:文件 {file_path} 在写入期间发生变化,拒绝覆盖。"
|
||||
"请重新读取文件并再次写入。"
|
||||
)
|
||||
except PermissionError:
|
||||
return f"错误:没有权限写入 {file_path}"
|
||||
except Exception as e:
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
import json
|
||||
import threading
|
||||
import uuid
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from app.agent.tools.base import ToolExecutionTimeoutError, format_tool_result_for_agent
|
||||
from app.agent.tools.factory import MoviePilotToolFactory
|
||||
from app.core.plugin import PluginManager
|
||||
from app.log import logger
|
||||
|
||||
|
||||
@@ -40,27 +42,59 @@ class MoviePilotToolsManager:
|
||||
self.session_id = session_id
|
||||
self.is_admin = is_admin
|
||||
self.tools: List[Any] = []
|
||||
self._tools_lock = threading.Lock()
|
||||
self._plugin_agent_tools_revision = -1
|
||||
self._load_tools()
|
||||
|
||||
def _load_tools(self):
|
||||
def _load_tools(self) -> None:
|
||||
"""
|
||||
加载所有MoviePilot工具
|
||||
"""
|
||||
try:
|
||||
# 创建工具实例
|
||||
self.tools = MoviePilotToolFactory.create_tools(
|
||||
session_id=self.session_id,
|
||||
user_id=self.user_id,
|
||||
channel=None,
|
||||
source="api",
|
||||
username="API Client",
|
||||
stream_handler=None,
|
||||
agent_context={"is_admin": self.is_admin},
|
||||
)
|
||||
plugin_manager = PluginManager()
|
||||
while True:
|
||||
plugin_tools_revision = (
|
||||
plugin_manager.get_plugin_agent_tools_revision()
|
||||
)
|
||||
tools = MoviePilotToolFactory.create_tools(
|
||||
session_id=self.session_id,
|
||||
user_id=self.user_id,
|
||||
channel=None,
|
||||
source="api",
|
||||
username="API Client",
|
||||
stream_handler=None,
|
||||
agent_context={"is_admin": self.is_admin},
|
||||
)
|
||||
if (
|
||||
plugin_tools_revision
|
||||
== plugin_manager.get_plugin_agent_tools_revision()
|
||||
):
|
||||
break
|
||||
self.tools = tools
|
||||
self._plugin_agent_tools_revision = plugin_tools_revision
|
||||
logger.info(f"成功加载 {len(self.tools)} 个工具")
|
||||
except Exception as e:
|
||||
logger.error(f"加载工具失败: {e}", exc_info=True)
|
||||
self.tools = []
|
||||
self._plugin_agent_tools_revision = -1
|
||||
|
||||
def _ensure_tools_current(self) -> None:
|
||||
"""
|
||||
在插件工具注册表变化后惰性刷新工具实例。
|
||||
"""
|
||||
plugin_manager = PluginManager()
|
||||
if (
|
||||
self._plugin_agent_tools_revision
|
||||
== plugin_manager.get_plugin_agent_tools_revision()
|
||||
):
|
||||
return
|
||||
with self._tools_lock:
|
||||
if (
|
||||
self._plugin_agent_tools_revision
|
||||
== plugin_manager.get_plugin_agent_tools_revision()
|
||||
):
|
||||
return
|
||||
self._load_tools()
|
||||
|
||||
def list_tools(self) -> List[ToolDefinition]:
|
||||
"""
|
||||
@@ -69,6 +103,7 @@ class MoviePilotToolsManager:
|
||||
Returns:
|
||||
工具定义列表
|
||||
"""
|
||||
self._ensure_tools_current()
|
||||
tools_list = []
|
||||
for tool in self.tools:
|
||||
if getattr(tool, "_require_admin", False) and not self.is_admin:
|
||||
@@ -102,6 +137,7 @@ class MoviePilotToolsManager:
|
||||
Returns:
|
||||
工具实例,如果未找到返回None
|
||||
"""
|
||||
self._ensure_tools_current()
|
||||
for tool in self.tools:
|
||||
if tool.name == tool_name:
|
||||
return tool
|
||||
|
||||
@@ -25,6 +25,7 @@ class ToolTag(str, Enum):
|
||||
Plugin = "plugin"
|
||||
Workflow = "workflow"
|
||||
Scheduler = "scheduler"
|
||||
AgentTask = "agent_task"
|
||||
File = "file"
|
||||
Directory = "directory"
|
||||
Web = "web"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api.endpoints import auth, login, user, webhook, message, agent, site, subscribe, \
|
||||
from app.api.endpoints import anilist, auth, login, user, webhook, message, agent, site, subscribe, \
|
||||
media, douban, search, plugin, tmdb, history, system, download, dashboard, \
|
||||
transfer, mediaserver, bangumi, storage, discover, recommend, workflow, torrent, mcp, mfa, openai, anthropic, llm, notification
|
||||
|
||||
@@ -29,6 +29,7 @@ api_router.include_router(storage.router, prefix="/storage", tags=["storage"])
|
||||
api_router.include_router(transfer.router, prefix="/transfer", tags=["transfer"])
|
||||
api_router.include_router(mediaserver.router, prefix="/mediaserver", tags=["mediaserver"])
|
||||
api_router.include_router(bangumi.router, prefix="/bangumi", tags=["bangumi"])
|
||||
api_router.include_router(anilist.router, prefix="/anilist", tags=["anilist"])
|
||||
api_router.include_router(discover.router, prefix="/discover", tags=["discover"])
|
||||
api_router.include_router(recommend.router, prefix="/recommend", tags=["recommend"])
|
||||
api_router.include_router(workflow.router, prefix="/workflow", tags=["workflow"])
|
||||
|
||||
7
app/api/apiv2.py
Normal file
7
app/api/apiv2.py
Normal file
@@ -0,0 +1,7 @@
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api.apiv1 import api_router
|
||||
|
||||
|
||||
api_router_v2 = APIRouter()
|
||||
api_router_v2.include_router(api_router)
|
||||
224
app/api/apiv2_utils.py
Normal file
224
app/api/apiv2_utils.py
Normal file
@@ -0,0 +1,224 @@
|
||||
import json
|
||||
from typing import Any, Awaitable, Callable
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
from fastapi.routing import APIRoute
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
from starlette.responses import Response as StarletteResponse
|
||||
|
||||
from app.schemas.response import Response
|
||||
|
||||
|
||||
API_V2_STR = "/api/v2"
|
||||
OPENAPI_V2_PATH = f"{API_V2_STR}/openapi.json"
|
||||
_PROTOCOL_PREFIXES = ("/openai", "/anthropic", "/mcp")
|
||||
_JSON_CONTENT_TYPES = ("application/json", "+json")
|
||||
|
||||
|
||||
def _is_protocol_path(path: str) -> bool:
|
||||
"""判断路径是否属于需要保留原始协议响应的接口。"""
|
||||
relative_path = path.removeprefix(API_V2_STR)
|
||||
return any(
|
||||
relative_path == prefix or relative_path.startswith(f"{prefix}/")
|
||||
for prefix in _PROTOCOL_PREFIXES
|
||||
)
|
||||
|
||||
|
||||
def _is_json_response(response: StarletteResponse) -> bool:
|
||||
"""判断响应是否为可安全解析的 JSON 响应。"""
|
||||
content_type = response.headers.get("content-type", "").split(";", 1)[0]
|
||||
return any(
|
||||
content_type == accepted_type or content_type.endswith(accepted_type)
|
||||
for accepted_type in _JSON_CONTENT_TYPES
|
||||
)
|
||||
|
||||
|
||||
def _is_response_payload(payload: Any) -> bool:
|
||||
"""判断响应内容是否已经符合通用 Response 结构。"""
|
||||
return isinstance(payload, dict) and {
|
||||
"success",
|
||||
"message",
|
||||
"data",
|
||||
}.issubset(payload)
|
||||
|
||||
|
||||
def _get_error_message(payload: Any) -> str:
|
||||
"""从旧版错误响应中提取统一的错误消息。"""
|
||||
if isinstance(payload, dict):
|
||||
detail = payload.get("detail")
|
||||
if isinstance(detail, str) and detail:
|
||||
return detail
|
||||
if isinstance(detail, list):
|
||||
messages = [
|
||||
item.get("msg")
|
||||
for item in detail
|
||||
if isinstance(item, dict) and isinstance(item.get("msg"), str)
|
||||
]
|
||||
if messages:
|
||||
return "; ".join(messages)
|
||||
if detail is not None:
|
||||
return json.dumps(detail, ensure_ascii=False)
|
||||
message = payload.get("message")
|
||||
if isinstance(message, str) and message:
|
||||
return message
|
||||
if isinstance(payload, str) and payload:
|
||||
return payload
|
||||
return "请求失败"
|
||||
|
||||
|
||||
def _copy_response_headers(source: StarletteResponse, target: StarletteResponse) -> None:
|
||||
"""复制适配前响应中仍然有效的头信息。"""
|
||||
for key, value in source.raw_headers:
|
||||
if key.lower() not in {b"content-length", b"content-type"}:
|
||||
target.raw_headers.append((key, value))
|
||||
|
||||
|
||||
def _restore_response_body(
|
||||
source: StarletteResponse,
|
||||
body: bytes,
|
||||
) -> StarletteResponse:
|
||||
"""在检查响应体后恢复原始响应内容和头信息。"""
|
||||
restored_response = StarletteResponse(
|
||||
content=body,
|
||||
status_code=source.status_code,
|
||||
background=source.background,
|
||||
)
|
||||
restored_response.raw_headers = list(source.raw_headers)
|
||||
return restored_response
|
||||
|
||||
|
||||
class V2ResponseMiddleware(BaseHTTPMiddleware):
|
||||
"""
|
||||
为 v2 REST 接口适配统一的 Response 响应结构。
|
||||
|
||||
已经返回项目 Response 模型的成功响应保持原样,避免改变既有接口语义;
|
||||
OpenAI、Anthropic 和 MCP 协议接口也保持原始协议响应。
|
||||
"""
|
||||
|
||||
async def dispatch(
|
||||
self,
|
||||
request: Request,
|
||||
call_next: Callable[[Request], Awaitable[StarletteResponse]],
|
||||
) -> StarletteResponse:
|
||||
"""处理 v2 请求并在必要时封装 JSON 响应。"""
|
||||
response = await call_next(request)
|
||||
if not request.url.path.startswith(f"{API_V2_STR}/"):
|
||||
return response
|
||||
if request.url.path == OPENAPI_V2_PATH:
|
||||
return response
|
||||
if _is_protocol_path(request.url.path):
|
||||
return response
|
||||
if response.status_code in {204, 304} or not _is_json_response(response):
|
||||
return response
|
||||
if response.headers.get("content-encoding"):
|
||||
return response
|
||||
|
||||
route = request.scope.get("route")
|
||||
route_response_model = getattr(route, "response_model", None)
|
||||
if response.status_code < 400 and route_response_model is Response:
|
||||
return response
|
||||
|
||||
body = b"".join([chunk async for chunk in response.body_iterator])
|
||||
if not body:
|
||||
return _restore_response_body(response, body)
|
||||
try:
|
||||
payload = json.loads(body)
|
||||
except (TypeError, ValueError):
|
||||
return _restore_response_body(response, body)
|
||||
|
||||
if _is_response_payload(payload):
|
||||
return _restore_response_body(response, body)
|
||||
|
||||
if response.status_code >= 400:
|
||||
content = {
|
||||
"success": False,
|
||||
"message": _get_error_message(payload),
|
||||
"data": {},
|
||||
}
|
||||
if isinstance(payload, dict) and isinstance(payload.get("detail_i18n"), str):
|
||||
content["message_i18n"] = payload["detail_i18n"]
|
||||
else:
|
||||
content = {
|
||||
"success": True,
|
||||
"message": "",
|
||||
"data": payload,
|
||||
}
|
||||
|
||||
wrapped_response = JSONResponse(
|
||||
content=content,
|
||||
status_code=response.status_code,
|
||||
background=response.background,
|
||||
)
|
||||
_copy_response_headers(response, wrapped_response)
|
||||
return wrapped_response
|
||||
|
||||
|
||||
def configure_v2_openapi(app: FastAPI) -> None:
|
||||
"""
|
||||
将 v2 普通 JSON 接口的 OpenAPI 响应模型改为通用 Response。
|
||||
|
||||
:param app: 已完成 v1/v2 路由注册的 FastAPI 应用
|
||||
"""
|
||||
if getattr(app, "_v2_openapi_configured", False):
|
||||
return
|
||||
|
||||
original_openapi = app.openapi
|
||||
|
||||
def custom_openapi() -> dict[str, Any]:
|
||||
"""生成包含 v2 通用响应模型的 OpenAPI 文档。"""
|
||||
schema = original_openapi()
|
||||
components = schema.setdefault("components", {}).setdefault("schemas", {})
|
||||
components["Response"] = Response.model_json_schema(
|
||||
ref_template="#/components/schemas/{model}"
|
||||
)
|
||||
|
||||
route_map = {
|
||||
(route.path, method.lower()): route
|
||||
for route in app.routes
|
||||
if isinstance(route, APIRoute)
|
||||
for method in route.methods
|
||||
}
|
||||
response_ref = {"$ref": "#/components/schemas/Response"}
|
||||
for path, path_item in schema.get("paths", {}).items():
|
||||
if not path.startswith(f"{API_V2_STR}/"):
|
||||
continue
|
||||
for method, operation in path_item.items():
|
||||
if method not in {
|
||||
"get",
|
||||
"post",
|
||||
"put",
|
||||
"patch",
|
||||
"delete",
|
||||
"options",
|
||||
"head",
|
||||
}:
|
||||
continue
|
||||
route = route_map.get((path, method))
|
||||
if (
|
||||
route is None
|
||||
or route.response_model is None
|
||||
or route.response_model is Any
|
||||
or route.response_model is Response
|
||||
or _is_protocol_path(path)
|
||||
):
|
||||
continue
|
||||
if route.status_code in {204, 304}:
|
||||
continue
|
||||
content_type = getattr(route.response_class, "media_type", None)
|
||||
if content_type and not (
|
||||
content_type == "application/json" or content_type.endswith("+json")
|
||||
):
|
||||
continue
|
||||
status_code = str(route.status_code or 200)
|
||||
response = operation.get("responses", {}).get(status_code)
|
||||
if response and "content" in response:
|
||||
json_content = response["content"].get("application/json")
|
||||
if json_content is not None:
|
||||
json_content["schema"] = response_ref
|
||||
|
||||
app.openapi_schema = schema
|
||||
return schema
|
||||
|
||||
app.openapi = custom_openapi
|
||||
app._v2_openapi_configured = True
|
||||
@@ -7,6 +7,7 @@ import shutil
|
||||
import subprocess
|
||||
import time
|
||||
import uuid
|
||||
from collections import deque
|
||||
from queue import Empty, Queue
|
||||
from pathlib import Path
|
||||
from threading import Lock
|
||||
@@ -20,6 +21,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from app import schemas
|
||||
from app.agent import MoviePilotAgent, ReplyMode, StreamingHandler, agent_manager
|
||||
from app.agent.llm.capability import AgentCapabilityManager
|
||||
from app.agent.mcp import agent_mcp_manager
|
||||
from app.chain.message import MessageChain
|
||||
from app.chain.site import site_interaction_manager
|
||||
from app.chain.skills import skills_interaction_manager
|
||||
@@ -34,6 +36,7 @@ from app.db.models.agentchat import AgentChat
|
||||
from app.db.user_oper import UserOper, get_current_active_user
|
||||
from app.helper.agent import attach_web_agent_edit_queue, detach_web_agent_edit_queue
|
||||
from app.helper.interaction import agent_interaction_manager, media_interaction_manager
|
||||
from app.helper.locale import LocaleHelper
|
||||
from app.log import logger
|
||||
from app.schemas.types import EventType, MessageChannel
|
||||
|
||||
@@ -48,6 +51,10 @@ WEB_AGENT_UPLOAD_CHUNK_SIZE = 1024 * 1024
|
||||
WEB_AGENT_BROWSER_AUDIO_SUFFIXES = {".aac", ".m4a", ".mp3", ".mp4", ".wav", ".wave"}
|
||||
WEB_AGENT_TRADITIONAL_IDLE_TIMEOUT_SECONDS = 2.0
|
||||
WEB_AGENT_TRADITIONAL_MAX_WAIT_SECONDS = 60.0
|
||||
WEB_AGENT_STREAM_COALESCE_SECONDS = 0.03
|
||||
WEB_AGENT_STREAM_COALESCE_MAX_CHARS = 256
|
||||
WEB_AGENT_STREAM_HEARTBEAT_SECONDS = 15.0
|
||||
WEB_AGENT_STREAM_QUEUE_MAX_SIZE = 64
|
||||
_WEB_AGENT_FILE_REGISTRY: dict[str, dict[str, Any]] = {}
|
||||
_WEB_AGENT_NOTICE_QUEUES: dict[str, list[Queue[schemas.Notification]]] = {}
|
||||
_WEB_AGENT_NOTICE_LOCK = Lock()
|
||||
@@ -55,6 +62,179 @@ _WEB_AGENT_NOTICE_LISTENER_REGISTERED = False
|
||||
_WEB_AGENT_BACKGROUND_TASKS: set[asyncio.Task] = set()
|
||||
|
||||
|
||||
class _WebAgentEventPublisher:
|
||||
"""合并 WebAgent 文本增量,并通过有界队列向 SSE 消费者提供事件。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._queue: asyncio.Queue[dict] = asyncio.Queue(
|
||||
maxsize=WEB_AGENT_STREAM_QUEUE_MAX_SIZE
|
||||
)
|
||||
self._pending_events: deque[dict] = deque()
|
||||
self._pending_signal = asyncio.Event()
|
||||
self._pending_delta = ""
|
||||
self._delta_timer: Optional[asyncio.TimerHandle] = None
|
||||
self._disposed = False
|
||||
self._max_depth = 0
|
||||
self._last_logged_depth = 0
|
||||
self._pump_task = asyncio.create_task(self._pump())
|
||||
|
||||
@property
|
||||
def max_depth(self) -> int:
|
||||
"""返回本轮发布器观测到的最大积压深度。"""
|
||||
return self._max_depth
|
||||
|
||||
def publish(self, event: dict) -> None:
|
||||
"""发布事件;相邻文本会按时间或长度边界合并。"""
|
||||
if self._disposed:
|
||||
return
|
||||
if event.get("type") == "delta":
|
||||
self._pending_delta += str(event.get("content") or "")
|
||||
if len(self._pending_delta) >= WEB_AGENT_STREAM_COALESCE_MAX_CHARS:
|
||||
self._flush_delta()
|
||||
elif self._delta_timer is None:
|
||||
loop = asyncio.get_running_loop()
|
||||
self._delta_timer = loop.call_later(
|
||||
WEB_AGENT_STREAM_COALESCE_SECONDS,
|
||||
self._flush_delta,
|
||||
)
|
||||
return
|
||||
|
||||
self._flush_delta()
|
||||
self._append_event(event)
|
||||
|
||||
async def get(self) -> dict:
|
||||
"""等待并返回下一条已排序事件。"""
|
||||
return await self._queue.get()
|
||||
|
||||
async def aclose(self) -> None:
|
||||
"""停止发布器并释放等待中的泵任务。"""
|
||||
if self._disposed:
|
||||
return
|
||||
self._disposed = True
|
||||
self._cancel_delta_timer()
|
||||
self._pending_delta = ""
|
||||
self._pending_events.clear()
|
||||
self._pump_task.cancel()
|
||||
try:
|
||||
await self._pump_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
def _cancel_delta_timer(self) -> None:
|
||||
"""取消尚未触发的文本合并计时器。"""
|
||||
if self._delta_timer is None:
|
||||
return
|
||||
self._delta_timer.cancel()
|
||||
self._delta_timer = None
|
||||
|
||||
def _flush_delta(self) -> None:
|
||||
"""把当前文本缓冲转换成一条增量事件。"""
|
||||
self._cancel_delta_timer()
|
||||
if not self._pending_delta or self._disposed:
|
||||
return
|
||||
content = self._pending_delta
|
||||
self._pending_delta = ""
|
||||
self._append_event({"type": "delta", "content": content})
|
||||
|
||||
def _append_event(self, event: dict) -> None:
|
||||
"""追加待发布事件,相邻文本在出口阻塞时继续合并。"""
|
||||
if (
|
||||
event.get("type") == "delta"
|
||||
and self._pending_events
|
||||
and self._pending_events[-1].get("type") == "delta"
|
||||
):
|
||||
self._pending_events[-1]["content"] += str(event.get("content") or "")
|
||||
else:
|
||||
self._pending_events.append(event)
|
||||
self._pending_signal.set()
|
||||
depth = self._queue.qsize() + len(self._pending_events)
|
||||
self._max_depth = max(self._max_depth, depth)
|
||||
if depth >= WEB_AGENT_STREAM_QUEUE_MAX_SIZE // 2 and depth > self._last_logged_depth:
|
||||
self._last_logged_depth = depth
|
||||
logger.debug(f"WebAgent SSE事件积压深度: {depth}")
|
||||
|
||||
async def _pump(self) -> None:
|
||||
"""按发布顺序把本地合并结果写入有界出口队列。"""
|
||||
while True:
|
||||
await self._pending_signal.wait()
|
||||
while self._pending_events:
|
||||
event = self._pending_events.popleft()
|
||||
await self._queue.put(event)
|
||||
self._pending_signal.clear()
|
||||
|
||||
|
||||
def _ensure_superuser(user: User) -> None:
|
||||
"""校验当前用户是否为超级管理员。"""
|
||||
if not getattr(user, "is_superuser", False):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Forbidden")
|
||||
|
||||
|
||||
@router.get("/mcp/servers", summary="查询 Agent MCP 服务器配置", response_model=schemas.Response)
|
||||
async def list_agent_mcp_servers(
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
) -> schemas.Response:
|
||||
"""
|
||||
查询 Agent 外部 MCP 服务器配置。
|
||||
"""
|
||||
_ensure_superuser(current_user)
|
||||
servers = agent_mcp_manager.get_servers()
|
||||
enabled_count = len([server for server in servers if server.enabled])
|
||||
return schemas.Response(
|
||||
success=True,
|
||||
data={
|
||||
"servers": [server.model_dump() for server in servers],
|
||||
"enabled_count": enabled_count,
|
||||
"total_count": len(servers),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@router.post("/mcp/servers", summary="保存 Agent MCP 服务器配置", response_model=schemas.Response)
|
||||
async def save_agent_mcp_servers(
|
||||
request: schemas.AgentMcpServersSaveRequest,
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
) -> schemas.Response:
|
||||
"""
|
||||
保存 Agent 外部 MCP 服务器配置。
|
||||
"""
|
||||
_ensure_superuser(current_user)
|
||||
success = await agent_mcp_manager.save_servers(request.servers)
|
||||
return schemas.Response(
|
||||
success=success,
|
||||
message="保存MCP配置成功" if success else "保存MCP配置失败",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/mcp/servers/test", summary="测试 Agent MCP 服务器", response_model=schemas.Response)
|
||||
async def test_agent_mcp_server(
|
||||
request: schemas.AgentMcpServerTestRequest,
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
) -> schemas.Response:
|
||||
"""
|
||||
测试 Agent 外部 MCP 服务器连接并读取工具列表。
|
||||
"""
|
||||
_ensure_superuser(current_user)
|
||||
try:
|
||||
result = await agent_mcp_manager.test_server(request.server)
|
||||
return schemas.Response(
|
||||
success=result.success,
|
||||
message=result.message,
|
||||
data=result.model_dump(),
|
||||
)
|
||||
except Exception as err:
|
||||
logger.warning(f"测试 Agent MCP 服务器失败: {err}")
|
||||
return schemas.Response(
|
||||
success=False,
|
||||
message=f"测试MCP服务器失败: {str(err)}",
|
||||
data={
|
||||
"success": False,
|
||||
"message": str(err),
|
||||
"tools": [],
|
||||
"tool_count": 0,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class _WebAgentStreamingHandler(StreamingHandler):
|
||||
"""
|
||||
Web 前端专用流式处理器,将工具提示和文本统一回调给 SSE。
|
||||
@@ -72,6 +252,20 @@ class _WebAgentStreamingHandler(StreamingHandler):
|
||||
"""
|
||||
self._on_emit = on_emit
|
||||
|
||||
def record_tool_call(
|
||||
self,
|
||||
tool_name: str,
|
||||
tool_message: Optional[str] = None,
|
||||
tool_kwargs: Optional[dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""记录并立即输出 Web 工具事件,避免汇总延迟到正文结束后。"""
|
||||
super().record_tool_call(
|
||||
tool_name=tool_name,
|
||||
tool_message=tool_message,
|
||||
tool_kwargs=tool_kwargs,
|
||||
)
|
||||
self.flush_pending_tool_summary()
|
||||
|
||||
def emit(self, token: str) -> str:
|
||||
"""追加 token 并同步通知 SSE 生产者。"""
|
||||
emitted = super().emit(token)
|
||||
@@ -145,7 +339,9 @@ class _WebAgentMoviePilotAgent(MoviePilotAgent):
|
||||
self.stream_handler = _WebAgentStreamingHandler(self._emit_output)
|
||||
|
||||
def _should_stream(self) -> bool:
|
||||
"""Web 面板需要实时输出,即使 Web 渠道本身不支持消息编辑。"""
|
||||
"""Web 对话实时输出,复用会话执行后台任务时改用非流式广播。"""
|
||||
if self.is_background:
|
||||
return False
|
||||
return True
|
||||
|
||||
def set_notification_callback(
|
||||
@@ -192,6 +388,18 @@ class _WebAgentMoviePilotAgent(MoviePilotAgent):
|
||||
"""文本输出交由 Web 流式处理器统一回调,避免重复增量。"""
|
||||
self.stream_handler.emit(text)
|
||||
|
||||
def _emit_output(self, text: str) -> None:
|
||||
"""保留完整输出状态,同时只把本次增量交给 Web SSE 回调。"""
|
||||
if not text:
|
||||
return
|
||||
self._streamed_output += text
|
||||
if not callable(self.output_callback):
|
||||
return
|
||||
try:
|
||||
self.output_callback(text)
|
||||
except Exception as e:
|
||||
logger.debug(f"Web智能体输出回调失败: {e}")
|
||||
|
||||
|
||||
def _build_web_agent_session_id(user: User, session_id: Optional[str]) -> str:
|
||||
"""
|
||||
@@ -242,16 +450,53 @@ async def _get_accessible_agent_chat(
|
||||
return chat
|
||||
|
||||
|
||||
def _append_web_agent_text_segment(assistant_message: dict, content: str) -> None:
|
||||
"""
|
||||
将文本增量追加到展示消息,并仅合并相邻文本片段。
|
||||
|
||||
:param assistant_message: 当前助手展示消息
|
||||
:param content: 新增文本
|
||||
"""
|
||||
if not content:
|
||||
return
|
||||
assistant_message["content"] = str(assistant_message.get("content") or "") + content
|
||||
segments = assistant_message.setdefault("segments", [])
|
||||
if segments and segments[-1].get("type") == "text":
|
||||
segments[-1]["content"] = str(segments[-1].get("content") or "") + content
|
||||
else:
|
||||
segments.append({"type": "text", "content": content})
|
||||
|
||||
|
||||
def _build_legacy_web_agent_segments(content: str, tools: list[dict]) -> list[dict]:
|
||||
"""
|
||||
为未携带有序片段的旧展示消息生成兼容布局。
|
||||
|
||||
:param content: 聚合后的助手文本
|
||||
:param tools: 工具提示列表
|
||||
:return: 按旧版工具在前、文本在后的顺序生成的片段
|
||||
"""
|
||||
segments = [
|
||||
{"type": "tool", "toolIndex": index}
|
||||
for index in range(len(tools))
|
||||
]
|
||||
if content:
|
||||
segments.append({"type": "text", "content": content})
|
||||
return segments
|
||||
|
||||
|
||||
def _apply_web_agent_display_event(event: dict, assistant_message: dict) -> None:
|
||||
"""
|
||||
将 WebAgent SSE 事件同步应用到服务端展示消息快照。
|
||||
"""
|
||||
event_type = event.get("type")
|
||||
if event_type == "delta":
|
||||
assistant_message["content"] += event.get("content") or ""
|
||||
_append_web_agent_text_segment(
|
||||
assistant_message, event.get("content") or ""
|
||||
)
|
||||
elif event_type == "tool":
|
||||
for tool in assistant_message["tools"]:
|
||||
tool["status"] = "done"
|
||||
tool_index = len(assistant_message["tools"])
|
||||
assistant_message["tools"].append(
|
||||
{
|
||||
"id": f"tool-{uuid.uuid4().hex}",
|
||||
@@ -259,6 +504,9 @@ def _apply_web_agent_display_event(event: dict, assistant_message: dict) -> None
|
||||
"status": "running",
|
||||
}
|
||||
)
|
||||
assistant_message.setdefault("segments", []).append(
|
||||
{"type": "tool", "toolIndex": tool_index}
|
||||
)
|
||||
elif event_type == "attachment" and event.get("attachment"):
|
||||
assistant_message["attachments"].append(event["attachment"])
|
||||
elif event_type == "choice" and event.get("choice"):
|
||||
@@ -270,14 +518,22 @@ def _apply_web_agent_display_event(event: dict, assistant_message: dict) -> None
|
||||
assistant_message["attachments"] = target_message.get("attachments") or []
|
||||
assistant_message["choices"] = target_message.get("choices") or []
|
||||
assistant_message["tools"] = target_message.get("tools") or []
|
||||
target_segments = target_message.get("segments")
|
||||
assistant_message["segments"] = (
|
||||
target_segments
|
||||
if isinstance(target_segments, list)
|
||||
else _build_legacy_web_agent_segments(
|
||||
assistant_message["content"], assistant_message["tools"]
|
||||
)
|
||||
)
|
||||
assistant_message["status"] = target_message.get("status") or "done"
|
||||
elif event_type == "error":
|
||||
assistant_message["status"] = "error"
|
||||
assistant_message["content"] = (
|
||||
assistant_message["content"]
|
||||
or event.get("message")
|
||||
or "智能助手响应失败"
|
||||
)
|
||||
if not assistant_message["content"]:
|
||||
_append_web_agent_text_segment(
|
||||
assistant_message,
|
||||
event.get("message") or "智能助手响应失败",
|
||||
)
|
||||
for tool in assistant_message["tools"]:
|
||||
tool["status"] = "done"
|
||||
elif event_type == "done":
|
||||
@@ -326,15 +582,25 @@ def _save_web_agent_display_snapshot(
|
||||
logger.debug(f"保存WebAgent展示历史失败: {e}")
|
||||
|
||||
|
||||
def _build_web_agent_sse(event_type: str, data: Optional[dict] = None) -> str:
|
||||
def _build_web_agent_sse(
|
||||
event_type: str,
|
||||
data: Optional[dict] = None,
|
||||
locale: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
构建 Web Agent SSE 消息。
|
||||
|
||||
:param event_type: 前端事件类型
|
||||
:param data: 事件数据
|
||||
:param locale: 当前请求语言
|
||||
:return: 符合 SSE 格式的字符串
|
||||
"""
|
||||
payload = {"type": event_type, **(data or {})}
|
||||
message = payload.get("message")
|
||||
if event_type == "error" and isinstance(message, str):
|
||||
payload["message_i18n"] = LocaleHelper.translate_text(
|
||||
message, locale=locale
|
||||
)
|
||||
return f"data: {json.dumps(payload, ensure_ascii=False)}\n\n"
|
||||
|
||||
|
||||
@@ -1597,6 +1863,7 @@ async def web_agent_stream(
|
||||
:return: SSE 流式响应
|
||||
"""
|
||||
prompt = payload.text.strip()
|
||||
locale = LocaleHelper.get_locale_from_request(request)
|
||||
display_prompt = (payload.display_text or payload.text).strip()
|
||||
is_traditional_message = (
|
||||
_is_web_agent_traditional_message(prompt)
|
||||
@@ -1610,6 +1877,7 @@ async def web_agent_stream(
|
||||
_build_web_agent_sse(
|
||||
"error",
|
||||
{"message": denied_message},
|
||||
locale=locale,
|
||||
)
|
||||
]),
|
||||
media_type="text/event-stream",
|
||||
@@ -1621,6 +1889,7 @@ async def web_agent_stream(
|
||||
_build_web_agent_sse(
|
||||
"error",
|
||||
{"message": unknown_command_message},
|
||||
locale=locale,
|
||||
)
|
||||
]),
|
||||
media_type="text/event-stream",
|
||||
@@ -1649,34 +1918,73 @@ async def web_agent_stream(
|
||||
"""
|
||||
生成传统消息链路的 WebAgent SSE 事件。
|
||||
"""
|
||||
yield _build_web_agent_sse("start", {"session_id": session_id})
|
||||
events = await _collect_web_agent_traditional_events(
|
||||
text=prompt,
|
||||
current_user=current_user,
|
||||
original_message_id=payload.original_message_id,
|
||||
original_chat_id=payload.original_chat_id,
|
||||
yield _build_web_agent_sse(
|
||||
"start",
|
||||
{"session_id": session_id},
|
||||
locale=locale,
|
||||
)
|
||||
collection_task = asyncio.create_task(
|
||||
_collect_web_agent_traditional_events(
|
||||
text=prompt,
|
||||
current_user=current_user,
|
||||
original_message_id=payload.original_message_id,
|
||||
original_chat_id=payload.original_chat_id,
|
||||
)
|
||||
)
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
events = await asyncio.wait_for(
|
||||
asyncio.shield(collection_task),
|
||||
timeout=WEB_AGENT_STREAM_HEARTBEAT_SECONDS,
|
||||
)
|
||||
break
|
||||
except asyncio.TimeoutError:
|
||||
if await request.is_disconnected():
|
||||
collection_task.cancel()
|
||||
return
|
||||
yield ": heartbeat\n\n"
|
||||
except asyncio.CancelledError:
|
||||
if not collection_task.done():
|
||||
collection_task.cancel()
|
||||
return
|
||||
|
||||
assistant_message = _build_web_agent_display_message_from_events(events)
|
||||
display_messages.append(assistant_message)
|
||||
|
||||
async def save_display_snapshot() -> None:
|
||||
"""后台保存传统消息展示快照,不阻塞 SSE 终态。"""
|
||||
try:
|
||||
await run_in_threadpool(
|
||||
_save_web_agent_display_snapshot,
|
||||
session_id=session_id,
|
||||
current_user=current_user,
|
||||
messages=display_messages,
|
||||
client_session_id=payload.session_id or session_id,
|
||||
)
|
||||
except Exception as err:
|
||||
logger.error(f"保存WebAgent传统消息快照失败: {str(err)}")
|
||||
|
||||
snapshot_task = asyncio.create_task(save_display_snapshot())
|
||||
_WEB_AGENT_BACKGROUND_TASKS.add(snapshot_task)
|
||||
snapshot_task.add_done_callback(_WEB_AGENT_BACKGROUND_TASKS.discard)
|
||||
await asyncio.sleep(0)
|
||||
for event in events:
|
||||
event_payload = copy.deepcopy(event)
|
||||
yield _build_web_agent_sse(event_payload.pop("type"), event_payload)
|
||||
yield _build_web_agent_sse(
|
||||
event_payload.pop("type"),
|
||||
event_payload,
|
||||
locale=locale,
|
||||
)
|
||||
if await request.is_disconnected():
|
||||
break
|
||||
await run_in_threadpool(
|
||||
_save_web_agent_display_snapshot,
|
||||
session_id=session_id,
|
||||
current_user=current_user,
|
||||
messages=display_messages,
|
||||
client_session_id=payload.session_id or session_id,
|
||||
)
|
||||
yield _build_web_agent_sse("done", {})
|
||||
return
|
||||
yield _build_web_agent_sse("done", {}, locale=locale)
|
||||
|
||||
return StreamingResponse(
|
||||
traditional_event_generator(),
|
||||
media_type="text/event-stream",
|
||||
headers={
|
||||
"Cache-Control": "no-cache",
|
||||
"Cache-Control": "no-cache, no-transform",
|
||||
"Connection": "keep-alive",
|
||||
"X-Accel-Buffering": "no",
|
||||
},
|
||||
@@ -1688,6 +1996,7 @@ async def web_agent_stream(
|
||||
_build_web_agent_sse(
|
||||
"error",
|
||||
{"message": "智能助手未启用,请先在系统设置中开启。"},
|
||||
locale=locale,
|
||||
)
|
||||
]),
|
||||
media_type="text/event-stream",
|
||||
@@ -1703,6 +2012,7 @@ async def web_agent_stream(
|
||||
_build_web_agent_sse(
|
||||
"error",
|
||||
{"message": "语音识别失败,请稍后重试。"},
|
||||
locale=locale,
|
||||
)
|
||||
]),
|
||||
media_type="text/event-stream",
|
||||
@@ -1713,6 +2023,7 @@ async def web_agent_stream(
|
||||
_build_web_agent_sse(
|
||||
"error",
|
||||
{"message": "请输入要发送给智能助手的内容或选择附件。"},
|
||||
locale=locale,
|
||||
)
|
||||
]),
|
||||
media_type="text/event-stream",
|
||||
@@ -1720,8 +2031,7 @@ async def web_agent_stream(
|
||||
|
||||
session_id = _build_web_agent_session_id(current_user, payload.session_id)
|
||||
MessageChain().bind_user_session(str(current_user.id), session_id)
|
||||
event_queue: asyncio.Queue = asyncio.Queue()
|
||||
last_output = ""
|
||||
event_publisher = _WebAgentEventPublisher()
|
||||
user_attachments = _build_web_agent_input_attachments(
|
||||
images=payload.images or [],
|
||||
files=[
|
||||
@@ -1746,16 +2056,13 @@ async def web_agent_stream(
|
||||
)
|
||||
display_messages.append(assistant_display_message)
|
||||
|
||||
def output_callback(output: str) -> None:
|
||||
def output_callback(delta: str) -> None:
|
||||
"""
|
||||
接收 Agent 累积输出并转成增量事件。
|
||||
接收 Agent 文本增量并转换成 SSE 事件。
|
||||
"""
|
||||
nonlocal last_output
|
||||
delta = output[len(last_output):] if output.startswith(last_output) else output
|
||||
last_output = output
|
||||
for item in _split_web_agent_output(delta):
|
||||
_apply_web_agent_display_event(item, assistant_display_message)
|
||||
event_queue.put_nowait(item)
|
||||
event_publisher.publish(item)
|
||||
|
||||
def notification_callback(notification: schemas.Notification) -> None:
|
||||
"""
|
||||
@@ -1763,7 +2070,7 @@ async def web_agent_stream(
|
||||
"""
|
||||
for item in _build_web_agent_notification_events(notification):
|
||||
_apply_web_agent_display_event(item, assistant_display_message)
|
||||
event_queue.put_nowait(item)
|
||||
event_publisher.publish(item)
|
||||
|
||||
async def event_generator() -> AsyncIterator[str]:
|
||||
"""
|
||||
@@ -1805,10 +2112,12 @@ async def web_agent_stream(
|
||||
"message": f"智能助手执行失败: {str(err)}",
|
||||
}
|
||||
_apply_web_agent_display_event(error_event, assistant_display_message)
|
||||
await event_queue.put(error_event)
|
||||
event_publisher.publish(error_event)
|
||||
finally:
|
||||
done_event = {"type": "done"}
|
||||
_apply_web_agent_display_event(done_event, assistant_display_message)
|
||||
# 终态先进入 SSE 队列,避免展示快照落库延迟前端结束动画。
|
||||
event_publisher.publish(done_event)
|
||||
await run_in_threadpool(
|
||||
_save_web_agent_display_snapshot,
|
||||
session_id=session_id,
|
||||
@@ -1816,35 +2125,51 @@ async def web_agent_stream(
|
||||
messages=display_messages,
|
||||
client_session_id=payload.session_id or session_id,
|
||||
)
|
||||
await event_queue.put(done_event)
|
||||
|
||||
task = asyncio.create_task(run_agent())
|
||||
_WEB_AGENT_BACKGROUND_TASKS.add(task)
|
||||
task.add_done_callback(_WEB_AGENT_BACKGROUND_TASKS.discard)
|
||||
disconnected = False
|
||||
terminal_sent = False
|
||||
try:
|
||||
yield _build_web_agent_sse(
|
||||
"start",
|
||||
{"session_id": session_id},
|
||||
locale=locale,
|
||||
)
|
||||
disconnected = False
|
||||
while not global_vars.is_system_stopped:
|
||||
if await request.is_disconnected():
|
||||
disconnected = True
|
||||
break
|
||||
event = await event_queue.get()
|
||||
yield _build_web_agent_sse(event.pop("type"), event)
|
||||
if task.done() and event_queue.empty():
|
||||
try:
|
||||
event = await asyncio.wait_for(
|
||||
event_publisher.get(),
|
||||
timeout=WEB_AGENT_STREAM_HEARTBEAT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
yield ": heartbeat\n\n"
|
||||
continue
|
||||
event_type = str(event.get("type") or "")
|
||||
if event_type == "done":
|
||||
terminal_sent = True
|
||||
yield _build_web_agent_sse(
|
||||
event_type,
|
||||
{key: value for key, value in event.items() if key != "type"},
|
||||
locale=locale,
|
||||
)
|
||||
if event_type == "done":
|
||||
break
|
||||
except asyncio.CancelledError:
|
||||
disconnected = True
|
||||
return
|
||||
finally:
|
||||
if not task.done() and not disconnected:
|
||||
if not task.done() and not disconnected and not terminal_sent:
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
await event_publisher.aclose()
|
||||
# 客户端退到后台导致 SSE 断开时,保留后台 Agent 继续执行;完成后会保存展示快照,
|
||||
# 前端恢复可见时可通过会话详情接口拉取最终状态。
|
||||
|
||||
@@ -1852,7 +2177,7 @@ async def web_agent_stream(
|
||||
event_generator(),
|
||||
media_type="text/event-stream",
|
||||
headers={
|
||||
"Cache-Control": "no-cache",
|
||||
"Cache-Control": "no-cache, no-transform",
|
||||
"Connection": "keep-alive",
|
||||
"X-Accel-Buffering": "no",
|
||||
},
|
||||
|
||||
169
app/api/endpoints/anilist.py
Normal file
169
app/api/endpoints/anilist.py
Normal file
@@ -0,0 +1,169 @@
|
||||
from typing import Annotated, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query
|
||||
|
||||
from app import schemas
|
||||
from app.chain.anilist import AniListChain
|
||||
from app.core.context import MediaInfo
|
||||
from app.core.security import verify_token
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
PageParam = Annotated[int, Query(ge=1)]
|
||||
CountParam = Annotated[int, Query(ge=1, le=50)]
|
||||
|
||||
|
||||
def _serialize_medias(medias: list[MediaInfo]) -> list[schemas.MediaInfo]:
|
||||
"""
|
||||
将内部媒体对象转换为 REST 响应模型。
|
||||
|
||||
:param medias: 统一媒体信息列表
|
||||
:return: REST 媒体响应列表
|
||||
"""
|
||||
return [schemas.MediaInfo(**media.to_dict()) for media in medias]
|
||||
|
||||
|
||||
@router.get(
|
||||
"/trending",
|
||||
summary="查询 AniList 当前趋势榜",
|
||||
response_model=list[schemas.MediaInfo],
|
||||
)
|
||||
async def anilist_trending(
|
||||
page: PageParam = 1,
|
||||
count: CountParam = 20,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> list[schemas.MediaInfo]:
|
||||
"""查询 AniList TRENDING NOW 榜单"""
|
||||
medias = await AniListChain().async_trending(page=page, count=count)
|
||||
return _serialize_medias(medias)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/popular-this-season",
|
||||
summary="查询 AniList 本季热门榜",
|
||||
response_model=list[schemas.MediaInfo],
|
||||
)
|
||||
async def anilist_popular_this_season(
|
||||
page: PageParam = 1,
|
||||
count: CountParam = 20,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> list[schemas.MediaInfo]:
|
||||
"""查询 AniList POPULAR THIS SEASON 榜单"""
|
||||
medias = await AniListChain().async_popular_this_season(page=page, count=count)
|
||||
return _serialize_medias(medias)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/discover",
|
||||
summary="探索 AniList 动画",
|
||||
response_model=list[schemas.MediaInfo],
|
||||
)
|
||||
async def anilist_discover(
|
||||
page: PageParam = 1,
|
||||
count: CountParam = 20,
|
||||
search: Optional[str] = None,
|
||||
genre: Optional[str] = None,
|
||||
media_format: Optional[str] = Query(None, alias="format"),
|
||||
season: Optional[str] = None,
|
||||
season_year: Optional[int] = None,
|
||||
status: Optional[str] = None,
|
||||
country: Optional[str] = None,
|
||||
sort: Optional[str] = None,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> list[schemas.MediaInfo]:
|
||||
"""按标题、类型、风格、季度、年份、状态、地区和排序探索 AniList 动画"""
|
||||
medias = await AniListChain().async_discover(
|
||||
page=page,
|
||||
count=count,
|
||||
search=search,
|
||||
genre=genre,
|
||||
media_format=media_format,
|
||||
season=season,
|
||||
season_year=season_year,
|
||||
status=status,
|
||||
country=country,
|
||||
sort=sort,
|
||||
)
|
||||
return _serialize_medias(medias)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/credits/{anilist_id}",
|
||||
summary="查询 AniList 配音演员",
|
||||
response_model=list[schemas.MediaPerson],
|
||||
)
|
||||
async def anilist_credits(
|
||||
anilist_id: int,
|
||||
page: PageParam = 1,
|
||||
count: CountParam = 20,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> list[schemas.MediaPerson]:
|
||||
"""查询 AniList 动画的日语配音演员"""
|
||||
return await AniListChain().async_credits(
|
||||
anilist_id=anilist_id, page=page, count=count
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/recommend/{anilist_id}",
|
||||
summary="查询 AniList 相关推荐",
|
||||
response_model=list[schemas.MediaInfo],
|
||||
)
|
||||
async def anilist_recommendations(
|
||||
anilist_id: int,
|
||||
page: PageParam = 1,
|
||||
count: CountParam = 20,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> list[schemas.MediaInfo]:
|
||||
"""查询 AniList 动画相关推荐"""
|
||||
medias = await AniListChain().async_recommendations(
|
||||
anilist_id=anilist_id, page=page, count=count
|
||||
)
|
||||
return _serialize_medias(medias)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/person/{person_id}",
|
||||
summary="查询 AniList 人物详情",
|
||||
response_model=schemas.MediaPerson,
|
||||
)
|
||||
async def anilist_person(
|
||||
person_id: int,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> Optional[schemas.MediaPerson]:
|
||||
"""根据 AniList 人物 ID 查询详情"""
|
||||
return await AniListChain().async_person_detail(person_id=person_id)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/person/credits/{person_id}",
|
||||
summary="查询 AniList 人物作品",
|
||||
response_model=list[schemas.MediaInfo],
|
||||
)
|
||||
async def anilist_person_credits(
|
||||
person_id: int,
|
||||
page: PageParam = 1,
|
||||
count: CountParam = 20,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> list[schemas.MediaInfo]:
|
||||
"""查询 AniList 人物参与的动画作品"""
|
||||
medias = await AniListChain().async_person_credits(
|
||||
person_id=person_id, page=page, count=count
|
||||
)
|
||||
return _serialize_medias(medias)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/{anilist_id}",
|
||||
summary="查询 AniList 动画详情",
|
||||
response_model=schemas.MediaInfo,
|
||||
)
|
||||
async def anilist_info(
|
||||
anilist_id: int,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> schemas.MediaInfo:
|
||||
"""根据 AniList 媒体 ID 查询动画详情"""
|
||||
info = await AniListChain().async_info(anilist_id)
|
||||
if not info:
|
||||
return schemas.MediaInfo()
|
||||
return schemas.MediaInfo(**MediaInfo(anilist_info=info).to_dict())
|
||||
@@ -39,6 +39,9 @@ def _anthropic_error_response(
|
||||
|
||||
|
||||
def _check_auth(api_key: Optional[str]) -> Optional[JSONResponse]:
|
||||
"""
|
||||
Anthropic 兼容接口以 API_TOKEN 认证受信客户端,认证通过即按管理员级 Agent 集成处理。
|
||||
"""
|
||||
if not api_key or api_key != settings.API_TOKEN:
|
||||
return _anthropic_error_response(
|
||||
"invalid x-api-key",
|
||||
@@ -122,6 +125,7 @@ async def messages(
|
||||
|
||||
session_seed = anthropic_version or "anthropic"
|
||||
session_id = build_session_id(f"{session_seed}:{uuid.uuid4().hex}", SESSION_PREFIX)
|
||||
# 兼容接口的 API_TOKEN 客户端按管理员级 MoviePilot Agent 集成处理。
|
||||
agent = _CollectingMoviePilotAgent(
|
||||
session_id=session_id,
|
||||
user_id=session_id,
|
||||
|
||||
@@ -4,7 +4,7 @@ from fastapi import APIRouter, HTTPException
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app import schemas
|
||||
from app.core.auth_bridge import build_token_response, consume_plugin_auth_ticket
|
||||
from app.core.auth import build_token_response, consume_plugin_auth_ticket
|
||||
from app.core.plugin import PluginManager
|
||||
from app.db.models.passkey import PassKey
|
||||
from app.db.models.user import User
|
||||
|
||||
@@ -7,6 +7,7 @@ from sqlalchemy.orm import Session
|
||||
from app import schemas
|
||||
from app.chain.dashboard import DashboardChain
|
||||
from app.chain.storage import StorageChain
|
||||
from app.core.config import settings
|
||||
from app.core.security import verify_apitoken
|
||||
from app.db import get_db
|
||||
from app.db.models.transferhistory import TransferHistory
|
||||
@@ -18,7 +19,7 @@ from app.utils.system import SystemUtils
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _build_statistic(name: Optional[str] = None) -> schemas.Statistic:
|
||||
def _build_statistic(db: Session, name: Optional[str] = None) -> schemas.Statistic:
|
||||
"""
|
||||
构建媒体数量统计信息。
|
||||
"""
|
||||
@@ -39,8 +40,14 @@ def _build_statistic(name: Optional[str] = None) -> schemas.Statistic:
|
||||
if not has_episode_count:
|
||||
# 所有媒体服务都未提供剧集统计时,返回 None 供前端展示“未获取”。
|
||||
ret_statistic.episode_count = None
|
||||
return ret_statistic
|
||||
return schemas.Statistic()
|
||||
else:
|
||||
ret_statistic = schemas.Statistic()
|
||||
|
||||
movie_count_month, tv_count_month, episode_count_month = TransferHistory.monthly_media_statistics(db)
|
||||
ret_statistic.movie_count_month = movie_count_month
|
||||
ret_statistic.tv_count_month = tv_count_month
|
||||
ret_statistic.episode_count_month = episode_count_month
|
||||
return ret_statistic
|
||||
|
||||
|
||||
def _build_storage() -> schemas.Storage:
|
||||
@@ -67,7 +74,8 @@ def _build_downloader(name: Optional[str] = None) -> schemas.DownloaderInfo:
|
||||
# 下载目录空间
|
||||
download_dirs = DirectoryHelper().get_local_download_dirs()
|
||||
_, free_space = SystemUtils.space_usage(
|
||||
[Path(d.download_path) for d in download_dirs]
|
||||
[Path(d.download_path) for d in download_dirs],
|
||||
btrfs_fsid_dedup=settings.BTRFS_FSID_DEDUP,
|
||||
)
|
||||
# 下载器信息
|
||||
downloader_info = schemas.DownloaderInfo()
|
||||
@@ -84,22 +92,27 @@ def _build_downloader(name: Optional[str] = None) -> schemas.DownloaderInfo:
|
||||
|
||||
@router.get("/statistic", summary="媒体数量统计", response_model=schemas.Statistic)
|
||||
def statistic(
|
||||
name: Optional[str] = None, _: Any = Depends(get_current_active_superuser)
|
||||
name: Optional[str] = None,
|
||||
db: Session = Depends(get_db),
|
||||
_: Any = Depends(get_current_active_superuser),
|
||||
) -> Any:
|
||||
"""
|
||||
查询媒体数量统计信息
|
||||
"""
|
||||
return _build_statistic(name)
|
||||
return _build_statistic(db, name)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/statistic2", summary="媒体数量统计(API_TOKEN)", response_model=schemas.Statistic
|
||||
)
|
||||
def statistic2(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
|
||||
def statistic2(
|
||||
_: Annotated[str, Depends(verify_apitoken)],
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
"""
|
||||
查询媒体数量统计信息 API_TOKEN认证(?token=xxx)
|
||||
"""
|
||||
return _build_statistic()
|
||||
return _build_statistic(db)
|
||||
|
||||
|
||||
@router.get("/storage", summary="本地存储空间", response_model=schemas.Storage)
|
||||
@@ -128,6 +141,14 @@ def processes(_: Any = Depends(get_current_active_superuser)) -> Any:
|
||||
return SystemUtils.processes()
|
||||
|
||||
|
||||
@router.get("/system", summary="系统摘要信息", response_model=schemas.DashboardSystemInfo)
|
||||
def system_info(_: Any = Depends(get_current_active_superuser)) -> Any:
|
||||
"""
|
||||
查询仪表板系统摘要信息
|
||||
"""
|
||||
return SystemUtils.dashboard_system_info()
|
||||
|
||||
|
||||
@router.get("/downloader", summary="下载器信息", response_model=schemas.DownloaderInfo)
|
||||
def downloader(
|
||||
name: Optional[str] = None, _: Any = Depends(get_current_active_superuser)
|
||||
@@ -158,6 +179,23 @@ async def schedule(_: Any = Depends(get_current_active_superuser)) -> Any:
|
||||
return Scheduler().list()
|
||||
|
||||
|
||||
@router.get(
|
||||
"/schedule/{job_id}/progress",
|
||||
summary="后台服务进度",
|
||||
response_model=schemas.Response,
|
||||
)
|
||||
async def schedule_progress(
|
||||
job_id: str, _: Any = Depends(get_current_active_superuser)
|
||||
) -> Any:
|
||||
"""
|
||||
查询指定后台服务的执行进度。
|
||||
"""
|
||||
progress = Scheduler().get_progress(job_id)
|
||||
if not progress:
|
||||
return schemas.Response(success=False, message="后台服务不存在")
|
||||
return schemas.Response(success=True, data=progress.model_dump())
|
||||
|
||||
|
||||
@router.get(
|
||||
"/schedule2",
|
||||
summary="后台服务(API_TOKEN)",
|
||||
@@ -170,6 +208,23 @@ async def schedule2(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
|
||||
return Scheduler().list()
|
||||
|
||||
|
||||
@router.get(
|
||||
"/schedule2/{job_id}/progress",
|
||||
summary="后台服务进度(API_TOKEN)",
|
||||
response_model=schemas.Response,
|
||||
)
|
||||
async def schedule_progress2(
|
||||
job_id: str, _: Annotated[str, Depends(verify_apitoken)]
|
||||
) -> Any:
|
||||
"""
|
||||
查询指定后台服务的执行进度 API_TOKEN认证(?token=xxx)
|
||||
"""
|
||||
progress = Scheduler().get_progress(job_id)
|
||||
if not progress:
|
||||
return schemas.Response(success=False, message="后台服务不存在")
|
||||
return schemas.Response(success=True, data=progress.model_dump())
|
||||
|
||||
|
||||
@router.get("/transfer", summary="文件整理统计", response_model=List[int])
|
||||
async def transfer(
|
||||
days: Optional[int] = 7,
|
||||
@@ -199,22 +254,26 @@ def cpu2(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
|
||||
return SystemUtils.cpu_usage()
|
||||
|
||||
|
||||
@router.get("/memory", summary="获取当前内存使用量和使用率", response_model=List[int])
|
||||
@router.get(
|
||||
"/memory",
|
||||
summary="获取当前应用与系统内存信息",
|
||||
response_model=schemas.DashboardMemoryInfo,
|
||||
)
|
||||
def memory(_: Any = Depends(get_current_active_superuser)) -> Any:
|
||||
"""
|
||||
获取当前内存使用率
|
||||
获取当前应用与系统内存信息
|
||||
"""
|
||||
return SystemUtils.memory_usage()
|
||||
|
||||
|
||||
@router.get(
|
||||
"/memory2",
|
||||
summary="获取当前内存使用量和使用率(API_TOKEN)",
|
||||
response_model=List[int],
|
||||
summary="获取当前应用与系统内存信息(API_TOKEN)",
|
||||
response_model=schemas.DashboardMemoryInfo,
|
||||
)
|
||||
def memory2(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
|
||||
"""
|
||||
获取当前内存使用率 API_TOKEN认证(?token=xxx)
|
||||
获取当前应用与系统内存信息 API_TOKEN认证(?token=xxx)
|
||||
"""
|
||||
return SystemUtils.memory_usage()
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import Any, List, Annotated, Optional
|
||||
from typing import Any, List, Annotated, Literal, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Body
|
||||
|
||||
@@ -9,12 +9,40 @@ from app.core.context import MediaInfo, Context, SubtitleInfo, TorrentInfo
|
||||
from app.core.metainfo import MetaInfo
|
||||
from app.core.security import verify_token
|
||||
from app.db.models.user import User
|
||||
from app.db.site_oper import SiteOper
|
||||
from app.db.systemconfig_oper import SystemConfigOper
|
||||
from app.db.user_oper import get_current_active_user
|
||||
from app.helper.directory import DirectoryHelper
|
||||
from app.schemas.types import SystemConfigKey
|
||||
from app.utils.security import SecurityUtils
|
||||
|
||||
router = APIRouter()
|
||||
MediaSource = Literal["themoviedb", "douban", "bangumi", "anilist"]
|
||||
|
||||
|
||||
def _prepare_subtitle_download(subtitle: SubtitleInfo) -> tuple[bool, str]:
|
||||
"""
|
||||
校验字幕下载签名,并用服务端站点配置覆盖请求凭据。
|
||||
"""
|
||||
if subtitle.site is None:
|
||||
return False, "字幕站点信息为空"
|
||||
|
||||
clean_url = SecurityUtils.verify_signed_url(
|
||||
subtitle.enclosure,
|
||||
purpose=SecurityUtils.subtitle_download_purpose(subtitle.site),
|
||||
)
|
||||
if not clean_url:
|
||||
return False, "字幕下载链接签名无效"
|
||||
|
||||
site = SiteOper().get(subtitle.site)
|
||||
if not site:
|
||||
return False, "字幕站点信息不存在"
|
||||
|
||||
subtitle.enclosure = clean_url
|
||||
subtitle.site_cookie = site.cookie
|
||||
subtitle.site_ua = site.ua
|
||||
subtitle.site_proxy = bool(site.proxy)
|
||||
return True, ""
|
||||
|
||||
|
||||
@router.get("/", summary="正在下载", response_model=List[schemas.DownloaderTorrent])
|
||||
@@ -70,6 +98,10 @@ def add(
|
||||
torrent_in: schemas.TorrentInfo,
|
||||
tmdbid: Annotated[int | None, Body()] = None,
|
||||
doubanid: Annotated[str | None, Body()] = None,
|
||||
bangumiid: Annotated[int | None, Body()] = None,
|
||||
anilistid: Annotated[int | None, Body()] = None,
|
||||
media_source: Annotated[MediaSource | None, Body()] = None,
|
||||
media_id: Annotated[str | None, Body()] = None,
|
||||
downloader: Annotated[str | None, Body()] = None,
|
||||
# 保存路径, 支持<storage>:<path>, 如rclone:/MP, smb:/server/share/Movies等
|
||||
save_path: Annotated[str | None, Body()] = None,
|
||||
@@ -81,15 +113,20 @@ def add(
|
||||
# 元数据
|
||||
metainfo = MetaInfo(title=torrent_in.title, subtitle=torrent_in.description)
|
||||
# 媒体信息
|
||||
if tmdbid or doubanid:
|
||||
if tmdbid or doubanid or bangumiid or anilistid or media_id:
|
||||
mediainfo = MediaChain().recognize_media(
|
||||
meta=metainfo,
|
||||
source=media_source,
|
||||
mediaid=media_id,
|
||||
tmdbid=tmdbid,
|
||||
doubanid=doubanid,
|
||||
bangumiid=bangumiid,
|
||||
anilistid=anilistid,
|
||||
)
|
||||
else:
|
||||
mediainfo = MediaChain().recognize_by_meta(
|
||||
metainfo,
|
||||
source=media_source,
|
||||
obtain_images=False,
|
||||
)
|
||||
if not mediainfo:
|
||||
@@ -119,6 +156,10 @@ def download_subtitle(
|
||||
subtitle_in: schemas.SubtitleInfo,
|
||||
tmdbid: Annotated[int | None, Body()] = None,
|
||||
doubanid: Annotated[str | None, Body()] = None,
|
||||
bangumiid: Annotated[int | None, Body()] = None,
|
||||
anilistid: Annotated[int | None, Body()] = None,
|
||||
media_source: Annotated[MediaSource | None, Body()] = None,
|
||||
media_id: Annotated[str | None, Body()] = None,
|
||||
save_path: Annotated[str | None, Body()] = None,
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
) -> Any:
|
||||
@@ -127,10 +168,18 @@ def download_subtitle(
|
||||
"""
|
||||
subtitle_info = SubtitleInfo()
|
||||
subtitle_info.from_dict(subtitle_in.model_dump())
|
||||
valid, message = _prepare_subtitle_download(subtitle_info)
|
||||
if not valid:
|
||||
return schemas.Response(success=False, message=message)
|
||||
|
||||
success, message, saved_files = DownloadChain().download_subtitle(
|
||||
subtitle=subtitle_info,
|
||||
media_source=media_source,
|
||||
media_id=media_id,
|
||||
tmdbid=tmdbid,
|
||||
doubanid=doubanid,
|
||||
bangumiid=bangumiid,
|
||||
anilistid=anilistid,
|
||||
save_path=save_path,
|
||||
username=current_user.name,
|
||||
)
|
||||
|
||||
@@ -22,8 +22,9 @@ from app.db.models import User
|
||||
from app.db.models.downloadhistory import DownloadHistory, DownloadFiles
|
||||
from app.db.models.transferhistory import TransferHistory
|
||||
from app.db.user_oper import (
|
||||
get_current_active_superuser_async,
|
||||
get_current_active_manage_user,
|
||||
get_current_active_superuser,
|
||||
get_current_active_superuser_async,
|
||||
)
|
||||
from app.helper.progress import ProgressHelper
|
||||
from app.schemas.types import EventType
|
||||
@@ -139,7 +140,7 @@ async def download_history(
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> Any:
|
||||
"""
|
||||
查询下载历史记录
|
||||
按下载时间倒序查询下载历史记录
|
||||
"""
|
||||
return await DownloadHistory.async_list_by_page(db, page, count)
|
||||
|
||||
@@ -223,7 +224,7 @@ def delete_transfer_history(
|
||||
deletesrc: Optional[bool] = False,
|
||||
deletedest: Optional[bool] = False,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
删除整理记录
|
||||
@@ -264,7 +265,7 @@ def delete_transfer_history(
|
||||
def ai_redo_transfer_history(
|
||||
history_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
手动触发单条历史记录的 AI 重新整理,并返回进度键。
|
||||
@@ -293,7 +294,7 @@ def ai_redo_transfer_history(
|
||||
def batch_ai_redo_transfer_history(
|
||||
payload: schemas.BatchTransferHistoryRedoRequest,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
手动触发多条历史记录的 AI 批量重新整理,并返回进度键。
|
||||
|
||||
@@ -36,7 +36,10 @@ class LlmTestRequest(BaseModel):
|
||||
base_url: Optional[str] = None
|
||||
base_url_preset: Optional[str] = None
|
||||
user_agent: Optional[str] = None
|
||||
temperature: Optional[float] = None
|
||||
use_proxy: Optional[bool] = None
|
||||
api_protocol: Optional[str] = None
|
||||
web_search_mode: Optional[str] = None
|
||||
|
||||
|
||||
class LlmProviderAuthStartRequest(BaseModel):
|
||||
@@ -48,7 +51,7 @@ class LlmProviderAuthStartRequest(BaseModel):
|
||||
method: str
|
||||
|
||||
|
||||
def _sanitize_llm_test_error(message: str, api_key: Optional[str] = None) -> str:
|
||||
def _sanitize_llm_error(message: str, api_key: Optional[str] = None) -> str:
|
||||
"""
|
||||
清理错误信息中的敏感字段,避免回显密钥。
|
||||
"""
|
||||
@@ -70,11 +73,14 @@ def _sanitize_llm_test_error(message: str, api_key: Optional[str] = None) -> str
|
||||
)
|
||||
|
||||
normalized_message = sanitized.lower().replace("_", "").replace(" ", "")
|
||||
if "str" in normalized_message and "modeldump" in normalized_message:
|
||||
if "str" in normalized_message and (
|
||||
"modeldump" in normalized_message
|
||||
or "setprivateattributes" in normalized_message
|
||||
):
|
||||
return (
|
||||
"服务返回内容不是兼容的模型响应,"
|
||||
"请检查基础地址是否填写为 API Base URL,不要填写网页地址或完整的 "
|
||||
"chat/completions 路径"
|
||||
"服务返回内容不是兼容的模型响应,请检查基础地址是否填写为 "
|
||||
"API Base URL,如果服务要求 /v1 等版本路径,请包含在基础地址中,"
|
||||
"不要填写网页地址或完整的 chat/completions 路径"
|
||||
)
|
||||
return sanitized
|
||||
|
||||
@@ -113,7 +119,10 @@ async def get_llm_models(
|
||||
},
|
||||
)
|
||||
except Exception as err:
|
||||
return schemas.Response(success=False, message=str(err))
|
||||
return schemas.Response(
|
||||
success=False,
|
||||
message=_sanitize_llm_error(str(err), api_key),
|
||||
)
|
||||
|
||||
|
||||
@router.get("/providers", summary="获取LLM提供商目录", response_model=schemas.Response)
|
||||
@@ -262,6 +271,8 @@ async def llm_test(
|
||||
base_url_preset=settings.LLM_BASE_URL_PRESET,
|
||||
user_agent=settings.LLM_USER_AGENT,
|
||||
use_proxy=settings.LLM_USE_PROXY,
|
||||
api_protocol=settings.LLM_API_PROTOCOL,
|
||||
web_search_mode=settings.LLM_WEB_SEARCH_MODE,
|
||||
)
|
||||
|
||||
if not payload.provider:
|
||||
@@ -286,16 +297,22 @@ async def llm_test(
|
||||
)
|
||||
|
||||
try:
|
||||
result = await LLMHelper.test_current_settings(
|
||||
provider=payload.provider,
|
||||
model=payload.model,
|
||||
thinking_level=payload.thinking_level,
|
||||
api_key=payload.api_key,
|
||||
base_url=payload.base_url,
|
||||
base_url_preset=payload.base_url_preset,
|
||||
user_agent=payload.user_agent,
|
||||
use_proxy=payload.use_proxy,
|
||||
)
|
||||
test_kwargs = {
|
||||
"provider": payload.provider,
|
||||
"model": payload.model,
|
||||
"thinking_level": payload.thinking_level,
|
||||
"api_key": payload.api_key,
|
||||
"base_url": payload.base_url,
|
||||
"base_url_preset": payload.base_url_preset,
|
||||
"user_agent": payload.user_agent,
|
||||
"use_proxy": payload.use_proxy,
|
||||
"api_protocol": payload.api_protocol,
|
||||
"web_search_mode": payload.web_search_mode,
|
||||
}
|
||||
if payload.temperature is not None:
|
||||
test_kwargs["temperature"] = payload.temperature
|
||||
|
||||
result = await LLMHelper.test_current_settings(**test_kwargs)
|
||||
if not result.get("reply_preview"):
|
||||
return schemas.Response(
|
||||
success=False,
|
||||
@@ -312,5 +329,5 @@ async def llm_test(
|
||||
except Exception as err:
|
||||
return schemas.Response(
|
||||
success=False,
|
||||
message=_sanitize_llm_test_error(str(err), payload.api_key),
|
||||
message=_sanitize_llm_error(str(err), payload.api_key),
|
||||
)
|
||||
|
||||
@@ -3,9 +3,10 @@ from typing import Any, List, Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, Form, HTTPException, Request, Response
|
||||
from fastapi.security import OAuth2PasswordRequestForm
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from app import schemas
|
||||
from app.chain.user import UserChain
|
||||
from app.chain.user import MfaRequired, UserChain
|
||||
from app.core import security
|
||||
from app.core.config import settings
|
||||
from app.db.systemconfig_oper import SystemConfigOper
|
||||
@@ -31,11 +32,14 @@ def login_access_token(
|
||||
)
|
||||
|
||||
if not success:
|
||||
# 如果是需要MFA验证,返回特殊标识
|
||||
if user_or_message == "MFA_REQUIRED":
|
||||
raise HTTPException(
|
||||
# 只有密码已经验证通过时才返回 MFA 方法,避免泄露账号安全配置。
|
||||
if isinstance(user_or_message, MfaRequired):
|
||||
return JSONResponse(
|
||||
status_code=401,
|
||||
detail="需要双重验证,请提供验证码或使用通行密钥",
|
||||
content={
|
||||
"detail": "需要二次验证",
|
||||
"mfa_methods": list(user_or_message.methods),
|
||||
},
|
||||
headers={"X-MFA-Required": "true"},
|
||||
)
|
||||
raise HTTPException(status_code=401, detail="用户名或密码错误")
|
||||
@@ -86,7 +90,7 @@ def wallpaper() -> Any:
|
||||
"""
|
||||
url = WallpaperHelper().get_wallpaper()
|
||||
if url:
|
||||
return schemas.Response(success=True, message=url)
|
||||
return schemas.Response(success=True, data=url)
|
||||
return schemas.Response(success=False)
|
||||
|
||||
|
||||
|
||||
@@ -19,6 +19,7 @@ router = APIRouter()
|
||||
# MCP 协议版本
|
||||
MCP_PROTOCOL_VERSIONS = ["2025-11-25", "2025-06-18", "2024-11-05"]
|
||||
MCP_PROTOCOL_VERSION = MCP_PROTOCOL_VERSIONS[0] # 默认使用最新版本
|
||||
# MCP 经 API_TOKEN / X-API-KEY 认证后是管理员级集成入口;隐藏工具只收敛暴露面,不构成权限边界。
|
||||
MCP_HIDDEN_TOOLS = {
|
||||
"execute_command",
|
||||
"search_web",
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from pathlib import Path
|
||||
from typing import List, Any, Union, Annotated, Optional
|
||||
from typing import Annotated, Any, List, Optional, Union
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
@@ -9,15 +9,84 @@ from app.chain.tmdb import TmdbChain
|
||||
from app.core.config import settings
|
||||
from app.core.context import Context
|
||||
from app.core.event import eventmanager
|
||||
from app.core.metainfo import MetaInfo
|
||||
from app.core.meta import MetaBase
|
||||
from app.core.metainfo import MetaInfo, MetaInfoPath
|
||||
from app.core.security import verify_token, verify_apitoken
|
||||
from app.db.models import User
|
||||
from app.db.user_oper import get_current_active_user, get_current_active_superuser
|
||||
from app.schemas import MediaType, MediaRecognizeConvertEventData
|
||||
from app.schemas.category import CategoryConfig
|
||||
from app.schemas.types import ChainEventType
|
||||
from app.utils.media import MEDIA_SOURCE_ID_FIELDS, parse_media_key
|
||||
|
||||
router = APIRouter()
|
||||
MediaSource = str
|
||||
|
||||
|
||||
def _build_recognize_metainfo(
|
||||
title: str,
|
||||
subtitle: Optional[str] = None,
|
||||
custom_words: Optional[str] = None,
|
||||
) -> MetaBase:
|
||||
"""构造标题识别元数据,并兼容第三方客户端传入媒体文件路径。"""
|
||||
custom_word_list = custom_words.split("\n") if custom_words else None
|
||||
normalized_title = title.replace("\\", "/")
|
||||
title_path = Path(normalized_title)
|
||||
if (
|
||||
("/" in title or "\\" in title)
|
||||
and "://" not in title
|
||||
and title_path.suffix.lower() in settings.RMT_MEDIAEXT
|
||||
):
|
||||
metainfo = MetaInfoPath(
|
||||
title_path,
|
||||
custom_words=custom_word_list,
|
||||
)
|
||||
metainfo.title = title
|
||||
return metainfo
|
||||
return MetaInfo(title, subtitle, custom_words=custom_word_list)
|
||||
|
||||
|
||||
def _build_media_seasons(
|
||||
mediainfo: Any, season: Optional[int] = None,
|
||||
) -> List[schemas.MediaSeason]:
|
||||
"""将任意数据源的统一媒体信息转换为季信息响应。"""
|
||||
seasons_info = []
|
||||
for item in mediainfo.season_info or []:
|
||||
season_number = item.get("season_number")
|
||||
if season is not None and season_number != season:
|
||||
continue
|
||||
seasons_info.append(schemas.MediaSeason(
|
||||
air_date=item.get("air_date"),
|
||||
episode_count=item.get("episode_count"),
|
||||
name=item.get("name"),
|
||||
overview=item.get("overview"),
|
||||
poster_path=item.get("poster_path") or mediainfo.poster_path,
|
||||
season_number=season_number,
|
||||
vote_average=item.get("vote_average"),
|
||||
))
|
||||
if seasons_info:
|
||||
return seasons_info
|
||||
|
||||
season_numbers = sorted((mediainfo.seasons or {}).keys())
|
||||
if season is not None:
|
||||
season_numbers = [season]
|
||||
elif not season_numbers:
|
||||
season_numbers = [mediainfo.season or 1]
|
||||
return [
|
||||
schemas.MediaSeason(
|
||||
season_number=season_number,
|
||||
poster_path=mediainfo.poster_path,
|
||||
name=f"第 {season_number} 季",
|
||||
air_date=mediainfo.release_date,
|
||||
overview=mediainfo.overview,
|
||||
vote_average=mediainfo.vote_average,
|
||||
episode_count=(
|
||||
len((mediainfo.seasons or {}).get(season_number) or [])
|
||||
or mediainfo.number_of_episodes
|
||||
),
|
||||
)
|
||||
for season_number in season_numbers
|
||||
]
|
||||
|
||||
|
||||
@router.get(
|
||||
@@ -26,14 +95,24 @@ router = APIRouter()
|
||||
async def recognize(
|
||||
title: str,
|
||||
subtitle: Optional[str] = None,
|
||||
custom_words: Optional[str] = None,
|
||||
source: Optional[MediaSource] = None,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> Any:
|
||||
"""
|
||||
根据标题、副标题识别媒体信息
|
||||
:param title: 标题
|
||||
:param subtitle: 副标题
|
||||
:param custom_words: 临时识别词(每行一条规则),传入时仅在本次识别中生效,不会保存到系统配置
|
||||
:param source: 请求级识别数据源
|
||||
:param _:
|
||||
"""
|
||||
# 识别媒体信息
|
||||
metainfo = MetaInfo(title, subtitle)
|
||||
mediainfo = await MediaChain().async_recognize_by_meta(metainfo)
|
||||
# 识别媒体信息,传入临时识别词时优先于系统配置的识别词生效
|
||||
metainfo = _build_recognize_metainfo(title, subtitle, custom_words)
|
||||
mediainfo = await MediaChain().async_recognize_by_meta(
|
||||
metainfo,
|
||||
source=source,
|
||||
)
|
||||
if mediainfo:
|
||||
return Context(meta_info=metainfo, media_info=mediainfo).to_dict()
|
||||
return schemas.Context()
|
||||
@@ -48,25 +127,29 @@ async def recognize2(
|
||||
_: Annotated[str, Depends(verify_apitoken)],
|
||||
title: str,
|
||||
subtitle: Optional[str] = None,
|
||||
custom_words: Optional[str] = None,
|
||||
source: Optional[MediaSource] = None,
|
||||
) -> Any:
|
||||
"""
|
||||
根据标题、副标题识别媒体信息 API_TOKEN认证(?token=xxx)
|
||||
"""
|
||||
# 识别媒体信息
|
||||
return await recognize(title, subtitle)
|
||||
return await recognize(title, subtitle, custom_words, source)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/recognize_file", summary="识别媒体信息(文件)", response_model=schemas.Context
|
||||
)
|
||||
async def recognize_file(
|
||||
path: str, _: schemas.TokenPayload = Depends(verify_token)
|
||||
path: str,
|
||||
source: Optional[MediaSource] = None,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> Any:
|
||||
"""
|
||||
根据文件路径识别媒体信息
|
||||
"""
|
||||
# 识别媒体信息
|
||||
context = await MediaChain().async_recognize_by_path(path)
|
||||
context = await MediaChain().async_recognize_by_path(path, source=source)
|
||||
if context:
|
||||
return context.to_dict()
|
||||
return schemas.Context()
|
||||
@@ -78,13 +161,15 @@ async def recognize_file(
|
||||
response_model=schemas.Context,
|
||||
)
|
||||
async def recognize_file2(
|
||||
path: str, _: Annotated[str, Depends(verify_apitoken)]
|
||||
path: str,
|
||||
_: Annotated[str, Depends(verify_apitoken)],
|
||||
source: Optional[MediaSource] = None,
|
||||
) -> Any:
|
||||
"""
|
||||
根据文件路径识别媒体信息 API_TOKEN认证(?token=xxx)
|
||||
"""
|
||||
# 识别媒体信息
|
||||
return await recognize_file(path)
|
||||
return await recognize_file(path, source)
|
||||
|
||||
|
||||
@router.get("/search", summary="搜索媒体/人物信息", response_model=List[dict])
|
||||
@@ -93,10 +178,19 @@ async def search(
|
||||
type: Optional[str] = "media",
|
||||
page: int = 1,
|
||||
count: int = 8,
|
||||
source: Optional[MediaSource] = None,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> Any:
|
||||
"""
|
||||
模糊搜索媒体/人物信息列表 media:媒体信息,person:人物信息
|
||||
模糊搜索媒体、合集或人物信息列表。
|
||||
|
||||
:param title: 搜索关键词
|
||||
:param type: 搜索类型,支持 media、collection、person
|
||||
:param page: 页码
|
||||
:param count: 每页数量
|
||||
:param source: 请求级搜索数据源
|
||||
:param _: Token校验
|
||||
:return: 搜索结果列表
|
||||
"""
|
||||
|
||||
def __get_source(obj: Union[schemas.MediaInfo, schemas.MediaPerson, dict]):
|
||||
@@ -109,15 +203,17 @@ async def search(
|
||||
|
||||
media_chain = MediaChain()
|
||||
if type == "media":
|
||||
_, medias = await media_chain.async_search(title=title)
|
||||
_, medias = await media_chain.async_search(title=title, source=source)
|
||||
result = [media.to_dict() for media in medias] if medias else []
|
||||
elif type == "collection":
|
||||
collections = await media_chain.async_search_collections(name=title)
|
||||
collections = await media_chain.async_search_collections(
|
||||
name=title, source=source
|
||||
)
|
||||
result = (
|
||||
[collection.to_dict() for collection in collections] if collections else []
|
||||
)
|
||||
else: # person
|
||||
persons = await media_chain.async_search_persons(name=title)
|
||||
persons = await media_chain.async_search_persons(name=title, source=source)
|
||||
result = [person.model_dump() for person in persons] if persons else []
|
||||
|
||||
if not result:
|
||||
@@ -137,26 +233,64 @@ async def search(
|
||||
def scrape(
|
||||
fileitem: schemas.FileItem,
|
||||
storage: Optional[str] = "local",
|
||||
media_source: Optional[MediaSource] = None,
|
||||
media_id: Optional[str] = None,
|
||||
type_name: Optional[MediaType] = None,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> Any:
|
||||
"""
|
||||
刮削媒体信息
|
||||
刮削媒体信息,可按请求指定媒体数据源及其原生ID
|
||||
|
||||
:param fileitem: 待刮削文件项
|
||||
:param storage: 文件所在存储
|
||||
:param media_source: 请求级媒体数据源
|
||||
:param media_id: 数据源原生ID
|
||||
:param type_name: 媒体类型
|
||||
:param _: Token校验
|
||||
"""
|
||||
if not fileitem or not fileitem.path:
|
||||
return schemas.Response(success=False, message="刮削路径无效")
|
||||
normalized_media_id = media_id.strip() if media_id else None
|
||||
if normalized_media_id and not media_source:
|
||||
return schemas.Response(
|
||||
success=False, message="指定媒体ID时必须同时指定媒体数据源"
|
||||
)
|
||||
if normalized_media_id and not normalized_media_id.isdigit():
|
||||
return schemas.Response(success=False, message="媒体ID格式无效")
|
||||
|
||||
chain = MediaChain()
|
||||
# 识别媒体信息
|
||||
context = chain.recognize_by_path(fileitem.path, obtain_images=True)
|
||||
if not context or not context.media_info:
|
||||
if normalized_media_id:
|
||||
meta_info = MetaInfoPath(Path(fileitem.path))
|
||||
media_info = chain.recognize_media(
|
||||
meta=meta_info,
|
||||
mtype=type_name,
|
||||
source=media_source,
|
||||
mediaid=normalized_media_id,
|
||||
)
|
||||
if media_info:
|
||||
media_info.scrape_source = media_source
|
||||
chain.obtain_images(mediainfo=media_info)
|
||||
else:
|
||||
context = chain.recognize_by_path(
|
||||
fileitem.path,
|
||||
source=media_source,
|
||||
obtain_images=True,
|
||||
)
|
||||
meta_info = context.meta_info if context else None
|
||||
media_info = context.media_info if context else None
|
||||
|
||||
if not media_info:
|
||||
return schemas.Response(success=False, message="刮削失败,无法识别媒体信息")
|
||||
if media_source:
|
||||
media_info.scrape_source = media_source
|
||||
if storage == "local":
|
||||
if not Path(fileitem.path).exists():
|
||||
return schemas.Response(success=False, message="刮削路径不存在")
|
||||
# 手动刮削 (暂时使用同步版本,可以后续优化为异步)
|
||||
chain.scrape_metadata(
|
||||
fileitem=fileitem,
|
||||
meta=context.meta_info,
|
||||
mediainfo=context.media_info,
|
||||
meta=meta_info,
|
||||
mediainfo=media_info,
|
||||
overwrite=True,
|
||||
)
|
||||
return schemas.Response(success=True, message=f"{fileitem.path} 刮削完成")
|
||||
@@ -237,13 +371,26 @@ async def seasons(
|
||||
查询媒体季信息
|
||||
"""
|
||||
if mediaid:
|
||||
if mediaid.startswith("tmdb:"):
|
||||
tmdbid = int(mediaid[5:])
|
||||
media_source, source_media_id = parse_media_key(mediaid)
|
||||
if media_source == "themoviedb":
|
||||
tmdbid = int(source_media_id)
|
||||
seasons_info = await TmdbChain().async_tmdb_seasons(tmdbid=tmdbid)
|
||||
if seasons_info:
|
||||
if season is not None:
|
||||
return [sea for sea in seasons_info if sea.season_number == season]
|
||||
return seasons_info
|
||||
elif media_source and source_media_id:
|
||||
mediainfo = await MediaChain().async_recognize_media(
|
||||
source=media_source,
|
||||
mediaid=source_media_id,
|
||||
mtype=MediaType.TV,
|
||||
cache=False,
|
||||
)
|
||||
if mediainfo:
|
||||
return _build_media_seasons(mediainfo, season)
|
||||
# 明确来源的查询不能按标题切换到默认识别源,避免辅助 TMDB 信息替换主身份。
|
||||
if media_source and source_media_id:
|
||||
return []
|
||||
if title:
|
||||
meta = MetaInfo(title)
|
||||
if year:
|
||||
@@ -254,7 +401,7 @@ async def seasons(
|
||||
obtain_images=False,
|
||||
)
|
||||
if mediainfo:
|
||||
if settings.RECOGNIZE_SOURCE == "themoviedb":
|
||||
if mediainfo.source == "themoviedb" and mediainfo.tmdb_id:
|
||||
seasons_info = await TmdbChain().async_tmdb_seasons(
|
||||
tmdbid=mediainfo.tmdb_id
|
||||
)
|
||||
@@ -264,19 +411,7 @@ async def seasons(
|
||||
sea for sea in seasons_info if sea.season_number == season
|
||||
]
|
||||
return seasons_info
|
||||
else:
|
||||
sea = season if season is not None else 1
|
||||
return [
|
||||
schemas.MediaSeason(
|
||||
season_number=sea,
|
||||
poster_path=mediainfo.poster_path,
|
||||
name=f"第 {sea} 季",
|
||||
air_date=mediainfo.release_date,
|
||||
overview=mediainfo.overview,
|
||||
vote_average=mediainfo.vote_average,
|
||||
episode_count=mediainfo.number_of_episodes,
|
||||
)
|
||||
]
|
||||
return _build_media_seasons(mediainfo, season)
|
||||
return []
|
||||
|
||||
|
||||
@@ -289,25 +424,22 @@ async def detail(
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> Any:
|
||||
"""
|
||||
根据媒体ID查询themoviedb或豆瓣媒体信息,type_name: 电影/电视剧
|
||||
根据带来源前缀的媒体ID查询媒体信息,type_name: 电影/电视剧
|
||||
"""
|
||||
mtype = MediaType(type_name)
|
||||
mediainfo = None
|
||||
mediachain = MediaChain()
|
||||
if mediaid.startswith("tmdb:"):
|
||||
media_source, source_media_id = parse_media_key(mediaid)
|
||||
if media_source and source_media_id:
|
||||
mediainfo = await mediachain.async_recognize_media(
|
||||
tmdbid=int(mediaid[5:]), mtype=mtype
|
||||
source=media_source,
|
||||
mediaid=source_media_id,
|
||||
mtype=mtype,
|
||||
)
|
||||
elif mediaid.startswith("douban:"):
|
||||
mediainfo = await mediachain.async_recognize_media(
|
||||
doubanid=mediaid[7:], mtype=mtype
|
||||
)
|
||||
elif mediaid.startswith("bangumi:"):
|
||||
mediainfo = await mediachain.async_recognize_media(
|
||||
bangumiid=int(mediaid[8:]), mtype=mtype
|
||||
)
|
||||
else:
|
||||
# 广播事件解析媒体信息
|
||||
if not mediainfo and (
|
||||
not media_source or media_source not in MEDIA_SOURCE_ID_FIELDS
|
||||
):
|
||||
# 旧探索插件可能只提供列表或转换事件,原生 ID 直识别失败后需保留原有兼容链路。
|
||||
event_data = MediaRecognizeConvertEventData(
|
||||
mediaid=mediaid, convert_type=settings.RECOGNIZE_SOURCE
|
||||
)
|
||||
@@ -318,15 +450,13 @@ async def detail(
|
||||
if event and event.event_data and event.event_data.media_dict:
|
||||
event_data: MediaRecognizeConvertEventData = event.event_data
|
||||
new_id = event_data.media_dict.get("id")
|
||||
if event_data.convert_type == "themoviedb":
|
||||
if new_id is not None and event_data.convert_type:
|
||||
mediainfo = await mediachain.async_recognize_media(
|
||||
tmdbid=new_id, mtype=mtype
|
||||
source=event_data.convert_type,
|
||||
mediaid=str(new_id),
|
||||
mtype=mtype,
|
||||
)
|
||||
elif event_data.convert_type == "douban":
|
||||
mediainfo = await mediachain.async_recognize_media(
|
||||
doubanid=new_id, mtype=mtype
|
||||
)
|
||||
elif title:
|
||||
if not mediainfo and title:
|
||||
# 使用名称识别兜底
|
||||
meta = MetaInfo(title)
|
||||
if year:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from typing import Any, List, Dict, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app import schemas
|
||||
@@ -16,10 +16,23 @@ from app.db.systemconfig_oper import SystemConfigOper
|
||||
from app.helper.mediaserver import MediaServerHelper
|
||||
from app.schemas import MediaType, NotExistMediaInfo
|
||||
from app.schemas.types import SystemConfigKey
|
||||
from app.utils.media import build_media_key, resolve_media_identity
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _require_mediaserver_result(result: Optional[List[Any]]) -> List[Any]:
|
||||
"""
|
||||
保留媒体服务器成功空列表,并把提供方失败转换为明确的网关错误。
|
||||
"""
|
||||
if result is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail="媒体服务器请求失败",
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
@router.get("/play/{itemid:path}", summary="在线播放")
|
||||
def play_item(
|
||||
itemid: str, _: schemas.TokenPayload = Depends(verify_token)
|
||||
@@ -130,7 +143,8 @@ def not_exists(
|
||||
exist_flag, no_exists = DownloadChain().get_no_exists_info(
|
||||
meta=meta, mediainfo=mediainfo
|
||||
)
|
||||
mediakey = mediainfo.tmdb_id or mediainfo.douban_id
|
||||
media_source, media_id = resolve_media_identity(media=mediainfo)
|
||||
mediakey = build_media_key(media_source, media_id)
|
||||
if mediainfo.type == MediaType.MOVIE:
|
||||
# 电影已存在时返回空列表,不存在时返回空对像列表
|
||||
return [] if exist_flag else [NotExistMediaInfo()]
|
||||
@@ -151,11 +165,12 @@ def latest(
|
||||
"""
|
||||
获取媒体服务器最新入库条目
|
||||
"""
|
||||
return (
|
||||
return _require_mediaserver_result(
|
||||
MediaServerChain().latest(
|
||||
server=server, count=count, username=userinfo.username
|
||||
server=server,
|
||||
count=count,
|
||||
username=userinfo.username,
|
||||
)
|
||||
or []
|
||||
)
|
||||
|
||||
|
||||
@@ -170,11 +185,12 @@ def playing(
|
||||
"""
|
||||
获取媒体服务器正在播放条目
|
||||
"""
|
||||
return (
|
||||
return _require_mediaserver_result(
|
||||
MediaServerChain().playing(
|
||||
server=server, count=count, username=userinfo.username
|
||||
server=server,
|
||||
count=count,
|
||||
username=userinfo.username,
|
||||
)
|
||||
or []
|
||||
)
|
||||
|
||||
|
||||
@@ -189,11 +205,12 @@ def library(
|
||||
"""
|
||||
获取媒体服务器媒体库列表
|
||||
"""
|
||||
return (
|
||||
return _require_mediaserver_result(
|
||||
MediaServerChain().librarys(
|
||||
server=server, username=userinfo.username, hidden=hidden
|
||||
server=server,
|
||||
username=userinfo.username,
|
||||
hidden=hidden,
|
||||
)
|
||||
or []
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -17,7 +17,7 @@ from app.db.message_oper import MessageOper
|
||||
from app.db.systemconfig_oper import SystemConfigOper
|
||||
from app.db.user_oper import get_current_active_superuser
|
||||
from app.helper.service import ServiceConfigHelper
|
||||
from app.helper.webpush import is_webpush_subscription_gone
|
||||
from app.helper.webpush import is_webpush_subscription_gone, webpush_options_for_endpoint
|
||||
from app.log import logger
|
||||
from app.modules.wechat.WXBizMsgCrypt3 import WXBizMsgCrypt
|
||||
from app.schemas.types import MessageChannel, SystemConfigKey
|
||||
@@ -316,6 +316,7 @@ def send_notification(
|
||||
data=json.dumps(payload.model_dump()),
|
||||
vapid_private_key=settings.VAPID.get("privateKey"),
|
||||
vapid_claims={"sub": settings.VAPID.get("subject")},
|
||||
**webpush_options_for_endpoint(sub.get("endpoint")),
|
||||
)
|
||||
except WebPushException as err:
|
||||
logger.error(f"WebPush发送失败: {str(err)}")
|
||||
|
||||
@@ -18,7 +18,12 @@ from app.db.models.passkey import PassKey
|
||||
from app.db.models.user import User
|
||||
from app.db.systemconfig_oper import SystemConfigOper
|
||||
from app.db.user_oper import get_current_active_user, get_current_active_user_async
|
||||
from app.helper.passkey import PassKeyHelper
|
||||
from app.helper.passkey import (
|
||||
PassKeyHelper,
|
||||
PassKeyRegistrationOriginMismatchError,
|
||||
PassKeyRegistrationVerificationError,
|
||||
PasskeyChallengeStore,
|
||||
)
|
||||
from app.log import logger
|
||||
from app.schemas.types import SystemConfigKey
|
||||
from app.utils.otp import OtpUtils
|
||||
@@ -83,17 +88,6 @@ def _verify_passkey_and_update(
|
||||
return success, new_sign_count
|
||||
|
||||
|
||||
async def _check_user_has_passkey(db: AsyncSession, user_id: int) -> bool:
|
||||
"""
|
||||
检查用户是否有 PassKey
|
||||
|
||||
:param db: 数据库会话
|
||||
:param user_id: 用户 ID
|
||||
:return: 是否有 PassKey
|
||||
"""
|
||||
return bool(await PassKey.async_get_by_user_id(db=db, user_id=user_id))
|
||||
|
||||
|
||||
# ==================== 请求模型 ====================
|
||||
|
||||
|
||||
@@ -122,12 +116,12 @@ class PassKeyDeleteRequest(schemas.BaseModel):
|
||||
|
||||
@router.get(
|
||||
"/status/{username}",
|
||||
summary="判断用户是否开启双重验证(MFA)",
|
||||
summary="判断用户是否开启二次验证",
|
||||
response_model=schemas.Response,
|
||||
)
|
||||
async def mfa_status(username: str, db: AsyncSession = Depends(get_async_db)) -> Any:
|
||||
"""
|
||||
检查指定用户是否启用了任何双重验证方式(OTP 或 PassKey)
|
||||
检查指定用户是否启用了二次验证
|
||||
"""
|
||||
user: User = await User.async_get_by_name(db, username)
|
||||
if not user:
|
||||
@@ -136,11 +130,7 @@ async def mfa_status(username: str, db: AsyncSession = Depends(get_async_db)) ->
|
||||
# 检查是否启用了OTP
|
||||
has_otp = user.is_otp
|
||||
|
||||
# 检查是否有PassKey
|
||||
has_passkey = await _check_user_has_passkey(db, user.id)
|
||||
|
||||
# 只要有任何一种验证方式,就需要双重验证
|
||||
return schemas.Response(success=(has_otp or has_passkey))
|
||||
return schemas.Response(success=has_otp)
|
||||
|
||||
|
||||
# ==================== OTP 相关接口 ====================
|
||||
@@ -181,14 +171,6 @@ async def otp_disable(
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""关闭当前用户的 OTP 验证功能"""
|
||||
# 安全检查:如果存在 PassKey,默认不允许关闭 OTP,除非配置允许
|
||||
has_passkey = await _check_user_has_passkey(db, current_user.id)
|
||||
if has_passkey and not settings.PASSKEY_ALLOW_REGISTER_WITHOUT_OTP:
|
||||
return schemas.Response(
|
||||
success=False,
|
||||
message="您已注册通行密钥,为了防止域名配置变更导致无法登录,请先删除所有通行密钥再关闭 OTP 验证",
|
||||
)
|
||||
|
||||
# 验证密码
|
||||
if not security.verify_password(data.password, str(current_user.hashed_password)):
|
||||
return schemas.Response(success=False, message="密码错误")
|
||||
@@ -209,7 +191,7 @@ class PassKeyRegistrationFinish(schemas.BaseModel):
|
||||
"""PassKey注册完成请求"""
|
||||
|
||||
credential: dict
|
||||
challenge: str
|
||||
transaction_token: str
|
||||
name: str = "通行密钥"
|
||||
|
||||
|
||||
@@ -223,7 +205,7 @@ class PassKeyAuthenticationFinish(schemas.BaseModel):
|
||||
"""PassKey认证完成请求"""
|
||||
|
||||
credential: dict
|
||||
challenge: str
|
||||
transaction_token: str
|
||||
|
||||
|
||||
@router.post(
|
||||
@@ -236,13 +218,6 @@ def passkey_register_start(
|
||||
) -> Any:
|
||||
"""开始注册 PassKey - 生成注册选项"""
|
||||
try:
|
||||
# 安全检查:默认需要先启用 OTP,除非配置允许在未启用 OTP 时注册
|
||||
if not current_user.is_otp and not settings.PASSKEY_ALLOW_REGISTER_WITHOUT_OTP:
|
||||
return schemas.Response(
|
||||
success=False,
|
||||
message="为了确保在域名配置错误时仍能找回访问权限,请先启用 OTP 验证码再注册通行密钥",
|
||||
)
|
||||
|
||||
# 获取用户已有的PassKey
|
||||
existing_passkeys = PassKey.get_by_user_id(db=None, user_id=current_user.id)
|
||||
existing_credentials = (
|
||||
@@ -259,8 +234,14 @@ def passkey_register_start(
|
||||
existing_credentials=existing_credentials,
|
||||
)
|
||||
|
||||
transaction_token = PasskeyChallengeStore.issue(
|
||||
challenge=challenge,
|
||||
purpose="registration",
|
||||
user_id=current_user.id,
|
||||
)
|
||||
return schemas.Response(
|
||||
success=True, data={"options": options_json, "challenge": challenge}
|
||||
success=True,
|
||||
data={"options": options_json, "transaction_token": transaction_token},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"生成PassKey注册选项失败: {e}")
|
||||
@@ -278,11 +259,21 @@ def passkey_register_finish(
|
||||
) -> Any:
|
||||
"""完成注册 PassKey - 验证并保存凭证"""
|
||||
try:
|
||||
challenge_state = PasskeyChallengeStore.consume(
|
||||
transaction_token=passkey_req.transaction_token,
|
||||
purpose="registration",
|
||||
)
|
||||
if not challenge_state or challenge_state.user_id != current_user.id:
|
||||
return schemas.Response(
|
||||
success=False,
|
||||
message="注册请求已失效,请重新发起注册",
|
||||
)
|
||||
|
||||
# 验证注册响应
|
||||
credential_id, public_key, sign_count, aaguid = (
|
||||
PassKeyHelper.verify_registration_response(
|
||||
credential=passkey_req.credential,
|
||||
expected_challenge=passkey_req.challenge,
|
||||
expected_challenge=challenge_state.challenge,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -309,9 +300,19 @@ def passkey_register_finish(
|
||||
logger.info(f"用户 {current_user.name} 成功注册PassKey: {passkey_req.name}")
|
||||
|
||||
return schemas.Response(success=True, message="通行密钥注册成功")
|
||||
except PassKeyRegistrationOriginMismatchError:
|
||||
return schemas.Response(
|
||||
success=False,
|
||||
message="访问域名与系统配置不一致,请使用配置的域名重试",
|
||||
)
|
||||
except PassKeyRegistrationVerificationError:
|
||||
return schemas.Response(
|
||||
success=False,
|
||||
message="通行密钥注册验证失败,请重新发起注册后重试",
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"注册PassKey失败: {e}")
|
||||
return schemas.Response(success=False, message=f"注册失败: {str(e)}")
|
||||
return schemas.Response(success=False, message="通行密钥注册失败,请稍后重试")
|
||||
|
||||
|
||||
@router.post(
|
||||
@@ -325,6 +326,7 @@ def passkey_authenticate_start(
|
||||
"""开始 PassKey 认证 - 生成认证选项"""
|
||||
try:
|
||||
existing_credentials = None
|
||||
user_id = None
|
||||
|
||||
# 如果指定了用户名,只允许该用户的PassKey
|
||||
if passkey_req.username:
|
||||
@@ -337,14 +339,21 @@ def passkey_authenticate_start(
|
||||
return schemas.Response(success=False, message="认证失败")
|
||||
|
||||
existing_credentials = _build_credential_list(existing_passkeys)
|
||||
user_id = user.id
|
||||
|
||||
# 生成认证选项
|
||||
options_json, challenge = PassKeyHelper.generate_authentication_options(
|
||||
existing_credentials=existing_credentials
|
||||
)
|
||||
|
||||
transaction_token = PasskeyChallengeStore.issue(
|
||||
challenge=challenge,
|
||||
purpose="authentication",
|
||||
user_id=user_id,
|
||||
)
|
||||
return schemas.Response(
|
||||
success=True, data={"options": options_json, "challenge": challenge}
|
||||
success=True,
|
||||
data={"options": options_json, "transaction_token": transaction_token},
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"生成PassKey认证选项失败: {e}")
|
||||
@@ -361,6 +370,13 @@ def passkey_authenticate_finish(
|
||||
) -> Any:
|
||||
"""完成 PassKey 认证 - 验证凭证并返回 token"""
|
||||
try:
|
||||
challenge_state = PasskeyChallengeStore.consume(
|
||||
transaction_token=passkey_req.transaction_token,
|
||||
purpose="authentication",
|
||||
)
|
||||
if not challenge_state:
|
||||
raise HTTPException(status_code=401, detail="认证请求已失效")
|
||||
|
||||
# 提取并标准化凭证ID
|
||||
try:
|
||||
credential_id = _extract_and_standardize_credential_id(
|
||||
@@ -375,11 +391,13 @@ def passkey_authenticate_finish(
|
||||
user = User.get_by_id(db=None, user_id=passkey.user_id) if passkey else None
|
||||
if not passkey or not user or not user.is_active:
|
||||
raise HTTPException(status_code=401, detail="认证失败")
|
||||
if challenge_state.user_id is not None and challenge_state.user_id != user.id:
|
||||
raise HTTPException(status_code=401, detail="认证失败")
|
||||
|
||||
# 验证认证响应并更新
|
||||
success, _ = _verify_passkey_and_update(
|
||||
credential=passkey_req.credential,
|
||||
challenge=passkey_req.challenge,
|
||||
challenge=challenge_state.challenge,
|
||||
passkey=passkey,
|
||||
)
|
||||
|
||||
@@ -493,46 +511,3 @@ async def passkey_delete(
|
||||
except Exception as e:
|
||||
logger.error(f"删除PassKey失败: {e}")
|
||||
return schemas.Response(success=False, message=f"删除失败: {str(e)}")
|
||||
|
||||
|
||||
@router.post(
|
||||
"/passkey/verify", summary="PassKey 二次验证", response_model=schemas.Response
|
||||
)
|
||||
def passkey_verify_mfa(
|
||||
passkey_req: PassKeyAuthenticationFinish,
|
||||
current_user: Annotated[User, Depends(get_current_active_user)],
|
||||
) -> Any:
|
||||
"""使用 PassKey 进行二次验证(MFA)"""
|
||||
try:
|
||||
# 提取并标准化凭证ID
|
||||
try:
|
||||
credential_id = _extract_and_standardize_credential_id(
|
||||
passkey_req.credential
|
||||
)
|
||||
except ValueError as e:
|
||||
logger.warning(f"PassKey二次验证失败,提供的凭证无效: {e}")
|
||||
return schemas.Response(success=False, message="验证失败")
|
||||
|
||||
# 查找PassKey(必须属于当前用户)
|
||||
passkey = PassKey.get_by_credential_id(db=None, credential_id=credential_id)
|
||||
if not passkey or passkey.user_id != current_user.id:
|
||||
return schemas.Response(
|
||||
success=False, message="通行密钥不存在或不属于当前用户"
|
||||
)
|
||||
|
||||
# 验证认证响应并更新
|
||||
success, _ = _verify_passkey_and_update(
|
||||
credential=passkey_req.credential,
|
||||
challenge=passkey_req.challenge,
|
||||
passkey=passkey,
|
||||
)
|
||||
|
||||
if not success:
|
||||
return schemas.Response(success=False, message="通行密钥验证失败")
|
||||
|
||||
logger.info(f"用户 {current_user.name} 通过PassKey二次验证成功")
|
||||
|
||||
return schemas.Response(success=True, message="二次验证成功")
|
||||
except Exception as e:
|
||||
logger.error(f"PassKey二次验证失败: {e}")
|
||||
return schemas.Response(success=False, message="验证失败")
|
||||
|
||||
@@ -231,6 +231,9 @@ def _error_response(
|
||||
def _check_auth(
|
||||
credentials: Optional[HTTPAuthorizationCredentials],
|
||||
) -> Optional[JSONResponse]:
|
||||
"""
|
||||
OpenAI 兼容接口以 API_TOKEN 认证受信客户端,认证通过即按管理员级 Agent 集成处理。
|
||||
"""
|
||||
if not credentials or credentials.scheme.lower() != "bearer":
|
||||
return _error_response(
|
||||
"Invalid bearer token.",
|
||||
@@ -317,6 +320,7 @@ async def chat_completions(
|
||||
|
||||
session_id = build_session_id(session_key, SESSION_PREFIX)
|
||||
username = str(payload.user or "openai-client")
|
||||
# 兼容接口的 API_TOKEN 客户端按管理员级 MoviePilot Agent 集成处理。
|
||||
agent = _CollectingMoviePilotAgent(
|
||||
session_id=session_id,
|
||||
user_id=session_key,
|
||||
@@ -409,6 +413,7 @@ async def responses(
|
||||
|
||||
session_key = str(payload.user or uuid.uuid4())
|
||||
session_id = build_session_id(session_key, SESSION_PREFIX)
|
||||
# 兼容接口的 API_TOKEN 客户端按管理员级 MoviePilot Agent 集成处理。
|
||||
agent = _CollectingMoviePilotAgent(
|
||||
session_id=session_id,
|
||||
user_id=session_key,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import asyncio
|
||||
import mimetypes
|
||||
import shutil
|
||||
from typing import Annotated, Any, List, Optional
|
||||
from typing import Annotated, Any, Dict, List, Optional
|
||||
|
||||
import aiofiles
|
||||
from anyio import Path as AsyncPath
|
||||
@@ -10,6 +11,7 @@ from starlette import status
|
||||
from starlette.responses import StreamingResponse
|
||||
|
||||
from app import schemas
|
||||
from app.api.apiv2_utils import API_V2_STR, OPENAPI_V2_PATH
|
||||
from app.command import Command
|
||||
from app.core.cache import async_fresh
|
||||
from app.core.config import settings
|
||||
@@ -35,10 +37,78 @@ from app.scheduler import Scheduler
|
||||
from app.schemas.event import PluginDataResetEventData
|
||||
from app.schemas.types import ChainEventType, SystemConfigKey
|
||||
|
||||
PROTECTED_ROUTES = {"/api/v1/openapi.json", "/docs", "/docs/oauth2-redirect", "/redoc"}
|
||||
PROTECTED_ROUTES = {
|
||||
"/api/v1/openapi.json",
|
||||
OPENAPI_V2_PATH,
|
||||
"/docs",
|
||||
"/docs/oauth2-redirect",
|
||||
"/redoc",
|
||||
}
|
||||
PLUGIN_PREFIX = f"{settings.API_V1_STR}/plugin"
|
||||
PLUGIN_V2_PREFIX = f"{API_V2_STR}/plugin"
|
||||
|
||||
router = APIRouter()
|
||||
_plugin_release_refresh_tasks: set[asyncio.Task] = set()
|
||||
|
||||
|
||||
async def _get_market_plugin_from_repo(
|
||||
plugin_manager: PluginManager,
|
||||
plugin_id: str,
|
||||
repo_url: str,
|
||||
force: bool,
|
||||
) -> Optional[schemas.Plugin]:
|
||||
"""
|
||||
只读取指定插件仓库的市场元数据,避免单插件详情触发全部市场刷新。
|
||||
"""
|
||||
market_plugins = await plugin_manager.async_get_plugins_from_market(
|
||||
repo_url, settings.VERSION_FLAG, force
|
||||
)
|
||||
market_plugin = next(
|
||||
(
|
||||
plugin
|
||||
for plugin in market_plugins or []
|
||||
if plugin.id == plugin_id
|
||||
),
|
||||
None,
|
||||
)
|
||||
if market_plugin or not settings.VERSION_FLAG:
|
||||
return market_plugin
|
||||
|
||||
compatible_plugins = await plugin_manager.async_get_plugins_from_market(
|
||||
repo_url, None, force
|
||||
)
|
||||
return next(
|
||||
(
|
||||
plugin
|
||||
for plugin in compatible_plugins or []
|
||||
if plugin.id == plugin_id
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
async def _refresh_plugin_release_versions(plugin_id: str, repo_url: str) -> None:
|
||||
"""
|
||||
后台强制刷新 Release 缓存,接口响应路径优先返回已有缓存。
|
||||
"""
|
||||
try:
|
||||
async with async_fresh(True):
|
||||
await PluginHelper().async_get_plugin_release_versions(plugin_id, repo_url)
|
||||
except Exception as e:
|
||||
logger.warning(f"后台刷新插件 {plugin_id} Release 列表失败:{e}")
|
||||
|
||||
|
||||
def _schedule_plugin_release_refresh(plugin_id: str, repo_url: str) -> None:
|
||||
"""
|
||||
保留后台任务引用,避免任务被回收,同时让 helper 负责同仓库强刷合并。
|
||||
"""
|
||||
task = asyncio.create_task(_refresh_plugin_release_versions(plugin_id, repo_url))
|
||||
_plugin_release_refresh_tasks.add(task)
|
||||
|
||||
def _discard_task(completed_task: asyncio.Task) -> None:
|
||||
_plugin_release_refresh_tasks.discard(completed_task)
|
||||
|
||||
task.add_done_callback(_discard_task)
|
||||
|
||||
|
||||
def register_plugin_api(plugin_id: Optional[str] = None):
|
||||
@@ -96,8 +166,11 @@ def _update_plugin_api_routes(plugin_id: Optional[str], action: str):
|
||||
elif Depends(verify_apikey) not in dependencies:
|
||||
dependencies.append(Depends(verify_apikey))
|
||||
app.add_api_route(**api, tags=["plugin"])
|
||||
v2_api = api.copy()
|
||||
v2_api["path"] = api_path.replace(PLUGIN_PREFIX, PLUGIN_V2_PREFIX, 1)
|
||||
app.add_api_route(**v2_api, tags=["plugin"])
|
||||
is_modified = True
|
||||
logger.debug(f"Added plugin route: {api_path}")
|
||||
logger.debug(f"Added plugin routes: {api_path}, {v2_api['path']}")
|
||||
except Exception as e:
|
||||
logger.error(f"Error adding plugin route {api_path}: {str(e)}")
|
||||
|
||||
@@ -115,8 +188,13 @@ def _remove_routes(plugin_id: str) -> bool:
|
||||
"""
|
||||
if not plugin_id:
|
||||
return False
|
||||
prefix = f"{PLUGIN_PREFIX}/{plugin_id}/"
|
||||
routes_to_remove = [route for route in app.routes if route.path.startswith(prefix)]
|
||||
prefixes = {
|
||||
f"{PLUGIN_PREFIX}/{plugin_id}/",
|
||||
f"{PLUGIN_V2_PREFIX}/{plugin_id}/",
|
||||
}
|
||||
routes_to_remove = [
|
||||
route for route in app.routes if any(route.path.startswith(prefix) for prefix in prefixes)
|
||||
]
|
||||
removed = False
|
||||
for route in routes_to_remove:
|
||||
try:
|
||||
@@ -239,6 +317,15 @@ async def _get_plugin_history_detail(
|
||||
if local_repo_plugin:
|
||||
return _merge_plugin_market_metadata(installed_plugin, local_repo_plugin)
|
||||
|
||||
if installed_plugin.repo_url:
|
||||
market_plugin = await _get_market_plugin_from_repo(
|
||||
plugin_manager, plugin_id, installed_plugin.repo_url, force
|
||||
)
|
||||
if not market_plugin:
|
||||
logger.debug(f"插件 {plugin_id} 未从来源仓库获取到更新说明,返回本地插件信息")
|
||||
return installed_plugin
|
||||
return _merge_plugin_market_metadata(installed_plugin, market_plugin)
|
||||
|
||||
market_plugin = next(
|
||||
(
|
||||
plugin
|
||||
@@ -359,30 +446,9 @@ async def plugin_releases(
|
||||
}
|
||||
|
||||
plugin_manager = PluginManager()
|
||||
market_plugins = await plugin_manager.async_get_plugins_from_market(
|
||||
repo_url, settings.VERSION_FLAG, force
|
||||
market_plugin = await _get_market_plugin_from_repo(
|
||||
plugin_manager, plugin_id, repo_url, force
|
||||
)
|
||||
market_plugin = next(
|
||||
(
|
||||
plugin
|
||||
for plugin in market_plugins or []
|
||||
if plugin.id == plugin_id
|
||||
),
|
||||
None,
|
||||
)
|
||||
if not market_plugin and settings.VERSION_FLAG:
|
||||
compatible_plugins = await plugin_manager.async_get_plugins_from_market(
|
||||
repo_url, None, force
|
||||
)
|
||||
market_plugin = next(
|
||||
(
|
||||
plugin
|
||||
for plugin in compatible_plugins or []
|
||||
if plugin.id == plugin_id
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
latest_version = market_plugin.plugin_version if market_plugin else None
|
||||
current_version = plugin_manager.get_local_plugin_version(plugin_id)
|
||||
if not getattr(market_plugin, "release", False):
|
||||
@@ -393,8 +459,15 @@ async def plugin_releases(
|
||||
"items": [],
|
||||
}
|
||||
|
||||
async with async_fresh(force):
|
||||
release_items = await PluginHelper().async_get_plugin_release_versions(plugin_id, repo_url)
|
||||
plugin_helper = PluginHelper()
|
||||
has_release_cache = (
|
||||
await plugin_helper.async_has_plugin_release_cache(repo_url)
|
||||
if force
|
||||
else False
|
||||
)
|
||||
release_items = await plugin_helper.async_get_plugin_release_versions(plugin_id, repo_url)
|
||||
if force and has_release_cache:
|
||||
_schedule_plugin_release_refresh(plugin_id, repo_url)
|
||||
items = []
|
||||
for item in release_items:
|
||||
version = item.get("version")
|
||||
@@ -419,6 +492,71 @@ async def statistic(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
|
||||
return await MoviePilotServerHelper.async_get_plugin_statistic()
|
||||
|
||||
|
||||
@router.get(
|
||||
"/rating",
|
||||
summary="批量查询插件评分",
|
||||
response_model=Dict[str, schemas.PluginRating],
|
||||
)
|
||||
async def plugin_ratings(
|
||||
plugin_ids: Optional[str] = None,
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
) -> Dict[str, schemas.PluginRating]:
|
||||
"""
|
||||
批量查询插件平均分、评分人数和当前安装实例评分。
|
||||
"""
|
||||
requested_ids = plugin_ids.split(",") if plugin_ids is not None else None
|
||||
ratings = await MoviePilotServerHelper.async_get_plugin_ratings(requested_ids)
|
||||
return {
|
||||
plugin_id: schemas.PluginRating.model_validate(rating)
|
||||
for plugin_id, rating in ratings.items()
|
||||
}
|
||||
|
||||
|
||||
@router.get(
|
||||
"/rating/{plugin_id}",
|
||||
summary="查询插件评分",
|
||||
response_model=schemas.PluginRating,
|
||||
)
|
||||
async def plugin_rating(
|
||||
plugin_id: str,
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
) -> schemas.PluginRating:
|
||||
"""
|
||||
查询单个插件平均分、评分人数和当前安装实例评分。
|
||||
"""
|
||||
rating = await MoviePilotServerHelper.async_get_plugin_rating(plugin_id)
|
||||
return schemas.PluginRating.model_validate(rating)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/rating/{plugin_id}",
|
||||
summary="提交插件评分",
|
||||
response_model=schemas.Response,
|
||||
)
|
||||
async def rate_plugin(
|
||||
plugin_id: str,
|
||||
payload: schemas.PluginRatingRequest,
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
) -> schemas.Response:
|
||||
"""
|
||||
为已安装插件新增或更新当前安装实例评分。
|
||||
"""
|
||||
installed_plugins = SystemConfigOper().get(SystemConfigKey.UserInstalledPlugins) or []
|
||||
if plugin_id not in installed_plugins:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"插件 {plugin_id} 未安装,无法评分",
|
||||
)
|
||||
|
||||
rating = await MoviePilotServerHelper.async_submit_plugin_rating(
|
||||
plugin_id,
|
||||
payload.rating,
|
||||
)
|
||||
if rating is None:
|
||||
return schemas.Response(success=False, message="连接MoviePilot服务器失败")
|
||||
return schemas.Response(success=True, data=rating)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/reload/{plugin_id}", summary="重新加载插件", response_model=schemas.Response
|
||||
)
|
||||
@@ -615,9 +753,9 @@ def reset_plugin(
|
||||
# 事件处理器需要运行中插件完成补偿;补偿后先停止插件,避免删除数据时仍有任务读写旧状态。
|
||||
plugin_manager.stop(plugin_id)
|
||||
# 删除配置
|
||||
plugin_manager.delete_plugin_config(plugin_id)
|
||||
plugin_manager.delete_plugin_config(plugin_id, force=True)
|
||||
# 删除插件所有数据
|
||||
plugin_manager.delete_plugin_data(plugin_id)
|
||||
plugin_manager.delete_plugin_data(plugin_id, force=True)
|
||||
# 重新加载插件
|
||||
reload_plugin(plugin_id)
|
||||
return schemas.Response(success=True)
|
||||
|
||||
@@ -1,17 +1,29 @@
|
||||
from typing import Any, List, Optional
|
||||
from typing import Any, Awaitable, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
|
||||
from app import schemas
|
||||
from app.chain.recommend import RecommendChain
|
||||
from app.core.event import eventmanager
|
||||
from app.core.security import verify_token
|
||||
from app.modules.themoviedb.tmdbv3api.exceptions import TMDbException
|
||||
from app.schemas import RecommendSourceEventData
|
||||
from app.schemas.types import ChainEventType
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
async def _require_tmdb_result(operation: Awaitable[List[Any]]) -> List[Any]:
|
||||
"""保留 TMDB 成功空列表,并把上游请求异常转换为明确的网关错误。"""
|
||||
try:
|
||||
return await operation
|
||||
except TMDbException as error:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail="TMDB请求失败",
|
||||
) from error
|
||||
|
||||
|
||||
@router.get(
|
||||
"/source",
|
||||
summary="获取推荐数据源",
|
||||
@@ -204,16 +216,19 @@ async def tmdb_movies(
|
||||
"""
|
||||
浏览TMDB电影信息
|
||||
"""
|
||||
return await RecommendChain().async_tmdb_movies(
|
||||
sort_by=sort_by,
|
||||
with_genres=with_genres,
|
||||
with_original_language=with_original_language,
|
||||
with_keywords=with_keywords,
|
||||
with_watch_providers=with_watch_providers,
|
||||
vote_average=vote_average,
|
||||
vote_count=vote_count,
|
||||
release_date=release_date,
|
||||
page=page,
|
||||
return await _require_tmdb_result(
|
||||
RecommendChain().async_tmdb_movies(
|
||||
sort_by=sort_by,
|
||||
with_genres=with_genres,
|
||||
with_original_language=with_original_language,
|
||||
with_keywords=with_keywords,
|
||||
with_watch_providers=with_watch_providers,
|
||||
vote_average=vote_average,
|
||||
vote_count=vote_count,
|
||||
release_date=release_date,
|
||||
page=page,
|
||||
raise_exception=True,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -233,16 +248,19 @@ async def tmdb_tvs(
|
||||
"""
|
||||
浏览TMDB剧集信息
|
||||
"""
|
||||
return await RecommendChain().async_tmdb_tvs(
|
||||
sort_by=sort_by,
|
||||
with_genres=with_genres,
|
||||
with_original_language=with_original_language,
|
||||
with_keywords=with_keywords,
|
||||
with_watch_providers=with_watch_providers,
|
||||
vote_average=vote_average,
|
||||
vote_count=vote_count,
|
||||
release_date=release_date,
|
||||
page=page,
|
||||
return await _require_tmdb_result(
|
||||
RecommendChain().async_tmdb_tvs(
|
||||
sort_by=sort_by,
|
||||
with_genres=with_genres,
|
||||
with_original_language=with_original_language,
|
||||
with_keywords=with_keywords,
|
||||
with_watch_providers=with_watch_providers,
|
||||
vote_average=vote_average,
|
||||
vote_count=vote_count,
|
||||
release_date=release_date,
|
||||
page=page,
|
||||
raise_exception=True,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -255,4 +273,6 @@ async def tmdb_trending(
|
||||
"""
|
||||
TMDB流行趋势
|
||||
"""
|
||||
return await RecommendChain().async_tmdb_trending(page=page)
|
||||
return await _require_tmdb_result(
|
||||
RecommendChain().async_tmdb_trending(page=page, raise_exception=True)
|
||||
)
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
import asyncio
|
||||
import json
|
||||
from typing import List, Any, Optional, AsyncIterator
|
||||
import time
|
||||
from typing import Any, AsyncIterator, Iterator, List, Optional
|
||||
from uuid import uuid4
|
||||
|
||||
from fastapi import APIRouter, Depends, Body, Request
|
||||
from fastapi.responses import StreamingResponse
|
||||
@@ -12,14 +14,23 @@ from app.core.config import settings
|
||||
from app.core.event import eventmanager
|
||||
from app.core.metainfo import MetaInfo
|
||||
from app.core.security import verify_resource_token, verify_token
|
||||
from app.helper.locale import LocaleHelper
|
||||
from app.log import logger
|
||||
from app.schemas import MediaRecognizeConvertEventData
|
||||
from app.schemas.types import MediaType, ChainEventType
|
||||
from app.utils.media import parse_media_key, resolve_media_identity
|
||||
from app.utils.security import SecurityUtils
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
_SSE_APPEND_FLUSH_INTERVAL = 1
|
||||
_SSE_APPEND_MAX_ITEMS = 48
|
||||
_SSE_HEARTBEAT_INTERVAL = 15
|
||||
_SSE_REPLACE_MAX_ITEMS = 48
|
||||
_SSE_RESPONSE_HEADERS = {
|
||||
"Cache-Control": "no-cache",
|
||||
"X-Accel-Buffering": "no",
|
||||
}
|
||||
|
||||
|
||||
def _parse_site_list(sites: Optional[str]) -> Optional[List[int]]:
|
||||
@@ -38,11 +49,126 @@ def _parse_media_type(mtype: Optional[str]) -> Optional[MediaType]:
|
||||
return MediaType.from_agent(mtype) or MediaType(mtype)
|
||||
|
||||
|
||||
def _sse_event(data: dict) -> str:
|
||||
def _resolve_media_season(
|
||||
explicit_season: Optional[int],
|
||||
recognized_season: Optional[int],
|
||||
) -> Optional[int]:
|
||||
"""合并显式季号与识别结果,显式值优先且季 0 属于有效业务值。"""
|
||||
return explicit_season if explicit_season is not None else recognized_season
|
||||
|
||||
|
||||
async def _resolve_media_search_params(
|
||||
mediaid: str,
|
||||
media_type: Optional[MediaType] = None,
|
||||
title: Optional[str] = None,
|
||||
year: Optional[str] = None,
|
||||
media_season: Optional[int] = None,
|
||||
) -> tuple[Optional[dict], str]:
|
||||
"""将任意来源媒体键解析为 SearchChain 可直接使用的识别参数。"""
|
||||
source, source_media_id = parse_media_key(mediaid)
|
||||
if source and source_media_id:
|
||||
if source in {"themoviedb", "bangumi", "anilist"} \
|
||||
and not source_media_id.isdigit():
|
||||
return None, "媒体ID格式错误"
|
||||
return {"source": source, "mediaid": source_media_id}, ""
|
||||
|
||||
event_data = MediaRecognizeConvertEventData(
|
||||
mediaid=mediaid, convert_type=settings.RECOGNIZE_SOURCE
|
||||
)
|
||||
event = await eventmanager.async_send_event(
|
||||
ChainEventType.MediaRecognizeConvert, event_data
|
||||
)
|
||||
if event and event.event_data and event.event_data.media_dict:
|
||||
event_data = event.event_data
|
||||
search_id = event_data.media_dict.get("id")
|
||||
if search_id is not None:
|
||||
return {
|
||||
"source": event_data.convert_type,
|
||||
"mediaid": str(search_id),
|
||||
}, ""
|
||||
|
||||
if not title:
|
||||
return None, "未知的媒体ID"
|
||||
|
||||
meta = MetaInfo(title)
|
||||
if year:
|
||||
meta.year = year
|
||||
if media_type:
|
||||
meta.type = media_type
|
||||
if media_season is not None:
|
||||
meta.type = MediaType.TV
|
||||
meta.begin_season = media_season
|
||||
mediainfo = await MediaChain().async_recognize_by_meta(
|
||||
meta,
|
||||
obtain_images=False,
|
||||
)
|
||||
if not mediainfo:
|
||||
return None, "未识别到媒体信息"
|
||||
source, source_media_id = resolve_media_identity(media=mediainfo)
|
||||
if not source or not source_media_id:
|
||||
return None, "媒体信息缺少有效ID"
|
||||
return {"source": source, "mediaid": source_media_id}, ""
|
||||
|
||||
|
||||
def _sse_event(data: dict, locale: Optional[str] = None) -> str:
|
||||
"""
|
||||
转换为SSE事件
|
||||
"""
|
||||
return f"data: {json.dumps(data, ensure_ascii=False)}\n\n"
|
||||
payload = data
|
||||
message = payload.get("message")
|
||||
text = payload.get("text")
|
||||
if isinstance(message, str) or isinstance(text, str):
|
||||
payload = data.copy()
|
||||
if isinstance(message, str):
|
||||
payload["message_i18n"] = LocaleHelper.translate_text(
|
||||
message, locale=locale
|
||||
)
|
||||
if isinstance(text, str):
|
||||
payload["text_i18n"] = LocaleHelper.translate_text(text, locale=locale)
|
||||
return f"data: {json.dumps(payload, ensure_ascii=False)}\n\n"
|
||||
|
||||
|
||||
def _serialize_signed_subtitle_result(subtitle: Any) -> dict:
|
||||
"""
|
||||
序列化字幕结果并签名下载链接,签名用途绑定站点 ID。
|
||||
"""
|
||||
data = subtitle.to_dict() if hasattr(subtitle, "to_dict") else dict(subtitle)
|
||||
enclosure = data.get("enclosure")
|
||||
if enclosure:
|
||||
data["enclosure"] = SecurityUtils.sign_url(
|
||||
enclosure,
|
||||
purpose=SecurityUtils.subtitle_download_purpose(data.get("site")),
|
||||
)
|
||||
return data
|
||||
|
||||
|
||||
def _serialize_signed_subtitle_results(subtitles: List[Any]) -> List[dict]:
|
||||
"""
|
||||
批量序列化字幕结果,确保返回给客户端的下载链接均已签名。
|
||||
"""
|
||||
return [_serialize_signed_subtitle_result(subtitle) for subtitle in subtitles]
|
||||
|
||||
|
||||
def _sign_subtitle_search_event(event: dict) -> dict:
|
||||
"""
|
||||
签名字幕搜索流事件中的下载链接。
|
||||
"""
|
||||
signed_event = dict(event)
|
||||
if "items" in signed_event:
|
||||
signed_event["items"] = _serialize_signed_subtitle_results(
|
||||
signed_event.get("items") or []
|
||||
)
|
||||
return signed_event
|
||||
|
||||
|
||||
async def _iter_signed_subtitle_search_events(
|
||||
event_source: AsyncIterator[dict],
|
||||
) -> AsyncIterator[dict]:
|
||||
"""
|
||||
输出仅包含签名字幕下载链接的搜索流事件。
|
||||
"""
|
||||
async for event in event_source:
|
||||
yield _sign_subtitle_search_event(event)
|
||||
|
||||
|
||||
def _merge_append_event(pending_event: Optional[dict], event: dict) -> dict:
|
||||
@@ -62,11 +188,40 @@ def _merge_append_event(pending_event: Optional[dict], event: dict) -> dict:
|
||||
return merged_event
|
||||
|
||||
|
||||
def _iter_replace_event_batches(event: dict) -> Iterator[dict]:
|
||||
"""
|
||||
将超大的最终替换事件拆成有序批次,避免单个 SSE 消息承载全部完整对象。
|
||||
"""
|
||||
items = event.get("items")
|
||||
if (
|
||||
event.get("type") != "replace"
|
||||
or not isinstance(items, list)
|
||||
or len(items) <= _SSE_REPLACE_MAX_ITEMS
|
||||
):
|
||||
yield event
|
||||
return
|
||||
|
||||
batch_count = (len(items) + _SSE_REPLACE_MAX_ITEMS - 1) // _SSE_REPLACE_MAX_ITEMS
|
||||
for batch_index in range(batch_count):
|
||||
start = batch_index * _SSE_REPLACE_MAX_ITEMS
|
||||
batch_event = dict(event)
|
||||
batch_event.update(
|
||||
{
|
||||
"type": "replace" if batch_index == 0 else "append",
|
||||
"items": items[start:start + _SSE_REPLACE_MAX_ITEMS],
|
||||
"replace_batch": True,
|
||||
"batch_index": batch_index,
|
||||
"batch_count": batch_count,
|
||||
}
|
||||
)
|
||||
yield batch_event
|
||||
|
||||
|
||||
async def _iter_batched_search_events(
|
||||
event_source: AsyncIterator[dict],
|
||||
) -> AsyncIterator[dict]:
|
||||
"""
|
||||
对搜索流事件做轻量批处理,避免站点结果集中返回时产生过密 SSE。
|
||||
对搜索流事件做轻量批处理,并在上游长时间静默时发送心跳。
|
||||
"""
|
||||
iterator = event_source.__aiter__()
|
||||
pending_append_event: Optional[dict] = None
|
||||
@@ -77,13 +232,19 @@ async def _iter_batched_search_events(
|
||||
if next_event_task is None:
|
||||
next_event_task = asyncio.create_task(anext(iterator))
|
||||
|
||||
timeout = _SSE_APPEND_FLUSH_INTERVAL if pending_append_event else None
|
||||
timeout = (
|
||||
_SSE_APPEND_FLUSH_INTERVAL
|
||||
if pending_append_event
|
||||
else _SSE_HEARTBEAT_INTERVAL
|
||||
)
|
||||
done, _ = await asyncio.wait({next_event_task}, timeout=timeout)
|
||||
|
||||
if not done:
|
||||
if pending_append_event:
|
||||
yield pending_append_event
|
||||
pending_append_event = None
|
||||
else:
|
||||
yield {"type": "heartbeat"}
|
||||
continue
|
||||
|
||||
try:
|
||||
@@ -109,7 +270,8 @@ async def _iter_batched_search_events(
|
||||
yield pending_append_event
|
||||
pending_append_event = None
|
||||
|
||||
yield event
|
||||
for batched_event in _iter_replace_event_batches(event):
|
||||
yield batched_event
|
||||
finally:
|
||||
if next_event_task and not next_event_task.done():
|
||||
next_event_task.cancel()
|
||||
@@ -121,12 +283,29 @@ async def _iter_batched_search_events(
|
||||
|
||||
async def _stream_search_events(request: Request, event_source: AsyncIterator[dict]):
|
||||
"""
|
||||
输出搜索SSE事件
|
||||
输出搜索 SSE 事件,并记录连接生命周期与传输规模。
|
||||
"""
|
||||
locale = LocaleHelper.get_locale_from_request(request)
|
||||
search_id = uuid4().hex[:12]
|
||||
request_path = getattr(getattr(request, "url", None), "path", "unknown")
|
||||
started_at = time.monotonic()
|
||||
event_count = 0
|
||||
transmitted_bytes = 0
|
||||
last_event_type = "none"
|
||||
last_stage = "none"
|
||||
termination_reason = "source_exhausted"
|
||||
logger.info(f"渐进式搜索流已建立,搜索ID:{search_id},路径:{request_path}")
|
||||
try:
|
||||
has_sent_final_replace = False
|
||||
async for event in _iter_batched_search_events(event_source):
|
||||
last_event_type = event.get("type") or "unknown"
|
||||
last_stage = event.get("stage") or last_stage
|
||||
if await request.is_disconnected():
|
||||
termination_reason = "client_disconnected"
|
||||
logger.warning(
|
||||
f"渐进式搜索客户端已断开,搜索ID:{search_id},路径:{request_path},"
|
||||
f"事件:{last_event_type},阶段:{last_stage}"
|
||||
)
|
||||
break
|
||||
# 精确搜索会先发送 replace,再发送 done。done 再带整包 items 只会重复占用带宽和前端内存。
|
||||
if event.get("type") == "replace" and event.get("items"):
|
||||
@@ -138,10 +317,36 @@ async def _stream_search_events(request: Request, event_source: AsyncIterator[di
|
||||
and event.get("items")
|
||||
):
|
||||
event = {key: value for key, value in event.items() if key != "items"}
|
||||
yield _sse_event(event)
|
||||
payload = _sse_event(event, locale=locale)
|
||||
event_count += 1
|
||||
transmitted_bytes += len(payload.encode("utf-8"))
|
||||
if event.get("type") == "done":
|
||||
termination_reason = "completed"
|
||||
yield payload
|
||||
except asyncio.CancelledError:
|
||||
termination_reason = "cancelled"
|
||||
logger.warning(
|
||||
f"渐进式搜索流已取消,搜索ID:{search_id},路径:{request_path},"
|
||||
f"事件:{last_event_type},阶段:{last_stage}"
|
||||
)
|
||||
raise
|
||||
except Exception as err:
|
||||
termination_reason = "error"
|
||||
logger.error(f"渐进式搜索出错:{err}", exc_info=True)
|
||||
yield _sse_event({"type": "error", "success": False, "message": str(err)})
|
||||
payload = _sse_event(
|
||||
{"type": "error", "success": False, "message": str(err)},
|
||||
locale=locale,
|
||||
)
|
||||
event_count += 1
|
||||
transmitted_bytes += len(payload.encode("utf-8"))
|
||||
yield payload
|
||||
finally:
|
||||
elapsed = time.monotonic() - started_at
|
||||
logger.info(
|
||||
f"渐进式搜索流结束,搜索ID:{search_id},路径:{request_path},"
|
||||
f"状态:{termination_reason},事件数:{event_count},"
|
||||
f"发送字节:{transmitted_bytes},耗时:{elapsed:.2f}秒"
|
||||
)
|
||||
|
||||
|
||||
@router.get("/last", summary="查询搜索结果", response_model=List[schemas.Context])
|
||||
@@ -168,7 +373,9 @@ async def search_latest_context(_: schemas.TokenPayload = Depends(verify_token))
|
||||
success=True,
|
||||
data={
|
||||
"params": params,
|
||||
"results": [result.to_dict() for result in results],
|
||||
"results": _serialize_signed_subtitle_results(results)
|
||||
if params.get("result_type") == "subtitle"
|
||||
else [result.to_dict() for result in results],
|
||||
},
|
||||
)
|
||||
|
||||
@@ -192,192 +399,35 @@ async def search_by_id_stream(
|
||||
media_type = _parse_media_type(mtype)
|
||||
media_season = int(season) if season else None
|
||||
site_list = _parse_site_list(sites)
|
||||
media_chain = MediaChain()
|
||||
search_chain = SearchChain()
|
||||
|
||||
async def event_source():
|
||||
nonlocal media_season
|
||||
torrents = None
|
||||
if mediaid.startswith("tmdb:"):
|
||||
tmdbid = int(mediaid.replace("tmdb:", ""))
|
||||
if settings.RECOGNIZE_SOURCE == "douban":
|
||||
doubaninfo = await media_chain.async_get_doubaninfo_by_tmdbid(
|
||||
tmdbid=tmdbid, mtype=media_type
|
||||
)
|
||||
if doubaninfo:
|
||||
torrents = search_chain.async_search_by_id_stream(
|
||||
doubanid=doubaninfo.get("id"),
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
else:
|
||||
yield {
|
||||
"type": "error",
|
||||
"success": False,
|
||||
"message": "未识别到豆瓣媒体信息",
|
||||
}
|
||||
return
|
||||
else:
|
||||
torrents = search_chain.async_search_by_id_stream(
|
||||
tmdbid=tmdbid,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
elif mediaid.startswith("douban:"):
|
||||
doubanid = mediaid.replace("douban:", "")
|
||||
if settings.RECOGNIZE_SOURCE == "themoviedb":
|
||||
tmdbinfo = await media_chain.async_get_tmdbinfo_by_doubanid(
|
||||
doubanid=doubanid, mtype=media_type
|
||||
)
|
||||
if tmdbinfo:
|
||||
if tmdbinfo.get("season") and not media_season:
|
||||
media_season = tmdbinfo.get("season")
|
||||
torrents = search_chain.async_search_by_id_stream(
|
||||
tmdbid=tmdbinfo.get("id"),
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
else:
|
||||
yield {
|
||||
"type": "error",
|
||||
"success": False,
|
||||
"message": "未识别到TMDB媒体信息",
|
||||
}
|
||||
return
|
||||
else:
|
||||
torrents = search_chain.async_search_by_id_stream(
|
||||
doubanid=doubanid,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
elif mediaid.startswith("bangumi:"):
|
||||
bangumiid = int(mediaid.replace("bangumi:", ""))
|
||||
if settings.RECOGNIZE_SOURCE == "themoviedb":
|
||||
tmdbinfo = await media_chain.async_get_tmdbinfo_by_bangumiid(
|
||||
bangumiid=bangumiid
|
||||
)
|
||||
if tmdbinfo:
|
||||
torrents = search_chain.async_search_by_id_stream(
|
||||
tmdbid=tmdbinfo.get("id"),
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
else:
|
||||
yield {
|
||||
"type": "error",
|
||||
"success": False,
|
||||
"message": "未识别到TMDB媒体信息",
|
||||
}
|
||||
return
|
||||
else:
|
||||
doubaninfo = await media_chain.async_get_doubaninfo_by_bangumiid(
|
||||
bangumiid=bangumiid
|
||||
)
|
||||
if doubaninfo:
|
||||
torrents = search_chain.async_search_by_id_stream(
|
||||
doubanid=doubaninfo.get("id"),
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
else:
|
||||
yield {
|
||||
"type": "error",
|
||||
"success": False,
|
||||
"message": "未识别到豆瓣媒体信息",
|
||||
}
|
||||
return
|
||||
else:
|
||||
event_data = MediaRecognizeConvertEventData(
|
||||
mediaid=mediaid, convert_type=settings.RECOGNIZE_SOURCE
|
||||
)
|
||||
event = await eventmanager.async_send_event(
|
||||
ChainEventType.MediaRecognizeConvert, event_data
|
||||
)
|
||||
if event and event.event_data:
|
||||
event_data = event.event_data
|
||||
if event_data.media_dict:
|
||||
search_id = event_data.media_dict.get("id")
|
||||
if event_data.convert_type == "themoviedb":
|
||||
torrents = search_chain.async_search_by_id_stream(
|
||||
tmdbid=search_id,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
elif event_data.convert_type == "douban":
|
||||
torrents = search_chain.async_search_by_id_stream(
|
||||
doubanid=search_id,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
else:
|
||||
if not title:
|
||||
yield {"type": "error", "success": False, "message": "未知的媒体ID"}
|
||||
return
|
||||
meta = MetaInfo(title)
|
||||
if year:
|
||||
meta.year = year
|
||||
if media_type:
|
||||
meta.type = media_type
|
||||
if media_season:
|
||||
meta.type = MediaType.TV
|
||||
meta.begin_season = media_season
|
||||
mediainfo = await media_chain.async_recognize_by_meta(
|
||||
meta,
|
||||
obtain_images=False,
|
||||
)
|
||||
if mediainfo:
|
||||
if settings.RECOGNIZE_SOURCE == "themoviedb":
|
||||
torrents = search_chain.async_search_by_id_stream(
|
||||
tmdbid=mediainfo.tmdb_id,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
else:
|
||||
torrents = search_chain.async_search_by_id_stream(
|
||||
doubanid=mediainfo.douban_id,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
|
||||
if not torrents:
|
||||
yield {"type": "error", "success": False, "message": "未搜索到任何资源"}
|
||||
"""解析媒体身份并输出精确搜索流事件。"""
|
||||
search_params, message = await _resolve_media_search_params(
|
||||
mediaid=mediaid,
|
||||
media_type=media_type,
|
||||
title=title,
|
||||
year=year,
|
||||
media_season=media_season,
|
||||
)
|
||||
if not search_params:
|
||||
yield {"type": "error", "success": False, "message": message}
|
||||
return
|
||||
|
||||
torrents = search_chain.async_search_by_id_stream(
|
||||
**search_params,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
async for event in torrents:
|
||||
yield event
|
||||
|
||||
return StreamingResponse(
|
||||
_stream_search_events(request, event_source()), media_type="text/event-stream"
|
||||
_stream_search_events(request, event_source()),
|
||||
media_type="text/event-stream",
|
||||
headers=_SSE_RESPONSE_HEADERS,
|
||||
)
|
||||
|
||||
|
||||
@@ -393,180 +443,32 @@ async def search_by_id(
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> Any:
|
||||
"""
|
||||
根据TMDBID/豆瓣ID精确搜索站点资源 tmdb:/douban:/bangumi:
|
||||
根据带来源前缀的媒体 ID 精确搜索站点资源。
|
||||
"""
|
||||
media_type = _parse_media_type(mtype)
|
||||
if season:
|
||||
media_season = int(season)
|
||||
else:
|
||||
media_season = None
|
||||
if sites:
|
||||
site_list = [int(site) for site in sites.split(",") if site]
|
||||
else:
|
||||
site_list = None
|
||||
torrents = None
|
||||
media_chain = MediaChain()
|
||||
search_chain = SearchChain()
|
||||
# 根据前缀识别媒体ID
|
||||
if mediaid.startswith("tmdb:"):
|
||||
tmdbid = int(mediaid.replace("tmdb:", ""))
|
||||
if settings.RECOGNIZE_SOURCE == "douban":
|
||||
# 通过TMDBID识别豆瓣ID
|
||||
doubaninfo = await media_chain.async_get_doubaninfo_by_tmdbid(
|
||||
tmdbid=tmdbid, mtype=media_type
|
||||
)
|
||||
if doubaninfo:
|
||||
torrents = await search_chain.async_search_by_id(
|
||||
doubanid=doubaninfo.get("id"),
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
else:
|
||||
return schemas.Response(success=False, message="未识别到豆瓣媒体信息")
|
||||
else:
|
||||
torrents = await search_chain.async_search_by_id(
|
||||
tmdbid=tmdbid,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
elif mediaid.startswith("douban:"):
|
||||
doubanid = mediaid.replace("douban:", "")
|
||||
if settings.RECOGNIZE_SOURCE == "themoviedb":
|
||||
# 通过豆瓣ID识别TMDBID
|
||||
tmdbinfo = await media_chain.async_get_tmdbinfo_by_doubanid(
|
||||
doubanid=doubanid, mtype=media_type
|
||||
)
|
||||
if tmdbinfo:
|
||||
if tmdbinfo.get("season") and not media_season:
|
||||
media_season = tmdbinfo.get("season")
|
||||
torrents = await search_chain.async_search_by_id(
|
||||
tmdbid=tmdbinfo.get("id"),
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
else:
|
||||
return schemas.Response(success=False, message="未识别到TMDB媒体信息")
|
||||
else:
|
||||
torrents = await search_chain.async_search_by_id(
|
||||
doubanid=doubanid,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
elif mediaid.startswith("bangumi:"):
|
||||
bangumiid = int(mediaid.replace("bangumi:", ""))
|
||||
if settings.RECOGNIZE_SOURCE == "themoviedb":
|
||||
# 通过BangumiID识别TMDBID
|
||||
tmdbinfo = await media_chain.async_get_tmdbinfo_by_bangumiid(
|
||||
bangumiid=bangumiid
|
||||
)
|
||||
if tmdbinfo:
|
||||
torrents = await search_chain.async_search_by_id(
|
||||
tmdbid=tmdbinfo.get("id"),
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
else:
|
||||
return schemas.Response(success=False, message="未识别到TMDB媒体信息")
|
||||
else:
|
||||
# 通过BangumiID识别豆瓣ID
|
||||
doubaninfo = await media_chain.async_get_doubaninfo_by_bangumiid(
|
||||
bangumiid=bangumiid
|
||||
)
|
||||
if doubaninfo:
|
||||
torrents = await search_chain.async_search_by_id(
|
||||
doubanid=doubaninfo.get("id"),
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=site_list,
|
||||
cache_local=True,
|
||||
)
|
||||
else:
|
||||
return schemas.Response(success=False, message="未识别到豆瓣媒体信息")
|
||||
else:
|
||||
# 未知前缀,广播事件解析媒体信息
|
||||
event_data = MediaRecognizeConvertEventData(
|
||||
mediaid=mediaid, convert_type=settings.RECOGNIZE_SOURCE
|
||||
)
|
||||
event = await eventmanager.async_send_event(
|
||||
ChainEventType.MediaRecognizeConvert, event_data
|
||||
)
|
||||
# 使用事件返回的上下文数据
|
||||
if event and event.event_data:
|
||||
event_data: MediaRecognizeConvertEventData = event.event_data
|
||||
if event_data.media_dict:
|
||||
search_id = event_data.media_dict.get("id")
|
||||
if event_data.convert_type == "themoviedb":
|
||||
torrents = await search_chain.async_search_by_id(
|
||||
tmdbid=search_id,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
cache_local=True,
|
||||
)
|
||||
elif event_data.convert_type == "douban":
|
||||
torrents = await search_chain.async_search_by_id(
|
||||
doubanid=search_id,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
cache_local=True,
|
||||
)
|
||||
else:
|
||||
if not title:
|
||||
return schemas.Response(success=False, message="未知的媒体ID")
|
||||
# 使用名称识别兜底
|
||||
meta = MetaInfo(title)
|
||||
if year:
|
||||
meta.year = year
|
||||
if media_type:
|
||||
meta.type = media_type
|
||||
if media_season:
|
||||
meta.type = MediaType.TV
|
||||
meta.begin_season = media_season
|
||||
mediainfo = await media_chain.async_recognize_by_meta(
|
||||
meta,
|
||||
obtain_images=False,
|
||||
)
|
||||
if mediainfo:
|
||||
if settings.RECOGNIZE_SOURCE == "themoviedb":
|
||||
torrents = await search_chain.async_search_by_id(
|
||||
tmdbid=mediainfo.tmdb_id,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
cache_local=True,
|
||||
)
|
||||
else:
|
||||
torrents = await search_chain.async_search_by_id(
|
||||
doubanid=mediainfo.douban_id,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
cache_local=True,
|
||||
)
|
||||
# 返回搜索结果
|
||||
media_season = int(season) if season else None
|
||||
search_params, message = await _resolve_media_search_params(
|
||||
mediaid=mediaid,
|
||||
media_type=media_type,
|
||||
title=title,
|
||||
year=year,
|
||||
media_season=media_season,
|
||||
)
|
||||
if not search_params:
|
||||
return schemas.Response(success=False, message=message)
|
||||
torrents = await SearchChain().async_search_by_id(
|
||||
**search_params,
|
||||
mtype=media_type,
|
||||
area=area,
|
||||
season=media_season,
|
||||
sites=_parse_site_list(sites),
|
||||
cache_local=True,
|
||||
)
|
||||
if not torrents:
|
||||
return schemas.Response(success=False, message="未搜索到任何资源")
|
||||
else:
|
||||
return schemas.Response(
|
||||
success=True, data=[torrent.to_dict() for torrent in torrents]
|
||||
)
|
||||
return schemas.Response(
|
||||
success=True, data=[torrent.to_dict() for torrent in torrents]
|
||||
)
|
||||
|
||||
|
||||
@router.get("/title/stream", summary="渐进式模糊搜索资源")
|
||||
@@ -585,7 +487,9 @@ async def search_by_title_stream(
|
||||
title=keyword, page=page, sites=_parse_site_list(sites), cache_local=True
|
||||
)
|
||||
return StreamingResponse(
|
||||
_stream_search_events(request, event_source), media_type="text/event-stream"
|
||||
_stream_search_events(request, event_source),
|
||||
media_type="text/event-stream",
|
||||
headers=_SSE_RESPONSE_HEADERS,
|
||||
)
|
||||
|
||||
|
||||
@@ -625,7 +529,12 @@ async def search_subtitle_by_title_stream(
|
||||
title=keyword, page=page, sites=_parse_site_list(sites), cache_local=True
|
||||
)
|
||||
return StreamingResponse(
|
||||
_stream_search_events(request, event_source), media_type="text/event-stream"
|
||||
_stream_search_events(
|
||||
request,
|
||||
_iter_signed_subtitle_search_events(event_source),
|
||||
),
|
||||
media_type="text/event-stream",
|
||||
headers=_SSE_RESPONSE_HEADERS,
|
||||
)
|
||||
|
||||
|
||||
@@ -645,7 +554,7 @@ async def search_subtitle_by_title(
|
||||
if not subtitles:
|
||||
return schemas.Response(success=False, message="未搜索到任何字幕")
|
||||
return schemas.Response(
|
||||
success=True, data=[subtitle.to_dict() for subtitle in subtitles]
|
||||
success=True, data=_serialize_signed_subtitle_results(subtitles)
|
||||
)
|
||||
|
||||
|
||||
@@ -666,7 +575,6 @@ async def _build_subtitle_search_source(
|
||||
media_season = int(season) if season else None
|
||||
media_episode = int(episode) if episode else None
|
||||
site_list = _parse_site_list(sites)
|
||||
media_chain = MediaChain()
|
||||
search_chain = SearchChain()
|
||||
|
||||
def call_search(**kwargs):
|
||||
@@ -685,80 +593,16 @@ async def _build_subtitle_search_source(
|
||||
return search_chain.async_search_subtitles_by_id_stream(**params)
|
||||
return search_chain.async_search_subtitles_by_id(**params)
|
||||
|
||||
if mediaid.startswith("tmdb:"):
|
||||
tmdbid = int(mediaid.replace("tmdb:", ""))
|
||||
if settings.RECOGNIZE_SOURCE == "douban":
|
||||
doubaninfo = await media_chain.async_get_doubaninfo_by_tmdbid(
|
||||
tmdbid=tmdbid, mtype=media_type
|
||||
)
|
||||
if not doubaninfo:
|
||||
return None, "未识别到豆瓣媒体信息"
|
||||
return call_search(doubanid=doubaninfo.get("id")), ""
|
||||
return call_search(tmdbid=tmdbid), ""
|
||||
|
||||
if mediaid.startswith("douban:"):
|
||||
doubanid = mediaid.replace("douban:", "")
|
||||
if settings.RECOGNIZE_SOURCE == "themoviedb":
|
||||
tmdbinfo = await media_chain.async_get_tmdbinfo_by_doubanid(
|
||||
doubanid=doubanid, mtype=media_type
|
||||
)
|
||||
if not tmdbinfo:
|
||||
return None, "未识别到TMDB媒体信息"
|
||||
if tmdbinfo.get("season") and not media_season:
|
||||
media_season = tmdbinfo.get("season")
|
||||
return call_search(tmdbid=tmdbinfo.get("id")), ""
|
||||
return call_search(doubanid=doubanid), ""
|
||||
|
||||
if mediaid.startswith("bangumi:"):
|
||||
bangumiid = int(mediaid.replace("bangumi:", ""))
|
||||
if settings.RECOGNIZE_SOURCE == "themoviedb":
|
||||
tmdbinfo = await media_chain.async_get_tmdbinfo_by_bangumiid(
|
||||
bangumiid=bangumiid
|
||||
)
|
||||
if not tmdbinfo:
|
||||
return None, "未识别到TMDB媒体信息"
|
||||
return call_search(tmdbid=tmdbinfo.get("id")), ""
|
||||
doubaninfo = await media_chain.async_get_doubaninfo_by_bangumiid(
|
||||
bangumiid=bangumiid
|
||||
)
|
||||
if not doubaninfo:
|
||||
return None, "未识别到豆瓣媒体信息"
|
||||
return call_search(doubanid=doubaninfo.get("id")), ""
|
||||
|
||||
event_data = MediaRecognizeConvertEventData(
|
||||
mediaid=mediaid, convert_type=settings.RECOGNIZE_SOURCE
|
||||
search_params, message = await _resolve_media_search_params(
|
||||
mediaid=mediaid,
|
||||
media_type=media_type,
|
||||
title=title,
|
||||
year=year,
|
||||
media_season=media_season,
|
||||
)
|
||||
event = await eventmanager.async_send_event(
|
||||
ChainEventType.MediaRecognizeConvert, event_data
|
||||
)
|
||||
if event and event.event_data and event.event_data.media_dict:
|
||||
event_data = event.event_data
|
||||
search_id = event_data.media_dict.get("id")
|
||||
if event_data.convert_type == "themoviedb":
|
||||
return call_search(tmdbid=search_id), ""
|
||||
if event_data.convert_type == "douban":
|
||||
return call_search(doubanid=search_id), ""
|
||||
|
||||
if not title:
|
||||
return None, "未知的媒体ID"
|
||||
|
||||
meta = MetaInfo(title)
|
||||
if year:
|
||||
meta.year = year
|
||||
if media_type:
|
||||
meta.type = media_type
|
||||
if media_season:
|
||||
meta.type = MediaType.TV
|
||||
meta.begin_season = media_season
|
||||
mediainfo = await media_chain.async_recognize_by_meta(
|
||||
meta,
|
||||
obtain_images=False,
|
||||
)
|
||||
if not mediainfo:
|
||||
return None, "未识别到媒体信息"
|
||||
if settings.RECOGNIZE_SOURCE == "themoviedb":
|
||||
return call_search(tmdbid=mediainfo.tmdb_id), ""
|
||||
return call_search(doubanid=mediainfo.douban_id), ""
|
||||
if not search_params:
|
||||
return None, message
|
||||
return call_search(**search_params), ""
|
||||
|
||||
|
||||
@router.get("/subtitle/media/{mediaid}/stream", summary="渐进式精确搜索字幕")
|
||||
@@ -774,7 +618,7 @@ async def search_subtitle_by_id_stream(
|
||||
_: schemas.TokenPayload = Depends(verify_resource_token),
|
||||
) -> Any:
|
||||
"""
|
||||
根据TMDBID/豆瓣ID渐进式精确搜索站点字幕资源,返回格式为SSE。
|
||||
根据带来源前缀的媒体 ID 渐进式精确搜索站点字幕资源,返回格式为SSE。
|
||||
"""
|
||||
subtitles, message = await _build_subtitle_search_source(
|
||||
mediaid=mediaid,
|
||||
@@ -798,7 +642,12 @@ async def search_subtitle_by_id_stream(
|
||||
yield event
|
||||
|
||||
return StreamingResponse(
|
||||
_stream_search_events(request, event_source()), media_type="text/event-stream"
|
||||
_stream_search_events(
|
||||
request,
|
||||
_iter_signed_subtitle_search_events(event_source()),
|
||||
),
|
||||
media_type="text/event-stream",
|
||||
headers=_SSE_RESPONSE_HEADERS,
|
||||
)
|
||||
|
||||
|
||||
@@ -814,7 +663,7 @@ async def search_subtitle_by_id(
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
) -> Any:
|
||||
"""
|
||||
根据TMDBID/豆瓣ID精确搜索站点字幕资源。
|
||||
根据带来源前缀的媒体 ID 精确搜索站点字幕资源。
|
||||
"""
|
||||
subtitles, message = await _build_subtitle_search_source(
|
||||
mediaid=mediaid,
|
||||
@@ -832,7 +681,7 @@ async def search_subtitle_by_id(
|
||||
if not subtitles:
|
||||
return schemas.Response(success=False, message="未搜索到任何字幕")
|
||||
return schemas.Response(
|
||||
success=True, data=[subtitle.to_dict() for subtitle in subtitles]
|
||||
success=True, data=_serialize_signed_subtitle_results(subtitles)
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -22,6 +22,8 @@ from app.db.models.siteuserdata import SiteUserData
|
||||
from app.db.site_oper import SiteOper
|
||||
from app.db.systemconfig_oper import SystemConfigOper
|
||||
from app.db.user_oper import (
|
||||
get_current_active_manage_user,
|
||||
get_current_active_manage_user_async,
|
||||
get_current_active_superuser,
|
||||
get_current_active_superuser_async,
|
||||
)
|
||||
@@ -37,7 +39,7 @@ router = APIRouter()
|
||||
@router.get("/", summary="所有站点", response_model=List[schemas.Site])
|
||||
async def read_sites(
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user_async),
|
||||
) -> List[dict]:
|
||||
"""
|
||||
获取站点列表
|
||||
@@ -50,7 +52,7 @@ async def add_site(
|
||||
*,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
site_in: schemas.Site,
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
新增站点
|
||||
@@ -89,7 +91,7 @@ async def update_site(
|
||||
*,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
site_in: schemas.Site,
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
更新站点信息
|
||||
@@ -150,7 +152,7 @@ def reset(
|
||||
async def update_sites_priority(
|
||||
priorities: List[dict],
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
_: User = Depends(get_current_active_manage_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
批量更新站点优先级
|
||||
@@ -203,7 +205,7 @@ def update_cookie_by_body(
|
||||
site_id: int,
|
||||
site_cookie_update: schemas.SiteCookieUpdate,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
使用请求体中的用户密码更新站点Cookie
|
||||
@@ -226,7 +228,7 @@ def update_cookie(
|
||||
password: str,
|
||||
code: Optional[str] = None,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
使用用户密码更新站点Cookie
|
||||
@@ -246,7 +248,7 @@ def update_cookie(
|
||||
def refresh_userdata(
|
||||
site_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
刷新站点用户数据
|
||||
@@ -273,7 +275,7 @@ def refresh_userdata(
|
||||
)
|
||||
async def read_userdata_latest(
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
_: User = Depends(get_current_active_manage_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
查询所有站点最新用户数据
|
||||
@@ -291,7 +293,7 @@ async def read_userdata(
|
||||
site_id: int,
|
||||
workdate: Optional[str] = None,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
_: User = Depends(get_current_active_manage_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
查询站点用户数据
|
||||
@@ -395,7 +397,7 @@ async def site_resource(
|
||||
cat: Optional[str] = None,
|
||||
page: Optional[int] = 0,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
_: User = Depends(get_current_active_manage_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
浏览站点资源
|
||||
@@ -543,7 +545,7 @@ async def support_sites(_: User = Depends(get_current_active_superuser_async)):
|
||||
async def read_site(
|
||||
site_id: int,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
_: User = Depends(get_current_active_manage_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
通过ID获取站点信息
|
||||
@@ -561,7 +563,7 @@ async def read_site(
|
||||
async def delete_site(
|
||||
site_id: int,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
_: User = Depends(get_current_active_manage_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
删除站点
|
||||
|
||||
@@ -15,6 +15,7 @@ from app.core.config import settings
|
||||
from app.core.security import verify_token
|
||||
from app.db.models import User
|
||||
from app.db.user_oper import (
|
||||
get_current_active_manage_user,
|
||||
get_current_active_superuser,
|
||||
get_current_active_superuser_async,
|
||||
)
|
||||
@@ -91,7 +92,7 @@ def list_files(
|
||||
fileitem: schemas.FileItem,
|
||||
sort: Optional[str] = "updated_at",
|
||||
keyword: Optional[str] = None,
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
查询当前目录下所有目录和文件
|
||||
@@ -117,7 +118,7 @@ def list_files(
|
||||
def mkdir(
|
||||
fileitem: schemas.FileItem,
|
||||
name: str,
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
创建目录
|
||||
@@ -135,7 +136,7 @@ def mkdir(
|
||||
|
||||
@router.post("/delete", summary="删除文件或目录", response_model=schemas.Response)
|
||||
def delete(
|
||||
fileitem: schemas.FileItem, _: User = Depends(get_current_active_superuser)
|
||||
fileitem: schemas.FileItem, _: User = Depends(get_current_active_manage_user)
|
||||
) -> Any:
|
||||
"""
|
||||
删除文件或目录
|
||||
@@ -150,7 +151,7 @@ def delete(
|
||||
|
||||
@router.post("/download", summary="下载文件")
|
||||
def download(
|
||||
fileitem: schemas.FileItem, _: User = Depends(get_current_active_superuser)
|
||||
fileitem: schemas.FileItem, _: User = Depends(get_current_active_manage_user)
|
||||
) -> Any:
|
||||
"""
|
||||
下载文件或目录
|
||||
@@ -166,7 +167,7 @@ def download(
|
||||
|
||||
@router.post("/image", summary="预览图片")
|
||||
def image(
|
||||
fileitem: schemas.FileItem, _: User = Depends(get_current_active_superuser)
|
||||
fileitem: schemas.FileItem, _: User = Depends(get_current_active_manage_user)
|
||||
) -> Any:
|
||||
"""
|
||||
下载文件或目录
|
||||
@@ -185,7 +186,7 @@ def rename(
|
||||
fileitem: schemas.FileItem,
|
||||
new_name: str,
|
||||
recursive: Optional[bool] = False,
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
重命名文件或目录
|
||||
|
||||
@@ -17,11 +17,13 @@ from app.db.models.subscribe import Subscribe
|
||||
from app.db.models.subscribehistory import SubscribeHistory
|
||||
from app.db.models.user import User
|
||||
from app.db.systemconfig_oper import SystemConfigOper
|
||||
from app.db.user_oper import get_current_active_user_async
|
||||
from app.db.user_oper import get_current_active_user, get_current_active_user_async
|
||||
from app.helper.server import MoviePilotServerHelper
|
||||
from app.log import logger
|
||||
from app.scheduler import Scheduler
|
||||
from app.schemas.event import SubscribeModifiedEventData
|
||||
from app.schemas.types import MediaType, EventType, SystemConfigKey
|
||||
from app.utils.media import normalize_media_source, parse_media_key
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -50,14 +52,98 @@ def build_subscribe_event_payload(subscribe: Subscribe) -> dict:
|
||||
return {column.name: values.get(column.name) for column in subscribe.__table__.columns}
|
||||
|
||||
|
||||
def can_access_subscribe(
|
||||
subscribe: Subscribe | SubscribeHistory | None, current_user: User
|
||||
) -> bool:
|
||||
"""
|
||||
判断当前用户是否可访问订阅及其历史记录。
|
||||
|
||||
超级用户拥有全局订阅管理能力;普通用户只能访问 username 精确匹配自己的订阅。
|
||||
空 username 表示无法归属的 legacy 订阅,只能由超级用户管理。
|
||||
"""
|
||||
if not subscribe:
|
||||
return False
|
||||
if current_user.is_superuser:
|
||||
return True
|
||||
username = subscribe.username
|
||||
return bool(username) and username == current_user.name
|
||||
|
||||
|
||||
async def get_accessible_subscribe(
|
||||
db: AsyncSession, subscribe_id: int, current_user: User
|
||||
) -> Subscribe | None:
|
||||
"""
|
||||
按订阅 ID 读取当前用户可访问的订阅行。
|
||||
"""
|
||||
subscribe = await Subscribe.async_get(db, subscribe_id)
|
||||
if can_access_subscribe(subscribe, current_user):
|
||||
return subscribe
|
||||
return None
|
||||
|
||||
|
||||
def get_accessible_subscribe_sync(
|
||||
db: Session, subscribe_id: int, current_user: User
|
||||
) -> Subscribe | None:
|
||||
"""
|
||||
同步读取当前用户可访问的订阅行。
|
||||
"""
|
||||
subscribe = Subscribe.get(db, subscribe_id)
|
||||
if can_access_subscribe(subscribe, current_user):
|
||||
return subscribe
|
||||
return None
|
||||
|
||||
|
||||
def select_accessible_subscribe(
|
||||
subscribes: List[Subscribe], current_user: User
|
||||
) -> Subscribe | None:
|
||||
"""
|
||||
从候选订阅中选择当前用户可访问的第一条记录。
|
||||
"""
|
||||
for subscribe in subscribes or []:
|
||||
if can_access_subscribe(subscribe, current_user):
|
||||
return subscribe
|
||||
return None
|
||||
|
||||
|
||||
async def list_subscribes_by_media_key(
|
||||
db: AsyncSession, media_key: str, season: Optional[int] = None,
|
||||
) -> List[Subscribe]:
|
||||
"""按统一媒体键查询订阅,并兼容迁移前的专用 ID 字段。"""
|
||||
source, media_id = parse_media_key(media_key)
|
||||
if not source or not media_id:
|
||||
return await Subscribe.async_list_by_mediaid(db, media_key)
|
||||
|
||||
subscribes = list(await Subscribe.async_list_by_media_identity(
|
||||
db, media_source=source, media_id=media_id
|
||||
))
|
||||
if source == "themoviedb" and media_id.isdigit():
|
||||
subscribes.extend(await Subscribe.async_get_by_tmdbid(db, int(media_id), season))
|
||||
elif source == "douban":
|
||||
subscribes.extend(await Subscribe.async_list_by_doubanid(db, media_id))
|
||||
elif source == "bangumi" and media_id.isdigit():
|
||||
subscribes.extend(await Subscribe.async_list_by_bangumiid(db, int(media_id)))
|
||||
elif source == "anilist" and media_id.isdigit():
|
||||
subscribes.extend(await Subscribe.async_list_by_anilistid(db, int(media_id)))
|
||||
|
||||
unique_subscribes = {subscribe.id: subscribe for subscribe in subscribes}
|
||||
if season is not None:
|
||||
return [
|
||||
subscribe for subscribe in unique_subscribes.values()
|
||||
if subscribe.season == season
|
||||
]
|
||||
return list(unique_subscribes.values())
|
||||
|
||||
|
||||
@router.get("/", summary="查询所有订阅", response_model=List[schemas.Subscribe])
|
||||
async def read_subscribes(
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
查询所有订阅
|
||||
"""
|
||||
if not current_user.is_superuser:
|
||||
return await Subscribe.async_list_by_username(db, current_user.name)
|
||||
return await Subscribe.async_list(db)
|
||||
|
||||
|
||||
@@ -68,7 +154,7 @@ async def list_subscribes(_: Annotated[str, Depends(verify_apitoken)]) -> Any:
|
||||
"""
|
||||
查询所有订阅 API_TOKEN认证(?token=xxx)
|
||||
"""
|
||||
return await read_subscribes()
|
||||
return await Subscribe.async_list()
|
||||
|
||||
|
||||
@router.post("/", summary="新增订阅", response_model=schemas.Response)
|
||||
@@ -85,8 +171,13 @@ async def create_subscribe(
|
||||
mtype = MediaType(subscribe_in.type)
|
||||
else:
|
||||
mtype = None
|
||||
# 豆瓣标理
|
||||
if subscribe_in.doubanid or subscribe_in.bangumiid:
|
||||
# 非 TMDB 来源的标题可能自带季标记,入库前统一拆分。
|
||||
if (
|
||||
subscribe_in.doubanid
|
||||
or subscribe_in.bangumiid
|
||||
or subscribe_in.anilistid
|
||||
or normalize_media_source(subscribe_in.media_source) not in (None, "themoviedb")
|
||||
):
|
||||
meta = MetaInfo(subscribe_in.name)
|
||||
subscribe_in.name = meta.name
|
||||
if subscribe_in.season is None:
|
||||
@@ -96,16 +187,14 @@ async def create_subscribe(
|
||||
title = subscribe_in.name
|
||||
else:
|
||||
title = None
|
||||
# 订阅用户
|
||||
subscribe_in.username = current_user.name
|
||||
# 转化为字典
|
||||
subscribe_dict = subscribe_in.model_dump()
|
||||
if subscribe_in.id:
|
||||
subscribe_dict.pop("id", None)
|
||||
# completed_episode 是响应派生字段,禁止写入持久层
|
||||
subscribe_dict.pop("completed_episode", None)
|
||||
subscribe_dict = subscribe_in.to_public_write_payload()
|
||||
subscribe_dict["username"] = current_user.name
|
||||
sid, message = await SubscribeChain().async_add(
|
||||
mtype=mtype, title=title, exist_ok=True, **subscribe_dict
|
||||
mtype=mtype,
|
||||
title=title,
|
||||
exist_ok=True,
|
||||
owner_scope=not current_user.is_superuser,
|
||||
**subscribe_dict,
|
||||
)
|
||||
return schemas.Response(success=bool(sid), message=message, data={"id": sid})
|
||||
|
||||
@@ -115,30 +204,22 @@ async def update_subscribe(
|
||||
*,
|
||||
subscribe_in: schemas.Subscribe,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
更新订阅信息
|
||||
"""
|
||||
subscribe = await Subscribe.async_get(db, subscribe_in.id)
|
||||
subscribe = await get_accessible_subscribe(db, subscribe_in.id, current_user)
|
||||
if not subscribe:
|
||||
return schemas.Response(success=False, message="订阅不存在")
|
||||
# 避免更新缺失集数
|
||||
old_subscribe_dict = subscribe.to_dict()
|
||||
subscribe_dict = subscribe_in.model_dump()
|
||||
if subscribe_in.episode_priority is None:
|
||||
subscribe_dict.pop("episode_priority", None)
|
||||
# completed_episode 是响应派生字段,禁止写入持久层
|
||||
subscribe_dict.pop("completed_episode", None)
|
||||
if not subscribe_in.lack_episode:
|
||||
# 没有缺失集数时,缺失集数清空,避免更新为0
|
||||
subscribe_dict.pop("lack_episode")
|
||||
elif subscribe_in.total_episode:
|
||||
# 总集数增加时,缺失集数也要增加
|
||||
if subscribe_in.total_episode > (subscribe.total_episode or 0):
|
||||
subscribe_dict["lack_episode"] = subscribe.lack_episode + (
|
||||
subscribe_in.total_episode - (subscribe.total_episode or 0)
|
||||
)
|
||||
subscribe_dict = subscribe_in.to_public_write_payload()
|
||||
subscribe_dict["username"] = subscribe.username
|
||||
if subscribe_in.total_episode and subscribe_in.total_episode > (subscribe.total_episode or 0):
|
||||
# 扩大目标范围时,新增加的集数尚无下载事实,应同步计入缺失集数。
|
||||
subscribe_dict["lack_episode"] = (subscribe.lack_episode or 0) + (
|
||||
subscribe_in.total_episode - (subscribe.total_episode or 0)
|
||||
)
|
||||
# 是否手动修改过总集数
|
||||
if subscribe_in.total_episode != subscribe.total_episode:
|
||||
subscribe_dict["manual_total_episode"] = 1
|
||||
@@ -149,11 +230,12 @@ async def update_subscribe(
|
||||
# 发送订阅调整事件
|
||||
await eventmanager.async_send_event(
|
||||
EventType.SubscribeModified,
|
||||
{
|
||||
"subscribe_id": subscribe_in.id,
|
||||
"old_subscribe_info": old_subscribe_dict,
|
||||
"subscribe_info": updated_subscribe.to_dict() if updated_subscribe else {},
|
||||
},
|
||||
SubscribeModifiedEventData(
|
||||
subscribe_id=subscribe_in.id,
|
||||
old_subscribe_info=old_subscribe_dict,
|
||||
subscribe_info=updated_subscribe.to_dict() if updated_subscribe else {},
|
||||
scene="update",
|
||||
).to_dict(),
|
||||
)
|
||||
return schemas.Response(success=True)
|
||||
|
||||
@@ -163,12 +245,12 @@ async def update_subscribe_status(
|
||||
subid: int,
|
||||
state: str,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
更新订阅状态
|
||||
"""
|
||||
subscribe = await Subscribe.async_get(db, subid)
|
||||
subscribe = await get_accessible_subscribe(db, subid, current_user)
|
||||
if not subscribe:
|
||||
return schemas.Response(success=False, message="订阅不存在")
|
||||
valid_states = ["R", "P", "S"]
|
||||
@@ -181,11 +263,12 @@ async def update_subscribe_status(
|
||||
# 发送订阅调整事件
|
||||
await eventmanager.async_send_event(
|
||||
EventType.SubscribeModified,
|
||||
{
|
||||
"subscribe_id": subid,
|
||||
"old_subscribe_info": old_subscribe_dict,
|
||||
"subscribe_info": updated_subscribe.to_dict() if updated_subscribe else {},
|
||||
},
|
||||
SubscribeModifiedEventData(
|
||||
subscribe_id=subid,
|
||||
old_subscribe_info=old_subscribe_dict,
|
||||
subscribe_info=updated_subscribe.to_dict() if updated_subscribe else {},
|
||||
scene="status",
|
||||
).to_dict(),
|
||||
)
|
||||
return schemas.Response(success=True)
|
||||
|
||||
@@ -196,52 +279,37 @@ async def subscribe_mediaid(
|
||||
season: Optional[int] = None,
|
||||
title: Optional[str] = None,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
根据 TMDBID/豆瓣ID/BangumiId 查询订阅 tmdb:/douban:
|
||||
根据 TMDB、豆瓣、Bangumi、AniList 或插件媒体键查询订阅。
|
||||
"""
|
||||
title_check = False
|
||||
if mediaid.startswith("tmdb:"):
|
||||
tmdbid = mediaid[5:]
|
||||
if not tmdbid or not str(tmdbid).isdigit():
|
||||
return Subscribe()
|
||||
result = await Subscribe.async_exists(db, tmdbid=int(tmdbid), season=season)
|
||||
elif mediaid.startswith("douban:"):
|
||||
doubanid = mediaid[7:]
|
||||
if not doubanid:
|
||||
return Subscribe()
|
||||
result = await Subscribe.async_get_by_doubanid(db, doubanid)
|
||||
if not result and title:
|
||||
title_check = True
|
||||
elif mediaid.startswith("bangumi:"):
|
||||
bangumiid = mediaid[8:]
|
||||
if not bangumiid or not str(bangumiid).isdigit():
|
||||
return Subscribe()
|
||||
result = await Subscribe.async_get_by_bangumiid(db, int(bangumiid))
|
||||
if not result and title:
|
||||
title_check = True
|
||||
else:
|
||||
result = await Subscribe.async_get_by_mediaid(db, mediaid)
|
||||
if not result and title:
|
||||
title_check = True
|
||||
subscribes = await list_subscribes_by_media_key(db, mediaid, season)
|
||||
result = select_accessible_subscribe(subscribes, current_user)
|
||||
source, _ = parse_media_key(mediaid)
|
||||
title_check = not result and bool(title) and source != "themoviedb"
|
||||
# 使用名称检查订阅
|
||||
if title_check and title:
|
||||
meta = MetaInfo(title)
|
||||
if season is not None:
|
||||
meta.begin_season = season
|
||||
result = await Subscribe.async_get_by_title(
|
||||
subscribes = await Subscribe.async_list_by_title(
|
||||
db, title=meta.name, season=meta.begin_season
|
||||
)
|
||||
result = select_accessible_subscribe(subscribes, current_user)
|
||||
|
||||
return result if result else Subscribe()
|
||||
|
||||
|
||||
@router.get("/refresh", summary="刷新订阅", response_model=schemas.Response)
|
||||
def refresh_subscribes(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
|
||||
def refresh_subscribes(
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
) -> Any:
|
||||
"""
|
||||
刷新所有订阅
|
||||
"""
|
||||
if not current_user.is_superuser:
|
||||
return schemas.Response(success=False, message="订阅不存在")
|
||||
Scheduler().start("subscribe_refresh")
|
||||
return schemas.Response(success=True)
|
||||
|
||||
@@ -250,12 +318,12 @@ def refresh_subscribes(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
|
||||
async def reset_subscribes(
|
||||
subid: int,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
重置订阅
|
||||
"""
|
||||
subscribe = await Subscribe.async_get(db, subid)
|
||||
subscribe = await get_accessible_subscribe(db, subid, current_user)
|
||||
if subscribe:
|
||||
# 在更新之前获取旧数据
|
||||
old_subscribe_dict = subscribe.to_dict()
|
||||
@@ -267,6 +335,8 @@ async def reset_subscribes(
|
||||
"lack_episode": subscribe.total_episode,
|
||||
"current_priority": None,
|
||||
"episode_priority": {},
|
||||
# 重置代表放弃手动总集数,后续订阅检查重新按 TMDB 集数更新。
|
||||
"manual_total_episode": 0,
|
||||
"state": "R",
|
||||
},
|
||||
)
|
||||
@@ -275,39 +345,57 @@ async def reset_subscribes(
|
||||
# 发送订阅调整事件
|
||||
await eventmanager.async_send_event(
|
||||
EventType.SubscribeModified,
|
||||
{
|
||||
"subscribe_id": subid,
|
||||
"old_subscribe_info": old_subscribe_dict,
|
||||
"subscribe_info": updated_subscribe.to_dict()
|
||||
SubscribeModifiedEventData(
|
||||
subscribe_id=subid,
|
||||
old_subscribe_info=old_subscribe_dict,
|
||||
subscribe_info=updated_subscribe.to_dict()
|
||||
if updated_subscribe
|
||||
else {},
|
||||
},
|
||||
scene="reset",
|
||||
).to_dict(),
|
||||
)
|
||||
return schemas.Response(success=True)
|
||||
return schemas.Response(success=False, message="订阅不存在")
|
||||
|
||||
|
||||
@router.get("/check", summary="刷新订阅 TMDB 信息", response_model=schemas.Response)
|
||||
def check_subscribes(_: schemas.TokenPayload = Depends(verify_token)) -> Any:
|
||||
def check_subscribes(
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
) -> Any:
|
||||
"""
|
||||
刷新订阅 TMDB 信息
|
||||
"""
|
||||
if not current_user.is_superuser:
|
||||
return schemas.Response(success=False, message="订阅不存在")
|
||||
Scheduler().start("subscribe_tmdb")
|
||||
return schemas.Response(success=True)
|
||||
|
||||
|
||||
@router.get("/search", summary="搜索所有订阅", response_model=schemas.Response)
|
||||
async def search_subscribes(
|
||||
background_tasks: BackgroundTasks, _: schemas.TokenPayload = Depends(verify_token)
|
||||
background_tasks: BackgroundTasks,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
搜索所有订阅
|
||||
"""
|
||||
background_tasks.add_task(
|
||||
Scheduler().start,
|
||||
job_id="subscribe_search",
|
||||
**{"sid": None, "state": "R", "manual": True},
|
||||
)
|
||||
if current_user.is_superuser:
|
||||
background_tasks.add_task(
|
||||
Scheduler().start,
|
||||
job_id="subscribe_search",
|
||||
**{"sid": None, "state": "R", "manual": True},
|
||||
)
|
||||
else:
|
||||
subscribes = await Subscribe.async_list_by_username(
|
||||
db, current_user.name, state="R"
|
||||
)
|
||||
for subscribe in subscribes:
|
||||
background_tasks.add_task(
|
||||
Scheduler().start,
|
||||
job_id="subscribe_search",
|
||||
**{"sid": subscribe.id, "state": None, "manual": True},
|
||||
)
|
||||
return schemas.Response(success=True)
|
||||
|
||||
|
||||
@@ -317,11 +405,15 @@ async def search_subscribes(
|
||||
async def search_subscribe(
|
||||
subscribe_id: int,
|
||||
background_tasks: BackgroundTasks,
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
根据订阅编号搜索订阅
|
||||
"""
|
||||
subscribe = await get_accessible_subscribe(db, subscribe_id, current_user)
|
||||
if not subscribe:
|
||||
return schemas.Response(success=False, message="订阅不存在")
|
||||
background_tasks.add_task(
|
||||
Scheduler().start,
|
||||
job_id="subscribe_search",
|
||||
@@ -335,31 +427,18 @@ async def delete_subscribe_by_mediaid(
|
||||
mediaid: str,
|
||||
season: Optional[int] = None,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
根据TMDBID或豆瓣ID删除订阅 tmdb:/douban:
|
||||
根据任意媒体数据源 ID 删除订阅。
|
||||
"""
|
||||
delete_subscribes = []
|
||||
if mediaid.startswith("tmdb:"):
|
||||
tmdbid = mediaid[5:]
|
||||
if not tmdbid or not str(tmdbid).isdigit():
|
||||
return schemas.Response(success=False)
|
||||
subscribes = await Subscribe.async_get_by_tmdbid(db, int(tmdbid), season)
|
||||
delete_subscribes.extend(subscribes)
|
||||
elif mediaid.startswith("douban:"):
|
||||
doubanid = mediaid[7:]
|
||||
if not doubanid:
|
||||
return schemas.Response(success=False)
|
||||
subscribe = await Subscribe.async_get_by_doubanid(db, doubanid)
|
||||
if subscribe:
|
||||
delete_subscribes.append(subscribe)
|
||||
else:
|
||||
subscribe = await Subscribe.async_get_by_mediaid(db, mediaid)
|
||||
if subscribe:
|
||||
delete_subscribes.append(subscribe)
|
||||
delete_subscribes = await list_subscribes_by_media_key(db, mediaid, season)
|
||||
delete_events = []
|
||||
for subscribe in delete_subscribes:
|
||||
for subscribe in [
|
||||
subscribe
|
||||
for subscribe in delete_subscribes
|
||||
if can_access_subscribe(subscribe, current_user)
|
||||
]:
|
||||
subscribe_info = build_subscribe_event_payload(subscribe)
|
||||
subscribe_id = subscribe_info.get("id")
|
||||
if not subscribe_id:
|
||||
@@ -425,7 +504,8 @@ async def seerr_subscribe(
|
||||
tmdbid=tmdbId,
|
||||
title=subject,
|
||||
year="",
|
||||
season=0,
|
||||
# 电影不传季号,避免被误判为剧集(S00)并污染通知标题
|
||||
season=None,
|
||||
username=user_name,
|
||||
)
|
||||
else:
|
||||
@@ -460,14 +540,19 @@ async def subscribe_history(
|
||||
page: Optional[int] = 1,
|
||||
count: Optional[int] = 30,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
查询电影/电视剧订阅历史
|
||||
"""
|
||||
histories = await SubscribeHistory.async_list_by_type(
|
||||
db, mtype=mtype, page=page, count=count
|
||||
)
|
||||
if current_user.is_superuser:
|
||||
histories = await SubscribeHistory.async_list_by_type(
|
||||
db, mtype=mtype, page=page, count=count
|
||||
)
|
||||
else:
|
||||
histories = await SubscribeHistory.async_list_by_type_and_username(
|
||||
db, mtype=mtype, username=current_user.name, page=page, count=count
|
||||
)
|
||||
result = []
|
||||
for history in histories:
|
||||
history_item = schemas.Subscribe.model_validate(history, from_attributes=True)
|
||||
@@ -484,12 +569,14 @@ async def subscribe_history(
|
||||
async def delete_subscribe_history(
|
||||
history_id: int,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
删除订阅历史
|
||||
"""
|
||||
await SubscribeHistory.async_delete(db, history_id)
|
||||
history = await SubscribeHistory.async_get(db, history_id)
|
||||
if can_access_subscribe(history, current_user):
|
||||
await SubscribeHistory.async_delete(db, history_id)
|
||||
return schemas.Response(success=True)
|
||||
|
||||
|
||||
@@ -534,7 +621,7 @@ async def popular_subscribes(
|
||||
# 处理标题
|
||||
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}季"
|
||||
@@ -542,6 +629,9 @@ async def popular_subscribes(
|
||||
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.media_id = sub.get("media_id")
|
||||
media.tvdb_id = sub.get("tvdbid")
|
||||
media.imdb_id = sub.get("imdbid")
|
||||
media.season = sub.get("season")
|
||||
@@ -561,11 +651,13 @@ async def popular_subscribes(
|
||||
async def user_subscribes(
|
||||
username: str,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
查询用户订阅
|
||||
"""
|
||||
if not current_user.is_superuser and username != current_user.name:
|
||||
return []
|
||||
return await Subscribe.async_list_by_username(db, username)
|
||||
|
||||
|
||||
@@ -577,12 +669,12 @@ async def user_subscribes(
|
||||
def subscribe_files(
|
||||
subscribe_id: int,
|
||||
db: Session = Depends(get_db),
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
) -> Any:
|
||||
"""
|
||||
订阅相关文件信息
|
||||
"""
|
||||
subscribe = Subscribe.get(db, subscribe_id)
|
||||
subscribe = get_accessible_subscribe_sync(db, subscribe_id, current_user)
|
||||
if subscribe:
|
||||
return SubscribeChain().subscribe_files_info(subscribe)
|
||||
return schemas.SubscrbieInfo()
|
||||
@@ -590,11 +682,16 @@ def subscribe_files(
|
||||
|
||||
@router.post("/share", summary="分享订阅", response_model=schemas.Response)
|
||||
async def subscribe_share(
|
||||
sub: schemas.SubscribeShare, _: schemas.TokenPayload = Depends(verify_token)
|
||||
sub: schemas.SubscribeShare,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
分享订阅
|
||||
"""
|
||||
subscribe = await get_accessible_subscribe(db, sub.subscribe_id, current_user)
|
||||
if not subscribe:
|
||||
return schemas.Response(success=False, message="订阅不存在")
|
||||
state, errmsg = await MoviePilotServerHelper.async_sub_share(
|
||||
subscribe_id=sub.subscribe_id,
|
||||
share_title=sub.share_title,
|
||||
@@ -724,26 +821,27 @@ async def subscribe_share_statistics(
|
||||
async def read_subscribe(
|
||||
subscribe_id: int,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
根据订阅编号查询订阅信息
|
||||
"""
|
||||
if not subscribe_id:
|
||||
return Subscribe()
|
||||
return await Subscribe.async_get(db, subscribe_id)
|
||||
subscribe = await get_accessible_subscribe(db, subscribe_id, current_user)
|
||||
return subscribe if subscribe else Subscribe()
|
||||
|
||||
|
||||
@router.delete("/{subscribe_id}", summary="删除订阅", response_model=schemas.Response)
|
||||
async def delete_subscribe(
|
||||
subscribe_id: int,
|
||||
db: AsyncSession = Depends(get_async_db),
|
||||
_: schemas.TokenPayload = Depends(verify_token),
|
||||
current_user: User = Depends(get_current_active_user_async),
|
||||
) -> Any:
|
||||
"""
|
||||
删除订阅信息
|
||||
"""
|
||||
subscribe = await Subscribe.async_get(db, subscribe_id)
|
||||
subscribe = await get_accessible_subscribe(db, subscribe_id, current_user)
|
||||
if subscribe:
|
||||
# 在删除之前获取订阅信息
|
||||
subscribe_info = build_subscribe_event_payload(subscribe)
|
||||
@@ -760,6 +858,14 @@ async def delete_subscribe(
|
||||
)
|
||||
# 统计订阅
|
||||
MoviePilotServerHelper.sub_done_async(
|
||||
{"tmdbid": subscribe_info.get("tmdbid"), "doubanid": subscribe_info.get("doubanid")}
|
||||
{
|
||||
"tmdbid": subscribe_info.get("tmdbid"),
|
||||
"doubanid": subscribe_info.get("doubanid"),
|
||||
"bangumiid": subscribe_info.get("bangumiid"),
|
||||
"anilistid": subscribe_info.get("anilistid"),
|
||||
"media_source": subscribe_info.get("media_source"),
|
||||
"media_id": subscribe_info.get("media_id"),
|
||||
"season": subscribe_info.get("season"),
|
||||
}
|
||||
)
|
||||
return schemas.Response(success=True)
|
||||
|
||||
@@ -35,10 +35,17 @@ from app.db.user_oper import (
|
||||
get_current_active_user_async,
|
||||
)
|
||||
from app.helper.image import ImageHelper
|
||||
from app.helper.locale import LocaleHelper
|
||||
from app.helper.market import (
|
||||
PLUGIN_MARKET_WIKI_URL,
|
||||
extract_plugin_market_repos_from_wiki,
|
||||
merge_plugin_market_repos,
|
||||
split_plugin_market_repo_urls,
|
||||
)
|
||||
from app.helper.message import MessageHelper
|
||||
from app.helper.server import MoviePilotServerHelper
|
||||
from app.helper.progress import ProgressHelper
|
||||
from app.helper.rule import RuleHelper
|
||||
from app.helper.server import MoviePilotServerHelper
|
||||
from app.helper.system import SystemHelper
|
||||
from app.log import logger
|
||||
from app.scheduler import Scheduler
|
||||
@@ -69,32 +76,47 @@ _PUBLIC_SYSTEM_CONFIG_KEYS = {
|
||||
_PUBLIC_SETTINGS_KEYS = {"PLUGIN_MARKET"}
|
||||
_LOG_DOWNLOAD_LIMIT = 10
|
||||
_LOG_DOWNLOAD_NAME_PATTERN = re.compile(r"^[A-Za-z0-9_-]+$")
|
||||
_PLUGIN_MARKET_WIKI_START = "<!-- plugin-market-repos:start -->"
|
||||
_PLUGIN_MARKET_WIKI_END = "<!-- plugin-market-repos:end -->"
|
||||
_PLUGIN_MARKET_WIKI_URL = "https://raw.githubusercontent.com/jxxghp/MoviePilot-Wiki/main/plugin.md"
|
||||
_PLUGIN_MARKET_REPO_PATTERN = re.compile(
|
||||
r"https?://github\.com/[A-Za-z0-9_.-]+/[A-Za-z0-9_.-]+(?:\.git)?/?",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def _normalize_plugin_market_repo_url(repo_url: str) -> Optional[str]:
|
||||
"""
|
||||
规范化插件仓库地址,便于跨来源合并去重。
|
||||
"""
|
||||
repo_url = (repo_url or "").strip().rstrip("/")
|
||||
if not repo_url:
|
||||
def _validate_llm_server_tool_config(env: dict) -> Optional[str]:
|
||||
"""校验强制服务端联网搜索配置,返回用户可读错误信息。"""
|
||||
from app.agent.llm.server_tools import (
|
||||
ServerToolRegistry,
|
||||
ServerToolUnavailableError,
|
||||
)
|
||||
|
||||
mode = ServerToolRegistry.normalize_web_search_mode(
|
||||
env.get(
|
||||
"LLM_WEB_SEARCH_MODE",
|
||||
getattr(settings, "LLM_WEB_SEARCH_MODE", "local"),
|
||||
)
|
||||
)
|
||||
if mode != "builtin":
|
||||
return None
|
||||
repo_url = repo_url.removesuffix(".git")
|
||||
parsed_url = urlparse(repo_url)
|
||||
if parsed_url.scheme not in {"http", "https"}:
|
||||
|
||||
provider = str(
|
||||
env.get("LLM_PROVIDER", getattr(settings, "LLM_PROVIDER", "")) or ""
|
||||
).strip()
|
||||
model = str(
|
||||
env.get("LLM_MODEL", getattr(settings, "LLM_MODEL", "")) or ""
|
||||
).strip()
|
||||
base_url = env.get("LLM_BASE_URL", getattr(settings, "LLM_BASE_URL", None))
|
||||
capability = ServerToolRegistry.get_capability(
|
||||
provider=provider,
|
||||
model=model,
|
||||
base_url=str(base_url or "").strip() or None,
|
||||
tool_id="web_search",
|
||||
)
|
||||
if capability:
|
||||
return None
|
||||
if (parsed_url.hostname or "").lower() != "github.com":
|
||||
return None
|
||||
paths = [item for item in parsed_url.path.split("/") if item]
|
||||
if len(paths) < 2:
|
||||
return None
|
||||
return f"https://github.com/{paths[0]}/{paths[1]}"
|
||||
|
||||
return str(
|
||||
ServerToolUnavailableError(
|
||||
provider=provider,
|
||||
model=model,
|
||||
tool_id="web_search",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _is_allowed_plugin_market_wiki_url(wiki_url: str) -> bool:
|
||||
@@ -114,55 +136,6 @@ def _is_allowed_plugin_market_wiki_url(wiki_url: str) -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _split_plugin_market_repo_urls(value: Optional[str]) -> list[str]:
|
||||
"""
|
||||
拆分插件市场仓库配置并保持原有顺序去重。
|
||||
"""
|
||||
repos: list[str] = []
|
||||
seen_repos = set()
|
||||
for item in re.split(r"[\n,,]+", value or ""):
|
||||
normalized_repo = _normalize_plugin_market_repo_url(item)
|
||||
if not normalized_repo or normalized_repo.lower() in seen_repos:
|
||||
continue
|
||||
repos.append(normalized_repo)
|
||||
seen_repos.add(normalized_repo.lower())
|
||||
return repos
|
||||
|
||||
|
||||
def _extract_plugin_market_repos_from_wiki(markdown: str) -> list[str]:
|
||||
"""
|
||||
从 Wiki 插件文档中提取插件仓库地址。
|
||||
"""
|
||||
content = markdown or ""
|
||||
if _PLUGIN_MARKET_WIKI_START in content and _PLUGIN_MARKET_WIKI_END in content:
|
||||
content = content.split(_PLUGIN_MARKET_WIKI_START, 1)[1].split(_PLUGIN_MARKET_WIKI_END, 1)[0]
|
||||
|
||||
repos: list[str] = []
|
||||
seen_repos = set()
|
||||
for item in _PLUGIN_MARKET_REPO_PATTERN.findall(content):
|
||||
normalized_repo = _normalize_plugin_market_repo_url(item)
|
||||
if not normalized_repo or normalized_repo.lower() in seen_repos:
|
||||
continue
|
||||
repos.append(normalized_repo)
|
||||
seen_repos.add(normalized_repo.lower())
|
||||
return repos
|
||||
|
||||
|
||||
def _merge_plugin_market_repos(local_repos: list[str], wiki_repos: list[str]) -> list[str]:
|
||||
"""
|
||||
合并本地与 Wiki 插件仓库地址,保留本地顺序并追加 Wiki 新地址。
|
||||
"""
|
||||
merged_repos: list[str] = []
|
||||
seen_repos = set()
|
||||
for repo in local_repos + wiki_repos:
|
||||
normalized_repo = _normalize_plugin_market_repo_url(repo)
|
||||
if not normalized_repo or normalized_repo.lower() in seen_repos:
|
||||
continue
|
||||
merged_repos.append(normalized_repo)
|
||||
seen_repos.add(normalized_repo.lower())
|
||||
return merged_repos
|
||||
|
||||
|
||||
def _match_nettest_prefix(url: str, prefix: str) -> bool:
|
||||
"""
|
||||
判断目标URL是否仍然落在允许的协议、主机、端口和路径前缀内。
|
||||
@@ -581,23 +554,27 @@ async def fetch_image(
|
||||
):
|
||||
return None
|
||||
|
||||
content = await ImageHelper().async_fetch_image(
|
||||
image_result = await ImageHelper().async_fetch_image_with_mime_type(
|
||||
url=fetch_url,
|
||||
proxy=proxy,
|
||||
use_cache=use_cache,
|
||||
cookies=cookies,
|
||||
)
|
||||
|
||||
if content:
|
||||
if image_result:
|
||||
content, media_type = image_result
|
||||
|
||||
# 检查 If-None-Match
|
||||
etag = HashUtils.md5(content)
|
||||
headers = RequestUtils.generate_cache_headers(etag, max_age=86400 * 7)
|
||||
headers["Content-Type"] = media_type
|
||||
headers["X-Content-Type-Options"] = "nosniff"
|
||||
if if_none_match == etag:
|
||||
return Response(status_code=304, headers=headers)
|
||||
# 返回缓存图片
|
||||
return Response(
|
||||
content=content,
|
||||
media_type=UrlUtils.get_mime_type(fetch_url, "image/jpeg"),
|
||||
media_type=media_type,
|
||||
headers=headers,
|
||||
)
|
||||
return None
|
||||
@@ -694,7 +671,6 @@ async def get_user_global_setting(_: User = Depends(get_current_active_user_asyn
|
||||
"RECOGNIZE_SOURCE",
|
||||
"SEARCH_SOURCE",
|
||||
"AI_RECOMMEND_ENABLED",
|
||||
"PASSKEY_ALLOW_REGISTER_WITHOUT_OTP",
|
||||
}
|
||||
)
|
||||
# 智能助手总开关未开启,智能推荐状态强制返回False
|
||||
@@ -759,6 +735,10 @@ async def set_env_setting(
|
||||
"""
|
||||
更新系统环境变量(仅管理员)
|
||||
"""
|
||||
validation_error = _validate_llm_server_tool_config(env)
|
||||
if validation_error:
|
||||
return schemas.Response(success=False, message=validation_error)
|
||||
|
||||
result = settings.update_settings(env=env)
|
||||
# 统计成功和失败的结果
|
||||
success_updates = {k: v for k, v in result.items() if v[0]}
|
||||
@@ -797,13 +777,14 @@ async def get_progress(
|
||||
实时获取处理进度,返回格式为SSE
|
||||
"""
|
||||
progress = ProgressHelper(process_type)
|
||||
locale = LocaleHelper.get_current_locale()
|
||||
|
||||
async def event_generator():
|
||||
try:
|
||||
while not global_vars.is_system_stopped:
|
||||
if await request.is_disconnected():
|
||||
break
|
||||
detail = progress.get()
|
||||
detail = progress.get(locale=locale)
|
||||
yield f"data: {json.dumps(detail)}\n\n"
|
||||
await asyncio.sleep(0.5)
|
||||
except asyncio.CancelledError:
|
||||
@@ -839,7 +820,7 @@ async def sync_plugin_market_from_wiki(
|
||||
"""
|
||||
从 Wiki 插件文档同步插件市场仓库地址。
|
||||
"""
|
||||
wiki_url = (request.wiki_url if request else None) or _PLUGIN_MARKET_WIKI_URL
|
||||
wiki_url = (request.wiki_url if request else None) or PLUGIN_MARKET_WIKI_URL
|
||||
wiki_url = wiki_url.strip()
|
||||
if not _is_allowed_plugin_market_wiki_url(wiki_url):
|
||||
return schemas.Response(success=False, message="不支持的 Wiki 同步地址")
|
||||
@@ -859,14 +840,14 @@ async def sync_plugin_market_from_wiki(
|
||||
message=f"访问 Wiki 插件仓库清单失败,状态码:{res.status_code}",
|
||||
)
|
||||
|
||||
wiki_repos = _extract_plugin_market_repos_from_wiki(res.text)
|
||||
wiki_repos = extract_plugin_market_repos_from_wiki(res.text)
|
||||
if not wiki_repos:
|
||||
return schemas.Response(success=False, message="未在 Wiki 中识别到插件仓库地址")
|
||||
|
||||
local_repos = _split_plugin_market_repo_urls(settings.PLUGIN_MARKET)
|
||||
local_repos = split_plugin_market_repo_urls(settings.PLUGIN_MARKET)
|
||||
local_repo_keys = {repo.lower() for repo in local_repos}
|
||||
added_count = len([repo for repo in wiki_repos if repo.lower() not in local_repo_keys])
|
||||
merged_repos = _merge_plugin_market_repos(local_repos, wiki_repos)
|
||||
merged_repos = merge_plugin_market_repos(local_repos, wiki_repos)
|
||||
merged_value = ",".join(merged_repos)
|
||||
|
||||
success, message = settings.update_setting("PLUGIN_MARKET", merged_value)
|
||||
@@ -1121,33 +1102,64 @@ def ruletest(
|
||||
"""
|
||||
过滤规则测试,规则类型 1-订阅,2-洗版,3-搜索
|
||||
"""
|
||||
metainfo = MetaInfo(title=title, subtitle=subtitle)
|
||||
torrent = schemas.TorrentInfo(
|
||||
title=title,
|
||||
description=subtitle,
|
||||
)
|
||||
# 查询规则组详情
|
||||
rulegroup = RuleHelper().get_rule_group(rulegroup_name)
|
||||
result_data = {
|
||||
"title": title,
|
||||
"subtitle": subtitle,
|
||||
"rulegroup_name": rulegroup_name,
|
||||
"rulegroup": rulegroup.model_dump() if rulegroup else None,
|
||||
"meta_info": metainfo.to_dict(),
|
||||
"media_info": None,
|
||||
"torrent_info": torrent.model_dump(),
|
||||
"priority": None,
|
||||
"matched": False,
|
||||
}
|
||||
if not rulegroup:
|
||||
return schemas.Response(
|
||||
success=False, message=f"过滤规则组 {rulegroup_name} 不存在!"
|
||||
success=False,
|
||||
message=f"过滤规则组 {rulegroup_name} 不存在!",
|
||||
data=result_data,
|
||||
)
|
||||
|
||||
# 根据标题查询媒体信息
|
||||
media_info = MediaChain().recognize_by_meta(
|
||||
MetaInfo(title=title, subtitle=subtitle),
|
||||
metainfo,
|
||||
obtain_images=False,
|
||||
)
|
||||
result_data["media_info"] = media_info.to_dict() if media_info else None
|
||||
if not media_info:
|
||||
return schemas.Response(success=False, message="未识别到媒体信息!")
|
||||
return schemas.Response(
|
||||
success=False,
|
||||
message="未识别到媒体信息!",
|
||||
data=result_data,
|
||||
)
|
||||
|
||||
# 过滤
|
||||
result = SearchChain().filter_torrents(
|
||||
rule_groups=[rulegroup.name], torrent_list=[torrent], mediainfo=media_info
|
||||
)
|
||||
if not result:
|
||||
return schemas.Response(success=False, message="不符合过滤规则!")
|
||||
return schemas.Response(
|
||||
success=False,
|
||||
message="不符合过滤规则!",
|
||||
data=result_data,
|
||||
)
|
||||
result_data.update(
|
||||
{
|
||||
"matched": True,
|
||||
"priority": 100 - result[0].pri_order + 1,
|
||||
"torrent_info": result[0].model_dump(),
|
||||
}
|
||||
)
|
||||
return schemas.Response(
|
||||
success=True, data={"priority": 100 - result[0].pri_order + 1}
|
||||
success=True,
|
||||
data=result_data,
|
||||
)
|
||||
|
||||
|
||||
@@ -1271,13 +1283,20 @@ def modulelist(_: schemas.TokenPayload = Depends(verify_token)):
|
||||
"""
|
||||
查询已加载的模块ID列表
|
||||
"""
|
||||
modules = [
|
||||
{
|
||||
"id": k,
|
||||
"name": v.get_name(),
|
||||
}
|
||||
for k, v in ModuleManager().get_modules().items()
|
||||
]
|
||||
modules = []
|
||||
for module_id, module in ModuleManager().get_modules().items():
|
||||
name = module.get_name()
|
||||
modules.append(
|
||||
{
|
||||
"id": module_id,
|
||||
"name": name,
|
||||
"name_i18n": LocaleHelper.translate(
|
||||
f"system.modules.{module_id}.name",
|
||||
default=name,
|
||||
),
|
||||
"name_key": f"system.modules.{module_id}.name",
|
||||
}
|
||||
)
|
||||
return schemas.Response(success=True, data={"modules": modules})
|
||||
|
||||
|
||||
@@ -1299,12 +1318,7 @@ def restart_system(_: User = Depends(get_current_active_superuser)):
|
||||
"""
|
||||
if not SystemHelper.can_restart():
|
||||
return schemas.Response(success=False, message="当前运行环境不支持重启操作!")
|
||||
# 标识停止事件
|
||||
global_vars.stop_system()
|
||||
# 执行重启
|
||||
ret, msg = SystemHelper.restart()
|
||||
if not ret:
|
||||
global_vars.resume_system()
|
||||
return schemas.Response(success=ret, message=msg)
|
||||
|
||||
|
||||
@@ -1322,11 +1336,7 @@ def upgrade_system(
|
||||
if not SystemHelper.can_restart():
|
||||
return schemas.Response(success=False, message="当前运行环境不支持升级操作!")
|
||||
|
||||
# 标识停止事件
|
||||
global_vars.stop_system()
|
||||
ret, msg = SystemHelper.upgrade(mode=mode or "release")
|
||||
if not ret:
|
||||
global_vars.resume_system()
|
||||
return schemas.Response(success=ret, message=msg)
|
||||
|
||||
|
||||
|
||||
@@ -4,12 +4,68 @@ from fastapi import APIRouter, Depends
|
||||
|
||||
from app import schemas
|
||||
from app.chain.tmdb import TmdbChain
|
||||
from app.core.config import settings
|
||||
from app.core.security import verify_token
|
||||
from app.schemas.types import MediaType
|
||||
from app.db.models.user import User
|
||||
from app.db.systemconfig_oper import SystemConfigOper
|
||||
from app.db.user_oper import get_current_active_superuser_async
|
||||
from app.modules.themoviedb.tmdb_cache import TmdbCache
|
||||
from app.schemas.types import MediaType, SystemConfigKey
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.get(
|
||||
"/cache", summary="查询 TheMovieDb 识别缓存", response_model=schemas.Response
|
||||
)
|
||||
async def tmdb_recognition_cache(
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
) -> schemas.Response:
|
||||
"""查询可管理的 TheMovieDb 识别缓存。"""
|
||||
cache_items = TmdbCache().list_items()
|
||||
recognized_count = sum(1 for item in cache_items if item["tmdb_id"])
|
||||
return schemas.Response(
|
||||
success=True,
|
||||
data={
|
||||
"count": len(cache_items),
|
||||
"recognized": recognized_count,
|
||||
"unrecognized": len(cache_items) - recognized_count,
|
||||
"shared_recognized": SystemConfigOper().get(
|
||||
SystemConfigKey.MediaRecognizeShareCount
|
||||
) or 0,
|
||||
"shared_recognize_enabled": settings.MEDIA_RECOGNIZE_SHARE,
|
||||
"data": cache_items,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/cache/{cache_key:path}",
|
||||
summary="删除指定 TheMovieDb 识别缓存",
|
||||
response_model=schemas.Response,
|
||||
)
|
||||
async def delete_tmdb_recognition_cache(
|
||||
cache_key: str,
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
) -> schemas.Response:
|
||||
"""按缓存键删除单条 TheMovieDb 识别缓存。"""
|
||||
deleted_item = TmdbCache().delete(cache_key)
|
||||
if not deleted_item:
|
||||
return schemas.Response(success=False, message="TheMovieDb 识别缓存不存在")
|
||||
return schemas.Response(success=True, message="TheMovieDb 识别缓存删除成功")
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/cache", summary="清空 TheMovieDb 识别缓存", response_model=schemas.Response
|
||||
)
|
||||
async def clear_tmdb_recognition_cache(
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
) -> schemas.Response:
|
||||
"""清空全部 TheMovieDb 识别缓存。"""
|
||||
TmdbCache().clear()
|
||||
return schemas.Response(success=True, message="TheMovieDb 识别缓存清理完成")
|
||||
|
||||
|
||||
@router.get(
|
||||
"/seasons/{tmdbid}", summary="TMDB所有季", response_model=List[schemas.TmdbSeason]
|
||||
)
|
||||
|
||||
@@ -174,6 +174,10 @@ async def reidentify_cache(
|
||||
torrent_hash: str,
|
||||
tmdbid: Optional[int] = None,
|
||||
doubanid: Optional[str] = None,
|
||||
bangumiid: Optional[int] = None,
|
||||
anilistid: Optional[int] = None,
|
||||
media_source: Optional[str] = None,
|
||||
media_id: Optional[str] = None,
|
||||
_: User = Depends(get_current_active_superuser_async),
|
||||
):
|
||||
"""
|
||||
@@ -182,6 +186,10 @@ async def reidentify_cache(
|
||||
:param torrent_hash: 种子hash(使用title+description的md5)
|
||||
:param tmdbid: 手动指定的TMDB ID
|
||||
:param doubanid: 手动指定的豆瓣ID
|
||||
:param bangumiid: 手动指定的 Bangumi ID
|
||||
:param anilistid: 手动指定的 AniList ID
|
||||
:param media_source: 媒体数据源
|
||||
:param media_id: 数据源原生 ID
|
||||
:param _: 当前用户,必须是超级用户
|
||||
"""
|
||||
|
||||
@@ -215,10 +223,16 @@ async def reidentify_cache(
|
||||
title=target_context.torrent_info.title,
|
||||
subtitle=target_context.torrent_info.description,
|
||||
)
|
||||
if tmdbid or doubanid:
|
||||
if tmdbid or doubanid or bangumiid or anilistid or media_source or media_id:
|
||||
# 手动指定媒体信息
|
||||
mediainfo = await media_chain.async_recognize_media(
|
||||
meta=meta, tmdbid=tmdbid, doubanid=doubanid
|
||||
meta=meta,
|
||||
tmdbid=tmdbid,
|
||||
doubanid=doubanid,
|
||||
bangumiid=bangumiid,
|
||||
anilistid=anilistid,
|
||||
source=media_source,
|
||||
mediaid=media_id,
|
||||
)
|
||||
else:
|
||||
# 自动重新识别
|
||||
|
||||
@@ -6,14 +6,16 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from app import schemas
|
||||
from app.chain.media import MediaChain
|
||||
from app.chain.storage import StorageChain
|
||||
from app.chain.transfer import TransferChain
|
||||
from app.core.config import settings, global_vars
|
||||
from app.core.security import verify_token, verify_apitoken
|
||||
from app.db import get_db
|
||||
from app.db.models import User
|
||||
from app.db.models.transferhistory import TransferHistory
|
||||
from app.db.user_oper import get_current_active_superuser
|
||||
from app.db.user_oper import (
|
||||
get_current_active_manage_user,
|
||||
get_current_active_superuser,
|
||||
)
|
||||
from app.helper.directory import DirectoryHelper
|
||||
from app.log import logger
|
||||
from app.schemas import (
|
||||
@@ -183,7 +185,7 @@ def _get_manual_transfer_target_key(
|
||||
def match_manual_transfer_target_path(
|
||||
transer_item: ManualTransferItem,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
根据源文件匹配手动整理目的路径。
|
||||
@@ -238,12 +240,46 @@ def match_manual_transfer_target_path(
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/manual/history",
|
||||
summary="查询手动转移成功历史",
|
||||
response_model=schemas.Response,
|
||||
)
|
||||
def query_manual_transfer_history(
|
||||
transer_item: ManualTransferItem,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
查询文件或目录命中的成功整理记录。
|
||||
|
||||
:param transer_item: 手工整理项
|
||||
:param db: 数据库
|
||||
:param _: Token校验
|
||||
"""
|
||||
src_fileitems, error_message = _resolve_manual_transfer_source_fileitems(
|
||||
transer_item=transer_item,
|
||||
db=db,
|
||||
)
|
||||
if error_message:
|
||||
return schemas.Response(success=False, message=error_message)
|
||||
|
||||
histories = TransferChain().get_manual_transfer_histories(
|
||||
_deduplicate_fileitems(src_fileitems)
|
||||
)
|
||||
history_info = schemas.ManualTransferHistoryInfo(
|
||||
reorganize=bool(histories),
|
||||
history_count=len(histories),
|
||||
)
|
||||
return schemas.Response(success=True, data=history_info.model_dump())
|
||||
|
||||
|
||||
@router.post("/manual", summary="手动转移", response_model=schemas.Response)
|
||||
def manual_transfer(
|
||||
transer_item: ManualTransferItem,
|
||||
background: Optional[bool] = False,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
手动转移,文件或历史记录,支持自定义剧集识别格式
|
||||
@@ -256,6 +292,7 @@ def manual_transfer(
|
||||
downloader = None
|
||||
download_hash = None
|
||||
src_fileitems: List[FileItem] = []
|
||||
cleanup_dest_fileitem: Optional[FileItem] = None
|
||||
target_path = Path(transer_item.target_path) if transer_item.target_path else None
|
||||
if transer_item.logid:
|
||||
# 查询历史记录
|
||||
@@ -266,23 +303,21 @@ def manual_transfer(
|
||||
)
|
||||
# 强制转移
|
||||
force = True
|
||||
downloader = history.downloader
|
||||
download_hash = history.download_hash
|
||||
# 下载器与 Hash 是同一组下载上下文,重新识别时由当前文件路径重新匹配。
|
||||
downloader = history.downloader if transer_item.from_history else None
|
||||
download_hash = history.download_hash if transer_item.from_history else None
|
||||
if history.status and ("move" in history.mode):
|
||||
# 重新整理成功的转移,则使用成功的 dest 做 in_path
|
||||
src_fileitems = [FileItem(**history.dest_fileitem)]
|
||||
else:
|
||||
# 源路径
|
||||
src_fileitems = [FileItem(**history.src_fileitem)]
|
||||
# 目的路径
|
||||
if history.dest_fileitem and not transer_item.preview:
|
||||
# 删除旧的已整理文件
|
||||
dest_fileitem = FileItem(**history.dest_fileitem)
|
||||
state = StorageChain().delete_media_file(dest_fileitem)
|
||||
if not state:
|
||||
return schemas.Response(
|
||||
success=False, message=f"{dest_fileitem.path} 删除失败"
|
||||
)
|
||||
if (
|
||||
history.dest_fileitem
|
||||
and not transer_item.preview
|
||||
and not transer_item.reorganize
|
||||
):
|
||||
cleanup_dest_fileitem = FileItem(**history.dest_fileitem)
|
||||
|
||||
# 从历史数据获取信息
|
||||
if transer_item.from_history:
|
||||
@@ -295,6 +330,14 @@ def manual_transfer(
|
||||
transer_item.doubanid = (
|
||||
str(history.doubanid) if history.doubanid else transer_item.doubanid
|
||||
)
|
||||
transer_item.bangumiid = history.bangumiid or transer_item.bangumiid
|
||||
transer_item.anilistid = history.anilistid or transer_item.anilistid
|
||||
transer_item.media_source = (
|
||||
history.media_source or transer_item.media_source
|
||||
)
|
||||
transer_item.media_id = (
|
||||
history.media_id or transer_item.media_id
|
||||
)
|
||||
transer_item.season = (
|
||||
int(str(history.seasons).replace("S", ""))
|
||||
if history.seasons
|
||||
@@ -412,6 +455,10 @@ def manual_transfer(
|
||||
target_path=target_path,
|
||||
tmdbid=transer_item.tmdbid,
|
||||
doubanid=transer_item.doubanid,
|
||||
bangumiid=transer_item.bangumiid,
|
||||
anilistid=transer_item.anilistid,
|
||||
media_source=transer_item.media_source,
|
||||
media_id=transer_item.media_id,
|
||||
mtype=mtype,
|
||||
season=transer_item.season,
|
||||
episode_group=transer_item.episode_group,
|
||||
@@ -426,7 +473,9 @@ def manual_transfer(
|
||||
downloader=downloader,
|
||||
download_hash=download_hash,
|
||||
preview=transer_item.preview,
|
||||
reorganize=transer_item.reorganize,
|
||||
sync_extra_files=False,
|
||||
cleanup_dest_fileitem=cleanup_dest_fileitem,
|
||||
)
|
||||
if transer_item.preview:
|
||||
if isinstance(errormsg, dict):
|
||||
@@ -493,6 +542,10 @@ def manual_transfer(
|
||||
target_path=target_path,
|
||||
tmdbid=transer_item.tmdbid,
|
||||
doubanid=transer_item.doubanid,
|
||||
bangumiid=transer_item.bangumiid,
|
||||
anilistid=transer_item.anilistid,
|
||||
media_source=transer_item.media_source,
|
||||
media_id=transer_item.media_id,
|
||||
mtype=mtype,
|
||||
season=transer_item.season,
|
||||
episode_group=transer_item.episode_group,
|
||||
@@ -507,7 +560,9 @@ def manual_transfer(
|
||||
downloader=downloader,
|
||||
download_hash=download_hash,
|
||||
preview=transer_item.preview,
|
||||
reorganize=transer_item.reorganize,
|
||||
sync_extra_files=True,
|
||||
cleanup_dest_fileitem=cleanup_dest_fileitem,
|
||||
)
|
||||
# 失败
|
||||
if not state:
|
||||
@@ -533,7 +588,7 @@ def manual_transfer(
|
||||
)
|
||||
def recommend_episode_format(
|
||||
recommend_item: EpisodeFormatRecommendItem,
|
||||
_: User = Depends(get_current_active_superuser),
|
||||
_: User = Depends(get_current_active_manage_user),
|
||||
) -> Any:
|
||||
"""
|
||||
根据目录样本推荐集数定位模板
|
||||
|
||||
@@ -119,7 +119,7 @@ async def upload_avatar(
|
||||
if not user:
|
||||
return schemas.Response(success=False, message="用户不存在")
|
||||
await user.async_update(db, {"avatar": f"data:image/ico;base64,{file_base64}"})
|
||||
return schemas.Response(success=True, message=file.filename)
|
||||
return schemas.Response(success=True, data={"filename": file.filename})
|
||||
|
||||
|
||||
@router.get("/config/{key}", summary="查询用户配置", response_model=schemas.Response)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user