test(agentcompat): add integration scenarios

Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
naiba
2026-07-20 04:51:20 +00:00
co-authored by naiba/CloudCode
parent b3b3d92895
commit 9af3720bf0
96 changed files with 13091 additions and 0 deletions
@@ -0,0 +1,21 @@
//go:build linux
package scenario
import (
"os"
"sigs.k8s.io/yaml"
)
func ReadConfigFile(path string) (AgentConfig, error) {
data, err := os.ReadFile(path)
if err != nil {
return AgentConfig{}, err
}
var config AgentConfig
if err := yaml.Unmarshal(data, &config); err != nil {
return AgentConfig{}, err
}
return config, nil
}
@@ -0,0 +1,101 @@
//go:build linux
package scenario
import (
"context"
"encoding/json"
"errors"
"reflect"
"github.com/nezhahq/nezha/integration/agentcompat/internal/evidence"
)
var ErrConfigIdentityChanged = errors.New("scenario: config identity changed")
type AgentConfigSnapshot struct {
Debug bool
ReportDelay uint32
ClientSecret string
UUID string
Server string
}
type AgentConfig struct {
Debug bool `json:"debug" yaml:"debug"`
Server string `json:"server" yaml:"server"`
ClientSecret string `json:"client_secret" yaml:"client_secret"`
UUID string `json:"uuid" yaml:"uuid"`
ReportDelay uint32 `json:"report_delay" yaml:"report_delay"`
TLS bool `json:"tls" yaml:"tls"`
InsecureTLS bool `json:"insecure_tls" yaml:"insecure_tls"`
}
type ConfigDiffResult struct {
DebugChanged bool
ReportDelayChanged bool
}
func ConfigDiff(original, updated AgentConfigSnapshot) (ConfigDiffResult, error) {
if original.ClientSecret != updated.ClientSecret || original.UUID != updated.UUID || original.Server != updated.Server {
return ConfigDiffResult{}, ErrConfigIdentityChanged
}
return ConfigDiffResult{DebugChanged: original.Debug != updated.Debug, ReportDelayChanged: original.ReportDelay != updated.ReportDelay}, nil
}
type Assertion struct {
Name string `json:"name"`
Passed bool `json:"passed"`
Details string `json:"details,omitempty"`
}
type Result struct {
Name string `json:"name"`
Passed bool `json:"passed"`
Assertions []Assertion `json:"assertions"`
CleanupOK bool `json:"cleanup_ok"`
Error string `json:"error,omitempty"`
}
type AssertionSet struct {
assertions []Assertion
}
func NewAssertionSet() *AssertionSet { return &AssertionSet{} }
func (set *AssertionSet) Record(name string, passed bool, details string) {
set.assertions = append(set.assertions, Assertion{Name: name, Passed: passed, Details: evidence.Redact(details)})
}
func (set *AssertionSet) Results() []Assertion {
return append([]Assertion(nil), set.assertions...)
}
func (set *AssertionSet) Run(run func(*AssertionSet) error) error { return run(set) }
func configSnapshot(config AgentConfig) AgentConfigSnapshot {
return AgentConfigSnapshot{Debug: config.Debug, ReportDelay: config.ReportDelay, ClientSecret: config.ClientSecret, UUID: config.UUID, Server: config.Server}
}
func changedOnlyDebugAndReportDelay(original, updated AgentConfig) error {
before := original
after := updated
before.Debug = after.Debug
before.ReportDelay = after.ReportDelay
if !reflect.DeepEqual(before, after) {
return errors.New("config update changed fields other than debug and report_delay")
}
return nil
}
func decodeAgentConfig(raw string) (AgentConfig, error) {
var config AgentConfig
if err := json.Unmarshal([]byte(raw), &config); err != nil {
return AgentConfig{}, err
}
return config, nil
}
type Context struct {
Context context.Context
}
@@ -0,0 +1,75 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"sync"
)
var (
ErrInvalidHeldCleanupAction = errors.New("held cleanup action is invalid")
ErrHeldCleanupClosed = errors.New("held cleanup stack is closed")
)
type heldCleanupAction struct {
name string
cleanup func(context.Context) error
}
type heldCleanupStackState uint8
const (
heldCleanupOpen heldCleanupStackState = iota
heldCleanupRunning
heldCleanupClosed
)
type heldCleanupStack struct {
mu sync.Mutex
state heldCleanupStackState
actions []heldCleanupAction
}
func newHeldCleanupStack() *heldCleanupStack {
return &heldCleanupStack{state: heldCleanupOpen}
}
func (stack *heldCleanupStack) Push(action heldCleanupAction) error {
if action.name == "" || action.cleanup == nil {
return ErrInvalidHeldCleanupAction
}
stack.mu.Lock()
defer stack.mu.Unlock()
if stack.state != heldCleanupOpen {
return ErrHeldCleanupClosed
}
stack.actions = append(stack.actions, action)
return nil
}
func (stack *heldCleanupStack) Run(ctx context.Context) error {
stack.mu.Lock()
if stack.state != heldCleanupOpen {
stack.mu.Unlock()
return ErrHeldCleanupClosed
}
stack.state = heldCleanupRunning
actions := append([]heldCleanupAction(nil), stack.actions...)
stack.mu.Unlock()
var joined error
for index := len(actions) - 1; index >= 0; index-- {
action := actions[index]
if err := action.cleanup(ctx); err != nil {
joined = errors.Join(joined, fmt.Errorf("cleanup %s: %w", action.name, err))
}
}
stack.mu.Lock()
stack.state = heldCleanupClosed
stack.mu.Unlock()
return joined
}
@@ -0,0 +1,175 @@
//go:build linux
package scenario
import (
"context"
"errors"
"reflect"
"strings"
"testing"
"github.com/stretchr/testify/require"
)
func TestHeldCleanupStackRunsActionsInReverseOrderAndJoinsErrors(t *testing.T) {
stack := newHeldCleanupStack()
var order []string
firstErr := errors.New("first cleanup failure")
secondErr := errors.New("second cleanup failure")
for _, action := range []heldCleanupAction{
{name: "first", cleanup: func(context.Context) error {
order = append(order, "first")
return firstErr
}},
{name: "second", cleanup: func(context.Context) error {
order = append(order, "second")
return secondErr
}},
} {
if err := stack.Push(action); err != nil {
t.Fatal(err)
}
}
err := stack.Run(context.Background())
if !errors.Is(err, firstErr) || !errors.Is(err, secondErr) || !strings.Contains(err.Error(), "first") || !strings.Contains(err.Error(), "second") {
t.Fatalf("joined cleanup error = %v", err)
}
if !reflect.DeepEqual(order, []string{"second", "first"}) {
t.Fatalf("cleanup order = %v", order)
}
}
func TestHeldTerminalCleanupOrderCancelsBeforeWaitingForAbsence(t *testing.T) {
// Given
stack := newHeldCleanupStack()
var order []string
for _, action := range []heldCleanupAction{
{name: "unregister", cleanup: func(context.Context) error { order = append(order, "unregister"); return nil }},
{name: "absence", cleanup: func(context.Context) error { order = append(order, "absence"); return nil }},
{name: "cancel", cleanup: func(context.Context) error { order = append(order, "cancel"); return nil }},
{name: "close", cleanup: func(context.Context) error { order = append(order, "close"); return nil }},
{name: "stop", cleanup: func(context.Context) error { order = append(order, "stop"); return nil }},
{name: "await", cleanup: func(context.Context) error { order = append(order, "await"); return nil }},
{name: "release", cleanup: func(context.Context) error { order = append(order, "release"); return nil }},
} {
if err := stack.Push(action); err != nil {
t.Fatal(err)
}
}
// When
require.NoError(t, stack.Run(context.Background()))
// Then
require.Equal(t, []string{"release", "await", "stop", "close", "cancel", "absence", "unregister"}, order)
}
func TestHeldFMCleanupOrderStopsTransportCancelsThenProvesAbsenceBeforeUnregister(t *testing.T) {
// Given
stack := newHeldCleanupStack()
var order []string
for _, action := range []heldCleanupAction{
{name: "remove fixture", cleanup: func(context.Context) error { order = append(order, "fixture"); return nil }},
{name: "unregister", cleanup: func(context.Context) error { order = append(order, "unregister"); return nil }},
{name: "absence", cleanup: func(context.Context) error { order = append(order, "absence"); return nil }},
{name: "cancel", cleanup: func(context.Context) error { order = append(order, "cancel"); return nil }},
{name: "transport", cleanup: func(context.Context) error { order = append(order, "transport"); return nil }},
} {
require.NoError(t, stack.Push(action))
}
// When
err := stack.Run(context.Background())
// Then
require.NoError(t, err)
require.Equal(t, []string{"transport", "cancel", "absence", "unregister", "fixture"}, order)
}
func TestHeldCleanupStackRunsEveryActionAndJoinsAllErrors(t *testing.T) {
// Given
stack := newHeldCleanupStack()
firstErr := errors.New("first")
secondErr := errors.New("second")
thirdErr := errors.New("third")
for name, actionErr := range map[string]error{"first": firstErr, "second": secondErr, "third": thirdErr} {
require.NoError(t, stack.Push(heldCleanupAction{name: name, cleanup: func(context.Context) error { return actionErr }}))
}
// When
err := stack.Run(context.Background())
// Then
require.ErrorIs(t, err, firstErr)
require.ErrorIs(t, err, secondErr)
require.ErrorIs(t, err, thirdErr)
}
func TestHeldLegacyFMCapabilityCleanupWiringRunsAllActionsInContractOrder(t *testing.T) {
// Given
stack := newHeldCleanupStack()
var order []string
originalErr := errors.New("original")
absenceErr := errors.New("absence")
cleanup := heldLegacyFMCapabilityCleanup{
Unregister: func(context.Context) error { order = append(order, "unregister"); return nil },
Absence: func(context.Context) error { order = append(order, "absence"); return absenceErr },
Cancel: func(context.Context) error { order = append(order, "cancel"); return originalErr },
}
require.NoError(t, pushHeldLegacyFMCapabilityCleanup(stack, cleanup))
// When
err := errors.Join(originalErr, stack.Run(context.Background()))
// Then
require.Equal(t, []string{"cancel", "absence", "unregister"}, order)
require.ErrorIs(t, err, originalErr)
require.ErrorIs(t, err, absenceErr)
}
func TestHeldTerminalCleanupGraceTimeoutLeavesFallbackBudget(t *testing.T) {
stack := newHeldCleanupStack()
graceReturned := make(chan struct{})
fallbackRan := make(chan struct{})
capabilityCleanupRan := make(chan struct{})
require.NoError(t, stack.Push(heldCleanupAction{name: "capability cleanup", cleanup: func(context.Context) error {
close(capabilityCleanupRan)
return nil
}}))
require.NoError(t, stack.Push(heldCleanupAction{name: "fallback", cleanup: func(context.Context) error {
close(fallbackRan)
return nil
}}))
require.NoError(t, stack.Push(heldCleanupAction{name: "graceful wait", cleanup: func(ctx context.Context) error {
<-ctx.Done()
close(graceReturned)
return ctx.Err()
}}))
cleanupContext, cancel := context.WithTimeout(context.Background(), 2*heldTerminalGracePeriod)
defer cancel()
err := stack.Run(cleanupContext)
require.ErrorIs(t, err, context.DeadlineExceeded)
<-graceReturned
<-fallbackRan
<-capabilityCleanupRan
}
func TestHeldCleanupStackRejectsInvalidAndLateActions(t *testing.T) {
stack := newHeldCleanupStack()
if err := stack.Push(heldCleanupAction{}); !errors.Is(err, ErrInvalidHeldCleanupAction) {
t.Fatalf("invalid push error = %v", err)
}
if err := stack.Run(context.Background()); err != nil {
t.Fatal(err)
}
if err := stack.Push(heldCleanupAction{name: "late", cleanup: func(context.Context) error { return nil }}); !errors.Is(err, ErrHeldCleanupClosed) {
t.Fatalf("late push error = %v", err)
}
if err := stack.Run(context.Background()); !errors.Is(err, ErrHeldCleanupClosed) {
t.Fatalf("second Run error = %v", err)
}
}
@@ -0,0 +1,217 @@
//go:build linux && agentcompat
package scenario
import (
"context"
"errors"
"fmt"
"net/http"
"os"
"path/filepath"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
)
type heldRealNATProfile struct {
ID uint64 `json:"id"`
}
type heldRealFixture struct {
dashboard *dashboard.Dashboard
agent *agent.Agent
dashboardPID int
agentPID int
closed bool
}
func (fixture *heldRealFixture) Close(ctx context.Context, sessionClosed, exactStreamGone, ownedResourceGone bool) (heldRealCleanup, error) {
if fixture.closed {
return heldRealCleanup{}, nil
}
fixture.closed = true
cleanupContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second)
defer cancel()
agentErr := fixture.agent.Stop(cleanupContext)
dashboardErr := fixture.dashboard.Stop(cleanupContext)
cleanup := heldRealCleanup{Agent: fixture.agent.CleanupReceipt(), Dashboard: fixture.dashboard.CleanupReceipt(), SessionClosed: sessionClosed, ExactStreamGone: exactStreamGone, OwnedResourceGone: ownedResourceGone, AgentPIDGone: heldRealPIDGone(fixture.agentPID), DashboardPIDGone: heldRealPIDGone(fixture.dashboardPID)}
return cleanup, errors.Join(agentErr, dashboardErr)
}
func requireHeldRealSources(t *testing.T) {
t.Helper()
if os.Getenv("AGENTCOMPAT_NEZHA_SOURCE") == "" || os.Getenv("AGENTCOMPAT_AGENT_SOURCE") == "" {
t.Skip("set AGENTCOMPAT_NEZHA_SOURCE and AGENTCOMPAT_AGENT_SOURCE")
}
}
func TestHeldLegacyFMSessionUsesExistingDashboardAndAgent(t *testing.T) {
// Given
requireHeldRealSources(t)
paths, err := contract.NewPaths(os.Getenv("AGENTCOMPAT_NEZHA_SOURCE"), os.Getenv("AGENTCOMPAT_AGENT_SOURCE"), t.TempDir())
require.NoError(t, err)
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Minute)
defer cancel()
realFixture, readiness, patClient := startHeldRealFixture(t, ctx, paths, "00000000-0000-0000-0000-000000000118", "held-fm")
dashboardInstance, agentInstance := realFixture.dashboard, realFixture.agent
t.Cleanup(func() { _, _ = realFixture.Close(context.Background(), false, false, false) })
plan := heldFMTestPlan(t, StressSessionFM)
plan.ID, err = NewStressSessionID(fmt.Sprintf("held-fm-real-%d", time.Now().UnixNano()))
require.NoError(t, err)
baseline, err := patClient.IOStreamState(ctx)
require.NoError(t, err)
dashboardPID, agentPID := dashboardInstance.PID(), agentInstance.PID()
// When
sessionCtx, sessionCancel := context.WithTimeout(ctx, 45*time.Second)
defer sessionCancel()
session, err := newHeldLegacyFMSession(sessionCtx, heldLegacyFMInput{Dashboard: dashboardInstance, PATClient: patClient, Agent: agentInstance, Readiness: readiness, Plan: plan})
require.NoError(t, err)
require.NoError(t, session.WaitLive(ctx))
require.True(t, session.ProtocolProved())
streamID, present := session.IOStreamID()
require.True(t, present)
fixtureRoot := filepath.Join(agentInstance.WorkspaceRoot(), "held-fm-"+heldLegacyFMRootName.ReplaceAllString(plan.ID.String(), "-"))
_, err = os.Stat(fixtureRoot)
require.NoError(t, err)
live, err := patClient.WaitForIOStreamState(ctx, client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(baseline.Count + 1), PresentStreamID: streamID})
require.NoError(t, err)
require.Equal(t, baseline.Count+1, live.Count)
require.Equal(t, dashboardPID, dashboardInstance.PID())
require.Equal(t, agentPID, agentInstance.PID())
dashboardPIDUnchanged := dashboardPID == dashboardInstance.PID()
agentPIDUnchanged := agentPID == agentInstance.PID()
// Then
require.NoError(t, session.Close(ctx))
require.NoError(t, session.WaitClosed(ctx))
closed, err := patClient.WaitForIOStreamState(ctx, client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(baseline.Count), AbsentStreamID: streamID})
require.NoError(t, err)
require.Equal(t, baseline.Count, closed.Count)
require.NotEmpty(t, streamID)
require.Equal(t, dashboardPID, dashboardInstance.PID())
require.Equal(t, agentPID, agentInstance.PID())
require.NotZero(t, dashboardPID)
require.NotZero(t, agentPID)
_, err = os.Stat(fixtureRoot)
require.ErrorIs(t, err, os.ErrNotExist)
cleanup, cleanupErr := realFixture.Close(ctx, true, closed.Count == baseline.Count, errors.Is(err, os.ErrNotExist))
require.NoError(t, cleanupErr)
require.True(t, heldRealCleanupOK(cleanup))
require.NoError(t, writeHeldRealEvidence("file-manager", heldRealEvidence{Kind: "file-manager", BaselineCount: baseline.Count, LiveCount: live.Count, ClosedCount: closed.Count, ExactIDPresent: true, ExactIDAbsent: true, ProtocolProved: session.ProtocolProved(), DashboardPIDUnchanged: dashboardPIDUnchanged, AgentPIDUnchanged: agentPIDUnchanged, CleanupOK: heldRealCleanupOK(cleanup)}))
}
func TestHeldNATSessionUsesExistingDashboardAndAgent(t *testing.T) {
// Given
requireHeldRealSources(t)
paths, err := contract.NewPaths(os.Getenv("AGENTCOMPAT_NEZHA_SOURCE"), os.Getenv("AGENTCOMPAT_AGENT_SOURCE"), t.TempDir())
require.NoError(t, err)
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Minute)
defer cancel()
realFixture, readiness, patClient := startHeldRealFixture(t, ctx, paths, "00000000-0000-0000-0000-000000000119", "held-nat")
dashboardInstance, agentInstance := realFixture.dashboard, realFixture.agent
t.Cleanup(func() { _, _ = realFixture.Close(context.Background(), false, false, false) })
plan := heldNATTestPlan(t)
plan.ID, err = NewStressSessionID(fmt.Sprintf("held-nat-real-%d", time.Now().UnixNano()))
require.NoError(t, err)
baseline, err := patClient.IOStreamState(ctx)
require.NoError(t, err)
dashboardPID, agentPID := dashboardInstance.PID(), agentInstance.PID()
// When
sessionCtx, sessionCancel := context.WithTimeout(ctx, 45*time.Second)
defer sessionCancel()
session, err := newHeldNATSession(sessionCtx, heldNATInput{Dashboard: dashboardInstance, PATClient: patClient, Agent: agentInstance, Readiness: readiness, Plan: plan})
require.NoError(t, err)
require.NoError(t, session.WaitLive(ctx))
require.True(t, session.ProtocolProved())
present, err := heldRealNATProfilePresent(ctx, dashboardInstance.Clients().REST, session.profileID)
require.NoError(t, err)
require.True(t, present)
streamID, present := session.IOStreamID()
require.True(t, present)
live, err := patClient.WaitForIOStreamState(ctx, client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(baseline.Count + 1), PresentStreamID: streamID})
require.NoError(t, err)
require.Equal(t, baseline.Count+1, live.Count)
require.Equal(t, dashboardPID, dashboardInstance.PID())
require.Equal(t, agentPID, agentInstance.PID())
dashboardPIDUnchanged := dashboardPID == dashboardInstance.PID()
agentPIDUnchanged := agentPID == agentInstance.PID()
require.Equal(t, http.MethodPatch, session.observed.Method)
require.Equal(t, "/held/"+plan.ID.String(), session.observed.Path)
domain, domainErr := heldNATDomain(plan.ID.String())
require.NoError(t, domainErr)
require.Equal(t, domain, session.observed.Host)
require.Equal(t, plan.ID.String(), session.observed.HeaderValue)
require.Equal(t, "held-body-"+plan.ID.String(), string(session.observed.Body))
require.False(t, session.observed.SensitiveHeadersPresent)
// Then
require.NoError(t, session.Close(ctx))
require.NoError(t, session.WaitClosed(ctx))
closed, err := patClient.WaitForIOStreamState(ctx, client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(baseline.Count), AbsentStreamID: streamID})
require.NoError(t, err)
require.Equal(t, baseline.Count, closed.Count)
require.NotEmpty(t, streamID)
require.Equal(t, dashboardPID, dashboardInstance.PID())
require.Equal(t, agentPID, agentInstance.PID())
require.NotZero(t, dashboardPID)
require.NotZero(t, agentPID)
require.False(t, session.observed.SensitiveHeadersPresent)
present, err = heldRealNATProfilePresent(ctx, dashboardInstance.Clients().REST, session.profileID)
require.NoError(t, err)
require.False(t, present)
cleanup, cleanupErr := realFixture.Close(ctx, true, closed.Count == baseline.Count, !present)
require.NoError(t, cleanupErr)
require.True(t, heldRealCleanupOK(cleanup))
require.NoError(t, writeHeldRealEvidence("nat", heldRealEvidence{Kind: "nat", BaselineCount: baseline.Count, LiveCount: live.Count, ClosedCount: closed.Count, ExactIDPresent: true, ExactIDAbsent: true, ProtocolProved: session.ProtocolProved(), SensitiveHeadersPresent: session.observed.SensitiveHeadersPresent, DashboardPIDUnchanged: dashboardPIDUnchanged, AgentPIDUnchanged: agentPIDUnchanged, CleanupOK: heldRealCleanupOK(cleanup)}))
}
func heldRealNATProfilePresent(ctx context.Context, admin *client.Client, profileID uint64) (bool, error) {
return heldRealNATProfilePresentWithQuery(ctx, profileID, func(queryContext context.Context) ([]heldRealNATProfile, error) {
return client.DoREST[struct{}, []heldRealNATProfile](queryContext, admin, client.RESTRequest[struct{}]{Method: http.MethodGet, Path: "/api/v1/nat"})
})
}
func heldRealNATProfilePresentWithQuery(ctx context.Context, profileID uint64, query func(context.Context) ([]heldRealNATProfile, error)) (bool, error) {
profiles, err := query(ctx)
if err != nil {
return false, err
}
for _, profile := range profiles {
if profile.ID == profileID {
return true, nil
}
}
return false, nil
}
func startHeldRealFixture(t *testing.T, ctx context.Context, paths contract.Paths, uuid, name string) (*heldRealFixture, agent.Readiness, *client.Client) {
t.Helper()
dashboardInstance, err := dashboard.Start(ctx, dashboard.StartConfig{SourceDir: paths.NezhaSource().String(), ReceiptGate: true})
require.NoError(t, err)
agentInstance, err := agent.Start(ctx, agent.AgentStartConfig{SourceDir: paths.AgentSource().String(), Endpoint: dashboardInstance.Endpoint(), Secret: dashboardInstance.AgentSecret(), UUID: uuid})
require.NoError(t, err)
require.NoError(t, dashboardInstance.WaitForReceiptAccepted(ctx))
require.NoError(t, dashboardInstance.ReleaseReceipt(ctx))
readiness, err := agentInstance.WaitReady(ctx, dashboardInstance)
require.NoError(t, err)
patClient, err := createTerminalPATClient(ctx, dashboardInstance, name, []string{"nezha:*"}, []uint64{readiness.ServerID})
require.NoError(t, err)
return &heldRealFixture{dashboard: dashboardInstance, agent: agentInstance, dashboardPID: dashboardInstance.PID(), agentPID: agentInstance.PID()}, readiness, patClient
}
func heldRealPIDGone(pid int) bool {
if pid < 1 {
return false
}
_, err := os.Stat(filepath.Join("/proc", fmt.Sprint(pid)))
return errors.Is(err, os.ErrNotExist)
}
@@ -0,0 +1,127 @@
//go:build linux
package scenario
import (
"context"
"errors"
"sync"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
type heldIOStreamCapabilityIdentity struct {
Purpose client.IOStreamCapabilityPurpose
ServerID uint64
ResourceID uint64
}
type heldIOStreamCapability struct {
client client.IOStreamCapabilityClient
identity heldIOStreamCapabilityIdentity
access client.IOStreamCapabilityAccessRequest
streamID string
mu sync.Mutex
waitOnce sync.Once
waitErr error
cancelOnce sync.Once
unregisterOnce sync.Once
cancelErr error
unregisterErr error
}
func registerHeldIOStreamCapability(ctx context.Context, transport *client.Client, identity heldIOStreamCapabilityIdentity) (*heldIOStreamCapability, error) {
registered, err := transport.IOStreamCapabilities().Register(ctx, client.IOStreamCapabilityRegisterRequest{
Purpose: identity.Purpose, ServerID: identity.ServerID, ResourceID: identity.ResourceID,
})
if err != nil {
return nil, err
}
return &heldIOStreamCapability{
client: transport.IOStreamCapabilities(), identity: identity,
access: client.IOStreamCapabilityAccessRequest{Capability: registered.Capability, Purpose: identity.Purpose, ServerID: identity.ServerID, ResourceID: identity.ResourceID},
}, nil
}
func (capability *heldIOStreamCapability) HeaderCapability() client.IOStreamCapability {
return capability.access.Capability
}
func (capability *heldIOStreamCapability) Wait(ctx context.Context) (string, error) {
capability.waitOnce.Do(func() {
response, err := capability.client.Wait(ctx, client.IOStreamCapabilityWaitRequest(capability.access))
if err != nil {
capability.waitErr = err
return
}
capability.streamID = response.StreamID.Value()
if capability.streamID == "" {
capability.waitErr = client.ErrIOStreamCapabilityUnavailable
}
})
capability.mu.Lock()
defer capability.mu.Unlock()
if capability.waitErr != nil {
return "", capability.waitErr
}
streamID := capability.streamID
return streamID, nil
}
func (capability *heldIOStreamCapability) Cancel(ctx context.Context) error {
capability.cancelOnce.Do(func() {
capability.mu.Lock()
defer capability.mu.Unlock()
capability.cancelErr = capability.client.Cancel(ctx, capability.access)
})
capability.mu.Lock()
defer capability.mu.Unlock()
return capability.cancelErr
}
func (capability *heldIOStreamCapability) Unregister(ctx context.Context) error {
capability.unregisterOnce.Do(func() {
capability.mu.Lock()
defer capability.mu.Unlock()
capability.unregisterErr = capability.client.Unregister(ctx, capability.access)
})
capability.mu.Lock()
defer capability.mu.Unlock()
return capability.unregisterErr
}
func (capability *heldIOStreamCapability) streamIDValue() string {
capability.mu.Lock()
defer capability.mu.Unlock()
return capability.streamID
}
func (capability *heldIOStreamCapability) waitExpectation(ctx context.Context, stateClient *client.Client, _ client.IOStreamState, absent bool) error {
streamID := capability.streamIDValue()
// Adapter-local ownership must not couple to the shared global count during concurrent construction.
expectation := client.IOStreamStateExpectation{}
if absent {
if streamID != "" {
expectation.AbsentStreamID = streamID
} else {
if _, waitErr := capability.Wait(ctx); waitErr == nil {
streamID = capability.streamIDValue()
expectation.AbsentStreamID = streamID
} else if ctx.Err() != nil {
return ctx.Err()
}
}
} else {
if streamID == "" {
return errors.New("held capability stream ID is missing")
}
expectation.PresentStreamID = streamID
}
_, err := stateClient.WaitForIOStreamState(ctx, expectation)
return err
}
func (capability *heldIOStreamCapability) WaitExpectation(ctx context.Context, stateClient *client.Client, baseline client.IOStreamState, absent bool) error {
return capability.waitExpectation(ctx, stateClient, baseline, absent)
}
@@ -0,0 +1,50 @@
//go:build linux
package scenario
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
func TestHeldIOStreamCapabilityWaitExpectationScopesToOwnedStream(t *testing.T) {
tests := []struct {
name string
absent bool
presentID string
absentID string
}{
{name: "live", presentID: "owned-stream"},
{name: "cleanup", absent: true, absentID: "owned-stream"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
var observed client.IOStreamStateExpectation
server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
require.NoError(t, json.NewDecoder(request.Body).Decode(&observed))
response.Header().Set("Content-Type", "application/json")
_, err := response.Write([]byte(`{"success":true,"data":{"count":99,"generation":1}}`))
require.NoError(t, err)
}))
defer server.Close()
stateClient, err := client.New(client.Config{BaseURL: server.URL})
require.NoError(t, err)
capability := &heldIOStreamCapability{streamID: "owned-stream"}
err = capability.waitExpectation(context.Background(), stateClient, client.IOStreamState{Count: 7}, test.absent)
require.NoError(t, err)
require.Nil(t, observed.ExpectedCount)
require.Equal(t, test.presentID, observed.PresentStreamID)
require.Equal(t, test.absentID, observed.AbsentStreamID)
})
}
}
@@ -0,0 +1,233 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
"regexp"
"strings"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
const (
heldLegacyFMCleanupTimeout = 10 * time.Second
heldLegacyFMPumpCapacity = 8
)
var (
ErrInvalidHeldLegacyFMInput = errors.New("held legacy FM input is invalid")
ErrHeldLegacyFMProtocol = errors.New("held legacy FM protocol proof failed")
heldLegacyFMRootName = regexp.MustCompile(`[^a-z0-9]+`)
)
type heldLegacyFMInput struct {
Dashboard *dashboard.Dashboard
PATClient *client.Client
Agent *agent.Agent
Readiness agent.Readiness
Plan StressSessionPlan
LifetimeContext context.Context
}
type heldLegacyFMSession struct {
lifecycle *heldSessionLifecycle
stack *heldCleanupStack
connection heldLegacyFMConnection
pump heldLegacyFMPump
protocol bool
}
func newHeldLegacyFMSession(ctx context.Context, input heldLegacyFMInput) (*heldLegacyFMSession, error) {
return newHeldLegacyFMSessionWithDependencies(ctx, input, defaultHeldLegacyFMDependencies())
}
func newHeldLegacyFMSessionWithDependencies(ctx context.Context, input heldLegacyFMInput, dependencies heldLegacyFMDependencies) (*heldLegacyFMSession, error) {
if err := validateHeldLegacyFMInput(ctx, input); err != nil {
return nil, err
}
if err := validateHeldPATClient(input.PATClient); err != nil {
return nil, err
}
if err := validateHeldReadiness(input.Agent, input.Readiness); err != nil {
return nil, err
}
stateClient := input.PATClient
baseline, err := dependencies.SnapshotState(ctx, stateClient)
if err != nil {
return nil, fmt.Errorf("snapshot FM IOStream state: %w", err)
}
rootName := heldLegacyFMRootName.ReplaceAllString(strings.ToLower(input.Plan.ID.String()), "-")
rootName = strings.Trim(rootName, "-")
if rootName == "" {
return nil, ErrInvalidHeldLegacyFMInput
}
root, err := fixture.NewAgentRoot(input.Agent.WorkspaceRoot(), "held-fm-"+rootName)
if err != nil {
return nil, fmt.Errorf("create held FM fixture root: %w", err)
}
stack := newHeldCleanupStack()
if err := stack.Push(heldCleanupAction{name: "remove FM fixture root", cleanup: func(cleanupContext context.Context) error {
return dependencies.RemoveFixture(cleanupContext, input.Agent.WorkspaceRoot(), root.Absolute())
}}); err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
listDirectory, err := root.Path("list")
if err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
if err := os.Mkdir(listDirectory.String(), 0o700); err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
if err := os.WriteFile(filepath.Join(listDirectory.String(), "entry.txt"), []byte("entry"), 0o600); err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
capability, err := dependencies.Register(ctx, input.PATClient, heldIOStreamCapabilityIdentity{Purpose: client.IOStreamCapabilityPurposeFileManager, ServerID: input.Readiness.ServerID})
if err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
if err := pushHeldLegacyFMCapabilityCleanup(stack, heldLegacyFMCapabilityCleanup{
Unregister: capability.Unregister,
Absence: func(cleanupContext context.Context) error {
return dependencies.WaitForState(cleanupContext, stateClient, baseline, capability, true)
},
Cancel: capability.Cancel,
}); err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
sessionID, err := dependencies.CreateSession(ctx, input.PATClient, input.Readiness.ServerID, capability.HeaderCapability())
if err != nil {
_, waitErr := capability.Wait(ctx)
return nil, rollbackHeldLegacyFM(ctx, stack, errors.Join(err, waitErr))
}
streamID, waitErr := capability.Wait(ctx)
if waitErr != nil || streamID != sessionID {
mismatchErr := error(nil)
if waitErr == nil {
mismatchErr = heldLegacyFMStreamMismatchError()
}
return nil, rollbackHeldLegacyFM(ctx, stack, errors.Join(waitErr, mismatchErr))
}
lifetimeContext := input.LifetimeContext
if lifetimeContext == nil {
lifetimeContext = ctx
}
lifecycle, err := newHeldSessionLifecycle(lifetimeContext, input.Plan, sessionID, heldLegacyFMCleanupTimeout)
if err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
connection, err := dependencies.DialWebSocket(ctx, input.PATClient, "/api/v1/ws/file/"+sessionID)
if err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
if err := stack.Push(heldCleanupAction{name: "close FM WebSocket", cleanup: func(context.Context) error { return connection.Close() }}); err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
pump, err := dependencies.NewPump(lifetimeContext, connection, heldLegacyFMPumpCapacity)
if err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
if err := stack.Push(heldCleanupAction{name: "stop FM WebSocket pump", cleanup: pump.Stop}); err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
dispatcher := legacyFMCommandDispatcher{writer: connection, root: root}
if err := dispatcher.list(ctx, "list"); err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
if err := proveHeldLegacyFMList(ctx, pump, listDirectory.String()); err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
if err := dependencies.WaitForState(ctx, stateClient, baseline, capability, false); err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
if err := lifecycle.markLive(nil); err != nil {
return nil, rollbackHeldLegacyFM(ctx, stack, err)
}
return &heldLegacyFMSession{lifecycle: lifecycle, stack: stack, connection: connection, pump: pump, protocol: true}, nil
}
func heldLegacyFMStreamMismatchError() error {
return fmt.Errorf("FM stream identity mismatch: %w", ErrHeldLegacyFMProtocol)
}
func validateHeldLegacyFMInput(ctx context.Context, input heldLegacyFMInput) error {
if ctx == nil || input.Dashboard == nil || input.Agent == nil || input.Plan.Kind != StressSessionFM || input.Plan.ID.String() == "" || input.Plan.Ordinal < 1 || input.Plan.Agent.Int() < 1 {
return ErrInvalidHeldLegacyFMInput
}
if err := validateHeldPATClient(input.PATClient); err != nil {
return err
}
return nil
}
func proveHeldLegacyFMList(ctx context.Context, pump heldLegacyFMPump, wantPath string) error {
select {
case frame, ok := <-pump.Events():
if !ok {
return errors.Join(pump.Err(), ErrHeldLegacyFMProtocol)
}
if frame.Type != client.FrameBinary {
return fmt.Errorf("FM list response frame type=%s: %w", frame.Type, ErrHeldLegacyFMProtocol)
}
parsed, err := parseLegacyFMList(frame.Payload)
if err != nil {
return err
}
if parsed.Path != wantPath || len(parsed.Entries) != 1 || parsed.Entries[0].Name != "entry.txt" || parsed.Entries[0].Dir {
return ErrHeldLegacyFMProtocol
}
return nil
case <-pump.Done():
return errors.Join(pump.Err(), ErrHeldLegacyFMProtocol)
case <-ctx.Done():
return ctx.Err()
}
}
func rollbackHeldLegacyFM(ctx context.Context, stack *heldCleanupStack, original error) error {
rollbackContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), heldLegacyFMCleanupTimeout)
defer cancel()
return errors.Join(original, stack.Run(rollbackContext))
}
func (session *heldLegacyFMSession) Plan() StressSessionPlan { return session.lifecycle.Plan() }
func (session *heldLegacyFMSession) WaitLive(ctx context.Context) error {
return session.lifecycle.WaitLive(ctx)
}
func (session *heldLegacyFMSession) IOStreamID() (string, bool) {
return session.lifecycle.IOStreamID()
}
func (session *heldLegacyFMSession) ProtocolProved() bool { return session.protocol }
func (session *heldLegacyFMSession) WaitClosed(ctx context.Context) error {
return session.lifecycle.WaitClosed(ctx)
}
func (session *heldLegacyFMSession) Done() <-chan struct{} { return session.lifecycle.Done() }
func (session *heldLegacyFMSession) CloseResult() error { return session.lifecycle.CloseResult() }
func (session *heldLegacyFMSession) Close(ctx context.Context) error {
owner, won := session.lifecycle.beginClose()
if !won {
return session.lifecycle.WaitClosed(ctx)
}
go func() {
cleanupContext, cancel := owner.cleanupContext()
cleanupErr := session.stack.Run(cleanupContext)
if cleanupContext.Err() != nil {
cleanupErr = errors.Join(cleanupErr, cleanupContext.Err())
}
cancel()
owner.markClosed(cleanupErr)
}()
return session.lifecycle.WaitClosed(ctx)
}
var _ heldSession = (*heldLegacyFMSession)(nil)
@@ -0,0 +1,21 @@
//go:build linux
package scenario
import "context"
type heldLegacyFMCapabilityCleanup struct {
Unregister func(context.Context) error
Absence func(context.Context) error
Cancel func(context.Context) error
}
func pushHeldLegacyFMCapabilityCleanup(stack *heldCleanupStack, cleanup heldLegacyFMCapabilityCleanup) error {
if err := stack.Push(heldCleanupAction{name: "unregister FM capability", cleanup: cleanup.Unregister}); err != nil {
return err
}
if err := stack.Push(heldCleanupAction{name: "restore FM IOStream baseline and absence", cleanup: cleanup.Absence}); err != nil {
return err
}
return stack.Push(heldCleanupAction{name: "cancel FM capability", cleanup: cleanup.Cancel})
}
@@ -0,0 +1,204 @@
//go:build linux
package scenario
import (
"context"
"errors"
"path/filepath"
"testing"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
)
type heldLegacyFMConstructorFault struct {
name string
failStage string
wantError error
}
type heldLegacyFMConstructorObservation struct {
order []string
expectation client.IOStreamStateExpectation
}
func TestNewHeldLegacyFMSessionConstructorFaultsRollbackInLIFOOrder(t *testing.T) {
tests := []heldLegacyFMConstructorFault{
{name: "create response", failStage: "create", wantError: errHeldLegacyFMConstructorCreate},
{name: "capability wait", failStage: "wait", wantError: errHeldLegacyFMConstructorWait},
{name: "response and wait mismatch", failStage: "mismatch", wantError: ErrHeldLegacyFMProtocol},
{name: "WebSocket dial", failStage: "dial", wantError: errHeldLegacyFMConstructorDial},
{name: "pump setup", failStage: "pump", wantError: errHeldLegacyFMConstructorPump},
{name: "list proof", failStage: "proof", wantError: errLegacyFMRemote},
}
for _, testCase := range tests {
t.Run(testCase.name, func(t *testing.T) {
input := heldLegacyFMConstructorInput(t)
observation := &heldLegacyFMConstructorObservation{}
dependencies := heldLegacyFMConstructorDependencies(input, observation, testCase.failStage)
_, err := newHeldLegacyFMSessionWithDependencies(context.Background(), input, dependencies)
require.ErrorIs(t, err, testCase.wantError)
for _, cleanupError := range []error{errHeldLegacyFMConstructorCancel, errHeldLegacyFMConstructorAbsence, errHeldLegacyFMConstructorUnregister, errHeldLegacyFMConstructorFixture} {
require.ErrorIs(t, err, cleanupError)
}
wantOrder := []string{"cancel", "absence", "unregister", "fixture"}
if testCase.failStage == "pump" {
wantOrder = []string{"close", "cancel", "absence", "unregister", "fixture"}
require.ErrorIs(t, err, errHeldLegacyFMConstructorClose)
}
if testCase.failStage == "proof" {
wantOrder = []string{"pump", "close", "cancel", "absence", "unregister", "fixture"}
require.ErrorIs(t, err, errHeldLegacyFMConstructorPumpStop)
require.ErrorIs(t, err, errHeldLegacyFMConstructorClose)
}
require.Equal(t, wantOrder, observation.order)
require.Nil(t, observation.expectation.ExpectedCount)
expectedAbsent := "stream-legacy-fm"
if testCase.failStage == "mismatch" {
expectedAbsent = "different-stream"
}
require.Equal(t, expectedAbsent, observation.expectation.AbsentStreamID)
})
}
}
var (
errHeldLegacyFMConstructorCreate = errors.New("held FM constructor create failed")
errHeldLegacyFMConstructorWait = errors.New("held FM constructor wait failed")
errHeldLegacyFMConstructorDial = errors.New("held FM constructor dial failed")
errHeldLegacyFMConstructorPump = errors.New("held FM constructor pump failed")
errHeldLegacyFMConstructorClose = errors.New("held FM constructor close cleanup failed")
errHeldLegacyFMConstructorPumpStop = errors.New("held FM constructor pump stop cleanup failed")
errHeldLegacyFMConstructorCancel = errors.New("held FM constructor cancel cleanup failed")
errHeldLegacyFMConstructorAbsence = errors.New("held FM constructor absence cleanup failed")
errHeldLegacyFMConstructorUnregister = errors.New("held FM constructor unregister cleanup failed")
errHeldLegacyFMConstructorFixture = errors.New("held FM constructor fixture cleanup failed")
)
type heldLegacyFMConstructorConnection struct {
order *[]string
}
func (connection *heldLegacyFMConstructorConnection) WriteFrame(context.Context, client.Frame) error {
return nil
}
func (connection *heldLegacyFMConstructorConnection) Close() error {
*connection.order = append(*connection.order, "close")
return errHeldLegacyFMConstructorClose
}
type heldLegacyFMConstructorPump struct {
events chan client.Frame
order *[]string
}
func (pump *heldLegacyFMConstructorPump) Events() <-chan client.Frame { return pump.events }
func (pump *heldLegacyFMConstructorPump) Done() <-chan struct{} { return make(chan struct{}) }
func (pump *heldLegacyFMConstructorPump) Err() error { return nil }
func (pump *heldLegacyFMConstructorPump) Stop(context.Context) error {
*pump.order = append(*pump.order, "pump")
return errHeldLegacyFMConstructorPumpStop
}
type heldLegacyFMConstructorCapability struct {
observation *heldLegacyFMConstructorObservation
streamID string
waitErr error
}
func (capability *heldLegacyFMConstructorCapability) HeaderCapability() client.IOStreamCapability {
parsed, _ := client.ParseIOStreamCapability("AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8")
return parsed
}
func (capability *heldLegacyFMConstructorCapability) Wait(context.Context) (string, error) {
return capability.streamID, capability.waitErr
}
func (capability *heldLegacyFMConstructorCapability) Cancel(context.Context) error {
capability.observation.order = append(capability.observation.order, "cancel")
return errHeldLegacyFMConstructorCancel
}
func (capability *heldLegacyFMConstructorCapability) Unregister(context.Context) error {
capability.observation.order = append(capability.observation.order, "unregister")
return errHeldLegacyFMConstructorUnregister
}
func (capability *heldLegacyFMConstructorCapability) WaitExpectation(_ context.Context, _ *client.Client, _ client.IOStreamState, absent bool) error {
if absent {
capability.observation.order = append(capability.observation.order, "absence")
capability.observation.expectation = client.IOStreamStateExpectation{AbsentStreamID: capability.streamID}
}
return errHeldLegacyFMConstructorAbsence
}
func heldLegacyFMConstructorDependencies(input heldLegacyFMInput, observation *heldLegacyFMConstructorObservation, failStage string) heldLegacyFMDependencies {
listPath := filepath.Join(input.Agent.WorkspaceRoot(), "held-fm-"+input.Plan.ID.String(), "list")
defaultDependencies := defaultHeldLegacyFMDependencies()
return heldLegacyFMDependencies{
RemoveFixture: func(ctx context.Context, workspaceRoot, fixtureRoot string) error {
observation.order = append(observation.order, "fixture")
_ = defaultDependencies.RemoveFixture(ctx, workspaceRoot, fixtureRoot)
return errHeldLegacyFMConstructorFixture
},
SnapshotState: func(context.Context, *client.Client) (client.IOStreamState, error) {
return client.IOStreamState{Count: 4}, nil
},
Register: func(context.Context, *client.Client, heldIOStreamCapabilityIdentity) (heldLegacyFMCapabilityHandle, error) {
capability := &heldLegacyFMConstructorCapability{observation: observation, streamID: "stream-legacy-fm"}
if failStage == "wait" {
capability.waitErr = errHeldLegacyFMConstructorWait
}
if failStage == "mismatch" {
capability.streamID = "different-stream"
}
return capability, nil
},
CreateSession: func(context.Context, *client.Client, uint64, client.IOStreamCapability) (string, error) {
if failStage == "create" {
return "", errHeldLegacyFMConstructorCreate
}
return "stream-legacy-fm", nil
},
DialWebSocket: func(context.Context, *client.Client, string) (heldLegacyFMConnection, error) {
if failStage == "dial" {
return nil, errHeldLegacyFMConstructorDial
}
return &heldLegacyFMConstructorConnection{order: &observation.order}, nil
},
NewPump: func(context.Context, heldLegacyFMConnection, int) (heldLegacyFMPump, error) {
if failStage == "pump" {
return nil, errHeldLegacyFMConstructorPump
}
frame := heldLegacyFMListFrame(listPath, "entry.txt", false)
if failStage == "proof" {
frame = []byte("NERRdenied")
}
events := make(chan client.Frame, 1)
events <- client.Frame{Type: client.FrameBinary, Payload: frame}
return &heldLegacyFMConstructorPump{events: events, order: &observation.order}, nil
},
WaitForState: func(ctx context.Context, stateClient *client.Client, baseline client.IOStreamState, capability heldLegacyFMCapabilityHandle, absent bool) error {
return capability.WaitExpectation(ctx, stateClient, baseline, absent)
},
}
}
func heldLegacyFMConstructorInput(t *testing.T) heldLegacyFMInput {
t.Helper()
agentInstance := newHeldReadinessTestAgent(t, "00000000-0000-0000-0000-000000000501")
return heldLegacyFMInput{
Dashboard: &dashboard.Dashboard{},
PATClient: &client.Client{},
Agent: agentInstance,
Readiness: completeHeldReadiness(agentInstance.UUID()),
Plan: heldFMTestPlan(t, StressSessionFM),
}
}
@@ -0,0 +1,84 @@
//go:build linux
package scenario
import (
"context"
"errors"
"os"
"path/filepath"
"strings"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
type heldLegacyFMConnection interface {
legacyFMFrameWriter
Close() error
}
type heldLegacyFMPump interface {
Events() <-chan client.Frame
Done() <-chan struct{}
Err() error
Stop(context.Context) error
}
type heldLegacyFMCapabilityHandle interface {
HeaderCapability() client.IOStreamCapability
Wait(context.Context) (string, error)
Cancel(context.Context) error
Unregister(context.Context) error
WaitExpectation(context.Context, *client.Client, client.IOStreamState, bool) error
}
type heldLegacyFMDependencies struct {
SnapshotState func(context.Context, *client.Client) (client.IOStreamState, error)
Register func(context.Context, *client.Client, heldIOStreamCapabilityIdentity) (heldLegacyFMCapabilityHandle, error)
CreateSession func(context.Context, *client.Client, uint64, client.IOStreamCapability) (string, error)
DialWebSocket func(context.Context, *client.Client, string) (heldLegacyFMConnection, error)
NewPump func(context.Context, heldLegacyFMConnection, int) (heldLegacyFMPump, error)
WaitForState func(context.Context, *client.Client, client.IOStreamState, heldLegacyFMCapabilityHandle, bool) error
RemoveFixture func(context.Context, string, string) error
}
func defaultHeldLegacyFMDependencies() heldLegacyFMDependencies {
return heldLegacyFMDependencies{
SnapshotState: func(ctx context.Context, transport *client.Client) (client.IOStreamState, error) {
return transport.IOStreamState(ctx)
},
Register: func(ctx context.Context, transport *client.Client, identity heldIOStreamCapabilityIdentity) (heldLegacyFMCapabilityHandle, error) {
return registerHeldIOStreamCapability(ctx, transport, identity)
},
CreateSession: func(ctx context.Context, transport *client.Client, serverID uint64, capability client.IOStreamCapability) (string, error) {
return createLegacyFMSession(ctx, transport, serverID, capability)
},
DialWebSocket: func(ctx context.Context, transport *client.Client, path string) (heldLegacyFMConnection, error) {
return transport.DialWebSocket(ctx, path)
},
NewPump: func(ctx context.Context, connection heldLegacyFMConnection, capacity int) (heldLegacyFMPump, error) {
concrete, ok := connection.(*client.WebSocketConnection)
if !ok {
return nil, ErrInvalidHeldFramePump
}
return newHeldWebSocketPump(ctx, concrete, capacity)
},
WaitForState: func(ctx context.Context, transport *client.Client, baseline client.IOStreamState, capability heldLegacyFMCapabilityHandle, absent bool) error {
return capability.WaitExpectation(ctx, transport, baseline, absent)
},
RemoveFixture: removeHeldLegacyFMFixture,
}
}
func removeHeldLegacyFMFixture(ctx context.Context, workspaceRoot, fixtureRoot string) error {
if err := ctx.Err(); err != nil {
return err
}
cleanWorkspace := filepath.Clean(workspaceRoot)
cleanFixture := filepath.Clean(fixtureRoot)
relative, err := filepath.Rel(cleanWorkspace, cleanFixture)
if err != nil || filepath.IsAbs(relative) || relative == "." || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) {
return errors.New("FM fixture root escaped Agent workspace")
}
return os.RemoveAll(cleanFixture)
}
@@ -0,0 +1,160 @@
//go:build linux
package scenario
import (
"context"
"errors"
"testing"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
)
var (
errHeldLegacyFMCreateRejected = errors.New("held FM create rejected")
errHeldLegacyFMResponseLost = errors.New("held FM response lost")
errHeldLegacyFMCapabilityWaitLost = errors.New("held FM capability wait lost")
)
type heldLegacyFMResponseLossCase struct {
name string
createError error
waitStreamID string
waitError error
wantAbsentStream string
wantCreateError error
wantWaitError error
}
func TestNewHeldLegacyFMSessionRecoversCreateResponseLossForExactCleanup(t *testing.T) {
tests := []heldLegacyFMResponseLossCase{
{
name: "ordinary rejection",
createError: errHeldLegacyFMCreateRejected,
waitError: client.ErrIOStreamCapabilityUnavailable,
wantCreateError: errHeldLegacyFMCreateRejected,
},
{
name: "response loss after stream creation",
createError: errHeldLegacyFMResponseLost,
waitStreamID: "stream-response-lost",
wantAbsentStream: "stream-response-lost",
wantCreateError: errHeldLegacyFMResponseLost,
},
{
name: "create and capability wait failure",
createError: errHeldLegacyFMResponseLost,
waitError: errHeldLegacyFMCapabilityWaitLost,
wantCreateError: errHeldLegacyFMResponseLost,
wantWaitError: errHeldLegacyFMCapabilityWaitLost,
},
}
for _, testCase := range tests {
t.Run(testCase.name, func(t *testing.T) {
input := heldLegacyFMResponseLossInput(t)
fixture := newHeldLegacyFMResponseLossFixture(testCase)
dependencies := heldLegacyFMResponseLossDependencies(input, fixture)
_, err := newHeldLegacyFMSessionWithDependencies(context.Background(), input, dependencies)
require.ErrorIs(t, err, testCase.wantCreateError)
if testCase.wantWaitError != nil {
require.ErrorIs(t, err, testCase.wantWaitError)
}
require.Equal(t, 1, fixture.waitCalls)
require.Equal(t, testCase.wantAbsentStream, fixture.absenceExpectation.AbsentStreamID)
require.Nil(t, fixture.absenceExpectation.ExpectedCount)
require.Equal(t, []string{"cancel", "absence", "unregister", "fixture"}, fixture.order)
})
}
}
type heldLegacyFMResponseLossFixture struct {
caseData heldLegacyFMResponseLossCase
order []string
waitCalls int
absenceExpectation client.IOStreamStateExpectation
}
func newHeldLegacyFMResponseLossFixture(caseData heldLegacyFMResponseLossCase) *heldLegacyFMResponseLossFixture {
return &heldLegacyFMResponseLossFixture{caseData: caseData}
}
type heldLegacyFMResponseLossCapability struct {
fixture *heldLegacyFMResponseLossFixture
waitOnce bool
streamID string
waitErr error
}
func (capability *heldLegacyFMResponseLossCapability) HeaderCapability() client.IOStreamCapability {
parsed, _ := client.ParseIOStreamCapability("AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8")
return parsed
}
func (capability *heldLegacyFMResponseLossCapability) Wait(context.Context) (string, error) {
if capability.waitOnce {
return capability.streamID, capability.waitErr
}
capability.waitOnce = true
capability.fixture.waitCalls++
capability.streamID = capability.fixture.caseData.waitStreamID
capability.waitErr = capability.fixture.caseData.waitError
return capability.streamID, capability.waitErr
}
func (capability *heldLegacyFMResponseLossCapability) Cancel(context.Context) error {
capability.fixture.order = append(capability.fixture.order, "cancel")
return nil
}
func (capability *heldLegacyFMResponseLossCapability) Unregister(context.Context) error {
capability.fixture.order = append(capability.fixture.order, "unregister")
return nil
}
func (capability *heldLegacyFMResponseLossCapability) WaitExpectation(_ context.Context, _ *client.Client, _ client.IOStreamState, absent bool) error {
if absent {
capability.fixture.order = append(capability.fixture.order, "absence")
streamID := capability.streamID
capability.fixture.absenceExpectation = client.IOStreamStateExpectation{AbsentStreamID: streamID}
}
return nil
}
func heldLegacyFMResponseLossDependencies(input heldLegacyFMInput, fixture *heldLegacyFMResponseLossFixture) heldLegacyFMDependencies {
defaults := defaultHeldLegacyFMDependencies()
return heldLegacyFMDependencies{
SnapshotState: func(context.Context, *client.Client) (client.IOStreamState, error) {
return client.IOStreamState{Count: 7}, nil
},
Register: func(context.Context, *client.Client, heldIOStreamCapabilityIdentity) (heldLegacyFMCapabilityHandle, error) {
return &heldLegacyFMResponseLossCapability{fixture: fixture}, nil
},
CreateSession: func(context.Context, *client.Client, uint64, client.IOStreamCapability) (string, error) {
return "", fixture.caseData.createError
},
WaitForState: func(ctx context.Context, stateClient *client.Client, baseline client.IOStreamState, capability heldLegacyFMCapabilityHandle, absent bool) error {
return capability.WaitExpectation(ctx, stateClient, baseline, absent)
},
RemoveFixture: func(ctx context.Context, workspaceRoot, fixtureRoot string) error {
fixture.order = append(fixture.order, "fixture")
return defaults.RemoveFixture(ctx, workspaceRoot, fixtureRoot)
},
}
}
func heldLegacyFMResponseLossInput(t *testing.T) heldLegacyFMInput {
t.Helper()
agentInstance := newHeldReadinessTestAgent(t, "00000000-0000-0000-0000-000000000601")
return heldLegacyFMInput{
Dashboard: &dashboard.Dashboard{},
PATClient: &client.Client{},
Agent: agentInstance,
Readiness: completeHeldReadiness(agentInstance.UUID()),
Plan: heldFMTestPlan(t, StressSessionFM),
}
}
@@ -0,0 +1,254 @@
//go:build linux
package scenario
import (
"context"
"encoding/binary"
"os"
"path/filepath"
"testing"
"time"
"github.com/gorilla/websocket"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
)
func TestNewHeldLegacyFMSessionRejectsInvalidTypedInputs(t *testing.T) {
// Given
plan := heldFMTestPlan(t, StressSessionFM)
agentInstance := newHeldReadinessTestAgent(t, "00000000-0000-0000-0000-000000000305")
readiness := completeHeldReadiness(agentInstance.UUID())
valid := heldLegacyFMInput{Dashboard: &dashboard.Dashboard{}, PATClient: &client.Client{}, Agent: agentInstance, Readiness: readiness, Plan: plan}
cases := []struct {
name string
input heldLegacyFMInput
}{
{name: "nil dashboard", input: heldLegacyFMInput{Agent: valid.Agent, Readiness: readiness, Plan: plan}},
{name: "nil agent", input: heldLegacyFMInput{Dashboard: valid.Dashboard, Readiness: readiness, Plan: plan}},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
// When
_, err := newHeldLegacyFMSession(context.Background(), testCase.input)
// Then
require.ErrorIs(t, err, ErrInvalidHeldLegacyFMInput)
})
}
}
func TestNewHeldLegacyFMSessionReturnsPreciseReadinessErrorForZeroServerID(t *testing.T) {
agentInstance := newHeldReadinessTestAgent(t, "00000000-0000-0000-0000-000000000306")
input := heldLegacyFMInput{Dashboard: &dashboard.Dashboard{}, PATClient: &client.Client{}, Agent: agentInstance, Readiness: completeHeldReadiness(agentInstance.UUID()), Plan: heldFMTestPlan(t, StressSessionFM)}
input.Readiness.ServerID = 0
_, err := newHeldLegacyFMSession(context.Background(), input)
require.ErrorIs(t, err, ErrHeldReadinessServerID)
require.NotErrorIs(t, err, ErrInvalidHeldLegacyFMInput)
}
func TestNewHeldLegacyFMSessionReturnsPreciseReadinessErrorForUUIDMismatch(t *testing.T) {
agentInstance := newHeldReadinessTestAgent(t, "00000000-0000-0000-0000-000000000307")
input := heldLegacyFMInput{Dashboard: &dashboard.Dashboard{}, PATClient: &client.Client{}, Agent: agentInstance, Readiness: completeHeldReadiness("00000000-0000-0000-0000-000000000308"), Plan: heldFMTestPlan(t, StressSessionFM)}
_, err := newHeldLegacyFMSession(context.Background(), input)
require.ErrorIs(t, err, ErrHeldReadinessAgentMismatch)
require.NotErrorIs(t, err, ErrInvalidHeldLegacyFMInput)
}
func TestHeldLegacyFMSessionImplementsHeldSession(t *testing.T) {
var _ heldSession = (*heldLegacyFMSession)(nil)
}
func TestNewHeldLegacyFMSessionRejectsNilContext(t *testing.T) {
// Given
input := heldLegacyFMInput{
Dashboard: &dashboard.Dashboard{}, PATClient: &client.Client{}, Agent: &agent.Agent{},
Readiness: agent.Readiness{ServerID: 9, UUID: "agent-uuid", Online: true}, Plan: heldFMTestPlan(t, StressSessionFM),
}
// When
_, err := newHeldLegacyFMSession(nil, input)
// Then
require.ErrorIs(t, err, ErrInvalidHeldLegacyFMInput)
}
func TestNewHeldLegacyFMSessionRejectsNilPATBeforeMutation(t *testing.T) {
// Given
agentInstance := newHeldReadinessTestAgent(t, "00000000-0000-0000-0000-000000000303")
input := heldLegacyFMInput{
Dashboard: &dashboard.Dashboard{},
PATClient: nil,
Agent: agentInstance,
Readiness: completeHeldReadiness(agentInstance.UUID()),
Plan: heldFMTestPlan(t, StressSessionFM),
}
workspaceRoot := agentInstance.WorkspaceRoot()
before, err := os.ReadDir(workspaceRoot)
require.NoError(t, err)
// When
var recovered any
func() {
defer func() { recovered = recover() }()
_, err = newHeldLegacyFMSession(context.Background(), input)
}()
// Then
require.Nil(t, recovered)
require.ErrorIs(t, err, ErrInvalidHeldPATClient)
after, readErr := os.ReadDir(workspaceRoot)
require.NoError(t, readErr)
require.Equal(t, before, after)
_, statErr := os.Stat(filepath.Join(workspaceRoot, "held-fm-held-fm-session"))
require.ErrorIs(t, statErr, os.ErrNotExist)
}
func TestHeldLegacyFMSessionRejectsNonFMPlanIdentity(t *testing.T) {
agentInstance := newHeldReadinessTestAgent(t, "00000000-0000-0000-0000-000000000309")
plan := heldFMTestPlan(t, StressSessionTerminal)
input := heldLegacyFMInput{
Dashboard: &dashboard.Dashboard{}, PATClient: &client.Client{}, Agent: agentInstance,
Readiness: completeHeldReadiness(agentInstance.UUID()), Plan: plan,
}
// When
_, err := newHeldLegacyFMSession(context.Background(), input)
// Then
require.Error(t, err)
require.ErrorIs(t, err, ErrInvalidHeldLegacyFMInput)
}
func TestHeldLegacyFMInputRejectsPATClientBeforeOtherValidation(t *testing.T) {
// Given
agentInstance := newHeldReadinessTestAgent(t, "00000000-0000-0000-0000-000000000304")
input := heldLegacyFMInput{
Dashboard: &dashboard.Dashboard{},
Agent: agentInstance,
Readiness: completeHeldReadiness(agentInstance.UUID()),
Plan: heldFMTestPlan(t, StressSessionFM),
}
// When
err := validateHeldLegacyFMInput(context.Background(), input)
// Then
require.ErrorIs(t, err, ErrInvalidHeldPATClient)
}
func TestHeldLegacyFMStreamMismatchErrorDoesNotExposeIdentifiers(t *testing.T) {
// Given
responseID := "response-secret-session"
capabilityID := "capability-secret-stream"
// When
err := heldLegacyFMStreamMismatchError()
message := err.Error()
// Then
require.ErrorIs(t, err, ErrHeldLegacyFMProtocol)
require.NotContains(t, message, responseID)
require.NotContains(t, message, capabilityID)
require.NotContains(t, message, "Authorization")
}
func TestHeldLegacyFMListProofRequiresBinaryExactNZFNPathAndEntry(t *testing.T) {
// Given
root := "/workspace/held-fm/list"
valid := heldLegacyFMListFrame(root, "entry.txt", false)
cases := []struct {
name string
frame client.Frame
}{
{name: "valid", frame: client.Frame{Type: client.FrameBinary, Payload: valid}},
{name: "text", frame: client.Frame{Type: client.FrameText, Payload: valid}},
{name: "wrong path", frame: client.Frame{Type: client.FrameBinary, Payload: heldLegacyFMListFrame("/other", "entry.txt", false)}},
{name: "directory entry", frame: client.Frame{Type: client.FrameBinary, Payload: heldLegacyFMListFrame(root, "entry.txt", true)}},
{name: "remote error", frame: client.Frame{Type: client.FrameBinary, Payload: []byte("NERRdenied")}},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
release := make(chan struct{})
server := heldPumpServer(t, func(connection *websocket.Conn) {
messageType := websocket.BinaryMessage
if testCase.frame.Type == client.FrameText {
messageType = websocket.TextMessage
}
require.NoError(t, connection.WriteMessage(messageType, testCase.frame.Payload))
<-release
})
connection := heldPumpConnection(t, server)
pump, err := newHeldWebSocketPump(context.Background(), connection, 1)
require.NoError(t, err)
err = proveHeldLegacyFMList(context.Background(), pump, root)
if testCase.name == "valid" {
require.NoError(t, err)
} else {
require.Error(t, err)
require.NotContains(t, err.Error(), root)
require.NotContains(t, err.Error(), "entry.txt")
}
require.NoError(t, pump.Stop(context.Background()))
close(release)
})
}
}
func TestHeldLegacyFMCanceledCloseWaiterRetainsCleanupResult(t *testing.T) {
// Given
lifecycle, err := newHeldSessionLifecycle(context.Background(), heldFMTestPlan(t, StressSessionFM), "held-fm-stream", time.Second)
require.NoError(t, err)
require.NoError(t, lifecycle.markLive(nil))
started := make(chan struct{})
release := make(chan struct{})
stack := newHeldCleanupStack()
require.NoError(t, stack.Push(heldCleanupAction{name: "blocked cleanup", cleanup: func(context.Context) error {
close(started)
<-release
return nil
}}))
session := &heldLegacyFMSession{lifecycle: lifecycle, stack: stack}
canceled, cancel := context.WithCancel(context.Background())
cancel()
// When
first := make(chan error, 1)
go func() { first <- session.Close(canceled) }()
<-started
// Then
require.ErrorIs(t, <-first, context.Canceled)
close(release)
require.NoError(t, session.Close(context.Background()))
}
func heldLegacyFMListFrame(path, name string, directory bool) []byte {
kind := byte(0)
if directory {
kind = 1
}
frame := make([]byte, 8, 8+len(path)+2+len(name))
copy(frame, []byte("NZFN"))
binary.BigEndian.PutUint32(frame[4:], uint32(len(path)))
frame = append(frame, []byte(path)...)
frame = append(frame, kind, byte(len(name)))
return append(frame, []byte(name)...)
}
func heldFMTestPlan(t *testing.T, kind StressSessionKind) StressSessionPlan {
t.Helper()
id, err := NewStressSessionID("held-fm-session")
require.NoError(t, err)
ordinal, err := NewStressAgentOrdinal(1)
require.NoError(t, err)
return StressSessionPlan{ID: id, Kind: kind, Ordinal: 1, Agent: ordinal}
}
@@ -0,0 +1,205 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"net"
"strings"
"sync"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
var ErrInvalidHeldNATSession = errors.New("held NAT session input is invalid")
type heldNATInput struct {
Dashboard *dashboard.Dashboard
PATClient *client.Client
Agent *agent.Agent
Readiness agent.Readiness
Plan StressSessionPlan
LifetimeContext context.Context
}
type heldNATSession struct {
lifecycle *heldSessionLifecycle
cleanup *heldCleanupStack
backend *fixture.NATHoldBackend
request *heldNATRequest
observed fixture.NATEchoRecord
profileID uint64
protocol bool
closeOnce sync.Once
}
func newHeldNATSession(ctx context.Context, input heldNATInput) (*heldNATSession, error) {
return newHeldNATSessionWithDependencies(ctx, input, activeHeldNATDependencies())
}
func newHeldNATSessionWithDependencies(ctx context.Context, input heldNATInput, dependencies heldNATDependencies) (*heldNATSession, error) {
if err := validateHeldPATClient(input.PATClient); err != nil {
return nil, err
}
if ctx == nil || input.Dashboard == nil || input.Agent == nil || input.Plan.Kind != StressSessionNAT || input.Plan.ID.String() == "" || input.Plan.Ordinal < 1 || input.Plan.Agent.Int() < 1 {
return nil, ErrInvalidHeldNATSession
}
if err := validateHeldReadiness(input.Agent, input.Readiness); err != nil {
return nil, err
}
stateClient := input.PATClient
baseline, err := dependencies.snapshotState(ctx, stateClient)
if err != nil {
return nil, fmt.Errorf("snapshot NAT IOStream baseline: %w", err)
}
lifetimeContext := input.LifetimeContext
if lifetimeContext == nil {
lifetimeContext = ctx
}
lifecycle, err := newHeldSessionLifecycle(lifetimeContext, input.Plan, "", 30*time.Second)
if err != nil {
return nil, fmt.Errorf("create held NAT lifecycle: %w", err)
}
backend, err := fixture.StartNATHoldBackend()
if err != nil {
return nil, fmt.Errorf("start held NAT backend: %w", err)
}
session := &heldNATSession{lifecycle: lifecycle, cleanup: newHeldCleanupStack(), backend: backend}
if err := session.cleanup.Push(heldCleanupAction{name: "close NAT backend", cleanup: func(context.Context) error { return dependencies.closeBackend(session.backend) }}); err != nil {
return nil, rollbackHeldNAT(session, err)
}
domain, err := heldNATDomain(input.Plan.ID.String())
if err != nil {
return nil, rollbackHeldNAT(session, err)
}
name := "agentcompat-held-nat-" + domain[:strings.Index(domain, ".")]
profileID, err := dependencies.createProfile(ctx, input.Dashboard, backend, input.Readiness.ServerID, name, domain)
if err != nil {
return nil, rollbackHeldNAT(session, fmt.Errorf("create held NAT profile: %w", err))
}
session.profileID = profileID
if err := session.cleanup.Push(heldCleanupAction{name: "NAT profile", cleanup: func(cleanupCtx context.Context) error {
return dependencies.deleteProfile(cleanupCtx, input.Dashboard, session.profileID)
}}); err != nil {
return nil, rollbackHeldNAT(session, err)
}
capability, err := dependencies.register(ctx, input.PATClient, heldIOStreamCapabilityIdentity{Purpose: client.IOStreamCapabilityPurposeNAT, ServerID: input.Readiness.ServerID, ResourceID: session.profileID})
if err != nil {
return nil, rollbackHeldNAT(session, err)
}
if err := session.cleanup.Push(heldCleanupAction{name: "unregister NAT capability", cleanup: func(cleanupCtx context.Context) error {
return dependencies.unregisterCapability(cleanupCtx, capability)
}}); err != nil {
return nil, rollbackHeldNAT(session, err)
}
if err := session.cleanup.Push(heldCleanupAction{name: "wait for NAT stream absence", cleanup: func(cleanupCtx context.Context) error {
return dependencies.waitExpectation(cleanupCtx, capability, stateClient, baseline, true)
}}); err != nil {
return nil, rollbackHeldNAT(session, err)
}
if err := session.cleanup.Push(heldCleanupAction{name: "cancel NAT capability", cleanup: func(cleanupCtx context.Context) error { return dependencies.cancelCapability(cleanupCtx, capability) }}); err != nil {
return nil, rollbackHeldNAT(session, err)
}
request, err := dependencies.startRequest(ctx, input.Dashboard.Endpoint(), domain, input.Plan.ID.String(), capability.HeaderCapability())
if err != nil {
return nil, rollbackHeldNAT(session, err)
}
session.request = request
if err := session.cleanup.Push(heldCleanupAction{name: "held NAT request", cleanup: func(cleanupCtx context.Context) error {
closeErr := dependencies.closeRequest(request)
requestErr := <-request.result
if errors.Is(requestErr, net.ErrClosed) {
requestErr = nil
}
return errors.Join(closeErr, requestErr, cleanupCtx.Err())
}}); err != nil {
return nil, rollbackHeldNAT(session, err)
}
if err := dependencies.waitRequestObserved(ctx, backend); err != nil {
return nil, rollbackHeldNAT(session, fmt.Errorf("prove held NAT request: %w", err))
}
observed, err := dependencies.waitRequest(ctx, backend)
if err != nil {
return nil, rollbackHeldNAT(session, fmt.Errorf("read held NAT request: %w", err))
}
if err := dependencies.proveRequest(observed, domain, input.Plan.ID.String()); err != nil {
return nil, rollbackHeldNAT(session, err)
}
session.observed = observed
session.protocol = true
streamID, err := dependencies.waitCapability(ctx, capability)
if err != nil {
return nil, rollbackHeldNAT(session, err)
}
if err := dependencies.setStreamID(lifecycle, streamID); err != nil {
return nil, rollbackHeldNAT(session, err)
}
if err := dependencies.waitExpectation(ctx, capability, stateClient, baseline, false); err != nil {
return nil, rollbackHeldNAT(session, fmt.Errorf("prove held NAT IOStream: %w", err))
}
if err := lifecycle.markLive(nil); err != nil {
return nil, rollbackHeldNAT(session, err)
}
return session, nil
}
func rollbackHeldNAT(session *heldNATSession, original error) error {
return errors.Join(original, session.rollback())
}
func (session *heldNATSession) rollback() error {
ctx, cancel := context.WithTimeout(context.WithoutCancel(session.lifecycle.baseContext), 30*time.Second)
defer cancel()
return session.cleanup.Run(ctx)
}
func (session *heldNATSession) Plan() StressSessionPlan { return session.lifecycle.Plan() }
func (session *heldNATSession) WaitLive(ctx context.Context) error {
return session.lifecycle.WaitLive(ctx)
}
func (session *heldNATSession) WaitClosed(ctx context.Context) error {
return session.lifecycle.WaitClosed(ctx)
}
func (session *heldNATSession) Done() <-chan struct{} { return session.lifecycle.Done() }
func (session *heldNATSession) CloseResult() error { return session.lifecycle.CloseResult() }
func (session *heldNATSession) IOStreamID() (string, bool) { return session.lifecycle.IOStreamID() }
func (session *heldNATSession) ProtocolProved() bool { return session.protocol }
func (session *heldNATSession) Close(ctx context.Context) error {
owner, won := session.lifecycle.beginClose()
if won {
session.closeOnce.Do(func() {
go func() {
cleanupCtx, cancel := owner.cleanupContext()
defer cancel()
owner.markClosed(session.cleanup.Run(cleanupCtx))
}()
})
}
if err := ctx.Err(); err != nil {
return err
}
return session.lifecycle.WaitClosed(ctx)
}
func heldNATDomain(identity string) (string, error) {
var builder strings.Builder
for _, character := range strings.ToLower(identity) {
if character >= 'a' && character <= 'z' || character >= '0' && character <= '9' || character == '-' {
builder.WriteRune(character)
}
}
if builder.Len() == 0 {
return "", ErrInvalidHeldNATSession
}
return builder.String() + ".agentcompat-nat.invalid", nil
}
var _ heldSession = (*heldNATSession)(nil)
@@ -0,0 +1,185 @@
//go:build linux
package scenario
import (
"context"
"errors"
"net"
"reflect"
"testing"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
type heldNATConstructorFault struct {
name string
failStage string
wantOrder []string
wantStreamID string
wantAbsentID string
wantPresent bool
wantOriginal error
wantRequest bool
wantProofData fixture.NATEchoRecord
wantBaseline int
}
func TestHeldNATConstructorFaultsRollbackRegisteredActions(t *testing.T) {
tests := []heldNATConstructorFault{
{name: "request start", failStage: "start", wantOrder: []string{"cancel", "absence", "unregister", "profile", "backend"}, wantAbsentID: "", wantBaseline: 9, wantOriginal: errHeldNATConstructorRequestStart},
{name: "request observation", failStage: "observe", wantOrder: []string{"request", "cancel", "absence", "unregister", "profile", "backend"}, wantAbsentID: "", wantBaseline: 9, wantOriginal: errHeldNATConstructorObservation, wantRequest: true},
{name: "sensitive header proof", failStage: "proof", wantOrder: []string{"request", "cancel", "absence", "unregister", "profile", "backend"}, wantAbsentID: "", wantBaseline: 9, wantOriginal: errHeldNATConstructorProof, wantRequest: true, wantProofData: fixture.NATEchoRecord{SensitiveHeadersPresent: true}},
{name: "capability wait", failStage: "wait", wantOrder: []string{"request", "cancel", "absence", "unregister", "profile", "backend"}, wantAbsentID: "", wantBaseline: 9, wantOriginal: errHeldNATConstructorWait, wantRequest: true},
{name: "exact stream ID assignment", failStage: "set-id", wantOrder: []string{"request", "cancel", "absence", "unregister", "profile", "backend"}, wantStreamID: "stream-401", wantAbsentID: "stream-401", wantBaseline: 9, wantOriginal: errHeldNATConstructorSetID, wantRequest: true},
{name: "present expectation", failStage: "present", wantOrder: []string{"request", "cancel", "absence", "unregister", "profile", "backend"}, wantStreamID: "stream-402", wantAbsentID: "stream-402", wantBaseline: 9, wantOriginal: errHeldNATConstructorPresent, wantRequest: true},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
var order []string
var observedAbsentID string
var observedPresent bool
var observedBaseline int
dependencies := constructorFaultDependencies(&order, &observedAbsentID, &observedPresent, &observedBaseline, test)
input := heldNATConstructorInput(t)
_, err := newHeldNATSessionWithDependencies(context.Background(), input, dependencies)
if !errors.Is(err, test.wantOriginal) || !errors.Is(err, errHeldNATConstructorCancel) || !errors.Is(err, errHeldNATConstructorAbsence) {
t.Fatalf("error=%v, want original plus cancel and absence failures", err)
}
if !errors.Is(err, errHeldNATConstructorUnregister) || !errors.Is(err, errHeldNATConstructorProfile) || !errors.Is(err, errHeldNATConstructorBackend) {
t.Fatalf("error=%v, want all later cleanup failures", err)
}
if test.wantRequest && !errors.Is(err, errHeldNATConstructorRequestClose) {
t.Fatalf("error=%v, want request close failure", err)
}
if !reflect.DeepEqual(order, test.wantOrder) {
t.Fatalf("cleanup order=%v, want %v", order, test.wantOrder)
}
if observedAbsentID != test.wantAbsentID || observedPresent != test.wantPresent {
t.Fatalf("absence expectation stream=%q present=%t, want stream=%q present=%t", observedAbsentID, observedPresent, test.wantAbsentID, test.wantPresent)
}
if observedBaseline != test.wantBaseline {
t.Fatalf("absence baseline=%d, want %d", observedBaseline, test.wantBaseline)
}
})
}
}
var (
errHeldNATConstructorRequestStart = errors.New("constructor request start failed")
errHeldNATConstructorObservation = errors.New("constructor request observation failed")
errHeldNATConstructorProof = errors.New("constructor request proof failed")
errHeldNATConstructorWait = errors.New("constructor capability wait failed")
errHeldNATConstructorSetID = errors.New("constructor stream ID assignment failed")
errHeldNATConstructorPresent = errors.New("constructor present expectation failed")
errHeldNATConstructorCancel = errors.New("constructor cancel cleanup failed")
errHeldNATConstructorAbsence = errors.New("constructor absence cleanup failed")
errHeldNATConstructorUnregister = errors.New("constructor unregister cleanup failed")
errHeldNATConstructorProfile = errors.New("constructor profile cleanup failed")
errHeldNATConstructorBackend = errors.New("constructor backend cleanup failed")
)
func constructorFaultDependencies(order *[]string, observedAbsentID *string, observedPresent *bool, observedBaseline *int, test heldNATConstructorFault) heldNATDependencies {
return heldNATDependencies{
snapshotState: func(context.Context, *client.Client) (client.IOStreamState, error) {
return client.IOStreamState{Count: 9}, nil
},
createProfile: func(context.Context, *dashboard.Dashboard, *fixture.NATHoldBackend, uint64, string, string) (uint64, error) {
return 401, nil
},
deleteProfile: func(context.Context, *dashboard.Dashboard, uint64) error {
*order = append(*order, "profile")
return errHeldNATConstructorProfile
},
register: func(context.Context, *client.Client, heldIOStreamCapabilityIdentity) (*heldIOStreamCapability, error) {
return &heldIOStreamCapability{}, nil
},
startRequest: func(context.Context, string, string, string, client.IOStreamCapability) (*heldNATRequest, error) {
if test.failStage == "start" {
return nil, errHeldNATConstructorRequestStart
}
return &heldNATRequest{connection: closeErrorConn{err: errHeldNATConstructorRequestClose}, result: closedResult()}, nil
},
waitRequestObserved: func(context.Context, *fixture.NATHoldBackend) error { return nil },
closeBackend: func(backend *fixture.NATHoldBackend) error {
*order = append(*order, "backend")
return errors.Join(backend.Close(), errHeldNATConstructorBackend)
},
closeRequest: func(request *heldNATRequest) error {
*order = append(*order, "request")
return errors.Join(request.close(), errHeldNATConstructorRequestClose)
},
cancelCapability: func(context.Context, *heldIOStreamCapability) error {
*order = append(*order, "cancel")
return errHeldNATConstructorCancel
},
unregisterCapability: func(context.Context, *heldIOStreamCapability) error {
*order = append(*order, "unregister")
return errHeldNATConstructorUnregister
},
waitExpectation: func(_ context.Context, capability *heldIOStreamCapability, _ *client.Client, baseline client.IOStreamState, absent bool) error {
*observedBaseline = baseline.Count
if absent {
*order = append(*order, "absence")
} else if test.failStage == "present" {
return errHeldNATConstructorPresent
}
if absent {
*observedAbsentID = capability.streamID
return errHeldNATConstructorAbsence
}
*observedPresent = true
return nil
},
waitRequest: func(context.Context, *fixture.NATHoldBackend) (fixture.NATEchoRecord, error) {
if test.failStage == "observe" {
return fixture.NATEchoRecord{}, errHeldNATConstructorObservation
}
return test.wantProofData, nil
},
proveRequest: func(fixture.NATEchoRecord, string, string) error {
if test.failStage == "proof" {
return errHeldNATConstructorProof
}
return nil
},
waitCapability: func(_ context.Context, capability *heldIOStreamCapability) (string, error) {
if test.failStage == "wait" {
return "", errHeldNATConstructorWait
}
capability.streamID = test.wantStreamID
return test.wantStreamID, nil
},
setStreamID: func(_ *heldSessionLifecycle, streamID string) error {
if test.failStage == "set-id" {
return errHeldNATConstructorSetID
}
return nil
},
}
}
var errHeldNATConstructorRequestClose = errors.New("constructor request close failed")
func closedResult() chan error {
result := make(chan error, 1)
result <- net.ErrClosed
return result
}
func heldNATConstructorInput(t *testing.T) heldNATInput {
t.Helper()
agentInstance := newHeldReadinessTestAgent(t, "00000000-0000-0000-0000-000000000402")
return heldNATInput{
Dashboard: &dashboard.Dashboard{},
PATClient: &client.Client{},
Agent: agentInstance,
Readiness: completeHeldReadiness(agentInstance.UUID()),
Plan: heldNATTestPlan(t),
}
}
@@ -0,0 +1,79 @@
//go:build linux
package scenario
import (
"context"
"net/http"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
type heldNATDependencies struct {
snapshotState func(context.Context, *client.Client) (client.IOStreamState, error)
createProfile func(context.Context, *dashboard.Dashboard, *fixture.NATHoldBackend, uint64, string, string) (uint64, error)
deleteProfile func(context.Context, *dashboard.Dashboard, uint64) error
register func(context.Context, *client.Client, heldIOStreamCapabilityIdentity) (*heldIOStreamCapability, error)
startRequest func(context.Context, string, string, string, client.IOStreamCapability) (*heldNATRequest, error)
waitRequestObserved func(context.Context, *fixture.NATHoldBackend) error
waitRequest func(context.Context, *fixture.NATHoldBackend) (fixture.NATEchoRecord, error)
proveRequest func(fixture.NATEchoRecord, string, string) error
waitCapability func(context.Context, *heldIOStreamCapability) (string, error)
setStreamID func(*heldSessionLifecycle, string) error
waitExpectation func(context.Context, *heldIOStreamCapability, *client.Client, client.IOStreamState, bool) error
closeBackend func(*fixture.NATHoldBackend) error
closeRequest func(*heldNATRequest) error
cancelCapability func(context.Context, *heldIOStreamCapability) error
unregisterCapability func(context.Context, *heldIOStreamCapability) error
}
func defaultHeldNATDependencies() heldNATDependencies {
return heldNATDependencies{
snapshotState: func(ctx context.Context, transport *client.Client) (client.IOStreamState, error) {
return transport.IOStreamState(ctx)
},
createProfile: func(ctx context.Context, dashboardInstance *dashboard.Dashboard, backend *fixture.NATHoldBackend, serverID uint64, name, domain string) (uint64, error) {
created, err := client.DoREST[natForm, natIDResponse](ctx, dashboardInstance.Clients().REST, client.RESTRequest[natForm]{Method: http.MethodPost, Path: "/api/v1/nat", Body: &natForm{Name: name, Enabled: true, ServerID: serverID, Host: backend.Address(), Domain: domain}})
return uint64(created), err
},
deleteProfile: func(ctx context.Context, dashboardInstance *dashboard.Dashboard, profileID uint64) error {
_, err := client.DoREST[[]uint64, struct{}](ctx, dashboardInstance.Clients().REST, client.RESTRequest[[]uint64]{Method: http.MethodPost, Path: "/api/v1/batch-delete/nat", Body: &[]uint64{profileID}})
return err
},
register: func(ctx context.Context, transport *client.Client, identity heldIOStreamCapabilityIdentity) (*heldIOStreamCapability, error) {
return registerHeldIOStreamCapability(ctx, transport, identity)
},
startRequest: startHeldNATRequest,
waitRequestObserved: func(ctx context.Context, backend *fixture.NATHoldBackend) error {
select {
case <-backend.RequestObserved():
return nil
case <-ctx.Done():
return ctx.Err()
}
},
waitRequest: func(ctx context.Context, backend *fixture.NATHoldBackend) (fixture.NATEchoRecord, error) {
return backend.WaitRequest(ctx)
},
proveRequest: proveHeldNATRequest,
waitCapability: func(ctx context.Context, capability *heldIOStreamCapability) (string, error) {
return capability.Wait(ctx)
},
setStreamID: func(lifecycle *heldSessionLifecycle, streamID string) error {
return lifecycle.setIOStreamID(streamID)
},
waitExpectation: func(ctx context.Context, capability *heldIOStreamCapability, stateClient *client.Client, baseline client.IOStreamState, absent bool) error {
return capability.waitExpectation(ctx, stateClient, baseline, absent)
},
closeBackend: func(backend *fixture.NATHoldBackend) error { return backend.Close() },
closeRequest: func(request *heldNATRequest) error { return request.close() },
cancelCapability: func(ctx context.Context, capability *heldIOStreamCapability) error { return capability.Cancel(ctx) },
unregisterCapability: func(ctx context.Context, capability *heldIOStreamCapability) error { return capability.Unregister(ctx) },
}
}
func activeHeldNATDependencies() heldNATDependencies {
return defaultHeldNATDependencies()
}
@@ -0,0 +1,17 @@
//go:build linux
package scenario
import (
"errors"
"net/http"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
func proveHeldNATRequest(observed fixture.NATEchoRecord, domain, identity string) error {
if observed.Method != http.MethodPatch || observed.Path != "/held/"+identity || observed.Host != domain || observed.HeaderValue != identity || string(observed.Body) != "held-body-"+identity || observed.SensitiveHeadersPresent {
return errors.New("held NAT request did not match exact protocol proof")
}
return nil
}
@@ -0,0 +1,56 @@
//go:build linux
package scenario
import (
"bufio"
"context"
"fmt"
"io"
"net"
"net/http"
"sync"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
)
type heldNATRequest struct {
connection net.Conn
result chan error
closeOnce sync.Once
closeErr error
}
func startHeldNATRequest(ctx context.Context, endpoint, domain, identity string, capability client.IOStreamCapability) (*heldNATRequest, error) {
connection, err := (&net.Dialer{}).DialContext(ctx, "tcp", endpoint)
if err != nil {
return nil, err
}
request := &heldNATRequest{connection: connection, result: make(chan error, 1)}
method := "PATCH"
path := "/held/" + identity
body := "held-body-" + identity
go func() {
wire := fmt.Sprintf("%s %s HTTP/1.1\r\nHost: %s\r\nX-AgentCompat-Echo: %s\r\n%s: %s\r\nContent-Length: %d\r\nConnection: close\r\n\r\n%s", method, path, domain, identity, agentcompatcontract.IOStreamCapabilityHeader, capability.Value(), len(body), body)
_, requestErr := io.WriteString(connection, wire)
if requestErr == nil {
response, readErr := http.ReadResponse(bufio.NewReader(connection), nil)
if readErr == nil {
_, readErr = io.Copy(io.Discard, response.Body)
closeErr := response.Body.Close()
if readErr == nil {
readErr = closeErr
}
}
requestErr = readErr
}
request.result <- requestErr
}()
return request, nil
}
func (request *heldNATRequest) close() error {
request.closeOnce.Do(func() { request.closeErr = request.connection.Close() })
return request.closeErr
}
@@ -0,0 +1,242 @@
//go:build linux
package scenario
import (
"context"
"errors"
"net"
"net/http"
"reflect"
"sync"
"testing"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
func TestHeldNATSessionRejectsInvalidInput(t *testing.T) {
plan := heldNATTestPlan(t)
patClient, err := client.New(client.Config{BaseURL: "http://127.0.0.1"})
if err != nil {
t.Fatal(err)
}
_, err = newHeldNATSession(context.Background(), heldNATInput{PATClient: patClient, Plan: plan})
if !errors.Is(err, ErrInvalidHeldNATSession) {
t.Fatalf("error=%v", err)
}
}
func TestHeldNATSessionRejectsNilPATBeforeRemoteMutation(t *testing.T) {
plan := heldNATTestPlan(t)
_, err := newHeldNATSession(context.Background(), heldNATInput{Plan: plan})
if !errors.Is(err, ErrInvalidHeldPATClient) {
t.Fatalf("error=%v, want ErrInvalidHeldPATClient before other validation", err)
}
}
func TestHeldNATSessionCloseBeforeLiveRetainsLifecycleError(t *testing.T) {
session := newTestHeldNATSession(t)
if err := session.Close(context.Background()); err != nil {
t.Fatal(err)
}
if err := session.WaitLive(context.Background()); !errors.Is(err, ErrHeldSessionClosedBeforeLive) {
t.Fatalf("WaitLive=%v", err)
}
if err := session.WaitClosed(context.Background()); err != nil {
t.Fatal(err)
}
}
func TestHeldNATSessionCanceledWaiterDoesNotCancelOwner(t *testing.T) {
session := newTestHeldNATSession(t)
if err := session.lifecycle.markLive(nil); err != nil {
t.Fatal(err)
}
canceled, cancel := context.WithCancel(context.Background())
cancel()
if err := session.Close(canceled); !errors.Is(err, context.Canceled) {
t.Fatalf("Close=%v", err)
}
if err := session.WaitClosed(context.Background()); err != nil {
t.Fatal(err)
}
}
func TestHeldNATProofRejectsSensitiveHeaders(t *testing.T) {
observed := fixture.NATEchoRecord{Method: http.MethodPatch, Path: "/held/held-nat", Host: "held-nat.agentcompat-nat.invalid", HeaderValue: "held-nat", Body: []byte("held-body-held-nat"), SensitiveHeadersPresent: true}
err := proveHeldNATRequest(observed, observed.Host, "held-nat")
if err == nil {
t.Fatal("proof accepted sensitive backend headers")
}
}
func TestHeldNATInputRejectsNilPATBeforeReadinessOrPlan(t *testing.T) {
_, err := newHeldNATSession(context.Background(), heldNATInput{PATClient: nil})
if !errors.Is(err, ErrInvalidHeldPATClient) {
t.Fatalf("error=%v, want PAT validation before readiness and plan validation", err)
}
}
func TestHeldNATCleanupOrderIsLIFOForRequiredResources(t *testing.T) {
stack := newHeldCleanupStack()
var order []string
for _, name := range []string{"baseline", "backend", "profile", "unregister", "absence", "cancel", "request"} {
name := name
if err := stack.Push(heldCleanupAction{name: name, cleanup: func(context.Context) error {
order = append(order, name)
return nil
}}); err != nil {
t.Fatal(err)
}
}
if err := stack.Run(context.Background()); err != nil {
t.Fatal(err)
}
want := []string{"request", "cancel", "absence", "unregister", "profile", "backend", "baseline"}
if !reflect.DeepEqual(order, want) {
t.Fatalf("cleanup order=%v, want %v", order, want)
}
}
func TestHeldNATRequestCloseRetainsConnectionError(t *testing.T) {
closeFailure := errors.New("request connection close failed")
request := &heldNATRequest{connection: closeErrorConn{err: closeFailure}, result: make(chan error, 1)}
if err := request.close(); !errors.Is(err, closeFailure) {
t.Fatalf("first close error=%v, want %v", err, closeFailure)
}
if err := request.close(); !errors.Is(err, closeFailure) {
t.Fatalf("repeated close error=%v, want %v", err, closeFailure)
}
}
func TestHeldNATSessionConcurrentCloseRetainsOneCleanupResult(t *testing.T) {
session := newTestHeldNATSession(t)
cleanupFailure := errors.New("NAT cleanup failed")
if err := session.cleanup.Push(heldCleanupAction{name: "request", cleanup: func(context.Context) error { return cleanupFailure }}); err != nil {
t.Fatal(err)
}
if err := session.lifecycle.markLive(nil); err != nil {
t.Fatal(err)
}
const callers = 8
errorsSeen := make(chan error, callers)
var group sync.WaitGroup
group.Add(callers)
for range callers {
go func() {
defer group.Done()
errorsSeen <- session.Close(context.Background())
}()
}
group.Wait()
for range callers {
if err := <-errorsSeen; !errors.Is(err, cleanupFailure) {
t.Fatalf("Close error=%v, want %v", err, cleanupFailure)
}
}
}
func TestHeldNATRollbackJoinsOriginalAndCleanupFailures(t *testing.T) {
session := newTestHeldNATSession(t)
original := errors.New("constructor failure")
rollbackFailure := errors.New("rollback failure")
if err := session.cleanup.Push(heldCleanupAction{name: "backend", cleanup: func(context.Context) error { return rollbackFailure }}); err != nil {
t.Fatal(err)
}
err := rollbackHeldNAT(session, original)
if !errors.Is(err, original) || !errors.Is(err, rollbackFailure) {
t.Fatalf("joined error=%v, want original and rollback failures", err)
}
}
func TestHeldNATConstructorRegistrationFailureRollsBackProfileBeforeBackend(t *testing.T) {
const profileID = uint64(77)
registrationFailure := errors.New("capability registration failed")
profileDeleteFailure := errors.New("profile deletion failed")
var cleanupOrder []string
agentInstance := newHeldReadinessTestAgent(t, "00000000-0000-0000-0000-000000000401")
plan := heldNATTestPlan(t)
plan.ID, _ = NewStressSessionID("constructor-rollback")
input := heldNATInput{
Dashboard: &dashboard.Dashboard{},
PATClient: &client.Client{},
Agent: agentInstance,
Readiness: completeHeldReadiness(agentInstance.UUID()),
Plan: plan,
}
dependencies := defaultHeldNATDependencies()
dependencies.snapshotState = func(context.Context, *client.Client) (client.IOStreamState, error) {
return client.IOStreamState{}, nil
}
dependencies.createProfile = func(context.Context, *dashboard.Dashboard, *fixture.NATHoldBackend, uint64, string, string) (uint64, error) {
return profileID, nil
}
dependencies.deleteProfile = func(context.Context, *dashboard.Dashboard, uint64) error {
cleanupOrder = append(cleanupOrder, "profile")
return profileDeleteFailure
}
dependencies.register = func(context.Context, *client.Client, heldIOStreamCapabilityIdentity) (*heldIOStreamCapability, error) {
return nil, registrationFailure
}
_, err := newHeldNATSessionWithDependencies(context.Background(), input, dependencies)
if !errors.Is(err, registrationFailure) || !errors.Is(err, profileDeleteFailure) {
t.Fatalf("constructor error=%v, want registration and profile cleanup failures", err)
}
if !reflect.DeepEqual(cleanupOrder, []string{"profile"}) {
t.Fatalf("cleanup order=%v, want profile before backend cleanup", cleanupOrder)
}
}
type closeErrorConn struct{ err error }
func (connection closeErrorConn) Read([]byte) (int, error) { return 0, net.ErrClosed }
func (connection closeErrorConn) Write([]byte) (int, error) { return 0, net.ErrClosed }
func (connection closeErrorConn) Close() error { return connection.err }
func (connection closeErrorConn) LocalAddr() net.Addr { return heldNATTestAddr{} }
func (connection closeErrorConn) RemoteAddr() net.Addr { return heldNATTestAddr{} }
func (connection closeErrorConn) SetDeadline(time.Time) error { return nil }
func (connection closeErrorConn) SetReadDeadline(time.Time) error { return nil }
func (connection closeErrorConn) SetWriteDeadline(time.Time) error { return nil }
type heldNATTestAddr struct{}
func (heldNATTestAddr) Network() string { return "held-nat-test" }
func (heldNATTestAddr) String() string { return "held-nat-test" }
func heldNATTestPlan(t *testing.T) StressSessionPlan {
t.Helper()
id, err := NewStressSessionID("held-nat")
if err != nil {
t.Fatal(err)
}
agent, err := NewStressAgentOrdinal(1)
if err != nil {
t.Fatal(err)
}
return StressSessionPlan{ID: id, Kind: StressSessionNAT, Ordinal: 1, Agent: agent}
}
func newTestHeldNATSession(t *testing.T) *heldNATSession {
t.Helper()
lifecycle, err := newHeldSessionLifecycle(context.Background(), heldNATTestPlan(t), "", time.Second)
if err != nil {
t.Fatal(err)
}
return &heldNATSession{lifecycle: lifecycle, cleanup: newHeldCleanupStack()}
}
@@ -0,0 +1,80 @@
//go:build linux
package scenario
import (
"errors"
"fmt"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
)
var (
ErrInvalidHeldReadiness = errors.New("held readiness is invalid")
ErrHeldReadinessServerID = errors.New("held readiness server ID is missing")
ErrHeldReadinessUUID = errors.New("held readiness UUID is missing")
ErrHeldReadinessAgentMismatch = errors.New("held readiness does not match Agent")
ErrHeldReadinessVersion = errors.New("held readiness version is missing")
ErrHeldReadinessOnline = errors.New("held readiness is offline")
ErrHeldReadinessVersionObserved = errors.New("held readiness version was not observed")
ErrHeldReadinessRequestTaskEstablished = errors.New("held readiness RequestTask was not established")
ErrHeldReadinessStateReceiptObserved = errors.New("held readiness state receipt was not observed")
)
type HeldReadinessValidationError struct {
Field string
cause error
}
func (validationError *HeldReadinessValidationError) Error() string {
return fmt.Sprintf("held readiness field %q: %s", validationError.Field, validationError.cause)
}
func (validationError *HeldReadinessValidationError) Is(target error) bool {
return target == ErrInvalidHeldReadiness || target == validationError.cause
}
func validateHeldReadiness(agentInstance *agent.Agent, readiness agent.Readiness) error {
if agentInstance == nil {
return ErrInvalidHeldReadiness
}
if readiness.ServerID == 0 {
return newHeldReadinessValidationError("server_id", ErrHeldReadinessServerID)
}
if readiness.UUID == "" {
return newHeldReadinessValidationError("uuid", ErrHeldReadinessUUID)
}
return validateHeldReadinessFacts(heldSessionAgentFacts{PID: agentInstance.PID(), UUID: agentInstance.UUID()}, readiness)
}
func validateHeldReadinessFacts(agentFacts heldSessionAgentFacts, readiness agent.Readiness) error {
if readiness.ServerID == 0 {
return newHeldReadinessValidationError("server_id", ErrHeldReadinessServerID)
}
if readiness.UUID == "" {
return newHeldReadinessValidationError("uuid", ErrHeldReadinessUUID)
}
if readiness.UUID != agentFacts.UUID {
return newHeldReadinessValidationError("uuid", ErrHeldReadinessAgentMismatch)
}
if readiness.Version == "" {
return newHeldReadinessValidationError("version", ErrHeldReadinessVersion)
}
if !readiness.Online {
return newHeldReadinessValidationError("online", ErrHeldReadinessOnline)
}
if !readiness.VersionObserved {
return newHeldReadinessValidationError("version_observed", ErrHeldReadinessVersionObserved)
}
if !readiness.RequestTaskEstablished {
return newHeldReadinessValidationError("request_task_established", ErrHeldReadinessRequestTaskEstablished)
}
if !readiness.StateReceiptObserved {
return newHeldReadinessValidationError("state_receipt_observed", ErrHeldReadinessStateReceiptObserved)
}
return nil
}
func newHeldReadinessValidationError(field string, cause error) error {
return &HeldReadinessValidationError{Field: field, cause: cause}
}
@@ -0,0 +1,134 @@
//go:build linux
package scenario
import (
"context"
"os"
"path/filepath"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
)
func TestHeldReadinessValidation_AcceptsCompleteEvidenceForActualAgent(t *testing.T) {
// Given
agentInstance := newHeldReadinessTestAgent(t, "00000000-0000-0000-0000-000000000301")
readiness := completeHeldReadiness(agentInstance.UUID())
// When
err := validateHeldReadiness(agentInstance, readiness)
// Then
require.NoError(t, err)
}
func TestHeldReadinessValidation_RejectsEachInvalidDimensionWithExactError(t *testing.T) {
// Given
agentInstance := newHeldReadinessTestAgent(t, "00000000-0000-0000-0000-000000000302")
valid := completeHeldReadiness(agentInstance.UUID())
allSentinels := []error{
ErrHeldReadinessServerID,
ErrHeldReadinessUUID,
ErrHeldReadinessAgentMismatch,
ErrHeldReadinessVersion,
ErrHeldReadinessOnline,
ErrHeldReadinessVersionObserved,
ErrHeldReadinessRequestTaskEstablished,
ErrHeldReadinessStateReceiptObserved,
}
tests := []struct {
name string
field string
wantError error
mutate func(*agent.Readiness)
}{
{name: "zero server ID", field: "server_id", wantError: ErrHeldReadinessServerID, mutate: func(readiness *agent.Readiness) { readiness.ServerID = 0 }},
{name: "empty UUID", field: "uuid", wantError: ErrHeldReadinessUUID, mutate: func(readiness *agent.Readiness) { readiness.UUID = "" }},
{name: "agent UUID mismatch", field: "uuid", wantError: ErrHeldReadinessAgentMismatch, mutate: func(readiness *agent.Readiness) { readiness.UUID = "00000000-0000-0000-0000-000000000399" }},
{name: "empty version", field: "version", wantError: ErrHeldReadinessVersion, mutate: func(readiness *agent.Readiness) { readiness.Version = "" }},
{name: "offline", field: "online", wantError: ErrHeldReadinessOnline, mutate: func(readiness *agent.Readiness) { readiness.Online = false }},
{name: "version not observed", field: "version_observed", wantError: ErrHeldReadinessVersionObserved, mutate: func(readiness *agent.Readiness) { readiness.VersionObserved = false }},
{name: "request task not established", field: "request_task_established", wantError: ErrHeldReadinessRequestTaskEstablished, mutate: func(readiness *agent.Readiness) { readiness.RequestTaskEstablished = false }},
{name: "state receipt not observed", field: "state_receipt_observed", wantError: ErrHeldReadinessStateReceiptObserved, mutate: func(readiness *agent.Readiness) { readiness.StateReceiptObserved = false }},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
readiness := valid
test.mutate(&readiness)
// When
err := validateHeldReadiness(agentInstance, readiness)
// Then
require.ErrorIs(t, err, ErrInvalidHeldReadiness)
require.ErrorIs(t, err, test.wantError)
var validationError *HeldReadinessValidationError
require.ErrorAs(t, err, &validationError)
require.Equal(t, test.field, validationError.Field)
require.NotContains(t, err.Error(), valid.UUID)
require.NotContains(t, err.Error(), valid.Version)
for _, sentinel := range allSentinels {
if sentinel != test.wantError {
require.NotErrorIs(t, err, sentinel)
}
}
})
}
}
func completeHeldReadiness(uuid string) agent.Readiness {
return agent.Readiness{
ServerID: 301,
UUID: uuid,
Version: "v2.1.0",
Online: true,
VersionObserved: true,
RequestTaskEstablished: true,
StateReceiptObserved: true,
}
}
func newHeldReadinessTestAgent(t *testing.T, uuid string) *agent.Agent {
t.Helper()
sourceDirectory := t.TempDir()
mainDirectory := filepath.Join(sourceDirectory, "cmd", "agent")
monitorDirectory := filepath.Join(sourceDirectory, "pkg", "monitor")
require.NoError(t, os.MkdirAll(mainDirectory, 0o700))
require.NoError(t, os.MkdirAll(monitorDirectory, 0o700))
require.NoError(t, os.WriteFile(filepath.Join(sourceDirectory, "go.mod"), []byte("module github.com/nezhahq/agent\n\ngo 1.26.3\n"), 0o600))
require.NoError(t, os.WriteFile(filepath.Join(monitorDirectory, "version.go"), []byte("package monitor\n\nvar Version = \"test\"\n"), 0o600))
require.NoError(t, os.WriteFile(filepath.Join(mainDirectory, "main.go"), []byte(`package main
import (
"os"
"os/signal"
"syscall"
"github.com/nezhahq/agent/pkg/monitor"
)
func main() {
_ = monitor.Version
signals := make(chan os.Signal, 1)
signal.Notify(signals, syscall.SIGINT, syscall.SIGTERM)
<-signals
}
`), 0o600))
agentInstance, err := agent.Start(t.Context(), agent.AgentStartConfig{
SourceDir: sourceDirectory,
Endpoint: "127.0.0.1:1",
UUID: uuid,
})
require.NoError(t, err)
t.Cleanup(func() {
cleanupContext, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
require.NoError(t, agentInstance.Stop(cleanupContext))
})
return agentInstance
}
@@ -0,0 +1,71 @@
//go:build linux && agentcompat
package scenario
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"slices"
processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process"
)
type heldRealEvidence struct {
Kind string `json:"kind"`
BaselineCount int `json:"baseline_count"`
LiveCount int `json:"live_count"`
ClosedCount int `json:"closed_count"`
ExactIDPresent bool `json:"exact_id_present"`
ExactIDAbsent bool `json:"exact_id_absent"`
ProtocolProved bool `json:"protocol_proved"`
SensitiveHeadersPresent bool `json:"sensitive_headers_present"`
DashboardPIDUnchanged bool `json:"dashboard_pid_unchanged"`
AgentPIDUnchanged bool `json:"agent_pid_unchanged"`
CleanupOK bool `json:"cleanup_ok"`
}
type heldRealCleanup struct {
Agent processharness.CleanupReceipt
Dashboard processharness.CleanupReceipt
SessionClosed bool
ExactStreamGone bool
OwnedResourceGone bool
AgentPIDGone bool
DashboardPIDGone bool
}
func heldRealArtifactKinds() []string {
return []string{"terminal", "file-manager", "nat"}
}
func heldRealCleanupOK(cleanup heldRealCleanup) bool {
return cleanup.Agent.Passed && cleanup.Dashboard.Passed && cleanup.SessionClosed && cleanup.ExactStreamGone && cleanup.OwnedResourceGone && cleanup.AgentPIDGone && cleanup.DashboardPIDGone
}
func heldRealArtifactKeys() []string {
return []string{
"kind", "baseline_count", "live_count", "closed_count", "exact_id_present", "exact_id_absent",
"protocol_proved", "sensitive_headers_present", "dashboard_pid_unchanged", "agent_pid_unchanged", "cleanup_ok",
}
}
func writeHeldRealEvidence(kind string, evidence heldRealEvidence) error {
if !slices.Contains(heldRealArtifactKinds(), kind) || evidence.Kind != kind {
return fmt.Errorf("held real evidence kind is invalid")
}
data, err := json.Marshal(evidence)
if err != nil {
return fmt.Errorf("encode held real evidence: %w", err)
}
root := "/tmp/nezha-held-real-sessions"
if err := os.MkdirAll(root, 0o700); err != nil {
return fmt.Errorf("create held real evidence directory: %w", err)
}
path := filepath.Join(root, kind+".json")
if err := os.WriteFile(path, data, 0o600); err != nil {
return fmt.Errorf("write held real evidence: %w", err)
}
return nil
}
@@ -0,0 +1,56 @@
//go:build linux && agentcompat
package scenario
import (
"context"
"encoding/json"
"errors"
"strings"
"testing"
"github.com/stretchr/testify/require"
processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process"
)
func TestHeldRealArtifactKindsIncludeEveryRealSessionKind(t *testing.T) {
require.ElementsMatch(t, []string{"terminal", "file-manager", "nat"}, heldRealArtifactKinds())
}
func TestHeldRealCleanupOKRejectsReceiptErrorAndRunningPID(t *testing.T) {
passedReceipt := processharness.NewCleanupReceipt([]processharness.CleanupRecord{{Name: "process", PID: 1}})
failedReceipt := processharness.NewCleanupReceipt([]processharness.CleanupRecord{{Name: "process", PID: 1, Error: "cleanup failed"}})
complete := heldRealCleanup{Agent: passedReceipt, Dashboard: passedReceipt, SessionClosed: true, ExactStreamGone: true, OwnedResourceGone: true, AgentPIDGone: true, DashboardPIDGone: true}
require.True(t, heldRealCleanupOK(complete))
complete.Agent = failedReceipt
require.False(t, heldRealCleanupOK(complete))
complete.Agent = passedReceipt
complete.AgentPIDGone = false
require.False(t, heldRealCleanupOK(complete))
}
func TestHeldRealNATProfileQueryPropagatesRESTError(t *testing.T) {
wantErr := errors.New("query failed")
present, err := heldRealNATProfilePresentWithQuery(context.Background(), 9, func(context.Context) ([]heldRealNATProfile, error) {
return nil, wantErr
})
require.ErrorIs(t, err, wantErr)
require.False(t, present)
}
func TestHeldRealEvidenceUsesOnlyRedactedSchema(t *testing.T) {
data, err := json.Marshal(heldRealEvidence{Kind: "terminal", SensitiveHeadersPresent: false})
require.NoError(t, err)
var fields map[string]json.RawMessage
require.NoError(t, json.Unmarshal(data, &fields))
for field := range fields {
require.Contains(t, heldRealArtifactKeys(), field)
}
require.NotContains(t, strings.ToLower(string(data)), "stream")
require.NotContains(t, strings.ToLower(string(data)), "token")
require.NotContains(t, strings.ToLower(string(data)), "authorization")
require.Error(t, writeHeldRealEvidence("unknown", heldRealEvidence{Kind: "unknown"}))
}
@@ -0,0 +1,176 @@
//go:build linux
package scenario
import (
"context"
"errors"
"sync"
"time"
)
var (
ErrInvalidHeldSessionPlan = errors.New("held session plan is invalid")
ErrHeldSessionClosedBeforeLive = errors.New("held session closed before live")
ErrHeldSessionLiveResolved = errors.New("held session live state already resolved")
)
type heldSession interface {
Plan() StressSessionPlan
WaitLive(context.Context) error
Close(context.Context) error
WaitClosed(context.Context) error
IOStreamID() (string, bool)
Done() <-chan struct{}
CloseResult() error
}
type heldSessionState uint8
const (
heldSessionConstructed heldSessionState = iota
heldSessionLive
heldSessionFailed
heldSessionClosing
heldSessionClosed
)
type heldSessionLifecycle struct {
baseContext context.Context
plan StressSessionPlan
ioStreamID string
cleanupTimeout time.Duration
mu sync.Mutex
state heldSessionState
liveResult error
liveDone chan struct{}
closedDone chan struct{}
closedResult error
}
type heldSessionCloseOwner struct {
lifecycle *heldSessionLifecycle
closeOnce sync.Once
}
func newHeldSessionLifecycle(baseContext context.Context, plan StressSessionPlan, ioStreamID string, cleanupTimeout time.Duration) (*heldSessionLifecycle, error) {
if baseContext == nil || plan.ID.String() == "" || !supportedHeldSessionKind(plan.Kind) || plan.Ordinal < 1 || plan.Agent.Int() < 1 || cleanupTimeout <= 0 {
return nil, ErrInvalidHeldSessionPlan
}
return &heldSessionLifecycle{
baseContext: baseContext,
plan: plan,
ioStreamID: ioStreamID,
cleanupTimeout: cleanupTimeout,
state: heldSessionConstructed,
liveDone: make(chan struct{}),
closedDone: make(chan struct{}),
}, nil
}
func supportedHeldSessionKind(kind StressSessionKind) bool {
switch kind {
case StressSessionTerminal, StressSessionNAT, StressSessionFM:
return true
default:
return false
}
}
func (lifecycle *heldSessionLifecycle) Plan() StressSessionPlan { return lifecycle.plan }
func (lifecycle *heldSessionLifecycle) IOStreamID() (string, bool) {
return lifecycle.ioStreamID, lifecycle.ioStreamID != ""
}
func (lifecycle *heldSessionLifecycle) setIOStreamID(streamID string) error {
if streamID == "" {
return errors.New("held session stream ID is empty")
}
lifecycle.mu.Lock()
defer lifecycle.mu.Unlock()
if lifecycle.state != heldSessionConstructed {
return ErrHeldSessionLiveResolved
}
lifecycle.ioStreamID = streamID
return nil
}
func (lifecycle *heldSessionLifecycle) markLive(err error) error {
lifecycle.mu.Lock()
defer lifecycle.mu.Unlock()
if lifecycle.state != heldSessionConstructed {
return ErrHeldSessionLiveResolved
}
lifecycle.liveResult = err
if err == nil {
lifecycle.state = heldSessionLive
} else {
lifecycle.state = heldSessionFailed
}
close(lifecycle.liveDone)
return nil
}
func (lifecycle *heldSessionLifecycle) WaitLive(ctx context.Context) error {
select {
case <-lifecycle.liveDone:
lifecycle.mu.Lock()
defer lifecycle.mu.Unlock()
return lifecycle.liveResult
case <-ctx.Done():
return ctx.Err()
}
}
func (lifecycle *heldSessionLifecycle) beginClose() (*heldSessionCloseOwner, bool) {
lifecycle.mu.Lock()
defer lifecycle.mu.Unlock()
if lifecycle.state == heldSessionConstructed {
lifecycle.liveResult = ErrHeldSessionClosedBeforeLive
lifecycle.state = heldSessionClosing
close(lifecycle.liveDone)
return &heldSessionCloseOwner{lifecycle: lifecycle}, true
}
if lifecycle.state == heldSessionLive || lifecycle.state == heldSessionFailed {
lifecycle.state = heldSessionClosing
return &heldSessionCloseOwner{lifecycle: lifecycle}, true
}
return nil, false
}
func (owner *heldSessionCloseOwner) cleanupContext() (context.Context, context.CancelFunc) {
// The deadline only signals cancellation; it cannot forcibly terminate arbitrary cleanup work.
return context.WithTimeout(context.WithoutCancel(owner.lifecycle.baseContext), owner.lifecycle.cleanupTimeout)
}
func (owner *heldSessionCloseOwner) markClosed(err error) {
owner.closeOnce.Do(func() {
lifecycle := owner.lifecycle
lifecycle.mu.Lock()
lifecycle.closedResult = err
lifecycle.state = heldSessionClosed
close(lifecycle.closedDone)
lifecycle.mu.Unlock()
})
}
func (lifecycle *heldSessionLifecycle) WaitClosed(ctx context.Context) error {
select {
case <-lifecycle.closedDone:
lifecycle.mu.Lock()
defer lifecycle.mu.Unlock()
return lifecycle.closedResult
case <-ctx.Done():
return ctx.Err()
}
}
func (lifecycle *heldSessionLifecycle) Done() <-chan struct{} { return lifecycle.closedDone }
func (lifecycle *heldSessionLifecycle) CloseResult() error {
lifecycle.mu.Lock()
defer lifecycle.mu.Unlock()
return lifecycle.closedResult
}
@@ -0,0 +1,80 @@
//go:build linux
package scenario
import (
"context"
"errors"
"sync/atomic"
"testing"
)
func TestHeldSessionAdapterCanceledWaiterRetainsCleanupForLaterCaller(t *testing.T) {
cleanupStarted := make(chan struct{})
cleanupRelease := make(chan struct{})
cleanupErr := errors.New("cleanup failed")
var cleanupCount atomic.Int32
lifecycle := heldTestLifecycle(t, "")
adapter := &heldSessionAdapter{
lifecycle: lifecycle,
cleanup: func(cleanupContext context.Context) error {
cleanupCount.Add(1)
if err := cleanupContext.Err(); err != nil {
t.Errorf("cleanup context canceled before release: %v", err)
}
close(cleanupStarted)
<-cleanupRelease
return cleanupErr
},
}
canceled, cancel := context.WithCancel(context.Background())
cancel()
firstResult := make(chan error, 1)
go func() { firstResult <- adapter.Close(canceled) }()
<-cleanupStarted
if err := <-firstResult; !errors.Is(err, context.Canceled) {
t.Fatalf("canceled Close error = %v", err)
}
close(cleanupRelease)
if err := adapter.Close(context.Background()); !errors.Is(err, cleanupErr) {
t.Fatalf("later Close error = %v", err)
}
if count := cleanupCount.Load(); count != 1 {
t.Fatalf("cleanup invocation count = %d, want 1", count)
}
}
type heldSessionAdapter struct {
lifecycle *heldSessionLifecycle
cleanup func(context.Context) error
}
func (adapter *heldSessionAdapter) Plan() StressSessionPlan { return adapter.lifecycle.Plan() }
func (adapter *heldSessionAdapter) WaitLive(ctx context.Context) error {
return adapter.lifecycle.WaitLive(ctx)
}
func (adapter *heldSessionAdapter) IOStreamID() (string, bool) { return adapter.lifecycle.IOStreamID() }
func (adapter *heldSessionAdapter) WaitClosed(ctx context.Context) error {
return adapter.lifecycle.WaitClosed(ctx)
}
func (adapter *heldSessionAdapter) Done() <-chan struct{} { return adapter.lifecycle.Done() }
func (adapter *heldSessionAdapter) CloseResult() error { return adapter.lifecycle.CloseResult() }
func (adapter *heldSessionAdapter) Close(ctx context.Context) error {
owner, won := adapter.lifecycle.beginClose()
if !won {
return adapter.lifecycle.WaitClosed(ctx)
}
go func() {
cleanupContext, cancel := owner.cleanupContext()
defer cancel()
owner.markClosed(adapter.cleanup(cleanupContext))
}()
return adapter.lifecycle.WaitClosed(ctx)
}
var _ heldSession = (*heldSessionAdapter)(nil)
@@ -0,0 +1,251 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"sync"
"golang.org/x/sync/errgroup"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
)
type heldSessionSet struct {
mu sync.Mutex
plans []StressSessionPlan
sessions []heldSession
state heldSessionSetStateObserver
dependencies HeldSessionSetDependencies
baseline client.IOStreamState
healthContext context.Context
healthMu sync.Mutex
healthError error
healthErrors []error
healthDoneOnce sync.Once
healthDone chan struct{}
healthStop chan struct{}
healthStopOnce sync.Once
healthShutdown chan struct{}
healthShutdownOnce sync.Once
healthShutdownRequests chan heldHealthShutdownRequest
healthWG sync.WaitGroup
healthCoordinatorWG sync.WaitGroup
healthCoordinatorDoneOnce sync.Once
healthEvents chan heldHealthMessage
healthSnapshots []chan heldHealthSnapshotRequest
coordinatorDone chan struct{}
healthSnapshotRequestHook func(int)
healthSnapshotReplyHook func(int)
healthSnapshotOverrideHook func(int, heldHealthSnapshotRequest) *heldHealthSnapshot
healthSnapshotSendHook func(int)
healthEventHook func(heldHealthMessage)
healthClosureObservedHook func(int)
healthSnapshotAcceptedHook func(heldHealthSnapshot)
healthShutdownAcceptedHook func()
healthShutdownAcknowledgedHook func()
healthWatcherDone []chan struct{}
closing bool
closeOnce sync.Once
closeDone chan struct{}
closeError error
}
func NewHeldSessionSet(ctx context.Context, input HeldSessionSetInput) (*heldSessionSet, error) {
if ctx == nil {
return nil, ErrInvalidHeldSessionSetTopology
}
plans, err := validateHeldSessionSetPlans(input.Plan)
if err != nil {
return nil, redactHeldSessionSetError(err)
}
topology, err := validateHeldSessionSetTopology(input, plans)
if err != nil {
return nil, redactHeldSessionSetError(err)
}
dependencies := input.Dependencies
if (dependencies.Terminal == nil) != (dependencies.NAT == nil) || (dependencies.Terminal == nil) != (dependencies.FM == nil) || (dependencies.Terminal == nil) != (dependencies.Snapshot == nil) || (dependencies.Terminal == nil) != (dependencies.WaitState == nil) || (dependencies.Terminal == nil) != (dependencies.InspectAgent == nil) || (dependencies.Terminal == nil) != (dependencies.ObserveState == nil) {
return nil, redactHeldSessionSetError(ErrInvalidHeldSessionSetTopology)
}
if dependencies.Terminal == nil {
dependencies = defaultHeldSessionSetDependencies()
}
baseline, err := dependencies.Snapshot(ctx, topology.stateClient)
if err != nil {
return nil, redactHeldSessionSetError(err)
}
set := newHeldSessionSet(plans, topology.stateClient, baseline, dependencies, ctx)
set.healthSnapshotRequestHook = input.testHealthSnapshotRequestHook
set.healthSnapshotReplyHook = input.testHealthSnapshotReplyHook
set.healthSnapshotOverrideHook = input.testHealthSnapshotOverrideHook
set.healthSnapshotSendHook = input.testHealthSnapshotSendHook
set.healthEventHook = input.testHealthEventHook
set.healthClosureObservedHook = input.testHealthClosureObservedHook
set.healthSnapshotAcceptedHook = input.testHealthSnapshotAcceptedHook
set.healthShutdownAcceptedHook = input.testHealthShutdownAcceptedHook
set.healthShutdownAcknowledgedHook = input.testHealthShutdownAcknowledgedHook
if err := set.construct(ctx, topology); err != nil {
return nil, redactHeldSessionSetError(err)
}
if err := set.waitLive(ctx); err != nil {
return nil, redactHeldSessionSetError(errors.Join(err, set.Close(context.WithoutCancel(ctx))))
}
set.startHealthWatchers()
return set, nil
}
func (set *heldSessionSet) construct(ctx context.Context, topology heldSessionSetTopology) error {
acquisitionContext, cancelAcquisition := context.WithCancel(ctx)
defer cancelAcquisition()
ready := make(chan struct{}, len(set.plans))
release := make(chan struct{})
var releaseOnce sync.Once
releaseWorkers := func() { releaseOnce.Do(func() { close(release) }) }
group, groupContext := errgroup.WithContext(acquisitionContext)
var mu sync.Mutex
var constructionErrors error
for index, plan := range set.plans {
index, plan := index, plan
group.Go(func() error {
select {
case ready <- struct{}{}:
case <-groupContext.Done():
return groupContext.Err()
}
select {
case <-release:
case <-groupContext.Done():
return groupContext.Err()
}
topologyAgent := topology.agents[plan.Agent.Int()]
// Acquisition cancellation unblocks siblings; successful sessions retain the outer lifetime context.
session, err := constructHeldSession(groupContext, ctx, topology.dashboard, topologyAgent, plan, set.dependencies)
if err != nil {
mu.Lock()
constructionErrors = errors.Join(constructionErrors, err)
mu.Unlock()
return err
}
mu.Lock()
set.sessions[index] = session
mu.Unlock()
return nil
})
}
for range set.plans {
select {
case <-ready:
case <-groupContext.Done():
releaseWorkers()
groupError := group.Wait()
return errors.Join(groupContext.Err(), groupError, constructionErrors, set.rollback(ctx))
}
}
releaseWorkers()
groupError := group.Wait()
if groupError != nil {
return errors.Join(groupError, constructionErrors, set.rollback(ctx))
}
return nil
}
func constructHeldSession(ctx, lifetimeContext context.Context, dashboardInstance *dashboard.Dashboard, topology HeldSessionAgent, plan StressSessionPlan, dependencies HeldSessionSetDependencies) (heldSession, error) {
switch plan.Kind {
case StressSessionTerminal:
return dependencies.Terminal(ctx, heldTerminalInput{Dashboard: dashboardInstance, PATClient: topology.PATClient, Agent: topology.Agent, Readiness: topology.Readiness, Plan: plan, LifetimeContext: lifetimeContext})
case StressSessionNAT:
return dependencies.NAT(ctx, heldNATInput{Dashboard: dashboardInstance, PATClient: topology.PATClient, Agent: topology.Agent, Readiness: topology.Readiness, Plan: plan, LifetimeContext: lifetimeContext})
case StressSessionFM:
return dependencies.FM(ctx, heldLegacyFMInput{Dashboard: dashboardInstance, PATClient: topology.PATClient, Agent: topology.Agent, Readiness: topology.Readiness, Plan: plan, LifetimeContext: lifetimeContext})
default:
return nil, fmt.Errorf("unsupported held session kind %q: %w", plan.Kind, ErrInvalidHeldSessionSetPlan)
}
}
func (set *heldSessionSet) waitLive(ctx context.Context) error {
group, groupContext := errgroup.WithContext(ctx)
for index, session := range set.sessions {
index, session := index, session
group.Go(func() error {
if session == nil {
return fmt.Errorf("session %d was not constructed", index)
}
if err := session.WaitLive(groupContext); err != nil {
return err
}
select {
case <-session.Done():
return errors.Join(ErrHeldSessionPrematureClose, session.CloseResult())
default:
}
if !sessionMatchesPlan(session, set.plans[index]) {
return fmt.Errorf("session %d does not match its plan", index)
}
return nil
})
}
if err := group.Wait(); err != nil {
return err
}
for _, session := range set.sessions {
if session == nil {
return errors.New("held session set has an unconstructed session")
}
}
streamIDs, err := set.streamIDs()
if err != nil {
return err
}
streamErr := waitHeldSessionSetStreams(ctx, set.state, set.baseline.Count+len(streamIDs), streamIDs, true, set.dependencies.WaitState)
_, aggregateErr := set.dependencies.WaitState(ctx, set.state, client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(set.baseline.Count + len(streamIDs))})
return errors.Join(streamErr, aggregateErr)
}
func (set *heldSessionSet) streamIDs() ([]string, error) {
ids := make([]string, len(set.sessions))
seen := make(map[string]struct{}, len(ids))
for index, session := range set.sessions {
streamID, present := session.IOStreamID()
if !present || streamID == "" {
return nil, errors.New("held session stream ID is empty")
}
if _, exists := seen[streamID]; exists {
return nil, fmt.Errorf("duplicate held session stream ID: %w", ErrInvalidHeldSessionSetPlan)
}
seen[streamID] = struct{}{}
ids[index] = streamID
}
return ids, nil
}
func sessionMatchesPlan(session heldSession, plan StressSessionPlan) bool {
return session.Plan() == plan
}
func waitHeldSessionSetStreams(ctx context.Context, state heldSessionSetStateObserver, expectedCount int, streamIDs []string, present bool, waitState func(context.Context, heldSessionSetStateObserver, client.IOStreamStateExpectation) (client.IOStreamState, error)) error {
group, groupContext := errgroup.WithContext(ctx)
var mu sync.Mutex
var joined error
for _, streamID := range streamIDs {
streamID := streamID
group.Go(func() error {
expectation := client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(expectedCount)}
if present {
expectation.PresentStreamID = streamID
} else {
expectation.AbsentStreamID = streamID
}
_, err := waitState(groupContext, state, expectation)
if err != nil {
mu.Lock()
joined = errors.Join(joined, err)
mu.Unlock()
}
return nil
})
}
return errors.Join(group.Wait(), joined)
}
@@ -0,0 +1,82 @@
//go:build linux
package scenario
import (
"context"
"errors"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
func newHeldSessionSet(plans []StressSessionPlan, state heldSessionSetStateObserver, baseline client.IOStreamState, dependencies HeldSessionSetDependencies, base context.Context) *heldSessionSet {
return &heldSessionSet{plans: plans, sessions: make([]heldSession, len(plans)), state: state, baseline: baseline, dependencies: dependencies, healthContext: base, healthDone: make(chan struct{}), healthStop: make(chan struct{}), healthShutdown: make(chan struct{}), healthShutdownRequests: make(chan heldHealthShutdownRequest), healthErrors: make([]error, len(plans)), coordinatorDone: make(chan struct{}), closeDone: make(chan struct{})}
}
func (set *heldSessionSet) rollback(ctx context.Context) error {
ownerContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second)
defer cancel()
return set.closeAll(ownerContext)
}
func (set *heldSessionSet) Close(ctx context.Context) error {
if err := ctx.Err(); err != nil {
return err
}
set.closeOnce.Do(func() {
set.beginOwnedClose()
if set.healthEvents == nil {
set.markHealthCoordinatorDone()
}
go func() {
ownerContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second)
defer cancel()
ack := make(chan struct{})
shutdownSent := false
select {
case set.healthShutdownRequests <- heldHealthShutdownRequest{ack: ack}:
shutdownSent = true
case <-set.coordinatorDone:
}
if shutdownSent {
select {
case <-ack:
case <-set.coordinatorDone:
}
}
<-set.coordinatorDone
set.healthWG.Wait()
set.closeError = redactHeldSessionSetError(set.closeAll(ownerContext))
close(set.closeDone)
}()
})
select {
case <-set.closeDone:
return set.closeError
case <-ctx.Done():
return ctx.Err()
}
}
func (set *heldSessionSet) closeAll(ctx context.Context) error {
var joined error
streamIDs := make([]string, 0, len(set.sessions))
for index := len(set.sessions) - 1; index >= 0; index-- {
session := set.sessions[index]
if session == nil {
continue
}
joined = errors.Join(joined, session.Close(ctx), session.WaitClosed(ctx))
streamID, present := session.IOStreamID()
if present && streamID != "" {
streamIDs = append(streamIDs, streamID)
}
}
if len(streamIDs) > 0 {
joined = errors.Join(joined, waitHeldSessionSetStreams(ctx, set.state, set.baseline.Count, streamIDs, false, set.dependencies.WaitState))
_, aggregateErr := set.dependencies.WaitState(ctx, set.state, client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(set.baseline.Count)})
joined = errors.Join(joined, aggregateErr)
}
return joined
}
@@ -0,0 +1,54 @@
//go:build linux
package scenario
import "context"
func (set *heldSessionSet) startHealthWatchers() {
set.healthEvents = make(chan heldHealthMessage, len(set.sessions))
set.healthSnapshots = make([]chan heldHealthSnapshotRequest, len(set.sessions))
set.healthWatcherDone = make([]chan struct{}, len(set.sessions))
for index := range set.sessions {
set.healthSnapshots[index] = make(chan heldHealthSnapshotRequest)
set.healthWatcherDone[index] = make(chan struct{})
}
set.healthCoordinatorWG.Add(1)
go set.runHealthCoordinator()
set.healthWG.Add(len(set.sessions))
for index, session := range set.sessions {
go set.watchHealth(index, session)
}
}
func (set *heldSessionSet) markHealthCoordinatorDone() {
set.healthCoordinatorDoneOnce.Do(func() { close(set.coordinatorDone) })
}
func (set *heldSessionSet) WaitHealthy(ctx context.Context) error {
select {
case <-set.healthDone:
return set.retainedHealthError()
case <-set.closeDone:
return set.retainedHealthError()
case <-ctx.Done():
return ctx.Err()
}
}
func (set *heldSessionSet) Done() <-chan struct{} { return set.healthDone }
func (set *heldSessionSet) beginOwnedClose() {
set.healthMu.Lock()
set.closing = true
set.healthMu.Unlock()
}
func (set *heldSessionSet) retainedHealthError() error {
set.healthMu.Lock()
defer set.healthMu.Unlock()
return set.healthError
}
func (set *heldSessionSet) stopHealthWatchers() {
set.healthStopOnce.Do(func() { close(set.healthStop) })
}
@@ -0,0 +1,234 @@
//go:build linux
package scenario
import "errors"
const heldSessionSetHealthMemberCount = 12
var ErrHeldSessionSetHealthProtocol = errors.New("held session set health protocol failed")
type heldHealthEvent struct {
index int
err error
}
type heldHealthSnapshotRequest struct {
epoch int
}
type heldHealthSnapshot struct {
epoch int
index int
event *heldHealthEvent
}
type heldHealthMessage struct {
epoch int
event *heldHealthEvent
snapshot *heldHealthSnapshot
}
type heldHealthShutdownRequest struct {
ack chan struct{}
}
type heldHealthEpochProtocol struct {
epoch int
seen [heldSessionSetHealthMemberCount]bool
replies int
err error
committed bool
}
func newHeldHealthEpochProtocol(epoch int) *heldHealthEpochProtocol {
return &heldHealthEpochProtocol{epoch: epoch}
}
func (protocol *heldHealthEpochProtocol) accept(snapshot heldHealthSnapshot) {
if protocol.err != nil || protocol.committed {
return
}
if snapshot.epoch != protocol.epoch || snapshot.index < 0 || snapshot.index >= heldSessionSetHealthMemberCount || protocol.seen[snapshot.index] {
protocol.err = ErrHeldSessionSetHealthProtocol
return
}
protocol.seen[snapshot.index] = true
protocol.replies++
if protocol.replies == heldSessionSetHealthMemberCount {
protocol.committed = true
}
}
func (set *heldSessionSet) watchHealth(index int, session heldSession) {
defer set.healthWG.Done()
defer close(set.healthWatcherDone[index])
var cached *heldHealthEvent
eventSent := false
closureObserved := false
for {
select {
case <-session.Done():
if !closureObserved && set.healthClosureObservedHook != nil {
set.healthClosureObservedHook(index)
closureObserved = true
}
if cached == nil {
cached = heldSessionHealthEvent(index, session)
}
if !eventSent {
select {
case set.healthEvents <- heldHealthMessage{epoch: 1, event: cached}:
if set.healthEventHook != nil {
set.healthEventHook(heldHealthMessage{epoch: 1, event: cached})
}
eventSent = true
case <-set.healthStop:
return
}
}
case request := <-set.healthSnapshots[index]:
if set.healthSnapshotRequestHook != nil {
set.healthSnapshotRequestHook(index)
}
select {
case <-session.Done():
if !closureObserved && set.healthClosureObservedHook != nil {
set.healthClosureObservedHook(index)
closureObserved = true
}
cached = heldSessionHealthEvent(index, session)
default:
}
message := &heldHealthSnapshot{epoch: request.epoch, index: index, event: cached}
if set.healthSnapshotOverrideHook != nil {
if override := set.healthSnapshotOverrideHook(index, request); override != nil {
message = override
}
}
if set.healthSnapshotSendHook != nil {
set.healthSnapshotSendHook(index)
}
select {
case set.healthEvents <- heldHealthMessage{epoch: request.epoch, snapshot: message}:
if set.healthEventHook != nil {
set.healthEventHook(heldHealthMessage{epoch: request.epoch, snapshot: message})
}
case <-set.healthStop:
return
}
if set.healthSnapshotReplyHook != nil {
set.healthSnapshotReplyHook(index)
}
case <-set.healthStop:
return
}
}
}
func heldSessionHealthEvent(index int, session heldSession) *heldHealthEvent {
errorValue := redactHeldSessionSetHealthError(session.CloseResult())
if errorValue == nil {
errorValue = redactHeldSessionSetHealthError(ErrHeldSessionPrematureClose)
}
return &heldHealthEvent{index: index, err: errorValue}
}
func (set *heldSessionSet) runHealthCoordinator() {
defer set.healthCoordinatorWG.Done()
defer set.markHealthCoordinatorDone()
for {
select {
case message := <-set.healthEvents:
if message.event != nil {
set.commitHealthEpoch(*message.event, 1, nil)
return
}
case request := <-set.healthShutdownRequests:
if set.healthShutdownAcceptedHook != nil {
set.healthShutdownAcceptedHook()
}
set.commitHealthEpoch(heldHealthEvent{index: heldSessionSetHealthMemberCount}, 1, request.ack)
return
case <-set.healthStop:
return
}
}
}
func (set *heldSessionSet) commitHealthEpoch(trigger heldHealthEvent, epoch int, shutdownAck chan struct{}) {
shutdownAcknowledgement := shutdownAck
acknowledgeShutdown := func() {
if shutdownAcknowledgement != nil {
close(shutdownAcknowledgement)
if set.healthShutdownAcknowledgedHook != nil {
set.healthShutdownAcknowledgedHook()
}
shutdownAcknowledgement = nil
}
}
for index := range set.healthSnapshots {
select {
case set.healthSnapshots[index] <- heldHealthSnapshotRequest{epoch: epoch}:
case <-set.healthStop:
acknowledgeShutdown()
return
}
}
best := trigger
protocol := newHeldHealthEpochProtocol(epoch)
acceptedSnapshots := [heldSessionSetHealthMemberCount]bool{}
for protocol.replies < len(set.healthSnapshots) {
select {
case message := <-set.healthEvents:
if message.event != nil {
if shutdownAcknowledgement != nil && !acceptedSnapshots[message.event.index] && message.event.index < best.index {
best = *message.event
}
continue
}
if message.snapshot == nil {
continue
}
protocol.accept(*message.snapshot)
if protocol.err != nil {
set.commitHealthError(redactHeldSessionSetHealthError(protocol.err))
set.stopHealthWatchers()
set.healthWG.Wait()
acknowledgeShutdown()
return
}
if set.healthSnapshotAcceptedHook != nil {
set.healthSnapshotAcceptedHook(*message.snapshot)
}
acceptedSnapshots[message.snapshot.index] = true
if message.snapshot.event != nil && message.snapshot.event.index < best.index {
best = *message.snapshot.event
}
case <-set.healthStop:
acknowledgeShutdown()
return
case request := <-set.healthShutdownRequests:
if set.healthShutdownAcceptedHook != nil {
set.healthShutdownAcceptedHook()
}
shutdownAcknowledgement = request.ack
}
}
set.commitHealthError(best.err)
set.stopHealthWatchers()
set.healthWG.Wait()
acknowledgeShutdown()
}
func (set *heldSessionSet) commitHealthError(err error) {
if err == nil {
return
}
set.healthMu.Lock()
defer set.healthMu.Unlock()
if set.healthError == nil {
set.healthError = err
set.healthDoneOnce.Do(func() { close(set.healthDone) })
}
}
@@ -0,0 +1,246 @@
//go:build linux
package scenario
import (
"context"
"errors"
"sort"
"testing"
"github.com/stretchr/testify/require"
)
func collectHealthIndexes(events <-chan int) []int {
indexes := make([]int, 0, heldSessionSetHealthMemberCount)
for index := 0; index < heldSessionSetHealthMemberCount; index++ {
indexes = append(indexes, <-events)
}
return indexes
}
func TestHeldSessionSetHealthActiveEpochCompletesEveryRequestBeforeShutdown(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
requestIndexes := make(chan int, heldSessionSetHealthMemberCount)
replyIndexes := make(chan int, heldSessionSetHealthMemberCount)
releaseBroadcast := make(chan struct{})
broadcastReached := make(chan struct{})
fixture.input.testHealthSnapshotRequestHook = func(index int) {
requestIndexes <- index
if index == 5 {
close(broadcastReached)
<-releaseBroadcast
}
}
fixture.input.testHealthSnapshotReplyHook = func(index int) { replyIndexes <- index }
set := fixture.returnedSet(t)
fixture.coordinator.sessions[11].closeResult = errors.New("active epoch trigger")
closeStarted := make(chan struct{})
fixture.coordinator.sessions[11].closeStarted = closeStarted
fixture.coordinator.sessions[11].prematureClose()
<-broadcastReached
closeResult := make(chan error, 1)
go func() { closeResult <- set.Close(context.Background()) }()
close(releaseBroadcast)
require.NoError(t, <-closeResult)
<-closeStarted
select {
case <-set.coordinatorDone:
default:
t.Fatal("member Close started before coordinator completion")
}
requests := collectHealthIndexes(requestIndexes)
replies := collectHealthIndexes(replyIndexes)
sort.Ints(requests)
sort.Ints(replies)
require.Equal(t, requests, replies)
require.Equal(t, []int{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11}, requests)
require.Error(t, set.WaitHealthy(context.Background()))
}
func TestHeldSessionSetHealthFinalCutRetainsAcceptedSnapshotAgainstLaterClosure(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
highEventSent := make(chan struct{})
shutdownAccepted := make(chan struct{})
shutdownAcknowledged := make(chan struct{})
indexZeroSendReady := make(chan struct{})
indexZeroSendAllowed := make(chan struct{})
indexZeroAccepted := make(chan struct{})
indexZeroClosureObserved := make(chan struct{})
releaseOtherSnapshots := make(chan struct{})
lateClosure := errors.New("late closure")
fixture.input.testHealthEventHook = func(message heldHealthMessage) {
if message.event != nil && message.event.index == 11 {
close(highEventSent)
}
}
fixture.input.testHealthShutdownAcceptedHook = func() { close(shutdownAccepted) }
fixture.input.testHealthShutdownAcknowledgedHook = func() { close(shutdownAcknowledged) }
fixture.input.testHealthSnapshotAcceptedHook = func(snapshot heldHealthSnapshot) {
if snapshot.index == 0 {
close(indexZeroAccepted)
}
}
fixture.input.testHealthSnapshotSendHook = func(index int) {
if index == 0 {
close(indexZeroSendReady)
<-indexZeroSendAllowed
} else {
<-releaseOtherSnapshots
}
}
fixture.input.testHealthClosureObservedHook = func(index int) {
if index == 0 {
close(indexZeroClosureObserved)
}
}
set := fixture.returnedSet(t)
triggerError := errors.New("trigger")
fixture.coordinator.sessions[11].closeResult = triggerError
fixture.coordinator.sessions[11].prematureClose()
<-highEventSent
closeResult := make(chan error, 1)
go func() { closeResult <- set.Close(context.Background()) }()
<-indexZeroSendReady
<-shutdownAccepted
close(indexZeroSendAllowed)
<-indexZeroAccepted
select {
case <-fixture.closeOrder:
t.Fatal("member cleanup started before remaining snapshots were released")
default:
}
fixture.coordinator.sessions[0].closeResult = lateClosure
fixture.coordinator.sessions[0].prematureClose()
<-indexZeroClosureObserved
select {
case <-fixture.closeOrder:
t.Fatal("member cleanup started before blocked snapshots were released")
default:
}
close(releaseOtherSnapshots)
require.NoError(t, <-closeResult)
require.ErrorIs(t, set.WaitHealthy(context.Background()), triggerError)
require.NotErrorIs(t, set.WaitHealthy(context.Background()), lateClosure)
<-shutdownAcknowledged
<-set.coordinatorDone
for _, watcherDone := range set.healthWatcherDone {
<-watcherDone
}
<-set.closeDone
for index := 11; index >= 0; index-- {
require.Equal(t, index, <-fixture.closeOrder)
require.Equal(t, index, <-fixture.waitOrder)
}
}
func TestHeldSessionSetHealthEpochRejectsDuplicateReply(t *testing.T) {
protocol := newHeldHealthEpochProtocol(1)
protocol.accept(heldHealthSnapshot{epoch: 1, index: 0})
protocol.accept(heldHealthSnapshot{epoch: 1, index: 0})
for index := 1; index < heldSessionSetHealthMemberCount; index++ {
protocol.accept(heldHealthSnapshot{epoch: 1, index: index})
}
require.ErrorIs(t, protocol.err, ErrHeldSessionSetHealthProtocol)
require.False(t, protocol.committed)
}
func TestHeldSessionSetHealthEpochRejectsOutOfRangeReply(t *testing.T) {
protocol := newHeldHealthEpochProtocol(1)
protocol.accept(heldHealthSnapshot{epoch: 1, index: heldSessionSetHealthMemberCount})
require.ErrorIs(t, protocol.err, ErrHeldSessionSetHealthProtocol)
require.False(t, protocol.committed)
}
func TestHeldSessionSetHealthEpochRejectsMissingIndexSubstitution(t *testing.T) {
protocol := newHeldHealthEpochProtocol(1)
for index := 0; index < heldSessionSetHealthMemberCount-1; index++ {
protocol.accept(heldHealthSnapshot{epoch: 1, index: index})
}
protocol.accept(heldHealthSnapshot{epoch: 1, index: heldSessionSetHealthMemberCount - 2})
require.ErrorIs(t, protocol.err, ErrHeldSessionSetHealthProtocol)
require.False(t, protocol.committed)
}
func TestHeldSessionSetHealthProtocolErrorAcknowledgesShutdownAndFinishes(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
requestReached := make(chan struct{})
releaseRequest := make(chan struct{})
fixture.input.testHealthSnapshotRequestHook = func(index int) {
if index != 4 {
return
}
close(requestReached)
<-releaseRequest
}
set := fixture.returnedSet(t)
fixture.coordinator.sessions[11].closeResult = sensitiveError("protocol-trigger")
fixture.coordinator.sessions[11].prematureClose()
<-requestReached
closeResult := make(chan error, 1)
go func() { closeResult <- set.Close(context.Background()) }()
set.healthEvents <- heldHealthMessage{snapshot: &heldHealthSnapshot{epoch: 1, index: heldSessionSetHealthMemberCount}}
close(releaseRequest)
require.NoError(t, <-closeResult)
require.ErrorIs(t, set.WaitHealthy(context.Background()), ErrHeldSessionSetHealthProtocol)
require.ErrorIs(t, set.WaitHealthy(context.Background()), ErrHeldSessionSetHealth)
require.Equal(t, "held session set health failed", set.WaitHealthy(context.Background()).Error())
<-set.coordinatorDone
set.healthWG.Wait()
}
func TestHeldSessionSetHealthEpochIncludesCloseBeforeOpenReply(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
fixture.input.testHealthSnapshotRequestHook = func(index int) {
if index == 0 {
fixture.coordinator.sessions[0].prematureClose()
}
}
fixture.input.testHealthSnapshotReplyHook = func(int) {}
set := fixture.returnedSet(t)
highError := sensitiveError("epoch-high")
lowError := sensitiveError("epoch-low")
fixture.coordinator.sessions[11].closeResult = highError
fixture.coordinator.sessions[0].closeResult = lowError
fixture.coordinator.sessions[11].prematureClose()
err := set.WaitHealthy(context.Background())
require.ErrorIs(t, err, lowError)
require.NotErrorIs(t, err, highError)
require.Equal(t, "held session set health failed", err.Error())
require.NotContains(t, err.Error(), "secret-token")
require.NoError(t, set.Close(context.Background()))
}
func TestHeldSessionSetHealthEpochExcludesCloseAfterOpenReply(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
fixture.input.testHealthSnapshotRequestHook = func(int) {}
fixture.input.testHealthSnapshotReplyHook = func(index int) {
if index == 0 {
fixture.coordinator.sessions[0].prematureClose()
}
}
set := fixture.returnedSet(t)
highError := sensitiveError("epoch-high")
lowError := sensitiveError("epoch-low")
fixture.coordinator.sessions[11].closeResult = highError
fixture.coordinator.sessions[0].closeResult = lowError
fixture.coordinator.sessions[11].prematureClose()
err := set.WaitHealthy(context.Background())
require.ErrorIs(t, err, highError)
require.NotErrorIs(t, err, lowError)
require.Equal(t, "held session set health failed", err.Error())
require.NoError(t, set.Close(context.Background()))
}
func TestHeldSessionSetHealthEpochCommitsWithoutUnrelatedDone(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
set := fixture.returnedSet(t)
triggerError := sensitiveError("epoch-trigger")
fixture.coordinator.sessions[11].closeResult = triggerError
fixture.coordinator.sessions[11].prematureClose()
err := set.WaitHealthy(context.Background())
require.ErrorIs(t, err, triggerError)
require.Equal(t, "held session set health failed", err.Error())
require.NoError(t, set.Close(context.Background()))
require.ErrorIs(t, set.WaitHealthy(context.Background()), triggerError)
}
@@ -0,0 +1,68 @@
//go:build linux
package scenario
import (
"context"
"errors"
"testing"
"github.com/stretchr/testify/require"
)
func TestHeldSessionSetPublicCloseJoinsReverseCleanupAndStateErrors(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
set := fixture.returnedSet(t)
closeError := sensitiveError("close")
waitClosedError := sensitiveError("wait-closed")
stateError := sensitiveError("absent")
for index, session := range fixture.coordinator.sessions {
if index%3 == 0 {
session.closeError = closeError
}
if index%4 == 0 {
session.waitClosedError = waitClosedError
}
}
fixture.state.absentAggregateError = stateError
err := set.Close(context.Background())
require.Error(t, err)
require.ErrorIs(t, err, closeError)
require.ErrorIs(t, err, waitClosedError)
require.ErrorIs(t, err, stateError)
require.Equal(t, "held session set operation failed", err.Error())
require.NotContains(t, err.Error(), "secret-token")
for index := 11; index >= 0; index-- {
require.Equal(t, index, <-fixture.closeOrder)
require.Equal(t, index, <-fixture.waitOrder)
}
require.Len(t, fixture.state.absent, 12)
}
func TestHeldSessionSetPublicConcurrentAndCanceledCloseShareOwnerResult(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
set := fixture.returnedSet(t)
closeError := errors.New("owner close error")
fixture.coordinator.sessions[0].closeError = closeError
closeStarted := make(chan struct{})
closeRelease := make(chan struct{})
fixture.coordinator.sessions[11].closeStarted = closeStarted
fixture.coordinator.sessions[11].closeRelease = closeRelease
ownerResult := make(chan error, 1)
go func() { ownerResult <- set.Close(context.Background()) }()
<-closeStarted
canceled, cancel := context.WithCancel(context.Background())
waiterResult := make(chan error, 1)
go func() { waiterResult <- set.Close(canceled) }()
cancel()
require.ErrorIs(t, <-waiterResult, context.Canceled)
close(closeRelease)
require.ErrorIs(t, <-ownerResult, closeError)
thirdResult := set.Close(context.Background())
require.ErrorIs(t, thirdResult, closeError)
require.Equal(t, "held session set operation failed", thirdResult.Error())
for index := 11; index >= 0; index-- {
require.Equal(t, index, <-fixture.closeOrder)
require.Equal(t, index, <-fixture.waitOrder)
}
}
@@ -0,0 +1,231 @@
//go:build linux
package scenario
import (
"context"
"errors"
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestHeldSessionSetPublicConstructorFailureRollsBackEverySuccess(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
firstError := sensitiveError("constructor-first")
secondError := sensitiveError("constructor-second")
fixture.coordinator.errors[fixture.plan.Sessions[2].ID.String()] = firstError
fixture.coordinator.errors[fixture.plan.Sessions[7].ID.String()] = secondError
fixture.coordinator.errorReady[2] = make(chan struct{})
fixture.coordinator.errorReady[7] = make(chan struct{})
result := make(chan error, 1)
go func() { _, err := NewHeldSessionSet(context.Background(), fixture.input); result <- err }()
for range fixture.plan.Sessions {
<-fixture.coordinator.ready
}
fixture.coordinator.releaseAll()
for _, index := range []int{0, 1, 3, 4, 5, 6, 8, 9, 10, 11} {
fixture.coordinator.release(index)
}
completed := make([]heldSessionConstructorCompletion, 0, len(fixture.plan.Sessions))
acquired := make([]int, 0, 10)
for range []int{0, 1, 3, 4, 5, 6, 8, 9, 10, 11} {
completion := <-fixture.coordinator.completed
completed = append(completed, completion)
acquired = append(acquired, <-fixture.coordinator.acquired)
}
fixture.coordinator.release(2)
fixture.coordinator.releaseError(2)
firstCompletion := <-fixture.coordinator.completed
completed = append(completed, firstCompletion)
require.Equal(t, firstError, firstCompletion.err)
fixture.coordinator.release(7)
fixture.coordinator.releaseError(7)
secondCompletion := <-fixture.coordinator.completed
completed = append(completed, secondCompletion)
require.Equal(t, secondError, secondCompletion.err)
err := <-result
require.Error(t, err)
require.ErrorIs(t, err, firstError)
require.ErrorIs(t, err, secondError)
require.NotContains(t, err.Error(), "secret-token")
require.NotContains(t, err.Error(), "secret-path")
require.ElementsMatch(t, []int{0, 1, 3, 4, 5, 6, 8, 9, 10, 11}, acquired)
wantRollback := []int{11, 10, 9, 8, 6, 5, 4, 3, 1, 0}
for _, index := range wantRollback {
require.Equal(t, index, <-fixture.closeOrder)
require.Equal(t, index, <-fixture.waitOrder)
}
require.ElementsMatch(t, acquiredStreamIDs(fixture, acquired), fixture.state.absent)
}
func TestHeldSessionSetPublicConstructorFailureCancelsSiblingAcquisition(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
trigger := sensitiveError("constructor-trigger")
blockedIndex := 7
fixture.coordinator.blockedIndex = blockedIndex
fixture.coordinator.blockedWaiting = make(chan struct{}, 1)
fixture.coordinator.blockedCanceled = make(chan struct{}, 1)
triggerIndex := 2
fixture.coordinator.errors[fixture.plan.Sessions[triggerIndex].ID.String()] = trigger
fixture.coordinator.errorReady[triggerIndex] = make(chan struct{})
outerContext := context.Background()
result := make(chan error, 1)
go func() { _, err := NewHeldSessionSet(outerContext, fixture.input); result <- err }()
for range fixture.plan.Sessions {
<-fixture.coordinator.ready
}
fixture.coordinator.releaseAll()
for _, index := range []int{0, 1, 3, 4, 5, 6, 8, 9, 10, 11} {
fixture.coordinator.release(index)
}
for range []int{0, 1, 3, 4, 5, 6, 8, 9, 10, 11} {
completion := <-fixture.coordinator.completed
require.NoError(t, completion.err)
<-fixture.coordinator.acquired
}
fixture.coordinator.release(blockedIndex)
select {
case <-fixture.coordinator.blockedWaiting:
case <-time.After(time.Second):
t.Fatal("blocked constructor did not start")
}
fixture.coordinator.release(triggerIndex)
fixture.coordinator.releaseError(triggerIndex)
triggerCompletion := <-fixture.coordinator.completed
require.ErrorIs(t, triggerCompletion.err, trigger)
select {
case <-fixture.coordinator.blockedCanceled:
case <-time.After(time.Second):
t.Fatal("sibling constructor was not canceled")
}
require.ErrorIs(t, <-result, trigger)
select {
case <-outerContext.Done():
t.Fatal("outer context was canceled")
default:
}
for _, index := range []int{11, 10, 9, 8, 6, 5, 4, 3, 1, 0} {
require.Equal(t, index, <-fixture.closeOrder)
require.Equal(t, index, <-fixture.waitOrder)
}
}
func TestHeldSessionSetPublicCanceledConstructionWaitsAndRollsBack(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
fixture.coordinator.lateSuccessIndex = 11
fixture.coordinator.lateSuccessWaiting = make(chan struct{}, 1)
fixture.coordinator.lateSuccessReady = make(chan struct{}, 1)
fixture.coordinator.lateSuccessRelease = make(chan struct{})
ctx, cancel := context.WithCancel(context.Background())
result := make(chan error, 1)
go func() { _, err := NewHeldSessionSet(ctx, fixture.input); result <- err }()
for range fixture.plan.Sessions {
<-fixture.coordinator.ready
}
for _, index := range []int{0, 4} {
fixture.coordinator.release(index)
}
fixture.coordinator.releaseAll()
for range []int{0, 4} {
<-fixture.coordinator.acquired
}
<-fixture.coordinator.lateSuccessWaiting
cancel()
fixture.coordinator.release(11)
<-fixture.coordinator.lateSuccessReady
close(fixture.coordinator.lateSuccessRelease)
err := <-result
require.ErrorIs(t, err, context.Canceled)
completions := make([]heldSessionConstructorCompletion, 0, len(fixture.plan.Sessions))
for range fixture.plan.Sessions {
completions = append(completions, <-fixture.coordinator.completed)
}
require.ElementsMatch(t, []int{0, 4, 11}, completionIndexes(completions, true))
require.ElementsMatch(t, []int{1, 2, 3, 5, 6, 7, 8, 9, 10}, completionIndexes(completions, false))
require.Equal(t, []int{11, 4, 0}, collectOrder(fixture.closeOrder, 3))
require.Equal(t, []int{11, 4, 0}, collectOrder(fixture.waitOrder, 3))
require.ElementsMatch(t, acquiredStreamIDs(fixture, []int{0, 4, 11}), fixture.state.absent)
}
func completionIndexes(completions []heldSessionConstructorCompletion, acquired bool) []int {
indexes := make([]int, 0, len(completions))
for _, completion := range completions {
if completion.acquired == acquired {
indexes = append(indexes, completion.index)
}
}
return indexes
}
func acquiredStreamIDs(fixture *heldSessionSetPublicFixture, indexes []int) []string {
ids := make([]string, 0, len(indexes))
for _, index := range indexes {
ids = append(ids, fixture.coordinator.sessions[index].streamID)
}
return ids
}
func collectOrder(events <-chan int, count int) []int {
order := make([]int, 0, count)
for index := 0; index < count; index++ {
order = append(order, <-events)
}
return order
}
func TestHeldSessionSetPublicAllLiveFailuresCloseMembers(t *testing.T) {
cases := []struct {
name string
mutate func(*heldSessionSetPublicFixture)
want error
}{
{name: "wait live", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.coordinator.sessions[0].waitLiveError = sensitiveError("live")
}, want: ErrHeldSessionSetOperation},
{name: "wrong plan", mutate: func(fixture *heldSessionSetPublicFixture) { fixture.coordinator.sessions[0].plan.Ordinal++ }, want: ErrHeldSessionSetOperation},
{name: "empty stream", mutate: func(fixture *heldSessionSetPublicFixture) { fixture.coordinator.sessions[0].streamID = "" }, want: ErrHeldSessionSetOperation},
{name: "duplicate stream", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.coordinator.sessions[1].streamID = fixture.coordinator.sessions[0].streamID
}, want: ErrHeldSessionSetOperation},
{name: "present state", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.state.presentAggregateError = sensitiveError("present")
}, want: ErrHeldSessionSetOperation},
{name: "aggregate state", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.state.presentAggregateError = sensitiveError("aggregate")
}, want: ErrHeldSessionSetOperation},
{name: "closed before return", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.coordinator.sessions[0].prematureClose()
fixture.coordinator.sessions[0].closeResult = sensitiveError("closed")
}, want: ErrHeldSessionSetOperation},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
testCase.mutate(fixture)
result := make(chan error, 1)
go func() { _, err := NewHeldSessionSet(context.Background(), fixture.input); result <- err }()
for range fixture.plan.Sessions {
<-fixture.coordinator.ready
}
fixture.releaseConstructors()
err := <-result
require.ErrorIs(t, err, testCase.want)
require.NotContains(t, err.Error(), "secret-token")
})
}
}
func TestHeldSessionSetPublicTopologyFailureHasNoSideEffects(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
fixture.input.ControlClient = nil
fixture.input.Dependencies = HeldSessionSetDependencies{Terminal: func(context.Context, heldTerminalInput) (heldSession, error) {
return nil, errors.New("constructor called")
}}
_, err := NewHeldSessionSet(context.Background(), fixture.input)
require.ErrorIs(t, err, ErrInvalidHeldSessionSetTopology)
require.Empty(t, fixture.state.present)
require.Empty(t, fixture.state.absent)
require.Empty(t, fixture.coordinator.ready)
}
@@ -0,0 +1,152 @@
//go:build linux
package scenario
import (
"context"
"errors"
"sync"
"testing"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
)
type heldSessionSetPublicFixture struct {
plan StressPlan
input HeldSessionSetInput
coordinator *heldSessionSetConstructorCoordinator
state *heldSessionSetStateFake
closeOrder chan int
waitOrder chan int
counters *heldSessionSetCallCounters
}
type heldSessionSetCallCounters struct {
mu sync.Mutex
inspect int
snapshot int
observe int
waitState int
terminal int
nat int
fm int
}
func (counters *heldSessionSetCallCounters) add(field *int) {
counters.mu.Lock()
(*field)++
counters.mu.Unlock()
}
func (counters *heldSessionSetCallCounters) values() heldSessionSetCallCounters {
counters.mu.Lock()
defer counters.mu.Unlock()
return heldSessionSetCallCounters{inspect: counters.inspect, snapshot: counters.snapshot, observe: counters.observe, waitState: counters.waitState, terminal: counters.terminal, nat: counters.nat, fm: counters.fm}
}
func (fixture *heldSessionSetPublicFixture) returnedSet(t *testing.T) *heldSessionSet {
t.Helper()
result := make(chan *heldSessionSet, 1)
errResult := make(chan error, 1)
go func() {
set, err := NewHeldSessionSet(context.Background(), fixture.input)
result <- set
errResult <- err
}()
for range fixture.plan.Sessions {
<-fixture.coordinator.ready
}
fixture.releaseConstructors()
set := <-result
if err := <-errResult; err != nil {
t.Fatal(err)
}
return set
}
func newHeldSessionSetPublicFixture(t *testing.T) *heldSessionSetPublicFixture {
t.Helper()
profile, err := contract.ProfileByName(string(contract.ProfilePRFull))
if err != nil {
t.Fatal(err)
}
plan, err := GenerateStressPlan(profile, contract.DefaultSeed)
if err != nil {
t.Fatal(err)
}
coordinator := newHeldSessionSetConstructorCoordinator()
closeOrder := make(chan int, len(plan.Sessions))
waitOrder := make(chan int, len(plan.Sessions))
for index, sessionPlan := range plan.Sessions {
coordinator.indices[sessionPlan.ID.String()] = index
coordinator.sessions[index] = newHeldSessionSetTestSession(index, sessionPlan, closeOrder, waitOrder)
}
topology := make([]HeldSessionAgent, heldSessionSetAgentCount)
agentFacts := make(map[*agent.Agent]heldSessionAgentFacts, heldSessionSetAgentCount)
controlServerIDs := make([]uint64, heldSessionSetAgentCount)
for index := range topology {
ordinal, ordinalErr := NewStressAgentOrdinal(index + 1)
if ordinalErr != nil {
t.Fatal(ordinalErr)
}
agentInstance := &agent.Agent{}
agentUUID := "uuid-" + string(rune('a'+index))
topology[index] = HeldSessionAgent{Ordinal: ordinal, Agent: agentInstance, PATClient: &client.Client{}, Readiness: agent.Readiness{ServerID: uint64(index + 1), UUID: agentUUID, Version: "test", Online: true, VersionObserved: true, RequestTaskEstablished: true, StateReceiptObserved: true}}
agentFacts[agentInstance] = heldSessionAgentFacts{PID: 1, UUID: agentUUID}
controlServerIDs[index] = uint64(index + 1)
}
state := &heldSessionSetStateFake{baseline: client.IOStreamState{Count: 7}, expectedPresentCount: 19, expectedAbsentCount: 7, presentStreamErrors: make(map[string]error), absentStreamErrors: make(map[string]error), aggregatePredicatesEmpty: true, waitCalls: make(chan client.IOStreamStateExpectation, 40)}
counters := &heldSessionSetCallCounters{}
fixture := &heldSessionSetPublicFixture{plan: plan, coordinator: coordinator, state: state, closeOrder: closeOrder, waitOrder: waitOrder, counters: counters}
dependencies := fixture.dependencies()
dependencies.InspectAgent = func(instance *agent.Agent) heldSessionAgentFacts {
counters.add(&counters.inspect)
return agentFacts[instance]
}
dependencies.ObserveState = func(*client.Client) heldSessionSetStateObserver { counters.add(&counters.observe); return state }
fixture.input = HeldSessionSetInput{Dashboard: &dashboard.Dashboard{}, Plan: plan, Topology: topology, ControlClient: &client.Client{}, ControlServerIDs: controlServerIDs, Dependencies: dependencies}
return fixture
}
func (fixture *heldSessionSetPublicFixture) dependencies() HeldSessionSetDependencies {
construct := func(ctx, lifetimeContext context.Context, plan StressSessionPlan) (heldSession, error) {
fixture.coordinator.contextSeen <- lifetimeContext
return fixture.coordinator.construct(ctx, plan)
}
return HeldSessionSetDependencies{
Terminal: func(ctx context.Context, input heldTerminalInput) (heldSession, error) {
fixture.counters.add(&fixture.counters.terminal)
return construct(ctx, input.LifetimeContext, input.Plan)
},
NAT: func(ctx context.Context, input heldNATInput) (heldSession, error) {
fixture.counters.add(&fixture.counters.nat)
return construct(ctx, input.LifetimeContext, input.Plan)
},
FM: func(ctx context.Context, input heldLegacyFMInput) (heldSession, error) {
fixture.counters.add(&fixture.counters.fm)
return construct(ctx, input.LifetimeContext, input.Plan)
},
Snapshot: func(context.Context, heldSessionSetStateObserver) (client.IOStreamState, error) {
fixture.counters.add(&fixture.counters.snapshot)
return fixture.state.baseline, fixture.state.snapshotError
},
WaitState: func(ctx context.Context, state heldSessionSetStateObserver, expectation client.IOStreamStateExpectation) (client.IOStreamState, error) {
fixture.counters.add(&fixture.counters.waitState)
return state.WaitForIOStreamState(ctx, expectation)
},
}
}
func (fixture *heldSessionSetPublicFixture) releaseConstructors() {
fixture.coordinator.releaseAll()
for index := range fixture.plan.Sessions {
fixture.coordinator.release(index)
}
}
func sensitiveError(label string) error {
return errors.New(label + " token=secret-token path=/secret/path uuid=secret-uuid stream=secret-stream")
}
@@ -0,0 +1,75 @@
//go:build linux
package scenario
import (
"context"
"errors"
"testing"
"github.com/stretchr/testify/require"
)
func TestHeldSessionSetPublicPrematureHealthIsRetainedAndCanonical(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
firstError := sensitiveError("health-first")
secondError := sensitiveError("health-second")
fixture.coordinator.sessions[0].closeResult = firstError
fixture.coordinator.sessions[1].closeResult = secondError
indexOneClosureObserved := make(chan struct{})
fixture.input.testHealthSnapshotRequestHook = func(index int) {
if index == 0 {
fixture.coordinator.sessions[0].prematureClose()
}
}
fixture.input.testHealthSnapshotReplyHook = func(index int) {
_ = index
}
fixture.input.testHealthClosureObservedHook = func(index int) {
if index == 1 {
close(indexOneClosureObserved)
}
}
set := fixture.returnedSet(t)
fixture.coordinator.sessions[1].prematureClose()
<-indexOneClosureObserved
fixture.input.testHealthSnapshotRequestHook = func(index int) {
if index == 0 {
fixture.coordinator.sessions[0].prematureClose()
}
}
err := set.WaitHealthy(context.Background())
require.ErrorIs(t, err, firstError)
require.NotErrorIs(t, err, secondError)
require.Equal(t, "held session set health failed", err.Error())
require.NotContains(t, err.Error(), "secret-token")
for index := 2; index < 12; index++ {
fixture.coordinator.sessions[index].prematureClose()
}
require.ErrorIs(t, set.WaitHealthy(context.Background()), firstError)
require.NoError(t, set.Close(context.Background()))
require.ErrorIs(t, set.WaitHealthy(context.Background()), firstError)
set.healthWG.Wait()
}
func TestHeldSessionSetPublicCanceledHealthWaiterRetainsLaterResult(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
set := fixture.returnedSet(t)
waitContext, cancel := context.WithCancel(context.Background())
cancel()
require.ErrorIs(t, set.WaitHealthy(waitContext), context.Canceled)
healthError := errors.New("health failure")
fixture.coordinator.sessions[3].closeResult = healthError
fixture.coordinator.sessions[3].prematureClose()
require.ErrorIs(t, set.WaitHealthy(context.Background()), healthError)
require.NoError(t, set.Close(context.Background()))
set.healthWG.Wait()
}
func TestHeldSessionSetPublicOwnedCloseDoesNotCreateHealthFailure(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
set := fixture.returnedSet(t)
require.NoError(t, set.Close(context.Background()))
require.NoError(t, set.WaitHealthy(context.Background()))
set.healthWG.Wait()
}
@@ -0,0 +1,81 @@
//go:build linux
package scenario
import (
"context"
"sort"
"testing"
"github.com/stretchr/testify/require"
)
func TestHeldSessionSetPublicCoordinatorRejectsMalformedSnapshots(t *testing.T) {
cases := []struct {
name string
target int
mutate func(heldHealthSnapshotRequest) heldHealthSnapshot
validFirst bool
}{
{name: "stale epoch", target: 0, mutate: func(request heldHealthSnapshotRequest) heldHealthSnapshot {
return heldHealthSnapshot{epoch: request.epoch - 1, index: 0}
}},
{name: "duplicate index", target: 1, mutate: func(request heldHealthSnapshotRequest) heldHealthSnapshot {
return heldHealthSnapshot{epoch: request.epoch, index: 0}
}},
{name: "out of range index", target: 0, mutate: func(request heldHealthSnapshotRequest) heldHealthSnapshot {
return heldHealthSnapshot{epoch: request.epoch, index: heldSessionSetHealthMemberCount}
}},
{name: "missing index substitution", target: 11, mutate: func(request heldHealthSnapshotRequest) heldHealthSnapshot {
return heldHealthSnapshot{epoch: request.epoch, index: heldSessionSetHealthMemberCount - 2}
}},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
requests := make(chan int, heldSessionSetHealthMemberCount)
overrideIndexes := make(chan int, 1)
malformedSent := make(chan struct{})
fixture.input.testHealthSnapshotRequestHook = func(index int) { requests <- index }
fixture.input.testHealthSnapshotOverrideHook = func(index int, request heldHealthSnapshotRequest) *heldHealthSnapshot {
if index != testCase.target {
return nil
}
select {
case <-malformedSent:
default:
close(malformedSent)
}
malformed := testCase.mutate(request)
overrideIndexes <- malformed.index
return &malformed
}
set := fixture.returnedSet(t)
closeResult := make(chan error, 1)
go func() { closeResult <- set.Close(context.Background()) }()
err := set.WaitHealthy(context.Background())
require.Error(t, err)
require.Equal(t, "held session set health failed", err.Error())
require.ErrorIs(t, err, ErrHeldSessionSetHealth)
require.ErrorIs(t, err, ErrHeldSessionSetHealthProtocol)
require.Equal(t, []int{testCase.mutate(heldHealthSnapshotRequest{epoch: 1}).index}, []int{<-overrideIndexes})
require.Equal(t, []int{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11}, sortedHealthIndexes(requests))
require.NoError(t, <-closeResult)
<-set.coordinatorDone
<-set.closeDone
for _, watcherDone := range set.healthWatcherDone {
<-watcherDone
}
for index := len(fixture.coordinator.sessions) - 1; index >= 0; index-- {
require.Equal(t, index, <-fixture.closeOrder)
require.Equal(t, index, <-fixture.waitOrder)
}
})
}
}
func sortedHealthIndexes(events <-chan int) []int {
indexes := collectHealthIndexes(events)
sort.Ints(indexes)
return indexes
}
@@ -0,0 +1,193 @@
//go:build linux
package scenario
import (
"context"
"errors"
"testing"
"github.com/stretchr/testify/require"
)
var heldSessionSetSensitiveFragments = []string{"secret-token", "/secret/path", "secret-uuid", "secret-stream", "server=991", "authorization="}
func requireHeldSessionSetRedacted(t *testing.T, err error, class error, causes ...error) {
t.Helper()
require.Error(t, err)
require.Equal(t, class.Error(), err.Error())
require.ErrorIs(t, err, class)
for _, cause := range causes {
require.ErrorIs(t, err, cause)
}
for _, fragment := range heldSessionSetSensitiveFragments {
require.NotContains(t, err.Error(), fragment, fragment)
}
}
func TestHeldSessionSetPublicRedactionInitialSnapshot(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
sentinel := sensitiveBoundaryError("snapshot")
fixture.state.snapshotError = sentinel
_, err := NewHeldSessionSet(context.Background(), fixture.input)
requireHeldSessionSetRedacted(t, err, ErrHeldSessionSetOperation, sentinel)
}
func TestHeldSessionSetPublicRedactionConstructorsByKind(t *testing.T) {
kinds := []StressSessionKind{StressSessionTerminal, StressSessionNAT, StressSessionFM}
for _, kind := range kinds {
t.Run(string(kind), func(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
index := sessionIndexByKind(fixture, kind)
sentinel := sensitiveBoundaryError(string(kind))
fixture.coordinator.errors[fixture.plan.Sessions[index].ID.String()] = sentinel
fixture.coordinator.errorReady[index] = make(chan struct{})
result := make(chan error, 1)
go func() {
_, err := NewHeldSessionSet(context.Background(), fixture.input)
result <- err
}()
for range fixture.plan.Sessions {
<-fixture.coordinator.ready
}
fixture.coordinator.releaseAll()
for member := range fixture.plan.Sessions {
if member != index {
fixture.coordinator.release(member)
}
}
fixture.coordinator.release(index)
fixture.coordinator.releaseError(index)
for range fixture.plan.Sessions {
<-fixture.coordinator.completed
}
requireHeldSessionSetRedacted(t, <-result, ErrHeldSessionSetOperation, sentinel)
})
}
}
func TestHeldSessionSetPublicRedactionWaitLiveAndStateBoundaries(t *testing.T) {
cases := []struct {
name string
closeSet bool
mutate func(*heldSessionSetPublicFixture, error)
}{
{name: "wait live", mutate: func(fixture *heldSessionSetPublicFixture, sentinel error) {
fixture.coordinator.sessions[0].waitLiveError = sentinel
}},
{name: "present per stream", mutate: func(fixture *heldSessionSetPublicFixture, sentinel error) {
fixture.state.presentStreamErrors[fixture.coordinator.sessions[0].streamID] = sentinel
}},
{name: "present aggregate", mutate: func(fixture *heldSessionSetPublicFixture, sentinel error) {
fixture.state.presentAggregateError = sentinel
}},
{name: "absent per stream", closeSet: true, mutate: func(fixture *heldSessionSetPublicFixture, sentinel error) {
fixture.state.absentStreamErrors[fixture.coordinator.sessions[0].streamID] = sentinel
}},
{name: "absent aggregate", closeSet: true, mutate: func(fixture *heldSessionSetPublicFixture, sentinel error) {
fixture.state.absentAggregateError = sentinel
}},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
sentinel := sensitiveBoundaryError(testCase.name)
testCase.mutate(fixture, sentinel)
var err error
if testCase.closeSet {
set := fixture.returnedSet(t)
err = set.Close(context.Background())
} else {
_, err = newHeldSessionSetWithReleasedConstructors(fixture)
}
requireHeldSessionSetRedacted(t, err, ErrHeldSessionSetOperation, sentinel)
})
}
}
func TestHeldSessionSetPublicRedactionHealthCloseAndAbsentBoundaries(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
set := fixture.returnedSet(t)
closeSentinel := sensitiveBoundaryError("close")
waitSentinel := sensitiveBoundaryError("wait-closed")
absentSentinel := sensitiveBoundaryError("absent")
fixture.coordinator.sessions[0].closeError = closeSentinel
fixture.coordinator.sessions[1].waitClosedError = waitSentinel
fixture.state.absentAggregateError = absentSentinel
err := set.Close(context.Background())
requireHeldSessionSetRedacted(t, err, ErrHeldSessionSetOperation, closeSentinel, waitSentinel, absentSentinel)
}
func TestHeldSessionSetPublicRedactionBaselineRestorationAggregate(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
set := fixture.returnedSet(t)
sentinel := sensitiveBoundaryError("baseline-restoration")
fixture.state.absentAggregateError = sentinel
err := set.Close(context.Background())
requireHeldSessionSetRedacted(t, err, ErrHeldSessionSetOperation, sentinel)
require.Equal(t, fixture.state.expectedAbsentCount, fixture.state.baseline.Count)
require.ElementsMatch(t, canonicalStreamIDs(fixture), fixture.state.absent)
}
func TestHeldSessionSetPublicRedactionPrematureHealthRetainsTypedErrors(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
set := fixture.returnedSet(t)
sentinel := sensitiveBoundaryError("premature")
fixture.coordinator.sessions[0].closeResult = sentinel
fixture.coordinator.sessions[0].prematureClose()
err := set.WaitHealthy(context.Background())
requireHeldSessionSetRedacted(t, err, ErrHeldSessionSetHealth, sentinel)
require.NoError(t, set.Close(context.Background()))
}
func TestHeldSessionSetPublicRedactionPrematureHealthNilCloseResult(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
set := fixture.returnedSet(t)
fixture.coordinator.sessions[0].prematureClose()
err := set.WaitHealthy(context.Background())
requireHeldSessionSetRedacted(t, err, ErrHeldSessionSetHealth, ErrHeldSessionPrematureClose)
require.NoError(t, set.Close(context.Background()))
}
func TestHeldSessionSetPublicRedactionProtocolHealth(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
fixture.input.testHealthSnapshotOverrideHook = func(index int, request heldHealthSnapshotRequest) *heldHealthSnapshot {
if index != 0 {
return nil
}
return &heldHealthSnapshot{epoch: request.epoch - 1, index: 0, event: &heldHealthEvent{err: sensitiveBoundaryError("protocol")}}
}
set := fixture.returnedSet(t)
fixture.coordinator.sessions[11].prematureClose()
err := set.WaitHealthy(context.Background())
requireHeldSessionSetRedacted(t, err, ErrHeldSessionSetHealth, ErrHeldSessionSetHealthProtocol)
require.NoError(t, set.Close(context.Background()))
}
func sensitiveBoundaryError(label string) error {
return errors.New(label + " token=secret-token path=/secret/path uuid=secret-uuid stream=secret-stream server=991 authorization=secret-token")
}
func sessionIndexByKind(fixture *heldSessionSetPublicFixture, kind StressSessionKind) int {
for index, session := range fixture.plan.Sessions {
if session.Kind == kind {
return index
}
}
panic("session kind not found")
}
func newHeldSessionSetWithReleasedConstructors(fixture *heldSessionSetPublicFixture) (*heldSessionSet, error) {
result := make(chan *heldSessionSet, 1)
errResult := make(chan error, 1)
go func() {
set, err := NewHeldSessionSet(context.Background(), fixture.input)
result <- set
errResult <- err
}()
for range fixture.plan.Sessions {
<-fixture.coordinator.ready
}
fixture.releaseConstructors()
return <-result, <-errResult
}
@@ -0,0 +1,66 @@
//go:build linux
package scenario
import (
"context"
"errors"
"testing"
"github.com/stretchr/testify/require"
)
func TestHeldSessionSetPublicPresentStateRetainsEveryDistinctSentinel(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
expected := make(map[string]error, len(fixture.plan.Sessions))
for index := range fixture.plan.Sessions {
sentinel := errors.New("present-state-" + fixture.coordinator.sessions[index].streamID)
expected[fixture.coordinator.sessions[index].streamID] = sentinel
fixture.state.presentStreamErrors[fixture.coordinator.sessions[index].streamID] = sentinel
}
aggregateSentinel := errors.New("present-state-aggregate")
fixture.state.presentAggregateError = aggregateSentinel
result := make(chan error, 1)
go func() {
_, err := NewHeldSessionSet(context.Background(), fixture.input)
result <- err
}()
for range fixture.plan.Sessions {
<-fixture.coordinator.ready
}
fixture.releaseConstructors()
err := <-result
require.Error(t, err)
for streamID, sentinel := range expected {
require.ErrorIs(t, err, sentinel, streamID)
}
require.ErrorIs(t, err, aggregateSentinel)
require.ElementsMatch(t, canonicalStreamIDs(fixture), fixture.state.present)
require.Equal(t, 12, len(fixture.state.present))
require.Equal(t, 26, fixture.counters.values().waitState)
require.Equal(t, 1, fixture.state.presentAggregateCalls)
require.True(t, fixture.state.aggregatePredicatesEmpty)
}
func TestHeldSessionSetPublicAbsentStateRetainsEveryDistinctSentinel(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
set := fixture.returnedSet(t)
expected := make(map[string]error, len(fixture.plan.Sessions))
for index := range fixture.plan.Sessions {
sentinel := errors.New("absent-state-" + fixture.coordinator.sessions[index].streamID)
expected[fixture.coordinator.sessions[index].streamID] = sentinel
fixture.state.absentStreamErrors[fixture.coordinator.sessions[index].streamID] = sentinel
}
aggregateSentinel := errors.New("absent-state-aggregate")
fixture.state.absentAggregateError = aggregateSentinel
err := set.Close(context.Background())
require.Error(t, err)
for streamID, sentinel := range expected {
require.ErrorIs(t, err, sentinel, streamID)
}
require.ErrorIs(t, err, aggregateSentinel)
require.ElementsMatch(t, canonicalStreamIDs(fixture), fixture.state.absent)
require.Equal(t, 26, fixture.counters.values().waitState)
require.Equal(t, 1, fixture.state.absentAggregateCalls)
require.True(t, fixture.state.aggregatePredicatesEmpty)
}
@@ -0,0 +1,124 @@
//go:build linux
package scenario
import (
"context"
"testing"
"github.com/stretchr/testify/require"
)
func TestHeldSessionSetPublicSuccessOrchestratesCanonicalLifecycle(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
result := make(chan *heldSessionSet, 1)
errResult := make(chan error, 1)
go func() {
set, err := NewHeldSessionSet(context.Background(), fixture.input)
result <- set
errResult <- err
}()
for range fixture.plan.Sessions {
<-fixture.coordinator.ready
}
permutation := []int{8, 1, 11, 3, 0, 10, 4, 7, 2, 9, 5, 6}
fixture.coordinator.releaseAll()
started := make([]string, 0, len(permutation))
completed := make([]heldSessionConstructorCompletion, 0, len(permutation))
acquired := make([]int, 0, len(permutation))
for _, index := range permutation {
fixture.coordinator.release(index)
started = append(started, <-fixture.coordinator.startEvents)
completion := <-fixture.coordinator.completed
completed = append(completed, completion)
require.Equal(t, index, completion.index)
require.NoError(t, completion.err)
require.True(t, completion.acquired)
acquired = append(acquired, <-fixture.coordinator.acquired)
require.Equal(t, index, acquired[len(acquired)-1])
}
set := <-result
require.NoError(t, <-errResult)
for range fixture.plan.Sessions {
constructorContext := <-fixture.coordinator.contextSeen
select {
case <-constructorContext.Done():
t.Fatal("constructor context was canceled after successful construction")
default:
}
}
counters := fixture.counters.values()
require.Equal(t, 8, counters.inspect)
require.Equal(t, 1, counters.snapshot)
require.Equal(t, 1, counters.observe)
require.Equal(t, 4, counters.terminal)
require.Equal(t, 4, counters.nat)
require.Equal(t, 4, counters.fm)
require.Len(t, set.sessions, 12)
for index, session := range set.sessions {
require.Equal(t, fixture.plan.Sessions[index], session.Plan())
}
require.Equal(t, permutation, acquired)
require.Equal(t, permutation, completedIndexes(completed))
require.Equal(t, canonicalPlanIndexes(set), []int{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11})
require.Equal(t, permutationPlanIDs(fixture, permutation), started)
require.NoError(t, set.Close(context.Background()))
require.NoError(t, set.WaitHealthy(context.Background()))
for index := 11; index >= 0; index-- {
require.Equal(t, index, <-fixture.closeOrder)
require.Equal(t, index, <-fixture.waitOrder)
}
require.Len(t, fixture.state.present, 12)
require.Len(t, fixture.state.absent, 12)
require.ElementsMatch(t, canonicalStreamIDs(fixture), fixture.state.present)
require.ElementsMatch(t, canonicalStreamIDs(fixture), fixture.state.absent)
require.Equal(t, 26, len(fixture.state.counts))
require.Equal(t, 13, countExpectedState(fixture.state.counts, 19))
require.Equal(t, 13, countExpectedState(fixture.state.counts, 7))
require.Equal(t, 1, fixture.state.presentAggregateCalls)
require.Equal(t, 1, fixture.state.absentAggregateCalls)
require.True(t, fixture.state.aggregatePredicatesEmpty)
set.healthWG.Wait()
}
func completedIndexes(completions []heldSessionConstructorCompletion) []int {
indexes := make([]int, 0, len(completions))
for _, completion := range completions {
indexes = append(indexes, completion.index)
}
return indexes
}
func canonicalPlanIndexes(set *heldSessionSet) []int {
indexes := make([]int, 0, len(set.sessions))
for index := range set.sessions {
indexes = append(indexes, index)
}
return indexes
}
func permutationPlanIDs(fixture *heldSessionSetPublicFixture, indexes []int) []string {
ids := make([]string, 0, len(indexes))
for _, index := range indexes {
ids = append(ids, fixture.plan.Sessions[index].ID.String())
}
return ids
}
func canonicalStreamIDs(fixture *heldSessionSetPublicFixture) []string {
ids := make([]string, 0, len(fixture.plan.Sessions))
for index := range fixture.plan.Sessions {
ids = append(ids, fixture.coordinator.sessions[index].streamID)
}
return ids
}
func countExpectedState(counts []int, expected int) int {
count := 0
for _, value := range counts {
if value == expected {
count++
}
}
return count
}
@@ -0,0 +1,116 @@
//go:build linux
package scenario
import (
"context"
"testing"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/stretchr/testify/require"
)
func TestHeldSessionSetPublicTopologyRejectsAuthorityMutationsWithoutSideEffects(t *testing.T) {
cases := []struct {
name string
mutate func(*heldSessionSetPublicFixture)
expectedInspect int
expectedError error
}{
{name: "wrong profile", mutate: func(fixture *heldSessionSetPublicFixture) { fixture.input.Plan.Profile = "missing-profile" }, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetPlan},
{name: "missing ordinal", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.Topology = fixture.input.Topology[:7]
}, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "duplicate ordinal", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.Topology[1].Ordinal = fixture.input.Topology[0].Ordinal
}, expectedInspect: 1, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "duplicate UUID", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.Topology[1].Readiness.UUID = fixture.input.Topology[0].Readiness.UUID
}, expectedInspect: 2, expectedError: ErrHeldReadinessAgentMismatch},
{name: "duplicate server ID", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.Topology[1].Readiness.ServerID = fixture.input.Topology[0].Readiness.ServerID
}, expectedInspect: 2, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "extra control server", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.ControlServerIDs = append(fixture.input.ControlServerIDs, 99)
}, expectedInspect: 8, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "duplicate control server", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.ControlServerIDs[1] = fixture.input.ControlServerIDs[0]
}, expectedInspect: 8, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "noncanonical plan", mutate: func(fixture *heldSessionSetPublicFixture) { fixture.input.Plan.Sessions[0].Ordinal++ }, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetPlan},
{name: "duplicate plan", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.Plan.Sessions[1].ID = fixture.input.Plan.Sessions[0].ID
}, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetPlan},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
testCase.mutate(fixture)
_, err := NewHeldSessionSet(context.Background(), fixture.input)
require.Error(t, err)
require.ErrorIs(t, err, ErrHeldSessionSetOperation)
require.ErrorIs(t, err, testCase.expectedError)
require.Empty(t, fixture.state.present)
require.Empty(t, fixture.state.absent)
require.Empty(t, fixture.coordinator.ready)
counters := fixture.counters.values()
require.Zero(t, counters.snapshot)
require.Zero(t, counters.observe)
require.Equal(t, testCase.expectedInspect, counters.inspect)
require.Zero(t, counters.terminal)
require.Zero(t, counters.nat)
require.Zero(t, counters.fm)
})
}
}
func TestHeldSessionSetPublicTopologyRejectsBoundaryInputsWithoutConstruction(t *testing.T) {
cases := []struct {
name string
mutate func(*heldSessionSetPublicFixture)
expectedInspect int
expectedError error
}{
{name: "nil dashboard", mutate: func(fixture *heldSessionSetPublicFixture) { fixture.input.Dashboard = nil }, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "nil control client", mutate: func(fixture *heldSessionSetPublicFixture) { fixture.input.ControlClient = nil }, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "nil agent", mutate: func(fixture *heldSessionSetPublicFixture) { fixture.input.Topology[0].Agent = nil }, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "nil pat", mutate: func(fixture *heldSessionSetPublicFixture) { fixture.input.Topology[0].PATClient = nil }, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "invalid pid", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.Dependencies.InspectAgent = func(*agent.Agent) heldSessionAgentFacts { return heldSessionAgentFacts{PID: 0, UUID: "uuid-a"} }
}, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "readiness mismatch", mutate: func(fixture *heldSessionSetPublicFixture) { fixture.input.Topology[0].Readiness.UUID = "different" }, expectedInspect: 1, expectedError: ErrInvalidHeldReadiness},
{name: "incomplete readiness", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.Topology[0].Readiness.VersionObserved = false
}, expectedInspect: 1, expectedError: ErrInvalidHeldReadiness},
{name: "zero control server", mutate: func(fixture *heldSessionSetPublicFixture) { fixture.input.ControlServerIDs[0] = 0 }, expectedInspect: 8, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "unknown control server", mutate: func(fixture *heldSessionSetPublicFixture) { fixture.input.ControlServerIDs[0] = 991 }, expectedInspect: 8, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "missing control server", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.ControlServerIDs = fixture.input.ControlServerIDs[:7]
}, expectedInspect: 8, expectedError: ErrInvalidHeldSessionSetTopology},
{name: "plan kind", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.Plan.Sessions[0].Kind = StressSessionKind("unknown")
}, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetPlan},
{name: "plan agent", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.Plan.Sessions[0].Agent = fixture.input.Plan.Sessions[1].Agent
}, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetPlan},
{name: "plan ordinal", mutate: func(fixture *heldSessionSetPublicFixture) { fixture.input.Plan.Sessions[0].Ordinal++ }, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetPlan},
{name: "plan duplicate id", mutate: func(fixture *heldSessionSetPublicFixture) {
fixture.input.Plan.Sessions[1].ID = fixture.input.Plan.Sessions[0].ID
}, expectedInspect: 0, expectedError: ErrInvalidHeldSessionSetPlan},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
fixture := newHeldSessionSetPublicFixture(t)
testCase.mutate(fixture)
_, err := NewHeldSessionSet(context.Background(), fixture.input)
require.Error(t, err)
require.ErrorIs(t, err, testCase.expectedError)
counters := fixture.counters.values()
require.Equal(t, testCase.expectedInspect, counters.inspect)
require.Zero(t, counters.snapshot)
require.Zero(t, counters.observe)
require.Zero(t, counters.terminal)
require.Zero(t, counters.nat)
require.Zero(t, counters.fm)
})
}
}
@@ -0,0 +1,168 @@
//go:build linux && agentcompat
package scenario
import (
"context"
"errors"
"os"
"syscall"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
)
func TestHeldSessionSetEightAgentFourFourFour(t *testing.T) {
requireHeldRealSources(t)
paths, err := contract.NewPaths(os.Getenv("AGENTCOMPAT_NEZHA_SOURCE"), os.Getenv("AGENTCOMPAT_AGENT_SOURCE"), t.TempDir())
require.NoError(t, err)
ctx, cancel := context.WithTimeout(t.Context(), 30*time.Minute)
defer cancel()
profile, err := contract.ProfileByName(string(contract.ProfilePRFull))
require.NoError(t, err)
plan, err := GenerateStressPlan(profile, contract.DefaultSeed)
require.NoError(t, err)
realFixture, err := startHeldSessionSetRealFixture(ctx, paths, plan)
require.NoError(t, err)
dashboardIdentity := realFixture.dashboard.RuntimeIdentity()
agentIdentities := make([]agent.ProcessIdentity, len(realFixture.agents))
workspaceRoots := make([]string, 0, len(realFixture.agents)+2)
workspaceRoots = append(workspaceRoots, realFixture.dashboard.WorkspaceRoot(), realFixture.preparedBinary.WorkspaceRoot())
for index, instance := range realFixture.agents {
agentIdentities[index] = instance.RuntimeIdentity()
workspaceRoots = append(workspaceRoots, instance.WorkspaceRoot())
}
t.Cleanup(func() { _ = realFixture.close(context.Background(), nil) })
input, err := realFixture.input(plan)
require.NoError(t, err)
baseline, err := realFixture.controlPAT.Client.IOStreamState(ctx)
require.NoError(t, err)
set, err := NewHeldSessionSet(ctx, input)
require.NoError(t, err)
require.Len(t, set.sessions, 12)
for index, session := range set.sessions {
require.Equal(t, plan.Sessions[index], session.Plan())
}
select {
case <-set.Done():
t.Fatal("held session set health completed while sessions were live")
default:
}
streamIDs := make([]string, len(set.sessions))
seen := make(map[string]struct{}, len(set.sessions))
protocolProved := true
for index, session := range set.sessions {
streamID, present := session.IOStreamID()
require.True(t, present)
require.NotEmpty(t, streamID)
require.NotContains(t, seen, streamID)
seen[streamID] = struct{}{}
streamIDs[index] = streamID
protocolProved = protocolProved && heldRealSessionProtocolProved(session)
}
require.True(t, protocolProved)
live, err := realFixture.controlPAT.Client.WaitForIOStreamState(ctx, client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(baseline.Count + 12)})
require.NoError(t, err)
require.Equal(t, baseline.Count+12, live.Count)
for _, streamID := range streamIDs {
_, err := realFixture.controlPAT.Client.WaitForIOStreamState(ctx, client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(baseline.Count + 12), PresentStreamID: streamID})
require.NoError(t, err)
}
require.Equal(t, dashboardIdentity, realFixture.dashboard.RuntimeIdentity())
for index, instance := range realFixture.agents {
require.Equal(t, agentIdentities[index], instance.RuntimeIdentity())
}
require.Equal(t, 4, countHeldSessionKind(plan.Sessions, StressSessionTerminal))
require.Equal(t, 4, countHeldSessionKind(plan.Sessions, StressSessionNAT))
require.Equal(t, 4, countHeldSessionKind(plan.Sessions, StressSessionFM))
require.NoError(t, set.Close(ctx))
require.NoError(t, set.WaitHealthy(ctx))
closed, err := realFixture.controlPAT.Client.WaitForIOStreamState(ctx, client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(baseline.Count)})
require.NoError(t, err)
require.Equal(t, baseline.Count, closed.Count)
for _, streamID := range streamIDs {
_, err := realFixture.controlPAT.Client.WaitForIOStreamState(ctx, client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(baseline.Count), AbsentStreamID: streamID})
require.NoError(t, err)
}
require.Equal(t, dashboardIdentity, realFixture.dashboard.RuntimeIdentity())
for index, instance := range realFixture.agents {
require.Equal(t, agentIdentities[index], instance.RuntimeIdentity())
}
resourcesAbsent := heldRealSessionResourcesAbsent(ctx, realFixture, set.sessions)
require.True(t, resourcesAbsent)
cleanupErr := realFixture.close(ctx, nil)
require.NoError(t, cleanupErr)
cleanupOK := realFixture.dashboard.CleanupReceipt().Passed && !realFixture.dashboard.CleanupReceipt().Forced
for _, instance := range realFixture.agents {
cleanupOK = cleanupOK && instance.CleanupReceipt().Passed && !instance.CleanupReceipt().Forced
}
processesClean := heldRealPIDGone(dashboardIdentity.PID) && heldRealGroupGone(dashboardIdentity.ProcessGroupID)
for _, identity := range agentIdentities {
processesClean = processesClean && heldRealPIDGone(identity.PID) && heldRealGroupGone(identity.ProcessGroupID)
}
workspacesClean := true
for _, root := range workspaceRoots {
workspacesClean = workspacesClean && heldSessionSetRealWorkspaceGone(root)
}
evidenceValue := heldSessionSetRealEvidence{Version: 1, Profile: string(plan.Profile), Seed: "4e5a4841", BaselineCount: baseline.Count, LiveCount: live.Count, ClosedCount: closed.Count, TerminalCount: 4, NATCount: 4, FMCount: 4, AgentOrdinals: []int{1, 2, 3, 4, 5, 6, 7, 8}, ProtocolProved: protocolProved, ExactIDsPresent: true, ExactIDsAbsent: true, PIDStable: true, ResourcesAbsent: resourcesAbsent, ProcessesClean: processesClean, WorkspacesClean: workspacesClean, CleanupOK: cleanupOK}
for index, instance := range realFixture.agents {
evidenceValue.AgentSummaries = append(evidenceValue.AgentSummaries, heldSessionSetRealAgentSummary{Ordinal: index + 1, ServerDigest: heldRealDigest(string(rune(realFixture.readiness[index].ServerID))), PATIdentity: realFixture.agentPATs[index].IdentitySeen, PATScopeExact: len(realFixture.agentPATs[index].ServerIDs) == 1 && realFixture.agentPATs[index].ServerIDs[0] == realFixture.readiness[index].ServerID})
_ = instance
}
for index, session := range set.sessions {
evidenceValue.SessionDigests = append(evidenceValue.SessionDigests, heldRealDigest(streamIDs[index]))
evidenceValue.SessionSummaries = append(evidenceValue.SessionSummaries, heldSessionSetRealSessionSummary{Ordinal: index + 1, Kind: string(session.Plan().Kind), AgentOrdinal: session.Plan().Agent.Int(), StreamDigest: heldRealDigest(streamIDs[index]), Present: true, Absent: true, Protocol: heldRealSessionProtocolProved(session)})
}
require.True(t, cleanupOK && processesClean && workspacesClean)
require.NoError(t, writeHeldSessionSetRealEvidence("/tmp/nezha-held-real-sessions", evidenceValue))
_, err = readHeldSessionSetRealEvidence("/tmp/nezha-held-real-sessions")
require.NoError(t, err)
}
func heldRealSessionProtocolProved(session heldSession) bool {
switch concrete := session.(type) {
case *heldTerminalSession:
return concrete.ProtocolProved()
case *heldNATSession:
return concrete.ProtocolProved()
case *heldLegacyFMSession:
return concrete.ProtocolProved()
default:
return false
}
}
func heldRealSessionResourcesAbsent(ctx context.Context, fixture *heldSessionSetRealFixture, sessions []heldSession) bool {
for _, session := range sessions {
switch concrete := session.(type) {
case *heldNATSession:
present, err := heldRealNATProfilePresent(ctx, fixture.dashboard.Clients().REST, concrete.profileID)
if err != nil || present {
return false
}
case *heldLegacyFMSession:
rootName := heldLegacyFMRootName.ReplaceAllString(session.Plan().ID.String(), "-")
if !heldSessionSetRealWorkspaceGone("" + fixture.agents[session.Plan().Agent.Int()-1].WorkspaceRoot() + "/held-fm-" + rootName) {
return false
}
case *heldTerminalSession:
_ = concrete
}
}
return true
}
func heldRealGroupGone(pgid int) bool {
if pgid < 1 {
return false
}
err := syscall.Kill(-pgid, 0)
return errors.Is(err, syscall.ESRCH)
}
@@ -0,0 +1,197 @@
//go:build linux && agentcompat
package scenario
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"slices"
"strings"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/evidence"
)
var ErrHeldSessionSetRealEvidenceInvalid = errors.New("held session set evidence is invalid")
type heldSessionSetRealEvidence struct {
Version int `json:"version"`
Profile string `json:"profile"`
Seed string `json:"seed"`
BaselineCount int `json:"baseline_count"`
LiveCount int `json:"live_count"`
ClosedCount int `json:"closed_count"`
TerminalCount int `json:"terminal_count"`
NATCount int `json:"nat_count"`
FMCount int `json:"fm_count"`
AgentOrdinals []int `json:"agent_ordinals"`
AgentSummaries []heldSessionSetRealAgentSummary `json:"agent_summaries"`
SessionSummaries []heldSessionSetRealSessionSummary `json:"session_summaries"`
SessionDigests []string `json:"session_digests"`
ProtocolProved bool `json:"protocol_proved"`
ExactIDsPresent bool `json:"exact_ids_present"`
ExactIDsAbsent bool `json:"exact_ids_absent"`
PIDStable bool `json:"pid_stable"`
ResourcesAbsent bool `json:"resources_absent"`
ProcessesClean bool `json:"processes_clean"`
WorkspacesClean bool `json:"workspaces_clean"`
CleanupOK bool `json:"cleanup_ok"`
}
type heldSessionSetRealAgentSummary struct {
Ordinal int `json:"ordinal"`
ServerDigest string `json:"server_digest"`
PATIdentity bool `json:"pat_identity"`
PATScopeExact bool `json:"pat_scope_exact"`
}
type heldSessionSetRealSessionSummary struct {
Ordinal int `json:"ordinal"`
Kind string `json:"kind"`
AgentOrdinal int `json:"agent_ordinal"`
StreamDigest string `json:"stream_digest"`
Present bool `json:"present"`
Absent bool `json:"absent"`
Protocol bool `json:"protocol"`
}
func validateHeldSessionSetRealEvidence(evidenceValue heldSessionSetRealEvidence) error {
if evidenceValue.Version != 1 {
return ErrHeldSessionSetRealEvidenceInvalid
}
if evidenceValue.Profile != string(contract.ProfilePRFull) || evidenceValue.Seed != "4e5a4841" || evidenceValue.BaselineCount < 0 || evidenceValue.LiveCount != evidenceValue.BaselineCount+12 || evidenceValue.ClosedCount != evidenceValue.BaselineCount {
return ErrHeldSessionSetRealEvidenceInvalid
}
if evidenceValue.TerminalCount != 4 || evidenceValue.NATCount != 4 || evidenceValue.FMCount != 4 || !slices.Equal(evidenceValue.AgentOrdinals, []int{1, 2, 3, 4, 5, 6, 7, 8}) || len(evidenceValue.AgentSummaries) != 8 || len(evidenceValue.SessionSummaries) != 12 || len(evidenceValue.SessionDigests) != 12 {
return ErrHeldSessionSetRealEvidenceInvalid
}
serverDigests := make(map[string]struct{}, len(evidenceValue.AgentSummaries))
for index, summary := range evidenceValue.AgentSummaries {
if summary.Ordinal != index+1 || !validHeldSessionSetRealDigest(summary.ServerDigest) || !summary.PATIdentity || !summary.PATScopeExact {
return ErrHeldSessionSetRealEvidenceInvalid
}
if _, exists := serverDigests[summary.ServerDigest]; exists {
return ErrHeldSessionSetRealEvidenceInvalid
}
serverDigests[summary.ServerDigest] = struct{}{}
}
profile, err := contract.ProfileByName(string(contract.ProfilePRFull))
if err != nil {
return ErrHeldSessionSetRealEvidenceInvalid
}
plan, err := GenerateStressPlan(profile, contract.DefaultSeed)
if err != nil || len(plan.Sessions) != 12 {
return ErrHeldSessionSetRealEvidenceInvalid
}
sessionDigests := make(map[string]struct{}, len(evidenceValue.SessionDigests))
for index, digest := range evidenceValue.SessionDigests {
if !validHeldSessionSetRealDigest(digest) {
return ErrHeldSessionSetRealEvidenceInvalid
}
if _, exists := sessionDigests[digest]; exists {
return ErrHeldSessionSetRealEvidenceInvalid
}
sessionDigests[digest] = struct{}{}
if digest != evidenceValue.SessionSummaries[index].StreamDigest {
return ErrHeldSessionSetRealEvidenceInvalid
}
}
kindCounts := map[StressSessionKind]int{}
for index, summary := range evidenceValue.SessionSummaries {
canonical := plan.Sessions[index]
if summary.Ordinal != index+1 || summary.Kind != string(canonical.Kind) || summary.AgentOrdinal != canonical.Agent.Int() || !validHeldSessionSetRealDigest(summary.StreamDigest) || !summary.Present || !summary.Absent || !summary.Protocol {
return ErrHeldSessionSetRealEvidenceInvalid
}
kindCounts[canonical.Kind]++
}
if kindCounts[StressSessionTerminal] != 4 || kindCounts[StressSessionNAT] != 4 || kindCounts[StressSessionFM] != 4 {
return ErrHeldSessionSetRealEvidenceInvalid
}
if !evidenceValue.ProtocolProved || !evidenceValue.ExactIDsPresent || !evidenceValue.ExactIDsAbsent || !evidenceValue.PIDStable || !evidenceValue.ResourcesAbsent || !evidenceValue.ProcessesClean || !evidenceValue.WorkspacesClean || !evidenceValue.CleanupOK {
return ErrHeldSessionSetRealEvidenceInvalid
}
return nil
}
func validHeldSessionSetRealDigest(value string) bool {
if len(value) != sha256.Size*2 || value != strings.ToLower(value) {
return false
}
_, err := hex.DecodeString(value)
return err == nil
}
func writeHeldSessionSetRealEvidence(root string, evidenceValue heldSessionSetRealEvidence) error {
if err := validateHeldSessionSetRealEvidence(evidenceValue); err != nil {
return err
}
info, err := os.Lstat(root)
if err != nil {
if !errors.Is(err, os.ErrNotExist) {
return err
}
if err := os.Mkdir(root, 0o700); err != nil {
return err
}
info, err = os.Lstat(root)
}
if err != nil || info.Mode()&os.ModeSymlink != 0 || !info.IsDir() || info.Mode().Perm() != 0o700 {
return errors.New("held session set evidence root is not a private directory")
}
path := filepath.Join(root, "held-session-set.json")
if stale, err := os.Lstat(path); err == nil {
if stale.Mode()&os.ModeSymlink != 0 || !stale.Mode().IsRegular() {
return errors.New("stale held session set evidence is not a regular file")
}
if err := os.Remove(path); err != nil {
return err
}
} else if !errors.Is(err, os.ErrNotExist) {
return err
}
data, err := json.Marshal(evidenceValue)
if err != nil {
return err
}
if evidence.Redact(string(data)) != string(data) {
return errors.New("held session set evidence requires redaction")
}
temporary, err := os.CreateTemp(root, ".held-session-set-*")
if err != nil {
return err
}
temporaryName := temporary.Name()
defer os.Remove(temporaryName)
if err := temporary.Chmod(0o600); err != nil {
_ = temporary.Close()
return err
}
if _, err := temporary.Write(data); err != nil {
_ = temporary.Close()
return err
}
if err := temporary.Close(); err != nil {
return err
}
return os.Rename(temporaryName, path)
}
func readHeldSessionSetRealEvidence(root string) (heldSessionSetRealEvidence, error) {
var result heldSessionSetRealEvidence
data, err := os.ReadFile(filepath.Join(root, "held-session-set.json"))
if err != nil {
return result, fmt.Errorf("read held session set evidence: %w", err)
}
if evidence.Redact(string(data)) != string(data) {
return result, errors.New("held session set evidence is not redacted")
}
if err := json.Unmarshal(data, &result); err != nil {
return result, err
}
return result, validateHeldSessionSetRealEvidence(result)
}
@@ -0,0 +1,142 @@
//go:build linux && agentcompat
package scenario
import (
"context"
"errors"
"fmt"
"os"
"slices"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
)
type heldSessionSetRealFixture struct {
dashboard *dashboard.Dashboard
preparedBinary *agent.PreparedBinary
agents []*agent.Agent
readiness []agent.Readiness
agentPATs []heldRealPATIdentity
plan StressPlan
controlPAT heldRealPATIdentity
controlServerIDs []uint64
closed bool
}
func startHeldSessionSetRealFixture(ctx context.Context, paths contract.Paths, plan StressPlan) (*heldSessionSetRealFixture, error) {
dashboardInstance, err := dashboard.Start(ctx, dashboard.StartConfig{SourceDir: paths.NezhaSource().String(), ReceiptGate: true})
if err != nil {
return nil, err
}
fixture := &heldSessionSetRealFixture{dashboard: dashboardInstance, plan: plan}
prepared, err := agent.PrepareBinary(ctx, paths.AgentSource().String())
if err != nil {
return nil, fixture.close(ctx, err)
}
fixture.preparedBinary = prepared
for ordinal := 1; ordinal <= heldSessionSetAgentCount; ordinal++ {
uuid := fmt.Sprintf("00000000-0000-0000-0000-%012d", 700+ordinal)
instance, startErr := agent.Start(ctx, agent.AgentStartConfig{PreparedBinary: prepared, Endpoint: dashboardInstance.Endpoint(), Secret: dashboardInstance.AgentSecret(), UUID: uuid})
if startErr != nil {
return nil, fixture.close(ctx, startErr)
}
fixture.agents = append(fixture.agents, instance)
}
if err := dashboardInstance.WaitForReceiptAccepted(ctx); err != nil {
return nil, fixture.close(ctx, err)
}
if err := dashboardInstance.ReleaseReceipt(ctx); err != nil {
return nil, fixture.close(ctx, err)
}
for index, instance := range fixture.agents {
serverID, infoErr := dashboardInstance.WaitForInfo2UUID(ctx, instance.UUID())
if infoErr != nil {
return nil, fixture.close(ctx, infoErr)
}
pat, patErr := mintHeldRealPAT(ctx, dashboardInstance, fmt.Sprintf("held-set-agent-%d", index+1), []uint64{serverID})
if patErr != nil {
return nil, fixture.close(ctx, patErr)
}
ready, readyErr := instance.WaitReadyEventDrivenWithClient(ctx, dashboardInstance, pat.Client)
if readyErr != nil {
return nil, fixture.close(ctx, readyErr)
}
fixture.readiness = append(fixture.readiness, ready)
fixture.agentPATs = append(fixture.agentPATs, pat)
}
for _, readiness := range fixture.readiness {
fixture.controlServerIDs = append(fixture.controlServerIDs, readiness.ServerID)
}
fixture.controlPAT, err = mintHeldRealPAT(ctx, dashboardInstance, "held-set-control", fixture.controlServerIDs)
if err != nil {
return nil, fixture.close(ctx, err)
}
if err := validateHeldSessionSetRealFixture(fixture, plan); err != nil {
return nil, fixture.close(ctx, err)
}
return fixture, nil
}
func validateHeldSessionSetRealFixture(fixture *heldSessionSetRealFixture, plan StressPlan) error {
if fixture == nil || fixture.dashboard == nil || fixture.preparedBinary == nil || len(fixture.agents) != heldSessionSetAgentCount || len(fixture.readiness) != heldSessionSetAgentCount || len(fixture.agentPATs) != heldSessionSetAgentCount || len(fixture.controlServerIDs) != heldSessionSetAgentCount {
return errors.New("held session set real fixture is incomplete")
}
if plan.Profile != contract.ProfilePRFull || len(plan.Sessions) != 12 {
return errors.New("held session set real fixture received noncanonical plan")
}
seenServers := make(map[uint64]struct{}, len(fixture.controlServerIDs))
for index, readiness := range fixture.readiness {
if readiness.ServerID == 0 || readiness.UUID != fixture.agents[index].UUID() || !fixture.agentPATs[index].IdentitySeen || !slices.Equal(fixture.agentPATs[index].ServerIDs, []uint64{readiness.ServerID}) {
return errors.New("held session set real fixture PAT mapping is invalid")
}
if _, exists := seenServers[readiness.ServerID]; exists {
return errors.New("held session set real fixture server IDs are not unique")
}
seenServers[readiness.ServerID] = struct{}{}
}
if !fixture.controlPAT.IdentitySeen || !slices.Equal(fixture.controlPAT.ServerIDs, fixture.controlServerIDs) {
return errors.New("held session set real fixture control PAT mapping is invalid")
}
return nil
}
func (fixture *heldSessionSetRealFixture) input(plan StressPlan) (HeldSessionSetInput, error) {
topology := make([]HeldSessionAgent, len(fixture.agents))
for index, instance := range fixture.agents {
ordinal, err := NewStressAgentOrdinal(index + 1)
if err != nil {
return HeldSessionSetInput{}, err
}
topology[index] = HeldSessionAgent{Ordinal: ordinal, Agent: instance, Readiness: fixture.readiness[index], PATClient: fixture.agentPATs[index].Client}
}
return HeldSessionSetInput{Dashboard: fixture.dashboard, Plan: plan, Topology: topology, ControlClient: fixture.controlPAT.Client, ControlServerIDs: append([]uint64(nil), fixture.controlServerIDs...), Dependencies: defaultHeldSessionSetDependencies()}, nil
}
func (fixture *heldSessionSetRealFixture) close(ctx context.Context, cause error) error {
if fixture == nil || fixture.closed {
return cause
}
fixture.closed = true
cleanupContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), 90*time.Second)
defer cancel()
joined := cause
for index := len(fixture.agents) - 1; index >= 0; index-- {
joined = errors.Join(joined, fixture.agents[index].Stop(cleanupContext))
}
if fixture.preparedBinary != nil {
joined = errors.Join(joined, fixture.preparedBinary.Close())
}
if fixture.dashboard != nil {
joined = errors.Join(joined, fixture.dashboard.Stop(cleanupContext))
}
return joined
}
func heldSessionSetRealWorkspaceGone(path string) bool {
_, err := os.Stat(path)
return errors.Is(err, os.ErrNotExist)
}
@@ -0,0 +1,69 @@
//go:build linux && agentcompat
package scenario
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"net/http"
"slices"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
)
func realHeldSessionSetAgentOrdinals(plan StressPlan) []int {
ordinals := make([]int, 0, heldSessionSetAgentCount)
for _, session := range plan.Sessions {
if !slices.Contains(ordinals, session.Agent.Int()) {
ordinals = append(ordinals, session.Agent.Int())
}
}
for ordinal := 1; ordinal <= heldSessionSetAgentCount; ordinal++ {
if !slices.Contains(ordinals, ordinal) {
ordinals = append(ordinals, ordinal)
}
}
return ordinals
}
type heldRealPATIdentity struct {
Client *client.Client
TokenID uint64
ServerIDs []uint64
IdentitySeen bool
}
func mintHeldRealPAT(ctx context.Context, dashboardInstance *dashboard.Dashboard, name string, serverIDs []uint64) (heldRealPATIdentity, error) {
pat, err := createTerminalPAT(ctx, dashboardInstance, name, serverIDs)
if err != nil {
return heldRealPATIdentity{}, err
}
identity, err := client.CallTool[struct{}, client.WhoAmIResult](ctx, pat, client.ToolCall[struct{}]{Name: "meta.whoami", Arguments: struct{}{}})
if err != nil {
return heldRealPATIdentity{}, err
}
whoami := identity.StructuredContent
if whoami.TokenID == 0 || whoami.TokenName != name || !slices.Equal(whoami.Scopes, []string{"nezha:*"}) || !slices.Equal(whoami.ServerIDs, serverIDs) {
return heldRealPATIdentity{}, errors.New("PAT identity or server allowlist mismatch")
}
return heldRealPATIdentity{Client: pat, TokenID: whoami.TokenID, ServerIDs: append([]uint64(nil), serverIDs...), IdentitySeen: true}, nil
}
func createTerminalPAT(ctx context.Context, dashboardInstance *dashboard.Dashboard, name string, serverIDs []uint64) (*client.Client, error) {
pat, err := client.DoREST[terminalPATRequest, terminalPATResponse](ctx, dashboardInstance.Clients().REST, client.RESTRequest[terminalPATRequest]{Method: http.MethodPost, Path: "/api/v1/api-tokens", Body: &terminalPATRequest{Name: name, Scopes: []string{"nezha:*"}, ServerIDs: serverIDs}})
if err != nil {
return nil, err
}
if pat.Token == "" {
return nil, errors.New("PAT response omitted token")
}
return dashboardInstance.AuthenticatedClient(pat.Token)
}
func heldRealDigest(value string) string {
digest := sha256.Sum256([]byte(value))
return hex.EncodeToString(digest[:])
}
@@ -0,0 +1,97 @@
//go:build linux && agentcompat
package scenario
import (
"os"
"slices"
"strings"
"testing"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
)
func validHeldSessionSetRealEvidenceFromPlan() heldSessionSetRealEvidence {
profile, _ := contract.ProfileByName(string(contract.ProfilePRFull))
plan, _ := GenerateStressPlan(profile, contract.DefaultSeed)
value := heldSessionSetRealEvidence{Version: 1, Profile: string(contract.ProfilePRFull), Seed: "4e5a4841", BaselineCount: 1, LiveCount: 13, ClosedCount: 1, TerminalCount: 4, NATCount: 4, FMCount: 4, AgentOrdinals: []int{1, 2, 3, 4, 5, 6, 7, 8}, ProtocolProved: true, ExactIDsPresent: true, ExactIDsAbsent: true, PIDStable: true, ResourcesAbsent: true, ProcessesClean: true, WorkspacesClean: true, CleanupOK: true}
for index := 1; index <= 8; index++ {
value.AgentSummaries = append(value.AgentSummaries, heldSessionSetRealAgentSummary{Ordinal: index, ServerDigest: heldRealDigest(string(rune('a' + index))), PATIdentity: true, PATScopeExact: true})
}
for index, session := range plan.Sessions {
digest := heldRealDigest(string(rune('A' + index)))
value.SessionDigests = append(value.SessionDigests, digest)
value.SessionSummaries = append(value.SessionSummaries, heldSessionSetRealSessionSummary{Ordinal: index + 1, Kind: string(session.Kind), AgentOrdinal: session.Agent.Int(), StreamDigest: digest, Present: true, Absent: true, Protocol: true})
}
return value
}
func TestHeldSessionSetRealPlanUsesCanonicalPRFullTopology(t *testing.T) {
profile, err := contract.ProfileByName(string(contract.ProfilePRFull))
require.NoError(t, err)
plan, err := GenerateStressPlan(profile, contract.DefaultSeed)
require.NoError(t, err)
require.Equal(t, "pr-full", string(plan.Profile))
require.Equal(t, uint64(0x4e5a4841), uint64(plan.Seed))
require.Len(t, plan.Sessions, 12)
require.Equal(t, 4, countHeldSessionKind(plan.Sessions, StressSessionTerminal))
require.Equal(t, 4, countHeldSessionKind(plan.Sessions, StressSessionNAT))
require.Equal(t, 4, countHeldSessionKind(plan.Sessions, StressSessionFM))
for ordinal := 1; ordinal <= heldSessionSetAgentCount; ordinal++ {
require.Contains(t, realHeldSessionSetAgentOrdinals(plan), ordinal)
}
}
func TestHeldSessionSetRealEvidenceRejectsIncompleteAndRedactsArtifact(t *testing.T) {
root := t.TempDir()
require.NoError(t, os.Chmod(root, 0o700))
evidence := validHeldSessionSetRealEvidence()
require.Error(t, validateHeldSessionSetRealEvidence(heldSessionSetRealEvidence{}))
require.NoError(t, validateHeldSessionSetRealEvidence(evidence))
require.NoError(t, writeHeldSessionSetRealEvidence(root, evidence))
readBack, err := readHeldSessionSetRealEvidence(root)
require.NoError(t, err)
require.Equal(t, evidence, readBack)
}
func TestHeldSessionSetRealEvidenceRejectsCanonicalMutations(t *testing.T) {
tests := []struct {
name string
mutate func(*heldSessionSetRealEvidence)
}{
{name: "version zero", mutate: func(value *heldSessionSetRealEvidence) { value.Version = 0 }},
{name: "wrong version", mutate: func(value *heldSessionSetRealEvidence) { value.Version = 2 }},
{name: "wrong seed", mutate: func(value *heldSessionSetRealEvidence) { value.Seed = "4e5a4842" }},
{name: "uppercase digest", mutate: func(value *heldSessionSetRealEvidence) {
value.SessionDigests[0] = strings.ToUpper(value.SessionDigests[0])
}},
{name: "non hex digest", mutate: func(value *heldSessionSetRealEvidence) { value.SessionDigests[0] = strings.Repeat("z", 64) }},
{name: "duplicate server digest", mutate: func(value *heldSessionSetRealEvidence) {
value.AgentSummaries[1].ServerDigest = value.AgentSummaries[0].ServerDigest
}},
{name: "duplicate session digest", mutate: func(value *heldSessionSetRealEvidence) {
value.SessionDigests[1] = value.SessionDigests[0]
value.SessionSummaries[1].StreamDigest = value.SessionDigests[0]
}},
{name: "summary mismatch", mutate: func(value *heldSessionSetRealEvidence) {
value.SessionSummaries[0].StreamDigest = heldRealDigest("different")
}},
{name: "digest order", mutate: func(value *heldSessionSetRealEvidence) { slices.Reverse(value.SessionDigests) }},
{name: "malformed length", mutate: func(value *heldSessionSetRealEvidence) { value.SessionDigests[0] = value.SessionDigests[0][:63] }},
}
for _, testCase := range tests {
t.Run(testCase.name, func(t *testing.T) {
value := validHeldSessionSetRealEvidence()
testCase.mutate(&value)
err := validateHeldSessionSetRealEvidence(value)
require.ErrorIs(t, err, ErrHeldSessionSetRealEvidenceInvalid)
require.NotContains(t, err.Error(), "secret")
})
}
}
func validHeldSessionSetRealEvidence() heldSessionSetRealEvidence {
return validHeldSessionSetRealEvidenceFromPlan()
}
@@ -0,0 +1,44 @@
//go:build linux
package scenario
import "errors"
var (
ErrHeldSessionSetOperation = errors.New("held session set operation failed")
ErrHeldSessionSetHealth = errors.New("held session set health failed")
ErrHeldSessionPrematureClose = errors.New("held session closed before set close")
)
type heldSessionSetClassifiedError struct {
class error
causes []error
}
func (err *heldSessionSetClassifiedError) Error() string { return err.class.Error() }
func (err *heldSessionSetClassifiedError) Is(target error) bool {
if errors.Is(err.class, target) {
return true
}
for _, cause := range err.causes {
if errors.Is(cause, target) {
return true
}
}
return false
}
func redactHeldSessionSetError(err error) error {
if err == nil {
return nil
}
return &heldSessionSetClassifiedError{class: ErrHeldSessionSetOperation, causes: []error{err}}
}
func redactHeldSessionSetHealthError(err error) error {
if err == nil {
return nil
}
return &heldSessionSetClassifiedError{class: ErrHeldSessionSetHealth, causes: []error{err}}
}
@@ -0,0 +1,284 @@
//go:build linux
package scenario
import (
"context"
"errors"
"sync"
"testing"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
type heldSessionHealthFake struct {
plan StressSessionPlan
done chan struct{}
closeError error
}
func (fake *heldSessionHealthFake) Plan() StressSessionPlan { return fake.plan }
func (fake *heldSessionHealthFake) WaitLive(context.Context) error { return nil }
func (fake *heldSessionHealthFake) Close(context.Context) error { return nil }
func (fake *heldSessionHealthFake) WaitClosed(context.Context) error { return nil }
func (fake *heldSessionHealthFake) IOStreamID() (string, bool) { return "health-fake", true }
func (fake *heldSessionHealthFake) Done() <-chan struct{} { return fake.done }
func (fake *heldSessionHealthFake) CloseResult() error { return fake.closeError }
func newHeldSessionHealthFake(t *testing.T, plan StressSessionPlan, closeError error) *heldSessionHealthFake {
t.Helper()
return &heldSessionHealthFake{plan: plan, done: make(chan struct{}), closeError: closeError}
}
type heldSessionSetTestSession struct {
mu sync.Mutex
plan StressSessionPlan
index int
streamID string
waitLiveError error
closeError error
waitClosedError error
closeResult error
done chan struct{}
closed bool
closeEvents chan int
waitClosedEvents chan int
closeStarted chan struct{}
closeRelease <-chan struct{}
}
func (session *heldSessionSetTestSession) Plan() StressSessionPlan { return session.plan }
func (session *heldSessionSetTestSession) WaitLive(context.Context) error {
return session.waitLiveError
}
func (session *heldSessionSetTestSession) IOStreamID() (string, bool) {
return session.streamID, session.streamID != ""
}
func (session *heldSessionSetTestSession) Done() <-chan struct{} { return session.done }
func (session *heldSessionSetTestSession) CloseResult() error { return session.closeResult }
func (session *heldSessionSetTestSession) prematureClose() {
session.mu.Lock()
defer session.mu.Unlock()
if !session.closed {
session.closed = true
close(session.done)
}
}
func (session *heldSessionSetTestSession) Close(context.Context) error {
if session.closeStarted != nil {
close(session.closeStarted)
session.closeStarted = nil
}
if session.closeRelease != nil {
<-session.closeRelease
}
session.mu.Lock()
if !session.closed {
session.closed = true
close(session.done)
}
if session.closeEvents != nil {
session.closeEvents <- session.index
}
session.mu.Unlock()
return session.closeError
}
func (session *heldSessionSetTestSession) WaitClosed(context.Context) error {
if session.waitClosedEvents != nil {
session.waitClosedEvents <- session.index
}
return session.waitClosedError
}
type heldSessionSetConstructorCoordinator struct {
mu sync.Mutex
startOnce sync.Once
ready chan int
started chan struct{}
startEvents chan string
startSlots []chan struct{}
completed chan heldSessionConstructorCompletion
acquired chan int
sessions map[int]*heldSessionSetTestSession
contextSeen chan context.Context
blockedIndex int
blockedWaiting chan struct{}
blockedCanceled chan struct{}
errors map[string]error
indices map[string]int
released []bool
lateSuccessIndex int
lateSuccessWaiting chan struct{}
lateSuccessReady chan struct{}
lateSuccessRelease chan struct{}
errorReady map[int]chan struct{}
}
type heldSessionConstructorCompletion struct {
index int
acquired bool
err error
}
func newHeldSessionSetConstructorCoordinator() *heldSessionSetConstructorCoordinator {
startSlots := make([]chan struct{}, 12)
for index := range startSlots {
startSlots[index] = make(chan struct{})
}
return &heldSessionSetConstructorCoordinator{ready: make(chan int, 12), started: make(chan struct{}), startEvents: make(chan string, 12), startSlots: startSlots, completed: make(chan heldSessionConstructorCompletion, 12), acquired: make(chan int, 12), sessions: make(map[int]*heldSessionSetTestSession), contextSeen: make(chan context.Context, 12), blockedIndex: -1, errors: make(map[string]error), indices: make(map[string]int), released: make([]bool, 12), lateSuccessIndex: -1, errorReady: make(map[int]chan struct{})}
}
func (coordinator *heldSessionSetConstructorCoordinator) construct(ctx context.Context, plan StressSessionPlan) (heldSession, error) {
index, exists := coordinator.indices[plan.ID.String()]
if !exists {
coordinator.completed <- heldSessionConstructorCompletion{index: -1, err: errors.New("constructor plan is not coordinated")}
return nil, errors.New("constructor plan is not coordinated")
}
coordinator.ready <- index
<-coordinator.started
if index == coordinator.blockedIndex {
coordinator.blockedWaiting <- struct{}{}
<-ctx.Done()
coordinator.blockedCanceled <- struct{}{}
coordinator.completed <- heldSessionConstructorCompletion{index: index, err: ctx.Err()}
return nil, ctx.Err()
}
if coordinator.errors[plan.ID.String()] != nil {
if ready := coordinator.errorReady[index]; ready != nil {
<-ready
}
coordinator.startEvents <- plan.ID.String()
err := coordinator.errors[plan.ID.String()]
coordinator.completed <- heldSessionConstructorCompletion{index: index, err: err}
return nil, err
}
if !coordinator.slotReleased(index) {
if index == coordinator.lateSuccessIndex {
coordinator.lateSuccessWaiting <- struct{}{}
<-coordinator.startSlots[index]
} else {
select {
case <-coordinator.startSlots[index]:
case <-ctx.Done():
err := ctx.Err()
coordinator.startEvents <- plan.ID.String()
coordinator.completed <- heldSessionConstructorCompletion{index: index, err: err}
return nil, err
}
}
}
coordinator.startEvents <- plan.ID.String()
if err := coordinator.errors[plan.ID.String()]; err != nil {
coordinator.completed <- heldSessionConstructorCompletion{index: index, err: err}
return nil, err
}
if index == coordinator.lateSuccessIndex {
coordinator.lateSuccessReady <- struct{}{}
<-coordinator.lateSuccessRelease
}
coordinator.completed <- heldSessionConstructorCompletion{index: index, acquired: true}
coordinator.acquired <- index
return coordinator.sessions[index], nil
}
func (coordinator *heldSessionSetConstructorCoordinator) releaseError(index int) {
close(coordinator.errorReady[index])
}
func (coordinator *heldSessionSetConstructorCoordinator) releaseAll() {
coordinator.startOnce.Do(func() { close(coordinator.started) })
}
func (coordinator *heldSessionSetConstructorCoordinator) release(index int) {
coordinator.mu.Lock()
coordinator.released[index] = true
coordinator.mu.Unlock()
close(coordinator.startSlots[index])
}
func (coordinator *heldSessionSetConstructorCoordinator) slotReleased(index int) bool {
coordinator.mu.Lock()
defer coordinator.mu.Unlock()
return coordinator.released[index]
}
type heldSessionSetStateFake struct {
mu sync.Mutex
state client.IOStreamState
baseline client.IOStreamState
snapshotError error
present []string
absent []string
counts []int
presentStreamErrors map[string]error
absentStreamErrors map[string]error
presentAggregateError error
absentAggregateError error
presentAggregateCalls int
absentAggregateCalls int
aggregatePredicatesEmpty bool
waitCalls chan client.IOStreamStateExpectation
expectedPresentCount int
expectedAbsentCount int
}
func (fake *heldSessionSetStateFake) IOStreamState(context.Context) (client.IOStreamState, error) {
if fake.baseline.Count == 0 && fake.state.Count != 0 {
return fake.state, fake.snapshotError
}
return fake.baseline, fake.snapshotError
}
func (fake *heldSessionSetStateFake) wait(ctx context.Context, expectation client.IOStreamStateExpectation) (client.IOStreamState, error) {
return fake.WaitForIOStreamState(ctx, expectation)
}
func (fake *heldSessionSetStateFake) WaitForIOStreamState(ctx context.Context, expectation client.IOStreamStateExpectation) (client.IOStreamState, error) {
fake.mu.Lock()
fake.counts = append(fake.counts, *expectation.ExpectedCount)
if expectation.PresentStreamID != "" {
fake.present = append(fake.present, expectation.PresentStreamID)
}
if expectation.AbsentStreamID != "" {
fake.absent = append(fake.absent, expectation.AbsentStreamID)
}
isAggregate := expectation.PresentStreamID == "" && expectation.AbsentStreamID == ""
err := fake.presentAggregateError
expectedCount := fake.expectedPresentCount
if expectation.AbsentStreamID != "" {
expectedCount = fake.expectedAbsentCount
err = fake.absentStreamErrors[expectation.AbsentStreamID]
} else if expectation.PresentStreamID != "" {
err = fake.presentStreamErrors[expectation.PresentStreamID]
} else if *expectation.ExpectedCount == fake.expectedAbsentCount {
expectedCount = fake.expectedAbsentCount
err = fake.absentAggregateError
}
if isAggregate {
fake.aggregatePredicatesEmpty = fake.aggregatePredicatesEmpty || expectation.PresentStreamID == "" && expectation.AbsentStreamID == ""
if *expectation.ExpectedCount == fake.expectedPresentCount {
fake.presentAggregateCalls++
} else if *expectation.ExpectedCount == fake.expectedAbsentCount {
fake.absentAggregateCalls++
} else {
err = errors.Join(err, ErrHeldSessionSetOperation)
}
}
if expectedCount != 0 && !isAggregate && *expectation.ExpectedCount != expectedCount {
err = errors.Join(err, ErrHeldSessionSetOperation)
}
if fake.waitCalls != nil {
fake.waitCalls <- expectation
}
fake.mu.Unlock()
if err := ctx.Err(); err != nil {
return client.IOStreamState{}, err
}
return fake.baseline, err
}
func newHeldSessionSetTestSession(index int, plan StressSessionPlan, closeEvents, waitClosedEvents chan int) *heldSessionSetTestSession {
return &heldSessionSetTestSession{plan: plan, index: index, streamID: "stream-" + plan.ID.String(), done: make(chan struct{}), closeEvents: closeEvents, waitClosedEvents: waitClosedEvents}
}
@@ -0,0 +1,67 @@
//go:build linux
package scenario
import (
"context"
"errors"
"testing"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
func TestHeldSessionSetErrorsRedactNestedIdentity(t *testing.T) {
// Given
secret := errors.New("stream=secret-stream uuid=secret-uuid server=991 authorization=secret-token")
// When
err := redactHeldSessionSetError(secret)
// Then
require.ErrorIs(t, err, secret)
require.Equal(t, "held session set operation failed", err.Error())
require.NotContains(t, err.Error(), "secret-stream")
require.NotContains(t, err.Error(), "secret-token")
}
func TestHeldSessionSetStateChecksUseEveryExactID(t *testing.T) {
// Given
observer := &heldSessionSetStateFake{state: client.IOStreamState{Count: 12}}
ids := []string{"stream-a", "stream-b", "stream-c"}
// When
err := waitHeldSessionSetStreams(context.Background(), observer, 12, ids, true, func(ctx context.Context, state heldSessionSetStateObserver, expectation client.IOStreamStateExpectation) (client.IOStreamState, error) {
return state.(*heldSessionSetStateFake).wait(ctx, expectation)
})
// Then
require.NoError(t, err)
require.ElementsMatch(t, ids, observer.present)
}
func TestHeldSessionSetHealthUsesCanonicalIndexWhenClosuresRace(t *testing.T) {
// Given
firstError := errors.New("first canonical health failure")
secondError := errors.New("second canonical health failure")
firstPlan := StressSessionPlan{Kind: StressSessionTerminal, Ordinal: 1}
secondPlan := StressSessionPlan{Kind: StressSessionNAT, Ordinal: 2}
first := newHeldSessionHealthFake(t, firstPlan, firstError)
second := newHeldSessionHealthFake(t, secondPlan, secondError)
set := newHeldSessionSet([]StressSessionPlan{firstPlan, secondPlan}, &heldSessionSetStateFake{state: client.IOStreamState{Count: 2}}, client.IOStreamState{Count: 0}, HeldSessionSetDependencies{}, context.Background())
set.sessions = []heldSession{first, second}
set.startHealthWatchers()
// When
close(second.done)
close(first.done)
err := set.WaitHealthy(context.Background())
set.stopHealthWatchers()
set.healthWG.Wait()
// Then
require.Error(t, err)
require.ErrorIs(t, err, firstError)
require.NotErrorIs(t, err, secondError)
}
@@ -0,0 +1,45 @@
//go:build linux
package scenario
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
)
func TestHeldSessionSetCanonicalPlanHasExactlyFourOfEachKind(t *testing.T) {
// Given
profile, err := contract.ProfileByName(string(contract.ProfilePRFull))
require.NoError(t, err)
plan, err := GenerateStressPlan(profile, contract.DefaultSeed)
require.NoError(t, err)
// When
validated, err := validateHeldSessionSetPlans(plan)
// Then
require.NoError(t, err)
require.Len(t, validated, 12)
require.Equal(t, 4, countHeldSessionKind(validated, StressSessionTerminal))
require.Equal(t, 4, countHeldSessionKind(validated, StressSessionNAT))
require.Equal(t, 4, countHeldSessionKind(validated, StressSessionFM))
}
func TestHeldSessionSetRejectsInvalidTopologyBeforeConstruction(t *testing.T) {
// Given
profile, err := contract.ProfileByName(string(contract.ProfilePRFull))
require.NoError(t, err)
plan, err := GenerateStressPlan(profile, contract.DefaultSeed)
require.NoError(t, err)
// When
_, err = NewHeldSessionSet(context.Background(), HeldSessionSetInput{Plan: plan})
// Then
require.Error(t, err)
require.ErrorIs(t, err, ErrInvalidHeldSessionSetTopology)
}
@@ -0,0 +1,233 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"reflect"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
)
var (
ErrInvalidHeldSessionSetTopology = errors.New("held session set topology is invalid")
ErrInvalidHeldSessionSetPlan = errors.New("held session set plan is invalid")
)
const heldSessionSetAgentCount = 8
type HeldSessionAgent struct {
Ordinal StressAgentOrdinal
Agent *agent.Agent
Readiness agent.Readiness
PATClient *client.Client
}
type HeldSessionSetInput struct {
Dashboard *dashboard.Dashboard
Plan StressPlan
Topology []HeldSessionAgent
Dependencies HeldSessionSetDependencies
ControlClient *client.Client
ControlServerIDs []uint64
testHealthSnapshotRequestHook func(int)
testHealthSnapshotReplyHook func(int)
testHealthSnapshotOverrideHook func(int, heldHealthSnapshotRequest) *heldHealthSnapshot
testHealthSnapshotSendHook func(int)
testHealthEventHook func(heldHealthMessage)
testHealthClosureObservedHook func(int)
testHealthSnapshotAcceptedHook func(heldHealthSnapshot)
testHealthShutdownAcceptedHook func()
testHealthShutdownAcknowledgedHook func()
}
type heldSessionSetTopology struct {
dashboard *dashboard.Dashboard
stateClient heldSessionSetStateObserver
agents map[int]HeldSessionAgent
}
type heldSessionSetStateObserver interface {
IOStreamState(context.Context) (client.IOStreamState, error)
WaitForIOStreamState(context.Context, client.IOStreamStateExpectation) (client.IOStreamState, error)
}
type heldSessionSetAuthorizedObserver struct {
client *client.Client
}
func (observer heldSessionSetAuthorizedObserver) IOStreamState(ctx context.Context) (client.IOStreamState, error) {
return observer.client.IOStreamState(ctx)
}
func (observer heldSessionSetAuthorizedObserver) WaitForIOStreamState(ctx context.Context, expectation client.IOStreamStateExpectation) (client.IOStreamState, error) {
return observer.client.WaitForIOStreamState(ctx, expectation)
}
type HeldSessionSetDependencies struct {
Terminal func(context.Context, heldTerminalInput) (heldSession, error)
NAT func(context.Context, heldNATInput) (heldSession, error)
FM func(context.Context, heldLegacyFMInput) (heldSession, error)
Snapshot func(context.Context, heldSessionSetStateObserver) (client.IOStreamState, error)
WaitState func(context.Context, heldSessionSetStateObserver, client.IOStreamStateExpectation) (client.IOStreamState, error)
InspectAgent func(*agent.Agent) heldSessionAgentFacts
ObserveState func(*client.Client) heldSessionSetStateObserver
}
type heldSessionAgentFacts struct {
PID int
UUID string
}
func defaultHeldSessionSetDependencies() HeldSessionSetDependencies {
return HeldSessionSetDependencies{
Terminal: func(ctx context.Context, input heldTerminalInput) (heldSession, error) {
return newHeldTerminalSession(ctx, input)
},
NAT: func(ctx context.Context, input heldNATInput) (heldSession, error) {
return newHeldNATSession(ctx, input)
},
FM: func(ctx context.Context, input heldLegacyFMInput) (heldSession, error) {
return newHeldLegacyFMSession(ctx, input)
},
Snapshot: func(ctx context.Context, stateClient heldSessionSetStateObserver) (client.IOStreamState, error) {
return stateClient.IOStreamState(ctx)
},
WaitState: func(ctx context.Context, stateClient heldSessionSetStateObserver, expectation client.IOStreamStateExpectation) (client.IOStreamState, error) {
return stateClient.WaitForIOStreamState(ctx, expectation)
},
InspectAgent: func(instance *agent.Agent) heldSessionAgentFacts {
return heldSessionAgentFacts{PID: instance.PID(), UUID: instance.UUID()}
},
ObserveState: func(controlClient *client.Client) heldSessionSetStateObserver {
return heldSessionSetAuthorizedObserver{client: controlClient}
},
}
}
func validateHeldSessionSetPlans(plan StressPlan) ([]StressSessionPlan, error) {
canonical, err := canonicalHeldSessionPlans(plan)
if err != nil {
return nil, errors.Join(ErrInvalidHeldSessionSetPlan, err)
}
if !reflect.DeepEqual(canonical, plan.Sessions) || len(plan.Sessions) != 12 {
return nil, ErrInvalidHeldSessionSetPlan
}
ids := make(map[string]struct{}, len(plan.Sessions))
for _, session := range plan.Sessions {
if session.ID.String() == "" {
return nil, fmt.Errorf("empty session plan ID: %w", ErrInvalidHeldSessionSetPlan)
}
if _, exists := ids[session.ID.String()]; exists {
return nil, fmt.Errorf("duplicate session plan ID: %s: %w", session.ID.String(), ErrInvalidHeldSessionSetPlan)
}
ids[session.ID.String()] = struct{}{}
}
if countHeldSessionKind(plan.Sessions, StressSessionTerminal) != 4 || countHeldSessionKind(plan.Sessions, StressSessionNAT) != 4 || countHeldSessionKind(plan.Sessions, StressSessionFM) != 4 {
return nil, ErrInvalidHeldSessionSetPlan
}
return append([]StressSessionPlan(nil), plan.Sessions...), nil
}
func canonicalHeldSessionPlans(plan StressPlan) ([]StressSessionPlan, error) {
if plan.Seed == 0 || plan.Profile == "" {
return nil, ErrInvalidHeldSessionSetPlan
}
profile, err := contract.ProfileByName(string(plan.Profile))
if err != nil {
return nil, err
}
canonical, err := GenerateStressPlan(profile, plan.Seed)
if err != nil {
return nil, err
}
return canonical.Sessions, nil
}
func countHeldSessionKind(plans []StressSessionPlan, kind StressSessionKind) int {
count := 0
for _, plan := range plans {
if plan.Kind == kind {
count++
}
}
return count
}
func validateHeldSessionSetTopology(input HeldSessionSetInput, plans []StressSessionPlan) (heldSessionSetTopology, error) {
if input.Dashboard == nil || input.ControlClient == nil {
return heldSessionSetTopology{}, ErrInvalidHeldSessionSetTopology
}
inspectAgent := input.Dependencies.InspectAgent
if inspectAgent == nil {
inspectAgent = defaultHeldSessionSetDependencies().InspectAgent
}
observeState := input.Dependencies.ObserveState
if observeState == nil {
observeState = defaultHeldSessionSetDependencies().ObserveState
}
profile, err := contract.ProfileByName(string(input.Plan.Profile))
if err != nil || profile.AgentCount() != heldSessionSetAgentCount || len(input.Topology) != heldSessionSetAgentCount {
return heldSessionSetTopology{}, ErrInvalidHeldSessionSetTopology
}
agents := make(map[int]HeldSessionAgent, len(input.Topology))
uuids := make(map[string]struct{}, len(input.Topology))
serverIDs := make(map[uint64]struct{}, len(input.Topology))
for _, topology := range input.Topology {
ordinal := topology.Ordinal.Int()
if ordinal < 1 || ordinal > heldSessionSetAgentCount || topology.Agent == nil || topology.PATClient == nil {
return heldSessionSetTopology{}, ErrInvalidHeldSessionSetTopology
}
if _, exists := agents[ordinal]; exists {
return heldSessionSetTopology{}, fmt.Errorf("duplicate agent ordinal %d: %w", ordinal, ErrInvalidHeldSessionSetTopology)
}
agentFacts := inspectAgent(topology.Agent)
if err := validateHeldReadinessFacts(agentFacts, topology.Readiness); err != nil {
return heldSessionSetTopology{}, err
}
if agentFacts.PID < 1 || topology.Readiness.ServerID == 0 || topology.Readiness.UUID != agentFacts.UUID {
return heldSessionSetTopology{}, ErrInvalidHeldSessionSetTopology
}
if _, exists := uuids[topology.Readiness.UUID]; exists {
return heldSessionSetTopology{}, ErrInvalidHeldSessionSetTopology
}
if _, exists := serverIDs[topology.Readiness.ServerID]; exists {
return heldSessionSetTopology{}, ErrInvalidHeldSessionSetTopology
}
uuids[topology.Readiness.UUID] = struct{}{}
serverIDs[topology.Readiness.ServerID] = struct{}{}
agents[ordinal] = topology
}
if len(input.ControlServerIDs) != heldSessionSetAgentCount {
return heldSessionSetTopology{}, ErrInvalidHeldSessionSetTopology
}
controlServerIDs := make(map[uint64]struct{}, len(input.ControlServerIDs))
for _, serverID := range input.ControlServerIDs {
if serverID == 0 {
return heldSessionSetTopology{}, ErrInvalidHeldSessionSetTopology
}
if _, exists := serverIDs[serverID]; !exists {
return heldSessionSetTopology{}, ErrInvalidHeldSessionSetTopology
}
if _, exists := controlServerIDs[serverID]; exists {
return heldSessionSetTopology{}, ErrInvalidHeldSessionSetTopology
}
controlServerIDs[serverID] = struct{}{}
}
for ordinal := 1; ordinal <= heldSessionSetAgentCount; ordinal++ {
if _, exists := agents[ordinal]; !exists {
return heldSessionSetTopology{}, ErrInvalidHeldSessionSetTopology
}
}
for _, plan := range plans {
if _, exists := agents[plan.Agent.Int()]; !exists {
return heldSessionSetTopology{}, fmt.Errorf("missing agent ordinal %d: %w", plan.Agent.Int(), ErrInvalidHeldSessionSetTopology)
}
}
return heldSessionSetTopology{dashboard: input.Dashboard, stateClient: observeState(input.ControlClient), agents: agents}, nil
}
@@ -0,0 +1,193 @@
//go:build linux
package scenario
import (
"context"
"errors"
"sync"
"testing"
"time"
)
func TestHeldSessionLifecycleRejectsInvalidConstruction(t *testing.T) {
validPlan := heldTestPlan(t)
cases := []struct {
name string
plan StressSessionPlan
timeout time.Duration
}{
{"zero-session-id", StressSessionPlan{Kind: validPlan.Kind, Ordinal: 1, Agent: validPlan.Agent}, time.Second},
{"unsupported-kind", StressSessionPlan{ID: validPlan.ID, Kind: StressSessionKind("unsupported"), Ordinal: 1, Agent: validPlan.Agent}, time.Second},
{"zero-ordinal", StressSessionPlan{ID: validPlan.ID, Kind: validPlan.Kind, Agent: validPlan.Agent}, time.Second},
{"zero-agent", StressSessionPlan{ID: validPlan.ID, Kind: validPlan.Kind, Ordinal: 1}, time.Second},
{"zero-timeout", validPlan, 0},
{"nil-base-context", validPlan, time.Second},
}
for _, testCase := range cases {
t.Run(testCase.name, func(t *testing.T) {
baseContext := context.Background()
if testCase.name == "nil-base-context" {
baseContext = nil
}
_, err := newHeldSessionLifecycle(baseContext, testCase.plan, "", testCase.timeout)
if !errors.Is(err, ErrInvalidHeldSessionPlan) {
t.Fatalf("construction error = %v", err)
}
})
}
}
func TestHeldSessionLifecycleRetainsLiveResultAndOptionalIOStreamID(t *testing.T) {
lifecycle := heldTestLifecycle(t, "io-stream-identity")
if err := lifecycle.markLive(nil); err != nil {
t.Fatal(err)
}
if err := lifecycle.WaitLive(context.Background()); err != nil {
t.Fatal(err)
}
streamID, present := lifecycle.IOStreamID()
if !present || streamID != "io-stream-identity" {
t.Fatalf("IOStream identity = %q, %v", streamID, present)
}
owner, won := lifecycle.beginClose()
if !won {
t.Fatal("beginClose did not return owner")
}
owner.markClosed(nil)
if err := lifecycle.WaitClosed(context.Background()); err != nil {
t.Fatal(err)
}
}
func TestHeldSessionLifecycleFailedLiveStateRetainsExactError(t *testing.T) {
liveErr := errors.New("session failed to become live")
lifecycle := heldTestLifecycle(t, "")
if err := lifecycle.markLive(liveErr); err != nil {
t.Fatal(err)
}
if lifecycle.state != heldSessionFailed {
t.Fatalf("state = %v, want failed", lifecycle.state)
}
if err := lifecycle.WaitLive(context.Background()); !errors.Is(err, liveErr) {
t.Fatalf("WaitLive error = %v", err)
}
if err := lifecycle.markLive(nil); !errors.Is(err, ErrHeldSessionLiveResolved) {
t.Fatalf("second markLive error = %v", err)
}
}
func TestHeldSessionLifecycleCloseBeforeLiveRetainsFailure(t *testing.T) {
lifecycle := heldTestLifecycle(t, "")
owner, won := lifecycle.beginClose()
if !won {
t.Fatal("beginClose did not return owner")
}
if err := lifecycle.WaitLive(context.Background()); !errors.Is(err, ErrHeldSessionClosedBeforeLive) {
t.Fatalf("WaitLive error = %v", err)
}
owner.markClosed(nil)
if err := lifecycle.WaitClosed(context.Background()); err != nil {
t.Fatal(err)
}
}
func TestHeldSessionLifecycleFailedLiveRetainsDistinctCleanupResult(t *testing.T) {
liveErr := errors.New("live failed")
cleanupErr := errors.New("cleanup failed")
lifecycle := heldTestLifecycle(t, "")
if err := lifecycle.markLive(liveErr); err != nil {
t.Fatal(err)
}
owner, won := lifecycle.beginClose()
if !won {
t.Fatal("beginClose did not return owner")
}
owner.markClosed(cleanupErr)
if err := lifecycle.WaitLive(context.Background()); !errors.Is(err, liveErr) {
t.Fatalf("live error = %v", err)
}
if err := lifecycle.WaitClosed(context.Background()); !errors.Is(err, cleanupErr) {
t.Fatalf("closed error = %v", err)
}
}
func TestHeldSessionLifecycleDoesNotImplementHeldSession(t *testing.T) {
lifecycle := heldTestLifecycle(t, "")
var candidate any = lifecycle
if _, ok := candidate.(heldSession); ok {
t.Fatal("lifecycle unexpectedly implements heldSession; cleanup ownership belongs to adapters")
}
}
func TestHeldSessionLifecycleBeginCloseHasSingleWinner(t *testing.T) {
lifecycle := heldTestLifecycle(t, "")
owners := make(chan *heldSessionCloseOwner, 2)
var waitGroup sync.WaitGroup
for range 2 {
waitGroup.Go(func() {
owner, won := lifecycle.beginClose()
if won {
owners <- owner
}
})
}
waitGroup.Wait()
close(owners)
var owner *heldSessionCloseOwner
for candidate := range owners {
if owner != nil {
t.Fatal("beginClose returned two owners")
}
owner = candidate
}
if owner == nil {
t.Fatal("beginClose returned no owner")
}
owner.markClosed(nil)
}
func TestHeldSessionCloseOwnerCleanupContextIgnoresParentCancellation(t *testing.T) {
parent, cancel := context.WithCancel(context.Background())
cancel()
lifecycle := heldTestLifecycleWithBase(t, parent, time.Second)
owner, won := lifecycle.beginClose()
if !won {
t.Fatal("beginClose did not return owner")
}
cleanupContext, cleanupCancel := owner.cleanupContext()
defer cleanupCancel()
if err := cleanupContext.Err(); err != nil {
t.Fatalf("cleanup context already canceled: %v", err)
}
owner.markClosed(nil)
}
func heldTestPlan(t *testing.T) StressSessionPlan {
t.Helper()
id, err := NewStressSessionID("held-session")
if err != nil {
t.Fatal(err)
}
agent, err := NewStressAgentOrdinal(1)
if err != nil {
t.Fatal(err)
}
return StressSessionPlan{ID: id, Kind: StressSessionTerminal, Ordinal: 1, Agent: agent}
}
func heldTestLifecycle(t *testing.T, streamID string) *heldSessionLifecycle {
return heldTestLifecycleWithBaseAndID(t, context.Background(), streamID, time.Second)
}
func heldTestLifecycleWithBase(t *testing.T, base context.Context, timeout time.Duration) *heldSessionLifecycle {
return heldTestLifecycleWithBaseAndID(t, base, "", timeout)
}
func heldTestLifecycleWithBaseAndID(t *testing.T, base context.Context, streamID string, timeout time.Duration) *heldSessionLifecycle {
lifecycle, err := newHeldSessionLifecycle(base, heldTestPlan(t), streamID, timeout)
if err != nil {
t.Fatal(err)
}
return lifecycle
}
@@ -0,0 +1,270 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"net/http"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
)
const (
heldTerminalCleanupTimeout = 10 * time.Second
heldTerminalPumpCapacity = 32
heldTerminalGracePeriod = time.Second
)
var (
ErrInvalidHeldTerminalInput = errors.New("held terminal input is invalid")
ErrHeldTerminalProtocol = errors.New("held terminal protocol proof failed")
ErrInvalidHeldPATClient = errors.New("held PAT client is invalid")
)
type heldTerminalInput struct {
Dashboard *dashboard.Dashboard
PATClient *client.Client
Agent *agent.Agent
Readiness agent.Readiness
Plan StressSessionPlan
LifetimeContext context.Context
}
type heldTerminalSession struct {
lifecycle *heldSessionLifecycle
stack *heldCleanupStack
connection heldTerminalConnection
pump heldTerminalPump
protocol bool
}
type heldTerminalConnection interface {
WriteFrame(context.Context, client.Frame) error
Close() error
}
type heldTerminalPump interface {
Events() <-chan client.Frame
Done() <-chan struct{}
Err() error
Stop(context.Context) error
Wait(context.Context) error
}
func newHeldTerminalSession(ctx context.Context, input heldTerminalInput) (*heldTerminalSession, error) {
if err := validateHeldTerminalInput(ctx, input); err != nil {
return nil, err
}
if err := validateHeldPATClient(input.PATClient); err != nil {
return nil, err
}
if err := validateHeldReadiness(input.Agent, input.Readiness); err != nil {
return nil, err
}
stateClient := input.PATClient
baseline, err := stateClient.IOStreamState(ctx)
if err != nil {
return nil, fmt.Errorf("snapshot terminal IOStream state: %w", err)
}
capability, err := registerHeldIOStreamCapability(ctx, input.PATClient, heldIOStreamCapabilityIdentity{Purpose: client.IOStreamCapabilityPurposeTerminal, ServerID: input.Readiness.ServerID})
if err != nil {
return nil, fmt.Errorf("register terminal capability: %w", err)
}
stack := newHeldCleanupStack()
if err := stack.Push(heldCleanupAction{name: "unregister terminal capability", cleanup: capability.Unregister}); err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
created, err := client.DoREST[terminalCreateRequest, terminalCreateResponse](ctx, input.PATClient, client.RESTRequest[terminalCreateRequest]{
Method: http.MethodPost, Path: "/api/v1/terminal", Body: &terminalCreateRequest{Protocol: "grpc", ServerID: input.Readiness.ServerID},
IOStreamCapability: capability.HeaderCapability(),
})
if err != nil {
return nil, rollbackHeldTerminal(ctx, stack, fmt.Errorf("create held terminal: %w", err))
}
streamID, err := capability.Wait(ctx)
if err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
if streamID != created.SessionID {
return nil, rollbackHeldTerminal(ctx, stack, fmt.Errorf("terminal stream mismatch: response_and_capability_ids_differ: %w", ErrHeldTerminalProtocol))
}
if err := stack.Push(heldCleanupAction{name: "wait for terminal stream absence", cleanup: func(cleanupContext context.Context) error {
return capability.waitExpectation(cleanupContext, stateClient, baseline, true)
}}); err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
if err := stack.Push(heldCleanupAction{name: "cancel terminal capability", cleanup: capability.Cancel}); err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
if err := validateHeldTerminalResponse(created, input.Readiness.ServerID); err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
lifetimeContext := input.LifetimeContext
if lifetimeContext == nil {
lifetimeContext = ctx
}
lifecycle, err := newHeldSessionLifecycle(lifetimeContext, input.Plan, streamID, heldTerminalCleanupTimeout)
if err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
connection, err := input.PATClient.DialWebSocket(ctx, "/api/v1/ws/terminal/"+created.SessionID)
if err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
resize, err := terminalResizeFrame(132, 43)
if err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
pump, err := newHeldWebSocketPump(lifetimeContext, connection, heldTerminalPumpCapacity)
if err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
if err := stack.Push(heldCleanupAction{name: "close terminal WebSocket", cleanup: func(context.Context) error { return connection.Close() }}); err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
if err := stack.Push(heldCleanupAction{name: "stop terminal WebSocket pump", cleanup: func(cleanupContext context.Context) error {
err := pump.Stop(cleanupContext)
if isExpectedHeldTerminalClose(err) {
return nil
}
return err
}}); err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
if err := stack.Push(heldCleanupAction{name: "await terminal stream release", cleanup: func(cleanupContext context.Context) error {
graceContext, cancelGrace := context.WithTimeout(cleanupContext, heldTerminalGracePeriod)
defer cancelGrace()
return pump.Wait(graceContext)
}}); err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
if err := connection.WriteFrame(ctx, resize); err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
command := heldTerminalCommand(input.Plan.ID.String())
proof := newHeldTerminalProof(input.Plan.ID.String())
if err := writeHeldTerminalCommandAfterFirstPumpFrame(ctx, connection, pump, proof, command); err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
if err := stack.Push(heldCleanupAction{name: "release held terminal command", cleanup: func(cleanupContext context.Context) error {
return connection.WriteFrame(cleanupContext, client.Frame{Type: client.FrameText, Payload: []byte("\n")})
}}); err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
for {
select {
case frame, ok := <-pump.Events():
if !ok {
return nil, rollbackHeldTerminal(ctx, stack, errors.Join(heldTerminalProofDiagnostics(proof, pump.Err()), proof.Failure(), ErrHeldTerminalProtocol))
}
proof.Consume(frame)
if proof.Complete() {
if err := capability.waitExpectation(ctx, stateClient, baseline, false); err != nil {
return nil, rollbackHeldTerminal(ctx, stack, err)
}
if markErr := lifecycle.markLive(nil); markErr != nil {
return nil, rollbackHeldTerminal(ctx, stack, markErr)
}
return &heldTerminalSession{lifecycle: lifecycle, stack: stack, connection: connection, pump: pump, protocol: true}, nil
}
case <-ctx.Done():
return nil, rollbackHeldTerminal(ctx, stack, errors.Join(heldTerminalProofDiagnostics(proof, pump.Err()), fmt.Errorf("terminal proof timeout: %w", ctx.Err())))
}
}
}
func isExpectedHeldTerminalClose(err error) bool {
var closeErr *client.WebSocketCloseError
return errors.As(err, &closeErr) && (closeErr.Code == 1000 || closeErr.Code == 1006)
}
func heldTerminalProofDiagnostics(proof *heldTerminalProof, pumpErr error) error {
return fmt.Errorf("terminal proof ended: frames=%d bytes=%d first_frame=%s marker=%t rows=%d cols=%d pump_error=%v", proof.FrameCount(), proof.ByteCount(), proof.FirstFrameType(), hasHeldTerminalMarker(proof.buffer, proof.marker), proof.rows, proof.cols, pumpErr)
}
func writeHeldTerminalCommandAfterFirstPumpFrame(ctx context.Context, connection heldTerminalConnection, pump heldTerminalPump, proof *heldTerminalProof, command string) error {
select {
case frame, ok := <-pump.Events():
if !ok {
return errors.Join(pump.Err(), proof.Failure(), ErrHeldTerminalProtocol)
}
proof.Consume(frame)
case <-ctx.Done():
return ctx.Err()
}
return connection.WriteFrame(ctx, client.Frame{Type: client.FrameText, Payload: []byte(command)})
}
func validateHeldTerminalInput(ctx context.Context, input heldTerminalInput) error {
if ctx == nil || input.Dashboard == nil || input.PATClient == nil || input.Agent == nil || input.Readiness.ServerID == 0 || input.Readiness.UUID == "" || input.Readiness.UUID != input.Agent.UUID() || input.Plan.Kind != StressSessionTerminal || input.Plan.ID.String() == "" || input.Plan.Ordinal < 1 || input.Plan.Agent.Int() < 1 {
return ErrInvalidHeldTerminalInput
}
return nil
}
func validateHeldPATClient(clientInstance *client.Client) error {
if clientInstance == nil {
return ErrInvalidHeldPATClient
}
return nil
}
func validateHeldTerminalResponse(response terminalCreateResponse, serverID uint64) error {
if response.SessionID == "" || response.ServerID != serverID {
return fmt.Errorf("created terminal identity is incomplete_or_wrong_server: %w", ErrHeldTerminalProtocol)
}
return nil
}
func rollbackHeldTerminal(ctx context.Context, stack *heldCleanupStack, original error) error {
rollbackContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), heldTerminalCleanupTimeout)
defer cancel()
return errors.Join(original, stack.Run(rollbackContext))
}
func (session *heldTerminalSession) Plan() StressSessionPlan { return session.lifecycle.Plan() }
func (session *heldTerminalSession) WaitLive(ctx context.Context) error {
return session.lifecycle.WaitLive(ctx)
}
func (session *heldTerminalSession) IOStreamID() (string, bool) {
return session.lifecycle.IOStreamID()
}
func (session *heldTerminalSession) ProtocolProved() bool { return session.protocol }
func (session *heldTerminalSession) WaitClosed(ctx context.Context) error {
return session.lifecycle.WaitClosed(ctx)
}
func (session *heldTerminalSession) Done() <-chan struct{} { return session.lifecycle.Done() }
func (session *heldTerminalSession) CloseResult() error { return session.lifecycle.CloseResult() }
func (session *heldTerminalSession) Close(ctx context.Context) error {
owner, won := session.lifecycle.beginClose()
if !won {
return session.lifecycle.WaitClosed(ctx)
}
go func() {
cleanupContext, cancel := owner.cleanupContext()
cleanupErr := session.stack.Run(cleanupContext)
if cleanupContext.Err() != nil {
cleanupErr = errors.Join(cleanupErr, cleanupContext.Err())
}
cancel()
owner.markClosed(cleanupErr)
}()
return session.lifecycle.WaitClosed(ctx)
}
func newHeldTerminalSessionForTest(lifecycle *heldSessionLifecycle, stack *heldCleanupStack) *heldTerminalSession {
return &heldTerminalSession{lifecycle: lifecycle, stack: stack}
}
var _ heldSession = (*heldTerminalSession)(nil)
@@ -0,0 +1,78 @@
//go:build linux
package scenario
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
type heldTerminalOrderConnection struct {
writes chan client.Frame
gate <-chan struct{}
}
func (connection *heldTerminalOrderConnection) WriteFrame(_ context.Context, frame client.Frame) error {
if frame.Type == client.FrameText {
<-connection.gate
}
connection.writes <- frame
return nil
}
func (connection *heldTerminalOrderConnection) Close() error { return nil }
type heldTerminalOrderPump struct {
events chan client.Frame
}
func (pump *heldTerminalOrderPump) Events() <-chan client.Frame { return pump.events }
func (pump *heldTerminalOrderPump) Done() <-chan struct{} { return nil }
func (pump *heldTerminalOrderPump) Err() error { return nil }
func (pump *heldTerminalOrderPump) Stop(context.Context) error { return nil }
func (pump *heldTerminalOrderPump) Wait(context.Context) error { return nil }
func TestHeldTerminalCommandWaitsForFirstPumpFrame(t *testing.T) {
// Given
firstFrame := make(chan client.Frame)
commandGate := make(chan struct{})
connection := &heldTerminalOrderConnection{writes: make(chan client.Frame, 1), gate: commandGate}
pump := &heldTerminalOrderPump{events: firstFrame}
proof := newHeldTerminalProof("marker")
// When
commandWritten := make(chan error, 1)
go func() {
commandWritten <- writeHeldTerminalCommandAfterFirstPumpFrame(context.Background(), connection, pump, proof, heldTerminalCommand("marker"))
}()
firstFrame <- client.Frame{Type: client.FrameText, Payload: []byte("first PTY output")}
select {
case err := <-commandWritten:
t.Fatalf("command completed before the first frame barrier was released: %v", err)
default:
}
close(commandGate)
command := <-connection.writes
// Then
require.Equal(t, client.FrameText, command.Type)
require.NoError(t, <-commandWritten)
}
func TestTerminalWireFormatsRemainAgentCompatible(t *testing.T) {
// Given
resize := mustTerminalResizeFrame(132, 43)
command := heldTerminalCommand("marker")
// When
// Then
require.Equal(t, client.FrameBinary, resize.Type)
require.Equal(t, byte(1), resize.Payload[0])
require.JSONEq(t, `{"Cols":132,"Rows":43}`, string(resize.Payload[1:]))
require.Equal(t, client.FrameText, client.Frame{Type: client.FrameText, Payload: []byte(command)}.Type)
require.NotEqual(t, byte(0), []byte(command)[0])
}
@@ -0,0 +1,114 @@
//go:build linux
package scenario
import (
"bytes"
"fmt"
"strconv"
"strings"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
const heldTerminalProofLimit = 64 << 10
type heldTerminalProof struct {
marker string
buffer []byte
rows uint32
cols uint32
framed bool
closed error
frames uint64
bytes uint64
first client.FrameType
}
func newHeldTerminalProof(marker string) *heldTerminalProof {
return &heldTerminalProof{marker: marker}
}
func (proof *heldTerminalProof) Consume(frame client.Frame) {
if proof.Complete() || proof.closed != nil {
return
}
proof.frames++
proof.bytes += uint64(len(frame.Payload))
if proof.first == "" {
proof.first = frame.Type
}
proof.buffer = append(proof.buffer, frame.Payload...)
if len(proof.buffer) > heldTerminalProofLimit {
proof.buffer = proof.buffer[len(proof.buffer)-heldTerminalProofLimit:]
}
proof.consumeRecord()
}
func (proof *heldTerminalProof) consumeRecord() {
for {
start := bytes.IndexByte(proof.buffer, 0x1e)
if start < 0 {
proof.buffer = nil
return
}
proof.buffer = proof.buffer[start:]
endOffset := bytes.IndexByte(proof.buffer[1:], 0x1f)
if endOffset < 0 {
return
}
end := 1 + endOffset
record := string(proof.buffer[1:end])
proof.buffer = proof.buffer[end+1:]
recordMarker, sizeText, ok := strings.Cut(record, "|")
if !ok || recordMarker != proof.marker {
continue
}
values := strings.Fields(sizeText)
if len(values) != 2 {
continue
}
rows, rowsErr := strconv.ParseUint(values[0], 10, 32)
cols, colsErr := strconv.ParseUint(values[1], 10, 32)
if rowsErr != nil || colsErr != nil || rows == 0 || cols == 0 {
continue
}
proof.rows, proof.cols, proof.framed = uint32(rows), uint32(cols), true
return
}
}
func (proof *heldTerminalProof) Closed(err error) error {
proof.closed = err
return fmt.Errorf("terminal proof closed: %w", err)
}
func (proof *heldTerminalProof) Failure() error {
if proof.closed == nil {
return nil
}
return fmt.Errorf("terminal proof closed: %w", proof.closed)
}
func (proof *heldTerminalProof) Complete() bool {
return proof.framed && proof.rows == 43 && proof.cols == 132
}
func (proof *heldTerminalProof) Rows() uint32 { return proof.rows }
func (proof *heldTerminalProof) Columns() uint32 { return proof.cols }
func (proof *heldTerminalProof) FrameCount() uint64 { return proof.frames }
func (proof *heldTerminalProof) ByteCount() uint64 { return proof.bytes }
func (proof *heldTerminalProof) FirstFrameType() client.FrameType { return proof.first }
func heldTerminalCommand(marker string) string {
part := len(marker) / 2
return fmt.Sprintf("printf '\\036%%s%%s|' '%s' '%s'; stty size; printf '\\037'; read -r held_terminal_release; exit\n", marker[:part], marker[part:])
}
func hasHeldTerminalMarker(output []byte, marker string) bool {
return strings.Contains(string(output), marker)
}
@@ -0,0 +1,56 @@
//go:build linux && agentcompat
package scenario
import (
"context"
"fmt"
"os"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
)
func TestHeldTerminalSessionUsesExistingDashboardAndAgent(t *testing.T) {
requireHeldRealSources(t)
paths, err := contract.NewPaths(os.Getenv("AGENTCOMPAT_NEZHA_SOURCE"), os.Getenv("AGENTCOMPAT_AGENT_SOURCE"), t.TempDir())
require.NoError(t, err)
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Minute)
defer cancel()
realFixture, readiness, patClient := startHeldRealFixture(t, ctx, paths, "00000000-0000-0000-0000-000000000118", "held-terminal")
dashboardInstance, agentInstance := realFixture.dashboard, realFixture.agent
t.Cleanup(func() { _, _ = realFixture.Close(context.Background(), false, false, false) })
plan := heldTestPlan(t)
plan.ID, err = NewStressSessionID(fmt.Sprintf("held-terminal-real-%d", time.Now().UnixNano()))
require.NoError(t, err)
baseline, err := patClient.IOStreamState(ctx)
require.NoError(t, err)
dashboardPID, agentPID := dashboardInstance.PID(), agentInstance.PID()
sessionCtx, sessionCancel := context.WithTimeout(ctx, 45*time.Second)
defer sessionCancel()
session, err := newHeldTerminalSession(sessionCtx, heldTerminalInput{Dashboard: dashboardInstance, PATClient: patClient, Agent: agentInstance, Readiness: readiness, Plan: plan})
require.NoError(t, err)
require.NoError(t, session.WaitLive(ctx))
streamID, present := session.IOStreamID()
require.True(t, present)
live, err := patClient.WaitForIOStreamState(ctx, client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(baseline.Count + 1), PresentStreamID: streamID})
require.NoError(t, err)
require.Equal(t, baseline.Count+1, live.Count)
require.Equal(t, dashboardPID, dashboardInstance.PID())
require.Equal(t, agentPID, agentInstance.PID())
require.NoError(t, session.Close(ctx))
require.NoError(t, session.WaitClosed(ctx))
closed, err := patClient.WaitForIOStreamState(ctx, client.IOStreamStateExpectation{ExpectedCount: client.ExpectedIOStreamCount(baseline.Count), AbsentStreamID: streamID})
require.NoError(t, err)
require.Equal(t, baseline.Count, closed.Count)
cleanup, cleanupErr := realFixture.Close(ctx, true, closed.Count == baseline.Count, true)
require.NoError(t, cleanupErr)
require.True(t, heldRealCleanupOK(cleanup))
require.NoError(t, writeHeldRealEvidence("terminal", heldRealEvidence{Kind: "terminal", BaselineCount: baseline.Count, LiveCount: live.Count, ClosedCount: closed.Count, ExactIDPresent: true, ExactIDAbsent: true, ProtocolProved: session.ProtocolProved(), DashboardPIDUnchanged: dashboardPID == dashboardInstance.PID(), AgentPIDUnchanged: agentPID == agentInstance.PID(), CleanupOK: heldRealCleanupOK(cleanup)}))
}
@@ -0,0 +1,257 @@
//go:build linux
package scenario
import (
"context"
"errors"
"testing"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
func TestHeldTerminalCommandSubmitsWithLineFeed(t *testing.T) {
command := heldTerminalCommand("marker-session")
require.Equal(t, byte('\n'), command[len(command)-1])
require.NotContains(t, command, "exit\\n")
}
func TestHeldTerminalProofRejectsEchoOnlyExactTokens(t *testing.T) {
proof := newHeldTerminalProof("marker-session")
proof.Consume(client.Frame{Type: client.FrameText, Payload: []byte("printf 'compat-size='; stty size; printf 'marker-session\\n'; read -r held_terminal_release; exit\\n\r\ncompat-size=43 132\r\nmarker-session\r\n")})
require.False(t, proof.Complete())
}
func TestHeldTerminalProofAcceptsMarkerAndExactSizeAcrossFrames(t *testing.T) {
// Given
proof := newHeldTerminalProof("marker-session")
// When
proof.Consume(client.Frame{Type: client.FrameText, Payload: []byte("prefix\r\n\x1emarker-session|43 ")})
proof.Consume(client.Frame{Type: client.FrameBinary, Payload: []byte("132\x1f\r\n")})
// Then
require.True(t, proof.Complete())
require.Equal(t, uint32(43), proof.Rows())
require.Equal(t, uint32(132), proof.Columns())
}
func TestHeldTerminalProofIgnoresMarkerInEchoedCommand(t *testing.T) {
// Given
proof := newHeldTerminalProof("marker-session")
// When
proof.Consume(client.Frame{Type: client.FrameText, Payload: []byte("printf '\\036%s%s|' 'marker-' 'session'; stty size; printf '\\037'\r\n")})
proof.Consume(client.Frame{Type: client.FrameText, Payload: []byte("\x1emarker-session|43 132\x1f\r\n")})
// Then
require.True(t, proof.Complete())
require.Equal(t, uint32(43), proof.Rows())
require.Equal(t, uint32(132), proof.Columns())
}
func TestHeldTerminalProofRejectsWrongMarkerAndWrongSize(t *testing.T) {
// Given
wrongMarker := newHeldTerminalProof("marker-session")
wrongSize := newHeldTerminalProof("marker-session")
// When
wrongMarker.Consume(client.Frame{Type: client.FrameText, Payload: []byte("\x1emarker-other|43 132\x1f\r\n")})
wrongSize.Consume(client.Frame{Type: client.FrameText, Payload: []byte("\x1emarker-session|42 132\x1f\r\n")})
// Then
require.False(t, wrongMarker.Complete())
require.False(t, wrongSize.Complete())
}
func TestHeldTerminalProofAcceptsValidRecordAfterWrongMarkerInSameFrame(t *testing.T) {
// Given
proof := newHeldTerminalProof("marker-session")
frame := []byte("\x1emarker-other|43 132\x1f\x1emarker-session|43 132\x1f")
// When
proof.Consume(client.Frame{Type: client.FrameText, Payload: frame})
// Then
require.True(t, proof.Complete())
}
func TestHeldTerminalProofAcceptsValidRecordAfterMalformedRecordInSameFrame(t *testing.T) {
// Given
proof := newHeldTerminalProof("marker-session")
frame := []byte("\x1emarker-session|43 nope\x1f\x1emarker-session|43 132\x1f")
// When
proof.Consume(client.Frame{Type: client.FrameBinary, Payload: frame})
// Then
require.True(t, proof.Complete())
}
func TestHeldTerminalProofScansMultipleInvalidRecordsBeforeValidRecord(t *testing.T) {
// Given
proof := newHeldTerminalProof("marker-session")
frame := []byte("noise\x1emarker-other|43 132\x1f\x1emarker-session|43 nope\x1f\x1emarker-session|43 132\x1f")
// When
proof.Consume(client.Frame{Type: client.FrameText, Payload: frame})
// Then
require.True(t, proof.Complete())
}
func TestHeldTerminalProofAcceptsFramedRecordAcrossFrames(t *testing.T) {
proof := newHeldTerminalProof("marker-session")
proof.Consume(client.Frame{Type: client.FrameBinary, Payload: []byte("prefix\x1emarker-session|43 ")})
proof.Consume(client.Frame{Type: client.FrameBinary, Payload: []byte("132\x1f\r\n")})
require.True(t, proof.Complete())
}
func TestHeldTerminalProofRejectsMalformedOrUnclosedRecord(t *testing.T) {
malformed := newHeldTerminalProof("marker-session")
unclosed := newHeldTerminalProof("marker-session")
malformed.Consume(client.Frame{Type: client.FrameBinary, Payload: []byte("\x1emarker-session|43 nope\x1f")})
unclosed.Consume(client.Frame{Type: client.FrameBinary, Payload: []byte("\x1emarker-session|43 132")})
require.False(t, malformed.Complete())
require.False(t, unclosed.Complete())
require.Equal(t, []byte("\x1emarker-session|43 132"), unclosed.buffer)
}
func TestHeldTerminalProofAcceptsEchoThenFramedRecord(t *testing.T) {
proof := newHeldTerminalProof("marker-session")
proof.Consume(client.Frame{Type: client.FrameText, Payload: []byte("printf '\\036%s%s|' 'marker-' 'session'; stty size; printf '\\037' exit\\n\r\n")})
require.False(t, proof.Complete())
proof.Consume(client.Frame{Type: client.FrameBinary, Payload: []byte("\x1emarker-session|43 132\x1f")})
require.True(t, proof.Complete())
}
func TestHeldTerminalProofRejectsClosedPumpBeforeProof(t *testing.T) {
// Given
proof := newHeldTerminalProof("marker-session")
// When
err := proof.Closed(errors.New("pump closed"))
// Then
require.Error(t, err)
require.ErrorContains(t, err, "pump closed")
}
func TestHeldTerminalProofBoundsAccumulator(t *testing.T) {
// Given
proof := newHeldTerminalProof("marker-session")
// When
proof.Consume(client.Frame{Type: client.FrameText, Payload: make([]byte, heldTerminalProofLimit+1)})
// Then
require.LessOrEqual(t, len(proof.buffer), heldTerminalProofLimit)
}
func TestHeldTerminalResponseRequiresExactSessionServerIdentity(t *testing.T) {
// Given
response := terminalCreateResponse{SessionID: "session", ServerID: 9}
// When
err := validateHeldTerminalResponse(response, 8)
// Then
require.ErrorIs(t, err, ErrHeldTerminalProtocol)
require.NoError(t, validateHeldTerminalResponse(response, 9))
}
func TestHeldTerminalInputRejectsMissingResourcesAndMismatchedReadiness(t *testing.T) {
// Given
plan := heldTestPlan(t)
input := heldTerminalInput{Plan: plan, Readiness: agent.Readiness{ServerID: 7, UUID: "agent"}}
// When
err := validateHeldTerminalInput(context.Background(), input)
// Then
require.ErrorIs(t, err, ErrInvalidHeldTerminalInput)
}
func TestHeldTerminalInputRejectsMissingPATClientBeforeRemoteMutation(t *testing.T) {
err := validateHeldPATClient(nil)
require.ErrorIs(t, err, ErrInvalidHeldPATClient)
}
func TestHeldTerminalCommandKeepsShellHeldUntilInput(t *testing.T) {
// Given
command := heldTerminalCommand("marker-session")
// Then
require.Contains(t, command, "marker-")
require.Contains(t, command, "session")
require.Contains(t, command, "stty size")
require.Contains(t, command, "read -r")
require.Contains(t, command, "exit")
require.Contains(t, command, "\\036")
require.Contains(t, command, "\\037")
require.NotEqual(t, "\n", command)
}
func TestHeldTerminalCleanupOwnerDoesNotUseCanceledWaiter(t *testing.T) {
// Given
cleanupStarted := make(chan struct{})
cleanupRelease := make(chan struct{})
lifecycle := heldTestLifecycle(t, "held-terminal-session")
stack := newHeldCleanupStack()
require.NoError(t, stack.Push(heldCleanupAction{name: "blocked", cleanup: func(cleanupContext context.Context) error {
require.NoError(t, cleanupContext.Err())
close(cleanupStarted)
<-cleanupRelease
return nil
}}))
session := newHeldTerminalSessionForTest(lifecycle, stack)
// When
canceled, cancel := context.WithCancel(context.Background())
cancel()
first := make(chan error, 1)
go func() { first <- session.Close(canceled) }()
<-cleanupStarted
// Then
require.ErrorIs(t, <-first, context.Canceled)
close(cleanupRelease)
require.NoError(t, session.Close(context.Background()))
}
func TestHeldTerminalCleanupStackReleasesHoldBeforeTransport(t *testing.T) {
// Given
stack := newHeldCleanupStack()
var order []string
require.NoError(t, stack.Push(heldCleanupAction{name: "absence", cleanup: func(context.Context) error {
order = append(order, "absence")
return nil
}}))
require.NoError(t, stack.Push(heldCleanupAction{name: "transport", cleanup: func(context.Context) error {
order = append(order, "transport")
return nil
}}))
require.NoError(t, stack.Push(heldCleanupAction{name: "release", cleanup: func(context.Context) error {
order = append(order, "release")
return nil
}}))
// When
require.NoError(t, stack.Run(context.Background()))
// Then
require.Equal(t, []string{"release", "transport", "absence"}, order)
}
@@ -0,0 +1,112 @@
//go:build linux
package scenario
import (
"context"
"errors"
"sync"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
var (
ErrInvalidHeldFramePump = errors.New("held WebSocket frame pump is invalid")
ErrHeldFrameBufferFull = errors.New("held WebSocket frame buffer is full")
)
type heldWebSocketPump struct {
connection *client.WebSocketConnection
ctx context.Context
cancel context.CancelFunc
events chan client.Frame
done chan struct{}
stopOnce sync.Once
stopDone chan struct{}
stopResult error
terminalMu sync.RWMutex
terminal error
}
func newHeldWebSocketPump(parent context.Context, connection *client.WebSocketConnection, capacity int) (*heldWebSocketPump, error) {
if parent == nil || connection == nil || capacity < 1 {
return nil, ErrInvalidHeldFramePump
}
pumpContext, cancel := context.WithCancel(parent)
pump := &heldWebSocketPump{connection: connection, ctx: pumpContext, cancel: cancel, events: make(chan client.Frame, capacity), done: make(chan struct{}), stopDone: make(chan struct{})}
go pump.readFrames()
return pump, nil
}
func (pump *heldWebSocketPump) Events() <-chan client.Frame { return pump.events }
func (pump *heldWebSocketPump) Done() <-chan struct{} { return pump.done }
func (pump *heldWebSocketPump) Err() error {
pump.terminalMu.RLock()
defer pump.terminalMu.RUnlock()
return pump.terminal
}
func (pump *heldWebSocketPump) readFrames() {
defer close(pump.events)
defer close(pump.done)
for {
frame, err := pump.connection.ReadFrameUntil(pump.ctx)
if err != nil {
if pump.ctx.Err() == nil {
pump.setTerminal(err)
}
return
}
select {
case pump.events <- frame:
case <-pump.ctx.Done():
return
default:
pump.setTerminal(ErrHeldFrameBufferFull)
pump.cancel()
_ = pump.connection.Close()
return
}
}
}
func (pump *heldWebSocketPump) setTerminal(err error) {
pump.terminalMu.Lock()
defer pump.terminalMu.Unlock()
if pump.terminal == nil {
pump.terminal = err
}
}
func (pump *heldWebSocketPump) Stop(ctx context.Context) error {
pump.stopOnce.Do(func() {
// The shutdown owner must outlive any individual caller's wait context.
go pump.stop()
})
select {
case <-pump.stopDone:
return pump.stopResult
case <-ctx.Done():
return ctx.Err()
}
}
func (pump *heldWebSocketPump) stop() {
pump.cancel()
closeErr := pump.connection.Close()
<-pump.done
// Publish only after the reader is joined so every waiter observes the same complete result.
pump.stopResult = errors.Join(pump.Err(), closeErr)
close(pump.stopDone)
}
func (pump *heldWebSocketPump) Wait(ctx context.Context) error {
select {
case <-pump.done:
return nil
case <-ctx.Done():
return ctx.Err()
}
}
@@ -0,0 +1,160 @@
//go:build linux
package scenario
import (
"context"
"errors"
"net"
"net/http/httptest"
"sync"
"testing"
"time"
"github.com/gorilla/websocket"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
func TestHeldWebSocketPumpStopRetainsPeerAndCloseErrorsAfterCanceledWaiter(t *testing.T) {
// Given
closeErr := errors.New("physical close failed")
server := heldPumpServer(t, func(connection *websocket.Conn) {
require.NoError(t, connection.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.ClosePolicyViolation, "peer shutdown")))
})
recordedConnection := newPumpRecordingConn(closeErr)
recordedConnection.closeEntered = make(chan struct{})
recordedConnection.allowClose = make(chan struct{})
t.Cleanup(func() { recordedConnection.releaseClose() })
connection := heldPumpConnectionWithConn(t, server, recordedConnection)
pump, err := newHeldWebSocketPump(context.Background(), connection, 1)
require.NoError(t, err)
require.NoError(t, pump.Wait(context.Background()))
peerErr := pump.Err()
require.Error(t, peerErr)
stopContext, cancel := context.WithCancel(context.Background())
cancel()
// When
firstStopErr := pump.Stop(stopContext)
recordedConnection.awaitClose(t)
recordedConnection.releaseClose()
laterStopErr := pump.Stop(context.Background())
// Then
require.ErrorIs(t, firstStopErr, context.Canceled)
require.ErrorIs(t, laterStopErr, peerErr)
require.ErrorIs(t, laterStopErr, closeErr)
require.Equal(t, 1, recordedConnection.closeCount())
require.Equal(t, laterStopErr, pump.Stop(context.Background()))
}
func TestHeldWebSocketPumpStopConcurrentCallersReceiveRetainedResult(t *testing.T) {
// Given
closeErr := errors.New("physical close failed")
serverReady := make(chan struct{})
server := heldPumpServer(t, func(connection *websocket.Conn) {
close(serverReady)
_, _, _ = connection.ReadMessage()
})
recordedConnection := newPumpRecordingConn(closeErr)
connection := heldPumpConnectionWithConn(t, server, recordedConnection)
pump, err := newHeldWebSocketPump(context.Background(), connection, 1)
require.NoError(t, err)
<-serverReady
// When
results := make(chan error, 8)
for range cap(results) {
go func() { results <- pump.Stop(context.Background()) }()
}
// Then
var retainedResult error
for range cap(results) {
result := <-results
require.ErrorIs(t, result, closeErr)
if retainedResult == nil {
retainedResult = result
continue
}
require.Equal(t, retainedResult, result)
}
require.Equal(t, 1, recordedConnection.closeCount())
require.Equal(t, retainedResult, pump.Stop(context.Background()))
select {
case <-pump.Done():
default:
t.Fatal("Stop returned before the reader joined")
}
}
type pumpRecordingConn struct {
net.Conn
mu sync.Mutex
closeErr error
closeCountValue int
closeEntered chan struct{}
allowClose chan struct{}
releaseOnce sync.Once
}
func newPumpRecordingConn(closeErr error) *pumpRecordingConn {
return &pumpRecordingConn{closeErr: closeErr}
}
func (connection *pumpRecordingConn) Close() error {
connection.mu.Lock()
connection.closeCountValue++
connection.mu.Unlock()
underlyingErr := connection.Conn.Close()
if connection.closeEntered != nil {
close(connection.closeEntered)
<-connection.allowClose
}
if connection.closeErr != nil {
return connection.closeErr
}
return underlyingErr
}
func (connection *pumpRecordingConn) awaitClose(t *testing.T) {
t.Helper()
select {
case <-connection.closeEntered:
case <-time.After(time.Second):
t.Fatal("pump owner did not start physical close")
}
}
func (connection *pumpRecordingConn) releaseClose() {
if connection.allowClose != nil {
connection.releaseOnce.Do(func() { close(connection.allowClose) })
}
}
func (connection *pumpRecordingConn) closeCount() int {
connection.mu.Lock()
defer connection.mu.Unlock()
return connection.closeCountValue
}
func heldPumpConnectionWithConn(t *testing.T, server *httptest.Server, recordedConnection *pumpRecordingConn) *client.WebSocketConnection {
t.Helper()
dialer := *websocket.DefaultDialer
dialer.NetDialContext = func(ctx context.Context, network, address string) (net.Conn, error) {
connection, err := (&net.Dialer{}).DialContext(ctx, network, address)
if err != nil {
return nil, err
}
recordedConnection.Conn = connection
return recordedConnection, nil
}
httpClient, err := client.New(client.Config{BaseURL: server.URL, RequestTimeout: time.Second, MaxResponseBytes: 1024, WebSocketDialer: &dialer})
require.NoError(t, err)
connection, err := httpClient.DialWebSocket(context.Background(), "/held")
require.NoError(t, err)
t.Cleanup(func() { _ = connection.Close() })
return connection
}
@@ -0,0 +1,156 @@
//go:build linux
package scenario
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gorilla/websocket"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
func TestHeldWebSocketPumpPreservesFrameOrderAndType(t *testing.T) {
serverReady := make(chan struct{})
serverRelease := make(chan struct{})
server := heldPumpServer(t, func(connection *websocket.Conn) {
close(serverReady)
require.NoError(t, connection.WriteMessage(websocket.TextMessage, []byte("text")))
require.NoError(t, connection.WriteMessage(websocket.BinaryMessage, []byte{1, 2}))
<-serverRelease
})
connection := heldPumpConnection(t, server)
pump, err := newHeldWebSocketPump(context.Background(), connection, 2)
require.NoError(t, err)
<-serverReady
select {
case first := <-pump.Events():
require.Equal(t, client.FrameText, first.Type)
require.Equal(t, []byte("text"), first.Payload)
case <-time.After(time.Second):
t.Fatal("missing text frame")
}
select {
case second := <-pump.Events():
require.Equal(t, client.FrameBinary, second.Type)
require.Equal(t, []byte{1, 2}, second.Payload)
case <-time.After(time.Second):
t.Fatal("missing binary frame")
}
require.NoError(t, pump.Stop(context.Background()))
close(serverRelease)
}
func TestHeldWebSocketPumpParentCancellationJoinsReader(t *testing.T) {
serverReady := make(chan struct{})
serverRelease := make(chan struct{})
server := heldPumpServer(t, func(*websocket.Conn) {
close(serverReady)
<-serverRelease
})
connection := heldPumpConnection(t, server)
parent, cancel := context.WithCancel(context.Background())
pump, err := newHeldWebSocketPump(parent, connection, 1)
require.NoError(t, err)
<-serverReady
cancel()
require.NoError(t, pump.Wait(context.Background()))
require.NoError(t, pump.Err())
_, ok := <-pump.Events()
require.False(t, ok)
require.NoError(t, pump.Stop(context.Background()))
require.NoError(t, pump.Stop(context.Background()))
require.NoError(t, pump.Err())
close(serverRelease)
}
func TestHeldWebSocketPumpStopJoinsBlockedRead(t *testing.T) {
serverReady := make(chan struct{})
serverRelease := make(chan struct{})
server := heldPumpServer(t, func(*websocket.Conn) {
close(serverReady)
<-serverRelease
})
connection := heldPumpConnection(t, server)
pump, err := newHeldWebSocketPump(context.Background(), connection, 1)
require.NoError(t, err)
<-serverReady
stopContext, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
require.NoError(t, pump.Stop(stopContext))
require.NoError(t, pump.Wait(context.Background()))
require.NoError(t, pump.Err())
_, ok := <-pump.Events()
require.False(t, ok)
close(serverRelease)
}
func TestHeldWebSocketPumpBufferFullFailsFast(t *testing.T) {
serverReady := make(chan struct{})
server := heldPumpServer(t, func(connection *websocket.Conn) {
close(serverReady)
for _, payload := range []string{"one", "two"} {
require.NoError(t, connection.WriteMessage(websocket.TextMessage, []byte(payload)))
}
})
connection := heldPumpConnection(t, server)
pump, err := newHeldWebSocketPump(context.Background(), connection, 1)
require.NoError(t, err)
<-serverReady
select {
case <-pump.Done():
case <-time.After(time.Second):
t.Fatal("buffer-full pump did not stop")
}
require.ErrorIs(t, pump.Err(), ErrHeldFrameBufferFull)
require.NoError(t, pump.Wait(context.Background()))
require.ErrorIs(t, pump.Stop(context.Background()), ErrHeldFrameBufferFull)
require.ErrorIs(t, pump.Stop(context.Background()), ErrHeldFrameBufferFull)
}
func TestHeldWebSocketPumpStopIsIdempotentAndRetainsPeerError(t *testing.T) {
server := heldPumpServer(t, func(connection *websocket.Conn) {
require.NoError(t, connection.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "peer")))
})
connection := heldPumpConnection(t, server)
pump, err := newHeldWebSocketPump(context.Background(), connection, 1)
require.NoError(t, err)
if err := pump.Wait(context.Background()); err != nil {
t.Fatal(err)
}
peerErr := pump.Err()
require.NotNil(t, peerErr)
require.ErrorIs(t, pump.Stop(context.Background()), peerErr)
require.ErrorIs(t, pump.Stop(context.Background()), peerErr)
}
func heldPumpServer(t *testing.T, serve func(*websocket.Conn)) *httptest.Server {
upgrader := websocket.Upgrader{}
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
connection, err := upgrader.Upgrade(writer, request, nil)
require.NoError(t, err)
defer connection.Close()
serve(connection)
}))
t.Cleanup(server.Close)
return server
}
func heldPumpConnection(t *testing.T, server *httptest.Server) *client.WebSocketConnection {
httpClient := newTestClient(t, server.URL)
connection, err := httpClient.DialWebSocket(context.Background(), "/held")
require.NoError(t, err)
t.Cleanup(func() { _ = connection.Close() })
return connection
}
func newTestClient(t *testing.T, baseURL string) *client.Client {
result, err := client.New(client.Config{BaseURL: baseURL, RequestTimeout: time.Second, MaxResponseBytes: 1024})
require.NoError(t, err)
return result
}
@@ -0,0 +1,267 @@
//go:build linux
package scenario
import (
"bytes"
"context"
"crypto/sha256"
"errors"
"fmt"
"os"
"path/filepath"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process"
)
type LegacyFMInput struct {
Paths contract.Paths
Fault contract.Fault
}
type LegacyFM struct{}
func (LegacyFM) Run(ctx context.Context, input LegacyFMInput) (result Result, runErr error) {
assertions := NewAssertionSet()
runID, err := newLegacyFMRunID()
if err != nil {
return Result{Name: "legacy-fm", Assertions: assertions.Results(), Error: err.Error()}, err
}
dashboardInstance, err := dashboard.Start(ctx, dashboard.StartConfig{SourceDir: input.Paths.NezhaSource().String(), ReceiptGate: true})
if err != nil {
return Result{Name: "legacy-fm", Assertions: assertions.Results(), Error: err.Error()}, err
}
result.CleanupOK = true
defer func() {
cleanupErr := dashboardInstance.Stop(context.Background())
result.CleanupOK = result.CleanupOK && cleanupErr == nil && dashboardInstance.CleanupReceipt().Passed
if cleanupErr != nil && runErr == nil {
runErr = cleanupErr
}
result.Passed = runErr == nil
result.Error = errorText(runErr)
}()
secret := dashboardInstance.AgentSecret()
agentInstance, err := agent.Start(ctx, agent.AgentStartConfig{
SourceDir: input.Paths.AgentSource().String(), Endpoint: dashboardInstance.Endpoint(),
Secret: secret, UUID: "00000000-0000-0000-0000-000000000115",
FMObserverRunID: runID,
})
if err != nil {
return Result{Name: "legacy-fm", Assertions: assertions.Results(), CleanupOK: true, Error: err.Error()}, err
}
defer func() {
cleanupErr := agentInstance.Stop(context.Background())
result.CleanupOK = result.CleanupOK && cleanupErr == nil && agentInstance.CleanupReceipt().Passed
if cleanupErr != nil && runErr == nil {
runErr = cleanupErr
}
}()
if input.Fault.String() == "agent-bad-secret" {
return finishLegacyFM(assertions, errors.New("fault injection agent-bad-secret"))
}
if err := dashboardInstance.WaitForReceiptAccepted(ctx); err != nil {
return finishLegacyFM(assertions, err)
}
if err := dashboardInstance.ReleaseReceipt(ctx); err != nil {
return finishLegacyFM(assertions, err)
}
if _, err := agentInstance.WaitReady(ctx, dashboardInstance); err != nil {
return finishLegacyFM(assertions, err)
}
serverID, err := findLegacyFMServerID(ctx, dashboardInstance, agentInstance.UUID())
if err != nil {
return finishLegacyFM(assertions, err)
}
root, err := fixture.NewAgentRoot(agentInstance.WorkspaceRoot(), "fm-files")
if err != nil {
return finishLegacyFM(assertions, err)
}
listPath, err := root.Path("legacy/list")
if err != nil {
return finishLegacyFM(assertions, err)
}
uploadPath, err := root.Path("legacy/upload.bin")
if err != nil {
return finishLegacyFM(assertions, err)
}
downloadPath, err := root.Path("legacy/download.bin")
if err != nil {
return finishLegacyFM(assertions, err)
}
payloadPattern := []byte("legacy-fm-exact-payload\x00\xff")
payload := bytes.Repeat(payloadPattern, (1<<20+257)/len(payloadPattern)+1)
payload = payload[:1<<20+257]
sentinel := []byte("outside-fm-root-sentinel")
sentinelPaths := []string{
filepath.Join(agentInstance.WorkspaceRoot(), "outside-fm-root-a.txt"),
filepath.Join(agentInstance.WorkspaceRoot(), "outside-fm-root-b.txt"),
}
for _, path := range sentinelPaths {
if err := os.WriteFile(path, sentinel, 0o600); err != nil {
return finishLegacyFM(assertions, err)
}
}
if err := os.WriteFile(downloadPath.String(), payload, 0o600); err != nil {
return finishLegacyFM(assertions, err)
}
if err := os.Mkdir(listPath.String(), 0o700); err != nil {
return finishLegacyFM(assertions, err)
}
if err := os.WriteFile(filepath.Join(listPath.String(), "entry.txt"), []byte("entry"), 0o600); err != nil {
return finishLegacyFM(assertions, err)
}
pathRejectionDispatches, err := probeLegacyFMRejectedPathDispatches(ctx, root)
assertions.Record("rejected paths dispatch zero FM frames", err == nil && pathRejectionDispatches == 0, fmt.Sprintf("path_rejections_dispatched: %d", pathRejectionDispatches))
if err != nil || pathRejectionDispatches != 0 {
return finishLegacyFM(assertions, errors.Join(err, errors.New("rejected FM path dispatched a frame")))
}
baselineSample, err := processharness.SampleProcess(agentInstance.PID())
if err != nil {
return finishLegacyFM(assertions, err)
}
admin := dashboardInstance.Clients().REST
session, err := createLegacyFMSession(ctx, admin, serverID)
if err != nil {
return finishLegacyFM(assertions, err)
}
ws, err := dashboardInstance.Clients().WebSocket.DialWebSocket(ctx, "/api/v1/ws/file/"+session)
if err != nil {
return finishLegacyFM(assertions, err)
}
defer ws.Close()
dispatcher := legacyFMCommandDispatcher{writer: ws, root: root}
if err := dispatcher.list(ctx, "legacy/list"); err != nil {
return finishLegacyFM(assertions, err)
}
listFrame, err := readBinaryFrame(ctx, ws)
if err != nil {
return finishLegacyFM(assertions, err)
}
parsedList, err := parseLegacyFMList(listFrame)
assertions.Record("list uses NZFN and exact entry", err == nil && parsedList.Path == listPath.String() && len(parsedList.Entries) == 1 && parsedList.Entries[0].Name == "entry.txt" && !parsedList.Entries[0].Dir, errorText(err))
if err != nil {
return finishLegacyFM(assertions, err)
}
if err := dispatcher.upload(ctx, "legacy/upload.bin", uint64(len(payload))); err != nil {
return finishLegacyFM(assertions, err)
}
if err := ws.WriteFrame(ctx, client.Frame{Type: client.FrameBinary, Payload: payload[:1<<20]}); err != nil {
return finishLegacyFM(assertions, err)
}
if err := ws.WriteFrame(ctx, client.Frame{Type: client.FrameBinary, Payload: payload[1<<20:]}); err != nil {
return finishLegacyFM(assertions, err)
}
completion, err := readBinaryFrame(ctx, ws)
if err != nil {
return finishLegacyFM(assertions, err)
}
if err := requireLegacyFMMarker(completion, "NZUP"); err != nil {
return finishLegacyFM(assertions, err)
}
written, err := os.ReadFile(uploadPath.String())
assertions.Record("upload returns NZUP and exact bytes", err == nil && bytes.Equal(written, payload), errorText(err))
if err != nil || !bytes.Equal(written, payload) {
return finishLegacyFM(assertions, errors.New("uploaded content mismatch"))
}
if err := dispatcher.download(ctx, "legacy/download.bin"); err != nil {
return finishLegacyFM(assertions, err)
}
producerAwaiter := newLegacyFMProducerAwaiter(agentInstance.FMProducerObserver(), runID, agentInstance.UUID(), session)
activeProducerSample, err := producerAwaiter.await(ctx, "active")
if err != nil {
return finishLegacyFM(assertions, err)
}
header, err := readBinaryFrame(ctx, ws)
if err != nil {
return finishLegacyFM(assertions, err)
}
downloadHeader, err := parseLegacyFMDownload(header)
if err != nil {
return finishLegacyFM(assertions, err)
}
if downloadHeader.Size != uint64(len(payload)) {
return finishLegacyFM(assertions, errors.New("download header size mismatch"))
}
downloaded, downloadFrameCount, err := readLegacyFMDownload(ctx, ws, downloadHeader.Size)
if err != nil {
return finishLegacyFM(assertions, err)
}
digest := sha256.Sum256(downloaded)
wantDigest := sha256.Sum256(payload)
assertions.Record("download returns NZTD across frames with exact hash/content", downloadFrameCount >= 2 && bytes.Equal(downloaded, payload) && digest == wantDigest, fmt.Sprintf("size=%d frames=%d", downloadHeader.Size, downloadFrameCount))
if downloadFrameCount < 2 || !bytes.Equal(downloaded, payload) {
return finishLegacyFM(assertions, errors.New("download framing or content mismatch"))
}
if err := dispatcher.download(ctx, "legacy/missing.bin"); err != nil {
return finishLegacyFM(assertions, err)
}
errorFrame, err := readBinaryFrame(ctx, ws)
if err != nil {
return finishLegacyFM(assertions, err)
}
_, nerr := parseLegacyFMDownload(errorFrame)
assertions.Record("missing download returns NERR", errors.Is(nerr, errLegacyFMRemote), errorText(nerr))
scopeChecksErr := verifyLegacyFMMissingScopes(ctx, dashboardInstance, serverID, session)
assertions.Record("each missing scope rejects FM creation and WebSocket", scopeChecksErr == nil, errorText(scopeChecksErr))
if scopeChecksErr != nil {
return finishLegacyFM(assertions, scopeChecksErr)
}
foreign, removeForeignUser, err := createForeignLegacyFMClient(ctx, dashboardInstance)
if err != nil {
return finishLegacyFM(assertions, err)
}
// Temporary users are scenario resources, so deletion failures must fail cleanup evidence.
defer func() {
cleanupErr := removeForeignUser()
result.CleanupOK = result.CleanupOK && cleanupErr == nil
if cleanupErr != nil && runErr == nil {
runErr = cleanupErr
}
}()
_, hijackErr := foreign.DialWebSocket(ctx, "/api/v1/ws/file/"+session)
assertions.Record("foreign PAT cannot hijack FM session", isLegacyFMSessionRejected(hijackErr), errorText(hijackErr))
if err := ws.Close(); err != nil {
return finishLegacyFM(assertions, err)
}
closedProducerSample, err := producerAwaiter.await(ctx, "closed")
if err != nil {
return finishLegacyFM(assertions, err)
}
residueProbe := legacyFMResidueProbe{
assertions: assertions, agentPID: agentInstance.PID(),
session: session, root: root, baseline: baselineSample, sessionClient: dashboardInstance.Clients().WebSocket,
producer: producerAwaiter.observation(activeProducerSample, closedProducerSample),
}
if err := residueProbe.run(ctx); err != nil {
return finishLegacyFM(assertions, err)
}
sentinelErr := verifyLegacyFMSentinels(sentinelPaths, sentinel)
assertions.Record("outside-root sentinels remain unchanged", sentinelErr == nil, errorText(sentinelErr))
if sentinelErr != nil {
return finishLegacyFM(assertions, sentinelErr)
}
filesystem := newMCPFilesystemClient(dashboardInstance.Clients().MCP, serverID, root)
fixtureCleanupErr := cleanupLegacyFMFixtures(ctx, filesystem)
assertions.Record("MCP cleanup removes FM fixture residue", fixtureCleanupErr == nil, errorText(fixtureCleanupErr))
if fixtureCleanupErr != nil {
return finishLegacyFM(assertions, fixtureCleanupErr)
}
return finishLegacyFM(assertions, nil)
}
@@ -0,0 +1,139 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process"
)
type legacyFMResidueProbe struct {
assertions *AssertionSet
agentPID int
session string
root fixture.AgentRoot
baseline processharness.Sample
sessionClient *client.Client
producer legacyFMProducerObservation
}
type legacyFMCountingWriter struct {
frameCount int
}
func (writer *legacyFMCountingWriter) WriteFrame(context.Context, client.Frame) error {
writer.frameCount++
return nil
}
func probeLegacyFMRejectedPathDispatches(ctx context.Context, root fixture.AgentRoot) (int, error) {
symlinkName := "rejected-symlink-parent"
symlinkPath := filepath.Join(root.Absolute(), symlinkName)
if err := os.Symlink(filepath.Dir(root.Absolute()), symlinkPath); err != nil {
return 0, fmt.Errorf("create rejected-path symlink: %w", err)
}
defer os.Remove(symlinkPath)
candidates := []string{
filepath.Join(filepath.Dir(root.Absolute()), "outside-fm-root"),
"../outside-fm-root",
".",
`C:\outside-fm-root`,
`inside\outside`,
symlinkName + "/file",
}
writer := &legacyFMCountingWriter{}
dispatcher := legacyFMCommandDispatcher{writer: writer, root: root}
for _, candidate := range candidates {
operations := []func() error{
func() error { return dispatcher.list(ctx, candidate) },
func() error { return dispatcher.upload(ctx, candidate, 1) },
func() error { return dispatcher.download(ctx, candidate) },
}
for _, operation := range operations {
var pathErr *fixture.AgentPathError
if err := operation(); !errors.As(err, &pathErr) {
return writer.frameCount, fmt.Errorf("rejected FM path %q crossed dispatch boundary: %w", candidate, err)
}
}
}
return writer.frameCount, nil
}
func countLegacyFMFixtureOpenFiles(pid int, root fixture.AgentRoot) (int, error) {
entries, err := os.ReadDir(fmt.Sprintf("/proc/%d/fd", pid))
if err != nil {
return 0, fmt.Errorf("read Agent file descriptors: %w", err)
}
count := 0
rootPath := filepath.Clean(root.Absolute())
for _, entry := range entries {
target, readErr := os.Readlink(filepath.Join("/proc", fmt.Sprint(pid), "fd", entry.Name()))
if readErr != nil {
if errors.Is(readErr, os.ErrNotExist) {
continue
}
return 0, fmt.Errorf("read Agent file descriptor %s: %w", entry.Name(), readErr)
}
target = strings.TrimSuffix(target, " (deleted)")
if target == rootPath || strings.HasPrefix(target, rootPath+string(filepath.Separator)) {
count++
}
}
return count, nil
}
func waitForLegacyFMFixtureOpenFilesClosed(ctx context.Context, pid int, root fixture.AgentRoot) (int, error) {
ticker := time.NewTicker(25 * time.Millisecond)
defer ticker.Stop()
for {
count, err := countLegacyFMFixtureOpenFiles(pid, root)
if err != nil || count == 0 {
return count, err
}
select {
case <-ctx.Done():
return count, ctx.Err()
case <-ticker.C:
}
}
}
func (probe legacyFMResidueProbe) run(ctx context.Context) error {
sessionResidueCount, cleanupErr := waitForLegacyFMSessionCleanup(ctx, probe.sessionClient, probe.session)
probe.assertions.Record("closed FM WebSocket removes session", cleanupErr == nil && sessionResidueCount == 0, fmt.Sprintf("fm_session_residue_count: %d; error=%s", sessionResidueCount, errorText(cleanupErr)))
if cleanupErr != nil {
return cleanupErr
}
producerErr := probe.producer.validate()
probe.assertions.Record("FM producer is active then exits", producerErr == nil, probe.producer.details())
if producerErr != nil {
return producerErr
}
openFileResidueCount, openFileErr := waitForLegacyFMFixtureOpenFilesClosed(ctx, probe.agentPID, probe.root)
probe.assertions.Record("FM closes fixture-root files", openFileErr == nil && openFileResidueCount == 0, fmt.Sprintf("fm_open_file_residue_count: %d", openFileResidueCount))
if openFileErr != nil {
return openFileErr
}
residueSample, err := processharness.SampleProcess(probe.agentPID)
if err != nil {
probe.assertions.Record("Agent process residue has no drift", false, errorText(err))
return err
}
processResidue := legacyFMProcessResidue{Baseline: probe.baseline, End: residueSample}
processErr := processResidue.validate()
probe.assertions.Record("Agent process residue has no drift", processErr == nil, fmt.Sprintf("baseline_non_stdio_fds=%d end_non_stdio_fds=%d baseline_descendants=%d end_descendants=%d baseline_tcp_listeners=%d end_tcp_listeners=%d baseline_tcp6_listeners=%d end_tcp6_listeners=%d", probe.baseline.NonStdioFDCount, residueSample.NonStdioFDCount, probe.baseline.DescendantCount, residueSample.DescendantCount, probe.baseline.TCPListenerCount, residueSample.TCPListenerCount, probe.baseline.TCP6ListenerCount, residueSample.TCP6ListenerCount))
return processErr
}
@@ -0,0 +1,115 @@
//go:build linux
package scenario
import (
"bytes"
"encoding/binary"
"errors"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
var (
errLegacyFMInvalidFrame = errors.New("legacy FM invalid frame")
errLegacyFMUnexpected = errors.New("legacy FM unexpected frame")
)
type legacyFMRemoteError struct{}
func (legacyFMRemoteError) Error() string { return "legacy FM agent error" }
var errLegacyFMRemote = legacyFMRemoteError{}
const (
legacyFMListOp byte = 0x00
legacyFMDownloadOp byte = 0x01
legacyFMUploadOp byte = 0x02
)
type legacyFMEntry struct {
Name string
Dir bool
}
type legacyFMList struct {
Path string
Entries []legacyFMEntry
}
type legacyFMDownloadHeader struct {
Size uint64
}
func buildLegacyFMList(path fixture.AgentPath) []byte {
return append([]byte{legacyFMListOp}, []byte(path.String())...)
}
func buildLegacyFMUpload(path fixture.AgentPath, size uint64) []byte {
frame := make([]byte, 1+8+len(path.String()))
frame[0] = legacyFMUploadOp
binary.BigEndian.PutUint64(frame[1:9], size)
copy(frame[9:], path.String())
return frame
}
func buildLegacyFMDownload(path fixture.AgentPath) []byte {
return append([]byte{legacyFMDownloadOp}, []byte(path.String())...)
}
func parseLegacyFMList(frame []byte) (legacyFMList, error) {
if message, ok := parseLegacyFMError(frame); ok {
return legacyFMList{}, message
}
if len(frame) < 8 || !bytes.Equal(frame[:4], []byte("NZFN")) {
return legacyFMList{}, errLegacyFMInvalidFrame
}
pathSize := binary.BigEndian.Uint32(frame[4:8])
if pathSize == 0 || uint64(pathSize)+8 > uint64(len(frame)) {
return legacyFMList{}, errLegacyFMInvalidFrame
}
pathEnd := 8 + int(pathSize)
result := legacyFMList{Path: string(frame[8:pathEnd])}
for cursor := pathEnd; cursor < len(frame); {
if len(frame)-cursor < 2 {
return legacyFMList{}, errLegacyFMInvalidFrame
}
entrySize := int(frame[cursor+1])
if entrySize == 0 || cursor+2+entrySize > len(frame) || frame[cursor] > 1 {
return legacyFMList{}, errLegacyFMInvalidFrame
}
result.Entries = append(result.Entries, legacyFMEntry{
Name: string(frame[cursor+2 : cursor+2+entrySize]),
Dir: frame[cursor] == 1,
})
cursor += 2 + entrySize
}
return result, nil
}
func parseLegacyFMDownload(frame []byte) (legacyFMDownloadHeader, error) {
if message, ok := parseLegacyFMError(frame); ok {
return legacyFMDownloadHeader{}, message
}
if len(frame) != 12 || !bytes.Equal(frame[:4], []byte("NZTD")) {
return legacyFMDownloadHeader{}, errLegacyFMInvalidFrame
}
return legacyFMDownloadHeader{Size: binary.BigEndian.Uint64(frame[4:12])}, nil
}
func parseLegacyFMError(frame []byte) (error, bool) {
if len(frame) < 4 || !bytes.Equal(frame[:4], []byte("NERR")) {
return nil, false
}
return errLegacyFMRemote, true
}
func requireLegacyFMMarker(frame []byte, marker string) error {
if message, ok := parseLegacyFMError(frame); ok {
return message
}
if !bytes.Equal(frame, []byte(marker)) {
return errLegacyFMUnexpected
}
return nil
}
@@ -0,0 +1,72 @@
//go:build linux
package scenario
import (
"encoding/binary"
"errors"
"testing"
"github.com/stretchr/testify/require"
)
func TestLegacyFMProtocol_ParsersRejectMalformedFrames(t *testing.T) {
tests := []struct {
name string
call func() error
}{
{name: "list magic", call: func() error { _, err := parseLegacyFMList([]byte("bad")); return err }},
{name: "list truncated path", call: func() error { return parseListFrame([]byte{'N', 'Z', 'F', 'N', 0, 0, 0, 4, 'x'}) }},
{name: "list invalid entry type", call: func() error { return parseListFrame([]byte{'N', 'Z', 'F', 'N', 0, 0, 0, 1, 'x', 2, 1, 'a'}) }},
{name: "download short", call: func() error { _, err := parseLegacyFMDownload([]byte("NZTD")); return err }},
{name: "marker mismatch", call: func() error { return requireLegacyFMMarker([]byte("NERR"), "NZUP") }},
{name: "NERR list response", call: func() error { _, err := parseLegacyFMList([]byte("NERRdenied")); return err }},
{name: "NERR download response", call: func() error { _, err := parseLegacyFMDownload([]byte("NERRmissing")); return err }},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if err := test.call(); err == nil {
t.Fatal("malformed frame accepted")
} else if !errors.Is(err, errLegacyFMInvalidFrame) && !errors.Is(err, errLegacyFMUnexpected) && !errors.Is(err, errLegacyFMRemote) {
t.Fatalf("unexpected error = %v", err)
}
})
}
}
func TestLegacyFMProtocol_RedactsRemotePayloadAndParsedListData(t *testing.T) {
sensitivePath := "/workspace/secret-fixture/pat-token"
sensitiveRemote := "NERRremote-pat-token"
listFrame := make([]byte, 8, 8+len(sensitivePath)+2)
copy(listFrame, []byte("NZFN"))
binary.BigEndian.PutUint32(listFrame[4:], uint32(len(sensitivePath)))
listFrame = append(listFrame, []byte(sensitivePath)...)
listFrame = append(listFrame, 0)
listErr := listErrorForTest(listFrame)
remoteErr := listErrorForTest([]byte(sensitiveRemote))
for _, err := range []error{listErr, remoteErr} {
require.Error(t, err)
require.NotContains(t, err.Error(), sensitivePath)
require.NotContains(t, err.Error(), "pat-token")
require.NotContains(t, err.Error(), "NERRremote")
}
}
func TestLegacyFMProtocol_RedactsMarkerMismatch(t *testing.T) {
err := requireLegacyFMMarker([]byte("wrong-secret-marker"), "expected-secret-marker")
require.ErrorIs(t, err, errLegacyFMUnexpected)
require.NotContains(t, err.Error(), "wrong-secret-marker")
require.NotContains(t, err.Error(), "expected-secret-marker")
}
func listErrorForTest(frame []byte) error {
_, err := parseLegacyFMList(frame)
return err
}
func parseListFrame(frame []byte) error {
_, err := parseLegacyFMList(frame)
return err
}
@@ -0,0 +1,225 @@
//go:build linux
package scenario
import (
"bytes"
"context"
"encoding/binary"
"errors"
"os"
"path/filepath"
"testing"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process"
)
type legacyFMRecordingWriter struct {
frames []client.Frame
}
func (writer *legacyFMRecordingWriter) WriteFrame(_ context.Context, frame client.Frame) error {
writer.frames = append(writer.frames, frame)
return nil
}
type legacyFMFilesystemWriter struct {
frames int
}
func (writer *legacyFMFilesystemWriter) WriteFrame(_ context.Context, frame client.Frame) error {
writer.frames++
switch frame.Payload[0] {
case 0:
_, err := os.ReadDir(string(frame.Payload[1:]))
return err
case 1:
_, err := os.ReadFile(string(frame.Payload[1:]))
return err
case 2:
return os.WriteFile(string(frame.Payload[9:]), []byte("uploaded"), 0o600)
default:
return errors.New("unexpected FM operation")
}
}
func TestLegacyFM_BuildersRequireAgentPath(t *testing.T) {
var _ func(fixture.AgentPath) []byte = buildLegacyFMList
var _ func(fixture.AgentPath) []byte = buildLegacyFMDownload
var _ func(fixture.AgentPath, uint64) []byte = buildLegacyFMUpload
root, err := fixture.NewAgentRoot(t.TempDir(), "fm-protocol")
if err != nil {
t.Fatal(err)
}
path, err := root.Path("wire/file.bin")
if err != nil {
t.Fatal(err)
}
if got, want := buildLegacyFMList(path), append([]byte{0x00}, []byte(path.String())...); !bytes.Equal(got, want) {
t.Fatalf("list frame = %x, want %x", got, want)
}
if got, want := buildLegacyFMDownload(path), append([]byte{0x01}, []byte(path.String())...); !bytes.Equal(got, want) {
t.Fatalf("download frame = %x, want %x", got, want)
}
wantUpload := append([]byte{0x02}, make([]byte, 8)...)
binary.BigEndian.PutUint64(wantUpload[1:9], 0x0102030405060708)
wantUpload = append(wantUpload, []byte(path.String())...)
if got := buildLegacyFMUpload(path, 0x0102030405060708); !bytes.Equal(got, wantUpload) {
t.Fatalf("upload frame = %x, want %x", got, wantUpload)
}
}
func TestLegacyFM_RejectedPathDispatchesNoFrame(t *testing.T) {
root, err := fixture.NewAgentRoot(t.TempDir(), "fm-path-boundary")
if err != nil {
t.Fatal(err)
}
symlinkTarget := t.TempDir()
if err := os.Symlink(symlinkTarget, filepath.Join(root.Absolute(), "linked")); err != nil {
t.Fatal(err)
}
tests := []struct {
name string
candidate string
wantReason fixture.PathRejectionReason
}{
{name: "absolute", candidate: filepath.Join(t.TempDir(), "outside"), wantReason: fixture.PathRejectionAbsolute},
{name: "parent", candidate: "../outside", wantReason: fixture.PathRejectionParent},
{name: "destructive root", candidate: ".", wantReason: fixture.PathRejectionDestructiveRoot},
{name: "volume", candidate: `C:\outside`, wantReason: fixture.PathRejectionVolume},
{name: "separator", candidate: `inside\outside`, wantReason: fixture.PathRejectionSeparator},
{name: "symlink parent", candidate: "linked/file", wantReason: fixture.PathRejectionSymlinkParent},
}
operations := []struct {
name string
run func(context.Context, legacyFMCommandDispatcher, string) error
}{
{name: "list", run: func(ctx context.Context, dispatcher legacyFMCommandDispatcher, candidate string) error {
return dispatcher.list(ctx, candidate)
}},
{name: "upload", run: func(ctx context.Context, dispatcher legacyFMCommandDispatcher, candidate string) error {
return dispatcher.upload(ctx, candidate, 1)
}},
{name: "download", run: func(ctx context.Context, dispatcher legacyFMCommandDispatcher, candidate string) error {
return dispatcher.download(ctx, candidate)
}},
}
for _, test := range tests {
for _, operation := range operations {
t.Run(test.name+"/"+operation.name, func(t *testing.T) {
writer := &legacyFMRecordingWriter{}
dispatcher := legacyFMCommandDispatcher{writer: writer, root: root}
pathErr := operation.run(t.Context(), dispatcher, test.candidate)
var agentPathErr *fixture.AgentPathError
if !errors.As(pathErr, &agentPathErr) || agentPathErr.Reason != test.wantReason {
t.Fatalf("rejected path error=%v, want reason %s", pathErr, test.wantReason)
}
if len(writer.frames) != 0 {
t.Fatalf("rejected path dispatched %d frames", len(writer.frames))
}
})
}
}
}
func TestLegacyFM_OutsideRootSentinelUnchanged(t *testing.T) {
parent := t.TempDir()
root, err := fixture.NewAgentRoot(parent, "fm-sentinel")
if err != nil {
t.Fatal(err)
}
sentinelPath := filepath.Join(parent, "outside-sentinel")
symlinkTargetDir := t.TempDir()
symlinkTargetPath := filepath.Join(symlinkTargetDir, "target-sentinel")
want := []byte("unchanged")
for _, path := range []string{sentinelPath, symlinkTargetPath} {
if err := os.WriteFile(path, want, 0o600); err != nil {
t.Fatal(err)
}
}
if err := os.Symlink(symlinkTargetDir, filepath.Join(root.Absolute(), "linked")); err != nil {
t.Fatal(err)
}
if err := os.Mkdir(filepath.Join(root.Absolute(), "inside-dir"), 0o700); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(root.Absolute(), "inside-download"), []byte("fixture"), 0o600); err != nil {
t.Fatal(err)
}
writer := &legacyFMFilesystemWriter{}
dispatcher := legacyFMCommandDispatcher{writer: writer, root: root}
coverage := legacyFMSentinelCoverage{}
coverage.ListRejected = dispatcher.list(t.Context(), "../outside-sentinel") != nil
coverage.UploadRejected = dispatcher.upload(t.Context(), "linked/target-sentinel", uint64(len(want))) != nil
coverage.DownloadRejected = dispatcher.download(t.Context(), sentinelPath) != nil
if writer.frames != 0 {
t.Fatalf("rejected sentinel paths dispatched %d frames", writer.frames)
}
for _, operation := range []func() error{
func() error { return dispatcher.list(t.Context(), "inside-dir") },
func() error { return dispatcher.upload(t.Context(), "inside-upload", 8) },
func() error { return dispatcher.download(t.Context(), "inside-download") },
} {
if err := operation(); err != nil {
t.Fatal(err)
}
coverage.SuccessCount++
}
if err := dispatcher.download(t.Context(), "inside-missing"); err != nil {
coverage.ErrorCount++
}
if err := coverage.validate(); err != nil {
t.Fatal(err)
}
for _, path := range []string{sentinelPath, symlinkTargetPath} {
got, readErr := os.ReadFile(path)
if readErr != nil {
t.Fatal(readErr)
}
if !bytes.Equal(got, want) {
t.Fatalf("outside-root sentinel %s = %q, want %q", path, got, want)
}
}
}
func TestLegacyFM_ProducerObservationRejectsHardcodedZero(t *testing.T) {
observation := legacyFMProducerObservation{
RunID: "run-a", AgentUUID: "agent-a", SessionID: "session-a",
Samples: []agent.FMProducerSample{{RunID: "run-a", AgentUUID: "agent-a", SessionID: "session-a", Phase: "closed", Active: 0}},
}
if err := observation.validate(); err == nil {
t.Fatal("producer observation accepted zero-only samples without a live active producer")
}
}
func TestLegacyFM_SentinelCoverageRejectsNoOp(t *testing.T) {
if err := (legacyFMSentinelCoverage{}).validate(); err == nil {
t.Fatal("sentinel coverage accepted a no-op test without rejected, successful, and failing FM operations")
}
}
func TestLegacyFM_ProcessResidueRejectsFDAndDescendantDrift(t *testing.T) {
tests := []struct {
name string
end processharness.Sample
}{
{name: "fd drift", end: processharness.Sample{NonStdioFDCount: 6, DescendantCount: 2}},
{name: "descendant drift", end: processharness.Sample{NonStdioFDCount: 5, DescendantCount: 3}},
}
baseline := processharness.Sample{NonStdioFDCount: 5, DescendantCount: 2}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if err := (legacyFMProcessResidue{Baseline: baseline, End: test.end}).validate(); err == nil {
t.Fatal("process residue accepted FD or descendant drift")
}
})
}
}
@@ -0,0 +1,248 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"net/http"
"strings"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
type legacyFMUserForm struct {
Role uint8 `json:"role"`
Username string `json:"username"`
Password string `json:"password"`
}
type legacyFMFrameWriter interface {
WriteFrame(context.Context, client.Frame) error
}
type legacyFMCommandDispatcher struct {
writer legacyFMFrameWriter
root fixture.AgentRoot
}
func (dispatcher legacyFMCommandDispatcher) list(ctx context.Context, relative string) error {
path, err := dispatcher.root.DestructivePath(relative)
if err != nil {
return err
}
return dispatcher.writer.WriteFrame(ctx, client.Frame{Type: client.FrameBinary, Payload: buildLegacyFMList(path)})
}
func (dispatcher legacyFMCommandDispatcher) upload(ctx context.Context, relative string, size uint64) error {
path, err := dispatcher.root.DestructivePath(relative)
if err != nil {
return err
}
return dispatcher.writer.WriteFrame(ctx, client.Frame{Type: client.FrameBinary, Payload: buildLegacyFMUpload(path, size)})
}
func (dispatcher legacyFMCommandDispatcher) download(ctx context.Context, relative string) error {
path, err := dispatcher.root.DestructivePath(relative)
if err != nil {
return err
}
return dispatcher.writer.WriteFrame(ctx, client.Frame{Type: client.FrameBinary, Payload: buildLegacyFMDownload(path)})
}
func createLegacyFMSession(ctx context.Context, dashboardClient *client.Client, serverID uint64, capabilities ...client.IOStreamCapability) (string, error) {
var capability client.IOStreamCapability
if len(capabilities) > 0 {
capability = capabilities[0]
}
response, err := client.DoREST[struct{}, struct {
SessionID string `json:"session_id"`
}](ctx, dashboardClient, client.RESTRequest[struct{}]{Method: http.MethodPost, Path: fmt.Sprintf("/api/v1/file?id=%d", serverID), IOStreamCapability: capability})
if err != nil {
return "", err
}
if response.SessionID == "" {
return "", errors.New("FM session response omitted session id")
}
return response.SessionID, nil
}
func verifyLegacyFMMissingScopes(ctx context.Context, dashboardInstance *dashboard.Dashboard, serverID uint64, session string) error {
incompleteScopeSets := [][]string{
{"nezha:server:write", "nezha:server:delete"},
{"nezha:server:read", "nezha:server:delete"},
{"nezha:server:read", "nezha:server:write"},
}
var scopeChecksErr error
for _, scopes := range incompleteScopeSets {
limited, err := createScopedClient(ctx, dashboardInstance, scopes)
if err != nil {
scopeChecksErr = errors.Join(scopeChecksErr, err)
continue
}
_, createErr := createLegacyFMSession(ctx, limited, serverID)
_, attachErr := limited.DialWebSocket(ctx, "/api/v1/ws/file/"+session)
if !isForbidden(createErr) || !isForbidden(attachErr) {
scopeChecksErr = errors.Join(scopeChecksErr, createErr, attachErr, errors.New("incomplete FM scopes were accepted"))
}
}
return scopeChecksErr
}
func findLegacyFMServerID(ctx context.Context, dashboardInstance *dashboard.Dashboard, uuid string) (uint64, error) {
type serverListArguments struct {
OnlineOnly bool `json:"online_only"`
}
type serverListResult struct {
Servers []struct {
ID uint64 `json:"id"`
UUID string `json:"uuid"`
} `json:"servers"`
}
result, err := client.CallTool[serverListArguments, serverListResult](ctx, dashboardInstance.Clients().MCP, client.ToolCall[serverListArguments]{Name: "server.list", Arguments: serverListArguments{OnlineOnly: true}})
if err != nil {
return 0, err
}
for _, server := range result.StructuredContent.Servers {
if server.UUID == uuid && server.ID != 0 {
return server.ID, nil
}
}
return 0, errors.New("server.list omitted online FM server")
}
func createForeignLegacyFMClient(ctx context.Context, dashboardInstance *dashboard.Dashboard) (*client.Client, func() error, error) {
const username = "agentcompat-fm-foreign"
const password = "agentcompat-fm-password"
admin := dashboardInstance.Clients().REST
userID, err := client.DoREST[legacyFMUserForm, uint64](ctx, admin, client.RESTRequest[legacyFMUserForm]{
Method: http.MethodPost,
Path: "/api/v1/user",
Body: &legacyFMUserForm{Role: 1, Username: username, Password: password},
})
if err != nil {
return nil, func() error { return nil }, err
}
cleanup := func() error {
cleanupContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), 10*time.Second)
defer cancel()
_, cleanupErr := client.DoREST[[]uint64, struct{}](cleanupContext, admin, client.RESTRequest[[]uint64]{Method: http.MethodPost, Path: "/api/v1/batch-delete/user", Body: &[]uint64{userID}})
return cleanupErr
}
loginClient, err := client.New(client.Config{BaseURL: dashboardInstance.URL()})
if err != nil {
return nil, func() error { return nil }, errors.Join(err, cleanup())
}
if _, err := loginClient.Login(ctx, client.LoginRequest{Username: username, Password: password}); err != nil {
return nil, func() error { return nil }, errors.Join(err, cleanup())
}
pat, err := client.DoREST[patRequest, patResponse](ctx, loginClient, client.RESTRequest[patRequest]{
Method: http.MethodPost,
Path: "/api/v1/api-tokens",
Body: &patRequest{Name: "agentcompat-fm-foreign", Scopes: []string{
"nezha:server:read", "nezha:server:write", "nezha:server:delete",
}},
})
if err != nil {
return nil, func() error { return nil }, errors.Join(err, cleanup())
}
foreign, err := dashboardInstance.AuthenticatedClient(pat.Token)
if err != nil {
return nil, func() error { return nil }, errors.Join(err, cleanup())
}
return foreign, cleanup, nil
}
func readLegacyFMDownload(ctx context.Context, connection *client.WebSocketConnection, size uint64) ([]byte, int, error) {
content := make([]byte, 0, size)
frameCount := 0
for uint64(len(content)) < size {
frame, err := readBinaryFrame(ctx, connection)
if err != nil {
return nil, frameCount, err
}
frameCount++
if message, ok := parseLegacyFMError(frame); ok {
return nil, frameCount, message
}
remaining := size - uint64(len(content))
if uint64(len(frame)) > remaining {
return nil, frameCount, errLegacyFMUnexpected
}
content = append(content, frame...)
}
return content, frameCount, nil
}
func waitForLegacyFMSessionCleanup(ctx context.Context, owner *client.Client, session string) (int, error) {
ticker := time.NewTicker(25 * time.Millisecond)
defer ticker.Stop()
for {
connection, err := owner.DialWebSocket(ctx, "/api/v1/ws/file/"+session)
if isLegacyFMSessionRejected(err) {
return 0, nil
}
if connection != nil {
_ = connection.Close()
}
select {
case <-ctx.Done():
return 1, ctx.Err()
case <-ticker.C:
}
}
}
func cleanupLegacyFMFixtures(ctx context.Context, filesystem mcpFilesystemClient) error {
deleted, err := filesystem.delete(ctx, "legacy", true)
if err != nil {
return err
}
if deleted.StructuredContent.DeletedCount != 5 {
return fmt.Errorf("FM cleanup deleted %d entries, want 5", deleted.StructuredContent.DeletedCount)
}
remaining, err := filesystem.list(ctx, ".", true)
if err != nil {
return err
}
if remaining.StructuredContent.Total != 0 || len(remaining.StructuredContent.Entries) != 0 {
return errors.New("FM fixture residue remains after cleanup")
}
return nil
}
func isLegacyFMSessionRejected(err error) bool {
if isForbidden(err) {
return true
}
var handshakeErr *client.WebSocketHandshakeError
return errors.As(err, &handshakeErr) && strings.Contains(handshakeErr.Message, "permission denied")
}
func readBinaryFrame(ctx context.Context, connection *client.WebSocketConnection) ([]byte, error) {
frame, err := connection.ReadFrame(ctx)
if err != nil {
return nil, err
}
if frame.Type != client.FrameBinary {
return nil, errLegacyFMUnexpected
}
return frame.Payload, nil
}
func finishLegacyFM(assertions *AssertionSet, runErr error) (Result, error) {
for _, assertion := range assertions.assertions {
if !assertion.Passed && runErr == nil {
runErr = fmt.Errorf("%s: %s", assertion.Name, assertion.Details)
}
}
result := Result{Name: "legacy-fm", Passed: runErr == nil, Assertions: assertions.Results(), CleanupOK: true}
if runErr != nil {
result.Error = errorText(runErr)
}
return result, runErr
}
@@ -0,0 +1,147 @@
//go:build linux
package scenario
import (
"bytes"
"context"
"crypto/rand"
"encoding/hex"
"errors"
"fmt"
"os"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process"
)
type legacyFMProducerObservation struct {
RunID string
AgentUUID string
SessionID string
Samples []agent.FMProducerSample
}
type legacyFMProducerAwaiter struct {
observer *agent.FMProducerObserver
identity legacyFMProducerObservation
}
func newLegacyFMProducerAwaiter(observer *agent.FMProducerObserver, runID, agentUUID, sessionID string) legacyFMProducerAwaiter {
return legacyFMProducerAwaiter{
observer: observer,
identity: legacyFMProducerObservation{RunID: runID, AgentUUID: agentUUID, SessionID: sessionID},
}
}
func (awaiter legacyFMProducerAwaiter) observation(active, closed agent.FMProducerSample) legacyFMProducerObservation {
result := awaiter.identity
result.Samples = []agent.FMProducerSample{active, closed}
return result
}
func (awaiter legacyFMProducerAwaiter) await(ctx context.Context, phase string) (agent.FMProducerSample, error) {
return awaiter.observer.Await(ctx, func(sample agent.FMProducerSample) bool {
matches := sample.RunID == awaiter.identity.RunID && sample.AgentUUID == awaiter.identity.AgentUUID && sample.SessionID == awaiter.identity.SessionID && sample.Phase == phase
return matches && (phase != "active" || sample.Active > 0)
})
}
func (observation legacyFMProducerObservation) validate() error {
if observation.RunID == "" || observation.AgentUUID == "" || observation.SessionID == "" {
return errors.New("FM producer observation identity is incomplete")
}
activeObserved := false
closedObserved := false
for _, sample := range observation.Samples {
if sample.RunID != observation.RunID || sample.AgentUUID != observation.AgentUUID || sample.SessionID != observation.SessionID {
return errors.New("FM producer observation identity mismatch")
}
switch sample.Phase {
case "active":
activeObserved = activeObserved || sample.Active > 0
case "idle":
continue
case "closed":
closedObserved = sample.Active == 0
}
}
if !activeObserved {
return errors.New("FM producer observation never saw an active producer")
}
if !closedObserved {
return errors.New("FM producer observation omitted closed zero state")
}
return nil
}
func (observation legacyFMProducerObservation) details() string {
active := int64(0)
closed := int64(-1)
for _, sample := range observation.Samples {
if sample.Phase == "active" && sample.Active > active {
active = sample.Active
}
if sample.Phase == "closed" {
closed = sample.Active
}
}
return fmt.Sprintf("fm_producer_active_count: %d; fm_producer_residue_count: %d; run_id=%s agent_uuid=%s session_id=%s source=live-agent-task", active, closed, observation.RunID, observation.AgentUUID, observation.SessionID)
}
type legacyFMSentinelCoverage struct {
ListRejected bool
UploadRejected bool
DownloadRejected bool
SuccessCount int
ErrorCount int
}
func (coverage legacyFMSentinelCoverage) validate() error {
if !coverage.ListRejected || !coverage.UploadRejected || !coverage.DownloadRejected {
return errors.New("sentinel coverage omitted a typed rejected dispatcher")
}
if coverage.SuccessCount == 0 || coverage.ErrorCount == 0 {
return errors.New("sentinel coverage omitted a real success or error operation")
}
return nil
}
type legacyFMProcessResidue struct {
Baseline processharness.Sample
End processharness.Sample
}
func (residue legacyFMProcessResidue) validate() error {
if residue.Baseline.NonStdioFDCount != residue.End.NonStdioFDCount {
return fmt.Errorf("Agent non-stdio FD drift: baseline=%d end=%d", residue.Baseline.NonStdioFDCount, residue.End.NonStdioFDCount)
}
if residue.Baseline.DescendantCount != residue.End.DescendantCount {
return fmt.Errorf("Agent descendant drift: baseline=%d end=%d", residue.Baseline.DescendantCount, residue.End.DescendantCount)
}
if residue.Baseline.TCPListenerCount != residue.End.TCPListenerCount || residue.Baseline.TCP6ListenerCount != residue.End.TCP6ListenerCount {
return errors.New("Agent listener count drift")
}
return nil
}
func newLegacyFMRunID() (string, error) {
bytes := make([]byte, 16)
if _, err := rand.Read(bytes); err != nil {
return "", fmt.Errorf("generate FM observation run id: %w", err)
}
return hex.EncodeToString(bytes), nil
}
func verifyLegacyFMSentinels(sentinelPaths []string, sentinel []byte) error {
for _, path := range sentinelPaths {
content, err := os.ReadFile(path)
if err != nil {
return err
}
if !bytes.Equal(content, sentinel) {
return errors.New("outside-root sentinel changed")
}
}
return nil
}
@@ -0,0 +1,235 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"os"
"slices"
"syscall"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/evidence"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
const mcpFilesystemScenarioName = "mcp-filesystem"
type MCPFilesystemInput struct {
Paths contract.Paths
}
type MCPFilesystem struct{}
type mcpFilesystemClient struct {
client *client.Client
serverID uint64
root fixture.AgentRoot
}
type mcpFilesystemWrite struct {
relative string
content string
encoding string
mode string
ifMatchSHA256 string
createDirs bool
}
func newMCPFilesystemClient(mcpClient *client.Client, serverID uint64, root fixture.AgentRoot) mcpFilesystemClient {
return mcpFilesystemClient{client: mcpClient, serverID: serverID, root: root}
}
func (filesystem mcpFilesystemClient) list(ctx context.Context, relative string, showHidden bool) (client.ToolCallResult[client.FsListResult], error) {
path, err := filesystem.root.Path(relative)
if err != nil {
return client.ToolCallResult[client.FsListResult]{}, err
}
arguments := client.FsListArguments{ServerID: filesystem.serverID, Path: path.String(), ShowHidden: showHidden}
return client.CallTool[client.FsListArguments, client.FsListResult](ctx, filesystem.client, client.ToolCall[client.FsListArguments]{Name: "fs.list", Arguments: arguments})
}
func (filesystem mcpFilesystemClient) read(ctx context.Context, relative string, offset, length int64, encoding string) (client.ToolCallResult[client.FsReadResult], error) {
path, err := filesystem.root.Path(relative)
if err != nil {
return client.ToolCallResult[client.FsReadResult]{}, err
}
arguments := client.FsReadArguments{ServerID: filesystem.serverID, Path: path.String(), Offset: offset, Length: length, Encoding: encoding}
return client.CallTool[client.FsReadArguments, client.FsReadResult](ctx, filesystem.client, client.ToolCall[client.FsReadArguments]{Name: "fs.read", Arguments: arguments})
}
func (filesystem mcpFilesystemClient) write(ctx context.Context, write mcpFilesystemWrite) (client.ToolCallResult[client.FsWriteResult], error) {
path, err := filesystem.root.Path(write.relative)
if err != nil {
return client.ToolCallResult[client.FsWriteResult]{}, err
}
arguments := client.FsWriteArguments{ServerID: filesystem.serverID, Path: path.String(), Content: write.content, Encoding: write.encoding, Mode: write.mode, IfMatchSHA256: write.ifMatchSHA256, CreateDirs: write.createDirs}
return client.CallTool[client.FsWriteArguments, client.FsWriteResult](ctx, filesystem.client, client.ToolCall[client.FsWriteArguments]{Name: "fs.write", Arguments: arguments})
}
func (filesystem mcpFilesystemClient) delete(ctx context.Context, relative string, recursive bool) (client.ToolCallResult[client.FsDeleteResult], error) {
path, err := filesystem.root.DestructivePath(relative)
if err != nil {
return client.ToolCallResult[client.FsDeleteResult]{}, err
}
arguments := client.FsDeleteArguments{ServerID: filesystem.serverID, Path: path.String(), Recursive: recursive}
return client.CallTool[client.FsDeleteArguments, client.FsDeleteResult](ctx, filesystem.client, client.ToolCall[client.FsDeleteArguments]{Name: "fs.delete", Arguments: arguments})
}
func (MCPFilesystem) Run(ctx context.Context, input MCPFilesystemInput) (result Result, runErr error) {
assertions := NewAssertionSet()
fixtureParent, err := os.MkdirTemp("", "agentcompat-mcp-filesystem-")
if err != nil {
return mcpFilesystemFinish(assertions, err)
}
root, err := fixture.NewAgentRoot(fixtureParent, "agent-filesystem")
if err != nil {
_ = os.RemoveAll(fixtureParent)
return mcpFilesystemFinish(assertions, err)
}
dashboardInstance, err := dashboard.Start(ctx, dashboard.StartConfig{SourceDir: input.Paths.NezhaSource().String(), ReceiptGate: true})
if err != nil {
_ = os.RemoveAll(fixtureParent)
return mcpFilesystemFinish(assertions, err)
}
defer func() {
cleanupErr := errors.Join(dashboardInstance.Stop(context.Background()), os.RemoveAll(fixtureParent))
result.CleanupOK = cleanupErr == nil && dashboardInstance.CleanupReceipt().Passed
if cleanupErr != nil && runErr == nil {
runErr = cleanupErr
result.Passed = false
result.Error = errorText(cleanupErr)
}
}()
agentInstance, err := agent.Start(ctx, agent.AgentStartConfig{SourceDir: input.Paths.AgentSource().String(), Endpoint: dashboardInstance.Endpoint(), Secret: dashboardInstance.AgentSecret(), UUID: "00000000-0000-0000-0000-000000000212"})
if err != nil {
return mcpFilesystemFinish(assertions, err)
}
defer func() {
cleanupErr := agentInstance.Stop(context.Background())
if cleanupErr != nil && runErr == nil {
runErr = cleanupErr
result.Passed = false
result.Error = errorText(cleanupErr)
}
}()
if err := dashboardInstance.WaitForReceiptAccepted(ctx); err != nil {
return mcpFilesystemFinish(assertions, err)
}
if err := dashboardInstance.ReleaseReceipt(ctx); err != nil {
return mcpFilesystemFinish(assertions, err)
}
readiness, err := agentInstance.WaitReady(ctx, dashboardInstance)
if err != nil {
return mcpFilesystemFinish(assertions, err)
}
mcpClient, err := createScopedClient(ctx, dashboardInstance, []string{"nezha:*"})
if err != nil {
return mcpFilesystemFinish(assertions, err)
}
initialize, err := mcpClient.Initialize(ctx)
assertions.Record("MCP initialize exact protocol and server", err == nil && initialize.ProtocolVersion == "2024-11-05" && initialize.ServerInfo.Name == "nezha-mcp" && initialize.ServerInfo.Version != "", errorText(err))
tools, err := mcpClient.ListTools(ctx)
toolNames := make([]string, 0, len(tools.Tools))
for _, tool := range tools.Tools {
toolNames = append(toolNames, tool.Name)
}
wantTools := []string{"fs.delete", "fs.list", "fs.read", "fs.write", "meta.whoami", "server.get", "server.list"}
assertions.Record("tools.list exposes filesystem identity and inventory tools", err == nil && containsAll(toolNames, wantTools), errorText(err))
whoami, err := client.CallTool[struct{}, client.WhoAmIResult](ctx, mcpClient, client.ToolCall[struct{}]{Name: "meta.whoami", Arguments: struct{}{}})
identity := whoami.StructuredContent
assertions.Record("meta.whoami exact administrator PAT identity", err == nil && identity.UserID != 0 && identity.IsAdmin && identity.TokenID != 0 && identity.TokenName == "agentcompat-scope-check" && slices.Equal(identity.Scopes, []string{"nezha:*"}) && len(identity.ServerIDs) == 0, errorText(err))
servers, err := client.CallTool[client.ServerListArguments, client.ServerListResult](ctx, mcpClient, client.ToolCall[client.ServerListArguments]{Name: "server.list", Arguments: client.ServerListArguments{OnlineOnly: true}})
serverID := uint64(0)
for _, server := range servers.StructuredContent.Servers {
if server.UUID == agentInstance.UUID() && server.Online {
serverID = server.ID
}
}
assertions.Record("server.list exact online UUID and count", err == nil && readiness.UUID == agentInstance.UUID() && serverID != 0 && servers.StructuredContent.Count == len(servers.StructuredContent.Servers), errorText(err))
server, err := client.CallTool[client.ServerGetArguments, client.ServerGetResult](ctx, mcpClient, client.ToolCall[client.ServerGetArguments]{Name: "server.get", Arguments: client.ServerGetArguments{ServerID: serverID}})
assertions.Record("server.get exact typed identity Host and State", err == nil && server.StructuredContent.ID == serverID && server.StructuredContent.UUID == agentInstance.UUID() && string(server.StructuredContent.Host) != "null" && string(server.StructuredContent.State) != "null", errorText(err))
if err != nil || serverID == 0 {
return mcpFilesystemFinish(assertions, errors.New("filesystem scenario inventory setup failed"))
}
filesystem := newMCPFilesystemClient(mcpClient, serverID, root)
if err := verifyMCPFilesystemPathGuards(ctx, assertions, filesystem, fixtureParent); err != nil {
return mcpFilesystemFinish(assertions, err)
}
const agentFixtureIdentity = 65534
permissionCredential := &syscall.Credential{Uid: agentFixtureIdentity, Gid: agentFixtureIdentity}
permissionAgent, err := agent.Start(ctx, agent.AgentStartConfig{SourceDir: input.Paths.AgentSource().String(), Endpoint: dashboardInstance.Endpoint(), Secret: dashboardInstance.AgentSecret(), UUID: "00000000-0000-0000-0000-000000000213", Credential: permissionCredential})
if err != nil {
return mcpFilesystemFinish(assertions, err)
}
defer func() {
cleanupErr := permissionAgent.Stop(context.Background())
if cleanupErr != nil && runErr == nil {
runErr = cleanupErr
result.Passed = false
result.Error = errorText(cleanupErr)
}
}()
if _, err := permissionAgent.WaitReady(ctx, dashboardInstance); err != nil {
return mcpFilesystemFinish(assertions, err)
}
permissionClient, err := createScopedClient(ctx, dashboardInstance, []string{"nezha:*"})
if err != nil {
return mcpFilesystemFinish(assertions, err)
}
permissionServers, err := client.CallTool[client.ServerListArguments, client.ServerListResult](ctx, permissionClient, client.ToolCall[client.ServerListArguments]{Name: "server.list", Arguments: client.ServerListArguments{OnlineOnly: true}})
if err != nil {
return mcpFilesystemFinish(assertions, err)
}
permissionServerID := uint64(0)
for _, server := range permissionServers.StructuredContent.Servers {
if server.UUID == permissionAgent.UUID() && server.Online {
permissionServerID = server.ID
}
}
if permissionServerID == 0 {
return mcpFilesystemFinish(assertions, errors.New("permission Agent is absent from server.list"))
}
permissionContract, err := observePermissionAgent(permissionAgent, permissionCredential, newMCPFilesystemClient(permissionClient, permissionServerID, root))
if err != nil {
return mcpFilesystemFinish(assertions, err)
}
if err := runMCPFilesystemOperations(ctx, assertions, filesystem, permissionContract, dashboardInstance); err != nil {
return mcpFilesystemFinish(assertions, err)
}
return mcpFilesystemFinish(assertions, nil)
}
func containsAll(values, required []string) bool {
for _, value := range required {
if !slices.Contains(values, value) {
return false
}
}
return true
}
func mcpFilesystemFinish(assertions *AssertionSet, runErr error) (Result, error) {
failedAssertion := false
for _, assertion := range assertions.assertions {
if !assertion.Passed {
failedAssertion = true
if runErr == nil {
runErr = fmt.Errorf("%s: %s", assertion.Name, assertion.Details)
}
}
}
if runErr != nil && !failedAssertion {
assertions.Record("scenario execution completed", false, errorText(runErr))
}
result := Result{Name: mcpFilesystemScenarioName, Passed: runErr == nil, Assertions: assertions.Results()}
if runErr != nil {
result.Error = evidence.Redact(runErr.Error())
}
return result, runErr
}
@@ -0,0 +1,59 @@
//go:build linux
package scenario
import (
"fmt"
"os"
"strconv"
"strings"
"syscall"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
)
type permissionAgentContract struct {
filesystem mcpFilesystemClient
uid uint32
gid uint32
processContract bool
}
func observePermissionAgent(agentInstance *agent.Agent, credential *syscall.Credential, filesystem mcpFilesystemClient) (permissionAgentContract, error) {
status, err := os.ReadFile(fmt.Sprintf("/proc/%d/status", agentInstance.PID()))
if err != nil {
return permissionAgentContract{}, fmt.Errorf("read permission Agent process status: %w", err)
}
uid, err := effectiveProcessIdentity(status, "Uid:")
if err != nil {
return permissionAgentContract{}, err
}
gid, err := effectiveProcessIdentity(status, "Gid:")
if err != nil {
return permissionAgentContract{}, err
}
return permissionAgentContract{
filesystem: filesystem,
uid: uid,
gid: gid,
processContract: uid == credential.Uid && gid == credential.Gid,
}, nil
}
func effectiveProcessIdentity(status []byte, field string) (uint32, error) {
for line := range strings.SplitSeq(string(status), "\n") {
if !strings.HasPrefix(line, field) {
continue
}
values := strings.Fields(strings.TrimPrefix(line, field))
if len(values) < 2 {
break
}
identity, err := strconv.ParseUint(values[1], 10, 32)
if err != nil {
return 0, fmt.Errorf("parse effective process %s: %w", field, err)
}
return uint32(identity), nil
}
return 0, fmt.Errorf("process status missing %s", field)
}
@@ -0,0 +1,148 @@
//go:build linux
package scenario
import (
"context"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"errors"
"fmt"
"net/http"
"os"
"slices"
"strings"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
)
func runMCPFilesystemOperations(ctx context.Context, assertions *AssertionSet, filesystem mcpFilesystemClient, permission permissionAgentContract, dashboardInstance *dashboard.Dashboard) error {
text := "typed filesystem payload"
textHash := sha256.Sum256([]byte(text))
textDigest := hex.EncodeToString(textHash[:])
written, err := filesystem.write(ctx, mcpFilesystemWrite{relative: "tree/text.txt", content: text, encoding: "utf8", mode: "0640", createDirs: true})
assertions.Record("fs.write create_dirs size and SHA", err == nil && written.StructuredContent.Size == int64(len(text)) && written.StructuredContent.SHA256 == textDigest, errorText(err))
if err != nil {
return err
}
tree, err := filesystem.list(ctx, "tree", false)
assertions.Record("fs.write applies exact file mode", err == nil && len(tree.StructuredContent.Entries) == 1 && tree.StructuredContent.Entries[0].Name == "text.txt" && tree.StructuredContent.Entries[0].Type == "file" && tree.StructuredContent.Entries[0].Size == int64(len(text)) && tree.StructuredContent.Entries[0].Mode == "0640", errorText(err))
_, err = filesystem.write(ctx, mcpFilesystemWrite{relative: ".hidden", content: "hidden", encoding: "utf8", mode: "0600"})
if err != nil {
return err
}
visible, err := filesystem.list(ctx, ".", false)
assertions.Record("fs.list hides dot entries with exact totals", err == nil && visible.StructuredContent.Total == 1 && entryNames(visible.StructuredContent.Entries) == "tree", errorText(err))
hidden, err := filesystem.list(ctx, ".", true)
assertions.Record("fs.list show_hidden includes exact dot entry", err == nil && hidden.StructuredContent.Total == 2 && entryNames(hidden.StructuredContent.Entries) == ".hidden,tree", errorText(err))
// Inventory and setup consume the PAT's fixed MCP call budget; rotate before content and CAS checks.
filesystem, err = refreshedMCPFilesystemClient(ctx, dashboardInstance, filesystem)
if err != nil {
return err
}
readText, err := filesystem.read(ctx, "tree/text.txt", 0, int64(len(text)), "utf8")
assertions.Record("fs.read utf8 exact content size SHA and truncation", err == nil && readText.StructuredContent.Content == text && readText.StructuredContent.Encoding == "utf8" && readText.StructuredContent.Size == int64(len(text)) && readText.StructuredContent.SHA256 == textDigest && !readText.StructuredContent.Truncated, errorText(err))
binary := []byte{0x00, 0x41, 0xff, 0x42}
binaryHash := sha256.Sum256(binary)
binaryDigest := hex.EncodeToString(binaryHash[:])
_, err = filesystem.write(ctx, mcpFilesystemWrite{relative: "tree/binary.bin", content: base64.StdEncoding.EncodeToString(binary), encoding: "base64", mode: "0600"})
if err != nil {
return err
}
readBinary, err := filesystem.read(ctx, "tree/binary.bin", 0, int64(len(binary)), "base64")
assertions.Record("fs.read base64 exact content size and SHA", err == nil && readBinary.StructuredContent.Content == base64.StdEncoding.EncodeToString(binary) && readBinary.StructuredContent.Encoding == "base64" && readBinary.StructuredContent.Size == int64(len(binary)) && readBinary.StructuredContent.SHA256 == binaryDigest, errorText(err))
updated := "CAS updated payload"
updatedHash := sha256.Sum256([]byte(updated))
updatedDigest := hex.EncodeToString(updatedHash[:])
cas, err := filesystem.write(ctx, mcpFilesystemWrite{relative: "tree/text.txt", content: updated, encoding: "utf8", mode: "0640", ifMatchSHA256: textDigest})
assertions.Record("fs.write CAS success exact size and SHA", err == nil && cas.StructuredContent.Size == int64(len(updated)) && cas.StructuredContent.SHA256 == updatedDigest, errorText(err))
_, casErr := filesystem.write(ctx, mcpFilesystemWrite{relative: "tree/text.txt", content: "must-not-land", encoding: "utf8", mode: "0640", ifMatchSHA256: textDigest})
unchanged, readErr := filesystem.read(ctx, "tree/text.txt", 0, int64(len(updated)), "utf8")
assertions.Record("fs.write CAS mismatch is typed and leaves content unchanged", toolFailureContains(casErr, "if_match precondition failed") && readErr == nil && unchanged.StructuredContent.Content == updated && unchanged.StructuredContent.SHA256 == updatedDigest, errorText(errors.Join(casErr, readErr)))
permissionDirectory, pathErr := filesystem.root.Path("permission-denied")
if pathErr != nil {
return pathErr
}
if err := os.Mkdir(permissionDirectory.String(), 0o500); err != nil {
return err
}
permissionResponse, permissionErr := permission.filesystem.write(ctx, mcpFilesystemWrite{relative: "permission-denied/file.txt", content: "denied", encoding: "utf8", mode: "0600"})
permissionDeniedObserved := permissionErr == nil && permissionResponse.StructuredContent.Error == "permission denied"
assertions.Record("fs.write Agent filesystem permission denial is typed", permission.processContract && permissionDeniedObserved, fmt.Sprintf("%s; uid=%d gid=%d process_contract=%t", permissionResponse.StructuredContent.Error, permission.uid, permission.gid, permission.processContract))
if err := os.Remove(permissionDirectory.String()); err != nil {
return err
}
filesystem, err = refreshedMCPFilesystemClient(ctx, dashboardInstance, filesystem)
if err != nil {
return err
}
_, missingErr := filesystem.read(ctx, "tree/missing.txt", 0, 1, "utf8")
assertions.Record("fs.read nonexistent path is typed", toolFailureContains(missingErr, "does not exist"), errorText(missingErr))
_, encodingErr := filesystem.read(ctx, "tree/text.txt", 0, 1, "rot13")
assertions.Record("fs.read invalid encoding is typed", toolFailureContains(encodingErr, "unknown encoding"), errorText(encodingErr))
_, writeEncodingErr := filesystem.write(ctx, mcpFilesystemWrite{relative: "tree/invalid-encoding.txt", content: "x", encoding: "rot13", mode: "0600"})
assertions.Record("fs.write invalid encoding is typed", toolFailureContains(writeEncodingErr, "unknown encoding"), errorText(writeEncodingErr))
_, modeErr := filesystem.write(ctx, mcpFilesystemWrite{relative: "tree/invalid-mode.txt", content: "x", encoding: "utf8", mode: "invalid"})
assertions.Record("fs.write invalid mode is typed", toolFailureContains(modeErr, "invalid mode"), errorText(modeErr))
filesystem, err = refreshedMCPFilesystemClient(ctx, dashboardInstance, filesystem)
if err != nil {
return err
}
oversize, oversizeErr := client.DoREST[agentcompatFsWriteContractRequest, agentcompatFsWriteContractResponse](ctx, filesystem.client, client.RESTRequest[agentcompatFsWriteContractRequest]{Method: http.MethodPost, Path: "/agentcompat/fs-write-contract", Body: &agentcompatFsWriteContractRequest{ServerID: filesystem.serverID, Operation: agentcompatFsWriteOperationOversize}})
oversizeHandlerObserved := oversize.AgentRPCResponse
productionMaxWriteCheck := oversizeErr == nil && oversize.Result.Error == "content exceeds max write size"
assertions.Record("fs.write Agent oversize contract is typed", oversizeHandlerObserved && productionMaxWriteCheck, fmt.Sprintf("%s; agent_handler=%t production_max_write_check=%t; agent_rpc_response=%t", oversize.Result.Error, oversizeHandlerObserved, productionMaxWriteCheck, oversize.AgentRPCResponse))
_, nonrecursiveErr := filesystem.delete(ctx, "tree", false)
assertions.Record("fs.delete nonrecursive rejects nonempty directory", toolFailureContains(nonrecursiveErr, "internal agent error"), errorText(nonrecursiveErr))
deleted, err := filesystem.delete(ctx, "tree", true)
assertions.Record("fs.delete recursive returns exact positive count", err == nil && deleted.StructuredContent.DeletedCount == 3, errorText(err))
final, err := filesystem.list(ctx, ".", true)
assertions.Record("fs.list final absence after recursive delete", err == nil && final.StructuredContent.Total == 1 && entryNames(final.StructuredContent.Entries) == ".hidden", errorText(err))
_, err = filesystem.delete(ctx, ".hidden", false)
if err != nil {
return err
}
empty, err := filesystem.list(ctx, ".", true)
assertions.Record("fs.list final fixture root is empty", err == nil && empty.StructuredContent.Total == 0 && len(empty.StructuredContent.Entries) == 0, errorText(err))
return err
}
type agentcompatFsWriteOperation string
const (
agentcompatFsWriteOperationOversize agentcompatFsWriteOperation = "oversize"
)
type agentcompatFsWriteContractRequest struct {
ServerID uint64 `json:"server_id"`
Operation agentcompatFsWriteOperation `json:"operation"`
}
type agentcompatFsWriteContractResponse struct {
Result client.FsWriteResult `json:"result"`
AgentRPCResponse bool `json:"agent_rpc_response"`
}
func refreshedMCPFilesystemClient(ctx context.Context, dashboardInstance *dashboard.Dashboard, filesystem mcpFilesystemClient) (mcpFilesystemClient, error) {
mcpClient, err := createScopedClient(ctx, dashboardInstance, []string{"nezha:*"})
if err != nil {
return mcpFilesystemClient{}, err
}
return newMCPFilesystemClient(mcpClient, filesystem.serverID, filesystem.root), nil
}
func entryNames(entries []client.FsEntry) string {
names := make([]string, 0, len(entries))
for _, entry := range entries {
names = append(names, entry.Name)
}
slices.Sort(names)
return strings.Join(names, ",")
}
func toolFailureContains(err error, text string) bool {
var failure *client.ToolFailure
return errors.As(err, &failure) && strings.Contains(failure.Message, text)
}
@@ -0,0 +1,96 @@
//go:build linux
package scenario
import (
"context"
"crypto/sha256"
"errors"
"fmt"
"os"
"path/filepath"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
func verifyMCPFilesystemPathGuards(ctx context.Context, assertions *AssertionSet, filesystem mcpFilesystemClient, fixtureParent string) error {
directTarget := filepath.Join(fixtureParent, "outside-sentinel")
symlinkParentTarget := filepath.Join(fixtureParent, "outside-directory", "parent-target.txt")
symlinkFinalTarget := filepath.Join(fixtureParent, "final-target.txt")
if err := os.Mkdir(filepath.Dir(symlinkParentTarget), 0o700); err != nil {
return err
}
for _, target := range []string{directTarget, symlinkParentTarget, symlinkFinalTarget} {
if err := os.WriteFile(target, []byte("unchanged:"+filepath.Base(target)), 0o600); err != nil {
return err
}
}
if err := os.Symlink(filepath.Dir(symlinkParentTarget), filepath.Join(filesystem.root.Absolute(), "linked-parent")); err != nil {
return err
}
if err := os.Symlink(symlinkFinalTarget, filepath.Join(filesystem.root.Absolute(), "linked-final")); err != nil {
return err
}
before, err := filesystemTargetHashes(directTarget, symlinkParentTarget, symlinkFinalTarget)
if err != nil {
return err
}
requestsBefore := filesystem.client.RequestCount()
rejections := []struct {
want fixture.PathRejectionReason
run func() error
}{
{fixture.PathRejectionAbsolute, func() error { _, err := filesystem.read(ctx, directTarget, 0, 1, "utf8"); return err }},
{fixture.PathRejectionParent, func() error {
_, err := filesystem.write(ctx, mcpFilesystemWrite{relative: "../outside-sentinel", content: "changed", encoding: "utf8", mode: "0600"})
return err
}},
{fixture.PathRejectionVolume, func() error { _, err := filesystem.list(ctx, `C:\outside.txt`, false); return err }},
{fixture.PathRejectionSeparator, func() error { _, err := filesystem.read(ctx, `inside\outside.txt`, 0, 1, "utf8"); return err }},
{fixture.PathRejectionDestructiveRoot, func() error { _, err := filesystem.delete(ctx, ".", true); return err }},
{fixture.PathRejectionSymlinkParent, func() error {
_, err := filesystem.write(ctx, mcpFilesystemWrite{relative: "linked-parent/parent-target.txt", content: "changed", encoding: "utf8", mode: "0600"})
return err
}},
{fixture.PathRejectionSymlinkFinal, func() error { _, err := filesystem.delete(ctx, "linked-final", false); return err }},
}
matched := 0
for _, rejection := range rejections {
var pathError *fixture.AgentPathError
if errors.As(rejection.run(), &pathError) && pathError.Reason == rejection.want {
matched++
}
}
requestsAfter := filesystem.client.RequestCount()
dispatched := int64(requestsAfter) - int64(requestsBefore)
after, err := filesystemTargetHashes(directTarget, symlinkParentTarget, symlinkFinalTarget)
if err != nil {
return err
}
if err := errors.Join(
os.Remove(filepath.Join(filesystem.root.Absolute(), "linked-parent")),
os.Remove(filepath.Join(filesystem.root.Absolute(), "linked-final")),
); err != nil {
return err
}
assertions.Record("fixture path rejections dispatch zero MCP HTTP requests", matched == len(rejections) && dispatched == 0, fmt.Sprintf("path_rejections_dispatched: %d; matched=%d total=%d requests_before=%d requests_after=%d", dispatched, matched, len(rejections), requestsBefore, requestsAfter))
assertions.Record("fixture path rejections leave outside and symlink targets unchanged", before == after, fmt.Sprintf("before=%x after=%x", before, after))
return nil
}
func filesystemTargetHashes(paths ...string) ([32]byte, error) {
hash := sha256.New()
for _, path := range paths {
content, err := os.ReadFile(path)
if err != nil {
return [32]byte{}, err
}
if _, err := hash.Write(content); err != nil {
return [32]byte{}, err
}
}
var digest [32]byte
copy(digest[:], hash.Sum(nil))
return digest, nil
}
@@ -0,0 +1,263 @@
//go:build linux && agentcompat
package scenario
import (
"context"
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
func TestMCPFilesystemScenario_RealFlow(t *testing.T) {
nezhaSource := os.Getenv("AGENTCOMPAT_NEZHA_SOURCE")
agentSource := os.Getenv("AGENTCOMPAT_AGENT_SOURCE")
if nezhaSource == "" || agentSource == "" {
t.Skip("set AGENTCOMPAT_NEZHA_SOURCE and AGENTCOMPAT_AGENT_SOURCE")
}
paths, err := contract.NewPaths(nezhaSource, agentSource, t.TempDir())
require.NoError(t, err)
result, err := (MCPFilesystem{}).Run(t.Context(), MCPFilesystemInput{Paths: paths})
for _, assertion := range result.Assertions {
t.Logf("assertion=%q passed=%t details=%q", assertion.Name, assertion.Passed, assertion.Details)
}
require.NoError(t, err)
require.True(t, result.Passed)
require.True(t, result.CleanupOK)
requireScenarioAssertions(t, result,
"fixture path rejections dispatch zero MCP HTTP requests",
"fixture path rejections leave outside and symlink targets unchanged",
"fs.write Agent filesystem permission denial is typed",
"fs.write Agent oversize contract is typed",
)
requireScenarioAssertionDetails(t, result, map[string]string{
"fixture path rejections dispatch zero MCP HTTP requests": "path_rejections_dispatched: 0",
"fs.write Agent filesystem permission denial is typed": "uid=65534 gid=65534 agent_handler=true",
"fs.write Agent oversize contract is typed": "agent_handler=true production_max_write_check=true",
})
requireScenarioAssertionDetails(t, result, map[string]string{
"fs.write Agent filesystem permission denial is typed": "process_contract=true agent_rpc_response=true",
"fs.write Agent oversize contract is typed": "agent_rpc_response=true",
})
}
func requireScenarioAssertions(t *testing.T, result Result, names ...string) {
t.Helper()
assertions := make(map[string]bool, len(result.Assertions))
for _, assertion := range result.Assertions {
assertions[assertion.Name] = assertion.Passed
}
for _, name := range names {
require.Truef(t, assertions[name], "required passing assertion %q is absent", name)
}
}
func requireScenarioAssertionDetails(t *testing.T, result Result, expected map[string]string) {
t.Helper()
details := make(map[string]string, len(result.Assertions))
for _, assertion := range result.Assertions {
details[assertion.Name] = assertion.Details
}
for name, exact := range expected {
require.Containsf(t, details[name], exact, "assertion %q lacks required computed evidence", name)
}
}
func TestMCPFilesystemClient_RejectsUnsafeFixturePathsBeforeDispatch(t *testing.T) {
t.Parallel()
tests := []struct {
name string
reason fixture.PathRejectionReason
configure func(t *testing.T, root fixture.AgentRoot) (string, string)
invoke func(context.Context, mcpFilesystemClient, string) error
}{
{
name: "absolute",
reason: fixture.PathRejectionAbsolute,
configure: func(t *testing.T, _ fixture.AgentRoot) (string, string) {
return filepath.Join(t.TempDir(), "outside"), ""
},
invoke: func(ctx context.Context, filesystem mcpFilesystemClient, candidate string) error {
_, err := filesystem.read(ctx, candidate, 0, 1, "utf8")
return err
},
},
{
name: "parent",
reason: fixture.PathRejectionParent,
configure: func(*testing.T, fixture.AgentRoot) (string, string) {
return "../outside", ""
},
invoke: func(ctx context.Context, filesystem mcpFilesystemClient, candidate string) error {
_, err := filesystem.write(ctx, mcpFilesystemWrite{relative: candidate, content: "changed", encoding: "utf8", mode: "0600", createDirs: true})
return err
},
},
{
name: "volume",
reason: fixture.PathRejectionVolume,
configure: func(*testing.T, fixture.AgentRoot) (string, string) {
return `C:\outside.txt`, ""
},
invoke: func(ctx context.Context, filesystem mcpFilesystemClient, candidate string) error {
_, err := filesystem.list(ctx, candidate, false)
return err
},
},
{
name: "separator",
reason: fixture.PathRejectionSeparator,
configure: func(*testing.T, fixture.AgentRoot) (string, string) {
return `inside\outside.txt`, ""
},
invoke: func(ctx context.Context, filesystem mcpFilesystemClient, candidate string) error {
_, err := filesystem.read(ctx, candidate, 0, 1, "base64")
return err
},
},
{
name: "destructive root",
reason: fixture.PathRejectionDestructiveRoot,
configure: func(*testing.T, fixture.AgentRoot) (string, string) {
return ".", ""
},
invoke: func(ctx context.Context, filesystem mcpFilesystemClient, candidate string) error {
_, err := filesystem.delete(ctx, candidate, true)
return err
},
},
{
name: "symlink parent",
reason: fixture.PathRejectionSymlinkParent,
configure: func(t *testing.T, root fixture.AgentRoot) (string, string) {
outside := t.TempDir()
target := filepath.Join(outside, "file.txt")
require.NoError(t, os.WriteFile(target, []byte("unchanged"), 0o600))
require.NoError(t, os.Symlink(outside, filepath.Join(root.Absolute(), "linked")))
return "linked/file.txt", target
},
invoke: func(ctx context.Context, filesystem mcpFilesystemClient, candidate string) error {
_, err := filesystem.write(ctx, mcpFilesystemWrite{relative: candidate, content: "changed", encoding: "utf8", mode: "0600", createDirs: true})
return err
},
},
{
name: "symlink final",
reason: fixture.PathRejectionSymlinkFinal,
configure: func(t *testing.T, root fixture.AgentRoot) (string, string) {
target := filepath.Join(t.TempDir(), "outside")
require.NoError(t, os.WriteFile(target, []byte("unchanged"), 0o600))
require.NoError(t, os.Symlink(target, filepath.Join(root.Absolute(), "linked.txt")))
return "linked.txt", target
},
invoke: func(ctx context.Context, filesystem mcpFilesystemClient, candidate string) error {
_, err := filesystem.delete(ctx, candidate, false)
return err
},
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
t.Parallel()
var requests atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
requests.Add(1)
}))
t.Cleanup(server.Close)
mcpClient, err := client.New(client.Config{BaseURL: server.URL, RequestTimeout: time.Second})
require.NoError(t, err)
parent := t.TempDir()
sentinel := filepath.Join(parent, "outside-sentinel")
require.NoError(t, os.WriteFile(sentinel, []byte("unchanged"), 0o600))
root, err := fixture.NewAgentRoot(parent, "mcp-filesystem")
require.NoError(t, err)
filesystem := newMCPFilesystemClient(mcpClient, 7, root)
candidate, symlinkTarget := test.configure(t, root)
err = test.invoke(t.Context(), filesystem, candidate)
var pathError *fixture.AgentPathError
require.ErrorAs(t, err, &pathError)
require.Equal(t, test.reason, pathError.Reason)
require.Zero(t, requests.Load())
content, readErr := os.ReadFile(sentinel)
require.NoError(t, readErr)
require.Equal(t, "unchanged", string(content))
if symlinkTarget != "" {
content, readErr = os.ReadFile(symlinkTarget)
require.NoError(t, readErr)
require.Equal(t, "unchanged", string(content))
}
})
}
}
func TestMCPFilesystemClient_UsesAgentPathForExactToolArguments(t *testing.T) {
t.Parallel()
requests := make(chan testFilesystemToolCall, 1)
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
call, err := decodeTestFilesystemToolCall(request.Body)
require.NoError(t, err)
requests <- call
writer.Header().Set("Content-Type", "application/json")
_, err = writer.Write([]byte(`{"jsonrpc":"2.0","id":1,"result":{"content":[{"type":"text","text":"ok"}],"structuredContent":{"size":7,"sha256":"239f59ed55e737c77147cf55ad0c1b030b6d7ee748a7426952f9b852d5a935e5"}}}`))
require.NoError(t, err)
}))
t.Cleanup(server.Close)
mcpClient, err := client.New(client.Config{BaseURL: server.URL, RequestTimeout: time.Second})
require.NoError(t, err)
root, err := fixture.NewAgentRoot(t.TempDir(), "mcp-filesystem")
require.NoError(t, err)
filesystem := newMCPFilesystemClient(mcpClient, 17, root)
result, err := filesystem.write(t.Context(), mcpFilesystemWrite{relative: "nested/payload.txt", content: "payload", encoding: "utf8", mode: "0640", createDirs: true})
require.NoError(t, err)
require.Equal(t, int64(7), result.StructuredContent.Size)
require.Equal(t, "239f59ed55e737c77147cf55ad0c1b030b6d7ee748a7426952f9b852d5a935e5", result.StructuredContent.SHA256)
call := <-requests
require.Equal(t, "fs.write", call.Name)
require.Equal(t, uint64(17), call.Arguments.ServerID)
require.True(t, filepath.IsAbs(call.Arguments.Path))
require.Equal(t, filepath.Join(root.Absolute(), "nested", "payload.txt"), call.Arguments.Path)
require.Equal(t, "payload", call.Arguments.Content)
require.Equal(t, "utf8", call.Arguments.Encoding)
require.Equal(t, "0640", call.Arguments.Mode)
require.True(t, call.Arguments.CreateDirs)
}
type testFilesystemToolCall struct {
Name string
Arguments client.FsWriteArguments
}
func decodeTestFilesystemToolCall(body io.Reader) (testFilesystemToolCall, error) {
var envelope struct {
Params struct {
Name string `json:"name"`
Arguments client.FsWriteArguments `json:"arguments"`
} `json:"params"`
}
if err := json.NewDecoder(body).Decode(&envelope); err != nil {
return testFilesystemToolCall{}, err
}
return testFilesystemToolCall{Name: envelope.Params.Name, Arguments: envelope.Params.Arguments}, nil
}
@@ -0,0 +1,211 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"io"
"net/http"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
type NATInput struct {
Paths contract.Paths
Fault contract.Fault
}
type NAT struct{}
type natForm struct {
Name string `json:"name"`
Enabled bool `json:"enabled"`
ServerID uint64 `json:"server_id"`
Host string `json:"host"`
Domain string `json:"domain"`
}
type natServerListRequest struct {
OnlineOnly bool `json:"online_only"`
}
type natServerListResponse struct {
Servers []struct {
ID uint64 `json:"id"`
UUID string `json:"uuid"`
Online bool `json:"online"`
} `json:"servers"`
}
type natIDResponse uint64
const natTestDomain = "agentcompat-nat.invalid"
const natHalfCloseTestDomain = "half-close.agentcompat-nat.invalid"
func (NAT) Run(ctx context.Context, input NATInput) (result Result, runErr error) {
assertions := NewAssertionSet()
dashboardInstance, err := dashboard.Start(ctx, dashboard.StartConfig{SourceDir: input.Paths.NezhaSource().String(), ReceiptGate: true})
if err != nil {
return Result{Name: "nat", Assertions: assertions.Results(), Error: errorText(err)}, err
}
var ordinaryBackend *fixture.NATEchoBackend
var halfCloseBackend *fixture.NATEchoBackend
var agentInstance *agent.Agent
defer func() {
cleanupContext, cancelCleanup := context.WithTimeout(context.Background(), 30*time.Second)
defer cancelCleanup()
var cleanupErr error
if agentInstance != nil {
cleanupErr = errors.Join(cleanupErr, agentInstance.Stop(cleanupContext))
}
if ordinaryBackend != nil {
cleanupErr = errors.Join(cleanupErr, ordinaryBackend.Close())
}
if halfCloseBackend != nil {
cleanupErr = errors.Join(cleanupErr, halfCloseBackend.Close())
}
cleanupErr = errors.Join(cleanupErr, dashboardInstance.Stop(cleanupContext))
agentCleanupPassed := agentInstance == nil || agentInstance.CleanupReceipt().Passed
result.CleanupOK = cleanupErr == nil && agentCleanupPassed && dashboardInstance.CleanupReceipt().Passed
if cleanupErr != nil && runErr == nil {
runErr = cleanupErr
result.Passed = false
result.Error = errorText(cleanupErr)
}
}()
ordinaryBackend, err = fixture.StartNATEchoBackend()
if err != nil {
return finishNAT(assertions, err)
}
halfCloseBackend, err = fixture.StartNATResponseHalfCloseEchoBackend()
if err != nil {
return finishNAT(assertions, err)
}
agentInstance, err = agent.Start(ctx, agent.AgentStartConfig{SourceDir: input.Paths.AgentSource().String(), Endpoint: dashboardInstance.Endpoint(), Secret: dashboardInstance.AgentSecret(), UUID: "00000000-0000-0000-0000-000000000114"})
if err != nil {
return finishNAT(assertions, err)
}
if err := dashboardInstance.WaitForReceiptAccepted(ctx); err != nil {
return finishNAT(assertions, err)
}
if err := dashboardInstance.ReleaseReceipt(ctx); err != nil {
return finishNAT(assertions, err)
}
if _, err := agentInstance.WaitReady(ctx, dashboardInstance); err != nil {
return finishNAT(assertions, err)
}
serverID, err := natPrimaryServerID(ctx, dashboardInstance, agentInstance.UUID())
if err != nil {
return finishNAT(assertions, err)
}
admin := dashboardInstance.Clients().REST
created, err := client.DoREST[natForm, natIDResponse](ctx, admin, client.RESTRequest[natForm]{Method: http.MethodPost, Path: "/api/v1/nat", Body: &natForm{Name: "agentcompat-nat", Enabled: true, ServerID: serverID, Host: ordinaryBackend.Address(), Domain: natTestDomain}})
if err != nil {
return finishNAT(assertions, err)
}
profileID := uint64(created)
ordinaryRequest := natHTTPRequestSpec{Endpoint: dashboardInstance.Endpoint(), Host: natTestDomain, Method: "PATCH", Path: "/nat?case=ordinary", Body: "ordinary"}
response, responseRecord, err := natHTTPRoundTrip(ctx, ordinaryRequest)
connectionErr := ordinaryBackend.WaitConnection(ctx)
if err == nil {
err = connectionErr
}
record, recordErr := ordinaryBackend.WaitRequest(ctx)
if err == nil {
err = recordErr
}
assertions.Record("enabled profile traverses exact HTTP request and response", err == nil && response.Status == http.StatusOK && response.PeerWriteClosed && response.LegacyCloseMarker == io.EOF.Error() && response.Body == natExpectedBody(ordinaryRequest) && natExactRequestObserved(ordinaryRequest, responseRecord, record), errorText(err))
if err != nil {
return finishNAT(assertions, err)
}
halfCloseProfile, err := client.DoREST[natForm, natIDResponse](ctx, admin, client.RESTRequest[natForm]{Method: http.MethodPost, Path: "/api/v1/nat", Body: &natForm{Name: "agentcompat-nat-half-close", Enabled: true, ServerID: serverID, Host: halfCloseBackend.Address(), Domain: natHalfCloseTestDomain}})
if err != nil {
return finishNAT(assertions, err)
}
// The fixture half-closes its response after writing it. Requiring a client
// request half-close would deadlock HTTP handling because Dashboard keeps the
// request side open while waiting for the backend response.
halfCloseRequest := natHTTPRequestSpec{Endpoint: dashboardInstance.Endpoint(), Host: natHalfCloseTestDomain, Method: "POST", Path: "/nat?case=half-close", Body: "half-closed"}
halfCloseResponse, halfCloseResponseRecord, err := natHTTPRoundTrip(ctx, halfCloseRequest)
halfCloseConnectionErr := halfCloseBackend.WaitConnection(ctx)
if err == nil {
err = halfCloseConnectionErr
}
halfCloseRecord, halfCloseRecordErr := halfCloseBackend.WaitRequest(ctx)
if err == nil {
err = halfCloseRecordErr
}
assertions.Record("backend half-close traverses response and exact request", err == nil && halfCloseResponse.Status == http.StatusOK && halfCloseResponse.PeerWriteClosed && halfCloseResponse.LegacyCloseMarker == io.EOF.Error() && halfCloseResponse.Body == natExpectedBody(halfCloseRequest) && natExactRequestObserved(halfCloseRequest, halfCloseResponseRecord, halfCloseRecord) && halfCloseRecord.ResponseHalfClosed, errorText(err))
if err != nil {
return finishNAT(assertions, err)
}
if _, err := client.DoREST[[]uint64, struct{}](ctx, admin, client.RESTRequest[[]uint64]{Method: http.MethodPost, Path: "/api/v1/batch-delete/nat", Body: &[]uint64{uint64(halfCloseProfile)}}); err != nil {
return finishNAT(assertions, err)
}
limited, err := createScopedClient(ctx, dashboardInstance, []string{"nezha:server:read"})
if err != nil {
return finishNAT(assertions, err)
}
_, unauthorizedErr := client.DoREST[natForm, natIDResponse](ctx, limited, client.RESTRequest[natForm]{Method: http.MethodPost, Path: "/api/v1/nat", Body: &natForm{Name: "unauthorized", Enabled: true, ServerID: serverID, Host: ordinaryBackend.Address(), Domain: "unauthorized." + natTestDomain}})
assertions.Record("unauthorized NAT profile is rejected", isForbidden(unauthorizedErr), errorText(unauthorizedErr))
disabled, err := client.DoREST[natForm, natIDResponse](ctx, admin, client.RESTRequest[natForm]{Method: http.MethodPost, Path: "/api/v1/nat", Body: &natForm{Name: "disabled", Enabled: false, ServerID: serverID, Host: ordinaryBackend.Address(), Domain: "disabled." + natTestDomain}})
if err != nil {
return finishNAT(assertions, err)
}
disabledResponse, err := natHTTPStatus(ctx, dashboardInstance.Endpoint(), "disabled."+natTestDomain)
disabledRouteErr := natAssertNoBackendConnection(ctx, ordinaryBackend)
assertions.Record("disabled NAT profile is blocked", err == nil && disabledRouteErr == nil && disabledResponse.Status == http.StatusForbidden, errorText(errors.Join(err, disabledRouteErr)))
if _, err := client.DoREST[[]uint64, struct{}](ctx, admin, client.RESTRequest[[]uint64]{Method: http.MethodPost, Path: "/api/v1/batch-delete/nat", Body: &[]uint64{uint64(disabled)}}); err != nil {
return finishNAT(assertions, err)
}
if _, err := client.DoREST[[]uint64, struct{}](ctx, admin, client.RESTRequest[[]uint64]{Method: http.MethodPost, Path: "/api/v1/batch-delete/nat", Body: &[]uint64{profileID}}); err != nil {
return finishNAT(assertions, err)
}
deletedResponse, err := natHTTPStatus(ctx, dashboardInstance.Endpoint(), natTestDomain)
deletedRouteErr := natAssertNoBackendConnection(ctx, ordinaryBackend)
assertions.Record("deleted NAT profile no longer routes", err == nil && natDeletedRouteObserved(deletedResponse, deletedRouteErr, natHTTPRequestSpec{Host: natTestDomain, Method: http.MethodGet, Path: "/"}), errorText(errors.Join(err, deletedRouteErr)))
return finishNAT(assertions, nil)
}
func natExactRequestObserved(request natHTTPRequestSpec, responseRecord, backendRecord fixture.NATEchoRecord) bool {
for _, record := range []fixture.NATEchoRecord{responseRecord, backendRecord} {
if record.Method != request.Method || record.Path != request.Path || record.Host != request.Host || record.HeaderValue != natEchoHeaderValue || string(record.Body) != request.Body {
return false
}
}
return true
}
func natDeletedRouteObserved(response natRawResponse, routeErr error, request natHTTPRequestSpec) bool {
return routeErr == nil && response.Status == http.StatusOK && response.Body != "" && response.Body != natExpectedBody(request)
}
func finishNAT(assertions *AssertionSet, runErr error) (Result, error) {
for _, assertion := range assertions.Results() {
if !assertion.Passed && runErr == nil {
runErr = fmt.Errorf("%s: %s", assertion.Name, assertion.Details)
}
}
result := Result{Name: "nat", Passed: runErr == nil, Assertions: assertions.Results(), Error: errorText(runErr), CleanupOK: false}
return result, runErr
}
func natPrimaryServerID(ctx context.Context, dashboardInstance *dashboard.Dashboard, uuid string) (uint64, error) {
response, err := client.CallTool[natServerListRequest, natServerListResponse](ctx, dashboardInstance.Clients().MCP, client.ToolCall[natServerListRequest]{Name: "server.list", Arguments: natServerListRequest{OnlineOnly: true}})
if err != nil {
return 0, err
}
for _, server := range response.StructuredContent.Servers {
if server.UUID == uuid && server.Online {
return server.ID, nil
}
}
return 0, errors.New("primary agent is not online")
}
@@ -0,0 +1,117 @@
//go:build linux
package scenario
import (
"bufio"
"context"
"errors"
"fmt"
"io"
"net"
"net/http"
"strings"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
type natRawResponse struct {
Status int
Body string
LegacyCloseMarker string
PeerWriteClosed bool
}
type natHTTPRequestSpec struct {
Endpoint string
Host string
Method string
Path string
Body string
}
const natEchoHeaderValue = "fixture"
func natHTTPStatus(ctx context.Context, endpoint, host string) (natRawResponse, error) {
return natHTTPRequest(ctx, natHTTPRequestSpec{Endpoint: endpoint, Host: host, Method: http.MethodGet, Path: "/"})
}
func natHTTPRoundTrip(ctx context.Context, request natHTTPRequestSpec) (natRawResponse, fixture.NATEchoRecord, error) {
response, err := natHTTPRequest(ctx, request)
if err != nil {
return natRawResponse{}, fixture.NATEchoRecord{}, err
}
record, err := parseNATResponse(response.Body)
return response, record, err
}
func natHTTPRequest(ctx context.Context, request natHTTPRequestSpec) (natRawResponse, error) {
connection, err := (&net.Dialer{}).DialContext(ctx, "tcp", request.Endpoint)
if err != nil {
return natRawResponse{}, err
}
defer connection.Close()
if deadline, ok := ctx.Deadline(); ok {
if err := connection.SetDeadline(deadline); err != nil {
return natRawResponse{}, err
}
}
wireRequest := fmt.Sprintf("%s %s HTTP/1.1\r\nHost: %s\r\nX-AgentCompat-Echo: %s\r\nContent-Length: %d\r\nConnection: close\r\n\r\n%s", request.Method, request.Path, request.Host, natEchoHeaderValue, len(request.Body), request.Body)
if _, err := io.WriteString(connection, wireRequest); err != nil {
return natRawResponse{}, err
}
reader := bufio.NewReader(connection)
response, err := http.ReadResponse(reader, nil)
if err != nil {
return natRawResponse{}, err
}
responseBody, err := io.ReadAll(response.Body)
if closeErr := response.Body.Close(); err == nil {
err = closeErr
}
if err != nil {
return natRawResponse{}, err
}
// Agent preserves its legacy wire contract by forwarding the local read
// result after the HTTP body. Reading through that marker to EOF proves the
// backend write-close reached Agent and the stream then terminated.
closeMarker, peerCloseErr := io.ReadAll(reader)
if peerCloseErr != nil {
return natRawResponse{}, fmt.Errorf("observe NAT peer write close: %w", peerCloseErr)
}
return natRawResponse{Status: response.StatusCode, Body: string(responseBody), LegacyCloseMarker: string(closeMarker), PeerWriteClosed: true}, nil
}
func natAssertNoBackendConnection(ctx context.Context, backend *fixture.NATEchoBackend) error {
deadline, cancel := context.WithTimeout(ctx, 300*time.Millisecond)
defer cancel()
err := backend.WaitConnection(deadline)
if errors.Is(err, context.DeadlineExceeded) {
return nil
}
if err == nil {
return errors.New("NAT backend received an unexpected connection")
}
return err
}
func parseNATResponse(body string) (fixture.NATEchoRecord, error) {
values := make(map[string]string)
for _, line := range strings.Split(strings.TrimSuffix(body, "\n"), "\n") {
key, value, ok := strings.Cut(line, "=")
if ok {
values[key] = value
}
}
for _, key := range []string{"method", "path", "host", "x-agentcompat-echo", "body"} {
if _, ok := values[key]; !ok {
return fixture.NATEchoRecord{}, fmt.Errorf("NAT response missing %s", key)
}
}
return fixture.NATEchoRecord{Method: values["method"], Path: values["path"], Host: values["host"], HeaderValue: values["x-agentcompat-echo"], Body: []byte(values["body"])}, nil
}
func natExpectedBody(request natHTTPRequestSpec) string {
return fmt.Sprintf("method=%s\npath=%s\nhost=%s\nx-agentcompat-echo=%s\nbody=%s\n", request.Method, request.Path, request.Host, natEchoHeaderValue, request.Body)
}
@@ -0,0 +1,154 @@
//go:build linux && agentcompat
package scenario
import (
"bufio"
"context"
"errors"
"fmt"
"io"
"net"
"net/http"
"testing"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
"github.com/stretchr/testify/require"
)
func TestNAT_HTTPRequestObservesPeerWriteCloseAfterDeclaredBody(t *testing.T) {
// Given
listener, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, listener.Close()) })
serverDone := make(chan error, 1)
go func() {
connection, acceptErr := listener.Accept()
if acceptErr != nil {
serverDone <- acceptErr
return
}
defer connection.Close()
request, readErr := http.ReadRequest(bufio.NewReader(connection))
if readErr != nil {
serverDone <- readErr
return
}
_, readErr = io.Copy(io.Discard, request.Body)
if closeErr := request.Body.Close(); readErr == nil {
readErr = closeErr
}
if readErr != nil {
serverDone <- readErr
return
}
body := "closed"
if _, writeErr := fmt.Fprintf(connection, "HTTP/1.1 200 OK\r\nContent-Length: %d\r\n\r\n%s%s", len(body), body, io.EOF.Error()); writeErr != nil {
serverDone <- writeErr
return
}
tcpConnection, ok := connection.(*net.TCPConn)
if !ok {
serverDone <- errors.New("test listener did not accept TCP connection")
return
}
serverDone <- tcpConnection.CloseWrite()
}()
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
// When
response, err := natHTTPRequest(ctx, natHTTPRequestSpec{Endpoint: listener.Addr().String(), Host: natTestDomain, Method: http.MethodGet, Path: "/"})
// Then
require.NoError(t, err)
require.Equal(t, http.StatusOK, response.Status)
require.Equal(t, "closed", response.Body)
require.Equal(t, io.EOF.Error(), response.LegacyCloseMarker)
require.True(t, response.PeerWriteClosed)
require.NoError(t, <-serverDone)
}
func TestNAT_ParseResponsePreservesExactRequestFields(t *testing.T) {
// Given
body := natExpectedBody(natHTTPRequestSpec{Host: natTestDomain, Method: http.MethodPatch, Path: "/nat?case=ordinary", Body: "ordinary"})
// When
record, err := parseNATResponse(body)
// Then
require.NoError(t, err)
require.Equal(t, http.MethodPatch, record.Method)
require.Equal(t, "/nat?case=ordinary", record.Path)
require.Equal(t, natTestDomain, record.Host)
require.Equal(t, "fixture", record.HeaderValue)
require.Equal(t, []byte("ordinary"), record.Body)
}
func TestNAT_ParseResponseRejectsMissingEvidence(t *testing.T) {
// Given
body := "method=GET\npath=/\n"
// When
_, err := parseNATResponse(body)
// Then
require.Error(t, err)
}
func TestNAT_ExactRequestObservedRejectsMismatchedHalfCloseEvidence(t *testing.T) {
// Given
request := natHTTPRequestSpec{Host: natHalfCloseTestDomain, Method: http.MethodPost, Path: "/nat?case=half-close", Body: "half-closed"}
exact := fixture.NATEchoRecord{Method: request.Method, Path: request.Path, Host: request.Host, HeaderValue: natEchoHeaderValue, Body: []byte(request.Body)}
mismatched := exact
mismatched.Host = natTestDomain
// When
observed := natExactRequestObserved(request, exact, mismatched)
// Then
require.False(t, observed)
}
func TestNAT_DeletedRouteObservedRequiresFallbackStatusAndBody(t *testing.T) {
// Given
request := natHTTPRequestSpec{Host: natTestDomain, Method: http.MethodGet, Path: "/"}
// When
observed := natDeletedRouteObserved(natRawResponse{Status: http.StatusOK, Body: "dashboard fallback"}, nil, request)
echoObserved := natDeletedRouteObserved(natRawResponse{Status: http.StatusOK, Body: natExpectedBody(request)}, nil, request)
rejectedObserved := natDeletedRouteObserved(natRawResponse{Status: http.StatusNotFound, Body: "not found"}, nil, request)
// Then
require.True(t, observed)
require.False(t, echoObserved)
require.False(t, rejectedObserved)
}
func TestNAT_FinishReturnsTypedFailedAssertion(t *testing.T) {
// Given
assertions := NewAssertionSet()
assertions.Record("deleted profile no longer routes", false, "backend connection observed")
// When
result, err := finishNAT(assertions, nil)
// Then
require.EqualError(t, err, "deleted profile no longer routes: backend connection observed")
require.Equal(t, "nat", result.Name)
require.False(t, result.Passed)
require.False(t, result.CleanupOK)
}
func TestNAT_FinishPreservesRuntimeError(t *testing.T) {
// Given
runtimeErr := errors.New("NAT runtime failed")
// When
result, err := finishNAT(NewAssertionSet(), runtimeErr)
// Then
require.ErrorIs(t, err, runtimeErr)
require.Equal(t, "NAT runtime failed", result.Error)
}
@@ -0,0 +1,75 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"slices"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
)
type ReconnectInput struct {
Paths contract.Paths
DashboardFault string
}
type ReconnectObservation struct {
ServerID uint64
UUID string
OldGeneration uint64
NewGeneration uint64
DisconnectAt time.Time
ReconnectAt time.Time
TaskIDs []uint64
ResultIDs []uint64
PostReconnect bool
AgentRestarted bool
}
type ReconnectResult struct {
Observation ReconnectObservation `json:"observation"`
CleanupOK bool `json:"cleanup_ok"`
Passed bool `json:"passed"`
Error string `json:"error,omitempty"`
}
type Reconnect struct{}
func (scenario Reconnect) Run(ctx context.Context, input ReconnectInput) (Result, error) {
result, _, err := scenario.RunWithEvidence(ctx, input)
return result, err
}
func (observation ReconnectObservation) Validate() error {
if observation.ServerID == 0 || observation.UUID == "" {
return errors.New("reconnect observation omitted server identity")
}
if observation.OldGeneration == 0 || observation.NewGeneration <= observation.OldGeneration {
return errors.New("reconnect generations are not strictly increasing")
}
if observation.DisconnectAt.IsZero() || observation.ReconnectAt.IsZero() || !observation.ReconnectAt.After(observation.DisconnectAt) {
return errors.New("reconnect timestamps are not ordered")
}
if len(observation.TaskIDs) == 0 || !slices.Equal(observation.TaskIDs, observation.ResultIDs) {
return errors.New("reconnect task and result IDs differ")
}
seenTaskIDs := make(map[uint64]struct{}, len(observation.TaskIDs))
for _, taskID := range observation.TaskIDs {
if _, exists := seenTaskIDs[taskID]; exists {
return fmt.Errorf("reconnect task ID %d was duplicated", taskID)
}
seenTaskIDs[taskID] = struct{}{}
}
if !observation.PostReconnect || !observation.AgentRestarted {
return errors.New("reconnect post-process checks were not completed")
}
return nil
}
func (observation ReconnectObservation) ReconnectInterval() time.Duration {
return observation.ReconnectAt.Sub(observation.DisconnectAt)
}
@@ -0,0 +1,82 @@
//go:build linux
package scenario
import (
"errors"
"fmt"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process"
)
type ReconnectFixtureEvidence struct {
Dashboard dashboard.FixtureIdentity `json:"dashboard"`
AgentRoot string `json:"agent_root"`
AgentConfigPath string `json:"agent_config_path"`
AgentBinaryPath string `json:"agent_binary_path"`
}
type ReconnectRuntimeEvidence struct {
DashboardBefore dashboard.RuntimeIdentity `json:"dashboard_before"`
DashboardAfter dashboard.RuntimeIdentity `json:"dashboard_after"`
AgentBefore agent.ProcessIdentity `json:"agent_before"`
AgentAfter agent.ProcessIdentity `json:"agent_after"`
StateGenerationBeforeAgentRestart uint64 `json:"state_generation_before_agent_restart"`
StateGenerationAfterAgentRestart uint64 `json:"state_generation_after_agent_restart"`
}
type ReconnectIdentityEvidence struct {
ServerID uint64 `json:"server_id"`
UUID string `json:"uuid"`
DashboardConfigUnchanged bool `json:"dashboard_config_unchanged"`
AgentConfigUnchanged bool `json:"agent_config_unchanged"`
DashboardFixtureUnchanged bool `json:"dashboard_fixture_unchanged"`
ClientsRecreated bool `json:"clients_recreated"`
BootstrapRecreated bool `json:"bootstrap_recreated"`
}
type ReconnectLifecycleEvidence struct {
DisconnectAt time.Time `json:"disconnect_at"`
ReconnectAt time.Time `json:"reconnect_at"`
ReconnectInterval time.Duration `json:"reconnect_interval"`
DashboardReceipts []dashboard.MCPReceiptPair `json:"dashboard_receipts"`
AgentReceipts []dashboard.MCPReceiptPair `json:"agent_receipts"`
StaleGenerationReceipts int `json:"stale_generation_receipts"`
DuplicateTaskIDs int `json:"duplicate_task_ids"`
LostResultIDs int `json:"lost_result_ids"`
OutsideRootSentinelUnchanged bool `json:"outside_root_sentinel_unchanged"`
}
type ReconnectEvidence struct {
Fixture ReconnectFixtureEvidence `json:"fixture"`
Runtime ReconnectRuntimeEvidence `json:"runtime"`
Identity ReconnectIdentityEvidence `json:"identity"`
Lifecycle ReconnectLifecycleEvidence `json:"lifecycle"`
Observation ReconnectObservation `json:"observation"`
AgentCleanup processharness.CleanupReceipt `json:"agent_cleanup"`
DashboardCleanup processharness.CleanupReceipt `json:"dashboard_cleanup"`
}
func (e ReconnectEvidence) Validate() error {
var validationErr error
validationErr = errors.Join(validationErr, e.Observation.Validate())
if e.Runtime.DashboardAfter.Generation <= e.Runtime.DashboardBefore.Generation || e.Runtime.DashboardAfter.PID == e.Runtime.DashboardBefore.PID {
validationErr = errors.Join(validationErr, errors.New("Dashboard runtime generation did not advance"))
}
if e.Runtime.AgentAfter.Generation <= e.Runtime.AgentBefore.Generation || e.Runtime.AgentAfter.PID == e.Runtime.AgentBefore.PID {
validationErr = errors.Join(validationErr, errors.New("Agent runtime generation did not advance"))
}
if e.Runtime.StateGenerationAfterAgentRestart <= e.Runtime.StateGenerationBeforeAgentRestart {
validationErr = errors.Join(validationErr, errors.New("Agent state stream generation did not advance"))
}
if !e.Identity.DashboardConfigUnchanged || !e.Identity.AgentConfigUnchanged || !e.Identity.DashboardFixtureUnchanged || !e.Identity.ClientsRecreated || !e.Identity.BootstrapRecreated {
validationErr = errors.Join(validationErr, errors.New("reconnect identity evidence is incomplete"))
}
if e.Lifecycle.ReconnectInterval <= 0 || e.Lifecycle.StaleGenerationReceipts != 0 || e.Lifecycle.DuplicateTaskIDs != 0 || e.Lifecycle.LostResultIDs != 0 || !e.Lifecycle.OutsideRootSentinelUnchanged {
validationErr = errors.Join(validationErr, fmt.Errorf("reconnect lifecycle evidence is invalid: interval=%s stale=%d duplicate=%d lost=%d sentinel=%t", e.Lifecycle.ReconnectInterval, e.Lifecycle.StaleGenerationReceipts, e.Lifecycle.DuplicateTaskIDs, e.Lifecycle.LostResultIDs, e.Lifecycle.OutsideRootSentinelUnchanged))
}
return validationErr
}
@@ -0,0 +1,106 @@
//go:build linux && agentcompat
package scenario
import (
"context"
"encoding/json"
"net"
"os"
"path/filepath"
"strconv"
"testing"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process"
"github.com/stretchr/testify/require"
)
func TestReconnectScenario_DashboardExitFaultCleansRealProcesses(t *testing.T) {
// Given
nezhaSource := os.Getenv("AGENTCOMPAT_NEZHA_SOURCE")
agentSource := os.Getenv("AGENTCOMPAT_AGENT_SOURCE")
if nezhaSource == "" || agentSource == "" {
t.Skip("set AGENTCOMPAT_NEZHA_SOURCE and AGENTCOMPAT_AGENT_SOURCE")
}
evidenceDirectory := os.Getenv("AGENTCOMPAT_RECONNECT_FAULT_EVIDENCE_DIR")
if evidenceDirectory == "" {
evidenceDirectory = t.TempDir()
}
require.NoError(t, os.MkdirAll(evidenceDirectory, 0o700))
paths, err := contract.NewPaths(nezhaSource, agentSource, evidenceDirectory)
require.NoError(t, err)
testContext, cancel := context.WithTimeout(t.Context(), 10*time.Minute)
defer cancel()
// When
result, reconnectEvidence, runErr := (Reconnect{}).RunWithEvidence(testContext, ReconnectInput{Paths: paths, DashboardFault: "dashboard-exit"})
// Then
require.ErrorIs(t, runErr, ErrReconnectDashboardExitFault)
require.Equal(t, reconnectScenarioName, result.Name)
require.False(t, result.Passed)
require.True(t, result.CleanupOK)
require.Contains(t, result.Error, ErrReconnectDashboardExitFault.Error())
require.Len(t, result.Assertions, 3)
require.True(t, result.Assertions[0].Passed)
require.Equal(t, "Dashboard disconnect barrier stopped generation one", result.Assertions[0].Name)
require.True(t, result.Assertions[1].Passed)
require.Equal(t, "outside-root sentinel remains unchanged", result.Assertions[1].Name)
require.True(t, result.Assertions[2].Passed)
require.Equal(t, "multi-generation process listener and workspace cleanup completed", result.Assertions[2].Name)
require.True(t, reconnectEvidence.Lifecycle.OutsideRootSentinelUnchanged)
require.True(t, reconnectEvidence.AgentCleanup.Passed)
require.False(t, reconnectEvidence.AgentCleanup.Forced)
require.Len(t, reconnectEvidence.AgentCleanup.Processes, 1)
require.True(t, reconnectEvidence.DashboardCleanup.Passed)
require.False(t, reconnectEvidence.DashboardCleanup.Forced)
require.Len(t, reconnectEvidence.DashboardCleanup.Processes, 1)
require.Zero(t, reconnectEvidence.Runtime.DashboardAfter)
require.Zero(t, reconnectEvidence.Runtime.AgentAfter)
require.NoDirExists(t, reconnectEvidence.Fixture.AgentRoot)
require.NoDirExists(t, reconnectEvidence.Fixture.Dashboard.WorkspaceRoot)
for _, cleanupReceipt := range []struct {
name string
processes []processharness.CleanupRecord
}{
{name: "Agent", processes: reconnectEvidence.AgentCleanup.Processes},
{name: "Dashboard", processes: reconnectEvidence.DashboardCleanup.Processes},
} {
for _, process := range cleanupReceipt.processes {
require.NoDirExists(t, filepath.Join("/proc", strconv.Itoa(process.PID)), "%s process %d survived cleanup", cleanupReceipt.name, process.PID)
}
}
for _, listenerIdentity := range []struct {
name string
address string
}{
{name: "Dashboard HTTP", address: reconnectEvidence.Fixture.Dashboard.HTTP.Address},
{name: "Dashboard receipt", address: reconnectEvidence.Fixture.Dashboard.Receipt.Address},
} {
listener, listenErr := net.Listen("tcp", listenerIdentity.address)
require.NoError(t, listenErr, "%s listener was not released", listenerIdentity.name)
require.NoError(t, listener.Close())
}
artifactPath := filepath.Join(evidenceDirectory, "reconnect-dashboard-exit-real-process.json")
type faultArtifact struct {
Result Result `json:"result"`
Evidence ReconnectEvidence `json:"evidence"`
Error string `json:"error"`
}
recordedArtifact := faultArtifact{Result: result, Evidence: reconnectEvidence, Error: runErr.Error()}
artifact, err := json.MarshalIndent(recordedArtifact, "", " ")
require.NoError(t, err)
require.NoError(t, os.WriteFile(artifactPath, append(artifact, '\n'), 0o600))
artifactInfo, err := os.Stat(artifactPath)
require.NoError(t, err)
require.Equal(t, os.FileMode(0o600), artifactInfo.Mode().Perm())
readArtifact, err := os.ReadFile(artifactPath)
require.NoError(t, err)
var decodedArtifact faultArtifact
require.NoError(t, json.Unmarshal(readArtifact, &decodedArtifact))
require.Equal(t, recordedArtifact, decodedArtifact)
t.Logf("reconnect Dashboard-exit fault artifact: %s", artifactPath)
}
@@ -0,0 +1,126 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"os"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
"github.com/nezhahq/nezha/model"
)
type reconnectExecArguments struct {
ServerID uint64 `json:"server_id"`
Cmd string `json:"cmd"`
Args []string `json:"args"`
}
type reconnectExecResult struct {
ExitCode int `json:"exit_code"`
Stdout string `json:"stdout"`
Stderr string `json:"stderr"`
Error string `json:"error"`
}
func runDashboardReconnectOperations(ctx context.Context, dashboardInstance *dashboard.Dashboard, serverID uint64, uuid, fixturePath string) ([]dashboard.MCPReceiptPair, error) {
cursor := dashboardInstance.MCPReceiptCursor()
server, err := client.CallTool[client.ServerGetArguments, client.ServerGetResult](ctx, dashboardInstance.Clients().MCP, client.ToolCall[client.ServerGetArguments]{Name: "server.get", Arguments: client.ServerGetArguments{ServerID: serverID}})
if err != nil || server.StructuredContent.ID != serverID || server.StructuredContent.UUID != uuid || string(server.StructuredContent.Host) == "null" || string(server.StructuredContent.State) == "null" {
return nil, errors.Join(errors.New("post-reconnect server.get identity mismatch"), err)
}
if err := runReconnectExec(ctx, dashboardInstance.Clients().MCP, serverID, "dashboard-reconnect"); err != nil {
return nil, err
}
write, err := client.CallTool[client.FsWriteArguments, client.FsWriteResult](ctx, dashboardInstance.Clients().MCP, client.ToolCall[client.FsWriteArguments]{Name: "fs.write", Arguments: client.FsWriteArguments{ServerID: serverID, Path: fixturePath, Content: "dashboard-generation-two", Encoding: "utf8", Mode: "0600", CreateDirs: true}})
if err != nil || write.StructuredContent.Size != int64(len("dashboard-generation-two")) || write.StructuredContent.Error != "" {
return nil, errors.Join(errors.New("post-reconnect fs.write mismatch"), err)
}
if err := runReconnectRead(ctx, dashboardInstance.Clients().MCP, serverID, fixturePath, "dashboard-generation-two"); err != nil {
return nil, err
}
generation := dashboardInstance.RuntimeIdentity().Generation
expectations, err := reconnectReceiptExpectations(dashboardInstance.MCPReceiptEventsAfter(cursor), generation, serverID, []uint64{model.TaskTypeExec, model.TaskTypeFsWrite, model.TaskTypeFsRead})
if err != nil {
return nil, err
}
return dashboardInstance.WaitForMCPReceiptSet(ctx, cursor, expectations)
}
func runAgentRestartOperations(ctx context.Context, dashboardInstance *dashboard.Dashboard, serverID uint64, fixturePath string) ([]dashboard.MCPReceiptPair, error) {
cursor := dashboardInstance.MCPReceiptCursor()
if err := runReconnectExec(ctx, dashboardInstance.Clients().MCP, serverID, "agent-restart"); err != nil {
return nil, err
}
if err := runReconnectRead(ctx, dashboardInstance.Clients().MCP, serverID, fixturePath, "dashboard-generation-two"); err != nil {
return nil, err
}
generation := dashboardInstance.RuntimeIdentity().Generation
expectations, err := reconnectReceiptExpectations(dashboardInstance.MCPReceiptEventsAfter(cursor), generation, serverID, []uint64{model.TaskTypeExec, model.TaskTypeFsRead})
if err != nil {
return nil, err
}
return dashboardInstance.WaitForMCPReceiptSet(ctx, cursor, expectations)
}
func reconnectReceiptExpectations(events []dashboard.MCPReceiptEvent, generation, serverID uint64, taskTypes []uint64) ([]dashboard.MCPReceiptExpectation, error) {
pending := append([]uint64(nil), taskTypes...)
expectations := make([]dashboard.MCPReceiptExpectation, 0, len(pending))
for _, event := range events {
if event.Kind != dashboard.MCPReceiptTask || event.DashboardGeneration != generation || event.ServerID != serverID {
continue
}
for index, taskType := range pending {
if taskType == event.TaskType {
expectations = append(expectations, dashboard.MCPReceiptExpectation{DashboardGeneration: event.DashboardGeneration, GateGeneration: event.GateGeneration, ServerID: event.ServerID, TaskID: event.TaskID, TaskType: event.TaskType})
pending = append(pending[:index], pending[index+1:]...)
break
}
}
}
if len(pending) != 0 {
return nil, errors.New("reconnect receipt task set is incomplete")
}
return expectations, nil
}
func runReconnectExec(ctx context.Context, mcpClient *client.Client, serverID uint64, marker string) error {
result, err := client.CallTool[reconnectExecArguments, reconnectExecResult](ctx, mcpClient, client.ToolCall[reconnectExecArguments]{Name: "server.exec", Arguments: reconnectExecArguments{ServerID: serverID, Cmd: "/bin/sh", Args: []string{"-c", "printf " + marker}}})
if err != nil {
return err
}
if result.StructuredContent.ExitCode != 0 || result.StructuredContent.Stdout != marker || result.StructuredContent.Stderr != "" || result.StructuredContent.Error != "" {
return fmt.Errorf("reconnect Exec mismatch: %+v", result.StructuredContent)
}
return nil
}
func runReconnectRead(ctx context.Context, mcpClient *client.Client, serverID uint64, path, expected string) error {
result, err := client.CallTool[client.FsReadArguments, client.FsReadResult](ctx, mcpClient, client.ToolCall[client.FsReadArguments]{Name: "fs.read", Arguments: client.FsReadArguments{ServerID: serverID, Path: path, Encoding: "utf8"}})
if err != nil {
return err
}
if result.StructuredContent.Content != expected || result.StructuredContent.Encoding != "utf8" || result.StructuredContent.Size != int64(len(expected)) || result.StructuredContent.Truncated {
return fmt.Errorf("reconnect fs.read mismatch: %+v", result.StructuredContent)
}
return nil
}
func prepareReconnectSentinel(agentRoot string) (fixturePath, sentinelPath string, err error) {
root, err := fixture.NewAgentRoot(agentRoot, "reconnect-files")
if err != nil {
return "", "", err
}
path, err := root.Path("runtime.txt")
if err != nil {
return "", "", err
}
fixturePath = path.String()
sentinelPath = agentRoot + "/outside-reconnect-sentinel"
err = os.WriteFile(sentinelPath, []byte("outside-reconnect-root-sentinel"), 0o600)
return fixturePath, sentinelPath, err
}
@@ -0,0 +1,83 @@
//go:build linux && agentcompat
package scenario
import (
"context"
"encoding/json"
"os"
"path/filepath"
"testing"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/stretchr/testify/require"
)
func TestReconnectScenario_RealDashboardAndAgentProcessRestarts(t *testing.T) {
// Given
nezhaSource := os.Getenv("AGENTCOMPAT_NEZHA_SOURCE")
agentSource := os.Getenv("AGENTCOMPAT_AGENT_SOURCE")
if nezhaSource == "" || agentSource == "" {
t.Skip("set AGENTCOMPAT_NEZHA_SOURCE and AGENTCOMPAT_AGENT_SOURCE")
}
evidenceDirectory := os.Getenv("AGENTCOMPAT_RECONNECT_EVIDENCE_DIR")
if evidenceDirectory == "" {
evidenceDirectory = t.TempDir()
}
require.NoError(t, os.MkdirAll(evidenceDirectory, 0o700))
paths, err := contract.NewPaths(nezhaSource, agentSource, evidenceDirectory)
require.NoError(t, err)
testContext, cancel := context.WithTimeout(t.Context(), 10*time.Minute)
defer cancel()
// When
result, reconnectEvidence, err := (Reconnect{}).RunWithEvidence(testContext, ReconnectInput{Paths: paths})
// Then
for _, assertion := range result.Assertions {
t.Logf("assertion=%q passed=%t details=%q", assertion.Name, assertion.Passed, assertion.Details)
}
t.Logf("agent cleanup=%+v dashboard cleanup=%+v", reconnectEvidence.AgentCleanup, reconnectEvidence.DashboardCleanup)
require.NoError(t, err)
require.True(t, result.Passed)
require.True(t, result.CleanupOK)
require.NoError(t, reconnectEvidence.Validate())
require.True(t, reconnectEvidence.AgentCleanup.Passed)
require.False(t, reconnectEvidence.AgentCleanup.Forced)
require.Len(t, reconnectEvidence.AgentCleanup.Processes, 2)
require.True(t, reconnectEvidence.DashboardCleanup.Passed)
require.False(t, reconnectEvidence.DashboardCleanup.Forced)
require.Len(t, reconnectEvidence.DashboardCleanup.Processes, 2)
require.True(t, reconnectEvidence.Identity.DashboardFixtureUnchanged)
require.NotZero(t, reconnectEvidence.Fixture.Dashboard.HTTP.Inode)
require.NotZero(t, reconnectEvidence.Fixture.Dashboard.Receipt.Inode)
require.Greater(t, reconnectEvidence.Runtime.DashboardAfter.Generation, reconnectEvidence.Runtime.DashboardBefore.Generation)
require.NotEqual(t, reconnectEvidence.Runtime.DashboardBefore.PID, reconnectEvidence.Runtime.DashboardAfter.PID)
require.Greater(t, reconnectEvidence.Runtime.AgentAfter.Generation, reconnectEvidence.Runtime.AgentBefore.Generation)
require.NotEqual(t, reconnectEvidence.Runtime.AgentBefore.PID, reconnectEvidence.Runtime.AgentAfter.PID)
require.Positive(t, reconnectEvidence.Lifecycle.ReconnectInterval)
require.Zero(t, reconnectEvidence.Lifecycle.StaleGenerationReceipts)
require.Zero(t, reconnectEvidence.Lifecycle.DuplicateTaskIDs)
require.Zero(t, reconnectEvidence.Lifecycle.LostResultIDs)
require.Len(t, reconnectEvidence.Observation.TaskIDs, 5)
require.Equal(t, reconnectEvidence.Observation.TaskIDs, reconnectEvidence.Observation.ResultIDs)
artifactPath := filepath.Join(evidenceDirectory, "reconnect-real-process.json")
artifact, err := json.MarshalIndent(struct {
Result Result `json:"result"`
Evidence ReconnectEvidence `json:"evidence"`
}{Result: result, Evidence: reconnectEvidence}, "", " ")
require.NoError(t, err)
require.NoError(t, os.WriteFile(artifactPath, append(artifact, '\n'), 0o600))
readArtifact, err := os.ReadFile(artifactPath)
require.NoError(t, err)
var recorded struct {
Result Result `json:"result"`
Evidence ReconnectEvidence `json:"evidence"`
}
require.NoError(t, json.Unmarshal(readArtifact, &recorded))
require.Equal(t, result, recorded.Result)
require.Equal(t, reconnectEvidence, recorded.Evidence)
t.Logf("reconnect evidence artifact: %s", artifactPath)
}
@@ -0,0 +1,35 @@
//go:build linux
package scenario
import "github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
func reconnectReceiptSummary(pairs ...[]dashboard.MCPReceiptPair) (taskIDs, resultIDs []uint64, duplicates, lost int) {
seen := make(map[uint64]struct{})
for _, group := range pairs {
for _, pair := range group {
taskIDs = append(taskIDs, pair.Task.TaskID)
resultIDs = append(resultIDs, pair.Result.TaskID)
if _, exists := seen[pair.Task.TaskID]; exists {
duplicates++
}
seen[pair.Task.TaskID] = struct{}{}
if pair.Task.TaskID == 0 || pair.Task.TaskID != pair.Result.TaskID {
lost++
}
}
}
return taskIDs, resultIDs, duplicates, lost
}
func staleReconnectReceiptCount(generation uint64, pairs ...[]dashboard.MCPReceiptPair) int {
count := 0
for _, group := range pairs {
for _, pair := range group {
if pair.Task.DashboardGeneration == generation || pair.Result.DashboardGeneration == generation {
count++
}
}
}
return count
}
@@ -0,0 +1,212 @@
//go:build linux
package scenario
import (
"bytes"
"context"
"errors"
"os"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process"
)
const reconnectScenarioName = "reconnect"
var ErrReconnectDashboardExitFault = errors.New("reconnect scenario: injected Dashboard exit")
func (Reconnect) RunWithEvidence(ctx context.Context, input ReconnectInput) (result Result, reconnectEvidence ReconnectEvidence, runErr error) {
assertions := NewAssertionSet()
dashboardInstance, err := dashboard.Start(ctx, dashboard.StartConfig{SourceDir: input.Paths.NezhaSource().String(), ReceiptGate: true})
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
dashboardRoot := dashboardInstance.WorkspaceRoot()
var agentInstance *agent.Agent
var agentRoot string
defer func() {
cleanupContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second)
defer cancel()
cleanupErr := stopTransferProcesses(cleanupContext, agentInstance, dashboardInstance)
reconnectEvidence.AgentCleanup = agentCleanupReceipt(agentInstance)
reconnectEvidence.DashboardCleanup = dashboardInstance.CleanupReceipt()
cleanupErr = errors.Join(cleanupErr, transferWorkspaceResidue(agentRoot, dashboardRoot))
assertions.Record("multi-generation process listener and workspace cleanup completed", cleanupErr == nil, errorText(cleanupErr))
result.CleanupOK = cleanupErr == nil
if cleanupErr != nil {
result, reconnectEvidence, runErr = reconnectFinish(assertions, errors.Join(runErr, cleanupErr), reconnectEvidence)
result.CleanupOK = false
return
}
result.Assertions = assertions.Results()
}()
const agentUUID = "00000000-0000-0000-0000-000000000217"
agentInstance, err = agent.Start(ctx, agent.AgentStartConfig{SourceDir: input.Paths.AgentSource().String(), Endpoint: dashboardInstance.Endpoint(), Secret: dashboardInstance.AgentSecret(), UUID: agentUUID})
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
agentRoot = agentInstance.WorkspaceRoot()
if err := dashboardInstance.WaitForReceiptAccepted(ctx); err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
if err := dashboardInstance.ReleaseReceipt(ctx); err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
readiness, err := agentInstance.WaitReady(ctx, dashboardInstance)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
serverID, err := transferServerID(ctx, dashboardInstance.Clients().MCP, readiness.UUID)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
fixturePath, sentinelPath, err := prepareReconnectSentinel(agentRoot)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
sentinelBytes, err := os.ReadFile(sentinelPath)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
dashboardConfig, err := os.ReadFile(dashboardInstance.ConfigPath())
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
agentConfig, err := os.ReadFile(agentInstance.ConfigPath())
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
fixtureBefore := dashboardInstance.FixtureIdentity()
runtimeBefore := dashboardInstance.RuntimeIdentity()
agentBefore := agentInstance.RuntimeIdentity()
clientsBefore := dashboardInstance.Clients()
bootstrapBefore := dashboardInstance.Bootstrap()
reconnectEvidence.Fixture = ReconnectFixtureEvidence{Dashboard: fixtureBefore, AgentRoot: agentRoot, AgentConfigPath: agentInstance.ConfigPath(), AgentBinaryPath: agentInstance.BinaryPath()}
reconnectEvidence.Runtime.DashboardBefore = runtimeBefore
reconnectEvidence.Runtime.AgentBefore = agentBefore
reconnectEvidence.Identity = ReconnectIdentityEvidence{ServerID: serverID, UUID: agentUUID}
stoppedRuntime, err := dashboardInstance.StopProcess(ctx)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
disconnectAt := time.Now().UTC()
assertions.Record("Dashboard disconnect barrier stopped generation one", stoppedRuntime == runtimeBefore && dashboardInstance.RuntimeIdentity().PID == 0, "")
if input.DashboardFault == "dashboard-exit" {
// This fault returns before the normal lifecycle evidence finalization below.
sentinelAfter, sentinelErr := os.ReadFile(sentinelPath)
sentinelUnchanged := sentinelErr == nil && bytes.Equal(sentinelBytes, sentinelAfter)
reconnectEvidence.Lifecycle.DisconnectAt = disconnectAt
reconnectEvidence.Lifecycle.OutsideRootSentinelUnchanged = sentinelUnchanged
assertions.Record("outside-root sentinel remains unchanged", sentinelUnchanged, errorText(sentinelErr))
if sentinelErr != nil {
return reconnectFinish(assertions, errors.Join(ErrReconnectDashboardExitFault, sentinelErr), reconnectEvidence)
}
if !sentinelUnchanged {
return reconnectFinish(assertions, errors.Join(ErrReconnectDashboardExitFault, errors.New("outside-root sentinel changed")), reconnectEvidence)
}
return reconnectFinish(assertions, ErrReconnectDashboardExitFault, reconnectEvidence)
}
runtimeAfter, err := dashboardInstance.StartProcess(ctx)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
if err := dashboardInstance.WaitForReceiptAccepted(ctx); err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
if err := dashboardInstance.ReleaseReceipt(ctx); err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
postDashboardReadiness, err := agentInstance.WaitReady(ctx, dashboardInstance)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
reconnectAt := time.Now().UTC()
serverIDAfter, err := transferServerID(ctx, dashboardInstance.Clients().MCP, postDashboardReadiness.UUID)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
dashboardConfigAfter, err := os.ReadFile(dashboardInstance.ConfigPath())
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
fixtureAfter := dashboardInstance.FixtureIdentity()
clientsAfter := dashboardInstance.Clients()
bootstrapAfter := dashboardInstance.Bootstrap()
reconnectEvidence.Runtime.DashboardAfter = runtimeAfter
reconnectEvidence.Identity.DashboardConfigUnchanged = bytes.Equal(dashboardConfig, dashboardConfigAfter)
reconnectEvidence.Identity.DashboardFixtureUnchanged = fixtureAfter == fixtureBefore
reconnectEvidence.Identity.ClientsRecreated = clientsAfter.REST != clientsBefore.REST && clientsAfter.MCP != clientsBefore.MCP && clientsAfter.WebSocket != clientsBefore.WebSocket
reconnectEvidence.Identity.BootstrapRecreated = bootstrapAfter.PATID != 0 && bootstrapAfter.PATID != bootstrapBefore.PATID && bootstrapAfter.LoginAuthenticated && bootstrapAfter.MCPToolCount > 0
assertions.Record("Dashboard generation two preserves fixture and recreates runtime clients", runtimeAfter.Generation > runtimeBefore.Generation && runtimeAfter.PID != runtimeBefore.PID && reconnectEvidence.Identity.DashboardFixtureUnchanged && reconnectEvidence.Identity.ClientsRecreated && reconnectEvidence.Identity.BootstrapRecreated, "")
assertions.Record("Agent reconnect preserves exact server ID and UUID", serverIDAfter == serverID && postDashboardReadiness.UUID == agentUUID, "")
dashboardPairs, err := runDashboardReconnectOperations(ctx, dashboardInstance, serverID, agentUUID, fixturePath)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
stateGenerationBefore := dashboardInstance.StateGeneration(serverID, agentUUID)
transition, err := agentInstance.RestartProcess(ctx)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
if err := dashboardInstance.WaitForStateGeneration(ctx, serverID, agentUUID, stateGenerationBefore+1, 1); err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
postAgentReadiness, err := agentInstance.WaitReady(ctx, dashboardInstance)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
agentConfigAfter, err := os.ReadFile(agentInstance.ConfigPath())
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
agentPairs, err := runAgentRestartOperations(ctx, dashboardInstance, serverID, fixturePath)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
taskIDs, resultIDs, duplicates, lost := reconnectReceiptSummary(dashboardPairs, agentPairs)
stale := staleReconnectReceiptCount(runtimeBefore.Generation, dashboardPairs, agentPairs)
sentinelAfter, err := os.ReadFile(sentinelPath)
if err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
reconnectEvidence.Runtime.AgentAfter = transition.Current
reconnectEvidence.Runtime.StateGenerationBeforeAgentRestart = stateGenerationBefore
reconnectEvidence.Runtime.StateGenerationAfterAgentRestart = dashboardInstance.StateGeneration(serverID, agentUUID)
reconnectEvidence.Identity.AgentConfigUnchanged = bytes.Equal(agentConfig, agentConfigAfter) && agentInstance.ConfigPath() == reconnectEvidence.Fixture.AgentConfigPath && agentInstance.BinaryPath() == reconnectEvidence.Fixture.AgentBinaryPath && agentInstance.WorkspaceRoot() == reconnectEvidence.Fixture.AgentRoot
reconnectEvidence.Lifecycle = ReconnectLifecycleEvidence{DisconnectAt: disconnectAt, ReconnectAt: reconnectAt, ReconnectInterval: reconnectAt.Sub(disconnectAt), DashboardReceipts: dashboardPairs, AgentReceipts: agentPairs, StaleGenerationReceipts: stale, DuplicateTaskIDs: duplicates, LostResultIDs: lost, OutsideRootSentinelUnchanged: bytes.Equal(sentinelBytes, sentinelAfter)}
reconnectEvidence.Observation = ReconnectObservation{ServerID: serverID, UUID: agentUUID, OldGeneration: runtimeBefore.Generation, NewGeneration: runtimeAfter.Generation, DisconnectAt: disconnectAt, ReconnectAt: reconnectAt, TaskIDs: taskIDs, ResultIDs: resultIDs, PostReconnect: true, AgentRestarted: transition.Previous == agentBefore && transition.Current.Generation > transition.Previous.Generation && postAgentReadiness.UUID == agentUUID}
assertions.Record("post-reconnect MCP task and result receipts are exactly once", duplicates == 0 && lost == 0 && len(taskIDs) == 5, "")
assertions.Record("stale Dashboard generation cannot receive new task receipts", stale == 0, "")
assertions.Record("Agent restart advances state stream and preserves config identity", reconnectEvidence.Runtime.StateGenerationAfterAgentRestart > stateGenerationBefore && reconnectEvidence.Identity.AgentConfigUnchanged && reconnectEvidence.Observation.AgentRestarted, "")
assertions.Record("outside-root sentinel remains unchanged", reconnectEvidence.Lifecycle.OutsideRootSentinelUnchanged, "")
if err := reconnectEvidence.Validate(); err != nil {
return reconnectFinish(assertions, err, reconnectEvidence)
}
return reconnectFinish(assertions, nil, reconnectEvidence)
}
func agentCleanupReceipt(agentInstance *agent.Agent) processharness.CleanupReceipt {
if agentInstance == nil {
return processharness.CleanupReceipt{}
}
return agentInstance.CleanupReceipt()
}
func reconnectFinish(assertions *AssertionSet, runErr error, reconnectEvidence ReconnectEvidence) (Result, ReconnectEvidence, error) {
for _, assertion := range assertions.assertions {
if !assertion.Passed && runErr == nil {
runErr = errors.New(assertion.Name + ": " + assertion.Details)
}
}
result := Result{Name: reconnectScenarioName, Passed: runErr == nil, Assertions: assertions.Results()}
if runErr != nil {
result.Error = errorText(runErr)
}
return result, reconnectEvidence, runErr
}
@@ -0,0 +1,84 @@
//go:build linux && agentcompat
package scenario
import (
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestReconnectObservation_RejectsNonIncreasingGenerations(t *testing.T) {
// Given
observation := ReconnectObservation{OldGeneration: 4, NewGeneration: 4}
// When
err := observation.Validate()
// Then
require.Error(t, err)
}
func TestReconnectObservation_RequiresUniqueCompleteTaskIDs(t *testing.T) {
// Given
observation := ReconnectObservation{
OldGeneration: 2,
NewGeneration: 3,
DisconnectAt: time.Unix(10, 0),
ReconnectAt: time.Unix(11, 0),
TaskIDs: []uint64{7, 7},
ResultIDs: []uint64{7},
}
// When
err := observation.Validate()
// Then
require.Error(t, err)
}
func TestReconnectObservation_RejectsNonAdjacentDuplicateTaskIDs(t *testing.T) {
// Given
observation := ReconnectObservation{
ServerID: 7,
UUID: "00000000-0000-0000-0000-000000000111",
OldGeneration: 2,
NewGeneration: 3,
DisconnectAt: time.Unix(10, 0),
ReconnectAt: time.Unix(11, 0),
TaskIDs: []uint64{7, 8, 7},
ResultIDs: []uint64{7, 8, 7},
PostReconnect: true,
AgentRestarted: true,
}
// When
err := observation.Validate()
// Then
require.Error(t, err)
}
func TestReconnectObservation_RecordsReconnectInterval(t *testing.T) {
// Given
observation := ReconnectObservation{
ServerID: 7,
UUID: "00000000-0000-0000-0000-000000000111",
OldGeneration: 2,
NewGeneration: 3,
DisconnectAt: time.Unix(10, 0),
ReconnectAt: time.Unix(11, 0),
TaskIDs: []uint64{7},
ResultIDs: []uint64{7},
PostReconnect: true,
AgentRestarted: true,
}
// When
err := observation.Validate()
// Then
require.NoError(t, err)
require.Equal(t, time.Second, observation.ReconnectInterval())
}
@@ -0,0 +1,293 @@
//go:build linux
package scenario
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/evidence"
)
type RegistrationConfigExecInput struct {
Paths contract.Paths
Fault contract.Fault
}
type RegistrationConfigExec struct{}
type serverListArguments struct {
OnlineOnly bool `json:"online_only"`
}
type serverListResult struct {
Servers []struct {
ID uint64 `json:"id"`
UUID string `json:"uuid"`
Online bool `json:"online"`
} `json:"servers"`
}
type serverGetArguments struct {
ServerID uint64 `json:"server_id"`
}
type serverGetResult struct {
UUID string `json:"uuid"`
Host json.RawMessage `json:"host"`
State json.RawMessage `json:"state"`
}
type execArguments struct {
ServerID uint64 `json:"server_id"`
Cmd string `json:"cmd"`
Args []string `json:"args"`
}
type execResult struct {
ExitCode int `json:"exit_code"`
Stdout string `json:"stdout"`
Stderr string `json:"stderr"`
StdoutTruncated bool `json:"stdout_truncated"`
TimedOut bool `json:"timed_out"`
Error string `json:"error"`
}
type configPostRequest struct {
Servers []uint64 `json:"servers"`
Config string `json:"config"`
}
type configPostResponse struct {
Success []uint64 `json:"success"`
Failure []uint64 `json:"failure"`
Offline []uint64 `json:"offline"`
}
type patRequest struct {
Name string `json:"name"`
Scopes []string `json:"scopes"`
ExpiresInDays int `json:"expires_in_days"`
}
type patResponse struct {
Token string `json:"token"`
}
func (RegistrationConfigExec) Run(ctx context.Context, input RegistrationConfigExecInput) (result Result, runErr error) {
assertions := NewAssertionSet()
dashboardInstance, err := dashboard.Start(ctx, dashboard.StartConfig{SourceDir: input.Paths.NezhaSource().String(), ReceiptGate: true})
if err != nil {
return Result{Name: "registration-config-exec", Assertions: assertions.Results(), Error: err.Error()}, err
}
defer func() {
cleanupErr := dashboardInstance.Stop(context.Background())
result.CleanupOK = cleanupErr == nil && dashboardInstance.CleanupReceipt().Passed
if cleanupErr != nil && runErr == nil {
runErr = cleanupErr
result.Passed = false
result.Error = errorText(cleanupErr)
}
}()
secret := dashboardInstance.AgentSecret()
agentSecret := secret
if input.Fault.String() == "agent-bad-secret" {
agentSecret = "wrong-agent-secret"
}
agentInstance, err := agent.Start(ctx, agent.AgentStartConfig{SourceDir: input.Paths.AgentSource().String(), Endpoint: dashboardInstance.Endpoint(), Secret: agentSecret, UUID: "00000000-0000-0000-0000-000000000111"})
if err != nil {
return Result{Name: "registration-config-exec", Assertions: assertions.Results(), Error: err.Error()}, err
}
defer func() {
cleanupErr := agentInstance.Stop(context.Background())
if cleanupErr != nil && runErr == nil {
runErr = cleanupErr
result.Passed = false
result.Error = errorText(cleanupErr)
}
}()
if input.Fault.String() == "agent-bad-secret" {
badContext, cancel := context.WithTimeout(ctx, 8*time.Second)
defer cancel()
err = agentInstance.AssertNeverOnline(badContext, dashboardInstance, 3*time.Second)
faultDetails := "invalid secret prevented readiness as expected"
if err != nil {
faultDetails = errorText(err)
}
assertions.Record("agent-bad-secret prevents readiness", false, faultDetails)
if err == nil {
err = errors.New("fault injection agent-bad-secret")
}
return finish(assertions, err)
}
if err := dashboardInstance.WaitForReceiptAccepted(ctx); err != nil {
return finish(assertions, err)
}
if err := dashboardInstance.ReleaseReceipt(ctx); err != nil {
return finish(assertions, err)
}
readiness, err := agentInstance.WaitReady(ctx, dashboardInstance)
assertions.Record("online inventory has exact UUID", err == nil && readiness.UUID == agentInstance.UUID() && readiness.Online, errorText(err))
assertions.Record("online inventory has Host and State", err == nil && len(readiness.Host) > 0 && len(readiness.State) > 0, errorText(err))
if err != nil {
return finish(assertions, err)
}
servers, err := client.CallTool[serverListArguments, serverListResult](ctx, dashboardInstance.Clients().MCP, client.ToolCall[serverListArguments]{Name: "server.list", Arguments: serverListArguments{OnlineOnly: true}})
if err != nil {
return finish(assertions, err)
}
serverID := uint64(0)
for _, server := range servers.StructuredContent.Servers {
if server.UUID == agentInstance.UUID() && server.Online {
serverID = server.ID
}
}
assertions.Record("server.list exact online UUID", serverID != 0, "")
if serverID == 0 {
return finish(assertions, errors.New("server.list did not return the agent UUID"))
}
server, err := client.CallTool[serverGetArguments, serverGetResult](ctx, dashboardInstance.Clients().MCP, client.ToolCall[serverGetArguments]{Name: "server.get", Arguments: serverGetArguments{ServerID: serverID}})
assertions.Record("server.get exact UUID and meaningful Host State", err == nil && server.StructuredContent.UUID == agentInstance.UUID() && string(server.StructuredContent.Host) != "null" && string(server.StructuredContent.State) != "null", errorText(err))
if err != nil {
return finish(assertions, err)
}
limited, err := createScopedClient(ctx, dashboardInstance, []string{"nezha:server:read"})
if err != nil {
return finish(assertions, err)
}
_, err = client.DoREST[struct{}, string](ctx, limited, client.RESTRequest[struct{}]{Method: http.MethodGet, Path: fmt.Sprintf("/api/v1/server/config/%d", serverID)})
assertions.Record("insufficient config scope denied", isForbidden(err), errorText(err))
configRaw, err := client.DoREST[struct{}, string](ctx, dashboardInstance.Clients().REST, client.RESTRequest[struct{}]{Method: http.MethodGet, Path: fmt.Sprintf("/api/v1/server/config/%d", serverID)})
if err != nil {
return finish(assertions, err)
}
original, err := decodeAgentConfig(configRaw)
assertions.Record("authorized config returns complete round-trip contract", err == nil && original.ClientSecret != "" && original.UUID == agentInstance.UUID() && original.Server != "", errorText(err))
if err != nil {
return finish(assertions, err)
}
updated := original
updated.Debug = !original.Debug
updated.ReportDelay = original.ReportDelay%4 + 1
configDiffErr := changedOnlyDebugAndReportDelay(original, updated)
assertions.Record("config diff changes only debug and report_delay", configDiffErr == nil, errorText(configDiffErr))
if configDiffErr != nil {
return finish(assertions, configDiffErr)
}
encoded, err := json.Marshal(updated)
if err != nil {
return finish(assertions, err)
}
response, err := client.DoREST[configPostRequest, configPostResponse](ctx, dashboardInstance.Clients().REST, client.RESTRequest[configPostRequest]{Method: http.MethodPost, Path: "/api/v1/server/config", Body: &configPostRequest{Servers: []uint64{serverID}, Config: string(encoded)}})
dispatchValid := err == nil && len(response.Success) == 1 && response.Success[0] == serverID
dispatchDetails := errorText(err)
if !dispatchValid && dispatchDetails == "" {
dispatchDetails = fmt.Sprintf("success=%v failure=%v offline=%v", response.Success, response.Failure, response.Offline)
}
assertions.Record("config update dispatched", dispatchValid, dispatchDetails)
if !dispatchValid {
if err != nil {
return finish(assertions, fmt.Errorf("config dispatch failed: %w", err))
}
return finish(assertions, errors.New("config dispatch returned no successful server"))
}
// Agent ApplyConfig commits after its deferred reload window, then reconnects;
// this state-generation event is the harness boundary that proves the new
// connection published state instead of merely accepting the task.
stateGeneration := dashboardInstance.StateGeneration(serverID, agentInstance.UUID())
if stateGeneration == 0 {
return finish(assertions, errors.New("state generation was not observed before config reload"))
}
if err := dashboardInstance.WaitForStateGeneration(ctx, serverID, agentInstance.UUID(), stateGeneration+1, 1); err != nil {
return finish(assertions, err)
}
persisted, err := waitForPersistedConfig(ctx, agentInstance.ConfigPath(), updated)
if err != nil {
return finish(assertions, err)
}
persistedMatches := persisted.Debug == updated.Debug && persisted.ReportDelay == updated.ReportDelay && persisted.ClientSecret == original.ClientSecret && persisted.UUID == original.UUID && persisted.Server == original.Server
assertions.Record("config reload persisted only requested changes", persistedMatches, fmt.Sprintf("debug=%t/%t report_delay=%d/%d uuid=%s/%s server=%s/%s", persisted.Debug, updated.Debug, persisted.ReportDelay, updated.ReportDelay, persisted.UUID, original.UUID, persisted.Server, original.Server))
postReload, err := agentInstance.WaitReady(ctx, dashboardInstance)
assertions.Record("post-reload online identity remains stable", err == nil && postReload.UUID == agentInstance.UUID() && postReload.Online, errorText(err))
if err != nil {
return finish(assertions, err)
}
exec, err := client.CallTool[execArguments, execResult](ctx, dashboardInstance.Clients().MCP, client.ToolCall[execArguments]{Name: "server.exec", Arguments: execArguments{ServerID: serverID, Cmd: "/bin/sh", Args: []string{"-c", "printf compat-exec"}}})
assertions.Record("valid Exec exact stdout exit and no truncation timeout", err == nil && exec.StructuredContent.ExitCode == 0 && exec.StructuredContent.Stdout == "compat-exec" && exec.StructuredContent.Error == "" && !exec.StructuredContent.StdoutTruncated && !exec.StructuredContent.TimedOut, errorText(err))
_, invalidErr := client.CallTool[execArguments, execResult](ctx, dashboardInstance.Clients().MCP, client.ToolCall[execArguments]{Name: "server.exec", Arguments: execArguments{ServerID: serverID, Cmd: "/definitely/missing/compat-command"}})
var toolFailure *client.ToolFailure
structuredFailure := errors.As(invalidErr, &toolFailure)
var invalidResult execResult
if structuredFailure {
decodeErr := json.Unmarshal(toolFailure.StructuredContent, &invalidResult)
structuredFailure = decodeErr == nil
if decodeErr != nil {
invalidErr = errors.Join(invalidErr, decodeErr)
}
}
assertions.Record("invalid Exec has typed nonzero semantics", structuredFailure && invalidResult.ExitCode != 0 && invalidResult.Error != "", errorText(invalidErr))
_, err = client.CallTool[serverListArguments, serverListResult](ctx, dashboardInstance.Clients().MCP, client.ToolCall[serverListArguments]{Name: "server.list", Arguments: serverListArguments{OnlineOnly: true}})
assertions.Record("MCP health continues after Exec", err == nil, errorText(err))
return finish(assertions, nil)
}
func waitForPersistedConfig(ctx context.Context, path string, want AgentConfig) (AgentConfig, error) {
deadline, cancel := context.WithTimeout(ctx, 20*time.Second)
defer cancel()
ticker := time.NewTicker(100 * time.Millisecond)
defer ticker.Stop()
for {
config, err := ReadConfigFile(path)
if err == nil && config.Debug == want.Debug && config.ReportDelay == want.ReportDelay {
return config, nil
}
select {
case <-ticker.C:
case <-deadline.Done():
if err != nil {
return AgentConfig{}, fmt.Errorf("wait for persisted config: %w", err)
}
return AgentConfig{}, fmt.Errorf("wait for persisted config: %w", deadline.Err())
}
}
}
func finish(assertions *AssertionSet, runErr error) (Result, error) {
for _, assertion := range assertions.assertions {
if !assertion.Passed && runErr == nil {
runErr = fmt.Errorf("%s: %s", assertion.Name, assertion.Details)
}
}
result := Result{Name: "registration-config-exec", Passed: runErr == nil, Assertions: assertions.Results(), CleanupOK: false}
if runErr != nil {
result.Error = evidence.Redact(runErr.Error())
}
return result, runErr
}
func createScopedClient(ctx context.Context, dashboardInstance *dashboard.Dashboard, scopes []string) (*client.Client, error) {
pat, err := client.DoREST[patRequest, patResponse](ctx, dashboardInstance.Clients().REST, client.RESTRequest[patRequest]{Method: http.MethodPost, Path: "/api/v1/api-tokens", Body: &patRequest{Name: "agentcompat-scope-check", Scopes: scopes}})
if err != nil {
return nil, err
}
return dashboardInstance.AuthenticatedClient(pat.Token)
}
func isForbidden(err error) bool {
var httpErr *client.HTTPError
if errors.As(err, &httpErr) {
return httpErr.StatusCode == http.StatusForbidden
}
var handshakeErr *client.WebSocketHandshakeError
return errors.As(err, &handshakeErr) && handshakeErr.StatusCode == http.StatusForbidden
}
func errorText(err error) string {
if err == nil {
return ""
}
return evidence.Redact(err.Error())
}
@@ -0,0 +1,66 @@
//go:build linux && agentcompat
package scenario
import (
"testing"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/evidence"
)
func TestConfigDiff_ChangesOnlyDebugAndReportDelay(t *testing.T) {
// Given
original := AgentConfigSnapshot{Debug: false, ReportDelay: 1, ClientSecret: "secret", UUID: "uuid", Server: "server"}
updated := original
updated.Debug = true
updated.ReportDelay = 2
// When
diff, err := ConfigDiff(original, updated)
// Then
require.NoError(t, err)
require.Equal(t, ConfigDiffResult{DebugChanged: true, ReportDelayChanged: true}, diff)
}
func TestConfigDiff_RejectsCredentialIdentityAndEndpointChanges(t *testing.T) {
// Given
original := AgentConfigSnapshot{ClientSecret: "secret", UUID: "uuid", Server: "server", ReportDelay: 1}
updated := original
updated.ClientSecret = "other-secret"
// When
_, err := ConfigDiff(original, updated)
// Then
require.ErrorIs(t, err, ErrConfigIdentityChanged)
}
func TestSensitiveEvidence_RedactsConfigCredentials(t *testing.T) {
// Given
secret := "0123456789abcdef0123456789abcdef"
config := `{"client_secret":"` + secret + `","server":"127.0.0.1:5555","debug":true}`
// When
redacted := evidence.Redact(config)
// Then
require.NotContains(t, redacted, secret)
require.Contains(t, redacted, "[REDACTED]")
}
func TestFinish_RecordsFailedAssertionAndErrorForCleanup(t *testing.T) {
assertions := NewAssertionSet()
assertions.Record("readiness", false, "invalid secret prevented readiness")
result, err := finish(assertions, nil)
require.EqualError(t, err, "readiness: invalid secret prevented readiness")
require.False(t, result.Passed)
require.False(t, result.CleanupOK)
require.Equal(t, "readiness: invalid secret prevented readiness", result.Error)
require.Len(t, result.Assertions, 1)
require.False(t, result.Assertions[0].Passed)
}
@@ -0,0 +1,209 @@
//go:build linux
package scenario
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/evidence"
processharness "github.com/nezhahq/nezha/integration/agentcompat/internal/process"
)
const (
terminalMarker = "compat-terminal"
terminalCommand = "printf 'compat-size='; stty size; printf 'compat-terminal\\n'; exit\n"
terminalShutdownContract = 2 * time.Second
terminalShutdownHarnessMargin = 500 * time.Millisecond
terminalAttachPATScope = "nezha:server:exec"
)
type TerminalInput struct {
Paths contract.Paths
Fault contract.Fault
}
type Terminal struct{}
type terminalCreateRequest struct {
Protocol string `json:"protocol"`
ServerID uint64 `json:"server_id"`
}
type terminalCreateResponse struct {
SessionID string `json:"session_id"`
ServerID uint64 `json:"server_id"`
}
type terminalUserRequest struct {
Role uint8 `json:"role"`
Username string `json:"username"`
Password string `json:"password"`
}
type terminalPATRequest struct {
Name string `json:"name"`
Scopes []string `json:"scopes"`
ServerIDs []uint64 `json:"server_ids,omitempty"`
}
type terminalPATResponse struct {
Token string `json:"token"`
}
func (Terminal) Run(ctx context.Context, input TerminalInput) (result Result, runErr error) {
assertions := NewAssertionSet()
dashboardInstance, err := dashboard.Start(ctx, dashboard.StartConfig{SourceDir: input.Paths.NezhaSource().String(), ReceiptGate: true})
if err != nil {
return terminalFinish(assertions, err)
}
result.CleanupOK = true
defer func() {
cleanupContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), 15*time.Second)
defer cancel()
cleanupErr := dashboardInstance.Stop(cleanupContext)
receipt := dashboardInstance.CleanupReceipt()
cleanupPassed := cleanupErr == nil && receipt.Passed && !receipt.Forced
result.CleanupOK = result.CleanupOK && cleanupPassed
if !cleanupPassed {
cleanupErr = errors.Join(cleanupErr, errors.New("dashboard cleanup receipt failed"))
result, runErr = terminalFinish(assertions, errors.Join(runErr, cleanupErr))
result.CleanupOK = false
}
}()
agentInstance, err := agent.Start(ctx, agent.AgentStartConfig{SourceDir: input.Paths.AgentSource().String(), Endpoint: dashboardInstance.Endpoint(), Secret: dashboardInstance.AgentSecret(), UUID: "00000000-0000-0000-0000-000000000113"})
if err != nil {
return terminalFinish(assertions, err)
}
defer func() {
cleanupContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), 15*time.Second)
defer cancel()
cleanupErr := agentInstance.Stop(cleanupContext)
receipt := agentInstance.CleanupReceipt()
cleanupPassed := cleanupErr == nil && receipt.Passed && !receipt.Forced
result.CleanupOK = result.CleanupOK && cleanupPassed
if !cleanupPassed {
cleanupErr = errors.Join(cleanupErr, errors.New("agent cleanup receipt failed"))
result, runErr = terminalFinish(assertions, errors.Join(runErr, cleanupErr))
result.CleanupOK = false
}
}()
if err := dashboardInstance.WaitForReceiptAccepted(ctx); err != nil {
return terminalFinish(assertions, err)
}
if err := dashboardInstance.ReleaseReceipt(ctx); err != nil {
return terminalFinish(assertions, err)
}
readiness, err := agentInstance.WaitReady(ctx, dashboardInstance)
if err != nil {
return terminalFinish(assertions, err)
}
serverID, err := terminalServerID(ctx, dashboardInstance.Clients().MCP, readiness.UUID)
if err != nil {
return terminalFinish(assertions, err)
}
baseline, err := processharness.SampleProcess(agentInstance.PID())
if err != nil {
return terminalFinish(assertions, err)
}
deniedClient, err := createTerminalPATClient(ctx, dashboardInstance, "terminal-denied", []string{"nezha:server:read"}, []uint64{serverID})
if err != nil {
return terminalFinish(assertions, err)
}
_, deniedErr := client.DoREST[terminalCreateRequest, terminalCreateResponse](ctx, deniedClient, client.RESTRequest[terminalCreateRequest]{Method: http.MethodPost, Path: "/api/v1/terminal", Body: &terminalCreateRequest{Protocol: "grpc", ServerID: serverID}})
assertions.Record("terminal denied PAT lacks exec scope", isForbidden(deniedErr), errorText(deniedErr))
terminal, err := client.DoREST[terminalCreateRequest, terminalCreateResponse](ctx, dashboardInstance.Clients().REST, client.RESTRequest[terminalCreateRequest]{Method: http.MethodPost, Path: "/api/v1/terminal", Body: &terminalCreateRequest{Protocol: "grpc", ServerID: serverID}})
if err != nil || terminal.SessionID == "" || terminal.ServerID != serverID {
return terminalFinish(assertions, errors.Join(err, errors.New("terminal creation returned incomplete session")))
}
foreignClient, cleanupForeignUser, err := createForeignTerminalPATClient(ctx, dashboardInstance)
if err != nil {
return terminalFinish(assertions, err)
}
hijackConnection, hijackErr := foreignClient.DialWebSocket(ctx, "/api/v1/ws/terminal/"+terminal.SessionID)
if hijackConnection != nil {
_ = hijackConnection.Close()
}
assertions.Record("foreign scoped PAT cannot hijack terminal session", isWebSocketDenied(hijackErr), fmt.Sprintf("scopes=[%s] denial=%s", terminalAttachPATScope, errorText(hijackErr)))
if err := cleanupForeignUser(); err != nil {
return terminalFinish(assertions, err)
}
connection, err := dashboardInstance.Clients().WebSocket.DialWebSocket(ctx, "/api/v1/ws/terminal/"+terminal.SessionID)
if err != nil {
return terminalFinish(assertions, err)
}
defer connection.Close()
if err := connection.WriteFrame(ctx, mustTerminalResizeFrame(132, 43)); err != nil {
return terminalFinish(assertions, err)
}
initialFrame, err := connection.ReadFrame(ctx)
if err != nil {
return terminalFinish(assertions, err)
}
active, err := processharness.SampleProcess(agentInstance.PID())
if err != nil {
return terminalFinish(assertions, err)
}
// Start the contract clock before sending exit so transport time is included.
exitSentAt := time.Now()
output, err := executeTerminalExit(ctx, terminalExitInput{InitialOutput: initialFrame.Payload, ExitSentAt: exitSentAt, Now: time.Now}, connection)
assertions.Record("terminal resize marker and bounded shell close observed", err == nil && output.MarkerObserved && output.SizeObserved && output.Rows == 43 && output.Cols == 132 && output.StreamClosed && terminalCloseWithinContract(output.CloseElapsed), terminalOutputDetails(output, err))
if err != nil {
return terminalFinish(assertions, err)
}
residue, err := processharness.SampleProcess(agentInstance.PID())
residueClean := err == nil && active.DescendantCount > baseline.DescendantCount && residue.DescendantCount == baseline.DescendantCount && residue.TCPListenerCount == baseline.TCPListenerCount && residue.TCP6ListenerCount == baseline.TCP6ListenerCount
assertions.Record("agent PTY child and listener residue cleared", residueClean, fmt.Sprintf("active_children=%d baseline_children=%d residue_children=%d baseline_listeners=%d/%d residue_listeners=%d/%d error=%s", active.DescendantCount, baseline.DescendantCount, residue.DescendantCount, baseline.TCPListenerCount, baseline.TCP6ListenerCount, residue.TCPListenerCount, residue.TCP6ListenerCount, errorText(err)))
_, staleErr := dashboardInstance.Clients().WebSocket.DialWebSocket(ctx, "/api/v1/ws/terminal/"+terminal.SessionID)
assertions.Record("terminal IOStream removed after shell exit", isWebSocketDenied(staleErr), errorText(staleErr))
_, invalidErr := dashboardInstance.Clients().WebSocket.DialWebSocket(ctx, "/api/v1/ws/terminal/invalid-session")
assertions.Record("invalid terminal session rejected", webSocketFailureContains(invalidErr, "permission denied"), errorText(invalidErr))
return terminalFinish(assertions, nil)
}
func terminalResizeFrame(cols, rows uint32) (client.Frame, error) {
payload, err := json.Marshal(struct {
Cols uint32
Rows uint32
}{Cols: cols, Rows: rows})
if err != nil {
return client.Frame{}, err
}
return client.Frame{Type: client.FrameBinary, Payload: append([]byte{1}, payload...)}, nil
}
func mustTerminalResizeFrame(cols, rows uint32) client.Frame {
frame, _ := terminalResizeFrame(cols, rows)
return frame
}
func terminalFinish(assertions *AssertionSet, runErr error) (Result, error) {
results := assertions.Results()
failedAssertion := false
for _, assertion := range results {
if !assertion.Passed && runErr == nil {
runErr = fmt.Errorf("%s: %s", assertion.Name, assertion.Details)
}
failedAssertion = failedAssertion || !assertion.Passed
}
if runErr != nil && !failedAssertion {
results = append(results, Assertion{Name: "terminal scenario completed", Passed: false, Details: evidence.Redact(runErr.Error())})
}
result := Result{Name: "terminal", Passed: runErr == nil, Assertions: results, CleanupOK: true}
if runErr != nil {
result.Error = evidence.Redact(runErr.Error())
}
return result, runErr
}
@@ -0,0 +1,112 @@
//go:build linux
package scenario
import (
"bytes"
"context"
"errors"
"fmt"
"regexp"
"strconv"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
var terminalSizePattern = regexp.MustCompile(`compat-size=([0-9]+) ([0-9]+)`)
type terminalFrameConnection interface {
WriteFrame(context.Context, client.Frame) error
ReadFrame(context.Context) (client.Frame, error)
}
type terminalExitInput struct {
InitialOutput []byte
ExitSentAt time.Time
Now func() time.Time
}
type terminalOutputReadInput struct {
InitialOutput []byte
ExitSentAt time.Time
Now func() time.Time
}
type terminalOutputResult struct {
Output string
MarkerObserved bool
SizeObserved bool
StreamClosed bool
Rows uint32
Cols uint32
CloseCode int
CloseElapsed time.Duration
}
func executeTerminalExit(ctx context.Context, input terminalExitInput, connection terminalFrameConnection) (terminalOutputResult, error) {
contractContext, cancelContract := context.WithDeadline(ctx, input.ExitSentAt.Add(terminalShutdownContract+terminalShutdownHarnessMargin))
defer cancelContract()
if err := connection.WriteFrame(contractContext, client.Frame{Type: client.FrameText, Payload: []byte(terminalCommand)}); err != nil {
return terminalOutputResult{}, err
}
return readTerminalOutput(contractContext, terminalOutputReadInput{InitialOutput: input.InitialOutput, ExitSentAt: input.ExitSentAt, Now: input.Now}, connection.ReadFrame)
}
func readTerminalOutput(ctx context.Context, input terminalOutputReadInput, read func(context.Context) (client.Frame, error)) (terminalOutputResult, error) {
var output bytes.Buffer
output.Write(input.InitialOutput)
result := terminalOutputResult{Output: output.String()}
observeTerminalOutput(&result, output.Bytes())
for {
frame, err := read(ctx)
if err != nil {
result.Output = output.String()
result.CloseCode = closeErrorCode(err)
var closeError *client.WebSocketCloseError
if result.MarkerObserved && result.SizeObserved && errors.As(err, &closeError) && terminalCloseCodeAccepted(closeError.Code) {
result.StreamClosed = true
result.CloseElapsed = input.Now().Sub(input.ExitSentAt)
return result, nil
}
return result, err
}
output.Write(frame.Payload)
observeTerminalOutput(&result, output.Bytes())
}
}
func observeTerminalOutput(result *terminalOutputResult, output []byte) {
result.MarkerObserved = bytes.Contains(output, []byte(terminalMarker))
matches := terminalSizePattern.FindSubmatch(output)
if len(matches) != 3 {
return
}
rows, rowsErr := strconv.ParseUint(string(matches[1]), 10, 32)
cols, colsErr := strconv.ParseUint(string(matches[2]), 10, 32)
if rowsErr == nil && colsErr == nil {
result.SizeObserved = true
result.Rows = uint32(rows)
result.Cols = uint32(cols)
}
}
func terminalCloseWithinContract(elapsed time.Duration) bool {
return elapsed <= terminalShutdownContract+terminalShutdownHarnessMargin
}
func terminalCloseCodeAccepted(code int) bool {
return code == 1000 || code == 1006
}
func closeErrorCode(err error) int {
var closeError *client.WebSocketCloseError
if errors.As(err, &closeError) {
return closeError.Code
}
return 0
}
func terminalOutputDetails(output terminalOutputResult, err error) string {
return fmt.Sprintf("marker=%t size_observed=%t rows=%d cols=%d closed=%t close_code=%d close_elapsed_ms=%d close_limit_ms=%d output=%q error=%s", output.MarkerObserved, output.SizeObserved, output.Rows, output.Cols, output.StreamClosed, output.CloseCode, output.CloseElapsed.Milliseconds(), (terminalShutdownContract + terminalShutdownHarnessMargin).Milliseconds(), output.Output, errorText(err))
}
@@ -0,0 +1,77 @@
//go:build linux
package scenario
import (
"context"
"errors"
"net/http"
"strings"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
)
func terminalServerID(ctx context.Context, mcpClient *client.Client, uuid string) (uint64, error) {
servers, err := client.CallTool[serverListArguments, serverListResult](ctx, mcpClient, client.ToolCall[serverListArguments]{Name: "server.list", Arguments: serverListArguments{OnlineOnly: true}})
if err != nil {
return 0, err
}
for _, server := range servers.StructuredContent.Servers {
if server.UUID == uuid && server.Online {
return server.ID, nil
}
}
return 0, errors.New("terminal agent server ID not found")
}
func createTerminalPATClient(ctx context.Context, dashboardInstance *dashboard.Dashboard, name string, scopes []string, serverIDs []uint64) (*client.Client, error) {
pat, err := client.DoREST[terminalPATRequest, terminalPATResponse](ctx, dashboardInstance.Clients().REST, client.RESTRequest[terminalPATRequest]{Method: http.MethodPost, Path: "/api/v1/api-tokens", Body: &terminalPATRequest{Name: name, Scopes: scopes, ServerIDs: serverIDs}})
if err != nil {
return nil, err
}
return dashboardInstance.AuthenticatedClient(pat.Token)
}
func createForeignTerminalPATClient(ctx context.Context, dashboardInstance *dashboard.Dashboard) (*client.Client, func() error, error) {
const username = "terminal-member"
const password = "terminal-member-password"
admin := dashboardInstance.Clients().REST
userID, err := client.DoREST[terminalUserRequest, uint64](ctx, admin, client.RESTRequest[terminalUserRequest]{Method: http.MethodPost, Path: "/api/v1/user", Body: &terminalUserRequest{Role: 1, Username: username, Password: password}})
if err != nil {
return nil, func() error { return nil }, err
}
cleanup := func() error {
cleanupContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), 10*time.Second)
defer cancel()
_, cleanupErr := client.DoREST[[]uint64, struct{}](cleanupContext, admin, client.RESTRequest[[]uint64]{Method: http.MethodPost, Path: "/api/v1/batch-delete/user", Body: &[]uint64{userID}})
return cleanupErr
}
member, err := client.New(client.Config{BaseURL: dashboardInstance.URL()})
if err != nil {
return nil, func() error { return nil }, errors.Join(err, cleanup())
}
if _, err := member.Login(ctx, client.LoginRequest{Username: username, Password: password}); err != nil {
return nil, func() error { return nil }, errors.Join(err, cleanup())
}
pat, err := client.DoREST[terminalPATRequest, terminalPATResponse](ctx, member, client.RESTRequest[terminalPATRequest]{Method: http.MethodPost, Path: "/api/v1/api-tokens", Body: &terminalPATRequest{Name: "terminal-foreign", Scopes: []string{terminalAttachPATScope}}})
if err != nil {
return nil, func() error { return nil }, errors.Join(err, cleanup())
}
foreign, err := dashboardInstance.AuthenticatedClient(pat.Token)
if err != nil {
return nil, func() error { return nil }, errors.Join(err, cleanup())
}
return foreign, cleanup, nil
}
func isWebSocketDenied(err error) bool {
var handshakeError *client.WebSocketHandshakeError
return errors.As(err, &handshakeError) && (handshakeError.StatusCode == http.StatusForbidden || webSocketFailureContains(err, "permission denied") || webSocketFailureContains(err, "ApiErrorUnauthorized"))
}
func webSocketFailureContains(err error, text string) bool {
var handshakeError *client.WebSocketHandshakeError
return errors.As(err, &handshakeError) && strings.Contains(handshakeError.Message, text)
}
@@ -0,0 +1,133 @@
//go:build linux && agentcompat
package scenario
import (
"context"
"errors"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
)
type terminalDeadlineProbe struct {
writeDeadline time.Time
}
func (probe *terminalDeadlineProbe) WriteFrame(ctx context.Context, _ client.Frame) error {
probe.writeDeadline, _ = ctx.Deadline()
return context.DeadlineExceeded
}
func (*terminalDeadlineProbe) ReadFrame(context.Context) (client.Frame, error) {
return client.Frame{}, errors.New("read must not run after blocked write")
}
func TestTerminalOutputReader_ReturnsMarkerAndCloseEvidence(t *testing.T) {
frames := make(chan client.Frame, 1)
frames <- client.Frame{Type: client.FrameBinary, Payload: []byte("shell prompt\r\ncompat-size=43 132\r\ncompat-terminal\r\n")}
close(frames)
result, err := readTerminalOutput(context.Background(), terminalOutputReadInput{ExitSentAt: time.Unix(0, 0), Now: func() time.Time { return time.Unix(0, int64(time.Second)) }}, func(context.Context) (client.Frame, error) {
frame, ok := <-frames
if !ok {
return client.Frame{}, &client.WebSocketCloseError{Code: 1000, Text: "normal closure"}
}
return frame, nil
})
require.NoError(t, err)
require.True(t, result.MarkerObserved)
require.True(t, result.StreamClosed)
require.Equal(t, 1000, result.CloseCode)
require.Equal(t, time.Second, result.CloseElapsed)
require.Contains(t, result.Output, "compat-terminal")
}
func TestTerminalOutputReader_ObservesRequestedPTYSize(t *testing.T) {
frames := make(chan client.Frame, 1)
frames <- client.Frame{Type: client.FrameBinary, Payload: []byte("compat-size=43 132\r\ncompat-terminal\r\n")}
close(frames)
result, err := readTerminalOutput(context.Background(), terminalOutputReadInput{ExitSentAt: time.Unix(0, 0), Now: func() time.Time { return time.Unix(0, int64(time.Second)) }}, func(context.Context) (client.Frame, error) {
frame, ok := <-frames
if !ok {
return client.Frame{}, &client.WebSocketCloseError{Code: 1000, Text: "normal closure"}
}
return frame, nil
})
require.NoError(t, err)
require.True(t, result.SizeObserved)
require.Equal(t, uint32(43), result.Rows)
require.Equal(t, uint32(132), result.Cols)
}
func TestTerminalCommand_ReportsSizeBeforeMarkerAndExit(t *testing.T) {
require.Equal(t, "printf 'compat-size='; stty size; printf 'compat-terminal\\n'; exit\n", terminalCommand)
}
func TestTerminalCloseContract_AllowsAgentTimeoutPlusHarnessMargin(t *testing.T) {
require.True(t, terminalCloseWithinContract(terminalShutdownContract+terminalShutdownHarnessMargin))
require.False(t, terminalCloseWithinContract(terminalShutdownContract+terminalShutdownHarnessMargin+time.Nanosecond))
}
func TestTerminalExit_UsesOneAbsoluteDeadlineForCommandWriteAndCloseRead(t *testing.T) {
exitSentAt := time.Unix(100, 0)
probe := &terminalDeadlineProbe{}
_, err := executeTerminalExit(context.Background(), terminalExitInput{ExitSentAt: exitSentAt, Now: func() time.Time { return exitSentAt }}, probe)
require.ErrorIs(t, err, context.DeadlineExceeded)
require.Equal(t, exitSentAt.Add(terminalShutdownContract+terminalShutdownHarnessMargin), probe.writeDeadline)
}
func TestForeignTerminalPATScopes_UseMinimumAttachScope(t *testing.T) {
require.Equal(t, "nezha:server:exec", terminalAttachPATScope)
}
func TestTerminalOutputReader_RejectsNonCloseErrorAfterMarker(t *testing.T) {
reads := 0
result, err := readTerminalOutput(context.Background(), terminalOutputReadInput{ExitSentAt: time.Unix(0, 0), Now: func() time.Time { return time.Unix(0, int64(time.Second)) }}, func(context.Context) (client.Frame, error) {
reads++
if reads == 1 {
return client.Frame{Type: client.FrameBinary, Payload: []byte("compat-size=43 132\r\ncompat-terminal\r\n")}, nil
}
return client.Frame{}, context.DeadlineExceeded
})
require.ErrorIs(t, err, context.DeadlineExceeded)
require.True(t, result.MarkerObserved)
require.False(t, result.StreamClosed)
}
func TestTerminalOutputReader_RejectsProtocolCloseAfterMarker(t *testing.T) {
reads := 0
result, err := readTerminalOutput(context.Background(), terminalOutputReadInput{ExitSentAt: time.Unix(0, 0), Now: func() time.Time { return time.Unix(0, int64(time.Second)) }}, func(context.Context) (client.Frame, error) {
reads++
if reads == 1 {
return client.Frame{Type: client.FrameBinary, Payload: []byte("compat-size=43 132\r\ncompat-terminal\r\n")}, nil
}
return client.Frame{}, &client.WebSocketCloseError{Code: 1002, Text: "protocol error"}
})
var closeError *client.WebSocketCloseError
require.ErrorAs(t, err, &closeError)
require.Equal(t, 1002, result.CloseCode)
require.True(t, result.MarkerObserved)
require.False(t, result.StreamClosed)
}
func TestTerminalResizeFrame_UsesAgentWireContract(t *testing.T) {
frame, err := terminalResizeFrame(132, 43)
require.NoError(t, err)
require.Equal(t, client.FrameBinary, frame.Type)
require.Equal(t, byte(1), frame.Payload[0])
require.JSONEq(t, `{"Cols":132,"Rows":43}`, string(frame.Payload[1:]))
}
@@ -0,0 +1,170 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"os"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/agent"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/dashboard"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
const transferScenarioName = "transfer-100mib"
var ErrTransferHashFault = errors.New("transfer scenario: injected hash mismatch")
type TransferInput struct {
Paths contract.Paths
Fault contract.Fault
}
type Transfer struct{}
func (scenario Transfer) Run(ctx context.Context, input TransferInput) (Result, error) {
result, _, err := scenario.RunWithEvidence(ctx, input)
return result, err
}
func (Transfer) RunWithEvidence(ctx context.Context, input TransferInput) (result Result, transferEvidence TransferEvidence, runErr error) {
assertions := NewAssertionSet()
dashboardInstance, err := dashboard.Start(ctx, dashboard.StartConfig{SourceDir: input.Paths.NezhaSource().String(), ReceiptGate: true})
if err != nil {
result, runErr = transferFinish(assertions, err)
return result, transferEvidence, runErr
}
var agentInstance *agent.Agent
var agentWorkspaceRoot string
dashboardWorkspaceRoot := dashboardInstance.WorkspaceRoot()
defer func() {
cleanupContext, cancel := context.WithTimeout(context.WithoutCancel(ctx), 30*time.Second)
defer cancel()
cleanupErr := stopTransferProcesses(cleanupContext, agentInstance, dashboardInstance)
residueErr := transferWorkspaceResidue(agentWorkspaceRoot, dashboardWorkspaceRoot)
cleanupErr = errors.Join(cleanupErr, residueErr)
assertions.Record("process listener and workspace cleanup completed", cleanupErr == nil, errorText(cleanupErr))
result.CleanupOK = cleanupErr == nil
if cleanupErr != nil {
result, runErr = transferFinish(assertions, errors.Join(runErr, cleanupErr))
result.CleanupOK = false
} else {
result.Assertions = assertions.Results()
}
}()
agentInstance, err = agent.Start(ctx, agent.AgentStartConfig{
SourceDir: input.Paths.AgentSource().String(),
Endpoint: dashboardInstance.Endpoint(),
Secret: dashboardInstance.AgentSecret(),
UUID: "00000000-0000-0000-0000-000000000216",
})
if err != nil {
result, runErr = transferFinish(assertions, err)
return result, transferEvidence, runErr
}
agentWorkspaceRoot = agentInstance.WorkspaceRoot()
if err := dashboardInstance.WaitForReceiptAccepted(ctx); err != nil {
result, runErr = transferFinish(assertions, err)
return result, transferEvidence, runErr
}
if err := dashboardInstance.ReleaseReceipt(ctx); err != nil {
result, runErr = transferFinish(assertions, err)
return result, transferEvidence, runErr
}
readiness, err := agentInstance.WaitReady(ctx, dashboardInstance)
if err != nil {
result, runErr = transferFinish(assertions, err)
return result, transferEvidence, runErr
}
serverID, err := transferServerID(ctx, dashboardInstance.Clients().MCP, readiness.UUID)
if err != nil {
result, runErr = transferFinish(assertions, err)
return result, transferEvidence, runErr
}
root, err := fixture.NewAgentRoot(agentInstance.WorkspaceRoot(), "transfer-files")
if err != nil {
result, runErr = transferFinish(assertions, err)
return result, transferEvidence, runErr
}
sentinels, err := newTransferSentinels(root, agentInstance.WorkspaceRoot())
if err != nil {
result, runErr = transferFinish(assertions, err)
return result, transferEvidence, runErr
}
execution := transferExecution{
client: dashboardInstance.Clients().MCP,
serverID: serverID,
root: root,
residueScope: transferResidueScope{AgentRoot: agentInstance.WorkspaceRoot(), DashboardPID: dashboardInstance.PID()},
sentinels: sentinels,
}
transferEvidence, err = execution.run(ctx, assertions, input.Fault)
result, runErr = transferFinish(assertions, err)
return result, transferEvidence, runErr
}
func transferWorkspaceResidue(workspaceRoots ...string) error {
var residueErr error
for _, workspaceRoot := range workspaceRoots {
if workspaceRoot == "" {
continue
}
if _, err := os.Stat(workspaceRoot); err == nil {
residueErr = errors.Join(residueErr, fmt.Errorf("workspace remains: %s", workspaceRoot))
} else if !errors.Is(err, os.ErrNotExist) {
residueErr = errors.Join(residueErr, fmt.Errorf("inspect workspace %s: %w", workspaceRoot, err))
}
}
return residueErr
}
func stopTransferProcesses(ctx context.Context, agentInstance *agent.Agent, dashboardInstance *dashboard.Dashboard) error {
var cleanupErr error
if agentInstance != nil {
stopErr := agentInstance.Stop(ctx)
receipt := agentInstance.CleanupReceipt()
if stopErr != nil || !receipt.Passed || receipt.Forced {
cleanupErr = errors.Join(cleanupErr, stopErr, errors.New("agent cleanup receipt failed"))
}
}
stopErr := dashboardInstance.Stop(ctx)
receipt := dashboardInstance.CleanupReceipt()
if stopErr != nil || !receipt.Passed || receipt.Forced {
cleanupErr = errors.Join(cleanupErr, stopErr, errors.New("dashboard cleanup receipt failed"))
}
return cleanupErr
}
func transferFinish(assertions *AssertionSet, runErr error) (Result, error) {
for _, assertion := range assertions.assertions {
if !assertion.Passed && runErr == nil {
runErr = errors.New(assertion.Name + ": " + assertion.Details)
}
}
result := Result{Name: transferScenarioName, Passed: runErr == nil, Assertions: assertions.Results()}
if runErr != nil {
result.Error = errorText(runErr)
}
return result, runErr
}
func transferServerID(ctx context.Context, mcpClient *client.Client, uuid string) (uint64, error) {
servers, err := client.CallTool[client.ServerListArguments, client.ServerListResult](ctx, mcpClient, client.ToolCall[client.ServerListArguments]{
Name: "server.list", Arguments: client.ServerListArguments{OnlineOnly: true},
})
if err != nil {
return 0, fmt.Errorf("list transfer Agents: %w", err)
}
for _, server := range servers.StructuredContent.Servers {
if server.UUID == uuid && server.Online {
return server.ID, nil
}
}
return 0, errors.New("online transfer Agent not found")
}
@@ -0,0 +1,69 @@
//go:build linux
package scenario
import (
"errors"
"fmt"
"strings"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
)
var (
ErrTransferEvidenceSizeMismatch = errors.New("transfer evidence size mismatch")
ErrTransferEvidenceHashMismatch = errors.New("transfer evidence hash mismatch")
ErrTransferEvidenceHeapBudget = errors.New("transfer evidence retained heap budget exceeded")
ErrTransferEvidenceMeasurement = errors.New("transfer evidence measurement missing")
ErrTransferEvidenceContract = errors.New("transfer evidence contract assertion failed")
)
type TransferEvidence struct {
WarmupUploadBytes uint64 `json:"warmup_upload_bytes"`
WarmupDownloadBytes uint64 `json:"warmup_download_bytes"`
WarmupSHA256 string `json:"warmup_sha256"`
WarmupDuration time.Duration `json:"warmup_duration"`
WarmupDeadlineRemaining time.Duration `json:"warmup_deadline_remaining"`
WarmupQuiescent bool `json:"warmup_quiescent"`
UploadBytes uint64 `json:"upload_bytes"`
DownloadBytes uint64 `json:"download_bytes"`
UploadSHA256 string `json:"upload_sha256"`
DownloadSHA256 string `json:"download_sha256"`
UploadChunks uint64 `json:"upload_chunks"`
DownloadChunks uint64 `json:"download_chunks"`
UploadDuration time.Duration `json:"upload_duration"`
DownloadDuration time.Duration `json:"download_duration"`
RetainedHeapBytes uint64 `json:"retained_heap_bytes"`
Mode string `json:"mode"`
CreateDirs bool `json:"create_dirs"`
UploadReplayRejected bool `json:"upload_replay_rejected"`
DownloadReplayRejected bool `json:"download_replay_rejected"`
OversizeRejected bool `json:"oversize_rejected"`
AgentTempResidue int `json:"agent_temp_residue"`
DashboardSpoolResidue int `json:"dashboard_spool_residue"`
OutsideRootSentinelsUnchanged bool `json:"outside_root_sentinels_unchanged"`
}
func (e TransferEvidence) Validate() error {
var validationErr error
if e.WarmupUploadBytes != transferWarmupBytes || e.WarmupDownloadBytes != transferWarmupBytes || e.WarmupSHA256 == "" || e.WarmupDuration <= 0 || e.WarmupDeadlineRemaining <= 0 || !e.WarmupQuiescent {
validationErr = errors.Join(validationErr, fmt.Errorf("%w: warmup_upload=%d warmup_download=%d warmup_sha=%q warmup_duration=%s warmup_deadline=%s warmup_quiescent=%t", ErrTransferEvidenceMeasurement, e.WarmupUploadBytes, e.WarmupDownloadBytes, e.WarmupSHA256, e.WarmupDuration, e.WarmupDeadlineRemaining, e.WarmupQuiescent))
}
if e.UploadBytes != contract.TransferBytes || e.DownloadBytes != contract.TransferBytes {
validationErr = errors.Join(validationErr, fmt.Errorf("%w: upload=%d download=%d want=%d", ErrTransferEvidenceSizeMismatch, e.UploadBytes, e.DownloadBytes, contract.TransferBytes))
}
if e.UploadSHA256 == "" || e.DownloadSHA256 == "" || !strings.EqualFold(e.UploadSHA256, e.DownloadSHA256) {
validationErr = errors.Join(validationErr, fmt.Errorf("%w: upload=%q download=%q", ErrTransferEvidenceHashMismatch, e.UploadSHA256, e.DownloadSHA256))
}
if e.UploadChunks == 0 || e.DownloadChunks == 0 || e.UploadDuration <= 0 || e.DownloadDuration <= 0 {
validationErr = errors.Join(validationErr, fmt.Errorf("%w: upload_chunks=%d download_chunks=%d upload_duration=%s download_duration=%s", ErrTransferEvidenceMeasurement, e.UploadChunks, e.DownloadChunks, e.UploadDuration, e.DownloadDuration))
}
if e.RetainedHeapBytes > contract.TransferHeapBytes {
validationErr = errors.Join(validationErr, fmt.Errorf("%w: retained=%d limit=%d", ErrTransferEvidenceHeapBudget, e.RetainedHeapBytes, contract.TransferHeapBytes))
}
if e.Mode != "0640" || !e.CreateDirs || !e.UploadReplayRejected || !e.DownloadReplayRejected || !e.OversizeRejected || e.AgentTempResidue != 0 || e.DashboardSpoolResidue != 0 || !e.OutsideRootSentinelsUnchanged {
validationErr = errors.Join(validationErr, fmt.Errorf("%w: mode=%q create_dirs=%t upload_replay=%t download_replay=%t oversize=%t agent_temp=%d dashboard_spool=%d sentinels=%t", ErrTransferEvidenceContract, e.Mode, e.CreateDirs, e.UploadReplayRejected, e.DownloadReplayRejected, e.OversizeRejected, e.AgentTempResidue, e.DashboardSpoolResidue, e.OutsideRootSentinelsUnchanged))
}
return validationErr
}
@@ -0,0 +1,157 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"os"
"strings"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
type transferExecution struct {
client *client.Client
serverID uint64
root fixture.AgentRoot
residueScope transferResidueScope
sentinels transferSentinels
}
func (execution transferExecution) run(ctx context.Context, assertions *AssertionSet, fault contract.Fault) (TransferEvidence, error) {
payload, err := fixture.NewPayload(contract.DefaultSeed, contract.TransferBytes)
if err != nil {
return TransferEvidence{}, err
}
stableDigest, err := fixture.VerifyPayload(payload.Reader(), contract.TransferBytes)
if err != nil {
return TransferEvidence{}, err
}
warmupEvidence, err := execution.runWarmup(ctx)
if err != nil {
return TransferEvidence{}, fmt.Errorf("transfer warm-up: %w", err)
}
quiescenceDeadline, err := confirmTransferQuiescence(ctx, execution.residueScope)
if err != nil {
return TransferEvidence{}, err
}
warmupEvidence.deadline = quiescenceDeadline
assertions.Record("small real upload and download warm-up precedes event and deadline quiescence", warmupEvidence.valid(), fmt.Sprintf("completion_event=download_response upload_bytes=%d download_bytes=%d sha256=%s duration=%s deadline_remaining=%s", warmupEvidence.uploadBytes, warmupEvidence.downloadBytes, warmupEvidence.sha256, warmupEvidence.duration, warmupEvidence.deadline))
uploadPath, err := execution.root.Path("measured/nested/upload.bin")
if err != nil {
return TransferEvidence{}, err
}
heapProbe := fixture.NewRetainedHeapProbe()
uploadEvidence, uploadURL, uploadErr := execution.upload(ctx, uploadPath, payload, stableDigest, fault)
if fault.String() == "transfer-hash" {
return execution.finishHashFault(ctx, assertions, uploadPath, warmupEvidence, uploadErr)
}
if uploadErr != nil {
return TransferEvidence{}, uploadErr
}
assertions.Record("exact 100MiB upload has size mode SHA and create_dirs", uploadEvidence.validUpload(stableDigest), uploadEvidence.details())
downloadEvidence, downloadURL, downloadErr := execution.download(ctx, uploadPath)
if downloadErr != nil {
return TransferEvidence{}, downloadErr
}
assertions.Record("exact 100MiB download has equal nonempty SHA", downloadEvidence.validDownload(stableDigest), downloadEvidence.details())
retainedHeapBytes := heapProbe.RetainedBytes()
uploadReplayErr := execution.replayUpload(ctx, uploadURL)
uploadReplayRejected := isTransferHTTPError(uploadReplayErr, 401, "already-used")
assertions.Record("upload token replay is typed unauthorized", uploadReplayRejected, errorText(uploadReplayErr))
downloadReplayErr := execution.replayDownload(ctx, downloadURL)
downloadReplayRejected := isTransferHTTPError(downloadReplayErr, 401, "")
assertions.Record("download token replay is typed unauthorized", downloadReplayRejected, errorText(downloadReplayErr))
oversizeErr := execution.probeOversize(ctx)
oversizeRejected := isTransferHTTPError(oversizeErr, 413, "transfer cap")
assertions.Record("100MiB plus one upload is typed too large", oversizeRejected, errorText(oversizeErr))
if _, err := confirmTransferQuiescence(ctx, execution.residueScope); err != nil {
return TransferEvidence{}, err
}
residue, err := transferResidue(execution.residueScope)
if err != nil {
return TransferEvidence{}, err
}
agentResidue, dashboardResidue := countTransferResidue(residue)
assertions.Record("Dashboard spool and Agent temp residue are zero", agentResidue == 0 && dashboardResidue == 0, fmt.Sprintf("agent_temp=%d dashboard_spool=%d", agentResidue, dashboardResidue))
sentinelsUnchanged, sentinelErr := execution.sentinels.unchanged()
assertions.Record("outside-root sentinels remain unchanged", sentinelErr == nil && sentinelsUnchanged, errorText(sentinelErr))
evidence := TransferEvidence{
WarmupUploadBytes: warmupEvidence.uploadBytes, WarmupDownloadBytes: warmupEvidence.downloadBytes,
WarmupSHA256: warmupEvidence.sha256, WarmupDuration: warmupEvidence.duration,
WarmupDeadlineRemaining: warmupEvidence.deadline, WarmupQuiescent: true,
UploadBytes: uploadEvidence.bytes, DownloadBytes: downloadEvidence.bytes,
UploadSHA256: uploadEvidence.sha256, DownloadSHA256: downloadEvidence.sha256,
UploadChunks: uploadEvidence.chunks, DownloadChunks: downloadEvidence.chunks,
UploadDuration: uploadEvidence.duration, DownloadDuration: downloadEvidence.duration,
RetainedHeapBytes: retainedHeapBytes, Mode: "0640", CreateDirs: true,
UploadReplayRejected: uploadReplayRejected, DownloadReplayRejected: downloadReplayRejected, OversizeRejected: oversizeRejected,
AgentTempResidue: agentResidue, DashboardSpoolResidue: dashboardResidue,
OutsideRootSentinelsUnchanged: sentinelsUnchanged,
}
heapErr := evidence.Validate()
assertions.Record("retained live heap stays within 16MiB", !errors.Is(heapErr, ErrTransferEvidenceHeapBudget), fmt.Sprintf("retained_heap_bytes=%d", evidence.RetainedHeapBytes))
if err := evidence.Validate(); err != nil {
return evidence, err
}
return evidence, nil
}
func (execution transferExecution) finishHashFault(ctx context.Context, assertions *AssertionSet, uploadPath fixture.AgentPath, warmup transferWarmupEvidence, uploadErr error) (TransferEvidence, error) {
assertions.Record("transfer-hash rejects upload with typed 502", isTransferHTTPError(uploadErr, 502, "sha256 mismatch"), errorText(uploadErr))
_, statErr := os.Stat(uploadPath.String())
assertions.Record("transfer-hash leaves target absent", errors.Is(statErr, os.ErrNotExist), errorText(statErr))
if _, err := confirmTransferQuiescence(ctx, execution.residueScope); err != nil {
return TransferEvidence{}, err
}
residue, err := transferResidue(execution.residueScope)
if err != nil {
return TransferEvidence{}, err
}
agentResidue, dashboardResidue := countTransferResidue(residue)
assertions.Record("Dashboard spool and Agent temp residue are zero", agentResidue == 0 && dashboardResidue == 0, fmt.Sprintf("agent_temp=%d dashboard_spool=%d", agentResidue, dashboardResidue))
unchanged, sentinelErr := execution.sentinels.unchanged()
assertions.Record("outside-root sentinels remain unchanged", sentinelErr == nil && unchanged, errorText(sentinelErr))
return TransferEvidence{
WarmupUploadBytes: warmup.uploadBytes, WarmupDownloadBytes: warmup.downloadBytes,
WarmupSHA256: warmup.sha256, WarmupDuration: warmup.duration,
WarmupDeadlineRemaining: warmup.deadline, WarmupQuiescent: true,
AgentTempResidue: agentResidue, DashboardSpoolResidue: dashboardResidue, OutsideRootSentinelsUnchanged: unchanged,
}, ErrTransferHashFault
}
func isTransferHTTPError(err error, status int, text string) bool {
var httpError *client.HTTPError
return errors.As(err, &httpError) && httpError.StatusCode == status && (text == "" || strings.Contains(httpError.Message, text))
}
type transferPathEvidence struct {
bytes uint64
sha256 string
chunks uint64
duration time.Duration
mode os.FileMode
}
func (evidence transferPathEvidence) details() string {
return fmt.Sprintf("bytes=%d sha256=%s chunks=%d duration=%s mode=%04o", evidence.bytes, evidence.sha256, evidence.chunks, evidence.duration, evidence.mode.Perm())
}
func (evidence transferPathEvidence) validUpload(want fixture.PayloadDigest) bool {
return evidence.validDownload(want) && evidence.mode.Perm() == 0o640
}
func (evidence transferPathEvidence) validDownload(want fixture.PayloadDigest) bool {
return evidence.bytes == contract.TransferBytes && evidence.sha256 != "" && evidence.sha256 == want.Hex()
}
@@ -0,0 +1,86 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
"time"
)
type transferResidueScope struct {
AgentRoot string
DashboardPID int
}
func confirmTransferQuiescence(ctx context.Context, scope transferResidueScope) (time.Duration, error) {
deadline, bounded := ctx.Deadline()
if !bounded {
return 0, errors.New("transfer quiescence requires a context deadline")
}
select {
case <-ctx.Done():
return 0, ctx.Err()
default:
}
residue, err := transferResidue(scope)
if err != nil {
return 0, err
}
if len(residue) > 0 {
return 0, fmt.Errorf("transfer completion left residue: %s", strings.Join(residue, ", "))
}
remaining := time.Until(deadline)
if remaining <= 0 {
return 0, context.DeadlineExceeded
}
return remaining, nil
}
func transferResidue(scope transferResidueScope) ([]string, error) {
var residue []string
err := filepath.WalkDir(scope.AgentRoot, func(path string, entry os.DirEntry, walkErr error) error {
if walkErr != nil {
return walkErr
}
if !entry.IsDir() && strings.HasPrefix(entry.Name(), ".mcp-xfer-") {
residue = append(residue, path)
}
return nil
})
if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, err
}
fdDirectory := filepath.Join("/proc", strconv.Itoa(scope.DashboardPID), "fd")
entries, err := os.ReadDir(fdDirectory)
if err != nil {
return nil, fmt.Errorf("read dashboard descriptors: %w", err)
}
for _, entry := range entries {
target, readErr := os.Readlink(filepath.Join(fdDirectory, entry.Name()))
if readErr != nil {
continue
}
if strings.Contains(target, "nz-mcp-xfer-") {
residue = append(residue, target)
}
}
return residue, nil
}
func countTransferResidue(residue []string) (agentTemp, dashboardSpool int) {
for _, path := range residue {
if strings.Contains(filepath.Base(path), ".mcp-xfer-") {
agentTemp++
}
if strings.Contains(path, "nz-mcp-xfer-") {
dashboardSpool++
}
}
return agentTemp, dashboardSpool
}
@@ -0,0 +1,183 @@
//go:build linux
package scenario
import (
"context"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"os"
"strings"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
func (execution transferExecution) upload(ctx context.Context, path fixture.AgentPath, payload fixture.Payload, digest fixture.PayloadDigest, fault contract.Fault) (transferPathEvidence, client.TransferURL, error) {
transferURL, err := client.RequestUploadURL(ctx, execution.client, client.UploadURLRequest{
ServerID: execution.serverID, Path: path.String(), TTLSeconds: 60, Mode: "0640", CreateDirs: true,
})
if err != nil {
return transferPathEvidence{}, client.TransferURL{}, err
}
transferClient, err := clientForTransferURL(transferURL)
if err != nil {
return transferPathEvidence{}, client.TransferURL{}, err
}
defer transferClient.Close()
measured := fixture.NewMeasuredReader(payload.Reader())
expectedSHA := digest.Hex()
if fault.String() == "transfer-hash" {
expectedSHA = strings.Repeat("0", 64)
}
started := time.Now()
result, err := transferClient.UploadTransfer(ctx, transferURL, client.UploadTransfer{
Body: measured, ContentLength: int64(contract.TransferBytes), SHA256: expectedSHA,
})
duration := time.Since(started)
measurement := measured.Measurement()
if err != nil {
return transferPathEvidence{bytes: measurement.Digest.Bytes, sha256: measurement.Digest.Hex(), chunks: measurement.Chunks, duration: duration}, transferURL, err
}
fileDigest, info, err := verifyUploadedTransfer(path)
if err != nil {
return transferPathEvidence{}, transferURL, err
}
if result.Size != info.Size() || result.SHA256 != fileDigest.Hex() {
return transferPathEvidence{}, transferURL, errors.New("upload result differs from Agent file")
}
return transferPathEvidence{bytes: uint64(result.Size), sha256: result.SHA256, chunks: measurement.Chunks, duration: duration, mode: info.Mode()}, transferURL, nil
}
func verifyUploadedTransfer(path fixture.AgentPath) (digest fixture.PayloadDigest, info os.FileInfo, err error) {
file, err := os.Open(path.String())
if err != nil {
return fixture.PayloadDigest{}, nil, fmt.Errorf("open uploaded transfer: %w", err)
}
defer func() { err = errors.Join(err, file.Close()) }()
digest, err = fixture.VerifyPayload(file, contract.TransferBytes)
if err != nil {
return fixture.PayloadDigest{}, nil, err
}
info, err = file.Stat()
if err != nil {
return fixture.PayloadDigest{}, nil, fmt.Errorf("stat uploaded transfer: %w", err)
}
return digest, info, nil
}
func (execution transferExecution) download(ctx context.Context, path fixture.AgentPath) (transferPathEvidence, client.TransferURL, error) {
transferURL, err := client.RequestDownloadURL(ctx, execution.client, client.DownloadURLRequest{
ServerID: execution.serverID, Path: path.String(), TTLSeconds: 60,
})
if err != nil {
return transferPathEvidence{}, client.TransferURL{}, err
}
transferClient, err := clientForTransferURL(transferURL)
if err != nil {
return transferPathEvidence{}, client.TransferURL{}, err
}
defer transferClient.Close()
measured := fixture.NewMeasuredWriter()
started := time.Now()
written, err := transferClient.DownloadTransfer(ctx, transferURL, measured)
duration := time.Since(started)
measurement := measured.Measurement()
if err != nil {
return transferPathEvidence{}, transferURL, err
}
if written != int64(measurement.Digest.Bytes) {
return transferPathEvidence{}, transferURL, errors.New("download byte count differs from measured digest")
}
return transferPathEvidence{bytes: measurement.Digest.Bytes, sha256: measurement.Digest.Hex(), chunks: measurement.Chunks, duration: duration}, transferURL, nil
}
func (execution transferExecution) replayUpload(ctx context.Context, transferURL client.TransferURL) error {
transferClient, err := clientForTransferURL(transferURL)
if err != nil {
return err
}
defer transferClient.Close()
payload, err := fixture.NewPayload(contract.DefaultSeed, 1)
if err != nil {
return err
}
digest, err := fixture.VerifyPayload(payload.Reader(), 1)
if err != nil {
return err
}
_, err = transferClient.UploadTransfer(ctx, transferURL, client.UploadTransfer{Body: payload.Reader(), ContentLength: 1, SHA256: digest.Hex()})
return err
}
func (execution transferExecution) replayDownload(ctx context.Context, transferURL client.TransferURL) error {
transferClient, err := clientForTransferURL(transferURL)
if err != nil {
return err
}
defer transferClient.Close()
_, err = transferClient.DownloadTransfer(ctx, transferURL, io.Discard)
return err
}
func (execution transferExecution) probeOversize(ctx context.Context) error {
path, err := execution.root.Path("oversize/rejected.bin")
if err != nil {
return err
}
transferURL, err := client.RequestUploadURL(ctx, execution.client, client.UploadURLRequest{
ServerID: execution.serverID, Path: path.String(), TTLSeconds: 60, CreateDirs: true,
})
if err != nil {
return err
}
transferClient, err := clientForTransferURL(transferURL)
if err != nil {
return err
}
defer transferClient.Close()
return transferClient.ProbeOversizeUpload(ctx, transferURL, client.OversizeUploadProbe{
Body: zeroReader{}, ContentLength: int64(contract.TransferBytes) + 1,
})
}
type ownedTransferClient struct {
*client.Client
transport *http.Transport
}
func clientForTransferURL(transferURL client.TransferURL) (*ownedTransferClient, error) {
parsed, err := url.Parse(transferURL.URL)
if err != nil {
return nil, fmt.Errorf("parse transfer URL origin: %w", err)
}
defaultTransport, ok := http.DefaultTransport.(*http.Transport)
if !ok {
return nil, errors.New("default HTTP transport is not cloneable")
}
transport := defaultTransport.Clone()
transferClient, err := client.New(client.Config{
BaseURL: parsed.Scheme + "://" + parsed.Host, HTTPClient: &http.Client{Transport: transport}, TransferTimeout: 5 * time.Minute,
})
if err != nil {
transport.CloseIdleConnections()
return nil, err
}
return &ownedTransferClient{Client: transferClient, transport: transport}, nil
}
func (client *ownedTransferClient) Close() {
client.transport.CloseIdleConnections()
}
type zeroReader struct{}
func (zeroReader) Read(destination []byte) (int, error) {
clear(destination)
return len(destination), nil
}
@@ -0,0 +1,149 @@
//go:build linux && agentcompat
package scenario
import (
"context"
"encoding/json"
"os"
"path/filepath"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
)
func TestTransferScenario_RealFlow(t *testing.T) {
nezhaSource := os.Getenv("AGENTCOMPAT_NEZHA_SOURCE")
agentSource := os.Getenv("AGENTCOMPAT_AGENT_SOURCE")
if nezhaSource == "" || agentSource == "" {
t.Skip("set AGENTCOMPAT_NEZHA_SOURCE and AGENTCOMPAT_AGENT_SOURCE")
}
evidenceDirectory := os.Getenv("AGENTCOMPAT_TRANSFER_EVIDENCE_DIR")
if evidenceDirectory == "" {
evidenceDirectory = t.TempDir()
}
require.NoError(t, os.MkdirAll(evidenceDirectory, 0o700))
paths, err := contract.NewPaths(nezhaSource, agentSource, evidenceDirectory)
require.NoError(t, err)
testContext, cancel := context.WithTimeout(t.Context(), 10*time.Minute)
defer cancel()
result, transferEvidence, err := (Transfer{}).RunWithEvidence(testContext, TransferInput{Paths: paths})
logTransferAssertions(t, result)
require.NoError(t, err)
require.True(t, result.Passed)
require.True(t, result.CleanupOK)
require.NoError(t, transferEvidence.Validate())
require.Equal(t, uint64(transferWarmupBytes), transferEvidence.WarmupUploadBytes)
require.Equal(t, transferEvidence.WarmupUploadBytes, transferEvidence.WarmupDownloadBytes)
require.NotEmpty(t, transferEvidence.WarmupSHA256)
require.Positive(t, transferEvidence.WarmupDuration)
require.Positive(t, transferEvidence.WarmupDeadlineRemaining)
require.True(t, transferEvidence.WarmupQuiescent)
require.Equal(t, uint64(104857600), transferEvidence.UploadBytes)
require.Equal(t, uint64(104857600), transferEvidence.DownloadBytes)
require.Equal(t, transferEvidence.UploadSHA256, transferEvidence.DownloadSHA256)
require.NotEmpty(t, transferEvidence.UploadSHA256)
require.Positive(t, transferEvidence.UploadChunks)
require.Positive(t, transferEvidence.DownloadChunks)
require.Positive(t, transferEvidence.UploadDuration)
require.Positive(t, transferEvidence.DownloadDuration)
require.LessOrEqual(t, transferEvidence.RetainedHeapBytes, uint64(16777216))
require.Equal(t, "0640", transferEvidence.Mode)
require.True(t, transferEvidence.CreateDirs)
require.True(t, transferEvidence.UploadReplayRejected)
require.True(t, transferEvidence.DownloadReplayRejected)
require.True(t, transferEvidence.OversizeRejected)
require.Zero(t, transferEvidence.AgentTempResidue)
require.Zero(t, transferEvidence.DashboardSpoolResidue)
require.True(t, transferEvidence.OutsideRootSentinelsUnchanged)
requireTransferAssertions(t, result,
"small real upload and download warm-up precedes event and deadline quiescence",
"exact 100MiB upload has size mode SHA and create_dirs",
"exact 100MiB download has equal nonempty SHA",
"upload token replay is typed unauthorized",
"download token replay is typed unauthorized",
"100MiB plus one upload is typed too large",
"retained live heap stays within 16MiB",
"outside-root sentinels remain unchanged",
"Dashboard spool and Agent temp residue are zero",
"process listener and workspace cleanup completed",
)
fault, err := contract.NewFault("transfer-hash")
require.NoError(t, err)
faultResult, faultEvidence, faultErr := (Transfer{}).RunWithEvidence(testContext, TransferInput{Paths: paths, Fault: fault})
logTransferAssertions(t, faultResult)
require.ErrorIs(t, faultErr, ErrTransferHashFault)
require.False(t, faultResult.Passed)
require.True(t, faultResult.CleanupOK)
require.Equal(t, uint64(transferWarmupBytes), faultEvidence.WarmupUploadBytes)
require.Equal(t, faultEvidence.WarmupUploadBytes, faultEvidence.WarmupDownloadBytes)
require.NotEmpty(t, faultEvidence.WarmupSHA256)
require.Positive(t, faultEvidence.WarmupDuration)
require.Positive(t, faultEvidence.WarmupDeadlineRemaining)
require.True(t, faultEvidence.WarmupQuiescent)
require.Zero(t, faultEvidence.AgentTempResidue)
require.Zero(t, faultEvidence.DashboardSpoolResidue)
require.True(t, faultEvidence.OutsideRootSentinelsUnchanged)
requireTransferAssertions(t, faultResult,
"small real upload and download warm-up precedes event and deadline quiescence",
"transfer-hash rejects upload with typed 502",
"transfer-hash leaves target absent",
"outside-root sentinels remain unchanged",
"Dashboard spool and Agent temp residue are zero",
"process listener and workspace cleanup completed",
)
artifactPath := filepath.Join(evidenceDirectory, "transfer-real-process.json")
artifact, err := json.MarshalIndent(struct {
SuccessResult Result `json:"success_result"`
SuccessEvidence TransferEvidence `json:"success_evidence"`
FaultResult Result `json:"fault_result"`
FaultEvidence TransferEvidence `json:"fault_evidence"`
FaultError string `json:"fault_error"`
}{SuccessResult: result, SuccessEvidence: transferEvidence, FaultResult: faultResult, FaultEvidence: faultEvidence, FaultError: faultErr.Error()}, "", " ")
require.NoError(t, err)
require.NoError(t, os.WriteFile(artifactPath, append(artifact, '\n'), 0o600))
readArtifact, err := os.ReadFile(artifactPath)
require.NoError(t, err)
var recorded struct {
SuccessResult Result `json:"success_result"`
SuccessEvidence TransferEvidence `json:"success_evidence"`
FaultResult Result `json:"fault_result"`
FaultEvidence TransferEvidence `json:"fault_evidence"`
FaultError string `json:"fault_error"`
}
require.NoError(t, json.Unmarshal(readArtifact, &recorded))
require.Equal(t, result, recorded.SuccessResult)
require.Equal(t, transferEvidence, recorded.SuccessEvidence)
require.Equal(t, faultResult, recorded.FaultResult)
require.Equal(t, faultEvidence, recorded.FaultEvidence)
require.Equal(t, faultErr.Error(), recorded.FaultError)
require.NoError(t, recorded.SuccessEvidence.Validate())
t.Logf("transfer evidence artifact: %s", artifactPath)
}
func logTransferAssertions(t *testing.T, result Result) {
t.Helper()
for _, assertion := range result.Assertions {
t.Logf("assertion=%q passed=%t details=%q", assertion.Name, assertion.Passed, assertion.Details)
}
}
func requireTransferAssertions(t *testing.T, result Result, names ...string) {
t.Helper()
byName := make(map[string]Assertion, len(result.Assertions))
for _, assertion := range result.Assertions {
byName[assertion.Name] = assertion
}
for _, name := range names {
assertion, exists := byName[name]
require.True(t, exists, "missing assertion %q", name)
require.True(t, assertion.Passed, "assertion %q failed: %s", name, assertion.Details)
}
}
@@ -0,0 +1,57 @@
//go:build linux
package scenario
import (
"bytes"
"errors"
"fmt"
"os"
"path/filepath"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
var transferSentinelContent = []byte("outside-transfer-root-sentinel")
type transferSentinels struct {
paths []string
}
func newTransferSentinels(root fixture.AgentRoot, workspaceRoot string) (transferSentinels, error) {
directPath := filepath.Join(workspaceRoot, "outside-transfer-sentinel")
symlinkDirectory := filepath.Join(workspaceRoot, "outside-transfer-directory")
symlinkTarget := filepath.Join(symlinkDirectory, "target-sentinel")
if err := os.Mkdir(symlinkDirectory, 0o700); err != nil {
return transferSentinels{}, fmt.Errorf("create transfer sentinel directory: %w", err)
}
for _, path := range []string{directPath, symlinkTarget} {
if err := os.WriteFile(path, transferSentinelContent, 0o600); err != nil {
return transferSentinels{}, fmt.Errorf("write transfer sentinel: %w", err)
}
}
symlinkPath := filepath.Join(root.Absolute(), "linked")
if err := os.Symlink(symlinkDirectory, symlinkPath); err != nil {
return transferSentinels{}, fmt.Errorf("create transfer sentinel symlink: %w", err)
}
if _, err := root.Path("../outside-transfer-sentinel"); err == nil {
return transferSentinels{}, errors.New("transfer AgentPath accepted parent escape")
}
if _, err := root.Path("linked/target-sentinel"); err == nil {
return transferSentinels{}, errors.New("transfer AgentPath accepted symlink parent")
}
return transferSentinels{paths: []string{directPath, symlinkTarget}}, nil
}
func (sentinels transferSentinels) unchanged() (bool, error) {
for _, path := range sentinels.paths {
content, err := os.ReadFile(path)
if err != nil {
return false, fmt.Errorf("read transfer sentinel: %w", err)
}
if !bytes.Equal(content, transferSentinelContent) {
return false, nil
}
}
return true, nil
}
@@ -0,0 +1,126 @@
//go:build linux
package scenario
import (
"context"
"testing"
"time"
"github.com/stretchr/testify/require"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
func TestTransferEvidence_RejectsEmptyOrMismatchedSHA(t *testing.T) {
t.Parallel()
evidence := validTransferEvidence()
evidence.UploadSHA256 = ""
evidence.DownloadSHA256 = "different"
err := evidence.Validate()
require.ErrorIs(t, err, ErrTransferEvidenceHashMismatch)
require.NotErrorIs(t, err, ErrTransferEvidenceSizeMismatch)
require.NotErrorIs(t, err, ErrTransferEvidenceMeasurement)
require.NotErrorIs(t, err, ErrTransferEvidenceHeapBudget)
require.NotErrorIs(t, err, ErrTransferEvidenceContract)
}
func TestTransferEvidence_RejectsHeapBudgetBreach(t *testing.T) {
t.Parallel()
evidence := validTransferEvidence()
evidence.RetainedHeapBytes = contract.TransferHeapBytes + 1
err := evidence.Validate()
require.ErrorIs(t, err, ErrTransferEvidenceHeapBudget)
require.NotErrorIs(t, err, ErrTransferEvidenceSizeMismatch)
require.NotErrorIs(t, err, ErrTransferEvidenceHashMismatch)
require.NotErrorIs(t, err, ErrTransferEvidenceMeasurement)
require.NotErrorIs(t, err, ErrTransferEvidenceContract)
}
func TestTransferEvidence_AcceptsExactNonemptyEqualHashes(t *testing.T) {
t.Parallel()
require.NoError(t, validTransferEvidence().Validate())
}
func TestTransferEvidence_RejectsMissingStreamingMeasurements(t *testing.T) {
t.Parallel()
evidence := validTransferEvidence()
evidence.UploadChunks = 0
evidence.DownloadChunks = 0
evidence.UploadDuration = 0
evidence.DownloadDuration = 0
err := evidence.Validate()
require.ErrorIs(t, err, ErrTransferEvidenceMeasurement)
require.NotErrorIs(t, err, ErrTransferEvidenceSizeMismatch)
require.NotErrorIs(t, err, ErrTransferEvidenceHashMismatch)
require.NotErrorIs(t, err, ErrTransferEvidenceHeapBudget)
require.NotErrorIs(t, err, ErrTransferEvidenceContract)
}
func TestTransferEvidence_ReportsAllTypedValidationErrors(t *testing.T) {
t.Parallel()
err := (TransferEvidence{}).Validate()
require.ErrorIs(t, err, ErrTransferEvidenceSizeMismatch)
require.ErrorIs(t, err, ErrTransferEvidenceHashMismatch)
require.ErrorIs(t, err, ErrTransferEvidenceMeasurement)
require.ErrorIs(t, err, ErrTransferEvidenceContract)
}
func TestTransferUploadEvidence_RejectsMissingMode(t *testing.T) {
t.Parallel()
digest := fixture.PayloadDigest{Bytes: contract.TransferBytes, SHA256: [32]byte{1}}
evidence := transferPathEvidence{
bytes: contract.TransferBytes, sha256: digest.Hex(), chunks: 1, duration: time.Nanosecond,
}
require.False(t, evidence.validUpload(digest))
require.True(t, evidence.validDownload(digest))
}
func TestConfirmTransferQuiescence_RequiresDeadline(t *testing.T) {
t.Parallel()
_, err := confirmTransferQuiescence(context.Background(), transferResidueScope{})
require.EqualError(t, err, "transfer quiescence requires a context deadline")
}
func validTransferEvidence() TransferEvidence {
return TransferEvidence{
WarmupUploadBytes: transferWarmupBytes,
WarmupDownloadBytes: transferWarmupBytes,
WarmupSHA256: "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb",
WarmupDuration: time.Nanosecond,
WarmupDeadlineRemaining: time.Second,
WarmupQuiescent: true,
UploadBytes: contract.TransferBytes,
DownloadBytes: contract.TransferBytes,
UploadSHA256: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa",
DownloadSHA256: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa",
UploadChunks: 1,
DownloadChunks: 1,
UploadDuration: time.Nanosecond,
DownloadDuration: time.Nanosecond,
RetainedHeapBytes: contract.TransferHeapBytes,
Mode: "0640",
CreateDirs: true,
UploadReplayRejected: true,
DownloadReplayRejected: true,
OversizeRejected: true,
OutsideRootSentinelsUnchanged: true,
}
}
@@ -0,0 +1,87 @@
//go:build linux
package scenario
import (
"context"
"errors"
"time"
"github.com/nezhahq/nezha/integration/agentcompat/internal/client"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
"github.com/nezhahq/nezha/integration/agentcompat/internal/fixture"
)
const transferWarmupBytes = 64 * 1024
type transferWarmupEvidence struct {
uploadBytes uint64
downloadBytes uint64
sha256 string
duration time.Duration
deadline time.Duration
}
func (execution transferExecution) runWarmup(ctx context.Context) (transferWarmupEvidence, error) {
payload, err := fixture.NewPayload(contract.DefaultSeed, transferWarmupBytes)
if err != nil {
return transferWarmupEvidence{}, err
}
digest, err := fixture.VerifyPayload(payload.Reader(), transferWarmupBytes)
if err != nil {
return transferWarmupEvidence{}, err
}
path, err := execution.root.Path("warmup/nested/payload.bin")
if err != nil {
return transferWarmupEvidence{}, err
}
uploadURL, err := client.RequestUploadURL(ctx, execution.client, client.UploadURLRequest{
ServerID: execution.serverID, Path: path.String(), TTLSeconds: 60, Mode: "0600", CreateDirs: true,
})
if err != nil {
return transferWarmupEvidence{}, err
}
uploadClient, err := clientForTransferURL(uploadURL)
if err != nil {
return transferWarmupEvidence{}, err
}
defer uploadClient.Close()
started := time.Now()
uploadResult, err := uploadClient.UploadTransfer(ctx, uploadURL, client.UploadTransfer{
Body: payload.Reader(), ContentLength: transferWarmupBytes, SHA256: digest.Hex(),
})
if err != nil {
return transferWarmupEvidence{}, err
}
if uploadResult.Size != transferWarmupBytes || uploadResult.SHA256 != digest.Hex() {
return transferWarmupEvidence{}, errors.New("warm-up upload evidence mismatch")
}
downloadURL, err := client.RequestDownloadURL(ctx, execution.client, client.DownloadURLRequest{
ServerID: execution.serverID, Path: path.String(), TTLSeconds: 60,
})
if err != nil {
return transferWarmupEvidence{}, err
}
downloadClient, err := clientForTransferURL(downloadURL)
if err != nil {
return transferWarmupEvidence{}, err
}
defer downloadClient.Close()
measured := fixture.NewMeasuredWriter()
written, err := downloadClient.DownloadTransfer(ctx, downloadURL, measured)
if err != nil {
return transferWarmupEvidence{}, err
}
measurement := measured.Measurement()
if written != transferWarmupBytes || measurement.Digest.Bytes != transferWarmupBytes || measurement.Digest.Hex() != digest.Hex() {
return transferWarmupEvidence{}, errors.New("warm-up download evidence mismatch")
}
return transferWarmupEvidence{
uploadBytes: uint64(uploadResult.Size), downloadBytes: measurement.Digest.Bytes,
sha256: digest.Hex(), duration: time.Since(started),
}, nil
}
func (evidence transferWarmupEvidence) valid() bool {
return evidence.uploadBytes == transferWarmupBytes && evidence.downloadBytes == transferWarmupBytes && evidence.sha256 != "" && evidence.duration > 0 && evidence.deadline > 0
}