package notify import ( "b612.me/stario" "context" "io" "math" "net" "sync" "sync/atomic" "testing" "time" ) func TestStreamRuntimeSeparatesPeerDataIDNamespaces(t *testing.T) { clientRuntime := newStreamRuntime("cstrm") serverRuntime := newStreamRuntime("sstrm") clientID := clientRuntime.nextDataID() serverID := serverRuntime.nextDataID() if clientID == serverID || clientID%2 != 1 || serverID%2 != 0 { t.Fatalf("client/server stream data ids = %d/%d, want disjoint odd/even namespaces", clientID, serverID) } } func TestStreamRuntimeKeepsLegacyZeroDataIDControlsCompatible(t *testing.T) { tests := []struct { name string rolePrefix string wantParity uint64 }{ {name: "client receives server stream", rolePrefix: "cstrm", wantParity: 0}, {name: "server receives client stream", rolePrefix: "sstrm", wantParity: 1}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { runtime := newStreamRuntime(tt.rolePrefix) scope := "legacy-zero" stream := newStreamHandle(context.Background(), runtime, scope, StreamOpenRequest{ StreamID: "legacy-stream", DataID: 0, }, 0, nil, nil, 0, nil, nil, nil, runtime.configSnapshot()) if err := runtime.registerInbound(scope, stream); err != nil { t.Fatalf("register legacy zero-DataID stream: %v", err) } defer stream.markReset(io.ErrClosedPipe) if dataID := stream.dataIDSnapshot(); dataID == 0 || dataID%2 != tt.wantParity { t.Fatalf("legacy assigned DataID = %d, want non-zero parity %d", dataID, tt.wantParity) } if got, ok := runtime.lookupControl(scope, stream.ID(), 0); !ok || got != stream { t.Fatalf("legacy DataID=0 control lookup = %p/%v, want %p/true", got, ok, stream) } }) } } func TestStreamOpenConcurrentlyFromBothPeers(t *testing.T) { server := NewServer().(*ServerCommon) if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { t.Fatalf("UseModernPSKServer failed: %v", err) } serverAccepted := make(chan StreamAcceptInfo, 1) server.SetStreamHandler(func(info StreamAcceptInfo) error { serverAccepted <- info return nil }) if err := server.Listen("tcp", "127.0.0.1:0"); err != nil { t.Fatalf("server Listen failed: %v", err) } defer func() { _ = server.Stop() }() client := NewClient().(*ClientCommon) if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil { t.Fatalf("UseModernPSKClient failed: %v", err) } clientAccepted := make(chan StreamAcceptInfo, 1) client.SetStreamHandler(func(info StreamAcceptInfo) error { clientAccepted <- info return nil }) if err := client.Connect("tcp", server.listener.Addr().String()); err != nil { t.Fatalf("client Connect failed: %v", err) } defer func() { _ = client.Stop() }() logical := waitForTransferControlLogicalConn(t, server, 2*time.Second) type openResult struct { stream Stream err error } clientResult := make(chan openResult, 1) serverResult := make(chan openResult, 1) start := make(chan struct{}) var ready sync.WaitGroup ready.Add(2) go func() { ready.Done() <-start stream, err := client.OpenStream(context.Background(), StreamOpenOptions{Channel: StreamDataChannel}) clientResult <- openResult{stream: stream, err: err} }() go func() { ready.Done() <-start stream, err := server.OpenStreamLogical(context.Background(), logical, StreamOpenOptions{Channel: StreamDataChannel}) serverResult <- openResult{stream: stream, err: err} }() ready.Wait() close(start) clientOpen := <-clientResult serverOpen := <-serverResult if clientOpen.err != nil || serverOpen.err != nil { t.Fatalf("concurrent stream opens failed: client=%v server=%v", clientOpen.err, serverOpen.err) } clientInbound := waitAcceptedStream(t, serverAccepted, 2*time.Second) serverInbound := waitAcceptedStream(t, clientAccepted, 2*time.Second) clientID := clientOpen.stream.(*streamHandle).dataIDSnapshot() serverID := serverOpen.stream.(*streamHandle).dataIDSnapshot() if clientID%2 != 1 || serverID%2 != 0 { t.Fatalf("local stream data ids = %d/%d, want odd/even", clientID, serverID) } if clientInbound.DataID != clientID || serverInbound.DataID != serverID { t.Fatalf("stream data ids differ across peers: client=%d/%d server=%d/%d", clientID, clientInbound.DataID, serverID, serverInbound.DataID) } if _, err := clientOpen.stream.Write([]byte("client")); err != nil { t.Fatalf("client stream write failed: %v", err) } readStreamExactly(t, clientInbound.Stream, "client", 2*time.Second) if _, err := serverOpen.stream.Write([]byte("server")); err != nil { t.Fatalf("server stream write failed: %v", err) } readStreamExactly(t, serverInbound.Stream, "server", 2*time.Second) _ = clientOpen.stream.Close() _ = clientInbound.Stream.Close() _ = serverOpen.stream.Close() _ = serverInbound.Stream.Close() } func TestStreamHandlerWriteBeforeOpenReplyIsDelivered(t *testing.T) { server := NewServer().(*ServerCommon) if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { t.Fatalf("UseModernPSKServer failed: %v", err) } server.SetStreamHandler(func(info StreamAcceptInfo) error { _, err := info.Stream.Write([]byte("early")) return err }) if err := server.Listen("tcp", "127.0.0.1:0"); err != nil { t.Fatalf("server Listen failed: %v", err) } defer func() { _ = server.Stop() }() client := NewClient().(*ClientCommon) if err := UseModernPSKClient(client, integrationSharedSecret, integrationModernPSKOptions()); err != nil { t.Fatalf("UseModernPSKClient failed: %v", err) } if err := client.Connect("tcp", server.listener.Addr().String()); err != nil { t.Fatalf("client Connect failed: %v", err) } defer func() { _ = client.Stop() }() stream, err := client.OpenStream(context.Background(), StreamOpenOptions{Channel: StreamDataChannel}) if err != nil { t.Fatalf("client OpenStream failed: %v", err) } readStreamExactly(t, stream, "early", 2*time.Second) _ = stream.Close() } func TestServerRejectsQueuedStreamOpenFromStaleTransport(t *testing.T) { server := NewServer().(*ServerCommon) UseLegacySecurityServer(server) var handlerCalls atomic.Int32 server.SetStreamHandler(func(StreamAcceptInfo) error { handlerCalls.Add(1) return nil }) firstLeft, firstRight := net.Pipe() defer firstRight.Close() logical := server.bootstrapAcceptedLogical("stale-inbound-stream-open", nil, firstLeft) if logical == nil { t.Fatal("bootstrapAcceptedLogical should return logical") } staleTransport := logical.CurrentTransportConn() secondLeft, secondRight := net.Pipe() defer secondRight.Close() if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { t.Fatalf("attach replacement transport: %v", err) } payload, err := encode(StreamOpenRequest{StreamID: "queued-stale-open", DataID: 1}) if err != nil { t.Fatalf("encode StreamOpenRequest: %v", err) } message := Message{ NetType: NET_SERVER, LogicalConn: logical, TransportConn: staleTransport, TransferMsg: TransferMsg{ Key: StreamOpenSignalKey, Value: payload, Type: MSG_ASYNC, }, } server.handleInboundStreamOpen(&message) if got := handlerCalls.Load(); got != 0 { t.Fatalf("stale StreamOpen handler calls = %d, want 0", got) } if stream, ok := server.getStreamRuntime().lookup(serverFileScope(logical), "queued-stale-open"); ok { t.Fatalf("stale StreamOpen registered runtime handle: %+v", stream.snapshot()) } } func TestClientRejectsStreamOpenWhenRouteReattachesDuringRegistration(t *testing.T) { client := NewClient().(*ClientCommon) UseLegacySecurityClient(client) var handlerCalls atomic.Int32 client.SetStreamHandler(func(StreamAcceptInfo) error { handlerCalls.Add(1) return nil }) stopCtx, stopFn := context.WithCancel(context.Background()) defer stopFn() queue := stario.NewQueueCtx(stopCtx, 4, math.MaxUint32) firstLeft, firstRight := net.Pipe() defer firstRight.Close() epoch := client.beginClientSessionEpoch() client.setClientSessionRuntime(newClientSessionRuntime(firstLeft, stopCtx, stopFn, queue, epoch)) client.markSessionStarted() defer client.markSessionStopped("test done", nil) payload, err := encode(StreamOpenRequest{StreamID: "client-stale-inbound-open", DataID: 2}) if err != nil { t.Fatalf("encode StreamOpenRequest: %v", err) } message := Message{ NetType: NET_CLIENT, ServerConn: client, clientRoute: client.clientSessionRouteSnapshot(), TransferMsg: TransferMsg{ Key: StreamOpenSignalKey, Value: payload, Type: MSG_ASYNC, }, } runtime := client.getStreamRuntime() runtime.mu.Lock() done := make(chan struct{}) go func() { defer close(done) client.handleInboundStreamOpen(&message) }() secondLeft, secondRight := net.Pipe() defer secondRight.Close() if err := client.attachClientSessionTransport(secondLeft); err != nil { runtime.mu.Unlock() t.Fatalf("attach client replacement transport: %v", err) } runtime.mu.Unlock() select { case <-done: case <-time.After(time.Second): t.Fatal("client stale StreamOpen handler did not return") } if got := handlerCalls.Load(); got != 0 { t.Fatalf("stale client StreamOpen handler calls = %d, want 0", got) } if stream, ok := runtime.lookup(clientFileScope(), "client-stale-inbound-open"); ok { t.Fatalf("stale client StreamOpen registered runtime handle: %+v", stream.snapshot()) } } func TestClientStaleStreamRouteCannotAffectReplacement(t *testing.T) { client := NewClient().(*ClientCommon) runtime := client.getStreamRuntime() epoch := client.beginClientSessionEpoch() staleLeft, staleRight := net.Pipe() defer staleLeft.Close() defer staleRight.Close() currentLeft, currentRight := net.Pipe() defer currentLeft.Close() defer currentRight.Close() staleRoute := clientSessionRoute{binding: newTransportBinding(staleLeft, nil), epoch: epoch} currentRoute := clientSessionRoute{binding: newTransportBinding(currentLeft, nil), epoch: epoch} stream := newStreamHandle(context.Background(), runtime, clientFileScope(), StreamOpenRequest{ StreamID: "replacement-stream", DataID: 2, Channel: StreamDataChannel, }, epoch, nil, nil, 0, nil, nil, nil, runtime.configSnapshot()) stream.setClientSessionRoute(currentRoute) if err := runtime.register(clientFileScope(), stream); err != nil { t.Fatalf("register replacement stream: %v", err) } closePayload, err := encode(StreamCloseRequest{StreamID: stream.ID(), Full: true}) if err != nil { t.Fatalf("encode close request: %v", err) } client.handleInboundStreamClose(&Message{ NetType: NET_CLIENT, ServerConn: client, clientRoute: staleRoute, TransferMsg: TransferMsg{Key: StreamCloseSignalKey, Value: closePayload, Type: MSG_ASYNC}, }) resetPayload, err := encode(StreamResetRequest{StreamID: stream.ID(), DataID: stream.dataIDSnapshot(), Error: "stale reset"}) if err != nil { t.Fatalf("encode reset request: %v", err) } client.handleInboundStreamReset(&Message{ NetType: NET_CLIENT, ServerConn: client, clientRoute: staleRoute, TransferMsg: TransferMsg{Key: StreamResetSignalKey, Value: resetPayload, Type: MSG_ASYNC}, }) mismatchedClose, err := encode(StreamCloseRequest{StreamID: stream.ID(), DataID: stream.dataIDSnapshot() + 2, Full: true}) if err != nil { t.Fatalf("encode mismatched close request: %v", err) } client.handleInboundStreamClose(&Message{ NetType: NET_CLIENT, ServerConn: client, clientRoute: currentRoute, TransferMsg: TransferMsg{Key: StreamCloseSignalKey, Value: mismatchedClose, Type: MSG_ASYNC}, }) mismatchedReset, err := encode(StreamResetRequest{StreamID: stream.ID(), DataID: stream.dataIDSnapshot() + 2, Error: "mismatched reset"}) if err != nil { t.Fatalf("encode mismatched reset request: %v", err) } client.handleInboundStreamReset(&Message{ NetType: NET_CLIENT, ServerConn: client, clientRoute: currentRoute, TransferMsg: TransferMsg{Key: StreamResetSignalKey, Value: mismatchedReset, Type: MSG_ASYNC}, }) client.dispatchStreamEnvelopeAtRoute(staleRoute, newStreamDataEnvelope(stream.ID(), []byte("stale-envelope"))) client.dispatchFastStreamDataWithOwnerAtRoute(staleRoute, streamFastDataFrame{ DataID: stream.dataIDSnapshot(), Seq: 1, Payload: []byte("stale-fast"), }, nil) stream.mu.Lock() defer stream.mu.Unlock() if stream.remoteClosed || stream.peerReadClosed || stream.resetErr != nil || len(stream.readQueue) != 0 || len(stream.readBuf.data) != 0 { t.Fatalf("stale route mutated replacement: remoteClosed=%v peerReadClosed=%v reset=%v queued=%d buffered=%d", stream.remoteClosed, stream.peerReadClosed, stream.resetErr, len(stream.readQueue), len(stream.readBuf.data)) } } func TestServerStaleStreamControlsCannotAffectReplacement(t *testing.T) { server := NewServer().(*ServerCommon) UseLegacySecurityServer(server) firstLeft, firstRight := net.Pipe() defer firstRight.Close() logical := server.bootstrapAcceptedLogical("stale-stream-control", nil, firstLeft) if logical == nil { t.Fatal("bootstrapAcceptedLogical should return logical") } staleTransport := logical.CurrentTransportConn() secondLeft, secondRight := net.Pipe() defer secondRight.Close() if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { t.Fatalf("attach replacement transport: %v", err) } currentTransport := logical.CurrentTransportConn() if currentTransport == nil || currentTransport == staleTransport { t.Fatal("replacement transport should be current") } runtime := server.getStreamRuntime() scope := serverFileScope(logical) stream := newStreamHandle(logical.stopContextSnapshot(), runtime, scope, StreamOpenRequest{ StreamID: "replacement-stream", DataID: 1, Channel: StreamDataChannel, }, 0, logical, currentTransport, currentTransport.TransportGeneration(), nil, nil, nil, runtime.configSnapshot()) if err := runtime.register(scope, stream); err != nil { t.Fatalf("register replacement stream: %v", err) } closePayload, err := encode(StreamCloseRequest{StreamID: stream.ID(), Full: true}) if err != nil { t.Fatalf("encode close request: %v", err) } server.handleInboundStreamClose(&Message{ NetType: NET_SERVER, LogicalConn: logical, TransportConn: staleTransport, TransferMsg: TransferMsg{Key: StreamCloseSignalKey, Value: closePayload, Type: MSG_ASYNC}, }) resetPayload, err := encode(StreamResetRequest{StreamID: stream.ID(), DataID: stream.dataIDSnapshot(), Error: "stale reset"}) if err != nil { t.Fatalf("encode reset request: %v", err) } server.handleInboundStreamReset(&Message{ NetType: NET_SERVER, LogicalConn: logical, TransportConn: staleTransport, TransferMsg: TransferMsg{Key: StreamResetSignalKey, Value: resetPayload, Type: MSG_ASYNC}, }) stream.mu.Lock() defer stream.mu.Unlock() if stream.remoteClosed || stream.peerReadClosed || stream.resetErr != nil { t.Fatalf("stale controls mutated replacement: remoteClosed=%v peerReadClosed=%v reset=%v", stream.remoteClosed, stream.peerReadClosed, stream.resetErr) } }