From f0caa4294e44581a25bdf4ddaadb5d2e0cabf2ad Mon Sep 17 00:00:00 2001 From: Nikolay Govorov Date: Wed, 1 Apr 2026 21:09:52 +0100 Subject: Move grpc handshake logic to protocol --- cmd/mirumd/main.go | 36 +++--- cmd/mirumw/client.go | 21 ++-- internal/protocol/handshake.go | 106 +++++++++++++++++ internal/protocol/handshake_test.go | 175 ++++++++++++++++++++++++++++ internal/protocol/hmac.go | 31 ----- internal/protocol/hmac_test.go | 87 -------------- proto/mirum.proto | 4 + 7 files changed, 315 insertions(+), 145 deletions(-) create mode 100644 internal/protocol/handshake.go create mode 100644 internal/protocol/handshake_test.go delete mode 100644 internal/protocol/hmac.go delete mode 100644 internal/protocol/hmac_test.go diff --git a/cmd/mirumd/main.go b/cmd/mirumd/main.go index 8bf098f..b53f567 100644 --- a/cmd/mirumd/main.go +++ b/cmd/mirumd/main.go @@ -213,6 +213,8 @@ func (s *mirumServer) Complete(ctx context.Context, result *pb.TaskResult) (*pb. } func (s *mirumServer) Handshake(stream pb.Mirum_HandshakeServer) error { + hs := protocol.NewServerHandshake(s.secret) + // Step 1: receive worker nonce in, err := stream.Recv() if err != nil { @@ -222,21 +224,17 @@ func (s *mirumServer) Handshake(stream pb.Mirum_HandshakeServer) error { if wc == nil { return fmt.Errorf("expected WorkerChallenge") } - workerNonce := wc.GetNonce() - if len(workerNonce) != protocol.NonceSize { - return fmt.Errorf("invalid nonce size: %d", len(workerNonce)) - } // Step 2: send server nonce + proof - serverNonce, err := protocol.GenerateNonce() + serverNonce, proof, err := hs.Challenge(wc.GetNonce()) if err != nil { - return fmt.Errorf("generate nonce: %w", err) + return err } if err := stream.Send(&pb.HandshakeOut{ Step: &pb.HandshakeOut_ServerChallenge{ ServerChallenge: &pb.ServerChallenge{ Nonce: serverNonce, - Proof: protocol.ComputeProof(s.secret, workerNonce, serverNonce), + Proof: proof, }, }, }); err != nil { @@ -253,21 +251,18 @@ func (s *mirumServer) Handshake(stream pb.Mirum_HandshakeServer) error { return fmt.Errorf("expected WorkerProof") } - if !protocol.VerifyProof(s.secret, serverNonce, workerNonce, wp.GetProof()) { - return s.reject(stream, "invalid secret") - } - - // Check clock skew wt := wp.GetWorkerTime() if wt == nil { return s.reject(stream, "worker_time is required") } - skew := time.Since(wt.AsTime()).Abs() - if skew > time.Minute { - return s.reject(stream, fmt.Sprintf("clock skew too large: %s", skew.Truncate(time.Second))) + + if err := hs.Verify(wp.GetProof(), wt.AsTime()); err != nil { + return s.reject(stream, err.Error()) } - if skew > 10*time.Second { - slog.Warn("clock skew", "worker", wp.GetName(), "skew", skew.Truncate(time.Second)) + + var warnings []string + if skew := time.Since(wt.AsTime()).Abs(); skew > 10*time.Second { + warnings = append(warnings, fmt.Sprintf("clock drift: %s", skew.Truncate(time.Second))) } slog.Info("worker connected", @@ -282,23 +277,24 @@ func (s *mirumServer) Handshake(stream pb.Mirum_HandshakeServer) error { if p, ok := peer.FromContext(stream.Context()); ok { s.authedPeers.Store(p.Addr.String(), true) } - return s.sendResult(stream, nil) + return s.sendResult(stream, nil, warnings) } -func (s *mirumServer) sendResult(stream pb.Mirum_HandshakeServer, errMsg *string) error { +func (s *mirumServer) sendResult(stream pb.Mirum_HandshakeServer, errMsg *string, warnings []string) error { return stream.Send(&pb.HandshakeOut{ Step: &pb.HandshakeOut_ServerResult{ ServerResult: &pb.ServerResult{ Error: errMsg, ServerVersion: protocol.VersionProto(), ServerTime: timestamppb.Now(), + Warnings: warnings, }, }, }) } func (s *mirumServer) reject(stream pb.Mirum_HandshakeServer, reason string) error { - if err := s.sendResult(stream, &reason); err != nil { + if err := s.sendResult(stream, &reason, nil); err != nil { return err } return fmt.Errorf("%s", reason) diff --git a/cmd/mirumw/client.go b/cmd/mirumw/client.go index b3be01f..3f7c0f0 100644 --- a/cmd/mirumw/client.go +++ b/cmd/mirumw/client.go @@ -103,8 +103,10 @@ func (c *client) handshake(ctx context.Context) error { return fmt.Errorf("open stream: %w", err) } + hs := protocol.NewClientHandshake([]byte(c.cfg.Secret)) + // Step 1: send worker nonce - workerNonce, err := protocol.GenerateNonce() + workerNonce, err := hs.Challenge() if err != nil { return fmt.Errorf("generate nonce: %w", err) } @@ -116,18 +118,19 @@ func (c *client) handshake(ctx context.Context) error { return fmt.Errorf("send challenge: %w", err) } - // Step 2: receive server challenge, verify server + // Step 2: receive server challenge, verify and compute proof out, err := stream.Recv() if err != nil { return fmt.Errorf("recv server challenge: %w", err) } - sc := out.GetServerChallenge() if sc == nil { return fmt.Errorf("expected ServerChallenge") } - if !protocol.VerifyProof([]byte(c.cfg.Secret), workerNonce, sc.GetNonce(), sc.GetProof()) { - return fmt.Errorf("server proof verification failed") + + workerProof, err := hs.Verify(sc.GetNonce(), sc.GetProof()) + if err != nil { + return err } // Step 3: send worker proof + metadata @@ -138,7 +141,7 @@ func (c *client) handshake(ctx context.Context) error { if err := stream.Send(&pb.HandshakeIn{ Step: &pb.HandshakeIn_WorkerProof{ WorkerProof: &pb.WorkerProof{ - Proof: protocol.ComputeProof([]byte(c.cfg.Secret), sc.GetNonce(), workerNonce), + Proof: workerProof, Name: name, Version: protocol.VersionProto(), Os: protocol.DetectOs(), @@ -161,7 +164,11 @@ func (c *client) handshake(ctx context.Context) error { return fmt.Errorf("expected ServerResult") } if sr.Error != nil { - return fmt.Errorf("server rejected: %s", *sr.Error) + return fmt.Errorf("%w: %s", protocol.ErrServerRejected, *sr.Error) + } + + for _, w := range sr.GetWarnings() { + slog.Warn("server warning", "msg", w) } v := sr.GetServerVersion() diff --git a/internal/protocol/handshake.go b/internal/protocol/handshake.go new file mode 100644 index 0000000..e14ac0b --- /dev/null +++ b/internal/protocol/handshake.go @@ -0,0 +1,106 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package protocol + +import ( + "crypto/hmac" + "crypto/rand" + "crypto/sha256" + "errors" + "time" +) + +var ( + ErrInvalidNonce = errors.New("invalid nonce size") + ErrInvalidProof = errors.New("invalid proof") + ErrClockSkew = errors.New("clock skew too large") + ErrServerRejected = errors.New("server rejected handshake") +) + +const NonceSize = 32 + +func generateNonce() ([]byte, error) { + nonce := make([]byte, NonceSize) + _, err := rand.Read(nonce) + return nonce, err +} + +// computeProof returns HMAC-SHA256(secret, first || second). +func computeProof(secret, first, second []byte) []byte { + mac := hmac.New(sha256.New, secret) + mac.Write(first) + mac.Write(second) + return mac.Sum(nil) +} + +func verifyProof(secret, first, second, proof []byte) bool { + expected := computeProof(secret, first, second) + return hmac.Equal(expected, proof) +} + +// ServerHandshake holds state for the server side of the handshake protocol. +type ServerHandshake struct { + secret []byte + workerNonce []byte + serverNonce []byte +} + +func NewServerHandshake(secret []byte) *ServerHandshake { + return &ServerHandshake{secret: secret} +} + +// Challenge validates the worker nonce and returns the server nonce + proof. +func (h *ServerHandshake) Challenge(workerNonce []byte) (serverNonce, proof []byte, err error) { + if len(workerNonce) != NonceSize { + return nil, nil, ErrInvalidNonce + } + + h.workerNonce = workerNonce + h.serverNonce, err = generateNonce() + if err != nil { + return nil, nil, err + } + + proof = computeProof(h.secret, workerNonce, h.serverNonce) + return h.serverNonce, proof, nil +} + +// Verify checks the worker's proof and clock skew. +func (h *ServerHandshake) Verify(proof []byte, workerTime time.Time) error { + if !verifyProof(h.secret, h.serverNonce, h.workerNonce, proof) { + return ErrInvalidProof + } + + skew := time.Since(workerTime).Abs() + if skew > time.Minute { + return ErrClockSkew + } + + return nil +} + +// ClientHandshake holds state for the client side of the handshake protocol. +type ClientHandshake struct { + secret []byte + workerNonce []byte +} + +func NewClientHandshake(secret []byte) *ClientHandshake { + return &ClientHandshake{secret: secret} +} + +// Challenge generates the worker nonce. +func (h *ClientHandshake) Challenge() (workerNonce []byte, err error) { + h.workerNonce, err = generateNonce() + return h.workerNonce, err +} + +// Verify checks the server proof and returns the worker proof. +func (h *ClientHandshake) Verify(serverNonce, serverProof []byte) (proof []byte, err error) { + if !verifyProof(h.secret, h.workerNonce, serverNonce, serverProof) { + return nil, ErrInvalidProof + } + proof = computeProof(h.secret, serverNonce, h.workerNonce) + return proof, nil +} diff --git a/internal/protocol/handshake_test.go b/internal/protocol/handshake_test.go new file mode 100644 index 0000000..702eebe --- /dev/null +++ b/internal/protocol/handshake_test.go @@ -0,0 +1,175 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package protocol + +import ( + "bytes" + "crypto/hmac" + "crypto/sha256" + "errors" + "testing" + "time" +) + +func TestGenerateNonce(t *testing.T) { + nonce, err := generateNonce() + if err != nil { + t.Fatal(err) + } + if len(nonce) != NonceSize { + t.Fatalf("len = %d, want %d", len(nonce), NonceSize) + } + + nonce2, _ := generateNonce() + if bytes.Equal(nonce, nonce2) { + t.Fatal("two nonces are identical") + } +} + +func TestComputeProof(t *testing.T) { + secret := []byte("secret") + first := []byte("first") + second := []byte("second") + + proof := computeProof(secret, first, second) + + // Golden value: HMAC-SHA256("secret", "first" || "second") + mac := hmac.New(sha256.New, secret) + mac.Write(first) + mac.Write(second) + want := mac.Sum(nil) + + if !bytes.Equal(proof, want) { + t.Fatalf("proof mismatch") + } + + // Different secret → different proof + other := computeProof([]byte("other"), first, second) + if bytes.Equal(proof, other) { + t.Fatal("different secrets produced same proof") + } + + // Order matters: (first, second) != (second, first) + reversed := computeProof(secret, second, first) + if bytes.Equal(proof, reversed) { + t.Fatal("argument order did not affect proof") + } +} + +func TestVerifyProof(t *testing.T) { + secret := []byte("secret") + first := []byte("first") + second := []byte("second") + proof := computeProof(secret, first, second) + + tests := []struct { + name string + secret []byte + first []byte + second []byte + proof []byte + want bool + }{ + {"valid", secret, first, second, proof, true}, + {"wrong proof", secret, first, second, []byte("wrong"), false}, + {"wrong secret", []byte("wrong"), first, second, proof, false}, + {"swapped args", secret, second, first, proof, false}, + {"empty proof", secret, first, second, nil, false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := verifyProof(tt.secret, tt.first, tt.second, tt.proof) + if got != tt.want { + t.Fatalf("VerifyProof = %v, want %v", got, tt.want) + } + }) + } +} + +// Full handshake: client and server with matching secrets. +func TestHandshake_Success(t *testing.T) { + secret := []byte("shared-secret") + server := NewServerHandshake(secret) + client := NewClientHandshake(secret) + + // Step 1: client generates nonce + workerNonce, err := client.Challenge() + if err != nil { + t.Fatal(err) + } + + // Step 2: server responds with its nonce + proof + serverNonce, serverProof, err := server.Challenge(workerNonce) + if err != nil { + t.Fatal(err) + } + + // Step 3: client verifies server, produces worker proof + workerProof, err := client.Verify(serverNonce, serverProof) + if err != nil { + t.Fatal(err) + } + + // Step 4: server verifies worker + if err := server.Verify(workerProof, time.Now()); err != nil { + t.Fatal(err) + } +} + +func TestHandshake_WrongSecret(t *testing.T) { + server := NewServerHandshake([]byte("server-secret")) + client := NewClientHandshake([]byte("wrong-secret")) + + workerNonce, _ := client.Challenge() + _, serverProof, _ := server.Challenge(workerNonce) + + // Client cannot verify server proof + _, err := client.Verify(server.serverNonce, serverProof) + if !errors.Is(err, ErrInvalidProof) { + t.Fatalf("err = %v, want ErrInvalidProof", err) + } +} + +func TestHandshake_BadNonce(t *testing.T) { + server := NewServerHandshake([]byte("secret")) + + _, _, err := server.Challenge([]byte("short")) + if !errors.Is(err, ErrInvalidNonce) { + t.Fatalf("err = %v, want ErrInvalidNonce", err) + } +} + +func TestHandshake_ClockSkew(t *testing.T) { + secret := []byte("secret") + server := NewServerHandshake(secret) + client := NewClientHandshake(secret) + + workerNonce, _ := client.Challenge() + serverNonce, serverProof, _ := server.Challenge(workerNonce) + workerProof, _ := client.Verify(serverNonce, serverProof) + + err := server.Verify(workerProof, time.Now().Add(-2*time.Minute)) + if !errors.Is(err, ErrClockSkew) { + t.Fatalf("err = %v, want ErrClockSkew", err) + } +} + +func TestHandshake_TamperedProof(t *testing.T) { + secret := []byte("secret") + server := NewServerHandshake(secret) + client := NewClientHandshake(secret) + + workerNonce, _ := client.Challenge() + serverNonce, serverProof, _ := server.Challenge(workerNonce) + workerProof, _ := client.Verify(serverNonce, serverProof) + + // Flip a byte + workerProof[0] ^= 0xff + + err := server.Verify(workerProof, time.Now()) + if !errors.Is(err, ErrInvalidProof) { + t.Fatalf("err = %v, want ErrInvalidProof", err) + } +} diff --git a/internal/protocol/hmac.go b/internal/protocol/hmac.go deleted file mode 100644 index 1b3db3e..0000000 --- a/internal/protocol/hmac.go +++ /dev/null @@ -1,31 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package protocol - -import ( - "crypto/hmac" - "crypto/rand" - "crypto/sha256" -) - -const NonceSize = 32 - -func GenerateNonce() ([]byte, error) { - nonce := make([]byte, NonceSize) - _, err := rand.Read(nonce) - return nonce, err -} - -// ComputeProof returns HMAC-SHA256(secret, first || second). -func ComputeProof(secret, first, second []byte) []byte { - mac := hmac.New(sha256.New, secret) - mac.Write(first) - mac.Write(second) - return mac.Sum(nil) -} - -func VerifyProof(secret, first, second, proof []byte) bool { - expected := ComputeProof(secret, first, second) - return hmac.Equal(expected, proof) -} diff --git a/internal/protocol/hmac_test.go b/internal/protocol/hmac_test.go deleted file mode 100644 index 49ffc79..0000000 --- a/internal/protocol/hmac_test.go +++ /dev/null @@ -1,87 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package protocol - -import ( - "bytes" - "crypto/hmac" - "crypto/sha256" - "testing" -) - -func TestGenerateNonce(t *testing.T) { - nonce, err := GenerateNonce() - if err != nil { - t.Fatal(err) - } - if len(nonce) != NonceSize { - t.Fatalf("len = %d, want %d", len(nonce), NonceSize) - } - - nonce2, _ := GenerateNonce() - if bytes.Equal(nonce, nonce2) { - t.Fatal("two nonces are identical") - } -} - -func TestComputeProof(t *testing.T) { - secret := []byte("secret") - first := []byte("first") - second := []byte("second") - - proof := ComputeProof(secret, first, second) - - // Golden value: HMAC-SHA256("secret", "first" || "second") - mac := hmac.New(sha256.New, secret) - mac.Write(first) - mac.Write(second) - want := mac.Sum(nil) - - if !bytes.Equal(proof, want) { - t.Fatalf("proof mismatch") - } - - // Different secret → different proof - other := ComputeProof([]byte("other"), first, second) - if bytes.Equal(proof, other) { - t.Fatal("different secrets produced same proof") - } - - // Order matters: (first, second) != (second, first) - reversed := ComputeProof(secret, second, first) - if bytes.Equal(proof, reversed) { - t.Fatal("argument order did not affect proof") - } -} - -func TestVerifyProof(t *testing.T) { - secret := []byte("secret") - first := []byte("first") - second := []byte("second") - proof := ComputeProof(secret, first, second) - - tests := []struct { - name string - secret []byte - first []byte - second []byte - proof []byte - want bool - }{ - {"valid", secret, first, second, proof, true}, - {"wrong proof", secret, first, second, []byte("wrong"), false}, - {"wrong secret", []byte("wrong"), first, second, proof, false}, - {"swapped args", secret, second, first, proof, false}, - {"empty proof", secret, first, second, nil, false}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := VerifyProof(tt.secret, tt.first, tt.second, tt.proof) - if got != tt.want { - t.Fatalf("VerifyProof = %v, want %v", got, tt.want) - } - }) - } -} diff --git a/proto/mirum.proto b/proto/mirum.proto index 80b0ea1..8f243fb 100644 --- a/proto/mirum.proto +++ b/proto/mirum.proto @@ -178,6 +178,10 @@ message ServerResult { Version server_version = 2; google.protobuf.Timestamp server_time = 3; + + // Non-fatal warnings for the worker to log (clock drift, upcoming + // deprecations, expiring secrets, known vulnerabilities, etc.). + repeated string warnings = 4; } message PollRequest { -- Gilti