280 lines
9.7 KiB
Go
280 lines
9.7 KiB
Go
|
|
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")
|
||
|
|
}
|
||
|
|
}
|