aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
Diffstat (limited to 'internal/protocol/handshake_test.go')
-rw-r--r--internal/protocol/handshake_test.go214+96 −118
1 files changed, 96 insertions, 118 deletions
diff --git a/internal/protocol/handshake_test.go b/internal/protocol/handshake_test.go
index 84eb5b4..d5e336b 100644
--- a/internal/protocol/handshake_test.go
+++ b/internal/protocol/handshake_test.go
@@ -1,4 +1,4 @@
-// SPDX-FileCopyrightText: 2026 Nikolay Govorov
+// Copyright (c) 2026 Nikolay Govorov
// SPDX-License-Identifier: AGPL-3.0-or-later
package protocol
@@ -6,178 +6,156 @@ package protocol
import (
"crypto/ed25519"
"crypto/rand"
- "crypto/x509"
- "encoding/pem"
- "os"
- "path/filepath"
+ "errors"
"testing"
"time"
)
-func writeTestKeyPair(t *testing.T) (privPath, pubPath string, pub ed25519.PublicKey) {
+func generateTestKey(t *testing.T) ed25519.PrivateKey {
t.Helper()
- pub, priv, err := ed25519.GenerateKey(rand.Reader)
+ _, priv, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
t.Fatal(err)
}
+ return priv
+}
- dir := t.TempDir()
+func TestHandshake_Success(t *testing.T) {
+ priv := generateTestKey(t)
+ pub := priv.Public().(ed25519.PublicKey)
- privDER, err := x509.MarshalPKCS8PrivateKey(priv)
- if err != nil {
- t.Fatal(err)
+ server := NewServerHandshake()
+ client := NewClientHandshake(priv)
+
+ if got := client.PublicKey(); len(got) != ed25519.PublicKeySize {
+ t.Fatalf("public key len = %d, want %d", len(got), ed25519.PublicKeySize)
}
- privPath = filepath.Join(dir, "test.key")
- if err := os.WriteFile(privPath, pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: privDER}), 0o600); err != nil {
+
+ serverNonce, err := server.Challenge(pub, nil)
+ if err != nil {
t.Fatal(err)
}
- pubDER, err := x509.MarshalPKIXPublicKey(pub)
+ signature, err := client.Sign(serverNonce, nil)
if err != nil {
t.Fatal(err)
}
- pubPath = filepath.Join(dir, "test.pub")
- if err := os.WriteFile(pubPath, pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: pubDER}), 0o644); err != nil {
+
+ if err := server.Verify(signature, time.Now()); err != nil {
t.Fatal(err)
}
-
- return privPath, pubPath, pub
}
-func TestLoadPrivateKey(t *testing.T) {
- privPath, _, wantPub := writeTestKeyPair(t)
+func TestHandshake_ChannelBinding(t *testing.T) {
+ priv := generateTestKey(t)
+ pub := priv.Public().(ed25519.PublicKey)
+ ekm := make([]byte, 32)
+ rand.Read(ekm)
- key, err := LoadPrivateKey(privPath)
+ server := NewServerHandshake()
+ client := NewClientHandshake(priv)
+
+ serverNonce, err := server.Challenge(pub, ekm)
if err != nil {
t.Fatal(err)
}
- gotPub := key.Public().(ed25519.PublicKey)
- if !gotPub.Equal(wantPub) {
- t.Fatal("public key mismatch")
+ signature, err := client.Sign(serverNonce, ekm)
+ if err != nil {
+ t.Fatal(err)
}
-}
-func TestLoadPrivateKey_NotFound(t *testing.T) {
- _, err := LoadPrivateKey("/nonexistent/path")
- if err == nil {
- t.Fatal("expected error")
+ if err := server.Verify(signature, time.Now()); err != nil {
+ t.Fatal(err)
}
}
-func TestLoadPrivateKey_NotPEM(t *testing.T) {
- path := filepath.Join(t.TempDir(), "bad.key")
- if err := os.WriteFile(path, []byte("not pem"), 0o600); err != nil {
- t.Fatal(err)
- }
+func TestHandshake_ChannelBinding_MismatchedEKM(t *testing.T) {
+ priv := generateTestKey(t)
+ pub := priv.Public().(ed25519.PublicKey)
- _, err := LoadPrivateKey(path)
- if err == nil {
- t.Fatal("expected error")
- }
-}
+ serverEKM := make([]byte, 32)
+ workerEKM := make([]byte, 32)
+ rand.Read(serverEKM)
+ rand.Read(workerEKM)
-func TestLoadPrivateKey_WrongKeyType(t *testing.T) {
- data := pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: []byte("garbage")})
- path := filepath.Join(t.TempDir(), "bad.key")
- if err := os.WriteFile(path, data, 0o600); err != nil {
- t.Fatal(err)
- }
+ server := NewServerHandshake()
+ client := NewClientHandshake(priv)
+
+ serverNonce, _ := server.Challenge(pub, serverEKM)
+ signature, _ := client.Sign(serverNonce, workerEKM)
- _, err := LoadPrivateKey(path)
- if err == nil {
- t.Fatal("expected error")
+ err := server.Verify(signature, time.Now())
+ if !errors.Is(err, ErrInvalidSignature) {
+ t.Fatalf("err = %v, want ErrInvalidSignature", err)
}
}
-func TestSelfSignedCert(t *testing.T) {
- _, priv, err := ed25519.GenerateKey(rand.Reader)
- if err != nil {
- t.Fatal(err)
- }
+func TestHandshake_WrongKey(t *testing.T) {
+ workerKey := generateTestKey(t)
+ otherKey := generateTestKey(t)
- meta := &WorkerMeta{
- Name: "test-worker",
- Version: "1.2.3",
- Os: "linux",
- Arch: "amd64",
- Runtime: "host",
- }
+ server := NewServerHandshake()
+ client := NewClientHandshake(workerKey)
- cert, err := SelfSignedCert(priv, meta)
- if err != nil {
- t.Fatal(err)
- }
+ serverNonce, _ := server.Challenge(otherKey.Public().(ed25519.PublicKey), nil)
+ signature, _ := client.Sign(serverNonce, nil)
- if len(cert.Certificate) != 1 {
- t.Fatalf("expected 1 cert, got %d", len(cert.Certificate))
+ err := server.Verify(signature, time.Now())
+ if !errors.Is(err, ErrInvalidSignature) {
+ t.Fatalf("err = %v, want ErrInvalidSignature", err)
}
+}
- parsed, err := x509.ParseCertificate(cert.Certificate[0])
- if err != nil {
- t.Fatal(err)
- }
+func TestHandshake_InvalidPublicKeyLength(t *testing.T) {
+ server := NewServerHandshake()
- // Public key matches
- pubKey, ok := parsed.PublicKey.(ed25519.PublicKey)
- if !ok {
- t.Fatal("certificate does not contain an ed25519 public key")
- }
- if !pubKey.Equal(priv.Public().(ed25519.PublicKey)) {
- t.Fatal("public key mismatch")
+ _, err := server.Challenge([]byte("short"), nil)
+ if !errors.Is(err, ErrInvalidPublicKey) {
+ t.Fatalf("err = %v, want ErrInvalidPublicKey", err)
}
+}
- // Key usage
- if parsed.KeyUsage&x509.KeyUsageDigitalSignature == 0 {
- t.Fatal("missing DigitalSignature key usage")
- }
- if len(parsed.ExtKeyUsage) != 1 || parsed.ExtKeyUsage[0] != x509.ExtKeyUsageClientAuth {
- t.Fatal("missing ClientAuth extended key usage")
- }
+func TestHandshake_TamperedSignature(t *testing.T) {
+ priv := generateTestKey(t)
+ pub := priv.Public().(ed25519.PublicKey)
- // NotBefore is recent (used for clock skew)
- if time.Since(parsed.NotBefore).Abs() > 5*time.Second {
- t.Fatalf("NotBefore too far from now: %v", parsed.NotBefore)
- }
-}
+ server := NewServerHandshake()
+ client := NewClientHandshake(priv)
-func TestSelfSignedCert_WorkerMeta(t *testing.T) {
- _, priv, err := ed25519.GenerateKey(rand.Reader)
- if err != nil {
- t.Fatal(err)
- }
+ serverNonce, _ := server.Challenge(pub, nil)
+ signature, _ := client.Sign(serverNonce, nil)
- want := &WorkerMeta{
- Name: "my-worker",
- Version: "0.5.1",
- Os: "darwin",
- Arch: "arm64",
- Runtime: "docker",
- }
+ signature[0] ^= 0xff
- cert, err := SelfSignedCert(priv, want)
- if err != nil {
- t.Fatal(err)
+ err := server.Verify(signature, time.Now())
+ if !errors.Is(err, ErrInvalidSignature) {
+ t.Fatalf("err = %v, want ErrInvalidSignature", err)
}
+}
- parsed, err := x509.ParseCertificate(cert.Certificate[0])
- if err != nil {
- t.Fatal(err)
- }
+func TestHandshake_ClockSkew(t *testing.T) {
+ priv := generateTestKey(t)
+ pub := priv.Public().(ed25519.PublicKey)
- got := ParseWorkerMeta(parsed)
- if got == nil {
- t.Fatal("ParseWorkerMeta returned nil")
- }
+ server := NewServerHandshake()
+ client := NewClientHandshake(priv)
+
+ serverNonce, _ := server.Challenge(pub, nil)
+ signature, _ := client.Sign(serverNonce, nil)
- if *got != *want {
- t.Fatalf("meta mismatch:\n got: %+v\nwant: %+v", got, want)
+ err := server.Verify(signature, time.Now().Add(-2*time.Minute))
+ if !errors.Is(err, ErrClockSkew) {
+ t.Fatalf("err = %v, want ErrClockSkew", err)
}
}
-func TestParseWorkerMeta_NoCert(t *testing.T) {
- cert := &x509.Certificate{}
- if meta := ParseWorkerMeta(cert); meta != nil {
- t.Fatalf("expected nil, got %+v", meta)
+func TestSign_InvalidNonceLength(t *testing.T) {
+ priv := generateTestKey(t)
+ client := NewClientHandshake(priv)
+
+ _, err := client.Sign([]byte("short"), nil)
+ if !errors.Is(err, ErrInvalidNonce) {
+ t.Fatalf("err = %v, want ErrInvalidNonce", err)
}
}