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:
179
modules/surf/admin_prefix_test.go
Normal file
179
modules/surf/admin_prefix_test.go
Normal 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
46
modules/surf/bodylimit.go
Normal 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
|
||||
}
|
||||
283
modules/surf/bodylimit_test.go
Normal file
283
modules/surf/bodylimit_test.go
Normal 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
105
modules/surf/clientip.go
Normal 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
|
||||
}
|
||||
59
modules/surf/clientip_test.go
Normal file
59
modules/surf/clientip_test.go
Normal 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
161
modules/surf/cors.go
Normal 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
|
||||
}
|
||||
92
modules/surf/cors_coverage_test.go
Normal file
92
modules/surf/cors_coverage_test.go
Normal 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
122
modules/surf/cors_test.go
Normal 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
182
modules/surf/limiter.go
Normal 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
|
||||
}
|
||||
92
modules/surf/limiter_coverage_test.go
Normal file
92
modules/surf/limiter_coverage_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
93
modules/surf/limiter_store.go
Normal file
93
modules/surf/limiter_store.go
Normal 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
|
||||
}
|
||||
492
modules/surf/limiter_test.go
Normal file
492
modules/surf/limiter_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
19
modules/surf/locale_from_principal.go
Normal file
19
modules/surf/locale_from_principal.go
Normal 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)
|
||||
})
|
||||
}
|
||||
53
modules/surf/locale_from_principal_test.go
Normal file
53
modules/surf/locale_from_principal_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
254
modules/surf/middleware_test.go
Normal file
254
modules/surf/middleware_test.go
Normal 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
98
modules/surf/params.go
Normal 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+":")
|
||||
}
|
||||
63
modules/surf/params_test.go
Normal file
63
modules/surf/params_test.go
Normal 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")
|
||||
}
|
||||
}
|
||||
39
modules/surf/routelist_command.go
Normal file
39
modules/surf/routelist_command.go
Normal 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
720
modules/surf/router.go
Normal 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
534
modules/surf/router_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
30
modules/surf/routetable.go
Normal file
30
modules/surf/routetable.go
Normal 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
|
||||
}
|
||||
83
modules/surf/routetable_coverage_test.go
Normal file
83
modules/surf/routetable_coverage_test.go
Normal 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())
|
||||
}
|
||||
}
|
||||
126
modules/surf/routetable_test.go
Normal file
126
modules/surf/routetable_test.go
Normal 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
99
modules/surf/serve.go
Normal 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
|
||||
}
|
||||
53
modules/surf/serve_test.go
Normal file
53
modules/surf/serve_test.go
Normal 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)
|
||||
}
|
||||
Reference in New Issue
Block a user