diff --git a/README.md b/README.md index 3a0744a..d9aa7b6 100644 --- a/README.md +++ b/README.md @@ -323,7 +323,7 @@ curl -sS -H "Authorization: Bearer $TOKEN" http://PUBLIC_IP:18780/settings/confi # equivalently ?token= on the query string ``` -If no relay is connected, the proxy returns `503 {"error":"search backend offline","code":"offline"}`. Tunnel protocol: TCP, `AUTH ` then length-prefixed JSON request/response frames (bodies base64). Latest tunnel connection wins. Hub tokens are not used in that handshake. +If no relay is connected, the proxy returns `503 {"error":"search backend offline","code":"offline"}`. Tunnel protocol: TCP, `AUTH ` then length-prefixed JSON request/response frames (bodies base64). Latest tunnel connection wins. Hub tokens are not used in that handshake. Both ends ping every 12s, treat a missing pong within 45s as dead, refresh a 60s read deadline on every frame (including ping/pong), enable 15s TCP keepalive, and the relay reconnects from 500ms up to 10s. Binaries: `go build -o search-proxy ./cmd/proxy` and `go build -o search-relay ./cmd/relay`. Copy them to the VPS; they are static-ish Go binaries (same module). diff --git a/cmd/proxy/main.go b/cmd/proxy/main.go index c46ec93..a58dc87 100644 --- a/cmd/proxy/main.go +++ b/cmd/proxy/main.go @@ -139,12 +139,10 @@ func (h *hub) handleTunnel(c net.Conn) { return } _ = c.SetDeadline(time.Time{}) - if tc, ok := c.(*net.TCPConn); ok { - _ = tc.SetKeepAlive(true) - _ = tc.SetKeepAlivePeriod(30 * time.Second) - } + tunnel.EnableTCPKeepAlive(c) s := newSession(c, br) + s.startKeepalive() h.mu.Lock() old := h.sess h.sess = s @@ -153,9 +151,10 @@ func (h *hub) handleTunnel(c net.Conn) { log.Printf("replacing previous tunnel from %s", old.remote) old.close(errReplaced) } - log.Printf("tunnel connected from %s", c.RemoteAddr()) + log.Printf("tunnel connected from %s (ping=%s pong_wait=%s read_idle=%s tcp_keepalive=%s)", + c.RemoteAddr(), tunnel.PingInterval, tunnel.PongTimeout, tunnel.ReadIdleTimeout, tunnel.TCPKeepAlivePeriod) s.readLoop() - log.Printf("tunnel disconnected from %s", c.RemoteAddr()) + log.Printf("tunnel disconnected from %s: %v", c.RemoteAddr(), s.closeReason()) h.mu.Lock() if h.sess == s { h.sess = nil @@ -397,18 +396,39 @@ type session struct { mu sync.Mutex pend map[string]chan tunnel.Frame closed bool + onPong func() + stopKA func() + reason error } func newSession(c net.Conn, br *bufio.Reader) *session { - return &session{conn: c, br: br, remote: c.RemoteAddr().String(), pend: map[string]chan tunnel.Frame{}} + return &session{ + conn: c, + br: br, + remote: c.RemoteAddr().String(), + pend: map[string]chan tunnel.Frame{}, + onPong: func() {}, + stopKA: func() {}, + } +} + +func (s *session) startKeepalive() { + s.startKeepaliveCfg(tunnel.DefaultKeepalive()) +} + +func (s *session) startKeepaliveCfg(cfg tunnel.KeepaliveConfig) { + s.onPong, s.stopKA = tunnel.StartKeepalive(context.Background(), cfg, s.write, func(err error) { + log.Printf("tunnel keepalive dead from %s: %v", s.remote, err) + s.close(err) + }) } func (s *session) write(f tunnel.Frame) error { s.wmu.Lock() defer s.wmu.Unlock() - _ = s.conn.SetWriteDeadline(time.Now().Add(30 * time.Second)) + _ = tunnel.SetWriteIdle(s.conn) err := tunnel.WriteFrame(s.conn, f) - _ = s.conn.SetWriteDeadline(time.Time{}) + tunnel.ClearWriteDeadline(s.conn) return err } @@ -437,19 +457,28 @@ func (s *session) roundTrip(ctx context.Context, f tunnel.Frame) (tunnel.Frame, } } +func (s *session) closeReason() error { + s.mu.Lock() + defer s.mu.Unlock() + if s.reason != nil { + return s.reason + } + return io.EOF +} + func (s *session) readLoop() { defer s.close(io.EOF) for { - _ = s.conn.SetReadDeadline(time.Now().Add(90 * time.Second)) var f tunnel.Frame - if err := tunnel.ReadFrame(s.br, &f); err != nil { + if err := tunnel.ReadFrameRefreshing(s.conn, s.br, &f); err != nil { + s.close(tunnel.ClassifyReadError(err)) return } switch f.Type { - case tunnel.TypePong, tunnel.TypePing: - if f.Type == tunnel.TypePing { - _ = s.write(tunnel.Frame{Type: tunnel.TypePong, ID: f.ID}) - } + case tunnel.TypePong: + s.onPong() + case tunnel.TypePing: + _ = s.write(tunnel.Frame{Type: tunnel.TypePong, ID: f.ID}) case tunnel.TypeResp, tunnel.TypeRespHead, tunnel.TypeRespChunk, tunnel.TypeRespEnd: s.mu.Lock() ch := s.pend[f.ID] @@ -475,6 +504,9 @@ func (s *session) close(err error) { return } s.closed = true + if err != nil { + s.reason = err + } for id, ch := range s.pend { select { case ch <- tunnel.Frame{Type: "resp", ID: id, Status: 503, Error: "offline"}: @@ -482,7 +514,11 @@ func (s *session) close(err error) { } delete(s.pend, id) } + stopKA := s.stopKA s.mu.Unlock() + if stopKA != nil { + stopKA() + } _ = s.conn.Close() } diff --git a/cmd/proxy/proxy_test.go b/cmd/proxy/proxy_test.go index 306f639..36b18bc 100644 --- a/cmd/proxy/proxy_test.go +++ b/cmd/proxy/proxy_test.go @@ -1,12 +1,18 @@ package main import ( + "bufio" "encoding/json" + "io" + "net" "net/http" "net/http/httptest" + "sync" "testing" + "time" "search-service/internal/proxyauth" + "search-service/internal/tunnel" ) func testProxy(t *testing.T) *hub { @@ -71,3 +77,82 @@ func TestProxyPapersSkipsBearerSoLoginCanRender(t *testing.T) { } } } + +func TestSessionEmitsAndAnswersPing(t *testing.T) { + client, server := net.Pipe() + defer client.Close() + defer server.Close() + + s := newSession(server, bufio.NewReader(server)) + s.startKeepaliveCfg(tunnel.KeepaliveConfig{ + Interval: 30 * time.Millisecond, + PongWait: 400 * time.Millisecond, + }) + done := make(chan struct{}) + go func() { + s.readLoop() + close(done) + }() + + var mu sync.Mutex + sawPong := false + sawPing := false + got := make(chan struct{}, 4) + go func() { + for { + _ = client.SetReadDeadline(time.Now().Add(2 * time.Second)) + var f tunnel.Frame + if err := tunnel.ReadFrame(client, &f); err != nil { + return + } + switch f.Type { + case tunnel.TypePong: + if f.ID == "peer" { + mu.Lock() + sawPong = true + mu.Unlock() + got <- struct{}{} + } + case tunnel.TypePing: + _ = tunnel.WriteFrame(client, tunnel.Frame{Type: tunnel.TypePong, ID: f.ID}) + mu.Lock() + first := !sawPing + sawPing = true + mu.Unlock() + if first { + got <- struct{}{} + } + } + } + }() + + select { + case <-got: + case <-time.After(2 * time.Second): + t.Fatal("proxy should emit its own ping") + } + if err := tunnel.WriteFrame(client, tunnel.Frame{Type: tunnel.TypePing, ID: "peer"}); err != nil { + t.Fatal(err) + } + select { + case <-got: + case <-time.After(2 * time.Second): + t.Fatal("proxy should answer peer ping") + } + mu.Lock() + okPing, okPong := sawPing, sawPong + mu.Unlock() + if !okPing { + t.Fatal("proxy should emit its own ping") + } + if !okPong { + t.Fatal("proxy should answer peer ping") + } + s.close(io.EOF) + _ = client.Close() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("readLoop did not exit") + } +} diff --git a/cmd/relay/main.go b/cmd/relay/main.go index 1528e3d..2de7628 100644 --- a/cmd/relay/main.go +++ b/cmd/relay/main.go @@ -8,7 +8,6 @@ import ( "fmt" "io" "log" - "net" "net/http" "net/url" "os" @@ -45,18 +44,20 @@ func main() { }, } - log.Printf("search-relay backend=%s tunnel=%s", backend, tun) - backoff := time.Second + log.Printf("search-relay backend=%s tunnel=%s keepalive ping=%s pong_wait=%s read_idle=%s tcp_keepalive=%s", + backend, tun, tunnel.PingInterval, tunnel.PongTimeout, tunnel.ReadIdleTimeout, tunnel.TCPKeepAlivePeriod) + backoff := tunnel.ReconnectMin for ctx.Err() == nil { connectedAt := time.Now() err := runOnce(ctx, tun, token, backend, client) if ctx.Err() != nil { break } - if time.Since(connectedAt) > 10*time.Second { - backoff = time.Second + alive := time.Since(connectedAt) + if alive > tunnel.HealthyResetAfter { + backoff = tunnel.ReconnectMin } - log.Printf("tunnel dropped: %v; reconnect in %s", err, backoff) + log.Printf("tunnel reconnect: reason=%v alive=%s backoff=%s", err, alive.Truncate(time.Millisecond), backoff) t := time.NewTimer(backoff) select { case <-ctx.Done(): @@ -64,12 +65,7 @@ func main() { return case <-t.C: } - if backoff < 30*time.Second { - backoff *= 2 - if backoff > 30*time.Second { - backoff = 30 * time.Second - } - } + backoff = tunnel.NextReconnectBackoff(backoff) } } @@ -106,12 +102,13 @@ func parseTunnelAddr(s string) (string, error) { } func runOnce(ctx context.Context, addr, token, backend string, client *http.Client) error { - d := net.Dialer{Timeout: 15 * time.Second, KeepAlive: 30 * time.Second} + d := tunnel.Dialer(15 * time.Second) c, err := d.DialContext(ctx, "tcp", addr) if err != nil { return err } defer c.Close() + tunnel.EnableTCPKeepAlive(c) _ = c.SetDeadline(time.Now().Add(15 * time.Second)) if _, err := c.Write([]byte("AUTH " + token + "\n")); err != nil { @@ -132,61 +129,62 @@ func runOnce(ctx context.Context, addr, token, backend string, client *http.Clie write := func(f tunnel.Frame) error { wmu.Lock() defer wmu.Unlock() - _ = c.SetWriteDeadline(time.Now().Add(30 * time.Second)) + _ = tunnel.SetWriteIdle(c) err := tunnel.WriteFrame(c, f) - _ = c.SetWriteDeadline(time.Time{}) + tunnel.ClearWriteDeadline(c) return err } - stopPing := make(chan struct{}) - go func() { - t := time.NewTicker(25 * time.Second) - defer t.Stop() - for { - select { - case <-ctx.Done(): - return - case <-stopPing: - return - case <-t.C: - if err := write(tunnel.Frame{Type: "ping"}); err != nil { - _ = c.Close() - return - } - } + var closeOnce sync.Once + closeConn := func() { closeOnce.Do(func() { _ = c.Close() }) } + defer closeConn() + + errCh := make(chan error, 1) + die := func(err error) { + select { + case errCh <- err: + default: } - }() - defer func() { close(stopPing); _ = c.Close() }() + closeConn() + } + + onPong, stopKA := tunnel.StartKeepalive(ctx, tunnel.DefaultKeepalive(), write, die) + defer stopKA() for { - select { - case <-ctx.Done(): + if ctx.Err() != nil { return ctx.Err() - default: } - _ = c.SetReadDeadline(time.Now().Add(90 * time.Second)) var f tunnel.Frame - if err := tunnel.ReadFrame(br, &f); err != nil { - return err + if err := tunnel.ReadFrameRefreshing(c, br, &f); err != nil { + select { + case kerr := <-errCh: + return tunnel.ClassifyReadError(kerr) + default: + } + if ctx.Err() != nil { + return ctx.Err() + } + return tunnel.ClassifyReadError(err) } switch f.Type { - case "ping": - _ = write(tunnel.Frame{Type: "pong", ID: f.ID}) - case "pong": - // keepalive + case tunnel.TypePing: + _ = write(tunnel.Frame{Type: tunnel.TypePong, ID: f.ID}) + case tunnel.TypePong: + onPong() case tunnel.TypeReq: go func(f tunnel.Frame) { if tunnel.PathNeedsStream(f.Path) { if err := streamBackend(ctx, client, backend, f, write); err != nil { log.Printf("stream download id=%s: %v", f.ID, err) - _ = c.Close() + die(fmt.Errorf("stream write: %w", err)) } return } resp := doBackend(ctx, client, backend, f) if err := write(resp); err != nil { log.Printf("write resp id=%s: %v", f.ID, err) - _ = c.Close() + die(fmt.Errorf("resp write: %w", err)) } }(f) } diff --git a/internal/tunnel/keepalive.go b/internal/tunnel/keepalive.go new file mode 100644 index 0000000..d7ef326 --- /dev/null +++ b/internal/tunnel/keepalive.go @@ -0,0 +1,248 @@ +package tunnel + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "strconv" + "sync" + "time" +) + +const ( + // PingInterval is how often each side emits an app-level ping. + PingInterval = 12 * time.Second + // PongTimeout is how long a ping may go unanswered before the link is dead. + PongTimeout = 45 * time.Second + // ReadIdleTimeout is the TCP read deadline, refreshed on every inbound frame. + ReadIdleTimeout = 60 * time.Second + // TCPKeepAlivePeriod is the OS-level TCP keepalive idle/probe interval. + TCPKeepAlivePeriod = 15 * time.Second + // ReconnectMin is the first reconnect delay after a drop. + ReconnectMin = 500 * time.Millisecond + // ReconnectMax caps exponential reconnect backoff. + ReconnectMax = 10 * time.Second + // HealthyResetAfter resets reconnect backoff once a session lasts this long. + HealthyResetAfter = 10 * time.Second + writeIdleTimeout = 30 * time.Second +) + +// ErrPongTimeout is returned when a keepalive ping is not answered in time. +var ErrPongTimeout = errors.New("tunnel keepalive: missing pong") + +// KeepaliveConfig tunes application-level ping/pong. +type KeepaliveConfig struct { + Interval time.Duration + PongWait time.Duration +} + +// DefaultKeepalive is the production ping interval and pong deadline. +func DefaultKeepalive() KeepaliveConfig { + return KeepaliveConfig{Interval: PingInterval, PongWait: PongTimeout} +} + +func (c KeepaliveConfig) normalized() KeepaliveConfig { + if c.Interval <= 0 { + c.Interval = PingInterval + } + if c.PongWait <= 0 { + c.PongWait = PongTimeout + } + return c +} + +// EnableTCPKeepAlive turns on OS probes with a short period (NAT-friendly). +func EnableTCPKeepAlive(c net.Conn) { + tc, ok := c.(*net.TCPConn) + if !ok { + return + } + _ = tc.SetKeepAlive(true) + _ = tc.SetKeepAlivePeriod(TCPKeepAlivePeriod) +} + +// Dialer returns a TCP dialer with a short OS keepalive period. +func Dialer(timeout time.Duration) net.Dialer { + if timeout <= 0 { + timeout = 15 * time.Second + } + return net.Dialer{ + Timeout: timeout, + KeepAlive: TCPKeepAlivePeriod, + KeepAliveConfig: net.KeepAliveConfig{ + Enable: true, + Idle: TCPKeepAlivePeriod, + Interval: TCPKeepAlivePeriod, + Count: 3, + }, + } +} + +// SetReadIdle sets an absolute read deadline idle in the future. +func SetReadIdle(c net.Conn, idle time.Duration) error { + if c == nil { + return nil + } + return c.SetReadDeadline(time.Now().Add(idle)) +} + +// RefreshReadDeadline extends the read deadline after any inbound frame +// (including ping/pong) so idle NAT mappings stay aligned with keepalive. +func RefreshReadDeadline(c net.Conn) error { + return SetReadIdle(c, ReadIdleTimeout) +} + +// ReadFrameRefreshing reads one frame and refreshes the read deadline before +// and after so ping/pong count as activity. +func ReadFrameRefreshing(c net.Conn, r io.Reader, f *Frame) error { + return ReadFrameRefreshingIdle(c, r, f, ReadIdleTimeout) +} + +// ReadFrameRefreshingIdle is ReadFrameRefreshing with a caller-chosen idle time. +func ReadFrameRefreshingIdle(c net.Conn, r io.Reader, f *Frame, idle time.Duration) error { + if err := SetReadIdle(c, idle); err != nil { + return err + } + if err := ReadFrame(r, f); err != nil { + return err + } + return SetReadIdle(c, idle) +} + +// ClassifyReadError annotates idle timeouts so reconnect logs are explicit. +func ClassifyReadError(err error) error { + if err == nil { + return nil + } + if errors.Is(err, ErrPongTimeout) { + return err + } + if IsTimeout(err) { + return fmt.Errorf("read idle timeout (%s): %w", ReadIdleTimeout, err) + } + return err +} + +// IsTimeout reports whether err is a network timeout (including deadlines). +func IsTimeout(err error) bool { + var ne net.Error + return errors.As(err, &ne) && ne.Timeout() +} + +// NextReconnectBackoff doubles cur, starting at ReconnectMin and capping at ReconnectMax. +func NextReconnectBackoff(cur time.Duration) time.Duration { + if cur < ReconnectMin { + return ReconnectMin + } + next := cur * 2 + if next > ReconnectMax { + return ReconnectMax + } + return next +} + +// Pinger tracks whether our keepalive pings have been answered. +type Pinger struct { + mu sync.Mutex + lastPong time.Time + pingSent bool +} + +// NewPinger starts a healthy keepalive window from now. +func NewPinger() *Pinger { + return &Pinger{lastPong: time.Now()} +} + +// OnPong records a pong (or equivalent proof that our ping was answered). +func (p *Pinger) OnPong() { + p.mu.Lock() + defer p.mu.Unlock() + p.lastPong = time.Now() + p.pingSent = false +} + +// MarkPing records that we have an outstanding ping. +func (p *Pinger) MarkPing() { + p.mu.Lock() + defer p.mu.Unlock() + p.pingSent = true +} + +// TimedOut reports whether a ping has gone unanswered longer than PongTimeout. +func (p *Pinger) TimedOut() bool { + return p.Overdue(PongTimeout) +} + +// Overdue is TimedOut with a caller-chosen wait (used by tests and StartKeepalive). +func (p *Pinger) Overdue(wait time.Duration) bool { + p.mu.Lock() + defer p.mu.Unlock() + if !p.pingSent { + return false + } + return time.Since(p.lastPong) > wait +} + +// StartKeepalive emits pings on Interval and calls die if no pong arrives +// within PongWait. onPong must be invoked when a pong frame is read. +func StartKeepalive(ctx context.Context, cfg KeepaliveConfig, write func(Frame) error, die func(error)) (onPong func(), stop func()) { + cfg = cfg.normalized() + p := NewPinger() + ctx, cancel := context.WithCancel(ctx) + var stopOnce sync.Once + stop = func() { stopOnce.Do(cancel) } + + sendPing := func() bool { + if p.Overdue(cfg.PongWait) { + die(fmt.Errorf("%w after %s", ErrPongTimeout, cfg.PongWait)) + return false + } + p.MarkPing() + if err := write(Frame{Type: TypePing, ID: pingID()}); err != nil { + die(fmt.Errorf("tunnel keepalive ping: %w", err)) + return false + } + return true + } + + go func() { + if !sendPing() { + return + } + ticker := time.NewTicker(cfg.Interval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + if !sendPing() { + return + } + } + } + }() + return p.OnPong, stop +} + +func pingID() string { + return strconv.FormatInt(time.Now().UnixNano(), 10) +} + +// SetWriteIdle is the short write deadline used for tunnel frames. +func SetWriteIdle(c net.Conn) error { + if c == nil { + return nil + } + return c.SetWriteDeadline(time.Now().Add(writeIdleTimeout)) +} + +// ClearWriteDeadline removes a write deadline after a successful frame write. +func ClearWriteDeadline(c net.Conn) { + if c == nil { + return + } + _ = c.SetWriteDeadline(time.Time{}) +} diff --git a/internal/tunnel/keepalive_test.go b/internal/tunnel/keepalive_test.go new file mode 100644 index 0000000..59b4b3f --- /dev/null +++ b/internal/tunnel/keepalive_test.go @@ -0,0 +1,229 @@ +package tunnel + +import ( + "context" + "errors" + "io" + "net" + "sync" + "testing" + "time" +) + +func TestNextReconnectBackoff(t *testing.T) { + if got := NextReconnectBackoff(0); got != ReconnectMin { + t.Fatalf("from 0: got %s want %s", got, ReconnectMin) + } + if got := NextReconnectBackoff(ReconnectMin); got != time.Second { + t.Fatalf("from min: got %s", got) + } + got := ReconnectMin + var steps []time.Duration + for i := 0; i < 8; i++ { + got = NextReconnectBackoff(got) + steps = append(steps, got) + if got > ReconnectMax { + t.Fatalf("backoff exceeded cap: %s", got) + } + } + if steps[len(steps)-1] != ReconnectMax { + t.Fatalf("did not reach cap, last=%s steps=%v", steps[len(steps)-1], steps) + } +} + +func TestPingerMissingPong(t *testing.T) { + p := NewPinger() + if p.TimedOut() { + t.Fatal("fresh pinger is not waiting on a ping") + } + p.MarkPing() + if p.TimedOut() { + t.Fatal("just-sent ping should not be overdue") + } + p.lastPong = time.Now().Add(-(PongTimeout + time.Second)) + if !p.TimedOut() { + t.Fatal("expected missing-pong timeout") + } + p.OnPong() + p.MarkPing() + if p.TimedOut() { + t.Fatal("pong should reset the watchdog") + } +} + +func TestReadDeadlineRefreshedOnPingPong(t *testing.T) { + a, b := net.Pipe() + defer a.Close() + defer b.Close() + + idle := 80 * time.Millisecond + if err := SetReadIdle(a, idle); err != nil { + t.Fatal(err) + } + + go func() { + time.Sleep(50 * time.Millisecond) + _ = WriteFrame(b, Frame{Type: TypePing, ID: "1"}) + time.Sleep(50 * time.Millisecond) // would exceed the original 80ms deadline + _ = WriteFrame(b, Frame{Type: TypePong, ID: "1"}) + time.Sleep(50 * time.Millisecond) + _ = WriteFrame(b, Frame{Type: TypePing, ID: "2"}) + }() + + var f Frame + if err := ReadFrameRefreshingIdle(a, a, &f, idle); err != nil { + t.Fatalf("first ping: %v", err) + } + if f.Type != TypePing || f.ID != "1" { + t.Fatalf("got %+v", f) + } + if err := ReadFrameRefreshingIdle(a, a, &f, idle); err != nil { + t.Fatalf("pong after refresh: %v", err) + } + if f.Type != TypePong { + t.Fatalf("got %+v", f) + } + if err := ReadFrameRefreshingIdle(a, a, &f, idle); err != nil { + t.Fatalf("second ping: %v", err) + } + if f.ID != "2" { + t.Fatalf("got %+v", f) + } +} + +func TestReadDeadlineExpiresWithoutFrames(t *testing.T) { + a, b := net.Pipe() + defer a.Close() + defer b.Close() + + errCh := make(chan error, 1) + go func() { + var f Frame + errCh <- ReadFrameRefreshingIdle(a, a, &f, 40*time.Millisecond) + }() + select { + case err := <-errCh: + if err == nil { + t.Fatal("expected idle timeout") + } + if !IsTimeout(err) { + t.Fatalf("want timeout, got %v", err) + } + if classified := ClassifyReadError(err); !IsTimeout(classified) { + t.Fatalf("classify: %v", classified) + } + case <-time.After(500 * time.Millisecond): + t.Fatal("ReadFrame did not return") + } +} + +func TestKeepaliveEmitsPingAndAcceptsPong(t *testing.T) { + a, b := net.Pipe() + defer a.Close() + defer b.Close() + + died := make(chan error, 1) + var wmu sync.Mutex + write := func(f Frame) error { + wmu.Lock() + defer wmu.Unlock() + return WriteFrame(a, f) + } + onPong, stop := StartKeepalive(context.Background(), KeepaliveConfig{ + Interval: 30 * time.Millisecond, + PongWait: 200 * time.Millisecond, + }, write, func(err error) { died <- err }) + defer stop() + + go func() { + for { + var f Frame + if err := ReadFrame(b, &f); err != nil { + return + } + if f.Type == TypePing { + // Local read-loop would record the matching pong here. + onPong() + } + } + }() + + select { + case err := <-died: + t.Fatalf("keepalive died: %v", err) + case <-time.After(300 * time.Millisecond): + } +} + +func TestKeepaliveDiesWithoutPong(t *testing.T) { + a, b := net.Pipe() + defer a.Close() + defer b.Close() + + died := make(chan error, 1) + write := func(f Frame) error { + return WriteFrame(a, f) + } + _, stop := StartKeepalive(context.Background(), KeepaliveConfig{ + Interval: 20 * time.Millisecond, + PongWait: 80 * time.Millisecond, + }, write, func(err error) { died <- err }) + defer stop() + + go func() { + for { + var f Frame + if err := ReadFrame(b, &f); err != nil { + return + } + } + }() + + select { + case err := <-died: + if !errors.Is(err, ErrPongTimeout) { + t.Fatalf("got %v", err) + } + case <-time.After(400 * time.Millisecond): + t.Fatal("expected pong timeout") + } +} + +func TestKeepaliveStopsWithoutDie(t *testing.T) { + a, b := net.Pipe() + defer a.Close() + defer b.Close() + + died := make(chan error, 1) + _, stop := StartKeepalive(context.Background(), KeepaliveConfig{ + Interval: 20 * time.Millisecond, + PongWait: time.Second, + }, func(f Frame) error { + return WriteFrame(a, f) + }, func(err error) { died <- err }) + + var f Frame + _ = b.SetReadDeadline(time.Now().Add(time.Second)) + if err := ReadFrame(b, &f); err != nil { + t.Fatalf("first ping: %v", err) + } + stop() + select { + case err := <-died: + t.Fatalf("stop should not die: %v", err) + case <-time.After(80 * time.Millisecond): + } +} + +func TestClassifyReadErrorPongTimeout(t *testing.T) { + err := ClassifyReadError(ErrPongTimeout) + if !errors.Is(err, ErrPongTimeout) { + t.Fatalf("%v", err) + } + if ClassifyReadError(nil) != nil { + t.Fatal("nil") + } + if ClassifyReadError(io.EOF) != io.EOF { + t.Fatal("eof passthrough") + } +}