mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 17:50:12 +00:00
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:
@@ -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}")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,135 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
)
|
||||
|
||||
// jwtOrPATAuthMiddleware 把 PAT 与 JWT 两条鉴权链组合到 /api/v1/* 入口。
|
||||
//
|
||||
// 处理顺序:
|
||||
// 1. apiTokenAuthMiddleware:识别 `Authorization: Bearer nzp_*`。命中(合法 PAT)
|
||||
// 把 user 挂到 ctx;非法 PAT 直接 abort 401。
|
||||
// 2. 如果 PAT 已挂 user → 跳过 JWT。
|
||||
// 3. 否则 → JWT 中间件接管,按现有 cookie / Bearer / query token 逻辑鉴权。
|
||||
//
|
||||
// 存量 JWT 客户端零感知;新 PAT 客户端可直接调 REST,但每个端点的 scope
|
||||
// 仍由 restScopeMiddleware 控制。
|
||||
func jwtOrPATAuthMiddleware(patMw, jwtMw gin.HandlerFunc) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
patMw(c)
|
||||
if c.IsAborted() {
|
||||
return
|
||||
}
|
||||
if APITokenFromContext(c) != nil {
|
||||
return
|
||||
}
|
||||
jwtMw(c)
|
||||
if c.IsAborted() {
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// patOrFallbackAuthMiddleware 是 optional 路由(ForceAuth=false 时也能匿名访问)
|
||||
// 的鉴权链:
|
||||
// 1. apiTokenAuthMiddleware:识别 PAT,命中后挂 user;非法 PAT 401 abort。
|
||||
// 2. 已挂 PAT → 跳过 JWT,restScopeMiddleware 会按 scope 收口。
|
||||
// 3. 未带 PAT → 走 fallbackJwtMw,存在 JWT 则挂 user,没有就匿名继续。
|
||||
//
|
||||
// 这是修复 ForceAuth=false 时 optional 路由完全不解析 PAT 的关键:
|
||||
// 之前直接用 fallbackAuthMw 会让 PAT 请求被当作 guest,scope 形同虚设。
|
||||
func patOrFallbackAuthMiddleware(patMw, fallbackJwtMw gin.HandlerFunc) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
patMw(c)
|
||||
if c.IsAborted() {
|
||||
return
|
||||
}
|
||||
if APITokenFromContext(c) != nil {
|
||||
return
|
||||
}
|
||||
fallbackJwtMw(c)
|
||||
}
|
||||
}
|
||||
|
||||
// restScopeMiddleware 在 /api/v1/* 路由上 enforce PAT scope。
|
||||
//
|
||||
// 行为:
|
||||
// - JWT 持有者(任何来源:cookie / Authorization Bearer 非 nzp_)→ 直接放行,
|
||||
// 沿用 JWT 模型的完整权限。
|
||||
// - PAT 持有者 → 必须命中给定 scope。命中后下游 handler 仍受 user 级权限
|
||||
// 检查(adminHandler / Server.HasPermission),scope 只能收窄不能放大。
|
||||
// - PAT 持有者遇到 scope=="" → 直接 403。空字符串作为 fail-closed 默认值,
|
||||
// 防止接入新路由时忘填 scope 把 PAT 静默放行。
|
||||
//
|
||||
// 因此"自我管理"端点(/profile、/api-tokens、/refresh-token 等)必须显式挂
|
||||
// restPATForbiddenMiddleware 来拒绝 PAT,而不是依赖空 scope 兜底。
|
||||
func restScopeMiddleware(scope string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
tok := APITokenFromContext(c)
|
||||
if tok == nil {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
if scope == "" || !tok.HasScope(scope) {
|
||||
c.AbortWithStatusJSON(http.StatusForbidden, model.CommonResponse[any]{
|
||||
Success: false,
|
||||
Error: "ApiErrorForbidden: api token lacks scope " + scope,
|
||||
})
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// restScopeAllOf is the multi-scope variant of restScopeMiddleware. It
|
||||
// gates on EVERY listed scope, used by routes whose semantics span more
|
||||
// than one capability — file-manager sessions read, write AND delete
|
||||
// files, so a PAT that only carries nezha:server:write must NOT be allowed
|
||||
// to open one. JWT callers pass through unchanged.
|
||||
func restScopeAllOf(scopes ...string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
tok := APITokenFromContext(c)
|
||||
if tok == nil {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
for _, scope := range scopes {
|
||||
if scope == "" || !tok.HasScope(scope) {
|
||||
c.AbortWithStatusJSON(http.StatusForbidden, model.CommonResponse[any]{
|
||||
Success: false,
|
||||
Error: "ApiErrorForbidden: api token lacks scope " + scope,
|
||||
})
|
||||
return
|
||||
}
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// serverConfigSensitiveScope 收紧 GET /server/config/:id 的 PAT scope 到
|
||||
// ScopeServerWrite:返回体里包含 client_secret 等下发到 agent 的凭据,单纯
|
||||
// nezha:server:read 不应足以读取。命名刻意带 Sensitive 而不是 Read,避免下
|
||||
// 个维护者把它当成普通 read scope 还原成 ScopeServerRead 重新打开提权链。
|
||||
func serverConfigSensitiveScope() string { return model.ScopeServerWrite }
|
||||
|
||||
// restPATForbiddenMiddleware 在「自我管理」类端点上显式拒绝 PAT。
|
||||
//
|
||||
// 这些端点(profile / api-tokens / oauth2 绑定 / refresh-token)一旦允许 PAT
|
||||
// 自调,就可形成提权链(PAT → 创建更高权限 PAT → ...)。
|
||||
// 显式 403 比静默放行更安全。
|
||||
func restPATForbiddenMiddleware() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
if APITokenFromContext(c) != nil {
|
||||
c.AbortWithStatusJSON(http.StatusForbidden, model.CommonResponse[any]{
|
||||
Success: false,
|
||||
Error: "ApiErrorForbidden: this endpoint is not accessible by api token",
|
||||
})
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
)
|
||||
|
||||
// restScopeMiddleware 的"空 scope"实际行为:PAT 调用方一律 403。
|
||||
// 这条测试把注释与实现的契约对齐:
|
||||
// 1. 实际行为:PAT + scope="" → 403。
|
||||
// 2. 文档约束:源码注释必须明确说出"空 scope 对 PAT 仍被拒绝",
|
||||
// 不能再保留"空字符串 = 放行"这种与实现相反的旧措辞。
|
||||
func TestRestScopeMiddleware_EmptyScopeRejectsPAT(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
r.GET("/x", func(c *gin.Context) {
|
||||
c.Set(apiTokenCtxKey, &model.APIToken{ID: 1})
|
||||
c.Set(model.CtxKeyAPIToken, &model.APIToken{ID: 1})
|
||||
c.Next()
|
||||
}, restScopeMiddleware(""), func(c *gin.Context) {
|
||||
c.Status(http.StatusOK)
|
||||
})
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest("GET", "/x", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
if w.Code != http.StatusForbidden {
|
||||
t.Fatalf("expected 403 when PAT hits restScopeMiddleware(\"\"); got %d body=%q", w.Code, w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestRestScopeMiddleware_DocReflectsEmptyScopeRejection(t *testing.T) {
|
||||
wd, err := os.Getwd()
|
||||
if err != nil {
|
||||
t.Fatalf("getwd: %v", err)
|
||||
}
|
||||
b, err := os.ReadFile(filepath.Join(wd, "api_token_scope.go"))
|
||||
if err != nil {
|
||||
t.Fatalf("read: %v", err)
|
||||
}
|
||||
src := string(b)
|
||||
idx := strings.Index(src, "func restScopeMiddleware(")
|
||||
if idx < 0 {
|
||||
t.Fatalf("restScopeMiddleware not found")
|
||||
}
|
||||
doc := src[:idx]
|
||||
if strings.Contains(doc, "空字符串)= 放行") || strings.Contains(doc, "空字符串) = 放行") {
|
||||
t.Fatalf("doc still claims empty scope means 放行; this contradicts the implementation which 403s PAT callers")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,223 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/service/singleton"
|
||||
)
|
||||
|
||||
// 在 /api/v1 风格的小 router 上重现 PAT + scope mw,验证 enforcement。
|
||||
func setupRESTScopeServer(t *testing.T) (*httptest.Server, string, func()) {
|
||||
t.Helper()
|
||||
cleanupBase, uid := setupMCPTest(t)
|
||||
|
||||
_, plain := mkToken(t, uid, []string{model.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")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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-JWT;ForceAuth=false 走
|
||||
// PAT-or-FallbackJWT,保证两种模式下 PAT 都会被解析,restScopeMiddleware
|
||||
// 才能按 scope 真实收口(否则匿名 PAT 请求会被当 guest,scope 失效)。
|
||||
optionalAuthMw := utils.IfOr(singleton.Conf.ForceAuth, authMw, patOrFallbackAuthMiddleware(patMw, fallbackAuthMw))
|
||||
|
||||
optionalAuth := api.Group("", optionalAuthMw)
|
||||
optionalAuth.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 资源族 scope(read/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 / 链接预览会以为站点挂了)。
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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)")
|
||||
}
|
||||
}
|
||||
@@ -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 404(body 还是 index.html)。
|
||||
// controller.go 旁边的注释明确说「新增前端路由时必须在 main.tsx 与这里同步加」。
|
||||
func TestFallbackToFrontend_APITokensRouteReturns200(t *testing.T) {
|
||||
t.Chdir(t.TempDir())
|
||||
router := newFrontendFallbackTestRouter(t)
|
||||
|
||||
w := performFrontendFallbackRequest(t, router, "/dashboard/settings/api-tokens")
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("/dashboard/settings/api-tokens fallback status = %d, want 200 "+
|
||||
"(front-end main.tsx registered the route — backend SPA fallback regex must mirror it)",
|
||||
w.Code)
|
||||
}
|
||||
if !strings.Contains(w.Body.String(), "admin index") {
|
||||
t.Fatalf("/dashboard/settings/api-tokens must serve admin index.html, got %q", w.Body.String())
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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;闸 2(PAT scope)由
|
||||
// mcpTool.RequiredScope 在 dispatch 之前过滤。
|
||||
package controller
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/service/rpc"
|
||||
"github.com/nezhahq/nezha/service/singleton"
|
||||
)
|
||||
|
||||
// --- JSON-RPC 2.0 wire types ---
|
||||
|
||||
type jsonRPCRequest struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID json.RawMessage `json:"id,omitempty"`
|
||||
Method string `json:"method"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
}
|
||||
|
||||
type jsonRPCResponse struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID json.RawMessage `json:"id,omitempty"`
|
||||
Result any `json:"result,omitempty"`
|
||||
Error *jsonRPCError `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type jsonRPCError struct {
|
||||
Code int `json:"code"`
|
||||
Message string `json:"message"`
|
||||
Data any `json:"data,omitempty"`
|
||||
}
|
||||
|
||||
const (
|
||||
// JSON-RPC 标准错误码
|
||||
rpcErrParse = -32700
|
||||
rpcErrInvalidRequest = -32600
|
||||
rpcErrMethodNotFound = -32601
|
||||
rpcErrInvalidParams = -32602
|
||||
rpcErrInternal = -32603
|
||||
// MCP 自定义错误码(>= -32000 高位段)
|
||||
rpcErrUnauthorized = -32001
|
||||
rpcErrForbidden = -32002
|
||||
)
|
||||
|
||||
// mcpJSONRPCMaxBodyBytes caps the JSON-RPC envelope size at the dashboard
|
||||
// edge. Real fs.write base64 content goes through fs.transfer (capped
|
||||
// separately by model.MCPFsTransferMaxSize) so tools/call params here are
|
||||
// always small. The cap is intentionally generous (8 MiB) to allow
|
||||
// per-request batched arguments while making OOM-via-decode impossible.
|
||||
const mcpJSONRPCMaxBodyBytes = 8 * 1024 * 1024
|
||||
|
||||
// --- MCP types ---
|
||||
|
||||
// mcpServerInfo MCP initialize 响应的 serverInfo 字段。
|
||||
type mcpServerInfo struct {
|
||||
Name string `json:"name"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
|
||||
type mcpInitializeResult struct {
|
||||
ProtocolVersion string `json:"protocolVersion"`
|
||||
Capabilities map[string]any `json:"capabilities"`
|
||||
ServerInfo mcpServerInfo `json:"serverInfo"`
|
||||
}
|
||||
|
||||
// mcpToolDescriptor 是 tools/list 返回的单条 tool 描述。
|
||||
type mcpToolDescriptor struct {
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
InputSchema map[string]any `json:"inputSchema"`
|
||||
}
|
||||
|
||||
// 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 discovery;JSON-RPC body 保留旧字段
|
||||
// 不打破 ScopeDenied 类内部断言。
|
||||
writeJSONRPCErrorWithStatus(c, nil, rpcErrUnauthorized, "missing or invalid API token", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
// MaxBytesReader 必须夹在 PAT 校验通过后、ShouldBindJSON 之前——
|
||||
// 校验前限流可能让攻击者用伪造 token 触发 audit;校验后限流既挡住合法
|
||||
// PAT 的 OOM,又不会让匿名请求走到 audit 路径。
|
||||
if c.Request != nil && c.Request.Body != nil {
|
||||
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, mcpJSONRPCMaxBodyBytes)
|
||||
}
|
||||
|
||||
// Consume the per-token budget before validating the request so malformed
|
||||
// envelopes and malformed tools/call params cannot flood the dashboard
|
||||
// without counting against the limiter. The outcome is applied after the
|
||||
// method is known so tools/call still surfaces the rate limit as a tool
|
||||
// error rather than a transport-level error.
|
||||
rateLimited := !mcpRateLimiterShared.Allow(tok.ID)
|
||||
|
||||
var req jsonRPCRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
if errors.Is(err, errors.New("http: request body too large")) || strings.Contains(err.Error(), "http: request body too large") {
|
||||
writeJSONRPCErrorWithStatus(c, nil, rpcErrInvalidRequest, "request body exceeds MCP envelope size limit", http.StatusRequestEntityTooLarge)
|
||||
return
|
||||
}
|
||||
// 限流优先:method 无从得知时,over-budget 请求即便 body 畸形也必须
|
||||
// 走 429,否则攻击者能用畸形 body 在不计入限额的情况下持续刷 parse error。
|
||||
if rateLimited {
|
||||
writeJSONRPCErrorWithStatus(c, nil, rpcErrForbidden, "rate limit exceeded for this token", http.StatusTooManyRequests)
|
||||
return
|
||||
}
|
||||
writeJSONRPCError(c, nil, rpcErrParse, "invalid json-rpc envelope: "+err.Error())
|
||||
return
|
||||
}
|
||||
if req.JSONRPC != "2.0" || req.Method == "" {
|
||||
if rateLimited {
|
||||
writeJSONRPCErrorWithStatus(c, req.ID, rpcErrForbidden, "rate limit exceeded for this token", http.StatusTooManyRequests)
|
||||
return
|
||||
}
|
||||
writeJSONRPCError(c, req.ID, rpcErrInvalidRequest, "invalid json-rpc envelope")
|
||||
return
|
||||
}
|
||||
|
||||
if rateLimited {
|
||||
if req.Method == "tools/call" {
|
||||
writeToolCallError(c, req.ID, model.MCPOutcomeRateLimited, "rate limit exceeded for this token")
|
||||
return
|
||||
}
|
||||
writeJSONRPCErrorWithStatus(c, req.ID, rpcErrForbidden, "rate limit exceeded for this token", http.StatusTooManyRequests)
|
||||
return
|
||||
}
|
||||
|
||||
switch req.Method {
|
||||
case "initialize":
|
||||
writeJSONRPCResult(c, req.ID, mcpInitializeResult{
|
||||
ProtocolVersion: "2024-11-05",
|
||||
Capabilities: map[string]any{
|
||||
"tools": map[string]any{"listChanged": false},
|
||||
},
|
||||
ServerInfo: mcpServerInfo{
|
||||
Name: "nezha-mcp",
|
||||
Version: singleton.Version,
|
||||
},
|
||||
})
|
||||
case "notifications/initialized", "ping":
|
||||
// 客户端通知或心跳;JSON-RPC 通知没有 id,但 ping 有 id 时返回空 result
|
||||
if len(req.ID) > 0 && string(req.ID) != "null" {
|
||||
writeJSONRPCResult(c, req.ID, struct{}{})
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusAccepted)
|
||||
case "tools/list":
|
||||
writeJSONRPCResult(c, req.ID, mcpToolsListResult{
|
||||
Tools: buildToolDescriptors(),
|
||||
})
|
||||
case "tools/call":
|
||||
handleToolsCall(c, &req, tok)
|
||||
default:
|
||||
writeJSONRPCError(c, req.ID, rpcErrMethodNotFound, "method not supported: "+req.Method)
|
||||
}
|
||||
}
|
||||
|
||||
func buildToolDescriptors() []mcpToolDescriptor {
|
||||
tools := listRegisteredMCPTools()
|
||||
out := make([]mcpToolDescriptor, 0, len(tools))
|
||||
for _, t := range tools {
|
||||
out = append(out, mcpToolDescriptor{
|
||||
Name: t.Name,
|
||||
Description: t.Description,
|
||||
InputSchema: t.InputSchema,
|
||||
})
|
||||
}
|
||||
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_id(best-effort,只用于审计字段)。
|
||||
func extractServerID(raw json.RawMessage) uint64 {
|
||||
if len(raw) == 0 {
|
||||
return 0
|
||||
}
|
||||
var probe struct {
|
||||
ServerID uint64 `json:"server_id"`
|
||||
}
|
||||
_ = json.Unmarshal(raw, &probe)
|
||||
return probe.ServerID
|
||||
}
|
||||
|
||||
func truncateString(s string, max int) string {
|
||||
if len(s) <= max {
|
||||
return s
|
||||
}
|
||||
return s[:max]
|
||||
}
|
||||
|
||||
// --- wire writers ---
|
||||
|
||||
func writeJSONRPCResult(c *gin.Context, id json.RawMessage, result any) {
|
||||
c.JSON(http.StatusOK, jsonRPCResponse{
|
||||
JSONRPC: "2.0",
|
||||
ID: id,
|
||||
Result: result,
|
||||
})
|
||||
}
|
||||
|
||||
func writeJSONRPCError(c *gin.Context, id json.RawMessage, code int, message string) {
|
||||
writeJSONRPCErrorWithStatus(c, id, code, message, http.StatusOK)
|
||||
}
|
||||
|
||||
func writeToolCallError(c *gin.Context, id json.RawMessage, errCode, errMsg string) {
|
||||
writeJSONRPCResult(c, id, mcpToolCallResult{
|
||||
Content: []mcpContent{{Type: "text", Text: errMsg}},
|
||||
IsError: true,
|
||||
StructuredContent: map[string]string{
|
||||
"error_code": errCode,
|
||||
"error": errMsg,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
func writeJSONRPCErrorWithStatus(c *gin.Context, id json.RawMessage, code int, message string, status int) {
|
||||
c.JSON(status, jsonRPCResponse{
|
||||
JSONRPC: "2.0",
|
||||
ID: id,
|
||||
Error: &jsonRPCError{Code: code, Message: message},
|
||||
})
|
||||
}
|
||||
|
||||
// --- 错误语义 ---
|
||||
|
||||
// mcpError 是 tool handler 可以返回的语义化错误。
|
||||
// dispatch 根据 Code 决定 audit outcome 与 JSON-RPC 错误码(如果命中 rpcErr* 域)。
|
||||
type mcpError struct {
|
||||
Code string
|
||||
Msg string
|
||||
}
|
||||
|
||||
func (e *mcpError) Error() string { return e.Msg }
|
||||
|
||||
func newMCPError(code, msg string) *mcpError { return &mcpError{Code: code, Msg: msg} }
|
||||
|
||||
// 预制错误
|
||||
var (
|
||||
errMCPInvalidArgs = func(s string) *mcpError { return newMCPError(model.MCPOutcomeInvalidArgs, s) }
|
||||
errMCPPermDenied = newMCPError(model.MCPOutcomePermDenied, "permission denied")
|
||||
errMCPScopeDenied = func(s string) *mcpError {
|
||||
return newMCPError(model.MCPOutcomeScopeDenied, "missing required scope: "+s)
|
||||
}
|
||||
errMCPServerOffline = newMCPError(model.MCPOutcomeServerOffline, "agent offline")
|
||||
errMCPAgentTimeout = newMCPError(model.MCPOutcomeAgentTimeout, "agent did not respond within timeout")
|
||||
errMCPUnsupported = newMCPError(model.MCPOutcomeUnsupportedAgent, "agent does not support this MCP capability; please upgrade the agent")
|
||||
)
|
||||
|
||||
// --- 共用工具 ---
|
||||
|
||||
var errNoToken = errors.New("no api token in context")
|
||||
|
||||
// decodeToolArgs 是 tool handler 用来反序列化 arguments 的辅助。
|
||||
func decodeToolArgs(raw json.RawMessage, out any) error {
|
||||
if len(raw) == 0 {
|
||||
return nil
|
||||
}
|
||||
if err := json.Unmarshal(raw, out); err != nil {
|
||||
return fmt.Errorf("invalid arguments: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// requireServerAccess 是 tool handler 共用的「闸 1 + 闸 2 服务器白名单」组合校验。
|
||||
// 通过返回 *model.Server;失败返回带语义 Code 的 mcpError,便于 dispatch 归类审计。
|
||||
func requireServerAccess(c *gin.Context, serverID uint64) (*model.Server, error) {
|
||||
if serverID == 0 {
|
||||
return nil, errMCPInvalidArgs("server_id required")
|
||||
}
|
||||
tok := APITokenFromContext(c)
|
||||
if tok != nil && !tok.CanAccessServer(serverID) {
|
||||
return nil, errMCPPermDenied
|
||||
}
|
||||
server, _ := singleton.ServerShared.Get(serverID)
|
||||
if server == nil {
|
||||
return nil, errMCPServerOffline
|
||||
}
|
||||
if !server.HasPermission(c) {
|
||||
return nil, errMCPPermDenied
|
||||
}
|
||||
return server, nil
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/service/singleton"
|
||||
)
|
||||
|
||||
// mcpAuditWrite 异步写一条 MCP 审计日志。失败仅 log,不阻塞业务。
|
||||
//
|
||||
// argsBytes:tool 的 raw JSON 参数(dispatcher 已经反序列化过)。
|
||||
// 只记录 sha256 全文哈希,不保留任何明文片段:server.exec 的 env/stdin、
|
||||
// fs.write 的 content 等字段会包含 token、密码、密钥、文件内容等敏感数据,
|
||||
// 任何长度的 peek 都可能让审计表本身成为 secret 仓库;以哈希做关联即可。
|
||||
//
|
||||
// 测试可以把 mcpAuditSync 置为 true 让写入同步,避免 goroutine 与测试 teardown
|
||||
// 形成竞态(不同测试 swap 全局 singleton.DB 时尤其明显)。
|
||||
func mcpAuditWrite(entry model.MCPAuditLog, argsBytes []byte) {
|
||||
if len(argsBytes) > 0 {
|
||||
sum := sha256.Sum256(argsBytes)
|
||||
entry.ArgsHash = hex.EncodeToString(sum[:])
|
||||
}
|
||||
entry.ArgsPeek = ""
|
||||
if entry.CreatedAt.IsZero() {
|
||||
entry.CreatedAt = time.Now()
|
||||
}
|
||||
db := singleton.DB
|
||||
write := func(e model.MCPAuditLog) {
|
||||
if db == nil {
|
||||
return
|
||||
}
|
||||
if err := db.Create(&e).Error; err != nil {
|
||||
log.Printf("NEZHA>> mcp audit write failed: %v", err)
|
||||
}
|
||||
}
|
||||
if mcpAuditSync {
|
||||
write(entry)
|
||||
return
|
||||
}
|
||||
go write(entry)
|
||||
}
|
||||
|
||||
// mcpAuditSync 仅供测试切换为同步写入,生产保持 false。
|
||||
var mcpAuditSync = false
|
||||
@@ -0,0 +1,72 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/service/singleton"
|
||||
)
|
||||
|
||||
// H7 regression: the MCP endpoint must cap incoming JSON-RPC body size
|
||||
// BEFORE decoding. Without this, a valid PAT can post a multi-GB body and
|
||||
// the dashboard exhausts memory in ShouldBindJSON. We assert the body
|
||||
// reader is wrapped in http.MaxBytesReader; the exact error path the
|
||||
// decoder takes is irrelevant as long as the cap is enforced.
|
||||
func TestMCPEndpoint_BodyIsCappedByMaxBytesReader(t *testing.T) {
|
||||
prevConf := singleton.Conf
|
||||
cfg := &model.Config{}
|
||||
cfg.SetMCPEnabled(true)
|
||||
singleton.Conf = &singleton.ConfigClass{Config: cfg}
|
||||
t.Cleanup(func() { singleton.Conf = prevConf })
|
||||
|
||||
tok := &model.APIToken{ID: 1, ScopesCSV: "nezha:server:read"}
|
||||
// 16 MiB of valid-JSON whitespace prefix forces the decoder to actually
|
||||
// stream past the limit, exercising MaxBytesReader.
|
||||
body := `{"jsonrpc":"2.0","id":1,"method":"initialize","params":"` +
|
||||
strings.Repeat("x", mcpJSONRPCMaxBodyBytes+1024) + `"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp", bytes.NewBufferString(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(w)
|
||||
c.Request = req
|
||||
c.Set(apiTokenCtxKey, tok)
|
||||
c.Set(model.CtxKeyAPIToken, tok)
|
||||
|
||||
mcpEndpoint(c)
|
||||
|
||||
if !strings.Contains(w.Body.String(), "request body") &&
|
||||
w.Code != http.StatusRequestEntityTooLarge {
|
||||
t.Fatalf("oversized body must be rejected with a body-size error, got code=%d body=%s",
|
||||
w.Code, w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestMCPEndpoint_AcceptsSmallBody(t *testing.T) {
|
||||
prevConf := singleton.Conf
|
||||
cfg := &model.Config{}
|
||||
cfg.SetMCPEnabled(true)
|
||||
singleton.Conf = &singleton.ConfigClass{Config: cfg}
|
||||
t.Cleanup(func() { singleton.Conf = prevConf })
|
||||
|
||||
tok := &model.APIToken{ID: 1, ScopesCSV: "nezha:server:read"}
|
||||
body := `{"jsonrpc":"2.0","id":1,"method":"initialize"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp", bytes.NewBufferString(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(w)
|
||||
c.Request = req
|
||||
c.Set(apiTokenCtxKey, tok)
|
||||
c.Set(model.CtxKeyAPIToken, tok)
|
||||
|
||||
mcpEndpoint(c)
|
||||
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("small valid body must succeed, got code=%d body=%s", w.Code, w.Body.String())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
)
|
||||
|
||||
// MCPMinAgentVersion 是支持 MCP 的最低 agent 版本。
|
||||
//
|
||||
// release 流程:在 agent ship 了 MCP handlers 后,把此值更新为该 release 的版本号。
|
||||
// 不变量:必须为非空。否则旧 agent 收到 TaskTypeExec/TaskTypeFs* 等新任务类型时
|
||||
// 走 default 分支不回 TaskResult,dashboard 要等 CallAgent 超时(30s)甚至更久
|
||||
// (fs.transfer 的 IOStream attach 30s)才能感知,这是 server-transfer 已经
|
||||
// 通过 MinServerTransferAgentVersion 修复过的同类问题。
|
||||
const MCPMinAgentVersion = "v2.1.0"
|
||||
|
||||
// requireAgentSupportsMCP 在 tool handler 调 CallAgent 之前快速失败不支持的 agent。
|
||||
// 仅作为 UX 优化:真正的安全/正确性由 agent 端 task switch 的 default 分支保障。
|
||||
func requireAgentSupportsMCP(server *model.Server) error {
|
||||
if MCPMinAgentVersion == "" || server == nil || 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
|
||||
// 分支不回 TaskResult;dashboard 必须在调 CallAgent 之前依据 Host.Version
|
||||
// 快速失败,否则用户要等到 30s/24h timeout 才知道 agent 不支持。
|
||||
// MinServerTransferAgentVersion 已为同类问题在 transfer 路径上确立了
|
||||
// release-time 必填的版本下限——这里把 MCP 也纳入同一不变量。
|
||||
func TestMCPMinAgentVersionIsPinnedToRelease(t *testing.T) {
|
||||
require.NotEmpty(t, MCPMinAgentVersion,
|
||||
"MCPMinAgentVersion must be set to the lowest agent build that ships MCP handlers; an empty string disables the gate and lets old agents hang dashboard requests until timeout")
|
||||
}
|
||||
|
||||
func TestRequireAgentSupportsMCPRejectsBelowMinVersion(t *testing.T) {
|
||||
old := &model.Server{Host: &model.Host{Version: "v0.0.1"}}
|
||||
err := requireAgentSupportsMCP(old)
|
||||
require.Error(t, err, "agents older than MCPMinAgentVersion must be rejected before CallAgent")
|
||||
require.True(t, errors.Is(err, errMCPUnsupported) || err.Error() == errMCPUnsupported.Error(),
|
||||
"expected the errMCPUnsupported sentinel, got %v", err)
|
||||
}
|
||||
|
||||
func TestRequireAgentSupportsMCPAcceptsCurrentVersion(t *testing.T) {
|
||||
current := &model.Server{Host: &model.Host{Version: MCPMinAgentVersion}}
|
||||
require.NoError(t, requireAgentSupportsMCP(current),
|
||||
"server reporting exactly MCPMinAgentVersion must be accepted")
|
||||
}
|
||||
|
||||
// 钉住「最近一个不带 MCP handler 的已发布 agent tag (v2.0.4) 必须被拒绝」。
|
||||
// v2.0.4 的 model/task.go 还没有 TaskTypeExec/TaskTypeFs* 常量,cmd/agent/
|
||||
// mcp_handlers.go 也不存在;如果版本门槛把它放行,dashboard 调 MCP tool
|
||||
// 后 agent 会走 default 分支不回 TaskResult,CallAgent 必须等 30s 超时。
|
||||
func TestRequireAgentSupportsMCPRejectsLastReleaseWithoutMCP(t *testing.T) {
|
||||
noMCP := &model.Server{Host: &model.Host{Version: "v2.0.4"}}
|
||||
err := requireAgentSupportsMCP(noMCP)
|
||||
require.Error(t, err,
|
||||
"v2.0.4 is the latest released agent tag that ships *without* MCP handlers; bumping MCPMinAgentVersion below the first MCP release re-introduces the silent-timeout bug")
|
||||
require.True(t, errors.Is(err, errMCPUnsupported) || err.Error() == errMCPUnsupported.Error(),
|
||||
"expected the errMCPUnsupported sentinel, got %v", err)
|
||||
}
|
||||
|
||||
func TestRequireAgentSupportsMCPDefersWhenAgentNeverReported(t *testing.T) {
|
||||
require.NoError(t, requireAgentSupportsMCP(&model.Server{Host: nil}),
|
||||
"Host==nil means agent never reported its build; defer the version decision to the CallAgent timeout layer")
|
||||
}
|
||||
@@ -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 喂给前端 fallback(HTML/404),
|
||||
// 真实 MCP 客户端在自动探测 SSE 时会卡住或拿到无效内容。
|
||||
//
|
||||
// 这条测试拼出和生产 routers() 一致的 /mcp 三件套,仅断言「非 POST 不返回 HTML」。
|
||||
type mcpFallbackDist struct{}
|
||||
|
||||
func (mcpFallbackDist) Open(string) (fs.File, error) { return nil, fs.ErrNotExist }
|
||||
|
||||
func setupMCPMethodRouter(t *testing.T) *gin.Engine {
|
||||
t.Helper()
|
||||
originalConf := singleton.Conf
|
||||
singleton.Conf = &singleton.ConfigClass{Config: &model.Config{
|
||||
ConfigDashboard: model.ConfigDashboard{
|
||||
AdminTemplate: "admin-dist",
|
||||
UserTemplate: "user-dist",
|
||||
},
|
||||
}}
|
||||
t.Cleanup(func() { singleton.Conf = originalConf })
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
r.POST("/mcp", mcpOriginGuard(), apiTokenAuthMiddleware(), mcpEndpoint)
|
||||
r.GET("/mcp", mcpMethodNotAllowed)
|
||||
r.DELETE("/mcp", mcpMethodNotAllowed)
|
||||
r.GET("/mcp/download/:token", mcpOriginGuard(), transferDownloadHandler)
|
||||
r.POST("/mcp/upload/:token", mcpOriginGuard(), transferUploadHandler)
|
||||
r.NoRoute(fallbackToFrontend(mcpFallbackDist{}))
|
||||
return r
|
||||
}
|
||||
|
||||
func TestMCP_GetReturnsMethodNotAllowed(t *testing.T) {
|
||||
t.Chdir(t.TempDir())
|
||||
r := setupMCPMethodRouter(t)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodGet, "/mcp", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusMethodNotAllowed {
|
||||
t.Fatalf("GET /mcp must return 405 per Streamable HTTP spec; got %d body=%q",
|
||||
w.Code, w.Body.String())
|
||||
}
|
||||
if strings.Contains(strings.ToLower(w.Body.String()), "<html") {
|
||||
t.Fatalf("GET /mcp must not fall back to the SPA index.html; body=%q", w.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestMCP_DeleteReturnsMethodNotAllowed(t *testing.T) {
|
||||
t.Chdir(t.TempDir())
|
||||
r := setupMCPMethodRouter(t)
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req := httptest.NewRequest(http.MethodDelete, "/mcp", nil)
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
if w.Code != http.StatusMethodNotAllowed {
|
||||
t.Fatalf("DELETE /mcp (session terminate) must return 405 when sessions are not implemented; got %d",
|
||||
w.Code)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,117 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/service/singleton"
|
||||
)
|
||||
|
||||
func mcpOriginGuard() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
origin := strings.TrimSpace(c.GetHeader("Origin"))
|
||||
if origin == "" {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
u, err := url.Parse(origin)
|
||||
if err != nil || u.Host == "" {
|
||||
abortOrigin(c)
|
||||
return
|
||||
}
|
||||
if !strings.EqualFold(u.Host, c.Request.Host) {
|
||||
abortOrigin(c)
|
||||
return
|
||||
}
|
||||
// DNS rebinding 防线:当 dashboard 明确仅暴露在 loopback 上时,浏览器只
|
||||
// 应通过 loopback hostname 命中本服务。任何带 Origin 的请求若 Host 不是
|
||||
// loopback 字面量,说明它解析自一个对外 DNS 名(典型攻击:evil.example
|
||||
// 解析到 127.0.0.1,Host == Origin == evil.example 等式成立但实际打的是
|
||||
// 用户本机 dashboard)。绑公网 IP / 未指定地址 / 公网 hostname 的部署
|
||||
// 跳过此检查 —— 那种部署下 Host 本来就是公网域名,再要求 loopback
|
||||
// 会把合法的同源前端请求一刀切。
|
||||
if dashboardListensOnLoopback() && !isLoopbackHostname(hostnameOnly(c.Request.Host)) {
|
||||
abortOrigin(c)
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// dashboardListensOnLoopback 判定 dashboard 是否仅暴露在 loopback 上。
|
||||
//
|
||||
// 之前把未指定地址(空字符串、`0.0.0.0`、`::`、`*`)也当 loopback 部署,理由是
|
||||
// 这些场景下浏览器仍可能通过 127.0.0.1 访问,rebinding 风险存在。问题是绝大多数
|
||||
// 生产部署就是 `listen_host=""` / `0.0.0.0`,配合公网域名访问;那种部署下浏览器
|
||||
// 同源 POST 的 Host 是公网域名而不是 loopback,旧逻辑会把合法的同源前端请求一刀切。
|
||||
//
|
||||
// 现在只在 ListenHost 明确写成 loopback 字面量时才开启 loopback 严格化:
|
||||
// - 显式 loopback IP(127.0.0.1 / ::1)或 localhost → 视为 loopback 部署,
|
||||
// 此时任何带 Origin 的请求 Host 不是 loopback 都拒掉(DNS rebinding 防线);
|
||||
// - 空 / `0.0.0.0` / `::` / `*` / 公网 IP / hostname → 不当 loopback 部署,
|
||||
// 仅做 Origin == Host 的同源校验,不再额外要求 Host 必须是 loopback。
|
||||
//
|
||||
// 调用方仍然在配额(Origin == Host)外用 isLoopbackHostname 校验 Host,所以
|
||||
// 当本机用户拿 127.0.0.1 访问绑公网 IP 的 dashboard 时这条防线依然生效(Host
|
||||
// 是 loopback,Origin 是公网域名,hostnameOnly 比较会失败)。
|
||||
func dashboardListensOnLoopback() bool {
|
||||
if singleton.Conf == nil {
|
||||
return false
|
||||
}
|
||||
host := strings.TrimSpace(singleton.Conf.ListenHost)
|
||||
if host == "" {
|
||||
return false
|
||||
}
|
||||
host = strings.Trim(host, "[]")
|
||||
switch host {
|
||||
case "0.0.0.0", "::", "*":
|
||||
return false
|
||||
}
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return ip.IsLoopback()
|
||||
}
|
||||
return strings.EqualFold(host, "localhost")
|
||||
}
|
||||
|
||||
func hostnameOnly(hostport string) string {
|
||||
if i := strings.LastIndex(hostport, ":"); i >= 0 {
|
||||
if !strings.Contains(hostport[i:], "]") {
|
||||
return hostport[:i]
|
||||
}
|
||||
}
|
||||
return hostport
|
||||
}
|
||||
|
||||
func isLoopbackHostname(host string) bool {
|
||||
host = strings.Trim(host, "[]")
|
||||
if host == "" {
|
||||
return false
|
||||
}
|
||||
if strings.EqualFold(host, "localhost") {
|
||||
return true
|
||||
}
|
||||
if ip := net.ParseIP(host); ip != nil {
|
||||
return ip.IsLoopback()
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func abortOrigin(c *gin.Context) {
|
||||
c.AbortWithStatusJSON(http.StatusForbidden, model.CommonResponse[any]{
|
||||
Success: false,
|
||||
Error: "ApiErrorForbidden: origin not allowed",
|
||||
})
|
||||
}
|
||||
|
||||
func mcpMethodNotAllowed(c *gin.Context) {
|
||||
c.Header("Allow", "POST")
|
||||
c.JSON(http.StatusMethodNotAllowed, model.CommonResponse[any]{
|
||||
Success: false,
|
||||
Error: "MCP endpoint only accepts POST (Streamable HTTP without standalone SSE / sessions)",
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
package controller
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/service/singleton"
|
||||
)
|
||||
|
||||
func setupMCPOriginRouter(t *testing.T) (*httptest.Server, string, func()) {
|
||||
t.Helper()
|
||||
cleanup, uid := setupMCPTest(t)
|
||||
_, plain := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
r.POST("/mcp", mcpOriginGuard(), apiTokenAuthMiddleware(), mcpEndpoint)
|
||||
r.GET("/mcp/download/:token", mcpOriginGuard(), transferDownloadHandler)
|
||||
r.POST("/mcp/upload/:token", mcpOriginGuard(), transferUploadHandler)
|
||||
ts := httptest.NewServer(r)
|
||||
return ts, plain, func() {
|
||||
ts.Close()
|
||||
cleanup()
|
||||
}
|
||||
}
|
||||
|
||||
func TestMCP_DisallowsCrossOriginRequest(t *testing.T) {
|
||||
ts, tok, cleanup := setupMCPOriginRouter(t)
|
||||
defer cleanup()
|
||||
|
||||
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader([]byte(`{"jsonrpc":"2.0","id":1,"method":"initialize"}`)))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+tok)
|
||||
req.Header.Set("Origin", "http://evil.example.com")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
require.Equal(t, http.StatusForbidden, resp.StatusCode)
|
||||
}
|
||||
|
||||
func TestMCP_AllowsRequestWithoutOriginHeader(t *testing.T) {
|
||||
ts, tok, cleanup := setupMCPOriginRouter(t)
|
||||
defer cleanup()
|
||||
|
||||
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader([]byte(`{"jsonrpc":"2.0","id":1,"method":"initialize"}`)))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+tok)
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
require.Equal(t, http.StatusOK, resp.StatusCode)
|
||||
}
|
||||
|
||||
func TestMCP_AllowsSameHostOrigin(t *testing.T) {
|
||||
ts, tok, cleanup := setupMCPOriginRouter(t)
|
||||
defer cleanup()
|
||||
|
||||
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader([]byte(`{"jsonrpc":"2.0","id":1,"method":"initialize"}`)))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+tok)
|
||||
req.Header.Set("Origin", "http://"+req.Host)
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
require.Equal(t, http.StatusOK, resp.StatusCode)
|
||||
}
|
||||
|
||||
// 公网部署回归:ListenHost 是 0.0.0.0/未指定时,前端会以公网 Host 同源访问。
|
||||
// 这条以前会被 dashboardListensOnLoopback 误判为 loopback 部署进而拒掉;
|
||||
// 现在必须放行,否则正常生产环境的 admin frontend MCP 入口直接 403。
|
||||
func TestMCP_PublicDeployment_AllowsPublicSameOrigin(t *testing.T) {
|
||||
ts, tok, cleanup := setupMCPOriginRouter(t)
|
||||
defer cleanup()
|
||||
|
||||
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader([]byte(`{"jsonrpc":"2.0","id":1,"method":"initialize"}`)))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+tok)
|
||||
req.Host = "dashboard.example.com"
|
||||
req.Header.Set("Origin", "https://dashboard.example.com")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
require.Equal(t, http.StatusOK, resp.StatusCode)
|
||||
}
|
||||
|
||||
// 显式绑 loopback 时仍然要执行 DNS rebinding 防线:Host 是公网域名 → 403。
|
||||
func TestMCP_LoopbackDeployment_RejectsPublicHost(t *testing.T) {
|
||||
ts, tok, cleanup := setupMCPOriginRouter(t)
|
||||
defer cleanup()
|
||||
prev := singleton.Conf.ListenHost
|
||||
singleton.Conf.ListenHost = "127.0.0.1"
|
||||
defer func() { singleton.Conf.ListenHost = prev }()
|
||||
|
||||
req, _ := http.NewRequest("POST", ts.URL+"/mcp", bytes.NewReader([]byte(`{"jsonrpc":"2.0","id":1,"method":"initialize"}`)))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+tok)
|
||||
req.Host = "dashboard.example.com"
|
||||
req.Header.Set("Origin", "https://dashboard.example.com")
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
require.Equal(t, http.StatusForbidden, resp.StatusCode)
|
||||
}
|
||||
@@ -0,0 +1,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 HTTP;GET SSE 通道未实现也不计划
|
||||
// 短期内上线(不需要 server→client 主动推送)。SDK 默认会试图发 GET,
|
||||
// 关掉 standalone SSE 即可严格互通。
|
||||
DisableStandaloneSSE: true,
|
||||
}
|
||||
}
|
||||
|
||||
func setupSDKCompat(t *testing.T) (string, string, func()) {
|
||||
t.Helper()
|
||||
cleanupBase, uid := setupMCPTest(t)
|
||||
|
||||
srv, _ := singleton.ServerShared.Get(7)
|
||||
srv.SetTaskStream(&e2eStream{dispatch: agentSim})
|
||||
|
||||
_, plain := mkToken(t, uid, []string{
|
||||
model.ScopeServerRead,
|
||||
model.ScopeServerWrite,
|
||||
model.ScopeServerDelete,
|
||||
model.ScopeServerExec,
|
||||
}, nil)
|
||||
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
r.POST("/mcp", apiTokenAuthMiddleware(), mcpEndpoint)
|
||||
ts := httptest.NewServer(r)
|
||||
return ts.URL + "/mcp", plain, func() {
|
||||
ts.Close()
|
||||
cleanupBase()
|
||||
}
|
||||
}
|
||||
|
||||
func TestSDKClient_InitializeHandshake(t *testing.T) {
|
||||
endpoint, token, cleanup := setupSDKCompat(t)
|
||||
defer cleanup()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
client := mcp.NewClient(&mcp.Implementation{Name: "nezha-it", Version: "v0"}, nil)
|
||||
session, err := client.Connect(ctx, sdkTransport(endpoint, token), nil)
|
||||
require.NoError(t, err, "official Go SDK must initialize against /mcp")
|
||||
defer session.Close()
|
||||
}
|
||||
|
||||
func TestSDKClient_ToolsList(t *testing.T) {
|
||||
endpoint, token, cleanup := setupSDKCompat(t)
|
||||
defer cleanup()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
client := mcp.NewClient(&mcp.Implementation{Name: "nezha-it", Version: "v0"}, nil)
|
||||
session, err := client.Connect(ctx, sdkTransport(endpoint, token), nil)
|
||||
require.NoError(t, err)
|
||||
defer session.Close()
|
||||
|
||||
lst, err := session.ListTools(ctx, nil)
|
||||
require.NoError(t, err)
|
||||
names := make(map[string]bool, len(lst.Tools))
|
||||
for _, tl := range lst.Tools {
|
||||
names[tl.Name] = true
|
||||
}
|
||||
for _, must := range []string{
|
||||
"meta.whoami",
|
||||
"server.list", "server.get", "server.exec",
|
||||
"fs.list", "fs.read", "fs.write", "fs.delete",
|
||||
"fs.download_url", "fs.upload_url",
|
||||
} {
|
||||
require.Truef(t, names[must], "tools/list missing %q", must)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSDKClient_Whoami(t *testing.T) {
|
||||
endpoint, token, cleanup := setupSDKCompat(t)
|
||||
defer cleanup()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
client := mcp.NewClient(&mcp.Implementation{Name: "nezha-it", Version: "v0"}, nil)
|
||||
session, err := client.Connect(ctx, sdkTransport(endpoint, token), nil)
|
||||
require.NoError(t, err)
|
||||
defer session.Close()
|
||||
|
||||
res, err := session.CallTool(ctx, &mcp.CallToolParams{
|
||||
Name: "meta.whoami",
|
||||
Arguments: map[string]any{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, res.IsError)
|
||||
tc, ok := res.Content[0].(*mcp.TextContent)
|
||||
require.True(t, ok)
|
||||
var payload map[string]any
|
||||
require.NoError(t, json.Unmarshal([]byte(tc.Text), &payload))
|
||||
require.NotZero(t, payload["user_id"])
|
||||
require.NotEmpty(t, payload["scopes"])
|
||||
}
|
||||
|
||||
func TestSDKClient_ServerExec(t *testing.T) {
|
||||
endpoint, token, cleanup := setupSDKCompat(t)
|
||||
defer cleanup()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
client := mcp.NewClient(&mcp.Implementation{Name: "nezha-it", Version: "v0"}, nil)
|
||||
session, err := client.Connect(ctx, sdkTransport(endpoint, token), nil)
|
||||
require.NoError(t, err)
|
||||
defer session.Close()
|
||||
|
||||
res, err := session.CallTool(ctx, &mcp.CallToolParams{
|
||||
Name: "server.exec",
|
||||
Arguments: map[string]any{
|
||||
"server_id": 7,
|
||||
"cmd": "echo",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, res.IsError, "exec failed: %v", res.Content)
|
||||
tc := res.Content[0].(*mcp.TextContent)
|
||||
require.Contains(t, tc.Text, "simulated")
|
||||
}
|
||||
|
||||
func TestSDKClient_FSLifecycle(t *testing.T) {
|
||||
endpoint, token, cleanup := setupSDKCompat(t)
|
||||
defer cleanup()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
client := mcp.NewClient(&mcp.Implementation{Name: "nezha-it", Version: "v0"}, nil)
|
||||
session, err := client.Connect(ctx, sdkTransport(endpoint, token), nil)
|
||||
require.NoError(t, err)
|
||||
defer session.Close()
|
||||
|
||||
path := t.TempDir() + "/sdk.txt"
|
||||
for _, step := range []struct {
|
||||
name string
|
||||
args map[string]any
|
||||
}{
|
||||
{"fs.write", map[string]any{"server_id": 7, "path": path, "content": "via-sdk", "encoding": "utf8"}},
|
||||
{"fs.read", map[string]any{"server_id": 7, "path": path}},
|
||||
{"fs.delete", map[string]any{"server_id": 7, "path": path}},
|
||||
} {
|
||||
res, err := session.CallTool(ctx, &mcp.CallToolParams{Name: step.name, Arguments: step.args})
|
||||
require.NoError(t, err, step.name)
|
||||
require.False(t, res.IsError, "%s failed: %v", step.name, res.Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSDKClient_BadPAT(t *testing.T) {
|
||||
endpoint, _, cleanup := setupSDKCompat(t)
|
||||
defer cleanup()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
client := mcp.NewClient(&mcp.Implementation{Name: "nezha-it", Version: "v0"}, nil)
|
||||
_, err := client.Connect(ctx, sdkTransport(endpoint, "nzp_invalid"), nil)
|
||||
require.Error(t, err)
|
||||
}
|
||||
@@ -0,0 +1,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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
// 闸 2:PAT 的 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 —— 协议帧和文件字节
|
||||
// 共用同一条 IOStream,HTTP 客户端不应收到“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 = ©Req
|
||||
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/"):]
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)")
|
||||
}
|
||||
}
|
||||
@@ -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 集合,再减去
|
||||
// servers(deny-list)。代表 CronCoverAll / ServiceCoverAll。受限 PAT
|
||||
// 必须确保 deny-list 已覆盖白名单外的全部 owner servers,否则 fan-out
|
||||
// 会跑到 PAT 白名单之外。
|
||||
coverModeAllMinusDeny
|
||||
|
||||
// coverModeAllowList: dispatch 时只在 servers(allow-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")
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -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 routes(controller.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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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.
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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 并塞两个用户,
|
||||
// 用户 10(member)和用户 999(foreign 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")
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user