package surf import ( "encoding/json" "net/http" "net/http/httptest" "strings" "testing" "git.golem15.com/golem15/summercms/backpack" "git.golem15.com/golem15/summercms/pact" "git.golem15.com/golem15/summercms/party" ) type routePlugin struct { id string mw map[string]pact.Middleware path string use []string } func (p routePlugin) ID() string { return p.id } func (p routePlugin) Requires() []string { return nil } func (p routePlugin) Register(any) error { return nil } func (p routePlugin) Boot(any) error { return nil } func (p routePlugin) Middlewares() map[string]pact.Middleware { return p.mw } func (p routePlugin) Routes(r pact.Router) error { r.Group("/api", Use(p.use...), func(g pact.Router) { g.Get(p.path, func(w http.ResponseWriter, r *http.Request) { w.Header().Set("X-Hit", "1") w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte(`{"ok":true}`)) }) }) return nil } func TestMissingMiddlewareNamesPluginAndName(t *testing.T) { r := New(nil) p := routePlugin{id: "golem15.demo", path: "/items", use: []string{"jwt.auth"}} r.BindPlugin(p.ID()) if err := p.Routes(r); err != nil { t.Fatal(err) } _, err := r.compile() if err == nil || !strings.Contains(err.Error(), "golem15.demo") || !strings.Contains(err.Error(), "jwt.auth") { t.Fatalf("want plugin and middleware in error, got %v", err) } } func TestRawGroupPanicBare500(t *testing.T) { r := New(nil) r.GroupRaw("/oauth", nil, func(g pact.Router) { g.Get("/panic", func(http.ResponseWriter, *http.Request) { panic("secret internals") }) }) r.Get("/panic", func(http.ResponseWriter, *http.Request) { panic("secret internals") }) h, err := r.compile() if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/oauth/panic", nil)) if rec.Code != http.StatusInternalServerError { t.Fatalf("raw status = %d", rec.Code) } if rec.Body.Len() != 0 { t.Fatalf("raw body = %q", rec.Body.String()) } if ct := rec.Header().Get("Content-Type"); ct != "" { t.Fatalf("raw Content-Type = %q", ct) } rec = httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/panic", nil)) if rec.Code != http.StatusInternalServerError { t.Fatalf("house status = %d", rec.Code) } if rec.Header().Get("Content-Type") != "application/json" { t.Fatalf("house Content-Type = %q", rec.Header().Get("Content-Type")) } var payload map[string]any if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil { t.Fatal(err) } if payload["error"] != true || payload["message"] != "Internal server error" { t.Fatalf("payload = %v", payload) } } func TestHouseMiddlewareDuplicateNameFailsBoot(t *testing.T) { identity := func(next http.Handler) http.Handler { return next } named := assemblePlugin{ id: "golem15.one", mw: map[string]pact.Middleware{"shared.mw": identity}, use: []string{}, } house := houseRoutePlugin{ id: "golem15.two", house: map[string]pact.Middleware{ "shared.mw": identity, }, } _, err := BuildRouter(backpack.New(nil), []party.Plugin{named, house}) if err == nil { t.Fatal("want duplicate-name error") } if !strings.Contains(err.Error(), "shared.mw") || !strings.Contains(err.Error(), "golem15.one") { t.Fatalf("want existing duplicate-name error, got %v", err) } } func TestRecoverReturnsOpaqueJSON500(t *testing.T) { r := New(nil) r.Get("/panic", func(http.ResponseWriter, *http.Request) { panic("secret internals") }) h, err := r.compile() if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/panic", nil)) if rec.Code != http.StatusInternalServerError { t.Fatalf("status = %d", rec.Code) } body := rec.Body.String() if strings.Contains(body, "secret") { t.Fatalf("leaked panic: %s", body) } var payload map[string]any if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil { t.Fatal(err) } if payload["error"] != true || payload["message"] != "Internal server error" { t.Fatalf("payload = %v", payload) } } func TestCORSPreflightBypassesNamedAuth(t *testing.T) { called := false r := New(nil) r.corsCfg = CORSConfig{ Paths: []string{"api/*"}, AllowedOrigins: []string{"http://localhost:3000"}, AllowedMethods: []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"}, AllowedHeaders: []string{"Authorization", "Content-Type", "Accept"}, } if err := r.RegisterMiddleware("golem15.user", "jwt.auth", func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { called = true next.ServeHTTP(w, req) }) }); err != nil { t.Fatal(err) } r.BindPlugin("golem15.demo") r.Group("/api", Use("jwt.auth"), func(g pact.Router) { g.Get("/items", func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) }) }) h, err := r.compile() if err != nil { t.Fatal(err) } req := httptest.NewRequest(http.MethodOptions, "/api/items", nil) req.Header.Set("Origin", "http://localhost:3000") rec := httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusNoContent { t.Fatalf("status = %d", rec.Code) } if called { t.Fatal("jwt.auth ran on preflight") } if rec.Header().Get("Access-Control-Allow-Origin") != "http://localhost:3000" { t.Fatalf("ACA origin = %q", rec.Header().Get("Access-Control-Allow-Origin")) } } func TestTypedIDRouteReturns404(t *testing.T) { r := New(nil) r.Get("/items/{id}", func(w http.ResponseWriter, req *http.Request) { id, ok := IntParam(req, "id") if !ok || id != 1 { http.NotFound(w, req) return } w.WriteHeader(http.StatusOK) }) r.Where("id", `[0-9]+`) h, err := r.compile() if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/nope", nil)) if rec.Code != http.StatusNotFound { t.Fatalf("malformed status = %d", rec.Code) } rec = httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/99", nil)) if rec.Code != http.StatusNotFound { t.Fatalf("unknown status = %d", rec.Code) } rec = httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/1", nil)) if rec.Code != http.StatusOK { t.Fatalf("known status = %d", rec.Code) } } func TestWhereInRejectsOutsideEnum(t *testing.T) { r := New(nil) r.Get("/kinds/{kind}", func(w http.ResponseWriter, req *http.Request) { w.WriteHeader(http.StatusOK) }) r.WhereIn("kind", "widget", "gadget") h, err := r.compile() if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/kinds/other", nil)) if rec.Code != http.StatusNotFound { t.Fatalf("status = %d", rec.Code) } rec = httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/kinds/widget", nil)) if rec.Code != http.StatusOK { t.Fatalf("status = %d", rec.Code) } } func TestWhereInvalidRegexFailsCompile(t *testing.T) { r := New(nil) r.Get("/items/{id}", func(http.ResponseWriter, *http.Request) {}) r.Where("id", "[") if _, err := r.compile(); err == nil { t.Fatal("want compile error for invalid regex") } } func TestWhereUnknownParamFailsCompile(t *testing.T) { r := New(nil) r.Get("/items/{id}", func(http.ResponseWriter, *http.Request) {}) r.Where("slug", `[a-z]+`) if _, err := r.compile(); err == nil || !strings.Contains(err.Error(), "slug") { t.Fatal("want unknown param error") } } func TestPostAndGetSamePathAreIndependent(t *testing.T) { r := New(nil) r.Get("/items", func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte("get")) }) r.Post("/items", func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusCreated) _, _ = w.Write([]byte("post")) }) h, err := r.compile() if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items", nil)) if rec.Code != http.StatusOK || rec.Body.String() != "get" { t.Fatalf("GET status=%d body=%q", rec.Code, rec.Body.String()) } rec = httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/items", nil)) if rec.Code != http.StatusCreated || rec.Body.String() != "post" { t.Fatalf("POST status=%d body=%q", rec.Code, rec.Body.String()) } } func TestVerbsOnRouterAndGroup(t *testing.T) { r := New(nil) r.Put("/r", func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte("put")) }) r.Patch("/r", func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte("patch")) }) r.Delete("/r", func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte("delete")) }) r.Group("/g", nil, func(g pact.Router) { g.Get("/x", func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte("gget")) }) g.Post("/x", func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte("gpost")) }) g.Put("/x", func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte("gput")) }) g.Patch("/x", func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte("gpatch")) }) g.Delete("/x", func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte("gdelete")) }) }) h, err := r.compile() if err != nil { t.Fatal(err) } cases := []struct { method, path, want string }{ {http.MethodPut, "/r", "put"}, {http.MethodPatch, "/r", "patch"}, {http.MethodDelete, "/r", "delete"}, {http.MethodGet, "/g/x", "gget"}, {http.MethodPost, "/g/x", "gpost"}, {http.MethodPut, "/g/x", "gput"}, {http.MethodPatch, "/g/x", "gpatch"}, {http.MethodDelete, "/g/x", "gdelete"}, } for _, tc := range cases { rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(tc.method, tc.path, nil)) if rec.Code != http.StatusOK || rec.Body.String() != tc.want { t.Fatalf("%s %s status=%d body=%q", tc.method, tc.path, rec.Code, rec.Body.String()) } } } func TestMiddlewareFactoryReceivesParam(t *testing.T) { r := New(nil) var gotParam string ran := false if err := r.RegisterMiddlewareFactory("golem15.fonoteka", "inv.scope", func(param string) pact.Middleware { gotParam = param return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { ran = true next.ServeHTTP(w, req) }) } }); err != nil { t.Fatal(err) } r.Get("/x", func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusNoContent) }, "inv.scope:write") h, err := r.compile() if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/x", nil)) if rec.Code != http.StatusNoContent { t.Fatalf("status = %d", rec.Code) } if gotParam != "write" { t.Fatalf("param = %q", gotParam) } if !ran { t.Fatal("factory middleware did not run") } } func TestDuplicateMiddlewareFactoryNamesPluginAndName(t *testing.T) { r := New(nil) fn := func(string) pact.Middleware { return func(next http.Handler) http.Handler { return next } } if err := r.RegisterMiddlewareFactory("golem15.fonoteka", "inv.scope", fn); err != nil { t.Fatal(err) } err := r.RegisterMiddlewareFactory("golem15.other", "inv.scope", fn) if err == nil || !strings.Contains(err.Error(), "golem15.fonoteka") || !strings.Contains(err.Error(), "inv.scope") { t.Fatalf("want plugin and factory name in error, got %v", err) } }