1
0

serverlist_client_test.go 3.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108
  1. package pia
  2. import (
  3. "context"
  4. "crypto"
  5. "crypto/rand"
  6. "crypto/rsa"
  7. "crypto/sha256"
  8. "crypto/x509"
  9. "encoding/base64"
  10. "encoding/pem"
  11. "net/http"
  12. "net/http/httptest"
  13. "strings"
  14. "testing"
  15. )
  16. func TestCatalogClientReturnsExplicitlyVerifiedSnapshot(t *testing.T) {
  17. payload := []byte(`{"version":6,"groups":{"wg":[]},"regions":[]}`)
  18. privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
  19. if err != nil {
  20. t.Fatal(err)
  21. }
  22. digest := sha256.Sum256(payload)
  23. signature, err := rsa.SignPKCS1v15(rand.Reader, privateKey, crypto.SHA256, digest[:])
  24. if err != nil {
  25. t.Fatal(err)
  26. }
  27. publicDER, err := x509.MarshalPKIXPublicKey(&privateKey.PublicKey)
  28. if err != nil {
  29. t.Fatal(err)
  30. }
  31. signed := append(append(append([]byte{}, payload...), '\n'), []byte(base64.StdEncoding.EncodeToString(signature))...)
  32. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
  33. w.Header().Set("Content-Type", "application/octet-stream")
  34. _, _ = w.Write(signed)
  35. }))
  36. defer server.Close()
  37. client := NewCatalogClient(server.URL+"/v6", pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: publicDER}))
  38. snapshot, err := client.Fetch(context.Background())
  39. if err != nil {
  40. t.Fatal(err)
  41. }
  42. if !snapshot.SignatureVerified || snapshot.SchemaHint != "6" || string(snapshot.Payload) != string(payload) {
  43. t.Fatalf("unexpected verified snapshot: %+v", snapshot)
  44. }
  45. }
  46. func TestCatalogClientRejectsUnsafeResponses(t *testing.T) {
  47. tests := []struct {
  48. name, contentType, body string
  49. maxBody int64
  50. wantCode string
  51. }{
  52. {name: "html", contentType: "text/html", body: "<html>maintenance</html>", wantCode: CodeCatalogSchemaUnsupported},
  53. {name: "oversized", contentType: "application/octet-stream", body: strings.Repeat("x", 65), maxBody: 64, wantCode: CodeCatalogUnavailable},
  54. {name: "unsigned", contentType: "application/json", body: `{"version":6,"groups":{},"regions":[]}`, wantCode: CodeCatalogSignatureInvalid},
  55. }
  56. for _, test := range tests {
  57. t.Run(test.name, func(t *testing.T) {
  58. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
  59. w.Header().Set("Content-Type", test.contentType)
  60. _, _ = w.Write([]byte(test.body))
  61. }))
  62. defer server.Close()
  63. client := NewCatalogClient(server.URL+"/v6", []byte("invalid public key"))
  64. if test.maxBody > 0 {
  65. client.MaxBody = test.maxBody
  66. }
  67. _, err := client.Fetch(context.Background())
  68. if CodeOf(err) != test.wantCode {
  69. t.Fatalf("got %s, want %s: %v", CodeOf(err), test.wantCode, err)
  70. }
  71. })
  72. }
  73. payload := []byte(`{"version":6,"groups":{"wg":[]},"regions":[]}`)
  74. signingKey, err := rsa.GenerateKey(rand.Reader, 2048)
  75. if err != nil {
  76. t.Fatal(err)
  77. }
  78. digest := sha256.Sum256(payload)
  79. signature, err := rsa.SignPKCS1v15(rand.Reader, signingKey, crypto.SHA256, digest[:])
  80. if err != nil {
  81. t.Fatal(err)
  82. }
  83. signed := append(append(append([]byte{}, payload...), '\n'), []byte(base64.StdEncoding.EncodeToString(signature))...)
  84. wrongKey, err := rsa.GenerateKey(rand.Reader, 2048)
  85. if err != nil {
  86. t.Fatal(err)
  87. }
  88. wrongDER, err := x509.MarshalPKIXPublicKey(&wrongKey.PublicKey)
  89. if err != nil {
  90. t.Fatal(err)
  91. }
  92. t.Run("valid signature from unpinned key", func(t *testing.T) {
  93. server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
  94. w.Header().Set("Content-Type", "application/octet-stream")
  95. _, _ = w.Write(signed)
  96. }))
  97. defer server.Close()
  98. client := NewCatalogClient(server.URL+"/v6", pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: wrongDER}))
  99. _, err := client.Fetch(context.Background())
  100. if CodeOf(err) != CodeCatalogSignatureInvalid {
  101. t.Fatalf("got %s, want %s: %v", CodeOf(err), CodeCatalogSignatureInvalid, err)
  102. }
  103. })
  104. }