From ab25662ddd6e4e75f843c4c4d0cb7eb1547389f9 Mon Sep 17 00:00:00 2001 From: naiba Date: Sat, 30 May 2026 15:56:44 +0000 Subject: [PATCH] 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 --- .gitignore | 1 + cmd/dashboard/controller/alertrule.go | 11 +- .../controller/alertrule_pat_fanout_test.go | 142 +++ cmd/dashboard/controller/api_token.go | 286 +++++ .../api_token_legacy_migration_test.go | 78 ++ .../api_token_optional_scope_test.go | 71 ++ .../controller/api_token_revoke_registry.go | 148 +++ .../api_token_revoke_registry_race_test.go | 46 + .../api_token_revoke_registry_test.go | 99 ++ cmd/dashboard/controller/api_token_scope.go | 135 +++ .../api_token_scope_empty_doc_test.go | 58 + .../controller/api_token_scope_test.go | 223 ++++ .../api_token_server_config_test.go | 55 + .../api_token_server_whitelist_test.go | 79 ++ cmd/dashboard/controller/api_token_test.go | 597 ++++++++++ .../batch_move_pat_whitelist_test.go | 71 ++ cmd/dashboard/controller/controller.go | 185 +-- cmd/dashboard/controller/cron.go | 51 +- .../controller/cron_cover_validation_test.go | 54 + .../controller/cron_dispatch_pat_test.go | 220 ++++ .../cron_list_pat_whitelist_test.go | 65 ++ .../controller/cron_pat_whitelist_test.go | 161 +++ .../controller/cron_service_cover_pat_test.go | 342 ++++++ .../controller/cron_update_owner_uid_test.go | 103 ++ cmd/dashboard/controller/csrf.go | 89 ++ cmd/dashboard/controller/csrf_test.go | 159 +++ cmd/dashboard/controller/fm.go | 10 +- .../frontend_fallback_api_tokens_test.go | 25 + cmd/dashboard/controller/jwt.go | 17 +- cmd/dashboard/controller/jwt_session_test.go | 117 ++ cmd/dashboard/controller/mcp.go | 493 ++++++++ cmd/dashboard/controller/mcp_audit.go | 48 + .../controller/mcp_body_limit_test.go | 72 ++ cmd/dashboard/controller/mcp_capability.go | 74 ++ .../controller/mcp_capability_test.go | 52 + .../controller/mcp_classify_disabled_test.go | 28 + .../controller/mcp_enable_flag_test.go | 96 ++ .../controller/mcp_end_to_end_test.go | 432 +++++++ .../controller/mcp_kill_switch_test.go | 226 ++++ .../controller/mcp_method_not_allowed_test.go | 77 ++ cmd/dashboard/controller/mcp_origin.go | 117 ++ cmd/dashboard/controller/mcp_origin_test.go | 109 ++ cmd/dashboard/controller/mcp_ratelimit.go | 90 ++ .../controller/mcp_ratelimit_bypass_test.go | 68 ++ .../controller/mcp_ratelimit_defaults_test.go | 18 + .../mcp_ratelimit_malformed_test.go | 74 ++ .../controller/mcp_ratelimit_prune_test.go | 62 + .../controller/mcp_sdk_compat_test.go | 195 ++++ cmd/dashboard/controller/mcp_test.go | 303 +++++ cmd/dashboard/controller/mcp_tools_exec.go | 117 ++ .../controller/mcp_tools_exec_error_test.go | 91 ++ .../controller/mcp_tools_exec_timeout_test.go | 63 + cmd/dashboard/controller/mcp_tools_fs.go | 248 ++++ cmd/dashboard/controller/mcp_tools_fs_test.go | 168 +++ cmd/dashboard/controller/mcp_tools_meta.go | 49 + cmd/dashboard/controller/mcp_tools_server.go | 148 +++ .../controller/mcp_tools_server_test.go | 113 ++ cmd/dashboard/controller/mcp_transfer.go | 1010 +++++++++++++++++ .../controller/mcp_transfer_audit_throttle.go | 72 ++ .../mcp_transfer_audit_throttle_test.go | 68 ++ .../controller/mcp_transfer_cancel_test.go | 96 ++ .../mcp_transfer_consume_authz_test.go | 66 ++ .../mcp_transfer_correctness_test.go | 590 ++++++++++ .../mcp_transfer_data_frame_collision_test.go | 30 + .../mcp_transfer_download_finalcheck_test.go | 124 ++ .../mcp_transfer_failure_audit_test.go | 147 +++ .../controller/mcp_transfer_gc_test.go | 53 + .../mcp_transfer_no_full_buffer_test.go | 113 ++ .../controller/mcp_transfer_path_cap_test.go | 38 + .../mcp_transfer_size_guard_test.go | 56 + .../controller/mcp_transfer_spool.go | 44 + .../mcp_transfer_upload_args_test.go | 94 ++ cmd/dashboard/controller/oauth2.go | 1 + cmd/dashboard/controller/oauth2_csrf_test.go | 28 + cmd/dashboard/controller/oauth2_test.go | 196 ++++ .../controller/pat_whitelist_view_test.go | 63 + cmd/dashboard/controller/permissions.go | 464 +++++++- .../permissions_cover_fanout_test.go | 182 +++ cmd/dashboard/controller/rest_scope_test.go | 202 ++++ cmd/dashboard/controller/scope_allof_test.go | 63 + cmd/dashboard/controller/scope_doc.go | 124 ++ .../controller/scope_doc_consistency_test.go | 161 +++ cmd/dashboard/controller/server.go | 49 +- cmd/dashboard/controller/server_group.go | 24 + .../server_group_update_pat_test.go | 117 ++ .../server_group_visibility_test.go | 80 ++ cmd/dashboard/controller/service.go | 47 +- .../controller/service_cache_key_test.go | 58 + .../controller/service_dispatch_pat_test.go | 185 +++ .../service_list_pat_whitelist_test.go | 66 ++ .../service_skip_enabled_only_test.go | 66 ++ .../controller/service_visibility_test.go | 37 + cmd/dashboard/controller/setting.go | 50 +- .../setting_enable_mcp_save_failure_test.go | 73 ++ .../controller/setting_enable_mcp_test.go | 84 ++ .../controller/stream_pat_authz_test.go | 90 ++ .../controller/tenant_isolation_test.go | 315 +++++ cmd/dashboard/controller/terminal.go | 10 +- cmd/dashboard/controller/transfer.go | 14 + .../controller/transfer_cancel_authz_test.go | 104 ++ .../controller/transfer_pat_whitelist_test.go | 156 +++ .../transfer_retry_pat_whitelist_test.go | 78 ++ .../controller/trigger_task_pat_scope_test.go | 99 ++ cmd/dashboard/controller/ws.go | 43 +- .../controller/ws_stream_visibility_test.go | 58 +- cmd/dashboard/main.go | 10 + cmd/dashboard/rpc/rpc.go | 45 +- go.mod | 9 +- go.sum | 12 + model/alertrule.go | 57 + model/alertrule_pat_whitelist_test.go | 227 ++++ model/api_token.go | 373 ++++++ model/api_token_migration_test.go | 112 ++ model/api_token_test.go | 148 +++ model/api_token_unified_scope_test.go | 82 ++ model/common.go | 5 + model/config.go | 23 + model/cron.go | 42 + model/cron_admin_pat_whitelist_test.go | 116 ++ model/cron_pat_whitelist_test.go | 102 ++ model/mcp_audit.go | 42 + model/mcp_enabled_atomic_test.go | 39 + model/nat.go | 19 + model/nat_pat_whitelist_test.go | 55 + model/server.go | 173 ++- model/server_transfer.go | 10 + model/service.go | 204 ++++ model/service_ignoreall_false_test.go | 51 + model/setting_api.go | 1 + pkg/grpcx/io_stream_wrapper.go | 89 +- .../io_stream_wrapper_concurrent_send_test.go | 71 ++ pkg/grpcx/io_stream_wrapper_test.go | 106 ++ service/rpc/io_stream.go | 183 ++- service/rpc/io_stream_race_test.go | 92 ++ service/rpc/mcp_cancel_double_close_test.go | 55 + service/rpc/mcp_kill_switch_race_test.go | 35 + .../mcp_kill_switch_registration_race_test.go | 80 ++ service/rpc/mcp_rpc.go | 351 ++++++ service/rpc/mcp_rpc_helper_doc_test.go | 61 + service/rpc/mcp_rpc_kill_switch_race_test.go | 95 ++ service/rpc/mcp_rpc_spoof_test.go | 149 +++ service/rpc/mcp_rpc_test.go | 165 +++ service/rpc/nezha.go | 79 +- .../rpc/request_task_missing_server_test.go | 25 + service/rpc/request_task_stale_stream_test.go | 46 + service/rpc/testdata_helper_test.go | 20 + service/rpc/wait_for_agent_revoke_test.go | 56 + service/singleton/crontask.go | 17 +- service/singleton/server.go | 37 +- .../singleton/server_delete_missing_test.go | 32 + service/singleton/server_transfer.go | 23 +- service/singleton/singleton.go | 12 +- service/singleton/testhelpers.go | 65 ++ 153 files changed, 16974 insertions(+), 244 deletions(-) create mode 100644 cmd/dashboard/controller/alertrule_pat_fanout_test.go create mode 100644 cmd/dashboard/controller/api_token.go create mode 100644 cmd/dashboard/controller/api_token_legacy_migration_test.go create mode 100644 cmd/dashboard/controller/api_token_optional_scope_test.go create mode 100644 cmd/dashboard/controller/api_token_revoke_registry.go create mode 100644 cmd/dashboard/controller/api_token_revoke_registry_race_test.go create mode 100644 cmd/dashboard/controller/api_token_revoke_registry_test.go create mode 100644 cmd/dashboard/controller/api_token_scope.go create mode 100644 cmd/dashboard/controller/api_token_scope_empty_doc_test.go create mode 100644 cmd/dashboard/controller/api_token_scope_test.go create mode 100644 cmd/dashboard/controller/api_token_server_config_test.go create mode 100644 cmd/dashboard/controller/api_token_server_whitelist_test.go create mode 100644 cmd/dashboard/controller/api_token_test.go create mode 100644 cmd/dashboard/controller/batch_move_pat_whitelist_test.go create mode 100644 cmd/dashboard/controller/cron_cover_validation_test.go create mode 100644 cmd/dashboard/controller/cron_dispatch_pat_test.go create mode 100644 cmd/dashboard/controller/cron_list_pat_whitelist_test.go create mode 100644 cmd/dashboard/controller/cron_pat_whitelist_test.go create mode 100644 cmd/dashboard/controller/cron_service_cover_pat_test.go create mode 100644 cmd/dashboard/controller/cron_update_owner_uid_test.go create mode 100644 cmd/dashboard/controller/csrf.go create mode 100644 cmd/dashboard/controller/csrf_test.go create mode 100644 cmd/dashboard/controller/frontend_fallback_api_tokens_test.go create mode 100644 cmd/dashboard/controller/mcp.go create mode 100644 cmd/dashboard/controller/mcp_audit.go create mode 100644 cmd/dashboard/controller/mcp_body_limit_test.go create mode 100644 cmd/dashboard/controller/mcp_capability.go create mode 100644 cmd/dashboard/controller/mcp_capability_test.go create mode 100644 cmd/dashboard/controller/mcp_classify_disabled_test.go create mode 100644 cmd/dashboard/controller/mcp_enable_flag_test.go create mode 100644 cmd/dashboard/controller/mcp_end_to_end_test.go create mode 100644 cmd/dashboard/controller/mcp_kill_switch_test.go create mode 100644 cmd/dashboard/controller/mcp_method_not_allowed_test.go create mode 100644 cmd/dashboard/controller/mcp_origin.go create mode 100644 cmd/dashboard/controller/mcp_origin_test.go create mode 100644 cmd/dashboard/controller/mcp_ratelimit.go create mode 100644 cmd/dashboard/controller/mcp_ratelimit_bypass_test.go create mode 100644 cmd/dashboard/controller/mcp_ratelimit_defaults_test.go create mode 100644 cmd/dashboard/controller/mcp_ratelimit_malformed_test.go create mode 100644 cmd/dashboard/controller/mcp_ratelimit_prune_test.go create mode 100644 cmd/dashboard/controller/mcp_sdk_compat_test.go create mode 100644 cmd/dashboard/controller/mcp_test.go create mode 100644 cmd/dashboard/controller/mcp_tools_exec.go create mode 100644 cmd/dashboard/controller/mcp_tools_exec_error_test.go create mode 100644 cmd/dashboard/controller/mcp_tools_exec_timeout_test.go create mode 100644 cmd/dashboard/controller/mcp_tools_fs.go create mode 100644 cmd/dashboard/controller/mcp_tools_fs_test.go create mode 100644 cmd/dashboard/controller/mcp_tools_meta.go create mode 100644 cmd/dashboard/controller/mcp_tools_server.go create mode 100644 cmd/dashboard/controller/mcp_tools_server_test.go create mode 100644 cmd/dashboard/controller/mcp_transfer.go create mode 100644 cmd/dashboard/controller/mcp_transfer_audit_throttle.go create mode 100644 cmd/dashboard/controller/mcp_transfer_audit_throttle_test.go create mode 100644 cmd/dashboard/controller/mcp_transfer_cancel_test.go create mode 100644 cmd/dashboard/controller/mcp_transfer_consume_authz_test.go create mode 100644 cmd/dashboard/controller/mcp_transfer_correctness_test.go create mode 100644 cmd/dashboard/controller/mcp_transfer_data_frame_collision_test.go create mode 100644 cmd/dashboard/controller/mcp_transfer_download_finalcheck_test.go create mode 100644 cmd/dashboard/controller/mcp_transfer_failure_audit_test.go create mode 100644 cmd/dashboard/controller/mcp_transfer_gc_test.go create mode 100644 cmd/dashboard/controller/mcp_transfer_no_full_buffer_test.go create mode 100644 cmd/dashboard/controller/mcp_transfer_path_cap_test.go create mode 100644 cmd/dashboard/controller/mcp_transfer_size_guard_test.go create mode 100644 cmd/dashboard/controller/mcp_transfer_spool.go create mode 100644 cmd/dashboard/controller/mcp_transfer_upload_args_test.go create mode 100644 cmd/dashboard/controller/oauth2_csrf_test.go create mode 100644 cmd/dashboard/controller/oauth2_test.go create mode 100644 cmd/dashboard/controller/pat_whitelist_view_test.go create mode 100644 cmd/dashboard/controller/permissions_cover_fanout_test.go create mode 100644 cmd/dashboard/controller/rest_scope_test.go create mode 100644 cmd/dashboard/controller/scope_allof_test.go create mode 100644 cmd/dashboard/controller/scope_doc.go create mode 100644 cmd/dashboard/controller/scope_doc_consistency_test.go create mode 100644 cmd/dashboard/controller/server_group_update_pat_test.go create mode 100644 cmd/dashboard/controller/service_cache_key_test.go create mode 100644 cmd/dashboard/controller/service_dispatch_pat_test.go create mode 100644 cmd/dashboard/controller/service_list_pat_whitelist_test.go create mode 100644 cmd/dashboard/controller/service_skip_enabled_only_test.go create mode 100644 cmd/dashboard/controller/setting_enable_mcp_save_failure_test.go create mode 100644 cmd/dashboard/controller/setting_enable_mcp_test.go create mode 100644 cmd/dashboard/controller/stream_pat_authz_test.go create mode 100644 cmd/dashboard/controller/tenant_isolation_test.go create mode 100644 cmd/dashboard/controller/transfer_cancel_authz_test.go create mode 100644 cmd/dashboard/controller/transfer_pat_whitelist_test.go create mode 100644 cmd/dashboard/controller/transfer_retry_pat_whitelist_test.go create mode 100644 cmd/dashboard/controller/trigger_task_pat_scope_test.go create mode 100644 model/alertrule_pat_whitelist_test.go create mode 100644 model/api_token.go create mode 100644 model/api_token_migration_test.go create mode 100644 model/api_token_test.go create mode 100644 model/api_token_unified_scope_test.go create mode 100644 model/cron_admin_pat_whitelist_test.go create mode 100644 model/cron_pat_whitelist_test.go create mode 100644 model/mcp_audit.go create mode 100644 model/mcp_enabled_atomic_test.go create mode 100644 model/nat_pat_whitelist_test.go create mode 100644 model/service_ignoreall_false_test.go create mode 100644 pkg/grpcx/io_stream_wrapper_concurrent_send_test.go create mode 100644 pkg/grpcx/io_stream_wrapper_test.go create mode 100644 service/rpc/io_stream_race_test.go create mode 100644 service/rpc/mcp_cancel_double_close_test.go create mode 100644 service/rpc/mcp_kill_switch_race_test.go create mode 100644 service/rpc/mcp_kill_switch_registration_race_test.go create mode 100644 service/rpc/mcp_rpc.go create mode 100644 service/rpc/mcp_rpc_helper_doc_test.go create mode 100644 service/rpc/mcp_rpc_kill_switch_race_test.go create mode 100644 service/rpc/mcp_rpc_spoof_test.go create mode 100644 service/rpc/mcp_rpc_test.go create mode 100644 service/rpc/request_task_missing_server_test.go create mode 100644 service/rpc/request_task_stale_stream_test.go create mode 100644 service/rpc/testdata_helper_test.go create mode 100644 service/rpc/wait_for_agent_revoke_test.go create mode 100644 service/singleton/server_delete_missing_test.go create mode 100644 service/singleton/testhelpers.go diff --git a/.gitignore b/.gitignore index fb3e1c7b..eeb498b6 100644 --- a/.gitignore +++ b/.gitignore @@ -24,3 +24,4 @@ /resource/template/theme-custom /resource/static/custom /cmd/dashboard/docs +.omo/ diff --git a/cmd/dashboard/controller/alertrule.go b/cmd/dashboard/controller/alertrule.go index 0d009252..3c54f4a2 100644 --- a/cmd/dashboard/controller/alertrule.go +++ b/cmd/dashboard/controller/alertrule.go @@ -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 diff --git a/cmd/dashboard/controller/alertrule_pat_fanout_test.go b/cmd/dashboard/controller/alertrule_pat_fanout_test.go new file mode 100644 index 00000000..e78f68af --- /dev/null +++ b/cmd/dashboard/controller/alertrule_pat_fanout_test.go @@ -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}") +} diff --git a/cmd/dashboard/controller/api_token.go b/cmd/dashboard/controller/api_token.go new file mode 100644 index 00000000..2f8d8e44 --- /dev/null +++ b/cmd/dashboard/controller/api_token.go @@ -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 +} diff --git a/cmd/dashboard/controller/api_token_legacy_migration_test.go b/cmd/dashboard/controller/api_token_legacy_migration_test.go new file mode 100644 index 00000000..8b68b366 --- /dev/null +++ b/cmd/dashboard/controller/api_token_legacy_migration_test.go @@ -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:*") +} diff --git a/cmd/dashboard/controller/api_token_optional_scope_test.go b/cmd/dashboard/controller/api_token_optional_scope_test.go new file mode 100644 index 00000000..1b276908 --- /dev/null +++ b/cmd/dashboard/controller/api_token_optional_scope_test.go @@ -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) + } +} diff --git a/cmd/dashboard/controller/api_token_revoke_registry.go b/cmd/dashboard/controller/api_token_revoke_registry.go new file mode 100644 index 00000000..62723593 --- /dev/null +++ b/cmd/dashboard/controller/api_token_revoke_registry.go @@ -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) +} diff --git a/cmd/dashboard/controller/api_token_revoke_registry_race_test.go b/cmd/dashboard/controller/api_token_revoke_registry_race_test.go new file mode 100644 index 00000000..3cb925cd --- /dev/null +++ b/cmd/dashboard/controller/api_token_revoke_registry_race_test.go @@ -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) + } +} diff --git a/cmd/dashboard/controller/api_token_revoke_registry_test.go b/cmd/dashboard/controller/api_token_revoke_registry_test.go new file mode 100644 index 00000000..d0a149e2 --- /dev/null +++ b/cmd/dashboard/controller/api_token_revoke_registry_test.go @@ -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 +} diff --git a/cmd/dashboard/controller/api_token_scope.go b/cmd/dashboard/controller/api_token_scope.go new file mode 100644 index 00000000..a9081ff7 --- /dev/null +++ b/cmd/dashboard/controller/api_token_scope.go @@ -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() + } +} diff --git a/cmd/dashboard/controller/api_token_scope_empty_doc_test.go b/cmd/dashboard/controller/api_token_scope_empty_doc_test.go new file mode 100644 index 00000000..415fc2e4 --- /dev/null +++ b/cmd/dashboard/controller/api_token_scope_empty_doc_test.go @@ -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") + } +} diff --git a/cmd/dashboard/controller/api_token_scope_test.go b/cmd/dashboard/controller/api_token_scope_test.go new file mode 100644 index 00000000..89971537 --- /dev/null +++ b/cmd/dashboard/controller/api_token_scope_test.go @@ -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", + ) +} diff --git a/cmd/dashboard/controller/api_token_server_config_test.go b/cmd/dashboard/controller/api_token_server_config_test.go new file mode 100644 index 00000000..b18a4352 --- /dev/null +++ b/cmd/dashboard/controller/api_token_server_config_test.go @@ -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) +} diff --git a/cmd/dashboard/controller/api_token_server_whitelist_test.go b/cmd/dashboard/controller/api_token_server_whitelist_test.go new file mode 100644 index 00000000..98350ddc --- /dev/null +++ b/cmd/dashboard/controller/api_token_server_whitelist_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/api_token_test.go b/cmd/dashboard/controller/api_token_test.go new file mode 100644 index 00000000..6e1904aa --- /dev/null +++ b/cmd/dashboard/controller/api_token_test.go @@ -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::* 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") +} diff --git a/cmd/dashboard/controller/batch_move_pat_whitelist_test.go b/cmd/dashboard/controller/batch_move_pat_whitelist_test.go new file mode 100644 index 00000000..0647f222 --- /dev/null +++ b/cmd/dashboard/controller/batch_move_pat_whitelist_test.go @@ -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) +} diff --git a/cmd/dashboard/controller/controller.go b/cmd/dashboard/controller/controller.go index a26f317a..6d4d8985 100644 --- a/cmd/dashboard/controller/controller.go +++ b/cmd/dashboard/controller/controller.go @@ -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 / 链接预览会以为站点挂了)。 diff --git a/cmd/dashboard/controller/cron.go b/cmd/dashboard/controller/cron.go index 31027bcc..a6a640a4 100644 --- a/cmd/dashboard/controller/cron.go +++ b/cmd/dashboard/controller/cron.go @@ -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) } diff --git a/cmd/dashboard/controller/cron_cover_validation_test.go b/cmd/dashboard/controller/cron_cover_validation_test.go new file mode 100644 index 00000000..ae5a6752 --- /dev/null +++ b/cmd/dashboard/controller/cron_cover_validation_test.go @@ -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) + } + }) + } +} diff --git a/cmd/dashboard/controller/cron_dispatch_pat_test.go b/cmd/dashboard/controller/cron_dispatch_pat_test.go new file mode 100644 index 00000000..ed442639 --- /dev/null +++ b/cmd/dashboard/controller/cron_dispatch_pat_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/cron_list_pat_whitelist_test.go b/cmd/dashboard/controller/cron_list_pat_whitelist_test.go new file mode 100644 index 00000000..274c04d9 --- /dev/null +++ b/cmd/dashboard/controller/cron_list_pat_whitelist_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/cron_pat_whitelist_test.go b/cmd/dashboard/controller/cron_pat_whitelist_test.go new file mode 100644 index 00000000..12f20c14 --- /dev/null +++ b/cmd/dashboard/controller/cron_pat_whitelist_test.go @@ -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) +} diff --git a/cmd/dashboard/controller/cron_service_cover_pat_test.go b/cmd/dashboard/controller/cron_service_cover_pat_test.go new file mode 100644 index 00000000..c181dc22 --- /dev/null +++ b/cmd/dashboard/controller/cron_service_cover_pat_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/cron_update_owner_uid_test.go b/cmd/dashboard/controller/cron_update_owner_uid_test.go new file mode 100644 index 00000000..3fc4df7a --- /dev/null +++ b/cmd/dashboard/controller/cron_update_owner_uid_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/csrf.go b/cmd/dashboard/controller/csrf.go new file mode 100644 index 00000000..0d3a498c --- /dev/null +++ b/cmd/dashboard/controller/csrf.go @@ -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() + } +} diff --git a/cmd/dashboard/controller/csrf_test.go b/cmd/dashboard/controller/csrf_test.go new file mode 100644 index 00000000..fec7275f --- /dev/null +++ b/cmd/dashboard/controller/csrf_test.go @@ -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)") + } +} diff --git a/cmd/dashboard/controller/fm.go b/cmd/dashboard/controller/fm.go index 4cb3b078..4444c724 100644 --- a/cmd/dashboard/controller/fm.go +++ b/cmd/dashboard/controller/fm.go @@ -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 { diff --git a/cmd/dashboard/controller/frontend_fallback_api_tokens_test.go b/cmd/dashboard/controller/frontend_fallback_api_tokens_test.go new file mode 100644 index 00000000..df22c630 --- /dev/null +++ b/cmd/dashboard/controller/frontend_fallback_api_tokens_test.go @@ -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()) + } +} diff --git a/cmd/dashboard/controller/jwt.go b/cmd/dashboard/controller/jwt.go index 522dcf70..03342e4a 100644 --- a/cmd/dashboard/controller/jwt.go +++ b/cmd/dashboard/controller/jwt.go @@ -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{ diff --git a/cmd/dashboard/controller/jwt_session_test.go b/cmd/dashboard/controller/jwt_session_test.go index 1278e93c..5019c93d 100644 --- a/cmd/dashboard/controller/jwt_session_test.go +++ b/cmd/dashboard/controller/jwt_session_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/mcp.go b/cmd/dashboard/controller/mcp.go new file mode 100644 index 00000000..4dfddd89 --- /dev/null +++ b/cmd/dashboard/controller/mcp.go @@ -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 +} diff --git a/cmd/dashboard/controller/mcp_audit.go b/cmd/dashboard/controller/mcp_audit.go new file mode 100644 index 00000000..6a6048a9 --- /dev/null +++ b/cmd/dashboard/controller/mcp_audit.go @@ -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 diff --git a/cmd/dashboard/controller/mcp_body_limit_test.go b/cmd/dashboard/controller/mcp_body_limit_test.go new file mode 100644 index 00000000..addfce33 --- /dev/null +++ b/cmd/dashboard/controller/mcp_body_limit_test.go @@ -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()) + } +} diff --git a/cmd/dashboard/controller/mcp_capability.go b/cmd/dashboard/controller/mcp_capability.go new file mode 100644 index 00000000..e743f530 --- /dev/null +++ b/cmd/dashboard/controller/mcp_capability.go @@ -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 +} diff --git a/cmd/dashboard/controller/mcp_capability_test.go b/cmd/dashboard/controller/mcp_capability_test.go new file mode 100644 index 00000000..b828c34b --- /dev/null +++ b/cmd/dashboard/controller/mcp_capability_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/mcp_classify_disabled_test.go b/cmd/dashboard/controller/mcp_classify_disabled_test.go new file mode 100644 index 00000000..7fe651fa --- /dev/null +++ b/cmd/dashboard/controller/mcp_classify_disabled_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/mcp_enable_flag_test.go b/cmd/dashboard/controller/mcp_enable_flag_test.go new file mode 100644 index 00000000..d404d3ad --- /dev/null +++ b/cmd/dashboard/controller/mcp_enable_flag_test.go @@ -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) +} diff --git a/cmd/dashboard/controller/mcp_end_to_end_test.go b/cmd/dashboard/controller/mcp_end_to_end_test.go new file mode 100644 index 00000000..ec1725a4 --- /dev/null +++ b/cmd/dashboard/controller/mcp_end_to_end_test.go @@ -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) +} diff --git a/cmd/dashboard/controller/mcp_kill_switch_test.go b/cmd/dashboard/controller/mcp_kill_switch_test.go new file mode 100644 index 00000000..d286d6d0 --- /dev/null +++ b/cmd/dashboard/controller/mcp_kill_switch_test.go @@ -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") + } +} diff --git a/cmd/dashboard/controller/mcp_method_not_allowed_test.go b/cmd/dashboard/controller/mcp_method_not_allowed_test.go new file mode 100644 index 00000000..3b0f923e --- /dev/null +++ b/cmd/dashboard/controller/mcp_method_not_allowed_test.go @@ -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()), "= 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)", + }) +} diff --git a/cmd/dashboard/controller/mcp_origin_test.go b/cmd/dashboard/controller/mcp_origin_test.go new file mode 100644 index 00000000..4daa086f --- /dev/null +++ b/cmd/dashboard/controller/mcp_origin_test.go @@ -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) +} diff --git a/cmd/dashboard/controller/mcp_ratelimit.go b/cmd/dashboard/controller/mcp_ratelimit.go new file mode 100644 index 00000000..0ef3c37f --- /dev/null +++ b/cmd/dashboard/controller/mcp_ratelimit.go @@ -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) diff --git a/cmd/dashboard/controller/mcp_ratelimit_bypass_test.go b/cmd/dashboard/controller/mcp_ratelimit_bypass_test.go new file mode 100644 index 00000000..447e67e2 --- /dev/null +++ b/cmd/dashboard/controller/mcp_ratelimit_bypass_test.go @@ -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) +} diff --git a/cmd/dashboard/controller/mcp_ratelimit_defaults_test.go b/cmd/dashboard/controller/mcp_ratelimit_defaults_test.go new file mode 100644 index 00000000..30350acf --- /dev/null +++ b/cmd/dashboard/controller/mcp_ratelimit_defaults_test.go @@ -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) + } +} diff --git a/cmd/dashboard/controller/mcp_ratelimit_malformed_test.go b/cmd/dashboard/controller/mcp_ratelimit_malformed_test.go new file mode 100644 index 00000000..cc55e6dd --- /dev/null +++ b/cmd/dashboard/controller/mcp_ratelimit_malformed_test.go @@ -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") + } +} diff --git a/cmd/dashboard/controller/mcp_ratelimit_prune_test.go b/cmd/dashboard/controller/mcp_ratelimit_prune_test.go new file mode 100644 index 00000000..3d29691d --- /dev/null +++ b/cmd/dashboard/controller/mcp_ratelimit_prune_test.go @@ -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) + } +} diff --git a/cmd/dashboard/controller/mcp_sdk_compat_test.go b/cmd/dashboard/controller/mcp_sdk_compat_test.go new file mode 100644 index 00000000..f4e4fd56 --- /dev/null +++ b/cmd/dashboard/controller/mcp_sdk_compat_test.go @@ -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) +} diff --git a/cmd/dashboard/controller/mcp_test.go b/cmd/dashboard/controller/mcp_test.go new file mode 100644 index 00000000..7bd3a0a6 --- /dev/null +++ b/cmd/dashboard/controller/mcp_test.go @@ -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) +} diff --git a/cmd/dashboard/controller/mcp_tools_exec.go b/cmd/dashboard/controller/mcp_tools_exec.go new file mode 100644 index 00000000..e54a2cb6 --- /dev/null +++ b/cmd/dashboard/controller/mcp_tools_exec.go @@ -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 +} diff --git a/cmd/dashboard/controller/mcp_tools_exec_error_test.go b/cmd/dashboard/controller/mcp_tools_exec_error_test.go new file mode 100644 index 00000000..89bea23d --- /dev/null +++ b/cmd/dashboard/controller/mcp_tools_exec_error_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/mcp_tools_exec_timeout_test.go b/cmd/dashboard/controller/mcp_tools_exec_timeout_test.go new file mode 100644 index 00000000..f87a50bc --- /dev/null +++ b/cmd/dashboard/controller/mcp_tools_exec_timeout_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/mcp_tools_fs.go b/cmd/dashboard/controller/mcp_tools_fs.go new file mode 100644 index 00000000..9c2081d2 --- /dev/null +++ b/cmd/dashboard/controller/mcp_tools_fs.go @@ -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 +} diff --git a/cmd/dashboard/controller/mcp_tools_fs_test.go b/cmd/dashboard/controller/mcp_tools_fs_test.go new file mode 100644 index 00000000..b34af70c --- /dev/null +++ b/cmd/dashboard/controller/mcp_tools_fs_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/mcp_tools_meta.go b/cmd/dashboard/controller/mcp_tools_meta.go new file mode 100644 index 00000000..4dd6d0b1 --- /dev/null +++ b/cmd/dashboard/controller/mcp_tools_meta.go @@ -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 +} diff --git a/cmd/dashboard/controller/mcp_tools_server.go b/cmd/dashboard/controller/mcp_tools_server.go new file mode 100644 index 00000000..c5c6103a --- /dev/null +++ b/cmd/dashboard/controller/mcp_tools_server.go @@ -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 +} diff --git a/cmd/dashboard/controller/mcp_tools_server_test.go b/cmd/dashboard/controller/mcp_tools_server_test.go new file mode 100644 index 00000000..34778d2e --- /dev/null +++ b/cmd/dashboard/controller/mcp_tools_server_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/mcp_transfer.go b/cmd/dashboard/controller/mcp_transfer.go new file mode 100644 index 00000000..cc2dad9d --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer.go @@ -0,0 +1,1010 @@ +package controller + +import ( + "bytes" + "context" + "crypto/hmac" + "crypto/sha256" + "encoding/binary" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "math" + "net/http" + "strconv" + "strings" + "sync" + "time" + + "github.com/gin-gonic/gin" + "github.com/hashicorp/go-uuid" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/utils" + pb "github.com/nezhahq/nezha/proto" + "github.com/nezhahq/nezha/service/rpc" + "github.com/nezhahq/nezha/service/singleton" +) + +// fs.download_url / fs.upload_url 旁路通道。 +// +// 设计目标:给 LLM 客户端一个不经 MCP 上下文的 URL,去用普通 HTTP 客户端 +// 上传/下载大文件(单文件 hard cap 100MiB,model.MCPFsTransferMaxSize)。 +// +// 传输实现:dashboard ↔ agent 走 gRPC IOStream 双向流(TaskTypeFsTransfer), +// dashboard 一边读 HTTP body 一边推给 agent;不再使用 base64/JSON 包装内容, +// 避免 gRPC 4MiB 单消息上限。 +// +// 安全机制: +// - 一次性 token,存内存 sync.Map,TTL 默认 300s,最多 600s +// - token 绑定 user_id + token_id + server_id + path + direction +// - 走 HMAC-SHA256 防篡改 +// - 命中后立即从内存删除,禁止重放 +// - revalidateTransferEntry 在 consume 时重新校验 PAT/scope/owner,应对 +// mint→consume 之间的权限变化 +// - 上传可选 ?sha256= 端到端校验;下载 NZTO 帧附 agent 计算的 sha +type transferDirection string + +const ( + transferDirDownload transferDirection = "download" + transferDirUpload transferDirection = "upload" + + transferTokenTTLDefault = 300 * time.Second + transferTokenTTLMax = 600 * time.Second + + maxTransferPathLen = 4096 +) + +func validateTransferPath(path string) error { + if path == "" { + return errMCPInvalidArgs("path required") + } + if len(path) > maxTransferPathLen { + return errMCPInvalidArgs("path too long") + } + return nil +} + +type transferEntry struct { + UserID uint64 + TokenID uint64 + ServerID uint64 + Path string + Direction transferDirection + ExpiresAt time.Time + + // Upload-only optional knobs carried from MCP fs.upload_url tool args + // to transferUploadHandler so the upload handler can forward them into + // FsTransferRequest. Empty/false values keep current behaviour for the + // download direction (these fields are simply ignored). + UploadMode string + UploadCreateDirs bool + UploadIfMatchSHA256 string +} + +var ( + transferEntries sync.Map + transferSecretMu sync.Mutex + transferSecretVal string +) + +// transferHMACSecret 返回进程内随机生成的 HMAC key。 +// 这是有意设计:transferEntries 本身也只活在内存 sync.Map 里,dashboard +// 重启等价于全部 token 失效;让 secret 也随进程随机,可以避免“secret 来自 +// 持久化 env 但 entries 已丢”这种半持久化状态,同时保证多副本部署不会 +// 意外互认对方签发的 token(每副本一份独立 secret)。 +func transferHMACSecret() string { + transferSecretMu.Lock() + defer transferSecretMu.Unlock() + if transferSecretVal != "" { + return transferSecretVal + } + transferSecretVal = utils.MustGenerateRandomString(64) + return transferSecretVal +} + +func mintTransferToken(e transferEntry) (string, error) { + id, err := utils.GenerateRandomString(24) + if err != nil { + return "", err + } + mac := hmac.New(sha256.New, []byte(transferHMACSecret())) + fmt.Fprintf(mac, "%s|%d|%d|%d|%s|%d", + e.Direction, e.UserID, e.TokenID, e.ServerID, e.Path, e.ExpiresAt.UnixNano()) + sig := hex.EncodeToString(mac.Sum(nil)) + tok := id + "." + sig + transferEntries.Store(tok, e) + return tok, nil +} + +func consumeTransferToken(tok string, dir transferDirection) (*transferEntry, error) { + raw, ok := transferEntries.LoadAndDelete(tok) + if !ok { + return nil, errors.New("invalid or already-used transfer token") + } + e, _ := raw.(transferEntry) + if e.Direction != dir { + return nil, errors.New("transfer token direction mismatch") + } + if time.Now().After(e.ExpiresAt) { + return nil, errors.New("transfer token expired") + } + return &e, nil +} + +// PurgeTransferEntries drops every minted-but-unconsumed transfer URL. +// EnableMCP=false invokes this so an admin pressing the kill switch +// invalidates the 5–10min trailing window of pre-signed download/upload +// URLs that consumeTransferToken would otherwise still honor. Returns the +// number of entries purged for audit. +func PurgeTransferEntries() int { + purged := 0 + transferEntries.Range(func(key, _ any) bool { + if _, ok := transferEntries.LoadAndDelete(key); ok { + purged++ + } + return true + }) + return purged +} + +// gcExpiredTransferEntries 按 ExpiresAt 删除所有已过期但从未被 consume +// 的 token。kickoffTransferGC 周期调度它,防止 transferEntries 在没人 +// 触发 kill switch 的情况下随时间无界增长。 +func gcExpiredTransferEntries(now time.Time) int { + removed := 0 + transferEntries.Range(func(key, raw any) bool { + e, ok := raw.(transferEntry) + if !ok { + transferEntries.Delete(key) + removed++ + return true + } + if now.After(e.ExpiresAt) { + if _, deleted := transferEntries.LoadAndDelete(key); deleted { + removed++ + } + } + return true + }) + return removed +} + +var transferGCStartOnce sync.Once + +// kickoffTransferGC 启动一个进程级 goroutine 定时回收过期 token, +// 避免每个 dashboard 启动都得 PurgeTransferEntries 才能把表清空。 +// 时间间隔取 transferTokenTTLDefault / 5,对默认 5min TTL 即 1min; +// 既能在 TTL 内多次扫到过期项,也不会让锁竞争变成热点。 +func kickoffTransferGC() { + transferGCStartOnce.Do(func() { + go func() { + ticker := time.NewTicker(transferTokenTTLDefault / 5) + defer ticker.Stop() + for range ticker.C { + gcExpiredTransferEntries(time.Now()) + } + }() + }) +} + +// --- tool: fs.download_url --- + +type fsDownloadURLArgs struct { + ServerID uint64 `json:"server_id"` + Path string `json:"path"` + TTLSeconds int `json:"ttl_seconds,omitempty"` +} + +// fsUploadURLArgs is the upload-side superset of fsDownloadURLArgs. agent's +// FsTransferRequest already supports per-upload Mode / CreateDirs / +// IfMatchSHA256, but fs.upload_url historically reused fsDownloadURLArgs and +// silently dropped these fields. Splitting the arg shape lets the MCP tool +// schema advertise them and mintTransferTool plumb them through to +// transferUploadHandler -> openFsTransferStream. +type fsUploadURLArgs struct { + ServerID uint64 `json:"server_id"` + Path string `json:"path"` + TTLSeconds int `json:"ttl_seconds,omitempty"` + Mode string `json:"mode,omitempty"` + CreateDirs bool `json:"create_dirs,omitempty"` + IfMatchSHA256 string `json:"if_match_sha256,omitempty"` +} + +func init() { + registerMCPTool(&mcpTool{ + Name: "fs.download_url", + Description: "Mint a one-time signed URL to stream a file (<=100MiB) via plain HTTP GET. Bypasses MCP context and uses gRPC IOStream end-to-end.", + InputSchema: map[string]any{ + "type": "object", + "properties": map[string]any{ + "server_id": map[string]any{"type": "integer"}, + "path": map[string]any{"type": "string"}, + "ttl_seconds": map[string]any{"type": "integer", "minimum": 30, "maximum": 600}, + }, + "required": []string{"server_id", "path"}, + }, + RequiredScope: model.ScopeServerRead, + Handler: handleFsDownloadURL, + }) + + registerMCPTool(&mcpTool{ + Name: "fs.upload_url", + Description: "Mint a one-time signed URL to stream a file (<=100MiB) via plain HTTP POST. Caller MUST send Content-Length; optional ?sha256= for end-to-end integrity. mode / create_dirs / if_match_sha256 are forwarded to the agent for atomic chmod / mkdir -p / optimistic concurrency.", + InputSchema: map[string]any{ + "type": "object", + "properties": map[string]any{ + "server_id": map[string]any{"type": "integer"}, + "path": map[string]any{"type": "string"}, + "ttl_seconds": map[string]any{"type": "integer", "minimum": 30, "maximum": 600}, + "mode": map[string]any{"type": "string", "description": "Octal mode like '0644'."}, + "create_dirs": map[string]any{"type": "boolean"}, + "if_match_sha256": map[string]any{"type": "string", "description": "64 hex chars; precondition checked by the agent before overwrite."}, + }, + "required": []string{"server_id", "path"}, + }, + RequiredScope: model.ScopeServerWrite, + Handler: handleFsUploadURL, + }) +} + +func handleFsDownloadURL(c *gin.Context, raw json.RawMessage) (any, error) { + var args fsDownloadURLArgs + if err := decodeToolArgs(raw, &args); err != nil { + return nil, err + } + return mintTransferTool(c, args.ServerID, args.Path, args.TTLSeconds, transferDirDownload, transferEntry{}) +} + +func handleFsUploadURL(c *gin.Context, raw json.RawMessage) (any, error) { + var args fsUploadURLArgs + if err := decodeToolArgs(raw, &args); err != nil { + return nil, err + } + if args.IfMatchSHA256 != "" && len(args.IfMatchSHA256) != 64 { + return nil, errMCPInvalidArgs("if_match_sha256 must be 64 hex chars") + } + return mintTransferTool(c, args.ServerID, args.Path, args.TTLSeconds, transferDirUpload, transferEntry{ + UploadMode: args.Mode, + UploadCreateDirs: args.CreateDirs, + UploadIfMatchSHA256: args.IfMatchSHA256, + }) +} + +func mintTransferTool(c *gin.Context, serverID uint64, path string, ttlSeconds int, dir transferDirection, uploadExtras transferEntry) (any, error) { + srv, err := requireServerAccess(c, serverID) + if err != nil { + return nil, err + } + if err := requireAgentSupportsMCP(srv); err != nil { + return nil, err + } + if err := validateTransferPath(path); err != nil { + return nil, err + } + ttl := time.Duration(ttlSeconds) * time.Second + if ttl <= 0 { + ttl = transferTokenTTLDefault + } + if ttl > transferTokenTTLMax { + ttl = transferTokenTTLMax + } + + tok := APITokenFromContext(c) + if tok == nil { + return nil, errNoToken + } + uid := uint64(0) + if u, ok := c.Get(model.CtxKeyAuthorizedUser); ok { + if user, ok := u.(*model.User); ok && user != nil { + uid = user.ID + } + } + entry := transferEntry{ + UserID: uid, + TokenID: tok.ID, + ServerID: serverID, + Path: path, + Direction: dir, + ExpiresAt: time.Now().Add(ttl), + UploadMode: uploadExtras.UploadMode, + UploadCreateDirs: uploadExtras.UploadCreateDirs, + UploadIfMatchSHA256: uploadExtras.UploadIfMatchSHA256, + } + t, err := mintTransferToken(entry) + if err != nil { + return nil, err + } + scheme := "https" + if c.Request.TLS == nil && c.Request.Header.Get("X-Forwarded-Proto") != "https" { + scheme = "http" + } + host := c.Request.Host + url := fmt.Sprintf("%s://%s/mcp/%s/%s", scheme, host, dir, t) + return map[string]any{ + "url": url, + "method": map[transferDirection]string{transferDirDownload: "GET", transferDirUpload: "POST"}[dir], + "expires_at": entry.ExpiresAt, + }, nil +} + +// --- HTTP handlers --- + +// transferDownloadHandler 处理 GET /mcp/download/:token。 +// 走 IOStream 双向流:dashboard 把 agent 推过来的 chunk 转发给 HTTP 客户端, +// 单文件 hard cap 100MiB(model.MCPFsTransferMaxSize)。 +// transferRevokableContext 把进行中的传输纳入 PAT 撤销注册表。返回的 ctx 在 +// 该 PAT 被 deleteAPIToken 撤销时取消,从而切断已开始的 upload/download; +// 否则只在传输自然结束时由 stop() 注销。stop() 必须 defer 调用。 +func transferRevokableContext(c *gin.Context, e *transferEntry) (context.Context, func()) { + ctx, cancel := context.WithCancel(c.Request.Context()) + dereg := patConnectionRegistryShared.register(e.TokenID, cancel) + return ctx, func() { + dereg() + cancel() + } +} + +func transferDownloadHandler(c *gin.Context) { + tok := c.Param("token") + entry, err := consumeTransferToken(tok, transferDirDownload) + if err != nil { + writeTransferFailureAudit(c, nil, "fs.download", classifyTransferConsumeError(err), err) + c.String(http.StatusUnauthorized, err.Error()) + return + } + if err := revalidateTransferEntry(entry); err != nil { + writeTransferFailureAudit(c, entry, "fs.download", classifyTransferRevalidateError(err), err) + c.String(http.StatusUnauthorized, err.Error()) + return + } + + ctx, stop := transferRevokableContext(c, entry) + defer stop() + + stream, cleanup, err := openFsTransferStream(ctx, entry.ServerID, &model.FsTransferRequest{ + Op: model.MCPFsTransferOpDownload, + Path: entry.Path, + }) + if err != nil { + writeTransferFailureAudit(c, entry, "fs.download", classifyTransferOpenStreamError(err), err) + c.String(http.StatusBadGateway, err.Error()) + return + } + defer cleanup() + + hdr, err := readXferFixedHeader(stream) + if err != nil { + writeTransferFailureAudit(c, entry, "fs.download", model.MCPOutcomeAgentTimeout, err) + c.String(http.StatusBadGateway, "agent did not return download header: "+err.Error()) + return + } + if hdr.IsErr() { + writeTransferFailureAudit(c, entry, "fs.download", model.MCPOutcomeAgentError, errors.New(hdr.ErrMsg)) + c.String(http.StatusBadGateway, hdr.ErrMsg) + return + } + if !bytes.Equal(hdr.Magic, model.MCPFsXferMagicDownloadHdr) { + writeTransferFailureAudit(c, entry, "fs.download", model.MCPOutcomeAgentError, errors.New("unexpected header magic")) + c.String(http.StatusBadGateway, "agent returned unexpected header magic") + return + } + if hdr.Size > model.MCPFsTransferMaxSize { + writeTransferFailureAudit(c, entry, "fs.download", model.MCPOutcomeAgentError, errors.New("file exceeds MCP transfer cap")) + c.String(http.StatusBadGateway, "file exceeds MCP transfer cap (100MiB)") + return + } + + if err := relayDownloadFrames(c, stream, hdr.Size); err != nil { + writeTransferFailureAudit(c, entry, "fs.download", model.MCPOutcomeAgentError, err) + return + } + + _ = singleton.DB.Create(&model.MCPAuditLog{ + UserID: entry.UserID, + TokenID: entry.TokenID, + Tool: "fs.download", + ServerID: entry.ServerID, + Outcome: model.MCPOutcomeOK, + IP: c.GetString(model.CtxKeyRealIPStr), + }).Error +} + +// transferUploadHandler 处理 POST /mcp/upload/:token;body 转发到 agent, +// 单文件 hard cap 100MiB。 +func transferUploadHandler(c *gin.Context) { + tok := c.Param("token") + entry, err := consumeTransferToken(tok, transferDirUpload) + if err != nil { + writeTransferFailureAudit(c, nil, "fs.upload", classifyTransferConsumeError(err), err) + c.String(http.StatusUnauthorized, err.Error()) + return + } + if err := revalidateTransferEntry(entry); err != nil { + writeTransferFailureAudit(c, entry, "fs.upload", classifyTransferRevalidateError(err), err) + c.String(http.StatusUnauthorized, err.Error()) + return + } + + // 1) 体积闸门:Content-Length 必须存在并且不超过 cap。流式上传时这是 + // 唯一能在打开 IOStream 之前就拒掉超大请求的依据,避免 agent 端 + // 拒绝时已经占了一个连接。 + if c.Request.ContentLength < 0 { + writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeInvalidArgs, errors.New("Content-Length required")) + c.String(http.StatusLengthRequired, "Content-Length required") + return + } + if c.Request.ContentLength > model.MCPFsTransferMaxSize { + writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeInvalidArgs, errors.New("body exceeds MCP transfer cap")) + c.String(http.StatusRequestEntityTooLarge, "body exceeds MCP transfer cap (100MiB)") + return + } + size := c.Request.ContentLength + + // 可选的端到端 sha256:通过 query 参数 sha256= 传入,agent 收到全部 + // 字节后会比对;不传则只回带 sha 但不强校验。 + expected := strings.ToLower(strings.TrimSpace(c.Query("sha256"))) + if expected != "" { + if _, decErr := hex.DecodeString(expected); decErr != nil || len(expected) != 64 { + writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeInvalidArgs, errors.New("sha256 must be 64 hex chars")) + c.String(http.StatusBadRequest, "sha256 must be 64 hex chars") + return + } + } + + // 上限再加一个字节做 MaxBytesReader 屏障:若客户端撒谎、实际 body 超过 + // Content-Length,HTTP 层会立即截断并报 413。 + c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, model.MCPFsTransferMaxSize+1) + + ctx, stop := transferRevokableContext(c, entry) + defer stop() + + stream, cleanup, err := openFsTransferStream(ctx, entry.ServerID, &model.FsTransferRequest{ + Op: model.MCPFsTransferOpUpload, + Path: entry.Path, + Size: size, + ExpectedSHA256: expected, + Mode: entry.UploadMode, + CreateDirs: entry.UploadCreateDirs, + IfMatchSHA256: entry.UploadIfMatchSHA256, + }) + if err != nil { + writeTransferFailureAudit(c, entry, "fs.upload", classifyTransferOpenStreamError(err), err) + c.String(http.StatusBadGateway, err.Error()) + return + } + defer cleanup() + + hdr, err := readXferFixedHeader(stream) + if err != nil { + writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentTimeout, err) + c.String(http.StatusBadGateway, "agent did not return upload ready frame: "+err.Error()) + return + } + if hdr.IsErr() { + writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentError, errors.New(hdr.ErrMsg)) + c.String(http.StatusBadGateway, hdr.ErrMsg) + return + } + if !bytes.Equal(hdr.Magic, model.MCPFsXferMagicUploadHdr) { + writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentError, errors.New("unexpected header magic")) + c.String(http.StatusBadGateway, "agent returned unexpected header magic") + return + } + if hdr.Size != size { + writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentError, errors.New("agent acknowledged unexpected size")) + c.String(http.StatusBadGateway, "agent acknowledged unexpected size") + return + } + + if _, copyErr := io.CopyN(stream, c.Request.Body, size); copyErr != nil { + writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentError, copyErr) + c.String(http.StatusBadGateway, "stream relay failed: "+copyErr.Error()) + return + } + + final, err := readXferFixedHeader(stream) + if err != nil { + writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentTimeout, err) + c.String(http.StatusBadGateway, "agent did not acknowledge upload: "+err.Error()) + return + } + if final.IsErr() { + writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentError, errors.New(final.ErrMsg)) + c.String(http.StatusBadGateway, final.ErrMsg) + return + } + if !bytes.Equal(final.Magic, model.MCPFsXferMagicOK) { + writeTransferFailureAudit(c, entry, "fs.upload", model.MCPOutcomeAgentError, errors.New("unexpected final magic")) + c.String(http.StatusBadGateway, "agent returned unexpected final magic") + return + } + + c.JSON(http.StatusOK, model.FsWriteResult{Size: int64(final.Size), SHA256: hex.EncodeToString(final.SHA256)}) + _ = singleton.DB.Create(&model.MCPAuditLog{ + UserID: entry.UserID, + TokenID: entry.TokenID, + Tool: "fs.upload", + ServerID: entry.ServerID, + Outcome: model.MCPOutcomeOK, + IP: c.GetString(model.CtxKeyRealIPStr), + }).Error +} + +// writeTransferFailureAudit 是 fs.upload / fs.download HTTP handler 失败路径 +// 共用的审计写入。outcome 必须用 model.MCPOutcome* 常量;entry 可以是 nil +// (token consume 阶段就失败时拿不到 entry,UserID/TokenID/ServerID 写 0)。 +// +// Anonymous failures (entry == nil) go through transferAnonAuditThrottleShared +// per-IP sampler so an unauthenticated attacker cannot flood mcp_audit_log +// by POSTing /mcp/upload/. Authenticated failures bypass the +// throttle so SIEM signal stays intact. +func writeTransferFailureAudit(c *gin.Context, entry *transferEntry, tool, outcome string, err error) { + ip := c.GetString(model.CtxKeyRealIPStr) + if entry == nil && !transferAnonAuditThrottleShared.shouldRecord(ip) { + return + } + entryLog := model.MCPAuditLog{ + Tool: tool, + Outcome: outcome, + IP: ip, + } + if entry != nil { + entryLog.UserID = entry.UserID + entryLog.TokenID = entry.TokenID + entryLog.ServerID = entry.ServerID + } + if err != nil { + msg := err.Error() + if len(msg) > 512 { + msg = msg[:512] + } + entryLog.ErrorMsg = msg + entryLog.ErrorCode = outcome + } + mcpAuditWrite(entryLog, nil) +} + +// classifyTransferConsumeError 把 consumeTransferToken 的错误映射成 outcome。 +// 让 SIEM 能区分“伪造/过期 token”与“direction 不匹配”等场景。 +func classifyTransferConsumeError(err error) string { + if err == nil { + return model.MCPOutcomeInternalError + } + msg := err.Error() + switch { + case strings.Contains(msg, "expired"): + return model.MCPOutcomeScopeDenied + case strings.Contains(msg, "direction mismatch"): + return model.MCPOutcomeInvalidArgs + default: + return model.MCPOutcomePermDenied + } +} + +// classifyTransferRevalidateError 把 revalidateTransferEntry 的错误映射成 +// outcome。最重要的一项是“MCP is disabled” → MCPOutcomeMCPDisabled,让 +// 运营在 audit 表里直接看出 kill switch 命中情况,而不是只看到 perm_denied。 +func classifyTransferRevalidateError(err error) string { + if err == nil { + return model.MCPOutcomeInternalError + } + msg := err.Error() + switch { + case strings.Contains(msg, "MCP is disabled"): + return model.MCPOutcomeMCPDisabled + case strings.Contains(msg, "no longer has required scope"): + return model.MCPOutcomeScopeDenied + case strings.Contains(msg, "no longer covers"): + return model.MCPOutcomeScopeDenied + case strings.Contains(msg, "expired"): + return model.MCPOutcomeScopeDenied + default: + return model.MCPOutcomePermDenied + } +} + +// classifyTransferOpenStreamError 把 openFsTransferStream 的失败映射成 +// outcome:offline / 30s attach 超时分别对应 ServerOffline / AgentTimeout。 +func classifyTransferOpenStreamError(err error) string { + if err == nil { + return model.MCPOutcomeInternalError + } + msg := err.Error() + switch { + case strings.Contains(msg, "server offline"): + return model.MCPOutcomeServerOffline + case strings.Contains(msg, "did not attach"): + return model.MCPOutcomeAgentTimeout + default: + return model.MCPOutcomeAgentError + } +} + +// frameReceiver is the frame-preserving subset of grpcx.IOStreamWrapper that +// the download relay needs. We accept the interface (not the concrete type) +// so test simulators can plug in a net.Pipe-backed stream without depending +// on the gRPC stack. +type frameReceiver interface { + RecvFrame() ([]byte, error) +} + +// relayDownloadFrames forwards declared-size payload from agent to HTTP +// client. The agent wraps every data chunk in an NZTC frame (4-byte magic + +// 8-byte big-endian length + payload) so payload that happens to begin with +// the same bytes as a control frame (NZTE / NZTO) cannot be misclassified. +// Control frames (NZTE error, NZTO success) sit on the same IOStream and +// are recognized by their magic; legitimate payload always arrives inside +// NZTC frames and is never matched against the control-frame magics. +// +// Payload is spooled to a per-request tmpfile rather than kept in a 100MiB +// memory buffer: a midstream NZTE must be able to switch the HTTP response +// to 502, which forces us to defer the body write until the final NZTO +// frame is observed; but we MUST NOT pay 100MiB of heap per concurrent +// download to do so. +func relayDownloadFrames(c *gin.Context, stream io.ReadWriteCloser, size int64) error { + spool, err := newTransferSpool() + if err != nil { + c.String(http.StatusInternalServerError, "transfer spool: "+err.Error()) + return err + } + defer spool.Close() + + // Hash the relayed bytes inline; we compare against the trailing + // NZTO declared sha256 in validateDownloadFinal so corrupt or + // truncated agent payloads can't reach the client. + hasher := sha256.New() + streamed := int64(0) + + remaining := size + header := make([]byte, 4+8) + for remaining > 0 { + if _, err := io.ReadFull(stream, header); err != nil { + c.String(http.StatusBadGateway, "stream relay failed: "+err.Error()) + return err + } + if bytes.HasPrefix(header, model.MCPFsXferMagicErr) { + msg := readMidstreamErrMsg(stream, header) + c.String(http.StatusBadGateway, msg) + return errMCPMidstreamAbort + } + if !bytes.HasPrefix(header, model.MCPFsXferMagicChunk) { + c.String(http.StatusBadGateway, "stream relay failed: expected NZTC chunk frame") + return errMCPMidstreamAbort + } + chunkLen := binary.BigEndian.Uint64(header[4:12]) + if chunkLen == 0 { + continue + } + if int64(chunkLen) > remaining { + c.String(http.StatusBadGateway, "agent oversent: more data bytes than declared size") + return errMCPMidstreamAbort + } + n, err := io.CopyN(io.MultiWriter(spool, hasher), stream, int64(chunkLen)) + if err != nil { + c.String(http.StatusBadGateway, "stream relay failed: "+err.Error()) + return err + } + streamed += n + remaining -= n + } + + final := make([]byte, 4+8+32) + if _, err := io.ReadFull(stream, final); err != nil { + c.String(http.StatusBadGateway, "agent did not send final transfer frame: "+err.Error()) + return errMCPMidstreamAbort + } + if bytes.HasPrefix(final, model.MCPFsXferMagicErr) { + msg := readMidstreamErrMsg(stream, final[:4+8]) + c.String(http.StatusBadGateway, msg) + return errMCPMidstreamAbort + } + if err := validateDownloadFinal(final, streamed, hasher.Sum(nil)); err != nil { + c.String(http.StatusBadGateway, err.Error()) + return errMCPMidstreamAbort + } + + if err := spool.Rewind(); err != nil { + c.String(http.StatusInternalServerError, "transfer spool rewind: "+err.Error()) + return err + } + c.Header("Content-Type", "application/octet-stream") + c.Header("Content-Length", strconv.FormatInt(size, 10)) + if _, writeErr := io.Copy(c.Writer, spool); writeErr != nil { + return writeErr + } + return nil +} + +// validateDownloadFinal cross-checks the trailing NZTO frame against the +// payload the dashboard actually relayed: +// - magic must be NZTO (defence-in-depth; the relay already checked). +// - frame must be the full 44 bytes (magic 4 + size 8 + sha256 32). +// - declared size must match streamed byte count exactly. +// - declared sha256 must match the streamed sha256, with one allowed +// "explicit skip" form: all-zero declared hash means agent could not +// compute a hash and we accept the size-only check. +// +// Without this gate a truncated or wrong-hash NZTO is silently accepted +// and the dashboard serves possibly-corrupt bytes to the HTTP client. +func validateDownloadFinal(final []byte, streamedSize int64, streamedSHA256 []byte) error { + if len(final) < 4 || !bytes.Equal(final[:4], model.MCPFsXferMagicOK) { + return errors.New("download final frame: unexpected magic") + } + if len(final) < 4+8+32 { + return errors.New("download final frame: truncated header (need size + sha256)") + } + declaredSize := binary.BigEndian.Uint64(final[4:12]) + if uint64(streamedSize) != declaredSize { + return errors.New("download final frame: declared size does not match streamed bytes") + } + declaredSHA := final[12:44] + allZero := true + for _, b := range declaredSHA { + if b != 0 { + allZero = false + break + } + } + if allZero { + return nil + } + if len(streamedSHA256) < 32 { + return errors.New("download final frame: streamed sha256 too short to compare") + } + if !bytes.Equal(declaredSHA, streamedSHA256[:32]) { + return errors.New("download final frame: declared sha256 does not match streamed bytes") + } + return nil +} + +// midstreamErrMsgCap 限制错误帧 payload 的累计读取量。错误消息只作 string +// 用,没有上限的话恶意/有 bug 的 agent 可在 NZTE 后持续发 256 字节块(且不 +// 关流),让 dashboard goroutine 内存无界增长或永久阻塞。读满 cap 即停止。 +const midstreamErrMsgCap = 8 << 10 + +func readMidstreamErrMsg(stream io.Reader, header []byte) string { + rest := make([]byte, 0, 256) + tail := make([]byte, 256) + for len(rest) < midstreamErrMsgCap { + n, err := stream.Read(tail) + if n > 0 { + room := midstreamErrMsgCap - len(rest) + if n > room { + n = room + } + rest = append(rest, tail[:n]...) + } + if err != nil || n < len(tail) { + break + } + } + return string(header[len(model.MCPFsXferMagicErr):]) + string(rest) +} + +var errMCPMidstreamAbort = errors.New("mcp transfer: aborted mid-stream by agent") + +// fsTransferXferHeader 是 dashboard 解析 NZTU/NZTD/NZTO/NZTE 后得到的统一 +// 结构。Magic 与 model.MCPFsXferMagic* 对照判断帧类型。 +type fsTransferXferHeader struct { + Magic []byte + Size int64 + SHA256 []byte + ErrMsg string +} + +func (h *fsTransferXferHeader) IsErr() bool { + return bytes.Equal(h.Magic, model.MCPFsXferMagicErr) +} + +// readXferFixedHeader 读取一帧 IOStream 数据并解析。每条 agent 控制帧都在 +// 单条 IOStreamData 内完整发送(agent 端用 stream.Send(buf) 整块写),所以 +// 一次 8KiB 缓冲即可拿到完整帧;不需要跨帧拼接。 +// +// 心跳帧(空 Data)由 agent 那侧的 ioStreamKeepAlive 周期性下发,io.Read +// 不会暴露空读,因此这里不必特殊跳过。 +func readXferFixedHeader(stream io.Reader) (*fsTransferXferHeader, error) { + buf := make([]byte, 4+8+32+512) + n, err := stream.Read(buf) + if err != nil { + return nil, err + } + return readXferFixedHeaderFromBytes(buf[:n]) +} + +// readXferFixedHeaderFromBytes parses a fully-received transfer control +// frame. Extracted so a malicious-input regression suite can pin the +// uint64→int64 overflow gate: raw u64 size > MCPFsTransferMaxSize or +// > MaxInt64 must be rejected BEFORE narrowing, otherwise the cap check +// later in the handler sees a wrapped negative value and lets the +// transfer through. +func readXferFixedHeaderFromBytes(raw []byte) (*fsTransferXferHeader, error) { + if len(raw) < 4 { + return nil, errors.New("frame too short") + } + magic := raw[:4] + out := &fsTransferXferHeader{Magic: append([]byte(nil), magic...)} + switch { + case bytes.Equal(magic, model.MCPFsXferMagicErr): + out.ErrMsg = string(raw[4:]) + return out, nil + case bytes.Equal(magic, model.MCPFsXferMagicUploadHdr): + if len(raw) < 4+8 { + return nil, errors.New("upload header too short") + } + size, err := xferSizeFromU64(binary.BigEndian.Uint64(raw[4:12])) + if err != nil { + return nil, err + } + out.Size = size + return out, nil + case bytes.Equal(magic, model.MCPFsXferMagicDownloadHdr): + if len(raw) < 4+8+32 { + return nil, errors.New("download header too short") + } + size, err := xferSizeFromU64(binary.BigEndian.Uint64(raw[4:12])) + if err != nil { + return nil, err + } + out.Size = size + out.SHA256 = append([]byte(nil), raw[12:44]...) + return out, nil + case bytes.Equal(magic, model.MCPFsXferMagicOK): + if len(raw) < 4+8+32 { + return nil, errors.New("ok header too short") + } + size, err := xferSizeFromU64(binary.BigEndian.Uint64(raw[4:12])) + if err != nil { + return nil, err + } + out.Size = size + out.SHA256 = append([]byte(nil), raw[12:44]...) + return out, nil + default: + return nil, errors.New("unexpected frame magic") + } +} + +// xferSizeFromU64 caps the raw u64 size carried by an NZTU/NZTD/NZTO frame +// at MCPFsTransferMaxSize AND math.MaxInt64. Both bounds matter: the cap +// keeps the protocol invariant, the MaxInt64 floor keeps later int64 +// arithmetic safe even if MCPFsTransferMaxSize is ever raised above +// MaxInt64 by accident. +func xferSizeFromU64(raw uint64) (int64, error) { + if raw > uint64(model.MCPFsTransferMaxSize) { + return 0, errors.New("declared size exceeds MCP transfer cap") + } + if raw > math.MaxInt64 { + return 0, errors.New("declared size overflows int64") + } + return int64(raw), nil +} + +// openFsTransferStream 走 IOStream 通道与目标 agent 建立一条专用大文件流。 +// 返回的 io.ReadWriteCloser 既能 Read(接收 agent→dashboard 字节)又能 Write +// (发送 dashboard→agent 字节);调用方通过 readXferFixedHeader 解析控制帧。 +// +// 内部步骤: +// 1. 分配 streamId,CreateStream(streamId, 0, serverID) 在 NezhaHandler 注册 +// 一个 ioStreamContext,targetServerID 用于 agent 侧 stream 归属校验。 +// 2. 通过 server 当前的 RequestTask 流发 TaskTypeFsTransfer,把 streamId + +// req JSON 下发给 agent。 +// 3. 等 agent 通过 IOStream() RPC 完成 magic 引导帧并 AgentConnected。 +// 4. 返回 agent 端流和 cleanup(CloseStream)。 +// +// 任何步骤失败都会 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 + + rpc.NezhaHandlerSingleton.CreateStreamWithPurpose(streamId, 0, serverID, rpc.PurposeMCPTransfer) + cleanup := func() { _ = rpc.NezhaHandlerSingleton.CloseStream(streamId) } + + body, err := json.Marshal(req) + if err != nil { + cleanup() + return nil, func() {}, err + } + // 关闭 entry-check 与 SendTask 之间的 TOCTOU:stream 已注册后再复查一次 + // 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 阶段的全部前置。 +// 这是 mint→consume 之间发生权限变化(PAT 吊销、scope/whitelist 收紧、 +// server 转手)时的兜底闸门:HMAC 签发与 sync.Map 一次性消费机制本身只能 +// 防伪造与防重放,无法感知后端状态。 +func revalidateTransferEntry(e *transferEntry) error { + if singleton.Conf == nil || !singleton.Conf.MCPEnabled() { + return errors.New("MCP is disabled by the dashboard administrator") + } + var tok model.APIToken + if err := singleton.DB.First(&tok, e.TokenID).Error; err != nil { + return errors.New("originating api token no longer exists") + } + if tok.IsExpired(time.Now()) { + return errors.New("originating api token expired") + } + wantScope := model.ScopeServerRead + if e.Direction == transferDirUpload { + wantScope = model.ScopeServerWrite + } + if !tok.HasScope(wantScope) { + return errors.New("originating api token no longer has required scope") + } + if !tok.CanAccessServer(e.ServerID) { + return errors.New("originating api token no longer covers target server") + } + srv, _ := singleton.ServerShared.Get(e.ServerID) + if srv == nil { + return errors.New("target server no longer exists") + } + var user model.User + if err := singleton.DB.First(&user, e.UserID).Error; err != nil { + return errors.New("originating user no longer exists") + } + if user.Role != model.RoleAdmin && srv.GetUserID() != e.UserID { + return errors.New("target server is no longer owned by the originating user") + } + return nil +} diff --git a/cmd/dashboard/controller/mcp_transfer_audit_throttle.go b/cmd/dashboard/controller/mcp_transfer_audit_throttle.go new file mode 100644 index 00000000..66a1895a --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_audit_throttle.go @@ -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/ 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) diff --git a/cmd/dashboard/controller/mcp_transfer_audit_throttle_test.go b/cmd/dashboard/controller/mcp_transfer_audit_throttle_test.go new file mode 100644 index 00000000..2b2f7088 --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_audit_throttle_test.go @@ -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") + } +} diff --git a/cmd/dashboard/controller/mcp_transfer_cancel_test.go b/cmd/dashboard/controller/mcp_transfer_cancel_test.go new file mode 100644 index 00000000..416564a7 --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_cancel_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/mcp_transfer_consume_authz_test.go b/cmd/dashboard/controller/mcp_transfer_consume_authz_test.go new file mode 100644 index 00000000..aab41729 --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_consume_authz_test.go @@ -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)) +} diff --git a/cmd/dashboard/controller/mcp_transfer_correctness_test.go b/cmd/dashboard/controller/mcp_transfer_correctness_test.go new file mode 100644 index 00000000..29976445 --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_correctness_test.go @@ -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 +} + + diff --git a/cmd/dashboard/controller/mcp_transfer_data_frame_collision_test.go b/cmd/dashboard/controller/mcp_transfer_data_frame_collision_test.go new file mode 100644 index 00000000..89b7edfe --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_data_frame_collision_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/mcp_transfer_download_finalcheck_test.go b/cmd/dashboard/controller/mcp_transfer_download_finalcheck_test.go new file mode 100644 index 00000000..d25b4277 --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_download_finalcheck_test.go @@ -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) + } +} diff --git a/cmd/dashboard/controller/mcp_transfer_failure_audit_test.go b/cmd/dashboard/controller/mcp_transfer_failure_audit_test.go new file mode 100644 index 00000000..113a9fa3 --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_failure_audit_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/mcp_transfer_gc_test.go b/cmd/dashboard/controller/mcp_transfer_gc_test.go new file mode 100644 index 00000000..36dfbb84 --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_gc_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/mcp_transfer_no_full_buffer_test.go b/cmd/dashboard/controller/mcp_transfer_no_full_buffer_test.go new file mode 100644 index 00000000..4872fa4e --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_no_full_buffer_test.go @@ -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) + } +} diff --git a/cmd/dashboard/controller/mcp_transfer_path_cap_test.go b/cmd/dashboard/controller/mcp_transfer_path_cap_test.go new file mode 100644 index 00000000..54061b41 --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_path_cap_test.go @@ -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") + } +} diff --git a/cmd/dashboard/controller/mcp_transfer_size_guard_test.go b/cmd/dashboard/controller/mcp_transfer_size_guard_test.go new file mode 100644 index 00000000..05bd2afd --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_size_guard_test.go @@ -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") + } +} diff --git a/cmd/dashboard/controller/mcp_transfer_spool.go b/cmd/dashboard/controller/mcp_transfer_spool.go new file mode 100644 index 00000000..c55aa62a --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_spool.go @@ -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 +} diff --git a/cmd/dashboard/controller/mcp_transfer_upload_args_test.go b/cmd/dashboard/controller/mcp_transfer_upload_args_test.go new file mode 100644 index 00000000..27eb3ae6 --- /dev/null +++ b/cmd/dashboard/controller/mcp_transfer_upload_args_test.go @@ -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/"):] +} diff --git a/cmd/dashboard/controller/oauth2.go b/cmd/dashboard/controller/oauth2.go index 1e6f3ee7..676099b1 100644 --- a/cmd/dashboard/controller/oauth2.go +++ b/cmd/dashboard/controller/oauth2.go @@ -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 diff --git a/cmd/dashboard/controller/oauth2_csrf_test.go b/cmd/dashboard/controller/oauth2_csrf_test.go new file mode 100644 index 00000000..1b9f5f10 --- /dev/null +++ b/cmd/dashboard/controller/oauth2_csrf_test.go @@ -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") + } +} diff --git a/cmd/dashboard/controller/oauth2_test.go b/cmd/dashboard/controller/oauth2_test.go new file mode 100644 index 00000000..b6d332b1 --- /dev/null +++ b/cmd/dashboard/controller/oauth2_test.go @@ -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) +} diff --git a/cmd/dashboard/controller/pat_whitelist_view_test.go b/cmd/dashboard/controller/pat_whitelist_view_test.go new file mode 100644 index 00000000..a823303d --- /dev/null +++ b/cmd/dashboard/controller/pat_whitelist_view_test.go @@ -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)") + } +} diff --git a/cmd/dashboard/controller/permissions.go b/cmd/dashboard/controller/permissions.go index 7c032de9..9fa2c322 100644 --- a/cmd/dashboard/controller/permissions.go +++ b/cmd/dashboard/controller/permissions.go @@ -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 { diff --git a/cmd/dashboard/controller/permissions_cover_fanout_test.go b/cmd/dashboard/controller/permissions_cover_fanout_test.go new file mode 100644 index 00000000..bc0676f2 --- /dev/null +++ b/cmd/dashboard/controller/permissions_cover_fanout_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/rest_scope_test.go b/cmd/dashboard/controller/rest_scope_test.go new file mode 100644 index 00000000..f8364d11 --- /dev/null +++ b/cmd/dashboard/controller/rest_scope_test.go @@ -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)) +} diff --git a/cmd/dashboard/controller/scope_allof_test.go b/cmd/dashboard/controller/scope_allof_test.go new file mode 100644 index 00000000..a8f4bc0c --- /dev/null +++ b/cmd/dashboard/controller/scope_allof_test.go @@ -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") + } + }) +} diff --git a/cmd/dashboard/controller/scope_doc.go b/cmd/dashboard/controller/scope_doc.go new file mode 100644 index 00000000..b59d2bd1 --- /dev/null +++ b/cmd/dashboard/controller/scope_doc.go @@ -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::* 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 diff --git a/cmd/dashboard/controller/scope_doc_consistency_test.go b/cmd/dashboard/controller/scope_doc_consistency_test.go new file mode 100644 index 00000000..94321e07 --- /dev/null +++ b/cmd/dashboard/controller/scope_doc_consistency_test.go @@ -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) + } + } +} diff --git a/cmd/dashboard/controller/server.go b/cmd/dashboard/controller/server.go index a56cb586..14e92801 100644 --- a/cmd/dashboard/controller/server.go +++ b/cmd/dashboard/controller/server.go @@ -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. diff --git a/cmd/dashboard/controller/server_group.go b/cmd/dashboard/controller/server_group.go index 4947aeee..36ee8558 100644 --- a/cmd/dashboard/controller/server_group.go +++ b/cmd/dashboard/controller/server_group.go @@ -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 diff --git a/cmd/dashboard/controller/server_group_update_pat_test.go b/cmd/dashboard/controller/server_group_update_pat_test.go new file mode 100644 index 00000000..9496d976 --- /dev/null +++ b/cmd/dashboard/controller/server_group_update_pat_test.go @@ -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 } +} diff --git a/cmd/dashboard/controller/server_group_visibility_test.go b/cmd/dashboard/controller/server_group_visibility_test.go index be5803f4..b748ceb7 100644 --- a/cmd/dashboard/controller/server_group_visibility_test.go +++ b/cmd/dashboard/controller/server_group_visibility_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/service.go b/cmd/dashboard/controller/service.go index 96574822..a2323573 100644 --- a/cmd/dashboard/controller/service.go +++ b/cmd/dashboard/controller/service.go @@ -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 diff --git a/cmd/dashboard/controller/service_cache_key_test.go b/cmd/dashboard/controller/service_cache_key_test.go new file mode 100644 index 00000000..3ef2af45 --- /dev/null +++ b/cmd/dashboard/controller/service_cache_key_test.go @@ -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) + } +} diff --git a/cmd/dashboard/controller/service_dispatch_pat_test.go b/cmd/dashboard/controller/service_dispatch_pat_test.go new file mode 100644 index 00000000..8f11856d --- /dev/null +++ b/cmd/dashboard/controller/service_dispatch_pat_test.go @@ -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) +} diff --git a/cmd/dashboard/controller/service_list_pat_whitelist_test.go b/cmd/dashboard/controller/service_list_pat_whitelist_test.go new file mode 100644 index 00000000..1c764ccb --- /dev/null +++ b/cmd/dashboard/controller/service_list_pat_whitelist_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/service_skip_enabled_only_test.go b/cmd/dashboard/controller/service_skip_enabled_only_test.go new file mode 100644 index 00000000..f3f63716 --- /dev/null +++ b/cmd/dashboard/controller/service_skip_enabled_only_test.go @@ -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") + } +} diff --git a/cmd/dashboard/controller/service_visibility_test.go b/cmd/dashboard/controller/service_visibility_test.go index 729b2d12..5ee44433 100644 --- a/cmd/dashboard/controller/service_visibility_test.go +++ b/cmd/dashboard/controller/service_visibility_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/setting.go b/cmd/dashboard/controller/setting.go index 4fd89061..1c0a270e 100644 --- a/cmd/dashboard/controller/setting.go +++ b/cmd/dashboard/controller/setting.go @@ -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 diff --git a/cmd/dashboard/controller/setting_enable_mcp_save_failure_test.go b/cmd/dashboard/controller/setting_enable_mcp_save_failure_test.go new file mode 100644 index 00000000..f296d46f --- /dev/null +++ b/cmd/dashboard/controller/setting_enable_mcp_save_failure_test.go @@ -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") + } +} diff --git a/cmd/dashboard/controller/setting_enable_mcp_test.go b/cmd/dashboard/controller/setting_enable_mcp_test.go new file mode 100644 index 00000000..7d4ada12 --- /dev/null +++ b/cmd/dashboard/controller/setting_enable_mcp_test.go @@ -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 +} diff --git a/cmd/dashboard/controller/stream_pat_authz_test.go b/cmd/dashboard/controller/stream_pat_authz_test.go new file mode 100644 index 00000000..86488e37 --- /dev/null +++ b/cmd/dashboard/controller/stream_pat_authz_test.go @@ -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") + } +} diff --git a/cmd/dashboard/controller/tenant_isolation_test.go b/cmd/dashboard/controller/tenant_isolation_test.go new file mode 100644 index 00000000..24b92959 --- /dev/null +++ b/cmd/dashboard/controller/tenant_isolation_test.go @@ -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") +} diff --git a/cmd/dashboard/controller/terminal.go b/cmd/dashboard/controller/terminal.go index 521795db..3d4bdf80 100644 --- a/cmd/dashboard/controller/terminal.go +++ b/cmd/dashboard/controller/terminal.go @@ -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 { diff --git a/cmd/dashboard/controller/transfer.go b/cmd/dashboard/controller/transfer.go index b3fbc6e5..0e8f3ca7 100644 --- a/cmd/dashboard/controller/transfer.go +++ b/cmd/dashboard/controller/transfer.go @@ -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) diff --git a/cmd/dashboard/controller/transfer_cancel_authz_test.go b/cmd/dashboard/controller/transfer_cancel_authz_test.go new file mode 100644 index 00000000..d3d52e9b --- /dev/null +++ b/cmd/dashboard/controller/transfer_cancel_authz_test.go @@ -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) +} diff --git a/cmd/dashboard/controller/transfer_pat_whitelist_test.go b/cmd/dashboard/controller/transfer_pat_whitelist_test.go new file mode 100644 index 00000000..725d1a51 --- /dev/null +++ b/cmd/dashboard/controller/transfer_pat_whitelist_test.go @@ -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)") +} diff --git a/cmd/dashboard/controller/transfer_retry_pat_whitelist_test.go b/cmd/dashboard/controller/transfer_retry_pat_whitelist_test.go new file mode 100644 index 00000000..c3d7c287 --- /dev/null +++ b/cmd/dashboard/controller/transfer_retry_pat_whitelist_test.go @@ -0,0 +1,78 @@ +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" +) + +// retryServerTransfer 必须像 cancel/list 一样受 PAT 的 server_ids 白名单收窄: +// admin 给自己签的 PAT 加上 ServerIDs={1} 之后,不能再用它 retry server 2 的 +// 历史 transfer 行。否则白名单只在 read/cancel 上生效,retry 路径仍是“admin +// 早返回 → 完全绕过白名单”,与 model.ServerTransfer.HasPermission 注释里 +// “PAT server_ids whitelist is evaluated FIRST, before the admin short- +// circuit” 直接冲突。 +func callRetryServerTransferWithAdminPAT(t *testing.T, transferID, callerID uint64, tok *model.APIToken) (commonResponseShape, int) { + 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.POST("/transfer/:id/retry", commonHandler(retryServerTransfer)) + + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, + "/transfer/"+strconv.FormatUint(transferID, 10)+"/retry", + bytes.NewReader(nil)) + r.ServeHTTP(w, req) + + var resp commonResponseShape + assert.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) + return resp, w.Code +} + +func TestRetryServerTransfer_AdminPATIsNarrowedByServerWhitelist(t *testing.T) { + cleanup := setupRetryServerTransferFixture(t) + defer cleanup() + + seedServer(t, 1, 300) + seedServer(t, 2, 300) + insideID := seedFailedTransfer(t, 1, 100, 200, 100) + outsideID := seedFailedTransfer(t, 2, 100, 200, 100) + + tok := &model.APIToken{ID: 18, UserID: 999} + tok.SetServerIDs([]uint64{1}) + + respOutside, statusOutside := callRetryServerTransferWithAdminPAT(t, outsideID, 999, tok) + assert.Equal(t, http.StatusOK, statusOutside) + assert.False(t, respOutside.Success, + "admin PAT scoped to {1} must NOT retry transfer of server 2 (admin early-return is no longer a bypass)") + assert.Contains(t, respOutside.Error, "permission denied") + + var count int64 + assert.NoError(t, singleton.DB.Model(&model.ServerTransfer{}). + Where("status = ?", model.ServerTransferStatusPending). + Count(&count).Error) + assert.Equal(t, int64(0), count, + "rejected retry must not create a Pending row") + + respInside, statusInside := callRetryServerTransferWithAdminPAT(t, insideID, 999, tok) + assert.Equal(t, http.StatusOK, statusInside) + assert.True(t, respInside.Success, + "admin PAT scoped to {1} must still be able to retry transfer of server 1: %s", + respInside.Error) +} diff --git a/cmd/dashboard/controller/trigger_task_pat_scope_test.go b/cmd/dashboard/controller/trigger_task_pat_scope_test.go new file mode 100644 index 00000000..9173afbf --- /dev/null +++ b/cmd/dashboard/controller/trigger_task_pat_scope_test.go @@ -0,0 +1,99 @@ +package controller + +import ( + "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 newTriggerTaskCtxWithPAT(viewer *model.User, tok *model.APIToken) *gin.Context { + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/service", http.NoBody) + if viewer != nil { + c.Set(model.CtxKeyAuthorizedUser, viewer) + } + if tok != nil { + c.Set(model.CtxKeyAPIToken, tok) + c.Set(apiTokenCtxKey, tok) + } + return c +} + +// 注册一个用户 1 拥有的触发任务,使 CronShared.CheckPermission 通过, +// 从而隔离出「PAT 缺少 cron:exec」这一条裁决路径。 +func registerOwnerTriggerTask(t *testing.T, id uint64) { + t.Helper() + singleton.CronShared.Update(&model.Cron{ + Common: model.Common{ID: id, UserID: 1}, + Name: "trigger", + TaskType: model.CronTypeTriggerTask, + Cover: model.CronCoverAll, + }) +} + +func TestValidateServersPATTriggerTaskRequiresCronExec(t *testing.T) { + setupAlertRuleFanoutFixture(t) + registerOwnerTriggerTask(t, 100) + + viewer := &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin} + + noExec := &model.APIToken{ID: 10, UserID: 1} + noExec.SetScopes([]string{model.ScopeServiceWrite}) + svc := &model.Service{ + Common: model.Common{UserID: 1}, + EnableTriggerTask: true, + FailTriggerTasks: []uint64{100}, + } + require.Error(t, validateServers(newTriggerTaskCtxWithPAT(viewer, noExec), svc), + "service:write PAT must not bind a trigger task without cron:exec") + + withExec := &model.APIToken{ID: 11, UserID: 1} + withExec.SetScopes([]string{model.ScopeServiceWrite, model.ScopeCronExec}) + require.NoError(t, validateServers(newTriggerTaskCtxWithPAT(viewer, withExec), svc), + "service:write + cron:exec PAT may bind a trigger task") +} + +func TestValidateRulePATTriggerTaskRequiresCronExec(t *testing.T) { + setupAlertRuleFanoutFixture(t) + registerOwnerTriggerTask(t, 200) + + viewer := &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin} + + noExec := &model.APIToken{ID: 12, UserID: 1} + noExec.SetScopes([]string{model.ScopeAlertRuleWrite}) + rule := &model.AlertRule{ + Common: model.Common{UserID: 1}, + Name: "r", + Rules: []*model.Rule{{Type: "offline", Cover: model.RuleCoverAll, Duration: 10, Ignore: map[uint64]bool{}}}, + RecoverTriggerTasks: []uint64{200}, + } + require.Error(t, validateRule(newTriggerTaskCtxWithPAT(viewer, noExec), rule), + "alertrule:write PAT must not bind a trigger task without cron:exec") + + withExec := &model.APIToken{ID: 13, UserID: 1} + withExec.SetScopes([]string{model.ScopeAlertRuleWrite, model.ScopeCronExec}) + require.NoError(t, validateRule(newTriggerTaskCtxWithPAT(viewer, withExec), rule), + "alertrule:write + cron:exec PAT may bind a trigger task") +} + +// JWT 调用者(无 PAT)不受 cron:exec 收口影响。 +func TestValidateServersJWTUnaffectedByTriggerTaskScope(t *testing.T) { + setupAlertRuleFanoutFixture(t) + registerOwnerTriggerTask(t, 300) + + viewer := &model.User{Common: model.Common{ID: 1}, Role: model.RoleAdmin} + svc := &model.Service{ + Common: model.Common{UserID: 1}, + EnableTriggerTask: true, + FailTriggerTasks: []uint64{300}, + } + require.NoError(t, validateServers(newTriggerTaskCtxWithPAT(viewer, nil), svc), + "JWT caller must not be gated by cron:exec") +} diff --git a/cmd/dashboard/controller/ws.go b/cmd/dashboard/controller/ws.go index ffcc15d8..89aacac4 100644 --- a/cmd/dashboard/controller/ws.go +++ b/cmd/dashboard/controller/ws.go @@ -5,6 +5,9 @@ import ( "net" "net/http" "net/url" + "slices" + "strconv" + "strings" "time" "unicode/utf8" @@ -115,6 +118,9 @@ func serverStream(c *gin.Context) (any, error) { } defer conn.Close() + deregisterPAT := registerPATConnection(c, func() { _ = conn.Close() }) + defer deregisterPAT() + userIp := c.GetString(model.CtxKeyRealIPStr) if userIp == "" { userIp = c.RemoteIP() @@ -130,6 +136,7 @@ func serverStream(c *gin.Context) (any, error) { userId = user.ID isAdmin = user.Role.IsAdmin() } + patAccessor, patCacheKey := patStreamContext(c) singleton.AddOnlineUser(connId, &model.OnlineUser{ UserID: userId, @@ -141,7 +148,7 @@ func serverStream(c *gin.Context) (any, error) { count := 0 for { - stat, err := getServerStat(count == 0, userId, isAdmin) + stat, err := getServerStat(count == 0, userId, isAdmin, patAccessor, patCacheKey) if err != nil { continue } @@ -167,12 +174,15 @@ var requestGroup singleflight.Group // depends on per-server ownership: prior to GHSA-hvv7-hfrh-7gxj this function // used a single isMember flag and leaked HideForGuest servers plus full Host // (PlatformVersion, agent Version, GPU) to every authenticated user. -func getServerStat(withPublicNote bool, viewerUserID uint64, viewerIsAdmin bool) ([]byte, error) { - cacheKey := fmt.Sprintf("serverStats::%t::%t::%d", withPublicNote, viewerIsAdmin, viewerUserID) +// +// patCacheKey distinguishes PATs with disjoint server_ids whitelists so two +// limited tokens for the same user do not share a singleflight projection. +func getServerStat(withPublicNote bool, viewerUserID uint64, viewerIsAdmin bool, pat model.APITokenAccessor, patCacheKey string) ([]byte, error) { + cacheKey := fmt.Sprintf("serverStats::%t::%t::%d::%s", withPublicNote, viewerIsAdmin, viewerUserID, patCacheKey) v, err, _ := requestGroup.Do(cacheKey, func() (any, error) { servers := filterServersForViewer( singleton.ServerShared.GetSortedList(), - viewerUserID, viewerIsAdmin, withPublicNote, + viewerUserID, viewerIsAdmin, withPublicNote, pat, ) return json.Marshal(model.StreamServerData{ Now: time.Now().Unix() * 1000, @@ -184,17 +194,40 @@ func getServerStat(withPublicNote bool, viewerUserID uint64, viewerIsAdmin bool) return v.([]byte), err } +// patStreamContext extracts the PAT accessor + a deterministic cache key +// fragment for the singleflight projection. Returns (nil, "jwt") for JWT +// requests so two callers from the same user collapse onto one frame. +func patStreamContext(c *gin.Context) (model.APITokenAccessor, string) { + tok := APITokenFromContext(c) + if tok == nil { + return nil, "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 tok, fmt.Sprintf("pat:%d:%s", tok.ID, strings.Join(parts, ",")) +} + // filterServersForViewer projects the global server list down to what a single // viewer is allowed to see. The rules are: // - HideForGuest servers are visible only to their owner and to admins. // - Non-owner / non-admin viewers (including authenticated members) get // Host.Filter() output, which drops PlatformVersion and agent Version. // - Admins are unconstrained. +// - A non-nil pat whitelist narrows visibility further; servers outside its +// allow-list are dropped even from admins/owners (a PAT scoped to a +// subset must never widen via its caller's role). // // viewerUserID == 0 represents an unauthenticated guest. -func filterServersForViewer(servers []*model.Server, viewerUserID uint64, viewerIsAdmin bool, withPublicNote bool) []model.StreamServer { +func filterServersForViewer(servers []*model.Server, viewerUserID uint64, viewerIsAdmin bool, withPublicNote bool, pat model.APITokenAccessor) []model.StreamServer { out := make([]model.StreamServer, 0, len(servers)) for _, server := range servers { + if pat != nil && !pat.CanAccessServer(server.ID) { + continue + } isOwnerOrAdmin := viewerIsAdmin || (viewerUserID != 0 && server.GetUserID() == viewerUserID) if server.HideForGuest && !isOwnerOrAdmin { continue diff --git a/cmd/dashboard/controller/ws_stream_visibility_test.go b/cmd/dashboard/controller/ws_stream_visibility_test.go index 9c006582..61a2745d 100644 --- a/cmd/dashboard/controller/ws_stream_visibility_test.go +++ b/cmd/dashboard/controller/ws_stream_visibility_test.go @@ -78,7 +78,7 @@ func findStreamServer(out []model.StreamServer, id uint64) *model.StreamServer { // Guest: no auth → skip every HideForGuest server, Host.Filter() drops // PlatformVersion and agent Version while keeping the rest (including GPU). func TestFilterServersForViewerGuestHidesPrivateAndRedactsHost(t *testing.T) { - out := filterServersForViewer(makeStreamTestServers(), 0, false, true) + out := filterServersForViewer(makeStreamTestServers(), 0, false, true, nil) assert.Len(t, out, 2) assert.Nil(t, findStreamServer(out, 2), "alice-hidden should be invisible to guests") @@ -98,7 +98,7 @@ func TestFilterServersForViewerGuestHidesPrivateAndRedactsHost(t *testing.T) { func TestFilterServersForViewerNonOwnerMemberMatchesGuest(t *testing.T) { servers := makeStreamTestServers() carolID := uint64(300) - out := filterServersForViewer(servers, carolID, false, true) + out := filterServersForViewer(servers, carolID, false, true, nil) assert.Len(t, out, 2) assert.Nil(t, findStreamServer(out, 2)) @@ -116,7 +116,7 @@ func TestFilterServersForViewerNonOwnerMemberMatchesGuest(t *testing.T) { func TestFilterServersForViewerOwnerSeesOwnHiddenAndFullHost(t *testing.T) { servers := makeStreamTestServers() aliceID := uint64(100) - out := filterServersForViewer(servers, aliceID, false, true) + out := filterServersForViewer(servers, aliceID, false, true, nil) assert.Len(t, out, 3, "alice sees her 2 servers + bob's 1 public server") assert.Nil(t, findStreamServer(out, 4), "alice must not see bob's hidden server") @@ -137,7 +137,7 @@ func TestFilterServersForViewerOwnerSeesOwnHiddenAndFullHost(t *testing.T) { // Admin: no restrictions — sees every server with full Host, regardless of owner or HideForGuest. func TestFilterServersForViewerAdminSeesAllWithFullHost(t *testing.T) { servers := makeStreamTestServers() - out := filterServersForViewer(servers, 999, true, true) + out := filterServersForViewer(servers, 999, true, true, nil) assert.Len(t, out, 4) for _, s := range out { @@ -149,9 +149,57 @@ func TestFilterServersForViewerAdminSeesAllWithFullHost(t *testing.T) { // First-tick frame includes PublicNote, subsequent frames omit it. // This must hold regardless of viewer. func TestFilterServersForViewerWithoutPublicNoteFlagOmitsNote(t *testing.T) { - out := filterServersForViewer(makeStreamTestServers(), 0, false, false) + out := filterServersForViewer(makeStreamTestServers(), 0, false, false, nil) for _, s := range out { assert.Empty(t, s.PublicNote, "follow-up frames must not include PublicNote") } } + +// patAllowList implements model.APITokenAccessor: empty = unrestricted. +type patAllowList []uint64 + +func (p patAllowList) CanAccessServer(id uint64) bool { + if len(p) == 0 { + return true + } + for _, allowed := range p { + if allowed == id { + return true + } + } + return false +} + +// ServerIDs lets model.DenyListSafeForLimitedPAT see the same "empty = +// unrestricted" convention CanAccessServer encodes. +func (p patAllowList) ServerIDs() []uint64 { + return []uint64(p) +} + +// PAT server_ids whitelist must narrow ws/server visibility even for admins; +// otherwise an admin-issued limited PAT still leaks every server's state. +func TestFilterServersForViewerPATWhitelistNarrowsAdmin(t *testing.T) { + out := filterServersForViewer(makeStreamTestServers(), 999, true, true, patAllowList{3}) + + if assert.Len(t, out, 1, "admin PAT scoped to server_ids=[3] must only see server 3") { + assert.Equal(t, uint64(3), out[0].ID) + } + assert.Nil(t, findStreamServer(out, 1), "admin PAT must not see server 1 outside whitelist") + assert.Nil(t, findStreamServer(out, 2), "admin PAT must not see server 2 outside whitelist") + assert.Nil(t, findStreamServer(out, 4), "admin PAT must not see server 4 outside whitelist") +} + +func TestFilterServersForViewerPATWhitelistNarrowsOwner(t *testing.T) { + out := filterServersForViewer(makeStreamTestServers(), 100, false, true, patAllowList{2}) + + if assert.Len(t, out, 1, "owner PAT scoped to {2} must only see server 2") { + assert.Equal(t, uint64(2), out[0].ID) + } + assert.Nil(t, findStreamServer(out, 1), "owner PAT must not see her own server 1 outside whitelist") +} + +func TestFilterServersForViewerNilPATKeepsLegacyVisibility(t *testing.T) { + withNil := filterServersForViewer(makeStreamTestServers(), 999, true, true, nil) + assert.Len(t, withNil, 4, "no PAT must keep admin-wide visibility") +} diff --git a/cmd/dashboard/main.go b/cmd/dashboard/main.go index 5ab636b8..3f389161 100644 --- a/cmd/dashboard/main.go +++ b/cmd/dashboard/main.go @@ -101,6 +101,13 @@ func initIDCodec() error { // @securityDefinitions.apikey BearerAuth // @in header // @name Authorization +// @description JWT session token. Browser/UI flow. Format: `Bearer ` or cookie `nz-jwt`. + +// @securityDefinitions.apikey APITokenAuth +// @in header +// @name Authorization +// @description Personal Access Token (PAT). Programmatic/CI/LLM flow. Format: `Bearer nzp_`. +// @description Each endpoint enforces a specific scope; see the `controller` package godoc for the authoritative scope table. // @externalDocs.description OpenAPI // @externalDocs.url https://swagger.io/resources/open-api/ @@ -140,6 +147,9 @@ func main() { singleton.CleanMonitorHistory() rpc.DispatchKeepalive() + rpc.SetMCPKillSwitchObserver(func() bool { + return singleton.Conf == nil || !singleton.Conf.MCPEnabled() + }) go rpc.DispatchTask(serviceSentinelDispatchBus) go singleton.AlertSentinelStart() diff --git a/cmd/dashboard/rpc/rpc.go b/cmd/dashboard/rpc/rpc.go index bd8e73cf..0057bae2 100644 --- a/cmd/dashboard/rpc/rpc.go +++ b/cmd/dashboard/rpc/rpc.go @@ -2,6 +2,7 @@ package rpc import ( "context" + "errors" "fmt" "log" "net/http" @@ -21,6 +22,13 @@ import ( "github.com/nezhahq/nezha/service/singleton" ) +// SetMCPKillSwitchObserver re-exports the service/rpc hook so cmd/dashboard +// can wire singleton.Conf.EnableMCP without importing the inner rpc package +// (cmd/dashboard already imports cmd/dashboard/rpc for ServeRPC). +func SetMCPKillSwitchObserver(fn func() bool) { + rpcService.SetMCPKillSwitchObserver(fn) +} + func ServeRPC() *grpc.Server { server := grpc.NewServer(grpc.ChainUnaryInterceptor(getRealIp, waf)) rpcService.NezhaHandlerSingleton = rpcService.NewNezhaHandler() @@ -96,27 +104,30 @@ func DispatchTask(serviceSentinelDispatchBus <-chan *model.Service) { if server == nil { continue } - stream := server.GetTaskStream() - if stream == nil { + if !canSendTaskToServer(task, server) { continue } - - if canSendTaskToServer(task, server) { - stream.Send(task.PB()) + // SendTask 走 holder-scoped send mutex,避免与 cron / + // server-transfer / MCP CallAgent / fs.transfer 等并发 + // SendMsg 同一 RequestTask stream。 + 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.Range { + // 快照后逐个 SendTask,不在 ServerShared 的 listMu.RLock 内做阻塞 + // gRPC:否则一个卡死 agent 会拖死需要写锁的 server 生命周期操作。 + for id, server := range singleton.ServerShared.GetList() { if server == nil || task.SkipServers[id] { continue } - stream := server.GetTaskStream() - if stream == nil { + if !canSendTaskToServer(task, server) { continue } - - if canSendTaskToServer(task, server) { - stream.Send(task.PB()) + if err := server.SendTask(task.PB()); err != nil && + !errors.Is(err, model.ErrTaskStreamOffline) { + log.Printf("NEZHA>> DispatchTask send error (server=%d): %v", id, err) } } } @@ -130,11 +141,10 @@ func DispatchKeepalive() { if s == nil { continue } - stream := s.GetTaskStream() - if stream == 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) } - stream.Send(&proto.Task{Type: model.TaskTypeKeepalive}) } }) } @@ -146,8 +156,7 @@ func ServeNAT(w http.ResponseWriter, r *http.Request, natConfig *model.NAT) { w.Write([]byte("server not found or not connected")) return } - stream := server.GetTaskStream() - if stream == nil { + if server.GetTaskStream() == nil { w.WriteHeader(http.StatusServiceUnavailable) w.Write([]byte("server not found or not connected")) return @@ -179,7 +188,7 @@ func ServeNAT(w http.ResponseWriter, r *http.Request, natConfig *model.NAT) { return } - if err := stream.Send(&proto.Task{ + if err := server.SendTask(&proto.Task{ Type: model.TaskTypeNAT, Data: string(taskData), }); err != nil { diff --git a/go.mod b/go.mod index 2cb3027e..a6fd6f67 100644 --- a/go.mod +++ b/go.mod @@ -22,11 +22,13 @@ require ( github.com/libdns/he v1.2.2 github.com/libdns/libdns v1.1.1 github.com/miekg/dns v1.1.72 + github.com/modelcontextprotocol/go-sdk v1.6.1 github.com/nezhahq/libdns-tencentcloud v0.0.0-20250501081622-bd293105845a github.com/ory/graceful v0.2.0 github.com/oschwald/maxminddb-golang v1.13.1 github.com/patrickmn/go-cache v2.1.0+incompatible github.com/robfig/cron/v3 v3.0.1 + github.com/sqids/sqids-go v0.4.1 github.com/stretchr/testify v1.11.1 github.com/swaggo/files v1.0.1 github.com/swaggo/gin-swagger v1.6.1 @@ -34,6 +36,7 @@ require ( github.com/tidwall/gjson v1.19.0 golang.org/x/crypto v0.52.0 golang.org/x/exp v0.0.0-20260508232706-74f9aab9d74a + golang.org/x/mod v0.36.0 golang.org/x/net v0.55.0 golang.org/x/oauth2 v0.36.0 golang.org/x/sync v0.20.0 @@ -75,6 +78,7 @@ require ( github.com/goccy/go-yaml v1.19.2 // indirect github.com/golang-jwt/jwt/v4 v4.5.2 // indirect github.com/golang/snappy v1.0.0 // indirect + github.com/google/jsonschema-go v0.4.3 // indirect github.com/jinzhu/inflection v1.0.0 // indirect github.com/jinzhu/now v1.1.5 // indirect github.com/json-iterator/go v1.1.12 // indirect @@ -93,7 +97,8 @@ require ( github.com/quic-go/qpack v0.6.0 // indirect github.com/quic-go/quic-go v0.59.1 // indirect github.com/rogpeppe/go-internal v1.14.1 // indirect - github.com/sqids/sqids-go v0.4.1 // indirect + github.com/segmentio/asm v1.1.3 // indirect + github.com/segmentio/encoding v0.5.4 // indirect github.com/tidwall/match v1.2.0 // indirect github.com/tidwall/pretty v1.2.1 // indirect github.com/tidwall/sjson v1.2.5 // indirect @@ -104,12 +109,12 @@ require ( github.com/valyala/gozstd v1.24.0 // indirect github.com/valyala/histogram v1.2.0 // indirect github.com/valyala/quicktemplate v1.8.0 // indirect + github.com/yosida95/uritemplate/v3 v3.0.2 // indirect github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 // indirect go.mongodb.org/mongo-driver/v2 v2.6.0 // indirect go.yaml.in/yaml/v2 v2.4.4 // indirect go.yaml.in/yaml/v3 v3.0.4 // indirect golang.org/x/arch v0.27.0 // indirect - golang.org/x/mod v0.36.0 // indirect golang.org/x/sys v0.45.0 // indirect golang.org/x/text v0.37.0 // indirect golang.org/x/time v0.15.0 // indirect diff --git a/go.sum b/go.sum index 6ee6145d..352e4475 100644 --- a/go.sum +++ b/go.sum @@ -91,6 +91,8 @@ github.com/goccy/go-yaml v1.19.2 h1:PmFC1S6h8ljIz6gMRBopkjP1TVT7xuwrButHID66PoM= github.com/goccy/go-yaml v1.19.2/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA= github.com/golang-jwt/jwt/v4 v4.5.2 h1:YtQM7lnr8iZ+j5q71MGKkNw9Mn7AjHM68uc9g5fXeUI= github.com/golang-jwt/jwt/v4 v4.5.2/go.mod h1:m21LjoU+eqJr34lmDMbreY2eSTRJ1cv77w39/MY0Ch0= +github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY= +github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= github.com/golang/snappy v1.0.0 h1:Oy607GVXHs7RtbggtPBnr2RmDArIsAefDwvrdWvRhGs= @@ -98,6 +100,8 @@ github.com/golang/snappy v1.0.0/go.mod h1:/XxbfmMg8lxefKM7IXC3fBNl/7bRcc72aCRzEW github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg= +github.com/google/jsonschema-go v0.4.3 h1:/DBOLZTfDow7pe2GmaJNhltueGTtDKICi8V8p+DQPd0= +github.com/google/jsonschema-go v0.4.3/go.mod h1:r5quNTdLOYEz95Ru18zA0ydNbBuYoo9tgaYcxEYhJVE= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= @@ -148,6 +152,8 @@ github.com/mitchellh/copystructure v1.2.0 h1:vpKXTN4ewci03Vljg/q9QvCGUDttBOGBIa1 github.com/mitchellh/copystructure v1.2.0/go.mod h1:qLl+cE2AmVv+CoeAwDPye/v+N2HKCj9FbZEVFJRxO9s= github.com/mitchellh/reflectwalk v1.0.2 h1:G2LzWKi524PWgd3mLHV8Y5k7s6XUvT0Gef6zxSIeXaQ= github.com/mitchellh/reflectwalk v1.0.2/go.mod h1:mSTlrgnPZtwu0c4WaC2kGObEpuNDbx0jmZXqmk4esnw= +github.com/modelcontextprotocol/go-sdk v1.6.1 h1:0zOSupjKUxPKSocPT1Wtago+mUHU2/uZ4xSOY0FGReU= +github.com/modelcontextprotocol/go-sdk v1.6.1/go.mod h1:kzm3kzFL1/+AziGOE0nUs3gvPoNxMCvkxokMkuFapXQ= github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd h1:TRLaZ9cD/w8PVh93nsPXa1VrQ6jlwL5oN8l14QlcNfg= github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= @@ -177,6 +183,10 @@ github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs= github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro= github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= +github.com/segmentio/asm v1.1.3 h1:WM03sfUOENvvKexOLp+pCqgb/WDjsi7EK8gIsICtzhc= +github.com/segmentio/asm v1.1.3/go.mod h1:Ld3L4ZXGNcSLRg4JBsZ3//1+f/TjYl0Mzen/DQy1EJg= +github.com/segmentio/encoding v0.5.4 h1:OW1VRern8Nw6ITAtwSZ7Idrl3MXCFwXHPgqESYfvNt0= +github.com/segmentio/encoding v0.5.4/go.mod h1:HS1ZKa3kSN32ZHVZ7ZLPLXWvOVIiZtyJnO1gPH1sKt0= github.com/sqids/sqids-go v0.4.1 h1:eQKYzmAZbLlRwHeHYPF35QhgxwZHLnlmVj9AkIj/rrw= github.com/sqids/sqids-go v0.4.1/go.mod h1:EMwHuPQgSNFS0A49jESTfIQS+066XQTVhukrzEPScl8= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= @@ -221,6 +231,8 @@ github.com/valyala/histogram v1.2.0 h1:wyYGAZZt3CpwUiIb9AU/Zbllg1llXyrtApRS815OL github.com/valyala/histogram v1.2.0/go.mod h1:Hb4kBwb4UxsaNbbbh+RRz8ZR6pdodR57tzWUS3BUzXY= github.com/valyala/quicktemplate v1.8.0 h1:zU0tjbIqTRgKQzFY1L42zq0qR3eh4WoQQdIdqCysW5k= github.com/valyala/quicktemplate v1.8.0/go.mod h1:qIqW8/igXt8fdrUln5kOSb+KWMaJ4Y8QUsfd1k6L2jM= +github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zIM+UJPGz4= +github.com/yosida95/uritemplate/v3 v3.0.2/go.mod h1:ILOh0sOhIJR3+L/8afwt/kE++YT040gmv5BQTMR2HP4= github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78 h1:ilQV1hzziu+LLM3zUTJ0trRztfwgjqKnBWNtSRkbmwM= github.com/youmark/pkcs8 v0.0.0-20240726163527-a2c0da244d78/go.mod h1:aL8wCCfTfSfmXjznFBSZNN13rSJjlIOI1fUNAtF7rmI= github.com/yuin/goldmark v1.4.13/go.mod h1:6yULJ656Px+3vBD8DxQVa3kxgyrAnzto9xy5taEt/CY= diff --git a/model/alertrule.go b/model/alertrule.go index f347044b..82bf2bbd 100644 --- a/model/alertrule.go +++ b/model/alertrule.go @@ -3,6 +3,7 @@ package model import ( "slices" + "github.com/gin-gonic/gin" "github.com/goccy/go-json" "gorm.io/gorm" ) @@ -63,6 +64,62 @@ func (r *AlertRule) Enabled() bool { return r.Enable != nil && *r.Enable } +// HasPermission extends the default owner/admin check with PAT +// server_ids whitelist enforcement. AlertRule.Snapshot fans out across +// every owner-visible server filtered only by Rule.Ignore semantics +// (RuleCoverAll: deny-list; RuleCoverIgnoreAll: allow-list). A +// server-limited PAT must therefore satisfy the same cover-fanout rule +// the cron / service paths use — otherwise it can create or update a +// rule that monitors servers outside its whitelist (admin owner: any +// server in the system). +// +// Unknown Rule.Cover is fail-closed: Snapshot's switch defaults to +// "monitor everything", so persisting it would defeat the PAT cover +// guard. createAlertRule / updateAlertRule should also reject unknown +// covers at write time; this method is the runtime safety net. +func (r *AlertRule) HasPermission(ctx *gin.Context) bool { + if !r.Common.HasPermission(ctx) { + return false + } + v, ok := ctx.Get(CtxKeyAPIToken) + if !ok { + return true + } + tok, _ := v.(APITokenAccessor) + if tok == nil { + return true + } + if wl, ok := tok.(APITokenWhitelistView); ok && len(wl.ServerIDs()) == 0 { + return true + } + for _, rule := range r.Rules { + if rule == nil { + continue + } + switch rule.Cover { + case RuleCoverAll: + denyIDs := make([]uint64, 0, len(rule.Ignore)) + for id, ignored := range rule.Ignore { + if ignored { + denyIDs = append(denyIDs, id) + } + } + if !DenyListSafeForLimitedPAT(tok, r.GetUserID(), denyIDs) { + return false + } + case RuleCoverIgnoreAll: + for id, monitored := range rule.Ignore { + if monitored && !tok.CanAccessServer(id) { + return false + } + } + default: + return false + } + } + return true +} + // Snapshot 对传入的Server进行该报警规则下所有type的检查 返回每项检查结果 func (r *AlertRule) Snapshot(cycleTransferStats *CycleTransferStats, server *Server, db *gorm.DB) []bool { point := make([]bool, len(r.Rules)) diff --git a/model/alertrule_pat_whitelist_test.go b/model/alertrule_pat_whitelist_test.go new file mode 100644 index 00000000..49fefd6c --- /dev/null +++ b/model/alertrule_pat_whitelist_test.go @@ -0,0 +1,227 @@ +package model + +import ( + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" +) + +// H4 regression: AlertRule had no AlertRule.HasPermission override, so +// limited PATs could create / list / update rules that fan out to every +// owner server. A RuleCoverAll + empty Ignore rule monitors every server +// the owner can reach (admin owner: every server in the system), which is +// exactly the cover-fanout the PAT whitelist is supposed to contain. +func TestAlertRuleHasPermission_DeniesRuleCoverAllEmptyIgnoreForLimitedPAT(t *testing.T) { + saved := OwnerServerIDsLookup + savedAdmin := OwnerIsAdminLookup + savedAll := AllServerIDsLookup + t.Cleanup(func() { + OwnerServerIDsLookup = saved + OwnerIsAdminLookup = savedAdmin + AllServerIDsLookup = savedAll + }) + OwnerServerIDsLookup = func(uid uint64) []uint64 { + if uid == 100 { + return []uint64{1, 2} + } + return nil + } + OwnerIsAdminLookup = func(uid uint64) bool { return false } + AllServerIDsLookup = func() []uint64 { return []uint64{1, 2, 3} } + + rule := &AlertRule{ + Common: Common{ID: 9, UserID: 100}, + Rules: []*Rule{{ + Type: "cpu", + Cover: RuleCoverAll, + Ignore: nil, + }}, + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) // doesn't cover 2 + + if rule.HasPermission(ctx) { + t.Fatal("server-limited PAT must not be allowed to operate on RuleCoverAll with empty Ignore — runtime fans out to owner server 2") + } +} + +func TestAlertRuleHasPermission_AllowsRuleCoverIgnoreAllEmptyIgnore(t *testing.T) { + saved := OwnerServerIDsLookup + t.Cleanup(func() { OwnerServerIDsLookup = saved }) + OwnerServerIDsLookup = func(uid uint64) []uint64 { return []uint64{1, 2} } + + rule := &AlertRule{ + Common: Common{ID: 10, UserID: 100}, + Rules: []*Rule{{ + Type: "cpu", + Cover: RuleCoverIgnoreAll, + Ignore: nil, // allow-list of zero ⇒ no-op + }}, + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) + + if !rule.HasPermission(ctx) { + t.Fatal("RuleCoverIgnoreAll + empty Ignore is a no-op rule; PAT must remain allowed") + } +} + +func TestAlertRuleHasPermission_DeniesUnknownCover(t *testing.T) { + saved := OwnerServerIDsLookup + t.Cleanup(func() { OwnerServerIDsLookup = saved }) + OwnerServerIDsLookup = func(uid uint64) []uint64 { return []uint64{1} } + + rule := &AlertRule{ + Common: Common{ID: 11, UserID: 100}, + Rules: []*Rule{{ + Type: "cpu", + Cover: 99, // unknown ⇒ runtime Snapshot falls through to "monitor everything" + }}, + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) + + if rule.HasPermission(ctx) { + t.Fatal("unknown Rule.Cover must fail-closed for limited PAT — Snapshot does not gate on it") + } +} + +// Regression: AlertRule.HasPermission built the deny-list from every key in +// Rule.Ignore regardless of its bool value, but Rule.Snapshot only skips a +// server when Ignore[id] == true. A limited PAT whitelisted to {1} could +// submit RuleCoverAll with Ignore{2: false}; the permission check treated 2 as +// denied (safe) while the runtime still monitored server 2. +func TestAlertRuleHasPermission_DeniesRuleCoverAllIgnoreFalseForLimitedPAT(t *testing.T) { + saved := OwnerServerIDsLookup + savedAdmin := OwnerIsAdminLookup + savedAll := AllServerIDsLookup + t.Cleanup(func() { + OwnerServerIDsLookup = saved + OwnerIsAdminLookup = savedAdmin + AllServerIDsLookup = savedAll + }) + OwnerServerIDsLookup = func(uid uint64) []uint64 { + if uid == 100 { + return []uint64{1, 2} + } + return nil + } + OwnerIsAdminLookup = func(uid uint64) bool { return false } + AllServerIDsLookup = func() []uint64 { return []uint64{1, 2, 3} } + + rule := &AlertRule{ + Common: Common{ID: 13, UserID: 100}, + Rules: []*Rule{{ + Type: "cpu", + Cover: RuleCoverAll, + Ignore: map[uint64]bool{2: false}, // key present but NOT actually ignored at runtime + }}, + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) // doesn't cover 2 + + if rule.HasPermission(ctx) { + t.Fatal("Ignore{2:false} does NOT exclude server 2 at runtime; limited PAT must be denied") + } +} + +// A genuine deny entry (value true) for every out-of-whitelist server keeps the +// rule contained and must remain allowed. +func TestAlertRuleHasPermission_AllowsRuleCoverAllIgnoreTrueCoversWhitelistGap(t *testing.T) { + saved := OwnerServerIDsLookup + savedAdmin := OwnerIsAdminLookup + savedAll := AllServerIDsLookup + t.Cleanup(func() { + OwnerServerIDsLookup = saved + OwnerIsAdminLookup = savedAdmin + AllServerIDsLookup = savedAll + }) + OwnerServerIDsLookup = func(uid uint64) []uint64 { return []uint64{1, 2} } + OwnerIsAdminLookup = func(uid uint64) bool { return false } + AllServerIDsLookup = func() []uint64 { return []uint64{1, 2} } + + rule := &AlertRule{ + Common: Common{ID: 14, UserID: 100}, + Rules: []*Rule{{ + Type: "cpu", + Cover: RuleCoverAll, + Ignore: map[uint64]bool{2: true}, // server 2 genuinely excluded + }}, + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) + + if !rule.HasPermission(ctx) { + t.Fatal("Ignore{2:true} excludes the only out-of-whitelist server; PAT must be allowed") + } +} + +// Regression: the RuleCoverIgnoreAll branch checked CanAccessServer for every +// key in Rule.Ignore, but Rule.Snapshot only monitors a server when +// Ignore[id] == true. A limited PAT whitelisted to {1} submitting +// RuleCoverIgnoreAll with Ignore{2: false} (server 2 is NOT monitored at +// runtime) was wrongly denied because the foreign key 2 failed the whitelist. +func TestAlertRuleHasPermission_AllowsRuleCoverIgnoreAllIgnoreFalseForLimitedPAT(t *testing.T) { + saved := OwnerServerIDsLookup + t.Cleanup(func() { OwnerServerIDsLookup = saved }) + OwnerServerIDsLookup = func(uid uint64) []uint64 { return []uint64{1, 2} } + + rule := &AlertRule{ + Common: Common{ID: 15, UserID: 100}, + Rules: []*Rule{{ + Type: "cpu", + Cover: RuleCoverIgnoreAll, + Ignore: map[uint64]bool{2: false}, // key present but server 2 is NOT monitored + }}, + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) // doesn't cover 2 + + if !rule.HasPermission(ctx) { + t.Fatal("Ignore{2:false} does NOT monitor server 2 at runtime; limited PAT must remain allowed") + } +} + +// The genuine allow entry (value true) for an out-of-whitelist server is the +// case that must still be denied. +func TestAlertRuleHasPermission_DeniesRuleCoverIgnoreAllIgnoreTrueForLimitedPAT(t *testing.T) { + saved := OwnerServerIDsLookup + t.Cleanup(func() { OwnerServerIDsLookup = saved }) + OwnerServerIDsLookup = func(uid uint64) []uint64 { return []uint64{1, 2} } + + rule := &AlertRule{ + Common: Common{ID: 16, UserID: 100}, + Rules: []*Rule{{ + Type: "cpu", + Cover: RuleCoverIgnoreAll, + Ignore: map[uint64]bool{2: true}, // server 2 IS monitored at runtime + }}, + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) // doesn't cover 2 + + if rule.HasPermission(ctx) { + t.Fatal("Ignore{2:true} monitors server 2; server-limited PAT must be denied") + } +} + +func TestAlertRuleHasPermission_NoPATPassesViaCommonHasPermission(t *testing.T) { + rule := &AlertRule{ + Common: Common{ID: 12, UserID: 100}, + Rules: []*Rule{{Type: "cpu", Cover: RuleCoverAll}}, + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + + if !rule.HasPermission(ctx) { + t.Fatal("owner without PAT must keep the existing owner/admin pass") + } +} diff --git a/model/api_token.go b/model/api_token.go new file mode 100644 index 00000000..030052b7 --- /dev/null +++ b/model/api_token.go @@ -0,0 +1,373 @@ +package model + +import ( + "crypto/sha256" + "encoding/hex" + "slices" + "strings" + "time" + + "gorm.io/gorm" +) + +// Scope 命名规范(唯一一套):nezha:{resource}:{verb} +// +// - resource: server / service / alertrule / cron / ddns / nat / +// notification / notification-group / transfer / admin +// - verb: read / write / delete / exec +// +// `*` 通配在 resource 或 verb 位均可: +// - nezha:server:* 给定资源的所有动作 +// - nezha:* admin-only 全权 +// +// 同一 scope 同时管 MCP tool 和 REST endpoint:例如 nezha:server:read 既允许 +// `server.list` MCP tool,也允许 `GET /api/v1/server`。 +// +// 历史上还有 mcp:* 一套,会被 HasScope 通过别名映射到 nezha:server:* 子集。 +// 由于 HasScope 同时服务 MCP tool 调度与 REST scope middleware,旧 mcp:fs:write +// 等会静默扩到所有 nezha:server:write REST 路由——这是命名分裂带来的提权漏洞。 +// 现在 mcp:* 不再在运行时被识别;createAPIToken 入口对老调用方做一次性归一化: +// 只读/exec 类(mcp:fs:read、mcp:server:read、mcp:server:exec)映射到对应 +// nezha:* read/exec scope;mcp:fs:write、mcp:fs:delete、mcp:* 一律拒签。 +// 数据库已有的危险旧 scope 由 MigrateLegacyMCPScopes 在启动迁移阶段清理。 +const ( + ScopeNezhaAll = "nezha:*" + + ScopeServerRead = "nezha:server:read" + ScopeServerWrite = "nezha:server:write" + ScopeServerDelete = "nezha:server:delete" + ScopeServerExec = "nezha:server:exec" + + ScopeServiceRead = "nezha:service:read" + ScopeServiceWrite = "nezha:service:write" + ScopeServiceDelete = "nezha:service:delete" + + ScopeAlertRuleRead = "nezha:alertrule:read" + ScopeAlertRuleWrite = "nezha:alertrule:write" + ScopeAlertRuleDelete = "nezha:alertrule:delete" + + ScopeCronRead = "nezha:cron:read" + ScopeCronWrite = "nezha:cron:write" + ScopeCronDelete = "nezha:cron:delete" + ScopeCronExec = "nezha:cron:exec" + + ScopeDDNSRead = "nezha:ddns:read" + ScopeDDNSWrite = "nezha:ddns:write" + ScopeDDNSDelete = "nezha:ddns:delete" + + ScopeNATRead = "nezha:nat:read" + ScopeNATWrite = "nezha:nat:write" + ScopeNATDelete = "nezha:nat:delete" + + ScopeNotificationRead = "nezha:notification:read" + ScopeNotificationWrite = "nezha:notification:write" + ScopeNotificationDelete = "nezha:notification:delete" + + ScopeNotificationGroupRead = "nezha:notification-group:read" + ScopeNotificationGroupWrite = "nezha:notification-group:write" + ScopeNotificationGroupDelete = "nezha:notification-group:delete" + + ScopeTransferRead = "nezha:transfer:read" + ScopeTransferWrite = "nezha:transfer:write" + ScopeTransferDelete = "nezha:transfer:delete" + + ScopeAdminAll = "nezha:admin:*" +) + +var AllScopes = []string{ + ScopeServerRead, ScopeServerWrite, ScopeServerDelete, ScopeServerExec, + ScopeServiceRead, ScopeServiceWrite, ScopeServiceDelete, + ScopeAlertRuleRead, ScopeAlertRuleWrite, ScopeAlertRuleDelete, + ScopeCronRead, ScopeCronWrite, ScopeCronDelete, ScopeCronExec, + ScopeDDNSRead, ScopeDDNSWrite, ScopeDDNSDelete, + ScopeNATRead, ScopeNATWrite, ScopeNATDelete, + ScopeNotificationRead, ScopeNotificationWrite, ScopeNotificationDelete, + ScopeNotificationGroupRead, ScopeNotificationGroupWrite, ScopeNotificationGroupDelete, + ScopeTransferRead, ScopeTransferWrite, ScopeTransferDelete, + + "nezha:server:*", + "nezha:service:*", + "nezha:alertrule:*", + "nezha:cron:*", + "nezha:ddns:*", + "nezha:nat:*", + "nezha:notification:*", + "nezha:notification-group:*", + "nezha:transfer:*", +} + +var AdminOnlyScopes = []string{ScopeNezhaAll, ScopeAdminAll} + +// legacyMCPReadOnlyRewrite 列出仍允许 createAPIToken 入口重写为 nezha:* 的旧 scope。 +// 只有只读/exec 类被接受;write/delete/wildcard 一律拒签——保留映射等于扩权。 +var legacyMCPReadOnlyRewrite = map[string]string{ + "mcp:server:read": ScopeServerRead, + "mcp:server:exec": ScopeServerExec, + "mcp:fs:read": ScopeServerRead, +} + +// NormalizeIncomingScope 把入参里的旧 mcp:* scope 重写到 nezha:* 命名。 +// 第二个返回值表示该 scope 是否被允许(false = 危险旧 scope,调用方应拒签)。 +func NormalizeIncomingScope(s string) (string, bool) { + if mapped, ok := legacyMCPReadOnlyRewrite[s]; ok { + return mapped, true + } + if strings.HasPrefix(s, "mcp:") { + return s, false + } + return s, true +} + +// APITokenPrefix 是明文 token 的人类可识别前缀。`nzp_` = nezha personal access token。 +const APITokenPrefix = "nzp_" + +// APIToken 是用户用于程序化访问的长期凭证。MCP 接入点 /mcp 用它做鉴权。 +// 双层鉴权:闸 1 用 UserID 复用 Server.HasPermission;闸 2 用 Scopes / ServerIDs。 +type APIToken struct { + ID uint64 `gorm:"primaryKey" json:"id,omitempty"` + UserID uint64 `gorm:"index" json:"user_id,omitempty"` + Name string `gorm:"type:varchar(128)" json:"name,omitempty"` + TokenHash string `gorm:"uniqueIndex;type:char(64)" json:"-"` + ScopesCSV string `gorm:"type:text" json:"-"` + ServersCSV string `gorm:"type:text" json:"-"` + ExpiresAt *time.Time `gorm:"index" json:"expires_at,omitempty"` + LastUsedAt *time.Time `json:"last_used_at,omitempty"` + LastUsedIP string `gorm:"type:varchar(64)" json:"last_used_ip,omitempty"` + CreatedAt time.Time `json:"created_at,omitempty"` + UpdatedAt time.Time `gorm:"autoUpdateTime" json:"updated_at,omitempty"` +} + +func (APIToken) TableName() string { + return "api_tokens" +} + +// HashAPIToken 计算明文 token 的存储哈希。 +func HashAPIToken(plaintext string) string { + sum := sha256.Sum256([]byte(plaintext)) + return hex.EncodeToString(sum[:]) +} + +// Scopes 解码逗号分隔的 scope 列表。 +func (t *APIToken) Scopes() []string { + if t.ScopesCSV == "" { + return nil + } + parts := strings.Split(t.ScopesCSV, ",") + out := parts[:0] + for _, p := range parts { + p = strings.TrimSpace(p) + if p != "" { + out = append(out, p) + } + } + return out +} + +// SetScopes 编码 scope 列表为 CSV。 +func (t *APIToken) SetScopes(scopes []string) { + t.ScopesCSV = strings.Join(scopes, ",") +} + +// ServerIDs 解码服务器 ID 白名单。空切片 = 不限制(继承用户原有权限)。 +func (t *APIToken) ServerIDs() []uint64 { + if t.ServersCSV == "" { + return nil + } + parts := strings.Split(t.ServersCSV, ",") + out := make([]uint64, 0, len(parts)) + for _, p := range parts { + p = strings.TrimSpace(p) + if p == "" { + continue + } + var id uint64 + for _, c := range p { + if c < '0' || c > '9' { + id = 0 + break + } + id = id*10 + uint64(c-'0') + } + if id != 0 { + out = append(out, id) + } + } + return out +} + +// SetServerIDs 编码服务器 ID 白名单。 +func (t *APIToken) SetServerIDs(ids []uint64) { + parts := make([]string, 0, len(ids)) + for _, id := range ids { + parts = append(parts, formatUint(id)) + } + t.ServersCSV = strings.Join(parts, ",") +} + +// HasScope 判定 token 是否携带某个 scope。 +// +// 匹配规则: +// - nezha:* 覆盖整个 nezha 命名空间 +// - 资源级通配:nezha:server:* 匹配所有 nezha:server:read/write/delete/exec +// - 精确匹配 +// +// 不再做 mcp:* 别名展开;任何遗留的 mcp:* scope 都视为无效(已被 +// MigrateLegacyMCPScopes 在启动迁移阶段清掉;运行时再遇到当作无权处理)。 +func (t *APIToken) HasScope(scope string) bool { + for _, s := range t.Scopes() { + if scopeMatches(s, scope) { + return true + } + } + return false +} + +// scopeMatches 判定 owned scope 是否覆盖 wanted scope。 +func scopeMatches(owned, wanted string) bool { + if owned == wanted { + return true + } + if owned == ScopeNezhaAll { + return strings.HasPrefix(wanted, "nezha:") + } + if strings.HasSuffix(owned, ":*") { + prefix := strings.TrimSuffix(owned, ":*") + return strings.HasPrefix(wanted, prefix+":") || wanted == prefix + } + return false +} + +// CanAccessServer 判定 token 是否被允许操作某 server(白名单层; +// 仍需上层调用 Server.HasPermission 做用户级权限校验)。 +func (t *APIToken) CanAccessServer(serverID uint64) bool { + ids := t.ServerIDs() + if len(ids) == 0 { + return true + } + return slices.Contains(ids, serverID) +} + +// IsExpired 判定 token 是否已过期。ExpiresAt 为 nil 表示永不过期。 +func (t *APIToken) IsExpired(now time.Time) bool { + return t.ExpiresAt != nil && now.After(*t.ExpiresAt) +} + +// BeforeCreate 在写入前强校验 TokenHash 必填,避免空哈希撞键。 +func (t *APIToken) BeforeCreate(tx *gorm.DB) error { + if t.TokenHash == "" { + return gorm.ErrInvalidData + } + return nil +} + +// MigrateLegacyMCPScopes 把数据库里残留的 mcp:* scope 一次性归一化: +// - 只读/exec 类映射到对应 nezha:* read/exec scope; +// - mcp:fs:write / mcp:fs:delete / mcp:* 会被剥掉(不再赋予 REST write/delete), +// 若 token 因此 scope 列表清空则整体删除——避免出现一张 0 scope 但仍能命中 +// auth middleware 的 PAT。 +// +// 返回 (rewrittenTokens, deletedTokens, err)。生产路径在启动时调用一次; +// 测试也会用它构造 fixture。 +func MigrateLegacyMCPScopes(db *gorm.DB) (int, int, error) { + if db == nil { + return 0, 0, nil + } + var rows []APIToken + if err := db.Where("scopes_csv LIKE ?", "%mcp:%").Find(&rows).Error; err != nil { + return 0, 0, err + } + rewritten, deleted := 0, 0 + for i := range rows { + tok := &rows[i] + old := tok.Scopes() + next := make([]string, 0, len(old)) + seen := make(map[string]struct{}, len(old)) + for _, s := range old { + mapped, ok := NormalizeIncomingScope(s) + if !ok { + continue + } + if _, dup := seen[mapped]; dup { + continue + } + seen[mapped] = struct{}{} + next = append(next, mapped) + } + if len(next) == 0 { + if err := db.Delete(&APIToken{}, tok.ID).Error; err != nil { + return rewritten, deleted, err + } + deleted++ + continue + } + joined := strings.Join(next, ",") + if joined == tok.ScopesCSV { + continue + } + if err := db.Model(&APIToken{}).Where("id = ?", tok.ID). + Update("scopes_csv", joined).Error; err != nil { + return rewritten, deleted, err + } + rewritten++ + } + return rewritten, deleted, nil +} + +// formatUint —— 小工具,避免引入 strconv。 +func formatUint(v uint64) string { + if v == 0 { + return "0" + } + var buf [20]byte + i := len(buf) + for v > 0 { + i-- + buf[i] = byte('0' + v%10) + v /= 10 + } + return string(buf[i:]) +} + +// APITokenCreateRequest 是创建 PAT 接口的入参。 +type APITokenCreateRequest struct { + Name string `json:"name" binding:"required,max=128"` + Scopes []string `json:"scopes" binding:"required,min=1,dive,max=64"` + ServerIDs []uint64 `json:"server_ids,omitempty"` + ExpiresInDays int `json:"expires_in_days,omitempty"` // 0 = 永不过期 +} + +// APITokenCreateResponse 创建 PAT 接口的出参;明文 token 仅在此刻返回一次。 +type APITokenCreateResponse struct { + ID uint64 `json:"id"` + Name string `json:"name"` + Token string `json:"token"` + Scopes []string `json:"scopes"` + ServerIDs []uint64 `json:"server_ids,omitempty"` + ExpiresAt *time.Time `json:"expires_at,omitempty"` +} + +// APITokenView 是 PAT 列表展示用的脱敏视图。 +type APITokenView struct { + ID uint64 `json:"id"` + Name string `json:"name"` + Scopes []string `json:"scopes"` + ServerIDs []uint64 `json:"server_ids,omitempty"` + ExpiresAt *time.Time `json:"expires_at,omitempty"` + LastUsedAt *time.Time `json:"last_used_at,omitempty"` + LastUsedIP string `json:"last_used_ip,omitempty"` + CreatedAt time.Time `json:"created_at"` +} + +// ToView 把数据库实体转为列表脱敏视图。 +func (t *APIToken) ToView() APITokenView { + return APITokenView{ + ID: t.ID, + Name: t.Name, + Scopes: t.Scopes(), + ServerIDs: t.ServerIDs(), + ExpiresAt: t.ExpiresAt, + LastUsedAt: t.LastUsedAt, + LastUsedIP: t.LastUsedIP, + CreatedAt: t.CreatedAt, + } +} diff --git a/model/api_token_migration_test.go b/model/api_token_migration_test.go new file mode 100644 index 00000000..7ddf1d09 --- /dev/null +++ b/model/api_token_migration_test.go @@ -0,0 +1,112 @@ +package model + +import ( + "strings" + "testing" + + "gorm.io/driver/sqlite" + "gorm.io/gorm" +) + +func newMigrationTestDB(t *testing.T) *gorm.DB { + t.Helper() + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + if err != nil { + t.Fatalf("open db: %v", err) + } + if err := db.AutoMigrate(&APIToken{}); err != nil { + t.Fatalf("migrate: %v", err) + } + return db +} + +func TestNormalizeIncomingScope_RewritesReadOnlyMCPVariants(t *testing.T) { + cases := map[string]string{ + "mcp:fs:read": ScopeServerRead, + "mcp:server:read": ScopeServerRead, + "mcp:server:exec": ScopeServerExec, + } + for in, want := range cases { + got, ok := NormalizeIncomingScope(in) + if !ok { + t.Fatalf("NormalizeIncomingScope(%q) ok=false; legacy read/exec must remain creatable", in) + } + if got != want { + t.Fatalf("NormalizeIncomingScope(%q) = %q, want %q", in, got, want) + } + } +} + +func TestNormalizeIncomingScope_RejectsDangerousLegacyVariants(t *testing.T) { + for _, in := range []string{"mcp:fs:write", "mcp:fs:delete", "mcp:*", "mcp:unknown"} { + if got, ok := NormalizeIncomingScope(in); ok { + t.Errorf("NormalizeIncomingScope(%q) = (%q, true); legacy write/delete/wildcard must be rejected", in, got) + } + } +} + +func TestNormalizeIncomingScope_PassesThroughNezhaScopes(t *testing.T) { + got, ok := NormalizeIncomingScope(ScopeServerWrite) + if !ok || got != ScopeServerWrite { + t.Fatalf("nezha:* must pass through unchanged: got (%q, %v)", got, ok) + } +} + +func TestMigrateLegacyMCPScopes_RewritesReadOnlyAndDropsDangerous(t *testing.T) { + db := newMigrationTestDB(t) + + tokens := []APIToken{ + {UserID: 1, Name: "read-only", TokenHash: HashAPIToken("nzp_a"), ScopesCSV: "mcp:fs:read"}, + {UserID: 2, Name: "mixed", TokenHash: HashAPIToken("nzp_b"), ScopesCSV: "mcp:server:read,mcp:fs:write"}, + {UserID: 3, Name: "purely-dangerous", TokenHash: HashAPIToken("nzp_c"), ScopesCSV: "mcp:fs:write,mcp:*"}, + {UserID: 4, Name: "modern", TokenHash: HashAPIToken("nzp_d"), ScopesCSV: ScopeServerRead}, + } + for i := range tokens { + if err := db.Create(&tokens[i]).Error; err != nil { + t.Fatalf("seed token %d: %v", i, err) + } + } + + rewritten, deleted, err := MigrateLegacyMCPScopes(db) + if err != nil { + t.Fatalf("MigrateLegacyMCPScopes: %v", err) + } + if rewritten < 2 { + t.Fatalf("expected >=2 rewrites (read-only + mixed); got %d", rewritten) + } + if deleted != 1 { + t.Fatalf("expected 1 deleted token (purely-dangerous); got %d", deleted) + } + + var got []APIToken + if err := db.Order("id ASC").Find(&got).Error; err != nil { + t.Fatalf("reload: %v", err) + } + if len(got) != 3 { + t.Fatalf("expected 3 surviving tokens; got %d", len(got)) + } + for _, tok := range got { + if strings.Contains(tok.ScopesCSV, "mcp:") { + t.Fatalf("token %d still carries legacy scope after migration: %q", tok.ID, tok.ScopesCSV) + } + } + + for _, tok := range got { + switch tok.UserID { + case 1: + if tok.ScopesCSV != ScopeServerRead { + t.Fatalf("uid=1 expected %q, got %q", ScopeServerRead, tok.ScopesCSV) + } + case 2: + if tok.ScopesCSV != ScopeServerRead { + t.Fatalf("uid=2 should keep only the safe read scope after drop; got %q", tok.ScopesCSV) + } + case 4: + if tok.ScopesCSV != ScopeServerRead { + t.Fatalf("uid=4 already-modern token must be untouched; got %q", tok.ScopesCSV) + } + default: + t.Fatalf("unexpected surviving token uid=%d", tok.UserID) + } + } +} diff --git a/model/api_token_test.go b/model/api_token_test.go new file mode 100644 index 00000000..6bdba54d --- /dev/null +++ b/model/api_token_test.go @@ -0,0 +1,148 @@ +package model + +import ( + "strings" + "testing" + "time" +) + +func TestHashAPIToken_DeterministicAndAvalanche(t *testing.T) { + a := HashAPIToken("nzp_alpha") + b := HashAPIToken("nzp_alpha") + if a != b { + t.Fatalf("hash must be deterministic for identical inputs") + } + c := HashAPIToken("nzp_alphb") + if a == c { + t.Fatalf("single-byte change in input must change hash") + } + if len(a) != 64 { + t.Fatalf("hash must be 64 hex chars (sha256), got %d", len(a)) + } +} + +func TestAPIToken_HasScope_AllAndExact(t *testing.T) { + tok := &APIToken{} + tok.SetScopes([]string{ScopeServerRead, ScopeServerExec}) + if !tok.HasScope(ScopeServerRead) { + t.Fatalf("explicit scope must pass") + } + if !tok.HasScope(ScopeServerExec) { + t.Fatalf("explicit scope must pass") + } + if tok.HasScope(ScopeServerWrite) { + t.Fatalf("missing scope must fail") + } + + tok.SetScopes([]string{ScopeNezhaAll}) + for _, s := range []string{ScopeServerRead, ScopeServerWrite, ScopeServerDelete, ScopeServerExec} { + if !tok.HasScope(s) { + t.Fatalf("nezha:* must cover %s", s) + } + } +} + +func TestAPIToken_HasScope_TrimsWhitespace(t *testing.T) { + tok := &APIToken{ScopesCSV: " nezha:server:read , nezha:server:exec "} + if !tok.HasScope(ScopeServerRead) { + t.Fatalf("scope with surrounding whitespace must be normalized") + } +} + +func TestAPIToken_CanAccessServer_EmptyMeansAll(t *testing.T) { + tok := &APIToken{} + if !tok.CanAccessServer(1) { + t.Fatalf("empty server list must allow any server") + } + if !tok.CanAccessServer(99999) { + t.Fatalf("empty server list must allow any server") + } +} + +func TestAPIToken_CanAccessServer_Whitelist(t *testing.T) { + tok := &APIToken{} + tok.SetServerIDs([]uint64{2, 5, 7}) + if !tok.CanAccessServer(5) { + t.Fatalf("listed server must be allowed") + } + if tok.CanAccessServer(6) { + t.Fatalf("unlisted server must be denied") + } +} + +func TestAPIToken_SetServerIDs_RoundTrip(t *testing.T) { + tok := &APIToken{} + tok.SetServerIDs([]uint64{10, 11, 12}) + got := tok.ServerIDs() + if len(got) != 3 || got[0] != 10 || got[2] != 12 { + t.Fatalf("round-trip failed: %v", got) + } +} + +func TestAPIToken_ServerIDs_SkipsGarbage(t *testing.T) { + tok := &APIToken{ServersCSV: "1,2,abc,3"} + got := tok.ServerIDs() + if len(got) != 3 || got[0] != 1 || got[1] != 2 || got[2] != 3 { + t.Fatalf("garbage entries must be skipped; got %v", got) + } +} + +func TestAPIToken_IsExpired(t *testing.T) { + tok := &APIToken{} + if tok.IsExpired(time.Now()) { + t.Fatalf("nil expiry must mean never expired") + } + past := time.Now().Add(-time.Hour) + tok.ExpiresAt = &past + if !tok.IsExpired(time.Now()) { + t.Fatalf("past expiry must mark expired") + } + future := time.Now().Add(time.Hour) + tok.ExpiresAt = &future + if tok.IsExpired(time.Now()) { + t.Fatalf("future expiry must not mark expired") + } +} + +func TestAPIToken_HashAPIToken_NoSecretLeak(t *testing.T) { + plaintext := "nzp_supersecret" + hash := HashAPIToken(plaintext) + if strings.Contains(hash, plaintext) { + t.Fatalf("hash must not contain plaintext") + } + if strings.Contains(hash, "super") { + t.Fatalf("hash must not contain secret substring") + } +} + +func TestAPIToken_BeforeCreate_RejectsEmptyHash(t *testing.T) { + tok := &APIToken{Name: "x"} + err := tok.BeforeCreate(nil) + if err == nil { + t.Fatalf("BeforeCreate must reject empty TokenHash") + } +} + +func TestAPIToken_BeforeCreate_AcceptsNonEmptyHash(t *testing.T) { + tok := &APIToken{Name: "x", TokenHash: HashAPIToken("nzp_xyz")} + if err := tok.BeforeCreate(nil); err != nil { + t.Fatalf("BeforeCreate must accept non-empty hash, got %v", err) + } +} + +func TestAPIToken_ToView_OmitsTokenHash(t *testing.T) { + tok := &APIToken{ + ID: 1, + UserID: 2, + Name: "x", + TokenHash: "DEADBEEF", + } + tok.SetScopes([]string{ScopeServerRead}) + v := tok.ToView() + if v.ID != 1 || v.Name != "x" { + t.Fatalf("view missing core fields") + } + if len(v.Scopes) != 1 || v.Scopes[0] != ScopeServerRead { + t.Fatalf("view missing scopes") + } +} diff --git a/model/api_token_unified_scope_test.go b/model/api_token_unified_scope_test.go new file mode 100644 index 00000000..cd151f70 --- /dev/null +++ b/model/api_token_unified_scope_test.go @@ -0,0 +1,82 @@ +package model + +import ( + "slices" + "testing" +) + +// 这些测试约束「scope 命名统一」契约: +// - 只有 nezha:* 一套是 first-class scope; +// - mcp:* 不再作为 HasScope 的别名(避免 mcp:fs:write 静默扩到 REST 的 nezha:server:write); +// - AllScopes / AdminOnlyScopes 不再包含 mcp:*,新建 token 不能再签发它们。 +// +// 旧 mcp:* 兼容由 createAPIToken 入口做一次性归一化(mcp:fs:read 等只读/exec 映射到 +// 对应的 nezha:* read/exec),但 write/delete 类不再映射;详见 controller.createAPIToken。 + +func TestAllScopes_DoesNotExposeLegacyMCPScopes(t *testing.T) { + legacy := []string{ + "mcp:*", + "mcp:server:read", + "mcp:server:exec", + "mcp:fs:read", + "mcp:fs:write", + "mcp:fs:delete", + } + for _, s := range legacy { + if slices.Contains(AllScopes, s) { + t.Errorf("AllScopes must not advertise legacy scope %q; only nezha:* is first-class", s) + } + if slices.Contains(AdminOnlyScopes, s) { + t.Errorf("AdminOnlyScopes must not advertise legacy scope %q", s) + } + } +} + +func TestHasScope_LegacyMCPNoLongerAliasesNezhaWrite(t *testing.T) { + // 旧 token 数据库里残留 mcp:fs:write,绝不允许覆盖 REST 的 nezha:server:write。 + tok := &APIToken{ScopesCSV: "mcp:fs:write"} + if tok.HasScope(ScopeServerWrite) { + t.Fatalf("legacy mcp:fs:write must NOT grant nezha:server:write via HasScope; " + + "REST routes (server/config, server/:id, batch-delete/server) would become reachable") + } + if tok.HasScope(ScopeServerDelete) { + t.Fatalf("legacy mcp:fs:write must NOT grant nezha:server:delete") + } +} + +func TestHasScope_LegacyMCPDeleteNoLongerAliasesNezhaDelete(t *testing.T) { + tok := &APIToken{ScopesCSV: "mcp:fs:delete"} + if tok.HasScope(ScopeServerDelete) { + t.Fatalf("legacy mcp:fs:delete must NOT grant nezha:server:delete via HasScope") + } +} + +func TestHasScope_LegacyMCPAllNoLongerWildcards(t *testing.T) { + tok := &APIToken{ScopesCSV: "mcp:*"} + for _, s := range []string{ScopeServerRead, ScopeServerWrite, ScopeServerDelete, ScopeServerExec} { + if tok.HasScope(s) { + t.Errorf("legacy mcp:* must not be treated as a nezha:* wildcard; granted %s", s) + } + } +} + +func TestHasScope_NezhaWildcardStillWorks(t *testing.T) { + tok := &APIToken{ScopesCSV: ScopeNezhaAll} + for _, s := range []string{ScopeServerRead, ScopeServerWrite, ScopeServerDelete, ScopeServerExec} { + if !tok.HasScope(s) { + t.Errorf("nezha:* wildcard must still cover %s", s) + } + } +} + +func TestHasScope_NezhaResourceWildcardStillWorks(t *testing.T) { + tok := &APIToken{ScopesCSV: "nezha:server:*"} + for _, s := range []string{ScopeServerRead, ScopeServerWrite, ScopeServerDelete, ScopeServerExec} { + if !tok.HasScope(s) { + t.Errorf("nezha:server:* must cover %s", s) + } + } + if tok.HasScope(ScopeServiceRead) { + t.Fatalf("nezha:server:* must NOT leak into nezha:service:* family") + } +} diff --git a/model/common.go b/model/common.go index 3d165f94..cda2b486 100644 --- a/model/common.go +++ b/model/common.go @@ -17,8 +17,13 @@ const ( CtxKeyAuthorizedUser = "ckau" CtxKeyRealIPStr = "ckri" CtxKeyIsIPMismatch = "ckipm" + CtxKeyAPIToken = "ckpat" ) +type APITokenAccessor interface { + CanAccessServer(uint64) bool +} + const ( CacheKeyOauth2State = "cko2s::" ) diff --git a/model/config.go b/model/config.go index c939a583..2990d961 100644 --- a/model/config.go +++ b/model/config.go @@ -6,6 +6,7 @@ import ( "path/filepath" "strconv" "strings" + "sync/atomic" "github.com/go-viper/mapstructure/v2" kmaps "github.com/knadh/koanf/maps" @@ -51,6 +52,8 @@ type ConfigDashboard struct { EnablePlainIPInNotification bool `koanf:"enable_plain_ip_in_notification" json:"enable_plain_ip_in_notification,omitempty"` // 通知信息IP不打码 + EnableMCP bool `koanf:"enable_mcp" json:"enable_mcp,omitempty"` // 是否启用 MCP 入口(默认关闭;启用前请审视 PAT scope/whitelist) + // IP变更提醒 EnableIPChangeNotification bool `koanf:"enable_ip_change_notification" json:"enable_ip_change_notification,omitempty"` IPChangeNotificationGroupID uint64 `koanf:"ip_change_notification_group_id" json:"ip_change_notification_group_id"` @@ -80,6 +83,11 @@ type Config struct { jwtSecretFromEnv bool `koanf:"-" json:"-" yaml:"-"` jwtSecretFromYAML bool `koanf:"-" json:"-" yaml:"-"` + // mcpEnabled:EnableMCP 的并发安全镜像,kill switch 跨 goroutine 读写走 + // MCPEnabled()/SetMCPEnabled()。放外层 Config 而非 ConfigDashboard,避免 + // SettingResponse 按值拷贝 ConfigDashboard 触发 copylocks。 + mcpEnabled atomic.Bool `koanf:"-" json:"-" yaml:"-"` + // oauth2 配置 Oauth2 map[string]*Oauth2Config `koanf:"oauth2" json:"oauth2,omitempty"` @@ -212,9 +220,23 @@ func (c *Config) Read(path string, frontendTemplates []FrontendTemplate) error { } } + c.mcpEnabled.Store(c.EnableMCP) + return nil } +// MCPEnabled 并发安全地读取 MCP kill switch 状态。 +func (c *Config) MCPEnabled() bool { + return c.mcpEnabled.Load() +} + +// SetMCPEnabled 并发安全地更新 MCP kill switch 状态。只写 atomic 镜像,不直接 +// 写 EnableMCP 明文字段——后者会与 listConfig 的 *singleton.Conf 整体拷贝读发生 +// 数据竞争。持久化由 save() 在 marshal 前从 atomic 同步明文字段完成。 +func (c *Config) SetMCPEnabled(v bool) { + c.mcpEnabled.Store(v) +} + // Save 保存配置文件 func (c *Config) Save() error { return c.save() @@ -282,6 +304,7 @@ func (c *Config) patchYAMLField(key string, value any) error { } func (c *Config) save() error { + c.EnableMCP = c.mcpEnabled.Load() data, err := yaml.Marshal(c) if err != nil { return err diff --git a/model/cron.go b/model/cron.go index 418c4215..0ea30304 100644 --- a/model/cron.go +++ b/model/cron.go @@ -3,6 +3,7 @@ package model import ( "time" + "github.com/gin-gonic/gin" "github.com/goccy/go-json" "github.com/robfig/cron/v3" "gorm.io/gorm" @@ -45,3 +46,44 @@ func (c *Cron) BeforeSave(tx *gorm.DB) error { func (c *Cron) AfterFind(tx *gorm.DB) error { return json.Unmarshal([]byte(c.ServersRaw), &c.Servers) } + +// HasPermission 扩展默认的 owner/admin 检查,使得 PAT 的 server_ids 白名单 +// 同样能收窄 cron 的列出、触发、删除路径。 +// +// 语义按 Cover 字段分流,与 dispatch 入口(CronTrigger)的 fan-out 规则严格 +// 对齐——Servers 字段在不同 Cover 下含义完全相反: +// +// - CronCoverIgnoreAll:Servers 是 allow-list;必须每个 server 都落在 PAT +// 白名单内。空 allow-list 是「matches nothing」的退化形态,安全。 +// - CronCoverAlertTrigger:Servers 是触发服务器 allow-list;与上同。 +// - CronCoverAll:Servers 是 deny-list。dispatch 时 fan out 到 owner 的 +// 全部 server 再减去这个 deny-list。受限 PAT 必须保证 deny-list 已经覆 +// 盖 owner 在白名单外的所有 servers——否则 CronTrigger 会把任务发到 +// PAT 没权限的 server 上。本方法和 controller 写侧 guard +// rejectImplicitCoverForLimitedPAT* / 运行时 guard +// enforcePATCronDispatchScope 共用 DenyListSafeForLimitedPAT,避免列表 +// 视图把越界历史/旁路写入行漏给受限 PAT。 +func (c *Cron) HasPermission(ctx *gin.Context) bool { + if !c.Common.HasPermission(ctx) { + return false + } + v, ok := ctx.Get(CtxKeyAPIToken) + if !ok { + return true + } + tok, _ := v.(APITokenAccessor) + if tok == nil { + return true + } + switch c.Cover { + case CronCoverAll: + return DenyListSafeForLimitedPAT(tok, c.GetUserID(), c.Servers) + default: + for _, id := range c.Servers { + if !tok.CanAccessServer(id) { + return false + } + } + return true + } +} diff --git a/model/cron_admin_pat_whitelist_test.go b/model/cron_admin_pat_whitelist_test.go new file mode 100644 index 00000000..941dcd1c --- /dev/null +++ b/model/cron_admin_pat_whitelist_test.go @@ -0,0 +1,116 @@ +package model + +import ( + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" +) + +// C1 regression: an admin-owned CronCoverAll fans out to every server in the +// system at runtime (CronTrigger gates on userIsAdmin(cr.UserID)). A +// server-limited PAT created by that admin must therefore only pass +// HasPermission when its deny-list covers EVERY server outside its whitelist +// system-wide — not just the admin's own servers, which is the visibly +// degenerate set that OwnerServerIDsLookup returns today. +// +// Without this regression, an admin with a PAT scoped to server X can create +// a CronCoverAll cron with deny-list = [X] and the dashboard will cheerfully +// dispatch the command to every OTHER user's server. +func TestCronHasPermission_AdminOwnerCoverAllDeniesUntilDenyListCoversAllOtherServers(t *testing.T) { + saved := OwnerServerIDsLookup + savedAdmin := OwnerIsAdminLookup + savedAll := AllServerIDsLookup + t.Cleanup(func() { + OwnerServerIDsLookup = saved + OwnerIsAdminLookup = savedAdmin + AllServerIDsLookup = savedAll + }) + + // Admin (uid=1) owns only server 1. Member (uid=200) owns server 2. + OwnerServerIDsLookup = func(ownerUID uint64) []uint64 { + switch ownerUID { + case 1: + return []uint64{1} + case 200: + return []uint64{2} + } + return nil + } + OwnerIsAdminLookup = func(uid uint64) bool { return uid == 1 } + AllServerIDsLookup = func() []uint64 { return []uint64{1, 2} } + + // Limited PAT (admin's) — scoped to server 1 only. + pat := &stubPATAccessor{ids: []uint64{1}} + + t.Run("deny_list_missing_other_owner_server_must_reject", func(t *testing.T) { + cron := &Cron{ + Common: Common{ID: 9, UserID: 1}, // admin-owned + Cover: CronCoverAll, + Servers: []uint64{1}, // deny self, but NOT server 2 (member's) + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 1}, Role: RoleAdmin}) + ctx.Set(CtxKeyAPIToken, pat) + + if cron.HasPermission(ctx) { + t.Fatal("admin-owned CronCoverAll fans out to ALL servers at runtime; " + + "deny-list missing server 2 must reject a PAT scoped to [1] — otherwise " + + "the cron will execute on a foreign user's server") + } + }) + + t.Run("deny_list_covers_all_other_servers_passes", func(t *testing.T) { + cron := &Cron{ + Common: Common{ID: 10, UserID: 1}, + Cover: CronCoverAll, + Servers: []uint64{2}, // deny the only server outside PAT whitelist + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 1}, Role: RoleAdmin}) + ctx.Set(CtxKeyAPIToken, pat) + + if !cron.HasPermission(ctx) { + t.Fatal("deny-list covering every non-whitelisted server must pass: fan-out " + + "is now contained inside the PAT whitelist") + } + }) +} + +// Companion: member-owned crons must NOT use the system-wide fan-out set. +// Runtime CronTrigger only ships to servers whose UserID matches the +// member-owner; HasPermission must mirror that to avoid false rejects on +// completely legitimate configs. +func TestCronHasPermission_MemberOwnerCoverAllStillUsesOwnerSet(t *testing.T) { + saved := OwnerServerIDsLookup + savedAdmin := OwnerIsAdminLookup + savedAll := AllServerIDsLookup + t.Cleanup(func() { + OwnerServerIDsLookup = saved + OwnerIsAdminLookup = savedAdmin + AllServerIDsLookup = savedAll + }) + + OwnerServerIDsLookup = func(ownerUID uint64) []uint64 { + if ownerUID == 100 { + return []uint64{1} + } + return nil + } + OwnerIsAdminLookup = func(uid uint64) bool { return false } + AllServerIDsLookup = func() []uint64 { return []uint64{1, 2, 3} } + + cron := &Cron{ + Common: Common{ID: 9, UserID: 100}, + Cover: CronCoverAll, + Servers: []uint64{}, // empty deny-list: fan-out = owner-set = [1] + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + // PAT whitelist = [1] — covers everything the member owner can fan out to. + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) + + if !cron.HasPermission(ctx) { + t.Fatal("member-owned CronCoverAll fans out to owner servers only; PAT [1] covers them all") + } +} diff --git a/model/cron_pat_whitelist_test.go b/model/cron_pat_whitelist_test.go new file mode 100644 index 00000000..f2702d44 --- /dev/null +++ b/model/cron_pat_whitelist_test.go @@ -0,0 +1,102 @@ +package model + +import ( + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" +) + +// stubPATAccessor 是只在测试里用的最小 APITokenAccessor,仅按 ids +// 字面包含判断。够用就行,不引入 *APIToken 在 model 包里转译 CSV。 +type stubPATAccessor struct { + ids []uint64 +} + +func (s *stubPATAccessor) CanAccessServer(id uint64) bool { + for _, x := range s.ids { + if x == id { + return true + } + } + return false +} + +// ServerIDs 暴露白名单,使 DenyListSafeForLimitedPAT 能区分「unscoped PAT」 +// 与「server-limited PAT」;缺这个方法时所有 stub 都会被当作不受限放行。 +func (s *stubPATAccessor) ServerIDs() []uint64 { + return s.ids +} + +// 钉死「server-limited PAT 不能通过 cover-all + 空 Servers 越过白名单」。 +// 老实现在 len(c.Servers)==0 时直接放行,但 CronCoverAll + 空 Servers 在 +// CronTrigger 里会 fan out 到 owner 的所有 server(包含白名单外的)。 +// HasPermission 是 cron 列表/手动触发/删除路径上唯一的 PAT 收口, +// 因此这里必须拒绝。 +func TestCronHasPermission_DeniesCoverAllEmptyServersForLimitedPAT(t *testing.T) { + cron := &Cron{ + Common: Common{ID: 9, UserID: 100}, + Cover: CronCoverAll, + Servers: nil, + } + + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) + + if cron.HasPermission(ctx) { + t.Fatal("server-limited PAT must not be allowed to operate on a CronCoverAll cron with empty Servers") + } +} + +// CoverIgnoreAll + 空 Servers 在 CronTrigger 里是 “allow-list of zero”, +// 不会 fan out。允许 PAT 继续看到/触发是无害的,但 HasPermission 的 +// 老语义在这一组合下仍是 return true,所以这条测试是「保持现状」的金线, +// 防止未来收紧时把这一无害情况也误拒。 +func TestCronHasPermission_AllowsCoverIgnoreAllEmptyServersForLimitedPAT(t *testing.T) { + cron := &Cron{ + Common: Common{ID: 10, UserID: 100}, + Cover: CronCoverIgnoreAll, + Servers: nil, + } + + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) + + if !cron.HasPermission(ctx) { + t.Fatal("CronCoverIgnoreAll + empty Servers is a no-op cron; server-limited PAT must remain allowed") + } +} + +// 现有 non-empty Servers 路径必须保持不变:白名单内允许、白名单外拒绝。 +// 这条用例钉死「修复 cover-all 路径时不能误改这条已有的金线」。 +func TestCronHasPermission_KeepsExistingNonEmptyServersSemantics(t *testing.T) { + t.Run("whitelisted", func(t *testing.T) { + cron := &Cron{ + Common: Common{ID: 11, UserID: 100}, + Cover: CronCoverIgnoreAll, + Servers: []uint64{1}, + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) + if !cron.HasPermission(ctx) { + t.Fatal("cron bound to whitelisted server 1 must remain accessible to PAT [1]") + } + }) + + t.Run("outside whitelist", func(t *testing.T) { + cron := &Cron{ + Common: Common{ID: 12, UserID: 100}, + Cover: CronCoverIgnoreAll, + Servers: []uint64{2}, + } + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) + if cron.HasPermission(ctx) { + t.Fatal("cron bound to non-whitelisted server 2 must be rejected for PAT [1]") + } + }) +} diff --git a/model/mcp_audit.go b/model/mcp_audit.go new file mode 100644 index 00000000..e4d57e13 --- /dev/null +++ b/model/mcp_audit.go @@ -0,0 +1,42 @@ +package model + +import "time" + +// MCPAuditLog 记录每一次 MCP tool 调用,用于事后追责与异常检测。 +// 写入是 best-effort:失败仅打日志,不阻塞业务请求。 +type MCPAuditLog struct { + ID uint64 `gorm:"primaryKey" json:"id"` + CreatedAt time.Time `gorm:"index" json:"created_at"` + UserID uint64 `gorm:"index" json:"user_id"` + TokenID uint64 `gorm:"index" json:"token_id"` + Tool string `gorm:"type:varchar(64);index" json:"tool"` + ServerID uint64 `gorm:"index" json:"server_id,omitempty"` + ArgsHash string `gorm:"type:char(64)" json:"args_hash"` + ArgsPeek string `gorm:"type:varchar(512)" json:"args_peek,omitempty"` + Outcome string `gorm:"type:varchar(32);index" json:"outcome"` + ErrorCode string `gorm:"type:varchar(32)" json:"error_code,omitempty"` + ErrorMsg string `gorm:"type:varchar(512)" json:"error_msg,omitempty"` + DurationMs int64 `json:"duration_ms"` + IP string `gorm:"type:varchar(64)" json:"ip"` +} + +func (MCPAuditLog) TableName() string { + return "mcp_audit_logs" +} + +const ( + MCPOutcomeOK = "ok" + MCPOutcomeScopeDenied = "scope_denied" + MCPOutcomePermDenied = "permission_denied" + MCPOutcomeServerOffline = "server_offline" + MCPOutcomeAgentTimeout = "agent_timeout" + MCPOutcomeAgentError = "agent_error" + // MCPOutcomeMCPDisabled 区分 “管理员按下 kill switch 把 MCP 关了” 与 + // “agent 真出故障” 两种语义:前者属 forbidden 类、不应该触发 agent + // 故障告警;详见 service/rpc/mcp_rpc.go 里 ErrMCPDisabled 的注释。 + MCPOutcomeMCPDisabled = "mcp_disabled" + MCPOutcomeInvalidArgs = "invalid_args" + MCPOutcomeRateLimited = "rate_limited" + MCPOutcomeUnsupportedAgent = "unsupported_agent" + MCPOutcomeInternalError = "internal_error" +) diff --git a/model/mcp_enabled_atomic_test.go b/model/mcp_enabled_atomic_test.go new file mode 100644 index 00000000..c2780778 --- /dev/null +++ b/model/mcp_enabled_atomic_test.go @@ -0,0 +1,39 @@ +package model + +import "testing" + +func TestSetMCPEnabledUsesAtomicAsSourceOfTruth(t *testing.T) { + c := &Config{} + if c.MCPEnabled() { + t.Fatal("zero-value Config must report MCP disabled") + } + c.SetMCPEnabled(true) + if !c.MCPEnabled() { + t.Fatal("MCPEnabled() must observe SetMCPEnabled(true)") + } + c.SetMCPEnabled(false) + if c.MCPEnabled() { + t.Fatal("MCPEnabled() must observe SetMCPEnabled(false)") + } +} + +// save() 在 marshal 前从 atomic 同步明文 EnableMCP 字段,因此持久化/JSON 仍拿到 +// 正确值;运行时 SetMCPEnabled 不直接写该字段以避免与 listConfig 的整体拷贝竞争。 +func TestSaveSyncsEnableMCPFieldFromAtomic(t *testing.T) { + c := &Config{} + c.filePath = t.TempDir() + "/config.yaml" + c.SetMCPEnabled(true) + if err := c.save(); err != nil { + t.Fatalf("save: %v", err) + } + if !c.EnableMCP { + t.Fatal("save() must sync EnableMCP field from the atomic mirror for persistence") + } + c.SetMCPEnabled(false) + if err := c.save(); err != nil { + t.Fatalf("save: %v", err) + } + if c.EnableMCP { + t.Fatal("save() must clear EnableMCP field when the atomic mirror is false") + } +} diff --git a/model/nat.go b/model/nat.go index c5cbec87..e395c68e 100644 --- a/model/nat.go +++ b/model/nat.go @@ -1,5 +1,7 @@ package model +import "github.com/gin-gonic/gin" + type NAT struct { Common Enabled bool `json:"enabled"` @@ -8,3 +10,20 @@ type NAT struct { Host string `json:"host"` Domain string `json:"domain" gorm:"unique"` } + +// HasPermission 在 owner/admin 之上叠加 PAT 的 server_ids 白名单, +// 与 Server/Service/Cron.HasPermission 一致,避免 server-limited PAT 越权。 +func (n *NAT) HasPermission(ctx *gin.Context) bool { + if !n.Common.HasPermission(ctx) { + return false + } + v, ok := ctx.Get(CtxKeyAPIToken) + if !ok { + return true + } + tok, ok := v.(APITokenAccessor) + if !ok || tok == nil { + return true + } + return tok.CanAccessServer(n.ServerID) +} diff --git a/model/nat_pat_whitelist_test.go b/model/nat_pat_whitelist_test.go new file mode 100644 index 00000000..b7214216 --- /dev/null +++ b/model/nat_pat_whitelist_test.go @@ -0,0 +1,55 @@ +package model + +import ( + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" +) + +// Regression: NAT had no HasPermission override, so it fell back to +// Common.HasPermission (owner/admin only). listHandler/CheckPermission gate +// on NAT.HasPermission, meaning a server-limited PAT could list/update/delete +// NAT records bound to off-whitelist servers of the same owner. NAT now +// applies CanAccessServer(NAT.ServerID) like Server/Service/Cron. +func TestNATHasPermission_DeniesOffWhitelistServerForLimitedPAT(t *testing.T) { + nat := &NAT{Common: Common{ID: 1, UserID: 100}, ServerID: 2} + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) // whitelist excludes server 2 + + if nat.HasPermission(ctx) { + t.Fatal("server-limited PAT must not reach a NAT bound to an off-whitelist server") + } +} + +func TestNATHasPermission_AllowsWhitelistedServerForLimitedPAT(t *testing.T) { + nat := &NAT{Common: Common{ID: 1, UserID: 100}, ServerID: 1} + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + ctx.Set(CtxKeyAPIToken, &stubPATAccessor{ids: []uint64{1}}) + + if !nat.HasPermission(ctx) { + t.Fatal("PAT whitelisted to the NAT's server must be allowed") + } +} + +func TestNATHasPermission_NoPATPassesViaCommonHasPermission(t *testing.T) { + nat := &NAT{Common: Common{ID: 1, UserID: 100}, ServerID: 2} + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 100}, Role: RoleMember}) + + if !nat.HasPermission(ctx) { + t.Fatal("owner without PAT must keep the existing owner/admin pass") + } +} + +func TestNATHasPermission_DeniesNonOwner(t *testing.T) { + nat := &NAT{Common: Common{ID: 1, UserID: 100}, ServerID: 1} + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 200}, Role: RoleMember}) + + if nat.HasPermission(ctx) { + t.Fatal("a different non-admin user must not reach another owner's NAT") + } +} diff --git a/model/server.go b/model/server.go index 3d25c027..a892f226 100644 --- a/model/server.go +++ b/model/server.go @@ -1,11 +1,14 @@ package model import ( + "errors" "log" "slices" + "sync" "sync/atomic" "time" + "github.com/gin-gonic/gin" "github.com/goccy/go-json" "gorm.io/gorm" @@ -38,7 +41,12 @@ type Server struct { // handler that reassigns the stream on every reconnect — a torn read of the // two-word interface header would panic on a subsequent .Send call. The // atomic.Pointer + holder struct lets us swap the stream lock-free while - // every reader observes a single, consistent value. + // every reader observes a single, consistent value. The holder also carries + // the send mutex so CopyFromRunningServer can share it across the old/new + // *Server objects that briefly co-exist during edit/transfer rotations — + // otherwise two *Server pointers would hold the same gRPC stream behind + // two independent mutexes, defeating the "one SendMsg goroutine per stream" + // invariant grpc-go requires. taskStream atomic.Pointer[taskStreamHolder] ConfigCache chan any `gorm:"-" json:"-"` @@ -51,8 +59,14 @@ type Server struct { // field `TaskStream pb.NezhaService_RequestTaskServer` was a plain interface // value: two words on the heap (type ptr + data ptr). Concurrent assignment // produced torn reads detectable by `go test -race` and crashable in production. +// +// sendMu lives on the holder (not on *Server) so it is bound to the stream +// itself: CopyFromRunningServer shares the same holder pointer with the new +// *Server, and SendTask locks via the holder, guaranteeing serialized SendMsg +// even when old/new *Server objects briefly co-exist during edit/transfer. type taskStreamHolder struct { - s pb.NezhaService_RequestTaskServer + s pb.NezhaService_RequestTaskServer + sendMu sync.Mutex } // SetTaskStream publishes the agent's RequestTask stream so other goroutines @@ -65,6 +79,13 @@ func (s *Server) SetTaskStream(stream pb.NezhaService_RequestTaskServer) { s.taskStream.Store(&taskStreamHolder{s: stream}) } +// adoptTaskStreamHolder publishes an existing holder verbatim. Used by +// CopyFromRunningServer so the new *Server shares the send mutex (and the +// underlying stream identity) with the old *Server. +func (s *Server) adoptTaskStreamHolder(h *taskStreamHolder) { + s.taskStream.Store(h) +} + // ClearTaskStreamIfCurrent detaches stream only if it is still the published // RequestTask stream. Disconnect cleanup uses this guard so an old stream // returning after a reconnect cannot erase the newer live stream. @@ -95,6 +116,31 @@ func (s *Server) GetTaskStream() pb.NezhaService_RequestTaskServer { return h.s } +// SendTask dispatches a task on the agent's RequestTask stream under the +// holder's sendMu so concurrent dispatchers (cron, server-transfer +// ApplyConfig, MCP CallAgent, MCP fs.transfer, force-update, report-config) +// cannot violate grpc-go's "one SendMsg goroutine per stream" rule. Returns +// ErrTaskStreamOffline if the agent has not published a stream yet; callers +// that need to distinguish offline from send failure should branch on that. +// +// The mutex is keyed by holder (= by stream) rather than by *Server so that +// edit/transfer rotations replacing *Server in the singleton map still share +// a single lock across the old and new objects pointing at the same stream. +func (s *Server) SendTask(task *pb.Task) error { + h := s.taskStream.Load() + if h == nil { + return ErrTaskStreamOffline + } + h.sendMu.Lock() + defer h.sendMu.Unlock() + return h.s.Send(task) +} + +// ErrTaskStreamOffline is returned by SendTask when the agent has no +// published RequestTask stream. Defined here (rather than in service/rpc) +// so model-layer callers can branch on it without an import cycle. +var ErrTaskStreamOffline = errors.New("agent task stream offline") + func InitServer(s *Server) { s.Host = &Host{} s.State = &HostState{} @@ -107,9 +153,12 @@ func (s *Server) CopyFromRunningServer(old *Server) { s.State = old.State s.GeoIP = old.GeoIP s.LastActive = old.LastActive - // taskStream is an atomic.Pointer; copy the published value rather than - // the field itself (atomic.Pointer is not safe to copy by value). - s.SetTaskStream(old.GetTaskStream()) + // Adopt the holder pointer verbatim so the new *Server shares the send + // mutex AND the stream identity with the old *Server; constructing a fresh + // holder via SetTaskStream(GetTaskStream()) would give the new object its + // own mutex, letting two *Server pointers race SendMsg on the same stream + // during the edit/transfer rotation window. + s.adoptTaskStreamHolder(old.taskStream.Load()) s.ConfigCache = old.ConfigCache s.PrevTransferInSnapshot = old.PrevTransferInSnapshot s.PrevTransferOutSnapshot = old.PrevTransferOutSnapshot @@ -147,6 +196,39 @@ type ServerOwnerInfo struct { // in tests / headless contexts so the JSON simply omits the owner field. var ServerOwnerLookup func(uid uint64) (ServerOwnerInfo, bool) +// OwnerServerIDsLookup is installed by singleton at startup to enumerate the +// IDs of every in-memory Server whose UserID == ownerUID. It exists so that +// Cron.HasPermission / Service.HasPermission can faithfully replay the +// dispatch-side "CoverAll deny-list must cover every PAT-whitelisted-out +// owner server" rule without depending on controller helpers (model must +// not import service/singleton — cycle). +// +// Left nil in tests / headless contexts; callers MUST treat a nil hook as +// "unknown owner topology" and fall back to a conservative decision (the +// existing model.Cron / model.Service code rejects non-trivial CoverAll +// configs for limited PATs when the hook is nil, matching the historical +// behaviour for empty deny-lists). +var OwnerServerIDsLookup func(ownerUID uint64) []uint64 + +// OwnerIsAdminLookup reports whether ownerUID is an admin user. When the +// owner is admin the runtime dispatch path (CronTrigger, DispatchTask) gates +// on userIsAdmin(cr.UserID) / userIsAdmin(svc.UserID) and fans out across +// EVERY in-memory server — not just the owner's. DenyListSafeForLimitedPAT +// must mirror that fan-out widening or a limited PAT can pass safety check +// with a deny-list that covers only the admin's own servers while the +// runtime still ships the task to foreign-owned servers. +// +// Left nil in tests / headless contexts; callers fall back to +// "owner-set only" which matches the pre-C1 behaviour. +var OwnerIsAdminLookup func(ownerUID uint64) bool + +// AllServerIDsLookup returns every in-memory server ID, regardless of +// owner. It is the system-wide fan-out set the runtime uses for +// admin-owned CoverAll cron/service dispatch and is the only correct +// containment set for a server-limited PAT operating on an admin-owned +// resource. Left nil in tests / headless contexts. +var AllServerIDsLookup func() []uint64 + type serverJSON Server type serverWithOwner struct { @@ -175,6 +257,87 @@ func (s *Server) MarshalJSON() ([]byte, error) { }) } +func (s *Server) HasPermission(ctx *gin.Context) bool { + if !s.Common.HasPermission(ctx) { + return false + } + v, ok := ctx.Get(CtxKeyAPIToken) + if !ok { + return true + } + tok, ok := v.(APITokenAccessor) + if !ok || tok == nil { + return true + } + return tok.CanAccessServer(s.GetID()) +} + +// APITokenWhitelistView is the optional shape an APITokenAccessor can +// implement so DenyListSafeForLimitedPAT can tell unscoped PATs (no +// whitelist → not limited) apart from server-limited ones. Accessors that +// do NOT expose ServerIDs() are treated as potentially limited; the safe +// dispatch path then requires denyList to cover every owner-visible server +// outside what the PAT can reach. +type APITokenWhitelistView interface { + ServerIDs() []uint64 +} + +// DenyListSafeForLimitedPAT reports whether a CoverAll/SkipServers deny-list +// keeps a server-limited PAT inside its server_ids whitelist. The runtime +// dispatch path (CronTrigger, DispatchTask) iterates every owner-visible +// server minus denyList; for the PAT to stay contained, every owner server +// outside its whitelist must already appear in denyList. JWT requests and +// PATs with no whitelist are unaffected. Nil OwnerServerIDsLookup forces +// the conservative "reject" branch instead of silently allowing a config +// the runtime would dispatch outside the whitelist. +func DenyListSafeForLimitedPAT(tok APITokenAccessor, ownerUID uint64, denyServers []uint64) bool { + if tok == nil { + return true + } + if wl, ok := tok.(APITokenWhitelistView); ok && len(wl.ServerIDs()) == 0 { + return true + } + fanout := ownerEffectiveFanoutServerIDs(ownerUID) + if fanout == nil { + return false + } + denySet := make(map[uint64]struct{}, len(denyServers)) + for _, id := range denyServers { + denySet[id] = struct{}{} + } + for _, id := range fanout { + if tok.CanAccessServer(id) { + continue + } + if _, denied := denySet[id]; !denied { + return false + } + } + return true +} + +// ownerEffectiveFanoutServerIDs returns the server set the runtime dispatch +// will actually fan out to for a resource owned by ownerUID. Admin owners +// short-circuit cronCanSendToServer / canSendServiceTask via userIsAdmin, +// so the safe containment set is the WHOLE system, not just the admin's +// own servers. Member owners stay bounded to their own server set. +// +// Returns nil to signal "topology unknown" — callers (DenyListSafeForLimitedPAT) +// fall back to fail-closed in that case, matching the historical conservative +// branch when OwnerServerIDsLookup was nil. +func ownerEffectiveFanoutServerIDs(ownerUID uint64) []uint64 { + if OwnerIsAdminLookup != nil && OwnerIsAdminLookup(ownerUID) { + if AllServerIDsLookup == nil { + return nil + } + return AllServerIDsLookup() + } + if OwnerServerIDsLookup == nil { + return nil + } + return OwnerServerIDsLookup(ownerUID) +} + func (s *Server) SplitList(x []*Server) ([]*Server, []*Server) { pri := func(s *Server) bool { return s.DisplayIndex == 0 diff --git a/model/server_transfer.go b/model/server_transfer.go index 90e49410..9ef0cb9d 100644 --- a/model/server_transfer.go +++ b/model/server_transfer.go @@ -81,11 +81,21 @@ type ServerTransfer struct { // admins, the source user, the destination user, and the initiator. Listing // uses this to filter what the caller can see; mutating endpoints (cancel, // retry) layer additional checks on top. +// +// PAT server_ids whitelist is evaluated FIRST, before the admin short- +// circuit, so an admin-issued PAT scoped to a subset of servers cannot +// widen reach by virtue of the caller being an admin. JWT callers (no PAT +// in context) skip the whitelist check. func (t *ServerTransfer) HasPermission(ctx *gin.Context) bool { auth, ok := ctx.Get(CtxKeyAuthorizedUser) if !ok { return false } + if v, ok := ctx.Get(CtxKeyAPIToken); ok { + if tok, _ := v.(APITokenAccessor); tok != nil && !tok.CanAccessServer(t.ServerID) { + return false + } + } user := *auth.(*User) if user.Role == RoleAdmin { return true diff --git a/model/service.go b/model/service.go index 26627205..cac80ac0 100644 --- a/model/service.go +++ b/model/service.go @@ -4,6 +4,7 @@ import ( "fmt" "log" + "github.com/gin-gonic/gin" "github.com/goccy/go-json" "github.com/robfig/cron/v3" "gorm.io/gorm" @@ -30,6 +31,158 @@ const ( // Pre-transfer agents do not recognise this type — dashboard MUST gate // transfers on agent capability before pushing. TaskTypeServerTransferApply + TaskTypeExec + TaskTypeFsList + TaskTypeFsRead + TaskTypeFsWrite + TaskTypeFsDelete + TaskTypeFsTransfer +) + +// IsMCPRPCResult 判定一个 TaskResult.Type 是否属于 MCP 走 RequestTask 通道的 +// 一次性 RPC 类型。dashboard 的 RequestTask 接收循环用它把这些回包路由到 +// Server.inflightRPC 等待方,而不是走 ServiceSentinel。 +// +// TaskTypeFsTransfer 走 IOStream 而不是 RequestTask 回包,故不在此列;agent +// 不会对它发 TaskResult。 +func IsMCPRPCResult(t uint64) bool { + switch t { + case TaskTypeExec, TaskTypeFsList, TaskTypeFsRead, TaskTypeFsWrite, TaskTypeFsDelete: + return true + } + return false +} + +// ExecRequest 是 server.exec 通过 Task.Data 下发到 agent 的载荷(JSON)。 +type ExecRequest struct { + 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"` +} + +// ExecResult 是 agent 通过 TaskResult.Data 回传的执行结果(JSON)。 +type ExecResult struct { + ExitCode int `json:"exit_code"` + Stdout string `json:"stdout"` + Stderr string `json:"stderr"` + DurationMs int64 `json:"duration_ms"` + StdoutTruncated bool `json:"stdout_truncated,omitempty"` + StderrTruncated bool `json:"stderr_truncated,omitempty"` + TimedOut bool `json:"timed_out,omitempty"` + Error string `json:"error,omitempty"` +} + +// FsListRequest fs.list 下发载荷。 +type FsListRequest struct { + Path string `json:"path"` + ShowHidden bool `json:"show_hidden,omitempty"` +} + +// FsEntry 单条目录元数据。 +type FsEntry struct { + Name string `json:"name"` + Type string `json:"type"` + Size int64 `json:"size"` + Mode string `json:"mode"` + ModTimeUnix int64 `json:"mtime"` + IsSymlink bool `json:"is_symlink,omitempty"` + LinkTarget string `json:"link_target,omitempty"` +} + +// FsListResult fs.list 回包。 +type FsListResult struct { + Entries []FsEntry `json:"entries"` + Truncated bool `json:"truncated,omitempty"` + Total int `json:"total,omitempty"` + Error string `json:"error,omitempty"` +} + +// FsReadRequest fs.read 下发载荷。Offset/Length 单位为字节;encoding 控制返回。 +type FsReadRequest struct { + Path string `json:"path"` + Offset int64 `json:"offset,omitempty"` + Length int64 `json:"length,omitempty"` + Encoding string `json:"encoding,omitempty"` +} + +// FsReadResult fs.read 回包。Content 按 encoding 编码(utf8 原文 / base64 二进制安全)。 +type FsReadResult struct { + Content string `json:"content"` + Encoding string `json:"encoding"` + Size int64 `json:"size"` + SHA256 string `json:"sha256,omitempty"` + Truncated bool `json:"truncated,omitempty"` + Error string `json:"error,omitempty"` +} + +// FsWriteRequest fs.write 下发载荷。Mode 用 unix 数字字符串如 "0644"。 +type FsWriteRequest struct { + 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"` +} + +// FsWriteResult fs.write 回包。 +type FsWriteResult struct { + Size int64 `json:"size"` + SHA256 string `json:"sha256"` + Error string `json:"error,omitempty"` +} + +// FsDeleteRequest fs.delete 下发载荷。 +type FsDeleteRequest struct { + Path string `json:"path"` + Recursive bool `json:"recursive,omitempty"` +} + +// FsDeleteResult fs.delete 回包。 +type FsDeleteResult struct { + DeletedCount int `json:"deleted_count"` + Error string `json:"error,omitempty"` +} + +const ( + // MCPFsTransferOpUpload / Download 区分 IOStream 内的数据流向。 + MCPFsTransferOpUpload = "upload" + MCPFsTransferOpDownload = "download" + + // MCPFsTransferMaxSize 单次传输硬上限,dashboard 和 agent 双方都拒绝 + // 超出大小的请求。设为 100MiB 与产品语义"~100MB 大文件"对齐。 + MCPFsTransferMaxSize = 100 * 1024 * 1024 +) + +// FsTransferRequest 通过 Task.Data 下发到 agent;agent 据此打开本地 +// IOStream,按 op 完成上/下行。streamId 用于 agent IOStream 引导帧。 +type FsTransferRequest struct { + StreamID string `json:"stream_id"` + Op string `json:"op"` + Path string `json:"path"` + Size int64 `json:"size,omitempty"` + Mode string `json:"mode,omitempty"` + CreateDirs bool `json:"create_dirs,omitempty"` + IfMatchSHA256 string `json:"if_match_sha256,omitempty"` + ExpectedSHA256 string `json:"expected_sha256,omitempty"` +} + +// 双向 IOStream 控制帧 magic(每帧第一帧的前 4 字节)。数据帧不带 magic。 +// +// 这套 magic 与 FM 协议(NZTD/NZFN/NERR/NZUP)共存而不冲突:NZTD 在 FM 表示 +// "file header",在 transfer 表示"download header",但两条协议通过不同的 +// task type(TaskTypeFM vs TaskTypeFsTransfer)分流,不会复用同一个 agent +// goroutine,所以 magic 撞名只是字面巧合,不会破坏解析。 +var ( + MCPFsXferMagicUploadHdr = []byte{0x4E, 0x5A, 0x54, 0x55} // NZTU + MCPFsXferMagicDownloadHdr = []byte{0x4E, 0x5A, 0x54, 0x44} // NZTD + MCPFsXferMagicOK = []byte{0x4E, 0x5A, 0x54, 0x4F} // NZTO + MCPFsXferMagicErr = []byte{0x4E, 0x5A, 0x54, 0x45} // NZTE + MCPFsXferMagicChunk = []byte{0x4E, 0x5A, 0x54, 0x43} // NZTC: download data chunk ) type TerminalTask struct { @@ -86,6 +239,57 @@ func (m *Service) PB() *pb.Task { } } +// HasPermission 扩展默认的 owner/admin 检查,让 PAT 的 server_ids 白名单 +// 同样能收窄 service monitor 的列出/删除/更新路径,语义与 Cron.HasPermission +// 对齐: +// - ServiceCoverAll:SkipServers 是 deny-set。DispatchTask 会探测 owner 在 +// deny-set 之外的所有 server,所以受限 PAT 必须保证 deny-set 已经覆盖 +// 白名单外的全部 owner servers。判定与 controller 的 +// enforcePATServiceDispatchScope / rejectImplicitServiceCoverForLimitedPAT +// 共用 denyListSafeForLimitedPAT。 +// - ServiceCoverIgnoreAll:SkipServers 是 allow-set,要求每个被覆盖的 +// server 都在 PAT 白名单内。 +// - 其它情况保留旧的“PAT 按 owner 关系判定”行为。 +func (m *Service) HasPermission(ctx *gin.Context) bool { + if !m.Common.HasPermission(ctx) { + return false + } + v, ok := ctx.Get(CtxKeyAPIToken) + if !ok { + return true + } + tok, _ := v.(APITokenAccessor) + if tok == nil { + return true + } + switch m.Cover { + case ServiceCoverAll: + return DenyListSafeForLimitedPAT(tok, m.GetUserID(), skipServersTrueIDs(m.SkipServers)) + case ServiceCoverIgnoreAll: + for _, id := range skipServersTrueIDs(m.SkipServers) { + if !tok.CanAccessServer(id) { + return false + } + } + return true + default: + return true + } +} + +func skipServersTrueIDs(skip map[uint64]bool) []uint64 { + if len(skip) == 0 { + return nil + } + out := make([]uint64, 0, len(skip)) + for id, mark := range skip { + if mark { + out = append(out, id) + } + } + return out +} + // CronSpec 返回服务监控请求间隔对应的 cron 表达式 func (m *Service) CronSpec() string { if m.Duration == 0 { diff --git a/model/service_ignoreall_false_test.go b/model/service_ignoreall_false_test.go new file mode 100644 index 00000000..07fc8087 --- /dev/null +++ b/model/service_ignoreall_false_test.go @@ -0,0 +1,51 @@ +package model + +import ( + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" +) + +func adminPATCtx(tok *APIToken) *gin.Context { + gin.SetMode(gin.TestMode) + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest("GET", "/", nil) + c.Set(CtxKeyAuthorizedUser, &User{Common: Common{ID: 1}, Role: RoleAdmin}) + c.Set(CtxKeyAPIToken, tok) + return c +} + +// ServiceCoverIgnoreAll permission must only consider SkipServers entries +// whose value is true (the allow-set actually dispatched at runtime). A +// `{2: false}` entry has no dispatch effect, so a PAT scoped to {1} must +// still be allowed to manage this service. +func TestServiceHasPermissionIgnoreAllSkipsFalseEntries(t *testing.T) { + tok := &APIToken{ID: 1, UserID: 1} + tok.SetServerIDs([]uint64{1}) + + svc := &Service{ + Common: Common{ID: 10, UserID: 1}, + Cover: ServiceCoverIgnoreAll, + SkipServers: map[uint64]bool{1: true, 2: false}, + } + + if !svc.HasPermission(adminPATCtx(tok)) { + t.Fatal("a `{2: false}` allow-set entry must not block a PAT scoped to {1}") + } +} + +func TestServiceHasPermissionIgnoreAllRejectsForeignTrueEntry(t *testing.T) { + tok := &APIToken{ID: 1, UserID: 1} + tok.SetServerIDs([]uint64{1}) + + svc := &Service{ + Common: Common{ID: 10, UserID: 1}, + Cover: ServiceCoverIgnoreAll, + SkipServers: map[uint64]bool{2: true}, + } + + if svc.HasPermission(adminPATCtx(tok)) { + t.Fatal("a true allow-set entry on server 2 must reject a PAT scoped to {1}") + } +} diff --git a/model/setting_api.go b/model/setting_api.go index d9ef7d91..9e136140 100644 --- a/model/setting_api.go +++ b/model/setting_api.go @@ -17,6 +17,7 @@ type SettingForm struct { AgentTLS bool `json:"tls,omitempty" validate:"optional"` EnableIPChangeNotification bool `json:"enable_ip_change_notification,omitempty" validate:"optional"` EnablePlainIPInNotification bool `json:"enable_plain_ip_in_notification,omitempty" validate:"optional"` + EnableMCP *bool `json:"enable_mcp,omitempty" validate:"optional"` } type Setting struct { diff --git a/pkg/grpcx/io_stream_wrapper.go b/pkg/grpcx/io_stream_wrapper.go index 44dce596..3103fdff 100644 --- a/pkg/grpcx/io_stream_wrapper.go +++ b/pkg/grpcx/io_stream_wrapper.go @@ -3,6 +3,7 @@ package grpcx import ( "context" "io" + "sync" "sync/atomic" "github.com/nezhahq/nezha/proto" @@ -16,8 +17,15 @@ type IOStream interface { Context() context.Context } +// IOStreamWrapper adapts a gRPC IOStream into an io.ReadWriteCloser and +// serializes every Send on the underlying stream. grpc-go forbids concurrent +// SendMsg on the same stream (Documentation/concurrency.md); the dashboard +// runs an IOStream keepalive goroutine alongside MCP fs.transfer / terminal / +// fm Writers, so all of them must funnel through this sendMu. The matching +// agent-side fix is serialIOStreamSender in agent/cmd/agent/mcp_fs_transfer.go. type IOStreamWrapper struct { IOStream + sendMu sync.Mutex dataBuf []byte closed *atomic.Bool closeCh chan struct{} @@ -31,21 +39,77 @@ func NewIOStreamWrapper(stream IOStream) *IOStreamWrapper { } } +// Send writes a single IOStreamData frame under the wrapper's send mutex. +// All goroutines that share this wrapper — keepalive ticker, Write callers, +// and any direct frame writer — MUST go through Send (or SendKeepalive) +// rather than touching the embedded IOStream.Send, otherwise grpc-go's +// concurrent-SendMsg invariant is violated and frames can corrupt or panic. +func (iw *IOStreamWrapper) Send(data *proto.IOStreamData) error { + iw.sendMu.Lock() + defer iw.sendMu.Unlock() + return iw.IOStream.Send(data) +} + +// SendKeepalive sends the dashboard's empty-payload heartbeat through the +// same sendMu as Send/Write so it cannot race the data path. +func (iw *IOStreamWrapper) SendKeepalive() error { + return iw.Send(&proto.IOStreamData{Data: []byte{}}) +} + +// RecvFrame returns the next non-empty IOStream frame as a single contiguous +// byte slice, preserving frame boundaries. Use this when a caller multiplexes +// control frames (magic + payload) and data frames over the same stream and +// must not let one frame's bytes spill into the next frame's parsing. +// +// The io.Reader path (Read) intentionally hides frame boundaries; callers that +// need them — e.g. MCP fs.transfer download where NZTE may interrupt NZTD +// payload mid-stream — call RecvFrame instead. +func (iw *IOStreamWrapper) RecvFrame() ([]byte, error) { + if len(iw.dataBuf) > 0 { + out := iw.dataBuf + iw.dataBuf = nil + return out, nil + } + for { + data, err := iw.Recv() + if err != nil { + return nil, err + } + if len(data.Data) == 0 { + continue + } + return data.Data, nil + } +} + func (iw *IOStreamWrapper) Read(p []byte) (n int, err error) { if len(iw.dataBuf) > 0 { n := copy(p, iw.dataBuf) iw.dataBuf = iw.dataBuf[n:] return n, nil } - var data *proto.IOStreamData - if data, err = iw.Recv(); err != nil { - return 0, err + // Skip zero-length heartbeat frames sent by ioStreamKeepAlive (see + // agent/cmd/agent/main.go ioStreamKeepAlive). protobuf treats an empty + // `bytes` field as a default value but still ships a valid Message, so + // Recv() returns a non-nil *IOStreamData whose Data is empty. Surfacing + // that as (0, nil) is legal io.Reader behaviour but every caller in the + // repo treats a 0-byte read as an unexpected control frame (e.g. + // mcp_transfer.readXferFixedHeader returns "frame too short"). Loop here + // until we get either real bytes or an error. + for { + var data *proto.IOStreamData + if data, err = iw.Recv(); err != nil { + return 0, err + } + if len(data.Data) == 0 { + continue + } + n = copy(p, data.Data) + if n < len(data.Data) { + iw.dataBuf = data.Data[n:] + } + return n, nil } - n = copy(p, data.Data) - if n < len(data.Data) { - iw.dataBuf = data.Data[n:] - } - return n, nil } func (iw *IOStreamWrapper) Write(p []byte) (n int, err error) { @@ -63,3 +127,12 @@ func (iw *IOStreamWrapper) Close() error { func (iw *IOStreamWrapper) Wait() { <-iw.closeCh } + +// Done exposes the wrapper's close signal as a read-only channel so callers +// that run alongside the wrapper (e.g. the dashboard's IOStream keepalive +// goroutine) can cancel cooperatively. Without this they would only stop on +// gRPC stream-context cancel or on their next failed Send, which can leave +// a goroutine waiting up to one keepalive tick after the wrapper was closed. +func (iw *IOStreamWrapper) Done() <-chan struct{} { + return iw.closeCh +} diff --git a/pkg/grpcx/io_stream_wrapper_concurrent_send_test.go b/pkg/grpcx/io_stream_wrapper_concurrent_send_test.go new file mode 100644 index 00000000..6e07659c --- /dev/null +++ b/pkg/grpcx/io_stream_wrapper_concurrent_send_test.go @@ -0,0 +1,71 @@ +package grpcx + +import ( + "context" + "sync" + "sync/atomic" + "testing" + + "github.com/nezhahq/nezha/proto" +) + +// sendObservingStream records the maximum number of goroutines that are +// inside Send at the same time. grpc-go's real server stream is NOT safe +// under concurrent Send, but this fake never blocks so any concurrent +// dispatch from the wrapper would surface here as maxInFlight > 1. +type sendObservingStream struct { + inFlight int32 + maxInFlight int32 +} + +func (s *sendObservingStream) Recv() (*proto.IOStreamData, error) { return nil, nil } +func (s *sendObservingStream) Context() context.Context { return context.Background() } +func (s *sendObservingStream) Send(*proto.IOStreamData) error { + cur := atomic.AddInt32(&s.inFlight, 1) + defer atomic.AddInt32(&s.inFlight, -1) + for { + prev := atomic.LoadInt32(&s.maxInFlight) + if cur <= prev || atomic.CompareAndSwapInt32(&s.maxInFlight, prev, cur) { + break + } + } + return nil +} + +// IOStreamWrapper.Send and SendKeepalive must be safe to call from many +// goroutines concurrently — this is the dashboard-side dual of the agent's +// serialIOStreamSender (see agent/cmd/agent/mcp_fs_transfer.go). Without the +// wrapper's sendMu, dashboard IOStream keepalive + MCP fs.transfer Write +// race the same gRPC stream, violating grpc-go's "no concurrent SendMsg" +// contract. We pin that with a stress test: many goroutines hammer Send / +// SendKeepalive / Write at once; the fake stream must NEVER observe more +// than one in-flight Send. +func TestIOStreamWrapper_SerializesConcurrentSends(t *testing.T) { + obs := &sendObservingStream{} + iw := NewIOStreamWrapper(obs) + + const workers = 16 + const opsPerWorker = 200 + var wg sync.WaitGroup + wg.Add(workers) + for i := 0; i < workers; i++ { + go func(seed int) { + defer wg.Done() + for j := 0; j < opsPerWorker; j++ { + switch (seed + j) % 3 { + case 0: + _ = iw.Send(&proto.IOStreamData{Data: []byte{byte(seed)}}) + case 1: + _ = iw.SendKeepalive() + case 2: + _, _ = iw.Write([]byte{byte(j)}) + } + } + }(i) + } + wg.Wait() + + if got := atomic.LoadInt32(&obs.maxInFlight); got != 1 { + t.Fatalf("IOStreamWrapper.Send must serialize through sendMu; observed max-in-flight=%d, want 1", got) + } +} diff --git a/pkg/grpcx/io_stream_wrapper_test.go b/pkg/grpcx/io_stream_wrapper_test.go new file mode 100644 index 00000000..b9d4913c --- /dev/null +++ b/pkg/grpcx/io_stream_wrapper_test.go @@ -0,0 +1,106 @@ +package grpcx + +import ( + "context" + "errors" + "io" + "testing" + "time" + + "github.com/nezhahq/nezha/proto" +) + +type fakeStream struct { + frames []*proto.IOStreamData + err error +} + +func (f *fakeStream) Recv() (*proto.IOStreamData, error) { + if len(f.frames) == 0 { + if f.err != nil { + return nil, f.err + } + return nil, io.EOF + } + frame := f.frames[0] + f.frames = f.frames[1:] + return frame, nil +} + +func (f *fakeStream) Send(*proto.IOStreamData) error { return nil } +func (f *fakeStream) Context() context.Context { return context.Background() } + +// Heartbeat frames sent by the agent (ioStreamKeepAlive in +// agent/cmd/agent/main.go) carry an empty Data. The previous wrapper +// surfaced them to callers as (n=0, nil), which made +// mcp_transfer.readXferFixedHeader return "frame too short". This test +// pins the contract: empty frames are transparently skipped and Read +// only returns when it has either real bytes or an error. +func TestIOStreamWrapper_ReadSkipsHeartbeats(t *testing.T) { + stream := &fakeStream{ + frames: []*proto.IOStreamData{ + {Data: []byte{}}, + {Data: []byte{}}, + {Data: []byte("hello")}, + }, + } + iw := NewIOStreamWrapper(stream) + buf := make([]byte, 16) + n, err := iw.Read(buf) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if n != 5 || string(buf[:n]) != "hello" { + t.Fatalf("expected 5 bytes 'hello', got n=%d data=%q", n, buf[:n]) + } +} + +// A stream that only ever sends heartbeats followed by an error must +// surface the error rather than spin forever or hand the caller (0, nil). +func TestIOStreamWrapper_ReadPropagatesErrorAfterHeartbeats(t *testing.T) { + wantErr := errors.New("stream closed") + stream := &fakeStream{ + frames: []*proto.IOStreamData{ + {Data: []byte{}}, + {Data: []byte{}}, + }, + err: wantErr, + } + iw := NewIOStreamWrapper(stream) + buf := make([]byte, 8) + n, err := iw.Read(buf) + if err == nil { + t.Fatalf("expected error after heartbeats + Recv failure") + } + if !errors.Is(err, wantErr) { + t.Fatalf("expected wrapped error %v, got %v", wantErr, err) + } + if n != 0 { + t.Fatalf("expected n=0 on error, got %d", n) + } +} + +// Close() must wake anything waiting on Done() immediately so co-running +// goroutines (e.g. the dashboard's IOStream keepalive ticker) can exit +// without waiting for the underlying gRPC stream context to cancel or for +// their next Send to fail. +func TestIOStreamWrapper_DoneFiresOnClose(t *testing.T) { + iw := NewIOStreamWrapper(&fakeStream{}) + select { + case <-iw.Done(): + t.Fatalf("Done() must not fire before Close()") + default: + } + if err := iw.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + select { + case <-iw.Done(): + case <-time.After(time.Second): + t.Fatalf("Done() did not fire after Close()") + } + // Idempotent: a second Close must not panic on the closed channel. + if err := iw.Close(); err != nil { + t.Fatalf("second Close: %v", err) + } +} diff --git a/service/rpc/io_stream.go b/service/rpc/io_stream.go index 6a61a298..3015850a 100644 --- a/service/rpc/io_stream.go +++ b/service/rpc/io_stream.go @@ -1,24 +1,39 @@ package rpc import ( + "context" "errors" "io" "sync" - "sync/atomic" "time" "github.com/nezhahq/nezha/service/singleton" ) +// StreamPurpose tags every IOStream with the feature that opened it so +// admin actions can drop only the relevant subset. Existing call sites +// (terminal / fm / NAT / server-transfer) keep PurposeLegacy and the +// previous semantics; only the new MCP fs.transfer path uses +// PurposeMCPTransfer, which is what EnableMCP=false revokes. +type StreamPurpose uint8 + +const ( + PurposeLegacy StreamPurpose = iota + PurposeMCPTransfer +) + type ioStreamContext struct { creatorUserID uint64 targetServerID uint64 + purpose StreamPurpose userIo io.ReadWriteCloser agentIo io.ReadWriteCloser userIoConnectCh chan struct{} agentIoConnectCh chan struct{} userIoChOnce sync.Once agentIoChOnce sync.Once + revokedCh chan struct{} + revokedOnce sync.Once } type bp struct { @@ -34,14 +49,20 @@ var bufPool = sync.Pool{ } func (s *NezhaHandler) CreateStream(streamId string, creatorUserID uint64, targetServerID uint64) { + s.CreateStreamWithPurpose(streamId, creatorUserID, targetServerID, PurposeLegacy) +} + +func (s *NezhaHandler) CreateStreamWithPurpose(streamId string, creatorUserID uint64, targetServerID uint64, purpose StreamPurpose) { s.ioStreamMutex.Lock() defer s.ioStreamMutex.Unlock() s.ioStreams[streamId] = &ioStreamContext{ creatorUserID: creatorUserID, targetServerID: targetServerID, + purpose: purpose, userIoConnectCh: make(chan struct{}), agentIoConnectCh: make(chan struct{}), + revokedCh: make(chan struct{}), } } @@ -64,6 +85,42 @@ func (s *NezhaHandler) IsStreamAuthorizedForAgent(streamId string, agentServerID return ctx.targetServerID != 0 && ctx.targetServerID == agentServerID } +// WaitForAgent 阻塞等待 agent 端通过 IOStream 接入并完成 AgentConnected。 +// dashboard 把 MCP 大文件传输的 task 派给 agent 后,需要等 agent dial 回来 +// 才能开始 Read/Write,这里以 timeout 内的轻量轮询暴露给 controller。 +// +// 同时返回 agent 端流(io.ReadWriteCloser)以便 controller 调 io.CopyN 转发 +// HTTP body;ok=false 表示超时或流已被关闭。 +func (s *NezhaHandler) WaitForAgent(ctx context.Context, streamId string, timeout time.Duration) (io.ReadWriteCloser, bool) { + deadline := time.NewTimer(timeout) + defer deadline.Stop() + for { + s.ioStreamMutex.RLock() + sc, ok := s.ioStreams[streamId] + if ok && sc.agentIo != nil { + s.ioStreamMutex.RUnlock() + return sc.agentIo, true + } + s.ioStreamMutex.RUnlock() + if !ok { + return nil, false + } + select { + case <-ctx.Done(): + return nil, false + case <-deadline.C: + return nil, false + case <-sc.revokedCh: + return nil, false + case <-sc.agentIoConnectCh: + s.ioStreamMutex.RLock() + ag := sc.agentIo + s.ioStreamMutex.RUnlock() + return ag, ag != nil + } + } +} + // IsStreamAuthorizedForUser checks whether the requesting user may attach to // the stream. A stream is reachable only by its creator or by an admin; any // other authenticated user must be rejected. Unknown streams are always @@ -108,6 +165,23 @@ func (s *NezhaHandler) StreamOwnership(streamId string) (uint64, bool) { return ctx.creatorUserID, true } +// StreamTarget returns the server ID the stream was opened against and +// whether the stream is still tracked. Callers MUST pass this through the +// requesting PAT's CanAccessServer check before allowing attachment — +// IsStreamAuthorizedForUser only knows about creator/admin, so without this +// dual gate an admin's server-limited PAT can hijack any stream by knowing +// the streamId. +func (s *NezhaHandler) StreamTarget(streamId string) (uint64, bool) { + s.ioStreamMutex.RLock() + defer s.ioStreamMutex.RUnlock() + + ctx, ok := s.ioStreams[streamId] + if !ok { + return 0, false + } + return ctx.targetServerID, true +} + func (s *NezhaHandler) GetStream(streamId string) (*ioStreamContext, error) { s.ioStreamMutex.RLock() defer s.ioStreamMutex.RUnlock() @@ -148,6 +222,37 @@ func (s *NezhaHandler) RevokeStreamsForServer(serverID uint64) { } } +// RevokeStreamsForPurpose tears down every IOStream tagged with the given +// purpose. Used as the IOStream half of the MCP kill switch: when the +// admin flips EnableMCP=false, any in-flight fs.transfer / fs.upload / +// fs.download must drop immediately rather than wait out the 5min IO +// timeout. Returns the number of streams revoked so the caller can log +// the blast radius. +func (s *NezhaHandler) RevokeStreamsForPurpose(purpose StreamPurpose) int { + s.ioStreamMutex.Lock() + defer s.ioStreamMutex.Unlock() + revoked := 0 + for streamId, ctx := range s.ioStreams { + if ctx.purpose != purpose { + continue + } + ctx.revokedOnce.Do(func() { + if ctx.revokedCh != nil { + close(ctx.revokedCh) + } + }) + if ctx.userIo != nil { + ctx.userIo.Close() + } + if ctx.agentIo != nil { + ctx.agentIo.Close() + } + delete(s.ioStreams, streamId) + revoked++ + } + return revoked +} + func (s *NezhaHandler) CloseStream(streamId string) error { s.ioStreamMutex.Lock() defer s.ioStreamMutex.Unlock() @@ -167,34 +272,50 @@ func (s *NezhaHandler) CloseStream(streamId string) error { +// UserConnected publishes the user-side IO under ioStreamMutex so concurrent +// Revoke* / WaitForAgent / StartStream see a consistent stream view. +// Without the lock, the bare assignment to stream.userIo races with the +// revoker's lock-protected read and triggers go-race. func (s *NezhaHandler) UserConnected(streamId string, userIo io.ReadWriteCloser) error { - stream, err := s.GetStream(streamId) - if err != nil { - return err + s.ioStreamMutex.Lock() + stream, ok := s.ioStreams[streamId] + if !ok { + s.ioStreamMutex.Unlock() + return errors.New("stream not found") } - stream.userIo = userIo + s.ioStreamMutex.Unlock() stream.userIoChOnce.Do(func() { close(stream.userIoConnectCh) }) - return nil } +// AgentConnected is the agent-side dual of UserConnected. Same locking +// rationale. func (s *NezhaHandler) AgentConnected(streamId string, agentIo io.ReadWriteCloser) error { - stream, err := s.GetStream(streamId) - if err != nil { - return err + s.ioStreamMutex.Lock() + stream, ok := s.ioStreams[streamId] + if !ok { + s.ioStreamMutex.Unlock() + return errors.New("stream not found") } - stream.agentIo = agentIo + s.ioStreamMutex.Unlock() stream.agentIoChOnce.Do(func() { close(stream.agentIoConnectCh) }) - return nil } +// streamEndpoints returns the user/agent IO under ioStreamMutex so callers +// never read the interface fields while UserConnected/AgentConnected write them. +func (s *NezhaHandler) streamEndpoints(stream *ioStreamContext) (userIo, agentIo io.ReadWriteCloser) { + s.ioStreamMutex.RLock() + defer s.ioStreamMutex.RUnlock() + return stream.userIo, stream.agentIo +} + func (s *NezhaHandler) StartStream(streamId string, timeout time.Duration) error { stream, err := s.GetStream(streamId) if err != nil { @@ -202,62 +323,50 @@ func (s *NezhaHandler) StartStream(streamId string, timeout time.Duration) error } timeoutTimer := time.NewTimer(timeout) + defer timeoutTimer.Stop() LOOP: for { select { case <-stream.userIoConnectCh: - if stream.agentIo != nil { - timeoutTimer.Stop() + if _, agentIo := s.streamEndpoints(stream); agentIo != nil { break LOOP } case <-stream.agentIoConnectCh: - if stream.userIo != nil { - timeoutTimer.Stop() + if userIo, _ := s.streamEndpoints(stream); userIo != nil { break LOOP } - case <-time.After(timeout): + case <-timeoutTimer.C: break LOOP } time.Sleep(time.Millisecond * 500) } - if stream.userIo == nil && stream.agentIo == nil { + userIo, agentIo := s.streamEndpoints(stream) + if userIo == nil && agentIo == nil { return singleton.Localizer.ErrorT("timeout: no connection established") } - if stream.userIo == nil { + if userIo == nil { return singleton.Localizer.ErrorT("timeout: user connection not established") } - if stream.agentIo == nil { + if agentIo == nil { return singleton.Localizer.ErrorT("timeout: agent connection not established") } - isDone := new(atomic.Bool) - endCh := make(chan struct{}) + errCh := make(chan error, 2) go func() { bp := bufPool.Get().(*bp) defer bufPool.Put(bp) - _, innerErr := io.CopyBuffer(stream.userIo, stream.agentIo, bp.buf) - if innerErr != nil { - err = innerErr - } - if isDone.CompareAndSwap(false, true) { - close(endCh) - } + _, innerErr := io.CopyBuffer(userIo, agentIo, bp.buf) + errCh <- innerErr }() go func() { bp := bufPool.Get().(*bp) defer bufPool.Put(bp) - _, innerErr := io.CopyBuffer(stream.agentIo, stream.userIo, bp.buf) - if innerErr != nil { - err = innerErr - } - if isDone.CompareAndSwap(false, true) { - close(endCh) - } + _, innerErr := io.CopyBuffer(agentIo, userIo, bp.buf) + errCh <- innerErr }() - <-endCh - return err + return <-errCh } diff --git a/service/rpc/io_stream_race_test.go b/service/rpc/io_stream_race_test.go new file mode 100644 index 00000000..e6fe8afe --- /dev/null +++ b/service/rpc/io_stream_race_test.go @@ -0,0 +1,92 @@ +package rpc + +import ( + "io" + "sync" + "testing" + "time" +) + +// nopRWC is a minimal io.ReadWriteCloser used for race tests; Close is a +// no-op so the racer goroutines do not panic on shared state. +type nopRWC struct{} + +func (nopRWC) Read(p []byte) (int, error) { return 0, io.EOF } +func (nopRWC) Write(p []byte) (int, error) { return len(p), nil } +func (nopRWC) Close() error { return nil } + +// H10 regression: UserConnected/AgentConnected mutate stream.userIo / +// stream.agentIo without holding ioStreamMutex, while WaitForAgent / +// RevokeStreamsForPurpose / RevokeStreamsForServer read & close the same +// fields under the lock. The go race detector catches it deterministically +// under -race; without the fix this test fails. +func TestIOStream_AgentConnectedIsRaceFreeUnderLock(t *testing.T) { + h := NewNezhaHandler() + const streamId = "race-test" + h.CreateStream(streamId, 1, 1) + t.Cleanup(func() { _ = h.CloseStream(streamId) }) + + var wg sync.WaitGroup + wg.Add(3) + + go func() { + defer wg.Done() + // repeatedly attach an agent + for i := 0; i < 200; i++ { + _ = h.AgentConnected(streamId, nopRWC{}) + } + }() + + go func() { + defer wg.Done() + // concurrently attach a user + for i := 0; i < 200; i++ { + _ = h.UserConnected(streamId, nopRWC{}) + } + }() + + go func() { + defer wg.Done() + // Revoker takes the write lock and reads the same userIo/agentIo + // fields the unsynchronised writers above are setting. Use a real + // targetServerID (1) so RevokeStreamsForServer actually inspects + // the entry's userIo/agentIo before deleting. + for i := 0; i < 200; i++ { + h.RevokeStreamsForServer(1) + h.CreateStream(streamId, 1, 1) + } + }() + + wg.Wait() +} + +// StartStream reads stream.userIo/agentIo while it waits for both endpoints. +// Those reads must be lock-protected against the concurrent writes done by +// UserConnected/AgentConnected; otherwise -race flags the data race on the +// interface fields. +func TestIOStream_StartStreamReadsAreRaceFree(t *testing.T) { + h := NewNezhaHandler() + const streamId = "startstream-race" + h.CreateStream(streamId, 2, 2) + t.Cleanup(func() { _ = h.CloseStream(streamId) }) + + var wg sync.WaitGroup + wg.Add(3) + + go func() { + defer wg.Done() + _ = h.StartStream(streamId, 50*time.Millisecond) + }() + go func() { + defer wg.Done() + time.Sleep(5 * time.Millisecond) + _ = h.AgentConnected(streamId, nopRWC{}) + }() + go func() { + defer wg.Done() + time.Sleep(5 * time.Millisecond) + _ = h.UserConnected(streamId, nopRWC{}) + }() + + wg.Wait() +} diff --git a/service/rpc/mcp_cancel_double_close_test.go b/service/rpc/mcp_cancel_double_close_test.go new file mode 100644 index 00000000..05eca6e2 --- /dev/null +++ b/service/rpc/mcp_cancel_double_close_test.go @@ -0,0 +1,55 @@ +package rpc + +import ( + "sync" + "sync/atomic" + "testing" + + pb "github.com/nezhahq/nezha/proto" +) + +// updateConfig has no serialization, so two concurrent admin PATCH /setting +// requests can both flip EnableMCP true->false and both invoke +// CancelAllMCPInflight concurrently. The sweep must close each entry's cancel +// channel at most once; a non-atomic check-then-close double-closes the same +// channel and panics, crashing the dashboard. +func TestCancelAllMCPInflight_ConcurrentSweepsDoNotDoubleClose(t *testing.T) { + mcpInflight.Range(func(key, _ any) bool { + mcpInflight.Delete(key) + return true + }) + t.Cleanup(func() { + mcpInflight.Range(func(key, _ any) bool { + mcpInflight.Delete(key) + return true + }) + }) + + const entries = 256 + for i := 0; i < entries; i++ { + mcpInflight.Store(uint64(i+1), &mcpInflightEntry{ + serverID: uint64(i + 1), + result: make(chan *pb.TaskResult, 1), + cancel: make(chan struct{}), + cancelled: new(atomic.Bool), + }) + } + + const sweepers = 8 + var wg sync.WaitGroup + wg.Add(sweepers) + for i := 0; i < sweepers; i++ { + go func() { + defer wg.Done() + // A double-close inside CancelAllMCPInflight panics here and + // fails the test (panic in a goroutine aborts the test binary). + CancelAllMCPInflight() + }() + } + wg.Wait() + + mcpInflight.Range(func(key, _ any) bool { + t.Fatalf("inflight entry %v survived the sweep", key) + return false + }) +} diff --git a/service/rpc/mcp_kill_switch_race_test.go b/service/rpc/mcp_kill_switch_race_test.go new file mode 100644 index 00000000..a7c3814f --- /dev/null +++ b/service/rpc/mcp_kill_switch_race_test.go @@ -0,0 +1,35 @@ +package rpc + +import ( + "context" + "testing" + "time" + + "github.com/nezhahq/nezha/model" +) + +// H9 regression: CallAgent must consult mcpKillSwitchObserved before any +// side-effects. Without this gate, mcpEndpoint's EnableMCP read and +// CancelAllMCPInflight race against a fresh CallAgent that registers +// AFTER the cancel sweep, surviving the disabled state. +func TestCallAgent_RefusesWhenKillSwitchObserved(t *testing.T) { + prevCheck := mcpKillSwitchObserver() + SetMCPKillSwitchObserver(func() bool { return true }) + t.Cleanup(func() { SetMCPKillSwitchObserver(prevCheck) }) + + _, err := CallAgent(context.Background(), 1, model.TaskTypeExec, struct{}{}, 50*time.Millisecond) + if err != ErrMCPDisabled { + t.Fatalf("CallAgent must short-circuit to ErrMCPDisabled when kill switch is observed, got %v", err) + } +} + +// The default hook must keep production behaviour disarmed so tests and +// unconfigured deployments do not short-circuit CallAgent. +func TestCallAgent_KillSwitchHookDefaultsToDisarmed(t *testing.T) { + if mcpKillSwitchObserver() == nil { + t.Fatal("mcpKillSwitchObserver must always return a non-nil probe so dashboard can wire it") + } + if mcpKillSwitchObserver()() { + t.Fatal("default hook must return false so unconfigured dashboards / tests don't accidentally short-circuit CallAgent") + } +} diff --git a/service/rpc/mcp_kill_switch_registration_race_test.go b/service/rpc/mcp_kill_switch_registration_race_test.go new file mode 100644 index 00000000..0b57c144 --- /dev/null +++ b/service/rpc/mcp_kill_switch_registration_race_test.go @@ -0,0 +1,80 @@ +package rpc + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + "github.com/nezhahq/nezha/model" +) + +// Registration-after-sweep race (review issue #1): CancelAllMCPInflight only +// cancels entries already present in the inflight map. A CallAgent that passes +// the upfront kill-switch check but has not yet Store()d its entry is invisible +// to the sweep, so without a post-registration re-check it goes on to SendTask +// a fresh exec/fs task to the agent AFTER EnableMCP=false. +// +// This test drives the worst-case interleaving deterministically: +// 1. CallAgent passes the upfront observer check (observer still false). +// 2. The operator flips the observer to "disabled" and runs the cancel sweep +// while CallAgent is paused between the check and Store. +// 3. CallAgent resumes; it MUST observe the kill switch on the post-Store +// re-check and return ErrMCPDisabled WITHOUT sending the task. +// +// With the race present, CallAgent sends the task and blocks until timeout +// (ErrAgentTimeout) — the agent received a fresh task past the kill switch. +func TestCallAgent_KillSwitchBeatsRegistrationAfterSweep(t *testing.T) { + const target uint64 = 7401 + + stream := newFakeStream() + cleanup := installFakeServer(t, target, stream) + defer cleanup() + + var killed bool + var mu sync.Mutex + prev := mcpKillSwitchObserver() + // Observer returns the operator-controlled flag. CallAgent reads it both + // before and (with the fix) after registering the inflight entry. + SetMCPKillSwitchObserver(func() bool { + mu.Lock() + defer mu.Unlock() + return killed + }) + t.Cleanup(func() { SetMCPKillSwitchObserver(prev) }) + + // Fail loudly if the agent ever receives a task: that means a fresh call + // leaked past the kill switch. + leaked := make(chan *struct{}, 1) + go func() { + select { + case <-stream.sent: + leaked <- nil + case <-time.After(2 * time.Second): + } + }() + + // 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 + // entry is missed by the sweep. + testKillSwitchAfterUpfrontCheck = func() { + mu.Lock() + killed = true + mu.Unlock() + CancelAllMCPInflight() + } + t.Cleanup(func() { testKillSwitchAfterUpfrontCheck = nil }) + + _, err := CallAgent(context.Background(), target, model.TaskTypeExec, + model.ExecRequest{Cmd: "x"}, 1*time.Second) + + if !errors.Is(err, ErrMCPDisabled) { + t.Fatalf("CallAgent must return ErrMCPDisabled when the kill switch fires during registration; got %v", err) + } + select { + case <-leaked: + t.Fatal("a fresh MCP task leaked to the agent past the kill switch") + default: + } +} diff --git a/service/rpc/mcp_rpc.go b/service/rpc/mcp_rpc.go new file mode 100644 index 00000000..435f5734 --- /dev/null +++ b/service/rpc/mcp_rpc.go @@ -0,0 +1,351 @@ +package rpc + +import ( + "context" + "encoding/json" + "errors" + "log" + "sync" + "sync/atomic" + "time" + + "github.com/nezhahq/nezha/model" + pb "github.com/nezhahq/nezha/proto" + "github.com/nezhahq/nezha/service/singleton" +) + +// MCP 的"调用-响应"模式复用了 RequestTask 双向流: +// - dashboard 发 Task(带新分配的 taskID + JSON params) +// - agent 执行后回 TaskResult(同 taskID + JSON result) +// - RequestTask 接收循环把这种 TaskType 识别后路由到 inflight 等待方 +// +// 不污染 model.Server 字段:用本包内的全局 inflight 表按 taskID 关联, +// 跨 server 共享单一命名空间。 + +var ( + mcpTaskIDCounter atomic.Uint64 + mcpInflight sync.Map // key: uint64 (taskID), value: chan *pb.TaskResult +) + +// ErrMCPDisabled 是 CallAgent 在 MCP kill switch 被触发时返回的哨兵错误。 +// 与 ErrAgentTimeout / ErrAgentOffline 平级,便于 controller 把它映射到 +// MCPOutcomeForbidden 之类的审计 code 而不是误报 agent 故障。 +var ErrMCPDisabled = errors.New("MCP is disabled by the dashboard administrator") + +// mcpKillSwitchObserved is a process-level hook the dashboard wires to +// singleton.Conf.EnableMCP. CallAgent consults it before any side-effects so +// the entry-check / cancel-sweep / registration race cannot leak a fresh +// call past EnableMCP=false. Defaults to "disarmed" so tests and headless +// builds are unaffected. +// +// Stored behind atomic.Pointer because SetMCPKillSwitchObserver (startup + +// tests) and CallAgent (any RPC goroutine) touch it concurrently; a plain +// func variable is a data race under -race. +var mcpKillSwitchObserved atomic.Pointer[func() bool] + +// disarmedKillSwitch is the default probe: never trips the kill switch. +var disarmedKillSwitch = func() bool { return false } + +// testKillSwitchAfterUpfrontCheck, when non-nil, runs inside CallAgent between +// the upfront kill-switch check and the inflight registration. Production +// leaves it nil; tests use it to drive the registration-after-sweep race +// deterministically. Guarded by the same atomic.Pointer for race-freedom. +var testKillSwitchAfterUpfrontCheck func() + +// SetMCPKillSwitchObserver installs the kill-switch probe the dashboard +// owns. Idempotent; the dashboard wires it at startup. Passing nil +// restores the default disarmed hook (used by tests to undo overrides). +func SetMCPKillSwitchObserver(fn func() bool) { + if fn == nil { + mcpKillSwitchObserved.Store(&disarmedKillSwitch) + return + } + mcpKillSwitchObserved.Store(&fn) +} + +// mcpKillSwitchObserver returns the currently installed probe, never nil. +func mcpKillSwitchObserver() func() bool { + if p := mcpKillSwitchObserved.Load(); p != nil { + return *p + } + return disarmedKillSwitch +} + +// allocateMCPTaskID 分配下一个 MCP 用的 task ID。 +// 取 1<<32 起步以与可能存在的 cron/transfer 等已有 ID 空间错开(cron.id 由 +// DB 自增,常量级,不会触及 1<<32)。 +func allocateMCPTaskID() uint64 { + const base uint64 = 1 << 32 + v := mcpTaskIDCounter.Add(1) + return base + v +} + +// CallAgent 给 serverID 对应的 agent 发一条 MCP-RPC 风格的 Task,并阻塞等待 TaskResult 回包。 +// +// taskType 必须是 model.IsMCPRPCResult 返回 true 的类型;params 会被 JSON 编码进 Task.Data。 +// 超时由调用方控制;触发超时后从 inflight 表移除等待 slot(晚到的回包会被丢弃)。 +// +// 错误语义: +// - server 未在线 / 未连接 task stream → ErrAgentOffline +// - 超时 → ctx.Err 或 ErrAgentTimeout +// - agent 回包 successful=false → 把 result.Data 当错误字符串返回 +// - CancelAllMCPInflight 期间被中断 → ErrMCPDisabled +// - 任何 send 失败、序列化失败 → 原始 error +// +// 返回的 raw JSON 是 agent 端 TaskResult.Data 的原文。 +func CallAgent(ctx context.Context, serverID uint64, taskType uint64, params any, timeout time.Duration) (json.RawMessage, error) { + if !model.IsMCPRPCResult(taskType) { + return nil, errors.New("CallAgent: task type is not registered as MCP RPC") + } + + killSwitch := mcpKillSwitchObserver() + if killSwitch() { + return nil, ErrMCPDisabled + } + + server, _ := singleton.ServerShared.Get(serverID) + if server == nil { + return nil, ErrAgentOffline + } + if server.GetTaskStream() == nil { + return nil, ErrAgentOffline + } + + body, err := json.Marshal(params) + if err != nil { + return nil, err + } + + taskID := allocateMCPTaskID() + resultCh := make(chan *pb.TaskResult, 1) + cancelCh := make(chan struct{}) + entry := &mcpInflightEntry{ + serverID: serverID, + result: resultCh, + cancel: cancelCh, + cancelled: new(atomic.Bool), + } + + if hook := testKillSwitchAfterUpfrontCheck; hook != nil { + hook() + } + + mcpInflight.Store(taskID, entry) + defer mcpInflight.Delete(taskID) + + // Close the registration-after-sweep window: a kill switch that fired + // between the upfront check and this Store is invisible to + // CancelAllMCPInflight (our entry was not in the map yet). Because the + // operator sets EnableMCP=false BEFORE running the sweep, re-reading the + // observer here after Store guarantees we either see it disabled, or the + // sweep saw our now-registered entry and flipped entry.cancelled. + if killSwitch() || entry.cancelled.Load() { + return nil, ErrMCPDisabled + } + + if err := server.SendTask(&pb.Task{ + Id: taskID, + Type: taskType, + Data: string(body), + }); err != nil { + if errors.Is(err, model.ErrTaskStreamOffline) { + return nil, ErrAgentOffline + } + return nil, err + } + + waitCtx := ctx + var cancel context.CancelFunc + if timeout > 0 { + waitCtx, cancel = context.WithTimeout(ctx, timeout) + defer cancel() + } + + select { + case res := <-resultCh: + // Cancel must beat a late agent reply: Go select picks a random + // ready case, so if CancelAllMCPInflight closed cancelCh after the + // agent already filled resultCh we could still surface success. + // Re-check the cancel flag and prefer ErrMCPDisabled, matching the + // contract documented above ("CancelAllMCPInflight 期间被中断 → + // ErrMCPDisabled") and what TestUpdateConfig_DisablingMCPInvokesKillSwitch + // expects. + if entry.cancelled.Load() { + return nil, ErrMCPDisabled + } + if res == nil { + return nil, errors.New("agent returned nil result") + } + if !res.GetSuccessful() { + if res.GetData() != "" { + return nil, errors.New(res.GetData()) + } + return nil, errors.New("agent returned unsuccessful result") + } + return json.RawMessage(res.GetData()), nil + case <-cancelCh: + return nil, ErrMCPDisabled + case <-waitCtx.Done(): + if errors.Is(waitCtx.Err(), context.DeadlineExceeded) { + return nil, ErrAgentTimeout + } + return nil, waitCtx.Err() + } +} + +// mcpInflightEntry binds an in-flight MCP call to its target serverID and +// pairs the result channel with a per-call cancel channel so the kill switch +// can break out of CallAgent without leaving the result channel dangling for +// the next late agent reply. The serverID is the authoritative reporter +// identity check at delivery time — without it deliverMCPResult would route +// purely by attacker-controlled TaskResult.Id (same bug class as commit +// 02129f1 in the cron path). +// +// cancelled flips to true the instant CancelAllMCPInflight observes the +// entry. Every code path that could complete the call — the CallAgent +// select on resultCh, deliverMCPResult, deliverMCPResultFromReporter — +// MUST consult it before treating an agent reply as authoritative, otherwise +// a TaskResult delivered concurrently with the kill switch can win Go's +// random select tiebreak and surface success after EnableMCP=false. +type mcpInflightEntry struct { + serverID uint64 + result chan *pb.TaskResult + cancel chan struct{} + cancelled *atomic.Bool + closeOnce sync.Once +} + +// closeCancel closes the entry's cancel channel exactly once. Concurrent +// CancelAllMCPInflight sweeps (two admin PATCH /setting requests both +// disabling MCP) would otherwise race a non-atomic check-then-close and +// panic on the second close. +func (e *mcpInflightEntry) closeCancel() { + e.closeOnce.Do(func() { close(e.cancel) }) +} + +// CancelAllMCPInflight closes every in-flight CallAgent so they return +// ErrMCPDisabled immediately. Used by the EnableMCP=false transition: by +// itself the inflight table holds the dashboard goroutine hostage until +// the agent replies (or the per-call timeout fires, up to ~305s for +// server.exec). Returns the number of calls cancelled for audit. +// +// Implementation notes: +// - Set the cancelled flag BEFORE closing cancelCh so any goroutine that +// already woke on resultCh observes it on the post-select re-check. +// - Delete the entry from mcpInflight immediately. Late agent replies via +// deliverMCPResult* would otherwise still find it (their own cancelled +// check covers concurrent delete, but evicting eagerly keeps the table +// small under repeated kill switch / re-enable cycles). +func CancelAllMCPInflight() int { + cancelled := 0 + mcpInflight.Range(func(key, value any) bool { + entry, ok := value.(*mcpInflightEntry) + if !ok { + return true + } + if entry.cancelled != nil { + entry.cancelled.Store(true) + } + entry.closeCancel() + mcpInflight.Delete(key) + cancelled++ + return true + }) + return cancelled +} + +// DeliverMCPResultForTest 暴露 deliverMCPResult 给跨包测试用:这是显式的 +// "信任路径 / 不做 reporter 校验"入口,专给不关心来源的旧测试用。 +// 安全敏感测试请用 DeliverMCPResultFromReporterForTest 并传入真实 reporterID。 +func DeliverMCPResultForTest(res *pb.TaskResult) { deliverMCPResult(res) } + +// DeliverMCPResultFromReporterForTest 暴露带 reporter 校验的投递入口给跨包 +// 测试用,与生产 RequestTask 路径同语义:reporterID 必须等于 inflight 条目 +// 登记的目标 serverID 才会投递。reporterID == 0 视为 "未知 reporter" 并被 +// 拒绝;要绕过 reporter 校验请改用 DeliverMCPResultForTest。 +func DeliverMCPResultFromReporterForTest(res *pb.TaskResult, reporterID uint64) { + deliverMCPResultFromReporter(res, reporterID) +} + +// inflightServerIDForTest 返回某个 taskID 当前挂载的目标 serverID。用于安全 +// 回归测试断言 inflight 条目确实把目标 server 绑进了路由表。 +// 未找到时返回 (0, false)。 +func inflightServerIDForTest(taskID uint64) (uint64, bool) { + v, ok := mcpInflight.Load(taskID) + if !ok { + return 0, false + } + entry, ok := v.(*mcpInflightEntry) + if !ok { + return 0, false + } + return entry.serverID, true +} + +// deliverMCPResult 把 RequestTask 收到的 MCP-RPC TaskResult 路由到等待方。 +// 找不到等待 slot(已超时被移除)则丢弃。 +// +// 此变体不做 reporter 校验,仅用于不关心 reporter 的内部/测试路径。生产 +// RequestTask 接收循环必须走 deliverMCPResultFromReporter,把 stream 上 +// 已认证的 clientID 作为 reporter 传入。 +func deliverMCPResult(res *pb.TaskResult) { + if res == nil { + return + } + v, ok := mcpInflight.Load(res.GetId()) + if !ok { + return + } + entry, ok := v.(*mcpInflightEntry) + if !ok { + return + } + if entry.cancelled != nil && entry.cancelled.Load() { + return + } + select { + case entry.result <- res: + default: + } +} + +// deliverMCPResultFromReporter 是生产路径的入口:要求 reporterID 与 inflight +// 条目登记的目标 serverID 一致才投递;否则丢弃并打日志。reporterID == 0 +// 视为“未知 reporter”,安全起见也丢弃。 +// +// 这条校验是必要的:mcpInflight 用全局递增 taskID 做键,跨 server 共享 +// 单一命名空间;如果不在投递时核对上报 agent 是 CallAgent 的目标 server, +// 任何已认证的恶意/失陷 agent 都能用猜到的 taskID 抢答其他 server 的 +// MCP 调用(resultCh 容量 1,先到者覆盖真正回包)——和 commit 02129f1 +// 在 cron 路径修过的攻击面同类。 +func deliverMCPResultFromReporter(res *pb.TaskResult, reporterID uint64) { + if res == nil { + return + } + v, ok := mcpInflight.Load(res.GetId()) + if !ok { + return + } + entry, ok := v.(*mcpInflightEntry) + if !ok { + return + } + if reporterID == 0 || entry.serverID != reporterID { + log.Printf("NEZHA>> MCP result ignored: taskID=%d targetServerID=%d reporterID=%d", + res.GetId(), entry.serverID, reporterID) + return + } + if entry.cancelled != nil && entry.cancelled.Load() { + return + } + select { + case entry.result <- res: + default: + } +} + +// 错误类型 +var ( + ErrAgentOffline = errors.New("agent offline or task stream not connected") + ErrAgentTimeout = errors.New("agent did not respond within timeout") +) diff --git a/service/rpc/mcp_rpc_helper_doc_test.go b/service/rpc/mcp_rpc_helper_doc_test.go new file mode 100644 index 00000000..69be8e0d --- /dev/null +++ b/service/rpc/mcp_rpc_helper_doc_test.go @@ -0,0 +1,61 @@ +package rpc + +import ( + "strings" + "sync/atomic" + "testing" + + pb "github.com/nezhahq/nezha/proto" +) + +// 把"测试 helper 的注释与运行时语义"钉成测试,避免 helper 文档骗读者: +// +// 1. DeliverMCPResultForTest 是显式的"信任路径 / 不做 reporter 校验"入口。 +// 2. DeliverMCPResultFromReporterForTest 是带 reporter 校验的入口; +// reporterID == 0 视为"未知 reporter",必须被拒绝,不能像旧注释暗示的 +// 那样当作"未知/不校验"放行。 +// +// 这条契约决定了任何安全敏感的跨包测试调用方式:要绕过 reporter, +// 必须用 DeliverMCPResultForTest,而不是 reporterID=0 通过 reporter 入口。 +func TestDeliverMCPResultFromReporterForTest_ZeroReporterIDIsRejected(t *testing.T) { + taskID := allocateMCPTaskID() + resultCh := make(chan *pb.TaskResult, 1) + cancelCh := make(chan struct{}) + mcpInflight.Store(taskID, &mcpInflightEntry{ + serverID: 7, + result: resultCh, + cancel: cancelCh, + cancelled: new(atomic.Bool), + }) + t.Cleanup(func() { mcpInflight.Delete(taskID) }) + + DeliverMCPResultFromReporterForTest(&pb.TaskResult{Id: taskID, Data: "x", Successful: true}, 0) + + select { + case <-resultCh: + t.Fatalf("reporterID==0 must be rejected by the reporter-checked helper; expected no delivery") + default: + } +} + +// 同时把"测试 helper 自身的文档约束"钉到代码里:注释必须明确说出 +// "reporterID == 0 视为未知 reporter 并被拒绝",否则未来维护者很容易看着 +// "不校验"的旧措辞写出绕过 reporter 的安全敏感测试。 +func TestDeliverMCPResultFromReporterForTest_DocStatesZeroIsRejected(t *testing.T) { + src := mustReadFile(t, "mcp_rpc.go") + if !strings.Contains(src, "DeliverMCPResultFromReporterForTest") { + t.Fatalf("expected helper to live in mcp_rpc.go") + } + // 提取 helper 上方的注释块:从 helper 名字往上找到第一段连续的 // 行。 + idx := strings.Index(src, "func DeliverMCPResultFromReporterForTest(") + if idx < 0 { + t.Fatalf("helper not found in source") + } + prefix := src[:idx] + if !strings.Contains(prefix, "reporterID == 0") { + t.Fatalf("doc must mention reporterID == 0 contract explicitly") + } + if strings.Contains(prefix, "不校验") { + t.Fatalf("doc still claims reporterID==0 is 不校验; this contradicts deliverMCPResultFromReporter which drops it") + } +} diff --git a/service/rpc/mcp_rpc_kill_switch_race_test.go b/service/rpc/mcp_rpc_kill_switch_race_test.go new file mode 100644 index 00000000..26cdf586 --- /dev/null +++ b/service/rpc/mcp_rpc_kill_switch_race_test.go @@ -0,0 +1,95 @@ +package rpc + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/nezhahq/nezha/model" + pb "github.com/nezhahq/nezha/proto" +) + +// Kill switch must beat a late agent reply. Without the cancelled-flag +// re-check in CallAgent, the following sequence surfaces success after +// EnableMCP=false: +// +// t0 agent puts TaskResult into resultCh (capacity 1, non-blocking) +// t1 admin flips EnableMCP=false → CancelAllMCPInflight closes cancelCh +// t2 CallAgent's select sees BOTH cases ready; Go picks one at random; +// if it picks resultCh, the call returns the agent's payload even +// though the operator's kill switch fired. +// +// The fix is to mark the entry cancelled BEFORE closing cancelCh and have +// the resultCh branch re-check that flag. This test pins the contract by +// driving the worst-case ordering: result is delivered FIRST, then the +// kill switch fires, then CallAgent observes both. With the race in place +// this would flake (random select); with the fix it always returns +// ErrMCPDisabled. +func TestCallAgent_KillSwitchBeatsConcurrentLateResult(t *testing.T) { + const target uint64 = 7301 + + stream := newFakeStream() + cleanup := installFakeServer(t, target, stream) + defer cleanup() + + delivered := make(chan struct{}) + go func() { + sent := <-stream.sent + deliverMCPResultFromReporter(&pb.TaskResult{ + Id: sent.GetId(), + Type: model.TaskTypeExec, + Successful: true, + Data: `{"exit_code":0,"stdout":"should-not-surface"}`, + }, target) + CancelAllMCPInflight() + close(delivered) + }() + + _, err := CallAgent(context.Background(), target, model.TaskTypeExec, + model.ExecRequest{Cmd: "x"}, 2*time.Second) + <-delivered + if !errors.Is(err, ErrMCPDisabled) { + t.Fatalf("kill switch must win the race with a late agent reply; want ErrMCPDisabled, got %v", err) + } +} + +// CancelAllMCPInflight must eagerly evict entries so a stale TaskResult +// that arrives after the kill switch cannot still land in resultCh. +// Without the cancelled flag this entry would still be reachable through +// deliverMCPResultFromReporter; the flag guarantees the late delivery is +// silently dropped even if the caller has not returned yet. +func TestCancelAllMCPInflight_LaterResultIsSwallowed(t *testing.T) { + const target uint64 = 7302 + + stream := newFakeStream() + cleanup := installFakeServer(t, target, stream) + defer cleanup() + + taskIDCh := make(chan uint64, 1) + go func() { + sent := <-stream.sent + taskIDCh <- sent.GetId() + }() + + resultCh := make(chan error, 1) + go func() { + _, err := CallAgent(context.Background(), target, model.TaskTypeExec, + model.ExecRequest{Cmd: "x"}, 5*time.Second) + resultCh <- err + }() + + taskID := <-taskIDCh + CancelAllMCPInflight() + + if err := <-resultCh; !errors.Is(err, ErrMCPDisabled) { + t.Fatalf("CallAgent must return ErrMCPDisabled after kill switch; got %v", err) + } + + deliverMCPResultFromReporter(&pb.TaskResult{ + Id: taskID, + Type: model.TaskTypeExec, + Successful: true, + Data: `{"exit_code":0}`, + }, target) +} diff --git a/service/rpc/mcp_rpc_spoof_test.go b/service/rpc/mcp_rpc_spoof_test.go new file mode 100644 index 00000000..d9bba569 --- /dev/null +++ b/service/rpc/mcp_rpc_spoof_test.go @@ -0,0 +1,149 @@ +package rpc + +import ( + "context" + "encoding/json" + "errors" + "testing" + "time" + + "github.com/nezhahq/nezha/model" + pb "github.com/nezhahq/nezha/proto" +) + +// These tests pin the security invariant that an MCP TaskResult delivered +// back through RequestTask must come from the SAME agent the CallAgent was +// targeted at. The receive loop in service/rpc/nezha.go has the authenticated +// clientID in scope; deliverMCPResult must consume it and reject mismatches. +// +// Why the invariant matters: mcpInflight is keyed by a globally increasing +// counter (allocateMCPTaskID) and the lookup table is shared across servers. +// Without binding the inflight entry to the target serverID and verifying it +// against the reporter clientID, any compromised agent A can race a forged +// TaskResult for server B's CallAgent (resultCh capacity is 1; first reply +// wins, real reply is dropped). The same class of attack motivated the cron +// path's CanReportCronResult and the transfer path's pending.ID == result.Id +// check in this very file's RequestTask switch. + +// TestDeliverMCPResult_RejectsForeignReporter is the security regression: a +// reporter that is NOT the call target must not be able to deliver into +// another server's inflight slot, even with a correctly-guessed taskID. +func TestDeliverMCPResult_RejectsForeignReporter(t *testing.T) { + const ( + targetServerID uint64 = 6101 + foreignAgentID uint64 = 6102 + ) + + stream := newFakeStream() + cleanup := installFakeServer(t, targetServerID, stream) + defer cleanup() + + captured := make(chan uint64, 1) + go func() { + sent := <-stream.sent + // Foreign agent racing a forged TaskResult with the right taskID. + DeliverMCPResultFromReporterForTest(&pb.TaskResult{ + Id: sent.GetId(), + Type: model.TaskTypeExec, + Successful: true, + Data: `{"exit_code":0,"stdout":"forged"}`, + }, foreignAgentID) + captured <- sent.GetId() + }() + + _, err := CallAgent(context.Background(), targetServerID, model.TaskTypeExec, + model.ExecRequest{Cmd: "x"}, 200*time.Millisecond) + if !errors.Is(err, ErrAgentTimeout) { + t.Fatalf("forged result from foreign reporter must NOT deliver; want ErrAgentTimeout, got %v", err) + } + select { + case <-captured: + case <-time.After(time.Second): + t.Fatalf("test stream never observed the dispatched task") + } +} + +// TestDeliverMCPResult_AcceptsMatchingReporter is the green companion: when +// the reporter clientID matches the inflight target, the result must still +// route correctly (we are not breaking the happy path). +func TestDeliverMCPResult_AcceptsMatchingReporter(t *testing.T) { + const targetServerID uint64 = 6103 + + stream := newFakeStream() + cleanup := installFakeServer(t, targetServerID, stream) + defer cleanup() + + want := model.ExecResult{ExitCode: 0, Stdout: "ok"} + payload, _ := json.Marshal(want) + + go func() { + sent := <-stream.sent + DeliverMCPResultFromReporterForTest(&pb.TaskResult{ + Id: sent.GetId(), + Type: model.TaskTypeExec, + Successful: true, + Data: string(payload), + }, targetServerID) + }() + + raw, err := CallAgent(context.Background(), targetServerID, model.TaskTypeExec, + model.ExecRequest{Cmd: "x"}, 2*time.Second) + if err != nil { + t.Fatalf("matching reporter must deliver, got %v", err) + } + var got model.ExecResult + if err := json.Unmarshal(raw, &got); err != nil { + t.Fatalf("bad result json: %v", err) + } + if got.Stdout != "ok" { + t.Fatalf("payload not propagated, got %+v", got) + } +} + +// TestDeliverMCPResult_InflightEntryBoundToServerID locks in the structural +// requirement that the inflight table records the target serverID. Without +// this binding deliverMCPResult cannot perform the reporter check above. +// Probing via reflection avoids exporting mcpInflight just for tests. +func TestDeliverMCPResult_InflightEntryBoundToServerID(t *testing.T) { + const targetServerID uint64 = 6104 + + stream := newFakeStream() + cleanup := installFakeServer(t, targetServerID, stream) + defer cleanup() + + gotEntry := make(chan struct { + taskID uint64 + serverID uint64 + found bool + }, 1) + go func() { + sent := <-stream.sent + taskID := sent.GetId() + serverID, ok := inflightServerIDForTest(taskID) + gotEntry <- struct { + taskID uint64 + serverID uint64 + found bool + }{taskID, serverID, ok} + // Unblock CallAgent so the inflight slot is cleaned up. + DeliverMCPResultFromReporterForTest(&pb.TaskResult{ + Id: taskID, + Type: model.TaskTypeExec, + Successful: true, + Data: "{}", + }, targetServerID) + }() + + _, err := CallAgent(context.Background(), targetServerID, model.TaskTypeExec, + model.ExecRequest{Cmd: "x"}, 2*time.Second) + if err != nil { + t.Fatalf("unexpected CallAgent error: %v", err) + } + probe := <-gotEntry + if !probe.found { + t.Fatalf("inflight entry for taskID=%d not found while CallAgent was blocking", probe.taskID) + } + if probe.serverID != targetServerID { + t.Fatalf("inflight entry must carry target serverID=%d, got %d", targetServerID, probe.serverID) + } +} diff --git a/service/rpc/mcp_rpc_test.go b/service/rpc/mcp_rpc_test.go new file mode 100644 index 00000000..eb423142 --- /dev/null +++ b/service/rpc/mcp_rpc_test.go @@ -0,0 +1,165 @@ +package rpc + +import ( + "context" + "encoding/json" + "errors" + "sync/atomic" + "testing" + "time" + + "google.golang.org/grpc/metadata" + + "github.com/nezhahq/nezha/model" + pb "github.com/nezhahq/nezha/proto" + "github.com/nezhahq/nezha/service/singleton" +) + +type fakeTaskStream struct { + sent chan *pb.Task + delay time.Duration +} + +func newFakeStream() *fakeTaskStream { + return &fakeTaskStream{sent: make(chan *pb.Task, 4)} +} + +func (s *fakeTaskStream) Send(t *pb.Task) error { s.sent <- t; return nil } +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() { + t.Helper() + original := singleton.ServerShared + sc := singleton.NewEmptyServerClassForTest() + srv := &model.Server{} + srv.ID = id + srv.SetTaskStream(stream) + sc.InsertForTest(srv) + singleton.ServerShared = sc + return func() { singleton.ServerShared = original } +} + +func TestCallAgent_RejectsNonMCPType(t *testing.T) { + _, err := CallAgent(context.Background(), 1, model.TaskTypeCommand, struct{}{}, time.Second) + if err == nil { + t.Fatalf("expected error for non-MCP type") + } +} + +func TestCallAgent_OfflineWhenNoStream(t *testing.T) { + original := singleton.ServerShared + sc := singleton.NewEmptyServerClassForTest() + srv := &model.Server{} + srv.ID = 7 + sc.InsertForTest(srv) + singleton.ServerShared = sc + t.Cleanup(func() { singleton.ServerShared = original }) + + _, err := CallAgent(context.Background(), 7, model.TaskTypeExec, struct{}{}, time.Second) + if !errors.Is(err, ErrAgentOffline) { + t.Fatalf("expected ErrAgentOffline, got %v", err) + } +} + +func TestCallAgent_HappyPath_DelivlersResultByTaskID(t *testing.T) { + stream := newFakeStream() + cleanup := installFakeServer(t, 42, stream) + defer cleanup() + + resultPayload, _ := json.Marshal(model.ExecResult{ExitCode: 0, Stdout: "hello"}) + + var captured atomic.Uint64 + done := make(chan struct{}) + go func() { + sent := <-stream.sent + captured.Store(sent.GetId()) + deliverMCPResult(&pb.TaskResult{ + Id: sent.GetId(), + Type: model.TaskTypeExec, + Data: string(resultPayload), + Successful: true, + }) + close(done) + }() + + raw, err := CallAgent(context.Background(), 42, model.TaskTypeExec, model.ExecRequest{Cmd: "x"}, 5*time.Second) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + <-done + if captured.Load() == 0 { + t.Fatalf("task id never captured") + } + var got model.ExecResult + if err := json.Unmarshal(raw, &got); err != nil { + t.Fatalf("bad result json: %v", err) + } + if got.Stdout != "hello" { + t.Fatalf("payload not propagated, got %+v", got) + } +} + +func TestCallAgent_Timeout(t *testing.T) { + stream := newFakeStream() + cleanup := installFakeServer(t, 43, stream) + defer cleanup() + + go func() { <-stream.sent }() + + _, err := CallAgent(context.Background(), 43, model.TaskTypeFsRead, model.FsReadRequest{Path: "/x"}, 50*time.Millisecond) + if !errors.Is(err, ErrAgentTimeout) { + t.Fatalf("expected ErrAgentTimeout, got %v", err) + } +} + +func TestCallAgent_LateResultIsDropped(t *testing.T) { + stream := newFakeStream() + cleanup := installFakeServer(t, 44, stream) + defer cleanup() + + var taskID uint64 + got := make(chan struct{}) + go func() { + sent := <-stream.sent + taskID = sent.GetId() + close(got) + }() + + _, err := CallAgent(context.Background(), 44, model.TaskTypeFsDelete, model.FsDeleteRequest{Path: "/x"}, 50*time.Millisecond) + if !errors.Is(err, ErrAgentTimeout) { + t.Fatalf("expected timeout") + } + <-got + + deliverMCPResult(&pb.TaskResult{Id: taskID, Type: model.TaskTypeFsDelete, Successful: true, Data: "{}"}) + if _, ok := mcpInflight.Load(taskID); ok { + t.Fatalf("inflight entry must be cleaned up after timeout") + } +} + +func TestCallAgent_UnsuccessfulIsError(t *testing.T) { + stream := newFakeStream() + cleanup := installFakeServer(t, 45, stream) + defer cleanup() + + go func() { + sent := <-stream.sent + deliverMCPResult(&pb.TaskResult{ + Id: sent.GetId(), + Type: sent.GetType(), + Successful: false, + Data: "agent says nope", + }) + }() + + _, err := CallAgent(context.Background(), 45, model.TaskTypeFsWrite, model.FsWriteRequest{Path: "/x", Content: "y"}, time.Second) + if err == nil || err.Error() != "agent says nope" { + t.Fatalf("expected agent error message, got %v", err) + } +} diff --git a/service/rpc/nezha.go b/service/rpc/nezha.go index 97325418..c006ee49 100644 --- a/service/rpc/nezha.go +++ b/service/rpc/nezha.go @@ -37,6 +37,32 @@ func NewNezhaHandler() *NezhaHandler { } } +// attachRequestTaskStream resolves the server for clientID and publishes the +// task stream. It mirrors the !ok || server == nil guard the other RPC entry +// points use: the server can be deleted between CheckRequestTask and this +// lookup, in which case Get returns a nil *Server and SetTaskStream would +// panic. +func attachRequestTaskStream(clientID uint64, stream pb.NezhaService_RequestTaskServer) (*model.Server, bool) { + server, ok := singleton.ServerShared.Get(clientID) + if !ok || server == nil { + return nil, false + } + server.SetTaskStream(stream) + return server, true +} + +// clearRequestTaskStream detaches the dropped stream from whichever *Server is +// currently published for clientID. Edit and transfer rotation publish a new +// *Server that adopts the same stream holder, so cleanup must target the live +// map entry; the captured server is only the fallback for a removed entry. +func clearRequestTaskStream(clientID uint64, captured *model.Server, stream pb.NezhaService_RequestTaskServer) { + if current, ok := singleton.ServerShared.Get(clientID); ok && current != nil { + current.ClearTaskStreamIfCurrent(stream) + return + } + captured.ClearTaskStreamIfCurrent(stream) +} + func (s *NezhaHandler) RequestTask(stream pb.NezhaService_RequestTaskServer) error { var clientID uint64 var err error @@ -44,9 +70,11 @@ func (s *NezhaHandler) RequestTask(stream pb.NezhaService_RequestTaskServer) err return err } - server, _ := singleton.ServerShared.Get(clientID) - server.SetTaskStream(stream) - defer server.ClearTaskStreamIfCurrent(stream) + server, ok := attachRequestTaskStream(clientID, stream) + if !ok { + return nil + } + defer clearRequestTaskStream(clientID, server, stream) // If a transfer is mid-flight for this server, the agent has just brought // up a fresh bidi stream — this is the moment to (re)deliver the // ApplyConfig task carrying the new owner's AgentSecret. Pushes from @@ -114,6 +142,10 @@ func (s *NezhaHandler) RequestTask(stream pb.NezhaService_RequestTaskServer) err log.Printf("NEZHA>> ServerTransfer MarkFailed(%d) failed: %v", result.GetId(), err) } default: + if model.IsMCPRPCResult(result.GetType()) { + deliverMCPResultFromReporter(result, clientID) + continue + } if model.IsServiceSentinelNeeded(result.GetType()) { singleton.ServiceSentinelShared.Dispatch(singleton.ReportData{ Data: result, @@ -265,24 +297,45 @@ func (s *NezhaHandler) IOStream(stream pb.NezhaService_IOStreamServer) error { return fmt.Errorf("stream not authorized for agent") } - go func() { - for { - if err := stream.Send(&pb.IOStreamData{Data: []byte{}}); err != nil { - log.Printf("NEZHA>> IOStream keepAlive error: %v\n", err) - return - } - time.Sleep(time.Second * 30) - } - }() - if _, err := s.GetStream(streamId); err != nil { return err } iw := grpcx.NewIOStreamWrapper(stream) + + // Keepalive MUST go through the wrapper so it shares the same sendMu as + // MCP fs.transfer / terminal / fm Writers. Calling stream.Send directly + // here used to race those Writers — grpc-go forbids concurrent SendMsg + // on the same stream. The wrapper's sendMu is the dashboard-side dual + // of agent/cmd/agent/mcp_fs_transfer.go's serialIOStreamSender. + keepaliveDone := make(chan struct{}) + go func() { + defer close(keepaliveDone) + ticker := time.NewTicker(time.Second * 30) + defer ticker.Stop() + for { + select { + case <-iw.Context().Done(): + return + case <-iw.Done(): + // 业务侧(CloseStream / RevokeStreamsForPurpose)调过 + // iw.Close()。即便底层 gRPC stream context 尚未取消,也 + // 必须立刻收手——否则要再等一整个 30s tick,handler 在 + // iw.Wait() 之后又得多等一拍 keepaliveDone 才能返回。 + return + case <-ticker.C: + if err := iw.SendKeepalive(); err != nil { + log.Printf("NEZHA>> IOStream keepAlive error: %v\n", err) + return + } + } + } + }() + if err := s.AgentConnected(streamId, iw); err != nil { return err } iw.Wait() + <-keepaliveDone return nil } diff --git a/service/rpc/request_task_missing_server_test.go b/service/rpc/request_task_missing_server_test.go new file mode 100644 index 00000000..a44350ba --- /dev/null +++ b/service/rpc/request_task_missing_server_test.go @@ -0,0 +1,25 @@ +package rpc + +import ( + "testing" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/singleton" +) + +func TestAttachRequestTaskStream_MissingServerDoesNotPanic(t *testing.T) { + reporter := requestTaskSecurityServer(7, 200, "eeeeeeee-eeee-eeee-eeee-eeeeeeeeeeee") + setupRequestTaskSecurityFixture(t, []*model.Server{reporter}, nil, map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember}, + }, map[string]uint64{"reporter-secret": 200}) + + singleton.ServerShared.Delete([]uint64{reporter.ID}) + + srv, ok := attachRequestTaskStream(reporter.ID, nil) + if ok { + t.Fatal("attach must report not-ok when the server was deleted between auth and lookup") + } + if srv != nil { + t.Fatalf("attach must return a nil server for a deleted id, got %#v", srv) + } +} diff --git a/service/rpc/request_task_stale_stream_test.go b/service/rpc/request_task_stale_stream_test.go new file mode 100644 index 00000000..8924c3d9 --- /dev/null +++ b/service/rpc/request_task_stale_stream_test.go @@ -0,0 +1,46 @@ +package rpc + +import ( + "context" + "errors" + "testing" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/singleton" +) + +// When a server is edited mid-session, updateServer swaps a new *Server into +// ServerShared that adopts the live stream holder. The agent's RequestTask +// cleanup must detach the stream from whichever *Server is currently published, +// not the stale object captured when the stream attached — otherwise the new +// object keeps reporting the agent as online on a dead stream. +func TestRequestTaskCleanupDetachesStreamFromCurrentServerAfterEdit(t *testing.T) { + reporter := requestTaskSecurityServer(7, 200, "ffffffff-ffff-ffff-ffff-ffffffffffff") + setupRequestTaskSecurityFixture(t, []*model.Server{reporter}, nil, map[uint64]model.UserInfo{ + 200: {Role: model.RoleMember}, + }, map[string]uint64{"reporter-secret": 200}) + + old, ok := singleton.ServerShared.Get(reporter.ID) + if !ok { + t.Fatalf("server %d not found", reporter.ID) + } + + stream := requestTaskSecurityAuthedStream("reporter-secret", reporter.UUID) + stream.onRecv = func() { + edited := &model.Server{Common: model.Common{ID: old.ID, UserID: old.UserID}, UUID: old.UUID, Name: "edited"} + edited.CopyFromRunningServer(old) + singleton.ServerShared.Update(edited, "") + } + + if err := NewNezhaHandler().RequestTask(stream); !errors.Is(err, context.Canceled) { + t.Fatalf("expected RequestTask to finish after Recv error, got %v", err) + } + + current, ok := singleton.ServerShared.Get(reporter.ID) + if !ok { + t.Fatalf("server %d not found after edit", reporter.ID) + } + if got := current.GetTaskStream(); got != nil { + t.Fatalf("edited server must report offline after the agent stream dropped, got %T", got) + } +} diff --git a/service/rpc/testdata_helper_test.go b/service/rpc/testdata_helper_test.go new file mode 100644 index 00000000..ee078958 --- /dev/null +++ b/service/rpc/testdata_helper_test.go @@ -0,0 +1,20 @@ +package rpc + +import ( + "os" + "path/filepath" + "testing" +) + +func mustReadFile(t *testing.T, name string) string { + t.Helper() + wd, err := os.Getwd() + if err != nil { + t.Fatalf("getwd: %v", err) + } + b, err := os.ReadFile(filepath.Join(wd, name)) + if err != nil { + t.Fatalf("read %s: %v", name, err) + } + return string(b) +} diff --git a/service/rpc/wait_for_agent_revoke_test.go b/service/rpc/wait_for_agent_revoke_test.go new file mode 100644 index 00000000..e247b7d6 --- /dev/null +++ b/service/rpc/wait_for_agent_revoke_test.go @@ -0,0 +1,56 @@ +package rpc + +import ( + "context" + "testing" + "time" +) + +// WaitForAgent 必须在 stream 被 RevokeStreamsForPurpose 强制下线时立刻返回, +// 否则 EnableMCP=false 的 kill switch 不能真的“立即”切断那些还卡在 +// “等待 agent attach” 阶段的 transfer 请求 —— 它们会一直等到 timeout +// (生产路径上是 30 秒)。 +// +// 期望行为:调用 RevokeStreamsForPurpose 后,WaitForAgent 在远小于 timeout +// 的时间内返回 (nil, false)。 +func TestWaitForAgent_RevokeWakesUpWaiter(t *testing.T) { + h := NewNezhaHandler() + const streamID = "kill-switch-wait" + h.CreateStreamWithPurpose(streamID, 0, 7, PurposeMCPTransfer) + + done := make(chan struct { + io any + ok bool + dur time.Duration + }, 1) + + start := time.Now() + go func() { + // 给一个明显大于 revoke 触发延时的 timeout;如果 revoke 没唤醒, + // WaitForAgent 会一直等到这里,下面的 assertion 就会失败。 + stream, ok := h.WaitForAgent(context.Background(), streamID, 5*time.Second) + done <- struct { + io any + ok bool + dur time.Duration + }{stream, ok, time.Since(start)} + }() + + // 让 WaitForAgent 真的进入 select 等待,再触发 kill switch。 + time.Sleep(50 * time.Millisecond) + if revoked := h.RevokeStreamsForPurpose(PurposeMCPTransfer); revoked != 1 { + t.Fatalf("expected to revoke exactly 1 MCP stream, got %d", revoked) + } + + select { + case res := <-done: + if res.ok { + t.Fatalf("WaitForAgent must return ok=false after revoke; got ok=true") + } + if res.dur > time.Second { + t.Fatalf("WaitForAgent did not wake up promptly after revoke (took %s); kill switch is not immediate", res.dur) + } + case <-time.After(2 * time.Second): + t.Fatalf("WaitForAgent never returned after revoke; kill switch did not wake the waiter") + } +} diff --git a/service/singleton/crontask.go b/service/singleton/crontask.go index 85418c2c..0d846ca2 100644 --- a/service/singleton/crontask.go +++ b/service/singleton/crontask.go @@ -264,13 +264,12 @@ func CronTrigger(cr *model.Cron, triggerServer ...uint64) func() { if !cronCanSendToServer(cr, s) { return } - stream := s.GetTaskStream() - if stream != nil { + if s.GetTaskStream() != nil { cronShared := CronShared if cronShared != nil { cronShared.reserveAlertTriggerCronResult(cr.ID, s.ID) } - if err := stream.Send(&pb.Task{ + if err := s.SendTask(&pb.Task{ Id: cr.ID, Data: cr.Command, Type: model.TaskTypeCommand, @@ -287,7 +286,13 @@ func CronTrigger(cr *model.Cron, triggerServer ...uint64) func() { return } - for _, s := range ServerShared.Range { + // 先在锁内快照 server 列表再逐个 SendTask:ServerShared.Range 会在整个 + // 回调期间持 listMu.RLock,而 SendTask 走阻塞 gRPC,一个卡死的 agent + // 会让需要写锁的 server 编辑/删除被拖死。GetList 克隆后即释放锁。 + for _, s := range ServerShared.GetList() { + if s == nil { + continue + } if !cronCanSendToServer(cr, s) { continue } @@ -297,8 +302,8 @@ func CronTrigger(cr *model.Cron, triggerServer ...uint64) func() { if cr.Cover == model.CronCoverIgnoreAll && !crIgnoreMap[s.ID] { continue } - if stream := s.GetTaskStream(); stream != nil { - stream.Send(&pb.Task{ + if s.GetTaskStream() != nil { + _ = s.SendTask(&pb.Task{ Id: cr.ID, Data: cr.Command, Type: model.TaskTypeCommand, diff --git a/service/singleton/server.go b/service/singleton/server.go index f17c9f2f..f2822502 100644 --- a/service/singleton/server.go +++ b/service/singleton/server.go @@ -38,9 +38,39 @@ func NewServerClass() *ServerClass { } sc.sortList() + model.OwnerServerIDsLookup = sc.ownerServerIDs + model.AllServerIDsLookup = sc.allServerIDs + model.OwnerIsAdminLookup = ownerIsAdmin + return sc } +func (c *ServerClass) ownerServerIDs(ownerUID uint64) []uint64 { + var ids []uint64 + c.Range(func(id uint64, s *model.Server) bool { + if s != nil && s.GetUserID() == ownerUID { + ids = append(ids, id) + } + return true + }) + return ids +} + +func (c *ServerClass) allServerIDs() []uint64 { + var ids []uint64 + c.Range(func(id uint64, s *model.Server) bool { + if s != nil { + ids = append(ids, id) + } + return true + }) + return ids +} + +func ownerIsAdmin(ownerUID uint64) bool { + return userIsAdmin(ownerUID) +} + func (c *ServerClass) Update(s *model.Server, uuid string) { c.listMu.Lock() @@ -64,8 +94,11 @@ func (c *ServerClass) Delete(idList []uint64) { c.listMu.Lock() for _, id := range idList { - serverUUID := c.list[id].UUID - delete(c.uuidToID, serverUUID) + s, ok := c.list[id] + if !ok { + continue + } + delete(c.uuidToID, s.UUID) delete(c.list, id) } diff --git a/service/singleton/server_delete_missing_test.go b/service/singleton/server_delete_missing_test.go new file mode 100644 index 00000000..5a40ccde --- /dev/null +++ b/service/singleton/server_delete_missing_test.go @@ -0,0 +1,32 @@ +package singleton + +import ( + "testing" + + "github.com/nezhahq/nezha/model" +) + +func TestServerClassDeleteMissingIDNoPanic(t *testing.T) { + c := &ServerClass{ + class: class[uint64, *model.Server]{ + list: map[uint64]*model.Server{ + 1: {Common: model.Common{ID: 1}, UUID: "uuid-1"}, + }, + }, + uuidToID: map[string]uint64{"uuid-1": 1}, + } + + c.Delete([]uint64{999999}) + + if _, ok := c.list[1]; !ok { + t.Fatalf("existing server 1 must remain after deleting a non-existent id") + } + + c.Delete([]uint64{1, 424242}) + if _, ok := c.list[1]; ok { + t.Fatalf("server 1 should be removed") + } + if _, ok := c.uuidToID["uuid-1"]; ok { + t.Fatalf("uuid mapping for server 1 should be removed") + } +} diff --git a/service/singleton/server_transfer.go b/service/singleton/server_transfer.go index 7508f776..901f802a 100644 --- a/service/singleton/server_transfer.go +++ b/service/singleton/server_transfer.go @@ -232,20 +232,6 @@ func NewServerTransferClass() *ServerTransferClass { } c.pending[t.ServerID] = &t } - for i := range pending { - t := pending[i] - // Skip ghost rows whose server has been deleted out from under - // the transfer (e.g. before OnServersDeleted existed, or because - // the row predates this branch). Loading them would resurrect a - // HasPending state that no longer corresponds to a real server - // and the timeout sweeper would log errors every 30s without - // being able to settle the row. - if s, ok := ServerShared.Get(t.ServerID); !ok || s == nil { - log.Printf("NEZHA>> ServerTransferClass: ignoring pending transfer %d for missing server %d (likely a leftover from before OnServersDeleted was wired)", t.ID, t.ServerID) - continue - } - c.pending[t.ServerID] = &t - } var reverted []model.ServerTransfer // acked_at IS NULL is non-negotiable: MarkRevertDelivered persists @@ -927,7 +913,14 @@ func (c *ServerTransferClass) sendApplyConfigTask(s *model.Server, stream pb.Nez // Keep Send synchronous under the per-server lock. A goroutine+timeout cannot // cancel grpc.ServerStream.Send; returning early would let a stale new-secret // ApplyConfig complete after a cancel/fail revert and overwrite the rollback. - if err := stream.Send(task); err != nil { + // + // Route through Server.SendTask so the holder-scoped send mutex is + // honoured: cron / MCP CallAgent / MCP fs.transfer dispatch on the same + // gRPC stream and would otherwise race grpc-go's one-SendMsg-per-stream + // invariant. The captured stream argument is still passed to + // ClearTaskStreamIfCurrent so a reconnect mid-Send cannot wipe a newer + // published stream when Send fails on the stale one. + if err := s.SendTask(task); err != nil { log.Printf("NEZHA>> ServerTransfer ApplyConfig send failed: serverID=%d transferID=%d: %v", s.ID, task.Id, err) s.ClearTaskStreamIfCurrent(stream) return err diff --git a/service/singleton/singleton.go b/service/singleton/singleton.go index 1fbc39c8..e4817d5d 100644 --- a/service/singleton/singleton.go +++ b/service/singleton/singleton.go @@ -94,11 +94,21 @@ func InitDBFromPath(path string) error { model.Notification{}, model.AlertRule{}, model.Service{}, model.NotificationGroupNotification{}, model.Cron{}, model.Transfer{}, model.ServerGroupServer{}, model.NAT{}, model.DDNSProfile{}, model.NotificationGroupNotification{}, - model.WAF{}, model.Oauth2Bind{}, model.ServerTransfer{}, model.JWTSession{}) + model.WAF{}, model.Oauth2Bind{}, model.ServerTransfer{}, model.JWTSession{}, + model.APIToken{}, model.MCPAuditLog{}) if err != nil { return err } + // 旧 mcp:* scope 与 nezha:* 并行了一段时间,HasScope 通过别名让 mcp:fs:write + // 静默扩到 REST nezha:server:write。统一命名后这里把残留旧 scope 一次性 + // 归一化(或在仅剩危险旧 scope 时整张 PAT 删除),保证运行时不再依赖别名。 + if rewritten, deleted, mErr := model.MigrateLegacyMCPScopes(DB); mErr != nil { + log.Printf("NEZHA>> MigrateLegacyMCPScopes failed: %v", mErr) + } else if rewritten > 0 || deleted > 0 { + log.Printf("NEZHA>> Migrated legacy mcp:* api token scopes: rewritten=%d deleted=%d", rewritten, deleted) + } + return nil } diff --git a/service/singleton/testhelpers.go b/service/singleton/testhelpers.go new file mode 100644 index 00000000..459a531a --- /dev/null +++ b/service/singleton/testhelpers.go @@ -0,0 +1,65 @@ +package singleton + +import "github.com/nezhahq/nezha/model" + +// NewEmptyServerClassForTest 构造一个不依赖 DB 的空 ServerClass,仅用于单测。 +// 生产路径请用 NewServerClass。 +func NewEmptyServerClassForTest() *ServerClass { + sc := &ServerClass{ + class: class[uint64, *model.Server]{ + list: make(map[uint64]*model.Server), + }, + uuidToID: make(map[string]uint64), + } + model.OwnerServerIDsLookup = sc.ownerServerIDs + model.AllServerIDsLookup = sc.allServerIDs + model.OwnerIsAdminLookup = ownerIsAdmin + return sc +} + +// InsertForTest 把一个 server 直接塞进内存表与排序快照,跳过 DB & InitServer 逻辑。 +// 调用方需保证 server.ID 已经设置。 +func (c *ServerClass) InsertForTest(s *model.Server) { + c.listMu.Lock() + c.list[s.ID] = s + if s.UUID != "" { + c.uuidToID[s.UUID] = s.ID + } + c.listMu.Unlock() + c.sortList() +} + +// NewEmptyDDNSClassForTest 构造一个不依赖 DB 的空 DDNSClass,仅用于单测。 +func NewEmptyDDNSClassForTest() *DDNSClass { + return &DDNSClass{ + class: class[uint64, *model.DDNSProfile]{ + list: make(map[uint64]*model.DDNSProfile), + }, + } +} + +// InsertForTest 把一个 DDNS profile 直接塞进内存表,跳过 DB。 +func (c *DDNSClass) InsertForTest(p *model.DDNSProfile) { + c.listMu.Lock() + c.list[p.ID] = p + c.listMu.Unlock() +} + +// NewEmptyNotificationClassForTest 构造空 NotificationClass。 +func NewEmptyNotificationClassForTest() *NotificationClass { + return &NotificationClass{ + class: class[uint64, *model.Notification]{ + list: make(map[uint64]*model.Notification), + }, + groupToIDList: make(map[uint64]map[uint64]*model.Notification), + idToGroupList: make(map[uint64]map[uint64]struct{}), + groupList: make(map[uint64]string), + } +} + +// InsertForTest 把一个 Notification 直接塞进内存表。 +func (c *NotificationClass) InsertForTest(n *model.Notification) { + c.listMu.Lock() + c.list[n.ID] = n + c.listMu.Unlock() +}