mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 09:40:12 +00:00
test(agentcompat): add integration scenarios
Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user