diff --git a/cmd/librenotes/serve.go b/cmd/librenotes/serve.go index 6b5cc62..aef66be 100644 --- a/cmd/librenotes/serve.go +++ b/cmd/librenotes/serve.go @@ -140,6 +140,11 @@ func runServe(args []string) error { apiHandler := api.Routes() root.Handle("/auth/", apiHandler) root.Handle("/api/", apiHandler) + // /healthz is mounted directly so the static fall-through handler + // below does not shadow it. The api.Routes() mux registers it for + // completeness but with apiHandler attached only at /auth/ and + // /api/, the route is otherwise unreachable from the public origin. + root.Handle("/healthz", apiHandler) pub, err := fs.Sub(publicFS, "web/public") if err != nil { diff --git a/cmd/librenotes/serve_test.go b/cmd/librenotes/serve_test.go new file mode 100644 index 0000000..50b2dce --- /dev/null +++ b/cmd/librenotes/serve_test.go @@ -0,0 +1,62 @@ +package main + +import ( + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +// TestServeMounts ensures the public origin exposes /healthz, /auth/*, +// and /api/* (auth-protected). It uses the same routing topology as +// runServe but skips the embedded file system, since the static +// fall-through is what shadowed /healthz before this test existed. +func TestServeMounts(t *testing.T) { + root := http.NewServeMux() + apiHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/healthz": + _, _ = io.WriteString(w, `{"status":"ok"}`) + case "/auth/login": + w.WriteHeader(http.StatusMethodNotAllowed) + case "/api/whoami": + w.WriteHeader(http.StatusUnauthorized) + default: + http.NotFound(w, r) + } + }) + root.Handle("/auth/", apiHandler) + root.Handle("/api/", apiHandler) + root.Handle("/healthz", apiHandler) + root.Handle("/", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.NotFound(w, r) + })) + + srv := httptest.NewServer(root) + defer srv.Close() + + cases := []struct { + path string + want int + }{ + {"/healthz", http.StatusOK}, + {"/auth/login", http.StatusMethodNotAllowed}, + {"/api/whoami", http.StatusUnauthorized}, + {"/does-not-exist", http.StatusNotFound}, + } + for _, tc := range cases { + resp, err := http.Get(srv.URL + tc.path) + if err != nil { + t.Fatalf("GET %s: %v", tc.path, err) + } + if resp.StatusCode != tc.want { + t.Errorf("%s: got %d, want %d", tc.path, resp.StatusCode, tc.want) + } + body, _ := io.ReadAll(resp.Body) + resp.Body.Close() + if tc.path == "/healthz" && !strings.Contains(string(body), `"status":"ok"`) { + t.Errorf("/healthz body = %q, want it to contain status:ok", body) + } + } +}