diff options
Diffstat (limited to 'internal/protocol/handshake_test.go')
| -rw-r--r-- | internal/protocol/handshake_test.go | 190 | +129 −61 |
1 files changed, 129 insertions, 61 deletions
diff --git a/internal/protocol/handshake_test.go b/internal/protocol/handshake_test.go index 84eb5b4..14f254a 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 @@ -8,6 +8,7 @@ import ( "crypto/rand" "crypto/x509" "encoding/pem" + "errors" "os" "path/filepath" "testing" @@ -58,6 +59,19 @@ func TestLoadPrivateKey(t *testing.T) { } } +func TestLoadPublicKey(t *testing.T) { + _, pubPath, wantPub := writeTestKeyPair(t) + + key, err := LoadPublicKey(pubPath) + if err != nil { + t.Fatal(err) + } + + if !key.Equal(wantPub) { + t.Fatal("public key mismatch") + } +} + func TestLoadPrivateKey_NotFound(t *testing.T) { _, err := LoadPrivateKey("/nonexistent/path") if err == nil { @@ -67,9 +81,7 @@ func TestLoadPrivateKey_NotFound(t *testing.T) { 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) - } + os.WriteFile(path, []byte("not pem"), 0o600) _, err := LoadPrivateKey(path) if err == nil { @@ -78,11 +90,10 @@ func TestLoadPrivateKey_NotPEM(t *testing.T) { } func TestLoadPrivateKey_WrongKeyType(t *testing.T) { + // Write a PEM block with garbage DER 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) - } + os.WriteFile(path, data, 0o600) _, err := LoadPrivateKey(path) if err == nil { @@ -90,94 +101,151 @@ func TestLoadPrivateKey_WrongKeyType(t *testing.T) { } } -func TestSelfSignedCert(t *testing.T) { +func generateTestKey(t *testing.T) ed25519.PrivateKey { + t.Helper() _, priv, err := ed25519.GenerateKey(rand.Reader) if err != nil { t.Fatal(err) } + return priv +} - meta := &WorkerMeta{ - Name: "test-worker", - Version: "1.2.3", - Os: "linux", - Arch: "amd64", - Runtime: "host", +func TestHandshake_Success(t *testing.T) { + priv := generateTestKey(t) + pub := priv.Public().(ed25519.PublicKey) + + 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) } - cert, err := SelfSignedCert(priv, meta) + serverNonce, err := server.Challenge(pub, nil) if err != nil { t.Fatal(err) } - if len(cert.Certificate) != 1 { - t.Fatalf("expected 1 cert, got %d", len(cert.Certificate)) - } - - parsed, err := x509.ParseCertificate(cert.Certificate[0]) + signature, err := client.Sign(serverNonce, nil) if err != nil { t.Fatal(err) } - // 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") + if err := server.Verify(signature, time.Now()); err != nil { + t.Fatal(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_ChannelBinding(t *testing.T) { + priv := generateTestKey(t) + pub := priv.Public().(ed25519.PublicKey) + ekm := make([]byte, 32) + rand.Read(ekm) - // 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) + serverNonce, err := server.Challenge(pub, ekm) if err != nil { t.Fatal(err) } - want := &WorkerMeta{ - Name: "my-worker", - Version: "0.5.1", - Os: "darwin", - Arch: "arm64", - Runtime: "docker", - } - - cert, err := SelfSignedCert(priv, want) + signature, err := client.Sign(serverNonce, ekm) if err != nil { t.Fatal(err) } - parsed, err := x509.ParseCertificate(cert.Certificate[0]) - if err != nil { + if err := server.Verify(signature, time.Now()); err != nil { t.Fatal(err) } +} + +func TestHandshake_ChannelBinding_MismatchedEKM(t *testing.T) { + priv := generateTestKey(t) + pub := priv.Public().(ed25519.PublicKey) - got := ParseWorkerMeta(parsed) - if got == nil { - t.Fatal("ParseWorkerMeta returned nil") + serverEKM := make([]byte, 32) + workerEKM := make([]byte, 32) + rand.Read(serverEKM) + rand.Read(workerEKM) + + server := NewServerHandshake() + client := NewClientHandshake(priv) + + serverNonce, _ := server.Challenge(pub, serverEKM) + signature, _ := client.Sign(serverNonce, workerEKM) + + err := server.Verify(signature, time.Now()) + if !errors.Is(err, ErrInvalidSignature) { + t.Fatalf("err = %v, want ErrInvalidSignature", err) } +} + +func TestHandshake_WrongKey(t *testing.T) { + workerKey := generateTestKey(t) + otherKey := generateTestKey(t) + + server := NewServerHandshake() + client := NewClientHandshake(workerKey) + + serverNonce, _ := server.Challenge(otherKey.Public().(ed25519.PublicKey), nil) + signature, _ := client.Sign(serverNonce, nil) + + err := server.Verify(signature, time.Now()) + if !errors.Is(err, ErrInvalidSignature) { + t.Fatalf("err = %v, want ErrInvalidSignature", err) + } +} + +func TestHandshake_InvalidPublicKeyLength(t *testing.T) { + server := NewServerHandshake() - if *got != *want { - t.Fatalf("meta mismatch:\n got: %+v\nwant: %+v", got, want) + _, err := server.Challenge([]byte("short"), nil) + if !errors.Is(err, ErrInvalidPublicKey) { + t.Fatalf("err = %v, want ErrInvalidPublicKey", err) } } -func TestParseWorkerMeta_NoCert(t *testing.T) { - cert := &x509.Certificate{} - if meta := ParseWorkerMeta(cert); meta != nil { - t.Fatalf("expected nil, got %+v", meta) +func TestHandshake_TamperedSignature(t *testing.T) { + priv := generateTestKey(t) + pub := priv.Public().(ed25519.PublicKey) + + server := NewServerHandshake() + client := NewClientHandshake(priv) + + serverNonce, _ := server.Challenge(pub, nil) + signature, _ := client.Sign(serverNonce, nil) + + signature[0] ^= 0xff + + err := server.Verify(signature, time.Now()) + if !errors.Is(err, ErrInvalidSignature) { + t.Fatalf("err = %v, want ErrInvalidSignature", err) + } +} + +func TestHandshake_ClockSkew(t *testing.T) { + priv := generateTestKey(t) + pub := priv.Public().(ed25519.PublicKey) + + server := NewServerHandshake() + client := NewClientHandshake(priv) + + serverNonce, _ := server.Challenge(pub, nil) + signature, _ := client.Sign(serverNonce, nil) + + err := server.Verify(signature, time.Now().Add(-2*time.Minute)) + if !errors.Is(err, ErrClockSkew) { + t.Fatalf("err = %v, want ErrClockSkew", err) + } +} + +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) } } |
