diff --git a/.codespellrc b/.codespellrc new file mode 100644 index 0000000..64f2a5d --- /dev/null +++ b/.codespellrc @@ -0,0 +1,2 @@ +[codespell] +ignore-words-list = te diff --git a/.gitignore b/.gitignore index 02fe8b7..12426ab 100644 --- a/.gitignore +++ b/.gitignore @@ -4,3 +4,4 @@ wicket *.swp .DS_Store .loq_cache +coverage.out diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 4b8c544..508604a 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -44,8 +44,8 @@ repos: pass_filenames: false - id: go-test - name: go test - entry: go test ./... + name: go test (100% coverage) + entry: ./check-coverage 100 language: system types: [go] pass_filenames: false diff --git a/broker.go b/broker.go new file mode 100644 index 0000000..ed8a5b6 --- /dev/null +++ b/broker.go @@ -0,0 +1,150 @@ +package main + +import ( + "strings" + "sync" +) + +type subscriber struct { + ch chan *Event + path string +} + +type Broker struct { + mu sync.RWMutex + subscribers map[string][]*subscriber + buffer *RingBuffer +} + +func NewBroker(bufferSize int) *Broker { + return &Broker{ + subscribers: make(map[string][]*subscriber), + buffer: NewRingBuffer(bufferSize), + } +} + +func (b *Broker) Publish(event *Event) { + b.buffer.Add(event) + + b.mu.RLock() + defer b.mu.RUnlock() + + for _, path := range publishPaths(event.Path) { + for _, sub := range b.subscribers[path] { + select { + case sub.ch <- event: + default: + } + } + } +} + +func (b *Broker) Subscribe(path string, lastEventID string) (<-chan *Event, func()) { + ch := make(chan *Event, 64) + sub := &subscriber{ch: ch, path: path} + + b.mu.Lock() + b.subscribers[path] = append(b.subscribers[path], sub) + b.mu.Unlock() + + if lastEventID != "" { + events := b.buffer.Since(lastEventID, path) + for _, e := range events { + ch <- e + } + } + + unsub := func() { + b.mu.Lock() + defer b.mu.Unlock() + subs := b.subscribers[path] + for i, s := range subs { + if s == sub { + b.subscribers[path] = append(subs[:i], subs[i+1:]...) + break + } + } + close(ch) + } + + return ch, unsub +} + +func publishPaths(eventPath string) []string { + paths := []string{eventPath} + for { + i := strings.LastIndex(eventPath, "/") + if i < 0 { + break + } + eventPath = eventPath[:i] + paths = append(paths, eventPath) + } + if paths[len(paths)-1] != "" { + paths = append(paths, "") + } + return paths +} + +type RingBuffer struct { + mu sync.RWMutex + buf []*Event + size int + write int + count int +} + +func NewRingBuffer(size int) *RingBuffer { + return &RingBuffer{ + buf: make([]*Event, size), + size: size, + } +} + +func (rb *RingBuffer) Add(event *Event) { + rb.mu.Lock() + defer rb.mu.Unlock() + rb.buf[rb.write%rb.size] = event + rb.write++ + if rb.count < rb.size { + rb.count++ + } +} + +func (rb *RingBuffer) Since(lastEventID string, subscribePath string) []*Event { + rb.mu.RLock() + defer rb.mu.RUnlock() + + start := rb.write - rb.count + found := false + foundIdx := start + + for i := start; i < rb.write; i++ { + e := rb.buf[i%rb.size] + if e.ID == lastEventID { + found = true + foundIdx = i + 1 + break + } + } + + if !found { + foundIdx = start + } + + var result []*Event + for i := foundIdx; i < rb.write; i++ { + e := rb.buf[i%rb.size] + if pathMatches(subscribePath, e.Path) { + result = append(result, e) + } + } + return result +} + +func pathMatches(subscribePath, eventPath string) bool { + if subscribePath == "" { + return true + } + return eventPath == subscribePath || strings.HasPrefix(eventPath, subscribePath+"/") +} diff --git a/broker_test.go b/broker_test.go new file mode 100644 index 0000000..b8bd803 --- /dev/null +++ b/broker_test.go @@ -0,0 +1,229 @@ +package main + +import ( + "fmt" + "testing" + "time" +) + +func newTestEvent(path string) *Event { + return &Event{ + ID: "test-" + path, + Timestamp: time.Now(), + Path: path, + Payload: map[string]any{"test": true}, + } +} + +func mustReceive(t *testing.T, ch <-chan *Event, timeout time.Duration) *Event { + t.Helper() + select { + case e := <-ch: + return e + case <-time.After(timeout): + t.Fatal("timed out waiting for event") + return nil + } +} + +func mustNotReceive(t *testing.T, ch <-chan *Event, timeout time.Duration) { + t.Helper() + select { + case e := <-ch: + t.Fatalf("expected no event, got %+v", e) + case <-time.After(timeout): + } +} + +func TestBroker_exactPathDelivery(t *testing.T) { + b := NewBroker(100) + ch, unsub := b.Subscribe("github.com/chrisguidry/docketeer", "") + defer unsub() + + event := newTestEvent("github.com/chrisguidry/docketeer") + b.Publish(event) + + got := mustReceive(t, ch, time.Second) + if got.ID != event.ID { + t.Errorf("expected event %s, got %s", event.ID, got.ID) + } +} + +func TestBroker_parentReceivesChildEvents(t *testing.T) { + b := NewBroker(100) + ch, unsub := b.Subscribe("github.com/chrisguidry", "") + defer unsub() + + event := newTestEvent("github.com/chrisguidry/docketeer") + b.Publish(event) + + got := mustReceive(t, ch, time.Second) + if got.ID != event.ID { + t.Errorf("expected event %s, got %s", event.ID, got.ID) + } +} + +func TestBroker_rootReceivesAll(t *testing.T) { + b := NewBroker(100) + ch, unsub := b.Subscribe("", "") + defer unsub() + + event := newTestEvent("github.com/chrisguidry/docketeer") + b.Publish(event) + + got := mustReceive(t, ch, time.Second) + if got.ID != event.ID { + t.Errorf("expected event %s, got %s", event.ID, got.ID) + } +} + +func TestBroker_unrelatedSubscriberDoesNotReceive(t *testing.T) { + b := NewBroker(100) + ch, unsub := b.Subscribe("gitlab.com", "") + defer unsub() + + b.Publish(newTestEvent("github.com/chrisguidry/docketeer")) + + mustNotReceive(t, ch, 50*time.Millisecond) +} + +func TestBroker_multipleSubscribersSamePath(t *testing.T) { + b := NewBroker(100) + ch1, unsub1 := b.Subscribe("github.com/chrisguidry/docketeer", "") + defer unsub1() + ch2, unsub2 := b.Subscribe("github.com/chrisguidry/docketeer", "") + defer unsub2() + + b.Publish(newTestEvent("github.com/chrisguidry/docketeer")) + + mustReceive(t, ch1, time.Second) + mustReceive(t, ch2, time.Second) +} + +func TestBroker_unsubscribeStopsDelivery(t *testing.T) { + b := NewBroker(100) + ch, unsub := b.Subscribe("github.com/chrisguidry/docketeer", "") + unsub() + + b.Publish(newTestEvent("github.com/chrisguidry/docketeer")) + + select { + case _, ok := <-ch: + if ok { + t.Fatal("expected channel to be closed, but received an event") + } + case <-time.After(50 * time.Millisecond): + t.Fatal("expected channel to be closed") + } +} + +func TestBroker_ringBufferWraps(t *testing.T) { + b := NewBroker(5) + for i := range 10 { + b.Publish(&Event{ + ID: fmt.Sprintf("event-%d", i), + Path: "test", + Payload: map[string]any{}, + }) + } + + ch, unsub := b.Subscribe("test", "event-4") + defer unsub() + + for i := 5; i < 10; i++ { + got := mustReceive(t, ch, time.Second) + expected := fmt.Sprintf("event-%d", i) + if got.ID != expected { + t.Errorf("expected %s, got %s", expected, got.ID) + } + } +} + +func TestBroker_lastEventIDReplay(t *testing.T) { + b := NewBroker(100) + b.Publish(&Event{ID: "e1", Path: "test", Payload: map[string]any{}}) + b.Publish(&Event{ID: "e2", Path: "test", Payload: map[string]any{}}) + b.Publish(&Event{ID: "e3", Path: "test", Payload: map[string]any{}}) + + ch, unsub := b.Subscribe("test", "e1") + defer unsub() + + got1 := mustReceive(t, ch, time.Second) + got2 := mustReceive(t, ch, time.Second) + if got1.ID != "e2" || got2.ID != "e3" { + t.Errorf("expected e2 and e3, got %s and %s", got1.ID, got2.ID) + } +} + +func TestBroker_lastEventIDRespectsPathHierarchy(t *testing.T) { + b := NewBroker(100) + b.Publish(&Event{ID: "e1", Path: "github.com/chrisguidry/docketeer", Payload: map[string]any{}}) + b.Publish(&Event{ID: "e2", Path: "gitlab.com/other/project", Payload: map[string]any{}}) + b.Publish(&Event{ID: "e3", Path: "github.com/chrisguidry/other", Payload: map[string]any{}}) + + ch, unsub := b.Subscribe("github.com/chrisguidry", "e1") + defer unsub() + + got := mustReceive(t, ch, time.Second) + if got.ID != "e3" { + t.Errorf("expected e3 (same prefix), got %s", got.ID) + } + mustNotReceive(t, ch, 50*time.Millisecond) +} + +func TestBroker_lastEventIDReplayExactPath(t *testing.T) { + b := NewBroker(100) + b.Publish(&Event{ID: "e1", Path: "exact/path", Payload: map[string]any{}}) + b.Publish(&Event{ID: "e2", Path: "exact/path/child", Payload: map[string]any{}}) + b.Publish(&Event{ID: "e3", Path: "exact/path", Payload: map[string]any{}}) + + ch, unsub := b.Subscribe("exact/path", "e1") + defer unsub() + + got1 := mustReceive(t, ch, time.Second) + got2 := mustReceive(t, ch, time.Second) + if got1.ID != "e2" { + t.Errorf("expected e2, got %s", got1.ID) + } + if got2.ID != "e3" { + t.Errorf("expected e3, got %s", got2.ID) + } +} + +func TestPathMatches(t *testing.T) { + tests := []struct { + subscribePath string + eventPath string + want bool + }{ + {"", "anything", true}, + {"exact", "exact", true}, + {"parent", "parent/child", true}, + {"parent", "parentchild", false}, + {"parent", "other", false}, + } + for _, tt := range tests { + got := pathMatches(tt.subscribePath, tt.eventPath) + if got != tt.want { + t.Errorf("pathMatches(%q, %q) = %v, want %v", tt.subscribePath, tt.eventPath, got, tt.want) + } + } +} + +func TestBroker_lastEventIDExpiredFromBuffer(t *testing.T) { + b := NewBroker(3) + b.Publish(&Event{ID: "old1", Path: "test", Payload: map[string]any{}}) + b.Publish(&Event{ID: "old2", Path: "test", Payload: map[string]any{}}) + b.Publish(&Event{ID: "old3", Path: "test", Payload: map[string]any{}}) + b.Publish(&Event{ID: "new1", Path: "test", Payload: map[string]any{}}) + b.Publish(&Event{ID: "new2", Path: "test", Payload: map[string]any{}}) + b.Publish(&Event{ID: "new3", Path: "test", Payload: map[string]any{}}) + + ch, unsub := b.Subscribe("test", "old1") + defer unsub() + + got := mustReceive(t, ch, time.Second) + if got.ID != "new1" { + t.Errorf("expected new1 (oldest in buffer), got %s", got.ID) + } +} diff --git a/check-coverage b/check-coverage new file mode 100755 index 0000000..220f78b --- /dev/null +++ b/check-coverage @@ -0,0 +1,18 @@ +#!/bin/bash +set -euo pipefail + +THRESHOLD="${1:-100}" +PROFILE=$(mktemp) +trap 'rm -f "$PROFILE"' EXIT + +go test -coverprofile="$PROFILE" -covermode=atomic ./... + +# Exclude main.go from coverage calculation (entrypoint, not unit-testable) +COVERAGE=$(grep -v "main\.go:" "$PROFILE" | go tool cover -func=/dev/stdin | grep ^total: | awk '{print $3}' | tr -d '%') + +if awk "BEGIN{exit(!($COVERAGE < $THRESHOLD))}"; then + echo "FAIL: coverage ${COVERAGE}% is below ${THRESHOLD}%" + exit 1 +fi + +echo "OK: coverage ${COVERAGE}% meets ${THRESHOLD}% threshold" diff --git a/configuration.go b/configuration.go new file mode 100644 index 0000000..6b17d99 --- /dev/null +++ b/configuration.go @@ -0,0 +1,59 @@ +package main + +import ( + "os" + "strings" + + "gopkg.in/yaml.v3" +) + +type Configuration struct { + Paths map[string]PathConfiguration `yaml:"paths"` +} + +type PathConfiguration struct { + Verify string `yaml:"verify"` + Secret string `yaml:"secret"` + SignatureHeader string `yaml:"signature_header"` + SubscribeSecret string `yaml:"subscribe_secret"` +} + +func LoadConfiguration(path string) (*Configuration, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, err + } + var cfg Configuration + if err := yaml.Unmarshal(data, &cfg); err != nil { + return nil, err + } + return &cfg, nil +} + +func (c *Configuration) LookupSubscribeSecret(path string) string { + if c == nil { + return "" + } + for { + if pc, ok := c.Paths[path]; ok && pc.SubscribeSecret != "" { + return pc.SubscribeSecret + } + i := strings.LastIndex(path, "/") + if i < 0 { + break + } + path = path[:i] + } + return "" +} + +func (c *Configuration) LookupVerification(path string) *PathConfiguration { + if c == nil { + return nil + } + pc, ok := c.Paths[path] + if !ok || pc.Verify == "" { + return nil + } + return &pc +} diff --git a/configuration_test.go b/configuration_test.go new file mode 100644 index 0000000..3a7aa6f --- /dev/null +++ b/configuration_test.go @@ -0,0 +1,148 @@ +package main + +import ( + "os" + "path/filepath" + "testing" +) + +func TestLoadConfiguration_valid(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "wicket.yaml") + os.WriteFile(path, []byte(` +paths: + github.com/chrisguidry/docketeer: + verify: hmac-sha256 + secret: "webhook-secret" + signature_header: X-Hub-Signature-256 + subscribe_secret: "sub-token" +`), 0644) + + cfg, err := LoadConfiguration(path) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + pc, ok := cfg.Paths["github.com/chrisguidry/docketeer"] + if !ok { + t.Fatal("expected path config for github.com/chrisguidry/docketeer") + } + if pc.Verify != "hmac-sha256" { + t.Errorf("expected hmac-sha256, got %s", pc.Verify) + } + if pc.Secret != "webhook-secret" { + t.Errorf("expected webhook-secret, got %s", pc.Secret) + } + if pc.SignatureHeader != "X-Hub-Signature-256" { + t.Errorf("expected X-Hub-Signature-256, got %s", pc.SignatureHeader) + } + if pc.SubscribeSecret != "sub-token" { + t.Errorf("expected sub-token, got %s", pc.SubscribeSecret) + } +} + +func TestLoadConfiguration_missingFile(t *testing.T) { + _, err := LoadConfiguration("/nonexistent/wicket.yaml") + if err == nil { + t.Fatal("expected error for missing file") + } +} + +func TestLoadConfiguration_empty(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "wicket.yaml") + os.WriteFile(path, []byte(""), 0644) + + cfg, err := LoadConfiguration(path) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if cfg.Paths == nil { + cfg.Paths = make(map[string]PathConfiguration) + } + if len(cfg.Paths) != 0 { + t.Errorf("expected no paths, got %d", len(cfg.Paths)) + } +} + +func TestLoadConfiguration_invalidYAML(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "wicket.yaml") + os.WriteFile(path, []byte("not: valid: yaml: [[["), 0644) + + _, err := LoadConfiguration(path) + if err == nil { + t.Fatal("expected error for invalid YAML") + } +} + +func TestLookupSubscribeSecret_exactMatch(t *testing.T) { + cfg := &Configuration{ + Paths: map[string]PathConfiguration{ + "github.com/chrisguidry/docketeer": {SubscribeSecret: "token-123"}, + }, + } + secret := cfg.LookupSubscribeSecret("github.com/chrisguidry/docketeer") + if secret != "token-123" { + t.Errorf("expected token-123, got %s", secret) + } +} + +func TestLookupSubscribeSecret_inheritsFromParent(t *testing.T) { + cfg := &Configuration{ + Paths: map[string]PathConfiguration{ + "github.com/chrisguidry": {SubscribeSecret: "parent-token"}, + }, + } + secret := cfg.LookupSubscribeSecret("github.com/chrisguidry/docketeer") + if secret != "parent-token" { + t.Errorf("expected parent-token, got %s", secret) + } +} + +func TestLookupSubscribeSecret_noConfig(t *testing.T) { + cfg := &Configuration{ + Paths: map[string]PathConfiguration{}, + } + secret := cfg.LookupSubscribeSecret("github.com/chrisguidry/docketeer") + if secret != "" { + t.Errorf("expected empty string, got %s", secret) + } +} + +func TestLookupPathConfiguration_exactMatchOnly(t *testing.T) { + cfg := &Configuration{ + Paths: map[string]PathConfiguration{ + "github.com/chrisguidry/docketeer": { + Verify: "hmac-sha256", + Secret: "webhook-secret", + SignatureHeader: "X-Hub-Signature-256", + }, + }, + } + + pc := cfg.LookupVerification("github.com/chrisguidry/docketeer") + if pc == nil { + t.Fatal("expected path config for exact match") + } + + pc = cfg.LookupVerification("github.com/chrisguidry/docketeer/subpath") + if pc != nil { + t.Fatal("verification should not inherit from parent") + } +} + +func TestLookupSubscribeSecret_nilConfig(t *testing.T) { + var cfg *Configuration + secret := cfg.LookupSubscribeSecret("any/path") + if secret != "" { + t.Errorf("expected empty string from nil config, got %s", secret) + } +} + +func TestLookupVerification_nilConfig(t *testing.T) { + var cfg *Configuration + pc := cfg.LookupVerification("any/path") + if pc != nil { + t.Error("expected nil from nil config") + } +} diff --git a/event.go b/event.go new file mode 100644 index 0000000..b1dae58 --- /dev/null +++ b/event.go @@ -0,0 +1,11 @@ +package main + +import "time" + +type Event struct { + ID string `json:"id"` + Timestamp time.Time `json:"timestamp"` + Path string `json:"path"` + Headers map[string]string `json:"headers"` + Payload any `json:"payload"` +} diff --git a/filter.go b/filter.go new file mode 100644 index 0000000..2205c18 --- /dev/null +++ b/filter.go @@ -0,0 +1,84 @@ +package main + +import ( + "fmt" + "net/url" + "strings" +) + +type Filter struct { + Path string + Value string +} + +func ParseFilters(query url.Values) []Filter { + raw := query["filter"] + if len(raw) == 0 { + return nil + } + filters := make([]Filter, 0, len(raw)) + for _, f := range raw { + path, value, ok := strings.Cut(f, ":") + if !ok { + continue + } + filters = append(filters, Filter{Path: path, Value: value}) + } + return filters +} + +func MatchAll(filters []Filter, event *Event) bool { + for _, f := range filters { + if !matchOne(f, event) { + return false + } + } + return true +} + +func matchOne(f Filter, event *Event) bool { + val, ok := resolveField(f.Path, event) + if !ok { + return false + } + return val == f.Value +} + +func resolveField(path string, event *Event) (string, bool) { + parts := strings.Split(path, ".") + + switch parts[0] { + case "id": + return event.ID, len(parts) == 1 + case "path": + return event.Path, len(parts) == 1 + case "timestamp": + return event.Timestamp.Format("2006-01-02T15:04:05Z07:00"), len(parts) == 1 + case "headers": + if len(parts) != 2 { + return "", false + } + val, ok := event.Headers[parts[1]] + return val, ok + case "payload": + return navigateJSON(event.Payload, parts[1:]) + default: + return "", false + } +} + +func navigateJSON(val any, parts []string) (string, bool) { + if len(parts) == 0 { + return fmt.Sprintf("%v", val), true + } + + m, ok := val.(map[string]any) + if !ok { + return "", false + } + next, ok := m[parts[0]] + if !ok { + return "", false + } + return navigateJSON(next, parts[1:]) +} diff --git a/filter_test.go b/filter_test.go new file mode 100644 index 0000000..cd7be02 --- /dev/null +++ b/filter_test.go @@ -0,0 +1,232 @@ +package main + +import ( + "net/url" + "testing" + "time" +) + +func TestParseFilters_empty(t *testing.T) { + filters := ParseFilters(url.Values{}) + if len(filters) != 0 { + t.Fatalf("expected no filters, got %d", len(filters)) + } +} + +func TestParseFilters_single(t *testing.T) { + v := url.Values{"filter": {"payload.ref:refs/heads/main"}} + filters := ParseFilters(v) + if len(filters) != 1 { + t.Fatalf("expected 1 filter, got %d", len(filters)) + } + if filters[0].Path != "payload.ref" { + t.Errorf("expected path payload.ref, got %s", filters[0].Path) + } + if filters[0].Value != "refs/heads/main" { + t.Errorf("expected value refs/heads/main, got %s", filters[0].Value) + } +} + +func TestParseFilters_multiple(t *testing.T) { + v := url.Values{"filter": { + "payload.ref:refs/heads/main", + "headers.X-GitHub-Event:push", + }} + filters := ParseFilters(v) + if len(filters) != 2 { + t.Fatalf("expected 2 filters, got %d", len(filters)) + } +} + +func TestParseFilters_colonInValue(t *testing.T) { + v := url.Values{"filter": {"payload.url:https://example.com"}} + filters := ParseFilters(v) + if len(filters) != 1 { + t.Fatalf("expected 1 filter, got %d", len(filters)) + } + if filters[0].Value != "https://example.com" { + t.Errorf("expected value with colon preserved, got %s", filters[0].Value) + } +} + +func TestMatchAll_emptyFilters(t *testing.T) { + event := &Event{ + Payload: map[string]any{"ref": "refs/heads/main"}, + } + if !MatchAll(nil, event) { + t.Error("empty filters should match everything") + } +} + +func TestMatchAll_singleMatch(t *testing.T) { + event := &Event{ + Payload: map[string]any{"ref": "refs/heads/main"}, + } + filters := []Filter{{Path: "payload.ref", Value: "refs/heads/main"}} + if !MatchAll(filters, event) { + t.Error("expected filter to match") + } +} + +func TestMatchAll_singleNoMatch(t *testing.T) { + event := &Event{ + Payload: map[string]any{"ref": "refs/heads/develop"}, + } + filters := []Filter{{Path: "payload.ref", Value: "refs/heads/main"}} + if MatchAll(filters, event) { + t.Error("expected filter not to match") + } +} + +func TestMatchAll_multipleAllMatch(t *testing.T) { + event := &Event{ + Headers: map[string]string{"X-GitHub-Event": "push"}, + Payload: map[string]any{"ref": "refs/heads/main"}, + } + filters := []Filter{ + {Path: "payload.ref", Value: "refs/heads/main"}, + {Path: "headers.X-GitHub-Event", Value: "push"}, + } + if !MatchAll(filters, event) { + t.Error("expected all filters to match") + } +} + +func TestMatchAll_multipleOneFails(t *testing.T) { + event := &Event{ + Headers: map[string]string{"X-GitHub-Event": "push"}, + Payload: map[string]any{"ref": "refs/heads/develop"}, + } + filters := []Filter{ + {Path: "payload.ref", Value: "refs/heads/main"}, + {Path: "headers.X-GitHub-Event", Value: "push"}, + } + if MatchAll(filters, event) { + t.Error("expected AND filter to fail when one doesn't match") + } +} + +func TestMatchAll_nestedDotPath(t *testing.T) { + event := &Event{ + Payload: map[string]any{ + "repository": map[string]any{ + "full_name": "chrisguidry/docketeer", + }, + }, + } + filters := []Filter{{Path: "payload.repository.full_name", Value: "chrisguidry/docketeer"}} + if !MatchAll(filters, event) { + t.Error("expected nested dot path to match") + } +} + +func TestMatchAll_missingField(t *testing.T) { + event := &Event{ + Payload: map[string]any{"ref": "refs/heads/main"}, + } + filters := []Filter{{Path: "payload.nonexistent.field", Value: "anything"}} + if MatchAll(filters, event) { + t.Error("missing field should not match") + } +} + +func TestMatchAll_topLevelFields(t *testing.T) { + event := &Event{ + ID: "abc-123", + Path: "github.com/chrisguidry/docketeer", + } + filters := []Filter{{Path: "path", Value: "github.com/chrisguidry/docketeer"}} + if !MatchAll(filters, event) { + t.Error("expected top-level path field to match") + } +} + +func TestMatchAll_idField(t *testing.T) { + event := &Event{ID: "abc-123"} + filters := []Filter{{Path: "id", Value: "abc-123"}} + if !MatchAll(filters, event) { + t.Error("expected id field to match") + } +} + +func TestMatchAll_timestampField(t *testing.T) { + event := &Event{Timestamp: time.Date(2026, 3, 4, 12, 0, 0, 0, time.UTC)} + filters := []Filter{{Path: "timestamp", Value: "2026-03-04T12:00:00Z"}} + if !MatchAll(filters, event) { + t.Error("expected timestamp field to match") + } +} + +func TestMatchAll_headersNeedsTwoParts(t *testing.T) { + event := &Event{Headers: map[string]string{"X-Foo": "bar"}} + filters := []Filter{{Path: "headers", Value: "anything"}} + if MatchAll(filters, event) { + t.Error("headers without key should not match") + } +} + +func TestMatchAll_unknownTopLevel(t *testing.T) { + event := &Event{} + filters := []Filter{{Path: "nonexistent", Value: "anything"}} + if MatchAll(filters, event) { + t.Error("unknown top-level field should not match") + } +} + +func TestMatchAll_emptyPath(t *testing.T) { + event := &Event{} + filters := []Filter{{Path: "", Value: "anything"}} + if MatchAll(filters, event) { + t.Error("empty path should not match") + } +} + +func TestMatchAll_payloadNonMapNavigation(t *testing.T) { + event := &Event{Payload: "just a string"} + filters := []Filter{{Path: "payload.deep.field", Value: "anything"}} + if MatchAll(filters, event) { + t.Error("navigating into non-map payload should not match") + } +} + +func TestParseFilters_invalidFormat(t *testing.T) { + v := url.Values{"filter": {"no-colon-here"}} + filters := ParseFilters(v) + if len(filters) != 0 { + t.Errorf("expected 0 filters for invalid format, got %d", len(filters)) + } +} + +func TestMatchAll_idWithSubpath(t *testing.T) { + event := &Event{ID: "abc-123"} + filters := []Filter{{Path: "id.sub", Value: "anything"}} + if MatchAll(filters, event) { + t.Error("id with subpath should not match") + } +} + +func TestMatchAll_pathWithSubpath(t *testing.T) { + event := &Event{Path: "some/path"} + filters := []Filter{{Path: "path.sub", Value: "anything"}} + if MatchAll(filters, event) { + t.Error("path with subpath should not match") + } +} + +func TestMatchAll_timestampWithSubpath(t *testing.T) { + event := &Event{} + filters := []Filter{{Path: "timestamp.sub", Value: "anything"}} + if MatchAll(filters, event) { + t.Error("timestamp with subpath should not match") + } +} + +func TestMatchAll_payloadLeafValue(t *testing.T) { + event := &Event{ + Payload: map[string]any{"count": 42}, + } + filters := []Filter{{Path: "payload.count", Value: "42"}} + if !MatchAll(filters, event) { + t.Error("expected numeric leaf value to match via fmt.Sprintf") + } +} diff --git a/go.mod b/go.mod index e932d2f..4a02013 100644 --- a/go.mod +++ b/go.mod @@ -1,3 +1,8 @@ module tangled.org/guid.foo/wicket go 1.24 + +require ( + github.com/google/uuid v1.6.0 + gopkg.in/yaml.v3 v3.0.1 +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..b4c5744 --- /dev/null +++ b/go.sum @@ -0,0 +1,6 @@ +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/main.go b/main.go index 09bb1a9..f0c55d7 100644 --- a/main.go +++ b/main.go @@ -1,16 +1,67 @@ package main import ( + "context" "flag" "fmt" + "log" + "net/http" "os" + "os/signal" + "sync/atomic" + "syscall" ) func main() { address := flag.String("address", ":8080", "listen address") - _ = flag.String("configuration", "", "path to configuration file") - _ = flag.Int("buffer-size", 1000, "event replay buffer size") + configPath := flag.String("configuration", "", "path to configuration file") + bufferSize := flag.Int("buffer-size", 1000, "event replay buffer size") flag.Parse() + var cfgPtr atomic.Pointer[Configuration] + if *configPath != "" { + cfg, err := LoadConfiguration(*configPath) + if err != nil { + log.Fatalf("loading configuration: %v", err) + } + cfgPtr.Store(cfg) + log.Printf("loaded configuration from %s", *configPath) + } + + broker := NewBroker(*bufferSize) + handler := NewServer(broker, &cfgPtr) + + server := &http.Server{ + Addr: *address, + Handler: handler, + } + + go func() { + sigs := make(chan os.Signal, 1) + signal.Notify(sigs, syscall.SIGHUP, syscall.SIGINT, syscall.SIGTERM) + for sig := range sigs { + switch sig { + case syscall.SIGHUP: + if *configPath == "" { + continue + } + cfg, err := LoadConfiguration(*configPath) + if err != nil { + log.Printf("reloading configuration: %v", err) + continue + } + cfgPtr.Store(cfg) + log.Printf("reloaded configuration from %s", *configPath) + case syscall.SIGINT, syscall.SIGTERM: + log.Printf("shutting down") + server.Shutdown(context.Background()) + return + } + } + }() + fmt.Fprintf(os.Stderr, "wicket listening on %s\n", *address) + if err := server.ListenAndServe(); err != http.ErrServerClosed { + log.Fatalf("server error: %v", err) + } } diff --git a/server.go b/server.go new file mode 100644 index 0000000..6f92706 --- /dev/null +++ b/server.go @@ -0,0 +1,160 @@ +package main + +import ( + "encoding/base64" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "sync/atomic" + "time" + + "github.com/google/uuid" +) + +type Server struct { + broker *Broker + config *atomic.Pointer[Configuration] +} + +func NewServer(broker *Broker, config *atomic.Pointer[Configuration]) http.Handler { + s := &Server{broker: broker, config: config} + return http.HandlerFunc(s.ServeHTTP) +} + +func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { + setCORSHeaders(w) + + switch r.Method { + case "OPTIONS": + w.WriteHeader(http.StatusNoContent) + case "POST": + s.handlePost(w, r) + case "GET": + if !strings.Contains(r.Header.Get("Accept"), "text/event-stream") { + http.NotFound(w, r) + return + } + s.handleSSE(w, r) + default: + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} + +func setCORSHeaders(w http.ResponseWriter) { + w.Header().Set("Access-Control-Allow-Origin", "*") + w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") + w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization, Last-Event-ID") +} + +func (s *Server) handlePost(w http.ResponseWriter, r *http.Request) { + path := strings.TrimPrefix(r.URL.Path, "/") + + body, _ := io.ReadAll(r.Body) + + cfg := s.config.Load() + if pc := cfg.LookupVerification(path); pc != nil { + verifier, err := NewVerifier(pc.Verify) + if err != nil { + http.Error(w, "server configuration error", http.StatusInternalServerError) + return + } + if err := verifier.Verify(body, r.Header, pc.Secret, pc.SignatureHeader); err != nil { + http.Error(w, "forbidden", http.StatusForbidden) + return + } + } + + var payload any + if strings.HasPrefix(r.Header.Get("Content-Type"), "application/json") { + if err := json.Unmarshal(body, &payload); err != nil { + http.Error(w, "invalid JSON", http.StatusBadRequest) + return + } + } else { + payload = base64.StdEncoding.EncodeToString(body) + } + + headers := extractHeaders(r.Header) + + event := &Event{ + ID: uuid.New().String(), + Timestamp: time.Now().UTC(), + Path: path, + Headers: headers, + Payload: payload, + } + + s.broker.Publish(event) + w.WriteHeader(http.StatusAccepted) +} + +func (s *Server) handleSSE(w http.ResponseWriter, r *http.Request) { + path := strings.TrimPrefix(r.URL.Path, "/") + + cfg := s.config.Load() + if secret := cfg.LookupSubscribeSecret(path); secret != "" { + auth := r.Header.Get("Authorization") + if !strings.HasPrefix(auth, "Bearer ") || strings.TrimPrefix(auth, "Bearer ") != secret { + http.Error(w, "unauthorized", http.StatusUnauthorized) + return + } + } + + flusher := w.(http.Flusher) + + filters := ParseFilters(r.URL.Query()) + lastEventID := r.Header.Get("Last-Event-ID") + + ch, unsub := s.broker.Subscribe(path, lastEventID) + defer unsub() + + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Connection", "keep-alive") + w.WriteHeader(http.StatusOK) + flusher.Flush() + + ctx := r.Context() + for { + select { + case <-ctx.Done(): + return + case event, ok := <-ch: + if !ok { + return + } + if !MatchAll(filters, event) { + continue + } + data, _ := json.Marshal(event) + fmt.Fprintf(w, "id: %s\ndata: %s\n\n", event.ID, data) + flusher.Flush() + } + } +} + +var hopByHopHeaders = map[string]bool{ + "Connection": true, + "Keep-Alive": true, + "Proxy-Authenticate": true, + "Proxy-Authorization": true, + "Te": true, + "Trailer": true, + "Transfer-Encoding": true, + "Upgrade": true, + "Host": true, + "Content-Length": true, +} + +func extractHeaders(h http.Header) map[string]string { + headers := make(map[string]string) + for name, values := range h { + if hopByHopHeaders[name] { + continue + } + headers[name] = values[0] + } + return headers +} diff --git a/server_test.go b/server_test.go new file mode 100644 index 0000000..84d9ac8 --- /dev/null +++ b/server_test.go @@ -0,0 +1,240 @@ +package main + +import ( + "crypto/hmac" + "crypto/sha256" + "encoding/hex" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" +) + +func newTestServer(cfg *Configuration) (*httptest.Server, *Broker) { + broker := NewBroker(100) + var cfgPtr atomic.Pointer[Configuration] + if cfg != nil { + cfgPtr.Store(cfg) + } + handler := NewServer(broker, &cfgPtr) + return httptest.NewServer(handler), broker +} + +func TestServer_postPublishesEvent(t *testing.T) { + ts, _ := newTestServer(nil) + defer ts.Close() + + resp, err := http.Post(ts.URL+"/test/topic", "application/json", strings.NewReader(`{"hello":"world"}`)) + if err != nil { + t.Fatalf("POST failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusAccepted { + t.Errorf("expected 202, got %d", resp.StatusCode) + } +} + +func TestServer_postWithValidHMAC(t *testing.T) { + cfg := &Configuration{ + Paths: map[string]PathConfiguration{ + "secure/path": { + Verify: "hmac-sha256", + Secret: "test-secret", + SignatureHeader: "X-Hub-Signature-256", + }, + }, + } + ts, _ := newTestServer(cfg) + defer ts.Close() + + body := `{"action":"push"}` + mac := hmac.New(sha256.New, []byte("test-secret")) + mac.Write([]byte(body)) + sig := "sha256=" + hex.EncodeToString(mac.Sum(nil)) + + req, _ := http.NewRequest("POST", ts.URL+"/secure/path", strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("X-Hub-Signature-256", sig) + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("POST failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusAccepted { + t.Errorf("expected 202, got %d", resp.StatusCode) + } +} + +func TestServer_postWithInvalidSignature(t *testing.T) { + cfg := &Configuration{ + Paths: map[string]PathConfiguration{ + "secure/path": { + Verify: "hmac-sha256", + Secret: "test-secret", + SignatureHeader: "X-Hub-Signature-256", + }, + }, + } + ts, _ := newTestServer(cfg) + defer ts.Close() + + req, _ := http.NewRequest("POST", ts.URL+"/secure/path", strings.NewReader(`{"bad":"data"}`)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("X-Hub-Signature-256", "sha256=deadbeef") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("POST failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusForbidden { + t.Errorf("expected 403, got %d", resp.StatusCode) + } +} + +func TestServer_postToUnconfiguredPath(t *testing.T) { + cfg := &Configuration{ + Paths: map[string]PathConfiguration{ + "secure/path": { + Verify: "hmac-sha256", + Secret: "test-secret", + SignatureHeader: "X-Hub-Signature-256", + }, + }, + } + ts, _ := newTestServer(cfg) + defer ts.Close() + + resp, err := http.Post(ts.URL+"/open/path", "application/json", strings.NewReader(`{"ok":true}`)) + if err != nil { + t.Fatalf("POST failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusAccepted { + t.Errorf("expected 202, got %d", resp.StatusCode) + } +} + +func TestServer_getWithoutSSEAccept(t *testing.T) { + ts, _ := newTestServer(nil) + defer ts.Close() + + resp, err := http.Get(ts.URL + "/test/topic") + if err != nil { + t.Fatalf("GET failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusNotFound { + t.Errorf("expected 404, got %d", resp.StatusCode) + } +} + +func TestServer_corsHeaders(t *testing.T) { + ts, _ := newTestServer(nil) + defer ts.Close() + + resp, err := http.Post(ts.URL+"/test", "application/json", strings.NewReader(`{}`)) + if err != nil { + t.Fatalf("POST failed: %v", err) + } + defer resp.Body.Close() + + if resp.Header.Get("Access-Control-Allow-Origin") != "*" { + t.Error("missing CORS Allow-Origin header") + } +} + +func TestServer_optionsPreflight(t *testing.T) { + ts, _ := newTestServer(nil) + defer ts.Close() + + req, _ := http.NewRequest("OPTIONS", ts.URL+"/test", nil) + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("OPTIONS failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusNoContent { + t.Errorf("expected 204, got %d", resp.StatusCode) + } + if resp.Header.Get("Access-Control-Allow-Methods") == "" { + t.Error("missing CORS Allow-Methods header") + } +} + +func TestServer_methodNotAllowed(t *testing.T) { + ts, _ := newTestServer(nil) + defer ts.Close() + + req, _ := http.NewRequest("DELETE", ts.URL+"/test", nil) + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("DELETE failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusMethodNotAllowed { + t.Errorf("expected 405, got %d", resp.StatusCode) + } +} + +func TestServer_postWithBadVerifierConfig(t *testing.T) { + cfg := &Configuration{ + Paths: map[string]PathConfiguration{ + "bad/path": { + Verify: "unknown-method", + Secret: "secret", + SignatureHeader: "X-Signature", + }, + }, + } + ts, _ := newTestServer(cfg) + defer ts.Close() + + req, _ := http.NewRequest("POST", ts.URL+"/bad/path", strings.NewReader(`{}`)) + req.Header.Set("Content-Type", "application/json") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("POST failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusInternalServerError { + t.Errorf("expected 500, got %d", resp.StatusCode) + } +} + +func TestServer_postInvalidJSON(t *testing.T) { + ts, _ := newTestServer(nil) + defer ts.Close() + + resp, err := http.Post(ts.URL+"/test", "application/json", strings.NewReader(`{not json`)) + if err != nil { + t.Fatalf("POST failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusBadRequest { + t.Errorf("expected 400, got %d", resp.StatusCode) + } +} + +func TestServer_postWithMissingSignatureOnSecuredPath(t *testing.T) { + cfg := &Configuration{ + Paths: map[string]PathConfiguration{ + "secure/path": { + Verify: "hmac-sha256", + Secret: "test-secret", + SignatureHeader: "X-Hub-Signature-256", + }, + }, + } + ts, _ := newTestServer(cfg) + defer ts.Close() + + resp, err := http.Post(ts.URL+"/secure/path", "application/json", strings.NewReader(`{}`)) + if err != nil { + t.Fatalf("POST failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusForbidden { + t.Errorf("expected 403, got %d", resp.StatusCode) + } +} diff --git a/sse_test.go b/sse_test.go new file mode 100644 index 0000000..a552288 --- /dev/null +++ b/sse_test.go @@ -0,0 +1,317 @@ +package main + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + "time" +) + +func sseSubscribe(ctx context.Context, url string, headers map[string]string) <-chan *Event { + events := make(chan *Event, 10) + go func() { + defer close(events) + req, _ := http.NewRequestWithContext(ctx, "GET", url, nil) + req.Header.Set("Accept", "text/event-stream") + for k, v := range headers { + req.Header.Set(k, v) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + return + } + defer resp.Body.Close() + scanner := bufio.NewScanner(resp.Body) + for scanner.Scan() { + line := scanner.Text() + if strings.HasPrefix(line, "data: ") { + var event Event + json.Unmarshal([]byte(strings.TrimPrefix(line, "data: ")), &event) + select { + case events <- &event: + case <-ctx.Done(): + return + } + } + } + }() + return events +} + +func TestServer_postAndSSEReceive(t *testing.T) { + ts, _ := newTestServer(nil) + defer ts.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + events := sseSubscribe(ctx, ts.URL+"/test/topic", nil) + time.Sleep(50 * time.Millisecond) + + http.Post(ts.URL+"/test/topic", "application/json", strings.NewReader(`{"hello":"world"}`)) + + select { + case event := <-events: + if event.Path != "test/topic" { + t.Errorf("expected path test/topic, got %s", event.Path) + } + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for SSE event") + } +} + +func TestServer_sseWithValidBearerToken(t *testing.T) { + cfg := &Configuration{ + Paths: map[string]PathConfiguration{ + "private/topic": {SubscribeSecret: "my-token"}, + }, + } + ts, _ := newTestServer(cfg) + defer ts.Close() + + req, _ := http.NewRequest("GET", ts.URL+"/private/topic", nil) + req.Header.Set("Accept", "text/event-stream") + req.Header.Set("Authorization", "Bearer my-token") + + client := &http.Client{Timeout: 500 * time.Millisecond} + resp, err := client.Do(req) + if err != nil { + t.Fatalf("GET failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Errorf("expected 200, got %d", resp.StatusCode) + } +} + +func TestServer_sseWithWrongToken(t *testing.T) { + cfg := &Configuration{ + Paths: map[string]PathConfiguration{ + "private/topic": {SubscribeSecret: "my-token"}, + }, + } + ts, _ := newTestServer(cfg) + defer ts.Close() + + req, _ := http.NewRequest("GET", ts.URL+"/private/topic", nil) + req.Header.Set("Accept", "text/event-stream") + req.Header.Set("Authorization", "Bearer wrong-token") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("GET failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusUnauthorized { + t.Errorf("expected 401, got %d", resp.StatusCode) + } +} + +func TestServer_sseToOpenPath(t *testing.T) { + ts, _ := newTestServer(nil) + defer ts.Close() + + req, _ := http.NewRequest("GET", ts.URL+"/open/topic", nil) + req.Header.Set("Accept", "text/event-stream") + + client := &http.Client{Timeout: 500 * time.Millisecond} + resp, err := client.Do(req) + if err != nil { + t.Fatalf("GET failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Errorf("expected 200, got %d", resp.StatusCode) + } + ct := resp.Header.Get("Content-Type") + if ct != "text/event-stream" { + t.Errorf("expected text/event-stream, got %s", ct) + } +} + +func TestServer_prefixSubscription(t *testing.T) { + ts, _ := newTestServer(nil) + defer ts.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + events := sseSubscribe(ctx, ts.URL+"/github.com/chrisguidry", nil) + time.Sleep(50 * time.Millisecond) + + http.Post(ts.URL+"/github.com/chrisguidry/docketeer", "application/json", strings.NewReader(`{"ref":"main"}`)) + + select { + case event := <-events: + if event.Path != "github.com/chrisguidry/docketeer" { + t.Errorf("expected child path, got %s", event.Path) + } + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for SSE event") + } +} + +func TestServer_lastEventIDReplay(t *testing.T) { + ts, broker := newTestServer(nil) + defer ts.Close() + + broker.Publish(&Event{ + ID: "replay-1", + Path: "test/topic", + Payload: map[string]any{"n": 1}, + }) + broker.Publish(&Event{ + ID: "replay-2", + Path: "test/topic", + Payload: map[string]any{"n": 2}, + }) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + events := sseSubscribe(ctx, ts.URL+"/test/topic", map[string]string{ + "Last-Event-ID": "replay-1", + }) + + select { + case event := <-events: + if event.ID != "replay-2" { + t.Errorf("expected replay-2, got %s", event.ID) + } + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for replayed event") + } +} + +func TestServer_filterQueryParam(t *testing.T) { + ts, _ := newTestServer(nil) + defer ts.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + events := sseSubscribe(ctx, ts.URL+"/test/topic?filter=payload.ref:refs/heads/main", nil) + time.Sleep(50 * time.Millisecond) + + http.Post(ts.URL+"/test/topic", "application/json", strings.NewReader(`{"ref":"refs/heads/develop"}`)) + time.Sleep(20 * time.Millisecond) + http.Post(ts.URL+"/test/topic", "application/json", strings.NewReader(`{"ref":"refs/heads/main"}`)) + + select { + case event := <-events: + payload := event.Payload.(map[string]any) + if payload["ref"] != "refs/heads/main" { + t.Errorf("expected filtered event with ref=main, got %v", payload["ref"]) + } + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for filtered event") + } +} + +func TestServer_sseWithMissingAuth(t *testing.T) { + cfg := &Configuration{ + Paths: map[string]PathConfiguration{ + "private/topic": {SubscribeSecret: "my-token"}, + }, + } + ts, _ := newTestServer(cfg) + defer ts.Close() + + req, _ := http.NewRequest("GET", ts.URL+"/private/topic", nil) + req.Header.Set("Accept", "text/event-stream") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("GET failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusUnauthorized { + t.Errorf("expected 401, got %d", resp.StatusCode) + } +} + +func TestServer_sseChannelClosed(t *testing.T) { + broker := NewBroker(100) + var cfgPtr atomic.Pointer[Configuration] + handler := NewServer(broker, &cfgPtr) + ts := httptest.NewServer(handler) + defer ts.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + events := sseSubscribe(ctx, ts.URL+"/test/topic", nil) + time.Sleep(50 * time.Millisecond) + + broker.mu.Lock() + for _, subs := range broker.subscribers { + for _, sub := range subs { + close(sub.ch) + } + } + broker.subscribers = make(map[string][]*subscriber) + broker.mu.Unlock() + + time.Sleep(50 * time.Millisecond) + + select { + case _, ok := <-events: + if ok { + t.Error("expected channel to be closed") + } + case <-time.After(time.Second): + t.Fatal("timed out waiting for SSE to close") + } +} + +func TestServer_nonJSONPayloadBase64(t *testing.T) { + ts, _ := newTestServer(nil) + defer ts.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + done := make(chan string, 1) + go func() { + req, _ := http.NewRequestWithContext(ctx, "GET", ts.URL+"/test/topic", nil) + req.Header.Set("Accept", "text/event-stream") + resp, err := http.DefaultClient.Do(req) + if err != nil { + return + } + defer resp.Body.Close() + scanner := bufio.NewScanner(resp.Body) + for scanner.Scan() { + line := scanner.Text() + if strings.HasPrefix(line, "data: ") { + var envelope struct { + Payload json.RawMessage `json:"payload"` + } + json.Unmarshal([]byte(strings.TrimPrefix(line, "data: ")), &envelope) + var s string + if json.Unmarshal(envelope.Payload, &s) == nil { + done <- s + return + } + } + } + }() + + time.Sleep(50 * time.Millisecond) + + resp, _ := http.Post(ts.URL+"/test/topic", "text/plain", bytes.NewReader([]byte("hello world"))) + resp.Body.Close() + + select { + case payload := <-done: + if payload != "aGVsbG8gd29ybGQ=" { + t.Errorf("expected base64 encoded payload, got %s", payload) + } + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for event") + } +} diff --git a/verify.go b/verify.go new file mode 100644 index 0000000..a03cb92 --- /dev/null +++ b/verify.go @@ -0,0 +1,59 @@ +package main + +import ( + "crypto/hmac" + "crypto/sha1" + "crypto/sha256" + "encoding/hex" + "errors" + "fmt" + "hash" + "net/http" + "strings" +) + +type Verifier interface { + Verify(body []byte, headers http.Header, secret string, signatureHeader string) error +} + +func NewVerifier(method string) (Verifier, error) { + switch method { + case "hmac-sha256": + return &hmacVerifier{prefix: "sha256=", newHash: sha256.New}, nil + case "hmac-sha1": + return &hmacVerifier{prefix: "sha1=", newHash: sha1.New}, nil + default: + return nil, fmt.Errorf("unknown verification method: %s", method) + } +} + +type hmacVerifier struct { + prefix string + newHash func() hash.Hash +} + +func (v *hmacVerifier) Verify(body []byte, headers http.Header, secret string, signatureHeader string) error { + sig := headers.Get(signatureHeader) + if sig == "" { + return errors.New("missing signature header") + } + + sigHex, ok := strings.CutPrefix(sig, v.prefix) + if !ok { + return fmt.Errorf("signature missing expected prefix %q", v.prefix) + } + + sigBytes, err := hex.DecodeString(sigHex) + if err != nil { + return fmt.Errorf("invalid hex in signature: %w", err) + } + + mac := hmac.New(v.newHash, []byte(secret)) + mac.Write(body) + expected := mac.Sum(nil) + + if !hmac.Equal(sigBytes, expected) { + return errors.New("signature mismatch") + } + return nil +} diff --git a/verify_test.go b/verify_test.go new file mode 100644 index 0000000..718c017 --- /dev/null +++ b/verify_test.go @@ -0,0 +1,120 @@ +package main + +import ( + "crypto/hmac" + "crypto/sha1" + "crypto/sha256" + "encoding/hex" + "net/http" + "testing" +) + +func computeHMACSHA256(secret, body string) string { + mac := hmac.New(sha256.New, []byte(secret)) + mac.Write([]byte(body)) + return "sha256=" + hex.EncodeToString(mac.Sum(nil)) +} + +func computeHMACSHA1(secret, body string) string { + mac := hmac.New(sha1.New, []byte(secret)) + mac.Write([]byte(body)) + return "sha1=" + hex.EncodeToString(mac.Sum(nil)) +} + +func TestNewVerifier_hmacSHA256(t *testing.T) { + v, err := NewVerifier("hmac-sha256") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if v == nil { + t.Fatal("expected verifier, got nil") + } +} + +func TestNewVerifier_hmacSHA1(t *testing.T) { + v, err := NewVerifier("hmac-sha1") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if v == nil { + t.Fatal("expected verifier, got nil") + } +} + +func TestNewVerifier_unknown(t *testing.T) { + _, err := NewVerifier("unknown-method") + if err == nil { + t.Fatal("expected error for unknown method") + } +} + +func TestHMACSHA256_validSignature(t *testing.T) { + secret := "test-secret" + body := `{"action":"push"}` + sig := computeHMACSHA256(secret, body) + + v, _ := NewVerifier("hmac-sha256") + headers := http.Header{"X-Hub-Signature-256": {sig}} + err := v.Verify([]byte(body), headers, secret, "X-Hub-Signature-256") + if err != nil { + t.Fatalf("expected valid signature, got error: %v", err) + } +} + +func TestHMACSHA256_invalidSignature(t *testing.T) { + v, _ := NewVerifier("hmac-sha256") + headers := http.Header{"X-Hub-Signature-256": {"sha256=deadbeef"}} + err := v.Verify([]byte("body"), headers, "secret", "X-Hub-Signature-256") + if err == nil { + t.Fatal("expected error for invalid signature") + } +} + +func TestHMACSHA256_missingHeader(t *testing.T) { + v, _ := NewVerifier("hmac-sha256") + headers := http.Header{} + err := v.Verify([]byte("body"), headers, "secret", "X-Hub-Signature-256") + if err == nil { + t.Fatal("expected error for missing signature header") + } +} + +func TestHMACSHA1_validSignature(t *testing.T) { + secret := "test-secret" + body := `{"action":"push"}` + sig := computeHMACSHA1(secret, body) + + v, _ := NewVerifier("hmac-sha1") + headers := http.Header{"X-Hub-Signature": {sig}} + err := v.Verify([]byte(body), headers, secret, "X-Hub-Signature") + if err != nil { + t.Fatalf("expected valid signature, got error: %v", err) + } +} + +func TestHMACSHA256_wrongPrefix(t *testing.T) { + v, _ := NewVerifier("hmac-sha256") + headers := http.Header{"X-Hub-Signature-256": {"sha1=abc123"}} + err := v.Verify([]byte("body"), headers, "secret", "X-Hub-Signature-256") + if err == nil { + t.Fatal("expected error for wrong prefix") + } +} + +func TestHMACSHA256_invalidHex(t *testing.T) { + v, _ := NewVerifier("hmac-sha256") + headers := http.Header{"X-Hub-Signature-256": {"sha256=not-hex!"}} + err := v.Verify([]byte("body"), headers, "secret", "X-Hub-Signature-256") + if err == nil { + t.Fatal("expected error for invalid hex") + } +} + +func TestHMACSHA1_invalidSignature(t *testing.T) { + v, _ := NewVerifier("hmac-sha1") + headers := http.Header{"X-Hub-Signature": {"sha1=deadbeef"}} + err := v.Verify([]byte("body"), headers, "secret", "X-Hub-Signature") + if err == nil { + t.Fatal("expected error for invalid signature") + } +}