114 lines
3.4 KiB
Go
114 lines
3.4 KiB
Go
package bouncer
|
|
|
|
import (
|
|
"errors"
|
|
"fmt"
|
|
"net/http"
|
|
"reflect"
|
|
)
|
|
|
|
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)
|
|
}
|
|
switch rv := reflect.ValueOf(g); rv.Kind() {
|
|
case reflect.Pointer, reflect.Map, reflect.Slice, reflect.Func, reflect.Chan, reflect.Interface:
|
|
if rv.IsNil() {
|
|
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
|
|
}
|
|
|
|
// Owner returns the plugin ID that registered the guard name. The second
|
|
// result is false when no guard has that name.
|
|
func (reg *Registry) Owner(name string) (string, bool) {
|
|
if reg == nil {
|
|
return "", false
|
|
}
|
|
ng, ok := reg.guards[name]
|
|
return ng.pluginID, ok
|
|
}
|
|
|
|
// 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")
|
|
}
|