From d376b1be2d85cf1c70bfe514529568519248e4fe Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Sat, 19 Sep 2026 18:59:12 +0200 Subject: [PATCH] feat(06-01): grow router verbs, factories, and guard registry - Add Post/Put/Patch/Delete on pact.Router and surf Router/Group - Resolve name:param middleware via RegisterMiddlewareFactory - Add bouncer.Registry with Guard, CredentialGuard, UnauthorizedWriter - Re-express jwt as NewJWTGuard without changing Middleware bodies --- bouncer/context.go | 22 ++++ bouncer/context_test.go | 21 ++++ bouncer/guard.go | 25 +++++ bouncer/jwt.go | 47 ++++++++ bouncer/registry.go | 96 ++++++++++++++++ bouncer/registry_test.go | 235 +++++++++++++++++++++++++++++++++++++++ pact/capabilities.go | 12 +- surf/router.go | 122 +++++++++++++++++--- surf/router_test.go | 112 +++++++++++++++++++ 9 files changed, 678 insertions(+), 14 deletions(-) create mode 100644 bouncer/context_test.go create mode 100644 bouncer/guard.go create mode 100644 bouncer/registry.go create mode 100644 bouncer/registry_test.go diff --git a/bouncer/context.go b/bouncer/context.go index 2f26a88..fc4d94c 100644 --- a/bouncer/context.go +++ b/bouncer/context.go @@ -26,3 +26,25 @@ func User(ctx context.Context) (*Principal, bool) { u, ok := ctx.Value(userKey{}).(*Principal) return u, ok && u != nil } + +type credentialKey struct{} + +// WithCredential stores the resolved credential (e.g. *ApiToken) on ctx. +func WithCredential(ctx context.Context, cred any) context.Context { + if ctx == nil { + ctx = context.Background() + } + return context.WithValue(ctx, credentialKey{}, cred) +} + +// Credential returns the resolved credential from ctx. +func Credential(ctx context.Context) (any, bool) { + if ctx == nil { + return nil, false + } + c := ctx.Value(credentialKey{}) + if c == nil { + return nil, false + } + return c, true +} diff --git a/bouncer/context_test.go b/bouncer/context_test.go new file mode 100644 index 0000000..af92aff --- /dev/null +++ b/bouncer/context_test.go @@ -0,0 +1,21 @@ +package bouncer + +import "testing" + +func TestContextCredentialRoundTrip(t *testing.T) { + if _, ok := Credential(t.Context()); ok { + t.Fatal("empty context must have no credential") + } + if _, ok := Credential(nil); ok { + t.Fatal("nil context must have no credential") + } + token := &struct{ Name string }{Name: "parity"} + got, ok := Credential(WithCredential(t.Context(), token)) + if !ok || got != token { + t.Fatalf("got %#v ok=%t", got, ok) + } + got, ok = Credential(WithCredential(nil, token)) + if !ok || got != token { + t.Fatalf("nil ctx got %#v ok=%t", got, ok) + } +} diff --git a/bouncer/guard.go b/bouncer/guard.go new file mode 100644 index 0000000..3f5d61e --- /dev/null +++ b/bouncer/guard.go @@ -0,0 +1,25 @@ +package bouncer + +import "net/http" + +// Guard resolves the caller's Principal for r, or an error describing why not. +type Guard interface { + Authenticate(r *http.Request) (*Principal, error) +} + +// CredentialGuard resolves the Principal AND its underlying credential (e.g. +// *models.ApiToken) in one pass -- a DB-backed guard must never verify twice +// per request (RESEARCH.md Pitfall 5: last_used_at must stamp once). +type CredentialGuard interface { + AuthenticateCredential(r *http.Request) (*Principal, any, error) +} + +// UnauthorizedWriter lets a guard write its own failure response. jwtGuard +// implements this (reusing write401's {"error":true,"message":...} shape). +// TokenGuard does NOT implement it: PHP's TokenScope, not ApiTokenGuard, owns +// the {"error":"Invalid token"} 401 body (D-08) -- Registry.Middleware must +// pass an unauthenticated request through untouched when a guard has no +// UnauthorizedWriter, leaving denial to downstream middleware. +type UnauthorizedWriter interface { + WriteUnauthorized(w http.ResponseWriter, err error) +} diff --git a/bouncer/jwt.go b/bouncer/jwt.go index 49f46a5..9643ffa 100644 --- a/bouncer/jwt.go +++ b/bouncer/jwt.go @@ -63,6 +63,53 @@ func Middleware(secret string, users UserProvider) func(http.Handler) http.Handl } } +type jwtGuard struct { + secret string + users UserProvider +} + +var ( + _ Guard = (*jwtGuard)(nil) + _ UnauthorizedWriter = (*jwtGuard)(nil) +) + +// NewJWTGuard adapts the existing bearerToken -> Verify -> users.FindByID +// chain (identical to Middleware's body) into a Guard + UnauthorizedWriter, +// so Registry.Middleware("jwt") is byte-identical to bouncer.Middleware. +func NewJWTGuard(secret string, users UserProvider) Guard { + return &jwtGuard{secret: secret, users: users} +} + +func (g *jwtGuard) Authenticate(r *http.Request) (*Principal, error) { + raw, err := bearerToken(r) + if err != nil { + return nil, err + } + sub, err := Verify(raw, g.secret) + if err != nil { + return nil, err + } + id, err := strconv.ParseUint(sub, 10, 64) + if err != nil || id == 0 { + return nil, errors.New(msgUserNotFound) + } + if g.users == nil { + return nil, errors.New(msgUserNotFound) + } + user, err := g.users.FindByID(r.Context(), uint(id)) + if err != nil { + return nil, errors.New("Authentication error") + } + if user == nil { + return nil, errors.New(msgUserNotFound) + } + return user, nil +} + +func (g *jwtGuard) WriteUnauthorized(w http.ResponseWriter, err error) { + write401(w, err.Error()) +} + // Verify parses a token with HS256 pinned and a required exp and sub. func Verify(tokenString, secret string) (string, error) { if strings.TrimSpace(secret) == "" { diff --git a/bouncer/registry.go b/bouncer/registry.go new file mode 100644 index 0000000..f5c0182 --- /dev/null +++ b/bouncer/registry.go @@ -0,0 +1,96 @@ +package bouncer + +import ( + "errors" + "fmt" + "net/http" +) + +type namedGuard struct { + pluginID string + g any +} + +// Registry stores named Guard / CredentialGuard implementations and derives +// auth middleware from them. +type Registry struct { + guards map[string]namedGuard +} + +// NewRegistry returns an empty named-guard registry. +func NewRegistry() *Registry { + return &Registry{guards: make(map[string]namedGuard)} +} + +// Register stores g under name. g must implement Guard or CredentialGuard. +// Empty name, nil g, a type implementing neither, or a duplicate name all +// fail with a "bouncer: ..." error naming pluginID and name. +func (reg *Registry) Register(pluginID, name string, g any) error { + if reg == nil { + return fmt.Errorf("bouncer: registry is nil") + } + if name == "" || g == nil { + return fmt.Errorf("bouncer: plugin %q registered empty guard %q", pluginID, name) + } + _, isGuard := g.(Guard) + _, isCred := g.(CredentialGuard) + if !isGuard && !isCred { + return fmt.Errorf("bouncer: plugin %q registered guard %q that implements neither Guard nor CredentialGuard", pluginID, name) + } + if existing, ok := reg.guards[name]; ok { + return fmt.Errorf("bouncer: guard %q already registered by %s", name, existing.pluginID) + } + if reg.guards == nil { + reg.guards = make(map[string]namedGuard) + } + reg.guards[name] = namedGuard{pluginID: pluginID, g: g} + return nil +} + +// Middleware derives an http middleware from a registered guard. Unknown +// names fail (fail boot, mirrors surf.RegisterMiddleware's contract). +// On Authenticate/AuthenticateCredential success: WithUser (+WithCredential +// if a credential was returned) then next.ServeHTTP. +// On failure: if the guard implements UnauthorizedWriter, it writes the +// response and the chain stops; otherwise next.ServeHTTP runs unauthenticated. +func (reg *Registry) Middleware(name string) (func(http.Handler) http.Handler, error) { + if reg == nil { + return nil, fmt.Errorf("bouncer: registry is nil") + } + ng, ok := reg.guards[name] + if !ok { + return nil, fmt.Errorf("bouncer: unknown guard %q", name) + } + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + principal, cred, err := authenticate(ng.g, r) + if err != nil || principal == nil { + if wtr, ok := ng.g.(UnauthorizedWriter); ok { + if err == nil { + err = errors.New("unauthenticated") + } + wtr.WriteUnauthorized(w, err) + return + } + next.ServeHTTP(w, r) + return + } + ctx := WithUser(r.Context(), principal) + if cred != nil { + ctx = WithCredential(ctx, cred) + } + next.ServeHTTP(w, r.WithContext(ctx)) + }) + }, nil +} + +func authenticate(g any, r *http.Request) (*Principal, any, error) { + if cg, ok := g.(CredentialGuard); ok { + return cg.AuthenticateCredential(r) + } + if gd, ok := g.(Guard); ok { + p, err := gd.Authenticate(r) + return p, nil, err + } + return nil, nil, fmt.Errorf("bouncer: guard implements neither Guard nor CredentialGuard") +} diff --git a/bouncer/registry_test.go b/bouncer/registry_test.go new file mode 100644 index 0000000..8f27daf --- /dev/null +++ b/bouncer/registry_test.go @@ -0,0 +1,235 @@ +package bouncer + +import ( + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +type writerGuard struct { + principal *Principal + err error + wrote *bool +} + +func (g writerGuard) Authenticate(*http.Request) (*Principal, error) { + return g.principal, g.err +} + +func (g writerGuard) WriteUnauthorized(w http.ResponseWriter, err error) { + if g.wrote != nil { + *g.wrote = true + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"from":"writer"}`)) +} + +type credOnlyGuard struct { + principal *Principal + cred any + err error +} + +func (g credOnlyGuard) AuthenticateCredential(*http.Request) (*Principal, any, error) { + return g.principal, g.cred, g.err +} + +type notAGuard struct{} + +func TestDuplicateGuardNameFailsWithPluginAndName(t *testing.T) { + reg := NewRegistry() + if err := reg.Register("golem15.user", "jwt", writerGuard{}); err != nil { + t.Fatal(err) + } + err := reg.Register("golem15.fonoteka", "jwt", writerGuard{}) + if err == nil || !strings.Contains(err.Error(), "golem15.user") || !strings.Contains(err.Error(), "jwt") { + t.Fatalf("want plugin and guard name in error, got %v", err) + } +} + +func TestUnknownGuardNameFails(t *testing.T) { + reg := NewRegistry() + _, err := reg.Middleware("jwt") + if err == nil || !strings.Contains(err.Error(), "jwt") { + t.Fatalf("want guard name in error, got %v", err) + } +} + +func TestRegisterNeitherInterfaceNamesPluginAndName(t *testing.T) { + reg := NewRegistry() + err := reg.Register("golem15.demo", "oops", notAGuard{}) + if err == nil || !strings.Contains(err.Error(), "golem15.demo") || !strings.Contains(err.Error(), "oops") { + t.Fatalf("want plugin and guard name in error, got %v", err) + } +} + +func TestWriterGuardFailureWritesOwnResponse(t *testing.T) { + wrote := false + reg := NewRegistry() + if err := reg.Register("golem15.user", "jwt", writerGuard{ + err: errors.New("nope"), + wrote: &wrote, + }); err != nil { + t.Fatal(err) + } + mw, err := reg.Middleware("jwt") + if err != nil { + t.Fatal(err) + } + called := false + h := mw(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + called = true + })) + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil)) + if !wrote { + t.Fatal("WriteUnauthorized was not called") + } + if called { + t.Fatal("next ran on guard failure") + } + if rec.Code != http.StatusUnauthorized || rec.Body.String() != `{"from":"writer"}` { + t.Fatalf("status=%d body=%q", rec.Code, rec.Body.String()) + } +} + +func TestCredentialGuardSoftFailAndSuccess(t *testing.T) { + reg := NewRegistry() + failing := credOnlyGuard{err: errors.New("bad token")} + if err := reg.Register("golem15.fonoteka", "inv_token", failing); err != nil { + t.Fatal(err) + } + mw, err := reg.Middleware("inv_token") + if err != nil { + t.Fatal(err) + } + + t.Run("failure-passes-through", func(t *testing.T) { + called := false + h := mw(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + called = true + if _, ok := User(r.Context()); ok { + t.Fatal("principal attached on failure") + } + if _, ok := Credential(r.Context()); ok { + t.Fatal("credential attached on failure") + } + w.WriteHeader(http.StatusNoContent) + })) + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil)) + if !called { + t.Fatal("next did not run on CredentialGuard failure") + } + if rec.Code != http.StatusNoContent { + t.Fatalf("status = %d", rec.Code) + } + }) + + okReg := NewRegistry() + token := &struct{ ID uint }{ID: 9} + principal := &Principal{ID: 3} + if err := okReg.Register("golem15.fonoteka", "inv_token", credOnlyGuard{ + principal: principal, + cred: token, + }); err != nil { + t.Fatal(err) + } + okMW, err := okReg.Middleware("inv_token") + if err != nil { + t.Fatal(err) + } + t.Run("success-attaches-both", func(t *testing.T) { + h := okMW(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + got, ok := User(r.Context()) + if !ok || got != principal { + t.Fatalf("user = %+v ok=%t", got, ok) + } + cred, ok := Credential(r.Context()) + if !ok || cred != token { + t.Fatalf("cred = %#v ok=%t", cred, ok) + } + w.WriteHeader(http.StatusNoContent) + })) + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil)) + if rec.Code != http.StatusNoContent { + t.Fatalf("status = %d", rec.Code) + } + }) +} + +func TestJWTGuardRegistryMatchesMiddleware(t *testing.T) { + users := memUsers{byID: map[uint]*Principal{1: {ID: 1}}} + direct := Middleware(secret, users)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("X-Hit", "1") + w.WriteHeader(http.StatusNoContent) + })) + reg := NewRegistry() + if err := reg.Register("golem15.user", "jwt", NewJWTGuard(secret, users)); err != nil { + t.Fatal(err) + } + viaReg, err := reg.Middleware("jwt") + if err != nil { + t.Fatal(err) + } + hReg := viaReg(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("X-Hit", "1") + w.WriteHeader(http.StatusNoContent) + })) + + expired := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ + "sub": "1", + "exp": time.Now().Add(-time.Hour).Unix(), + }, []byte(secret)) + badSig := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ + "sub": "1", + "exp": time.Now().Add(time.Hour).Unix(), + }, []byte("other-secret")) + unknown := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ + "sub": "99", + "exp": time.Now().Add(time.Hour).Unix(), + }, []byte(secret)) + + cases := []struct { + name string + header string + }{ + {name: "missing"}, + {name: "expired", header: "Bearer " + expired}, + {name: "bad-signature", header: "Bearer " + badSig}, + {name: "unknown-user", header: "Bearer " + unknown}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + req1 := httptest.NewRequest(http.MethodGet, "/", nil) + req2 := httptest.NewRequest(http.MethodGet, "/", nil) + if tc.header != "" { + req1.Header.Set("Authorization", tc.header) + req2.Header.Set("Authorization", tc.header) + } + rec1 := httptest.NewRecorder() + rec2 := httptest.NewRecorder() + direct.ServeHTTP(rec1, req1) + hReg.ServeHTTP(rec2, req2) + if rec1.Code != rec2.Code { + t.Fatalf("status direct=%d registry=%d", rec1.Code, rec2.Code) + } + if rec1.Body.String() != rec2.Body.String() { + t.Fatalf("body direct=%q registry=%q", rec1.Body.String(), rec2.Body.String()) + } + if rec1.Header().Get("Content-Type") != rec2.Header().Get("Content-Type") { + t.Fatalf("content-type direct=%q registry=%q", rec1.Header().Get("Content-Type"), rec2.Header().Get("Content-Type")) + } + if rec1.Header().Get("X-Hit") != "" || rec2.Header().Get("X-Hit") != "" { + t.Fatal("handler ran") + } + }) + } +} diff --git a/pact/capabilities.go b/pact/capabilities.go index 8820257..aa2d528 100644 --- a/pact/capabilities.go +++ b/pact/capabilities.go @@ -36,10 +36,20 @@ type HasMiddleware interface { Middlewares() map[string]Middleware } +// HasMiddlewareFactories is implemented by plugins that register +// parameterized middleware ("name:param"), resolved at wrap time. +type HasMiddlewareFactories interface { + MiddlewareFactories() map[string]func(param string) Middleware +} + // Router is the Laravel-like group builder implemented by surf. type Router interface { Group(prefix string, middleware []string, fn func(Router)) Get(path string, handler http.HandlerFunc, middleware ...string) + Post(path string, handler http.HandlerFunc, middleware ...string) + Put(path string, handler http.HandlerFunc, middleware ...string) + Patch(path string, handler http.HandlerFunc, middleware ...string) + Delete(path string, handler http.HandlerFunc, middleware ...string) Where(param, pattern string) WhereIn(param string, values ...string) } @@ -117,4 +127,4 @@ type OptionalMessage interface { // // The kernel type-asserts HasConfig (party, before Register), HasCommands // (generated app main, after Boot), HasMigrations (lagoon migrate), and -// HasMiddleware/HasRoutes (surf assemble). +// HasMiddleware/HasMiddlewareFactories/HasRoutes (surf assemble). diff --git a/surf/router.go b/surf/router.go index 2f04741..8af14ae 100644 --- a/surf/router.go +++ b/surf/router.go @@ -21,6 +21,11 @@ type namedMiddleware struct { fn pact.Middleware } +type namedMiddlewareFactory struct { + pluginID string + fn func(param string) pact.Middleware +} + type route struct { pluginID string method string @@ -36,6 +41,7 @@ type Router struct { prefix string middleware []string named map[string]namedMiddleware + factories map[string]namedMiddlewareFactory routes []route seen map[string]string origins []string @@ -58,9 +64,10 @@ type Group struct { // New returns an empty router. func New(origins []string) *Router { return &Router{ - named: make(map[string]namedMiddleware), - seen: make(map[string]string), - origins: origins, + named: make(map[string]namedMiddleware), + factories: make(map[string]namedMiddlewareFactory), + seen: make(map[string]string), + origins: origins, } } @@ -79,6 +86,23 @@ func (r *Router) RegisterMiddleware(pluginID, name string, fn pact.Middleware) e return nil } +// RegisterMiddlewareFactory stores a parameterized middleware builder. +// At wrap time a name not found in r.named is split on the first ':' and +// the base is looked up here. Duplicate factory names fail. +func (r *Router) RegisterMiddlewareFactory(pluginID, name string, fn func(param string) pact.Middleware) error { + if r == nil { + return fmt.Errorf("surf: router is nil") + } + if name == "" || fn == nil { + return fmt.Errorf("surf: plugin %q registered empty middleware factory", pluginID) + } + if existing, ok := r.factories[name]; ok { + return fmt.Errorf("surf: middleware factory %q already registered by %s", name, existing.pluginID) + } + r.factories[name] = namedMiddlewareFactory{pluginID: pluginID, fn: fn} + return nil +} + // BindPlugin records the plugin declaring subsequent routes. func (r *Router) BindPlugin(id string) { if r != nil { @@ -104,7 +128,35 @@ func (r *Router) Get(path string, handler http.HandlerFunc, middleware ...string if r == nil { return } - r.add(r.pluginID, r.prefix, r.middleware, path, handler, middleware) + r.add(r.pluginID, r.prefix, r.middleware, http.MethodGet, path, handler, middleware) +} + +func (r *Router) Post(path string, handler http.HandlerFunc, middleware ...string) { + if r == nil { + return + } + r.add(r.pluginID, r.prefix, r.middleware, http.MethodPost, path, handler, middleware) +} + +func (r *Router) Put(path string, handler http.HandlerFunc, middleware ...string) { + if r == nil { + return + } + r.add(r.pluginID, r.prefix, r.middleware, http.MethodPut, path, handler, middleware) +} + +func (r *Router) Patch(path string, handler http.HandlerFunc, middleware ...string) { + if r == nil { + return + } + r.add(r.pluginID, r.prefix, r.middleware, http.MethodPatch, path, handler, middleware) +} + +func (r *Router) Delete(path string, handler http.HandlerFunc, middleware ...string) { + if r == nil { + return + } + r.add(r.pluginID, r.prefix, r.middleware, http.MethodDelete, path, handler, middleware) } func (g *Group) Group(prefix string, middleware []string, fn func(pact.Router)) { @@ -125,7 +177,35 @@ func (g *Group) Get(path string, handler http.HandlerFunc, middleware ...string) if g == nil || g.router == nil { return } - g.router.add(g.pluginID, g.prefix, g.middleware, path, handler, middleware) + g.router.add(g.pluginID, g.prefix, g.middleware, http.MethodGet, path, handler, middleware) +} + +func (g *Group) Post(path string, handler http.HandlerFunc, middleware ...string) { + if g == nil || g.router == nil { + return + } + g.router.add(g.pluginID, g.prefix, g.middleware, http.MethodPost, path, handler, middleware) +} + +func (g *Group) Put(path string, handler http.HandlerFunc, middleware ...string) { + if g == nil || g.router == nil { + return + } + g.router.add(g.pluginID, g.prefix, g.middleware, http.MethodPut, path, handler, middleware) +} + +func (g *Group) Patch(path string, handler http.HandlerFunc, middleware ...string) { + if g == nil || g.router == nil { + return + } + g.router.add(g.pluginID, g.prefix, g.middleware, http.MethodPatch, path, handler, middleware) +} + +func (g *Group) Delete(path string, handler http.HandlerFunc, middleware ...string) { + if g == nil || g.router == nil { + return + } + g.router.add(g.pluginID, g.prefix, g.middleware, http.MethodDelete, path, handler, middleware) } // Where attaches a compiled regex constraint to the last route, matching PHP ->where(). @@ -184,9 +264,9 @@ func (r *Router) addConstraint(c Constraint) { last.constraints = append(last.constraints, c) } -func (r *Router) add(pluginID, prefix string, groupMW []string, path string, handler http.HandlerFunc, extra []string) { +func (r *Router) add(pluginID, prefix string, groupMW []string, method, path string, handler http.HandlerFunc, extra []string) { full := joinPath(prefix, path) - key := "GET " + full + key := method + " " + full if prev, ok := r.seen[key]; ok { r.compileErr = fmt.Errorf("surf: duplicate route %s registered by %s and %s", key, prev, pluginID) return @@ -196,7 +276,7 @@ func (r *Router) add(pluginID, prefix string, groupMW []string, path string, han mw = append(mw, extra...) r.routes = append(r.routes, route{ pluginID: pluginID, - method: "GET", + method: method, path: full, handler: handler, middleware: mw, @@ -213,7 +293,7 @@ func (r *Router) compile() (http.Handler, error) { if err != nil { return nil, err } - mux.Handle("GET "+rt.path, h) + mux.Handle(rt.method+" "+rt.path, h) } return recoverJSON(cors(r.origins, mux)), nil } @@ -224,11 +304,18 @@ func (r *Router) wrap(rt route) (http.Handler, error) { h = orgSlot(h) for i := len(rt.middleware) - 1; i >= 0; i-- { name := rt.middleware[i] - named, ok := r.named[name] - if !ok { - return nil, fmt.Errorf("surf: plugin %q references unknown middleware %q", rt.pluginID, name) + if named, ok := r.named[name]; ok { + h = named.fn(h) + continue } - h = named.fn(h) + base, param, hasParam := strings.Cut(name, ":") + if hasParam { + if factory, ok := r.factories[base]; ok { + h = factory.fn(param)(h) + continue + } + } + return nil, fmt.Errorf("surf: plugin %q references unknown middleware %q", rt.pluginID, name) } h = locale(h) return h, nil @@ -246,6 +333,15 @@ func Assemble(app *backpack.App, plugins []party.Plugin) (http.Handler, error) { } } } + for _, p := range plugins { + if hf, ok := p.(pact.HasMiddlewareFactories); ok { + for name, fn := range hf.MiddlewareFactories() { + if err := r.RegisterMiddlewareFactory(p.ID(), name, fn); err != nil { + return nil, err + } + } + } + } for _, p := range plugins { r.BindPlugin(p.ID()) if hr, ok := p.(pact.HasRoutes); ok { diff --git a/surf/router_test.go b/surf/router_test.go index 494391a..ebdb595 100644 --- a/surf/router_test.go +++ b/surf/router_test.go @@ -180,3 +180,115 @@ func TestWhereUnknownParamFailsCompile(t *testing.T) { 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) + } +}