From a0377ed2b236239b6845da0162a5519637eba70e Mon Sep 17 00:00:00 2001 From: uweltman Date: Tue, 7 Jul 2026 23:08:41 -0700 Subject: [PATCH 1/5] Add opt-in TLS for inter-service gRPC connections (cherry picked from commit fe3d95764) --- packages/shared/pkg/grpc/server.go | 51 ++++++++++++++++++++++++++++-- 1 file changed, 49 insertions(+), 2 deletions(-) diff --git a/packages/shared/pkg/grpc/server.go b/packages/shared/pkg/grpc/server.go index 1e2e039b77..1e12c6bd4a 100644 --- a/packages/shared/pkg/grpc/server.go +++ b/packages/shared/pkg/grpc/server.go @@ -2,14 +2,18 @@ package grpc import ( "context" + "crypto/tls" "time" + "go.uber.org/zap" + "github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/logging" "github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/recovery" "github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/selector" "go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc" "go.opentelemetry.io/otel/attribute" "google.golang.org/grpc" + "google.golang.org/grpc/credentials" "google.golang.org/grpc/keepalive" "google.golang.org/grpc/metadata" @@ -22,6 +26,10 @@ type ServerOption func(*serverOptions) type serverOptions struct { withSandboxResumeMetrics bool + certFile string + keyFile string + certPEM []byte + keyPEM []byte } // WithSandboxResumeMetrics adds sandbox.resume attribute to otelgrpc metrics, @@ -30,6 +38,22 @@ func WithSandboxResumeMetrics() ServerOption { return func(o *serverOptions) { o.withSandboxResumeMetrics = true } } +// WithTLS configures server-side TLS using the given certificate and key files. +func WithTLS(certFile, keyFile string) ServerOption { + return func(o *serverOptions) { + o.certFile = certFile + o.keyFile = keyFile + } +} + +// WithTLSFromPEM configures server-side TLS using in-memory PEM-encoded certificate and key. +func WithTLSFromPEM(certPEM, keyPEM []byte) ServerOption { + return func(o *serverOptions) { + o.certPEM = certPEM + o.keyPEM = keyPEM + } +} + func NewGRPCServer(tel *telemetry.Client, opts ...ServerOption) *grpc.Server { var cfg serverOptions for _, o := range opts { @@ -57,7 +81,7 @@ func NewGRPCServer(tel *telemetry.Client, opts ...ServerOption) *grpc.Server { otelOpts = append(otelOpts, otelgrpc.WithMetricAttributesFn(extractSandboxResumeAttrs)) } - return grpc.NewServer( + serverOpts := []grpc.ServerOption{ grpc.KeepaliveEnforcementPolicy(keepalive.EnforcementPolicy{ MinTime: 5 * time.Second, PermitWithoutStream: true, @@ -82,7 +106,30 @@ func NewGRPCServer(tel *telemetry.Client, opts ...ServerOption) *grpc.Server { ignoredLoggingRoutes, ), ), - ) + } + + if cfg.certFile != "" && cfg.keyFile != "" { + creds, err := credentials.NewServerTLSFromFile(cfg.certFile, cfg.keyFile) + if err != nil { + logger.L().Fatal(context.Background(), "failed to load gRPC TLS credentials", + zap.String("certFile", cfg.certFile), + zap.String("keyFile", cfg.keyFile), + zap.Error(err), + ) + } + + serverOpts = append(serverOpts, grpc.Creds(creds)) + } else if len(cfg.certPEM) > 0 && len(cfg.keyPEM) > 0 { + cert, err := tls.X509KeyPair(cfg.certPEM, cfg.keyPEM) + if err != nil { + logger.L().Fatal(context.Background(), "failed to parse gRPC TLS PEM credentials", zap.Error(err)) + } + + creds := credentials.NewServerTLSFromCert(&cert) + serverOpts = append(serverOpts, grpc.Creds(creds)) + } + + return grpc.NewServer(serverOpts...) } // extractSandboxResumeAttrs reads sandbox.resume from gRPC metadata set by the From 3038b5cecf48ed5620689d5c1e49eefc6df4ebce Mon Sep 17 00:00:00 2001 From: uweltman Date: Thu, 9 Jul 2026 12:29:20 -0700 Subject: [PATCH 2/5] Add opt-in TLS for orchestrator proxy/Hyperloop (port 5007) (cherry picked from commit 141d45766) --- packages/shared/pkg/proxy/pool/client.go | 3 ++ packages/shared/pkg/proxy/pool/pool.go | 6 +++- packages/shared/pkg/proxy/proxy.go | 43 ++++++++++++++++++++++++ 3 files changed, 51 insertions(+), 1 deletion(-) diff --git a/packages/shared/pkg/proxy/pool/client.go b/packages/shared/pkg/proxy/pool/client.go index c5a7fb4806..5b748d053d 100644 --- a/packages/shared/pkg/proxy/pool/client.go +++ b/packages/shared/pkg/proxy/pool/client.go @@ -2,6 +2,7 @@ package pool import ( "context" + "crypto/tls" "errors" "log" "net" @@ -36,6 +37,7 @@ func newProxyClient( currentConnsCounter *atomic.Int64, l *log.Logger, disableKeepAlives bool, + tlsConfig *tls.Config, ) *ProxyClient { activeConnections := smap.New[*tracking.Connection]() @@ -49,6 +51,7 @@ func newProxyClient( ResponseHeaderTimeout: 0, DisableKeepAlives: disableKeepAlives, ForceAttemptHTTP2: false, + TLSClientConfig: tlsConfig, // TCP configuration DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) { var conn net.Conn diff --git a/packages/shared/pkg/proxy/pool/pool.go b/packages/shared/pkg/proxy/pool/pool.go index e2a6221fc7..2bbcd7308d 100644 --- a/packages/shared/pkg/proxy/pool/pool.go +++ b/packages/shared/pkg/proxy/pool/pool.go @@ -2,6 +2,7 @@ package pool import ( "context" + "crypto/tls" "sync/atomic" "time" @@ -30,15 +31,17 @@ type ProxyPool struct { totalConnsCounter atomic.Uint64 currentConnsCounter atomic.Int64 disableKeepAlives bool + tlsConfig *tls.Config } -func New(maxClientConns int, maxConnectionAttempts int, idleTimeout time.Duration, disableKeepAlives bool) *ProxyPool { +func New(maxClientConns int, maxConnectionAttempts int, idleTimeout time.Duration, disableKeepAlives bool, tlsConfig *tls.Config) *ProxyPool { return &ProxyPool{ pool: smap.New[*ProxyClient](), maxClientConns: maxClientConns, maxConnectionAttempts: maxConnectionAttempts, idleTimeout: idleTimeout, disableKeepAlives: disableKeepAlives, + tlsConfig: tlsConfig, } } @@ -84,6 +87,7 @@ func (p *ProxyPool) Get(ctx context.Context, d *Destination) *ProxyClient { &p.currentConnsCounter, stdLogger, p.disableKeepAlives, + p.tlsConfig, ) }) } diff --git a/packages/shared/pkg/proxy/proxy.go b/packages/shared/pkg/proxy/proxy.go index 5ef0a72a0f..cacd8e6273 100644 --- a/packages/shared/pkg/proxy/proxy.go +++ b/packages/shared/pkg/proxy/proxy.go @@ -2,6 +2,7 @@ package proxy import ( "context" + "crypto/tls" "fmt" "net" "net/http" @@ -36,6 +37,18 @@ type Proxy struct { currentServerConnsCounter atomic.Int64 } +type options struct { + upstreamTLS *tls.Config +} + +type Option func(*options) + +func WithUpstreamTLS(cfg *tls.Config) Option { + return func(o *options) { + o.upstreamTLS = cfg + } +} + type MaxConnectionAttempts int const ( @@ -50,12 +63,19 @@ func New( getDestination func(r *http.Request) (*pool.Destination, error), connLimitConfig *ConnectionLimitConfig, disableKeepAlives bool, + opts ...Option, ) *Proxy { + var cfg options + for _, o := range opts { + o(&cfg) + } + p := pool.New( maxClientConns, int(maxConnectionAttempts), idleTimeout, disableKeepAlives, + cfg.upstreamTLS, ) proxy := &Proxy{ @@ -108,6 +128,29 @@ func (p *Proxy) ListenAndServe(ctx context.Context) error { return p.Serve(l) } +func (p *Proxy) ListenAndServeTLS(ctx context.Context, certFile, keyFile string) error { + cert, err := tls.LoadX509KeyPair(certFile, keyFile) + if err != nil { + return fmt.Errorf("load proxy TLS cert: %w", err) + } + + tlsCfg := &tls.Config{ + Certificates: []tls.Certificate{cert}, + MinVersion: tls.VersionTLS12, + NextProtos: []string{"h2", "http/1.1"}, + } + + var lisCfg net.ListenConfig + l, err := lisCfg.Listen(ctx, "tcp", p.Addr) + if err != nil { + return err + } + + tlsListener := tls.NewListener(l, tlsCfg) + + return p.Serve(tlsListener) +} + func (p *Proxy) Serve(l net.Listener) error { return p.Server.Serve(tracking.NewListener(l, &p.currentServerConnsCounter)) } From 1ea650f58b9297f562f2e2602dfeef460ecab3d5 Mon Sep 17 00:00:00 2001 From: Ulf Weltman Date: Fri, 10 Jul 2026 16:37:16 -0700 Subject: [PATCH 3/5] Add unit tests for TLS server and client paths (cherry picked from commit ac770b2c5) --- packages/shared/pkg/grpc/tls_test.go | 264 +++++++++++++++++ packages/shared/pkg/proxy/tls_test.go | 411 ++++++++++++++++++++++++++ 2 files changed, 675 insertions(+) create mode 100644 packages/shared/pkg/grpc/tls_test.go create mode 100644 packages/shared/pkg/proxy/tls_test.go diff --git a/packages/shared/pkg/grpc/tls_test.go b/packages/shared/pkg/grpc/tls_test.go new file mode 100644 index 0000000000..a8f3361dd3 --- /dev/null +++ b/packages/shared/pkg/grpc/tls_test.go @@ -0,0 +1,264 @@ +package grpc + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "net" + "os" + "path/filepath" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "google.golang.org/grpc" + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials/insecure" + healthpb "google.golang.org/grpc/health/grpc_health_v1" + "google.golang.org/grpc/status" + + "github.com/e2b-dev/infra/packages/shared/pkg/telemetry" +) + +type testCA struct { + cert *x509.Certificate + key *ecdsa.PrivateKey + certPEM []byte + pool *x509.CertPool +} + +type testCert struct { + certPEM []byte + keyPEM []byte + certFile string + keyFile string +} + +func newTestCA(t *testing.T) *testCA { + t.Helper() + + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + + template := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{Organization: []string{"Test"}, CommonName: "Test CA"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageCRLSign, + BasicConstraintsValid: true, + IsCA: true, + } + + certDER, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key) + require.NoError(t, err) + + cert, err := x509.ParseCertificate(certDER) + require.NoError(t, err) + + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER}) + + certPool := x509.NewCertPool() + certPool.AddCert(cert) + + return &testCA{cert: cert, key: key, certPEM: certPEM, pool: certPool} +} + +func (ca *testCA) issueCert(t *testing.T, hosts ...string) *testCert { + t.Helper() + + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + + template := &x509.Certificate{ + SerialNumber: big.NewInt(2), + Subject: pkix.Name{CommonName: "test-server"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + } + + for _, h := range hosts { + if ip := net.ParseIP(h); ip != nil { + template.IPAddresses = append(template.IPAddresses, ip) + } else { + template.DNSNames = append(template.DNSNames, h) + } + } + + certDER, err := x509.CreateCertificate(rand.Reader, template, ca.cert, &key.PublicKey, ca.key) + require.NoError(t, err) + + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER}) + + keyDER, err := x509.MarshalECPrivateKey(key) + require.NoError(t, err) + keyPEM := pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}) + + dir := t.TempDir() + certFile := filepath.Join(dir, "cert.pem") + keyFile := filepath.Join(dir, "key.pem") + require.NoError(t, os.WriteFile(certFile, certPEM, 0600)) + require.NoError(t, os.WriteFile(keyFile, keyPEM, 0600)) + + return &testCert{certPEM: certPEM, keyPEM: keyPEM, certFile: certFile, keyFile: keyFile} +} + +func TestWithTLS_ServerAcceptsTLSConnections(t *testing.T) { + t.Parallel() + + ca := newTestCA(t) + cert := ca.issueCert(t, "127.0.0.1") + + tel := &telemetry.Client{} + srv := NewGRPCServer(tel, WithTLS(cert.certFile, cert.keyFile)) + healthpb.RegisterHealthServer(srv, &healthpb.UnimplementedHealthServer{}) + + lis, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + go srv.Serve(lis) + t.Cleanup(func() { srv.Stop() }) + + creds := credentials.NewTLS(&tls.Config{RootCAs: ca.pool}) + conn, err := grpc.NewClient(lis.Addr().String(), grpc.WithTransportCredentials(creds)) + require.NoError(t, err) + defer conn.Close() + + client := healthpb.NewHealthClient(conn) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + _, err = client.Check(ctx, &healthpb.HealthCheckRequest{}) + // Unimplemented means the TLS handshake and gRPC framing succeeded + assertGRPCReachable(t, err) +} + +func TestWithTLSFromPEM_ServerAcceptsTLSConnections(t *testing.T) { + t.Parallel() + + ca := newTestCA(t) + cert := ca.issueCert(t, "127.0.0.1") + + tel := &telemetry.Client{} + srv := NewGRPCServer(tel, WithTLSFromPEM(cert.certPEM, cert.keyPEM)) + healthpb.RegisterHealthServer(srv, &healthpb.UnimplementedHealthServer{}) + + lis, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + go srv.Serve(lis) + t.Cleanup(func() { srv.Stop() }) + + creds := credentials.NewTLS(&tls.Config{RootCAs: ca.pool}) + conn, err := grpc.NewClient(lis.Addr().String(), grpc.WithTransportCredentials(creds)) + require.NoError(t, err) + defer conn.Close() + + client := healthpb.NewHealthClient(conn) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + _, err = client.Check(ctx, &healthpb.HealthCheckRequest{}) + assertGRPCReachable(t, err) +} + +func TestWithTLS_WrongCA_Rejected(t *testing.T) { + t.Parallel() + + ca := newTestCA(t) + wrongCA := newTestCA(t) + cert := ca.issueCert(t, "127.0.0.1") + + tel := &telemetry.Client{} + srv := NewGRPCServer(tel, WithTLS(cert.certFile, cert.keyFile)) + healthpb.RegisterHealthServer(srv, &healthpb.UnimplementedHealthServer{}) + + lis, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + go srv.Serve(lis) + t.Cleanup(func() { srv.Stop() }) + + creds := credentials.NewTLS(&tls.Config{RootCAs: wrongCA.pool}) + conn, err := grpc.NewClient(lis.Addr().String(), grpc.WithTransportCredentials(creds)) + require.NoError(t, err) + defer conn.Close() + + client := healthpb.NewHealthClient(conn) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + _, err = client.Check(ctx, &healthpb.HealthCheckRequest{}) + require.Error(t, err) + assert.Contains(t, err.Error(), "certificate signed by unknown authority") +} + +func TestNoTLS_PlaintextFallback(t *testing.T) { + t.Parallel() + + tel := &telemetry.Client{} + srv := NewGRPCServer(tel) + healthpb.RegisterHealthServer(srv, &healthpb.UnimplementedHealthServer{}) + + lis, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + go srv.Serve(lis) + t.Cleanup(func() { srv.Stop() }) + + conn, err := grpc.NewClient(lis.Addr().String(), grpc.WithTransportCredentials(insecure.NewCredentials())) + require.NoError(t, err) + defer conn.Close() + + client := healthpb.NewHealthClient(conn) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + _, err = client.Check(ctx, &healthpb.HealthCheckRequest{}) + assertGRPCReachable(t, err) +} + +func TestWithTLS_InsecureClient_Fails(t *testing.T) { + t.Parallel() + + ca := newTestCA(t) + cert := ca.issueCert(t, "127.0.0.1") + + tel := &telemetry.Client{} + srv := NewGRPCServer(tel, WithTLS(cert.certFile, cert.keyFile)) + healthpb.RegisterHealthServer(srv, &healthpb.UnimplementedHealthServer{}) + + lis, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + go srv.Serve(lis) + t.Cleanup(func() { srv.Stop() }) + + conn, err := grpc.NewClient(lis.Addr().String(), grpc.WithTransportCredentials(insecure.NewCredentials())) + require.NoError(t, err) + defer conn.Close() + + client := healthpb.NewHealthClient(conn) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + _, err = client.Check(ctx, &healthpb.HealthCheckRequest{}) + require.Error(t, err) +} + +func assertGRPCReachable(t *testing.T, err error) { + t.Helper() + if err == nil { + return + } + st, ok := status.FromError(err) + if ok && st.Code() == codes.Unimplemented { + return + } + t.Fatalf("expected reachable gRPC server, got: %v", err) +} diff --git a/packages/shared/pkg/proxy/tls_test.go b/packages/shared/pkg/proxy/tls_test.go new file mode 100644 index 0000000000..453c33ceb3 --- /dev/null +++ b/packages/shared/pkg/proxy/tls_test.go @@ -0,0 +1,411 @@ +package proxy + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "fmt" + "io" + "math/big" + "net" + "net/http" + "net/url" + "os" + "path/filepath" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/e2b-dev/infra/packages/shared/pkg/logger" + "github.com/e2b-dev/infra/packages/shared/pkg/proxy/pool" +) + +type testCA struct { + cert *x509.Certificate + key *ecdsa.PrivateKey + certPEM []byte + pool *x509.CertPool +} + +type testCert struct { + certPEM []byte + keyPEM []byte + certFile string + keyFile string +} + +func newTestCA(t *testing.T) *testCA { + t.Helper() + + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + + template := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{Organization: []string{"Test"}, CommonName: "Test CA"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageCRLSign, + BasicConstraintsValid: true, + IsCA: true, + } + + certDER, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key) + require.NoError(t, err) + + cert, err := x509.ParseCertificate(certDER) + require.NoError(t, err) + + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER}) + + certPool := x509.NewCertPool() + certPool.AddCert(cert) + + return &testCA{cert: cert, key: key, certPEM: certPEM, pool: certPool} +} + +func (ca *testCA) issueCert(t *testing.T, hosts ...string) *testCert { + t.Helper() + + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + + template := &x509.Certificate{ + SerialNumber: big.NewInt(2), + Subject: pkix.Name{CommonName: "test-server"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + } + + for _, h := range hosts { + if ip := net.ParseIP(h); ip != nil { + template.IPAddresses = append(template.IPAddresses, ip) + } else { + template.DNSNames = append(template.DNSNames, h) + } + } + + certDER, err := x509.CreateCertificate(rand.Reader, template, ca.cert, &key.PublicKey, ca.key) + require.NoError(t, err) + + certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER}) + + keyDER, err := x509.MarshalECPrivateKey(key) + require.NoError(t, err) + keyPEM := pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}) + + dir := t.TempDir() + certFile := filepath.Join(dir, "cert.pem") + keyFile := filepath.Join(dir, "key.pem") + require.NoError(t, os.WriteFile(certFile, certPEM, 0600)) + require.NoError(t, os.WriteFile(keyFile, keyPEM, 0600)) + + return &testCert{certPEM: certPEM, keyPEM: keyPEM, certFile: certFile, keyFile: keyFile} +} + +func TestListenAndServeTLS(t *testing.T) { + t.Parallel() + + ca := newTestCA(t) + cert := ca.issueCert(t, "127.0.0.1", "localhost") + + backend := startPlaintextBackend(t) + + proxy := New( + 0, + ClientProxyRetries, + 20*time.Second, + func(*http.Request) (*pool.Destination, error) { + return &pool.Destination{ + Url: backend, + SandboxId: "test", + RequestLogger: logger.NewNopLogger(), + ConnectionKey: "test", + }, nil + }, + nil, + false, + ) + + var lisCfg net.ListenConfig + l, err := lisCfg.Listen(t.Context(), "tcp", "127.0.0.1:0") + require.NoError(t, err) + port := l.Addr().(*net.TCPAddr).Port + l.Close() + + proxy.Addr = fmt.Sprintf("127.0.0.1:%d", port) + + go func() { + _ = proxy.ListenAndServeTLS(t.Context(), cert.certFile, cert.keyFile) + }() + t.Cleanup(func() { proxy.Close() }) + + waitForPort(t, port) + + client := &http.Client{ + Transport: &http.Transport{ + TLSClientConfig: &tls.Config{RootCAs: ca.pool}, + }, + } + + resp, err := client.Get(fmt.Sprintf("https://127.0.0.1:%d/hello", port)) + require.NoError(t, err) + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + assert.Equal(t, http.StatusOK, resp.StatusCode) + assert.Equal(t, "ok", string(body)) +} + +func TestListenAndServeTLS_WrongCA_Rejected(t *testing.T) { + t.Parallel() + + ca := newTestCA(t) + wrongCA := newTestCA(t) + cert := ca.issueCert(t, "127.0.0.1") + + backend := startPlaintextBackend(t) + + proxy := New( + 0, + ClientProxyRetries, + 20*time.Second, + func(*http.Request) (*pool.Destination, error) { + return &pool.Destination{ + Url: backend, + SandboxId: "test", + RequestLogger: logger.NewNopLogger(), + ConnectionKey: "test", + }, nil + }, + nil, + false, + ) + + var lisCfg net.ListenConfig + l, err := lisCfg.Listen(t.Context(), "tcp", "127.0.0.1:0") + require.NoError(t, err) + port := l.Addr().(*net.TCPAddr).Port + l.Close() + + proxy.Addr = fmt.Sprintf("127.0.0.1:%d", port) + + go func() { + _ = proxy.ListenAndServeTLS(t.Context(), cert.certFile, cert.keyFile) + }() + t.Cleanup(func() { proxy.Close() }) + + waitForPort(t, port) + + client := &http.Client{ + Transport: &http.Transport{ + TLSClientConfig: &tls.Config{RootCAs: wrongCA.pool}, + }, + } + + _, err = client.Get(fmt.Sprintf("https://127.0.0.1:%d/hello", port)) + require.Error(t, err) + assert.Contains(t, err.Error(), "certificate signed by unknown authority") +} + +func TestWithUpstreamTLS(t *testing.T) { + t.Parallel() + + ca := newTestCA(t) + cert := ca.issueCert(t, "127.0.0.1") + + backend := startTLSBackend(t, cert) + + proxy := New( + 0, + ClientProxyRetries, + 20*time.Second, + func(*http.Request) (*pool.Destination, error) { + return &pool.Destination{ + Url: backend, + SandboxId: "test", + RequestLogger: logger.NewNopLogger(), + ConnectionKey: "test", + }, nil + }, + nil, + false, + WithUpstreamTLS(&tls.Config{RootCAs: ca.pool, MinVersion: tls.VersionTLS12}), + ) + + var lisCfg net.ListenConfig + l, err := lisCfg.Listen(t.Context(), "tcp", "127.0.0.1:0") + require.NoError(t, err) + port := l.Addr().(*net.TCPAddr).Port + + go func() { + _ = proxy.Serve(l) + }() + t.Cleanup(func() { proxy.Close() }) + + resp, err := http.Get(fmt.Sprintf("http://127.0.0.1:%d/hello", port)) + require.NoError(t, err) + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + assert.Equal(t, http.StatusOK, resp.StatusCode) + assert.Equal(t, "ok", string(body)) +} + +func TestWithUpstreamTLS_WrongCA_Fails(t *testing.T) { + t.Parallel() + + ca := newTestCA(t) + wrongCA := newTestCA(t) + cert := ca.issueCert(t, "127.0.0.1") + + backend := startTLSBackend(t, cert) + + proxy := New( + 0, + ClientProxyRetries, + 20*time.Second, + func(*http.Request) (*pool.Destination, error) { + return &pool.Destination{ + Url: backend, + SandboxId: "test", + RequestLogger: logger.NewNopLogger(), + ConnectionKey: "test", + }, nil + }, + nil, + false, + WithUpstreamTLS(&tls.Config{RootCAs: wrongCA.pool, MinVersion: tls.VersionTLS12}), + ) + + var lisCfg net.ListenConfig + l, err := lisCfg.Listen(t.Context(), "tcp", "127.0.0.1:0") + require.NoError(t, err) + port := l.Addr().(*net.TCPAddr).Port + + go func() { + _ = proxy.Serve(l) + }() + t.Cleanup(func() { proxy.Close() }) + + resp, err := http.Get(fmt.Sprintf("http://127.0.0.1:%d/hello", port)) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusBadGateway, resp.StatusCode) +} + +func TestPlaintextFallbackWhenNoTLSConfigured(t *testing.T) { + t.Parallel() + + backend := startPlaintextBackend(t) + + proxy := New( + 0, + ClientProxyRetries, + 20*time.Second, + func(*http.Request) (*pool.Destination, error) { + return &pool.Destination{ + Url: backend, + SandboxId: "test", + RequestLogger: logger.NewNopLogger(), + ConnectionKey: "test", + }, nil + }, + nil, + false, + ) + + var lisCfg net.ListenConfig + l, err := lisCfg.Listen(t.Context(), "tcp", "127.0.0.1:0") + require.NoError(t, err) + port := l.Addr().(*net.TCPAddr).Port + + go func() { + _ = proxy.Serve(l) + }() + t.Cleanup(func() { proxy.Close() }) + + resp, err := http.Get(fmt.Sprintf("http://127.0.0.1:%d/hello", port)) + require.NoError(t, err) + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + assert.Equal(t, http.StatusOK, resp.StatusCode) + assert.Equal(t, "ok", string(body)) +} + +func startPlaintextBackend(t *testing.T) *url.URL { + t.Helper() + + var lisCfg net.ListenConfig + l, err := lisCfg.Listen(context.Background(), "tcp", "127.0.0.1:0") + require.NoError(t, err) + + srv := &http.Server{ + Handler: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + w.Write([]byte("ok")) + }), + } + go srv.Serve(l) + t.Cleanup(func() { srv.Close() }) + + u, _ := url.Parse(fmt.Sprintf("http://%s", l.Addr().String())) + return u +} + +func startTLSBackend(t *testing.T, cert *testCert) *url.URL { + t.Helper() + + tlsCert, err := tls.X509KeyPair(cert.certPEM, cert.keyPEM) + require.NoError(t, err) + + var lisCfg net.ListenConfig + l, err := lisCfg.Listen(context.Background(), "tcp", "127.0.0.1:0") + require.NoError(t, err) + + tlsListener := tls.NewListener(l, &tls.Config{ + Certificates: []tls.Certificate{tlsCert}, + MinVersion: tls.VersionTLS12, + }) + + srv := &http.Server{ + Handler: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + w.Write([]byte("ok")) + }), + } + go srv.Serve(tlsListener) + t.Cleanup(func() { srv.Close() }) + + u, _ := url.Parse(fmt.Sprintf("https://%s", l.Addr().String())) + return u +} + +func waitForPort(t *testing.T, port int) { + t.Helper() + + addr := fmt.Sprintf("127.0.0.1:%d", port) + for range 50 { + conn, err := net.DialTimeout("tcp", addr, 50*time.Millisecond) + if err == nil { + conn.Close() + return + } + time.Sleep(20 * time.Millisecond) + } + t.Fatalf("port %d did not become available", port) +} From 9cb8317a42892b2e1bfbd5f706b944404a33943e Mon Sep 17 00:00:00 2001 From: Jakub Novak Date: Mon, 27 Jul 2026 14:34:23 +0000 Subject: [PATCH 4/5] feat(proxy): add dedicated TLS listener Keep plaintext listener unchanged while serving TLS on an explicit address. --- packages/shared/pkg/proxy/proxy.go | 27 +++++++++----- packages/shared/pkg/proxy/tls_test.go | 51 +++++++++++++++++++++------ 2 files changed, 59 insertions(+), 19 deletions(-) diff --git a/packages/shared/pkg/proxy/proxy.go b/packages/shared/pkg/proxy/proxy.go index cacd8e6273..076f96c51d 100644 --- a/packages/shared/pkg/proxy/proxy.go +++ b/packages/shared/pkg/proxy/proxy.go @@ -129,8 +129,24 @@ func (p *Proxy) ListenAndServe(ctx context.Context) error { } func (p *Proxy) ListenAndServeTLS(ctx context.Context, certFile, keyFile string) error { + return p.ListenAndServeTLSOn(ctx, p.Addr, certFile, keyFile) +} + +func (p *Proxy) ListenAndServeTLSOn(ctx context.Context, addr, certFile, keyFile string) error { + var lisCfg net.ListenConfig + l, err := lisCfg.Listen(ctx, "tcp", addr) + if err != nil { + return err + } + + return p.ServeTLS(l, certFile, keyFile) +} + +func (p *Proxy) ServeTLS(l net.Listener, certFile, keyFile string) error { cert, err := tls.LoadX509KeyPair(certFile, keyFile) if err != nil { + l.Close() + return fmt.Errorf("load proxy TLS cert: %w", err) } @@ -140,15 +156,10 @@ func (p *Proxy) ListenAndServeTLS(ctx context.Context, certFile, keyFile string) NextProtos: []string{"h2", "http/1.1"}, } - var lisCfg net.ListenConfig - l, err := lisCfg.Listen(ctx, "tcp", p.Addr) - if err != nil { - return err - } - - tlsListener := tls.NewListener(l, tlsCfg) + trackedListener := tracking.NewListener(l, &p.currentServerConnsCounter) + tlsListener := tls.NewListener(trackedListener, tlsCfg) - return p.Serve(tlsListener) + return p.Server.Serve(tlsListener) } func (p *Proxy) Serve(l net.Listener) error { diff --git a/packages/shared/pkg/proxy/tls_test.go b/packages/shared/pkg/proxy/tls_test.go index 453c33ceb3..2dac992ca6 100644 --- a/packages/shared/pkg/proxy/tls_test.go +++ b/packages/shared/pkg/proxy/tls_test.go @@ -35,8 +35,8 @@ type testCA struct { } type testCert struct { - certPEM []byte - keyPEM []byte + certPEM []byte + keyPEM []byte certFile string keyFile string } @@ -112,7 +112,7 @@ func (ca *testCA) issueCert(t *testing.T, hosts ...string) *testCert { return &testCert{certPEM: certPEM, keyPEM: keyPEM, certFile: certFile, keyFile: keyFile} } -func TestListenAndServeTLS(t *testing.T) { +func TestListenAndServeTLSOn_UsesSeparatePort(t *testing.T) { t.Parallel() ca := newTestCA(t) @@ -137,34 +137,63 @@ func TestListenAndServeTLS(t *testing.T) { ) var lisCfg net.ListenConfig - l, err := lisCfg.Listen(t.Context(), "tcp", "127.0.0.1:0") + plainListener, err := lisCfg.Listen(t.Context(), "tcp", "127.0.0.1:0") require.NoError(t, err) - port := l.Addr().(*net.TCPAddr).Port - l.Close() + plainPort := plainListener.Addr().(*net.TCPAddr).Port + plainListener.Close() - proxy.Addr = fmt.Sprintf("127.0.0.1:%d", port) + tlsListener, err := lisCfg.Listen(t.Context(), "tcp", "127.0.0.1:0") + require.NoError(t, err) + tlsPort := tlsListener.Addr().(*net.TCPAddr).Port + tlsListener.Close() + + proxy.Addr = fmt.Sprintf("127.0.0.1:%d", plainPort) go func() { - _ = proxy.ListenAndServeTLS(t.Context(), cert.certFile, cert.keyFile) + _ = proxy.ListenAndServe(t.Context()) + }() + go func() { + tlsAddr := fmt.Sprintf("127.0.0.1:%d", tlsPort) + _ = proxy.ListenAndServeTLSOn(t.Context(), tlsAddr, cert.certFile, cert.keyFile) }() t.Cleanup(func() { proxy.Close() }) - waitForPort(t, port) + waitForPort(t, plainPort) + waitForPort(t, tlsPort) client := &http.Client{ Transport: &http.Transport{ - TLSClientConfig: &tls.Config{RootCAs: ca.pool}, + TLSClientConfig: &tls.Config{RootCAs: ca.pool}, + ForceAttemptHTTP2: true, }, } - resp, err := client.Get(fmt.Sprintf("https://127.0.0.1:%d/hello", port)) + resp, err := client.Get(fmt.Sprintf("https://127.0.0.1:%d/hello", tlsPort)) require.NoError(t, err) defer resp.Body.Close() body, err := io.ReadAll(resp.Body) require.NoError(t, err) assert.Equal(t, http.StatusOK, resp.StatusCode) + assert.Equal(t, 2, resp.ProtoMajor) assert.Equal(t, "ok", string(body)) + + plaintextResp, err := http.Get(fmt.Sprintf("http://127.0.0.1:%d/hello", plainPort)) + require.NoError(t, err) + defer plaintextResp.Body.Close() + + plaintextBody, err := io.ReadAll(plaintextResp.Body) + require.NoError(t, err) + assert.Equal(t, http.StatusOK, plaintextResp.StatusCode) + assert.Equal(t, "ok", string(plaintextBody)) + + plaintextOnTLSPort, err := http.Get(fmt.Sprintf("http://127.0.0.1:%d/hello", tlsPort)) + require.NoError(t, err) + defer plaintextOnTLSPort.Body.Close() + assert.Equal(t, http.StatusBadRequest, plaintextOnTLSPort.StatusCode) + + _, err = client.Get(fmt.Sprintf("https://127.0.0.1:%d/hello", plainPort)) + require.Error(t, err) } func TestListenAndServeTLS_WrongCA_Rejected(t *testing.T) { From db107cc03d005a96f253ac4fd90a1d4e35414263 Mon Sep 17 00:00:00 2001 From: "github-actions[bot]" Date: Tue, 28 Jul 2026 14:08:59 +0000 Subject: [PATCH 5/5] chore: auto-commit generated changes --- packages/shared/pkg/grpc/server.go | 3 +-- packages/shared/pkg/grpc/tls_test.go | 6 +++--- packages/shared/pkg/proxy/tls_test.go | 4 ++-- 3 files changed, 6 insertions(+), 7 deletions(-) diff --git a/packages/shared/pkg/grpc/server.go b/packages/shared/pkg/grpc/server.go index 1e12c6bd4a..db3e92866f 100644 --- a/packages/shared/pkg/grpc/server.go +++ b/packages/shared/pkg/grpc/server.go @@ -5,13 +5,12 @@ import ( "crypto/tls" "time" - "go.uber.org/zap" - "github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/logging" "github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/recovery" "github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/selector" "go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc" "go.opentelemetry.io/otel/attribute" + "go.uber.org/zap" "google.golang.org/grpc" "google.golang.org/grpc/credentials" "google.golang.org/grpc/keepalive" diff --git a/packages/shared/pkg/grpc/tls_test.go b/packages/shared/pkg/grpc/tls_test.go index a8f3361dd3..eeb75c8f06 100644 --- a/packages/shared/pkg/grpc/tls_test.go +++ b/packages/shared/pkg/grpc/tls_test.go @@ -19,8 +19,8 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "google.golang.org/grpc" - "google.golang.org/grpc/credentials" "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials" "google.golang.org/grpc/credentials/insecure" healthpb "google.golang.org/grpc/health/grpc_health_v1" "google.golang.org/grpc/status" @@ -107,8 +107,8 @@ func (ca *testCA) issueCert(t *testing.T, hosts ...string) *testCert { dir := t.TempDir() certFile := filepath.Join(dir, "cert.pem") keyFile := filepath.Join(dir, "key.pem") - require.NoError(t, os.WriteFile(certFile, certPEM, 0600)) - require.NoError(t, os.WriteFile(keyFile, keyPEM, 0600)) + require.NoError(t, os.WriteFile(certFile, certPEM, 0o600)) + require.NoError(t, os.WriteFile(keyFile, keyPEM, 0o600)) return &testCert{certPEM: certPEM, keyPEM: keyPEM, certFile: certFile, keyFile: keyFile} } diff --git a/packages/shared/pkg/proxy/tls_test.go b/packages/shared/pkg/proxy/tls_test.go index 2dac992ca6..69fc838c03 100644 --- a/packages/shared/pkg/proxy/tls_test.go +++ b/packages/shared/pkg/proxy/tls_test.go @@ -106,8 +106,8 @@ func (ca *testCA) issueCert(t *testing.T, hosts ...string) *testCert { dir := t.TempDir() certFile := filepath.Join(dir, "cert.pem") keyFile := filepath.Join(dir, "key.pem") - require.NoError(t, os.WriteFile(certFile, certPEM, 0600)) - require.NoError(t, os.WriteFile(keyFile, keyPEM, 0600)) + require.NoError(t, os.WriteFile(certFile, certPEM, 0o600)) + require.NoError(t, os.WriteFile(keyFile, keyPEM, 0o600)) return &testCert{certPEM: certPEM, keyPEM: keyPEM, certFile: certFile, keyFile: keyFile} }