package surf import ( "bytes" "fmt" "math" "net/http" "strings" "time" "git.golem15.com/golem15/summercms/modules/backpack" "git.golem15.com/golem15/summercms/modules/cabana" "git.golem15.com/golem15/summercms/modules/pact" "git.golem15.com/golem15/summercms/modules/party" "git.golem15.com/golem15/summercms/modules/towel" "git.golem15.com/golem15/summercms/modules/wire" ) // Use names a middleware list for Group, matching PHP ->middleware(). func Use(names ...string) []string { return names } type namedMiddleware struct { pluginID string fn pact.Middleware houseTagged bool } type namedMiddlewareFactory struct { pluginID string fn func(param string) pact.Middleware houseTagged bool } type route struct { pluginID string method string path string handler http.Handler middleware []string constraints []Constraint raw bool } // Router compiles group declarations onto net/http ServeMux. type Router struct { pluginID string prefix string middleware []string named map[string]namedMiddleware factories map[string]namedMiddlewareFactory routes []route seen map[string]string origins []string compileErr error limiter *FixedWindowLimiter corsCfg CORSConfig defaultBytes int64 uploadBytes int64 built map[string]pact.Middleware } var ( _ pact.Router = (*Router)(nil) _ pact.Router = (*Group)(nil) ) // Group is a prefixed route collection. type Group struct { router *Router pluginID string prefix string middleware []string raw bool } // New returns an empty router. func New(origins []string) *Router { return &Router{ named: make(map[string]namedMiddleware), factories: make(map[string]namedMiddlewareFactory), seen: make(map[string]string), built: make(map[string]pact.Middleware), origins: origins, } } // RegisterMiddleware stores a named wrapper. Duplicate names fail. func (r *Router) RegisterMiddleware(pluginID, name string, fn pact.Middleware) error { return r.registerNamed(pluginID, name, false, fn) } // RegisterHouseMiddleware stores a named wrapper tagged as house-envelope/error // handling. Raw groups refuse these names at wrap/Assemble time. Called only // from BuildRouter's HasHouseMiddleware plugin loop. func (r *Router) RegisterHouseMiddleware(pluginID, name string, fn pact.Middleware) error { return r.registerNamed(pluginID, name, true, fn) } func (r *Router) registerNamed(pluginID, name string, tagged bool, fn 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", pluginID) } if existing, ok := r.named[name]; ok { return fmt.Errorf("surf: middleware %q already registered by %s", name, existing.pluginID) } r.named[name] = namedMiddleware{pluginID: pluginID, fn: fn, houseTagged: tagged} 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 { return r.registerNamedFactory(pluginID, name, false, fn) } // RegisterHouseMiddlewareFactory stores a parameterized house-envelope/error // factory. No plugin-facing capability wires this yet; it exists as // infrastructure for a future house-tagged parameterized consumer. func (r *Router) RegisterHouseMiddlewareFactory(pluginID, name string, fn func(param string) pact.Middleware) error { return r.registerNamedFactory(pluginID, name, true, fn) } func (r *Router) registerNamedFactory(pluginID, name string, tagged bool, 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, houseTagged: tagged} return nil } // BindPlugin records the plugin declaring subsequent routes. func (r *Router) BindPlugin(id string) { if r != nil { r.pluginID = id } } func (r *Router) Group(prefix string, middleware []string, fn func(pact.Router)) { r.openGroup(prefix, middleware, false, fn) } func (r *Router) GroupRaw(prefix string, middleware []string, fn func(pact.Router)) { r.openGroup(prefix, middleware, true, fn) } func (r *Router) openGroup(prefix string, middleware []string, raw bool, fn func(pact.Router)) { if r == nil || fn == nil { return } g := &Group{ router: r, pluginID: r.pluginID, prefix: joinPath(r.prefix, prefix), middleware: append([]string{}, r.middleware...), raw: raw, } g.middleware = append(g.middleware, middleware...) fn(g) } func (r *Router) Get(path string, handler http.HandlerFunc, middleware ...string) { if r == nil { return } r.add(r.pluginID, r.prefix, r.middleware, http.MethodGet, path, handler, middleware, false) } 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, false) } 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, false) } 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, false) } 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, false) } func (g *Group) Group(prefix string, middleware []string, fn func(pact.Router)) { g.openGroup(prefix, middleware, g.raw, fn) } func (g *Group) GroupRaw(prefix string, middleware []string, fn func(pact.Router)) { g.openGroup(prefix, middleware, true, fn) } func (g *Group) openGroup(prefix string, middleware []string, raw bool, fn func(pact.Router)) { if g == nil || g.router == nil || fn == nil { return } next := &Group{ router: g.router, pluginID: g.pluginID, prefix: joinPath(g.prefix, prefix), middleware: append([]string{}, g.middleware...), raw: raw, } next.middleware = append(next.middleware, middleware...) fn(next) } 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, http.MethodGet, path, handler, middleware, g.raw) } 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, g.raw) } 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, g.raw) } 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, g.raw) } 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, g.raw) } // Where attaches a compiled regex constraint to the last route, matching PHP ->where(). func (r *Router) Where(param, pattern string) { if r == nil { return } c, err := Regex(param, pattern) if err != nil { r.compileErr = err return } r.addConstraint(c) } func (g *Group) Where(param, pattern string) { if g == nil || g.router == nil { return } g.router.Where(param, pattern) } // WhereIn attaches an allow-listed enum constraint to the last route. func (r *Router) WhereIn(param string, values ...string) { if r == nil { return } c, err := Enum(param, values...) if err != nil { r.compileErr = err return } r.addConstraint(c) } func (g *Group) WhereIn(param string, values ...string) { if g == nil || g.router == nil { return } g.router.WhereIn(param, values...) } func (r *Router) addConstraint(c Constraint) { if r.compileErr != nil { return } if len(r.routes) == 0 { r.compileErr = fmt.Errorf("surf: Where(%q) with no route", c.param) return } last := &r.routes[len(r.routes)-1] if !pathHasParam(last.path, c.param) { r.compileErr = fmt.Errorf("surf: Where(%q) is not a path parameter of %s", c.param, last.path) return } last.constraints = append(last.constraints, c) } func (r *Router) add(pluginID, prefix string, groupMW []string, method, path string, handler http.HandlerFunc, extra []string, raw bool) { full := joinPath(prefix, path) 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 } r.seen[key] = pluginID mw := append([]string{}, groupMW...) mw = append(mw, extra...) r.routes = append(r.routes, route{ pluginID: pluginID, method: method, path: full, handler: handler, middleware: mw, raw: raw, }) } func (r *Router) compile() (http.Handler, error) { if r.compileErr != nil { return nil, r.compileErr } mux := http.NewServeMux() for _, rt := range r.routes { h, err := r.wrap(rt) if err != nil { return nil, err } if err := handleRoute(mux, rt, h); err != nil { return nil, err } } return pathScopedCORS(r.corsCfg, mux), nil } // handleRoute registers a route, converting a ServeMux conflict panic into an error. func handleRoute(mux *http.ServeMux, rt route, h http.Handler) (err error) { defer func() { if rec := recover(); rec != nil { err = fmt.Errorf("surf: route conflict for %s %s (plugin %q): %v", rt.method, rt.path, rt.pluginID, rec) } }() mux.Handle(rt.method+" "+rt.path, h) return nil } func (r *Router) wrap(rt route) (http.Handler, error) { h := constrain(rt.handler, rt.constraints) limit, err := routeBodyLimit(rt, r.defaultBytes) if err != nil { return nil, fmt.Errorf("surf: plugin %q: %w", rt.pluginID, err) } h = orgSlot(h) for i := len(rt.middleware) - 1; i >= 0; i-- { name := rt.middleware[i] if named, ok := r.named[name]; ok { if rt.raw && named.houseTagged { return nil, fmt.Errorf("surf: raw group cannot use house-envelope middleware %q (plugin %q)", name, rt.pluginID) } h = named.fn(h) continue } base, param, hasParam := strings.Cut(name, ":") if hasParam { if factory, ok := r.factories[base]; ok { if rt.raw && factory.houseTagged { return nil, fmt.Errorf("surf: raw group cannot use house-envelope middleware %q (plugin %q)", name, rt.pluginID) } if base == "throttle" && r.limiter != nil { if err := r.limiter.ValidateThrottle(param); err != nil { return nil, fmt.Errorf("surf: plugin %q: %w", rt.pluginID, err) } } if base == "body.limit" { if _, err := parseBodyLimit(param); err != nil { return nil, fmt.Errorf("surf: plugin %q: %w", rt.pluginID, err) } } mw, ok := r.built[name] if !ok { mw = factory.fn(param) r.built[name] = mw } h = mw(h) continue } } return nil, fmt.Errorf("surf: plugin %q references unknown middleware %q", rt.pluginID, name) } // Body cap wraps every named/factory middleware but stays inside recovery. if limit > 0 { h = bodyLimit(limit)(h) } h = locale(h) if rt.raw { h = recoverBare(h) } else { h = recoverJSON(h) } return h, nil } // Assemble registers plugin middleware and routes, then compiles ServeMux. func Assemble(app *backpack.App, plugins []party.Plugin) (http.Handler, error) { r, err := BuildRouter(app, plugins) if err != nil { return nil, err } return r.compile() } // BuildRouter registers plugin middleware and routes without compiling ServeMux, // so callers (route:list) can inspect Routes() after a successful boot. func BuildRouter(app *backpack.App, plugins []party.Plugin) (*Router, error) { r := New(corsOrigins(app)) trusted := TrustedProxies(nil) if app != nil { trusted = TrustedProxies(app.Config) } // longest bucket decay is 1 minute; sweep at 2x lim := NewFixedWindowLimiter(NewMemoryStore(2*time.Minute), trusted) r.limiter = lim if err := r.RegisterMiddlewareFactory("surf", "throttle", func(param string) pact.Middleware { return lim.Middleware(param) }); err != nil { return nil, err } if err := r.RegisterMiddlewareFactory("surf", "body.limit", func(param string) pact.Middleware { // Limit is applied outermost-inside-recovery in wrap(); the factory only occupies the name. return func(next http.Handler) http.Handler { return next } }); err != nil { return nil, err } if err := r.RegisterMiddleware("surf", "locale.from-principal", LocaleFromPrincipal); err != nil { return nil, err } if app != nil { corsCfg, err := LoadCORSConfig(app.Config) if err != nil { return nil, err } r.corsCfg = corsCfg if app.Config != nil { d, err := requiredBytes(app, "http.body_limits.default_bytes") if err != nil { return nil, err } u, err := requiredBytes(app, "http.body_limits.upload_bytes") if err != nil { return nil, err } r.defaultBytes, r.uploadBytes = d, u } } for _, p := range plugins { if hm, ok := p.(pact.HasMiddleware); ok { for name, fn := range hm.Middlewares() { if err := r.RegisterMiddleware(p.ID(), name, fn); err != nil { return nil, err } } } } 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 { if hh, ok := p.(pact.HasHouseMiddleware); ok { for name, fn := range hh.HouseMiddlewares() { if err := r.RegisterHouseMiddleware(p.ID(), name, fn); err != nil { return nil, err } } } } for _, p := range plugins { if bp, ok := p.(BucketProvider); ok { for name, b := range bp.Buckets() { if err := lim.RegisterBucket(p.ID(), name, b); err != nil { return nil, err } } } } for _, p := range plugins { r.BindPlugin(p.ID()) if hr, ok := p.(pact.HasRoutes); ok { if err := hr.Routes(r); err != nil { return nil, fmt.Errorf("surf: routes %s: %w", p.ID(), err) } } } admin, err := cabana.Activate(app, plugins) if err != nil { return nil, err } if admin != nil { if err := r.RegisterMiddleware("summercms.cabana", "backend", admin.Middleware); err != nil { return nil, err } r.BindPlugin("summercms.cabana") admin.Mount(r) if err := r.checkAdminPrefix(admin.Prefix); err != nil { return nil, err } } for _, rt := range r.routes { if _, err := r.wrap(rt); err != nil { return nil, err } } return r, nil } // checkAdminPrefix fails boot when a plugin other than cabana owns a route at // or under the admin prefix: the admin SPA and API own that whole subtree. func (r *Router) checkAdminPrefix(prefix string) error { if prefix == "" { return nil } for _, rt := range r.routes { if rt.pluginID == "summercms.cabana" { continue } if rt.path == prefix || strings.HasPrefix(rt.path, prefix+"/") { return fmt.Errorf("surf: route %s %s (plugin %q) is under the admin prefix backend.uri %s", rt.method, rt.path, rt.pluginID, prefix) } } return nil } func requiredBytes(app *backpack.App, key string) (int64, error) { raw, ok := app.Config.Lookup(key) if !ok { return 0, fmt.Errorf("surf: config %s is required", key) } var n int64 switch v := raw.(type) { case int: n = int64(v) case int64: n = v case uint64: if v > math.MaxInt64 { return 0, fmt.Errorf("surf: config %s out of range", key) } n = int64(v) case float64: if v != math.Trunc(v) || v > 9007199254740992 { return 0, fmt.Errorf("surf: config %s must be a whole number", key) } n = int64(v) default: return 0, fmt.Errorf("surf: config %s must be numeric, got %T", key, raw) } if n < 1 { return 0, fmt.Errorf("surf: config %s must be >= 1", key) } return n, nil } func corsOrigins(app *backpack.App) []string { if app == nil || app.Config == nil { return nil } raw, ok := app.Config.Lookup("http.cors.allowed_origins") if !ok { return nil } switch v := raw.(type) { case []string: return v case []any: out := make([]string, 0, len(v)) for _, item := range v { s, _ := item.(string) if s != "" { out = append(out, s) } } return out default: return nil } } func recoverJSON(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { buffered := newBufferedResponse() defer func() { if rec := recover(); rec != nil { wire.WriteOpaque500(w) return } buffered.commit(w) }() next.ServeHTTP(buffered, r) }) } func recoverBare(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { buffered := newBufferedResponse() defer func() { if rec := recover(); rec != nil { w.WriteHeader(http.StatusInternalServerError) return } buffered.commit(w) }() next.ServeHTTP(buffered, r) }) } type bufferedResponse struct { header http.Header status int body bytes.Buffer } func newBufferedResponse() *bufferedResponse { return &bufferedResponse{header: make(http.Header)} } func (w *bufferedResponse) Header() http.Header { return w.header } func (w *bufferedResponse) WriteHeader(status int) { if w.status == 0 { w.status = status } } func (w *bufferedResponse) Write(p []byte) (int, error) { if w.status == 0 { w.status = http.StatusOK } return w.body.Write(p) } // Flush deliberately does not expose or commit the destination writer. Route // output becomes visible only after the handler returns successfully. func (*bufferedResponse) Flush() {} func (w *bufferedResponse) commit(dst http.ResponseWriter) { for key, values := range w.header { dst.Header()[key] = append([]string(nil), values...) } status := w.status if status == 0 { status = http.StatusOK } dst.WriteHeader(status) _, _ = dst.Write(w.body.Bytes()) } func locale(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { next.ServeHTTP(w, r.WithContext(towel.WithLocale(r.Context(), r.Header.Get("Accept-Language")))) }) } func orgSlot(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { next.ServeHTTP(w, r.WithContext(towel.WithOrganization(r.Context(), ""))) }) } // Limiter wraps handlers. It is a retained Phase 3 seam with no implementer // after Phase 6; rate limiting is provided by FixedWindowLimiter through // the parameterized throttle middleware. type Limiter interface { Wrap(http.Handler) http.Handler } func joinPath(prefix, path string) string { prefix = strings.TrimSuffix(prefix, "/") path = strings.TrimSpace(path) if path == "" || path == "/" { if prefix == "" { return "/" } return prefix } if !strings.HasPrefix(path, "/") { path = "/" + path } if prefix == "" { return path } return prefix + path }