From 4ee4c4a2fc27d644cdfac9131cb369e8617df861 Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Thu, 17 Sep 2026 20:04:13 +0200 Subject: [PATCH] feat(03-01): add ServeMux groups, JWT verifier, and serve command Named middleware resolves at boot, HS256 tokens are pinned with required exp/sub, and both binaries expose a signal-aware serve command. Co-authored-by: Cursor --- bouncer/context.go | 28 ++ bouncer/jwt.go | 139 ++++++++++ bouncer/jwt_test.go | 134 ++++++++++ cmd/summer/main.go | 1 + cmd/summer/main_test.go | 2 +- examples/hello/hello_test.go | 39 +++ examples/hello/main.go | 2 + examples/hello/plugins/greeter/plugin.go | 26 +- go.mod | 1 + go.sum | 2 + internal/build/build.go | 2 + internal/build/build_test.go | 3 + lagoon/commands.go | 2 +- pact/capabilities.go | 34 ++- pact/capabilities_test.go | 29 ++ surf/router.go | 326 +++++++++++++++++++++++ surf/router_test.go | 143 ++++++++++ surf/serve.go | 75 ++++++ 18 files changed, 976 insertions(+), 12 deletions(-) create mode 100644 bouncer/context.go create mode 100644 bouncer/jwt.go create mode 100644 bouncer/jwt_test.go create mode 100644 surf/router.go create mode 100644 surf/router_test.go create mode 100644 surf/serve.go diff --git a/bouncer/context.go b/bouncer/context.go new file mode 100644 index 0000000..2f26a88 --- /dev/null +++ b/bouncer/context.go @@ -0,0 +1,28 @@ +package bouncer + +import "context" + +type userKey struct{} + +// Principal is the authenticated identity stored on the request context. +type Principal struct { + ID uint + MustChangePassword bool +} + +// WithUser stores the verified principal on ctx. +func WithUser(ctx context.Context, user *Principal) context.Context { + if ctx == nil { + ctx = context.Background() + } + return context.WithValue(ctx, userKey{}, user) +} + +// User returns the verified principal from ctx. +func User(ctx context.Context) (*Principal, bool) { + if ctx == nil { + return nil, false + } + u, ok := ctx.Value(userKey{}).(*Principal) + return u, ok && u != nil +} diff --git a/bouncer/jwt.go b/bouncer/jwt.go new file mode 100644 index 0000000..49f46a5 --- /dev/null +++ b/bouncer/jwt.go @@ -0,0 +1,139 @@ +package bouncer + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "strconv" + "strings" + + "github.com/golang-jwt/jwt/v5" +) + +const ( + msgTokenNotProvided = "Token not provided" + msgTokenExpired = "Token has expired" + msgUserNotFound = "User not found" + msgBadSignature = "Token Signature could not be verified." + msgMalformed = "Wrong number of segments" + msgRequiredClaims = "JWT payload does not contain the required claims" +) + +// UserProvider loads a persisted user by JWT subject. +type UserProvider interface { + FindByID(ctx context.Context, id uint) (*Principal, error) +} + +// Middleware validates a pinned HS256 bearer token and loads the user. +func Middleware(secret string, users UserProvider) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + raw, err := bearerToken(r) + if err != nil { + write401(w, err.Error()) + return + } + sub, err := Verify(raw, secret) + if err != nil { + write401(w, err.Error()) + return + } + id, err := strconv.ParseUint(sub, 10, 64) + if err != nil || id == 0 { + write401(w, msgUserNotFound) + return + } + if users == nil { + write401(w, msgUserNotFound) + return + } + user, err := users.FindByID(r.Context(), uint(id)) + if err != nil { + write401(w, "Authentication error") + return + } + if user == nil { + write401(w, msgUserNotFound) + return + } + next.ServeHTTP(w, r.WithContext(WithUser(r.Context(), user))) + }) + } +} + +// Verify parses a token with HS256 pinned and a required exp and sub. +func Verify(tokenString, secret string) (string, error) { + if strings.TrimSpace(secret) == "" { + return "", fmt.Errorf("bouncer: jwt secret is empty") + } + parser := jwt.NewParser(jwt.WithValidMethods([]string{"HS256"}), jwt.WithExpirationRequired()) + claims := jwt.MapClaims{} + _, err := parser.ParseWithClaims(tokenString, claims, func(t *jwt.Token) (any, error) { + return []byte(secret), nil + }) + if err != nil { + return "", mapJWTError(err) + } + sub := subject(claims) + if sub == "" { + return "", errors.New(msgRequiredClaims) + } + return sub, nil +} + +func bearerToken(r *http.Request) (string, error) { + h := strings.TrimSpace(r.Header.Get("Authorization")) + if h == "" { + return "", errors.New(msgTokenNotProvided) + } + token, ok := strings.CutPrefix(h, "Bearer ") + token = strings.TrimSpace(token) + if !ok || token == "" { + return "", errors.New(msgTokenNotProvided) + } + return token, nil +} + +func subject(claims jwt.MapClaims) string { + switch v := claims["sub"].(type) { + case string: + return strings.TrimSpace(v) + case float64: + if v <= 0 { + return "" + } + return strconv.FormatInt(int64(v), 10) + case json.Number: + return strings.TrimSpace(v.String()) + default: + return "" + } +} + +func mapJWTError(err error) error { + switch { + case errors.Is(err, jwt.ErrTokenExpired): + return errors.New(msgTokenExpired) + case errors.Is(err, jwt.ErrTokenSignatureInvalid), errors.Is(err, jwt.ErrTokenUnverifiable): + return errors.New(msgBadSignature) + case errors.Is(err, jwt.ErrTokenMalformed): + return errors.New(msgMalformed) + default: + msg := err.Error() + if strings.Contains(strings.ToLower(msg), "expired") { + return errors.New(msgTokenExpired) + } + if strings.Contains(strings.ToLower(msg), "malformed") || strings.Contains(msg, "segment") { + return errors.New(msgMalformed) + } + return errors.New(msgBadSignature) + } +} + +func write401(w http.ResponseWriter, message string) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusUnauthorized) + _ = json.NewEncoder(w).Encode(map[string]any{"error": true, "message": message}) +} diff --git a/bouncer/jwt_test.go b/bouncer/jwt_test.go new file mode 100644 index 0000000..39e5476 --- /dev/null +++ b/bouncer/jwt_test.go @@ -0,0 +1,134 @@ +package bouncer + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +const secret = "test-secret" + +type memUsers struct { + byID map[uint]*Principal + err error +} + +func (m memUsers) FindByID(ctx context.Context, id uint) (*Principal, error) { + if m.err != nil { + return nil, m.err + } + return m.byID[id], nil +} + +func sign(t *testing.T, method jwt.SigningMethod, claims jwt.MapClaims, key []byte) string { + t.Helper() + tok := jwt.NewWithClaims(method, claims) + s, err := tok.SignedString(key) + if err != nil { + t.Fatal(err) + } + return s +} + +func TestVerifyRejectsBadTokens(t *testing.T) { + valid := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ + "sub": "1", + "exp": time.Now().Add(time.Hour).Unix(), + }, []byte(secret)) + if sub, err := Verify(valid, secret); err != nil || sub != "1" { + t.Fatalf("valid token: %s %v", sub, err) + } + + expired := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ + "sub": "1", + "exp": time.Now().Add(-time.Hour).Unix(), + }, []byte(secret)) + if _, err := Verify(expired, secret); err == nil || err.Error() != msgTokenExpired { + t.Fatalf("expired: %v", err) + } + + none := sign(t, jwt.SigningMethodHS384, jwt.MapClaims{ + "sub": "1", + "exp": time.Now().Add(time.Hour).Unix(), + }, []byte(secret)) + if _, err := Verify(none, secret); err == nil || err.Error() != msgBadSignature { + t.Fatalf("wrong alg: %v", err) + } + + badSig := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ + "sub": "1", + "exp": time.Now().Add(time.Hour).Unix(), + }, []byte("other-secret")) + if _, err := Verify(badSig, secret); err == nil || err.Error() != msgBadSignature { + t.Fatalf("bad sig: %v", err) + } + + if _, err := Verify("not-a-jwt", secret); err == nil || err.Error() != msgMalformed { + t.Fatalf("malformed: %v", err) + } + + missingSub := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ + "exp": time.Now().Add(time.Hour).Unix(), + }, []byte(secret)) + if _, err := Verify(missingSub, secret); err == nil || err.Error() != msgRequiredClaims { + t.Fatalf("missing sub: %v", err) + } +} + +func TestMiddlewareStatusBodies(t *testing.T) { + users := memUsers{byID: map[uint]*Principal{1: {ID: 1}}} + h := Middleware(secret, users)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("X-Hit", "1") + w.WriteHeader(http.StatusNoContent) + })) + + assert401 := func(t *testing.T, req *http.Request, msg string) { + t.Helper() + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + if rec.Code != http.StatusUnauthorized { + t.Fatalf("status = %d body %s", rec.Code, rec.Body.String()) + } + if rec.Header().Get("X-Hit") != "" { + t.Fatal("handler ran") + } + var body map[string]any + if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { + t.Fatal(err) + } + if body["error"] != true || body["message"] != msg { + t.Fatalf("body = %v want %s", body, msg) + } + } + + t.Run("missing", func(t *testing.T) { + assert401(t, httptest.NewRequest(http.MethodGet, "/", nil), msgTokenNotProvided) + }) + t.Run("unknown-user", func(t *testing.T) { + tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ + "sub": "99", + "exp": time.Now().Add(time.Hour).Unix(), + }, []byte(secret)) + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.Header.Set("Authorization", "Bearer "+tok) + assert401(t, req, msgUserNotFound) + }) + t.Run("valid", func(t *testing.T) { + tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ + "sub": "1", + "exp": time.Now().Add(time.Hour).Unix(), + }, []byte(secret)) + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.Header.Set("Authorization", "Bearer "+tok) + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + if rec.Code != http.StatusNoContent || rec.Header().Get("X-Hit") != "1" { + t.Fatalf("status=%d hit=%s", rec.Code, rec.Header().Get("X-Hit")) + } + }) +} diff --git a/cmd/summer/main.go b/cmd/summer/main.go index db44eca..514157c 100644 --- a/cmd/summer/main.go +++ b/cmd/summer/main.go @@ -35,6 +35,7 @@ func toolCommands() []bonfire.Command { delegateCommand("migrate", "Run plugin migrations in the app binary"), delegateRollbackCommand(), delegateCommand("migrate:status", "Show per-plugin migration history in the app binary"), + delegateCommand("serve", "Run the app HTTP server"), } } diff --git a/cmd/summer/main_test.go b/cmd/summer/main_test.go index 4452e9b..03e5ec4 100644 --- a/cmd/summer/main_test.go +++ b/cmd/summer/main_test.go @@ -19,7 +19,7 @@ func TestToolCommandNames(t *testing.T) { for _, c := range toolCommands() { names = append(names, c.Name) } - for _, want := range []string{"build", "make:plugin", "plugin:add", "dev", "migrate", "migrate:rollback", "migrate:status"} { + for _, want := range []string{"build", "make:plugin", "plugin:add", "dev", "migrate", "migrate:rollback", "migrate:status", "serve"} { if !slices.Contains(names, want) { t.Fatalf("missing %s in %v", want, names) } diff --git a/examples/hello/hello_test.go b/examples/hello/hello_test.go index 77b2664..1c7b899 100644 --- a/examples/hello/hello_test.go +++ b/examples/hello/hello_test.go @@ -3,6 +3,8 @@ package main import ( "bytes" "crypto/sha256" + "net/http" + "net/http/httptest" "os" "os/exec" "strings" @@ -13,6 +15,7 @@ import ( "git.golem15.com/golem15/summercms/compass" "git.golem15.com/golem15/summercms/pact" "git.golem15.com/golem15/summercms/party" + "git.golem15.com/golem15/summercms/surf" ) func TestGreeterHelloPrintsLayeredConfig(t *testing.T) { @@ -104,6 +107,42 @@ func TestBuiltBinaryGreeterHello(t *testing.T) { assertGreeting(t, got, want) } +func TestTypedItemRoute(t *testing.T) { + cfg, err := compass.Load("config") + if err != nil { + t.Fatal(err) + } + application := backpack.New(cfg) + plugins, err := party.Activate(application, PluginIDs) + if err != nil { + t.Fatal(err) + } + h, err := surf.Assemble(application, plugins) + 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 id status = %d", rec.Code) + } + rec = httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/items/12", nil)) + if rec.Code != http.StatusOK || rec.Body.String() != `{"id":12}` { + t.Fatalf("got %d %s", rec.Code, rec.Body.String()) + } +} + +func TestHelpListsServe(t *testing.T) { + var buf bytes.Buffer + if err := run([]string{"--help"}, &buf); err != nil { + t.Fatal(err) + } + if !strings.Contains(buf.String(), "serve") { + t.Fatalf("help missing serve:\n%s", buf.String()) + } +} + func TestMigrateWithoutDSNFailsLoudly(t *testing.T) { var buf bytes.Buffer err := run([]string{"migrate"}, &buf) diff --git a/examples/hello/main.go b/examples/hello/main.go index 9e00808..96b7d95 100644 --- a/examples/hello/main.go +++ b/examples/hello/main.go @@ -13,6 +13,7 @@ import ( "git.golem15.com/golem15/summercms/lagoon" "git.golem15.com/golem15/summercms/pact" "git.golem15.com/golem15/summercms/party" + "git.golem15.com/golem15/summercms/surf" ) func main() { @@ -33,6 +34,7 @@ func run(args []string, out io.Writer) error { return err } commands := lagoon.RuntimeCommands(app, plugins) + commands = append(commands, surf.ServeCommand(app, plugins)) for _, plugin := range plugins { if hasCommands, ok := plugin.(pact.HasCommands); ok { commands = append(commands, hasCommands.Commands()...) diff --git a/examples/hello/plugins/greeter/plugin.go b/examples/hello/plugins/greeter/plugin.go index a5a94e9..d964981 100644 --- a/examples/hello/plugins/greeter/plugin.go +++ b/examples/hello/plugins/greeter/plugin.go @@ -2,17 +2,24 @@ package greeter import ( "context" + "fmt" + "net/http" "git.golem15.com/golem15/summercms/backpack" "git.golem15.com/golem15/summercms/bonfire" "git.golem15.com/golem15/summercms/festival" "git.golem15.com/golem15/summercms/pact" "git.golem15.com/golem15/summercms/party" + "git.golem15.com/golem15/summercms/surf" "git.golem15.com/golem15/summercms/towel" ) -var _ festival.Collectable = (*HelloEvent)(nil) -var _ festival.Handleable = (*HelloEvent)(nil) +var ( + _ festival.Collectable = (*HelloEvent)(nil) + _ festival.Handleable = (*HelloEvent)(nil) + _ pact.HasCommands = (*Plugin)(nil) + _ pact.HasRoutes = (*Plugin)(nil) +) // Plugin is the golem15.greeter plugin. It requires golem15.hello. type Plugin struct { @@ -49,6 +56,19 @@ func (p *Plugin) Boot(app *backpack.App) error { return nil } +func (p *Plugin) Routes(r pact.Router) error { + r.Get("/items/{id}", func(w http.ResponseWriter, req *http.Request) { + id, ok := surf.IntParam(req, "id") + if !ok { + http.NotFound(w, req) + return + } + w.Header().Set("Content-Type", "application/json") + fmt.Fprintf(w, `{"id":%d}`, id) + }) + return nil +} + func (p *Plugin) Commands() []bonfire.Command { return []bonfire.Command{{ Name: "greeter:hello", @@ -110,8 +130,6 @@ func (e *HelloEvent) IsHandled() bool { return e != nil && e.handled } -var _ pact.HasCommands = (*Plugin)(nil) - func init() { party.Register(&Plugin{}) } diff --git a/go.mod b/go.mod index 5d04d7d..d2fc88e 100644 --- a/go.mod +++ b/go.mod @@ -8,6 +8,7 @@ require ( github.com/fsnotify/fsnotify v1.10.1 github.com/go-gormigrate/gormigrate/v2 v2.1.7 github.com/goccy/go-yaml v1.19.2 + github.com/golang-jwt/jwt/v5 v5.3.1 github.com/jackc/pgx/v5 v5.10.0 github.com/knadh/koanf/parsers/yaml v1.1.1 github.com/knadh/koanf/providers/confmap v1.0.1 diff --git a/go.sum b/go.sum index 48b85a0..8ba3db4 100644 --- a/go.sum +++ b/go.sum @@ -10,6 +10,8 @@ github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9L github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= github.com/goccy/go-yaml v1.19.2 h1:PmFC1S6h8ljIz6gMRBopkjP1TVT7xuwrButHID66PoM= github.com/goccy/go-yaml v1.19.2/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA= +github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY= +github.com/golang-jwt/jwt/v5 v5.3.1/go.mod h1:fxCRLWMO43lRc8nhHWY6LGqRcf+1gQWArsqaEUEa5bE= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= diff --git a/internal/build/build.go b/internal/build/build.go index 5b1abf6..929d3b1 100644 --- a/internal/build/build.go +++ b/internal/build/build.go @@ -88,6 +88,7 @@ func generateMain(m Manifest) ([]byte, error) { b.WriteString("\t\"git.golem15.com/golem15/summercms/lagoon\"\n") b.WriteString("\t\"git.golem15.com/golem15/summercms/pact\"\n") b.WriteString("\t\"git.golem15.com/golem15/summercms/party\"\n") + b.WriteString("\t\"git.golem15.com/golem15/summercms/surf\"\n") b.WriteString(")\n\n") b.WriteString("func main() {\n") b.WriteString("\tif err := run(os.Args[1:], os.Stdout); err != nil {\n") @@ -106,6 +107,7 @@ func generateMain(m Manifest) ([]byte, error) { b.WriteString("\t\treturn err\n") b.WriteString("\t}\n") b.WriteString("\tcommands := lagoon.RuntimeCommands(app, plugins)\n") + b.WriteString("\tcommands = append(commands, surf.ServeCommand(app, plugins))\n") b.WriteString("\tfor _, plugin := range plugins {\n") b.WriteString("\t\tif hasCommands, ok := plugin.(pact.HasCommands); ok {\n") b.WriteString("\t\t\tcommands = append(commands, hasCommands.Commands()...)\n") diff --git a/internal/build/build_test.go b/internal/build/build_test.go index c4c1ff9..5f95391 100644 --- a/internal/build/build_test.go +++ b/internal/build/build_test.go @@ -97,6 +97,9 @@ func TestGenerateStableQuotedImportsInManifestOrder(t *testing.T) { if !bytes.Contains(mainSrc, []byte("lagoon.RuntimeCommands")) { t.Fatalf("main does not register lagoon runtime commands:\n%s", mainSrc) } + if !bytes.Contains(mainSrc, []byte("surf.ServeCommand")) { + t.Fatalf("main does not register serve:\n%s", mainSrc) + } if bytes.Contains(mainSrc, []byte("examples/hello")) { t.Fatal("generated main hard-codes examples/hello") } diff --git a/lagoon/commands.go b/lagoon/commands.go index 016e7bd..5ddb1fc 100644 --- a/lagoon/commands.go +++ b/lagoon/commands.go @@ -11,7 +11,7 @@ import ( ) // RuntimeCommands returns migrate, migrate:rollback and migrate:status. -// Serve is registered once the HTTP layer exists. +// Serve is registered separately via surf.ServeCommand. func RuntimeCommands(app *backpack.App, plugins []party.Plugin) []bonfire.Command { return []bonfire.Command{ { diff --git a/pact/capabilities.go b/pact/capabilities.go index e06ff36..1567b3b 100644 --- a/pact/capabilities.go +++ b/pact/capabilities.go @@ -2,6 +2,7 @@ package pact import ( "io/fs" + "net/http" "git.golem15.com/golem15/summercms/bonfire" "github.com/go-gormigrate/gormigrate/v2" @@ -26,6 +27,30 @@ type HasMigrations interface { Migrations() []*gormigrate.Migration } +// Middleware is a named HTTP wrapper registered by a plugin. +type Middleware func(http.Handler) http.Handler + +// HasMiddleware is implemented by plugins that register named middleware. +type HasMiddleware interface { + Middlewares() map[string]Middleware +} + +// Router is the Laravel-like group builder implemented by surf. +type Router interface { + Group(prefix string, middleware []string, fn func(Router)) + Get(path string, handler http.HandlerFunc, middleware ...string) +} + +// HasRoutes is implemented by plugins that declare HTTP routes. +type HasRoutes interface { + Routes(Router) error +} + +// HasModels is implemented by plugins that expose GORM models. +type HasModels interface { + Models() []any +} + // OptionalMessage is a service an optional plugin may publish so other // plugins can integrate without importing that plugin's package. type OptionalMessage interface { @@ -33,12 +58,8 @@ type OptionalMessage interface { } // Future capability families are type-asserted when their first consumer -// packages exist. Method payloads for HTTP/model surfaces are declared in -// the phase that ships surf: +// packages exist: // -// HasModels -// HasRoutes -// HasMiddleware // HasJobs // HasListeners // HasAdminControllers @@ -49,4 +70,5 @@ type OptionalMessage interface { // HasLang // // The kernel type-asserts HasConfig (party, before Register), HasCommands -// (generated app main, after Boot), and HasMigrations (lagoon migrate). +// (generated app main, after Boot), HasMigrations (lagoon migrate), and +// HasMiddleware/HasRoutes (surf assemble). diff --git a/pact/capabilities_test.go b/pact/capabilities_test.go index 42c9dcc..ac94465 100644 --- a/pact/capabilities_test.go +++ b/pact/capabilities_test.go @@ -117,3 +117,32 @@ func TestOptionalMessageContract(t *testing.T) { type extraMessage struct{ s string } func (e extraMessage) Message() string { return e.s } + +type routesOnly struct{} + +func (routesOnly) Routes(Router) error { return nil } + +type middlewareOnly struct{} + +func (middlewareOnly) Middlewares() map[string]Middleware { + return map[string]Middleware{"jwt.auth": nil} +} + +type modelsOnly struct{} + +func (modelsOnly) Models() []any { return []any{struct{}{}} } + +func TestHTTPCapabilitiesDiscoveredByTypeAssertion(t *testing.T) { + var r HasRoutes = routesOnly{} + if err := r.Routes(nil); err != nil { + t.Fatal(err) + } + var m HasMiddleware = middlewareOnly{} + if _, ok := m.Middlewares()["jwt.auth"]; !ok { + t.Fatal("missing jwt.auth") + } + var models HasModels = modelsOnly{} + if len(models.Models()) != 1 { + t.Fatalf("Models = %+v", models.Models()) + } +} diff --git a/surf/router.go b/surf/router.go new file mode 100644 index 0000000..edb0b0a --- /dev/null +++ b/surf/router.go @@ -0,0 +1,326 @@ +package surf + +import ( + "fmt" + "net/http" + "strconv" + "strings" + + "git.golem15.com/golem15/summercms/backpack" + "git.golem15.com/golem15/summercms/pact" + "git.golem15.com/golem15/summercms/party" + "git.golem15.com/golem15/summercms/towel" +) + +// Use names a middleware list for Group, matching PHP ->middleware(). +func Use(names ...string) []string { + return names +} + +type namedMiddleware struct { + pluginID string + fn pact.Middleware +} + +type route struct { + pluginID string + method string + path string + handler http.Handler + middleware []string +} + +// Router compiles group declarations onto net/http ServeMux. +type Router struct { + pluginID string + prefix string + middleware []string + named map[string]namedMiddleware + routes []route + seen map[string]string + origins []string + compileErr error +} + +var ( + _ pact.Router = (*Router)(nil) + _ pact.Router = (*Group)(nil) +) + +// Group is a prefixed route collection. +type Group struct { + router *Router + pluginID string + prefix string + middleware []string +} + +// New returns an empty router. +func New(origins []string) *Router { + return &Router{ + named: make(map[string]namedMiddleware), + seen: make(map[string]string), + origins: origins, + } +} + +// RegisterMiddleware stores a named wrapper. Duplicate names fail. +func (r *Router) RegisterMiddleware(pluginID, name string, fn pact.Middleware) error { + if r == nil { + return fmt.Errorf("surf: router is nil") + } + if name == "" || fn == nil { + return fmt.Errorf("surf: plugin %q registered empty middleware", pluginID) + } + if existing, ok := r.named[name]; ok { + return fmt.Errorf("surf: middleware %q already registered by %s", name, existing.pluginID) + } + r.named[name] = namedMiddleware{pluginID: pluginID, fn: fn} + return nil +} + +// BindPlugin records the plugin declaring subsequent routes. +func (r *Router) BindPlugin(id string) { + if r != nil { + r.pluginID = id + } +} + +func (r *Router) Group(prefix string, middleware []string, fn func(pact.Router)) { + if r == nil || fn == nil { + return + } + g := &Group{ + router: r, + pluginID: r.pluginID, + prefix: joinPath(r.prefix, prefix), + middleware: append([]string{}, r.middleware...), + } + g.middleware = append(g.middleware, middleware...) + fn(g) +} + +func (r *Router) Get(path string, handler http.HandlerFunc, middleware ...string) { + if r == nil { + return + } + r.add(r.pluginID, r.prefix, r.middleware, path, handler, middleware) +} + +func (g *Group) Group(prefix string, middleware []string, fn func(pact.Router)) { + if g == nil || g.router == nil || fn == nil { + return + } + next := &Group{ + router: g.router, + pluginID: g.pluginID, + prefix: joinPath(g.prefix, prefix), + middleware: append([]string{}, g.middleware...), + } + next.middleware = append(next.middleware, middleware...) + fn(next) +} + +func (g *Group) Get(path string, handler http.HandlerFunc, middleware ...string) { + if g == nil || g.router == nil { + return + } + g.router.add(g.pluginID, g.prefix, g.middleware, path, handler, middleware) +} + +func (r *Router) add(pluginID, prefix string, groupMW []string, path string, handler http.HandlerFunc, extra []string) { + full := joinPath(prefix, path) + key := "GET " + full + if prev, ok := r.seen[key]; ok { + r.compileErr = fmt.Errorf("surf: duplicate route %s registered by %s and %s", key, prev, pluginID) + return + } + r.seen[key] = pluginID + mw := append([]string{}, groupMW...) + mw = append(mw, extra...) + r.routes = append(r.routes, route{ + pluginID: pluginID, + method: "GET", + path: full, + handler: handler, + middleware: mw, + }) +} + +func (r *Router) compile() (http.Handler, error) { + if r.compileErr != nil { + return nil, r.compileErr + } + mux := http.NewServeMux() + for _, rt := range r.routes { + h, err := r.wrap(rt) + if err != nil { + return nil, err + } + mux.Handle("GET "+rt.path, h) + } + return recoverJSON(cors(r.origins, mux)), nil +} + +func (r *Router) wrap(rt route) (http.Handler, error) { + h := rt.handler + h = noOpLimit(h) + h = orgSlot(h) + for i := len(rt.middleware) - 1; i >= 0; i-- { + name := rt.middleware[i] + named, ok := r.named[name] + if !ok { + return nil, fmt.Errorf("surf: plugin %q references unknown middleware %q", rt.pluginID, name) + } + h = named.fn(h) + } + h = locale(h) + return h, nil +} + +// Assemble registers plugin middleware and routes, then compiles ServeMux. +func Assemble(app *backpack.App, plugins []party.Plugin) (http.Handler, error) { + r := New(corsOrigins(app)) + for _, p := range plugins { + if hm, ok := p.(pact.HasMiddleware); ok { + for name, fn := range hm.Middlewares() { + if err := r.RegisterMiddleware(p.ID(), name, fn); err != nil { + return nil, err + } + } + } + } + for _, p := range plugins { + r.BindPlugin(p.ID()) + if hr, ok := p.(pact.HasRoutes); ok { + if err := hr.Routes(r); err != nil { + return nil, fmt.Errorf("surf: routes %s: %w", p.ID(), err) + } + } + } + for _, rt := range r.routes { + if _, err := r.wrap(rt); err != nil { + return nil, err + } + } + return r.compile() +} + +func corsOrigins(app *backpack.App) []string { + if app == nil || app.Config == nil { + return nil + } + raw, ok := app.Config.Lookup("http.cors.allowed_origins") + if !ok { + return nil + } + switch v := raw.(type) { + case []string: + return v + case []any: + out := make([]string, 0, len(v)) + for _, item := range v { + s, _ := item.(string) + if s != "" { + out = append(out, s) + } + } + return out + default: + return nil + } +} + +func recoverJSON(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + defer func() { + if rec := recover(); rec != nil { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte(`{"error":true,"message":"Internal server error"}`)) + } + }() + next.ServeHTTP(w, r) + }) +} + +func cors(origins []string, next http.Handler) http.Handler { + allowed := make(map[string]struct{}, len(origins)) + for _, o := range origins { + allowed[o] = struct{}{} + } + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + origin := r.Header.Get("Origin") + if _, ok := allowed[origin]; ok && origin != "" { + w.Header().Set("Access-Control-Allow-Origin", origin) + w.Header().Set("Vary", "Origin") + w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type, Accept") + w.Header().Set("Access-Control-Allow-Methods", "GET, POST, PUT, PATCH, DELETE, OPTIONS") + } + if r.Method == http.MethodOptions { + w.WriteHeader(http.StatusNoContent) + return + } + next.ServeHTTP(w, r) + }) +} + +func locale(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + next.ServeHTTP(w, r.WithContext(towel.WithLocale(r.Context(), r.Header.Get("Accept-Language")))) + }) +} + +func orgSlot(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + next.ServeHTTP(w, r.WithContext(towel.WithOrganization(r.Context(), ""))) + }) +} + +// Limiter wraps handlers. Phase 6 replaces the no-op with named buckets. +type Limiter interface { + Wrap(http.Handler) http.Handler +} + +type noopLimiter struct{} + +func (noopLimiter) Wrap(next http.Handler) http.Handler { return next } + +func noOpLimit(next http.Handler) http.Handler { + return noopLimiter{}.Wrap(next) +} + +// IntParam returns a positive integer path value. Missing or malformed ids +// are false so callers can 404 both unknown and non-integer values. +func IntParam(r *http.Request, name string) (int64, bool) { + if r == nil { + return 0, false + } + raw := r.PathValue(name) + if raw == "" { + return 0, false + } + n, err := strconv.ParseInt(raw, 10, 64) + if err != nil || n < 1 { + return 0, false + } + return n, true +} + +func joinPath(prefix, path string) string { + prefix = strings.TrimSuffix(prefix, "/") + path = strings.TrimSpace(path) + if path == "" || path == "/" { + if prefix == "" { + return "/" + } + return prefix + } + if !strings.HasPrefix(path, "/") { + path = "/" + path + } + if prefix == "" { + return path + } + return prefix + path +} diff --git a/surf/router_test.go b/surf/router_test.go new file mode 100644 index 0000000..6a72adf --- /dev/null +++ b/surf/router_test.go @@ -0,0 +1,143 @@ +package surf + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "git.golem15.com/golem15/summercms/pact" +) + +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 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([]string{"http://localhost:3000"}) + 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 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) + } +} + +func TestTypedIDRouteReturns404(t *testing.T) { + r := New(nil) + r.Get("/items/{id}", func(w http.ResponseWriter, req *http.Request) { + if _, ok := IntParam(req, "id"); !ok { + http.NotFound(w, req) + return + } + w.WriteHeader(http.StatusOK) + }) + 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("status = %d", rec.Code) + } +} diff --git a/surf/serve.go b/surf/serve.go new file mode 100644 index 0000000..ddcf36e --- /dev/null +++ b/surf/serve.go @@ -0,0 +1,75 @@ +package surf + +import ( + "context" + "fmt" + "net" + "net/http" + "os" + "os/signal" + "strings" + "syscall" + "time" + + "git.golem15.com/golem15/summercms/backpack" + "git.golem15.com/golem15/summercms/bonfire" + "git.golem15.com/golem15/summercms/lagoon" + "git.golem15.com/golem15/summercms/party" +) + +// 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 + } + 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 + } + }, + } +}