feat(agentcompat): expose dashboard capability routes

Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
naiba
2026-07-20 04:30:50 +00:00
co-authored by naiba/CloudCode
parent 0c6eb416a5
commit 2640b86d3a
42 changed files with 4083 additions and 271 deletions
@@ -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
}
@@ -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)
}
@@ -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
}
@@ -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)
}
@@ -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))
}
@@ -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
}
@@ -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)
}
@@ -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
}
@@ -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)
}
@@ -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
}
@@ -0,0 +1,7 @@
//go:build !agentcompat
package controller
import "github.com/gin-gonic/gin"
func registerAgentcompatRoutes(*gin.Engine) {}
@@ -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
}
@@ -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)
}
}
@@ -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")
}
}
@@ -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}
}
@@ -0,0 +1,7 @@
//go:build agentcompat && !linux
package controller
import "github.com/gin-gonic/gin"
func registerAgentcompatSQLiteHoldRoutes(*gin.Engine, gin.HandlerFunc) {}
@@ -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)
}
+1
View File
@@ -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)
+15 -15
View File
@@ -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)
@@ -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))
}
+15 -15
View File
@@ -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)
@@ -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
}
@@ -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))
})
}
}
@@ -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
}
@@ -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
}
@@ -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
}
@@ -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}
}
@@ -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())
}
@@ -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))))
}
}
+8 -3
View File
@@ -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