package surf import ( "fmt" "net/http" "strconv" "strings" "git.golem15.com/golem15/summercms/backpack" "git.golem15.com/golem15/summercms/pact" "git.golem15.com/golem15/summercms/party" "git.golem15.com/golem15/summercms/towel" ) // 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 } type route struct { pluginID string method string path string handler http.Handler middleware []string } // Router compiles group declarations onto net/http ServeMux. type Router struct { pluginID string prefix string middleware []string named map[string]namedMiddleware routes []route seen map[string]string origins []string compileErr error } 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 } // New returns an empty router. func New(origins []string) *Router { return &Router{ named: make(map[string]namedMiddleware), seen: make(map[string]string), origins: origins, } } // RegisterMiddleware stores a named wrapper. Duplicate names fail. func (r *Router) RegisterMiddleware(pluginID, name string, 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} 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)) { if r == nil || fn == nil { return } g := &Group{ router: r, pluginID: r.pluginID, prefix: joinPath(r.prefix, prefix), middleware: append([]string{}, r.middleware...), } 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, path, handler, middleware) } func (g *Group) Group(prefix string, middleware []string, 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...), } 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, path, handler, middleware) } func (r *Router) add(pluginID, prefix string, groupMW []string, path string, handler http.HandlerFunc, extra []string) { full := joinPath(prefix, path) key := "GET " + 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: "GET", path: full, handler: handler, middleware: mw, }) } 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 } mux.Handle("GET "+rt.path, h) } return recoverJSON(cors(r.origins, mux)), nil } func (r *Router) wrap(rt route) (http.Handler, error) { h := rt.handler h = noOpLimit(h) 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) } h = named.fn(h) } h = locale(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 := New(corsOrigins(app)) 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 { 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) } } } for _, rt := range r.routes { if _, err := r.wrap(rt); err != nil { return nil, err } } return r.compile() } 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) { defer func() { if rec := recover(); rec != nil { w.Header().Set("Content-Type", "application/json") w.WriteHeader(http.StatusInternalServerError) _, _ = w.Write([]byte(`{"error":true,"message":"Internal server error"}`)) } }() next.ServeHTTP(w, r) }) } func cors(origins []string, next http.Handler) http.Handler { allowed := make(map[string]struct{}, len(origins)) for _, o := range origins { allowed[o] = struct{}{} } return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { origin := r.Header.Get("Origin") if _, ok := allowed[origin]; ok && origin != "" { w.Header().Set("Access-Control-Allow-Origin", origin) w.Header().Set("Vary", "Origin") w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type, Accept") w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, PATCH, DELETE, OPTIONS") } if r.Method == http.MethodOptions { w.WriteHeader(http.StatusNoContent) return } next.ServeHTTP(w, r) }) } 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. Phase 6 replaces the no-op with named buckets. type Limiter interface { Wrap(http.Handler) http.Handler } type noopLimiter struct{} func (noopLimiter) Wrap(next http.Handler) http.Handler { return next } func noOpLimit(next http.Handler) http.Handler { return noopLimiter{}.Wrap(next) } // IntParam returns a positive integer path value. Missing or malformed ids // are false so callers can 404 both unknown and non-integer values. func IntParam(r *http.Request, name string) (int64, bool) { if r == nil { return 0, false } raw := r.PathValue(name) if raw == "" { return 0, false } n, err := strconv.ParseInt(raw, 10, 64) if err != nil || n < 1 { return 0, false } return n, true } 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 }