package notify import ( "context" "errors" "math" "net" "os" "sync/atomic" "testing" "time" "b612.me/stario" ) func TestBulkRecoveryQueueRetriesWithoutConcurrentWorkers(t *testing.T) { var attempts atomic.Int32 var active atomic.Int32 var maxActive atomic.Int32 done := make(chan struct{}) q := newBulkRecoveryQueue(func(error) { t.Errorf("recovery should succeed") }) if !q.enqueue(func(context.Context) error { current := active.Add(1) for { previous := maxActive.Load() if current <= previous || maxActive.CompareAndSwap(previous, current) { break } } defer active.Add(-1) attempt := attempts.Add(1) if attempt < 3 { return errors.New("retry") } close(done) return nil }) { t.Fatal("enqueue unexpectedly rejected task") } select { case <-done: case <-time.After(3 * time.Second): t.Fatal("timed out waiting for recovery retries") } if got := attempts.Load(); got != 3 { t.Fatalf("attempts = %d, want 3", got) } if got := maxActive.Load(); got != 1 { t.Fatalf("max concurrent recovery workers = %d, want 1", got) } } func TestBulkRecoveryQueueStopsWorkerWhenIdleAndRestarts(t *testing.T) { q := newBulkRecoveryQueue(func(error) { t.Fatal("unexpected recovery error") }) done := make(chan struct{}) if !q.enqueue(func(context.Context) error { close(done) return nil }) { t.Fatal("enqueue unexpectedly rejected task") } select { case <-done: case <-time.After(time.Second): t.Fatal("timed out waiting for recovery task") } deadline := time.Now().Add(time.Second) for time.Now().Before(deadline) { q.mu.Lock() running := q.workerRunning q.mu.Unlock() if !running { break } time.Sleep(time.Millisecond) } q.mu.Lock() running := q.workerRunning q.mu.Unlock() if running { t.Fatal("recovery worker remained alive after queue drained") } secondDone := make(chan struct{}) if !q.enqueue(func(context.Context) error { close(secondDone) return nil }) { t.Fatal("enqueue after idle unexpectedly rejected task") } select { case <-secondDone: case <-time.After(time.Second): t.Fatal("recovery queue did not restart after idle") } } func TestBulkRecoveryQueueDoesNotRetryDetachedSession(t *testing.T) { var attempts atomic.Int32 q := newBulkRecoveryQueue(nil) q.run(func(context.Context) error { attempts.Add(1) return transportDetachedSessionEpochError() }) if got := attempts.Load(); got != 1 { t.Fatalf("detached recovery attempts = %d, want 1", got) } } func TestBulkRecoveryContextCancellationStopsRetryLoop(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() started := make(chan struct{}) var attempts atomic.Int32 task := func(taskCtx context.Context) error { attempts.Add(1) close(started) <-taskCtx.Done() return taskCtx.Err() } done := make(chan error, 1) go func() { done <- runBulkRecoveryTaskContext(ctx, task) }() select { case <-started: case <-time.After(time.Second): t.Fatal("recovery task did not start") } cancel() select { case err := <-done: if !errors.Is(err, context.Canceled) { t.Fatalf("recovery error = %v, want context.Canceled", err) } case <-time.After(time.Second): t.Fatal("recovery did not stop after parent cancellation") } if got := attempts.Load(); got != 1 { t.Fatalf("recovery attempts = %d, want 1", got) } } func TestClientBulkRecoveryQueueOverflowDoesNotBlockCaller(t *testing.T) { release := make(chan struct{}) started := make(chan struct{}, bulkRecoveryQueueSize+16) blockingTask := func(ctx context.Context) error { started <- struct{}{} select { case <-release: return nil case <-ctx.Done(): return ctx.Err() } } q := newBulkRecoveryQueue(nil) if !q.enqueue(blockingTask) { t.Fatal("initial recovery task was rejected") } select { case <-started: case <-time.After(time.Second): t.Fatal("recovery worker did not start") } for q.enqueue(blockingTask) { } client := NewClient().(*ClientCommon) client.bulkRecovery = q epoch := client.beginClientSessionEpoch() returned := make(chan struct{}) go func() { client.bestEffortBulkResetAtEpoch(epoch, BulkResetRequest{BulkID: "overflow"}) close(returned) }() select { case <-returned: close(release) case <-time.After(100 * time.Millisecond): close(release) <-returned t.Fatal("queue-full recovery blocked the caller") } } func TestClientBulkResetRecoveryTaskRejectsStaleSession(t *testing.T) { client := NewClient().(*ClientCommon) epoch := client.beginClientSessionEpoch() client.beginClientSessionEpoch() task := newClientBulkResetRecoveryTask(client, epoch, BulkResetRequest{BulkID: "stale"}) err := task(context.Background()) if !errors.Is(err, errTransportDetached) { t.Fatalf("stale recovery task error = %v, want transport detached", err) } } func TestBulkOpenAutoWaitsForDedicatedResetBeforeSharedFallback(t *testing.T) { server := NewServer().(*ServerCommon) if err := UseModernPSKServer(server, integrationSharedSecret, integrationModernPSKOptions()); err != nil { t.Fatalf("UseModernPSKServer failed: %v", err) } accepted := make(chan BulkAcceptInfo, 2) server.SetBulkHandler(func(info BulkAcceptInfo) error { accepted <- info return nil }) if err := server.Listen("tcp", "127.0.0.1:0"); err != nil { t.Fatalf("server Listen failed: %v", err) } defer 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 client.Stop() client.setClientConnectSource(newClientFactoryConnectSource(func(context.Context) (net.Conn, error) { return nil, errors.New("forced attach dial failure") })) // Occupy every asynchronous recovery worker. The old Auto path queued its // reset behind these tasks and immediately sent shared open, reproducing the // explicit-ID race. The fixed path performs the reset synchronously. release := make(chan struct{}) defer close(release) started := make(chan struct{}, bulkRecoveryWorkers) q := newBulkRecoveryQueue(nil) blockingTask := func(ctx context.Context) error { started <- struct{}{} select { case <-release: return nil case <-ctx.Done(): return ctx.Err() } } for i := 0; i < bulkRecoveryWorkers; i++ { if !q.enqueue(blockingTask) { t.Fatalf("enqueue blocking recovery task %d failed", i) } } for i := 0; i < bulkRecoveryWorkers; i++ { select { case <-started: case <-time.After(time.Second): t.Fatal("recovery worker did not become occupied") } } client.bulkRecovery = q bulk, err := client.OpenBulk(context.Background(), BulkOpenOptions{ Mode: BulkOpenModeAuto, ID: "auto-reset-order", Range: BulkRange{ Offset: 0, Length: 128, }, }) if err != nil { t.Fatalf("Auto fallback failed while reset workers were occupied: %v", err) } if bulk.Snapshot().Dedicated { t.Fatal("Auto fallback returned a dedicated bulk") } defer bulk.Close() seenShared := false deadline := time.After(2 * time.Second) for !seenShared { select { case info := <-accepted: if info.Bulk != nil { defer info.Bulk.Close() } if info.ID == bulk.ID() { seenShared = !info.Dedicated } case <-deadline: t.Fatal("timed out waiting for shared fallback accept") } } } func TestServerBulkResetRecoveryTaskRejectsStaleTransport(t *testing.T) { server := NewServer().(*ServerCommon) UseLegacySecurityServer(server) runtimeCtx, runtimeCancel := context.WithCancel(context.Background()) defer runtimeCancel() queue := stario.NewQueueCtx(runtimeCtx, 4, math.MaxUint32) server.setServerSessionRuntime(&serverSessionRuntime{ stopCtx: runtimeCtx, stopFn: runtimeCancel, queue: queue, }) server.markSessionStarted() defer server.markSessionStopped("test done", nil) firstLeft, firstRight := net.Pipe() defer firstRight.Close() logical, _, _ := newRegisteredServerLogicalForTest(t, server, "bulk-recovery-stale", firstLeft, runtimeCtx, runtimeCancel) firstTransport := logical.CurrentTransportConn() if firstTransport == nil { t.Fatal("first transport snapshot should exist") } secondLeft, secondRight := net.Pipe() defer secondRight.Close() if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { t.Fatalf("attachClientConnSessionTransport failed: %v", err) } if firstTransport.IsCurrent() { t.Fatal("first transport should be stale after reattach") } task := newServerBulkResetRecoveryTask(server, logical, firstTransport, BulkResetRequest{BulkID: "stale"}) err := task(context.Background()) if !errors.Is(err, errTransportDetached) { t.Fatalf("stale recovery task error = %v, want transport detached", err) } } func TestClientBulkOpenDoesNotCrossTransportReattach(t *testing.T) { client := NewClient().(*ClientCommon) UseLegacySecurityClient(client) 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) runtime := client.getBulkRuntime() runtime.mu.Lock() result := make(chan error, 1) go func() { ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() _, err := client.OpenBulk(ctx, BulkOpenOptions{Mode: BulkOpenModeShared}) result <- err }() waitForBulkOpenRouteCapture(t, runtime) 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() if err := <-result; !errors.Is(err, errTransportDetached) { t.Fatalf("bulk open error = %v, want transport detached", err) } assertNoPipeWrite(t, secondRight, "client bulk open crossed onto replacement transport") } func TestServerBulkOpenDoesNotCrossTransportReattach(t *testing.T) { server := NewServer().(*ServerCommon) UseLegacySecurityServer(server) runtimeCtx, runtimeCancel := context.WithCancel(context.Background()) defer runtimeCancel() queue := stario.NewQueueCtx(runtimeCtx, 4, math.MaxUint32) server.setServerSessionRuntime(&serverSessionRuntime{stopCtx: runtimeCtx, stopFn: runtimeCancel, queue: queue}) server.markSessionStarted() defer server.markSessionStopped("test done", nil) firstLeft, firstRight := net.Pipe() defer firstRight.Close() logical, _, _ := newRegisteredServerLogicalForTest(t, server, "bulk-open-reattach", firstLeft, runtimeCtx, runtimeCancel) runtime := server.getBulkRuntime() runtime.mu.Lock() result := make(chan error, 1) go func() { ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() _, err := server.OpenBulkLogical(ctx, logical, BulkOpenOptions{Mode: BulkOpenModeShared}) result <- err }() waitForBulkOpenRouteCapture(t, runtime) secondLeft, secondRight := net.Pipe() defer secondRight.Close() if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { runtime.mu.Unlock() t.Fatalf("attach server replacement transport: %v", err) } runtime.mu.Unlock() if err := <-result; !errors.Is(err, errTransportDetached) { t.Fatalf("bulk open error = %v, want transport detached", err) } assertNoPipeWrite(t, secondRight, "server bulk open crossed onto replacement transport") } func TestServerRejectsQueuedBulkOpenFromStaleTransport(t *testing.T) { server := NewServer().(*ServerCommon) UseLegacySecurityServer(server) var handlerCalls atomic.Int32 server.SetBulkHandler(func(BulkAcceptInfo) error { handlerCalls.Add(1) return nil }) firstLeft, firstRight := net.Pipe() defer firstRight.Close() logical := server.bootstrapAcceptedLogical("stale-inbound-bulk-open", nil, firstLeft) if logical == nil { t.Fatal("bootstrapAcceptedLogical should return logical") } staleTransport := logical.CurrentTransportConn() if staleTransport == nil { t.Fatal("initial transport snapshot should exist") } secondLeft, secondRight := net.Pipe() defer secondRight.Close() if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { t.Fatalf("attach replacement transport: %v", err) } payload, err := encode(BulkOpenRequest{BulkID: "queued-stale-open", DataID: 1}) if err != nil { t.Fatalf("encode BulkOpenRequest: %v", err) } message := Message{ NetType: NET_SERVER, LogicalConn: logical, TransportConn: staleTransport, TransferMsg: TransferMsg{ Key: BulkOpenSignalKey, Value: payload, Type: MSG_ASYNC, }, } server.handleInboundBulkOpen(&message) if got := handlerCalls.Load(); got != 0 { t.Fatalf("stale BulkOpen handler calls = %d, want 0", got) } if bulk, ok := server.getBulkRuntime().lookup(serverFileScope(logical), "queued-stale-open"); ok { t.Fatalf("stale BulkOpen registered runtime handle: %+v", bulk.snapshot()) } } func TestServerTransportReattachResetsExistingBulk(t *testing.T) { server := NewServer().(*ServerCommon) UseLegacySecurityServer(server) firstLeft, firstRight := net.Pipe() defer firstRight.Close() logical := server.bootstrapAcceptedLogical("bulk-reset-on-reattach", nil, firstLeft) if logical == nil { t.Fatal("bootstrapAcceptedLogical should return logical") } firstTransport := logical.CurrentTransportConn() if firstTransport == nil { t.Fatal("initial transport snapshot should exist") } runtime := server.getBulkRuntime() bulk := newBulkHandle(logical.stopContextSnapshot(), runtime, serverFileScope(logical), BulkOpenRequest{ BulkID: "bulk-before-reattach", DataID: 2, }, 0, logical, firstTransport, firstTransport.TransportGeneration(), nil, nil, nil, nil, nil) if err := runtime.register(serverFileScope(logical), bulk); err != nil { t.Fatalf("register bulk: %v", err) } secondLeft, secondRight := net.Pipe() defer secondRight.Close() if err := server.attachAcceptedLogicalTransport(logical, secondLeft.RemoteAddr(), secondLeft); err != nil { t.Fatalf("attach replacement transport: %v", err) } select { case <-bulk.Context().Done(): case <-time.After(time.Second): t.Fatal("old-generation bulk remained active after transport reattach") } if err := bulk.resetErrSnapshot(); !errors.Is(err, errTransportDetached) { t.Fatalf("old-generation bulk reset error = %v, want transport detached", err) } if registered, ok := runtime.lookup(serverFileScope(logical), bulk.ID()); ok { t.Fatalf("old-generation bulk remained registered: %+v", registered.snapshot()) } if firstTransport.IsCurrent() { t.Fatal("old transport remained current after replacement") } } func TestClientRejectsBulkOpenWhenRouteReattachesDuringRegistration(t *testing.T) { client := NewClient().(*ClientCommon) UseLegacySecurityClient(client) var handlerCalls atomic.Int32 client.SetBulkHandler(func(BulkAcceptInfo) 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(BulkOpenRequest{BulkID: "client-stale-inbound-open", DataID: 2}) if err != nil { t.Fatalf("encode BulkOpenRequest: %v", err) } message := Message{ NetType: NET_CLIENT, ServerConn: client, clientRoute: client.clientSessionRouteSnapshot(), TransferMsg: TransferMsg{ Key: BulkOpenSignalKey, Value: payload, Type: MSG_ASYNC, }, } runtime := client.getBulkRuntime() runtime.mu.Lock() done := make(chan struct{}) go func() { defer close(done) client.handleInboundBulkOpen(&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 BulkOpen handler did not return") } if got := handlerCalls.Load(); got != 0 { t.Fatalf("stale client BulkOpen handler calls = %d, want 0", got) } if bulk, ok := runtime.lookup(clientFileScope(), "client-stale-inbound-open"); ok { t.Fatalf("stale client BulkOpen registered runtime handle: %+v", bulk.snapshot()) } } func TestClientBulkReadyDoesNotCrossTransportReattach(t *testing.T) { client := NewClient().(*ClientCommon) UseLegacySecurityClient(client) 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) route := client.clientSessionRouteSnapshot() bulk := newBulkHandle(stopCtx, nil, clientFileScope(), BulkOpenRequest{ BulkID: "client-ready-reattach", DataID: 1, }, epoch, nil, nil, 0, nil, nil, nil, nil, nil) bulk.setClientSessionRoute(route) defer bulk.finalize() secondLeft, secondRight := net.Pipe() defer secondRight.Close() if err := client.attachClientSessionTransport(secondLeft); err != nil { t.Fatalf("attach client replacement transport: %v", err) } done := make(chan struct{}) go func() { client.clientBulkAcceptReadyNotifier(bulk)(nil) close(done) }() assertNoPipeWrite(t, secondRight, "client bulk ready crossed onto replacement transport") select { case <-done: case <-time.After(time.Second): t.Fatal("client bulk ready notifier did not reject stale route") } } func TestServerBulkReadyDoesNotCrossTransportReattach(t *testing.T) { server := NewServer().(*ServerCommon) UseLegacySecurityServer(server) runtimeCtx, runtimeCancel := context.WithCancel(context.Background()) defer runtimeCancel() queue := stario.NewQueueCtx(runtimeCtx, 4, math.MaxUint32) server.setServerSessionRuntime(&serverSessionRuntime{stopCtx: runtimeCtx, stopFn: runtimeCancel, queue: queue}) server.markSessionStarted() defer server.markSessionStopped("test done", nil) firstLeft, firstRight := net.Pipe() defer firstRight.Close() logical, _, _ := newRegisteredServerLogicalForTest(t, server, "bulk-ready-reattach", firstLeft, runtimeCtx, runtimeCancel) logical.applyClientConnAttachmentProfile(0, 100*time.Millisecond, server.defaultMsgEn, server.defaultMsgDe, server.handshakeRsaKey, server.SecretKey) firstTransport := logical.CurrentTransportConn() if firstTransport == nil { t.Fatal("first transport snapshot should exist") } bulk := newBulkHandle(runtimeCtx, nil, serverFileScope(logical), BulkOpenRequest{ BulkID: "server-ready-reattach", DataID: 1, }, 0, logical, firstTransport, firstTransport.TransportGeneration(), nil, nil, nil, nil, nil) defer bulk.finalize() secondLeft, secondRight := net.Pipe() defer secondRight.Close() if err := logical.attachClientConnSessionTransport(secondLeft); err != nil { t.Fatalf("attach server replacement transport: %v", err) } result := make(chan error, 1) go func() { ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() result <- sendBulkReadyServer(ctx, server, logical, firstTransport, BulkReadyRequest{ BulkID: bulk.ID(), DataID: bulk.dataIDSnapshot(), }) }() assertNoPipeWrite(t, secondRight, "server bulk ready crossed onto replacement transport") select { case err := <-result: if !errors.Is(err, errTransportDetached) { t.Fatalf("server bulk ready error = %v, want transport detached", err) } case <-time.After(time.Second): t.Fatal("server bulk ready send did not reject stale transport") } } func waitForBulkOpenRouteCapture(t *testing.T, runtime *bulkRuntime) { t.Helper() deadline := time.Now().Add(time.Second) for time.Now().Before(deadline) { if runtime.seq.Load() != 0 { return } time.Sleep(time.Millisecond) } runtime.mu.Unlock() t.Fatal("timed out waiting for bulk open to capture its route") } func assertNoPipeWrite(t *testing.T, conn net.Conn, message string) { t.Helper() if err := conn.SetReadDeadline(time.Now().Add(50 * time.Millisecond)); err != nil { t.Fatalf("set pipe read deadline: %v", err) } buf := make([]byte, 1) if n, err := conn.Read(buf); n != 0 || !errors.Is(err, os.ErrDeadlineExceeded) { t.Fatalf("%s: read=%d err=%v", message, n, err) } } func TestClientInboundParserDoesNotJoinFramesAcrossTransportReattach(t *testing.T) { client := NewClient().(*ClientCommon) UseLegacySecurityClient(client) stopCtx, stopFn := context.WithCancel(context.Background()) defer stopFn() queue := stario.NewQueueCtx(stopCtx, 4, math.MaxUint32) firstBinding := newTransportBinding(nil, queue) secondBinding := newTransportBinding(nil, queue) secondRuntime := prepareClientSessionRuntime(&clientSessionRuntime{ transport: secondBinding, transportAttached: true, stopCtx: stopCtx, stopFn: stopFn, queue: queue, inboundDispatcher: newInboundDispatcher(), epoch: 2, }) client.setClientSessionRuntime(secondRuntime) received := make(chan Message, 2) client.SetLink("reattach-frame", func(msg *Message) { received <- *msg }) env, err := wrapTransferMsgEnvelope(TransferMsg{ID: 91, Key: "reattach-frame", Value: MsgVal("payload"), Type: MSG_ASYNC}, client.sequenceEn) if err != nil { t.Fatalf("wrap transfer envelope: %v", err) } wire, err := client.encodeEnvelope(env) if err != nil { t.Fatalf("encode envelope: %v", err) } cut := len(wire) / 2 firstRoute := clientSessionRoute{binding: firstBinding, epoch: 1, sessionStopCtx: stopCtx, transportStopCtx: stopCtx} secondRoute := clientSessionRouteFromRuntime(secondRuntime) client.pushMessageFastAtRoute(firstRoute, queue, wire[:cut], secondRuntime.inboundDispatcher) client.pushMessageFastAtRoute(secondRoute, queue, wire[cut:], secondRuntime.inboundDispatcher) select { case msg := <-received: t.Fatalf("split frame crossed transports: %+v", msg.TransferMsg) case <-time.After(20 * time.Millisecond): } client.pushMessageFastAtRoute(secondRoute, queue, wire, secondRuntime.inboundDispatcher) select { case msg := <-received: if msg.Key != "reattach-frame" || string(msg.Value) != "payload" { t.Fatalf("decoded message = %+v", msg.TransferMsg) } case <-time.After(time.Second): t.Fatal("complete replacement-transport frame was not dispatched") } }