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" "git.golem15.com/golem15/summercms/towel" ) type assemblePlugin struct { id string mw map[string]pact.Middleware use []string path string hit *bool } func (p assemblePlugin) ID() string { return p.id } func (p assemblePlugin) Requires() []string { return nil } func (p assemblePlugin) Register(*backpack.App) error { return nil } func (p assemblePlugin) Boot(*backpack.App) error { return nil } func (p assemblePlugin) Middlewares() map[string]pact.Middleware { return p.mw } func (p assemblePlugin) Routes(r pact.Router) error { path := p.path if path == "" { path = "/items" } r.Group("/api", Use(p.use...), func(g pact.Router) { g.Get(path, func(w http.ResponseWriter, r *http.Request) { if p.hit != nil { *p.hit = true } w.WriteHeader(http.StatusOK) }) }) return nil } func TestAssembleMissingMiddlewareFailsBoot(t *testing.T) { p := assemblePlugin{id: "golem15.demo", use: []string{"jwt.auth"}} _, err := Assemble(backpack.New(nil), []party.Plugin{p}) 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 TestUnauthenticatedNamedGuardDoesNotReachHandler(t *testing.T) { hit := false p := assemblePlugin{ id: "golem15.demo", use: []string{"jwt.auth"}, hit: &hit, mw: map[string]pact.Middleware{ "jwt.auth": func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusUnauthorized) _, _ = w.Write([]byte(`{"error":true,"message":"Token not provided"}`)) }) }, }, } h, err := Assemble(backpack.New(nil), []party.Plugin{p}) if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/items", nil)) if rec.Code != http.StatusUnauthorized { t.Fatalf("status = %d", rec.Code) } if hit { t.Fatal("unauthenticated request reached the handler") } } func TestPipelineOrderRecoverCORSLocaleAuthPasswordOrgRateHandler(t *testing.T) { var order []string record := func(name string) pact.Middleware { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { order = append(order, name) next.ServeHTTP(w, r) }) } } 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", record("jwt.auth")); err != nil { t.Fatal(err) } if err := r.RegisterMiddleware("golem15.fonoteka", "inv.must-change-password", record("inv.must-change-password")); err != nil { t.Fatal(err) } r.BindPlugin("golem15.demo") r.Group("/api", Use("jwt.auth", "inv.must-change-password"), func(g pact.Router) { g.Get("/items", func(w http.ResponseWriter, req *http.Request) { loc, ok := towel.Locale(req.Context()) if !ok || loc != "pl" { t.Errorf("locale = %q ok=%t, want pl from Accept-Language", loc, ok) } if _, ok := towel.Organization(req.Context()); !ok { t.Error("org slot must run before the handler") } order = append(order, "handler") w.WriteHeader(http.StatusNoContent) }) }) h, err := r.compile() if err != nil { t.Fatal(err) } t.Run("preflight-before-auth", func(t *testing.T) { order = nil 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 rec.Header().Get("Access-Control-Allow-Origin") != "http://localhost:3000" { t.Fatalf("ACA origin = %q", rec.Header().Get("Access-Control-Allow-Origin")) } if len(order) != 0 { t.Fatalf("named stages ran on preflight: %v", order) } }) t.Run("get-order", func(t *testing.T) { order = nil req := httptest.NewRequest(http.MethodGet, "/api/items", nil) req.Header.Set("Accept-Language", "pl") 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) } want := []string{"jwt.auth", "inv.must-change-password", "handler"} if strings.Join(order, ",") != strings.Join(want, ",") { t.Fatalf("order = %v want %v", order, want) } }) t.Run("panic-still-opaque", func(t *testing.T) { panicRouter := New(nil) panicRouter.Get("/boom", func(http.ResponseWriter, *http.Request) { panic("stack-trace-secret") }) ph, err := panicRouter.compile() if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() ph.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/boom", nil)) if rec.Code != http.StatusInternalServerError { t.Fatalf("status = %d", rec.Code) } body := rec.Body.String() if strings.Contains(body, "stack-trace-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 TestGroupAndPerRouteMiddlewareCompose(t *testing.T) { var order []string record := func(name string) pact.Middleware { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { order = append(order, name) next.ServeHTTP(w, r) }) } } r := New(nil) if err := r.RegisterMiddleware("golem15.user", "jwt.auth", record("jwt.auth")); err != nil { t.Fatal(err) } if err := r.RegisterMiddleware("golem15.demo", "audit", record("audit")); err != nil { t.Fatal(err) } r.BindPlugin("golem15.demo") r.Group("/api", Use("jwt.auth"), func(g pact.Router) { g.Get("/plain", func(w http.ResponseWriter, r *http.Request) { order = append(order, "plain") w.WriteHeader(http.StatusNoContent) }) g.Get("/audited", func(w http.ResponseWriter, r *http.Request) { order = append(order, "audited") w.WriteHeader(http.StatusNoContent) }, "audit") }) h, err := r.compile() if err != nil { t.Fatal(err) } order = nil rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/plain", nil)) if rec.Code != http.StatusNoContent || strings.Join(order, ",") != "jwt.auth,plain" { t.Fatalf("plain = %d %v", rec.Code, order) } order = nil rec = httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/audited", nil)) if rec.Code != http.StatusNoContent || strings.Join(order, ",") != "jwt.auth,audit,audited" { t.Fatalf("audited = %d %v", rec.Code, order) } } func TestServeMuxRejectsWrongMethod(t *testing.T) { r := New(nil) r.Get("/only-get", func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusOK) }) h, err := r.compile() if err != nil { t.Fatal(err) } rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/only-get", nil)) if rec.Code != http.StatusMethodNotAllowed && rec.Code != http.StatusNotFound { t.Fatalf("POST status = %d", rec.Code) } rec = httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/only-get", nil)) if rec.Code != http.StatusOK { t.Fatalf("GET status = %d", rec.Code) } }