From 897f30fd36e3166b89ab9d09ced81503f7e9894e Mon Sep 17 00:00:00 2001 From: Justin Bradford Date: Wed, 8 Apr 2026 02:52:42 -0700 Subject: [PATCH] feat: add client/server version check mechanism to gRPC calls (#260) * Add version check mechanism to gRPC calls * Use semver not semver/v3 * Give dev builds a special "infinite" version number (999.0.0-dev) * Handle no metadata on grpc call correctly for version check * Move versioncheck package to root pkg/ from pkg/api/ since it is shared by both pkg/api/ and pkg/client/ * Only show "no daemon version" warning once * Append version headers to metadata, not overwrite... * Unit tests on versioncheck logic * Move SetHeader to the ServerStream in ServerStreamInterceptor * Go modernizer nits: interface{} -> any * Use more conventional gRPC header names for version/min-versions * Add TODO notes on checkDaemonVersionInResponse and related code that can be removed eventually after transition to version checking client/daemons * Use testify for testing assertions * Add explanatory comments on MinCLIVersion and MinDaemonVersion --------- Co-authored-by: Pasha Sviderski --- AGENTS.md | 1 + internal/machine/machine.go | 5 + internal/version/version.go | 3 + pkg/client/connector/ssh.go | 3 + pkg/client/connector/sshcli.go | 3 + pkg/client/connector/tcp.go | 3 + pkg/client/connector/unix.go | 3 + pkg/client/connector/wireguard.go | 3 + pkg/versioncheck/interceptor.go | 179 ++++++++++++++++++++++ pkg/versioncheck/interceptor_test.go | 216 +++++++++++++++++++++++++++ 10 files changed, 419 insertions(+) create mode 100644 pkg/versioncheck/interceptor.go create mode 100644 pkg/versioncheck/interceptor_test.go diff --git a/AGENTS.md b/AGENTS.md index 951835cd..08fe7d4d 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -197,6 +197,7 @@ uc context use # Switch context - Integration tests in `test/e2e/` - Test fixtures in `test/fixtures/` - Use table driven tests whenever possible +- Use the `testify` library for assertions (e.g., `require.Equal`, `assert.Nil`) ### Dependencies diff --git a/internal/machine/machine.go b/internal/machine/machine.go index 0c0f730c..2c3bfa03 100644 --- a/internal/machine/machine.go +++ b/internal/machine/machine.go @@ -34,6 +34,7 @@ import ( "github.com/psviderski/uncloud/internal/machine/network" "github.com/psviderski/uncloud/internal/machine/store" "github.com/psviderski/uncloud/pkg/api" + versionpkg "github.com/psviderski/uncloud/pkg/versioncheck" "github.com/psviderski/unregistry" "github.com/siderolabs/grpc-proxy/proxy" "golang.org/x/sync/errgroup" @@ -272,6 +273,8 @@ func NewMachine(config *Config) (*Machine, error) { proxyDirector := apiproxy.NewDirector(config.MachineSockPath, constants.MachineAPIPort) localProxyServer := grpc.NewServer( grpc.ForceServerCodecV2(proxy.Codec()), + grpc.UnaryInterceptor(versionpkg.ServerUnaryInterceptor), + grpc.StreamInterceptor(versionpkg.ServerStreamInterceptor), grpc.UnknownServiceHandler( proxy.TransparentHandler(proxyDirector.Director), ), @@ -425,6 +428,8 @@ func (m *Machine) Run(ctx context.Context) error { m.proxyDirector.UpdateLocalAddress(m.state.Network.ManagementIP.String()) proxyServer := grpc.NewServer( grpc.ForceServerCodecV2(proxy.Codec()), + grpc.UnaryInterceptor(versionpkg.ServerUnaryInterceptor), + grpc.StreamInterceptor(versionpkg.ServerStreamInterceptor), grpc.UnknownServiceHandler( proxy.TransparentHandler(m.proxyDirector.Director), ), diff --git a/internal/version/version.go b/internal/version/version.go index c05c1c24..8b5850d9 100644 --- a/internal/version/version.go +++ b/internal/version/version.go @@ -3,5 +3,8 @@ package version var version string func String() string { + if version == "" { + return "999.0.0-dev" + } return version } diff --git a/pkg/client/connector/ssh.go b/pkg/client/connector/ssh.go index dfd06595..0dc31886 100644 --- a/pkg/client/connector/ssh.go +++ b/pkg/client/connector/ssh.go @@ -9,6 +9,7 @@ import ( "github.com/psviderski/uncloud/internal/machine" "github.com/psviderski/uncloud/internal/sshexec" + "github.com/psviderski/uncloud/pkg/versioncheck" "golang.org/x/crypto/ssh" "golang.org/x/net/proxy" "google.golang.org/grpc" @@ -74,6 +75,8 @@ func (c *SSHConnector) Connect(ctx context.Context) (*grpc.ClientConn, error) { "unix://"+sockPath, grpc.WithTransportCredentials(insecure.NewCredentials()), grpc.WithDefaultServiceConfig(defaultServiceConfig), + grpc.WithUnaryInterceptor(versioncheck.ClientUnaryInterceptor), + grpc.WithStreamInterceptor(versioncheck.ClientStreamInterceptor), grpc.WithContextDialer( func(ctx context.Context, addr string) (net.Conn, error) { addr = strings.TrimPrefix(addr, "unix://") diff --git a/pkg/client/connector/sshcli.go b/pkg/client/connector/sshcli.go index f32474d9..ddfdd4c6 100644 --- a/pkg/client/connector/sshcli.go +++ b/pkg/client/connector/sshcli.go @@ -11,6 +11,7 @@ import ( "strings" "github.com/docker/cli/cli/connhelper/commandconn" + "github.com/psviderski/uncloud/pkg/versioncheck" "golang.org/x/net/proxy" "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" @@ -77,6 +78,8 @@ func (c *SSHCLIConnector) Connect(ctx context.Context) (*grpc.ClientConn, error) "passthrough:///", // Dummy target since we're using a custom dialer. grpc.WithTransportCredentials(insecure.NewCredentials()), grpc.WithDefaultServiceConfig(defaultServiceConfig), + grpc.WithUnaryInterceptor(versioncheck.ClientUnaryInterceptor), + grpc.WithStreamInterceptor(versioncheck.ClientStreamInterceptor), grpc.WithContextDialer(func(ctx context.Context, _ string) (net.Conn, error) { dialArgs := append(c.buildSSHArgs(), "uncloudd", "dial-stdio") if c.config.SockPath != "" { diff --git a/pkg/client/connector/tcp.go b/pkg/client/connector/tcp.go index 087072f8..1e1d24e2 100644 --- a/pkg/client/connector/tcp.go +++ b/pkg/client/connector/tcp.go @@ -5,6 +5,7 @@ import ( "fmt" "net/netip" + "github.com/psviderski/uncloud/pkg/versioncheck" "golang.org/x/net/proxy" "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" @@ -24,6 +25,8 @@ func (c *TCPConnector) Connect(_ context.Context) (*grpc.ClientConn, error) { c.apiAddr.String(), grpc.WithTransportCredentials(insecure.NewCredentials()), grpc.WithDefaultServiceConfig(defaultServiceConfig), + grpc.WithUnaryInterceptor(versioncheck.ClientUnaryInterceptor), + grpc.WithStreamInterceptor(versioncheck.ClientStreamInterceptor), ) if err != nil { return nil, fmt.Errorf("create machine API client: %w", err) diff --git a/pkg/client/connector/unix.go b/pkg/client/connector/unix.go index 6a94abd9..b03e6013 100644 --- a/pkg/client/connector/unix.go +++ b/pkg/client/connector/unix.go @@ -4,6 +4,7 @@ import ( "context" "fmt" + "github.com/psviderski/uncloud/pkg/versioncheck" "golang.org/x/net/proxy" "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" @@ -26,6 +27,8 @@ func (c *UnixConnector) Connect(_ context.Context) (*grpc.ClientConn, error) { target, grpc.WithTransportCredentials(insecure.NewCredentials()), grpc.WithDefaultServiceConfig(defaultServiceConfig), + grpc.WithUnaryInterceptor(versioncheck.ClientUnaryInterceptor), + grpc.WithStreamInterceptor(versioncheck.ClientStreamInterceptor), ) if err != nil { return nil, fmt.Errorf("create machine API client: %w", err) diff --git a/pkg/client/connector/wireguard.go b/pkg/client/connector/wireguard.go index a3898489..e02b9c8c 100644 --- a/pkg/client/connector/wireguard.go +++ b/pkg/client/connector/wireguard.go @@ -12,6 +12,7 @@ import ( "github.com/psviderski/uncloud/internal/machine/network" "github.com/psviderski/uncloud/internal/machine/network/tunnel" "github.com/psviderski/uncloud/pkg/client" + "github.com/psviderski/uncloud/pkg/versioncheck" "golang.org/x/net/proxy" "google.golang.org/grpc" "google.golang.org/grpc/credentials/insecure" @@ -70,6 +71,8 @@ func (c *WireGuardConnector) Connect(ctx context.Context) (*grpc.ClientConn, err grpc.WithContextDialer(func(ctx context.Context, addr string) (net.Conn, error) { return c.tun.DialContext(ctx, "tcp", addr) }), + grpc.WithUnaryInterceptor(versioncheck.ClientUnaryInterceptor), + grpc.WithStreamInterceptor(versioncheck.ClientStreamInterceptor), ) if err != nil { return nil, fmt.Errorf("connect to machine API through WireGuard tunnel: %w", err) diff --git a/pkg/versioncheck/interceptor.go b/pkg/versioncheck/interceptor.go new file mode 100644 index 00000000..893d85c2 --- /dev/null +++ b/pkg/versioncheck/interceptor.go @@ -0,0 +1,179 @@ +package versioncheck + +import ( + "context" + "fmt" + "os" + + "github.com/Masterminds/semver" + internalVersion "github.com/psviderski/uncloud/internal/version" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/metadata" + "google.golang.org/grpc/status" +) + +const ( + MetadataKeyCLIVersion = "uncloud-client-version" + MetadataKeyMinDaemonVersion = "uncloud-min-server-version" + MetadataKeyDaemonVersion = "uncloud-server-version" + + // MinCLIVersion is the minimum client version the daemon accepts. The daemon + // rejects requests from older clients, forcing them to upgrade. This provides + // a clean cut-off for dropping support for old clients. + // + // MinDaemonVersion is the minimum daemon version the client requires. The client + // sends this with each request so the daemon can immediately reject if it's too old, + // avoiding the need for a preflight request. This is useful when a new client feature + // requires daemon capabilities that didn't exist in older versions. + // + // The two minimums are independent: a client might require a newer daemon for new + // features, while that same daemon could still handle requests from older clients. + MinCLIVersion = "0.0.0" + MinDaemonVersion = "0.0.0" + + ReleaseURL = "https://github.com/psviderski/uncloud/releases/latest" +) + +var ( + // currentVersion is the version of this binary (CLI or daemon) + currentVersion = semver.MustParse(internalVersion.String()) + // zeroVersion is used when no version is specified (treated as 0.0.0) + zeroVersion = semver.MustParse("0.0.0") + // Pre-parsed minimum versions for comparison + minCLIVersion = semver.MustParse(MinCLIVersion) + minDaemonVersion = semver.MustParse(MinDaemonVersion) + + // warned tracks if we've already printed the daemon version warning + // TODO: remove when checkDaemonVersionInResponse is no longer needed (see below) + warned bool +) + +func extractVersion(md metadata.MD, key string) *semver.Version { + if md == nil { + return zeroVersion + } + values := md.Get(key) + if len(values) == 0 || values[0] == "" { + return zeroVersion + } + sv, err := semver.NewVersion(values[0]) + if err != nil { + return zeroVersion + } + return sv +} + +func checkClientVersionHeaders(ctx context.Context) error { + md, _ := metadata.FromIncomingContext(ctx) + + actualCLIVersion := extractVersion(md, MetadataKeyCLIVersion) + if actualCLIVersion.LessThan(minCLIVersion) { + return status.Errorf(codes.FailedPrecondition, + "version check failed: client version is below minimum %s. Please upgrade: %s", + minCLIVersion, ReleaseURL) + } + + requiredMinDaemon := extractVersion(md, MetadataKeyMinDaemonVersion) + if currentVersion.LessThan(requiredMinDaemon) { + return status.Errorf(codes.FailedPrecondition, + "version check failed: daemon version %s is below client's minimum required version %s. Please upgrade the daemon: %s", + currentVersion, requiredMinDaemon, ReleaseURL) + } + + return nil +} + +func ServerUnaryInterceptor(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { + if err := checkClientVersionHeaders(ctx); err != nil { + return nil, err + } + if err := grpc.SetHeader(ctx, metadata.Pairs(MetadataKeyDaemonVersion, currentVersion.String())); err != nil { + return nil, err + } + return handler(ctx, req) +} + +func ServerStreamInterceptor(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error { + if err := checkClientVersionHeaders(ss.Context()); err != nil { + return err + } + if err := ss.SetHeader(metadata.Pairs(MetadataKeyDaemonVersion, currentVersion.String())); err != nil { + return err + } + return handler(srv, ss) +} + +func ClientUnaryInterceptor(ctx context.Context, method string, req, reply any, cc *grpc.ClientConn, invoker grpc.UnaryInvoker, opts ...grpc.CallOption) error { + ctx = metadata.AppendToOutgoingContext(ctx, + MetadataKeyCLIVersion, currentVersion.String(), + MetadataKeyMinDaemonVersion, MinDaemonVersion, + ) + + // TODO: remove when checkDaemonVersionInResponse is no longer needed, + // as we'll no longer need to extract headers from the response here. + var respMD metadata.MD + opts = append(opts, grpc.Header(&respMD)) + + err := invoker(ctx, method, req, reply, cc, opts...) + if err != nil { + return err + } + + // TODO: Remove eventually (see note on method below) + checkDaemonVersionInResponse(respMD) + + return nil +} + +// This is just needed as a warning during the transition to version checking +// releases. It warns the user when they just communicated with a daemon that did +// not check the version requirements. +// TODO: Remove this in some later release, after users have upgraded. +func checkDaemonVersionInResponse(md metadata.MD) { + daemonVersion := extractVersion(md, MetadataKeyDaemonVersion) + if daemonVersion.LessThan(minDaemonVersion) { + if warned { + return + } + warned = true + + msg := fmt.Sprintf("daemon version is below minimum required version %s. The daemon did not verify this CLI's minimum version requirement, so the operation may not have behaved as intended. Please upgrade the daemon: %s", + minDaemonVersion, ReleaseURL) + fmt.Fprintf(os.Stderr, "WARNING: %s\n", msg) + } +} + +func ClientStreamInterceptor(ctx context.Context, desc *grpc.StreamDesc, cc *grpc.ClientConn, method string, streamer grpc.Streamer, opts ...grpc.CallOption) (grpc.ClientStream, error) { + ctx = metadata.AppendToOutgoingContext(ctx, + MetadataKeyCLIVersion, currentVersion.String(), + MetadataKeyMinDaemonVersion, MinDaemonVersion, + ) + + stream, err := streamer(ctx, desc, cc, method, opts...) + if err != nil { + return nil, err + } + + // TODO: Wrapping the stream in versionedClientStream will no longer + // be necessary when we are ready to remove the temporary, transition + // safety check checkDaemonVersionInResponse (see note on method above) + return &versionedClientStream{ClientStream: stream}, nil +} + +// TODO: remove when checkDaemonVersionInResponse is no longer needed +type versionedClientStream struct { + grpc.ClientStream +} + +// TODO: remove when checkDaemonVersionInResponse is no longer needed +func (s *versionedClientStream) Header() (metadata.MD, error) { + md, err := s.ClientStream.Header() + if err != nil { + return nil, err + } + + checkDaemonVersionInResponse(md) + + return md, nil +} diff --git a/pkg/versioncheck/interceptor_test.go b/pkg/versioncheck/interceptor_test.go new file mode 100644 index 00000000..876a4015 --- /dev/null +++ b/pkg/versioncheck/interceptor_test.go @@ -0,0 +1,216 @@ +package versioncheck + +import ( + "bytes" + "context" + "io" + "os" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/metadata" + "google.golang.org/grpc/status" +) + +func TestExtractVersion(t *testing.T) { + tests := []struct { + name string + md metadata.MD + key string + expected string + }{ + { + name: "nil metadata", + md: nil, + key: MetadataKeyCLIVersion, + expected: "0.0.0", + }, + { + name: "missing key", + md: metadata.MD{}, + key: MetadataKeyCLIVersion, + expected: "0.0.0", + }, + { + name: "empty value", + md: metadata.Pairs(MetadataKeyCLIVersion, ""), + key: MetadataKeyCLIVersion, + expected: "0.0.0", + }, + { + name: "invalid version", + md: metadata.Pairs(MetadataKeyCLIVersion, "not-a-version"), + key: MetadataKeyCLIVersion, + expected: "0.0.0", + }, + { + name: "valid version", + md: metadata.Pairs(MetadataKeyCLIVersion, "1.2.3"), + key: MetadataKeyCLIVersion, + expected: "1.2.3", + }, + { + name: "version with prerelease", + md: metadata.Pairs(MetadataKeyCLIVersion, "0.0.0-dev"), + key: MetadataKeyCLIVersion, + expected: "0.0.0-dev", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := extractVersion(tt.md, tt.key) + assert.Equal(t, tt.expected, got.String()) + }) + } +} + +func TestCheckClientVersionHeaders(t *testing.T) { + tests := []struct { + name string + md metadata.MD + wantErr bool + errCode codes.Code + errContain string + }{ + { + name: "cli version below minimum", + md: metadata.Pairs( + MetadataKeyCLIVersion, "0.0.0-dev", + ), + wantErr: true, + errCode: codes.FailedPrecondition, + errContain: "client version is below minimum", + }, + { + name: "cli version above minimum", + md: metadata.Pairs( + MetadataKeyCLIVersion, "999.0.0", + ), + wantErr: false, + }, + { + name: "min daemon version above current daemon", + md: metadata.Pairs( + MetadataKeyCLIVersion, "999.0.0", + MetadataKeyMinDaemonVersion, "999.0.0", + ), + wantErr: true, + errCode: codes.FailedPrecondition, + errContain: "daemon version", + }, + { + name: "min daemon version below current daemon", + md: metadata.Pairs( + MetadataKeyCLIVersion, "999.0.0", + MetadataKeyMinDaemonVersion, "0.0.1", + ), + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := context.Background() + if tt.md != nil { + ctx = metadata.NewIncomingContext(ctx, tt.md) + } + + err := checkClientVersionHeaders(ctx) + + if tt.wantErr { + require.Error(t, err) + st, ok := status.FromError(err) + require.True(t, ok, "expected gRPC status error, got %T", err) + assert.Equal(t, tt.errCode, st.Code()) + assert.True(t, strings.Contains(st.Message(), tt.errContain), + "error message = %q, want to contain %q", st.Message(), tt.errContain) + } else { + assert.NoError(t, err) + } + }) + } +} + +func TestCheckDaemonVersionInResponse(t *testing.T) { + tests := []struct { + name string + md metadata.MD + wantWarning bool + }{ + { + name: "daemon version below minimum", + md: metadata.Pairs(MetadataKeyDaemonVersion, "0.0.0-dev"), + wantWarning: true, + }, + { + name: "daemon version above minimum", + md: metadata.Pairs(MetadataKeyDaemonVersion, "999.0.0"), + wantWarning: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Reset warned flag for each test + warned = false + + // Capture stderr + old := os.Stderr + r, w, _ := os.Pipe() + os.Stderr = w + + checkDaemonVersionInResponse(tt.md) + + w.Close() + var buf bytes.Buffer + io.Copy(&buf, r) + os.Stderr = old + + output := buf.String() + + if tt.wantWarning { + assert.True(t, strings.Contains(output, "WARNING"), "expected warning output, got none") + } else { + assert.Equal(t, "", output) + } + }) + } +} + +func TestCheckDaemonVersionInResponse_WarnOnce(t *testing.T) { + // Reset warned flag + warned = false + + md := metadata.Pairs(MetadataKeyDaemonVersion, "0.0.0-dev") + + // First call - should warn + old := os.Stderr + r, w, _ := os.Pipe() + os.Stderr = w + + checkDaemonVersionInResponse(md) + + w.Close() + var buf bytes.Buffer + io.Copy(&buf, r) + os.Stderr = old + + assert.True(t, strings.Contains(buf.String(), "WARNING"), "first call should warn") + + // Second call - should NOT warn (warned flag is now true) + r2, w2, _ := os.Pipe() + os.Stderr = w2 + + checkDaemonVersionInResponse(md) + + w2.Close() + var buf2 bytes.Buffer + io.Copy(&buf2, r2) + os.Stderr = old + + assert.Equal(t, "", buf2.String(), "second call should not warn") +}