diff --git a/config.go b/config.go index 731b6e0..c41c07b 100644 --- a/config.go +++ b/config.go @@ -3,6 +3,7 @@ package grpc import ( "crypto/tls" stderr "errors" + "fmt" "math" "os" "strings" @@ -11,6 +12,7 @@ import ( "github.com/bmatcuk/doublestar/v4" "github.com/roadrunner-server/errors" "github.com/roadrunner-server/pool/v2/pool" + "github.com/roadrunner-server/tcplisten" ) type ClientAuthType string @@ -24,8 +26,9 @@ const ( ) type Config struct { - Listen string `mapstructure:"listen"` - Proto []string `mapstructure:"proto"` + Listen string `mapstructure:"listen"` + UnixSocket *tcplisten.UnixSocketOptions `mapstructure:"unix_socket"` + Proto []string `mapstructure:"proto"` TLS *TLS `mapstructure:"tls"` @@ -64,6 +67,9 @@ func (c *Config) InitDefaults() error { //nolint:gocyclo,gocognit if !strings.Contains(c.Listen, ":") { return errors.E(op, errors.Errorf("malformed grpc address, provided: %s", c.Listen)) } + if err := c.UnixSocket.Validate(c.Listen); err != nil { + return errors.E(op, fmt.Errorf("grpc.unix_socket: %w", err)) + } protos := make([]string, 0, len(c.Proto)) for _, path := range c.Proto { diff --git a/go.mod b/go.mod index d4c1794..6842790 100644 --- a/go.mod +++ b/go.mod @@ -15,7 +15,7 @@ require ( github.com/roadrunner-server/errors v1.5.0 github.com/roadrunner-server/goridge/v4 v4.0.0-beta.3 github.com/roadrunner-server/pool/v2 v2.0.0-beta.1 - github.com/roadrunner-server/tcplisten v1.5.2 + github.com/roadrunner-server/tcplisten v1.6.0 github.com/stretchr/testify v1.12.1 go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.71.0 go.opentelemetry.io/contrib/propagators/jaeger v1.46.0 diff --git a/go.sum b/go.sum index 2d2dcc3..db41348 100644 --- a/go.sum +++ b/go.sum @@ -44,8 +44,8 @@ github.com/roadrunner-server/goridge/v4 v4.0.0-beta.3 h1:+kUw00/fpqwdMWrPMYW+OZH github.com/roadrunner-server/goridge/v4 v4.0.0-beta.3/go.mod h1:1aHppV68y/VqRED/AsfNg59sft9aQOhqgr5Z5n49jbM= github.com/roadrunner-server/pool/v2 v2.0.0-beta.1 h1:jpYXFtdD6QGAdAGPgMxrNi3j1CegCRpb2y+A+3GnXFA= github.com/roadrunner-server/pool/v2 v2.0.0-beta.1/go.mod h1:Bo1wT7RtL3eyQHXBUohNhtj/yAmRt6Rq8smuBg5pWkY= -github.com/roadrunner-server/tcplisten v1.5.2 h1:nn8yXYrhRDkfQ9AAu4V075uT4fZRmOnpxkawgE+bWPA= -github.com/roadrunner-server/tcplisten v1.5.2/go.mod h1:DufGBz7Dlx2KrNe/4RukEvGMTqZKB0Uve1GztwcyyR8= +github.com/roadrunner-server/tcplisten v1.6.0 h1:xfFeA2PZTmwJdwc/InhJGq200ew/lfTDReF3oa4AyI4= +github.com/roadrunner-server/tcplisten v1.6.0/go.mod h1:M01BcmhsBiek8WfkiRQwVXwVamgZ5YV36Wa0hz937dA= github.com/shirou/gopsutil v3.21.11+incompatible h1:+1+c1VGhc88SSonWP6foOcLhvnKlUeu/erjjvaPEYiI= github.com/shirou/gopsutil v3.21.11+incompatible/go.mod h1:5b4v6he4MtMOwMlS0TUMTu2PcXUg8+E1lC7eC3UO/RA= github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= diff --git a/plugin.go b/plugin.go index 1453db0..15f5988 100644 --- a/plugin.go +++ b/plugin.go @@ -86,7 +86,6 @@ func (p *Plugin) Init(cfg api.Configurer, log api.Logger, server api.Server) err if err != nil { return errors.E(op, err) } - err = p.config.InitDefaults() if err != nil { return errors.E(op, err) @@ -168,7 +167,7 @@ func (p *Plugin) Serve() chan error { return errCh } - l, err := tcplisten.CreateListener(p.config.Listen) + l, err := tcplisten.CreateListenerWithOptions(p.config.Listen, p.config.UnixSocket) if err != nil { errCh <- errors.E(op, err) return errCh diff --git a/schema.json b/schema.json index bc630d7..c682480 100644 --- a/schema.json +++ b/schema.json @@ -20,6 +20,9 @@ "tcp://127.0.0.1:${TCP_PORT}" ] }, + "unix_socket": { + "$ref": "https://raw.githubusercontent.com/roadrunner-server/tcplisten/v1.6.0/schema.json" + }, "proto": { "type": "array", "minItems": 1, diff --git a/tests/go.mod b/tests/go.mod index 1b7caae..fb4b391 100644 --- a/tests/go.mod +++ b/tests/go.mod @@ -52,7 +52,7 @@ require ( github.com/roadrunner-server/errors v1.5.0 // indirect github.com/roadrunner-server/events v1.0.1 // indirect github.com/roadrunner-server/pool/v2 v2.0.0-beta.1 // indirect - github.com/roadrunner-server/tcplisten v1.5.2 // indirect + github.com/roadrunner-server/tcplisten v1.6.0 // indirect github.com/sagikazarmark/locafero v0.12.0 // indirect github.com/shirou/gopsutil v3.21.11+incompatible // indirect github.com/spf13/afero v1.15.0 // indirect diff --git a/tests/go.sum b/tests/go.sum index 73ac0af..c42cc08 100644 --- a/tests/go.sum +++ b/tests/go.sum @@ -92,8 +92,8 @@ github.com/roadrunner-server/server/v6 v6.0.0-beta.7 h1:EiRKdWFPOYLoYy53xoLbyU88 github.com/roadrunner-server/server/v6 v6.0.0-beta.7/go.mod h1:uq0yIZgp1v80BGIHPZHKHFyYIWTVJJvofThaC8QWf7w= github.com/roadrunner-server/status/v6 v6.0.0-beta.8 h1:Ya1/vZnPgRUc4eWDTVihLbptS5z9olBUypwXSnH6olw= github.com/roadrunner-server/status/v6 v6.0.0-beta.8/go.mod h1:y/TqFnBItSHmhvnvZD+xqX5hTcIRGcQBoVrltizy9Pg= -github.com/roadrunner-server/tcplisten v1.5.2 h1:nn8yXYrhRDkfQ9AAu4V075uT4fZRmOnpxkawgE+bWPA= -github.com/roadrunner-server/tcplisten v1.5.2/go.mod h1:DufGBz7Dlx2KrNe/4RukEvGMTqZKB0Uve1GztwcyyR8= +github.com/roadrunner-server/tcplisten v1.6.0 h1:xfFeA2PZTmwJdwc/InhJGq200ew/lfTDReF3oa4AyI4= +github.com/roadrunner-server/tcplisten v1.6.0/go.mod h1:M01BcmhsBiek8WfkiRQwVXwVamgZ5YV36Wa0hz937dA= github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ= github.com/rogpeppe/go-internal v1.14.1/go.mod h1:MaRKkUm5W0goXpeCfT7UZI6fk/L7L7so1lCWt35ZSgc= github.com/sagikazarmark/locafero v0.12.0 h1:/NQhBAkUb4+fH1jivKHWusDYFjMOOKU88eegjfxfHb4= diff --git a/tests/helpers/rr.go b/tests/helpers/rr.go index 03b48cb..516be88 100644 --- a/tests/helpers/rr.go +++ b/tests/helpers/rr.go @@ -29,6 +29,7 @@ const ( // bootCfg holds the options applied to a container before it is started. type bootCfg struct { version string + flags []string logLevel slog.Level logger loggerKind probe func(ctx context.Context) bool @@ -50,6 +51,11 @@ func WithConfigVersion(v string) Option { return func(b *bootCfg) { b.version = v } } +// WithConfigFlags applies configuration overrides. +func WithConfigFlags(flags ...string) Option { + return func(b *bootCfg) { b.flags = flags } +} + // WithLogLevel sets the endure container log level (debug by default). func WithLogLevel(l slog.Level) Option { return func(b *bootCfg) { b.logLevel = l } @@ -171,7 +177,7 @@ func newContainer(t *testing.T, cfgPath string, plugins []any, opts []Option) (* o(bc) } - cfg := &config.Plugin{Version: bc.version, Path: cfgPath} + cfg := &config.Plugin{Version: bc.version, Path: cfgPath, Flags: bc.flags} rr := &RR{} all := []any{cfg} diff --git a/tests/unix_socket_test.go b/tests/unix_socket_test.go new file mode 100644 index 0000000..2e4d57b --- /dev/null +++ b/tests/unix_socket_test.go @@ -0,0 +1,135 @@ +//go:build linux || darwin || freebsd + +package grpc_test + +import ( + "fmt" + "log/slog" + "os" + "path/filepath" + "testing" + + "tests/helpers" + mocklogger "tests/mock" + + "github.com/roadrunner-server/config/v6" + grpcPlugin "github.com/roadrunner-server/grpc/v6" + "github.com/stretchr/testify/require" +) + +func TestUnixSocketConfigRejectsInvalidOptions(t *testing.T) { + log := mocklogger.NewLogger(slog.New(slog.DiscardHandler)) + + for _, tc := range []struct { + name string + listen string + options string + wantErr string + }{ + {name: "TCP options", listen: "tcp://127.0.0.1:0", options: `{mode: "0600"}`, wantErr: "filesystem unix:// address"}, + {name: "invalid mode", listen: "unix://grpc.sock", options: `{mode: "0780"}`, wantErr: "invalid unix socket mode"}, + {name: "unquoted mode", listen: "unix://grpc.sock", options: "{mode: 0600}", wantErr: "invalid unix socket mode"}, + {name: "empty socket path", listen: "unix://", options: `{mode: "0600"}`, wantErr: "filesystem unix:// address"}, + {name: "scalar options", listen: "unix://grpc.sock", options: "false", wantErr: "expected a map"}, + {name: "negative UID", listen: "unix://grpc.sock", options: "{uid: -1}", wantErr: "invalid unix socket uid"}, + {name: "negative GID", listen: "unix://grpc.sock", options: "{gid: -1}", wantErr: "invalid unix socket gid"}, + {name: "reserved UID", listen: "unix://grpc.sock", options: "{uid: 4294967295}", wantErr: "invalid unix socket uid"}, + {name: "reserved GID", listen: "unix://grpc.sock", options: "{gid: 4294967295}", wantErr: "invalid unix socket gid"}, + } { + t.Run(tc.name, func(t *testing.T) { + cfg := &config.Plugin{Path: unixSocketConfig(t, tc.listen, tc.options)} + require.NoError(t, cfg.Init()) + p := &grpcPlugin.Plugin{} + require.ErrorContains(t, p.Init(cfg, log, nil), tc.wantErr) + }) + } +} + +func TestUnixSocketMode(t *testing.T) { + for _, tc := range []struct { + name string + fileMode string + flags []string + wantMode os.FileMode + }{ + {name: "quoted 0600", fileMode: "0600", wantMode: 0o600}, + {name: "quoted 0640", fileMode: "0640", wantMode: 0o640}, + {name: "mode override", fileMode: "0600", flags: []string{"grpc.unix_socket.mode=0640"}, wantMode: 0o640}, + } { + t.Run(tc.name, func(t *testing.T) { + socket := unixSocketPath(t) + cfgPath := unixSocketConfig(t, "unix://"+socket, fmt.Sprintf(`{mode: %q}`, tc.fileMode)) + helpers.Start(t, cfgPath, grpcPlugins(), helpers.WithConfigFlags(tc.flags...)) + + info, err := os.Stat(socket) + require.NoError(t, err) + require.Equal(t, tc.wantMode, info.Mode().Perm()) + }) + } +} + +func TestUnixSocketPing(t *testing.T) { + socket := unixSocketPath(t) + cfgPath := unixSocketConfig(t, "unix://"+socket, `{mode: "0600"}`) + helpers.Start(t, cfgPath, grpcPlugins()) + + got, err := ping(t, helpers.Dial(t, "unix://"+socket), "TOST") + + require.NoError(t, err) + require.Equal(t, "TOST", got) +} + +func TestUnixSocketStopRemovesListener(t *testing.T) { + socket := unixSocketPath(t) + cfgPath := unixSocketConfig(t, "unix://"+socket, `{mode: "0600"}`) + _, stop := helpers.Start(t, cfgPath, grpcPlugins()) + require.FileExists(t, socket) + + stop() + + require.NoFileExists(t, socket) +} + +func TestUnixSocketOwnershipErrorRemovesListener(t *testing.T) { + if os.Geteuid() == 0 { + t.Skip("Requires an unprivileged process.") + } + + socket := unixSocketPath(t) + cfgPath := unixSocketConfig(t, "unix://"+socket, "{uid: 0}") + err := helpers.StartExpectServeError(t, cfgPath, grpcPlugins()) + require.ErrorContains(t, err, "chown unix socket") + + _, err = os.Stat(socket) + require.ErrorIs(t, err, os.ErrNotExist) +} + +func unixSocketPath(t *testing.T) string { + t.Helper() + + // Keep the socket path below the macOS length limit. + dir, err := os.MkdirTemp("", "rr-grpc-") + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, os.RemoveAll(dir)) }) + return filepath.Join(dir, "grpc.sock") +} + +func unixSocketConfig(t *testing.T, listen, options string) string { + t.Helper() + + contents := fmt.Sprintf(`version: "3" +server: + command: "php php_test_files/worker-grpc.php" +grpc: + listen: %q + proto: + - "proto/service/service.proto" + unix_socket: %s + pool: + debug: true + destroy_timeout: 5s +`, listen, options) + path := filepath.Join(t.TempDir(), ".rr.yaml") + require.NoError(t, os.WriteFile(path, []byte(contents), 0o600)) + return path +}