package transport import ( "context" "errors" "net" "strconv" "strings" "testing" "time" toki "git.toki-labs.com/toki/proto-socket/go" "go.uber.org/zap" "google.golang.org/protobuf/proto" iop "iop/proto/gen/iop" ) type recordingConn struct { writeDeadlines []time.Time } func (c *recordingConn) Read(_ []byte) (int, error) { return 0, nil } func (c *recordingConn) Write(b []byte) (int, error) { return len(b), nil } func (c *recordingConn) Close() error { return nil } func (c *recordingConn) LocalAddr() net.Addr { return noopAddr("local") } func (c *recordingConn) RemoteAddr() net.Addr { return noopAddr("remote") } func (c *recordingConn) SetDeadline(_ time.Time) error { return nil } func (c *recordingConn) SetReadDeadline(_ time.Time) error { return nil } func (c *recordingConn) SetWriteDeadline(t time.Time) error { c.writeDeadlines = append(c.writeDeadlines, t) return nil } type noopAddr string func (a noopAddr) Network() string { return string(a) } func (a noopAddr) String() string { return string(a) } // closedLocalAddr returns a loopback address that has no listener, so a dial to // it is refused deterministically. func closedLocalAddr(t *testing.T) string { t.Helper() l, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } addr := l.Addr().String() _ = l.Close() return addr } func rejectingEdgeParserMap() toki.ParserMap { return toki.ParserMap{ toki.TypeNameOf(&iop.RegisterRequest{}): func(b []byte) (proto.Message, error) { m := &iop.RegisterRequest{} return m, proto.Unmarshal(b, m) }, } } // startRejectingEdge starts a mock edge that answers RegisterRequest with // Accepted=false and the given reason. func startRejectingEdge(t *testing.T, ctx context.Context, reason string) string { t.Helper() l, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } addr := l.Addr().String() host, portStr, _ := net.SplitHostPort(addr) port, _ := strconv.Atoi(portStr) _ = l.Close() server := toki.NewTcpServer(host, port, func(conn net.Conn) *toki.TcpClient { client := toki.NewTcpClient(conn, 30, 10, rejectingEdgeParserMap()) toki.AddRequestListenerTyped[*iop.RegisterRequest, *iop.RegisterResponse]( &client.Communicator, func(_ *iop.RegisterRequest) (*iop.RegisterResponse, error) { return &iop.RegisterResponse{Accepted: false, Reason: reason}, nil }, ) return client }) if err := server.Start(ctx); err != nil { t.Fatalf("start rejecting edge: %v", err) } t.Cleanup(func() { server.Stop() }) return addr } // TestDialEdgeClassifiesConnectFailures verifies that DialEdge classifies its // failures into fatal (invalid address, authoritative registration rejection) // and retryable (Edge unavailable) while preserving the underlying cause. REFACTOR-2. func TestDialEdgeClassifiesConnectFailures(t *testing.T) { logger := zap.NewNop() t.Run("invalid address is fatal", func(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() _, err := DialEdge(ctx, "missing-port", "tok", logger) if err == nil { t.Fatal("expected error for invalid address") } if !IsFatalConnectError(err) { t.Fatalf("invalid address must be fatal, got retryable: %v", err) } var ce *ConnectError if !errors.As(err, &ce) || ce.Retryable() { t.Fatalf("expected non-retryable ConnectError, got %v", err) } }) t.Run("non-numeric port is fatal", func(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() _, err := DialEdge(ctx, "127.0.0.1:not-a-port", "tok", logger) if err == nil { t.Fatal("expected error for invalid port") } if !IsFatalConnectError(err) { t.Fatalf("invalid port must be fatal, got retryable: %v", err) } var numErr *strconv.NumError if !errors.As(err, &numErr) || !errors.Is(err, strconv.ErrSyntax) { t.Fatalf("expected preserved strconv syntax cause, got %v", err) } }) for _, port := range []string{"0", "-1", "65536"} { port := port t.Run("out-of-range port "+port+" is fatal", func(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() _, err := DialEdge(ctx, net.JoinHostPort("127.0.0.1", port), "tok", logger) if err == nil { t.Fatalf("expected error for out-of-range port %s", port) } if !IsFatalConnectError(err) { t.Fatalf("out-of-range port %s must be fatal, got retryable: %v", port, err) } var ce *ConnectError if !errors.As(err, &ce) || ce.Retryable() { t.Fatalf("expected non-retryable ConnectError, got %v", err) } var numErr *strconv.NumError if !errors.As(err, &numErr) || !errors.Is(err, strconv.ErrRange) { t.Fatalf("expected preserved strconv range cause, got %v", err) } if numErr.Num != port { t.Fatalf("range cause port = %q, want %q", numErr.Num, port) } }) } t.Run("connection refused is retryable and preserves cause", func(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() _, err := DialEdge(ctx, closedLocalAddr(t), "tok", logger) if err == nil { t.Fatal("expected error for refused connection") } if IsFatalConnectError(err) { t.Fatalf("refused connection must be retryable, got fatal: %v", err) } var ce *ConnectError if !errors.As(err, &ce) || !ce.Retryable() { t.Fatalf("expected retryable ConnectError, got %v", err) } // The underlying network error must remain reachable through the wrap. var opErr *net.OpError if !errors.As(err, &opErr) { t.Fatalf("expected wrapped *net.OpError cause, got %v", err) } }) t.Run("registration rejected is fatal and preserves reason", func(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() addr := startRejectingEdge(t, ctx, "bad-token") _, err := DialEdge(ctx, addr, "tok", logger) if err == nil { t.Fatal("expected error for rejected registration") } if !IsFatalConnectError(err) { t.Fatalf("rejected registration must be fatal, got retryable: %v", err) } if !strings.Contains(err.Error(), "bad-token") { t.Fatalf("expected rejection reason preserved in error, got %v", err) } }) t.Run("unclassified error defaults to retryable", func(t *testing.T) { if IsFatalConnectError(errors.New("plain")) { t.Fatal("unclassified error must not be fatal") } }) } func TestWriteDeadlineConnSetsAndClearsDeadline(t *testing.T) { base := &recordingConn{} conn := &writeDeadlineConn{Conn: base, timeout: 25 * time.Millisecond} n, err := conn.Write([]byte("ping")) if err != nil { t.Fatalf("Write: %v", err) } if n != 4 { t.Fatalf("Write bytes: got %d want 4", n) } if len(base.writeDeadlines) != 2 { t.Fatalf("write deadline calls: got %d want 2", len(base.writeDeadlines)) } if base.writeDeadlines[0].IsZero() { t.Fatal("first write deadline must be non-zero") } if !base.writeDeadlines[1].IsZero() { t.Fatalf("second write deadline must clear the deadline, got %v", base.writeDeadlines[1]) } }