diff options
Diffstat
| -rw-r--r-- | cmd/mirumd/main.go | 36 | +16 −20 |
| -rw-r--r-- | cmd/mirumw/client.go | 21 | +14 −7 |
| -rw-r--r-- | internal/protocol/handshake.go | 106 | +106 −0 |
| -rw-r--r-- | internal/protocol/handshake_test.go (renamed from internal/protocol/hmac_test.go) | 102 | +95 −7 |
| -rw-r--r-- | internal/protocol/hmac.go | 31 | +0 −31 |
| -rw-r--r-- | proto/mirum.proto | 4 | +4 −0 |
6 files changed, 235 insertions, 65 deletions
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 --- /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/hmac_test.go b/internal/protocol/handshake_test.go index 49ffc79..702eebe 100644 --- a/internal/protocol/hmac_test.go +++ b/internal/protocol/handshake_test.go @@ -7,11 +7,13 @@ import ( "bytes" "crypto/hmac" "crypto/sha256" + "errors" "testing" + "time" ) func TestGenerateNonce(t *testing.T) { - nonce, err := GenerateNonce() + nonce, err := generateNonce() if err != nil { t.Fatal(err) } @@ -19,7 +21,7 @@ func TestGenerateNonce(t *testing.T) { t.Fatalf("len = %d, want %d", len(nonce), NonceSize) } - nonce2, _ := GenerateNonce() + nonce2, _ := generateNonce() if bytes.Equal(nonce, nonce2) { t.Fatal("two nonces are identical") } @@ -30,7 +32,7 @@ func TestComputeProof(t *testing.T) { first := []byte("first") second := []byte("second") - proof := ComputeProof(secret, first, second) + proof := computeProof(secret, first, second) // Golden value: HMAC-SHA256("secret", "first" || "second") mac := hmac.New(sha256.New, secret) @@ -43,13 +45,13 @@ func TestComputeProof(t *testing.T) { } // Different secret → different proof - other := ComputeProof([]byte("other"), first, second) + 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) + reversed := computeProof(secret, second, first) if bytes.Equal(proof, reversed) { t.Fatal("argument order did not affect proof") } @@ -59,7 +61,7 @@ func TestVerifyProof(t *testing.T) { secret := []byte("secret") first := []byte("first") second := []byte("second") - proof := ComputeProof(secret, first, second) + proof := computeProof(secret, first, second) tests := []struct { name string @@ -78,10 +80,96 @@ func TestVerifyProof(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - got := VerifyProof(tt.secret, tt.first, tt.second, tt.proof) + 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 --- 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/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 { |
