fix(mcp): harden dashboard dispatch lifecycle

Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
naiba
2026-07-20 04:32:16 +00:00
co-authored by naiba/CloudCode
parent 2640b86d3a
commit 3f6de265a0
48 changed files with 1332 additions and 265 deletions
@@ -33,6 +33,8 @@ func setupAlertRuleFanoutFixture(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err) require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.Server{}, &model.AlertRule{}, &model.Cron{}, &model.User{})) require.NoError(t, db.AutoMigrate(&model.Server{}, &model.AlertRule{}, &model.Cron{}, &model.User{}))
singleton.DB = db singleton.DB = db
@@ -50,6 +52,8 @@ func setupAlertRuleFanoutFixture(t *testing.T) {
singleton.CronShared = singleton.NewCronClass() singleton.CronShared = singleton.NewCronClass()
t.Cleanup(func() { t.Cleanup(func() {
singleton.CronShared.Close()
_ = sqlDB.Close()
singleton.DB = originalDB singleton.DB = originalDB
singleton.Cache = originalCache singleton.Cache = originalCache
singleton.Loc = originalLoc singleton.Loc = originalLoc
+18 -8
View File
@@ -17,9 +17,10 @@ import (
) )
const ( const (
apiTokenSecretLength = 32 // 明文 token 随机部分长度(hex 编码前) apiTokenSecretLength = 32 // 明文 token 随机部分长度(hex 编码前)
apiTokenCtxKey = "nz_api_token" // #nosec G101 -- gin context key name, not a credential apiTokenCtxKey = "nz_api_token" // #nosec G101 -- gin context key name, not a credential
apiTokenLastUsedCtxKey = "nz_api_token_used_marker" // #nosec G101 -- gin context key name, not a credential apiTokenLastUsedCtxKey = "nz_api_token_used_marker" // #nosec G101 -- gin context key name, not a credential
apiTokenReadOnlyCtxKey = "nz_api_token_read_only" // #nosec G101 -- gin context key name, not a credential
apiTokenAuthSchemePrefix = "Bearer " apiTokenAuthSchemePrefix = "Bearer "
) )
@@ -204,6 +205,7 @@ func deleteAPIToken(c *gin.Context) (any, error) {
// 命中但 token 无效:直接 401 并 abort,不再走到 JWT。 // 命中但 token 无效:直接 401 并 abort,不再走到 JWT。
func apiTokenAuthMiddleware() gin.HandlerFunc { func apiTokenAuthMiddleware() gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
readOnly := apiTokenAuthReadOnly(c)
raw := strings.TrimSpace(c.GetHeader("Authorization")) raw := strings.TrimSpace(c.GetHeader("Authorization"))
if raw == "" { if raw == "" {
return return
@@ -223,7 +225,9 @@ func apiTokenAuthMiddleware() gin.HandlerFunc {
err := singleton.DB.Where("token_hash = ?", model.HashAPIToken(plaintext)).First(&tok).Error err := singleton.DB.Where("token_hash = ?", model.HashAPIToken(plaintext)).First(&tok).Error
if err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken) if !readOnly {
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken)
}
abortAPITokenUnauthorized(c, "invalid api token") abortAPITokenUnauthorized(c, "invalid api token")
return return
} }
@@ -232,19 +236,25 @@ func apiTokenAuthMiddleware() gin.HandlerFunc {
} }
now := time.Now() now := time.Now()
if tok.IsExpired(now) { if tok.IsExpired(now) {
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken) if !readOnly {
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken)
}
abortAPITokenUnauthorized(c, "api token expired") abortAPITokenUnauthorized(c, "api token expired")
return return
} }
var user model.User var user model.User
if err := singleton.DB.First(&user, tok.UserID).Error; err != nil { if err := singleton.DB.First(&user, tok.UserID).Error; err != nil {
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken) if !readOnly {
model.BlockIP(singleton.DB, realIP, model.WAFBlockReasonTypeBruteForceToken, model.BlockIDToken)
}
abortAPITokenUnauthorized(c, "owner of api token not found") abortAPITokenUnauthorized(c, "owner of api token not found")
return return
} }
model.UnblockIP(singleton.DB, realIP, model.BlockIDToken) if !readOnly {
model.UnblockIP(singleton.DB, realIP, model.BlockIDToken)
}
c.Set(model.CtxKeyAuthorizedUser, &user) c.Set(model.CtxKeyAuthorizedUser, &user)
c.Set(apiTokenCtxKey, &tok) c.Set(apiTokenCtxKey, &tok)
@@ -253,7 +263,7 @@ func apiTokenAuthMiddleware() gin.HandlerFunc {
// last_used 同步更新:开销极低(一行 UPDATE),异步路径在 // last_used 同步更新:开销极低(一行 UPDATE),异步路径在
// 多连接 sqlite 测试场景下会和测试 teardown 形成竞态,并把 // 多连接 sqlite 测试场景下会和测试 teardown 形成竞态,并把
// `last_used_*` 写丢到不可见的 :memory: 实例。生产路径上等价。 // `last_used_*` 写丢到不可见的 :memory: 实例。生产路径上等价。
if v, ok := c.Get(apiTokenLastUsedCtxKey); !ok || v != true { if !readOnly && apiTokenLastUsedOnce(c) {
c.Set(apiTokenLastUsedCtxKey, true) c.Set(apiTokenLastUsedCtxKey, true)
_ = singleton.DB.Model(&model.APIToken{}). _ = singleton.DB.Model(&model.APIToken{}).
Where("id = ?", tok.ID). Where("id = ?", tok.ID).
@@ -0,0 +1,19 @@
package controller
import "github.com/gin-gonic/gin"
func suppressAPITokenAuthWrites(c *gin.Context) { c.Set(apiTokenReadOnlyCtxKey, true) }
func apiTokenAuthReadOnly(c *gin.Context) bool {
value, ok := c.Get(apiTokenReadOnlyCtxKey)
return ok && value == true
}
func apiTokenLastUsedOnce(c *gin.Context) bool {
value, seen := c.Get(apiTokenLastUsedCtxKey)
if seen && value == true {
return false
}
c.Set(apiTokenLastUsedCtxKey, true)
return true
}
@@ -15,9 +15,9 @@ import (
// restScopeMiddleware 的"空 scope"实际行为:PAT 调用方一律 403。 // restScopeMiddleware 的"空 scope"实际行为:PAT 调用方一律 403。
// 这条测试把注释与实现的契约对齐: // 这条测试把注释与实现的契约对齐:
// 1. 实际行为:PAT + scope="" → 403。 // 1. 实际行为:PAT + scope="" → 403。
// 2. 文档约束:源码注释必须明确说出"空 scope 对 PAT 仍被拒绝" // 2. 文档约束:源码注释必须明确说出"空 scope 对 PAT 仍被拒绝"
// 不能再保留"空字符串 = 放行"这种与实现相反的旧措辞。 // 不能再保留"空字符串 = 放行"这种与实现相反的旧措辞。
func TestRestScopeMiddleware_EmptyScopeRejectsPAT(t *testing.T) { func TestRestScopeMiddleware_EmptyScopeRejectsPAT(t *testing.T) {
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
r := gin.New() r := gin.New()
@@ -23,9 +23,12 @@ func setupAPITokenTest(t *testing.T) func() {
originalDB := singleton.DB originalDB := singleton.DB
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err) require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.User{}, &model.APIToken{}, &model.Server{})) require.NoError(t, db.AutoMigrate(&model.User{}, &model.APIToken{}, &model.Server{}))
singleton.DB = db singleton.DB = db
return func() { return func() {
_ = sqlDB.Close()
singleton.DB = originalDB singleton.DB = originalDB
} }
} }
@@ -27,9 +27,9 @@ func callBatchMoveWithPAT(t *testing.T, callerID uint64, role model.Role, tok *m
r.ServeHTTP(w, req) r.ServeHTTP(w, req)
var resp struct { var resp struct {
Success bool `json:"success"` Success bool `json:"success"`
Error string `json:"error"` Error string `json:"error"`
Data []model.BatchMoveServerResult `json:"data"` Data []model.BatchMoveServerResult `json:"data"`
} }
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
return resp.Data, resp.Success, resp.Error return resp.Data, resp.Success, resp.Error
@@ -46,13 +46,14 @@ func setupCronDispatchPATFixture(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err) require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{}, &model.User{}, &model.NotificationGroup{}, &model.Notification{})) require.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{}, &model.User{}, &model.NotificationGroup{}, &model.Notification{}))
singleton.DB = db singleton.DB = db
singleton.Loc = time.UTC singleton.Loc = time.UTC
singleton.Cache = cache.New(time.Minute, time.Minute) singleton.Cache = cache.New(time.Minute, time.Minute)
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations) singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
singleton.CronShared = singleton.NewCronClass()
// CronTrigger 的 fan-out 路径会在 server 没接入 task-stream 时调 // CronTrigger 的 fan-out 路径会在 server 没接入 task-stream 时调
// NotificationShared.SendNotification 上报「离线」;这里给出一个空的 // NotificationShared.SendNotification 上报「离线」;这里给出一个空的
// notification class,避免 nil deref。本测试不验证通知内容。 // notification class,避免 nil deref。本测试不验证通知内容。
@@ -66,12 +67,16 @@ func setupCronDispatchPATFixture(t *testing.T) {
sc.InsertForTest(s) sc.InsertForTest(s)
} }
singleton.ServerShared = sc singleton.ServerShared = sc
singleton.CronShared = singleton.NewCronClass()
singleton.UserLock.Lock() singleton.UserLock.Lock()
singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}} singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}}
singleton.UserLock.Unlock() singleton.UserLock.Unlock()
t.Cleanup(func() { t.Cleanup(func() {
// Test-owned cron jobs must be joined before restoring process-global singleton dependencies.
singleton.CronShared.Close()
_ = sqlDB.Close()
singleton.DB = originalDB singleton.DB = originalDB
singleton.Cache = originalCache singleton.Cache = originalCache
singleton.Loc = originalLoc singleton.Loc = originalLoc
@@ -30,14 +30,16 @@ func setupCronManualTriggerFixture(t *testing.T) {
originalUserInfo := singleton.UserInfoMap originalUserInfo := singleton.UserInfoMap
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err) require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{}, &model.User{})) require.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{}, &model.User{}))
singleton.DB = db singleton.DB = db
singleton.Loc = time.UTC singleton.Loc = time.UTC
singleton.Cache = cache.New(time.Minute, time.Minute) singleton.Cache = cache.New(time.Minute, time.Minute)
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations) singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
singleton.CronShared = singleton.NewCronClass()
singleton.ServerShared = singleton.NewServerClass() singleton.ServerShared = singleton.NewServerClass()
singleton.CronShared = singleton.NewCronClass()
singleton.UserLock.Lock() singleton.UserLock.Lock()
singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}} singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}}
singleton.UserLock.Unlock() singleton.UserLock.Unlock()
@@ -53,6 +55,8 @@ func setupCronManualTriggerFixture(t *testing.T) {
singleton.CronShared.Update(cr) singleton.CronShared.Update(cr)
t.Cleanup(func() { t.Cleanup(func() {
singleton.CronShared.Close()
_ = sqlDB.Close()
singleton.DB = originalDB singleton.DB = originalDB
singleton.Cache = originalCache singleton.Cache = originalCache
singleton.Loc = originalLoc singleton.Loc = originalLoc
@@ -33,14 +33,16 @@ func setupCronPATWhitelistFixture(t *testing.T) (cronID7, cronID8 uint64) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err) require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{}, &model.User{})) require.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{}, &model.User{}))
singleton.DB = db singleton.DB = db
singleton.Loc = time.UTC singleton.Loc = time.UTC
singleton.Cache = cache.New(time.Minute, time.Minute) singleton.Cache = cache.New(time.Minute, time.Minute)
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations) singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
singleton.CronShared = singleton.NewCronClass()
singleton.ServerShared = singleton.NewServerClass() singleton.ServerShared = singleton.NewServerClass()
singleton.CronShared = singleton.NewCronClass()
singleton.UserLock.Lock() singleton.UserLock.Lock()
singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}} singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}}
singleton.UserLock.Unlock() singleton.UserLock.Unlock()
@@ -67,6 +69,8 @@ func setupCronPATWhitelistFixture(t *testing.T) (cronID7, cronID8 uint64) {
singleton.CronShared.Update(cr8) singleton.CronShared.Update(cr8)
t.Cleanup(func() { t.Cleanup(func() {
singleton.CronShared.Close()
_ = sqlDB.Close()
singleton.DB = originalDB singleton.DB = originalDB
singleton.Cache = originalCache singleton.Cache = originalCache
singleton.Loc = originalLoc singleton.Loc = originalLoc
@@ -50,27 +50,28 @@ func setupCoverPATFixture(t *testing.T) {
originalCron := singleton.CronShared originalCron := singleton.CronShared
originalServer := singleton.ServerShared originalServer := singleton.ServerShared
originalUserInfo := singleton.UserInfoMap originalUserInfo := singleton.UserInfoMap
originalNotification := singleton.NotificationShared
originalSentinel := singleton.ServiceSentinelShared
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err) require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{}, &model.User{}, &model.Service{}, &model.NotificationGroup{}, &model.ServiceHistory{})) require.NoError(t, db.AutoMigrate(&model.Cron{}, &model.Server{}, &model.User{}, &model.Service{}, &model.NotificationGroup{}, &model.ServiceHistory{}))
singleton.DB = db singleton.DB = db
singleton.Loc = time.UTC singleton.Loc = time.UTC
singleton.Cache = cache.New(time.Minute, time.Minute) singleton.Cache = cache.New(time.Minute, time.Minute)
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations) singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
singleton.NotificationShared = singleton.NewEmptyNotificationClassForTest()
sc := singleton.NewEmptyServerClassForTest()
singleton.ServerShared = sc
singleton.CronShared = singleton.NewCronClass() singleton.CronShared = singleton.NewCronClass()
originalSentinel := singleton.ServiceSentinelShared
sentinel, err := singleton.NewServiceSentinel(make(chan *model.Service, 4)) sentinel, err := singleton.NewServiceSentinel(make(chan *model.Service, 4))
require.NoError(t, err) require.NoError(t, err)
singleton.ServiceSentinelShared = sentinel singleton.ServiceSentinelShared = sentinel
t.Cleanup(func() {
sentinel.Close()
singleton.ServiceSentinelShared = originalSentinel
})
sc := singleton.NewEmptyServerClassForTest()
for _, id := range []uint64{1, 2} { for _, id := range []uint64{1, 2} {
s := &model.Server{} s := &model.Server{}
s.ID = id s.ID = id
@@ -84,12 +85,18 @@ func setupCoverPATFixture(t *testing.T) {
singleton.UserLock.Unlock() singleton.UserLock.Unlock()
t.Cleanup(func() { t.Cleanup(func() {
// Background components must be joined before restoring process globals.
sentinel.Close()
singleton.CronShared.Close()
_ = sqlDB.Close()
singleton.DB = originalDB singleton.DB = originalDB
singleton.Cache = originalCache singleton.Cache = originalCache
singleton.Loc = originalLoc singleton.Loc = originalLoc
singleton.Localizer = originalLocalizer singleton.Localizer = originalLocalizer
singleton.CronShared = originalCron singleton.CronShared = originalCron
singleton.ServerShared = originalServer singleton.ServerShared = originalServer
singleton.NotificationShared = originalNotification
singleton.ServiceSentinelShared = originalSentinel
singleton.UserLock.Lock() singleton.UserLock.Lock()
singleton.UserInfoMap = originalUserInfo singleton.UserInfoMap = originalUserInfo
singleton.UserLock.Unlock() singleton.UserLock.Unlock()
@@ -109,7 +109,7 @@ func TestForceUpdateServerOnlineForeignIDIndistinguishableFromUnknown(t *testing
defer reset() defer reset()
const bobID = uint64(200) const bobID = uint64(200)
foreignResp := decodeForceUpdate(t, runForceUpdate(t, bobID, []uint64{1})) // alice's online foreignResp := decodeForceUpdate(t, runForceUpdate(t, bobID, []uint64{1})) // alice's online
unknownResp := decodeForceUpdate(t, runForceUpdate(t, bobID, []uint64{9999})) // does not exist unknownResp := decodeForceUpdate(t, runForceUpdate(t, bobID, []uint64{9999})) // does not exist
assert.Equal(t, foreignResp.Success, unknownResp.Success, assert.Equal(t, foreignResp.Success, unknownResp.Success,
+4 -4
View File
@@ -64,10 +64,10 @@ func issueJWTSession(c *gin.Context, user *model.User, jwtTimeoutHours int) (map
func initParams() *jwt.GinJWTMiddleware { func initParams() *jwt.GinJWTMiddleware {
return &jwt.GinJWTMiddleware{ return &jwt.GinJWTMiddleware{
Realm: singleton.Conf.SiteName, Realm: singleton.Conf.SiteName,
Key: []byte(singleton.Conf.JWTSecretKey), Key: []byte(singleton.Conf.JWTSecretKey),
CookieName: "nz-jwt", CookieName: "nz-jwt",
SendCookie: true, SendCookie: true,
// Pin the signing algorithm so a future library default change (or an // Pin the signing algorithm so a future library default change (or an
// `alg: none` confusion attempt) cannot weaken token validation. // `alg: none` confusion attempt) cannot weaken token validation.
SigningAlgorithm: "HS256", SigningAlgorithm: "HS256",
@@ -30,6 +30,8 @@ func setupJWTSessionTest(t *testing.T) (cleanup func()) {
originalConf := singleton.Conf originalConf := singleton.Conf
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err) require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.User{}, &model.JWTSession{}, &model.WAF{})) require.NoError(t, db.AutoMigrate(&model.User{}, &model.JWTSession{}, &model.WAF{}))
singleton.DB = db singleton.DB = db
singleton.Conf = &singleton.ConfigClass{Config: &model.Config{JWTTimeout: 1}} singleton.Conf = &singleton.ConfigClass{Config: &model.Config{JWTTimeout: 1}}
@@ -42,6 +44,7 @@ func setupJWTSessionTest(t *testing.T) (cleanup func()) {
}).Error) }).Error)
return func() { return func() {
_ = sqlDB.Close()
singleton.DB = originalDB singleton.DB = originalDB
singleton.Conf = originalConf singleton.Conf = originalConf
} }
+5 -1
View File
@@ -135,6 +135,8 @@ func newValidationContext(t *testing.T, userID uint64, role model.Role) *gin.Con
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
assert.NoError(t, err) assert.NoError(t, err)
sqlDB, err := db.DB()
assert.NoError(t, err)
assert.NoError(t, db.AutoMigrate( assert.NoError(t, db.AutoMigrate(
&model.Cron{}, &model.Cron{},
&model.Server{}, &model.Server{},
@@ -169,8 +171,8 @@ func newValidationContext(t *testing.T, userID uint64, role model.Role) *gin.Con
singleton.DB = db singleton.DB = db
singleton.Loc = time.Local singleton.Loc = time.Local
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations) singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations)
singleton.CronShared = singleton.NewCronClass()
singleton.ServerShared = singleton.NewServerClass() singleton.ServerShared = singleton.NewServerClass()
singleton.CronShared = singleton.NewCronClass()
singleton.UserLock.Lock() singleton.UserLock.Lock()
singleton.UserInfoMap = map[uint64]model.UserInfo{ singleton.UserInfoMap = map[uint64]model.UserInfo{
1: {Role: model.RoleAdmin}, 1: {Role: model.RoleAdmin},
@@ -178,6 +180,8 @@ func newValidationContext(t *testing.T, userID uint64, role model.Role) *gin.Con
} }
singleton.UserLock.Unlock() singleton.UserLock.Unlock()
t.Cleanup(func() { t.Cleanup(func() {
singleton.CronShared.Close()
_ = sqlDB.Close()
singleton.DB = originalDB singleton.DB = originalDB
singleton.Loc = originalLoc singleton.Loc = originalLoc
singleton.Localizer = originalLocalizer singleton.Localizer = originalLocalizer
+25 -15
View File
@@ -307,6 +307,20 @@ func handleToolsCall(c *gin.Context, req *jsonRPCRequest, tok *model.APIToken) {
} }
finish := func(outcome, errCode, errMsg string, result any) { finish := func(outcome, errCode, errMsg string, result any) {
if outcome == model.MCPOutcomeOK {
textPayload, err := marshalMCPToolResult(result)
if err != nil {
outcome = model.MCPOutcomeAgentError
errCode = model.MCPOutcomeAgentError
errMsg = "failed to encode tool result: " + err.Error()
result = nil
} else {
writeJSONRPCResult(c, req.ID, mcpToolCallResult{
Content: []mcpContent{{Type: "text", Text: textPayload}},
StructuredContent: result,
})
}
}
audit.Outcome = outcome audit.Outcome = outcome
audit.ErrorCode = errCode audit.ErrorCode = errCode
audit.ErrorMsg = truncateString(errMsg, 512) audit.ErrorMsg = truncateString(errMsg, 512)
@@ -315,24 +329,15 @@ func handleToolsCall(c *gin.Context, req *jsonRPCRequest, tok *model.APIToken) {
mcpAuditWrite(audit, p.Arguments) mcpAuditWrite(audit, p.Arguments)
if outcome == model.MCPOutcomeOK { 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 return
} }
// 错误结果不带 structuredContent:严格客户端会拿它去校验工具声明的 // Semantic tool failures may retain a typed structured result (notably
// outputSchema(要求 exit_code/stdout/... 等成功字段),缺字段就整条 // server.exec) so clients can distinguish a non-zero command outcome from
// 响应报 -32602,把真正的 isError 文本掩盖掉。错误信息走 content[].text。 // transport, authorization, or deadline failures.
writeJSONRPCResult(c, req.ID, mcpToolCallResult{ writeJSONRPCResult(c, req.ID, mcpToolCallResult{
Content: []mcpContent{{Type: "text", Text: errMsg}}, Content: []mcpContent{{Type: "text", Text: errMsg}},
IsError: true, StructuredContent: result,
IsError: true,
}) })
} }
@@ -356,6 +361,11 @@ func handleToolsCall(c *gin.Context, req *jsonRPCRequest, tok *model.APIToken) {
result, err := tool.Handler(c, p.Arguments) result, err := tool.Handler(c, p.Arguments)
if err != nil { if err != nil {
code, msg := classifyToolError(err) code, msg := classifyToolError(err)
var structuredErr interface{ StructuredResult() any }
if errors.As(err, &structuredErr) {
finish(code, code, msg, structuredErr.StructuredResult())
return
}
finish(code, code, msg, nil) finish(code, code, msg, nil)
return return
} }
+6 -2
View File
@@ -18,10 +18,14 @@ const MCPMinAgentVersion = "v2.1.0"
// requireAgentSupportsMCP 在 tool handler 调 CallAgent 之前快速失败不支持的 agent。 // requireAgentSupportsMCP 在 tool handler 调 CallAgent 之前快速失败不支持的 agent。
// 仅作为 UX 优化:真正的安全/正确性由 agent 端 task switch 的 default 分支保障。 // 仅作为 UX 优化:真正的安全/正确性由 agent 端 task switch 的 default 分支保障。
func requireAgentSupportsMCP(server *model.Server) error { func requireAgentSupportsMCP(server *model.Server) error {
if MCPMinAgentVersion == "" || server == nil || server.Host == nil { if MCPMinAgentVersion == "" || server == nil {
return nil return nil
} }
if compareSemver(server.Host.Version, MCPMinAgentVersion) < 0 { runtime := server.RuntimeSnapshot()
if runtime.Host == nil {
return nil
}
if compareSemver(runtime.Host.Version, MCPMinAgentVersion) < 0 {
return errMCPUnsupported return errMCPUnsupported
} }
return nil return nil
@@ -10,10 +10,7 @@ import (
"github.com/nezhahq/nezha/service/singleton" "github.com/nezhahq/nezha/service/singleton"
) )
// A tool error must not ship structuredContent: strict clients validate it func TestServerExec_ErrorResult_PreservesStructuredContent(t *testing.T) {
// against the tool's outputSchema (which requires exec/fs result fields) and
// reject the whole response with -32602, masking the real isError text.
func TestServerExec_ErrorResult_OmitsStructuredContent(t *testing.T) {
cleanup, uid := setupMCPTest(t) cleanup, uid := setupMCPTest(t)
defer cleanup() defer cleanup()
@@ -40,8 +37,12 @@ func TestServerExec_ErrorResult_OmitsStructuredContent(t *testing.T) {
require.True(t, tcr.IsError) require.True(t, tcr.IsError)
require.Contains(t, tcr.Content[0].Text, "agent disabled command execution", require.Contains(t, tcr.Content[0].Text, "agent disabled command execution",
"error text must carry the real cause") "error text must carry the real cause")
require.Nil(t, tcr.StructuredContent, var result model.ExecResult
"error responses must omit structuredContent so strict clients don't validate it against outputSchema and mask the real error") structuredJSON, err := json.Marshal(tcr.StructuredContent)
require.NoError(t, err)
require.NoError(t, json.Unmarshal(structuredJSON, &result))
require.Equal(t, -1, result.ExitCode)
require.Equal(t, "agent disabled command execution", result.Error)
} }
func TestScopeDenied_OmitsStructuredContent(t *testing.T) { func TestScopeDenied_OmitsStructuredContent(t *testing.T) {
@@ -62,6 +63,5 @@ func TestScopeDenied_OmitsStructuredContent(t *testing.T) {
_, tcr := decodeRPC(w) _, tcr := decodeRPC(w)
require.NotNil(t, tcr) require.NotNil(t, tcr)
require.True(t, tcr.IsError) require.True(t, tcr.IsError)
require.Nil(t, tcr.StructuredContent, require.Nil(t, tcr.StructuredContent)
"scope-denied error must also omit structuredContent")
} }
+9 -1
View File
@@ -17,6 +17,7 @@ type MCPRateLimiter struct {
secLimit int secLimit int
minLimit int minLimit int
lastPrune time.Time lastPrune time.Time
clock func() time.Time
} }
type tokenWindow struct { type tokenWindow struct {
@@ -34,10 +35,15 @@ type tokenWindow struct {
const mcpRateLimiterPruneInterval = time.Minute const mcpRateLimiterPruneInterval = time.Minute
func newMCPRateLimiter(secLimit, minLimit int) *MCPRateLimiter { func newMCPRateLimiter(secLimit, minLimit int) *MCPRateLimiter {
return newMCPRateLimiterWithClock(secLimit, minLimit, time.Now)
}
func newMCPRateLimiterWithClock(secLimit, minLimit int, clock func() time.Time) *MCPRateLimiter {
return &MCPRateLimiter{ return &MCPRateLimiter{
perToken: make(map[uint64]*tokenWindow), perToken: make(map[uint64]*tokenWindow),
secLimit: secLimit, secLimit: secLimit,
minLimit: minLimit, minLimit: minLimit,
clock: clock,
} }
} }
@@ -58,9 +64,11 @@ func (r *MCPRateLimiter) Allow(tokenID uint64) bool {
if tokenID == 0 { if tokenID == 0 {
return true return true
} }
now := time.Now()
r.mu.Lock() r.mu.Lock()
defer r.mu.Unlock() defer r.mu.Unlock()
// Sampling while holding the state lock linearizes the clock read with the
// bucket update, so a delayed sample cannot commit after a newer rollover.
now := r.clock()
if now.Sub(r.lastPrune) >= mcpRateLimiterPruneInterval { if now.Sub(r.lastPrune) >= mcpRateLimiterPruneInterval {
r.pruneStaleLocked(now) r.pruneStaleLocked(now)
r.lastPrune = now r.lastPrune = now
@@ -0,0 +1,51 @@
//go:build agentcompat
package controller
import (
"time"
"github.com/gin-gonic/gin"
)
type agentcompatMCPRateLimitProbeRequest struct {
}
type agentcompatMCPRateLimitProbeResponse struct {
SecondAllowedCount int `json:"second_allowed_count"`
SecondRejectedAtCount int `json:"second_rejected_at_count"`
MinuteAllowedCount int `json:"minute_allowed_count"`
MinuteRejectedAtCount int `json:"minute_rejected_at_count"`
}
func agentcompatMCPRateLimitProbeRoute(context *gin.Context) (agentcompatMCPRateLimitProbeResponse, error) {
var request agentcompatMCPRateLimitProbeRequest
if err := decodeAgentcompatJSON(context, &request); err != nil {
return agentcompatMCPRateLimitProbeResponse{}, err
}
return runAgentcompatMCPRateLimitProbe(request)
}
func runAgentcompatMCPRateLimitProbe(request agentcompatMCPRateLimitProbeRequest) (agentcompatMCPRateLimitProbeResponse, error) {
response := agentcompatMCPRateLimitProbeResponse{}
secondNow := time.Unix(1_700_000_000, 0)
secondLimiter := newMCPRateLimiterWithClock(10, 120, func() time.Time { return secondNow })
for requestNumber := 1; requestNumber <= 11; requestNumber++ {
if secondLimiter.Allow(1) {
response.SecondAllowedCount++
} else if response.SecondRejectedAtCount == 0 {
response.SecondRejectedAtCount = requestNumber
}
}
minuteNow := time.Unix(1_700_000_000, 0)
minuteLimiter := newMCPRateLimiterWithClock(10_000, 120, func() time.Time { return minuteNow })
for requestNumber := 1; requestNumber <= 121; requestNumber++ {
if minuteLimiter.Allow(1) {
response.MinuteAllowedCount++
} else if response.MinuteRejectedAtCount == 0 {
response.MinuteRejectedAtCount = requestNumber
}
minuteNow = minuteNow.Add(100 * time.Millisecond)
}
return response, nil
}
@@ -0,0 +1,102 @@
//go:build agentcompat
package controller
import (
"bytes"
"encoding/json"
"net/http/httptest"
"sync"
"testing"
"github.com/gin-gonic/gin"
)
func TestAgentcompatMCPRateLimitProbe_returnsTypedBoundaryCountsWithoutSharedState(t *testing.T) {
// Given
request := agentcompatMCPRateLimitProbeRequest{}
originalLimiter := mcpRateLimiterShared
// When
response, err := runAgentcompatMCPRateLimitProbe(request)
// Then
if err != nil {
t.Fatalf("probe returned error: %v", err)
}
if response.SecondAllowedCount != 10 || response.SecondRejectedAtCount != 11 || response.MinuteAllowedCount != 120 || response.MinuteRejectedAtCount != 121 {
t.Fatalf("probe result = %+v, want second=10/11 minute=120/121", response)
}
t.Logf("typed probe result: second=%d/%d minute=%d/%d", response.SecondAllowedCount, response.SecondRejectedAtCount, response.MinuteAllowedCount, response.MinuteRejectedAtCount)
if mcpRateLimiterShared != originalLimiter {
t.Fatal("probe mutated the shared production limiter")
}
}
func TestAgentcompatMCPRateLimitProbe_isRepeatableAndConcurrentSafe(t *testing.T) {
// Given
request := agentcompatMCPRateLimitProbeRequest{}
results := make(chan agentcompatMCPRateLimitProbeResponse, 8)
errors := make(chan error, 8)
// When
var waitGroup sync.WaitGroup
for probeNumber := 0; probeNumber < 8; probeNumber++ {
waitGroup.Add(1)
go func() {
defer waitGroup.Done()
response, err := runAgentcompatMCPRateLimitProbe(request)
if err != nil {
errors <- err
return
}
results <- response
}()
}
waitGroup.Wait()
close(results)
close(errors)
// Then
for err := range errors {
t.Fatalf("concurrent probe returned error: %v", err)
}
for response := range results {
if response.SecondAllowedCount != 10 || response.SecondRejectedAtCount != 11 || response.MinuteAllowedCount != 120 || response.MinuteRejectedAtCount != 121 {
t.Fatalf("concurrent probe result = %+v, want second=10/11 minute=120/121", response)
}
}
}
func TestAgentcompatMCPRateLimitProbe_rejectsCallerControlledParameters(t *testing.T) {
// Given
gin.SetMode(gin.TestMode)
context, _ := gin.CreateTestContext(httptest.NewRecorder())
context.Request = httptest.NewRequest("POST", "/", bytes.NewBufferString(`{"token_id":7}`))
context.Request.Header.Set("Content-Type", "application/json")
// When
var request agentcompatMCPRateLimitProbeRequest
err := decodeAgentcompatJSON(context, &request)
// Then
if err == nil {
t.Fatal("caller-controlled rate probe parameters must be rejected")
}
context.Request = httptest.NewRequest("POST", "/", bytes.NewBufferString(`{}`))
if err := decodeAgentcompatJSON(context, &request); err != nil {
t.Fatalf("canonical empty probe request must be accepted: %v", err)
}
if _, err := json.Marshal(request); err != nil {
t.Fatalf("canonical request must remain JSON encodable: %v", err)
}
}
func TestAgentcompatMCPRateLimitProbeRejectsTrailingJSON(t *testing.T) {
context, _ := gin.CreateTestContext(httptest.NewRecorder())
context.Request = httptest.NewRequest("POST", "/", bytes.NewBufferString(`{}{}`))
var request agentcompatMCPRateLimitProbeRequest
if err := decodeAgentcompatJSON(context, &request); err == nil {
t.Fatal("trailing JSON values must be rejected")
}
}
@@ -0,0 +1,197 @@
package controller
import (
"sync"
"sync/atomic"
"testing"
"time"
)
type mcpRateLimitFakeClock struct {
now time.Time
}
func (clock *mcpRateLimitFakeClock) Now() time.Time {
return clock.now
}
func (clock *mcpRateLimitFakeClock) Advance(duration time.Duration) {
clock.now = clock.now.Add(duration)
}
func TestMCPRateLimiter_baselinePreservesAnonymousBypassAndBucketReset(t *testing.T) {
// Given
limiter := newMCPRateLimiter(3, 3)
// When
anonymousAllowed := limiter.Allow(0)
firstAllowed := limiter.Allow(7)
secondAllowed := limiter.Allow(7)
thirdAllowed := limiter.Allow(7)
fourthRejected := limiter.Allow(7)
// Then
if !anonymousAllowed || !firstAllowed || !secondAllowed || !thirdAllowed || fourthRejected {
t.Fatalf("baseline limiter semantics changed: anonymous=%t first=%t second=%t third=%t rejected=%t", anonymousAllowed, firstAllowed, secondAllowed, thirdAllowed, fourthRejected)
}
}
func TestMCPRateLimiter_allowsTenAndRejectsElevenWithinOneSecond(t *testing.T) {
// Given
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
limiter := newMCPRateLimiterWithClock(10, 120, clock.Now)
// When
allowed := 0
for requestNumber := 1; requestNumber <= 11; requestNumber++ {
if limiter.Allow(7) {
allowed++
}
}
// Then
if allowed != 10 {
t.Fatalf("allowed count = %d, want 10", allowed)
}
if limiter.Allow(7) {
t.Fatal("twelfth request must remain rejected in the same second")
}
}
func TestMCPRateLimiter_allowsOneAfterSecondBucketRollover(t *testing.T) {
// Given
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
limiter := newMCPRateLimiterWithClock(10, 120, clock.Now)
for requestNumber := 1; requestNumber <= 10; requestNumber++ {
if !limiter.Allow(7) {
t.Fatalf("request %d must be allowed before rollover", requestNumber)
}
}
// When
clock.Advance(time.Second)
// Then
if !limiter.Allow(7) {
t.Fatal("request after one-second bucket rollover must be allowed")
}
}
func TestMCPRateLimiter_allows120AndRejects121WithinOneMinute(t *testing.T) {
// Given
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
limiter := newMCPRateLimiterWithClock(10_000, 120, clock.Now)
// When
allowed := 0
for requestNumber := 1; requestNumber <= 121; requestNumber++ {
if limiter.Allow(7) {
allowed++
}
clock.Advance(100 * time.Millisecond)
}
// Then
if allowed != 120 {
t.Fatalf("allowed count = %d, want 120", allowed)
}
}
func TestMCPRateLimiter_independentTokensKeepIndependentBudgets(t *testing.T) {
// Given
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
limiter := newMCPRateLimiterWithClock(1, 120, clock.Now)
// When
firstTokenAllowed := limiter.Allow(7)
firstTokenRejected := limiter.Allow(7)
secondTokenAllowed := limiter.Allow(8)
// Then
if !firstTokenAllowed || firstTokenRejected || !secondTokenAllowed {
t.Fatalf("token budgets are not independent: first allowed=%t rejected=%t second allowed=%t", firstTokenAllowed, firstTokenRejected, secondTokenAllowed)
}
}
func TestMCPRateLimiter_clockControlledCallsDoNotMutateSharedLimiter(t *testing.T) {
// Given
originalLimiter := mcpRateLimiterShared
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
limiter := newMCPRateLimiterWithClock(1, 1, clock.Now)
// When
if !limiter.Allow(7) {
t.Fatal("isolated limiter request must be allowed")
}
// Then
if mcpRateLimiterShared != originalLimiter {
t.Fatal("isolated limiter changed shared production limiter ownership")
}
}
func TestMCPRateLimiter_concurrentCallsAtSecondBoundaryAllowExactlyLimit(t *testing.T) {
// Given
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
limiter := newMCPRateLimiterWithClock(10, 120, clock.Now)
var allowed atomic.Int32
var waitGroup sync.WaitGroup
// When
for requestNumber := 0; requestNumber < 20; requestNumber++ {
waitGroup.Add(1)
go func() {
defer waitGroup.Done()
if limiter.Allow(7) {
allowed.Add(1)
}
}()
}
waitGroup.Wait()
// Then
if got := allowed.Load(); got != 10 {
t.Fatalf("concurrent allowed count = %d, want 10", got)
}
}
func TestMCPRateLimiter_minuteBucketRolloverRestoresBudget(t *testing.T) {
// Given
clock := &mcpRateLimitFakeClock{now: time.Unix(1_700_000_000, 0)}
limiter := newMCPRateLimiterWithClock(10_000, 1, clock.Now)
if !limiter.Allow(7) {
t.Fatal("first request must be allowed")
}
if limiter.Allow(7) {
t.Fatal("second request must be rejected before minute rollover")
}
// When
clock.Advance(time.Minute)
// Then
if !limiter.Allow(7) {
t.Fatal("request after minute rollover must be allowed")
}
}
func TestMCPRateLimiter_samplesClockWhileHoldingStateLock(t *testing.T) {
// Given
now := time.Unix(1_700_000_000, 0)
var limiter *MCPRateLimiter
clock := func() time.Time {
if limiter.mu.TryLock() {
limiter.mu.Unlock()
t.Fatal("clock callback acquired limiter state lock; Allow sampled before locking")
}
return now
}
limiter = newMCPRateLimiterWithClock(1, 1, clock)
// When
allowed := limiter.Allow(7)
// Then
if !allowed {
t.Fatal("initial request must be allowed")
}
}
@@ -7,22 +7,39 @@ import (
"testing" "testing"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/model" "github.com/nezhahq/nezha/model"
"github.com/nezhahq/nezha/service/singleton" "github.com/nezhahq/nezha/service/singleton"
) )
func mcpEndpointTestCtx(t *testing.T, tok *model.APIToken, body any) (*gin.Context, *httptest.ResponseRecorder) { func mcpEndpointTestCtx(t *testing.T, tok *model.APIToken, body any) (*gin.Context, *httptest.ResponseRecorder, error) {
t.Helper() t.Helper()
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
w := httptest.NewRecorder() w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w) c, engine := gin.CreateTestContext(w)
raw, _ := json.Marshal(body) require.NotNil(t, engine)
raw, err := json.Marshal(body)
if err != nil {
return nil, nil, err
}
c.Request = httptest.NewRequest("POST", "/mcp", bytes.NewReader(raw)) c.Request = httptest.NewRequest("POST", "/mcp", bytes.NewReader(raw))
c.Request.Header.Set("Content-Type", "application/json") c.Request.Header.Set("Content-Type", "application/json")
c.Set(apiTokenCtxKey, tok) c.Set(apiTokenCtxKey, tok)
c.Set(model.CtxKeyAPIToken, tok) c.Set(model.CtxKeyAPIToken, tok)
return c, w return c, w, nil
}
func TestMCPEndpointTestCtx_surfacesJSONMarshalError(t *testing.T) {
// Given
tok := &model.APIToken{ID: 4242, UserID: 1}
unsupportedBody := map[string]any{"unsupported": func() {}}
// When
_, _, err := mcpEndpointTestCtx(t, tok, unsupportedBody)
// Then
require.Error(t, err)
} }
// An unknown tool name in tools/call must still consume the per-token rate // An unknown tool name in tools/call must still consume the per-token rate
@@ -53,7 +70,8 @@ func TestMCPUnknownToolCountsAgainstRateLimit(t *testing.T) {
var lastStatus int var lastStatus int
for i := 0; i < 5; i++ { for i := 0; i < 5; i++ {
c, w := mcpEndpointTestCtx(t, tok, body) c, w, err := mcpEndpointTestCtx(t, tok, body)
require.NoError(t, err)
mcpEndpoint(c) mcpEndpoint(c)
lastStatus = w.Code lastStatus = w.Code
} }
@@ -0,0 +1,45 @@
package controller
import (
"sync"
"testing"
"time"
)
type linearizationClock struct {
mu sync.Mutex
times []time.Time
latest time.Time
}
func newLinearizationClock(oldTime, newTime time.Time) *linearizationClock {
return &linearizationClock{
times: []time.Time{oldTime, oldTime, newTime},
latest: newTime,
}
}
func (clock *linearizationClock) Now() time.Time {
clock.mu.Lock()
now := clock.latest
if len(clock.times) > 0 {
now = clock.times[0]
clock.times = clock.times[1:]
}
clock.mu.Unlock()
return now
}
func TestLinearizationClockReturnsLatestTimeAfterSequenceIsConsumed(t *testing.T) {
oldTime := time.Unix(1_700_000_000, 0)
newTime := oldTime.Add(time.Second)
clock := newLinearizationClock(oldTime, newTime)
for range 3 {
clock.Now()
}
if got := clock.Now(); !got.Equal(newTime) {
t.Fatalf("clock returned %s after sequence was consumed, want %s", got, newTime)
}
}
@@ -6,6 +6,7 @@ import (
"testing" "testing"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/model" "github.com/nezhahq/nezha/model"
"github.com/nezhahq/nezha/service/singleton" "github.com/nezhahq/nezha/service/singleton"
@@ -15,7 +16,8 @@ func mcpEndpointRawCtx(t *testing.T, tok *model.APIToken, raw []byte) (*gin.Cont
t.Helper() t.Helper()
gin.SetMode(gin.TestMode) gin.SetMode(gin.TestMode)
w := httptest.NewRecorder() w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w) c, engine := gin.CreateTestContext(w)
require.NotNil(t, engine)
c.Request = httptest.NewRequest("POST", "/mcp", bytes.NewReader(raw)) c.Request = httptest.NewRequest("POST", "/mcp", bytes.NewReader(raw))
c.Request.Header.Set("Content-Type", "application/json") c.Request.Header.Set("Content-Type", "application/json")
c.Set(apiTokenCtxKey, tok) c.Set(apiTokenCtxKey, tok)
@@ -47,8 +49,9 @@ func TestMCPMalformedToolsCallParamsCountsAgainstRateLimit(t *testing.T) {
raw := []byte(`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":"not-an-object"}`) raw := []byte(`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":"not-an-object"}`)
for i := 0; i < 5; i++ { for i := 0; i < 5; i++ {
c, _ := mcpEndpointRawCtx(t, tok, raw) c, w := mcpEndpointRawCtx(t, tok, raw)
mcpEndpoint(c) mcpEndpoint(c)
require.NotEmpty(t, w.Body.Bytes())
} }
if mcpRateLimiterShared.Allow(tok.ID) { if mcpRateLimiterShared.Allow(tok.ID) {
@@ -64,8 +67,9 @@ func TestMCPMalformedEnvelopeCountsAgainstRateLimit(t *testing.T) {
raw := []byte(`{not valid json`) raw := []byte(`{not valid json`)
for i := 0; i < 5; i++ { for i := 0; i < 5; i++ {
c, _ := mcpEndpointRawCtx(t, tok, raw) c, w := mcpEndpointRawCtx(t, tok, raw)
mcpEndpoint(c) mcpEndpoint(c)
require.NotEmpty(t, w.Body.Bytes())
} }
if mcpRateLimiterShared.Allow(tok.ID) { if mcpRateLimiterShared.Allow(tok.ID) {
@@ -0,0 +1,14 @@
package controller
import "encoding/json"
func marshalMCPToolResult(result any) (string, error) {
if result == nil {
return "{}", nil
}
b, err := json.Marshal(result)
if err != nil {
return "", err
}
return string(b), nil
}
@@ -0,0 +1,47 @@
package controller
import (
"encoding/json"
"strings"
"sync"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/model"
)
var registerUnmarshalableMCPTool sync.Once
func TestMCPToolCall_unmarshalableSuccessResultReturnsExplicitToolError(t *testing.T) {
// Given
cleanup, uid := setupMCPTest(t)
defer cleanup()
registerUnmarshalableMCPTool.Do(func() {
registerMCPTool(&mcpTool{
Name: "test.unmarshalable-success",
RequiredScope: "",
Handler: func(*gin.Context, json.RawMessage) (any, error) {
return map[string]any{"unsupported": func() {}}, nil
},
})
})
tok, plainToken := mkToken(t, uid, []string{model.ScopeServerRead}, nil)
require.NotEmpty(t, plainToken)
c, w := mcpCallCtx(t, tok, uid, jsonRPCRequest{
JSONRPC: "2.0", ID: json.RawMessage("1"), Method: "tools/call",
Params: jsonObj(t, toolCallParams{Name: "test.unmarshalable-success", Arguments: json.RawMessage("{}")}),
})
// When
mcpEndpoint(c)
// Then
_, result := decodeRPC(w)
require.NotNil(t, result)
require.True(t, result.IsError)
require.Nil(t, result.StructuredContent)
require.Len(t, result.Content, 1)
require.True(t, strings.Contains(result.Content[0].Text, "encode"), result.Content[0].Text)
}
+3
View File
@@ -35,6 +35,8 @@ func setupMCPTest(t *testing.T) (func(), uint64) {
patConnectionRegistryShared = newPATConnectionRegistry() patConnectionRegistryShared = newPATConnectionRegistry()
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err) require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.User{}, &model.APIToken{}, &model.MCPAuditLog{}, &model.Server{}, &model.WAF{})) require.NoError(t, db.AutoMigrate(&model.User{}, &model.APIToken{}, &model.MCPAuditLog{}, &model.Server{}, &model.WAF{}))
singleton.DB = db singleton.DB = db
singleton.Conf = &singleton.ConfigClass{Config: &model.Config{JWTTimeout: 1}} singleton.Conf = &singleton.ConfigClass{Config: &model.Config{JWTTimeout: 1}}
@@ -52,6 +54,7 @@ func setupMCPTest(t *testing.T) (func(), uint64) {
singleton.ServerShared = sc singleton.ServerShared = sc
cleanup := func() { cleanup := func() {
_ = sqlDB.Close()
singleton.DB = originalDB singleton.DB = originalDB
singleton.ServerShared = originalServer singleton.ServerShared = originalServer
singleton.Conf = originalConf singleton.Conf = originalConf
+13 -2
View File
@@ -2,7 +2,7 @@ package controller
import ( import (
"encoding/json" "encoding/json"
"errors" "fmt"
"time" "time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -115,11 +115,22 @@ func handleServerExec(c *gin.Context, raw json.RawMessage) (any, error) {
// handlers do, so MCP isError=true and audit outcome=agent_error. Non-zero // handlers do, so MCP isError=true and audit outcome=agent_error. Non-zero
// ExitCode alone is a normal command outcome, not a tool error. // ExitCode alone is a normal command outcome, not a tool error.
if res.Error != "" { if res.Error != "" {
return nil, errors.New(res.Error) return nil, &execToolError{mcpError: mcpError{Code: model.MCPOutcomeAgentError, Msg: res.Error}, result: res}
} }
return res, nil return res, nil
} }
type execToolError struct {
mcpError
result model.ExecResult
}
func (err *execToolError) StructuredResult() any { return err.result }
func (err *execToolError) Error() string {
return fmt.Sprintf("%s: %s", err.Code, err.Msg)
}
// callAgentTimeout 给 dashboard 侧 CallAgent 计算等待上限。 // callAgentTimeout 给 dashboard 侧 CallAgent 计算等待上限。
// 在用户请求的 timeout 基础上加 5s buffer,让 agent 端的 hard timeout 先触发, // 在用户请求的 timeout 基础上加 5s buffer,让 agent 端的 hard timeout 先触发,
// 这样 dashboard 收到的总是结构化结果(包含 timed_out=true), // 这样 dashboard 收到的总是结构化结果(包含 timed_out=true),
@@ -78,6 +78,12 @@ func TestServerExec_AgentReportedErrorBecomesToolError(t *testing.T) {
"agent ExecResult.Error must propagate as MCP tool error; got %+v", tcr) "agent ExecResult.Error must propagate as MCP tool error; got %+v", tcr)
require.Contains(t, tcr.Content[0].Text, "agent disabled command execution", require.Contains(t, tcr.Content[0].Text, "agent disabled command execution",
"tool error text must surface the agent-reported error message") "tool error text must surface the agent-reported error message")
var result model.ExecResult
structuredJSON, err := json.Marshal(tcr.StructuredContent)
require.NoError(t, err)
require.NoError(t, json.Unmarshal(structuredJSON, &result))
require.Equal(t, -1, result.ExitCode)
require.Equal(t, "agent disabled command execution", result.Error)
require.Eventually(t, func() bool { require.Eventually(t, func() bool {
var got model.MCPAuditLog var got model.MCPAuditLog
+10 -8
View File
@@ -136,7 +136,8 @@ func handleServerList(c *gin.Context, raw json.RawMessage) (any, error) {
if !tok.CanAccessServer(s.ID) { if !tok.CanAccessServer(s.ID) {
continue continue
} }
online := !s.LastActive.IsZero() && now.Sub(s.LastActive) < onlineWindow runtime := s.RuntimeSnapshot()
online := !runtime.LastActive.IsZero() && now.Sub(runtime.LastActive) < onlineWindow
if args.OnlineOnly && !online { if args.OnlineOnly && !online {
continue continue
} }
@@ -145,11 +146,11 @@ func handleServerList(c *gin.Context, raw json.RawMessage) (any, error) {
Name: s.Name, Name: s.Name,
UUID: s.UUID, UUID: s.UUID,
Online: online, Online: online,
LastActive: s.LastActive, LastActive: runtime.LastActive,
} }
if s.Host != nil { if runtime.Host != nil {
item.Platform = s.Host.Platform item.Platform = runtime.Host.Platform
item.Arch = s.Host.Arch item.Arch = runtime.Host.Arch
} }
if s.GeoIP != nil { if s.GeoIP != nil {
item.IPv4 = s.GeoIP.IP.IPv4Addr item.IPv4 = s.GeoIP.IP.IPv4Addr
@@ -187,15 +188,16 @@ func handleServerGet(c *gin.Context, raw json.RawMessage) (any, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
runtime := s.RuntimeSnapshot()
return map[string]any{ return map[string]any{
"id": s.ID, "id": s.ID,
"name": s.Name, "name": s.Name,
"uuid": s.UUID, "uuid": s.UUID,
"note": s.Note, "note": s.Note,
"public_note": s.PublicNote, "public_note": s.PublicNote,
"host": s.Host, "host": runtime.Host,
"state": s.State, "state": runtime.State,
"geoip": s.GeoIP, "geoip": s.GeoIP,
"last_active": s.LastActive, "last_active": runtime.LastActive,
}, nil }, nil
} }
-94
View File
@@ -19,12 +19,8 @@ import (
"time" "time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/hashicorp/go-uuid"
"github.com/nezhahq/nezha/model" "github.com/nezhahq/nezha/model"
"github.com/nezhahq/nezha/pkg/utils" "github.com/nezhahq/nezha/pkg/utils"
pb "github.com/nezhahq/nezha/proto"
"github.com/nezhahq/nezha/service/rpc"
"github.com/nezhahq/nezha/service/singleton" "github.com/nezhahq/nezha/service/singleton"
) )
@@ -931,96 +927,6 @@ func xferSizeFromU64(raw uint64) (int64, error) {
return int64(raw), nil return int64(raw), nil
} }
// openFsTransferStream 走 IOStream 通道与目标 agent 建立一条专用大文件流。
// 返回的 io.ReadWriteCloser 既能 Read(接收 agent→dashboard 字节)又能 Write
// (发送 dashboard→agent 字节);调用方通过 readXferFixedHeader 解析控制帧。
//
// 内部步骤:
// 1. 分配 streamIdCreateStream(streamId, 0, serverID) 在 NezhaHandler 注册
// 一个 ioStreamContexttargetServerID 用于 agent 侧 stream 归属校验。
// 2. 通过 server 当前的 RequestTask 流发 TaskTypeFsTransfer,把 streamId +
// req JSON 下发给 agent。
// 3. 等 agent 通过 IOStream() RPC 完成 magic 引导帧并 AgentConnected。
// 4. 返回 agent 端流和 cleanupCloseStream)。
//
// 任何步骤失败都会 CloseStream 释放资源;调用方只需在 defer cleanup() 即可。
func openFsTransferStream(ctx context.Context, serverID uint64, req *model.FsTransferRequest) (io.ReadWriteCloser, func(), error) {
if singleton.Conf == nil || !singleton.Conf.MCPEnabled() {
return nil, func() {}, errors.New("MCP is disabled by the dashboard administrator")
}
server, _ := singleton.ServerShared.Get(serverID)
if server == nil {
return nil, func() {}, errors.New("server offline")
}
if server.GetTaskStream() == nil {
return nil, func() {}, errors.New("server offline")
}
streamId, err := uuid.GenerateUUID()
if err != nil {
return nil, func() {}, err
}
req.StreamID = streamId
if err := rpc.NezhaHandlerSingleton.CreateStreamWithPurpose(streamId, 0, serverID, rpc.PurposeMCPTransfer); err != nil {
return nil, func() {}, err
}
cleanup := func() { _ = rpc.NezhaHandlerSingleton.CloseStream(streamId) }
body, err := json.Marshal(req)
if err != nil {
cleanup()
return nil, func() {}, err
}
// 关闭 entry-check 与 SendTask 之间的 TOCTOUstream 已注册后再复查一次
// kill switch / ctx 取消,确保 disable sweep 要么扫到这条已注册 stream、
// 要么这里读到 disabled,绝不会在禁用/吊销后仍把 transfer 任务发给 agent。
if singleton.Conf == nil || !singleton.Conf.MCPEnabled() {
cleanup()
return nil, func() {}, errors.New("MCP is disabled by the dashboard administrator")
}
if err := ctx.Err(); err != nil {
cleanup()
return nil, func() {}, err
}
if err := server.SendTask(&pb.Task{
Type: model.TaskTypeFsTransfer,
Data: string(body),
}); err != nil {
cleanup()
if errors.Is(err, model.ErrTaskStreamOffline) {
return nil, func() {}, errors.New("server offline")
}
return nil, func() {}, err
}
agentStream, ok := rpc.NezhaHandlerSingleton.WaitForAgent(ctx, streamId, 30*time.Second)
if !ok {
cleanup()
return nil, func() {}, errors.New("agent did not attach within 30s")
}
// After attach, the relay blocks in IOStreamWrapper.Read, which only
// honours the gRPC stream context — not this per-transfer ctx. Without
// the watcher below, a PAT revocation (deleteAPIToken cancels ctx) or a
// client disconnect would leave a stalled/compromised agent pinning this
// goroutine + IOStream until process restart. Closing the stream on
// ctx.Done() unblocks the agent-side handler (iw.Wait) so Read returns.
watcherDone := make(chan struct{})
go func() {
select {
case <-ctx.Done():
_ = rpc.NezhaHandlerSingleton.CloseStream(streamId)
case <-watcherDone:
}
}()
wrappedCleanup := func() {
close(watcherDone)
cleanup()
}
return agentStream, wrappedCleanup, nil
}
// revalidateTransferEntry 在消费一次性 URL 时重新检查 mint 阶段的全部前置。 // revalidateTransferEntry 在消费一次性 URL 时重新检查 mint 阶段的全部前置。
// 这是 mint→consume 之间发生权限变化(PAT 吊销、scope/whitelist 收紧、 // 这是 mint→consume 之间发生权限变化(PAT 吊销、scope/whitelist 收紧、
// server 转手)时的兜底闸门:HMAC 签发与 sync.Map 一次性消费机制本身只能 // server 转手)时的兜底闸门:HMAC 签发与 sync.Map 一次性消费机制本身只能
@@ -3,6 +3,9 @@ package controller
import ( import (
"context" "context"
"encoding/json" "encoding/json"
"fmt"
"sync"
"sync/atomic"
"testing" "testing"
"time" "time"
@@ -23,16 +26,36 @@ func newFakeAgentIO() *grpcx.IOStreamWrapper {
// wrapper is closed, exactly the post-attach state where nothing watches the // wrapper is closed, exactly the post-attach state where nothing watches the
// per-transfer context. // per-transfer context.
type fakeAgentStream struct { type fakeAgentStream struct {
closed chan struct{} closed chan struct{}
closeOnce sync.Once
closeSeen chan struct{}
recvDone chan struct{}
closeCall atomic.Int32
} }
func (f *fakeAgentStream) Recv() (*pb.IOStreamData, error) { func (f *fakeAgentStream) Recv() (*pb.IOStreamData, error) {
<-f.closed <-f.closed
close(f.recvDone)
return nil, context.Canceled return nil, context.Canceled
} }
func (f *fakeAgentStream) Send(*pb.IOStreamData) error { return nil } func (f *fakeAgentStream) Send(*pb.IOStreamData) error { return nil }
func (f *fakeAgentStream) Context() context.Context { return context.Background() } func (f *fakeAgentStream) Context() context.Context { return context.Background() }
func (f *fakeAgentStream) closeEndpoint() {
f.closeOnce.Do(func() { close(f.closeSeen) })
}
func (f *fakeAgentStream) Close() error {
f.closeCall.Add(1)
f.closeEndpoint()
select {
case <-f.closed:
default:
close(f.closed)
}
return nil
}
// transferRevokableContext only cancels a context; the post-attach relay // transferRevokableContext only cancels a context; the post-attach relay
// (readXferFixedHeader / relayDownloadFrames / io.CopyN) and IOStreamWrapper.Read // (readXferFixedHeader / relayDownloadFrames / io.CopyN) and IOStreamWrapper.Read
// do not watch it. A revoked PAT (or a disconnected HTTP client) must still // do not watch it. A revoked PAT (or a disconnected HTTP client) must still
@@ -59,18 +82,17 @@ func TestOpenFsTransferStream_CancelClosesAttachedStream(t *testing.T) {
t.Cleanup(func() { singleton.ServerShared = originalShared }) t.Cleanup(func() { singleton.ServerShared = originalShared })
streamIDCh := make(chan string, 1) streamIDCh := make(chan string, 1)
agentStreamCh := make(chan *fakeAgentStream, 1)
attachReady := make(chan struct{})
go func() { go func() {
task := <-stream.sent task := <-stream.sent
var req model.FsTransferRequest var req model.FsTransferRequest
_ = json.Unmarshal([]byte(task.GetData()), &req) require.NoError(t, json.Unmarshal([]byte(task.GetData()), &req))
streamIDCh <- req.StreamID streamIDCh <- req.StreamID
deadline := time.Now().Add(2 * time.Second) fakeAgent := &fakeAgentStream{closed: make(chan struct{}), closeSeen: make(chan struct{}), recvDone: make(chan struct{})}
for time.Now().Before(deadline) { agentStreamCh <- fakeAgent
if err := rpc.NezhaHandlerSingleton.AgentConnected(req.StreamID, newFakeAgentIO()); err == nil { require.NoError(t, rpc.NezhaHandlerSingleton.AgentConnected(req.StreamID, grpcx.NewIOStreamWrapper(fakeAgent)))
return close(attachReady)
}
time.Sleep(5 * time.Millisecond)
}
}() }()
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
@@ -83,14 +105,65 @@ func TestOpenFsTransferStream_CancelClosesAttachedStream(t *testing.T) {
defer cleanup() defer cleanup()
streamID := <-streamIDCh streamID := <-streamIDCh
<-attachReady
_, getErr := rpc.NezhaHandlerSingleton.GetStream(streamID) _, getErr := rpc.NezhaHandlerSingleton.GetStream(streamID)
require.NoError(t, getErr, "stream must be live before cancel") require.NoError(t, getErr, "stream must be live before cancel")
agentStream := <-agentStreamCh
readDone := make(chan error, 1)
go func() {
_, readErr := streamIO.Read(make([]byte, 16))
readDone <- readErr
}()
cancel() cancel()
require.Eventually(t, func() bool { select {
_, e := rpc.NezhaHandlerSingleton.GetStream(streamID) case <-readDone:
return e != nil case <-time.After(time.Second):
}, 2*time.Second, 10*time.Millisecond, t.Fatal("cancelling the transfer context must unblock the actual streamIO.Read")
"cancelling the transfer context must tear down the attached IOStream") }
select {
case <-agentStream.closeSeen:
case <-time.After(time.Second):
t.Fatal("cancelling the transfer context must close the fake agent endpoint")
}
select {
case <-agentStream.closed:
case <-time.After(time.Second):
t.Fatal("cancelling the transfer context must close the handler endpoint")
}
select {
case <-agentStream.recvDone:
case <-time.After(time.Second):
t.Fatal("cancelling the transfer context must let the fake handler exit")
}
require.Equal(t, int32(1), agentStream.closeCall.Load(), "attached endpoint must be closed exactly once")
require.Equal(t, 0, rpc.NezhaHandlerSingleton.StreamCount())
for index := 0; index < 40; index++ {
require.NoError(t, rpc.NezhaHandlerSingleton.CreateStream(fmt.Sprintf("cancel-reuse-%d", index), 0, 7))
}
require.ErrorIs(t, rpc.NezhaHandlerSingleton.CreateStream("cancel-reuse-over", 0, 7), rpc.ErrTooManyStreamsForServer)
for index := 0; index < 40; index++ {
require.NoError(t, rpc.NezhaHandlerSingleton.CloseStream(fmt.Sprintf("cancel-reuse-%d", index)))
}
require.Equal(t, 0, rpc.NezhaHandlerSingleton.StreamCount())
cleanupDone := make(chan struct{})
go func() {
var wg sync.WaitGroup
for range 32 {
wg.Add(1)
go func() {
defer wg.Done()
cleanup()
}()
}
wg.Wait()
close(cleanupDone)
}()
select {
case <-cleanupDone:
case <-time.After(2 * time.Second):
t.Fatal("concurrent cleanup calls must complete")
}
} }
@@ -586,5 +586,3 @@ func (b *bigReader) Read(p []byte) (int, error) {
b.remaining -= int64(n) b.remaining -= int64(n)
return n, nil return n, nil
} }
@@ -0,0 +1,84 @@
package controller
import (
"context"
"encoding/json"
"errors"
"io"
"sync"
"time"
"github.com/hashicorp/go-uuid"
"github.com/nezhahq/nezha/model"
pb "github.com/nezhahq/nezha/proto"
"github.com/nezhahq/nezha/service/rpc"
"github.com/nezhahq/nezha/service/singleton"
)
// openFsTransferStream owns the task-to-agent IOStream lifecycle. The returned
// cleanup is safe for concurrent callers and shares ownership with cancellation.
func openFsTransferStream(ctx context.Context, serverID uint64, req *model.FsTransferRequest) (io.ReadWriteCloser, func(), error) {
if singleton.Conf == nil || !singleton.Conf.MCPEnabled() {
return nil, func() {}, errors.New("MCP is disabled by the dashboard administrator")
}
server, _ := singleton.ServerShared.Get(serverID)
if server == nil || server.GetTaskStream() == nil {
return nil, func() {}, errors.New("server offline")
}
handler := rpc.NezhaHandlerSingleton
streamID, err := uuid.GenerateUUID()
if err != nil {
return nil, func() {}, err
}
req.StreamID = streamID
if err := handler.CreateStreamWithPurpose(streamID, 0, serverID, rpc.PurposeMCPTransfer); err != nil {
return nil, func() {}, err
}
var cleanupOnce sync.Once
cleanup := func() { cleanupOnce.Do(func() { _ = handler.CloseStream(streamID) }) }
body, err := json.Marshal(req)
if err != nil {
cleanup()
return nil, func() {}, err
}
// The stream is owned by cleanup until the caller receives it; every failure path releases it.
if singleton.Conf == nil || !singleton.Conf.MCPEnabled() {
cleanup()
return nil, func() {}, errors.New("MCP is disabled by the dashboard administrator")
}
if err := ctx.Err(); err != nil {
cleanup()
return nil, func() {}, err
}
if err := server.SendTask(&pb.Task{Type: model.TaskTypeFsTransfer, Data: string(body)}); err != nil {
cleanup()
if errors.Is(err, model.ErrTaskStreamOffline) {
return nil, func() {}, errors.New("server offline")
}
return nil, func() {}, err
}
agentStream, ok := handler.WaitForAgent(ctx, streamID, 30*time.Second)
if !ok {
cleanup()
return nil, func() {}, errors.New("agent did not attach within 30s")
}
watcherDone := make(chan struct{})
var watcherOnce sync.Once
go func() {
select {
case <-ctx.Done():
cleanup()
case <-watcherDone:
}
}()
wrappedCleanup := func() {
watcherOnce.Do(func() { close(watcherDone) })
cleanup()
}
return agentStream, wrappedCleanup, nil
}
+7 -5
View File
@@ -9,7 +9,6 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/goccy/go-json" "github.com/goccy/go-json"
"github.com/jinzhu/copier"
"gorm.io/gorm" "gorm.io/gorm"
"github.com/nezhahq/nezha/model" "github.com/nezhahq/nezha/model"
@@ -31,10 +30,13 @@ import (
// @Router /server [get] // @Router /server [get]
func listServer(c *gin.Context) ([]*model.Server, error) { func listServer(c *gin.Context) ([]*model.Server, error) {
slist := singleton.ServerShared.GetSortedList() slist := singleton.ServerShared.GetSortedList()
ssl := make([]*model.Server, 0, len(slist))
var ssl []*model.Server for _, server := range slist {
if err := copier.Copy(&ssl, &slist); err != nil { if server == nil {
return nil, err continue
}
runtime := server.RuntimeSnapshot()
ssl = append(ssl, server.RuntimeCopy(runtime))
} }
return ssl, nil return ssl, nil
} }
@@ -32,6 +32,8 @@ func setupServerGroupVisibilityFixture(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err) require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.Server{}, &model.ServerGroup{}, &model.ServerGroupServer{}, &model.User{})) require.NoError(t, db.AutoMigrate(&model.Server{}, &model.ServerGroup{}, &model.ServerGroupServer{}, &model.User{}))
singleton.DB = db singleton.DB = db
@@ -55,6 +57,7 @@ func setupServerGroupVisibilityFixture(t *testing.T) {
singleton.ServerShared = singleton.NewServerClass() singleton.ServerShared = singleton.NewServerClass()
t.Cleanup(func() { t.Cleanup(func() {
_ = sqlDB.Close()
singleton.DB = originalDB singleton.DB = originalDB
singleton.Cache = originalCache singleton.Cache = originalCache
singleton.Loc = originalLoc singleton.Loc = originalLoc
@@ -39,20 +39,14 @@ func setupServiceDispatchPATFixture(t *testing.T) {
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err) require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
require.NoError(t, db.AutoMigrate(&model.Service{}, &model.Server{}, &model.User{}, &model.ServiceHistory{})) require.NoError(t, db.AutoMigrate(&model.Service{}, &model.Server{}, &model.User{}, &model.ServiceHistory{}))
singleton.DB = db singleton.DB = db
singleton.Loc = time.UTC singleton.Loc = time.UTC
singleton.Cache = cache.New(time.Minute, time.Minute) singleton.Cache = cache.New(time.Minute, time.Minute)
singleton.Localizer = i18n.NewLocalizer("en_US", "nezha", "translations", i18n.Translations) 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() sc := singleton.NewEmptyServerClassForTest()
for _, id := range []uint64{1, 2} { for _, id := range []uint64{1, 2} {
s := &model.Server{} s := &model.Server{}
@@ -61,6 +55,13 @@ func setupServiceDispatchPATFixture(t *testing.T) {
sc.InsertForTest(s) sc.InsertForTest(s)
} }
singleton.ServerShared = sc singleton.ServerShared = sc
// ServiceSentinel 在构造时会调 CronShared.AddFunc 注册每日/每周维护任务,
// 必须先于 NewServiceSentinel 装配。
singleton.CronShared = singleton.NewCronClass()
sentinel, err := singleton.NewServiceSentinel(make(chan *model.Service, 4))
require.NoError(t, err)
singleton.ServiceSentinelShared = sentinel
singleton.UserLock.Lock() singleton.UserLock.Lock()
singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}} singleton.UserInfoMap = map[uint64]model.UserInfo{100: {Role: model.RoleMember}}
@@ -68,6 +69,8 @@ func setupServiceDispatchPATFixture(t *testing.T) {
t.Cleanup(func() { t.Cleanup(func() {
sentinel.Close() sentinel.Close()
singleton.CronShared.Close()
_ = sqlDB.Close()
singleton.ServiceSentinelShared = originalSentinel singleton.ServiceSentinelShared = originalSentinel
singleton.CronShared = originalCron singleton.CronShared = originalCron
singleton.DB = originalDB singleton.DB = originalDB
@@ -32,6 +32,8 @@ func setupTenancyTest(t *testing.T) func() {
} }
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
require.NoError(t, err) require.NoError(t, err)
sqlDB, err := db.DB()
require.NoError(t, err)
require.NoError(t, db.AutoMigrate( require.NoError(t, db.AutoMigrate(
&model.User{}, &model.User{},
&model.Cron{}, &model.Cron{},
@@ -47,6 +49,7 @@ func setupTenancyTest(t *testing.T) func() {
singleton.DDNSShared = singleton.NewEmptyDDNSClassForTest() singleton.DDNSShared = singleton.NewEmptyDDNSClassForTest()
singleton.NotificationShared = singleton.NewEmptyNotificationClassForTest() singleton.NotificationShared = singleton.NewEmptyNotificationClassForTest()
return func() { return func() {
_ = sqlDB.Close()
singleton.DB = originalDB singleton.DB = originalDB
singleton.Localizer = originalLocalizer singleton.Localizer = originalLocalizer
singleton.ServerShared = originalServer singleton.ServerShared = originalServer
@@ -118,16 +121,16 @@ func TestTenancy_CreateDDNS_InjectedUserIDIgnored(t *testing.T) {
defer setupTenancyTest(t)() defer setupTenancyTest(t)()
body := map[string]any{ body := map[string]any{
"name": "evil-ddns", "name": "evil-ddns",
"provider": "webhook", "provider": "webhook",
"access_id": "x", "access_id": "x",
"access_secret": "y", "access_secret": "y",
"webhook_url": "http://127.0.0.1/", "webhook_url": "http://127.0.0.1/",
"webhook_method": "GET", "webhook_method": "GET",
"webhook_request_type": "json", "webhook_request_type": "json",
"webhook_request_body": "", "webhook_request_body": "",
"webhook_headers": "", "webhook_headers": "",
"user_id": 999, // attacker "user_id": 999, // attacker
} }
c := ctxAsMemberWithBody(10, body) c := ctxAsMemberWithBody(10, body)
_, err := createDDNS(c) _, err := createDDNS(c)
@@ -155,12 +158,12 @@ func TestTenancy_UpdateDDNS_ForeignOwnerRejected(t *testing.T) {
require.NoError(t, singleton.DB.Create(&foreign).Error) require.NoError(t, singleton.DB.Create(&foreign).Error)
c := ctxAsMemberWithBody(10, map[string]any{ c := ctxAsMemberWithBody(10, map[string]any{
"name": "hijacked", "name": "hijacked",
"provider": "webhook", "provider": "webhook",
"access_id": "x", "access_id": "x",
"access_secret": "y", "access_secret": "y",
"webhook_url": "http://attacker/", "webhook_url": "http://attacker/",
"webhook_method": "GET", "webhook_method": "GET",
"webhook_request_type": "json", "webhook_request_type": "json",
}) })
c.Params = gin.Params{{Key: "id", Value: itoa(foreign.ID)}} c.Params = gin.Params{{Key: "id", Value: itoa(foreign.ID)}}
@@ -237,7 +240,7 @@ func TestTenancy_UpdateNotificationGroup_ForeignOwnerRejected(t *testing.T) {
require.NoError(t, singleton.DB.Create(&foreign).Error) require.NoError(t, singleton.DB.Create(&foreign).Error)
c := ctxAsMemberWithBody(10, map[string]any{ c := ctxAsMemberWithBody(10, map[string]any{
"name": "hijacked", "name": "hijacked",
"notifications": []uint64{}, "notifications": []uint64{},
}) })
c.Params = gin.Params{{Key: "id", Value: itoa(foreign.ID)}} c.Params = gin.Params{{Key: "id", Value: itoa(foreign.ID)}}
@@ -0,0 +1,79 @@
package controller
import (
"context"
"sync"
"time"
"github.com/gorilla/websocket"
)
type websocketPingWriter interface {
WriteMessage(messageType int, data []byte) error
}
type websocketPingConnection interface {
websocketPingWriter
Close() error
}
type websocketPingTransport struct {
websocketPingWriter
closeOnce sync.Once
closeErr error
close func() error
}
func newWebsocketPingTransport(writer websocketPingWriter, close func() error) *websocketPingTransport {
return &websocketPingTransport{websocketPingWriter: writer, close: close}
}
func (transport *websocketPingTransport) Close() error {
transport.closeOnce.Do(func() { transport.closeErr = transport.close() })
return transport.closeErr
}
func websocketPingLoop(ctx context.Context, ticks <-chan time.Time, writer websocketPingWriter) error {
for {
select {
case <-ctx.Done():
return nil
case _, ok := <-ticks:
if !ok {
return nil
}
}
select {
case <-ctx.Done():
return nil
default:
}
if err := writer.WriteMessage(websocket.PingMessage, []byte{}); err != nil {
return err
}
}
}
func startWebsocketPing(ctx context.Context, ticks <-chan time.Time, connection websocketPingConnection) func() {
workerContext, cancel := context.WithCancel(ctx)
workerDone := make(chan struct{})
go func() {
defer close(workerDone)
_ = websocketPingLoop(workerContext, ticks, connection)
}()
return func() {
_ = connection.Close()
cancel()
<-workerDone
}
}
func startWebsocketPingTicker(ctx context.Context, interval time.Duration, connection websocketPingConnection) func() {
ticker := time.NewTicker(interval)
stop := startWebsocketPing(ctx, ticker.C, connection)
return func() {
ticker.Stop()
stop()
}
}
@@ -0,0 +1,152 @@
package controller
import (
"context"
"errors"
"sync"
"testing"
"time"
"github.com/stretchr/testify/require"
)
type pingWriterFake struct {
mu sync.Mutex
writeCalls int
writeErr error
writeStarted chan struct{}
continueWrite chan struct{}
}
func (writer *pingWriterFake) Close() error { return nil }
type permanentlyBlockedPingWriter struct {
writeStarted chan struct{}
transportClosed chan struct{}
}
func (writer *permanentlyBlockedPingWriter) WriteMessage(int, []byte) error {
close(writer.writeStarted)
<-writer.transportClosed
return errors.New("transport closed")
}
func (writer *permanentlyBlockedPingWriter) Close() error {
close(writer.transportClosed)
return nil
}
func (writer *pingWriterFake) WriteMessage(int, []byte) error {
writer.mu.Lock()
writer.writeCalls++
writer.mu.Unlock()
if writer.writeStarted != nil {
close(writer.writeStarted)
<-writer.continueWrite
}
return writer.writeErr
}
func (writer *pingWriterFake) calls() int {
writer.mu.Lock()
defer writer.mu.Unlock()
return writer.writeCalls
}
func TestWebsocketPingLoop_stopsAndJoinsWithoutWritingAfterStop(t *testing.T) {
// Given
ticks := make(chan time.Time, 1)
writer := &pingWriterFake{writeStarted: make(chan struct{}), continueWrite: make(chan struct{})}
stop := startWebsocketPing(context.Background(), ticks, writer)
// When
ticks <- time.Time{}
<-writer.writeStarted
close(writer.continueWrite)
stop()
ticks <- time.Time{}
// Then
require.Equal(t, 1, writer.calls())
}
func TestWebsocketPingLoop_exitsOnWriteError(t *testing.T) {
// Given
ticks := make(chan time.Time, 1)
writer := &pingWriterFake{writeErr: errors.New("closed")}
done := make(chan error, 1)
go func() { done <- websocketPingLoop(context.Background(), ticks, writer) }()
// When
ticks <- time.Time{}
// Then
require.Error(t, <-done)
ticks <- time.Time{}
require.Equal(t, 1, writer.calls())
}
func TestWebsocketPingLoop_cleanupOverlapsTickAndJoinsWriter(t *testing.T) {
// Given
ticks := make(chan time.Time, 1)
writer := &pingWriterFake{
writeStarted: make(chan struct{}),
continueWrite: make(chan struct{}),
}
stop := startWebsocketPing(context.Background(), ticks, writer)
ticks <- time.Time{}
<-writer.writeStarted
// When
stopped := make(chan struct{})
go func() {
stop()
close(stopped)
}()
select {
case <-stopped:
require.Fail(t, "ping worker stop returned before the in-flight write joined")
default:
}
close(writer.continueWrite)
<-stopped
ticks <- time.Time{}
// Then
require.Equal(t, 1, writer.calls())
}
func TestWebsocketPingStop_unblocksPermanentlyBlockedWriteBeforeJoin(t *testing.T) {
// Given
ticks := make(chan time.Time, 1)
writer := &permanentlyBlockedPingWriter{
writeStarted: make(chan struct{}),
transportClosed: make(chan struct{}),
}
stop := startWebsocketPing(context.Background(), ticks, writer)
ticks <- time.Time{}
<-writer.writeStarted
// When
stopped := make(chan struct{})
go func() {
stop()
close(stopped)
}()
// Then
deadline := time.NewTimer(time.Second)
defer deadline.Stop()
select {
case <-writer.transportClosed:
case <-deadline.C:
require.Fail(t, "stop did not close the blocked ping transport")
}
select {
case <-stopped:
case <-deadline.C:
require.Fail(t, "stop did not join after closing the blocked ping transport")
}
}
var _ websocketPingWriter = (*pingWriterFake)(nil)
+77
View File
@@ -0,0 +1,77 @@
package rpc
import (
"errors"
"log"
"github.com/nezhahq/nezha/model"
"github.com/nezhahq/nezha/proto"
"github.com/nezhahq/nezha/service/singleton"
)
func DispatchTask(serviceSentinelDispatchBus <-chan *model.Service) {
for task := range serviceSentinelDispatchBus {
if task == nil {
continue
}
switch task.Cover {
case model.ServiceCoverIgnoreAll:
for id, enabled := range task.SkipServers {
if !enabled {
continue
}
server, _ := singleton.ServerShared.Get(id)
if server == nil {
continue
}
if !canSendTaskToServer(task, server) {
continue
}
if err := server.SendTask(task.PB()); err != nil && !errors.Is(err, model.ErrTaskStreamOffline) {
log.Printf("NEZHA>> DispatchTask send error (server=%d): %v", id, err)
}
}
case model.ServiceCoverAll:
for id, server := range singleton.ServerShared.GetList() {
if server == nil || task.SkipServers[id] {
continue
}
if !canSendTaskToServer(task, server) {
continue
}
if err := server.SendTask(task.PB()); err != nil && !errors.Is(err, model.ErrTaskStreamOffline) {
log.Printf("NEZHA>> DispatchTask send error (server=%d): %v", id, err)
}
}
}
}
}
func DispatchKeepalive() {
singleton.CronShared.AddFunc("@every 20s", func() {
list := singleton.ServerShared.GetSortedList()
for _, s := range list {
if s == nil {
continue
}
if err := s.SendTask(&proto.Task{Type: model.TaskTypeKeepalive}); err != nil && !errors.Is(err, model.ErrTaskStreamOffline) {
log.Printf("NEZHA>> Keepalive send error (server=%d): %v", s.ID, err)
}
}
})
}
func canSendTaskToServer(task *model.Service, server *model.Server) bool {
var role model.Role
singleton.UserLock.RLock()
if u, ok := singleton.UserInfoMap[task.UserID]; !ok {
role = model.RoleMember
} else {
role = u.Role
}
singleton.UserLock.RUnlock()
return task.UserID == server.GetUserID() || role.IsAdmin()
}
+10 -10
View File
@@ -25,18 +25,18 @@ func (MCPAuditLog) TableName() string {
} }
const ( const (
MCPOutcomeOK = "ok" MCPOutcomeOK = "ok"
MCPOutcomeScopeDenied = "scope_denied" MCPOutcomeScopeDenied = "scope_denied"
MCPOutcomePermDenied = "permission_denied" MCPOutcomePermDenied = "permission_denied"
MCPOutcomeServerOffline = "server_offline" MCPOutcomeServerOffline = "server_offline"
MCPOutcomeAgentTimeout = "agent_timeout" MCPOutcomeAgentTimeout = "agent_timeout"
MCPOutcomeAgentError = "agent_error" MCPOutcomeAgentError = "agent_error"
// MCPOutcomeMCPDisabled 区分 “管理员按下 kill switch 把 MCP 关了” 与 // MCPOutcomeMCPDisabled 区分 “管理员按下 kill switch 把 MCP 关了” 与
// “agent 真出故障” 两种语义:前者属 forbidden 类、不应该触发 agent // “agent 真出故障” 两种语义:前者属 forbidden 类、不应该触发 agent
// 故障告警;详见 service/rpc/mcp_rpc.go 里 ErrMCPDisabled 的注释。 // 故障告警;详见 service/rpc/mcp_rpc.go 里 ErrMCPDisabled 的注释。
MCPOutcomeMCPDisabled = "mcp_disabled" MCPOutcomeMCPDisabled = "mcp_disabled"
MCPOutcomeInvalidArgs = "invalid_args" MCPOutcomeInvalidArgs = "invalid_args"
MCPOutcomeRateLimited = "rate_limited" MCPOutcomeRateLimited = "rate_limited"
MCPOutcomeUnsupportedAgent = "unsupported_agent" MCPOutcomeUnsupportedAgent = "unsupported_agent"
MCPOutcomeInternalError = "internal_error" MCPOutcomeInternalError = "internal_error"
) )
+3 -3
View File
@@ -16,9 +16,9 @@ type SettingForm struct {
AgentRealIPHeader string `json:"agent_real_ip_header,omitempty" validate:"optional"` // Agent真实IP AgentRealIPHeader string `json:"agent_real_ip_header,omitempty" validate:"optional"` // Agent真实IP
UserTemplate string `json:"user_template,omitempty" validate:"optional"` UserTemplate string `json:"user_template,omitempty" validate:"optional"`
AgentTLS bool `json:"tls,omitempty" validate:"optional"` AgentTLS bool `json:"tls,omitempty" validate:"optional"`
EnableIPChangeNotification bool `json:"enable_ip_change_notification,omitempty" validate:"optional"` EnableIPChangeNotification bool `json:"enable_ip_change_notification,omitempty" validate:"optional"`
EnablePlainIPInNotification bool `json:"enable_plain_ip_in_notification,omitempty" validate:"optional"` EnablePlainIPInNotification bool `json:"enable_plain_ip_in_notification,omitempty" validate:"optional"`
EnableMCP *bool `json:"enable_mcp,omitempty" validate:"optional"` EnableMCP *bool `json:"enable_mcp,omitempty" validate:"optional"`
} }
@@ -17,11 +17,11 @@ import (
// a fresh exec/fs task to the agent AFTER EnableMCP=false. // a fresh exec/fs task to the agent AFTER EnableMCP=false.
// //
// This test drives the worst-case interleaving deterministically: // This test drives the worst-case interleaving deterministically:
// 1. CallAgent passes the upfront observer check (observer still false). // 1. CallAgent passes the upfront observer check (observer still false).
// 2. The operator flips the observer to "disabled" and runs the cancel sweep // 2. The operator flips the observer to "disabled" and runs the cancel sweep
// while CallAgent is paused between the check and Store. // while CallAgent is paused between the check and Store.
// 3. CallAgent resumes; it MUST observe the kill switch on the post-Store // 3. CallAgent resumes; it MUST observe the kill switch on the post-Store
// re-check and return ErrMCPDisabled WITHOUT sending the task. // re-check and return ErrMCPDisabled WITHOUT sending the task.
// //
// With the race present, CallAgent sends the task and blocks until timeout // With the race present, CallAgent sends the task and blocks until timeout
// (ErrAgentTimeout) — the agent received a fresh task past the kill switch. // (ErrAgentTimeout) — the agent received a fresh task past the kill switch.
@@ -58,13 +58,14 @@ func TestCallAgent_KillSwitchBeatsRegistrationAfterSweep(t *testing.T) {
// Arrange the interleaving: hook fires once CallAgent is about to register, // Arrange the interleaving: hook fires once CallAgent is about to register,
// flipping the kill switch and running the cancel sweep so the not-yet-Stored // flipping the kill switch and running the cancel sweep so the not-yet-Stored
// entry is missed by the sweep. // entry is missed by the sweep.
testKillSwitchAfterUpfrontCheck = func() { hook := func() {
mu.Lock() mu.Lock()
killed = true killed = true
mu.Unlock() mu.Unlock()
CancelAllMCPInflight() CancelAllMCPInflight()
} }
t.Cleanup(func() { testKillSwitchAfterUpfrontCheck = nil }) testKillSwitchAfterUpfrontCheck.Store(&hook)
t.Cleanup(func() { testKillSwitchAfterUpfrontCheck.Store(nil) })
_, err := CallAgent(context.Background(), target, model.TaskTypeExec, _, err := CallAgent(context.Background(), target, model.TaskTypeExec,
model.ExecRequest{Cmd: "x"}, 1*time.Second) model.ExecRequest{Cmd: "x"}, 1*time.Second)
+50 -17
View File
@@ -49,8 +49,13 @@ var disarmedKillSwitch = func() bool { return false }
// testKillSwitchAfterUpfrontCheck, when non-nil, runs inside CallAgent between // testKillSwitchAfterUpfrontCheck, when non-nil, runs inside CallAgent between
// the upfront kill-switch check and the inflight registration. Production // the upfront kill-switch check and the inflight registration. Production
// leaves it nil; tests use it to drive the registration-after-sweep race // leaves it nil; tests use it to drive the registration-after-sweep race
// deterministically. Guarded by the same atomic.Pointer for race-freedom. // deterministically.
var testKillSwitchAfterUpfrontCheck func() var testKillSwitchAfterUpfrontCheck atomic.Pointer[func()]
var (
testMCPResultBeforeCancellationCheck atomic.Pointer[func()]
testMCPResultAfterCancellationCheck atomic.Pointer[func()]
)
// SetMCPKillSwitchObserver installs the kill-switch probe the dashboard // SetMCPKillSwitchObserver installs the kill-switch probe the dashboard
// owns. Idempotent; the dashboard wires it at startup. Passing nil // owns. Idempotent; the dashboard wires it at startup. Passing nil
@@ -126,8 +131,8 @@ func CallAgent(ctx context.Context, serverID uint64, taskType uint64, params any
cancelled: new(atomic.Bool), cancelled: new(atomic.Bool),
} }
if hook := testKillSwitchAfterUpfrontCheck; hook != nil { if hook := testKillSwitchAfterUpfrontCheck.Load(); hook != nil {
hook() (*hook)()
} }
mcpInflight.Store(taskID, entry) mcpInflight.Store(taskID, entry)
@@ -153,6 +158,7 @@ func CallAgent(ctx context.Context, serverID uint64, taskType uint64, params any
} }
return nil, err return nil, err
} }
notifyMCPTaskDispatched(serverID, taskID, taskType)
waitCtx := ctx waitCtx := ctx
var cancel context.CancelFunc var cancel context.CancelFunc
@@ -163,6 +169,9 @@ func CallAgent(ctx context.Context, serverID uint64, taskType uint64, params any
select { select {
case res := <-resultCh: case res := <-resultCh:
if hook := testMCPResultBeforeCancellationCheck.Load(); hook != nil {
(*hook)()
}
// Cancel must beat a late agent reply: Go select picks a random // Cancel must beat a late agent reply: Go select picks a random
// ready case, so if CancelAllMCPInflight closed cancelCh after the // ready case, so if CancelAllMCPInflight closed cancelCh after the
// agent already filled resultCh we could still surface success. // agent already filled resultCh we could still surface success.
@@ -170,9 +179,13 @@ func CallAgent(ctx context.Context, serverID uint64, taskType uint64, params any
// contract documented above ("CancelAllMCPInflight 期间被中断 → // contract documented above ("CancelAllMCPInflight 期间被中断 →
// ErrMCPDisabled") and what TestUpdateConfig_DisablingMCPInvokesKillSwitch // ErrMCPDisabled") and what TestUpdateConfig_DisablingMCPInvokesKillSwitch
// expects. // expects.
if entry.cancelled.Load() { if !entry.claimResult() {
return nil, ErrMCPDisabled return nil, ErrMCPDisabled
} }
notifyMCPTaskResultAccepted(entry.serverID, res.GetId(), res.GetType())
if hook := testMCPResultAfterCancellationCheck.Load(); hook != nil {
(*hook)()
}
if res == nil { if res == nil {
return nil, errors.New("agent returned nil result") return nil, errors.New("agent returned nil result")
} }
@@ -201,20 +214,39 @@ func CallAgent(ctx context.Context, serverID uint64, taskType uint64, params any
// purely by attacker-controlled TaskResult.Id (same bug class as commit // purely by attacker-controlled TaskResult.Id (same bug class as commit
// 02129f1 in the cron path). // 02129f1 in the cron path).
// //
// cancelled flips to true the instant CancelAllMCPInflight observes the // cancelled flips to true when CancelAllMCPInflight wins the entry lock before
// entry. Every code path that could complete the call — the CallAgent // the result is claimed. Every code path that could complete the call — the
// select on resultCh, deliverMCPResult, deliverMCPResultFromReporter // CallAgent select on resultCh, deliverMCPResult, deliverMCPResultFromReporter
// MUST consult it before treating an agent reply as authoritative, otherwise // MUST consult it before treating an agent reply as authoritative.
// a TaskResult delivered concurrently with the kill switch can win Go's
// random select tiebreak and surface success after EnableMCP=false.
type mcpInflightEntry struct { type mcpInflightEntry struct {
serverID uint64 serverID uint64
result chan *pb.TaskResult result chan *pb.TaskResult
cancel chan struct{} cancel chan struct{}
cancelled *atomic.Bool cancelled *atomic.Bool
mu sync.Mutex
claimed bool
closeOnce sync.Once closeOnce sync.Once
} }
func (e *mcpInflightEntry) claimResult() bool {
e.mu.Lock()
defer e.mu.Unlock()
if e.cancelled.Load() {
return false
}
e.claimed = true
return true
}
func (e *mcpInflightEntry) cancelCall() {
e.mu.Lock()
if !e.claimed {
e.cancelled.Store(true)
}
e.mu.Unlock()
e.closeCancel()
}
// closeCancel closes the entry's cancel channel exactly once. Concurrent // closeCancel closes the entry's cancel channel exactly once. Concurrent
// CancelAllMCPInflight sweeps (two admin PATCH /setting requests both // CancelAllMCPInflight sweeps (two admin PATCH /setting requests both
// disabling MCP) would otherwise race a non-atomic check-then-close and // disabling MCP) would otherwise race a non-atomic check-then-close and
@@ -243,10 +275,7 @@ func CancelAllMCPInflight() int {
if !ok { if !ok {
return true return true
} }
if entry.cancelled != nil { entry.cancelCall()
entry.cancelled.Store(true)
}
entry.closeCancel()
mcpInflight.Delete(key) mcpInflight.Delete(key)
cancelled++ cancelled++
return true return true
@@ -300,7 +329,9 @@ func deliverMCPResult(res *pb.TaskResult) {
if !ok { if !ok {
return return
} }
if entry.cancelled != nil && entry.cancelled.Load() { entry.mu.Lock()
defer entry.mu.Unlock()
if entry.cancelled.Load() {
return return
} }
select { select {
@@ -335,7 +366,9 @@ func deliverMCPResultFromReporter(res *pb.TaskResult, reporterID uint64) {
res.GetId(), entry.serverID, reporterID) res.GetId(), entry.serverID, reporterID)
return return
} }
if entry.cancelled != nil && entry.cancelled.Load() { entry.mu.Lock()
defer entry.mu.Unlock()
if entry.cancelled.Load() {
return return
} }
select { select {
+64 -3
View File
@@ -3,6 +3,7 @@ package rpc
import ( import (
"context" "context"
"errors" "errors"
"sync/atomic"
"testing" "testing"
"time" "time"
@@ -33,6 +34,17 @@ func TestCallAgent_KillSwitchBeatsConcurrentLateResult(t *testing.T) {
cleanup := installFakeServer(t, target, stream) cleanup := installFakeServer(t, target, stream)
defer cleanup() defer cleanup()
resultSelected := make(chan struct{})
resumeResult := make(chan struct{})
var resultHook atomic.Pointer[func()]
hook := func() {
close(resultSelected)
<-resumeResult
}
resultHook.Store(&hook)
testMCPResultBeforeCancellationCheck.Store(resultHook.Load())
t.Cleanup(func() { testMCPResultBeforeCancellationCheck.Store(nil) })
delivered := make(chan struct{}) delivered := make(chan struct{})
go func() { go func() {
sent := <-stream.sent sent := <-stream.sent
@@ -42,18 +54,67 @@ func TestCallAgent_KillSwitchBeatsConcurrentLateResult(t *testing.T) {
Successful: true, Successful: true,
Data: `{"exit_code":0,"stdout":"should-not-surface"}`, Data: `{"exit_code":0,"stdout":"should-not-surface"}`,
}, target) }, target)
CancelAllMCPInflight()
close(delivered) close(delivered)
}() }()
_, err := CallAgent(context.Background(), target, model.TaskTypeExec, errCh := make(chan error, 1)
model.ExecRequest{Cmd: "x"}, 2*time.Second) go func() {
_, err := CallAgent(context.Background(), target, model.TaskTypeExec,
model.ExecRequest{Cmd: "x"}, 2*time.Second)
errCh <- err
}()
<-resultSelected
CancelAllMCPInflight()
close(resumeResult)
<-delivered <-delivered
err := <-errCh
if !errors.Is(err, ErrMCPDisabled) { if !errors.Is(err, ErrMCPDisabled) {
t.Fatalf("kill switch must win the race with a late agent reply; want ErrMCPDisabled, got %v", err) t.Fatalf("kill switch must win the race with a late agent reply; want ErrMCPDisabled, got %v", err)
} }
} }
func TestCallAgent_ResultBeforeKillSwitchReturnsSuccess(t *testing.T) {
const target uint64 = 7303
stream := newFakeStream()
cleanup := installFakeServer(t, target, stream)
defer cleanup()
resultClaimed := make(chan struct{})
resumeResult := make(chan struct{})
var resultHook atomic.Pointer[func()]
hook := func() {
close(resultClaimed)
<-resumeResult
}
resultHook.Store(&hook)
testMCPResultAfterCancellationCheck.Store(resultHook.Load())
t.Cleanup(func() { testMCPResultAfterCancellationCheck.Store(nil) })
go func() {
sent := <-stream.sent
deliverMCPResultFromReporter(&pb.TaskResult{
Id: sent.GetId(),
Type: model.TaskTypeExec,
Successful: true,
Data: `{"exit_code":0}`,
}, target)
}()
errCh := make(chan error, 1)
go func() {
_, err := CallAgent(context.Background(), target, model.TaskTypeExec,
model.ExecRequest{Cmd: "x"}, 2*time.Second)
errCh <- err
}()
<-resultClaimed
CancelAllMCPInflight()
close(resumeResult)
if err := <-errCh; err != nil {
t.Fatalf("result claimed before kill switch must succeed, got %v", err)
}
}
// CancelAllMCPInflight must eagerly evict entries so a stale TaskResult // CancelAllMCPInflight must eagerly evict entries so a stale TaskResult
// that arrives after the kill switch cannot still land in resultCh. // that arrives after the kill switch cannot still land in resultCh.
// Without the cancelled flag this entry would still be reachable through // Without the cancelled flag this entry would still be reachable through
+15 -8
View File
@@ -18,20 +18,27 @@ import (
type fakeTaskStream struct { type fakeTaskStream struct {
sent chan *pb.Task sent chan *pb.Task
delay time.Duration delay time.Duration
err error
} }
func newFakeStream() *fakeTaskStream { func newFakeStream() *fakeTaskStream {
return &fakeTaskStream{sent: make(chan *pb.Task, 4)} return &fakeTaskStream{sent: make(chan *pb.Task, 4)}
} }
func (s *fakeTaskStream) Send(t *pb.Task) error { s.sent <- t; return nil } func (s *fakeTaskStream) Send(t *pb.Task) error {
func (s *fakeTaskStream) Recv() (*pb.TaskResult, error) { return nil, context.Canceled } if s.err != nil {
func (s *fakeTaskStream) SetHeader(metadata.MD) error { return nil } return s.err
func (s *fakeTaskStream) SendHeader(metadata.MD) error { return nil } }
func (s *fakeTaskStream) SetTrailer(metadata.MD) {} s.sent <- t
func (s *fakeTaskStream) Context() context.Context { return context.Background() } return nil
func (s *fakeTaskStream) SendMsg(any) error { return nil } }
func (s *fakeTaskStream) RecvMsg(any) error { return context.Canceled } func (s *fakeTaskStream) Recv() (*pb.TaskResult, error) { return nil, context.Canceled }
func (s *fakeTaskStream) SetHeader(metadata.MD) error { return nil }
func (s *fakeTaskStream) SendHeader(metadata.MD) error { return nil }
func (s *fakeTaskStream) SetTrailer(metadata.MD) {}
func (s *fakeTaskStream) Context() context.Context { return context.Background() }
func (s *fakeTaskStream) SendMsg(any) error { return nil }
func (s *fakeTaskStream) RecvMsg(any) error { return context.Canceled }
func installFakeServer(t *testing.T, id uint64, stream pb.NezhaService_RequestTaskServer) func() { func installFakeServer(t *testing.T, id uint64, stream pb.NezhaService_RequestTaskServer) func() {
t.Helper() t.Helper()
+6 -6
View File
@@ -120,12 +120,12 @@ func TestIsReservedDashboardHostCollapsesEquivalentForms(t *testing.T) {
}) })
reserved := []string{ reserved := []string{
"panel.example.com.", // trailing dot, no port "panel.example.com.", // trailing dot, no port
"panel.example.com.:8008", // trailing dot with port "panel.example.com.:8008", // trailing dot with port
"PANEL.EXAMPLE.COM.", // trailing dot, mixed case "PANEL.EXAMPLE.COM.", // trailing dot, mixed case
"[0:0:0:0:0:0:0:1]:8008", // IPv6 expanded form of ::1 "[0:0:0:0:0:0:0:1]:8008", // IPv6 expanded form of ::1
"::1", // IPv6 compressed, bare "::1", // IPv6 compressed, bare
"[::1]", // IPv6 compressed, bracketed "[::1]", // IPv6 compressed, bracketed
} }
for _, d := range reserved { for _, d := range reserved {
if !IsReservedDashboardHost(d) { if !IsReservedDashboardHost(d) {