Merge branch 'upstream/master' into master and preserve domain extensions

This commit is contained in:
Bot
2026-08-31 02:46:42 +08:00
758 changed files with 89269 additions and 1366 deletions
+21
View File
@@ -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:
- "*"
+19 -14
View File
@@ -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
View File
@@ -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
View File
@@ -26,4 +26,6 @@
/cmd/dashboard/docs /cmd/dashboard/docs
/data/* /data/*
app app
dashboard dashboard
.omo/
+5 -6
View File
@@ -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
View File
@@ -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)
}
+22 -2
View File
@@ -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}")
}
+294
View File
@@ -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
}
+135
View File
@@ -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 → 跳过 JWTrestScopeMiddleware 会按 scope 收口。
// 3. 未带 PAT → 走 fallbackJwtMw,存在 JWT 则挂 user,没有就匿名继续。
//
// 这是修复 ForceAuth=false 时 optional 路由完全不解析 PAT 的关键:
// 之前直接用 fallbackAuthMw 会让 PAT 请求被当作 guestscope 形同虚设。
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")
}
+615
View File
@@ -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)
}
+180 -90
View File
@@ -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-JWTForceAuth=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 资源族 scoperead/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")
}
+49 -5
View File
@@ -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")
}
+149
View File
@@ -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)
}
}
+162
View File
@@ -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)")
}
}
+15 -2
View File
@@ -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
+30 -16
View File
@@ -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 404body 还是 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
View File
@@ -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")
}
+123
View File
@@ -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
}
+505
View File
@@ -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;闸 2PAT 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 discoveryJSON-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_idbest-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
}
+48
View File
@@ -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,不阻塞业务。
//
// argsBytestool 的 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 分支不回 TaskResultdashboard 要等 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
// 分支不回 TaskResultdashboard 必须在调 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 分支不回 TaskResultCallAgent 必须等 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 喂给前端 fallbackHTML/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)
}
}
+117
View File
@@ -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.1Host == 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 IP127.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
// 是 loopbackOrigin 是公网域名,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)",
})
}
+109
View File
@@ -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)
}
+98
View File
@@ -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 HTTPGET 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)
}
+306
View File
@@ -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)
}
+144
View File
@@ -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-runWindows 需 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")
}
+311
View File
@@ -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
}
// 闸 2PAT 的 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")
}
+974
View File
@@ -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 100MiBmodel.MCPFsTransferMaxSize)。
//
// 传输实现:dashboard ↔ agent 走 gRPC IOStream 双向流(TaskTypeFsTransfer),
// dashboard 一边读 HTTP body 一边推给 agent;不再使用 base64/JSON 包装内容,
// 避免 gRPC 4MiB 单消息上限。
//
// 安全机制:
// - 一次性 token,存内存 sync.MapTTL 默认 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)
// 校验 HMACtoken 形如 id.sigsig 必须等于 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 510min 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 100MiBmodel.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/:tokenbody 转发到 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-LengthHTTP 层会立即截断并报 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 阶段就失败时拿不到 entryUserID/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 的失败映射成
// outcomeoffline / 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