diff --git a/packages/shared/pkg/grpc/server.go b/packages/shared/pkg/grpc/server.go index 1e2e039b77..db3e92866f 100644 --- a/packages/shared/pkg/grpc/server.go +++ b/packages/shared/pkg/grpc/server.go @@ -2,6 +2,7 @@ package grpc import ( "context" + "crypto/tls" "time" "github.com/grpc-ecosystem/go-grpc-middleware/v2/interceptors/logging" @@ -9,7 +10,9 @@ import ( "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" "google.golang.org/grpc/metadata" @@ -22,6 +25,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 +37,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 +80,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 +105,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 diff --git a/packages/shared/pkg/grpc/tls_test.go b/packages/shared/pkg/grpc/tls_test.go new file mode 100644 index 0000000000..eeb75c8f06 --- /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/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" + + "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, 0o600)) + require.NoError(t, os.WriteFile(keyFile, keyPEM, 0o600)) + + 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/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..076f96c51d 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,40 @@ func (p *Proxy) ListenAndServe(ctx context.Context) error { return p.Serve(l) } +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) + } + + tlsCfg := &tls.Config{ + Certificates: []tls.Certificate{cert}, + MinVersion: tls.VersionTLS12, + NextProtos: []string{"h2", "http/1.1"}, + } + + trackedListener := tracking.NewListener(l, &p.currentServerConnsCounter) + tlsListener := tls.NewListener(trackedListener, tlsCfg) + + return p.Server.Serve(tlsListener) +} + func (p *Proxy) Serve(l net.Listener) error { return p.Server.Serve(tracking.NewListener(l, &p.currentServerConnsCounter)) } diff --git a/packages/shared/pkg/proxy/tls_test.go b/packages/shared/pkg/proxy/tls_test.go new file mode 100644 index 0000000000..69fc838c03 --- /dev/null +++ b/packages/shared/pkg/proxy/tls_test.go @@ -0,0 +1,440 @@ +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, 0o600)) + require.NoError(t, os.WriteFile(keyFile, keyPEM, 0o600)) + + return &testCert{certPEM: certPEM, keyPEM: keyPEM, certFile: certFile, keyFile: keyFile} +} + +func TestListenAndServeTLSOn_UsesSeparatePort(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 + plainListener, err := lisCfg.Listen(t.Context(), "tcp", "127.0.0.1:0") + require.NoError(t, err) + plainPort := plainListener.Addr().(*net.TCPAddr).Port + plainListener.Close() + + 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.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, plainPort) + waitForPort(t, tlsPort) + + client := &http.Client{ + Transport: &http.Transport{ + TLSClientConfig: &tls.Config{RootCAs: ca.pool}, + ForceAttemptHTTP2: true, + }, + } + + 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) { + 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) +}