diff --git a/AGENTS.md b/AGENTS.md index 8682ba4..a9e87a1 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -13,13 +13,15 @@ Module path: `github.com/tphakala/et-go`. Go version: see `go.mod`. ## Layout -| Path | Role | -| ------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| `cmd/et` | Entry point and flag parsing. | -| `internal/protocol` | Wire messages generated from upstream's `.proto` files (opaque API), plus `Header`, `Version` and `Packet`. Regenerate with `go generate ./internal/protocol` (needs protoc 3.21.12). | -| `internal/seal` | One direction of the libsodium-compatible encrypted stream: secretbox with a counter nonce. | -| `internal/wire` | Handshake message framing, stream frame framing, packet layout and size limits. | -| `rules/` | ruleguard matchers used by golangci-lint (build tag `ruleguard`). | +| Path | Role | +| ----------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `cmd/et` | Entry point and flag parsing. | +| `internal/etcp` | Reliable, ordered, encrypted packet connection over replaceable TCP links: replay ring, recover exchange, liveness probes, reconnect backoff, write backpressure. | +| `internal/etservertest` | Test-only fake etserver (written independently from upstream semantics) and an in-memory `net.Pipe` network with cut and refuse controls, for synctest-driven etcp tests. | +| `internal/protocol` | Wire messages generated from upstream's `.proto` files (opaque API), plus `Header`, `Version` and `Packet`. Regenerate with `go generate ./internal/protocol` (needs protoc 3.21.12). | +| `internal/seal` | One direction of the libsodium-compatible encrypted stream: secretbox with a counter nonce. | +| `internal/wire` | Handshake message framing, stream frame framing, packet layout and size limits. | +| `rules/` | ruleguard matchers used by golangci-lint (build tag `ruleguard`). | Update this table when a package is added. diff --git a/internal/etcp/backoff.go b/internal/etcp/backoff.go new file mode 100644 index 0000000..4facc35 --- /dev/null +++ b/internal/etcp/backoff.go @@ -0,0 +1,39 @@ +package etcp + +import ( + "math/rand/v2" + "time" +) + +const ( + backoffBase = 250 * time.Millisecond + backoffMax = 5 * time.Second + backoffReset = 30 * time.Second // a link that lived this long resets the schedule +) + +// backoff yields reconnect delays: none before the first attempt, then 250 ms +// doubling to a 5 s cap, each with +/-20% jitter, never above the cap. +type backoff struct { + attempt int + jitter func() float64 // returns [0, 1); nil means math/rand/v2 +} + +func (b *backoff) next() time.Duration { + n := b.attempt + b.attempt++ + if n == 0 { + return 0 + } + // The shift is clamped because backoffBase<<36 overflows int64 and would + // yield a negative delay, a hot redial loop; 5 is the first shift at + // which backoffBase exceeds backoffMax, so the clamp never lowers a delay. + d := min(backoffBase< backoffMax { + t.Fatalf("delay %v outside (0, %v]", got, backoffMax) + } + } + }) + } +} diff --git a/internal/etcp/backpressure_internal_test.go b/internal/etcp/backpressure_internal_test.go new file mode 100644 index 0000000..f0774d0 --- /dev/null +++ b/internal/etcp/backpressure_internal_test.go @@ -0,0 +1,123 @@ +package etcp + +import ( + "context" + "strings" + "testing" + "testing/synctest" + "time" + + "github.com/tphakala/et-go/internal/protocol" +) + +// blockedConn returns a Conn with no link whose backlog is over its limit, so +// every WritePacket parks until the test makes room. +func blockedConn(t *testing.T) *Conn { + t.Helper() + var d Dialer + c := d.newConn("et.example:2022", "XXXtestclient001", strings.Repeat("k", 32)) + c.mu.Lock() + c.unsent = c.limit + 1 + c.mu.Unlock() + return c +} + +// makeRoom empties the backlog and hands out one wakeup, as a link writer +// does after a drain. +func makeRoom(c *Conn) { + c.mu.Lock() + c.unsent = 0 + c.mu.Unlock() + signal(c.space) +} + +func waitWrite(t *testing.T, name string, done <-chan error) { + t.Helper() + select { + case err := <-done: + if err != nil { + t.Fatalf("writer %s: %v", name, err) + } + case <-time.After(time.Minute): + t.Fatalf("writer %s still blocked a minute after room was made", name) + } +} + +// One drain wakes one blocked writer, which passes the turn on, so every +// writer that fits proceeds. +func TestBlockedWritersAllProceed(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + c := blockedConn(t) + defer c.cancel(nil) + + doneA := make(chan error, 1) + doneB := make(chan error, 1) + go func() { doneA <- c.WritePacket(t.Context(), protocol.Packet{}) }() + synctest.Wait() + go func() { doneB <- c.WritePacket(t.Context(), protocol.Packet{}) }() + synctest.Wait() + + makeRoom(c) + waitWrite(t, "A", doneA) + waitWrite(t, "B", doneB) + }) +} + +// A writer that is handed the wakeup and whose context is cancelled before it +// runs must not drop the wakeup: the other blocked writer would then stay +// parked with room available. Channel receivers are served in order, so the +// wakeup goes to A, the first writer to park. +func TestWakeupNotLostOnCancel(t *testing.T) { + for range 20 { + synctest.Test(t, func(t *testing.T) { + c := blockedConn(t) + defer c.cancel(nil) + + ctxA, cancelA := context.WithCancel(t.Context()) + doneA := make(chan error, 1) + doneB := make(chan error, 1) + go func() { doneA <- c.WritePacket(ctxA, protocol.Packet{}) }() + synctest.Wait() + go func() { doneB <- c.WritePacket(t.Context(), protocol.Packet{}) }() + synctest.Wait() + + makeRoom(c) // A is handed the wakeup + cancelA() // and is cancelled before it gets to run + select { // A may enqueue or return its cause; either is fine + case <-doneA: + case <-time.After(time.Minute): + t.Fatal("writer A still blocked a minute after it was cancelled") + } + waitWrite(t, "B", doneB) + }) + } +} + +// A write that reaches the queue lock after the Conn ended is refused, not +// queued on the dead Conn. The test holds the lock so the writer parks on +// it, then ends the Conn. It uses real time because synctest cannot wait +// for a goroutine blocked on a mutex; a writer that has not reached the lock +// in time sees the ended Conn anyway, so the test cannot fail spuriously. +func TestWritePacketRefusedOnceConnEnded(t *testing.T) { + var d Dialer + c := d.newConn("et.example:2022", "XXXtestclient001", strings.Repeat("k", 32)) + c.mu.Lock() + done := make(chan error, 1) + go func() { done <- c.WritePacket(t.Context(), protocol.Packet{}) }() + time.Sleep(50 * time.Millisecond) // let the writer park on c.mu + c.cancel(errClosed) + c.mu.Unlock() + select { + case err := <-done: + if err == nil { + t.Fatal("WritePacket = nil after the Conn ended, want its cause") + } + case <-time.After(time.Minute): + t.Fatal("WritePacket did not return") + } + c.mu.Lock() + defer c.mu.Unlock() + if n := c.ring.next(); n != 0 { + t.Fatalf("%d packets queued on the ended Conn, want 0", n) + } +} diff --git a/internal/etcp/backpressure_test.go b/internal/etcp/backpressure_test.go new file mode 100644 index 0000000..31ae658 --- /dev/null +++ b/internal/etcp/backpressure_test.go @@ -0,0 +1,91 @@ +package etcp_test + +import ( + "context" + "errors" + "net" + "testing" + "testing/synctest" + "time" + + "github.com/tphakala/et-go/internal/etcp" +) + +// A writer blocked on a full backlog is released by its own context, with +// that context's cause, and by Close, with net.ErrClosed. An already +// cancelled context is refused without queueing. +func TestBackpressureWaitEnds(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{ReplayLimit: 1 << 10}) + defer h.close() + + synctest.Wait() + h.net.SetRefuse(true) + h.net.CutAll() + synctest.Wait() + errCause := errors.New("caller gave up") + ctx, cancel := context.WithCancelCause(t.Context()) + cancel(errCause) + if err := h.conn.WritePacket(ctx, numbered(0, 10)); !errors.Is(err, errCause) { + t.Fatalf("write with a cancelled context = %v, want %v", err, errCause) + } + if err := h.conn.WritePacket(t.Context(), numbered(0, 2048)); err != nil { + t.Fatalf("first write: %v", err) + } + + ctx, cancel = context.WithCancelCause(t.Context()) + done := make(chan error, 1) + go func() { done <- h.conn.WritePacket(ctx, numbered(1, 10)) }() + synctest.Wait() + cancel(errCause) + if err := within(t, done); !errors.Is(err, errCause) { + t.Fatalf("blocked write after cancel = %v, want %v", err, errCause) + } + + go func() { done <- h.conn.WritePacket(t.Context(), numbered(1, 10)) }() + synctest.Wait() + if err := h.conn.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if err := within(t, done); !errors.Is(err, net.ErrClosed) { + t.Fatalf("blocked write after Close = %v, want net.ErrClosed", err) + } + }) +} + +func TestBackpressureWhileDisconnected(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{ReplayLimit: 4 << 10}) + defer h.close() + + synctest.Wait() + h.net.SetRefuse(true) + h.net.CutAll() + synctest.Wait() // the dead link is noticed and the supervisor is backing off + + // Each packet seals to 1024+2+16 = 1042 bytes. Writes are admitted + // while the unsent backlog is at most 4 KiB, so the fourth takes it + // past the limit and the fifth must wait. + for i := range 4 { + if err := h.conn.WritePacket(t.Context(), numbered(i, 1024)); err != nil { + t.Fatalf("WritePacket %d: %v", i, err) + } + } + done := make(chan error, 1) + go func() { done <- h.conn.WritePacket(t.Context(), numbered(4, 1024)) }() + synctest.Sleep(time.Minute) + select { + case err := <-done: + t.Fatalf("fifth write returned %v while the backlog was full", err) + default: + } + + h.net.SetRefuse(false) + if err := within(t, done); err != nil { + t.Fatalf("fifth write after reconnect: %v", err) + } + if err := expectNumbered(t.Context(), 5, h.srv.Recv); err != nil { + t.Fatal(err) + } + }) +} diff --git a/internal/etcp/conn.go b/internal/etcp/conn.go new file mode 100644 index 0000000..3c6c8ad --- /dev/null +++ b/internal/etcp/conn.go @@ -0,0 +1,205 @@ +package etcp + +import ( + "context" + "fmt" + "log/slog" + "sync" + "time" + + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/seal" + "github.com/tphakala/et-go/internal/wire" + "golang.org/x/crypto/nacl/secretbox" +) + +const inboxSize = 64 + +// maxPayload is the largest packet payload whose sealed, serialized packet +// still fits in one wire frame. +const maxPayload = wire.MaxFrameSize - 2 - secretbox.Overhead + +// Conn is a reliable, ordered, encrypted packet connection that survives +// reconnects. Its methods are safe for concurrent use. +// +// Waiting happens only on channels and timers, never on a held mutex, so +// testing/synctest can drive every blocking path. +type Conn struct { + netDialer netDialer + addr string + id string + keepAlive time.Duration + probe protocol.Packet + limit int + logger *slog.Logger + + ctx context.Context // lifetime of the Conn; its cause is what callers see + cancel context.CancelCauseFunc + wg sync.WaitGroup + + mu sync.Mutex // guards the outbound state below + out *seal.Stream // client to server + ring ring + flushed int64 // next sequence number the current link writer sends + unsent int // bytes of entries at or after flushed + // lastProbe is the sequence number of the latest probe, or -1. + lastProbe int64 + + // in and recvSeq belong to whichever goroutine reads: a link reader, the + // supervisor during recovery, or Dial's caller during a first-connect + // recovery, never two at once. + in *seal.Stream + recvSeq int64 + // pending holds packets already opened and counted in recvSeq whose link + // ended before ReadPacket took them; the next link reader hands them + // over first. Owned like recvSeq until readerDone closes, then by + // ReadPacket under tailMu. Pending packets are always newer than every + // packet in inbox: a reader hands nothing on while pending is non-empty. + pending []protocol.Packet + readerDone chan struct{} // closed when the supervisor, and so every reader, has exited + tailMu sync.Mutex + + inbox chan protocol.Packet + wake chan struct{} // cap 1: new outbound data for the link writer + space chan struct{} // cap 1: the unsent backlog fell to limit or below +} + +// WritePacket seals p and queues it for delivery. It returns once p is queued, +// not when it is sent. It blocks only when the unsent backlog exceeds +// ReplayLimit, until a link drains it, ctx ends or the Conn fails. A payload +// too large for one sealed frame (just under wire.MaxFrameSize) is refused +// with an error wrapping wire.ErrTooLarge. A ctx already done on entry is +// refused; one that ends while the write is blocked returns its cause, unless +// the write was woken for room at the same moment, in which case it may still +// be queued, as a net.Conn write racing its deadline may complete. +func (c *Conn) WritePacket(ctx context.Context, p protocol.Packet) error { + if len(p.Payload) > maxPayload { + return fmt.Errorf("etcp: packet payload of %d bytes: %w", len(p.Payload), wire.ErrTooLarge) + } + // ctx is checked here and in the select below, never between waking on + // c.space and the room check: a writer that took the wakeup and then + // returned would leave the next blocked writer parked with room available. + if err := ctx.Err(); err != nil { + return context.Cause(ctx) + } + for { + c.mu.Lock() + // Checked under the lock: a write that waited for it while the + // Conn ended must not be queued on the dead Conn. + if c.ctx.Err() != nil { + c.mu.Unlock() + return context.Cause(c.ctx) + } + if c.unsent <= c.limit { + c.enqueueLocked(p) + room := c.unsent <= c.limit + c.mu.Unlock() + signal(c.wake) + if room { + signal(c.space) // pass the turn to another blocked writer + } + return nil + } + c.mu.Unlock() + select { + case <-c.space: + case <-ctx.Done(): + return context.Cause(ctx) + case <-c.ctx.Done(): + return context.Cause(c.ctx) + } + } +} + +// ReadPacket returns the next packet from the server, blocking until one +// arrives, ctx ends or the Conn fails. Packets are delivered exactly once, in +// order. Packets that arrived before the Conn failed are still returned +// first, so a shell's last output is not lost when the session ends; that +// includes a packet a dead link was holding for a caller that had stopped +// reading. +func (c *Conn) ReadPacket(ctx context.Context) (protocol.Packet, error) { + select { + case p := <-c.inbox: + return p, nil + case <-ctx.Done(): + return protocol.Packet{}, context.Cause(ctx) + case <-c.ctx.Done(): + // Wait for the readers to stop, so inbox and pending are final; + // then return the inbox, then pending, which is newer. + select { + case <-c.readerDone: + case <-ctx.Done(): + return protocol.Packet{}, context.Cause(ctx) + } + select { + case p := <-c.inbox: + return p, nil + default: + } + c.tailMu.Lock() + defer c.tailMu.Unlock() + if len(c.pending) > 0 { + p := c.pending[0] + c.pending[0] = protocol.Packet{} + c.pending = c.pending[1:] + return p, nil + } + return protocol.Packet{}, context.Cause(c.ctx) + } +} + +// Close shuts the connection down and waits for its goroutines. Afterwards +// WritePacket returns net.ErrClosed, and ReadPacket returns any packets that +// had already arrived and then net.ErrClosed; if the Conn had already +// failed, both keep returning the error it failed with. The server keeps the +// session; a later client cannot resume it (the replay state is gone). +func (c *Conn) Close() error { + c.cancel(errClosed) + c.wg.Wait() + return nil +} + +// sealedNone reports whether no packet has been sealed yet, probes included. +func (c *Conn) sealedNone() bool { + c.mu.Lock() + defer c.mu.Unlock() + return c.ring.next() == 0 +} + +// probeOnce queues the liveness probe regardless of the backlog limit, +// unless an earlier probe is still unwritten: that one draws the same +// answer, and since probes bypass ReplayLimit, one per quiet keepAlive +// period behind a stuck writer would grow the ring without bound. +func (c *Conn) probeOnce() { + c.mu.Lock() + if c.lastProbe >= c.flushed { + c.mu.Unlock() + return + } + c.lastProbe = c.ring.next() + c.enqueueLocked(c.probe) + c.mu.Unlock() + signal(c.wake) +} + +// enqueueLocked seals p exactly once, so nonce order equals sequence order and +// a replay resends identical bytes. +func (c *Conn) enqueueLocked(p protocol.Packet) { + buf := make([]byte, 0, 2+secretbox.Overhead+len(p.Payload)) + data := c.out.Seal(wire.AppendPacket(buf, true, p.Header, nil), p.Payload) + c.ring.push(data) + c.unsent += len(data) +} + +// fail ends the Conn with err and returns it. +func (c *Conn) fail(err error) error { + c.cancel(err) + return err +} + +func signal(ch chan struct{}) { + select { + case ch <- struct{}{}: + default: + } +} diff --git a/internal/etcp/conn_test.go b/internal/etcp/conn_test.go new file mode 100644 index 0000000..8d37fa5 --- /dev/null +++ b/internal/etcp/conn_test.go @@ -0,0 +1,478 @@ +package etcp_test + +import ( + "context" + "errors" + "fmt" + "net" + "strings" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/tphakala/et-go/internal/etcp" + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/wire" +) + +func TestDialAndExchange(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + for i := range 3 { + if err := h.conn.WritePacket(t.Context(), numbered(i, 100)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + if err := h.srv.Send(t.Context(), numbered(i, 100)); err != nil { + t.Fatalf("Send: %v", err) + } + } + if err := expectNumbered(t.Context(), 3, h.srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + if err := expectNumbered(t.Context(), 3, h.conn.ReadPacket); err != nil { + t.Fatalf("client side: %v", err) + } + }) +} + +func TestCloseEndsReadsAndWrites(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + if err := h.conn.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + if _, err := readPacket(t, h.conn); !errors.Is(err, net.ErrClosed) { + t.Fatalf("ReadPacket after Close = %v, want net.ErrClosed", err) + } + if err := h.conn.WritePacket(t.Context(), numbered(0, 10)); !errors.Is(err, net.ErrClosed) { + t.Fatalf("WritePacket after Close = %v, want net.ErrClosed", err) + } + }) +} + +func TestDialErrors(t *testing.T) { + tests := []struct { + name string + status protocol.ConnectStatus + want error + }{ + {name: "unregistered id", status: protocol.ConnectStatus_INVALID_KEY, want: etcp.ErrRejected}, + {name: "protocol mismatch", status: protocol.ConnectStatus_MISMATCHED_PROTOCOL, want: etcp.ErrVersion}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(_ int, c *rawServer) { _ = c.respond(tt.status) }} + d := etcp.Dialer{NetDialer: s} + _, err := d.Dial(t.Context(), testAddr, testID, testKey) + if !errors.Is(err, tt.want) { + t.Fatalf("Dial error = %v, want %v", err, tt.want) + } + s.wg.Wait() + }) + }) + } +} + +// Options outside their documented range are refused before dialing: a +// negative KeepAlive makes the watcher fire at once (spinning before the +// first packet, then redialing forever), a negative ReplayLimit blocks every +// write, and a ReplayLimit above +// upstream's 64 MiB is refused as documented. The boundary values themselves +// are accepted. +func TestDialRejectsBadOptions(t *testing.T) { + tests := []struct { + name string + d etcp.Dialer + ok bool // the boundary values themselves are accepted + }{ + {name: "negative KeepAlive", d: etcp.Dialer{KeepAlive: -time.Second}}, + {name: "tiny KeepAlive", d: etcp.Dialer{KeepAlive: time.Nanosecond}}, + {name: "KeepAlive at its minimum", d: etcp.Dialer{KeepAlive: 100 * time.Millisecond}, ok: true}, + {name: "negative ReplayLimit", d: etcp.Dialer{ReplayLimit: -1}}, + {name: "ReplayLimit above 64 MiB", d: etcp.Dialer{ReplayLimit: 64<<20 + 1}}, + {name: "ReplayLimit at 64 MiB", d: etcp.Dialer{ReplayLimit: 64 << 20}, ok: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(_ int, c *rawServer) { + if c.respond(protocol.ConnectStatus_NEW_CLIENT) == nil { + c.drain() + } + }} + d := tt.d + d.NetDialer = s + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err == nil { + _ = conn.Close() + } + s.wg.Wait() + if tt.ok { + if err != nil { + t.Fatalf("Dial = %v; want the boundary value accepted", err) + } + return + } + if err == nil || s.dials.Load() != 0 { + t.Fatalf("Dial = %v after %d dials; want an error before dialing", err, s.dials.Load()) + } + }) + }) + } +} + +// When the caller's context ends during the handshake, Dial reports the +// context's cause, not the closed-connection error that ending it produced. +func TestDialContextCauseMidHandshake(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(_ int, c *rawServer) { + var req protocol.ConnectRequest + if wire.ReadMessage(c.br, &req) == nil { + c.drain() // never answer + } + }} + d := etcp.Dialer{NetDialer: s} + ctx, cancel := context.WithTimeout(t.Context(), 2*time.Second) + defer cancel() + conn, err := d.Dial(ctx, testAddr, testID, testKey) + if err == nil { + _ = conn.Close() + } + s.wg.Wait() + if !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("Dial = %v, want an error wrapping context.DeadlineExceeded", err) + } + }) +} + +// cancelOnReplyConn cancels its context inside the Read that delivers the +// first handshake reply's body (past its 8-byte length), so the cancellation +// deterministically lands just as the server's answer arrives. +type cancelOnReplyConn struct { + net.Conn + read int + cancel context.CancelCauseFunc + cause error +} + +func (c *cancelOnReplyConn) Read(p []byte) (int, error) { + n, err := c.Conn.Read(p) + c.read += n + if c.read > 8 { + c.cancel(c.cause) + } + return n, err +} + +type cancelOnReplyDialer struct { + inner *scripted + cancel context.CancelCauseFunc + cause error +} + +func (d cancelOnReplyDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + c, err := d.inner.DialContext(ctx, network, address) + if err != nil { + return nil, err + } + return &cancelOnReplyConn{Conn: c, cancel: d.cancel, cause: d.cause}, nil +} + +// A definitive answer from the server that arrives as the caller's context +// ends is reported, not hidden behind the context's cause. +func TestDialKeepsDefinitiveAnswerOverCause(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(_ int, c *rawServer) { _ = c.respond(protocol.ConnectStatus_INVALID_KEY) }} + errCause := errors.New("caller gave up") + ctx, cancel := context.WithCancelCause(t.Context()) + defer cancel(nil) + d := etcp.Dialer{NetDialer: cancelOnReplyDialer{inner: s, cancel: cancel, cause: errCause}} + _, err := d.Dial(ctx, testAddr, testID, testKey) + s.wg.Wait() + if ctx.Err() == nil { + t.Fatal("the context did not end during the handshake, so the test proved nothing") + } + if !errors.Is(err, etcp.ErrRejected) || errors.Is(err, errCause) { + t.Fatalf("Dial = %v, want ErrRejected, not the context's cause", err) + } + }) +} + +// blockingDialer never connects: it waits for ctx and returns its error, as +// net.Dialer does for a dial that has not finished. +type blockingDialer struct{} + +func (blockingDialer) DialContext(ctx context.Context, _, _ string) (net.Conn, error) { + <-ctx.Done() + return nil, ctx.Err() +} + +// The cause is reported the same way when ctx ends while the TCP dial itself +// is still in progress, not only during the handshake. +func TestDialContextCauseDuringDial(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + errCause := errors.New("caller gave up") + ctx, cancel := context.WithCancelCause(t.Context()) + go func() { + time.Sleep(time.Second) + cancel(errCause) + }() + d := etcp.Dialer{NetDialer: blockingDialer{}} + if _, err := d.Dial(ctx, testAddr, testID, testKey); !errors.Is(err, errCause) { + t.Fatalf("Dial = %v, want an error wrapping the context's cause", err) + } + }) +} + +func TestDialRejectsShortPasskey(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(_ int, c *rawServer) { + if c.respond(protocol.ConnectStatus_NEW_CLIENT) == nil { + c.drain() + } + }} + d := etcp.Dialer{NetDialer: s} + conn, err := d.Dial(t.Context(), testAddr, testID, "short") + if err == nil { + _ = conn.Close() + } + s.wg.Wait() + if err == nil || s.dials.Load() != 0 { + t.Fatalf("Dial with a 5-byte passkey = %v after %d dials; want an error before dialing", err, s.dials.Load()) + } + }) +} + +// On etserver 7.0.0 a shell exit closes the link and the redial gets +// INVALID_KEY. The shell's last output must still reach the reader first. +func TestSessionEndDeliversLastOutput(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + for i := range 20 { + if err := h.srv.Send(t.Context(), numbered(i, 50)); err != nil { + t.Fatalf("Send: %v", err) + } + } + h.srv.EndSession() + synctest.Wait() // the Conn has already failed when reading starts + if err := expectNumbered(t.Context(), 20, h.conn.ReadPacket); err != nil { + t.Fatalf("last output: %v", err) + } + if _, err := readPacket(t, h.conn); !errors.Is(err, etcp.ErrSessionEnded) { + t.Fatalf("after last output: %v, want ErrSessionEnded", err) + } + }) +} + +func TestReconnectAfterCut(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + for i := range 10 { + if i == 5 { + synctest.Wait() + h.net.CutAll() + } + if err := h.conn.WritePacket(t.Context(), numbered(i, 200)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + if err := h.srv.Send(t.Context(), numbered(i, 200)); err != nil { + t.Fatalf("Send: %v", err) + } + } + if err := expectNumbered(t.Context(), 10, h.srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + if err := expectNumbered(t.Context(), 10, h.conn.ReadPacket); err != nil { + t.Fatalf("client side: %v", err) + } + if got := h.net.Dials(); got != 2 { + t.Fatalf("Dials() = %d, want 2", got) + } + }) +} + +// Cutting the replacement connection after n bytes lands the cut inside the +// connect handshake, the sequence exchange, a catchup message, or a frame. +func TestReconnectSurvivesCutsAnywhere(t *testing.T) { + for _, n := range []int64{1, 9, 20, 40, 60, 100, 300, 1000, 5000} { + t.Run(fmt.Sprintf("bytes=%d", n), func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + for i := range 20 { + if i == 10 { + synctest.Wait() + h.net.CutAfter(n) + h.net.CutAll() + } + if err := h.conn.WritePacket(t.Context(), numbered(i, 300)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + if err := h.srv.Send(t.Context(), numbered(i, 300)); err != nil { + t.Fatalf("Send: %v", err) + } + } + ctx, cancel := context.WithTimeout(t.Context(), time.Hour) + defer cancel() + if err := expectNumbered(ctx, 20, h.srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + if err := expectNumbered(ctx, 20, h.conn.ReadPacket); err != nil { + t.Fatalf("client side: %v", err) + } + // The initial link (which CutAll ends), the budgeted link the + // cut ends, and the one after it: fewer than three dials means + // the budgeted cut never landed. + if got := h.net.Dials(); got < 3 { + t.Fatalf("Dials() = %d, want at least 3: the cut after %d bytes never fired", got, n) + } + }) + }) + } +} + +// A backlog larger than one link write batch is sent across several batches +// with nothing skipped or repeated. +// +// The server holds its first read until every packet is queued, so the +// writer's first Write blocks while the rest pile up behind it and at least +// one later batch must stop short of the backlog (100 frames of 2070 bytes +// do not fit one 64 KiB batch), whatever the scheduling. +func TestBacklogSpansWriteBatches(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const n = 100 + release := make(chan struct{}) + got := make(chan []int, 1) + s := &scripted{handle: func(i int, c *rawServer) { + if i != 0 || c.respond(protocol.ConnectStatus_NEW_CLIENT) != nil { + return + } + <-release + var nums []int + for len(nums) < n { + frame, err := wire.ReadFrame(c.br, nil) + if err != nil { + break + } + p, err := c.open(frame) + if err != nil { + break + } + k, err := number(p) + if err != nil { + break + } + nums = append(nums, k) + } + got <- nums + c.drain() + }} + d := etcp.Dialer{NetDialer: s} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { + _ = conn.Close() + s.wg.Wait() + }() + var releaseOnce sync.Once + releaseServer := func() { releaseOnce.Do(func() { close(release) }) } + defer releaseServer() // runs first, so an early failure cannot strand s.wg.Wait + ctx, cancel := context.WithTimeout(t.Context(), time.Hour) + defer cancel() + for i := range n { + if err := conn.WritePacket(ctx, numbered(i, 2048)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + } + synctest.Wait() + releaseServer() + + select { + case nums := <-got: + for i, k := range nums { + if k != i { + t.Fatalf("server got packet %d at position %d (lost, duplicated or reordered)", k, i) + } + } + if len(nums) != n { + t.Fatalf("server got %d packets, want %d", len(nums), n) + } + case <-time.After(time.Minute): + t.Fatal("server never received the backlog") + } + if got := s.dials.Load(); got != 1 { + t.Fatalf("dials = %d, want 1", got) + } + }) +} + +// A payload whose sealed frame cannot fit wire.MaxFrameSize is refused at +// once instead of being queued where no link could ever send it. +func TestWritePacketTooLarge(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + big := protocol.Packet{Header: protocol.HeaderTerminalBuffer, Payload: make([]byte, wire.MaxFrameSize)} + if err := h.conn.WritePacket(t.Context(), big); !errors.Is(err, wire.ErrTooLarge) { + t.Fatalf("WritePacket(%d bytes) = %v, want wire.ErrTooLarge", len(big.Payload), err) + } + if err := h.conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket after the refusal: %v", err) + } + if err := expectNumbered(t.Context(), 1, h.srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + }) +} + +// The session passkey never appears in Dial's errors. +func TestDialErrorsHidePasskey(t *testing.T) { + short := testKey[:31] + var d etcp.Dialer + if _, err := d.Dial(t.Context(), testAddr, testID, short); err == nil || strings.Contains(err.Error(), short) { + t.Fatalf("Dial with a 31-byte passkey = %v; want an error that does not quote it", err) + } + // The server's error text is peer-controlled: a server that knows the + // passkey could echo it, or send terminal escape sequences, so none of + // it may reach the error. + peerText := "rejected " + testKey + " \x1b]0;pwned\x07" + for _, status := range []protocol.ConnectStatus{ + protocol.ConnectStatus_INVALID_KEY, + protocol.ConnectStatus_MISMATCHED_PROTOCOL, + } { + t.Run(status.String(), func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(_ int, c *rawServer) { + var req protocol.ConnectRequest + if wire.ReadMessage(c.br, &req) != nil { + return + } + resp := &protocol.ConnectResponse{} + resp.SetStatus(status) + resp.SetError(peerText) + _ = wire.WriteMessage(c.conn, resp) + }} + d := etcp.Dialer{NetDialer: s} + _, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err == nil || strings.Contains(err.Error(), testKey) || strings.Contains(err.Error(), "\x1b") { + t.Fatalf("rejected Dial = %q; want an error without the server's text", err) + } + s.wg.Wait() + }) + }) + } +} diff --git a/internal/etcp/dialer.go b/internal/etcp/dialer.go new file mode 100644 index 0000000..af79354 --- /dev/null +++ b/internal/etcp/dialer.go @@ -0,0 +1,297 @@ +// Package etcp is a reliable, ordered, encrypted packet connection to an +// etserver that survives reconnects. Upstream calls this layer EternalTCP +// (src/base/BackedReader.cpp, BackedWriter.cpp, Connection.cpp). +package etcp + +import ( + "cmp" + "context" + "errors" + "fmt" + "io" + "log/slog" + "net" + "time" + + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/seal" + "github.com/tphakala/et-go/internal/wire" +) + +const ( + // defaultKeepAlive is upstream's maximum (src/base/Headers.hpp:180 at + // et-v7.0.0). + defaultKeepAlive = 5 * time.Second + // defaultReplayLimit is upstream's MAX_BACKUP_BYTES + // (src/base/BackedWriter.hpp:32 at et-v7.0.0). + defaultReplayLimit = 64 << 20 + dialTimeout = 10 * time.Second + + // minKeepAlive is the smallest non-zero KeepAlive Dial accepts, a policy + // floor: much below it the watcher would probe and drop links faster + // than a typical round trip, and a nanosecond value would spin. + minKeepAlive = 100 * time.Millisecond + // maxReplayLimit is the largest ReplayLimit Dial accepts: upstream's own + // MAX_BACKUP_BYTES. It does not by itself keep a catchup under the + // message limit (a disconnected ring can hold about twice the limit); + // writeRecover checks the size and fails with ErrReplayExceeded. + maxReplayLimit = defaultReplayLimit +) + +// ErrSessionEnded reports that the server ended the session, normally because +// the remote shell exited (a redial got INVALID_KEY). It wraps io.EOF, the +// end-of-stream convention of net.Conn, so errors.Is(err, io.EOF) holds. +var ErrSessionEnded = fmt.Errorf("etcp: session ended by server: %w", io.EOF) + +var ( + // ErrVersion reports that the server speaks another protocol version + // (MISMATCHED_PROTOCOL), on the first connect or a redial. + ErrVersion = errors.New("etcp: protocol version mismatch") + // ErrIntegrity reports a packet that failed authentication or could not + // be parsed, a frame above wire.MaxFrameSize, or a handshake message + // that is oversized or does not decode: the stream is out of step or + // tampered with and cannot be resumed. + ErrIntegrity = errors.New("etcp: stream integrity failure") + // ErrReplayExceeded reports that the peer needs packets no longer + // retained, is ahead of what was sent, or needs a catchup too large to + // send in one message. + ErrReplayExceeded = errors.New("etcp: peer needs data beyond the replay window") + // ErrRejected reports that the server refused the session: INVALID_KEY + // on the first connect, NEW_CLIENT on a redial, or an unknown status. + ErrRejected = errors.New("etcp: server rejected the session") +) + +// A Dialer contains options for connecting to an etserver. The zero value is usable. +type Dialer struct { + // NetDialer dials each TCP link. Nil means a net.Dialer with a 10 s timeout + // and TCP keepalive (KeepAliveConfig: idle 15 s, interval 5 s, count 3). + NetDialer interface { + DialContext(ctx context.Context, network, address string) (net.Conn, error) + } + // KeepAlive is the quiet period after which a probe is sent; after two + // quiet periods the link is declared dead. Only inbound frames count as + // proof of life, and the probe queues behind any unsent backlog, so an + // upload that takes longer than two periods to drain while the server + // sends nothing costs a reconnect (the data survives it). Zero means 5 s + // (upstream's maximum, src/base/Headers.hpp:180 at et-v7.0.0); Dial + // refuses a negative value or one below 100 ms. Probing starts only + // after the first WritePacket, because etserver aborts when a session's + // first packet is not INITIAL_PAYLOAD + // (src/terminal/TerminalServer.cpp:429-439 at et-v7.0.0); before that, a + // dead link is detected only by a read error or by TCP keepalive when + // the NetDialer enables it (the default one does). + KeepAlive time.Duration + // Probe is the packet sent as a liveness probe. The zero value sends header 0 + // with no payload, which is KEEP_ALIVE, echoed by etserver once the session + // runs (src/terminal/TerminalServer.cpp:389-393 at et-v7.0.0). + Probe protocol.Packet + // ReplayLimit bounds the sealed bytes kept for replay; Dial refuses a + // negative value or one above 64 MiB, and zero means 64 MiB. It is a soft + // bound applied separately to the two kinds of retained packets: packets + // already written to a socket are trimmed down to it, and WritePacket + // blocks while the not-yet-written backlog exceeds it (a single packet + // may take the backlog over, and probes are never blocked). While + // disconnected the ring can therefore hold about twice ReplayLimit plus + // one packet. Packets count as sent once written to the socket, and + // written packets are trimmed first, so a window smaller than what the + // kernel and the network can hold in flight turns a reconnect into + // ErrReplayExceeded; so does a catchup too large for one handshake + // message. Small values are for tests only. + ReplayLimit int + // Logger receives connection events. Nil discards them. + Logger *slog.Logger +} + +// Dial connects and completes the first handshake. It expects NEW_CLIENT; a +// RETURNING_CLIENT answer (upstream allows it when a first attempt died after +// registering, src/base/ClientConnection.cpp:35-39 at et-v7.0.0) runs the recover exchange +// with empty state. INVALID_KEY yields ErrRejected and MISMATCHED_PROTOCOL +// yields ErrVersion. A passkey that is not 32 bytes, a Probe payload too +// large for one sealed frame, or a KeepAlive or ReplayLimit outside the range +// its field documents is refused before dialing. +func (d *Dialer) Dial(ctx context.Context, addr, id, passkey string) (*Conn, error) { + if len(passkey) != 32 { + return nil, fmt.Errorf("etcp: passkey must be 32 bytes, got %d", len(passkey)) + } + if len(d.Probe.Payload) > maxPayload { + return nil, fmt.Errorf("etcp: probe payload of %d bytes: %w", len(d.Probe.Payload), wire.ErrTooLarge) + } + if d.KeepAlive < 0 || (d.KeepAlive > 0 && d.KeepAlive < minKeepAlive) { + return nil, fmt.Errorf("etcp: KeepAlive %v: must be 0 (the default) or at least %v", d.KeepAlive, minKeepAlive) + } + if d.ReplayLimit < 0 || d.ReplayLimit > maxReplayLimit { + return nil, fmt.Errorf("etcp: ReplayLimit %d: must be between 0 (the default) and %d", d.ReplayLimit, maxReplayLimit) + } + c := d.newConn(addr, id, passkey) + nc, catchup, err := c.connect(ctx, true) + if err != nil { + c.cancel(err) + return nil, err + } + c.wg.Go(func() { + defer close(c.readerDone) + c.supervise(nc, catchup) + }) + return c, nil +} + +func (d *Dialer) newConn(addr, id, passkey string) *Conn { + var nd netDialer = d.NetDialer + if nd == nil { + nd = &net.Dialer{ + Timeout: dialTimeout, + KeepAliveConfig: net.KeepAliveConfig{ + Enable: true, + Idle: 15 * time.Second, + Interval: 5 * time.Second, + Count: 3, + }, + } + } + logger := d.Logger + if logger == nil { + logger = slog.New(slog.DiscardHandler) + } + ctx, cancel := context.WithCancelCause(context.Background()) + c := &Conn{ + netDialer: nd, + addr: addr, + id: id, + keepAlive: cmp.Or(d.KeepAlive, defaultKeepAlive), + probe: d.Probe, + logger: logger, + ctx: ctx, + cancel: cancel, + inbox: make(chan protocol.Packet, inboxSize), + readerDone: make(chan struct{}), + wake: make(chan struct{}, 1), + space: make(chan struct{}, 1), + } + c.lastProbe = -1 + c.limit = cmp.Or(d.ReplayLimit, defaultReplayLimit) + c.ring.limit = c.limit + var key [32]byte + copy(key[:], passkey) + c.out = seal.New(&key, seal.ClientToServer) + c.in = seal.New(&key, seal.ServerToClient) + return c +} + +type netDialer interface { + DialContext(ctx context.Context, network, address string) (net.Conn, error) +} + +// connect dials one TCP link and runs the connect handshake on it (plus the +// recover exchange for a returning client). It returns the ready connection +// and the peer's catchup entries, which the new link delivers first. If ctx +// ends during the dial or the handshake, the error wraps context.Cause(ctx), +// unless a definitive error (one isFatal accepts, such as the server's +// rejection) came first. A handshake message that is oversized or does not +// decode yields ErrIntegrity. +func (c *Conn) connect(ctx context.Context, first bool) (net.Conn, [][]byte, error) { + dctx, cancel := context.WithTimeout(ctx, dialTimeout) + defer cancel() + nc, err := c.netDialer.DialContext(dctx, "tcp", c.addr) + if err != nil { + if ctx.Err() != nil { + // As for the handshake below: the caller needs its cause, not + // the dialer's generic context error. + return nil, nil, fmt.Errorf("etcp: dial: %w", context.Cause(ctx)) + } + return nil, nil, fmt.Errorf("etcp: dial: %w", err) + } + stop := context.AfterFunc(ctx, func() { _ = nc.Close() }) + defer stop() + catchup, err := c.handshake(idleConn{Conn: nc, timeout: handshakeIdle}, first) + if err == nil && !stop() { + // ctx ended as the handshake finished and the AfterFunc has + // closed nc: the link is unusable. + err = context.Cause(ctx) + } + if err == nil { + err = nc.SetDeadline(time.Time{}) + } + if err != nil { + _ = nc.Close() + if errors.Is(err, wire.ErrTooLarge) || errors.Is(err, wire.ErrMalformed) { + // A handshake message from the server that is oversized or + // does not decode means a broken or hostile peer (a cut stream + // yields an EOF error instead); as on the stream, the session + // cannot continue, and redialing would meet it again. Our own + // catchup cannot reach here: writeRecover checks its size + // first and fails with ErrReplayExceeded. + err = fmt.Errorf("%w: %w", ErrIntegrity, err) + } + if ctx.Err() != nil && !isFatal(err) { + // Ending ctx closes nc, so the handshake reports a closed + // connection; the cause is what the caller needs. A definitive + // answer that arrived just before (a rejection, a version + // mismatch) is kept: it says more than the deadline does. + return nil, nil, fmt.Errorf("etcp: connect: %w", context.Cause(ctx)) + } + return nil, nil, err + } + return nc, catchup, nil +} + +func (c *Conn) handshake(conn net.Conn, first bool) ([][]byte, error) { + req := &protocol.ConnectRequest{} + req.SetClientId(c.id) + req.SetVersion(protocol.Version) + if err := wire.WriteMessage(conn, req); err != nil { + return nil, fmt.Errorf("etcp: connect request: %w", err) + } + var resp protocol.ConnectResponse + if err := wire.ReadMessage(conn, &resp); err != nil { + return nil, fmt.Errorf("etcp: connect response: %w", err) + } + switch resp.GetStatus() { + case protocol.ConnectStatus_NEW_CLIENT: + if first { + return nil, nil + } + // Deliberately stricter than upstream, whose client closes the + // socket and keeps redialing on any status other than INVALID_KEY + // and RETURNING_CLIENT (src/base/ClientConnection.cpp:113-127 at + // et-v7.0.0). This applies to NEW_CLIENT here and equally to + // MISMATCHED_PROTOCOL and unknown statuses below: a server that + // has lost the session's state or changed protocol will not answer + // differently on the next redial. + return nil, fmt.Errorf("%w: server answered NEW_CLIENT to a returning client", ErrRejected) + case protocol.ConnectStatus_RETURNING_CLIENT: + // Also possible on the first connect, when an earlier attempt died + // after the server registered it. Upstream's client accepts that + // answer but then skips the recover exchange + // (src/base/ClientConnection.cpp:35-62 at et-v7.0.0), although the + // server always runs it after RETURNING_CLIENT + // (src/base/ServerConnection.cpp:117-123). + // Running it here with empty state (nothing sent, nothing received) + // matches what the server expects. + return c.recover(conn) + // The response's error text is never quoted: it is peer-controlled, so + // it could carry the passkey or terminal escape sequences, and upstream + // only puts fixed sentences there (src/base/ServerConnection.cpp:55-59, + // 93 at et-v7.0.0). + case protocol.ConnectStatus_INVALID_KEY: + if first { + return nil, fmt.Errorf("%w: server does not know this client", ErrRejected) + } + return nil, ErrSessionEnded + case protocol.ConnectStatus_MISMATCHED_PROTOCOL: + return nil, fmt.Errorf("%w: server does not speak protocol version %d", ErrVersion, protocol.Version) + default: + return nil, fmt.Errorf("%w: unexpected status %v", ErrRejected, resp.GetStatus()) + } +} + +// isFatal reports whether a reconnect error ends the Conn instead of +// triggering a retry. Integrity failures on the stream never reach here (the +// reader ends the Conn directly); ones in a handshake message do. +func isFatal(err error) bool { + for _, target := range []error{ErrSessionEnded, ErrVersion, ErrReplayExceeded, ErrRejected, ErrIntegrity} { + if errors.Is(err, target) { + return true + } + } + return false +} diff --git a/internal/etcp/drain_bench_test.go b/internal/etcp/drain_bench_test.go new file mode 100644 index 0000000..d4e65c5 --- /dev/null +++ b/internal/etcp/drain_bench_test.go @@ -0,0 +1,44 @@ +package etcp + +import ( + "context" + "strings" + "testing" +) + +// countingWriter cancels once it has taken want bytes. +type countingWriter struct { + n, want int + done context.CancelFunc +} + +func (w *countingWriter) Write(p []byte) (int, error) { + w.n += len(p) + if w.n >= w.want { + w.done() + } + return len(p), nil +} + +// BenchmarkDrainBacklog drains a backlog of minimum-size entries, the shape a +// long outage with small interactive packets leaves, and fails only if the +// writer stops before the backlog is drained. It measures one backlog size; +// compare runs before and after a writer change to see its cost (taking the +// whole backlog on every writer pass made the drain quadratic). +func BenchmarkDrainBacklog(b *testing.B) { + const entries = 500_000 + for b.Loop() { + var d Dialer + c := d.newConn("et.example:2022", "XXXtestclient001", strings.Repeat("k", 32)) + for range entries { + c.ring.push(make([]byte, 18)) + c.unsent += 18 + } + ctx, cancel := context.WithCancel(b.Context()) + w := &countingWriter{want: entries * (4 + 18), done: cancel} + if err := c.writeLoop(ctx, w); err == nil || ctx.Err() == nil { + b.Fatalf("writeLoop = %v before draining %d bytes (took %d)", err, w.want, w.n) + } + c.cancel(nil) + } +} diff --git a/internal/etcp/export_test.go b/internal/etcp/export_test.go new file mode 100644 index 0000000..016b75a --- /dev/null +++ b/internal/etcp/export_test.go @@ -0,0 +1,24 @@ +package etcp + +// SetMaxCatchupSize lowers the catchup size limit for a test, so the limit +// can be reached without writing 128 MiB, and returns a function that +// restores it. +func SetMaxCatchupSize(n int) (restore func()) { + old := maxCatchupSize + maxCatchupSize = n + return func() { maxCatchupSize = old } +} + +// RingBytes reports how many bytes of sealed packets c retains for replay. +func RingBytes(c *Conn) int { + c.mu.Lock() + defer c.mu.Unlock() + return c.ring.bytes +} + +// Unsent reports the bytes of sealed packets c has not yet written to a link. +func Unsent(c *Conn) int { + c.mu.Lock() + defer c.mu.Unlock() + return c.unsent +} diff --git a/internal/etcp/handshake_test.go b/internal/etcp/handshake_test.go new file mode 100644 index 0000000..e65d3b4 --- /dev/null +++ b/internal/etcp/handshake_test.go @@ -0,0 +1,451 @@ +package etcp_test + +import ( + "context" + "encoding/binary" + "errors" + "fmt" + "net" + "sync/atomic" + "syscall" + "testing" + "testing/synctest" + "time" + + "github.com/tphakala/et-go/internal/etcp" + "github.com/tphakala/et-go/internal/etservertest" + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/wire" + "golang.org/x/crypto/nacl/secretbox" +) + +// The largest payload whose sealed frame fits wire.MaxFrameSize is sent +// intact; one byte more is refused before it is queued. +func TestWritePacketPayloadBoundary(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + largest := wire.MaxFrameSize - 2 - secretbox.Overhead + over := protocol.Packet{Header: protocol.HeaderTerminalBuffer, Payload: make([]byte, largest+1)} + if err := h.conn.WritePacket(t.Context(), over); !errors.Is(err, wire.ErrTooLarge) { + t.Fatalf("WritePacket(%d bytes) = %v, want wire.ErrTooLarge", largest+1, err) + } + fits := protocol.Packet{Header: protocol.HeaderTerminalBuffer, Payload: make([]byte, largest)} + fits.Payload[largest-1] = 0x5a + if err := h.conn.WritePacket(t.Context(), fits); err != nil { + t.Fatalf("WritePacket(%d bytes): %v", largest, err) + } + ctx, cancel := context.WithTimeout(t.Context(), time.Minute) + defer cancel() + p, err := h.srv.Recv(ctx) + if err != nil { + t.Fatalf("Recv: %v", err) + } + if len(p.Payload) != largest || p.Payload[largest-1] != 0x5a { + t.Fatalf("server got %d bytes, want %d intact", len(p.Payload), largest) + } + }) +} + +// A RETURNING_CLIENT answer to the first connect (a first attempt that died +// after the server registered it) runs the recover exchange with empty +// state, in upstream's order, and the session then works. +func TestDialFirstReturningClient(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + got := make(chan int, 1) + s := &scripted{handle: func(i int, c *rawServer) { + if i != 0 { + return + } + if c.respond(protocol.ConnectStatus_RETURNING_CLIENT) != nil { + return + } + // Upstream writes each message before reading the peer's + // (src/base/Connection.cpp:105-143 at et-v7.0.0). + if wire.WriteMessage(c.conn, &protocol.SequenceHeader{}) != nil { + return + } + var theirSeq protocol.SequenceHeader + if wire.ReadMessage(c.br, &theirSeq) != nil || theirSeq.GetSequenceNumber() != 0 { + return + } + if wire.WriteMessage(c.conn, &protocol.CatchupBuffer{}) != nil { + return + } + var theirs protocol.CatchupBuffer + if wire.ReadMessage(c.br, &theirs) != nil || len(theirs.GetBuffer()) != 0 { + return + } + frame, err := wire.ReadFrame(c.br, nil) + if err != nil { + return + } + if p, err := c.open(frame); err == nil { + if n, err := number(p); err == nil { + got <- n + } + } + c.drain() + }} + d := etcp.Dialer{NetDialer: s} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial after RETURNING_CLIENT: %v", err) + } + defer func() { + _ = conn.Close() + s.wg.Wait() + }() + if err := conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + select { + case n := <-got: + if n != 0 { + t.Fatalf("server got packet %d, want 0", n) + } + case <-time.After(time.Minute): + t.Fatal("no packet after a first-connect recover exchange") + } + }) +} + +// On a redial, NEW_CLIENT means the server lost the session and +// MISMATCHED_PROTOCOL a server upgrade; neither can be recovered, so the +// Conn ends instead of redialing. +func TestRedialStatusIsFatal(t *testing.T) { + tests := []struct { + name string + status protocol.ConnectStatus + want error + }{ + {name: "new client", status: protocol.ConnectStatus_NEW_CLIENT, want: etcp.ErrRejected}, + {name: "protocol mismatch", status: protocol.ConnectStatus_MISMATCHED_PROTOCOL, want: etcp.ErrVersion}, + {name: "unknown status", status: protocol.ConnectStatus(99), want: etcp.ErrRejected}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(i int, c *rawServer) { + if i == 0 { + // Accept, then drop the link after the test's first + // packet: closing at once would race Dial's own + // SetDeadline on the pipe. + c.acceptFrames(1) + return + } + _ = c.respond(tt.status) + }} + d := etcp.Dialer{NetDialer: s} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { + _ = conn.Close() + s.wg.Wait() + }() + if err := conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + ctx, cancel := context.WithTimeout(t.Context(), time.Minute) + defer cancel() + if _, err := conn.ReadPacket(ctx); !errors.Is(err, tt.want) { + t.Fatalf("ReadPacket = %v, want %v", err, tt.want) + } + if got := s.dials.Load(); got != 2 { + t.Fatalf("dials = %d, want 2: the redial status must not be retried", got) + } + }) + }) + } +} + +// The zero-value Dialer dials real TCP with net.Dialer. A loopback listener +// served by the fake server exercises that path end to end. +func TestDefaultNetDialer(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("Listen: %v", err) + } + defer func() { _ = ln.Close() }() + srv := etservertest.NewServer(testID, testKey) + // Real sockets and real time: a failure waits out this bound. + ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) + defer cancel() + served := make(chan struct{}) + go func() { + defer close(served) + c, err := ln.Accept() + if err != nil { + return + } + _ = srv.Serve(ctx, c) + }() + + var d etcp.Dialer + conn, err := d.Dial(ctx, ln.Addr().String(), testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { + _ = conn.Close() + cancel() + <-served + }() + if err := conn.WritePacket(ctx, numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + if err := expectNumbered(ctx, 1, srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } +} + +// tcpStyleConn reports use of its own closed end as net.ErrClosed, as a TCP +// conn does, instead of net.Pipe's io.ErrClosedPipe. With reset set, a failed +// Write reports a connection reset instead, as a TCP write can when the peer +// that sent a bad message also dropped the connection. +type tcpStyleConn struct { + net.Conn + closed atomic.Bool + reset bool +} + +func (c *tcpStyleConn) Close() error { + c.closed.Store(true) + return c.Conn.Close() +} + +func (c *tcpStyleConn) Read(p []byte) (int, error) { + n, err := c.Conn.Read(p) + if err != nil && c.closed.Load() { + err = fmt.Errorf("tcp-style read: %w", net.ErrClosed) + } + return n, err +} + +func (c *tcpStyleConn) Write(p []byte) (int, error) { + n, err := c.Conn.Write(p) + switch { + case err != nil && c.reset: + err = fmt.Errorf("tcp-style write: %w", syscall.ECONNRESET) + case err != nil && c.closed.Load(): + err = fmt.Errorf("tcp-style write: %w", net.ErrClosed) + } + return n, err +} + +// The deadline setters are relabelled too, for fidelity only: idleConn calls +// them around every Read and Write, but a recover test whose stuck operation +// is a Write sees the relabelled Write error, not a setter's. +func (c *tcpStyleConn) SetDeadline(t time.Time) error { + return c.relabel(c.Conn.SetDeadline(t)) +} + +func (c *tcpStyleConn) SetReadDeadline(t time.Time) error { + return c.relabel(c.Conn.SetReadDeadline(t)) +} + +func (c *tcpStyleConn) SetWriteDeadline(t time.Time) error { + return c.relabel(c.Conn.SetWriteDeadline(t)) +} + +func (c *tcpStyleConn) relabel(err error) error { + if err != nil && c.closed.Load() { + return fmt.Errorf("tcp-style set deadline: %w", net.ErrClosed) + } + return err +} + +type tcpStyleDialer struct { + inner *scripted + reset bool +} + +func (d tcpStyleDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + c, err := d.inner.DialContext(ctx, network, address) + if err != nil { + return nil, err + } + return &tcpStyleConn{Conn: c, reset: d.reset}, nil +} + +// When our half of the recover exchange completes and only then the server's +// CatchupBuffer turns out bad, the exchange still fails with ErrIntegrity: +// reporting success would hand connect a conn the reader already closed, so +// the attempt would fail as an ordinary I/O error and the Conn would redial +// into the same bad message forever. +func TestRecoverReaderFailsAfterWriterSucceeds(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + bad := append(binary.LittleEndian.AppendUint64(nil, 3), 0xff, 0xff, 0xff) // does not decode + s := &scripted{handle: func(i int, c *rawServer) { + if i == 0 { + c.acceptFrames(1) + return + } + if c.respond(protocol.ConnectStatus_RETURNING_CLIENT) != nil { + return + } + var seq protocol.SequenceHeader + if wire.ReadMessage(c.br, &seq) != nil { + return + } + if wire.WriteMessage(c.conn, &protocol.SequenceHeader{}) != nil { + return + } + var theirs protocol.CatchupBuffer + if wire.ReadMessage(c.br, &theirs) != nil { // our whole half is now read + return + } + if _, err := c.conn.Write(bad); err == nil { + c.drain() + } + }} + d := etcp.Dialer{NetDialer: s} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { + _ = conn.Close() + s.wg.Wait() + }() + if err := conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + if _, err := readPacket(t, conn); !errors.Is(err, etcp.ErrIntegrity) { + t.Fatalf("ReadPacket = %v, want ErrIntegrity", err) + } + if got := s.dials.Load(); got != 2 { + t.Fatalf("dials = %d, want 2", got) + } + }) +} + +// As TestBadHandshakeMessageIsFatal below shows for the ConnectResponse, a +// bad SequenceHeader or CatchupBuffer in the recover exchange ends the Conn, +// even when our own write is still blocked because the server has not read +// it: in these rows the server never reads, so the stuck write is our +// SequenceHeader (upstream reads it before writing its catchup, +// src/base/Connection.cpp:116-131 at et-v7.0.0, but the reader must not +// depend on that). +func TestBadRecoverMessageIsFatal(t *testing.T) { + undecodable := append(binary.LittleEndian.AppendUint64(nil, 3), 0xff, 0xff, 0xff) + oversized := binary.LittleEndian.AppendUint64(nil, wire.MaxMessageSize+1) + emptySeq := binary.LittleEndian.AppendUint64(nil, 0) + tests := []struct { + name string + good [][]byte // valid messages sent before the bad one + bad []byte + tcp bool // report a local close as net.ErrClosed, as a TCP conn does + rst bool // report our failed write as a connection reset instead + }{ + {name: "sequence header", bad: undecodable}, + {name: "catchup buffer", good: [][]byte{emptySeq}, bad: undecodable}, + {name: "sequence header, TCP-style close", bad: undecodable, tcp: true}, + {name: "catchup buffer, TCP-style close", good: [][]byte{emptySeq}, bad: undecodable, tcp: true}, + {name: "catchup buffer, write reset", good: [][]byte{emptySeq}, bad: undecodable, tcp: true, rst: true}, + {name: "oversized catchup buffer, write reset", good: [][]byte{emptySeq}, bad: oversized, tcp: true, rst: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + quit := make(chan struct{}) + s := &scripted{handle: func(i int, c *rawServer) { + if i == 0 { + c.acceptFrames(1) + return + } + if c.respond(protocol.ConnectStatus_RETURNING_CLIENT) != nil { + return + } + for _, m := range append(tt.good, tt.bad) { + if _, err := c.conn.Write(m); err != nil { + return + } + } + <-quit // never read the client's messages + }} + d := etcp.Dialer{NetDialer: s} + if tt.tcp { + d.NetDialer = tcpStyleDialer{inner: s, reset: tt.rst} + } + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { + _ = conn.Close() + close(quit) + s.wg.Wait() + }() + start := time.Now() + if err := conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + if _, err := readPacket(t, conn); !errors.Is(err, etcp.ErrIntegrity) { + t.Fatalf("ReadPacket = %v, want ErrIntegrity", err) + } + if got := s.dials.Load(); got != 2 { + t.Fatalf("dials = %d, want 2: a bad recover message must not be retried", got) + } + // The reader's failure unblocks our stuck write at once (no + // fake time passes), not after the handshake idle timeout. + if elapsed := time.Since(start); elapsed >= time.Second { + t.Fatalf("recover failed after %v, want at once, not after the idle timeout", elapsed) + } + }) + }) + } +} + +// A redial whose ConnectResponse is oversized or does not decode can only +// come from a broken or hostile server; as on the stream, the Conn ends +// with ErrIntegrity instead of redialing forever. +func TestBadHandshakeMessageIsFatal(t *testing.T) { + tests := []struct { + name string + msg []byte + }{ + {name: "length above the limit", msg: binary.LittleEndian.AppendUint64(nil, wire.MaxMessageSize+1)}, + {name: "body that does not decode", msg: append(binary.LittleEndian.AppendUint64(nil, 3), 0xff, 0xff, 0xff)}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(i int, c *rawServer) { + if i == 0 { + c.acceptFrames(1) // drop the link after the test's first packet + return + } + var req protocol.ConnectRequest + if wire.ReadMessage(c.br, &req) != nil { + return + } + if _, err := c.conn.Write(tt.msg); err == nil { + c.drain() + } + }} + d := etcp.Dialer{NetDialer: s} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { + _ = conn.Close() + s.wg.Wait() + }() + if err := conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + if _, err := readPacket(t, conn); !errors.Is(err, etcp.ErrIntegrity) { + t.Fatalf("ReadPacket = %v, want ErrIntegrity", err) + } + if got := s.dials.Load(); got != 2 { + t.Fatalf("dials = %d, want 2: a bad handshake message must not be retried", got) + } + }) + }) + } +} diff --git a/internal/etcp/helpers_test.go b/internal/etcp/helpers_test.go new file mode 100644 index 0000000..016d1e2 --- /dev/null +++ b/internal/etcp/helpers_test.go @@ -0,0 +1,299 @@ +package etcp_test + +import ( + "bufio" + "context" + "errors" + "fmt" + "net" + "strconv" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/tphakala/et-go/internal/etcp" + "github.com/tphakala/et-go/internal/etservertest" + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/seal" + "github.com/tphakala/et-go/internal/wire" +) + +const ( + testID = "XXXtestclient001" + testKey = "0123456789abcdef0123456789abcdef" + testAddr = "et.example:2022" +) + +// harness is one session against the fake server over the in-memory network. +type harness struct { + srv *etservertest.Server + net *etservertest.Network + conn *etcp.Conn +} + +// newHarness dials a fresh session. Callers defer close, which closes the +// Conn and then the network, so every goroutine has exited before the +// synctest bubble ends. +func newHarness(t *testing.T, d etcp.Dialer) *harness { + t.Helper() + srv := etservertest.NewServer(testID, testKey) + nw := etservertest.NewNetwork(srv) + if d.NetDialer == nil { + d.NetDialer = nw + } + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + nw.Close() + t.Fatalf("Dial: %v", err) + } + return &harness{srv: srv, net: nw, conn: conn} +} + +func (h *harness) close() { + _ = h.conn.Close() + h.net.Close() +} + +// numbered builds a TERMINAL_BUFFER whose payload starts with i, so receivers +// can check order and completeness. +func numbered(i, size int) protocol.Packet { + payload := fmt.Appendf(nil, "%08d", i) + for len(payload) < size { + payload = append(payload, byte('a'+i%26)) + } + return protocol.Packet{Header: protocol.HeaderTerminalBuffer, Payload: payload} +} + +func number(p protocol.Packet) (int, error) { + if len(p.Payload) < 8 { + return 0, fmt.Errorf("payload %q too short", p.Payload) + } + return strconv.Atoi(string(p.Payload[:8])) +} + +// expectNumbered reads packets, skipping keepalives, and checks that 0..n-1 +// arrive in order, none lost or duplicated up to n-1. It stops there, so a +// duplicate after the last packet needs expectNothingMore. Without a deadline +// on ctx it waits at most an hour: in a synctest bubble a reconnect loop keeps +// fake time moving, so an unbounded wait would hang the suite instead of +// failing the test. +func expectNumbered(ctx context.Context, n int, recv func(context.Context) (protocol.Packet, error)) error { + if _, ok := ctx.Deadline(); !ok { + var cancel context.CancelFunc + ctx, cancel = context.WithTimeout(ctx, time.Hour) + defer cancel() + } + for want := 0; want < n; { + p, err := recv(ctx) + if err != nil { + return fmt.Errorf("after %d of %d packets: %w", want, n, err) + } + if p.Header == protocol.HeaderKeepAlive { + continue + } + got, err := number(p) + if err != nil { + return err + } + if got != want { + return fmt.Errorf("got packet %d, want %d (lost, duplicated or reordered)", got, want) + } + want++ + } + return nil +} + +// expectNothingMore reads for a minute of fake time and reports any packet +// other than a keepalive: a packet delivered again after the expected ones. +func expectNothingMore(ctx context.Context, recv func(context.Context) (protocol.Packet, error)) error { + ctx, cancel := context.WithTimeout(ctx, time.Minute) + defer cancel() + for { + p, err := recv(ctx) + if err != nil { + if errors.Is(err, context.DeadlineExceeded) { + return nil + } + return err + } + if p.Header != protocol.HeaderKeepAlive { + n, _ := number(p) + return fmt.Errorf("unexpected packet %d after the last one (duplicated)", n) + } + } +} + +// within returns the next value from ch, failing the test if none arrives in +// an hour of fake time, so a regression that strands a goroutine fails on an +// assertion instead of hanging the suite. +func within[T any](t *testing.T, ch <-chan T) T { + t.Helper() + select { + case v := <-ch: + return v + case <-time.After(time.Hour): + t.Fatal("no result within an hour") + var zero T + return zero + } +} + +// readPacket is conn.ReadPacket bounded to an hour of fake time. +func readPacket(t *testing.T, conn *etcp.Conn) (protocol.Packet, error) { + t.Helper() + ctx, cancel := context.WithTimeout(t.Context(), time.Hour) + defer cancel() + return conn.ReadPacket(ctx) +} + +// scripted is a NetDialer whose server side is a test function, for server +// behaviour the fake server does not produce on purpose. +type scripted struct { + wg sync.WaitGroup + dials atomic.Int32 + handle func(i int, c *rawServer) +} + +func (s *scripted) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + i := int(s.dials.Add(1)) - 1 + client, server := net.Pipe() + s.wg.Go(func() { + defer func() { _ = server.Close() }() + var key [32]byte + copy(key[:], testKey) + s.handle(i, &rawServer{ + conn: server, + br: bufio.NewReader(server), + out: seal.New(&key, seal.ServerToClient), + in: seal.New(&key, seal.ClientToServer), + }) + }) + return client, nil +} + +// rawServer is the server end of a scripted connection. +type rawServer struct { + conn net.Conn + br *bufio.Reader + out *seal.Stream + in *seal.Stream +} + +func (s *rawServer) respond(status protocol.ConnectStatus) error { + var req protocol.ConnectRequest + if err := wire.ReadMessage(s.br, &req); err != nil { + return err + } + resp := &protocol.ConnectResponse{} + resp.SetStatus(status) + return wire.WriteMessage(s.conn, resp) +} + +// drain reads and discards frames until the connection closes. +func (s *rawServer) drain() { + var buf []byte + for { + frame, err := wire.ReadFrame(s.br, buf) + if err != nil { + return + } + buf = frame + } +} + +// acceptFrames answers NEW_CLIENT, reads n frames, then returns so the link +// drops. +func (s *rawServer) acceptFrames(n int) { + if s.respond(protocol.ConnectStatus_NEW_CLIENT) != nil { + return + } + var buf []byte + for range n { + frame, err := wire.ReadFrame(s.br, buf) + if err != nil { + return + } + buf = frame + } +} + +// claimSequence answers RETURNING_CLIENT and claims to have received peer +// packets, then drains. +func (s *rawServer) claimSequence(peer int32) { + if s.respond(protocol.ConnectStatus_RETURNING_CLIENT) != nil { + return + } + var mine protocol.SequenceHeader + if wire.ReadMessage(s.br, &mine) != nil { + return + } + sh := &protocol.SequenceHeader{} + sh.SetSequenceNumber(peer) + if wire.WriteMessage(s.conn, sh) != nil { + return + } + s.drain() +} + +// pausedRecover answers RETURNING_CLIENT, claims to have received nothing, +// reads the client's catchup (so its snapshot is taken), then closes snapped +// and waits for resume before finishing the exchange. It reports the number +// of the first data packet the new link carries on got. +func (s *rawServer) pausedRecover(snapped chan<- struct{}, resume <-chan struct{}, got chan<- int) { + if s.respond(protocol.ConnectStatus_RETURNING_CLIENT) != nil { + return + } + var mine protocol.SequenceHeader + if wire.ReadMessage(s.br, &mine) != nil { + return + } + if wire.WriteMessage(s.conn, &protocol.SequenceHeader{}) != nil { + return + } + var theirs protocol.CatchupBuffer + if wire.ReadMessage(s.br, &theirs) != nil { + return + } + for _, b := range theirs.GetBuffer() { + if _, err := s.open(b); err != nil { // keeps the nonce in step + return + } + } + close(snapped) + <-resume + if wire.WriteMessage(s.conn, &protocol.CatchupBuffer{}) != nil { + return + } + for { + frame, err := wire.ReadFrame(s.br, nil) + if err != nil { + return + } + p, err := s.open(frame) + if err != nil { + return + } + if p.Header == protocol.HeaderKeepAlive { + continue + } + if n, err := number(p); err == nil { + got <- n + } + s.drain() + return + } +} + +// open parses and opens one sealed packet from the client. +func (s *rawServer) open(b []byte) (protocol.Packet, error) { + _, h, payload, err := wire.ParsePacket(b) + if err != nil { + return protocol.Packet{}, err + } + plain, err := s.in.Open(nil, payload) + if err != nil { + return protocol.Packet{}, err + } + return protocol.Packet{Header: h, Payload: plain}, nil +} diff --git a/internal/etcp/idleconn.go b/internal/etcp/idleconn.go new file mode 100644 index 0000000..d3852ce --- /dev/null +++ b/internal/etcp/idleconn.go @@ -0,0 +1,62 @@ +package etcp + +import ( + "net" + "time" +) + +const ( + handshakeIdle = 30 * time.Second + idleChunk = 64 << 10 +) + +// idleConn gives handshake I/O an idle timeout that resets on progress, like +// upstream's SocketHandler (src/base/SocketHandler.cpp:6,14,39,68 at +// et-v7.0.0). Every Read gets a fresh deadline, and Write sends in idleChunk +// pieces with a fresh deadline before each: net.Conn.Write transfers the whole +// slice in one call, so a single deadline set before it would limit the entire +// transfer. Write progress is seen per chunk, so a peer that drains fewer than +// idleChunk bytes per timeout still times out. +// +// Progress in either direction also pushes the other direction's deadline. In +// the recover exchange upstream writes its whole catchup before it reads ours +// (src/base/Connection.cpp:125-135 at et-v7.0.0), so once the socket +// buffers fill, our catchup write cannot progress while the server's catchup +// is still arriving; that write must not time out while those bytes flow. +// The connection is idle only when neither direction moves. +type idleConn struct { + net.Conn + timeout time.Duration +} + +func (c idleConn) Read(p []byte) (int, error) { + if err := c.SetReadDeadline(time.Now().Add(c.timeout)); err != nil { + return 0, err + } + n, err := c.Conn.Read(p) + if n > 0 { + // An error here means the conn is closed; the next operation reports it. + _ = c.SetWriteDeadline(time.Now().Add(c.timeout)) + } + return n, err +} + +func (c idleConn) Write(p []byte) (int, error) { + var n int + for len(p) > 0 { + if err := c.SetWriteDeadline(time.Now().Add(c.timeout)); err != nil { + return n, err + } + m, err := c.Conn.Write(p[:min(len(p), idleChunk)]) + n += m + if m > 0 { + // As in Read: a closed conn surfaces on the next operation. + _ = c.SetReadDeadline(time.Now().Add(c.timeout)) + } + if err != nil { + return n, err + } + p = p[m:] + } + return n, nil +} diff --git a/internal/etcp/idleconn_test.go b/internal/etcp/idleconn_test.go new file mode 100644 index 0000000..31bd25f --- /dev/null +++ b/internal/etcp/idleconn_test.go @@ -0,0 +1,132 @@ +package etcp + +import ( + "errors" + "io" + "net" + "os" + "testing" + "testing/synctest" + "time" +) + +// A slow but steady peer must never trip the idle timeout, however long the +// whole transfer takes. Before chunking, one deadline covered the entire +// Write and this failed after handshakeIdle. +func TestIdleConnSlowSteadyPeer(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + client, server := net.Pipe() + defer func() { _ = client.Close() }() + defer func() { _ = server.Close() }() + + const total = 1 << 20 // 1 MiB at 8 KiB/s takes 128 s, far past the 30 s timeout + done := make(chan error, 1) + go func() { + buf := make([]byte, 8<<10) + for read := 0; read < total; { + n, err := server.Read(buf) + if err != nil { + done <- err + return + } + read += n + time.Sleep(time.Second) + } + done <- nil + }() + + start := time.Now() + ic := idleConn{Conn: client, timeout: handshakeIdle} + if _, err := ic.Write(make([]byte, total)); err != nil { + t.Fatalf("Write: %v", err) + } + if err := <-done; err != nil { + t.Fatalf("reader: %v", err) + } + if elapsed := time.Since(start); elapsed <= handshakeIdle { + t.Fatalf("transfer took %v; the test needs it to outlast the %v timeout", elapsed, handshakeIdle) + } + }) +} + +// Write progress also keeps a blocked read alive: in the recover exchange +// the reader may wait for the peer's catchup while our own catchup is still +// flowing out. A read that outlasts the idle timeout must not fail while +// writes progress, and must fail once they stop. +func TestIdleConnWriteProgressKeepsReadAlive(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + client, server := net.Pipe() + defer func() { _ = client.Close() }() + defer func() { _ = server.Close() }() + + ic := idleConn{Conn: client, timeout: handshakeIdle} + readErr := make(chan error, 1) + go func() { + _, err := io.ReadFull(ic, make([]byte, 1)) // the peer never writes + readErr <- err + }() + stop := make(chan struct{}) + drained := make(chan struct{}) + defer func() { // no goroutine outlives a failed assertion + close(stop) + _ = server.Close() // ends a drainer blocked in Read + <-drained + }() + go func() { // the peer drains 8 KiB/s + defer close(drained) + buf := make([]byte, 8<<10) + for { + if _, err := server.Read(buf); err != nil { + return + } + select { + case <-stop: + return + case <-time.After(time.Second): + } + } + }() + + const total = 512 << 10 // 64 s at 8 KiB/s, past the 30 s timeout + if _, err := ic.Write(make([]byte, total)); err != nil { + t.Fatalf("Write: %v", err) + } + select { + case err := <-readErr: + t.Fatalf("read failed with %v while writes were still progressing", err) + default: + } + select { + case err := <-readErr: + if !errors.Is(err, os.ErrDeadlineExceeded) { + t.Fatalf("read error = %v, want os.ErrDeadlineExceeded once writes stop", err) + } + case <-time.After(time.Minute): + t.Fatal("read never timed out after writes stopped") + } + }) +} + +func TestIdleConnStalledPeer(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + client, server := net.Pipe() + defer func() { _ = client.Close() }() + defer func() { _ = server.Close() }() + + ic := idleConn{Conn: client, timeout: handshakeIdle} + start := time.Now() + _, err := ic.Write([]byte("nobody reads this")) + if !errors.Is(err, os.ErrDeadlineExceeded) { + t.Fatalf("Write error = %v, want os.ErrDeadlineExceeded", err) + } + if elapsed := time.Since(start); elapsed != handshakeIdle { + t.Fatalf("timed out after %v, want %v", elapsed, handshakeIdle) + } + + start = time.Now() + _, err = io.ReadFull(ic, make([]byte, 1)) + if !errors.Is(err, os.ErrDeadlineExceeded) || time.Since(start) != handshakeIdle { + t.Fatalf("Read error = %v after %v, want os.ErrDeadlineExceeded after %v", err, time.Since(start), handshakeIdle) + } + }) +} diff --git a/internal/etcp/lifecycle_test.go b/internal/etcp/lifecycle_test.go new file mode 100644 index 0000000..ca78792 --- /dev/null +++ b/internal/etcp/lifecycle_test.go @@ -0,0 +1,314 @@ +package etcp_test + +import ( + "bytes" + "context" + "encoding/binary" + "errors" + "log/slog" + "net" + "strings" + "testing" + "testing/synctest" + "time" + + "github.com/tphakala/et-go/internal/etcp" + "github.com/tphakala/et-go/internal/etservertest" + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/wire" +) + +// A link that lived past the reset period restarts the backoff schedule, so +// the first redial after it drops is immediate even when an earlier outage +// had pushed the delay to its cap. +func TestBackoffResetsAfterLongLink(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + srv := etservertest.NewServer(testID, testKey) + nw := etservertest.NewNetwork(srv) + defer nw.Close() + clock := &dialClock{inner: nw} + d := etcp.Dialer{NetDialer: clock} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { _ = conn.Close() }() + if err := conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + + synctest.Wait() + nw.SetRefuse(true) + nw.CutAll() + synctest.Sleep(time.Minute) // the backoff reaches its 5 s cap + nw.SetRefuse(false) + synctest.Sleep(10 * time.Second) // one capped delay: reconnected + synctest.Sleep(31 * time.Second) // the new link outlives the 30 s reset period + + clock.mu.Lock() + before := len(clock.times) + clock.mu.Unlock() + nw.CutAll() + // No fake time passes in synctest.Wait, so a redial seen here + // happened at the moment of the cut, with no backoff delay. + synctest.Wait() + + clock.mu.Lock() + defer clock.mu.Unlock() + if len(clock.times) <= before { + t.Fatal("no immediate redial after the long-lived link was cut: the backoff was not reset") + } + }) +} + +// Close returns even while the reader is blocked handing a packet to a +// caller that has stopped reading, and ReadPacket still returns every packet +// that had arrived, the reader's included, before net.ErrClosed. +func TestCloseWithFullInbox(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + for i := range 200 { // more than the Conn buffers, so its reader blocks + if err := h.srv.Send(t.Context(), numbered(i, 10)); err != nil { + t.Fatalf("Send: %v", err) + } + } + synctest.Wait() + + done := make(chan error, 1) + go func() { done <- h.conn.Close() }() + if err := within(t, done); err != nil { + t.Fatalf("Close: %v", err) + } + // The inbox holds packets 0 to 63 and the reader had packet 64. + if err := expectNumbered(t.Context(), 65, h.conn.ReadPacket); err != nil { + t.Fatalf("after Close: %v", err) + } + if _, err := readPacket(t, h.conn); !errors.Is(err, net.ErrClosed) { + t.Fatalf("ReadPacket after the last packet = %v, want net.ErrClosed", err) + } + }) +} + +// A blocked ReadPacket returns its context's cause when the caller gives up. +func TestReadPacketContextCancel(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + errCause := errors.New("caller gave up") + ctx, cancel := context.WithCancelCause(t.Context()) + done := make(chan error, 1) + go func() { + _, err := h.conn.ReadPacket(ctx) + done <- err + }() + synctest.Wait() + cancel(errCause) + if err := within(t, done); !errors.Is(err, errCause) { + t.Fatalf("ReadPacket = %v, want %v", err, errCause) + } + }) +} + +// A frame length above wire.MaxFrameSize can only come from a broken or +// hostile peer; the Conn ends with ErrIntegrity instead of redialing. +func TestOversizedServerFrameIsFatal(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(_ int, c *rawServer) { + if c.respond(protocol.ConnectStatus_NEW_CLIENT) != nil { + return + } + _, _ = c.conn.Write(binary.BigEndian.AppendUint32(nil, wire.MaxFrameSize+1)) + c.drain() + }} + d := etcp.Dialer{NetDialer: s} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { + _ = conn.Close() + s.wg.Wait() + }() + if _, err := readPacket(t, conn); !errors.Is(err, etcp.ErrIntegrity) { + t.Fatalf("ReadPacket = %v, want ErrIntegrity", err) + } + synctest.Sleep(time.Minute) + if got := s.dials.Load(); got != 1 { + t.Fatalf("dials = %d, want 1: an oversized frame must not be retried", got) + } + }) +} + +// The Logger receives link events, and nothing logged carries the passkey. +func TestLoggerRecordsLinkEventsWithoutPasskey(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + var buf bytes.Buffer + logger := slog.New(slog.NewTextHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug})) + h := newHarness(t, etcp.Dialer{Logger: logger}) + defer h.close() + if err := h.conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + synctest.Wait() + h.net.CutAll() + if err := expectNumbered(t.Context(), 1, h.srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + synctest.Wait() + // Close joins every goroutine that logs, so buf is safe to read. + if err := h.conn.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + out := buf.String() + for _, want := range []string{"etcp: link lost", "etcp: link restored"} { + if !strings.Contains(out, want) { + t.Errorf("log lacks %q:\n%s", want, out) + } + } + if strings.Contains(out, testKey) { + t.Fatalf("log contains the passkey:\n%s", out) + } + }) +} + +// A write whose context was already cancelled is refused and never queued: +// the server's first packet is the next one written. +func TestCancelledWriteNotQueued(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + ctx, cancel := context.WithCancel(t.Context()) + cancel() + if err := h.conn.WritePacket(ctx, numbered(99, 10)); !errors.Is(err, context.Canceled) { + t.Fatalf("WritePacket with a cancelled context = %v, want context.Canceled", err) + } + if err := h.conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + if err := expectNumbered(t.Context(), 1, h.srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + }) +} + +// A write is admitted while the unsent backlog is at most ReplayLimit: with +// the backlog exactly at the limit, the next write still goes in. +func TestBackpressureBoundary(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const sealed = 1024 + 2 + 16 // payload, packet header, MAC + h := newHarness(t, etcp.Dialer{ReplayLimit: 3 * sealed}) + defer h.close() + + synctest.Wait() + h.net.SetRefuse(true) + h.net.CutAll() + synctest.Wait() + for i := range 3 { // the backlog is now exactly ReplayLimit + if err := h.conn.WritePacket(t.Context(), numbered(i, 1024)); err != nil { + t.Fatalf("WritePacket %d: %v", i, err) + } + } + done := make(chan error, 1) + go func() { done <- h.conn.WritePacket(t.Context(), numbered(3, 1024)) }() + synctest.Wait() + select { + case err := <-done: + if err != nil { + t.Fatalf("write at the limit: %v", err) + } + default: + t.Fatal("a write with the backlog exactly at ReplayLimit blocked") + } + h.net.SetRefuse(false) + if err := expectNumbered(t.Context(), 4, h.srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + }) +} + +// A packet a dead link left pending is still returned when the Conn then +// ends before another link takes it over: here the session ends during the +// outage, so the reconnect meets INVALID_KEY. It arrived before the Conn +// failed, so ReadPacket returns it, after the inbox and before the error. +func TestPendingPacketSurvivesSessionEnd(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + const n = 100 // more than the inbox holds + if err := h.conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket 0: %v", err) + } + for i := range n { + if err := h.srv.Send(t.Context(), numbered(i, 10)); err != nil { + t.Fatalf("Send %d: %v", i, err) + } + } + synctest.Wait() // the inbox is full and the reader holds the next packet + h.net.CutAll() + h.srv.EndSession() + synctest.Wait() + if err := h.conn.WritePacket(t.Context(), numbered(1, 10)); err != nil { + t.Fatalf("WritePacket 1: %v", err) + } + synctest.Wait() // the redial met INVALID_KEY and the Conn ended + + // The inbox holds packets 0 to 63 and the reader had packet 64. + if err := expectNumbered(t.Context(), 65, h.conn.ReadPacket); err != nil { + t.Fatalf("client side: %v", err) + } + if _, err := readPacket(t, h.conn); !errors.Is(err, etcp.ErrSessionEnded) { + t.Fatalf("ReadPacket after the last packet = %v, want ErrSessionEnded", err) + } + }) +} + +// A caller that has stopped reading must not stop the Conn from replacing a +// dead link: the reader, blocked handing a packet to a full inbox, has to +// let the link go so writes reach the server over the next one. The packet +// in its hand is already counted as received, so it must still be +// delivered, once and in order. +func TestReconnectWhileCallerNotReading(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + const n = 100 // more than the inbox holds + if err := h.conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket 0: %v", err) + } + for i := range n { + if err := h.srv.Send(t.Context(), numbered(i, 10)); err != nil { + t.Fatalf("Send %d: %v", i, err) + } + } + synctest.Wait() // the inbox is full and the reader is blocked + h.net.CutAll() + synctest.Wait() + if err := h.conn.WritePacket(t.Context(), numbered(1, 10)); err != nil { + t.Fatalf("WritePacket 1: %v", err) + } + if err := expectNumbered(t.Context(), 2, h.srv.Recv); err != nil { + t.Fatalf("server side, caller not reading: %v", err) + } + // Cut the replacement link too, while its reader is still trying + // to hand over the packet the first link left pending. + synctest.Wait() + h.net.CutAll() + synctest.Wait() + if err := h.conn.WritePacket(t.Context(), numbered(2, 10)); err != nil { + t.Fatalf("WritePacket 2: %v", err) + } + if p, err := h.srv.Recv(t.Context()); err != nil { + t.Fatalf("server side, second cut: %v", err) + } else if got, err := number(p); err != nil || got != 2 { + t.Fatalf("server side, second cut: got packet %d (%v), want 2", got, err) + } + if err := expectNumbered(t.Context(), n, h.conn.ReadPacket); err != nil { + t.Fatalf("client side: %v", err) + } + }) +} diff --git a/internal/etcp/link.go b/internal/etcp/link.go new file mode 100644 index 0000000..8fb212f --- /dev/null +++ b/internal/etcp/link.go @@ -0,0 +1,234 @@ +package etcp + +import ( + "bufio" + "context" + "errors" + "fmt" + "io" + "net" + "sync" + "sync/atomic" + "time" + + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/wire" + "golang.org/x/crypto/nacl/secretbox" +) + +const ( + readBufSize = 64 << 10 + writeBufSize = 64 << 10 + + // maxBatchEntries bounds how many ring entries one writer pass takes: a + // framed entry is at least 4+2+secretbox.Overhead bytes, so no pass can + // frame more than this many into writeBufSize. Taking the whole backlog + // instead would copy it under c.mu on every pass, quadratic in its length. + maxBatchEntries = writeBufSize/(4+2+secretbox.Overhead) + 1 +) + +var ( + errClosed = net.ErrClosed + errLinkDead = errors.New("etcp: no traffic from server") +) + +// link is the per-TCP-connection state shared by its goroutines. +type link struct { + alive chan struct{} // cap 1: a frame arrived + delivering atomic.Bool // the reader is blocked handing a packet to ReadPacket +} + +// runLink runs a reader, a writer and a liveness watcher on nc until one of +// them fails, then joins all three and returns the cause. +func (c *Conn) runLink(nc net.Conn, catchup [][]byte) error { + ctx, cancel := context.WithCancelCause(c.ctx) + defer cancel(nil) + stop := context.AfterFunc(ctx, func() { _ = nc.Close() }) + defer stop() + + l := &link{alive: make(chan struct{}, 1)} + var wg sync.WaitGroup + wg.Go(func() { cancel(c.readLoop(ctx, l, nc, catchup)) }) + wg.Go(func() { cancel(c.writeLoop(ctx, nc)) }) + wg.Go(func() { cancel(c.watch(ctx, l)) }) + wg.Wait() + _ = nc.Close() + return context.Cause(ctx) +} + +// readLoop hands over packets an earlier link left pending, then delivers the +// peer's catchup, then frames from the link. +func (c *Conn) readLoop(ctx context.Context, l *link, r io.Reader, catchup [][]byte) error { + for len(c.pending) > 0 { + if err := c.handOver(ctx, l, c.pending[0]); err != nil { + return err // pending stays as it is for the next link + } + c.pending[0] = protocol.Packet{} + c.pending = c.pending[1:] + } + c.pending = nil + for _, b := range catchup { + if err := c.deliver(ctx, l, b); err != nil { + return err + } + } + br := bufio.NewReaderSize(r, readBufSize) + var buf []byte + for { + frame, err := wire.ReadFrame(br, buf) + if err != nil { + if errors.Is(err, wire.ErrTooLarge) { + return c.fail(fmt.Errorf("%w: %w", ErrIntegrity, err)) + } + return fmt.Errorf("etcp: read: %w", err) + } + buf = frame + signal(l.alive) + if err := c.deliver(ctx, l, frame); err != nil { + return err + } + } +} + +// deliver opens one sealed packet and hands it to ReadPacket. A malformed or +// unauthenticated packet ends the Conn: the streams are out of sync or tampered +// with, and reconnecting would replay into the same state. After a failed +// Open the inbound nonce has advanced, so the stream cannot be resumed; +// upstream treats a failed decrypt as fatal too +// (src/base/CryptoHandler.cpp:40-42 at et-v7.0.0). +func (c *Conn) deliver(ctx context.Context, l *link, b []byte) error { + // Live frames are bounded by wire.ReadFrame; catchup entries arrive in + // one handshake message, so bound them here to the same limit. + if len(b) > wire.MaxFrameSize { + return c.fail(fmt.Errorf("%w: packet of %d bytes: %w", ErrIntegrity, len(b), wire.ErrTooLarge)) + } + encrypted, h, payload, err := wire.ParsePacket(b) + if err != nil { + return c.fail(fmt.Errorf("%w: %w", ErrIntegrity, err)) + } + if !encrypted { + return c.fail(fmt.Errorf("%w: unencrypted packet", ErrIntegrity)) + } + plain, err := c.in.Open(nil, payload) + if err != nil { + return c.fail(fmt.Errorf("%w: %w", ErrIntegrity, err)) + } + c.recvSeq++ + p := protocol.Packet{Header: h, Payload: plain} + if err := c.handOver(ctx, l, p); err != nil { + // Counted in recvSeq, so the server will not resend it: keep it + // for the next link reader rather than drop it. + c.pending = append(c.pending, p) + return err + } + return nil +} + +// handOver gives p to ReadPacket, or fails when the link ends first (ctx is +// the link's context). Waiting on the Conn alone would keep a dead link +// alive, and with it every write, until the caller read again. +func (c *Conn) handOver(ctx context.Context, l *link, p protocol.Packet) error { + l.delivering.Store(true) + defer l.delivering.Store(false) + select { + case c.inbox <- p: + return nil + case <-ctx.Done(): + return context.Cause(ctx) + } +} + +// writeLoop sends ring entries from flushed onwards, trimming the ring and +// releasing blocked writers as the backlog drains. It frames a batch of +// entries (at least one, then up to writeBufSize bytes) into one reused +// buffer and sends it with a single Write, so the steady state allocates +// nothing per packet. Entries are immutable once sealed, so they are framed +// outside the lock. +func (c *Conn) writeLoop(ctx context.Context, w io.Writer) error { + var ( + batch [][]byte + buf []byte + ) + for { + c.mu.Lock() + batch = c.ring.appendRange(batch[:0], c.flushed, min(c.ring.next(), c.flushed+maxBatchEntries)) + c.mu.Unlock() + if len(batch) == 0 { + select { + case <-c.wake: + continue + case <-ctx.Done(): + return context.Cause(ctx) + } + } + buf = buf[:0] + sent, n := 0, 0 + for _, f := range batch { + if sent > 0 && len(buf)+4+len(f) > writeBufSize { + break + } + var err error + // WritePacket refuses packets above wire.MaxFrameSize, so an + // error here means the ring is corrupt and no link can help. + if buf, err = wire.AppendFrame(buf, f); err != nil { + return c.fail(fmt.Errorf("%w: %w", ErrIntegrity, err)) + } + sent++ + n += len(f) + } + clear(batch) // drop references so trimmed entries can be collected + // A short count without an error breaks the io.Writer contract; + // counting the whole batch as sent would drop its tail from replay. + if m, err := w.Write(buf); err != nil || m != len(buf) { + if err == nil { + err = io.ErrShortWrite + } + return fmt.Errorf("etcp: write: %w", err) + } + c.mu.Lock() + c.flushed += int64(sent) + c.unsent -= n + c.ring.trim(c.flushed, c.unsent) + room := c.unsent <= c.limit + c.mu.Unlock() + if room { + signal(c.space) + } + } +} + +// watch declares the link dead after two quiet keepAlive periods, sending the +// probe after the first. Time the reader spends blocked on a slow ReadPacket +// caller does not count as silence, and neither does the time before the +// caller's first packet, when no probe may be sent. +func (c *Conn) watch(ctx context.Context, l *link) error { + t := time.NewTimer(c.keepAlive) + defer t.Stop() + probed := false + for { + select { + case <-ctx.Done(): + return context.Cause(ctx) + case <-l.alive: + probed = false + case <-t.C: + switch { + case l.delivering.Load(): + probed = false + case c.sealedNone(): + // No probe before the caller's first packet: etserver 7.0.0 + // aborts when a session's first packet is not INITIAL_PAYLOAD + // (src/terminal/TerminalServer.cpp:429-439 at et-v7.0.0). + // Until then a dead link is found only by a read error or + // by TCP keepalive when the NetDialer enables it. + probed = false + case probed: + return errLinkDead + default: + c.probeOnce() + probed = true + } + } + t.Reset(c.keepAlive) + } +} diff --git a/internal/etcp/link_internal_test.go b/internal/etcp/link_internal_test.go new file mode 100644 index 0000000..d0c7e49 --- /dev/null +++ b/internal/etcp/link_internal_test.go @@ -0,0 +1,122 @@ +package etcp + +import ( + "context" + "errors" + "io" + "strings" + "testing" + "testing/synctest" + "time" + + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/seal" + "github.com/tphakala/et-go/internal/wire" +) + +// shortWriter reports one byte less than it was given and no error, which +// breaks the io.Writer contract. +type shortWriter struct{} + +func (shortWriter) Write(p []byte) (int, error) { return max(len(p)-1, 0), nil } + +// A writer that reports a short count without an error must fail the link +// rather than count the whole batch as sent: the unsent tail would otherwise +// be lost from replay. +func TestWriteLoopShortWrite(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + var d Dialer + c := d.newConn("et.example:2022", "XXXtestclient001", strings.Repeat("k", 32)) + defer c.cancel(nil) + c.ring.push(make([]byte, 18)) + c.unsent += 18 + + // Bounded: a writeLoop that accepted the short write would wait for + // more data forever. + ctx, cancel := context.WithTimeout(t.Context(), time.Minute) + defer cancel() + err := c.writeLoop(ctx, shortWriter{}) + if !errors.Is(err, io.ErrShortWrite) { + t.Fatalf("writeLoop = %v, want io.ErrShortWrite", err) + } + if c.flushed != 0 { + t.Fatalf("flushed = %d after a short write, want 0", c.flushed) + } + }) +} + +// A packet larger than any frame the stream allows ends the Conn with +// ErrIntegrity even when it authenticates: catchup entries arrive inside one +// handshake message, so they must not bypass the frame limit that +// wire.ReadFrame applies to live frames. (Called directly: pushing a 16 MiB +// catchup through the fake network takes minutes under the race detector.) +func TestDeliverRefusesOversizedPacket(t *testing.T) { + var d Dialer + c := d.newConn("et.example:2022", "XXXtestclient001", strings.Repeat("k", 32)) + defer c.cancel(nil) + var key [32]byte + copy(key[:], strings.Repeat("k", 32)) + sealed := seal.New(&key, seal.ServerToClient).Seal(nil, make([]byte, wire.MaxFrameSize)) + b := wire.AppendPacket(nil, true, protocol.HeaderTerminalBuffer, sealed) + + err := c.deliver(t.Context(), &link{alive: make(chan struct{}, 1)}, b) + if !errors.Is(err, ErrIntegrity) { + t.Fatalf("deliver(%d-byte packet) = %v, want ErrIntegrity", len(b), err) + } +} + +// backlogWriter accepts its first Write, during which the caller fills the +// unsent backlog as a racing WritePacket would, then blocks until ctx ends. +type backlogWriter struct { + ctx context.Context + c *Conn + writes int +} + +func (w *backlogWriter) Write(p []byte) (int, error) { + w.writes++ + if w.writes > 1 { + <-w.ctx.Done() + return 0, w.ctx.Err() + } + w.c.mu.Lock() + for range 5 { + w.c.ring.push(make([]byte, 20)) + w.c.unsent += 20 + } + w.c.mu.Unlock() + return len(p), nil +} + +// The replay limit applies to written packets and to the unsent backlog +// separately: a full backlog must not trim the replay copies of packets just +// written, which may still be in flight and are needed after a cut. +func TestWriteLoopKeepsWrittenWhileBacklogFull(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + d := Dialer{ReplayLimit: 100} + c := d.newConn("et.example:2022", "XXXtestclient001", strings.Repeat("k", 32)) + defer c.cancel(nil) + for range 5 { // 100 bytes to write, exactly the limit + c.ring.push(make([]byte, 20)) + c.unsent += 20 + } + + ctx, cancel := context.WithTimeout(t.Context(), time.Minute) + defer cancel() + done := make(chan error, 1) + go func() { done <- c.writeLoop(ctx, &backlogWriter{ctx: ctx, c: c}) }() + synctest.Wait() // the first batch is written, the second is stuck + + c.mu.Lock() + first, flushed, unsent := c.ring.first, c.flushed, c.unsent + c.mu.Unlock() + if flushed != 5 || unsent != 100 { + t.Fatalf("flushed, unsent = %d, %d; want 5, 100", flushed, unsent) + } + if first != 0 { + t.Fatalf("ring.first = %d, want 0: the 100 written bytes are within the limit", first) + } + cancel() + <-done + }) +} diff --git a/internal/etcp/liveness_test.go b/internal/etcp/liveness_test.go new file mode 100644 index 0000000..2f4a6fd --- /dev/null +++ b/internal/etcp/liveness_test.go @@ -0,0 +1,220 @@ +package etcp_test + +import ( + "context" + "errors" + "testing" + "testing/synctest" + "time" + + "github.com/tphakala/et-go/internal/etcp" + "github.com/tphakala/et-go/internal/etservertest" + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/wire" + "golang.org/x/crypto/nacl/secretbox" +) + +// A probe still waiting to be written is not followed by another: probes +// bypass ReplayLimit, so while the writer is stuck on a server that has +// stopped reading, a new probe every quiet keepAlive period (the server's +// occasional output keeps the link from being declared dead) would grow the +// replay ring without bound. +func TestProbesDoNotPileUpBehindStuckWriter(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(i int, c *rawServer) { + if i > 0 || c.respond(protocol.ConnectStatus_NEW_CLIENT) != nil { + return + } + if _, err := wire.ReadFrame(c.br, nil); err != nil { // packet 0 + return + } + // Stop reading, but send a little output just slower than + // the keepalive period. + for { + time.Sleep(6 * time.Second) + sealed := c.out.Seal(nil, []byte("x")) + if wire.WriteFrame(c.conn, wire.AppendPacket(nil, true, protocol.HeaderTerminalBuffer, sealed)) != nil { + return + } + } + }} + d := etcp.Dialer{NetDialer: s} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + done := make(chan struct{}) + defer func() { + _ = conn.Close() + <-done + s.wg.Wait() + }() + go func() { // the caller keeps reading + defer close(done) + for { + if _, err := conn.ReadPacket(t.Context()); err != nil { + return + } + } + }() + if err := conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + time.Sleep(20 * time.Minute) + // One probe, sealed with an empty payload, may wait unwritten. + if got, want := etcp.Unsent(conn), 2+secretbox.Overhead; got > want { + t.Fatalf("unsent = %d bytes after 20 minutes behind a stuck writer, want at most %d (one probe)", got, want) + } + }) +} + +// etserver 7.0.0 aborts the whole server when a session's first packet is not +// INITIAL_PAYLOAD (src/terminal/TerminalServer.cpp:429-439 at et-v7.0.0), so +// no probe may go out before the caller's first packet, however long the link +// stays quiet, and that quiet is not taken for a dead link. +func TestNoProbeBeforeFirstPacket(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + synctest.Sleep(time.Minute) + if got := h.net.Dials(); got != 1 { + t.Fatalf("Dials() = %d after a quiet minute before any write, want 1", got) + } + if err := h.conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + ctx, cancel := context.WithTimeout(t.Context(), time.Minute) + defer cancel() + p, err := h.srv.Recv(ctx) + if err != nil { + t.Fatalf("Recv: %v", err) + } + if n, err := number(p); p.Header != protocol.HeaderTerminalBuffer || err != nil || n != 0 { + t.Fatalf("first packet the server got = header %v %q, want the caller's packet 0", p.Header, p.Payload) + } + }) +} + +func TestLivenessKeepsHealthyLink(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + if err := h.conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + + synctest.Sleep(time.Minute) + if got := h.net.Dials(); got != 1 { + t.Fatalf("Dials() = %d after a quiet minute with echoes, want 1", got) + } + ctx, cancel := context.WithTimeout(t.Context(), time.Minute) + defer cancel() + if err := expectNumbered(ctx, 1, h.srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + p, err := h.srv.Recv(ctx) + if err != nil || p.Header != protocol.HeaderKeepAlive { + t.Fatalf("server got %v, %v; want a KEEP_ALIVE probe after the packet", p.Header, err) + } + }) +} + +func TestLivenessDetectsDeadLink(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + h.srv.EchoKeepAlive(false) + if err := h.conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + + // Probe at 5 s, dead at 10 s, immediate redial. + synctest.Sleep(9 * time.Second) + if got := h.net.Dials(); got != 1 { + t.Fatalf("Dials() = %d at 9 s, want 1", got) + } + synctest.Sleep(2 * time.Second) + if got := h.net.Dials(); got != 2 { + t.Fatalf("Dials() = %d at 11 s, want 2", got) + } + }) +} + +// While the reader is blocked on a caller that is not reading, the link is +// not declared dead: that silence is ours, not the network's. +func TestLivenessIgnoresSlowReader(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + for i := range 200 { // more than the Conn buffers, so its reader blocks + if err := h.srv.Send(t.Context(), numbered(i, 10)); err != nil { + t.Fatalf("Send: %v", err) + } + } + synctest.Sleep(time.Minute) + if err := expectNumbered(t.Context(), 200, h.conn.ReadPacket); err != nil { + t.Fatal(err) + } + synctest.Wait() + if got := h.net.Dials(); got != 1 { + t.Fatalf("Dials() = %d, want 1: a slow reader is not a dead link", got) + } + }) +} + +// A probe that could never be sent is refused up front, not discovered when +// the first quiet period ends the Conn. +func TestDialRejectsOversizedProbe(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(_ int, c *rawServer) { + if c.respond(protocol.ConnectStatus_NEW_CLIENT) == nil { + c.drain() + } + }} + d := etcp.Dialer{ + NetDialer: s, + Probe: protocol.Packet{Payload: make([]byte, wire.MaxFrameSize)}, + } + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err == nil { + _ = conn.Close() + } + s.wg.Wait() + if !errors.Is(err, wire.ErrTooLarge) || s.dials.Load() != 0 { + t.Fatalf("Dial = %v after %d dials; want wire.ErrTooLarge before dialing", err, s.dials.Load()) + } + }) +} + +// A long upload over a slow uplink queues the probe behind the backlog, so +// its echo comes late while the server itself sends nothing, and the +// watcher may drop the link (the KeepAlive doc states this limit). Whatever +// reconnects that costs, every packet must still arrive exactly once and in +// order. +func TestLivenessSlowUploadDeliversEverything(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + srv := etservertest.NewServer(testID, testKey) + nw := etservertest.NewNetwork(srv) + defer nw.Close() + d := etcp.Dialer{NetDialer: throttledDialer{inner: nw}} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { _ = conn.Close() }() + + const n = 400 // 400 KiB at 8 KiB/s: 50 s, ten keepalive periods + for i := range n { + if err := conn.WritePacket(t.Context(), numbered(i, 1024)); err != nil { + t.Fatalf("WritePacket %d: %v", i, err) + } + } + ctx, cancel := context.WithTimeout(t.Context(), time.Hour) + defer cancel() + if err := expectNumbered(ctx, n, srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + }) +} diff --git a/internal/etcp/outage_test.go b/internal/etcp/outage_test.go new file mode 100644 index 0000000..9cb0bc6 --- /dev/null +++ b/internal/etcp/outage_test.go @@ -0,0 +1,88 @@ +package etcp_test + +import ( + "context" + "net" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/tphakala/et-go/internal/etcp" + "github.com/tphakala/et-go/internal/etservertest" +) + +// dialClock records when each dial happens. +type dialClock struct { + inner interface { + DialContext(ctx context.Context, network, address string) (net.Conn, error) + } + mu sync.Mutex + times []time.Time +} + +func (d *dialClock) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + d.mu.Lock() + d.times = append(d.times, time.Now()) + d.mu.Unlock() + return d.inner.DialContext(ctx, network, address) +} + +// A long outage (laptop asleep, network down). Dials keep +// failing, the delay between them never exceeds 5 s, and once the network +// returns the session resumes with nothing lost. +func TestLongOutage(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + srv := etservertest.NewServer(testID, testKey) + nw := etservertest.NewNetwork(srv) + defer nw.Close() + clock := &dialClock{inner: nw} + d := etcp.Dialer{NetDialer: clock} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { _ = conn.Close() }() + + synctest.Wait() + nw.SetRefuse(true) + nw.CutAll() + synctest.Wait() + if err := conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + if err := srv.Send(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("Send: %v", err) + } + // 500 refused dials take about 40 minutes of fake time at the 5 s + // backoff cap: a long outage by any measure. The deadline turns a + // backoff that stopped dialing into a failure instead of a hang. + deadline := time.Now().Add(2 * time.Hour) + for nw.Dials() < 501 { + if time.Now().After(deadline) { + t.Fatalf("only %d dials in two hours of outage", nw.Dials()) + } + synctest.Sleep(5 * time.Second) + } + nw.SetRefuse(false) + if err := expectNumbered(t.Context(), 1, srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + if err := expectNumbered(t.Context(), 1, conn.ReadPacket); err != nil { + t.Fatalf("client side: %v", err) + } + + clock.mu.Lock() + defer clock.mu.Unlock() + var longest time.Duration + for i := 1; i < len(clock.times); i++ { + longest = max(longest, clock.times[i].Sub(clock.times[i-1])) + } + if longest > 5*time.Second { + t.Fatalf("longest gap between dials = %v, want at most 5s", longest) + } + if longest < 4*time.Second { + t.Fatalf("longest gap between dials = %v; backoff never reached its cap", longest) + } + }) +} diff --git a/internal/etcp/property_test.go b/internal/etcp/property_test.go new file mode 100644 index 0000000..23b7c55 --- /dev/null +++ b/internal/etcp/property_test.go @@ -0,0 +1,109 @@ +package etcp_test + +import ( + "context" + "fmt" + "math/rand/v2" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/tphakala/et-go/internal/etcp" +) + +// Random traffic in both directions while the network is cut at random +// moments, including in the middle of frames, handshakes and catchups. Every +// packet must arrive exactly once and in order on both sides. +func TestPropertyExactlyOnceUnderRandomCuts(t *testing.T) { + for seed := range uint64(20) { + t.Run(fmt.Sprintf("seed=%d", seed), func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + runProperty(t, seed) + }) + }) + } +} + +func runProperty(t *testing.T, seed uint64) { + t.Helper() + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + const n = 300 + ctx, cancel := context.WithTimeout(t.Context(), time.Hour) + defer cancel() + + var wg sync.WaitGroup + wg.Go(func() { // client to server + rng := rand.New(rand.NewPCG(seed, 1)) + for i := range n { + if err := h.conn.WritePacket(ctx, numbered(i, rng.IntN(2048))); err != nil { + t.Errorf("WritePacket %d: %v", i, err) + return + } + time.Sleep(time.Duration(rng.IntN(5)) * time.Millisecond) + } + }) + wg.Go(func() { // server to client + rng := rand.New(rand.NewPCG(seed, 2)) + for i := range n { + if err := h.srv.Send(ctx, numbered(i, rng.IntN(2048))); err != nil { + t.Errorf("Send %d: %v", i, err) + return + } + time.Sleep(time.Duration(rng.IntN(5)) * time.Millisecond) + } + }) + wg.Go(func() { // the network + rng := rand.New(rand.NewPCG(seed, 3)) + for range 15 { + time.Sleep(time.Duration(rng.IntN(100)) * time.Millisecond) + if rng.IntN(2) == 0 { + h.net.CutAll() + } else { + h.net.CutAfter(rng.Int64N(20_000)) + h.net.CutAll() + } + } + }) + + errs := make(chan error, 2) + wg.Go(func() { errs <- expectNumbered(ctx, n, h.srv.Recv) }) + wg.Go(func() { errs <- expectNumbered(ctx, n, h.conn.ReadPacket) }) + for range 2 { + if err := <-errs; err != nil { + t.Error(err) + } + } + cancel() + wg.Wait() + + if got := h.net.Dials(); got < 2 { + t.Errorf("Dials() = %d: no cut forced a reconnect, so the run proved nothing", got) + } + + // expectNumbered stops at packet n-1, so a recover after the last packet + // is otherwise never exercised. Force one now that both sides hold + // everything: it must replay nothing and must settle, with no further + // redial across the quiet minute. With counter nonces a resent packet + // fails the receiver's authentication: the client ends its Conn with + // ErrIntegrity (a read error below), while the fake server just drops + // the link, which shows up as a redial loop in the dial count. + beforeCut := h.net.Dials() + h.net.CutAll() + synctest.Sleep(time.Minute) + settled := h.net.Dials() + if settled == beforeCut { + t.Errorf("the late cut forced no reconnect (Dials() stayed %d)", settled) + } + if err := expectNothingMore(t.Context(), h.srv.Recv); err != nil { + t.Errorf("server side after the late recover: %v", err) + } + if err := expectNothingMore(t.Context(), h.conn.ReadPacket); err != nil { + t.Errorf("client side after the late recover: %v", err) + } + if got := h.net.Dials(); got != settled { + t.Errorf("Dials() went from %d to %d with no traffic and no cuts: a redial loop", settled, got) + } +} diff --git a/internal/etcp/recover.go b/internal/etcp/recover.go new file mode 100644 index 0000000..195c53f --- /dev/null +++ b/internal/etcp/recover.go @@ -0,0 +1,183 @@ +package etcp + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "sync" + "time" + + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/wire" + "google.golang.org/protobuf/proto" +) + +// supervise owns reconnects: it runs a link until it dies, then redials with +// backoff and runs the next one, until the Conn ends. +func (c *Conn) supervise(nc net.Conn, catchup [][]byte) { + var b backoff + for { + started := time.Now() + cause := c.runLink(nc, catchup) + if c.ctx.Err() != nil { + return + } + c.logger.Info("etcp: link lost", "cause", cause) + if time.Since(started) >= backoffReset { + b.reset() + } + var err error + nc, catchup, err = c.reconnect(&b) + if err != nil { + c.cancel(err) + return + } + c.logger.Info("etcp: link restored") + } +} + +// reconnect dials until a link is recovered, the error is fatal, or the Conn ends. +func (c *Conn) reconnect(b *backoff) (net.Conn, [][]byte, error) { + for { + if d := b.next(); d > 0 { + t := time.NewTimer(d) + select { + case <-t.C: + case <-c.ctx.Done(): + t.Stop() + return nil, nil, context.Cause(c.ctx) + } + } + nc, catchup, err := c.connect(c.ctx, false) + if err == nil { + return nc, catchup, nil + } + if isFatal(err) || c.ctx.Err() != nil { + return nil, nil, err + } + c.logger.Debug("etcp: reconnect attempt failed", "err", err) + } +} + +// recover runs the client side of the recover exchange and returns the peer's +// catchup entries. Upstream has each side write its SequenceHeader, read the +// peer's, write its CatchupBuffer, then read the peer's +// (src/base/Connection.cpp:105-143). Because both sides write before they +// read at every step, this client reads the peer's two messages concurrently +// with writing its own: done in sequence, the exchange deadlocks as soon as +// the connection cannot buffer what both sides write (at once over net.Pipe, +// and over TCP once both catchups exceed the socket buffers). Whichever side +// fails first closes conn to unblock the other; readerErrWins decides which +// error is reported. +func (c *Conn) recover(conn net.Conn) ([][]byte, error) { + var ( + peer protocol.SequenceHeader + theirs protocol.CatchupBuffer + gotSeq = make(chan error, 1) + gotAll = make(chan error, 1) + wg sync.WaitGroup + ) + wg.Go(func() { + err := wire.ReadMessage(conn, &peer) + gotSeq <- err + if err == nil { + err = wire.ReadMessage(conn, &theirs) + } + if err != nil { + // Unblock the writer: upstream reads our messages only after + // writing its own, so ours may be stuck until the idle timeout. + _ = conn.Close() + } + gotAll <- err + }) + + snap, err := c.writeRecover(conn, gotSeq, &peer) + if err != nil { + _ = conn.Close() // unblock the reader + } + rerr := <-gotAll + wg.Wait() + if rerr != nil && (err == nil || readerErrWins(err, rerr)) { + return nil, fmt.Errorf("etcp: read recover message: %w", rerr) + } + if err != nil { + return nil, err + } + + c.mu.Lock() + c.unsent -= c.ring.bytesBetween(c.flushed, snap) + c.flushed = snap + // Trim as writeLoop does after a Write: a link that dies before its + // first Write would otherwise leave the catchup in the ring while + // WritePacket admits another ReplayLimit. + c.ring.trim(c.flushed, c.unsent) + room := c.unsent <= c.limit + c.mu.Unlock() + if room { + signal(c.space) + } + return theirs.GetBuffer(), nil +} + +// readerErrWins reports whether the recover reader's error rerr, rather than +// the writer's err, explains a failed exchange. It does when rerr marks a +// bad message from the server (which connect turns into ErrIntegrity): over +// TCP the same bad peer may also reset the connection, and the writer's +// reset error can arrive before the reader's close takes effect. It also +// does when the writer merely hit the conn that the reader's failure +// closed. Otherwise the writer failed on its own (an I/O error, or +// ErrReplayExceeded, which is fatal either way) and its error stands. +func readerErrWins(err, rerr error) bool { + if errors.Is(rerr, wire.ErrTooLarge) || errors.Is(rerr, wire.ErrMalformed) { + return true + } + return errors.Is(err, net.ErrClosed) || errors.Is(err, io.ErrClosedPipe) +} + +// maxCatchupSize is the largest CatchupBuffer we send: etserver refuses +// handshake messages above 128 MiB (src/base/SocketHandler.hpp:60 at +// et-v7.0.0), the same bound as wire.MaxMessageSize. It is a variable only +// so a test can lower it. +var maxCatchupSize = wire.MaxMessageSize + +// writeRecover writes our half of the recover exchange and returns the +// sequence number the new link starts sending from. It fails with +// ErrReplayExceeded when the peer's position is outside the retained window +// or our catchup is too large for one message. +func (c *Conn) writeRecover(conn net.Conn, gotSeq <-chan error, peer *protocol.SequenceHeader) (int64, error) { + mine := &protocol.SequenceHeader{} + mine.SetSequenceNumber(int32(c.recvSeq)) + if err := wire.WriteMessage(conn, mine); err != nil { + return 0, fmt.Errorf("etcp: write sequence: %w", err) + } + if err := <-gotSeq; err != nil { + return 0, fmt.Errorf("etcp: read sequence: %w", err) + } + + // Read ring.next(), the send sequence, exactly once. Packets written + // after this snapshot are not in our catchup; the new link sends them + // because it starts at snap. + c.mu.Lock() + snap := c.ring.next() + ours, ok := c.ring.since(int64(peer.GetSequenceNumber()), snap) + first := c.ring.first + c.mu.Unlock() + if !ok { + return 0, fmt.Errorf("%w: peer is at %d, retained window is [%d, %d]", + ErrReplayExceeded, peer.GetSequenceNumber(), first, snap) + } + cb := &protocol.CatchupBuffer{} + cb.SetBuffer(ours) + // A catchup the server cannot accept would fail the same way on every + // redial, so it ends the Conn instead. + if size := proto.Size(cb); size > maxCatchupSize { + return 0, fmt.Errorf("%w: catchup of %d bytes exceeds the %d-byte message limit", + ErrReplayExceeded, size, maxCatchupSize) + } + if err := wire.WriteMessage(conn, cb); err != nil { + return 0, fmt.Errorf("etcp: write catchup: %w", err) + } + return snap, nil +} diff --git a/internal/etcp/recover_test.go b/internal/etcp/recover_test.go new file mode 100644 index 0000000..6630c48 --- /dev/null +++ b/internal/etcp/recover_test.go @@ -0,0 +1,387 @@ +package etcp_test + +import ( + "context" + "errors" + "io" + "sync" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + "github.com/tphakala/et-go/internal/etcp" + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/wire" + "golang.org/x/crypto/nacl/secretbox" +) + +func TestSessionEndedIsEOF(t *testing.T) { + if !errors.Is(etcp.ErrSessionEnded, io.EOF) { + t.Fatal("ErrSessionEnded does not wrap io.EOF") + } + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + h.srv.EndSession() + _, err := readPacket(t, h.conn) + if !errors.Is(err, etcp.ErrSessionEnded) || !errors.Is(err, io.EOF) { + t.Fatalf("ReadPacket = %v, want ErrSessionEnded wrapping io.EOF", err) + } + }) +} + +// Both sides have a catchup to replay. The fake server writes its whole +// catchup before reading ours, as upstream does, and net.Pipe buffers nothing, +// so this completes only if the client reads concurrently. +func TestCatchupBothWays(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + synctest.Wait() + h.net.SetRefuse(true) + h.net.CutAll() + synctest.Wait() + for i := range 50 { + if err := h.conn.WritePacket(t.Context(), numbered(i, 4096)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + if err := h.srv.Send(t.Context(), numbered(i, 4096)); err != nil { + t.Fatalf("Send: %v", err) + } + } + h.net.SetRefuse(false) + + ctx, cancel := context.WithTimeout(t.Context(), time.Minute) + defer cancel() + if err := expectNumbered(ctx, 50, h.srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + if err := expectNumbered(ctx, 50, h.conn.ReadPacket); err != nil { + t.Fatalf("client side: %v", err) + } + }) +} + +// Writes keep arriving across repeated cuts and reconnects; none may be +// skipped or sent twice. Recovery over net.Pipe takes no fake time, so the +// cuts rarely land inside a recover exchange; TestWritePacketRacingRecovery +// pins that race deterministically. +func TestWriteDuringRecovery(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + h := newHarness(t, etcp.Dialer{}) + defer h.close() + + const n = 2000 + var wg sync.WaitGroup + wg.Go(func() { + for i := range n { + if err := h.conn.WritePacket(t.Context(), numbered(i, 64)); err != nil { + t.Errorf("WritePacket %d: %v", i, err) + return + } + if i%100 == 0 { + time.Sleep(10 * time.Millisecond) + } + } + }) + wg.Go(func() { + for range 10 { + time.Sleep(15 * time.Millisecond) + h.net.CutAll() + } + }) + if err := expectNumbered(t.Context(), n, h.srv.Recv); err != nil { + t.Fatal(err) + } + wg.Wait() + }) +} + +// The server claims to have received less than we still hold (or more than +// we ever sent): the session cannot be recovered and the Conn ends. +func TestReplayWindowExceeded(t *testing.T) { + tests := []struct { + name string + peer int32 + }{ + {name: "peer behind the window", peer: 0}, + {name: "peer ahead of us", peer: 1_000_000}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(i int, c *rawServer) { + if i == 0 { + c.acceptFrames(100) // the client writes and trims its ring + return + } + c.claimSequence(tt.peer) + }} + d := etcp.Dialer{NetDialer: s, ReplayLimit: 1 << 10} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { + _ = conn.Close() + s.wg.Wait() + }() + for i := range 100 { // 100 KiB through a 1 KiB window + if err := conn.WritePacket(t.Context(), numbered(i, 1024)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + } + ctx, cancel := context.WithTimeout(t.Context(), time.Minute) + defer cancel() + if _, err := conn.ReadPacket(ctx); !errors.Is(err, etcp.ErrReplayExceeded) { + t.Fatalf("ReadPacket = %v, want ErrReplayExceeded", err) + } + }) + }) + } +} + +// A catchup too large for one handshake message (etserver refuses messages +// above wire.MaxMessageSize, src/base/SocketHandler.hpp:60 at et-v7.0.0) can +// never be sent, so the Conn ends with ErrReplayExceeded instead of redialing +// forever. The limit is lowered so a few KiB reach it. +func TestOversizedCatchupIsFatal(t *testing.T) { + defer etcp.SetMaxCatchupSize(1 << 10)() + synctest.Test(t, func(t *testing.T) { + release := make(chan struct{}) + s := &scripted{handle: func(i int, c *rawServer) { + if i == 0 { + // Accept the link but read nothing until released, so the + // packets below stay unsent and land in the catchup. + if c.respond(protocol.ConnectStatus_NEW_CLIENT) == nil { + <-release + } + return + } + c.claimSequence(0) + }} + d := etcp.Dialer{NetDialer: s} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { + _ = conn.Close() + s.wg.Wait() + }() + var releaseOnce sync.Once + releaseServer := func() { releaseOnce.Do(func() { close(release) }) } + defer releaseServer() // runs first, so an early failure cannot strand s.wg.Wait + for i := range 8 { // about 8 KiB of catchup against a 1 KiB limit + if err := conn.WritePacket(t.Context(), numbered(i, 1024)); err != nil { + t.Fatalf("WritePacket: %v", err) + } + } + synctest.Wait() + releaseServer() + + ctx, cancel := context.WithTimeout(t.Context(), time.Minute) + defer cancel() + if _, err := conn.ReadPacket(ctx); !errors.Is(err, etcp.ErrReplayExceeded) { + t.Fatalf("ReadPacket = %v, want ErrReplayExceeded", err) + } + if got := s.dials.Load(); got != 2 { + t.Fatalf("dials = %d, want 2: an oversized catchup must not be retried", got) + } + }) +} + +// A packet that fails authentication ends the Conn; it is not retried. +func TestIntegrityFailureIsFatal(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + s := &scripted{handle: func(_ int, c *rawServer) { + if c.respond(protocol.ConnectStatus_NEW_CLIENT) != nil { + return + } + garbage := wire.AppendPacket(nil, true, protocol.HeaderTerminalBuffer, make([]byte, 40)) + _ = wire.WriteFrame(c.conn, garbage) + c.drain() + }} + d := etcp.Dialer{NetDialer: s} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { + _ = conn.Close() + s.wg.Wait() + }() + if _, err := readPacket(t, conn); !errors.Is(err, etcp.ErrIntegrity) { + t.Fatalf("ReadPacket = %v, want ErrIntegrity", err) + } + synctest.Sleep(time.Minute) + if got := s.dials.Load(); got != 1 { + t.Fatalf("dials = %d, want 1: integrity failures must not be retried", got) + } + }) +} + +// A packet written after the recover snapshot is not in our catchup, so the +// new link must send it. Reading ring.next() a second time at the end of the +// exchange would skip it. +// +// In the second row the packets written during recovery exceed ReplayLimit, +// so recover's trim must stop at the snapshot: trimming past it would drop +// packets no link has sent. +func TestWritePacketRacingRecovery(t *testing.T) { + tests := []struct { + name string + limit int + during int // packets written between the snapshot and resume + }{ + {name: "one packet", during: 1}, + {name: "more than ReplayLimit", limit: 64, during: 3}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + snapped := make(chan struct{}) + resume := make(chan struct{}) + got := make(chan int, 1) + s := &scripted{handle: func(i int, c *rawServer) { + switch i { + case 0: + c.acceptFrames(1) // take packet 0 off the wire, then drop the link + case 1: + c.pausedRecover(snapped, resume, got) + default: + // Later links (the watcher may redial, since drain + // never echoes probes) end at once; got already holds + // packet 1's number, so they cannot change the verdict. + } + }} + d := etcp.Dialer{NetDialer: s, ReplayLimit: tt.limit} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { + _ = conn.Close() + s.wg.Wait() + }() + if err := conn.WritePacket(t.Context(), numbered(0, 10)); err != nil { + t.Fatalf("WritePacket 0: %v", err) + } + <-snapped + for i := 1; i <= tt.during; i++ { + if err := conn.WritePacket(t.Context(), numbered(i, 10)); err != nil { + t.Fatalf("WritePacket %d: %v", i, err) + } + } + close(resume) + + ctx, cancel := context.WithTimeout(t.Context(), time.Minute) + defer cancel() + select { + case n := <-got: + if n != 1 { + t.Fatalf("new link sent packet %d first, want 1", n) + } + case <-ctx.Done(): + t.Fatal("packet written during recovery never arrived") + } + // The new link has sent everything, so recover must have + // counted the packets written during it exactly once. + synctest.Wait() + if got := etcp.Unsent(conn); got != 0 { + t.Fatalf("unsent = %d bytes after the link drained, want 0", got) + } + }) + }) + } +} + +// A link that recovers and then dies before writing anything must not let the +// replay ring grow past ReplayLimit: recover counts our catchup as sent, which +// frees WritePacket to admit another ReplayLimit of packets, so recover must +// also trim what it has now written. The server here reports its +// true received count in each SequenceHeader and drops every link right after +// the exchange, while the caller keeps writing. +func TestFlappingLinkKeepsRingBounded(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + const ( + limit = 1 << 10 + payload = 100 + sealed = 2 + secretbox.Overhead + payload // one ring entry + ) + var received, caughtUp, lastCatchup atomic.Int32 + s := &scripted{handle: func(i int, c *rawServer) { + if i == 0 { + if c.respond(protocol.ConnectStatus_NEW_CLIENT) != nil { + return + } + if _, err := wire.ReadFrame(c.br, nil); err == nil { + received.Store(1) + } + return + } + if c.respond(protocol.ConnectStatus_RETURNING_CLIENT) != nil { + return + } + var mine protocol.SequenceHeader + if wire.ReadMessage(c.br, &mine) != nil { + return + } + sh := &protocol.SequenceHeader{} + sh.SetSequenceNumber(received.Load()) + if wire.WriteMessage(c.conn, sh) != nil { + return + } + var theirs protocol.CatchupBuffer + if wire.ReadMessage(c.br, &theirs) != nil { + return + } + received.Add(int32(len(theirs.GetBuffer()))) + caughtUp.Add(int32(len(theirs.GetBuffer()))) + lastCatchup.Store(int32(len(theirs.GetBuffer()))) + _ = wire.WriteMessage(c.conn, &protocol.CatchupBuffer{}) + }} + d := etcp.Dialer{NetDialer: s, ReplayLimit: limit} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + var wg sync.WaitGroup + defer func() { + _ = conn.Close() + wg.Wait() + s.wg.Wait() + }() + wg.Go(func() { + for i := 0; ; i++ { + if conn.WritePacket(t.Context(), numbered(i, payload)) != nil { + return + } + } + }) + // Cap the wait in fake time: a regression that ends the Conn stops + // the redials, and an uncapped loop would hang the suite. + deadline := time.Now().Add(time.Hour) + for s.dials.Load() < 12 { + if time.Now().After(deadline) { + t.Fatalf("only %d dials after an hour", s.dials.Load()) + } + time.Sleep(time.Second) + } + if caughtUp.Load() == 0 { + t.Fatal("no catchup reached the server; the test exercises nothing") + } + // The caller writes between every pair of links, so an empty latest + // catchup means WritePacket stayed blocked: recover must release it. + if lastCatchup.Load() == 0 { + t.Fatal("writer stalled: the last recover carried no catchup") + } + // At most ReplayLimit of written entries survive a trim, plus the + // unsent backlog WritePacket admits: ReplayLimit and one packet. + if got, bound := etcp.RingBytes(conn), 2*limit+sealed; got > bound { + t.Fatalf("ring holds %d bytes after %d links, want at most %d", got, s.dials.Load(), bound) + } + }) +} diff --git a/internal/etcp/ring.go b/internal/etcp/ring.go new file mode 100644 index 0000000..061a359 --- /dev/null +++ b/internal/etcp/ring.go @@ -0,0 +1,63 @@ +package etcp + +import "slices" + +// ring holds sealed, serialized outbound packets by sequence number, for the +// link writer and for replay after a reconnect. Sequence numbers are +// contiguous: entries[i] has sequence number first+i. +type ring struct { + first int64 + entries [][]byte + bytes int + limit int +} + +// next is the sequence number the next pushed packet gets, which is also the +// number of packets sealed so far. +func (r *ring) next() int64 { return r.first + int64(len(r.entries)) } + +func (r *ring) push(data []byte) { + r.entries = append(r.entries, data) + r.bytes += len(data) +} + +// trim drops the oldest written entries while they hold more than limit +// bytes, never dropping an entry at or after keep (not yet written to any +// link). unsent is the size of the entries from keep on; the limit applies +// to written entries alone, so a full backlog cannot squeeze out the replay +// copies of packets still in flight. +func (r *ring) trim(keep int64, unsent int) { + n := 0 + for n < len(r.entries) && r.bytes-unsent > r.limit && r.first+int64(n) < keep { + r.bytes -= len(r.entries[n]) + r.entries[n] = nil + n++ + } + r.entries = r.entries[n:] + r.first += int64(n) +} + +// appendRange appends entries [from, to) to dst. The caller guarantees +// first <= from <= to <= next(). +func (r *ring) appendRange(dst [][]byte, from, to int64) [][]byte { + return append(dst, r.entries[from-r.first:to-r.first]...) +} + +// since returns a copy of entries [from, to), or false when from is outside +// the retained window [first, to] or to is past next(). +func (r *ring) since(from, to int64) ([][]byte, bool) { + if from < r.first || from > to || to > r.next() { + return nil, false + } + return slices.Clone(r.entries[from-r.first : to-r.first]), true +} + +// bytesBetween sums the sizes of entries [from, to). The caller guarantees +// first <= from <= to <= next(). +func (r *ring) bytesBetween(from, to int64) int { + var n int + for _, e := range r.entries[from-r.first : to-r.first] { + n += len(e) + } + return n +} diff --git a/internal/etcp/ring_test.go b/internal/etcp/ring_test.go new file mode 100644 index 0000000..a3f5d28 --- /dev/null +++ b/internal/etcp/ring_test.go @@ -0,0 +1,75 @@ +package etcp + +import ( + "bytes" + "testing" +) + +func filledRing(limit, n, size int) *ring { + r := &ring{limit: limit} + for i := range n { + r.push(bytes.Repeat([]byte{byte(i)}, size)) + } + return r +} + +func TestRingNextCountsPushes(t *testing.T) { + r := filledRing(1<<20, 5, 10) + if got := r.next(); got != 5 { + t.Fatalf("next() = %d, want 5", got) + } +} + +func TestRingTrimKeepsUnsent(t *testing.T) { + r := filledRing(25, 10, 10) // 100 bytes held, limit 25 + r.trim(4, 60) // entries 0..3 were written; 4..9 (60 bytes) were not + if r.first != 2 { + t.Fatalf("first = %d, want 2 (trim written bytes to the limit, ignoring the backlog)", r.first) + } + r.trim(4, 60) // the written 20 bytes are already within the limit + if r.first != 2 { + t.Fatalf("first = %d after a second trim, want 2", r.first) + } + r.trim(10, 0) + if r.first != 8 || r.bytes != 20 { + t.Fatalf("first, bytes = %d, %d; want 8, 20", r.first, r.bytes) + } +} + +func TestRingSince(t *testing.T) { + r := filledRing(25, 10, 10) + r.trim(10, 0) // retains 8 and 9 + tests := []struct { + name string + from int64 + want int + ok bool + }{ + {name: "whole window", from: 8, want: 2, ok: true}, + {name: "peer up to date", from: 10, want: 0, ok: true}, + {name: "peer behind the window", from: 7, ok: false}, + {name: "peer ahead of us", from: 11, ok: false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, ok := r.since(tt.from, r.next()) + if ok != tt.ok || len(got) != tt.want { + t.Fatalf("since(%d) = %d entries, %v; want %d, %v", tt.from, len(got), ok, tt.want, tt.ok) + } + if ok && tt.want > 0 && got[0][0] != byte(tt.from) { + t.Fatalf("first entry holds %d, want %d", got[0][0], tt.from) + } + }) + } +} + +func TestRingBytesBetweenAndRange(t *testing.T) { + r := filledRing(1<<20, 6, 10) + if got := r.bytesBetween(2, 5); got != 30 { + t.Fatalf("bytesBetween(2, 5) = %d, want 30", got) + } + got := r.appendRange(nil, 1, 3) + if len(got) != 2 || got[0][0] != 1 || got[1][0] != 2 { + t.Fatalf("appendRange(1, 3) = %v", got) + } +} diff --git a/internal/etcp/throttle_test.go b/internal/etcp/throttle_test.go new file mode 100644 index 0000000..24ea488 --- /dev/null +++ b/internal/etcp/throttle_test.go @@ -0,0 +1,156 @@ +package etcp_test + +import ( + "context" + "net" + "testing" + "testing/synctest" + "time" + + "github.com/tphakala/et-go/internal/etcp" + "github.com/tphakala/et-go/internal/etservertest" +) + +// throttledDialer slows every client-side write to 1 KiB per 125 ms of fake +// time (8 KiB/s), like a congested uplink, or with downlink set every +// client-side read instead, like a congested downlink. +type throttledDialer struct { + inner interface { + DialContext(ctx context.Context, network, address string) (net.Conn, error) + } + downlink bool +} + +func (d throttledDialer) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + c, err := d.inner.DialContext(ctx, network, address) + if err != nil { + return nil, err + } + if d.downlink { + return slowReadConn{c}, nil + } + return throttledConn{c}, nil +} + +// slowReadConn reads at most 1 KiB per 125 ms of fake time. +type slowReadConn struct{ net.Conn } + +func (c slowReadConn) Read(p []byte) (int, error) { + n, err := c.Conn.Read(p[:min(len(p), 1<<10)]) + if n > 0 { + time.Sleep(125 * time.Millisecond) + } + return n, err +} + +type throttledConn struct{ net.Conn } + +func (c throttledConn) Write(p []byte) (int, error) { + var n int + for len(p) > 0 { + m, err := c.Conn.Write(p[:min(len(p), 1<<10)]) + n += m + if err != nil { + return n, err + } + p = p[m:] + time.Sleep(125 * time.Millisecond) + } + return n, nil +} + +// A catchup larger than one idle timeout's worth of bytes on +// a throttled link must not trip the 30 s handshake idle timeout while bytes +// keep flowing. About 1 MiB at 8 KiB/s takes over two minutes of fake time. +func TestThrottledCatchupDoesNotTimeOut(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + srv := etservertest.NewServer(testID, testKey) + nw := etservertest.NewNetwork(srv) + defer nw.Close() + d := etcp.Dialer{NetDialer: throttledDialer{inner: nw}} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { _ = conn.Close() }() + + synctest.Wait() + nw.SetRefuse(true) + nw.CutAll() + synctest.Wait() + const n = 1024 // 1024 packets of 1 KiB: the catchup is about 1 MiB + for i := range n { + if err := conn.WritePacket(t.Context(), numbered(i, 1024)); err != nil { + t.Fatalf("WritePacket %d: %v", i, err) + } + } + before := nw.Dials() + start := time.Now() + nw.SetRefuse(false) + + ctx, cancel := context.WithTimeout(t.Context(), time.Hour) + defer cancel() + if err := expectNumbered(ctx, n, srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + if elapsed := time.Since(start); elapsed <= 30*time.Second { + t.Fatalf("recovery took %v; the catchup must outlast the 30s idle timeout for this test to mean anything", elapsed) + } + if got := nw.Dials() - before; got != 1 { + t.Fatalf("recovery used %d dials, want 1: the idle timeout fired while bytes were flowing", got) + } + }) +} + +// Upstream writes its whole catchup before reading ours +// (src/base/Connection.cpp:125-135 at et-v7.0.0), so while a slow downlink +// carries the server's catchup, our own catchup write cannot progress. That +// write must not time out while the server's bytes are still arriving: +// progress in either direction keeps the handshake alive. +func TestRecoverSurvivesSlowServerCatchup(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + srv := etservertest.NewServer(testID, testKey) + nw := etservertest.NewNetwork(srv) + defer nw.Close() + d := etcp.Dialer{NetDialer: throttledDialer{inner: nw, downlink: true}} + conn, err := d.Dial(t.Context(), testAddr, testID, testKey) + if err != nil { + t.Fatalf("Dial: %v", err) + } + defer func() { _ = conn.Close() }() + + synctest.Wait() + nw.SetRefuse(true) + nw.CutAll() + synctest.Wait() + const serverPackets = 512 // 512 KiB at 8 KiB/s: over a minute, past the 30 s idle timeout + for i := range serverPackets { + if err := srv.Send(t.Context(), numbered(i, 1024)); err != nil { + t.Fatalf("Send %d: %v", i, err) + } + } + for i := range 4 { // a non-empty client catchup, which blocks until the server reads it + if err := conn.WritePacket(t.Context(), numbered(i, 1024)); err != nil { + t.Fatalf("WritePacket %d: %v", i, err) + } + } + before := nw.Dials() + start := time.Now() + nw.SetRefuse(false) + + ctx, cancel := context.WithTimeout(t.Context(), time.Hour) + defer cancel() + if err := expectNumbered(ctx, 4, srv.Recv); err != nil { + t.Fatalf("server side: %v", err) + } + if err := expectNumbered(ctx, serverPackets, conn.ReadPacket); err != nil { + t.Fatalf("client side: %v", err) + } + if elapsed := time.Since(start); elapsed <= 30*time.Second { + t.Fatalf("recovery took %v; the server's catchup must outlast the 30s idle timeout for this test to mean anything", elapsed) + } + if got := nw.Dials() - before; got != 1 { + t.Fatalf("recovery used %d dials, want 1: our blocked catchup write timed out while the server's catchup was arriving", got) + } + }) +} diff --git a/internal/etservertest/network.go b/internal/etservertest/network.go new file mode 100644 index 0000000..9aa9e3b --- /dev/null +++ b/internal/etservertest/network.go @@ -0,0 +1,163 @@ +package etservertest + +import ( + "context" + "errors" + "io" + "maps" + "net" + "slices" + "sync" +) + +var errRefused = errors.New("connection refused") + +// Network is an in-memory network that connects etcp to a Server over +// net.Pipe, with fault injection. It implements etcp.Dialer.NetDialer. +// net.Pipe has no buffering, so every write blocks until the peer reads it. +type Network struct { + srv *Server + ctx context.Context + cancel context.CancelFunc + wg sync.WaitGroup + + mu sync.Mutex + pairs map[*pair]struct{} + refuse bool + dials int + cutAfter int64 // byte budget for the next connection; negative means none +} + +// NewNetwork returns a network whose connections are served by s. +func NewNetwork(s *Server) *Network { + ctx, cancel := context.WithCancel(context.Background()) + return &Network{ + srv: s, + ctx: ctx, + cancel: cancel, + pairs: make(map[*pair]struct{}), + cutAfter: -1, + } +} + +// DialContext connects to the server. It fails with a *net.OpError while +// SetRefuse(true) is in effect and after Close, and with a *net.OpError +// wrapping ctx.Err() when ctx has already ended, as net.Dialer does. +func (n *Network) DialContext(ctx context.Context, network, address string) (net.Conn, error) { + n.mu.Lock() + defer n.mu.Unlock() + n.dials++ + if err := ctx.Err(); err != nil { + return nil, &net.OpError{Op: "dial", Net: network, Err: err} + } + if n.refuse || n.ctx.Err() != nil { + return nil, &net.OpError{Op: "dial", Net: network, Err: errRefused} + } + budget := n.cutAfter + n.cutAfter = -1 + client, server := net.Pipe() + p := &pair{remaining: budget} + p.client = &pipeEnd{Conn: client, p: p} + p.server = &pipeEnd{Conn: server, p: p} + n.pairs[p] = struct{}{} + // Started under n.mu, so Close, whose CutAll takes n.mu before its + // wg.Wait, cannot wait on the group before this goroutine is added. + n.wg.Go(func() { + _ = n.srv.Serve(n.ctx, p.server) + p.close() + n.mu.Lock() + delete(n.pairs, p) + n.mu.Unlock() + }) + return p.client, nil +} + +// CutAll closes every live connection, both ends. +func (n *Network) CutAll() { + n.mu.Lock() + pairs := slices.Collect(maps.Keys(n.pairs)) + n.mu.Unlock() + for _, p := range pairs { + p.close() + } +} + +// CutAfter makes the next connection die after bytes have crossed it, counted +// over both directions. A cut can land in the middle of a frame or message. +func (n *Network) CutAfter(bytes int64) { + n.mu.Lock() + n.cutAfter = bytes + n.mu.Unlock() +} + +// SetRefuse makes dials fail while on is true. +func (n *Network) SetRefuse(on bool) { + n.mu.Lock() + n.refuse = on + n.mu.Unlock() +} + +// Dials returns the number of dial attempts so far, refused ones included. +func (n *Network) Dials() int { + n.mu.Lock() + defer n.mu.Unlock() + return n.dials +} + +// Close cuts everything, refuses further dials and waits for the server +// goroutines it started. +func (n *Network) Close() { + n.cancel() + n.CutAll() + n.wg.Wait() +} + +// pair is the two ends of one connection and its shared byte budget. +type pair struct { + client, server *pipeEnd + once sync.Once + mu sync.Mutex + remaining int64 // negative means unlimited +} + +func (p *pair) close() { + p.once.Do(func() { + _ = p.client.Conn.Close() + _ = p.server.Conn.Close() + }) +} + +// take reserves up to n bytes of the budget and reports whether the +// connection must be cut after writing them. +func (p *pair) take(n int) (allowed int, cut bool) { + p.mu.Lock() + defer p.mu.Unlock() + if p.remaining < 0 { + return n, false + } + allowed = int(min(int64(n), p.remaining)) + p.remaining -= int64(allowed) + return allowed, p.remaining == 0 +} + +type pipeEnd struct { + net.Conn + p *pair +} + +func (e *pipeEnd) Write(b []byte) (int, error) { + allowed, cut := e.p.take(len(b)) + n, err := e.Conn.Write(b[:allowed]) + if cut { + e.p.close() + if err == nil && n < len(b) { + err = io.ErrClosedPipe + } + } + return n, err +} + +func (e *pipeEnd) Close() error { + e.p.close() + return nil +} diff --git a/internal/etservertest/server.go b/internal/etservertest/server.go new file mode 100644 index 0000000..234c6a6 --- /dev/null +++ b/internal/etservertest/server.go @@ -0,0 +1,340 @@ +// Package etservertest provides an in-process fake etserver and an in-memory +// network, so etcp can be tested without sockets and under testing/synctest. +// +// The server is written independently of etcp, from upstream EternalTerminal +// semantics (tag et-v7.0.0), so the two implementations check each other. +// In particular it writes its whole catchup before reading the client's, +// exactly like upstream's Connection::recover (src/base/Connection.cpp:105-143). +package etservertest + +import ( + "bufio" + "context" + "errors" + "fmt" + "io" + "net" + "slices" + "sync" + + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/seal" + "github.com/tphakala/et-go/internal/wire" +) + +var errSessionEnded = errors.New("etservertest: session ended") + +// Server is an in-process fake etserver for one session. +type Server struct { + id string + key [32]byte + + mu sync.Mutex + registered bool // false after EndSession: connects get INVALID_KEY + known bool // a client connected before: the next connect is RETURNING_CLIENT + echo bool + out *seal.Stream // server to client + sent [][]byte // every sealed packet sent; index is the sequence number + in *seal.Stream // client to server + recvSeq int64 + queue []protocol.Packet // received, not yet returned by Recv + cur *serverLink + + notify chan struct{} // cap 1: queue grew + wake chan struct{} // cap 1: sent grew or the session ended +} + +type serverLink struct { + conn net.Conn + done chan struct{} // closed when Serve for this link returns +} + +// NewServer returns a server that accepts the session id with passkey. It +// panics if passkey is not 32 bytes, which is a bug in the test. +func NewServer(id, passkey string) *Server { + if len(passkey) != 32 { + panic("etservertest: passkey must be 32 bytes") + } + s := &Server{ + id: id, + registered: true, + echo: true, + notify: make(chan struct{}, 1), + wake: make(chan struct{}, 1), + } + copy(s.key[:], passkey) + s.out = seal.New(&s.key, seal.ServerToClient) + s.in = seal.New(&s.key, seal.ClientToServer) + return s +} + +// Serve runs one link on c: the connect handshake (and the recover exchange +// for a returning client), then the encrypted stream, until c fails, the +// session ends or ctx ends. A link replaces the one before it when it takes +// over, as upstream closes the old socket before recovering on the new one +// (src/base/ServerClientConnection.cpp:27-35 at et-v7.0.0). Takeover follows +// the order links reach it, not the order they connected, so a test that +// connects twice waits for the first link to settle. +func (s *Server) Serve(ctx context.Context, c net.Conn) error { + defer func() { _ = c.Close() }() + stop := context.AfterFunc(ctx, func() { _ = c.Close() }) + defer stop() + + var req protocol.ConnectRequest + if err := wire.ReadMessage(c, &req); err != nil { + return fmt.Errorf("etservertest: connect request: %w", err) + } + status, text := s.admit(&req) + resp := &protocol.ConnectResponse{} + resp.SetStatus(status) + if text != "" { + resp.SetError(text) + } + if err := wire.WriteMessage(c, resp); err != nil { + return fmt.Errorf("etservertest: connect response: %w", err) + } + if status != protocol.ConnectStatus_NEW_CLIENT && status != protocol.ConnectStatus_RETURNING_CLIENT { + return fmt.Errorf("etservertest: rejected: %s", text) + } + + done := s.takeOver(c) + defer close(done) + + var flushed int64 + if status == protocol.ConnectStatus_RETURNING_CLIENT { + f, err := s.recover(c) + if err != nil { + return err + } + flushed = f + } + + lctx, cancel := context.WithCancelCause(ctx) + defer cancel(nil) + stopClose := context.AfterFunc(lctx, func() { _ = c.Close() }) + defer stopClose() + + var wg sync.WaitGroup + wg.Go(func() { cancel(s.readLoop(c)) }) + wg.Go(func() { cancel(s.writeLoop(lctx, c, flushed)) }) + wg.Wait() + return context.Cause(lctx) +} + +// admit decides the ConnectResponse status for req. +func (s *Server) admit(req *protocol.ConnectRequest) (status protocol.ConnectStatus, text string) { + s.mu.Lock() + defer s.mu.Unlock() + switch { + case req.GetVersion() != protocol.Version: + // Upstream's text (src/base/ServerConnection.cpp:55-59 at et-v7.0.0). + return protocol.ConnectStatus_MISMATCHED_PROTOCOL, fmt.Sprintf( + "Mismatched protocol versions. Your client & server must be on the same version of ET. Client: %d != Server: %d", + req.GetVersion(), protocol.Version) + case req.GetClientId() != s.id || !s.registered: + // MEASURED against etserver 7.0.0: the error text for an unknown or + // ended session is "Client is not registered". + return protocol.ConnectStatus_INVALID_KEY, "Client is not registered" + case !s.known: + s.known = true + return protocol.ConnectStatus_NEW_CLIENT, "" + default: + return protocol.ConnectStatus_RETURNING_CLIENT, "" + } +} + +// takeOver makes c the active link, closing the previous one and waiting for +// its Serve to finish so the two never touch the stream state at once. +func (s *Server) takeOver(c net.Conn) chan struct{} { + done := make(chan struct{}) + s.mu.Lock() + prev := s.cur + s.cur = &serverLink{conn: c, done: done} + s.mu.Unlock() + if prev != nil { + _ = prev.conn.Close() + <-prev.done + } + return done +} + +// recover runs the server side of the recover exchange in upstream order and +// returns the sequence number the new link's writer starts from. +func (s *Server) recover(c io.ReadWriter) (int64, error) { + s.mu.Lock() + mine := s.recvSeq + s.mu.Unlock() + + sh := &protocol.SequenceHeader{} + sh.SetSequenceNumber(int32(mine)) + if err := wire.WriteMessage(c, sh); err != nil { + return 0, fmt.Errorf("etservertest: write sequence: %w", err) + } + var peer protocol.SequenceHeader + if err := wire.ReadMessage(c, &peer); err != nil { + return 0, fmt.Errorf("etservertest: read sequence: %w", err) + } + + s.mu.Lock() + from := int64(peer.GetSequenceNumber()) + if from < 0 || from > int64(len(s.sent)) { + s.mu.Unlock() + return 0, fmt.Errorf("etservertest: client sequence %d outside [0, %d]", from, len(s.sent)) + } + catchup := slices.Clone(s.sent[from:]) + flushed := int64(len(s.sent)) + s.mu.Unlock() + + // Upstream writes its whole catchup before reading ours + // (src/base/Connection.cpp:125-135). Over net.Pipe, which has no buffer, + // this deadlocks unless the client reads concurrently. + cb := &protocol.CatchupBuffer{} + cb.SetBuffer(catchup) + if err := wire.WriteMessage(c, cb); err != nil { + return 0, fmt.Errorf("etservertest: write catchup: %w", err) + } + var theirs protocol.CatchupBuffer + if err := wire.ReadMessage(c, &theirs); err != nil { + return 0, fmt.Errorf("etservertest: read catchup: %w", err) + } + for _, b := range theirs.GetBuffer() { + if err := s.accept(b); err != nil { + return 0, err + } + } + return flushed, nil +} + +func (s *Server) readLoop(r io.Reader) error { + br := bufio.NewReader(r) + var buf []byte + for { + frame, err := wire.ReadFrame(br, buf) + if err != nil { + return fmt.Errorf("etservertest: read: %w", err) + } + buf = frame + if err := s.accept(frame); err != nil { + return err + } + } +} + +// accept opens one sealed packet from the client and queues it for Recv. +func (s *Server) accept(b []byte) error { + encrypted, h, payload, err := wire.ParsePacket(b) + if err != nil { + return fmt.Errorf("etservertest: %w", err) + } + if !encrypted { + return errors.New("etservertest: unencrypted packet") + } + s.mu.Lock() + defer s.mu.Unlock() + plain, err := s.in.Open(nil, payload) + if err != nil { + return fmt.Errorf("etservertest: %w", err) + } + s.recvSeq++ + s.queue = append(s.queue, protocol.Packet{Header: h, Payload: plain}) + signal(s.notify) + // Upstream echoes KEEP_ALIVE (src/terminal/TerminalServer.cpp:389-393), but + // only once the session runs: a first packet other than INITIAL_PAYLOAD + // aborts etserver (TerminalServer.cpp:429-439). This fake echoes at any + // time and does not model that abort; etcp's TestNoProbeBeforeFirstPacket + // pins the client side instead. + if h == protocol.HeaderKeepAlive && s.echo { + s.enqueueLocked(protocol.Packet{Header: protocol.HeaderKeepAlive}) + signal(s.wake) + } + return nil +} + +func (s *Server) writeLoop(ctx context.Context, w io.Writer, flushed int64) error { + for { + s.mu.Lock() + pending := s.sent[flushed:] + ended := !s.registered + s.mu.Unlock() + if len(pending) == 0 { + if ended { + return errSessionEnded + } + select { + case <-s.wake: + continue + case <-ctx.Done(): + return context.Cause(ctx) + } + } + for _, f := range pending { + if err := wire.WriteFrame(w, f); err != nil { + return fmt.Errorf("etservertest: write: %w", err) + } + } + flushed += int64(len(pending)) + } +} + +func (s *Server) enqueueLocked(p protocol.Packet) { + sealed := s.out.Seal(nil, p.Payload) + s.sent = append(s.sent, wire.AppendPacket(nil, true, p.Header, sealed)) +} + +// Recv returns the next packet the client sent, exactly once and in order, +// including KEEP_ALIVE probes. +func (s *Server) Recv(ctx context.Context) (protocol.Packet, error) { + for { + s.mu.Lock() + if len(s.queue) > 0 { + p := s.queue[0] + s.queue[0] = protocol.Packet{} + s.queue = s.queue[1:] + s.mu.Unlock() + return p, nil + } + s.mu.Unlock() + select { + case <-s.notify: + case <-ctx.Done(): + return protocol.Packet{}, context.Cause(ctx) + } + } +} + +// Send queues a packet to the client. It is sealed at once and replayed +// across reconnects, like upstream's BackedWriter. +func (s *Server) Send(ctx context.Context, p protocol.Packet) error { + if err := ctx.Err(); err != nil { + return context.Cause(ctx) + } + s.mu.Lock() + s.enqueueLocked(p) + s.mu.Unlock() + signal(s.wake) + return nil +} + +// EndSession drops the session like a shell exit on etserver 7.0.0: the live +// link flushes what was already sent and closes, and later connects get +// INVALID_KEY (src/terminal/TerminalServer.cpp:329-332,422-426). +func (s *Server) EndSession() { + s.mu.Lock() + s.registered = false + s.mu.Unlock() + signal(s.wake) +} + +// EchoKeepAlive controls whether KEEP_ALIVE packets are echoed (default true). +func (s *Server) EchoKeepAlive(on bool) { + s.mu.Lock() + s.echo = on + s.mu.Unlock() +} + +func signal(c chan struct{}) { + select { + case c <- struct{}{}: + default: + } +} diff --git a/internal/etservertest/server_test.go b/internal/etservertest/server_test.go new file mode 100644 index 0000000..8f0df3b --- /dev/null +++ b/internal/etservertest/server_test.go @@ -0,0 +1,276 @@ +package etservertest_test + +import ( + "bufio" + "context" + "errors" + "net" + "testing" + "testing/synctest" + "time" + + "github.com/tphakala/et-go/internal/etservertest" + "github.com/tphakala/et-go/internal/protocol" + "github.com/tphakala/et-go/internal/seal" + "github.com/tphakala/et-go/internal/wire" +) + +const ( + testID = "XXXtestclient001" + testKey = "0123456789abcdef0123456789abcdef" +) + +// rawClient speaks just enough of the protocol to exercise the server +// without etcp, so the two implementations stay independent. +type rawClient struct { + conn net.Conn + br *bufio.Reader + out *seal.Stream + in *seal.Stream +} + +func connect(t *testing.T, n *etservertest.Network, id string, version int32) (*rawClient, *protocol.ConnectResponse) { + t.Helper() + conn, err := n.DialContext(t.Context(), "tcp", "et:2022") + if err != nil { + t.Fatalf("DialContext: %v", err) + } + req := &protocol.ConnectRequest{} + req.SetClientId(id) + req.SetVersion(version) + if err := wire.WriteMessage(conn, req); err != nil { + t.Fatalf("WriteMessage: %v", err) + } + resp := &protocol.ConnectResponse{} + if err := wire.ReadMessage(conn, resp); err != nil { + t.Fatalf("ReadMessage: %v", err) + } + var key [32]byte + copy(key[:], testKey) + return &rawClient{ + conn: conn, + br: bufio.NewReader(conn), + out: seal.New(&key, seal.ClientToServer), + in: seal.New(&key, seal.ServerToClient), + }, resp +} + +func (c *rawClient) send(t *testing.T, h protocol.Header, payload string) { + t.Helper() + frame := wire.AppendPacket(nil, true, h, c.out.Seal(nil, []byte(payload))) + if err := wire.WriteFrame(c.conn, frame); err != nil { + t.Fatalf("WriteFrame: %v", err) + } +} + +func (c *rawClient) recv(t *testing.T) protocol.Packet { + t.Helper() + frame, err := wire.ReadFrame(c.br, nil) + if err != nil { + t.Fatalf("ReadFrame: %v", err) + } + _, h, payload, err := wire.ParsePacket(frame) + if err != nil { + t.Fatalf("ParsePacket: %v", err) + } + plain, err := c.in.Open(nil, payload) + if err != nil { + t.Fatalf("Open: %v", err) + } + return protocol.Packet{Header: h, Payload: plain} +} + +func TestServerNewClientExchange(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + srv := etservertest.NewServer(testID, testKey) + n := etservertest.NewNetwork(srv) + defer n.Close() + + c, resp := connect(t, n, testID, protocol.Version) + if got := resp.GetStatus(); got != protocol.ConnectStatus_NEW_CLIENT { + t.Fatalf("status = %v, want NEW_CLIENT", got) + } + c.send(t, protocol.HeaderTerminalBuffer, "ls\n") + p, err := srv.Recv(t.Context()) + if err != nil || p.Header != protocol.HeaderTerminalBuffer || string(p.Payload) != "ls\n" { + t.Fatalf("Recv = %v %q, %v", p.Header, p.Payload, err) + } + if err := srv.Send(t.Context(), protocol.Packet{Header: protocol.HeaderTerminalBuffer, Payload: []byte("out")}); err != nil { + t.Fatalf("Send: %v", err) + } + if got := c.recv(t); string(got.Payload) != "out" { + t.Fatalf("client got %q, want %q", got.Payload, "out") + } + }) +} + +func TestServerEchoesKeepAlive(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + srv := etservertest.NewServer(testID, testKey) + n := etservertest.NewNetwork(srv) + defer n.Close() + + c, _ := connect(t, n, testID, protocol.Version) + c.send(t, protocol.HeaderKeepAlive, "") + if got := c.recv(t); got.Header != protocol.HeaderKeepAlive { + t.Fatalf("echo header = %v, want KEEP_ALIVE", got.Header) + } + }) +} + +func TestServerRejects(t *testing.T) { + tests := []struct { + name string + id string + version int32 + end bool + want protocol.ConnectStatus + }{ + {name: "unknown id", id: "XXXsomeoneelse01", version: protocol.Version, want: protocol.ConnectStatus_INVALID_KEY}, + {name: "old protocol", id: testID, version: 5, want: protocol.ConnectStatus_MISMATCHED_PROTOCOL}, + {name: "ended session", id: testID, version: protocol.Version, end: true, want: protocol.ConnectStatus_INVALID_KEY}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + srv := etservertest.NewServer(testID, testKey) + n := etservertest.NewNetwork(srv) + defer n.Close() + if tt.end { + srv.EndSession() + } + _, resp := connect(t, n, tt.id, tt.version) + if got := resp.GetStatus(); got != tt.want { + t.Fatalf("status = %v, want %v", got, tt.want) + } + }) + }) + } +} + +// Upstream writes each recover message before reading the peer's +// (src/base/Connection.cpp:105-143 at et-v7.0.0); etcp's deadlock test +// (TestCatchupBothWays) only means something if the fake does the same. A +// client that reads each server message before writing its own must get +// both; a fake that read first would leave these reads to time out. +func TestServerRecoverWritesCatchupFirst(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + srv := etservertest.NewServer(testID, testKey) + n := etservertest.NewNetwork(srv) + defer n.Close() + + connect(t, n, testID, protocol.Version) // registers the client + // Let the first link's Serve take over before the second connects: + // otherwise its late takeOver can close the second link. + synctest.Wait() + if err := srv.Send(t.Context(), protocol.Packet{Header: protocol.HeaderTerminalBuffer, Payload: []byte("owed")}); err != nil { + t.Fatalf("Send: %v", err) + } + n.CutAll() + + c, resp := connect(t, n, testID, protocol.Version) + if got := resp.GetStatus(); got != protocol.ConnectStatus_RETURNING_CLIENT { + t.Fatalf("status = %v, want RETURNING_CLIENT", got) + } + if err := c.conn.SetReadDeadline(time.Now().Add(time.Minute)); err != nil { + t.Fatalf("SetReadDeadline: %v", err) + } + var seq protocol.SequenceHeader + if err := wire.ReadMessage(c.br, &seq); err != nil { + t.Fatalf("server SequenceHeader before ours: %v", err) + } + if err := wire.WriteMessage(c.conn, &protocol.SequenceHeader{}); err != nil { + t.Fatalf("write SequenceHeader: %v", err) + } + var cb protocol.CatchupBuffer + if err := wire.ReadMessage(c.br, &cb); err != nil { + t.Fatalf("server CatchupBuffer before ours: %v", err) + } + if got := len(cb.GetBuffer()); got != 1 { + t.Fatalf("server catchup holds %d packets, want the 1 it owes", got) + } + if err := wire.WriteMessage(c.conn, &protocol.CatchupBuffer{}); err != nil { + t.Fatalf("write CatchupBuffer: %v", err) + } + }) +} + +func TestServerEndSessionFlushesThenCloses(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + srv := etservertest.NewServer(testID, testKey) + n := etservertest.NewNetwork(srv) + defer n.Close() + + c, _ := connect(t, n, testID, protocol.Version) + if err := srv.Send(t.Context(), protocol.Packet{Header: protocol.HeaderTerminalBuffer, Payload: []byte("bye")}); err != nil { + t.Fatalf("Send: %v", err) + } + srv.EndSession() + if got := c.recv(t); string(got.Payload) != "bye" { + t.Fatalf("last output = %q, want %q", got.Payload, "bye") + } + if _, err := wire.ReadFrame(c.br, nil); err == nil { + t.Fatal("link still open after EndSession") + } + }) +} + +// A dial whose context has already ended fails the way net.Dialer does: a +// *net.OpError wrapping the context's error. +func TestNetworkDialEndedContext(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + n := etservertest.NewNetwork(etservertest.NewServer(testID, testKey)) + defer n.Close() + + ctx, cancel := context.WithCancel(t.Context()) + cancel() + _, err := n.DialContext(ctx, "tcp", "et:2022") + if _, ok := errors.AsType[*net.OpError](err); !ok || !errors.Is(err, context.Canceled) { + t.Fatalf("dial with an ended context = %v, want a *net.OpError wrapping context.Canceled", err) + } + }) +} + +func TestNetworkRefuseAndCount(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + n := etservertest.NewNetwork(etservertest.NewServer(testID, testKey)) + defer n.Close() + + n.SetRefuse(true) + _, err := n.DialContext(t.Context(), "tcp", "et:2022") + if _, ok := errors.AsType[*net.OpError](err); !ok { + t.Fatalf("refused dial error = %v, want *net.OpError", err) + } + n.SetRefuse(false) + conn, err := n.DialContext(t.Context(), "tcp", "et:2022") + if err != nil { + t.Fatalf("DialContext: %v", err) + } + _ = conn.Close() + if got := n.Dials(); got != 2 { + t.Fatalf("Dials() = %d, want 2", got) + } + }) +} + +func TestNetworkCutAfter(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + n := etservertest.NewNetwork(etservertest.NewServer(testID, testKey)) + defer n.Close() + + n.CutAfter(5) + conn, err := n.DialContext(t.Context(), "tcp", "et:2022") + if err != nil { + t.Fatalf("DialContext: %v", err) + } + // An 8-byte write (the size of a handshake length prefix) exceeds + // the 5-byte budget, so the cut lands inside it. + wrote, err := conn.Write(make([]byte, 8)) + if wrote != 5 || err == nil { + t.Fatalf("Write = %d, %v; want 5 bytes and an error", wrote, err) + } + if _, err := conn.Read(make([]byte, 1)); err == nil { + t.Fatal("Read after cut succeeded") + } + }) +} diff --git a/internal/wire/frame.go b/internal/wire/frame.go index 668083f..e26fc76 100644 --- a/internal/wire/frame.go +++ b/internal/wire/frame.go @@ -5,48 +5,80 @@ import ( "errors" "fmt" "io" + "slices" ) +// AppendFrame appends frame to dst with its 4-byte big-endian length prefix +// and returns the extended slice. A frame above MaxFrameSize returns dst +// unchanged and an error wrapping ErrTooLarge. Appending into a reused dst +// lets a writer batch several frames into one Write without allocating per +// frame. frame must not lie in dst's spare capacity, which the prefix +// overwrites. +func AppendFrame(dst, frame []byte) ([]byte, error) { + if len(frame) > MaxFrameSize { + return dst, fmt.Errorf("wire: frame of %d bytes: %w", len(frame), ErrTooLarge) + } + dst = binary.BigEndian.AppendUint32(dst, uint32(len(frame))) + return append(dst, frame...), nil +} + // WriteFrame writes frame with a 4-byte big-endian length prefix, in a single -// Write call. +// Write call. A writer that reports fewer bytes than it was given yields +// io.ErrShortWrite. func WriteFrame(w io.Writer, frame []byte) error { if len(frame) > MaxFrameSize { - return fmt.Errorf("wire: write frame of %d bytes: %w", len(frame), ErrTooLarge) + // Checked before allocating the output buffer, so an oversized + // frame costs nothing. + return fmt.Errorf("wire: frame of %d bytes: %w", len(frame), ErrTooLarge) + } + buf, err := AppendFrame(make([]byte, 0, 4+len(frame)), frame) + if err != nil { + return err } - buf := make([]byte, 4, 4+len(frame)) - binary.BigEndian.PutUint32(buf, uint32(len(frame))) - buf = append(buf, frame...) - if _, err := w.Write(buf); err != nil { + if err := writeAll(w, buf); err != nil { return fmt.Errorf("wire: write frame: %w", err) } return nil } // ReadFrame reads one frame and returns its body, reusing buf's capacity when -// it is large enough. The result aliases buf; callers that keep it past the -// next ReadFrame must copy it. A length above MaxFrameSize returns ErrTooLarge +// it is large enough; the 4-byte length is read into buf's storage too, so a +// reused buf makes ReadFrame allocate nothing. When buf is too small the body +// buffer grows as bytes arrive rather than being allocated at the declared +// length up front. The result aliases buf, and every call, even a failed one, +// may overwrite buf's contents; callers that keep a result past the next +// ReadFrame must copy it. A length above MaxFrameSize returns ErrTooLarge // without reading the body. io.EOF is returned unwrapped only when the stream // ends cleanly before a frame starts; a stream that ends anywhere after the // first length byte, including right after a complete length, yields a // wrapped io.ErrUnexpectedEOF. func ReadFrame(r io.Reader, buf []byte) ([]byte, error) { - var hdr [4]byte - if _, err := io.ReadFull(r, hdr[:]); err != nil { - if errors.Is(err, io.EOF) { + // A local array for the length would escape to the heap through the + // io.Reader call, one allocation per frame. + buf = slices.Grow(buf[:0], 4)[:4] + if m, err := io.ReadFull(r, buf); err != nil { + if err = headerErr(m, err); errors.Is(err, io.EOF) { return nil, io.EOF } return nil, fmt.Errorf("wire: read frame length: %w", err) } - n := binary.BigEndian.Uint32(hdr[:]) + n := binary.BigEndian.Uint32(buf) if n > MaxFrameSize { return nil, fmt.Errorf("wire: read frame length %d: %w", n, ErrTooLarge) } - if cap(buf) < int(n) { - buf = make([]byte, n) - } - buf = buf[:n] - if _, err := io.ReadFull(r, buf); err != nil { - return nil, fmt.Errorf("wire: read frame body: %w", bodyErr(err)) + buf, err := readBody(r, buf[:0], int(n)) + if err != nil { + return nil, fmt.Errorf("wire: read frame body: %w", err) } return buf, nil } + +// writeAll issues one Write of buf and turns a short count without an error +// into io.ErrShortWrite, as io.Copy does. +func writeAll(w io.Writer, buf []byte) error { + n, err := w.Write(buf) + if err == nil && n != len(buf) { + err = io.ErrShortWrite + } + return err +} diff --git a/internal/wire/frame_test.go b/internal/wire/frame_test.go index 97d5cbf..22c1e34 100644 --- a/internal/wire/frame_test.go +++ b/internal/wire/frame_test.go @@ -106,9 +106,10 @@ func TestWriteFrameWriterError(t *testing.T) { // TestWriteFrameSingleWrite pins the single-Write contract on the success // path: WriteFrame must build the length prefix and body in one buffer -// before writing, not write the header and body separately. etcp relies on -// this to keep a frame from interleaving with another goroutine's write on -// the same connection. +// before writing, not write the header and body separately, so a frame is +// never split by another goroutine's write on a connection whose Write is +// safe for concurrent use. etcp's link writer frames with AppendFrame and +// issues its own single Write per batch. func TestWriteFrameSingleWrite(t *testing.T) { w := &countingWriter{} if err := WriteFrame(w, []byte("hello")); err != nil { @@ -121,7 +122,13 @@ func TestWriteFrameSingleWrite(t *testing.T) { func TestWriteFrameLimit(t *testing.T) { var buf bytes.Buffer - err := WriteFrame(&buf, make([]byte, MaxFrameSize+1)) + frame := make([]byte, MaxFrameSize+1) + var err error + // The size is checked before the output buffer is allocated, so an + // oversized frame costs no 16 MiB allocation. + if got := allocatedBytes(func() { err = WriteFrame(&buf, frame) }); got > 1<<20 { + t.Fatalf("WriteFrame allocated %d bytes to reject an oversized frame", got) + } if !errors.Is(err, ErrTooLarge) { t.Fatalf("WriteFrame = %v, want ErrTooLarge", err) } @@ -159,14 +166,32 @@ func TestReadFrameTruncated(t *testing.T) { } } +// FuzzReadFrame checks ReadFrame against the framing itself: on success the +// body is exactly the bytes after the 4-byte length, and reading the same +// input into a reused buffer, a small one or one that already holds bytes, +// gives the same body or the same error. func FuzzReadFrame(f *testing.F) { - f.Add([]byte{0, 0, 0, 3, 1, 2, 3}) - f.Add([]byte{0xff, 0xff, 0xff, 0xff}) - f.Add([]byte{}) - f.Fuzz(func(t *testing.T, in []byte) { + f.Add([]byte{0, 0, 0, 3, 1, 2, 3}, 0) + f.Add([]byte{0xff, 0xff, 0xff, 0xff}, 2) + f.Add([]byte{}, 64) + // 96 KiB declared: past the first 64 KiB grow chunk, with a body pattern + // that shows any bytes written at the wrong offset. + f.Add(append([]byte{0, 1, 0x80, 0}, bytes.Repeat([]byte{1, 2, 3, 4, 5, 6, 7}, 14100)...), 100) + f.Fuzz(func(t *testing.T, in []byte, capacity int) { got, err := ReadFrame(bytes.NewReader(in), nil) - if err == nil && len(got) > MaxFrameSize { - t.Fatalf("ReadFrame returned %d bytes, above MaxFrameSize", len(got)) + if err == nil { + if len(in) < 4 { + t.Fatalf("ReadFrame succeeded on %d input bytes", len(in)) + } + n := int(binary.BigEndian.Uint32(in)) + if n > MaxFrameSize || !bytes.Equal(got, in[4:4+n]) { + t.Fatalf("ReadFrame body of %d bytes does not match the %d declared", len(got), n) + } + } + scratch := bytes.Repeat([]byte{0xee}, max(capacity, 0)%(1<<17)) + again, err2 := ReadFrame(bytes.NewReader(in), scratch) + if (err == nil) != (err2 == nil) || !bytes.Equal(got, again) { + t.Fatalf("reused buffer (cap %d): got %d bytes, %v; fresh: %d bytes, %v", cap(scratch), len(again), err2, len(got), err) } }) } @@ -187,3 +212,109 @@ func BenchmarkFrameRoundTrip(b *testing.B) { } } } + +func TestAppendFrame(t *testing.T) { + dst := []byte{0xee} + got, err := AppendFrame(dst, []byte{0xaa, 0xbb}) + if err != nil { + t.Fatalf("AppendFrame: %v", err) + } + want := []byte{0xee, 0, 0, 0, 2, 0xaa, 0xbb} + if !bytes.Equal(got, want) { + t.Fatalf("AppendFrame = % x, want % x", got, want) + } + + got, err = AppendFrame(dst, make([]byte, MaxFrameSize+1)) + if !errors.Is(err, ErrTooLarge) { + t.Fatalf("AppendFrame(oversized) error = %v, want ErrTooLarge", err) + } + if !bytes.Equal(got, dst) { + t.Fatalf("AppendFrame(oversized) = % x, want dst unchanged", got) + } +} + +// A writer that reports a short count without an error must not silently +// truncate the frame. +func TestWriteFrameShortWrite(t *testing.T) { + if err := WriteFrame(shortWriter{}, []byte("hello")); !errors.Is(err, io.ErrShortWrite) { + t.Fatalf("WriteFrame = %v, want io.ErrShortWrite", err) + } +} + +// A reader that returns the first bytes of a length together with a wrapped +// io.EOF has cut the stream mid-header: that is not a clean end. With no +// byte read at all, the same wrapped io.EOF is a clean end, reported as a +// plain io.EOF. +func TestReadFrameWrappedEOF(t *testing.T) { + tests := []struct { + name string + in []byte + clean bool + }{ + {name: "nothing read", in: nil, clean: true}, + {name: "partial length", in: []byte{0, 0}}, + {name: "partial body", in: []byte{0, 0, 0, 5, 'a'}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := ReadFrame(&wrappedEOFReader{data: tt.in}, nil) + if tt.clean { + if err != io.EOF { //nolint:errorlint // a clean end must be the plain sentinel + t.Fatalf("ReadFrame = %v, want plain io.EOF", err) + } + return + } + if !errors.Is(err, io.ErrUnexpectedEOF) { + t.Fatalf("ReadFrame = %v, want io.ErrUnexpectedEOF", err) + } + }) + } +} + +// A peer that declares a large frame and sends only a few bytes must not +// make ReadFrame allocate the declared length up front. +func TestReadFrameGrowsWithData(t *testing.T) { + in := []byte("\x01\x00\x00\x00only ten b") // 16 MiB declared, 10 bytes sent + var err error + got := allocatedBytes(func() { + _, err = ReadFrame(bytes.NewReader(in), nil) + }) + if !errors.Is(err, io.ErrUnexpectedEOF) { + t.Fatalf("ReadFrame = %v, want io.ErrUnexpectedEOF", err) + } + if got > 1<<20 { + t.Fatalf("ReadFrame allocated %d bytes for a 10-byte body, want under 1 MiB", got) + } +} + +// ReadFrame into a buffer that is large enough, and AppendFrame into one, +// allocate nothing: etcp's read and write loops run them for every packet. +func TestFrameSteadyStateAllocs(t *testing.T) { + var enc bytes.Buffer + if err := WriteFrame(&enc, make([]byte, 1024)); err != nil { + t.Fatalf("WriteFrame: %v", err) + } + data := enc.Bytes() + r := bytes.NewReader(data) + buf := make([]byte, 0, 2048) + if n := testing.AllocsPerRun(100, func() { + r.Reset(data) + var err error + if buf, err = ReadFrame(r, buf); err != nil { + t.Fatalf("ReadFrame: %v", err) + } + }); n != 0 { + t.Errorf("ReadFrame with a reused buffer: %v allocs, want 0", n) + } + + frame := make([]byte, 1024) + out := make([]byte, 0, 2048) + if n := testing.AllocsPerRun(100, func() { + var err error + if out, err = AppendFrame(out[:0], frame); err != nil { + t.Fatalf("AppendFrame: %v", err) + } + }); n != 0 { + t.Errorf("AppendFrame into a reused buffer: %v allocs, want 0", n) + } +} diff --git a/internal/wire/message.go b/internal/wire/message.go index 455eabf..b0f9cb9 100644 --- a/internal/wire/message.go +++ b/internal/wire/message.go @@ -9,7 +9,8 @@ import ( ) // WriteMessage writes m as an 8-byte little-endian length followed by its -// protobuf encoding, in a single Write call. +// protobuf encoding, in a single Write call. A writer that reports fewer +// bytes than it was given yields io.ErrShortWrite. func WriteMessage(w io.Writer, m proto.Message) error { body, err := proto.Marshal(m) if err != nil { @@ -21,7 +22,7 @@ func WriteMessage(w io.Writer, m proto.Message) error { buf := make([]byte, 8, 8+len(body)) binary.LittleEndian.PutUint64(buf, uint64(len(body))) buf = append(buf, body...) - if _, err := w.Write(buf); err != nil { + if err := writeAll(w, buf); err != nil { return fmt.Errorf("wire: write %T: %w", m, err) } return nil @@ -30,24 +31,26 @@ func WriteMessage(w io.Writer, m proto.Message) error { // ReadMessage reads one length-prefixed message into m. A zero length yields // the message's zero value, as upstream does. A length above MaxMessageSize // (or negative as a signed int64) returns ErrTooLarge without reading the body. -// A stream that ends before the first length byte yields a wrapped io.EOF; -// one that ends anywhere later, including right after a complete length, -// yields a wrapped io.ErrUnexpectedEOF. +// The body buffer grows as bytes arrive rather than being allocated at the +// declared length up front. A stream that ends before the first length byte +// yields a wrapped io.EOF; one that ends anywhere later, including right +// after a complete length, yields a wrapped io.ErrUnexpectedEOF. A complete +// body that does not decode yields an error wrapping ErrMalformed. func ReadMessage(r io.Reader, m proto.Message) error { var hdr [8]byte - if _, err := io.ReadFull(r, hdr[:]); err != nil { - return fmt.Errorf("wire: read %T length: %w", m, err) + if k, err := io.ReadFull(r, hdr[:]); err != nil { + return fmt.Errorf("wire: read %T length: %w", m, headerErr(k, err)) } n := int64(binary.LittleEndian.Uint64(hdr[:])) if n < 0 || n > MaxMessageSize { return fmt.Errorf("wire: read %T length %d: %w", m, n, ErrTooLarge) } - body := make([]byte, n) - if _, err := io.ReadFull(r, body); err != nil { - return fmt.Errorf("wire: read %T body: %w", m, bodyErr(err)) + body, err := readBody(r, nil, int(n)) + if err != nil { + return fmt.Errorf("wire: read %T body: %w", m, err) } if err := proto.Unmarshal(body, m); err != nil { - return fmt.Errorf("wire: unmarshal %T: %w", m, err) + return fmt.Errorf("wire: unmarshal %T: %w: %w", m, ErrMalformed, err) } return nil } diff --git a/internal/wire/message_test.go b/internal/wire/message_test.go index f5856d9..d3f27b0 100644 --- a/internal/wire/message_test.go +++ b/internal/wire/message_test.go @@ -152,9 +152,11 @@ func TestWriteMessageWriterError(t *testing.T) { // TestWriteMessageSingleWrite pins the single-Write contract on the success // path: WriteMessage must build the length prefix and marshaled body in one -// buffer before writing, not write the header and body separately. etcp -// relies on this to keep a message from interleaving with another -// goroutine's write on the same connection. +// buffer before writing, not write the header and body separately, so that +// on a connection whose Write is safe for concurrent use a message is never +// split by another goroutine's write. etcp does not depend on it: each of +// its connections has one writer at a time, and its handshake conn splits +// large writes into chunks anyway. func TestWriteMessageSingleWrite(t *testing.T) { sh := &protocol.SequenceHeader{} sh.SetSequenceNumber(1) @@ -172,8 +174,8 @@ func TestReadMessageGarbage(t *testing.T) { // Length 2, then bytes that are not a valid protobuf (field 0 is illegal). in := []byte{2, 0, 0, 0, 0, 0, 0, 0, 0x00, 0x00} var sh protocol.SequenceHeader - if err := ReadMessage(bytes.NewReader(in), &sh); err == nil { - t.Fatal("ReadMessage(garbage) = nil, want an unmarshal error") + if err := ReadMessage(bytes.NewReader(in), &sh); !errors.Is(err, ErrMalformed) { + t.Fatalf("ReadMessage(garbage) = %v, want ErrMalformed", err) } } @@ -201,3 +203,38 @@ func FuzzReadMessage(f *testing.F) { } }) } + +func TestWriteMessageShortWrite(t *testing.T) { + sh := &protocol.SequenceHeader{} + sh.SetSequenceNumber(1) + if err := WriteMessage(shortWriter{}, sh); !errors.Is(err, io.ErrShortWrite) { + t.Fatalf("WriteMessage = %v, want io.ErrShortWrite", err) + } +} + +// See TestReadFrameWrappedEOF: bytes followed by a wrapped io.EOF are a cut. +func TestReadMessageWrappedEOF(t *testing.T) { + var sh protocol.SequenceHeader + err := ReadMessage(&wrappedEOFReader{data: []byte{1, 0, 0}}, &sh) + if !errors.Is(err, io.ErrUnexpectedEOF) { + t.Fatalf("ReadMessage = %v, want io.ErrUnexpectedEOF", err) + } +} + +// A peer that declares a message of MaxMessageSize and sends only a few +// bytes must not make ReadMessage allocate the declared length up front. +func TestReadMessageGrowsWithData(t *testing.T) { + in := binary.LittleEndian.AppendUint64(nil, MaxMessageSize) + in = append(in, 0x08, 0x01) + var sh protocol.SequenceHeader + var err error + got := allocatedBytes(func() { + err = ReadMessage(bytes.NewReader(in), &sh) + }) + if !errors.Is(err, io.ErrUnexpectedEOF) { + t.Fatalf("ReadMessage = %v, want io.ErrUnexpectedEOF", err) + } + if got > 1<<20 { + t.Fatalf("ReadMessage allocated %d bytes for a 2-byte body, want under 1 MiB", got) + } +} diff --git a/internal/wire/wire.go b/internal/wire/wire.go index f8e05bd..2c6a9cf 100644 --- a/internal/wire/wire.go +++ b/internal/wire/wire.go @@ -16,6 +16,7 @@ package wire import ( "errors" "io" + "slices" ) // Size limits for lengths read from the network. @@ -33,16 +34,66 @@ const ( // negative handshake length. var ErrTooLarge = errors.New("wire: length exceeds limit") +// ErrMalformed reports a handshake message whose complete body does not +// decode: the peer sent garbage, since a cut stream yields +// io.ErrUnexpectedEOF before decoding is attempted. +var ErrMalformed = errors.New("wire: malformed message") + // ErrShortPacket reports a serialized packet without its 2-byte header. var ErrShortPacket = errors.New("wire: packet shorter than its 2-byte header") -// bodyErr maps io.EOF from a body read to io.ErrUnexpectedEOF. io.ReadFull -// returns io.EOF when no byte at all was read, which for a body means the -// stream was cut after a complete length prefix: a broken link, not a clean -// end of stream. +// growChunk is the first allocation for a body that does not fit the +// caller's buffer; the buffer then doubles as bytes arrive. +const growChunk = 64 << 10 + +// bodyErr maps an io.EOF (possibly wrapped) from a body read to +// io.ErrUnexpectedEOF. A body read happens only after a complete length +// prefix, so an end of stream there, whether before the first body byte or +// between two growth chunks, is a broken link, not a clean end of stream. func bodyErr(err error) error { if errors.Is(err, io.EOF) { return io.ErrUnexpectedEOF } return err } + +// headerErr classifies a failed length-prefix read that got n bytes. Only a +// read that got no byte at all is a clean end, reported as a plain io.EOF. +// io.ReadFull maps only an unwrapped io.EOF after a partial read to +// io.ErrUnexpectedEOF, so a reader that returns bytes together with a +// wrapped io.EOF is mapped here. +func headerErr(n int, err error) error { + if !errors.Is(err, io.EOF) { + return err + } + if n == 0 { + return io.EOF + } + return io.ErrUnexpectedEOF +} + +// readBody reads exactly n bytes into buf's storage and returns buf[:n], +// reusing buf's capacity when it is large enough. Otherwise the buffer grows +// as bytes arrive, starting at growChunk and doubling, so a peer that +// declares a large length and sends little costs memory in proportion to +// what it sent, not to what it declared. +func readBody(r io.Reader, buf []byte, n int) ([]byte, error) { + if cap(buf) >= n { + buf = buf[:n] + if _, err := io.ReadFull(r, buf); err != nil { + return nil, bodyErr(err) + } + return buf, nil + } + buf = buf[:0] + for len(buf) < n { + next := min(n, max(2*len(buf), growChunk)) + buf = slices.Grow(buf, next-len(buf)) + m, err := io.ReadFull(r, buf[len(buf):next]) + buf = buf[:len(buf)+m] + if err != nil { + return nil, bodyErr(err) + } + } + return buf, nil +} diff --git a/internal/wire/wire_test.go b/internal/wire/wire_test.go index 3ace957..1f3ad01 100644 --- a/internal/wire/wire_test.go +++ b/internal/wire/wire_test.go @@ -3,6 +3,9 @@ package wire import ( "bytes" "errors" + "fmt" + "io" + "runtime" ) // errWriter is an io.Writer that always fails, counting how many times @@ -24,8 +27,8 @@ var errWriterSentinel = errors.New("wire_test: writer error") // countingWriter is an io.Writer that always succeeds, counting how many // times Write was called. WriteFrame and WriteMessage each build their // whole output (length prefix and body) in one buffer and issue a single -// Write; etcp relies on that, since a partial write on a real connection -// could otherwise interleave with another goroutine's frame. +// Write, so a frame or message is never split by another goroutine's write +// on a connection whose Write is safe for concurrent use. type countingWriter struct { buf bytes.Buffer calls int @@ -46,3 +49,33 @@ func (w *countingWriter) Write(p []byte) (int, error) { w.calls++ return w.buf.Write(p) } + +// shortWriter reports writing one byte less than it was given and no error, +// which breaks the io.Writer contract. +type shortWriter struct{} + +func (shortWriter) Write(p []byte) (int, error) { + return max(len(p)-1, 0), nil +} + +// wrappedEOFReader returns all of data in one Read together with an io.EOF +// wrapped in another error, as some io.Reader implementations do. +type wrappedEOFReader struct { + data []byte +} + +func (r *wrappedEOFReader) Read(p []byte) (int, error) { + n := copy(p, r.data) + r.data = r.data[n:] + return n, fmt.Errorf("wrappedEOFReader: %w", io.EOF) +} + +// allocatedBytes returns the heap bytes f allocates, measured as the +// difference in runtime.MemStats.TotalAlloc. +func allocatedBytes(f func()) uint64 { + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + f() + runtime.ReadMemStats(&after) + return after.TotalAlloc - before.TotalAlloc +}