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
This commit is contained in:
Jakub Zych
2026-09-19 18:59:12 +02:00
parent ecab09c37a
commit d376b1be2d
9 changed files with 678 additions and 14 deletions

View File

@@ -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
}

21
bouncer/context_test.go Normal file
View File

@@ -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)
}
}

25
bouncer/guard.go Normal file
View File

@@ -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)
}

View File

@@ -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) == "" {

96
bouncer/registry.go Normal file
View File

@@ -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")
}

235
bouncer/registry_test.go Normal file
View File

@@ -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")
}
})
}
}

View File

@@ -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).

View File

@@ -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 {

View File

@@ -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)
}
}