mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 09:40:12 +00:00
feat(agentcompat): expose dashboard capability routes
Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
@@ -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)
|
||||||
|
}
|
||||||
@@ -69,6 +69,7 @@ func routers(r *gin.Engine, frontendDist fs.FS) {
|
|||||||
r.DELETE("/mcp", mcpOriginGuard(), mcpMethodNotAllowed)
|
r.DELETE("/mcp", mcpOriginGuard(), mcpMethodNotAllowed)
|
||||||
r.GET("/mcp/download/:token", mcpOriginGuard(), transferDownloadHandler)
|
r.GET("/mcp/download/:token", mcpOriginGuard(), transferDownloadHandler)
|
||||||
r.POST("/mcp/upload/:token", mcpOriginGuard(), transferUploadHandler)
|
r.POST("/mcp/upload/:token", mcpOriginGuard(), transferUploadHandler)
|
||||||
|
registerAgentcompatRoutes(r)
|
||||||
|
|
||||||
api := r.Group("api/v1")
|
api := r.Group("api/v1")
|
||||||
api.POST("/login", authMiddleware.LoginHandler)
|
api.POST("/login", authMiddleware.LoginHandler)
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ import (
|
|||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/goccy/go-json"
|
"github.com/goccy/go-json"
|
||||||
"github.com/gorilla/websocket"
|
|
||||||
"github.com/hashicorp/go-uuid"
|
"github.com/hashicorp/go-uuid"
|
||||||
|
|
||||||
"github.com/nezhahq/nezha/model"
|
"github.com/nezhahq/nezha/model"
|
||||||
@@ -26,6 +25,7 @@ import (
|
|||||||
// @Success 200 {object} model.CreateFMResponse
|
// @Success 200 {object} model.CreateFMResponse
|
||||||
// @Router /file [post]
|
// @Router /file [post]
|
||||||
func createFM(c *gin.Context) (*model.CreateFMResponse, error) {
|
func createFM(c *gin.Context) (*model.CreateFMResponse, error) {
|
||||||
|
prepareAgentcompatCapabilityHeader(c)
|
||||||
idStr := c.Query("id")
|
idStr := c.Query("id")
|
||||||
id, err := strconv.ParseUint(idStr, 10, 64)
|
id, err := strconv.ParseUint(idStr, 10, 64)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -49,17 +49,24 @@ func createFM(c *gin.Context) (*model.CreateFMResponse, error) {
|
|||||||
return nil, err
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
fmData, _ := json.Marshal(&model.TaskFM{
|
fmData, err := json.Marshal(&model.TaskFM{
|
||||||
StreamID: streamId,
|
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{
|
if err := server.SendTask(&proto.Task{
|
||||||
Type: model.TaskTypeFM,
|
Type: model.TaskTypeFM,
|
||||||
Data: string(fmData),
|
Data: string(fmData),
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
|
cleanup()
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -92,21 +99,14 @@ func fmStream(c *gin.Context) (any, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, newWsError("%v", err)
|
return nil, newWsError("%v", err)
|
||||||
}
|
}
|
||||||
defer wsConn.Close()
|
|
||||||
conn := websocketx.NewConn(wsConn)
|
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()
|
defer deregisterPAT()
|
||||||
|
// Join the ping worker before PAT and WebSocket cleanup can close its writer.
|
||||||
go func() {
|
defer stopPing()
|
||||||
// PING 保活
|
|
||||||
for {
|
|
||||||
if err = conn.WriteMessage(websocket.PingMessage, []byte{}); err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
time.Sleep(time.Second * 10)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
if err = rpc.NezhaHandlerSingleton.UserConnected(streamId, conn); err != nil {
|
if err = rpc.NezhaHandlerSingleton.UserConnected(streamId, conn); err != nil {
|
||||||
return nil, newWsError("%v", err)
|
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))
|
||||||
|
}
|
||||||
@@ -5,7 +5,6 @@ import (
|
|||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/goccy/go-json"
|
"github.com/goccy/go-json"
|
||||||
"github.com/gorilla/websocket"
|
|
||||||
"github.com/hashicorp/go-uuid"
|
"github.com/hashicorp/go-uuid"
|
||||||
|
|
||||||
"github.com/nezhahq/nezha/model"
|
"github.com/nezhahq/nezha/model"
|
||||||
@@ -25,6 +24,7 @@ import (
|
|||||||
// @Success 200 {object} model.CreateTerminalResponse
|
// @Success 200 {object} model.CreateTerminalResponse
|
||||||
// @Router /terminal [post]
|
// @Router /terminal [post]
|
||||||
func createTerminal(c *gin.Context) (*model.CreateTerminalResponse, error) {
|
func createTerminal(c *gin.Context) (*model.CreateTerminalResponse, error) {
|
||||||
|
prepareAgentcompatCapabilityHeader(c)
|
||||||
var createTerminalReq model.TerminalForm
|
var createTerminalReq model.TerminalForm
|
||||||
if err := c.ShouldBind(&createTerminalReq); err != nil {
|
if err := c.ShouldBind(&createTerminalReq); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
@@ -47,17 +47,24 @@ func createTerminal(c *gin.Context) (*model.CreateTerminalResponse, error) {
|
|||||||
return nil, err
|
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
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
terminalData, _ := json.Marshal(&model.TerminalTask{
|
terminalData, err := json.Marshal(&model.TerminalTask{
|
||||||
StreamID: streamId,
|
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{
|
if err := server.SendTask(&proto.Task{
|
||||||
Type: model.TaskTypeTerminalGRPC,
|
Type: model.TaskTypeTerminalGRPC,
|
||||||
Data: string(terminalData),
|
Data: string(terminalData),
|
||||||
}); err != nil {
|
}); err != nil {
|
||||||
|
cleanup()
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -93,21 +100,14 @@ func terminalStream(c *gin.Context) (any, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, newWsError("%v", err)
|
return nil, newWsError("%v", err)
|
||||||
}
|
}
|
||||||
defer wsConn.Close()
|
|
||||||
conn := websocketx.NewConn(wsConn)
|
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()
|
defer deregisterPAT()
|
||||||
|
// Join the ping worker before PAT and WebSocket cleanup can close its writer.
|
||||||
go func() {
|
defer stopPing()
|
||||||
// PING 保活
|
|
||||||
for {
|
|
||||||
if err = conn.WriteMessage(websocket.PingMessage, []byte{}); err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
time.Sleep(time.Second * 10)
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
if err = rpc.NezhaHandlerSingleton.UserConnected(streamId, conn); err != nil {
|
if err = rpc.NezhaHandlerSingleton.UserConnected(streamId, conn); err != nil {
|
||||||
return nil, newWsError("%v", err)
|
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))))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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 {
|
func filterServersForViewer(servers []*model.Server, viewerUserID uint64, viewerIsAdmin bool, withPublicNote bool, pat model.APITokenAccessor) []model.StreamServer {
|
||||||
out := make([]model.StreamServer, 0, len(servers))
|
out := make([]model.StreamServer, 0, len(servers))
|
||||||
for _, server := range servers {
|
for _, server := range servers {
|
||||||
|
runtime := server.RuntimeSnapshot()
|
||||||
if pat != nil && !pat.CanAccessServer(server.ID) {
|
if pat != nil && !pat.CanAccessServer(server.ID) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -236,15 +237,19 @@ func filterServersForViewer(servers []*model.Server, viewerUserID uint64, viewer
|
|||||||
if server.GeoIP != nil {
|
if server.GeoIP != nil {
|
||||||
countryCode = server.GeoIP.CountryCode
|
countryCode = server.GeoIP.CountryCode
|
||||||
}
|
}
|
||||||
|
publicHost := runtime.Host
|
||||||
|
if publicHost != nil && !isOwnerOrAdmin {
|
||||||
|
publicHost = publicHost.Filter()
|
||||||
|
}
|
||||||
out = append(out, model.StreamServer{
|
out = append(out, model.StreamServer{
|
||||||
ID: server.ID,
|
ID: server.ID,
|
||||||
Name: server.Name,
|
Name: server.Name,
|
||||||
PublicNote: utils.IfOr(withPublicNote, server.PublicNote, ""),
|
PublicNote: utils.IfOr(withPublicNote, server.PublicNote, ""),
|
||||||
DisplayIndex: server.DisplayIndex,
|
DisplayIndex: server.DisplayIndex,
|
||||||
Host: utils.IfOr(isOwnerOrAdmin, server.Host, server.Host.Filter()),
|
Host: publicHost,
|
||||||
State: server.State,
|
State: runtime.State,
|
||||||
CountryCode: countryCode,
|
CountryCode: countryCode,
|
||||||
LastActive: server.LastActive,
|
LastActive: runtime.LastActive,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
return out
|
return out
|
||||||
|
|||||||
@@ -0,0 +1,94 @@
|
|||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"log"
|
||||||
|
"net/netip"
|
||||||
|
|
||||||
|
"google.golang.org/grpc"
|
||||||
|
"google.golang.org/grpc/metadata"
|
||||||
|
"google.golang.org/grpc/peer"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/utils"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func ctxWithRealIP(ctx context.Context) (context.Context, error) {
|
||||||
|
var ip, connectingIp string
|
||||||
|
p, ok := peer.FromContext(ctx)
|
||||||
|
if ok {
|
||||||
|
addrPort, err := netip.ParseAddrPort(p.Addr.String())
|
||||||
|
if err == nil {
|
||||||
|
connectingIp = addrPort.Addr().String()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
ctx = context.WithValue(ctx, model.CtxKeyConnectingIP{}, connectingIp)
|
||||||
|
|
||||||
|
if singleton.Conf.AgentRealIPHeader == "" {
|
||||||
|
return ctx, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if singleton.Conf.AgentRealIPHeader == model.ConfigUsePeerIP {
|
||||||
|
if connectingIp == "" {
|
||||||
|
return ctx, fmt.Errorf("connecting ip not found")
|
||||||
|
}
|
||||||
|
ip = connectingIp
|
||||||
|
} else {
|
||||||
|
vals := metadata.ValueFromIncomingContext(ctx, singleton.Conf.AgentRealIPHeader)
|
||||||
|
if len(vals) == 0 {
|
||||||
|
return ctx, fmt.Errorf("real ip header not found")
|
||||||
|
}
|
||||||
|
var err error
|
||||||
|
ip, err = utils.GetIPFromHeader(vals[0])
|
||||||
|
if err != nil {
|
||||||
|
return ctx, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if singleton.Conf.Debug {
|
||||||
|
log.Printf("NEZHA>> gRPC Agent Real IP: %s, connecting IP: %s\n", ip, connectingIp)
|
||||||
|
}
|
||||||
|
|
||||||
|
return context.WithValue(ctx, model.CtxKeyRealIP{}, ip), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func waf(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) {
|
||||||
|
realip, _ := ctx.Value(model.CtxKeyRealIP{}).(string)
|
||||||
|
if err := model.CheckIP(singleton.DB, realip); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return handler(ctx, req)
|
||||||
|
}
|
||||||
|
|
||||||
|
func getRealIp(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) {
|
||||||
|
ctx, err := ctxWithRealIP(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return handler(ctx, req)
|
||||||
|
}
|
||||||
|
|
||||||
|
type realIPServerStream struct {
|
||||||
|
grpc.ServerStream
|
||||||
|
ctx context.Context
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *realIPServerStream) Context() context.Context { return s.ctx }
|
||||||
|
|
||||||
|
func getRealIpStream(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
|
||||||
|
ctx, err := ctxWithRealIP(ss.Context())
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return handler(srv, &realIPServerStream{ServerStream: ss, ctx: ctx})
|
||||||
|
}
|
||||||
|
|
||||||
|
func wafStream(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
|
||||||
|
realip, _ := ss.Context().Value(model.CtxKeyRealIP{}).(string)
|
||||||
|
if err := model.CheckIP(singleton.DB, realip); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return handler(srv, ss)
|
||||||
|
}
|
||||||
@@ -0,0 +1,115 @@
|
|||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"log"
|
||||||
|
"net/http"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/goccy/go-json"
|
||||||
|
|
||||||
|
"github.com/hashicorp/go-uuid"
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/utils"
|
||||||
|
"github.com/nezhahq/nezha/proto"
|
||||||
|
serviceRPC "github.com/nezhahq/nezha/service/rpc"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
)
|
||||||
|
|
||||||
|
func ServeNAT(w http.ResponseWriter, r *http.Request, natConfig *model.NAT) {
|
||||||
|
capabilityLease, capabilityErr := prepareNATCapability(r, natConfig)
|
||||||
|
if capabilityErr != nil {
|
||||||
|
w.WriteHeader(http.StatusServiceUnavailable)
|
||||||
|
w.Write([]byte("NAT capability unavailable"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
handler := serviceRPC.NezhaHandlerSingleton
|
||||||
|
streamId := ""
|
||||||
|
legacyStreamOwned := false
|
||||||
|
relayCompleted := false
|
||||||
|
defer func() {
|
||||||
|
if capabilityLease.active {
|
||||||
|
if !relayCompleted {
|
||||||
|
capabilityLease.cleanup(handler)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if legacyStreamOwned {
|
||||||
|
_ = handler.CloseStream(streamId)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
server, _ := singleton.ServerShared.Get(natConfig.ServerID)
|
||||||
|
if server == nil {
|
||||||
|
w.WriteHeader(http.StatusServiceUnavailable)
|
||||||
|
w.Write([]byte("server not found or not connected"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if server.GetTaskStream() == nil {
|
||||||
|
w.WriteHeader(http.StatusServiceUnavailable)
|
||||||
|
w.Write([]byte("server not found or not connected"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
streamId, err := uuid.GenerateUUID()
|
||||||
|
if err != nil {
|
||||||
|
w.WriteHeader(http.StatusServiceUnavailable)
|
||||||
|
w.Write(fmt.Appendf(nil, "stream id error: %v", err))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if capabilityLease.active {
|
||||||
|
capabilityLease.streamLease, err = handler.CreateAgentCompatNATStream(capabilityLease.handle, streamId)
|
||||||
|
} else {
|
||||||
|
err = handler.CreateStream(streamId, 0, server.ID)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
w.WriteHeader(http.StatusTooManyRequests)
|
||||||
|
w.Write(fmt.Appendf(nil, "stream limit: %v", err))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
legacyStreamOwned = !capabilityLease.active
|
||||||
|
taskData, err := json.Marshal(model.TaskNAT{StreamID: streamId, Host: natConfig.Host})
|
||||||
|
if err != nil {
|
||||||
|
w.WriteHeader(http.StatusServiceUnavailable)
|
||||||
|
w.Write(fmt.Appendf(nil, "task data error: %v", err))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := server.SendTask(&proto.Task{Type: model.TaskTypeNAT, Data: string(taskData)}); err != nil {
|
||||||
|
w.WriteHeader(http.StatusServiceUnavailable)
|
||||||
|
w.Write(fmt.Appendf(nil, "send task error: %v", err))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
// Authorization authenticates Dashboard access before NAT ingress; it must not become an origin credential for the configured NAT backend.
|
||||||
|
wWrapped, err := utils.NewRequestWrapper(r, w)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := handler.UserConnected(streamId, wWrapped); err != nil {
|
||||||
|
if closeErr := wWrapped.Close(); closeErr != nil {
|
||||||
|
log.Printf("NEZHA>> NAT request wrapper close error after user connection failure: %v", closeErr)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if capabilityLease.active {
|
||||||
|
if err := handler.PublishAgentCompatNATStream(capabilityLease.handle, serviceRPC.AgentCompatNATPublication{
|
||||||
|
Purpose: serviceRPC.AgentCompatCapabilityNAT,
|
||||||
|
TargetServerID: server.ID,
|
||||||
|
ResourceID: natConfig.ID,
|
||||||
|
StreamID: streamId,
|
||||||
|
}); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if capabilityLease.active {
|
||||||
|
capabilityLease.publicationOwned, err = handler.StartAgentCompatNATStream(capabilityLease.handle, time.Second*10)
|
||||||
|
if err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
} else if err := handler.StartStream(streamId, time.Second*10); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
relayCompleted = true
|
||||||
|
}
|
||||||
@@ -0,0 +1,48 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||||
|
serviceRPC "github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
type natCapabilityLease struct {
|
||||||
|
active bool
|
||||||
|
publicationOwned bool
|
||||||
|
streamLease *serviceRPC.AgentCompatNATStreamLease
|
||||||
|
access serviceRPC.AgentCompatCapabilityAccess
|
||||||
|
handle serviceRPC.AgentCompatNATPublishHandle
|
||||||
|
}
|
||||||
|
|
||||||
|
func prepareNATCapability(request *http.Request, natConfig *model.NAT) (natCapabilityLease, error) {
|
||||||
|
values := request.Header.Values(agentcompatcontract.IOStreamCapabilityHeader)
|
||||||
|
if len(values) == 0 {
|
||||||
|
request.Header.Del("Authorization")
|
||||||
|
return natCapabilityLease{}, nil
|
||||||
|
}
|
||||||
|
request.Header.Del(agentcompatcontract.IOStreamCapabilityHeader)
|
||||||
|
if len(values) != 1 || values[0] == "" {
|
||||||
|
request.Header.Del("Authorization")
|
||||||
|
return natCapabilityLease{}, errors.New("invalid NAT capability")
|
||||||
|
}
|
||||||
|
access, handle, err := serviceRPC.NezhaHandlerSingleton.ConsumeAgentCompatNATCapabilityForProfile(values[0], natConfig.ServerID, natConfig.ID)
|
||||||
|
request.Header.Del("Authorization")
|
||||||
|
if err != nil {
|
||||||
|
return natCapabilityLease{}, errors.New("invalid NAT capability")
|
||||||
|
}
|
||||||
|
return natCapabilityLease{active: true, access: access, handle: handle}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (lease natCapabilityLease) cleanup(handler *serviceRPC.NezhaHandler) {
|
||||||
|
if !lease.active {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
_ = handler.CancelAgentCompatIOStreamCapability(lease.access)
|
||||||
|
_ = handler.CloseAgentCompatNATStreamLease(lease.streamLease)
|
||||||
|
_ = handler.UnregisterAgentCompatIOStreamCapability(lease.access)
|
||||||
|
}
|
||||||
@@ -0,0 +1,82 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"log"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||||
|
serviceRPC "github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestPrepareNATCapabilityConsumesAndRemovesHeaderBeforeNATWork(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
handler := serviceRPC.NewNezhaHandler()
|
||||||
|
original := serviceRPC.NezhaHandlerSingleton
|
||||||
|
serviceRPC.NezhaHandlerSingleton = handler
|
||||||
|
t.Cleanup(func() { serviceRPC.NezhaHandlerSingleton = original })
|
||||||
|
request := &http.Request{Header: make(http.Header)}
|
||||||
|
request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "malformed")
|
||||||
|
|
||||||
|
// When
|
||||||
|
lease, err := prepareNATCapability(request, &model.NAT{Common: model.Common{ID: 91}, ServerID: 81})
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("malformed capability unexpectedly accepted")
|
||||||
|
}
|
||||||
|
if lease.active {
|
||||||
|
t.Fatal("malformed capability unexpectedly activated")
|
||||||
|
}
|
||||||
|
if _, present := request.Header[agentcompatcontract.IOStreamCapabilityHeader]; present {
|
||||||
|
t.Fatal("capability header remained after hook")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPrepareNATCapabilityRejectsDuplicateHeaderAfterRemovingAllValues(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
request := &http.Request{Header: make(http.Header)}
|
||||||
|
request.Header.Add(agentcompatcontract.IOStreamCapabilityHeader, "one")
|
||||||
|
request.Header.Add(agentcompatcontract.IOStreamCapabilityHeader, "two")
|
||||||
|
|
||||||
|
// When
|
||||||
|
lease, err := prepareNATCapability(request, &model.NAT{Common: model.Common{ID: 92}, ServerID: 82})
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("duplicate capability header unexpectedly accepted")
|
||||||
|
}
|
||||||
|
if lease.active {
|
||||||
|
t.Fatal("duplicate capability header unexpectedly activated")
|
||||||
|
}
|
||||||
|
if _, present := request.Header[agentcompatcontract.IOStreamCapabilityHeader]; present {
|
||||||
|
t.Fatal("duplicate capability header remained after hook")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServeNATAgentCompatSensitiveHeadersStayOutOfErrorsAndLogs(t *testing.T) {
|
||||||
|
request := &http.Request{Header: make(http.Header)}
|
||||||
|
request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "capability-secret")
|
||||||
|
request.Header.Set("Authorization", "Bearer deterministic-secret")
|
||||||
|
writer := &serveNATResponseWriter{}
|
||||||
|
var logs bytes.Buffer
|
||||||
|
originalOutput := log.Writer()
|
||||||
|
log.SetOutput(&logs)
|
||||||
|
t.Cleanup(func() { log.SetOutput(originalOutput) })
|
||||||
|
|
||||||
|
ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 91}, ServerID: 81})
|
||||||
|
|
||||||
|
if writer.status != http.StatusServiceUnavailable {
|
||||||
|
t.Fatalf("status = %d, want %d", writer.status, http.StatusServiceUnavailable)
|
||||||
|
}
|
||||||
|
for _, output := range []string{writer.body, logs.String()} {
|
||||||
|
if strings.Contains(output, "capability-secret") || strings.Contains(output, "deterministic-secret") {
|
||||||
|
t.Fatalf("sensitive value leaked in %q", output)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,25 @@
|
|||||||
|
//go:build !agentcompat
|
||||||
|
|
||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
serviceRPC "github.com/nezhahq/nezha/service/rpc"
|
||||||
|
)
|
||||||
|
|
||||||
|
type natCapabilityLease struct {
|
||||||
|
active bool
|
||||||
|
publicationOwned bool
|
||||||
|
streamLease *serviceRPC.AgentCompatNATStreamLease
|
||||||
|
access serviceRPC.AgentCompatCapabilityAccess
|
||||||
|
handle serviceRPC.AgentCompatNATPublishHandle
|
||||||
|
}
|
||||||
|
|
||||||
|
func prepareNATCapability(request *http.Request, _ *model.NAT) (natCapabilityLease, error) {
|
||||||
|
request.Header.Del("Authorization")
|
||||||
|
return natCapabilityLease{}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (lease natCapabilityLease) cleanup(*serviceRPC.NezhaHandler) {}
|
||||||
@@ -0,0 +1,64 @@
|
|||||||
|
//go:build !agentcompat
|
||||||
|
|
||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/goccy/go-json"
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||||
|
"github.com/nezhahq/nezha/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestServeNATDefaultForwardsCapabilityHeaderAsOrdinaryData(t *testing.T) {
|
||||||
|
fixture := newServeNATFixture(t)
|
||||||
|
connection := newServeNATConn()
|
||||||
|
agent := &serveNATAgent{readErr: io.EOF, writeDone: make(chan struct{})}
|
||||||
|
request := &http.Request{
|
||||||
|
Method: http.MethodPost,
|
||||||
|
URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"},
|
||||||
|
Header: make(http.Header),
|
||||||
|
Body: &serveNATBody{},
|
||||||
|
}
|
||||||
|
request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "ordinary-data")
|
||||||
|
request.Header.Set("Authorization", "Bearer deterministic-secret")
|
||||||
|
fixture.taskStream.onSend = func(task *proto.Task) error {
|
||||||
|
var nat model.TaskNAT
|
||||||
|
if err := json.Unmarshal([]byte(task.Data), &nat); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return fixture.handler.AgentConnected(nat.StreamID, agent)
|
||||||
|
}
|
||||||
|
done := make(chan struct{})
|
||||||
|
deadline, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
go func() {
|
||||||
|
ServeNAT(&serveNATResponseWriter{conn: connection}, request, &model.NAT{ServerID: fixture.server.ID, Host: "target.example"})
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
select {
|
||||||
|
case <-agent.writeDone:
|
||||||
|
case <-deadline.Done():
|
||||||
|
t.Fatal("agent did not receive ordinary NAT request")
|
||||||
|
}
|
||||||
|
_ = connection.Close()
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-deadline.Done():
|
||||||
|
t.Fatal("ServeNAT did not finish")
|
||||||
|
}
|
||||||
|
forwarded := strings.ToLower(string(agent.writtenBytes()))
|
||||||
|
if !strings.Contains(forwarded, strings.ToLower(agentcompatcontract.IOStreamCapabilityHeader)+": ordinary-data") {
|
||||||
|
t.Fatal("default NAT flow did not forward capability header as ordinary request data")
|
||||||
|
}
|
||||||
|
if strings.Contains(forwarded, "Authorization:") || strings.Contains(forwarded, "deterministic-secret") {
|
||||||
|
t.Fatal("default NAT flow forwarded Authorization credentials")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,35 @@
|
|||||||
|
//go:build !agentcompat
|
||||||
|
|
||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestPrepareNATCapabilityDefaultPreservesHeaderAsOrdinaryRequestData(t *testing.T) {
|
||||||
|
// Given
|
||||||
|
request := &http.Request{Header: make(http.Header)}
|
||||||
|
request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "ordinary-data")
|
||||||
|
request.Header.Set("Authorization", "Bearer deterministic-secret")
|
||||||
|
|
||||||
|
// When
|
||||||
|
lease, err := prepareNATCapability(request, &model.NAT{Common: model.Common{ID: 91}, ServerID: 81})
|
||||||
|
|
||||||
|
// Then
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("default NAT capability hook returned error: %v", err)
|
||||||
|
}
|
||||||
|
if lease.active {
|
||||||
|
t.Fatal("default NAT capability hook unexpectedly activated")
|
||||||
|
}
|
||||||
|
if got := request.Header.Get(agentcompatcontract.IOStreamCapabilityHeader); got != "ordinary-data" {
|
||||||
|
t.Fatalf("default NAT capability hook changed header to %q", got)
|
||||||
|
}
|
||||||
|
if got := request.Header.Get("Authorization"); got != "" {
|
||||||
|
t.Fatalf("default NAT capability hook retained Authorization %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,245 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/goccy/go-json"
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||||
|
"github.com/nezhahq/nezha/proto"
|
||||||
|
serviceRPC "github.com/nezhahq/nezha/service/rpc"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestServeNATAgentCompatReleasesCapabilityBeforeServerDispatch(t *testing.T) {
|
||||||
|
fixture := newServeNATFixture(t)
|
||||||
|
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 101)
|
||||||
|
originalServers := singleton.ServerShared
|
||||||
|
singleton.ServerShared = singleton.NewServerClass()
|
||||||
|
t.Cleanup(func() { singleton.ServerShared = originalServers })
|
||||||
|
request := capabilityRequest(capability)
|
||||||
|
|
||||||
|
ServeNAT(&serveNATResponseWriter{}, request, &model.NAT{Common: model.Common{ID: 101}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||||
|
|
||||||
|
assertNATCapabilitySlotReleased(t, fixture.handler, access)
|
||||||
|
require.Equal(t, 0, fixture.handler.StreamCount())
|
||||||
|
require.Empty(t, fixture.taskStream.sent)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServeNATAgentCompatTaskStreamUnavailableReleasesCapabilityBeforeActivity(t *testing.T) {
|
||||||
|
fixture := newServeNATFixture(t)
|
||||||
|
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 105)
|
||||||
|
fixture.server.SetTaskStream(nil)
|
||||||
|
request := capabilityRequest(capability)
|
||||||
|
|
||||||
|
ServeNAT(&serveNATResponseWriter{}, request, &model.NAT{Common: model.Common{ID: 105}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||||
|
|
||||||
|
assertNATCapabilitySlotReleased(t, fixture.handler, access)
|
||||||
|
require.Equal(t, 0, fixture.handler.StreamCount())
|
||||||
|
require.Empty(t, fixture.taskStream.sent)
|
||||||
|
if request.Header.Get(agentcompatcontract.IOStreamCapabilityHeader) != "" {
|
||||||
|
t.Fatal("capability header remained after task stream unavailable")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServeNATAgentCompatCreateQuotaFailurePreservesExistingStreams(t *testing.T) {
|
||||||
|
fixture := newServeNATFixture(t)
|
||||||
|
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 106)
|
||||||
|
const existingStreamCount = 40
|
||||||
|
for index := 0; index < existingStreamCount; index++ {
|
||||||
|
require.NoError(t, fixture.handler.CreateStreamWithPurpose("quota-existing-"+string(rune(index)), 0, fixture.server.ID, serviceRPC.PurposeNAT))
|
||||||
|
}
|
||||||
|
t.Cleanup(func() {
|
||||||
|
for index := 0; index < existingStreamCount; index++ {
|
||||||
|
_ = fixture.handler.CloseStream("quota-existing-" + string(rune(index)))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
request := capabilityRequest(capability)
|
||||||
|
|
||||||
|
ServeNAT(&serveNATResponseWriter{}, request, &model.NAT{Common: model.Common{ID: 106}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||||
|
|
||||||
|
assertNATCapabilitySlotReleased(t, fixture.handler, access)
|
||||||
|
require.Equal(t, existingStreamCount, fixture.handler.StreamCount())
|
||||||
|
require.Empty(t, fixture.taskStream.sent)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServeNATAgentCompatPublishCancellationClosesRegistryOwnedEndpoints(t *testing.T) {
|
||||||
|
fixture := newServeNATFixture(t)
|
||||||
|
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 107)
|
||||||
|
connection := newServeNATConn()
|
||||||
|
agent := &serveNATAgent{readErr: io.EOF, writeDone: make(chan struct{})}
|
||||||
|
observerEntered := make(chan struct{})
|
||||||
|
observerRelease := make(chan struct{})
|
||||||
|
var observerOnce sync.Once
|
||||||
|
fixture.handler.SetAgentCompatCapabilityPublishObserverForTest(func() {
|
||||||
|
observerOnce.Do(func() { close(observerEntered) })
|
||||||
|
<-observerRelease
|
||||||
|
})
|
||||||
|
t.Cleanup(func() { fixture.handler.SetAgentCompatCapabilityPublishObserverForTest(nil) })
|
||||||
|
fixture.taskStream.onSend = func(task *proto.Task) error {
|
||||||
|
var nat model.TaskNAT
|
||||||
|
require.NoError(t, json.Unmarshal([]byte(task.Data), &nat))
|
||||||
|
return fixture.handler.AgentConnected(nat.StreamID, agent)
|
||||||
|
}
|
||||||
|
request := capabilityRequest(capability)
|
||||||
|
writer := &serveNATResponseWriter{conn: connection}
|
||||||
|
serveDone := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 107}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||||
|
close(serveDone)
|
||||||
|
}()
|
||||||
|
select {
|
||||||
|
case <-observerEntered:
|
||||||
|
case <-serveDone:
|
||||||
|
t.Fatal("ServeNAT returned before publish pause")
|
||||||
|
}
|
||||||
|
fixture.handler.CreateStreamWithPurpose("publish-replacement", 0, fixture.server.ID, serviceRPC.PurposeNAT)
|
||||||
|
require.NoError(t, fixture.handler.CancelAgentCompatIOStreamCapability(access))
|
||||||
|
close(observerRelease)
|
||||||
|
completionContext, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||||
|
defer cancel()
|
||||||
|
select {
|
||||||
|
case <-serveDone:
|
||||||
|
case <-completionContext.Done():
|
||||||
|
t.Fatal("ServeNAT did not finish after publish cancellation")
|
||||||
|
}
|
||||||
|
assertCapabilityReleased(t, fixture.handler, access)
|
||||||
|
require.Equal(t, int32(1), connection.closeCount.Load())
|
||||||
|
require.Equal(t, int32(1), agent.closeCount.Load())
|
||||||
|
require.Equal(t, 1, fixture.handler.StreamCount())
|
||||||
|
select {
|
||||||
|
case <-agent.writeDone:
|
||||||
|
t.Fatal("agent endpoint was used after publish cancellation")
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServeNATAgentCompatSendTaskFailureReleasesStreamAndCapabilityBeforeWrapper(t *testing.T) {
|
||||||
|
fixture := newServeNATFixture(t)
|
||||||
|
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 102)
|
||||||
|
fixture.taskStream.sendErr = io.ErrClosedPipe
|
||||||
|
connection := newServeNATConn()
|
||||||
|
request := capabilityRequest(capability)
|
||||||
|
request.Body = &serveNATBody{}
|
||||||
|
writer := &serveNATResponseWriter{conn: connection}
|
||||||
|
|
||||||
|
ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 102}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||||
|
|
||||||
|
assertCapabilityReleased(t, fixture.handler, access)
|
||||||
|
require.Equal(t, 0, fixture.handler.StreamCount())
|
||||||
|
require.Equal(t, int32(0), connection.closeCount.Load())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServeNATAgentCompatWrapperFailureReleasesCapabilityWithoutSecondResponse(t *testing.T) {
|
||||||
|
fixture := newServeNATFixture(t)
|
||||||
|
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 103)
|
||||||
|
request := capabilityRequest(capability)
|
||||||
|
writer := &nonHijackingResponseWriter{}
|
||||||
|
|
||||||
|
ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 103}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||||
|
|
||||||
|
assertCapabilityReleased(t, fixture.handler, access)
|
||||||
|
require.Equal(t, 0, fixture.handler.StreamCount())
|
||||||
|
require.Equal(t, 0, writer.writes)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServeNATAgentCompatStartFailureCancelsPublishedCapabilityAndClosesEndpoints(t *testing.T) {
|
||||||
|
fixture := newServeNATFixture(t)
|
||||||
|
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 104)
|
||||||
|
connection := newServeNATConn()
|
||||||
|
agent := &serveNATAgent{readErr: io.EOF, writeDone: make(chan struct{})}
|
||||||
|
fixture.taskStream.onSend = func(task *proto.Task) error {
|
||||||
|
var nat model.TaskNAT
|
||||||
|
if err := json.Unmarshal([]byte(task.Data), &nat); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return fixture.handler.AgentConnected(nat.StreamID, agent)
|
||||||
|
}
|
||||||
|
request := capabilityRequest(capability)
|
||||||
|
writer := &serveNATResponseWriter{conn: connection}
|
||||||
|
|
||||||
|
ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 104}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||||
|
|
||||||
|
assertCapabilityReleased(t, fixture.handler, access)
|
||||||
|
require.Equal(t, 0, fixture.handler.StreamCount())
|
||||||
|
require.Equal(t, int32(1), connection.closeCount.Load())
|
||||||
|
require.Equal(t, int32(1), agent.closeCount.Load())
|
||||||
|
}
|
||||||
|
|
||||||
|
func registerNATCapabilityForTest(t *testing.T, handler *serviceRPC.NezhaHandler, serverID, resourceID uint64) (string, serviceRPC.AgentCompatCapabilityAccess) {
|
||||||
|
t.Helper()
|
||||||
|
capability, err := handler.RegisterAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityRegistration{
|
||||||
|
Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: resourceID, UserID: resourceID + 1}, Purpose: serviceRPC.AgentCompatCapabilityNAT,
|
||||||
|
TargetServerID: serverID, ResourceID: resourceID, ServerAccessAllowed: true,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
parsed, err := serviceRPC.ParseAgentCompatIOStreamCapability(capability.String())
|
||||||
|
require.NoError(t, err)
|
||||||
|
return capability.String(), serviceRPC.AgentCompatCapabilityAccess{
|
||||||
|
Capability: parsed, Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: resourceID, UserID: resourceID + 1},
|
||||||
|
Purpose: serviceRPC.AgentCompatCapabilityNAT, TargetServerID: serverID, ResourceID: resourceID, ServerAccessAllowed: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func capabilityRequest(value string) *http.Request {
|
||||||
|
request := &http.Request{Method: http.MethodPost, URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"}, Header: make(http.Header), Body: &serveNATBody{}}
|
||||||
|
request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, value)
|
||||||
|
return request
|
||||||
|
}
|
||||||
|
|
||||||
|
func assertCapabilityReleased(t *testing.T, handler *serviceRPC.NezhaHandler, access serviceRPC.AgentCompatCapabilityAccess) {
|
||||||
|
t.Helper()
|
||||||
|
_, _, err := handler.ConsumeAgentCompatNATCapabilityForProfile(access.Capability.String(), access.TargetServerID, access.ResourceID)
|
||||||
|
require.ErrorIs(t, err, serviceRPC.ErrAgentCompatCapabilityHidden)
|
||||||
|
}
|
||||||
|
|
||||||
|
func assertNATCapabilitySlotReleased(t *testing.T, handler *serviceRPC.NezhaHandler, access serviceRPC.AgentCompatCapabilityAccess) {
|
||||||
|
t.Helper()
|
||||||
|
assertCapabilityReleased(t, handler, access)
|
||||||
|
type activeCapability struct {
|
||||||
|
capability serviceRPC.AgentCompatIOStreamCapability
|
||||||
|
resourceID uint64
|
||||||
|
}
|
||||||
|
active := make([]activeCapability, 0, 16)
|
||||||
|
for index := uint64(0); index < 16; index++ {
|
||||||
|
capability, err := handler.RegisterAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityRegistration{
|
||||||
|
Owner: access.Owner, Purpose: serviceRPC.AgentCompatCapabilityNAT, TargetServerID: access.TargetServerID,
|
||||||
|
ResourceID: access.ResourceID + index + 1000, ServerAccessAllowed: true,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
active = append(active, activeCapability{capability: capability, resourceID: access.ResourceID + index + 1000})
|
||||||
|
}
|
||||||
|
_, err := handler.RegisterAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityRegistration{
|
||||||
|
Owner: access.Owner, Purpose: serviceRPC.AgentCompatCapabilityNAT, TargetServerID: access.TargetServerID,
|
||||||
|
ResourceID: access.ResourceID + 2000, ServerAccessAllowed: true,
|
||||||
|
})
|
||||||
|
require.ErrorIs(t, err, serviceRPC.ErrAgentCompatCapabilityUnavailable)
|
||||||
|
for _, item := range active {
|
||||||
|
parsed, parseErr := serviceRPC.ParseAgentCompatIOStreamCapability(item.capability.String())
|
||||||
|
require.NoError(t, parseErr)
|
||||||
|
require.NoError(t, handler.CancelAgentCompatIOStreamCapability(serviceRPC.AgentCompatCapabilityAccess{
|
||||||
|
Capability: parsed, Owner: access.Owner, Purpose: serviceRPC.AgentCompatCapabilityNAT,
|
||||||
|
TargetServerID: access.TargetServerID, ResourceID: item.resourceID, ServerAccessAllowed: true,
|
||||||
|
}))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type nonHijackingResponseWriter struct {
|
||||||
|
writes int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (writer *nonHijackingResponseWriter) Header() http.Header { return make(http.Header) }
|
||||||
|
func (writer *nonHijackingResponseWriter) Write(data []byte) (int, error) {
|
||||||
|
writer.writes += len(data)
|
||||||
|
return len(data), nil
|
||||||
|
}
|
||||||
|
func (*nonHijackingResponseWriter) WriteHeader(int) {}
|
||||||
@@ -0,0 +1,151 @@
|
|||||||
|
//go:build agentcompat
|
||||||
|
|
||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/goccy/go-json"
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||||
|
"github.com/nezhahq/nezha/proto"
|
||||||
|
serviceRPC "github.com/nezhahq/nezha/service/rpc"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestServeNATAgentCompatPublishesExactStreamAfterRequestTransfer(t *testing.T) {
|
||||||
|
fixture := newServeNATFixture(t)
|
||||||
|
capability, err := fixture.handler.RegisterAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityRegistration{
|
||||||
|
Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: 1, UserID: 2},
|
||||||
|
Purpose: serviceRPC.AgentCompatCapabilityNAT,
|
||||||
|
TargetServerID: fixture.server.ID,
|
||||||
|
ResourceID: 91,
|
||||||
|
ServerAccessAllowed: true,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
access, err := serviceRPC.ParseAgentCompatIOStreamCapability(capability.String())
|
||||||
|
require.NoError(t, err)
|
||||||
|
connection := newServeNATConn()
|
||||||
|
writer := &serveNATResponseWriter{conn: connection}
|
||||||
|
agent := &orderedNATAgent{readStarted: make(chan struct{}), readRelease: make(chan struct{}), writeStarted: make(chan struct{})}
|
||||||
|
taskStreamIDReady := make(chan string, 1)
|
||||||
|
request := &http.Request{
|
||||||
|
Method: http.MethodPost,
|
||||||
|
URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"},
|
||||||
|
Header: make(http.Header),
|
||||||
|
Body: io.NopCloser(strings.NewReader("payload")),
|
||||||
|
}
|
||||||
|
request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability.String())
|
||||||
|
request.Header.Set("Authorization", "Bearer deterministic-secret")
|
||||||
|
request.Header.Set("X-Ordinary-NAT", "ordinary")
|
||||||
|
fixture.taskStream.onSend = func(task *proto.Task) error {
|
||||||
|
var nat model.TaskNAT
|
||||||
|
if err := json.Unmarshal([]byte(task.Data), &nat); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := fixture.handler.AgentConnected(nat.StreamID, agent); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
taskStreamIDReady <- nat.StreamID
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
serveDone := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 91}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||||
|
close(serveDone)
|
||||||
|
}()
|
||||||
|
streamID, err := fixture.handler.WaitAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityAccess{
|
||||||
|
Capability: access,
|
||||||
|
Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: 1, UserID: 2},
|
||||||
|
Purpose: serviceRPC.AgentCompatCapabilityNAT,
|
||||||
|
TargetServerID: fixture.server.ID,
|
||||||
|
ResourceID: 91,
|
||||||
|
ServerAccessAllowed: true,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
taskStreamID := <-taskStreamIDReady
|
||||||
|
require.Equal(t, taskStreamID, streamID)
|
||||||
|
select {
|
||||||
|
case <-agent.readStarted:
|
||||||
|
case <-serveDone:
|
||||||
|
t.Fatal("ServeNAT returned before StartStream read")
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-agent.writeStarted:
|
||||||
|
case <-serveDone:
|
||||||
|
t.Fatal("ServeNAT returned before request reached agent")
|
||||||
|
}
|
||||||
|
close(agent.readRelease)
|
||||||
|
require.NoError(t, connection.Close())
|
||||||
|
completionContext, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
select {
|
||||||
|
case <-serveDone:
|
||||||
|
case <-completionContext.Done():
|
||||||
|
t.Fatal("ServeNAT did not complete after StartStream")
|
||||||
|
}
|
||||||
|
require.NoError(t, fixture.handler.CancelAgentCompatIOStreamCapability(serviceRPC.AgentCompatCapabilityAccess{
|
||||||
|
Capability: access,
|
||||||
|
Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: 1, UserID: 2},
|
||||||
|
Purpose: serviceRPC.AgentCompatCapabilityNAT,
|
||||||
|
TargetServerID: fixture.server.ID,
|
||||||
|
ResourceID: 91,
|
||||||
|
ServerAccessAllowed: true,
|
||||||
|
}))
|
||||||
|
require.NotContains(t, string(agent.writtenBytes()), capability.String())
|
||||||
|
forwarded := string(agent.writtenBytes())
|
||||||
|
for _, expected := range []string{"POST /nat HTTP/1.1", "Host: example.test", "X-Ordinary-Nat: ordinary", "payload"} {
|
||||||
|
require.Contains(t, forwarded, expected)
|
||||||
|
}
|
||||||
|
require.NotContains(t, forwarded, agentcompatcontract.IOStreamCapabilityHeader)
|
||||||
|
require.NotContains(t, forwarded, "Authorization:")
|
||||||
|
require.NotContains(t, forwarded, "deterministic-secret")
|
||||||
|
}
|
||||||
|
|
||||||
|
type orderedNATAgent struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
written bytes.Buffer
|
||||||
|
readStarted chan struct{}
|
||||||
|
readRelease chan struct{}
|
||||||
|
writeStarted chan struct{}
|
||||||
|
readOnce sync.Once
|
||||||
|
writeOnce sync.Once
|
||||||
|
closeCount int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (agent *orderedNATAgent) Read([]byte) (int, error) {
|
||||||
|
agent.readOnce.Do(func() { close(agent.readStarted) })
|
||||||
|
<-agent.readRelease
|
||||||
|
return 0, io.EOF
|
||||||
|
}
|
||||||
|
|
||||||
|
func (agent *orderedNATAgent) Write(data []byte) (int, error) {
|
||||||
|
agent.mu.Lock()
|
||||||
|
count, err := agent.written.Write(data)
|
||||||
|
agent.mu.Unlock()
|
||||||
|
agent.writeOnce.Do(func() { close(agent.writeStarted) })
|
||||||
|
return count, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (agent *orderedNATAgent) Close() error {
|
||||||
|
agent.mu.Lock()
|
||||||
|
defer agent.mu.Unlock()
|
||||||
|
agent.closeCount++
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (agent *orderedNATAgent) writtenBytes() []byte {
|
||||||
|
agent.mu.Lock()
|
||||||
|
defer agent.mu.Unlock()
|
||||||
|
return append([]byte(nil), agent.written.Bytes()...)
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ io.ReadWriteCloser = (*orderedNATAgent)(nil)
|
||||||
@@ -0,0 +1,93 @@
|
|||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/goccy/go-json"
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/proto"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestServeNATClosesHijackedResourcesWhenStreamDisappearsBeforeUserTransfer(t *testing.T) {
|
||||||
|
fixture := newServeNATFixture(t)
|
||||||
|
conn := newServeNATConn()
|
||||||
|
body := &serveNATBody{}
|
||||||
|
writer := &serveNATResponseWriter{conn: conn}
|
||||||
|
request := &http.Request{Method: http.MethodPost, URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"}, Header: make(http.Header), Body: body, ContentLength: 0}
|
||||||
|
request.Header.Set("X-NAT-Test", "failure")
|
||||||
|
fixture.taskStream.onSend = func(task *proto.Task) error {
|
||||||
|
var nat model.TaskNAT
|
||||||
|
if err := json.Unmarshal([]byte(task.Data), &nat); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return fixture.handler.CloseStream(nat.StreamID)
|
||||||
|
}
|
||||||
|
|
||||||
|
ServeNAT(writer, request, &model.NAT{ServerID: fixture.server.ID, Host: "target.example"})
|
||||||
|
|
||||||
|
if got := body.closeCount.Load(); got == 0 {
|
||||||
|
t.Fatal("request body was not closed after failed transfer")
|
||||||
|
}
|
||||||
|
if got := conn.closeCount.Load(); got != 1 {
|
||||||
|
t.Fatalf("hijacked connection close count = %d, want 1", got)
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-conn.readDone:
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("hijacked connection read remained blocked after failed transfer")
|
||||||
|
}
|
||||||
|
if writer.status != 0 || writer.writes != 0 {
|
||||||
|
t.Fatalf("hijacked failure wrote HTTP response: status=%d writes=%d", writer.status, writer.writes)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestServeNATTransfersRequestAndRegistryOwnsSuccessfulCleanup(t *testing.T) {
|
||||||
|
fixture := newServeNATFixture(t)
|
||||||
|
conn := newServeNATConn()
|
||||||
|
body := &serveNATBody{}
|
||||||
|
agent := &serveNATAgent{readErr: io.EOF, writeDone: make(chan struct{})}
|
||||||
|
writer := &serveNATResponseWriter{conn: conn}
|
||||||
|
request := &http.Request{Method: http.MethodPost, URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"}, Header: make(http.Header), Body: body, ContentLength: 0}
|
||||||
|
request.Header.Set("X-NAT-Test", "success")
|
||||||
|
fixture.taskStream.onSend = func(task *proto.Task) error {
|
||||||
|
var nat model.TaskNAT
|
||||||
|
if err := json.Unmarshal([]byte(task.Data), &nat); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return fixture.handler.AgentConnected(nat.StreamID, agent)
|
||||||
|
}
|
||||||
|
|
||||||
|
serveDone := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
ServeNAT(writer, request, &model.NAT{ServerID: fixture.server.ID, Host: "target.example"})
|
||||||
|
close(serveDone)
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-agent.writeDone:
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("agent did not receive the transferred NAT request")
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-serveDone:
|
||||||
|
case <-time.After(time.Second):
|
||||||
|
t.Fatal("ServeNAT did not finish after the successful stream was closed")
|
||||||
|
}
|
||||||
|
if got := body.closeCount.Load(); got == 0 {
|
||||||
|
t.Fatal("registry-owned request body was not closed")
|
||||||
|
}
|
||||||
|
if got := conn.closeCount.Load(); got != 1 {
|
||||||
|
t.Fatalf("registry-owned hijacked connection close count = %d, want 1", got)
|
||||||
|
}
|
||||||
|
if got := agent.closeCount.Load(); got != 1 {
|
||||||
|
t.Fatalf("registry-owned agent close count = %d, want 1", got)
|
||||||
|
}
|
||||||
|
if got := strings.ToLower(string(agent.writtenBytes())); !strings.Contains(got, "post /nat http/1.1") || !strings.Contains(got, "x-nat-test: success") {
|
||||||
|
t.Fatalf("agent received incomplete NAT request: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,171 @@
|
|||||||
|
package rpc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/nezhahq/nezha/model"
|
||||||
|
"github.com/nezhahq/nezha/proto"
|
||||||
|
rpcService "github.com/nezhahq/nezha/service/rpc"
|
||||||
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"google.golang.org/grpc/metadata"
|
||||||
|
"gorm.io/driver/sqlite"
|
||||||
|
"gorm.io/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
type serveNATFixture struct {
|
||||||
|
handler *rpcService.NezhaHandler
|
||||||
|
server *model.Server
|
||||||
|
taskStream *serveNATTaskStream
|
||||||
|
}
|
||||||
|
|
||||||
|
func newServeNATFixture(t *testing.T) serveNATFixture {
|
||||||
|
t.Helper()
|
||||||
|
originalDB, originalServerShared, originalHandler := singleton.DB, singleton.ServerShared, rpcService.NezhaHandlerSingleton
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.NoError(t, db.AutoMigrate(model.Server{}))
|
||||||
|
server := &model.Server{Common: model.Common{ID: 7}, UUID: "serve-nat-test", Name: "serve-nat-test"}
|
||||||
|
require.NoError(t, db.Create(server).Error)
|
||||||
|
singleton.DB = db
|
||||||
|
singleton.ServerShared = singleton.NewServerClass()
|
||||||
|
handler := rpcService.NewNezhaHandler()
|
||||||
|
rpcService.NezhaHandlerSingleton = handler
|
||||||
|
taskStream := &serveNATTaskStream{}
|
||||||
|
server, ok := singleton.ServerShared.Get(server.ID)
|
||||||
|
require.True(t, ok)
|
||||||
|
server.SetTaskStream(taskStream)
|
||||||
|
taskStream.server = server
|
||||||
|
t.Cleanup(func() {
|
||||||
|
rpcService.NezhaHandlerSingleton, singleton.ServerShared, singleton.DB = originalHandler, originalServerShared, originalDB
|
||||||
|
if dbSQL, dbErr := db.DB(); dbErr == nil {
|
||||||
|
_ = dbSQL.Close()
|
||||||
|
}
|
||||||
|
})
|
||||||
|
return serveNATFixture{handler: handler, server: server, taskStream: taskStream}
|
||||||
|
}
|
||||||
|
|
||||||
|
type serveNATTaskStream struct {
|
||||||
|
server *model.Server
|
||||||
|
onSend func(*proto.Task) error
|
||||||
|
sendErr error
|
||||||
|
sent []*proto.Task
|
||||||
|
}
|
||||||
|
|
||||||
|
func (stream *serveNATTaskStream) Send(task *proto.Task) error {
|
||||||
|
stream.sent = append(stream.sent, task)
|
||||||
|
if stream.sendErr != nil {
|
||||||
|
return stream.sendErr
|
||||||
|
}
|
||||||
|
if stream.onSend != nil {
|
||||||
|
return stream.onSend(task)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (*serveNATTaskStream) Recv() (*proto.TaskResult, error) { return nil, io.EOF }
|
||||||
|
func (*serveNATTaskStream) SetHeader(metadata.MD) error { return nil }
|
||||||
|
func (*serveNATTaskStream) SendHeader(metadata.MD) error { return nil }
|
||||||
|
func (*serveNATTaskStream) SetTrailer(metadata.MD) {}
|
||||||
|
func (*serveNATTaskStream) Context() context.Context { return context.Background() }
|
||||||
|
func (*serveNATTaskStream) SendMsg(any) error { return nil }
|
||||||
|
func (*serveNATTaskStream) RecvMsg(any) error { return io.EOF }
|
||||||
|
|
||||||
|
type serveNATResponseWriter struct {
|
||||||
|
conn *serveNATConn
|
||||||
|
header http.Header
|
||||||
|
status, writes int
|
||||||
|
body string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (writer *serveNATResponseWriter) Header() http.Header {
|
||||||
|
if writer.header == nil {
|
||||||
|
writer.header = make(http.Header)
|
||||||
|
}
|
||||||
|
return writer.header
|
||||||
|
}
|
||||||
|
func (writer *serveNATResponseWriter) Write(data []byte) (int, error) {
|
||||||
|
writer.writes += len(data)
|
||||||
|
writer.body += string(data)
|
||||||
|
return len(data), nil
|
||||||
|
}
|
||||||
|
func (writer *serveNATResponseWriter) WriteHeader(status int) { writer.status = status }
|
||||||
|
func (writer *serveNATResponseWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
|
||||||
|
go func() { _, _ = writer.conn.Read(make([]byte, 1)) }()
|
||||||
|
return writer.conn, bufio.NewReadWriter(bufio.NewReader(writer.conn), bufio.NewWriter(writer.conn)), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type serveNATBody struct{ closeCount atomic.Int32 }
|
||||||
|
|
||||||
|
func (*serveNATBody) Read([]byte) (int, error) { return 0, io.EOF }
|
||||||
|
func (body *serveNATBody) Close() error { body.closeCount.Add(1); return nil }
|
||||||
|
|
||||||
|
type serveNATConn struct {
|
||||||
|
closed, readDone chan struct{}
|
||||||
|
readDoneOnce, closeOnce sync.Once
|
||||||
|
closeCount atomic.Int32
|
||||||
|
}
|
||||||
|
|
||||||
|
func newServeNATConn() *serveNATConn {
|
||||||
|
return &serveNATConn{closed: make(chan struct{}), readDone: make(chan struct{})}
|
||||||
|
}
|
||||||
|
func (conn *serveNATConn) Read([]byte) (int, error) {
|
||||||
|
<-conn.closed
|
||||||
|
conn.readDoneOnce.Do(func() { close(conn.readDone) })
|
||||||
|
return 0, io.EOF
|
||||||
|
}
|
||||||
|
func (*serveNATConn) Write(data []byte) (int, error) { return len(data), nil }
|
||||||
|
func (conn *serveNATConn) Close() error {
|
||||||
|
conn.closeCount.Add(1)
|
||||||
|
conn.closeOnce.Do(func() { close(conn.closed) })
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (*serveNATConn) LocalAddr() net.Addr { return serveNATAddr("local") }
|
||||||
|
func (*serveNATConn) RemoteAddr() net.Addr { return serveNATAddr("remote") }
|
||||||
|
func (*serveNATConn) SetDeadline(time.Time) error { return nil }
|
||||||
|
func (*serveNATConn) SetReadDeadline(time.Time) error { return nil }
|
||||||
|
func (*serveNATConn) SetWriteDeadline(time.Time) error { return nil }
|
||||||
|
|
||||||
|
type serveNATAddr string
|
||||||
|
|
||||||
|
func (addr serveNATAddr) Network() string { return "test" }
|
||||||
|
func (addr serveNATAddr) String() string { return string(addr) }
|
||||||
|
|
||||||
|
type serveNATAgent struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
written bytes.Buffer
|
||||||
|
readErr error
|
||||||
|
writeDone chan struct{}
|
||||||
|
writeOnce sync.Once
|
||||||
|
closeCount atomic.Int32
|
||||||
|
}
|
||||||
|
|
||||||
|
func (agent *serveNATAgent) Read([]byte) (int, error) { return 0, agent.readErr }
|
||||||
|
func (agent *serveNATAgent) Write(data []byte) (int, error) {
|
||||||
|
agent.mu.Lock()
|
||||||
|
defer agent.mu.Unlock()
|
||||||
|
count, err := agent.written.Write(data)
|
||||||
|
if agent.writeDone != nil {
|
||||||
|
agent.writeOnce.Do(func() { close(agent.writeDone) })
|
||||||
|
}
|
||||||
|
return count, err
|
||||||
|
}
|
||||||
|
func (agent *serveNATAgent) Close() error { agent.closeCount.Add(1); return nil }
|
||||||
|
func (agent *serveNATAgent) writtenBytes() []byte {
|
||||||
|
agent.mu.Lock()
|
||||||
|
defer agent.mu.Unlock()
|
||||||
|
return append([]byte(nil), agent.written.Bytes()...)
|
||||||
|
}
|
||||||
|
|
||||||
|
var _ proto.NezhaService_RequestTaskServer = (*serveNATTaskStream)(nil)
|
||||||
|
var _ http.Hijacker = (*serveNATResponseWriter)(nil)
|
||||||
|
var _ net.Conn = (*serveNATConn)(nil)
|
||||||
|
var _ io.ReadWriteCloser = (*serveNATAgent)(nil)
|
||||||
+9
-238
@@ -1,27 +1,23 @@
|
|||||||
package rpc
|
package rpc
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"net"
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"log"
|
|
||||||
"net/http"
|
|
||||||
"net/netip"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/goccy/go-json"
|
|
||||||
"google.golang.org/grpc"
|
"google.golang.org/grpc"
|
||||||
"google.golang.org/grpc/metadata"
|
|
||||||
"google.golang.org/grpc/peer"
|
|
||||||
|
|
||||||
"github.com/hashicorp/go-uuid"
|
|
||||||
"github.com/nezhahq/nezha/model"
|
|
||||||
"github.com/nezhahq/nezha/pkg/utils"
|
|
||||||
"github.com/nezhahq/nezha/proto"
|
"github.com/nezhahq/nezha/proto"
|
||||||
rpcService "github.com/nezhahq/nezha/service/rpc"
|
rpcService "github.com/nezhahq/nezha/service/rpc"
|
||||||
"github.com/nezhahq/nezha/service/singleton"
|
"github.com/nezhahq/nezha/service/singleton"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func SetReceiptGateListener(listener net.Listener) {
|
||||||
|
rpcService.SetReceiptGateListener(listener)
|
||||||
|
}
|
||||||
|
|
||||||
|
func CloseReceiptGate() {
|
||||||
|
rpcService.CloseReceiptGate()
|
||||||
|
}
|
||||||
|
|
||||||
// SetMCPKillSwitchObserver re-exports the service/rpc hook so cmd/dashboard
|
// SetMCPKillSwitchObserver re-exports the service/rpc hook so cmd/dashboard
|
||||||
// can wire singleton.Conf.EnableMCP without importing the inner rpc package
|
// can wire singleton.Conf.EnableMCP without importing the inner rpc package
|
||||||
// (cmd/dashboard already imports cmd/dashboard/rpc for ServeRPC).
|
// (cmd/dashboard already imports cmd/dashboard/rpc for ServeRPC).
|
||||||
@@ -46,228 +42,3 @@ func ServeRPC() *grpc.Server {
|
|||||||
proto.RegisterNezhaServiceServer(server, rpcService.NezhaHandlerSingleton)
|
proto.RegisterNezhaServiceServer(server, rpcService.NezhaHandlerSingleton)
|
||||||
return server
|
return server
|
||||||
}
|
}
|
||||||
|
|
||||||
func ctxWithRealIP(ctx context.Context) (context.Context, error) {
|
|
||||||
var ip, connectingIp string
|
|
||||||
p, ok := peer.FromContext(ctx)
|
|
||||||
if ok {
|
|
||||||
addrPort, err := netip.ParseAddrPort(p.Addr.String())
|
|
||||||
if err == nil {
|
|
||||||
connectingIp = addrPort.Addr().String()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
ctx = context.WithValue(ctx, model.CtxKeyConnectingIP{}, connectingIp)
|
|
||||||
|
|
||||||
if singleton.Conf.AgentRealIPHeader == "" {
|
|
||||||
return ctx, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
if singleton.Conf.AgentRealIPHeader == model.ConfigUsePeerIP {
|
|
||||||
if connectingIp == "" {
|
|
||||||
return ctx, fmt.Errorf("connecting ip not found")
|
|
||||||
}
|
|
||||||
// Peer-IP mode: peer IP is the real IP. Leaving ip="" makes
|
|
||||||
// CheckIP/BlockIP short-circuit on empty IP, disabling the WAF.
|
|
||||||
ip = connectingIp
|
|
||||||
} else {
|
|
||||||
vals := metadata.ValueFromIncomingContext(ctx, singleton.Conf.AgentRealIPHeader)
|
|
||||||
if len(vals) == 0 {
|
|
||||||
return ctx, fmt.Errorf("real ip header not found")
|
|
||||||
}
|
|
||||||
var err error
|
|
||||||
ip, err = utils.GetIPFromHeader(vals[0])
|
|
||||||
if err != nil {
|
|
||||||
return ctx, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if singleton.Conf.Debug {
|
|
||||||
log.Printf("NEZHA>> gRPC Agent Real IP: %s, connecting IP: %s\n", ip, connectingIp)
|
|
||||||
}
|
|
||||||
|
|
||||||
return context.WithValue(ctx, model.CtxKeyRealIP{}, ip), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func waf(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) {
|
|
||||||
realip, _ := ctx.Value(model.CtxKeyRealIP{}).(string)
|
|
||||||
if err := model.CheckIP(singleton.DB, realip); err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return handler(ctx, req)
|
|
||||||
}
|
|
||||||
|
|
||||||
func getRealIp(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) {
|
|
||||||
ctx, err := ctxWithRealIP(ctx)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
return handler(ctx, req)
|
|
||||||
}
|
|
||||||
|
|
||||||
// realIPServerStream overrides Context() so stream handlers and
|
|
||||||
// authHandler.check observe the resolved real IP, like the unary path.
|
|
||||||
type realIPServerStream struct {
|
|
||||||
grpc.ServerStream
|
|
||||||
ctx context.Context
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *realIPServerStream) Context() context.Context { return s.ctx }
|
|
||||||
|
|
||||||
func getRealIpStream(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
|
|
||||||
ctx, err := ctxWithRealIP(ss.Context())
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return handler(srv, &realIPServerStream{ServerStream: ss, ctx: ctx})
|
|
||||||
}
|
|
||||||
|
|
||||||
func wafStream(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
|
|
||||||
realip, _ := ss.Context().Value(model.CtxKeyRealIP{}).(string)
|
|
||||||
if err := model.CheckIP(singleton.DB, realip); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
return handler(srv, ss)
|
|
||||||
}
|
|
||||||
|
|
||||||
func DispatchTask(serviceSentinelDispatchBus <-chan *model.Service) {
|
|
||||||
for task := range serviceSentinelDispatchBus {
|
|
||||||
if task == nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
switch task.Cover {
|
|
||||||
case model.ServiceCoverIgnoreAll:
|
|
||||||
for id, enabled := range task.SkipServers {
|
|
||||||
if !enabled {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
server, _ := singleton.ServerShared.Get(id)
|
|
||||||
if server == nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if !canSendTaskToServer(task, server) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
// SendTask 走 holder-scoped send mutex,避免与 cron /
|
|
||||||
// server-transfer / MCP CallAgent / fs.transfer 等并发
|
|
||||||
// SendMsg 同一 RequestTask stream。
|
|
||||||
if err := server.SendTask(task.PB()); err != nil &&
|
|
||||||
!errors.Is(err, model.ErrTaskStreamOffline) {
|
|
||||||
log.Printf("NEZHA>> DispatchTask send error (server=%d): %v", id, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
case model.ServiceCoverAll:
|
|
||||||
// 快照后逐个 SendTask,不在 ServerShared 的 listMu.RLock 内做阻塞
|
|
||||||
// gRPC:否则一个卡死 agent 会拖死需要写锁的 server 生命周期操作。
|
|
||||||
for id, server := range singleton.ServerShared.GetList() {
|
|
||||||
if server == nil || task.SkipServers[id] {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if !canSendTaskToServer(task, server) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if err := server.SendTask(task.PB()); err != nil &&
|
|
||||||
!errors.Is(err, model.ErrTaskStreamOffline) {
|
|
||||||
log.Printf("NEZHA>> DispatchTask send error (server=%d): %v", id, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func DispatchKeepalive() {
|
|
||||||
singleton.CronShared.AddFunc("@every 20s", func() {
|
|
||||||
list := singleton.ServerShared.GetSortedList()
|
|
||||||
for _, s := range list {
|
|
||||||
if s == nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if err := s.SendTask(&proto.Task{Type: model.TaskTypeKeepalive}); err != nil &&
|
|
||||||
!errors.Is(err, model.ErrTaskStreamOffline) {
|
|
||||||
log.Printf("NEZHA>> Keepalive send error (server=%d): %v", s.ID, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func ServeNAT(w http.ResponseWriter, r *http.Request, natConfig *model.NAT) {
|
|
||||||
server, _ := singleton.ServerShared.Get(natConfig.ServerID)
|
|
||||||
if server == nil {
|
|
||||||
w.WriteHeader(http.StatusServiceUnavailable)
|
|
||||||
w.Write([]byte("server not found or not connected"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if server.GetTaskStream() == nil {
|
|
||||||
w.WriteHeader(http.StatusServiceUnavailable)
|
|
||||||
w.Write([]byte("server not found or not connected"))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
streamId, err := uuid.GenerateUUID()
|
|
||||||
if err != nil {
|
|
||||||
w.WriteHeader(http.StatusServiceUnavailable)
|
|
||||||
w.Write(fmt.Appendf(nil, "stream id error: %v", err))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// NAT streams are anonymous HTTP-facing tunnels; they are NOT reachable
|
|
||||||
// via /ws/terminal or /ws/file (which check stream ownership), so the
|
|
||||||
// creator user ID does not need to identify a real user. The targetServerID
|
|
||||||
// IS required though — the receiving agent must prove it is the server the
|
|
||||||
// NAT config addressed, otherwise any agent that snoops the streamId can
|
|
||||||
// answer NAT traffic on behalf of an unrelated host.
|
|
||||||
if err := rpcService.NezhaHandlerSingleton.CreateStream(streamId, 0, server.ID); err != nil {
|
|
||||||
w.WriteHeader(http.StatusTooManyRequests)
|
|
||||||
w.Write(fmt.Appendf(nil, "stream limit: %v", err))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
defer rpcService.NezhaHandlerSingleton.CloseStream(streamId)
|
|
||||||
|
|
||||||
taskData, err := json.Marshal(model.TaskNAT{
|
|
||||||
StreamID: streamId,
|
|
||||||
Host: natConfig.Host,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
w.WriteHeader(http.StatusServiceUnavailable)
|
|
||||||
w.Write(fmt.Appendf(nil, "task data error: %v", err))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := server.SendTask(&proto.Task{
|
|
||||||
Type: model.TaskTypeNAT,
|
|
||||||
Data: string(taskData),
|
|
||||||
}); err != nil {
|
|
||||||
w.WriteHeader(http.StatusServiceUnavailable)
|
|
||||||
w.Write(fmt.Appendf(nil, "send task error: %v", err))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
wWrapped, err := utils.NewRequestWrapper(r, w)
|
|
||||||
if err != nil {
|
|
||||||
w.WriteHeader(http.StatusServiceUnavailable)
|
|
||||||
w.Write(fmt.Appendf(nil, "request wrapper error: %v", err))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := rpcService.NezhaHandlerSingleton.UserConnected(streamId, wWrapped); err != nil {
|
|
||||||
w.WriteHeader(http.StatusServiceUnavailable)
|
|
||||||
w.Write(fmt.Appendf(nil, "user connected error: %v", err))
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
rpcService.NezhaHandlerSingleton.StartStream(streamId, time.Second*10)
|
|
||||||
}
|
|
||||||
|
|
||||||
func canSendTaskToServer(task *model.Service, server *model.Server) bool {
|
|
||||||
var role model.Role
|
|
||||||
singleton.UserLock.RLock()
|
|
||||||
if u, ok := singleton.UserInfoMap[task.UserID]; !ok {
|
|
||||||
role = model.RoleMember
|
|
||||||
} else {
|
|
||||||
role = u.Role
|
|
||||||
}
|
|
||||||
singleton.UserLock.RUnlock()
|
|
||||||
|
|
||||||
return task.UserID == server.GetUserID() || role.IsAdmin()
|
|
||||||
}
|
|
||||||
|
|||||||
Reference in New Issue
Block a user