mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 09:40:12 +00:00
Merge branch 'upstream/master' into master and preserve domain extensions
This commit is contained in:
@@ -0,0 +1,21 @@
|
|||||||
|
version: 2
|
||||||
|
updates:
|
||||||
|
- package-ecosystem: "gomod"
|
||||||
|
directory: "/"
|
||||||
|
schedule:
|
||||||
|
interval: "weekly"
|
||||||
|
open-pull-requests-limit: 10
|
||||||
|
groups:
|
||||||
|
go-dependencies:
|
||||||
|
patterns:
|
||||||
|
- "*"
|
||||||
|
|
||||||
|
- package-ecosystem: "github-actions"
|
||||||
|
directory: "/"
|
||||||
|
schedule:
|
||||||
|
interval: "weekly"
|
||||||
|
open-pull-requests-limit: 5
|
||||||
|
groups:
|
||||||
|
github-actions:
|
||||||
|
patterns:
|
||||||
|
- "*"
|
||||||
@@ -36,7 +36,7 @@ jobs:
|
|||||||
|
|
||||||
- run: git config --global --add safe.directory /__w/nezha/nezha
|
- run: git config --global --add safe.directory /__w/nezha/nezha
|
||||||
|
|
||||||
- uses: actions/checkout@v4
|
- uses: actions/checkout@v7.0.1
|
||||||
|
|
||||||
- name: Prepare frontends' dists
|
- name: Prepare frontends' dists
|
||||||
run: |
|
run: |
|
||||||
@@ -50,7 +50,7 @@ jobs:
|
|||||||
wget -qO pkg/geoip/geoip.db https://ipinfo.io/data/free/country.mmdb?token=${IPINFO_TOKEN}
|
wget -qO pkg/geoip/geoip.db https://ipinfo.io/data/free/country.mmdb?token=${IPINFO_TOKEN}
|
||||||
|
|
||||||
- name: Set up Go
|
- name: Set up Go
|
||||||
uses: actions/setup-go@v5
|
uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: "1.26.x"
|
go-version: "1.26.x"
|
||||||
|
|
||||||
@@ -63,7 +63,7 @@ jobs:
|
|||||||
- name: Cache zstd for s390x
|
- name: Cache zstd for s390x
|
||||||
if: matrix.goarch == 's390x'
|
if: matrix.goarch == 's390x'
|
||||||
id: cache-zstd
|
id: cache-zstd
|
||||||
uses: actions/cache@v4
|
uses: actions/cache@v6
|
||||||
with:
|
with:
|
||||||
path: /tmp/zstd-s390x
|
path: /tmp/zstd-s390x
|
||||||
key: zstd-s390x-v1.5.7
|
key: zstd-s390x-v1.5.7
|
||||||
@@ -99,7 +99,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Build with tag
|
- name: Build with tag
|
||||||
if: contains(github.ref, 'refs/tags/')
|
if: contains(github.ref, 'refs/tags/')
|
||||||
uses: goreleaser/goreleaser-action@v6
|
uses: goreleaser/goreleaser-action@v7
|
||||||
env:
|
env:
|
||||||
GOOS: ${{ matrix.goos }}
|
GOOS: ${{ matrix.goos }}
|
||||||
GOARCH: ${{ matrix.goarch }}
|
GOARCH: ${{ matrix.goarch }}
|
||||||
@@ -111,7 +111,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Build snapshot
|
- name: Build snapshot
|
||||||
if: contains(github.ref, 'refs/tags/') == false
|
if: contains(github.ref, 'refs/tags/') == false
|
||||||
uses: goreleaser/goreleaser-action@v6
|
uses: goreleaser/goreleaser-action@v7
|
||||||
env:
|
env:
|
||||||
GOOS: ${{ matrix.goos }}
|
GOOS: ${{ matrix.goos }}
|
||||||
GOARCH: ${{ matrix.goarch }}
|
GOARCH: ${{ matrix.goarch }}
|
||||||
@@ -122,7 +122,7 @@ jobs:
|
|||||||
args: build --single-target --clean --skip=validate --snapshot
|
args: build --single-target --clean --skip=validate --snapshot
|
||||||
|
|
||||||
- name: Upload artifacts
|
- name: Upload artifacts
|
||||||
uses: actions/upload-artifact@v4
|
uses: actions/upload-artifact@v7
|
||||||
with:
|
with:
|
||||||
name: dashboard-${{ matrix.goos }}-${{ matrix.goarch }}
|
name: dashboard-${{ matrix.goos }}-${{ matrix.goarch }}
|
||||||
path: |
|
path: |
|
||||||
@@ -135,7 +135,7 @@ jobs:
|
|||||||
name: Release
|
name: Release
|
||||||
steps:
|
steps:
|
||||||
- name: Download artifacts
|
- name: Download artifacts
|
||||||
uses: actions/download-artifact@v4
|
uses: actions/download-artifact@v8
|
||||||
with:
|
with:
|
||||||
path: ./assets
|
path: ./assets
|
||||||
|
|
||||||
@@ -149,7 +149,7 @@ jobs:
|
|||||||
done
|
done
|
||||||
|
|
||||||
- name: Release
|
- name: Release
|
||||||
uses: softprops/action-gh-release@v2
|
uses: softprops/action-gh-release@v3
|
||||||
with:
|
with:
|
||||||
files: "assets/*/*/*.zip"
|
files: "assets/*/*/*.zip"
|
||||||
generate_release_notes: true
|
generate_release_notes: true
|
||||||
@@ -181,10 +181,10 @@ jobs:
|
|||||||
needs: build
|
needs: build
|
||||||
name: Release Docker images
|
name: Release Docker images
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v4
|
- uses: actions/checkout@v7.0.1
|
||||||
|
|
||||||
- name: Download artifacts
|
- name: Download artifacts
|
||||||
uses: actions/download-artifact@v4
|
uses: actions/download-artifact@v8
|
||||||
with:
|
with:
|
||||||
path: ./assets
|
path: ./assets
|
||||||
|
|
||||||
@@ -220,10 +220,10 @@ jobs:
|
|||||||
password: ${{ secrets.ALI_PAT }}
|
password: ${{ secrets.ALI_PAT }}
|
||||||
|
|
||||||
- name: Set up QEMU
|
- name: Set up QEMU
|
||||||
uses: docker/setup-qemu-action@v3
|
uses: docker/setup-qemu-action@v4
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@v3
|
uses: docker/setup-buildx-action@v4
|
||||||
|
|
||||||
- name: Set up image name
|
- name: Set up image name
|
||||||
run: |
|
run: |
|
||||||
@@ -238,11 +238,16 @@ jobs:
|
|||||||
|
|
||||||
- name: Build dasbboard image And Push with tag
|
- name: Build dasbboard image And Push with tag
|
||||||
if: contains(github.ref, 'refs/tags/')
|
if: contains(github.ref, 'refs/tags/')
|
||||||
uses: docker/build-push-action@v5
|
uses: docker/build-push-action@v7
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: ./Dockerfile
|
file: ./Dockerfile
|
||||||
platforms: linux/amd64,linux/arm64,linux/s390x
|
platforms: linux/amd64,linux/arm64,linux/s390x
|
||||||
|
# Aliyun Container Registry rejects BuildKit attestation manifests
|
||||||
|
# (application/vnd.oci.empty.v1+json). Both registries share this
|
||||||
|
# multi-platform push, so publish a plain OCI image index here.
|
||||||
|
provenance: false
|
||||||
|
sbom: false
|
||||||
push: true
|
push: true
|
||||||
tags: |
|
tags: |
|
||||||
${{ steps.image-name.outputs.GHCR_IMAGE_NAME }}:latest
|
${{ steps.image-name.outputs.GHCR_IMAGE_NAME }}:latest
|
||||||
@@ -252,7 +257,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Build dasbboard image And Push snapshot
|
- name: Build dasbboard image And Push snapshot
|
||||||
if: contains(github.ref, 'refs/tags/') == false
|
if: contains(github.ref, 'refs/tags/') == false
|
||||||
uses: docker/build-push-action@v5
|
uses: docker/build-push-action@v7
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: ./Dockerfile
|
file: ./Dockerfile
|
||||||
|
|||||||
+117
-26
@@ -2,54 +2,145 @@ name: Run Tests
|
|||||||
|
|
||||||
on:
|
on:
|
||||||
push:
|
push:
|
||||||
paths:
|
branches:
|
||||||
- "**.go"
|
- master
|
||||||
- "go.mod"
|
|
||||||
- "go.sum"
|
|
||||||
- "resource/**"
|
|
||||||
- ".github/workflows/test.yml"
|
|
||||||
pull_request:
|
pull_request:
|
||||||
branches:
|
branches:
|
||||||
- master
|
- master
|
||||||
|
merge_group:
|
||||||
|
|
||||||
|
permissions:
|
||||||
|
contents: read
|
||||||
|
|
||||||
|
concurrency:
|
||||||
|
group: nezha-quality-${{ github.workflow }}-${{ github.ref }}
|
||||||
|
cancel-in-progress: true
|
||||||
|
|
||||||
jobs:
|
jobs:
|
||||||
tests:
|
tests:
|
||||||
|
name: Ordinary tests and build (${{ matrix.os }})
|
||||||
strategy:
|
strategy:
|
||||||
fail-fast: true
|
fail-fast: false
|
||||||
matrix:
|
matrix:
|
||||||
os: [ubuntu, windows, macos]
|
os: [ubuntu-latest, windows-latest, macos-latest]
|
||||||
|
runs-on: ${{ matrix.os }}
|
||||||
runs-on: ${{ matrix.os }}-latest
|
timeout-minutes: 30
|
||||||
env:
|
|
||||||
GO111MODULE: on
|
|
||||||
FORCE_JAVASCRIPT_ACTIONS_TO_NODE24: true
|
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v4
|
- uses: actions/checkout@v7.0.1
|
||||||
|
with:
|
||||||
- uses: actions/setup-go@v5
|
persist-credentials: false
|
||||||
|
- uses: actions/setup-go@v7
|
||||||
with:
|
with:
|
||||||
go-version: "1.26.x"
|
go-version: "1.26.x"
|
||||||
|
cache: false
|
||||||
|
|
||||||
- name: generate swagger docs
|
- name: Generate Swagger docs
|
||||||
shell: bash
|
|
||||||
run: |
|
run: |
|
||||||
go install github.com/swaggo/swag/cmd/swag@latest
|
go install github.com/swaggo/swag/cmd/swag@v1.16.6
|
||||||
mkdir -p ./cmd/dashboard/user-dist ./cmd/dashboard/admin-dist
|
|
||||||
touch ./cmd/dashboard/user-dist/a
|
touch ./cmd/dashboard/user-dist/a
|
||||||
touch ./cmd/dashboard/admin-dist/a
|
touch ./cmd/dashboard/admin-dist/a
|
||||||
swag init --pd -d cmd/dashboard -g main.go -o cmd/dashboard/docs
|
swag init --pd -d cmd/dashboard -g main.go -o cmd/dashboard/docs
|
||||||
|
|
||||||
- name: Unit test
|
- name: Unit test
|
||||||
run: |
|
run: go test -mod=readonly -count=1 ./...
|
||||||
go test -v ./...
|
|
||||||
|
|
||||||
- name: Build test
|
- name: Build dashboard
|
||||||
run: go build -v ./cmd/dashboard
|
run: go build -v ./cmd/dashboard
|
||||||
|
|
||||||
|
linux-race-quality:
|
||||||
|
name: Linux race and quality
|
||||||
|
runs-on: ubuntu-24.04
|
||||||
|
timeout-minutes: 45
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v7.0.1
|
||||||
|
with:
|
||||||
|
persist-credentials: false
|
||||||
|
- uses: actions/setup-go@v7
|
||||||
|
with:
|
||||||
|
go-version: "1.26.x"
|
||||||
|
cache: false
|
||||||
|
|
||||||
|
- name: Generate Swagger docs
|
||||||
|
run: |
|
||||||
|
go install github.com/swaggo/swag/cmd/swag@v1.16.6
|
||||||
|
touch ./cmd/dashboard/user-dist/a
|
||||||
|
touch ./cmd/dashboard/admin-dist/a
|
||||||
|
swag init --pd -d cmd/dashboard -g main.go -o cmd/dashboard/docs
|
||||||
|
|
||||||
|
- name: Race and shuffle tests
|
||||||
|
run: go test -mod=readonly -race -shuffle=on -count=1 ./...
|
||||||
|
|
||||||
|
- name: Vet
|
||||||
|
run: go vet ./...
|
||||||
|
|
||||||
|
- name: Check formatting
|
||||||
|
shell: bash
|
||||||
|
run: test -z "$(git ls-files -co --exclude-standard '*.go' -z | xargs -0 gofmt -l)"
|
||||||
|
|
||||||
|
- name: Build dashboard
|
||||||
|
run: go build ./cmd/dashboard
|
||||||
|
|
||||||
- name: Run Gosec Security Scanner
|
- name: Run Gosec Security Scanner
|
||||||
if: runner.os == 'Linux'
|
shell: bash
|
||||||
uses: securego/gosec@master
|
|
||||||
env:
|
env:
|
||||||
GOTOOLCHAIN: auto
|
GOTOOLCHAIN: auto
|
||||||
|
run: |
|
||||||
|
go install github.com/securego/gosec/v2/cmd/gosec@v2.27.1
|
||||||
|
gosec --exclude=G104,G115,G117,G203,G402,G703,G704 ./...
|
||||||
|
|
||||||
|
agentcompat-stress:
|
||||||
|
name: Linux agent compatibility stress
|
||||||
|
runs-on: ubuntu-24.04
|
||||||
|
timeout-minutes: 75
|
||||||
|
steps:
|
||||||
|
- name: Checkout Nezha revision
|
||||||
|
uses: actions/checkout@v7.0.1
|
||||||
with:
|
with:
|
||||||
args: --exclude=G103,G104,G107,G115,G117,G203,G402,G703,G704 ./...
|
path: nezha
|
||||||
|
persist-credentials: false
|
||||||
|
- name: Checkout Agent repository
|
||||||
|
uses: actions/checkout@v7.0.1
|
||||||
|
with:
|
||||||
|
repository: nezhahq/agent
|
||||||
|
path: agent
|
||||||
|
persist-credentials: false
|
||||||
|
- name: Set up Go
|
||||||
|
uses: actions/setup-go@v7
|
||||||
|
with:
|
||||||
|
go-version: "1.26.x"
|
||||||
|
cache: false
|
||||||
|
- name: Prepare Dashboard build inputs
|
||||||
|
working-directory: nezha
|
||||||
|
run: |
|
||||||
|
go install github.com/swaggo/swag/cmd/swag@v1.16.6
|
||||||
|
mkdir -p cmd/dashboard/user-dist cmd/dashboard/admin-dist
|
||||||
|
printf 'placeholder\n' > cmd/dashboard/user-dist/placeholder.txt
|
||||||
|
printf 'placeholder\n' > cmd/dashboard/admin-dist/placeholder.txt
|
||||||
|
swag init --pd -d cmd/dashboard -g main.go -o cmd/dashboard/docs
|
||||||
|
- name: Require named stress test
|
||||||
|
working-directory: nezha
|
||||||
|
run: go test -mod=readonly -tags=agentcompat -list '^TestStressPRFullEightAgentExactlyOnce$' ./integration/agentcompat/internal/scenario | grep -Fx 'TestStressPRFullEightAgentExactlyOnce'
|
||||||
|
- name: Run PR-full agent compatibility stress
|
||||||
|
working-directory: nezha
|
||||||
|
env:
|
||||||
|
AGENTCOMPAT_NEZHA_SOURCE: ${{ github.workspace }}/nezha
|
||||||
|
AGENTCOMPAT_AGENT_SOURCE: ${{ github.workspace }}/agent
|
||||||
|
run: go test -mod=readonly -tags=agentcompat -run '^TestStressPRFullEightAgentExactlyOnce$' -count=1 -v ./integration/agentcompat/internal/scenario
|
||||||
|
|
||||||
|
nezha-quality-required:
|
||||||
|
name: nezha-quality-required
|
||||||
|
if: ${{ always() }}
|
||||||
|
needs:
|
||||||
|
- tests
|
||||||
|
- linux-race-quality
|
||||||
|
- agentcompat-stress
|
||||||
|
runs-on: ubuntu-24.04
|
||||||
|
timeout-minutes: 5
|
||||||
|
steps:
|
||||||
|
- name: Require all blocking jobs to pass
|
||||||
|
shell: bash
|
||||||
|
run: |
|
||||||
|
test "${{ needs.tests.result }}" = success
|
||||||
|
test "${{ needs.linux-race-quality.result }}" = success
|
||||||
|
test "${{ needs.agentcompat-stress.result }}" = success
|
||||||
|
|
||||||
|
|||||||
+3
-1
@@ -26,4 +26,6 @@
|
|||||||
/cmd/dashboard/docs
|
/cmd/dashboard/docs
|
||||||
/data/*
|
/data/*
|
||||||
app
|
app
|
||||||
dashboard
|
dashboard
|
||||||
|
.omo/
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
# 哪吒面板 (Nezha Dashboard) - 个人编译指南
|
# 哪吒面板 (Nezha Dashboard) - 域名增强定制版
|
||||||
这份文档记录了在 Arch Linux 环境下,使用 VS Code Dev Containers 编译自定义主题版本的哪吒面板的完整流程。
|
本项目集成了域名管理(WHOIS/RDAP、Nazhumi 价格同步、到期提醒)、自定义通知系统与 Telegram Bot 交互、自定义 Branding、VPS 到期解析与可视化配置生成等核心功能。
|
||||||
|
|
||||||
|
|
||||||
## ⚠️ 核心前置条件 (不做会卡死)
|
## ⚠️ 核心前置条件 (不做会卡死)
|
||||||
网络环境:必须开启 TUN 模式 (透明代理)。
|
网络环境:必须开启 TUN 模式 (透明代理)。
|
||||||
@@ -53,7 +54,6 @@ cp -r ./admin-frontend-domain/dist/* ./nezha_domains/cmd/dashboard/admin-dist/
|
|||||||
确保 API 文档和 gRPC 代码是最新的。
|
确保 API 文档和 gRPC 代码是最新的。
|
||||||
|
|
||||||
```Bash
|
```Bash
|
||||||
|
|
||||||
# 生成 Swagger 文档
|
# 生成 Swagger 文档
|
||||||
swag init --pd -d . -g ./cmd/dashboard/main.go -o ./cmd/dashboard/docs --requiredByDefault
|
swag init --pd -d . -g ./cmd/dashboard/main.go -o ./cmd/dashboard/docs --requiredByDefault
|
||||||
|
|
||||||
@@ -64,7 +64,6 @@ protoc --go-grpc_out="require_unimplemented_servers=false:." --go_out="." proto/
|
|||||||
生成可执行文件 dashboard。
|
生成可执行文件 dashboard。
|
||||||
|
|
||||||
```Bash
|
```Bash
|
||||||
|
|
||||||
# -s -w: 去除调试符号,减小体积
|
# -s -w: 去除调试符号,减小体积
|
||||||
go build -ldflags="-s -w" -o dashboard cmd/dashboard/main.go
|
go build -ldflags="-s -w" -o dashboard cmd/dashboard/main.go
|
||||||
```
|
```
|
||||||
@@ -72,10 +71,10 @@ go build -ldflags="-s -w" -o dashboard cmd/dashboard/main.go
|
|||||||
编译完成后,运行以下命令测试:
|
编译完成后,运行以下命令测试:
|
||||||
|
|
||||||
```Bash
|
```Bash
|
||||||
|
|
||||||
# 运行面板
|
# 运行面板
|
||||||
./dashboard
|
./dashboard
|
||||||
|
|
||||||
# 如果能看到 Logo 输出,或者提示 config.yaml 不存在,说明编译成功。
|
# 如果能看到 Logo 输出,或者提示 config.yaml 不存在,说明编译成功。
|
||||||
# 如果提示 user-dist 404 之类的,说明前端文件没复制对。
|
# 如果提示 user-dist 404 之类的,说明前端文件没复制对。
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|||||||
+1
-1
@@ -6,4 +6,4 @@ Code in `master` branch.
|
|||||||
|
|
||||||
## Reporting a Vulnerability
|
## Reporting a Vulnerability
|
||||||
|
|
||||||
Thank you for your contribution to open source security, please email hi@nai.ba with details of the vulnerability.
|
Thank you for your contribution to open source security, please submit on https://github.com/nezhahq/nezha/security
|
||||||
|
|||||||
@@ -0,0 +1,86 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
type agentcompatCapabilityIdentity struct {
|
||||||
|
purpose rpc.AgentCompatCapabilityPurpose
|
||||||
|
serverID uint64
|
||||||
|
resourceID uint64
|
||||||
|
}
|
||||||
|
|
||||||
|
type agentcompatCapabilityProof struct {
|
||||||
|
owner rpc.AgentCompatCapabilityOwner
|
||||||
|
identity agentcompatCapabilityIdentity
|
||||||
|
}
|
||||||
|
|
||||||
|
func currentAgentcompatCapabilityProof(context *gin.Context, identity agentcompatCapabilityIdentity) (agentcompatCapabilityProof, error) {
|
||||||
|
if err := validateAgentcompatCapabilityIdentity(identity.purpose, identity.serverID, identity.resourceID); err != nil {
|
||||||
|
return agentcompatCapabilityProof{}, err
|
||||||
|
}
|
||||||
|
token := APITokenFromContext(context)
|
||||||
|
authorized, present := context.Get(model.CtxKeyAuthorizedUser)
|
||||||
|
user, validUser := authorized.(*model.User)
|
||||||
|
if token == nil || token.ID == 0 || !present || !validUser || user == nil || user.ID == 0 {
|
||||||
|
return agentcompatCapabilityProof{}, errAgentcompatCapabilityUnavailable
|
||||||
|
}
|
||||||
|
if singleton.ServerShared == nil {
|
||||||
|
return agentcompatCapabilityProof{}, errAgentcompatCapabilityUnavailable
|
||||||
|
}
|
||||||
|
server, exists := singleton.ServerShared.Get(identity.serverID)
|
||||||
|
if !exists || server == nil || !server.HasPermission(context) || !patAllowsServer(context, identity.serverID) {
|
||||||
|
return agentcompatCapabilityProof{}, errAgentcompatCapabilityUnavailable
|
||||||
|
}
|
||||||
|
if identity.purpose == rpc.AgentCompatCapabilityNAT && !currentAgentcompatNATPermission(context, identity) {
|
||||||
|
return agentcompatCapabilityProof{}, errAgentcompatCapabilityUnavailable
|
||||||
|
}
|
||||||
|
return agentcompatCapabilityProof{
|
||||||
|
owner: rpc.AgentCompatCapabilityOwner{PATID: token.ID, UserID: user.ID, IsAdmin: user.Role.IsAdmin()},
|
||||||
|
identity: identity,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func currentAgentcompatNATPermission(context *gin.Context, identity agentcompatCapabilityIdentity) bool {
|
||||||
|
if singleton.NATShared == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
domain := singleton.NATShared.GetDomain(identity.resourceID)
|
||||||
|
if domain == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
profile := singleton.NATShared.GetNATConfigByDomain(domain)
|
||||||
|
return profile != nil && profile.ID == identity.resourceID && profile.ServerID == identity.serverID && profile.HasPermission(context)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (proof agentcompatCapabilityProof) registration() rpc.AgentCompatCapabilityRegistration {
|
||||||
|
return rpc.AgentCompatCapabilityRegistration{
|
||||||
|
Owner: proof.owner, Purpose: proof.identity.purpose, TargetServerID: proof.identity.serverID,
|
||||||
|
ResourceID: proof.identity.resourceID, ServerAccessAllowed: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (proof agentcompatCapabilityProof) access(capability rpc.AgentCompatIOStreamCapability) rpc.AgentCompatCapabilityAccess {
|
||||||
|
return rpc.AgentCompatCapabilityAccess{
|
||||||
|
Capability: capability, Owner: proof.owner, Purpose: proof.identity.purpose,
|
||||||
|
TargetServerID: proof.identity.serverID, ResourceID: proof.identity.resourceID, ServerAccessAllowed: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func agentcompatCapabilityIdentityFromWire(purpose string, serverID, resourceID uint64) (agentcompatCapabilityIdentity, error) {
|
||||||
|
parsedPurpose, err := parseAgentcompatCapabilityPurpose(purpose)
|
||||||
|
if err != nil {
|
||||||
|
return agentcompatCapabilityIdentity{}, err
|
||||||
|
}
|
||||||
|
identity := agentcompatCapabilityIdentity{purpose: parsedPurpose, serverID: serverID, resourceID: resourceID}
|
||||||
|
if err := validateAgentcompatCapabilityIdentity(identity.purpose, identity.serverID, identity.resourceID); err != nil {
|
||||||
|
return agentcompatCapabilityIdentity{}, err
|
||||||
|
}
|
||||||
|
return identity, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,193 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityBoundaryAcceptsExactly512Bytes(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
|
||||||
|
body := padAgentcompatCapabilityBody(t, `{"purpose":"terminal","server_id":7}`)
|
||||||
|
status, responseBody := postAgentcompatBoundary(t, server.URL+agentcompatCapabilityRegisterPath, token, body, true)
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
var envelope model.CommonResponse[agentcompatCapabilityRegisterResponse]
|
||||||
|
require.NoError(t, json.Unmarshal([]byte(responseBody), &envelope))
|
||||||
|
require.True(t, envelope.Success)
|
||||||
|
require.NotEmpty(t, envelope.Data.Capability)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityBoundaryRejects513BytesWithoutDisclosure(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
body := padAgentcompatCapabilityBody(t, `{"purpose":"terminal","server_id":7}`) + " "
|
||||||
|
|
||||||
|
for _, route := range []struct {
|
||||||
|
name string
|
||||||
|
path string
|
||||||
|
wantError bool
|
||||||
|
wantSuccess bool
|
||||||
|
}{
|
||||||
|
{name: "register", path: agentcompatCapabilityRegisterPath, wantError: true},
|
||||||
|
{name: "wait", path: agentcompatCapabilityWaitPath, wantError: true},
|
||||||
|
{name: "cancel", path: agentcompatCapabilityCancelPath, wantSuccess: true},
|
||||||
|
{name: "unregister", path: agentcompatCapabilityUnregisterPath, wantSuccess: true},
|
||||||
|
} {
|
||||||
|
t.Run(route.name, func(t *testing.T) {
|
||||||
|
status, responseBody := postAgentcompatBoundary(t, server.URL+route.path, token, body, true)
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.NotContains(t, responseBody, "terminal")
|
||||||
|
require.NotContains(t, responseBody, "server_id")
|
||||||
|
if route.wantError {
|
||||||
|
var envelope model.CommonResponse[agentcompatCapabilityRegisterResponse]
|
||||||
|
require.NoError(t, json.Unmarshal([]byte(responseBody), &envelope))
|
||||||
|
require.False(t, envelope.Success)
|
||||||
|
require.Equal(t, errAgentcompatCapabilityInvalid.Error(), envelope.Error)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
require.Equal(t, `{"success":true,"data":{}}`, responseBody)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityBoundaryRejectsChunkedOversizeWithoutDisclosure(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
body := padAgentcompatCapabilityBody(t, `{"purpose":"terminal","server_id":7}`) + " "
|
||||||
|
|
||||||
|
status, responseBody := postAgentcompatBoundary(t, server.URL+agentcompatCapabilityRegisterPath, token, body, false)
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.Contains(t, responseBody, errAgentcompatCapabilityInvalid.Error())
|
||||||
|
require.NotContains(t, responseBody, "terminal")
|
||||||
|
require.NotContains(t, responseBody, "server_id")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityBoundaryRejectsDuplicateKeysInEitherOrder(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
|
||||||
|
keys := []string{"purpose", "server_id", "resource_id"}
|
||||||
|
for _, key := range keys {
|
||||||
|
for _, body := range []string{
|
||||||
|
`{"` + key + `":"terminal","` + key + `":"terminal","purpose":"terminal","server_id":7}`,
|
||||||
|
`{"` + key + `":0,"purpose":"terminal","server_id":7,"` + key + `":"terminal"}`,
|
||||||
|
} {
|
||||||
|
status, responseBody := postAgentcompatBoundary(t, server.URL+agentcompatCapabilityRegisterPath, token, body, true)
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.Contains(t, responseBody, errAgentcompatCapabilityInvalid.Error())
|
||||||
|
require.NotContains(t, responseBody, "terminal")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityBoundaryRejectsAccessGrammar(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
|
||||||
|
cases := []string{
|
||||||
|
`{"capability":"x","purpose":"terminal"}`,
|
||||||
|
`{"capability":"x","server_id":7}`,
|
||||||
|
`{"capability":"x","purpose":"terminal","server_id":7,"capability":"y"}`,
|
||||||
|
`{"capability":"x","purpose":"terminal","server_id":0}`,
|
||||||
|
`{"capability":"x","purpose":"terminal","server_id":7,"unknown":"private"}`,
|
||||||
|
`{"capability":"x","purpose":"terminal","server_id":7} {}`,
|
||||||
|
`{"capability":null,"purpose":"terminal","server_id":7}`,
|
||||||
|
`{"capability":7,"purpose":"terminal","server_id":7}`,
|
||||||
|
`{"capability":"x","purpose":7,"server_id":7}`,
|
||||||
|
`{"capability":"x","purpose":"terminal","server_id":"7"}`,
|
||||||
|
`{"capability":"x","purpose":"terminal","server_id":-1}`,
|
||||||
|
`{"capability":"x","purpose":"terminal","server_id":1.5}`,
|
||||||
|
`{"capability":"x","purpose":"terminal","server_id":1e2}`,
|
||||||
|
`null`, `[]`, `{`,
|
||||||
|
}
|
||||||
|
for _, body := range cases {
|
||||||
|
status, responseBody := postAgentcompatBoundary(t, server.URL+agentcompatCapabilityWaitPath, token, body, true)
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.Contains(t, responseBody, errAgentcompatCapabilityInvalid.Error())
|
||||||
|
require.NotContains(t, responseBody, "private")
|
||||||
|
require.NotContains(t, responseBody, "x")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityBoundaryRejectsMissingRegisterFields(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
|
||||||
|
for _, body := range []string{`{"server_id":7}`, `{"purpose":"terminal"}`, `{}`} {
|
||||||
|
status, responseBody := postAgentcompatBoundary(t, server.URL+agentcompatCapabilityRegisterPath, token, body, true)
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.Contains(t, responseBody, errAgentcompatCapabilityInvalid.Error())
|
||||||
|
require.NotContains(t, responseBody, "terminal")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func padAgentcompatCapabilityBody(t *testing.T, body string) string {
|
||||||
|
t.Helper()
|
||||||
|
require.LessOrEqual(t, len(body), 512)
|
||||||
|
return body + strings.Repeat(" ", 512-len(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
func postAgentcompatBoundary(t *testing.T, path, token, body string, contentLength bool) (int, string) {
|
||||||
|
t.Helper()
|
||||||
|
requestContext, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
var reader io.Reader = bytes.NewBufferString(body)
|
||||||
|
request, err := http.NewRequestWithContext(requestContext, http.MethodPost, path, reader)
|
||||||
|
require.NoError(t, err)
|
||||||
|
if !contentLength {
|
||||||
|
request.Body = io.NopCloser(strings.NewReader(body))
|
||||||
|
request.ContentLength = -1
|
||||||
|
}
|
||||||
|
request.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
request.Header.Set("Content-Type", "application/json")
|
||||||
|
response, err := http.DefaultClient.Do(request)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer response.Body.Close()
|
||||||
|
responseBody, err := io.ReadAll(response.Body)
|
||||||
|
require.NoError(t, err)
|
||||||
|
return response.StatusCode, string(responseBody)
|
||||||
|
}
|
||||||
@@ -0,0 +1,199 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
agentcompatCapabilityRequestMaxBytes = 512
|
||||||
|
|
||||||
|
agentcompatCapabilityRegisterPath = "/agentcompat/io-stream-capability/register"
|
||||||
|
agentcompatCapabilityWaitPath = "/agentcompat/io-stream-capability/wait"
|
||||||
|
agentcompatCapabilityCancelPath = "/agentcompat/io-stream-capability/cancel"
|
||||||
|
agentcompatCapabilityUnregisterPath = "/agentcompat/io-stream-capability/unregister"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
agentcompatCapabilityPurposeTerminal = "terminal"
|
||||||
|
agentcompatCapabilityPurposeFileManager = "file_manager"
|
||||||
|
agentcompatCapabilityPurposeNAT = "nat"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
errAgentcompatCapabilityInvalid = errors.New("agentcompat capability request is invalid")
|
||||||
|
errAgentcompatCapabilityUnavailable = errors.New("agentcompat capability is not available")
|
||||||
|
errAgentcompatCapabilityConflict = errors.New("agentcompat capability is active")
|
||||||
|
errAgentcompatCapabilityCleanup = errors.New("agentcompat capability cleanup failed")
|
||||||
|
)
|
||||||
|
|
||||||
|
type agentcompatCapabilityRegisterRequest struct {
|
||||||
|
Purpose string `json:"purpose"`
|
||||||
|
ServerID uint64 `json:"server_id"`
|
||||||
|
ResourceID uint64 `json:"resource_id,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type agentcompatCapabilityRegisterResponse struct {
|
||||||
|
Capability string `json:"capability"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type agentcompatCapabilityAccessRequest struct {
|
||||||
|
Capability string `json:"capability"`
|
||||||
|
Purpose string `json:"purpose"`
|
||||||
|
ServerID uint64 `json:"server_id"`
|
||||||
|
ResourceID uint64 `json:"resource_id,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type agentcompatCapabilityWaitRequest = agentcompatCapabilityAccessRequest
|
||||||
|
type agentcompatCapabilityCancelRequest = agentcompatCapabilityAccessRequest
|
||||||
|
type agentcompatCapabilityUnregisterRequest = agentcompatCapabilityAccessRequest
|
||||||
|
|
||||||
|
type agentcompatCapabilityWaitResponse struct {
|
||||||
|
StreamID string `json:"stream_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type agentcompatCapabilityEmptyResponse struct{}
|
||||||
|
|
||||||
|
func decodeAgentcompatCapabilityRequest[T any](context *gin.Context, destination *T) error {
|
||||||
|
context.Request.Body = http.MaxBytesReader(context.Writer, context.Request.Body, agentcompatCapabilityRequestMaxBytes)
|
||||||
|
decoder := json.NewDecoder(context.Request.Body)
|
||||||
|
fields, err := decodeAgentcompatCapabilityObject(decoder)
|
||||||
|
if err != nil {
|
||||||
|
return errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
switch request := any(destination).(type) {
|
||||||
|
case *agentcompatCapabilityRegisterRequest:
|
||||||
|
request.Purpose, err = decodeAgentcompatCapabilityString(fields.purpose)
|
||||||
|
if err == nil {
|
||||||
|
request.ServerID, err = decodeAgentcompatCapabilityUint64(fields.serverID)
|
||||||
|
}
|
||||||
|
if err == nil && fields.resourceID != nil {
|
||||||
|
request.ResourceID, err = decodeAgentcompatCapabilityUint64(fields.resourceID)
|
||||||
|
}
|
||||||
|
if err != nil || fields.purpose == nil || fields.serverID == nil {
|
||||||
|
return errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
case *agentcompatCapabilityAccessRequest:
|
||||||
|
request.Capability, err = decodeAgentcompatCapabilityString(fields.capability)
|
||||||
|
if err == nil {
|
||||||
|
request.Purpose, err = decodeAgentcompatCapabilityString(fields.purpose)
|
||||||
|
}
|
||||||
|
if err == nil {
|
||||||
|
request.ServerID, err = decodeAgentcompatCapabilityUint64(fields.serverID)
|
||||||
|
}
|
||||||
|
if err == nil && fields.resourceID != nil {
|
||||||
|
request.ResourceID, err = decodeAgentcompatCapabilityUint64(fields.resourceID)
|
||||||
|
}
|
||||||
|
if err != nil || fields.capability == nil || fields.purpose == nil || fields.serverID == nil {
|
||||||
|
return errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type agentcompatCapabilityFields struct {
|
||||||
|
purpose json.RawMessage
|
||||||
|
serverID json.RawMessage
|
||||||
|
resourceID json.RawMessage
|
||||||
|
capability json.RawMessage
|
||||||
|
seen uint8
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeAgentcompatCapabilityObject(decoder *json.Decoder) (agentcompatCapabilityFields, error) {
|
||||||
|
var fields agentcompatCapabilityFields
|
||||||
|
token, err := decoder.Token()
|
||||||
|
if err != nil || token != json.Delim('{') {
|
||||||
|
return fields, errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
for decoder.More() {
|
||||||
|
keyToken, err := decoder.Token()
|
||||||
|
key, validKey := keyToken.(string)
|
||||||
|
if err != nil || !validKey {
|
||||||
|
return fields, errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
var bit uint8
|
||||||
|
var target *json.RawMessage
|
||||||
|
switch key {
|
||||||
|
case "purpose":
|
||||||
|
bit, target = 1, &fields.purpose
|
||||||
|
case "server_id":
|
||||||
|
bit, target = 2, &fields.serverID
|
||||||
|
case "resource_id":
|
||||||
|
bit, target = 4, &fields.resourceID
|
||||||
|
case "capability":
|
||||||
|
bit, target = 8, &fields.capability
|
||||||
|
default:
|
||||||
|
return fields, errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
if fields.seen&bit != 0 || decoder.Decode(target) != nil {
|
||||||
|
return fields, errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
fields.seen |= bit
|
||||||
|
}
|
||||||
|
if token, err = decoder.Token(); err != nil || token != json.Delim('}') {
|
||||||
|
return fields, errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
if _, err = decoder.Token(); !errors.Is(err, io.EOF) {
|
||||||
|
return fields, errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
return fields, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeAgentcompatCapabilityString(raw json.RawMessage) (string, error) {
|
||||||
|
if strings.TrimSpace(string(raw)) == "null" {
|
||||||
|
return "", errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
var value string
|
||||||
|
if err := json.Unmarshal(raw, &value); err != nil {
|
||||||
|
return "", errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
return value, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeAgentcompatCapabilityUint64(raw json.RawMessage) (uint64, error) {
|
||||||
|
return strconv.ParseUint(strings.TrimSpace(string(raw)), 10, 64)
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseAgentcompatCapabilityPurpose(value string) (rpc.AgentCompatCapabilityPurpose, error) {
|
||||||
|
switch value {
|
||||||
|
case agentcompatCapabilityPurposeTerminal:
|
||||||
|
return rpc.AgentCompatCapabilityTerminal, nil
|
||||||
|
case agentcompatCapabilityPurposeFileManager:
|
||||||
|
return rpc.AgentCompatCapabilityFileManager, nil
|
||||||
|
case agentcompatCapabilityPurposeNAT:
|
||||||
|
return rpc.AgentCompatCapabilityNAT, nil
|
||||||
|
default:
|
||||||
|
return 0, errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func validateAgentcompatCapabilityIdentity(purpose rpc.AgentCompatCapabilityPurpose, serverID, resourceID uint64) error {
|
||||||
|
if serverID == 0 {
|
||||||
|
return errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
switch purpose {
|
||||||
|
case rpc.AgentCompatCapabilityTerminal, rpc.AgentCompatCapabilityFileManager:
|
||||||
|
if resourceID != 0 {
|
||||||
|
return errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
case rpc.AgentCompatCapabilityNAT:
|
||||||
|
if resourceID == 0 {
|
||||||
|
return errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,24 @@
|
|||||||
|
//go:build !agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestDefaultBuildDoesNotRegisterAgentcompatCapabilityRoutes(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
router := gin.New()
|
||||||
|
registerAgentcompatRoutes(router)
|
||||||
|
request := httptest.NewRequest(http.MethodPost, "/agentcompat/io-stream-capability/register", nil)
|
||||||
|
response := httptest.NewRecorder()
|
||||||
|
|
||||||
|
router.ServeHTTP(response, request)
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusNotFound, response.Code)
|
||||||
|
}
|
||||||
@@ -0,0 +1,85 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
type agentcompatCapabilityCloseFailure struct {
|
||||||
|
closed atomic.Int32
|
||||||
|
}
|
||||||
|
|
||||||
|
func (*agentcompatCapabilityCloseFailure) Read([]byte) (int, error) { return 0, io.EOF }
|
||||||
|
func (*agentcompatCapabilityCloseFailure) Write(data []byte) (int, error) { return len(data), nil }
|
||||||
|
func (failure *agentcompatCapabilityCloseFailure) Close() error {
|
||||||
|
failure.closed.Add(1)
|
||||||
|
return errors.New("private-stream private-capability private-server")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityBoundUnregisterAndCancelCleanupErrorsAreGeneric(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
handler := rpc.NewNezhaHandler()
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = handler
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
token, plaintext := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
capability := registerAgentcompatCapability(t, server.URL, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "file_manager", ServerID: 7})
|
||||||
|
parsed, err := rpc.ParseAgentCompatIOStreamCapability(capability)
|
||||||
|
require.NoError(t, err)
|
||||||
|
access := rpc.AgentCompatCapabilityAccess{
|
||||||
|
Capability: parsed, Owner: rpc.AgentCompatCapabilityOwner{PATID: token.ID, UserID: userID},
|
||||||
|
Purpose: rpc.AgentCompatCapabilityFileManager, TargetServerID: 7, ServerAccessAllowed: true,
|
||||||
|
}
|
||||||
|
require.NoError(t, handler.CreateStreamWithPurpose("private-cleanup-stream", userID, 7, rpc.PurposeFileManager))
|
||||||
|
require.NoError(t, handler.BindAgentCompatIOStreamCapability(rpc.AgentCompatCapabilityBinding{AgentCompatCapabilityAccess: access, StreamID: "private-cleanup-stream"}))
|
||||||
|
|
||||||
|
unregisterStatus, unregisterBody := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityUnregisterPath, plaintext, agentcompatCapabilityUnregisterRequest{Capability: capability, Purpose: "file_manager", ServerID: 7})
|
||||||
|
require.Equal(t, http.StatusOK, unregisterStatus)
|
||||||
|
require.Contains(t, unregisterBody, errAgentcompatCapabilityConflict.Error())
|
||||||
|
require.NotContains(t, unregisterBody, capability)
|
||||||
|
require.NotContains(t, unregisterBody, "private-cleanup-stream")
|
||||||
|
|
||||||
|
failure := &agentcompatCapabilityCloseFailure{}
|
||||||
|
require.NoError(t, handler.UserConnected("private-cleanup-stream", failure))
|
||||||
|
cancelStatus, cancelBody := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityCancelPath, plaintext, agentcompatCapabilityCancelRequest{Capability: capability, Purpose: "file_manager", ServerID: 7})
|
||||||
|
require.Equal(t, http.StatusOK, cancelStatus)
|
||||||
|
require.Contains(t, cancelBody, errAgentcompatCapabilityCleanup.Error())
|
||||||
|
require.NotContains(t, cancelBody, capability)
|
||||||
|
require.NotContains(t, cancelBody, "private-cleanup-stream")
|
||||||
|
require.Equal(t, int32(1), failure.closed.Load())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityWaitHonorsRequestCancellation(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, plaintext := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
capability := registerAgentcompatCapability(t, server.URL, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "terminal", ServerID: 7})
|
||||||
|
requestContext, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
|
||||||
|
defer cancel()
|
||||||
|
body := strings.NewReader(capabilityAccessJSON(capability, "terminal", 7, 0))
|
||||||
|
request, err := http.NewRequestWithContext(requestContext, http.MethodPost, server.URL+agentcompatCapabilityWaitPath, body)
|
||||||
|
require.NoError(t, err)
|
||||||
|
request.Header.Set("Authorization", "Bearer "+plaintext)
|
||||||
|
request.Header.Set("Content-Type", "application/json")
|
||||||
|
_, err = server.Client().Do(request)
|
||||||
|
require.True(t, errors.Is(err, context.DeadlineExceeded))
|
||||||
|
}
|
||||||
@@ -0,0 +1,114 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
func registerAgentcompatCapabilityRoutes(router *gin.Engine, patAuth gin.HandlerFunc) {
|
||||||
|
router.POST(agentcompatCapabilityRegisterPath, patAuth, commonHandler(agentcompatCapabilityRegister))
|
||||||
|
router.POST(agentcompatCapabilityWaitPath, patAuth, commonHandler(agentcompatCapabilityWait))
|
||||||
|
router.POST(agentcompatCapabilityCancelPath, patAuth, commonHandler(agentcompatCapabilityCancel))
|
||||||
|
router.POST(agentcompatCapabilityUnregisterPath, patAuth, commonHandler(agentcompatCapabilityUnregister))
|
||||||
|
}
|
||||||
|
|
||||||
|
func agentcompatCapabilityRegister(context *gin.Context) (agentcompatCapabilityRegisterResponse, error) {
|
||||||
|
var request agentcompatCapabilityRegisterRequest
|
||||||
|
if err := decodeAgentcompatCapabilityRequest(context, &request); err != nil {
|
||||||
|
return agentcompatCapabilityRegisterResponse{}, err
|
||||||
|
}
|
||||||
|
identity, err := agentcompatCapabilityIdentityFromWire(request.Purpose, request.ServerID, request.ResourceID)
|
||||||
|
if err != nil {
|
||||||
|
return agentcompatCapabilityRegisterResponse{}, err
|
||||||
|
}
|
||||||
|
proof, err := currentAgentcompatCapabilityProof(context, identity)
|
||||||
|
if err != nil || rpc.NezhaHandlerSingleton == nil {
|
||||||
|
return agentcompatCapabilityRegisterResponse{}, errAgentcompatCapabilityUnavailable
|
||||||
|
}
|
||||||
|
capability, err := rpc.NezhaHandlerSingleton.RegisterAgentCompatIOStreamCapability(context.Request.Context(), proof.registration())
|
||||||
|
if err != nil {
|
||||||
|
if contextError := context.Request.Context().Err(); contextError != nil {
|
||||||
|
return agentcompatCapabilityRegisterResponse{}, contextError
|
||||||
|
}
|
||||||
|
return agentcompatCapabilityRegisterResponse{}, errAgentcompatCapabilityUnavailable
|
||||||
|
}
|
||||||
|
return agentcompatCapabilityRegisterResponse{Capability: capability.String()}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func agentcompatCapabilityWait(context *gin.Context) (agentcompatCapabilityWaitResponse, error) {
|
||||||
|
var request agentcompatCapabilityWaitRequest
|
||||||
|
if err := decodeAgentcompatCapabilityRequest(context, &request); err != nil {
|
||||||
|
return agentcompatCapabilityWaitResponse{}, err
|
||||||
|
}
|
||||||
|
access, err := currentAgentcompatCapabilityAccess(context, request)
|
||||||
|
if errors.Is(err, errAgentcompatCapabilityInvalid) {
|
||||||
|
return agentcompatCapabilityWaitResponse{}, errAgentcompatCapabilityInvalid
|
||||||
|
}
|
||||||
|
if err != nil || rpc.NezhaHandlerSingleton == nil {
|
||||||
|
return agentcompatCapabilityWaitResponse{}, errAgentcompatCapabilityUnavailable
|
||||||
|
}
|
||||||
|
streamID, err := rpc.NezhaHandlerSingleton.WaitAgentCompatIOStreamCapability(context.Request.Context(), access)
|
||||||
|
if err != nil {
|
||||||
|
if contextError := context.Request.Context().Err(); contextError != nil {
|
||||||
|
return agentcompatCapabilityWaitResponse{}, contextError
|
||||||
|
}
|
||||||
|
return agentcompatCapabilityWaitResponse{}, errAgentcompatCapabilityUnavailable
|
||||||
|
}
|
||||||
|
return agentcompatCapabilityWaitResponse{StreamID: streamID}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func agentcompatCapabilityCancel(context *gin.Context) (agentcompatCapabilityEmptyResponse, error) {
|
||||||
|
access, available := inertAgentcompatCapabilityAccess(context)
|
||||||
|
if !available || rpc.NezhaHandlerSingleton == nil {
|
||||||
|
return agentcompatCapabilityEmptyResponse{}, nil
|
||||||
|
}
|
||||||
|
if err := rpc.NezhaHandlerSingleton.CancelAgentCompatIOStreamCapability(access); err != nil {
|
||||||
|
return agentcompatCapabilityEmptyResponse{}, errAgentcompatCapabilityCleanup
|
||||||
|
}
|
||||||
|
return agentcompatCapabilityEmptyResponse{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func agentcompatCapabilityUnregister(context *gin.Context) (agentcompatCapabilityEmptyResponse, error) {
|
||||||
|
access, available := inertAgentcompatCapabilityAccess(context)
|
||||||
|
if !available || rpc.NezhaHandlerSingleton == nil {
|
||||||
|
return agentcompatCapabilityEmptyResponse{}, nil
|
||||||
|
}
|
||||||
|
err := rpc.NezhaHandlerSingleton.UnregisterAgentCompatIOStreamCapability(access)
|
||||||
|
if errors.Is(err, rpc.ErrAgentCompatCapabilityBound) {
|
||||||
|
return agentcompatCapabilityEmptyResponse{}, errAgentcompatCapabilityConflict
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return agentcompatCapabilityEmptyResponse{}, errAgentcompatCapabilityUnavailable
|
||||||
|
}
|
||||||
|
return agentcompatCapabilityEmptyResponse{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func currentAgentcompatCapabilityAccess(context *gin.Context, request agentcompatCapabilityAccessRequest) (rpc.AgentCompatCapabilityAccess, error) {
|
||||||
|
identity, err := agentcompatCapabilityIdentityFromWire(request.Purpose, request.ServerID, request.ResourceID)
|
||||||
|
if err != nil {
|
||||||
|
return rpc.AgentCompatCapabilityAccess{}, err
|
||||||
|
}
|
||||||
|
capability, err := rpc.ParseAgentCompatIOStreamCapability(request.Capability)
|
||||||
|
if err != nil {
|
||||||
|
return rpc.AgentCompatCapabilityAccess{}, errAgentcompatCapabilityUnavailable
|
||||||
|
}
|
||||||
|
proof, err := currentAgentcompatCapabilityProof(context, identity)
|
||||||
|
if err != nil {
|
||||||
|
return rpc.AgentCompatCapabilityAccess{}, errAgentcompatCapabilityUnavailable
|
||||||
|
}
|
||||||
|
return proof.access(capability), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func inertAgentcompatCapabilityAccess(context *gin.Context) (rpc.AgentCompatCapabilityAccess, bool) {
|
||||||
|
var request agentcompatCapabilityAccessRequest
|
||||||
|
if decodeAgentcompatCapabilityRequest(context, &request) != nil {
|
||||||
|
return rpc.AgentCompatCapabilityAccess{}, false
|
||||||
|
}
|
||||||
|
access, err := currentAgentcompatCapabilityAccess(context, request)
|
||||||
|
return access, err == nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,109 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/http"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAgentcompatNATCapabilityUsesExactCurrentProfilePermission(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
handler := rpc.NewNezhaHandler()
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
originalNAT := singleton.NATShared
|
||||||
|
rpc.NezhaHandlerSingleton = handler
|
||||||
|
t.Cleanup(func() {
|
||||||
|
rpc.NezhaHandlerSingleton = originalHandler
|
||||||
|
singleton.NATShared = originalNAT
|
||||||
|
})
|
||||||
|
require.NoError(t, singleton.DB.AutoMigrate(&model.NAT{}))
|
||||||
|
secondServer := &model.Server{Common: model.Common{ID: 8}}
|
||||||
|
secondServer.SetUserID(userID)
|
||||||
|
singleton.ServerShared.InsertForTest(secondServer)
|
||||||
|
profile := &model.NAT{Common: model.Common{ID: 41, UserID: userID}, Name: "private-profile", ServerID: 7, Domain: "private-profile.example"}
|
||||||
|
foreignProfile := &model.NAT{Common: model.Common{ID: 42, UserID: 999}, Name: "foreign-profile", ServerID: 7, Domain: "foreign-profile.example"}
|
||||||
|
require.NoError(t, singleton.DB.Create(profile).Error)
|
||||||
|
require.NoError(t, singleton.DB.Create(foreignProfile).Error)
|
||||||
|
singleton.NATShared = singleton.NewNATClass()
|
||||||
|
token, plaintext := mkToken(t, userID, []string{model.ScopeNATRead}, []uint64{7, 8})
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
|
||||||
|
capability := registerAgentcompatCapability(t, server.URL, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "nat", ServerID: 7, ResourceID: 41})
|
||||||
|
status, mismatchBody := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityRegisterPath, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "nat", ServerID: 8, ResourceID: 41})
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.Contains(t, mismatchBody, errAgentcompatCapabilityUnavailable.Error())
|
||||||
|
foreignStatus, foreignBody := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityRegisterPath, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "nat", ServerID: 7, ResourceID: 42})
|
||||||
|
require.Equal(t, http.StatusOK, foreignStatus)
|
||||||
|
require.Contains(t, foreignBody, errAgentcompatCapabilityUnavailable.Error())
|
||||||
|
|
||||||
|
parsed, err := rpc.ParseAgentCompatIOStreamCapability(capability)
|
||||||
|
require.NoError(t, err)
|
||||||
|
access := rpc.AgentCompatCapabilityAccess{
|
||||||
|
Capability: parsed, Owner: rpc.AgentCompatCapabilityOwner{PATID: token.ID, UserID: userID},
|
||||||
|
Purpose: rpc.AgentCompatCapabilityNAT, TargetServerID: 7, ResourceID: 41, ServerAccessAllowed: true,
|
||||||
|
}
|
||||||
|
handle, err := handler.ConsumeAgentCompatNATCapability(access)
|
||||||
|
require.NoError(t, err)
|
||||||
|
lease, err := handler.CreateAgentCompatNATStream(handle, "private-nat-stream")
|
||||||
|
require.NoError(t, err)
|
||||||
|
t.Cleanup(func() { require.NoError(t, handler.CloseAgentCompatNATStreamLease(lease)) })
|
||||||
|
require.NoError(t, handler.PublishAgentCompatNATStream(handle, rpc.AgentCompatNATPublication{Purpose: rpc.AgentCompatCapabilityNAT, TargetServerID: 7, ResourceID: 41, StreamID: "private-nat-stream"}))
|
||||||
|
|
||||||
|
profile.ServerID = 8
|
||||||
|
singleton.NATShared.Update(profile)
|
||||||
|
_, unavailableErr := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](context.Background(), server.URL+agentcompatCapabilityWaitPath, plaintext, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "nat", ServerID: 7, ResourceID: 41})
|
||||||
|
require.Error(t, unavailableErr)
|
||||||
|
require.NotContains(t, unavailableErr.Error(), capability)
|
||||||
|
require.NotContains(t, unavailableErr.Error(), "private-nat-stream")
|
||||||
|
profile.ServerID = 7
|
||||||
|
singleton.NATShared.Update(profile)
|
||||||
|
waitContext, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
waited, err := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](waitContext, server.URL+agentcompatCapabilityWaitPath, plaintext, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "nat", ServerID: 7, ResourceID: 41})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, "private-nat-stream", waited.StreamID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityOwnerIncludesCurrentAdminRole(t *testing.T) {
|
||||||
|
cleanup, _ := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
handler := rpc.NewNezhaHandler()
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = handler
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
admin := &model.User{Common: model.Common{ID: 300}, Username: "cap-admin", Role: model.RoleAdmin}
|
||||||
|
require.NoError(t, singleton.DB.Create(admin).Error)
|
||||||
|
target := &model.Server{Common: model.Common{ID: 8}}
|
||||||
|
target.SetUserID(999)
|
||||||
|
singleton.ServerShared.InsertForTest(target)
|
||||||
|
token, plaintext := mkDistinctCapabilityToken(t, admin.ID, "admin")
|
||||||
|
token.SetServerIDs([]uint64{8})
|
||||||
|
require.NoError(t, singleton.DB.Model(token).Update("servers_csv", token.ServersCSV).Error)
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
capability := registerAgentcompatCapability(t, server.URL, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "terminal", ServerID: 8})
|
||||||
|
parsed, err := rpc.ParseAgentCompatIOStreamCapability(capability)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, handler.CreateStreamWithPurpose("admin-private-stream", admin.ID, 8, rpc.PurposeTerminal))
|
||||||
|
require.NoError(t, handler.BindAgentCompatIOStreamCapability(rpc.AgentCompatCapabilityBinding{
|
||||||
|
AgentCompatCapabilityAccess: rpc.AgentCompatCapabilityAccess{
|
||||||
|
Capability: parsed, Owner: rpc.AgentCompatCapabilityOwner{PATID: token.ID, UserID: admin.ID, IsAdmin: true},
|
||||||
|
Purpose: rpc.AgentCompatCapabilityTerminal, TargetServerID: 8, ServerAccessAllowed: true,
|
||||||
|
},
|
||||||
|
StreamID: "admin-private-stream",
|
||||||
|
}))
|
||||||
|
waitContext, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
waited, err := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](waitContext, server.URL+agentcompatCapabilityWaitPath, plaintext, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "terminal", ServerID: 8})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, "admin-private-stream", waited.StreamID)
|
||||||
|
}
|
||||||
@@ -0,0 +1,175 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityRoutesRequirePATAndDoNotAcceptOwnerFields(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
router := gin.New()
|
||||||
|
registerAgentcompatRoutes(router)
|
||||||
|
server := httptest.NewServer(router)
|
||||||
|
t.Cleanup(server.Close)
|
||||||
|
|
||||||
|
requestBody := `{"purpose":"terminal","server_id":7,"pat_id":999,"user_id":999,"is_admin":true}`
|
||||||
|
requestContext, cancelRequests := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancelRequests()
|
||||||
|
for _, authorization := range []string{"", "Bearer jwt-looking-value", "Bearer " + token} {
|
||||||
|
request, err := http.NewRequestWithContext(requestContext, http.MethodPost, server.URL+agentcompatCapabilityRegisterPath, strings.NewReader(requestBody))
|
||||||
|
require.NoError(t, err)
|
||||||
|
request.Header.Set("Content-Type", "application/json")
|
||||||
|
if authorization != "" {
|
||||||
|
request.Header.Set("Authorization", authorization)
|
||||||
|
}
|
||||||
|
response, err := server.Client().Do(request)
|
||||||
|
require.NoError(t, err)
|
||||||
|
body, readErr := io.ReadAll(response.Body)
|
||||||
|
response.Body.Close()
|
||||||
|
require.NoError(t, readErr)
|
||||||
|
if authorization == "Bearer "+token {
|
||||||
|
require.Equal(t, http.StatusOK, response.StatusCode)
|
||||||
|
var envelope model.CommonResponse[agentcompatCapabilityRegisterResponse]
|
||||||
|
require.NoError(t, json.Unmarshal(body, &envelope))
|
||||||
|
require.False(t, envelope.Success)
|
||||||
|
require.NotContains(t, string(body), "999")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
require.Equal(t, http.StatusUnauthorized, response.StatusCode)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityRoutesRegisterWaitCancelUnregisterTypedLifecycle(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
tok, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
router := gin.New()
|
||||||
|
registerAgentcompatRoutes(router)
|
||||||
|
server := httptest.NewServer(router)
|
||||||
|
t.Cleanup(server.Close)
|
||||||
|
|
||||||
|
capability := registerAgentcompatCapability(t, server.URL, token, agentcompatCapabilityRegisterRequest{Purpose: "terminal", ServerID: 7})
|
||||||
|
require.NoError(t, rpc.NezhaHandlerSingleton.CreateStreamWithPurpose("private-stream", userID, 7, rpc.PurposeTerminal))
|
||||||
|
parsed, err := rpc.ParseAgentCompatIOStreamCapability(capability)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, rpc.NezhaHandlerSingleton.BindAgentCompatIOStreamCapability(rpc.AgentCompatCapabilityBinding{
|
||||||
|
AgentCompatCapabilityAccess: rpc.AgentCompatCapabilityAccess{
|
||||||
|
Capability: parsed,
|
||||||
|
Owner: rpc.AgentCompatCapabilityOwner{PATID: tok.ID, UserID: userID},
|
||||||
|
Purpose: rpc.AgentCompatCapabilityTerminal, TargetServerID: 7,
|
||||||
|
ServerAccessAllowed: true,
|
||||||
|
},
|
||||||
|
StreamID: "private-stream",
|
||||||
|
}))
|
||||||
|
waitContext, cancelWait := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancelWait()
|
||||||
|
result, err := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](waitContext, server.URL+agentcompatCapabilityWaitPath, token, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "terminal", ServerID: 7})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, "private-stream", result.StreamID)
|
||||||
|
|
||||||
|
status, body := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityUnregisterPath, token, agentcompatCapabilityUnregisterRequest{Capability: capability, Purpose: "terminal", ServerID: 7})
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.NotContains(t, body, capability)
|
||||||
|
|
||||||
|
status, body = postAgentcompatEmpty(t, server.URL+agentcompatCapabilityCancelPath, token, agentcompatCapabilityCancelRequest{Capability: capability, Purpose: "terminal", ServerID: 7})
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.NotContains(t, body, capability)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityRoutesWaitCancellationDoesNotLeakSecrets(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
router := gin.New()
|
||||||
|
registerAgentcompatRoutes(router)
|
||||||
|
server := httptest.NewServer(router)
|
||||||
|
t.Cleanup(server.Close)
|
||||||
|
|
||||||
|
request := agentcompatCapabilityWaitRequest{Capability: strings.Repeat("a", 43), Purpose: "terminal", ServerID: 7}
|
||||||
|
status, body := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityWaitPath, token, request)
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.NotContains(t, body, request.Capability)
|
||||||
|
require.NotContains(t, body, "private-stream")
|
||||||
|
}
|
||||||
|
|
||||||
|
func registerAgentcompatCapability(t *testing.T, baseURL, token string, request agentcompatCapabilityRegisterRequest) string {
|
||||||
|
t.Helper()
|
||||||
|
requestContext, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
response, err := postAgentcompatCapability[agentcompatCapabilityRegisterRequest, agentcompatCapabilityRegisterResponse](requestContext, baseURL+agentcompatCapabilityRegisterPath, token, request)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotEmpty(t, response.Capability)
|
||||||
|
return response.Capability
|
||||||
|
}
|
||||||
|
|
||||||
|
func postAgentcompatEmpty(t *testing.T, path, token string, body any) (int, string) {
|
||||||
|
t.Helper()
|
||||||
|
encoded, err := json.Marshal(body)
|
||||||
|
require.NoError(t, err)
|
||||||
|
requestContext, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
request, err := http.NewRequestWithContext(requestContext, http.MethodPost, path, bytes.NewReader(encoded))
|
||||||
|
require.NoError(t, err)
|
||||||
|
request.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
request.Header.Set("Content-Type", "application/json")
|
||||||
|
response, err := http.DefaultClient.Do(request)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer response.Body.Close()
|
||||||
|
bodyBytes, err := io.ReadAll(response.Body)
|
||||||
|
require.NoError(t, err)
|
||||||
|
return response.StatusCode, string(bodyBytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
func postAgentcompatCapability[Request, Response any](ctx context.Context, path, token string, body Request) (Response, error) {
|
||||||
|
var zero Response
|
||||||
|
encoded, err := json.Marshal(body)
|
||||||
|
if err != nil {
|
||||||
|
return zero, err
|
||||||
|
}
|
||||||
|
request, err := http.NewRequestWithContext(ctx, http.MethodPost, path, bytes.NewReader(encoded))
|
||||||
|
if err != nil {
|
||||||
|
return zero, err
|
||||||
|
}
|
||||||
|
request.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
request.Header.Set("Content-Type", "application/json")
|
||||||
|
response, err := http.DefaultClient.Do(request)
|
||||||
|
if err != nil {
|
||||||
|
return zero, err
|
||||||
|
}
|
||||||
|
defer response.Body.Close()
|
||||||
|
var envelope model.CommonResponse[Response]
|
||||||
|
if err := json.NewDecoder(response.Body).Decode(&envelope); err != nil {
|
||||||
|
return zero, err
|
||||||
|
}
|
||||||
|
if !envelope.Success {
|
||||||
|
return zero, errors.New(envelope.Error)
|
||||||
|
}
|
||||||
|
return envelope.Data, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,238 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityCancelAndUnregisterAreUniformAndNonMutating(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
handler := rpc.NewNezhaHandler()
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = handler
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
ownerToken, ownerPAT := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
_, foreignPAT := mkDistinctCapabilityToken(t, userID, "foreign")
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
capability := registerAgentcompatCapability(t, server.URL, ownerPAT, agentcompatCapabilityRegisterRequest{Purpose: agentcompatCapabilityPurposeTerminal, ServerID: 7})
|
||||||
|
bindAgentcompatTerminalCapability(t, handler, capability, ownerToken.ID, userID, 7, "private-stream")
|
||||||
|
start := handler.SnapshotIOStreamState()
|
||||||
|
unknown := base64.RawURLEncoding.EncodeToString([]byte(strings.Repeat("u", 32)))
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
path string
|
||||||
|
pat string
|
||||||
|
body string
|
||||||
|
}{
|
||||||
|
{name: "malformed cancel", path: agentcompatCapabilityCancelPath, pat: ownerPAT, body: `{"capability":`},
|
||||||
|
{name: "oversize cancel", path: agentcompatCapabilityCancelPath, pat: ownerPAT, body: strings.Repeat("x", 513)},
|
||||||
|
{name: "duplicate cancel", path: agentcompatCapabilityCancelPath, pat: ownerPAT, body: `{"capability":"x","capability":"y","purpose":"terminal","server_id":7}`},
|
||||||
|
{name: "unknown cancel", path: agentcompatCapabilityCancelPath, pat: ownerPAT, body: capabilityAccessJSON(unknown, "terminal", 7, 0)},
|
||||||
|
{name: "foreign cancel", path: agentcompatCapabilityCancelPath, pat: foreignPAT, body: capabilityAccessJSON(capability, "terminal", 7, 0)},
|
||||||
|
{name: "purpose cancel", path: agentcompatCapabilityCancelPath, pat: ownerPAT, body: capabilityAccessJSON(capability, "file_manager", 7, 0)},
|
||||||
|
{name: "server cancel", path: agentcompatCapabilityCancelPath, pat: ownerPAT, body: capabilityAccessJSON(capability, "terminal", 999999, 0)},
|
||||||
|
{name: "invalid unregister", path: agentcompatCapabilityUnregisterPath, pat: ownerPAT, body: capabilityAccessJSON("invalid", "terminal", 7, 0)},
|
||||||
|
{name: "oversize unregister", path: agentcompatCapabilityUnregisterPath, pat: ownerPAT, body: strings.Repeat("x", 513)},
|
||||||
|
{name: "duplicate unregister", path: agentcompatCapabilityUnregisterPath, pat: ownerPAT, body: `{"capability":"x","purpose":"terminal","server_id":7,"server_id":8}`},
|
||||||
|
{name: "foreign unregister", path: agentcompatCapabilityUnregisterPath, pat: foreignPAT, body: capabilityAccessJSON(capability, "terminal", 7, 0)},
|
||||||
|
{name: "resource unregister", path: agentcompatCapabilityUnregisterPath, pat: ownerPAT, body: capabilityAccessJSON(capability, "terminal", 7, 41)},
|
||||||
|
}
|
||||||
|
var uniformBody string
|
||||||
|
for _, testCase := range cases {
|
||||||
|
t.Run(testCase.name, func(t *testing.T) {
|
||||||
|
status, body := postAgentcompatRaw(t, server.URL+testCase.path, testCase.pat, testCase.body)
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
if uniformBody == "" {
|
||||||
|
uniformBody = body
|
||||||
|
}
|
||||||
|
require.Equal(t, uniformBody, body)
|
||||||
|
require.Equal(t, start, handler.SnapshotIOStreamState())
|
||||||
|
require.NotContains(t, body, capability)
|
||||||
|
require.NotContains(t, body, "private-stream")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
waitContext, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
waited, err := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](waitContext, server.URL+agentcompatCapabilityWaitPath, ownerPAT, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "terminal", ServerID: 7})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, "private-stream", waited.StreamID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityPermissionRevocationIsRecheckedWithoutMutation(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
handler := rpc.NewNezhaHandler()
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = handler
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
token, plaintext := mkToken(t, userID, []string{model.ScopeServerRead}, []uint64{7})
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
capability := registerAgentcompatCapability(t, server.URL, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "terminal", ServerID: 7})
|
||||||
|
bindAgentcompatTerminalCapability(t, handler, capability, token.ID, userID, 7, "revoked-private-stream")
|
||||||
|
start := handler.SnapshotIOStreamState()
|
||||||
|
token.SetServerIDs([]uint64{99})
|
||||||
|
require.NoError(t, singleton.DB.Model(token).Update("servers_csv", token.ServersCSV).Error)
|
||||||
|
|
||||||
|
cancelStatus, cancelBody := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityCancelPath, plaintext, agentcompatCapabilityCancelRequest{Capability: capability, Purpose: "terminal", ServerID: 7})
|
||||||
|
unregisterStatus, unregisterBody := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityUnregisterPath, plaintext, agentcompatCapabilityUnregisterRequest{Capability: capability, Purpose: "terminal", ServerID: 7})
|
||||||
|
require.Equal(t, http.StatusOK, cancelStatus)
|
||||||
|
require.Equal(t, http.StatusOK, unregisterStatus)
|
||||||
|
require.Equal(t, cancelBody, unregisterBody)
|
||||||
|
require.Equal(t, start, handler.SnapshotIOStreamState())
|
||||||
|
|
||||||
|
token.SetServerIDs(nil)
|
||||||
|
require.NoError(t, singleton.DB.Model(token).Update("servers_csv", "").Error)
|
||||||
|
waitContext, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
waited, err := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](waitContext, server.URL+agentcompatCapabilityWaitPath, plaintext, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "terminal", ServerID: 7})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, "revoked-private-stream", waited.StreamID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityForeignWaitAndRevokedPATCannotObserveBinding(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
handler := rpc.NewNezhaHandler()
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = handler
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
ownerToken, ownerPAT := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
_, foreignPAT := mkDistinctCapabilityToken(t, userID, "foreign-wait")
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
capability := registerAgentcompatCapability(t, server.URL, ownerPAT, agentcompatCapabilityRegisterRequest{Purpose: "terminal", ServerID: 7})
|
||||||
|
bindAgentcompatTerminalCapability(t, handler, capability, ownerToken.ID, userID, 7, "foreign-private-stream")
|
||||||
|
start := handler.SnapshotIOStreamState()
|
||||||
|
|
||||||
|
foreignContext, cancelForeign := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancelForeign()
|
||||||
|
_, foreignErr := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](foreignContext, server.URL+agentcompatCapabilityWaitPath, foreignPAT, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "terminal", ServerID: 7})
|
||||||
|
require.Error(t, foreignErr)
|
||||||
|
require.NotContains(t, foreignErr.Error(), capability)
|
||||||
|
require.NotContains(t, foreignErr.Error(), "foreign-private-stream")
|
||||||
|
require.Equal(t, start, handler.SnapshotIOStreamState())
|
||||||
|
|
||||||
|
require.NoError(t, singleton.DB.Delete(&model.APIToken{}, ownerToken.ID).Error)
|
||||||
|
status, body := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityWaitPath, ownerPAT, agentcompatCapabilityAccessRequest{Capability: capability, Purpose: "terminal", ServerID: 7})
|
||||||
|
require.Equal(t, http.StatusUnauthorized, status)
|
||||||
|
require.NotContains(t, body, capability)
|
||||||
|
require.NotContains(t, body, "foreign-private-stream")
|
||||||
|
require.Equal(t, start, handler.SnapshotIOStreamState())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityRegisterRequiresCurrentServerWhitelist(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, plaintext := mkToken(t, userID, []string{model.ScopeServerRead}, []uint64{99})
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
|
||||||
|
status, body := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityRegisterPath, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "file_manager", ServerID: 7})
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.Contains(t, body, errAgentcompatCapabilityUnavailable.Error())
|
||||||
|
require.NotContains(t, body, "99")
|
||||||
|
require.NotContains(t, body, plaintext)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatCapabilityRegisterRejectsInvalidIdentityWithoutDisclosure(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, plaintext := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
privateServer := "777777777"
|
||||||
|
cases := []string{
|
||||||
|
`{`,
|
||||||
|
`{"purpose":"unknown","server_id":7}`,
|
||||||
|
`{"purpose":"terminal","server_id":0}`,
|
||||||
|
`{"purpose":"terminal","server_id":7,"resource_id":41}`,
|
||||||
|
`{"purpose":"nat","server_id":7}`,
|
||||||
|
`{"purpose":"terminal","server_id":` + privateServer + `}`,
|
||||||
|
}
|
||||||
|
for _, body := range cases {
|
||||||
|
status, responseBody := postAgentcompatRaw(t, server.URL+agentcompatCapabilityRegisterPath, plaintext, body)
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.NotContains(t, responseBody, privateServer)
|
||||||
|
require.NotContains(t, responseBody, plaintext)
|
||||||
|
var envelope model.CommonResponse[agentcompatCapabilityRegisterResponse]
|
||||||
|
require.NoError(t, json.Unmarshal([]byte(responseBody), &envelope))
|
||||||
|
require.False(t, envelope.Success)
|
||||||
|
require.Empty(t, envelope.Data.Capability)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAgentcompatCapabilityServer(t *testing.T) *httptest.Server {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
router := gin.New()
|
||||||
|
registerAgentcompatRoutes(router)
|
||||||
|
server := httptest.NewServer(router)
|
||||||
|
t.Cleanup(server.Close)
|
||||||
|
return server
|
||||||
|
}
|
||||||
|
|
||||||
|
func mkDistinctCapabilityToken(t *testing.T, userID uint64, suffix string) (*model.APIToken, string) {
|
||||||
|
t.Helper()
|
||||||
|
plaintext := "nzp_" + strings.Repeat("z", 32) + "_" + suffix
|
||||||
|
token := &model.APIToken{UserID: userID, Name: suffix, TokenHash: model.HashAPIToken(plaintext)}
|
||||||
|
token.SetScopes([]string{model.ScopeServerRead})
|
||||||
|
require.NoError(t, singleton.DB.Create(token).Error)
|
||||||
|
return token, plaintext
|
||||||
|
}
|
||||||
|
|
||||||
|
func bindAgentcompatTerminalCapability(t *testing.T, handler *rpc.NezhaHandler, rawCapability string, tokenID, userID, serverID uint64, streamID string) {
|
||||||
|
t.Helper()
|
||||||
|
capability, err := rpc.ParseAgentCompatIOStreamCapability(rawCapability)
|
||||||
|
require.NoError(t, err)
|
||||||
|
access := rpc.AgentCompatCapabilityAccess{
|
||||||
|
Capability: capability, Owner: rpc.AgentCompatCapabilityOwner{PATID: tokenID, UserID: userID},
|
||||||
|
Purpose: rpc.AgentCompatCapabilityTerminal, TargetServerID: serverID, ServerAccessAllowed: true,
|
||||||
|
}
|
||||||
|
require.NoError(t, handler.CreateStreamWithPurpose(streamID, userID, serverID, rpc.PurposeTerminal))
|
||||||
|
require.NoError(t, handler.BindAgentCompatIOStreamCapability(rpc.AgentCompatCapabilityBinding{AgentCompatCapabilityAccess: access, StreamID: streamID}))
|
||||||
|
}
|
||||||
|
|
||||||
|
func capabilityAccessJSON(capability, purpose string, serverID, resourceID uint64) string {
|
||||||
|
body, _ := json.Marshal(agentcompatCapabilityAccessRequest{Capability: capability, Purpose: purpose, ServerID: serverID, ResourceID: resourceID})
|
||||||
|
return string(body)
|
||||||
|
}
|
||||||
|
|
||||||
|
func postAgentcompatRaw(t *testing.T, path, token, body string) (int, string) {
|
||||||
|
t.Helper()
|
||||||
|
requestContext, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
request, err := http.NewRequestWithContext(requestContext, http.MethodPost, path, bytes.NewBufferString(body))
|
||||||
|
require.NoError(t, err)
|
||||||
|
request.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
request.Header.Set("Content-Type", "application/json")
|
||||||
|
response, err := http.DefaultClient.Do(request)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer response.Body.Close()
|
||||||
|
responseBody, err := io.ReadAll(response.Body)
|
||||||
|
require.NoError(t, err)
|
||||||
|
return response.StatusCode, string(responseBody)
|
||||||
|
}
|
||||||
@@ -0,0 +1,147 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
agentcompatOversizeWriteSentinel = "agentcompat:oversize-write-contract"
|
||||||
|
agentcompatOversizeWriteServerPath = "/tmp/agentcompat-oversize-contract.txt"
|
||||||
|
agentcompatFsWriteOperationOversize = "oversize"
|
||||||
|
)
|
||||||
|
|
||||||
|
type agentcompatFsWriteContractRequest struct {
|
||||||
|
ServerID uint64 `json:"server_id"`
|
||||||
|
Operation agentcompatFsWriteOperation `json:"operation"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type agentcompatFsWriteOperation string
|
||||||
|
|
||||||
|
type agentcompatFsWriteContractResponse struct {
|
||||||
|
Result model.FsWriteResult `json:"result"`
|
||||||
|
AgentRPCResponse bool `json:"agent_rpc_response"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func registerAgentcompatRoutes(router *gin.Engine) {
|
||||||
|
patAuth := requiredAgentcompatPAT(apiTokenAuthMiddleware())
|
||||||
|
router.POST("/agentcompat/fs-write-contract", patAuth, commonHandler(agentcompatFsWriteContract))
|
||||||
|
router.POST("/agentcompat/mcp-rate-limit-probe", patAuth, commonHandler(agentcompatMCPRateLimitProbeRoute))
|
||||||
|
router.GET("/agentcompat/io-stream-state", patAuth, commonHandler(agentcompatIOStreamSnapshot))
|
||||||
|
router.POST("/agentcompat/io-stream-state", patAuth, commonHandler(agentcompatIOStreamWait))
|
||||||
|
router.POST("/agentcompat/io-stream-quota-probe", patAuth, commonHandler(agentcompatIOStreamQuotaProbeRoute))
|
||||||
|
registerAgentcompatCapabilityRoutes(router, patAuth)
|
||||||
|
registerAgentcompatSQLiteHoldRoutes(router, patAuth)
|
||||||
|
}
|
||||||
|
|
||||||
|
type agentcompatIOStreamQuotaProbeResponse struct {
|
||||||
|
UserAccepted int `json:"user_accepted"`
|
||||||
|
UserRejected int `json:"user_rejected"`
|
||||||
|
ServerAccepted int `json:"server_accepted"`
|
||||||
|
ServerRejected int `json:"server_rejected"`
|
||||||
|
Clean bool `json:"clean"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func agentcompatIOStreamQuotaProbeRoute(context *gin.Context) (agentcompatIOStreamQuotaProbeResponse, error) {
|
||||||
|
if rpc.NezhaHandlerSingleton == nil {
|
||||||
|
return agentcompatIOStreamQuotaProbeResponse{}, errors.New("IOStream handler is unavailable")
|
||||||
|
}
|
||||||
|
result := rpc.RunIOStreamQuotaProbe(context.Request.Context())
|
||||||
|
if result.Err != nil {
|
||||||
|
return agentcompatIOStreamQuotaProbeResponse{}, result.Err
|
||||||
|
}
|
||||||
|
return agentcompatIOStreamQuotaProbeResponse{UserAccepted: result.UserAccepted, UserRejected: result.UserRejected, ServerAccepted: result.ServerAccepted, ServerRejected: result.ServerRejected, Clean: result.TrackedStreams == 0}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func requiredAgentcompatPAT(auth gin.HandlerFunc) gin.HandlerFunc {
|
||||||
|
return func(context *gin.Context) {
|
||||||
|
auth(context)
|
||||||
|
if context.IsAborted() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if APITokenFromContext(context) == nil {
|
||||||
|
abortAPITokenUnauthorized(context, "api token required")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func agentcompatIOStreamSnapshot(*gin.Context) (rpc.IOStreamState, error) {
|
||||||
|
if rpc.NezhaHandlerSingleton == nil {
|
||||||
|
return rpc.IOStreamState{}, errors.New("IOStream handler is unavailable")
|
||||||
|
}
|
||||||
|
return rpc.NezhaHandlerSingleton.SnapshotIOStreamState(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func agentcompatIOStreamWait(context *gin.Context) (rpc.IOStreamState, error) {
|
||||||
|
if rpc.NezhaHandlerSingleton == nil {
|
||||||
|
return rpc.IOStreamState{}, errors.New("IOStream handler is unavailable")
|
||||||
|
}
|
||||||
|
var expectation rpc.IOStreamStateExpectation
|
||||||
|
if err := context.ShouldBindJSON(&expectation); err != nil {
|
||||||
|
return rpc.IOStreamState{}, err
|
||||||
|
}
|
||||||
|
return rpc.NezhaHandlerSingleton.WaitForIOStreamState(context.Request.Context(), expectation)
|
||||||
|
}
|
||||||
|
|
||||||
|
func agentcompatFsWriteContract(context *gin.Context) (agentcompatFsWriteContractResponse, error) {
|
||||||
|
var request agentcompatFsWriteContractRequest
|
||||||
|
if err := decodeAgentcompatJSON(context, &request); err != nil {
|
||||||
|
return agentcompatFsWriteContractResponse{}, err
|
||||||
|
}
|
||||||
|
if request.ServerID == 0 || request.Operation == "" {
|
||||||
|
return agentcompatFsWriteContractResponse{}, errors.New("server_id and operation are required")
|
||||||
|
}
|
||||||
|
if request.Operation != agentcompatFsWriteOperation(agentcompatFsWriteOperationOversize) {
|
||||||
|
return agentcompatFsWriteContractResponse{}, errors.New("unknown filesystem write operation")
|
||||||
|
}
|
||||||
|
token := APITokenFromContext(context)
|
||||||
|
if token == nil || !token.HasScope(model.ScopeServerWrite) {
|
||||||
|
return agentcompatFsWriteContractResponse{}, errors.New("missing required scope: " + model.ScopeServerWrite)
|
||||||
|
}
|
||||||
|
server, err := requireServerAccess(context, request.ServerID)
|
||||||
|
if err != nil {
|
||||||
|
return agentcompatFsWriteContractResponse{}, err
|
||||||
|
}
|
||||||
|
content := "denied"
|
||||||
|
if request.Operation == agentcompatFsWriteOperation(agentcompatFsWriteOperationOversize) {
|
||||||
|
content = agentcompatOversizeWriteSentinel
|
||||||
|
}
|
||||||
|
// This tagged probe owns both path and payload; accepting either from callers would make it an arbitrary-write endpoint.
|
||||||
|
raw, err := rpc.CallAgent(context.Request.Context(), server.ID, model.TaskTypeFsWrite, model.FsWriteRequest{
|
||||||
|
Path: agentcompatOversizeWriteServerPath, Content: content, Encoding: "utf8", Mode: "0600",
|
||||||
|
}, 30*time.Second)
|
||||||
|
if err != nil {
|
||||||
|
return agentcompatFsWriteContractResponse{}, err
|
||||||
|
}
|
||||||
|
var result model.FsWriteResult
|
||||||
|
if err := json.Unmarshal(raw, &result); err != nil {
|
||||||
|
return agentcompatFsWriteContractResponse{}, err
|
||||||
|
}
|
||||||
|
return agentcompatFsWriteContractResponse{Result: result, AgentRPCResponse: true}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeAgentcompatJSON(context *gin.Context, value any) error {
|
||||||
|
decoder := json.NewDecoder(context.Request.Body)
|
||||||
|
decoder.DisallowUnknownFields()
|
||||||
|
if err := decoder.Decode(value); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
var trailing any
|
||||||
|
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
|
||||||
|
if err == nil {
|
||||||
|
return fmt.Errorf("trailing JSON values are not allowed")
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
//go:build !agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import "github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
func registerAgentcompatRoutes(*gin.Engine) {}
|
||||||
@@ -0,0 +1,69 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAgentcompatFsWriteContractRejectsCallerPayloadAndForeignServer(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
_, plaintext := mkToken(t, userID, []string{model.ScopeServerWrite}, []uint64{99})
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
|
||||||
|
status, body := postAgentcompatRaw(t, server.URL+"/agentcompat/fs-write-contract", plaintext, `{"server_id":7,"path":"/tmp/foreign","operation":"oversize","content":"attacker"}`)
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.NotContains(t, body, "attacker")
|
||||||
|
require.Contains(t, body, "unknown field")
|
||||||
|
|
||||||
|
status, body = postAgentcompatRaw(t, server.URL+"/agentcompat/fs-write-contract", plaintext, `{"server_id":7,"operation":"oversize"}`)
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.Contains(t, body, "permission denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatFsWriteContractRequiresWriteScopeBeforeRPC(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
_, plaintext := mkToken(t, userID, []string{model.ScopeServerRead}, nil)
|
||||||
|
server := newAgentcompatCapabilityServer(t)
|
||||||
|
|
||||||
|
status, body := postAgentcompatRaw(t, server.URL+"/agentcompat/fs-write-contract", plaintext, `{"server_id":7,"operation":"oversize"}`)
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.Contains(t, body, "missing required scope")
|
||||||
|
require.NotContains(t, body, `"agent_rpc_response":true`)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatProbeJSONRejectsUnknownFieldsAndTrailingValues(t *testing.T) {
|
||||||
|
tests := []string{
|
||||||
|
`{"server_id":7,"operation":"oversize","path":"caller"}`,
|
||||||
|
`{"server_id":7,"operation":"oversize","content":"caller"}`,
|
||||||
|
`{"server_id":7,"operation":"oversize"}{}`,
|
||||||
|
}
|
||||||
|
for _, body := range tests {
|
||||||
|
t.Run(body, func(t *testing.T) {
|
||||||
|
context := newAgentcompatJSONContext(t, body)
|
||||||
|
var request agentcompatFsWriteContractRequest
|
||||||
|
require.Error(t, decodeAgentcompatJSON(context, &request))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAgentcompatJSONContext(t *testing.T, body string) *gin.Context {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
context, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||||
|
context.Request = httptest.NewRequest(http.MethodPost, "/", bytes.NewBufferString(body))
|
||||||
|
context.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
return context
|
||||||
|
}
|
||||||
@@ -0,0 +1,40 @@
|
|||||||
|
//go:build agentcompat && linux
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAgentcompatSQLiteHoldErrorUsesFixedRedactedMessages(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
cause error
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{agentcompatSQLiteHoldInvalidRequest{}, "agentcompat sqlite hold request is invalid"},
|
||||||
|
{singleton.ErrSQLiteHoldSessionActive, "agentcompat sqlite hold is active"},
|
||||||
|
{singleton.ErrSQLiteHoldStaleSession, "agentcompat sqlite hold receipt is stale"},
|
||||||
|
{singleton.ErrSQLiteHoldFinalizationNotStarted, "agentcompat sqlite hold is not ready"},
|
||||||
|
{singleton.ErrSQLiteHoldFinalizationStarted, "agentcompat sqlite hold is not ready"},
|
||||||
|
{singleton.ErrSQLiteHoldUnexpectedSelection, "agentcompat sqlite hold was aborted"},
|
||||||
|
{singleton.ErrSQLiteHoldAmbiguousCandidate, "agentcompat sqlite hold was aborted"},
|
||||||
|
{singleton.ErrSQLiteHoldAborted, "agentcompat sqlite hold was aborted"},
|
||||||
|
{context.Canceled, "agentcompat sqlite hold wait was canceled"},
|
||||||
|
{context.DeadlineExceeded, "agentcompat sqlite hold wait was canceled"},
|
||||||
|
{errors.New("path=/secret token=nzp_secret"), "agentcompat sqlite hold control is unavailable"},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, test := range tests {
|
||||||
|
// When
|
||||||
|
err := agentcompatSQLiteHoldError(test.cause)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
require.EqualError(t, err, test.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,168 @@
|
|||||||
|
//go:build agentcompat && linux
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
const agentcompatSQLiteHoldPath = "/agentcompat/sqlite-hold/"
|
||||||
|
|
||||||
|
type agentcompatSQLiteHoldControl interface {
|
||||||
|
ArmNextSQLiteHold() (singleton.SQLiteHoldReceipt, error)
|
||||||
|
WaitSQLiteHoldSelected(context.Context, singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error)
|
||||||
|
WaitSQLiteHoldFinalizing(context.Context, singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error)
|
||||||
|
SnapshotSQLiteHold(singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error)
|
||||||
|
ReleaseSQLiteHold(singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error)
|
||||||
|
AbortSQLiteHold(singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
type agentcompatSQLiteHoldFacade struct{}
|
||||||
|
|
||||||
|
func (agentcompatSQLiteHoldFacade) ArmNextSQLiteHold() (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
return singleton.ArmNextSQLiteHold()
|
||||||
|
}
|
||||||
|
func (agentcompatSQLiteHoldFacade) WaitSQLiteHoldSelected(ctx context.Context, receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
return singleton.WaitSQLiteHoldSelected(ctx, receipt)
|
||||||
|
}
|
||||||
|
func (agentcompatSQLiteHoldFacade) WaitSQLiteHoldFinalizing(ctx context.Context, receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
return singleton.WaitSQLiteHoldFinalizing(ctx, receipt)
|
||||||
|
}
|
||||||
|
func (agentcompatSQLiteHoldFacade) SnapshotSQLiteHold(receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
return singleton.SnapshotSQLiteHold(receipt)
|
||||||
|
}
|
||||||
|
func (agentcompatSQLiteHoldFacade) ReleaseSQLiteHold(receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
return singleton.ReleaseSQLiteHold(receipt)
|
||||||
|
}
|
||||||
|
func (agentcompatSQLiteHoldFacade) AbortSQLiteHold(receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
return singleton.AbortSQLiteHold(receipt)
|
||||||
|
}
|
||||||
|
|
||||||
|
var agentcompatSQLiteHoldController agentcompatSQLiteHoldControl = agentcompatSQLiteHoldFacade{}
|
||||||
|
|
||||||
|
type agentcompatSQLiteHoldRequest struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
State singleton.SQLiteHoldControlState `json:"state"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func registerAgentcompatSQLiteHoldRoutes(router *gin.Engine, patAuth gin.HandlerFunc) {
|
||||||
|
readOnlyPAT := func(c *gin.Context) { suppressAPITokenAuthWrites(c) }
|
||||||
|
router.POST(agentcompatSQLiteHoldPath+"arm", readOnlyPAT, patAuth, commonHandler(agentcompatSQLiteHoldArm))
|
||||||
|
router.POST(agentcompatSQLiteHoldPath+"wait", readOnlyPAT, patAuth, commonHandler(agentcompatSQLiteHoldWait))
|
||||||
|
router.POST(agentcompatSQLiteHoldPath+"snapshot", readOnlyPAT, patAuth, commonHandler(agentcompatSQLiteHoldSnapshot))
|
||||||
|
router.POST(agentcompatSQLiteHoldPath+"release", readOnlyPAT, patAuth, commonHandler(agentcompatSQLiteHoldRelease))
|
||||||
|
router.POST(agentcompatSQLiteHoldPath+"abort", readOnlyPAT, patAuth, commonHandler(agentcompatSQLiteHoldAbort))
|
||||||
|
}
|
||||||
|
|
||||||
|
func agentcompatSQLiteHoldArm(c *gin.Context) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
var request struct{}
|
||||||
|
if err := decodeAgentcompatSQLiteHoldRequest(c, &request); err != nil {
|
||||||
|
return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldError(err)
|
||||||
|
}
|
||||||
|
receipt, err := agentcompatSQLiteHoldController.ArmNextSQLiteHold()
|
||||||
|
return receipt, agentcompatSQLiteHoldError(err)
|
||||||
|
}
|
||||||
|
func agentcompatSQLiteHoldWait(c *gin.Context) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
receipt, err := decodeAgentcompatSQLiteHoldReceipt(c, true)
|
||||||
|
if err != nil {
|
||||||
|
return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldError(err)
|
||||||
|
}
|
||||||
|
switch receipt.State {
|
||||||
|
case singleton.SQLiteHoldControlStateSelected:
|
||||||
|
receipt, err = agentcompatSQLiteHoldController.WaitSQLiteHoldSelected(c.Request.Context(), receipt)
|
||||||
|
case singleton.SQLiteHoldControlStateFinalizing:
|
||||||
|
receipt, err = agentcompatSQLiteHoldController.WaitSQLiteHoldFinalizing(c.Request.Context(), receipt)
|
||||||
|
default:
|
||||||
|
return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldInvalidRequest{}
|
||||||
|
}
|
||||||
|
return receipt, agentcompatSQLiteHoldError(err)
|
||||||
|
}
|
||||||
|
func agentcompatSQLiteHoldSnapshot(c *gin.Context) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
receipt, err := decodeAgentcompatSQLiteHoldReceipt(c, false)
|
||||||
|
if err != nil {
|
||||||
|
return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldError(err)
|
||||||
|
}
|
||||||
|
result, err := agentcompatSQLiteHoldController.SnapshotSQLiteHold(receipt)
|
||||||
|
return result, agentcompatSQLiteHoldError(err)
|
||||||
|
}
|
||||||
|
func agentcompatSQLiteHoldRelease(c *gin.Context) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
receipt, err := decodeAgentcompatSQLiteHoldReceipt(c, false)
|
||||||
|
if err != nil {
|
||||||
|
return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldError(err)
|
||||||
|
}
|
||||||
|
result, err := agentcompatSQLiteHoldController.ReleaseSQLiteHold(receipt)
|
||||||
|
return result, agentcompatSQLiteHoldError(err)
|
||||||
|
}
|
||||||
|
func agentcompatSQLiteHoldAbort(c *gin.Context) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
receipt, err := decodeAgentcompatSQLiteHoldReceipt(c, false)
|
||||||
|
if err != nil {
|
||||||
|
return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldError(err)
|
||||||
|
}
|
||||||
|
result, err := agentcompatSQLiteHoldController.AbortSQLiteHold(receipt)
|
||||||
|
return result, agentcompatSQLiteHoldError(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
type agentcompatSQLiteHoldInvalidRequest struct{}
|
||||||
|
|
||||||
|
func (agentcompatSQLiteHoldInvalidRequest) Error() string {
|
||||||
|
return "agentcompat sqlite hold request is invalid"
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeAgentcompatSQLiteHoldRequest(c *gin.Context, value any) error {
|
||||||
|
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, 512)
|
||||||
|
decoder := json.NewDecoder(c.Request.Body)
|
||||||
|
decoder.DisallowUnknownFields()
|
||||||
|
if err := decoder.Decode(value); err != nil {
|
||||||
|
return agentcompatSQLiteHoldInvalidRequest{}
|
||||||
|
}
|
||||||
|
var trailing any
|
||||||
|
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
|
||||||
|
return agentcompatSQLiteHoldInvalidRequest{}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func decodeAgentcompatSQLiteHoldReceipt(c *gin.Context, requireState bool) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
var request agentcompatSQLiteHoldRequest
|
||||||
|
if err := decodeAgentcompatSQLiteHoldRequest(c, &request); err != nil {
|
||||||
|
return singleton.SQLiteHoldReceipt{}, err
|
||||||
|
}
|
||||||
|
decoded, err := base64.RawURLEncoding.DecodeString(request.ID)
|
||||||
|
if err != nil || len(request.ID) != 43 || len(decoded) != 32 {
|
||||||
|
return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldInvalidRequest{}
|
||||||
|
}
|
||||||
|
if !requireState && request.State != "" {
|
||||||
|
return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldInvalidRequest{}
|
||||||
|
}
|
||||||
|
return singleton.SQLiteHoldReceipt{ID: request.ID, State: request.State}, nil
|
||||||
|
}
|
||||||
|
func agentcompatSQLiteHoldError(err error) error {
|
||||||
|
if err == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if errors.As(err, new(agentcompatSQLiteHoldInvalidRequest)) {
|
||||||
|
return agentcompatSQLiteHoldInvalidRequest{}
|
||||||
|
}
|
||||||
|
switch {
|
||||||
|
case errors.Is(err, singleton.ErrSQLiteHoldSessionActive):
|
||||||
|
return errors.New("agentcompat sqlite hold is active")
|
||||||
|
case errors.Is(err, singleton.ErrSQLiteHoldStaleSession):
|
||||||
|
return errors.New("agentcompat sqlite hold receipt is stale")
|
||||||
|
case errors.Is(err, singleton.ErrSQLiteHoldFinalizationNotStarted), errors.Is(err, singleton.ErrSQLiteHoldFinalizationStarted):
|
||||||
|
return errors.New("agentcompat sqlite hold is not ready")
|
||||||
|
case errors.Is(err, singleton.ErrSQLiteHoldUnexpectedSelection), errors.Is(err, singleton.ErrSQLiteHoldAmbiguousCandidate), errors.Is(err, singleton.ErrSQLiteHoldAborted):
|
||||||
|
return errors.New("agentcompat sqlite hold was aborted")
|
||||||
|
case errors.Is(err, context.Canceled), errors.Is(err, context.DeadlineExceeded):
|
||||||
|
return errors.New("agentcompat sqlite hold wait was canceled")
|
||||||
|
default:
|
||||||
|
return errors.New("agentcompat sqlite hold control is unavailable")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,201 @@
|
|||||||
|
//go:build agentcompat && linux
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
type agentcompatSQLiteHoldControlProbe struct {
|
||||||
|
err error
|
||||||
|
armCalls int
|
||||||
|
lastCall string
|
||||||
|
waitContextError error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (probe *agentcompatSQLiteHoldControlProbe) ArmNextSQLiteHold() (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
probe.armCalls++
|
||||||
|
probe.lastCall = "arm"
|
||||||
|
return agentcompatSQLiteHoldTestReceipt(singleton.SQLiteHoldControlStateArmed), probe.err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (probe *agentcompatSQLiteHoldControlProbe) WaitSQLiteHoldSelected(ctx context.Context, receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
probe.lastCall = "wait-selected"
|
||||||
|
probe.waitContextError = ctx.Err()
|
||||||
|
receipt.State = singleton.SQLiteHoldControlStateSelected
|
||||||
|
return receipt, probe.err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (probe *agentcompatSQLiteHoldControlProbe) WaitSQLiteHoldFinalizing(ctx context.Context, receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
probe.lastCall = "wait-finalizing"
|
||||||
|
probe.waitContextError = ctx.Err()
|
||||||
|
receipt.State = singleton.SQLiteHoldControlStateFinalizing
|
||||||
|
return receipt, probe.err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (probe *agentcompatSQLiteHoldControlProbe) SnapshotSQLiteHold(receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
probe.lastCall = "snapshot"
|
||||||
|
receipt.State = singleton.SQLiteHoldControlStateSelected
|
||||||
|
return receipt, probe.err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (probe *agentcompatSQLiteHoldControlProbe) ReleaseSQLiteHold(receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
probe.lastCall = "release"
|
||||||
|
receipt.State = singleton.SQLiteHoldControlStateReleased
|
||||||
|
return receipt, probe.err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (probe *agentcompatSQLiteHoldControlProbe) AbortSQLiteHold(receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) {
|
||||||
|
probe.lastCall = "abort"
|
||||||
|
receipt.State = singleton.SQLiteHoldControlStateAborted
|
||||||
|
return receipt, probe.err
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatSQLiteHoldRoutesRequirePATWithoutScopeOrAuthWrites(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
router, token, storedToken, probe := setupAgentcompatSQLiteHoldRouteTest(t)
|
||||||
|
|
||||||
|
// When
|
||||||
|
status, response := requestAgentcompatSQLiteHold(t, router, context.Background(), "arm", token, `{}`)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.True(t, response.Success)
|
||||||
|
require.Equal(t, agentcompatSQLiteHoldTestReceipt(singleton.SQLiteHoldControlStateArmed), response.Data)
|
||||||
|
require.Equal(t, 1, probe.armCalls)
|
||||||
|
var refreshed model.APIToken
|
||||||
|
require.NoError(t, singleton.DB.First(&refreshed, storedToken.ID).Error)
|
||||||
|
require.Nil(t, refreshed.LastUsedAt)
|
||||||
|
require.Empty(t, refreshed.LastUsedIP)
|
||||||
|
|
||||||
|
for _, authorization := range []string{"", "jwt-looking-value", "nzp_invalid"} {
|
||||||
|
status, _ = requestAgentcompatSQLiteHold(t, router, context.Background(), "arm", authorization, `{}`)
|
||||||
|
require.Equal(t, http.StatusUnauthorized, status)
|
||||||
|
}
|
||||||
|
var wafRows int64
|
||||||
|
require.NoError(t, singleton.DB.Model(&model.WAF{}).Count(&wafRows).Error)
|
||||||
|
require.Zero(t, wafRows)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatSQLiteHoldRoutesRejectInvalidBodiesBeforeControl(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
router, token, _, probe := setupAgentcompatSQLiteHoldRouteTest(t)
|
||||||
|
tests := []string{"", `{`, `{"unknown":true}`, `{} {}`, strings.Repeat(" ", 513) + `{}`}
|
||||||
|
|
||||||
|
for _, body := range tests {
|
||||||
|
// When
|
||||||
|
status, response := requestAgentcompatSQLiteHold(t, router, context.Background(), "arm", token, body)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.False(t, response.Success)
|
||||||
|
require.Equal(t, "agentcompat sqlite hold request is invalid", response.Error)
|
||||||
|
}
|
||||||
|
require.Zero(t, probe.armCalls)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatSQLiteHoldRoutesDispatchTypedLifecycleAndPropagateContext(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
router, token, _, probe := setupAgentcompatSQLiteHoldRouteTest(t)
|
||||||
|
receiptID := agentcompatSQLiteHoldTestReceipt("").ID
|
||||||
|
tests := []struct {
|
||||||
|
path string
|
||||||
|
body string
|
||||||
|
wantCall string
|
||||||
|
wantState singleton.SQLiteHoldControlState
|
||||||
|
}{
|
||||||
|
{"wait", `{"id":"` + receiptID + `","state":"selected"}`, "wait-selected", singleton.SQLiteHoldControlStateSelected},
|
||||||
|
{"wait", `{"id":"` + receiptID + `","state":"finalizing"}`, "wait-finalizing", singleton.SQLiteHoldControlStateFinalizing},
|
||||||
|
{"snapshot", `{"id":"` + receiptID + `"}`, "snapshot", singleton.SQLiteHoldControlStateSelected},
|
||||||
|
{"release", `{"id":"` + receiptID + `"}`, "release", singleton.SQLiteHoldControlStateReleased},
|
||||||
|
{"abort", `{"id":"` + receiptID + `"}`, "abort", singleton.SQLiteHoldControlStateAborted},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, test := range tests {
|
||||||
|
// When
|
||||||
|
status, response := requestAgentcompatSQLiteHold(t, router, context.Background(), test.path, token, test.body)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.True(t, response.Success)
|
||||||
|
require.Equal(t, test.wantCall, probe.lastCall)
|
||||||
|
require.Equal(t, test.wantState, response.Data.State)
|
||||||
|
}
|
||||||
|
|
||||||
|
canceledContext, cancel := context.WithCancel(context.Background())
|
||||||
|
cancel()
|
||||||
|
status, response := requestAgentcompatSQLiteHold(t, router, canceledContext, "wait", token, `{"id":"`+receiptID+`","state":"selected"}`)
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.True(t, response.Success)
|
||||||
|
require.ErrorIs(t, probe.waitContextError, context.Canceled)
|
||||||
|
|
||||||
|
status, response = requestAgentcompatSQLiteHold(t, router, context.Background(), "wait", token, `{"id":"`+receiptID+`","state":"armed"}`)
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.False(t, response.Success)
|
||||||
|
require.Equal(t, "agentcompat sqlite hold request is invalid", response.Error)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatSQLiteHoldRoutesRedactControlErrors(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
router, token, _, probe := setupAgentcompatSQLiteHoldRouteTest(t)
|
||||||
|
probe.err = errors.New("token=nzp_secret path=/root/private dashboard.sqlite-journal")
|
||||||
|
|
||||||
|
// When
|
||||||
|
status, response := requestAgentcompatSQLiteHold(t, router, context.Background(), "arm", token, `{}`)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
require.Equal(t, http.StatusOK, status)
|
||||||
|
require.False(t, response.Success)
|
||||||
|
require.Equal(t, "agentcompat sqlite hold control is unavailable", response.Error)
|
||||||
|
encoded, err := json.Marshal(response)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotContains(t, string(encoded), "nzp_secret")
|
||||||
|
require.NotContains(t, string(encoded), "/root/private")
|
||||||
|
require.NotContains(t, string(encoded), "journal")
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupAgentcompatSQLiteHoldRouteTest(t *testing.T) (*gin.Engine, string, *model.APIToken, *agentcompatSQLiteHoldControlProbe) {
|
||||||
|
t.Helper()
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
t.Cleanup(cleanup)
|
||||||
|
storedToken, token := mkToken(t, userID, nil, nil)
|
||||||
|
probe := &agentcompatSQLiteHoldControlProbe{}
|
||||||
|
originalControl := agentcompatSQLiteHoldController
|
||||||
|
agentcompatSQLiteHoldController = probe
|
||||||
|
t.Cleanup(func() { agentcompatSQLiteHoldController = originalControl })
|
||||||
|
router := gin.New()
|
||||||
|
registerAgentcompatSQLiteHoldRoutes(router, requiredAgentcompatPAT(apiTokenAuthMiddleware()))
|
||||||
|
return router, token, storedToken, probe
|
||||||
|
}
|
||||||
|
|
||||||
|
func requestAgentcompatSQLiteHold(t *testing.T, router *gin.Engine, ctx context.Context, path, authorization, body string) (int, model.CommonResponse[singleton.SQLiteHoldReceipt]) {
|
||||||
|
t.Helper()
|
||||||
|
request := httptest.NewRequestWithContext(ctx, http.MethodPost, agentcompatSQLiteHoldPath+path, bytes.NewBufferString(body))
|
||||||
|
request.Header.Set("Content-Type", "application/json")
|
||||||
|
if authorization != "" {
|
||||||
|
request.Header.Set("Authorization", "Bearer "+authorization)
|
||||||
|
}
|
||||||
|
responseRecorder := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(responseRecorder, request)
|
||||||
|
var response model.CommonResponse[singleton.SQLiteHoldReceipt]
|
||||||
|
require.NoError(t, json.Unmarshal(responseRecorder.Body.Bytes(), &response))
|
||||||
|
return responseRecorder.Code, response
|
||||||
|
}
|
||||||
|
|
||||||
|
func agentcompatSQLiteHoldTestReceipt(state singleton.SQLiteHoldControlState) singleton.SQLiteHoldReceipt {
|
||||||
|
return singleton.SQLiteHoldReceipt{ID: base64.RawURLEncoding.EncodeToString(bytes.Repeat([]byte{17}, 32)), State: state}
|
||||||
|
}
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
//go:build agentcompat && !linux
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import "github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
func registerAgentcompatSQLiteHoldRoutes(*gin.Engine, gin.HandlerFunc) {}
|
||||||
@@ -0,0 +1,27 @@
|
|||||||
|
//go:build !agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAgentcompatSQLiteHoldRoutesAreAbsentWithoutBuildTag(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
router := gin.New()
|
||||||
|
registerAgentcompatRoutes(router)
|
||||||
|
request := httptest.NewRequest(http.MethodPost, "/agentcompat/sqlite-hold/arm", nil)
|
||||||
|
response := httptest.NewRecorder()
|
||||||
|
|
||||||
|
// When
|
||||||
|
router.ServeHTTP(response, request)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
require.Equal(t, http.StatusNotFound, response.Code)
|
||||||
|
}
|
||||||
@@ -1,7 +1,7 @@
|
|||||||
package controller
|
package controller
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"maps"
|
"slices"
|
||||||
"strconv"
|
"strconv"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -167,9 +167,14 @@ func batchDeleteAlertRule(c *gin.Context) (any, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func validateRule(c *gin.Context, r *model.AlertRule) error {
|
func validateRule(c *gin.Context, r *model.AlertRule) error {
|
||||||
|
if !r.HasPermission(c) {
|
||||||
|
return singleton.Localizer.ErrorT("permission denied")
|
||||||
|
}
|
||||||
if len(r.Rules) > 0 {
|
if len(r.Rules) > 0 {
|
||||||
for _, rule := range r.Rules {
|
for _, rule := range r.Rules {
|
||||||
if !singleton.ServerShared.CheckPermission(c, maps.Keys(rule.Ignore)) {
|
switch rule.Cover {
|
||||||
|
case model.RuleCoverAll, model.RuleCoverIgnoreAll:
|
||||||
|
default:
|
||||||
return singleton.Localizer.ErrorT("permission denied")
|
return singleton.Localizer.ErrorT("permission denied")
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -192,5 +197,20 @@ func validateRule(c *gin.Context, r *model.AlertRule) error {
|
|||||||
} else {
|
} else {
|
||||||
return singleton.Localizer.ErrorT("need to configure at least a single rule")
|
return singleton.Localizer.ErrorT("need to configure at least a single rule")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if !singleton.CronShared.CheckPermission(c, slices.Values(r.FailTriggerTasks)) {
|
||||||
|
return singleton.Localizer.ErrorT("permission denied")
|
||||||
|
}
|
||||||
|
if !singleton.CronShared.CheckPermission(c, slices.Values(r.RecoverTriggerTasks)) {
|
||||||
|
return singleton.Localizer.ErrorT("permission denied")
|
||||||
|
}
|
||||||
|
if err := enforcePATTriggerTaskScope(c, r.FailTriggerTasks, r.RecoverTriggerTasks); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := assertOwnsNotificationGroup(c, r.NotificationGroupID); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,146 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/patrickmn/go-cache"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/i18n"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func setupAlertRuleFanoutFixture(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
originalDB := singleton.DB
|
||||||
|
originalCache := singleton.Cache
|
||||||
|
originalLoc := singleton.Loc
|
||||||
|
originalLocalizer := singleton.Localizer
|
||||||
|
originalServer := singleton.ServerShared
|
||||||
|
originalCron := singleton.CronShared
|
||||||
|
originalUserInfo := singleton.UserInfoMap
|
||||||
|
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
sqlDB, err := db.DB()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, db.AutoMigrate(&model.Server{}, &model.AlertRule{}, &model.Cron{}, &model.User{}))
|
||||||
|
|
||||||
|
singleton.DB = db
|
||||||
|
singleton.Loc = time.UTC
|
||||||
|
singleton.Cache = cache.New(time.Minute, time.Minute)
|
||||||
|
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
|
||||||
|
singleton.UserLock.Lock()
|
||||||
|
singleton.UserInfoMap = map[uint64]model.UserInfo{1: {Role: model.RoleAdmin}}
|
||||||
|
singleton.UserLock.Unlock()
|
||||||
|
|
||||||
|
require.NoError(t, db.Create(&model.Server{Common: model.Common{ID: 1, UserID: 1}, Name: "s1", UUID: "s1"}).Error)
|
||||||
|
require.NoError(t, db.Create(&model.Server{Common: model.Common{ID: 2, UserID: 1}, Name: "s2", UUID: "s2"}).Error)
|
||||||
|
|
||||||
|
singleton.ServerShared = singleton.NewServerClass()
|
||||||
|
singleton.CronShared = singleton.NewCronClass()
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
singleton.CronShared.Close()
|
||||||
|
_ = sqlDB.Close()
|
||||||
|
singleton.DB = originalDB
|
||||||
|
singleton.Cache = originalCache
|
||||||
|
singleton.Loc = originalLoc
|
||||||
|
singleton.Localizer = originalLocalizer
|
||||||
|
singleton.ServerShared = originalServer
|
||||||
|
singleton.CronShared = originalCron
|
||||||
|
singleton.UserLock.Lock()
|
||||||
|
singleton.UserInfoMap = originalUserInfo
|
||||||
|
singleton.UserLock.Unlock()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAlertRuleCtxWithPAT(t *testing.T, viewer *model.User, tok *model.APIToken, body any) *gin.Context {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||||
|
raw, _ := json.Marshal(body)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/alert-rule", bytes.NewReader(raw))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
if viewer != nil {
|
||||||
|
c.Set(model.CtxKeyAuthorizedUser, viewer)
|
||||||
|
}
|
||||||
|
if tok != nil {
|
||||||
|
c.Set(model.CtxKeyAPIToken, tok)
|
||||||
|
}
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// A server-limited PAT must not be able to create a RuleCoverAll rule with an
|
||||||
|
// empty Ignore (deny-list). Empty deny-list means "monitor every owner-visible
|
||||||
|
// server", which escapes the PAT's server_ids whitelist.
|
||||||
|
func TestCreateAlertRulePATCoverAllEmptyIgnoreRejected(t *testing.T) {
|
||||||
|
setupAlertRuleFanoutFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 5, UserID: 1}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
form := map[string]any{
|
||||||
|
"name": "all-servers",
|
||||||
|
"enable": false,
|
||||||
|
"rules": []map[string]any{
|
||||||
|
{"type": "offline", "cover": model.RuleCoverAll, "duration": 10},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
c := newAlertRuleCtxWithPAT(t, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin}, tok, form)
|
||||||
|
_, err := createAlertRule(c)
|
||||||
|
require.Error(t, err, "CoverAll + empty Ignore must be rejected for a PAT scoped to {1}")
|
||||||
|
|
||||||
|
var count int64
|
||||||
|
require.NoError(t, singleton.DB.Model(&model.AlertRule{}).Count(&count).Error)
|
||||||
|
assert.Equal(t, int64(0), count, "no alert rule should be persisted")
|
||||||
|
}
|
||||||
|
|
||||||
|
// The same PAT may create a RuleCoverAll rule when it explicitly denies every
|
||||||
|
// server outside its whitelist (here: server 2), since the fan-out is then
|
||||||
|
// confined to server 1. Exercised at validateRule to avoid the alert-sentinel
|
||||||
|
// side effects of a full createAlertRule.
|
||||||
|
func TestValidateRulePATCoverAllDenyingOutsideServersAllowed(t *testing.T) {
|
||||||
|
setupAlertRuleFanoutFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 6, UserID: 1}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
c := newAlertRuleCtxWithPAT(t, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin}, tok, nil)
|
||||||
|
|
||||||
|
r := &model.AlertRule{
|
||||||
|
Common: model.Common{UserID: 1},
|
||||||
|
Name: "only-server-1",
|
||||||
|
Rules: []*model.Rule{{Type: "offline", Cover: model.RuleCoverAll, Duration: 10, Ignore: map[uint64]bool{2: true}}},
|
||||||
|
}
|
||||||
|
require.NoError(t, validateRule(c, r), "CoverAll denying every out-of-whitelist server must be allowed")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Empty deny-list at validateRule level must also be rejected.
|
||||||
|
func TestValidateRulePATCoverAllEmptyIgnoreRejected(t *testing.T) {
|
||||||
|
setupAlertRuleFanoutFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 7, UserID: 1}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
c := newAlertRuleCtxWithPAT(t, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin}, tok, nil)
|
||||||
|
|
||||||
|
r := &model.AlertRule{
|
||||||
|
Common: model.Common{UserID: 1},
|
||||||
|
Name: "all",
|
||||||
|
Rules: []*model.Rule{{Type: "offline", Cover: model.RuleCoverAll, Duration: 10}},
|
||||||
|
}
|
||||||
|
require.Error(t, validateRule(c, r), "CoverAll + empty Ignore must be rejected for PAT scoped to {1}")
|
||||||
|
}
|
||||||
@@ -0,0 +1,294 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"net/http"
|
||||||
|
"slices"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/utils"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
apiTokenSecretLength = 32 // 明文 token 随机部分长度(hex 编码前)
|
||||||
|
apiTokenCtxKey = "nz_api_token" // #nosec G101 -- gin context key name, not a credential
|
||||||
|
apiTokenLastUsedCtxKey = "nz_api_token_used_marker" // #nosec G101 -- gin context key name, not a credential
|
||||||
|
apiTokenReadOnlyCtxKey = "nz_api_token_read_only" // #nosec G101 -- gin context key name, not a credential
|
||||||
|
apiTokenAuthSchemePrefix = "Bearer "
|
||||||
|
)
|
||||||
|
|
||||||
|
// listAPITokens 列出当前用户的所有 PAT(脱敏,不含 token 明文)。
|
||||||
|
// @Summary List API tokens
|
||||||
|
// @Tags auth required
|
||||||
|
// @Produce json
|
||||||
|
// @Success 200 {object} model.CommonResponse[[]model.APITokenView]
|
||||||
|
// @Router /api-tokens [get]
|
||||||
|
func listAPITokens(c *gin.Context) ([]model.APITokenView, error) {
|
||||||
|
uid := getUid(c)
|
||||||
|
var rows []model.APIToken
|
||||||
|
if err := singleton.DB.Where("user_id = ?", uid).Order("id DESC").Find(&rows).Error; err != nil {
|
||||||
|
return nil, newGormError("%v", err)
|
||||||
|
}
|
||||||
|
out := make([]model.APITokenView, 0, len(rows))
|
||||||
|
for i := range rows {
|
||||||
|
out = append(out, rows[i].ToView())
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// createAPIToken 创建一个 PAT。明文 token 仅在响应中返回一次。
|
||||||
|
// @Summary Create API token
|
||||||
|
// @Tags auth required
|
||||||
|
// @Accept json
|
||||||
|
// @Param body body model.APITokenCreateRequest true "request"
|
||||||
|
// @Produce json
|
||||||
|
// @Success 200 {object} model.CommonResponse[model.APITokenCreateResponse]
|
||||||
|
// @Router /api-tokens [post]
|
||||||
|
func createAPIToken(c *gin.Context) (*model.APITokenCreateResponse, error) {
|
||||||
|
var req model.APITokenCreateRequest
|
||||||
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
req.Name = strings.TrimSpace(req.Name)
|
||||||
|
if req.Name == "" {
|
||||||
|
return nil, errors.New("name required")
|
||||||
|
}
|
||||||
|
if len(req.Name) > 128 {
|
||||||
|
return nil, errors.New("name too long (max 128 chars)")
|
||||||
|
}
|
||||||
|
if req.ExpiresInDays < 0 {
|
||||||
|
return nil, errors.New("expires_in_days must be >= 0")
|
||||||
|
}
|
||||||
|
if req.ExpiresInDays > 3650 {
|
||||||
|
return nil, errors.New("expires_in_days too large (max 3650, i.e. 10 years)")
|
||||||
|
}
|
||||||
|
if len(req.Scopes) > 32 {
|
||||||
|
return nil, errors.New("too many scopes (max 32)")
|
||||||
|
}
|
||||||
|
if len(req.ServerIDs) > 1000 {
|
||||||
|
return nil, errors.New("too many server_ids (max 1000)")
|
||||||
|
}
|
||||||
|
|
||||||
|
allowed := append(append([]string{}, model.AllScopes...), model.AdminOnlyScopes...)
|
||||||
|
seen := make(map[string]struct{}, len(req.Scopes))
|
||||||
|
cleaned := make([]string, 0, len(req.Scopes))
|
||||||
|
for _, s := range req.Scopes {
|
||||||
|
s = strings.TrimSpace(s)
|
||||||
|
if s == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
normalized, ok := model.NormalizeIncomingScope(s)
|
||||||
|
if !ok {
|
||||||
|
return nil, errors.New("unknown scope: " + s)
|
||||||
|
}
|
||||||
|
if !slices.Contains(allowed, normalized) {
|
||||||
|
return nil, errors.New("unknown scope: " + s)
|
||||||
|
}
|
||||||
|
if _, dup := seen[normalized]; dup {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen[normalized] = struct{}{}
|
||||||
|
cleaned = append(cleaned, normalized)
|
||||||
|
}
|
||||||
|
if len(cleaned) == 0 {
|
||||||
|
return nil, errors.New("at least one scope required")
|
||||||
|
}
|
||||||
|
|
||||||
|
if !callerIsAdmin(c) {
|
||||||
|
for _, s := range cleaned {
|
||||||
|
if slices.Contains(model.AdminOnlyScopes, s) {
|
||||||
|
return nil, errors.New("only admin can issue scope: " + s)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(req.ServerIDs) > 0 {
|
||||||
|
seenSrv := make(map[uint64]struct{}, len(req.ServerIDs))
|
||||||
|
deduped := make([]uint64, 0, len(req.ServerIDs))
|
||||||
|
for _, sid := range req.ServerIDs {
|
||||||
|
if sid == 0 {
|
||||||
|
return nil, errors.New("server_id 0 is invalid")
|
||||||
|
}
|
||||||
|
if _, dup := seenSrv[sid]; dup {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seenSrv[sid] = struct{}{}
|
||||||
|
deduped = append(deduped, sid)
|
||||||
|
server, _ := singleton.ServerShared.Get(sid)
|
||||||
|
if server == nil {
|
||||||
|
return nil, errors.New("server not found")
|
||||||
|
}
|
||||||
|
if !callerIsAdmin(c) && !server.HasPermission(c) {
|
||||||
|
return nil, errors.New("permission denied on server")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
req.ServerIDs = deduped
|
||||||
|
}
|
||||||
|
|
||||||
|
secret, err := utils.GenerateRandomString(apiTokenSecretLength)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
plaintext := model.APITokenPrefix + secret
|
||||||
|
|
||||||
|
tok := model.APIToken{
|
||||||
|
UserID: getUid(c),
|
||||||
|
Name: req.Name,
|
||||||
|
TokenHash: model.HashAPIToken(plaintext),
|
||||||
|
}
|
||||||
|
tok.SetScopes(cleaned)
|
||||||
|
if len(req.ServerIDs) > 0 {
|
||||||
|
tok.SetServerIDs(req.ServerIDs)
|
||||||
|
}
|
||||||
|
if req.ExpiresInDays > 0 {
|
||||||
|
exp := time.Now().Add(time.Duration(req.ExpiresInDays) * 24 * time.Hour)
|
||||||
|
tok.ExpiresAt = &exp
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := singleton.DB.Create(&tok).Error; err != nil {
|
||||||
|
return nil, newGormError("%v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &model.APITokenCreateResponse{
|
||||||
|
ID: tok.ID,
|
||||||
|
Name: tok.Name,
|
||||||
|
Token: plaintext,
|
||||||
|
Scopes: tok.Scopes(),
|
||||||
|
ServerIDs: tok.ServerIDs(),
|
||||||
|
ExpiresAt: tok.ExpiresAt,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// deleteAPIToken 吊销一个 PAT。
|
||||||
|
// @Summary Revoke API token
|
||||||
|
// @Tags auth required
|
||||||
|
// @Param id path uint true "token id"
|
||||||
|
// @Produce json
|
||||||
|
// @Success 200 {object} model.CommonResponse[any]
|
||||||
|
// @Router /api-tokens/{id} [delete]
|
||||||
|
func deleteAPIToken(c *gin.Context) (any, error) {
|
||||||
|
idStr := c.Param("id")
|
||||||
|
id, err := strconv.ParseUint(idStr, 10, 64)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
q := singleton.DB.Where("id = ?", id)
|
||||||
|
if !callerIsAdmin(c) {
|
||||||
|
q = q.Where("user_id = ?", getUid(c))
|
||||||
|
}
|
||||||
|
res := q.Delete(&model.APIToken{})
|
||||||
|
if res.Error != nil {
|
||||||
|
return nil, newGormError("%v", res.Error)
|
||||||
|
}
|
||||||
|
if res.RowsAffected == 0 {
|
||||||
|
return nil, errors.New("not found")
|
||||||
|
}
|
||||||
|
// Fan out the revocation to any active long-lived connection that
|
||||||
|
// carries this PAT — ws/server, ws/transfer, terminal, FM. Without
|
||||||
|
// this hook a deleted PAT keeps streaming until the underlying
|
||||||
|
// connection naturally drops.
|
||||||
|
patConnectionRegistryShared.revokeToken(id)
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// apiTokenAuthMiddleware 解析 `Authorization: Bearer nzp_xxx`,
|
||||||
|
// 命中后把 *model.User 挂到 ctx 上,使下游一切 Server.HasPermission/getUid 复用 JWT 路径。
|
||||||
|
//
|
||||||
|
// 不命中(无 Authorization 头或前缀不是 nzp_):放行下一个中间件(例如 JWT)。
|
||||||
|
// 命中但 token 无效:直接 401 并 abort,不再走到 JWT。
|
||||||
|
func apiTokenAuthMiddleware() gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
readOnly := apiTokenAuthReadOnly(c)
|
||||||
|
raw := strings.TrimSpace(c.GetHeader("Authorization"))
|
||||||
|
if raw == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !strings.HasPrefix(raw, apiTokenAuthSchemePrefix) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
plaintext := strings.TrimSpace(strings.TrimPrefix(raw, apiTokenAuthSchemePrefix))
|
||||||
|
if !strings.HasPrefix(plaintext, model.APITokenPrefix) {
|
||||||
|
// 既然有 Bearer 但不是 PAT 前缀,交给后续 JWT 中间件处理
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
realIP := c.GetString(model.CtxKeyRealIPStr)
|
||||||
|
|
||||||
|
var tok model.APIToken
|
||||||
|
err := singleton.DB.Where("token_hash = ?", model.HashAPIToken(plaintext)).First(&tok).Error
|
||||||
|
if err != nil {
|
||||||
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||||
|
if !readOnly {
|
||||||
|
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken)
|
||||||
|
}
|
||||||
|
abortAPITokenUnauthorized(c, "invalid api token")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
abortAPITokenUnauthorized(c, "api token lookup failed")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
now := time.Now()
|
||||||
|
if tok.IsExpired(now) {
|
||||||
|
if !readOnly {
|
||||||
|
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken)
|
||||||
|
}
|
||||||
|
abortAPITokenUnauthorized(c, "api token expired")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
var user model.User
|
||||||
|
if err := singleton.DB.First(&user, tok.UserID).Error; err != nil {
|
||||||
|
if !readOnly {
|
||||||
|
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken)
|
||||||
|
}
|
||||||
|
abortAPITokenUnauthorized(c, "owner of api token not found")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if !readOnly {
|
||||||
|
model.UnblockIP(singleton.DB, realIP, model.BlockIDToken)
|
||||||
|
}
|
||||||
|
|
||||||
|
c.Set(model.CtxKeyAuthorizedUser, &user)
|
||||||
|
c.Set(apiTokenCtxKey, &tok)
|
||||||
|
c.Set(model.CtxKeyAPIToken, &tok)
|
||||||
|
|
||||||
|
// last_used 同步更新:开销极低(一行 UPDATE),异步路径在
|
||||||
|
// 多连接 sqlite 测试场景下会和测试 teardown 形成竞态,并把
|
||||||
|
// `last_used_*` 写丢到不可见的 :memory: 实例。生产路径上等价。
|
||||||
|
if !readOnly && apiTokenLastUsedOnce(c) {
|
||||||
|
c.Set(apiTokenLastUsedCtxKey, true)
|
||||||
|
_ = singleton.DB.Model(&model.APIToken{}).
|
||||||
|
Where("id = ?", tok.ID).
|
||||||
|
Updates(map[string]any{
|
||||||
|
"last_used_at": now,
|
||||||
|
"last_used_ip": c.GetString(model.CtxKeyRealIPStr),
|
||||||
|
}).Error
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func abortAPITokenUnauthorized(c *gin.Context, reason string) {
|
||||||
|
c.AbortWithStatusJSON(http.StatusUnauthorized, model.CommonResponse[any]{
|
||||||
|
Success: false,
|
||||||
|
Error: "ApiErrorUnauthorized: " + reason,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// APITokenFromContext 取当前请求关联的 PAT,未命中返回 nil。
|
||||||
|
// MCP tool 中间件用它做 scope 校验(闸 2)。
|
||||||
|
func APITokenFromContext(c *gin.Context) *model.APIToken {
|
||||||
|
v, ok := c.Get(apiTokenCtxKey)
|
||||||
|
if !ok {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
t, _ := v.(*model.APIToken)
|
||||||
|
return t
|
||||||
|
}
|
||||||
@@ -0,0 +1,78 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
// createAPIToken 是「旧 mcp:* → 新 nezha:*」唯一的归一化入口:
|
||||||
|
// - mcp:fs:read / mcp:server:read 归一化为 nezha:server:read;
|
||||||
|
// - mcp:server:exec 归一化为 nezha:server:exec;
|
||||||
|
// - mcp:fs:write / mcp:fs:delete / mcp:* 不再可签发——它们历史上覆盖范围
|
||||||
|
// 比 nezha:server:write/delete 窄(只跑 MCP fs 工具),静默映射会扩权。
|
||||||
|
//
|
||||||
|
// 这样老调用方传旧 scope 还能创建只读 PAT,但拿不到 write/delete 提权。
|
||||||
|
|
||||||
|
func TestCreateAPIToken_RewritesLegacyReadScopeToNezhaRead(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "legacy-reader",
|
||||||
|
Scopes: []string{"mcp:fs:read"},
|
||||||
|
})
|
||||||
|
res, err := createAPIToken(c)
|
||||||
|
require.NoError(t, err, "legacy mcp:fs:read must be accepted at create time and rewritten")
|
||||||
|
require.Equal(t, []string{model.ScopeServerRead}, res.Scopes,
|
||||||
|
"create response must reflect the new unified scope name, not the legacy alias")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_RewritesLegacyExecScopeToNezhaExec(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "legacy-exec",
|
||||||
|
Scopes: []string{"mcp:server:exec"},
|
||||||
|
})
|
||||||
|
res, err := createAPIToken(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, []string{model.ScopeServerExec}, res.Scopes)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_RejectsLegacyMCPWriteScope(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "legacy-writer",
|
||||||
|
Scopes: []string{"mcp:fs:write"},
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err,
|
||||||
|
"mcp:fs:write must be rejected: silently mapping to nezha:server:write would expand the original "+
|
||||||
|
"MCP-only write capability to every REST server mutation route")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_RejectsLegacyMCPDeleteScope(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "legacy-deleter",
|
||||||
|
Scopes: []string{"mcp:fs:delete"},
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_RejectsLegacyMCPWildcardScope(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(1, model.RoleAdmin)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "legacy-admin",
|
||||||
|
Scopes: []string{"mcp:*"},
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err,
|
||||||
|
"mcp:* must be rejected even for admin: the new unified namespace is nezha:* / nezha:admin:*")
|
||||||
|
}
|
||||||
@@ -0,0 +1,71 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
func setupOptionalAuthRouter(t *testing.T, plainToken string) *httptest.Server {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
|
||||||
|
jwtMw := func(c *gin.Context) {
|
||||||
|
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "no jwt"})
|
||||||
|
}
|
||||||
|
patMw := apiTokenAuthMiddleware()
|
||||||
|
authMw := jwtOrPATAuthMiddleware(patMw, jwtMw)
|
||||||
|
|
||||||
|
stub := func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) }
|
||||||
|
optionalAuth := r.Group("/api/v1", authMw)
|
||||||
|
optionalAuth.GET("/server-group", restScopeMiddleware(model.ScopeInventoryRead), stub)
|
||||||
|
optionalAuth.GET("/service", restScopeMiddleware(model.ScopeServiceRead), stub)
|
||||||
|
optionalAuth.GET("/server/:id/metrics", restScopeMiddleware(model.ScopeServerRead), stub)
|
||||||
|
|
||||||
|
ts := httptest.NewServer(r)
|
||||||
|
t.Cleanup(ts.Close)
|
||||||
|
_ = plainToken
|
||||||
|
return ts
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOptionalAuth_PATWithoutScopeIsDenied(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
_, plain := mkToken(t, uid, []string{model.ScopeNotificationRead}, nil)
|
||||||
|
ts := setupOptionalAuthRouter(t, plain)
|
||||||
|
|
||||||
|
for _, path := range []string{
|
||||||
|
"/api/v1/server-group",
|
||||||
|
"/api/v1/service",
|
||||||
|
"/api/v1/server/7/metrics",
|
||||||
|
} {
|
||||||
|
resp := doReq(t, ts, "GET", path, plain)
|
||||||
|
resp.Body.Close()
|
||||||
|
require.Equal(t, http.StatusForbidden, resp.StatusCode, "PAT lacking required scope must be denied for %s", path)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOptionalAuth_PATWithMatchingScopeAllowed(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
_, plain := mkToken(t, uid, []string{model.ScopeInventoryRead, model.ScopeServerRead, model.ScopeServiceRead}, nil)
|
||||||
|
ts := setupOptionalAuthRouter(t, plain)
|
||||||
|
|
||||||
|
for _, path := range []string{
|
||||||
|
"/api/v1/server-group",
|
||||||
|
"/api/v1/service",
|
||||||
|
"/api/v1/server/7/metrics",
|
||||||
|
} {
|
||||||
|
resp := doReq(t, ts, "GET", path, plain)
|
||||||
|
resp.Body.Close()
|
||||||
|
require.Equal(t, http.StatusOK, resp.StatusCode, "PAT with matching scope must pass for %s", path)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,19 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import "github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
func suppressAPITokenAuthWrites(c *gin.Context) { c.Set(apiTokenReadOnlyCtxKey, true) }
|
||||||
|
|
||||||
|
func apiTokenAuthReadOnly(c *gin.Context) bool {
|
||||||
|
value, ok := c.Get(apiTokenReadOnlyCtxKey)
|
||||||
|
return ok && value == true
|
||||||
|
}
|
||||||
|
|
||||||
|
func apiTokenLastUsedOnce(c *gin.Context) bool {
|
||||||
|
value, seen := c.Get(apiTokenLastUsedCtxKey)
|
||||||
|
if seen && value == true {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
c.Set(apiTokenLastUsedCtxKey, true)
|
||||||
|
return true
|
||||||
|
}
|
||||||
@@ -0,0 +1,148 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
// revokeTombstoneTTL bounds how long a revoked token id is remembered to
|
||||||
|
// close the revoke->register race. The race window is a single request's
|
||||||
|
// auth-to-register gap (sub-second); minutes of slack is ample. Without a
|
||||||
|
// TTL the tombstone set grows unbounded over the process lifetime as PATs
|
||||||
|
// are created and deleted.
|
||||||
|
const revokeTombstoneTTL = 10 * time.Minute
|
||||||
|
|
||||||
|
// patConnectionRegistry tracks active long-lived connections (terminal,
|
||||||
|
// FM, ws/server, ws/transfer, etc.) per PAT id so that deleteAPIToken can
|
||||||
|
// cancel them immediately on revocation. Without this, a deleted PAT
|
||||||
|
// keeps streaming until the underlying connection naturally drops.
|
||||||
|
//
|
||||||
|
// The registry deliberately holds no goroutines — it only stores cancel
|
||||||
|
// hooks the connection setup already owns. Handlers register on entry
|
||||||
|
// and deregister on exit; revokeToken walks the per-token slice and
|
||||||
|
// invokes every hook under the lock.
|
||||||
|
type patConnectionRegistry struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
byToken map[uint64]map[uint64]func()
|
||||||
|
// revoked is a tombstone set closing the revoke->register race: a
|
||||||
|
// connection can pass apiTokenAuthMiddleware (token cached in ctx) and
|
||||||
|
// only register its cancel hook AFTER deleteAPIToken already walked the
|
||||||
|
// registry. Without the tombstone that late registration would survive
|
||||||
|
// revocation. register consults revoked under the same lock and cancels
|
||||||
|
// immediately when the id is already gone.
|
||||||
|
revoked map[uint64]time.Time
|
||||||
|
nextID uint64
|
||||||
|
}
|
||||||
|
|
||||||
|
func newPATConnectionRegistry() *patConnectionRegistry {
|
||||||
|
return &patConnectionRegistry{
|
||||||
|
byToken: make(map[uint64]map[uint64]func()),
|
||||||
|
revoked: make(map[uint64]time.Time),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// pruneRevokedLocked drops tombstones older than revokeTombstoneTTL. Caller
|
||||||
|
// must hold r.mu. Bounds the tombstone set to recently-revoked ids.
|
||||||
|
func (r *patConnectionRegistry) pruneRevokedLocked(now time.Time) {
|
||||||
|
for id, at := range r.revoked {
|
||||||
|
if now.Sub(at) > revokeTombstoneTTL {
|
||||||
|
delete(r.revoked, id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// register stores cancel under tokenID and returns a deregister hook the
|
||||||
|
// caller MUST invoke when the connection ends. Returning a closure
|
||||||
|
// (rather than exposing an id) prevents callers from forgetting to clean
|
||||||
|
// up and avoids leaking entries past connection lifetime.
|
||||||
|
//
|
||||||
|
// If tokenID was already revoked, register does NOT store the hook; it
|
||||||
|
// cancels immediately and returns a no-op deregister, so a connection that
|
||||||
|
// raced past revocation is torn down at once.
|
||||||
|
func (r *patConnectionRegistry) register(tokenID uint64, cancel func()) func() {
|
||||||
|
r.mu.Lock()
|
||||||
|
now := time.Now()
|
||||||
|
r.pruneRevokedLocked(now)
|
||||||
|
if at, dead := r.revoked[tokenID]; dead && now.Sub(at) <= revokeTombstoneTTL {
|
||||||
|
r.mu.Unlock()
|
||||||
|
cancel()
|
||||||
|
return func() {}
|
||||||
|
}
|
||||||
|
r.nextID++
|
||||||
|
id := r.nextID
|
||||||
|
conns, ok := r.byToken[tokenID]
|
||||||
|
if !ok {
|
||||||
|
conns = make(map[uint64]func())
|
||||||
|
r.byToken[tokenID] = conns
|
||||||
|
}
|
||||||
|
conns[id] = cancel
|
||||||
|
r.mu.Unlock()
|
||||||
|
|
||||||
|
var once sync.Once
|
||||||
|
return func() {
|
||||||
|
once.Do(func() {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
if m, ok := r.byToken[tokenID]; ok {
|
||||||
|
delete(m, id)
|
||||||
|
if len(m) == 0 {
|
||||||
|
delete(r.byToken, tokenID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// revokeToken cancels every active connection registered under tokenID,
|
||||||
|
// clears the entry, and records a tombstone so any connection still racing
|
||||||
|
// toward register is cancelled on arrival. Safe to call on an unknown id.
|
||||||
|
func (r *patConnectionRegistry) revokeToken(tokenID uint64) {
|
||||||
|
r.mu.Lock()
|
||||||
|
conns := r.byToken[tokenID]
|
||||||
|
delete(r.byToken, tokenID)
|
||||||
|
now := time.Now()
|
||||||
|
r.pruneRevokedLocked(now)
|
||||||
|
r.revoked[tokenID] = now
|
||||||
|
r.mu.Unlock()
|
||||||
|
|
||||||
|
for _, cancel := range conns {
|
||||||
|
cancel()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// countForToken returns the number of active connections registered
|
||||||
|
// under tokenID. Intended for tests + future SIEM exposure; callers MUST
|
||||||
|
// NOT use it for policy decisions because the count can change the
|
||||||
|
// instant the lock is released.
|
||||||
|
func (r *patConnectionRegistry) countForToken(tokenID uint64) int {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
return len(r.byToken[tokenID])
|
||||||
|
}
|
||||||
|
|
||||||
|
var patConnectionRegistryShared = newPATConnectionRegistry()
|
||||||
|
|
||||||
|
// registerPATConnection wires the request-bound PAT (if any) into the
|
||||||
|
// process-wide revocation registry. Returns a deregister hook the
|
||||||
|
// handler MUST defer. For JWT-authenticated requests the hook is a
|
||||||
|
// no-op so call sites stay portable.
|
||||||
|
//
|
||||||
|
// Long-lived endpoints (terminal, FM, ws/server, ws/transfer) call
|
||||||
|
// this on entry and pass a cancel function that drops their websocket
|
||||||
|
// or relay loop. deleteAPIToken then revokes every active hook
|
||||||
|
// registered under the deleted token id.
|
||||||
|
func registerPATConnection(c interface {
|
||||||
|
Get(any) (any, bool)
|
||||||
|
}, cancel func()) func() {
|
||||||
|
v, ok := c.Get(apiTokenCtxKey)
|
||||||
|
if !ok {
|
||||||
|
return func() {}
|
||||||
|
}
|
||||||
|
tok, ok := v.(*model.APIToken)
|
||||||
|
if !ok || tok == nil {
|
||||||
|
return func() {}
|
||||||
|
}
|
||||||
|
return patConnectionRegistryShared.register(tok.ID, cancel)
|
||||||
|
}
|
||||||
@@ -0,0 +1,46 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 撤销发生在 register 之前时,迟到的连接必须被立即取消,而不是存活下来。
|
||||||
|
func TestRegisterAfterRevokeCancelsImmediately(t *testing.T) {
|
||||||
|
r := newPATConnectionRegistry()
|
||||||
|
r.revokeToken(42)
|
||||||
|
|
||||||
|
var cancelled atomic.Bool
|
||||||
|
dereg := r.register(42, func() { cancelled.Store(true) })
|
||||||
|
|
||||||
|
require.True(t, cancelled.Load(), "late registration on a revoked token must cancel at once")
|
||||||
|
require.Equal(t, 0, r.countForToken(42), "revoked token must not retain connections")
|
||||||
|
dereg() // must be a safe no-op
|
||||||
|
}
|
||||||
|
|
||||||
|
// 并发 revoke/register 下不得有连接逃过撤销。
|
||||||
|
func TestRevokeRegisterNoSurvivor(t *testing.T) {
|
||||||
|
for iter := 0; iter < 200; iter++ {
|
||||||
|
r := newPATConnectionRegistry()
|
||||||
|
var cancelled atomic.Bool
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(2)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
r.register(7, func() { cancelled.Store(true) })
|
||||||
|
}()
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
r.revokeToken(7)
|
||||||
|
}()
|
||||||
|
wg.Wait()
|
||||||
|
|
||||||
|
// 无论谁先跑:要么 register 先(被 revoke 取消),要么 revoke 先
|
||||||
|
// (register 在 tombstone 上立即取消)。两种顺序都不能留下活连接。
|
||||||
|
require.True(t, cancelled.Load(), "iter %d: connection survived revocation", iter)
|
||||||
|
require.Equal(t, 0, r.countForToken(7), "iter %d: registry must be empty after revoke", iter)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,99 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// M7 regression: long-lived PAT-authenticated handlers (ws/server,
|
||||||
|
// ws/transfer, terminal, FM) must register a cancel hook so that
|
||||||
|
// deleteAPIToken can close active connections immediately. Without this,
|
||||||
|
// a revoked PAT keeps streaming until the connection naturally drops.
|
||||||
|
func TestPATConnectionRegistry_CancelsOnRevoke(t *testing.T) {
|
||||||
|
registry := newPATConnectionRegistry()
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
deregister := registry.register(42, cancel)
|
||||||
|
defer deregister()
|
||||||
|
|
||||||
|
registry.revokeToken(42)
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("revokeToken must cancel the registered context within 1s")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPATConnectionRegistry_DoesNotCancelOtherTokens(t *testing.T) {
|
||||||
|
registry := newPATConnectionRegistry()
|
||||||
|
ctxA, cancelA := context.WithCancel(context.Background())
|
||||||
|
ctxB, cancelB := context.WithCancel(context.Background())
|
||||||
|
defer cancelA()
|
||||||
|
defer cancelB()
|
||||||
|
|
||||||
|
deregisterA := registry.register(1, cancelA)
|
||||||
|
deregisterB := registry.register(2, cancelB)
|
||||||
|
defer deregisterA()
|
||||||
|
defer deregisterB()
|
||||||
|
|
||||||
|
registry.revokeToken(1)
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ctxA.Done():
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("token 1's connection must be cancelled")
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ctxB.Done():
|
||||||
|
t.Fatal("token 2's connection must NOT be cancelled (separate token)")
|
||||||
|
case <-time.After(50 * time.Millisecond):
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPATConnectionRegistry_DeregisterClearsEntry(t *testing.T) {
|
||||||
|
registry := newPATConnectionRegistry()
|
||||||
|
_, cancel := context.WithCancel(context.Background())
|
||||||
|
deregister := registry.register(1, cancel)
|
||||||
|
|
||||||
|
deregister()
|
||||||
|
|
||||||
|
if got := registry.countForToken(1); got != 0 {
|
||||||
|
t.Fatalf("after deregister, count for token 1 must be 0, got %d", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPATConnectionRegistry_MultipleConnsPerToken(t *testing.T) {
|
||||||
|
registry := newPATConnectionRegistry()
|
||||||
|
ctx1, c1 := context.WithCancel(context.Background())
|
||||||
|
ctx2, c2 := context.WithCancel(context.Background())
|
||||||
|
defer c1()
|
||||||
|
defer c2()
|
||||||
|
|
||||||
|
d1 := registry.register(7, c1)
|
||||||
|
d2 := registry.register(7, c2)
|
||||||
|
defer d1()
|
||||||
|
defer d2()
|
||||||
|
|
||||||
|
if got := registry.countForToken(7); got != 2 {
|
||||||
|
t.Fatalf("expected 2 connections for token 7, got %d", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
registry.revokeToken(7)
|
||||||
|
|
||||||
|
for _, ctx := range []context.Context{ctx1, ctx2} {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("all connections for the revoked token must be cancelled")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPATConnectionRegistry_RevokeUnknownTokenIsNoOp(t *testing.T) {
|
||||||
|
registry := newPATConnectionRegistry()
|
||||||
|
registry.revokeToken(999) // must not panic
|
||||||
|
}
|
||||||
@@ -0,0 +1,135 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
// jwtOrPATAuthMiddleware 把 PAT 与 JWT 两条鉴权链组合到 /api/v1/* 入口。
|
||||||
|
//
|
||||||
|
// 处理顺序:
|
||||||
|
// 1. apiTokenAuthMiddleware:识别 `Authorization: Bearer nzp_*`。命中(合法 PAT)
|
||||||
|
// 把 user 挂到 ctx;非法 PAT 直接 abort 401。
|
||||||
|
// 2. 如果 PAT 已挂 user → 跳过 JWT。
|
||||||
|
// 3. 否则 → JWT 中间件接管,按现有 cookie / Bearer / query token 逻辑鉴权。
|
||||||
|
//
|
||||||
|
// 存量 JWT 客户端零感知;新 PAT 客户端可直接调 REST,但每个端点的 scope
|
||||||
|
// 仍由 restScopeMiddleware 控制。
|
||||||
|
func jwtOrPATAuthMiddleware(patMw, jwtMw gin.HandlerFunc) gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
patMw(c)
|
||||||
|
if c.IsAborted() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if APITokenFromContext(c) != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
jwtMw(c)
|
||||||
|
if c.IsAborted() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// patOrFallbackAuthMiddleware 是 optional 路由(ForceAuth=false 时也能匿名访问)
|
||||||
|
// 的鉴权链:
|
||||||
|
// 1. apiTokenAuthMiddleware:识别 PAT,命中后挂 user;非法 PAT 401 abort。
|
||||||
|
// 2. 已挂 PAT → 跳过 JWT,restScopeMiddleware 会按 scope 收口。
|
||||||
|
// 3. 未带 PAT → 走 fallbackJwtMw,存在 JWT 则挂 user,没有就匿名继续。
|
||||||
|
//
|
||||||
|
// 这是修复 ForceAuth=false 时 optional 路由完全不解析 PAT 的关键:
|
||||||
|
// 之前直接用 fallbackAuthMw 会让 PAT 请求被当作 guest,scope 形同虚设。
|
||||||
|
func patOrFallbackAuthMiddleware(patMw, fallbackJwtMw gin.HandlerFunc) gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
patMw(c)
|
||||||
|
if c.IsAborted() {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if APITokenFromContext(c) != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
fallbackJwtMw(c)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// restScopeMiddleware 在 /api/v1/* 路由上 enforce PAT scope。
|
||||||
|
//
|
||||||
|
// 行为:
|
||||||
|
// - JWT 持有者(任何来源:cookie / Authorization Bearer 非 nzp_)→ 直接放行,
|
||||||
|
// 沿用 JWT 模型的完整权限。
|
||||||
|
// - PAT 持有者 → 必须命中给定 scope。命中后下游 handler 仍受 user 级权限
|
||||||
|
// 检查(adminHandler / Server.HasPermission),scope 只能收窄不能放大。
|
||||||
|
// - PAT 持有者遇到 scope=="" → 直接 403。空字符串作为 fail-closed 默认值,
|
||||||
|
// 防止接入新路由时忘填 scope 把 PAT 静默放行。
|
||||||
|
//
|
||||||
|
// 因此"自我管理"端点(/profile、/api-tokens、/refresh-token 等)必须显式挂
|
||||||
|
// restPATForbiddenMiddleware 来拒绝 PAT,而不是依赖空 scope 兜底。
|
||||||
|
func restScopeMiddleware(scope string) gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
tok := APITokenFromContext(c)
|
||||||
|
if tok == nil {
|
||||||
|
c.Next()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if scope == "" || !tok.HasScope(scope) {
|
||||||
|
c.AbortWithStatusJSON(http.StatusForbidden, model.CommonResponse[any]{
|
||||||
|
Success: false,
|
||||||
|
Error: "ApiErrorForbidden: api token lacks scope " + scope,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.Next()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// restScopeAllOf is the multi-scope variant of restScopeMiddleware. It
|
||||||
|
// gates on EVERY listed scope, used by routes whose semantics span more
|
||||||
|
// than one capability — file-manager sessions read, write AND delete
|
||||||
|
// files, so a PAT that only carries nezha:server:write must NOT be allowed
|
||||||
|
// to open one. JWT callers pass through unchanged.
|
||||||
|
func restScopeAllOf(scopes ...string) gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
tok := APITokenFromContext(c)
|
||||||
|
if tok == nil {
|
||||||
|
c.Next()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
for _, scope := range scopes {
|
||||||
|
if scope == "" || !tok.HasScope(scope) {
|
||||||
|
c.AbortWithStatusJSON(http.StatusForbidden, model.CommonResponse[any]{
|
||||||
|
Success: false,
|
||||||
|
Error: "ApiErrorForbidden: api token lacks scope " + scope,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
c.Next()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// serverConfigSensitiveScope 收紧 GET /server/config/:id 的 PAT scope 到
|
||||||
|
// ScopeServerWrite:返回体里包含 client_secret 等下发到 agent 的凭据,单纯
|
||||||
|
// nezha:server:read 不应足以读取。命名刻意带 Sensitive 而不是 Read,避免下
|
||||||
|
// 个维护者把它当成普通 read scope 还原成 ScopeServerRead 重新打开提权链。
|
||||||
|
func serverConfigSensitiveScope() string { return model.ScopeServerWrite }
|
||||||
|
|
||||||
|
// restPATForbiddenMiddleware 在「自我管理」类端点上显式拒绝 PAT。
|
||||||
|
//
|
||||||
|
// 这些端点(profile / api-tokens / oauth2 绑定 / refresh-token)一旦允许 PAT
|
||||||
|
// 自调,就可形成提权链(PAT → 创建更高权限 PAT → ...)。
|
||||||
|
// 显式 403 比静默放行更安全。
|
||||||
|
func restPATForbiddenMiddleware() gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
if APITokenFromContext(c) != nil {
|
||||||
|
c.AbortWithStatusJSON(http.StatusForbidden, model.CommonResponse[any]{
|
||||||
|
Success: false,
|
||||||
|
Error: "ApiErrorForbidden: this endpoint is not accessible by api token",
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.Next()
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
// restScopeMiddleware 的"空 scope"实际行为:PAT 调用方一律 403。
|
||||||
|
// 这条测试把注释与实现的契约对齐:
|
||||||
|
// 1. 实际行为:PAT + scope="" → 403。
|
||||||
|
// 2. 文档约束:源码注释必须明确说出"空 scope 对 PAT 仍被拒绝",
|
||||||
|
// 不能再保留"空字符串 = 放行"这种与实现相反的旧措辞。
|
||||||
|
func TestRestScopeMiddleware_EmptyScopeRejectsPAT(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.GET("/x", func(c *gin.Context) {
|
||||||
|
c.Set(apiTokenCtxKey, &model.APIToken{ID: 1})
|
||||||
|
c.Set(model.CtxKeyAPIToken, &model.APIToken{ID: 1})
|
||||||
|
c.Next()
|
||||||
|
}, restScopeMiddleware(""), func(c *gin.Context) {
|
||||||
|
c.Status(http.StatusOK)
|
||||||
|
})
|
||||||
|
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest("GET", "/x", nil)
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
if w.Code != http.StatusForbidden {
|
||||||
|
t.Fatalf("expected 403 when PAT hits restScopeMiddleware(\"\"); got %d body=%q", w.Code, w.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRestScopeMiddleware_DocReflectsEmptyScopeRejection(t *testing.T) {
|
||||||
|
wd, err := os.Getwd()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("getwd: %v", err)
|
||||||
|
}
|
||||||
|
b, err := os.ReadFile(filepath.Join(wd, "api_token_scope.go"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read: %v", err)
|
||||||
|
}
|
||||||
|
src := string(b)
|
||||||
|
idx := strings.Index(src, "func restScopeMiddleware(")
|
||||||
|
if idx < 0 {
|
||||||
|
t.Fatalf("restScopeMiddleware not found")
|
||||||
|
}
|
||||||
|
doc := src[:idx]
|
||||||
|
if strings.Contains(doc, "空字符串)= 放行") || strings.Contains(doc, "空字符串) = 放行") {
|
||||||
|
t.Fatalf("doc still claims empty scope means 放行; this contradicts the implementation which 403s PAT callers")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,223 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 在 /api/v1 风格的小 router 上重现 PAT + scope mw,验证 enforcement。
|
||||||
|
func setupRESTScopeServer(t *testing.T) (*httptest.Server, string, func()) {
|
||||||
|
t.Helper()
|
||||||
|
cleanupBase, uid := setupMCPTest(t)
|
||||||
|
|
||||||
|
_, plain := mkToken(t, uid, []string{model.ScopeInventoryRead}, nil)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
pat := apiTokenAuthMiddleware()
|
||||||
|
r.GET("/api/v1/server",
|
||||||
|
pat,
|
||||||
|
restScopeMiddleware(model.ScopeInventoryRead),
|
||||||
|
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
|
||||||
|
)
|
||||||
|
r.POST("/api/v1/server/config",
|
||||||
|
pat,
|
||||||
|
restScopeMiddleware(model.ScopeServerWrite),
|
||||||
|
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
|
||||||
|
)
|
||||||
|
r.GET("/api/v1/profile",
|
||||||
|
pat,
|
||||||
|
restPATForbiddenMiddleware(),
|
||||||
|
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
|
||||||
|
)
|
||||||
|
|
||||||
|
ts := httptest.NewServer(r)
|
||||||
|
return ts, plain, func() {
|
||||||
|
ts.Close()
|
||||||
|
cleanupBase()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func httpGetWithToken(t *testing.T, ts *httptest.Server, path, token string) (int, map[string]any) {
|
||||||
|
t.Helper()
|
||||||
|
req, _ := http.NewRequest("GET", ts.URL+path, nil)
|
||||||
|
if token != "" {
|
||||||
|
req.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
}
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
var out map[string]any
|
||||||
|
_ = json.Unmarshal(body, &out)
|
||||||
|
return resp.StatusCode, out
|
||||||
|
}
|
||||||
|
|
||||||
|
func httpPostWithToken(t *testing.T, ts *httptest.Server, path, token string) (int, map[string]any) {
|
||||||
|
t.Helper()
|
||||||
|
req, _ := http.NewRequest("POST", ts.URL+path, strings.NewReader("{}"))
|
||||||
|
if token != "" {
|
||||||
|
req.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
}
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
var out map[string]any
|
||||||
|
_ = json.Unmarshal(body, &out)
|
||||||
|
return resp.StatusCode, out
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRESTScope_PATWithReadCanGET(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupRESTScopeServer(t)
|
||||||
|
defer cleanup()
|
||||||
|
code, body := httpGetWithToken(t, ts, "/api/v1/server", tok)
|
||||||
|
require.Equal(t, 200, code)
|
||||||
|
require.True(t, body["ok"].(bool))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRESTScope_PATWithReadCannotWrite(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupRESTScopeServer(t)
|
||||||
|
defer cleanup()
|
||||||
|
code, body := httpPostWithToken(t, ts, "/api/v1/server/config", tok)
|
||||||
|
require.Equal(t, 403, code)
|
||||||
|
require.Contains(t, body["error"], "nezha:server:write")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRESTScope_NoTokenIsTransparentToScopeMW(t *testing.T) {
|
||||||
|
ts, _, cleanup := setupRESTScopeServer(t)
|
||||||
|
defer cleanup()
|
||||||
|
code, _ := httpGetWithToken(t, ts, "/api/v1/server", "")
|
||||||
|
require.Equal(t, 200, code, "scope mw is PAT-only enforcement; JWT flow is gated by jwtOrPATAuthMiddleware before this layer. In this minimal router there is no JWT mw, so no token = handler runs (security is enforced upstream)")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRESTScope_PATForbiddenOnSelfManagement(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupRESTScopeServer(t)
|
||||||
|
defer cleanup()
|
||||||
|
code, body := httpGetWithToken(t, ts, "/api/v1/profile", tok)
|
||||||
|
require.Equal(t, 403, code)
|
||||||
|
require.Contains(t, body["error"], "not accessible by api token")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRESTScope_JWTUserSkipsScope(t *testing.T) {
|
||||||
|
cleanupBase, uid := setupMCPTest(t)
|
||||||
|
defer cleanupBase()
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
pat := apiTokenAuthMiddleware()
|
||||||
|
r.GET("/api/v1/server",
|
||||||
|
pat,
|
||||||
|
func(c *gin.Context) {
|
||||||
|
c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: uid}, Role: model.RoleMember})
|
||||||
|
c.Next()
|
||||||
|
},
|
||||||
|
restScopeMiddleware(model.ScopeInventoryRead),
|
||||||
|
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
|
||||||
|
)
|
||||||
|
ts := httptest.NewServer(r)
|
||||||
|
defer ts.Close()
|
||||||
|
|
||||||
|
code, body := httpGetWithToken(t, ts, "/api/v1/server", "")
|
||||||
|
require.Equal(t, 200, code, "JWT-attached request (no PAT in ctx) must bypass scope check")
|
||||||
|
require.True(t, body["ok"].(bool))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRESTScope_NezhaAllUnlocksEverything(t *testing.T) {
|
||||||
|
cleanupBase, uid := setupMCPTest(t)
|
||||||
|
defer cleanupBase()
|
||||||
|
_, plain := mkToken(t, uid, []string{model.ScopeNezhaAll}, nil)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
pat := apiTokenAuthMiddleware()
|
||||||
|
r.POST("/api/v1/server/config",
|
||||||
|
pat,
|
||||||
|
restScopeMiddleware(model.ScopeServerWrite),
|
||||||
|
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
|
||||||
|
)
|
||||||
|
r.POST("/api/v1/batch-delete/server",
|
||||||
|
pat,
|
||||||
|
restScopeMiddleware(model.ScopeInventoryDelete),
|
||||||
|
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
|
||||||
|
)
|
||||||
|
ts := httptest.NewServer(r)
|
||||||
|
defer ts.Close()
|
||||||
|
|
||||||
|
for _, path := range []string{"/api/v1/server/config", "/api/v1/batch-delete/server"} {
|
||||||
|
code, _ := httpPostWithToken(t, ts, path, plain)
|
||||||
|
require.Equalf(t, 200, code, "nezha:* must unlock %s", path)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- WAF brute force ---
|
||||||
|
|
||||||
|
func TestRESTScope_BadPATIncrementsWAFCounter(t *testing.T) {
|
||||||
|
cleanupBase, _ := setupMCPTest(t)
|
||||||
|
defer cleanupBase()
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.Use(func(c *gin.Context) {
|
||||||
|
c.Set(model.CtxKeyRealIPStr, "203.0.113.7")
|
||||||
|
c.Next()
|
||||||
|
})
|
||||||
|
r.GET("/api/v1/server",
|
||||||
|
apiTokenAuthMiddleware(),
|
||||||
|
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
|
||||||
|
)
|
||||||
|
ts := httptest.NewServer(r)
|
||||||
|
defer ts.Close()
|
||||||
|
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
code, _ := httpGetWithToken(t, ts, "/api/v1/server", "nzp_invalid_token_xxx")
|
||||||
|
require.Equal(t, 401, code)
|
||||||
|
}
|
||||||
|
|
||||||
|
var w model.WAF
|
||||||
|
require.NoError(t, singleton.DB.Where("block_identifier = ?", model.BlockIDToken).First(&w).Error)
|
||||||
|
require.GreaterOrEqual(t, w.Count, uint64(3))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRESTScope_GoodPATClearsWAFCounter(t *testing.T) {
|
||||||
|
cleanupBase, uid := setupMCPTest(t)
|
||||||
|
defer cleanupBase()
|
||||||
|
|
||||||
|
_, plain := mkToken(t, uid, []string{model.ScopeInventoryRead}, nil)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.Use(func(c *gin.Context) {
|
||||||
|
c.Set(model.CtxKeyRealIPStr, "198.51.100.5")
|
||||||
|
c.Next()
|
||||||
|
})
|
||||||
|
r.GET("/api/v1/server",
|
||||||
|
apiTokenAuthMiddleware(),
|
||||||
|
restScopeMiddleware(model.ScopeInventoryRead),
|
||||||
|
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
|
||||||
|
)
|
||||||
|
ts := httptest.NewServer(r)
|
||||||
|
defer ts.Close()
|
||||||
|
|
||||||
|
_, _ = httpGetWithToken(t, ts, "/api/v1/server", "nzp_invalid_token_xxx")
|
||||||
|
var w model.WAF
|
||||||
|
require.NoError(t, singleton.DB.Where("block_identifier = ?", model.BlockIDToken).First(&w).Error)
|
||||||
|
require.GreaterOrEqual(t, w.Count, uint64(1))
|
||||||
|
|
||||||
|
code, _ := httpGetWithToken(t, ts, "/api/v1/server", plain)
|
||||||
|
require.Equal(t, 200, code)
|
||||||
|
require.ErrorContains(t,
|
||||||
|
singleton.DB.Where("block_identifier = ?", model.BlockIDToken).First(&model.WAF{}).Error,
|
||||||
|
"record not found",
|
||||||
|
)
|
||||||
|
}
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestREST_ServerConfigRequiresWriteScope(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
_, plain := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.GET("/api/v1/server/config/:id",
|
||||||
|
apiTokenAuthMiddleware(),
|
||||||
|
restScopeMiddleware(serverConfigSensitiveScope()),
|
||||||
|
func(c *gin.Context) { c.JSON(http.StatusOK, gin.H{"ok": true}) },
|
||||||
|
)
|
||||||
|
ts := httptest.NewServer(r)
|
||||||
|
defer ts.Close()
|
||||||
|
|
||||||
|
resp := doReq(t, ts, "GET", "/api/v1/server/config/7", plain)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
require.Equal(t, http.StatusForbidden, resp.StatusCode,
|
||||||
|
"nezha:server:read must not be sufficient to read agent config (contains client_secret)")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestREST_ServerConfigGrantedByWriteScope(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
_, plain := mkToken(t, uid, []string{model.ScopeServerWrite}, nil)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.GET("/api/v1/server/config/:id",
|
||||||
|
apiTokenAuthMiddleware(),
|
||||||
|
restScopeMiddleware(serverConfigSensitiveScope()),
|
||||||
|
func(c *gin.Context) { c.JSON(http.StatusOK, gin.H{"ok": true}) },
|
||||||
|
)
|
||||||
|
ts := httptest.NewServer(r)
|
||||||
|
defer ts.Close()
|
||||||
|
|
||||||
|
resp := doReq(t, ts, "GET", "/api/v1/server/config/7", plain)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
require.Equal(t, http.StatusOK, resp.StatusCode)
|
||||||
|
}
|
||||||
@@ -0,0 +1,79 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func patRequestCtx(t *testing.T, tok *model.APIToken, uid uint64, method, path string, body any) (*gin.Context, *httptest.ResponseRecorder) {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
var rdr *bytes.Reader
|
||||||
|
if body != nil {
|
||||||
|
b, _ := json.Marshal(body)
|
||||||
|
rdr = bytes.NewReader(b)
|
||||||
|
} else {
|
||||||
|
rdr = bytes.NewReader(nil)
|
||||||
|
}
|
||||||
|
c.Request = httptest.NewRequest(method, path, rdr)
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: uid}, Role: model.RoleMember})
|
||||||
|
c.Set(apiTokenCtxKey, tok)
|
||||||
|
c.Set(model.CtxKeyAPIToken, tok)
|
||||||
|
return c, w
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestREST_PATServerWhitelistBlocksOtherServer(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, []uint64{99})
|
||||||
|
|
||||||
|
srv, _ := singleton.ServerShared.Get(7)
|
||||||
|
require.NotNil(t, srv)
|
||||||
|
require.Equal(t, uid, srv.GetUserID())
|
||||||
|
|
||||||
|
c, _ := patRequestCtx(t, tok, uid, "GET", "/api/v1/server/config/7", nil)
|
||||||
|
c.Params = gin.Params{{Key: "id", Value: "7"}}
|
||||||
|
|
||||||
|
_, err := getServerConfig(c)
|
||||||
|
require.Error(t, err, "PAT not in server whitelist must be rejected")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestREST_PATServerWhitelistBlocksSetConfig(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerWrite}, []uint64{99})
|
||||||
|
|
||||||
|
c, _ := patRequestCtx(t, tok, uid, "POST", "/api/v1/server/config", model.ServerConfigForm{
|
||||||
|
Servers: []uint64{7},
|
||||||
|
Config: "{}",
|
||||||
|
})
|
||||||
|
_, err := setServerConfig(c)
|
||||||
|
require.Error(t, err, "setServerConfig must reject non-whitelisted server")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestREST_PATServerWhitelistAllowsListedServer(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, []uint64{7})
|
||||||
|
|
||||||
|
c, _ := patRequestCtx(t, tok, uid, "GET", "/api/v1/server/config/7", nil)
|
||||||
|
c.Params = gin.Params{{Key: "id", Value: "7"}}
|
||||||
|
|
||||||
|
data, err := getServerConfig(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, "", data, "no agent stream connected so handler should return empty")
|
||||||
|
}
|
||||||
@@ -0,0 +1,615 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func setupAPITokenTest(t *testing.T) func() {
|
||||||
|
t.Helper()
|
||||||
|
originalDB := singleton.DB
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
sqlDB, err := db.DB()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, db.AutoMigrate(&model.User{}, &model.APIToken{}, &model.Server{}))
|
||||||
|
singleton.DB = db
|
||||||
|
return func() {
|
||||||
|
_ = sqlDB.Close()
|
||||||
|
singleton.DB = originalDB
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func ctxAsUser(uid uint64, role model.Role) *gin.Context {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||||
|
c.Request = httptest.NewRequest("POST", "/", nil)
|
||||||
|
c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: uid}, Role: role})
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
func bindJSON(c *gin.Context, body any) {
|
||||||
|
b, _ := json.Marshal(body)
|
||||||
|
c.Request = httptest.NewRequest("POST", "/", bytes.NewReader(b))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_MemberCanCreateExplicitScope(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "claude",
|
||||||
|
Scopes: []string{model.ScopeServerRead},
|
||||||
|
})
|
||||||
|
res, err := createAPIToken(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NotEmpty(t, res.Token)
|
||||||
|
require.True(t, strings.HasPrefix(res.Token, model.APITokenPrefix))
|
||||||
|
require.Greater(t, res.ID, uint64(0))
|
||||||
|
|
||||||
|
var stored model.APIToken
|
||||||
|
require.NoError(t, singleton.DB.First(&stored, res.ID).Error)
|
||||||
|
require.Equal(t, model.HashAPIToken(res.Token), stored.TokenHash)
|
||||||
|
require.Equal(t, uint64(10), stored.UserID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_MemberCannotIssueWildcard(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "x",
|
||||||
|
Scopes: []string{model.ScopeNezhaAll},
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Contains(t, err.Error(), "admin")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_AdminCanIssueWildcard(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(1, model.RoleAdmin)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "ops-script",
|
||||||
|
Scopes: []string{model.ScopeNezhaAll},
|
||||||
|
})
|
||||||
|
res, err := createAPIToken(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Contains(t, res.Scopes, model.ScopeNezhaAll)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_RejectsUnknownScope(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(1, model.RoleAdmin)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "x",
|
||||||
|
Scopes: []string{"mcp:hack:everything"},
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Contains(t, err.Error(), "unknown scope")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_RejectsEmptyScopes(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(1, model.RoleAdmin)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{Name: "x", Scopes: []string{}})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_RejectsTooManyServerIDs(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(1, model.RoleAdmin)
|
||||||
|
ids := make([]uint64, 1001)
|
||||||
|
for i := range ids {
|
||||||
|
ids[i] = uint64(i + 1)
|
||||||
|
}
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "x",
|
||||||
|
Scopes: []string{model.ScopeServerRead},
|
||||||
|
ServerIDs: ids,
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Contains(t, err.Error(), "too many server_ids")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_RejectsExpirationOutOfRange(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(1, model.RoleAdmin)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "x",
|
||||||
|
Scopes: []string{model.ScopeServerRead},
|
||||||
|
ExpiresInDays: -1,
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err)
|
||||||
|
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "x",
|
||||||
|
Scopes: []string{model.ScopeServerRead},
|
||||||
|
ExpiresInDays: 10000,
|
||||||
|
})
|
||||||
|
_, err = createAPIToken(c)
|
||||||
|
require.Error(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDeleteAPIToken_OnlyOwnerOrAdmin(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
tok := model.APIToken{UserID: 10, Name: "x", TokenHash: model.HashAPIToken("nzp_x")}
|
||||||
|
tok.SetScopes([]string{model.ScopeServerRead})
|
||||||
|
require.NoError(t, singleton.DB.Create(&tok).Error)
|
||||||
|
|
||||||
|
c := ctxAsUser(11, model.RoleMember)
|
||||||
|
c.Params = gin.Params{{Key: "id", Value: itoa(tok.ID)}}
|
||||||
|
_, err := deleteAPIToken(c)
|
||||||
|
require.Error(t, err, "other member must not delete")
|
||||||
|
|
||||||
|
c = ctxAsUser(10, model.RoleMember)
|
||||||
|
c.Params = gin.Params{{Key: "id", Value: itoa(tok.ID)}}
|
||||||
|
_, err = deleteAPIToken(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDeleteAPIToken_AdminCanDeleteAny(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
tok := model.APIToken{UserID: 10, Name: "x", TokenHash: model.HashAPIToken("nzp_y")}
|
||||||
|
tok.SetScopes([]string{model.ScopeServerRead})
|
||||||
|
require.NoError(t, singleton.DB.Create(&tok).Error)
|
||||||
|
|
||||||
|
c := ctxAsUser(1, model.RoleAdmin)
|
||||||
|
c.Params = gin.Params{{Key: "id", Value: itoa(tok.ID)}}
|
||||||
|
_, err := deleteAPIToken(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// installServerForAPIToken 在 ServerShared 里塞一台属于 ownerUID 的 server,
|
||||||
|
// 仅用于 PAT 创建路径的 server_ids 权限校验测试。
|
||||||
|
func installServerForAPIToken(t *testing.T, serverID, ownerUID uint64) func() {
|
||||||
|
t.Helper()
|
||||||
|
original := singleton.ServerShared
|
||||||
|
sc := singleton.NewEmptyServerClassForTest()
|
||||||
|
srv := &model.Server{}
|
||||||
|
srv.ID = serverID
|
||||||
|
srv.SetUserID(ownerUID)
|
||||||
|
sc.InsertForTest(srv)
|
||||||
|
singleton.ServerShared = sc
|
||||||
|
return func() { singleton.ServerShared = original }
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_MemberCannotIncludeForeignServerID(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
defer installServerForAPIToken(t, 42, 999)() // server 42 owned by user 999
|
||||||
|
|
||||||
|
c := ctxAsUser(10, model.RoleMember) // attacker is user 10
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "evil",
|
||||||
|
Scopes: []string{model.ScopeServerRead},
|
||||||
|
ServerIDs: []uint64{42},
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err, "member must not be able to bind foreign server_id into a PAT")
|
||||||
|
require.Contains(t, err.Error(), "permission denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_MemberCannotIncludeNonexistentServerID(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
cleanup := installServerForAPIToken(t, 1, 10)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "evil2",
|
||||||
|
Scopes: []string{model.ScopeServerRead},
|
||||||
|
ServerIDs: []uint64{9999}, // never-existed server
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Contains(t, err.Error(), "server not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_AdminCanIncludeAnyServerID(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
defer installServerForAPIToken(t, 77, 999)() // foreign server
|
||||||
|
|
||||||
|
c := ctxAsUser(1, model.RoleAdmin)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "ops-script",
|
||||||
|
Scopes: []string{model.ScopeServerRead},
|
||||||
|
ServerIDs: []uint64{77},
|
||||||
|
})
|
||||||
|
res, err := createAPIToken(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, []uint64{77}, res.ServerIDs)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_AdminCannotIncludeNonexistentServerID(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
defer installServerForAPIToken(t, 77, 999)() // only server 77 exists
|
||||||
|
|
||||||
|
c := ctxAsUser(1, model.RoleAdmin)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "ops-future-bind",
|
||||||
|
Scopes: []string{model.ScopeServerRead},
|
||||||
|
ServerIDs: []uint64{4242}, // never-existed server id
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err, "admin must not bind a PAT to a nonexistent server_id; a later-created server with that id would auto-inherit the grant")
|
||||||
|
require.Contains(t, err.Error(), "server not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_MemberOwnServerIDIsAccepted(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
defer installServerForAPIToken(t, 55, 10)() // user 10 owns server 55
|
||||||
|
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "self",
|
||||||
|
Scopes: []string{model.ScopeServerRead},
|
||||||
|
ServerIDs: []uint64{55},
|
||||||
|
})
|
||||||
|
res, err := createAPIToken(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, []uint64{55}, res.ServerIDs)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPITokenAuthMW_ExpiredTokenRejected(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
plain := "nzp_" + strings.Repeat("e", 32)
|
||||||
|
past := time.Now().Add(-time.Hour)
|
||||||
|
tok := model.APIToken{
|
||||||
|
UserID: 10,
|
||||||
|
Name: "expired",
|
||||||
|
TokenHash: model.HashAPIToken(plain),
|
||||||
|
ExpiresAt: &past,
|
||||||
|
}
|
||||||
|
tok.SetScopes([]string{model.ScopeServerRead})
|
||||||
|
require.NoError(t, singleton.DB.Create(&tok).Error)
|
||||||
|
require.NoError(t, singleton.DB.Create(&model.User{Common: model.Common{ID: 10}}).Error)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||||
|
c.Request.Header.Set("Authorization", "Bearer "+plain)
|
||||||
|
apiTokenAuthMiddleware()(c)
|
||||||
|
require.True(t, c.IsAborted(), "expired token must abort the request")
|
||||||
|
require.Equal(t, http.StatusUnauthorized, w.Code)
|
||||||
|
require.Contains(t, w.Body.String(), "expired")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPITokenAuthMW_OwnerDeletedRejected(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
plain := "nzp_" + strings.Repeat("o", 32)
|
||||||
|
tok := model.APIToken{
|
||||||
|
UserID: 999,
|
||||||
|
Name: "orphan",
|
||||||
|
TokenHash: model.HashAPIToken(plain),
|
||||||
|
}
|
||||||
|
tok.SetScopes([]string{model.ScopeServerRead})
|
||||||
|
require.NoError(t, singleton.DB.Create(&tok).Error)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||||
|
c.Request.Header.Set("Authorization", "Bearer "+plain)
|
||||||
|
apiTokenAuthMiddleware()(c)
|
||||||
|
require.True(t, c.IsAborted(), "owner-less PAT must abort")
|
||||||
|
require.Equal(t, http.StatusUnauthorized, w.Code)
|
||||||
|
require.Contains(t, w.Body.String(), "owner")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPITokenAuthMW_HappyPathSetsUserContext(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
plain := "nzp_" + strings.Repeat("g", 32)
|
||||||
|
tok := model.APIToken{
|
||||||
|
UserID: 10,
|
||||||
|
Name: "good",
|
||||||
|
TokenHash: model.HashAPIToken(plain),
|
||||||
|
}
|
||||||
|
tok.SetScopes([]string{model.ScopeServerRead})
|
||||||
|
require.NoError(t, singleton.DB.Create(&tok).Error)
|
||||||
|
require.NoError(t, singleton.DB.Create(&model.User{Common: model.Common{ID: 10}, Username: "alice"}).Error)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||||
|
c.Request.Header.Set("Authorization", "Bearer "+plain)
|
||||||
|
apiTokenAuthMiddleware()(c)
|
||||||
|
require.False(t, c.IsAborted())
|
||||||
|
|
||||||
|
user, ok := c.Get(model.CtxKeyAuthorizedUser)
|
||||||
|
require.True(t, ok)
|
||||||
|
require.Equal(t, uint64(10), user.(*model.User).ID)
|
||||||
|
|
||||||
|
got := APITokenFromContext(c)
|
||||||
|
require.NotNil(t, got)
|
||||||
|
require.Equal(t, "good", got.Name)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPITokenAuthMW_NonNZPBearerPassesThrough(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||||
|
c.Request.Header.Set("Authorization", "Bearer some-jwt-here")
|
||||||
|
apiTokenAuthMiddleware()(c)
|
||||||
|
require.False(t, c.IsAborted(), "non-nzp Bearer must pass through to JWT middleware")
|
||||||
|
require.Nil(t, APITokenFromContext(c))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPITokenAuthMW_EmptyNZPBodyRejected(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||||
|
c.Request.Header.Set("Authorization", "Bearer nzp_")
|
||||||
|
apiTokenAuthMiddleware()(c)
|
||||||
|
require.True(t, c.IsAborted(),
|
||||||
|
"Bearer with empty nzp_ body must be rejected (would otherwise hash an empty string and look it up)")
|
||||||
|
require.Equal(t, http.StatusUnauthorized, w.Code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPITokenAuthMW_RevokedTokenRejected(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
plain := "nzp_" + strings.Repeat("r", 32)
|
||||||
|
tok := model.APIToken{
|
||||||
|
UserID: 10,
|
||||||
|
Name: "to-revoke",
|
||||||
|
TokenHash: model.HashAPIToken(plain),
|
||||||
|
}
|
||||||
|
tok.SetScopes([]string{model.ScopeServerRead})
|
||||||
|
require.NoError(t, singleton.DB.Create(&tok).Error)
|
||||||
|
require.NoError(t, singleton.DB.Create(&model.User{Common: model.Common{ID: 10}}).Error)
|
||||||
|
|
||||||
|
require.NoError(t, singleton.DB.Delete(&model.APIToken{}, tok.ID).Error)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||||
|
c.Request.Header.Set("Authorization", "Bearer "+plain)
|
||||||
|
apiTokenAuthMiddleware()(c)
|
||||||
|
require.True(t, c.IsAborted(), "revoked PAT must abort")
|
||||||
|
require.Equal(t, http.StatusUnauthorized, w.Code)
|
||||||
|
require.Contains(t, w.Body.String(), "invalid api token")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPITokenAuthMW_OversizedTokenIsRejected(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||||
|
c.Request.Header.Set("Authorization", "Bearer nzp_"+strings.Repeat("X", 100*1024))
|
||||||
|
apiTokenAuthMiddleware()(c)
|
||||||
|
require.True(t, c.IsAborted(), "huge nzp_ body must still abort (no DoS via lookup)")
|
||||||
|
require.Equal(t, http.StatusUnauthorized, w.Code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPITokenAuthMW_NonASCIITokenIsRejected(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||||
|
c.Request.Header.Set("Authorization", "Bearer nzp_中文😀💀")
|
||||||
|
apiTokenAuthMiddleware()(c)
|
||||||
|
require.True(t, c.IsAborted(), "non-ASCII nzp_ body must be rejected by hash lookup")
|
||||||
|
require.Equal(t, http.StatusUnauthorized, w.Code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPITokenAuthMW_SQLInjectionAttemptIsRejected(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
plain := "nzp_" + strings.Repeat("s", 32)
|
||||||
|
tok := model.APIToken{
|
||||||
|
UserID: 10,
|
||||||
|
Name: "real",
|
||||||
|
TokenHash: model.HashAPIToken(plain),
|
||||||
|
}
|
||||||
|
tok.SetScopes([]string{model.ScopeServerRead})
|
||||||
|
require.NoError(t, singleton.DB.Create(&tok).Error)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||||
|
c.Request.Header.Set("Authorization", "Bearer nzp_'; DROP TABLE api_tokens; --")
|
||||||
|
apiTokenAuthMiddleware()(c)
|
||||||
|
require.True(t, c.IsAborted(), "SQL-injection-shaped token must just look up and miss")
|
||||||
|
require.Equal(t, http.StatusUnauthorized, w.Code)
|
||||||
|
|
||||||
|
var stored model.APIToken
|
||||||
|
require.NoError(t, singleton.DB.First(&stored, tok.ID).Error,
|
||||||
|
"real token row must survive — GORM uses prepared statements")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPITokenAuthMW_LowercaseBearerPassesThrough(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
plain := "nzp_" + strings.Repeat("l", 32)
|
||||||
|
tok := model.APIToken{UserID: 10, Name: "x", TokenHash: model.HashAPIToken(plain)}
|
||||||
|
tok.SetScopes([]string{model.ScopeServerRead})
|
||||||
|
require.NoError(t, singleton.DB.Create(&tok).Error)
|
||||||
|
require.NoError(t, singleton.DB.Create(&model.User{Common: model.Common{ID: 10}}).Error)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||||
|
c.Request.Header.Set("Authorization", "bearer "+plain)
|
||||||
|
apiTokenAuthMiddleware()(c)
|
||||||
|
require.False(t, c.IsAborted(),
|
||||||
|
"lowercase 'bearer ' is not the canonical scheme; PAT mw must skip it (RFC 7235 says scheme is case-insensitive, "+
|
||||||
|
"but we deliberately match GitHub/AWS behaviour of strict 'Bearer ' to keep PAT/JWT lookup paths predictable)")
|
||||||
|
require.Nil(t, APITokenFromContext(c), "lowercase bearer must not register as PAT")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPITokenAuthMW_TrailingWhitespaceTolerated(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
plain := "nzp_" + strings.Repeat("w", 32)
|
||||||
|
tok := model.APIToken{UserID: 10, Name: "trim", TokenHash: model.HashAPIToken(plain)}
|
||||||
|
tok.SetScopes([]string{model.ScopeServerRead})
|
||||||
|
require.NoError(t, singleton.DB.Create(&tok).Error)
|
||||||
|
require.NoError(t, singleton.DB.Create(&model.User{Common: model.Common{ID: 10}}).Error)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/", nil)
|
||||||
|
c.Request.Header.Set("Authorization", "Bearer "+plain+" ")
|
||||||
|
apiTokenAuthMiddleware()(c)
|
||||||
|
require.False(t, c.IsAborted(),
|
||||||
|
"trailing/leading whitespace around PAT must be trimmed (curl users often paste with newlines)")
|
||||||
|
require.NotNil(t, APITokenFromContext(c))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_NameTooLongRejected(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: strings.Repeat("X", 129),
|
||||||
|
Scopes: []string{model.ScopeServerRead},
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err, "name >128 chars must be rejected (binding tag max=128 or handler check)")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateAPIToken_EmptyNameRejected(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: " ",
|
||||||
|
Scopes: []string{model.ScopeServerRead},
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Contains(t, err.Error(), "name required")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPIToken_DuplicateHashViolatesUniqueIndex(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
plain := "nzp_" + strings.Repeat("a", 32)
|
||||||
|
for _, uid := range []uint64{10, 11} {
|
||||||
|
tok := model.APIToken{
|
||||||
|
UserID: uid,
|
||||||
|
Name: "dup",
|
||||||
|
TokenHash: model.HashAPIToken(plain),
|
||||||
|
}
|
||||||
|
tok.SetScopes([]string{model.ScopeServerRead})
|
||||||
|
err := singleton.DB.Create(&tok).Error
|
||||||
|
if uid == 10 {
|
||||||
|
require.NoError(t, err, "first insert must succeed")
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
require.Error(t, err, "duplicate hash must violate unique index (defense against forged tokens)")
|
||||||
|
require.Contains(t, err.Error(), "UNIQUE")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListAPITokens_ReturnsOnlyOwn(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
for i, uid := range []uint64{10, 10, 20} {
|
||||||
|
tok := model.APIToken{UserID: uid, Name: "n", TokenHash: model.HashAPIToken("nzp_unique_" + itoa(uint64(i)))}
|
||||||
|
tok.SetScopes([]string{model.ScopeServerRead})
|
||||||
|
require.NoError(t, singleton.DB.Create(&tok).Error)
|
||||||
|
}
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
got, err := listAPITokens(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, got, 2)
|
||||||
|
}
|
||||||
|
|
||||||
|
func itoa(v uint64) string {
|
||||||
|
return strings.TrimSpace(jsonNum(v))
|
||||||
|
}
|
||||||
|
|
||||||
|
func jsonNum(v uint64) string {
|
||||||
|
b, _ := json.Marshal(v)
|
||||||
|
return string(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
// scope_doc.go and HasScope advertise nezha:<resource>:* as a first-class
|
||||||
|
// scope shape, and rest_scope_test.go pins runtime support for it. The
|
||||||
|
// create-API-token endpoint must accept those wildcards too — otherwise
|
||||||
|
// the documented surface is unreachable via the only endpoint that can
|
||||||
|
// issue PATs.
|
||||||
|
func TestCreateAPIToken_AcceptsResourceWildcardScopes(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
|
||||||
|
cases := []string{
|
||||||
|
"nezha:server:*",
|
||||||
|
"nezha:service:*",
|
||||||
|
"nezha:cron:*",
|
||||||
|
"nezha:transfer:*",
|
||||||
|
}
|
||||||
|
for _, scope := range cases {
|
||||||
|
t.Run(scope, func(t *testing.T) {
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "wildcard-" + scope,
|
||||||
|
Scopes: []string{scope},
|
||||||
|
})
|
||||||
|
res, err := createAPIToken(c)
|
||||||
|
require.NoError(t, err, "resource wildcard %q must be issuable", scope)
|
||||||
|
require.Contains(t, res.Scopes, scope)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// nezha:admin:* is admin-only and already on AdminOnlyScopes; this test
|
||||||
|
// ensures the new wildcard acceptance does NOT widen admin-only scopes
|
||||||
|
// to members.
|
||||||
|
func TestCreateAPIToken_ResourceWildcardStillRejectsAdminOnly(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(10, model.RoleMember)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "member-tries-admin-wildcard",
|
||||||
|
Scopes: []string{"nezha:admin:*"},
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Contains(t, err.Error(), "admin")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unknown resources must still be rejected even with a wildcard verb so
|
||||||
|
// that nezha:bogus:* does not become a forward-compat blank cheque.
|
||||||
|
func TestCreateAPIToken_RejectsUnknownResourceWildcard(t *testing.T) {
|
||||||
|
defer setupAPITokenTest(t)()
|
||||||
|
c := ctxAsUser(1, model.RoleAdmin)
|
||||||
|
bindJSON(c, model.APITokenCreateRequest{
|
||||||
|
Name: "bogus",
|
||||||
|
Scopes: []string{"nezha:bogus:*"},
|
||||||
|
})
|
||||||
|
_, err := createAPIToken(c)
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Contains(t, err.Error(), "unknown scope")
|
||||||
|
}
|
||||||
@@ -0,0 +1,71 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func callBatchMoveWithPAT(t *testing.T, callerID uint64, role model.Role, tok *model.APIToken, body string) ([]model.BatchMoveServerResult, bool, string) {
|
||||||
|
t.Helper()
|
||||||
|
r := gin.New()
|
||||||
|
r.Use(newPATCtxSetter(callerID, role, tok))
|
||||||
|
r.POST("/batch-move/server", commonHandler(batchMoveServer))
|
||||||
|
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/batch-move/server", bytes.NewReader([]byte(body)))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
var resp struct {
|
||||||
|
Success bool `json:"success"`
|
||||||
|
Error string `json:"error"`
|
||||||
|
Data []model.BatchMoveServerResult `json:"data"`
|
||||||
|
}
|
||||||
|
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
|
||||||
|
return resp.Data, resp.Success, resp.Error
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBatchMoveServer_AdminPATScopeNarrowsServerIDs(t *testing.T) {
|
||||||
|
cleanup := setupRetryServerTransferFixture(t)
|
||||||
|
defer cleanup()
|
||||||
|
seedServer(t, 1, 999)
|
||||||
|
seedServer(t, 2, 999)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 18, UserID: 999}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
data, ok, errStr := callBatchMoveWithPAT(t, 999, model.RoleAdmin, tok,
|
||||||
|
`{"ids":[1,2],"to_user":200}`)
|
||||||
|
assert.True(t, ok, "batch-move call must succeed at the request layer: %s", errStr)
|
||||||
|
require.Len(t, data, 2)
|
||||||
|
|
||||||
|
resultByID := map[uint64]model.BatchMoveServerResult{}
|
||||||
|
for _, r := range data {
|
||||||
|
resultByID[r.ServerID] = r
|
||||||
|
}
|
||||||
|
|
||||||
|
assert.NotEqual(t, model.BatchMoveServerResultPending, resultByID[2].Status,
|
||||||
|
"admin PAT scoped to {1} MUST NOT be able to move server 2; got status=%q error=%q",
|
||||||
|
resultByID[2].Status, resultByID[2].Error)
|
||||||
|
|
||||||
|
var pending int64
|
||||||
|
require.NoError(t, singleton.DB.Model(&model.ServerTransfer{}).
|
||||||
|
Where("server_id = ? AND status = ?", 2, model.ServerTransferStatusPending).
|
||||||
|
Count(&pending).Error)
|
||||||
|
assert.Equal(t, int64(0), pending,
|
||||||
|
"rejected batch-move of server 2 must not create a Pending row")
|
||||||
|
|
||||||
|
assert.Equal(t, model.BatchMoveServerResultPending, resultByID[1].Status,
|
||||||
|
"admin PAT scoped to {1} must still be able to move server 1; got status=%q error=%q",
|
||||||
|
resultByID[1].Status, resultByID[1].Error)
|
||||||
|
}
|
||||||
@@ -8,7 +8,6 @@ import (
|
|||||||
"log"
|
"log"
|
||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
"os"
|
||||||
"path"
|
|
||||||
"regexp"
|
"regexp"
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -45,6 +44,8 @@ func ServeWeb(frontendDist fs.FS) http.Handler {
|
|||||||
|
|
||||||
routers(r, frontendDist)
|
routers(r, frontendDist)
|
||||||
|
|
||||||
|
kickoffTransferGC()
|
||||||
|
|
||||||
return r
|
return r
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -56,6 +57,20 @@ func routers(r *gin.Engine, frontendDist fs.FS) {
|
|||||||
if err := authMiddleware.MiddlewareInit(); err != nil {
|
if err := authMiddleware.MiddlewareInit(); err != nil {
|
||||||
log.Fatal("authMiddleware.MiddlewareInit Error:" + err.Error())
|
log.Fatal("authMiddleware.MiddlewareInit Error:" + err.Error())
|
||||||
}
|
}
|
||||||
|
// /mcp — Model Context Protocol endpoint, authenticated by PAT only (闸 1 + 闸 2)。
|
||||||
|
// 不放在 /api/v1 下:MCP client 配置 URL 更短,且 MCP transport 协议演进与 REST API
|
||||||
|
// 解耦。鉴权一律走 apiTokenAuthMiddleware;不接受 JWT 以避免浏览器误触。
|
||||||
|
// mcpOriginGuard 防止 DNS rebinding / 浏览器跨站调用。
|
||||||
|
r.POST("/mcp", mcpOriginGuard(), apiTokenAuthMiddleware(), mcpEndpoint)
|
||||||
|
// Streamable HTTP 规范要求:不实现 standalone SSE / session 时,GET / DELETE
|
||||||
|
// 必须显式返回 405,让客户端走 POST-only 路径并跳过 session 终止流程。
|
||||||
|
// 不显式注册时,Gin 会走 NoRoute → fallbackToFrontend,对 MCP 客户端是 HTML/404。
|
||||||
|
r.GET("/mcp", mcpOriginGuard(), mcpMethodNotAllowed)
|
||||||
|
r.DELETE("/mcp", mcpOriginGuard(), mcpMethodNotAllowed)
|
||||||
|
r.GET("/mcp/download/:token", mcpOriginGuard(), transferDownloadHandler)
|
||||||
|
r.POST("/mcp/upload/:token", mcpOriginGuard(), transferUploadHandler)
|
||||||
|
registerAgentcompatRoutes(r)
|
||||||
|
|
||||||
api := r.Group("api/v1")
|
api := r.Group("api/v1")
|
||||||
api.POST("/login", authMiddleware.LoginHandler)
|
api.POST("/login", authMiddleware.LoginHandler)
|
||||||
api.GET("/oauth2/:provider", commonHandler(oauth2redirect))
|
api.GET("/oauth2/:provider", commonHandler(oauth2redirect))
|
||||||
@@ -68,92 +83,113 @@ func routers(r *gin.Engine, frontendDist fs.FS) {
|
|||||||
fallbackAuth.GET("/setting", commonHandler(listConfig))
|
fallbackAuth.GET("/setting", commonHandler(listConfig))
|
||||||
fallbackAuth.GET("/oauth2/callback", commonHandler(oauth2callback(authMiddleware)))
|
fallbackAuth.GET("/oauth2/callback", commonHandler(oauth2callback(authMiddleware)))
|
||||||
|
|
||||||
authMw := authMiddleware.MiddlewareFunc()
|
jwtMw := authMiddleware.MiddlewareFunc()
|
||||||
optionalAuthMw := utils.IfOr(singleton.Conf.ForceAuth, authMw, fallbackAuthMw)
|
patMw := apiTokenAuthMiddleware()
|
||||||
|
authMw := jwtOrPATAuthMiddleware(patMw, jwtMw)
|
||||||
|
// optional 路由:ForceAuth=true 走严格 PAT-or-JWT;ForceAuth=false 走
|
||||||
|
// PAT-or-FallbackJWT,保证两种模式下 PAT 都会被解析,restScopeMiddleware
|
||||||
|
// 才能按 scope 真实收口(否则匿名 PAT 请求会被当 guest,scope 失效)。
|
||||||
|
optionalAuthMw := utils.IfOr(singleton.Conf.ForceAuth, authMw, patOrFallbackAuthMiddleware(patMw, fallbackAuthMw))
|
||||||
|
|
||||||
optionalAuth := api.Group("", optionalAuthMw)
|
optionalAuth := api.Group("", optionalAuthMw)
|
||||||
optionalAuth.GET("/ws/server", commonHandler(serverStream))
|
optionalAuth.GET("/ws/server", restScopeMiddleware(model.ScopeInventoryRead), commonHandler(serverStream))
|
||||||
optionalAuth.GET("/server-group", commonHandler(listServerGroup))
|
optionalAuth.GET("/server-group", restScopeMiddleware(model.ScopeInventoryRead), commonHandler(listServerGroup))
|
||||||
|
|
||||||
optionalAuth.GET("/service", commonHandler(showService))
|
optionalAuth.GET("/service", restScopeMiddleware(model.ScopeServiceRead), commonHandler(showService))
|
||||||
optionalAuth.GET("/service/server", commonHandler(listServerWithServices))
|
optionalAuth.GET("/service/server", restScopeMiddleware(model.ScopeServiceRead), commonHandler(listServerWithServices))
|
||||||
optionalAuth.GET("/domains", commonHandler(GetDomainList))
|
optionalAuth.GET("/domains", commonHandler(GetDomainList))
|
||||||
optionalAuth.GET("/service/:id/history", commonHandler(getServiceHistory))
|
optionalAuth.GET("/service/:id/history", restScopeMiddleware(model.ScopeServiceRead), commonHandler(getServiceHistory))
|
||||||
optionalAuth.GET("/server/:id/service", commonHandler(listServerServices))
|
optionalAuth.GET("/server/:id/service", restScopeMiddleware(model.ScopeServiceRead), commonHandler(listServerServices))
|
||||||
optionalAuth.GET("/server/:id/metrics", commonHandler(getServerMetrics))
|
optionalAuth.GET("/server/:id/metrics", restScopeMiddleware(model.ScopeServerRead), commonHandler(getServerMetrics))
|
||||||
|
|
||||||
auth := api.Group("", authMw)
|
// CSRF middleware applies group-wide. Safe methods short-circuit and
|
||||||
|
// PAT bearer requests bypass — so the only callers gated are
|
||||||
|
// cookie-JWT POST/PATCH/PUT/DELETE, which is exactly the H6 surface.
|
||||||
|
auth := api.Group("", authMw, csrfMiddleware())
|
||||||
|
|
||||||
auth.GET("/refresh-token", authMiddleware.RefreshHandler)
|
// 「自我管理」类端点 — 显式禁止 PAT 访问(避免 PAT 自我提权链)。
|
||||||
|
patForbidden := restPATForbiddenMiddleware()
|
||||||
|
auth.POST("/refresh-token", patForbidden, authMiddleware.RefreshHandler)
|
||||||
|
auth.GET("/profile", patForbidden, commonHandler(getProfile))
|
||||||
|
auth.POST("/profile", patForbidden, commonHandler(updateProfile))
|
||||||
|
auth.POST("/oauth2/:provider/unbind", patForbidden, commonHandler(unbindOauth2))
|
||||||
|
auth.GET("/api-tokens", patForbidden, commonHandler(listAPITokens))
|
||||||
|
auth.POST("/api-tokens", patForbidden, commonHandler(createAPIToken))
|
||||||
|
auth.DELETE("/api-tokens/:id", patForbidden, commonHandler(deleteAPIToken))
|
||||||
|
|
||||||
auth.GET("/file", commonHandler(createFM))
|
// 资源族划分:
|
||||||
auth.GET("/ws/file/:id", commonHandler(fmStream))
|
// - nezha:inventory:* —— 对“服务器台账”的枚举与删除(列出 server / server-group、
|
||||||
|
// 删除 server / server-group)。这是管理后台清单管理动作。
|
||||||
|
// - nezha:server:* —— 对已知 server 的运行态操作(文件读写、编辑配置、
|
||||||
|
// force-update、batch-move)。(Web Terminal 按安全要求已移除)
|
||||||
|
auth.POST("/file", restScopeAllOf(model.ScopeServerRead, model.ScopeServerWrite, model.ScopeServerDelete), commonHandler(createFM))
|
||||||
|
auth.GET("/ws/file/:id", restScopeAllOf(model.ScopeServerRead, model.ScopeServerWrite, model.ScopeServerDelete), commonHandler(fmStream))
|
||||||
|
auth.GET("/server", restScopeMiddleware(model.ScopeInventoryRead), listHandler(listServer))
|
||||||
|
auth.PATCH("/server/:id", restScopeMiddleware(model.ScopeServerWrite), commonHandler(updateServer))
|
||||||
|
auth.GET("/server/config/:id", restScopeMiddleware(serverConfigSensitiveScope()), commonHandler(getServerConfig))
|
||||||
|
auth.POST("/server/config", restScopeMiddleware(model.ScopeServerWrite), commonHandler(setServerConfig))
|
||||||
|
auth.POST("/batch-delete/server", restScopeMiddleware(model.ScopeInventoryDelete), commonHandler(batchDeleteServer))
|
||||||
|
auth.POST("/batch-move/server", restScopeMiddleware(model.ScopeServerWrite), commonHandler(batchMoveServer))
|
||||||
|
auth.POST("/force-update/server", restScopeMiddleware(model.ScopeServerWrite), commonHandler(forceUpdateServer))
|
||||||
|
auth.POST("/server-group", restScopeMiddleware(model.ScopeServerWrite), commonHandler(createServerGroup))
|
||||||
|
auth.PATCH("/server-group/:id", restScopeMiddleware(model.ScopeServerWrite), commonHandler(updateServerGroup))
|
||||||
|
auth.POST("/batch-delete/server-group", restScopeMiddleware(model.ScopeInventoryDelete), commonHandler(batchDeleteServerGroup))
|
||||||
|
|
||||||
auth.GET("/profile", commonHandler(getProfile))
|
// transfer — 严格使用 nezha:transfer 资源族 scope(read/write/delete)。
|
||||||
auth.POST("/profile", commonHandler(updateProfile))
|
auth.GET("/transfer", restScopeMiddleware(model.ScopeTransferRead), listHandler(listServerTransfer))
|
||||||
auth.POST("/oauth2/:provider/unbind", commonHandler(unbindOauth2))
|
auth.POST("/transfer/:id/cancel", restScopeMiddleware(model.ScopeTransferWrite), commonHandler(cancelServerTransfer))
|
||||||
|
auth.POST("/transfer/:id/retry", restScopeMiddleware(model.ScopeTransferWrite), commonHandler(retryServerTransfer))
|
||||||
|
auth.GET("/ws/transfer", restScopeMiddleware(model.ScopeTransferRead), commonHandler(transferStream))
|
||||||
|
|
||||||
auth.GET("/user", adminHandler(listUser))
|
|
||||||
auth.POST("/user", adminHandler(createUser))
|
|
||||||
auth.POST("/batch-delete/user", adminHandler(batchDeleteUser))
|
|
||||||
|
|
||||||
auth.GET("/service/list", listHandler(listService))
|
// service monitor
|
||||||
auth.POST("/service", commonHandler(createService))
|
auth.GET("/service/list", restScopeMiddleware(model.ScopeServiceRead), listHandler(listService))
|
||||||
auth.PATCH("/service/:id", commonHandler(updateService))
|
auth.POST("/service", restScopeMiddleware(model.ScopeServiceWrite), commonHandler(createService))
|
||||||
auth.POST("/batch-delete/service", commonHandler(batchDeleteService))
|
auth.PATCH("/service/:id", restScopeMiddleware(model.ScopeServiceWrite), commonHandler(updateService))
|
||||||
|
auth.POST("/batch-delete/service", restScopeMiddleware(model.ScopeServiceDelete), commonHandler(batchDeleteService))
|
||||||
|
|
||||||
auth.POST("/server-group", commonHandler(createServerGroup))
|
auth.GET("/notification-group", restScopeMiddleware(model.ScopeNotificationGroupRead), commonHandler(listNotificationGroup))
|
||||||
auth.PATCH("/server-group/:id", commonHandler(updateServerGroup))
|
auth.POST("/notification-group", restScopeMiddleware(model.ScopeNotificationGroupWrite), commonHandler(createNotificationGroup))
|
||||||
auth.POST("/batch-delete/server-group", commonHandler(batchDeleteServerGroup))
|
auth.PATCH("/notification-group/:id", restScopeMiddleware(model.ScopeNotificationGroupWrite), commonHandler(updateNotificationGroup))
|
||||||
|
auth.POST("/batch-delete/notification-group", restScopeMiddleware(model.ScopeNotificationGroupDelete), commonHandler(batchDeleteNotificationGroup))
|
||||||
|
|
||||||
auth.GET("/notification-group", commonHandler(listNotificationGroup))
|
auth.GET("/notification", restScopeMiddleware(model.ScopeNotificationRead), listHandler(listNotification))
|
||||||
auth.POST("/notification-group", commonHandler(createNotificationGroup))
|
auth.POST("/notification", restScopeMiddleware(model.ScopeNotificationWrite), commonHandler(createNotification))
|
||||||
auth.PATCH("/notification-group/:id", commonHandler(updateNotificationGroup))
|
auth.PATCH("/notification/:id", restScopeMiddleware(model.ScopeNotificationWrite), commonHandler(updateNotification))
|
||||||
auth.POST("/batch-delete/notification-group", commonHandler(batchDeleteNotificationGroup))
|
auth.POST("/batch-delete/notification", restScopeMiddleware(model.ScopeNotificationDelete), commonHandler(batchDeleteNotification))
|
||||||
|
|
||||||
auth.GET("/server", listHandler(listServer))
|
auth.GET("/alert-rule", restScopeMiddleware(model.ScopeAlertRuleRead), listHandler(listAlertRule))
|
||||||
auth.PATCH("/server/:id", commonHandler(updateServer))
|
auth.POST("/alert-rule", restScopeMiddleware(model.ScopeAlertRuleWrite), commonHandler(createAlertRule))
|
||||||
auth.GET("/server/config/:id", commonHandler(getServerConfig))
|
auth.PATCH("/alert-rule/:id", restScopeMiddleware(model.ScopeAlertRuleWrite), commonHandler(updateAlertRule))
|
||||||
auth.POST("/server/config", commonHandler(setServerConfig))
|
auth.POST("/batch-delete/alert-rule", restScopeMiddleware(model.ScopeAlertRuleDelete), commonHandler(batchDeleteAlertRule))
|
||||||
auth.POST("/batch-delete/server", commonHandler(batchDeleteServer))
|
|
||||||
auth.POST("/batch-move/server", commonHandler(batchMoveServer))
|
|
||||||
auth.POST("/force-update/server", commonHandler(forceUpdateServer))
|
|
||||||
|
|
||||||
auth.GET("/notification", listHandler(listNotification))
|
auth.GET("/cron", restScopeMiddleware(model.ScopeCronRead), listHandler(listCron))
|
||||||
auth.POST("/notification", commonHandler(createNotification))
|
auth.POST("/cron", restScopeMiddleware(model.ScopeCronWrite), commonHandler(createCron))
|
||||||
auth.PATCH("/notification/:id", commonHandler(updateNotification))
|
auth.PATCH("/cron/:id", restScopeMiddleware(model.ScopeCronWrite), commonHandler(updateCron))
|
||||||
auth.POST("/batch-delete/notification", commonHandler(batchDeleteNotification))
|
auth.POST("/cron/:id/manual", restScopeMiddleware(model.ScopeCronExec), commonHandler(manualTriggerCron))
|
||||||
|
auth.POST("/batch-delete/cron", restScopeMiddleware(model.ScopeCronDelete), commonHandler(batchDeleteCron))
|
||||||
|
|
||||||
auth.GET("/alert-rule", listHandler(listAlertRule))
|
auth.GET("/ddns", restScopeMiddleware(model.ScopeDDNSRead), listHandler(listDDNS))
|
||||||
auth.POST("/alert-rule", commonHandler(createAlertRule))
|
auth.GET("/ddns/providers", restScopeMiddleware(model.ScopeDDNSRead), commonHandler(listProviders))
|
||||||
auth.PATCH("/alert-rule/:id", commonHandler(updateAlertRule))
|
auth.POST("/ddns", restScopeMiddleware(model.ScopeDDNSWrite), commonHandler(createDDNS))
|
||||||
auth.POST("/batch-delete/alert-rule", commonHandler(batchDeleteAlertRule))
|
auth.PATCH("/ddns/:id", restScopeMiddleware(model.ScopeDDNSWrite), commonHandler(updateDDNS))
|
||||||
|
auth.POST("/batch-delete/ddns", restScopeMiddleware(model.ScopeDDNSDelete), commonHandler(batchDeleteDDNS))
|
||||||
|
|
||||||
auth.GET("/cron", listHandler(listCron))
|
auth.GET("/nat", restScopeMiddleware(model.ScopeNATRead), listHandler(listNAT))
|
||||||
auth.POST("/cron", commonHandler(createCron))
|
auth.POST("/nat", restScopeMiddleware(model.ScopeNATWrite), commonHandler(createNAT))
|
||||||
auth.PATCH("/cron/:id", commonHandler(updateCron))
|
auth.PATCH("/nat/:id", restScopeMiddleware(model.ScopeNATWrite), commonHandler(updateNAT))
|
||||||
auth.GET("/cron/:id/manual", commonHandler(manualTriggerCron))
|
auth.POST("/batch-delete/nat", restScopeMiddleware(model.ScopeNATDelete), commonHandler(batchDeleteNAT))
|
||||||
auth.POST("/batch-delete/cron", commonHandler(batchDeleteCron))
|
|
||||||
|
|
||||||
auth.GET("/ddns", listHandler(listDDNS))
|
// 管理员资源 — 仅 nezha:* / nezha:admin:* 持有者可调(adminHandler 进一步校验 user.Role)。
|
||||||
auth.GET("/ddns/providers", commonHandler(listProviders))
|
auth.GET("/user", restScopeMiddleware(model.ScopeAdminAll), adminHandler(listUser))
|
||||||
auth.POST("/ddns", commonHandler(createDDNS))
|
auth.POST("/user", restScopeMiddleware(model.ScopeAdminAll), adminHandler(createUser))
|
||||||
auth.PATCH("/ddns/:id", commonHandler(updateDDNS))
|
auth.POST("/batch-delete/user", restScopeMiddleware(model.ScopeAdminAll), adminHandler(batchDeleteUser))
|
||||||
auth.POST("/batch-delete/ddns", commonHandler(batchDeleteDDNS))
|
auth.GET("/waf", restScopeMiddleware(model.ScopeAdminAll), pAdminHandler(listBlockedAddress))
|
||||||
|
auth.POST("/batch-delete/waf", restScopeMiddleware(model.ScopeAdminAll), adminHandler(batchDeleteBlockedAddress))
|
||||||
auth.GET("/nat", listHandler(listNAT))
|
auth.GET("/online-user", restScopeMiddleware(model.ScopeAdminAll), pAdminHandler(listOnlineUser))
|
||||||
auth.POST("/nat", commonHandler(createNAT))
|
auth.POST("/online-user/batch-block", restScopeMiddleware(model.ScopeAdminAll), adminHandler(batchBlockOnlineUser))
|
||||||
auth.PATCH("/nat/:id", commonHandler(updateNAT))
|
auth.PATCH("/setting", restScopeMiddleware(model.ScopeAdminAll), adminHandler(updateConfig))
|
||||||
auth.POST("/batch-delete/nat", commonHandler(batchDeleteNAT))
|
auth.POST("/maintenance", restScopeMiddleware(model.ScopeAdminAll), adminHandler(runMaintenance))
|
||||||
|
|
||||||
auth.GET("/waf", pCommonHandler(listBlockedAddress))
|
|
||||||
auth.POST("/batch-delete/waf", adminHandler(batchDeleteBlockedAddress))
|
|
||||||
|
|
||||||
auth.GET("/online-user", pCommonHandler(listOnlineUser))
|
|
||||||
auth.POST("/online-user/batch-block", adminHandler(batchBlockOnlineUser))
|
|
||||||
|
|
||||||
auth.PATCH("/setting", adminHandler(updateConfig))
|
|
||||||
auth.POST("/maintenance", adminHandler(runMaintenance))
|
|
||||||
|
|
||||||
auth.POST("/domains", commonHandler(AddDomain))
|
auth.POST("/domains", commonHandler(AddDomain))
|
||||||
auth.POST("/domains/:id/verify", commonHandler(VerifyDomain))
|
auth.POST("/domains/:id/verify", commonHandler(VerifyDomain))
|
||||||
@@ -293,6 +329,29 @@ func pCommonHandler[S ~[]E, E any](handler pHandlerFunc[S, E]) func(*gin.Context
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func pAdminHandler[S ~[]E, E any](handler pHandlerFunc[S, E]) func(*gin.Context) {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
auth, ok := c.Get(model.CtxKeyAuthorizedUser)
|
||||||
|
if !ok {
|
||||||
|
c.JSON(http.StatusOK, newErrorResponse(singleton.Localizer.ErrorT("unauthorized")))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
user := *auth.(*model.User)
|
||||||
|
if !user.Role.IsAdmin() {
|
||||||
|
c.JSON(http.StatusOK, newErrorResponse(singleton.Localizer.ErrorT("permission denied")))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := handler(c)
|
||||||
|
if err != nil {
|
||||||
|
c.JSON(http.StatusOK, newErrorResponse(err))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
c.JSON(http.StatusOK, model.PaginatedResponse[S, E]{Success: true, Data: data})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func filter[S ~[]E, E model.CommonInterface](ctx *gin.Context, s S) S {
|
func filter[S ~[]E, E model.CommonInterface](ctx *gin.Context, s S) S {
|
||||||
return slices.DeleteFunc(s, func(e E) bool {
|
return slices.DeleteFunc(s, func(e E) bool {
|
||||||
return !e.HasPermission(ctx)
|
return !e.HasPermission(ctx)
|
||||||
@@ -305,27 +364,52 @@ func getUid(c *gin.Context) uint64 {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func fallbackToFrontend(frontendDist fs.FS) func(*gin.Context) {
|
func fallbackToFrontend(frontendDist fs.FS) func(*gin.Context) {
|
||||||
checkLocalFileOrFs := func(c *gin.Context, fs fs.FS, path string, customStatusCode int) bool {
|
serveFile := func(c *gin.Context, name string, file fs.File, customStatusCode int) bool {
|
||||||
if _, err := os.Stat(path); err == nil {
|
defer file.Close()
|
||||||
http.ServeFile(utils.NewGinCustomWriter(c, customStatusCode), c.Request, path)
|
fileStat, err := file.Stat()
|
||||||
return true
|
|
||||||
}
|
|
||||||
f, err := fs.Open(path)
|
|
||||||
if err != nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
defer f.Close()
|
|
||||||
fileStat, err := f.Stat()
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
if fileStat.IsDir() {
|
if fileStat.IsDir() {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
http.ServeContent(utils.NewGinCustomWriter(c, customStatusCode), c.Request, path, fileStat.ModTime(), f.(io.ReadSeeker))
|
readSeeker, ok := file.(io.ReadSeeker)
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
http.ServeContent(utils.NewGinCustomWriter(c, customStatusCode), c.Request, name, fileStat.ModTime(), readSeeker)
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
checkLocalFileOrFs := func(c *gin.Context, frontendFS fs.FS, templateRoot, filePath string, customStatusCode int) bool {
|
||||||
|
if filePath != "" {
|
||||||
|
localRoot, err := os.OpenRoot(templateRoot)
|
||||||
|
if err == nil {
|
||||||
|
defer localRoot.Close()
|
||||||
|
// URL paths must stay inside the selected template root; never join them against the process cwd.
|
||||||
|
if file, err := localRoot.Open(filePath); err == nil && serveFile(c, filePath, file, customStatusCode) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if !fs.ValidPath(filePath) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
templateFS, err := fs.Sub(frontendFS, templateRoot)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
file, err := templateFS.Open(filePath)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if serveFile(c, filePath, file, customStatusCode) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
frontendPageUrlRegistry := []*regexp.Regexp{
|
frontendPageUrlRegistry := []*regexp.Regexp{
|
||||||
// official user frontend
|
// official user frontend
|
||||||
regexp.MustCompile(`^/$`),
|
regexp.MustCompile(`^/$`),
|
||||||
@@ -346,6 +430,12 @@ func fallbackToFrontend(frontendDist fs.FS) func(*gin.Context) {
|
|||||||
regexp.MustCompile(`^/dashboard/settings/user$`),
|
regexp.MustCompile(`^/dashboard/settings/user$`),
|
||||||
regexp.MustCompile(`^/dashboard/settings/online-user$`),
|
regexp.MustCompile(`^/dashboard/settings/online-user$`),
|
||||||
regexp.MustCompile(`^/dashboard/settings/waf$`),
|
regexp.MustCompile(`^/dashboard/settings/waf$`),
|
||||||
|
regexp.MustCompile(`^/dashboard/settings/api-tokens$`),
|
||||||
|
// 注意:这里的白名单决定哪些 URL 走 index.html fallback;漏一条就会把
|
||||||
|
// 直接刷新该页面变成 404(HTTP 状态码层面,body 仍是 index.html,所以
|
||||||
|
// 浏览器内 SPA 看起来正常,但 monitoring / 链接预览会以为站点挂了)。
|
||||||
|
// 新增前端路由时必须在 admin-frontend/src/main.tsx 与这里同步加。
|
||||||
|
regexp.MustCompile(`^/dashboard/transfer$`),
|
||||||
}
|
}
|
||||||
|
|
||||||
getFallbackStatusCode := func(path string) int {
|
getFallbackStatusCode := func(path string) int {
|
||||||
@@ -370,22 +460,22 @@ func fallbackToFrontend(frontendDist fs.FS) func(*gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fallbackStatusCode := getFallbackStatusCode(c.Request.URL.Path)
|
fallbackStatusCode := getFallbackStatusCode(c.Request.URL.Path)
|
||||||
if strings.HasPrefix(c.Request.URL.Path, "/dashboard") {
|
// Only /dashboard/ belongs to the admin frontend; /dashboard.. must not be trimmed into ../.
|
||||||
stripPath := strings.TrimPrefix(c.Request.URL.Path, "/dashboard")
|
if strings.HasPrefix(c.Request.URL.Path, "/dashboard/") {
|
||||||
localFilePath := path.Join(singleton.Conf.AdminTemplate, stripPath)
|
stripPath := strings.TrimPrefix(c.Request.URL.Path, "/dashboard/")
|
||||||
if checkLocalFileOrFs(c, frontendDist, localFilePath, http.StatusOK) {
|
if checkLocalFileOrFs(c, frontendDist, singleton.Conf.AdminTemplate, stripPath, http.StatusOK) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if !checkLocalFileOrFs(c, frontendDist, singleton.Conf.AdminTemplate+"/index.html", fallbackStatusCode) {
|
if !checkLocalFileOrFs(c, frontendDist, singleton.Conf.AdminTemplate, "index.html", fallbackStatusCode) {
|
||||||
c.JSON(http.StatusNotFound, newErrorResponse(errors.New("404 Not Found")))
|
c.JSON(http.StatusNotFound, newErrorResponse(errors.New("404 Not Found")))
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
localFilePath := path.Join(singleton.Conf.UserTemplate, c.Request.URL.Path)
|
stripPath := strings.TrimPrefix(c.Request.URL.Path, "/")
|
||||||
if checkLocalFileOrFs(c, frontendDist, localFilePath, http.StatusOK) {
|
if checkLocalFileOrFs(c, frontendDist, singleton.Conf.UserTemplate, stripPath, http.StatusOK) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if !checkLocalFileOrFs(c, frontendDist, singleton.Conf.UserTemplate+"/index.html", fallbackStatusCode) {
|
if !checkLocalFileOrFs(c, frontendDist, singleton.Conf.UserTemplate, "index.html", fallbackStatusCode) {
|
||||||
c.JSON(http.StatusNotFound, newErrorResponse(errors.New("404 Not Found")))
|
c.JSON(http.StatusNotFound, newErrorResponse(errors.New("404 Not Found")))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,269 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"golang.org/x/crypto/bcrypt"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestGetProfile_RedactsPasswordHash(t *testing.T) {
|
||||||
|
defer setupTenancyTest(t)()
|
||||||
|
require.NoError(t, singleton.DB.AutoMigrate(&model.Oauth2Bind{}))
|
||||||
|
|
||||||
|
passwordHash := "$2a$10$profilePasswordHashMustNotLeak012345678901"
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/profile", nil)
|
||||||
|
c.Set(model.CtxKeyAuthorizedUser, &model.User{
|
||||||
|
Common: model.Common{ID: 10},
|
||||||
|
Username: "alice",
|
||||||
|
Password: passwordHash,
|
||||||
|
Role: model.RoleMember,
|
||||||
|
})
|
||||||
|
|
||||||
|
commonHandler(getProfile)(c)
|
||||||
|
|
||||||
|
require.Equal(t, http.StatusOK, w.Code)
|
||||||
|
require.NotContains(t, w.Body.String(), passwordHash)
|
||||||
|
|
||||||
|
var body map[string]any
|
||||||
|
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body))
|
||||||
|
require.Equal(t, true, body["success"])
|
||||||
|
data, ok := body["data"].(map[string]any)
|
||||||
|
require.True(t, ok)
|
||||||
|
require.Equal(t, "alice", data["username"])
|
||||||
|
require.NotContains(t, data, "password")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateProfile_UsesStoredPasswordHashAfterRedaction(t *testing.T) {
|
||||||
|
defer setupTenancyTest(t)()
|
||||||
|
require.NoError(t, singleton.DB.AutoMigrate(&model.Oauth2Bind{}, &model.JWTSession{}))
|
||||||
|
|
||||||
|
originalUserInfoMap := singleton.UserInfoMap
|
||||||
|
originalAgentSecretToUserID := singleton.AgentSecretToUserId
|
||||||
|
singleton.UserInfoMap = make(map[uint64]model.UserInfo)
|
||||||
|
singleton.AgentSecretToUserId = make(map[string]uint64)
|
||||||
|
defer func() {
|
||||||
|
singleton.UserInfoMap = originalUserInfoMap
|
||||||
|
singleton.AgentSecretToUserId = originalAgentSecretToUserID
|
||||||
|
}()
|
||||||
|
|
||||||
|
oldHash, err := bcrypt.GenerateFromPassword([]byte("old-password"), bcrypt.MinCost)
|
||||||
|
require.NoError(t, err)
|
||||||
|
user := model.User{
|
||||||
|
Common: model.Common{ID: 10},
|
||||||
|
Username: "alice",
|
||||||
|
Password: string(oldHash),
|
||||||
|
Role: model.RoleMember,
|
||||||
|
AgentSecret: "profile-update-agent-secret",
|
||||||
|
}
|
||||||
|
require.NoError(t, singleton.DB.Create(&user).Error)
|
||||||
|
|
||||||
|
_, err = updateProfile(profileUpdateContext(t, user, model.ProfileForm{
|
||||||
|
OriginalPassword: "wrong-password",
|
||||||
|
NewUsername: "mallory",
|
||||||
|
NewPassword: "new-password",
|
||||||
|
}))
|
||||||
|
require.Error(t, err)
|
||||||
|
|
||||||
|
var unchanged model.User
|
||||||
|
require.NoError(t, singleton.DB.First(&unchanged, user.ID).Error)
|
||||||
|
require.Equal(t, "alice", unchanged.Username)
|
||||||
|
require.Equal(t, string(oldHash), unchanged.Password)
|
||||||
|
|
||||||
|
_, err = updateProfile(profileUpdateContext(t, user, model.ProfileForm{
|
||||||
|
OriginalPassword: "old-password",
|
||||||
|
NewUsername: "alice-renamed",
|
||||||
|
NewPassword: "new-password",
|
||||||
|
}))
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
var after model.User
|
||||||
|
require.NoError(t, singleton.DB.First(&after, user.ID).Error)
|
||||||
|
require.Equal(t, "alice-renamed", after.Username)
|
||||||
|
require.NoError(t, bcrypt.CompareHashAndPassword([]byte(after.Password), []byte("new-password")))
|
||||||
|
}
|
||||||
|
|
||||||
|
func profileUpdateContext(t *testing.T, user model.User, form model.ProfileForm) *gin.Context {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
body, err := json.Marshal(form)
|
||||||
|
require.NoError(t, err)
|
||||||
|
c := ctxAs(user.ID, user.Role)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/profile", bytes.NewReader(body))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
c.Set(model.CtxKeyAuthorizedUser, &user)
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListDDNS_RedactsCredentials(t *testing.T) {
|
||||||
|
defer setupTenancyTest(t)()
|
||||||
|
|
||||||
|
p := model.DDNSProfile{
|
||||||
|
Common: model.Common{UserID: 10},
|
||||||
|
Name: "cf",
|
||||||
|
Provider: "cloudflare",
|
||||||
|
AccessID: "id",
|
||||||
|
AccessSecret: "super-secret-token",
|
||||||
|
WebhookHeaders: `{"Authorization":"Bearer xxx"}`,
|
||||||
|
}
|
||||||
|
require.NoError(t, singleton.DB.Create(&p).Error)
|
||||||
|
singleton.DDNSShared.InsertForTest(&p)
|
||||||
|
|
||||||
|
c := ctxAs(10, model.RoleAdmin)
|
||||||
|
out, err := listDDNS(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, out, 1)
|
||||||
|
require.Empty(t, out[0].AccessSecret, "access_secret must be redacted in list response")
|
||||||
|
require.Empty(t, out[0].WebhookHeaders, "webhook_headers must be redacted in list response")
|
||||||
|
require.Equal(t, "id", out[0].AccessID, "non-secret fields must be preserved")
|
||||||
|
|
||||||
|
var stored model.DDNSProfile
|
||||||
|
require.NoError(t, singleton.DB.First(&stored, p.ID).Error)
|
||||||
|
require.Equal(t, "super-secret-token", stored.AccessSecret, "redaction must not mutate stored data")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListNotification_RedactsCredentials(t *testing.T) {
|
||||||
|
defer setupTenancyTest(t)()
|
||||||
|
|
||||||
|
n := model.Notification{
|
||||||
|
Common: model.Common{UserID: 10},
|
||||||
|
Name: "slack",
|
||||||
|
URL: "https://hooks.slack.com/services/T/B/secret",
|
||||||
|
RequestHeader: `{"Authorization":"Bearer xxx"}`,
|
||||||
|
RequestBody: `{"text":"#NEZHA#"}`,
|
||||||
|
}
|
||||||
|
require.NoError(t, singleton.DB.Create(&n).Error)
|
||||||
|
singleton.NotificationShared.InsertForTest(&n)
|
||||||
|
|
||||||
|
c := ctxAs(10, model.RoleAdmin)
|
||||||
|
out, err := listNotification(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, out, 1)
|
||||||
|
require.Empty(t, out[0].URL, "url must be redacted in list response")
|
||||||
|
require.Empty(t, out[0].RequestHeader, "request_header must be redacted in list response")
|
||||||
|
require.Empty(t, out[0].RequestBody, "request_body must be redacted in list response")
|
||||||
|
require.Equal(t, "slack", out[0].Name, "non-secret fields must be preserved")
|
||||||
|
|
||||||
|
var stored model.Notification
|
||||||
|
require.NoError(t, singleton.DB.First(&stored, n.ID).Error)
|
||||||
|
require.Equal(t, "https://hooks.slack.com/services/T/B/secret", stored.URL,
|
||||||
|
"redaction must not mutate stored data")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateDDNS_EmptySecretPreservesStored(t *testing.T) {
|
||||||
|
defer setupTenancyTest(t)()
|
||||||
|
|
||||||
|
existing := model.DDNSProfile{
|
||||||
|
Common: model.Common{UserID: 10},
|
||||||
|
Name: "cf",
|
||||||
|
Provider: "webhook",
|
||||||
|
AccessID: "id",
|
||||||
|
AccessSecret: "keep-me",
|
||||||
|
WebhookURL: "http://127.0.0.1/",
|
||||||
|
WebhookMethod: 1,
|
||||||
|
WebhookHeaders: `{"X-Token":"keep-header"}`,
|
||||||
|
}
|
||||||
|
require.NoError(t, singleton.DB.Create(&existing).Error)
|
||||||
|
singleton.DDNSShared.InsertForTest(&existing)
|
||||||
|
|
||||||
|
c := ctxAsMemberWithBody(10, map[string]any{
|
||||||
|
"name": "cf-renamed",
|
||||||
|
"provider": "webhook",
|
||||||
|
"access_id": "id",
|
||||||
|
"access_secret": "",
|
||||||
|
"webhook_url": "http://127.0.0.1/",
|
||||||
|
"webhook_method": 1,
|
||||||
|
"webhook_request_type": 1,
|
||||||
|
"webhook_headers": "",
|
||||||
|
"max_retries": 3,
|
||||||
|
})
|
||||||
|
c.Params = gin.Params{{Key: "id", Value: itoa(existing.ID)}}
|
||||||
|
_, err := updateDDNS(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
var after model.DDNSProfile
|
||||||
|
require.NoError(t, singleton.DB.First(&after, existing.ID).Error)
|
||||||
|
require.Equal(t, "cf-renamed", after.Name, "non-secret edits must apply")
|
||||||
|
require.Equal(t, "keep-me", after.AccessSecret, "empty submitted secret must preserve stored value")
|
||||||
|
require.Equal(t, `{"X-Token":"keep-header"}`, after.WebhookHeaders,
|
||||||
|
"empty submitted headers must preserve stored value")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateDDNS_NonEmptySecretOverwrites(t *testing.T) {
|
||||||
|
defer setupTenancyTest(t)()
|
||||||
|
|
||||||
|
existing := model.DDNSProfile{
|
||||||
|
Common: model.Common{UserID: 10},
|
||||||
|
Name: "cf",
|
||||||
|
Provider: "webhook",
|
||||||
|
AccessSecret: "old",
|
||||||
|
WebhookURL: "http://127.0.0.1/",
|
||||||
|
WebhookMethod: 1,
|
||||||
|
}
|
||||||
|
require.NoError(t, singleton.DB.Create(&existing).Error)
|
||||||
|
singleton.DDNSShared.InsertForTest(&existing)
|
||||||
|
|
||||||
|
c := ctxAsMemberWithBody(10, map[string]any{
|
||||||
|
"name": "cf",
|
||||||
|
"provider": "webhook",
|
||||||
|
"access_secret": "new-secret",
|
||||||
|
"webhook_url": "http://127.0.0.1/",
|
||||||
|
"webhook_method": 1,
|
||||||
|
"webhook_request_type": 1,
|
||||||
|
"max_retries": 3,
|
||||||
|
})
|
||||||
|
c.Params = gin.Params{{Key: "id", Value: itoa(existing.ID)}}
|
||||||
|
_, err := updateDDNS(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
var after model.DDNSProfile
|
||||||
|
require.NoError(t, singleton.DB.First(&after, existing.ID).Error)
|
||||||
|
require.Equal(t, "new-secret", after.AccessSecret, "non-empty submitted secret must overwrite")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateNotification_EmptyFieldsPreserveStored(t *testing.T) {
|
||||||
|
defer setupTenancyTest(t)()
|
||||||
|
|
||||||
|
existing := model.Notification{
|
||||||
|
Common: model.Common{UserID: 10},
|
||||||
|
Name: "slack",
|
||||||
|
URL: "https://hooks.slack.com/services/keep",
|
||||||
|
RequestMethod: model.NotificationRequestMethodGET,
|
||||||
|
RequestType: model.NotificationRequestTypeJSON,
|
||||||
|
RequestHeader: `{"Authorization":"keep"}`,
|
||||||
|
RequestBody: "",
|
||||||
|
}
|
||||||
|
require.NoError(t, singleton.DB.Create(&existing).Error)
|
||||||
|
singleton.NotificationShared.InsertForTest(&existing)
|
||||||
|
|
||||||
|
c := ctxAsMemberWithBody(10, map[string]any{
|
||||||
|
"name": "slack-renamed",
|
||||||
|
"url": "",
|
||||||
|
"request_method": model.NotificationRequestMethodGET,
|
||||||
|
"request_type": model.NotificationRequestTypeJSON,
|
||||||
|
"request_header": "",
|
||||||
|
"request_body": "",
|
||||||
|
"skip_check": true,
|
||||||
|
})
|
||||||
|
c.Params = gin.Params{{Key: "id", Value: itoa(existing.ID)}}
|
||||||
|
_, err := updateNotification(c)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
var after model.Notification
|
||||||
|
require.NoError(t, singleton.DB.First(&after, existing.ID).Error)
|
||||||
|
require.Equal(t, "slack-renamed", after.Name, "non-secret edits must apply")
|
||||||
|
require.Equal(t, "https://hooks.slack.com/services/keep", after.URL,
|
||||||
|
"empty submitted url must preserve stored value")
|
||||||
|
require.Equal(t, `{"Authorization":"keep"}`, after.RequestHeader,
|
||||||
|
"empty submitted request_header must preserve stored value")
|
||||||
|
}
|
||||||
@@ -50,10 +50,22 @@ func createCron(c *gin.Context) (uint64, error) {
|
|||||||
return 0, err
|
return 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
if !singleton.ServerShared.CheckPermission(c, slices.Values(cf.Servers)) {
|
if !isValidCronCover(cf.Cover) {
|
||||||
return 0, singleton.Localizer.ErrorT("permission denied")
|
return 0, singleton.Localizer.ErrorT("permission denied")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if err := checkCronServerListPermission(c, cf.Cover, cf.Servers, getUid(c)); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := rejectImplicitCoverForLimitedPAT(c, cf.Cover, cf.Servers); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := assertOwnsNotificationGroup(c, cf.NotificationGroupID); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
|
||||||
cr.UserID = getUid(c)
|
cr.UserID = getUid(c)
|
||||||
cr.TaskType = cf.TaskType
|
cr.TaskType = cf.TaskType
|
||||||
cr.Name = cf.Name
|
cr.Name = cf.Name
|
||||||
@@ -68,7 +80,6 @@ func createCron(c *gin.Context) (uint64, error) {
|
|||||||
return 0, singleton.Localizer.ErrorT("scheduled tasks cannot be triggered by alarms")
|
return 0, singleton.Localizer.ErrorT("scheduled tasks cannot be triggered by alarms")
|
||||||
}
|
}
|
||||||
|
|
||||||
// 对于计划任务类型,需要更新CronJob
|
|
||||||
var err error
|
var err error
|
||||||
if cf.TaskType == model.CronTypeCronTask {
|
if cf.TaskType == model.CronTypeCronTask {
|
||||||
if cr.CronJobID, err = singleton.CronShared.AddFunc(cr.Scheduler, singleton.CronTrigger(&cr)); err != nil {
|
if cr.CronJobID, err = singleton.CronShared.AddFunc(cr.Scheduler, singleton.CronTrigger(&cr)); err != nil {
|
||||||
@@ -108,8 +119,8 @@ func updateCron(c *gin.Context) (any, error) {
|
|||||||
return 0, err
|
return 0, err
|
||||||
}
|
}
|
||||||
|
|
||||||
if !singleton.ServerShared.CheckPermission(c, slices.Values(cf.Servers)) {
|
if !isValidCronCover(cf.Cover) {
|
||||||
return 0, singleton.Localizer.ErrorT("permission denied")
|
return nil, singleton.Localizer.ErrorT("permission denied")
|
||||||
}
|
}
|
||||||
|
|
||||||
var cr model.Cron
|
var cr model.Cron
|
||||||
@@ -121,6 +132,18 @@ func updateCron(c *gin.Context) (any, error) {
|
|||||||
return nil, singleton.Localizer.ErrorT("permission denied")
|
return nil, singleton.Localizer.ErrorT("permission denied")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if err := checkCronServerListPermission(c, cf.Cover, cf.Servers, cr.GetUserID()); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := rejectImplicitCoverForLimitedPATWithOwner(c, cf.Cover, cf.Servers, cr.GetUserID()); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := assertOwnsNotificationGroup(c, cf.NotificationGroupID); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
cr.TaskType = cf.TaskType
|
cr.TaskType = cf.TaskType
|
||||||
cr.Name = cf.Name
|
cr.Name = cf.Name
|
||||||
cr.Scheduler = cf.Scheduler
|
cr.Scheduler = cf.Scheduler
|
||||||
@@ -159,7 +182,7 @@ func updateCron(c *gin.Context) (any, error) {
|
|||||||
// @param id path uint true "Task ID"
|
// @param id path uint true "Task ID"
|
||||||
// @Produce json
|
// @Produce json
|
||||||
// @Success 200 {object} model.CommonResponse[any]
|
// @Success 200 {object} model.CommonResponse[any]
|
||||||
// @Router /cron/{id}/manual [get]
|
// @Router /cron/{id}/manual [post]
|
||||||
func manualTriggerCron(c *gin.Context) (any, error) {
|
func manualTriggerCron(c *gin.Context) (any, error) {
|
||||||
idStr := c.Param("id")
|
idStr := c.Param("id")
|
||||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
id, err := strconv.ParseUint(idStr, 10, 64)
|
||||||
@@ -176,6 +199,14 @@ func manualTriggerCron(c *gin.Context) (any, error) {
|
|||||||
return nil, singleton.Localizer.ErrorT("permission denied")
|
return nil, singleton.Localizer.ErrorT("permission denied")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 运行时回放写侧 rejectImplicitCoverForLimitedPAT* 同一条 PAT 收口:
|
||||||
|
// 历史脏数据 / 旁路写入的 cron 仍可能携带「CronCoverAll + 不充分 deny-list」
|
||||||
|
// 的配置;CronTrigger 没有 PAT 上下文,manualTrigger 这里是唯一阻止
|
||||||
|
// 受限 PAT 触发 fan-out 到白名单外 owner servers 的同步入口。
|
||||||
|
if err := enforcePATCronDispatchScope(c, cr); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
singleton.ManualTrigger(cr)
|
singleton.ManualTrigger(cr)
|
||||||
return nil, nil
|
return nil, nil
|
||||||
}
|
}
|
||||||
@@ -201,6 +232,19 @@ func batchDeleteCron(c *gin.Context) (any, error) {
|
|||||||
return nil, singleton.Localizer.ErrorT("permission denied")
|
return nil, singleton.Localizer.ErrorT("permission denied")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 与 manualTriggerCron 对称:删除会改变 fan-out 范围本身,受限 PAT 不
|
||||||
|
// 应通过删除一个白名单内的「掩护」cron 间接放大对白名单外 owner servers
|
||||||
|
// 的影响。回放同一条 cover-fanout 收口。
|
||||||
|
for _, id := range cr {
|
||||||
|
existing, ok := singleton.CronShared.Get(id)
|
||||||
|
if !ok || existing == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if err := enforcePATCronDispatchScope(c, existing); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if err := singleton.DB.Unscoped().Delete(&model.Cron{}, "id in (?)", cr).Error; err != nil {
|
if err := singleton.DB.Unscoped().Delete(&model.Cron{}, "id in (?)", cr).Error; err != nil {
|
||||||
return nil, newGormError("%v", err)
|
return nil, newGormError("%v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,54 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
// C2 regression: writes must reject unknown Cover values so dirty configs
|
||||||
|
// cannot be persisted. CronTrigger has no PAT context on the periodic
|
||||||
|
// scheduler path, so unknown Cover sails past every PAT guard and dispatches
|
||||||
|
// via the default branch in CronTrigger (no CoverAll/IgnoreAll match → still
|
||||||
|
// reaches every server that passes cronCanSendToServer).
|
||||||
|
func TestIsValidCronCover_RejectsUnknown(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
cover uint8
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"CoverIgnoreAll", model.CronCoverIgnoreAll, true},
|
||||||
|
{"CoverAll", model.CronCoverAll, true},
|
||||||
|
{"CoverAlertTrigger", model.CronCoverAlertTrigger, true},
|
||||||
|
{"unknown_99", 99, false},
|
||||||
|
{"unknown_max", 255, false},
|
||||||
|
{"unknown_3", 3, false},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
if got := isValidCronCover(tc.cover); got != tc.want {
|
||||||
|
t.Fatalf("isValidCronCover(%d) = %v, want %v", tc.cover, got, tc.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIsValidServiceCover_RejectsUnknown(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
cover uint8
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"ServiceCoverAll", model.ServiceCoverAll, true},
|
||||||
|
{"ServiceCoverIgnoreAll", model.ServiceCoverIgnoreAll, true},
|
||||||
|
{"unknown_99", 99, false},
|
||||||
|
{"unknown_max", 255, false},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
if got := isValidServiceCover(tc.cover); got != tc.want {
|
||||||
|
t.Fatalf("isValidServiceCover(%d) = %v, want %v", tc.cover, got, tc.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,225 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
// 回归 cron 运行时入口 (manualTriggerCron / batchDeleteCron) 上的 PAT
|
||||||
|
// cover-fanout 收口。配合 permissions_cover_fanout_test.go 的底座单测,
|
||||||
|
// 形成「共享底座 ↔ 资源专用入口」两层钉子,任何后续重构(例如把 guard
|
||||||
|
// 拆出 controller、把 cover 模式合并/拆分)都必须保留:
|
||||||
|
// - 受限 PAT 不能通过 manualTrigger 触发一个 deny-list 不充分的
|
||||||
|
// CronCoverAll → 否则 CronTrigger fan out 到白名单外 owner servers。
|
||||||
|
// - 受限 PAT 不能通过 batchDelete 删除同样形态的 cron → 否则相当于
|
||||||
|
// 间接操作白名单外 owner servers 的调度策略。
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strconv"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/patrickmn/go-cache"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/i18n"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// setupCronDispatchPATFixture 与 setupCoverPATFixture 同一拓扑:alice
|
||||||
|
// (uid=100) 拥有 server 1 / server 2;下游测试给出 PAT server_ids=[1]。
|
||||||
|
func setupCronDispatchPATFixture(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
originalDB := singleton.DB
|
||||||
|
originalCache := singleton.Cache
|
||||||
|
originalLoc := singleton.Loc
|
||||||
|
originalLocalizer := singleton.Localizer
|
||||||
|
originalCron := singleton.CronShared
|
||||||
|
originalServer := singleton.ServerShared
|
||||||
|
originalUserInfo := singleton.UserInfoMap
|
||||||
|
originalNotification := singleton.NotificationShared
|
||||||
|
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
sqlDB, err := db.DB()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{}, &model.User{}, &model.NotificationGroup{}, &model.Notification{}))
|
||||||
|
|
||||||
|
singleton.DB = db
|
||||||
|
singleton.Loc = time.UTC
|
||||||
|
singleton.Cache = cache.New(time.Minute, time.Minute)
|
||||||
|
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
|
||||||
|
// CronTrigger 的 fan-out 路径会在 server 没接入 task-stream 时调
|
||||||
|
// NotificationShared.SendNotification 上报「离线」;这里给出一个空的
|
||||||
|
// notification class,避免 nil deref。本测试不验证通知内容。
|
||||||
|
singleton.NotificationShared = singleton.NewNotificationClass()
|
||||||
|
|
||||||
|
sc := singleton.NewEmptyServerClassForTest()
|
||||||
|
for _, id := range []uint64{1, 2} {
|
||||||
|
s := &model.Server{}
|
||||||
|
s.ID = id
|
||||||
|
s.SetUserID(100)
|
||||||
|
sc.InsertForTest(s)
|
||||||
|
}
|
||||||
|
singleton.ServerShared = sc
|
||||||
|
singleton.CronShared = singleton.NewCronClass()
|
||||||
|
|
||||||
|
singleton.UserLock.Lock()
|
||||||
|
singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}}
|
||||||
|
singleton.UserLock.Unlock()
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
// Test-owned cron jobs must be joined before restoring process-global singleton dependencies.
|
||||||
|
singleton.CronShared.Close()
|
||||||
|
_ = sqlDB.Close()
|
||||||
|
singleton.DB = originalDB
|
||||||
|
singleton.Cache = originalCache
|
||||||
|
singleton.Loc = originalLoc
|
||||||
|
singleton.Localizer = originalLocalizer
|
||||||
|
singleton.CronShared = originalCron
|
||||||
|
singleton.ServerShared = originalServer
|
||||||
|
singleton.NotificationShared = originalNotification
|
||||||
|
singleton.UserLock.Lock()
|
||||||
|
singleton.UserInfoMap = originalUserInfo
|
||||||
|
singleton.UserLock.Unlock()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func insertCronForDispatchTest(t *testing.T, cover uint8, servers []uint64) uint64 {
|
||||||
|
t.Helper()
|
||||||
|
cr := &model.Cron{
|
||||||
|
Common: model.Common{UserID: 100},
|
||||||
|
Name: "dispatch-fixture",
|
||||||
|
TaskType: model.CronTypeCronTask,
|
||||||
|
Command: "echo dispatch",
|
||||||
|
Servers: servers,
|
||||||
|
Cover: cover,
|
||||||
|
}
|
||||||
|
require.NoError(t, singleton.DB.Create(cr).Error)
|
||||||
|
singleton.CronShared.Update(cr)
|
||||||
|
return cr.ID
|
||||||
|
}
|
||||||
|
|
||||||
|
func newCronDispatchRouter(t *testing.T, tok *model.APIToken) *gin.Engine {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.Use(func(c *gin.Context) {
|
||||||
|
setAuthUser(c, 100, model.RoleMember)
|
||||||
|
if tok != nil {
|
||||||
|
c.Set(model.CtxKeyAPIToken, tok)
|
||||||
|
c.Set(apiTokenCtxKey, tok)
|
||||||
|
}
|
||||||
|
c.Next()
|
||||||
|
})
|
||||||
|
r.POST("/api/v1/cron/:id/manual", commonHandler(manualTriggerCron))
|
||||||
|
r.POST("/api/v1/batch-delete/cron", commonHandler(batchDeleteCron))
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestManualTriggerCron_RejectsCoverAllWithInsufficientDenyForLimitedPAT(t *testing.T) {
|
||||||
|
setupCronDispatchPATFixture(t)
|
||||||
|
cronID := insertCronForDispatchTest(t, model.CronCoverAll, []uint64{1})
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 17, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := newCronDispatchRouter(t, tok)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost,
|
||||||
|
"/api/v1/cron/"+strconv.FormatUint(cronID, 10)+"/manual", nil)
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.False(t, success,
|
||||||
|
"PAT [1] must NOT manually trigger a CronCoverAll cron whose deny-list does not cover owner server 2; CronTrigger would fan out to it")
|
||||||
|
assert.Contains(t, errMsg, "permission denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestManualTriggerCron_AllowsCoverAllWhenDenyCoversNonWhitelisted(t *testing.T) {
|
||||||
|
setupCronDispatchPATFixture(t)
|
||||||
|
cronID := insertCronForDispatchTest(t, model.CronCoverAll, []uint64{2})
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 18, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := newCronDispatchRouter(t, tok)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost,
|
||||||
|
"/api/v1/cron/"+strconv.FormatUint(cronID, 10)+"/manual", nil)
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.True(t, success,
|
||||||
|
"CronCoverAll whose deny-list covers every non-whitelisted owner server must remain triggerable: error=%s", errMsg)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestManualTriggerCron_AllowsCoverIgnoreAllInsideWhitelist(t *testing.T) {
|
||||||
|
setupCronDispatchPATFixture(t)
|
||||||
|
cronID := insertCronForDispatchTest(t, model.CronCoverIgnoreAll, []uint64{1})
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 19, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := newCronDispatchRouter(t, tok)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost,
|
||||||
|
"/api/v1/cron/"+strconv.FormatUint(cronID, 10)+"/manual", nil)
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.True(t, success,
|
||||||
|
"CronCoverIgnoreAll allow-list inside PAT whitelist must trigger normally: error=%s", errMsg)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBatchDeleteCron_RejectsCoverAllWithInsufficientDenyForLimitedPAT(t *testing.T) {
|
||||||
|
setupCronDispatchPATFixture(t)
|
||||||
|
cronID := insertCronForDispatchTest(t, model.CronCoverAll, []uint64{1})
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 21, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := newCronDispatchRouter(t, tok)
|
||||||
|
body, _ := json.Marshal([]uint64{cronID})
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/batch-delete/cron", bytes.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.False(t, success,
|
||||||
|
"PAT [1] must NOT batch-delete a CronCoverAll cron whose deny-list does not cover owner server 2")
|
||||||
|
assert.Contains(t, errMsg, "permission denied")
|
||||||
|
|
||||||
|
var rows []model.Cron
|
||||||
|
require.NoError(t, singleton.DB.Find(&rows).Error)
|
||||||
|
assert.Len(t, rows, 1, "cron row must still exist when the delete call is rejected")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBatchDeleteCron_AllowsCoverAllWhenDenyCoversNonWhitelisted(t *testing.T) {
|
||||||
|
setupCronDispatchPATFixture(t)
|
||||||
|
cronID := insertCronForDispatchTest(t, model.CronCoverAll, []uint64{2})
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 22, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := newCronDispatchRouter(t, tok)
|
||||||
|
body, _ := json.Marshal([]uint64{cronID})
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/batch-delete/cron", bytes.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.True(t, success,
|
||||||
|
"deny-list covering every non-whitelisted owner server must allow batch-delete: error=%s", errMsg)
|
||||||
|
|
||||||
|
var rows []model.Cron
|
||||||
|
require.NoError(t, singleton.DB.Find(&rows).Error)
|
||||||
|
assert.Empty(t, rows, "cron row must be deleted when the call succeeds")
|
||||||
|
}
|
||||||
@@ -0,0 +1,65 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newCronListPATRouter(t *testing.T, tok *model.APIToken) *gin.Engine {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.Use(func(c *gin.Context) {
|
||||||
|
setAuthUser(c, 100, model.RoleMember)
|
||||||
|
if tok != nil {
|
||||||
|
c.Set(model.CtxKeyAPIToken, tok)
|
||||||
|
c.Set(apiTokenCtxKey, tok)
|
||||||
|
}
|
||||||
|
c.Next()
|
||||||
|
})
|
||||||
|
r.GET("/api/v1/cron", listHandler(listCron))
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
// GET /api/v1/cron must replay the same deny-list rule the dispatch guards
|
||||||
|
// use; otherwise a stale or out-of-band-written CronCoverAll row whose
|
||||||
|
// Servers deny-list does not cover the non-whitelisted owner server still
|
||||||
|
// shows up in the limited PAT's list view.
|
||||||
|
func TestListCron_HidesCoverAllWithInsufficientDenyForLimitedPAT(t *testing.T) {
|
||||||
|
setupCronDispatchPATFixture(t)
|
||||||
|
insufficient := insertCronForDispatchTest(t, model.CronCoverAll, []uint64{1})
|
||||||
|
sufficient := insertCronForDispatchTest(t, model.CronCoverAll, []uint64{2})
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 23, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := newCronListPATRouter(t, tok)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/cron", nil)
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
var resp struct {
|
||||||
|
Success bool `json:"success"`
|
||||||
|
Error string `json:"error"`
|
||||||
|
Data []*model.Cron `json:"data"`
|
||||||
|
}
|
||||||
|
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
|
||||||
|
require.True(t, resp.Success, resp.Error)
|
||||||
|
|
||||||
|
seen := map[uint64]bool{}
|
||||||
|
for _, c := range resp.Data {
|
||||||
|
seen[c.ID] = true
|
||||||
|
}
|
||||||
|
assert.False(t, seen[insufficient],
|
||||||
|
"PAT [1] must NOT see a CronCoverAll whose deny-list does not cover owner server 2 (rows=%+v)", resp.Data)
|
||||||
|
assert.True(t, seen[sufficient],
|
||||||
|
"PAT [1] must still see a CronCoverAll whose deny-list already covers every non-whitelisted owner server")
|
||||||
|
}
|
||||||
@@ -0,0 +1,109 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/patrickmn/go-cache"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/i18n"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func setupCronManualTriggerFixture(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
originalDB := singleton.DB
|
||||||
|
originalCache := singleton.Cache
|
||||||
|
originalLoc := singleton.Loc
|
||||||
|
originalLocalizer := singleton.Localizer
|
||||||
|
originalCron := singleton.CronShared
|
||||||
|
originalServer := singleton.ServerShared
|
||||||
|
originalUserInfo := singleton.UserInfoMap
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
sqlDB, err := db.DB()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{}, &model.User{}))
|
||||||
|
|
||||||
|
singleton.DB = db
|
||||||
|
singleton.Loc = time.UTC
|
||||||
|
singleton.Cache = cache.New(time.Minute, time.Minute)
|
||||||
|
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
|
||||||
|
singleton.ServerShared = singleton.NewServerClass()
|
||||||
|
singleton.CronShared = singleton.NewCronClass()
|
||||||
|
singleton.UserLock.Lock()
|
||||||
|
singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}}
|
||||||
|
singleton.UserLock.Unlock()
|
||||||
|
|
||||||
|
cr := &model.Cron{
|
||||||
|
Common: model.Common{ID: 7, UserID: 100},
|
||||||
|
Name: "victim cron",
|
||||||
|
TaskType: model.CronTypeCronTask,
|
||||||
|
Command: "echo csrf-poc",
|
||||||
|
Cover: model.CronCoverIgnoreAll,
|
||||||
|
}
|
||||||
|
require.NoError(t, db.Create(cr).Error)
|
||||||
|
singleton.CronShared.Update(cr)
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
singleton.CronShared.Close()
|
||||||
|
_ = sqlDB.Close()
|
||||||
|
singleton.DB = originalDB
|
||||||
|
singleton.Cache = originalCache
|
||||||
|
singleton.Loc = originalLoc
|
||||||
|
singleton.Localizer = originalLocalizer
|
||||||
|
singleton.CronShared = originalCron
|
||||||
|
singleton.ServerShared = originalServer
|
||||||
|
singleton.UserLock.Lock()
|
||||||
|
singleton.UserInfoMap = originalUserInfo
|
||||||
|
singleton.UserLock.Unlock()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func newCronManualRouter() *gin.Engine {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.Use(func(c *gin.Context) {
|
||||||
|
c.Set(model.CtxKeyAuthorizedUser, &model.User{
|
||||||
|
Common: model.Common{ID: 100},
|
||||||
|
Role: model.RoleMember,
|
||||||
|
})
|
||||||
|
c.Next()
|
||||||
|
})
|
||||||
|
r.POST("/api/v1/cron/:id/manual", commonHandler(manualTriggerCron))
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCronManualTriggerRejectsCrossSiteGET(t *testing.T) {
|
||||||
|
setupCronManualTriggerFixture(t)
|
||||||
|
r := newCronManualRouter()
|
||||||
|
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/cron/7/manual", nil)
|
||||||
|
req.Header.Set("Origin", "https://attacker.example")
|
||||||
|
req.Header.Set("Sec-Fetch-Site", "cross-site")
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
assert.Equal(t, http.StatusNotFound, w.Code, "manual trigger must reject cross-site GET — the route is POST-only after the CSRF fix")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCronManualTriggerAcceptsSameSitePOST(t *testing.T) {
|
||||||
|
setupCronManualTriggerFixture(t)
|
||||||
|
r := newCronManualRouter()
|
||||||
|
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/cron/7/manual", nil)
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.True(t, success, "owner POST must succeed: error=%q", errMsg)
|
||||||
|
}
|
||||||
@@ -0,0 +1,165 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strconv"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/patrickmn/go-cache"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/i18n"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func setupCronPATWhitelistFixture(t *testing.T) (cronID7, cronID8 uint64) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
originalDB := singleton.DB
|
||||||
|
originalCache := singleton.Cache
|
||||||
|
originalLoc := singleton.Loc
|
||||||
|
originalLocalizer := singleton.Localizer
|
||||||
|
originalCron := singleton.CronShared
|
||||||
|
originalServer := singleton.ServerShared
|
||||||
|
originalUserInfo := singleton.UserInfoMap
|
||||||
|
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
sqlDB, err := db.DB()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{}, &model.User{}))
|
||||||
|
|
||||||
|
singleton.DB = db
|
||||||
|
singleton.Loc = time.UTC
|
||||||
|
singleton.Cache = cache.New(time.Minute, time.Minute)
|
||||||
|
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
|
||||||
|
singleton.ServerShared = singleton.NewServerClass()
|
||||||
|
singleton.CronShared = singleton.NewCronClass()
|
||||||
|
singleton.UserLock.Lock()
|
||||||
|
singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}}
|
||||||
|
singleton.UserLock.Unlock()
|
||||||
|
|
||||||
|
cr7 := &model.Cron{
|
||||||
|
Common: model.Common{UserID: 100},
|
||||||
|
Name: "cron-on-server-1",
|
||||||
|
TaskType: model.CronTypeCronTask,
|
||||||
|
Command: "echo s1",
|
||||||
|
Servers: []uint64{1},
|
||||||
|
Cover: model.CronCoverIgnoreAll,
|
||||||
|
}
|
||||||
|
cr8 := &model.Cron{
|
||||||
|
Common: model.Common{UserID: 100},
|
||||||
|
Name: "cron-on-server-2",
|
||||||
|
TaskType: model.CronTypeCronTask,
|
||||||
|
Command: "echo s2",
|
||||||
|
Servers: []uint64{2},
|
||||||
|
Cover: model.CronCoverIgnoreAll,
|
||||||
|
}
|
||||||
|
require.NoError(t, db.Create(cr7).Error)
|
||||||
|
require.NoError(t, db.Create(cr8).Error)
|
||||||
|
singleton.CronShared.Update(cr7)
|
||||||
|
singleton.CronShared.Update(cr8)
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
singleton.CronShared.Close()
|
||||||
|
_ = sqlDB.Close()
|
||||||
|
singleton.DB = originalDB
|
||||||
|
singleton.Cache = originalCache
|
||||||
|
singleton.Loc = originalLoc
|
||||||
|
singleton.Localizer = originalLocalizer
|
||||||
|
singleton.CronShared = originalCron
|
||||||
|
singleton.ServerShared = originalServer
|
||||||
|
singleton.UserLock.Lock()
|
||||||
|
singleton.UserInfoMap = originalUserInfo
|
||||||
|
singleton.UserLock.Unlock()
|
||||||
|
})
|
||||||
|
|
||||||
|
return cr7.ID, cr8.ID
|
||||||
|
}
|
||||||
|
|
||||||
|
func newCronPATRouter(tok *model.APIToken) *gin.Engine {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.Use(func(c *gin.Context) {
|
||||||
|
setAuthUser(c, 100, model.RoleMember)
|
||||||
|
if tok != nil {
|
||||||
|
c.Set(model.CtxKeyAPIToken, tok)
|
||||||
|
c.Set(apiTokenCtxKey, tok)
|
||||||
|
}
|
||||||
|
c.Next()
|
||||||
|
})
|
||||||
|
r.POST("/api/v1/cron/:id/manual", commonHandler(manualTriggerCron))
|
||||||
|
r.GET("/api/v1/cron", listHandler(listCron))
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCronManualTrigger_DeniesServerOutsidePATWhitelist(t *testing.T) {
|
||||||
|
_, cron8 := setupCronPATWhitelistFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 17, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := newCronPATRouter(tok)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost,
|
||||||
|
"/api/v1/cron/"+strconv.FormatUint(cron8, 10)+"/manual", nil)
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.False(t, success,
|
||||||
|
"PAT whitelist [1] must not allow triggering a cron bound to server 2")
|
||||||
|
assert.Contains(t, errMsg, "permission denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCronManualTrigger_AllowsServerInsidePATWhitelist(t *testing.T) {
|
||||||
|
cron7, _ := setupCronPATWhitelistFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 17, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := newCronPATRouter(tok)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost,
|
||||||
|
"/api/v1/cron/"+strconv.FormatUint(cron7, 10)+"/manual", nil)
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.True(t, success,
|
||||||
|
"PAT whitelist [1] must still allow triggering a cron bound to server 1: error=%s", errMsg)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestListCron_HidesRowsForServersOutsidePATWhitelist(t *testing.T) {
|
||||||
|
cron7, cron8 := setupCronPATWhitelistFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 17, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := newCronPATRouter(tok)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/api/v1/cron", nil)
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
var resp struct {
|
||||||
|
Success bool `json:"success"`
|
||||||
|
Error string `json:"error"`
|
||||||
|
Data []*model.Cron `json:"data"`
|
||||||
|
}
|
||||||
|
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
|
||||||
|
assert.True(t, resp.Success, resp.Error)
|
||||||
|
|
||||||
|
seen := map[uint64]bool{}
|
||||||
|
for _, c := range resp.Data {
|
||||||
|
seen[c.ID] = true
|
||||||
|
}
|
||||||
|
assert.True(t, seen[cron7], "cron bound to whitelisted server 1 must remain visible")
|
||||||
|
assert.False(t, seen[cron8],
|
||||||
|
"cron bound to non-whitelisted server 2 must be hidden from PAT view (rows=%+v)", resp.Data)
|
||||||
|
}
|
||||||
@@ -0,0 +1,349 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
// Regression tests for the implicit-cover PAT bypass classes.
|
||||||
|
//
|
||||||
|
// Background: ServerShared.CheckPermission iterates an idList and returns true
|
||||||
|
// for an empty list — it can only veto explicit IDs. createCron / createService
|
||||||
|
// both pipe cf.Servers (cron) and ss.SkipServers (service) through that helper.
|
||||||
|
// But under cover=CronCoverAll the cron's Servers slice is a *deny list* (and
|
||||||
|
// empty → fan out to every server owned by the user); under cover=ServiceCoverAll
|
||||||
|
// the service's SkipServers map is the equivalent deny set. A PAT scoped to
|
||||||
|
// server_ids=[1] can therefore craft a "cover all, deny none" config and force
|
||||||
|
// dashboard to dispatch cron commands / service probes to servers outside the
|
||||||
|
// PAT whitelist.
|
||||||
|
//
|
||||||
|
// These tests are deliberately end-to-end through commonHandler so a future
|
||||||
|
// refactor that moves the guard to a different layer still has to satisfy the
|
||||||
|
// "PAT can't escape its whitelist via cover semantics" invariant.
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/patrickmn/go-cache"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/i18n"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// setupCoverPATFixture builds a member-owned, two-server universe.
|
||||||
|
// alice (uid=100) owns server 1 and server 2. The caller PAT below will be
|
||||||
|
// scoped to server_ids=[1] only, so cover-all configs that fan out to
|
||||||
|
// server 2 must be rejected at the create/update boundary.
|
||||||
|
func setupCoverPATFixture(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
originalDB := singleton.DB
|
||||||
|
originalCache := singleton.Cache
|
||||||
|
originalLoc := singleton.Loc
|
||||||
|
originalLocalizer := singleton.Localizer
|
||||||
|
originalCron := singleton.CronShared
|
||||||
|
originalServer := singleton.ServerShared
|
||||||
|
originalUserInfo := singleton.UserInfoMap
|
||||||
|
originalNotification := singleton.NotificationShared
|
||||||
|
originalSentinel := singleton.ServiceSentinelShared
|
||||||
|
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
sqlDB, err := db.DB()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{}, &model.User{}, &model.Service{}, &model.NotificationGroup{}, &model.ServiceHistory{}))
|
||||||
|
|
||||||
|
singleton.DB = db
|
||||||
|
singleton.Loc = time.UTC
|
||||||
|
singleton.Cache = cache.New(time.Minute, time.Minute)
|
||||||
|
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
|
||||||
|
singleton.NotificationShared = singleton.NewEmptyNotificationClassForTest()
|
||||||
|
sc := singleton.NewEmptyServerClassForTest()
|
||||||
|
singleton.ServerShared = sc
|
||||||
|
singleton.CronShared = singleton.NewCronClass()
|
||||||
|
|
||||||
|
sentinel, err := singleton.NewServiceSentinel(make(chan *model.Service, 4))
|
||||||
|
require.NoError(t, err)
|
||||||
|
singleton.ServiceSentinelShared = sentinel
|
||||||
|
|
||||||
|
for _, id := range []uint64{1, 2} {
|
||||||
|
s := &model.Server{}
|
||||||
|
s.ID = id
|
||||||
|
s.SetUserID(100)
|
||||||
|
sc.InsertForTest(s)
|
||||||
|
}
|
||||||
|
singleton.ServerShared = sc
|
||||||
|
|
||||||
|
singleton.UserLock.Lock()
|
||||||
|
singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}}
|
||||||
|
singleton.UserLock.Unlock()
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
// Background components must be joined before restoring process globals.
|
||||||
|
sentinel.Close()
|
||||||
|
singleton.CronShared.Close()
|
||||||
|
_ = sqlDB.Close()
|
||||||
|
singleton.DB = originalDB
|
||||||
|
singleton.Cache = originalCache
|
||||||
|
singleton.Loc = originalLoc
|
||||||
|
singleton.Localizer = originalLocalizer
|
||||||
|
singleton.CronShared = originalCron
|
||||||
|
singleton.ServerShared = originalServer
|
||||||
|
singleton.NotificationShared = originalNotification
|
||||||
|
singleton.ServiceSentinelShared = originalSentinel
|
||||||
|
singleton.UserLock.Lock()
|
||||||
|
singleton.UserInfoMap = originalUserInfo
|
||||||
|
singleton.UserLock.Unlock()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func coverPATRouter(t *testing.T, tok *model.APIToken, handler func(*gin.Context)) *gin.Engine {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.Use(func(c *gin.Context) {
|
||||||
|
setAuthUser(c, 100, model.RoleMember)
|
||||||
|
if tok != nil {
|
||||||
|
c.Set(model.CtxKeyAPIToken, tok)
|
||||||
|
c.Set(apiTokenCtxKey, tok)
|
||||||
|
}
|
||||||
|
c.Next()
|
||||||
|
})
|
||||||
|
r.POST("/api/v1/cron", handler)
|
||||||
|
r.POST("/api/v1/service", handler)
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateCron_RejectsCoverAllForServerLimitedPAT(t *testing.T) {
|
||||||
|
setupCoverPATFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 17, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := coverPATRouter(t, tok, commonHandler(createCron))
|
||||||
|
|
||||||
|
body, _ := json.Marshal(model.CronForm{
|
||||||
|
TaskType: model.CronTypeCronTask,
|
||||||
|
Name: "evil cover-all",
|
||||||
|
Scheduler: "@every 1m",
|
||||||
|
Command: "echo pwned",
|
||||||
|
Servers: nil,
|
||||||
|
Cover: model.CronCoverAll,
|
||||||
|
})
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/cron", bytes.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.False(t, success,
|
||||||
|
"PAT scoped to server_ids=[1] must NOT be able to create a CronCoverAll cron with no Servers — that fans out to server 2 outside the whitelist")
|
||||||
|
assert.Contains(t, errMsg, "permission denied")
|
||||||
|
|
||||||
|
var rows []model.Cron
|
||||||
|
require.NoError(t, singleton.DB.Find(&rows).Error)
|
||||||
|
assert.Empty(t, rows, "no cron row must be persisted when the create call is rejected")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateCron_RejectsCoverIgnoreAllWithEmptyServersForLimitedPAT(t *testing.T) {
|
||||||
|
// CoverIgnoreAll + empty Servers is "allow-list of zero" → effectively a
|
||||||
|
// no-op cron. We still reject it because it normalises away the
|
||||||
|
// whitelist hint a curious caller might attempt next ("just flip cover
|
||||||
|
// to All and we'll get fan-out"). Defence-in-depth: any cover-mode that
|
||||||
|
// implies dispatch beyond the literal Servers slice must require the
|
||||||
|
// PAT to cover at least one whitelisted server explicitly.
|
||||||
|
setupCoverPATFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 18, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := coverPATRouter(t, tok, commonHandler(createCron))
|
||||||
|
|
||||||
|
body, _ := json.Marshal(model.CronForm{
|
||||||
|
TaskType: model.CronTypeCronTask,
|
||||||
|
Name: "ambiguous-cover",
|
||||||
|
Scheduler: "@every 1m",
|
||||||
|
Command: "echo",
|
||||||
|
Servers: nil,
|
||||||
|
Cover: model.CronCoverIgnoreAll,
|
||||||
|
})
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/cron", bytes.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
// Empty Servers + IgnoreAll is the degenerate "matches nothing" case;
|
||||||
|
// it must succeed (it cannot escape) so legitimate API consumers
|
||||||
|
// who serialise a 0-server allow-list aren't blocked.
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.True(t, success, "CoverIgnoreAll with no Servers is a no-op; not a bypass: error=%s", errMsg)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateService_AllowsCoverIgnoreAllEmptySkipForLimitedPAT(t *testing.T) {
|
||||||
|
// ServiceCoverIgnoreAll + empty SkipServers is the degenerate "matches
|
||||||
|
// nothing" case: DispatchTask iterates only entries marked true in
|
||||||
|
// SkipServers, so an empty map causes zero fan-out. Pin the no-op
|
||||||
|
// classification so a future refactor that broadens IgnoreAll's
|
||||||
|
// semantics has to update this test (and the dispatch-side guard) in
|
||||||
|
// lock-step with the writer-side guard.
|
||||||
|
setupCoverPATFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 21, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := coverPATRouter(t, tok, commonHandler(createService))
|
||||||
|
|
||||||
|
body, _ := json.Marshal(model.ServiceForm{
|
||||||
|
Name: "no-op monitor",
|
||||||
|
Target: "example.invalid:80",
|
||||||
|
Type: model.TaskTypeTCPPing,
|
||||||
|
Cover: model.ServiceCoverIgnoreAll,
|
||||||
|
SkipServers: nil,
|
||||||
|
Duration: 30,
|
||||||
|
})
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/service", bytes.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.True(t, success, "CoverIgnoreAll with no SkipServers is a no-op; not a bypass: error=%s", errMsg)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCreateService_RejectsCoverAllForServerLimitedPAT(t *testing.T) {
|
||||||
|
setupCoverPATFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 19, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := coverPATRouter(t, tok, commonHandler(createService))
|
||||||
|
|
||||||
|
body, _ := json.Marshal(model.ServiceForm{
|
||||||
|
Name: "evil cover-all monitor",
|
||||||
|
Target: "example.invalid:443",
|
||||||
|
Type: model.TaskTypeTCPPing,
|
||||||
|
Cover: model.ServiceCoverAll,
|
||||||
|
SkipServers: nil,
|
||||||
|
Duration: 30,
|
||||||
|
})
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/service", bytes.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.False(t, success,
|
||||||
|
"PAT scoped to server_ids=[1] must NOT be able to create a ServiceCoverAll monitor with no SkipServers — DispatchTask fans out to server 2 outside the whitelist")
|
||||||
|
assert.Contains(t, errMsg, "permission denied")
|
||||||
|
|
||||||
|
var rows []model.Service
|
||||||
|
require.NoError(t, singleton.DB.Find(&rows).Error)
|
||||||
|
assert.Empty(t, rows, "no service row must be persisted when the create call is rejected")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Threat: PAT server_ids=[1] + Cover=CronCoverAll + Servers=[1] (deny-list)
|
||||||
|
// passes the writer-side guard (len(Servers)>0), then CronTrigger iterates all
|
||||||
|
// owner servers, skips the whitelisted server 1, and dispatches to server 2 —
|
||||||
|
// outside the whitelist. CronTrigger has no PAT context, so the write-time
|
||||||
|
// guard is the only enforcement point.
|
||||||
|
func TestCreateCron_RejectsCoverAllWithDenyListCoveringOnlyWhitelistedServers(t *testing.T) {
|
||||||
|
setupCoverPATFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 31, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := coverPATRouter(t, tok, commonHandler(createCron))
|
||||||
|
|
||||||
|
body, _ := json.Marshal(model.CronForm{
|
||||||
|
TaskType: model.CronTypeCronTask,
|
||||||
|
Name: "cover-all deny-only-whitelisted",
|
||||||
|
Scheduler: "@every 1m",
|
||||||
|
Command: "echo pwned-via-server-2",
|
||||||
|
Servers: []uint64{1},
|
||||||
|
Cover: model.CronCoverAll,
|
||||||
|
})
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/cron", bytes.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.False(t, success,
|
||||||
|
"PAT [1] must NOT create a CronCoverAll whose deny-list only contains whitelisted servers; CronTrigger would fan out to server 2")
|
||||||
|
assert.Contains(t, errMsg, "permission denied")
|
||||||
|
|
||||||
|
var rows []model.Cron
|
||||||
|
require.NoError(t, singleton.DB.Find(&rows).Error)
|
||||||
|
assert.Empty(t, rows, "no cron row must be persisted when the create call is rejected")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Positive case: a server-limited PAT IS allowed to create CronCoverAll when
|
||||||
|
// the deny-list already covers every owner-visible server outside its
|
||||||
|
// whitelist. Pinning this prevents future "just block all CoverAll for PATs"
|
||||||
|
// over-corrections that would break a legitimate "schedule on whitelisted
|
||||||
|
// servers only, via deny-list" workflow.
|
||||||
|
func TestCreateCron_AllowsCoverAllWhenDenyListCoversAllNonWhitelistedServers(t *testing.T) {
|
||||||
|
setupCoverPATFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 41, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := coverPATRouter(t, tok, commonHandler(createCron))
|
||||||
|
|
||||||
|
body, _ := json.Marshal(model.CronForm{
|
||||||
|
TaskType: model.CronTypeCronTask,
|
||||||
|
Name: "legit cover-all",
|
||||||
|
Scheduler: "@every 1m",
|
||||||
|
Command: "echo s1-only",
|
||||||
|
Servers: []uint64{2},
|
||||||
|
Cover: model.CronCoverAll,
|
||||||
|
})
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/cron", bytes.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.True(t, success,
|
||||||
|
"CronCoverAll with deny-list covering every non-whitelisted server must succeed for a server-limited PAT: error=%s", errMsg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Service-monitor analogue of the cron deny-list bypass: ServiceCoverAll +
|
||||||
|
// SkipServers={1:true} passes the writer-side guard (skipCount>0), then
|
||||||
|
// DispatchTask probes server 2. Same write-time enforcement requirement.
|
||||||
|
func TestCreateService_RejectsCoverAllWithSkipListCoveringOnlyWhitelistedServers(t *testing.T) {
|
||||||
|
setupCoverPATFixture(t)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 32, UserID: 100}
|
||||||
|
tok.SetServerIDs([]uint64{1})
|
||||||
|
|
||||||
|
r := coverPATRouter(t, tok, commonHandler(createService))
|
||||||
|
|
||||||
|
body, _ := json.Marshal(model.ServiceForm{
|
||||||
|
Name: "cover-all skip-only-whitelisted monitor",
|
||||||
|
Target: "example.invalid:8443",
|
||||||
|
Type: model.TaskTypeTCPPing,
|
||||||
|
Cover: model.ServiceCoverAll,
|
||||||
|
SkipServers: map[uint64]bool{1: true},
|
||||||
|
Duration: 30,
|
||||||
|
})
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/api/v1/service", bytes.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
assert.False(t, success,
|
||||||
|
"PAT [1] must NOT create a ServiceCoverAll whose SkipServers only marks whitelisted servers; DispatchTask would probe server 2")
|
||||||
|
assert.Contains(t, errMsg, "permission denied")
|
||||||
|
|
||||||
|
var rows []model.Service
|
||||||
|
require.NoError(t, singleton.DB.Find(&rows).Error)
|
||||||
|
assert.Empty(t, rows, "no service row must be persisted when the create call is rejected")
|
||||||
|
}
|
||||||
@@ -0,0 +1,103 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/patrickmn/go-cache"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/i18n"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func setupCronUpdateOwnerUIDFixture(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
originalCache := singleton.Cache
|
||||||
|
originalLoc := singleton.Loc
|
||||||
|
originalLocalizer := singleton.Localizer
|
||||||
|
originalServer := singleton.ServerShared
|
||||||
|
|
||||||
|
singleton.Loc = time.UTC
|
||||||
|
singleton.Cache = cache.New(time.Minute, time.Minute)
|
||||||
|
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
|
||||||
|
|
||||||
|
sc := singleton.NewEmptyServerClassForTest()
|
||||||
|
for _, id := range []uint64{1, 2} {
|
||||||
|
s := &model.Server{}
|
||||||
|
s.ID = id
|
||||||
|
s.SetUserID(100)
|
||||||
|
sc.InsertForTest(s)
|
||||||
|
}
|
||||||
|
adminServer := &model.Server{}
|
||||||
|
adminServer.ID = 5
|
||||||
|
adminServer.SetUserID(200)
|
||||||
|
sc.InsertForTest(adminServer)
|
||||||
|
singleton.ServerShared = sc
|
||||||
|
|
||||||
|
t.Cleanup(func() {
|
||||||
|
singleton.Cache = originalCache
|
||||||
|
singleton.Loc = originalLoc
|
||||||
|
singleton.Localizer = originalLocalizer
|
||||||
|
singleton.ServerShared = originalServer
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func newCtxAsAdminWithLimitedPAT(t *testing.T, callerUID uint64, whitelist []uint64) *gin.Context {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||||
|
c.Set(model.CtxKeyAuthorizedUser, &model.User{
|
||||||
|
Common: model.Common{ID: callerUID},
|
||||||
|
Role: model.RoleAdmin,
|
||||||
|
})
|
||||||
|
tok := &model.APIToken{ID: 33, UserID: callerUID}
|
||||||
|
tok.SetServerIDs(whitelist)
|
||||||
|
c.Set(apiTokenCtxKey, tok)
|
||||||
|
c.Set(model.CtxKeyAPIToken, tok)
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
// Threat: updateCron currently calls
|
||||||
|
//
|
||||||
|
// rejectImplicitCoverForLimitedPAT(c, cf.Cover, cf.Servers)
|
||||||
|
//
|
||||||
|
// which internally resolves the owner UID via getUid(c) (caller id). When an
|
||||||
|
// admin uses a server-limited PAT to flip a *foreign* cron to CoverAll with
|
||||||
|
// an under-specified deny-list, the helper validates the deny-list against
|
||||||
|
// the admin's own servers, not the cron owner's. The admin's only owned
|
||||||
|
// server is 5 and it's already in the whitelist, so the guard returns nil
|
||||||
|
// even though CronTrigger will fan out to the cron owner's servers 1 and 2
|
||||||
|
// — both outside the PAT whitelist. The correct owner is the existing
|
||||||
|
// cron.UserID, not the caller. This test calls the helper directly with the
|
||||||
|
// cron owner uid and pins the safe behaviour.
|
||||||
|
func TestRejectImplicitCoverForLimitedPAT_RejectsCallerWhenCronOwnerHasUncoveredServers(t *testing.T) {
|
||||||
|
setupCronUpdateOwnerUIDFixture(t)
|
||||||
|
|
||||||
|
c := newCtxAsAdminWithLimitedPAT(t, 200, []uint64{5})
|
||||||
|
|
||||||
|
const cronOwnerUID = uint64(100)
|
||||||
|
err := rejectImplicitCoverForLimitedPATWithOwner(c, model.CronCoverAll, nil, cronOwnerUID)
|
||||||
|
require.Error(t, err,
|
||||||
|
"limited PAT must NOT pass cover-all check when the cron owner has servers outside the PAT whitelist; caller uid must not be used as owner")
|
||||||
|
assert.Contains(t, err.Error(), "permission denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Pins the safe path: same helper, but caller uid happens to equal the cron
|
||||||
|
// owner and the deny-list covers every owner-visible server outside the
|
||||||
|
// whitelist. Prevents regressing the helper into a blanket "always deny
|
||||||
|
// limited PAT" form.
|
||||||
|
func TestRejectImplicitCoverForLimitedPAT_AllowsCallerWhenDenyListCoversEveryOwnerServerOutsideWhitelist(t *testing.T) {
|
||||||
|
setupCronUpdateOwnerUIDFixture(t)
|
||||||
|
|
||||||
|
c := newCtxAsAdminWithLimitedPAT(t, 100, []uint64{1})
|
||||||
|
|
||||||
|
err := rejectImplicitCoverForLimitedPATWithOwner(c, model.CronCoverAll, []uint64{2}, 100)
|
||||||
|
require.NoError(t, err,
|
||||||
|
"deny-list [2] covers every server uid 100 owns outside the PAT whitelist [1]; must pass")
|
||||||
|
}
|
||||||
@@ -0,0 +1,149 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/hmac"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// issueCSRFToken mints a signed double-submit token (nonce.HMAC-SHA256 keyed
|
||||||
|
// by the JWT secret). Signing defeats sibling-subdomain cookie tossing: a
|
||||||
|
// naive double-submit trusts any header==cookie pair, but an injected cookie
|
||||||
|
// carries no valid HMAC and fails validateCSRFToken. Returns "" pre-init
|
||||||
|
// (no secret); callers treat that as "no cookie minted".
|
||||||
|
func issueCSRFToken() string {
|
||||||
|
secret := csrfSigningSecret()
|
||||||
|
if secret == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
var b [32]byte
|
||||||
|
if _, err := rand.Read(b[:]); err != nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
nonce := hex.EncodeToString(b[:])
|
||||||
|
return nonce + "." + csrfSign(nonce, secret)
|
||||||
|
}
|
||||||
|
|
||||||
|
func csrfSigningSecret() string {
|
||||||
|
if singleton.Conf == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return singleton.Conf.JWTSecretKey
|
||||||
|
}
|
||||||
|
|
||||||
|
func csrfSign(nonce, secret string) string {
|
||||||
|
mac := hmac.New(sha256.New, []byte(secret))
|
||||||
|
mac.Write([]byte(nonce))
|
||||||
|
return hex.EncodeToString(mac.Sum(nil))
|
||||||
|
}
|
||||||
|
|
||||||
|
// validateCSRFToken reports whether value is a well-formed nonce.signature
|
||||||
|
// pair whose signature verifies under the current server secret. Constant
|
||||||
|
// -time comparison guards against signature-probing side channels.
|
||||||
|
func validateCSRFToken(value string) bool {
|
||||||
|
secret := csrfSigningSecret()
|
||||||
|
if secret == "" || value == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
idx := strings.LastIndex(value, ".")
|
||||||
|
if idx <= 0 || idx == len(value)-1 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
nonce, sig := value[:idx], value[idx+1:]
|
||||||
|
return hmac.Equal([]byte(sig), []byte(csrfSign(nonce, secret)))
|
||||||
|
}
|
||||||
|
|
||||||
|
// setCSRFCookie issues a fresh signed CSRF token cookie. Called by login +
|
||||||
|
// refresh handlers so the frontend always has a paired value to mirror back
|
||||||
|
// into the X-CSRF-Token header. The cookie is intentionally HttpOnly=false —
|
||||||
|
// SPA JS must be able to read it. SameSite=Strict here (not Lax) because
|
||||||
|
// the cookie's sole purpose is the same-origin double-submit check and we
|
||||||
|
// don't want it leaking on cross-site GET navigation either.
|
||||||
|
func setCSRFCookie(c *gin.Context) {
|
||||||
|
token := issueCSRFToken()
|
||||||
|
if token == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// Secure is set only when the request arrives over HTTPS, mirroring
|
||||||
|
// writeOauth2StateCookie. On plain HTTP (e.g. intranet deployments) a
|
||||||
|
// Secure cookie would be dropped by the browser, breaking the
|
||||||
|
// double-submit pair, so we must not force it unconditionally.
|
||||||
|
secure := c.Request.URL.Scheme == "https" || c.Request.TLS != nil
|
||||||
|
c.SetSameSite(http.SameSiteStrictMode)
|
||||||
|
c.SetCookie(csrfCookieName, token, 0, "/", "", secure, false)
|
||||||
|
}
|
||||||
|
|
||||||
|
const (
|
||||||
|
csrfCookieName = "nz-csrf"
|
||||||
|
csrfHeaderName = "X-CSRF-Token"
|
||||||
|
)
|
||||||
|
|
||||||
|
// csrfMiddleware enforces a double-submit-cookie CSRF gate on unsafe
|
||||||
|
// HTTP methods for cookie-authenticated requests.
|
||||||
|
//
|
||||||
|
// Why: SameSite=Lax on the nz-jwt cookie blocks the simplest cross-site
|
||||||
|
// form POST, but it does not stop same-site XSS-pivot CSRF, header method
|
||||||
|
// override, redirect-leaking auth helpers, or carefully chained sub-domain
|
||||||
|
// attacks. The double-submit pattern (server sets a JS-readable nz-csrf
|
||||||
|
// cookie, client mirrors the value into X-CSRF-Token) closes the gap
|
||||||
|
// without coupling auth state to a server-side session.
|
||||||
|
//
|
||||||
|
// Bypass conditions:
|
||||||
|
// - Safe methods (GET/HEAD/OPTIONS): no state mutation, no CSRF risk.
|
||||||
|
// - Bearer-token PAT requests (`Authorization: Bearer nzp_*`): stateless,
|
||||||
|
// no ambient cookie, so a CSRF attack cannot induce them.
|
||||||
|
//
|
||||||
|
// Reject conditions:
|
||||||
|
// - Missing or empty X-CSRF-Token header.
|
||||||
|
// - Missing or empty nz-csrf cookie.
|
||||||
|
// - Header value != cookie value.
|
||||||
|
// - Cookie value not signed by the server (validateCSRFToken fails).
|
||||||
|
//
|
||||||
|
// The middleware DOES NOT set the csrf cookie on its own — that is the
|
||||||
|
// JWT login / refresh handler's job, since those are the only places that
|
||||||
|
// know when to mint a fresh value. The pair just has to exist by the time
|
||||||
|
// any unsafe call reaches here.
|
||||||
|
func csrfMiddleware() gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
switch c.Request.Method {
|
||||||
|
case http.MethodGet, http.MethodHead, http.MethodOptions:
|
||||||
|
// Self-heal sessions that predate the CSRF cookie (or whose
|
||||||
|
// nz-csrf expired): seed a fresh value on a safe method so the
|
||||||
|
// double-submit pair exists before the next unsafe call. Safe
|
||||||
|
// methods mutate nothing, so minting here carries no CSRF risk.
|
||||||
|
if cookie, err := c.Cookie(csrfCookieName); err != nil || cookie == "" {
|
||||||
|
setCSRFCookie(c)
|
||||||
|
}
|
||||||
|
c.Next()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// A PAT request carries no ambient cookie, so CSRF cannot induce it.
|
||||||
|
// The exemption must check the authenticated PAT identity resolved by
|
||||||
|
// apiTokenAuthMiddleware, not a forgeable Authorization header value.
|
||||||
|
if APITokenFromContext(c) != nil {
|
||||||
|
c.Next()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
header := c.GetHeader(csrfHeaderName)
|
||||||
|
cookie, err := c.Cookie(csrfCookieName)
|
||||||
|
// Both halves must be present, mirror each other, AND carry a valid
|
||||||
|
// server signature. The signature check is what stops a cookie-tossed
|
||||||
|
// pair from a sibling subdomain.
|
||||||
|
if err != nil || cookie == "" || header == "" || header != cookie || !validateCSRFToken(cookie) {
|
||||||
|
c.AbortWithStatusJSON(http.StatusForbidden, model.CommonResponse[any]{
|
||||||
|
Success: false,
|
||||||
|
Error: "ApiErrorForbidden: missing or invalid CSRF token",
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.Next()
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,84 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func withCSRFSecret(t *testing.T, secret string) {
|
||||||
|
t.Helper()
|
||||||
|
prev := singleton.Conf
|
||||||
|
singleton.Conf = &singleton.ConfigClass{Config: &model.Config{}}
|
||||||
|
singleton.Conf.JWTSecretKey = secret
|
||||||
|
t.Cleanup(func() { singleton.Conf = prev })
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sibling-subdomain cookie tossing: an attacker who can set nz-csrf for the
|
||||||
|
// parent domain injects an attacker-chosen value and mirrors it into the
|
||||||
|
// header. A naive double-submit accepts header==cookie. The signed
|
||||||
|
// double-submit must reject it because the injected value carries no valid
|
||||||
|
// server HMAC.
|
||||||
|
func TestCSRFMiddleware_RejectsUnsignedInjectedPair(t *testing.T) {
|
||||||
|
withCSRFSecret(t, "test-jwt-secret")
|
||||||
|
mw := csrfMiddleware()
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("POST", "/api/v1/profile", nil)
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: "nz-jwt", Value: "session"})
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: csrfCookieName, Value: "attacker-chosen"})
|
||||||
|
c.Request.Header.Set(csrfHeaderName, "attacker-chosen")
|
||||||
|
mw(c)
|
||||||
|
if !c.IsAborted() || w.Code != http.StatusForbidden {
|
||||||
|
t.Fatalf("an unsigned (cookie-tossed) csrf pair must be rejected, got aborted=%v code=%d", c.IsAborted(), w.Code)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A token minted by the server (issueCSRFToken) must pass when mirrored
|
||||||
|
// correctly into the header.
|
||||||
|
func TestCSRFMiddleware_AcceptsServerSignedPair(t *testing.T) {
|
||||||
|
withCSRFSecret(t, "test-jwt-secret")
|
||||||
|
token := issueCSRFToken()
|
||||||
|
if token == "" {
|
||||||
|
t.Fatal("issueCSRFToken must produce a value")
|
||||||
|
}
|
||||||
|
mw := csrfMiddleware()
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("POST", "/api/v1/profile", nil)
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: "nz-jwt", Value: "session"})
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: csrfCookieName, Value: token})
|
||||||
|
c.Request.Header.Set(csrfHeaderName, token)
|
||||||
|
mw(c)
|
||||||
|
if c.IsAborted() {
|
||||||
|
t.Fatal("a correctly mirrored server-signed csrf pair must pass")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Even with header==cookie, a value whose signature does not verify under the
|
||||||
|
// server secret must be rejected (forged signature segment).
|
||||||
|
func TestCSRFMiddleware_RejectsForgedSignature(t *testing.T) {
|
||||||
|
withCSRFSecret(t, "test-jwt-secret")
|
||||||
|
token := issueCSRFToken()
|
||||||
|
forged := token[:len(token)-1]
|
||||||
|
if token[len(token)-1] == 'a' {
|
||||||
|
forged += "b"
|
||||||
|
} else {
|
||||||
|
forged += "a"
|
||||||
|
}
|
||||||
|
mw := csrfMiddleware()
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("POST", "/api/v1/profile", nil)
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: csrfCookieName, Value: forged})
|
||||||
|
c.Request.Header.Set(csrfHeaderName, forged)
|
||||||
|
mw(c)
|
||||||
|
if !c.IsAborted() || w.Code != http.StatusForbidden {
|
||||||
|
t.Fatalf("a forged-signature csrf pair must be rejected, got aborted=%v code=%d", c.IsAborted(), w.Code)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,162 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
// H6 regression: cookie-JWT unsafe-method routes need a real CSRF gate.
|
||||||
|
// SameSite=Lax blocks the obvious cross-site form-POST but does not stop
|
||||||
|
// same-site siblings, sub-domain XSS pivots, header method override quirks,
|
||||||
|
// or the various legacy edge cases. The middleware below is the double
|
||||||
|
// -submit cookie pattern: require X-CSRF-Token whose value matches the
|
||||||
|
// `nz-csrf` cookie. PAT bearer requests bypass the gate (they don't carry
|
||||||
|
// the cookie at all and are already authenticated stateless).
|
||||||
|
func TestCSRFMiddleware_AllowsSafeMethodsWithoutToken(t *testing.T) {
|
||||||
|
mw := csrfMiddleware()
|
||||||
|
for _, m := range []string{"GET", "HEAD", "OPTIONS"} {
|
||||||
|
t.Run(m, func(t *testing.T) {
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest(m, "/api/v1/profile", nil)
|
||||||
|
mw(c)
|
||||||
|
if c.IsAborted() {
|
||||||
|
t.Fatalf("%s must pass without csrf token", m)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Self-heal: a session created before the CSRF cookie existed (or one whose
|
||||||
|
// nz-csrf expired) only carries nz-jwt. Without seeding a fresh nz-csrf on a
|
||||||
|
// safe GET, every subsequent unsafe call — including the auto refresh-token
|
||||||
|
// POST — would 403 forever and force a manual re-login. GET carries no CSRF
|
||||||
|
// risk, so the middleware mints the cookie when it is absent.
|
||||||
|
func TestCSRFMiddleware_SeedsCookieOnSafeMethodWhenMissing(t *testing.T) {
|
||||||
|
withCSRFSecret(t, "test-jwt-secret")
|
||||||
|
mw := csrfMiddleware()
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/api/v1/profile", nil)
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: "nz-jwt", Value: "session"})
|
||||||
|
mw(c)
|
||||||
|
if c.IsAborted() {
|
||||||
|
t.Fatal("safe GET must never abort")
|
||||||
|
}
|
||||||
|
var seeded bool
|
||||||
|
for _, sc := range w.Result().Cookies() {
|
||||||
|
if sc.Name == csrfCookieName && sc.Value != "" {
|
||||||
|
seeded = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !seeded {
|
||||||
|
t.Fatal("missing nz-csrf must be seeded on a safe GET so the SPA can self-heal")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCSRFMiddleware_DoesNotReseedWhenCookiePresent(t *testing.T) {
|
||||||
|
mw := csrfMiddleware()
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/api/v1/profile", nil)
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: csrfCookieName, Value: "existing"})
|
||||||
|
mw(c)
|
||||||
|
for _, sc := range w.Result().Cookies() {
|
||||||
|
if sc.Name == csrfCookieName {
|
||||||
|
t.Fatal("existing nz-csrf must not be rotated on every GET")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCSRFMiddleware_BlocksUnsafeMethodWithoutToken(t *testing.T) {
|
||||||
|
mw := csrfMiddleware()
|
||||||
|
for _, m := range []string{"POST", "PATCH", "PUT", "DELETE"} {
|
||||||
|
t.Run(m, func(t *testing.T) {
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest(m, "/api/v1/profile", nil)
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: "nz-jwt", Value: "anything"})
|
||||||
|
mw(c)
|
||||||
|
if !c.IsAborted() || w.Code != http.StatusForbidden {
|
||||||
|
t.Fatalf("%s without csrf token must abort 403, got aborted=%v code=%d", m, c.IsAborted(), w.Code)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCSRFMiddleware_AcceptsMatchingHeaderAndCookie(t *testing.T) {
|
||||||
|
withCSRFSecret(t, "test-jwt-secret")
|
||||||
|
token := issueCSRFToken()
|
||||||
|
mw := csrfMiddleware()
|
||||||
|
for _, m := range []string{"POST", "PATCH", "PUT", "DELETE"} {
|
||||||
|
t.Run(m, func(t *testing.T) {
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest(m, "/api/v1/profile", nil)
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: "nz-jwt", Value: "anything"})
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: "nz-csrf", Value: token})
|
||||||
|
c.Request.Header.Set("X-CSRF-Token", token)
|
||||||
|
mw(c)
|
||||||
|
if c.IsAborted() {
|
||||||
|
t.Fatalf("%s with matching signed csrf header+cookie must pass", m)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCSRFMiddleware_RejectsMismatchedHeader(t *testing.T) {
|
||||||
|
mw := csrfMiddleware()
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("POST", "/api/v1/profile", nil)
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: "nz-csrf", Value: "value-a"})
|
||||||
|
c.Request.Header.Set("X-CSRF-Token", "value-b")
|
||||||
|
mw(c)
|
||||||
|
if !c.IsAborted() || w.Code != http.StatusForbidden {
|
||||||
|
t.Fatalf("mismatched csrf token must abort 403, got aborted=%v code=%d", c.IsAborted(), w.Code)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCSRFMiddleware_AuthenticatedPATBypassesCheck(t *testing.T) {
|
||||||
|
mw := csrfMiddleware()
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("POST", "/mcp", strings.NewReader("{}"))
|
||||||
|
c.Set(apiTokenCtxKey, &model.APIToken{ID: 1})
|
||||||
|
mw(c)
|
||||||
|
if c.IsAborted() {
|
||||||
|
t.Fatal("authenticated PAT must bypass CSRF — stateless auth")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCSRFMiddleware_ForgedBearerHeaderDoesNotBypass(t *testing.T) {
|
||||||
|
mw := csrfMiddleware()
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("POST", "/api/v1/profile", nil)
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: "nz-jwt", Value: "session"})
|
||||||
|
c.Request.Header.Set("Authorization", "Bearer "+model.APITokenPrefix+"never-authenticated")
|
||||||
|
mw(c)
|
||||||
|
if !c.IsAborted() || w.Code != http.StatusForbidden {
|
||||||
|
t.Fatalf("a Bearer nzp_* header that never authenticated must not skip CSRF, got aborted=%v code=%d", c.IsAborted(), w.Code)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCSRFMiddleware_EmptyTokenRejected(t *testing.T) {
|
||||||
|
mw := csrfMiddleware()
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("POST", "/api/v1/profile", nil)
|
||||||
|
c.Request.AddCookie(&http.Cookie{Name: "nz-csrf", Value: ""})
|
||||||
|
c.Request.Header.Set("X-CSRF-Token", "")
|
||||||
|
mw(c)
|
||||||
|
if !c.IsAborted() {
|
||||||
|
t.Fatal("empty csrf token must not satisfy the check (defeats the gate entirely)")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -30,6 +30,13 @@ func listDDNS(c *gin.Context) ([]*model.DDNSProfile, error) {
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 列表端点不回显写入态凭据:ddnsProfiles 是 copier 复制出的副本,置零安全,
|
||||||
|
// 不影响 singleton 内原始数据。
|
||||||
|
for _, p := range ddnsProfiles {
|
||||||
|
p.AccessSecret = ""
|
||||||
|
p.WebhookHeaders = ""
|
||||||
|
}
|
||||||
|
|
||||||
return ddnsProfiles, nil
|
return ddnsProfiles, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -137,12 +144,18 @@ func updateDDNS(c *gin.Context) (any, error) {
|
|||||||
p.Provider = df.Provider
|
p.Provider = df.Provider
|
||||||
p.Domains = df.Domains
|
p.Domains = df.Domains
|
||||||
p.AccessID = df.AccessID
|
p.AccessID = df.AccessID
|
||||||
p.AccessSecret = df.AccessSecret
|
|
||||||
p.WebhookURL = df.WebhookURL
|
p.WebhookURL = df.WebhookURL
|
||||||
p.WebhookMethod = df.WebhookMethod
|
p.WebhookMethod = df.WebhookMethod
|
||||||
p.WebhookRequestType = df.WebhookRequestType
|
p.WebhookRequestType = df.WebhookRequestType
|
||||||
p.WebhookRequestBody = df.WebhookRequestBody
|
p.WebhookRequestBody = df.WebhookRequestBody
|
||||||
p.WebhookHeaders = df.WebhookHeaders
|
|
||||||
|
// 凭据在列表接口已脱敏,前端无法回填;空值视为"不修改",保留旧值避免误清空。
|
||||||
|
if df.AccessSecret != "" {
|
||||||
|
p.AccessSecret = df.AccessSecret
|
||||||
|
}
|
||||||
|
if df.WebhookHeaders != "" {
|
||||||
|
p.WebhookHeaders = df.WebhookHeaders
|
||||||
|
}
|
||||||
|
|
||||||
for n, domain := range p.Domains {
|
for n, domain := range p.Domains {
|
||||||
// IDN to ASCII
|
// IDN to ASCII
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ import (
|
|||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/goccy/go-json"
|
"github.com/goccy/go-json"
|
||||||
"github.com/gorilla/websocket"
|
|
||||||
"github.com/hashicorp/go-uuid"
|
"github.com/hashicorp/go-uuid"
|
||||||
|
|
||||||
"github.com/nezhahq/nezha/model"
|
"github.com/nezhahq/nezha/model"
|
||||||
@@ -24,8 +23,9 @@ import (
|
|||||||
// @Param id query uint true "Server ID"
|
// @Param id query uint true "Server ID"
|
||||||
// @Produce json
|
// @Produce json
|
||||||
// @Success 200 {object} model.CreateFMResponse
|
// @Success 200 {object} model.CreateFMResponse
|
||||||
// @Router /file [get]
|
// @Router /file [post]
|
||||||
func createFM(c *gin.Context) (*model.CreateFMResponse, error) {
|
func createFM(c *gin.Context) (*model.CreateFMResponse, error) {
|
||||||
|
prepareAgentcompatCapabilityHeader(c)
|
||||||
idStr := c.Query("id")
|
idStr := c.Query("id")
|
||||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
id, err := strconv.ParseUint(idStr, 10, 64)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -33,7 +33,10 @@ func createFM(c *gin.Context) (*model.CreateFMResponse, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
server, _ := singleton.ServerShared.Get(id)
|
server, _ := singleton.ServerShared.Get(id)
|
||||||
if server == nil || server.TaskStream == nil {
|
if server == nil {
|
||||||
|
return nil, singleton.Localizer.ErrorT("server not found or not connected")
|
||||||
|
}
|
||||||
|
if server.GetTaskStream() == nil {
|
||||||
return nil, singleton.Localizer.ErrorT("server not found or not connected")
|
return nil, singleton.Localizer.ErrorT("server not found or not connected")
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -46,15 +49,24 @@ func createFM(c *gin.Context) (*model.CreateFMResponse, error) {
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
rpc.NezhaHandlerSingleton.CreateStream(streamId)
|
cleanup, err := createIOStreamWithAgentcompatCapability(c, streamId, getUid(c), server.ID, rpc.AgentCompatCapabilityFileManager)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
fmData, _ := json.Marshal(&model.TaskFM{
|
fmData, err := json.Marshal(&model.TaskFM{
|
||||||
StreamID: streamId,
|
StreamID: streamId,
|
||||||
})
|
})
|
||||||
if err := server.TaskStream.Send(&proto.Task{
|
if err != nil {
|
||||||
|
// A stream is owned by the caller only after this function succeeds.
|
||||||
|
cleanup()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if err := server.SendTask(&proto.Task{
|
||||||
Type: model.TaskTypeFM,
|
Type: model.TaskTypeFM,
|
||||||
Data: string(fmData),
|
Data: string(fmData),
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
|
cleanup()
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -72,6 +84,12 @@ func createFM(c *gin.Context) (*model.CreateFMResponse, error) {
|
|||||||
// @Router /ws/file/{id} [get]
|
// @Router /ws/file/{id} [get]
|
||||||
func fmStream(c *gin.Context) (any, error) {
|
func fmStream(c *gin.Context) (any, error) {
|
||||||
streamId := c.Param("id")
|
streamId := c.Param("id")
|
||||||
|
// GHSA-style fix: io_stream sessions must be reachable only by their creator
|
||||||
|
// (or an admin). Without this, any authenticated user who learns a stream
|
||||||
|
// UUID can hijack a live file-manager session on the target server.
|
||||||
|
if !streamAttachAllowedForRequest(c, streamId) {
|
||||||
|
return nil, singleton.Localizer.ErrorT("permission denied")
|
||||||
|
}
|
||||||
if _, err := rpc.NezhaHandlerSingleton.GetStream(streamId); err != nil {
|
if _, err := rpc.NezhaHandlerSingleton.GetStream(streamId); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -81,18 +99,14 @@ func fmStream(c *gin.Context) (any, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, newWsError("%v", err)
|
return nil, newWsError("%v", err)
|
||||||
}
|
}
|
||||||
defer wsConn.Close()
|
|
||||||
conn := websocketx.NewConn(wsConn)
|
conn := websocketx.NewConn(wsConn)
|
||||||
|
pingTransport := newWebsocketPingTransport(conn, wsConn.Close)
|
||||||
|
stopPing := startWebsocketPingTicker(c.Request.Context(), time.Second*10, pingTransport)
|
||||||
|
|
||||||
go func() {
|
deregisterPAT := registerPATConnection(c, func() { _ = pingTransport.Close() })
|
||||||
// PING 保活
|
defer deregisterPAT()
|
||||||
for {
|
// Join the ping worker before PAT and WebSocket cleanup can close its writer.
|
||||||
if err = conn.WriteMessage(websocket.PingMessage, []byte{}); err != nil {
|
defer stopPing()
|
||||||
return
|
|
||||||
}
|
|
||||||
time.Sleep(time.Second * 10)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
if err = rpc.NezhaHandlerSingleton.UserConnected(streamId, conn); err != nil {
|
if err = rpc.NezhaHandlerSingleton.UserConnected(streamId, conn); err != nil {
|
||||||
return nil, newWsError("%v", err)
|
return nil, newWsError("%v", err)
|
||||||
|
|||||||
@@ -0,0 +1,149 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/i18n"
|
||||||
|
pb "github.com/nezhahq/nezha/proto"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// fakeTaskStream is the minimum stub of pb.NezhaService_RequestTaskServer
|
||||||
|
// required to make a server look "online" to forceUpdateServer. Only Send is
|
||||||
|
// called; we capture its argument so the test can verify the upgrade task is
|
||||||
|
// NOT dispatched for foreign IDs.
|
||||||
|
type fakeTaskStream struct {
|
||||||
|
pb.NezhaService_RequestTaskServer
|
||||||
|
sentTasks []*pb.Task
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeTaskStream) Send(t *pb.Task) error {
|
||||||
|
f.sentTasks = append(f.sentTasks, t)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// setupServerOwnershipFixture seeds an in-memory DB with alice's server
|
||||||
|
// (UserID=100, ID=1). The returned stream is wired in so the server is
|
||||||
|
// "online" — i.e. exercises the path that previously returned permission
|
||||||
|
// denied for foreign callers (the actual leak channel).
|
||||||
|
func setupServerOwnershipFixture(t *testing.T) (stream *fakeTaskStream, reset func()) {
|
||||||
|
t.Helper()
|
||||||
|
if singleton.Localizer == nil {
|
||||||
|
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
|
||||||
|
}
|
||||||
|
originalDB := singleton.DB
|
||||||
|
originalShared := singleton.ServerShared
|
||||||
|
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NoError(t, db.AutoMigrate(&model.Server{}))
|
||||||
|
assert.NoError(t, db.Create(&model.Server{
|
||||||
|
Common: model.Common{ID: 1, UserID: 100},
|
||||||
|
Name: "alice-online",
|
||||||
|
}).Error)
|
||||||
|
singleton.DB = db
|
||||||
|
singleton.ServerShared = singleton.NewServerClass()
|
||||||
|
|
||||||
|
alice, _ := singleton.ServerShared.Get(1)
|
||||||
|
stream = &fakeTaskStream{}
|
||||||
|
alice.SetTaskStream(stream)
|
||||||
|
|
||||||
|
return stream, func() {
|
||||||
|
singleton.DB = originalDB
|
||||||
|
singleton.ServerShared = originalShared
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func runForceUpdate(t *testing.T, callerID uint64, ids []uint64) []byte {
|
||||||
|
t.Helper()
|
||||||
|
r := gin.New()
|
||||||
|
r.Use(func(c *gin.Context) {
|
||||||
|
setAuthUser(c, callerID, model.RoleMember)
|
||||||
|
c.Next()
|
||||||
|
})
|
||||||
|
r.POST("/force-update/server", commonHandler(forceUpdateServer))
|
||||||
|
|
||||||
|
body, _ := json.Marshal(ids)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/force-update/server", bytes.NewReader(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
return w.Body.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
type forceUpdateBody struct {
|
||||||
|
Success bool `json:"success"`
|
||||||
|
Error string `json:"error"`
|
||||||
|
Data struct {
|
||||||
|
Offline []uint64 `json:"offline"`
|
||||||
|
Success []uint64 `json:"success"`
|
||||||
|
Failure []uint64 `json:"failure"`
|
||||||
|
} `json:"data"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeForceUpdate(t *testing.T, body []byte) forceUpdateBody {
|
||||||
|
t.Helper()
|
||||||
|
var resp forceUpdateBody
|
||||||
|
assert.NoError(t, json.Unmarshal(body, &resp))
|
||||||
|
return resp
|
||||||
|
}
|
||||||
|
|
||||||
|
// Core regression: bob submitting alice's online server ID must NOT produce
|
||||||
|
// a distinct response from bob submitting an unknown ID. The original code
|
||||||
|
// returned "permission denied" for the former and a structured success/Offline
|
||||||
|
// response for the latter — that delta is the enumeration oracle for server
|
||||||
|
// IDs and online state.
|
||||||
|
func TestForceUpdateServerOnlineForeignIDIndistinguishableFromUnknown(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
_, reset := setupServerOwnershipFixture(t)
|
||||||
|
defer reset()
|
||||||
|
|
||||||
|
const bobID = uint64(200)
|
||||||
|
foreignResp := decodeForceUpdate(t, runForceUpdate(t, bobID, []uint64{1})) // alice's online
|
||||||
|
unknownResp := decodeForceUpdate(t, runForceUpdate(t, bobID, []uint64{9999})) // does not exist
|
||||||
|
|
||||||
|
assert.Equal(t, foreignResp.Success, unknownResp.Success,
|
||||||
|
"top-level success flag must not differ between foreign-online and unknown IDs")
|
||||||
|
assert.Equal(t, foreignResp.Error, unknownResp.Error,
|
||||||
|
"error string must not differ — distinct error reveals existence/state of foreign servers")
|
||||||
|
assert.Equal(t, foreignResp.Data.Success, unknownResp.Data.Success)
|
||||||
|
assert.Equal(t, foreignResp.Data.Failure, unknownResp.Data.Failure)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Submitting a foreign online server must NOT actually trigger the upgrade
|
||||||
|
// task on it — that would be a write primitive on someone else's machine.
|
||||||
|
func TestForceUpdateServerForeignOnlineDoesNotDispatchUpgrade(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
stream, reset := setupServerOwnershipFixture(t)
|
||||||
|
defer reset()
|
||||||
|
|
||||||
|
_ = runForceUpdate(t, 200, []uint64{1}) // bob hits alice's online server
|
||||||
|
assert.Empty(t, stream.sentTasks,
|
||||||
|
"foreign server must not receive the upgrade task even when online")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Sanity: owner submitting their own online server must still get the upgrade
|
||||||
|
// dispatched and a structured success response — the hardening must not
|
||||||
|
// regress the legitimate case.
|
||||||
|
func TestForceUpdateServerOwnerOnlineStillDispatches(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
stream, reset := setupServerOwnershipFixture(t)
|
||||||
|
defer reset()
|
||||||
|
|
||||||
|
resp := decodeForceUpdate(t, runForceUpdate(t, 100, []uint64{1})) // alice on her own server
|
||||||
|
assert.True(t, resp.Success)
|
||||||
|
assert.Equal(t, []uint64{1}, resp.Data.Success)
|
||||||
|
assert.Empty(t, resp.Data.Offline)
|
||||||
|
assert.Empty(t, resp.Data.Failure)
|
||||||
|
assert.Len(t, stream.sentTasks, 1, "owner's own server must receive the upgrade task exactly once")
|
||||||
|
}
|
||||||
@@ -0,0 +1,25 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 前端在 main.tsx 注册了 /dashboard/settings/api-tokens,但后端 fallback 白名单
|
||||||
|
// 漏加这条会让用户直接刷新该页面拿到 HTTP 404(body 还是 index.html)。
|
||||||
|
// controller.go 旁边的注释明确说「新增前端路由时必须在 main.tsx 与这里同步加」。
|
||||||
|
func TestFallbackToFrontend_APITokensRouteReturns200(t *testing.T) {
|
||||||
|
t.Chdir(t.TempDir())
|
||||||
|
router := newFrontendFallbackTestRouter(t)
|
||||||
|
|
||||||
|
w := performFrontendFallbackRequest(t, router, "/dashboard/settings/api-tokens")
|
||||||
|
if w.Code != http.StatusOK {
|
||||||
|
t.Fatalf("/dashboard/settings/api-tokens fallback status = %d, want 200 "+
|
||||||
|
"(front-end main.tsx registered the route — backend SPA fallback regex must mirror it)",
|
||||||
|
w.Code)
|
||||||
|
}
|
||||||
|
if !strings.Contains(w.Body.String(), "admin index") {
|
||||||
|
t.Fatalf("/dashboard/settings/api-tokens must serve admin index.html, got %q", w.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,133 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io/fs"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newFrontendFallbackTestRouter(t *testing.T) *gin.Engine {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
originalConf := singleton.Conf
|
||||||
|
singleton.Conf = &singleton.ConfigClass{Config: &model.Config{
|
||||||
|
ConfigDashboard: model.ConfigDashboard{
|
||||||
|
AdminTemplate: "admin-dist",
|
||||||
|
UserTemplate: "user-dist",
|
||||||
|
},
|
||||||
|
}}
|
||||||
|
t.Cleanup(func() { singleton.Conf = originalConf })
|
||||||
|
|
||||||
|
writeFrontendFallbackTestFile(t, "admin-dist/index.html", "<html>admin index</html>")
|
||||||
|
writeFrontendFallbackTestFile(t, "admin-dist/assets/app.js", "console.log('admin asset')")
|
||||||
|
writeFrontendFallbackTestFile(t, "user-dist/index.html", "<html>user index</html>")
|
||||||
|
writeFrontendFallbackTestFile(t, "data/config.yaml", "jwt_secret_key: traversal-secret")
|
||||||
|
|
||||||
|
r := gin.New()
|
||||||
|
r.NoRoute(fallbackToFrontend(testFrontendDist{}))
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeFrontendFallbackTestFile(t *testing.T, name, content string) {
|
||||||
|
t.Helper()
|
||||||
|
if err := os.MkdirAll(filepath.Dir(name), 0o755); err != nil {
|
||||||
|
t.Fatalf("create fixture directory: %v", err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(name, []byte(content), 0o644); err != nil {
|
||||||
|
t.Fatalf("write fixture file: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type testFrontendDist struct{}
|
||||||
|
|
||||||
|
func (testFrontendDist) Open(string) (fs.File, error) {
|
||||||
|
return nil, fs.ErrNotExist
|
||||||
|
}
|
||||||
|
|
||||||
|
func performFrontendFallbackRequest(t *testing.T, router *gin.Engine, target string) *httptest.ResponseRecorder {
|
||||||
|
t.Helper()
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, target, nil)
|
||||||
|
router.ServeHTTP(w, req)
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallbackToFrontendBlocksDashboardTraversal(t *testing.T) {
|
||||||
|
t.Chdir(t.TempDir())
|
||||||
|
router := newFrontendFallbackTestRouter(t)
|
||||||
|
|
||||||
|
tests := []string{
|
||||||
|
"/dashboard../data/config.yaml",
|
||||||
|
"/dashboard%2e%2e/data/config.yaml",
|
||||||
|
"/dashboard%2e%2e%2fdata%2fconfig.yaml",
|
||||||
|
"/dashboard/../data/config.yaml",
|
||||||
|
"/dashboard/%2e%2e/data/config.yaml",
|
||||||
|
"/dashboard../assets/app.js",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, target := range tests {
|
||||||
|
t.Run(target, func(t *testing.T) {
|
||||||
|
w := performFrontendFallbackRequest(t, router, target)
|
||||||
|
body := w.Body.String()
|
||||||
|
if strings.Contains(body, "traversal-secret") || strings.Contains(body, "jwt_secret_key") || strings.Contains(body, "admin asset") {
|
||||||
|
t.Fatalf("%s leaked protected content with status %d: %q", target, w.Code, body)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallbackToFrontendBlocksUserTraversal(t *testing.T) {
|
||||||
|
t.Chdir(t.TempDir())
|
||||||
|
router := newFrontendFallbackTestRouter(t)
|
||||||
|
|
||||||
|
tests := []string{
|
||||||
|
"/../data/config.yaml",
|
||||||
|
"/%2e%2e/data/config.yaml",
|
||||||
|
"/%2e%2e%2fdata%2fconfig.yaml",
|
||||||
|
"/../admin-dist/assets/app.js",
|
||||||
|
"/%2e%2e/admin-dist/assets/app.js",
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, target := range tests {
|
||||||
|
t.Run(target, func(t *testing.T) {
|
||||||
|
w := performFrontendFallbackRequest(t, router, target)
|
||||||
|
body := w.Body.String()
|
||||||
|
if strings.Contains(body, "traversal-secret") || strings.Contains(body, "jwt_secret_key") || strings.Contains(body, "admin asset") {
|
||||||
|
t.Fatalf("%s leaked protected content with status %d: %q", target, w.Code, body)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFallbackToFrontendPreservesDashboardRoutes(t *testing.T) {
|
||||||
|
t.Chdir(t.TempDir())
|
||||||
|
router := newFrontendFallbackTestRouter(t)
|
||||||
|
|
||||||
|
w := performFrontendFallbackRequest(t, router, "/dashboard")
|
||||||
|
if w.Code != http.StatusMovedPermanently {
|
||||||
|
t.Fatalf("/dashboard status = %d, want %d", w.Code, http.StatusMovedPermanently)
|
||||||
|
}
|
||||||
|
if location := w.Header().Get("Location"); location != "/dashboard/" {
|
||||||
|
t.Fatalf("/dashboard Location = %q, want /dashboard/", location)
|
||||||
|
}
|
||||||
|
|
||||||
|
w = performFrontendFallbackRequest(t, router, "/dashboard/")
|
||||||
|
if w.Code != http.StatusOK || !strings.Contains(w.Body.String(), "admin index") {
|
||||||
|
t.Fatalf("/dashboard/ status = %d body = %q, want admin index", w.Code, w.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
w = performFrontendFallbackRequest(t, router, "/dashboard/assets/app.js")
|
||||||
|
if w.Code != http.StatusOK || !strings.Contains(w.Body.String(), "admin asset") {
|
||||||
|
t.Fatalf("/dashboard/assets/app.js status = %d body = %q, want admin asset", w.Code, w.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,214 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAgentcompatIOStreamStateRoutesRequirePATAndReturnTypedState(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, token := mkToken(t, userID, []string{model.ScopeInventoryRead}, nil)
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
router := gin.New()
|
||||||
|
registerAgentcompatRoutes(router)
|
||||||
|
server := httptest.NewServer(router)
|
||||||
|
t.Cleanup(server.Close)
|
||||||
|
|
||||||
|
unauthenticated, err := http.NewRequest(http.MethodGet, server.URL+"/agentcompat/io-stream-state", nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
response, err := server.Client().Do(unauthenticated)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Equal(t, http.StatusUnauthorized, response.StatusCode)
|
||||||
|
response.Body.Close()
|
||||||
|
|
||||||
|
request, err := http.NewRequest(http.MethodGet, server.URL+"/agentcompat/io-stream-state", nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
request.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
response, err = server.Client().Do(request)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer response.Body.Close()
|
||||||
|
var envelope model.CommonResponse[rpc.IOStreamState]
|
||||||
|
responseBody, readErr := io.ReadAll(response.Body)
|
||||||
|
require.NoError(t, readErr)
|
||||||
|
require.NotContains(t, string(responseBody), "route-wait")
|
||||||
|
require.NoError(t, json.Unmarshal(responseBody, &envelope))
|
||||||
|
require.Equal(t, http.StatusOK, response.StatusCode)
|
||||||
|
require.True(t, envelope.Success)
|
||||||
|
require.Equal(t, rpc.IOStreamState{}, envelope.Data)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatIOStreamStateWaitRouteWakesAfterClose(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, token := mkToken(t, userID, []string{model.ScopeInventoryRead}, nil)
|
||||||
|
require.NoError(t, rpc.NezhaHandlerSingleton.CreateStream("route-wait", 1, 7))
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
router := gin.New()
|
||||||
|
registerAgentcompatRoutes(router)
|
||||||
|
server := httptest.NewServer(router)
|
||||||
|
t.Cleanup(server.Close)
|
||||||
|
body, err := json.Marshal(rpc.IOStreamStateExpectation{ExpectedCount: rpc.ExpectedIOStreamCount(0), AbsentStreamID: "route-wait"})
|
||||||
|
require.NoError(t, err)
|
||||||
|
requestContext, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
request, err := http.NewRequestWithContext(requestContext, http.MethodPost, server.URL+"/agentcompat/io-stream-state", bytes.NewReader(body))
|
||||||
|
require.NoError(t, err)
|
||||||
|
request.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
request.Header.Set("Content-Type", "application/json")
|
||||||
|
result := make(chan *http.Response, 1)
|
||||||
|
go func() {
|
||||||
|
response, requestErr := server.Client().Do(request)
|
||||||
|
if requestErr == nil {
|
||||||
|
result <- response
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
rpc.NezhaHandlerSingleton.CloseStream("route-wait")
|
||||||
|
response := <-result
|
||||||
|
defer response.Body.Close()
|
||||||
|
var envelope model.CommonResponse[rpc.IOStreamState]
|
||||||
|
responseBody, readErr := io.ReadAll(response.Body)
|
||||||
|
require.NoError(t, readErr)
|
||||||
|
require.NotContains(t, string(responseBody), "route-wait")
|
||||||
|
require.NoError(t, json.Unmarshal(responseBody, &envelope))
|
||||||
|
require.True(t, envelope.Success)
|
||||||
|
require.Equal(t, 0, envelope.Data.Count)
|
||||||
|
require.Equal(t, uint64(2), envelope.Data.Generation)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatIOStreamStateCreateWaitRouteWakesAfterCreate(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, token := mkToken(t, userID, []string{model.ScopeInventoryRead}, nil)
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
router := gin.New()
|
||||||
|
waitCaptured := make(chan struct{})
|
||||||
|
var observerOnce sync.Once
|
||||||
|
rpc.NezhaHandlerSingleton.SetIOStreamStateWaitObserverForAgentcompat(func() {
|
||||||
|
observerOnce.Do(func() { close(waitCaptured) })
|
||||||
|
})
|
||||||
|
t.Cleanup(func() {
|
||||||
|
rpc.NezhaHandlerSingleton.SetIOStreamStateWaitObserverForAgentcompat(nil)
|
||||||
|
})
|
||||||
|
registerAgentcompatRoutes(router)
|
||||||
|
server := httptest.NewServer(router)
|
||||||
|
t.Cleanup(server.Close)
|
||||||
|
body := bytes.NewBufferString(`{"expected_count":1}`)
|
||||||
|
requestContext, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
request, err := http.NewRequestWithContext(requestContext, http.MethodPost, server.URL+"/agentcompat/io-stream-state", body)
|
||||||
|
require.NoError(t, err)
|
||||||
|
request.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
request.Header.Set("Content-Type", "application/json")
|
||||||
|
result := make(chan *http.Response, 1)
|
||||||
|
go func() {
|
||||||
|
response, requestErr := server.Client().Do(request)
|
||||||
|
if requestErr == nil {
|
||||||
|
result <- response
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
select {
|
||||||
|
case <-waitCaptured:
|
||||||
|
case <-requestContext.Done():
|
||||||
|
t.Fatal(requestContext.Err())
|
||||||
|
}
|
||||||
|
require.NoError(t, rpc.NezhaHandlerSingleton.CreateStream("route-create", 1, 7))
|
||||||
|
var response *http.Response
|
||||||
|
select {
|
||||||
|
case response = <-result:
|
||||||
|
case <-requestContext.Done():
|
||||||
|
t.Fatal(requestContext.Err())
|
||||||
|
}
|
||||||
|
defer response.Body.Close()
|
||||||
|
responseBody, err := io.ReadAll(response.Body)
|
||||||
|
require.NoError(t, err)
|
||||||
|
var envelope model.CommonResponse[rpc.IOStreamState]
|
||||||
|
require.NoError(t, json.Unmarshal(responseBody, &envelope))
|
||||||
|
require.True(t, envelope.Success)
|
||||||
|
require.Equal(t, 1, envelope.Data.Count)
|
||||||
|
require.Equal(t, uint64(1), envelope.Data.Generation)
|
||||||
|
require.NotContains(t, string(responseBody), "route-create")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatIOStreamStateRejectsInvalidExpectationAndCancellation(t *testing.T) {
|
||||||
|
cleanup, userID := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
_, token := mkToken(t, userID, []string{model.ScopeInventoryRead}, nil)
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
router := gin.New()
|
||||||
|
registerAgentcompatRoutes(router)
|
||||||
|
server := httptest.NewServer(router)
|
||||||
|
t.Cleanup(server.Close)
|
||||||
|
for _, rawBody := range []string{`{}`, `{"expected_count":null}`, `{"expected_count":-1,"absent_stream_id":"private-stream-id"}`} {
|
||||||
|
body := bytes.NewBufferString(rawBody)
|
||||||
|
request, err := http.NewRequest(http.MethodPost, server.URL+"/agentcompat/io-stream-state", body)
|
||||||
|
require.NoError(t, err)
|
||||||
|
request.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
request.Header.Set("Content-Type", "application/json")
|
||||||
|
response, err := server.Client().Do(request)
|
||||||
|
require.NoError(t, err)
|
||||||
|
responseBody, readErr := io.ReadAll(response.Body)
|
||||||
|
response.Body.Close()
|
||||||
|
require.NoError(t, readErr)
|
||||||
|
var envelope model.CommonResponse[rpc.IOStreamState]
|
||||||
|
require.NoError(t, json.Unmarshal(responseBody, &envelope))
|
||||||
|
require.False(t, envelope.Success)
|
||||||
|
require.NotEmpty(t, envelope.Error)
|
||||||
|
require.NotContains(t, string(responseBody), "private-stream-id")
|
||||||
|
}
|
||||||
|
|
||||||
|
body := bytes.NewBufferString(`{"absent_stream_id":"absence-only"}`)
|
||||||
|
request, err := http.NewRequest(http.MethodPost, server.URL+"/agentcompat/io-stream-state", body)
|
||||||
|
require.NoError(t, err)
|
||||||
|
request.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
request.Header.Set("Content-Type", "application/json")
|
||||||
|
response, err := server.Client().Do(request)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer response.Body.Close()
|
||||||
|
var envelope model.CommonResponse[rpc.IOStreamState]
|
||||||
|
require.NoError(t, json.NewDecoder(response.Body).Decode(&envelope))
|
||||||
|
require.True(t, envelope.Success)
|
||||||
|
require.Equal(t, 0, envelope.Data.Count)
|
||||||
|
|
||||||
|
unauthenticatedBody := bytes.NewBufferString(`{"expected_count":0}`)
|
||||||
|
unauthenticatedPost, err := http.NewRequest(http.MethodPost, server.URL+"/agentcompat/io-stream-state", unauthenticatedBody)
|
||||||
|
require.NoError(t, err)
|
||||||
|
unauthenticatedPost.Header.Set("Content-Type", "application/json")
|
||||||
|
unauthenticatedResponse, err := server.Client().Do(unauthenticatedPost)
|
||||||
|
require.NoError(t, err)
|
||||||
|
unauthenticatedResponse.Body.Close()
|
||||||
|
require.Equal(t, http.StatusUnauthorized, unauthenticatedResponse.StatusCode)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
cancel()
|
||||||
|
_, waitErr := rpc.NezhaHandlerSingleton.WaitForIOStreamState(ctx, rpc.IOStreamStateExpectation{ExpectedCount: rpc.ExpectedIOStreamCount(1)})
|
||||||
|
require.True(t, errors.Is(waitErr, context.Canceled))
|
||||||
|
}
|
||||||
+119
-27
@@ -1,6 +1,8 @@
|
|||||||
package controller
|
package controller
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
"net/http"
|
"net/http"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -12,30 +14,87 @@ import (
|
|||||||
|
|
||||||
"github.com/nezhahq/nezha/cmd/dashboard/controller/waf"
|
"github.com/nezhahq/nezha/cmd/dashboard/controller/waf"
|
||||||
"github.com/nezhahq/nezha/model"
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/idcodec"
|
||||||
"github.com/nezhahq/nezha/pkg/utils"
|
"github.com/nezhahq/nezha/pkg/utils"
|
||||||
"github.com/nezhahq/nezha/service/singleton"
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
jwtClaimUserID = "uid"
|
||||||
|
jwtClaimKeyID = "keyId"
|
||||||
|
jwtKeyIDBytes = 32
|
||||||
|
)
|
||||||
|
|
||||||
|
func uaHash(c *gin.Context) string {
|
||||||
|
sum := sha256.Sum256([]byte(c.Request.UserAgent()))
|
||||||
|
return hex.EncodeToString(sum[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
func issueJWTSession(c *gin.Context, user *model.User, jwtTimeoutHours int) (map[string]interface{}, error) {
|
||||||
|
keyID, err := utils.GenerateRandomString(jwtKeyIDBytes)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
// encodedUID is reversible Sqids obfuscation keyed by JWTSecretKey, NOT
|
||||||
|
// a one-way hash. It exists to defeat enumeration on the wire, not to
|
||||||
|
// keep the uid confidential — see L2 note in idcodec docs.
|
||||||
|
encodedUID, err := idcodec.Encode(user.ID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
now := time.Now()
|
||||||
|
sess := model.JWTSession{
|
||||||
|
KeyID: keyID,
|
||||||
|
UserID: user.ID,
|
||||||
|
IP: c.GetString(model.CtxKeyRealIPStr),
|
||||||
|
UAHash: uaHash(c),
|
||||||
|
TokenVersion: user.TokenVersion,
|
||||||
|
ExpiresAt: now.Add(time.Hour * time.Duration(jwtTimeoutHours)),
|
||||||
|
CreatedAt: now,
|
||||||
|
LastUsedAt: now,
|
||||||
|
}
|
||||||
|
if err := singleton.DB.Create(&sess).Error; err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return map[string]interface{}{
|
||||||
|
jwtClaimUserID: encodedUID,
|
||||||
|
jwtClaimKeyID: keyID,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
func initParams() *jwt.GinJWTMiddleware {
|
func initParams() *jwt.GinJWTMiddleware {
|
||||||
return &jwt.GinJWTMiddleware{
|
return &jwt.GinJWTMiddleware{
|
||||||
Realm: singleton.Conf.SiteName,
|
Realm: singleton.Conf.SiteName,
|
||||||
Key: []byte(singleton.Conf.JWTSecretKey),
|
Key: []byte(singleton.Conf.JWTSecretKey),
|
||||||
CookieName: "nz-jwt",
|
CookieName: "nz-jwt",
|
||||||
SendCookie: true,
|
SendCookie: true,
|
||||||
Timeout: time.Hour * time.Duration(singleton.Conf.JWTTimeout),
|
// Pin the signing algorithm so a future library default change (or an
|
||||||
MaxRefresh: time.Hour * time.Duration(singleton.Conf.JWTTimeout),
|
// `alg: none` confusion attempt) cannot weaken token validation.
|
||||||
IdentityKey: model.CtxKeyAuthorizedUser,
|
SigningAlgorithm: "HS256",
|
||||||
PayloadFunc: payloadFunc(),
|
// Lax keeps OAuth callback redirects (top-level GET navigations from
|
||||||
|
// the provider domain) working while blocking cross-site POST CSRF.
|
||||||
|
// HttpOnly/Secure are intentionally left default: the frontend reads
|
||||||
|
// `!!document.cookie` for login-state display and many deployments
|
||||||
|
// terminate TLS at a proxy upstream — both warrant a separate change.
|
||||||
|
CookieSameSite: http.SameSiteLaxMode,
|
||||||
|
Timeout: time.Hour * time.Duration(singleton.Conf.JWTTimeout),
|
||||||
|
MaxRefresh: time.Hour * time.Duration(singleton.Conf.JWTTimeout),
|
||||||
|
IdentityKey: model.CtxKeyAuthorizedUser,
|
||||||
|
PayloadFunc: payloadFunc(),
|
||||||
|
|
||||||
IdentityHandler: identityHandler(),
|
IdentityHandler: identityHandler(),
|
||||||
Authenticator: authenticator(),
|
Authenticator: authenticator(),
|
||||||
Authorizator: authorizator(),
|
Authorizator: authorizator(),
|
||||||
Unauthorized: unauthorized(),
|
Unauthorized: unauthorized(),
|
||||||
TokenLookup: "header: Authorization, query: token, cookie: nz-jwt",
|
// query: token still accepted because the WebSocket browser API
|
||||||
TokenHeadName: "Bearer",
|
// cannot set Authorization headers; removing it would break the
|
||||||
TimeFunc: time.Now,
|
// /ws/* routes until the frontend migrates to cookie auth.
|
||||||
|
TokenLookup: "header: Authorization, query: token, cookie: nz-jwt",
|
||||||
|
TokenHeadName: "Bearer",
|
||||||
|
TimeFunc: time.Now,
|
||||||
|
|
||||||
LoginResponse: func(c *gin.Context, code int, token string, expire time.Time) {
|
LoginResponse: func(c *gin.Context, code int, token string, expire time.Time) {
|
||||||
|
setCSRFCookie(c)
|
||||||
c.JSON(http.StatusOK, model.CommonResponse[model.LoginResponse]{
|
c.JSON(http.StatusOK, model.CommonResponse[model.LoginResponse]{
|
||||||
Success: true,
|
Success: true,
|
||||||
Data: model.LoginResponse{
|
Data: model.LoginResponse{
|
||||||
@@ -61,28 +120,56 @@ func identityHandler() func(c *gin.Context) any {
|
|||||||
return func(c *gin.Context) any {
|
return func(c *gin.Context) any {
|
||||||
claims := jwt.ExtractClaims(c)
|
claims := jwt.ExtractClaims(c)
|
||||||
|
|
||||||
userId, ok := claims["user_id"].(string)
|
keyID, ok := claims[jwtClaimKeyID].(string)
|
||||||
if !ok {
|
if !ok || keyID == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
encodedUID, ok := claims[jwtClaimUserID].(string)
|
||||||
|
if !ok || encodedUID == "" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
claimUID, err := idcodec.Decode(encodedUID)
|
||||||
|
if err != nil {
|
||||||
|
realIP := c.GetString(model.CtxKeyRealIPStr)
|
||||||
|
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
tokenIP, ok := claims["ip"].(string)
|
var sess model.JWTSession
|
||||||
if !ok {
|
if err := singleton.DB.First(&sess, "key_id = ?", keyID).Error; err != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if sess.RevokedAt != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
now := time.Now()
|
||||||
|
if now.After(sess.ExpiresAt) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if claimUID != sess.UserID {
|
||||||
|
realIP := c.GetString(model.CtxKeyRealIPStr)
|
||||||
|
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
currentIP := c.GetString(model.CtxKeyRealIPStr)
|
currentIP := c.GetString(model.CtxKeyRealIPStr)
|
||||||
|
if sess.IP != currentIP {
|
||||||
if tokenIP != currentIP {
|
|
||||||
// IP地址不匹配,token无效
|
|
||||||
c.Set(model.CtxKeyIsIPMismatch, true)
|
c.Set(model.CtxKeyIsIPMismatch, true)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
var user model.User
|
var user model.User
|
||||||
if err := singleton.DB.First(&user, userId).Error; err != nil {
|
if err := singleton.DB.First(&user, sess.UserID).Error; err != nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
if user.TokenVersion != sess.TokenVersion {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = singleton.DB.Model(&model.JWTSession{}).
|
||||||
|
Where("key_id = ?", keyID).
|
||||||
|
Update("last_used_at", now).Error
|
||||||
|
|
||||||
|
c.Set(jwtClaimKeyID, keyID)
|
||||||
return &user
|
return &user
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -106,7 +193,7 @@ func authenticator() func(c *gin.Context) (any, error) {
|
|||||||
var user model.User
|
var user model.User
|
||||||
realip := c.GetString(model.CtxKeyRealIPStr)
|
realip := c.GetString(model.CtxKeyRealIPStr)
|
||||||
|
|
||||||
if err := singleton.DB.Select("id", "password", "reject_password").Where("username = ?", loginVals.Username).First(&user).Error; err != nil {
|
if err := singleton.DB.Select("id", "password", "reject_password", "token_version").Where("username = ?", loginVals.Username).First(&user).Error; err != nil {
|
||||||
if err == gorm.ErrRecordNotFound {
|
if err == gorm.ErrRecordNotFound {
|
||||||
model.BlockIP(singleton.DB, realip, model.WAFBlockReasonTypeLoginFail, model.BlockIDUnknownUser)
|
model.BlockIP(singleton.DB, realip, model.WAFBlockReasonTypeLoginFail, model.BlockIDUnknownUser)
|
||||||
}
|
}
|
||||||
@@ -126,11 +213,7 @@ func authenticator() func(c *gin.Context) (any, error) {
|
|||||||
model.UnblockIP(singleton.DB, realip, model.BlockIDUnknownUser)
|
model.UnblockIP(singleton.DB, realip, model.BlockIDUnknownUser)
|
||||||
model.UnblockIP(singleton.DB, realip, int64(user.ID))
|
model.UnblockIP(singleton.DB, realip, int64(user.ID))
|
||||||
|
|
||||||
// 返回用户ID和IP地址的组合,用于在payloadFunc中设置JWT claims
|
return issueJWTSession(c, &user, singleton.Conf.JWTTimeout)
|
||||||
return map[string]interface{}{
|
|
||||||
"user_id": utils.Itoa(user.ID),
|
|
||||||
"ip": realip,
|
|
||||||
}, nil
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -158,8 +241,17 @@ func unauthorized() func(c *gin.Context, code int, message string) {
|
|||||||
// @Tags auth required
|
// @Tags auth required
|
||||||
// @Produce json
|
// @Produce json
|
||||||
// @Success 200 {object} model.CommonResponse[model.LoginResponse]
|
// @Success 200 {object} model.CommonResponse[model.LoginResponse]
|
||||||
// @Router /refresh-token [get]
|
// @Router /refresh-token [post]
|
||||||
func refreshResponse(c *gin.Context, code int, token string, expire time.Time) {
|
func refreshResponse(c *gin.Context, code int, token string, expire time.Time) {
|
||||||
|
if keyID := c.GetString(jwtClaimKeyID); keyID != "" {
|
||||||
|
_ = singleton.DB.Model(&model.JWTSession{}).
|
||||||
|
Where("key_id = ?", keyID).
|
||||||
|
Updates(map[string]interface{}{
|
||||||
|
"expires_at": expire,
|
||||||
|
"last_used_at": time.Now(),
|
||||||
|
}).Error
|
||||||
|
}
|
||||||
|
setCSRFCookie(c)
|
||||||
c.JSON(http.StatusOK, model.CommonResponse[model.LoginResponse]{
|
c.JSON(http.StatusOK, model.CommonResponse[model.LoginResponse]{
|
||||||
Success: true,
|
Success: true,
|
||||||
Data: model.LoginResponse{
|
Data: model.LoginResponse{
|
||||||
|
|||||||
@@ -0,0 +1,366 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
jwt "github.com/appleboy/gin-jwt/v2"
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"golang.org/x/crypto/bcrypt"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/idcodec"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
const jwtSessionTestMasterKey = "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef"
|
||||||
|
|
||||||
|
func setupJWTSessionTest(t *testing.T) (cleanup func()) {
|
||||||
|
t.Helper()
|
||||||
|
require.NoError(t, idcodec.Init([]byte(jwtSessionTestMasterKey)))
|
||||||
|
|
||||||
|
originalDB := singleton.DB
|
||||||
|
originalConf := singleton.Conf
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
sqlDB, err := db.DB()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, db.AutoMigrate(&model.User{}, &model.JWTSession{}, &model.WAF{}))
|
||||||
|
singleton.DB = db
|
||||||
|
singleton.Conf = &singleton.ConfigClass{Config: &model.Config{JWTTimeout: 1}}
|
||||||
|
|
||||||
|
require.NoError(t, db.Create(&model.User{
|
||||||
|
Common: model.Common{ID: 100},
|
||||||
|
Username: "victim",
|
||||||
|
Role: model.RoleMember,
|
||||||
|
TokenVersion: 7,
|
||||||
|
}).Error)
|
||||||
|
|
||||||
|
return func() {
|
||||||
|
_ = sqlDB.Close()
|
||||||
|
singleton.DB = originalDB
|
||||||
|
singleton.Conf = originalConf
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newCtxForUser(userID uint64, ip, ua string) *gin.Context {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||||
|
ctx.Request = httptest.NewRequest("GET", "/", nil)
|
||||||
|
ctx.Request.Header.Set("User-Agent", ua)
|
||||||
|
ctx.Set(model.CtxKeyRealIPStr, ip)
|
||||||
|
if userID != 0 {
|
||||||
|
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: userID}})
|
||||||
|
}
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIssueJWTSessionWritesRow(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
ctx := newCtxForUser(0, "1.2.3.4", "test-ua")
|
||||||
|
user := model.User{Common: model.Common{ID: 100}, TokenVersion: 7}
|
||||||
|
claims, err := issueJWTSession(ctx, &user, 1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
hashUID, _ := claims[jwtClaimUserID].(string)
|
||||||
|
keyID, _ := claims[jwtClaimKeyID].(string)
|
||||||
|
assert.NotEqual(t, "100", hashUID, "uid claim must be obfuscated, not raw integer")
|
||||||
|
got, err := idcodec.Decode(hashUID)
|
||||||
|
require.NoError(t, err)
|
||||||
|
assert.Equal(t, uint64(100), got)
|
||||||
|
|
||||||
|
var sess model.JWTSession
|
||||||
|
require.NoError(t, singleton.DB.First(&sess, "key_id = ?", keyID).Error)
|
||||||
|
assert.Equal(t, uint64(100), sess.UserID)
|
||||||
|
assert.Equal(t, "1.2.3.4", sess.IP)
|
||||||
|
assert.Equal(t, uint64(7), sess.TokenVersion)
|
||||||
|
assert.True(t, sess.ExpiresAt.After(time.Now()))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIdentityHandlerHappyPath(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
ctx := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
user := model.User{Common: model.Common{ID: 100}, TokenVersion: 7}
|
||||||
|
claims, err := issueJWTSession(ctx, &user, 1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
verify := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
verify.Set("JWT_PAYLOAD", jwt.MapClaims{
|
||||||
|
jwtClaimUserID: claims[jwtClaimUserID],
|
||||||
|
jwtClaimKeyID: claims[jwtClaimKeyID],
|
||||||
|
})
|
||||||
|
|
||||||
|
identity := identityHandler()(verify)
|
||||||
|
require.NotNil(t, identity, "happy path must return user identity")
|
||||||
|
u := identity.(*model.User)
|
||||||
|
assert.Equal(t, uint64(100), u.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIdentityHandlerRejectsMismatchedClaimUID(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
ctx := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
user := model.User{Common: model.Common{ID: 100}, TokenVersion: 7}
|
||||||
|
claims, err := issueJWTSession(ctx, &user, 1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
forgedUID, err := idcodec.Encode(999)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
verify := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
verify.Set("JWT_PAYLOAD", jwt.MapClaims{
|
||||||
|
jwtClaimUserID: forgedUID,
|
||||||
|
jwtClaimKeyID: claims[jwtClaimKeyID],
|
||||||
|
})
|
||||||
|
|
||||||
|
identity := identityHandler()(verify)
|
||||||
|
assert.Nil(t, identity, "claim uid not matching session.user_id must reject")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIdentityHandlerRejectsRevokedSession(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
ctx := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
user := model.User{Common: model.Common{ID: 100}, TokenVersion: 7}
|
||||||
|
claims, err := issueJWTSession(ctx, &user, 1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
keyID := claims[jwtClaimKeyID].(string)
|
||||||
|
require.NoError(t, singleton.RevokeJWTSession(keyID))
|
||||||
|
|
||||||
|
verify := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
verify.Set("JWT_PAYLOAD", jwt.MapClaims{
|
||||||
|
jwtClaimUserID: claims[jwtClaimUserID],
|
||||||
|
jwtClaimKeyID: claims[jwtClaimKeyID],
|
||||||
|
})
|
||||||
|
|
||||||
|
identity := identityHandler()(verify)
|
||||||
|
assert.Nil(t, identity, "revoked session must reject")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIdentityHandlerRejectsTokenVersionBump(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
ctx := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
user := model.User{Common: model.Common{ID: 100}, TokenVersion: 7}
|
||||||
|
claims, err := issueJWTSession(ctx, &user, 1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
require.NoError(t, singleton.DB.Model(&model.User{}).
|
||||||
|
Where("id = ?", 100).
|
||||||
|
Update("token_version", 8).Error)
|
||||||
|
|
||||||
|
verify := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
verify.Set("JWT_PAYLOAD", jwt.MapClaims{
|
||||||
|
jwtClaimUserID: claims[jwtClaimUserID],
|
||||||
|
jwtClaimKeyID: claims[jwtClaimKeyID],
|
||||||
|
})
|
||||||
|
|
||||||
|
identity := identityHandler()(verify)
|
||||||
|
assert.Nil(t, identity, "session whose TokenVersion is stale must reject")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIdentityHandlerFlagsIPMismatch(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
ctx := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
user := model.User{Common: model.Common{ID: 100}, TokenVersion: 7}
|
||||||
|
claims, err := issueJWTSession(ctx, &user, 1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
verify := newCtxForUser(0, "9.9.9.9", "ua")
|
||||||
|
verify.Set("JWT_PAYLOAD", jwt.MapClaims{
|
||||||
|
jwtClaimUserID: claims[jwtClaimUserID],
|
||||||
|
jwtClaimKeyID: claims[jwtClaimKeyID],
|
||||||
|
})
|
||||||
|
|
||||||
|
identity := identityHandler()(verify)
|
||||||
|
assert.Nil(t, identity, "IP mismatch must reject")
|
||||||
|
assert.True(t, verify.GetBool(model.CtxKeyIsIPMismatch))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIdentityHandlerRejectsUnknownKeyID(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
hashUID, err := idcodec.Encode(100)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
verify := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
verify.Set("JWT_PAYLOAD", jwt.MapClaims{
|
||||||
|
jwtClaimUserID: hashUID,
|
||||||
|
jwtClaimKeyID: "this-key-id-was-never-issued",
|
||||||
|
})
|
||||||
|
|
||||||
|
identity := identityHandler()(verify)
|
||||||
|
assert.Nil(t, identity, "key id absent from sessions table must reject (no oracle to confirm secret)")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthenticatorPersistsCurrentTokenVersion(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
pw, err := bcrypt.GenerateFromPassword([]byte("correct horse"), bcrypt.MinCost)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, singleton.DB.Model(&model.User{}).
|
||||||
|
Where("id = ?", 100).
|
||||||
|
Update("password", string(pw)).Error)
|
||||||
|
|
||||||
|
ctx := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
body, _ := json.Marshal(model.LoginRequest{Username: "victim", Password: "correct horse"})
|
||||||
|
ctx.Request = httptest.NewRequest("POST", "/api/v1/login", bytes.NewReader(body))
|
||||||
|
ctx.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
ctx.Request.Header.Set("User-Agent", "ua")
|
||||||
|
ctx.Set(model.CtxKeyRealIPStr, "1.2.3.4")
|
||||||
|
|
||||||
|
data, err := authenticator()(ctx)
|
||||||
|
require.NoError(t, err)
|
||||||
|
claims, ok := data.(map[string]interface{})
|
||||||
|
require.True(t, ok, "authenticator must return claims map")
|
||||||
|
keyID, _ := claims[jwtClaimKeyID].(string)
|
||||||
|
require.NotEmpty(t, keyID)
|
||||||
|
|
||||||
|
var sess model.JWTSession
|
||||||
|
require.NoError(t, singleton.DB.First(&sess, "key_id = ?", keyID).Error)
|
||||||
|
assert.Equal(t, uint64(7), sess.TokenVersion,
|
||||||
|
"session must record the user's current token_version, otherwise identityHandler will reject the freshly-issued token")
|
||||||
|
|
||||||
|
verify := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
verify.Set("JWT_PAYLOAD", jwt.MapClaims{
|
||||||
|
jwtClaimUserID: claims[jwtClaimUserID],
|
||||||
|
jwtClaimKeyID: claims[jwtClaimKeyID],
|
||||||
|
})
|
||||||
|
assert.NotNil(t, identityHandler()(verify),
|
||||||
|
"the very next request with the freshly-issued token must authenticate")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthenticator_BadPasswordReturnsFailedAuth(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
pw, err := bcrypt.GenerateFromPassword([]byte("correct horse"), bcrypt.MinCost)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, singleton.DB.Model(&model.User{}).
|
||||||
|
Where("id = ?", 100).
|
||||||
|
Update("password", string(pw)).Error)
|
||||||
|
|
||||||
|
ctx := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
body, _ := json.Marshal(model.LoginRequest{Username: "victim", Password: "wrong"})
|
||||||
|
ctx.Request = httptest.NewRequest("POST", "/api/v1/login", bytes.NewReader(body))
|
||||||
|
ctx.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
ctx.Request.Header.Set("User-Agent", "ua")
|
||||||
|
ctx.Set(model.CtxKeyRealIPStr, "1.2.3.4")
|
||||||
|
|
||||||
|
_, err = authenticator()(ctx)
|
||||||
|
require.Error(t, err, "wrong password must fail authentication")
|
||||||
|
require.Equal(t, jwt.ErrFailedAuthentication, err)
|
||||||
|
|
||||||
|
var w model.WAF
|
||||||
|
require.NoError(t, singleton.DB.Where("block_identifier = ?", int64(100)).First(&w).Error,
|
||||||
|
"bad password must increment WAF counter under user-specific BlockID")
|
||||||
|
require.GreaterOrEqual(t, w.Count, uint64(1))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthenticator_UnknownUserReturnsFailedAuth(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
ctx := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
body, _ := json.Marshal(model.LoginRequest{Username: "ghost", Password: "anything"})
|
||||||
|
ctx.Request = httptest.NewRequest("POST", "/api/v1/login", bytes.NewReader(body))
|
||||||
|
ctx.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
ctx.Set(model.CtxKeyRealIPStr, "1.2.3.4")
|
||||||
|
|
||||||
|
_, err := authenticator()(ctx)
|
||||||
|
require.Error(t, err)
|
||||||
|
require.Equal(t, jwt.ErrFailedAuthentication, err)
|
||||||
|
|
||||||
|
var w model.WAF
|
||||||
|
require.NoError(t, singleton.DB.Where("block_identifier = ?", int64(model.BlockIDUnknownUser)).First(&w).Error,
|
||||||
|
"unknown user must increment WAF counter under BlockIDUnknownUser")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthenticator_RejectPasswordUserRefused(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
pw, err := bcrypt.GenerateFromPassword([]byte("ok"), bcrypt.MinCost)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, singleton.DB.Model(&model.User{}).
|
||||||
|
Where("id = ?", 100).
|
||||||
|
Updates(map[string]any{"password": string(pw), "reject_password": true}).Error)
|
||||||
|
|
||||||
|
ctx := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
body, _ := json.Marshal(model.LoginRequest{Username: "victim", Password: "ok"})
|
||||||
|
ctx.Request = httptest.NewRequest("POST", "/api/v1/login", bytes.NewReader(body))
|
||||||
|
ctx.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
ctx.Set(model.CtxKeyRealIPStr, "1.2.3.4")
|
||||||
|
|
||||||
|
_, err = authenticator()(ctx)
|
||||||
|
require.Equal(t, jwt.ErrFailedAuthentication, err,
|
||||||
|
"users with reject_password=true must not be able to log in via password even with correct one")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestIdentityHandler_ExpiredSessionRejected(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
ctx := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
user := model.User{Common: model.Common{ID: 100}, TokenVersion: 7}
|
||||||
|
claims, err := issueJWTSession(ctx, &user, 1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
keyID := claims[jwtClaimKeyID].(string)
|
||||||
|
|
||||||
|
require.NoError(t, singleton.DB.Model(&model.JWTSession{}).
|
||||||
|
Where("key_id = ?", keyID).
|
||||||
|
Update("expires_at", time.Now().Add(-time.Hour)).Error)
|
||||||
|
|
||||||
|
verify := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
verify.Set("JWT_PAYLOAD", jwt.MapClaims{
|
||||||
|
jwtClaimUserID: claims[jwtClaimUserID],
|
||||||
|
jwtClaimKeyID: claims[jwtClaimKeyID],
|
||||||
|
})
|
||||||
|
|
||||||
|
identity := identityHandler()(verify)
|
||||||
|
require.Nil(t, identity, "session whose expires_at is in the past must reject")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRefreshResponse_UpdatesSessionExpires(t *testing.T) {
|
||||||
|
cleanup := setupJWTSessionTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
ctx := newCtxForUser(0, "1.2.3.4", "ua")
|
||||||
|
user := model.User{Common: model.Common{ID: 100}, TokenVersion: 7}
|
||||||
|
claims, err := issueJWTSession(ctx, &user, 1)
|
||||||
|
require.NoError(t, err)
|
||||||
|
keyID := claims[jwtClaimKeyID].(string)
|
||||||
|
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/api/v1/refresh-token", nil)
|
||||||
|
c.Set(jwtClaimKeyID, keyID)
|
||||||
|
|
||||||
|
newExpire := time.Now().Add(2 * time.Hour).Truncate(time.Second)
|
||||||
|
refreshResponse(c, 200, "fake-token", newExpire)
|
||||||
|
|
||||||
|
var sess model.JWTSession
|
||||||
|
require.NoError(t, singleton.DB.First(&sess, "key_id = ?", keyID).Error)
|
||||||
|
require.WithinDuration(t, newExpire, sess.ExpiresAt, time.Second,
|
||||||
|
"refreshResponse must extend the session's expires_at to the new expiry")
|
||||||
|
require.WithinDuration(t, time.Now(), sess.LastUsedAt, 5*time.Second,
|
||||||
|
"refreshResponse must touch last_used_at")
|
||||||
|
}
|
||||||
@@ -1,12 +1,18 @@
|
|||||||
package controller
|
package controller
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"net/http/httptest"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
jwt "github.com/appleboy/gin-jwt/v2"
|
jwt "github.com/appleboy/gin-jwt/v2"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/i18n"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestPayloadFunc(t *testing.T) {
|
func TestPayloadFunc(t *testing.T) {
|
||||||
@@ -77,3 +83,120 @@ func TestIPBinding(t *testing.T) {
|
|||||||
assert.Nil(t, claims["ip"])
|
assert.Nil(t, claims["ip"])
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestValidateRuleRejectsForeignTriggerTasks(t *testing.T) {
|
||||||
|
ctx := newMemberValidationContext(t)
|
||||||
|
|
||||||
|
alertRule := &model.AlertRule{
|
||||||
|
Common: model.Common{UserID: 200},
|
||||||
|
Name: "member alert",
|
||||||
|
Rules: []*model.Rule{{Type: "offline", Duration: 3}},
|
||||||
|
FailTriggerTasks: []uint64{42},
|
||||||
|
RecoverTriggerTasks: []uint64{42},
|
||||||
|
}
|
||||||
|
|
||||||
|
assert.Error(t, validateRule(ctx, alertRule))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidateServersRejectsForeignTriggerTasks(t *testing.T) {
|
||||||
|
ctx := newMemberValidationContext(t)
|
||||||
|
|
||||||
|
service := &model.Service{
|
||||||
|
Common: model.Common{UserID: 200},
|
||||||
|
Name: "member service",
|
||||||
|
EnableTriggerTask: true,
|
||||||
|
FailTriggerTasks: []uint64{42},
|
||||||
|
RecoverTriggerTasks: []uint64{42},
|
||||||
|
SkipServers: map[uint64]bool{},
|
||||||
|
}
|
||||||
|
|
||||||
|
assert.Error(t, validateServers(ctx, service))
|
||||||
|
}
|
||||||
|
|
||||||
|
func newMemberValidationContext(t *testing.T) *gin.Context {
|
||||||
|
t.Helper()
|
||||||
|
return newValidationContext(t, 200, model.RoleMember)
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAdminValidationContext(t *testing.T) *gin.Context {
|
||||||
|
t.Helper()
|
||||||
|
return newValidationContext(t, 1, model.RoleAdmin)
|
||||||
|
}
|
||||||
|
|
||||||
|
func newValidationContext(t *testing.T, userID uint64, role model.Role) *gin.Context {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
originalDB := singleton.DB
|
||||||
|
originalLoc := singleton.Loc
|
||||||
|
originalLocalizer := singleton.Localizer
|
||||||
|
originalCronShared := singleton.CronShared
|
||||||
|
originalServerShared := singleton.ServerShared
|
||||||
|
originalUserInfo := singleton.UserInfoMap
|
||||||
|
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
assert.NoError(t, err)
|
||||||
|
sqlDB, err := db.DB()
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.NoError(t, db.AutoMigrate(
|
||||||
|
&model.Cron{},
|
||||||
|
&model.Server{},
|
||||||
|
&model.NotificationGroup{},
|
||||||
|
&model.NotificationGroupNotification{},
|
||||||
|
&model.ServerGroup{},
|
||||||
|
&model.ServerGroupServer{},
|
||||||
|
))
|
||||||
|
assert.NoError(t, db.Create(&model.Cron{
|
||||||
|
Common: model.Common{ID: 42, UserID: 1},
|
||||||
|
Name: "foreign trigger task",
|
||||||
|
Command: "admin-maintenance",
|
||||||
|
TaskType: model.CronTypeTriggerTask,
|
||||||
|
Cover: model.CronCoverAlertTrigger,
|
||||||
|
}).Error)
|
||||||
|
assert.NoError(t, db.Create(&model.Cron{
|
||||||
|
Common: model.Common{ID: 43, UserID: 200},
|
||||||
|
Name: "member trigger task",
|
||||||
|
Command: "member-task",
|
||||||
|
TaskType: model.CronTypeTriggerTask,
|
||||||
|
Cover: model.CronCoverAlertTrigger,
|
||||||
|
}).Error)
|
||||||
|
assert.NoError(t, db.Create(&model.NotificationGroup{
|
||||||
|
Common: model.Common{ID: 7, UserID: 1},
|
||||||
|
Name: "admin group",
|
||||||
|
}).Error)
|
||||||
|
assert.NoError(t, db.Create(&model.NotificationGroup{
|
||||||
|
Common: model.Common{ID: 8, UserID: 200},
|
||||||
|
Name: "member group",
|
||||||
|
}).Error)
|
||||||
|
|
||||||
|
singleton.DB = db
|
||||||
|
singleton.Loc = time.Local
|
||||||
|
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
|
||||||
|
singleton.ServerShared = singleton.NewServerClass()
|
||||||
|
singleton.CronShared = singleton.NewCronClass()
|
||||||
|
singleton.UserLock.Lock()
|
||||||
|
singleton.UserInfoMap = map[uint64]model.UserInfo{
|
||||||
|
1: {Role: model.RoleAdmin},
|
||||||
|
200: {Role: model.RoleMember},
|
||||||
|
}
|
||||||
|
singleton.UserLock.Unlock()
|
||||||
|
t.Cleanup(func() {
|
||||||
|
singleton.CronShared.Close()
|
||||||
|
_ = sqlDB.Close()
|
||||||
|
singleton.DB = originalDB
|
||||||
|
singleton.Loc = originalLoc
|
||||||
|
singleton.Localizer = originalLocalizer
|
||||||
|
singleton.CronShared = originalCronShared
|
||||||
|
singleton.ServerShared = originalServerShared
|
||||||
|
singleton.UserLock.Lock()
|
||||||
|
singleton.UserInfoMap = originalUserInfo
|
||||||
|
singleton.UserLock.Unlock()
|
||||||
|
})
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||||
|
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{
|
||||||
|
Common: model.Common{ID: userID},
|
||||||
|
Role: role,
|
||||||
|
})
|
||||||
|
return ctx
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,505 @@
|
|||||||
|
// Package controller — MCP (Model Context Protocol) server.
|
||||||
|
//
|
||||||
|
// 落地约束:
|
||||||
|
// - 仅支持 Streamable HTTP transport 的 POST 半边(请求-响应、无 SSE)。
|
||||||
|
// 首版面向 LLM 工具调用,不需要 server→client 主动推送。后续要做 GET SSE
|
||||||
|
// 长连接(resource subscription)时再补;客户端兼容 fallback 到普通 POST。
|
||||||
|
// - JSON-RPC 2.0 编解码内嵌于本文件,未引入第三方 MCP SDK:MCP 协议表面足够小
|
||||||
|
// (initialize / tools/list / tools/call),自实现可控、零额外依赖。
|
||||||
|
// - 双层鉴权:闸 1(用户对 server 的所有权)由各 tool handler 调
|
||||||
|
// singleton.ServerShared.Get + Server.HasPermission;闸 2(PAT scope)由
|
||||||
|
// mcpTool.RequiredScope 在 dispatch 之前过滤。
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// --- JSON-RPC 2.0 wire types ---
|
||||||
|
|
||||||
|
type jsonRPCRequest struct {
|
||||||
|
JSONRPC string `json:"jsonrpc"`
|
||||||
|
ID json.RawMessage `json:"id,omitempty"`
|
||||||
|
Method string `json:"method"`
|
||||||
|
Params json.RawMessage `json:"params,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type jsonRPCResponse struct {
|
||||||
|
JSONRPC string `json:"jsonrpc"`
|
||||||
|
ID json.RawMessage `json:"id,omitempty"`
|
||||||
|
Result any `json:"result,omitempty"`
|
||||||
|
Error *jsonRPCError `json:"error,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type jsonRPCError struct {
|
||||||
|
Code int `json:"code"`
|
||||||
|
Message string `json:"message"`
|
||||||
|
Data any `json:"data,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
const (
|
||||||
|
// JSON-RPC 标准错误码
|
||||||
|
rpcErrParse = -32700
|
||||||
|
rpcErrInvalidRequest = -32600
|
||||||
|
rpcErrMethodNotFound = -32601
|
||||||
|
rpcErrInvalidParams = -32602
|
||||||
|
rpcErrInternal = -32603
|
||||||
|
// MCP 自定义错误码(>= -32000 高位段)
|
||||||
|
rpcErrUnauthorized = -32001
|
||||||
|
rpcErrForbidden = -32002
|
||||||
|
)
|
||||||
|
|
||||||
|
// mcpJSONRPCMaxBodyBytes caps the JSON-RPC envelope size at the dashboard
|
||||||
|
// edge. Real fs.write base64 content goes through fs.transfer (capped
|
||||||
|
// separately by model.MCPFsTransferMaxSize) so tools/call params here are
|
||||||
|
// always small. The cap is intentionally generous (8 MiB) to allow
|
||||||
|
// per-request batched arguments while making OOM-via-decode impossible.
|
||||||
|
const mcpJSONRPCMaxBodyBytes = 8 * 1024 * 1024
|
||||||
|
|
||||||
|
// --- MCP types ---
|
||||||
|
|
||||||
|
// mcpServerInfo MCP initialize 响应的 serverInfo 字段。
|
||||||
|
type mcpServerInfo struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Version string `json:"version"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type mcpInitializeResult struct {
|
||||||
|
ProtocolVersion string `json:"protocolVersion"`
|
||||||
|
Capabilities map[string]any `json:"capabilities"`
|
||||||
|
ServerInfo mcpServerInfo `json:"serverInfo"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// mcpToolDescriptor 是 tools/list 返回的单条 tool 描述。
|
||||||
|
type mcpToolDescriptor struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
InputSchema map[string]any `json:"inputSchema"`
|
||||||
|
OutputSchema map[string]any `json:"outputSchema,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// mcpToolsListResult tools/list 响应。
|
||||||
|
type mcpToolsListResult struct {
|
||||||
|
Tools []mcpToolDescriptor `json:"tools"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// mcpContent 是 tools/call 响应里 content[] 的元素。
|
||||||
|
// 仅实现 text 类型;嵌入对象的结构化数据放在外层 structuredContent。
|
||||||
|
type mcpContent struct {
|
||||||
|
Type string `json:"type"`
|
||||||
|
Text string `json:"text,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// mcpToolCallResult tools/call 响应。
|
||||||
|
type mcpToolCallResult struct {
|
||||||
|
Content []mcpContent `json:"content"`
|
||||||
|
StructuredContent any `json:"structuredContent,omitempty"`
|
||||||
|
IsError bool `json:"isError,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- tool 注册框架 ---
|
||||||
|
|
||||||
|
// mcpToolHandler 实际业务逻辑:拿到 raw params + gin ctx,返回任意可序列化结构。
|
||||||
|
type mcpToolHandler func(c *gin.Context, params json.RawMessage) (any, error)
|
||||||
|
|
||||||
|
// mcpTool 是注册表里的单元:声明 + scope 要求 + 处理函数。
|
||||||
|
type mcpTool struct {
|
||||||
|
Name string
|
||||||
|
Description string
|
||||||
|
InputSchema map[string]any
|
||||||
|
OutputSchema map[string]any // 可选;声明 structuredContent 形状,供严格客户端校验
|
||||||
|
RequiredScope string // 闸 2 入口;空字符串 = 任意 PAT 都能调(如 meta.whoami)
|
||||||
|
Handler mcpToolHandler
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
mcpToolsMu sync.RWMutex
|
||||||
|
mcpTools = map[string]*mcpTool{}
|
||||||
|
)
|
||||||
|
|
||||||
|
// registerMCPTool 把一个 tool 加进全局注册表。建议各 tool 文件在 init() 里调用。
|
||||||
|
func registerMCPTool(t *mcpTool) {
|
||||||
|
if t == nil || t.Name == "" || t.Handler == nil {
|
||||||
|
panic("registerMCPTool: invalid tool")
|
||||||
|
}
|
||||||
|
mcpToolsMu.Lock()
|
||||||
|
defer mcpToolsMu.Unlock()
|
||||||
|
if _, dup := mcpTools[t.Name]; dup {
|
||||||
|
panic("registerMCPTool: duplicate name " + t.Name)
|
||||||
|
}
|
||||||
|
mcpTools[t.Name] = t
|
||||||
|
}
|
||||||
|
|
||||||
|
// listRegisteredMCPTools 拷贝一份当前注册表(按名字稳定排序逻辑放在调用方)。
|
||||||
|
func listRegisteredMCPTools() []*mcpTool {
|
||||||
|
mcpToolsMu.RLock()
|
||||||
|
defer mcpToolsMu.RUnlock()
|
||||||
|
out := make([]*mcpTool, 0, len(mcpTools))
|
||||||
|
for _, t := range mcpTools {
|
||||||
|
out = append(out, t)
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- 入口 handler ---
|
||||||
|
|
||||||
|
// mcpEndpoint 处理 POST /mcp。
|
||||||
|
// 鉴权:上游 apiTokenAuthMiddleware 已经把 PAT 解析到 CtxKeyAuthorizedUser,
|
||||||
|
// 此处只要确认有 PAT 即可(不接受裸 JWT,避免浏览器误触)。
|
||||||
|
func mcpEndpoint(c *gin.Context) {
|
||||||
|
if singleton.Conf == nil || !singleton.Conf.MCPEnabled() {
|
||||||
|
writeJSONRPCError(c, nil, rpcErrForbidden, "MCP is disabled by the dashboard administrator")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
tok := APITokenFromContext(c)
|
||||||
|
if tok == nil {
|
||||||
|
// 同时返回 HTTP 401 + JSON-RPC error:标准 MCP HTTP client 依赖
|
||||||
|
// HTTP 401 触发 auth 重试/OAuth discovery;JSON-RPC body 保留旧字段
|
||||||
|
// 不打破 ScopeDenied 类内部断言。
|
||||||
|
writeJSONRPCErrorWithStatus(c, nil, rpcErrUnauthorized, "missing or invalid API token", http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// MaxBytesReader 必须夹在 PAT 校验通过后、ShouldBindJSON 之前——
|
||||||
|
// 校验前限流可能让攻击者用伪造 token 触发 audit;校验后限流既挡住合法
|
||||||
|
// PAT 的 OOM,又不会让匿名请求走到 audit 路径。
|
||||||
|
if c.Request != nil && c.Request.Body != nil {
|
||||||
|
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, mcpJSONRPCMaxBodyBytes)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Consume the per-token budget before validating the request so malformed
|
||||||
|
// envelopes and malformed tools/call params cannot flood the dashboard
|
||||||
|
// without counting against the limiter. The outcome is applied after the
|
||||||
|
// method is known so tools/call still surfaces the rate limit as a tool
|
||||||
|
// error rather than a transport-level error.
|
||||||
|
rateLimited := !mcpRateLimiterShared.Allow(tok.ID)
|
||||||
|
|
||||||
|
var req jsonRPCRequest
|
||||||
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
|
if errors.Is(err, errors.New("http: request body too large")) || strings.Contains(err.Error(), "http: request body too large") {
|
||||||
|
writeJSONRPCErrorWithStatus(c, nil, rpcErrInvalidRequest, "request body exceeds MCP envelope size limit", http.StatusRequestEntityTooLarge)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// 限流优先:method 无从得知时,over-budget 请求即便 body 畸形也必须
|
||||||
|
// 走 429,否则攻击者能用畸形 body 在不计入限额的情况下持续刷 parse error。
|
||||||
|
if rateLimited {
|
||||||
|
writeJSONRPCErrorWithStatus(c, nil, rpcErrForbidden, "rate limit exceeded for this token", http.StatusTooManyRequests)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
writeJSONRPCError(c, nil, rpcErrParse, "invalid json-rpc envelope: "+err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if req.JSONRPC != "2.0" || req.Method == "" {
|
||||||
|
if rateLimited {
|
||||||
|
writeJSONRPCErrorWithStatus(c, req.ID, rpcErrForbidden, "rate limit exceeded for this token", http.StatusTooManyRequests)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
writeJSONRPCError(c, req.ID, rpcErrInvalidRequest, "invalid json-rpc envelope")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if rateLimited {
|
||||||
|
if req.Method == "tools/call" {
|
||||||
|
writeToolCallError(c, req.ID, model.MCPOutcomeRateLimited, "rate limit exceeded for this token")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
writeJSONRPCErrorWithStatus(c, req.ID, rpcErrForbidden, "rate limit exceeded for this token", http.StatusTooManyRequests)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
switch req.Method {
|
||||||
|
case "initialize":
|
||||||
|
writeJSONRPCResult(c, req.ID, mcpInitializeResult{
|
||||||
|
ProtocolVersion: "2024-11-05",
|
||||||
|
Capabilities: map[string]any{
|
||||||
|
"tools": map[string]any{"listChanged": false},
|
||||||
|
},
|
||||||
|
ServerInfo: mcpServerInfo{
|
||||||
|
Name: "nezha-mcp",
|
||||||
|
Version: singleton.Version,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
case "notifications/initialized", "ping":
|
||||||
|
// 客户端通知或心跳;JSON-RPC 通知没有 id,但 ping 有 id 时返回空 result
|
||||||
|
if len(req.ID) > 0 && string(req.ID) != "null" {
|
||||||
|
writeJSONRPCResult(c, req.ID, struct{}{})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.Status(http.StatusAccepted)
|
||||||
|
case "tools/list":
|
||||||
|
writeJSONRPCResult(c, req.ID, mcpToolsListResult{
|
||||||
|
Tools: buildToolDescriptors(),
|
||||||
|
})
|
||||||
|
case "tools/call":
|
||||||
|
handleToolsCall(c, &req, tok)
|
||||||
|
default:
|
||||||
|
writeJSONRPCError(c, req.ID, rpcErrMethodNotFound, "method not supported: "+req.Method)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func buildToolDescriptors() []mcpToolDescriptor {
|
||||||
|
tools := listRegisteredMCPTools()
|
||||||
|
out := make([]mcpToolDescriptor, 0, len(tools))
|
||||||
|
for _, t := range tools {
|
||||||
|
out = append(out, mcpToolDescriptor{
|
||||||
|
Name: t.Name,
|
||||||
|
Description: t.Description,
|
||||||
|
InputSchema: t.InputSchema,
|
||||||
|
OutputSchema: t.OutputSchema,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
// toolCallParams 是 tools/call 的 params 结构。
|
||||||
|
type toolCallParams struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Arguments json.RawMessage `json:"arguments,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleToolsCall(c *gin.Context, req *jsonRPCRequest, tok *model.APIToken) {
|
||||||
|
var p toolCallParams
|
||||||
|
if len(req.Params) > 0 {
|
||||||
|
if err := json.Unmarshal(req.Params, &p); err != nil {
|
||||||
|
writeJSONRPCError(c, req.ID, rpcErrInvalidParams, "invalid arguments: "+err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if p.Name == "" {
|
||||||
|
writeJSONRPCError(c, req.ID, rpcErrInvalidParams, "tool name required")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
mcpToolsMu.RLock()
|
||||||
|
tool, ok := mcpTools[p.Name]
|
||||||
|
mcpToolsMu.RUnlock()
|
||||||
|
if !ok {
|
||||||
|
writeJSONRPCError(c, req.ID, rpcErrMethodNotFound, "unknown tool: "+p.Name)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
uid := uint64(0)
|
||||||
|
if u, ok := c.Get(model.CtxKeyAuthorizedUser); ok {
|
||||||
|
if user, ok := u.(*model.User); ok && user != nil {
|
||||||
|
uid = user.ID
|
||||||
|
}
|
||||||
|
}
|
||||||
|
startedAt := time.Now()
|
||||||
|
audit := model.MCPAuditLog{
|
||||||
|
UserID: uid,
|
||||||
|
TokenID: tok.ID,
|
||||||
|
Tool: p.Name,
|
||||||
|
IP: c.GetString(model.CtxKeyRealIPStr),
|
||||||
|
}
|
||||||
|
|
||||||
|
finish := func(outcome, errCode, errMsg string, result any) {
|
||||||
|
if outcome == model.MCPOutcomeOK {
|
||||||
|
textPayload, err := marshalMCPToolResult(result)
|
||||||
|
if err != nil {
|
||||||
|
outcome = model.MCPOutcomeAgentError
|
||||||
|
errCode = model.MCPOutcomeAgentError
|
||||||
|
errMsg = "failed to encode tool result: " + err.Error()
|
||||||
|
result = nil
|
||||||
|
} else {
|
||||||
|
writeJSONRPCResult(c, req.ID, mcpToolCallResult{
|
||||||
|
Content: []mcpContent{{Type: "text", Text: textPayload}},
|
||||||
|
StructuredContent: result,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
audit.Outcome = outcome
|
||||||
|
audit.ErrorCode = errCode
|
||||||
|
audit.ErrorMsg = truncateString(errMsg, 512)
|
||||||
|
audit.DurationMs = time.Since(startedAt).Milliseconds()
|
||||||
|
audit.ServerID = extractServerID(p.Arguments)
|
||||||
|
mcpAuditWrite(audit, p.Arguments)
|
||||||
|
|
||||||
|
if outcome == model.MCPOutcomeOK {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// Semantic tool failures may retain a typed structured result (notably
|
||||||
|
// server.exec) so clients can distinguish a non-zero command outcome from
|
||||||
|
// transport, authorization, or deadline failures.
|
||||||
|
writeJSONRPCResult(c, req.ID, mcpToolCallResult{
|
||||||
|
Content: []mcpContent{{Type: "text", Text: errMsg}},
|
||||||
|
StructuredContent: result,
|
||||||
|
IsError: true,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
if tool.RequiredScope != "" && !tok.HasScope(tool.RequiredScope) {
|
||||||
|
finish(model.MCPOutcomeScopeDenied, model.MCPOutcomeScopeDenied,
|
||||||
|
"missing required scope: "+tool.RequiredScope, nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// 让 PAT 吊销能立即中断进行中的 tools/call(如 server.exec 最长 ~305s):
|
||||||
|
// 派生一个可取消 ctx 注入 c.Request,下游 CallAgent 用 c.Request.Context()
|
||||||
|
// 即会观察到取消;cancel 注册进吊销表,deleteAPIToken 会立刻触发它。
|
||||||
|
if c.Request != nil {
|
||||||
|
callCtx, cancel := context.WithCancel(c.Request.Context())
|
||||||
|
defer cancel()
|
||||||
|
deregister := registerPATConnection(c, cancel)
|
||||||
|
defer deregister()
|
||||||
|
c.Request = c.Request.WithContext(callCtx)
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := tool.Handler(c, p.Arguments)
|
||||||
|
if err != nil {
|
||||||
|
code, msg := classifyToolError(err)
|
||||||
|
var structuredErr interface{ StructuredResult() any }
|
||||||
|
if errors.As(err, &structuredErr) {
|
||||||
|
finish(code, code, msg, structuredErr.StructuredResult())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
finish(code, code, msg, nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
finish(model.MCPOutcomeOK, "", "", result)
|
||||||
|
}
|
||||||
|
|
||||||
|
// classifyToolError 把任何 handler 返回的 error 归类成审计 outcome + 安全错误消息。
|
||||||
|
// 优先匹配 mcpError 自带的 Code;否则匹配已知的 rpc.ErrAgent* 类型,最后回退 internal。
|
||||||
|
func classifyToolError(err error) (code, msg string) {
|
||||||
|
if me, ok := err.(*mcpError); ok {
|
||||||
|
return me.Code, me.Msg
|
||||||
|
}
|
||||||
|
if errors.Is(err, rpc.ErrAgentOffline) {
|
||||||
|
return model.MCPOutcomeServerOffline, "agent offline"
|
||||||
|
}
|
||||||
|
if errors.Is(err, rpc.ErrAgentTimeout) {
|
||||||
|
return model.MCPOutcomeAgentTimeout, "agent did not respond within timeout"
|
||||||
|
}
|
||||||
|
if errors.Is(err, rpc.ErrMCPDisabled) {
|
||||||
|
// kill switch 触发的中断必须独立成 outcome,避免审计/SIEM 把
|
||||||
|
// “管理员关了 MCP”误报成 agent 故障;错误文本透传原始原因。
|
||||||
|
return model.MCPOutcomeMCPDisabled, err.Error()
|
||||||
|
}
|
||||||
|
return model.MCPOutcomeAgentError, err.Error()
|
||||||
|
}
|
||||||
|
|
||||||
|
// extractServerID 从 raw arguments JSON 里提取 server_id(best-effort,只用于审计字段)。
|
||||||
|
func extractServerID(raw json.RawMessage) uint64 {
|
||||||
|
if len(raw) == 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
var probe struct {
|
||||||
|
ServerID uint64 `json:"server_id"`
|
||||||
|
}
|
||||||
|
_ = json.Unmarshal(raw, &probe)
|
||||||
|
return probe.ServerID
|
||||||
|
}
|
||||||
|
|
||||||
|
func truncateString(s string, max int) string {
|
||||||
|
if len(s) <= max {
|
||||||
|
return s
|
||||||
|
}
|
||||||
|
return s[:max]
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- wire writers ---
|
||||||
|
|
||||||
|
func writeJSONRPCResult(c *gin.Context, id json.RawMessage, result any) {
|
||||||
|
c.JSON(http.StatusOK, jsonRPCResponse{
|
||||||
|
JSONRPC: "2.0",
|
||||||
|
ID: id,
|
||||||
|
Result: result,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeJSONRPCError(c *gin.Context, id json.RawMessage, code int, message string) {
|
||||||
|
writeJSONRPCErrorWithStatus(c, id, code, message, http.StatusOK)
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeToolCallError(c *gin.Context, id json.RawMessage, errCode, errMsg string) {
|
||||||
|
writeJSONRPCResult(c, id, mcpToolCallResult{
|
||||||
|
Content: []mcpContent{{Type: "text", Text: errMsg}},
|
||||||
|
IsError: true,
|
||||||
|
StructuredContent: map[string]string{
|
||||||
|
"error_code": errCode,
|
||||||
|
"error": errMsg,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeJSONRPCErrorWithStatus(c *gin.Context, id json.RawMessage, code int, message string, status int) {
|
||||||
|
c.JSON(status, jsonRPCResponse{
|
||||||
|
JSONRPC: "2.0",
|
||||||
|
ID: id,
|
||||||
|
Error: &jsonRPCError{Code: code, Message: message},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- 错误语义 ---
|
||||||
|
|
||||||
|
// mcpError 是 tool handler 可以返回的语义化错误。
|
||||||
|
// dispatch 根据 Code 决定 audit outcome 与 JSON-RPC 错误码(如果命中 rpcErr* 域)。
|
||||||
|
type mcpError struct {
|
||||||
|
Code string
|
||||||
|
Msg string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *mcpError) Error() string { return e.Msg }
|
||||||
|
|
||||||
|
func newMCPError(code, msg string) *mcpError { return &mcpError{Code: code, Msg: msg} }
|
||||||
|
|
||||||
|
// 预制错误
|
||||||
|
var (
|
||||||
|
errMCPInvalidArgs = func(s string) *mcpError { return newMCPError(model.MCPOutcomeInvalidArgs, s) }
|
||||||
|
errMCPPermDenied = newMCPError(model.MCPOutcomePermDenied, "permission denied")
|
||||||
|
errMCPScopeDenied = func(s string) *mcpError {
|
||||||
|
return newMCPError(model.MCPOutcomeScopeDenied, "missing required scope: "+s)
|
||||||
|
}
|
||||||
|
errMCPServerOffline = newMCPError(model.MCPOutcomeServerOffline, "agent offline")
|
||||||
|
errMCPAgentTimeout = newMCPError(model.MCPOutcomeAgentTimeout, "agent did not respond within timeout")
|
||||||
|
errMCPUnsupported = newMCPError(model.MCPOutcomeUnsupportedAgent, "agent does not support this MCP capability; please upgrade the agent")
|
||||||
|
)
|
||||||
|
|
||||||
|
// --- 共用工具 ---
|
||||||
|
|
||||||
|
var errNoToken = errors.New("no api token in context")
|
||||||
|
|
||||||
|
// decodeToolArgs 是 tool handler 用来反序列化 arguments 的辅助。
|
||||||
|
func decodeToolArgs(raw json.RawMessage, out any) error {
|
||||||
|
if len(raw) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(raw, out); err != nil {
|
||||||
|
return fmt.Errorf("invalid arguments: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// requireServerAccess 是 tool handler 共用的「闸 1 + 闸 2 服务器白名单」组合校验。
|
||||||
|
// 通过返回 *model.Server;失败返回带语义 Code 的 mcpError,便于 dispatch 归类审计。
|
||||||
|
func requireServerAccess(c *gin.Context, serverID uint64) (*model.Server, error) {
|
||||||
|
if serverID == 0 {
|
||||||
|
return nil, errMCPInvalidArgs("server_id required")
|
||||||
|
}
|
||||||
|
tok := APITokenFromContext(c)
|
||||||
|
if tok != nil && !tok.CanAccessServer(serverID) {
|
||||||
|
return nil, errMCPPermDenied
|
||||||
|
}
|
||||||
|
server, _ := singleton.ServerShared.Get(serverID)
|
||||||
|
if server == nil {
|
||||||
|
return nil, errMCPServerOffline
|
||||||
|
}
|
||||||
|
if !server.HasPermission(c) {
|
||||||
|
return nil, errMCPPermDenied
|
||||||
|
}
|
||||||
|
return server, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,48 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
|
"log"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// mcpAuditWrite 异步写一条 MCP 审计日志。失败仅 log,不阻塞业务。
|
||||||
|
//
|
||||||
|
// argsBytes:tool 的 raw JSON 参数(dispatcher 已经反序列化过)。
|
||||||
|
// 只记录 sha256 全文哈希,不保留任何明文片段:server.exec 的 env/stdin、
|
||||||
|
// fs.write 的 content 等字段会包含 token、密码、密钥、文件内容等敏感数据,
|
||||||
|
// 任何长度的 peek 都可能让审计表本身成为 secret 仓库;以哈希做关联即可。
|
||||||
|
//
|
||||||
|
// 测试可以把 mcpAuditSync 置为 true 让写入同步,避免 goroutine 与测试 teardown
|
||||||
|
// 形成竞态(不同测试 swap 全局 singleton.DB 时尤其明显)。
|
||||||
|
func mcpAuditWrite(entry model.MCPAuditLog, argsBytes []byte) {
|
||||||
|
if len(argsBytes) > 0 {
|
||||||
|
sum := sha256.Sum256(argsBytes)
|
||||||
|
entry.ArgsHash = hex.EncodeToString(sum[:])
|
||||||
|
}
|
||||||
|
entry.ArgsPeek = ""
|
||||||
|
if entry.CreatedAt.IsZero() {
|
||||||
|
entry.CreatedAt = time.Now()
|
||||||
|
}
|
||||||
|
db := singleton.DB
|
||||||
|
write := func(e model.MCPAuditLog) {
|
||||||
|
if db == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := db.Create(&e).Error; err != nil {
|
||||||
|
log.Printf("NEZHA>> mcp audit write failed: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if mcpAuditSync {
|
||||||
|
write(entry)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
go write(entry)
|
||||||
|
}
|
||||||
|
|
||||||
|
// mcpAuditSync 仅供测试切换为同步写入,生产保持 false。
|
||||||
|
var mcpAuditSync = false
|
||||||
@@ -0,0 +1,72 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// H7 regression: the MCP endpoint must cap incoming JSON-RPC body size
|
||||||
|
// BEFORE decoding. Without this, a valid PAT can post a multi-GB body and
|
||||||
|
// the dashboard exhausts memory in ShouldBindJSON. We assert the body
|
||||||
|
// reader is wrapped in http.MaxBytesReader; the exact error path the
|
||||||
|
// decoder takes is irrelevant as long as the cap is enforced.
|
||||||
|
func TestMCPEndpoint_BodyIsCappedByMaxBytesReader(t *testing.T) {
|
||||||
|
prevConf := singleton.Conf
|
||||||
|
cfg := &model.Config{}
|
||||||
|
cfg.SetMCPEnabled(true)
|
||||||
|
singleton.Conf = &singleton.ConfigClass{Config: cfg}
|
||||||
|
t.Cleanup(func() { singleton.Conf = prevConf })
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 1, ScopesCSV: "nezha:server:read"}
|
||||||
|
// 16 MiB of valid-JSON whitespace prefix forces the decoder to actually
|
||||||
|
// stream past the limit, exercising MaxBytesReader.
|
||||||
|
body := `{"jsonrpc":"2.0","id":1,"method":"initialize","params":"` +
|
||||||
|
strings.Repeat("x", mcpJSONRPCMaxBodyBytes+1024) + `"}`
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/mcp", bytes.NewBufferString(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = req
|
||||||
|
c.Set(apiTokenCtxKey, tok)
|
||||||
|
c.Set(model.CtxKeyAPIToken, tok)
|
||||||
|
|
||||||
|
mcpEndpoint(c)
|
||||||
|
|
||||||
|
if !strings.Contains(w.Body.String(), "request body") &&
|
||||||
|
w.Code != http.StatusRequestEntityTooLarge {
|
||||||
|
t.Fatalf("oversized body must be rejected with a body-size error, got code=%d body=%s",
|
||||||
|
w.Code, w.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPEndpoint_AcceptsSmallBody(t *testing.T) {
|
||||||
|
prevConf := singleton.Conf
|
||||||
|
cfg := &model.Config{}
|
||||||
|
cfg.SetMCPEnabled(true)
|
||||||
|
singleton.Conf = &singleton.ConfigClass{Config: cfg}
|
||||||
|
t.Cleanup(func() { singleton.Conf = prevConf })
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 1, ScopesCSV: "nezha:server:read"}
|
||||||
|
body := `{"jsonrpc":"2.0","id":1,"method":"initialize"}`
|
||||||
|
req := httptest.NewRequest(http.MethodPost, "/mcp", bytes.NewBufferString(body))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = req
|
||||||
|
c.Set(apiTokenCtxKey, tok)
|
||||||
|
c.Set(model.CtxKeyAPIToken, tok)
|
||||||
|
|
||||||
|
mcpEndpoint(c)
|
||||||
|
|
||||||
|
if w.Code != http.StatusOK {
|
||||||
|
t.Fatalf("small valid body must succeed, got code=%d body=%s", w.Code, w.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,71 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
// MCPMinAgentVersion 是支持 MCP 的最低 agent 版本。
|
||||||
|
//
|
||||||
|
// release 流程:在 agent ship 了 MCP handlers 后,把此值更新为该 release 的版本号。
|
||||||
|
// 不变量:必须为非空。否则旧 agent 收到 TaskTypeExec/TaskTypeFs* 等新任务类型时
|
||||||
|
// 走 default 分支不回 TaskResult,dashboard 要等 CallAgent 超时(30s)甚至更久
|
||||||
|
// (fs.transfer 的 IOStream attach 30s)才能感知,这是 server-transfer 已经
|
||||||
|
// 通过 MinServerTransferAgentVersion 修复过的同类问题。
|
||||||
|
const MCPMinAgentVersion = "v2.1.0"
|
||||||
|
|
||||||
|
// requireAgentSupportsMCP 在 tool handler 调 CallAgent 之前快速失败不支持的 agent。
|
||||||
|
// 仅作为 UX 优化:真正的安全/正确性由 agent 端 task switch 的 default 分支保障。
|
||||||
|
func requireAgentSupportsMCP(server *model.Server) error {
|
||||||
|
if MCPMinAgentVersion == "" || server == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
runtime := server.RuntimeSnapshot()
|
||||||
|
if runtime.Host == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if compareSemver(runtime.Host.Version, MCPMinAgentVersion) < 0 {
|
||||||
|
return errMCPUnsupported
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// compareSemver 比较两个 "MAJOR.MINOR.PATCH[-suffix]" 字符串,返回 -1/0/1。
|
||||||
|
// semverParts 已剥掉可选的 "v" 前缀与 "-/+" 后缀,所以只按三段数字定序:
|
||||||
|
// 数字段相等即视为相等版本。绝不能回退到字符串字典序——agent 上报 "2.1.0"
|
||||||
|
// 而门槛常量是 "v2.1.0",'2'(0x32) < 'v'(0x76) 会把相等版本误判为更旧,
|
||||||
|
// 导致所有 agent 被错误地判为不支持 MCP。
|
||||||
|
func compareSemver(a, b string) int {
|
||||||
|
aparts := semverParts(a)
|
||||||
|
bparts := semverParts(b)
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
if aparts[i] < bparts[i] {
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
if aparts[i] > bparts[i] {
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func semverParts(v string) [3]int {
|
||||||
|
v = strings.TrimPrefix(v, "v")
|
||||||
|
if i := strings.IndexAny(v, "-+"); i >= 0 {
|
||||||
|
v = v[:i]
|
||||||
|
}
|
||||||
|
var out [3]int
|
||||||
|
parts := strings.Split(v, ".")
|
||||||
|
for i := 0; i < 3 && i < len(parts); i++ {
|
||||||
|
n := 0
|
||||||
|
for _, c := range parts[i] {
|
||||||
|
if c < '0' || c > '9' {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
n = n*10 + int(c-'0')
|
||||||
|
}
|
||||||
|
out[i] = n
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
@@ -0,0 +1,77 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 旧 agent 不识别 TaskTypeExec/TaskTypeFs* 等新 task type,会走 default
|
||||||
|
// 分支不回 TaskResult;dashboard 必须在调 CallAgent 之前依据 Host.Version
|
||||||
|
// 快速失败,否则用户要等到 30s/24h timeout 才知道 agent 不支持。
|
||||||
|
// MinServerTransferAgentVersion 已为同类问题在 transfer 路径上确立了
|
||||||
|
// release-time 必填的版本下限——这里把 MCP 也纳入同一不变量。
|
||||||
|
func TestMCPMinAgentVersionIsPinnedToRelease(t *testing.T) {
|
||||||
|
require.NotEmpty(t, MCPMinAgentVersion,
|
||||||
|
"MCPMinAgentVersion must be set to the lowest agent build that ships MCP handlers; an empty string disables the gate and lets old agents hang dashboard requests until timeout")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRequireAgentSupportsMCPRejectsBelowMinVersion(t *testing.T) {
|
||||||
|
old := &model.Server{Host: &model.Host{Version: "v0.0.1"}}
|
||||||
|
err := requireAgentSupportsMCP(old)
|
||||||
|
require.Error(t, err, "agents older than MCPMinAgentVersion must be rejected before CallAgent")
|
||||||
|
require.True(t, errors.Is(err, errMCPUnsupported) || err.Error() == errMCPUnsupported.Error(),
|
||||||
|
"expected the errMCPUnsupported sentinel, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRequireAgentSupportsMCPAcceptsCurrentVersion(t *testing.T) {
|
||||||
|
current := &model.Server{Host: &model.Host{Version: MCPMinAgentVersion}}
|
||||||
|
require.NoError(t, requireAgentSupportsMCP(current),
|
||||||
|
"server reporting exactly MCPMinAgentVersion must be accepted")
|
||||||
|
}
|
||||||
|
|
||||||
|
// 钉住「最近一个不带 MCP handler 的已发布 agent tag (v2.0.4) 必须被拒绝」。
|
||||||
|
// v2.0.4 的 model/task.go 还没有 TaskTypeExec/TaskTypeFs* 常量,cmd/agent/
|
||||||
|
// mcp_handlers.go 也不存在;如果版本门槛把它放行,dashboard 调 MCP tool
|
||||||
|
// 后 agent 会走 default 分支不回 TaskResult,CallAgent 必须等 30s 超时。
|
||||||
|
func TestRequireAgentSupportsMCPRejectsLastReleaseWithoutMCP(t *testing.T) {
|
||||||
|
noMCP := &model.Server{Host: &model.Host{Version: "v2.0.4"}}
|
||||||
|
err := requireAgentSupportsMCP(noMCP)
|
||||||
|
require.Error(t, err,
|
||||||
|
"v2.0.4 is the latest released agent tag that ships *without* MCP handlers; bumping MCPMinAgentVersion below the first MCP release re-introduces the silent-timeout bug")
|
||||||
|
require.True(t, errors.Is(err, errMCPUnsupported) || err.Error() == errMCPUnsupported.Error(),
|
||||||
|
"expected the errMCPUnsupported sentinel, got %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRequireAgentSupportsMCPDefersWhenAgentNeverReported(t *testing.T) {
|
||||||
|
require.NoError(t, requireAgentSupportsMCP(&model.Server{Host: nil}),
|
||||||
|
"Host==nil means agent never reported its build; defer the version decision to the CallAgent timeout layer")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompareSemverIgnoresVPrefixMismatch(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
a, b string
|
||||||
|
want int
|
||||||
|
}{
|
||||||
|
{"2.1.0", "v2.1.0", 0},
|
||||||
|
{"v2.1.0", "2.1.0", 0},
|
||||||
|
{"2.1.0", "2.1.0", 0},
|
||||||
|
{"2.1.0", "v2.1.1", -1},
|
||||||
|
{"v2.1.2", "2.1.0", 1},
|
||||||
|
{"2.2.0", "v2.1.9", 1},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
require.Equalf(t, c.want, compareSemver(c.a, c.b),
|
||||||
|
"compareSemver(%q,%q): the only difference is the optional 'v' prefix and/or numeric ordering; "+
|
||||||
|
"a bare lexical fallback wrongly orders %q < %q", c.a, c.b, c.a, c.b)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRequireAgentSupportsMCPAcceptsBarePrefixReport(t *testing.T) {
|
||||||
|
agent := &model.Server{Host: &model.Host{Version: "2.1.0"}}
|
||||||
|
require.NoError(t, requireAgentSupportsMCP(agent),
|
||||||
|
"agents report Host.Version without the 'v' prefix (e.g. \"2.1.0\"); it must compare equal to MCPMinAgentVersion \"v2.1.0\" and be accepted")
|
||||||
|
}
|
||||||
@@ -0,0 +1,28 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
// classifyToolError 必须把 rpc.ErrMCPDisabled 归类成 forbidden 类 outcome,而不是
|
||||||
|
// 当作 agent_error。
|
||||||
|
//
|
||||||
|
// ErrMCPDisabled 是 dashboard 主动按下 kill switch 的语义信号(见
|
||||||
|
// service/rpc/mcp_rpc.go 注释),controller 把它揉进 agent_error 等于把“管理员
|
||||||
|
// 关了 MCP”和“agent 真出故障”混在一起:审计日志、SIEM 告警和 MCP 客户端的
|
||||||
|
// structuredContent.error_code 都会错配。
|
||||||
|
func TestClassifyToolError_MCPDisabledMapsToForbidden(t *testing.T) {
|
||||||
|
code, msg := classifyToolError(rpc.ErrMCPDisabled)
|
||||||
|
|
||||||
|
assert.Equal(t, model.MCPOutcomeMCPDisabled, code,
|
||||||
|
"rpc.ErrMCPDisabled must map to MCPOutcomeMCPDisabled, not agent_error")
|
||||||
|
assert.NotEqual(t, model.MCPOutcomeAgentError, code,
|
||||||
|
"kill-switch errors must not be reported as agent_error in audit/SIEM")
|
||||||
|
assert.Contains(t, msg, "MCP is disabled",
|
||||||
|
"error text should preserve the kill switch reason")
|
||||||
|
}
|
||||||
@@ -0,0 +1,96 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// installTestConfig swaps singleton.Conf with one backed by a tmp file so
|
||||||
|
// updateConfig's Conf.Save() write-through has a real target. The caller's
|
||||||
|
// setupMCPTest will restore the original Conf when its cleanup runs.
|
||||||
|
func installTestConfig(t *testing.T) {
|
||||||
|
t.Helper()
|
||||||
|
dir := t.TempDir()
|
||||||
|
cfg := &model.Config{}
|
||||||
|
require.NoError(t, cfg.Read(filepath.Join(dir, "config.yaml"), nil))
|
||||||
|
singleton.Conf = &singleton.ConfigClass{Config: cfg}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateConfig_PersistsEnableMCPFlag(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
installTestConfig(t)
|
||||||
|
|
||||||
|
origTemplates := singleton.FrontendTemplates
|
||||||
|
singleton.FrontendTemplates = []model.FrontendTemplate{
|
||||||
|
{Path: "user-dist", IsAdmin: false},
|
||||||
|
}
|
||||||
|
defer func() { singleton.FrontendTemplates = origTemplates }()
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.Use(func(c *gin.Context) {
|
||||||
|
setAuthUser(c, uid, model.RoleAdmin)
|
||||||
|
c.Next()
|
||||||
|
})
|
||||||
|
r.PATCH("/api/v1/setting", commonHandler(updateConfig))
|
||||||
|
|
||||||
|
body := map[string]any{
|
||||||
|
"site_name": "test",
|
||||||
|
"language": "en_US",
|
||||||
|
"user_template": "user-dist",
|
||||||
|
"enable_mcp": true,
|
||||||
|
}
|
||||||
|
raw, _ := json.Marshal(body)
|
||||||
|
req := httptest.NewRequest(http.MethodPatch, "/api/v1/setting", bytes.NewReader(raw))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
require.True(t, success, "PATCH /setting must succeed: %s", errMsg)
|
||||||
|
require.True(t, singleton.Conf.EnableMCP,
|
||||||
|
"enable_mcp=true in body must flip singleton.Conf.EnableMCP")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPEndpoint_RefusesWhenDisabled(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
singleton.Conf.SetMCPEnabled(false)
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "initialize",
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
var env jsonRPCResponse
|
||||||
|
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &env))
|
||||||
|
require.NotNil(t, env.Error, "MCP must return JSON-RPC error when disabled; body=%s", w.Body.String())
|
||||||
|
require.Equal(t, rpcErrForbidden, env.Error.Code,
|
||||||
|
"disabled MCP must surface as rpcErrForbidden so callers can distinguish from auth failure")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPEndpoint_AllowsWhenEnabled(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
singleton.Conf.SetMCPEnabled(true)
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "initialize",
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
var env jsonRPCResponse
|
||||||
|
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &env))
|
||||||
|
require.Nil(t, env.Error, "MCP must process requests when enabled; got error=%+v", env.Error)
|
||||||
|
}
|
||||||
@@ -0,0 +1,433 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
|
"encoding/binary"
|
||||||
|
"encoding/json"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"google.golang.org/grpc/metadata"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
pb "github.com/nezhahq/nezha/proto"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
type e2eStream struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
dispatch func(*pb.Task) *pb.TaskResult
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *e2eStream) Send(t *pb.Task) error {
|
||||||
|
// fs.upload_url / fs.download_url 走 IOStream 路径,由独立 mux 处理;
|
||||||
|
// 这里的 RPC-style dispatch 只覆盖 fs.read/fs.write/fs.list/fs.delete/server.exec。
|
||||||
|
if t.GetType() == model.TaskTypeFsTransfer {
|
||||||
|
return e2eHandleFsTransfer(t)
|
||||||
|
}
|
||||||
|
s.mu.Lock()
|
||||||
|
d := s.dispatch
|
||||||
|
s.mu.Unlock()
|
||||||
|
if d == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
go func(task *pb.Task) {
|
||||||
|
if res := d(task); res != nil {
|
||||||
|
rpc.DeliverMCPResultForTest(res)
|
||||||
|
}
|
||||||
|
}(t)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// e2eHandleFsTransfer 模拟真实 agent 收到 TaskTypeFsTransfer:把本地文件
|
||||||
|
// 系统作为后端,按 op 跑完整协议帧并复制字节。和真 agent 不同:
|
||||||
|
// - 不做 sha256 强校验(测试侧用 NZTO 中的 32 字节固定 0 占位)。
|
||||||
|
// - 复用 net.Pipe + rpc.NezhaHandlerSingleton.AgentConnected 注入 dashboard 端。
|
||||||
|
func e2eHandleFsTransfer(t *pb.Task) error {
|
||||||
|
var req model.FsTransferRequest
|
||||||
|
if err := json.Unmarshal([]byte(t.GetData()), &req); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
dashboardSide, agentSide := net.Pipe()
|
||||||
|
if err := rpc.NezhaHandlerSingleton.AgentConnected(req.StreamID, dashboardSide); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
go func() {
|
||||||
|
defer agentSide.Close()
|
||||||
|
switch req.Op {
|
||||||
|
case model.MCPFsTransferOpDownload:
|
||||||
|
data, err := os.ReadFile(req.Path)
|
||||||
|
if err != nil {
|
||||||
|
buf := append([]byte(nil), model.MCPFsXferMagicErr...)
|
||||||
|
buf = append(buf, err.Error()...)
|
||||||
|
_, _ = agentSide.Write(buf)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
hdr := append([]byte(nil), model.MCPFsXferMagicDownloadHdr...)
|
||||||
|
sz := make([]byte, 8)
|
||||||
|
binary.BigEndian.PutUint64(sz, uint64(len(data)))
|
||||||
|
hdr = append(hdr, sz...)
|
||||||
|
hdr = append(hdr, make([]byte, 32)...)
|
||||||
|
if _, err := agentSide.Write(hdr); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if len(data) > 0 {
|
||||||
|
chunk := append([]byte(nil), model.MCPFsXferMagicChunk...)
|
||||||
|
chunkLen := make([]byte, 8)
|
||||||
|
binary.BigEndian.PutUint64(chunkLen, uint64(len(data)))
|
||||||
|
chunk = append(chunk, chunkLen...)
|
||||||
|
chunk = append(chunk, data...)
|
||||||
|
if _, err := agentSide.Write(chunk); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
ok := append([]byte(nil), model.MCPFsXferMagicOK...)
|
||||||
|
ok = append(ok, sz...)
|
||||||
|
ok = append(ok, make([]byte, 32)...)
|
||||||
|
_, _ = agentSide.Write(ok)
|
||||||
|
case model.MCPFsTransferOpUpload:
|
||||||
|
hdr := append([]byte(nil), model.MCPFsXferMagicUploadHdr...)
|
||||||
|
sz := make([]byte, 8)
|
||||||
|
binary.BigEndian.PutUint64(sz, uint64(req.Size))
|
||||||
|
hdr = append(hdr, sz...)
|
||||||
|
if _, err := agentSide.Write(hdr); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
buf := make([]byte, req.Size)
|
||||||
|
if req.Size > 0 {
|
||||||
|
if _, err := io.ReadFull(agentSide, buf); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(req.Path, buf, 0o644); err != nil {
|
||||||
|
errBuf := append([]byte(nil), model.MCPFsXferMagicErr...)
|
||||||
|
errBuf = append(errBuf, err.Error()...)
|
||||||
|
_, _ = agentSide.Write(errBuf)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ok := append([]byte(nil), model.MCPFsXferMagicOK...)
|
||||||
|
okSz := make([]byte, 8)
|
||||||
|
binary.BigEndian.PutUint64(okSz, uint64(len(buf)))
|
||||||
|
ok = append(ok, okSz...)
|
||||||
|
ok = append(ok, make([]byte, 32)...)
|
||||||
|
_, _ = agentSide.Write(ok)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *e2eStream) Recv() (*pb.TaskResult, error) { return nil, context.Canceled }
|
||||||
|
func (s *e2eStream) SetHeader(metadata.MD) error { return nil }
|
||||||
|
func (s *e2eStream) SendHeader(metadata.MD) error { return nil }
|
||||||
|
func (s *e2eStream) SetTrailer(metadata.MD) {}
|
||||||
|
func (s *e2eStream) Context() context.Context { return context.Background() }
|
||||||
|
func (s *e2eStream) SendMsg(any) error { return nil }
|
||||||
|
func (s *e2eStream) RecvMsg(any) error { return context.Canceled }
|
||||||
|
|
||||||
|
func agentSim(task *pb.Task) *pb.TaskResult {
|
||||||
|
res := &pb.TaskResult{Id: task.GetId(), Type: task.GetType(), Successful: true}
|
||||||
|
switch task.GetType() {
|
||||||
|
case model.TaskTypeFsList:
|
||||||
|
var req model.FsListRequest
|
||||||
|
_ = json.Unmarshal([]byte(task.GetData()), &req)
|
||||||
|
entries, err := os.ReadDir(req.Path)
|
||||||
|
if err != nil {
|
||||||
|
b, _ := json.Marshal(model.FsListResult{Error: err.Error()})
|
||||||
|
res.Data = string(b)
|
||||||
|
return res
|
||||||
|
}
|
||||||
|
out := make([]model.FsEntry, 0, len(entries))
|
||||||
|
for _, e := range entries {
|
||||||
|
info, _ := e.Info()
|
||||||
|
out = append(out, model.FsEntry{Name: e.Name(), Type: "file", Size: info.Size()})
|
||||||
|
}
|
||||||
|
b, _ := json.Marshal(model.FsListResult{Entries: out, Total: len(out)})
|
||||||
|
res.Data = string(b)
|
||||||
|
case model.TaskTypeFsRead:
|
||||||
|
var req model.FsReadRequest
|
||||||
|
_ = json.Unmarshal([]byte(task.GetData()), &req)
|
||||||
|
data, err := os.ReadFile(req.Path)
|
||||||
|
if err != nil {
|
||||||
|
b, _ := json.Marshal(model.FsReadResult{Error: err.Error()})
|
||||||
|
res.Data = string(b)
|
||||||
|
return res
|
||||||
|
}
|
||||||
|
encoding := req.Encoding
|
||||||
|
if encoding == "" {
|
||||||
|
encoding = "utf8"
|
||||||
|
}
|
||||||
|
var content string
|
||||||
|
switch encoding {
|
||||||
|
case "base64":
|
||||||
|
content = base64.StdEncoding.EncodeToString(data)
|
||||||
|
default:
|
||||||
|
content = string(data)
|
||||||
|
}
|
||||||
|
b, _ := json.Marshal(model.FsReadResult{Content: content, Encoding: encoding, Size: int64(len(data))})
|
||||||
|
res.Data = string(b)
|
||||||
|
case model.TaskTypeFsWrite:
|
||||||
|
var req model.FsWriteRequest
|
||||||
|
_ = json.Unmarshal([]byte(task.GetData()), &req)
|
||||||
|
data := []byte(req.Content)
|
||||||
|
if req.Encoding == "base64" {
|
||||||
|
decoded, decErr := base64.StdEncoding.DecodeString(req.Content)
|
||||||
|
if decErr != nil {
|
||||||
|
b, _ := json.Marshal(model.FsWriteResult{Error: decErr.Error()})
|
||||||
|
res.Data = string(b)
|
||||||
|
return res
|
||||||
|
}
|
||||||
|
data = decoded
|
||||||
|
}
|
||||||
|
_ = os.WriteFile(req.Path, data, 0o644)
|
||||||
|
b, _ := json.Marshal(model.FsWriteResult{Size: int64(len(data))})
|
||||||
|
res.Data = string(b)
|
||||||
|
case model.TaskTypeFsDelete:
|
||||||
|
var req model.FsDeleteRequest
|
||||||
|
_ = json.Unmarshal([]byte(task.GetData()), &req)
|
||||||
|
_ = os.RemoveAll(req.Path)
|
||||||
|
b, _ := json.Marshal(model.FsDeleteResult{DeletedCount: 1})
|
||||||
|
res.Data = string(b)
|
||||||
|
case model.TaskTypeExec:
|
||||||
|
b, _ := json.Marshal(model.ExecResult{ExitCode: 0, Stdout: "simulated"})
|
||||||
|
res.Data = string(b)
|
||||||
|
default:
|
||||||
|
res.Successful = false
|
||||||
|
res.Data = "unsupported task"
|
||||||
|
}
|
||||||
|
return res
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupEndToEnd(t *testing.T) (*httptest.Server, string, func()) {
|
||||||
|
t.Helper()
|
||||||
|
cleanupBase, uid := setupMCPTest(t)
|
||||||
|
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
|
||||||
|
stream := &e2eStream{dispatch: agentSim}
|
||||||
|
srv, _ := singleton.ServerShared.Get(7)
|
||||||
|
srv.SetTaskStream(stream)
|
||||||
|
|
||||||
|
prevCleanup := cleanupBase
|
||||||
|
cleanupBase = func() {
|
||||||
|
rpc.NezhaHandlerSingleton = originalHandler
|
||||||
|
prevCleanup()
|
||||||
|
}
|
||||||
|
|
||||||
|
_, plain := mkToken(t, uid, []string{
|
||||||
|
model.ScopeInventoryRead,
|
||||||
|
model.ScopeInventoryDelete,
|
||||||
|
model.ScopeServerRead,
|
||||||
|
model.ScopeServerExec,
|
||||||
|
model.ScopeServerWrite,
|
||||||
|
model.ScopeServerDelete,
|
||||||
|
}, nil)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.POST("/mcp", apiTokenAuthMiddleware(), mcpEndpoint)
|
||||||
|
r.GET("/mcp/download/:token", transferDownloadHandler)
|
||||||
|
r.POST("/mcp/upload/:token", transferUploadHandler)
|
||||||
|
ts := httptest.NewServer(r)
|
||||||
|
|
||||||
|
return ts, plain, func() {
|
||||||
|
ts.Close()
|
||||||
|
cleanupBase()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func e2eCall(t *testing.T, ts *httptest.Server, token, method, toolName string, args any) map[string]any {
|
||||||
|
t.Helper()
|
||||||
|
body := map[string]any{"jsonrpc": "2.0", "id": 1, "method": method}
|
||||||
|
if method == "tools/call" {
|
||||||
|
argsRaw, _ := json.Marshal(args)
|
||||||
|
body["params"] = map[string]any{"name": toolName, "arguments": json.RawMessage(argsRaw)}
|
||||||
|
}
|
||||||
|
b, _ := json.Marshal(body)
|
||||||
|
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader(b))
|
||||||
|
req.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
out, _ := io.ReadAll(resp.Body)
|
||||||
|
var env map[string]any
|
||||||
|
require.NoError(t, json.Unmarshal(out, &env))
|
||||||
|
return env
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestE2E_Initialize(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupEndToEnd(t)
|
||||||
|
defer cleanup()
|
||||||
|
env := e2eCall(t, ts, tok, "initialize", "", nil)
|
||||||
|
require.Nil(t, env["error"])
|
||||||
|
info := env["result"].(map[string]any)["serverInfo"].(map[string]any)
|
||||||
|
require.Equal(t, "nezha-mcp", info["name"])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestE2E_ToolsList(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupEndToEnd(t)
|
||||||
|
defer cleanup()
|
||||||
|
env := e2eCall(t, ts, tok, "tools/list", "", nil)
|
||||||
|
require.Nil(t, env["error"])
|
||||||
|
tools := env["result"].(map[string]any)["tools"].([]any)
|
||||||
|
require.GreaterOrEqual(t, len(tools), 9)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestE2E_WhoamiAndServerList(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupEndToEnd(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
env := e2eCall(t, ts, tok, "tools/call", "meta.whoami", map[string]any{})
|
||||||
|
res := env["result"].(map[string]any)
|
||||||
|
require.False(t, res["isError"] == true)
|
||||||
|
|
||||||
|
env = e2eCall(t, ts, tok, "tools/call", "server.list", map[string]any{})
|
||||||
|
res = env["result"].(map[string]any)
|
||||||
|
require.False(t, res["isError"] == true)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestE2E_ServerExec(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupEndToEnd(t)
|
||||||
|
defer cleanup()
|
||||||
|
env := e2eCall(t, ts, tok, "tools/call", "server.exec", map[string]any{
|
||||||
|
"server_id": 7, "cmd": "echo",
|
||||||
|
})
|
||||||
|
res := env["result"].(map[string]any)
|
||||||
|
require.False(t, res["isError"] == true, "exec failed: %v", res)
|
||||||
|
struc := res["structuredContent"].(map[string]any)
|
||||||
|
require.Equal(t, "simulated", struc["stdout"])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestE2E_FsLifecycle(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupEndToEnd(t)
|
||||||
|
defer cleanup()
|
||||||
|
dir := t.TempDir()
|
||||||
|
p := filepath.Join(dir, "e2e.txt")
|
||||||
|
|
||||||
|
env := e2eCall(t, ts, tok, "tools/call", "fs.write", map[string]any{
|
||||||
|
"server_id": 7, "path": p, "content": "ohi", "encoding": "utf8",
|
||||||
|
})
|
||||||
|
require.False(t, env["result"].(map[string]any)["isError"] == true)
|
||||||
|
|
||||||
|
env = e2eCall(t, ts, tok, "tools/call", "fs.read", map[string]any{"server_id": 7, "path": p})
|
||||||
|
res := env["result"].(map[string]any)
|
||||||
|
struc := res["structuredContent"].(map[string]any)
|
||||||
|
require.Equal(t, "ohi", struc["content"])
|
||||||
|
|
||||||
|
env = e2eCall(t, ts, tok, "tools/call", "fs.delete", map[string]any{"server_id": 7, "path": p})
|
||||||
|
require.False(t, env["result"].(map[string]any)["isError"] == true)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestE2E_DownloadUploadURL(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupEndToEnd(t)
|
||||||
|
defer cleanup()
|
||||||
|
dir := t.TempDir()
|
||||||
|
p := filepath.Join(dir, "blob.txt")
|
||||||
|
require.NoError(t, os.WriteFile(p, []byte("payload"), 0o644))
|
||||||
|
|
||||||
|
env := e2eCall(t, ts, tok, "tools/call", "fs.download_url", map[string]any{
|
||||||
|
"server_id": 7, "path": p, "ttl_seconds": 60,
|
||||||
|
})
|
||||||
|
res := env["result"].(map[string]any)
|
||||||
|
require.False(t, res["isError"] == true, "download_url failed: %v", res)
|
||||||
|
url := res["structuredContent"].(map[string]any)["url"].(string)
|
||||||
|
url = ts.URL + url[strings.Index(url, "/mcp/"):]
|
||||||
|
resp, err := http.Get(url)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
require.Equal(t, 200, resp.StatusCode)
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
require.Equal(t, "payload", string(body))
|
||||||
|
|
||||||
|
upPath := filepath.Join(dir, "up.txt")
|
||||||
|
env = e2eCall(t, ts, tok, "tools/call", "fs.upload_url", map[string]any{
|
||||||
|
"server_id": 7, "path": upPath, "ttl_seconds": 60,
|
||||||
|
})
|
||||||
|
res = env["result"].(map[string]any)
|
||||||
|
require.False(t, res["isError"] == true)
|
||||||
|
upURL := res["structuredContent"].(map[string]any)["url"].(string)
|
||||||
|
upURL = ts.URL + upURL[strings.Index(upURL, "/mcp/"):]
|
||||||
|
|
||||||
|
upReq, _ := http.NewRequest("POST", upURL, bytes.NewReader([]byte("hello-upload")))
|
||||||
|
upResp, err := http.DefaultClient.Do(upReq)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer upResp.Body.Close()
|
||||||
|
require.Equal(t, 200, upResp.StatusCode)
|
||||||
|
got, _ := os.ReadFile(upPath)
|
||||||
|
require.Equal(t, "hello-upload", string(got))
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestE2E_DownloadUploadURL_100MiB 走完整 mint→IOStream→relay 路径,验证
|
||||||
|
// 大文件能跨越旧 4MiB gRPC 上限,并且字节序保持不变。
|
||||||
|
func TestE2E_DownloadUploadURL_100MiB(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupEndToEnd(t)
|
||||||
|
defer cleanup()
|
||||||
|
dir := t.TempDir()
|
||||||
|
src := filepath.Join(dir, "src.bin")
|
||||||
|
|
||||||
|
want := make([]byte, model.MCPFsTransferMaxSize)
|
||||||
|
for i := range want {
|
||||||
|
want[i] = byte(i % 251)
|
||||||
|
}
|
||||||
|
require.NoError(t, os.WriteFile(src, want, 0o644))
|
||||||
|
|
||||||
|
env := e2eCall(t, ts, tok, "tools/call", "fs.download_url", map[string]any{
|
||||||
|
"server_id": 7, "path": src, "ttl_seconds": 60,
|
||||||
|
})
|
||||||
|
res := env["result"].(map[string]any)
|
||||||
|
require.False(t, res["isError"] == true, "download_url failed: %v", res)
|
||||||
|
url := res["structuredContent"].(map[string]any)["url"].(string)
|
||||||
|
url = ts.URL + url[strings.Index(url, "/mcp/"):]
|
||||||
|
resp, err := http.Get(url)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
require.Equal(t, 200, resp.StatusCode)
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
require.Equal(t, len(want), len(body), "100MiB body length mismatch")
|
||||||
|
require.True(t, bytes.Equal(want, body), "100MiB body content mismatch")
|
||||||
|
|
||||||
|
upPath := filepath.Join(dir, "up.bin")
|
||||||
|
env = e2eCall(t, ts, tok, "tools/call", "fs.upload_url", map[string]any{
|
||||||
|
"server_id": 7, "path": upPath, "ttl_seconds": 60,
|
||||||
|
})
|
||||||
|
res = env["result"].(map[string]any)
|
||||||
|
require.False(t, res["isError"] == true, "upload_url failed: %v", res)
|
||||||
|
upURL := res["structuredContent"].(map[string]any)["url"].(string)
|
||||||
|
upURL = ts.URL + upURL[strings.Index(upURL, "/mcp/"):]
|
||||||
|
req, _ := http.NewRequest("POST", upURL, bytes.NewReader(want))
|
||||||
|
req.ContentLength = int64(len(want))
|
||||||
|
upResp, err := http.DefaultClient.Do(req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer upResp.Body.Close()
|
||||||
|
require.Equal(t, 200, upResp.StatusCode)
|
||||||
|
got, _ := os.ReadFile(upPath)
|
||||||
|
require.Equal(t, len(want), len(got))
|
||||||
|
require.True(t, bytes.Equal(want, got))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestE2E_AuditRowsAreWritten(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupEndToEnd(t)
|
||||||
|
defer cleanup()
|
||||||
|
_ = e2eCall(t, ts, tok, "tools/call", "meta.whoami", map[string]any{})
|
||||||
|
_ = e2eCall(t, ts, tok, "tools/call", "server.list", map[string]any{})
|
||||||
|
|
||||||
|
require.Eventually(t, func() bool {
|
||||||
|
var cnt int64
|
||||||
|
_ = singleton.DB.Model(&model.MCPAuditLog{}).Count(&cnt).Error
|
||||||
|
return cnt >= 2
|
||||||
|
}, 3*time.Second, 20*time.Millisecond)
|
||||||
|
}
|
||||||
@@ -0,0 +1,67 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestServerExec_ErrorResult_PreservesStructuredContent(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
srv, _ := singleton.ServerShared.Get(7)
|
||||||
|
srv.SetTaskStream(&execErrorStream{errMsg: "agent disabled command execution"})
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerExec}, nil)
|
||||||
|
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{
|
||||||
|
Name: "server.exec",
|
||||||
|
Arguments: jsonRaw(map[string]any{
|
||||||
|
"server_id": 7,
|
||||||
|
"cmd": "whoami",
|
||||||
|
"timeout_seconds": 2,
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.NotNil(t, tcr)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "agent disabled command execution",
|
||||||
|
"error text must carry the real cause")
|
||||||
|
var result model.ExecResult
|
||||||
|
structuredJSON, err := json.Marshal(tcr.StructuredContent)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, json.Unmarshal(structuredJSON, &result))
|
||||||
|
require.Equal(t, -1, result.ExitCode)
|
||||||
|
require.Equal(t, "agent disabled command execution", result.Error)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestScopeDenied_OmitsStructuredContent(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{
|
||||||
|
Name: "server.exec",
|
||||||
|
Arguments: jsonRaw(map[string]any{"server_id": 7, "cmd": "echo"}),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.NotNil(t, tcr)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Nil(t, tcr.StructuredContent)
|
||||||
|
}
|
||||||
@@ -0,0 +1,226 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"google.golang.org/grpc/metadata"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
pb "github.com/nezhahq/nezha/proto"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// killSwitchStream is a minimal RequestTask stream that just records sent
|
||||||
|
// tasks; it never replies. CallAgent under this stream blocks until the
|
||||||
|
// kill switch wakes it up, which is exactly the behaviour these tests
|
||||||
|
// pin down.
|
||||||
|
type killSwitchStream struct {
|
||||||
|
sent chan *pb.Task
|
||||||
|
}
|
||||||
|
|
||||||
|
func newKillSwitchStream() *killSwitchStream {
|
||||||
|
return &killSwitchStream{sent: make(chan *pb.Task, 4)}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *killSwitchStream) Send(t *pb.Task) error { s.sent <- t; return nil }
|
||||||
|
func (s *killSwitchStream) Recv() (*pb.TaskResult, error) { return nil, context.Canceled }
|
||||||
|
func (s *killSwitchStream) SetHeader(metadata.MD) error { return nil }
|
||||||
|
func (s *killSwitchStream) SendHeader(metadata.MD) error { return nil }
|
||||||
|
func (s *killSwitchStream) SetTrailer(metadata.MD) {}
|
||||||
|
func (s *killSwitchStream) Context() context.Context { return context.Background() }
|
||||||
|
func (s *killSwitchStream) SendMsg(any) error { return nil }
|
||||||
|
func (s *killSwitchStream) RecvMsg(any) error { return context.Canceled }
|
||||||
|
|
||||||
|
func TestRevalidateTransferEntry_BlocksWhenMCPDisabled(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
singleton.Conf.SetMCPEnabled(false)
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
entry := &transferEntry{
|
||||||
|
UserID: uid,
|
||||||
|
TokenID: tok.ID,
|
||||||
|
ServerID: 7,
|
||||||
|
Path: "/srv/file",
|
||||||
|
Direction: transferDirDownload,
|
||||||
|
ExpiresAt: time.Now().Add(5 * time.Minute),
|
||||||
|
}
|
||||||
|
|
||||||
|
err := revalidateTransferEntry(entry)
|
||||||
|
require.Error(t, err, "revalidate must reject when EnableMCP=false")
|
||||||
|
require.Contains(t, err.Error(), "MCP is disabled",
|
||||||
|
"error message must surface kill switch reason, not look like a transient agent fault")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPurgeTransferEntries_DropsMintedTokens(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
_, err := mintTransferToken(transferEntry{
|
||||||
|
UserID: uid,
|
||||||
|
TokenID: tok.ID,
|
||||||
|
ServerID: 7,
|
||||||
|
Path: "/srv/blob",
|
||||||
|
Direction: transferDirDownload,
|
||||||
|
ExpiresAt: time.Now().Add(5 * time.Minute),
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
purged := PurgeTransferEntries()
|
||||||
|
require.GreaterOrEqual(t, purged, 3, "all minted entries must be dropped")
|
||||||
|
|
||||||
|
count := 0
|
||||||
|
transferEntries.Range(func(_, _ any) bool { count++; return true })
|
||||||
|
require.Equal(t, 0, count, "transferEntries must be empty after purge")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRevokeStreamsForPurpose_OnlyTouchesMatchingPurpose(t *testing.T) {
|
||||||
|
h := rpc.NewNezhaHandler()
|
||||||
|
h.CreateStreamWithPurpose("legacy-1", 0, 1, rpc.PurposeLegacy)
|
||||||
|
h.CreateStreamWithPurpose("mcp-1", 0, 1, rpc.PurposeMCPTransfer)
|
||||||
|
h.CreateStreamWithPurpose("mcp-2", 0, 2, rpc.PurposeMCPTransfer)
|
||||||
|
|
||||||
|
revoked := h.RevokeStreamsForPurpose(rpc.PurposeMCPTransfer)
|
||||||
|
require.Equal(t, 2, revoked, "kill switch must take down both MCP streams")
|
||||||
|
|
||||||
|
_, legacyErr := h.GetStream("legacy-1")
|
||||||
|
require.NoError(t, legacyErr,
|
||||||
|
"legacy purpose streams (terminal/fm/nat) must NOT be revoked by the MCP kill switch")
|
||||||
|
_, mcp1Err := h.GetStream("mcp-1")
|
||||||
|
require.Error(t, mcp1Err, "mcp-1 must be gone after revoke")
|
||||||
|
_, mcp2Err := h.GetStream("mcp-2")
|
||||||
|
require.Error(t, mcp2Err, "mcp-2 must be gone after revoke")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCancelAllMCPInflight_UnblocksCallAgent(t *testing.T) {
|
||||||
|
stream := newKillSwitchStream()
|
||||||
|
original := singleton.ServerShared
|
||||||
|
sc := singleton.NewEmptyServerClassForTest()
|
||||||
|
srv := &model.Server{}
|
||||||
|
srv.ID = 88
|
||||||
|
srv.SetTaskStream(stream)
|
||||||
|
sc.InsertForTest(srv)
|
||||||
|
singleton.ServerShared = sc
|
||||||
|
t.Cleanup(func() { singleton.ServerShared = original })
|
||||||
|
|
||||||
|
done := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
_, err := rpc.CallAgent(context.Background(), 88, model.TaskTypeExec,
|
||||||
|
model.ExecRequest{Cmd: "sleep"}, 30*time.Second)
|
||||||
|
done <- err
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-stream.sent:
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatalf("CallAgent never reached stream.Send within 1s")
|
||||||
|
}
|
||||||
|
|
||||||
|
rpc.CancelAllMCPInflight()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case err := <-done:
|
||||||
|
require.ErrorIs(t, err, rpc.ErrMCPDisabled,
|
||||||
|
"CallAgent must surface ErrMCPDisabled when kill switch fires; got %v", err)
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatalf("CallAgent did not return after CancelAllMCPInflight; kill switch is broken")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateConfig_DisablingMCPInvokesKillSwitch(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
installTestConfig(t)
|
||||||
|
singleton.Conf.SetMCPEnabled(true)
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
_, err := mintTransferToken(transferEntry{
|
||||||
|
UserID: uid,
|
||||||
|
TokenID: tok.ID,
|
||||||
|
ServerID: 7,
|
||||||
|
Path: "/srv/blob",
|
||||||
|
Direction: transferDirDownload,
|
||||||
|
ExpiresAt: time.Now().Add(5 * time.Minute),
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
rpc.NezhaHandlerSingleton.CreateStreamWithPurpose("mcp-active", 0, 7, rpc.PurposeMCPTransfer)
|
||||||
|
|
||||||
|
stream := newKillSwitchStream()
|
||||||
|
sc := singleton.NewEmptyServerClassForTest()
|
||||||
|
srv := &model.Server{}
|
||||||
|
srv.ID = 7
|
||||||
|
srv.SetTaskStream(stream)
|
||||||
|
sc.InsertForTest(srv)
|
||||||
|
originalShared := singleton.ServerShared
|
||||||
|
singleton.ServerShared = sc
|
||||||
|
t.Cleanup(func() { singleton.ServerShared = originalShared })
|
||||||
|
|
||||||
|
rpcDone := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
_, err := rpc.CallAgent(context.Background(), 7, model.TaskTypeFsRead,
|
||||||
|
model.FsReadRequest{Path: "/x"}, 30*time.Second)
|
||||||
|
rpcDone <- err
|
||||||
|
}()
|
||||||
|
select {
|
||||||
|
case <-stream.sent:
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatalf("background CallAgent never reached the stream")
|
||||||
|
}
|
||||||
|
|
||||||
|
origTemplates := singleton.FrontendTemplates
|
||||||
|
singleton.FrontendTemplates = []model.FrontendTemplate{{Path: "user-dist", IsAdmin: false}}
|
||||||
|
defer func() { singleton.FrontendTemplates = origTemplates }()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.Use(func(c *gin.Context) {
|
||||||
|
setAuthUser(c, uid, model.RoleAdmin)
|
||||||
|
c.Next()
|
||||||
|
})
|
||||||
|
r.PATCH("/api/v1/setting", commonHandler(updateConfig))
|
||||||
|
settingBody := map[string]any{
|
||||||
|
"site_name": "test",
|
||||||
|
"language": "en_US",
|
||||||
|
"user_template": "user-dist",
|
||||||
|
"enable_mcp": false,
|
||||||
|
}
|
||||||
|
raw, _ := json.Marshal(settingBody)
|
||||||
|
req := httptest.NewRequest(http.MethodPatch, "/api/v1/setting", bytes.NewReader(raw))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
success, errMsg := decodeCommonResponseError(t, w.Body.Bytes())
|
||||||
|
require.True(t, success, "PATCH /setting must succeed: %s", errMsg)
|
||||||
|
require.False(t, singleton.Conf.EnableMCP, "config must reflect kill switch state")
|
||||||
|
|
||||||
|
count := 0
|
||||||
|
transferEntries.Range(func(_, _ any) bool { count++; return true })
|
||||||
|
require.Equal(t, 0, count, "unconsumed transfer URLs must be purged")
|
||||||
|
|
||||||
|
_, streamErr := rpc.NezhaHandlerSingleton.GetStream("mcp-active")
|
||||||
|
require.Error(t, streamErr, "active MCP IOStream must be revoked")
|
||||||
|
|
||||||
|
select {
|
||||||
|
case err := <-rpcDone:
|
||||||
|
require.True(t, errors.Is(err, rpc.ErrMCPDisabled),
|
||||||
|
"in-flight CallAgent must wake up with ErrMCPDisabled, got %v", err)
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatalf("in-flight CallAgent did not wake up after kill switch")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,77 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io/fs"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Streamable HTTP 规范(modelcontextprotocol.io /basic/transports)要求:
|
||||||
|
// 服务端如果不提供 standalone SSE,必须对 GET /mcp 返回 405 Method Not Allowed。
|
||||||
|
// 现状是 Gin 的 NoRoute fallback 会把 GET /mcp 喂给前端 fallback(HTML/404),
|
||||||
|
// 真实 MCP 客户端在自动探测 SSE 时会卡住或拿到无效内容。
|
||||||
|
//
|
||||||
|
// 这条测试拼出和生产 routers() 一致的 /mcp 三件套,仅断言「非 POST 不返回 HTML」。
|
||||||
|
type mcpFallbackDist struct{}
|
||||||
|
|
||||||
|
func (mcpFallbackDist) Open(string) (fs.File, error) { return nil, fs.ErrNotExist }
|
||||||
|
|
||||||
|
func setupMCPMethodRouter(t *testing.T) *gin.Engine {
|
||||||
|
t.Helper()
|
||||||
|
originalConf := singleton.Conf
|
||||||
|
singleton.Conf = &singleton.ConfigClass{Config: &model.Config{
|
||||||
|
ConfigDashboard: model.ConfigDashboard{
|
||||||
|
AdminTemplate: "admin-dist",
|
||||||
|
UserTemplate: "user-dist",
|
||||||
|
},
|
||||||
|
}}
|
||||||
|
t.Cleanup(func() { singleton.Conf = originalConf })
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.POST("/mcp", mcpOriginGuard(), apiTokenAuthMiddleware(), mcpEndpoint)
|
||||||
|
r.GET("/mcp", mcpMethodNotAllowed)
|
||||||
|
r.DELETE("/mcp", mcpMethodNotAllowed)
|
||||||
|
r.GET("/mcp/download/:token", mcpOriginGuard(), transferDownloadHandler)
|
||||||
|
r.POST("/mcp/upload/:token", mcpOriginGuard(), transferUploadHandler)
|
||||||
|
r.NoRoute(fallbackToFrontend(mcpFallbackDist{}))
|
||||||
|
return r
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_GetReturnsMethodNotAllowed(t *testing.T) {
|
||||||
|
t.Chdir(t.TempDir())
|
||||||
|
r := setupMCPMethodRouter(t)
|
||||||
|
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/mcp", nil)
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
if w.Code != http.StatusMethodNotAllowed {
|
||||||
|
t.Fatalf("GET /mcp must return 405 per Streamable HTTP spec; got %d body=%q",
|
||||||
|
w.Code, w.Body.String())
|
||||||
|
}
|
||||||
|
if strings.Contains(strings.ToLower(w.Body.String()), "<html") {
|
||||||
|
t.Fatalf("GET /mcp must not fall back to the SPA index.html; body=%q", w.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_DeleteReturnsMethodNotAllowed(t *testing.T) {
|
||||||
|
t.Chdir(t.TempDir())
|
||||||
|
r := setupMCPMethodRouter(t)
|
||||||
|
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodDelete, "/mcp", nil)
|
||||||
|
r.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
if w.Code != http.StatusMethodNotAllowed {
|
||||||
|
t.Fatalf("DELETE /mcp (session terminate) must return 405 when sessions are not implemented; got %d",
|
||||||
|
w.Code)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,117 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func mcpOriginGuard() gin.HandlerFunc {
|
||||||
|
return func(c *gin.Context) {
|
||||||
|
origin := strings.TrimSpace(c.GetHeader("Origin"))
|
||||||
|
if origin == "" {
|
||||||
|
c.Next()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
u, err := url.Parse(origin)
|
||||||
|
if err != nil || u.Host == "" {
|
||||||
|
abortOrigin(c)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !strings.EqualFold(u.Host, c.Request.Host) {
|
||||||
|
abortOrigin(c)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
// DNS rebinding 防线:当 dashboard 明确仅暴露在 loopback 上时,浏览器只
|
||||||
|
// 应通过 loopback hostname 命中本服务。任何带 Origin 的请求若 Host 不是
|
||||||
|
// loopback 字面量,说明它解析自一个对外 DNS 名(典型攻击:evil.example
|
||||||
|
// 解析到 127.0.0.1,Host == Origin == evil.example 等式成立但实际打的是
|
||||||
|
// 用户本机 dashboard)。绑公网 IP / 未指定地址 / 公网 hostname 的部署
|
||||||
|
// 跳过此检查 —— 那种部署下 Host 本来就是公网域名,再要求 loopback
|
||||||
|
// 会把合法的同源前端请求一刀切。
|
||||||
|
if dashboardListensOnLoopback() && !isLoopbackHostname(hostnameOnly(c.Request.Host)) {
|
||||||
|
abortOrigin(c)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.Next()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// dashboardListensOnLoopback 判定 dashboard 是否仅暴露在 loopback 上。
|
||||||
|
//
|
||||||
|
// 之前把未指定地址(空字符串、`0.0.0.0`、`::`、`*`)也当 loopback 部署,理由是
|
||||||
|
// 这些场景下浏览器仍可能通过 127.0.0.1 访问,rebinding 风险存在。问题是绝大多数
|
||||||
|
// 生产部署就是 `listen_host=""` / `0.0.0.0`,配合公网域名访问;那种部署下浏览器
|
||||||
|
// 同源 POST 的 Host 是公网域名而不是 loopback,旧逻辑会把合法的同源前端请求一刀切。
|
||||||
|
//
|
||||||
|
// 现在只在 ListenHost 明确写成 loopback 字面量时才开启 loopback 严格化:
|
||||||
|
// - 显式 loopback IP(127.0.0.1 / ::1)或 localhost → 视为 loopback 部署,
|
||||||
|
// 此时任何带 Origin 的请求 Host 不是 loopback 都拒掉(DNS rebinding 防线);
|
||||||
|
// - 空 / `0.0.0.0` / `::` / `*` / 公网 IP / hostname → 不当 loopback 部署,
|
||||||
|
// 仅做 Origin == Host 的同源校验,不再额外要求 Host 必须是 loopback。
|
||||||
|
//
|
||||||
|
// 调用方仍然在配额(Origin == Host)外用 isLoopbackHostname 校验 Host,所以
|
||||||
|
// 当本机用户拿 127.0.0.1 访问绑公网 IP 的 dashboard 时这条防线依然生效(Host
|
||||||
|
// 是 loopback,Origin 是公网域名,hostnameOnly 比较会失败)。
|
||||||
|
func dashboardListensOnLoopback() bool {
|
||||||
|
if singleton.Conf == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
host := strings.TrimSpace(singleton.Conf.ListenHost)
|
||||||
|
if host == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
host = strings.Trim(host, "[]")
|
||||||
|
switch host {
|
||||||
|
case "0.0.0.0", "::", "*":
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if ip := net.ParseIP(host); ip != nil {
|
||||||
|
return ip.IsLoopback()
|
||||||
|
}
|
||||||
|
return strings.EqualFold(host, "localhost")
|
||||||
|
}
|
||||||
|
|
||||||
|
func hostnameOnly(hostport string) string {
|
||||||
|
if i := strings.LastIndex(hostport, ":"); i >= 0 {
|
||||||
|
if !strings.Contains(hostport[i:], "]") {
|
||||||
|
return hostport[:i]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return hostport
|
||||||
|
}
|
||||||
|
|
||||||
|
func isLoopbackHostname(host string) bool {
|
||||||
|
host = strings.Trim(host, "[]")
|
||||||
|
if host == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if strings.EqualFold(host, "localhost") {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if ip := net.ParseIP(host); ip != nil {
|
||||||
|
return ip.IsLoopback()
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func abortOrigin(c *gin.Context) {
|
||||||
|
c.AbortWithStatusJSON(http.StatusForbidden, model.CommonResponse[any]{
|
||||||
|
Success: false,
|
||||||
|
Error: "ApiErrorForbidden: origin not allowed",
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func mcpMethodNotAllowed(c *gin.Context) {
|
||||||
|
c.Header("Allow", "POST")
|
||||||
|
c.JSON(http.StatusMethodNotAllowed, model.CommonResponse[any]{
|
||||||
|
Success: false,
|
||||||
|
Error: "MCP endpoint only accepts POST (Streamable HTTP without standalone SSE / sessions)",
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -0,0 +1,109 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func setupMCPOriginRouter(t *testing.T) (*httptest.Server, string, func()) {
|
||||||
|
t.Helper()
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
_, plain := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.POST("/mcp", mcpOriginGuard(), apiTokenAuthMiddleware(), mcpEndpoint)
|
||||||
|
r.GET("/mcp/download/:token", mcpOriginGuard(), transferDownloadHandler)
|
||||||
|
r.POST("/mcp/upload/:token", mcpOriginGuard(), transferUploadHandler)
|
||||||
|
ts := httptest.NewServer(r)
|
||||||
|
return ts, plain, func() {
|
||||||
|
ts.Close()
|
||||||
|
cleanup()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_DisallowsCrossOriginRequest(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupMCPOriginRouter(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader([]byte(`{"jsonrpc":"2.0","id":1,"method":"initialize"}`)))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("Authorization", "Bearer "+tok)
|
||||||
|
req.Header.Set("Origin", "http://evil.example.com")
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
require.Equal(t, http.StatusForbidden, resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_AllowsRequestWithoutOriginHeader(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupMCPOriginRouter(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader([]byte(`{"jsonrpc":"2.0","id":1,"method":"initialize"}`)))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("Authorization", "Bearer "+tok)
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
require.Equal(t, http.StatusOK, resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_AllowsSameHostOrigin(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupMCPOriginRouter(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader([]byte(`{"jsonrpc":"2.0","id":1,"method":"initialize"}`)))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("Authorization", "Bearer "+tok)
|
||||||
|
req.Header.Set("Origin", "http://"+req.Host)
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
require.Equal(t, http.StatusOK, resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 公网部署回归:ListenHost 是 0.0.0.0/未指定时,前端会以公网 Host 同源访问。
|
||||||
|
// 这条以前会被 dashboardListensOnLoopback 误判为 loopback 部署进而拒掉;
|
||||||
|
// 现在必须放行,否则正常生产环境的 admin frontend MCP 入口直接 403。
|
||||||
|
func TestMCP_PublicDeployment_AllowsPublicSameOrigin(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupMCPOriginRouter(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader([]byte(`{"jsonrpc":"2.0","id":1,"method":"initialize"}`)))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("Authorization", "Bearer "+tok)
|
||||||
|
req.Host = "dashboard.example.com"
|
||||||
|
req.Header.Set("Origin", "https://dashboard.example.com")
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
require.Equal(t, http.StatusOK, resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 显式绑 loopback 时仍然要执行 DNS rebinding 防线:Host 是公网域名 → 403。
|
||||||
|
func TestMCP_LoopbackDeployment_RejectsPublicHost(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupMCPOriginRouter(t)
|
||||||
|
defer cleanup()
|
||||||
|
prev := singleton.Conf.ListenHost
|
||||||
|
singleton.Conf.ListenHost = "127.0.0.1"
|
||||||
|
defer func() { singleton.Conf.ListenHost = prev }()
|
||||||
|
|
||||||
|
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader([]byte(`{"jsonrpc":"2.0","id":1,"method":"initialize"}`)))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
req.Header.Set("Authorization", "Bearer "+tok)
|
||||||
|
req.Host = "dashboard.example.com"
|
||||||
|
req.Header.Set("Origin", "https://dashboard.example.com")
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
require.Equal(t, http.StatusForbidden, resp.StatusCode)
|
||||||
|
}
|
||||||
@@ -0,0 +1,98 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// MCPRateLimiter 实现按 token 的双层 token bucket:
|
||||||
|
// - 秒级:默认 10 req/s,应对单个 LLM 突发
|
||||||
|
// - 分钟级:默认 120 req/min,应对长时间刷
|
||||||
|
//
|
||||||
|
// 实现走简单 sliding window(按桶截断的 counter),轻量、O(1);
|
||||||
|
// 进程内即可,未持久化——重启等价于配额刷新,可接受。
|
||||||
|
type MCPRateLimiter struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
perToken map[uint64]*tokenWindow
|
||||||
|
secLimit int
|
||||||
|
minLimit int
|
||||||
|
lastPrune time.Time
|
||||||
|
clock func() time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
type tokenWindow struct {
|
||||||
|
secBucketStart time.Time
|
||||||
|
secCount int
|
||||||
|
minBucketStart time.Time
|
||||||
|
minCount int
|
||||||
|
}
|
||||||
|
|
||||||
|
// mcpRateLimiterPruneInterval bounds how often Allow sweeps the map. Without
|
||||||
|
// eviction the map kept one entry per token ID ever seen, so PAT churn grew
|
||||||
|
// it without bound. A token idle longer than its minute bucket carries no
|
||||||
|
// live budget, so dropping it is lossless; the interval keeps the sweep
|
||||||
|
// amortized O(1) per call instead of O(map) every call.
|
||||||
|
const mcpRateLimiterPruneInterval = time.Minute
|
||||||
|
|
||||||
|
func newMCPRateLimiter(secLimit, minLimit int) *MCPRateLimiter {
|
||||||
|
return newMCPRateLimiterWithClock(secLimit, minLimit, time.Now)
|
||||||
|
}
|
||||||
|
|
||||||
|
func newMCPRateLimiterWithClock(secLimit, minLimit int, clock func() time.Time) *MCPRateLimiter {
|
||||||
|
return &MCPRateLimiter{
|
||||||
|
perToken: make(map[uint64]*tokenWindow),
|
||||||
|
secLimit: secLimit,
|
||||||
|
minLimit: minLimit,
|
||||||
|
clock: clock,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// pruneStaleLocked drops windows whose minute bucket started more than one
|
||||||
|
// minute ago: such a token has no accumulated budget left, so removing it
|
||||||
|
// cannot change any future Allow decision. Caller must hold r.mu.
|
||||||
|
func (r *MCPRateLimiter) pruneStaleLocked(now time.Time) {
|
||||||
|
for id, w := range r.perToken {
|
||||||
|
if now.Sub(w.minBucketStart) >= time.Minute {
|
||||||
|
delete(r.perToken, id)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Allow 返回是否允许本次调用。被拒返回 false。
|
||||||
|
// tokenID = 0 时不限流(管理路径或匿名)。
|
||||||
|
func (r *MCPRateLimiter) Allow(tokenID uint64) bool {
|
||||||
|
if tokenID == 0 {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
// Sampling while holding the state lock linearizes the clock read with the
|
||||||
|
// bucket update, so a delayed sample cannot commit after a newer rollover.
|
||||||
|
now := r.clock()
|
||||||
|
if now.Sub(r.lastPrune) >= mcpRateLimiterPruneInterval {
|
||||||
|
r.pruneStaleLocked(now)
|
||||||
|
r.lastPrune = now
|
||||||
|
}
|
||||||
|
w, ok := r.perToken[tokenID]
|
||||||
|
if !ok {
|
||||||
|
w = &tokenWindow{secBucketStart: now, minBucketStart: now}
|
||||||
|
r.perToken[tokenID] = w
|
||||||
|
}
|
||||||
|
if now.Sub(w.secBucketStart) >= time.Second {
|
||||||
|
w.secBucketStart = now
|
||||||
|
w.secCount = 0
|
||||||
|
}
|
||||||
|
if now.Sub(w.minBucketStart) >= time.Minute {
|
||||||
|
w.minBucketStart = now
|
||||||
|
w.minCount = 0
|
||||||
|
}
|
||||||
|
if w.secCount >= r.secLimit || w.minCount >= r.minLimit {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
w.secCount++
|
||||||
|
w.minCount++
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// 全局单例,参数固定(生产可观测后再考虑配置化)。
|
||||||
|
var mcpRateLimiterShared = newMCPRateLimiter(10, 120)
|
||||||
@@ -0,0 +1,51 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
type agentcompatMCPRateLimitProbeRequest struct {
|
||||||
|
}
|
||||||
|
|
||||||
|
type agentcompatMCPRateLimitProbeResponse struct {
|
||||||
|
SecondAllowedCount int `json:"second_allowed_count"`
|
||||||
|
SecondRejectedAtCount int `json:"second_rejected_at_count"`
|
||||||
|
MinuteAllowedCount int `json:"minute_allowed_count"`
|
||||||
|
MinuteRejectedAtCount int `json:"minute_rejected_at_count"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func agentcompatMCPRateLimitProbeRoute(context *gin.Context) (agentcompatMCPRateLimitProbeResponse, error) {
|
||||||
|
var request agentcompatMCPRateLimitProbeRequest
|
||||||
|
if err := decodeAgentcompatJSON(context, &request); err != nil {
|
||||||
|
return agentcompatMCPRateLimitProbeResponse{}, err
|
||||||
|
}
|
||||||
|
return runAgentcompatMCPRateLimitProbe(request)
|
||||||
|
}
|
||||||
|
|
||||||
|
func runAgentcompatMCPRateLimitProbe(request agentcompatMCPRateLimitProbeRequest) (agentcompatMCPRateLimitProbeResponse, error) {
|
||||||
|
response := agentcompatMCPRateLimitProbeResponse{}
|
||||||
|
secondNow := time.Unix(1_700_000_000, 0)
|
||||||
|
secondLimiter := newMCPRateLimiterWithClock(10, 120, func() time.Time { return secondNow })
|
||||||
|
for requestNumber := 1; requestNumber <= 11; requestNumber++ {
|
||||||
|
if secondLimiter.Allow(1) {
|
||||||
|
response.SecondAllowedCount++
|
||||||
|
} else if response.SecondRejectedAtCount == 0 {
|
||||||
|
response.SecondRejectedAtCount = requestNumber
|
||||||
|
}
|
||||||
|
}
|
||||||
|
minuteNow := time.Unix(1_700_000_000, 0)
|
||||||
|
minuteLimiter := newMCPRateLimiterWithClock(10_000, 120, func() time.Time { return minuteNow })
|
||||||
|
for requestNumber := 1; requestNumber <= 121; requestNumber++ {
|
||||||
|
if minuteLimiter.Allow(1) {
|
||||||
|
response.MinuteAllowedCount++
|
||||||
|
} else if response.MinuteRejectedAtCount == 0 {
|
||||||
|
response.MinuteRejectedAtCount = requestNumber
|
||||||
|
}
|
||||||
|
minuteNow = minuteNow.Add(100 * time.Millisecond)
|
||||||
|
}
|
||||||
|
return response, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,102 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http/httptest"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAgentcompatMCPRateLimitProbe_returnsTypedBoundaryCountsWithoutSharedState(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
request := agentcompatMCPRateLimitProbeRequest{}
|
||||||
|
originalLimiter := mcpRateLimiterShared
|
||||||
|
|
||||||
|
// When
|
||||||
|
response, err := runAgentcompatMCPRateLimitProbe(request)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("probe returned error: %v", err)
|
||||||
|
}
|
||||||
|
if response.SecondAllowedCount != 10 || response.SecondRejectedAtCount != 11 || response.MinuteAllowedCount != 120 || response.MinuteRejectedAtCount != 121 {
|
||||||
|
t.Fatalf("probe result = %+v, want second=10/11 minute=120/121", response)
|
||||||
|
}
|
||||||
|
t.Logf("typed probe result: second=%d/%d minute=%d/%d", response.SecondAllowedCount, response.SecondRejectedAtCount, response.MinuteAllowedCount, response.MinuteRejectedAtCount)
|
||||||
|
if mcpRateLimiterShared != originalLimiter {
|
||||||
|
t.Fatal("probe mutated the shared production limiter")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatMCPRateLimitProbe_isRepeatableAndConcurrentSafe(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
request := agentcompatMCPRateLimitProbeRequest{}
|
||||||
|
results := make(chan agentcompatMCPRateLimitProbeResponse, 8)
|
||||||
|
errors := make(chan error, 8)
|
||||||
|
|
||||||
|
// When
|
||||||
|
var waitGroup sync.WaitGroup
|
||||||
|
for probeNumber := 0; probeNumber < 8; probeNumber++ {
|
||||||
|
waitGroup.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer waitGroup.Done()
|
||||||
|
response, err := runAgentcompatMCPRateLimitProbe(request)
|
||||||
|
if err != nil {
|
||||||
|
errors <- err
|
||||||
|
return
|
||||||
|
}
|
||||||
|
results <- response
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
waitGroup.Wait()
|
||||||
|
close(results)
|
||||||
|
close(errors)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
for err := range errors {
|
||||||
|
t.Fatalf("concurrent probe returned error: %v", err)
|
||||||
|
}
|
||||||
|
for response := range results {
|
||||||
|
if response.SecondAllowedCount != 10 || response.SecondRejectedAtCount != 11 || response.MinuteAllowedCount != 120 || response.MinuteRejectedAtCount != 121 {
|
||||||
|
t.Fatalf("concurrent probe result = %+v, want second=10/11 minute=120/121", response)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatMCPRateLimitProbe_rejectsCallerControlledParameters(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
context, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||||
|
context.Request = httptest.NewRequest("POST", "/", bytes.NewBufferString(`{"token_id":7}`))
|
||||||
|
context.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
// When
|
||||||
|
var request agentcompatMCPRateLimitProbeRequest
|
||||||
|
err := decodeAgentcompatJSON(context, &request)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("caller-controlled rate probe parameters must be rejected")
|
||||||
|
}
|
||||||
|
context.Request = httptest.NewRequest("POST", "/", bytes.NewBufferString(`{}`))
|
||||||
|
if err := decodeAgentcompatJSON(context, &request); err != nil {
|
||||||
|
t.Fatalf("canonical empty probe request must be accepted: %v", err)
|
||||||
|
}
|
||||||
|
if _, err := json.Marshal(request); err != nil {
|
||||||
|
t.Fatalf("canonical request must remain JSON encodable: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentcompatMCPRateLimitProbeRejectsTrailingJSON(t *testing.T) {
|
||||||
|
context, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||||
|
context.Request = httptest.NewRequest("POST", "/", bytes.NewBufferString(`{}{}`))
|
||||||
|
var request agentcompatMCPRateLimitProbeRequest
|
||||||
|
if err := decodeAgentcompatJSON(context, &request); err == nil {
|
||||||
|
t.Fatal("trailing JSON values must be rejected")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,197 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
type mcpRateLimitFakeClock struct {
|
||||||
|
now time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
func (clock *mcpRateLimitFakeClock) Now() time.Time {
|
||||||
|
return clock.now
|
||||||
|
}
|
||||||
|
|
||||||
|
func (clock *mcpRateLimitFakeClock) Advance(duration time.Duration) {
|
||||||
|
clock.now = clock.now.Add(duration)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPRateLimiter_baselinePreservesAnonymousBypassAndBucketReset(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
limiter := newMCPRateLimiter(3, 3)
|
||||||
|
|
||||||
|
// When
|
||||||
|
anonymousAllowed := limiter.Allow(0)
|
||||||
|
firstAllowed := limiter.Allow(7)
|
||||||
|
secondAllowed := limiter.Allow(7)
|
||||||
|
thirdAllowed := limiter.Allow(7)
|
||||||
|
fourthRejected := limiter.Allow(7)
|
||||||
|
// Then
|
||||||
|
if !anonymousAllowed || !firstAllowed || !secondAllowed || !thirdAllowed || fourthRejected {
|
||||||
|
t.Fatalf("baseline limiter semantics changed: anonymous=%t first=%t second=%t third=%t rejected=%t", anonymousAllowed, firstAllowed, secondAllowed, thirdAllowed, fourthRejected)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPRateLimiter_allowsTenAndRejectsElevenWithinOneSecond(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
|
||||||
|
limiter := newMCPRateLimiterWithClock(10, 120, clock.Now)
|
||||||
|
|
||||||
|
// When
|
||||||
|
allowed := 0
|
||||||
|
for requestNumber := 1; requestNumber <= 11; requestNumber++ {
|
||||||
|
if limiter.Allow(7) {
|
||||||
|
allowed++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if allowed != 10 {
|
||||||
|
t.Fatalf("allowed count = %d, want 10", allowed)
|
||||||
|
}
|
||||||
|
if limiter.Allow(7) {
|
||||||
|
t.Fatal("twelfth request must remain rejected in the same second")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPRateLimiter_allowsOneAfterSecondBucketRollover(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
|
||||||
|
limiter := newMCPRateLimiterWithClock(10, 120, clock.Now)
|
||||||
|
for requestNumber := 1; requestNumber <= 10; requestNumber++ {
|
||||||
|
if !limiter.Allow(7) {
|
||||||
|
t.Fatalf("request %d must be allowed before rollover", requestNumber)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// When
|
||||||
|
clock.Advance(time.Second)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if !limiter.Allow(7) {
|
||||||
|
t.Fatal("request after one-second bucket rollover must be allowed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPRateLimiter_allows120AndRejects121WithinOneMinute(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
|
||||||
|
limiter := newMCPRateLimiterWithClock(10_000, 120, clock.Now)
|
||||||
|
|
||||||
|
// When
|
||||||
|
allowed := 0
|
||||||
|
for requestNumber := 1; requestNumber <= 121; requestNumber++ {
|
||||||
|
if limiter.Allow(7) {
|
||||||
|
allowed++
|
||||||
|
}
|
||||||
|
clock.Advance(100 * time.Millisecond)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if allowed != 120 {
|
||||||
|
t.Fatalf("allowed count = %d, want 120", allowed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPRateLimiter_independentTokensKeepIndependentBudgets(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
|
||||||
|
limiter := newMCPRateLimiterWithClock(1, 120, clock.Now)
|
||||||
|
|
||||||
|
// When
|
||||||
|
firstTokenAllowed := limiter.Allow(7)
|
||||||
|
firstTokenRejected := limiter.Allow(7)
|
||||||
|
secondTokenAllowed := limiter.Allow(8)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if !firstTokenAllowed || firstTokenRejected || !secondTokenAllowed {
|
||||||
|
t.Fatalf("token budgets are not independent: first allowed=%t rejected=%t second allowed=%t", firstTokenAllowed, firstTokenRejected, secondTokenAllowed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPRateLimiter_clockControlledCallsDoNotMutateSharedLimiter(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
originalLimiter := mcpRateLimiterShared
|
||||||
|
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
|
||||||
|
limiter := newMCPRateLimiterWithClock(1, 1, clock.Now)
|
||||||
|
|
||||||
|
// When
|
||||||
|
if !limiter.Allow(7) {
|
||||||
|
t.Fatal("isolated limiter request must be allowed")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if mcpRateLimiterShared != originalLimiter {
|
||||||
|
t.Fatal("isolated limiter changed shared production limiter ownership")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPRateLimiter_concurrentCallsAtSecondBoundaryAllowExactlyLimit(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
|
||||||
|
limiter := newMCPRateLimiterWithClock(10, 120, clock.Now)
|
||||||
|
var allowed atomic.Int32
|
||||||
|
var waitGroup sync.WaitGroup
|
||||||
|
|
||||||
|
// When
|
||||||
|
for requestNumber := 0; requestNumber < 20; requestNumber++ {
|
||||||
|
waitGroup.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer waitGroup.Done()
|
||||||
|
if limiter.Allow(7) {
|
||||||
|
allowed.Add(1)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
waitGroup.Wait()
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if got := allowed.Load(); got != 10 {
|
||||||
|
t.Fatalf("concurrent allowed count = %d, want 10", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPRateLimiter_minuteBucketRolloverRestoresBudget(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
|
||||||
|
limiter := newMCPRateLimiterWithClock(10_000, 1, clock.Now)
|
||||||
|
if !limiter.Allow(7) {
|
||||||
|
t.Fatal("first request must be allowed")
|
||||||
|
}
|
||||||
|
if limiter.Allow(7) {
|
||||||
|
t.Fatal("second request must be rejected before minute rollover")
|
||||||
|
}
|
||||||
|
|
||||||
|
// When
|
||||||
|
clock.Advance(time.Minute)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if !limiter.Allow(7) {
|
||||||
|
t.Fatal("request after minute rollover must be allowed")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPRateLimiter_samplesClockWhileHoldingStateLock(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
now := time.Unix(1_700_000_000, 0)
|
||||||
|
var limiter *MCPRateLimiter
|
||||||
|
clock := func() time.Time {
|
||||||
|
if limiter.mu.TryLock() {
|
||||||
|
limiter.mu.Unlock()
|
||||||
|
t.Fatal("clock callback acquired limiter state lock; Allow sampled before locking")
|
||||||
|
}
|
||||||
|
return now
|
||||||
|
}
|
||||||
|
limiter = newMCPRateLimiterWithClock(1, 1, clock)
|
||||||
|
|
||||||
|
// When
|
||||||
|
allowed := limiter.Allow(7)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if !allowed {
|
||||||
|
t.Fatal("initial request must be allowed")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,86 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func mcpEndpointTestCtx(t *testing.T, tok *model.APIToken, body any) (*gin.Context, *httptest.ResponseRecorder, error) {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, engine := gin.CreateTestContext(w)
|
||||||
|
require.NotNil(t, engine)
|
||||||
|
raw, err := json.Marshal(body)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
|
c.Request = httptest.NewRequest("POST", "/mcp", bytes.NewReader(raw))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
c.Set(apiTokenCtxKey, tok)
|
||||||
|
c.Set(model.CtxKeyAPIToken, tok)
|
||||||
|
return c, w, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPEndpointTestCtx_surfacesJSONMarshalError(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
tok := &model.APIToken{ID: 4242, UserID: 1}
|
||||||
|
unsupportedBody := map[string]any{"unsupported": func() {}}
|
||||||
|
|
||||||
|
// When
|
||||||
|
_, _, err := mcpEndpointTestCtx(t, tok, unsupportedBody)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
require.Error(t, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// An unknown tool name in tools/call must still consume the per-token rate
|
||||||
|
// budget; otherwise a valid PAT can flood /mcp with does-not-exist tools and
|
||||||
|
// bypass the limiter entirely.
|
||||||
|
func TestMCPUnknownToolCountsAgainstRateLimit(t *testing.T) {
|
||||||
|
originalConf := singleton.Conf
|
||||||
|
originalLimiter := mcpRateLimiterShared
|
||||||
|
t.Cleanup(func() {
|
||||||
|
singleton.Conf = originalConf
|
||||||
|
mcpRateLimiterShared = originalLimiter
|
||||||
|
})
|
||||||
|
|
||||||
|
cfg := &model.Config{}
|
||||||
|
cfg.SetMCPEnabled(true)
|
||||||
|
singleton.Conf = &singleton.ConfigClass{Config: cfg}
|
||||||
|
|
||||||
|
mcpRateLimiterShared = newMCPRateLimiter(2, 2)
|
||||||
|
|
||||||
|
tok := &model.APIToken{ID: 4242, UserID: 1}
|
||||||
|
|
||||||
|
body := map[string]any{
|
||||||
|
"jsonrpc": "2.0",
|
||||||
|
"id": 1,
|
||||||
|
"method": "tools/call",
|
||||||
|
"params": map[string]any{"name": "does.not.exist", "arguments": map[string]any{}},
|
||||||
|
}
|
||||||
|
|
||||||
|
var lastStatus int
|
||||||
|
for i := 0; i < 5; i++ {
|
||||||
|
c, w, err := mcpEndpointTestCtx(t, tok, body)
|
||||||
|
require.NoError(t, err)
|
||||||
|
mcpEndpoint(c)
|
||||||
|
lastStatus = w.Code
|
||||||
|
}
|
||||||
|
|
||||||
|
if !mcpRateLimiterSaturated(tok.ID) {
|
||||||
|
t.Fatalf("after 5 unknown-tool calls with a budget of 2, the limiter must be saturated (last status %d)", lastStatus)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func mcpRateLimiterSaturated(tokenID uint64) bool {
|
||||||
|
return !mcpRateLimiterShared.Allow(tokenID)
|
||||||
|
}
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
type linearizationClock struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
times []time.Time
|
||||||
|
latest time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
func newLinearizationClock(oldTime, newTime time.Time) *linearizationClock {
|
||||||
|
return &linearizationClock{
|
||||||
|
times: []time.Time{oldTime, oldTime, newTime},
|
||||||
|
latest: newTime,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (clock *linearizationClock) Now() time.Time {
|
||||||
|
clock.mu.Lock()
|
||||||
|
now := clock.latest
|
||||||
|
if len(clock.times) > 0 {
|
||||||
|
now = clock.times[0]
|
||||||
|
clock.times = clock.times[1:]
|
||||||
|
}
|
||||||
|
clock.mu.Unlock()
|
||||||
|
return now
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLinearizationClockReturnsLatestTimeAfterSequenceIsConsumed(t *testing.T) {
|
||||||
|
oldTime := time.Unix(1_700_000_000, 0)
|
||||||
|
newTime := oldTime.Add(time.Second)
|
||||||
|
clock := newLinearizationClock(oldTime, newTime)
|
||||||
|
|
||||||
|
for range 3 {
|
||||||
|
clock.Now()
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := clock.Now(); !got.Equal(newTime) {
|
||||||
|
t.Fatalf("clock returned %s after sequence was consumed, want %s", got, newTime)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,18 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
// TestMCPRateLimiter_DefaultsAreDoubledFromInitialBaseline pins the
|
||||||
|
// production-side per-token budget. The first iteration shipped 5/s + 60/min,
|
||||||
|
// which gated legitimate LLM bursts more aggressively than the audit /
|
||||||
|
// concurrency story required. Doubling to 10/s + 120/min keeps the bucket
|
||||||
|
// shape (same window, same per-token bookkeeping) so observed behavior
|
||||||
|
// regressions stay attributable to budget rather than algorithm changes.
|
||||||
|
func TestMCPRateLimiter_DefaultsAreDoubledFromInitialBaseline(t *testing.T) {
|
||||||
|
if mcpRateLimiterShared.secLimit != 10 {
|
||||||
|
t.Fatalf("default per-second limit = %d, want 10", mcpRateLimiterShared.secLimit)
|
||||||
|
}
|
||||||
|
if mcpRateLimiterShared.minLimit != 120 {
|
||||||
|
t.Fatalf("default per-minute limit = %d, want 120", mcpRateLimiterShared.minLimit)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,78 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func mcpEndpointRawCtx(t *testing.T, tok *model.APIToken, raw []byte) (*gin.Context, *httptest.ResponseRecorder) {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, engine := gin.CreateTestContext(w)
|
||||||
|
require.NotNil(t, engine)
|
||||||
|
c.Request = httptest.NewRequest("POST", "/mcp", bytes.NewReader(raw))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
c.Set(apiTokenCtxKey, tok)
|
||||||
|
c.Set(model.CtxKeyAPIToken, tok)
|
||||||
|
return c, w
|
||||||
|
}
|
||||||
|
|
||||||
|
func mcpRateLimitTestSetup(t *testing.T, budget int) *model.APIToken {
|
||||||
|
t.Helper()
|
||||||
|
originalConf := singleton.Conf
|
||||||
|
originalLimiter := mcpRateLimiterShared
|
||||||
|
t.Cleanup(func() {
|
||||||
|
singleton.Conf = originalConf
|
||||||
|
mcpRateLimiterShared = originalLimiter
|
||||||
|
})
|
||||||
|
cfg := &model.Config{}
|
||||||
|
cfg.SetMCPEnabled(true)
|
||||||
|
singleton.Conf = &singleton.ConfigClass{Config: cfg}
|
||||||
|
mcpRateLimiterShared = newMCPRateLimiter(budget, budget)
|
||||||
|
return &model.APIToken{ID: 4243, UserID: 1}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A flood of tools/call requests whose params fail to parse must still consume
|
||||||
|
// the per-token budget. Otherwise a valid PAT bypasses the limiter by always
|
||||||
|
// sending malformed arguments.
|
||||||
|
func TestMCPMalformedToolsCallParamsCountsAgainstRateLimit(t *testing.T) {
|
||||||
|
tok := mcpRateLimitTestSetup(t, 2)
|
||||||
|
|
||||||
|
raw := []byte(`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":"not-an-object"}`)
|
||||||
|
|
||||||
|
for i := 0; i < 5; i++ {
|
||||||
|
c, w := mcpEndpointRawCtx(t, tok, raw)
|
||||||
|
mcpEndpoint(c)
|
||||||
|
require.NotEmpty(t, w.Body.Bytes())
|
||||||
|
}
|
||||||
|
|
||||||
|
if mcpRateLimiterShared.Allow(tok.ID) {
|
||||||
|
t.Fatal("malformed tools/call params must still consume the rate budget; limiter not saturated")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A flood of unparseable JSON-RPC envelopes from an authenticated PAT must also
|
||||||
|
// consume the budget.
|
||||||
|
func TestMCPMalformedEnvelopeCountsAgainstRateLimit(t *testing.T) {
|
||||||
|
tok := mcpRateLimitTestSetup(t, 2)
|
||||||
|
|
||||||
|
raw := []byte(`{not valid json`)
|
||||||
|
|
||||||
|
for i := 0; i < 5; i++ {
|
||||||
|
c, w := mcpEndpointRawCtx(t, tok, raw)
|
||||||
|
mcpEndpoint(c)
|
||||||
|
require.NotEmpty(t, w.Body.Bytes())
|
||||||
|
}
|
||||||
|
|
||||||
|
if mcpRateLimiterShared.Allow(tok.ID) {
|
||||||
|
t.Fatal("malformed JSON-RPC envelope must still consume the rate budget; limiter not saturated")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// The per-token limiter map had no eviction: every distinct token ID ever
|
||||||
|
// seen left a permanent entry. A user churning PATs (create/use/delete in a
|
||||||
|
// loop) grows the map without bound. Allow must opportunistically prune
|
||||||
|
// windows idle past the minute bucket so memory stays proportional to the
|
||||||
|
// active token set, not the historical one.
|
||||||
|
func TestMCPRateLimiter_PrunesStaleTokenWindows(t *testing.T) {
|
||||||
|
rl := newMCPRateLimiter(10, 120)
|
||||||
|
|
||||||
|
stale := time.Now().Add(-10 * time.Minute)
|
||||||
|
for i := uint64(1); i <= 500; i++ {
|
||||||
|
rl.mu.Lock()
|
||||||
|
rl.perToken[i] = &tokenWindow{
|
||||||
|
secBucketStart: stale,
|
||||||
|
minBucketStart: stale,
|
||||||
|
}
|
||||||
|
rl.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// A fresh request triggers a prune sweep of idle windows.
|
||||||
|
if !rl.Allow(99999) {
|
||||||
|
t.Fatal("fresh token must be allowed")
|
||||||
|
}
|
||||||
|
|
||||||
|
rl.mu.Lock()
|
||||||
|
size := len(rl.perToken)
|
||||||
|
rl.mu.Unlock()
|
||||||
|
|
||||||
|
// Only the just-active token (99999) should remain; the 500 stale ones
|
||||||
|
// must have been evicted.
|
||||||
|
if size > 1 {
|
||||||
|
t.Fatalf("stale token windows were not pruned: map still holds %d entries", size)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Pruning must NOT evict tokens that are still within their active window,
|
||||||
|
// otherwise an in-flight client loses its accumulated count and effectively
|
||||||
|
// resets its budget.
|
||||||
|
func TestMCPRateLimiter_KeepsActiveTokenWindows(t *testing.T) {
|
||||||
|
rl := newMCPRateLimiter(10, 120)
|
||||||
|
|
||||||
|
if !rl.Allow(1) {
|
||||||
|
t.Fatal("token 1 must be allowed")
|
||||||
|
}
|
||||||
|
if !rl.Allow(2) {
|
||||||
|
t.Fatal("token 2 must be allowed")
|
||||||
|
}
|
||||||
|
|
||||||
|
rl.mu.Lock()
|
||||||
|
size := len(rl.perToken)
|
||||||
|
rl.mu.Unlock()
|
||||||
|
|
||||||
|
if size != 2 {
|
||||||
|
t.Fatalf("active token windows must be retained, got %d entries", size)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,14 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import "encoding/json"
|
||||||
|
|
||||||
|
func marshalMCPToolResult(result any) (string, error) {
|
||||||
|
if result == nil {
|
||||||
|
return "{}", nil
|
||||||
|
}
|
||||||
|
b, err := json.Marshal(result)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
return string(b), nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,47 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
var registerUnmarshalableMCPTool sync.Once
|
||||||
|
|
||||||
|
func TestMCPToolCall_unmarshalableSuccessResultReturnsExplicitToolError(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
registerUnmarshalableMCPTool.Do(func() {
|
||||||
|
registerMCPTool(&mcpTool{
|
||||||
|
Name: "test.unmarshalable-success",
|
||||||
|
RequiredScope: "",
|
||||||
|
Handler: func(*gin.Context, json.RawMessage) (any, error) {
|
||||||
|
return map[string]any{"unsupported": func() {}}, nil
|
||||||
|
},
|
||||||
|
})
|
||||||
|
})
|
||||||
|
tok, plainToken := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
require.NotEmpty(t, plainToken)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "test.unmarshalable-success", Arguments: json.RawMessage("{}")}),
|
||||||
|
})
|
||||||
|
|
||||||
|
// When
|
||||||
|
mcpEndpoint(c)
|
||||||
|
|
||||||
|
// Then
|
||||||
|
_, result := decodeRPC(w)
|
||||||
|
require.NotNil(t, result)
|
||||||
|
require.True(t, result.IsError)
|
||||||
|
require.Nil(t, result.StructuredContent)
|
||||||
|
require.Len(t, result.Content, 1)
|
||||||
|
require.True(t, strings.Contains(result.Content[0].Text, "encode"), result.Content[0].Text)
|
||||||
|
}
|
||||||
@@ -0,0 +1,197 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/modelcontextprotocol/go-sdk/mcp"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// MCP 协议兼容性集成测试:用 modelcontextprotocol/go-sdk 官方 Go MCP client
|
||||||
|
// 对 dashboard /mcp 跑完整 initialize + tools/list + tools/call。
|
||||||
|
// 协议层用官方 SDK 严格编解码 — 任何与 MCP spec 的偏差都会被立即报错。
|
||||||
|
|
||||||
|
type sdkPATRoundTripper struct {
|
||||||
|
base http.RoundTripper
|
||||||
|
token string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (rt *sdkPATRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||||
|
req.Header.Set("Authorization", "Bearer "+rt.token)
|
||||||
|
return rt.base.RoundTrip(req)
|
||||||
|
}
|
||||||
|
|
||||||
|
func sdkTransport(endpoint, token string) *mcp.StreamableClientTransport {
|
||||||
|
return &mcp.StreamableClientTransport{
|
||||||
|
Endpoint: endpoint,
|
||||||
|
HTTPClient: &http.Client{
|
||||||
|
Transport: &sdkPATRoundTripper{base: http.DefaultTransport, token: token},
|
||||||
|
Timeout: 5 * time.Second,
|
||||||
|
},
|
||||||
|
// /mcp 当前只实现 POST 半边 Streamable HTTP;GET SSE 通道未实现也不计划
|
||||||
|
// 短期内上线(不需要 server→client 主动推送)。SDK 默认会试图发 GET,
|
||||||
|
// 关掉 standalone SSE 即可严格互通。
|
||||||
|
DisableStandaloneSSE: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupSDKCompat(t *testing.T) (string, string, func()) {
|
||||||
|
t.Helper()
|
||||||
|
cleanupBase, uid := setupMCPTest(t)
|
||||||
|
|
||||||
|
srv, _ := singleton.ServerShared.Get(7)
|
||||||
|
srv.SetTaskStream(&e2eStream{dispatch: agentSim})
|
||||||
|
|
||||||
|
_, plain := mkToken(t, uid, []string{
|
||||||
|
model.ScopeInventoryRead,
|
||||||
|
model.ScopeInventoryDelete,
|
||||||
|
model.ScopeServerRead,
|
||||||
|
model.ScopeServerWrite,
|
||||||
|
model.ScopeServerDelete,
|
||||||
|
model.ScopeServerExec,
|
||||||
|
}, nil)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
r := gin.New()
|
||||||
|
r.POST("/mcp", apiTokenAuthMiddleware(), mcpEndpoint)
|
||||||
|
ts := httptest.NewServer(r)
|
||||||
|
return ts.URL + "/mcp", plain, func() {
|
||||||
|
ts.Close()
|
||||||
|
cleanupBase()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSDKClient_InitializeHandshake(t *testing.T) {
|
||||||
|
endpoint, token, cleanup := setupSDKCompat(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
client := mcp.NewClient(&mcp.Implementation{Name: "nezha-it", Version: "v0"}, nil)
|
||||||
|
session, err := client.Connect(ctx, sdkTransport(endpoint, token), nil)
|
||||||
|
require.NoError(t, err, "official Go SDK must initialize against /mcp")
|
||||||
|
defer session.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSDKClient_ToolsList(t *testing.T) {
|
||||||
|
endpoint, token, cleanup := setupSDKCompat(t)
|
||||||
|
defer cleanup()
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
client := mcp.NewClient(&mcp.Implementation{Name: "nezha-it", Version: "v0"}, nil)
|
||||||
|
session, err := client.Connect(ctx, sdkTransport(endpoint, token), nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer session.Close()
|
||||||
|
|
||||||
|
lst, err := session.ListTools(ctx, nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
names := make(map[string]bool, len(lst.Tools))
|
||||||
|
for _, tl := range lst.Tools {
|
||||||
|
names[tl.Name] = true
|
||||||
|
}
|
||||||
|
for _, must := range []string{
|
||||||
|
"meta.whoami",
|
||||||
|
"server.list", "server.get", "server.exec",
|
||||||
|
"fs.list", "fs.read", "fs.write", "fs.delete",
|
||||||
|
"fs.download_url", "fs.upload_url",
|
||||||
|
} {
|
||||||
|
require.Truef(t, names[must], "tools/list missing %q", must)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSDKClient_Whoami(t *testing.T) {
|
||||||
|
endpoint, token, cleanup := setupSDKCompat(t)
|
||||||
|
defer cleanup()
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
client := mcp.NewClient(&mcp.Implementation{Name: "nezha-it", Version: "v0"}, nil)
|
||||||
|
session, err := client.Connect(ctx, sdkTransport(endpoint, token), nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer session.Close()
|
||||||
|
|
||||||
|
res, err := session.CallTool(ctx, &mcp.CallToolParams{
|
||||||
|
Name: "meta.whoami",
|
||||||
|
Arguments: map[string]any{},
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.False(t, res.IsError)
|
||||||
|
tc, ok := res.Content[0].(*mcp.TextContent)
|
||||||
|
require.True(t, ok)
|
||||||
|
var payload map[string]any
|
||||||
|
require.NoError(t, json.Unmarshal([]byte(tc.Text), &payload))
|
||||||
|
require.NotZero(t, payload["user_id"])
|
||||||
|
require.NotEmpty(t, payload["scopes"])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSDKClient_ServerExec(t *testing.T) {
|
||||||
|
endpoint, token, cleanup := setupSDKCompat(t)
|
||||||
|
defer cleanup()
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
client := mcp.NewClient(&mcp.Implementation{Name: "nezha-it", Version: "v0"}, nil)
|
||||||
|
session, err := client.Connect(ctx, sdkTransport(endpoint, token), nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer session.Close()
|
||||||
|
|
||||||
|
res, err := session.CallTool(ctx, &mcp.CallToolParams{
|
||||||
|
Name: "server.exec",
|
||||||
|
Arguments: map[string]any{
|
||||||
|
"server_id": 7,
|
||||||
|
"cmd": "echo",
|
||||||
|
},
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.False(t, res.IsError, "exec failed: %v", res.Content)
|
||||||
|
tc := res.Content[0].(*mcp.TextContent)
|
||||||
|
require.Contains(t, tc.Text, "simulated")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSDKClient_FSLifecycle(t *testing.T) {
|
||||||
|
endpoint, token, cleanup := setupSDKCompat(t)
|
||||||
|
defer cleanup()
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
client := mcp.NewClient(&mcp.Implementation{Name: "nezha-it", Version: "v0"}, nil)
|
||||||
|
session, err := client.Connect(ctx, sdkTransport(endpoint, token), nil)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer session.Close()
|
||||||
|
|
||||||
|
path := t.TempDir() + "/sdk.txt"
|
||||||
|
for _, step := range []struct {
|
||||||
|
name string
|
||||||
|
args map[string]any
|
||||||
|
}{
|
||||||
|
{"fs.write", map[string]any{"server_id": 7, "path": path, "content": "via-sdk", "encoding": "utf8"}},
|
||||||
|
{"fs.read", map[string]any{"server_id": 7, "path": path}},
|
||||||
|
{"fs.delete", map[string]any{"server_id": 7, "path": path}},
|
||||||
|
} {
|
||||||
|
res, err := session.CallTool(ctx, &mcp.CallToolParams{Name: step.name, Arguments: step.args})
|
||||||
|
require.NoError(t, err, step.name)
|
||||||
|
require.False(t, res.IsError, "%s failed: %v", step.name, res.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSDKClient_BadPAT(t *testing.T) {
|
||||||
|
endpoint, _, cleanup := setupSDKCompat(t)
|
||||||
|
defer cleanup()
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
client := mcp.NewClient(&mcp.Implementation{Name: "nezha-it", Version: "v0"}, nil)
|
||||||
|
_, err := client.Connect(ctx, sdkTransport(endpoint, "nzp_invalid"), nil)
|
||||||
|
require.Error(t, err)
|
||||||
|
}
|
||||||
@@ -0,0 +1,306 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/i18n"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func setupMCPTest(t *testing.T) (func(), uint64) {
|
||||||
|
t.Helper()
|
||||||
|
originalDB := singleton.DB
|
||||||
|
originalServer := singleton.ServerShared
|
||||||
|
originalConf := singleton.Conf
|
||||||
|
originalAuditSync := mcpAuditSync
|
||||||
|
originalLimiter := mcpRateLimiterShared
|
||||||
|
originalLocalizer := singleton.Localizer
|
||||||
|
originalPATRegistry := patConnectionRegistryShared
|
||||||
|
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
|
||||||
|
mcpAuditSync = true
|
||||||
|
mcpRateLimiterShared = newMCPRateLimiter(1000, 10000)
|
||||||
|
// Fresh per test: the DB resets token IDs to 1 each run, so a stale
|
||||||
|
// revoke tombstone from a prior test would otherwise cancel a reused id.
|
||||||
|
patConnectionRegistryShared = newPATConnectionRegistry()
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
sqlDB, err := db.DB()
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, db.AutoMigrate(&model.User{}, &model.APIToken{}, &model.MCPAuditLog{}, &model.Server{}, &model.WAF{}))
|
||||||
|
singleton.DB = db
|
||||||
|
singleton.Conf = &singleton.ConfigClass{Config: &model.Config{JWTTimeout: 1}}
|
||||||
|
singleton.Conf.SetMCPEnabled(true)
|
||||||
|
|
||||||
|
user := model.User{Common: model.Common{ID: 100}, Username: "alice", Role: model.RoleMember}
|
||||||
|
require.NoError(t, db.Create(&user).Error)
|
||||||
|
|
||||||
|
sc := singleton.NewEmptyServerClassForTest()
|
||||||
|
srv := &model.Server{}
|
||||||
|
srv.ID = 7
|
||||||
|
srv.Name = "alpha"
|
||||||
|
srv.SetUserID(100)
|
||||||
|
sc.InsertForTest(srv)
|
||||||
|
singleton.ServerShared = sc
|
||||||
|
|
||||||
|
cleanup := func() {
|
||||||
|
_ = sqlDB.Close()
|
||||||
|
singleton.DB = originalDB
|
||||||
|
singleton.ServerShared = originalServer
|
||||||
|
singleton.Conf = originalConf
|
||||||
|
singleton.Localizer = originalLocalizer
|
||||||
|
mcpAuditSync = originalAuditSync
|
||||||
|
mcpRateLimiterShared = originalLimiter
|
||||||
|
patConnectionRegistryShared = originalPATRegistry
|
||||||
|
}
|
||||||
|
return cleanup, user.ID
|
||||||
|
}
|
||||||
|
|
||||||
|
func mkToken(t *testing.T, uid uint64, scopes []string, serverIDs []uint64) (*model.APIToken, string) {
|
||||||
|
t.Helper()
|
||||||
|
plain := "nzp_" + strings.Repeat("a", 32) + "_" + ctoa(uid)
|
||||||
|
tok := model.APIToken{UserID: uid, Name: "t", TokenHash: model.HashAPIToken(plain)}
|
||||||
|
tok.SetScopes(scopes)
|
||||||
|
if len(serverIDs) > 0 {
|
||||||
|
tok.SetServerIDs(serverIDs)
|
||||||
|
}
|
||||||
|
require.NoError(t, singleton.DB.Create(&tok).Error)
|
||||||
|
return &tok, plain
|
||||||
|
}
|
||||||
|
|
||||||
|
func mcpCallCtx(t *testing.T, tok *model.APIToken, uid uint64, body any) (*gin.Context, *httptest.ResponseRecorder) {
|
||||||
|
t.Helper()
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
b, _ := json.Marshal(body)
|
||||||
|
c.Request = httptest.NewRequest("POST", "/mcp", bytes.NewReader(b))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: uid}, Role: model.RoleMember})
|
||||||
|
c.Set(apiTokenCtxKey, tok)
|
||||||
|
c.Set(model.CtxKeyAPIToken, tok)
|
||||||
|
return c, w
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeRPC(w *httptest.ResponseRecorder) (jsonRPCResponse, *mcpToolCallResult) {
|
||||||
|
var env jsonRPCResponse
|
||||||
|
_ = json.Unmarshal(w.Body.Bytes(), &env)
|
||||||
|
if env.Result == nil {
|
||||||
|
return env, nil
|
||||||
|
}
|
||||||
|
rb, _ := json.Marshal(env.Result)
|
||||||
|
var tcr mcpToolCallResult
|
||||||
|
_ = json.Unmarshal(rb, &tcr)
|
||||||
|
return env, &tcr
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_RejectsMissingToken(t *testing.T) {
|
||||||
|
cleanup, _ := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
body, _ := json.Marshal(jsonRPCRequest{JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "initialize"})
|
||||||
|
c.Request = httptest.NewRequest("POST", "/mcp", bytes.NewReader(body))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
mcpEndpoint(c)
|
||||||
|
var env jsonRPCResponse
|
||||||
|
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &env))
|
||||||
|
require.NotNil(t, env.Error)
|
||||||
|
require.Equal(t, rpcErrUnauthorized, env.Error.Code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_Initialize_ReturnsServerInfo(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "initialize"})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
env, _ := decodeRPC(w)
|
||||||
|
require.Nil(t, env.Error)
|
||||||
|
rb, _ := json.Marshal(env.Result)
|
||||||
|
require.Contains(t, string(rb), "nezha-mcp")
|
||||||
|
require.Contains(t, string(rb), "protocolVersion")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_ToolsList_IncludesRegisteredTools(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/list"})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
env, _ := decodeRPC(w)
|
||||||
|
require.Nil(t, env.Error)
|
||||||
|
rb, _ := json.Marshal(env.Result)
|
||||||
|
for _, name := range []string{"meta.whoami", "server.list", "server.exec", "fs.list", "fs.read", "fs.write", "fs.delete", "fs.download_url", "fs.upload_url"} {
|
||||||
|
require.Contains(t, string(rb), name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_Whoami_HappyPath(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead, model.ScopeServerRead}, []uint64{7, 8})
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "meta.whoami", Arguments: json.RawMessage("{}")}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.NotNil(t, tcr)
|
||||||
|
require.False(t, tcr.IsError, "got error content: %v", tcr.Content)
|
||||||
|
scb, _ := json.Marshal(tcr.StructuredContent)
|
||||||
|
require.Contains(t, string(scb), "user_id")
|
||||||
|
require.Contains(t, string(scb), "scopes")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_ScopeDenied(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "server.exec", Arguments: jsonRaw(map[string]any{"server_id": 7, "cmd": "echo"})}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.NotNil(t, tcr)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "missing required scope")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_PermissionDenied_WhenWrongUserOwnsServer(t *testing.T) {
|
||||||
|
cleanup, _ := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
require.NoError(t, singleton.DB.Create(&model.User{Common: model.Common{ID: 200}, Username: "bob", Role: model.RoleMember}).Error)
|
||||||
|
tok, _ := mkToken(t, 200, []string{model.ScopeServerRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, 200, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "fs.list", Arguments: jsonRaw(map[string]any{"server_id": 7, "path": "/tmp"})}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.NotNil(t, tcr)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "permission denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_ServerWhitelist_DenyOutsideList(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, []uint64{99})
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "fs.list", Arguments: jsonRaw(map[string]any{"server_id": 7, "path": "/tmp"})}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.NotNil(t, tcr)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_UnknownTool(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "does.not.exist", Arguments: json.RawMessage("{}")}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
env, _ := decodeRPC(w)
|
||||||
|
require.NotNil(t, env.Error)
|
||||||
|
require.Equal(t, rpcErrMethodNotFound, env.Error.Code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_InvalidJSONEnvelope(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("POST", "/mcp", bytes.NewReader([]byte("garbage")))
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: uid}, Role: model.RoleMember})
|
||||||
|
c.Set(apiTokenCtxKey, tok)
|
||||||
|
mcpEndpoint(c)
|
||||||
|
var env jsonRPCResponse
|
||||||
|
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &env))
|
||||||
|
require.NotNil(t, env.Error)
|
||||||
|
require.Equal(t, rpcErrParse, env.Error.Code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_AuditRowIsWritten(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
c, _ := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "meta.whoami", Arguments: json.RawMessage("{}")}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
|
||||||
|
require.Eventually(t, func() bool {
|
||||||
|
var cnt int64
|
||||||
|
_ = singleton.DB.Model(&model.MCPAuditLog{}).Where("token_id = ?", tok.ID).Count(&cnt).Error
|
||||||
|
return cnt == 1
|
||||||
|
}, 2*time.Second, 20*time.Millisecond, "audit row never appeared")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCP_RateLimit(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
original := mcpRateLimiterShared
|
||||||
|
mcpRateLimiterShared = newMCPRateLimiter(2, 100)
|
||||||
|
defer func() { mcpRateLimiterShared = original }()
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
|
||||||
|
for i := 0; i < 2; i++ {
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "meta.whoami", Arguments: json.RawMessage("{}")}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.False(t, tcr.IsError)
|
||||||
|
}
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "meta.whoami", Arguments: json.RawMessage("{}")}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "rate limit")
|
||||||
|
}
|
||||||
|
|
||||||
|
func jsonObj(t *testing.T, v any) json.RawMessage {
|
||||||
|
t.Helper()
|
||||||
|
b, err := json.Marshal(v)
|
||||||
|
require.NoError(t, err)
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
func jsonRaw(v map[string]any) json.RawMessage {
|
||||||
|
b, _ := json.Marshal(v)
|
||||||
|
return b
|
||||||
|
}
|
||||||
|
|
||||||
|
func ctoa(v uint64) string {
|
||||||
|
b, _ := json.Marshal(v)
|
||||||
|
return string(b)
|
||||||
|
}
|
||||||
@@ -0,0 +1,144 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
// server.exec — 非交互一次性命令。
|
||||||
|
// 协议约束(agent 端强制):
|
||||||
|
// - 不开 pty
|
||||||
|
// - 默认 30s 超时,硬上限 300s
|
||||||
|
// - stdout/stderr 各自最多 64KB(默认),硬上限 1MB
|
||||||
|
// - 受 agent 配置 DisableCommandExecute 影响
|
||||||
|
// - 命令返回或超时时,agent 会回收整个进程组/JobObject:'cmd &'、nohup、
|
||||||
|
// disown 这类普通后台进程都会被一并杀掉。要留长驻进程必须脱离会话
|
||||||
|
// (setsid / screen -dmS / tmux new -d / systemd-run;Windows 需 breakaway)。
|
||||||
|
//
|
||||||
|
// LLM 要用 shell 特性(管道、重定向)必须显式传 cmd="sh" args=["-c","..."],
|
||||||
|
// 这样审计日志能完整记录被执行的指令。
|
||||||
|
const mcpExecMaxTimeoutSec uint32 = 300
|
||||||
|
|
||||||
|
type execArgs struct {
|
||||||
|
ServerID uint64 `json:"server_id"`
|
||||||
|
Cmd string `json:"cmd"`
|
||||||
|
Args []string `json:"args,omitempty"`
|
||||||
|
Cwd string `json:"cwd,omitempty"`
|
||||||
|
Env map[string]string `json:"env,omitempty"`
|
||||||
|
TimeoutSeconds uint32 `json:"timeout_seconds,omitempty"`
|
||||||
|
Stdin string `json:"stdin,omitempty"`
|
||||||
|
MaxOutputBytes uint32 `json:"max_output_bytes,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
registerMCPTool(&mcpTool{
|
||||||
|
Name: "server.exec",
|
||||||
|
Description: "Run a non-interactive command on the target server and return stdout/stderr/exit_code. No pty. Use cmd='sh' args=['-c', '...'] for shell features. The entire process tree is killed when the command returns or times out, so plain background jobs ('cmd &', nohup, disown) do NOT survive; to leave a process running after the call, fully detach it from the session (e.g. setsid, 'screen -dmS', 'tmux new -d', systemd-run; on Windows it must break away from the job object).",
|
||||||
|
InputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"server_id": map[string]any{"type": "integer"},
|
||||||
|
"cmd": map[string]any{"type": "string"},
|
||||||
|
"args": map[string]any{"type": "array", "items": map[string]any{"type": "string"}},
|
||||||
|
"cwd": map[string]any{"type": "string"},
|
||||||
|
"env": map[string]any{"type": "object"},
|
||||||
|
"timeout_seconds": map[string]any{"type": "integer", "minimum": 1, "maximum": 300},
|
||||||
|
"stdin": map[string]any{"type": "string"},
|
||||||
|
"max_output_bytes": map[string]any{"type": "integer"},
|
||||||
|
},
|
||||||
|
"required": []string{"server_id", "cmd"},
|
||||||
|
},
|
||||||
|
OutputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"exit_code": map[string]any{"type": "integer"},
|
||||||
|
"stdout": map[string]any{"type": "string"},
|
||||||
|
"stderr": map[string]any{"type": "string"},
|
||||||
|
"duration_ms": map[string]any{"type": "integer"},
|
||||||
|
"stdout_truncated": map[string]any{"type": "boolean"},
|
||||||
|
"stderr_truncated": map[string]any{"type": "boolean"},
|
||||||
|
"timed_out": map[string]any{"type": "boolean"},
|
||||||
|
},
|
||||||
|
"required": []string{"exit_code", "stdout", "stderr", "duration_ms"},
|
||||||
|
},
|
||||||
|
RequiredScope: model.ScopeServerExec,
|
||||||
|
Handler: handleServerExec,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleServerExec(c *gin.Context, raw json.RawMessage) (any, error) {
|
||||||
|
var args execArgs
|
||||||
|
if err := decodeToolArgs(raw, &args); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if args.TimeoutSeconds > mcpExecMaxTimeoutSec {
|
||||||
|
return nil, errMCPInvalidArgs("timeout_seconds out of range; must be 1..300")
|
||||||
|
}
|
||||||
|
srv, err := requireServerAccess(c, args.ServerID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if err := requireAgentSupportsMCP(srv); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if args.Cmd == "" {
|
||||||
|
return nil, errMCPInvalidArgs("cmd required")
|
||||||
|
}
|
||||||
|
|
||||||
|
req := model.ExecRequest{
|
||||||
|
Cmd: args.Cmd,
|
||||||
|
Args: args.Args,
|
||||||
|
Cwd: args.Cwd,
|
||||||
|
Env: args.Env,
|
||||||
|
TimeoutSeconds: args.TimeoutSeconds,
|
||||||
|
Stdin: args.Stdin,
|
||||||
|
MaxOutputBytes: args.MaxOutputBytes,
|
||||||
|
}
|
||||||
|
|
||||||
|
timeout := callAgentTimeout(args.TimeoutSeconds, 30)
|
||||||
|
raw2, err := rpc.CallAgent(c.Request.Context(), args.ServerID, model.TaskTypeExec, req, timeout)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var res model.ExecResult
|
||||||
|
if err := json.Unmarshal(raw2, &res); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
// ExecResult.Error means the agent refused / failed to run the command
|
||||||
|
// (disabled, empty cmd, Start/process-group failure). Surface it like fs.*
|
||||||
|
// handlers do, so MCP isError=true and audit outcome=agent_error. Non-zero
|
||||||
|
// ExitCode alone is a normal command outcome, not a tool error.
|
||||||
|
if res.Error != "" {
|
||||||
|
return nil, &execToolError{mcpError: mcpError{Code: model.MCPOutcomeAgentError, Msg: res.Error}, result: res}
|
||||||
|
}
|
||||||
|
return res, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type execToolError struct {
|
||||||
|
mcpError
|
||||||
|
result model.ExecResult
|
||||||
|
}
|
||||||
|
|
||||||
|
func (err *execToolError) StructuredResult() any { return err.result }
|
||||||
|
|
||||||
|
func (err *execToolError) Error() string {
|
||||||
|
return fmt.Sprintf("%s: %s", err.Code, err.Msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// callAgentTimeout 给 dashboard 侧 CallAgent 计算等待上限。
|
||||||
|
// 在用户请求的 timeout 基础上加 5s buffer,让 agent 端的 hard timeout 先触发,
|
||||||
|
// 这样 dashboard 收到的总是结构化结果(包含 timed_out=true),
|
||||||
|
// 而不是 ErrAgentTimeout。
|
||||||
|
func callAgentTimeout(reqTimeoutSec uint32, defaultSec uint32) time.Duration {
|
||||||
|
t := reqTimeoutSec
|
||||||
|
if t == 0 {
|
||||||
|
t = defaultSec
|
||||||
|
}
|
||||||
|
return time.Duration(t+5) * time.Second
|
||||||
|
}
|
||||||
@@ -0,0 +1,97 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"google.golang.org/grpc/metadata"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
pb "github.com/nezhahq/nezha/proto"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// execErrorStream replies every Task with a TaskResult whose Successful=true
|
||||||
|
// but whose Data carries a model.ExecResult{Error: ...}. This is exactly what
|
||||||
|
// the real agent does for "agent disabled command execution" / "cmd required"
|
||||||
|
// / pre-start failures.
|
||||||
|
type execErrorStream struct {
|
||||||
|
errMsg string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *execErrorStream) Send(t *pb.Task) error {
|
||||||
|
go func(taskID uint64) {
|
||||||
|
b, _ := json.Marshal(model.ExecResult{ExitCode: -1, Error: s.errMsg})
|
||||||
|
rpc.DeliverMCPResultForTest(&pb.TaskResult{
|
||||||
|
Id: taskID,
|
||||||
|
Type: model.TaskTypeExec,
|
||||||
|
Successful: true,
|
||||||
|
Data: string(b),
|
||||||
|
})
|
||||||
|
}(t.GetId())
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *execErrorStream) Recv() (*pb.TaskResult, error) { return nil, context.Canceled }
|
||||||
|
func (s *execErrorStream) SetHeader(metadata.MD) error { return nil }
|
||||||
|
func (s *execErrorStream) SendHeader(metadata.MD) error { return nil }
|
||||||
|
func (s *execErrorStream) SetTrailer(metadata.MD) {}
|
||||||
|
func (s *execErrorStream) Context() context.Context { return context.Background() }
|
||||||
|
func (s *execErrorStream) SendMsg(any) error { return nil }
|
||||||
|
func (s *execErrorStream) RecvMsg(any) error { return context.Canceled }
|
||||||
|
|
||||||
|
// TestServerExec_AgentReportedErrorBecomesToolError pins the protocol contract
|
||||||
|
// that fs.* tools already obey: when the agent returns a structured result
|
||||||
|
// with a non-empty Error field, MCP tools/call must surface isError=true and
|
||||||
|
// audit must record agent_error — not MCPOutcomeOK with a quietly-failed
|
||||||
|
// structuredContent. The previous handler ignored ExecResult.Error and
|
||||||
|
// returned res, nil, which made the LLM and the audit log both believe the
|
||||||
|
// command succeeded while the agent had actually refused it.
|
||||||
|
func TestServerExec_AgentReportedErrorBecomesToolError(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
srv, _ := singleton.ServerShared.Get(7)
|
||||||
|
srv.SetTaskStream(&execErrorStream{errMsg: "agent disabled command execution"})
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerExec}, nil)
|
||||||
|
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{
|
||||||
|
Name: "server.exec",
|
||||||
|
Arguments: jsonRaw(map[string]any{
|
||||||
|
"server_id": 7,
|
||||||
|
"cmd": "echo",
|
||||||
|
"timeout_seconds": 2,
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.NotNil(t, tcr, "tools/call must return a tool result envelope")
|
||||||
|
require.True(t, tcr.IsError,
|
||||||
|
"agent ExecResult.Error must propagate as MCP tool error; got %+v", tcr)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "agent disabled command execution",
|
||||||
|
"tool error text must surface the agent-reported error message")
|
||||||
|
var result model.ExecResult
|
||||||
|
structuredJSON, err := json.Marshal(tcr.StructuredContent)
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, json.Unmarshal(structuredJSON, &result))
|
||||||
|
require.Equal(t, -1, result.ExitCode)
|
||||||
|
require.Equal(t, "agent disabled command execution", result.Error)
|
||||||
|
|
||||||
|
require.Eventually(t, func() bool {
|
||||||
|
var got model.MCPAuditLog
|
||||||
|
err := singleton.DB.Where("token_id = ?", tok.ID).First(&got).Error
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return got.Outcome == model.MCPOutcomeAgentError
|
||||||
|
}, 2*time.Second, 20*time.Millisecond,
|
||||||
|
"audit row must record agent_error, not ok, when ExecResult.Error is set")
|
||||||
|
}
|
||||||
@@ -0,0 +1,63 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestServerExec_RejectsOutOfRangeTimeoutSeconds(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerExec}, nil)
|
||||||
|
|
||||||
|
// timeout_seconds is documented as 1..300 in the tool schema; sending
|
||||||
|
// 1_000_000 would otherwise let the dashboard wait ~1e6s on rpc.CallAgent
|
||||||
|
// when the agent is unreachable / old, turning one MCP call into a
|
||||||
|
// long-lived goroutine + connection occupation. Handler must reject
|
||||||
|
// before touching the RPC layer.
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{
|
||||||
|
Name: "server.exec",
|
||||||
|
Arguments: jsonRaw(map[string]any{
|
||||||
|
"server_id": 7,
|
||||||
|
"cmd": "echo",
|
||||||
|
"timeout_seconds": 1_000_000,
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.NotNil(t, tcr)
|
||||||
|
require.True(t, tcr.IsError, "expected isError=true for out-of-range timeout, got %+v", tcr)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "timeout_seconds")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServerExec_RejectsZeroLikeNegativeTimeoutBoundary(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerExec}, nil)
|
||||||
|
|
||||||
|
// 301s sits one above the documented maximum. The previous handler
|
||||||
|
// happily forwarded it as-is and added +5s to the dashboard-side wait,
|
||||||
|
// so any client could ignore the schema bound. Pin the rejection.
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{
|
||||||
|
Name: "server.exec",
|
||||||
|
Arguments: jsonRaw(map[string]any{
|
||||||
|
"server_id": 7,
|
||||||
|
"cmd": "echo",
|
||||||
|
"timeout_seconds": 301,
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "timeout_seconds")
|
||||||
|
}
|
||||||
@@ -0,0 +1,311 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
const fsCallTimeout = 30 * time.Second
|
||||||
|
|
||||||
|
func fsEntrySchema() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"name": map[string]any{"type": "string"},
|
||||||
|
"type": map[string]any{"type": "string"},
|
||||||
|
"size": map[string]any{"type": "integer"},
|
||||||
|
"mode": map[string]any{"type": "string"},
|
||||||
|
"mtime": map[string]any{"type": "integer"},
|
||||||
|
"is_symlink": map[string]any{"type": "boolean"},
|
||||||
|
"link_target": map[string]any{"type": "string"},
|
||||||
|
},
|
||||||
|
"required": []string{"name", "type", "size", "mode", "mtime"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
registerMCPTool(&mcpTool{
|
||||||
|
Name: "fs.list",
|
||||||
|
Description: "List entries of a directory on the target server.",
|
||||||
|
InputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"server_id": map[string]any{"type": "integer"},
|
||||||
|
"path": map[string]any{"type": "string", "description": "Absolute path."},
|
||||||
|
"show_hidden": map[string]any{"type": "boolean"},
|
||||||
|
},
|
||||||
|
"required": []string{"server_id", "path"},
|
||||||
|
},
|
||||||
|
OutputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"entries": map[string]any{"type": "array", "items": fsEntrySchema()},
|
||||||
|
"truncated": map[string]any{"type": "boolean"},
|
||||||
|
"total": map[string]any{"type": "integer"},
|
||||||
|
},
|
||||||
|
"required": []string{"entries"},
|
||||||
|
},
|
||||||
|
RequiredScope: model.ScopeServerRead,
|
||||||
|
Handler: handleFsList,
|
||||||
|
})
|
||||||
|
|
||||||
|
registerMCPTool(&mcpTool{
|
||||||
|
Name: "fs.read",
|
||||||
|
Description: "Read a file. Default max 1MB; use offset/length for larger ranges, or fs.download_url for streaming up to 100MiB out-of-band.",
|
||||||
|
InputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"server_id": map[string]any{"type": "integer"},
|
||||||
|
"path": map[string]any{"type": "string"},
|
||||||
|
"offset": map[string]any{"type": "integer", "minimum": 0},
|
||||||
|
"length": map[string]any{"type": "integer", "minimum": 1},
|
||||||
|
"encoding": map[string]any{"type": "string", "enum": []string{"utf8", "base64"}},
|
||||||
|
},
|
||||||
|
"required": []string{"server_id", "path"},
|
||||||
|
},
|
||||||
|
OutputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"content": map[string]any{"type": "string"},
|
||||||
|
"encoding": map[string]any{"type": "string"},
|
||||||
|
"size": map[string]any{"type": "integer"},
|
||||||
|
"sha256": map[string]any{"type": "string"},
|
||||||
|
"truncated": map[string]any{"type": "boolean"},
|
||||||
|
},
|
||||||
|
"required": []string{"content", "encoding", "size"},
|
||||||
|
},
|
||||||
|
RequiredScope: model.ScopeServerRead,
|
||||||
|
Handler: handleFsRead,
|
||||||
|
})
|
||||||
|
|
||||||
|
registerMCPTool(&mcpTool{
|
||||||
|
Name: "fs.write",
|
||||||
|
Description: "Atomic write to a file. Supports utf8 / base64 content, optional sha256 optimistic lock, create_dirs.",
|
||||||
|
InputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"server_id": map[string]any{"type": "integer"},
|
||||||
|
"path": map[string]any{"type": "string"},
|
||||||
|
"content": map[string]any{"type": "string"},
|
||||||
|
"encoding": map[string]any{"type": "string", "enum": []string{"utf8", "base64"}},
|
||||||
|
"mode": map[string]any{"type": "string", "description": "Octal mode like '0644'."},
|
||||||
|
"if_match_sha256": map[string]any{"type": "string"},
|
||||||
|
"create_dirs": map[string]any{"type": "boolean"},
|
||||||
|
},
|
||||||
|
"required": []string{"server_id", "path", "content"},
|
||||||
|
},
|
||||||
|
OutputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"size": map[string]any{"type": "integer"},
|
||||||
|
"sha256": map[string]any{"type": "string"},
|
||||||
|
},
|
||||||
|
"required": []string{"size", "sha256"},
|
||||||
|
},
|
||||||
|
RequiredScope: model.ScopeServerWrite,
|
||||||
|
Handler: handleFsWrite,
|
||||||
|
})
|
||||||
|
|
||||||
|
registerMCPTool(&mcpTool{
|
||||||
|
Name: "fs.delete",
|
||||||
|
Description: "Delete a file or directory. Pass recursive=true for non-empty directories.",
|
||||||
|
InputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"server_id": map[string]any{"type": "integer"},
|
||||||
|
"path": map[string]any{"type": "string"},
|
||||||
|
"recursive": map[string]any{"type": "boolean"},
|
||||||
|
},
|
||||||
|
"required": []string{"server_id", "path"},
|
||||||
|
},
|
||||||
|
OutputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"deleted_count": map[string]any{"type": "integer"},
|
||||||
|
},
|
||||||
|
"required": []string{"deleted_count"},
|
||||||
|
},
|
||||||
|
RequiredScope: model.ScopeServerDelete,
|
||||||
|
Handler: handleFsDelete,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
type fsListArgs struct {
|
||||||
|
ServerID uint64 `json:"server_id"`
|
||||||
|
Path string `json:"path"`
|
||||||
|
ShowHidden bool `json:"show_hidden,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleFsList(c *gin.Context, raw json.RawMessage) (any, error) {
|
||||||
|
var args fsListArgs
|
||||||
|
if err := decodeToolArgs(raw, &args); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
srv, err := requireServerAccess(c, args.ServerID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if err := requireAgentSupportsMCP(srv); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if args.Path == "" {
|
||||||
|
return nil, errMCPInvalidArgs("path required")
|
||||||
|
}
|
||||||
|
out, err := rpc.CallAgent(c.Request.Context(), args.ServerID, model.TaskTypeFsList,
|
||||||
|
model.FsListRequest{Path: args.Path, ShowHidden: args.ShowHidden}, fsCallTimeout)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var res model.FsListResult
|
||||||
|
if err := json.Unmarshal(out, &res); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if res.Error != "" {
|
||||||
|
return nil, errors.New(res.Error)
|
||||||
|
}
|
||||||
|
return res, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type fsReadArgs struct {
|
||||||
|
ServerID uint64 `json:"server_id"`
|
||||||
|
Path string `json:"path"`
|
||||||
|
Offset int64 `json:"offset,omitempty"`
|
||||||
|
Length int64 `json:"length,omitempty"`
|
||||||
|
Encoding string `json:"encoding,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleFsRead(c *gin.Context, raw json.RawMessage) (any, error) {
|
||||||
|
var args fsReadArgs
|
||||||
|
if err := decodeToolArgs(raw, &args); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
srv, err := requireServerAccess(c, args.ServerID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if err := requireAgentSupportsMCP(srv); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if args.Path == "" {
|
||||||
|
return nil, errMCPInvalidArgs("path required")
|
||||||
|
}
|
||||||
|
if args.Offset < 0 {
|
||||||
|
return nil, errMCPInvalidArgs("offset must be >= 0")
|
||||||
|
}
|
||||||
|
if args.Length < 0 {
|
||||||
|
return nil, errMCPInvalidArgs("length must be >= 0")
|
||||||
|
}
|
||||||
|
out, err := rpc.CallAgent(c.Request.Context(), args.ServerID, model.TaskTypeFsRead,
|
||||||
|
model.FsReadRequest{
|
||||||
|
Path: args.Path,
|
||||||
|
Offset: args.Offset,
|
||||||
|
Length: args.Length,
|
||||||
|
Encoding: args.Encoding,
|
||||||
|
}, fsCallTimeout)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var res model.FsReadResult
|
||||||
|
if err := json.Unmarshal(out, &res); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if res.Error != "" {
|
||||||
|
return nil, errors.New(res.Error)
|
||||||
|
}
|
||||||
|
return res, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type fsWriteArgs struct {
|
||||||
|
ServerID uint64 `json:"server_id"`
|
||||||
|
Path string `json:"path"`
|
||||||
|
Content string `json:"content"`
|
||||||
|
Encoding string `json:"encoding,omitempty"`
|
||||||
|
Mode string `json:"mode,omitempty"`
|
||||||
|
IfMatchSHA256 string `json:"if_match_sha256,omitempty"`
|
||||||
|
CreateDirs bool `json:"create_dirs,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleFsWrite(c *gin.Context, raw json.RawMessage) (any, error) {
|
||||||
|
var args fsWriteArgs
|
||||||
|
if err := decodeToolArgs(raw, &args); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
srv, err := requireServerAccess(c, args.ServerID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if err := requireAgentSupportsMCP(srv); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if args.Path == "" {
|
||||||
|
return nil, errMCPInvalidArgs("path required")
|
||||||
|
}
|
||||||
|
if args.IfMatchSHA256 != "" {
|
||||||
|
if _, decErr := hex.DecodeString(args.IfMatchSHA256); decErr != nil || len(args.IfMatchSHA256) != 64 {
|
||||||
|
return nil, errMCPInvalidArgs("if_match_sha256 must be 64 hex chars")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
out, err := rpc.CallAgent(c.Request.Context(), args.ServerID, model.TaskTypeFsWrite,
|
||||||
|
model.FsWriteRequest{
|
||||||
|
Path: args.Path,
|
||||||
|
Content: args.Content,
|
||||||
|
Encoding: args.Encoding,
|
||||||
|
Mode: args.Mode,
|
||||||
|
IfMatchSHA256: args.IfMatchSHA256,
|
||||||
|
CreateDirs: args.CreateDirs,
|
||||||
|
}, fsCallTimeout)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var res model.FsWriteResult
|
||||||
|
if err := json.Unmarshal(out, &res); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if res.Error != "" {
|
||||||
|
return nil, errors.New(res.Error)
|
||||||
|
}
|
||||||
|
return res, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type fsDeleteArgs struct {
|
||||||
|
ServerID uint64 `json:"server_id"`
|
||||||
|
Path string `json:"path"`
|
||||||
|
Recursive bool `json:"recursive,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleFsDelete(c *gin.Context, raw json.RawMessage) (any, error) {
|
||||||
|
var args fsDeleteArgs
|
||||||
|
if err := decodeToolArgs(raw, &args); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
srv, err := requireServerAccess(c, args.ServerID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if err := requireAgentSupportsMCP(srv); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if args.Path == "" {
|
||||||
|
return nil, errMCPInvalidArgs("path required")
|
||||||
|
}
|
||||||
|
out, err := rpc.CallAgent(c.Request.Context(), args.ServerID, model.TaskTypeFsDelete,
|
||||||
|
model.FsDeleteRequest{Path: args.Path, Recursive: args.Recursive}, fsCallTimeout)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var res model.FsDeleteResult
|
||||||
|
if err := json.Unmarshal(out, &res); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if res.Error != "" {
|
||||||
|
return nil, errors.New(res.Error)
|
||||||
|
}
|
||||||
|
return res, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,168 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// fs.* 跨租户拒绝测试:member token 调 fs.list/read/write/delete 时,如果
|
||||||
|
// server.UserID != caller.ID 必须 isError 返回,且不会触达 agent。
|
||||||
|
//
|
||||||
|
// 这些用例不依赖 agent simulator —— 它们要验证的就是 requireServerAccess 在
|
||||||
|
// agent 调用前拦截。如果错误发生在 agent CallAgent,说明权限漏失。
|
||||||
|
|
||||||
|
func makeForeignServerMCP(t *testing.T, id, ownerUID uint64) {
|
||||||
|
t.Helper()
|
||||||
|
srv := &model.Server{}
|
||||||
|
srv.ID = id
|
||||||
|
srv.SetUserID(ownerUID)
|
||||||
|
singleton.ServerShared.InsertForTest(srv)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPFs_List_ForeignServerRejected(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
makeForeignServerMCP(t, 200, 999)
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{
|
||||||
|
Name: "fs.list",
|
||||||
|
Arguments: jsonRaw(map[string]any{"server_id": 200, "path": "/etc"}),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.NotNil(t, tcr)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "permission denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPFs_Read_ForeignServerRejected(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
makeForeignServerMCP(t, 201, 999)
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{
|
||||||
|
Name: "fs.read",
|
||||||
|
Arguments: jsonRaw(map[string]any{"server_id": 201, "path": "/etc/passwd"}),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "permission denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPFs_Write_ForeignServerRejected(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
makeForeignServerMCP(t, 202, 999)
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerWrite}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{
|
||||||
|
Name: "fs.write",
|
||||||
|
Arguments: jsonRaw(map[string]any{
|
||||||
|
"server_id": 202, "path": "/tmp/evil", "content": "x",
|
||||||
|
}),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "permission denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPFs_Delete_ForeignServerRejected(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
makeForeignServerMCP(t, 203, 999)
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerDelete}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{
|
||||||
|
Name: "fs.delete",
|
||||||
|
Arguments: jsonRaw(map[string]any{"server_id": 203, "path": "/tmp/foo"}),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "permission denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPFs_DownloadURL_ForeignServerRejected(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
makeForeignServerMCP(t, 204, 999)
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{
|
||||||
|
Name: "fs.download_url",
|
||||||
|
Arguments: jsonRaw(map[string]any{"server_id": 204, "path": "/etc/shadow"}),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "permission denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPFs_UploadURL_ForeignServerRejected(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
makeForeignServerMCP(t, 205, 999)
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerWrite}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{
|
||||||
|
Name: "fs.upload_url",
|
||||||
|
Arguments: jsonRaw(map[string]any{"server_id": 205, "path": "/tmp/up"}),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "permission denied")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPFs_PATServerWhitelistFiltersFs(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
// Both servers owned by the same user, but the PAT only whitelists server 300.
|
||||||
|
// fs.list against server 7 (in setupMCPTest) must be denied even though
|
||||||
|
// caller user owns it, because the PAT was minted for server 300 only.
|
||||||
|
srv := &model.Server{}
|
||||||
|
srv.ID = 300
|
||||||
|
srv.SetUserID(uid)
|
||||||
|
singleton.ServerShared.InsertForTest(srv)
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, []uint64{300})
|
||||||
|
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{
|
||||||
|
Name: "fs.list",
|
||||||
|
Arguments: jsonRaw(map[string]any{"server_id": 7, "path": "/etc"}),
|
||||||
|
}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "permission denied")
|
||||||
|
}
|
||||||
@@ -0,0 +1,61 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
)
|
||||||
|
|
||||||
|
// meta.whoami 让 LLM 启动时知道自己拿的是哪张 PAT、能干什么、能动哪些服务器。
|
||||||
|
// 不要求任何 scope(任意有效 PAT 均可调用)。
|
||||||
|
type whoamiResult struct {
|
||||||
|
UserID uint64 `json:"user_id"`
|
||||||
|
IsAdmin bool `json:"is_admin"`
|
||||||
|
TokenID uint64 `json:"token_id"`
|
||||||
|
TokenName string `json:"token_name"`
|
||||||
|
Scopes []string `json:"scopes"`
|
||||||
|
ServerIDs []uint64 `json:"server_ids,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
registerMCPTool(&mcpTool{
|
||||||
|
Name: "meta.whoami",
|
||||||
|
Description: "Return the identity, scopes and accessible server IDs of the current API token.",
|
||||||
|
InputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{},
|
||||||
|
},
|
||||||
|
OutputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"user_id": map[string]any{"type": "integer"},
|
||||||
|
"is_admin": map[string]any{"type": "boolean"},
|
||||||
|
"token_id": map[string]any{"type": "integer"},
|
||||||
|
"token_name": map[string]any{"type": "string"},
|
||||||
|
"scopes": map[string]any{"type": "array", "items": map[string]any{"type": "string"}},
|
||||||
|
"server_ids": map[string]any{"type": "array", "items": map[string]any{"type": "integer"}},
|
||||||
|
},
|
||||||
|
"required": []string{"user_id", "is_admin", "token_id", "scopes"},
|
||||||
|
},
|
||||||
|
RequiredScope: "",
|
||||||
|
Handler: handleMetaWhoami,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleMetaWhoami(c *gin.Context, _ json.RawMessage) (any, error) {
|
||||||
|
tok := APITokenFromContext(c)
|
||||||
|
if tok == nil {
|
||||||
|
return nil, errNoToken
|
||||||
|
}
|
||||||
|
user, _ := c.MustGet(model.CtxKeyAuthorizedUser).(*model.User)
|
||||||
|
return whoamiResult{
|
||||||
|
UserID: user.ID,
|
||||||
|
IsAdmin: user.Role.IsAdmin(),
|
||||||
|
TokenID: tok.ID,
|
||||||
|
TokenName: tok.Name,
|
||||||
|
Scopes: tok.Scopes(),
|
||||||
|
ServerIDs: tok.ServerIDs(),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,203 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// server.list 返回当前 PAT 可见的服务器精简列表。
|
||||||
|
//
|
||||||
|
// 输出字段刻意保持小:LLM context 很贵,列 100 台机器时不要把整张 Host/State 表
|
||||||
|
// 全塞进去。需要细节时再调 server.get。
|
||||||
|
type serverListItem struct {
|
||||||
|
ID uint64 `json:"id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
UUID string `json:"uuid,omitempty"`
|
||||||
|
IPv4 string `json:"ipv4,omitempty"`
|
||||||
|
IPv6 string `json:"ipv6,omitempty"`
|
||||||
|
Online bool `json:"online"`
|
||||||
|
Platform string `json:"platform,omitempty"`
|
||||||
|
Arch string `json:"arch,omitempty"`
|
||||||
|
LastActive time.Time `json:"last_active,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type serverListArgs struct {
|
||||||
|
OnlineOnly bool `json:"online_only,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// serverListResult 是 server.list 的返回外壳。MCP 2025-06-18 规定
|
||||||
|
// structuredContent 必须是 JSON object,不能是裸数组/标量,否则严格客户端
|
||||||
|
// (官方 TS/Python SDK)会拒绝整条 tools/call 结果。因此这里把列表包进对象,
|
||||||
|
// 不再直接返回 []serverListItem。
|
||||||
|
type serverListResult struct {
|
||||||
|
Servers []serverListItem `json:"servers"`
|
||||||
|
Count int `json:"count"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func serverListItemSchema() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"id": map[string]any{"type": "integer"},
|
||||||
|
"name": map[string]any{"type": "string"},
|
||||||
|
"uuid": map[string]any{"type": "string"},
|
||||||
|
"ipv4": map[string]any{"type": "string"},
|
||||||
|
"ipv6": map[string]any{"type": "string"},
|
||||||
|
"online": map[string]any{"type": "boolean"},
|
||||||
|
"platform": map[string]any{"type": "string"},
|
||||||
|
"arch": map[string]any{"type": "string"},
|
||||||
|
"last_active": map[string]any{"type": "string", "format": "date-time"},
|
||||||
|
},
|
||||||
|
"required": []string{"id", "name", "online"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
registerMCPTool(&mcpTool{
|
||||||
|
Name: "server.list",
|
||||||
|
Description: "List servers visible to the current API token. Returns minimal metadata; call server.get for full details.",
|
||||||
|
InputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"online_only": map[string]any{
|
||||||
|
"type": "boolean",
|
||||||
|
"description": "If true, only return servers that have reported within the last 30s.",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
OutputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"servers": map[string]any{
|
||||||
|
"type": "array",
|
||||||
|
"items": serverListItemSchema(),
|
||||||
|
},
|
||||||
|
"count": map[string]any{"type": "integer"},
|
||||||
|
},
|
||||||
|
"required": []string{"servers", "count"},
|
||||||
|
},
|
||||||
|
RequiredScope: model.ScopeInventoryRead,
|
||||||
|
Handler: handleServerList,
|
||||||
|
})
|
||||||
|
|
||||||
|
registerMCPTool(&mcpTool{
|
||||||
|
Name: "server.get",
|
||||||
|
Description: "Return full Host/State snapshot for a single server.",
|
||||||
|
InputSchema: serverGetSchema(),
|
||||||
|
OutputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"id": map[string]any{"type": "integer"},
|
||||||
|
"name": map[string]any{"type": "string"},
|
||||||
|
"uuid": map[string]any{"type": "string"},
|
||||||
|
"note": map[string]any{"type": "string"},
|
||||||
|
"public_note": map[string]any{"type": "string"},
|
||||||
|
"host": map[string]any{"type": "object"},
|
||||||
|
"state": map[string]any{"type": "object"},
|
||||||
|
"geoip": map[string]any{"type": "object"},
|
||||||
|
"last_active": map[string]any{"type": "string", "format": "date-time"},
|
||||||
|
},
|
||||||
|
"required": []string{"id"},
|
||||||
|
},
|
||||||
|
RequiredScope: model.ScopeServerRead,
|
||||||
|
Handler: handleServerGet,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleServerList(c *gin.Context, raw json.RawMessage) (any, error) {
|
||||||
|
var args serverListArgs
|
||||||
|
if err := decodeToolArgs(raw, &args); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
tok := APITokenFromContext(c)
|
||||||
|
if tok == nil {
|
||||||
|
return nil, errNoToken
|
||||||
|
}
|
||||||
|
|
||||||
|
slist := singleton.ServerShared.GetSortedList()
|
||||||
|
now := time.Now()
|
||||||
|
const onlineWindow = 30 * time.Second
|
||||||
|
|
||||||
|
out := make([]serverListItem, 0, len(slist))
|
||||||
|
for _, s := range slist {
|
||||||
|
if s == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// 闸 1:复用现有用户权限过滤
|
||||||
|
if !s.HasPermission(c) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
// 闸 2:PAT 的 server 白名单(若已设置)
|
||||||
|
if !tok.CanAccessServer(s.ID) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
runtime := s.RuntimeSnapshot()
|
||||||
|
online := !runtime.LastActive.IsZero() && now.Sub(runtime.LastActive) < onlineWindow
|
||||||
|
if args.OnlineOnly && !online {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
item := serverListItem{
|
||||||
|
ID: s.ID,
|
||||||
|
Name: s.Name,
|
||||||
|
UUID: s.UUID,
|
||||||
|
Online: online,
|
||||||
|
LastActive: runtime.LastActive,
|
||||||
|
}
|
||||||
|
if runtime.Host != nil {
|
||||||
|
item.Platform = runtime.Host.Platform
|
||||||
|
item.Arch = runtime.Host.Arch
|
||||||
|
}
|
||||||
|
if s.GeoIP != nil {
|
||||||
|
item.IPv4 = s.GeoIP.IP.IPv4Addr
|
||||||
|
item.IPv6 = s.GeoIP.IP.IPv6Addr
|
||||||
|
}
|
||||||
|
out = append(out, item)
|
||||||
|
}
|
||||||
|
return serverListResult{Servers: out, Count: len(out)}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// server.get
|
||||||
|
type serverGetArgs struct {
|
||||||
|
ServerID uint64 `json:"server_id"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func serverGetSchema() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"server_id": map[string]any{
|
||||||
|
"type": "integer",
|
||||||
|
"description": "Target server ID.",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"required": []string{"server_id"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleServerGet(c *gin.Context, raw json.RawMessage) (any, error) {
|
||||||
|
var args serverGetArgs
|
||||||
|
if err := decodeToolArgs(raw, &args); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
s, err := requireServerAccess(c, args.ServerID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
runtime := s.RuntimeSnapshot()
|
||||||
|
return map[string]any{
|
||||||
|
"id": s.ID,
|
||||||
|
"name": s.Name,
|
||||||
|
"uuid": s.UUID,
|
||||||
|
"note": s.Note,
|
||||||
|
"public_note": s.PublicNote,
|
||||||
|
"host": runtime.Host,
|
||||||
|
"state": runtime.State,
|
||||||
|
"geoip": s.GeoIP,
|
||||||
|
"last_active": runtime.LastActive,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,122 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestServerList_FiltersByPermission(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
srv2 := &model.Server{}
|
||||||
|
srv2.ID = 8
|
||||||
|
srv2.Name = "beta"
|
||||||
|
srv2.SetUserID(999)
|
||||||
|
singleton.ServerShared.InsertForTest(srv2)
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeInventoryRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "server.list", Arguments: json.RawMessage("{}")}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.False(t, tcr.IsError)
|
||||||
|
|
||||||
|
rows := decodeServerListRows(t, tcr.StructuredContent)
|
||||||
|
require.Len(t, rows, 1, "must filter out non-owned server")
|
||||||
|
require.EqualValues(t, 7, rows[0]["id"])
|
||||||
|
}
|
||||||
|
|
||||||
|
// decodeServerListRows unwraps the {servers,count} object that server.list now
|
||||||
|
// returns (MCP requires structuredContent to be an object, not a bare array)
|
||||||
|
// and asserts count stays in sync with the servers slice length.
|
||||||
|
func decodeServerListRows(t *testing.T, structured any) []map[string]any {
|
||||||
|
t.Helper()
|
||||||
|
rb, _ := json.Marshal(structured)
|
||||||
|
var res struct {
|
||||||
|
Servers []map[string]any `json:"servers"`
|
||||||
|
Count int `json:"count"`
|
||||||
|
}
|
||||||
|
require.NoError(t, json.Unmarshal(rb, &res))
|
||||||
|
require.Equal(t, len(res.Servers), res.Count, "count must match servers length")
|
||||||
|
return res.Servers
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServerList_ServerWhitelistFurtherFiltering(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
srv2 := &model.Server{}
|
||||||
|
srv2.ID = 8
|
||||||
|
srv2.Name = "beta"
|
||||||
|
srv2.SetUserID(uid)
|
||||||
|
singleton.ServerShared.InsertForTest(srv2)
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeInventoryRead}, []uint64{8})
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "server.list", Arguments: json.RawMessage("{}")}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.False(t, tcr.IsError)
|
||||||
|
rows := decodeServerListRows(t, tcr.StructuredContent)
|
||||||
|
require.Len(t, rows, 1)
|
||||||
|
require.EqualValues(t, 8, rows[0]["id"])
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServerList_OnlineOnlyFilter(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
srv, _ := singleton.ServerShared.Get(7)
|
||||||
|
require.NotNil(t, srv)
|
||||||
|
srv.LastActive = time.Now()
|
||||||
|
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeInventoryRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "server.list", Arguments: jsonRaw(map[string]any{"online_only": true})}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.False(t, tcr.IsError)
|
||||||
|
rows := decodeServerListRows(t, tcr.StructuredContent)
|
||||||
|
require.Len(t, rows, 1)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServerGet_RequiresServerID(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "server.get", Arguments: json.RawMessage("{}")}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "server_id required")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServerExec_ScopeMissing(t *testing.T) {
|
||||||
|
cleanup, uid := setupMCPTest(t)
|
||||||
|
defer cleanup()
|
||||||
|
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||||
|
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
|
||||||
|
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
|
||||||
|
Params: jsonObj(t, toolCallParams{Name: "server.exec", Arguments: jsonRaw(map[string]any{"server_id": 7, "cmd": "echo"})}),
|
||||||
|
})
|
||||||
|
mcpEndpoint(c)
|
||||||
|
_, tcr := decodeRPC(w)
|
||||||
|
require.True(t, tcr.IsError)
|
||||||
|
require.Contains(t, tcr.Content[0].Text, "nezha:server:exec")
|
||||||
|
}
|
||||||
@@ -0,0 +1,974 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"crypto/hmac"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/binary"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"math"
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/utils"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
// fs.download_url / fs.upload_url 旁路通道。
|
||||||
|
//
|
||||||
|
// 设计目标:给 LLM 客户端一个不经 MCP 上下文的 URL,去用普通 HTTP 客户端
|
||||||
|
// 上传/下载大文件(单文件 hard cap 100MiB,model.MCPFsTransferMaxSize)。
|
||||||
|
//
|
||||||
|
// 传输实现:dashboard ↔ agent 走 gRPC IOStream 双向流(TaskTypeFsTransfer),
|
||||||
|
// dashboard 一边读 HTTP body 一边推给 agent;不再使用 base64/JSON 包装内容,
|
||||||
|
// 避免 gRPC 4MiB 单消息上限。
|
||||||
|
//
|
||||||
|
// 安全机制:
|
||||||
|
// - 一次性 token,存内存 sync.Map,TTL 默认 300s,最多 600s
|
||||||
|
// - token 绑定 user_id + token_id + server_id + path + direction
|
||||||
|
// - consume 时重算并以常数时间比对 entry 的 HMAC-SHA256,防篡改
|
||||||
|
// - 命中后立即从内存删除,禁止重放
|
||||||
|
// - revalidateTransferEntry 在 consume 时重新校验 PAT/scope/owner,应对
|
||||||
|
// mint→consume 之间的权限变化
|
||||||
|
// - 上传可选 ?sha256=<hex> 端到端校验;下载 NZTO 帧附 agent 计算的 sha
|
||||||
|
type transferDirection string
|
||||||
|
|
||||||
|
const (
|
||||||
|
transferDirDownload transferDirection = "download"
|
||||||
|
transferDirUpload transferDirection = "upload"
|
||||||
|
|
||||||
|
transferTokenTTLDefault = 300 * time.Second
|
||||||
|
transferTokenTTLMax = 600 * time.Second
|
||||||
|
|
||||||
|
// maxTransferDuration bounds a single upload/download once the agent has
|
||||||
|
// attached. 100MiB over a slow link still completes well within this;
|
||||||
|
// anything longer is treated as a stalled/abusive transfer and cancelled.
|
||||||
|
maxTransferDuration = 10 * time.Minute
|
||||||
|
|
||||||
|
maxTransferPathLen = 4096
|
||||||
|
)
|
||||||
|
|
||||||
|
func validateTransferPath(path string) error {
|
||||||
|
if path == "" {
|
||||||
|
return errMCPInvalidArgs("path required")
|
||||||
|
}
|
||||||
|
if len(path) > maxTransferPathLen {
|
||||||
|
return errMCPInvalidArgs("path too long")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type transferEntry struct {
|
||||||
|
UserID uint64
|
||||||
|
TokenID uint64
|
||||||
|
ServerID uint64
|
||||||
|
Path string
|
||||||
|
Direction transferDirection
|
||||||
|
ExpiresAt time.Time
|
||||||
|
|
||||||
|
// Upload-only optional knobs carried from MCP fs.upload_url tool args
|
||||||
|
// to transferUploadHandler so the upload handler can forward them into
|
||||||
|
// FsTransferRequest. Empty/false values keep current behaviour for the
|
||||||
|
// download direction (these fields are simply ignored).
|
||||||
|
UploadMode string
|
||||||
|
UploadCreateDirs bool
|
||||||
|
UploadIfMatchSHA256 string
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
transferEntries sync.Map
|
||||||
|
transferSecretMu sync.Mutex
|
||||||
|
transferSecretVal string
|
||||||
|
)
|
||||||
|
|
||||||
|
// transferHMACSecret 返回进程内随机生成的 HMAC key。
|
||||||
|
// 这是有意设计:transferEntries 本身也只活在内存 sync.Map 里,dashboard
|
||||||
|
// 重启等价于全部 token 失效;让 secret 也随进程随机,可以避免“secret 来自
|
||||||
|
// 持久化 env 但 entries 已丢”这种半持久化状态,同时保证多副本部署不会
|
||||||
|
// 意外互认对方签发的 token(每副本一份独立 secret)。
|
||||||
|
func transferHMACSecret() string {
|
||||||
|
transferSecretMu.Lock()
|
||||||
|
defer transferSecretMu.Unlock()
|
||||||
|
if transferSecretVal != "" {
|
||||||
|
return transferSecretVal
|
||||||
|
}
|
||||||
|
transferSecretVal = utils.MustGenerateRandomString(64)
|
||||||
|
return transferSecretVal
|
||||||
|
}
|
||||||
|
|
||||||
|
// transferTokenSig 计算 entry 的 HMAC-SHA256 签名(hex);mint 与 consume 共用。
|
||||||
|
func transferTokenSig(e transferEntry) string {
|
||||||
|
mac := hmac.New(sha256.New, []byte(transferHMACSecret()))
|
||||||
|
fmt.Fprintf(mac, "%s|%d|%d|%d|%s|%d",
|
||||||
|
e.Direction, e.UserID, e.TokenID, e.ServerID, e.Path, e.ExpiresAt.UnixNano())
|
||||||
|
return hex.EncodeToString(mac.Sum(nil))
|
||||||
|
}
|
||||||
|
|
||||||
|
func mintTransferToken(e transferEntry) (string, error) {
|
||||||
|
id, err := utils.GenerateRandomString(24)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
tok := id + "." + transferTokenSig(e)
|
||||||
|
transferEntries.Store(tok, e)
|
||||||
|
return tok, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func consumeTransferToken(tok string, dir transferDirection) (*transferEntry, error) {
|
||||||
|
raw, ok := transferEntries.LoadAndDelete(tok)
|
||||||
|
if !ok {
|
||||||
|
return nil, errors.New("invalid or already-used transfer token")
|
||||||
|
}
|
||||||
|
e, _ := raw.(transferEntry)
|
||||||
|
// 校验 HMAC:token 形如 id.sig,sig 必须等于 entry 字段在进程 secret 下的
|
||||||
|
// HMAC-SHA256。仅靠 sync.Map key 随机性不构成完整性保护——一旦 entry 被
|
||||||
|
// 持久化/跨副本共享/从 token 解码,缺这一步即认证绕过。常数时间比较防侧信道。
|
||||||
|
idx := strings.LastIndex(tok, ".")
|
||||||
|
if idx < 0 {
|
||||||
|
return nil, errors.New("malformed transfer token")
|
||||||
|
}
|
||||||
|
if !hmac.Equal([]byte(tok[idx+1:]), []byte(transferTokenSig(e))) {
|
||||||
|
return nil, errors.New("transfer token signature mismatch")
|
||||||
|
}
|
||||||
|
if e.Direction != dir {
|
||||||
|
return nil, errors.New("transfer token direction mismatch")
|
||||||
|
}
|
||||||
|
if time.Now().After(e.ExpiresAt) {
|
||||||
|
return nil, errors.New("transfer token expired")
|
||||||
|
}
|
||||||
|
return &e, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// PurgeTransferEntries drops every minted-but-unconsumed transfer URL.
|
||||||
|
// EnableMCP=false invokes this so an admin pressing the kill switch
|
||||||
|
// invalidates the 5–10min trailing window of pre-signed download/upload
|
||||||
|
// URLs that consumeTransferToken would otherwise still honor. Returns the
|
||||||
|
// number of entries purged for audit.
|
||||||
|
func PurgeTransferEntries() int {
|
||||||
|
purged := 0
|
||||||
|
transferEntries.Range(func(key, _ any) bool {
|
||||||
|
if _, ok := transferEntries.LoadAndDelete(key); ok {
|
||||||
|
purged++
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
return purged
|
||||||
|
}
|
||||||
|
|
||||||
|
// gcExpiredTransferEntries 按 ExpiresAt 删除所有已过期但从未被 consume
|
||||||
|
// 的 token。kickoffTransferGC 周期调度它,防止 transferEntries 在没人
|
||||||
|
// 触发 kill switch 的情况下随时间无界增长。
|
||||||
|
func gcExpiredTransferEntries(now time.Time) int {
|
||||||
|
removed := 0
|
||||||
|
transferEntries.Range(func(key, raw any) bool {
|
||||||
|
e, ok := raw.(transferEntry)
|
||||||
|
if !ok {
|
||||||
|
transferEntries.Delete(key)
|
||||||
|
removed++
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if now.After(e.ExpiresAt) {
|
||||||
|
if _, deleted := transferEntries.LoadAndDelete(key); deleted {
|
||||||
|
removed++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
})
|
||||||
|
return removed
|
||||||
|
}
|
||||||
|
|
||||||
|
var transferGCStartOnce sync.Once
|
||||||
|
|
||||||
|
// kickoffTransferGC 启动一个进程级 goroutine 定时回收过期 token,
|
||||||
|
// 避免每个 dashboard 启动都得 PurgeTransferEntries 才能把表清空。
|
||||||
|
// 时间间隔取 transferTokenTTLDefault / 5,对默认 5min TTL 即 1min;
|
||||||
|
// 既能在 TTL 内多次扫到过期项,也不会让锁竞争变成热点。
|
||||||
|
func kickoffTransferGC() {
|
||||||
|
transferGCStartOnce.Do(func() {
|
||||||
|
go func() {
|
||||||
|
ticker := time.NewTicker(transferTokenTTLDefault / 5)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for range ticker.C {
|
||||||
|
gcExpiredTransferEntries(time.Now())
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- tool: fs.download_url ---
|
||||||
|
|
||||||
|
type fsDownloadURLArgs struct {
|
||||||
|
ServerID uint64 `json:"server_id"`
|
||||||
|
Path string `json:"path"`
|
||||||
|
TTLSeconds int `json:"ttl_seconds,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// fsUploadURLArgs is the upload-side superset of fsDownloadURLArgs. agent's
|
||||||
|
// FsTransferRequest already supports per-upload Mode / CreateDirs /
|
||||||
|
// IfMatchSHA256, but fs.upload_url historically reused fsDownloadURLArgs and
|
||||||
|
// silently dropped these fields. Splitting the arg shape lets the MCP tool
|
||||||
|
// schema advertise them and mintTransferTool plumb them through to
|
||||||
|
// transferUploadHandler -> openFsTransferStream.
|
||||||
|
type fsUploadURLArgs struct {
|
||||||
|
ServerID uint64 `json:"server_id"`
|
||||||
|
Path string `json:"path"`
|
||||||
|
TTLSeconds int `json:"ttl_seconds,omitempty"`
|
||||||
|
Mode string `json:"mode,omitempty"`
|
||||||
|
CreateDirs bool `json:"create_dirs,omitempty"`
|
||||||
|
IfMatchSHA256 string `json:"if_match_sha256,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
registerMCPTool(&mcpTool{
|
||||||
|
Name: "fs.download_url",
|
||||||
|
Description: "Mint a one-time signed URL to stream a file (<=100MiB) via plain HTTP GET. Bypasses MCP context and uses gRPC IOStream end-to-end.",
|
||||||
|
InputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"server_id": map[string]any{"type": "integer"},
|
||||||
|
"path": map[string]any{"type": "string"},
|
||||||
|
"ttl_seconds": map[string]any{"type": "integer", "minimum": 30, "maximum": 600},
|
||||||
|
},
|
||||||
|
"required": []string{"server_id", "path"},
|
||||||
|
},
|
||||||
|
OutputSchema: transferURLOutputSchema(),
|
||||||
|
RequiredScope: model.ScopeServerRead,
|
||||||
|
Handler: handleFsDownloadURL,
|
||||||
|
})
|
||||||
|
|
||||||
|
registerMCPTool(&mcpTool{
|
||||||
|
Name: "fs.upload_url",
|
||||||
|
Description: "Mint a one-time signed URL to stream a file (<=100MiB) via plain HTTP POST. Caller MUST send Content-Length; optional ?sha256=<hex> for end-to-end integrity. mode / create_dirs / if_match_sha256 are forwarded to the agent for atomic chmod / mkdir -p / optimistic concurrency.",
|
||||||
|
InputSchema: map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"server_id": map[string]any{"type": "integer"},
|
||||||
|
"path": map[string]any{"type": "string"},
|
||||||
|
"ttl_seconds": map[string]any{"type": "integer", "minimum": 30, "maximum": 600},
|
||||||
|
"mode": map[string]any{"type": "string", "description": "Octal mode like '0644'."},
|
||||||
|
"create_dirs": map[string]any{"type": "boolean"},
|
||||||
|
"if_match_sha256": map[string]any{"type": "string", "description": "64 hex chars; precondition checked by the agent before overwrite."},
|
||||||
|
},
|
||||||
|
"required": []string{"server_id", "path"},
|
||||||
|
},
|
||||||
|
OutputSchema: transferURLOutputSchema(),
|
||||||
|
RequiredScope: model.ScopeServerWrite,
|
||||||
|
Handler: handleFsUploadURL,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func transferURLOutputSchema() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
"properties": map[string]any{
|
||||||
|
"url": map[string]any{"type": "string"},
|
||||||
|
"method": map[string]any{"type": "string"},
|
||||||
|
"expires_at": map[string]any{"type": "string", "format": "date-time"},
|
||||||
|
},
|
||||||
|
"required": []string{"url", "method", "expires_at"},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleFsDownloadURL(c *gin.Context, raw json.RawMessage) (any, error) {
|
||||||
|
var args fsDownloadURLArgs
|
||||||
|
if err := decodeToolArgs(raw, &args); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return mintTransferTool(c, args.ServerID, args.Path, args.TTLSeconds, transferDirDownload, transferEntry{})
|
||||||
|
}
|
||||||
|
|
||||||
|
func handleFsUploadURL(c *gin.Context, raw json.RawMessage) (any, error) {
|
||||||
|
var args fsUploadURLArgs
|
||||||
|
if err := decodeToolArgs(raw, &args); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if args.IfMatchSHA256 != "" {
|
||||||
|
if _, decErr := hex.DecodeString(args.IfMatchSHA256); decErr != nil || len(args.IfMatchSHA256) != 64 {
|
||||||
|
return nil, errMCPInvalidArgs("if_match_sha256 must be 64 hex chars")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return mintTransferTool(c, args.ServerID, args.Path, args.TTLSeconds, transferDirUpload, transferEntry{
|
||||||
|
UploadMode: args.Mode,
|
||||||
|
UploadCreateDirs: args.CreateDirs,
|
||||||
|
UploadIfMatchSHA256: args.IfMatchSHA256,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func mintTransferTool(c *gin.Context, serverID uint64, path string, ttlSeconds int, dir transferDirection, uploadExtras transferEntry) (any, error) {
|
||||||
|
srv, err := requireServerAccess(c, serverID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if err := requireAgentSupportsMCP(srv); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if err := validateTransferPath(path); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
ttl := time.Duration(ttlSeconds) * time.Second
|
||||||
|
if ttl <= 0 {
|
||||||
|
ttl = transferTokenTTLDefault
|
||||||
|
}
|
||||||
|
if ttl > transferTokenTTLMax {
|
||||||
|
ttl = transferTokenTTLMax
|
||||||
|
}
|
||||||
|
|
||||||
|
tok := APITokenFromContext(c)
|
||||||
|
if tok == nil {
|
||||||
|
return nil, errNoToken
|
||||||
|
}
|
||||||
|
uid := uint64(0)
|
||||||
|
if u, ok := c.Get(model.CtxKeyAuthorizedUser); ok {
|
||||||
|
if user, ok := u.(*model.User); ok && user != nil {
|
||||||
|
uid = user.ID
|
||||||
|
}
|
||||||
|
}
|
||||||
|
entry := transferEntry{
|
||||||
|
UserID: uid,
|
||||||
|
TokenID: tok.ID,
|
||||||
|
ServerID: serverID,
|
||||||
|
Path: path,
|
||||||
|
Direction: dir,
|
||||||
|
ExpiresAt: time.Now().Add(ttl),
|
||||||
|
UploadMode: uploadExtras.UploadMode,
|
||||||
|
UploadCreateDirs: uploadExtras.UploadCreateDirs,
|
||||||
|
UploadIfMatchSHA256: uploadExtras.UploadIfMatchSHA256,
|
||||||
|
}
|
||||||
|
t, err := mintTransferToken(entry)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
scheme := "https"
|
||||||
|
if c.Request.TLS == nil && c.Request.Header.Get("X-Forwarded-Proto") != "https" {
|
||||||
|
scheme = "http"
|
||||||
|
}
|
||||||
|
host := c.Request.Host
|
||||||
|
url := fmt.Sprintf("%s://%s/mcp/%s/%s", scheme, host, dir, t)
|
||||||
|
return map[string]any{
|
||||||
|
"url": url,
|
||||||
|
"method": map[transferDirection]string{transferDirDownload: "GET", transferDirUpload: "POST"}[dir],
|
||||||
|
"expires_at": entry.ExpiresAt,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// --- HTTP handlers ---
|
||||||
|
|
||||||
|
// transferDownloadHandler 处理 GET /mcp/download/:token。
|
||||||
|
// 走 IOStream 双向流:dashboard 把 agent 推过来的 chunk 转发给 HTTP 客户端,
|
||||||
|
// 单文件 hard cap 100MiB(model.MCPFsTransferMaxSize)。
|
||||||
|
// transferRevokableContext 把进行中的传输纳入 PAT 撤销注册表。返回的 ctx 在
|
||||||
|
// 该 PAT 被 deleteAPIToken 撤销时取消,从而切断已开始的 upload/download;
|
||||||
|
// 否则只在传输自然结束时由 stop() 注销。stop() 必须 defer 调用。
|
||||||
|
func transferRevokableContext(c *gin.Context, e *transferEntry) (context.Context, func()) {
|
||||||
|
// Cap the whole transfer with a hard deadline. After the agent attaches,
|
||||||
|
// the relay blocks in IOStreamWrapper.Read, which only honours this ctx
|
||||||
|
// (openFsTransferStream closes the stream on ctx.Done). Without the
|
||||||
|
// deadline a stalled or malicious agent that attaches but never sends a
|
||||||
|
// complete header/chunk/final frame pins this goroutine, the IOStream and
|
||||||
|
// the spool tmpfile until the client disconnects, allowing concurrent
|
||||||
|
// hung transfers to exhaust resources within the rate limit.
|
||||||
|
ctx, cancel := context.WithTimeout(c.Request.Context(), maxTransferDuration)
|
||||||
|
dereg := patConnectionRegistryShared.register(e.TokenID, cancel)
|
||||||
|
return ctx, func() {
|
||||||
|
dereg()
|
||||||
|
cancel()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func transferDownloadHandler(c *gin.Context) {
|
||||||
|
tok := c.Param("token")
|
||||||
|
entry, err := consumeTransferToken(tok, transferDirDownload)
|
||||||
|
if err != nil {
|
||||||
|
writeTransferFailureAudit(c, nil, "fs.download", classifyTransferConsumeError(err), err)
|
||||||
|
c.String(http.StatusUnauthorized, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := revalidateTransferEntry(entry); err != nil {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.download", classifyTransferRevalidateError(err), err)
|
||||||
|
c.String(http.StatusUnauthorized, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ctx, stop := transferRevokableContext(c, entry)
|
||||||
|
defer stop()
|
||||||
|
|
||||||
|
stream, cleanup, err := openFsTransferStream(ctx, entry.ServerID, &model.FsTransferRequest{
|
||||||
|
Op: model.MCPFsTransferOpDownload,
|
||||||
|
Path: entry.Path,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.download", classifyTransferOpenStreamError(err), err)
|
||||||
|
c.String(http.StatusBadGateway, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
hdr, err := readXferFixedHeader(stream)
|
||||||
|
if err != nil {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.download", model.MCPOutcomeAgentTimeout, err)
|
||||||
|
c.String(http.StatusBadGateway, "agent did not return download header: "+err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if hdr.IsErr() {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.download", model.MCPOutcomeAgentError, errors.New(hdr.ErrMsg))
|
||||||
|
c.String(http.StatusBadGateway, hdr.ErrMsg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !bytes.Equal(hdr.Magic, model.MCPFsXferMagicDownloadHdr) {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.download", model.MCPOutcomeAgentError, errors.New("unexpected header magic"))
|
||||||
|
c.String(http.StatusBadGateway, "agent returned unexpected header magic")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if hdr.Size > model.MCPFsTransferMaxSize {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.download", model.MCPOutcomeAgentError, errors.New("file exceeds MCP transfer cap"))
|
||||||
|
c.String(http.StatusBadGateway, "file exceeds MCP transfer cap (100MiB)")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := relayDownloadFrames(c, stream, hdr.Size); err != nil {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.download", model.MCPOutcomeAgentError, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
_ = singleton.DB.Create(&model.MCPAuditLog{
|
||||||
|
UserID: entry.UserID,
|
||||||
|
TokenID: entry.TokenID,
|
||||||
|
Tool: "fs.download",
|
||||||
|
ServerID: entry.ServerID,
|
||||||
|
Outcome: model.MCPOutcomeOK,
|
||||||
|
IP: c.GetString(model.CtxKeyRealIPStr),
|
||||||
|
}).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// transferUploadHandler 处理 POST /mcp/upload/:token;body 转发到 agent,
|
||||||
|
// 单文件 hard cap 100MiB。
|
||||||
|
func transferUploadHandler(c *gin.Context) {
|
||||||
|
tok := c.Param("token")
|
||||||
|
entry, err := consumeTransferToken(tok, transferDirUpload)
|
||||||
|
if err != nil {
|
||||||
|
writeTransferFailureAudit(c, nil, "fs.upload", classifyTransferConsumeError(err), err)
|
||||||
|
c.String(http.StatusUnauthorized, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := revalidateTransferEntry(entry); err != nil {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", classifyTransferRevalidateError(err), err)
|
||||||
|
c.String(http.StatusUnauthorized, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// 1) 体积闸门:Content-Length 必须存在并且不超过 cap。流式上传时这是
|
||||||
|
// 唯一能在打开 IOStream 之前就拒掉超大请求的依据,避免 agent 端
|
||||||
|
// 拒绝时已经占了一个连接。
|
||||||
|
if c.Request.ContentLength < 0 {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeInvalidArgs, errors.New("Content-Length required"))
|
||||||
|
c.String(http.StatusLengthRequired, "Content-Length required")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if c.Request.ContentLength > model.MCPFsTransferMaxSize {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeInvalidArgs, errors.New("body exceeds MCP transfer cap"))
|
||||||
|
c.String(http.StatusRequestEntityTooLarge, "body exceeds MCP transfer cap (100MiB)")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
size := c.Request.ContentLength
|
||||||
|
|
||||||
|
// 可选的端到端 sha256:通过 query 参数 sha256=<hex> 传入,agent 收到全部
|
||||||
|
// 字节后会比对;不传则只回带 sha 但不强校验。
|
||||||
|
expected := strings.ToLower(strings.TrimSpace(c.Query("sha256")))
|
||||||
|
if expected != "" {
|
||||||
|
if _, decErr := hex.DecodeString(expected); decErr != nil || len(expected) != 64 {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeInvalidArgs, errors.New("sha256 must be 64 hex chars"))
|
||||||
|
c.String(http.StatusBadRequest, "sha256 must be 64 hex chars")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 上限再加一个字节做 MaxBytesReader 屏障:若客户端撒谎、实际 body 超过
|
||||||
|
// Content-Length,HTTP 层会立即截断并报 413。
|
||||||
|
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, model.MCPFsTransferMaxSize+1)
|
||||||
|
|
||||||
|
ctx, stop := transferRevokableContext(c, entry)
|
||||||
|
defer stop()
|
||||||
|
|
||||||
|
stream, cleanup, err := openFsTransferStream(ctx, entry.ServerID, &model.FsTransferRequest{
|
||||||
|
Op: model.MCPFsTransferOpUpload,
|
||||||
|
Path: entry.Path,
|
||||||
|
Size: size,
|
||||||
|
ExpectedSHA256: expected,
|
||||||
|
Mode: entry.UploadMode,
|
||||||
|
CreateDirs: entry.UploadCreateDirs,
|
||||||
|
IfMatchSHA256: entry.UploadIfMatchSHA256,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", classifyTransferOpenStreamError(err), err)
|
||||||
|
c.String(http.StatusBadGateway, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
hdr, err := readXferFixedHeader(stream)
|
||||||
|
if err != nil {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentTimeout, err)
|
||||||
|
c.String(http.StatusBadGateway, "agent did not return upload ready frame: "+err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if hdr.IsErr() {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentError, errors.New(hdr.ErrMsg))
|
||||||
|
c.String(http.StatusBadGateway, hdr.ErrMsg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !bytes.Equal(hdr.Magic, model.MCPFsXferMagicUploadHdr) {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentError, errors.New("unexpected header magic"))
|
||||||
|
c.String(http.StatusBadGateway, "agent returned unexpected header magic")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if hdr.Size != size {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentError, errors.New("agent acknowledged unexpected size"))
|
||||||
|
c.String(http.StatusBadGateway, "agent acknowledged unexpected size")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, copyErr := io.CopyN(stream, c.Request.Body, size); copyErr != nil {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentError, copyErr)
|
||||||
|
c.String(http.StatusBadGateway, "stream relay failed: "+copyErr.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
final, err := readXferFixedHeader(stream)
|
||||||
|
if err != nil {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentTimeout, err)
|
||||||
|
c.String(http.StatusBadGateway, "agent did not acknowledge upload: "+err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if final.IsErr() {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentError, errors.New(final.ErrMsg))
|
||||||
|
c.String(http.StatusBadGateway, final.ErrMsg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if !bytes.Equal(final.Magic, model.MCPFsXferMagicOK) {
|
||||||
|
writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentError, errors.New("unexpected final magic"))
|
||||||
|
c.String(http.StatusBadGateway, "agent returned unexpected final magic")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
c.JSON(http.StatusOK, model.FsWriteResult{Size: int64(final.Size), SHA256: hex.EncodeToString(final.SHA256)})
|
||||||
|
_ = singleton.DB.Create(&model.MCPAuditLog{
|
||||||
|
UserID: entry.UserID,
|
||||||
|
TokenID: entry.TokenID,
|
||||||
|
Tool: "fs.upload",
|
||||||
|
ServerID: entry.ServerID,
|
||||||
|
Outcome: model.MCPOutcomeOK,
|
||||||
|
IP: c.GetString(model.CtxKeyRealIPStr),
|
||||||
|
}).Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// writeTransferFailureAudit 是 fs.upload / fs.download HTTP handler 失败路径
|
||||||
|
// 共用的审计写入。outcome 必须用 model.MCPOutcome* 常量;entry 可以是 nil
|
||||||
|
// (token consume 阶段就失败时拿不到 entry,UserID/TokenID/ServerID 写 0)。
|
||||||
|
//
|
||||||
|
// Anonymous failures (entry == nil) go through transferAnonAuditThrottleShared
|
||||||
|
// per-IP sampler so an unauthenticated attacker cannot flood mcp_audit_log
|
||||||
|
// by POSTing /mcp/upload/<random>. Authenticated failures bypass the
|
||||||
|
// throttle so SIEM signal stays intact.
|
||||||
|
func writeTransferFailureAudit(c *gin.Context, entry *transferEntry, tool, outcome string, err error) {
|
||||||
|
ip := c.GetString(model.CtxKeyRealIPStr)
|
||||||
|
if entry == nil && !transferAnonAuditThrottleShared.shouldRecord(ip) {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
entryLog := model.MCPAuditLog{
|
||||||
|
Tool: tool,
|
||||||
|
Outcome: outcome,
|
||||||
|
IP: ip,
|
||||||
|
}
|
||||||
|
if entry != nil {
|
||||||
|
entryLog.UserID = entry.UserID
|
||||||
|
entryLog.TokenID = entry.TokenID
|
||||||
|
entryLog.ServerID = entry.ServerID
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
msg := err.Error()
|
||||||
|
if len(msg) > 512 {
|
||||||
|
msg = msg[:512]
|
||||||
|
}
|
||||||
|
entryLog.ErrorMsg = msg
|
||||||
|
entryLog.ErrorCode = outcome
|
||||||
|
}
|
||||||
|
mcpAuditWrite(entryLog, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
// classifyTransferConsumeError 把 consumeTransferToken 的错误映射成 outcome。
|
||||||
|
// 让 SIEM 能区分“伪造/过期 token”与“direction 不匹配”等场景。
|
||||||
|
func classifyTransferConsumeError(err error) string {
|
||||||
|
if err == nil {
|
||||||
|
return model.MCPOutcomeInternalError
|
||||||
|
}
|
||||||
|
msg := err.Error()
|
||||||
|
switch {
|
||||||
|
case strings.Contains(msg, "expired"):
|
||||||
|
return model.MCPOutcomeScopeDenied
|
||||||
|
case strings.Contains(msg, "direction mismatch"):
|
||||||
|
return model.MCPOutcomeInvalidArgs
|
||||||
|
default:
|
||||||
|
return model.MCPOutcomePermDenied
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// classifyTransferRevalidateError 把 revalidateTransferEntry 的错误映射成
|
||||||
|
// outcome。最重要的一项是“MCP is disabled” → MCPOutcomeMCPDisabled,让
|
||||||
|
// 运营在 audit 表里直接看出 kill switch 命中情况,而不是只看到 perm_denied。
|
||||||
|
func classifyTransferRevalidateError(err error) string {
|
||||||
|
if err == nil {
|
||||||
|
return model.MCPOutcomeInternalError
|
||||||
|
}
|
||||||
|
msg := err.Error()
|
||||||
|
switch {
|
||||||
|
case strings.Contains(msg, "MCP is disabled"):
|
||||||
|
return model.MCPOutcomeMCPDisabled
|
||||||
|
case strings.Contains(msg, "no longer has required scope"):
|
||||||
|
return model.MCPOutcomeScopeDenied
|
||||||
|
case strings.Contains(msg, "no longer covers"):
|
||||||
|
return model.MCPOutcomeScopeDenied
|
||||||
|
case strings.Contains(msg, "expired"):
|
||||||
|
return model.MCPOutcomeScopeDenied
|
||||||
|
default:
|
||||||
|
return model.MCPOutcomePermDenied
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// classifyTransferOpenStreamError 把 openFsTransferStream 的失败映射成
|
||||||
|
// outcome:offline / 30s attach 超时分别对应 ServerOffline / AgentTimeout。
|
||||||
|
func classifyTransferOpenStreamError(err error) string {
|
||||||
|
if err == nil {
|
||||||
|
return model.MCPOutcomeInternalError
|
||||||
|
}
|
||||||
|
msg := err.Error()
|
||||||
|
switch {
|
||||||
|
case strings.Contains(msg, "server offline"):
|
||||||
|
return model.MCPOutcomeServerOffline
|
||||||
|
case strings.Contains(msg, "did not attach"):
|
||||||
|
return model.MCPOutcomeAgentTimeout
|
||||||
|
default:
|
||||||
|
return model.MCPOutcomeAgentError
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// frameReceiver is the frame-preserving subset of grpcx.IOStreamWrapper that
|
||||||
|
// the download relay needs. We accept the interface (not the concrete type)
|
||||||
|
// so test simulators can plug in a net.Pipe-backed stream without depending
|
||||||
|
// on the gRPC stack.
|
||||||
|
type frameReceiver interface {
|
||||||
|
RecvFrame() ([]byte, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// relayDownloadFrames forwards declared-size payload from agent to HTTP
|
||||||
|
// client. The agent wraps every data chunk in an NZTC frame (4-byte magic +
|
||||||
|
// 8-byte big-endian length + payload) so payload that happens to begin with
|
||||||
|
// the same bytes as a control frame (NZTE / NZTO) cannot be misclassified.
|
||||||
|
// Control frames (NZTE error, NZTO success) sit on the same IOStream and
|
||||||
|
// are recognized by their magic; legitimate payload always arrives inside
|
||||||
|
// NZTC frames and is never matched against the control-frame magics.
|
||||||
|
//
|
||||||
|
// Payload is spooled to a per-request tmpfile rather than kept in a 100MiB
|
||||||
|
// memory buffer: a midstream NZTE must be able to switch the HTTP response
|
||||||
|
// to 502, which forces us to defer the body write until the final NZTO
|
||||||
|
// frame is observed; but we MUST NOT pay 100MiB of heap per concurrent
|
||||||
|
// download to do so.
|
||||||
|
func relayDownloadFrames(c *gin.Context, stream io.ReadWriteCloser, size int64) error {
|
||||||
|
spool, err := newTransferSpool()
|
||||||
|
if err != nil {
|
||||||
|
c.String(http.StatusInternalServerError, "transfer spool: "+err.Error())
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer spool.Close()
|
||||||
|
|
||||||
|
// Hash the relayed bytes inline; we compare against the trailing
|
||||||
|
// NZTO declared sha256 in validateDownloadFinal so corrupt or
|
||||||
|
// truncated agent payloads can't reach the client.
|
||||||
|
hasher := sha256.New()
|
||||||
|
streamed := int64(0)
|
||||||
|
|
||||||
|
remaining := size
|
||||||
|
header := make([]byte, 4+8)
|
||||||
|
for remaining > 0 {
|
||||||
|
if _, err := io.ReadFull(stream, header); err != nil {
|
||||||
|
c.String(http.StatusBadGateway, "stream relay failed: "+err.Error())
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if bytes.HasPrefix(header, model.MCPFsXferMagicErr) {
|
||||||
|
msg := readMidstreamErrMsg(stream, header)
|
||||||
|
c.String(http.StatusBadGateway, msg)
|
||||||
|
return errMCPMidstreamAbort
|
||||||
|
}
|
||||||
|
if !bytes.HasPrefix(header, model.MCPFsXferMagicChunk) {
|
||||||
|
c.String(http.StatusBadGateway, "stream relay failed: expected NZTC chunk frame")
|
||||||
|
return errMCPMidstreamAbort
|
||||||
|
}
|
||||||
|
chunkLen := binary.BigEndian.Uint64(header[4:12])
|
||||||
|
if chunkLen == 0 {
|
||||||
|
// A zero-length data frame makes no progress toward `remaining`.
|
||||||
|
// Treating it as a no-op `continue` lets a malicious or buggy
|
||||||
|
// agent stream an unbounded run of zero-length NZTC frames,
|
||||||
|
// pinning this goroutine, the gRPC stream and the spool tmpfile
|
||||||
|
// forever (the final NZTO is never reached). Reject it: a real
|
||||||
|
// transfer that still owes bytes never needs an empty data frame.
|
||||||
|
c.String(http.StatusBadGateway, "stream relay failed: zero-length data frame while payload incomplete")
|
||||||
|
return errMCPMidstreamAbort
|
||||||
|
}
|
||||||
|
if int64(chunkLen) > remaining {
|
||||||
|
c.String(http.StatusBadGateway, "agent oversent: more data bytes than declared size")
|
||||||
|
return errMCPMidstreamAbort
|
||||||
|
}
|
||||||
|
n, err := io.CopyN(io.MultiWriter(spool, hasher), stream, int64(chunkLen))
|
||||||
|
if err != nil {
|
||||||
|
c.String(http.StatusBadGateway, "stream relay failed: "+err.Error())
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
streamed += n
|
||||||
|
remaining -= n
|
||||||
|
}
|
||||||
|
|
||||||
|
final := make([]byte, 4+8+32)
|
||||||
|
if _, err := io.ReadFull(stream, final); err != nil {
|
||||||
|
c.String(http.StatusBadGateway, "agent did not send final transfer frame: "+err.Error())
|
||||||
|
return errMCPMidstreamAbort
|
||||||
|
}
|
||||||
|
if bytes.HasPrefix(final, model.MCPFsXferMagicErr) {
|
||||||
|
msg := readMidstreamErrMsg(stream, final[:4+8])
|
||||||
|
c.String(http.StatusBadGateway, msg)
|
||||||
|
return errMCPMidstreamAbort
|
||||||
|
}
|
||||||
|
if err := validateDownloadFinal(final, streamed, hasher.Sum(nil)); err != nil {
|
||||||
|
c.String(http.StatusBadGateway, err.Error())
|
||||||
|
return errMCPMidstreamAbort
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := spool.Rewind(); err != nil {
|
||||||
|
c.String(http.StatusInternalServerError, "transfer spool rewind: "+err.Error())
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
c.Header("Content-Type", "application/octet-stream")
|
||||||
|
c.Header("Content-Length", strconv.FormatInt(size, 10))
|
||||||
|
if _, writeErr := io.Copy(c.Writer, spool); writeErr != nil {
|
||||||
|
return writeErr
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// validateDownloadFinal cross-checks the trailing NZTO frame against the
|
||||||
|
// payload the dashboard actually relayed:
|
||||||
|
// - magic must be NZTO (defence-in-depth; the relay already checked).
|
||||||
|
// - frame must be the full 44 bytes (magic 4 + size 8 + sha256 32).
|
||||||
|
// - declared size must match streamed byte count exactly.
|
||||||
|
// - declared sha256 must match the streamed sha256, with one allowed
|
||||||
|
// "explicit skip" form: all-zero declared hash means agent could not
|
||||||
|
// compute a hash and we accept the size-only check.
|
||||||
|
//
|
||||||
|
// Without this gate a truncated or wrong-hash NZTO is silently accepted
|
||||||
|
// and the dashboard serves possibly-corrupt bytes to the HTTP client.
|
||||||
|
func validateDownloadFinal(final []byte, streamedSize int64, streamedSHA256 []byte) error {
|
||||||
|
if len(final) < 4 || !bytes.Equal(final[:4], model.MCPFsXferMagicOK) {
|
||||||
|
return errors.New("download final frame: unexpected magic")
|
||||||
|
}
|
||||||
|
if len(final) < 4+8+32 {
|
||||||
|
return errors.New("download final frame: truncated header (need size + sha256)")
|
||||||
|
}
|
||||||
|
declaredSize := binary.BigEndian.Uint64(final[4:12])
|
||||||
|
if uint64(streamedSize) != declaredSize {
|
||||||
|
return errors.New("download final frame: declared size does not match streamed bytes")
|
||||||
|
}
|
||||||
|
declaredSHA := final[12:44]
|
||||||
|
allZero := true
|
||||||
|
for _, b := range declaredSHA {
|
||||||
|
if b != 0 {
|
||||||
|
allZero = false
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if allZero {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if len(streamedSHA256) < 32 {
|
||||||
|
return errors.New("download final frame: streamed sha256 too short to compare")
|
||||||
|
}
|
||||||
|
if !bytes.Equal(declaredSHA, streamedSHA256[:32]) {
|
||||||
|
return errors.New("download final frame: declared sha256 does not match streamed bytes")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// midstreamErrMsgCap 限制错误帧 payload 的累计读取量。错误消息只作 string
|
||||||
|
// 用,没有上限的话恶意/有 bug 的 agent 可在 NZTE 后持续发 256 字节块(且不
|
||||||
|
// 关流),让 dashboard goroutine 内存无界增长或永久阻塞。读满 cap 即停止。
|
||||||
|
const midstreamErrMsgCap = 8 << 10
|
||||||
|
|
||||||
|
func readMidstreamErrMsg(stream io.Reader, header []byte) string {
|
||||||
|
rest := make([]byte, 0, 256)
|
||||||
|
tail := make([]byte, 256)
|
||||||
|
for len(rest) < midstreamErrMsgCap {
|
||||||
|
n, err := stream.Read(tail)
|
||||||
|
if n > 0 {
|
||||||
|
room := midstreamErrMsgCap - len(rest)
|
||||||
|
if n > room {
|
||||||
|
n = room
|
||||||
|
}
|
||||||
|
rest = append(rest, tail[:n]...)
|
||||||
|
}
|
||||||
|
if err != nil || n < len(tail) {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return string(header[len(model.MCPFsXferMagicErr):]) + string(rest)
|
||||||
|
}
|
||||||
|
|
||||||
|
var errMCPMidstreamAbort = errors.New("mcp transfer: aborted mid-stream by agent")
|
||||||
|
|
||||||
|
// fsTransferXferHeader 是 dashboard 解析 NZTU/NZTD/NZTO/NZTE 后得到的统一
|
||||||
|
// 结构。Magic 与 model.MCPFsXferMagic* 对照判断帧类型。
|
||||||
|
type fsTransferXferHeader struct {
|
||||||
|
Magic []byte
|
||||||
|
Size int64
|
||||||
|
SHA256 []byte
|
||||||
|
ErrMsg string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *fsTransferXferHeader) IsErr() bool {
|
||||||
|
return bytes.Equal(h.Magic, model.MCPFsXferMagicErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
// readXferFixedHeader 读取一帧 IOStream 数据并解析。每条 agent 控制帧都在
|
||||||
|
// 单条 IOStreamData 内完整发送(agent 端用 stream.Send(buf) 整块写),所以
|
||||||
|
// 一次 8KiB 缓冲即可拿到完整帧;不需要跨帧拼接。
|
||||||
|
//
|
||||||
|
// 心跳帧(空 Data)由 agent 那侧的 ioStreamKeepAlive 周期性下发,io.Read
|
||||||
|
// 不会暴露空读,因此这里不必特殊跳过。
|
||||||
|
func readXferFixedHeader(stream io.Reader) (*fsTransferXferHeader, error) {
|
||||||
|
buf := make([]byte, 4+8+32+512)
|
||||||
|
n, err := stream.Read(buf)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return readXferFixedHeaderFromBytes(buf[:n])
|
||||||
|
}
|
||||||
|
|
||||||
|
// readXferFixedHeaderFromBytes parses a fully-received transfer control
|
||||||
|
// frame. Extracted so a malicious-input regression suite can pin the
|
||||||
|
// uint64→int64 overflow gate: raw u64 size > MCPFsTransferMaxSize or
|
||||||
|
// > MaxInt64 must be rejected BEFORE narrowing, otherwise the cap check
|
||||||
|
// later in the handler sees a wrapped negative value and lets the
|
||||||
|
// transfer through.
|
||||||
|
func readXferFixedHeaderFromBytes(raw []byte) (*fsTransferXferHeader, error) {
|
||||||
|
if len(raw) < 4 {
|
||||||
|
return nil, errors.New("frame too short")
|
||||||
|
}
|
||||||
|
magic := raw[:4]
|
||||||
|
out := &fsTransferXferHeader{Magic: append([]byte(nil), magic...)}
|
||||||
|
switch {
|
||||||
|
case bytes.Equal(magic, model.MCPFsXferMagicErr):
|
||||||
|
out.ErrMsg = string(raw[4:])
|
||||||
|
return out, nil
|
||||||
|
case bytes.Equal(magic, model.MCPFsXferMagicUploadHdr):
|
||||||
|
if len(raw) < 4+8 {
|
||||||
|
return nil, errors.New("upload header too short")
|
||||||
|
}
|
||||||
|
size, err := xferSizeFromU64(binary.BigEndian.Uint64(raw[4:12]))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
out.Size = size
|
||||||
|
return out, nil
|
||||||
|
case bytes.Equal(magic, model.MCPFsXferMagicDownloadHdr):
|
||||||
|
if len(raw) < 4+8+32 {
|
||||||
|
return nil, errors.New("download header too short")
|
||||||
|
}
|
||||||
|
size, err := xferSizeFromU64(binary.BigEndian.Uint64(raw[4:12]))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
out.Size = size
|
||||||
|
out.SHA256 = append([]byte(nil), raw[12:44]...)
|
||||||
|
return out, nil
|
||||||
|
case bytes.Equal(magic, model.MCPFsXferMagicOK):
|
||||||
|
if len(raw) < 4+8+32 {
|
||||||
|
return nil, errors.New("ok header too short")
|
||||||
|
}
|
||||||
|
size, err := xferSizeFromU64(binary.BigEndian.Uint64(raw[4:12]))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
out.Size = size
|
||||||
|
out.SHA256 = append([]byte(nil), raw[12:44]...)
|
||||||
|
return out, nil
|
||||||
|
default:
|
||||||
|
return nil, errors.New("unexpected frame magic")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// xferSizeFromU64 caps the raw u64 size carried by an NZTU/NZTD/NZTO frame
|
||||||
|
// at MCPFsTransferMaxSize AND math.MaxInt64. Both bounds matter: the cap
|
||||||
|
// keeps the protocol invariant, the MaxInt64 floor keeps later int64
|
||||||
|
// arithmetic safe even if MCPFsTransferMaxSize is ever raised above
|
||||||
|
// MaxInt64 by accident.
|
||||||
|
func xferSizeFromU64(raw uint64) (int64, error) {
|
||||||
|
if raw > uint64(model.MCPFsTransferMaxSize) {
|
||||||
|
return 0, errors.New("declared size exceeds MCP transfer cap")
|
||||||
|
}
|
||||||
|
if raw > math.MaxInt64 {
|
||||||
|
return 0, errors.New("declared size overflows int64")
|
||||||
|
}
|
||||||
|
return int64(raw), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// revalidateTransferEntry 在消费一次性 URL 时重新检查 mint 阶段的全部前置。
|
||||||
|
// 这是 mint→consume 之间发生权限变化(PAT 吊销、scope/whitelist 收紧、
|
||||||
|
// server 转手)时的兜底闸门:HMAC 签发与 sync.Map 一次性消费机制本身只能
|
||||||
|
// 防伪造与防重放,无法感知后端状态。
|
||||||
|
func revalidateTransferEntry(e *transferEntry) error {
|
||||||
|
if singleton.Conf == nil || !singleton.Conf.MCPEnabled() {
|
||||||
|
return errors.New("MCP is disabled by the dashboard administrator")
|
||||||
|
}
|
||||||
|
var tok model.APIToken
|
||||||
|
if err := singleton.DB.First(&tok, e.TokenID).Error; err != nil {
|
||||||
|
return errors.New("originating api token no longer exists")
|
||||||
|
}
|
||||||
|
// Bind the reloaded token back to the minting user. If the original PAT
|
||||||
|
// was deleted and its numeric primary key reused by a different user's
|
||||||
|
// token, the row would still load here; without this check the stale
|
||||||
|
// one-time URL would be revalidated against an unrelated token.
|
||||||
|
if tok.UserID != e.UserID {
|
||||||
|
return errors.New("originating api token no longer exists")
|
||||||
|
}
|
||||||
|
if tok.IsExpired(time.Now()) {
|
||||||
|
return errors.New("originating api token expired")
|
||||||
|
}
|
||||||
|
wantScope := model.ScopeServerRead
|
||||||
|
if e.Direction == transferDirUpload {
|
||||||
|
wantScope = model.ScopeServerWrite
|
||||||
|
}
|
||||||
|
if !tok.HasScope(wantScope) {
|
||||||
|
return errors.New("originating api token no longer has required scope")
|
||||||
|
}
|
||||||
|
if !tok.CanAccessServer(e.ServerID) {
|
||||||
|
return errors.New("originating api token no longer covers target server")
|
||||||
|
}
|
||||||
|
srv, _ := singleton.ServerShared.Get(e.ServerID)
|
||||||
|
if srv == nil {
|
||||||
|
return errors.New("target server no longer exists")
|
||||||
|
}
|
||||||
|
var user model.User
|
||||||
|
if err := singleton.DB.First(&user, e.UserID).Error; err != nil {
|
||||||
|
return errors.New("originating user no longer exists")
|
||||||
|
}
|
||||||
|
if user.Role != model.RoleAdmin && srv.GetUserID() != e.UserID {
|
||||||
|
return errors.New("target server is no longer owned by the originating user")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,72 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// transferAnonAuditThrottle caps the number of audit rows written per
|
||||||
|
// source IP within a sliding window for transfer requests that failed
|
||||||
|
// before a valid `entry` could be loaded (bogus/expired/replayed token).
|
||||||
|
//
|
||||||
|
// Without this cap an unauthenticated attacker can POST millions of
|
||||||
|
// /mcp/upload/<random> requests; every miss invokes
|
||||||
|
// writeTransferFailureAudit which inserts into mcp_audit_log. The
|
||||||
|
// throttle keeps a small per-IP token bucket in memory and drops audit
|
||||||
|
// rows past the budget — successful and authenticated failures (entry
|
||||||
|
// != nil) bypass this gate entirely so SIEM signal is unaffected.
|
||||||
|
type transferAnonAuditThrottle struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
window time.Duration
|
||||||
|
limit int
|
||||||
|
hits map[string]*anonHitBucket
|
||||||
|
clock func() time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
type anonHitBucket struct {
|
||||||
|
firstAt time.Time
|
||||||
|
count int
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTransferAnonAuditThrottle(window time.Duration, perWindow int) *transferAnonAuditThrottle {
|
||||||
|
return &transferAnonAuditThrottle{
|
||||||
|
window: window,
|
||||||
|
limit: perWindow,
|
||||||
|
hits: make(map[string]*anonHitBucket),
|
||||||
|
clock: time.Now,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// shouldRecord reports whether the anonymous failure for this IP should
|
||||||
|
// land in the audit table. Empty ip is treated as "always record" since
|
||||||
|
// suppressing it would silently lose signal in test/headless contexts.
|
||||||
|
func (t *transferAnonAuditThrottle) shouldRecord(ip string) bool {
|
||||||
|
if ip == "" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
t.mu.Lock()
|
||||||
|
defer t.mu.Unlock()
|
||||||
|
now := t.clock()
|
||||||
|
t.pruneLocked(now)
|
||||||
|
|
||||||
|
b, ok := t.hits[ip]
|
||||||
|
if !ok || now.Sub(b.firstAt) >= t.window {
|
||||||
|
t.hits[ip] = &anonHitBucket{firstAt: now, count: 1}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
if b.count >= t.limit {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
b.count++
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *transferAnonAuditThrottle) pruneLocked(now time.Time) {
|
||||||
|
for ip, b := range t.hits {
|
||||||
|
if now.Sub(b.firstAt) >= t.window {
|
||||||
|
delete(t.hits, ip)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var transferAnonAuditThrottleShared = newTransferAnonAuditThrottle(time.Minute, 5)
|
||||||
@@ -0,0 +1,68 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// H8 regression: anonymous transfer failures (entry=nil — token was bogus
|
||||||
|
// or already consumed) must be sampled, not written to audit one-for-one.
|
||||||
|
// Otherwise any unauthenticated attacker can flood the audit table by
|
||||||
|
// repeatedly POSTing /mcp/upload/garbage.
|
||||||
|
func TestTransferAnonAuditThrottle_FirstRequestPasses(t *testing.T) {
|
||||||
|
th := newTransferAnonAuditThrottle(10*time.Second, 5)
|
||||||
|
|
||||||
|
if !th.shouldRecord("1.2.3.4") {
|
||||||
|
t.Fatal("first anon failure from an IP must be recorded")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTransferAnonAuditThrottle_BurstCappedPerWindow(t *testing.T) {
|
||||||
|
th := newTransferAnonAuditThrottle(time.Minute, 3)
|
||||||
|
const ip = "5.6.7.8"
|
||||||
|
|
||||||
|
recorded := 0
|
||||||
|
for i := 0; i < 20; i++ {
|
||||||
|
if th.shouldRecord(ip) {
|
||||||
|
recorded++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if recorded > 3 {
|
||||||
|
t.Fatalf("burst of 20 anon failures must be capped at 3 per window, got %d", recorded)
|
||||||
|
}
|
||||||
|
if recorded == 0 {
|
||||||
|
t.Fatal("burst must record at least one sample")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTransferAnonAuditThrottle_IndependentPerIP(t *testing.T) {
|
||||||
|
th := newTransferAnonAuditThrottle(time.Minute, 1)
|
||||||
|
|
||||||
|
if !th.shouldRecord("a") {
|
||||||
|
t.Fatal("first request from IP a must be recorded")
|
||||||
|
}
|
||||||
|
if !th.shouldRecord("b") {
|
||||||
|
t.Fatal("first request from a different IP must not share IP a's budget")
|
||||||
|
}
|
||||||
|
if th.shouldRecord("a") {
|
||||||
|
t.Fatal("second request from IP a within window must be dropped")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTransferAnonAuditThrottle_WindowResets(t *testing.T) {
|
||||||
|
th := newTransferAnonAuditThrottle(20*time.Millisecond, 1)
|
||||||
|
const ip = "9.9.9.9"
|
||||||
|
|
||||||
|
if !th.shouldRecord(ip) {
|
||||||
|
t.Fatal("first request must be recorded")
|
||||||
|
}
|
||||||
|
if th.shouldRecord(ip) {
|
||||||
|
t.Fatal("second request inside window must be dropped")
|
||||||
|
}
|
||||||
|
|
||||||
|
time.Sleep(40 * time.Millisecond)
|
||||||
|
|
||||||
|
if !th.shouldRecord(ip) {
|
||||||
|
t.Fatal("request after window expiry must be recorded again")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,169 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/grpcx"
|
||||||
|
pb "github.com/nezhahq/nezha/proto"
|
||||||
|
"github.com/nezhahq/nezha/service/rpc"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newFakeAgentIO() *grpcx.IOStreamWrapper {
|
||||||
|
return grpcx.NewIOStreamWrapper(&fakeAgentStream{closed: make(chan struct{})})
|
||||||
|
}
|
||||||
|
|
||||||
|
// fakeAgentStream mimics an attached-but-silent agent: Recv blocks until the
|
||||||
|
// wrapper is closed, exactly the post-attach state where nothing watches the
|
||||||
|
// per-transfer context.
|
||||||
|
type fakeAgentStream struct {
|
||||||
|
closed chan struct{}
|
||||||
|
closeOnce sync.Once
|
||||||
|
closeSeen chan struct{}
|
||||||
|
recvDone chan struct{}
|
||||||
|
closeCall atomic.Int32
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeAgentStream) Recv() (*pb.IOStreamData, error) {
|
||||||
|
<-f.closed
|
||||||
|
close(f.recvDone)
|
||||||
|
return nil, context.Canceled
|
||||||
|
}
|
||||||
|
func (f *fakeAgentStream) Send(*pb.IOStreamData) error { return nil }
|
||||||
|
func (f *fakeAgentStream) Context() context.Context { return context.Background() }
|
||||||
|
|
||||||
|
func (f *fakeAgentStream) closeEndpoint() {
|
||||||
|
f.closeOnce.Do(func() { close(f.closeSeen) })
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeAgentStream) Close() error {
|
||||||
|
f.closeCall.Add(1)
|
||||||
|
f.closeEndpoint()
|
||||||
|
select {
|
||||||
|
case <-f.closed:
|
||||||
|
default:
|
||||||
|
close(f.closed)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// transferRevokableContext only cancels a context; the post-attach relay
|
||||||
|
// (readXferFixedHeader / relayDownloadFrames / io.CopyN) and IOStreamWrapper.Read
|
||||||
|
// do not watch it. A revoked PAT (or a disconnected HTTP client) must still
|
||||||
|
// tear down the attached stream, else a stalled/compromised agent pins a
|
||||||
|
// dashboard goroutine + IOStream until restart. openFsTransferStream must wire
|
||||||
|
// ctx cancellation to CloseStream.
|
||||||
|
func TestOpenFsTransferStream_CancelClosesAttachedStream(t *testing.T) {
|
||||||
|
cleanupMCP, _ := setupMCPTest(t)
|
||||||
|
defer cleanupMCP()
|
||||||
|
singleton.Conf.SetMCPEnabled(true)
|
||||||
|
|
||||||
|
originalHandler := rpc.NezhaHandlerSingleton
|
||||||
|
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
|
||||||
|
t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler })
|
||||||
|
|
||||||
|
stream := newKillSwitchStream()
|
||||||
|
sc := singleton.NewEmptyServerClassForTest()
|
||||||
|
srv := &model.Server{}
|
||||||
|
srv.ID = 7
|
||||||
|
srv.SetTaskStream(stream)
|
||||||
|
sc.InsertForTest(srv)
|
||||||
|
originalShared := singleton.ServerShared
|
||||||
|
singleton.ServerShared = sc
|
||||||
|
t.Cleanup(func() { singleton.ServerShared = originalShared })
|
||||||
|
|
||||||
|
streamIDCh := make(chan string, 1)
|
||||||
|
agentStreamCh := make(chan *fakeAgentStream, 1)
|
||||||
|
attachReady := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
task := <-stream.sent
|
||||||
|
var req model.FsTransferRequest
|
||||||
|
require.NoError(t, json.Unmarshal([]byte(task.GetData()), &req))
|
||||||
|
streamIDCh <- req.StreamID
|
||||||
|
fakeAgent := &fakeAgentStream{closed: make(chan struct{}), closeSeen: make(chan struct{}), recvDone: make(chan struct{})}
|
||||||
|
agentStreamCh <- fakeAgent
|
||||||
|
require.NoError(t, rpc.NezhaHandlerSingleton.AgentConnected(req.StreamID, grpcx.NewIOStreamWrapper(fakeAgent)))
|
||||||
|
close(attachReady)
|
||||||
|
}()
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
streamIO, cleanup, err := openFsTransferStream(ctx, 7, &model.FsTransferRequest{
|
||||||
|
Op: model.MCPFsTransferOpDownload,
|
||||||
|
Path: "/srv/file",
|
||||||
|
})
|
||||||
|
require.NoError(t, err, "agent must attach so openFsTransferStream returns a live stream")
|
||||||
|
require.NotNil(t, streamIO)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
streamID := <-streamIDCh
|
||||||
|
<-attachReady
|
||||||
|
_, getErr := rpc.NezhaHandlerSingleton.GetStream(streamID)
|
||||||
|
require.NoError(t, getErr, "stream must be live before cancel")
|
||||||
|
agentStream := <-agentStreamCh
|
||||||
|
readDone := make(chan error, 1)
|
||||||
|
go func() {
|
||||||
|
_, readErr := streamIO.Read(make([]byte, 16))
|
||||||
|
readDone <- readErr
|
||||||
|
}()
|
||||||
|
|
||||||
|
cancel()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-readDone:
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("cancelling the transfer context must unblock the actual streamIO.Read")
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-agentStream.closeSeen:
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("cancelling the transfer context must close the fake agent endpoint")
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-agentStream.closed:
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("cancelling the transfer context must close the handler endpoint")
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-agentStream.recvDone:
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("cancelling the transfer context must let the fake handler exit")
|
||||||
|
}
|
||||||
|
require.Equal(t, int32(1), agentStream.closeCall.Load(), "attached endpoint must be closed exactly once")
|
||||||
|
require.Equal(t, 0, rpc.NezhaHandlerSingleton.StreamCount())
|
||||||
|
for index := 0; index < 40; index++ {
|
||||||
|
require.NoError(t, rpc.NezhaHandlerSingleton.CreateStream(fmt.Sprintf("cancel-reuse-%d", index), 0, 7))
|
||||||
|
}
|
||||||
|
require.ErrorIs(t, rpc.NezhaHandlerSingleton.CreateStream("cancel-reuse-over", 0, 7), rpc.ErrTooManyStreamsForServer)
|
||||||
|
for index := 0; index < 40; index++ {
|
||||||
|
require.NoError(t, rpc.NezhaHandlerSingleton.CloseStream(fmt.Sprintf("cancel-reuse-%d", index)))
|
||||||
|
}
|
||||||
|
require.Equal(t, 0, rpc.NezhaHandlerSingleton.StreamCount())
|
||||||
|
|
||||||
|
cleanupDone := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for range 32 {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
cleanup()
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
wg.Wait()
|
||||||
|
close(cleanupDone)
|
||||||
|
}()
|
||||||
|
select {
|
||||||
|
case <-cleanupDone:
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatal("concurrent cleanup calls must complete")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,66 @@
|
|||||||
|
package controller
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestTransferConsume_RevokedTokenIsRejected(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupTransferTest(t, xferAgentDownloadSend([]byte("ok")))
|
||||||
|
defer cleanup()
|
||||||
|
url := mintDownloadURL(t, ts, tok, "/srv/file")
|
||||||
|
|
||||||
|
if err := singleton.DB.Where("token_hash = ?", model.HashAPIToken(tok)).
|
||||||
|
Delete(&model.APIToken{}).Error; err != nil {
|
||||||
|
t.Fatalf("revoke token: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := http.Get(url)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
require.Equalf(t, http.StatusUnauthorized, resp.StatusCode,
|
||||||
|
"download URL must return 401 after the originating PAT is revoked; body=%s", string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTransferConsume_NarrowedServerWhitelistIsRejected(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupTransferTest(t, xferAgentDownloadSend([]byte("ok")))
|
||||||
|
defer cleanup()
|
||||||
|
url := mintDownloadURL(t, ts, tok, "/srv/file")
|
||||||
|
|
||||||
|
var stored model.APIToken
|
||||||
|
require.NoError(t, singleton.DB.Where("token_hash = ?", model.HashAPIToken(tok)).
|
||||||
|
First(&stored).Error)
|
||||||
|
stored.SetServerIDs([]uint64{999})
|
||||||
|
require.NoError(t, singleton.DB.Save(&stored).Error)
|
||||||
|
|
||||||
|
resp, err := http.Get(url)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
require.Equalf(t, http.StatusUnauthorized, resp.StatusCode,
|
||||||
|
"download URL must return 401 after PAT server_ids no longer cover the target; body=%s", string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTransferConsume_ServerOwnershipChangeIsRejected(t *testing.T) {
|
||||||
|
ts, tok, cleanup := setupTransferTest(t, xferAgentDownloadSend([]byte("ok")))
|
||||||
|
defer cleanup()
|
||||||
|
url := mintDownloadURL(t, ts, tok, "/srv/file")
|
||||||
|
|
||||||
|
srv, _ := singleton.ServerShared.Get(7)
|
||||||
|
require.NotNil(t, srv)
|
||||||
|
srv.SetUserID(99999)
|
||||||
|
|
||||||
|
resp, err := http.Get(url)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
require.Equalf(t, http.StatusUnauthorized, resp.StatusCode,
|
||||||
|
"download URL must return 401 after server is transferred away from the minting user; body=%s", string(body))
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user