catalog_test.go 3.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125
  1. package pia
  2. import (
  3. "context"
  4. "os"
  5. "path/filepath"
  6. "sync"
  7. "sync/atomic"
  8. "testing"
  9. "time"
  10. )
  11. type fakeServerListSource struct {
  12. snapshot ServerListSnapshot
  13. err error
  14. calls int
  15. }
  16. func (f *fakeServerListSource) Fetch(context.Context) (ServerListSnapshot, error) {
  17. f.calls++
  18. return f.snapshot, f.err
  19. }
  20. func TestCatalogCachesOnlyVerifiedParsedSnapshots(t *testing.T) {
  21. raw, err := os.ReadFile(filepath.Join("testdata", "serverlist", "v6_valid.json"))
  22. if err != nil {
  23. t.Fatal(err)
  24. }
  25. source := &fakeServerListSource{snapshot: ServerListSnapshot{Payload: raw, SchemaHint: "6", SignatureVerified: true}}
  26. now := time.Unix(1_700_000_000, 0)
  27. catalog := NewCatalog(source)
  28. catalog.CacheTTL = 30 * time.Minute
  29. catalog.Now = func() time.Time { return now }
  30. first, schema, err := catalog.ListRegions(context.Background())
  31. if err != nil || schema != "v6" || len(first) != 1 {
  32. t.Fatalf("unexpected first result: schema=%q regions=%v err=%v", schema, first, err)
  33. }
  34. first[0].WireGuard[0].Hostname = "mutated-by-caller"
  35. second, _, err := catalog.ListRegions(context.Background())
  36. if err != nil || source.calls != 1 {
  37. t.Fatalf("verified snapshot was not cached: calls=%d err=%v", source.calls, err)
  38. }
  39. if second[0].WireGuard[0].Hostname == "mutated-by-caller" {
  40. t.Fatal("catalog returned mutable cached storage")
  41. }
  42. now = now.Add(-time.Second)
  43. if _, _, err := catalog.ListRegions(context.Background()); err != nil || source.calls != 2 {
  44. t.Fatalf("backward clock movement incorrectly extended the cache: calls=%d err=%v", source.calls, err)
  45. }
  46. now = now.Add(catalog.CacheTTL + time.Second)
  47. if _, _, err := catalog.ListRegions(context.Background()); err != nil || source.calls != 3 {
  48. t.Fatalf("expired snapshot was not refreshed: calls=%d err=%v", source.calls, err)
  49. }
  50. }
  51. type gatedServerListSource struct {
  52. snapshot ServerListSnapshot
  53. started chan struct{}
  54. release chan struct{}
  55. startOnce sync.Once
  56. calls atomic.Int32
  57. }
  58. func (g *gatedServerListSource) Fetch(context.Context) (ServerListSnapshot, error) {
  59. g.calls.Add(1)
  60. g.startOnce.Do(func() { close(g.started) })
  61. <-g.release
  62. return g.snapshot, nil
  63. }
  64. func TestCatalogCoalescesConcurrentRefresh(t *testing.T) {
  65. raw, err := os.ReadFile(filepath.Join("testdata", "serverlist", "v6_valid.json"))
  66. if err != nil {
  67. t.Fatal(err)
  68. }
  69. source := &gatedServerListSource{
  70. snapshot: ServerListSnapshot{Payload: raw, SchemaHint: "6", SignatureVerified: true},
  71. started: make(chan struct{}),
  72. release: make(chan struct{}),
  73. }
  74. catalog := NewCatalog(source)
  75. catalog.CacheTTL = time.Hour
  76. errc := make(chan error, 2)
  77. go func() {
  78. _, _, err := catalog.ListRegions(context.Background())
  79. errc <- err
  80. }()
  81. <-source.started
  82. go func() {
  83. _, _, err := catalog.ListRegions(context.Background())
  84. errc <- err
  85. }()
  86. deadline := time.Now().Add(200 * time.Millisecond)
  87. for time.Now().Before(deadline) {
  88. if source.calls.Load() > 1 {
  89. close(source.release)
  90. t.Fatalf("concurrent refresh issued %d fetches, want 1", source.calls.Load())
  91. }
  92. time.Sleep(time.Millisecond)
  93. }
  94. close(source.release)
  95. for i := 0; i < 2; i++ {
  96. if err := <-errc; err != nil {
  97. t.Fatal(err)
  98. }
  99. }
  100. if source.calls.Load() != 1 {
  101. t.Fatalf("concurrent refresh issued %d fetches, want 1", source.calls.Load())
  102. }
  103. }
  104. func TestCatalogRejectsUnverifiedSnapshot(t *testing.T) {
  105. source := &fakeServerListSource{snapshot: ServerListSnapshot{
  106. Payload: []byte(`{"version":6,"groups":{"wg":[]},"regions":[]}`), SchemaHint: "6", SignatureVerified: false,
  107. }}
  108. catalog := NewCatalog(source)
  109. _, _, err := catalog.ListRegions(context.Background())
  110. if CodeOf(err) != CodeCatalogSignatureInvalid {
  111. t.Fatalf("unverified snapshot returned %s, want %s: %v", CodeOf(err), CodeCatalogSignatureInvalid, err)
  112. }
  113. }