package notify import ( "context" "errors" "math" "net" "os" "testing" "time" "b612.me/stario" ) func newServerBlackholeTransport(t *testing.T, id string) (*ServerCommon, *LogicalConn, *TransportConn, net.Conn, net.Conn) { t.Helper() server := NewServer().(*ServerCommon) UseLegacySecurityServer(server) stopCtx, stopFn := context.WithCancel(context.Background()) server.setServerSessionRuntime(&serverSessionRuntime{ stopCtx: stopCtx, stopFn: stopFn, queue: stario.NewQueueCtx(stopCtx, 4, math.MaxUint32), }) server.markSessionStarted() left, right := net.Pipe() logical, _, _ := newRegisteredServerLogicalForTest(t, server, id, left, stopCtx, stopFn) logical.applyAttachmentProfile(0, 0, server.defaultMsgEn, server.defaultMsgDe, server.defaultFastStreamEncode, server.defaultFastBulkEncode, server.defaultFastPlainEncode, server.handshakeRsaKey, server.SecretKey) transport := logical.CurrentTransportConn() if transport == nil { t.Fatal("server transport is nil") } t.Cleanup(func() { _ = left.Close() _ = right.Close() server.markSessionStopped("test done", nil) }) return server, logical, transport, left, right } func requireBoundedServerWrite(t *testing.T, timeout time.Duration, write func(context.Context) error) { t.Helper() ctx, cancel := context.WithTimeout(context.Background(), timeout) defer cancel() done := make(chan error, 1) started := time.Now() go func() { done <- write(ctx) }() select { case err := <-done: if err == nil { t.Fatal("blackhole server write returned nil") } if !errors.Is(err, context.DeadlineExceeded) && !errors.Is(err, os.ErrDeadlineExceeded) { var netErr net.Error if !errors.As(err, &netErr) || !netErr.Timeout() { t.Fatalf("blackhole server write error=%v, want deadline error", err) } } if elapsed := time.Since(started); elapsed > 5*timeout { t.Fatalf("blackhole server write returned after %v, want bounded by %v", elapsed, timeout) } case <-time.After(10 * timeout): t.Fatalf("blackhole server write ignored its %v context deadline", timeout) } } func TestServerSharedBulkWriteHonorsContextDeadline(t *testing.T) { server, logical, transport, _, _ := newServerBlackholeTransport(t, "shared-bulk-write-context") requireBoundedServerWrite(t, 30*time.Millisecond, func(ctx context.Context) error { return server.sendFastBulkDataTransport(ctx, logical, transport, 1, 0, []byte("bulk"), bulkFastPathVersionV1) }) } func TestServerSharedStreamWriteHonorsContextDeadline(t *testing.T) { server, logical, transport, _, _ := newServerBlackholeTransport(t, "shared-stream-write-context") stream := newStreamHandle(context.Background(), nil, serverFileScope(logical), StreamOpenRequest{ StreamID: "shared-stream-write-context", DataID: 1, Channel: StreamDataChannel, FastPathVersion: streamFastPathVersionV1, }, 0, logical, transport, transport.TransportGeneration(), nil, nil, nil, defaultStreamConfig()) requireBoundedServerWrite(t, 30*time.Millisecond, func(ctx context.Context) error { return server.sendFastStreamDataTransport(ctx, logical, transport, stream, []byte("stream")) }) } func TestMessageReplyUsesConfiguredDefaultWriteTimeout(t *testing.T) { server, logical, transport, left, _ := newServerBlackholeTransport(t, "reply-default-write-timeout") server.SetReplyWriteTimeout(35 * time.Millisecond) message := Message{ NetType: NET_SERVER, LogicalConn: logical, TransportConn: transport, TransferMsg: TransferMsg{ ID: 1, Key: "reply-default-write-timeout", Type: MSG_SYNC_ASK, }, inboundConn: left, } started := time.Now() err := message.Reply([]byte("reply")) if err == nil { t.Fatal("blackhole Message.Reply returned nil") } if elapsed := time.Since(started); elapsed > 250*time.Millisecond { t.Fatalf("Message.Reply returned after %v, want configured default write bound", elapsed) } } func TestMessageReplyCtxUsesEarlierCallerDeadline(t *testing.T) { server, logical, transport, left, _ := newServerBlackholeTransport(t, "reply-caller-write-timeout") server.SetReplyWriteTimeout(time.Second) message := Message{ NetType: NET_SERVER, LogicalConn: logical, TransportConn: transport, TransferMsg: TransferMsg{ ID: 2, Key: "reply-caller-write-timeout", Type: MSG_SYNC_ASK, }, inboundConn: left, } ctx, cancel := context.WithTimeout(context.Background(), 25*time.Millisecond) defer cancel() started := time.Now() err := message.ReplyCtx(ctx, []byte("reply")) if err == nil { t.Fatal("blackhole Message.ReplyCtx returned nil") } if elapsed := time.Since(started); elapsed > 200*time.Millisecond { t.Fatalf("Message.ReplyCtx returned after %v, want caller deadline", elapsed) } canceled, cancelNow := context.WithCancel(context.Background()) cancelNow() if err := message.ReplyObjCtx(canceled, "ok"); !errors.Is(err, context.Canceled) { t.Fatalf("Message.ReplyObjCtx canceled error=%v, want context canceled", err) } } type closeTrackingWriteConn struct { closed bool writes int } func (c *closeTrackingWriteConn) Read([]byte) (int, error) { return 0, net.ErrClosed } func (c *closeTrackingWriteConn) Close() error { c.closed = true; return nil } func (c *closeTrackingWriteConn) LocalAddr() net.Addr { return nil } func (c *closeTrackingWriteConn) RemoteAddr() net.Addr { return nil } func (c *closeTrackingWriteConn) SetDeadline(time.Time) error { return nil } func (c *closeTrackingWriteConn) SetReadDeadline(time.Time) error { return nil } func (c *closeTrackingWriteConn) SetWriteDeadline(time.Time) error { return nil } func (c *closeTrackingWriteConn) Write(data []byte) (int, error) { c.writes++ return len(data), nil } func TestServerRawWriteGateWaitTimeoutDoesNotCloseConnection(t *testing.T) { server, logical, transport, _, _ := newServerBlackholeTransport(t, "raw-write-gate-timeout") conn := &closeTrackingWriteConn{} gateRef := retainRawConnWriteGate(conn) <-gateRef.gate t.Cleanup(func() { gateRef.gate <- struct{}{} releaseRawConnWriteGate(conn, gateRef) }) ctx, cancel := context.WithTimeout(context.Background(), 25*time.Millisecond) defer cancel() err := server.writeEnvelopePayloadContextTimeout(ctx, logical, transport, conn, []byte("reply"), time.Second) if !errors.Is(err, context.DeadlineExceeded) { t.Fatalf("raw write gate wait error=%v, want context deadline exceeded", err) } if conn.closed { t.Fatal("raw write gate wait timeout closed a connection before physical write started") } if conn.writes != 0 { t.Fatalf("raw write gate wait performed %d physical writes", conn.writes) } } func TestServerUDPWriteGateWaitHonorsWriteTimeout(t *testing.T) { server := NewServer().(*ServerCommon) stopCtx, stopFn := context.WithCancel(context.Background()) t.Cleanup(stopFn) sender, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 0}) if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = sender.Close() }) receiver, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 0}) if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = receiver.Close() }) server.setServerSessionRuntime(&serverSessionRuntime{ stopCtx: stopCtx, stopFn: stopFn, queue: stario.NewQueueCtx(stopCtx, 4, math.MaxUint32), udpListener: sender, }) logical := newServerLogicalConn(server, "udp-write-gate-timeout", receiver.LocalAddr()) transport := &TransportConn{logical: logical, remoteAddr: receiver.LocalAddr(), attached: true} gateHeld := make(chan struct{}) releaseGate := make(chan struct{}) firstDone := make(chan error, 1) go func() { firstDone <- server.withUDPWriteLock(context.Background(), func() error { close(gateHeld) <-releaseGate return nil }) }() select { case <-gateHeld: case <-time.After(time.Second): t.Fatal("first UDP writer did not acquire the write gate") } started := time.Now() writeDone := make(chan error, 1) go func() { writeDone <- server.writeEnvelopePayloadContextTimeout(context.Background(), logical, transport, nil, []byte("reply"), 30*time.Millisecond) }() select { case err = <-writeDone: case <-time.After(200 * time.Millisecond): close(releaseGate) <-firstDone <-writeDone t.Fatal("UDP write gate wait ignored the configured write timeout") } if !errors.Is(err, context.DeadlineExceeded) { close(releaseGate) <-firstDone t.Fatalf("UDP write gate wait error=%v, want context deadline exceeded", err) } if elapsed := time.Since(started); elapsed > 200*time.Millisecond { close(releaseGate) <-firstDone t.Fatalf("UDP write gate wait returned after %v, want configured write bound", elapsed) } close(releaseGate) if err := <-firstDone; err != nil { t.Fatalf("first UDP writer failed: %v", err) } } func TestTransportWriteGateRegistryReleasesBindingsAndRawWrites(t *testing.T) { conn := &serializedWriteTestConn{} first := newTransportBinding(conn, stario.NewQueue()) second := newTransportBinding(conn, stario.NewQueue()) if first.writeGateSnapshot() != second.writeGateSnapshot() { t.Fatal("same physical connection did not share one write gate") } entry, ok := transportConnWriteGates.Load(conn) if !ok { t.Fatal("shared write gate was not registered") } first.stopBackgroundWorkers() if current, ok := transportConnWriteGates.Load(conn); !ok || current != entry { t.Fatal("stopping one shared binding removed the live write gate") } second.stopBackgroundWorkers() if _, ok := transportConnWriteGates.Load(conn); ok { t.Fatal("last binding release retained the historical connection write gate") } rawConn := &serializedWriteTestConn{} if err := writeFullToConn(rawConn, []byte("raw")); err != nil { t.Fatalf("raw write failed: %v", err) } if _, ok := transportConnWriteGates.Load(rawConn); ok { t.Fatal("completed raw write retained a temporary write gate reference") } }