feat(auth): add PAT auth, scoped REST/MCP access, CSRF, and tenant isolation

Introduce Personal Access Tokens (nzp_*) as a stateless auth path alongside
JWT, gated per-endpoint by a scope middleware (nezha:{resource}:{verb}) with
fail-closed empty-scope defaults and a server-id whitelist. Self-management
endpoints (profile, api-tokens, oauth2 bind, refresh-token) explicitly reject
PATs to block privilege-escalation chains. A revoke registry tears down active
long-lived connections (terminal, fm, ws, transfer, mcp) the moment a PAT is
deleted, with a tombstone closing the revoke->register race.

Add an MCP endpoint that proxies tool calls (exec, fs read/write/delete,
transfer) to agents over gRPC, guarded by origin/DNS-rebinding checks, a
per-token rate limiter, audit logging, and a kill switch. Serialize all
sends through the IOStream wrapper to honour grpc-go's concurrency contract.

Add CSRF double-submit protection on unsafe cookie-authenticated methods,
exempting authenticated PAT requests by context identity (not a forgeable
Authorization header). Apply visibility/whitelist filtering consistently
across list, get-by-id, and mutate paths to enforce tenant isolation.

Migrate legacy mcp:* scopes: rewrite read/exec to nezha:* equivalents and
drop dangerous write/delete/wildcard grants.

Co-authored-by: cloudcode <cloudcode@users.noreply.github.com>
This commit is contained in:
naiba
2026-05-30 15:56:44 +00:00
co-authored by cloudcode
parent 029695344c
commit e8dabf5bc6
153 changed files with 16974 additions and 244 deletions
+9 -2
View File
@@ -1,7 +1,6 @@
package controller
import (
"maps"
"slices"
"strconv"
"time"
@@ -168,9 +167,14 @@ func batchDeleteAlertRule(c *gin.Context) (any, 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 {
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")
}
@@ -200,6 +204,9 @@ func validateRule(c *gin.Context, r *model.AlertRule) error {
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
@@ -0,0 +1,142 @@
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)
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.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}")
}
+286
View File
@@ -0,0 +1,286 @@
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" // gin context 里存 *model.APIToken 的 key
apiTokenLastUsedCtxKey = "nz_api_token_used_marker" // 标记是否需要异步更新 last_used
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)
if !callerIsAdmin(c) {
server, _ := singleton.ServerShared.Get(sid)
if server == nil {
return nil, errors.New("server not found")
}
if !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) {
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) {
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) {
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 {
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken)
abortAPITokenUnauthorized(c, "owner of api token not found")
return
}
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 v, ok := c.Get(apiTokenLastUsedCtxKey); !ok || v != true {
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.ScopeServerRead), 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.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,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.ScopeServerRead}, nil)
gin.SetMode(gin.TestMode)
r := gin.New()
pat := apiTokenAuthMiddleware()
r.GET("/api/v1/server",
pat,
restScopeMiddleware(model.ScopeServerRead),
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.ScopeServerRead),
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.ScopeServerDelete),
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.ScopeServerRead}, 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.ScopeServerRead),
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")
}
+597
View File
@@ -0,0 +1,597 @@
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)
require.NoError(t, db.AutoMigrate(&model.User{}, &model.APIToken{}, &model.Server{}))
singleton.DB = db
return func() {
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_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)
}
+107 -78
View File
@@ -44,6 +44,8 @@ func ServeWeb(frontendDist fs.FS) http.Handler {
routers(r, frontendDist)
kickoffTransferGC()
return r
}
@@ -55,6 +57,19 @@ func routers(r *gin.Engine, frontendDist fs.FS) {
if err := authMiddleware.MiddlewareInit(); err != nil {
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)
api := r.Group("api/v1")
api.POST("/login", authMiddleware.LoginHandler)
api.GET("/oauth2/:provider", commonHandler(oauth2redirect))
@@ -64,99 +79,112 @@ func routers(r *gin.Engine, frontendDist fs.FS) {
fallbackAuth.GET("/setting", commonHandler(listConfig))
fallbackAuth.GET("/oauth2/callback", commonHandler(oauth2callback(authMiddleware)))
authMw := authMiddleware.MiddlewareFunc()
optionalAuthMw := utils.IfOr(singleton.Conf.ForceAuth, authMw, fallbackAuthMw)
jwtMw := authMiddleware.MiddlewareFunc()
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.GET("/ws/server", commonHandler(serverStream))
optionalAuth.GET("/server-group", commonHandler(listServerGroup))
optionalAuth.GET("/ws/server", restScopeMiddleware(model.ScopeServerRead), commonHandler(serverStream))
optionalAuth.GET("/server-group", restScopeMiddleware(model.ScopeServerRead), commonHandler(listServerGroup))
optionalAuth.GET("/service", commonHandler(showService))
optionalAuth.GET("/service/server", commonHandler(listServerWithServices))
optionalAuth.GET("/service/:id/history", commonHandler(getServiceHistory))
optionalAuth.GET("/server/:id/service", commonHandler(listServerServices))
optionalAuth.GET("/server/:id/metrics", commonHandler(getServerMetrics))
optionalAuth.GET("/service", restScopeMiddleware(model.ScopeServiceRead), commonHandler(showService))
optionalAuth.GET("/service/server", restScopeMiddleware(model.ScopeServiceRead), commonHandler(listServerWithServices))
optionalAuth.GET("/service/:id/history", restScopeMiddleware(model.ScopeServiceRead), commonHandler(getServiceHistory))
optionalAuth.GET("/server/:id/service", restScopeMiddleware(model.ScopeServiceRead), commonHandler(listServerServices))
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.POST("/terminal", commonHandler(createTerminal))
auth.GET("/ws/terminal/:id", commonHandler(terminalStream))
// server / terminal / fm / transfer 共享 nezha:server:* 资源族
auth.POST("/terminal", restScopeMiddleware(model.ScopeServerExec), commonHandler(createTerminal))
auth.GET("/ws/terminal/:id", restScopeMiddleware(model.ScopeServerExec), commonHandler(terminalStream))
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.ScopeServerRead), 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.ScopeServerDelete), 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.ScopeServerDelete), commonHandler(batchDeleteServerGroup))
auth.POST("/file", commonHandler(createFM))
auth.GET("/ws/file/:id", commonHandler(fmStream))
// transfer — 严格使用 nezha:transfer 资源族 scoperead/write/delete)。
// 注意:曾经计划让 nezha:server:read 兼听只读 transfer,但 restScopeMiddleware
// / APIToken.HasScope 不做 server↔transfer 别名展开,前端 SCOPE_OPTIONS 也已经
// 单独暴露 nezha:transfer:read,所以这里维持精确匹配语义。
auth.GET("/transfer", restScopeMiddleware(model.ScopeTransferRead), listHandler(listServerTransfer))
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("/profile", commonHandler(getProfile))
auth.POST("/profile", commonHandler(updateProfile))
auth.POST("/oauth2/:provider/unbind", commonHandler(unbindOauth2))
// service monitor
auth.GET("/service/list", restScopeMiddleware(model.ScopeServiceRead), listHandler(listService))
auth.POST("/service", restScopeMiddleware(model.ScopeServiceWrite), commonHandler(createService))
auth.PATCH("/service/:id", restScopeMiddleware(model.ScopeServiceWrite), commonHandler(updateService))
auth.POST("/batch-delete/service", restScopeMiddleware(model.ScopeServiceDelete), commonHandler(batchDeleteService))
auth.GET("/user", adminHandler(listUser))
auth.POST("/user", adminHandler(createUser))
auth.POST("/batch-delete/user", adminHandler(batchDeleteUser))
auth.GET("/notification-group", restScopeMiddleware(model.ScopeNotificationGroupRead), commonHandler(listNotificationGroup))
auth.POST("/notification-group", restScopeMiddleware(model.ScopeNotificationGroupWrite), commonHandler(createNotificationGroup))
auth.PATCH("/notification-group/:id", restScopeMiddleware(model.ScopeNotificationGroupWrite), commonHandler(updateNotificationGroup))
auth.POST("/batch-delete/notification-group", restScopeMiddleware(model.ScopeNotificationGroupDelete), commonHandler(batchDeleteNotificationGroup))
auth.GET("/service/list", listHandler(listService))
auth.POST("/service", commonHandler(createService))
auth.PATCH("/service/:id", commonHandler(updateService))
auth.POST("/batch-delete/service", commonHandler(batchDeleteService))
auth.GET("/notification", restScopeMiddleware(model.ScopeNotificationRead), listHandler(listNotification))
auth.POST("/notification", restScopeMiddleware(model.ScopeNotificationWrite), commonHandler(createNotification))
auth.PATCH("/notification/:id", restScopeMiddleware(model.ScopeNotificationWrite), commonHandler(updateNotification))
auth.POST("/batch-delete/notification", restScopeMiddleware(model.ScopeNotificationDelete), commonHandler(batchDeleteNotification))
auth.POST("/server-group", commonHandler(createServerGroup))
auth.PATCH("/server-group/:id", commonHandler(updateServerGroup))
auth.POST("/batch-delete/server-group", commonHandler(batchDeleteServerGroup))
auth.GET("/alert-rule", restScopeMiddleware(model.ScopeAlertRuleRead), listHandler(listAlertRule))
auth.POST("/alert-rule", restScopeMiddleware(model.ScopeAlertRuleWrite), commonHandler(createAlertRule))
auth.PATCH("/alert-rule/:id", restScopeMiddleware(model.ScopeAlertRuleWrite), commonHandler(updateAlertRule))
auth.POST("/batch-delete/alert-rule", restScopeMiddleware(model.ScopeAlertRuleDelete), commonHandler(batchDeleteAlertRule))
auth.GET("/notification-group", commonHandler(listNotificationGroup))
auth.POST("/notification-group", commonHandler(createNotificationGroup))
auth.PATCH("/notification-group/:id", commonHandler(updateNotificationGroup))
auth.POST("/batch-delete/notification-group", commonHandler(batchDeleteNotificationGroup))
auth.GET("/cron", restScopeMiddleware(model.ScopeCronRead), listHandler(listCron))
auth.POST("/cron", restScopeMiddleware(model.ScopeCronWrite), commonHandler(createCron))
auth.PATCH("/cron/:id", restScopeMiddleware(model.ScopeCronWrite), commonHandler(updateCron))
auth.POST("/cron/:id/manual", restScopeMiddleware(model.ScopeCronExec), commonHandler(manualTriggerCron))
auth.POST("/batch-delete/cron", restScopeMiddleware(model.ScopeCronDelete), commonHandler(batchDeleteCron))
auth.GET("/server", listHandler(listServer))
auth.PATCH("/server/:id", commonHandler(updateServer))
auth.GET("/server/config/:id", commonHandler(getServerConfig))
auth.POST("/server/config", commonHandler(setServerConfig))
auth.POST("/batch-delete/server", commonHandler(batchDeleteServer))
auth.POST("/batch-move/server", commonHandler(batchMoveServer))
auth.POST("/force-update/server", commonHandler(forceUpdateServer))
auth.GET("/ddns", restScopeMiddleware(model.ScopeDDNSRead), listHandler(listDDNS))
auth.GET("/ddns/providers", restScopeMiddleware(model.ScopeDDNSRead), commonHandler(listProviders))
auth.POST("/ddns", restScopeMiddleware(model.ScopeDDNSWrite), commonHandler(createDDNS))
auth.PATCH("/ddns/:id", restScopeMiddleware(model.ScopeDDNSWrite), commonHandler(updateDDNS))
auth.POST("/batch-delete/ddns", restScopeMiddleware(model.ScopeDDNSDelete), commonHandler(batchDeleteDDNS))
auth.GET("/transfer", listHandler(listServerTransfer))
auth.POST("/transfer/:id/cancel", commonHandler(cancelServerTransfer))
auth.POST("/transfer/:id/retry", commonHandler(retryServerTransfer))
auth.GET("/ws/transfer", commonHandler(transferStream))
auth.GET("/nat", restScopeMiddleware(model.ScopeNATRead), listHandler(listNAT))
auth.POST("/nat", restScopeMiddleware(model.ScopeNATWrite), commonHandler(createNAT))
auth.PATCH("/nat/:id", restScopeMiddleware(model.ScopeNATWrite), commonHandler(updateNAT))
auth.POST("/batch-delete/nat", restScopeMiddleware(model.ScopeNATDelete), commonHandler(batchDeleteNAT))
auth.GET("/notification", listHandler(listNotification))
auth.POST("/notification", commonHandler(createNotification))
auth.PATCH("/notification/:id", commonHandler(updateNotification))
auth.POST("/batch-delete/notification", commonHandler(batchDeleteNotification))
auth.GET("/alert-rule", listHandler(listAlertRule))
auth.POST("/alert-rule", commonHandler(createAlertRule))
auth.PATCH("/alert-rule/:id", commonHandler(updateAlertRule))
auth.POST("/batch-delete/alert-rule", commonHandler(batchDeleteAlertRule))
auth.GET("/cron", listHandler(listCron))
auth.POST("/cron", commonHandler(createCron))
auth.PATCH("/cron/:id", commonHandler(updateCron))
auth.POST("/cron/:id/manual", commonHandler(manualTriggerCron))
auth.POST("/batch-delete/cron", commonHandler(batchDeleteCron))
auth.GET("/ddns", listHandler(listDDNS))
auth.GET("/ddns/providers", commonHandler(listProviders))
auth.POST("/ddns", commonHandler(createDDNS))
auth.PATCH("/ddns/:id", commonHandler(updateDDNS))
auth.POST("/batch-delete/ddns", commonHandler(batchDeleteDDNS))
auth.GET("/nat", listHandler(listNAT))
auth.POST("/nat", commonHandler(createNAT))
auth.PATCH("/nat/:id", commonHandler(updateNAT))
auth.POST("/batch-delete/nat", commonHandler(batchDeleteNAT))
auth.GET("/waf", pAdminHandler(listBlockedAddress))
auth.POST("/batch-delete/waf", adminHandler(batchDeleteBlockedAddress))
auth.GET("/online-user", pAdminHandler(listOnlineUser))
auth.POST("/online-user/batch-block", adminHandler(batchBlockOnlineUser))
auth.PATCH("/setting", adminHandler(updateConfig))
auth.POST("/maintenance", adminHandler(runMaintenance))
// 管理员资源 — 仅 nezha:* / nezha:admin:* 持有者可调(adminHandler 进一步校验 user.Role)。
auth.GET("/user", restScopeMiddleware(model.ScopeAdminAll), adminHandler(listUser))
auth.POST("/user", restScopeMiddleware(model.ScopeAdminAll), adminHandler(createUser))
auth.POST("/batch-delete/user", restScopeMiddleware(model.ScopeAdminAll), adminHandler(batchDeleteUser))
auth.GET("/waf", restScopeMiddleware(model.ScopeAdminAll), pAdminHandler(listBlockedAddress))
auth.POST("/batch-delete/waf", restScopeMiddleware(model.ScopeAdminAll), adminHandler(batchDeleteBlockedAddress))
auth.GET("/online-user", restScopeMiddleware(model.ScopeAdminAll), pAdminHandler(listOnlineUser))
auth.POST("/online-user/batch-block", restScopeMiddleware(model.ScopeAdminAll), adminHandler(batchBlockOnlineUser))
auth.PATCH("/setting", restScopeMiddleware(model.ScopeAdminAll), adminHandler(updateConfig))
auth.POST("/maintenance", restScopeMiddleware(model.ScopeAdminAll), adminHandler(runMaintenance))
r.NoRoute(fallbackToFrontend(frontendDist))
}
@@ -390,6 +418,7 @@ func fallbackToFrontend(frontendDist fs.FS) func(*gin.Context) {
regexp.MustCompile(`^/dashboard/settings/user$`),
regexp.MustCompile(`^/dashboard/settings/online-user$`),
regexp.MustCompile(`^/dashboard/settings/waf$`),
regexp.MustCompile(`^/dashboard/settings/api-tokens$`),
// 注意:这里的白名单决定哪些 URL 走 index.html fallback;漏一条就会把
// 直接刷新该页面变成 404(HTTP 状态码层面,body 仍是 index.html,所以
// 浏览器内 SPA 看起来正常,但 monitoring / 链接预览会以为站点挂了)。
+44 -7
View File
@@ -50,10 +50,18 @@ func createCron(c *gin.Context) (uint64, error) {
return 0, err
}
if !singleton.ServerShared.CheckPermission(c, slices.Values(cf.Servers)) {
if !isValidCronCover(cf.Cover) {
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
}
@@ -111,12 +119,8 @@ func updateCron(c *gin.Context) (any, error) {
return 0, err
}
if !singleton.ServerShared.CheckPermission(c, slices.Values(cf.Servers)) {
return 0, singleton.Localizer.ErrorT("permission denied")
}
if err := assertOwnsNotificationGroup(c, cf.NotificationGroupID); err != nil {
return nil, err
if !isValidCronCover(cf.Cover) {
return nil, singleton.Localizer.ErrorT("permission denied")
}
var cr model.Cron
@@ -128,6 +132,18 @@ func updateCron(c *gin.Context) (any, error) {
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.Name = cf.Name
cr.Scheduler = cf.Scheduler
@@ -183,6 +199,14 @@ func manualTriggerCron(c *gin.Context) (any, error) {
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)
return nil, nil
}
@@ -208,6 +232,19 @@ func batchDeleteCron(c *gin.Context) (any, error) {
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 {
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,220 @@
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)
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)
singleton.CronShared = singleton.NewCronClass()
// 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.UserLock.Lock()
singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}}
singleton.UserLock.Unlock()
t.Cleanup(func() {
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,161 @@
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)
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.CronShared = singleton.NewCronClass()
singleton.ServerShared = singleton.NewServerClass()
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.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,342 @@
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
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
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.CronShared = singleton.NewCronClass()
originalSentinel := singleton.ServiceSentinelShared
sentinel, err := singleton.NewServiceSentinel(make(chan *model.Service, 4))
require.NoError(t, err)
singleton.ServiceSentinelShared = sentinel
t.Cleanup(func() {
sentinel.Close()
singleton.ServiceSentinelShared = originalSentinel
})
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.UserLock.Lock()
singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}}
singleton.UserLock.Unlock()
t.Cleanup(func() {
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 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")
}
+89
View File
@@ -0,0 +1,89 @@
package controller
import (
"crypto/rand"
"encoding/hex"
"net/http"
"github.com/gin-gonic/gin"
"github.com/nezhahq/nezha/model"
)
// setCSRFCookie issues a fresh 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) {
var b [32]byte
if _, err := rand.Read(b[:]); err != nil {
return
}
c.SetSameSite(http.SameSiteStrictMode)
c.SetCookie(csrfCookieName, hex.EncodeToString(b[:]), 0, "/", "", false, 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.
//
// 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)
if err != nil || cookie == "" || header == "" || header != cookie {
c.AbortWithStatusJSON(http.StatusForbidden, model.CommonResponse[any]{
Success: false,
Error: "ApiErrorForbidden: missing or invalid CSRF token",
})
return
}
c.Next()
}
}
+159
View File
@@ -0,0 +1,159 @@
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) {
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) {
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: "matching-token"})
c.Request.Header.Set("X-CSRF-Token", "matching-token")
mw(c)
if c.IsAborted() {
t.Fatalf("%s with matching 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)")
}
}
+6 -4
View File
@@ -36,8 +36,7 @@ func createFM(c *gin.Context) (*model.CreateFMResponse, error) {
if server == nil {
return nil, singleton.Localizer.ErrorT("server not found or not connected")
}
stream := server.GetTaskStream()
if stream == nil {
if server.GetTaskStream() == nil {
return nil, singleton.Localizer.ErrorT("server not found or not connected")
}
@@ -55,7 +54,7 @@ func createFM(c *gin.Context) (*model.CreateFMResponse, error) {
fmData, _ := json.Marshal(&model.TaskFM{
StreamID: streamId,
})
if err := stream.Send(&proto.Task{
if err := server.SendTask(&proto.Task{
Type: model.TaskTypeFM,
Data: string(fmData),
}); err != nil {
@@ -79,7 +78,7 @@ func fmStream(c *gin.Context) (any, error) {
// 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 !rpc.NezhaHandlerSingleton.IsStreamAuthorizedForUser(streamId, getUid(c), callerIsAdmin(c)) {
if !streamAttachAllowedForRequest(c, streamId) {
return nil, singleton.Localizer.ErrorT("permission denied")
}
if _, err := rpc.NezhaHandlerSingleton.GetStream(streamId); err != nil {
@@ -94,6 +93,9 @@ func fmStream(c *gin.Context) (any, error) {
defer wsConn.Close()
conn := websocketx.NewConn(wsConn)
deregisterPAT := registerPATConnection(c, func() { _ = wsConn.Close() })
defer deregisterPAT()
go func() {
// PING 保活
for {
@@ -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())
}
}
+11 -6
View File
@@ -35,7 +35,10 @@ func issueJWTSession(c *gin.Context, user *model.User, jwtTimeoutHours int) (map
if err != nil {
return nil, err
}
hashUID, err := idcodec.Encode(user.ID)
// 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
}
@@ -54,7 +57,7 @@ func issueJWTSession(c *gin.Context, user *model.User, jwtTimeoutHours int) (map
return nil, err
}
return map[string]interface{}{
jwtClaimUserID: hashUID,
jwtClaimUserID: encodedUID,
jwtClaimKeyID: keyID,
}, nil
}
@@ -91,6 +94,7 @@ func initParams() *jwt.GinJWTMiddleware {
TimeFunc: time.Now,
LoginResponse: func(c *gin.Context, code int, token string, expire time.Time) {
setCSRFCookie(c)
c.JSON(http.StatusOK, model.CommonResponse[model.LoginResponse]{
Success: true,
Data: model.LoginResponse{
@@ -120,11 +124,11 @@ func identityHandler() func(c *gin.Context) any {
if !ok || keyID == "" {
return nil
}
hashUID, ok := claims[jwtClaimUserID].(string)
if !ok || hashUID == "" {
encodedUID, ok := claims[jwtClaimUserID].(string)
if !ok || encodedUID == "" {
return nil
}
claimUID, err := idcodec.Decode(hashUID)
claimUID, err := idcodec.Decode(encodedUID)
if err != nil {
realIP := c.GetString(model.CtxKeyRealIPStr)
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken)
@@ -237,7 +241,7 @@ func unauthorized() func(c *gin.Context, code int, message string) {
// @Tags auth required
// @Produce json
// @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) {
if keyID := c.GetString(jwtClaimKeyID); keyID != "" {
_ = singleton.DB.Model(&model.JWTSession{}).
@@ -247,6 +251,7 @@ func refreshResponse(c *gin.Context, code int, token string, expire time.Time) {
"last_used_at": time.Now(),
}).Error
}
setCSRFCookie(c)
c.JSON(http.StatusOK, model.CommonResponse[model.LoginResponse]{
Success: true,
Data: model.LoginResponse{
@@ -244,3 +244,120 @@ func TestAuthenticatorPersistsCurrentTokenVersion(t *testing.T) {
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")
}
+493
View File
@@ -0,0 +1,493 @@
// 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"`
}
// 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
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,
})
}
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) {
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 {
textPayload := "{}"
if result != nil {
if b, err := json.Marshal(result); err == nil {
textPayload = string(b)
}
}
writeJSONRPCResult(c, req.ID, mcpToolCallResult{
Content: []mcpContent{{Type: "text", Text: textPayload}},
StructuredContent: result,
})
return
}
writeJSONRPCResult(c, req.ID, mcpToolCallResult{
Content: []mcpContent{{Type: "text", Text: errMsg}},
IsError: true,
StructuredContent: map[string]string{
"error_code": errCode,
"error": errMsg,
},
})
}
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)
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,74 @@
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 || server.Host == nil {
return nil
}
if compareSemver(server.Host.Version, MCPMinAgentVersion) < 0 {
return errMCPUnsupported
}
return nil
}
// compareSemver 比较两个 "MAJOR.MINOR.PATCH[-suffix]" 字符串。
// 返回 -1/0/1。无法解析时按字符串字典序比较,保证全序但可能不精确——
// 对于 "agent 太老" 的快速失败用途已经足够。
func compareSemver(a, b string) int {
if a == b {
return 0
}
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
}
}
if a < b {
return -1
}
if a > b {
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,52 @@
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")
}
@@ -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,432 @@
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.ScopeServerRead,
model.ScopeServerExec,
model.ScopeServerRead,
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,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)
}
+90
View File
@@ -0,0 +1,90 @@
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
}
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 &MCPRateLimiter{
perToken: make(map[uint64]*tokenWindow),
secLimit: secLimit,
minLimit: minLimit,
}
}
// 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
}
now := time.Now()
r.mu.Lock()
defer r.mu.Unlock()
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,68 @@
package controller
import (
"bytes"
"encoding/json"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"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) {
t.Helper()
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
raw, _ := json.Marshal(body)
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
}
// 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 := mcpEndpointTestCtx(t, tok, body)
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,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,74 @@
package controller
import (
"bytes"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"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, _ := gin.CreateTestContext(w)
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, _ := mcpEndpointRawCtx(t, tok, raw)
mcpEndpoint(c)
}
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, _ := mcpEndpointRawCtx(t, tok, raw)
mcpEndpoint(c)
}
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,195 @@
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.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)
}
+303
View File
@@ -0,0 +1,303 @@
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)
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() {
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)
}
+117
View File
@@ -0,0 +1,117 @@
package controller
import (
"encoding/json"
"errors"
"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 影响
//
// 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.",
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"},
},
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, errors.New(res.Error)
}
return res, nil
}
// 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,91 @@
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")
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")
}
+248
View File
@@ -0,0 +1,248 @@
package controller
import (
"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 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"},
},
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"},
},
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"},
},
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"},
},
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")
}
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")
}
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,49 @@
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{},
},
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,148 @@
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"`
}
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.",
},
},
},
RequiredScope: model.ScopeServerRead,
Handler: handleServerList,
})
registerMCPTool(&mcpTool{
Name: "server.get",
Description: "Return full Host/State snapshot for a single server.",
InputSchema: serverGetSchema(),
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
}
online := !s.LastActive.IsZero() && now.Sub(s.LastActive) < onlineWindow
if args.OnlineOnly && !online {
continue
}
item := serverListItem{
ID: s.ID,
Name: s.Name,
UUID: s.UUID,
Online: online,
LastActive: s.LastActive,
}
if s.Host != nil {
item.Platform = s.Host.Platform
item.Arch = s.Host.Arch
}
if s.GeoIP != nil {
item.IPv4 = s.GeoIP.IP.IPv4Addr
item.IPv6 = s.GeoIP.IP.IPv6Addr
}
out = append(out, item)
}
return 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
}
return map[string]any{
"id": s.ID,
"name": s.Name,
"uuid": s.UUID,
"note": s.Note,
"public_note": s.PublicNote,
"host": s.Host,
"state": s.State,
"geoip": s.GeoIP,
"last_active": s.LastActive,
}, nil
}
@@ -0,0 +1,113 @@
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.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.list", Arguments: json.RawMessage("{}")}),
})
mcpEndpoint(c)
_, tcr := decodeRPC(w)
require.False(t, tcr.IsError)
rb, _ := json.Marshal(tcr.StructuredContent)
var rows []map[string]any
require.NoError(t, json.Unmarshal(rb, &rows))
require.Len(t, rows, 1, "must filter out non-owned server")
require.EqualValues(t, 7, rows[0]["id"])
}
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.ScopeServerRead}, []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)
rb, _ := json.Marshal(tcr.StructuredContent)
var rows []map[string]any
require.NoError(t, json.Unmarshal(rb, &rows))
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.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.list", Arguments: jsonRaw(map[string]any{"online_only": true})}),
})
mcpEndpoint(c)
_, tcr := decodeRPC(w)
require.False(t, tcr.IsError)
rb, _ := json.Marshal(tcr.StructuredContent)
var rows []map[string]any
require.NoError(t, json.Unmarshal(rb, &rows))
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")
}
File diff suppressed because it is too large Load Diff
@@ -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,96 @@
package controller
import (
"context"
"encoding/json"
"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{}
}
func (f *fakeAgentStream) Recv() (*pb.IOStreamData, error) {
<-f.closed
return nil, context.Canceled
}
func (f *fakeAgentStream) Send(*pb.IOStreamData) error { return nil }
func (f *fakeAgentStream) Context() context.Context { return context.Background() }
// 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)
go func() {
task := <-stream.sent
var req model.FsTransferRequest
_ = json.Unmarshal([]byte(task.GetData()), &req)
streamIDCh <- req.StreamID
deadline := time.Now().Add(2 * time.Second)
for time.Now().Before(deadline) {
if err := rpc.NezhaHandlerSingleton.AgentConnected(req.StreamID, newFakeAgentIO()); err == nil {
return
}
time.Sleep(5 * time.Millisecond)
}
}()
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
_, getErr := rpc.NezhaHandlerSingleton.GetStream(streamID)
require.NoError(t, getErr, "stream must be live before cancel")
cancel()
require.Eventually(t, func() bool {
_, e := rpc.NezhaHandlerSingleton.GetStream(streamID)
return e != nil
}, 2*time.Second, 10*time.Millisecond,
"cancelling the transfer context must tear down the attached IOStream")
}
@@ -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))
}
@@ -0,0 +1,590 @@
package controller
import (
"bytes"
"context"
"encoding/binary"
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
"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"
)
// xferAgentSim 模拟 agent 在收到 TaskTypeFsTransfer 后的整个 IOStream 行为:
// - 通过 net.Pipe 拿到一个 in-memory 全双工流;
// - 把 dashboard 侧那一端塞进 rpc.NezhaHandlerSingleton.AgentConnected
// - 在 agent 侧 goroutine 里跑 upload/download 的协议帧逻辑。
//
// 该函数把"如果是真 agent 会做什么"全部就地展开,使 dashboard 端 transfer
// handler 能在没有真实 gRPC 链路的情况下完整跑过:mint→consume→stream→OK。
type xferStreamMux struct {
agent func(req *model.FsTransferRequest, dashboardSide io.ReadWriteCloser) ([]byte, error)
}
func (m *xferStreamMux) Send(t *pb.Task) error {
if t.GetType() != model.TaskTypeFsTransfer {
return nil
}
var req model.FsTransferRequest
if err := json.Unmarshal([]byte(t.GetData()), &req); err != nil {
return err
}
dashboardSide, agentSide := newFramedPipe()
if err := rpc.NezhaHandlerSingleton.AgentConnected(req.StreamID, dashboardSide); err != nil {
return err
}
go func() {
defer agentSide.Close()
_, _ = m.agent(&req, agentSide)
}()
return nil
}
// framedPipe is a frame-preserving full-duplex in-memory stream pair used by
// the MCP transfer tests in place of net.Pipe. Each Write on one side becomes
// exactly one frame on the other side, so RecvFrame on dashboardSide observes
// the same frame boundaries production code sees via grpcx.IOStreamWrapper.
// net.Pipe coalesces bytes and would let an NZTE control frame's bytes spill
// into a previous data frame's parse — the very bug we are testing for.
type framedPipe struct {
in chan []byte
out chan []byte
closed chan struct{}
once *sync.Once
rest []byte
}
func newFramedPipe() (*framedPipe, *framedPipe) {
closeCh := make(chan struct{})
once := new(sync.Once)
a := make(chan []byte, 64)
b := make(chan []byte, 64)
return &framedPipe{in: a, out: b, closed: closeCh, once: once},
&framedPipe{in: b, out: a, closed: closeCh, once: once}
}
func (p *framedPipe) Write(buf []byte) (int, error) {
frame := append([]byte(nil), buf...)
select {
case p.out <- frame:
return len(buf), nil
case <-p.closed:
return 0, io.ErrClosedPipe
}
}
func (p *framedPipe) Read(buf []byte) (int, error) {
if len(p.rest) > 0 {
n := copy(buf, p.rest)
p.rest = p.rest[n:]
return n, nil
}
select {
case frame, ok := <-p.in:
if !ok {
return 0, io.EOF
}
n := copy(buf, frame)
if n < len(frame) {
p.rest = frame[n:]
}
return n, nil
default:
}
select {
case frame, ok := <-p.in:
if !ok {
return 0, io.EOF
}
n := copy(buf, frame)
if n < len(frame) {
p.rest = frame[n:]
}
return n, nil
case <-p.closed:
return 0, io.EOF
}
}
func (p *framedPipe) RecvFrame() ([]byte, error) {
if len(p.rest) > 0 {
out := p.rest
p.rest = nil
return out, nil
}
select {
case frame, ok := <-p.in:
if !ok {
return nil, io.EOF
}
return frame, nil
default:
}
select {
case frame, ok := <-p.in:
if !ok {
return nil, io.EOF
}
return frame, nil
case <-p.closed:
return nil, io.EOF
}
}
func (p *framedPipe) Close() error {
p.once.Do(func() { close(p.closed) })
return nil
}
func (m *xferStreamMux) Recv() (*pb.TaskResult, error) { return nil, context.Canceled }
func (m *xferStreamMux) SetHeader(metadata.MD) error { return nil }
func (m *xferStreamMux) SendHeader(metadata.MD) error { return nil }
func (m *xferStreamMux) SetTrailer(metadata.MD) {}
func (m *xferStreamMux) Context() context.Context { return context.Background() }
func (m *xferStreamMux) SendMsg(any) error { return nil }
func (m *xferStreamMux) RecvMsg(any) error { return context.Canceled }
// xferAgentUploadAccept 实现 NZTU + 接收 size 字节 + NZTO 的完整握手。把读到
// 的原始字节作为返回值,方便测试断言。
func xferAgentUploadAccept(req *model.FsTransferRequest, stream io.ReadWriteCloser) ([]byte, error) {
got, err := xferAgentUploadRead(req, stream)
if err != nil {
return got, err
}
return got, xferAgentUploadAck(stream, uint64(len(got)))
}
func xferAgentUploadRead(req *model.FsTransferRequest, stream io.ReadWriteCloser) ([]byte, error) {
header := append([]byte(nil), model.MCPFsXferMagicUploadHdr...)
sz := make([]byte, 8)
binary.BigEndian.PutUint64(sz, uint64(req.Size))
header = append(header, sz...)
if _, err := stream.Write(header); err != nil {
return nil, err
}
got := make([]byte, 0, req.Size)
if req.Size > 0 {
buf := make([]byte, req.Size)
if _, err := io.ReadFull(stream, buf); err != nil {
return nil, err
}
got = buf
}
return got, nil
}
func xferAgentUploadAck(stream io.ReadWriteCloser, size uint64) error {
ok := append([]byte(nil), model.MCPFsXferMagicOK...)
finalSize := make([]byte, 8)
binary.BigEndian.PutUint64(finalSize, size)
ok = append(ok, finalSize...)
ok = append(ok, make([]byte, 32)...)
_, err := stream.Write(ok)
return err
}
// xferAgentDownloadSend 模拟 agent 向 dashboard 推 payload:发 NZTD、NZTC(chunk)
// 包装的 payload、最后 NZTO。NZTC 包装是 dashboard 区分数据帧与控制帧
// (NZTE/NZTO)的唯一依据;离开它后 dashboard 没办法把首字节恰好等于 NZTE 的
// 合法文件内容与真错误帧区分开。
func xferAgentDownloadSend(payload []byte) func(*model.FsTransferRequest, io.ReadWriteCloser) ([]byte, error) {
return func(_ *model.FsTransferRequest, stream io.ReadWriteCloser) ([]byte, error) {
hdr := append([]byte(nil), model.MCPFsXferMagicDownloadHdr...)
sz := make([]byte, 8)
binary.BigEndian.PutUint64(sz, uint64(len(payload)))
hdr = append(hdr, sz...)
hdr = append(hdr, make([]byte, 32)...)
if _, err := stream.Write(hdr); err != nil {
return nil, err
}
if len(payload) > 0 {
chunk := append([]byte(nil), model.MCPFsXferMagicChunk...)
chunkLen := make([]byte, 8)
binary.BigEndian.PutUint64(chunkLen, uint64(len(payload)))
chunk = append(chunk, chunkLen...)
chunk = append(chunk, payload...)
if _, err := stream.Write(chunk); err != nil {
return nil, err
}
}
ok := append([]byte(nil), model.MCPFsXferMagicOK...)
ok = append(ok, sz...)
ok = append(ok, make([]byte, 32)...)
_, err := stream.Write(ok)
return payload, err
}
}
// xferAgentError 模拟 agent 直接发 NZTE 拒绝。
func xferAgentError(msg string) func(*model.FsTransferRequest, io.ReadWriteCloser) ([]byte, error) {
return func(_ *model.FsTransferRequest, stream io.ReadWriteCloser) ([]byte, error) {
buf := append([]byte(nil), model.MCPFsXferMagicErr...)
buf = append(buf, msg...)
_, err := stream.Write(buf)
return nil, err
}
}
func setupTransferTest(t *testing.T, agent func(*model.FsTransferRequest, io.ReadWriteCloser) ([]byte, error)) (*httptest.Server, string, func()) {
t.Helper()
cleanupBase, uid := setupMCPTest(t)
originalHandler := rpc.NezhaHandlerSingleton
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
stream := &xferStreamMux{agent: agent}
srv, _ := singleton.ServerShared.Get(7)
srv.SetTaskStream(stream)
_, plain := mkToken(t, uid, []string{
model.ScopeServerRead, model.ScopeServerWrite,
}, 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()
rpc.NezhaHandlerSingleton = originalHandler
cleanupBase()
}
}
func mintDownloadURL(t *testing.T, ts *httptest.Server, tok, path string) string {
t.Helper()
body := map[string]any{
"jsonrpc": "2.0", "id": 1, "method": "tools/call",
"params": map[string]any{
"name": "fs.download_url",
"arguments": map[string]any{"server_id": 7, "path": path, "ttl_seconds": 60},
},
}
b, _ := json.Marshal(body)
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader(b))
req.Header.Set("Authorization", "Bearer "+tok)
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))
res, _ := env["result"].(map[string]any)
struc, _ := res["structuredContent"].(map[string]any)
url, _ := struc["url"].(string)
require.NotEmpty(t, url, "fs.download_url did not return url: %v", env)
return ts.URL + url[strings.Index(url, "/mcp/"):]
}
// /mcp/download 必须把 agent 推过来的原始字节一字不差地交给 HTTP 客户端。
func TestTransferDownload_ReturnsRawBinaryBytes(t *testing.T) {
want := []byte{0x00, 0x01, 0xFF, 0xAB, 'h', 'i'}
ts, tok, cleanup := setupTransferTest(t, xferAgentDownloadSend(want))
defer cleanup()
url := mintDownloadURL(t, ts, tok, "/srv/blob")
resp, err := http.Get(url)
require.NoError(t, err)
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
require.Equal(t, http.StatusOK, resp.StatusCode, "status=%d body=%q", resp.StatusCode, string(body))
require.Equal(t, want, body, "client must receive raw file bytes")
}
func mintUploadURL(t *testing.T, ts *httptest.Server, tok, path string) string {
t.Helper()
body := map[string]any{
"jsonrpc": "2.0", "id": 1, "method": "tools/call",
"params": map[string]any{
"name": "fs.upload_url",
"arguments": map[string]any{"server_id": 7, "path": path, "ttl_seconds": 60},
},
}
b, _ := json.Marshal(body)
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader(b))
req.Header.Set("Authorization", "Bearer "+tok)
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))
res, _ := env["result"].(map[string]any)
struc, _ := res["structuredContent"].(map[string]any)
url, _ := struc["url"].(string)
require.NotEmpty(t, url, "fs.upload_url did not return url: %v", env)
return ts.URL + url[strings.Index(url, "/mcp/"):]
}
func TestTransferUpload_PreservesArbitraryBinary(t *testing.T) {
binary := []byte{0x00, 0x01, 0xC3, 0x28, 0xFF, 0xFE, 'h', 'i'}
var captured []byte
var capturedMu sync.Mutex
agent := func(req *model.FsTransferRequest, stream io.ReadWriteCloser) ([]byte, error) {
got, err := xferAgentUploadRead(req, stream)
capturedMu.Lock()
captured = got
capturedMu.Unlock()
if err != nil {
return got, err
}
return got, xferAgentUploadAck(stream, uint64(len(got)))
}
ts, tok, cleanup := setupTransferTest(t, agent)
defer cleanup()
url := mintUploadURL(t, ts, tok, "/srv/upload.bin")
upResp, err := http.Post(url, "application/octet-stream", bytes.NewReader(binary))
require.NoError(t, err)
defer upResp.Body.Close()
require.Equal(t, http.StatusOK, upResp.StatusCode)
capturedMu.Lock()
defer capturedMu.Unlock()
require.Equal(t, binary, captured, "agent must receive byte-for-byte body")
}
// agent 发完声明的 payload 后没有发任何最终控制帧就关掉 stream 时,
// dashboard 不能把这个未确认的传输当成成功:因为协议规定下载完成由 NZTO
// 帧承载 size/SHA256,缺失最终帧意味着 agent 没有正向确认整段数据。
func TestTransferDownload_MissingFinalOKFrameMustFail(t *testing.T) {
payload := []byte("partial-but-no-final-ok")
agent := func(_ *model.FsTransferRequest, stream io.ReadWriteCloser) ([]byte, error) {
hdr := append([]byte(nil), model.MCPFsXferMagicDownloadHdr...)
sz := make([]byte, 8)
binary.BigEndian.PutUint64(sz, uint64(len(payload)))
hdr = append(hdr, sz...)
hdr = append(hdr, make([]byte, 32)...)
if _, err := stream.Write(hdr); err != nil {
return nil, err
}
chunk := append([]byte(nil), model.MCPFsXferMagicChunk...)
chunkLen := make([]byte, 8)
binary.BigEndian.PutUint64(chunkLen, uint64(len(payload)))
chunk = append(chunk, chunkLen...)
chunk = append(chunk, payload...)
if _, err := stream.Write(chunk); err != nil {
return nil, err
}
// 故意不发任何最终帧:直接由 setup 的 defer agentSide.Close() 关闭。
return payload, nil
}
ts, tok, cleanup := setupTransferTest(t, agent)
defer cleanup()
url := mintDownloadURL(t, ts, tok, "/srv/missing-final.bin")
resp, err := http.Get(url)
require.NoError(t, err)
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
require.NotEqualf(t, http.StatusOK, resp.StatusCode,
"download without a final NZTO must not be reported as 200 OK; body=%q", string(body))
}
// agent 在 payload 之后写了一个非 NZTO 也非 NZTE 的乱码 4 字节 magic 时,
// dashboard 必须把它当作协议错误,而不是默默成功。
func TestTransferDownload_NonOKNonErrFinalMagicMustFail(t *testing.T) {
payload := []byte("ok-bytes-but-bogus-tail")
agent := func(_ *model.FsTransferRequest, stream io.ReadWriteCloser) ([]byte, error) {
hdr := append([]byte(nil), model.MCPFsXferMagicDownloadHdr...)
sz := make([]byte, 8)
binary.BigEndian.PutUint64(sz, uint64(len(payload)))
hdr = append(hdr, sz...)
hdr = append(hdr, make([]byte, 32)...)
if _, err := stream.Write(hdr); err != nil {
return nil, err
}
chunk := append([]byte(nil), model.MCPFsXferMagicChunk...)
chunkLen := make([]byte, 8)
binary.BigEndian.PutUint64(chunkLen, uint64(len(payload)))
chunk = append(chunk, chunkLen...)
chunk = append(chunk, payload...)
if _, err := stream.Write(chunk); err != nil {
return nil, err
}
// 12 字节、非 NZTO/NZTE 的乱码最终帧。
bogus := []byte{'X', 'X', 'X', 'X', 0, 0, 0, 0, 0, 0, 0, 0}
if _, err := stream.Write(bogus); err != nil {
return nil, err
}
return payload, nil
}
ts, tok, cleanup := setupTransferTest(t, agent)
defer cleanup()
url := mintDownloadURL(t, ts, tok, "/srv/bogus-final.bin")
resp, err := http.Get(url)
require.NoError(t, err)
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
require.NotEqualf(t, http.StatusOK, resp.StatusCode,
"download with a non-NZTO non-NZTE final frame must not be reported as 200 OK; body=%q", string(body))
}
// agent 拒绝(NZTE)时 dashboard 必须把错误透出去,不能假装 200。
func TestTransferDownload_SurfacesAgentError(t *testing.T) {
ts, tok, cleanup := setupTransferTest(t, xferAgentError("file too large"))
defer cleanup()
url := mintDownloadURL(t, ts, tok, "/srv/huge")
resp, err := http.Get(url)
require.NoError(t, err)
defer resp.Body.Close()
require.NotEqual(t, http.StatusOK, resp.StatusCode,
"agent NZTE must surface as non-200 to client")
}
// 下载途中 agent 发现源被截断并切到 NZTE 错误帧时,dashboard 绝不能
// 把那个错误帧的字节当成文件正文塞进 HTTP body —— 协议帧和文件字节
// 共用同一条 IOStreamHTTP 客户端不应收到“200 OK + 截断后混入 NZTE
// magic + agent 错误文本”。
func TestTransferDownload_MidStreamErrorDoesNotCorruptBody(t *testing.T) {
declared := []byte("HELLO-WORLD!")
partial := declared[:5]
agent := func(_ *model.FsTransferRequest, stream io.ReadWriteCloser) ([]byte, error) {
hdr := append([]byte(nil), model.MCPFsXferMagicDownloadHdr...)
sz := make([]byte, 8)
binary.BigEndian.PutUint64(sz, uint64(len(declared)))
hdr = append(hdr, sz...)
hdr = append(hdr, make([]byte, 32)...)
if _, err := stream.Write(hdr); err != nil {
return nil, err
}
chunk := append([]byte(nil), model.MCPFsXferMagicChunk...)
chunkLen := make([]byte, 8)
binary.BigEndian.PutUint64(chunkLen, uint64(len(partial)))
chunk = append(chunk, chunkLen...)
chunk = append(chunk, partial...)
if _, err := stream.Write(chunk); err != nil {
return nil, err
}
errFrame := append([]byte(nil), model.MCPFsXferMagicErr...)
errFrame = append(errFrame, []byte("source truncated mid-transfer")...)
_, err := stream.Write(errFrame)
return partial, err
}
ts, tok, cleanup := setupTransferTest(t, agent)
defer cleanup()
url := mintDownloadURL(t, ts, tok, "/srv/blob")
resp, err := http.Get(url)
require.NoError(t, err)
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
if resp.StatusCode == http.StatusOK {
require.Failf(t, "mid-stream NZTE leaked into HTTP body",
"expected non-200 once agent switched to NZTE; got 200 with body=%q (len=%d, declared=%d)",
string(body), len(body), len(declared))
}
require.NotContains(t, string(body), string(model.MCPFsXferMagicErr),
"NZTE control frame magic must never appear in the HTTP body")
}
// 上传时 Content-Length 超过 100MiB 必须直接 413,不进 IOStream。
func TestTransferUpload_RejectsOversizedBody(t *testing.T) {
ts, tok, cleanup := setupTransferTest(t, xferAgentUploadAccept)
defer cleanup()
url := mintUploadURL(t, ts, tok, "/srv/big.bin")
body := &bigReader{remaining: model.MCPFsTransferMaxSize + 1}
req, _ := http.NewRequest("POST", url, body)
req.ContentLength = int64(model.MCPFsTransferMaxSize + 1)
resp, err := http.DefaultClient.Do(req)
require.NoError(t, err)
defer resp.Body.Close()
require.Equal(t, http.StatusRequestEntityTooLarge, resp.StatusCode)
}
// dashboard 必须接受 ?sha256=<64hex> 形式并把 32B sha 透传给 agent。这个测试
// 不模拟失败,仅锁定 query 透传 + agent 正常返回 NZTO 时整链路 200。SHA256
// 真不匹配走的是下面 TestTransferUpload_SHA256MismatchReturns502。
func TestTransferUpload_AcceptsSHA256Query(t *testing.T) {
want := []byte("ohi")
var sawExpected string
var sawMu sync.Mutex
agent := func(req *model.FsTransferRequest, stream io.ReadWriteCloser) ([]byte, error) {
sawMu.Lock()
sawExpected = req.ExpectedSHA256
sawMu.Unlock()
return xferAgentUploadAccept(req, stream)
}
ts, tok, cleanup := setupTransferTest(t, agent)
defer cleanup()
want64 := strings.Repeat("0", 64)
url := mintUploadURL(t, ts, tok, "/srv/up.bin")
url += "?sha256=" + want64
resp, err := http.Post(url, "application/octet-stream", bytes.NewReader(want))
require.NoError(t, err)
defer resp.Body.Close()
require.Equal(t, http.StatusOK, resp.StatusCode)
sawMu.Lock()
defer sawMu.Unlock()
require.Equal(t, want64, sawExpected, "dashboard must forward ?sha256 to agent verbatim")
}
// SHA256 不匹配时 agent 会用 NZTE 拒绝;dashboard 必须把 NZTE 透传成 502 而不是
// 因为 io.CopyN 已经写完 body 就返回 200。原版测试用 xferAgentUploadAccept 模拟
// 成功握手,错误返回值被 dashboard 忽略,最终断言 200,把这条 integrity 错误
// 路径假阳性 pin 住了。此处用 xferAgentError 真正模拟 agent NZTE。
func TestTransferUpload_SHA256MismatchReturns502(t *testing.T) {
want := []byte("ohi")
agent := func(req *model.FsTransferRequest, stream io.ReadWriteCloser) ([]byte, error) {
if _, err := xferAgentUploadRead(req, stream); err != nil {
return nil, err
}
return xferAgentError("sha256 mismatch")(req, stream)
}
ts, tok, cleanup := setupTransferTest(t, agent)
defer cleanup()
url := mintUploadURL(t, ts, tok, "/srv/up.bin")
url += "?sha256=" + strings.Repeat("0", 64)
resp, err := http.Post(url, "application/octet-stream", bytes.NewReader(want))
require.NoError(t, err)
defer resp.Body.Close()
require.Equal(t, http.StatusBadGateway, resp.StatusCode)
body, _ := io.ReadAll(resp.Body)
require.Contains(t, string(body), "sha256 mismatch")
}
type bigReader struct{ remaining int64 }
func (b *bigReader) Read(p []byte) (int, error) {
if b.remaining <= 0 {
return 0, io.EOF
}
n := len(p)
if int64(n) > b.remaining {
n = int(b.remaining)
}
for i := 0; i < n; i++ {
p[i] = 0
}
b.remaining -= int64(n)
return n, nil
}
@@ -0,0 +1,30 @@
package controller
import (
"io"
"net/http"
"testing"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/model"
)
func TestTransferDownload_DataFrameBeginningWithErrMagicIsNotMisclassified(t *testing.T) {
collide := append([]byte(nil), model.MCPFsXferMagicErr...)
collide = append(collide, []byte("xx-real-file-bytes-xx")...)
ts, tok, cleanup := setupTransferTest(t, xferAgentDownloadSend(collide))
defer cleanup()
url := mintDownloadURL(t, ts, tok, "/srv/collide.bin")
resp, err := http.Get(url)
require.NoError(t, err)
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
require.Equal(t, http.StatusOK, resp.StatusCode,
"file content starting with the NZTE magic must not be misread as an agent error; status=%d body=%q",
resp.StatusCode, string(body))
require.Equal(t, collide, body,
"client must receive the raw file bytes byte-for-byte even when they start with NZTE")
}
@@ -0,0 +1,124 @@
package controller
import (
"bytes"
"crypto/sha256"
"encoding/binary"
"encoding/hex"
"testing"
"github.com/nezhahq/nezha/model"
)
// M2 regression: download finalisation must validate the trailing NZTO
// frame's declared size AND sha256 against what was actually streamed.
// The old relay only checked the 4-byte magic, so a truncated NZTO (no
// hash) or a wrong-hash payload was silently accepted.
func TestValidateDownloadFinal_RejectsTruncatedNZTO(t *testing.T) {
buf := make([]byte, 4)
copy(buf, model.MCPFsXferMagicOK)
if err := validateDownloadFinal(buf, 0, sha256.New().Sum(nil)); err == nil {
t.Fatal("a 4-byte NZTO (magic only, no size+sha) must be rejected")
}
}
func TestValidateDownloadFinal_RejectsSizeMismatch(t *testing.T) {
h := sha256.New()
h.Write([]byte("payload"))
sum := h.Sum(nil)
buf := make([]byte, 4+8+32)
copy(buf[:4], model.MCPFsXferMagicOK)
binary.BigEndian.PutUint64(buf[4:12], 999) // declared size 999
copy(buf[12:44], sum)
if err := validateDownloadFinal(buf, int64(len("payload")), sum); err == nil {
t.Fatal("declared size mismatch with actual streamed bytes must be rejected")
}
}
func TestValidateDownloadFinal_RejectsHashMismatch(t *testing.T) {
declared := []byte("declared")
streamed := []byte("streamed-something-else")
h := sha256.New()
h.Write(declared)
declaredHash := h.Sum(nil)
streamedH := sha256.New()
streamedH.Write(streamed)
streamedHash := streamedH.Sum(nil)
buf := make([]byte, 4+8+32)
copy(buf[:4], model.MCPFsXferMagicOK)
binary.BigEndian.PutUint64(buf[4:12], uint64(len(streamed)))
copy(buf[12:44], declaredHash)
if err := validateDownloadFinal(buf, int64(len(streamed)), streamedHash); err == nil {
t.Fatal("declared sha256 != streamed sha256 must be rejected")
}
}
func TestValidateDownloadFinal_AcceptsMatchingSizeAndHash(t *testing.T) {
payload := []byte("hello world")
h := sha256.New()
h.Write(payload)
sum := h.Sum(nil)
buf := make([]byte, 4+8+32)
copy(buf[:4], model.MCPFsXferMagicOK)
binary.BigEndian.PutUint64(buf[4:12], uint64(len(payload)))
copy(buf[12:44], sum)
if err := validateDownloadFinal(buf, int64(len(payload)), sum); err != nil {
t.Fatalf("matching final header must pass, got %v", err)
}
}
// Hash skip: agent may omit the sha when the source filesystem can't
// produce one (e.g. live device). Encode as all-zero sha256; that's a
// legal but explicit "no hash" signal. Size must still match.
func TestValidateDownloadFinal_AllowsAllZeroHashAsExplicitSkip(t *testing.T) {
payload := []byte("nothash")
streamedHash, _ := hex.DecodeString("0000000000000000000000000000000000000000000000000000000000000000")
buf := make([]byte, 4+8+32)
copy(buf[:4], model.MCPFsXferMagicOK)
binary.BigEndian.PutUint64(buf[4:12], uint64(len(payload)))
// declared bytes 12-44 already zero by make()
if err := validateDownloadFinal(buf, int64(len(payload)), streamedHash); err != nil {
t.Fatalf("all-zero declared hash with matching size must pass (explicit skip), got %v", err)
}
}
// Defence-in-depth: the magic must still match. validateDownloadFinal is
// reached after the relay already checked it, but a second check costs
// nothing and survives future refactors that split the parsing.
func TestValidateDownloadFinal_RejectsWrongMagic(t *testing.T) {
buf := make([]byte, 4+8+32)
copy(buf[:4], []byte("XXXX"))
if err := validateDownloadFinal(buf, 0, sha256.New().Sum(nil)); err == nil {
t.Fatal("non-NZTO magic must be rejected")
}
}
// Defence-in-depth: bytes.Compare of slices of different length still
// returns non-zero, but Go semantics for hex.EncodeToString are wider
// than 32 bytes. Pin that the validator only inspects the first 32 hash
// bytes.
func TestValidateDownloadFinal_OnlyConsiders32HashBytes(t *testing.T) {
payload := []byte("X")
h := sha256.New()
h.Write(payload)
sum := h.Sum(nil)
buf := make([]byte, 4+8+32)
copy(buf[:4], model.MCPFsXferMagicOK)
binary.BigEndian.PutUint64(buf[4:12], uint64(len(payload)))
copy(buf[12:44], sum)
streamedHashExtra := append(bytes.Clone(sum), 0xAA, 0xBB)
if err := validateDownloadFinal(buf, int64(len(payload)), streamedHashExtra); err != nil {
t.Fatalf("validator must compare exactly the first 32 streamed hash bytes, got %v", err)
}
}
@@ -0,0 +1,147 @@
package controller
import (
"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/singleton"
)
// fs.upload / fs.download 失败路径必须写一条 MCPAuditLog,否则审计表只能看到
// 成功调用,运营无法发现"PAT 被吊销后仍有人尝试消费 URL"、"agent 拒绝执行"、
// "kill switch 已开却仍有调用打进来"这类信号。成功路径已经在写审计,这里把
// 失败路径的契约钉死。
func countAuditRows(t *testing.T, tool, outcome string) int64 {
t.Helper()
var cnt int64
q := singleton.DB.Model(&model.MCPAuditLog{}).Where("tool = ?", tool)
if outcome != "" {
q = q.Where("outcome = ?", outcome)
}
require.NoError(t, q.Count(&cnt).Error)
return cnt
}
func newTransferRouter(t *testing.T) *gin.Engine {
t.Helper()
gin.SetMode(gin.TestMode)
r := gin.New()
r.GET("/mcp/download/:token", transferDownloadHandler)
r.POST("/mcp/upload/:token", transferUploadHandler)
return r
}
func TestTransferDownload_AuditsTokenExpired(t *testing.T) {
cleanup, uid := setupMCPTest(t)
defer cleanup()
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
r := newTransferRouter(t)
url, err := mintTransferToken(transferEntry{
UserID: uid,
TokenID: tok.ID,
ServerID: 7,
Path: "/srv/blob",
Direction: transferDirDownload,
ExpiresAt: time.Now().Add(-time.Second),
})
require.NoError(t, err)
w := httptest.NewRecorder()
req := httptest.NewRequest("GET", "/mcp/download/"+url, nil)
r.ServeHTTP(w, req)
require.Equal(t, http.StatusUnauthorized, w.Code,
"expired token must surface as 401 to client")
require.Equal(t, int64(1), countAuditRows(t, "fs.download", ""),
"failed download must still produce an audit row so SIEM can observe the rejection")
}
func TestTransferDownload_AuditsRevalidateFailureWhenMCPDisabled(t *testing.T) {
cleanup, uid := setupMCPTest(t)
defer cleanup()
tok, _ := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
r := newTransferRouter(t)
url, 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)
singleton.Conf.SetMCPEnabled(false)
w := httptest.NewRecorder()
req := httptest.NewRequest("GET", "/mcp/download/"+url, nil)
r.ServeHTTP(w, req)
require.Equal(t, http.StatusUnauthorized, w.Code)
require.Equal(t, int64(1),
countAuditRows(t, "fs.download", model.MCPOutcomeMCPDisabled),
"kill switch must be observable in audit log with outcome=mcp_disabled, not silently swallowed")
}
func TestTransferUpload_AuditsTokenExpired(t *testing.T) {
cleanup, uid := setupMCPTest(t)
defer cleanup()
tok, _ := mkToken(t, uid, []string{model.ScopeServerWrite}, nil)
r := newTransferRouter(t)
url, err := mintTransferToken(transferEntry{
UserID: uid,
TokenID: tok.ID,
ServerID: 7,
Path: "/srv/blob",
Direction: transferDirUpload,
ExpiresAt: time.Now().Add(-time.Second),
})
require.NoError(t, err)
w := httptest.NewRecorder()
req := httptest.NewRequest("POST", "/mcp/upload/"+url, strings.NewReader(""))
req.ContentLength = 0
r.ServeHTTP(w, req)
require.Equal(t, http.StatusUnauthorized, w.Code)
require.Equal(t, int64(1), countAuditRows(t, "fs.upload", ""),
"failed upload must still produce an audit row")
}
func TestTransferUpload_AuditsRevalidateFailureWhenMCPDisabled(t *testing.T) {
cleanup, uid := setupMCPTest(t)
defer cleanup()
tok, _ := mkToken(t, uid, []string{model.ScopeServerWrite}, nil)
r := newTransferRouter(t)
url, err := mintTransferToken(transferEntry{
UserID: uid,
TokenID: tok.ID,
ServerID: 7,
Path: "/srv/blob",
Direction: transferDirUpload,
ExpiresAt: time.Now().Add(5 * time.Minute),
})
require.NoError(t, err)
singleton.Conf.SetMCPEnabled(false)
w := httptest.NewRecorder()
req := httptest.NewRequest("POST", "/mcp/upload/"+url, strings.NewReader(""))
req.ContentLength = 0
r.ServeHTTP(w, req)
require.Equal(t, http.StatusUnauthorized, w.Code)
require.Equal(t, int64(1),
countAuditRows(t, "fs.upload", model.MCPOutcomeMCPDisabled),
"upload kill switch must be observable in audit log")
}
@@ -0,0 +1,53 @@
package controller
import (
"testing"
"time"
"github.com/stretchr/testify/require"
)
// transferEntries 是 mint→consume 的内存表。
// mintTransferToken 把 token Store 进去,consumeTransferToken 命中后才删,
// PurgeTransferEntries 是 kill switch 的全量清理。这条路径目前缺一个
// 按 ExpiresAt 的过期回收:从未被 consume 的 token 会一直留到下一次
// kill switch 才被清掉。
//
// 这条测试钉死「过期项必须被 gcExpiredTransferEntries() 清掉,
// 且未过期项必须保留」。
func TestGCExpiredTransferEntries_RemovesOnlyExpired(t *testing.T) {
// 隔离全局状态,避免被其它测试遗留的 entry 干扰。
PurgeTransferEntries()
t.Cleanup(func() { PurgeTransferEntries() })
now := time.Now()
expiredTok, err := mintTransferToken(transferEntry{
UserID: 1,
TokenID: 1,
ServerID: 1,
Path: "/srv/expired",
Direction: transferDirDownload,
ExpiresAt: now.Add(-time.Second),
})
require.NoError(t, err)
freshTok, err := mintTransferToken(transferEntry{
UserID: 1,
TokenID: 1,
ServerID: 1,
Path: "/srv/fresh",
Direction: transferDirDownload,
ExpiresAt: now.Add(5 * time.Minute),
})
require.NoError(t, err)
removed := gcExpiredTransferEntries(now)
require.Equal(t, 1, removed, "exactly one expired entry must be removed")
_, expiredStillThere := transferEntries.Load(expiredTok)
require.False(t, expiredStillThere, "expired entry must be gone after GC")
_, freshStillThere := transferEntries.Load(freshTok)
require.True(t, freshStillThere, "fresh entry must survive GC")
// 二次 GC 不应误删未过期项,也不应报告假阳性。
require.Equal(t, 0, gcExpiredTransferEntries(now), "second GC must be a no-op for non-expired entries")
}
@@ -0,0 +1,113 @@
package controller
import (
"bufio"
"bytes"
"encoding/binary"
"errors"
"net"
"net/http"
"runtime"
"testing"
"github.com/gin-gonic/gin"
"github.com/nezhahq/nezha/model"
)
type fixedSizeFrameStream struct {
buf bytes.Buffer
}
func (s *fixedSizeFrameStream) Read(p []byte) (int, error) { return s.buf.Read(p) }
func (s *fixedSizeFrameStream) Write(p []byte) (int, error) { return len(p), nil }
func (s *fixedSizeFrameStream) Close() error { return nil }
func writeChunkFrame(out *bytes.Buffer, chunk []byte) {
out.Write(model.MCPFsXferMagicChunk)
var sz [8]byte
binary.BigEndian.PutUint64(sz[:], uint64(len(chunk)))
out.Write(sz[:])
out.Write(chunk)
}
func writeOKFrame(out *bytes.Buffer, size uint64) {
out.Write(model.MCPFsXferMagicOK)
var sz [8]byte
binary.BigEndian.PutUint64(sz[:], size)
out.Write(sz[:])
out.Write(make([]byte, 32))
}
// countingDiscardWriter satisfies http.ResponseWriter but throws bytes away
// after counting them, so the test can measure relayDownloadFrames heap
// pressure without httptest.ResponseRecorder caching 100MiB of body in
// memory and dominating the measurement.
type countingDiscardWriter struct {
header http.Header
written int64
status int
}
func newCountingDiscardWriter() *countingDiscardWriter {
return &countingDiscardWriter{header: make(http.Header)}
}
func (w *countingDiscardWriter) Header() http.Header { return w.header }
func (w *countingDiscardWriter) Write(p []byte) (int, error) {
w.written += int64(len(p))
return len(p), nil
}
func (w *countingDiscardWriter) WriteHeader(status int) { w.status = status }
func (w *countingDiscardWriter) Flush() {}
func (w *countingDiscardWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
return nil, nil, errors.New("not hijackable")
}
func TestRelayDownloadFrames_DoesNotBufferEntirePayloadInMemory(t *testing.T) {
const size = int64(model.MCPFsTransferMaxSize)
const chunk = 1 * 1024 * 1024
var src bytes.Buffer
payload := make([]byte, chunk)
for i := range payload {
payload[i] = byte(i % 251)
}
remaining := size
for remaining > 0 {
toWrite := int64(chunk)
if toWrite > remaining {
toWrite = remaining
}
writeChunkFrame(&src, payload[:toWrite])
remaining -= toWrite
}
writeOKFrame(&src, uint64(size))
stream := &fixedSizeFrameStream{buf: src}
sink := newCountingDiscardWriter()
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(sink)
runtime.GC()
var before runtime.MemStats
runtime.ReadMemStats(&before)
if err := relayDownloadFrames(c, stream, size); err != nil {
t.Fatalf("relayDownloadFrames returned err: %v", err)
}
var after runtime.MemStats
runtime.ReadMemStats(&after)
delta := int64(after.HeapAlloc) - int64(before.HeapAlloc)
const allow = 16 * 1024 * 1024
if delta > allow {
t.Fatalf("relayDownloadFrames retained %d bytes in heap after a %d-byte transfer (allow <= %d). 100MiB 旁路通道不应整文件缓在 dashboard 内存里。",
delta, size, allow)
}
if sink.written != size {
t.Fatalf("expected %d bytes forwarded to HTTP client, got %d", size, sink.written)
}
}
@@ -0,0 +1,38 @@
package controller
import (
"context"
"strings"
"testing"
"github.com/nezhahq/nezha/model"
"github.com/nezhahq/nezha/service/singleton"
)
func TestValidateTransferPathRejectsOversizedPath(t *testing.T) {
if err := validateTransferPath(strings.Repeat("a", maxTransferPathLen+1)); err == nil {
t.Fatal("path longer than maxTransferPathLen must be rejected to bound transferEntry memory")
}
if err := validateTransferPath(""); err == nil {
t.Fatal("empty path must be rejected")
}
if err := validateTransferPath("/etc/hostname"); err != nil {
t.Fatalf("a normal path must be accepted, got %v", err)
}
}
// openFsTransferStream must refuse to start a new transfer (and never reach
// SendTask) once the administrator has disabled MCP, closing the race window
// between revalidateTransferEntry and stream creation.
func TestOpenFsTransferStreamRefusesWhenMCPDisabled(t *testing.T) {
originalConf := singleton.Conf
t.Cleanup(func() { singleton.Conf = originalConf })
cfg := &model.Config{}
cfg.SetMCPEnabled(false)
singleton.Conf = &singleton.ConfigClass{Config: cfg}
_, _, err := openFsTransferStream(context.Background(), 1, &model.FsTransferRequest{})
if err == nil {
t.Fatal("transfer stream must not open while MCP is disabled")
}
}
@@ -0,0 +1,56 @@
package controller
import (
"encoding/binary"
"testing"
"github.com/nezhahq/nezha/model"
)
// M3 regression: a malicious or corrupt agent can put `> MaxInt64` into the
// size field of an NZTU/NZTD/NZTO frame. Direct uint64→int64 cast wraps to
// a negative value, which bypasses the `hdr.Size > MCPFsTransferMaxSize`
// check (a negative is always less). The guarded reader must reject
// oversize raw u64 BEFORE narrowing.
func TestReadXferFixedHeader_RejectsOversizedUploadSize(t *testing.T) {
buf := make([]byte, 4+8)
copy(buf[:4], model.MCPFsXferMagicUploadHdr)
binary.BigEndian.PutUint64(buf[4:12], uint64(model.MCPFsTransferMaxSize+1))
_, err := readXferFixedHeaderFromBytes(buf)
if err == nil {
t.Fatal("size > MCPFsTransferMaxSize must be rejected; otherwise int64 narrowing lets the upload through with a negative size")
}
}
func TestReadXferFixedHeader_RejectsOverflowingUploadSize(t *testing.T) {
buf := make([]byte, 4+8)
copy(buf[:4], model.MCPFsXferMagicUploadHdr)
binary.BigEndian.PutUint64(buf[4:12], ^uint64(0))
_, err := readXferFixedHeaderFromBytes(buf)
if err == nil {
t.Fatal("raw u64=MaxUint64 must be rejected before int64 cast wraps it to -1")
}
}
func TestReadXferFixedHeader_AcceptsLegalUploadSize(t *testing.T) {
buf := make([]byte, 4+8)
copy(buf[:4], model.MCPFsXferMagicUploadHdr)
binary.BigEndian.PutUint64(buf[4:12], 1024)
hdr, err := readXferFixedHeaderFromBytes(buf)
if err != nil {
t.Fatalf("legal size must pass, got %v", err)
}
if hdr.Size != 1024 {
t.Fatalf("want size=1024, got %d", hdr.Size)
}
}
func TestReadXferFixedHeader_RejectsOversizedDownloadSize(t *testing.T) {
buf := make([]byte, 4+8+32)
copy(buf[:4], model.MCPFsXferMagicDownloadHdr)
binary.BigEndian.PutUint64(buf[4:12], uint64(model.MCPFsTransferMaxSize+1))
_, err := readXferFixedHeaderFromBytes(buf)
if err == nil {
t.Fatal("download size > MCPFsTransferMaxSize must be rejected")
}
}
@@ -0,0 +1,44 @@
package controller
import (
"io"
"os"
"runtime"
)
type transferSpool struct {
f *os.File
}
func newTransferSpool() (*transferSpool, error) {
f, err := os.CreateTemp("", "nz-mcp-xfer-*")
if err != nil {
return nil, err
}
// 提前 unlink,文件句柄关掉就回收磁盘;Windows 不支持就留到 Close 兜底。
if runtime.GOOS != "windows" {
_ = os.Remove(f.Name())
}
return &transferSpool{f: f}, nil
}
func (s *transferSpool) Write(p []byte) (int, error) { return s.f.Write(p) }
func (s *transferSpool) Read(p []byte) (int, error) { return s.f.Read(p) }
func (s *transferSpool) Rewind() error {
_, err := s.f.Seek(0, io.SeekStart)
return err
}
func (s *transferSpool) Close() {
if s.f == nil {
return
}
name := s.f.Name()
_ = s.f.Close()
if runtime.GOOS == "windows" {
_ = os.Remove(name)
}
s.f = nil
}
@@ -0,0 +1,94 @@
package controller
import (
"bytes"
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/model"
)
// fs.upload_url 必须把 agent 已经支持的上传语义(mode / create_dirs /
// if_match_sha256)从 MCP tool arguments 透传到 agent 的 FsTransferRequest。
// 当前实现复用了 fs.download_url 的 fsDownloadURLArgs,只解析 server_id /
// path / ttl_seconds,导致这些字段被静默丢弃,跨仓 wire model + agent 能力
// 与 MCP 工具调用面失联。
func TestMintFsUploadURL_PropagatesModeCreateDirsAndIfMatchToAgent(t *testing.T) {
var captured *model.FsTransferRequest
var mu sync.Mutex
agent := func(req *model.FsTransferRequest, stream io.ReadWriteCloser) ([]byte, error) {
mu.Lock()
copyReq := *req
captured = &copyReq
mu.Unlock()
got, err := xferAgentUploadRead(req, stream)
if err != nil {
return got, err
}
return got, xferAgentUploadAck(stream, uint64(len(got)))
}
ts, tok, cleanup := setupTransferTest(t, agent)
defer cleanup()
url := mintFsUploadURLWithOptions(t, ts, tok, "/srv/upload.bin", map[string]any{
"mode": "0640",
"create_dirs": true,
"if_match_sha256": strings.Repeat("a", 64),
})
upResp, err := http.Post(url, "application/octet-stream", bytes.NewReader([]byte("hello")))
require.NoError(t, err)
defer upResp.Body.Close()
require.Equal(t, http.StatusOK, upResp.StatusCode)
mu.Lock()
defer mu.Unlock()
require.NotNil(t, captured, "agent must have received the FsTransferRequest")
require.Equal(t, "0640", captured.Mode,
"fs.upload_url must forward mode to the agent FsTransferRequest")
require.True(t, captured.CreateDirs,
"fs.upload_url must forward create_dirs to the agent FsTransferRequest")
require.Equal(t, strings.Repeat("a", 64), captured.IfMatchSHA256,
"fs.upload_url must forward if_match_sha256 to the agent FsTransferRequest")
}
func mintFsUploadURLWithOptions(t *testing.T, ts *httptest.Server, tok, path string, extra map[string]any) string {
t.Helper()
args := map[string]any{
"server_id": 7,
"path": path,
"ttl_seconds": 60,
}
for k, v := range extra {
args[k] = v
}
body := map[string]any{
"jsonrpc": "2.0", "id": 1, "method": "tools/call",
"params": map[string]any{
"name": "fs.upload_url",
"arguments": args,
},
}
b, _ := json.Marshal(body)
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader(b))
req.Header.Set("Authorization", "Bearer "+tok)
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))
res, _ := env["result"].(map[string]any)
struc, _ := res["structuredContent"].(map[string]any)
url, _ := struc["url"].(string)
require.NotEmptyf(t, url, "fs.upload_url did not return url: %v", env)
return ts.URL + url[strings.Index(url, "/mcp/"):]
}
+1
View File
@@ -201,6 +201,7 @@ func oauth2callback(jwtConfig *jwt.GinJWTMiddleware) func(c *gin.Context) (any,
}
jwtConfig.SetCookie(c, tokenString)
setCSRFCookie(c)
c.Redirect(http.StatusFound, utils.IfOr(state.Action == model.RTypeBind, "/dashboard/profile?oauth2=true", "/dashboard/login?oauth2=true"))
return nil, errNoop
@@ -0,0 +1,28 @@
package controller
import (
"net/http/httptest"
"strings"
"testing"
"github.com/gin-gonic/gin"
)
// setCSRFCookie must mint a readable nz-csrf cookie; OAuth2 callback relies on
// it so OAuth-only sessions can satisfy the double-submit CSRF gate.
func TestSetCSRFCookieIssuesReadableToken(t *testing.T) {
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest("GET", "/", nil)
setCSRFCookie(c)
setCookie := w.Header().Get("Set-Cookie")
if !strings.Contains(setCookie, csrfCookieName+"=") {
t.Fatalf("expected %s cookie, got %q", csrfCookieName, setCookie)
}
if strings.Contains(strings.ToLower(setCookie), "httponly") {
t.Fatal("CSRF cookie must be JS-readable (not HttpOnly) for the SPA to mirror it")
}
}
+196
View File
@@ -0,0 +1,196 @@
package controller
import (
"fmt"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"github.com/patrickmn/go-cache"
"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"
)
// OAuth2 callback 测试核心安全语义:state CSRF、provider 校验、解绑权限。
//
// 这些测试用 verifyState 的私有路径直接构造场景,因为 callback 的完整链路涉及
// 真实 IdP HTTP 调用;safety-critical 的 state 校验本身可以单测。
func setupOAuth2Test(t *testing.T) func() {
t.Helper()
originalDB := singleton.DB
originalConf := singleton.Conf
originalCache := singleton.Cache
originalLocalizer := singleton.Localizer
if singleton.Localizer == nil {
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
}
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.User{}, &model.Oauth2Bind{}, &model.WAF{}))
singleton.DB = db
singleton.Conf = &singleton.ConfigClass{Config: &model.Config{
Oauth2: map[string]*model.Oauth2Config{
"github": {ClientID: "x", ClientSecret: "y"},
},
}}
singleton.Cache = cache.New(time.Minute, time.Minute)
return func() {
singleton.DB = originalDB
singleton.Conf = originalConf
singleton.Cache = originalCache
singleton.Localizer = originalLocalizer
}
}
func newOAuth2Ctx(t *testing.T) (*gin.Context, *httptest.ResponseRecorder) {
t.Helper()
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest("GET", "/oauth2/callback", nil)
c.Set(model.CtxKeyRealIPStr, "1.2.3.4")
return c, w
}
func TestOAuth2_VerifyState_RejectsMissingCookie(t *testing.T) {
defer setupOAuth2Test(t)()
c, _ := newOAuth2Ctx(t)
_, err := verifyState(c, "any-state-value")
require.Error(t, err, "missing nz-o2s cookie must be rejected")
}
func TestOAuth2_VerifyState_RejectsUnknownCookie(t *testing.T) {
defer setupOAuth2Test(t)()
c, _ := newOAuth2Ctx(t)
c.Request.AddCookie(&http.Cookie{Name: "nz-o2s", Value: "never-issued-key"})
_, err := verifyState(c, "any-state")
require.Error(t, err, "unknown state key (no cache entry) must be rejected")
}
func TestOAuth2_VerifyState_RejectsStateMismatch(t *testing.T) {
defer setupOAuth2Test(t)()
c, _ := newOAuth2Ctx(t)
stateKey := "k-1"
singleton.Cache.Set(
fmt.Sprintf("%s%s", model.CacheKeyOauth2State, stateKey),
&model.Oauth2State{State: "real-state", Provider: "github"},
cache.DefaultExpiration,
)
c.Request.AddCookie(&http.Cookie{Name: "nz-o2s", Value: stateKey})
_, err := verifyState(c, "forged-state")
require.Error(t, err, "attacker-supplied state that differs from cached must be rejected (CSRF defense)")
}
func TestOAuth2_VerifyState_HappyPath(t *testing.T) {
defer setupOAuth2Test(t)()
c, _ := newOAuth2Ctx(t)
stateKey := "k-ok"
singleton.Cache.Set(
fmt.Sprintf("%s%s", model.CacheKeyOauth2State, stateKey),
&model.Oauth2State{State: "good-state", Provider: "github", Action: model.RTypeBind},
cache.DefaultExpiration,
)
c.Request.AddCookie(&http.Cookie{Name: "nz-o2s", Value: stateKey})
st, err := verifyState(c, "good-state")
require.NoError(t, err)
require.Equal(t, "github", st.Provider)
require.Equal(t, model.RTypeBind, st.Action)
}
func TestOAuth2_Unbind_UnknownProviderRejected(t *testing.T) {
defer setupOAuth2Test(t)()
c, _ := newOAuth2Ctx(t)
c.Params = gin.Params{{Key: "provider", Value: "unknown-provider"}}
c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 1}})
_, err := unbindOauth2(c)
require.Error(t, err)
require.Contains(t, err.Error(), "provider not found")
}
func TestOAuth2_Unbind_BlocksLastBindWhenRejectPassword(t *testing.T) {
defer setupOAuth2Test(t)()
require.NoError(t, singleton.DB.Create(&model.Oauth2Bind{
UserID: 42,
Provider: "github",
OpenID: "openid-only-one",
}).Error)
c, _ := newOAuth2Ctx(t)
c.Params = gin.Params{{Key: "provider", Value: "github"}}
c.Set(model.CtxKeyAuthorizedUser, &model.User{
Common: model.Common{ID: 42},
RejectPassword: true,
})
_, err := unbindOauth2(c)
require.Error(t, err,
"user with reject_password=true must NOT be able to unbind their last OAuth2 provider (would lock them out)")
}
func TestOAuth2_Unbind_AllowsWhenPasswordLoginPossible(t *testing.T) {
defer setupOAuth2Test(t)()
require.NoError(t, singleton.DB.Create(&model.Oauth2Bind{
UserID: 42,
Provider: "github",
OpenID: "openid-1",
}).Error)
c, _ := newOAuth2Ctx(t)
c.Params = gin.Params{{Key: "provider", Value: "github"}}
c.Set(model.CtxKeyAuthorizedUser, &model.User{
Common: model.Common{ID: 42},
RejectPassword: false,
})
_, err := unbindOauth2(c)
require.NoError(t, err)
var cnt int64
require.NoError(t, singleton.DB.Model(&model.Oauth2Bind{}).
Where("user_id = ? AND provider = ?", 42, "github").Count(&cnt).Error)
require.Equal(t, int64(0), cnt, "binding must be deleted")
}
func TestOAuth2_Unbind_OnlyAffectsOwnBindings(t *testing.T) {
defer setupOAuth2Test(t)()
require.NoError(t, singleton.DB.Create(&model.Oauth2Bind{
UserID: 42, Provider: "github", OpenID: "mine",
}).Error)
require.NoError(t, singleton.DB.Create(&model.Oauth2Bind{
UserID: 999, Provider: "github", OpenID: "victim",
}).Error)
c, _ := newOAuth2Ctx(t)
c.Params = gin.Params{{Key: "provider", Value: "github"}}
c.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 42}})
_, err := unbindOauth2(c)
require.NoError(t, err)
var victim model.Oauth2Bind
require.NoError(t, singleton.DB.
Where("user_id = ? AND provider = ?", 999, "github").
First(&victim).Error,
"another user's binding must not be touched")
require.Equal(t, "victim", victim.OpenID)
}
@@ -0,0 +1,63 @@
package controller
import (
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/nezhahq/nezha/model"
)
// L1 regression: patHasServerWhitelist used to type-assert specifically to
// *model.APIToken, so any other APITokenAccessor that ALSO implements
// CanAccessServer/ServerIDs (test stubs, future wrappers) was silently
// treated as "not limited" by the cover-fanout guard. The check must use
// the APITokenWhitelistView interface instead.
type viewOnlyPAT struct {
ids []uint64
}
func (v *viewOnlyPAT) CanAccessServer(id uint64) bool {
for _, x := range v.ids {
if x == id {
return true
}
}
return false
}
func (v *viewOnlyPAT) ServerIDs() []uint64 { return v.ids }
func TestPatHasServerWhitelist_RecognisesNonAPITokenWhitelistViewImplementor(t *testing.T) {
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAPIToken, &viewOnlyPAT{ids: []uint64{1}})
if !patHasServerWhitelist(ctx) {
t.Fatal("any APITokenWhitelistView implementor with non-empty ServerIDs must be flagged as limited; otherwise non-*model.APIToken wrappers silently escape the cover-fanout guard")
}
}
func TestPatHasServerWhitelist_EmptyWhitelistViaInterfaceCountsAsUnlimited(t *testing.T) {
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAPIToken, &viewOnlyPAT{ids: nil})
if patHasServerWhitelist(ctx) {
t.Fatal("empty whitelist = unlimited (existing semantics); must continue to return false")
}
}
func TestPatHasServerWhitelist_NoPATReturnsFalse(t *testing.T) {
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
if patHasServerWhitelist(ctx) {
t.Fatal("JWT requests (no PAT) must return false — there's no whitelist to escape")
}
}
func TestPatHasServerWhitelist_RealAPITokenStillWorks(t *testing.T) {
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAPIToken, &model.APIToken{ServersCSV: "1,2"})
if !patHasServerWhitelist(ctx) {
t.Fatal("real *model.APIToken with ServersCSV must still be flagged as limited (regression backstop)")
}
}
+458 -6
View File
@@ -1,12 +1,31 @@
package controller
import (
"slices"
"github.com/gin-gonic/gin"
"github.com/nezhahq/nezha/model"
"github.com/nezhahq/nezha/service/rpc"
"github.com/nezhahq/nezha/service/singleton"
)
// streamAttachAllowedForRequest combines the existing creator/admin check
// with a per-request PAT whitelist gate against the stream's target server.
// Terminal and FM endpoints attach to a long-lived stream and inherit any
// authority the creator held — without the second gate an admin's PAT
// scoped to [X] could hijack a stream targeting server Y.
func streamAttachAllowedForRequest(c *gin.Context, streamId string) bool {
if !rpc.NezhaHandlerSingleton.IsStreamAuthorizedForUser(streamId, getUid(c), callerIsAdmin(c)) {
return false
}
target, ok := rpc.NezhaHandlerSingleton.StreamTarget(streamId)
if !ok {
return false
}
return patAllowsServer(c, target)
}
func callerIsAdmin(c *gin.Context) bool {
auth, ok := c.Get(model.CtxKeyAuthorizedUser)
if !ok {
@@ -19,10 +38,441 @@ func callerIsAdmin(c *gin.Context) bool {
return user.Role.IsAdmin()
}
// patAllowsServer reports whether the caller's PAT (if any) is allowed to
// touch serverID. JWT callers (no PAT in context) always pass. Used as an
// extra guard before the admin / owner short-circuits so a PAT scoped to
// a server_ids whitelist cannot widen reach via the caller's admin role.
func patAllowsServer(c *gin.Context, serverID uint64) bool {
v, ok := c.Get(model.CtxKeyAPIToken)
if !ok {
return true
}
tok, _ := v.(model.APITokenAccessor)
if tok == nil {
return true
}
return tok.CanAccessServer(serverID)
}
// patHasServerWhitelist reports whether the caller is authenticated by a PAT
// that carries a non-empty server_ids whitelist. Cover-all semantics in
// Cron (CronCoverAll / CronCoverIgnoreAll-with-empty-Servers) and Service
// (ServiceCoverAll-with-empty-SkipServers) intentionally fan out to every
// server the cron/service's owner has — so a whitelisted PAT cannot create
// or update such configs without escaping its own whitelist. JWT callers
// and unscoped PATs have no whitelist to escape and pass through.
//
// This is the gate that turns the implicit-cover bypass at
// /api/v1/{cron,service} POST/PATCH into a 403; the dispatch side
// (CronTrigger, DispatchTask) does not re-check PAT context, so the only
// safe place to enforce it is at write time.
func patHasServerWhitelist(c *gin.Context) bool {
v, ok := c.Get(model.CtxKeyAPIToken)
if !ok {
return false
}
wl, ok := v.(model.APITokenWhitelistView)
if !ok || wl == nil {
return false
}
return len(wl.ServerIDs()) > 0
}
// patAccessorFromContext returns the request's PAT viewed as an
// APITokenAccessor, or nil for JWT requests. Routes that need to project
// server-keyed data through the PAT whitelist (server-group, ws/server,
// future stream/list endpoints) use this instead of poking c.Get directly.
func patAccessorFromContext(c *gin.Context) model.APITokenAccessor {
v, ok := c.Get(model.CtxKeyAPIToken)
if !ok {
return nil
}
tok, _ := v.(model.APITokenAccessor)
if tok == nil {
return nil
}
return tok
}
// checkCronServerListPermission validates the cron's Servers field. Under
// CronCoverIgnoreAll / CronCoverAlertTrigger the field is an allow-list and
// must satisfy Server.HasPermission (owner + PAT whitelist). Under
// CronCoverAll the field is a deny-list expressing exclusion; the caller
// only needs to own each listed server (PAT whitelist intersection is
// enforced separately by assertPATCoverFanoutWithinWhitelist).
func checkCronServerListPermission(c *gin.Context, cover uint8, servers []uint64, ownerUID uint64) error {
if cover == model.CronCoverAll {
denySet := make(map[uint64]bool, len(servers))
for _, id := range servers {
denySet[id] = true
}
if !denyListOwnedByCaller(ownerUID, denySet) {
return singleton.Localizer.ErrorT("permission denied")
}
return nil
}
if !singleton.ServerShared.CheckPermission(c, slices.Values(servers)) {
return singleton.Localizer.ErrorT("permission denied")
}
return nil
}
// checkServiceSkipServerPermission is the service-monitor analogue.
// ServiceCoverAll → SkipServers is a deny-set, only ownership required.
// ServiceCoverIgnoreAll → SkipServers is an allow-set, full Server.HasPermission.
//
// Runtime DispatchTask + skipServersToDenyList only consult entries whose
// bool value is true; false entries are no-ops. Filtering to true-only
// here keeps the write-side permission check aligned with the runtime
// fan-out (a member touching `{2: false}` for a foreign-owned server 2
// has no dispatch effect, so rejecting the request is over-restrictive
// and inconsistent with what listing / runtime see).
func checkServiceSkipServerPermission(c *gin.Context, cover uint8, skip map[uint64]bool, ownerUID uint64) error {
effective := make(map[uint64]bool, len(skip))
for id, enabled := range skip {
if enabled {
effective[id] = true
}
}
if cover == model.ServiceCoverAll {
if !denyListOwnedByCaller(ownerUID, effective) {
return singleton.Localizer.ErrorT("permission denied")
}
return nil
}
ids := make([]uint64, 0, len(effective))
for id := range effective {
ids = append(ids, id)
}
if !singleton.ServerShared.CheckPermission(c, slices.Values(ids)) {
return singleton.Localizer.ErrorT("permission denied")
}
return nil
}
// denyListOwnedByCaller verifies every id in denyList refers to a server
// owned by ownerUID. Under *CoverAll the deny-list expresses exclusion, not
// access, so it must not point at someone else's servers.
//
// Admin owners are special: runtime CronTrigger / DispatchTask fans out
// across the WHOLE system via userIsAdmin(owner), so a safe deny-list for
// an admin-owned resource must be allowed to include foreign-owned servers
// — that's the only way a limited PAT can contain the fan-out. We still
// require each id to refer to a real server, just not to be owned by the
// admin specifically.
func denyListOwnedByCaller(ownerUID uint64, denyList map[uint64]bool) bool {
ownerIsAdmin := model.OwnerIsAdminLookup != nil && model.OwnerIsAdminLookup(ownerUID)
for id := range denyList {
s, found := singleton.ServerShared.Get(id)
if !found || s == nil {
return false
}
if ownerIsAdmin {
continue
}
if s.GetUserID() != ownerUID {
return false
}
}
return true
}
// denyListCoversAllOwnerServersOutsidePATWhitelist reports whether every
// server visible to the cron/service owner that is NOT in the caller PAT's
// server_ids whitelist also appears in denyList. Under *CoverAll semantics
// the runtime dispatch (CronTrigger / DispatchTask) fans out to ServerShared
// minus denyList; the only way a server-limited PAT can stay inside its
// whitelist is if denyList already covers every owner-visible server outside
// that whitelist. Returning true means the configuration is safe.
func denyListCoversAllOwnerServersOutsidePATWhitelist(c *gin.Context, ownerUID uint64, denyList map[uint64]bool) bool {
tok := patAccessorFromContext(c)
if tok == nil {
return true
}
denyIDs := make([]uint64, 0, len(denyList))
for id, mark := range denyList {
if mark {
denyIDs = append(denyIDs, id)
}
}
return model.DenyListSafeForLimitedPAT(tok, ownerUID, denyIDs)
}
// coverMode 抽象「cover 字段在 dispatch 时如何解读 servers 字段」。
//
// 写侧 rejectImplicit* 与运行时 manual/batch-delete 入口共用同一条 PAT 收口
// 路径(assertPATCoverFanoutWithinWhitelist),靠它把两边的规则对齐。新增任
// 何带 cover 概念的资源时,只需在自己的资源专用入口里把 Cover 枚举翻译成
// 这三档之一即可。
type coverMode uint8
const (
// coverModePinnedByCaller: dispatch 阶段不按 servers 字段做 fan-out
// 真实目标在 fire 时由外部信号(如告警触发者 server)钉死。代表:
// CronCoverAlertTrigger。PAT 在这里不做额外收口。
coverModePinnedByCaller coverMode = iota
// coverModeAllMinusDeny: dispatch 时取 owner 全量 server 集合,再减去
// serversdeny-list)。代表 CronCoverAll / ServiceCoverAll。受限 PAT
// 必须确保 deny-list 已覆盖白名单外的全部 owner servers,否则 fan-out
// 会跑到 PAT 白名单之外。
coverModeAllMinusDeny
// coverModeAllowList: dispatch 时只在 serversallow-list)内 fan-out。
// 代表 CronCoverIgnoreAll / ServiceCoverIgnoreAll。受限 PAT 必须能访
// 问 allow-list 中的每一个 server。空 allow-list 是「matches nothing」
// 的退化形态,安全。
coverModeAllowList
)
// assertPATCoverFanoutWithinWhitelist 是 cover-all / cover-ignore-all 两类
// 「按 owner 全量 fan-out」资源的 PAT 收口。
//
// 任何会按「owner servers 减 denyList」或「allowList 自身」展开的资源都必须
// 在 dispatch 入口(manual 触发 / batch-delete / mutation)调用它;写侧
// rejectImplicit* 也走同一条路径,从根上保证两边不漂移。
//
// JWT 请求或不带 server 白名单的 PAT 直接放行——它们没有「白名单」可越过。
//
// 失败时统一返回 i18n "permission denied",与既有写侧 guard 行为一致。
func assertPATCoverFanoutWithinWhitelist(c *gin.Context, ownerUID uint64, mode coverMode, servers []uint64) error {
if !patHasServerWhitelist(c) {
return nil
}
switch mode {
case coverModePinnedByCaller:
return nil
case coverModeAllMinusDeny:
denySet := make(map[uint64]bool, len(servers))
for _, id := range servers {
denySet[id] = true
}
if !denyListCoversAllOwnerServersOutsidePATWhitelist(c, ownerUID, denySet) {
return singleton.Localizer.ErrorT("permission denied")
}
return nil
case coverModeAllowList:
tok := patAccessorFromContext(c)
if tok == nil {
return nil
}
for _, id := range servers {
if !tok.CanAccessServer(id) {
return singleton.Localizer.ErrorT("permission denied")
}
}
return nil
default:
// 未识别 cover 模式按拒绝处理;新增 coverMode 必须显式 wire 到
// 资源专用入口里,不允许沉默放行。
return singleton.Localizer.ErrorT("permission denied")
}
}
// coverModeUnknown 表示 Cron/Service 持久化里出现了当前代码不认识的 cover
// 常量。这一档专门让 assertPATCoverFanoutWithinWhitelist 走 default 分支
// fail-closed,保证「未知 cover 必须显式 wire,否则拒绝」的不变量。
const coverModeUnknown coverMode = 255
// patGroupMembershipAccessAllowed returns false when the caller's PAT
// carries a server_ids whitelist that does not cover every current member
// of groupID. JWT requests and unscoped PATs always pass. Used by
// updateServerGroup before the transactional DELETE+INSERT — otherwise a
// PAT scoped to [X] could indirectly remove server Y from a shared group.
func patGroupMembershipAccessAllowed(c *gin.Context, groupID uint64) bool {
tok := patAccessorFromContext(c)
if tok == nil || !patHasServerWhitelist(c) {
return true
}
var members []model.ServerGroupServer
if err := singleton.DB.Where("server_group_id = ?", groupID).Find(&members).Error; err != nil {
return false
}
for _, m := range members {
if !tok.CanAccessServer(m.ServerId) {
return false
}
}
return true
}
// isValidCronCover reports whether cover is one of the runtime-recognised
// Cron Cover constants. Unknown values must be rejected at write time —
// CronTrigger's periodic scheduler path has no PAT context, so any dirty
// row persisted with an unrecognised Cover still fans out via the default
// branch (no CoverAll/IgnoreAll match → broadcast to every server passing
// cronCanSendToServer). The same allowlist applies for batch-delete and
// manual-trigger guard wiring.
func isValidCronCover(cover uint8) bool {
switch cover {
case model.CronCoverIgnoreAll, model.CronCoverAll, model.CronCoverAlertTrigger:
return true
}
return false
}
// isValidServiceCover is the service-monitor analogue. ServiceCoverAll and
// ServiceCoverIgnoreAll are the only branches DispatchTask + Snapshot
// recognise; anything else degrades to "default fan-out" which silently
// escapes the PAT cover-fanout guard.
func isValidServiceCover(cover uint8) bool {
switch cover {
case model.ServiceCoverAll, model.ServiceCoverIgnoreAll:
return true
}
return false
}
// cronCoverMode 把 model.CronCover* 翻译成共享底座认识的 coverMode。
//
// 未来引入新的 Cron Cover 常量时必须在这里显式 wire,否则
// assertPATCoverFanoutWithinWhitelist 会按 default 分支拒绝,避免悄悄绕过。
func cronCoverMode(cover uint8) coverMode {
switch cover {
case model.CronCoverAll:
return coverModeAllMinusDeny
case model.CronCoverIgnoreAll:
return coverModeAllowList
case model.CronCoverAlertTrigger:
return coverModePinnedByCaller
default:
// 未识别 cover 不能降级成 pinned——pinned 会被 assert 直接放行,
// 让受限 PAT 借未知 cover 绕过 fan-out 收口。统一报告 unknown
// 由 assert 的 default 分支 fail-closed。
return coverModeUnknown
}
}
// serviceCoverMode 是 cronCoverMode 在 service monitor 侧的对照。Service 没
// 有 alert-trigger 这一档,只有 All 与 IgnoreAll。
func serviceCoverMode(cover uint8) coverMode {
switch cover {
case model.ServiceCoverAll:
return coverModeAllMinusDeny
case model.ServiceCoverIgnoreAll:
return coverModeAllowList
default:
// 同 cronCoverMode:未识别 cover 不允许借 pinned 旁路 PAT 收口。
return coverModeUnknown
}
}
// rejectImplicitCoverForLimitedPAT enforces the cover-all PAT guard for the
// cron write path. cf.Servers is the literal allow/deny list; under
// CronCoverAll it is a deny-list, under CronCoverIgnoreAll it is an
// allow-list, and under CronCoverAlertTrigger it does not gate dispatch at
// all (the alert trigger pins the target server at fire time). A PAT that
// carries a server_ids whitelist must therefore either (a) leave the deny-list
// empty under non-CoverAll modes — that's allow-list semantics, safe — or
// (b) under CronCoverAll, supply a deny-list that already covers every
// owner-visible server outside the PAT whitelist, otherwise CronTrigger fans
// out to those servers. Alert triggers stay unrestricted because their
// dispatch boundary is enforced by Cron.HasPermission against the trigger
// server id.
func rejectImplicitCoverForLimitedPAT(c *gin.Context, cover uint8, denyServers []uint64) error {
return rejectImplicitCoverForLimitedPATWithOwner(c, cover, denyServers, getUid(c))
}
// rejectImplicitCoverForLimitedPATWithOwner is the explicit-owner variant
// of rejectImplicitCoverForLimitedPAT. updateCron MUST use this with the
// existing cron's UserID — not the caller — because CronTrigger fans out
// to the cron OWNER's servers at dispatch time, regardless of who issued
// the PATCH. Defaulting to getUid(c) (as rejectImplicitCoverForLimitedPAT
// does for createCron) is only safe when the caller is the owner-to-be,
// i.e. the cron is being created with cr.UserID = getUid(c).
//
// 实现层只是把参数翻译到共享底座 assertPATCoverFanoutWithinWhitelist 上;
// 写侧/运行时入口共用同一裁决,避免两边语义漂移。
func rejectImplicitCoverForLimitedPATWithOwner(c *gin.Context, cover uint8, denyServers []uint64, ownerUID uint64) error {
// 写侧只关心 CronCoverAll 的 deny-list 是否充分——CoverIgnoreAll 的
// allow-list 在 checkCronServerListPermission 已经过 Server.HasPermission
// 收口;CoverAlertTrigger 在 fire 时再校验。保留这条提前 return 与
// 老语义完全一致,避免重复 403。
if cover != model.CronCoverAll {
return nil
}
return assertPATCoverFanoutWithinWhitelist(c, ownerUID, coverModeAllMinusDeny, denyServers)
}
// rejectImplicitServiceCoverForLimitedPAT is the service-monitor analogue.
// ServiceCoverAll treats SkipServers as a deny-set: DispatchTask iterates
// ServerShared.Range and probes every server owned by the service owner that
// is NOT marked true in SkipServers. A server-limited PAT must therefore mark
// every owner-visible server outside its whitelist as skipped.
//
// 同样靠 assertPATCoverFanoutWithinWhitelist 落地,与 cron 写侧/运行时入口
// 共用一条裁决路径。
func rejectImplicitServiceCoverForLimitedPAT(c *gin.Context, cover uint8, skipServers map[uint64]bool, ownerUID uint64) error {
if cover != model.ServiceCoverAll {
return nil
}
denyServers := skipServersToDenyList(skipServers)
return assertPATCoverFanoutWithinWhitelist(c, ownerUID, coverModeAllMinusDeny, denyServers)
}
// skipServersToDenyList 把 service monitor 用的 SkipServers map 展平成
// 共享底座需要的切片形态,并按 true 过滤。写侧/运行时入口共用,避免重复
// 写遍历逻辑。
func skipServersToDenyList(skip map[uint64]bool) []uint64 {
out := make([]uint64, 0, len(skip))
for id, mark := range skip {
if mark {
out = append(out, id)
}
}
return out
}
// enforcePATCronDispatchScope 是 cron 运行时入口(manualTriggerCron /
// batchDeleteCron)的 PAT 收口。把 cr.Cover / cr.Servers 翻译成 coverMode
// 后交给共享底座;语义与写侧 rejectImplicitCoverForLimitedPAT* 严格对齐,
// 闭合「写时拦下 / 运行时回放同一条规则」的不变量,避免历史脏数据 + 受
// 限 PAT 形成越权 fan-out。
func enforcePATCronDispatchScope(c *gin.Context, cr *model.Cron) error {
if cr == nil {
return nil
}
return assertPATCoverFanoutWithinWhitelist(c, cr.GetUserID(), cronCoverMode(cr.Cover), cr.Servers)
}
// enforcePATServiceDispatchScope 是 service monitor 运行时入口
// batchDeleteService 等)的 PAT 收口。SkipServers 是 map[uint64]bool
// 这里展开成 deny-list 切片喂给共享底座;语义与
// rejectImplicitServiceCoverForLimitedPAT 严格对齐。
func enforcePATServiceDispatchScope(c *gin.Context, svc *model.Service) error {
if svc == nil {
return nil
}
return assertPATCoverFanoutWithinWhitelist(c, svc.GetUserID(), serviceCoverMode(svc.Cover), skipServersToDenyList(svc.SkipServers))
}
// enforcePATTriggerTaskScope 阻止 service:write / alertrule:write 的 PAT 通过绑定
// trigger task 越权执行 cron。运行时 alertsentinel/servicesentinel 触发
// CronShared.SendTriggerTasks 时没有 PAT 上下文,CheckPermission 也只校验
// ownership/白名单而非 scope,所以必须在写侧对 PAT 额外要求 ScopeCronExec。
func enforcePATTriggerTaskScope(c *gin.Context, failTasks, recoverTasks []uint64) error {
if len(failTasks) == 0 && len(recoverTasks) == 0 {
return nil
}
tok := APITokenFromContext(c)
if tok == nil {
return nil
}
if !tok.HasScope(model.ScopeCronExec) {
return singleton.Localizer.ErrorT("permission denied")
}
return nil
}
func userCanViewServer(c *gin.Context, server *model.Server) bool {
if server == nil {
return false
}
// PAT 白名单优先于 admin/owner 早返回:admin 自己签发的 server_ids 受限 PAT
// 必须只能看见白名单里的 server,否则给自己设的硬边界形同虚设。
if !patAllowsServer(c, server.GetID()) {
return false
}
if callerIsAdmin(c) {
return true
}
@@ -39,16 +489,18 @@ func userCanViewService(c *gin.Context, service *model.Service) bool {
if service == nil {
return false
}
// EnableShowInService 是显式公开旗标:guest 都可看,PAT 白名单不收窄
// 公开视图(公开 service 本来就不绑特定 server)。其它分支才走 PAT。
if service.EnableShowInService {
return true
}
if callerIsAdmin(c) {
return true
if _, isMember := c.Get(model.CtxKeyAuthorizedUser); !isMember {
return false
}
if _, isMember := c.Get(model.CtxKeyAuthorizedUser); isMember {
return service.HasPermission(c)
}
return false
// 关键:必须先让 Service.HasPermission 跑 PAT 白名单收口,再让 admin
// 身份在没有 PAT 的请求上短路放行。否则 admin 自己签发的 server_ids
// 受限 PAT 会被 admin 早返回直接放过,绕过 list/history 入口的 PAT 边界。
return service.HasPermission(c)
}
func assertOwnsNotificationGroup(c *gin.Context, groupID uint64) error {
@@ -0,0 +1,182 @@
package controller
// 共享底座 assertPATCoverFanoutWithinWhitelist 的单元测试。
//
// 这一层不知道 cron / service,只知道三种 coverMode;测试矩阵覆盖
// {JWT / 无白名单 PAT / 有白名单 PAT × 充分 deny / 不充分 deny / allow-list
// 内 / 越界},钉死「写侧 rejectImplicit* 与运行时 enforce* 必须共用同一裁
// 决路径」这条不变量。任何后续重构改动了规则但忘了同步两侧,这里会先于
// 资源专用入口测试暴露问题。
import (
"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 setupCoverFanoutFixture(t *testing.T) {
t.Helper()
gin.SetMode(gin.TestMode)
ensureLocalizerForStreamTests(t)
originalServer := singleton.ServerShared
sc := singleton.NewEmptyServerClassForTest()
for _, id := range []uint64{1, 2, 3} {
s := &model.Server{}
s.ID = id
s.SetUserID(100)
sc.InsertForTest(s)
}
other := &model.Server{}
other.ID = 9
other.SetUserID(200)
sc.InsertForTest(other)
singleton.ServerShared = sc
t.Cleanup(func() { singleton.ServerShared = originalServer })
}
func ctxWithPAT(t *testing.T, tok *model.APIToken) *gin.Context {
t.Helper()
c, _ := gin.CreateTestContext(nil)
if tok != nil {
c.Set(model.CtxKeyAPIToken, tok)
c.Set(apiTokenCtxKey, tok)
}
return c
}
func TestAssertPATCoverFanout_JWTAlwaysPasses(t *testing.T) {
setupCoverFanoutFixture(t)
c := ctxWithPAT(t, nil)
require.NoError(t, assertPATCoverFanoutWithinWhitelist(c, 100, coverModeAllMinusDeny, nil))
require.NoError(t, assertPATCoverFanoutWithinWhitelist(c, 100, coverModeAllowList, []uint64{2, 3}))
require.NoError(t, assertPATCoverFanoutWithinWhitelist(c, 100, coverModePinnedByCaller, []uint64{2, 3}))
}
func TestAssertPATCoverFanout_UnscopedPATAlwaysPasses(t *testing.T) {
setupCoverFanoutFixture(t)
tok := &model.APIToken{ID: 1, UserID: 100}
c := ctxWithPAT(t, tok)
require.NoError(t, assertPATCoverFanoutWithinWhitelist(c, 100, coverModeAllMinusDeny, nil),
"PAT without server whitelist must not be restricted by cover-fanout guard")
require.NoError(t, assertPATCoverFanoutWithinWhitelist(c, 100, coverModeAllowList, []uint64{2, 3}))
}
func TestAssertPATCoverFanout_AllMinusDeny_RejectsInsufficientDeny(t *testing.T) {
setupCoverFanoutFixture(t)
tok := &model.APIToken{ID: 1, UserID: 100}
tok.SetServerIDs([]uint64{1})
c := ctxWithPAT(t, tok)
err := assertPATCoverFanoutWithinWhitelist(c, 100, coverModeAllMinusDeny, []uint64{1})
assert.Error(t, err, "deny-list covering only whitelisted server 1 still fans out to owner servers 2/3")
err = assertPATCoverFanoutWithinWhitelist(c, 100, coverModeAllMinusDeny, []uint64{2})
assert.Error(t, err, "deny-list missing owner server 3 must be rejected")
}
func TestAssertPATCoverFanout_AllMinusDeny_AcceptsSufficientDeny(t *testing.T) {
setupCoverFanoutFixture(t)
tok := &model.APIToken{ID: 1, UserID: 100}
tok.SetServerIDs([]uint64{1})
c := ctxWithPAT(t, tok)
err := assertPATCoverFanoutWithinWhitelist(c, 100, coverModeAllMinusDeny, []uint64{2, 3})
assert.NoError(t, err, "deny-list covers every owner server outside the PAT whitelist; must pass")
}
func TestAssertPATCoverFanout_AllowList_RejectsOutsideWhitelist(t *testing.T) {
setupCoverFanoutFixture(t)
tok := &model.APIToken{ID: 1, UserID: 100}
tok.SetServerIDs([]uint64{1})
c := ctxWithPAT(t, tok)
err := assertPATCoverFanoutWithinWhitelist(c, 100, coverModeAllowList, []uint64{1, 2})
assert.Error(t, err, "allow-list containing non-whitelisted server 2 must be rejected")
}
func TestAssertPATCoverFanout_AllowList_AcceptsInsideWhitelist(t *testing.T) {
setupCoverFanoutFixture(t)
tok := &model.APIToken{ID: 1, UserID: 100}
tok.SetServerIDs([]uint64{1})
c := ctxWithPAT(t, tok)
require.NoError(t, assertPATCoverFanoutWithinWhitelist(c, 100, coverModeAllowList, []uint64{1}))
require.NoError(t, assertPATCoverFanoutWithinWhitelist(c, 100, coverModeAllowList, nil),
"empty allow-list is the degenerate matches-nothing case; not a bypass")
}
func TestAssertPATCoverFanout_PinnedByCaller_PassesAlways(t *testing.T) {
setupCoverFanoutFixture(t)
tok := &model.APIToken{ID: 1, UserID: 100}
tok.SetServerIDs([]uint64{1})
c := ctxWithPAT(t, tok)
require.NoError(t, assertPATCoverFanoutWithinWhitelist(c, 100, coverModePinnedByCaller, []uint64{2, 3}),
"alert-trigger dispatch pins the target server at fire time; assertPATCoverFanoutWithinWhitelist must not pre-judge")
}
func TestCronCoverMode_KnownValues(t *testing.T) {
assert.Equal(t, coverModeAllMinusDeny, cronCoverMode(model.CronCoverAll))
assert.Equal(t, coverModeAllowList, cronCoverMode(model.CronCoverIgnoreAll))
assert.Equal(t, coverModePinnedByCaller, cronCoverMode(model.CronCoverAlertTrigger))
}
func TestServiceCoverMode_KnownValues(t *testing.T) {
assert.Equal(t, coverModeAllMinusDeny, serviceCoverMode(model.ServiceCoverAll))
assert.Equal(t, coverModeAllowList, serviceCoverMode(model.ServiceCoverIgnoreAll))
}
func TestSkipServersToDenyList_FiltersOnlyTrue(t *testing.T) {
got := skipServersToDenyList(map[uint64]bool{1: true, 2: false, 3: true})
assert.ElementsMatch(t, []uint64{1, 3}, got,
"only true entries are real skips; false-valued entries must not be promoted to deny-list")
}
// 资源专用入口在底座上薄包装的契约:cron-runtime 与 service-runtime 必须
// 调底座,因此底座在「不充分 deny-list」时返回的 error 必须穿透到入口。
func TestEnforcePATCronDispatchScope_RelaysBaseDecision(t *testing.T) {
setupCoverFanoutFixture(t)
tok := &model.APIToken{ID: 1, UserID: 100}
tok.SetServerIDs([]uint64{1})
c := ctxWithPAT(t, tok)
cr := &model.Cron{
Common: model.Common{UserID: 100},
Cover: model.CronCoverAll,
Servers: []uint64{1},
}
err := enforcePATCronDispatchScope(c, cr)
assert.Error(t, err, "cover-all cron whose deny-list only covers whitelisted server must be rejected")
cr.Servers = []uint64{2, 3}
require.NoError(t, enforcePATCronDispatchScope(c, cr),
"deny-list covering every non-whitelisted owner server must pass")
}
func TestEnforcePATServiceDispatchScope_RelaysBaseDecision(t *testing.T) {
setupCoverFanoutFixture(t)
tok := &model.APIToken{ID: 1, UserID: 100}
tok.SetServerIDs([]uint64{1})
c := ctxWithPAT(t, tok)
svc := &model.Service{
Common: model.Common{UserID: 100},
Cover: model.ServiceCoverAll,
SkipServers: map[uint64]bool{1: true},
}
err := enforcePATServiceDispatchScope(c, svc)
assert.Error(t, err, "cover-all service whose SkipServers only marks whitelisted servers must be rejected")
svc.SkipServers = map[uint64]bool{2: true, 3: true}
require.NoError(t, enforcePATServiceDispatchScope(c, svc),
"SkipServers covering every non-whitelisted owner server must pass")
}
+202
View File
@@ -0,0 +1,202 @@
package controller
import (
"bytes"
"encoding/json"
"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"
)
// setupRESTScopeTest 准备一个 PAT + 一个最小路由表,用于测 REST scope enforce。
func setupRESTScopeTest(t *testing.T) (*httptest.Server, *model.APIToken, string, func()) {
t.Helper()
cleanupBase, uid := setupMCPTest(t)
tok, plain := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
gin.SetMode(gin.TestMode)
r := gin.New()
patMw := apiTokenAuthMiddleware()
r.GET("/server",
patMw,
restScopeMiddleware(model.ScopeServerRead),
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
)
r.POST("/server/config",
patMw,
restScopeMiddleware(model.ScopeServerWrite),
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
)
r.POST("/server-group",
patMw,
restScopeMiddleware(model.ScopeServerWrite),
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
)
r.GET("/api-tokens",
patMw,
restPATForbiddenMiddleware(),
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
)
ts := httptest.NewServer(r)
return ts, tok, plain, func() {
ts.Close()
cleanupBase()
}
}
func doReq(t *testing.T, ts *httptest.Server, method, path, token string) *http.Response {
t.Helper()
req, _ := http.NewRequest(method, ts.URL+path, bytes.NewReader([]byte("{}")))
if token != "" {
req.Header.Set("Authorization", "Bearer "+token)
}
req.Header.Set("Content-Type", "application/json")
resp, err := http.DefaultClient.Do(req)
require.NoError(t, err)
return resp
}
func TestREST_PATWithMatchingScopeAllowed(t *testing.T) {
ts, _, tok, cleanup := setupRESTScopeTest(t)
defer cleanup()
resp := doReq(t, ts, "GET", "/server", tok)
defer resp.Body.Close()
require.Equal(t, http.StatusOK, resp.StatusCode)
}
func TestREST_PATWithoutScopeDenied(t *testing.T) {
ts, _, tok, cleanup := setupRESTScopeTest(t)
defer cleanup()
resp := doReq(t, ts, "POST", "/server/config", tok)
defer resp.Body.Close()
require.Equal(t, http.StatusForbidden, resp.StatusCode)
var body model.CommonResponse[any]
require.NoError(t, json.NewDecoder(resp.Body).Decode(&body))
require.False(t, body.Success)
require.Contains(t, body.Error, "nezha:server:write")
}
func TestREST_SelfManagementForbidsPAT(t *testing.T) {
ts, _, tok, cleanup := setupRESTScopeTest(t)
defer cleanup()
resp := doReq(t, ts, "GET", "/api-tokens", tok)
defer resp.Body.Close()
require.Equal(t, http.StatusForbidden, resp.StatusCode)
}
func TestREST_PATWildcardCoversAllVerbs(t *testing.T) {
cleanupBase, uid := setupMCPTest(t)
defer cleanupBase()
tok, plain := mkToken(t, uid, []string{"nezha:server:*"}, nil)
_ = tok
gin.SetMode(gin.TestMode)
r := gin.New()
patMw := apiTokenAuthMiddleware()
r.GET("/server", patMw, restScopeMiddleware(model.ScopeServerRead),
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) })
r.POST("/server/config", patMw, restScopeMiddleware(model.ScopeServerWrite),
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) })
r.POST("/batch-delete/server", patMw, restScopeMiddleware(model.ScopeServerDelete),
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) })
ts := httptest.NewServer(r)
defer ts.Close()
for _, tc := range []struct {
method, path string
}{
{"GET", "/server"},
{"POST", "/server/config"},
{"POST", "/batch-delete/server"},
} {
resp := doReq(t, ts, tc.method, tc.path, plain)
resp.Body.Close()
require.Equal(t, http.StatusOK, resp.StatusCode, "%s %s should be allowed by nezha:server:*", tc.method, tc.path)
}
}
func TestREST_NezhaAllGrantsEverything(t *testing.T) {
cleanupBase, uid := setupMCPTest(t)
defer cleanupBase()
_, plain := mkToken(t, uid, []string{model.ScopeNezhaAll}, nil)
gin.SetMode(gin.TestMode)
r := gin.New()
r.POST("/maintenance",
apiTokenAuthMiddleware(),
restScopeMiddleware(model.ScopeAdminAll),
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
)
ts := httptest.NewServer(r)
defer ts.Close()
resp := doReq(t, ts, "POST", "/maintenance", plain)
defer resp.Body.Close()
require.Equal(t, http.StatusOK, resp.StatusCode)
}
func TestREST_NoAuthGoesToJWTChain(t *testing.T) {
cleanupBase, _ := setupMCPTest(t)
defer cleanupBase()
jwtCalled := false
fakeJwt := func(c *gin.Context) {
jwtCalled = true
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "no jwt"})
}
gin.SetMode(gin.TestMode)
r := gin.New()
r.GET("/server",
jwtOrPATAuthMiddleware(apiTokenAuthMiddleware(), fakeJwt),
restScopeMiddleware(model.ScopeServerRead),
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
)
ts := httptest.NewServer(r)
defer ts.Close()
resp := doReq(t, ts, "GET", "/server", "")
resp.Body.Close()
require.True(t, jwtCalled, "JWT mw must be invoked when no PAT")
require.Equal(t, http.StatusUnauthorized, resp.StatusCode)
}
func TestREST_BadPATShortCircuitsBeforeJWT(t *testing.T) {
cleanupBase, _ := setupMCPTest(t)
defer cleanupBase()
jwtCalled := false
fakeJwt := func(c *gin.Context) { jwtCalled = true }
gin.SetMode(gin.TestMode)
r := gin.New()
r.Use(func(c *gin.Context) {
c.Set(model.CtxKeyRealIPStr, "203.0.113.99")
c.Next()
})
r.GET("/server",
jwtOrPATAuthMiddleware(apiTokenAuthMiddleware(), fakeJwt),
restScopeMiddleware(model.ScopeServerRead),
func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) },
)
ts := httptest.NewServer(r)
defer ts.Close()
resp := doReq(t, ts, "GET", "/server", "nzp_bogus_token_value")
resp.Body.Close()
require.False(t, jwtCalled, "JWT mw must NOT run after bad PAT abort")
require.Equal(t, http.StatusUnauthorized, resp.StatusCode)
var blocked model.WAF
err := singleton.DB.Where("block_identifier = ?", model.BlockIDToken).First(&blocked).Error
require.NoError(t, err, "bad PAT must trigger WAF BlockIP")
require.GreaterOrEqual(t, blocked.Count, uint64(1))
}
@@ -0,0 +1,63 @@
package controller
import (
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/nezhahq/nezha/model"
)
// H3 regression: file-manager sessions read/write/delete files, but the route
// only requires nezha:server:write. PAT scopes are advertised as fine-grained
// (read / write / delete / exec); allowing a write-only PAT to open an FM
// session that can list & remove files silently widens the scope.
func TestRestScopeAllOf_RequiresEveryScope(t *testing.T) {
mw := restScopeAllOf(model.ScopeServerRead, model.ScopeServerWrite, model.ScopeServerDelete)
t.Run("rejects_token_missing_delete", func(t *testing.T) {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
tok := &model.APIToken{ScopesCSV: "nezha:server:read,nezha:server:write"}
c.Set(apiTokenCtxKey, tok)
c.Set(model.CtxKeyAPIToken, tok)
mw(c)
if !c.IsAborted() || w.Code != 403 {
t.Fatalf("missing delete scope must abort with 403, got aborted=%v code=%d", c.IsAborted(), w.Code)
}
})
t.Run("accepts_token_with_all_scopes", func(t *testing.T) {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
tok := &model.APIToken{ScopesCSV: "nezha:server:read,nezha:server:write,nezha:server:delete"}
c.Set(apiTokenCtxKey, tok)
c.Set(model.CtxKeyAPIToken, tok)
mw(c)
if c.IsAborted() {
t.Fatal("token carrying all required scopes must pass")
}
})
t.Run("jwt_callers_skip_check", func(t *testing.T) {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
mw(c)
if c.IsAborted() {
t.Fatal("JWT (no PAT) must pass through restScopeAllOf unchanged")
}
})
t.Run("wildcard_resource_scope_satisfies_all", func(t *testing.T) {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
tok := &model.APIToken{ScopesCSV: "nezha:server:*"}
c.Set(apiTokenCtxKey, tok)
c.Set(model.CtxKeyAPIToken, tok)
mw(c)
if c.IsAborted() {
t.Fatal("nezha:server:* must satisfy server:read+write+delete")
}
})
}
+124
View File
@@ -0,0 +1,124 @@
// Package controller — scope reference table for REST + MCP.
//
// Each REST endpoint under /api/v1/* and each MCP tool under /mcp requires a
// specific scope when authenticated via PAT (`Authorization: Bearer nzp_*`).
// JWT-authenticated requests skip scope enforcement.
//
// This file is the authoritative human + LLM-readable index. The actual
// enforcement lives in controller.go (REST) and mcp_tools_*.go (MCP). When
// you change an endpoint's scope requirement, update this table.
//
// # Scope naming
//
// nezha:{resource}:{verb}
// resource: server | service | alertrule | cron | ddns | nat |
// notification | notification-group | transfer | admin
// verb: read | write | delete | exec
//
// nezha:* Admin-only superuser
// nezha:admin:* Admin-only user/waf/setting/online-user management
// nezha:<res>:* All actions on a resource
//
// # MCP tools (POST /mcp tools/call)
//
// meta.whoami — (any scope)
// server.list nezha:server:read
// server.get nezha:server:read
// server.exec nezha:server:exec
// fs.list nezha:server:read
// fs.read nezha:server:read
// fs.write nezha:server:write
// fs.delete nezha:server:delete
// fs.download_url nezha:server:read
// fs.upload_url nezha:server:write
//
// # REST endpoints (PAT required scope)
//
// GET /api/v1/server nezha:server:read
// PATCH /api/v1/server/{id} nezha:server:write
// GET /api/v1/server/config/{id} nezha:server:write
// POST /api/v1/server/config nezha:server:write
// POST /api/v1/batch-delete/server nezha:server:delete
// POST /api/v1/batch-move/server nezha:server:write
// POST /api/v1/force-update/server nezha:server:write
// POST /api/v1/server-group nezha:server:write
// PATCH /api/v1/server-group/{id} nezha:server:write
// POST /api/v1/batch-delete/server-group nezha:server:delete
// POST /api/v1/terminal nezha:server:exec
// GET /api/v1/ws/terminal/{id} nezha:server:exec
// POST /api/v1/file nezha:server:write
// GET /api/v1/ws/file/{id} nezha:server:write
// GET /api/v1/ws/server nezha:server:read
// GET /api/v1/server-group nezha:server:read
// GET /api/v1/service nezha:service:read
// GET /api/v1/service/server nezha:service:read
// GET /api/v1/service/{id}/history nezha:service:read
// GET /api/v1/server/{id}/service nezha:service:read
// GET /api/v1/server/{id}/metrics nezha:server:read
//
// GET /api/v1/transfer nezha:transfer:read
// POST /api/v1/transfer/{id}/cancel nezha:transfer:write
// POST /api/v1/transfer/{id}/retry nezha:transfer:write
// GET /api/v1/ws/transfer nezha:transfer:read
//
// GET /api/v1/service/list nezha:service:read
// POST /api/v1/service nezha:service:write
// PATCH /api/v1/service/{id} nezha:service:write
// POST /api/v1/batch-delete/service nezha:service:delete
//
// GET /api/v1/alert-rule nezha:alertrule:read
// POST /api/v1/alert-rule nezha:alertrule:write
// PATCH /api/v1/alert-rule/{id} nezha:alertrule:write
// POST /api/v1/batch-delete/alert-rule nezha:alertrule:delete
//
// GET /api/v1/cron nezha:cron:read
// POST /api/v1/cron nezha:cron:write
// PATCH /api/v1/cron/{id} nezha:cron:write
// POST /api/v1/cron/{id}/manual nezha:cron:exec
// POST /api/v1/batch-delete/cron nezha:cron:delete
//
// GET /api/v1/ddns nezha:ddns:read
// GET /api/v1/ddns/providers nezha:ddns:read
// POST /api/v1/ddns nezha:ddns:write
// PATCH /api/v1/ddns/{id} nezha:ddns:write
// POST /api/v1/batch-delete/ddns nezha:ddns:delete
//
// GET /api/v1/nat nezha:nat:read
// POST /api/v1/nat nezha:nat:write
// PATCH /api/v1/nat/{id} nezha:nat:write
// POST /api/v1/batch-delete/nat nezha:nat:delete
//
// GET /api/v1/notification nezha:notification:read
// POST /api/v1/notification nezha:notification:write
// PATCH /api/v1/notification/{id} nezha:notification:write
// POST /api/v1/batch-delete/notification nezha:notification:delete
//
// GET /api/v1/notification-group nezha:notification-group:read
// POST /api/v1/notification-group nezha:notification-group:write
// PATCH /api/v1/notification-group/{id} nezha:notification-group:write
// POST /api/v1/batch-delete/notification-group nezha:notification-group:delete
//
// GET /api/v1/user nezha:admin:*
// POST /api/v1/user nezha:admin:*
// POST /api/v1/batch-delete/user nezha:admin:*
// GET /api/v1/waf nezha:admin:*
// POST /api/v1/batch-delete/waf nezha:admin:*
// GET /api/v1/online-user nezha:admin:*
// POST /api/v1/online-user/batch-block nezha:admin:*
// PATCH /api/v1/setting nezha:admin:*
// POST /api/v1/maintenance nezha:admin:*
//
// # Endpoints permanently forbidden to PAT
//
// These are personal-account-management endpoints; a PAT must never call them
// (would allow self-elevation chains: PAT → mint stronger PAT → ...).
// `restPATForbiddenMiddleware` returns 403 to PAT-authenticated requests.
//
// POST /api/v1/refresh-token
// GET /api/v1/profile
// POST /api/v1/profile
// POST /api/v1/oauth2/{provider}/unbind
// GET /api/v1/api-tokens
// POST /api/v1/api-tokens
// DELETE /api/v1/api-tokens/{id}
package controller
@@ -0,0 +1,161 @@
package controller
// Pins the human-readable PAT scope table in scope_doc.go to the actual
// REST routes registered in controller.go. Without this, the table drifts
// silently every time a route is added/changed — and the table is the
// source LLM clients (and our frontend SCOPE_OPTIONS copy) read from.
//
// The check is intentionally textual: scope_doc.go is a doc-only file with
// no runtime hooks, and routers() bakes scopes into closures at boot, so
// there is no cheap way to reflect them at test time without an invasive
// refactor. Instead we maintain a single canonical (method, path, scope)
// list here and assert both directions:
// - every entry appears verbatim in scope_doc.go
// - every scope-bearing line in scope_doc.go appears in the table
// Adding a new scoped route must update both files, and forgetting either
// is a compile-on-demand failure.
import (
"os"
"regexp"
"strings"
"testing"
"github.com/stretchr/testify/require"
)
type scopedRoute struct {
Method string
Path string
Scope string
}
func canonicalRoutes() []scopedRoute {
return []scopedRoute{
{"GET", "/api/v1/server", "nezha:server:read"},
{"PATCH", "/api/v1/server/{id}", "nezha:server:write"},
{"GET", "/api/v1/server/config/{id}", "nezha:server:write"},
{"POST", "/api/v1/server/config", "nezha:server:write"},
{"POST", "/api/v1/batch-delete/server", "nezha:server:delete"},
{"POST", "/api/v1/batch-move/server", "nezha:server:write"},
{"POST", "/api/v1/force-update/server", "nezha:server:write"},
{"POST", "/api/v1/server-group", "nezha:server:write"},
{"PATCH", "/api/v1/server-group/{id}", "nezha:server:write"},
{"POST", "/api/v1/batch-delete/server-group", "nezha:server:delete"},
{"POST", "/api/v1/terminal", "nezha:server:exec"},
{"GET", "/api/v1/ws/terminal/{id}", "nezha:server:exec"},
{"POST", "/api/v1/file", "nezha:server:write"},
{"GET", "/api/v1/ws/file/{id}", "nezha:server:write"},
// optional-auth scoped routescontroller.go:91-98)。这些 GET 端点既支持
// 未登录访客,也接受 PAT;当走 PAT 路径时 restScopeMiddleware 会强制对应的
// read scope。漏掉这一段会让 scope_doc.go 与实际 router 漂移而测试不报错。
{"GET", "/api/v1/ws/server", "nezha:server:read"},
{"GET", "/api/v1/server-group", "nezha:server:read"},
{"GET", "/api/v1/service", "nezha:service:read"},
{"GET", "/api/v1/service/server", "nezha:service:read"},
{"GET", "/api/v1/service/{id}/history", "nezha:service:read"},
{"GET", "/api/v1/server/{id}/service", "nezha:service:read"},
{"GET", "/api/v1/server/{id}/metrics", "nezha:server:read"},
{"GET", "/api/v1/transfer", "nezha:transfer:read"},
{"POST", "/api/v1/transfer/{id}/cancel", "nezha:transfer:write"},
{"POST", "/api/v1/transfer/{id}/retry", "nezha:transfer:write"},
{"GET", "/api/v1/ws/transfer", "nezha:transfer:read"},
{"GET", "/api/v1/service/list", "nezha:service:read"},
{"POST", "/api/v1/service", "nezha:service:write"},
{"PATCH", "/api/v1/service/{id}", "nezha:service:write"},
{"POST", "/api/v1/batch-delete/service", "nezha:service:delete"},
{"GET", "/api/v1/alert-rule", "nezha:alertrule:read"},
{"POST", "/api/v1/alert-rule", "nezha:alertrule:write"},
{"PATCH", "/api/v1/alert-rule/{id}", "nezha:alertrule:write"},
{"POST", "/api/v1/batch-delete/alert-rule", "nezha:alertrule:delete"},
{"GET", "/api/v1/cron", "nezha:cron:read"},
{"POST", "/api/v1/cron", "nezha:cron:write"},
{"PATCH", "/api/v1/cron/{id}", "nezha:cron:write"},
{"POST", "/api/v1/cron/{id}/manual", "nezha:cron:exec"},
{"POST", "/api/v1/batch-delete/cron", "nezha:cron:delete"},
{"GET", "/api/v1/ddns", "nezha:ddns:read"},
{"GET", "/api/v1/ddns/providers", "nezha:ddns:read"},
{"POST", "/api/v1/ddns", "nezha:ddns:write"},
{"PATCH", "/api/v1/ddns/{id}", "nezha:ddns:write"},
{"POST", "/api/v1/batch-delete/ddns", "nezha:ddns:delete"},
{"GET", "/api/v1/nat", "nezha:nat:read"},
{"POST", "/api/v1/nat", "nezha:nat:write"},
{"PATCH", "/api/v1/nat/{id}", "nezha:nat:write"},
{"POST", "/api/v1/batch-delete/nat", "nezha:nat:delete"},
{"GET", "/api/v1/notification", "nezha:notification:read"},
{"POST", "/api/v1/notification", "nezha:notification:write"},
{"PATCH", "/api/v1/notification/{id}", "nezha:notification:write"},
{"POST", "/api/v1/batch-delete/notification", "nezha:notification:delete"},
{"GET", "/api/v1/notification-group", "nezha:notification-group:read"},
{"POST", "/api/v1/notification-group", "nezha:notification-group:write"},
{"PATCH", "/api/v1/notification-group/{id}", "nezha:notification-group:write"},
{"POST", "/api/v1/batch-delete/notification-group", "nezha:notification-group:delete"},
{"GET", "/api/v1/user", "nezha:admin:*"},
{"POST", "/api/v1/user", "nezha:admin:*"},
{"POST", "/api/v1/batch-delete/user", "nezha:admin:*"},
{"GET", "/api/v1/waf", "nezha:admin:*"},
{"POST", "/api/v1/batch-delete/waf", "nezha:admin:*"},
{"GET", "/api/v1/online-user", "nezha:admin:*"},
{"POST", "/api/v1/online-user/batch-block", "nezha:admin:*"},
{"PATCH", "/api/v1/setting", "nezha:admin:*"},
{"POST", "/api/v1/maintenance", "nezha:admin:*"},
}
}
var scopeDocLineRE = regexp.MustCompile(`^(GET|POST|PATCH|DELETE|PUT)\s+(/api/v1/\S+)\s+(nezha:\S+)$`)
func extractScopeDocEntries(t *testing.T) map[string]scopedRoute {
t.Helper()
raw, err := os.ReadFile("scope_doc.go")
require.NoError(t, err)
entries := map[string]scopedRoute{}
for _, line := range strings.Split(string(raw), "\n") {
stripped := strings.TrimPrefix(line, "//")
stripped = strings.TrimSpace(stripped)
stripped = strings.Join(strings.Fields(stripped), " ")
m := scopeDocLineRE.FindStringSubmatch(stripped)
if m == nil {
continue
}
r := scopedRoute{Method: m[1], Path: m[2], Scope: m[3]}
entries[r.Method+" "+r.Path] = r
}
return entries
}
func TestScopeDocMatchesCanonicalRoutes(t *testing.T) {
doc := extractScopeDocEntries(t)
for _, want := range canonicalRoutes() {
key := want.Method + " " + want.Path
got, ok := doc[key]
if !ok {
t.Errorf("scope_doc.go missing entry: %s %s (expected scope %s)", want.Method, want.Path, want.Scope)
continue
}
if got.Scope != want.Scope {
t.Errorf("scope_doc.go scope mismatch for %s %s: doc=%s code=%s", want.Method, want.Path, got.Scope, want.Scope)
}
}
}
func TestCanonicalRoutesCoverScopeDoc(t *testing.T) {
doc := extractScopeDocEntries(t)
canonical := map[string]scopedRoute{}
for _, r := range canonicalRoutes() {
canonical[r.Method+" "+r.Path] = r
}
for key, entry := range doc {
if _, ok := canonical[key]; !ok {
t.Errorf("scope_doc.go has %s %s (scope %s) with no canonical route — stale doc or missing test entry", entry.Method, entry.Path, entry.Scope)
}
}
}
+34 -15
View File
@@ -21,8 +21,9 @@ import (
// List server
// @Summary List server
// @Security BearerAuth
// @Security APITokenAuth
// @Schemes
// @Description List server
// @Description List server. PAT scope required: nezha:server:read.
// @Tags auth required
// @Param id query uint false "Resource ID"
// @Produce json
@@ -196,11 +197,15 @@ func forceUpdateServer(c *gin.Context) (*model.ServerTaskResponse, error) {
forceUpdateResp.Offline = append(forceUpdateResp.Offline, sid)
continue
}
if stream := server.GetTaskStream(); stream != nil {
if err := stream.Send(&pb.Task{
if server.GetTaskStream() != nil {
if err := server.SendTask(&pb.Task{
Type: model.TaskTypeUpgrade,
}); err != nil {
forceUpdateResp.Failure = append(forceUpdateResp.Failure, sid)
if errors.Is(err, model.ErrTaskStreamOffline) {
forceUpdateResp.Offline = append(forceUpdateResp.Offline, sid)
} else {
forceUpdateResp.Failure = append(forceUpdateResp.Failure, sid)
}
} else {
forceUpdateResp.Success = append(forceUpdateResp.Success, sid)
}
@@ -232,18 +237,19 @@ func getServerConfig(c *gin.Context) (string, error) {
if !ok {
return "", nil
}
stream := s.GetTaskStream()
if stream == nil {
return "", nil
}
if !s.HasPermission(c) {
return "", singleton.Localizer.ErrorT("permission denied")
}
if s.GetTaskStream() == nil {
return "", nil
}
if err := stream.Send(&pb.Task{
if err := s.SendTask(&pb.Task{
Type: model.TaskTypeReportConfig,
}); err != nil {
if errors.Is(err, model.ErrTaskStreamOffline) {
return "", nil
}
return "", err
}
@@ -308,21 +314,23 @@ func setServerConfig(c *gin.Context) (*model.ServerTaskResponse, error) {
go func(srvGroup []*model.Server) {
defer wg.Done()
for _, s := range srvGroup {
// Create and send the task.
task := &pb.Task{
Type: model.TaskTypeApplyConfig,
Data: configForm.Config,
}
stream := s.GetTaskStream()
if stream == nil {
if s.GetTaskStream() == nil {
respMu.Lock()
resp.Offline = append(resp.Offline, s.ID)
respMu.Unlock()
continue
}
if err := stream.Send(task); err != nil {
if err := s.SendTask(task); err != nil {
respMu.Lock()
resp.Failure = append(resp.Failure, s.ID)
if errors.Is(err, model.ErrTaskStreamOffline) {
resp.Offline = append(resp.Offline, s.ID)
} else {
resp.Failure = append(resp.Failure, s.ID)
}
respMu.Unlock()
continue
}
@@ -393,6 +401,17 @@ func batchMoveServer(c *gin.Context) ([]model.BatchMoveServerResult, error) {
continue
}
// PAT server_ids 白名单优先于 admin/owner 早返回:admin 给自己签发的
// 限定 server_ids PAT 必须只能 move 白名单内 server。前面的 admin/owner
// 检查只看 currentOwner,不会触达白名单,这里显式补一道。返回
// ServerNotFound 与未知/外部 server 的语义对齐,避免泄露白名单外
// server 是否存在。
if !patAllowsServer(c, sid) {
res.Status = model.BatchMoveServerResultServerNotFound
results = append(results, res)
continue
}
// Per-server permission: admin or current owner. We do NOT use the
// bulk CheckPermission because we want a partial-success response
// rather than rejecting the whole batch on the first unauthorized id.
+24
View File
@@ -28,6 +28,8 @@ func listServerGroup(c *gin.Context) ([]*model.ServerGroupResponseItem, error) {
_, isMember := c.Get(model.CtxKeyAuthorizedUser)
isAdmin := isMember && callerIsAdmin(c)
pat := patAccessorFromContext(c)
patLimited := pat != nil && patHasServerWhitelist(c)
visibleServerIDs := make(map[uint64]struct{})
if !isMember {
@@ -47,6 +49,9 @@ func listServerGroup(c *gin.Context) ([]*model.ServerGroupResponseItem, error) {
continue
}
}
if pat != nil && !pat.CanAccessServer(s.ServerId) {
continue
}
if _, ok := groupServers[s.ServerGroupId]; !ok {
groupServers[s.ServerGroupId] = make([]uint64, 0)
}
@@ -61,6 +66,9 @@ func listServerGroup(c *gin.Context) ([]*model.ServerGroupResponseItem, error) {
if !isMember && len(groupServers[s.ID]) == 0 {
continue
}
if patLimited && len(groupServers[s.ID]) == 0 {
continue
}
sgRes = append(sgRes, &model.ServerGroupResponseItem{
Group: s,
Servers: groupServers[s.ID],
@@ -169,6 +177,10 @@ func updateServerGroup(c *gin.Context) (any, error) {
return nil, singleton.Localizer.ErrorT("unauthorized")
}
if !patGroupMembershipAccessAllowed(c, sgDB.ID) {
return nil, singleton.Localizer.ErrorT("permission denied")
}
sgDB.Name = sg.Name
var count int64
@@ -237,6 +249,18 @@ func batchDeleteServerGroup(c *gin.Context) (any, error) {
}
}
if pat := patAccessorFromContext(c); pat != nil && patHasServerWhitelist(c) {
var members []model.ServerGroupServer
if err := singleton.DB.Where("server_group_id in (?)", sgs).Find(&members).Error; err != nil {
return nil, err
}
for _, m := range members {
if !pat.CanAccessServer(m.ServerId) {
return nil, singleton.Localizer.ErrorT("permission denied")
}
}
}
err := singleton.DB.Transaction(func(tx *gorm.DB) error {
if err := tx.Unscoped().Delete(&model.ServerGroup{}, "id in (?)", sgs).Error; err != nil {
return err
@@ -0,0 +1,117 @@
package controller
import (
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"github.com/nezhahq/nezha/model"
"github.com/nezhahq/nezha/service/singleton"
)
// H1 regression: updateServerGroup must reject a server-limited PAT whose
// whitelist does not cover the group's CURRENT membership. Today the
// handler only checks the incoming sg.Servers list and then unconditionally
// `DELETE FROM server_group_server WHERE server_group_id = ?`, so a PAT
// scoped to [X] can remove server Y (owned by another tenant or just
// outside the whitelist) from a group it shares with X.
func TestPatHasGroupMembershipAccess_DeniesGroupContainingOutsideServer(t *testing.T) {
db := newTestDB(t)
swap := swapSingletonDB(t, db)
defer swap()
if err := db.Create(&model.ServerGroupServer{
Common: model.Common{ID: 1, UserID: 1},
ServerGroupId: 42,
ServerId: 9, // outside whitelist
}).Error; err != nil {
t.Fatal(err)
}
if err := db.Create(&model.ServerGroupServer{
Common: model.Common{ID: 2, UserID: 1},
ServerGroupId: 42,
ServerId: 1, // inside whitelist
}).Error; err != nil {
t.Fatal(err)
}
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin})
ctx.Set(model.CtxKeyAPIToken, &model.APIToken{ServersCSV: "1"})
if patGroupMembershipAccessAllowed(ctx, 42) {
t.Fatal("PAT scoped to [1] must NOT be allowed to mutate a group whose current members include server 9; " +
"transactional DELETE+INSERT would drop server 9 from the group")
}
}
func TestPatHasGroupMembershipAccess_AllowsGroupFullyInsideWhitelist(t *testing.T) {
db := newTestDB(t)
swap := swapSingletonDB(t, db)
defer swap()
if err := db.Create(&model.ServerGroupServer{
Common: model.Common{ID: 1, UserID: 1},
ServerGroupId: 7,
ServerId: 1,
}).Error; err != nil {
t.Fatal(err)
}
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin})
ctx.Set(model.CtxKeyAPIToken, &model.APIToken{ServersCSV: "1,2"})
if !patGroupMembershipAccessAllowed(ctx, 7) {
t.Fatal("PAT whitelist [1,2] covers all current members → must allow update")
}
}
func TestPatHasGroupMembershipAccess_JWTAlwaysAllowed(t *testing.T) {
db := newTestDB(t)
swap := swapSingletonDB(t, db)
defer swap()
if err := db.Create(&model.ServerGroupServer{
Common: model.Common{ID: 1, UserID: 1},
ServerGroupId: 100,
ServerId: 99,
}).Error; err != nil {
t.Fatal(err)
}
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin})
if !patGroupMembershipAccessAllowed(ctx, 100) {
t.Fatal("JWT requests (no PAT) must always pass — the existing admin/owner check stands")
}
}
func newTestDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(&model.ServerGroupServer{}); err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
sqlDB, _ := db.DB()
if sqlDB != nil {
_ = sqlDB.Close()
}
})
return db
}
func swapSingletonDB(t *testing.T, db *gorm.DB) func() {
t.Helper()
original := singleton.DB
singleton.DB = db
return func() { singleton.DB = original }
}
@@ -1,6 +1,9 @@
package controller
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
@@ -120,3 +123,80 @@ func TestListServerGroupAdminSeesAllGroupsIncludingEmpty(t *testing.T) {
assert.ElementsMatch(t, []string{"Public Group", "Empty Group"}, names,
"admin must keep full visibility, including empty groups")
}
func newServerGroupCtxWithPAT(viewer *model.User, tok *model.APIToken) *gin.Context {
c := newServerGroupCtx(viewer)
if tok != nil {
c.Set(model.CtxKeyAPIToken, tok)
}
return c
}
// PAT scoped to server_ids must hide groups whose membership is entirely
// outside the whitelist and must strip out-of-whitelist server IDs from
// remaining groups. Otherwise admin-issued limited PATs still enumerate
// every group name + server id via /api/v1/server-group.
func TestListServerGroupPATWhitelistFiltersGroupsAndServerIDs(t *testing.T) {
setupServerGroupVisibilityFixture(t)
require.NoError(t, singleton.DB.Create(&model.ServerGroupServer{
Common: model.Common{UserID: 1}, ServerGroupId: 10, ServerId: 2,
}).Error)
tok := &model.APIToken{ID: 77, UserID: 1}
tok.SetServerIDs([]uint64{1})
items, err := listServerGroup(newServerGroupCtxWithPAT(&model.User{
Common: model.Common{ID: 1}, Role: model.RoleAdmin,
}, tok))
require.NoError(t, err)
names := collectGroupNames(items)
assert.ElementsMatch(t, []string{"Public Group"}, names,
"PAT scoped to {1} must drop the empty group and not surface group names containing only server 2")
if assert.Len(t, items, 1) {
assert.ElementsMatch(t, []uint64{1}, items[0].Servers,
"server IDs outside the PAT whitelist must be redacted from the response")
}
}
func TestListServerGroupPATWithDisjointWhitelistReturnsEmpty(t *testing.T) {
setupServerGroupVisibilityFixture(t)
tok := &model.APIToken{ID: 78, UserID: 1}
tok.SetServerIDs([]uint64{9999})
items, err := listServerGroup(newServerGroupCtxWithPAT(&model.User{
Common: model.Common{ID: 1}, Role: model.RoleAdmin,
}, tok))
require.NoError(t, err)
assert.Empty(t, items, "PAT scoped to a server it cannot reach must see no groups, not all of them")
}
// batchDeleteServerGroup must refuse to delete a group whose members are not
// entirely covered by the PAT whitelist; otherwise an admin's limited PAT can
// drop groups that touch servers outside its scope.
func TestBatchDeleteServerGroupRejectsPATOutsideWhitelist(t *testing.T) {
setupServerGroupVisibilityFixture(t)
require.NoError(t, singleton.DB.Create(&model.ServerGroupServer{
Common: model.Common{UserID: 1}, ServerGroupId: 10, ServerId: 2,
}).Error)
tok := &model.APIToken{ID: 79, UserID: 1}
tok.SetServerIDs([]uint64{1})
c := newServerGroupCtxWithPAT(&model.User{
Common: model.Common{ID: 1}, Role: model.RoleAdmin,
}, tok)
body, _ := json.Marshal([]uint64{10})
c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/batch-delete/server-group", bytes.NewReader(body))
c.Request.Header.Set("Content-Type", "application/json")
_, err := batchDeleteServerGroup(c)
require.Error(t, err, "PAT scoped to {1} must not delete group 10 which still contains server 2")
var remaining int64
require.NoError(t, singleton.DB.Model(&model.ServerGroup{}).Where("id = ?", 10).Count(&remaining).Error)
assert.Equal(t, int64(1), remaining, "group 10 must remain after refused PAT delete")
}
+43 -4
View File
@@ -2,7 +2,6 @@ package controller
import (
"fmt"
"maps"
"slices"
"strconv"
"strings"
@@ -56,7 +55,18 @@ func serviceResponseCacheKey(c *gin.Context) string {
if !ok || user == nil {
return "list-service::guest"
}
return fmt.Sprintf("list-service::%t::%d", user.Role.IsAdmin(), user.ID)
base := fmt.Sprintf("list-service::%t::%d", user.Role.IsAdmin(), user.ID)
tok := APITokenFromContext(c)
if tok == nil {
return base + "::jwt"
}
ids := tok.ServerIDs()
slices.Sort(ids)
parts := make([]string, 0, len(ids))
for _, id := range ids {
parts = append(parts, strconv.FormatUint(id, 10))
}
return fmt.Sprintf("%s::pat:%d::servers:%s", base, tok.ID, strings.Join(parts, ","))
}
func filterCycleTransferStatsForViewer(c *gin.Context, stats map[uint64]model.CycleTransferStats) map[uint64]model.CycleTransferStats {
@@ -459,6 +469,10 @@ func createService(c *gin.Context) (uint64, error) {
return 0, err
}
if !isValidServiceCover(mf.Cover) {
return 0, singleton.Localizer.ErrorT("permission denied")
}
uid := getUid(c)
var m model.Service
@@ -518,6 +532,11 @@ func updateService(c *gin.Context) (any, error) {
if err := c.ShouldBindJSON(&mf); err != nil {
return nil, err
}
if !isValidServiceCover(mf.Cover) {
return nil, singleton.Localizer.ErrorT("permission denied")
}
var m model.Service
if err := singleton.DB.First(&m, id).Error; err != nil {
return nil, singleton.Localizer.ErrorT("service id %d does not exist", id)
@@ -581,6 +600,19 @@ func batchDeleteService(c *gin.Context) (any, error) {
return nil, singleton.Localizer.ErrorT("permission denied")
}
// 与 batchDeleteCron 对称:DispatchTask 没有 PAT 上下文,这里是阻止
// 受限 PAT 通过删除 ServiceCoverAll + 不充分 SkipServers 间接影响
// 白名单外 owner servers 探测状态的唯一同步入口。
for _, id := range ids {
existing, ok := singleton.ServiceSentinelShared.Get(id)
if !ok || existing == nil {
continue
}
if err := enforcePATServiceDispatchScope(c, existing); err != nil {
return nil, err
}
}
err := singleton.DB.Transaction(func(tx *gorm.DB) error {
return tx.Unscoped().Delete(&model.Service{}, "id in (?)", ids).Error
})
@@ -593,8 +625,12 @@ func batchDeleteService(c *gin.Context) (any, error) {
}
func validateServers(c *gin.Context, ss *model.Service) error {
if !singleton.ServerShared.CheckPermission(c, maps.Keys(ss.SkipServers)) {
return singleton.Localizer.ErrorT("permission denied")
if err := checkServiceSkipServerPermission(c, ss.Cover, ss.SkipServers, ss.GetUserID()); err != nil {
return err
}
if err := rejectImplicitServiceCoverForLimitedPAT(c, ss.Cover, ss.SkipServers, ss.GetUserID()); err != nil {
return err
}
if !singleton.CronShared.CheckPermission(c, slices.Values(ss.FailTriggerTasks)) {
@@ -603,6 +639,9 @@ func validateServers(c *gin.Context, ss *model.Service) error {
if !singleton.CronShared.CheckPermission(c, slices.Values(ss.RecoverTriggerTasks)) {
return singleton.Localizer.ErrorT("permission denied")
}
if err := enforcePATTriggerTaskScope(c, ss.FailTriggerTasks, ss.RecoverTriggerTasks); err != nil {
return err
}
if err := assertOwnsNotificationGroup(c, ss.NotificationGroupID); err != nil {
return err
@@ -0,0 +1,58 @@
package controller
import (
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/nezhahq/nezha/model"
)
func newCacheKeyCtx(t *testing.T, user *model.User, tok *model.APIToken) *gin.Context {
t.Helper()
gin.SetMode(gin.TestMode)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest("GET", "/api/v1/service", nil)
if user != nil {
c.Set(model.CtxKeyAuthorizedUser, user)
}
if tok != nil {
c.Set(model.CtxKeyAPIToken, tok)
c.Set(apiTokenCtxKey, tok)
}
return c
}
func TestServiceResponseCacheKey_DistinguishesPATsWithDifferentServerWhitelist(t *testing.T) {
user := &model.User{Common: model.Common{ID: 100}, Role: model.RoleMember}
tokA := &model.APIToken{ID: 1, UserID: 100}
tokA.SetServerIDs([]uint64{7})
tokB := &model.APIToken{ID: 2, UserID: 100}
tokB.SetServerIDs([]uint64{8})
keyA := serviceResponseCacheKey(newCacheKeyCtx(t, user, tokA))
keyB := serviceResponseCacheKey(newCacheKeyCtx(t, user, tokB))
if keyA == keyB {
t.Fatalf("singleflight key must differ across PATs with disjoint server_ids; got %q for both",
keyA)
}
}
func TestServiceResponseCacheKey_DistinguishesPATFromJWT(t *testing.T) {
user := &model.User{Common: model.Common{ID: 100}, Role: model.RoleMember}
tok := &model.APIToken{ID: 1, UserID: 100}
tok.SetServerIDs([]uint64{7})
keyPAT := serviceResponseCacheKey(newCacheKeyCtx(t, user, tok))
keyJWT := serviceResponseCacheKey(newCacheKeyCtx(t, user, nil))
if keyPAT == keyJWT {
t.Fatalf("PAT-shaped key must not collide with the JWT-shaped key; got %q for both", keyPAT)
}
}
@@ -0,0 +1,185 @@
package controller
// 回归 service monitor 运行时入口 (batchDeleteService) 上的 PAT
// cover-fanout 收口。与 cron_dispatch_pat_test.go 对称,钉死写侧
// rejectImplicitServiceCoverForLimitedPAT 与运行时
// enforcePATServiceDispatchScope 共用同一裁决路径。
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 setupServiceDispatchPATFixture(t *testing.T) {
t.Helper()
originalDB := singleton.DB
originalCache := singleton.Cache
originalLoc := singleton.Loc
originalLocalizer := singleton.Localizer
originalServer := singleton.ServerShared
originalUserInfo := singleton.UserInfoMap
originalSentinel := singleton.ServiceSentinelShared
originalCron := singleton.CronShared
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.Service{}, &model.Server{}, &model.User{}, &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)
// ServiceSentinel 在构造时会调 CronShared.AddFunc 注册每日/每周维护任务,
// 必须先于 NewServiceSentinel 装配。
singleton.CronShared = singleton.NewCronClass()
sentinel, err := singleton.NewServiceSentinel(make(chan *model.Service, 4))
require.NoError(t, err)
singleton.ServiceSentinelShared = sentinel
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.UserLock.Lock()
singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}}
singleton.UserLock.Unlock()
t.Cleanup(func() {
sentinel.Close()
singleton.ServiceSentinelShared = originalSentinel
singleton.CronShared = originalCron
singleton.DB = originalDB
singleton.Cache = originalCache
singleton.Loc = originalLoc
singleton.Localizer = originalLocalizer
singleton.ServerShared = originalServer
singleton.UserLock.Lock()
singleton.UserInfoMap = originalUserInfo
singleton.UserLock.Unlock()
})
}
func insertServiceForDispatchTest(t *testing.T, cover uint8, skip map[uint64]bool) uint64 {
t.Helper()
svc := &model.Service{
Common: model.Common{UserID: 100},
Name: "dispatch-svc-fixture",
Type: model.TaskTypeTCPPing,
Target: "example.invalid:80",
Duration: 30,
Cover: cover,
SkipServers: skip,
}
require.NoError(t, singleton.DB.Create(svc).Error)
require.NoError(t, singleton.ServiceSentinelShared.Update(svc))
singleton.ServiceSentinelShared.UpdateServiceList()
return svc.ID
}
func newServiceDispatchRouter(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/batch-delete/service", commonHandler(batchDeleteService))
return r
}
func TestBatchDeleteService_RejectsCoverAllWithInsufficientSkipForLimitedPAT(t *testing.T) {
setupServiceDispatchPATFixture(t)
svcID := insertServiceForDispatchTest(t, model.ServiceCoverAll, map[uint64]bool{1: true})
tok := &model.APIToken{ID: 31, UserID: 100}
tok.SetServerIDs([]uint64{1})
r := newServiceDispatchRouter(t, tok)
body, _ := json.Marshal([]uint64{svcID})
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/v1/batch-delete/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 batch-delete a ServiceCoverAll monitor whose SkipServers only marks whitelisted servers; DispatchTask still probes server 2")
assert.Contains(t, errMsg, "permission denied")
var rows []model.Service
require.NoError(t, singleton.DB.Find(&rows).Error)
assert.Len(t, rows, 1, "service row must still exist when the delete call is rejected")
}
func TestBatchDeleteService_AllowsCoverAllWhenSkipCoversNonWhitelisted(t *testing.T) {
setupServiceDispatchPATFixture(t)
svcID := insertServiceForDispatchTest(t, model.ServiceCoverAll, map[uint64]bool{2: true})
tok := &model.APIToken{ID: 32, UserID: 100}
tok.SetServerIDs([]uint64{1})
r := newServiceDispatchRouter(t, tok)
body, _ := json.Marshal([]uint64{svcID})
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/v1/batch-delete/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,
"SkipServers covering every non-whitelisted owner server must allow batch-delete: error=%s", errMsg)
var rows []model.Service
require.NoError(t, singleton.DB.Find(&rows).Error)
assert.Empty(t, rows, "service row must be deleted when the call succeeds")
}
func TestBatchDeleteService_AllowsCoverIgnoreAllInsideWhitelist(t *testing.T) {
setupServiceDispatchPATFixture(t)
svcID := insertServiceForDispatchTest(t, model.ServiceCoverIgnoreAll, map[uint64]bool{1: true})
tok := &model.APIToken{ID: 33, UserID: 100}
tok.SetServerIDs([]uint64{1})
r := newServiceDispatchRouter(t, tok)
body, _ := json.Marshal([]uint64{svcID})
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/api/v1/batch-delete/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,
"ServiceCoverIgnoreAll allow-list inside PAT whitelist must allow batch-delete: error=%s", errMsg)
var rows []model.Service
require.NoError(t, singleton.DB.Find(&rows).Error)
assert.Empty(t, rows)
}
@@ -0,0 +1,66 @@
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 newServiceListPATRouter(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/service/list", listHandler(listService))
return r
}
// GET /api/v1/service/list must hide ServiceCoverAll rows whose SkipServers
// deny-set does not cover every owner server outside the PAT whitelist.
// DispatchTask would still probe those servers, so leaking the row to the
// list view (and exposing target/credentials/triggers) is a real PAT scope
// escape.
func TestListService_HidesCoverAllWithInsufficientSkipForLimitedPAT(t *testing.T) {
setupServiceDispatchPATFixture(t)
insufficient := insertServiceForDispatchTest(t, model.ServiceCoverAll, map[uint64]bool{1: true})
sufficient := insertServiceForDispatchTest(t, model.ServiceCoverAll, map[uint64]bool{2: true})
tok := &model.APIToken{ID: 34, UserID: 100}
tok.SetServerIDs([]uint64{1})
r := newServiceListPATRouter(t, tok)
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/v1/service/list", nil)
r.ServeHTTP(w, req)
var resp struct {
Success bool `json:"success"`
Error string `json:"error"`
Data []*model.Service `json:"data"`
}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
require.True(t, resp.Success, resp.Error)
seen := map[uint64]bool{}
for _, s := range resp.Data {
seen[s.ID] = true
}
assert.False(t, seen[insufficient],
"PAT [1] must NOT see a ServiceCoverAll whose SkipServers does not cover owner server 2 (rows=%+v)", resp.Data)
assert.True(t, seen[sufficient],
"PAT [1] must still see a ServiceCoverAll whose SkipServers already covers every non-whitelisted owner server")
}
@@ -0,0 +1,66 @@
package controller
import (
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/nezhahq/nezha/model"
"github.com/nezhahq/nezha/pkg/i18n"
"github.com/nezhahq/nezha/service/singleton"
)
func ensureLocalizerForServiceTest(t *testing.T) {
t.Helper()
if singleton.Localizer == nil {
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
}
}
// M13 regression: checkServiceSkipServerPermission must treat SkipServers
// as a typed map[uint64]bool where only `true` entries actually skip
// at runtime (DispatchTask only consults true keys). Entries with value
// false carry no dispatch meaning, so requiring HasPermission on them
// rejects perfectly legitimate updates from members whose PAT does not
// own the no-op `{2: false}` server.
func TestCheckServiceSkipServerPermission_IgnoresFalseEntries(t *testing.T) {
ensureLocalizerForServiceTest(t)
saved := singleton.ServerShared
t.Cleanup(func() { singleton.ServerShared = saved })
sc := singleton.NewEmptyServerClassForTest()
sc.InsertForTest(&model.Server{Common: model.Common{ID: 1, UserID: 100}})
sc.InsertForTest(&model.Server{Common: model.Common{ID: 2, UserID: 999}})
singleton.ServerShared = sc
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 100}, Role: model.RoleMember})
skip := map[uint64]bool{
1: true, // member owns it — legal allow-list entry
2: false, // no-op entry; member doesn't own server 2 but it's not actually skipped
}
if err := checkServiceSkipServerPermission(ctx, model.ServiceCoverIgnoreAll, skip, 100); err != nil {
t.Fatalf("`{2: false}` must NOT trigger permission denied — it has no runtime dispatch effect, got %v", err)
}
}
func TestCheckServiceSkipServerPermission_RejectsForeignTrueEntries(t *testing.T) {
ensureLocalizerForServiceTest(t)
saved := singleton.ServerShared
t.Cleanup(func() { singleton.ServerShared = saved })
sc := singleton.NewEmptyServerClassForTest()
sc.InsertForTest(&model.Server{Common: model.Common{ID: 1, UserID: 100}})
sc.InsertForTest(&model.Server{Common: model.Common{ID: 2, UserID: 999}})
singleton.ServerShared = sc
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 100}, Role: model.RoleMember})
skip := map[uint64]bool{
2: true, // member doesn't own server 2 — true entry IS the allow-list, must reject
}
if err := checkServiceSkipServerPermission(ctx, model.ServiceCoverIgnoreAll, skip, 100); err == nil {
t.Fatal("true entry pointing at foreign-owned server must still be rejected — pre-existing safety invariant")
}
}
@@ -46,3 +46,40 @@ func TestUserCanViewServiceHiddenServiceAllowsAdmin(t *testing.T) {
admin := &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin}
assert.True(t, userCanViewService(newServiceVisibilityCtx(admin), hidden), "admin must be able to see any hidden service")
}
// 钉死 admin 自己签发的 server_ids 受限 PAT 不能借助 admin 身份在
// service 可见性入口绕过白名单:与 userCanViewServer 的 PAT-first 收口
// 保持对称,避免 hidden service 通过 admin 早返回泄漏给受限 PAT。
func TestUserCanViewServiceLimitedPATShouldDenyAdminWhenOutsideWhitelist(t *testing.T) {
hidden := &model.Service{
Common: model.Common{ID: 1, UserID: 100},
Cover: model.ServiceCoverIgnoreAll,
SkipServers: map[uint64]bool{2: true},
}
admin := &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin}
tok := &model.APIToken{ID: 7, UserID: 1}
tok.SetServerIDs([]uint64{1})
ctx := newServiceVisibilityCtx(admin)
ctx.Set(model.CtxKeyAPIToken, tok)
assert.False(t, userCanViewService(ctx, hidden),
"admin caller using a server_ids=[1] PAT must NOT see a CoverIgnoreAll service whose only target is the non-whitelisted server 2")
}
func TestUserCanViewServiceLimitedPATAllowsAdminInsideWhitelist(t *testing.T) {
visible := &model.Service{
Common: model.Common{ID: 2, UserID: 100},
Cover: model.ServiceCoverIgnoreAll,
SkipServers: map[uint64]bool{1: true},
}
admin := &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin}
tok := &model.APIToken{ID: 7, UserID: 1}
tok.SetServerIDs([]uint64{1})
ctx := newServiceVisibilityCtx(admin)
ctx.Set(model.CtxKeyAPIToken, tok)
assert.True(t, userCanViewService(ctx, visible),
"admin caller using a server_ids=[1] PAT must still see a CoverIgnoreAll service bound to whitelisted server 1")
}
+49 -1
View File
@@ -2,11 +2,13 @@ package controller
import (
"errors"
"log"
"strings"
"github.com/gin-gonic/gin"
"github.com/nezhahq/nezha/model"
"github.com/nezhahq/nezha/service/rpc"
"github.com/nezhahq/nezha/service/singleton"
)
@@ -107,8 +109,15 @@ func updateConfig(c *gin.Context) (any, error) {
singleton.Conf.AgentRealIPHeader = sf.AgentRealIPHeader
singleton.Conf.AgentTLS = sf.AgentTLS
singleton.Conf.UserTemplate = sf.UserTemplate
mcpWasEnabled := singleton.Conf.MCPEnabled()
mcpNext := resolveSettingEnableMCP(sf.EnableMCP, mcpWasEnabled)
if err := singleton.Conf.Save(); err != nil {
if err := applyEnableMCPTransition(
mcpWasEnabled, mcpNext,
singleton.Conf.SetMCPEnabled,
singleton.Conf.Save,
fireMCPKillSwitch,
); err != nil {
return nil, newGormError("%v", err)
}
@@ -116,6 +125,45 @@ func updateConfig(c *gin.Context) (any, error) {
return nil, nil
}
// applyEnableMCPTransition commits the new EnableMCP value and persists it,
// guaranteeing the in-memory flag and the kill-switch cleanup stay consistent
// with what actually reached durable storage:
// - setVal(next) is applied so Save serialises the new value.
// - If save fails, the flag is rolled back to prev and no cleanup runs, so a
// failed disable cannot leave the dashboard half-disabled (new requests
// rejected while in-flight RPC/streams/URLs are never revoked).
// - cleanup runs only on a persisted enabled->disabled transition.
func applyEnableMCPTransition(prev, next bool, setVal func(bool), save func() error, cleanup func()) error {
setVal(next)
if err := save(); err != nil {
setVal(prev)
return err
}
if prev && !next {
cleanup()
}
return nil
}
func fireMCPKillSwitch() {
purgedURLs := PurgeTransferEntries()
revokedStreams := rpc.NezhaHandlerSingleton.RevokeStreamsForPurpose(rpc.PurposeMCPTransfer)
cancelledRPC := rpc.CancelAllMCPInflight()
log.Printf("NEZHA>> MCP kill switch fired: purged=%d urls, revoked=%d streams, cancelled=%d rpc",
purgedURLs, revokedStreams, cancelledRPC)
}
// resolveSettingEnableMCP picks the effective EnableMCP value for the
// update. A nil form pointer means "field absent" so we MUST preserve
// the current config to avoid accidentally tripping the kill switch on
// partial PATCH calls that omit enable_mcp.
func resolveSettingEnableMCP(formValue *bool, current bool) bool {
if formValue == nil {
return current
}
return *formValue
}
// Perform maintenance
// @Summary Perform maintenance
// @Security BearerAuth
@@ -0,0 +1,73 @@
package controller
import (
"errors"
"testing"
)
// Review issue #3: when persisting the new EnableMCP value fails, updateConfig
// must NOT leave the dashboard in a half-disabled state where in-memory
// EnableMCP=false (new requests rejected) but the kill-switch cleanup
// (PurgeTransferEntries / RevokeStreamsForPurpose / CancelAllMCPInflight)
// never ran. applyEnableMCPTransition owns that invariant: on save failure it
// rolls the in-memory flag back to its previous value and runs no cleanup.
func TestApplyEnableMCPTransition_SaveFailureRollsBackAndSkipsCleanup(t *testing.T) {
current := true
cleanupRan := false
setVal := func(v bool) { current = v }
saveErr := errors.New("disk full")
save := func() error { return saveErr }
cleanup := func() { cleanupRan = true }
err := applyEnableMCPTransition(true /*prev*/, false /*next*/, setVal, save, cleanup)
if !errors.Is(err, saveErr) {
t.Fatalf("expected the save error to propagate, got %v", err)
}
if current != true {
t.Fatalf("in-memory EnableMCP must roll back to its previous value on save failure; got %v", current)
}
if cleanupRan {
t.Fatal("kill-switch cleanup must NOT run when the new value was never persisted")
}
}
func TestApplyEnableMCPTransition_DisableSuccessRunsCleanup(t *testing.T) {
current := true
cleanupRan := false
setVal := func(v bool) { current = v }
save := func() error { return nil }
cleanup := func() { cleanupRan = true }
if err := applyEnableMCPTransition(true, false, setVal, save, cleanup); err != nil {
t.Fatalf("unexpected error: %v", err)
}
if current != false {
t.Fatalf("EnableMCP must be committed to false after a successful save; got %v", current)
}
if !cleanupRan {
t.Fatal("kill-switch cleanup must run when MCP transitions enabled->disabled and the save succeeds")
}
}
func TestApplyEnableMCPTransition_EnableSuccessSkipsCleanup(t *testing.T) {
current := false
cleanupRan := false
setVal := func(v bool) { current = v }
save := func() error { return nil }
cleanup := func() { cleanupRan = true }
if err := applyEnableMCPTransition(false, true, setVal, save, cleanup); err != nil {
t.Fatalf("unexpected error: %v", err)
}
if current != true {
t.Fatalf("EnableMCP must be committed to true; got %v", current)
}
if cleanupRan {
t.Fatal("cleanup must only run on the enabled->disabled transition, not when enabling")
}
}
@@ -0,0 +1,84 @@
package controller
import (
"bytes"
"encoding/json"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/nezhahq/nezha/model"
)
// M10 regression: a PATCH /setting payload that omits "enable_mcp" must
// preserve the current value. With EnableMCP as a plain bool + omitempty,
// any partial update silently set EnableMCP=false and tripped the MCP
// kill switch (PurgeTransferEntries + RevokeStreamsForPurpose +
// CancelAllMCPInflight). Switching to *bool makes "field absent" a real
// signal at decode time.
func TestSettingForm_OmittedEnableMCPLeavesConfigUnchanged(t *testing.T) {
body := []byte(`{"site_name":"X"}`)
var sf model.SettingForm
if err := json.NewDecoder(bytes.NewReader(body)).Decode(&sf); err != nil {
t.Fatal(err)
}
if sf.EnableMCP != nil {
t.Fatalf("EnableMCP must be nil when JSON omits the key, got %v", sf.EnableMCP)
}
}
func TestSettingForm_ExplicitEnableMCPTrueDecodes(t *testing.T) {
body := []byte(`{"enable_mcp":true}`)
var sf model.SettingForm
if err := json.NewDecoder(bytes.NewReader(body)).Decode(&sf); err != nil {
t.Fatal(err)
}
if sf.EnableMCP == nil || !*sf.EnableMCP {
t.Fatalf("EnableMCP must be *true, got %v", sf.EnableMCP)
}
}
func TestSettingForm_ExplicitEnableMCPFalseDecodes(t *testing.T) {
body := []byte(`{"enable_mcp":false}`)
var sf model.SettingForm
if err := json.NewDecoder(bytes.NewReader(body)).Decode(&sf); err != nil {
t.Fatal(err)
}
if sf.EnableMCP == nil || *sf.EnableMCP {
t.Fatalf("EnableMCP must be *false, got %v", sf.EnableMCP)
}
}
// updateMCPEnableFromForm is the resolver helper: nil = keep current,
// non-nil = use the explicit value. Kept as a small pure function so the
// kill-switch wiring stays trivial to audit.
func TestUpdateMCPEnableFromForm_NilKeepsCurrent(t *testing.T) {
_, w := newRecorderCtxForMCPSettingTest(t)
prev := true
got := resolveSettingEnableMCP(nil, prev)
if got != prev {
t.Fatalf("nil form value must keep current=%v, got %v", prev, got)
}
if w.Code != 200 {
t.Fatal("resolver must not write to response")
}
}
func TestUpdateMCPEnableFromForm_NonNilOverrides(t *testing.T) {
f := false
if got := resolveSettingEnableMCP(&f, true); got != false {
t.Fatal("explicit *false must override current=true")
}
tr := true
if got := resolveSettingEnableMCP(&tr, false); got != true {
t.Fatal("explicit *true must override current=false")
}
}
func newRecorderCtxForMCPSettingTest(t *testing.T) (*gin.Context, *httptest.ResponseRecorder) {
t.Helper()
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
return c, w
}
@@ -0,0 +1,90 @@
package controller
import (
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
"github.com/nezhahq/nezha/model"
"github.com/nezhahq/nezha/service/rpc"
)
func ensureNezhaSingleton(t *testing.T) {
t.Helper()
if rpc.NezhaHandlerSingleton == nil {
rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler()
}
}
// H2 regression: terminal/FM stream attachment must respect the caller PAT's
// server_ids whitelist. The existing IsStreamAuthorizedForUser only gates on
// creator-id / admin role, so an admin's server-limited PAT could attach to
// a stream targeting any server simply by knowing the streamId.
func TestStreamAttachAllowedForRequest_DeniesPATOutsideWhitelist(t *testing.T) {
ensureNezhaSingleton(t)
streamId := "stream-h2-deny"
rpc.NezhaHandlerSingleton.CreateStream(streamId, 1, 99)
t.Cleanup(func() { _ = rpc.NezhaHandlerSingleton.CloseStream(streamId) })
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin})
ctx.Set(model.CtxKeyAPIToken, &model.APIToken{ServersCSV: "1"}) // does NOT include 99
if streamAttachAllowedForRequest(ctx, streamId) {
t.Fatal("admin PAT scoped to [1] must NOT attach to a stream targeting server 99")
}
}
func TestStreamAttachAllowedForRequest_AllowsPATInsideWhitelist(t *testing.T) {
ensureNezhaSingleton(t)
streamId := "stream-h2-allow"
rpc.NezhaHandlerSingleton.CreateStream(streamId, 1, 5)
t.Cleanup(func() { _ = rpc.NezhaHandlerSingleton.CloseStream(streamId) })
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin})
ctx.Set(model.CtxKeyAPIToken, &model.APIToken{ServersCSV: "5"})
if !streamAttachAllowedForRequest(ctx, streamId) {
t.Fatal("PAT scoped to [5] must attach to a stream targeting server 5")
}
}
func TestStreamAttachAllowedForRequest_JWTAdminUnchanged(t *testing.T) {
ensureNezhaSingleton(t)
streamId := "stream-h2-jwt"
rpc.NezhaHandlerSingleton.CreateStream(streamId, 1, 99)
t.Cleanup(func() { _ = rpc.NezhaHandlerSingleton.CloseStream(streamId) })
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin})
if !streamAttachAllowedForRequest(ctx, streamId) {
t.Fatal("JWT admin (no PAT) must continue to attach via the existing admin branch")
}
}
func TestStreamAttachAllowedForRequest_DeniesNonCreatorMember(t *testing.T) {
ensureNezhaSingleton(t)
streamId := "stream-h2-foreign"
rpc.NezhaHandlerSingleton.CreateStream(streamId, 1, 5)
t.Cleanup(func() { _ = rpc.NezhaHandlerSingleton.CloseStream(streamId) })
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 2}, Role: model.RoleMember})
if streamAttachAllowedForRequest(ctx, streamId) {
t.Fatal("non-creator non-admin member must remain denied (pre-existing GHSA gate)")
}
}
func TestStreamAttachAllowedForRequest_UnknownStreamRejected(t *testing.T) {
ensureNezhaSingleton(t)
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin})
if streamAttachAllowedForRequest(ctx, "does-not-exist") {
t.Fatal("unknown streamId must remain rejected")
}
}
@@ -0,0 +1,315 @@
package controller
import (
"bytes"
"encoding/json"
"net/http/httptest"
"testing"
"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"
)
// 通用租户隔离测试夹具:在 in-memory DB 上挂载所需 model 并塞两个用户,
// 用户 10member)和用户 999foreign owner)。
//
// 每个测试在两条路径上验证 member 不能跨租户:
// - create 时即使请求体里包含 user_id 字段也不会越权
// - update / delete 时不会改写或读取到 foreign owner 的资源
func setupTenancyTest(t *testing.T) func() {
t.Helper()
originalDB := singleton.DB
originalLocalizer := singleton.Localizer
originalServer := singleton.ServerShared
if singleton.Localizer == nil {
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
}
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(
&model.User{},
&model.Cron{},
&model.DDNSProfile{},
&model.Notification{},
&model.AlertRule{},
&model.NotificationGroup{},
))
originalDDNS := singleton.DDNSShared
originalNotif := singleton.NotificationShared
singleton.DB = db
singleton.ServerShared = singleton.NewEmptyServerClassForTest()
singleton.DDNSShared = singleton.NewEmptyDDNSClassForTest()
singleton.NotificationShared = singleton.NewEmptyNotificationClassForTest()
return func() {
singleton.DB = originalDB
singleton.Localizer = originalLocalizer
singleton.ServerShared = originalServer
singleton.DDNSShared = originalDDNS
singleton.NotificationShared = originalNotif
}
}
func ctxAs(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 ctxAsMemberWithBody(uid uint64, body any) *gin.Context {
c := ctxAs(uid, model.RoleMember)
b, _ := json.Marshal(body)
c.Request = httptest.NewRequest("POST", "/", bytes.NewReader(b))
c.Request.Header.Set("Content-Type", "application/json")
return c
}
// 设计说明:create 路径的"手工 user_id 注入"防护通过两点联合保证:
// 1. CronForm/DDNSForm/NotificationForm 等 form struct 不嵌入 Common
// 绑定时不会 unmarshal "user_id" 字段
// 2. handler 第一行 `xxx.UserID = getUid(c)` 显式覆盖
// 因为 create 路径还会依赖 ServerShared / Localizer 等外部 singleton
// 在单元测试中难以无副作用地完整运行;改用代码静态约束:在 form_no_userid_test.go
// 里用 reflect 验证所有 *Form 结构无 UserID 字段(next step)。
// 这里只测真正的所有权防线:update / delete。
// ---------- Cron ----------
func TestTenancy_UpdateCron_ForeignOwnerRejected(t *testing.T) {
defer setupTenancyTest(t)()
foreign := model.Cron{
Common: model.Common{UserID: 999},
Name: "foreign-cron",
TaskType: model.CronTypeCronTask,
Scheduler: "@every 5m",
Command: "echo",
}
require.NoError(t, singleton.DB.Create(&foreign).Error)
c := ctxAsMemberWithBody(10, map[string]any{
"name": "hijacked",
"task_type": model.CronTypeCronTask,
"scheduler": "@every 1m",
"command": "echo pwned",
"servers": []uint64{},
"cover": model.CronCoverAll,
})
c.Params = gin.Params{{Key: "id", Value: itoa(foreign.ID)}}
_, err := updateCron(c)
require.Error(t, err, "member 10 must not be able to update foreign-owned cron")
var after model.Cron
require.NoError(t, singleton.DB.First(&after, foreign.ID).Error)
require.Equal(t, "foreign-cron", after.Name, "foreign cron must not be modified")
require.Equal(t, uint64(999), after.UserID, "ownership must remain")
}
// ---------- DDNS ----------
func TestTenancy_CreateDDNS_InjectedUserIDIgnored(t *testing.T) {
defer setupTenancyTest(t)()
body := map[string]any{
"name": "evil-ddns",
"provider": "webhook",
"access_id": "x",
"access_secret": "y",
"webhook_url": "http://127.0.0.1/",
"webhook_method": "GET",
"webhook_request_type": "json",
"webhook_request_body": "",
"webhook_headers": "",
"user_id": 999, // attacker
}
c := ctxAsMemberWithBody(10, body)
_, err := createDDNS(c)
if err == nil {
var stored model.DDNSProfile
require.NoError(t, singleton.DB.First(&stored, "name = ?", "evil-ddns").Error)
require.Equal(t, uint64(10), stored.UserID,
"createDDNS must overwrite UserID with caller")
}
}
func TestTenancy_UpdateDDNS_ForeignOwnerRejected(t *testing.T) {
defer setupTenancyTest(t)()
foreign := model.DDNSProfile{
Common: model.Common{UserID: 999},
Name: "foreign-ddns",
Provider: "webhook",
AccessID: "x",
AccessSecret: "y",
WebhookURL: "http://127.0.0.1/",
WebhookMethod: 1,
WebhookRequestType: 1,
}
require.NoError(t, singleton.DB.Create(&foreign).Error)
c := ctxAsMemberWithBody(10, map[string]any{
"name": "hijacked",
"provider": "webhook",
"access_id": "x",
"access_secret": "y",
"webhook_url": "http://attacker/",
"webhook_method": "GET",
"webhook_request_type": "json",
})
c.Params = gin.Params{{Key: "id", Value: itoa(foreign.ID)}}
_, err := updateDDNS(c)
require.Error(t, err, "member must not be able to update foreign-owned DDNS")
var after model.DDNSProfile
require.NoError(t, singleton.DB.First(&after, foreign.ID).Error)
require.Equal(t, "foreign-ddns", after.Name, "foreign DDNS must not be modified")
require.Equal(t, "http://127.0.0.1/", after.WebhookURL, "webhook URL must not be hijacked")
}
func TestTenancy_DeleteDDNS_ForeignOwnerRejected(t *testing.T) {
defer setupTenancyTest(t)()
foreign := model.DDNSProfile{
Common: model.Common{UserID: 999},
Name: "foreign-ddns-del",
Provider: "webhook",
WebhookURL: "http://127.0.0.1/",
WebhookMethod: 1,
WebhookRequestType: 1,
}
require.NoError(t, singleton.DB.Create(&foreign).Error)
singleton.DDNSShared.InsertForTest(&foreign)
c := ctxAsMemberWithBody(10, []uint64{foreign.ID})
_, err := batchDeleteDDNS(c)
require.Error(t, err, "member must not be able to batch-delete foreign DDNS")
var after model.DDNSProfile
require.NoErrorf(t, singleton.DB.First(&after, foreign.ID).Error,
"foreign DDNS must still exist after member's failed batch-delete (handler err=%v)", err)
}
// ---------- Notification ----------
func TestTenancy_UpdateNotification_ForeignOwnerRejected(t *testing.T) {
defer setupTenancyTest(t)()
foreign := model.Notification{
Common: model.Common{UserID: 999},
Name: "foreign-notify",
URL: "http://127.0.0.1/",
RequestMethod: 1,
RequestType: 1,
}
require.NoError(t, singleton.DB.Create(&foreign).Error)
c := ctxAsMemberWithBody(10, map[string]any{
"name": "hijacked",
"url": "http://attacker/",
"request_method": 1,
"request_type": 1,
})
c.Params = gin.Params{{Key: "id", Value: itoa(foreign.ID)}}
_, err := updateNotification(c)
require.Error(t, err)
var after model.Notification
require.NoError(t, singleton.DB.First(&after, foreign.ID).Error)
require.Equal(t, "http://127.0.0.1/", after.URL)
}
// ---------- NotificationGroup ----------
func TestTenancy_UpdateNotificationGroup_ForeignOwnerRejected(t *testing.T) {
defer setupTenancyTest(t)()
foreign := model.NotificationGroup{
Common: model.Common{UserID: 999},
Name: "foreign-ng",
}
require.NoError(t, singleton.DB.Create(&foreign).Error)
c := ctxAsMemberWithBody(10, map[string]any{
"name": "hijacked",
"notifications": []uint64{},
})
c.Params = gin.Params{{Key: "id", Value: itoa(foreign.ID)}}
_, err := updateNotificationGroup(c)
require.Error(t, err)
var after model.NotificationGroup
require.NoError(t, singleton.DB.First(&after, foreign.ID).Error)
require.Equal(t, "foreign-ng", after.Name)
}
// ---------- AlertRule ----------
func TestTenancy_UpdateAlertRule_ForeignOwnerRejected(t *testing.T) {
defer setupTenancyTest(t)()
foreign := model.AlertRule{
Common: model.Common{UserID: 999},
Name: "foreign-rule",
}
require.NoError(t, singleton.DB.Create(&foreign).Error)
c := ctxAsMemberWithBody(10, map[string]any{
"name": "hijacked",
})
c.Params = gin.Params{{Key: "id", Value: itoa(foreign.ID)}}
_, err := updateAlertRule(c)
require.Error(t, err)
var after model.AlertRule
require.NoError(t, singleton.DB.First(&after, foreign.ID).Error)
require.Equal(t, "foreign-rule", after.Name)
}
func TestTenancy_BatchDeleteAlertRule_ForeignOwnerSilentlySkipped(t *testing.T) {
defer setupTenancyTest(t)()
foreign := model.AlertRule{Common: model.Common{UserID: 999}, Name: "foreign-rule"}
require.NoError(t, singleton.DB.Create(&foreign).Error)
c := ctxAsMemberWithBody(10, []uint64{foreign.ID})
_, err := batchDeleteAlertRule(c)
_ = err
var after model.AlertRule
require.NoError(t, singleton.DB.First(&after, foreign.ID).Error,
"member's batch-delete must not be able to remove foreign alert rule")
require.Equal(t, uint64(999), after.UserID)
}
// Cron batch-delete 的所有权保护与 updateCron 共用 cr.HasPermission 检查路径
// cron.go:127 vs cron.go:207),updateCron 用例已经覆盖该路径。这里不复测
// 是因为 batchDeleteCron 调 CronShared.CheckPermission,需要完整 CronShared
// 在内存中注册,会让单测 fixture 显著膨胀,性价比低。
// ---------- Notification batch-delete ----------
func TestTenancy_BatchDeleteNotification_ForeignOwnerSilentlySkipped(t *testing.T) {
defer setupTenancyTest(t)()
foreign := model.Notification{
Common: model.Common{UserID: 999},
Name: "foreign-notify-del",
URL: "http://127.0.0.1/",
}
require.NoError(t, singleton.DB.Create(&foreign).Error)
singleton.NotificationShared.InsertForTest(&foreign)
c := ctxAsMemberWithBody(10, []uint64{foreign.ID})
_, _ = batchDeleteNotification(c)
var after model.Notification
require.NoError(t, singleton.DB.First(&after, foreign.ID).Error,
"member must not be able to batch-delete foreign notification")
}
+6 -4
View File
@@ -34,8 +34,7 @@ func createTerminal(c *gin.Context) (*model.CreateTerminalResponse, error) {
if server == nil {
return nil, singleton.Localizer.ErrorT("server not found or not connected")
}
stream := server.GetTaskStream()
if stream == nil {
if server.GetTaskStream() == nil {
return nil, singleton.Localizer.ErrorT("server not found or not connected")
}
@@ -53,7 +52,7 @@ func createTerminal(c *gin.Context) (*model.CreateTerminalResponse, error) {
terminalData, _ := json.Marshal(&model.TerminalTask{
StreamID: streamId,
})
if err := stream.Send(&proto.Task{
if err := server.SendTask(&proto.Task{
Type: model.TaskTypeTerminalGRPC,
Data: string(terminalData),
}); err != nil {
@@ -80,7 +79,7 @@ func terminalStream(c *gin.Context) (any, error) {
// (or an admin). Without this, any authenticated user who learns a stream
// UUID — via Referer leak, access logs, browser history — can hijack a live
// terminal and gain shell access to the target server.
if !rpc.NezhaHandlerSingleton.IsStreamAuthorizedForUser(streamId, getUid(c), callerIsAdmin(c)) {
if !streamAttachAllowedForRequest(c, streamId) {
return nil, singleton.Localizer.ErrorT("permission denied")
}
if _, err := rpc.NezhaHandlerSingleton.GetStream(streamId); err != nil {
@@ -95,6 +94,9 @@ func terminalStream(c *gin.Context) (any, error) {
defer wsConn.Close()
conn := websocketx.NewConn(wsConn)
deregisterPAT := registerPATConnection(c, func() { _ = wsConn.Close() })
defer deregisterPAT()
go func() {
// PING 保活
for {
+14
View File
@@ -78,6 +78,9 @@ func cancelServerTransfer(c *gin.Context) (*model.ServerTransfer, error) {
if err := q.First(&t, tid).Error; err != nil {
return nil, singleton.Localizer.ErrorT("permission denied")
}
if !t.HasPermission(c) {
return nil, singleton.Localizer.ErrorT("permission denied")
}
updated, err := singleton.ServerTransferShared.Cancel(tid)
if err != nil {
@@ -132,6 +135,14 @@ func retryServerTransfer(c *gin.Context) (*model.ServerTransfer, error) {
return nil, newGormError("%v", err)
}
// PAT server_ids 白名单必须在 admin short-circuit 之后再收一次,否则
// admin 给自己签发的“仅 server_ids={X}”PAT 仍能 retry 任意历史 transfer
// 行,与 ServerTransfer.HasPermission 注释和 cancelServerTransfer 已有
// 的复核语义直接冲突。
if !prev.HasPermission(c) {
return nil, singleton.Localizer.ErrorT("permission denied")
}
return singleton.ServerTransferShared.Retry(&prev, getUid(c))
}
@@ -166,6 +177,9 @@ func transferStream(c *gin.Context) (any, error) {
}
defer conn.Close()
deregisterPAT := registerPATConnection(c, func() { _ = conn.Close() })
defer deregisterPAT()
subID, ch := singleton.ServerTransferShared.Subscribe()
defer singleton.ServerTransferShared.Unsubscribe(subID)
@@ -0,0 +1,104 @@
package controller
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"strconv"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/nezhahq/nezha/model"
"github.com/nezhahq/nezha/service/singleton"
)
// cancelServerTransfer 的核心租户安全语义:
// - admin 可以取消任意 transfer 行
// - member 只能取消自己作为 FromUserID 的 transfer
// - 行不存在 vs 行存在但调用者不是 FromUserID 必须返回**相同的** "permission denied"
// 避免通过响应差异枚举 transfer ID 是否存在
//
// 该 handler 已有保护(transfer.go:73-80),但此前没有任何测试盯住它。
func seedPendingTransfer(t *testing.T, serverID, fromUID, toUID, initUID uint64) uint64 {
t.Helper()
tr := &model.ServerTransfer{
ServerID: serverID,
FromUserID: fromUID,
ToUserID: toUID,
InitiatorID: initUID,
Status: model.ServerTransferStatusPending,
}
assert.NoError(t, singleton.DB.Create(tr).Error)
singleton.ServerTransferShared.Register(tr)
return tr.ID
}
func callCancelServerTransfer(t *testing.T, transferID, callerID uint64, role model.Role) (commonResponseShape, int) {
t.Helper()
r := gin.New()
r.Use(func(c *gin.Context) {
setAuthUser(c, callerID, role)
c.Next()
})
r.POST("/transfer/:id/cancel", commonHandler(cancelServerTransfer))
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost,
"/transfer/"+strconv.FormatUint(transferID, 10)+"/cancel",
bytes.NewReader(nil))
r.ServeHTTP(w, req)
var resp commonResponseShape
assert.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
return resp, w.Code
}
func TestCancelServerTransfer_MemberCancelsOwnTransfer(t *testing.T) {
cleanup := setupRetryServerTransferFixture(t)
defer cleanup()
seedServer(t, 1, 100)
id := seedPendingTransfer(t, 1, 100, 200, 100)
resp, status := callCancelServerTransfer(t, id, 100, model.RoleMember)
assert.Equal(t, http.StatusOK, status)
assert.True(t, resp.Success, "FromUserID member must be able to cancel own transfer: %s", resp.Error)
}
func TestCancelServerTransfer_MemberCannotCancelOthers(t *testing.T) {
cleanup := setupRetryServerTransferFixture(t)
defer cleanup()
seedServer(t, 1, 100)
id := seedPendingTransfer(t, 1, 100, 200, 100)
resp, status := callCancelServerTransfer(t, id, 200, model.RoleMember)
assert.Equal(t, http.StatusOK, status)
assert.False(t, resp.Success, "ToUserID member must NOT be able to cancel another user's transfer")
assert.Contains(t, resp.Error, "permission denied")
}
func TestCancelServerTransfer_MemberCannotEnumerateNonexistentIDs(t *testing.T) {
cleanup := setupRetryServerTransferFixture(t)
defer cleanup()
resp, status := callCancelServerTransfer(t, 99999, 100, model.RoleMember)
assert.Equal(t, http.StatusOK, status)
assert.False(t, resp.Success)
assert.Contains(t, resp.Error, "permission denied",
"nonexistent transfer must return the SAME error as 'not your transfer', "+
"so an attacker can't probe which transfer IDs exist")
}
func TestCancelServerTransfer_AdminCancelsAny(t *testing.T) {
cleanup := setupRetryServerTransferFixture(t)
defer cleanup()
seedServer(t, 1, 100)
id := seedPendingTransfer(t, 1, 100, 200, 100)
resp, status := callCancelServerTransfer(t, id, 999, model.RoleAdmin)
assert.Equal(t, http.StatusOK, status)
assert.True(t, resp.Success, "admin must be able to cancel any transfer: %s", resp.Error)
}
@@ -0,0 +1,156 @@
package controller
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"strconv"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/model"
)
func newPATCtxSetter(callerID uint64, role model.Role, tok *model.APIToken) gin.HandlerFunc {
return func(c *gin.Context) {
setAuthUser(c, callerID, role)
if tok != nil {
c.Set(model.CtxKeyAPIToken, tok)
c.Set(apiTokenCtxKey, tok)
}
c.Next()
}
}
func callListTransferWithPAT(t *testing.T, callerID uint64, tok *model.APIToken) ([]*model.ServerTransfer, bool, string) {
t.Helper()
r := gin.New()
r.Use(newPATCtxSetter(callerID, model.RoleMember, tok))
r.GET("/transfer", listHandler(listServerTransfer))
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/transfer", nil)
r.ServeHTTP(w, req)
var resp struct {
Success bool `json:"success"`
Error string `json:"error"`
Data []*model.ServerTransfer `json:"data"`
}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
return resp.Data, resp.Success, resp.Error
}
func callCancelTransferWithPAT(t *testing.T, transferID, callerID uint64, tok *model.APIToken) (commonResponseShape, int) {
t.Helper()
r := gin.New()
r.Use(newPATCtxSetter(callerID, model.RoleMember, tok))
r.POST("/transfer/:id/cancel", commonHandler(cancelServerTransfer))
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost,
"/transfer/"+strconv.FormatUint(transferID, 10)+"/cancel",
bytes.NewReader(nil))
r.ServeHTTP(w, req)
var resp commonResponseShape
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
return resp, w.Code
}
func TestListServerTransfer_HidesRowsForServersOutsidePATWhitelist(t *testing.T) {
cleanup := setupRetryServerTransferFixture(t)
defer cleanup()
seedServer(t, 1, 100)
seedServer(t, 2, 100)
insideID := seedPendingTransfer(t, 1, 100, 200, 100)
outsideID := seedPendingTransfer(t, 2, 100, 200, 100)
tok := &model.APIToken{ID: 17, UserID: 100}
tok.SetServerIDs([]uint64{1})
rows, ok, errStr := callListTransferWithPAT(t, 100, tok)
assert.True(t, ok, "list call must succeed: %s", errStr)
seen := map[uint64]bool{}
for _, r := range rows {
seen[r.ID] = true
}
assert.True(t, seen[insideID],
"transfer of whitelisted server 1 must still be visible (got %d rows)", len(rows))
assert.False(t, seen[outsideID],
"transfer of non-whitelisted server 2 must be hidden from PAT view (rows=%+v)", rows)
}
func TestCancelServerTransfer_DeniesServerOutsidePATWhitelist(t *testing.T) {
cleanup := setupRetryServerTransferFixture(t)
defer cleanup()
seedServer(t, 1, 100)
seedServer(t, 2, 100)
_ = seedPendingTransfer(t, 1, 100, 200, 100)
outsideID := seedPendingTransfer(t, 2, 100, 200, 100)
tok := &model.APIToken{ID: 17, UserID: 100}
tok.SetServerIDs([]uint64{1})
resp, status := callCancelTransferWithPAT(t, outsideID, 100, tok)
assert.Equal(t, http.StatusOK, status)
assert.False(t, resp.Success,
"PAT whitelist [1] must not allow cancelling transfer of server 2 (FromUserID match alone is not enough)")
assert.Contains(t, resp.Error, "permission denied")
}
// admin PAT 同样必须受 server_ids 收窄:admin 给自己签的 PAT 加上 ServerIDs={1}
// 后,列表/取消都不能再触达白名单外的 server。这是修复 ServerTransfer.HasPermission
// 在 admin 早返回前未检查 PAT 的回归用例。
func callListTransferWithAdminPAT(t *testing.T, callerID uint64, tok *model.APIToken) ([]*model.ServerTransfer, bool, string) {
t.Helper()
r := gin.New()
r.Use(func(c *gin.Context) {
setAuthUser(c, callerID, model.RoleAdmin)
if tok != nil {
c.Set(model.CtxKeyAPIToken, tok)
c.Set(apiTokenCtxKey, tok)
}
c.Next()
})
r.GET("/transfer", listHandler(listServerTransfer))
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/transfer", nil)
r.ServeHTTP(w, req)
var resp struct {
Success bool `json:"success"`
Error string `json:"error"`
Data []*model.ServerTransfer `json:"data"`
}
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
return resp.Data, resp.Success, resp.Error
}
func TestListServerTransfer_AdminPATIsAlsoNarrowedByWhitelist(t *testing.T) {
cleanup := setupRetryServerTransferFixture(t)
defer cleanup()
seedServer(t, 1, 100)
seedServer(t, 2, 100)
insideID := seedPendingTransfer(t, 1, 100, 200, 100)
outsideID := seedPendingTransfer(t, 2, 100, 200, 100)
tok := &model.APIToken{ID: 18, UserID: 999}
tok.SetServerIDs([]uint64{1})
rows, ok, errStr := callListTransferWithAdminPAT(t, 999, tok)
assert.True(t, ok, "list call must succeed: %s", errStr)
seen := map[uint64]bool{}
for _, r := range rows {
seen[r.ID] = true
}
assert.True(t, seen[insideID], "admin PAT scoped to {1} must still see transfer of server 1")
assert.False(t, seen[outsideID],
"admin PAT scoped to {1} must NOT see transfer of server 2 (admin early-return is no longer a bypass)")
}

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