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:
@@ -26,3 +26,25 @@ func User(ctx context.Context) (*Principal, bool) {
|
|||||||
u, ok := ctx.Value(userKey{}).(*Principal)
|
u, ok := ctx.Value(userKey{}).(*Principal)
|
||||||
return u, ok && u != nil
|
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
21
bouncer/context_test.go
Normal 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
25
bouncer/guard.go
Normal 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)
|
||||||
|
}
|
||||||
@@ -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.
|
// Verify parses a token with HS256 pinned and a required exp and sub.
|
||||||
func Verify(tokenString, secret string) (string, error) {
|
func Verify(tokenString, secret string) (string, error) {
|
||||||
if strings.TrimSpace(secret) == "" {
|
if strings.TrimSpace(secret) == "" {
|
||||||
|
|||||||
96
bouncer/registry.go
Normal file
96
bouncer/registry.go
Normal 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
235
bouncer/registry_test.go
Normal 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")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -36,10 +36,20 @@ type HasMiddleware interface {
|
|||||||
Middlewares() map[string]Middleware
|
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.
|
// Router is the Laravel-like group builder implemented by surf.
|
||||||
type Router interface {
|
type Router interface {
|
||||||
Group(prefix string, middleware []string, fn func(Router))
|
Group(prefix string, middleware []string, fn func(Router))
|
||||||
Get(path string, handler http.HandlerFunc, middleware ...string)
|
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)
|
Where(param, pattern string)
|
||||||
WhereIn(param string, values ...string)
|
WhereIn(param string, values ...string)
|
||||||
}
|
}
|
||||||
@@ -117,4 +127,4 @@ type OptionalMessage interface {
|
|||||||
//
|
//
|
||||||
// The kernel type-asserts HasConfig (party, before Register), HasCommands
|
// The kernel type-asserts HasConfig (party, before Register), HasCommands
|
||||||
// (generated app main, after Boot), HasMigrations (lagoon migrate), and
|
// (generated app main, after Boot), HasMigrations (lagoon migrate), and
|
||||||
// HasMiddleware/HasRoutes (surf assemble).
|
// HasMiddleware/HasMiddlewareFactories/HasRoutes (surf assemble).
|
||||||
|
|||||||
122
surf/router.go
122
surf/router.go
@@ -21,6 +21,11 @@ type namedMiddleware struct {
|
|||||||
fn pact.Middleware
|
fn pact.Middleware
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type namedMiddlewareFactory struct {
|
||||||
|
pluginID string
|
||||||
|
fn func(param string) pact.Middleware
|
||||||
|
}
|
||||||
|
|
||||||
type route struct {
|
type route struct {
|
||||||
pluginID string
|
pluginID string
|
||||||
method string
|
method string
|
||||||
@@ -36,6 +41,7 @@ type Router struct {
|
|||||||
prefix string
|
prefix string
|
||||||
middleware []string
|
middleware []string
|
||||||
named map[string]namedMiddleware
|
named map[string]namedMiddleware
|
||||||
|
factories map[string]namedMiddlewareFactory
|
||||||
routes []route
|
routes []route
|
||||||
seen map[string]string
|
seen map[string]string
|
||||||
origins []string
|
origins []string
|
||||||
@@ -58,9 +64,10 @@ type Group struct {
|
|||||||
// New returns an empty router.
|
// New returns an empty router.
|
||||||
func New(origins []string) *Router {
|
func New(origins []string) *Router {
|
||||||
return &Router{
|
return &Router{
|
||||||
named: make(map[string]namedMiddleware),
|
named: make(map[string]namedMiddleware),
|
||||||
seen: make(map[string]string),
|
factories: make(map[string]namedMiddlewareFactory),
|
||||||
origins: origins,
|
seen: make(map[string]string),
|
||||||
|
origins: origins,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -79,6 +86,23 @@ func (r *Router) RegisterMiddleware(pluginID, name string, fn pact.Middleware) e
|
|||||||
return nil
|
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.
|
// BindPlugin records the plugin declaring subsequent routes.
|
||||||
func (r *Router) BindPlugin(id string) {
|
func (r *Router) BindPlugin(id string) {
|
||||||
if r != nil {
|
if r != nil {
|
||||||
@@ -104,7 +128,35 @@ func (r *Router) Get(path string, handler http.HandlerFunc, middleware ...string
|
|||||||
if r == nil {
|
if r == nil {
|
||||||
return
|
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)) {
|
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 {
|
if g == nil || g.router == nil {
|
||||||
return
|
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().
|
// 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)
|
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)
|
full := joinPath(prefix, path)
|
||||||
key := "GET " + full
|
key := method + " " + full
|
||||||
if prev, ok := r.seen[key]; ok {
|
if prev, ok := r.seen[key]; ok {
|
||||||
r.compileErr = fmt.Errorf("surf: duplicate route %s registered by %s and %s", key, prev, pluginID)
|
r.compileErr = fmt.Errorf("surf: duplicate route %s registered by %s and %s", key, prev, pluginID)
|
||||||
return
|
return
|
||||||
@@ -196,7 +276,7 @@ func (r *Router) add(pluginID, prefix string, groupMW []string, path string, han
|
|||||||
mw = append(mw, extra...)
|
mw = append(mw, extra...)
|
||||||
r.routes = append(r.routes, route{
|
r.routes = append(r.routes, route{
|
||||||
pluginID: pluginID,
|
pluginID: pluginID,
|
||||||
method: "GET",
|
method: method,
|
||||||
path: full,
|
path: full,
|
||||||
handler: handler,
|
handler: handler,
|
||||||
middleware: mw,
|
middleware: mw,
|
||||||
@@ -213,7 +293,7 @@ func (r *Router) compile() (http.Handler, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
mux.Handle("GET "+rt.path, h)
|
mux.Handle(rt.method+" "+rt.path, h)
|
||||||
}
|
}
|
||||||
return recoverJSON(cors(r.origins, mux)), nil
|
return recoverJSON(cors(r.origins, mux)), nil
|
||||||
}
|
}
|
||||||
@@ -224,11 +304,18 @@ func (r *Router) wrap(rt route) (http.Handler, error) {
|
|||||||
h = orgSlot(h)
|
h = orgSlot(h)
|
||||||
for i := len(rt.middleware) - 1; i >= 0; i-- {
|
for i := len(rt.middleware) - 1; i >= 0; i-- {
|
||||||
name := rt.middleware[i]
|
name := rt.middleware[i]
|
||||||
named, ok := r.named[name]
|
if named, ok := r.named[name]; ok {
|
||||||
if !ok {
|
h = named.fn(h)
|
||||||
return nil, fmt.Errorf("surf: plugin %q references unknown middleware %q", rt.pluginID, name)
|
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)
|
h = locale(h)
|
||||||
return h, nil
|
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 {
|
for _, p := range plugins {
|
||||||
r.BindPlugin(p.ID())
|
r.BindPlugin(p.ID())
|
||||||
if hr, ok := p.(pact.HasRoutes); ok {
|
if hr, ok := p.(pact.HasRoutes); ok {
|
||||||
|
|||||||
@@ -180,3 +180,115 @@ func TestWhereUnknownParamFailsCompile(t *testing.T) {
|
|||||||
t.Fatal("want unknown param error")
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user