Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,7 @@ require (
k8s.io/apimachinery v0.36.1
k8s.io/client-go v0.36.1
k8s.io/metrics v0.36.1
k8s.io/streaming v0.36.1
k8s.io/utils v0.0.0-20260319190234-28399d86e0b5
sigs.k8s.io/controller-runtime v0.24.1
sigs.k8s.io/yaml v1.6.0
Expand Down Expand Up @@ -186,7 +187,6 @@ require (
gotest.tools/v3 v3.5.2 // indirect
k8s.io/klog/v2 v2.140.0 // indirect
k8s.io/kube-openapi v0.0.0-20260317180543-43fb72c5454a // indirect
k8s.io/streaming v0.36.1 // indirect
sigs.k8s.io/json v0.0.0-20250730193827-2d320260d730 // indirect
sigs.k8s.io/randfill v1.0.0 // indirect
sigs.k8s.io/structured-merge-diff/v6 v6.3.2 // indirect
Expand Down
58 changes: 8 additions & 50 deletions internal/ateclient/builder.go
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ import (
"crypto/tls"
"crypto/x509"
"fmt"
"io"
"net"
"net/http"
"os"
"sync"
Expand All @@ -40,7 +40,6 @@ import (
"k8s.io/client-go/kubernetes"
"k8s.io/client-go/rest"
"k8s.io/client-go/tools/clientcmd"
"k8s.io/client-go/tools/portforward"
"k8s.io/client-go/transport/spdy"
metricsv1beta1 "k8s.io/metrics/pkg/client/clientset/versioned"
)
Expand Down Expand Up @@ -190,59 +189,22 @@ func dialPortForward(ctx context.Context, kubeconfigPath, k8sContext string, tra
return nil, fmt.Errorf("failed to create SPDY transport: %w", err)
}

dialer := spdy.NewDialer(upgrader, &http.Client{Transport: transport}, http.MethodPost, req.URL())

stopCh := make(chan struct{})
readyCh := make(chan struct{})

ports := []string{"0:443"} // Port 0 asks OS for a random available local port

fw, err := portforward.New(dialer, ports, stopCh, readyCh, io.Discard, io.Discard)
if err != nil {
return nil, fmt.Errorf("failed to create port forwarder: %w", err)
}

errCh := make(chan error, 1)
var wg sync.WaitGroup
wg.Add(1)
go func() {
defer wg.Done()
if err := fw.ForwardPorts(); err != nil {
errCh <- fmt.Errorf("port forwarding failed: %w", err)
}
}()

// Wait for the tunnel to be ready, an error, or context cancellation
select {
case <-readyCh:
// Tunnel is ready!
case err := <-errCh:
return nil, err
case <-ctx.Done():
return nil, ctx.Err()
}

forwardedPorts, err := fw.GetPorts()
if err != nil || len(forwardedPorts) == 0 {
close(stopCh)
return nil, fmt.Errorf("failed to get forwarded ports: %w", err)
}

localPort := forwardedPorts[0].Local
localEndpoint := fmt.Sprintf("127.0.0.1:%d", localPort)
dialer := newSPDYPortForwardDialer(upgrader, &http.Client{Transport: transport}, http.MethodPost, req.URL())

tlsCfg, err := serverTLSConfig(ctx, clientset)
if err != nil {
close(stopCh)
return nil, err
}

var opts []grpc.DialOption
opts = append(opts, grpc.WithTransportCredentials(credentials.NewTLS(tlsCfg)))
opts = append(opts, grpc.WithStatsHandler(otelgrpc.NewClientHandler()))
opts = append(opts, grpc.WithAuthority(apiServerName))
opts = append(opts, grpc.WithContextDialer(func(ctx context.Context, _ string) (net.Conn, error) {
return dialPodPort(ctx, dialer, 443)
}))
tokenOpt, err := bearerTokenDialOption(ctx, clientset)
if err != nil {
close(stopCh)
return nil, err
}
opts = append(opts, tokenOpt)
Expand All @@ -251,20 +213,16 @@ func dialPortForward(ctx context.Context, kubeconfigPath, k8sContext string, tra
opts = append(opts, grpc.WithUnaryInterceptor(newTraceInterceptor()))
}

conn, err := grpc.NewClient(localEndpoint, opts...)
conn, err := grpc.NewClient("passthrough:///"+apiServerName+":443", opts...)
if err != nil {
close(stopCh)
return nil, fmt.Errorf("failed to dial gRPC over tunnel: %w", err)
}

return &Client{
ControlClient: ateapipb.NewControlClient(conn),
DebugClient: ateapipb.NewDebugClient(conn),
conn: conn,
cancel: func() {
close(stopCh)
wg.Wait()
},
cancel: func() {},
}, nil
}

Expand Down
197 changes: 197 additions & 0 deletions internal/ateclient/portforward.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,197 @@
// Copyright 2021 The Kubernetes Authors.
// Copyright 2026 Google LLC
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.

package ateclient

import (
"context"
"errors"
"fmt"
"io"
"log/slog"
"net"
"net/http"
"net/url"
"strconv"
"sync"
"sync/atomic"
"time"

corev1 "k8s.io/api/core/v1"
"k8s.io/client-go/tools/portforward"
"k8s.io/client-go/transport/spdy"
"k8s.io/streaming/pkg/httpstream"
)

type portForwardDialer func(context.Context, ...string) (httpstream.Connection, string, error)

func newSPDYPortForwardDialer(upgrader spdy.Upgrader, client *http.Client, method string, target *url.URL) portForwardDialer {
return func(ctx context.Context, protocols ...string) (httpstream.Connection, string, error) {
req, err := http.NewRequestWithContext(ctx, method, target.String(), nil)
if err != nil {
return nil, "", fmt.Errorf("creating port-forward request: %w", err)
}
return spdy.NegotiateStreaming(upgrader, client, req, protocols...)
}
}

// dialPodPort exposes a Kubernetes port-forward data stream directly as a
// net.Conn. This avoids an intermediate localhost TCP listener and its noisy
// connection-reset errors when the gRPC client shuts down.
func dialPodPort(ctx context.Context, dialer portForwardDialer, remotePort int) (conn net.Conn, finalErr error) {
if err := ctx.Err(); err != nil {
return nil, err
}

streamConn, protocol, err := dialer(ctx, portforward.PortForwardProtocolV1Name)
if err != nil {
return nil, fmt.Errorf("dialing port-forward connection: %w", err)
}
defer func() {
if finalErr != nil {
_ = streamConn.Close()
}
}()

setupDone := make(chan struct{})
monitorDone := make(chan struct{})
var setupFinished atomic.Bool
go func() {
defer close(monitorDone)
select {
case <-ctx.Done():
if setupFinished.CompareAndSwap(false, true) {
_ = streamConn.Close()
}
case <-setupDone:
}
}()
defer func() {
close(setupDone)
<-monitorDone
}()

if err := ctx.Err(); err != nil {
return nil, err
}
if protocol != portforward.PortForwardProtocolV1Name {
return nil, fmt.Errorf("port-forward protocol mismatch: server selected %q", protocol)
}

headers := http.Header{}
headers.Set(corev1.StreamType, corev1.StreamTypeError)
headers.Set(corev1.PortHeader, strconv.Itoa(remotePort))
headers.Set(corev1.PortForwardRequestIDHeader, "0")

errorStream, err := streamConn.CreateStream(headers)
if err != nil {
return nil, fmt.Errorf("creating port-forward error stream: %w", err)
}
// The error stream is read-only from the client's perspective.
_ = errorStream.Close()

headers.Set(corev1.StreamType, corev1.StreamTypeData)
dataStream, err := streamConn.CreateStream(headers)
if err != nil {
return nil, fmt.Errorf("creating port-forward data stream: %w", err)
}

stream := &portForwardStream{
Stream: dataStream,
errorStream: errorStream,
streamConn: streamConn,
remotePort: remotePort,
}

if err := ctx.Err(); err != nil {
return nil, err
}
if !setupFinished.CompareAndSwap(false, true) {
return nil, ctx.Err()
}
go stream.watchErrors()
return stream, nil
}

type portForwardStream struct {
httpstream.Stream
errorStream httpstream.Stream
streamConn httpstream.Connection
remotePort int
closeOnce sync.Once
closed atomic.Bool
closeErr error
}

var _ net.Conn = (*portForwardStream)(nil)

func (s *portForwardStream) watchErrors() {
message, err := io.ReadAll(s.errorStream)
if s.closed.Load() {
return
}
switch {
case err != nil:
slog.Error("Error reading from port-forward error stream",
slog.Int("remotePort", s.remotePort), slog.Any("err", err))
case len(message) > 0:
slog.Error("Port-forward connection failed",
slog.Int("remotePort", s.remotePort), slog.String("error", string(message)))
default:
return
}
_ = s.Close()
}

func (s *portForwardStream) Close() error {
s.closeOnce.Do(func() {
s.closed.Store(true)
s.closeErr = errors.Join(s.Stream.Close(), s.streamConn.Close())
})
return s.closeErr
}

func (s *portForwardStream) LocalAddr() net.Addr {
return portForwardAddr("kubernetes-api")
}

func (s *portForwardStream) RemoteAddr() net.Addr {
return portForwardAddr(fmt.Sprintf("pod:%d", s.remotePort))
}

// SPDY streams do not expose deadline controls. gRPC applies RPC deadlines at
// the transport layer, so these methods intentionally match Kubernetes'
// direct port-forward net.Conn adapter and remain no-ops.
func (s *portForwardStream) SetDeadline(time.Time) error {
return nil
}

func (s *portForwardStream) SetReadDeadline(time.Time) error {
return nil
}

func (s *portForwardStream) SetWriteDeadline(time.Time) error {
return nil
}

type portForwardAddr string

func (a portForwardAddr) Network() string {
return "port-forward"
}

func (a portForwardAddr) String() string {
return string(a)
}
Loading