register_test.go 9.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260
  1. package pia
  2. import (
  3. "context"
  4. "crypto/rand"
  5. "crypto/rsa"
  6. "crypto/tls"
  7. "crypto/x509"
  8. "crypto/x509/pkix"
  9. "encoding/base64"
  10. "encoding/pem"
  11. "math/big"
  12. "net"
  13. "net/http"
  14. "net/http/httptest"
  15. "net/netip"
  16. "os"
  17. "path/filepath"
  18. "strings"
  19. "sync/atomic"
  20. "testing"
  21. "time"
  22. )
  23. func testPubKey() string {
  24. raw := make([]byte, 32)
  25. raw[0] = 1
  26. return base64.StdEncoding.EncodeToString(raw)
  27. }
  28. func TestParseRegistrationFixture(t *testing.T) {
  29. raw, err := os.ReadFile(filepath.Join("testdata", "addkey", "success.json"))
  30. if err != nil {
  31. t.Fatal(err)
  32. }
  33. result, err := parseRegistration(raw)
  34. if err != nil {
  35. t.Fatal(err)
  36. }
  37. if result.PeerIP.String() != "10.42.0.2/32" || result.ServerPort != 51820 || result.ServerIP.String() != "198.51.100.10" || len(result.DNSServers) != 2 {
  38. t.Fatalf("unexpected registration result: %+v", result)
  39. }
  40. prefixed := []byte(`{"status":"OK","peer_ip":"10.42.0.3/32","server_key":"AgAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=","server_ip":"198.51.100.10","server_port":51820,"dns_servers":["10.0.0.242"]}`)
  41. prefixedResult, err := parseRegistration(prefixed)
  42. if err != nil || prefixedResult.PeerIP.String() != "10.42.0.3/32" {
  43. t.Fatalf("expected an explicit /32 peer address to remain supported: result=%+v err=%v", prefixedResult, err)
  44. }
  45. missing, err := os.ReadFile(filepath.Join("testdata", "addkey", "missing_port.json"))
  46. if err != nil {
  47. t.Fatal(err)
  48. }
  49. if _, err := parseRegistration(missing); err == nil || CodeOf(err) != CodeRegistrationInvalid {
  50. t.Fatalf("expected missing server port to be rejected: %v", err)
  51. }
  52. dnsRaw, err := os.ReadFile(filepath.Join("testdata", "addkey", "invalid_dns.json"))
  53. if err != nil {
  54. t.Fatal(err)
  55. }
  56. dnsResult, err := parseRegistration(dnsRaw)
  57. if err != nil || len(dnsResult.DNSServers) != 0 {
  58. t.Fatalf("invalid dns_servers must be ignored: result=%+v err=%v", dnsResult, err)
  59. }
  60. invalidFixtures := []struct{ file, wantCode string }{
  61. {"status_error.json", CodeRegistrationRejected},
  62. {"invalid_peer_ip.json", CodeRegistrationInvalid},
  63. {"invalid_peer_prefix.json", CodeRegistrationInvalid},
  64. {"invalid_server_key.json", CodeRegistrationInvalid},
  65. {"invalid_server_ip.json", CodeRegistrationInvalid},
  66. {"invalid_port.json", CodeRegistrationInvalid},
  67. }
  68. for _, test := range invalidFixtures {
  69. t.Run(test.file, func(t *testing.T) {
  70. raw, readErr := os.ReadFile(filepath.Join("testdata", "addkey", test.file))
  71. if readErr != nil {
  72. t.Fatal(readErr)
  73. }
  74. if _, parseErr := parseRegistration(raw); CodeOf(parseErr) != test.wantCode {
  75. t.Fatalf("expected %s, got %s: %v", test.wantCode, CodeOf(parseErr), parseErr)
  76. }
  77. })
  78. }
  79. }
  80. func TestRegistrationTLSHostnameAndCA(t *testing.T) {
  81. fixture, err := os.ReadFile(filepath.Join("testdata", "addkey", "success.json"))
  82. if err != nil {
  83. t.Fatal(err)
  84. }
  85. const token = "test-token-value-that-is-long-enough"
  86. server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  87. if r.URL.Query().Get("pt") != token || r.URL.Query().Get("pubkey") == "" {
  88. t.Errorf("registration query is missing required values")
  89. }
  90. w.Header().Set("Content-Type", "application/json")
  91. _, _ = w.Write(fixture)
  92. }))
  93. server.StartTLS()
  94. defer server.Close()
  95. certificate := server.Certificate()
  96. caPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certificate.Raw})
  97. host, portText, err := net.SplitHostPort(server.Listener.Addr().String())
  98. if err != nil {
  99. t.Fatal(err)
  100. }
  101. port, err := net.LookupPort("tcp", portText)
  102. if err != nil {
  103. t.Fatal(err)
  104. }
  105. client := NewRegistrationClient(caPEM)
  106. client.Port = uint16(port)
  107. key := testPubKey()
  108. t.Setenv("HTTPS_PROXY", "http://127.0.0.1:1")
  109. _, err = client.RegisterKey(context.Background(), WireGuardServer{Hostname: "example.com", IP: netip.MustParseAddr(host)}, token, key)
  110. if err != nil {
  111. t.Fatalf("expected TLS registration to succeed: %v", err)
  112. }
  113. _, err = client.RegisterKey(context.Background(), WireGuardServer{Hostname: "example.com", IP: netip.MustParseAddr(host)}, token, "")
  114. if CodeOf(err) != CodeInvalidInput {
  115. t.Fatalf("zero public key returned %s, want %s", CodeOf(err), CodeInvalidInput)
  116. }
  117. _, err = client.RegisterKey(context.Background(), WireGuardServer{Hostname: "wrong.example", IP: netip.MustParseAddr(host)}, token, key)
  118. if CodeOf(err) != CodeTLSValidation {
  119. t.Fatalf("wrong hostname returned %s, want %s: %v", CodeOf(err), CodeTLSValidation, err)
  120. }
  121. client = NewRegistrationClient([]byte("-----BEGIN CERTIFICATE-----\ninvalid\n-----END CERTIFICATE-----"))
  122. client.Port = uint16(port)
  123. _, err = client.RegisterKey(context.Background(), WireGuardServer{Hostname: "example.com", IP: netip.MustParseAddr(host)}, token, key)
  124. if CodeOf(err) != CodeTLSValidation {
  125. t.Fatalf("wrong CA returned %s, want %s", CodeOf(err), CodeTLSValidation)
  126. }
  127. expiredServer, expiredCA := newExpiredTLSServer(t, fixture)
  128. defer expiredServer.Close()
  129. expiredHost, expiredPortText, err := net.SplitHostPort(expiredServer.Listener.Addr().String())
  130. if err != nil {
  131. t.Fatal(err)
  132. }
  133. expiredPort, err := net.LookupPort("tcp", expiredPortText)
  134. if err != nil {
  135. t.Fatal(err)
  136. }
  137. client = NewRegistrationClient(expiredCA)
  138. client.Port = uint16(expiredPort)
  139. _, err = client.RegisterKey(context.Background(), WireGuardServer{Hostname: "example.com", IP: netip.MustParseAddr(expiredHost)}, token, key)
  140. if CodeOf(err) != CodeTLSValidation {
  141. t.Fatalf("expired certificate returned %s, want %s: %v", CodeOf(err), CodeTLSValidation, err)
  142. }
  143. }
  144. func TestRegistrationResponseGuardsAndRedirect(t *testing.T) {
  145. var destinationHits atomic.Int32
  146. destination := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
  147. destinationHits.Add(1)
  148. w.WriteHeader(http.StatusOK)
  149. }))
  150. defer destination.Close()
  151. tests := []struct {
  152. name, contentType, body, redirect string
  153. maxBody int64
  154. delay time.Duration
  155. wantCode string
  156. }{
  157. {name: "HTML", contentType: "text/html", body: "<html>maintenance</html>", wantCode: CodeRegistrationInvalid},
  158. {name: "oversized", contentType: "application/json", body: strings.Repeat("x", 65), maxBody: 64, wantCode: CodeRegistrationInvalid},
  159. {name: "redirect", contentType: "application/json", redirect: destination.URL, wantCode: CodeNetworkUnavailable},
  160. {name: "timeout", contentType: "application/json", body: `{"status":"OK"}`, delay: 100 * time.Millisecond, wantCode: CodeTimeout},
  161. }
  162. const token = "test-token-value-that-is-long-enough"
  163. for _, test := range tests {
  164. t.Run(test.name, func(t *testing.T) {
  165. server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  166. if test.delay > 0 {
  167. time.Sleep(test.delay)
  168. }
  169. if test.redirect != "" {
  170. http.Redirect(w, r, test.redirect, http.StatusTemporaryRedirect)
  171. return
  172. }
  173. w.Header().Set("Content-Type", test.contentType)
  174. _, _ = w.Write([]byte(test.body))
  175. }))
  176. server.StartTLS()
  177. defer server.Close()
  178. caPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: server.Certificate().Raw})
  179. host, portText, err := net.SplitHostPort(server.Listener.Addr().String())
  180. if err != nil {
  181. t.Fatal(err)
  182. }
  183. port, err := net.LookupPort("tcp", portText)
  184. if err != nil {
  185. t.Fatal(err)
  186. }
  187. client := NewRegistrationClient(caPEM)
  188. client.Port = uint16(port)
  189. if test.maxBody > 0 {
  190. client.MaxBody = test.maxBody
  191. }
  192. if test.delay > 0 {
  193. client.Timeout = 25 * time.Millisecond
  194. }
  195. _, err = client.RegisterKey(
  196. context.Background(),
  197. WireGuardServer{Hostname: "example.com", IP: netip.MustParseAddr(host)},
  198. token,
  199. testPubKey(),
  200. )
  201. if err == nil {
  202. t.Fatal("expected registration error")
  203. }
  204. if CodeOf(err) != test.wantCode {
  205. t.Fatalf("got %s, want %s: %v", CodeOf(err), test.wantCode, err)
  206. }
  207. if containsSecret(err.Error(), token) {
  208. t.Fatalf("token leaked in error: %v", err)
  209. }
  210. })
  211. }
  212. if destinationHits.Load() != 0 {
  213. t.Fatal("registration request followed a redirect and exposed secrets")
  214. }
  215. }
  216. func newExpiredTLSServer(t *testing.T, response []byte) (*httptest.Server, []byte) {
  217. t.Helper()
  218. key, err := rsa.GenerateKey(rand.Reader, 2048)
  219. if err != nil {
  220. t.Fatal(err)
  221. }
  222. template := &x509.Certificate{
  223. SerialNumber: big.NewInt(1),
  224. Subject: pkix.Name{CommonName: "example.com"},
  225. DNSNames: []string{"example.com"},
  226. NotBefore: time.Now().Add(-48 * time.Hour),
  227. NotAfter: time.Now().Add(-24 * time.Hour),
  228. KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment | x509.KeyUsageCertSign,
  229. ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
  230. IsCA: true,
  231. BasicConstraintsValid: true,
  232. }
  233. der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key)
  234. if err != nil {
  235. t.Fatal(err)
  236. }
  237. certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})
  238. keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(key)})
  239. certificate, err := tls.X509KeyPair(certPEM, keyPEM)
  240. if err != nil {
  241. t.Fatal(err)
  242. }
  243. server := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
  244. w.Header().Set("Content-Type", "application/json")
  245. _, _ = w.Write(response)
  246. }))
  247. server.TLS = &tls.Config{Certificates: []tls.Certificate{certificate}, MinVersion: tls.VersionTLS12}
  248. server.StartTLS()
  249. return server, certPEM
  250. }