Skip to content
Closed
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
50 changes: 48 additions & 2 deletions packages/shared/pkg/grpc/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,14 +2,17 @@ package grpc

import (
"context"
"crypto/tls"
"time"

"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"
"google.golang.org/grpc/metadata"

Expand All @@ -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,
Expand All @@ -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 {
Expand Down Expand Up @@ -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,
Expand All @@ -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))
}

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Incomplete TLS silently disables encryption

Medium Severity

WithTLS / WithTLSFromPEM only enable credentials when both cert and key are non-empty. If either side is missing or empty, NewGRPCServer falls through and serves plaintext with no error, so a misconfigured TLS enablement path can silently leave traffic unencrypted.

Fix in Cursor Fix in Web

Reviewed by Cursor Bugbot for commit db107cc. Configure here.


return grpc.NewServer(serverOpts...)
}

// extractSandboxResumeAttrs reads sandbox.resume from gRPC metadata set by the
Expand Down
264 changes: 264 additions & 0 deletions packages/shared/pkg/grpc/tls_test.go
Original file line number Diff line number Diff line change
@@ -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")

Check failure on line 126 in packages/shared/pkg/grpc/tls_test.go

View workflow job for this annotation

GitHub Actions / lint / golangci-lint (/home/runner/work/infra/infra/packages/shared)

net.Listen must not be called. use (*net.ListenConfig).Listen (noctx)
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")

Check failure on line 155 in packages/shared/pkg/grpc/tls_test.go

View workflow job for this annotation

GitHub Actions / lint / golangci-lint (/home/runner/work/infra/infra/packages/shared)

net.Listen must not be called. use (*net.ListenConfig).Listen (noctx)
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")

Check failure on line 184 in packages/shared/pkg/grpc/tls_test.go

View workflow job for this annotation

GitHub Actions / lint / golangci-lint (/home/runner/work/infra/infra/packages/shared)

net.Listen must not be called. use (*net.ListenConfig).Listen (noctx)
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")

Check failure on line 210 in packages/shared/pkg/grpc/tls_test.go

View workflow job for this annotation

GitHub Actions / lint / golangci-lint (/home/runner/work/infra/infra/packages/shared)

net.Listen must not be called. use (*net.ListenConfig).Listen (noctx)
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")

Check failure on line 237 in packages/shared/pkg/grpc/tls_test.go

View workflow job for this annotation

GitHub Actions / lint / golangci-lint (/home/runner/work/infra/infra/packages/shared)

net.Listen must not be called. use (*net.ListenConfig).Listen (noctx)
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)
}
3 changes: 3 additions & 0 deletions packages/shared/pkg/proxy/pool/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package pool

import (
"context"
"crypto/tls"
"errors"
"log"
"net"
Expand Down Expand Up @@ -36,6 +37,7 @@ func newProxyClient(
currentConnsCounter *atomic.Int64,
l *log.Logger,
disableKeepAlives bool,
tlsConfig *tls.Config,
) *ProxyClient {
activeConnections := smap.New[*tracking.Connection]()

Expand All @@ -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
Expand Down
Loading
Loading