From 2640b86d3ac7b205ff66d1679c99a1686fa0ee4a Mon Sep 17 00:00:00 2001 From: naiba Date: Mon, 20 Jul 2026 04:30:50 +0000 Subject: [PATCH] feat(agentcompat): expose dashboard capability routes Co-authored-by: naiba/CloudCode --- .../agentcompat_capability_access.go | 86 ++++++ .../agentcompat_capability_boundary_test.go | 193 ++++++++++++++ .../agentcompat_capability_contract.go | 199 ++++++++++++++ .../agentcompat_capability_default_test.go | 24 ++ .../agentcompat_capability_error_test.go | 85 ++++++ .../agentcompat_capability_handlers.go | 114 ++++++++ .../agentcompat_capability_nat_test.go | 109 ++++++++ .../agentcompat_capability_routes_test.go | 175 +++++++++++++ .../agentcompat_capability_security_test.go | 238 +++++++++++++++++ .../controller/agentcompat_routes.go | 147 +++++++++++ .../controller/agentcompat_routes_default.go | 7 + .../agentcompat_routes_security_test.go | 69 +++++ ...qlite_hold_error_agentcompat_linux_test.go | 40 +++ ...at_sqlite_hold_routes_agentcompat_linux.go | 168 ++++++++++++ ...lite_hold_routes_agentcompat_linux_test.go | 201 ++++++++++++++ ...sqlite_hold_routes_agentcompat_nonlinux.go | 7 + ...tcompat_sqlite_hold_routes_default_test.go | 27 ++ cmd/dashboard/controller/controller.go | 1 + cmd/dashboard/controller/fm.go | 30 +-- .../io_stream_state_agentcompat_test.go | 214 +++++++++++++++ cmd/dashboard/controller/terminal.go | 30 +-- .../controller/terminal_fm_agentcompat.go | 64 +++++ .../terminal_fm_agentcompat_cleanup_test.go | 130 +++++++++ .../terminal_fm_agentcompat_default.go | 18 ++ .../terminal_fm_agentcompat_default_test.go | 69 +++++ .../terminal_fm_agentcompat_rejection_test.go | 101 +++++++ .../terminal_fm_agentcompat_test.go | 213 +++++++++++++++ ...minal_fm_agentcompat_wrong_purpose_test.go | 92 +++++++ .../controller/terminal_fm_lifecycle_test.go | 122 +++++++++ cmd/dashboard/controller/ws.go | 11 +- cmd/dashboard/rpc/grpc_interceptor.go | 94 +++++++ cmd/dashboard/rpc/nat.go | 115 ++++++++ .../rpc/nat_capability_agentcompat.go | 48 ++++ .../rpc/nat_capability_agentcompat_test.go | 82 ++++++ cmd/dashboard/rpc/nat_capability_default.go | 25 ++ .../rpc/nat_capability_default_flow_test.go | 64 +++++ .../rpc/nat_capability_default_test.go | 35 +++ ...at_capability_failures_agentcompat_test.go | 245 +++++++++++++++++ .../nat_capability_flow_agentcompat_test.go | 151 +++++++++++ cmd/dashboard/rpc/nat_test.go | 93 +++++++ cmd/dashboard/rpc/nat_test_support_test.go | 171 ++++++++++++ cmd/dashboard/rpc/rpc.go | 247 +----------------- 42 files changed, 4083 insertions(+), 271 deletions(-) create mode 100644 cmd/dashboard/controller/agentcompat_capability_access.go create mode 100644 cmd/dashboard/controller/agentcompat_capability_boundary_test.go create mode 100644 cmd/dashboard/controller/agentcompat_capability_contract.go create mode 100644 cmd/dashboard/controller/agentcompat_capability_default_test.go create mode 100644 cmd/dashboard/controller/agentcompat_capability_error_test.go create mode 100644 cmd/dashboard/controller/agentcompat_capability_handlers.go create mode 100644 cmd/dashboard/controller/agentcompat_capability_nat_test.go create mode 100644 cmd/dashboard/controller/agentcompat_capability_routes_test.go create mode 100644 cmd/dashboard/controller/agentcompat_capability_security_test.go create mode 100644 cmd/dashboard/controller/agentcompat_routes.go create mode 100644 cmd/dashboard/controller/agentcompat_routes_default.go create mode 100644 cmd/dashboard/controller/agentcompat_routes_security_test.go create mode 100644 cmd/dashboard/controller/agentcompat_sqlite_hold_error_agentcompat_linux_test.go create mode 100644 cmd/dashboard/controller/agentcompat_sqlite_hold_routes_agentcompat_linux.go create mode 100644 cmd/dashboard/controller/agentcompat_sqlite_hold_routes_agentcompat_linux_test.go create mode 100644 cmd/dashboard/controller/agentcompat_sqlite_hold_routes_agentcompat_nonlinux.go create mode 100644 cmd/dashboard/controller/agentcompat_sqlite_hold_routes_default_test.go create mode 100644 cmd/dashboard/controller/io_stream_state_agentcompat_test.go create mode 100644 cmd/dashboard/controller/terminal_fm_agentcompat.go create mode 100644 cmd/dashboard/controller/terminal_fm_agentcompat_cleanup_test.go create mode 100644 cmd/dashboard/controller/terminal_fm_agentcompat_default.go create mode 100644 cmd/dashboard/controller/terminal_fm_agentcompat_default_test.go create mode 100644 cmd/dashboard/controller/terminal_fm_agentcompat_rejection_test.go create mode 100644 cmd/dashboard/controller/terminal_fm_agentcompat_test.go create mode 100644 cmd/dashboard/controller/terminal_fm_agentcompat_wrong_purpose_test.go create mode 100644 cmd/dashboard/controller/terminal_fm_lifecycle_test.go create mode 100644 cmd/dashboard/rpc/grpc_interceptor.go create mode 100644 cmd/dashboard/rpc/nat.go create mode 100644 cmd/dashboard/rpc/nat_capability_agentcompat.go create mode 100644 cmd/dashboard/rpc/nat_capability_agentcompat_test.go create mode 100644 cmd/dashboard/rpc/nat_capability_default.go create mode 100644 cmd/dashboard/rpc/nat_capability_default_flow_test.go create mode 100644 cmd/dashboard/rpc/nat_capability_default_test.go create mode 100644 cmd/dashboard/rpc/nat_capability_failures_agentcompat_test.go create mode 100644 cmd/dashboard/rpc/nat_capability_flow_agentcompat_test.go create mode 100644 cmd/dashboard/rpc/nat_test.go create mode 100644 cmd/dashboard/rpc/nat_test_support_test.go diff --git a/cmd/dashboard/controller/agentcompat_capability_access.go b/cmd/dashboard/controller/agentcompat_capability_access.go new file mode 100644 index 00000000..453fd6be --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_capability_access.go @@ -0,0 +1,86 @@ +//go:build agentcompat + +package controller + +import ( + "github.com/gin-gonic/gin" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/rpc" + "github.com/nezhahq/nezha/service/singleton" +) + +type agentcompatCapabilityIdentity struct { + purpose rpc.AgentCompatCapabilityPurpose + serverID uint64 + resourceID uint64 +} + +type agentcompatCapabilityProof struct { + owner rpc.AgentCompatCapabilityOwner + identity agentcompatCapabilityIdentity +} + +func currentAgentcompatCapabilityProof(context *gin.Context, identity agentcompatCapabilityIdentity) (agentcompatCapabilityProof, error) { + if err := validateAgentcompatCapabilityIdentity(identity.purpose, identity.serverID, identity.resourceID); err != nil { + return agentcompatCapabilityProof{}, err + } + token := APITokenFromContext(context) + authorized, present := context.Get(model.CtxKeyAuthorizedUser) + user, validUser := authorized.(*model.User) + if token == nil || token.ID == 0 || !present || !validUser || user == nil || user.ID == 0 { + return agentcompatCapabilityProof{}, errAgentcompatCapabilityUnavailable + } + if singleton.ServerShared == nil { + return agentcompatCapabilityProof{}, errAgentcompatCapabilityUnavailable + } + server, exists := singleton.ServerShared.Get(identity.serverID) + if !exists || server == nil || !server.HasPermission(context) || !patAllowsServer(context, identity.serverID) { + return agentcompatCapabilityProof{}, errAgentcompatCapabilityUnavailable + } + if identity.purpose == rpc.AgentCompatCapabilityNAT && !currentAgentcompatNATPermission(context, identity) { + return agentcompatCapabilityProof{}, errAgentcompatCapabilityUnavailable + } + return agentcompatCapabilityProof{ + owner: rpc.AgentCompatCapabilityOwner{PATID: token.ID, UserID: user.ID, IsAdmin: user.Role.IsAdmin()}, + identity: identity, + }, nil +} + +func currentAgentcompatNATPermission(context *gin.Context, identity agentcompatCapabilityIdentity) bool { + if singleton.NATShared == nil { + return false + } + domain := singleton.NATShared.GetDomain(identity.resourceID) + if domain == "" { + return false + } + profile := singleton.NATShared.GetNATConfigByDomain(domain) + return profile != nil && profile.ID == identity.resourceID && profile.ServerID == identity.serverID && profile.HasPermission(context) +} + +func (proof agentcompatCapabilityProof) registration() rpc.AgentCompatCapabilityRegistration { + return rpc.AgentCompatCapabilityRegistration{ + Owner: proof.owner, Purpose: proof.identity.purpose, TargetServerID: proof.identity.serverID, + ResourceID: proof.identity.resourceID, ServerAccessAllowed: true, + } +} + +func (proof agentcompatCapabilityProof) access(capability rpc.AgentCompatIOStreamCapability) rpc.AgentCompatCapabilityAccess { + return rpc.AgentCompatCapabilityAccess{ + Capability: capability, Owner: proof.owner, Purpose: proof.identity.purpose, + TargetServerID: proof.identity.serverID, ResourceID: proof.identity.resourceID, ServerAccessAllowed: true, + } +} + +func agentcompatCapabilityIdentityFromWire(purpose string, serverID, resourceID uint64) (agentcompatCapabilityIdentity, error) { + parsedPurpose, err := parseAgentcompatCapabilityPurpose(purpose) + if err != nil { + return agentcompatCapabilityIdentity{}, err + } + identity := agentcompatCapabilityIdentity{purpose: parsedPurpose, serverID: serverID, resourceID: resourceID} + if err := validateAgentcompatCapabilityIdentity(identity.purpose, identity.serverID, identity.resourceID); err != nil { + return agentcompatCapabilityIdentity{}, err + } + return identity, nil +} diff --git a/cmd/dashboard/controller/agentcompat_capability_boundary_test.go b/cmd/dashboard/controller/agentcompat_capability_boundary_test.go new file mode 100644 index 00000000..b58cba55 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_capability_boundary_test.go @@ -0,0 +1,193 @@ +//go:build agentcompat + +package controller + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net/http" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/rpc" +) + +func TestAgentcompatCapabilityBoundaryAcceptsExactly512Bytes(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + server := newAgentcompatCapabilityServer(t) + + body := padAgentcompatCapabilityBody(t, `{"purpose":"terminal","server_id":7}`) + status, responseBody := postAgentcompatBoundary(t, server.URL+agentcompatCapabilityRegisterPath, token, body, true) + require.Equal(t, http.StatusOK, status) + var envelope model.CommonResponse[agentcompatCapabilityRegisterResponse] + require.NoError(t, json.Unmarshal([]byte(responseBody), &envelope)) + require.True(t, envelope.Success) + require.NotEmpty(t, envelope.Data.Capability) +} + +func TestAgentcompatCapabilityBoundaryRejects513BytesWithoutDisclosure(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + server := newAgentcompatCapabilityServer(t) + body := padAgentcompatCapabilityBody(t, `{"purpose":"terminal","server_id":7}`) + " " + + for _, route := range []struct { + name string + path string + wantError bool + wantSuccess bool + }{ + {name: "register", path: agentcompatCapabilityRegisterPath, wantError: true}, + {name: "wait", path: agentcompatCapabilityWaitPath, wantError: true}, + {name: "cancel", path: agentcompatCapabilityCancelPath, wantSuccess: true}, + {name: "unregister", path: agentcompatCapabilityUnregisterPath, wantSuccess: true}, + } { + t.Run(route.name, func(t *testing.T) { + status, responseBody := postAgentcompatBoundary(t, server.URL+route.path, token, body, true) + require.Equal(t, http.StatusOK, status) + require.NotContains(t, responseBody, "terminal") + require.NotContains(t, responseBody, "server_id") + if route.wantError { + var envelope model.CommonResponse[agentcompatCapabilityRegisterResponse] + require.NoError(t, json.Unmarshal([]byte(responseBody), &envelope)) + require.False(t, envelope.Success) + require.Equal(t, errAgentcompatCapabilityInvalid.Error(), envelope.Error) + return + } + require.Equal(t, `{"success":true,"data":{}}`, responseBody) + }) + } +} + +func TestAgentcompatCapabilityBoundaryRejectsChunkedOversizeWithoutDisclosure(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + server := newAgentcompatCapabilityServer(t) + body := padAgentcompatCapabilityBody(t, `{"purpose":"terminal","server_id":7}`) + " " + + status, responseBody := postAgentcompatBoundary(t, server.URL+agentcompatCapabilityRegisterPath, token, body, false) + require.Equal(t, http.StatusOK, status) + require.Contains(t, responseBody, errAgentcompatCapabilityInvalid.Error()) + require.NotContains(t, responseBody, "terminal") + require.NotContains(t, responseBody, "server_id") +} + +func TestAgentcompatCapabilityBoundaryRejectsDuplicateKeysInEitherOrder(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + server := newAgentcompatCapabilityServer(t) + + keys := []string{"purpose", "server_id", "resource_id"} + for _, key := range keys { + for _, body := range []string{ + `{"` + key + `":"terminal","` + key + `":"terminal","purpose":"terminal","server_id":7}`, + `{"` + key + `":0,"purpose":"terminal","server_id":7,"` + key + `":"terminal"}`, + } { + status, responseBody := postAgentcompatBoundary(t, server.URL+agentcompatCapabilityRegisterPath, token, body, true) + require.Equal(t, http.StatusOK, status) + require.Contains(t, responseBody, errAgentcompatCapabilityInvalid.Error()) + require.NotContains(t, responseBody, "terminal") + } + } +} + +func TestAgentcompatCapabilityBoundaryRejectsAccessGrammar(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + server := newAgentcompatCapabilityServer(t) + + cases := []string{ + `{"capability":"x","purpose":"terminal"}`, + `{"capability":"x","server_id":7}`, + `{"capability":"x","purpose":"terminal","server_id":7,"capability":"y"}`, + `{"capability":"x","purpose":"terminal","server_id":0}`, + `{"capability":"x","purpose":"terminal","server_id":7,"unknown":"private"}`, + `{"capability":"x","purpose":"terminal","server_id":7} {}`, + `{"capability":null,"purpose":"terminal","server_id":7}`, + `{"capability":7,"purpose":"terminal","server_id":7}`, + `{"capability":"x","purpose":7,"server_id":7}`, + `{"capability":"x","purpose":"terminal","server_id":"7"}`, + `{"capability":"x","purpose":"terminal","server_id":-1}`, + `{"capability":"x","purpose":"terminal","server_id":1.5}`, + `{"capability":"x","purpose":"terminal","server_id":1e2}`, + `null`, `[]`, `{`, + } + for _, body := range cases { + status, responseBody := postAgentcompatBoundary(t, server.URL+agentcompatCapabilityWaitPath, token, body, true) + require.Equal(t, http.StatusOK, status) + require.Contains(t, responseBody, errAgentcompatCapabilityInvalid.Error()) + require.NotContains(t, responseBody, "private") + require.NotContains(t, responseBody, "x") + } +} + +func TestAgentcompatCapabilityBoundaryRejectsMissingRegisterFields(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + server := newAgentcompatCapabilityServer(t) + + for _, body := range []string{`{"server_id":7}`, `{"purpose":"terminal"}`, `{}`} { + status, responseBody := postAgentcompatBoundary(t, server.URL+agentcompatCapabilityRegisterPath, token, body, true) + require.Equal(t, http.StatusOK, status) + require.Contains(t, responseBody, errAgentcompatCapabilityInvalid.Error()) + require.NotContains(t, responseBody, "terminal") + } +} + +func padAgentcompatCapabilityBody(t *testing.T, body string) string { + t.Helper() + require.LessOrEqual(t, len(body), 512) + return body + strings.Repeat(" ", 512-len(body)) +} + +func postAgentcompatBoundary(t *testing.T, path, token, body string, contentLength bool) (int, string) { + t.Helper() + requestContext, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + var reader io.Reader = bytes.NewBufferString(body) + request, err := http.NewRequestWithContext(requestContext, http.MethodPost, path, reader) + require.NoError(t, err) + if !contentLength { + request.Body = io.NopCloser(strings.NewReader(body)) + request.ContentLength = -1 + } + request.Header.Set("Authorization", "Bearer "+token) + request.Header.Set("Content-Type", "application/json") + response, err := http.DefaultClient.Do(request) + require.NoError(t, err) + defer response.Body.Close() + responseBody, err := io.ReadAll(response.Body) + require.NoError(t, err) + return response.StatusCode, string(responseBody) +} diff --git a/cmd/dashboard/controller/agentcompat_capability_contract.go b/cmd/dashboard/controller/agentcompat_capability_contract.go new file mode 100644 index 00000000..cc027a01 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_capability_contract.go @@ -0,0 +1,199 @@ +//go:build agentcompat + +package controller + +import ( + "encoding/json" + "errors" + "io" + "net/http" + "strconv" + "strings" + + "github.com/gin-gonic/gin" + + "github.com/nezhahq/nezha/service/rpc" +) + +const ( + agentcompatCapabilityRequestMaxBytes = 512 + + agentcompatCapabilityRegisterPath = "/agentcompat/io-stream-capability/register" + agentcompatCapabilityWaitPath = "/agentcompat/io-stream-capability/wait" + agentcompatCapabilityCancelPath = "/agentcompat/io-stream-capability/cancel" + agentcompatCapabilityUnregisterPath = "/agentcompat/io-stream-capability/unregister" +) + +const ( + agentcompatCapabilityPurposeTerminal = "terminal" + agentcompatCapabilityPurposeFileManager = "file_manager" + agentcompatCapabilityPurposeNAT = "nat" +) + +var ( + errAgentcompatCapabilityInvalid = errors.New("agentcompat capability request is invalid") + errAgentcompatCapabilityUnavailable = errors.New("agentcompat capability is not available") + errAgentcompatCapabilityConflict = errors.New("agentcompat capability is active") + errAgentcompatCapabilityCleanup = errors.New("agentcompat capability cleanup failed") +) + +type agentcompatCapabilityRegisterRequest struct { + Purpose string `json:"purpose"` + ServerID uint64 `json:"server_id"` + ResourceID uint64 `json:"resource_id,omitempty"` +} + +type agentcompatCapabilityRegisterResponse struct { + Capability string `json:"capability"` +} + +type agentcompatCapabilityAccessRequest struct { + Capability string `json:"capability"` + Purpose string `json:"purpose"` + ServerID uint64 `json:"server_id"` + ResourceID uint64 `json:"resource_id,omitempty"` +} + +type agentcompatCapabilityWaitRequest = agentcompatCapabilityAccessRequest +type agentcompatCapabilityCancelRequest = agentcompatCapabilityAccessRequest +type agentcompatCapabilityUnregisterRequest = agentcompatCapabilityAccessRequest + +type agentcompatCapabilityWaitResponse struct { + StreamID string `json:"stream_id"` +} + +type agentcompatCapabilityEmptyResponse struct{} + +func decodeAgentcompatCapabilityRequest[T any](context *gin.Context, destination *T) error { + context.Request.Body = http.MaxBytesReader(context.Writer, context.Request.Body, agentcompatCapabilityRequestMaxBytes) + decoder := json.NewDecoder(context.Request.Body) + fields, err := decodeAgentcompatCapabilityObject(decoder) + if err != nil { + return errAgentcompatCapabilityInvalid + } + switch request := any(destination).(type) { + case *agentcompatCapabilityRegisterRequest: + request.Purpose, err = decodeAgentcompatCapabilityString(fields.purpose) + if err == nil { + request.ServerID, err = decodeAgentcompatCapabilityUint64(fields.serverID) + } + if err == nil && fields.resourceID != nil { + request.ResourceID, err = decodeAgentcompatCapabilityUint64(fields.resourceID) + } + if err != nil || fields.purpose == nil || fields.serverID == nil { + return errAgentcompatCapabilityInvalid + } + case *agentcompatCapabilityAccessRequest: + request.Capability, err = decodeAgentcompatCapabilityString(fields.capability) + if err == nil { + request.Purpose, err = decodeAgentcompatCapabilityString(fields.purpose) + } + if err == nil { + request.ServerID, err = decodeAgentcompatCapabilityUint64(fields.serverID) + } + if err == nil && fields.resourceID != nil { + request.ResourceID, err = decodeAgentcompatCapabilityUint64(fields.resourceID) + } + if err != nil || fields.capability == nil || fields.purpose == nil || fields.serverID == nil { + return errAgentcompatCapabilityInvalid + } + default: + return errAgentcompatCapabilityInvalid + } + return nil +} + +type agentcompatCapabilityFields struct { + purpose json.RawMessage + serverID json.RawMessage + resourceID json.RawMessage + capability json.RawMessage + seen uint8 +} + +func decodeAgentcompatCapabilityObject(decoder *json.Decoder) (agentcompatCapabilityFields, error) { + var fields agentcompatCapabilityFields + token, err := decoder.Token() + if err != nil || token != json.Delim('{') { + return fields, errAgentcompatCapabilityInvalid + } + for decoder.More() { + keyToken, err := decoder.Token() + key, validKey := keyToken.(string) + if err != nil || !validKey { + return fields, errAgentcompatCapabilityInvalid + } + var bit uint8 + var target *json.RawMessage + switch key { + case "purpose": + bit, target = 1, &fields.purpose + case "server_id": + bit, target = 2, &fields.serverID + case "resource_id": + bit, target = 4, &fields.resourceID + case "capability": + bit, target = 8, &fields.capability + default: + return fields, errAgentcompatCapabilityInvalid + } + if fields.seen&bit != 0 || decoder.Decode(target) != nil { + return fields, errAgentcompatCapabilityInvalid + } + fields.seen |= bit + } + if token, err = decoder.Token(); err != nil || token != json.Delim('}') { + return fields, errAgentcompatCapabilityInvalid + } + if _, err = decoder.Token(); !errors.Is(err, io.EOF) { + return fields, errAgentcompatCapabilityInvalid + } + return fields, nil +} + +func decodeAgentcompatCapabilityString(raw json.RawMessage) (string, error) { + if strings.TrimSpace(string(raw)) == "null" { + return "", errAgentcompatCapabilityInvalid + } + var value string + if err := json.Unmarshal(raw, &value); err != nil { + return "", errAgentcompatCapabilityInvalid + } + return value, nil +} + +func decodeAgentcompatCapabilityUint64(raw json.RawMessage) (uint64, error) { + return strconv.ParseUint(strings.TrimSpace(string(raw)), 10, 64) +} + +func parseAgentcompatCapabilityPurpose(value string) (rpc.AgentCompatCapabilityPurpose, error) { + switch value { + case agentcompatCapabilityPurposeTerminal: + return rpc.AgentCompatCapabilityTerminal, nil + case agentcompatCapabilityPurposeFileManager: + return rpc.AgentCompatCapabilityFileManager, nil + case agentcompatCapabilityPurposeNAT: + return rpc.AgentCompatCapabilityNAT, nil + default: + return 0, errAgentcompatCapabilityInvalid + } +} + +func validateAgentcompatCapabilityIdentity(purpose rpc.AgentCompatCapabilityPurpose, serverID, resourceID uint64) error { + if serverID == 0 { + return errAgentcompatCapabilityInvalid + } + switch purpose { + case rpc.AgentCompatCapabilityTerminal, rpc.AgentCompatCapabilityFileManager: + if resourceID != 0 { + return errAgentcompatCapabilityInvalid + } + case rpc.AgentCompatCapabilityNAT: + if resourceID == 0 { + return errAgentcompatCapabilityInvalid + } + default: + return errAgentcompatCapabilityInvalid + } + return nil +} diff --git a/cmd/dashboard/controller/agentcompat_capability_default_test.go b/cmd/dashboard/controller/agentcompat_capability_default_test.go new file mode 100644 index 00000000..a60c6109 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_capability_default_test.go @@ -0,0 +1,24 @@ +//go:build !agentcompat + +package controller + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestDefaultBuildDoesNotRegisterAgentcompatCapabilityRoutes(t *testing.T) { + gin.SetMode(gin.TestMode) + router := gin.New() + registerAgentcompatRoutes(router) + request := httptest.NewRequest(http.MethodPost, "/agentcompat/io-stream-capability/register", nil) + response := httptest.NewRecorder() + + router.ServeHTTP(response, request) + + require.Equal(t, http.StatusNotFound, response.Code) +} diff --git a/cmd/dashboard/controller/agentcompat_capability_error_test.go b/cmd/dashboard/controller/agentcompat_capability_error_test.go new file mode 100644 index 00000000..941956de --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_capability_error_test.go @@ -0,0 +1,85 @@ +//go:build agentcompat + +package controller + +import ( + "context" + "errors" + "io" + "net/http" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/rpc" +) + +type agentcompatCapabilityCloseFailure struct { + closed atomic.Int32 +} + +func (*agentcompatCapabilityCloseFailure) Read([]byte) (int, error) { return 0, io.EOF } +func (*agentcompatCapabilityCloseFailure) Write(data []byte) (int, error) { return len(data), nil } +func (failure *agentcompatCapabilityCloseFailure) Close() error { + failure.closed.Add(1) + return errors.New("private-stream private-capability private-server") +} + +func TestAgentcompatCapabilityBoundUnregisterAndCancelCleanupErrorsAreGeneric(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + handler := rpc.NewNezhaHandler() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = handler + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + token, plaintext := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + server := newAgentcompatCapabilityServer(t) + capability := registerAgentcompatCapability(t, server.URL, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "file_manager", ServerID: 7}) + parsed, err := rpc.ParseAgentCompatIOStreamCapability(capability) + require.NoError(t, err) + access := rpc.AgentCompatCapabilityAccess{ + Capability: parsed, Owner: rpc.AgentCompatCapabilityOwner{PATID: token.ID, UserID: userID}, + Purpose: rpc.AgentCompatCapabilityFileManager, TargetServerID: 7, ServerAccessAllowed: true, + } + require.NoError(t, handler.CreateStreamWithPurpose("private-cleanup-stream", userID, 7, rpc.PurposeFileManager)) + require.NoError(t, handler.BindAgentCompatIOStreamCapability(rpc.AgentCompatCapabilityBinding{AgentCompatCapabilityAccess: access, StreamID: "private-cleanup-stream"})) + + unregisterStatus, unregisterBody := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityUnregisterPath, plaintext, agentcompatCapabilityUnregisterRequest{Capability: capability, Purpose: "file_manager", ServerID: 7}) + require.Equal(t, http.StatusOK, unregisterStatus) + require.Contains(t, unregisterBody, errAgentcompatCapabilityConflict.Error()) + require.NotContains(t, unregisterBody, capability) + require.NotContains(t, unregisterBody, "private-cleanup-stream") + + failure := &agentcompatCapabilityCloseFailure{} + require.NoError(t, handler.UserConnected("private-cleanup-stream", failure)) + cancelStatus, cancelBody := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityCancelPath, plaintext, agentcompatCapabilityCancelRequest{Capability: capability, Purpose: "file_manager", ServerID: 7}) + require.Equal(t, http.StatusOK, cancelStatus) + require.Contains(t, cancelBody, errAgentcompatCapabilityCleanup.Error()) + require.NotContains(t, cancelBody, capability) + require.NotContains(t, cancelBody, "private-cleanup-stream") + require.Equal(t, int32(1), failure.closed.Load()) +} + +func TestAgentcompatCapabilityWaitHonorsRequestCancellation(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, plaintext := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + server := newAgentcompatCapabilityServer(t) + capability := registerAgentcompatCapability(t, server.URL, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "terminal", ServerID: 7}) + requestContext, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + body := strings.NewReader(capabilityAccessJSON(capability, "terminal", 7, 0)) + request, err := http.NewRequestWithContext(requestContext, http.MethodPost, server.URL+agentcompatCapabilityWaitPath, body) + require.NoError(t, err) + request.Header.Set("Authorization", "Bearer "+plaintext) + request.Header.Set("Content-Type", "application/json") + _, err = server.Client().Do(request) + require.True(t, errors.Is(err, context.DeadlineExceeded)) +} diff --git a/cmd/dashboard/controller/agentcompat_capability_handlers.go b/cmd/dashboard/controller/agentcompat_capability_handlers.go new file mode 100644 index 00000000..e7bc7ba6 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_capability_handlers.go @@ -0,0 +1,114 @@ +//go:build agentcompat + +package controller + +import ( + "errors" + + "github.com/gin-gonic/gin" + + "github.com/nezhahq/nezha/service/rpc" +) + +func registerAgentcompatCapabilityRoutes(router *gin.Engine, patAuth gin.HandlerFunc) { + router.POST(agentcompatCapabilityRegisterPath, patAuth, commonHandler(agentcompatCapabilityRegister)) + router.POST(agentcompatCapabilityWaitPath, patAuth, commonHandler(agentcompatCapabilityWait)) + router.POST(agentcompatCapabilityCancelPath, patAuth, commonHandler(agentcompatCapabilityCancel)) + router.POST(agentcompatCapabilityUnregisterPath, patAuth, commonHandler(agentcompatCapabilityUnregister)) +} + +func agentcompatCapabilityRegister(context *gin.Context) (agentcompatCapabilityRegisterResponse, error) { + var request agentcompatCapabilityRegisterRequest + if err := decodeAgentcompatCapabilityRequest(context, &request); err != nil { + return agentcompatCapabilityRegisterResponse{}, err + } + identity, err := agentcompatCapabilityIdentityFromWire(request.Purpose, request.ServerID, request.ResourceID) + if err != nil { + return agentcompatCapabilityRegisterResponse{}, err + } + proof, err := currentAgentcompatCapabilityProof(context, identity) + if err != nil || rpc.NezhaHandlerSingleton == nil { + return agentcompatCapabilityRegisterResponse{}, errAgentcompatCapabilityUnavailable + } + capability, err := rpc.NezhaHandlerSingleton.RegisterAgentCompatIOStreamCapability(context.Request.Context(), proof.registration()) + if err != nil { + if contextError := context.Request.Context().Err(); contextError != nil { + return agentcompatCapabilityRegisterResponse{}, contextError + } + return agentcompatCapabilityRegisterResponse{}, errAgentcompatCapabilityUnavailable + } + return agentcompatCapabilityRegisterResponse{Capability: capability.String()}, nil +} + +func agentcompatCapabilityWait(context *gin.Context) (agentcompatCapabilityWaitResponse, error) { + var request agentcompatCapabilityWaitRequest + if err := decodeAgentcompatCapabilityRequest(context, &request); err != nil { + return agentcompatCapabilityWaitResponse{}, err + } + access, err := currentAgentcompatCapabilityAccess(context, request) + if errors.Is(err, errAgentcompatCapabilityInvalid) { + return agentcompatCapabilityWaitResponse{}, errAgentcompatCapabilityInvalid + } + if err != nil || rpc.NezhaHandlerSingleton == nil { + return agentcompatCapabilityWaitResponse{}, errAgentcompatCapabilityUnavailable + } + streamID, err := rpc.NezhaHandlerSingleton.WaitAgentCompatIOStreamCapability(context.Request.Context(), access) + if err != nil { + if contextError := context.Request.Context().Err(); contextError != nil { + return agentcompatCapabilityWaitResponse{}, contextError + } + return agentcompatCapabilityWaitResponse{}, errAgentcompatCapabilityUnavailable + } + return agentcompatCapabilityWaitResponse{StreamID: streamID}, nil +} + +func agentcompatCapabilityCancel(context *gin.Context) (agentcompatCapabilityEmptyResponse, error) { + access, available := inertAgentcompatCapabilityAccess(context) + if !available || rpc.NezhaHandlerSingleton == nil { + return agentcompatCapabilityEmptyResponse{}, nil + } + if err := rpc.NezhaHandlerSingleton.CancelAgentCompatIOStreamCapability(access); err != nil { + return agentcompatCapabilityEmptyResponse{}, errAgentcompatCapabilityCleanup + } + return agentcompatCapabilityEmptyResponse{}, nil +} + +func agentcompatCapabilityUnregister(context *gin.Context) (agentcompatCapabilityEmptyResponse, error) { + access, available := inertAgentcompatCapabilityAccess(context) + if !available || rpc.NezhaHandlerSingleton == nil { + return agentcompatCapabilityEmptyResponse{}, nil + } + err := rpc.NezhaHandlerSingleton.UnregisterAgentCompatIOStreamCapability(access) + if errors.Is(err, rpc.ErrAgentCompatCapabilityBound) { + return agentcompatCapabilityEmptyResponse{}, errAgentcompatCapabilityConflict + } + if err != nil { + return agentcompatCapabilityEmptyResponse{}, errAgentcompatCapabilityUnavailable + } + return agentcompatCapabilityEmptyResponse{}, nil +} + +func currentAgentcompatCapabilityAccess(context *gin.Context, request agentcompatCapabilityAccessRequest) (rpc.AgentCompatCapabilityAccess, error) { + identity, err := agentcompatCapabilityIdentityFromWire(request.Purpose, request.ServerID, request.ResourceID) + if err != nil { + return rpc.AgentCompatCapabilityAccess{}, err + } + capability, err := rpc.ParseAgentCompatIOStreamCapability(request.Capability) + if err != nil { + return rpc.AgentCompatCapabilityAccess{}, errAgentcompatCapabilityUnavailable + } + proof, err := currentAgentcompatCapabilityProof(context, identity) + if err != nil { + return rpc.AgentCompatCapabilityAccess{}, errAgentcompatCapabilityUnavailable + } + return proof.access(capability), nil +} + +func inertAgentcompatCapabilityAccess(context *gin.Context) (rpc.AgentCompatCapabilityAccess, bool) { + var request agentcompatCapabilityAccessRequest + if decodeAgentcompatCapabilityRequest(context, &request) != nil { + return rpc.AgentCompatCapabilityAccess{}, false + } + access, err := currentAgentcompatCapabilityAccess(context, request) + return access, err == nil +} diff --git a/cmd/dashboard/controller/agentcompat_capability_nat_test.go b/cmd/dashboard/controller/agentcompat_capability_nat_test.go new file mode 100644 index 00000000..e48d1444 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_capability_nat_test.go @@ -0,0 +1,109 @@ +//go:build agentcompat + +package controller + +import ( + "context" + "net/http" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/rpc" + "github.com/nezhahq/nezha/service/singleton" +) + +func TestAgentcompatNATCapabilityUsesExactCurrentProfilePermission(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + handler := rpc.NewNezhaHandler() + originalHandler := rpc.NezhaHandlerSingleton + originalNAT := singleton.NATShared + rpc.NezhaHandlerSingleton = handler + t.Cleanup(func() { + rpc.NezhaHandlerSingleton = originalHandler + singleton.NATShared = originalNAT + }) + require.NoError(t, singleton.DB.AutoMigrate(&model.NAT{})) + secondServer := &model.Server{Common: model.Common{ID: 8}} + secondServer.SetUserID(userID) + singleton.ServerShared.InsertForTest(secondServer) + profile := &model.NAT{Common: model.Common{ID: 41, UserID: userID}, Name: "private-profile", ServerID: 7, Domain: "private-profile.example"} + foreignProfile := &model.NAT{Common: model.Common{ID: 42, UserID: 999}, Name: "foreign-profile", ServerID: 7, Domain: "foreign-profile.example"} + require.NoError(t, singleton.DB.Create(profile).Error) + require.NoError(t, singleton.DB.Create(foreignProfile).Error) + singleton.NATShared = singleton.NewNATClass() + token, plaintext := mkToken(t, userID, []string{model.ScopeNATRead}, []uint64{7, 8}) + server := newAgentcompatCapabilityServer(t) + + capability := registerAgentcompatCapability(t, server.URL, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "nat", ServerID: 7, ResourceID: 41}) + status, mismatchBody := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityRegisterPath, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "nat", ServerID: 8, ResourceID: 41}) + require.Equal(t, http.StatusOK, status) + require.Contains(t, mismatchBody, errAgentcompatCapabilityUnavailable.Error()) + foreignStatus, foreignBody := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityRegisterPath, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "nat", ServerID: 7, ResourceID: 42}) + require.Equal(t, http.StatusOK, foreignStatus) + require.Contains(t, foreignBody, errAgentcompatCapabilityUnavailable.Error()) + + parsed, err := rpc.ParseAgentCompatIOStreamCapability(capability) + require.NoError(t, err) + access := rpc.AgentCompatCapabilityAccess{ + Capability: parsed, Owner: rpc.AgentCompatCapabilityOwner{PATID: token.ID, UserID: userID}, + Purpose: rpc.AgentCompatCapabilityNAT, TargetServerID: 7, ResourceID: 41, ServerAccessAllowed: true, + } + handle, err := handler.ConsumeAgentCompatNATCapability(access) + require.NoError(t, err) + lease, err := handler.CreateAgentCompatNATStream(handle, "private-nat-stream") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, handler.CloseAgentCompatNATStreamLease(lease)) }) + require.NoError(t, handler.PublishAgentCompatNATStream(handle, rpc.AgentCompatNATPublication{Purpose: rpc.AgentCompatCapabilityNAT, TargetServerID: 7, ResourceID: 41, StreamID: "private-nat-stream"})) + + profile.ServerID = 8 + singleton.NATShared.Update(profile) + _, unavailableErr := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](context.Background(), server.URL+agentcompatCapabilityWaitPath, plaintext, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "nat", ServerID: 7, ResourceID: 41}) + require.Error(t, unavailableErr) + require.NotContains(t, unavailableErr.Error(), capability) + require.NotContains(t, unavailableErr.Error(), "private-nat-stream") + profile.ServerID = 7 + singleton.NATShared.Update(profile) + waitContext, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + waited, err := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](waitContext, server.URL+agentcompatCapabilityWaitPath, plaintext, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "nat", ServerID: 7, ResourceID: 41}) + require.NoError(t, err) + require.Equal(t, "private-nat-stream", waited.StreamID) +} + +func TestAgentcompatCapabilityOwnerIncludesCurrentAdminRole(t *testing.T) { + cleanup, _ := setupMCPTest(t) + t.Cleanup(cleanup) + handler := rpc.NewNezhaHandler() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = handler + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + admin := &model.User{Common: model.Common{ID: 300}, Username: "cap-admin", Role: model.RoleAdmin} + require.NoError(t, singleton.DB.Create(admin).Error) + target := &model.Server{Common: model.Common{ID: 8}} + target.SetUserID(999) + singleton.ServerShared.InsertForTest(target) + token, plaintext := mkDistinctCapabilityToken(t, admin.ID, "admin") + token.SetServerIDs([]uint64{8}) + require.NoError(t, singleton.DB.Model(token).Update("servers_csv", token.ServersCSV).Error) + server := newAgentcompatCapabilityServer(t) + capability := registerAgentcompatCapability(t, server.URL, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "terminal", ServerID: 8}) + parsed, err := rpc.ParseAgentCompatIOStreamCapability(capability) + require.NoError(t, err) + require.NoError(t, handler.CreateStreamWithPurpose("admin-private-stream", admin.ID, 8, rpc.PurposeTerminal)) + require.NoError(t, handler.BindAgentCompatIOStreamCapability(rpc.AgentCompatCapabilityBinding{ + AgentCompatCapabilityAccess: rpc.AgentCompatCapabilityAccess{ + Capability: parsed, Owner: rpc.AgentCompatCapabilityOwner{PATID: token.ID, UserID: admin.ID, IsAdmin: true}, + Purpose: rpc.AgentCompatCapabilityTerminal, TargetServerID: 8, ServerAccessAllowed: true, + }, + StreamID: "admin-private-stream", + })) + waitContext, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + waited, err := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](waitContext, server.URL+agentcompatCapabilityWaitPath, plaintext, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "terminal", ServerID: 8}) + require.NoError(t, err) + require.Equal(t, "admin-private-stream", waited.StreamID) +} diff --git a/cmd/dashboard/controller/agentcompat_capability_routes_test.go b/cmd/dashboard/controller/agentcompat_capability_routes_test.go new file mode 100644 index 00000000..eb6ed3b6 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_capability_routes_test.go @@ -0,0 +1,175 @@ +//go:build agentcompat + +package controller + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/rpc" +) + +func TestAgentcompatCapabilityRoutesRequirePATAndDoNotAcceptOwnerFields(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + router := gin.New() + registerAgentcompatRoutes(router) + server := httptest.NewServer(router) + t.Cleanup(server.Close) + + requestBody := `{"purpose":"terminal","server_id":7,"pat_id":999,"user_id":999,"is_admin":true}` + requestContext, cancelRequests := context.WithTimeout(context.Background(), 5*time.Second) + defer cancelRequests() + for _, authorization := range []string{"", "Bearer jwt-looking-value", "Bearer " + token} { + request, err := http.NewRequestWithContext(requestContext, http.MethodPost, server.URL+agentcompatCapabilityRegisterPath, strings.NewReader(requestBody)) + require.NoError(t, err) + request.Header.Set("Content-Type", "application/json") + if authorization != "" { + request.Header.Set("Authorization", authorization) + } + response, err := server.Client().Do(request) + require.NoError(t, err) + body, readErr := io.ReadAll(response.Body) + response.Body.Close() + require.NoError(t, readErr) + if authorization == "Bearer "+token { + require.Equal(t, http.StatusOK, response.StatusCode) + var envelope model.CommonResponse[agentcompatCapabilityRegisterResponse] + require.NoError(t, json.Unmarshal(body, &envelope)) + require.False(t, envelope.Success) + require.NotContains(t, string(body), "999") + continue + } + require.Equal(t, http.StatusUnauthorized, response.StatusCode) + } +} + +func TestAgentcompatCapabilityRoutesRegisterWaitCancelUnregisterTypedLifecycle(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + tok, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + router := gin.New() + registerAgentcompatRoutes(router) + server := httptest.NewServer(router) + t.Cleanup(server.Close) + + capability := registerAgentcompatCapability(t, server.URL, token, agentcompatCapabilityRegisterRequest{Purpose: "terminal", ServerID: 7}) + require.NoError(t, rpc.NezhaHandlerSingleton.CreateStreamWithPurpose("private-stream", userID, 7, rpc.PurposeTerminal)) + parsed, err := rpc.ParseAgentCompatIOStreamCapability(capability) + require.NoError(t, err) + require.NoError(t, rpc.NezhaHandlerSingleton.BindAgentCompatIOStreamCapability(rpc.AgentCompatCapabilityBinding{ + AgentCompatCapabilityAccess: rpc.AgentCompatCapabilityAccess{ + Capability: parsed, + Owner: rpc.AgentCompatCapabilityOwner{PATID: tok.ID, UserID: userID}, + Purpose: rpc.AgentCompatCapabilityTerminal, TargetServerID: 7, + ServerAccessAllowed: true, + }, + StreamID: "private-stream", + })) + waitContext, cancelWait := context.WithTimeout(context.Background(), time.Second) + defer cancelWait() + result, err := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](waitContext, server.URL+agentcompatCapabilityWaitPath, token, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "terminal", ServerID: 7}) + require.NoError(t, err) + require.Equal(t, "private-stream", result.StreamID) + + status, body := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityUnregisterPath, token, agentcompatCapabilityUnregisterRequest{Capability: capability, Purpose: "terminal", ServerID: 7}) + require.Equal(t, http.StatusOK, status) + require.NotContains(t, body, capability) + + status, body = postAgentcompatEmpty(t, server.URL+agentcompatCapabilityCancelPath, token, agentcompatCapabilityCancelRequest{Capability: capability, Purpose: "terminal", ServerID: 7}) + require.Equal(t, http.StatusOK, status) + require.NotContains(t, body, capability) +} + +func TestAgentcompatCapabilityRoutesWaitCancellationDoesNotLeakSecrets(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, token := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + router := gin.New() + registerAgentcompatRoutes(router) + server := httptest.NewServer(router) + t.Cleanup(server.Close) + + request := agentcompatCapabilityWaitRequest{Capability: strings.Repeat("a", 43), Purpose: "terminal", ServerID: 7} + status, body := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityWaitPath, token, request) + require.Equal(t, http.StatusOK, status) + require.NotContains(t, body, request.Capability) + require.NotContains(t, body, "private-stream") +} + +func registerAgentcompatCapability(t *testing.T, baseURL, token string, request agentcompatCapabilityRegisterRequest) string { + t.Helper() + requestContext, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + response, err := postAgentcompatCapability[agentcompatCapabilityRegisterRequest, agentcompatCapabilityRegisterResponse](requestContext, baseURL+agentcompatCapabilityRegisterPath, token, request) + require.NoError(t, err) + require.NotEmpty(t, response.Capability) + return response.Capability +} + +func postAgentcompatEmpty(t *testing.T, path, token string, body any) (int, string) { + t.Helper() + encoded, err := json.Marshal(body) + require.NoError(t, err) + requestContext, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + request, err := http.NewRequestWithContext(requestContext, http.MethodPost, path, bytes.NewReader(encoded)) + require.NoError(t, err) + request.Header.Set("Authorization", "Bearer "+token) + request.Header.Set("Content-Type", "application/json") + response, err := http.DefaultClient.Do(request) + require.NoError(t, err) + defer response.Body.Close() + bodyBytes, err := io.ReadAll(response.Body) + require.NoError(t, err) + return response.StatusCode, string(bodyBytes) +} + +func postAgentcompatCapability[Request, Response any](ctx context.Context, path, token string, body Request) (Response, error) { + var zero Response + encoded, err := json.Marshal(body) + if err != nil { + return zero, err + } + request, err := http.NewRequestWithContext(ctx, http.MethodPost, path, bytes.NewReader(encoded)) + if err != nil { + return zero, err + } + request.Header.Set("Authorization", "Bearer "+token) + request.Header.Set("Content-Type", "application/json") + response, err := http.DefaultClient.Do(request) + if err != nil { + return zero, err + } + defer response.Body.Close() + var envelope model.CommonResponse[Response] + if err := json.NewDecoder(response.Body).Decode(&envelope); err != nil { + return zero, err + } + if !envelope.Success { + return zero, errors.New(envelope.Error) + } + return envelope.Data, nil +} diff --git a/cmd/dashboard/controller/agentcompat_capability_security_test.go b/cmd/dashboard/controller/agentcompat_capability_security_test.go new file mode 100644 index 00000000..cdddde99 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_capability_security_test.go @@ -0,0 +1,238 @@ +//go:build agentcompat + +package controller + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/rpc" + "github.com/nezhahq/nezha/service/singleton" +) + +func TestAgentcompatCapabilityCancelAndUnregisterAreUniformAndNonMutating(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + handler := rpc.NewNezhaHandler() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = handler + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + ownerToken, ownerPAT := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + _, foreignPAT := mkDistinctCapabilityToken(t, userID, "foreign") + server := newAgentcompatCapabilityServer(t) + capability := registerAgentcompatCapability(t, server.URL, ownerPAT, agentcompatCapabilityRegisterRequest{Purpose: agentcompatCapabilityPurposeTerminal, ServerID: 7}) + bindAgentcompatTerminalCapability(t, handler, capability, ownerToken.ID, userID, 7, "private-stream") + start := handler.SnapshotIOStreamState() + unknown := base64.RawURLEncoding.EncodeToString([]byte(strings.Repeat("u", 32))) + + cases := []struct { + name string + path string + pat string + body string + }{ + {name: "malformed cancel", path: agentcompatCapabilityCancelPath, pat: ownerPAT, body: `{"capability":`}, + {name: "oversize cancel", path: agentcompatCapabilityCancelPath, pat: ownerPAT, body: strings.Repeat("x", 513)}, + {name: "duplicate cancel", path: agentcompatCapabilityCancelPath, pat: ownerPAT, body: `{"capability":"x","capability":"y","purpose":"terminal","server_id":7}`}, + {name: "unknown cancel", path: agentcompatCapabilityCancelPath, pat: ownerPAT, body: capabilityAccessJSON(unknown, "terminal", 7, 0)}, + {name: "foreign cancel", path: agentcompatCapabilityCancelPath, pat: foreignPAT, body: capabilityAccessJSON(capability, "terminal", 7, 0)}, + {name: "purpose cancel", path: agentcompatCapabilityCancelPath, pat: ownerPAT, body: capabilityAccessJSON(capability, "file_manager", 7, 0)}, + {name: "server cancel", path: agentcompatCapabilityCancelPath, pat: ownerPAT, body: capabilityAccessJSON(capability, "terminal", 999999, 0)}, + {name: "invalid unregister", path: agentcompatCapabilityUnregisterPath, pat: ownerPAT, body: capabilityAccessJSON("invalid", "terminal", 7, 0)}, + {name: "oversize unregister", path: agentcompatCapabilityUnregisterPath, pat: ownerPAT, body: strings.Repeat("x", 513)}, + {name: "duplicate unregister", path: agentcompatCapabilityUnregisterPath, pat: ownerPAT, body: `{"capability":"x","purpose":"terminal","server_id":7,"server_id":8}`}, + {name: "foreign unregister", path: agentcompatCapabilityUnregisterPath, pat: foreignPAT, body: capabilityAccessJSON(capability, "terminal", 7, 0)}, + {name: "resource unregister", path: agentcompatCapabilityUnregisterPath, pat: ownerPAT, body: capabilityAccessJSON(capability, "terminal", 7, 41)}, + } + var uniformBody string + for _, testCase := range cases { + t.Run(testCase.name, func(t *testing.T) { + status, body := postAgentcompatRaw(t, server.URL+testCase.path, testCase.pat, testCase.body) + require.Equal(t, http.StatusOK, status) + if uniformBody == "" { + uniformBody = body + } + require.Equal(t, uniformBody, body) + require.Equal(t, start, handler.SnapshotIOStreamState()) + require.NotContains(t, body, capability) + require.NotContains(t, body, "private-stream") + }) + } + + waitContext, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + waited, err := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](waitContext, server.URL+agentcompatCapabilityWaitPath, ownerPAT, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "terminal", ServerID: 7}) + require.NoError(t, err) + require.Equal(t, "private-stream", waited.StreamID) +} + +func TestAgentcompatCapabilityPermissionRevocationIsRecheckedWithoutMutation(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + handler := rpc.NewNezhaHandler() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = handler + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + token, plaintext := mkToken(t, userID, []string{model.ScopeServerRead}, []uint64{7}) + server := newAgentcompatCapabilityServer(t) + capability := registerAgentcompatCapability(t, server.URL, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "terminal", ServerID: 7}) + bindAgentcompatTerminalCapability(t, handler, capability, token.ID, userID, 7, "revoked-private-stream") + start := handler.SnapshotIOStreamState() + token.SetServerIDs([]uint64{99}) + require.NoError(t, singleton.DB.Model(token).Update("servers_csv", token.ServersCSV).Error) + + cancelStatus, cancelBody := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityCancelPath, plaintext, agentcompatCapabilityCancelRequest{Capability: capability, Purpose: "terminal", ServerID: 7}) + unregisterStatus, unregisterBody := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityUnregisterPath, plaintext, agentcompatCapabilityUnregisterRequest{Capability: capability, Purpose: "terminal", ServerID: 7}) + require.Equal(t, http.StatusOK, cancelStatus) + require.Equal(t, http.StatusOK, unregisterStatus) + require.Equal(t, cancelBody, unregisterBody) + require.Equal(t, start, handler.SnapshotIOStreamState()) + + token.SetServerIDs(nil) + require.NoError(t, singleton.DB.Model(token).Update("servers_csv", "").Error) + waitContext, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + waited, err := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](waitContext, server.URL+agentcompatCapabilityWaitPath, plaintext, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "terminal", ServerID: 7}) + require.NoError(t, err) + require.Equal(t, "revoked-private-stream", waited.StreamID) +} + +func TestAgentcompatCapabilityForeignWaitAndRevokedPATCannotObserveBinding(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + handler := rpc.NewNezhaHandler() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = handler + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + ownerToken, ownerPAT := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + _, foreignPAT := mkDistinctCapabilityToken(t, userID, "foreign-wait") + server := newAgentcompatCapabilityServer(t) + capability := registerAgentcompatCapability(t, server.URL, ownerPAT, agentcompatCapabilityRegisterRequest{Purpose: "terminal", ServerID: 7}) + bindAgentcompatTerminalCapability(t, handler, capability, ownerToken.ID, userID, 7, "foreign-private-stream") + start := handler.SnapshotIOStreamState() + + foreignContext, cancelForeign := context.WithTimeout(context.Background(), time.Second) + defer cancelForeign() + _, foreignErr := postAgentcompatCapability[agentcompatCapabilityWaitRequest, agentcompatCapabilityWaitResponse](foreignContext, server.URL+agentcompatCapabilityWaitPath, foreignPAT, agentcompatCapabilityWaitRequest{Capability: capability, Purpose: "terminal", ServerID: 7}) + require.Error(t, foreignErr) + require.NotContains(t, foreignErr.Error(), capability) + require.NotContains(t, foreignErr.Error(), "foreign-private-stream") + require.Equal(t, start, handler.SnapshotIOStreamState()) + + require.NoError(t, singleton.DB.Delete(&model.APIToken{}, ownerToken.ID).Error) + status, body := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityWaitPath, ownerPAT, agentcompatCapabilityAccessRequest{Capability: capability, Purpose: "terminal", ServerID: 7}) + require.Equal(t, http.StatusUnauthorized, status) + require.NotContains(t, body, capability) + require.NotContains(t, body, "foreign-private-stream") + require.Equal(t, start, handler.SnapshotIOStreamState()) +} + +func TestAgentcompatCapabilityRegisterRequiresCurrentServerWhitelist(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, plaintext := mkToken(t, userID, []string{model.ScopeServerRead}, []uint64{99}) + server := newAgentcompatCapabilityServer(t) + + status, body := postAgentcompatEmpty(t, server.URL+agentcompatCapabilityRegisterPath, plaintext, agentcompatCapabilityRegisterRequest{Purpose: "file_manager", ServerID: 7}) + + require.Equal(t, http.StatusOK, status) + require.Contains(t, body, errAgentcompatCapabilityUnavailable.Error()) + require.NotContains(t, body, "99") + require.NotContains(t, body, plaintext) +} + +func TestAgentcompatCapabilityRegisterRejectsInvalidIdentityWithoutDisclosure(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, plaintext := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + server := newAgentcompatCapabilityServer(t) + privateServer := "777777777" + cases := []string{ + `{`, + `{"purpose":"unknown","server_id":7}`, + `{"purpose":"terminal","server_id":0}`, + `{"purpose":"terminal","server_id":7,"resource_id":41}`, + `{"purpose":"nat","server_id":7}`, + `{"purpose":"terminal","server_id":` + privateServer + `}`, + } + for _, body := range cases { + status, responseBody := postAgentcompatRaw(t, server.URL+agentcompatCapabilityRegisterPath, plaintext, body) + require.Equal(t, http.StatusOK, status) + require.NotContains(t, responseBody, privateServer) + require.NotContains(t, responseBody, plaintext) + var envelope model.CommonResponse[agentcompatCapabilityRegisterResponse] + require.NoError(t, json.Unmarshal([]byte(responseBody), &envelope)) + require.False(t, envelope.Success) + require.Empty(t, envelope.Data.Capability) + } +} + +func newAgentcompatCapabilityServer(t *testing.T) *httptest.Server { + t.Helper() + gin.SetMode(gin.TestMode) + router := gin.New() + registerAgentcompatRoutes(router) + server := httptest.NewServer(router) + t.Cleanup(server.Close) + return server +} + +func mkDistinctCapabilityToken(t *testing.T, userID uint64, suffix string) (*model.APIToken, string) { + t.Helper() + plaintext := "nzp_" + strings.Repeat("z", 32) + "_" + suffix + token := &model.APIToken{UserID: userID, Name: suffix, TokenHash: model.HashAPIToken(plaintext)} + token.SetScopes([]string{model.ScopeServerRead}) + require.NoError(t, singleton.DB.Create(token).Error) + return token, plaintext +} + +func bindAgentcompatTerminalCapability(t *testing.T, handler *rpc.NezhaHandler, rawCapability string, tokenID, userID, serverID uint64, streamID string) { + t.Helper() + capability, err := rpc.ParseAgentCompatIOStreamCapability(rawCapability) + require.NoError(t, err) + access := rpc.AgentCompatCapabilityAccess{ + Capability: capability, Owner: rpc.AgentCompatCapabilityOwner{PATID: tokenID, UserID: userID}, + Purpose: rpc.AgentCompatCapabilityTerminal, TargetServerID: serverID, ServerAccessAllowed: true, + } + require.NoError(t, handler.CreateStreamWithPurpose(streamID, userID, serverID, rpc.PurposeTerminal)) + require.NoError(t, handler.BindAgentCompatIOStreamCapability(rpc.AgentCompatCapabilityBinding{AgentCompatCapabilityAccess: access, StreamID: streamID})) +} + +func capabilityAccessJSON(capability, purpose string, serverID, resourceID uint64) string { + body, _ := json.Marshal(agentcompatCapabilityAccessRequest{Capability: capability, Purpose: purpose, ServerID: serverID, ResourceID: resourceID}) + return string(body) +} + +func postAgentcompatRaw(t *testing.T, path, token, body string) (int, string) { + t.Helper() + requestContext, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + request, err := http.NewRequestWithContext(requestContext, http.MethodPost, path, bytes.NewBufferString(body)) + require.NoError(t, err) + request.Header.Set("Authorization", "Bearer "+token) + request.Header.Set("Content-Type", "application/json") + response, err := http.DefaultClient.Do(request) + require.NoError(t, err) + defer response.Body.Close() + responseBody, err := io.ReadAll(response.Body) + require.NoError(t, err) + return response.StatusCode, string(responseBody) +} diff --git a/cmd/dashboard/controller/agentcompat_routes.go b/cmd/dashboard/controller/agentcompat_routes.go new file mode 100644 index 00000000..c262d4d2 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_routes.go @@ -0,0 +1,147 @@ +//go:build agentcompat + +package controller + +import ( + "encoding/json" + "errors" + "fmt" + "io" + "time" + + "github.com/gin-gonic/gin" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/rpc" +) + +const ( + agentcompatOversizeWriteSentinel = "agentcompat:oversize-write-contract" + agentcompatOversizeWriteServerPath = "/tmp/agentcompat-oversize-contract.txt" + agentcompatFsWriteOperationOversize = "oversize" +) + +type agentcompatFsWriteContractRequest struct { + ServerID uint64 `json:"server_id"` + Operation agentcompatFsWriteOperation `json:"operation"` +} + +type agentcompatFsWriteOperation string + +type agentcompatFsWriteContractResponse struct { + Result model.FsWriteResult `json:"result"` + AgentRPCResponse bool `json:"agent_rpc_response"` +} + +func registerAgentcompatRoutes(router *gin.Engine) { + patAuth := requiredAgentcompatPAT(apiTokenAuthMiddleware()) + router.POST("/agentcompat/fs-write-contract", patAuth, commonHandler(agentcompatFsWriteContract)) + router.POST("/agentcompat/mcp-rate-limit-probe", patAuth, commonHandler(agentcompatMCPRateLimitProbeRoute)) + router.GET("/agentcompat/io-stream-state", patAuth, commonHandler(agentcompatIOStreamSnapshot)) + router.POST("/agentcompat/io-stream-state", patAuth, commonHandler(agentcompatIOStreamWait)) + router.POST("/agentcompat/io-stream-quota-probe", patAuth, commonHandler(agentcompatIOStreamQuotaProbeRoute)) + registerAgentcompatCapabilityRoutes(router, patAuth) + registerAgentcompatSQLiteHoldRoutes(router, patAuth) +} + +type agentcompatIOStreamQuotaProbeResponse struct { + UserAccepted int `json:"user_accepted"` + UserRejected int `json:"user_rejected"` + ServerAccepted int `json:"server_accepted"` + ServerRejected int `json:"server_rejected"` + Clean bool `json:"clean"` +} + +func agentcompatIOStreamQuotaProbeRoute(context *gin.Context) (agentcompatIOStreamQuotaProbeResponse, error) { + if rpc.NezhaHandlerSingleton == nil { + return agentcompatIOStreamQuotaProbeResponse{}, errors.New("IOStream handler is unavailable") + } + result := rpc.RunIOStreamQuotaProbe(context.Request.Context()) + if result.Err != nil { + return agentcompatIOStreamQuotaProbeResponse{}, result.Err + } + return agentcompatIOStreamQuotaProbeResponse{UserAccepted: result.UserAccepted, UserRejected: result.UserRejected, ServerAccepted: result.ServerAccepted, ServerRejected: result.ServerRejected, Clean: result.TrackedStreams == 0}, nil +} + +func requiredAgentcompatPAT(auth gin.HandlerFunc) gin.HandlerFunc { + return func(context *gin.Context) { + auth(context) + if context.IsAborted() { + return + } + if APITokenFromContext(context) == nil { + abortAPITokenUnauthorized(context, "api token required") + } + } +} + +func agentcompatIOStreamSnapshot(*gin.Context) (rpc.IOStreamState, error) { + if rpc.NezhaHandlerSingleton == nil { + return rpc.IOStreamState{}, errors.New("IOStream handler is unavailable") + } + return rpc.NezhaHandlerSingleton.SnapshotIOStreamState(), nil +} + +func agentcompatIOStreamWait(context *gin.Context) (rpc.IOStreamState, error) { + if rpc.NezhaHandlerSingleton == nil { + return rpc.IOStreamState{}, errors.New("IOStream handler is unavailable") + } + var expectation rpc.IOStreamStateExpectation + if err := context.ShouldBindJSON(&expectation); err != nil { + return rpc.IOStreamState{}, err + } + return rpc.NezhaHandlerSingleton.WaitForIOStreamState(context.Request.Context(), expectation) +} + +func agentcompatFsWriteContract(context *gin.Context) (agentcompatFsWriteContractResponse, error) { + var request agentcompatFsWriteContractRequest + if err := decodeAgentcompatJSON(context, &request); err != nil { + return agentcompatFsWriteContractResponse{}, err + } + if request.ServerID == 0 || request.Operation == "" { + return agentcompatFsWriteContractResponse{}, errors.New("server_id and operation are required") + } + if request.Operation != agentcompatFsWriteOperation(agentcompatFsWriteOperationOversize) { + return agentcompatFsWriteContractResponse{}, errors.New("unknown filesystem write operation") + } + token := APITokenFromContext(context) + if token == nil || !token.HasScope(model.ScopeServerWrite) { + return agentcompatFsWriteContractResponse{}, errors.New("missing required scope: " + model.ScopeServerWrite) + } + server, err := requireServerAccess(context, request.ServerID) + if err != nil { + return agentcompatFsWriteContractResponse{}, err + } + content := "denied" + if request.Operation == agentcompatFsWriteOperation(agentcompatFsWriteOperationOversize) { + content = agentcompatOversizeWriteSentinel + } + // This tagged probe owns both path and payload; accepting either from callers would make it an arbitrary-write endpoint. + raw, err := rpc.CallAgent(context.Request.Context(), server.ID, model.TaskTypeFsWrite, model.FsWriteRequest{ + Path: agentcompatOversizeWriteServerPath, Content: content, Encoding: "utf8", Mode: "0600", + }, 30*time.Second) + if err != nil { + return agentcompatFsWriteContractResponse{}, err + } + var result model.FsWriteResult + if err := json.Unmarshal(raw, &result); err != nil { + return agentcompatFsWriteContractResponse{}, err + } + return agentcompatFsWriteContractResponse{Result: result, AgentRPCResponse: true}, nil +} + +func decodeAgentcompatJSON(context *gin.Context, value any) error { + decoder := json.NewDecoder(context.Request.Body) + decoder.DisallowUnknownFields() + if err := decoder.Decode(value); err != nil { + return err + } + var trailing any + if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) { + if err == nil { + return fmt.Errorf("trailing JSON values are not allowed") + } + return err + } + return nil +} diff --git a/cmd/dashboard/controller/agentcompat_routes_default.go b/cmd/dashboard/controller/agentcompat_routes_default.go new file mode 100644 index 00000000..9815f6e3 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_routes_default.go @@ -0,0 +1,7 @@ +//go:build !agentcompat + +package controller + +import "github.com/gin-gonic/gin" + +func registerAgentcompatRoutes(*gin.Engine) {} diff --git a/cmd/dashboard/controller/agentcompat_routes_security_test.go b/cmd/dashboard/controller/agentcompat_routes_security_test.go new file mode 100644 index 00000000..2dee0e78 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_routes_security_test.go @@ -0,0 +1,69 @@ +//go:build agentcompat + +package controller + +import ( + "bytes" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/model" +) + +func TestAgentcompatFsWriteContractRejectsCallerPayloadAndForeignServer(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + _, plaintext := mkToken(t, userID, []string{model.ScopeServerWrite}, []uint64{99}) + server := newAgentcompatCapabilityServer(t) + + status, body := postAgentcompatRaw(t, server.URL+"/agentcompat/fs-write-contract", plaintext, `{"server_id":7,"path":"/tmp/foreign","operation":"oversize","content":"attacker"}`) + + require.Equal(t, http.StatusOK, status) + require.NotContains(t, body, "attacker") + require.Contains(t, body, "unknown field") + + status, body = postAgentcompatRaw(t, server.URL+"/agentcompat/fs-write-contract", plaintext, `{"server_id":7,"operation":"oversize"}`) + require.Equal(t, http.StatusOK, status) + require.Contains(t, body, "permission denied") +} + +func TestAgentcompatFsWriteContractRequiresWriteScopeBeforeRPC(t *testing.T) { + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + _, plaintext := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + server := newAgentcompatCapabilityServer(t) + + status, body := postAgentcompatRaw(t, server.URL+"/agentcompat/fs-write-contract", plaintext, `{"server_id":7,"operation":"oversize"}`) + + require.Equal(t, http.StatusOK, status) + require.Contains(t, body, "missing required scope") + require.NotContains(t, body, `"agent_rpc_response":true`) +} + +func TestAgentcompatProbeJSONRejectsUnknownFieldsAndTrailingValues(t *testing.T) { + tests := []string{ + `{"server_id":7,"operation":"oversize","path":"caller"}`, + `{"server_id":7,"operation":"oversize","content":"caller"}`, + `{"server_id":7,"operation":"oversize"}{}`, + } + for _, body := range tests { + t.Run(body, func(t *testing.T) { + context := newAgentcompatJSONContext(t, body) + var request agentcompatFsWriteContractRequest + require.Error(t, decodeAgentcompatJSON(context, &request)) + }) + } +} + +func newAgentcompatJSONContext(t *testing.T, body string) *gin.Context { + t.Helper() + gin.SetMode(gin.TestMode) + context, _ := gin.CreateTestContext(httptest.NewRecorder()) + context.Request = httptest.NewRequest(http.MethodPost, "/", bytes.NewBufferString(body)) + context.Request.Header.Set("Content-Type", "application/json") + return context +} diff --git a/cmd/dashboard/controller/agentcompat_sqlite_hold_error_agentcompat_linux_test.go b/cmd/dashboard/controller/agentcompat_sqlite_hold_error_agentcompat_linux_test.go new file mode 100644 index 00000000..46953c06 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_sqlite_hold_error_agentcompat_linux_test.go @@ -0,0 +1,40 @@ +//go:build agentcompat && linux + +package controller + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/service/singleton" +) + +func TestAgentcompatSQLiteHoldErrorUsesFixedRedactedMessages(t *testing.T) { + tests := []struct { + cause error + want string + }{ + {agentcompatSQLiteHoldInvalidRequest{}, "agentcompat sqlite hold request is invalid"}, + {singleton.ErrSQLiteHoldSessionActive, "agentcompat sqlite hold is active"}, + {singleton.ErrSQLiteHoldStaleSession, "agentcompat sqlite hold receipt is stale"}, + {singleton.ErrSQLiteHoldFinalizationNotStarted, "agentcompat sqlite hold is not ready"}, + {singleton.ErrSQLiteHoldFinalizationStarted, "agentcompat sqlite hold is not ready"}, + {singleton.ErrSQLiteHoldUnexpectedSelection, "agentcompat sqlite hold was aborted"}, + {singleton.ErrSQLiteHoldAmbiguousCandidate, "agentcompat sqlite hold was aborted"}, + {singleton.ErrSQLiteHoldAborted, "agentcompat sqlite hold was aborted"}, + {context.Canceled, "agentcompat sqlite hold wait was canceled"}, + {context.DeadlineExceeded, "agentcompat sqlite hold wait was canceled"}, + {errors.New("path=/secret token=nzp_secret"), "agentcompat sqlite hold control is unavailable"}, + } + + for _, test := range tests { + // When + err := agentcompatSQLiteHoldError(test.cause) + + // Then + require.EqualError(t, err, test.want) + } +} diff --git a/cmd/dashboard/controller/agentcompat_sqlite_hold_routes_agentcompat_linux.go b/cmd/dashboard/controller/agentcompat_sqlite_hold_routes_agentcompat_linux.go new file mode 100644 index 00000000..f966eb22 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_sqlite_hold_routes_agentcompat_linux.go @@ -0,0 +1,168 @@ +//go:build agentcompat && linux + +package controller + +import ( + "context" + "encoding/base64" + "encoding/json" + "errors" + "io" + "net/http" + + "github.com/gin-gonic/gin" + + "github.com/nezhahq/nezha/service/singleton" +) + +const agentcompatSQLiteHoldPath = "/agentcompat/sqlite-hold/" + +type agentcompatSQLiteHoldControl interface { + ArmNextSQLiteHold() (singleton.SQLiteHoldReceipt, error) + WaitSQLiteHoldSelected(context.Context, singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) + WaitSQLiteHoldFinalizing(context.Context, singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) + SnapshotSQLiteHold(singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) + ReleaseSQLiteHold(singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) + AbortSQLiteHold(singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) +} + +type agentcompatSQLiteHoldFacade struct{} + +func (agentcompatSQLiteHoldFacade) ArmNextSQLiteHold() (singleton.SQLiteHoldReceipt, error) { + return singleton.ArmNextSQLiteHold() +} +func (agentcompatSQLiteHoldFacade) WaitSQLiteHoldSelected(ctx context.Context, receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) { + return singleton.WaitSQLiteHoldSelected(ctx, receipt) +} +func (agentcompatSQLiteHoldFacade) WaitSQLiteHoldFinalizing(ctx context.Context, receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) { + return singleton.WaitSQLiteHoldFinalizing(ctx, receipt) +} +func (agentcompatSQLiteHoldFacade) SnapshotSQLiteHold(receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) { + return singleton.SnapshotSQLiteHold(receipt) +} +func (agentcompatSQLiteHoldFacade) ReleaseSQLiteHold(receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) { + return singleton.ReleaseSQLiteHold(receipt) +} +func (agentcompatSQLiteHoldFacade) AbortSQLiteHold(receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) { + return singleton.AbortSQLiteHold(receipt) +} + +var agentcompatSQLiteHoldController agentcompatSQLiteHoldControl = agentcompatSQLiteHoldFacade{} + +type agentcompatSQLiteHoldRequest struct { + ID string `json:"id"` + State singleton.SQLiteHoldControlState `json:"state"` +} + +func registerAgentcompatSQLiteHoldRoutes(router *gin.Engine, patAuth gin.HandlerFunc) { + readOnlyPAT := func(c *gin.Context) { suppressAPITokenAuthWrites(c) } + router.POST(agentcompatSQLiteHoldPath+"arm", readOnlyPAT, patAuth, commonHandler(agentcompatSQLiteHoldArm)) + router.POST(agentcompatSQLiteHoldPath+"wait", readOnlyPAT, patAuth, commonHandler(agentcompatSQLiteHoldWait)) + router.POST(agentcompatSQLiteHoldPath+"snapshot", readOnlyPAT, patAuth, commonHandler(agentcompatSQLiteHoldSnapshot)) + router.POST(agentcompatSQLiteHoldPath+"release", readOnlyPAT, patAuth, commonHandler(agentcompatSQLiteHoldRelease)) + router.POST(agentcompatSQLiteHoldPath+"abort", readOnlyPAT, patAuth, commonHandler(agentcompatSQLiteHoldAbort)) +} + +func agentcompatSQLiteHoldArm(c *gin.Context) (singleton.SQLiteHoldReceipt, error) { + var request struct{} + if err := decodeAgentcompatSQLiteHoldRequest(c, &request); err != nil { + return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldError(err) + } + receipt, err := agentcompatSQLiteHoldController.ArmNextSQLiteHold() + return receipt, agentcompatSQLiteHoldError(err) +} +func agentcompatSQLiteHoldWait(c *gin.Context) (singleton.SQLiteHoldReceipt, error) { + receipt, err := decodeAgentcompatSQLiteHoldReceipt(c, true) + if err != nil { + return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldError(err) + } + switch receipt.State { + case singleton.SQLiteHoldControlStateSelected: + receipt, err = agentcompatSQLiteHoldController.WaitSQLiteHoldSelected(c.Request.Context(), receipt) + case singleton.SQLiteHoldControlStateFinalizing: + receipt, err = agentcompatSQLiteHoldController.WaitSQLiteHoldFinalizing(c.Request.Context(), receipt) + default: + return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldInvalidRequest{} + } + return receipt, agentcompatSQLiteHoldError(err) +} +func agentcompatSQLiteHoldSnapshot(c *gin.Context) (singleton.SQLiteHoldReceipt, error) { + receipt, err := decodeAgentcompatSQLiteHoldReceipt(c, false) + if err != nil { + return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldError(err) + } + result, err := agentcompatSQLiteHoldController.SnapshotSQLiteHold(receipt) + return result, agentcompatSQLiteHoldError(err) +} +func agentcompatSQLiteHoldRelease(c *gin.Context) (singleton.SQLiteHoldReceipt, error) { + receipt, err := decodeAgentcompatSQLiteHoldReceipt(c, false) + if err != nil { + return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldError(err) + } + result, err := agentcompatSQLiteHoldController.ReleaseSQLiteHold(receipt) + return result, agentcompatSQLiteHoldError(err) +} +func agentcompatSQLiteHoldAbort(c *gin.Context) (singleton.SQLiteHoldReceipt, error) { + receipt, err := decodeAgentcompatSQLiteHoldReceipt(c, false) + if err != nil { + return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldError(err) + } + result, err := agentcompatSQLiteHoldController.AbortSQLiteHold(receipt) + return result, agentcompatSQLiteHoldError(err) +} + +type agentcompatSQLiteHoldInvalidRequest struct{} + +func (agentcompatSQLiteHoldInvalidRequest) Error() string { + return "agentcompat sqlite hold request is invalid" +} + +func decodeAgentcompatSQLiteHoldRequest(c *gin.Context, value any) error { + c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, 512) + decoder := json.NewDecoder(c.Request.Body) + decoder.DisallowUnknownFields() + if err := decoder.Decode(value); err != nil { + return agentcompatSQLiteHoldInvalidRequest{} + } + var trailing any + if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) { + return agentcompatSQLiteHoldInvalidRequest{} + } + return nil +} +func decodeAgentcompatSQLiteHoldReceipt(c *gin.Context, requireState bool) (singleton.SQLiteHoldReceipt, error) { + var request agentcompatSQLiteHoldRequest + if err := decodeAgentcompatSQLiteHoldRequest(c, &request); err != nil { + return singleton.SQLiteHoldReceipt{}, err + } + decoded, err := base64.RawURLEncoding.DecodeString(request.ID) + if err != nil || len(request.ID) != 43 || len(decoded) != 32 { + return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldInvalidRequest{} + } + if !requireState && request.State != "" { + return singleton.SQLiteHoldReceipt{}, agentcompatSQLiteHoldInvalidRequest{} + } + return singleton.SQLiteHoldReceipt{ID: request.ID, State: request.State}, nil +} +func agentcompatSQLiteHoldError(err error) error { + if err == nil { + return nil + } + if errors.As(err, new(agentcompatSQLiteHoldInvalidRequest)) { + return agentcompatSQLiteHoldInvalidRequest{} + } + switch { + case errors.Is(err, singleton.ErrSQLiteHoldSessionActive): + return errors.New("agentcompat sqlite hold is active") + case errors.Is(err, singleton.ErrSQLiteHoldStaleSession): + return errors.New("agentcompat sqlite hold receipt is stale") + case errors.Is(err, singleton.ErrSQLiteHoldFinalizationNotStarted), errors.Is(err, singleton.ErrSQLiteHoldFinalizationStarted): + return errors.New("agentcompat sqlite hold is not ready") + case errors.Is(err, singleton.ErrSQLiteHoldUnexpectedSelection), errors.Is(err, singleton.ErrSQLiteHoldAmbiguousCandidate), errors.Is(err, singleton.ErrSQLiteHoldAborted): + return errors.New("agentcompat sqlite hold was aborted") + case errors.Is(err, context.Canceled), errors.Is(err, context.DeadlineExceeded): + return errors.New("agentcompat sqlite hold wait was canceled") + default: + return errors.New("agentcompat sqlite hold control is unavailable") + } +} diff --git a/cmd/dashboard/controller/agentcompat_sqlite_hold_routes_agentcompat_linux_test.go b/cmd/dashboard/controller/agentcompat_sqlite_hold_routes_agentcompat_linux_test.go new file mode 100644 index 00000000..bf03787c --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_sqlite_hold_routes_agentcompat_linux_test.go @@ -0,0 +1,201 @@ +//go:build agentcompat && linux + +package controller + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/singleton" +) + +type agentcompatSQLiteHoldControlProbe struct { + err error + armCalls int + lastCall string + waitContextError error +} + +func (probe *agentcompatSQLiteHoldControlProbe) ArmNextSQLiteHold() (singleton.SQLiteHoldReceipt, error) { + probe.armCalls++ + probe.lastCall = "arm" + return agentcompatSQLiteHoldTestReceipt(singleton.SQLiteHoldControlStateArmed), probe.err +} + +func (probe *agentcompatSQLiteHoldControlProbe) WaitSQLiteHoldSelected(ctx context.Context, receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) { + probe.lastCall = "wait-selected" + probe.waitContextError = ctx.Err() + receipt.State = singleton.SQLiteHoldControlStateSelected + return receipt, probe.err +} + +func (probe *agentcompatSQLiteHoldControlProbe) WaitSQLiteHoldFinalizing(ctx context.Context, receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) { + probe.lastCall = "wait-finalizing" + probe.waitContextError = ctx.Err() + receipt.State = singleton.SQLiteHoldControlStateFinalizing + return receipt, probe.err +} + +func (probe *agentcompatSQLiteHoldControlProbe) SnapshotSQLiteHold(receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) { + probe.lastCall = "snapshot" + receipt.State = singleton.SQLiteHoldControlStateSelected + return receipt, probe.err +} + +func (probe *agentcompatSQLiteHoldControlProbe) ReleaseSQLiteHold(receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) { + probe.lastCall = "release" + receipt.State = singleton.SQLiteHoldControlStateReleased + return receipt, probe.err +} + +func (probe *agentcompatSQLiteHoldControlProbe) AbortSQLiteHold(receipt singleton.SQLiteHoldReceipt) (singleton.SQLiteHoldReceipt, error) { + probe.lastCall = "abort" + receipt.State = singleton.SQLiteHoldControlStateAborted + return receipt, probe.err +} + +func TestAgentcompatSQLiteHoldRoutesRequirePATWithoutScopeOrAuthWrites(t *testing.T) { + // Given + router, token, storedToken, probe := setupAgentcompatSQLiteHoldRouteTest(t) + + // When + status, response := requestAgentcompatSQLiteHold(t, router, context.Background(), "arm", token, `{}`) + + // Then + require.Equal(t, http.StatusOK, status) + require.True(t, response.Success) + require.Equal(t, agentcompatSQLiteHoldTestReceipt(singleton.SQLiteHoldControlStateArmed), response.Data) + require.Equal(t, 1, probe.armCalls) + var refreshed model.APIToken + require.NoError(t, singleton.DB.First(&refreshed, storedToken.ID).Error) + require.Nil(t, refreshed.LastUsedAt) + require.Empty(t, refreshed.LastUsedIP) + + for _, authorization := range []string{"", "jwt-looking-value", "nzp_invalid"} { + status, _ = requestAgentcompatSQLiteHold(t, router, context.Background(), "arm", authorization, `{}`) + require.Equal(t, http.StatusUnauthorized, status) + } + var wafRows int64 + require.NoError(t, singleton.DB.Model(&model.WAF{}).Count(&wafRows).Error) + require.Zero(t, wafRows) +} + +func TestAgentcompatSQLiteHoldRoutesRejectInvalidBodiesBeforeControl(t *testing.T) { + // Given + router, token, _, probe := setupAgentcompatSQLiteHoldRouteTest(t) + tests := []string{"", `{`, `{"unknown":true}`, `{} {}`, strings.Repeat(" ", 513) + `{}`} + + for _, body := range tests { + // When + status, response := requestAgentcompatSQLiteHold(t, router, context.Background(), "arm", token, body) + + // Then + require.Equal(t, http.StatusOK, status) + require.False(t, response.Success) + require.Equal(t, "agentcompat sqlite hold request is invalid", response.Error) + } + require.Zero(t, probe.armCalls) +} + +func TestAgentcompatSQLiteHoldRoutesDispatchTypedLifecycleAndPropagateContext(t *testing.T) { + // Given + router, token, _, probe := setupAgentcompatSQLiteHoldRouteTest(t) + receiptID := agentcompatSQLiteHoldTestReceipt("").ID + tests := []struct { + path string + body string + wantCall string + wantState singleton.SQLiteHoldControlState + }{ + {"wait", `{"id":"` + receiptID + `","state":"selected"}`, "wait-selected", singleton.SQLiteHoldControlStateSelected}, + {"wait", `{"id":"` + receiptID + `","state":"finalizing"}`, "wait-finalizing", singleton.SQLiteHoldControlStateFinalizing}, + {"snapshot", `{"id":"` + receiptID + `"}`, "snapshot", singleton.SQLiteHoldControlStateSelected}, + {"release", `{"id":"` + receiptID + `"}`, "release", singleton.SQLiteHoldControlStateReleased}, + {"abort", `{"id":"` + receiptID + `"}`, "abort", singleton.SQLiteHoldControlStateAborted}, + } + + for _, test := range tests { + // When + status, response := requestAgentcompatSQLiteHold(t, router, context.Background(), test.path, token, test.body) + + // Then + require.Equal(t, http.StatusOK, status) + require.True(t, response.Success) + require.Equal(t, test.wantCall, probe.lastCall) + require.Equal(t, test.wantState, response.Data.State) + } + + canceledContext, cancel := context.WithCancel(context.Background()) + cancel() + status, response := requestAgentcompatSQLiteHold(t, router, canceledContext, "wait", token, `{"id":"`+receiptID+`","state":"selected"}`) + require.Equal(t, http.StatusOK, status) + require.True(t, response.Success) + require.ErrorIs(t, probe.waitContextError, context.Canceled) + + status, response = requestAgentcompatSQLiteHold(t, router, context.Background(), "wait", token, `{"id":"`+receiptID+`","state":"armed"}`) + require.Equal(t, http.StatusOK, status) + require.False(t, response.Success) + require.Equal(t, "agentcompat sqlite hold request is invalid", response.Error) +} + +func TestAgentcompatSQLiteHoldRoutesRedactControlErrors(t *testing.T) { + // Given + router, token, _, probe := setupAgentcompatSQLiteHoldRouteTest(t) + probe.err = errors.New("token=nzp_secret path=/root/private dashboard.sqlite-journal") + + // When + status, response := requestAgentcompatSQLiteHold(t, router, context.Background(), "arm", token, `{}`) + + // Then + require.Equal(t, http.StatusOK, status) + require.False(t, response.Success) + require.Equal(t, "agentcompat sqlite hold control is unavailable", response.Error) + encoded, err := json.Marshal(response) + require.NoError(t, err) + require.NotContains(t, string(encoded), "nzp_secret") + require.NotContains(t, string(encoded), "/root/private") + require.NotContains(t, string(encoded), "journal") +} + +func setupAgentcompatSQLiteHoldRouteTest(t *testing.T) (*gin.Engine, string, *model.APIToken, *agentcompatSQLiteHoldControlProbe) { + t.Helper() + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + storedToken, token := mkToken(t, userID, nil, nil) + probe := &agentcompatSQLiteHoldControlProbe{} + originalControl := agentcompatSQLiteHoldController + agentcompatSQLiteHoldController = probe + t.Cleanup(func() { agentcompatSQLiteHoldController = originalControl }) + router := gin.New() + registerAgentcompatSQLiteHoldRoutes(router, requiredAgentcompatPAT(apiTokenAuthMiddleware())) + return router, token, storedToken, probe +} + +func requestAgentcompatSQLiteHold(t *testing.T, router *gin.Engine, ctx context.Context, path, authorization, body string) (int, model.CommonResponse[singleton.SQLiteHoldReceipt]) { + t.Helper() + request := httptest.NewRequestWithContext(ctx, http.MethodPost, agentcompatSQLiteHoldPath+path, bytes.NewBufferString(body)) + request.Header.Set("Content-Type", "application/json") + if authorization != "" { + request.Header.Set("Authorization", "Bearer "+authorization) + } + responseRecorder := httptest.NewRecorder() + router.ServeHTTP(responseRecorder, request) + var response model.CommonResponse[singleton.SQLiteHoldReceipt] + require.NoError(t, json.Unmarshal(responseRecorder.Body.Bytes(), &response)) + return responseRecorder.Code, response +} + +func agentcompatSQLiteHoldTestReceipt(state singleton.SQLiteHoldControlState) singleton.SQLiteHoldReceipt { + return singleton.SQLiteHoldReceipt{ID: base64.RawURLEncoding.EncodeToString(bytes.Repeat([]byte{17}, 32)), State: state} +} diff --git a/cmd/dashboard/controller/agentcompat_sqlite_hold_routes_agentcompat_nonlinux.go b/cmd/dashboard/controller/agentcompat_sqlite_hold_routes_agentcompat_nonlinux.go new file mode 100644 index 00000000..07112ad9 --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_sqlite_hold_routes_agentcompat_nonlinux.go @@ -0,0 +1,7 @@ +//go:build agentcompat && !linux + +package controller + +import "github.com/gin-gonic/gin" + +func registerAgentcompatSQLiteHoldRoutes(*gin.Engine, gin.HandlerFunc) {} diff --git a/cmd/dashboard/controller/agentcompat_sqlite_hold_routes_default_test.go b/cmd/dashboard/controller/agentcompat_sqlite_hold_routes_default_test.go new file mode 100644 index 00000000..f8e9760e --- /dev/null +++ b/cmd/dashboard/controller/agentcompat_sqlite_hold_routes_default_test.go @@ -0,0 +1,27 @@ +//go:build !agentcompat + +package controller + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestAgentcompatSQLiteHoldRoutesAreAbsentWithoutBuildTag(t *testing.T) { + // Given + gin.SetMode(gin.TestMode) + router := gin.New() + registerAgentcompatRoutes(router) + request := httptest.NewRequest(http.MethodPost, "/agentcompat/sqlite-hold/arm", nil) + response := httptest.NewRecorder() + + // When + router.ServeHTTP(response, request) + + // Then + require.Equal(t, http.StatusNotFound, response.Code) +} diff --git a/cmd/dashboard/controller/controller.go b/cmd/dashboard/controller/controller.go index 5ce81998..1823d2c4 100644 --- a/cmd/dashboard/controller/controller.go +++ b/cmd/dashboard/controller/controller.go @@ -69,6 +69,7 @@ func routers(r *gin.Engine, frontendDist fs.FS) { r.DELETE("/mcp", mcpOriginGuard(), mcpMethodNotAllowed) r.GET("/mcp/download/:token", mcpOriginGuard(), transferDownloadHandler) r.POST("/mcp/upload/:token", mcpOriginGuard(), transferUploadHandler) + registerAgentcompatRoutes(r) api := r.Group("api/v1") api.POST("/login", authMiddleware.LoginHandler) diff --git a/cmd/dashboard/controller/fm.go b/cmd/dashboard/controller/fm.go index 0e11f834..dfdb2728 100644 --- a/cmd/dashboard/controller/fm.go +++ b/cmd/dashboard/controller/fm.go @@ -6,7 +6,6 @@ import ( "github.com/gin-gonic/gin" "github.com/goccy/go-json" - "github.com/gorilla/websocket" "github.com/hashicorp/go-uuid" "github.com/nezhahq/nezha/model" @@ -26,6 +25,7 @@ import ( // @Success 200 {object} model.CreateFMResponse // @Router /file [post] func createFM(c *gin.Context) (*model.CreateFMResponse, error) { + prepareAgentcompatCapabilityHeader(c) idStr := c.Query("id") id, err := strconv.ParseUint(idStr, 10, 64) if err != nil { @@ -49,17 +49,24 @@ func createFM(c *gin.Context) (*model.CreateFMResponse, error) { return nil, err } - if err := rpc.NezhaHandlerSingleton.CreateStream(streamId, getUid(c), server.ID); err != nil { + cleanup, err := createIOStreamWithAgentcompatCapability(c, streamId, getUid(c), server.ID, rpc.AgentCompatCapabilityFileManager) + if err != nil { return nil, err } - fmData, _ := json.Marshal(&model.TaskFM{ + fmData, err := json.Marshal(&model.TaskFM{ StreamID: streamId, }) + if err != nil { + // A stream is owned by the caller only after this function succeeds. + cleanup() + return nil, err + } if err := server.SendTask(&proto.Task{ Type: model.TaskTypeFM, Data: string(fmData), }); err != nil { + cleanup() return nil, err } @@ -92,21 +99,14 @@ func fmStream(c *gin.Context) (any, error) { if err != nil { return nil, newWsError("%v", err) } - defer wsConn.Close() conn := websocketx.NewConn(wsConn) + pingTransport := newWebsocketPingTransport(conn, wsConn.Close) + stopPing := startWebsocketPingTicker(c.Request.Context(), time.Second*10, pingTransport) - deregisterPAT := registerPATConnection(c, func() { _ = wsConn.Close() }) + deregisterPAT := registerPATConnection(c, func() { _ = pingTransport.Close() }) defer deregisterPAT() - - go func() { - // PING 保活 - for { - if err = conn.WriteMessage(websocket.PingMessage, []byte{}); err != nil { - return - } - time.Sleep(time.Second * 10) - } - }() + // Join the ping worker before PAT and WebSocket cleanup can close its writer. + defer stopPing() if err = rpc.NezhaHandlerSingleton.UserConnected(streamId, conn); err != nil { return nil, newWsError("%v", err) diff --git a/cmd/dashboard/controller/io_stream_state_agentcompat_test.go b/cmd/dashboard/controller/io_stream_state_agentcompat_test.go new file mode 100644 index 00000000..78299f21 --- /dev/null +++ b/cmd/dashboard/controller/io_stream_state_agentcompat_test.go @@ -0,0 +1,214 @@ +//go:build agentcompat + +package controller + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "sync" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/service/rpc" +) + +func TestAgentcompatIOStreamStateRoutesRequirePATAndReturnTypedState(t *testing.T) { + cleanup, userID := setupMCPTest(t) + defer cleanup() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, token := mkToken(t, userID, []string{model.ScopeInventoryRead}, nil) + gin.SetMode(gin.TestMode) + router := gin.New() + registerAgentcompatRoutes(router) + server := httptest.NewServer(router) + t.Cleanup(server.Close) + + unauthenticated, err := http.NewRequest(http.MethodGet, server.URL+"/agentcompat/io-stream-state", nil) + require.NoError(t, err) + response, err := server.Client().Do(unauthenticated) + require.NoError(t, err) + require.Equal(t, http.StatusUnauthorized, response.StatusCode) + response.Body.Close() + + request, err := http.NewRequest(http.MethodGet, server.URL+"/agentcompat/io-stream-state", nil) + require.NoError(t, err) + request.Header.Set("Authorization", "Bearer "+token) + response, err = server.Client().Do(request) + require.NoError(t, err) + defer response.Body.Close() + var envelope model.CommonResponse[rpc.IOStreamState] + responseBody, readErr := io.ReadAll(response.Body) + require.NoError(t, readErr) + require.NotContains(t, string(responseBody), "route-wait") + require.NoError(t, json.Unmarshal(responseBody, &envelope)) + require.Equal(t, http.StatusOK, response.StatusCode) + require.True(t, envelope.Success) + require.Equal(t, rpc.IOStreamState{}, envelope.Data) +} + +func TestAgentcompatIOStreamStateWaitRouteWakesAfterClose(t *testing.T) { + cleanup, userID := setupMCPTest(t) + defer cleanup() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, token := mkToken(t, userID, []string{model.ScopeInventoryRead}, nil) + require.NoError(t, rpc.NezhaHandlerSingleton.CreateStream("route-wait", 1, 7)) + gin.SetMode(gin.TestMode) + router := gin.New() + registerAgentcompatRoutes(router) + server := httptest.NewServer(router) + t.Cleanup(server.Close) + body, err := json.Marshal(rpc.IOStreamStateExpectation{ExpectedCount: rpc.ExpectedIOStreamCount(0), AbsentStreamID: "route-wait"}) + require.NoError(t, err) + requestContext, cancel := context.WithCancel(context.Background()) + defer cancel() + request, err := http.NewRequestWithContext(requestContext, http.MethodPost, server.URL+"/agentcompat/io-stream-state", bytes.NewReader(body)) + require.NoError(t, err) + request.Header.Set("Authorization", "Bearer "+token) + request.Header.Set("Content-Type", "application/json") + result := make(chan *http.Response, 1) + go func() { + response, requestErr := server.Client().Do(request) + if requestErr == nil { + result <- response + } + }() + rpc.NezhaHandlerSingleton.CloseStream("route-wait") + response := <-result + defer response.Body.Close() + var envelope model.CommonResponse[rpc.IOStreamState] + responseBody, readErr := io.ReadAll(response.Body) + require.NoError(t, readErr) + require.NotContains(t, string(responseBody), "route-wait") + require.NoError(t, json.Unmarshal(responseBody, &envelope)) + require.True(t, envelope.Success) + require.Equal(t, 0, envelope.Data.Count) + require.Equal(t, uint64(2), envelope.Data.Generation) +} + +func TestAgentcompatIOStreamStateCreateWaitRouteWakesAfterCreate(t *testing.T) { + cleanup, userID := setupMCPTest(t) + defer cleanup() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, token := mkToken(t, userID, []string{model.ScopeInventoryRead}, nil) + gin.SetMode(gin.TestMode) + router := gin.New() + waitCaptured := make(chan struct{}) + var observerOnce sync.Once + rpc.NezhaHandlerSingleton.SetIOStreamStateWaitObserverForAgentcompat(func() { + observerOnce.Do(func() { close(waitCaptured) }) + }) + t.Cleanup(func() { + rpc.NezhaHandlerSingleton.SetIOStreamStateWaitObserverForAgentcompat(nil) + }) + registerAgentcompatRoutes(router) + server := httptest.NewServer(router) + t.Cleanup(server.Close) + body := bytes.NewBufferString(`{"expected_count":1}`) + requestContext, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + request, err := http.NewRequestWithContext(requestContext, http.MethodPost, server.URL+"/agentcompat/io-stream-state", body) + require.NoError(t, err) + request.Header.Set("Authorization", "Bearer "+token) + request.Header.Set("Content-Type", "application/json") + result := make(chan *http.Response, 1) + go func() { + response, requestErr := server.Client().Do(request) + if requestErr == nil { + result <- response + } + }() + select { + case <-waitCaptured: + case <-requestContext.Done(): + t.Fatal(requestContext.Err()) + } + require.NoError(t, rpc.NezhaHandlerSingleton.CreateStream("route-create", 1, 7)) + var response *http.Response + select { + case response = <-result: + case <-requestContext.Done(): + t.Fatal(requestContext.Err()) + } + defer response.Body.Close() + responseBody, err := io.ReadAll(response.Body) + require.NoError(t, err) + var envelope model.CommonResponse[rpc.IOStreamState] + require.NoError(t, json.Unmarshal(responseBody, &envelope)) + require.True(t, envelope.Success) + require.Equal(t, 1, envelope.Data.Count) + require.Equal(t, uint64(1), envelope.Data.Generation) + require.NotContains(t, string(responseBody), "route-create") +} + +func TestAgentcompatIOStreamStateRejectsInvalidExpectationAndCancellation(t *testing.T) { + cleanup, userID := setupMCPTest(t) + defer cleanup() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + _, token := mkToken(t, userID, []string{model.ScopeInventoryRead}, nil) + gin.SetMode(gin.TestMode) + router := gin.New() + registerAgentcompatRoutes(router) + server := httptest.NewServer(router) + t.Cleanup(server.Close) + for _, rawBody := range []string{`{}`, `{"expected_count":null}`, `{"expected_count":-1,"absent_stream_id":"private-stream-id"}`} { + body := bytes.NewBufferString(rawBody) + request, err := http.NewRequest(http.MethodPost, server.URL+"/agentcompat/io-stream-state", body) + require.NoError(t, err) + request.Header.Set("Authorization", "Bearer "+token) + request.Header.Set("Content-Type", "application/json") + response, err := server.Client().Do(request) + require.NoError(t, err) + responseBody, readErr := io.ReadAll(response.Body) + response.Body.Close() + require.NoError(t, readErr) + var envelope model.CommonResponse[rpc.IOStreamState] + require.NoError(t, json.Unmarshal(responseBody, &envelope)) + require.False(t, envelope.Success) + require.NotEmpty(t, envelope.Error) + require.NotContains(t, string(responseBody), "private-stream-id") + } + + body := bytes.NewBufferString(`{"absent_stream_id":"absence-only"}`) + request, err := http.NewRequest(http.MethodPost, server.URL+"/agentcompat/io-stream-state", body) + require.NoError(t, err) + request.Header.Set("Authorization", "Bearer "+token) + request.Header.Set("Content-Type", "application/json") + response, err := server.Client().Do(request) + require.NoError(t, err) + defer response.Body.Close() + var envelope model.CommonResponse[rpc.IOStreamState] + require.NoError(t, json.NewDecoder(response.Body).Decode(&envelope)) + require.True(t, envelope.Success) + require.Equal(t, 0, envelope.Data.Count) + + unauthenticatedBody := bytes.NewBufferString(`{"expected_count":0}`) + unauthenticatedPost, err := http.NewRequest(http.MethodPost, server.URL+"/agentcompat/io-stream-state", unauthenticatedBody) + require.NoError(t, err) + unauthenticatedPost.Header.Set("Content-Type", "application/json") + unauthenticatedResponse, err := server.Client().Do(unauthenticatedPost) + require.NoError(t, err) + unauthenticatedResponse.Body.Close() + require.Equal(t, http.StatusUnauthorized, unauthenticatedResponse.StatusCode) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + _, waitErr := rpc.NezhaHandlerSingleton.WaitForIOStreamState(ctx, rpc.IOStreamStateExpectation{ExpectedCount: rpc.ExpectedIOStreamCount(1)}) + require.True(t, errors.Is(waitErr, context.Canceled)) +} diff --git a/cmd/dashboard/controller/terminal.go b/cmd/dashboard/controller/terminal.go index 1661dab0..0a78c65b 100644 --- a/cmd/dashboard/controller/terminal.go +++ b/cmd/dashboard/controller/terminal.go @@ -5,7 +5,6 @@ import ( "github.com/gin-gonic/gin" "github.com/goccy/go-json" - "github.com/gorilla/websocket" "github.com/hashicorp/go-uuid" "github.com/nezhahq/nezha/model" @@ -25,6 +24,7 @@ import ( // @Success 200 {object} model.CreateTerminalResponse // @Router /terminal [post] func createTerminal(c *gin.Context) (*model.CreateTerminalResponse, error) { + prepareAgentcompatCapabilityHeader(c) var createTerminalReq model.TerminalForm if err := c.ShouldBind(&createTerminalReq); err != nil { return nil, err @@ -47,17 +47,24 @@ func createTerminal(c *gin.Context) (*model.CreateTerminalResponse, error) { return nil, err } - if err := rpc.NezhaHandlerSingleton.CreateStream(streamId, getUid(c), server.ID); err != nil { + cleanup, err := createIOStreamWithAgentcompatCapability(c, streamId, getUid(c), server.ID, rpc.AgentCompatCapabilityTerminal) + if err != nil { return nil, err } - terminalData, _ := json.Marshal(&model.TerminalTask{ + terminalData, err := json.Marshal(&model.TerminalTask{ StreamID: streamId, }) + if err != nil { + // A stream is owned by the caller only after this function succeeds. + cleanup() + return nil, err + } if err := server.SendTask(&proto.Task{ Type: model.TaskTypeTerminalGRPC, Data: string(terminalData), }); err != nil { + cleanup() return nil, err } @@ -93,21 +100,14 @@ func terminalStream(c *gin.Context) (any, error) { if err != nil { return nil, newWsError("%v", err) } - defer wsConn.Close() conn := websocketx.NewConn(wsConn) + pingTransport := newWebsocketPingTransport(conn, wsConn.Close) + stopPing := startWebsocketPingTicker(c.Request.Context(), time.Second*10, pingTransport) - deregisterPAT := registerPATConnection(c, func() { _ = wsConn.Close() }) + deregisterPAT := registerPATConnection(c, func() { _ = pingTransport.Close() }) defer deregisterPAT() - - go func() { - // PING 保活 - for { - if err = conn.WriteMessage(websocket.PingMessage, []byte{}); err != nil { - return - } - time.Sleep(time.Second * 10) - } - }() + // Join the ping worker before PAT and WebSocket cleanup can close its writer. + defer stopPing() if err = rpc.NezhaHandlerSingleton.UserConnected(streamId, conn); err != nil { return nil, newWsError("%v", err) diff --git a/cmd/dashboard/controller/terminal_fm_agentcompat.go b/cmd/dashboard/controller/terminal_fm_agentcompat.go new file mode 100644 index 00000000..39a7299d --- /dev/null +++ b/cmd/dashboard/controller/terminal_fm_agentcompat.go @@ -0,0 +1,64 @@ +//go:build agentcompat + +package controller + +import ( + "github.com/gin-gonic/gin" + + "github.com/nezhahq/nezha/pkg/agentcompatcontract" + "github.com/nezhahq/nezha/service/rpc" +) + +type agentcompatCapabilityHeaderContextKey struct{} + +func prepareAgentcompatCapabilityHeader(c *gin.Context) { + values := c.Request.Header.Values(agentcompatcontract.IOStreamCapabilityHeader) + if len(values) == 0 { + return + } + c.Request.Header.Del(agentcompatcontract.IOStreamCapabilityHeader) + c.Set(agentcompatCapabilityHeaderContextKey{}, values) +} + +func createIOStreamWithAgentcompatCapability(c *gin.Context, streamID string, creatorUserID, serverID uint64, purpose rpc.AgentCompatCapabilityPurpose) (func(), error) { + identity := agentcompatCapabilityIdentity{purpose: purpose, serverID: serverID} + rawValues, _ := c.Get(agentcompatCapabilityHeaderContextKey{}) + values, _ := rawValues.([]string) + if len(values) == 0 { + if err := rpc.NezhaHandlerSingleton.CreateStream(streamID, creatorUserID, serverID); err != nil { + return nil, err + } + return func() { _ = rpc.NezhaHandlerSingleton.CloseStream(streamID) }, nil + } + if len(values) != 1 || values[0] == "" { + return nil, errAgentcompatCapabilityUnavailable + } + capability, err := rpc.ParseAgentCompatIOStreamCapability(values[0]) + if err != nil { + return nil, errAgentcompatCapabilityUnavailable + } + proof, err := currentAgentcompatCapabilityProof(c, identity) + if err != nil { + return nil, errAgentcompatCapabilityUnavailable + } + access := proof.access(capability) + if err := rpc.NezhaHandlerSingleton.CreateStreamWithPurpose(streamID, creatorUserID, serverID, agentcompatStreamPurpose(purpose)); err != nil { + _ = rpc.NezhaHandlerSingleton.UnregisterAgentCompatIOStreamCapability(access) + return nil, err + } + // Bind before dispatch preserves exact cancellation when the create response is lost. + if err := rpc.NezhaHandlerSingleton.BindAgentCompatIOStreamCapability(rpc.AgentCompatCapabilityBinding{AgentCompatCapabilityAccess: access, StreamID: streamID}); err != nil { + _ = rpc.NezhaHandlerSingleton.CloseStream(streamID) + return nil, errAgentcompatCapabilityUnavailable + } + return func() { + _ = rpc.NezhaHandlerSingleton.CancelAgentCompatIOStreamCapability(access) + }, nil +} + +func agentcompatStreamPurpose(purpose rpc.AgentCompatCapabilityPurpose) rpc.StreamPurpose { + if purpose == rpc.AgentCompatCapabilityTerminal { + return rpc.PurposeTerminal + } + return rpc.PurposeFileManager +} diff --git a/cmd/dashboard/controller/terminal_fm_agentcompat_cleanup_test.go b/cmd/dashboard/controller/terminal_fm_agentcompat_cleanup_test.go new file mode 100644 index 00000000..eab6b84f --- /dev/null +++ b/cmd/dashboard/controller/terminal_fm_agentcompat_cleanup_test.go @@ -0,0 +1,130 @@ +//go:build agentcompat + +package controller + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/agentcompatcontract" + "github.com/nezhahq/nezha/service/rpc" + "github.com/nezhahq/nezha/service/singleton" +) + +func TestCreateAgentcompatNoHeaderPreservesLegacyPurposeForTerminalAndFM(t *testing.T) { + for _, testCase := range []struct { + name string + purpose rpc.AgentCompatCapabilityPurpose + target string + body any + }{ + {name: "terminal", purpose: rpc.AgentCompatCapabilityTerminal, target: "/terminal", body: model.TerminalForm{ServerID: 7}}, + {name: "file manager", purpose: rpc.AgentCompatCapabilityFileManager, target: "/file?id=7"}, + } { + t.Run(testCase.name, func(t *testing.T) { + handler, token, request := newAgentcompatCreateFixture(t, "POST", testCase.target, testCase.body, testCase.purpose) + capability := registerAgentcompatForCreate(t, handler, token, testCase.purpose) + probe := &agentcompatTaskProbe{test: t} + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + server.SetTaskStream(probe) + var responseStream string + if testCase.purpose == rpc.AgentCompatCapabilityTerminal { + response, err := createTerminal(request) + require.NoError(t, err) + responseStream = response.SessionID + } else { + response, err := createFM(request) + require.NoError(t, err) + responseStream = response.SessionID + } + require.Equal(t, 1, probe.calls()) + access := agentcompatAccessForCreate(t, handler, token, testCase.purpose, capability) + require.ErrorIs(t, handler.BindAgentCompatIOStreamCapability(rpc.AgentCompatCapabilityBinding{AgentCompatCapabilityAccess: access, StreamID: responseStream}), rpc.ErrAgentCompatCapabilityHidden) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(access)) + require.NoError(t, handler.CloseStream(responseStream)) + }) + } +} + +func TestCreateAgentcompatSendFailureReleasesExactPATBoundaryForTerminalAndFM(t *testing.T) { + for _, testCase := range []struct { + name string + purpose rpc.AgentCompatCapabilityPurpose + target string + body any + }{ + {name: "terminal", purpose: rpc.AgentCompatCapabilityTerminal, target: "/terminal", body: model.TerminalForm{ServerID: 7}}, + {name: "file manager", purpose: rpc.AgentCompatCapabilityFileManager, target: "/file?id=7"}, + } { + t.Run(testCase.name, func(t *testing.T) { + handler, token, request := newAgentcompatCreateFixture(t, "POST", testCase.target, testCase.body, testCase.purpose) + failed := registerAgentcompatForCreate(t, handler, token, testCase.purpose) + request.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, failed) + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + server.SetTaskStream(&agentcompatTaskProbe{test: t, sendErr: errors.New("dispatch failed")}) + if testCase.purpose == rpc.AgentCompatCapabilityTerminal { + _, err := createTerminal(request) + require.Error(t, err) + } else { + _, err := createFM(request) + require.Error(t, err) + } + capabilities := make([]string, 0, 16) + for range 16 { + capabilities = append(capabilities, registerAgentcompatForCreate(t, handler, token, testCase.purpose)) + } + _, err := handler.RegisterAgentCompatIOStreamCapability(context.Background(), rpc.AgentCompatCapabilityRegistration{Owner: rpc.AgentCompatCapabilityOwner{PATID: token.ID, UserID: token.UserID}, Purpose: testCase.purpose, TargetServerID: 7, ServerAccessAllowed: true}) + require.ErrorIs(t, err, rpc.ErrAgentCompatCapabilityUnavailable) + for _, raw := range capabilities { + access := agentcompatAccessForCreate(t, handler, token, testCase.purpose, raw) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(access)) + } + }) + } +} + +func TestCreateAgentcompatResponseLossWaitCancelForTerminalAndFM(t *testing.T) { + for _, testCase := range []struct { + name string + purpose rpc.AgentCompatCapabilityPurpose + target string + body any + }{ + {name: "terminal", purpose: rpc.AgentCompatCapabilityTerminal, target: "/terminal", body: model.TerminalForm{ServerID: 7}}, + {name: "file manager", purpose: rpc.AgentCompatCapabilityFileManager, target: "/file?id=7"}, + } { + t.Run(testCase.name, func(t *testing.T) { + handler, token, request := newAgentcompatCreateFixture(t, "POST", testCase.target, testCase.body, testCase.purpose) + capability := registerAgentcompatForCreate(t, handler, token, testCase.purpose) + request.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability) + probe := &agentcompatTaskProbe{test: t} + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + server.SetTaskStream(probe) + var streamID string + if testCase.purpose == rpc.AgentCompatCapabilityTerminal { + response, err := createTerminal(request) + require.NoError(t, err) + streamID = response.SessionID + } else { + response, err := createFM(request) + require.NoError(t, err) + streamID = response.SessionID + } + access := agentcompatAccessForCreate(t, handler, token, testCase.purpose, capability) + waited, err := handler.WaitAgentCompatIOStreamCapability(context.Background(), access) + require.NoError(t, err) + require.Equal(t, streamID, waited) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(access)) + require.Equal(t, 0, handler.StreamCount()) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(access)) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(access)) + }) + } +} diff --git a/cmd/dashboard/controller/terminal_fm_agentcompat_default.go b/cmd/dashboard/controller/terminal_fm_agentcompat_default.go new file mode 100644 index 00000000..8d424609 --- /dev/null +++ b/cmd/dashboard/controller/terminal_fm_agentcompat_default.go @@ -0,0 +1,18 @@ +//go:build !agentcompat + +package controller + +import ( + "github.com/gin-gonic/gin" + + "github.com/nezhahq/nezha/service/rpc" +) + +func prepareAgentcompatCapabilityHeader(*gin.Context) {} + +func createIOStreamWithAgentcompatCapability(_ *gin.Context, streamID string, creatorUserID, serverID uint64, _ rpc.AgentCompatCapabilityPurpose) (func(), error) { + if err := rpc.NezhaHandlerSingleton.CreateStream(streamID, creatorUserID, serverID); err != nil { + return nil, err + } + return func() { _ = rpc.NezhaHandlerSingleton.CloseStream(streamID) }, nil +} diff --git a/cmd/dashboard/controller/terminal_fm_agentcompat_default_test.go b/cmd/dashboard/controller/terminal_fm_agentcompat_default_test.go new file mode 100644 index 00000000..c1257a10 --- /dev/null +++ b/cmd/dashboard/controller/terminal_fm_agentcompat_default_test.go @@ -0,0 +1,69 @@ +//go:build !agentcompat + +package controller + +import ( + "errors" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/agentcompatcontract" + "github.com/nezhahq/nezha/service/rpc" + "github.com/nezhahq/nezha/service/singleton" +) + +func TestDefaultCreateTerminalPreservesCapabilityHeaderAndLegacyDispatch(t *testing.T) { + // Given + handler, _, request := newDefaultCreateFixture(t, "POST", "/terminal", model.TerminalForm{ServerID: 7}) + request.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "malformed") + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + stream := &failingRequestTaskStream{err: errors.New("stop after dispatch")} + server.SetTaskStream(stream) + + // When + _, err := createTerminal(request) + + // Then + require.ErrorIs(t, err, stream.err) + require.Equal(t, "malformed", request.Request.Header.Get(agentcompatcontract.IOStreamCapabilityHeader)) + require.Equal(t, 1, stream.calls()) + require.Equal(t, 0, handler.StreamCount()) +} + +func TestDefaultCreateFMPreservesCapabilityHeaderAndLegacyDispatch(t *testing.T) { + // Given + handler, _, request := newDefaultCreateFixture(t, "POST", "/file?id=7", nil) + request.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "malformed") + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + stream := &failingRequestTaskStream{err: errors.New("stop after dispatch")} + server.SetTaskStream(stream) + + // When + _, err := createFM(request) + + // Then + require.ErrorIs(t, err, stream.err) + require.Equal(t, "malformed", request.Request.Header.Get(agentcompatcontract.IOStreamCapabilityHeader)) + require.Equal(t, 1, stream.calls()) + require.Equal(t, 0, handler.StreamCount()) +} + +func newDefaultCreateFixture(t *testing.T, method, target string, body any) (*rpc.NezhaHandler, *model.APIToken, *gin.Context) { + t.Helper() + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + handler := rpc.NewNezhaHandler() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = handler + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + token, _ := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + request := newAuthorizedControllerContext(t, method, target, body) + request.Set(apiTokenCtxKey, token) + request.Set(model.CtxKeyAPIToken, token) + return handler, token, request +} diff --git a/cmd/dashboard/controller/terminal_fm_agentcompat_rejection_test.go b/cmd/dashboard/controller/terminal_fm_agentcompat_rejection_test.go new file mode 100644 index 00000000..323f33b1 --- /dev/null +++ b/cmd/dashboard/controller/terminal_fm_agentcompat_rejection_test.go @@ -0,0 +1,101 @@ +//go:build agentcompat + +package controller + +import ( + "testing" + + "github.com/gin-gonic/gin" + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/agentcompatcontract" + "github.com/nezhahq/nezha/service/rpc" + "github.com/nezhahq/nezha/service/singleton" + "github.com/stretchr/testify/require" +) + +func TestCreateIOStreamAgentcompatRejectsAdversarialHeadersSymmetrically(t *testing.T) { + cases := []struct { + name string + configure func(*testing.T, *rpc.NezhaHandler, *model.APIToken, *gin.Context, string) + purpose rpc.AgentCompatCapabilityPurpose + request string + wantHeader bool + }{ + {name: "malformed terminal", purpose: rpc.AgentCompatCapabilityTerminal, request: "bad"}, + {name: "empty terminal", purpose: rpc.AgentCompatCapabilityTerminal, request: ""}, + {name: "duplicate terminal", purpose: rpc.AgentCompatCapabilityTerminal, request: "duplicate"}, + {name: "jwt only terminal", purpose: rpc.AgentCompatCapabilityTerminal, request: "jwt"}, + {name: "foreign pat terminal", purpose: rpc.AgentCompatCapabilityTerminal, request: "foreign"}, + {name: "whitelist terminal", purpose: rpc.AgentCompatCapabilityTerminal, request: "whitelist"}, + {name: "malformed file manager", purpose: rpc.AgentCompatCapabilityFileManager, request: "bad"}, + {name: "empty file manager", purpose: rpc.AgentCompatCapabilityFileManager, request: ""}, + {name: "duplicate file manager", purpose: rpc.AgentCompatCapabilityFileManager, request: "duplicate"}, + {name: "jwt only file manager", purpose: rpc.AgentCompatCapabilityFileManager, request: "jwt"}, + {name: "foreign pat file manager", purpose: rpc.AgentCompatCapabilityFileManager, request: "foreign"}, + {name: "whitelist file manager", purpose: rpc.AgentCompatCapabilityFileManager, request: "whitelist"}, + } + for _, testCase := range cases { + t.Run(testCase.name, func(t *testing.T) { + // Given + handler, token, request := newAgentcompatCreateFixture(t, "POST", createAgentcompatTarget(testCase.purpose), createAgentcompatBody(testCase.purpose), testCase.purpose) + capability := registerAgentcompatForCreate(t, handler, token, testCase.purpose) + request.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability) + if testCase.request == "duplicate" { + request.Request.Header.Add(agentcompatcontract.IOStreamCapabilityHeader, capability) + } + if testCase.request == "bad" || testCase.request == "" { + request.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, testCase.request) + } + if testCase.request == "jwt" { + request.Set(apiTokenCtxKey, nil) + request.Set(model.CtxKeyAPIToken, nil) + } + if testCase.request == "foreign" { + foreign, _ := mkDistinctCapabilityToken(t, token.UserID, testCase.name) + request.Set(apiTokenCtxKey, foreign) + request.Set(model.CtxKeyAPIToken, foreign) + } + if testCase.request == "whitelist" { + token.SetServerIDs([]uint64{99}) + require.NoError(t, singleton.DB.Model(token).Update("servers_csv", token.ServersCSV).Error) + } + if testCase.configure != nil { + testCase.configure(t, handler, token, request, capability) + } + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + probe := &agentcompatTaskProbe{test: t} + server.SetTaskStream(probe) + + // When + var err error + if testCase.purpose == rpc.AgentCompatCapabilityTerminal { + _, err = createTerminal(request) + } else { + _, err = createFM(request) + } + + // Then + require.Error(t, err) + require.Equal(t, 0, probe.calls()) + require.Equal(t, 0, handler.StreamCount()) + require.Empty(t, request.Request.Header.Values(agentcompatcontract.IOStreamCapabilityHeader)) + ownerAccess := agentcompatAccessForCreate(t, handler, token, testCase.purpose, capability) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(ownerAccess)) + }) + } +} + +func createAgentcompatTarget(purpose rpc.AgentCompatCapabilityPurpose) string { + if purpose == rpc.AgentCompatCapabilityTerminal { + return "/terminal" + } + return "/file?id=7" +} + +func createAgentcompatBody(purpose rpc.AgentCompatCapabilityPurpose) any { + if purpose == rpc.AgentCompatCapabilityTerminal { + return model.TerminalForm{ServerID: 7} + } + return nil +} diff --git a/cmd/dashboard/controller/terminal_fm_agentcompat_test.go b/cmd/dashboard/controller/terminal_fm_agentcompat_test.go new file mode 100644 index 00000000..b2a2bc48 --- /dev/null +++ b/cmd/dashboard/controller/terminal_fm_agentcompat_test.go @@ -0,0 +1,213 @@ +//go:build agentcompat + +package controller + +import ( + "context" + "encoding/json" + "errors" + "sync" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/agentcompatcontract" + pb "github.com/nezhahq/nezha/proto" + "github.com/nezhahq/nezha/service/rpc" + "github.com/nezhahq/nezha/service/singleton" +) + +type agentcompatTaskProbe struct { + pb.NezhaService_RequestTaskServer + mu sync.Mutex + test *testing.T + check func(*testing.T) + sendErr error + sendCall int + task *pb.Task +} + +func (probe *agentcompatTaskProbe) Send(task *pb.Task) error { + probe.mu.Lock() + probe.sendCall++ + probe.task = task + check := probe.check + err := probe.sendErr + probe.mu.Unlock() + if check != nil { + check(probe.test) + } + return err +} + +func (probe *agentcompatTaskProbe) calls() int { + probe.mu.Lock() + defer probe.mu.Unlock() + return probe.sendCall +} + +func (probe *agentcompatTaskProbe) Context() context.Context { return context.Background() } + +func TestCreateTerminalAgentcompatBindsBeforeDispatchAndRemovesHeader(t *testing.T) { + // Given + handler, token, request := newAgentcompatCreateFixture(t, "POST", "/terminal", model.TerminalForm{ServerID: 7}, rpc.AgentCompatCapabilityTerminal) + capability := registerAgentcompatForCreate(t, handler, token, rpc.AgentCompatCapabilityTerminal) + request.Request.Header.Add(agentcompatcontract.IOStreamCapabilityHeader, capability) + var waited string + probe := &agentcompatTaskProbe{test: t} + probe.check = func(t *testing.T) { + access := agentcompatAccessForCreate(t, handler, token, rpc.AgentCompatCapabilityTerminal, capability) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + streamID, err := handler.WaitAgentCompatIOStreamCapability(ctx, access) + require.NoError(t, err) + waited = streamID + } + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + server.SetTaskStream(probe) + + // When + response, err := createTerminal(request) + + // Then + require.NoError(t, err) + require.Equal(t, response.SessionID, waited) + require.Empty(t, request.Request.Header.Values(agentcompatcontract.IOStreamCapabilityHeader)) + access := agentcompatAccessForCreate(t, handler, token, rpc.AgentCompatCapabilityTerminal, capability) + _, waitErr := handler.WaitAgentCompatIOStreamCapability(context.Background(), access) + require.NoError(t, waitErr) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(access)) + require.Equal(t, 0, handler.StreamCount()) +} + +func TestCreateFMAgentcompatRejectsTerminalCapabilityWithoutDispatch(t *testing.T) { + // Given + handler, token, request := newAgentcompatCreateFixture(t, "POST", "/file?id=7", nil, rpc.AgentCompatCapabilityFileManager) + capability := registerAgentcompatForCreate(t, handler, token, rpc.AgentCompatCapabilityTerminal) + request.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability) + probe := &agentcompatTaskProbe{test: t} + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + server.SetTaskStream(probe) + + // When + response, err := createFM(request) + + // Then + require.Error(t, err) + require.Nil(t, response) + require.Equal(t, 0, probe.sendCall) + require.Equal(t, 0, handler.StreamCount()) +} + +func TestCreateTerminalAgentcompatRejectsFileManagerCapabilityWithoutDispatch(t *testing.T) { + // Given + handler, token, request := newAgentcompatCreateFixture(t, "POST", "/terminal", model.TerminalForm{ServerID: 7}, rpc.AgentCompatCapabilityTerminal) + capability := registerAgentcompatForCreate(t, handler, token, rpc.AgentCompatCapabilityFileManager) + request.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability) + probe := &agentcompatTaskProbe{test: t} + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + server.SetTaskStream(probe) + + // When + response, err := createTerminal(request) + + // Then + require.Error(t, err) + require.Nil(t, response) + require.Equal(t, 0, probe.calls()) + require.Equal(t, 0, handler.StreamCount()) + access := agentcompatAccessForCreate(t, handler, token, rpc.AgentCompatCapabilityFileManager, capability) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(access)) +} + +func TestCreateFMAgentcompatBindsBeforeDispatchWithExactTaskStreamID(t *testing.T) { + // Given + handler, token, request := newAgentcompatCreateFixture(t, "POST", "/file?id=7", nil, rpc.AgentCompatCapabilityFileManager) + capability := registerAgentcompatForCreate(t, handler, token, rpc.AgentCompatCapabilityFileManager) + request.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability) + probe := &agentcompatTaskProbe{test: t} + probe.check = func(t *testing.T) { + access := agentcompatAccessForCreate(t, handler, token, rpc.AgentCompatCapabilityFileManager, capability) + streamID, err := handler.WaitAgentCompatIOStreamCapability(context.Background(), access) + require.NoError(t, err) + require.NotEmpty(t, streamID) + } + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + server.SetTaskStream(probe) + + // When + response, err := createFM(request) + + // Then + require.NoError(t, err) + require.Equal(t, 1, probe.calls()) + require.NotNil(t, probe.task) + require.Equal(t, uint64(model.TaskTypeFM), probe.task.Type) + var task model.TaskFM + require.NoError(t, json.Unmarshal([]byte(probe.task.Data), &task)) + require.Equal(t, response.SessionID, task.StreamID) + require.Empty(t, request.Request.Header.Values(agentcompatcontract.IOStreamCapabilityHeader)) + access := agentcompatAccessForCreate(t, handler, token, rpc.AgentCompatCapabilityFileManager, capability) + waited, err := handler.WaitAgentCompatIOStreamCapability(context.Background(), access) + require.NoError(t, err) + require.Equal(t, response.SessionID, waited) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(access)) +} + +func TestCreateTerminalAgentcompatSendFailureReleasesCapabilityAndStream(t *testing.T) { + // Given + handler, token, request := newAgentcompatCreateFixture(t, "POST", "/terminal", model.TerminalForm{ServerID: 7}, rpc.AgentCompatCapabilityTerminal) + capability := registerAgentcompatForCreate(t, handler, token, rpc.AgentCompatCapabilityTerminal) + request.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability) + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + server.SetTaskStream(&agentcompatTaskProbe{test: t, sendErr: errors.New("dispatch failed")}) + + // When + response, err := createTerminal(request) + + // Then + require.Error(t, err) + require.Nil(t, response) + require.Equal(t, 0, handler.StreamCount()) + access := agentcompatAccessForCreate(t, handler, token, rpc.AgentCompatCapabilityTerminal, capability) + _, waitErr := handler.WaitAgentCompatIOStreamCapability(context.Background(), access) + require.Error(t, waitErr) +} + +func newAgentcompatCreateFixture(t *testing.T, method, target string, body any, _ rpc.AgentCompatCapabilityPurpose) (*rpc.NezhaHandler, *model.APIToken, *gin.Context) { + t.Helper() + cleanup, userID := setupMCPTest(t) + t.Cleanup(cleanup) + handler := rpc.NewNezhaHandler() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = handler + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + token, _ := mkToken(t, userID, []string{model.ScopeServerRead}, nil) + request := newAuthorizedControllerContext(t, method, target, body) + request.Set(apiTokenCtxKey, token) + request.Set(model.CtxKeyAPIToken, token) + return handler, token, request +} + +func registerAgentcompatForCreate(t *testing.T, handler *rpc.NezhaHandler, token *model.APIToken, purpose rpc.AgentCompatCapabilityPurpose) string { + t.Helper() + capability, err := handler.RegisterAgentCompatIOStreamCapability(context.Background(), rpc.AgentCompatCapabilityRegistration{ + Owner: rpc.AgentCompatCapabilityOwner{PATID: token.ID, UserID: token.UserID}, Purpose: purpose, TargetServerID: 7, ServerAccessAllowed: true, + }) + require.NoError(t, err) + return capability.String() +} + +func agentcompatAccessForCreate(t *testing.T, handler *rpc.NezhaHandler, token *model.APIToken, purpose rpc.AgentCompatCapabilityPurpose, raw string) rpc.AgentCompatCapabilityAccess { + t.Helper() + capability, err := rpc.ParseAgentCompatIOStreamCapability(raw) + require.NoError(t, err) + return rpc.AgentCompatCapabilityAccess{Capability: capability, Owner: rpc.AgentCompatCapabilityOwner{PATID: token.ID, UserID: token.UserID}, Purpose: purpose, TargetServerID: 7, ServerAccessAllowed: true} +} diff --git a/cmd/dashboard/controller/terminal_fm_agentcompat_wrong_purpose_test.go b/cmd/dashboard/controller/terminal_fm_agentcompat_wrong_purpose_test.go new file mode 100644 index 00000000..363a5ed5 --- /dev/null +++ b/cmd/dashboard/controller/terminal_fm_agentcompat_wrong_purpose_test.go @@ -0,0 +1,92 @@ +//go:build agentcompat + +package controller + +import ( + "context" + "encoding/json" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/agentcompatcontract" + "github.com/nezhahq/nezha/service/rpc" + "github.com/nezhahq/nezha/service/singleton" +) + +func TestCreateFMWrongRoutePreservesTerminalCapability(t *testing.T) { + // Given + handler, token, wrongRequest := newAgentcompatCreateFixture(t, "POST", "/file?id=7", nil, rpc.AgentCompatCapabilityFileManager) + capability := registerAgentcompatForCreate(t, handler, token, rpc.AgentCompatCapabilityTerminal) + wrongRequest.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability) + probe := &agentcompatTaskProbe{test: t} + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + server.SetTaskStream(probe) + + // When + response, err := createFM(wrongRequest) + + // Then + require.Error(t, err) + require.Nil(t, response) + require.Empty(t, wrongRequest.Request.Header.Values(agentcompatcontract.IOStreamCapabilityHeader)) + require.Equal(t, 0, probe.calls()) + require.Equal(t, 0, handler.StreamCount()) + + correctRequest := newAuthorizedControllerContext(t, "POST", "/terminal", model.TerminalForm{ServerID: 7}) + correctRequest.Set(apiTokenCtxKey, token) + correctRequest.Set(model.CtxKeyAPIToken, token) + correctRequest.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability) + terminalResponse, err := createTerminal(correctRequest) + require.NoError(t, err) + require.Equal(t, 1, probe.calls()) + var task model.TerminalTask + require.NoError(t, json.Unmarshal([]byte(probe.task.Data), &task)) + require.Equal(t, terminalResponse.SessionID, task.StreamID) + access := agentcompatAccessForCreate(t, handler, token, rpc.AgentCompatCapabilityTerminal, capability) + waited, err := handler.WaitAgentCompatIOStreamCapability(context.Background(), access) + require.NoError(t, err) + require.Equal(t, terminalResponse.SessionID, waited) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(access)) + require.Equal(t, 0, handler.StreamCount()) +} + +func TestCreateTerminalWrongRoutePreservesFileManagerCapability(t *testing.T) { + // Given + handler, token, wrongRequest := newAgentcompatCreateFixture(t, "POST", "/terminal", model.TerminalForm{ServerID: 7}, rpc.AgentCompatCapabilityTerminal) + capability := registerAgentcompatForCreate(t, handler, token, rpc.AgentCompatCapabilityFileManager) + wrongRequest.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability) + probe := &agentcompatTaskProbe{test: t} + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + server.SetTaskStream(probe) + + // When + response, err := createTerminal(wrongRequest) + + // Then + require.Error(t, err) + require.Nil(t, response) + require.Empty(t, wrongRequest.Request.Header.Values(agentcompatcontract.IOStreamCapabilityHeader)) + require.Equal(t, 0, probe.calls()) + require.Equal(t, 0, handler.StreamCount()) + + correctRequest := newAuthorizedControllerContext(t, "POST", "/file?id=7", nil) + correctRequest.Set(apiTokenCtxKey, token) + correctRequest.Set(model.CtxKeyAPIToken, token) + correctRequest.Request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability) + fmResponse, err := createFM(correctRequest) + require.NoError(t, err) + require.Equal(t, 1, probe.calls()) + var task model.TaskFM + require.NoError(t, json.Unmarshal([]byte(probe.task.Data), &task)) + require.Equal(t, fmResponse.SessionID, task.StreamID) + access := agentcompatAccessForCreate(t, handler, token, rpc.AgentCompatCapabilityFileManager, capability) + waited, err := handler.WaitAgentCompatIOStreamCapability(context.Background(), access) + require.NoError(t, err) + require.Equal(t, fmResponse.SessionID, waited) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(access)) + require.Equal(t, 0, handler.StreamCount()) +} diff --git a/cmd/dashboard/controller/terminal_fm_lifecycle_test.go b/cmd/dashboard/controller/terminal_fm_lifecycle_test.go new file mode 100644 index 00000000..13eddd68 --- /dev/null +++ b/cmd/dashboard/controller/terminal_fm_lifecycle_test.go @@ -0,0 +1,122 @@ +package controller + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "net/http/httptest" + "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" +) + +type failingRequestTaskStream struct { + pb.NezhaService_RequestTaskServer + mu sync.Mutex + sendCalls int + err error +} + +func (stream *failingRequestTaskStream) Send(*pb.Task) error { + stream.mu.Lock() + defer stream.mu.Unlock() + stream.sendCalls++ + return stream.err +} + +func (stream *failingRequestTaskStream) calls() int { + stream.mu.Lock() + defer stream.mu.Unlock() + return stream.sendCalls +} + +func (stream *failingRequestTaskStream) Context() context.Context { return context.Background() } +func (stream *failingRequestTaskStream) SetHeader(metadata.MD) error { return nil } +func (stream *failingRequestTaskStream) SendHeader(metadata.MD) error { return nil } +func (stream *failingRequestTaskStream) SetTrailer(metadata.MD) {} +func (stream *failingRequestTaskStream) SendMsg(any) error { return nil } +func (stream *failingRequestTaskStream) RecvMsg(any) error { return nil } + +func newAuthorizedControllerContext(t *testing.T, method, target string, body any) *gin.Context { + t.Helper() + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + context, _ := gin.CreateTestContext(recorder) + encoded, err := json.Marshal(body) + require.NoError(t, err) + context.Request = httptest.NewRequest(method, target, bytes.NewReader(encoded)) + context.Request.Header.Set("Content-Type", "application/json") + context.Set(model.CtxKeyAuthorizedUser, &model.User{Common: model.Common{ID: 100}, Role: model.RoleMember}) + return context +} + +func TestCreateTerminalReturnsSendErrorAndReleasesStreamCapacity(t *testing.T) { + cleanupFixture, _ := setupMCPTest(t) + defer cleanupFixture() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + + sendError := errors.New("terminal task send failed") + stream := &failingRequestTaskStream{err: sendError} + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + server.SetTaskStream(stream) + + request := newAuthorizedControllerContext(t, "POST", "/terminal", model.TerminalForm{ServerID: 7}) + response, err := createTerminal(request) + require.ErrorIs(t, err, sendError) + require.Nil(t, response) + require.Equal(t, 1, stream.calls()) + assertStreamCapacityReusable(t, rpc.NezhaHandlerSingleton, 100, 7, "terminal-reused") +} + +func TestCreateFMReturnsSendErrorAndReleasesStreamCapacity(t *testing.T) { + cleanupFixture, _ := setupMCPTest(t) + defer cleanupFixture() + originalHandler := rpc.NezhaHandlerSingleton + rpc.NezhaHandlerSingleton = rpc.NewNezhaHandler() + t.Cleanup(func() { rpc.NezhaHandlerSingleton = originalHandler }) + + sendError := errors.New("FM task send failed") + stream := &failingRequestTaskStream{err: sendError} + server, ok := singleton.ServerShared.Get(7) + require.True(t, ok) + server.SetTaskStream(stream) + + request := newAuthorizedControllerContext(t, "POST", "/file?id=7", nil) + response, err := createFM(request) + require.ErrorIs(t, err, sendError) + require.Nil(t, response) + require.Equal(t, 1, stream.calls()) + assertStreamCapacityReusable(t, rpc.NezhaHandlerSingleton, 100, 7, "fm-reused") +} + +func assertStreamCapacityReusable(t *testing.T, handler *rpc.NezhaHandler, userID, serverID uint64, streamID string) { + t.Helper() + _, tracked := handler.StreamOwnership(streamID) + require.False(t, tracked, "failed task dispatch must not leave the replacement stream tracked") + for index := 0; index < 20; index++ { + require.NoError(t, handler.CreateStream(streamID+"-user-"+ctoa(uint64(index)), userID, serverID+uint64(index))) + } + require.ErrorIs(t, handler.CreateStream(streamID+"-user-over", userID, serverID+100), rpc.ErrTooManyStreamsForUser) + for index := 0; index < 20; index++ { + require.NoError(t, handler.CloseStream(streamID+"-user-"+ctoa(uint64(index)))) + } + for index := 0; index < 40; index++ { + require.NoError(t, handler.CreateStream(streamID+"-server-"+ctoa(uint64(index)), userID+uint64(index)+1000, serverID+1000)) + } + require.ErrorIs(t, handler.CreateStream(streamID+"-server-over", userID+1000, serverID+1000), rpc.ErrTooManyStreamsForServer) + for index := 0; index < 40; index++ { + require.NoError(t, handler.CloseStream(streamID+"-server-"+ctoa(uint64(index)))) + } +} diff --git a/cmd/dashboard/controller/ws.go b/cmd/dashboard/controller/ws.go index 89aacac4..67d1b44b 100644 --- a/cmd/dashboard/controller/ws.go +++ b/cmd/dashboard/controller/ws.go @@ -225,6 +225,7 @@ func patStreamContext(c *gin.Context) (model.APITokenAccessor, string) { 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 { + runtime := server.RuntimeSnapshot() if pat != nil && !pat.CanAccessServer(server.ID) { continue } @@ -236,15 +237,19 @@ func filterServersForViewer(servers []*model.Server, viewerUserID uint64, viewer if server.GeoIP != nil { countryCode = server.GeoIP.CountryCode } + publicHost := runtime.Host + if publicHost != nil && !isOwnerOrAdmin { + publicHost = publicHost.Filter() + } out = append(out, model.StreamServer{ ID: server.ID, Name: server.Name, PublicNote: utils.IfOr(withPublicNote, server.PublicNote, ""), DisplayIndex: server.DisplayIndex, - Host: utils.IfOr(isOwnerOrAdmin, server.Host, server.Host.Filter()), - State: server.State, + Host: publicHost, + State: runtime.State, CountryCode: countryCode, - LastActive: server.LastActive, + LastActive: runtime.LastActive, }) } return out diff --git a/cmd/dashboard/rpc/grpc_interceptor.go b/cmd/dashboard/rpc/grpc_interceptor.go new file mode 100644 index 00000000..68b80a06 --- /dev/null +++ b/cmd/dashboard/rpc/grpc_interceptor.go @@ -0,0 +1,94 @@ +package rpc + +import ( + "context" + "fmt" + "log" + "net/netip" + + "google.golang.org/grpc" + "google.golang.org/grpc/metadata" + "google.golang.org/grpc/peer" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/utils" + "github.com/nezhahq/nezha/service/singleton" +) + +func ctxWithRealIP(ctx context.Context) (context.Context, error) { + var ip, connectingIp string + p, ok := peer.FromContext(ctx) + if ok { + addrPort, err := netip.ParseAddrPort(p.Addr.String()) + if err == nil { + connectingIp = addrPort.Addr().String() + } + } + ctx = context.WithValue(ctx, model.CtxKeyConnectingIP{}, connectingIp) + + if singleton.Conf.AgentRealIPHeader == "" { + return ctx, nil + } + + if singleton.Conf.AgentRealIPHeader == model.ConfigUsePeerIP { + if connectingIp == "" { + return ctx, fmt.Errorf("connecting ip not found") + } + ip = connectingIp + } else { + vals := metadata.ValueFromIncomingContext(ctx, singleton.Conf.AgentRealIPHeader) + if len(vals) == 0 { + return ctx, fmt.Errorf("real ip header not found") + } + var err error + ip, err = utils.GetIPFromHeader(vals[0]) + if err != nil { + return ctx, err + } + } + + if singleton.Conf.Debug { + log.Printf("NEZHA>> gRPC Agent Real IP: %s, connecting IP: %s\n", ip, connectingIp) + } + + return context.WithValue(ctx, model.CtxKeyRealIP{}, ip), nil +} + +func waf(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { + realip, _ := ctx.Value(model.CtxKeyRealIP{}).(string) + if err := model.CheckIP(singleton.DB, realip); err != nil { + return nil, err + } + return handler(ctx, req) +} + +func getRealIp(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { + ctx, err := ctxWithRealIP(ctx) + if err != nil { + return nil, err + } + return handler(ctx, req) +} + +type realIPServerStream struct { + grpc.ServerStream + ctx context.Context +} + +func (s *realIPServerStream) Context() context.Context { return s.ctx } + +func getRealIpStream(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error { + ctx, err := ctxWithRealIP(ss.Context()) + if err != nil { + return err + } + return handler(srv, &realIPServerStream{ServerStream: ss, ctx: ctx}) +} + +func wafStream(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error { + realip, _ := ss.Context().Value(model.CtxKeyRealIP{}).(string) + if err := model.CheckIP(singleton.DB, realip); err != nil { + return err + } + return handler(srv, ss) +} diff --git a/cmd/dashboard/rpc/nat.go b/cmd/dashboard/rpc/nat.go new file mode 100644 index 00000000..d5b3acb6 --- /dev/null +++ b/cmd/dashboard/rpc/nat.go @@ -0,0 +1,115 @@ +package rpc + +import ( + "fmt" + "log" + "net/http" + "time" + + "github.com/goccy/go-json" + + "github.com/hashicorp/go-uuid" + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/utils" + "github.com/nezhahq/nezha/proto" + serviceRPC "github.com/nezhahq/nezha/service/rpc" + "github.com/nezhahq/nezha/service/singleton" +) + +func ServeNAT(w http.ResponseWriter, r *http.Request, natConfig *model.NAT) { + capabilityLease, capabilityErr := prepareNATCapability(r, natConfig) + if capabilityErr != nil { + w.WriteHeader(http.StatusServiceUnavailable) + w.Write([]byte("NAT capability unavailable")) + return + } + handler := serviceRPC.NezhaHandlerSingleton + streamId := "" + legacyStreamOwned := false + relayCompleted := false + defer func() { + if capabilityLease.active { + if !relayCompleted { + capabilityLease.cleanup(handler) + } + return + } + if legacyStreamOwned { + _ = handler.CloseStream(streamId) + } + }() + server, _ := singleton.ServerShared.Get(natConfig.ServerID) + if server == nil { + w.WriteHeader(http.StatusServiceUnavailable) + w.Write([]byte("server not found or not connected")) + return + } + if server.GetTaskStream() == nil { + w.WriteHeader(http.StatusServiceUnavailable) + w.Write([]byte("server not found or not connected")) + return + } + + streamId, err := uuid.GenerateUUID() + if err != nil { + w.WriteHeader(http.StatusServiceUnavailable) + w.Write(fmt.Appendf(nil, "stream id error: %v", err)) + return + } + if capabilityLease.active { + capabilityLease.streamLease, err = handler.CreateAgentCompatNATStream(capabilityLease.handle, streamId) + } else { + err = handler.CreateStream(streamId, 0, server.ID) + } + if err != nil { + w.WriteHeader(http.StatusTooManyRequests) + w.Write(fmt.Appendf(nil, "stream limit: %v", err)) + return + } + legacyStreamOwned = !capabilityLease.active + taskData, err := json.Marshal(model.TaskNAT{StreamID: streamId, Host: natConfig.Host}) + if err != nil { + w.WriteHeader(http.StatusServiceUnavailable) + w.Write(fmt.Appendf(nil, "task data error: %v", err)) + return + } + + if err := server.SendTask(&proto.Task{Type: model.TaskTypeNAT, Data: string(taskData)}); err != nil { + w.WriteHeader(http.StatusServiceUnavailable) + w.Write(fmt.Appendf(nil, "send task error: %v", err)) + return + } + + // Authorization authenticates Dashboard access before NAT ingress; it must not become an origin credential for the configured NAT backend. + wWrapped, err := utils.NewRequestWrapper(r, w) + if err != nil { + return + } + + if err := handler.UserConnected(streamId, wWrapped); err != nil { + if closeErr := wWrapped.Close(); closeErr != nil { + log.Printf("NEZHA>> NAT request wrapper close error after user connection failure: %v", closeErr) + } + return + } + + if capabilityLease.active { + if err := handler.PublishAgentCompatNATStream(capabilityLease.handle, serviceRPC.AgentCompatNATPublication{ + Purpose: serviceRPC.AgentCompatCapabilityNAT, + TargetServerID: server.ID, + ResourceID: natConfig.ID, + StreamID: streamId, + }); err != nil { + return + } + } + if capabilityLease.active { + capabilityLease.publicationOwned, err = handler.StartAgentCompatNATStream(capabilityLease.handle, time.Second*10) + if err != nil { + return + } + } else if err := handler.StartStream(streamId, time.Second*10); err != nil { + return + } + relayCompleted = true +} diff --git a/cmd/dashboard/rpc/nat_capability_agentcompat.go b/cmd/dashboard/rpc/nat_capability_agentcompat.go new file mode 100644 index 00000000..d346ae77 --- /dev/null +++ b/cmd/dashboard/rpc/nat_capability_agentcompat.go @@ -0,0 +1,48 @@ +//go:build agentcompat + +package rpc + +import ( + "errors" + "net/http" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/agentcompatcontract" + serviceRPC "github.com/nezhahq/nezha/service/rpc" +) + +type natCapabilityLease struct { + active bool + publicationOwned bool + streamLease *serviceRPC.AgentCompatNATStreamLease + access serviceRPC.AgentCompatCapabilityAccess + handle serviceRPC.AgentCompatNATPublishHandle +} + +func prepareNATCapability(request *http.Request, natConfig *model.NAT) (natCapabilityLease, error) { + values := request.Header.Values(agentcompatcontract.IOStreamCapabilityHeader) + if len(values) == 0 { + request.Header.Del("Authorization") + return natCapabilityLease{}, nil + } + request.Header.Del(agentcompatcontract.IOStreamCapabilityHeader) + if len(values) != 1 || values[0] == "" { + request.Header.Del("Authorization") + return natCapabilityLease{}, errors.New("invalid NAT capability") + } + access, handle, err := serviceRPC.NezhaHandlerSingleton.ConsumeAgentCompatNATCapabilityForProfile(values[0], natConfig.ServerID, natConfig.ID) + request.Header.Del("Authorization") + if err != nil { + return natCapabilityLease{}, errors.New("invalid NAT capability") + } + return natCapabilityLease{active: true, access: access, handle: handle}, nil +} + +func (lease natCapabilityLease) cleanup(handler *serviceRPC.NezhaHandler) { + if !lease.active { + return + } + _ = handler.CancelAgentCompatIOStreamCapability(lease.access) + _ = handler.CloseAgentCompatNATStreamLease(lease.streamLease) + _ = handler.UnregisterAgentCompatIOStreamCapability(lease.access) +} diff --git a/cmd/dashboard/rpc/nat_capability_agentcompat_test.go b/cmd/dashboard/rpc/nat_capability_agentcompat_test.go new file mode 100644 index 00000000..94910c40 --- /dev/null +++ b/cmd/dashboard/rpc/nat_capability_agentcompat_test.go @@ -0,0 +1,82 @@ +//go:build agentcompat + +package rpc + +import ( + "bytes" + "log" + "net/http" + "strings" + "testing" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/agentcompatcontract" + serviceRPC "github.com/nezhahq/nezha/service/rpc" +) + +func TestPrepareNATCapabilityConsumesAndRemovesHeaderBeforeNATWork(t *testing.T) { + // Given + handler := serviceRPC.NewNezhaHandler() + original := serviceRPC.NezhaHandlerSingleton + serviceRPC.NezhaHandlerSingleton = handler + t.Cleanup(func() { serviceRPC.NezhaHandlerSingleton = original }) + request := &http.Request{Header: make(http.Header)} + request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "malformed") + + // When + lease, err := prepareNATCapability(request, &model.NAT{Common: model.Common{ID: 91}, ServerID: 81}) + + // Then + if err == nil { + t.Fatal("malformed capability unexpectedly accepted") + } + if lease.active { + t.Fatal("malformed capability unexpectedly activated") + } + if _, present := request.Header[agentcompatcontract.IOStreamCapabilityHeader]; present { + t.Fatal("capability header remained after hook") + } +} + +func TestPrepareNATCapabilityRejectsDuplicateHeaderAfterRemovingAllValues(t *testing.T) { + // Given + request := &http.Request{Header: make(http.Header)} + request.Header.Add(agentcompatcontract.IOStreamCapabilityHeader, "one") + request.Header.Add(agentcompatcontract.IOStreamCapabilityHeader, "two") + + // When + lease, err := prepareNATCapability(request, &model.NAT{Common: model.Common{ID: 92}, ServerID: 82}) + + // Then + if err == nil { + t.Fatal("duplicate capability header unexpectedly accepted") + } + if lease.active { + t.Fatal("duplicate capability header unexpectedly activated") + } + if _, present := request.Header[agentcompatcontract.IOStreamCapabilityHeader]; present { + t.Fatal("duplicate capability header remained after hook") + } +} + +func TestServeNATAgentCompatSensitiveHeadersStayOutOfErrorsAndLogs(t *testing.T) { + request := &http.Request{Header: make(http.Header)} + request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "capability-secret") + request.Header.Set("Authorization", "Bearer deterministic-secret") + writer := &serveNATResponseWriter{} + var logs bytes.Buffer + originalOutput := log.Writer() + log.SetOutput(&logs) + t.Cleanup(func() { log.SetOutput(originalOutput) }) + + ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 91}, ServerID: 81}) + + if writer.status != http.StatusServiceUnavailable { + t.Fatalf("status = %d, want %d", writer.status, http.StatusServiceUnavailable) + } + for _, output := range []string{writer.body, logs.String()} { + if strings.Contains(output, "capability-secret") || strings.Contains(output, "deterministic-secret") { + t.Fatalf("sensitive value leaked in %q", output) + } + } +} diff --git a/cmd/dashboard/rpc/nat_capability_default.go b/cmd/dashboard/rpc/nat_capability_default.go new file mode 100644 index 00000000..8773477b --- /dev/null +++ b/cmd/dashboard/rpc/nat_capability_default.go @@ -0,0 +1,25 @@ +//go:build !agentcompat + +package rpc + +import ( + "net/http" + + "github.com/nezhahq/nezha/model" + serviceRPC "github.com/nezhahq/nezha/service/rpc" +) + +type natCapabilityLease struct { + active bool + publicationOwned bool + streamLease *serviceRPC.AgentCompatNATStreamLease + access serviceRPC.AgentCompatCapabilityAccess + handle serviceRPC.AgentCompatNATPublishHandle +} + +func prepareNATCapability(request *http.Request, _ *model.NAT) (natCapabilityLease, error) { + request.Header.Del("Authorization") + return natCapabilityLease{}, nil +} + +func (lease natCapabilityLease) cleanup(*serviceRPC.NezhaHandler) {} diff --git a/cmd/dashboard/rpc/nat_capability_default_flow_test.go b/cmd/dashboard/rpc/nat_capability_default_flow_test.go new file mode 100644 index 00000000..c3698afa --- /dev/null +++ b/cmd/dashboard/rpc/nat_capability_default_flow_test.go @@ -0,0 +1,64 @@ +//go:build !agentcompat + +package rpc + +import ( + "context" + "io" + "net/http" + "net/url" + "strings" + "testing" + "time" + + "github.com/goccy/go-json" + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/agentcompatcontract" + "github.com/nezhahq/nezha/proto" +) + +func TestServeNATDefaultForwardsCapabilityHeaderAsOrdinaryData(t *testing.T) { + fixture := newServeNATFixture(t) + connection := newServeNATConn() + agent := &serveNATAgent{readErr: io.EOF, writeDone: make(chan struct{})} + request := &http.Request{ + Method: http.MethodPost, + URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"}, + Header: make(http.Header), + Body: &serveNATBody{}, + } + request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "ordinary-data") + request.Header.Set("Authorization", "Bearer deterministic-secret") + fixture.taskStream.onSend = func(task *proto.Task) error { + var nat model.TaskNAT + if err := json.Unmarshal([]byte(task.Data), &nat); err != nil { + return err + } + return fixture.handler.AgentConnected(nat.StreamID, agent) + } + done := make(chan struct{}) + deadline, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + go func() { + ServeNAT(&serveNATResponseWriter{conn: connection}, request, &model.NAT{ServerID: fixture.server.ID, Host: "target.example"}) + close(done) + }() + select { + case <-agent.writeDone: + case <-deadline.Done(): + t.Fatal("agent did not receive ordinary NAT request") + } + _ = connection.Close() + select { + case <-done: + case <-deadline.Done(): + t.Fatal("ServeNAT did not finish") + } + forwarded := strings.ToLower(string(agent.writtenBytes())) + if !strings.Contains(forwarded, strings.ToLower(agentcompatcontract.IOStreamCapabilityHeader)+": ordinary-data") { + t.Fatal("default NAT flow did not forward capability header as ordinary request data") + } + if strings.Contains(forwarded, "Authorization:") || strings.Contains(forwarded, "deterministic-secret") { + t.Fatal("default NAT flow forwarded Authorization credentials") + } +} diff --git a/cmd/dashboard/rpc/nat_capability_default_test.go b/cmd/dashboard/rpc/nat_capability_default_test.go new file mode 100644 index 00000000..7ff5e686 --- /dev/null +++ b/cmd/dashboard/rpc/nat_capability_default_test.go @@ -0,0 +1,35 @@ +//go:build !agentcompat + +package rpc + +import ( + "net/http" + "testing" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/agentcompatcontract" +) + +func TestPrepareNATCapabilityDefaultPreservesHeaderAsOrdinaryRequestData(t *testing.T) { + // Given + request := &http.Request{Header: make(http.Header)} + request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "ordinary-data") + request.Header.Set("Authorization", "Bearer deterministic-secret") + + // When + lease, err := prepareNATCapability(request, &model.NAT{Common: model.Common{ID: 91}, ServerID: 81}) + + // Then + if err != nil { + t.Fatalf("default NAT capability hook returned error: %v", err) + } + if lease.active { + t.Fatal("default NAT capability hook unexpectedly activated") + } + if got := request.Header.Get(agentcompatcontract.IOStreamCapabilityHeader); got != "ordinary-data" { + t.Fatalf("default NAT capability hook changed header to %q", got) + } + if got := request.Header.Get("Authorization"); got != "" { + t.Fatalf("default NAT capability hook retained Authorization %q", got) + } +} diff --git a/cmd/dashboard/rpc/nat_capability_failures_agentcompat_test.go b/cmd/dashboard/rpc/nat_capability_failures_agentcompat_test.go new file mode 100644 index 00000000..822bafe8 --- /dev/null +++ b/cmd/dashboard/rpc/nat_capability_failures_agentcompat_test.go @@ -0,0 +1,245 @@ +//go:build agentcompat + +package rpc + +import ( + "context" + "io" + "net/http" + "net/url" + "sync" + "testing" + "time" + + "github.com/goccy/go-json" + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/agentcompatcontract" + "github.com/nezhahq/nezha/proto" + serviceRPC "github.com/nezhahq/nezha/service/rpc" + "github.com/nezhahq/nezha/service/singleton" + "github.com/stretchr/testify/require" +) + +func TestServeNATAgentCompatReleasesCapabilityBeforeServerDispatch(t *testing.T) { + fixture := newServeNATFixture(t) + capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 101) + originalServers := singleton.ServerShared + singleton.ServerShared = singleton.NewServerClass() + t.Cleanup(func() { singleton.ServerShared = originalServers }) + request := capabilityRequest(capability) + + ServeNAT(&serveNATResponseWriter{}, request, &model.NAT{Common: model.Common{ID: 101}, ServerID: fixture.server.ID, Host: "target.example"}) + + assertNATCapabilitySlotReleased(t, fixture.handler, access) + require.Equal(t, 0, fixture.handler.StreamCount()) + require.Empty(t, fixture.taskStream.sent) +} + +func TestServeNATAgentCompatTaskStreamUnavailableReleasesCapabilityBeforeActivity(t *testing.T) { + fixture := newServeNATFixture(t) + capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 105) + fixture.server.SetTaskStream(nil) + request := capabilityRequest(capability) + + ServeNAT(&serveNATResponseWriter{}, request, &model.NAT{Common: model.Common{ID: 105}, ServerID: fixture.server.ID, Host: "target.example"}) + + assertNATCapabilitySlotReleased(t, fixture.handler, access) + require.Equal(t, 0, fixture.handler.StreamCount()) + require.Empty(t, fixture.taskStream.sent) + if request.Header.Get(agentcompatcontract.IOStreamCapabilityHeader) != "" { + t.Fatal("capability header remained after task stream unavailable") + } +} + +func TestServeNATAgentCompatCreateQuotaFailurePreservesExistingStreams(t *testing.T) { + fixture := newServeNATFixture(t) + capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 106) + const existingStreamCount = 40 + for index := 0; index < existingStreamCount; index++ { + require.NoError(t, fixture.handler.CreateStreamWithPurpose("quota-existing-"+string(rune(index)), 0, fixture.server.ID, serviceRPC.PurposeNAT)) + } + t.Cleanup(func() { + for index := 0; index < existingStreamCount; index++ { + _ = fixture.handler.CloseStream("quota-existing-" + string(rune(index))) + } + }) + request := capabilityRequest(capability) + + ServeNAT(&serveNATResponseWriter{}, request, &model.NAT{Common: model.Common{ID: 106}, ServerID: fixture.server.ID, Host: "target.example"}) + + assertNATCapabilitySlotReleased(t, fixture.handler, access) + require.Equal(t, existingStreamCount, fixture.handler.StreamCount()) + require.Empty(t, fixture.taskStream.sent) +} + +func TestServeNATAgentCompatPublishCancellationClosesRegistryOwnedEndpoints(t *testing.T) { + fixture := newServeNATFixture(t) + capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 107) + connection := newServeNATConn() + agent := &serveNATAgent{readErr: io.EOF, writeDone: make(chan struct{})} + observerEntered := make(chan struct{}) + observerRelease := make(chan struct{}) + var observerOnce sync.Once + fixture.handler.SetAgentCompatCapabilityPublishObserverForTest(func() { + observerOnce.Do(func() { close(observerEntered) }) + <-observerRelease + }) + t.Cleanup(func() { fixture.handler.SetAgentCompatCapabilityPublishObserverForTest(nil) }) + fixture.taskStream.onSend = func(task *proto.Task) error { + var nat model.TaskNAT + require.NoError(t, json.Unmarshal([]byte(task.Data), &nat)) + return fixture.handler.AgentConnected(nat.StreamID, agent) + } + request := capabilityRequest(capability) + writer := &serveNATResponseWriter{conn: connection} + serveDone := make(chan struct{}) + go func() { + ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 107}, ServerID: fixture.server.ID, Host: "target.example"}) + close(serveDone) + }() + select { + case <-observerEntered: + case <-serveDone: + t.Fatal("ServeNAT returned before publish pause") + } + fixture.handler.CreateStreamWithPurpose("publish-replacement", 0, fixture.server.ID, serviceRPC.PurposeNAT) + require.NoError(t, fixture.handler.CancelAgentCompatIOStreamCapability(access)) + close(observerRelease) + completionContext, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + select { + case <-serveDone: + case <-completionContext.Done(): + t.Fatal("ServeNAT did not finish after publish cancellation") + } + assertCapabilityReleased(t, fixture.handler, access) + require.Equal(t, int32(1), connection.closeCount.Load()) + require.Equal(t, int32(1), agent.closeCount.Load()) + require.Equal(t, 1, fixture.handler.StreamCount()) + select { + case <-agent.writeDone: + t.Fatal("agent endpoint was used after publish cancellation") + default: + } +} + +func TestServeNATAgentCompatSendTaskFailureReleasesStreamAndCapabilityBeforeWrapper(t *testing.T) { + fixture := newServeNATFixture(t) + capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 102) + fixture.taskStream.sendErr = io.ErrClosedPipe + connection := newServeNATConn() + request := capabilityRequest(capability) + request.Body = &serveNATBody{} + writer := &serveNATResponseWriter{conn: connection} + + ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 102}, ServerID: fixture.server.ID, Host: "target.example"}) + + assertCapabilityReleased(t, fixture.handler, access) + require.Equal(t, 0, fixture.handler.StreamCount()) + require.Equal(t, int32(0), connection.closeCount.Load()) +} + +func TestServeNATAgentCompatWrapperFailureReleasesCapabilityWithoutSecondResponse(t *testing.T) { + fixture := newServeNATFixture(t) + capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 103) + request := capabilityRequest(capability) + writer := &nonHijackingResponseWriter{} + + ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 103}, ServerID: fixture.server.ID, Host: "target.example"}) + + assertCapabilityReleased(t, fixture.handler, access) + require.Equal(t, 0, fixture.handler.StreamCount()) + require.Equal(t, 0, writer.writes) +} + +func TestServeNATAgentCompatStartFailureCancelsPublishedCapabilityAndClosesEndpoints(t *testing.T) { + fixture := newServeNATFixture(t) + capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 104) + connection := newServeNATConn() + agent := &serveNATAgent{readErr: io.EOF, writeDone: make(chan struct{})} + fixture.taskStream.onSend = func(task *proto.Task) error { + var nat model.TaskNAT + if err := json.Unmarshal([]byte(task.Data), &nat); err != nil { + return err + } + return fixture.handler.AgentConnected(nat.StreamID, agent) + } + request := capabilityRequest(capability) + writer := &serveNATResponseWriter{conn: connection} + + ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 104}, ServerID: fixture.server.ID, Host: "target.example"}) + + assertCapabilityReleased(t, fixture.handler, access) + require.Equal(t, 0, fixture.handler.StreamCount()) + require.Equal(t, int32(1), connection.closeCount.Load()) + require.Equal(t, int32(1), agent.closeCount.Load()) +} + +func registerNATCapabilityForTest(t *testing.T, handler *serviceRPC.NezhaHandler, serverID, resourceID uint64) (string, serviceRPC.AgentCompatCapabilityAccess) { + t.Helper() + capability, err := handler.RegisterAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityRegistration{ + Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: resourceID, UserID: resourceID + 1}, Purpose: serviceRPC.AgentCompatCapabilityNAT, + TargetServerID: serverID, ResourceID: resourceID, ServerAccessAllowed: true, + }) + require.NoError(t, err) + parsed, err := serviceRPC.ParseAgentCompatIOStreamCapability(capability.String()) + require.NoError(t, err) + return capability.String(), serviceRPC.AgentCompatCapabilityAccess{ + Capability: parsed, Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: resourceID, UserID: resourceID + 1}, + Purpose: serviceRPC.AgentCompatCapabilityNAT, TargetServerID: serverID, ResourceID: resourceID, ServerAccessAllowed: true, + } +} + +func capabilityRequest(value string) *http.Request { + request := &http.Request{Method: http.MethodPost, URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"}, Header: make(http.Header), Body: &serveNATBody{}} + request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, value) + return request +} + +func assertCapabilityReleased(t *testing.T, handler *serviceRPC.NezhaHandler, access serviceRPC.AgentCompatCapabilityAccess) { + t.Helper() + _, _, err := handler.ConsumeAgentCompatNATCapabilityForProfile(access.Capability.String(), access.TargetServerID, access.ResourceID) + require.ErrorIs(t, err, serviceRPC.ErrAgentCompatCapabilityHidden) +} + +func assertNATCapabilitySlotReleased(t *testing.T, handler *serviceRPC.NezhaHandler, access serviceRPC.AgentCompatCapabilityAccess) { + t.Helper() + assertCapabilityReleased(t, handler, access) + type activeCapability struct { + capability serviceRPC.AgentCompatIOStreamCapability + resourceID uint64 + } + active := make([]activeCapability, 0, 16) + for index := uint64(0); index < 16; index++ { + capability, err := handler.RegisterAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityRegistration{ + Owner: access.Owner, Purpose: serviceRPC.AgentCompatCapabilityNAT, TargetServerID: access.TargetServerID, + ResourceID: access.ResourceID + index + 1000, ServerAccessAllowed: true, + }) + require.NoError(t, err) + active = append(active, activeCapability{capability: capability, resourceID: access.ResourceID + index + 1000}) + } + _, err := handler.RegisterAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityRegistration{ + Owner: access.Owner, Purpose: serviceRPC.AgentCompatCapabilityNAT, TargetServerID: access.TargetServerID, + ResourceID: access.ResourceID + 2000, ServerAccessAllowed: true, + }) + require.ErrorIs(t, err, serviceRPC.ErrAgentCompatCapabilityUnavailable) + for _, item := range active { + parsed, parseErr := serviceRPC.ParseAgentCompatIOStreamCapability(item.capability.String()) + require.NoError(t, parseErr) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(serviceRPC.AgentCompatCapabilityAccess{ + Capability: parsed, Owner: access.Owner, Purpose: serviceRPC.AgentCompatCapabilityNAT, + TargetServerID: access.TargetServerID, ResourceID: item.resourceID, ServerAccessAllowed: true, + })) + } +} + +type nonHijackingResponseWriter struct { + writes int +} + +func (writer *nonHijackingResponseWriter) Header() http.Header { return make(http.Header) } +func (writer *nonHijackingResponseWriter) Write(data []byte) (int, error) { + writer.writes += len(data) + return len(data), nil +} +func (*nonHijackingResponseWriter) WriteHeader(int) {} diff --git a/cmd/dashboard/rpc/nat_capability_flow_agentcompat_test.go b/cmd/dashboard/rpc/nat_capability_flow_agentcompat_test.go new file mode 100644 index 00000000..2ed5a003 --- /dev/null +++ b/cmd/dashboard/rpc/nat_capability_flow_agentcompat_test.go @@ -0,0 +1,151 @@ +//go:build agentcompat + +package rpc + +import ( + "bytes" + "context" + "io" + "net/http" + "net/url" + "strings" + "sync" + "testing" + "time" + + "github.com/goccy/go-json" + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/pkg/agentcompatcontract" + "github.com/nezhahq/nezha/proto" + serviceRPC "github.com/nezhahq/nezha/service/rpc" + "github.com/stretchr/testify/require" +) + +func TestServeNATAgentCompatPublishesExactStreamAfterRequestTransfer(t *testing.T) { + fixture := newServeNATFixture(t) + capability, err := fixture.handler.RegisterAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityRegistration{ + Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: 1, UserID: 2}, + Purpose: serviceRPC.AgentCompatCapabilityNAT, + TargetServerID: fixture.server.ID, + ResourceID: 91, + ServerAccessAllowed: true, + }) + require.NoError(t, err) + access, err := serviceRPC.ParseAgentCompatIOStreamCapability(capability.String()) + require.NoError(t, err) + connection := newServeNATConn() + writer := &serveNATResponseWriter{conn: connection} + agent := &orderedNATAgent{readStarted: make(chan struct{}), readRelease: make(chan struct{}), writeStarted: make(chan struct{})} + taskStreamIDReady := make(chan string, 1) + request := &http.Request{ + Method: http.MethodPost, + URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"}, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader("payload")), + } + request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability.String()) + request.Header.Set("Authorization", "Bearer deterministic-secret") + request.Header.Set("X-Ordinary-NAT", "ordinary") + fixture.taskStream.onSend = func(task *proto.Task) error { + var nat model.TaskNAT + if err := json.Unmarshal([]byte(task.Data), &nat); err != nil { + return err + } + if err := fixture.handler.AgentConnected(nat.StreamID, agent); err != nil { + return err + } + taskStreamIDReady <- nat.StreamID + return nil + } + serveDone := make(chan struct{}) + go func() { + ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 91}, ServerID: fixture.server.ID, Host: "target.example"}) + close(serveDone) + }() + streamID, err := fixture.handler.WaitAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityAccess{ + Capability: access, + Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: 1, UserID: 2}, + Purpose: serviceRPC.AgentCompatCapabilityNAT, + TargetServerID: fixture.server.ID, + ResourceID: 91, + ServerAccessAllowed: true, + }) + require.NoError(t, err) + taskStreamID := <-taskStreamIDReady + require.Equal(t, taskStreamID, streamID) + select { + case <-agent.readStarted: + case <-serveDone: + t.Fatal("ServeNAT returned before StartStream read") + } + select { + case <-agent.writeStarted: + case <-serveDone: + t.Fatal("ServeNAT returned before request reached agent") + } + close(agent.readRelease) + require.NoError(t, connection.Close()) + completionContext, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + select { + case <-serveDone: + case <-completionContext.Done(): + t.Fatal("ServeNAT did not complete after StartStream") + } + require.NoError(t, fixture.handler.CancelAgentCompatIOStreamCapability(serviceRPC.AgentCompatCapabilityAccess{ + Capability: access, + Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: 1, UserID: 2}, + Purpose: serviceRPC.AgentCompatCapabilityNAT, + TargetServerID: fixture.server.ID, + ResourceID: 91, + ServerAccessAllowed: true, + })) + require.NotContains(t, string(agent.writtenBytes()), capability.String()) + forwarded := string(agent.writtenBytes()) + for _, expected := range []string{"POST /nat HTTP/1.1", "Host: example.test", "X-Ordinary-Nat: ordinary", "payload"} { + require.Contains(t, forwarded, expected) + } + require.NotContains(t, forwarded, agentcompatcontract.IOStreamCapabilityHeader) + require.NotContains(t, forwarded, "Authorization:") + require.NotContains(t, forwarded, "deterministic-secret") +} + +type orderedNATAgent struct { + mu sync.Mutex + written bytes.Buffer + readStarted chan struct{} + readRelease chan struct{} + writeStarted chan struct{} + readOnce sync.Once + writeOnce sync.Once + closeCount int +} + +func (agent *orderedNATAgent) Read([]byte) (int, error) { + agent.readOnce.Do(func() { close(agent.readStarted) }) + <-agent.readRelease + return 0, io.EOF +} + +func (agent *orderedNATAgent) Write(data []byte) (int, error) { + agent.mu.Lock() + count, err := agent.written.Write(data) + agent.mu.Unlock() + agent.writeOnce.Do(func() { close(agent.writeStarted) }) + return count, err +} + +func (agent *orderedNATAgent) Close() error { + agent.mu.Lock() + defer agent.mu.Unlock() + agent.closeCount++ + return nil +} + +func (agent *orderedNATAgent) writtenBytes() []byte { + agent.mu.Lock() + defer agent.mu.Unlock() + return append([]byte(nil), agent.written.Bytes()...) +} + +var _ io.ReadWriteCloser = (*orderedNATAgent)(nil) diff --git a/cmd/dashboard/rpc/nat_test.go b/cmd/dashboard/rpc/nat_test.go new file mode 100644 index 00000000..ed5ccde0 --- /dev/null +++ b/cmd/dashboard/rpc/nat_test.go @@ -0,0 +1,93 @@ +package rpc + +import ( + "io" + "net/http" + "net/url" + "strings" + "testing" + "time" + + "github.com/goccy/go-json" + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/proto" +) + +func TestServeNATClosesHijackedResourcesWhenStreamDisappearsBeforeUserTransfer(t *testing.T) { + fixture := newServeNATFixture(t) + conn := newServeNATConn() + body := &serveNATBody{} + writer := &serveNATResponseWriter{conn: conn} + request := &http.Request{Method: http.MethodPost, URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"}, Header: make(http.Header), Body: body, ContentLength: 0} + request.Header.Set("X-NAT-Test", "failure") + fixture.taskStream.onSend = func(task *proto.Task) error { + var nat model.TaskNAT + if err := json.Unmarshal([]byte(task.Data), &nat); err != nil { + return err + } + return fixture.handler.CloseStream(nat.StreamID) + } + + ServeNAT(writer, request, &model.NAT{ServerID: fixture.server.ID, Host: "target.example"}) + + if got := body.closeCount.Load(); got == 0 { + t.Fatal("request body was not closed after failed transfer") + } + if got := conn.closeCount.Load(); got != 1 { + t.Fatalf("hijacked connection close count = %d, want 1", got) + } + select { + case <-conn.readDone: + case <-time.After(time.Second): + t.Fatal("hijacked connection read remained blocked after failed transfer") + } + if writer.status != 0 || writer.writes != 0 { + t.Fatalf("hijacked failure wrote HTTP response: status=%d writes=%d", writer.status, writer.writes) + } +} + +func TestServeNATTransfersRequestAndRegistryOwnsSuccessfulCleanup(t *testing.T) { + fixture := newServeNATFixture(t) + conn := newServeNATConn() + body := &serveNATBody{} + agent := &serveNATAgent{readErr: io.EOF, writeDone: make(chan struct{})} + writer := &serveNATResponseWriter{conn: conn} + request := &http.Request{Method: http.MethodPost, URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"}, Header: make(http.Header), Body: body, ContentLength: 0} + request.Header.Set("X-NAT-Test", "success") + fixture.taskStream.onSend = func(task *proto.Task) error { + var nat model.TaskNAT + if err := json.Unmarshal([]byte(task.Data), &nat); err != nil { + return err + } + return fixture.handler.AgentConnected(nat.StreamID, agent) + } + + serveDone := make(chan struct{}) + go func() { + ServeNAT(writer, request, &model.NAT{ServerID: fixture.server.ID, Host: "target.example"}) + close(serveDone) + }() + + select { + case <-agent.writeDone: + case <-time.After(time.Second): + t.Fatal("agent did not receive the transferred NAT request") + } + select { + case <-serveDone: + case <-time.After(time.Second): + t.Fatal("ServeNAT did not finish after the successful stream was closed") + } + if got := body.closeCount.Load(); got == 0 { + t.Fatal("registry-owned request body was not closed") + } + if got := conn.closeCount.Load(); got != 1 { + t.Fatalf("registry-owned hijacked connection close count = %d, want 1", got) + } + if got := agent.closeCount.Load(); got != 1 { + t.Fatalf("registry-owned agent close count = %d, want 1", got) + } + if got := strings.ToLower(string(agent.writtenBytes())); !strings.Contains(got, "post /nat http/1.1") || !strings.Contains(got, "x-nat-test: success") { + t.Fatalf("agent received incomplete NAT request: %q", got) + } +} diff --git a/cmd/dashboard/rpc/nat_test_support_test.go b/cmd/dashboard/rpc/nat_test_support_test.go new file mode 100644 index 00000000..d4a62c17 --- /dev/null +++ b/cmd/dashboard/rpc/nat_test_support_test.go @@ -0,0 +1,171 @@ +package rpc + +import ( + "bufio" + "bytes" + "context" + "io" + "net" + "net/http" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/nezhahq/nezha/model" + "github.com/nezhahq/nezha/proto" + rpcService "github.com/nezhahq/nezha/service/rpc" + "github.com/nezhahq/nezha/service/singleton" + "github.com/stretchr/testify/require" + "google.golang.org/grpc/metadata" + "gorm.io/driver/sqlite" + "gorm.io/gorm" +) + +type serveNATFixture struct { + handler *rpcService.NezhaHandler + server *model.Server + taskStream *serveNATTaskStream +} + +func newServeNATFixture(t *testing.T) serveNATFixture { + t.Helper() + originalDB, originalServerShared, originalHandler := singleton.DB, singleton.ServerShared, rpcService.NezhaHandlerSingleton + db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + require.NoError(t, err) + require.NoError(t, db.AutoMigrate(model.Server{})) + server := &model.Server{Common: model.Common{ID: 7}, UUID: "serve-nat-test", Name: "serve-nat-test"} + require.NoError(t, db.Create(server).Error) + singleton.DB = db + singleton.ServerShared = singleton.NewServerClass() + handler := rpcService.NewNezhaHandler() + rpcService.NezhaHandlerSingleton = handler + taskStream := &serveNATTaskStream{} + server, ok := singleton.ServerShared.Get(server.ID) + require.True(t, ok) + server.SetTaskStream(taskStream) + taskStream.server = server + t.Cleanup(func() { + rpcService.NezhaHandlerSingleton, singleton.ServerShared, singleton.DB = originalHandler, originalServerShared, originalDB + if dbSQL, dbErr := db.DB(); dbErr == nil { + _ = dbSQL.Close() + } + }) + return serveNATFixture{handler: handler, server: server, taskStream: taskStream} +} + +type serveNATTaskStream struct { + server *model.Server + onSend func(*proto.Task) error + sendErr error + sent []*proto.Task +} + +func (stream *serveNATTaskStream) Send(task *proto.Task) error { + stream.sent = append(stream.sent, task) + if stream.sendErr != nil { + return stream.sendErr + } + if stream.onSend != nil { + return stream.onSend(task) + } + return nil +} +func (*serveNATTaskStream) Recv() (*proto.TaskResult, error) { return nil, io.EOF } +func (*serveNATTaskStream) SetHeader(metadata.MD) error { return nil } +func (*serveNATTaskStream) SendHeader(metadata.MD) error { return nil } +func (*serveNATTaskStream) SetTrailer(metadata.MD) {} +func (*serveNATTaskStream) Context() context.Context { return context.Background() } +func (*serveNATTaskStream) SendMsg(any) error { return nil } +func (*serveNATTaskStream) RecvMsg(any) error { return io.EOF } + +type serveNATResponseWriter struct { + conn *serveNATConn + header http.Header + status, writes int + body string +} + +func (writer *serveNATResponseWriter) Header() http.Header { + if writer.header == nil { + writer.header = make(http.Header) + } + return writer.header +} +func (writer *serveNATResponseWriter) Write(data []byte) (int, error) { + writer.writes += len(data) + writer.body += string(data) + return len(data), nil +} +func (writer *serveNATResponseWriter) WriteHeader(status int) { writer.status = status } +func (writer *serveNATResponseWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { + go func() { _, _ = writer.conn.Read(make([]byte, 1)) }() + return writer.conn, bufio.NewReadWriter(bufio.NewReader(writer.conn), bufio.NewWriter(writer.conn)), nil +} + +type serveNATBody struct{ closeCount atomic.Int32 } + +func (*serveNATBody) Read([]byte) (int, error) { return 0, io.EOF } +func (body *serveNATBody) Close() error { body.closeCount.Add(1); return nil } + +type serveNATConn struct { + closed, readDone chan struct{} + readDoneOnce, closeOnce sync.Once + closeCount atomic.Int32 +} + +func newServeNATConn() *serveNATConn { + return &serveNATConn{closed: make(chan struct{}), readDone: make(chan struct{})} +} +func (conn *serveNATConn) Read([]byte) (int, error) { + <-conn.closed + conn.readDoneOnce.Do(func() { close(conn.readDone) }) + return 0, io.EOF +} +func (*serveNATConn) Write(data []byte) (int, error) { return len(data), nil } +func (conn *serveNATConn) Close() error { + conn.closeCount.Add(1) + conn.closeOnce.Do(func() { close(conn.closed) }) + return nil +} +func (*serveNATConn) LocalAddr() net.Addr { return serveNATAddr("local") } +func (*serveNATConn) RemoteAddr() net.Addr { return serveNATAddr("remote") } +func (*serveNATConn) SetDeadline(time.Time) error { return nil } +func (*serveNATConn) SetReadDeadline(time.Time) error { return nil } +func (*serveNATConn) SetWriteDeadline(time.Time) error { return nil } + +type serveNATAddr string + +func (addr serveNATAddr) Network() string { return "test" } +func (addr serveNATAddr) String() string { return string(addr) } + +type serveNATAgent struct { + mu sync.Mutex + written bytes.Buffer + readErr error + writeDone chan struct{} + writeOnce sync.Once + closeCount atomic.Int32 +} + +func (agent *serveNATAgent) Read([]byte) (int, error) { return 0, agent.readErr } +func (agent *serveNATAgent) Write(data []byte) (int, error) { + agent.mu.Lock() + defer agent.mu.Unlock() + count, err := agent.written.Write(data) + if agent.writeDone != nil { + agent.writeOnce.Do(func() { close(agent.writeDone) }) + } + return count, err +} +func (agent *serveNATAgent) Close() error { agent.closeCount.Add(1); return nil } +func (agent *serveNATAgent) writtenBytes() []byte { + agent.mu.Lock() + defer agent.mu.Unlock() + return append([]byte(nil), agent.written.Bytes()...) +} + +var _ proto.NezhaService_RequestTaskServer = (*serveNATTaskStream)(nil) +var _ http.Hijacker = (*serveNATResponseWriter)(nil) +var _ net.Conn = (*serveNATConn)(nil) +var _ io.ReadWriteCloser = (*serveNATAgent)(nil) diff --git a/cmd/dashboard/rpc/rpc.go b/cmd/dashboard/rpc/rpc.go index ba9ddeb0..ff95df9b 100644 --- a/cmd/dashboard/rpc/rpc.go +++ b/cmd/dashboard/rpc/rpc.go @@ -1,27 +1,23 @@ package rpc import ( - "context" - "errors" - "fmt" - "log" - "net/http" - "net/netip" - "time" + "net" - "github.com/goccy/go-json" "google.golang.org/grpc" - "google.golang.org/grpc/metadata" - "google.golang.org/grpc/peer" - "github.com/hashicorp/go-uuid" - "github.com/nezhahq/nezha/model" - "github.com/nezhahq/nezha/pkg/utils" "github.com/nezhahq/nezha/proto" rpcService "github.com/nezhahq/nezha/service/rpc" "github.com/nezhahq/nezha/service/singleton" ) +func SetReceiptGateListener(listener net.Listener) { + rpcService.SetReceiptGateListener(listener) +} + +func CloseReceiptGate() { + rpcService.CloseReceiptGate() +} + // 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). @@ -46,228 +42,3 @@ func ServeRPC() *grpc.Server { proto.RegisterNezhaServiceServer(server, rpcService.NezhaHandlerSingleton) return server } - -func ctxWithRealIP(ctx context.Context) (context.Context, error) { - var ip, connectingIp string - p, ok := peer.FromContext(ctx) - if ok { - addrPort, err := netip.ParseAddrPort(p.Addr.String()) - if err == nil { - connectingIp = addrPort.Addr().String() - } - } - ctx = context.WithValue(ctx, model.CtxKeyConnectingIP{}, connectingIp) - - if singleton.Conf.AgentRealIPHeader == "" { - return ctx, nil - } - - if singleton.Conf.AgentRealIPHeader == model.ConfigUsePeerIP { - if connectingIp == "" { - return ctx, fmt.Errorf("connecting ip not found") - } - // Peer-IP mode: peer IP is the real IP. Leaving ip="" makes - // CheckIP/BlockIP short-circuit on empty IP, disabling the WAF. - ip = connectingIp - } else { - vals := metadata.ValueFromIncomingContext(ctx, singleton.Conf.AgentRealIPHeader) - if len(vals) == 0 { - return ctx, fmt.Errorf("real ip header not found") - } - var err error - ip, err = utils.GetIPFromHeader(vals[0]) - if err != nil { - return ctx, err - } - } - - if singleton.Conf.Debug { - log.Printf("NEZHA>> gRPC Agent Real IP: %s, connecting IP: %s\n", ip, connectingIp) - } - - return context.WithValue(ctx, model.CtxKeyRealIP{}, ip), nil -} - -func waf(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { - realip, _ := ctx.Value(model.CtxKeyRealIP{}).(string) - if err := model.CheckIP(singleton.DB, realip); err != nil { - return nil, err - } - return handler(ctx, req) -} - -func getRealIp(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { - ctx, err := ctxWithRealIP(ctx) - if err != nil { - return nil, err - } - return handler(ctx, req) -} - -// realIPServerStream overrides Context() so stream handlers and -// authHandler.check observe the resolved real IP, like the unary path. -type realIPServerStream struct { - grpc.ServerStream - ctx context.Context -} - -func (s *realIPServerStream) Context() context.Context { return s.ctx } - -func getRealIpStream(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error { - ctx, err := ctxWithRealIP(ss.Context()) - if err != nil { - return err - } - return handler(srv, &realIPServerStream{ServerStream: ss, ctx: ctx}) -} - -func wafStream(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error { - realip, _ := ss.Context().Value(model.CtxKeyRealIP{}).(string) - if err := model.CheckIP(singleton.DB, realip); err != nil { - return err - } - return handler(srv, ss) -} - -func DispatchTask(serviceSentinelDispatchBus <-chan *model.Service) { - for task := range serviceSentinelDispatchBus { - if task == nil { - continue - } - - switch task.Cover { - case model.ServiceCoverIgnoreAll: - for id, enabled := range task.SkipServers { - if !enabled { - continue - } - - server, _ := singleton.ServerShared.Get(id) - if server == nil { - continue - } - if !canSendTaskToServer(task, server) { - continue - } - // 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: - // 快照后逐个 SendTask,不在 ServerShared 的 listMu.RLock 内做阻塞 - // gRPC:否则一个卡死 agent 会拖死需要写锁的 server 生命周期操作。 - for id, server := range singleton.ServerShared.GetList() { - if server == nil || task.SkipServers[id] { - continue - } - if !canSendTaskToServer(task, server) { - continue - } - if err := server.SendTask(task.PB()); err != nil && - !errors.Is(err, model.ErrTaskStreamOffline) { - log.Printf("NEZHA>> DispatchTask send error (server=%d): %v", id, err) - } - } - } - } -} - -func DispatchKeepalive() { - singleton.CronShared.AddFunc("@every 20s", func() { - list := singleton.ServerShared.GetSortedList() - for _, s := range list { - if s == nil { - continue - } - if err := s.SendTask(&proto.Task{Type: model.TaskTypeKeepalive}); err != nil && - !errors.Is(err, model.ErrTaskStreamOffline) { - log.Printf("NEZHA>> Keepalive send error (server=%d): %v", s.ID, err) - } - } - }) -} - -func ServeNAT(w http.ResponseWriter, r *http.Request, natConfig *model.NAT) { - server, _ := singleton.ServerShared.Get(natConfig.ServerID) - if server == nil { - w.WriteHeader(http.StatusServiceUnavailable) - w.Write([]byte("server not found or not connected")) - return - } - if server.GetTaskStream() == nil { - w.WriteHeader(http.StatusServiceUnavailable) - w.Write([]byte("server not found or not connected")) - return - } - - streamId, err := uuid.GenerateUUID() - if err != nil { - w.WriteHeader(http.StatusServiceUnavailable) - w.Write(fmt.Appendf(nil, "stream id error: %v", err)) - return - } - - // NAT streams are anonymous HTTP-facing tunnels; they are NOT reachable - // via /ws/terminal or /ws/file (which check stream ownership), so the - // creator user ID does not need to identify a real user. The targetServerID - // IS required though — the receiving agent must prove it is the server the - // NAT config addressed, otherwise any agent that snoops the streamId can - // answer NAT traffic on behalf of an unrelated host. - if err := rpcService.NezhaHandlerSingleton.CreateStream(streamId, 0, server.ID); err != nil { - w.WriteHeader(http.StatusTooManyRequests) - w.Write(fmt.Appendf(nil, "stream limit: %v", err)) - return - } - defer rpcService.NezhaHandlerSingleton.CloseStream(streamId) - - taskData, err := json.Marshal(model.TaskNAT{ - StreamID: streamId, - Host: natConfig.Host, - }) - if err != nil { - w.WriteHeader(http.StatusServiceUnavailable) - w.Write(fmt.Appendf(nil, "task data error: %v", err)) - return - } - - if err := server.SendTask(&proto.Task{ - Type: model.TaskTypeNAT, - Data: string(taskData), - }); err != nil { - w.WriteHeader(http.StatusServiceUnavailable) - w.Write(fmt.Appendf(nil, "send task error: %v", err)) - return - } - - wWrapped, err := utils.NewRequestWrapper(r, w) - if err != nil { - w.WriteHeader(http.StatusServiceUnavailable) - w.Write(fmt.Appendf(nil, "request wrapper error: %v", err)) - return - } - - if err := rpcService.NezhaHandlerSingleton.UserConnected(streamId, wWrapped); err != nil { - w.WriteHeader(http.StatusServiceUnavailable) - w.Write(fmt.Appendf(nil, "user connected error: %v", err)) - return - } - - rpcService.NezhaHandlerSingleton.StartStream(streamId, time.Second*10) -} - -func canSendTaskToServer(task *model.Service, server *model.Server) bool { - var role model.Role - singleton.UserLock.RLock() - if u, ok := singleton.UserInfoMap[task.UserID]; !ok { - role = model.RoleMember - } else { - role = u.Role - } - singleton.UserLock.RUnlock() - - return task.UserID == server.GetUserID() || role.IsAdmin() -}