refactor(10.2-01): nest framework packages under modules

- Move remaining beach packages and embedded admin assets\n- Rewrite framework, example, build, and gate paths
This commit is contained in:
Jakub Zych
2026-09-28 02:21:02 +02:00
parent ac1f6d14f4
commit 5e50b166ef
277 changed files with 303 additions and 303 deletions

View File

@@ -0,0 +1,179 @@
package surf
import (
"io/fs"
"net/http"
"os"
"path/filepath"
"strings"
"testing"
"testing/fstest"
"git.golem15.com/golem15/summercms/modules/backpack"
"git.golem15.com/golem15/summercms/modules/compass"
"git.golem15.com/golem15/summercms/modules/pact"
"git.golem15.com/golem15/summercms/modules/party"
)
// TestPhase10AdminPrefixCollision fails boot when a plugin other than cabana
// registers a route at or under the admin prefix (backend.uri), which would
// otherwise shadow or be shadowed by the admin SPA and API.
func TestPhase10AdminPrefixCollision(t *testing.T) {
for _, tc := range []struct {
name string
method string
path string
raw bool
}{
{"exact prefix", http.MethodGet, "/acme-admin", false},
{"under prefix", http.MethodPost, "/acme-admin/hook", false},
{"raw under api", http.MethodGet, "/acme-admin/api/v1/extra", true},
{"deeper path", http.MethodGet, "/acme-admin/a/b/c", false},
{"exact prefix raw", http.MethodPost, "/acme-admin", true},
} {
t.Run(tc.name, func(t *testing.T) {
app := adminPrefixApp(t)
plugins := []party.Plugin{
adminPrefixPlugin{},
prefixRoutePlugin{id: "acme.intruder", method: tc.method, path: tc.path, raw: tc.raw},
}
_, err := BuildRouter(app, plugins)
if err == nil {
t.Fatalf("BuildRouter accepted %s %s under the admin prefix", tc.method, tc.path)
}
for _, want := range []string{tc.method, tc.path, "acme.intruder", "/acme-admin"} {
if !strings.Contains(err.Error(), want) {
t.Fatalf("error %q does not name %q", err, want)
}
}
})
}
t.Run("sibling path is allowed", func(t *testing.T) {
app := adminPrefixApp(t)
plugins := []party.Plugin{
adminPrefixPlugin{},
prefixRoutePlugin{id: "acme.neighbour", method: http.MethodGet, path: "/acme-adminx"},
}
if _, err := BuildRouter(app, plugins); err != nil {
t.Fatalf("sibling path rejected: %v", err)
}
})
t.Run("default /backend prefix", func(t *testing.T) {
for _, tc := range []struct {
path string
collide bool
}{
{"/backend", true},
{"/backend/deep/hook", true},
{"/backendx", false},
{"/back", false},
{"/api/backend", false},
} {
app := adminPrefixAppWith(t, "")
plugins := []party.Plugin{
adminPrefixPlugin{},
prefixRoutePlugin{id: "acme.intruder", method: http.MethodGet, path: tc.path},
}
_, err := BuildRouter(app, plugins)
if tc.collide && (err == nil || !strings.Contains(err.Error(), "/backend")) {
t.Fatalf("%s under the default prefix: err=%v", tc.path, err)
}
if !tc.collide && err != nil {
t.Fatalf("%s rejected next to the default prefix: %v", tc.path, err)
}
}
})
t.Run("no admin means no prefix check", func(t *testing.T) {
app := adminPrefixApp(t)
plugins := []party.Plugin{prefixRoutePlugin{id: "acme.site", method: http.MethodGet, path: "/acme-admin"}}
if _, err := BuildRouter(app, plugins); err != nil {
t.Fatalf("a route at the prefix without admin controllers was rejected: %v", err)
}
if err := (&Router{}).checkAdminPrefix(""); err != nil {
t.Fatalf("empty prefix: %v", err)
}
})
}
func adminPrefixApp(t *testing.T) *backpack.App {
t.Helper()
return adminPrefixAppWith(t, "/acme-admin")
}
// adminPrefixAppWith configures backend.uri; an empty uri keeps the default.
func adminPrefixAppWith(t *testing.T, uri string) *backpack.App {
t.Helper()
dir := t.TempDir()
if err := os.WriteFile(filepath.Join(dir, "app.yaml"), []byte("name: admin-prefix\n"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(dir, "http.yaml"), []byte("body_limits:\n default_bytes: 1024\n upload_bytes: 1024\n"), 0o644); err != nil {
t.Fatal(err)
}
if uri != "" {
if err := os.WriteFile(filepath.Join(dir, "backend.yaml"), []byte("uri: "+uri+"\n"), 0o644); err != nil {
t.Fatal(err)
}
}
cfg, err := compass.Open(compass.Options{
Dir: dir,
Environ: []string{"SUMMER_ENV=development", "SUMMER_ADMIN__JWT__SECRET=summercms-test-only-admin-hs256-secret"},
})
if err != nil {
t.Fatal(err)
}
return backpack.New(cfg)
}
type adminPrefixPlugin struct{}
func (adminPrefixPlugin) ID() string { return "acme.demo" }
func (adminPrefixPlugin) Requires() []string { return nil }
func (adminPrefixPlugin) Register(*backpack.App) error { return nil }
func (adminPrefixPlugin) Boot(*backpack.App) error { return nil }
func (adminPrefixPlugin) AdminControllers() []pact.AdminController {
return []pact.AdminController{adminPrefixController{}}
}
func (adminPrefixPlugin) AdminFS() fs.FS {
return fstest.MapFS{
"controllers/widgets/config_list.yaml": &fstest.MapFile{Data: []byte("list: ~/plugins/acme/demo/models/widget/columns.yaml\nmodelClass: Widget\nrecordsPerPage: 20\n")},
"models/widget/columns.yaml": &fstest.MapFile{Data: []byte("columns:\n name:\n label: Name\n")},
}
}
type adminPrefixController struct{}
func (adminPrefixController) ID() string { return "acme.demo.widgets" }
func (adminPrefixController) ModelName() string { return "Widget" }
func (adminPrefixController) ConfigDir() string { return "controllers/widgets" }
type prefixRoutePlugin struct {
id string
method string
path string
raw bool
}
func (p prefixRoutePlugin) ID() string { return p.id }
func (p prefixRoutePlugin) Requires() []string { return nil }
func (p prefixRoutePlugin) Register(*backpack.App) error { return nil }
func (p prefixRoutePlugin) Boot(*backpack.App) error { return nil }
func (p prefixRoutePlugin) Routes(r pact.Router) error {
h := func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusNoContent) }
open := r.Group
if p.raw {
open = r.GroupRaw
}
open("", nil, func(g pact.Router) {
switch p.method {
case http.MethodPost:
g.Post(p.path, h)
default:
g.Get(p.path, h)
}
})
return nil
}

46
modules/surf/bodylimit.go Normal file
View File

@@ -0,0 +1,46 @@
package surf
import (
"fmt"
"net/http"
"strconv"
"strings"
)
func bodyLimit(n int64) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if n > 0 && r.Body != nil {
r.Body = http.MaxBytesReader(w, r.Body, n)
}
next.ServeHTTP(w, r)
})
}
}
func parseBodyLimit(param string) (int64, error) {
n, err := strconv.ParseInt(param, 10, 64)
if err != nil || n <= 0 {
return 0, fmt.Errorf("invalid body.limit %q", param)
}
return n, nil
}
func routeBodyLimit(rt route, defaultBytes int64) (int64, error) {
if rt.raw {
return 0, nil
}
limit := defaultBytes
for _, name := range rt.middleware {
base, param, ok := strings.Cut(name, ":")
if !ok || base != "body.limit" {
continue
}
n, err := parseBodyLimit(param)
if err != nil {
return 0, err
}
limit = n
}
return limit, nil
}

View File

@@ -0,0 +1,283 @@
package surf
import (
"errors"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"git.golem15.com/golem15/summercms/modules/backpack"
"git.golem15.com/golem15/summercms/modules/pact"
"git.golem15.com/golem15/summercms/modules/party"
)
func TestBodyLimitDefaultRejectsOversizedBody(t *testing.T) {
cfg := writeHTTPConfig(t, `
body_limits:
default_bytes: 32
upload_bytes: 64
`)
p := bodyEchoPlugin{id: "golem15.demo"}
h, err := Assemble(backpack.New(cfg), []party.Plugin{p})
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/echo", strings.NewReader(strings.Repeat("a", 64)))
h.ServeHTTP(rec, req)
if rec.Code != http.StatusRequestEntityTooLarge {
t.Fatalf("status = %d body=%q (want 413 from MaxBytesReader)", rec.Code, rec.Body.String())
}
rec = httptest.NewRecorder()
req = httptest.NewRequest(http.MethodPost, "/echo", strings.NewReader("ok"))
h.ServeHTTP(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("small body status = %d", rec.Code)
}
}
func TestBodyLimitRawExempt(t *testing.T) {
cfg := writeHTTPConfig(t, `
body_limits:
default_bytes: 8
upload_bytes: 8
`)
p := bodyEchoPlugin{id: "golem15.demo", raw: true}
h, err := Assemble(backpack.New(cfg), []party.Plugin{p})
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/echo", strings.NewReader(strings.Repeat("a", 64)))
h.ServeHTTP(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("raw route should not apply default body limit, status = %d", rec.Code)
}
}
func TestBodyLimitOverride(t *testing.T) {
cfg := writeHTTPConfig(t, `
body_limits:
default_bytes: 8
upload_bytes: 64
`)
p := bodyEchoPlugin{id: "golem15.demo", extra: []string{"body.limit:64"}}
h, err := Assemble(backpack.New(cfg), []party.Plugin{p})
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/echo", strings.NewReader(strings.Repeat("a", 32)))
h.ServeHTTP(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("override should allow 32 bytes, status = %d", rec.Code)
}
}
func TestBodyLimitLoadsConfigValues(t *testing.T) {
cfg := writeHTTPConfig(t, `
body_limits:
default_bytes: 8388608
upload_bytes: 2097152
`)
r, err := BuildRouter(backpack.New(cfg), nil)
if err != nil {
t.Fatal(err)
}
if r.defaultBytes != 8388608 || r.uploadBytes != 2097152 {
t.Fatalf("limits = %d / %d", r.defaultBytes, r.uploadBytes)
}
}
type bodyEchoPlugin struct {
id string
raw bool
extra []string
}
func (p bodyEchoPlugin) ID() string { return p.id }
func (p bodyEchoPlugin) Requires() []string { return nil }
func (p bodyEchoPlugin) Register(*backpack.App) error { return nil }
func (p bodyEchoPlugin) Boot(*backpack.App) error { return nil }
func (p bodyEchoPlugin) Routes(r pact.Router) error {
h := func(w http.ResponseWriter, req *http.Request) {
_, err := io.Copy(io.Discard, req.Body)
if err != nil {
var maxErr *http.MaxBytesError
if errors.As(err, &maxErr) {
w.WriteHeader(http.StatusRequestEntityTooLarge)
return
}
w.WriteHeader(http.StatusInternalServerError)
return
}
w.WriteHeader(http.StatusOK)
}
open := r.Group
if p.raw {
open = r.GroupRaw
}
open("/", Use(p.extra...), func(g pact.Router) {
g.Post("/echo", h)
})
return nil
}
type mwReadPlugin struct{ seen *error }
func (p mwReadPlugin) ID() string { return "golem15.mwread" }
func (p mwReadPlugin) Requires() []string { return nil }
func (p mwReadPlugin) Register(*backpack.App) error { return nil }
func (p mwReadPlugin) Boot(*backpack.App) error { return nil }
func (p mwReadPlugin) Middlewares() map[string]pact.Middleware {
return map[string]pact.Middleware{"reader": func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, *p.seen = io.ReadAll(r.Body)
next.ServeHTTP(w, r)
})
}}
}
func (p mwReadPlugin) Routes(r pact.Router) error {
r.Group("/", Use("reader"), func(g pact.Router) {
g.Post("/x", func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) })
})
return nil
}
func TestBodyLimitBoundsNamedMiddleware(t *testing.T) {
cfg := writeHTTPConfig(t, "body_limits:\n default_bytes: 4\n upload_bytes: 8\n")
var seen error
h, err := Assemble(backpack.New(cfg), []party.Plugin{mwReadPlugin{seen: &seen}})
if err != nil {
t.Fatal(err)
}
req := httptest.NewRequest(http.MethodPost, "/x", strings.NewReader(strings.Repeat("a", 10)))
h.ServeHTTP(httptest.NewRecorder(), req)
var maxErr *http.MaxBytesError
if !errors.As(seen, &maxErr) {
t.Fatalf("middleware read err = %v, want MaxBytesError", seen)
}
}
func TestBodyLimitMissingConfigFailsBoot(t *testing.T) {
cfg := writeHTTPConfig(t, "body_limits:\n upload_bytes: 8\n")
if _, err := BuildRouter(backpack.New(cfg), nil); err == nil || !strings.Contains(err.Error(), "default_bytes") {
t.Fatalf("err = %v", err)
}
}
func TestCompileRouteConflictReturnsError(t *testing.T) {
r := New(nil)
ok := func(w http.ResponseWriter, _ *http.Request) {}
r.BindPlugin("a")
r.Get("/a/{x}", ok)
r.Get("/a/{y}", ok)
if _, err := r.compile(); err == nil {
t.Fatal("expected conflict error")
}
}
type bodyProbePlugin struct {
id string
raw bool
use []string
mw map[string]pact.Middleware
routes func(pact.Router)
}
func (p bodyProbePlugin) ID() string { return p.id }
func (p bodyProbePlugin) Requires() []string { return nil }
func (p bodyProbePlugin) Register(*backpack.App) error { return nil }
func (p bodyProbePlugin) Boot(*backpack.App) error { return nil }
func (p bodyProbePlugin) Middlewares() map[string]pact.Middleware { return p.mw }
func (p bodyProbePlugin) Routes(r pact.Router) error {
open := r.Group
if p.raw {
open = r.GroupRaw
}
open("/", Use(p.use...), func(g pact.Router) {
g.Post("/x", func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) })
})
return nil
}
func TestBodyLimitBoundsBodyConsumingMiddleware(t *testing.T) {
type result struct {
n int
err error
}
run := func(t *testing.T, raw bool, use []string, bodyLen int, after func()) (result, *httptest.ResponseRecorder) {
t.Helper()
var res result
reader := func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
b, err := io.ReadAll(r.Body)
res = result{n: len(b), err: err}
if after != nil {
after()
}
next.ServeHTTP(w, r)
})
}
cfg := writeHTTPConfig(t, "body_limits:\n default_bytes: 4\n upload_bytes: 8\n")
h, err := Assemble(backpack.New(cfg), []party.Plugin{bodyProbePlugin{
id: "golem15.probe", raw: raw, use: use,
mw: map[string]pact.Middleware{"reader": reader},
}})
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/x", strings.NewReader(strings.Repeat("a", bodyLen))))
return res, rec
}
var maxErr *http.MaxBytesError
t.Run("default limit", func(t *testing.T) {
res, _ := run(t, false, []string{"reader"}, 10, nil)
if !errors.As(res.err, &maxErr) || res.n > 4 {
t.Fatalf("read %d bytes, err = %v", res.n, res.err)
}
})
t.Run("body.limit override raises cap", func(t *testing.T) {
res, _ := run(t, false, []string{"reader", "body.limit:6"}, 5, nil)
if res.err != nil || res.n != 5 {
t.Fatalf("read %d bytes, err = %v", res.n, res.err)
}
})
t.Run("body.limit override still bounds", func(t *testing.T) {
res, _ := run(t, false, []string{"reader", "body.limit:6"}, 10, nil)
if !errors.As(res.err, &maxErr) || res.n > 6 {
t.Fatalf("read %d bytes, err = %v", res.n, res.err)
}
})
t.Run("raw route unaffected", func(t *testing.T) {
res, rec := run(t, true, []string{"reader"}, 10, nil)
if res.err != nil || res.n != 10 || rec.Code != http.StatusOK {
t.Fatalf("read %d bytes, err = %v, code %d", res.n, res.err, rec.Code)
}
})
t.Run("panic after read still clean 500", func(t *testing.T) {
_, rec := run(t, false, []string{"reader"}, 10, func() { panic("boom") })
if rec.Code != http.StatusInternalServerError || strings.Contains(rec.Body.String(), "boom") {
t.Fatalf("code %d body %q", rec.Code, rec.Body.String())
}
})
}
func TestBodyLimitInvalidParamFailsBoot(t *testing.T) {
for _, param := range []string{"abc", "0", "-5", ""} {
t.Run(param, func(t *testing.T) {
cfg := writeHTTPConfig(t, "body_limits:\n default_bytes: 4\n upload_bytes: 8\n")
_, err := Assemble(backpack.New(cfg), []party.Plugin{bodyProbePlugin{
id: "golem15.probe", use: []string{"body.limit:" + param},
}})
if err == nil || !strings.Contains(err.Error(), "body.limit") {
t.Fatalf("err = %v", err)
}
})
}
}

105
modules/surf/clientip.go Normal file
View File

@@ -0,0 +1,105 @@
package surf
import (
"net"
"net/http"
"net/netip"
"strings"
"git.golem15.com/golem15/summercms/modules/compass"
)
// ClientIP is the single source of client IP for limiter keys (D-04).
// RemoteAddr is used unless it parses as being inside one of trusted;
// in that case the rightmost X-Forwarded-For hop NOT inside any trusted
// prefix is used. An empty trusted list means RemoteAddr only.
func ClientIP(r *http.Request, trusted []netip.Prefix) string {
if r == nil {
return ""
}
remote := parseIP(r.RemoteAddr)
if len(trusted) == 0 || remote == (netip.Addr{}) || !addrTrusted(remote, trusted) {
if remote == (netip.Addr{}) {
return ""
}
return remote.String()
}
xff := r.Header.Get("X-Forwarded-For")
if xff == "" {
return remote.String()
}
hops := strings.Split(xff, ",")
for i := len(hops) - 1; i >= 0; i-- {
hop := strings.TrimSpace(hops[i])
if hop == "" {
continue
}
addr, err := netip.ParseAddr(hop)
if err != nil {
continue
}
addr = addr.Unmap()
if !addrTrusted(addr, trusted) {
return addr.String()
}
}
return remote.String()
}
// TrustedProxies reads http.trusted_proxies (a []string of CIDRs) from cfg
// and parses it once into []netip.Prefix. A malformed entry is skipped, not
// fatal (logged by the caller if desired).
func TrustedProxies(cfg *compass.Config) []netip.Prefix {
if cfg == nil {
return nil
}
raw, ok := cfg.Lookup("http.trusted_proxies")
if !ok {
return nil
}
var entries []string
switch v := raw.(type) {
case []string:
entries = v
case []any:
for _, item := range v {
s, _ := item.(string)
if s != "" {
entries = append(entries, s)
}
}
default:
return nil
}
var out []netip.Prefix
for _, e := range entries {
p, err := netip.ParsePrefix(strings.TrimSpace(e))
if err != nil {
continue
}
out = append(out, p)
}
return out
}
func parseIP(remoteAddr string) netip.Addr {
host, _, err := net.SplitHostPort(remoteAddr)
if err != nil {
host = remoteAddr
}
host = strings.Trim(host, "[]")
addr, err := netip.ParseAddr(host)
if err != nil {
return netip.Addr{}
}
return addr.Unmap()
}
func addrTrusted(addr netip.Addr, trusted []netip.Prefix) bool {
for _, p := range trusted {
if p.Contains(addr) {
return true
}
}
return false
}

View File

@@ -0,0 +1,59 @@
package surf
import (
"net/http"
"net/http/httptest"
"net/netip"
"testing"
)
func TestClientIPEmptyTrustedUsesRemoteAddr(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.RemoteAddr = "203.0.113.9:1234"
req.Header.Set("X-Forwarded-For", "198.51.100.1")
got := ClientIP(req, nil)
if got != "203.0.113.9" {
t.Fatalf("got %q", got)
}
}
func TestClientIPRejectsSpoofedXFF(t *testing.T) {
trusted := []netip.Prefix{mustPrefix("10.0.0.0/8")}
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.RemoteAddr = "203.0.113.9:1234"
req.Header.Set("X-Forwarded-For", "198.51.100.1")
got := ClientIP(req, trusted)
if got != "203.0.113.9" {
t.Fatalf("untrusted RemoteAddr must ignore X-Forwarded-For, got %q", got)
}
}
func TestClientIPRightmostUntrustedHop(t *testing.T) {
trusted := []netip.Prefix{mustPrefix("10.0.0.0/8")}
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.RemoteAddr = "10.0.0.1:443"
req.Header.Set("X-Forwarded-For", "198.51.100.7, 203.0.113.10, 10.0.0.2")
got := ClientIP(req, trusted)
if got != "203.0.113.10" {
t.Fatalf("rightmost untrusted hop = %q", got)
}
}
func TestClientIPAllHopsTrustedFallsBack(t *testing.T) {
trusted := []netip.Prefix{mustPrefix("10.0.0.0/8")}
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.RemoteAddr = "10.0.0.1:443"
req.Header.Set("X-Forwarded-For", "10.0.0.8, 10.0.0.9")
got := ClientIP(req, trusted)
if got != "10.0.0.1" {
t.Fatalf("fallback RemoteAddr = %q", got)
}
}
func mustPrefix(s string) netip.Prefix {
p, err := netip.ParsePrefix(s)
if err != nil {
panic(err)
}
return p
}

161
modules/surf/cors.go Normal file
View File

@@ -0,0 +1,161 @@
package surf
import (
"net/http"
"regexp"
"strconv"
"strings"
"git.golem15.com/golem15/summercms/modules/compass"
)
// CORSConfig matches Laravel config/cors.php keys.
type CORSConfig struct {
Paths []string `koanf:"paths"`
AllowedMethods []string `koanf:"allowed_methods"`
AllowedOrigins []string `koanf:"allowed_origins"`
AllowedOriginsPatterns []string `koanf:"allowed_origins_patterns"`
AllowedHeaders []string `koanf:"allowed_headers"`
ExposedHeaders []string `koanf:"exposed_headers"`
MaxAge int `koanf:"max_age"`
SupportsCredentials bool `koanf:"supports_credentials"`
}
// LoadCORSConfig reads http.cors. Missing section yields a zero config (no
// path matches, so no CORS headers).
func LoadCORSConfig(cfg *compass.Config) (CORSConfig, error) {
var out CORSConfig
if cfg == nil || !cfg.Has("http.cors") {
return out, nil
}
if err := cfg.LoadSection("http.cors", &out); err != nil {
return out, err
}
return out, nil
}
func pathScopedCORS(cfg CORSConfig, next http.Handler) http.Handler {
globs := make([]*regexp.Regexp, 0, len(cfg.Paths))
for _, p := range cfg.Paths {
if re := compileLaravelGlob(p); re != nil {
globs = append(globs, re)
}
}
originPats := make([]*regexp.Regexp, 0, len(cfg.AllowedOriginsPatterns))
for _, p := range cfg.AllowedOriginsPatterns {
re, err := regexp.Compile(p)
if err != nil {
continue
}
originPats = append(originPats, re)
}
allowAnyOrigin := containsStar(cfg.AllowedOrigins)
allowAnyMethod := containsStar(cfg.AllowedMethods)
allowAnyHeader := containsStar(cfg.AllowedHeaders)
origins := make(map[string]struct{}, len(cfg.AllowedOrigins))
for _, o := range cfg.AllowedOrigins {
if o != "" && o != "*" {
origins[o] = struct{}{}
}
}
methods := strings.Join(cfg.AllowedMethods, ", ")
headers := strings.Join(cfg.AllowedHeaders, ", ")
exposed := strings.Join(cfg.ExposedHeaders, ", ")
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if !pathMatchesCORS(globs, r.URL.Path) {
next.ServeHTTP(w, r)
return
}
origin := r.Header.Get("Origin")
allowed, value := corsAllowOrigin(allowAnyOrigin, origins, originPats, origin)
if allowed {
w.Header().Set("Access-Control-Allow-Origin", value)
if value != "*" {
w.Header().Set("Vary", "Origin")
}
if allowAnyMethod {
w.Header().Set("Access-Control-Allow-Methods", "*")
} else if methods != "" {
w.Header().Set("Access-Control-Allow-Methods", methods)
}
if allowAnyHeader {
w.Header().Set("Access-Control-Allow-Headers", "*")
} else if headers != "" {
w.Header().Set("Access-Control-Allow-Headers", headers)
}
if exposed != "" {
w.Header().Set("Access-Control-Expose-Headers", exposed)
}
if cfg.MaxAge > 0 {
w.Header().Set("Access-Control-Max-Age", strconv.Itoa(cfg.MaxAge))
}
if cfg.SupportsCredentials {
w.Header().Set("Access-Control-Allow-Credentials", "true")
}
}
if r.Method == http.MethodOptions {
w.WriteHeader(http.StatusNoContent)
return
}
next.ServeHTTP(w, r)
})
}
func corsAllowOrigin(allowAny bool, origins map[string]struct{}, pats []*regexp.Regexp, origin string) (bool, string) {
if allowAny {
return true, "*"
}
if origin == "" {
return false, ""
}
if _, ok := origins[origin]; ok {
return true, origin
}
for _, re := range pats {
if re.MatchString(origin) {
return true, origin
}
}
return false, ""
}
func pathMatchesCORS(globs []*regexp.Regexp, urlPath string) bool {
trimmed := strings.TrimPrefix(urlPath, "/")
for _, re := range globs {
if re.MatchString(trimmed) {
return true
}
}
return false
}
func compileLaravelGlob(pattern string) *regexp.Regexp {
pattern = strings.TrimPrefix(pattern, "/")
var b strings.Builder
b.WriteString("(?s)^")
for i := 0; i < len(pattern); i++ {
switch pattern[i] {
case '*':
b.WriteString(".*")
case '?':
b.WriteByte('.')
default:
b.WriteString(regexp.QuoteMeta(pattern[i : i+1]))
}
}
b.WriteByte('$')
re, err := regexp.Compile(b.String())
if err != nil {
return nil
}
return re
}
func containsStar(vals []string) bool {
for _, v := range vals {
if v == "*" {
return true
}
}
return false
}

View File

@@ -0,0 +1,92 @@
package surf
import (
"net/http"
"net/http/httptest"
"testing"
)
// Gap (d): pathScopedCORS with a path matching NONE of the configured
// globs. TestCORSPathScopedHeaders already covers the two named acme
// groups; this fixture is framework-only (/healthz vs api/*).
func TestCORSPathScopedNoMatchIndependentOfAcme(t *testing.T) {
cfg := CORSConfig{
Paths: []string{"api/*", "oauth/mcp/*"},
AllowedMethods: []string{"*"},
AllowedOrigins: []string{"*"},
AllowedHeaders: []string{"*"},
}
inner := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusNoContent)
})
h := pathScopedCORS(cfg, inner)
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/healthz", nil))
if rec.Code != http.StatusNoContent {
t.Fatalf("status = %d", rec.Code)
}
if got := rec.Header().Get("Access-Control-Allow-Origin"); got != "" {
t.Fatalf("unmatched path must not set ACAO, got %q", got)
}
}
func TestCORSAllowOriginExactAndPattern(t *testing.T) {
cfg := CORSConfig{
Paths: []string{"api/*"},
AllowedMethods: []string{"GET", "POST"},
AllowedOrigins: []string{"https://app.example.test"},
AllowedOriginsPatterns: []string{`^https://.*\.example\.test$`},
AllowedHeaders: []string{"Authorization", "Content-Type"},
ExposedHeaders: []string{"X-RateLimit-Limit"},
MaxAge: 600,
SupportsCredentials: true,
}
inner := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
})
h := pathScopedCORS(cfg, inner)
t.Run("exact origin", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/api/v1/items", nil)
req.Header.Set("Origin", "https://app.example.test")
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Header().Get("Access-Control-Allow-Origin") != "https://app.example.test" {
t.Fatalf("ACAO = %q", rec.Header().Get("Access-Control-Allow-Origin"))
}
if rec.Header().Get("Vary") != "Origin" {
t.Fatalf("Vary = %q", rec.Header().Get("Vary"))
}
if rec.Header().Get("Access-Control-Allow-Credentials") != "true" {
t.Fatal("missing credentials header")
}
if rec.Header().Get("Access-Control-Max-Age") != "600" {
t.Fatalf("Max-Age = %q", rec.Header().Get("Access-Control-Max-Age"))
}
})
t.Run("pattern origin", func(t *testing.T) {
req := httptest.NewRequest(http.MethodOptions, "/api/v1/items", nil)
req.Header.Set("Origin", "https://admin.example.test")
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusNoContent {
t.Fatalf("preflight status = %d", rec.Code)
}
if rec.Header().Get("Access-Control-Allow-Origin") != "https://admin.example.test" {
t.Fatalf("ACAO = %q", rec.Header().Get("Access-Control-Allow-Origin"))
}
})
t.Run("disallowed origin", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/api/v1/items", nil)
req.Header.Set("Origin", "https://evil.test")
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if got := rec.Header().Get("Access-Control-Allow-Origin"); got != "" {
t.Fatalf("disallowed origin ACAO = %q", got)
}
})
}

122
modules/surf/cors_test.go Normal file
View File

@@ -0,0 +1,122 @@
package surf
import (
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"testing"
"git.golem15.com/golem15/summercms/modules/backpack"
"git.golem15.com/golem15/summercms/modules/compass"
"git.golem15.com/golem15/summercms/modules/party"
)
func TestCORSLaravelGlobMatchesNestedPaths(t *testing.T) {
re := compileLaravelGlob("api/*")
if re == nil || !re.MatchString("api/v1/acme/genres") {
t.Fatal("api/* must match api/v1/acme/genres (Laravel Str::is, not Go path.Match)")
}
if re.MatchString("_acme/api/v1/genres") {
t.Fatal("api/* must not match _acme/api/v1/genres")
}
mcp := compileLaravelGlob("oauth/mcp/*")
if mcp == nil || !mcp.MatchString("oauth/mcp/token") {
t.Fatal("oauth/mcp/* must match oauth/mcp/token")
}
}
func TestCORSPathScopedHeaders(t *testing.T) {
cfg := CORSConfig{
Paths: []string{"api/*", "oauth/mcp/*"},
AllowedMethods: []string{"*"},
AllowedOrigins: []string{"*"},
AllowedHeaders: []string{"*"},
}
mux := http.NewServeMux()
mux.HandleFunc("GET /api/v1/acme/genres", func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
})
mux.HandleFunc("GET /_acme/api/v1/genres", func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
})
h := pathScopedCORS(cfg, mux)
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/v1/acme/genres", nil))
if rec.Header().Get("Access-Control-Allow-Origin") != "*" {
t.Fatalf("token group ACAO = %q", rec.Header().Get("Access-Control-Allow-Origin"))
}
rec = httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/_acme/api/v1/genres", nil))
if got := rec.Header().Get("Access-Control-Allow-Origin"); got != "" {
t.Fatalf("JWT group ACAO = %q, want empty", got)
}
}
func TestCORSConfigLoadedFromHTTPSection(t *testing.T) {
cfg := writeHTTPConfig(t, `
cors:
paths: ["api/*"]
allowed_methods: ["*"]
allowed_origins: ["*"]
allowed_origins_patterns: []
allowed_headers: ["*"]
exposed_headers: []
max_age: 0
supports_credentials: false
`)
got, err := LoadCORSConfig(cfg)
if err != nil {
t.Fatal(err)
}
if len(got.Paths) != 1 || got.Paths[0] != "api/*" {
t.Fatalf("paths = %v", got.Paths)
}
if !containsStar(got.AllowedOrigins) {
t.Fatalf("origins = %v", got.AllowedOrigins)
}
}
func writeHTTPConfig(t *testing.T, httpYAML string) *compass.Config {
t.Helper()
dir := t.TempDir()
if err := os.WriteFile(filepath.Join(dir, "app.yaml"), []byte("name: cors-test\n"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(dir, "http.yaml"), []byte(httpYAML), 0o644); err != nil {
t.Fatal(err)
}
cfg, err := compass.Open(compass.Options{
Dir: dir,
Environ: []string{"SUMMER_ENV=development"},
})
if err != nil {
t.Fatal(err)
}
return cfg
}
func TestCORSAssembleUsesConfig(t *testing.T) {
cfg := writeHTTPConfig(t, `
body_limits:
default_bytes: 1048576
upload_bytes: 1048576
cors:
paths: ["api/*"]
allowed_methods: ["*"]
allowed_origins: ["*"]
allowed_headers: ["*"]
`)
p := assemblePlugin{id: "golem15.demo", path: "/genres"}
h, err := Assemble(backpack.New(cfg), []party.Plugin{p})
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/genres", nil))
if rec.Header().Get("Access-Control-Allow-Origin") != "*" {
t.Fatalf("ACAO = %q", rec.Header().Get("Access-Control-Allow-Origin"))
}
}

182
modules/surf/limiter.go Normal file
View File

@@ -0,0 +1,182 @@
package surf
import (
"fmt"
"math"
"net/http"
"net/netip"
"strconv"
"strings"
"sync"
"time"
"git.golem15.com/golem15/summercms/modules/bouncer"
"git.golem15.com/golem15/summercms/modules/pact"
)
const tooManyAttemptsBody = `{"message":"Too Many Attempts."}`
// Bucket is one named rate-limit definition. Key composes the limiter key
// from the request (token id, IP, route param -- D-01's per-bucket rule).
type Bucket struct {
Name string
Max int
Decay time.Duration
Key func(r *http.Request) string
}
// BucketProvider is implemented by plugins that declare named buckets (not a
// pact interface: it lives in surf and is type-asserted directly in
// Assemble/BuildRouter, since pact cannot import surf without a cycle).
type BucketProvider interface {
Buckets() map[string]Bucket
}
// FixedWindowLimiter is the concrete rate limiter (distinct from the
// pre-existing surf.Limiter interface). It owns the Store, the named-bucket
// table, and the inline-throttle parser, and produces the "throttle"
// middleware factory's per-route pact.Middleware.
type FixedWindowLimiter struct {
store Store
trusted []netip.Prefix
mu sync.Mutex
buckets map[string]Bucket
owners map[string]string
inline map[string]Bucket
}
// NewFixedWindowLimiter is the ONE constructor signature for this type --
// trusted is required at construction (not a later setter) because both
// named-bucket Key closures (registered later via RegisterBucket) and the
// inline "N,M" throttle's own key resolver need the same trusted-proxy list.
func NewFixedWindowLimiter(store Store, trusted []netip.Prefix) *FixedWindowLimiter {
cp := make([]netip.Prefix, len(trusted))
copy(cp, trusted)
return &FixedWindowLimiter{
store: store,
trusted: cp,
buckets: make(map[string]Bucket),
owners: make(map[string]string),
inline: make(map[string]Bucket),
}
}
// RegisterBucket stores a named bucket. Duplicate names fail.
func (l *FixedWindowLimiter) RegisterBucket(pluginID, name string, b Bucket) error {
if l == nil {
return fmt.Errorf("surf: limiter is nil")
}
if name == "" {
return fmt.Errorf("surf: plugin %q registered empty bucket", pluginID)
}
switch {
case l.store == nil:
return fmt.Errorf("surf: plugin %q bucket %q: limiter has no store", pluginID, name)
case b.Key == nil:
return fmt.Errorf("surf: plugin %q bucket %q has nil Key", pluginID, name)
case b.Max < 1:
return fmt.Errorf("surf: plugin %q bucket %q has Max < 1", pluginID, name)
case b.Decay <= 0:
return fmt.Errorf("surf: plugin %q bucket %q has non-positive Decay", pluginID, name)
}
l.mu.Lock()
defer l.mu.Unlock()
if existing, ok := l.owners[name]; ok {
return fmt.Errorf("surf: bucket %q already registered by %s", name, existing)
}
b.Name = name
l.buckets[name] = b
l.owners[name] = pluginID
return nil
}
// Middleware builds the throttle: factory body. param is either a
// registered bucket name or a literal "N,M" pair (inline throttle).
func (l *FixedWindowLimiter) Middleware(param string) pact.Middleware {
failClosed := func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
})
}
if l == nil {
return failClosed
}
b, err := l.resolve(param)
if err != nil || b.Key == nil || l.store == nil {
return failClosed
}
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
key := b.Key(r)
allowed, attempts, retryAfter := l.store.Attempt(key, b.Max, b.Decay)
if !allowed {
secs := int(retryAfter / time.Second)
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Retry-After", strconv.Itoa(secs))
w.Header().Set("X-RateLimit-Reset", strconv.FormatInt(time.Now().Add(retryAfter).Unix(), 10))
w.Header().Set("X-RateLimit-Limit", strconv.Itoa(b.Max))
w.Header().Set("X-RateLimit-Remaining", "0")
w.WriteHeader(http.StatusTooManyRequests)
_, _ = w.Write([]byte(tooManyAttemptsBody))
return
}
remaining := b.Max - attempts
if remaining < 0 {
remaining = 0
}
w.Header().Set("X-RateLimit-Limit", strconv.Itoa(b.Max))
w.Header().Set("X-RateLimit-Remaining", strconv.Itoa(remaining))
next.ServeHTTP(w, r)
})
}
}
// ValidateThrottle is called once per route at Assemble time so a malformed
// inline "N,M" or an unregistered bucket name fails boot instead of the
// first live request.
func (l *FixedWindowLimiter) ValidateThrottle(param string) error {
if l == nil {
return fmt.Errorf("surf: limiter is nil")
}
if l.store == nil {
return fmt.Errorf("surf: limiter has no store")
}
_, err := l.resolve(param)
return err
}
func (l *FixedWindowLimiter) resolve(param string) (Bucket, error) {
l.mu.Lock()
defer l.mu.Unlock()
if b, ok := l.buckets[param]; ok {
return b, nil
}
if b, ok := l.inline[param]; ok {
return b, nil
}
nStr, mStr, ok := strings.Cut(param, ",")
if !ok {
return Bucket{}, fmt.Errorf("surf: unknown throttle %q", param)
}
n, errN := strconv.Atoi(strings.TrimSpace(nStr))
m, errM := strconv.Atoi(strings.TrimSpace(mStr))
if errN != nil || errM != nil || n < 1 || m < 1 || int64(m) > math.MaxInt64/int64(time.Minute) {
return Bucket{}, fmt.Errorf("surf: malformed throttle %q", param)
}
trusted := l.trusted
b := Bucket{
Name: param,
Max: n,
Decay: time.Duration(m) * time.Minute,
Key: func(r *http.Request) string {
if r != nil {
if u, ok := bouncer.User(r.Context()); ok {
return "u:" + strconv.FormatUint(uint64(u.ID), 10)
}
}
return "inline:domainless|" + ClientIP(r, trusted)
},
}
l.inline[param] = b
return b, nil
}

View File

@@ -0,0 +1,92 @@
package surf
import (
"net/http"
"net/http/httptest"
"net/netip"
"os"
"path/filepath"
"testing"
"time"
"git.golem15.com/golem15/summercms/modules/compass"
)
// Gap (b): MemoryStore's sweep goroutine (purge), not Attempt's lazy expiry.
// Construct with a short sweep and assert the internal map drops
// an expired entry without a later Attempt.
func TestMemoryStoreSweepRemovesExpiredEntry(t *testing.T) {
s := NewMemoryStore(15 * time.Millisecond)
t.Cleanup(func() { close(s.stop) })
s.Attempt("k", 1, 25*time.Millisecond)
s.mu.Lock()
n := len(s.entries)
s.mu.Unlock()
if n != 1 {
t.Fatalf("after Attempt, entries = %d", n)
}
deadline := time.Now().Add(200 * time.Millisecond)
for time.Now().Before(deadline) {
s.mu.Lock()
n = len(s.entries)
s.mu.Unlock()
if n == 0 {
return
}
time.Sleep(10 * time.Millisecond)
}
t.Fatalf("sweep did not drop expired entry, count=%d", n)
}
func TestMemoryStoreAttemptRetryAfterAndExpiry(t *testing.T) {
s := NewMemoryStore(0)
allowed, attempts, retryAfter := s.Attempt("k", 1, 20*time.Millisecond)
if !allowed || attempts != 1 || retryAfter <= 0 {
t.Fatalf("first attempt = allowed %v, attempts %d, retryAfter %s", allowed, attempts, retryAfter)
}
allowed, attempts, retryAfter = s.Attempt("k", 1, 20*time.Millisecond)
if allowed || attempts != 1 || retryAfter <= 0 {
t.Fatalf("denied attempt = allowed %v, attempts %d, retryAfter %s", allowed, attempts, retryAfter)
}
time.Sleep(30 * time.Millisecond)
allowed, attempts, retryAfter = s.Attempt("k", 1, 20*time.Millisecond)
if !allowed || attempts != 1 || retryAfter <= 0 {
t.Fatalf("expired attempt = allowed %v, attempts %d, retryAfter %s", allowed, attempts, retryAfter)
}
}
func TestTrustedProxiesParsesConfig(t *testing.T) {
if TrustedProxies(nil) != nil {
t.Fatal("nil cfg must return nil")
}
dir := t.TempDir()
if err := os.WriteFile(filepath.Join(dir, "app.yaml"), []byte("name: t\n"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(dir, "http.yaml"), []byte("trusted_proxies:\n - 10.0.0.0/8\n - not-a-cidr\n - 192.168.0.0/16\n"), 0o644); err != nil {
t.Fatal(err)
}
cfg, err := compass.Open(compass.Options{Dir: dir, Environ: []string{}})
if err != nil {
t.Fatal(err)
}
got := TrustedProxies(cfg)
if len(got) != 2 || got[0].String() != "10.0.0.0/8" || got[1].String() != "192.168.0.0/16" {
t.Fatalf("TrustedProxies = %v", got)
}
}
func TestClientIPNilRequestAndEmptyXFF(t *testing.T) {
if got := ClientIP(nil, nil); got != "" {
t.Fatalf("nil request = %q", got)
}
trusted := []netip.Prefix{mustPrefix("10.0.0.0/8")}
req := httptest.NewRequest(http.MethodGet, "/", nil)
req.RemoteAddr = "10.0.0.1:443"
if got := ClientIP(req, trusted); got != "10.0.0.1" {
t.Fatalf("trusted RemoteAddr, empty XFF = %q", got)
}
}

View File

@@ -0,0 +1,93 @@
package surf
import (
"sync"
"time"
)
// Store owns fixed-window admission as one atomic operation. Attempt performs
// lazy expiry, threshold comparison, and an admitted increment together so
// concurrent callers cannot pass a split check-then-increment boundary.
type Store interface {
Attempt(key string, max int, decay time.Duration) (allowed bool, attempts int, retryAfter time.Duration)
}
type counterEntry struct {
count int
resetAt time.Time
}
// MemoryStore is an in-process, mutex-guarded Store.
type MemoryStore struct {
mu sync.Mutex
entries map[string]*counterEntry
sweep time.Duration
stop chan struct{}
}
// NewMemoryStore returns an in-process, mutex-guarded Store. sweep controls
// the background expired-entry cleanup interval (memory hygiene only --
// correctness does not depend on it, since expiry is checked lazily).
// A non-positive sweep disables the background goroutine.
func NewMemoryStore(sweep time.Duration) *MemoryStore {
s := &MemoryStore{
entries: make(map[string]*counterEntry),
sweep: sweep,
stop: make(chan struct{}),
}
if sweep > 0 {
go s.loop()
}
return s
}
func (s *MemoryStore) loop() {
ticker := time.NewTicker(s.sweep)
defer ticker.Stop()
for {
select {
case <-ticker.C:
s.purge()
case <-s.stop:
return
}
}
}
func (s *MemoryStore) purge() {
s.mu.Lock()
defer s.mu.Unlock()
now := time.Now()
for k, e := range s.entries {
if now.After(e.resetAt) {
delete(s.entries, k)
}
}
}
// Attempt admits and counts one request when key is below max. The first
// attempt opens the window; later attempts never extend it. A denied attempt
// leaves the exhausted count unchanged.
func (s *MemoryStore) Attempt(key string, max int, decay time.Duration) (bool, int, time.Duration) {
s.mu.Lock()
defer s.mu.Unlock()
now := time.Now()
e, ok := s.entries[key]
if ok && now.After(e.resetAt) {
delete(s.entries, key)
ok = false
}
if !ok {
e = &counterEntry{count: 0, resetAt: now.Add(decay)}
s.entries[key] = e
}
retryAfter := e.resetAt.Sub(now)
if retryAfter < 0 {
retryAfter = 0
}
if e.count >= max {
return false, e.count, retryAfter
}
e.count++
return true, e.count, retryAfter
}

View File

@@ -0,0 +1,492 @@
package surf
import (
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"git.golem15.com/golem15/summercms/modules/backpack"
"git.golem15.com/golem15/summercms/modules/bouncer"
"git.golem15.com/golem15/summercms/modules/pact"
"git.golem15.com/golem15/summercms/modules/party"
)
func TestMemoryStoreAtomicAttempt(t *testing.T) {
s := NewMemoryStore(0)
decay := 30 * time.Millisecond
allowed, attempts, retryAfter := s.Attempt("k", 1, decay)
if !allowed || attempts != 1 || retryAfter <= 0 {
t.Fatalf("first attempt = allowed %v, attempts %d, retryAfter %s", allowed, attempts, retryAfter)
}
allowed, attempts, retryAfter = s.Attempt("k", 1, decay)
if allowed || attempts != 1 || retryAfter <= 0 {
t.Fatalf("denied attempt = allowed %v, attempts %d, retryAfter %s", allowed, attempts, retryAfter)
}
time.Sleep(decay + 10*time.Millisecond)
allowed, attempts, retryAfter = s.Attempt("k", 1, decay)
if !allowed || attempts != 1 || retryAfter <= 0 {
t.Fatalf("fresh-window attempt = allowed %v, attempts %d, retryAfter %s", allowed, attempts, retryAfter)
}
}
func TestMemoryStoreConcurrentAttempt(t *testing.T) {
const workers = 32
s := NewMemoryStore(0)
ready := sync.WaitGroup{}
ready.Add(workers)
start := make(chan struct{})
results := make(chan bool, workers)
var workersDone sync.WaitGroup
workersDone.Add(workers)
for range workers {
go func() {
defer workersDone.Done()
ready.Done()
<-start
allowed, attempts, _ := s.Attempt("shared", 1, time.Minute)
if attempts != 1 {
t.Errorf("attempts = %d, want 1", attempts)
}
results <- allowed
}()
}
ready.Wait()
close(start)
workersDone.Wait()
close(results)
allowed := 0
for result := range results {
if result {
allowed++
}
}
if allowed != 1 {
t.Fatalf("allowed = %d, want 1", allowed)
}
}
func TestFixedWindowLimiterConcurrentMaxOne(t *testing.T) {
const workers = 32
lim := NewFixedWindowLimiter(NewMemoryStore(0), nil)
var handlerCalls atomic.Int32
h := lim.Middleware("1,1")(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
handlerCalls.Add(1)
w.WriteHeader(http.StatusNoContent)
}))
ready := sync.WaitGroup{}
ready.Add(workers)
start := make(chan struct{})
statuses := make(chan int, workers)
var workersDone sync.WaitGroup
workersDone.Add(workers)
for range workers {
go func() {
defer workersDone.Done()
req := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil)
req.RemoteAddr = "192.0.2.1:1234"
ready.Done()
<-start
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code == http.StatusTooManyRequests && rec.Body.String() != tooManyAttemptsBody {
t.Errorf("429 body = %q", rec.Body.String())
}
statuses <- rec.Code
}()
}
ready.Wait()
close(start)
workersDone.Wait()
close(statuses)
successes, denied := 0, 0
for status := range statuses {
switch status {
case http.StatusNoContent:
successes++
case http.StatusTooManyRequests:
denied++
default:
t.Errorf("unexpected status %d", status)
}
}
if successes != 1 || denied != workers-1 || handlerCalls.Load() != 1 {
t.Fatalf("successes=%d denied=%d handlerCalls=%d", successes, denied, handlerCalls.Load())
}
}
func TestFixedWindowLimiterMemoryStoreWindow(t *testing.T) {
s := NewMemoryStore(0)
decay := 80 * time.Millisecond
if allowed, attempts, _ := s.Attempt("k", 3, decay); !allowed || attempts != 1 {
t.Fatalf("attempt 1 = allowed %v, attempts %d", allowed, attempts)
}
if allowed, attempts, _ := s.Attempt("k", 3, decay); !allowed || attempts != 2 {
t.Fatalf("attempt 2 = allowed %v, attempts %d", allowed, attempts)
}
if allowed, attempts, _ := s.Attempt("k", 3, decay); !allowed || attempts != 3 {
t.Fatalf("attempt 3 = allowed %v, attempts %d", allowed, attempts)
}
if allowed, attempts, _ := s.Attempt("k", 3, decay); allowed || attempts != 3 {
t.Fatalf("denied attempt = allowed %v, attempts %d", allowed, attempts)
}
time.Sleep(decay + 20*time.Millisecond)
if allowed, attempts, _ := s.Attempt("k", 3, decay); !allowed || attempts != 1 {
t.Fatalf("fresh-window attempt = allowed %v, attempts %d", allowed, attempts)
}
}
func TestFixedWindowLimiterMemoryStoreFirstHitWins(t *testing.T) {
s := NewMemoryStore(0)
decay := 200 * time.Millisecond
if allowed, _, _ := s.Attempt("k", 2, decay); !allowed {
t.Fatal("first attempt denied")
}
time.Sleep(120 * time.Millisecond)
if allowed, _, _ := s.Attempt("k", 2, decay); !allowed {
t.Fatal("second attempt denied")
}
time.Sleep(100 * time.Millisecond)
if allowed, attempts, _ := s.Attempt("k", 1, decay); !allowed || attempts != 1 {
t.Fatal("window was extended by a later hit")
}
}
func TestFixedWindowLimiterSuccessHeaders(t *testing.T) {
lim := NewFixedWindowLimiter(NewMemoryStore(0), nil)
h := lim.Middleware("3,1")(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{"ok":true}`))
}))
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil)
h.ServeHTTP(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("status = %d", rec.Code)
}
if rec.Header().Get("X-RateLimit-Limit") != "3" {
t.Fatalf("limit = %q", rec.Header().Get("X-RateLimit-Limit"))
}
if rec.Header().Get("X-RateLimit-Remaining") != "2" {
t.Fatalf("remaining = %q", rec.Header().Get("X-RateLimit-Remaining"))
}
if rec.Header().Get("Retry-After") != "" || rec.Header().Get("X-RateLimit-Reset") != "" {
t.Fatal("Retry-After / X-RateLimit-Reset on success")
}
}
func TestFixedWindowLimiterTooManyAttemptsHeaders(t *testing.T) {
lim := NewFixedWindowLimiter(NewMemoryStore(0), nil)
h := lim.Middleware("1,1")(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
}))
req := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil)
ok := httptest.NewRecorder()
h.ServeHTTP(ok, req)
if ok.Code != http.StatusOK {
t.Fatalf("first status = %d", ok.Code)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusTooManyRequests {
t.Fatalf("status = %d", rec.Code)
}
if rec.Header().Get("Retry-After") == "" {
t.Fatal("missing Retry-After")
}
if rec.Header().Get("X-RateLimit-Reset") == "" {
t.Fatal("missing X-RateLimit-Reset")
}
if rec.Header().Get("X-RateLimit-Limit") != "1" {
t.Fatalf("limit = %q", rec.Header().Get("X-RateLimit-Limit"))
}
if rec.Header().Get("X-RateLimit-Remaining") != "0" {
t.Fatalf("remaining = %q", rec.Header().Get("X-RateLimit-Remaining"))
}
if rec.Body.String() != `{"message":"Too Many Attempts."}` {
t.Fatalf("body = %q", rec.Body.String())
}
var payload map[string]any
if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil {
t.Fatal(err)
}
if payload["message"] != "Too Many Attempts." {
t.Fatalf("payload = %v", payload)
}
}
func TestFixedWindowLimiterStackedBuckets(t *testing.T) {
okHandler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
})
t.Run("exhaust A", func(t *testing.T) {
lim := NewFixedWindowLimiter(NewMemoryStore(0), nil)
if err := lim.RegisterBucket("demo", "bucket-a", Bucket{
Max: 1, Decay: time.Minute,
Key: func(*http.Request) string { return "a" },
}); err != nil {
t.Fatal(err)
}
if err := lim.RegisterBucket("demo", "bucket-b", Bucket{
Max: 100, Decay: time.Minute,
Key: func(*http.Request) string { return "b" },
}); err != nil {
t.Fatal(err)
}
h := stackThrottle(lim, okHandler, "bucket-a", "bucket-b")
req := httptest.NewRequest(http.MethodGet, "/x", nil)
first := httptest.NewRecorder()
h.ServeHTTP(first, req)
if first.Code != http.StatusOK {
t.Fatalf("first = %d", first.Code)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusTooManyRequests {
t.Fatalf("exhausted A status = %d", rec.Code)
}
})
t.Run("exhaust B", func(t *testing.T) {
lim := NewFixedWindowLimiter(NewMemoryStore(0), nil)
if err := lim.RegisterBucket("demo", "bucket-a", Bucket{
Max: 100, Decay: time.Minute,
Key: func(*http.Request) string { return "a-fresh" },
}); err != nil {
t.Fatal(err)
}
if err := lim.RegisterBucket("demo", "bucket-b", Bucket{
Max: 1, Decay: time.Minute,
Key: func(*http.Request) string { return "b-only" },
}); err != nil {
t.Fatal(err)
}
h := stackThrottle(lim, okHandler, "bucket-a", "bucket-b")
req := httptest.NewRequest(http.MethodGet, "/x", nil)
first := httptest.NewRecorder()
h.ServeHTTP(first, req)
if first.Code != http.StatusOK {
t.Fatalf("first = %d", first.Code)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusTooManyRequests {
t.Fatalf("exhausted B status = %d", rec.Code)
}
})
}
func TestFixedWindowLimiterInlineThrottleKeys(t *testing.T) {
okHandler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
})
t.Run("principals differ", func(t *testing.T) {
lim := NewFixedWindowLimiter(NewMemoryStore(0), nil)
h := lim.Middleware("1,1")(okHandler)
reqA := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil)
reqA.RemoteAddr = "192.0.2.1:1"
reqA = reqA.WithContext(bouncer.WithUser(reqA.Context(), &bouncer.Principal{ID: 1}))
reqB := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil)
reqB.RemoteAddr = "192.0.2.1:1"
reqB = reqB.WithContext(bouncer.WithUser(reqB.Context(), &bouncer.Principal{ID: 2}))
a1 := httptest.NewRecorder()
h.ServeHTTP(a1, reqA)
a2 := httptest.NewRecorder()
h.ServeHTTP(a2, reqA)
if a2.Code != http.StatusTooManyRequests {
t.Fatalf("user 1 second status = %d", a2.Code)
}
b1 := httptest.NewRecorder()
h.ServeHTTP(b1, reqB)
if b1.Code != http.StatusOK {
t.Fatalf("user 2 should have a distinct key, status = %d", b1.Code)
}
})
t.Run("anonymous same IP different Host", func(t *testing.T) {
lim := NewFixedWindowLimiter(NewMemoryStore(0), nil)
h := lim.Middleware("1,1")(okHandler)
req1 := httptest.NewRequest(http.MethodGet, "http://first.example/x", nil)
req1.RemoteAddr = "192.0.2.1:1"
req2 := httptest.NewRequest(http.MethodGet, "http://second.example/x", nil)
req2.RemoteAddr = "192.0.2.1:9"
first := httptest.NewRecorder()
h.ServeHTTP(first, req1)
if first.Code != http.StatusOK {
t.Fatalf("first = %d", first.Code)
}
second := httptest.NewRecorder()
h.ServeHTTP(second, req2)
if second.Code != http.StatusTooManyRequests {
t.Fatalf("anonymous same IP with rotated Host should share key, status = %d", second.Code)
}
})
t.Run("anonymous inline policies share a domainless key", func(t *testing.T) {
lim := NewFixedWindowLimiter(NewMemoryStore(0), nil)
twoPerMinute := lim.Middleware("2,1")(okHandler)
onePerMinute := lim.Middleware("1,1")(okHandler)
request := func(h http.Handler) *httptest.ResponseRecorder {
req := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil)
req.RemoteAddr = "192.0.2.9:1234"
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
return rec
}
if rec := request(twoPerMinute); rec.Code != http.StatusOK {
t.Fatalf("first throttle:2,1 status = %d", rec.Code)
}
if rec := request(twoPerMinute); rec.Code != http.StatusOK {
t.Fatalf("second throttle:2,1 status = %d", rec.Code)
}
if rec := request(onePerMinute); rec.Code != http.StatusTooManyRequests {
t.Fatalf("throttle:1,1 after shared exhaustion status = %d, want 429", rec.Code)
}
})
}
func TestFixedWindowLimiterUnknownBucketFailsAssemble(t *testing.T) {
p := assemblePlugin{id: "golem15.demo", use: []string{"throttle:missing"}}
_, err := Assemble(backpack.New(nil), []party.Plugin{p})
if err == nil || !strings.Contains(err.Error(), "missing") {
t.Fatalf("want unknown throttle in error, got %v", err)
}
}
func TestFixedWindowLimiterDuplicateBucket(t *testing.T) {
lim := NewFixedWindowLimiter(NewMemoryStore(0), nil)
b := Bucket{Max: 1, Decay: time.Minute, Key: func(*http.Request) string { return "k" }}
if err := lim.RegisterBucket("one", "shared", b); err != nil {
t.Fatal(err)
}
err := lim.RegisterBucket("two", "shared", b)
if err == nil || !strings.Contains(err.Error(), "shared") || !strings.Contains(err.Error(), "one") {
t.Fatalf("got %v", err)
}
}
func stackThrottle(lim *FixedWindowLimiter, next http.Handler, names ...string) http.Handler {
h := next
for i := len(names) - 1; i >= 0; i-- {
h = lim.Middleware(names[i])(h)
}
return h
}
var _ pact.Router = (*Router)(nil)
func TestRegisterBucketRejectsInvalidDefinitions(t *testing.T) {
l := NewFixedWindowLimiter(NewMemoryStore(time.Minute), nil)
key := func(*http.Request) string { return "k" }
if err := l.RegisterBucket("p", "n", Bucket{Max: 1, Decay: 0, Key: key}); err == nil {
t.Fatal("zero decay accepted")
}
if err := l.RegisterBucket("p", "n", Bucket{Max: 1, Decay: time.Minute}); err == nil {
t.Fatal("nil key accepted")
}
if err := l.ValidateThrottle("1,9223372036854775807"); err == nil {
t.Fatal("overflowing minutes accepted")
}
}
func TestRegisterBucketRejectsInvalid(t *testing.T) {
key := func(*http.Request) string { return "k" }
cases := []struct {
name string
store Store
b Bucket
}{
{"nil key", NewMemoryStore(time.Minute), Bucket{Max: 1, Decay: time.Minute}},
{"max zero", NewMemoryStore(time.Minute), Bucket{Max: 0, Decay: time.Minute, Key: key}},
{"max negative", NewMemoryStore(time.Minute), Bucket{Max: -1, Decay: time.Minute, Key: key}},
{"decay zero", NewMemoryStore(time.Minute), Bucket{Max: 1, Key: key}},
{"decay negative", NewMemoryStore(time.Minute), Bucket{Max: 1, Decay: -time.Second, Key: key}},
{"nil store", nil, Bucket{Max: 1, Decay: time.Minute, Key: key}},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
l := NewFixedWindowLimiter(tc.store, nil)
err := l.RegisterBucket("golem15.p", "bkt", tc.b)
if err == nil || !strings.Contains(err.Error(), "golem15.p") || !strings.Contains(err.Error(), "bkt") {
t.Fatalf("err = %v", err)
}
})
}
}
func TestValidateThrottleRejectsOverflowAndNilStore(t *testing.T) {
l := NewFixedWindowLimiter(NewMemoryStore(time.Minute), nil)
for _, p := range []string{"1,9223372036854775807", "0,1", "1,0", "-1,1", "x,y", "nope"} {
if err := l.ValidateThrottle(p); err == nil {
t.Errorf("%q accepted", p)
}
}
if err := l.ValidateThrottle("5,1"); err != nil {
t.Fatal(err)
}
if err := NewFixedWindowLimiter(nil, nil).ValidateThrottle("5,1"); err == nil {
t.Fatal("nil store accepted")
}
var nilLim *FixedWindowLimiter
if err := nilLim.ValidateThrottle("5,1"); err == nil {
t.Fatal("nil limiter accepted")
}
}
func TestMiddlewareFailsClosed(t *testing.T) {
var nilLim *FixedWindowLimiter
cases := map[string]*FixedWindowLimiter{
"unknown bucket": NewFixedWindowLimiter(NewMemoryStore(time.Minute), nil),
"nil store": NewFixedWindowLimiter(nil, nil),
"nil limiter": nilLim,
}
params := map[string]string{"unknown bucket": "missing", "nil store": "5,1", "nil limiter": "5,1"}
for name, l := range cases {
t.Run(name, func(t *testing.T) {
called := false
h := l.Middleware(params[name])(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
called = true
w.WriteHeader(http.StatusNoContent)
}))
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil))
if called || rec.Code != http.StatusInternalServerError {
t.Fatalf("called=%v code=%d", called, rec.Code)
}
})
}
}
func TestRegisterBucketRejectsNilLimiterEmptyNameAndDuplicate(t *testing.T) {
key := func(*http.Request) string { return "k" }
good := Bucket{Max: 1, Decay: time.Minute, Key: key}
var nilLim *FixedWindowLimiter
if err := nilLim.RegisterBucket("p", "n", good); err == nil {
t.Fatal("nil limiter accepted")
}
l := NewFixedWindowLimiter(NewMemoryStore(time.Minute), nil)
if err := l.RegisterBucket("p", "", good); err == nil {
t.Fatal("empty name accepted")
}
if err := l.RegisterBucket("p", "n", good); err != nil {
t.Fatal(err)
}
if err := l.RegisterBucket("q", "n", good); err == nil || !strings.Contains(err.Error(), "p") {
t.Fatalf("duplicate: %v", err)
}
}

View File

@@ -0,0 +1,19 @@
package surf
import (
"net/http"
"git.golem15.com/golem15/summercms/modules/bouncer"
"git.golem15.com/golem15/summercms/modules/towel"
)
// LocaleFromPrincipal overrides the header locale with the authenticated
// principal's PreferredLocale when that value is non-empty.
func LocaleFromPrincipal(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if p, ok := bouncer.User(r.Context()); ok && p.PreferredLocale != "" {
r = r.WithContext(towel.WithLocale(r.Context(), p.PreferredLocale))
}
next.ServeHTTP(w, r)
})
}

View File

@@ -0,0 +1,53 @@
package surf
import (
"net/http"
"net/http/httptest"
"testing"
"git.golem15.com/golem15/summercms/modules/bouncer"
"git.golem15.com/golem15/summercms/modules/towel"
)
func TestLocaleFromPrincipal(t *testing.T) {
var got string
var ok bool
h := LocaleFromPrincipal(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
got, ok = towel.Locale(r.Context())
}))
ctx := towel.WithLocale(t.Context(), "en")
ctx = bouncer.WithUser(ctx, &bouncer.Principal{PreferredLocale: "pl"})
req := httptest.NewRequest(http.MethodGet, "/", nil).WithContext(ctx)
h.ServeHTTP(httptest.NewRecorder(), req)
if !ok || got != "pl" {
t.Fatalf("locale = %q ok=%t", got, ok)
}
got, ok = "", false
ctx = towel.WithLocale(t.Context(), "en")
ctx = bouncer.WithUser(ctx, &bouncer.Principal{})
req = httptest.NewRequest(http.MethodGet, "/", nil).WithContext(ctx)
h.ServeHTTP(httptest.NewRecorder(), req)
if !ok || got != "en" {
t.Fatalf("empty preferred locale = %q ok=%t", got, ok)
}
got, ok = "", false
req = httptest.NewRequest(http.MethodGet, "/", nil).WithContext(towel.WithLocale(t.Context(), "de"))
h.ServeHTTP(httptest.NewRecorder(), req)
if !ok || got != "de" {
t.Fatalf("no principal locale = %q ok=%t", got, ok)
}
}
func TestBuildRouterRegistersLocaleFromPrincipal(t *testing.T) {
r, err := BuildRouter(nil, nil)
if err != nil {
t.Fatal(err)
}
mw, ok := r.named["locale.from-principal"]
if !ok || mw.pluginID != "surf" || mw.fn == nil {
t.Fatalf("registration ok=%t mw=%+v", ok, mw)
}
}

View File

@@ -0,0 +1,254 @@
package surf
import (
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"git.golem15.com/golem15/summercms/modules/backpack"
"git.golem15.com/golem15/summercms/modules/pact"
"git.golem15.com/golem15/summercms/modules/party"
"git.golem15.com/golem15/summercms/modules/towel"
)
type assemblePlugin struct {
id string
mw map[string]pact.Middleware
use []string
path string
hit *bool
}
func (p assemblePlugin) ID() string { return p.id }
func (p assemblePlugin) Requires() []string { return nil }
func (p assemblePlugin) Register(*backpack.App) error { return nil }
func (p assemblePlugin) Boot(*backpack.App) error { return nil }
func (p assemblePlugin) Middlewares() map[string]pact.Middleware { return p.mw }
func (p assemblePlugin) Routes(r pact.Router) error {
path := p.path
if path == "" {
path = "/items"
}
r.Group("/api", Use(p.use...), func(g pact.Router) {
g.Get(path, func(w http.ResponseWriter, r *http.Request) {
if p.hit != nil {
*p.hit = true
}
w.WriteHeader(http.StatusOK)
})
})
return nil
}
func TestAssembleMissingMiddlewareFailsBoot(t *testing.T) {
p := assemblePlugin{id: "golem15.demo", use: []string{"jwt.auth"}}
_, err := Assemble(backpack.New(nil), []party.Plugin{p})
if err == nil || !strings.Contains(err.Error(), "golem15.demo") || !strings.Contains(err.Error(), "jwt.auth") {
t.Fatalf("want plugin and middleware in error, got %v", err)
}
}
func TestUnauthenticatedNamedGuardDoesNotReachHandler(t *testing.T) {
hit := false
p := assemblePlugin{
id: "golem15.demo",
use: []string{"jwt.auth"},
hit: &hit,
mw: map[string]pact.Middleware{
"jwt.auth": func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusUnauthorized)
_, _ = w.Write([]byte(`{"error":true,"message":"Token not provided"}`))
})
},
},
}
h, err := Assemble(backpack.New(nil), []party.Plugin{p})
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/items", nil))
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d", rec.Code)
}
if hit {
t.Fatal("unauthenticated request reached the handler")
}
}
func TestPipelineOrderRecoverCORSLocaleAuthPasswordOrgRateHandler(t *testing.T) {
var order []string
record := func(name string) pact.Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
order = append(order, name)
next.ServeHTTP(w, r)
})
}
}
r := New(nil)
r.corsCfg = CORSConfig{
Paths: []string{"api/*"},
AllowedOrigins: []string{"http://localhost:3000"},
AllowedMethods: []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"},
AllowedHeaders: []string{"Authorization", "Content-Type", "Accept"},
}
if err := r.RegisterMiddleware("golem15.user", "jwt.auth", record("jwt.auth")); err != nil {
t.Fatal(err)
}
if err := r.RegisterMiddleware("golem15.acme", "inv.must-change-password", record("inv.must-change-password")); err != nil {
t.Fatal(err)
}
r.BindPlugin("golem15.demo")
r.Group("/api", Use("jwt.auth", "inv.must-change-password"), func(g pact.Router) {
g.Get("/items", func(w http.ResponseWriter, req *http.Request) {
loc, ok := towel.Locale(req.Context())
if !ok || loc != "pl" {
t.Errorf("locale = %q ok=%t, want pl from Accept-Language", loc, ok)
}
if _, ok := towel.Organization(req.Context()); !ok {
t.Error("org slot must run before the handler")
}
order = append(order, "handler")
w.WriteHeader(http.StatusNoContent)
})
})
h, err := r.compile()
if err != nil {
t.Fatal(err)
}
t.Run("preflight-before-auth", func(t *testing.T) {
order = nil
req := httptest.NewRequest(http.MethodOptions, "/api/items", nil)
req.Header.Set("Origin", "http://localhost:3000")
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusNoContent {
t.Fatalf("status = %d", rec.Code)
}
if rec.Header().Get("Access-Control-Allow-Origin") != "http://localhost:3000" {
t.Fatalf("ACA origin = %q", rec.Header().Get("Access-Control-Allow-Origin"))
}
if len(order) != 0 {
t.Fatalf("named stages ran on preflight: %v", order)
}
})
t.Run("get-order", func(t *testing.T) {
order = nil
req := httptest.NewRequest(http.MethodGet, "/api/items", nil)
req.Header.Set("Accept-Language", "pl")
req.Header.Set("Origin", "http://localhost:3000")
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusNoContent {
t.Fatalf("status = %d", rec.Code)
}
want := []string{"jwt.auth", "inv.must-change-password", "handler"}
if strings.Join(order, ",") != strings.Join(want, ",") {
t.Fatalf("order = %v want %v", order, want)
}
})
t.Run("panic-still-opaque", func(t *testing.T) {
panicRouter := New(nil)
panicRouter.Get("/boom", func(http.ResponseWriter, *http.Request) {
panic("stack-trace-secret")
})
ph, err := panicRouter.compile()
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
ph.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/boom", nil))
if rec.Code != http.StatusInternalServerError {
t.Fatalf("status = %d", rec.Code)
}
body := rec.Body.String()
if strings.Contains(body, "stack-trace-secret") {
t.Fatalf("leaked panic: %s", body)
}
var payload map[string]any
if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil {
t.Fatal(err)
}
if payload["error"] != true || payload["message"] != "Internal server error" {
t.Fatalf("payload = %v", payload)
}
})
}
func TestGroupAndPerRouteMiddlewareCompose(t *testing.T) {
var order []string
record := func(name string) pact.Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
order = append(order, name)
next.ServeHTTP(w, r)
})
}
}
r := New(nil)
if err := r.RegisterMiddleware("golem15.user", "jwt.auth", record("jwt.auth")); err != nil {
t.Fatal(err)
}
if err := r.RegisterMiddleware("golem15.demo", "audit", record("audit")); err != nil {
t.Fatal(err)
}
r.BindPlugin("golem15.demo")
r.Group("/api", Use("jwt.auth"), func(g pact.Router) {
g.Get("/plain", func(w http.ResponseWriter, r *http.Request) {
order = append(order, "plain")
w.WriteHeader(http.StatusNoContent)
})
g.Get("/audited", func(w http.ResponseWriter, r *http.Request) {
order = append(order, "audited")
w.WriteHeader(http.StatusNoContent)
}, "audit")
})
h, err := r.compile()
if err != nil {
t.Fatal(err)
}
order = nil
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/plain", nil))
if rec.Code != http.StatusNoContent || strings.Join(order, ",") != "jwt.auth,plain" {
t.Fatalf("plain = %d %v", rec.Code, order)
}
order = nil
rec = httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/audited", nil))
if rec.Code != http.StatusNoContent || strings.Join(order, ",") != "jwt.auth,audit,audited" {
t.Fatalf("audited = %d %v", rec.Code, order)
}
}
func TestServeMuxRejectsWrongMethod(t *testing.T) {
r := New(nil)
r.Get("/only-get", func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
})
h, err := r.compile()
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/only-get", nil))
if rec.Code != http.StatusMethodNotAllowed && rec.Code != http.StatusNotFound {
t.Fatalf("POST status = %d", rec.Code)
}
rec = httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/only-get", nil))
if rec.Code != http.StatusOK {
t.Fatalf("GET status = %d", rec.Code)
}
}

98
modules/surf/params.go Normal file
View File

@@ -0,0 +1,98 @@
package surf
import (
"fmt"
"net/http"
"regexp"
"strconv"
"strings"
)
// Constraint is a path-parameter restriction matching PHP ->where() /
// ->whereIn(). Patterns and enum sets are compiled at route registration;
// request text is only matched, never interpolated into SQL or regex.
type Constraint struct {
param string
re *regexp.Regexp
enum map[string]struct{}
}
// 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
}
// Regex compiles a PHP-style where() pattern for param at registration.
func Regex(param, pattern string) (Constraint, error) {
if param == "" {
return Constraint{}, fmt.Errorf("surf: constraint param is empty")
}
if pattern == "" {
return Constraint{}, fmt.Errorf("surf: regex for %q is empty", param)
}
re, err := regexp.Compile("^(?:" + pattern + ")$")
if err != nil {
return Constraint{}, fmt.Errorf("surf: invalid regex for %q: %w", param, err)
}
return Constraint{param: param, re: re}, nil
}
// Enum allow-lists exact path values for param, matching PHP whereIn().
func Enum(param string, values ...string) (Constraint, error) {
if param == "" {
return Constraint{}, fmt.Errorf("surf: constraint param is empty")
}
if len(values) == 0 {
return Constraint{}, fmt.Errorf("surf: enum for %q is empty", param)
}
set := make(map[string]struct{}, len(values))
for _, v := range values {
if v == "" {
return Constraint{}, fmt.Errorf("surf: enum for %q contains an empty value", param)
}
set[v] = struct{}{}
}
return Constraint{param: param, enum: set}, nil
}
func (c Constraint) match(value string) bool {
if c.enum != nil {
_, ok := c.enum[value]
return ok
}
if c.re != nil {
return c.re.MatchString(value)
}
return false
}
func constrain(next http.Handler, cs []Constraint) http.Handler {
if len(cs) == 0 {
return next
}
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
for _, c := range cs {
if !c.match(r.PathValue(c.param)) {
http.NotFound(w, r)
return
}
}
next.ServeHTTP(w, r)
})
}
func pathHasParam(path, param string) bool {
return strings.Contains(path, "{"+param+"}") || strings.Contains(path, "{"+param+":")
}

View File

@@ -0,0 +1,63 @@
package surf
import (
"net/http"
"net/http/httptest"
"testing"
)
func TestIntParamMalformedIsFalse(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/items/abc", nil)
req.SetPathValue("id", "abc")
if _, ok := IntParam(req, "id"); ok {
t.Fatal("malformed id must be false")
}
req.SetPathValue("id", "12")
n, ok := IntParam(req, "id")
if !ok || n != 12 {
t.Fatalf("got %d %t", n, ok)
}
req.SetPathValue("id", "0")
if _, ok := IntParam(req, "id"); ok {
t.Fatal("non-positive id must be false")
}
req.SetPathValue("id", "")
if _, ok := IntParam(req, "id"); ok {
t.Fatal("missing id must be false")
}
if _, ok := IntParam(nil, "id"); ok {
t.Fatal("nil request must be false")
}
}
func TestRegexCompilesAtRegistration(t *testing.T) {
c, err := Regex("id", `[0-9]+`)
if err != nil {
t.Fatal(err)
}
if !c.match("12") || c.match("nope") || c.match("") {
t.Fatal("regex match")
}
if _, err := Regex("id", "["); err == nil {
t.Fatal("want invalid regex error")
}
if _, err := Regex("", `[0-9]+`); err == nil {
t.Fatal("want empty param error")
}
}
func TestEnumAllowListDoesNotBuildRegexFromValues(t *testing.T) {
c, err := Enum("kind", "widget", "gadget")
if err != nil {
t.Fatal(err)
}
if c.re != nil {
t.Fatal("enum must not compile a regex")
}
if !c.match("widget") || !c.match("gadget") || c.match("other") || c.match("widget|gadget") {
t.Fatal("enum match")
}
if _, err := Enum("kind"); err == nil {
t.Fatal("want empty enum error")
}
}

View File

@@ -0,0 +1,39 @@
package surf
import (
"context"
"strconv"
"strings"
"git.golem15.com/golem15/summercms/modules/backpack"
"git.golem15.com/golem15/summercms/modules/bonfire"
"git.golem15.com/golem15/summercms/modules/party"
)
// RouteListCommand builds the router the same way ServeCommand does, then
// renders Routes() as a table instead of serving.
func RouteListCommand(app *backpack.App, plugins []party.Plugin) bonfire.Command {
return bonfire.Command{
Name: "route:list",
Description: "List registered HTTP routes",
Run: func(ctx context.Context, in bonfire.Input, out bonfire.Output) error {
r, err := BuildRouter(app, plugins)
if err != nil {
return err
}
infos := r.Routes()
rows := make([][]string, 0, len(infos))
for _, rt := range infos {
rows = append(rows, []string{
rt.Method,
rt.Pattern,
rt.PluginID,
strings.Join(rt.Middleware, ","),
strconv.FormatBool(rt.Raw),
})
}
out.Table([]string{"Method", "Pattern", "Plugin", "Middleware", "Raw"}, rows)
return nil
},
}
}

720
modules/surf/router.go Normal file
View File

@@ -0,0 +1,720 @@
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
}

534
modules/surf/router_test.go Normal file
View File

@@ -0,0 +1,534 @@
package surf
import (
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"git.golem15.com/golem15/summercms/modules/backpack"
"git.golem15.com/golem15/summercms/modules/pact"
"git.golem15.com/golem15/summercms/modules/party"
)
type routePlugin struct {
id string
mw map[string]pact.Middleware
path string
use []string
}
func (p routePlugin) ID() string { return p.id }
func (p routePlugin) Requires() []string { return nil }
func (p routePlugin) Register(any) error { return nil }
func (p routePlugin) Boot(any) error { return nil }
func (p routePlugin) Middlewares() map[string]pact.Middleware { return p.mw }
func (p routePlugin) Routes(r pact.Router) error {
r.Group("/api", Use(p.use...), func(g pact.Router) {
g.Get(p.path, func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Hit", "1")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{"ok":true}`))
})
})
return nil
}
func TestMissingMiddlewareNamesPluginAndName(t *testing.T) {
r := New(nil)
p := routePlugin{id: "golem15.demo", path: "/items", use: []string{"jwt.auth"}}
r.BindPlugin(p.ID())
if err := p.Routes(r); err != nil {
t.Fatal(err)
}
_, err := r.compile()
if err == nil || !strings.Contains(err.Error(), "golem15.demo") || !strings.Contains(err.Error(), "jwt.auth") {
t.Fatalf("want plugin and middleware in error, got %v", err)
}
}
func TestRawGroupPanicBare500(t *testing.T) {
r := New(nil)
r.GroupRaw("/oauth", nil, func(g pact.Router) {
g.Get("/panic", func(http.ResponseWriter, *http.Request) {
panic("secret internals")
})
})
r.Get("/panic", func(http.ResponseWriter, *http.Request) {
panic("secret internals")
})
h, err := r.compile()
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/oauth/panic", nil))
if rec.Code != http.StatusInternalServerError {
t.Fatalf("raw status = %d", rec.Code)
}
if rec.Body.Len() != 0 {
t.Fatalf("raw body = %q", rec.Body.String())
}
if ct := rec.Header().Get("Content-Type"); ct != "" {
t.Fatalf("raw Content-Type = %q", ct)
}
rec = httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/panic", nil))
if rec.Code != http.StatusInternalServerError {
t.Fatalf("house status = %d", rec.Code)
}
if rec.Header().Get("Content-Type") != "application/json" {
t.Fatalf("house Content-Type = %q", rec.Header().Get("Content-Type"))
}
var payload map[string]any
if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil {
t.Fatal(err)
}
if payload["error"] != true || payload["message"] != "Internal server error" {
t.Fatalf("payload = %v", payload)
}
}
func TestRecoverDiscardsPartialResponse(t *testing.T) {
const secretBody = "secret-partial"
t.Run("house", func(t *testing.T) {
r := New(nil)
r.Get("/panic", func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("X-Partial", "secret")
w.WriteHeader(http.StatusAccepted)
_, _ = w.Write([]byte(secretBody))
panic("secret internals")
})
h, err := r.compile()
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/panic", nil))
if rec.Code != http.StatusInternalServerError {
t.Fatalf("status = %d", rec.Code)
}
if got := rec.Header().Get("Content-Type"); got != "application/json" {
t.Fatalf("Content-Type = %q", got)
}
if got := rec.Header().Get("X-Partial"); got != "" {
t.Fatalf("X-Partial leaked: %q", got)
}
const want = `{"error":true,"message":"Internal server error"}`
if got := rec.Body.String(); got != want {
t.Fatalf("body = %q, want %q", got, want)
}
if strings.Contains(rec.Body.String(), secretBody) {
t.Fatalf("partial body leaked: %q", rec.Body.String())
}
})
t.Run("raw", func(t *testing.T) {
r := New(nil)
r.GroupRaw("/oauth", nil, func(g pact.Router) {
g.Get("/panic", func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("X-Partial", "secret")
w.WriteHeader(http.StatusAccepted)
_, _ = w.Write([]byte(secretBody))
panic("secret internals")
})
})
h, err := r.compile()
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/oauth/panic", nil))
if rec.Code != http.StatusInternalServerError {
t.Fatalf("status = %d", rec.Code)
}
if rec.Body.Len() != 0 {
t.Fatalf("body = %q", rec.Body.String())
}
if got := rec.Header().Get("Content-Type"); got != "" {
t.Fatalf("Content-Type = %q", got)
}
if got := rec.Header().Get("X-Partial"); got != "" {
t.Fatalf("X-Partial leaked: %q", got)
}
if strings.Contains(rec.Body.String(), secretBody) {
t.Fatalf("partial body leaked: %q", rec.Body.String())
}
})
}
func TestBufferedResponseCommitsSuccessfulOutput(t *testing.T) {
r := New(nil)
r.Get("/explicit", func(w http.ResponseWriter, _ *http.Request) {
w.Header().Add("X-Result", "one")
w.Header().Add("X-Result", "two")
w.WriteHeader(http.StatusCreated)
w.WriteHeader(http.StatusTeapot)
_, _ = w.Write([]byte("created"))
})
r.Get("/implicit", func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("X-Result", "implicit")
_, _ = w.Write([]byte("ok"))
})
h, err := r.compile()
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/explicit", nil))
if rec.Code != http.StatusCreated {
t.Fatalf("explicit status = %d", rec.Code)
}
if got := rec.Header().Values("X-Result"); len(got) != 2 || got[0] != "one" || got[1] != "two" {
t.Fatalf("explicit X-Result = %q", got)
}
if got := rec.Body.String(); got != "created" {
t.Fatalf("explicit body = %q", got)
}
rec = httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/implicit", nil))
if rec.Code != http.StatusOK {
t.Fatalf("implicit status = %d", rec.Code)
}
if got := rec.Header().Get("X-Result"); got != "implicit" {
t.Fatalf("implicit X-Result = %q", got)
}
if got := rec.Body.String(); got != "ok" {
t.Fatalf("implicit body = %q", got)
}
}
func TestHouseMiddlewareDuplicateNameFailsBoot(t *testing.T) {
identity := func(next http.Handler) http.Handler { return next }
named := assemblePlugin{
id: "golem15.one",
mw: map[string]pact.Middleware{"shared.mw": identity},
use: []string{},
}
house := houseRoutePlugin{
id: "golem15.two",
house: map[string]pact.Middleware{
"shared.mw": identity,
},
}
_, err := BuildRouter(backpack.New(nil), []party.Plugin{named, house})
if err == nil {
t.Fatal("want duplicate-name error")
}
if !strings.Contains(err.Error(), "shared.mw") || !strings.Contains(err.Error(), "golem15.one") {
t.Fatalf("want existing duplicate-name error, got %v", err)
}
}
func TestRecoverReturnsOpaqueJSON500(t *testing.T) {
r := New(nil)
r.Get("/panic", func(http.ResponseWriter, *http.Request) {
panic("secret internals")
})
h, err := r.compile()
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/panic", nil))
if rec.Code != http.StatusInternalServerError {
t.Fatalf("status = %d", rec.Code)
}
body := rec.Body.String()
if strings.Contains(body, "secret") {
t.Fatalf("leaked panic: %s", body)
}
var payload map[string]any
if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil {
t.Fatal(err)
}
if payload["error"] != true || payload["message"] != "Internal server error" {
t.Fatalf("payload = %v", payload)
}
}
func TestCORSPreflightBypassesNamedAuth(t *testing.T) {
called := false
r := New(nil)
r.corsCfg = CORSConfig{
Paths: []string{"api/*"},
AllowedOrigins: []string{"http://localhost:3000"},
AllowedMethods: []string{"GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"},
AllowedHeaders: []string{"Authorization", "Content-Type", "Accept"},
}
if err := r.RegisterMiddleware("golem15.user", "jwt.auth", func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
called = true
next.ServeHTTP(w, req)
})
}); err != nil {
t.Fatal(err)
}
r.BindPlugin("golem15.demo")
r.Group("/api", Use("jwt.auth"), func(g pact.Router) {
g.Get("/items", func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
})
})
h, err := r.compile()
if err != nil {
t.Fatal(err)
}
req := httptest.NewRequest(http.MethodOptions, "/api/items", nil)
req.Header.Set("Origin", "http://localhost:3000")
rec := httptest.NewRecorder()
h.ServeHTTP(rec, req)
if rec.Code != http.StatusNoContent {
t.Fatalf("status = %d", rec.Code)
}
if called {
t.Fatal("jwt.auth ran on preflight")
}
if rec.Header().Get("Access-Control-Allow-Origin") != "http://localhost:3000" {
t.Fatalf("ACA origin = %q", rec.Header().Get("Access-Control-Allow-Origin"))
}
}
func TestTypedIDRouteReturns404(t *testing.T) {
r := New(nil)
r.Get("/items/{id}", func(w http.ResponseWriter, req *http.Request) {
id, ok := IntParam(req, "id")
if !ok || id != 1 {
http.NotFound(w, req)
return
}
w.WriteHeader(http.StatusOK)
})
r.Where("id", `[0-9]+`)
h, err := r.compile()
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/nope", nil))
if rec.Code != http.StatusNotFound {
t.Fatalf("malformed status = %d", rec.Code)
}
rec = httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/99", nil))
if rec.Code != http.StatusNotFound {
t.Fatalf("unknown status = %d", rec.Code)
}
rec = httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/1", nil))
if rec.Code != http.StatusOK {
t.Fatalf("known status = %d", rec.Code)
}
}
func TestWhereInRejectsOutsideEnum(t *testing.T) {
r := New(nil)
r.Get("/kinds/{kind}", func(w http.ResponseWriter, req *http.Request) {
w.WriteHeader(http.StatusOK)
})
r.WhereIn("kind", "widget", "gadget")
h, err := r.compile()
if err != nil {
t.Fatal(err)
}
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/kinds/other", nil))
if rec.Code != http.StatusNotFound {
t.Fatalf("status = %d", rec.Code)
}
rec = httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/kinds/widget", nil))
if rec.Code != http.StatusOK {
t.Fatalf("status = %d", rec.Code)
}
}
func TestWhereInvalidRegexFailsCompile(t *testing.T) {
r := New(nil)
r.Get("/items/{id}", func(http.ResponseWriter, *http.Request) {})
r.Where("id", "[")
if _, err := r.compile(); err == nil {
t.Fatal("want compile error for invalid regex")
}
}
func TestWhereUnknownParamFailsCompile(t *testing.T) {
r := New(nil)
r.Get("/items/{id}", func(http.ResponseWriter, *http.Request) {})
r.Where("slug", `[a-z]+`)
if _, err := r.compile(); err == nil || !strings.Contains(err.Error(), "slug") {
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.acme", "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.acme", "inv.scope", fn); err != nil {
t.Fatal(err)
}
err := r.RegisterMiddlewareFactory("golem15.other", "inv.scope", fn)
if err == nil || !strings.Contains(err.Error(), "golem15.acme") || !strings.Contains(err.Error(), "inv.scope") {
t.Fatalf("want plugin and factory name in error, got %v", err)
}
}
func TestBuildRouterFailsOnMissingBodyConfig(t *testing.T) {
bad := map[string]string{
"missing default": "body_limits:\n upload_bytes: 8\n",
"missing upload": "body_limits:\n default_bytes: 8\n",
"zero": "body_limits:\n default_bytes: 0\n upload_bytes: 8\n",
"negative": "body_limits:\n default_bytes: 4\n upload_bytes: -1\n",
"non-numeric": "body_limits:\n default_bytes: lots\n upload_bytes: 8\n",
}
for name, y := range bad {
t.Run(name, func(t *testing.T) {
if _, err := BuildRouter(backpack.New(writeHTTPConfig(t, y)), nil); err == nil {
t.Fatal("expected error")
}
})
}
if _, err := BuildRouter(backpack.New(writeHTTPConfig(t, "body_limits:\n default_bytes: 4\n upload_bytes: 8\n")), nil); err != nil {
t.Fatal(err)
}
}
type factoryCountPlugin struct{ calls *map[string]int }
func (p factoryCountPlugin) ID() string { return "golem15.count" }
func (p factoryCountPlugin) Requires() []string { return nil }
func (p factoryCountPlugin) Register(*backpack.App) error { return nil }
func (p factoryCountPlugin) Boot(*backpack.App) error { return nil }
func (p factoryCountPlugin) MiddlewareFactories() map[string]func(string) pact.Middleware {
return map[string]func(string) pact.Middleware{"count": func(param string) pact.Middleware {
(*p.calls)[param]++
return func(next http.Handler) http.Handler { return next }
}}
}
func (p factoryCountPlugin) Routes(r pact.Router) error {
r.Group("/", Use("count:a"), func(g pact.Router) {
g.Get("/one", func(http.ResponseWriter, *http.Request) {})
g.Get("/two", func(http.ResponseWriter, *http.Request) {})
g.Get("/three", func(http.ResponseWriter, *http.Request) {}, "count:b")
})
return nil
}
func TestFactoriesBuiltOncePerName(t *testing.T) {
calls := map[string]int{}
cfg := writeHTTPConfig(t, "body_limits:\n default_bytes: 4\n upload_bytes: 8\n")
if _, err := Assemble(backpack.New(cfg), []party.Plugin{factoryCountPlugin{calls: &calls}}); err != nil {
t.Fatal(err)
}
if calls["a"] != 1 || calls["b"] != 1 {
t.Fatalf("factory invocations = %v, want each name:param built once", calls)
}
}

View File

@@ -0,0 +1,30 @@
package surf
// RouteInfo is a read-only snapshot of one registered route after Assemble
// or BuildRouter.
type RouteInfo struct {
Method string
Pattern string
PluginID string
Middleware []string
Raw bool
}
// Routes returns a defensive copy of every registered route.
func (r *Router) Routes() []RouteInfo {
if r == nil {
return nil
}
out := make([]RouteInfo, len(r.routes))
for i, rt := range r.routes {
mw := append([]string{}, rt.middleware...)
out[i] = RouteInfo{
Method: rt.method,
Pattern: rt.path,
PluginID: rt.pluginID,
Middleware: mw,
Raw: rt.raw,
}
}
return out
}

View File

@@ -0,0 +1,83 @@
package surf
import (
"net/http"
"strings"
"testing"
"git.golem15.com/golem15/summercms/modules/pact"
)
// Gap (c): RegisterHouseMiddlewareFactory duplicate-name failure. The
// non-house RegisterMiddlewareFactory duplicate path is already covered
// by TestDuplicateMiddlewareFactoryNamesPluginAndName.
func TestDuplicateHouseMiddlewareFactoryNamesPluginAndName(t *testing.T) {
r := New(nil)
fn := func(string) pact.Middleware {
return func(next http.Handler) http.Handler { return next }
}
if err := r.RegisterHouseMiddlewareFactory("golem15.demo", "house.param", fn); err != nil {
t.Fatal(err)
}
err := r.RegisterHouseMiddlewareFactory("golem15.other", "house.param", fn)
if err == nil || !strings.Contains(err.Error(), "golem15.demo") || !strings.Contains(err.Error(), "house.param") {
t.Fatalf("want plugin and factory name in error, got %v", err)
}
}
func TestRegisterHouseMiddlewareFactoryEmptyName(t *testing.T) {
r := New(nil)
err := r.RegisterHouseMiddlewareFactory("golem15.demo", "", func(string) pact.Middleware {
return func(next http.Handler) http.Handler { return next }
})
if err == nil || !strings.Contains(err.Error(), "golem15.demo") {
t.Fatalf("empty factory name: %v", err)
}
}
// Gap (g): Router.Routes() on an empty router, and an empty raw group
// (zero routes). Raw:true is inspectable on the Group itself; Routes()
// only lists registered handlers, so an empty raw group produces no
// RouteInfo — that is structural, not a missing assertion.
func TestRoutesEmptyRouter(t *testing.T) {
r := New(nil)
got := r.Routes()
if got == nil {
t.Fatal("empty router Routes() must return a non-nil empty slice")
}
if len(got) != 0 {
t.Fatalf("empty router Routes() = %#v", got)
}
if (*Router)(nil).Routes() != nil {
t.Fatal("nil router Routes() must return nil")
}
}
func TestEmptyRawGroupRawFlagOnGroupNotRouteInfo(t *testing.T) {
r := New(nil)
var inner *Group
r.GroupRaw("/oauth", nil, func(g pact.Router) {
inner = g.(*Group)
})
if inner == nil || !inner.raw {
t.Fatal("empty GroupRaw must keep raw=true on Group")
}
if len(r.Routes()) != 0 {
t.Fatalf("empty raw group leaked RouteInfo: %v", r.Routes())
}
// Group.GroupRaw (nested) is a distinct method from Router.GroupRaw.
r.Group("/wrap", nil, func(g pact.Router) {
g.GroupRaw("/inner", nil, func(c pact.Router) {
inner = c.(*Group)
})
})
if inner == nil || !inner.raw {
t.Fatal("nested GroupRaw must be raw")
}
if len(r.Routes()) != 0 {
t.Fatalf("nested empty raw group leaked RouteInfo: %v", r.Routes())
}
}

View File

@@ -0,0 +1,126 @@
package surf
import (
"net/http"
"strings"
"testing"
"git.golem15.com/golem15/summercms/modules/backpack"
"git.golem15.com/golem15/summercms/modules/pact"
"git.golem15.com/golem15/summercms/modules/party"
)
func TestRouteTableRawFlagAndStickyInheritance(t *testing.T) {
r := New(nil)
r.BindPlugin("golem15.demo")
r.Get("/plain", func(http.ResponseWriter, *http.Request) {})
r.Group("/g", nil, func(g pact.Router) {
g.Get("/nested", func(http.ResponseWriter, *http.Request) {})
})
r.GroupRaw("/raw", nil, func(g pact.Router) {
g.Get("/a", func(http.ResponseWriter, *http.Request) {})
g.Group("/child", nil, func(c pact.Router) {
c.Get("/b", func(http.ResponseWriter, *http.Request) {})
})
})
byPattern := map[string]RouteInfo{}
for _, rt := range r.Routes() {
byPattern[rt.Pattern] = rt
}
want := map[string]bool{
"/plain": false,
"/g/nested": false,
"/raw/a": true,
"/raw/child/b": true,
}
for pattern, raw := range want {
got, ok := byPattern[pattern]
if !ok {
t.Fatalf("missing route %s in %v", pattern, keys(byPattern))
}
if got.Raw != raw {
t.Fatalf("%s Raw = %v, want %v", pattern, got.Raw, raw)
}
if got.Method != http.MethodGet {
t.Fatalf("%s Method = %s", pattern, got.Method)
}
if got.PluginID != "golem15.demo" {
t.Fatalf("%s PluginID = %s", pattern, got.PluginID)
}
}
}
func keys(m map[string]RouteInfo) []string {
out := make([]string, 0, len(m))
for k := range m {
out = append(out, k)
}
return out
}
func TestRouteTableDoesNotAliasInternalSlice(t *testing.T) {
r := New(nil)
r.BindPlugin("golem15.demo")
r.Group("/api", Use("jwt.auth"), func(g pact.Router) {
g.Get("/items", func(http.ResponseWriter, *http.Request) {})
})
first := r.Routes()
if len(first) != 1 || len(first[0].Middleware) != 1 {
t.Fatalf("got %+v", first)
}
first[0].Middleware[0] = "mutated"
first[0].Pattern = "/changed"
second := r.Routes()
if second[0].Middleware[0] != "jwt.auth" {
t.Fatalf("internal middleware aliased: %v", second[0].Middleware)
}
if second[0].Pattern != "/api/items" {
t.Fatalf("internal pattern aliased: %s", second[0].Pattern)
}
}
func TestRawGroupHouseMiddlewareRefusedAtBuild(t *testing.T) {
identity := func(next http.Handler) http.Handler { return next }
p := houseRoutePlugin{
id: "golem15.demo",
house: map[string]pact.Middleware{
"house.err": identity,
},
rawUse: []string{"house.err"},
}
_, err := BuildRouter(backpack.New(nil), []party.Plugin{p})
if err == nil {
t.Fatal("want BuildRouter error")
}
if !strings.Contains(err.Error(), "house.err") {
t.Fatalf("want middleware name in error, got %v", err)
}
if !strings.Contains(err.Error(), "golem15.demo") {
t.Fatalf("want plugin in error, got %v", err)
}
if !strings.Contains(err.Error(), "raw group cannot use house-envelope middleware") {
t.Fatalf("want raw-group refusal, got %v", err)
}
}
type houseRoutePlugin struct {
id string
house map[string]pact.Middleware
mw map[string]pact.Middleware
rawUse []string
}
func (p houseRoutePlugin) ID() string { return p.id }
func (p houseRoutePlugin) Requires() []string { return nil }
func (p houseRoutePlugin) Register(*backpack.App) error { return nil }
func (p houseRoutePlugin) Boot(*backpack.App) error { return nil }
func (p houseRoutePlugin) Middlewares() map[string]pact.Middleware { return p.mw }
func (p houseRoutePlugin) HouseMiddlewares() map[string]pact.Middleware { return p.house }
func (p houseRoutePlugin) Routes(r pact.Router) error {
r.GroupRaw("/oauth", Use(p.rawUse...), func(g pact.Router) {
g.Get("/x", func(http.ResponseWriter, *http.Request) {})
})
return nil
}

99
modules/surf/serve.go Normal file
View File

@@ -0,0 +1,99 @@
package surf
import (
"context"
"fmt"
"net"
"net/http"
"os"
"os/signal"
"strings"
"syscall"
"time"
"git.golem15.com/golem15/summercms/modules/backpack"
"git.golem15.com/golem15/summercms/modules/bonfire"
"git.golem15.com/golem15/summercms/modules/lagoon"
"git.golem15.com/golem15/summercms/modules/lagoon/attach"
"git.golem15.com/golem15/summercms/modules/party"
"gocloud.dev/blob"
)
// ServeCommand starts a signal-aware HTTP server on the assembled router.
func ServeCommand(app *backpack.App, plugins []party.Plugin) bonfire.Command {
return bonfire.Command{
Name: "serve",
Description: "Serve the application HTTP API",
Flags: []bonfire.Flag{{
Name: "addr",
Description: "Listen address",
Default: ":8080",
}},
Run: func(ctx context.Context, in bonfire.Input, out bonfire.Output) error {
addr, _ := in.Flag("addr")
if strings.TrimSpace(addr) == "" {
addr = ":8080"
}
sqlDB, gdb, err := lagoon.OpenFromApp(ctx, app)
if err != nil {
return err
}
defer sqlDB.Close()
if err := lagoon.Publish(app, sqlDB, gdb); err != nil {
return err
}
bucket, err := publishUploads(ctx, app)
if err != nil {
return err
}
defer bucket.Close()
h, err := Assemble(app, plugins)
if err != nil {
return err
}
ctx, stop := signal.NotifyContext(ctx, os.Interrupt, syscall.SIGTERM)
defer stop()
srv := &http.Server{Addr: addr, Handler: h, BaseContext: func(net.Listener) context.Context { return ctx }}
errCh := make(chan error, 1)
go func() {
out.Info(fmt.Sprintf("listening on %s", addr))
errCh <- srv.ListenAndServe()
}()
select {
case <-ctx.Done():
shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
if err := srv.Shutdown(shutdownCtx); err != nil {
return err
}
err := <-errCh
if err == nil || err == http.ErrServerClosed {
return nil
}
return err
case err := <-errCh:
if err == nil || err == http.ErrServerClosed {
return nil
}
return err
}
},
}
}
// publishUploads opens storage.uploads.bucket_url and publishes *blob.Bucket.
// An empty URL fails boot the same way an empty JWT secret does.
func publishUploads(ctx context.Context, app *backpack.App) (*blob.Bucket, error) {
if app == nil {
return nil, fmt.Errorf("surf: app is nil")
}
bucket, err := attach.OpenBucket(ctx, app.Config)
if err != nil {
return nil, err
}
if err := attach.Publish(app, bucket); err != nil {
_ = bucket.Close()
return nil, err
}
return bucket, nil
}

View File

@@ -0,0 +1,53 @@
package surf
import (
"os"
"path/filepath"
"strings"
"testing"
"git.golem15.com/golem15/summercms/modules/backpack"
"git.golem15.com/golem15/summercms/modules/compass"
"gocloud.dev/blob"
)
func TestPublishUploadsRequiresURL(t *testing.T) {
app := serveTestApp(t, "uploads:\n bucket_url: \"\"\n")
_, err := publishUploads(t.Context(), app)
if err == nil || !strings.Contains(err.Error(), "bucket_url") {
t.Fatalf("got %v, want bucket_url error", err)
}
}
func TestPublishUploadsStoresBucket(t *testing.T) {
app := serveTestApp(t, "uploads:\n bucket_url: \"mem://\"\n public_path_prefix: \"/storage/uploads\"\n")
bucket, err := publishUploads(t.Context(), app)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = bucket.Close() })
got, ok := app.Lookup[*blob.Bucket]()
if !ok || got != bucket {
t.Fatal("publishUploads must store the opened *blob.Bucket")
}
}
func serveTestApp(t *testing.T, storageYAML string) *backpack.App {
t.Helper()
dir := t.TempDir()
if err := os.WriteFile(filepath.Join(dir, "app.yaml"), []byte("name: serve-uploads\n"), 0o644); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(dir, "storage.yaml"), []byte(storageYAML), 0o644); err != nil {
t.Fatal(err)
}
cfg, err := compass.Open(compass.Options{
Dir: dir,
Env: "development",
Environ: []string{"SUMMER_ENV=development"},
})
if err != nil {
t.Fatal(err)
}
return backpack.New(cfg)
}