aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorNikolay Govorov <me@govorov.online>2026-04-01 21:09:52 +0100
committerNikolay Govorov <me@govorov.online>2026-04-01 21:09:52 +0100
commitf0caa4294e44581a25bdf4ddaadb5d2e0cabf2ad (patch)
tree8edf0bce7fa9619e3f1e7d02c2d12526ed7522f2
parentde6dc17b4ec92dfdf66b992e429e186f3f2c279e (diff)
downloadtar
tar.gz
tar.bz2
tar.lz
tar.xz
tar.zst
zip
Move grpc handshake logic to protocol
Diffstat
-rw-r--r--cmd/mirumd/main.go36+16 −20
-rw-r--r--cmd/mirumw/client.go21+14 −7
-rw-r--r--internal/protocol/handshake.go106+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.go31+0 −31
-rw-r--r--proto/mirum.proto4+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 {