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