From 26e92da33b91cac9cfcc36dfe774eff8cbff1c00 Mon Sep 17 00:00:00 2001 From: naiba Date: Mon, 20 Jul 2026 04:27:33 +0000 Subject: [PATCH] feat(agentcompat): add scoped IO stream capabilities Co-authored-by: naiba/CloudCode --- .../io_stream_capability_agentcompat_test.go | 229 ++++++++++++++++ .../io_stream_capability_bind_agentcompat.go | 74 ++++++ ...eam_capability_binding_agentcompat_test.go | 193 ++++++++++++++ ...io_stream_capability_cancel_agentcompat.go | 99 +++++++ service/rpc/io_stream_capability_default.go | 45 ++++ .../rpc/io_stream_capability_default_test.go | 64 +++++ .../io_stream_capability_nat_agentcompat.go | 84 ++++++ ..._stream_capability_nat_agentcompat_test.go | 169 ++++++++++++ ...ility_nat_atomic_start_agentcompat_test.go | 180 +++++++++++++ ...capability_nat_barrier_agentcompat_test.go | 117 ++++++++ ...ity_nat_handle_binding_agentcompat_test.go | 75 ++++++ ...ty_nat_lease_authority_agentcompat_test.go | 202 ++++++++++++++ ...capability_nat_profile_agentcompat_test.go | 99 +++++++ ...ream_capability_publication_agentcompat.go | 66 +++++ ...o_stream_capability_publication_default.go | 9 + ...ility_quota_accounting_agentcompat_test.go | 230 ++++++++++++++++ ...ility_quota_boundaries_agentcompat_test.go | 94 +++++++ ...ity_quota_test_helpers_agentcompat_test.go | 40 +++ ..._stream_capability_register_agentcompat.go | 96 +++++++ ...am_capability_security_agentcompat_test.go | 234 ++++++++++++++++ .../io_stream_capability_state_agentcompat.go | 78 ++++++ .../rpc/io_stream_capability_state_default.go | 9 + ...apability_test_helpers_agentcompat_test.go | 38 +++ service/rpc/io_stream_capability_types.go | 93 +++++++ service/rpc/io_stream_quota_agentcompat.go | 144 ++++++++++ .../rpc/io_stream_quota_agentcompat_test.go | 77 ++++++ service/rpc/io_stream_state_agentcompat.go | 11 + service/rpc/mcp_receipt_agentcompat_test.go | 197 ++++++++++++++ service/rpc/receipt_gate_agentcompat.go | 251 ++++++++++++++++++ service/rpc/receipt_gate_agentcompat_test.go | 224 ++++++++++++++++ service/rpc/receipt_gate_default.go | 19 ++ 31 files changed, 3540 insertions(+) create mode 100644 service/rpc/io_stream_capability_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_bind_agentcompat.go create mode 100644 service/rpc/io_stream_capability_binding_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_cancel_agentcompat.go create mode 100644 service/rpc/io_stream_capability_default.go create mode 100644 service/rpc/io_stream_capability_default_test.go create mode 100644 service/rpc/io_stream_capability_nat_agentcompat.go create mode 100644 service/rpc/io_stream_capability_nat_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_nat_atomic_start_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_nat_barrier_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_nat_handle_binding_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_nat_lease_authority_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_nat_profile_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_publication_agentcompat.go create mode 100644 service/rpc/io_stream_capability_publication_default.go create mode 100644 service/rpc/io_stream_capability_quota_accounting_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_quota_boundaries_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_quota_test_helpers_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_register_agentcompat.go create mode 100644 service/rpc/io_stream_capability_security_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_state_agentcompat.go create mode 100644 service/rpc/io_stream_capability_state_default.go create mode 100644 service/rpc/io_stream_capability_test_helpers_agentcompat_test.go create mode 100644 service/rpc/io_stream_capability_types.go create mode 100644 service/rpc/io_stream_quota_agentcompat.go create mode 100644 service/rpc/io_stream_quota_agentcompat_test.go create mode 100644 service/rpc/io_stream_state_agentcompat.go create mode 100644 service/rpc/mcp_receipt_agentcompat_test.go create mode 100644 service/rpc/receipt_gate_agentcompat.go create mode 100644 service/rpc/receipt_gate_agentcompat_test.go create mode 100644 service/rpc/receipt_gate_default.go diff --git a/service/rpc/io_stream_capability_agentcompat_test.go b/service/rpc/io_stream_capability_agentcompat_test.go new file mode 100644 index 00000000..45d01fb7 --- /dev/null +++ b/service/rpc/io_stream_capability_agentcompat_test.go @@ -0,0 +1,229 @@ +//go:build agentcompat + +package rpc + +import ( + "context" + "encoding/base64" + "errors" + "io" + "sync" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/require" +) + +func capabilityOwner(patID, userID uint64) AgentCompatCapabilityOwner { + return AgentCompatCapabilityOwner{PATID: patID, UserID: userID, IsAdmin: false} +} + +func capabilityRegistration(owner AgentCompatCapabilityOwner, purpose AgentCompatCapabilityPurpose, serverID, resourceID uint64) AgentCompatCapabilityRegistration { + return AgentCompatCapabilityRegistration{ + Owner: owner, Purpose: purpose, TargetServerID: serverID, ResourceID: resourceID, ServerAccessAllowed: true, + } +} + +func capabilityAccess(capability AgentCompatIOStreamCapability, registration AgentCompatCapabilityRegistration) AgentCompatCapabilityAccess { + return AgentCompatCapabilityAccess{ + Capability: capability, Owner: registration.Owner, Purpose: registration.Purpose, + TargetServerID: registration.TargetServerID, ResourceID: registration.ResourceID, ServerAccessAllowed: true, + } +} + +func TestAgentCompatCapabilityMintUsesURLSafe256BitTokens(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(1, 2), AgentCompatCapabilityTerminal, 3, 0) + + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + + require.NoError(t, err) + raw, err := base64.RawURLEncoding.DecodeString(capability.String()) + require.NoError(t, err) + require.Len(t, raw, 32) + parsed, err := ParseAgentCompatIOStreamCapability(capability.String()) + require.NoError(t, err) + require.Equal(t, capability, parsed) + _, err = ParseAgentCompatIOStreamCapability("not-a-capability") + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) +} + +func TestAgentCompatCapabilityMintRetriesActiveAndUsedCollisions(t *testing.T) { + handler := NewNezhaHandler() + first := make([]byte, 32) + second := make([]byte, 32) + third := make([]byte, 32) + first[0], second[0], third[0] = 1, 2, 3 + var calls atomic.Int32 + handler.setAgentCompatCapabilityTokenSourceForTest(func(destination []byte) error { + switch calls.Add(1) { + case 1, 2, 4: + copy(destination, first) + return nil + case 3: + copy(destination, second) + return nil + default: + copy(destination, third) + return nil + } + }) + registration := capabilityRegistration(capabilityOwner(1, 2), AgentCompatCapabilityTerminal, 3, 0) + firstCapability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + activeCollisionCapability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + require.NotEqual(t, firstCapability, activeCollisionCapability) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(firstCapability, registration))) + + tombstoneCollisionCapability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + + require.NoError(t, err) + require.NotEqual(t, firstCapability, tombstoneCollisionCapability) + require.Equal(t, int32(5), calls.Load()) +} + +func TestAgentCompatCapabilityRegistrationRequiresServerAccessProof(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(1, 2), AgentCompatCapabilityTerminal, 3, 0) + registration.ServerAccessAllowed = false + + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) +} + +func TestAgentCompatCapabilityWaitRequiresExactOwnerAndRetainsBinding(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(10, 20), AgentCompatCapabilityTerminal, 30, 0) + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + require.NoError(t, handler.CreateStreamWithPurpose("terminal-bound", 20, 30, PurposeTerminal)) + require.NoError(t, handler.BindAgentCompatIOStreamCapability(AgentCompatCapabilityBinding{ + AgentCompatCapabilityAccess: capabilityAccess(capability, registration), StreamID: "terminal-bound", + })) + + streamID, err := handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccess(capability, registration)) + require.NoError(t, err) + require.Equal(t, "terminal-bound", streamID) + foreign := capabilityAccess(capability, registration) + foreign.Owner.PATID++ + _, err = handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), foreign) + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) +} + +type capabilityCloseEndpoint struct { + handler *NezhaHandler + streamID string + err error + closed atomic.Int32 +} + +func (endpoint *capabilityCloseEndpoint) Read([]byte) (int, error) { return 0, io.EOF } +func (endpoint *capabilityCloseEndpoint) Write(data []byte) (int, error) { return len(data), nil } +func (endpoint *capabilityCloseEndpoint) Close() error { + endpoint.closed.Add(1) + endpoint.handler.StreamOwnership(endpoint.streamID) + return endpoint.err +} + +func TestAgentCompatCapabilityCancelClosesOutsideLockAndJoinsErrors(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityFileManager, "fm-close", 41) + firstErr := errors.New("user close") + secondErr := errors.New("agent close") + first := &capabilityCloseEndpoint{handler: handler, streamID: "fm-close", err: firstErr} + second := &capabilityCloseEndpoint{handler: handler, streamID: "fm-close", err: secondErr} + require.NoError(t, handler.UserConnected("fm-close", first)) + require.NoError(t, handler.AgentConnected("fm-close", second)) + + err := handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration)) + + require.ErrorIs(t, err, firstErr) + require.ErrorIs(t, err, secondErr) + require.Equal(t, int32(1), first.closed.Load()) + require.Equal(t, int32(1), second.closed.Load()) +} + +func boundCapabilityFixture(t *testing.T, purpose AgentCompatCapabilityPurpose, streamID string, serverID uint64) (*NezhaHandler, AgentCompatCapabilityRegistration, AgentCompatIOStreamCapability) { + t.Helper() + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(11, 21), purpose, serverID, 0) + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + require.NoError(t, handler.CreateStreamWithPurpose(streamID, 21, serverID, purpose.streamPurpose())) + require.NoError(t, handler.BindAgentCompatIOStreamCapability(AgentCompatCapabilityBinding{ + AgentCompatCapabilityAccess: capabilityAccess(capability, registration), StreamID: streamID, + })) + return handler, registration, capability +} + +func TestAgentCompatCapabilityCancelRacingCloseChangesGenerationOnce(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityTerminal, "race-close", 51) + endpoint := &capabilityCloseEndpoint{handler: handler, streamID: "race-close"} + require.NoError(t, handler.AgentConnected("race-close", endpoint)) + start := handler.SnapshotIOStreamState() + ready := make(chan struct{}) + raceCtx := agentCompatCapabilityTestContext(t) + var waitGroup sync.WaitGroup + waitGroup.Add(2) + go func() { + defer waitGroup.Done() + select { + case <-ready: + _ = handler.CloseStream("race-close") + case <-raceCtx.Done(): + } + }() + go func() { + defer waitGroup.Done() + select { + case <-ready: + _ = handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration)) + case <-raceCtx.Done(): + } + }() + close(ready) + raceDone := make(chan struct{}) + go func() { + waitGroup.Wait() + close(raceDone) + }() + awaitAgentCompatCapabilitySignal(t, raceDone, "cancel/close race did not complete") + require.NoError(t, raceCtx.Err()) + + state := handler.SnapshotIOStreamState() + require.Equal(t, start.Generation+1, state.Generation) + require.Equal(t, int32(1), endpoint.closed.Load()) +} + +func TestAgentCompatCapabilityWaitTimeoutKeepsRegistration(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(1, 2), AgentCompatCapabilityTerminal, 3, 0) + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err = handler.WaitAgentCompatIOStreamCapability(ctx, capabilityAccess(capability, registration)) + require.ErrorIs(t, err, context.Canceled) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) +} + +func TestAgentCompatCapabilityWaitWakesAfterUnregister(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(1, 2), AgentCompatCapabilityTerminal, 3, 0) + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + result := make(chan error, 1) + started := make(chan struct{}) + handler.setAgentCompatCapabilityWaitObserverForTest(func() { close(started) }) + waitCtx := agentCompatCapabilityTestContext(t) + go func() { + _, waitErr := handler.WaitAgentCompatIOStreamCapability(waitCtx, capabilityAccess(capability, registration)) + result <- waitErr + }() + awaitAgentCompatCapabilitySignal(t, started, "wait observer did not start") + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + + require.ErrorIs(t, receiveAgentCompatCapabilityError(t, result, "unregister did not wake waiter"), ErrAgentCompatCapabilityHidden) +} diff --git a/service/rpc/io_stream_capability_bind_agentcompat.go b/service/rpc/io_stream_capability_bind_agentcompat.go new file mode 100644 index 00000000..1620a432 --- /dev/null +++ b/service/rpc/io_stream_capability_bind_agentcompat.go @@ -0,0 +1,74 @@ +//go:build agentcompat + +package rpc + +import "context" + +func (s *NezhaHandler) BindAgentCompatIOStreamCapability(binding AgentCompatCapabilityBinding) error { + s.ioStreamMutex.Lock() + defer s.ioStreamMutex.Unlock() + registration, allowed := s.agentCompatRegistrationLocked(binding.AgentCompatCapabilityAccess) + if !allowed || registration.phase != agentCompatCapabilityRegistered || registration.registration.Purpose == AgentCompatCapabilityNAT { + return ErrAgentCompatCapabilityHidden + } + stream, exists := s.ioStreams[binding.StreamID] + stored := registration.registration + if !exists || binding.StreamID == "" || stream.creatorUserID != stored.Owner.UserID || + stream.targetServerID != stored.TargetServerID || stream.purpose != stored.Purpose.streamPurpose() { + return ErrAgentCompatCapabilityHidden + } + if registration.stream != nil { + if registration.stream == stream && registration.streamID == binding.StreamID { + return nil + } + return ErrAgentCompatCapabilityConflict + } + registration.streamID = binding.StreamID + registration.stream = stream + registration.publishLocked() + return nil +} + +func (s *NezhaHandler) WaitAgentCompatIOStreamCapability(ctx context.Context, access AgentCompatCapabilityAccess) (string, error) { + for { + s.ioStreamMutex.RLock() + registration, allowed := s.agentCompatRegistrationLocked(access) + if !allowed { + s.ioStreamMutex.RUnlock() + return "", ErrAgentCompatCapabilityHidden + } + if registration.streamID != "" { + streamID := registration.streamID + stream := registration.stream + stored := registration.registration + current, live := s.ioStreams[streamID] + if stored.Purpose == AgentCompatCapabilityNAT && registration.phase == agentCompatCapabilityPublished && stream != nil { + s.ioStreamMutex.RUnlock() + return streamID, nil + } + // A reused StreamID must not turn a retained capability into authority over a replacement stream. + creatorMatches := stream != nil && stream.creatorUserID == stored.Owner.UserID + if stored.Purpose == AgentCompatCapabilityNAT { + creatorMatches = stream != nil && stream.creatorUserID == 0 + } + valid := live && current == stream && creatorMatches && + stream.targetServerID == stored.TargetServerID && stream.purpose == stored.Purpose.streamPurpose() + s.ioStreamMutex.RUnlock() + if !valid { + return "", ErrAgentCompatCapabilityHidden + } + return streamID, nil + } + notify := registration.notify + observer := s.agentCompatCapabilities.waitObserver + s.ioStreamMutex.RUnlock() + if observer != nil { + observer() + } + select { + case <-ctx.Done(): + return "", ctx.Err() + case <-notify: + } + } +} diff --git a/service/rpc/io_stream_capability_binding_agentcompat_test.go b/service/rpc/io_stream_capability_binding_agentcompat_test.go new file mode 100644 index 00000000..2fd1c81a --- /dev/null +++ b/service/rpc/io_stream_capability_binding_agentcompat_test.go @@ -0,0 +1,193 @@ +//go:build agentcompat + +package rpc + +import ( + "errors" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestAgentCompatCapabilityCancelLostCreateResponseDeletesOnlyExactStream(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityTerminal, "lost-response", 61) + require.NoError(t, handler.CreateStreamWithPurpose("other-stream", 21, 61, PurposeTerminal)) + start := handler.SnapshotIOStreamState() + + err := handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration)) + + require.NoError(t, err) + state := handler.SnapshotIOStreamState() + require.Equal(t, start.Generation+1, state.Generation) + require.Equal(t, 1, state.Count) + _, found := handler.StreamOwnership("other-stream") + require.True(t, found) +} + +func TestAgentCompatCapabilityCancelOneOfConcurrentCapabilitiesKeepsOthers(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(12, 22), AgentCompatCapabilityTerminal, 62, 0) + first, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + second, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + for streamID, capability := range map[string]AgentCompatIOStreamCapability{"first": first, "second": second} { + require.NoError(t, handler.CreateStreamWithPurpose(streamID, 22, 62, PurposeTerminal)) + require.NoError(t, handler.BindAgentCompatIOStreamCapability(AgentCompatCapabilityBinding{ + AgentCompatCapabilityAccess: capabilityAccess(capability, registration), StreamID: streamID, + })) + } + + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(capabilityAccess(first, registration))) + + streamID, err := handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccess(second, registration)) + require.NoError(t, err) + require.Equal(t, "second", streamID) +} + +func TestAgentCompatCapabilityBindValidatesStoredIdentityAndStream(t *testing.T) { + tests := []struct { + name string + mutateAccess func(*AgentCompatCapabilityAccess) + streamOwner uint64 + streamServer uint64 + streamPurpose StreamPurpose + }{ + {name: "foreign PAT", mutateAccess: func(access *AgentCompatCapabilityAccess) { access.Owner.PATID++ }, streamOwner: 23, streamServer: 63, streamPurpose: PurposeTerminal}, + {name: "user mismatch", mutateAccess: func(access *AgentCompatCapabilityAccess) { access.Owner.UserID++ }, streamOwner: 23, streamServer: 63, streamPurpose: PurposeTerminal}, + {name: "admin mismatch", mutateAccess: func(access *AgentCompatCapabilityAccess) { access.Owner.IsAdmin = true }, streamOwner: 23, streamServer: 63, streamPurpose: PurposeTerminal}, + {name: "purpose mismatch", mutateAccess: func(access *AgentCompatCapabilityAccess) { access.Purpose = AgentCompatCapabilityFileManager }, streamOwner: 23, streamServer: 63, streamPurpose: PurposeTerminal}, + {name: "target mismatch", mutateAccess: func(access *AgentCompatCapabilityAccess) { access.TargetServerID++ }, streamOwner: 23, streamServer: 63, streamPurpose: PurposeTerminal}, + {name: "resource mismatch", mutateAccess: func(access *AgentCompatCapabilityAccess) { access.ResourceID++ }, streamOwner: 23, streamServer: 63, streamPurpose: PurposeTerminal}, + {name: "access denied", mutateAccess: func(access *AgentCompatCapabilityAccess) { access.ServerAccessAllowed = false }, streamOwner: 23, streamServer: 63, streamPurpose: PurposeTerminal}, + {name: "stream creator mismatch", mutateAccess: func(*AgentCompatCapabilityAccess) {}, streamOwner: 24, streamServer: 63, streamPurpose: PurposeTerminal}, + {name: "stream server mismatch", mutateAccess: func(*AgentCompatCapabilityAccess) {}, streamOwner: 23, streamServer: 64, streamPurpose: PurposeTerminal}, + {name: "stream purpose mismatch", mutateAccess: func(*AgentCompatCapabilityAccess) {}, streamOwner: 23, streamServer: 63, streamPurpose: PurposeFileManager}, + } + for _, testCase := range tests { + t.Run(testCase.name, func(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(13, 23), AgentCompatCapabilityTerminal, 63, 0) + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + require.NoError(t, handler.CreateStreamWithPurpose("candidate", testCase.streamOwner, testCase.streamServer, testCase.streamPurpose)) + access := capabilityAccess(capability, registration) + testCase.mutateAccess(&access) + + err = handler.BindAgentCompatIOStreamCapability(AgentCompatCapabilityBinding{AgentCompatCapabilityAccess: access, StreamID: "candidate"}) + + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) + }) + } +} + +func TestAgentCompatCapabilityBindIsIdempotentButRejectsConflict(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityTerminal, "original", 64) + binding := AgentCompatCapabilityBinding{AgentCompatCapabilityAccess: capabilityAccess(capability, registration), StreamID: "original"} + require.NoError(t, handler.BindAgentCompatIOStreamCapability(binding)) + require.NoError(t, handler.CreateStreamWithPurpose("conflict", 21, 64, PurposeTerminal)) + binding.StreamID = "conflict" + + err := handler.BindAgentCompatIOStreamCapability(binding) + + require.ErrorIs(t, err, ErrAgentCompatCapabilityConflict) + streamID, err := handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccess(capability, registration)) + require.NoError(t, err) + require.Equal(t, "original", streamID) +} + +func TestAgentCompatCapabilityCancelMismatchOrReplacementDoesNotDetach(t *testing.T) { + t.Run("target mismatch", func(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityTerminal, "target-mismatch", 65) + start := handler.SnapshotIOStreamState() + access := capabilityAccess(capability, registration) + access.TargetServerID++ + + err := handler.CancelAgentCompatIOStreamCapability(access) + + require.NoError(t, err) + require.Equal(t, start, handler.SnapshotIOStreamState()) + }) + t.Run("entry replacement", func(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityTerminal, "replaced", 66) + require.NoError(t, handler.CloseStream("replaced")) + require.NoError(t, handler.CreateStreamWithPurpose("replaced", 21, 66, PurposeTerminal)) + start := handler.SnapshotIOStreamState() + + err := handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration)) + + require.NoError(t, err) + require.Equal(t, start, handler.SnapshotIOStreamState()) + }) +} + +func TestAgentCompatCapabilityCancelIsIdentityHidingIdempotentForAbsentAndUnbound(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(14, 24), AgentCompatCapabilityTerminal, 67, 0) + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + start := handler.SnapshotIOStreamState() + + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(AgentCompatCapabilityAccess{})) + require.Equal(t, start, handler.SnapshotIOStreamState()) +} + +func TestAgentCompatCapabilityCancelAfterNormalCloseIsIdempotent(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityTerminal, "normally-closed", 69) + require.NoError(t, handler.CloseStream("normally-closed")) + start := handler.SnapshotIOStreamState() + + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + + require.Equal(t, start, handler.SnapshotIOStreamState()) +} + +func TestAgentCompatCapabilityForeignCancelDoesNotMutate(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityTerminal, "foreign-cancel", 70) + start := handler.SnapshotIOStreamState() + access := capabilityAccess(capability, registration) + access.Owner.PATID++ + + err := handler.CancelAgentCompatIOStreamCapability(access) + + require.NoError(t, err) + require.Equal(t, start, handler.SnapshotIOStreamState()) + _, found := handler.StreamOwnership("foreign-cancel") + require.True(t, found) +} + +func TestAgentCompatCapabilityUnregisterRejectsBoundLiveStream(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityTerminal, "bound-unregister", 68) + start := handler.SnapshotIOStreamState() + + err := handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(capability, registration)) + + require.ErrorIs(t, err, ErrAgentCompatCapabilityBound) + require.Equal(t, start, handler.SnapshotIOStreamState()) +} + +func TestAgentCompatCapabilityUnregisterRequiresSamePATAndIsIdempotent(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(15, 25), AgentCompatCapabilityTerminal, 71, 0) + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + foreign := capabilityAccess(capability, registration) + foreign.Owner.PATID++ + + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(foreign)) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) +} + +func TestAgentCompatCapabilityTokenSourceErrorIsVisible(t *testing.T) { + handler := NewNezhaHandler() + sourceErr := errors.New("token source failed") + handler.setAgentCompatCapabilityTokenSourceForTest(func([]byte) error { return sourceErr }) + + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityRegistration(capabilityOwner(1, 2), AgentCompatCapabilityTerminal, 3, 0)) + + require.ErrorIs(t, err, sourceErr) +} diff --git a/service/rpc/io_stream_capability_cancel_agentcompat.go b/service/rpc/io_stream_capability_cancel_agentcompat.go new file mode 100644 index 00000000..2b28cc69 --- /dev/null +++ b/service/rpc/io_stream_capability_cancel_agentcompat.go @@ -0,0 +1,99 @@ +//go:build agentcompat + +package rpc + +import ( + "errors" + "io" +) + +func (s *NezhaHandler) CancelAgentCompatIOStreamCapability(access AgentCompatCapabilityAccess) error { + s.ioStreamMutex.Lock() + registration, exists := s.agentCompatCapabilities.active[access.Capability.value] + if !exists { + s.ioStreamMutex.Unlock() + return nil + } + if !agentCompatAccessMatches(access, registration) { + s.ioStreamMutex.Unlock() + // Foreign and absent capabilities intentionally share the same inert result to prevent enumeration. + return nil + } + if registration.stream == nil || registration.streamID == "" { + s.removeAgentCompatCapabilityLocked(access.Capability.value, registration) + s.ioStreamMutex.Unlock() + return nil + } + stream := registration.stream + stored := registration.registration + current, live := s.ioStreams[registration.streamID] + if !live { + s.removeAgentCompatCapabilityLocked(access.Capability.value, registration) + s.ioStreamMutex.Unlock() + return nil + } + creatorMatches := stream.creatorUserID == stored.Owner.UserID + if stored.Purpose == AgentCompatCapabilityNAT { + creatorMatches = stream.creatorUserID == 0 + } + if !access.ServerAccessAllowed || current != stream || !creatorMatches || + stream.targetServerID != stored.TargetServerID || stream.purpose != stored.Purpose.streamPurpose() { + s.ioStreamMutex.Unlock() + return nil + } + stream.revoke() + endpoints := make([]io.ReadWriteCloser, 0, 2) + if stream.userIo != nil { + endpoints = append(endpoints, stream.userIo) + } + if stream.agentIo != nil { + endpoints = append(endpoints, stream.agentIo) + } + delete(s.ioStreams, registration.streamID) + s.publishIOStreamStateChangeLocked() + s.removeAgentCompatCapabilityLocked(access.Capability.value, registration) + s.ioStreamMutex.Unlock() + + closeErrors := make([]error, 0, len(endpoints)) + for _, endpoint := range endpoints { + if err := endpoint.Close(); err != nil { + closeErrors = append(closeErrors, err) + } + } + return errors.Join(closeErrors...) +} + +func (s *NezhaHandler) UnregisterAgentCompatIOStreamCapability(access AgentCompatCapabilityAccess) error { + s.ioStreamMutex.Lock() + defer s.ioStreamMutex.Unlock() + registration, exists := s.agentCompatCapabilities.active[access.Capability.value] + if !exists { + return nil + } + if registration.registration.Owner.PATID != access.Owner.PATID || !agentCompatAccessMatches(access, registration) { + return nil + } + if registration.stream != nil { + if current, live := s.ioStreams[registration.streamID]; live && current == registration.stream { + return ErrAgentCompatCapabilityBound + } + } + s.removeAgentCompatCapabilityLocked(access.Capability.value, registration) + return nil +} + +func (s *NezhaHandler) removeAgentCompatCapabilityLocked(capability string, registration *agentCompatCapabilityRegistration) { + current, active := s.agentCompatCapabilities.active[capability] + if !active || current != registration { + return + } + delete(s.agentCompatCapabilities.active, capability) + patID := registration.registration.Owner.PATID + remaining := s.agentCompatCapabilities.activeByPAT[patID] - 1 + if remaining == 0 { + delete(s.agentCompatCapabilities.activeByPAT, patID) + } else { + s.agentCompatCapabilities.activeByPAT[patID] = remaining + } + registration.publishLocked() +} diff --git a/service/rpc/io_stream_capability_default.go b/service/rpc/io_stream_capability_default.go new file mode 100644 index 00000000..cb3ca9b9 --- /dev/null +++ b/service/rpc/io_stream_capability_default.go @@ -0,0 +1,45 @@ +//go:build !agentcompat + +package rpc + +import "context" + +func (*NezhaHandler) RegisterAgentCompatIOStreamCapability(context.Context, AgentCompatCapabilityRegistration) (AgentCompatIOStreamCapability, error) { + return AgentCompatIOStreamCapability{}, ErrAgentCompatCapabilityUnavailable +} + +func (*NezhaHandler) BindAgentCompatIOStreamCapability(AgentCompatCapabilityBinding) error { + return ErrAgentCompatCapabilityUnavailable +} + +func (*NezhaHandler) ConsumeAgentCompatNATCapability(AgentCompatCapabilityAccess) (AgentCompatNATPublishHandle, error) { + return AgentCompatNATPublishHandle{}, ErrAgentCompatCapabilityUnavailable +} + +func (*NezhaHandler) ConsumeAgentCompatNATCapabilityForProfile(string, uint64, uint64) (AgentCompatCapabilityAccess, AgentCompatNATPublishHandle, error) { + return AgentCompatCapabilityAccess{}, AgentCompatNATPublishHandle{}, ErrAgentCompatCapabilityUnavailable +} + +func (*NezhaHandler) PublishAgentCompatNATStream(AgentCompatNATPublishHandle, AgentCompatNATPublication) error { + return ErrAgentCompatCapabilityUnavailable +} + +func (*NezhaHandler) WaitAgentCompatIOStreamCapability(context.Context, AgentCompatCapabilityAccess) (string, error) { + return "", ErrAgentCompatCapabilityUnavailable +} + +func (*NezhaHandler) CancelAgentCompatIOStreamCapability(AgentCompatCapabilityAccess) error { + return nil +} + +func (*NezhaHandler) UnregisterAgentCompatIOStreamCapability(AgentCompatCapabilityAccess) error { + return nil +} + +func (*NezhaHandler) CreateAgentCompatNATStream(AgentCompatNATPublishHandle, string) (*AgentCompatNATStreamLease, error) { + return nil, ErrAgentCompatCapabilityUnavailable +} + +func (*NezhaHandler) CloseAgentCompatNATStreamLease(*AgentCompatNATStreamLease) error { + return nil +} diff --git a/service/rpc/io_stream_capability_default_test.go b/service/rpc/io_stream_capability_default_test.go new file mode 100644 index 00000000..6a2644bc --- /dev/null +++ b/service/rpc/io_stream_capability_default_test.go @@ -0,0 +1,64 @@ +//go:build !agentcompat + +package rpc + +import ( + "context" + "errors" + "reflect" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestAgentCompatCapabilityDefaultBuildHasNoRegistryState(t *testing.T) { + handler := NewNezhaHandler() + state := reflect.ValueOf(handler.agentCompatCapabilities) + + require.Equal(t, 0, state.NumField()) +} + +func TestAgentCompatCapabilityDefaultBuildUsesStableUnavailableAndNoopContracts(t *testing.T) { + handler := NewNezhaHandler() + registration := AgentCompatCapabilityRegistration{} + capability, err := handler.RegisterAgentCompatIOStreamCapability(context.Background(), registration) + require.Empty(t, capability.String()) + require.ErrorIs(t, err, ErrAgentCompatCapabilityUnavailable) + access := AgentCompatCapabilityAccess{} + require.ErrorIs(t, handler.BindAgentCompatIOStreamCapability(AgentCompatCapabilityBinding{}), ErrAgentCompatCapabilityUnavailable) + _, err = handler.ConsumeAgentCompatNATCapability(access) + require.ErrorIs(t, err, ErrAgentCompatCapabilityUnavailable) + require.ErrorIs(t, handler.PublishAgentCompatNATStream(AgentCompatNATPublishHandle{}, AgentCompatNATPublication{}), ErrAgentCompatCapabilityUnavailable) + _, err = handler.WaitAgentCompatIOStreamCapability(context.Background(), access) + require.True(t, errors.Is(err, ErrAgentCompatCapabilityUnavailable)) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(access)) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(access)) +} + +func TestAgentCompatNATCapabilityForProfileDefaultBuildIsUnavailableAndNoOp(t *testing.T) { + handler := NewNezhaHandler() + access, handle, err := handler.ConsumeAgentCompatNATCapabilityForProfile("not-a-capability", 1, 2) + + require.Equal(t, AgentCompatCapabilityAccess{}, access) + require.Equal(t, AgentCompatNATPublishHandle{}, handle) + require.ErrorIs(t, err, ErrAgentCompatCapabilityUnavailable) +} + +func TestAgentCompatNATAtomicStartDefaultBuildIsUnavailableAndStateless(t *testing.T) { + handler := NewNezhaHandler() + publicationOwned, err := handler.StartAgentCompatNATStream(AgentCompatNATPublishHandle{}, 0) + + require.False(t, publicationOwned) + require.ErrorIs(t, err, ErrAgentCompatCapabilityUnavailable) + require.Equal(t, 0, handler.StreamCount()) +} + +func TestAgentCompatNATLeaseDefaultBuildIsUnavailableAndStateless(t *testing.T) { + handler := NewNezhaHandler() + lease, err := handler.CreateAgentCompatNATStream(AgentCompatNATPublishHandle{}, "known") + + require.Nil(t, lease) + require.ErrorIs(t, err, ErrAgentCompatCapabilityUnavailable) + require.NoError(t, handler.CloseAgentCompatNATStreamLease(nil)) + require.Equal(t, 0, handler.StreamCount()) +} diff --git a/service/rpc/io_stream_capability_nat_agentcompat.go b/service/rpc/io_stream_capability_nat_agentcompat.go new file mode 100644 index 00000000..144b07dc --- /dev/null +++ b/service/rpc/io_stream_capability_nat_agentcompat.go @@ -0,0 +1,84 @@ +//go:build agentcompat + +package rpc + +func (s *NezhaHandler) ConsumeAgentCompatNATCapability(access AgentCompatCapabilityAccess) (AgentCompatNATPublishHandle, error) { + s.ioStreamMutex.Lock() + defer s.ioStreamMutex.Unlock() + registration, allowed := s.agentCompatRegistrationLocked(access) + if !allowed { + return AgentCompatNATPublishHandle{}, ErrAgentCompatCapabilityHidden + } + return s.consumeAgentCompatNATCapabilityLocked(registration, access.Capability.value) +} + +func (s *NezhaHandler) ConsumeAgentCompatNATCapabilityForProfile(value string, targetServerID, resourceID uint64) (AgentCompatCapabilityAccess, AgentCompatNATPublishHandle, error) { + capability, err := ParseAgentCompatIOStreamCapability(value) + if err != nil { + return AgentCompatCapabilityAccess{}, AgentCompatNATPublishHandle{}, ErrAgentCompatCapabilityHidden + } + + s.ioStreamMutex.Lock() + defer s.ioStreamMutex.Unlock() + registration, exists := s.agentCompatCapabilities.active[capability.value] + if !exists || registration.registration.Purpose != AgentCompatCapabilityNAT || + registration.registration.TargetServerID != targetServerID || registration.registration.ResourceID != resourceID { + return AgentCompatCapabilityAccess{}, AgentCompatNATPublishHandle{}, ErrAgentCompatCapabilityHidden + } + handle, err := s.consumeAgentCompatNATCapabilityLocked(registration, capability.value) + if err != nil { + return AgentCompatCapabilityAccess{}, AgentCompatNATPublishHandle{}, err + } + stored := registration.registration + return AgentCompatCapabilityAccess{ + Capability: capability, Owner: stored.Owner, Purpose: stored.Purpose, + TargetServerID: stored.TargetServerID, ResourceID: stored.ResourceID, + ServerAccessAllowed: stored.ServerAccessAllowed, + }, handle, nil +} + +func (s *NezhaHandler) consumeAgentCompatNATCapabilityLocked(registration *agentCompatCapabilityRegistration, capability string) (AgentCompatNATPublishHandle, error) { + if registration == nil || registration.registration.Purpose != AgentCompatCapabilityNAT || registration.phase != agentCompatCapabilityRegistered { + return AgentCompatNATPublishHandle{}, ErrAgentCompatCapabilityHidden + } + registration.phase = agentCompatCapabilityConsumed + return AgentCompatNATPublishHandle{ + registration: registration, generation: registration.generation, + capability: capability, + }, nil +} + +func (s *NezhaHandler) PublishAgentCompatNATStream(handle AgentCompatNATPublishHandle, publication AgentCompatNATPublication) error { + s.ioStreamMutex.RLock() + publishObserver := s.agentCompatCapabilities.publishObserver + s.ioStreamMutex.RUnlock() + if publishObserver != nil { + publishObserver() + } + s.ioStreamMutex.Lock() + defer s.ioStreamMutex.Unlock() + registration := handle.registration + // Pointer identity plus generation makes a late publisher inert after unregister/cancel. + if registration == nil || registration.generation != handle.generation { + return nil + } + current, active := s.agentCompatCapabilities.active[handle.capability] + if !active || current != registration { + return nil + } + if registration.phase == agentCompatCapabilityPublished { + return nil + } + stored := registration.registration + stream := registration.stream + exists := publication.StreamID != "" && registration.streamID == publication.StreamID && stream != nil && s.ioStreams[publication.StreamID] == stream + if registration.phase != agentCompatCapabilityConsumed || publication.Purpose != stored.Purpose || + publication.TargetServerID != stored.TargetServerID || publication.ResourceID != stored.ResourceID || + !exists || stream.creatorUserID != 0 || + stream.targetServerID != stored.TargetServerID || stream.purpose != PurposeNAT { + return ErrAgentCompatCapabilityHidden + } + registration.phase = agentCompatCapabilityPublished + registration.publishLocked() + return nil +} diff --git a/service/rpc/io_stream_capability_nat_agentcompat_test.go b/service/rpc/io_stream_capability_nat_agentcompat_test.go new file mode 100644 index 00000000..02fd4265 --- /dev/null +++ b/service/rpc/io_stream_capability_nat_agentcompat_test.go @@ -0,0 +1,169 @@ +//go:build agentcompat + +package rpc + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func natCapabilityFixture(t *testing.T, patID, userID, serverID, profileID uint64) (*NezhaHandler, AgentCompatCapabilityRegistration, AgentCompatIOStreamCapability) { + t.Helper() + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(patID, userID), AgentCompatCapabilityNAT, serverID, profileID) + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + return handler, registration, capability +} + +func TestAgentCompatNATCapabilityTransitionsAndRetainsFirstPublication(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 21, 31, 71, 81) + handle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(capability, registration)) + require.NoError(t, err) + _, err = handler.ConsumeAgentCompatNATCapability(capabilityAccess(capability, registration)) + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) + _, err = handler.CreateAgentCompatNATStream(handle, "nat-first") + require.NoError(t, err) + require.NoError(t, handler.CreateStreamWithPurpose("nat-second", 0, 71, PurposeNAT)) + publication := AgentCompatNATPublication{Purpose: AgentCompatCapabilityNAT, TargetServerID: 71, ResourceID: 81, StreamID: "nat-first"} + require.NoError(t, handler.PublishAgentCompatNATStream(handle, publication)) + publication.StreamID = "nat-second" + require.NoError(t, handler.PublishAgentCompatNATStream(handle, publication)) + + streamID, err := handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccess(capability, registration)) + require.NoError(t, err) + require.Equal(t, "nat-first", streamID) +} + +func TestAgentCompatNATCapabilityPublicationBeforeWaitWorks(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 22, 32, 72, 82) + handle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(capability, registration)) + require.NoError(t, err) + _, err = handler.CreateAgentCompatNATStream(handle, "nat-published") + require.NoError(t, err) + require.NoError(t, handler.PublishAgentCompatNATStream(handle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 72, ResourceID: 82, StreamID: "nat-published", + })) + + streamID, err := handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccess(capability, registration)) + require.NoError(t, err) + require.Equal(t, "nat-published", streamID) +} + +func TestAgentCompatNATCapabilityValidatesConsumeAndPublishIdentity(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 23, 33, 73, 83) + access := capabilityAccess(capability, registration) + access.ResourceID++ + _, err := handler.ConsumeAgentCompatNATCapability(access) + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) + handle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(capability, registration)) + require.NoError(t, err) + _, err = handler.CreateAgentCompatNATStream(handle, "nat-identity") + require.NoError(t, err) + + err = handler.PublishAgentCompatNATStream(handle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 73, ResourceID: 84, StreamID: "nat-identity", + }) + + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) +} + +func TestAgentCompatNATCapabilityLatePublishAfterUnregisterIsIgnored(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 24, 34, 74, 84) + handle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(capability, registration)) + require.NoError(t, err) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + require.NoError(t, handler.CreateStreamWithPurpose("nat-late", 0, 74, PurposeNAT)) + + err = handler.PublishAgentCompatNATStream(handle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 74, ResourceID: 84, StreamID: "nat-late", + }) + + require.NoError(t, err) + _, err = handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccess(capability, registration)) + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) +} + +func TestAgentCompatNATCapabilityLatePublishAfterCancelIsIgnored(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 27, 37, 77, 87) + handle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(capability, registration)) + require.NoError(t, err) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + require.NoError(t, handler.CreateStreamWithPurpose("nat-after-cancel", 0, 77, PurposeNAT)) + + err = handler.PublishAgentCompatNATStream(handle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 77, ResourceID: 87, StreamID: "nat-after-cancel", + }) + + require.NoError(t, err) + _, found := handler.StreamOwnership("nat-after-cancel") + require.True(t, found) +} + +func TestAgentCompatNATCapabilityReusedTokenCannotBindAnotherStream(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 28, 38, 78, 88) + handle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(capability, registration)) + require.NoError(t, err) + _, err = handler.CreateAgentCompatNATStream(handle, "nat-original") + require.NoError(t, err) + require.NoError(t, handler.PublishAgentCompatNATStream(handle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 78, ResourceID: 88, StreamID: "nat-original", + })) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + require.NoError(t, handler.CreateStreamWithPurpose("nat-reuse", 0, 78, PurposeNAT)) + + _, err = handler.ConsumeAgentCompatNATCapability(capabilityAccess(capability, registration)) + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) + require.NoError(t, handler.PublishAgentCompatNATStream(handle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 78, ResourceID: 88, StreamID: "nat-reuse", + })) + _, found := handler.StreamOwnership("nat-reuse") + require.True(t, found) +} + +func TestAgentCompatNATCapabilityCancelDetachesPublishedStream(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 25, 35, 75, 85) + handle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(capability, registration)) + require.NoError(t, err) + _, err = handler.CreateAgentCompatNATStream(handle, "nat-cancel") + require.NoError(t, err) + require.NoError(t, handler.PublishAgentCompatNATStream(handle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 75, ResourceID: 85, StreamID: "nat-cancel", + })) + start := handler.SnapshotIOStreamState() + + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + + state := handler.SnapshotIOStreamState() + require.Equal(t, start.Generation+1, state.Generation) + require.Equal(t, 0, state.Count) +} + +func TestAgentCompatNATCapabilitiesRemainSeparatedAcrossProfiles(t *testing.T) { + handler := NewNezhaHandler() + owner := capabilityOwner(26, 36) + firstRegistration := capabilityRegistration(owner, AgentCompatCapabilityNAT, 76, 86) + secondRegistration := capabilityRegistration(owner, AgentCompatCapabilityNAT, 76, 87) + first, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), firstRegistration) + require.NoError(t, err) + second, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), secondRegistration) + require.NoError(t, err) + firstHandle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(first, firstRegistration)) + require.NoError(t, err) + secondHandle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(second, secondRegistration)) + require.NoError(t, err) + _, err = handler.CreateAgentCompatNATStream(firstHandle, "nat-profile-first") + require.NoError(t, err) + _, err = handler.CreateAgentCompatNATStream(secondHandle, "nat-profile-second") + require.NoError(t, err) + require.NoError(t, handler.PublishAgentCompatNATStream(firstHandle, AgentCompatNATPublication{Purpose: AgentCompatCapabilityNAT, TargetServerID: 76, ResourceID: 86, StreamID: "nat-profile-first"})) + require.NoError(t, handler.PublishAgentCompatNATStream(secondHandle, AgentCompatNATPublication{Purpose: AgentCompatCapabilityNAT, TargetServerID: 76, ResourceID: 87, StreamID: "nat-profile-second"})) + + firstStream, err := handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccess(first, firstRegistration)) + require.NoError(t, err) + secondStream, err := handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccess(second, secondRegistration)) + require.NoError(t, err) + require.Equal(t, "nat-profile-first", firstStream) + require.Equal(t, "nat-profile-second", secondStream) +} diff --git a/service/rpc/io_stream_capability_nat_atomic_start_agentcompat_test.go b/service/rpc/io_stream_capability_nat_atomic_start_agentcompat_test.go new file mode 100644 index 00000000..3b6a8187 --- /dev/null +++ b/service/rpc/io_stream_capability_nat_atomic_start_agentcompat_test.go @@ -0,0 +1,180 @@ +//go:build agentcompat + +package rpc + +import ( + "bytes" + "io" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +type atomicNATEndpoint struct { + handler *NezhaHandler + data *bytes.Reader + written bytes.Buffer + mu sync.Mutex + closed atomic.Int32 + readSeen atomic.Int32 + writeSeen chan struct{} +} + +func (endpoint *atomicNATEndpoint) Read(data []byte) (int, error) { + endpoint.readSeen.Add(1) + endpoint.handler.SnapshotIOStreamState() + return endpoint.data.Read(data) +} + +func (endpoint *atomicNATEndpoint) Write(data []byte) (int, error) { + endpoint.mu.Lock() + defer endpoint.mu.Unlock() + endpoint.handler.SnapshotIOStreamState() + n, err := endpoint.written.Write(data) + if endpoint.writeSeen != nil { + select { + case <-endpoint.writeSeen: + default: + close(endpoint.writeSeen) + } + } + return n, err +} + +func (endpoint *atomicNATEndpoint) Close() error { + endpoint.closed.Add(1) + endpoint.handler.SnapshotIOStreamState() + return nil +} + +func newAtomicNATEndpoint(handler *NezhaHandler, payload string) *atomicNATEndpoint { + return &atomicNATEndpoint{handler: handler, data: bytes.NewReader([]byte(payload)), writeSeen: make(chan struct{})} +} + +func publishAtomicNATStream(t *testing.T, streamID string) (*NezhaHandler, AgentCompatCapabilityAccess, AgentCompatNATPublishHandle, AgentCompatCapabilityRegistration) { + t.Helper() + handler, registration, capability := natCapabilityFixture(t, 301, 302, 303, 304) + access := capabilityAccess(capability, registration) + handle, err := handler.ConsumeAgentCompatNATCapability(access) + require.NoError(t, err) + _, err = handler.CreateAgentCompatNATStream(handle, streamID) + require.NoError(t, err) + require.NoError(t, handler.PublishAgentCompatNATStream(handle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 303, ResourceID: 304, StreamID: streamID, + })) + return handler, access, handle, registration +} + +func TestAgentCompatNATAtomicStartWhenCanceledBeforeCaptureDoesNotTouchReplacement(t *testing.T) { + handler, access, handle, _ := publishAtomicNATStream(t, "atomic-replacement-before-capture") + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(access)) + require.NoError(t, handler.CreateStreamWithPurpose("atomic-replacement-before-capture", 0, 303, PurposeNAT)) + replacement := newAtomicNATEndpoint(handler, "replacement") + require.NoError(t, handler.UserConnected("atomic-replacement-before-capture", replacement)) + require.NoError(t, handler.AgentConnected("atomic-replacement-before-capture", replacement)) + + publicationOwned, err := handler.StartAgentCompatNATStream(handle, time.Millisecond) + + require.True(t, publicationOwned) + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) + require.Equal(t, int32(0), replacement.readSeen.Load()) + require.Equal(t, int32(0), replacement.closed.Load()) + _, found := handler.StreamOwnership("atomic-replacement-before-capture") + require.True(t, found) + t.Logf("replacement after cancel-before-capture: read=%d write=%d close=%d registered=%t", replacement.readSeen.Load(), replacement.written.Len(), replacement.closed.Load(), found) +} + +func TestAgentCompatNATAtomicStartWhenCanceledAfterCaptureDoesNotCloseReplacement(t *testing.T) { + handler, access, handle, _ := publishAtomicNATStream(t, "atomic-replacement-after-capture") + old := newAtomicNATEndpoint(handler, "old") + require.NoError(t, handler.UserConnected("atomic-replacement-after-capture", old)) + + result := make(chan error, 1) + go func() { + _, err := handler.StartAgentCompatNATStream(handle, time.Second) + result <- err + }() + stream := mustGetStream(t, handler, "atomic-replacement-after-capture") + select { + case <-stream.startCaptureCh: + case <-time.After(time.Second): + t.Fatal("atomic start did not capture retained stream") + } + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(access)) + require.NoError(t, handler.CreateStreamWithPurpose("atomic-replacement-after-capture", 0, 303, PurposeNAT)) + replacement := newAtomicNATEndpoint(handler, "replacement") + require.NoError(t, handler.UserConnected("atomic-replacement-after-capture", replacement)) + require.NoError(t, handler.AgentConnected("atomic-replacement-after-capture", replacement)) + + require.EqualError(t, receiveAtomicNATError(t, result), "stream revoked") + require.Equal(t, int32(1), old.closed.Load()) + require.Equal(t, int32(0), replacement.readSeen.Load()) + require.Equal(t, int32(0), replacement.closed.Load()) + _, found := handler.StreamOwnership("atomic-replacement-after-capture") + require.True(t, found) + t.Logf("replacement after cancel-after-capture: read=%d write=%d close=%d registered=%t", replacement.readSeen.Load(), replacement.written.Len(), replacement.closed.Load(), found) +} + +func TestAgentCompatNATAtomicStartDetachesOnlyRetainedStreamAfterRelay(t *testing.T) { + handler, _, handle, registration := publishAtomicNATStream(t, "atomic-normal-completion") + user := newAtomicNATEndpoint(handler, "request-bytes") + agent := newAtomicNATEndpoint(handler, "") + require.NoError(t, handler.UserConnected("atomic-normal-completion", user)) + require.NoError(t, handler.AgentConnected("atomic-normal-completion", agent)) + + result := make(chan error, 1) + var publicationOwned bool + go func() { + var err error + publicationOwned, err = handler.StartAgentCompatNATStream(handle, time.Second) + result <- err + }() + select { + case <-agent.writeSeen: + case <-time.After(time.Second): + t.Fatal("atomic relay did not transfer request bytes") + } + err := receiveAtomicNATError(t, result) + require.True(t, publicationOwned) + require.NoError(t, err) + require.Equal(t, int32(1), user.closed.Load()) + require.Equal(t, int32(1), agent.closed.Load()) + require.Equal(t, "request-bytes", agent.written.String()) + streamID, waitErr := handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccessFromRegistration(handle, registration)) + require.NoError(t, waitErr) + require.Equal(t, "atomic-normal-completion", streamID) + require.NoError(t, handler.CreateStreamWithPurpose("atomic-normal-completion", 0, 303, PurposeNAT)) + replacement := newAtomicNATEndpoint(handler, "replacement") + require.NoError(t, handler.UserConnected("atomic-normal-completion", replacement)) + require.NoError(t, handler.AgentConnected("atomic-normal-completion", replacement)) + require.Equal(t, int32(0), replacement.readSeen.Load()) + require.Equal(t, int32(0), replacement.closed.Load()) + t.Logf("replacement after normal retained teardown: read=%d write=%d close=%d registered=true", replacement.readSeen.Load(), replacement.written.Len(), replacement.closed.Load()) +} + +func capabilityAccessFromRegistration(handle AgentCompatNATPublishHandle, registration AgentCompatCapabilityRegistration) AgentCompatCapabilityAccess { + return AgentCompatCapabilityAccess{Capability: AgentCompatIOStreamCapability{value: handle.capability}, Owner: registration.Owner, Purpose: registration.Purpose, TargetServerID: registration.TargetServerID, ResourceID: registration.ResourceID, ServerAccessAllowed: registration.ServerAccessAllowed} +} + +func mustGetStream(t *testing.T, handler *NezhaHandler, streamID string) *ioStreamContext { + t.Helper() + stream, err := handler.GetStream(streamID) + require.NoError(t, err) + return stream +} + +func receiveAtomicNATError(t *testing.T, result <-chan error) error { + t.Helper() + select { + case err := <-result: + return err + case <-time.After(time.Second): + t.Fatal("atomic start did not return") + return nil + } +} + +var _ io.ReadWriteCloser = (*atomicNATEndpoint)(nil) diff --git a/service/rpc/io_stream_capability_nat_barrier_agentcompat_test.go b/service/rpc/io_stream_capability_nat_barrier_agentcompat_test.go new file mode 100644 index 00000000..e57861a9 --- /dev/null +++ b/service/rpc/io_stream_capability_nat_barrier_agentcompat_test.go @@ -0,0 +1,117 @@ +//go:build agentcompat + +package rpc + +import ( + "sync" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestAgentCompatNATCapabilityUnregisterBarrierMakesQueuedPublishInert(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 38, 48, 60, 70) + handle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(capability, registration)) + require.NoError(t, err) + _, err = handler.CreateAgentCompatNATStream(handle, "nat-barrier") + require.NoError(t, err) + require.NoError(t, handler.detachExactStream("nat-barrier", handle.registration.stream)) + stateBeforeRace := handler.SnapshotIOStreamState() + + publishEntered := make(chan struct{}) + publishRelease := make(chan struct{}) + publishObserverCtx := agentCompatCapabilityTestContext(t) + var observeOnce sync.Once + handler.setAgentCompatCapabilityPublishObserverForTest(func() { + observeOnce.Do(func() { + close(publishEntered) + select { + case <-publishRelease: + case <-publishObserverCtx.Done(): + } + }) + }) + t.Cleanup(func() { handler.setAgentCompatCapabilityPublishObserverForTest(nil) }) + publishResult := make(chan error, 1) + go func() { + publishResult <- handler.PublishAgentCompatNATStream(handle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 60, ResourceID: 70, StreamID: "nat-barrier", + }) + }() + awaitAgentCompatCapabilitySignal(t, publishEntered, "publish did not enter production path before unregister") + + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + close(publishRelease) + require.NoError(t, receiveAgentCompatCapabilityError(t, publishResult, "queued publish did not return after release")) + _, err = handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccess(capability, registration)) + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) + require.Equal(t, stateBeforeRace, handler.SnapshotIOStreamState()) + _, found := handler.StreamOwnership("nat-barrier") + require.False(t, found) + handler.ioStreamMutex.RLock() + _, active := handler.agentCompatCapabilities.active[capability.value] + retainedStreamID := handle.registration.streamID + retainedStream := handle.registration.stream + handler.ioStreamMutex.RUnlock() + require.False(t, active) + require.Equal(t, "nat-barrier", retainedStreamID) + require.NotNil(t, retainedStream) +} + +func TestAgentCompatNATCapabilityPublishObserverCanReenterRegistry(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 39, 49, 61, 71) + handle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(capability, registration)) + require.NoError(t, err) + _, err = handler.CreateAgentCompatNATStream(handle, "nat-observer-reentry") + require.NoError(t, err) + observerEntered := make(chan struct{}) + var observeOnce sync.Once + handler.setAgentCompatCapabilityPublishObserverForTest(func() { + handler.SnapshotIOStreamState() + observeOnce.Do(func() { close(observerEntered) }) + }) + t.Cleanup(func() { handler.setAgentCompatCapabilityPublishObserverForTest(nil) }) + publishResult := make(chan error, 1) + go func() { + publishResult <- handler.PublishAgentCompatNATStream(handle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 61, ResourceID: 71, StreamID: "nat-observer-reentry", + }) + }() + + awaitAgentCompatCapabilitySignal(t, observerEntered, "publish observer did not reenter registry") + require.NoError(t, receiveAgentCompatCapabilityError(t, publishResult, "publish observer reentry deadlocked")) + streamID, err := handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccess(capability, registration)) + require.NoError(t, err) + require.Equal(t, "nat-observer-reentry", streamID) +} + +func TestAgentCompatNATCapabilityPublishObserverIsHandlerScoped(t *testing.T) { + first, firstRegistration, firstCapability := natCapabilityFixture(t, 40, 50, 62, 72) + second, secondRegistration, secondCapability := natCapabilityFixture(t, 41, 51, 63, 73) + firstHandle, err := first.ConsumeAgentCompatNATCapability(capabilityAccess(firstCapability, firstRegistration)) + require.NoError(t, err) + secondHandle, err := second.ConsumeAgentCompatNATCapability(capabilityAccess(secondCapability, secondRegistration)) + require.NoError(t, err) + _, err = first.CreateAgentCompatNATStream(firstHandle, "nat-scoped-first") + require.NoError(t, err) + _, err = second.CreateAgentCompatNATStream(secondHandle, "nat-scoped-second") + require.NoError(t, err) + firstObserved := make(chan struct{}) + secondObserved := make(chan struct{}) + first.setAgentCompatCapabilityPublishObserverForTest(func() { close(firstObserved) }) + second.setAgentCompatCapabilityPublishObserverForTest(func() { close(secondObserved) }) + + require.NoError(t, first.PublishAgentCompatNATStream(firstHandle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 62, ResourceID: 72, StreamID: "nat-scoped-first", + })) + awaitAgentCompatCapabilitySignal(t, firstObserved, "first handler observer did not run") + select { + case <-secondObserved: + t.Fatal("second handler observer ran for first handler publish") + default: + } + require.NoError(t, second.PublishAgentCompatNATStream(secondHandle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 63, ResourceID: 73, StreamID: "nat-scoped-second", + })) + awaitAgentCompatCapabilitySignal(t, secondObserved, "second handler observer did not run") +} diff --git a/service/rpc/io_stream_capability_nat_handle_binding_agentcompat_test.go b/service/rpc/io_stream_capability_nat_handle_binding_agentcompat_test.go new file mode 100644 index 00000000..8ea5045f --- /dev/null +++ b/service/rpc/io_stream_capability_nat_handle_binding_agentcompat_test.go @@ -0,0 +1,75 @@ +//go:build agentcompat + +package rpc + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestAgentCompatNATHandleCreationBindsEachHandleToItsOwnStream(t *testing.T) { + handler := NewNezhaHandler() + firstRegistration := capabilityRegistration(capabilityOwner(501, 502), AgentCompatCapabilityNAT, 503, 504) + secondRegistration := capabilityRegistration(capabilityOwner(505, 506), AgentCompatCapabilityNAT, 503, 507) + firstCapability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), firstRegistration) + require.NoError(t, err) + secondCapability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), secondRegistration) + require.NoError(t, err) + firstHandle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(firstCapability, firstRegistration)) + require.NoError(t, err) + secondHandle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(secondCapability, secondRegistration)) + require.NoError(t, err) + + _, err = handler.CreateAgentCompatNATStream(firstHandle, "handle-bound-first") + require.NoError(t, err) + _, err = handler.CreateAgentCompatNATStream(secondHandle, "handle-bound-second") + require.NoError(t, err) + beforeCrossPublish := snapshotAgentCompatNATCreationState(handler, firstRegistration.Owner.PATID, secondRegistration.Owner.PATID) + + require.ErrorIs(t, handler.PublishAgentCompatNATStream(firstHandle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 503, ResourceID: 504, StreamID: "handle-bound-second", + }), ErrAgentCompatCapabilityHidden) + require.ErrorIs(t, handler.PublishAgentCompatNATStream(secondHandle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 503, ResourceID: 507, StreamID: "handle-bound-first", + }), ErrAgentCompatCapabilityHidden) + requireUnchangedAgentCompatNATCreationState(t, handler, beforeCrossPublish) + requireNATHandleBindingsIntact(t, handler, firstHandle, "handle-bound-first") + requireNATHandleBindingsIntact(t, handler, secondHandle, "handle-bound-second") + require.NoError(t, handler.PublishAgentCompatNATStream(firstHandle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 503, ResourceID: 504, StreamID: "handle-bound-first", + })) + require.NoError(t, handler.PublishAgentCompatNATStream(secondHandle, AgentCompatNATPublication{ + Purpose: AgentCompatCapabilityNAT, TargetServerID: 503, ResourceID: 507, StreamID: "handle-bound-second", + })) + firstLease, err := handler.CreateAgentCompatNATStream(firstHandle, "handle-bound-again") + require.Nil(t, firstLease) + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) + require.Equal(t, 2, handler.StreamCount()) +} + +func TestAgentCompatNATHandleCreationCancelReleasesOnlyItsBoundStream(t *testing.T) { + handler := NewNezhaHandler() + firstRegistration := capabilityRegistration(capabilityOwner(508, 509), AgentCompatCapabilityNAT, 510, 511) + secondRegistration := capabilityRegistration(capabilityOwner(512, 513), AgentCompatCapabilityNAT, 510, 514) + firstCapability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), firstRegistration) + require.NoError(t, err) + secondCapability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), secondRegistration) + require.NoError(t, err) + firstAccess := capabilityAccess(firstCapability, firstRegistration) + secondAccess := capabilityAccess(secondCapability, secondRegistration) + firstHandle, err := handler.ConsumeAgentCompatNATCapability(firstAccess) + require.NoError(t, err) + secondHandle, err := handler.ConsumeAgentCompatNATCapability(secondAccess) + require.NoError(t, err) + _, err = handler.CreateAgentCompatNATStream(firstHandle, "bound-first") + require.NoError(t, err) + _, err = handler.CreateAgentCompatNATStream(secondHandle, "bound-second") + require.NoError(t, err) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(firstAccess)) + _, firstFound := handler.StreamOwnership("bound-first") + _, secondFound := handler.StreamOwnership("bound-second") + require.False(t, firstFound) + require.True(t, secondFound) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(secondAccess)) +} diff --git a/service/rpc/io_stream_capability_nat_lease_authority_agentcompat_test.go b/service/rpc/io_stream_capability_nat_lease_authority_agentcompat_test.go new file mode 100644 index 00000000..ea4bd5ec --- /dev/null +++ b/service/rpc/io_stream_capability_nat_lease_authority_agentcompat_test.go @@ -0,0 +1,202 @@ +//go:build agentcompat + +package rpc + +import ( + "errors" + "strconv" + "testing" + + "github.com/stretchr/testify/require" +) + +type agentCompatNATCreationState struct { + streamState IOStreamState + active int + used int + patState map[uint64]agentCompatNATPATState +} + +type agentCompatNATPATState struct { + active uint16 + exists bool +} + +func snapshotAgentCompatNATCreationState(handler *NezhaHandler, patIDs ...uint64) agentCompatNATCreationState { + active, used := agentCompatCapabilityRegistryCounts(handler) + patState := make(map[uint64]agentCompatNATPATState, len(patIDs)) + for _, patID := range patIDs { + patActive, patExists := agentCompatCapabilityActiveForPAT(handler, patID) + patState[patID] = agentCompatNATPATState{active: patActive, exists: patExists} + } + return agentCompatNATCreationState{ + streamState: handler.SnapshotIOStreamState(), + active: active, + used: used, + patState: patState, + } +} + +func requireUnchangedAgentCompatNATCreationState(t *testing.T, handler *NezhaHandler, before agentCompatNATCreationState) { + t.Helper() + ids := make([]uint64, 0, len(before.patState)) + for patID := range before.patState { + ids = append(ids, patID) + } + after := snapshotAgentCompatNATCreationState(handler, ids...) + require.Equal(t, before, after) +} + +func requireHiddenCreateAgentCompatNATStream(t *testing.T, handler *NezhaHandler, handle AgentCompatNATPublishHandle, streamID string, before agentCompatNATCreationState) { + t.Helper() + lease, err := handler.CreateAgentCompatNATStream(handle, streamID) + require.Nil(t, lease) + require.True(t, errors.Is(err, ErrAgentCompatCapabilityHidden) || errors.Is(err, ErrAgentCompatCapabilityUnavailable)) + requireUnchangedAgentCompatNATCreationState(t, handler, before) +} + +func requireNATHandleBindingsIntact(t *testing.T, handler *NezhaHandler, handle AgentCompatNATPublishHandle, streamID string) { + t.Helper() + handler.ioStreamMutex.RLock() + defer handler.ioStreamMutex.RUnlock() + registration := handle.registration + require.NotNil(t, registration) + require.Equal(t, agentCompatCapabilityConsumed, registration.phase) + require.Equal(t, streamID, registration.streamID) + stream, exists := handler.ioStreams[streamID] + require.True(t, exists) + require.Same(t, stream, registration.stream) +} + +func TestAgentCompatNATHandleCreationAuthorityIsBoundToExactRegistration(t *testing.T) { + handler := NewNezhaHandler() + firstRegistration := capabilityRegistration(capabilityOwner(601, 602), AgentCompatCapabilityNAT, 603, 604) + secondRegistration := capabilityRegistration(capabilityOwner(605, 606), AgentCompatCapabilityNAT, 603, 607) + firstCapability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), firstRegistration) + if err != nil { + t.Fatal("first capability registration failed") + } + secondCapability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), secondRegistration) + if err != nil { + t.Fatal("second capability registration failed") + } + firstHandle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(firstCapability, firstRegistration)) + if err != nil { + t.Fatal("first capability consume failed") + } + secondHandle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(secondCapability, secondRegistration)) + if err != nil { + t.Fatal("second capability consume failed") + } + firstLease, err := handler.CreateAgentCompatNATStream(firstHandle, "bound-first") + if err != nil || firstLease == nil { + t.Fatal("first capability did not create its stream") + } + secondLease, err := handler.CreateAgentCompatNATStream(secondHandle, "bound-second") + if err != nil || secondLease == nil { + t.Fatal("second capability did not create its stream") + } + beforeCrossPublish := snapshotAgentCompatNATCreationState(handler, firstRegistration.Owner.PATID, secondRegistration.Owner.PATID) + if err := handler.PublishAgentCompatNATStream(firstHandle, AgentCompatNATPublication{Purpose: AgentCompatCapabilityNAT, TargetServerID: 603, ResourceID: 604, StreamID: "bound-second"}); err != ErrAgentCompatCapabilityHidden { + t.Fatal("first capability published the second stream") + } + if err := handler.PublishAgentCompatNATStream(secondHandle, AgentCompatNATPublication{Purpose: AgentCompatCapabilityNAT, TargetServerID: 603, ResourceID: 607, StreamID: "bound-first"}); err != ErrAgentCompatCapabilityHidden { + t.Fatal("second capability published the first stream") + } + requireUnchangedAgentCompatNATCreationState(t, handler, beforeCrossPublish) + requireNATHandleBindingsIntact(t, handler, firstHandle, "bound-first") + requireNATHandleBindingsIntact(t, handler, secondHandle, "bound-second") + require.NoError(t, handler.PublishAgentCompatNATStream(firstHandle, AgentCompatNATPublication{Purpose: AgentCompatCapabilityNAT, TargetServerID: 603, ResourceID: 604, StreamID: "bound-first"})) + publishedBeforeRepeat := snapshotAgentCompatNATCreationState(handler, firstRegistration.Owner.PATID) + requireHiddenCreateAgentCompatNATStream(t, handler, firstHandle, "bound-after-publish", publishedBeforeRepeat) + if lease, err := handler.CreateAgentCompatNATStream(firstHandle, "bound-again"); lease != nil || err != ErrAgentCompatCapabilityHidden { + t.Fatal("repeated creation mutated first capability state") + } + if handler.StreamCount() != 2 { + t.Fatal("repeated creation changed stream accounting") + } + if err := handler.CancelAgentCompatIOStreamCapability(capabilityAccess(firstCapability, firstRegistration)); err != nil { + t.Fatal("first capability cancellation failed") + } + if _, found := handler.StreamOwnership("bound-first"); found { + t.Fatal("first stream remained after cancellation") + } + if _, found := handler.StreamOwnership("bound-second"); !found { + t.Fatal("second stream was affected by first cancellation") + } + if err := handler.CloseAgentCompatNATStreamLease(secondLease); err != nil { + t.Fatal("second exact lease close failed") + } +} + +func TestAgentCompatNATHandleCreationRejectsInvalidAuthorityWithoutMutation(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 701, 702, 703, 704) + access := capabilityAccess(capability, registration) + before := snapshotAgentCompatNATCreationState(handler, registration.Owner.PATID) + requireHiddenCreateAgentCompatNATStream(t, handler, AgentCompatNATPublishHandle{}, "invalid-zero", before) + + foreignHandler, foreignRegistration, foreignCapability := natCapabilityFixture(t, 705, 706, 703, 707) + foreignHandle, err := foreignHandler.ConsumeAgentCompatNATCapability(capabilityAccess(foreignCapability, foreignRegistration)) + require.NoError(t, err) + foreignBefore := snapshotAgentCompatNATCreationState(foreignHandler, foreignRegistration.Owner.PATID) + requireHiddenCreateAgentCompatNATStream(t, handler, foreignHandle, "invalid-foreign", before) + requireUnchangedAgentCompatNATCreationState(t, foreignHandler, foreignBefore) + + registeredCapability := registerAgentCompatCapability(t, handler, capabilityRegistration(capabilityOwner(708, 709), AgentCompatCapabilityNAT, 703, 710)) + registeredParsed, err := ParseAgentCompatIOStreamCapability(registeredCapability.String()) + require.NoError(t, err) + registeredHandle := AgentCompatNATPublishHandle{capability: registeredParsed.value} + handler.ioStreamMutex.RLock() + registeredHandle.registration = handler.agentCompatCapabilities.active[registeredParsed.value] + registeredHandle.generation = registeredHandle.registration.generation + handler.ioStreamMutex.RUnlock() + registeredBefore := snapshotAgentCompatNATCreationState(handler, registration.Owner.PATID, 708) + requireHiddenCreateAgentCompatNATStream(t, handler, registeredHandle, "invalid-registered", registeredBefore) + + staleHandle, err := handler.ConsumeAgentCompatNATCapability(access) + require.NoError(t, err) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(access)) + staleBefore := snapshotAgentCompatNATCreationState(handler, registration.Owner.PATID) + requireHiddenCreateAgentCompatNATStream(t, handler, staleHandle, "invalid-unregistered", staleBefore) + + cancelRegistration := capabilityRegistration(capabilityOwner(711, 712), AgentCompatCapabilityNAT, 703, 713) + cancelCapability := registerAgentCompatCapability(t, handler, cancelRegistration) + cancelAccess := capabilityAccess(cancelCapability, cancelRegistration) + cancelHandle, err := handler.ConsumeAgentCompatNATCapability(cancelAccess) + require.NoError(t, err) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(cancelAccess)) + cancelledBefore := snapshotAgentCompatNATCreationState(handler, registration.Owner.PATID, cancelRegistration.Owner.PATID) + requireHiddenCreateAgentCompatNATStream(t, handler, cancelHandle, "invalid-cancelled", cancelledBefore) + + wrongPurposeRegistration := capabilityRegistration(capabilityOwner(714, 715), AgentCompatCapabilityTerminal, 703, 0) + wrongPurposeCapability := registerAgentCompatCapability(t, handler, wrongPurposeRegistration) + wrongPurposeParsed, err := ParseAgentCompatIOStreamCapability(wrongPurposeCapability.String()) + require.NoError(t, err) + handler.ioStreamMutex.RLock() + wrongPurposeRegistrationState := handler.agentCompatCapabilities.active[wrongPurposeParsed.value] + handler.ioStreamMutex.RUnlock() + wrongPurposeHandle := AgentCompatNATPublishHandle{registration: wrongPurposeRegistrationState, generation: wrongPurposeRegistrationState.generation, capability: wrongPurposeParsed.value} + wrongPurposeBefore := snapshotAgentCompatNATCreationState(handler, registration.Owner.PATID, wrongPurposeRegistration.Owner.PATID) + requireHiddenCreateAgentCompatNATStream(t, handler, wrongPurposeHandle, "invalid-purpose", wrongPurposeBefore) + +} + +func TestAgentCompatNATHandleCreationRepeatedCreatePreservesAccountingAndQuota(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 721, 722, 723, 724) + handle, err := handler.ConsumeAgentCompatNATCapability(capabilityAccess(capability, registration)) + require.NoError(t, err) + lease, err := handler.CreateAgentCompatNATStream(handle, "repeated-create-first") + require.NoError(t, err) + stateBeforeRepeat := snapshotAgentCompatNATCreationState(handler, registration.Owner.PATID) + requireHiddenCreateAgentCompatNATStream(t, handler, handle, "repeated-create-second", stateBeforeRepeat) + + for index := 0; index < maxStreamsPerServer-1; index++ { + require.NoError(t, handler.CreateStreamWithPurpose("quota-boundary-"+strconv.Itoa(index), 0, registration.TargetServerID, PurposeNAT)) + } + require.ErrorIs(t, handler.CreateStreamWithPurpose("quota-boundary-overflow", 0, registration.TargetServerID, PurposeNAT), ErrTooManyStreamsForServer) + require.NoError(t, handler.CloseAgentCompatNATStreamLease(lease)) + for index := 0; index < maxStreamsPerServer-1; index++ { + require.NoError(t, handler.CloseStream("quota-boundary-"+strconv.Itoa(index))) + } + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) +} diff --git a/service/rpc/io_stream_capability_nat_profile_agentcompat_test.go b/service/rpc/io_stream_capability_nat_profile_agentcompat_test.go new file mode 100644 index 00000000..bfedd0e3 --- /dev/null +++ b/service/rpc/io_stream_capability_nat_profile_agentcompat_test.go @@ -0,0 +1,99 @@ +//go:build agentcompat + +package rpc + +import ( + "errors" + "strings" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestAgentCompatNATCapabilityForProfileConsumesStoredRegistration(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 61, 71, 81, 91) + + access, handle, err := handler.ConsumeAgentCompatNATCapabilityForProfile(capability.String(), 81, 91) + + require.NoError(t, err) + require.Equal(t, registration.Owner, access.Owner) + require.Equal(t, registration.Purpose, access.Purpose) + require.Equal(t, registration.TargetServerID, access.TargetServerID) + require.Equal(t, registration.ResourceID, access.ResourceID) + require.True(t, access.ServerAccessAllowed) + require.NotEmpty(t, handle.capability) +} + +func TestAgentCompatNATCapabilityForProfileHidesMalformedUnknownAndForeignTuples(t *testing.T) { + handler, _, capability := natCapabilityFixture(t, 62, 72, 82, 92) + terminalRegistration := capabilityRegistration(capabilityOwner(66, 76), AgentCompatCapabilityTerminal, 82, 0) + terminalCapability := registerAgentCompatCapability(t, handler, terminalRegistration) + cases := []struct { + name string + value string + serverID uint64 + resourceID uint64 + }{ + {name: "malformed", value: "not-a-capability", serverID: 82, resourceID: 92}, + {name: "unknown", value: strings.Repeat("a", 43), serverID: 82, resourceID: 92}, + {name: "wrong server", value: capability.String(), serverID: 83, resourceID: 92}, + {name: "wrong profile", value: capability.String(), serverID: 82, resourceID: 93}, + {name: "wrong purpose", value: terminalCapability.String(), serverID: 82, resourceID: 0}, + } + for _, testCase := range cases { + t.Run(testCase.name, func(t *testing.T) { + _, _, err := handler.ConsumeAgentCompatNATCapabilityForProfile(testCase.value, testCase.serverID, testCase.resourceID) + require.True(t, errors.Is(err, ErrAgentCompatCapabilityHidden)) + require.NotContains(t, err.Error(), testCase.value) + }) + } +} + +func TestAgentCompatNATCapabilityForProfileHidesRepeatedAndInactiveConsume(t *testing.T) { + handler, _, capability := natCapabilityFixture(t, 63, 73, 83, 93) + activeBefore, usedBefore := agentCompatCapabilityRegistryCounts(handler) + _, _, err := handler.ConsumeAgentCompatNATCapabilityForProfile(capability.String(), 83, 93) + require.NoError(t, err) + activeAfter, usedAfter := agentCompatCapabilityRegistryCounts(handler) + require.Equal(t, activeBefore, activeAfter) + require.Equal(t, usedBefore, usedAfter) + + _, _, err = handler.ConsumeAgentCompatNATCapabilityForProfile(capability.String(), 83, 93) + require.True(t, errors.Is(err, ErrAgentCompatCapabilityHidden)) + activeAfterRepeat, usedAfterRepeat := agentCompatCapabilityRegistryCounts(handler) + require.Equal(t, activeAfter, activeAfterRepeat) + require.Equal(t, usedAfter, usedAfterRepeat) + + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(AgentCompatCapabilityAccess{})) +} + +func TestAgentCompatNATCapabilityForProfileHidesCancelledAndUnregistered(t *testing.T) { + tests := []struct { + name string + cleanup func(*NezhaHandler, AgentCompatCapabilityAccess) error + }{ + {name: "cancelled", cleanup: (*NezhaHandler).CancelAgentCompatIOStreamCapability}, + {name: "unregistered", cleanup: (*NezhaHandler).UnregisterAgentCompatIOStreamCapability}, + } + for _, testCase := range tests { + t.Run(testCase.name, func(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 64, 74, 84, 94) + access := capabilityAccess(capability, registration) + require.NoError(t, testCase.cleanup(handler, access)) + + _, _, err := handler.ConsumeAgentCompatNATCapabilityForProfile(capability.String(), 84, 94) + require.True(t, errors.Is(err, ErrAgentCompatCapabilityHidden)) + }) + } +} + +func TestAgentCompatNATCapabilityForProfileDoesNotLeakSensitiveValues(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 65, 75, 85, 95) + _, _, err := handler.ConsumeAgentCompatNATCapabilityForProfile(capability.String(), 86, 95) + require.Error(t, err) + message := err.Error() + for _, sensitive := range []string{capability.String(), "65", "75", "85", "95", "nat"} { + require.NotContains(t, message, sensitive) + } + require.Equal(t, AgentCompatCapabilityNAT, registration.Purpose) +} diff --git a/service/rpc/io_stream_capability_publication_agentcompat.go b/service/rpc/io_stream_capability_publication_agentcompat.go new file mode 100644 index 00000000..39880b7e --- /dev/null +++ b/service/rpc/io_stream_capability_publication_agentcompat.go @@ -0,0 +1,66 @@ +//go:build agentcompat + +package rpc + +import ( + "errors" + "time" +) + +func (s *NezhaHandler) CreateAgentCompatNATStream(handle AgentCompatNATPublishHandle, streamID string) (*AgentCompatNATStreamLease, error) { + s.ioStreamMutex.Lock() + defer s.ioStreamMutex.Unlock() + registration := handle.registration + if registration == nil || registration.generation != handle.generation { + return nil, ErrAgentCompatCapabilityHidden + } + currentRegistration, active := s.agentCompatCapabilities.active[handle.capability] + stored := registration.registration + if !active || currentRegistration != registration || registration.phase != agentCompatCapabilityConsumed || + stored.Purpose != AgentCompatCapabilityNAT || registration.stream != nil || streamID == "" { + return nil, ErrAgentCompatCapabilityHidden + } + if err := s.createStreamLocked(streamID, 0, stored.TargetServerID, PurposeNAT); err != nil { + if err == ErrStreamAlreadyExists { + return nil, ErrAgentCompatCapabilityHidden + } + return nil, err + } + stream := s.ioStreams[streamID] + registration.streamID = streamID + registration.stream = stream + return &AgentCompatNATStreamLease{streamID: streamID, stream: stream}, nil +} + +func (s *NezhaHandler) CloseAgentCompatNATStreamLease(lease *AgentCompatNATStreamLease) error { + if lease == nil { + return nil + } + return s.detachExactStream(lease.streamID, lease.stream) +} + +func (s *NezhaHandler) StartAgentCompatNATStream(handle AgentCompatNATPublishHandle, timeout time.Duration) (bool, error) { + s.ioStreamMutex.RLock() + registration := handle.registration + publicationOwned := registration != nil && registration.generation == handle.generation && + registration.phase == agentCompatCapabilityPublished && registration.streamID != "" && registration.stream != nil + if registration == nil || registration.generation != handle.generation { + s.ioStreamMutex.RUnlock() + return publicationOwned, ErrAgentCompatCapabilityHidden + } + current, active := s.agentCompatCapabilities.active[handle.capability] + stored := registration.registration + streamID := registration.streamID + stream := registration.stream + valid := active && current == registration && registration.phase == agentCompatCapabilityPublished && + streamID != "" && stream != nil && s.ioStreams[streamID] == stream && + stream.creatorUserID == 0 && stream.targetServerID == stored.TargetServerID && + stream.purpose == PurposeNAT && stored.Purpose == AgentCompatCapabilityNAT + s.ioStreamMutex.RUnlock() + if !valid { + return publicationOwned, ErrAgentCompatCapabilityHidden + } + startErr := s.startStreamContext(streamID, stream, timeout) + closeErr := s.detachExactStream(streamID, stream) + return publicationOwned, errors.Join(startErr, closeErr) +} diff --git a/service/rpc/io_stream_capability_publication_default.go b/service/rpc/io_stream_capability_publication_default.go new file mode 100644 index 00000000..26ad82de --- /dev/null +++ b/service/rpc/io_stream_capability_publication_default.go @@ -0,0 +1,9 @@ +//go:build !agentcompat + +package rpc + +import "time" + +func (*NezhaHandler) StartAgentCompatNATStream(AgentCompatNATPublishHandle, time.Duration) (bool, error) { + return false, ErrAgentCompatCapabilityUnavailable +} diff --git a/service/rpc/io_stream_capability_quota_accounting_agentcompat_test.go b/service/rpc/io_stream_capability_quota_accounting_agentcompat_test.go new file mode 100644 index 00000000..6573b714 --- /dev/null +++ b/service/rpc/io_stream_capability_quota_accounting_agentcompat_test.go @@ -0,0 +1,230 @@ +//go:build agentcompat + +package rpc + +import ( + "errors" + "sync" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestAgentCompatCapabilityConcurrentRegistrationEnforcesExactPerPATQuota(t *testing.T) { + handler := NewNezhaHandler() + setUniqueAgentCompatCapabilityTokens(handler) + registration := capabilityRegistration(capabilityOwner(104, 204), AgentCompatCapabilityTerminal, 305, 0) + start := make(chan struct{}) + results := make(chan error, 64) + var waitGroup sync.WaitGroup + waitGroup.Add(64) + for range 64 { + go func() { + defer waitGroup.Done() + <-start + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + results <- err + }() + } + close(start) + waitGroup.Wait() + close(results) + + succeeded, unavailable := 0, 0 + for err := range results { + switch { + case err == nil: + succeeded++ + case errors.Is(err, ErrAgentCompatCapabilityUnavailable): + unavailable++ + default: + t.Fatalf("unexpected registration error: %v", err) + } + } + require.Equal(t, 16, succeeded) + require.Equal(t, 48, unavailable) + active, used := agentCompatCapabilityRegistryCounts(handler) + require.Equal(t, 16, active) + require.Equal(t, 16, used) +} + +func TestAgentCompatCapabilityConcurrentRegistrationEnforcesExactGlobalQuota(t *testing.T) { + handler := NewNezhaHandler() + setUniqueAgentCompatCapabilityTokens(handler) + start := make(chan struct{}) + results := make(chan error, 256) + var waitGroup sync.WaitGroup + waitGroup.Add(256) + for index := range 256 { + go func() { + defer waitGroup.Done() + <-start + registration := capabilityRegistration(capabilityOwner(uint64(index+1), uint64(index+1001)), AgentCompatCapabilityTerminal, 306, 0) + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + results <- err + }() + } + close(start) + waitGroup.Wait() + close(results) + + succeeded, unavailable := 0, 0 + for err := range results { + if err == nil { + succeeded++ + continue + } + require.ErrorIs(t, err, ErrAgentCompatCapabilityUnavailable) + unavailable++ + } + require.Equal(t, 128, succeeded) + require.Equal(t, 128, unavailable) + active, used := agentCompatCapabilityRegistryCounts(handler) + require.Equal(t, 128, active) + require.Equal(t, 128, used) +} + +func TestAgentCompatCapabilityConcurrentRegistrationEnforcesExactProcessMintQuota(t *testing.T) { + handler := NewNezhaHandler() + setUniqueAgentCompatCapabilityTokens(handler) + handler.ioStreamMutex.Lock() + for index := range agentCompatCapabilityMaxProcessMints - 1 { + handler.agentCompatCapabilities.used[string(rune(index+1))] = struct{}{} + } + handler.ioStreamMutex.Unlock() + start := make(chan struct{}) + results := make(chan error, 2) + var waitGroup sync.WaitGroup + waitGroup.Add(2) + for index := range 2 { + go func() { + defer waitGroup.Done() + <-start + registration := capabilityRegistration(capabilityOwner(uint64(index+201), uint64(index+301)), AgentCompatCapabilityTerminal, 312, 0) + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + results <- err + }() + } + close(start) + waitGroup.Wait() + close(results) + + succeeded, unavailable := 0, 0 + for err := range results { + if err == nil { + succeeded++ + continue + } + require.ErrorIs(t, err, ErrAgentCompatCapabilityUnavailable) + unavailable++ + } + require.Equal(t, 1, succeeded) + require.Equal(t, 1, unavailable) + active, used := agentCompatCapabilityRegistryCounts(handler) + require.Equal(t, 1, active) + require.Equal(t, agentCompatCapabilityMaxProcessMints, used) +} + +func TestAgentCompatCapabilityRemovalRequiresExactActiveRegistration(t *testing.T) { + handler := NewNezhaHandler() + setUniqueAgentCompatCapabilityTokens(handler) + registration := capabilityRegistration(capabilityOwner(105, 205), AgentCompatCapabilityTerminal, 307, 0) + capability := registerAgentCompatCapability(t, handler, registration) + handler.ioStreamMutex.Lock() + activeRegistration := handler.agentCompatCapabilities.active[capability.value] + staleRegistration := &agentCompatCapabilityRegistration{registration: registration, notify: make(chan struct{})} + handler.removeAgentCompatCapabilityLocked(capability.value, staleRegistration) + handler.ioStreamMutex.Unlock() + for range 15 { + registerAgentCompatCapability(t, handler, registration) + } + + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.ErrorIs(t, err, ErrAgentCompatCapabilityUnavailable) + active, used := agentCompatCapabilityRegistryCounts(handler) + require.Equal(t, 16, active) + require.Equal(t, 16, used) + + handler.ioStreamMutex.Lock() + handler.removeAgentCompatCapabilityLocked(capability.value, activeRegistration) + handler.removeAgentCompatCapabilityLocked(capability.value, activeRegistration) + handler.ioStreamMutex.Unlock() + + replacement := registerAgentCompatCapability(t, handler, registration) + require.NotEmpty(t, replacement.String()) + active, used = agentCompatCapabilityRegistryCounts(handler) + require.Equal(t, 16, active) + require.Equal(t, 17, used) +} + +func TestAgentCompatCapabilityForeignRemovalDoesNotReleasePerPATQuota(t *testing.T) { + handler := NewNezhaHandler() + setUniqueAgentCompatCapabilityTokens(handler) + registration := capabilityRegistration(capabilityOwner(106, 206), AgentCompatCapabilityTerminal, 308, 0) + capability := registerAgentCompatCapability(t, handler, registration) + for range 15 { + registerAgentCompatCapability(t, handler, registration) + } + foreign := capabilityAccess(capability, registration) + foreign.Owner.PATID++ + + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(foreign)) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(foreign)) + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(AgentCompatCapabilityAccess{})) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(AgentCompatCapabilityAccess{})) + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.ErrorIs(t, err, ErrAgentCompatCapabilityUnavailable) + active, used := agentCompatCapabilityRegistryCounts(handler) + require.Equal(t, 16, active) + require.Equal(t, 16, used) +} + +func TestAgentCompatCapabilityCancelReleasesQuotaBeforeEndpointCloseFailure(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityTerminal, "quota-close-failure", 309) + setUniqueAgentCompatCapabilityTokens(handler) + closeErr := errors.New("endpoint close failed") + endpoint := &capabilityCloseEndpoint{handler: handler, streamID: "quota-close-failure", err: closeErr} + require.NoError(t, handler.AgentConnected("quota-close-failure", endpoint)) + for range 15 { + registerAgentCompatCapability(t, handler, registration) + } + + err := handler.CancelAgentCompatIOStreamCapability(capabilityAccess(capability, registration)) + + require.ErrorIs(t, err, closeErr) + replacement := registerAgentCompatCapability(t, handler, registration) + require.NotEmpty(t, replacement.String()) + activeForPAT, exists := agentCompatCapabilityActiveForPAT(handler, registration.Owner.PATID) + require.True(t, exists) + require.Equal(t, uint16(16), activeForPAT) +} + +func TestAgentCompatCapabilityBoundUnregisterConflictRetainsQuota(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityTerminal, "quota-bound-conflict", 310) + setUniqueAgentCompatCapabilityTokens(handler) + for range 15 { + registerAgentCompatCapability(t, handler, registration) + } + + err := handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(capability, registration)) + _, registerErr := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + + require.ErrorIs(t, err, ErrAgentCompatCapabilityBound) + require.ErrorIs(t, registerErr, ErrAgentCompatCapabilityUnavailable) + activeForPAT, exists := agentCompatCapabilityActiveForPAT(handler, registration.Owner.PATID) + require.True(t, exists) + require.Equal(t, uint16(16), activeForPAT) +} + +func TestAgentCompatCapabilityLastRemovalDeletesPerPATAccountingEntry(t *testing.T) { + handler := NewNezhaHandler() + setUniqueAgentCompatCapabilityTokens(handler) + registration := capabilityRegistration(capabilityOwner(107, 207), AgentCompatCapabilityTerminal, 311, 0) + capability := registerAgentCompatCapability(t, handler, registration) + + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + + activeForPAT, exists := agentCompatCapabilityActiveForPAT(handler, registration.Owner.PATID) + require.False(t, exists) + require.Zero(t, activeForPAT) +} diff --git a/service/rpc/io_stream_capability_quota_boundaries_agentcompat_test.go b/service/rpc/io_stream_capability_quota_boundaries_agentcompat_test.go new file mode 100644 index 00000000..1d1c13e8 --- /dev/null +++ b/service/rpc/io_stream_capability_quota_boundaries_agentcompat_test.go @@ -0,0 +1,94 @@ +//go:build agentcompat + +package rpc + +import ( + "errors" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestAgentCompatCapabilityRegistrationEnforcesPerPATActiveQuotaAndReusesReleasedSlot(t *testing.T) { + handler := NewNezhaHandler() + issued := setUniqueAgentCompatCapabilityTokens(handler) + registration := capabilityRegistration(capabilityOwner(101, 201), AgentCompatCapabilityTerminal, 301, 0) + capabilities := make([]AgentCompatIOStreamCapability, 0, 16) + for range 16 { + capabilities = append(capabilities, registerAgentCompatCapability(t, handler, registration)) + } + + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.ErrorIs(t, err, ErrAgentCompatCapabilityUnavailable) + require.Equal(t, uint64(16), issued.Load()) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(capabilities[0], registration))) + + replacement := registerAgentCompatCapability(t, handler, registration) + require.NotEmpty(t, replacement.String()) + active, used := agentCompatCapabilityRegistryCounts(handler) + require.Equal(t, 16, active) + require.Equal(t, 17, used) +} + +func TestAgentCompatCapabilityRegistrationEnforcesGlobalActiveQuotaAndReusesReleasedSlot(t *testing.T) { + handler := NewNezhaHandler() + issued := setUniqueAgentCompatCapabilityTokens(handler) + registrations := make([]AgentCompatCapabilityRegistration, 0, 128) + capabilities := make([]AgentCompatIOStreamCapability, 0, 128) + for index := range 128 { + registration := capabilityRegistration(capabilityOwner(uint64(index+1), uint64(index+1001)), AgentCompatCapabilityTerminal, 302, 0) + registrations = append(registrations, registration) + capabilities = append(capabilities, registerAgentCompatCapability(t, handler, registration)) + } + overflow := capabilityRegistration(capabilityOwner(10000, 20000), AgentCompatCapabilityTerminal, 302, 0) + + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), overflow) + require.ErrorIs(t, err, ErrAgentCompatCapabilityUnavailable) + require.Equal(t, uint64(128), issued.Load()) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(capabilities[0], registrations[0]))) + + replacement := registerAgentCompatCapability(t, handler, overflow) + require.NotEmpty(t, replacement.String()) + active, used := agentCompatCapabilityRegistryCounts(handler) + require.Equal(t, 128, active) + require.Equal(t, 129, used) +} + +func TestAgentCompatCapabilityRegistrationEnforcesProcessLifetimeMintQuota(t *testing.T) { + handler := NewNezhaHandler() + issued := setUniqueAgentCompatCapabilityTokens(handler) + registration := capabilityRegistration(capabilityOwner(102, 202), AgentCompatCapabilityTerminal, 303, 0) + for range 4096 { + capability := registerAgentCompatCapability(t, handler, registration) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(capability, registration))) + } + + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + + require.ErrorIs(t, err, ErrAgentCompatCapabilityUnavailable) + require.Equal(t, uint64(4096), issued.Load()) + active, used := agentCompatCapabilityRegistryCounts(handler) + require.Zero(t, active) + require.Equal(t, 4096, used) +} + +func TestAgentCompatCapabilityCollisionRetriesDoNotConsumeQuota(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(103, 203), AgentCompatCapabilityTerminal, 304, 0) + fixedToken := make([]byte, 32) + fixedToken[0] = 1 + handler.setAgentCompatCapabilityTokenSourceForTest(func(destination []byte) error { + copy(destination, fixedToken) + return nil + }) + first := registerAgentCompatCapability(t, handler, registration) + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(capabilityAccess(first, registration))) + + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + + require.ErrorIs(t, err, ErrAgentCompatCapabilityTokenExhausted) + active, used := agentCompatCapabilityRegistryCounts(handler) + require.Zero(t, active) + require.Equal(t, 1, used) + require.False(t, errors.Is(err, ErrAgentCompatCapabilityUnavailable)) +} diff --git a/service/rpc/io_stream_capability_quota_test_helpers_agentcompat_test.go b/service/rpc/io_stream_capability_quota_test_helpers_agentcompat_test.go new file mode 100644 index 00000000..398ba01f --- /dev/null +++ b/service/rpc/io_stream_capability_quota_test_helpers_agentcompat_test.go @@ -0,0 +1,40 @@ +//go:build agentcompat + +package rpc + +import ( + "encoding/binary" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/require" +) + +func setUniqueAgentCompatCapabilityTokens(handler *NezhaHandler) *atomic.Uint64 { + var issued atomic.Uint64 + handler.setAgentCompatCapabilityTokenSourceForTest(func(destination []byte) error { + binary.LittleEndian.PutUint64(destination, issued.Add(1)) + return nil + }) + return &issued +} + +func registerAgentCompatCapability(t *testing.T, handler *NezhaHandler, registration AgentCompatCapabilityRegistration) AgentCompatIOStreamCapability { + t.Helper() + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + return capability +} + +func agentCompatCapabilityRegistryCounts(handler *NezhaHandler) (active, used int) { + handler.ioStreamMutex.RLock() + defer handler.ioStreamMutex.RUnlock() + return len(handler.agentCompatCapabilities.active), len(handler.agentCompatCapabilities.used) +} + +func agentCompatCapabilityActiveForPAT(handler *NezhaHandler, patID uint64) (uint16, bool) { + handler.ioStreamMutex.RLock() + defer handler.ioStreamMutex.RUnlock() + active, exists := handler.agentCompatCapabilities.activeByPAT[patID] + return active, exists +} diff --git a/service/rpc/io_stream_capability_register_agentcompat.go b/service/rpc/io_stream_capability_register_agentcompat.go new file mode 100644 index 00000000..f92e6034 --- /dev/null +++ b/service/rpc/io_stream_capability_register_agentcompat.go @@ -0,0 +1,96 @@ +//go:build agentcompat + +package rpc + +import ( + "context" + "encoding/base64" +) + +const agentCompatCapabilityTokenAttempts = 32 + +func validAgentCompatRegistration(registration AgentCompatCapabilityRegistration) bool { + if !registration.ServerAccessAllowed || registration.Owner.PATID == 0 || registration.Owner.UserID == 0 || registration.TargetServerID == 0 { + return false + } + switch registration.Purpose { + case AgentCompatCapabilityTerminal, AgentCompatCapabilityFileManager: + return registration.ResourceID == 0 + case AgentCompatCapabilityNAT: + return registration.ResourceID != 0 + default: + return false + } +} + +func (s *NezhaHandler) RegisterAgentCompatIOStreamCapability(ctx context.Context, registration AgentCompatCapabilityRegistration) (AgentCompatIOStreamCapability, error) { + if !validAgentCompatRegistration(registration) { + return AgentCompatIOStreamCapability{}, ErrAgentCompatCapabilityHidden + } + if err := ctx.Err(); err != nil { + return AgentCompatIOStreamCapability{}, err + } + s.ioStreamMutex.RLock() + tokenSource := s.agentCompatCapabilities.tokenSource + quotaAvailable := s.agentCompatCapabilityQuotaAvailableLocked(registration.Owner.PATID) + s.ioStreamMutex.RUnlock() + if !quotaAvailable { + return AgentCompatIOStreamCapability{}, ErrAgentCompatCapabilityUnavailable + } + for range agentCompatCapabilityTokenAttempts { + if err := ctx.Err(); err != nil { + return AgentCompatIOStreamCapability{}, err + } + // Token generation may block or reenter the registry, so it must never run under ioStreamMutex. + raw := make([]byte, 32) + if err := tokenSource(raw); err != nil { + return AgentCompatIOStreamCapability{}, err + } + capability := AgentCompatIOStreamCapability{value: base64.RawURLEncoding.EncodeToString(raw)} + + if err := ctx.Err(); err != nil { + return AgentCompatIOStreamCapability{}, err + } + s.ioStreamMutex.Lock() + // Recheck every quota under the insertion lock so concurrent mints cannot oversubscribe any bound. + if !s.agentCompatCapabilityQuotaAvailableLocked(registration.Owner.PATID) { + s.ioStreamMutex.Unlock() + return AgentCompatIOStreamCapability{}, ErrAgentCompatCapabilityUnavailable + } + if _, used := s.agentCompatCapabilities.used[capability.value]; used { + s.ioStreamMutex.Unlock() + continue + } + s.agentCompatCapabilities.used[capability.value] = struct{}{} + s.agentCompatCapabilities.nextIdentity++ + s.agentCompatCapabilities.activeByPAT[registration.Owner.PATID]++ + s.agentCompatCapabilities.active[capability.value] = &agentCompatCapabilityRegistration{ + registration: registration, phase: agentCompatCapabilityRegistered, + generation: s.agentCompatCapabilities.nextIdentity, notify: make(chan struct{}), + } + s.ioStreamMutex.Unlock() + return capability, nil + } + return AgentCompatIOStreamCapability{}, ErrAgentCompatCapabilityTokenExhausted +} + +func (s *NezhaHandler) agentCompatCapabilityQuotaAvailableLocked(patID uint64) bool { + return s.agentCompatCapabilities.activeByPAT[patID] < agentCompatCapabilityMaxActivePerPAT && + len(s.agentCompatCapabilities.active) < agentCompatCapabilityMaxActiveGlobal && + len(s.agentCompatCapabilities.used) < agentCompatCapabilityMaxProcessMints +} + +func sameAgentCompatOwner(left, right AgentCompatCapabilityOwner) bool { + return left == right +} + +func agentCompatAccessMatches(access AgentCompatCapabilityAccess, registration *agentCompatCapabilityRegistration) bool { + stored := registration.registration + return access.ServerAccessAllowed && sameAgentCompatOwner(access.Owner, stored.Owner) && + access.Purpose == stored.Purpose && access.TargetServerID == stored.TargetServerID && access.ResourceID == stored.ResourceID +} + +func (s *NezhaHandler) agentCompatRegistrationLocked(access AgentCompatCapabilityAccess) (*agentCompatCapabilityRegistration, bool) { + registration, exists := s.agentCompatCapabilities.active[access.Capability.value] + return registration, exists && agentCompatAccessMatches(access, registration) +} diff --git a/service/rpc/io_stream_capability_security_agentcompat_test.go b/service/rpc/io_stream_capability_security_agentcompat_test.go new file mode 100644 index 00000000..bcbe5933 --- /dev/null +++ b/service/rpc/io_stream_capability_security_agentcompat_test.go @@ -0,0 +1,234 @@ +//go:build agentcompat + +package rpc + +import ( + "context" + "errors" + "sync" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestAgentCompatCapabilityRegistrationExhaustsPermanentCollisionWithoutBlockingRegistry(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(31, 41), AgentCompatCapabilityTerminal, 51, 0) + fixedToken := make([]byte, 32) + fixedToken[0] = 1 + handler.setAgentCompatCapabilityTokenSourceForTest(func(destination []byte) error { + copy(destination, fixedToken) + return nil + }) + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + + entered := make(chan struct{}) + release := make(chan struct{}) + releaseCtx := agentCompatCapabilityTestContext(t) + var once sync.Once + handler.setAgentCompatCapabilityTokenSourceForTest(func(destination []byte) error { + once.Do(func() { + close(entered) + select { + case <-release: + case <-releaseCtx.Done(): + } + }) + copy(destination, fixedToken) + return nil + }) + result := make(chan error, 1) + go func() { + _, registerErr := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + result <- registerErr + }() + awaitAgentCompatCapabilitySignal(t, entered, "token source did not enter") + + registryRead := make(chan struct{}) + go func() { + handler.SnapshotIOStreamState() + close(registryRead) + }() + awaitAgentCompatCapabilitySignal(t, registryRead, "token source blocked unrelated registry operation") + close(release) + require.ErrorIs(t, receiveAgentCompatCapabilityError(t, result, "permanent token collision did not terminate"), ErrAgentCompatCapabilityTokenExhausted) +} + +func TestAgentCompatCapabilityRegistrationPreservesCanceledContext(t *testing.T) { + handler := NewNezhaHandler() + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err := handler.RegisterAgentCompatIOStreamCapability(ctx, capabilityRegistration(capabilityOwner(32, 42), AgentCompatCapabilityTerminal, 52, 0)) + + require.ErrorIs(t, err, context.Canceled) +} + +func TestAgentCompatCapabilityTokenSourceCanReenterRegistry(t *testing.T) { + handler := NewNezhaHandler() + handler.setAgentCompatCapabilityTokenSourceForTest(func(destination []byte) error { + handler.SnapshotIOStreamState() + destination[0] = 1 + return nil + }) + result := make(chan error, 1) + go func() { + _, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityRegistration(capabilityOwner(33, 43), AgentCompatCapabilityTerminal, 53, 0)) + result <- err + }() + + require.NoError(t, receiveAgentCompatCapabilityError(t, result, "reentrant token source deadlocked")) +} + +func TestAgentCompatCapabilityWaitObserverCanReenterRegistry(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(34, 44), AgentCompatCapabilityTerminal, 54, 0) + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + access := capabilityAccess(capability, registration) + handler.setAgentCompatCapabilityWaitObserverForTest(func() { + require.NoError(t, handler.UnregisterAgentCompatIOStreamCapability(access)) + }) + result := make(chan error, 1) + waitCtx := agentCompatCapabilityTestContext(t) + go func() { + _, waitErr := handler.WaitAgentCompatIOStreamCapability(waitCtx, access) + result <- waitErr + }() + + require.ErrorIs(t, receiveAgentCompatCapabilityError(t, result, "reentrant wait observer deadlocked"), ErrAgentCompatCapabilityHidden) +} + +func TestAgentCompatCapabilityWaitRejectsSameIDReplacement(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityTerminal, "reused-stream-id", 55) + require.NoError(t, handler.CloseStream("reused-stream-id")) + require.NoError(t, handler.CreateStreamWithPurpose("reused-stream-id", 21, 55, PurposeTerminal)) + + _, err := handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccess(capability, registration)) + + require.ErrorIs(t, err, ErrAgentCompatCapabilityHidden) +} + +func TestAgentCompatCapabilityCancelAndUnregisterDoNotEnumerateForeignIdentity(t *testing.T) { + operations := []struct { + name string + run func(*NezhaHandler, AgentCompatCapabilityAccess) error + }{ + {name: "cancel", run: (*NezhaHandler).CancelAgentCompatIOStreamCapability}, + {name: "unregister", run: (*NezhaHandler).UnregisterAgentCompatIOStreamCapability}, + } + for _, operation := range operations { + t.Run(operation.name, func(t *testing.T) { + handler, registration, capability := boundCapabilityFixture(t, AgentCompatCapabilityTerminal, operation.name+"-foreign", 56) + foreign := capabilityAccess(capability, registration) + foreign.Owner.PATID++ + before := handler.SnapshotIOStreamState() + + foreignErr := operation.run(handler, foreign) + unknownErr := operation.run(handler, AgentCompatCapabilityAccess{}) + + require.NoError(t, foreignErr) + require.NoError(t, unknownErr) + require.Equal(t, before, handler.SnapshotIOStreamState()) + _, found := handler.StreamOwnership(operation.name + "-foreign") + require.True(t, found) + }) + } +} + +func TestAgentCompatCapabilityAccessMismatchMatrixIsHiddenOrInert(t *testing.T) { + mutations := []struct { + name string + mutate func(*AgentCompatCapabilityAccess) + }{ + {name: "PAT", mutate: func(access *AgentCompatCapabilityAccess) { access.Owner.PATID++ }}, + {name: "user", mutate: func(access *AgentCompatCapabilityAccess) { access.Owner.UserID++ }}, + {name: "admin", mutate: func(access *AgentCompatCapabilityAccess) { access.Owner.IsAdmin = !access.Owner.IsAdmin }}, + {name: "purpose", mutate: func(access *AgentCompatCapabilityAccess) { access.Purpose = AgentCompatCapabilityFileManager }}, + {name: "resource", mutate: func(access *AgentCompatCapabilityAccess) { access.ResourceID++ }}, + {name: "server", mutate: func(access *AgentCompatCapabilityAccess) { access.TargetServerID++ }}, + {name: "access proof", mutate: func(access *AgentCompatCapabilityAccess) { access.ServerAccessAllowed = false }}, + } + for _, mutation := range mutations { + t.Run(mutation.name, func(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(35, 45), AgentCompatCapabilityTerminal, 57, 0) + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + require.NoError(t, handler.CreateStreamWithPurpose("matrix", 45, 57, PurposeTerminal)) + access := capabilityAccess(capability, registration) + mutation.mutate(&access) + + _, waitErr := handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), access) + bindErr := handler.BindAgentCompatIOStreamCapability(AgentCompatCapabilityBinding{AgentCompatCapabilityAccess: access, StreamID: "matrix"}) + cancelErr := handler.CancelAgentCompatIOStreamCapability(access) + unregisterErr := handler.UnregisterAgentCompatIOStreamCapability(access) + + require.ErrorIs(t, waitErr, ErrAgentCompatCapabilityHidden) + require.ErrorIs(t, bindErr, ErrAgentCompatCapabilityHidden) + require.NoError(t, cancelErr) + require.NoError(t, unregisterErr) + require.NoError(t, handler.BindAgentCompatIOStreamCapability(AgentCompatCapabilityBinding{ + AgentCompatCapabilityAccess: capabilityAccess(capability, registration), StreamID: "matrix", + })) + streamID, err := handler.WaitAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), capabilityAccess(capability, registration)) + require.NoError(t, err) + require.Equal(t, "matrix", streamID) + }) + } +} + +func TestAgentCompatNATCapabilityConsumeMismatchMatrixIsHidden(t *testing.T) { + mutations := []struct { + name string + mutate func(*AgentCompatCapabilityAccess) + }{ + {name: "PAT", mutate: func(access *AgentCompatCapabilityAccess) { access.Owner.PATID++ }}, + {name: "user", mutate: func(access *AgentCompatCapabilityAccess) { access.Owner.UserID++ }}, + {name: "admin", mutate: func(access *AgentCompatCapabilityAccess) { access.Owner.IsAdmin = !access.Owner.IsAdmin }}, + {name: "purpose", mutate: func(access *AgentCompatCapabilityAccess) { access.Purpose = AgentCompatCapabilityTerminal }}, + {name: "resource", mutate: func(access *AgentCompatCapabilityAccess) { access.ResourceID++ }}, + {name: "server", mutate: func(access *AgentCompatCapabilityAccess) { access.TargetServerID++ }}, + {name: "access proof", mutate: func(access *AgentCompatCapabilityAccess) { access.ServerAccessAllowed = false }}, + } + for _, mutation := range mutations { + t.Run(mutation.name, func(t *testing.T) { + handler, registration, capability := natCapabilityFixture(t, 36, 46, 58, 68) + access := capabilityAccess(capability, registration) + mutation.mutate(&access) + + _, err := handler.ConsumeAgentCompatNATCapability(access) + + require.True(t, errors.Is(err, ErrAgentCompatCapabilityHidden)) + }) + } +} + +func TestAgentCompatCapabilityCancelBeforeBindWakesWaiterAndPreventsBind(t *testing.T) { + handler := NewNezhaHandler() + registration := capabilityRegistration(capabilityOwner(37, 47), AgentCompatCapabilityTerminal, 59, 0) + capability, err := handler.RegisterAgentCompatIOStreamCapability(agentCompatCapabilityTestContext(t), registration) + require.NoError(t, err) + access := capabilityAccess(capability, registration) + started := make(chan struct{}) + var observed atomic.Bool + handler.setAgentCompatCapabilityWaitObserverForTest(func() { + if observed.CompareAndSwap(false, true) { + close(started) + } + }) + result := make(chan error, 1) + waitCtx := agentCompatCapabilityTestContext(t) + go func() { + _, waitErr := handler.WaitAgentCompatIOStreamCapability(waitCtx, access) + result <- waitErr + }() + awaitAgentCompatCapabilitySignal(t, started, "wait observer did not start") + + require.NoError(t, handler.CancelAgentCompatIOStreamCapability(access)) + require.ErrorIs(t, receiveAgentCompatCapabilityError(t, result, "canceled waiter did not return"), ErrAgentCompatCapabilityHidden) + require.NoError(t, handler.CreateStreamWithPurpose("after-cancel", 47, 59, PurposeTerminal)) + require.ErrorIs(t, handler.BindAgentCompatIOStreamCapability(AgentCompatCapabilityBinding{AgentCompatCapabilityAccess: access, StreamID: "after-cancel"}), ErrAgentCompatCapabilityHidden) +} diff --git a/service/rpc/io_stream_capability_state_agentcompat.go b/service/rpc/io_stream_capability_state_agentcompat.go new file mode 100644 index 00000000..fe2a4797 --- /dev/null +++ b/service/rpc/io_stream_capability_state_agentcompat.go @@ -0,0 +1,78 @@ +//go:build agentcompat + +package rpc + +import ( + "crypto/rand" +) + +type agentCompatCapabilityPhase uint8 + +const ( + agentCompatCapabilityRegistered agentCompatCapabilityPhase = iota + 1 + agentCompatCapabilityConsumed + agentCompatCapabilityPublished +) + +const ( + agentCompatCapabilityMaxActivePerPAT = 16 + agentCompatCapabilityMaxActiveGlobal = 128 + agentCompatCapabilityMaxProcessMints = 4096 +) + +type agentCompatCapabilityRegistration struct { + registration AgentCompatCapabilityRegistration + phase agentCompatCapabilityPhase + generation uint64 + streamID string + stream *ioStreamContext + notify chan struct{} +} + +type agentCompatCapabilityState struct { + active map[string]*agentCompatCapabilityRegistration + activeByPAT map[uint64]uint16 + // Used tokens are process-lifetime tombstones; deletion never makes a capability reusable. + used map[string]struct{} + tokenSource func([]byte) error + nextIdentity uint64 + waitObserver func() + publishObserver func() +} + +func (s *NezhaHandler) initializeAgentCompatCapabilities() { + s.agentCompatCapabilities.active = make(map[string]*agentCompatCapabilityRegistration) + s.agentCompatCapabilities.activeByPAT = make(map[uint64]uint16) + s.agentCompatCapabilities.used = make(map[string]struct{}) + s.agentCompatCapabilities.tokenSource = func(destination []byte) error { + _, err := rand.Read(destination) + return err + } +} + +func (s *NezhaHandler) setAgentCompatCapabilityTokenSourceForTest(source func([]byte) error) { + s.ioStreamMutex.Lock() + defer s.ioStreamMutex.Unlock() + s.agentCompatCapabilities.tokenSource = source +} + +func (s *NezhaHandler) setAgentCompatCapabilityWaitObserverForTest(observer func()) { + s.ioStreamMutex.Lock() + defer s.ioStreamMutex.Unlock() + s.agentCompatCapabilities.waitObserver = observer +} + +func (s *NezhaHandler) setAgentCompatCapabilityPublishObserverForTest(observer func()) { + s.ioStreamMutex.Lock() + defer s.ioStreamMutex.Unlock() + s.agentCompatCapabilities.publishObserver = observer +} + +func (s *NezhaHandler) SetAgentCompatCapabilityPublishObserverForTest(observer func()) { + s.setAgentCompatCapabilityPublishObserverForTest(observer) +} + +func (registration *agentCompatCapabilityRegistration) publishLocked() { + close(registration.notify) + registration.notify = make(chan struct{}) +} diff --git a/service/rpc/io_stream_capability_state_default.go b/service/rpc/io_stream_capability_state_default.go new file mode 100644 index 00000000..f1243812 --- /dev/null +++ b/service/rpc/io_stream_capability_state_default.go @@ -0,0 +1,9 @@ +//go:build !agentcompat + +package rpc + +type agentCompatCapabilityState struct{} + +type agentCompatCapabilityRegistration struct{} + +func (*NezhaHandler) initializeAgentCompatCapabilities() {} diff --git a/service/rpc/io_stream_capability_test_helpers_agentcompat_test.go b/service/rpc/io_stream_capability_test_helpers_agentcompat_test.go new file mode 100644 index 00000000..48404b95 --- /dev/null +++ b/service/rpc/io_stream_capability_test_helpers_agentcompat_test.go @@ -0,0 +1,38 @@ +//go:build agentcompat + +package rpc + +import ( + "context" + "testing" + "time" +) + +const agentCompatCapabilityTestTimeout = 5 * time.Second + +func agentCompatCapabilityTestContext(t *testing.T) context.Context { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), agentCompatCapabilityTestTimeout) + t.Cleanup(cancel) + return ctx +} + +func awaitAgentCompatCapabilitySignal(t *testing.T, signal <-chan struct{}, failureMessage string) { + t.Helper() + select { + case <-signal: + case <-agentCompatCapabilityTestContext(t).Done(): + t.Fatal(failureMessage) + } +} + +func receiveAgentCompatCapabilityError(t *testing.T, result <-chan error, failureMessage string) error { + t.Helper() + select { + case err := <-result: + return err + case <-agentCompatCapabilityTestContext(t).Done(): + t.Fatal(failureMessage) + return nil + } +} diff --git a/service/rpc/io_stream_capability_types.go b/service/rpc/io_stream_capability_types.go new file mode 100644 index 00000000..2593bccf --- /dev/null +++ b/service/rpc/io_stream_capability_types.go @@ -0,0 +1,93 @@ +package rpc + +import ( + "encoding/base64" + "errors" +) + +var ( + ErrAgentCompatCapabilityUnavailable = errors.New("agentcompat IOStream capability unavailable") + ErrAgentCompatCapabilityHidden = errors.New("agentcompat IOStream capability unavailable") + ErrAgentCompatCapabilityConflict = errors.New("agentcompat IOStream capability conflict") + ErrAgentCompatCapabilityBound = errors.New("agentcompat IOStream capability has a live bound stream") + ErrAgentCompatCapabilityTokenExhausted = errors.New("agentcompat IOStream capability token attempts exhausted") +) + +type AgentCompatCapabilityPurpose uint8 + +const ( + AgentCompatCapabilityTerminal AgentCompatCapabilityPurpose = iota + 1 + AgentCompatCapabilityFileManager + AgentCompatCapabilityNAT +) + +func (purpose AgentCompatCapabilityPurpose) streamPurpose() StreamPurpose { + switch purpose { + case AgentCompatCapabilityTerminal: + return PurposeTerminal + case AgentCompatCapabilityFileManager: + return PurposeFileManager + case AgentCompatCapabilityNAT: + return PurposeNAT + default: + return PurposeLegacy + } +} + +type AgentCompatCapabilityOwner struct { + PATID uint64 + UserID uint64 + IsAdmin bool +} + +type AgentCompatCapabilityRegistration struct { + Owner AgentCompatCapabilityOwner + Purpose AgentCompatCapabilityPurpose + TargetServerID uint64 + ResourceID uint64 + ServerAccessAllowed bool +} + +type AgentCompatIOStreamCapability struct{ value string } + +func (capability AgentCompatIOStreamCapability) String() string { return capability.value } + +func ParseAgentCompatIOStreamCapability(value string) (AgentCompatIOStreamCapability, error) { + raw, err := base64.RawURLEncoding.DecodeString(value) + if err != nil || len(raw) != 32 { + return AgentCompatIOStreamCapability{}, ErrAgentCompatCapabilityHidden + } + return AgentCompatIOStreamCapability{value: value}, nil +} + +type AgentCompatCapabilityAccess struct { + Capability AgentCompatIOStreamCapability + Owner AgentCompatCapabilityOwner + Purpose AgentCompatCapabilityPurpose + TargetServerID uint64 + ResourceID uint64 + ServerAccessAllowed bool +} + +type AgentCompatCapabilityBinding struct { + AgentCompatCapabilityAccess + StreamID string +} + +type AgentCompatNATPublishHandle struct { + registration *agentCompatCapabilityRegistration + generation uint64 + capability string +} + +type AgentCompatNATStreamLease struct { + streamID string + stream *ioStreamContext +} + +type AgentCompatNATPublication struct { + Purpose AgentCompatCapabilityPurpose + TargetServerID uint64 + ResourceID uint64 + StreamID string +} diff --git a/service/rpc/io_stream_quota_agentcompat.go b/service/rpc/io_stream_quota_agentcompat.go new file mode 100644 index 00000000..8966292d --- /dev/null +++ b/service/rpc/io_stream_quota_agentcompat.go @@ -0,0 +1,144 @@ +//go:build agentcompat + +package rpc + +import ( + "context" + "errors" + "fmt" + "time" +) + +type IOStreamQuotaProbeResult struct { + UserAccepted int + UserRejected int + ServerAccepted int + ServerRejected int + TrackedStreams int + WaitForAgentWokeOnClose bool + UserSlotReused bool + UserBoundaryError error + ServerBoundaryError error + Err error +} + +func RunIOStreamQuotaProbe(ctx context.Context) IOStreamQuotaProbeResult { + if err := ctx.Err(); err != nil { + return IOStreamQuotaProbeResult{Err: err} + } + h := NewNezhaHandler() + result := IOStreamQuotaProbeResult{} + defer func() { + h.ioStreamMutex.RLock() + streamIDs := make([]string, 0, len(h.ioStreams)) + for streamID := range h.ioStreams { + streamIDs = append(streamIDs, streamID) + } + h.ioStreamMutex.RUnlock() + for _, streamID := range streamIDs { + _ = h.CloseStream(streamID) + } + }() + for i := 0; i < maxStreamsPerUser; i++ { + if err := h.CreateStream(fmt.Sprintf("probe-user-%d", i), 101, uint64(i+1)); err != nil { + result.Err = fmt.Errorf("create user stream %d: %w", i, err) + return result + } + result.UserAccepted++ + } + result.UserBoundaryError = h.CreateStream("probe-user-over", 101, 500) + if !errors.Is(result.UserBoundaryError, ErrTooManyStreamsForUser) { + result.Err = fmt.Errorf("user boundary returned %v", result.UserBoundaryError) + return result + } + result.UserRejected = 1 + + for i := 0; i < maxStreamsPerServer; i++ { + if err := h.CreateStream(fmt.Sprintf("probe-server-%d", i), uint64(i+1000), 700); err != nil { + result.Err = fmt.Errorf("create server stream %d: %w", i, err) + return result + } + result.ServerAccepted++ + } + result.ServerBoundaryError = h.CreateStream("probe-server-over", 2000, 700) + if !errors.Is(result.ServerBoundaryError, ErrTooManyStreamsForServer) { + result.Err = fmt.Errorf("server boundary returned %v", result.ServerBoundaryError) + return result + } + result.ServerRejected = 1 + + if err := h.CloseStream("probe-user-0"); err != nil { + result.Err = fmt.Errorf("close stale user slot: %w", err) + return result + } + if err := h.CreateStream("probe-user-reused", 101, 501); err != nil { + result.Err = fmt.Errorf("reuse stale user slot: %w", err) + return result + } + result.UserSlotReused = true + + if err := h.CreateStream("probe-wait", 0, 502); err != nil { + result.Err = fmt.Errorf("create cancellation probe stream: %w", err) + return result + } + waitStream, err := h.GetStream("probe-wait") + if err != nil { + result.Err = fmt.Errorf("get cancellation probe stream: %w", err) + return result + } + waitResult := make(chan bool, 1) + go func() { + _, ok := h.WaitForAgent(ctx, "probe-wait", 30*time.Second) + waitResult <- ok + }() + select { + case <-waitStream.waitStartedCh: + case <-ctx.Done(): + result.Err = ctx.Err() + return result + } + if err := h.CloseStream("probe-wait"); err != nil { + result.Err = fmt.Errorf("close cancellation probe stream: %w", err) + return result + } + select { + case ok := <-waitResult: + result.WaitForAgentWokeOnClose = !ok + case <-ctx.Done(): + result.Err = ctx.Err() + return result + } + if !result.WaitForAgentWokeOnClose { + result.Err = errors.New("WaitForAgent did not wake after stream close") + return result + } + + for i := 0; i < maxStreamsPerUser; i++ { + if err := h.CloseStream(fmt.Sprintf("probe-user-%d", i)); err != nil { + result.Err = fmt.Errorf("close user stream %d: %w", i, err) + return result + } + if err := h.CloseStream(fmt.Sprintf("probe-user-%d", i)); err != nil { + result.Err = fmt.Errorf("repeat close user stream %d: %w", i, err) + return result + } + } + if err := h.CloseStream("probe-user-reused"); err != nil { + result.Err = fmt.Errorf("close reused user slot: %w", err) + return result + } + for i := 0; i < maxStreamsPerServer; i++ { + if err := h.CloseStream(fmt.Sprintf("probe-server-%d", i)); err != nil { + result.Err = fmt.Errorf("close server stream %d: %w", i, err) + return result + } + } + if err := ctx.Err(); err != nil { + result.Err = err + return result + } + h.ioStreamMutex.RLock() + result.TrackedStreams = len(h.ioStreams) + h.ioStreamMutex.RUnlock() + return result +} diff --git a/service/rpc/io_stream_quota_agentcompat_test.go b/service/rpc/io_stream_quota_agentcompat_test.go new file mode 100644 index 00000000..84060e0f --- /dev/null +++ b/service/rpc/io_stream_quota_agentcompat_test.go @@ -0,0 +1,77 @@ +//go:build agentcompat + +package rpc + +import ( + "context" + "errors" + "fmt" + "sync" + "testing" +) + +func TestAgentCompatIOStreamQuotaProbe(t *testing.T) { + result := RunIOStreamQuotaProbe(context.Background()) + if result.Err != nil { + t.Fatalf("quota probe failed: %v", result.Err) + } + if result.UserAccepted != maxStreamsPerUser || result.UserRejected != 1 { + t.Fatalf("unexpected user boundary counts: accepted=%d rejected=%d", result.UserAccepted, result.UserRejected) + } + if result.ServerAccepted != maxStreamsPerServer || result.ServerRejected != 1 { + t.Fatalf("unexpected server boundary counts: accepted=%d rejected=%d", result.ServerAccepted, result.ServerRejected) + } + if result.TrackedStreams != 0 { + t.Fatalf("probe left tracked streams: %d", result.TrackedStreams) + } + if !result.WaitForAgentWokeOnClose { + t.Fatal("probe did not prove WaitForAgent wakes after real stream close") + } + if !result.UserSlotReused { + t.Fatal("probe did not prove a released user slot was reusable") + } +} + +func TestAgentCompatIOStreamQuotaProbeUsesProductionSeam(t *testing.T) { + result := RunIOStreamQuotaProbe(context.Background()) + if !errors.Is(result.UserBoundaryError, ErrTooManyStreamsForUser) { + t.Fatalf("user rejection must preserve production error, got %v", result.UserBoundaryError) + } + if !errors.Is(result.ServerBoundaryError, ErrTooManyStreamsForServer) { + t.Fatalf("server rejection must preserve production error, got %v", result.ServerBoundaryError) + } +} + +func TestAgentCompatIOStreamQuotaProbeConcurrentBoundaryCalls(t *testing.T) { + h := NewNezhaHandler() + const userID, serverID = uint64(701), uint64(901) + var wg sync.WaitGroup + results := make(chan error, maxStreamsPerUser+1) + for i := 0; i < maxStreamsPerUser+1; i++ { + wg.Add(1) + go func(index int) { + defer wg.Done() + results <- h.CreateStream(fmt.Sprintf("concurrent-user-%d", index), userID, serverID+uint64(index)) + }(i) + } + wg.Wait() + close(results) + accepted, rejected := 0, 0 + for err := range results { + if err == nil { + accepted++ + continue + } + if errors.Is(err, ErrTooManyStreamsForUser) { + rejected++ + continue + } + t.Fatalf("unexpected concurrent boundary error: %v", err) + } + if accepted != maxStreamsPerUser || rejected != 1 { + t.Fatalf("unexpected concurrent boundary counts: accepted=%d rejected=%d", accepted, rejected) + } + for i := 0; i < maxStreamsPerUser+1; i++ { + _ = h.CloseStream(fmt.Sprintf("concurrent-user-%d", i)) + } +} diff --git a/service/rpc/io_stream_state_agentcompat.go b/service/rpc/io_stream_state_agentcompat.go new file mode 100644 index 00000000..65bb719e --- /dev/null +++ b/service/rpc/io_stream_state_agentcompat.go @@ -0,0 +1,11 @@ +//go:build agentcompat + +package rpc + +// SetIOStreamStateWaitObserverForAgentcompat installs a deterministic harness +// seam for observing that a waiter captured its notification channel. +func (s *NezhaHandler) SetIOStreamStateWaitObserverForAgentcompat(observer func()) { + s.ioStreamMutex.Lock() + defer s.ioStreamMutex.Unlock() + s.ioStreamWaitLockedHook = observer +} diff --git a/service/rpc/mcp_receipt_agentcompat_test.go b/service/rpc/mcp_receipt_agentcompat_test.go new file mode 100644 index 00000000..c8dde517 --- /dev/null +++ b/service/rpc/mcp_receipt_agentcompat_test.go @@ -0,0 +1,197 @@ +//go:build agentcompat + +package rpc + +import ( + "bufio" + "context" + "errors" + "fmt" + "net" + "strconv" + "strings" + "testing" + "time" + + "github.com/nezhahq/nezha/model" + pb "github.com/nezhahq/nezha/proto" + "github.com/stretchr/testify/require" +) + +func TestMCPReceiptGate_FormatsTaskAndResultWithGeneration(t *testing.T) { + // Given + serverConn, clientConn := net.Pipe() + defer clientConn.Close() + gate := installReceiptGateForTest(serverConn) + defer clearReceiptGateForTest() + reader := bufio.NewReader(clientConn) + + // When + taskDone := make(chan struct{}) + go func() { + notifyMCPTaskDispatched(7, 9, model.TaskTypeExec) + close(taskDone) + }() + taskLine := mustReadLine(t, reader) + <-taskDone + resultDone := make(chan struct{}) + go func() { + notifyMCPTaskResultAccepted(7, 9, model.TaskTypeExec) + close(resultDone) + }() + resultLine := mustReadLine(t, reader) + <-resultDone + + // Then + require.Equal(t, "task "+itoa(gate.generation)+" 7 9 "+itoa(model.TaskTypeExec)+"\n", taskLine) + require.Equal(t, "result "+itoa(gate.generation)+" 7 9 "+itoa(model.TaskTypeExec)+"\n", resultLine) +} + +func TestCallAgent_EmitsOneTaskAndOneAcceptedResult(t *testing.T) { + // Given + serverConn, clientConn := net.Pipe() + defer clientConn.Close() + gate := installReceiptGateForTest(serverConn) + defer clearReceiptGateForTest() + stream := newFakeStream() + cleanup := installFakeServer(t, 801, stream) + defer cleanup() + reader := bufio.NewReader(clientConn) + lines := make(chan string, 2) + go func() { + line, err := reader.ReadString('\n') + if err != nil { + return + } + lines <- line + line, err = reader.ReadString('\n') + if err == nil { + lines <- line + } + }() + + go func() { + sent := <-stream.sent + deliverMCPResult(&pb.TaskResult{Id: sent.GetId(), Type: sent.GetType(), Successful: true, Data: "{}"}) + deliverMCPResult(&pb.TaskResult{Id: sent.GetId(), Type: sent.GetType(), Successful: true, Data: "{}"}) + }() + + // When + _, err := CallAgent(context.Background(), 801, model.TaskTypeExec, model.ExecRequest{Cmd: "x"}, time.Second) + + // Then + require.NoError(t, err) + taskLine := <-lines + resultLine := <-lines + require.Equal(t, "task "+itoa(gate.generation)+" 801 "+itoa(parseReceiptTaskID(taskLine))+" "+itoa(model.TaskTypeExec)+"\n", taskLine) + require.Equal(t, "result "+itoa(gate.generation)+" 801 "+itoa(parseReceiptTaskID(resultLine))+" "+itoa(model.TaskTypeExec)+"\n", resultLine) + require.Equal(t, parseReceiptTaskID(taskLine), parseReceiptTaskID(resultLine)) + clientConn.SetReadDeadline(time.Now().Add(20 * time.Millisecond)) + _, readErr := reader.ReadString('\n') + require.Error(t, readErr) +} + +func TestCallAgent_SendFailureEmitsNoTaskReceipt(t *testing.T) { + // Given + serverConn, clientConn := net.Pipe() + defer clientConn.Close() + installReceiptGateForTest(serverConn) + defer clearReceiptGateForTest() + stream := &fakeTaskStream{sent: make(chan *pb.Task, 1), err: errors.New("send failed")} + cleanup := installFakeServer(t, 802, stream) + defer cleanup() + + // When + _, err := CallAgent(context.Background(), 802, model.TaskTypeExec, model.ExecRequest{Cmd: "x"}, time.Second) + + // Then + require.EqualError(t, err, "send failed") + clientConn.SetReadDeadline(time.Now().Add(20 * time.Millisecond)) + _, readErr := bufio.NewReader(clientConn).ReadString('\n') + require.Error(t, readErr) +} + +func TestCallAgent_LateDuplicateAndCancelledResultsEmitNoAcceptedReceipt(t *testing.T) { + // Given + serverConn, clientConn := net.Pipe() + defer clientConn.Close() + installReceiptGateForTest(serverConn) + defer clearReceiptGateForTest() + stream := newFakeStream() + cleanup := installFakeServer(t, 803, stream) + defer cleanup() + reader := bufio.NewReader(clientConn) + taskLineCh := make(chan string, 1) + + // When + taskIDCh := make(chan uint64, 1) + go func() { + sent := <-stream.sent + taskIDCh <- sent.GetId() + line, _ := reader.ReadString('\n') + taskLineCh <- line + }() + _, err := CallAgent(context.Background(), 803, model.TaskTypeFsRead, model.FsReadRequest{Path: "/x"}, 20*time.Millisecond) + require.ErrorIs(t, err, ErrAgentTimeout) + taskID := <-taskIDCh + deliverMCPResult(&pb.TaskResult{Id: taskID, Type: model.TaskTypeFsRead, Successful: true, Data: "{}"}) + deliverMCPResult(&pb.TaskResult{Id: taskID, Type: model.TaskTypeFsRead, Successful: true, Data: "{}"}) + + // Then + require.Contains(t, <-taskLineCh, "task ") + clientConn.SetReadDeadline(time.Now().Add(20 * time.Millisecond)) + _, readErr := reader.ReadString('\n') + require.Error(t, readErr) +} + +func TestCallAgent_CancelledResultEmitsNoAcceptedReceipt(t *testing.T) { + // Given + serverConn, clientConn := net.Pipe() + defer clientConn.Close() + installReceiptGateForTest(serverConn) + defer clearReceiptGateForTest() + stream := newFakeStream() + cleanup := installFakeServer(t, 804, stream) + defer cleanup() + reader := bufio.NewReader(clientConn) + taskLineCh := make(chan string, 1) + go func() { + line, _ := reader.ReadString('\n') + taskLineCh <- line + }() + taskIDCh := make(chan uint64, 1) + go func() { + sent := <-stream.sent + taskIDCh <- sent.GetId() + }() + + // When + errCh := make(chan error, 1) + go func() { + _, err := CallAgent(context.Background(), 804, model.TaskTypeExec, model.ExecRequest{Cmd: "x"}, time.Second) + errCh <- err + }() + taskID := <-taskIDCh + CancelAllMCPInflight() + deliverMCPResult(&pb.TaskResult{Id: taskID, Type: model.TaskTypeExec, Successful: true, Data: "{}"}) + + // Then + require.ErrorIs(t, <-errCh, ErrMCPDisabled) + require.Contains(t, <-taskLineCh, "task ") + clientConn.SetReadDeadline(time.Now().Add(20 * time.Millisecond)) + _, readErr := reader.ReadString('\n') + require.Error(t, readErr) +} + +func itoa(value uint64) string { + return strconv.FormatUint(value, 10) +} + +func parseReceiptTaskID(line string) uint64 { + fields := strings.Fields(line) + value, err := strconv.ParseUint(fields[3], 10, 64) + if err != nil { + panic(fmt.Sprintf("invalid receipt line %q: %v", line, err)) + } + return value +} diff --git a/service/rpc/receipt_gate_agentcompat.go b/service/rpc/receipt_gate_agentcompat.go new file mode 100644 index 00000000..0ba57a86 --- /dev/null +++ b/service/rpc/receipt_gate_agentcompat.go @@ -0,0 +1,251 @@ +//go:build agentcompat + +package rpc + +import ( + "bufio" + "context" + "errors" + "fmt" + "net" + "strings" + "sync" + "time" +) + +const receiptGateCommandTimeout = 30 * time.Second + +type receiptGate struct { + conn net.Conn + read *bufio.Reader + generation uint64 + stateMu sync.Mutex + ioMu sync.Mutex + closeOnce sync.Once + context context.Context + cancel context.CancelFunc + hold bool + acceptedCount uint64 +} + +var activeReceiptGate *receiptGate +var activeReceiptGateMu sync.RWMutex +var receiptGateListener net.Listener +var receiptGateGeneration uint64 +var receiptGateCancel context.CancelFunc +var receiptGateWaitGroup sync.WaitGroup + +func newReceiptGate(conn net.Conn, generation uint64) *receiptGate { + ctx, cancel := context.WithCancel(context.Background()) + return &receiptGate{conn: conn, read: bufio.NewReader(conn), generation: generation, context: ctx, cancel: cancel, hold: true} +} + +func SetReceiptGateListener(listener net.Listener) { + if listener == nil { + return + } + activeReceiptGateMu.Lock() + previousListener := receiptGateListener + previousCancel := receiptGateCancel + receiptGateListener = listener + listenerContext, cancel := context.WithCancel(context.Background()) + receiptGateCancel = cancel + activeReceiptGateMu.Unlock() + if previousCancel != nil { + previousCancel() + } + if previousListener != nil { + _ = previousListener.Close() + } + receiptGateWaitGroup.Add(1) + go acceptReceiptGateConnections(listenerContext, listener) +} + +func acceptReceiptGateConnections(ctx context.Context, listener net.Listener) { + defer receiptGateWaitGroup.Done() + for { + connection, err := listener.Accept() + if err != nil { + select { + case <-ctx.Done(): + return + default: + } + return + } + activeReceiptGateMu.Lock() + receiptGateGeneration++ + generation := receiptGateGeneration + previous := activeReceiptGate + gate := newReceiptGate(connection, generation) + activeReceiptGate = gate + activeReceiptGateMu.Unlock() + if previous != nil { + previous.close() + } + if err := connection.SetWriteDeadline(time.Now().Add(receiptGateCommandTimeout)); err != nil { + resetReceiptGate(gate) + continue + } + if _, err := fmt.Fprintln(connection, "ready"); err != nil { + resetReceiptGate(gate) + continue + } + _ = connection.SetWriteDeadline(time.Time{}) + } +} + +func (gate *receiptGate) close() { + gate.closeOnce.Do(func() { + gate.cancel() + _ = gate.conn.Close() + }) +} + +func CloseReceiptGate() { + activeReceiptGateMu.Lock() + listener := receiptGateListener + cancel := receiptGateCancel + gate := activeReceiptGate + receiptGateListener = nil + receiptGateCancel = nil + activeReceiptGate = nil + activeReceiptGateMu.Unlock() + if cancel != nil { + cancel() + } + if listener != nil { + _ = listener.Close() + } + if gate != nil { + gate.close() + } + receiptGateWaitGroup.Wait() +} + +func resetReceiptGate(gate *receiptGate) { + activeReceiptGateMu.Lock() + if activeReceiptGate == gate { + activeReceiptGate = nil + } + activeReceiptGateMu.Unlock() + gate.close() +} + +func currentReceiptGate() *receiptGate { + activeReceiptGateMu.RLock() + defer activeReceiptGateMu.RUnlock() + return activeReceiptGate +} + +func (gate *receiptGate) sendAccepted(serverID uint64, uuid string, generation, count uint64) error { + gate.stateMu.Lock() + gate.acceptedCount++ + count = gate.acceptedCount + hold := gate.hold + gate.stateMu.Unlock() + gate.ioMu.Lock() + defer gate.ioMu.Unlock() + if err := gate.conn.SetDeadline(time.Now().Add(receiptGateCommandTimeout)); err != nil { + resetReceiptGate(gate) + return err + } + if _, err := fmt.Fprintf(gate.conn, "accepted %d %s %d %d %d\n", serverID, uuid, gate.generation, generation, count); err != nil { + resetReceiptGate(gate) + return err + } + if !hold { + _ = gate.conn.SetDeadline(time.Time{}) + return nil + } + command, err := gate.read.ReadString('\n') + if err != nil { + resetReceiptGate(gate) + return err + } + if strings.TrimSpace(command) != "release" { + err := errors.New("receipt gate received unexpected command") + resetReceiptGate(gate) + return err + } + gate.stateMu.Lock() + gate.hold = false + gate.stateMu.Unlock() + if err := gate.conn.SetDeadline(time.Time{}); err != nil { + resetReceiptGate(gate) + return err + } + return nil +} + +func notifyReceiptAccepted(serverID uint64, uuid string, generation, count uint64) error { + gate := currentReceiptGate() + if gate == nil { + return nil + } + return gate.sendAccepted(serverID, uuid, generation, count) +} + +func notifyStateReceived(serverID uint64, uuid string, generation, count uint64) error { + gate := currentReceiptGate() + if gate == nil { + return nil + } + gate.ioMu.Lock() + defer gate.ioMu.Unlock() + if err := gate.conn.SetWriteDeadline(time.Now().Add(receiptGateCommandTimeout)); err != nil { + resetReceiptGate(gate) + return err + } + if _, err := fmt.Fprintf(gate.conn, "state %d %s %d %d\n", serverID, uuid, generation, count); err != nil { + resetReceiptGate(gate) + return err + } + return gate.conn.SetWriteDeadline(time.Time{}) +} + +func notifyInfo2(serverID uint64, uuid string) error { + gate := currentReceiptGate() + if gate == nil { + return nil + } + gate.ioMu.Lock() + defer gate.ioMu.Unlock() + if err := gate.conn.SetWriteDeadline(time.Now().Add(receiptGateCommandTimeout)); err != nil { + resetReceiptGate(gate) + return err + } + if _, err := fmt.Fprintf(gate.conn, "info2 %d %d %s\n", gate.generation, serverID, uuid); err != nil { + resetReceiptGate(gate) + return err + } + return gate.conn.SetWriteDeadline(time.Time{}) +} + +func notifyMCPTaskDispatched(serverID, taskID, taskType uint64) { + notifyMCPReceipt("task", serverID, taskID, taskType) +} + +func notifyMCPTaskResultAccepted(serverID, taskID, taskType uint64) { + notifyMCPReceipt("result", serverID, taskID, taskType) +} + +func notifyMCPReceipt(kind string, serverID, taskID, taskType uint64) { + gate := currentReceiptGate() + if gate == nil { + return + } + gate.ioMu.Lock() + defer gate.ioMu.Unlock() + if err := gate.conn.SetWriteDeadline(time.Now().Add(receiptGateCommandTimeout)); err != nil { + resetReceiptGate(gate) + return + } + if _, err := fmt.Fprintf(gate.conn, "%s %d %d %d %d\n", kind, gate.generation, serverID, taskID, taskType); err != nil { + resetReceiptGate(gate) + return + } + if err := gate.conn.SetWriteDeadline(time.Time{}); err != nil { + resetReceiptGate(gate) + } +} diff --git a/service/rpc/receipt_gate_agentcompat_test.go b/service/rpc/receipt_gate_agentcompat_test.go new file mode 100644 index 00000000..00402286 --- /dev/null +++ b/service/rpc/receipt_gate_agentcompat_test.go @@ -0,0 +1,224 @@ +//go:build agentcompat + +package rpc + +import ( + "bufio" + "fmt" + "net" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func installReceiptGateForTest(conn net.Conn) *receiptGate { + activeReceiptGateMu.Lock() + receiptGateGeneration++ + generation := receiptGateGeneration + activeReceiptGateMu.Unlock() + gate := newReceiptGate(conn, generation) + activeReceiptGateMu.Lock() + activeReceiptGate = gate + activeReceiptGateMu.Unlock() + return gate +} + +func clearReceiptGateForTest() { + activeReceiptGateMu.Lock() + gate := activeReceiptGate + activeReceiptGate = nil + activeReceiptGateMu.Unlock() + if gate != nil { + gate.close() + } +} + +func TestReceiptGate_EOFResetsGate(t *testing.T) { + // Given + serverConn, clientConn := net.Pipe() + defer clientConn.Close() + installReceiptGateForTest(serverConn) + defer clearReceiptGateForTest() + gate := currentReceiptGate() + require.NotNil(t, gate) + go func() { + reader := bufio.NewReader(clientConn) + _, _ = reader.ReadString('\n') + _ = clientConn.Close() + }() + + // When + err := notifyReceiptAccepted(7, "uuid", 1, 1) + + // Then + require.Error(t, err) + activeReceiptGateMu.RLock() + active := activeReceiptGate + activeReceiptGateMu.RUnlock() + require.Nil(t, active) +} + +func TestReceiptGate_ListenerAcceptsAndReplacesConnections(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer CloseReceiptGate() + SetReceiptGateListener(listener) + + oldClient, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + oldReader := bufio.NewReader(oldClient) + require.Equal(t, "ready\n", mustReadLine(t, oldReader)) + + newClient, err := net.Dial("tcp", listener.Addr().String()) + require.NoError(t, err) + defer newClient.Close() + newReader := bufio.NewReader(newClient) + require.Equal(t, "ready\n", mustReadLine(t, newReader)) + _ = oldClient.SetReadDeadline(time.Now().Add(time.Second)) + _, oldErr := oldReader.ReadString('\n') + require.Error(t, oldErr) +} + +func TestReceiptGate_CloseInterruptsHeldReadAndQueuedWrite(t *testing.T) { + // Given + serverConn, clientConn := net.Pipe() + defer clientConn.Close() + installReceiptGateForTest(serverConn) + acceptedStarted := make(chan struct{}) + acceptedDone := make(chan error, 1) + go func() { + close(acceptedStarted) + acceptedDone <- notifyReceiptAccepted(7, "uuid", 1, 1) + }() + <-acceptedStarted + reader := bufio.NewReader(clientConn) + require.Equal(t, "accepted 7 uuid "+fmt.Sprint(currentReceiptGate().generation)+" 1 1\n", mustReadLine(t, reader)) + + infoStarted := make(chan struct{}) + infoDone := make(chan error, 1) + go func() { + close(infoStarted) + infoDone <- notifyInfo2(9, "held") + }() + <-infoStarted + + // When + CloseReceiptGate() + + // Then + select { + case <-acceptedDone: + case <-time.After(time.Second): + t.Fatal("held receipt read was not interrupted") + } + select { + case <-infoDone: + case <-time.After(time.Second): + t.Fatal("queued notification write was not released") + } +} + +func mustReadLine(t *testing.T, reader *bufio.Reader) string { + t.Helper() + line, err := reader.ReadString('\n') + require.NoError(t, err) + return line +} + +func TestReceiptGate_MalformedCommandResetsGate(t *testing.T) { + // Given + serverConn, clientConn := net.Pipe() + defer clientConn.Close() + installReceiptGateForTest(serverConn) + defer clearReceiptGateForTest() + go func() { + reader := bufio.NewReader(clientConn) + _, _ = reader.ReadString('\n') + _, _ = clientConn.Write([]byte("hold\n")) + }() + + // When + err := notifyReceiptAccepted(7, "uuid", 1, 1) + + // Then + require.EqualError(t, err, "receipt gate received unexpected command") + activeReceiptGateMu.RLock() + active := activeReceiptGate + activeReceiptGateMu.RUnlock() + require.Nil(t, active) +} + +func TestReceiptGate_ReplacementClosesOldConnection(t *testing.T) { + t.Run("replacement closes old connection", func(t *testing.T) { + // Given + oldServer, oldClient := net.Pipe() + newServer, newClient := net.Pipe() + t.Cleanup(func() { require.NoError(t, oldClient.Close()) }) + t.Cleanup(func() { require.NoError(t, newClient.Close()) }) + oldGate := installReceiptGateForTest(oldServer) + t.Cleanup(oldGate.close) + t.Cleanup(clearReceiptGateForTest) + newGate := newReceiptGate(newServer, oldGate.generation+1) + activeReceiptGateMu.Lock() + activeReceiptGate = newGate + activeReceiptGateMu.Unlock() + oldDone := make(chan error, 1) + go func() { oldDone <- oldGate.sendAccepted(7, "uuid", 1, 1) }() + reader := bufio.NewReader(oldClient) + _, _ = reader.ReadString('\n') + + // When + oldGate.close() + + // Then + select { + case err := <-oldDone: + require.Error(t, err) + case <-time.After(time.Second): + t.Fatal("old receipt gate remained blocked after replacement") + } + }) + + require.Nil(t, currentReceiptGate()) +} + +func TestReceiptGate_Info2AndReceiptNotificationsSerialize(t *testing.T) { + // Given + serverConn, clientConn := net.Pipe() + defer clientConn.Close() + installReceiptGateForTest(serverConn) + defer clearReceiptGateForTest() + gate := currentReceiptGate() + require.NotNil(t, gate) + lines := make(chan string, 2) + go func() { + reader := bufio.NewReader(clientConn) + line, err := reader.ReadString('\n') + if err != nil { + return + } + lines <- strings.TrimSpace(line) + _, _ = clientConn.Write([]byte("release\n")) + line, err = reader.ReadString('\n') + if err == nil { + lines <- strings.TrimSpace(line) + } + }() + + // When + acceptedDone := make(chan error, 1) + go func() { acceptedDone <- notifyReceiptAccepted(7, "uuid", 1, 1) }() + select { + case err := <-acceptedDone: + require.NoError(t, err) + case <-time.After(time.Second): + require.NoError(t, <-acceptedDone) + } + require.NoError(t, notifyInfo2(7, "uuid")) + + // Then + require.Equal(t, "accepted 7 uuid "+fmt.Sprint(gate.generation)+" 1 1", <-lines) + require.Equal(t, "info2 "+fmt.Sprint(gate.generation)+" 7 uuid", <-lines) +} diff --git a/service/rpc/receipt_gate_default.go b/service/rpc/receipt_gate_default.go new file mode 100644 index 00000000..c6476201 --- /dev/null +++ b/service/rpc/receipt_gate_default.go @@ -0,0 +1,19 @@ +//go:build !agentcompat + +package rpc + +import "net" + +func SetReceiptGateListener(net.Listener) {} + +func CloseReceiptGate() {} + +func notifyReceiptAccepted(uint64, string, uint64, uint64) error { return nil } + +func notifyStateReceived(uint64, string, uint64, uint64) error { return nil } + +func notifyInfo2(uint64, string) error { return nil } + +func notifyMCPTaskDispatched(uint64, uint64, uint64) {} + +func notifyMCPTaskResultAccepted(uint64, uint64, uint64) {}