diff options
Diffstat (limited to 'internal/protocol/handshake_test.go')
| -rw-r--r-- | internal/protocol/handshake_test.go | 214 | +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) } } |
