diff --git a/integration/agentcompat/internal/scenario/config_file.go b/integration/agentcompat/internal/scenario/config_file.go new file mode 100644 index 00000000..2de52452 --- /dev/null +++ b/integration/agentcompat/internal/scenario/config_file.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/foundation.go b/integration/agentcompat/internal/scenario/foundation.go new file mode 100644 index 00000000..bd6b9508 --- /dev/null +++ b/integration/agentcompat/internal/scenario/foundation.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_cleanup.go b/integration/agentcompat/internal/scenario/held_cleanup.go new file mode 100644 index 00000000..5dc06c52 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_cleanup.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_cleanup_test.go b/integration/agentcompat/internal/scenario/held_cleanup_test.go new file mode 100644 index 00000000..3ef26e90 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_cleanup_test.go @@ -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) + } +} diff --git a/integration/agentcompat/internal/scenario/held_fm_nat_real_test.go b/integration/agentcompat/internal/scenario/held_fm_nat_real_test.go new file mode 100644 index 00000000..084f9429 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_fm_nat_real_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_io_stream_capability.go b/integration/agentcompat/internal/scenario/held_io_stream_capability.go new file mode 100644 index 00000000..6395d695 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_io_stream_capability.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_io_stream_capability_test.go b/integration/agentcompat/internal/scenario/held_io_stream_capability_test.go new file mode 100644 index 00000000..1571632c --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_io_stream_capability_test.go @@ -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) + }) + } +} diff --git a/integration/agentcompat/internal/scenario/held_legacy_fm.go b/integration/agentcompat/internal/scenario/held_legacy_fm.go new file mode 100644 index 00000000..a94a7913 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_legacy_fm.go @@ -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) diff --git a/integration/agentcompat/internal/scenario/held_legacy_fm_cleanup.go b/integration/agentcompat/internal/scenario/held_legacy_fm_cleanup.go new file mode 100644 index 00000000..0421ad92 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_legacy_fm_cleanup.go @@ -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}) +} diff --git a/integration/agentcompat/internal/scenario/held_legacy_fm_constructor_fault_test.go b/integration/agentcompat/internal/scenario/held_legacy_fm_constructor_fault_test.go new file mode 100644 index 00000000..e2a3b773 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_legacy_fm_constructor_fault_test.go @@ -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), + } +} diff --git a/integration/agentcompat/internal/scenario/held_legacy_fm_dependencies.go b/integration/agentcompat/internal/scenario/held_legacy_fm_dependencies.go new file mode 100644 index 00000000..d97a94fd --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_legacy_fm_dependencies.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_legacy_fm_response_loss_test.go b/integration/agentcompat/internal/scenario/held_legacy_fm_response_loss_test.go new file mode 100644 index 00000000..b3eb96d8 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_legacy_fm_response_loss_test.go @@ -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), + } +} diff --git a/integration/agentcompat/internal/scenario/held_legacy_fm_test.go b/integration/agentcompat/internal/scenario/held_legacy_fm_test.go new file mode 100644 index 00000000..ad8f64a2 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_legacy_fm_test.go @@ -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} +} diff --git a/integration/agentcompat/internal/scenario/held_nat.go b/integration/agentcompat/internal/scenario/held_nat.go new file mode 100644 index 00000000..dbd4ce88 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_nat.go @@ -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) diff --git a/integration/agentcompat/internal/scenario/held_nat_constructor_fault_test.go b/integration/agentcompat/internal/scenario/held_nat_constructor_fault_test.go new file mode 100644 index 00000000..8b7075a5 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_nat_constructor_fault_test.go @@ -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), + } +} diff --git a/integration/agentcompat/internal/scenario/held_nat_dependencies.go b/integration/agentcompat/internal/scenario/held_nat_dependencies.go new file mode 100644 index 00000000..ddd8e569 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_nat_dependencies.go @@ -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() +} diff --git a/integration/agentcompat/internal/scenario/held_nat_proof.go b/integration/agentcompat/internal/scenario/held_nat_proof.go new file mode 100644 index 00000000..c2fb4691 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_nat_proof.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_nat_request.go b/integration/agentcompat/internal/scenario/held_nat_request.go new file mode 100644 index 00000000..d3ef252f --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_nat_request.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_nat_test.go b/integration/agentcompat/internal/scenario/held_nat_test.go new file mode 100644 index 00000000..deb6e147 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_nat_test.go @@ -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()} +} diff --git a/integration/agentcompat/internal/scenario/held_readiness.go b/integration/agentcompat/internal/scenario/held_readiness.go new file mode 100644 index 00000000..019475b0 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_readiness.go @@ -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} +} diff --git a/integration/agentcompat/internal/scenario/held_readiness_test.go b/integration/agentcompat/internal/scenario/held_readiness_test.go new file mode 100644 index 00000000..74943c9a --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_readiness_test.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_real_evidence.go b/integration/agentcompat/internal/scenario/held_real_evidence.go new file mode 100644 index 00000000..e5ab2551 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_real_evidence.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_real_semantics_test.go b/integration/agentcompat/internal/scenario/held_real_semantics_test.go new file mode 100644 index 00000000..ff668feb --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_real_semantics_test.go @@ -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"})) +} diff --git a/integration/agentcompat/internal/scenario/held_session.go b/integration/agentcompat/internal/scenario/held_session.go new file mode 100644 index 00000000..71065b34 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_session_adapter_test.go b/integration/agentcompat/internal/scenario/held_session_adapter_test.go new file mode 100644 index 00000000..ace91a4e --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_adapter_test.go @@ -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) diff --git a/integration/agentcompat/internal/scenario/held_session_set.go b/integration/agentcompat/internal/scenario/held_session_set.go new file mode 100644 index 00000000..d75f794c --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_close.go b/integration/agentcompat/internal/scenario/held_session_set_close.go new file mode 100644 index 00000000..76286b8d --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_close.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_health.go b/integration/agentcompat/internal/scenario/held_session_set_health.go new file mode 100644 index 00000000..347002a9 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_health.go @@ -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) }) +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_health_coordinator.go b/integration/agentcompat/internal/scenario/held_session_set_health_coordinator.go new file mode 100644 index 00000000..de43057b --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_health_coordinator.go @@ -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) }) + } +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_health_epoch_test.go b/integration/agentcompat/internal/scenario/held_session_set_health_epoch_test.go new file mode 100644 index 00000000..5d4f197f --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_health_epoch_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_public_close_test.go b/integration/agentcompat/internal/scenario/held_session_set_public_close_test.go new file mode 100644 index 00000000..27dc004c --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_public_close_test.go @@ -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) + } +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_public_failures_test.go b/integration/agentcompat/internal/scenario/held_session_set_public_failures_test.go new file mode 100644 index 00000000..c9403561 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_public_failures_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_public_fixture_test.go b/integration/agentcompat/internal/scenario/held_session_set_public_fixture_test.go new file mode 100644 index 00000000..dfbf869e --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_public_fixture_test.go @@ -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") +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_public_health_test.go b/integration/agentcompat/internal/scenario/held_session_set_public_health_test.go new file mode 100644 index 00000000..9436b183 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_public_health_test.go @@ -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() +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_public_protocol_test.go b/integration/agentcompat/internal/scenario/held_session_set_public_protocol_test.go new file mode 100644 index 00000000..2e4f0e26 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_public_protocol_test.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_public_redaction_test.go b/integration/agentcompat/internal/scenario/held_session_set_public_redaction_test.go new file mode 100644 index 00000000..6aca1a0d --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_public_redaction_test.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_public_state_test.go b/integration/agentcompat/internal/scenario/held_session_set_public_state_test.go new file mode 100644 index 00000000..0a7b6d30 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_public_state_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_public_success_test.go b/integration/agentcompat/internal/scenario/held_session_set_public_success_test.go new file mode 100644 index 00000000..9226a205 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_public_success_test.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_public_topology_test.go b/integration/agentcompat/internal/scenario/held_session_set_public_topology_test.go new file mode 100644 index 00000000..c6ee5205 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_public_topology_test.go @@ -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) + }) + } +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_real_e2e_test.go b/integration/agentcompat/internal/scenario/held_session_set_real_e2e_test.go new file mode 100644 index 00000000..8f82bd47 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_real_e2e_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_real_evidence.go b/integration/agentcompat/internal/scenario/held_session_set_real_evidence.go new file mode 100644 index 00000000..4f6dceb1 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_real_evidence.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_real_fixture.go b/integration/agentcompat/internal/scenario/held_session_set_real_fixture.go new file mode 100644 index 00000000..056c6d65 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_real_fixture.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_real_plan.go b/integration/agentcompat/internal/scenario/held_session_set_real_plan.go new file mode 100644 index 00000000..8997695a --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_real_plan.go @@ -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[:]) +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_real_support_test.go b/integration/agentcompat/internal/scenario/held_session_set_real_support_test.go new file mode 100644 index 00000000..53bd6211 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_real_support_test.go @@ -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() +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_redaction.go b/integration/agentcompat/internal/scenario/held_session_set_redaction.go new file mode 100644 index 00000000..3d2f35a9 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_redaction.go @@ -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}} +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_source_fake_test.go b/integration/agentcompat/internal/scenario/held_session_set_source_fake_test.go new file mode 100644 index 00000000..641fa979 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_source_fake_test.go @@ -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} +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_source_test.go b/integration/agentcompat/internal/scenario/held_session_set_source_test.go new file mode 100644 index 00000000..839337e2 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_source_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_test.go b/integration/agentcompat/internal/scenario/held_session_set_test.go new file mode 100644 index 00000000..595417a7 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_session_set_topology.go b/integration/agentcompat/internal/scenario/held_session_set_topology.go new file mode 100644 index 00000000..c76c0bcc --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_set_topology.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_session_test.go b/integration/agentcompat/internal/scenario/held_session_test.go new file mode 100644 index 00000000..5a4bfd51 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_session_test.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_terminal.go b/integration/agentcompat/internal/scenario/held_terminal.go new file mode 100644 index 00000000..bf7d689e --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_terminal.go @@ -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) diff --git a/integration/agentcompat/internal/scenario/held_terminal_order_test.go b/integration/agentcompat/internal/scenario/held_terminal_order_test.go new file mode 100644 index 00000000..6d367be5 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_terminal_order_test.go @@ -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]) +} diff --git a/integration/agentcompat/internal/scenario/held_terminal_protocol.go b/integration/agentcompat/internal/scenario/held_terminal_protocol.go new file mode 100644 index 00000000..6a5ccb13 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_terminal_protocol.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_terminal_real_test.go b/integration/agentcompat/internal/scenario/held_terminal_real_test.go new file mode 100644 index 00000000..31b73e0f --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_terminal_real_test.go @@ -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)})) +} diff --git a/integration/agentcompat/internal/scenario/held_terminal_test.go b/integration/agentcompat/internal/scenario/held_terminal_test.go new file mode 100644 index 00000000..6d029f9d --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_terminal_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/held_websocket_pump.go b/integration/agentcompat/internal/scenario/held_websocket_pump.go new file mode 100644 index 00000000..2349e461 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_websocket_pump.go @@ -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() + } +} diff --git a/integration/agentcompat/internal/scenario/held_websocket_pump_stop_contract_test.go b/integration/agentcompat/internal/scenario/held_websocket_pump_stop_contract_test.go new file mode 100644 index 00000000..719a6dfa --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_websocket_pump_stop_contract_test.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/held_websocket_pump_test.go b/integration/agentcompat/internal/scenario/held_websocket_pump_test.go new file mode 100644 index 00000000..cb7cc190 --- /dev/null +++ b/integration/agentcompat/internal/scenario/held_websocket_pump_test.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/legacy_fm.go b/integration/agentcompat/internal/scenario/legacy_fm.go new file mode 100644 index 00000000..6703f56e --- /dev/null +++ b/integration/agentcompat/internal/scenario/legacy_fm.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/legacy_fm_observation.go b/integration/agentcompat/internal/scenario/legacy_fm_observation.go new file mode 100644 index 00000000..e3446860 --- /dev/null +++ b/integration/agentcompat/internal/scenario/legacy_fm_observation.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/legacy_fm_protocol.go b/integration/agentcompat/internal/scenario/legacy_fm_protocol.go new file mode 100644 index 00000000..21e55a29 --- /dev/null +++ b/integration/agentcompat/internal/scenario/legacy_fm_protocol.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/legacy_fm_protocol_error_test.go b/integration/agentcompat/internal/scenario/legacy_fm_protocol_error_test.go new file mode 100644 index 00000000..4881c314 --- /dev/null +++ b/integration/agentcompat/internal/scenario/legacy_fm_protocol_error_test.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/legacy_fm_protocol_test.go b/integration/agentcompat/internal/scenario/legacy_fm_protocol_test.go new file mode 100644 index 00000000..42931e62 --- /dev/null +++ b/integration/agentcompat/internal/scenario/legacy_fm_protocol_test.go @@ -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") + } + }) + } +} diff --git a/integration/agentcompat/internal/scenario/legacy_fm_support.go b/integration/agentcompat/internal/scenario/legacy_fm_support.go new file mode 100644 index 00000000..3ebeb9a9 --- /dev/null +++ b/integration/agentcompat/internal/scenario/legacy_fm_support.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/legacy_fm_verification.go b/integration/agentcompat/internal/scenario/legacy_fm_verification.go new file mode 100644 index 00000000..d6cfe472 --- /dev/null +++ b/integration/agentcompat/internal/scenario/legacy_fm_verification.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/mcp_filesystem.go b/integration/agentcompat/internal/scenario/mcp_filesystem.go new file mode 100644 index 00000000..7d7a1ac7 --- /dev/null +++ b/integration/agentcompat/internal/scenario/mcp_filesystem.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/mcp_filesystem_agent_evidence.go b/integration/agentcompat/internal/scenario/mcp_filesystem_agent_evidence.go new file mode 100644 index 00000000..e466c32d --- /dev/null +++ b/integration/agentcompat/internal/scenario/mcp_filesystem_agent_evidence.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/mcp_filesystem_operations.go b/integration/agentcompat/internal/scenario/mcp_filesystem_operations.go new file mode 100644 index 00000000..36ecbe41 --- /dev/null +++ b/integration/agentcompat/internal/scenario/mcp_filesystem_operations.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/mcp_filesystem_safety.go b/integration/agentcompat/internal/scenario/mcp_filesystem_safety.go new file mode 100644 index 00000000..40b2d50e --- /dev/null +++ b/integration/agentcompat/internal/scenario/mcp_filesystem_safety.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/mcp_filesystem_test.go b/integration/agentcompat/internal/scenario/mcp_filesystem_test.go new file mode 100644 index 00000000..9d6ce5d5 --- /dev/null +++ b/integration/agentcompat/internal/scenario/mcp_filesystem_test.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/nat.go b/integration/agentcompat/internal/scenario/nat.go new file mode 100644 index 00000000..c19f48e7 --- /dev/null +++ b/integration/agentcompat/internal/scenario/nat.go @@ -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") +} diff --git a/integration/agentcompat/internal/scenario/nat_http.go b/integration/agentcompat/internal/scenario/nat_http.go new file mode 100644 index 00000000..407e9a69 --- /dev/null +++ b/integration/agentcompat/internal/scenario/nat_http.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/nat_test.go b/integration/agentcompat/internal/scenario/nat_test.go new file mode 100644 index 00000000..50997f33 --- /dev/null +++ b/integration/agentcompat/internal/scenario/nat_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/reconnect.go b/integration/agentcompat/internal/scenario/reconnect.go new file mode 100644 index 00000000..3bdc0575 --- /dev/null +++ b/integration/agentcompat/internal/scenario/reconnect.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/reconnect_evidence.go b/integration/agentcompat/internal/scenario/reconnect_evidence.go new file mode 100644 index 00000000..0e69cc75 --- /dev/null +++ b/integration/agentcompat/internal/scenario/reconnect_evidence.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/reconnect_fault_real_test.go b/integration/agentcompat/internal/scenario/reconnect_fault_real_test.go new file mode 100644 index 00000000..ac3f3d74 --- /dev/null +++ b/integration/agentcompat/internal/scenario/reconnect_fault_real_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/reconnect_operations.go b/integration/agentcompat/internal/scenario/reconnect_operations.go new file mode 100644 index 00000000..2d5c6947 --- /dev/null +++ b/integration/agentcompat/internal/scenario/reconnect_operations.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/reconnect_real_test.go b/integration/agentcompat/internal/scenario/reconnect_real_test.go new file mode 100644 index 00000000..310105ad --- /dev/null +++ b/integration/agentcompat/internal/scenario/reconnect_real_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/reconnect_receipts.go b/integration/agentcompat/internal/scenario/reconnect_receipts.go new file mode 100644 index 00000000..ed1937d6 --- /dev/null +++ b/integration/agentcompat/internal/scenario/reconnect_receipts.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/reconnect_run.go b/integration/agentcompat/internal/scenario/reconnect_run.go new file mode 100644 index 00000000..19564084 --- /dev/null +++ b/integration/agentcompat/internal/scenario/reconnect_run.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/reconnect_test.go b/integration/agentcompat/internal/scenario/reconnect_test.go new file mode 100644 index 00000000..57450313 --- /dev/null +++ b/integration/agentcompat/internal/scenario/reconnect_test.go @@ -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()) +} diff --git a/integration/agentcompat/internal/scenario/registration_config_exec.go b/integration/agentcompat/internal/scenario/registration_config_exec.go new file mode 100644 index 00000000..caef89b5 --- /dev/null +++ b/integration/agentcompat/internal/scenario/registration_config_exec.go @@ -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()) +} diff --git a/integration/agentcompat/internal/scenario/registration_config_exec_test.go b/integration/agentcompat/internal/scenario/registration_config_exec_test.go new file mode 100644 index 00000000..7e40a81e --- /dev/null +++ b/integration/agentcompat/internal/scenario/registration_config_exec_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/terminal.go b/integration/agentcompat/internal/scenario/terminal.go new file mode 100644 index 00000000..8545c48c --- /dev/null +++ b/integration/agentcompat/internal/scenario/terminal.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/terminal_observation.go b/integration/agentcompat/internal/scenario/terminal_observation.go new file mode 100644 index 00000000..91f2568d --- /dev/null +++ b/integration/agentcompat/internal/scenario/terminal_observation.go @@ -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)) +} diff --git a/integration/agentcompat/internal/scenario/terminal_support.go b/integration/agentcompat/internal/scenario/terminal_support.go new file mode 100644 index 00000000..0148d5e2 --- /dev/null +++ b/integration/agentcompat/internal/scenario/terminal_support.go @@ -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) +} diff --git a/integration/agentcompat/internal/scenario/terminal_test.go b/integration/agentcompat/internal/scenario/terminal_test.go new file mode 100644 index 00000000..d492a1b2 --- /dev/null +++ b/integration/agentcompat/internal/scenario/terminal_test.go @@ -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:])) +} diff --git a/integration/agentcompat/internal/scenario/transfer.go b/integration/agentcompat/internal/scenario/transfer.go new file mode 100644 index 00000000..bc787762 --- /dev/null +++ b/integration/agentcompat/internal/scenario/transfer.go @@ -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") +} diff --git a/integration/agentcompat/internal/scenario/transfer_evidence.go b/integration/agentcompat/internal/scenario/transfer_evidence.go new file mode 100644 index 00000000..cd6da3b3 --- /dev/null +++ b/integration/agentcompat/internal/scenario/transfer_evidence.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/transfer_execution.go b/integration/agentcompat/internal/scenario/transfer_execution.go new file mode 100644 index 00000000..5d29a2c8 --- /dev/null +++ b/integration/agentcompat/internal/scenario/transfer_execution.go @@ -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() +} diff --git a/integration/agentcompat/internal/scenario/transfer_observation.go b/integration/agentcompat/internal/scenario/transfer_observation.go new file mode 100644 index 00000000..6353e4bd --- /dev/null +++ b/integration/agentcompat/internal/scenario/transfer_observation.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/transfer_operations.go b/integration/agentcompat/internal/scenario/transfer_operations.go new file mode 100644 index 00000000..6772fc88 --- /dev/null +++ b/integration/agentcompat/internal/scenario/transfer_operations.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/transfer_real_test.go b/integration/agentcompat/internal/scenario/transfer_real_test.go new file mode 100644 index 00000000..dfbd6586 --- /dev/null +++ b/integration/agentcompat/internal/scenario/transfer_real_test.go @@ -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) + } +} diff --git a/integration/agentcompat/internal/scenario/transfer_sentinels.go b/integration/agentcompat/internal/scenario/transfer_sentinels.go new file mode 100644 index 00000000..89832fdb --- /dev/null +++ b/integration/agentcompat/internal/scenario/transfer_sentinels.go @@ -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 +} diff --git a/integration/agentcompat/internal/scenario/transfer_test.go b/integration/agentcompat/internal/scenario/transfer_test.go new file mode 100644 index 00000000..b8ebcfc8 --- /dev/null +++ b/integration/agentcompat/internal/scenario/transfer_test.go @@ -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, + } +} diff --git a/integration/agentcompat/internal/scenario/transfer_warmup.go b/integration/agentcompat/internal/scenario/transfer_warmup.go new file mode 100644 index 00000000..b2547ec7 --- /dev/null +++ b/integration/agentcompat/internal/scenario/transfer_warmup.go @@ -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 +}