diff --git a/server.go b/server.go index 6f92706..044cc4b 100644 --- a/server.go +++ b/server.go @@ -42,6 +42,10 @@ func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { } } +func normalizePath(raw string) string { + return strings.TrimRight(strings.TrimPrefix(raw, "/"), "/") +} + func setCORSHeaders(w http.ResponseWriter) { w.Header().Set("Access-Control-Allow-Origin", "*") w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") @@ -49,7 +53,7 @@ func setCORSHeaders(w http.ResponseWriter) { } func (s *Server) handlePost(w http.ResponseWriter, r *http.Request) { - path := strings.TrimPrefix(r.URL.Path, "/") + path := normalizePath(r.URL.Path) body, _ := io.ReadAll(r.Body) @@ -91,7 +95,7 @@ func (s *Server) handlePost(w http.ResponseWriter, r *http.Request) { } func (s *Server) handleSSE(w http.ResponseWriter, r *http.Request) { - path := strings.TrimPrefix(r.URL.Path, "/") + path := normalizePath(r.URL.Path) cfg := s.config.Load() if secret := cfg.LookupSubscribeSecret(path); secret != "" { diff --git a/server_test.go b/server_test.go index 84d9ac8..59e3a38 100644 --- a/server_test.go +++ b/server_test.go @@ -9,6 +9,7 @@ import ( "strings" "sync/atomic" "testing" + "time" ) func newTestServer(cfg *Configuration) (*httptest.Server, *Broker) { @@ -35,6 +36,25 @@ func TestServer_postPublishesEvent(t *testing.T) { } } +func TestServer_postTrailingSlashNormalized(t *testing.T) { + ts, broker := newTestServer(nil) + defer ts.Close() + + ch, unsub := broker.Subscribe("test/topic", "") + defer unsub() + + http.Post(ts.URL+"/test/topic/", "application/json", strings.NewReader(`{}`)) + + select { + case event := <-ch: + if event.Path != "test/topic" { + t.Errorf("expected normalized path test/topic, got %s", event.Path) + } + case <-time.After(time.Second): + t.Fatal("timed out: POST with trailing slash should deliver to normalized path") + } +} + func TestServer_postWithValidHMAC(t *testing.T) { cfg := &Configuration{ Paths: map[string]PathConfiguration{ diff --git a/sse_test.go b/sse_test.go index a552288..fd96d9b 100644 --- a/sse_test.go +++ b/sse_test.go @@ -156,6 +156,28 @@ func TestServer_prefixSubscription(t *testing.T) { } } +func TestServer_prefixSubscriptionWithTrailingSlash(t *testing.T) { + ts, _ := newTestServer(nil) + defer ts.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + events := sseSubscribe(ctx, ts.URL+"/test/", 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 test/topic, got %s", event.Path) + } + case <-time.After(2 * time.Second): + t.Fatal("timed out: trailing slash subscribe should receive child events") + } +} + func TestServer_lastEventIDReplay(t *testing.T) { ts, broker := newTestServer(nil) defer ts.Close()