diff --git a/.env.example b/.env.example index 4f6f661..63eaa66 100755 --- a/.env.example +++ b/.env.example @@ -1,24 +1,72 @@ +# ============================================================================== +# Tusk — example environment +# Copy to .env and fill in. The server fails fast on startup if anything +# required here is missing or invalid, rather than misbehaving later. +# ============================================================================== + APP_NAME=Tusk +APP_VERSION=1.0.0 +# dev | staging | production. Production changes several defaults — see below. APP_MODE=dev -SERVER_PORT=8081 -ALLOWED_ORIGINS=http://localhost:3000,http://example.com +SERVER_PORT=8080 + +# Comma-separated. Leave empty to allow no cross-origin browser requests. +ALLOWED_ORIGINS=http://localhost:3000 +# --- Security ----------------------------------------------------------------- +# REQUIRED. Minimum 32 characters — HS256 signatures are only as strong as this +# key, so a short one lets an attacker mint valid tokens for any user. +# Generate with: openssl rand -base64 48 +JWT_SECRET= +# How long an access token stays valid. Refresh tokens last 7 days. +ACCESS_TOKEN_TTL=1h -DB_DRIVER=mysql +# --- Database ----------------------------------------------------------------- +# PostgreSQL is the supported database. mysql and sqlite drivers still load but +# the shipped migrations target Postgres. +DB_DRIVER=postgres DB_HOST=localhost -DB_PORT=3306 +DB_PORT=5432 DB_NAME=tusk -DB_USER=root -DB_PASS=toor +DB_USER=postgres +DB_PASS= + +# disable | require | verify-ca | verify-full +# Defaults to "require" when APP_MODE=production, "disable" otherwise. +# DB_SSLMODE=require + +# Connection pool +DB_MAX_IDLE_CONNS=10 +DB_MAX_OPEN_CONNS=100 +DB_CONN_MAX_LIFETIME=60 + +# --- API documentation -------------------------------------------------------- +# Serves /docs and /openapi.json. Defaults to false when APP_MODE=production, +# true otherwise. API endpoints are unaffected either way — this only controls +# whether the route map and schemas are published. +# DOCS_ENABLED=false +# --- HTTP server limits ------------------------------------------------------- +# Raise READ_TIMEOUT if clients legitimately send large bodies over slow links. +READ_TIMEOUT=30s +WRITE_TIMEOUT=60s +IDLE_TIMEOUT=120s +SHUTDOWN_TIMEOUT=30s +# Maximum request body, in bytes. Default 10 MiB. +MAX_REQUEST_BODY_BYTES=10485760 -ACCESS_TOKEN_TTL= "3600s" +# --- Rate limiting (per client IP) -------------------------------------------- +# Token bucket: BURST is the capacity, RPS the sustained refill rate. +RATE_LIMIT_BURST=30 +RATE_LIMIT_RPS=10 +# --- Logging ------------------------------------------------------------------ +LOG_LEVEL=debug -# --- Mailer Configuration --- -MAIL_HOST=smtp.mailtrap.io # Example: smtp.gmail.com, smtp.mailtrap.io -MAIL_PORT=2525 # Example: 587 (TLS), 465 (SSL), 2525 (Mailtrap) -MAIL_USERNAME=your_username # Your SMTP username -MAIL_PASSWORD=your_password # Your SMTP password -MAIL_SENDER=no-reply@tusk.com # The "From" email address \ No newline at end of file +# --- Mailer ------------------------------------------------------------------- +MAIL_HOST=smtp.mailtrap.io +MAIL_PORT=2525 +MAIL_USERNAME= +MAIL_PASSWORD= +MAIL_SENDER=no-reply@tusk.com diff --git a/.github/workflows/go.yml b/.github/workflows/go.yml index 170d0c9..4246dbc 100755 --- a/.github/workflows/go.yml +++ b/.github/workflows/go.yml @@ -12,14 +12,38 @@ permissions: jobs: build: runs-on: ubuntu-latest + + services: + # Integration tests run against a real PostgreSQL. SQLite is not used as a + # stand-in: it diverges on row-level security, partial indexes and ON + # CONFLICT semantics — precisely the behaviour worth testing. + postgres: + image: postgres:15 + env: + POSTGRES_USER: postgres + POSTGRES_PASSWORD: postgres + POSTGRES_DB: tusk_test + ports: + - 5432:5432 + # Without a health check the test step can start before Postgres accepts + # connections, producing an intermittent failure that looks like a bug in + # the code rather than a race in the workflow. + options: >- + --health-cmd pg_isready + --health-interval 10s + --health-timeout 5s + --health-retries 5 + steps: - name: Checkout code uses: actions/checkout@v4 + # Reading the version from go.mod keeps CI and the module in lockstep. + # A hardcoded version silently diverges the moment go.mod is bumped. - name: Set up Go uses: actions/setup-go@v5 with: - go-version: '1.24.x' + go-version-file: 'go.mod' check-latest: true - name: Verify dependencies @@ -28,5 +52,14 @@ jobs: - name: Build packages run: go build -v ./... + - name: Vet + run: go vet ./... + - name: Run tests + env: + TEST_DATABASE_URL: "host=127.0.0.1 port=5432 user=postgres password=postgres dbname=tusk_test sslmode=disable TimeZone=UTC" + # Integration tests skip when no database is reachable, which is correct + # locally but disastrous here — a green build would mean nothing was + # tested. This turns a skip into a failure. + REQUIRE_DB_TESTS: "1" run: go test -v -race ./... diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 0000000..01a7c71 --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1,118 @@ +# Tusk + +Opinionated Go backend framework for REST APIs. Huma v2 (OpenAPI 3.1) + chi + GORM + PostgreSQL. +Module path: `github.com/codetheuri/tusk` + +Full guides live in [`docs/`](docs/). This file covers what you cannot infer from reading the code. + +--- + +## Commands + +```bash +make dev # hot-reload server (Air) +make run # run without hot reload +make build # production binary → ./bin/api +make test # all tests; integration tests skip without a database +make test-integration # same, but a missing database is a failure +make test-db-setup # create the local tusk_test database (one-off) +make vet # static analysis +make migrate-up # apply migrations +make migrate-down # roll back one +make auth-sync # push code-declared permissions into the DB +``` + +`make auth-sync` is **required after adding or renaming any permission**. Permissions are declared in code and only reach the database through this command — a new permission simply won't authorize anything until it is run. + +--- + +## Architecture — one direction, no exceptions + +``` +Handler → Service → Repository +``` + +- **Handlers** bind and validate input, call a service, return a DTO. No business rules. No SQL. No direct DB access. +- **Services** hold business rules. They must never import `net/http` or know about status codes — they return domain errors. This is what makes them testable without a request. +- **Repositories** touch storage and nothing else. Always accept `context.Context`. No business rules. + +Every dependency arrives through a constructor (`NewService(repo, cfg)`). The only global is the `authz` permission registry, and that is populated at `init()` and read-only afterwards. + +If you find yourself wanting to call a repository from a handler, or check a permission inside a service, the design has drifted — say so rather than working around it. + +--- + +## Design philosophy + +Tusk should feel like **Go**, not Laravel, Spring Boot, or ASP.NET. + +- No framework magic. No hidden behaviour. No reflection unless genuinely unavoidable. +- *"A little copying is better than a little dependency."* Prefer duplication over premature abstraction. +- Before introducing any abstraction, ask: **does this reduce complexity, or merely hide it?** +- Never build an abstraction for a future possibility. Build it when the second real case arrives. +- Every package has one clear responsibility. Keep exported APIs small and intentional. +- If a design becomes hard to explain, it is probably too complicated. + +Prefer the standard library. Return errors, don't panic — except for unrecoverable startup failures. Never use global mutable state. + +--- + +## Adding a module + +Domains live in `internal//`: + +``` +model.go GORM entities +dto.go Huma request/response shapes — kept separate from models on purpose +permissions.go permission constants + []authz.Permission + init() registration +repository_*.go storage +service_*.go business rules +handler_*.go HTTP boundary +router.go huma.Register calls, wiring, guards +``` + +`internal/auth` is the reference implementation — read it before writing a new module. Register the module's routes from `cmd/api/main.go`, against `application.API()` — `pkg/app` deliberately registers none itself. Then run `make auth-sync`. + +Protect a route with `guard.Protected(op, PermSomething)`. Public routes simply omit it. + +--- + +## Conventions worth stating + +**Responses** are uniform: `{success, message, data}` on success, `{success, message, errors}` on failure, built through `pkg/response`. Never expose internal error text to clients — log the detail, return something a user can act on. + +**Documentation** — every exported type, function, and package carries a doc comment explaining *why it exists*, not just what it does. + +**Security** — validate all input, never trust client data, use prepared queries, never log secrets, bcrypt or Argon2 for passwords. + +**Performance** — prefer clarity. Avoid premature optimization, unnecessary allocation, reflection, and goroutine leaks. Benchmark before optimizing. + +--- + +## Current state + +- **PostgreSQL only.** `LoadConfig` rejects any other driver at startup. Do not add dialect branching. +- **Primary keys are UUIDv7**, generated by `pkg/id`, never by the database. New models take `uuid.UUID` with `gorm:"type:uuid;primaryKey"` and a `BeforeCreate` hook that assigns **only when the ID is zero** — a client that minted its own ID offline must keep it. +- **Multi-tenancy is opt-in per model.** A model joins by implementing `tenant.TenantColumn() string`; everything else is untouched, and an application that opts nothing in behaves as though `pkg/tenant` were not installed. The scoping callbacks are registered for every connection by `database.Connect`, so a tenanted model cannot be added later and silently go unscoped. A tenanted query with no tenant in context **errors** rather than returning every row — use `tenant.Unscoped(db)` to mean it. See `docs/roadmap.md` §3.2. +- **Server-rendered pages** use `pkg/view` (layout inheritance, buffered rendering, parse-at-startup) and `pkg/session` (database-backed cookie sessions, hashed tokens, rotation). Never render straight to a `ResponseWriter`, and never store a session token unhashed. See `docs/roadmap.md` §3.3. +- **One PostgreSQL driver: pgx.** `gorm.io/driver/postgres` uses it, so config, the migrate CLI and the test harness all use it too. Do not reintroduce `lib/pq` — the two disagree on placeholder handling, and a schema built by one driver and used by another is how that surfaced. +- Production hardening work is tracked in [`docs/roadmap.md`](docs/roadmap.md). Phases 1–4 are complete; §3 lists what is planned next. +- **Integration tests need PostgreSQL.** They connect via `internal/testdb`, defaulting to `127.0.0.1:5434/tusk_test` and overridable with `TEST_DATABASE_URL`. They skip when no database is reachable, so `go test ./...` always works — but CI sets `REQUIRE_DB_TESTS=1` to turn a skip into a failure. The harness serialises database tests across packages with an advisory lock, since `go test ./...` runs packages in parallel and each one truncates the schema. Anything touching a query belongs in an integration test: GORM builds SQL at runtime, so a broken join compiles and vets cleanly. + +--- + +## Working with me on this project + +I am using Tusk to learn advanced Go architecture. When you improve my code, **do not just rewrite it**. Explain: + +- why the original was less idiomatic +- why the new design is preferred, and what Go principle it follows +- whether it improves readability, simplicity, testability, or maintainability + +Teaching me matters more than producing code quickly. + +--- + +## Other agent tooling + +`.agents/rules/` and `.agents/workflows/` hold the same philosophy in Windsurf's format. They are kept deliberately — this file is authoritative for Claude Code, `.agents/` for Windsurf. They overlap by design rather than through a generator, per the copying-over-dependency rule above. If you change a rule here that contradicts `.agents/`, update both. diff --git a/Makefile b/Makefile index ca78a9d..80dbdbb 100644 --- a/Makefile +++ b/Makefile @@ -1,4 +1,4 @@ -.PHONY: dev run build test coverage vet migrate-up migrate-down migrate-status clean help +.PHONY: dev run build test test-unit test-integration test-db-setup coverage vet migrate-up migrate-down migrate-reset migrate-status auth-sync auth-sync-prune clean help # ============================================================================== # Development commands @@ -18,9 +18,24 @@ build: go build -o ./bin/api ./cmd/api/main.go @echo "Built ./bin/api successfully!" -## test: run unit tests across all packages +## test: run all tests (integration tests skip if no database is reachable) test: - go test -v -race ./... + go test -race ./... + +## test-unit: run only tests that need no database +test-unit: + go test -race -short ./config/... ./pkg/... ./internal/middleware/... + +## test-integration: run all tests and FAIL if the test database is unreachable +## Requires PostgreSQL. Override the target with TEST_DATABASE_URL. +test-integration: + REQUIRE_DB_TESTS=1 go test -race -count=1 ./... + +## test-db-setup: create the local test database (one-off) +test-db-setup: + @psql "$${TEST_ADMIN_URL:-postgres://root:root@127.0.0.1:5434/postgres}" \ + -c "CREATE DATABASE tusk_test" 2>/dev/null && echo "Created tusk_test" \ + || echo "tusk_test already exists (or psql is unavailable)" ## coverage: run tests and generate coverage report coverage: @@ -43,6 +58,10 @@ migrate-up: migrate-down: go run ./cmd/migrate/main.go down +## migrate-reset: revert every migration, dropping the schema +migrate-reset: + go run ./cmd/migrate/main.go reset + ## migrate-status: check the status of database migrations migrate-status: go run ./cmd/migrate/main.go status diff --git a/README.md b/README.md index 5804f40..c04d844 100755 --- a/README.md +++ b/README.md @@ -20,6 +20,7 @@ Explore the full documentation guides in the [`docs/`](docs/) directory: - 🗄️ **[Database & Migrations](docs/database-and-migrations.md)** - GORM connectivity, seeder tools, and schema migration CLI (`cmd/migrate`). - 🔍 **[Querying, Filtering & Pagination](docs/querying-and-pagination.md)** - Dynamic searching, sorting, field filtering, and metadata envelopes (`pkg/query`). - 📬 **[Standardized Responses & Error Handling](docs/responses-and-errors.md)** - Uniform JSON response structure (`pkg/response`) and status code conventions. +- 🗺️ **[Roadmap & Known Gaps](docs/roadmap.md)** - Production hardening backlog, planned UUIDv7 keys, optional multi-tenancy, and server-rendered page support. --- diff --git a/cmd/api/main.go b/cmd/api/main.go index d10317d..4d5869b 100755 --- a/cmd/api/main.go +++ b/cmd/api/main.go @@ -1,30 +1,67 @@ +// Command api is the Tusk HTTP server entrypoint. package main import ( - "os" + "github.com/danielgtaylor/huma/v2" - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/internal/app" - "github.com/codetheuri/tusk/pkg/logger" + "github.com/codetheuri/tusk/v2/config" + "github.com/codetheuri/tusk/v2/internal/auth" + "github.com/codetheuri/tusk/v2/pkg/app" + "github.com/codetheuri/tusk/v2/pkg/logger" ) func main() { - log := logger.NewConsoleLogger() + // A bootstrap logger, because configuration must be loaded before we know + // which logger the environment wants — and a configuration failure still + // needs somewhere to be reported. + log := logger.NewTextLogger("info") cfg, err := config.LoadConfig() if err != nil { log.Fatal("Failed to load configuration", err) - os.Exit(1) } - application, err := app.New(cfg, log) + // Now that the environment is known: JSON in production, readable text + // locally. + log = logger.New(cfg.IsProduction(), cfg.LOG_LEVEL) + + application, err := app.New(app.Options{ + Config: cfg, + Logger: log, + Title: "Tusk Backend API", + Version: "1.0.0", + Description: "## Official Tusk Enterprise Backend API Documentation\n\nWelcome to the developer documentation for Tusk. Explore Identity, Authentication, and RBAC endpoints below.", + Contact: &huma.Contact{ + Name: "API Support", + Email: "theurij113@gmail.com", + }, + Tags: []*huma.Tag{ + {Name: "Authentication", Description: "User registration, login, token refresh, logout, and self profile operations"}, + {Name: "Users", Description: "User account management and listing"}, + {Name: "Roles", Description: "Security role management (CRUD)"}, + {Name: "Role Permissions", Description: "Attaching and detaching permission strings to/from roles"}, + {Name: "User Roles", Description: "Assigning and revoking security roles to/from users"}, + {Name: "Permissions", Description: "System permission catalog listing"}, + }, + TagGroups: []app.TagGroup{{ + Name: "IAM", + Tags: []string{ + "Authentication", "Users", "Roles", + "Role Permissions", "User Roles", "Permissions", + }, + }}, + }) if err != nil { log.Fatal("Application setup failed", err) - os.Exit(1) } + // Routes are registered here rather than inside app.New. Tusk's own auth + // module is just the first consumer of the framework, not a privileged part + // of it — a service that wants different authentication registers its own + // module in exactly this place. + auth.RegisterRoutes(application.API(), application.DB(), cfg, log) + if err := application.Run(); err != nil { log.Fatal("Server exited with error", err) - os.Exit(1) } } diff --git a/cmd/migrate/main.go b/cmd/migrate/main.go index b72d7de..53fc1d1 100755 --- a/cmd/migrate/main.go +++ b/cmd/migrate/main.go @@ -1,3 +1,6 @@ +// Command migrate applies, reverts, and reports database schema migrations. +// +// Usage: go run ./cmd/migrate [up|down|status] package main import ( @@ -7,12 +10,11 @@ import ( "log" "os" - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/database" - "github.com/pressly/goose/v3" - // Ensure drivers are loaded - _ "github.com/lib/pq" - _ "github.com/go-sql-driver/mysql" + // PostgreSQL driver, registered for its side effects. + _ "github.com/jackc/pgx/v5/stdlib" + + "github.com/codetheuri/tusk/v2/config" + "github.com/codetheuri/tusk/v2/database" ) func main() { @@ -20,10 +22,9 @@ func main() { args := flag.Args() if len(args) < 1 { - fmt.Println("Usage: go run ./cmd/migrate/main.go [up|down|status]") + fmt.Println("Usage: go run ./cmd/migrate [up|down|status]") os.Exit(1) } - command := args[0] cfg, err := config.LoadConfig() @@ -31,41 +32,43 @@ func main() { log.Fatalf("failed to load configuration: %v", err) } - driver := cfg.DBDriver - if driver == "pgsql" { - driver = "postgres" - } - - db, err := sql.Open(driver, cfg.DbURL) + // pgx, whatever DB_DRIVER is spelled as. Tusk is PostgreSQL-only (config + // rejects anything else), and gorm.io/driver/postgres opens the running + // application's connections through pgx — so applying migrations through a + // second driver would mean the schema is built by one and used by another. + db, err := sql.Open("pgx", cfg.DbURL) if err != nil { - log.Fatalf("Failed to connect to the database: %v", err) + log.Fatalf("failed to connect to the database: %v", err) } defer db.Close() - // Tell goose to use the embedded migrations - goose.SetBaseFS(database.EmbedMigrations) - - if err := goose.SetDialect(driver); err != nil { - log.Fatalf("Failed to set dialect: %v", err) + // Fail here rather than at the first query: sql.Open does not actually + // connect, so without a ping a bad DSN surfaces as a confusing migration + // error instead of a connection error. + if err := db.Ping(); err != nil { + log.Fatalf("failed to reach the database: %v", err) } + // The goose wiring — dialect, embedded FS, and which directory belongs to + // which driver — lives in package database so this CLI and the application + // cannot disagree about where migrations come from. switch command { case "up": - if err := goose.Up(db, "migrations"); err != nil { - log.Fatalf("Goose up failed: %v", err) - } + err = database.RunMigrations(db, cfg.DBDriver) case "down": - if err := goose.Down(db, "migrations"); err != nil { - log.Fatalf("Goose down failed: %v", err) - } + err = database.RollbackMigration(db, cfg.DBDriver) + case "reset": + err = database.ResetMigrations(db, cfg.DBDriver) case "status": - if err := goose.Status(db, "migrations"); err != nil { - log.Fatalf("Goose status failed: %v", err) - } + err = database.MigrationStatus(db, cfg.DBDriver) default: fmt.Printf("Unknown command: %s\n", command) os.Exit(1) } - log.Printf("Migration command '%s' completed successfully.", command) + if err != nil { + log.Fatalf("migration %q failed: %v", command, err) + } + + log.Printf("Migration command %q completed successfully.", command) } diff --git a/cmd/tusk/main.go b/cmd/tusk/main.go index 9ba5aab..bca1d8e 100644 --- a/cmd/tusk/main.go +++ b/cmd/tusk/main.go @@ -7,12 +7,12 @@ import ( "os" "time" - "github.com/codetheuri/tusk/config" + "github.com/codetheuri/tusk/v2/config" // Import modules to trigger explicit permission registration in init() - _ "github.com/codetheuri/tusk/internal/auth" - appDatabase "github.com/codetheuri/tusk/internal/platform/database" - "github.com/codetheuri/tusk/pkg/authz" - "github.com/codetheuri/tusk/pkg/logger" + "github.com/codetheuri/tusk/v2/database" + _ "github.com/codetheuri/tusk/v2/internal/auth" + "github.com/codetheuri/tusk/v2/pkg/authz" + "github.com/codetheuri/tusk/v2/pkg/logger" ) func main() { @@ -54,7 +54,7 @@ func handleAuthCommand(args []string) { log.Fatal("Failed to load configuration", err) } - db, err := appDatabase.NewGoRMDB(cfg, log) + db, err := database.Connect(cfg, log) if err != nil { log.Fatal("Failed to connect to database", err) } diff --git a/config/config.go b/config/config.go index bafa9da..3387230 100755 --- a/config/config.go +++ b/config/config.go @@ -7,11 +7,8 @@ import ( "strings" "time" - - _ "github.com/go-sql-driver/mysql" // MySQL driver + _ "github.com/jackc/pgx/v5/stdlib" // the one PostgreSQL driver, matching gorm.io/driver/postgres "github.com/joho/godotenv" - _ "github.com/lib/pq" // PostgreSQL driver - _ "gorm.io/driver/sqlite" // SQLite driver ) type Config struct { @@ -24,23 +21,41 @@ type Config struct { ServerPort int LOG_LEVEL string JWTSecret string - AccessTokenTTL time.Duration + AccessTokenTTL time.Duration AppName string AppVersion string AppMode string DbURL string + DBSSLMode string DBMaxIdleConns int DBMaxOpenConns int DBConnMaxLifetime int - CORSOrigins []string + CORSOrigins []string + + // HTTP server limits. + // ReadTimeout must accommodate the largest request a client may legitimately + // send over a slow connection — a mobile upload on 2G, for example. Set it too + // low and such requests are severed mid-flight with no useful error. + ReadTimeout time.Duration + WriteTimeout time.Duration + IdleTimeout time.Duration + ShutdownTimeout time.Duration + MaxRequestBodyBytes int64 + + // Rate limiting, per client IP. + RateLimitBurst float64 // bucket capacity — how large a spike is tolerated + RateLimitRPS float64 // sustained refill rate, requests per second + + // DocsEnabled controls whether /docs and /openapi.json are served. + // Defaults to false in production. API endpoints are unaffected either way. + DocsEnabled bool //mailer config MailerHost string MailerPort int - MailerUsername string + MailerUsername string MailerPassword string MailerSender string - } func LoadConfig() (*Config, error) { @@ -49,15 +64,13 @@ func LoadConfig() (*Config, error) { return nil, fmt.Errorf("error loading .env file: %w", err) } cfg := &Config{ - DBUser: os.Getenv("DB_USER"), - DBPass: os.Getenv("DB_PASS"), - DBHost: os.Getenv("DB_HOST"), - // DBPort: os.Getenv("DB_PORT"), + DBUser: os.Getenv("DB_USER"), + DBPass: os.Getenv("DB_PASS"), + DBHost: os.Getenv("DB_HOST"), DBName: os.Getenv("DB_NAME"), DBDriver: os.Getenv("DB_DRIVER"), LOG_LEVEL: os.Getenv("LOG_LEVEL"), JWTSecret: os.Getenv("JWT_SECRET"), - // AccessTokenTTL: os.Getenv("ACCESS_TOKEN_TTL"), AppName: os.Getenv("APP_NAME"), AppVersion: os.Getenv("APP_VERSION"), AppMode: os.Getenv("APP_MODE"), @@ -70,25 +83,28 @@ func LoadConfig() (*Config, error) { MailerUsername: os.Getenv("MAIL_USERNAME"), MailerPassword: os.Getenv("MAIL_PASSWORD"), MailerSender: os.Getenv("MAIL_SENDER"), + } + // A short secret is not a weak secret in the way a short password is — it is a + // forgeable one. HS256 signatures are only as strong as the key, so anything + // brute-forceable lets an attacker mint valid tokens for any user. 32 bytes is + // the practical floor for HMAC-SHA256. + if len(cfg.JWTSecret) < minJWTSecretLength { + if cfg.JWTSecret == "" { + return nil, fmt.Errorf("JWT_SECRET not set in .env") + } + return nil, fmt.Errorf("JWT_SECRET must be at least %d characters (got %d)", minJWTSecretLength, len(cfg.JWTSecret)) + } + accessTokenTTLStr := os.Getenv("ACCESS_TOKEN_TTL") + if accessTokenTTLStr == "" { - - + accessTokenTTLStr = "24h" } - JWTSecret := os.Getenv("JWT_SECRET") - if JWTSecret == "" { - return nil, fmt.Errorf("JWT_SECRET not set in .env") + // Parse the duration string (e.g., "3600s", "1h", "24h") + parsedTTL, err := time.ParseDuration(accessTokenTTLStr) + if err != nil { + return nil, fmt.Errorf("invalid ACCESS_TOKEN_TTL value: %s, error: %w", accessTokenTTLStr, err) } - accessTokenTTLStr := os.Getenv("ACCESS_TOKEN_TTL") - if accessTokenTTLStr == "" { - - accessTokenTTLStr = "24h" - } - // Parse the duration string (e.g., "3600s", "1h", "24h") - parsedTTL, err := time.ParseDuration(accessTokenTTLStr) - if err != nil { - return nil, fmt.Errorf("invalid ACCESS_TOKEN_TTL value: %s, error: %w", accessTokenTTLStr, err) - } - cfg.AccessTokenTTL = parsedTTL + cfg.AccessTokenTTL = parsedTTL if cfg.DBDriver == "" { return nil, fmt.Errorf("DB_DRIVER not set in .env") @@ -115,9 +131,9 @@ func LoadConfig() (*Config, error) { return nil, fmt.Errorf("invalid SERVER_PORT value %s: %w", serverPortStr, err) } cfg.ServerPort = serverPort - //mail port + //mail port mailerPortStr := os.Getenv("MAIL_PORT") - if mailerPortStr != "" { + if mailerPortStr != "" { mailPort, err := strconv.Atoi(mailerPortStr) if err != nil { return nil, fmt.Errorf("invalid MAIL_PORT value %s: %w", mailerPortStr, err) @@ -125,13 +141,26 @@ func LoadConfig() (*Config, error) { cfg.MailerPort = mailPort } - //basic validation - if cfg.DBDriver != "sqlite" && (cfg.DBUser == "" || cfg.DBPass == "" || cfg.DBHost == "" || cfg.DBName == "") { - return nil, fmt.Errorf("missing required database configuration") + // Name the variables that are actually missing. + // + // An error like "missing required database configuration" is true but useless: + // it sends the reader to re-read a whole file looking for which of four keys is + // blank. Listing them turns a hunt into a fix — and this is exactly where a + // typo in a key name (DB_PASSWORD instead of DB_PASS) surfaces. + required := map[string]string{ + "DB_HOST": cfg.DBHost, + "DB_NAME": cfg.DBName, + "DB_USER": cfg.DBUser, + "DB_PASS": cfg.DBPass, + } + var missing []string + for _, key := range []string{"DB_HOST", "DB_NAME", "DB_USER", "DB_PASS"} { + if strings.TrimSpace(required[key]) == "" { + missing = append(missing, key) + } } - //sqlite - if cfg.DBDriver == "sqlite" && cfg.DBName == "" { - return nil, fmt.Errorf("DB_NAME not set for sqlite driver (should be file path)") + if len(missing) > 0 { + return nil, fmt.Errorf("missing required database configuration: %s", strings.Join(missing, ", ")) } if val := os.Getenv("DB_MAX_IDLE_CONNS"); val != "" { if i, err := strconv.Atoi(val); err == nil { @@ -155,30 +184,105 @@ func LoadConfig() (*Config, error) { } else { cfg.CORSOrigins = []string{} } - //dsn based on DB driver + + // --- HTTP server limits --- + cfg.ReadTimeout = envDuration("READ_TIMEOUT", 30*time.Second) + cfg.WriteTimeout = envDuration("WRITE_TIMEOUT", 60*time.Second) + cfg.IdleTimeout = envDuration("IDLE_TIMEOUT", 120*time.Second) + cfg.ShutdownTimeout = envDuration("SHUTDOWN_TIMEOUT", 30*time.Second) + cfg.MaxRequestBodyBytes = int64(envInt("MAX_REQUEST_BODY_BYTES", 10<<20)) // 10 MiB + + // --- Rate limiting --- + cfg.RateLimitBurst = float64(envInt("RATE_LIMIT_BURST", 30)) + cfg.RateLimitRPS = float64(envInt("RATE_LIMIT_RPS", 10)) + + // --- API documentation --- + // Off by default in production; on by default everywhere else. An explicit + // DOCS_ENABLED always wins, so production docs remain possible deliberately. + cfg.DocsEnabled = envBool("DOCS_ENABLED", !cfg.IsProduction()) + + // --- Database TLS --- + // Defaults to require in production: a hardcoded "disable" silently sends + // credentials and data over an unencrypted socket with no way to turn it on. + defaultSSL := "disable" + if cfg.IsProduction() { + defaultSSL = "require" + } + cfg.DBSSLMode = getEnvOr("DB_SSLMODE", defaultSSL) + // DSN. PostgreSQL only — see the default case for why. switch cfg.DBDriver { - case "mysql": - cfg.DbURL = fmt.Sprintf("%s:%s@tcp(%s:%s)/%s?charset=utf8mb4&parseTime=True&loc=Local", - cfg.DBUser, - cfg.DBPass, - cfg.DBHost, - cfg.DBPort, - cfg.DBName, - ) case "postgres", "pgsql": - cfg.DbURL = fmt.Sprintf("host=%s user=%s password=%s dbname=%s port=%s sslmode=disable TimeZone=UTC", + cfg.DbURL = fmt.Sprintf("host=%s user=%s password=%s dbname=%s port=%s sslmode=%s TimeZone=UTC", cfg.DBHost, cfg.DBUser, cfg.DBPass, cfg.DBName, cfg.DBPort, + cfg.DBSSLMode, ) - case "sqlite": - cfg.DbURL = cfg.DBName - default: - return nil, fmt.Errorf("unsupported DB_DRIVER: %s", cfg.DBDriver) + // Fail here rather than at the first query. MySQL and SQLite support was + // dropped when primary keys became UUIDs: neither has a native UUID type, + // so their schema could no longer describe the same rows as the Go models. + // A driver that connects but cannot read its own tables is worse than one + // that refuses to start. + return nil, fmt.Errorf( + "unsupported DB_DRIVER %q: Tusk targets PostgreSQL only. Set DB_DRIVER=postgres", + cfg.DBDriver) } return cfg, nil } + +// minJWTSecretLength is the practical floor for an HS256 signing key. +const minJWTSecretLength = 32 + +// IsProduction reports whether the application is running in production mode. +// Behaviour that should differ between environments — documentation exposure, +// log formatting, error verbosity — keys off this rather than inspecting +// APP_MODE in scattered places. +func (c *Config) IsProduction() bool { + mode := strings.ToLower(strings.TrimSpace(c.AppMode)) + return mode == "production" || mode == "prod" +} + +// envInt reads an integer environment variable, falling back to def when unset +// or unparseable. Configuration should never fail closed on a typo in an +// optional tuning knob; required values are validated explicitly instead. +func envInt(key string, def int) int { + if raw := os.Getenv(key); raw != "" { + if v, err := strconv.Atoi(strings.TrimSpace(raw)); err == nil { + return v + } + } + return def +} + +// envDuration reads a Go duration string such as "30s" or "2m". +func envDuration(key string, def time.Duration) time.Duration { + if raw := os.Getenv(key); raw != "" { + if v, err := time.ParseDuration(strings.TrimSpace(raw)); err == nil { + return v + } + } + return def +} + +// envBool reads a boolean environment variable, accepting the forms +// strconv.ParseBool understands ("1", "true", "TRUE", "0", "false", …). +func envBool(key string, def bool) bool { + if raw := os.Getenv(key); raw != "" { + if v, err := strconv.ParseBool(strings.TrimSpace(raw)); err == nil { + return v + } + } + return def +} + +// getEnvOr returns the environment variable value, or def when unset or empty. +func getEnvOr(key, def string) string { + if v := strings.TrimSpace(os.Getenv(key)); v != "" { + return v + } + return def +} diff --git a/config/config_test.go b/config/config_test.go index 824faca..667a418 100644 --- a/config/config_test.go +++ b/config/config_test.go @@ -1,19 +1,30 @@ package config import ( - "os" + "strings" "testing" + "time" ) +// validSecret is long enough to satisfy minJWTSecretLength. +const validSecret = "test-secret-key-that-is-long-enough-for-hs256" + +// baseEnv sets the minimum required configuration for a successful load. +// t.Setenv restores the previous value automatically when the test ends, which +// is why these tests do not manage cleanup by hand. +func baseEnv(t *testing.T) { + t.Helper() + t.Setenv("JWT_SECRET", validSecret) + t.Setenv("DB_DRIVER", "postgres") + t.Setenv("DB_HOST", "localhost") + t.Setenv("DB_PORT", "5432") + t.Setenv("DB_NAME", "tusk_test") + t.Setenv("DB_USER", "postgres") + t.Setenv("DB_PASS", "postgres") +} + func TestLoadConfig_Defaults(t *testing.T) { - os.Setenv("JWT_SECRET", "test-secret-key-12345") - os.Setenv("DB_DRIVER", "sqlite") - os.Setenv("DB_NAME", ":memory:") - defer func() { - os.Unsetenv("JWT_SECRET") - os.Unsetenv("DB_DRIVER") - os.Unsetenv("DB_NAME") - }() + baseEnv(t) cfg, err := LoadConfig() if err != nil { @@ -23,7 +34,197 @@ func TestLoadConfig_Defaults(t *testing.T) { if cfg.ServerPort != 8080 { t.Errorf("expected default ServerPort 8080, got %d", cfg.ServerPort) } - if cfg.DBDriver != "sqlite" { - t.Errorf("expected DBDriver sqlite, got %s", cfg.DBDriver) + if cfg.DBDriver != "postgres" { + t.Errorf("expected DBDriver postgres, got %s", cfg.DBDriver) + } + if cfg.ReadTimeout != 30*time.Second { + t.Errorf("expected default ReadTimeout 30s, got %s", cfg.ReadTimeout) + } + if cfg.MaxRequestBodyBytes != 10<<20 { + t.Errorf("expected default body limit 10MiB, got %d", cfg.MaxRequestBodyBytes) + } + if cfg.RateLimitRPS != 10 || cfg.RateLimitBurst != 30 { + t.Errorf("expected default rate limit 10rps/30burst, got %v/%v", cfg.RateLimitRPS, cfg.RateLimitBurst) + } +} + +// A short signing key is forgeable, not merely weak — loading must fail rather +// than start a server that will happily accept minted tokens. +func TestLoadConfig_RejectsShortJWTSecret(t *testing.T) { + baseEnv(t) + t.Setenv("JWT_SECRET", "too-short") + + _, err := LoadConfig() + if err == nil { + t.Fatal("expected short JWT_SECRET to be rejected, got nil error") + } + if !strings.Contains(err.Error(), "at least") { + t.Errorf("expected error to state the minimum length, got: %v", err) + } +} + +func TestLoadConfig_RejectsMissingJWTSecret(t *testing.T) { + baseEnv(t) + t.Setenv("JWT_SECRET", "") + + if _, err := LoadConfig(); err == nil { + t.Fatal("expected missing JWT_SECRET to be rejected, got nil error") + } +} + +func TestIsProduction(t *testing.T) { + cases := map[string]bool{ + "production": true, + "prod": true, + "PRODUCTION": true, + " prod ": true, + "dev": false, + "staging": false, + "": false, + } + for mode, want := range cases { + cfg := &Config{AppMode: mode} + if got := cfg.IsProduction(); got != want { + t.Errorf("IsProduction(%q) = %v, want %v", mode, got, want) + } + } +} + +// Documentation exposes the full route and schema map, so it defaults off in +// production and on everywhere else — unless explicitly overridden. +func TestDocsEnabled_DefaultsByMode(t *testing.T) { + t.Run("off in production", func(t *testing.T) { + baseEnv(t) + t.Setenv("APP_MODE", "production") + + cfg, err := LoadConfig() + if err != nil { + t.Fatalf("load failed: %v", err) + } + if cfg.DocsEnabled { + t.Error("expected docs disabled by default in production") + } + }) + + t.Run("on in development", func(t *testing.T) { + baseEnv(t) + t.Setenv("APP_MODE", "dev") + + cfg, err := LoadConfig() + if err != nil { + t.Fatalf("load failed: %v", err) + } + if !cfg.DocsEnabled { + t.Error("expected docs enabled by default outside production") + } + }) + + t.Run("explicit override wins in production", func(t *testing.T) { + baseEnv(t) + t.Setenv("APP_MODE", "production") + t.Setenv("DOCS_ENABLED", "true") + + cfg, err := LoadConfig() + if err != nil { + t.Fatalf("load failed: %v", err) + } + if !cfg.DocsEnabled { + t.Error("expected explicit DOCS_ENABLED=true to override the production default") + } + }) +} + +// Database traffic must not silently fall back to an unencrypted connection in +// production just because DB_SSLMODE was left unset. +func TestDBSSLMode_DefaultsToRequireInProduction(t *testing.T) { + baseEnv(t) + t.Setenv("APP_MODE", "production") + + cfg, err := LoadConfig() + if err != nil { + t.Fatalf("load failed: %v", err) + } + if cfg.DBSSLMode != "require" { + t.Errorf("expected DBSSLMode 'require' in production, got %q", cfg.DBSSLMode) + } +} + +func TestEnvHelpers_FallBackOnGarbage(t *testing.T) { + t.Setenv("TUSK_TEST_INT", "not-a-number") + if got := envInt("TUSK_TEST_INT", 42); got != 42 { + t.Errorf("envInt fallback = %d, want 42", got) + } + + t.Setenv("TUSK_TEST_DUR", "not-a-duration") + if got := envDuration("TUSK_TEST_DUR", time.Minute); got != time.Minute { + t.Errorf("envDuration fallback = %s, want 1m", got) + } + + t.Setenv("TUSK_TEST_BOOL", "maybe") + if got := envBool("TUSK_TEST_BOOL", true); !got { + t.Error("envBool fallback = false, want true") + } +} + +// Unsupported drivers must be refused at startup, not at the first query. +// MySQL and SQLite were dropped when primary keys became UUIDs — neither has a +// native UUID type, so their schema can no longer describe the same rows as the +// Go models. A process that connects but cannot read its own tables is worse +// than one that refuses to start. +func TestLoadConfig_RejectsUnsupportedDrivers(t *testing.T) { + for _, driver := range []string{"mysql", "sqlite", "mssql", ""} { + t.Run(driver, func(t *testing.T) { + baseEnv(t) + t.Setenv("DB_DRIVER", driver) + + if _, err := LoadConfig(); err == nil { + t.Fatalf("expected driver %q to be rejected", driver) + } + }) + } +} + +func TestLoadConfig_AcceptsPgsqlAlias(t *testing.T) { + baseEnv(t) + t.Setenv("DB_DRIVER", "pgsql") + + cfg, err := LoadConfig() + if err != nil { + t.Fatalf("expected pgsql to be accepted as an alias: %v", err) + } + if !strings.Contains(cfg.DbURL, "dbname=tusk_test") { + t.Errorf("DSN does not name the database: %s", cfg.DbURL) + } +} + +// A configuration error should name what is wrong. The generic form of this +// message cost a real debugging round-trip when DB_PASSWORD was set instead of +// DB_PASS: every key looked present, and the error pointed at nothing. +func TestLoadConfig_NamesMissingDatabaseKeys(t *testing.T) { + baseEnv(t) + t.Setenv("DB_PASS", "") + + _, err := LoadConfig() + if err == nil { + t.Fatal("expected missing DB_PASS to be rejected") + } + if !strings.Contains(err.Error(), "DB_PASS") { + t.Errorf("error must name the missing key, got: %v", err) + } +} + +func TestLoadConfig_NamesAllMissingKeys(t *testing.T) { + baseEnv(t) + t.Setenv("DB_USER", "") + t.Setenv("DB_PASS", "") + + _, err := LoadConfig() + if err == nil { + t.Fatal("expected missing keys to be rejected") + } + for _, key := range []string{"DB_USER", "DB_PASS"} { + if !strings.Contains(err.Error(), key) { + t.Errorf("error should list %s, got: %v", key, err) + } } } diff --git a/internal/platform/database/gorm.go b/database/connect.go similarity index 69% rename from internal/platform/database/gorm.go rename to database/connect.go index a9c719f..19e76c9 100755 --- a/internal/platform/database/gorm.go +++ b/database/connect.go @@ -1,33 +1,34 @@ package database +// This file owns connecting to PostgreSQL. It sits beside the migration runner +// deliberately: both are "the database layer", and the previous +// internal/platform/ tier existed to hold exactly one file. Being here rather +// than under internal/ also means an application built on Tusk can open a +// connection configured the same way the framework does. + import ( "context" "fmt" "time" - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/pkg/logger" + "github.com/codetheuri/tusk/v2/config" + "github.com/codetheuri/tusk/v2/pkg/logger" + "github.com/codetheuri/tusk/v2/pkg/tenant" - "gorm.io/driver/mysql" "gorm.io/driver/postgres" - "gorm.io/driver/sqlite" "gorm.io/gorm" gormlogger "gorm.io/gorm/logger" ) -func NewGoRMDB(cfg *config.Config, log logger.Logger) (*gorm.DB, error) { +func Connect(cfg *config.Config, log logger.Logger) (*gorm.DB, error) { newLogger := NewGormLogger(log) var db *gorm.DB var err error switch cfg.DBDriver { - case "mysql": - db, err = gorm.Open(mysql.Open(cfg.DbURL), &gorm.Config{}) case "postgres", "pgsql": db, err = gorm.Open(postgres.Open(cfg.DbURL), &gorm.Config{}) - case "sqlite": - db, err = gorm.Open(sqlite.Open(cfg.DbURL), &gorm.Config{}) default: return nil, fmt.Errorf("unsupported DB_DRIVER: %s", cfg.DBDriver) } @@ -52,6 +53,17 @@ func NewGoRMDB(cfg *config.Config, log logger.Logger) (*gorm.DB, error) { log.Error("database is unreachable", err) return nil, fmt.Errorf("database is unreachable: %w", err) } + + // Installed for every application, single- and multi-tenant alike. The + // callbacks return immediately for models that do not implement + // tenant.Tenanted, so an application that opts nothing in pays nothing and + // behaves exactly as it did before. Registering here rather than leaving it + // to each application means a tenanted model cannot be added later and + // silently go unscoped because someone forgot a line of wiring. + if err := tenant.Register(db); err != nil { + return nil, fmt.Errorf("failed to register tenant callbacks: %w", err) + } + log.Info("Database connected successfully ") return db, nil } diff --git a/database/migrate.go b/database/migrate.go index 0d04eed..db7c5df 100644 --- a/database/migrate.go +++ b/database/migrate.go @@ -1,32 +1,183 @@ +// Package database owns schema migrations and seed data. package database import ( "database/sql" "embed" "fmt" + "io/fs" + "sync" "github.com/pressly/goose/v3" ) -//go:embed migrations/*.sql +// EmbedMigrations carries the SQL files into the binary so a deployed artefact +// can migrate itself with no accompanying files on disk. +// +// The pattern reaches into the dialect subdirectory; `migrations/*.sql` would +// silently match nothing and produce a no-op migration run. +// +//go:embed migrations/postgres/*.sql var EmbedMigrations embed.FS -// RunMigrations executes all embedded SQL migrations against the database. -func RunMigrations(db *sql.DB, dialect string) error { - goose.SetBaseFS(EmbedMigrations) - - // Ensure dialect is compatible with goose (e.g., pgsql -> postgres) - if dialect == "pgsql" { - dialect = "postgres" +// migrationDir maps a driver name to its migration directory. +// +// PostgreSQL is the only supported dialect. The MySQL migration was removed when +// primary keys became UUIDs: MySQL has no native UUID type, so the schema and the +// Go models could no longer describe the same thing. Rather than ship a migration +// that produces tables the application cannot read, the path is gone. The SQL +// remains in git history for anyone who needs it. +func migrationDir(dialect string) (string, string, error) { + switch dialect { + case "postgres", "pgsql": + return "postgres", "migrations/postgres", nil + default: + return "", "", fmt.Errorf( + "no migrations available for driver %q: Tusk targets PostgreSQL only "+ + "(set DB_DRIVER=postgres)", dialect) } +} - if err := goose.SetDialect(dialect); err != nil { +// Migrator applies one set of migrations. +// +// It exists so a service built on Tusk can run its own schema through the same +// machinery. Tusk's built-in migrations create Tusk's identity tables, which an +// application with different ones — users who log in by phone and have no email, +// say — does not want. Such a service supplies its own embedded filesystem here. +// +// Migrations are deliberately not shareable as library code. A migration is a +// record of a change already applied to a specific database; if it could change +// underneath you when a dependency updated, it would no longer be a record of +// anything. Each service owns its own. +type Migrator struct { + fsys fs.FS + dir string +} + +// NewMigrator returns a Migrator over a caller-supplied migration set. +// +// //go:embed migrations/*.sql +// var migrations embed.FS +// +// m := database.NewMigrator(migrations, "migrations") +// err := m.Up(db) +func NewMigrator(fsys fs.FS, dir string) *Migrator { + return &Migrator{fsys: fsys, dir: dir} +} + +// gooseMu guards goose's package-level configuration. +// +// SetBaseFS and SetDialect are global, so two Migrators running at once would +// otherwise apply one's migrations from the other's filesystem. Rare in a server, +// entirely plausible in a test binary. +var gooseMu sync.Mutex + +// prepare points goose at this Migrator's files. Callers hold gooseMu. +func (m *Migrator) prepare() error { + goose.SetBaseFS(m.fsys) + + // Always postgres: Tusk targets it exclusively, and config rejects anything + // else long before this runs. + if err := goose.SetDialect("postgres"); err != nil { return fmt.Errorf("failed to set goose dialect: %w", err) } + return nil +} - if err := goose.Up(db, "migrations"); err != nil { +// Up applies every pending migration. +func (m *Migrator) Up(db *sql.DB) error { + gooseMu.Lock() + defer gooseMu.Unlock() + + if err := m.prepare(); err != nil { + return err + } + if err := goose.Up(db, m.dir); err != nil { return fmt.Errorf("failed to run migrations: %w", err) } - return nil } + +// Down reverts the most recently applied migration. +func (m *Migrator) Down(db *sql.DB) error { + gooseMu.Lock() + defer gooseMu.Unlock() + + if err := m.prepare(); err != nil { + return err + } + return goose.Down(db, m.dir) +} + +// Reset rolls every migration back, in reverse order. +// +// Destructive by design: it drops the schema. Provided because "down" alone +// reverts a single step, so verifying that a rollback path actually works — that +// each Down truly reverses its Up, rather than merely appearing to — requires +// unwinding the whole stack and rebuilding it. +func (m *Migrator) Reset(db *sql.DB) error { + gooseMu.Lock() + defer gooseMu.Unlock() + + if err := m.prepare(); err != nil { + return err + } + return goose.DownTo(db, m.dir, 0) +} + +// Status prints which migrations have been applied. +func (m *Migrator) Status(db *sql.DB) error { + gooseMu.Lock() + defer gooseMu.Unlock() + + if err := m.prepare(); err != nil { + return err + } + return goose.Status(db, m.dir) +} + +// builtin returns a Migrator over Tusk's own migrations, rejecting any dialect +// Tusk does not support. +func builtin(dialect string) (*Migrator, error) { + _, dir, err := migrationDir(dialect) + if err != nil { + return nil, err + } + return NewMigrator(EmbedMigrations, dir), nil +} + +// RunMigrations applies all pending Tusk migrations for the configured driver. +func RunMigrations(db *sql.DB, dialect string) error { + m, err := builtin(dialect) + if err != nil { + return err + } + return m.Up(db) +} + +// RollbackMigration reverts the most recently applied Tusk migration. +func RollbackMigration(db *sql.DB, dialect string) error { + m, err := builtin(dialect) + if err != nil { + return err + } + return m.Down(db) +} + +// ResetMigrations rolls every Tusk migration back. +func ResetMigrations(db *sql.DB, dialect string) error { + m, err := builtin(dialect) + if err != nil { + return err + } + return m.Reset(db) +} + +// MigrationStatus prints the state of Tusk's migrations. +func MigrationStatus(db *sql.DB, dialect string) error { + m, err := builtin(dialect) + if err != nil { + return err + } + return m.Status(db) +} diff --git a/database/migrations/00001_identity_auth_rbac.mysql.sql b/database/migrations/00001_identity_auth_rbac.mysql.sql deleted file mode 100644 index 53a171d..0000000 --- a/database/migrations/00001_identity_auth_rbac.mysql.sql +++ /dev/null @@ -1,86 +0,0 @@ --- +goose Up -CREATE TABLE IF NOT EXISTS users ( - id BIGINT UNSIGNED AUTO_INCREMENT PRIMARY KEY, - username VARCHAR(191) NOT NULL UNIQUE, - email VARCHAR(191) NOT NULL UNIQUE, - phone VARCHAR(50) NULL UNIQUE, - password VARCHAR(255) NOT NULL, - is_super_user BOOLEAN DEFAULT FALSE, - is_active BOOLEAN DEFAULT TRUE, - is_verified BOOLEAN DEFAULT FALSE, - failed_login_attempts INT DEFAULT 0, - locked_until DATETIME(3) NULL, - last_login_at DATETIME(3) NULL, - created_at DATETIME(3) NULL, - updated_at DATETIME(3) NULL -); - -CREATE INDEX idx_users_username ON users(username); -CREATE INDEX idx_users_email ON users(email); -CREATE INDEX idx_users_phone ON users(phone); - -CREATE TABLE IF NOT EXISTS user_profiles ( - user_id BIGINT UNSIGNED PRIMARY KEY, - first_name VARCHAR(100) DEFAULT '', - last_name VARCHAR(100) DEFAULT '', - avatar VARCHAR(255) DEFAULT '', - bio TEXT, - created_at DATETIME(3) NULL, - updated_at DATETIME(3) NULL, - CONSTRAINT fk_user_profiles_user FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE -); - -CREATE TABLE IF NOT EXISTS refresh_tokens ( - id BIGINT UNSIGNED AUTO_INCREMENT PRIMARY KEY, - user_id BIGINT UNSIGNED NOT NULL, - token_hash VARCHAR(255) NOT NULL UNIQUE, - expires_at DATETIME(3) NOT NULL, - revoked_at DATETIME(3) NULL, - created_at DATETIME(3) NULL, - CONSTRAINT fk_refresh_tokens_user FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE -); - -CREATE INDEX idx_refresh_tokens_user ON refresh_tokens(user_id); -CREATE INDEX idx_refresh_tokens_hash ON refresh_tokens(token_hash); - -CREATE TABLE IF NOT EXISTS permissions ( - name VARCHAR(191) PRIMARY KEY, - description TEXT, - created_at DATETIME(3) NULL, - updated_at DATETIME(3) NULL -); - -CREATE TABLE IF NOT EXISTS roles ( - id BIGINT UNSIGNED AUTO_INCREMENT PRIMARY KEY, - name VARCHAR(191) NOT NULL UNIQUE, - description TEXT, - created_at DATETIME(3) NULL, - updated_at DATETIME(3) NULL -); - -CREATE TABLE IF NOT EXISTS role_permissions ( - role_id BIGINT UNSIGNED NOT NULL, - permission_name VARCHAR(191) NOT NULL, - created_at DATETIME(3) NULL, - PRIMARY KEY (role_id, permission_name), - CONSTRAINT fk_role_permissions_role FOREIGN KEY (role_id) REFERENCES roles(id) ON DELETE CASCADE, - CONSTRAINT fk_role_permissions_permission FOREIGN KEY (permission_name) REFERENCES permissions(name) ON DELETE CASCADE -); - -CREATE TABLE IF NOT EXISTS user_roles ( - user_id BIGINT UNSIGNED NOT NULL, - role_id BIGINT UNSIGNED NOT NULL, - created_at DATETIME(3) NULL, - PRIMARY KEY (user_id, role_id), - CONSTRAINT fk_user_roles_user FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE, - CONSTRAINT fk_user_roles_role FOREIGN KEY (role_id) REFERENCES roles(id) ON DELETE CASCADE -); - --- +goose Down -DROP TABLE IF EXISTS user_roles; -DROP TABLE IF EXISTS role_permissions; -DROP TABLE IF EXISTS roles; -DROP TABLE IF EXISTS permissions; -DROP TABLE IF EXISTS refresh_tokens; -DROP TABLE IF EXISTS user_profiles; -DROP TABLE IF EXISTS users; diff --git a/database/migrations/postgres/00001_identity_auth_rbac.sql b/database/migrations/postgres/00001_identity_auth_rbac.sql new file mode 100644 index 0000000..936a0a4 --- /dev/null +++ b/database/migrations/postgres/00001_identity_auth_rbac.sql @@ -0,0 +1,107 @@ +-- Identity, authentication and RBAC schema. +-- +-- Primary keys are UUIDv7, generated in Go by pkg/id rather than by the database. +-- Two reasons there is no DEFAULT here: +-- +-- 1. PostgreSQL only gained a native uuidv7() in version 18, and Tusk supports +-- earlier versions. gen_random_uuid() would produce v4, losing the index +-- locality that motivated v7 in the first place. +-- 2. More fundamentally, an offline-capable client must be able to create a row +-- and know its identifier before the database has ever seen it. A server-side +-- default cannot serve that case, so the application is the right place to +-- mint identifiers regardless of server version. +-- +-- Other notes on this schema: +-- * TIMESTAMPTZ throughout. A wall-clock time with no zone is a latent bug the +-- moment a server, a client, or a DST boundary disagrees. +-- * No redundant indexes on UNIQUE columns — PostgreSQL already creates a +-- B-tree to enforce every UNIQUE constraint. Only non-unique columns that are +-- actually queried get an explicit index. +-- * VARCHAR(191) is retained to match the GORM model tags. The length is a +-- MySQL utf8mb4 index-limit artefact and carries no meaning here. + +-- +goose Up +CREATE TABLE IF NOT EXISTS users ( + id UUID PRIMARY KEY, + username VARCHAR(191) NOT NULL UNIQUE, + email VARCHAR(191) NOT NULL UNIQUE, + phone VARCHAR(50) UNIQUE, + password VARCHAR(255) NOT NULL, + is_super_user BOOLEAN NOT NULL DEFAULT FALSE, + is_active BOOLEAN NOT NULL DEFAULT TRUE, + is_verified BOOLEAN NOT NULL DEFAULT FALSE, + failed_login_attempts INTEGER NOT NULL DEFAULT 0, + locked_until TIMESTAMPTZ, + last_login_at TIMESTAMPTZ, + created_at TIMESTAMPTZ, + updated_at TIMESTAMPTZ +); + +CREATE TABLE IF NOT EXISTS user_profiles ( + user_id UUID PRIMARY KEY REFERENCES users(id) ON DELETE CASCADE, + first_name VARCHAR(100) NOT NULL DEFAULT '', + last_name VARCHAR(100) NOT NULL DEFAULT '', + avatar VARCHAR(255) NOT NULL DEFAULT '', + bio TEXT, + created_at TIMESTAMPTZ, + updated_at TIMESTAMPTZ +); + +CREATE TABLE IF NOT EXISTS refresh_tokens ( + id UUID PRIMARY KEY, + user_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE, + token_hash VARCHAR(255) NOT NULL UNIQUE, + expires_at TIMESTAMPTZ NOT NULL, + revoked_at TIMESTAMPTZ, + created_at TIMESTAMPTZ +); + +-- user_id is not unique (a user may hold several active sessions), so this index +-- does real work: revoking every session for a user is a lookup by user_id. +CREATE INDEX IF NOT EXISTS idx_refresh_tokens_user ON refresh_tokens(user_id); + +CREATE TABLE IF NOT EXISTS permissions ( + name VARCHAR(191) PRIMARY KEY, + description TEXT, + created_at TIMESTAMPTZ, + updated_at TIMESTAMPTZ +); + +CREATE TABLE IF NOT EXISTS roles ( + id UUID PRIMARY KEY, + name VARCHAR(191) NOT NULL UNIQUE, + description TEXT, + created_at TIMESTAMPTZ, + updated_at TIMESTAMPTZ +); + +CREATE TABLE IF NOT EXISTS role_permissions ( + role_id UUID NOT NULL REFERENCES roles(id) ON DELETE CASCADE, + permission_name VARCHAR(191) NOT NULL REFERENCES permissions(name) ON DELETE CASCADE, + created_at TIMESTAMPTZ, + PRIMARY KEY (role_id, permission_name) +); + +-- The composite primary key already indexes (role_id, permission_name), covering +-- lookups by role. The reverse — "which roles grant this permission?" — is not, +-- so it gets its own index. +CREATE INDEX IF NOT EXISTS idx_role_permissions_permission ON role_permissions(permission_name); + +CREATE TABLE IF NOT EXISTS user_roles ( + user_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE, + role_id UUID NOT NULL REFERENCES roles(id) ON DELETE CASCADE, + created_at TIMESTAMPTZ, + PRIMARY KEY (user_id, role_id) +); + +-- Same reasoning: the primary key covers user_id, this covers "who holds this role?" +CREATE INDEX IF NOT EXISTS idx_user_roles_role ON user_roles(role_id); + +-- +goose Down +DROP TABLE IF EXISTS user_roles; +DROP TABLE IF EXISTS role_permissions; +DROP TABLE IF EXISTS roles; +DROP TABLE IF EXISTS permissions; +DROP TABLE IF EXISTS refresh_tokens; +DROP TABLE IF EXISTS user_profiles; +DROP TABLE IF EXISTS users; diff --git a/database/migrations/postgres/00002_sessions.sql b/database/migrations/postgres/00002_sessions.sql new file mode 100644 index 0000000..ba048a4 --- /dev/null +++ b/database/migrations/postgres/00002_sessions.sql @@ -0,0 +1,42 @@ +-- +goose Up + +-- Browser sessions for server-rendered pages, separate from the JWT path the API +-- uses. A session is a database row so that deleting the row ends access — the +-- thing a stateless token cannot do, and the reason an operator console needs +-- this table rather than another kind of token. +CREATE TABLE IF NOT EXISTS sessions ( + id UUID PRIMARY KEY, + + -- The SHA-256 digest of the token, never the token. A database dump + -- therefore does not hand over live sessions. 64 hex characters. + token_hash VARCHAR(64) NOT NULL UNIQUE, + + -- Which audience the session belongs to. An application with both an + -- operator console and ordinary user sessions gives them different kinds, so + -- a cookie from one can never resolve against the other. + subject_kind VARCHAR(64) NOT NULL, + + -- Who is signed in. Deliberately no foreign key: the subject may live in + -- Tusk's users table or in an application's own operators table, and this + -- table has no business deciding which. + subject_id UUID NOT NULL, + + -- Recorded for an operator reviewing active sessions. Never used to make an + -- access decision: behind a proxy both values are supplied by someone else. + ip VARCHAR(45) NOT NULL DEFAULT '', + user_agent VARCHAR(255) NOT NULL DEFAULT '', + + created_at TIMESTAMPTZ NOT NULL, + last_seen_at TIMESTAMPTZ NOT NULL, + expires_at TIMESTAMPTZ NOT NULL +); + +-- Covers "every session for this person", which is what signing someone out +-- everywhere has to do. +CREATE INDEX IF NOT EXISTS idx_sessions_subject ON sessions(subject_kind, subject_id); + +-- Covers the periodic sweep of expired rows. Nothing else deletes them. +CREATE INDEX IF NOT EXISTS idx_sessions_expires ON sessions(expires_at); + +-- +goose Down +DROP TABLE IF EXISTS sessions; diff --git a/database/migrations_test.go b/database/migrations_test.go new file mode 100644 index 0000000..ce5c955 --- /dev/null +++ b/database/migrations_test.go @@ -0,0 +1,129 @@ +package database_test + +import ( + "database/sql" + "regexp" + "testing" + + _ "github.com/jackc/pgx/v5/stdlib" + + "github.com/codetheuri/tusk/v2/database" + "github.com/codetheuri/tusk/v2/pkg/testdb" +) + +// migrationTestDB is created and dropped by this test. +// +// It deliberately does not share the database used by the other integration +// tests: this one drops every table, and Go runs different packages in parallel, +// so sharing would tear the schema out from under a concurrently running suite. +const migrationTestDB = "tusk_test_migrations" + +var dbNamePattern = regexp.MustCompile(`dbname=\S+`) + +// dsnFor rewrites the configured test DSN to point at a different database on the +// same server, so credentials and host stay in one place. +func dsnFor(name string) string { + return dbNamePattern.ReplaceAllString(testdb.DSN(), "dbname="+name) +} + +// TestMigrations_UpDownUp verifies the schema can be built, torn down, and built +// again. +// +// The Down direction is the half nobody exercises until the day a deploy has to +// be rolled back — which is the worst possible moment to discover a typo in it. +func TestMigrations_UpDownUp(t *testing.T) { + admin, err := sql.Open("pgx", dsnFor("postgres")) + if err != nil { + t.Skipf("cannot open admin connection, skipping: %v", err) + } + defer admin.Close() + + if err := admin.Ping(); err != nil { + // Same policy as testdb.Connect: skip locally, fail in CI. + if testdb.Required() { + t.Fatalf("REQUIRE_DB_TESTS is set but no database is reachable: %v", err) + } + t.Skipf("no test database reachable, skipping: %v", err) + } + + // A previous crashed run may have left the database behind. + if _, err := admin.Exec("DROP DATABASE IF EXISTS " + migrationTestDB); err != nil { + t.Fatalf("failed to drop stale test database: %v", err) + } + if _, err := admin.Exec("CREATE DATABASE " + migrationTestDB); err != nil { + t.Fatalf("failed to create test database: %v", err) + } + t.Cleanup(func() { + // Connections must be closed before the database can be dropped. + _, _ = admin.Exec("DROP DATABASE IF EXISTS " + migrationTestDB) + }) + + db, err := sql.Open("pgx", dsnFor(migrationTestDB)) + if err != nil { + t.Fatalf("failed to connect to the migration test database: %v", err) + } + + expected := []string{ + "users", "user_profiles", "refresh_tokens", + "permissions", "roles", "role_permissions", "user_roles", + "sessions", + } + + // --- Up --- + if err := database.RunMigrations(db, "postgres"); err != nil { + t.Fatalf("migrate up failed: %v", err) + } + for _, table := range expected { + if !tableExists(t, db, table) { + t.Errorf("after migrate up, table %q is missing", table) + } + } + + // --- Down --- + // Every migration, not just the most recent: rolling back one step would + // leave earlier migrations applied and prove nothing about their Down. + if err := database.ResetMigrations(db, "postgres"); err != nil { + t.Fatalf("migrate reset failed: %v", err) + } + for _, table := range expected { + if tableExists(t, db, table) { + t.Errorf("after migrate down, table %q still exists", table) + } + } + + // --- Up again --- + // Re-applying proves the Down left the schema in a state the Up can rebuild, + // rather than merely appearing to succeed. + if err := database.RunMigrations(db, "postgres"); err != nil { + t.Fatalf("migrate up after rollback failed: %v", err) + } + for _, table := range expected { + if !tableExists(t, db, table) { + t.Errorf("after re-applying migrations, table %q is missing", table) + } + } + + db.Close() +} + +// TestMigrations_UnknownDriverIsRejected guards the dialect routing: an +// unsupported driver must fail loudly rather than silently running no migrations. +func TestMigrations_UnknownDriverIsRejected(t *testing.T) { + err := database.RunMigrations(nil, "cassandra") + if err == nil { + t.Fatal("expected an unsupported driver to be rejected") + } +} + +func tableExists(t *testing.T, db *sql.DB, name string) bool { + t.Helper() + var exists bool + err := db.QueryRow( + `SELECT EXISTS (SELECT 1 FROM pg_tables WHERE schemaname='public' AND tablename=$1)`, + name, + ).Scan(&exists) + if err != nil { + t.Fatalf("failed to check for table %q: %v", name, err) + } + return exists +} diff --git a/docs/architecture.md b/docs/architecture.md index d719e10..36f3700 100644 --- a/docs/architecture.md +++ b/docs/architecture.md @@ -122,21 +122,22 @@ Infrastructure drivers that are decoupled from specific business logic belong in --- -### 2. Cross-Cutting Middleware (`internal/middleware/`) -HTTP request lifecycle handlers and security controls belong in `internal/middleware/`: +### 2. Cross-Cutting Middleware (`pkg/middleware/`) +HTTP request lifecycle handlers and security controls belong in `pkg/middleware/`. Under `pkg/` rather than `internal/` because a service built on Tusk needs the same chain, and Go's `internal/` rule would make it unreachable from another module: -- **Rate Limiter**: `internal/middleware/ratelimit.go` (uses `pkg/cache` or Redis sliding window token bucket). -- **Tracing / OpenTelemetry Middleware**: `internal/middleware/tracing.go` (attaches trace context to incoming HTTP requests). -- **Metrics Middleware**: `internal/middleware/metrics.go` (promhttp / request latency metrics). +- **Rate Limiter**: `pkg/middleware/ratelimit.go` (uses `pkg/cache` or Redis sliding window token bucket). +- **Tracing / OpenTelemetry Middleware**: `pkg/middleware/tracing.go` (attaches trace context to incoming HTTP requests). +- **Metrics Middleware**: `pkg/middleware/metrics.go` (promhttp / request latency metrics). --- -### 3. Core Server Lifecycle & Platform (`internal/platform/`) -Third-party client initializers and server lifecycle management belong in `internal/platform/`: +### 3. Core Server Lifecycle (`pkg/app/`, `database/`) +Server assembly and third-party client initialization: -- **Redis Client Initialization**: `internal/platform/redis/redis.go` -- **Graceful Shutdown**: `internal/app/app.go` (trapping `SIGINT`/`SIGTERM` to gracefully stop HTTP server, queue workers, and close DB/Redis pools). -- **OpenTelemetry Provider Init**: `internal/platform/telemetry/tracer.go` +- **Application shell**: `pkg/app/app.go` — router, middleware chain, health endpoints, Huma configuration, graceful shutdown (`SIGINT`/`SIGTERM`). It registers **no routes of its own**; callers register modules against `App.API()`. +- **Database connection**: `database/connect.go`, beside the migration runner. +- **Redis Client Initialization**: `pkg/redis/redis.go` +- **OpenTelemetry Provider Init**: `pkg/telemetry/tracer.go` --- diff --git a/docs/authorization.md b/docs/authorization.md index 9cb3e50..f756f1f 100644 --- a/docs/authorization.md +++ b/docs/authorization.md @@ -21,7 +21,7 @@ Permissions are defined in domain packages (e.g. `internal/auth/permissions.go`) ```go package auth -import "github.com/codetheuri/tusk/pkg/authz" +import "github.com/codetheuri/tusk/v2/pkg/authz" const ( PermUsersRead = "users.read" diff --git a/docs/database-and-migrations.md b/docs/database-and-migrations.md index 41957de..b13d069 100644 --- a/docs/database-and-migrations.md +++ b/docs/database-and-migrations.md @@ -6,7 +6,7 @@ Tusk uses **[GORM](https://gorm.io/)** for ORM operations and database connectiv ## Database Connection Management -Database connections are initialized in `internal/platform/database/gorm.go`. Connection settings (driver, host, port, pool sizes) are loaded directly from `.env`. +Database connections are initialized in `database/connect.go`. Connection settings (driver, host, port, pool sizes) are loaded directly from `.env`. --- diff --git a/docs/roadmap.md b/docs/roadmap.md new file mode 100644 index 0000000..ba6d04b --- /dev/null +++ b/docs/roadmap.md @@ -0,0 +1,193 @@ +# Tusk Roadmap & Known Gaps + +**Last reviewed:** 2026-09-07 +**Status:** Phase 1 hardening complete. Phase 2 (UUIDv7 keys) complete. + +This document tracks known production gaps and planned work. Items marked ✅ are resolved; treat anything else as unhandled. + +--- + +## 1. Production hardening + +Ordered by risk. + +### High — all resolved + +| # | Gap | Location | +|---|---|---| +| H1 ✅ | **Rate limiter is implemented but never installed.** `middleware/ratelimit.go` is complete; `app.New()` never calls it. Every endpoint currently accepts unlimited requests from a single IP — including login, which makes password brute-forcing free. | `internal/app/app.go` | +| H2 ✅ | **A database query runs on every authenticated request.** `HumaAuthenticate` issues `SELECT is_super_user FROM users WHERE id = ?` purely to read one boolean. Move the flag into the JWT so it arrives with the request. Accept that a revoked super-user keeps the flag until the token expires — which is precisely why fine-grained *permissions* must stay in the database rather than the token. | `internal/middleware/jwt.go` | +| H3 ✅ | **Invalid tokens are treated as absent tokens.** An expired or malformed token currently falls through unauthenticated and fails later at the permission guard, so the client receives `403 Forbidden` when the truth is `401 Unauthorized`. Clients cannot distinguish "you lack permission" from "your session expired", so they never know to refresh. | `internal/middleware/jwt.go` | +| H4 ✅ | **Raw-string context keys.** `context.WithValue(ctx, "user_id", …)` is set alongside the typed `ContextKeyUserID`, and `pkg/authz` reads the raw-string one — so the unsafe key is the load-bearing one. Any dependency using the same string silently overwrites it, at runtime, with no compile error. Contradicts `.agents/rules/gostyle.md`. | `internal/middleware/jwt.go`, `pkg/authz/middleware.go` | +| H5 ✅ | **No request body size limit.** Unbounded request bodies are a trivial memory-exhaustion vector. | `internal/app/app.go` | +| H6 ✅ | **Migrations are MySQL-only.** `00001_identity_auth_rbac.mysql.sql` uses MySQL syntax and will fail against PostgreSQL, which is now the supported database. | `database/migrations/` | + +### Medium — all resolved + +| # | Gap | Location | +|---|---|---| +| M1 ✅ | **Config performs no validation.** No minimum length enforced on `JWT_SECRET`, connection-pool sizes hardcoded rather than environment-driven, dead commented-out assignments, and no `IsProduction()` helper. Configuration errors should fail loudly at startup, not surface as confusing runtime behaviour. | `config/config.go` | +| M2 ✅ | **Server timeouts are too tight for large payloads.** `ReadTimeout: 5s` / `WriteTimeout: 10s`, with no `ReadHeaderTimeout` or `MaxHeaderBytes`. Any endpoint accepting a substantial upload over a slow connection will be cut off mid-request. | `internal/app/app.go` | +| M3 ✅ | **API documentation is always public.** `/docs` and `/openapi.json` are served unconditionally, including in production. See §2. | `internal/app/app.go` | +| M4 ✅ | **Two parallel authentication middlewares.** `Authenticate` (net/http) and `HumaAuthenticate` (Huma) duplicate JWT parsing and will drift apart. Consolidate on one implementation. | `internal/middleware/jwt.go` | +| M5 ✅ | **Logging is unstructured.** A custom console logger with no JSON output and no request or trace correlation. Migrate to `log/slog` with a JSON handler in production and `request_id` on every line. | `pkg/logger/` | +| M6 ✅ | **No database integration tests.** Config, JWT parsing, and logging now have unit coverage; `pkg/*` has unit coverage; the auth module and the database layer have none. Tests should run against a real throwaway PostgreSQL — not SQLite, which diverges on exactly the behaviour worth testing. | repo-wide | + +### Low — resolved + +| # | Gap | Location | +|---|---|---| +| L1 ✅ | Health endpoints hand-format JSON with `fmt.Sprintf` rather than encoding it, and there is no `/live` distinct from `/ready`. | `internal/app/app.go` | + +**Already correct, no action needed:** `.env` is gitignored and no secrets are tracked in git. + +--- + +## 2. Gating API documentation in production + +Endpoints stay live; only the documentation surface disappears. + +Huma supports this natively — setting `Config.DocsPath` and `Config.OpenAPIPath` to `""` removes those routes entirely rather than merely hiding them. + +Planned design: a `DOCS_ENABLED` environment variable defaulting to **off** when `APP_MODE=production`. An optional later refinement is to serve documentation behind a super-admin session, so it remains reachable in production for those authorized to see it. + +--- + +## 3. Planned capabilities + +### 3.1 UUIDv7 primary keys ✅ *(shipped 2026-09-08)* + +Move framework tables from `uint` auto-increment to UUIDv7 (`uuid.NewV7()`, already available via `google/uuid` v1.6.0). + +**Rationale.** Auto-increment keys cannot be generated outside the database, which rules out offline-capable clients, dataset merges, and multi-writer topologies. The cost of choosing UUID where `BIGINT` would have sufficed is 8 extra bytes per key; the cost of choosing `BIGINT` where UUID was needed is rewriting every table and foreign key in the system. UUIDv7 specifically — rather than v4 — because its leading millisecond timestamp keeps new keys sorting to the end of the B-tree, preserving insert locality that random v4 keys destroy. + +**Known tradeoff.** UUIDv7 embeds a creation timestamp, so an identifier reveals when its record was created and any two identifiers reveal their relative ordering. For ordinary business records this is unremarkable. Applications where creation-time correlation is sensitive should expose a separate random public identifier and keep v7 internal. + +An application with a measured reason to use `BIGINT` for its own high-volume tables can simply declare them that way — GORM is indifferent, and no generic `Repository[T, ID]` machinery is needed. + +### 3.2 Optional multi-tenancy (`pkg/tenant`) ✅ *(shipped 2026-09-08)* + +Shared schema, tenant column, automatic query scoping, with PostgreSQL row-level security as a safety net beneath the application's own filtering. + +**Tenancy must be opt-in per model.** Most applications are single-tenant, and a framework that imposes a tenant column on every table is unusable for them. A model opts in by implementing a single method: + +```go +func (Customer) TenantColumn() string { return "business_id" } +``` + +The query-scoping callback skips every model that does not implement it. Returning the *column name* rather than assuming `tenant_id` lets each application use its own domain vocabulary. Row-level security policies are written only for opted-in tables, and the tenant middleware stays inert until an application registers a resolver. + +Framework tables: `users` becomes tenant-scoped; `roles` and `permissions` stay global, since role names generally mean the same thing across tenants. Per-model opt-in means `roles` can join later without disturbing anything else. + +**As shipped.** Three layers, each independently sufficient: + +| Layer | Mechanism | Behaviour when it is the only one left | +|---|---|---| +| Context | `tenant.Middleware` + a `Resolver` reading a verified claim | Handlers cannot read a tenant from anywhere a caller controls | +| Query scoping | GORM callbacks on create, query, row, update, delete | A tenanted model queried with no tenant errors rather than returning everything | +| Row-level security | `PolicyStatements` DDL + `tenant.Transaction` setting `app.tenant_id` | The database refuses foreign rows even when the application asks for them | + +Callbacks are registered by `database.Connect` for every application, not opted into per project — a model added later cannot go unscoped because someone forgot a line of wiring. `tenant.Unscoped(db)` suspends layer 2 deliberately and greppably. + +Two failure modes worth knowing about, both covered by tests: + +- **An `OR` in a caller's own conditions can swallow the tenant predicate.** `Where(a).Or(b)` with a naively appended `AND` parses as `a OR (b AND tenant)`, so everything matching `a` leaks. The caller's expressions are grouped before the predicate is added. +- **Row-level security does not apply to superusers or roles holding `BYPASSRLS`**, silently. Policies install, appear in `\d`, and are never consulted. `tenant.VerifyEnforcement(db)` reports this; call it at startup in any deployment relying on the RLS layer. Development databases very often run as a superuser, which is exactly when someone concludes the layer works. + +### 3.3 Server-rendered pages (`pkg/view`, `pkg/session`) ✅ *(shipped 2026-09-08)* + +Tusk is API-first, but admin panels are a recurring need and previously had no supported answer. + +- `pkg/view` — `html/template` with layout inheritance, hot reload in development, caching in production. +- `pkg/session` — database-backed cookie sessions, HTTP-only + Secure + SameSite, rotated on privilege change. + +This lets a Huma API and an HTML admin panel coexist in one binary without either becoming an afterthought. + +**As shipped.** Both packages exist to make three recurring mistakes impossible rather than merely discouraged. + +`pkg/view` renders into a buffer and writes nothing until the whole page succeeds. `html/template` streams output as it executes, so rendering straight to the `ResponseWriter` means an error halfway down a page has already sent `200` and half the document — and the handler can no longer choose a status. It also parses every page at construction, so a malformed template stops the process starting instead of surfacing as a 500 to whoever visits that page first. Each page gets its own template set, so two pages may define `"content"` without colliding. + +`pkg/session` stores only the SHA-256 digest of a 256-bit token, never the token, so a leaked backup does not hand over live sessions. SHA-256 rather than bcrypt is deliberate and the opposite reasoning to passwords: the token is unguessable regardless of hash speed, and it is verified on every request, so a slow hash would be a denial-of-service surface. `Middleware` loads a session; `Require` refuses without one. Keeping those separate is what prevents the classic redirect loop, where a login page that redirects on the mere *presence* of a cookie bounces a visitor holding a stale one between `/login` and `/dashboard` forever. + +`Options.Kind` separates audiences sharing one table, so an operator-console cookie can never resolve against an ordinary user session. `Rotate` inherits the absolute expiry rather than restarting it — rotation is a security measure, not a way to hold a session open indefinitely. + +--- + +## 4. Changelog + +### 2026-09-08 — pkg/query: LIKE was case-sensitive on Postgres + +Found while building Salio's customer search. `Apply`'s multi-column search used `LIKE`, which is case-sensitive on PostgreSQL — the only database Tusk targets (D6). A search for "wanjiku" silently missed "Wanjiku". Fixed to `ILIKE`. `internal/auth`'s user search (`AllowedSearches: []string{"users.username", "users.email", "users.phone"}`) had the identical defect and is fixed by the same change. + +`pkg/query` had no test coverage before this — added `pkg/query/query_test.go` (5 tests, real PostgreSQL): case-insensitive search, an unlisted-column search is ignored, pagination metadata across pages, `MaxPerPage` enforcement, and sort-column whitelisting. + +### 2026-09-08 — Layout: what a downstream service can import + +Tusk is meant to be built on, and Go's `internal/` rule makes a package unreachable from any other module. The pieces a service actually needs moved out: + +| Was | Now | +|---|---| +| `internal/middleware` | `pkg/middleware` | +| `internal/app` | `pkg/app` | +| `internal/platform/database/gorm.go` | `database/connect.go` (beside the migration runner; `NewGoRMDB` → `Connect`) | +| `internal/auth` | unchanged — still internal | + +`pkg/app.New` now takes an `Options` struct and **registers no routes**. It returns an `App` exposing `API()`, `Router()`, `DB()` and `Config()`; callers register their own modules. Tusk's own auth module is registered from `cmd/api/main.go` like any other consumer — it is the framework's first user, not a privileged part of it. + +`internal/auth` stays internal deliberately. A service wanting different authentication — phone-only login, no email — writes its own module against `pkg/authz` rather than bending Tusk's. + +Paths in the §2 defect table are left as they were when those defects were fixed; git holds the history. + + +### 2026-09-08 — Phase 4: server-rendered pages + +- New `pkg/view` (9 tests) and `pkg/session` (13 tests, real PostgreSQL), plus migration `00002_sessions`. +- New `database.ResetMigrations` / `make migrate-reset` / `migrate reset`, unwinding every migration rather than one step. +- **The migration round-trip test was asserting less than it claimed.** It called `RollbackMigration`, which reverts a single step; with only one migration in the tree that happened to drop everything, so the test appeared to verify the whole rollback path. Adding a second migration exposed it. It now resets fully. +- **`lib/pq` is gone.** `config`, `cmd/migrate` and the tests all registered it while `gorm.io/driver/postgres` used pgx, so the schema was being built by one driver and used by another. Everything is pgx now, and lib/pq is no longer a dependency. + +### 2026-09-08 — Phase 3: optional multi-tenancy + +- New `pkg/tenant`: `Tenanted` opt-in interface, context propagation, GORM scoping callbacks, `Unscoped` escape hatch, RLS policy DDL, `Transaction`, `VerifyEnforcement`, and HTTP + Huma middleware. +- `Claims` gains an omitempty `tenant_id`; `middleware.TenantResolver` feeds it to the middleware. +- `database.Connect` and `internal/testdb` both register the callbacks, so tests exercise the production configuration. +- **`internal/testdb` now connects through pgx rather than lib/pq.** `gorm.io/driver/postgres` opens production connections with `stdlib.OpenDB`, so the harness had been testing a driver the application never uses — and they differ: GORM's migrator introspection reuses a placeholder, which pgx accepts and lib/pq rejects outright. +- **Database tests now take a PostgreSQL advisory lock.** `go test ./...` runs packages concurrently and every database test TRUNCATEs the whole schema; with a second database-using package this produced failures that looked like application bugs and reproduced only sometimes. The lock lives in the harness rather than in a `-p 1` Makefile flag, because a plain `go test ./...` has to be correct too. + +### 2026-09-07 — Phase 1 hardening + +Resolved H1–H6, M1–M5 and L1. Notable behaviour changes for anyone upgrading: + +- **`JWT_SECRET` must now be at least 32 characters.** Startup fails otherwise. This is a hard break for installations running a short key — rotate it before deploying. +- **`/docs` and `/openapi.json` no longer appear when `APP_MODE=production`.** Set `DOCS_ENABLED=true` to restore them. API endpoints are unchanged. +- **An invalid or expired token now returns 401, not 403.** Clients that treated 403 as "session expired" should be checked. +- **`is_super_user` is read from the JWT** rather than the database. Existing tokens lack the claim and will be treated as non-super until reissued — it fails closed, so no privilege is granted by accident. +- **Migrations moved to `database/migrations//`.** PostgreSQL is primary and receives all new work; the MySQL migration is preserved but frozen. +- **`DB_SSLMODE` is new**, defaulting to `require` in production. Previously the DSN hardcoded `sslmode=disable`, so production database traffic was unencrypted with no way to change it. +- Rate limiting and a request body cap are now actually installed; both are configurable. +- `swaggo/swag` and its dependency tree were dropped — leftovers from before Huma. + +**Integration testing (M6).** Tests now run against a real PostgreSQL via `internal/testdb`: + +- `make test` — everything; integration tests **skip** if no database is reachable. +- `make test-integration` — the same, but a missing database is a **failure**. +- `make test-db-setup` — creates the local `tusk_test` database once. +- Override the target with `TEST_DATABASE_URL`; default is `127.0.0.1:5434/tusk_test`. + +CI sets `REQUIRE_DB_TESTS=1`, so a green build can never mean the integration suite was silently skipped. The workflow now also reads its Go version from `go.mod` — it was pinned to 1.24.x while the module required 1.25.7 — and runs `go vet`. + +Covered: registration, login by username or email, account lockout and reset, refresh-token rotation and replay rejection, logout revocation, unique-constraint enforcement, RBAC evaluation across the three-table join, super-user bypass, permission-sync idempotency, and a migration up/down/up round trip in its own throwaway database. + +### 2026-09-08 — UUIDv7 primary keys + +Framework tables moved from `uint` auto-increment to UUIDv7. Identifiers are generated in Go by the new `pkg/id`, never by the database. + +**Breaking changes:** + +- **PostgreSQL is now the only supported driver.** MySQL and SQLite were dropped, and `LoadConfig` refuses anything else at startup rather than failing at the first query. Neither has a native UUID type, so their schema could no longer describe the same rows as the models. The MySQL migration was deleted; the SQL remains in git history. +- **Every primary and foreign key is now `UUID`.** Migration `00001` was rewritten rather than layered with a conversion — there was no production data to preserve. An existing database must be recreated. +- `Claims.UserID`, `authz.Subject.UserID`, and every repository, service and DTO signature now take `uuid.UUID`. Tokens issued before this change carry a numeric `user_id` and will fail to parse, so all sessions are invalidated. + +**No server-side default.** The migration deliberately omits `DEFAULT`. PostgreSQL gained `uuidv7()` only in version 18, and `gen_random_uuid()` would yield v4 — losing the index locality that motivated v7. More fundamentally, an offline client must be able to create a row and know its ID before the server has seen it, which no server-side default can serve. + +**Client-supplied identifiers are preserved.** Each model's `BeforeCreate` assigns an ID only when one is absent. A device that generated its own ID while disconnected keeps it, because other records it created may already reference it. Both directions are covered by tests. diff --git a/go.mod b/go.mod index e567b15..e00306b 100644 --- a/go.mod +++ b/go.mod @@ -1,11 +1,11 @@ -module github.com/codetheuri/tusk +module github.com/codetheuri/tusk/v2 go 1.25.7 require ( github.com/danielgtaylor/huma/v2 v2.39.0 - github.com/go-sql-driver/mysql v1.10.0 github.com/google/uuid v1.6.0 + github.com/jackc/pgx/v5 v5.10.0 github.com/joho/godotenv v1.5.1 github.com/pressly/goose/v3 v3.27.2 ) @@ -19,23 +19,17 @@ require ( require ( github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect - github.com/jackc/pgx/v5 v5.10.0 // indirect github.com/jackc/puddle/v2 v2.2.2 // indirect github.com/jinzhu/inflection v1.0.0 // indirect github.com/jinzhu/now v1.1.5 // indirect - github.com/mattn/go-sqlite3 v1.14.22 // indirect golang.org/x/sync v0.21.0 // indirect ) require ( - filippo.io/edwards25519 v1.2.0 // indirect github.com/go-chi/chi/v5 v5.3.0 github.com/golang-jwt/jwt/v5 v5.3.1 - github.com/lib/pq v1.12.3 golang.org/x/crypto v0.52.0 golang.org/x/text v0.37.0 // indirect - gorm.io/driver/mysql v1.6.0 gorm.io/driver/postgres v1.6.0 - gorm.io/driver/sqlite v1.6.0 gorm.io/gorm v1.31.2 ) diff --git a/go.sum b/go.sum old mode 100644 new mode 100755 index f0ef1ac..7e0c909 --- a/go.sum +++ b/go.sum @@ -1,5 +1,3 @@ -filippo.io/edwards25519 v1.2.0 h1:crnVqOiS4jqYleHd9vaKZ+HKtHfllngJIiOpNpoJsjo= -filippo.io/edwards25519 v1.2.0/go.mod h1:xzAOLCNug/yB62zG1bQ8uziwrIqIuxhctzJT18Q77mc= github.com/danielgtaylor/huma/v2 v2.39.0 h1:YiXbzhJBSeQVkKbhn8adZR48Ei4XFx/K6jShQ3O92qU= github.com/danielgtaylor/huma/v2 v2.39.0/go.mod h1:pGstQdMhQnP9ZBnrqPRb9goqOWs1HU1uQewKWmkJOAY= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= @@ -11,8 +9,6 @@ github.com/fxamacker/cbor/v2 v2.9.2 h1:X4Ksno9+x3cz0TZv69ec1hxP/+tymuR8PXQJyDwfh github.com/fxamacker/cbor/v2 v2.9.2/go.mod h1:vM4b+DJCtHn+zz7h3FFp/hDAI9WNWCsZj23V5ytsSxQ= github.com/go-chi/chi/v5 v5.3.0 h1:halUjDxhshgXHMrao5bB8eNBXo/rnzwr8m5m36glehM= github.com/go-chi/chi/v5 v5.3.0/go.mod h1:R+tYY2hNuVUUjxoPtqUdgBqevM9s9njzkTLutVsOCto= -github.com/go-sql-driver/mysql v1.10.0 h1:Q+1LV8DkHJvSYAdR83XzuhDaTykuDx0l6fkXxoWCWfw= -github.com/go-sql-driver/mysql v1.10.0/go.mod h1:M+cqaI7+xxXGG9swrdeUIoPG3Y3KCkF0pZej+SK+nWk= 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/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= @@ -31,8 +27,6 @@ github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ= github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8= github.com/joho/godotenv v1.5.1 h1:7eLL/+HRGLY0ldzfGMeQkb7vMd0as4CfYvUVzLqw0N0= github.com/joho/godotenv v1.5.1/go.mod h1:f4LDr5Voq0i2e/R5DDNOoa2zzDfwtkZa6DnEwAbqwq4= -github.com/lib/pq v1.12.3 h1:tTWxr2YLKwIvK90ZXEw8GP7UFHtcbTtty8zsI+YjrfQ= -github.com/lib/pq v1.12.3/go.mod h1:/p+8NSbOcwzAEI7wiMXFlgydTwcgTr3OSKMsD2BitpA= github.com/mattn/go-isatty v0.0.22 h1:j8l17JJ9i6VGPUFUYoTUKPSgKe/83EYU2zBC7YNKMw4= github.com/mattn/go-isatty v0.0.22/go.mod h1:ZXfXG4SQHsB/w3ZeOYbR0PrPwLy+n6xiMrJlRFqopa4= github.com/mattn/go-sqlite3 v1.14.22 h1:2gZY6PC6kBnID23Tichd1K+Z0oS6nE/XwU+Vz/5o4kU= @@ -70,8 +64,6 @@ gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8 gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= -gorm.io/driver/mysql v1.6.0 h1:eNbLmNTpPpTOVZi8MMxCi2aaIm0ZpInbORNXDwyLGvg= -gorm.io/driver/mysql v1.6.0/go.mod h1:D/oCC2GWK3M/dqoLxnOlaNKmXz8WNTfcS9y5ovaSqKo= gorm.io/driver/postgres v1.6.0 h1:2dxzU8xJ+ivvqTRph34QX+WrRaJlmfyPqXmoGVjMBa4= gorm.io/driver/postgres v1.6.0/go.mod h1:vUw0mrGgrTK+uPHEhAdV4sfFELrByKVGnaVRkXDhtWo= gorm.io/driver/sqlite v1.6.0 h1:WHRRrIiulaPiPFmDcod6prc4l2VGVWHz80KspNsxSfQ= diff --git a/internal/app/app.go b/internal/app/app.go deleted file mode 100644 index 58e24ca..0000000 --- a/internal/app/app.go +++ /dev/null @@ -1,177 +0,0 @@ -package app - -import ( - "context" - "fmt" - "net" - "net/http" - "os" - "os/signal" - "syscall" - "time" - - "github.com/danielgtaylor/huma/v2" - "github.com/danielgtaylor/huma/v2/adapters/humachi" - "github.com/go-chi/chi/v5" - - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/internal/auth" - "github.com/codetheuri/tusk/internal/middleware" - - appDatabase "github.com/codetheuri/tusk/internal/platform/database" - - "github.com/codetheuri/tusk/pkg/logger" - "github.com/codetheuri/tusk/pkg/response" -) - -type App struct { - cfg *config.Config - router *chi.Mux - log logger.Logger -} - -func New(cfg *config.Config, log logger.Logger) (*App, error) { - db, err := appDatabase.NewGoRMDB(cfg, log) - if err != nil { - return nil, fmt.Errorf("database connection failed: %w", err) - } - - r := chi.NewRouter() - - // Middlewares - r.Use(middleware.RequestID()) - r.Use(middleware.Logger(log)) - r.Use(middleware.Recovery(log)) - r.Use(middleware.CORS(cfg.CORSOrigins, log)) - r.Use(middleware.SecurityHeaders) - - r.Get("/health", func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusOK) - w.Write([]byte(fmt.Sprintf(`{"status":"ok","app_name":"%s","version":"%s","time":"%s"}`, cfg.AppName, cfg.AppVersion, time.Now().Format(time.RFC3339)))) - }) - - r.Get("/ready", func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "application/json") - sqlDB, err := db.DB() - if err != nil || sqlDB.PingContext(r.Context()) != nil { - w.WriteHeader(http.StatusServiceUnavailable) - w.Write([]byte(`{"status":"unavailable","database":"disconnected"}`)) - return - } - w.WriteHeader(http.StatusOK) - w.Write([]byte(`{"status":"ready","database":"connected"}`)) - }) - - // Initialize Huma custom formatting - response.SetupHuma() - - humaConfig := huma.DefaultConfig("Tusk Backend API", "1.0.0") - humaConfig.Info.Description = "## Official Tusk Enterprise Backend API Documentation\n\nWelcome to the developer documentation for Tusk. Explore Identity, Authentication, and RBAC endpoints below." - - humaConfig.Tags = []*huma.Tag{ - {Name: "Authentication", Description: "User registration, login, token refresh, logout, and self profile operations"}, - {Name: "Users", Description: "User account management and listing"}, - {Name: "Roles", Description: "Security role management (CRUD)"}, - {Name: "Role Permissions", Description: "Attaching and detaching permission strings to/from roles"}, - {Name: "User Roles", Description: "Assigning and revoking security roles to/from users"}, - {Name: "Permissions", Description: "System permission catalog listing"}, - } - - humaConfig.Info.Contact = &huma.Contact{ - Name: "API Support", - Email: "theurij113@gmail.com", - } - - // Add Bearer JWT Security Scheme to OpenAPI docs - humaConfig.Components = &huma.Components{ - SecuritySchemes: map[string]*huma.SecurityScheme{ - "bearerAuth": { - Type: "http", - Scheme: "bearer", - BearerFormat: "JWT", - Description: "Enter token as: Bearer ", - }, - }, - } - - // Disable $schema from showing up in OpenAPI docs/responses - humaConfig.CreateHooks = []func(huma.Config) huma.Config{ - func(c huma.Config) huma.Config { - c.SchemasPath = "" - return c - }, - } - api := humachi.New(r, humaConfig) - - // Configure OpenAPI x-tagGroups extension for Scalar / Redoc UI sidebar grouping under IAM - openAPI := api.OpenAPI() - if openAPI.Extensions == nil { - openAPI.Extensions = map[string]any{} - } - openAPI.Extensions["x-tagGroups"] = []map[string]any{ - { - "name": "IAM", - "tags": []string{ - "Authentication", - "Users", - "Roles", - "Role Permissions", - "User Roles", - "Permissions", - }, - }, - } - - // Global Huma JWT Authentication Middleware - api.UseMiddleware(middleware.HumaAuthenticate(api, cfg.JWTSecret, db)) - - // Routes - auth.RegisterRoutes(api, db, cfg, log) - - return &App{ - cfg: cfg, - router: r, - log: log, - }, nil -} - -func (a *App) Run() error { - srv := &http.Server{ - Addr: fmt.Sprintf(":%d", a.cfg.ServerPort), - Handler: a.router, - ReadTimeout: 5 * time.Second, - WriteTimeout: 10 * time.Second, - IdleTimeout: 60 * time.Second, - } - - ln, err := net.Listen("tcp", srv.Addr) - if err != nil { - return fmt.Errorf("failed to start listener: %w", err) - } - - actualAddr := ln.Addr().(*net.TCPAddr) - a.log.Info(fmt.Sprintf("Server is listening on port %d. Docs at http://localhost:%d/docs", actualAddr.Port, actualAddr.Port)) - - go func() { - if err := srv.Serve(ln); err != nil && err != http.ErrServerClosed { - a.log.Fatal("Server failed to listen or serve", err) - } - }() - - quit := make(chan os.Signal, 1) - signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM) - sig := <-quit - a.log.Warn("Received shutdown signal", "signal", sig.String()) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - a.log.Info("Attempting to shut down gracefully...") - if err := srv.Shutdown(ctx); err != nil { - return fmt.Errorf("server shutdown failed: %w", err) - } - - a.log.Info("Server shut down gracefully.") - return nil -} diff --git a/internal/auth/dto.go b/internal/auth/dto.go index 50290d7..e79798e 100644 --- a/internal/auth/dto.go +++ b/internal/auth/dto.go @@ -1,9 +1,11 @@ package auth import ( - "github.com/codetheuri/tusk/pkg/authz" - "github.com/codetheuri/tusk/pkg/query" - "github.com/codetheuri/tusk/pkg/response" + "github.com/google/uuid" + + "github.com/codetheuri/tusk/v2/pkg/authz" + "github.com/codetheuri/tusk/v2/pkg/query" + "github.com/codetheuri/tusk/v2/pkg/response" ) // ------------------------------------------------------------- @@ -108,7 +110,7 @@ type CreateRoleInput struct { } type UpdateRoleInput struct { - ID uint `path:"id" doc:"Role ID"` + ID uuid.UUID `path:"id" doc:"Role ID"` Body struct { Name *string `json:"name,omitempty" minLength:"2" doc:"Updated name of the role"` Description *string `json:"description,omitempty" doc:"Updated description"` @@ -116,7 +118,7 @@ type UpdateRoleInput struct { } type RoleIDInput struct { - ID uint `path:"id" doc:"Role ID"` + ID uuid.UUID `path:"id" doc:"Role ID"` } type RoleData struct { @@ -138,27 +140,27 @@ type RolesOutput struct { } type AddRolePermissionInput struct { - ID uint `path:"id" doc:"Role ID"` + ID uuid.UUID `path:"id" doc:"Role ID"` Body struct { PermissionName string `json:"permission_name" doc:"Permission string to attach (e.g., users.read)"` } } type RemoveRolePermissionInput struct { - ID uint `path:"id" doc:"Role ID"` - PermissionName string `path:"permission_name" doc:"Permission string to remove"` + ID uuid.UUID `path:"id" doc:"Role ID"` + PermissionName string `path:"permission_name" doc:"Permission string to remove"` } type AssignUserRoleInput struct { - UserID uint `path:"user_id" doc:"User ID"` + UserID uuid.UUID `path:"user_id" doc:"User ID"` Body struct { - RoleID uint `json:"role_id" doc:"Role ID to assign"` + RoleID uuid.UUID `json:"role_id" doc:"Role ID to assign"` } } type RemoveUserRoleInput struct { - UserID uint `path:"user_id" doc:"User ID"` - RoleID uint `path:"role_id" doc:"Role ID to revoke"` + UserID uuid.UUID `path:"user_id" doc:"User ID"` + RoleID uuid.UUID `path:"role_id" doc:"Role ID to revoke"` } type MessageOutput struct { diff --git a/internal/auth/handler_auth.go b/internal/auth/handler_auth.go index 4b32192..bb5fe6d 100644 --- a/internal/auth/handler_auth.go +++ b/internal/auth/handler_auth.go @@ -6,9 +6,10 @@ import ( "github.com/danielgtaylor/huma/v2" - "github.com/codetheuri/tusk/pkg/logger" - "github.com/codetheuri/tusk/pkg/query" - "github.com/codetheuri/tusk/pkg/response" + "github.com/codetheuri/tusk/v2/pkg/authz" + "github.com/codetheuri/tusk/v2/pkg/logger" + "github.com/codetheuri/tusk/v2/pkg/query" + "github.com/codetheuri/tusk/v2/pkg/response" ) type Handler struct { @@ -98,10 +99,11 @@ func (h *Handler) Logout(ctx context.Context, input *RefreshTokenInput) (*Messag // Me returns the profile and permission list of the authenticated user. func (h *Handler) Me(ctx context.Context, input *MeInput) (*ProfileOutput, error) { - userID, ok := ctx.Value("user_id").(uint) + sub, ok := authz.SubjectFromContext(ctx) if !ok { return nil, huma.Error401Unauthorized("Authentication required") } + userID := sub.UserID user, perms, err := h.service.GetCurrentUser(ctx, userID) if err != nil { @@ -118,10 +120,11 @@ func (h *Handler) Me(ctx context.Context, input *MeInput) (*ProfileOutput, error // UpdateProfile updates the personal profile of the authenticated user. func (h *Handler) UpdateProfile(ctx context.Context, input *UpdateProfileInput) (*ProfileOutput, error) { - userID, ok := ctx.Value("user_id").(uint) + sub, ok := authz.SubjectFromContext(ctx) if !ok { return nil, huma.Error401Unauthorized("Authentication required") } + userID := sub.UserID _, err := h.service.UpdateProfile(ctx, userID, input.Body.FirstName, input.Body.LastName, input.Body.Avatar, input.Body.Bio) if err != nil { diff --git a/internal/auth/handler_role.go b/internal/auth/handler_role.go index 77019b0..8b9b2d6 100644 --- a/internal/auth/handler_role.go +++ b/internal/auth/handler_role.go @@ -5,7 +5,7 @@ import ( "github.com/danielgtaylor/huma/v2" - "github.com/codetheuri/tusk/pkg/response" + "github.com/codetheuri/tusk/v2/pkg/response" ) // CreateRole handles creating a new RBAC role (name & description). diff --git a/internal/auth/integration_test.go b/internal/auth/integration_test.go new file mode 100644 index 0000000..cb8eb5b --- /dev/null +++ b/internal/auth/integration_test.go @@ -0,0 +1,425 @@ +package auth + +import ( + "context" + "testing" + "time" + + "gorm.io/gorm" + + "github.com/codetheuri/tusk/v2/config" + "github.com/codetheuri/tusk/v2/pkg/authz" + "github.com/codetheuri/tusk/v2/pkg/id" + "github.com/codetheuri/tusk/v2/pkg/logger" + "github.com/codetheuri/tusk/v2/pkg/testdb" +) + +// newTestService wires the real repository and service against a real database. +// Nothing is mocked: the point of these tests is the SQL, which a mock would +// replace with the very assumption under test. +func newTestService(t *testing.T) (*Service, *Repository, *gorm.DB) { + t.Helper() + + db := testdb.Connect(t) + if db == nil { + return nil, nil, nil // Connect already skipped the test + } + + cfg := &config.Config{ + JWTSecret: "integration-test-secret-key-long-enough", + AccessTokenTTL: time.Hour, + } + repo := NewRepository(db, logger.NewTextLogger("error")) + return NewService(repo, cfg), repo, db +} + +func mustRegister(t *testing.T, svc *Service, username, email, password string) *User { + t.Helper() + user, err := svc.Register(context.Background(), &RegisterRequest{ + Username: username, + Email: email, + Password: password, + PasswordConfirm: password, + FirstName: "Test", + LastName: "User", + }) + if err != nil { + t.Fatalf("register %q failed: %v", username, err) + } + return user +} + +func TestRegisterAndLogin(t *testing.T) { + svc, _, _ := newTestService(t) + if svc == nil { + return + } + ctx := context.Background() + + user := mustRegister(t, svc, "alice", "alice@example.com", "correct-horse-battery") + if id.IsZero(user.ID) { + t.Fatal("expected an identifier to be assigned") + } + + // The stored password must be a hash, never the plaintext. + if user.Password == "correct-horse-battery" { + t.Fatal("password was stored in plaintext") + } + + tokens, err := svc.Login(ctx, &LoginRequest{Login: "alice", Password: "correct-horse-battery"}) + if err != nil { + t.Fatalf("login failed: %v", err) + } + if tokens.AccessToken == "" || tokens.RefreshToken == "" { + t.Error("expected both an access and a refresh token") + } + + if _, err := svc.Login(ctx, &LoginRequest{Login: "alice", Password: "wrong"}); err == nil { + t.Error("expected login with a wrong password to fail") + } +} + +// Login accepts a username, an email, or a phone in one field. That flexibility +// is a query concern, so it can only be verified against a real database. +func TestLogin_AcceptsUsernameOrEmail(t *testing.T) { + svc, _, _ := newTestService(t) + if svc == nil { + return + } + ctx := context.Background() + + mustRegister(t, svc, "bob", "bob@example.com", "a-good-password") + + for _, login := range []string{"bob", "bob@example.com"} { + if _, err := svc.Login(ctx, &LoginRequest{Login: login, Password: "a-good-password"}); err != nil { + t.Errorf("login with %q failed: %v", login, err) + } + } +} + +// Five failed attempts must lock the account, and the lock must survive a +// subsequent attempt with the *correct* password — otherwise it is not a lock. +func TestLogin_LocksAccountAfterFailedAttempts(t *testing.T) { + svc, repo, _ := newTestService(t) + if svc == nil { + return + } + ctx := context.Background() + + user := mustRegister(t, svc, "carol", "carol@example.com", "the-real-password") + + for i := 0; i < 5; i++ { + if _, err := svc.Login(ctx, &LoginRequest{Login: "carol", Password: "wrong"}); err == nil { + t.Fatalf("attempt %d: expected failure with a wrong password", i+1) + } + } + + stored, err := repo.FindByID(ctx, user.ID) + if err != nil { + t.Fatalf("failed to reload user: %v", err) + } + if stored.FailedLoginAttempts < 5 { + t.Errorf("FailedLoginAttempts = %d, want >= 5", stored.FailedLoginAttempts) + } + if stored.LockedUntil == nil || !stored.LockedUntil.After(time.Now()) { + t.Fatalf("expected the account to be locked, LockedUntil = %v", stored.LockedUntil) + } + + if _, err := svc.Login(ctx, &LoginRequest{Login: "carol", Password: "the-real-password"}); err == nil { + t.Error("expected a locked account to reject even the correct password") + } +} + +func TestLogin_ResetsFailureCountOnSuccess(t *testing.T) { + svc, repo, _ := newTestService(t) + if svc == nil { + return + } + ctx := context.Background() + + user := mustRegister(t, svc, "dave", "dave@example.com", "the-real-password") + + for i := 0; i < 3; i++ { + _, _ = svc.Login(ctx, &LoginRequest{Login: "dave", Password: "wrong"}) + } + if _, err := svc.Login(ctx, &LoginRequest{Login: "dave", Password: "the-real-password"}); err != nil { + t.Fatalf("login with the correct password failed: %v", err) + } + + stored, err := repo.FindByID(ctx, user.ID) + if err != nil { + t.Fatalf("failed to reload user: %v", err) + } + if stored.FailedLoginAttempts != 0 { + t.Errorf("FailedLoginAttempts = %d after a successful login, want 0", stored.FailedLoginAttempts) + } + if stored.LastLoginAt == nil { + t.Error("expected LastLoginAt to be recorded") + } +} + +// Refresh tokens rotate: presenting one must invalidate it. Replaying a rotated +// token is the signal of a stolen credential, so it must not succeed. +func TestRefreshToken_RotatesAndRevokes(t *testing.T) { + svc, _, _ := newTestService(t) + if svc == nil { + return + } + ctx := context.Background() + + mustRegister(t, svc, "erin", "erin@example.com", "a-good-password") + tokens, err := svc.Login(ctx, &LoginRequest{Login: "erin", Password: "a-good-password"}) + if err != nil { + t.Fatalf("login failed: %v", err) + } + + rotated, err := svc.RefreshToken(ctx, tokens.RefreshToken) + if err != nil { + t.Fatalf("refresh failed: %v", err) + } + if rotated.RefreshToken == tokens.RefreshToken { + t.Error("expected a new refresh token, got the same one back") + } + + if _, err := svc.RefreshToken(ctx, tokens.RefreshToken); err == nil { + t.Error("expected the rotated-away refresh token to be rejected on reuse") + } +} + +func TestLogout_RevokesRefreshToken(t *testing.T) { + svc, _, _ := newTestService(t) + if svc == nil { + return + } + ctx := context.Background() + + mustRegister(t, svc, "frank", "frank@example.com", "a-good-password") + tokens, err := svc.Login(ctx, &LoginRequest{Login: "frank", Password: "a-good-password"}) + if err != nil { + t.Fatalf("login failed: %v", err) + } + + if err := svc.Logout(ctx, tokens.RefreshToken); err != nil { + t.Fatalf("logout failed: %v", err) + } + if _, err := svc.RefreshToken(ctx, tokens.RefreshToken); err == nil { + t.Error("expected a revoked refresh token to be unusable") + } +} + +// The unique constraints live in the migration, so only a real database proves +// they are actually enforced. +func TestRegister_RejectsDuplicates(t *testing.T) { + svc, _, _ := newTestService(t) + if svc == nil { + return + } + ctx := context.Background() + + mustRegister(t, svc, "grace", "grace@example.com", "a-good-password") + + _, err := svc.Register(ctx, &RegisterRequest{ + Username: "grace", Email: "different@example.com", + Password: "a-good-password", PasswordConfirm: "a-good-password", + }) + if err == nil { + t.Error("expected a duplicate username to be rejected") + } + + _, err = svc.Register(ctx, &RegisterRequest{ + Username: "different", Email: "grace@example.com", + Password: "a-good-password", PasswordConfirm: "a-good-password", + }) + if err == nil { + t.Error("expected a duplicate email to be rejected") + } +} + +// The RBAC check walks users → user_roles → role_permissions. That is a +// three-table join built at runtime, which no unit test can exercise. +func TestRBAC_PermissionEvaluation(t *testing.T) { + svc, _, db := newTestService(t) + if svc == nil { + return + } + ctx := context.Background() + + user := mustRegister(t, svc, "heidi", "heidi@example.com", "a-good-password") + + role, err := svc.CreateRole(ctx, &CreateRoleRequest{Name: "editor", Description: "May edit users"}) + if err != nil { + t.Fatalf("create role failed: %v", err) + } + + evaluator := authz.NewEvaluator(db) + subject := authz.Subject{UserID: user.ID} + + // Before anything is granted. + allowed, err := evaluator.IsAuthorized(ctx, subject, authz.RequirePermissionPolicy{Permission: PermUsersUpdate}) + if err != nil { + t.Fatalf("authorization check failed: %v", err) + } + if allowed { + t.Error("a user with no roles must not hold any permission") + } + + if err := svc.AddRolePermission(ctx, role.ID, PermUsersUpdate); err != nil { + t.Fatalf("attaching permission to role failed: %v", err) + } + + // The role holds the permission, but the user does not hold the role yet. + allowed, _ = evaluator.IsAuthorized(ctx, subject, authz.RequirePermissionPolicy{Permission: PermUsersUpdate}) + if allowed { + t.Error("permission leaked to a user who was never assigned the role") + } + + if err := svc.AssignUserRole(ctx, user.ID, role.ID); err != nil { + t.Fatalf("assigning role to user failed: %v", err) + } + + allowed, err = evaluator.IsAuthorized(ctx, subject, authz.RequirePermissionPolicy{Permission: PermUsersUpdate}) + if err != nil { + t.Fatalf("authorization check failed: %v", err) + } + if !allowed { + t.Error("expected the permission to be granted through the assigned role") + } + + // A permission the role was never given must still be refused. + allowed, _ = evaluator.IsAuthorized(ctx, subject, authz.RequirePermissionPolicy{Permission: PermUsersDelete}) + if allowed { + t.Error("expected an ungranted permission to be refused") + } + + // Revoking the role revokes the permission with it. + if err := svc.RemoveUserRole(ctx, user.ID, role.ID); err != nil { + t.Fatalf("removing role failed: %v", err) + } + allowed, _ = evaluator.IsAuthorized(ctx, subject, authz.RequirePermissionPolicy{Permission: PermUsersUpdate}) + if allowed { + t.Error("expected permission to be revoked along with the role") + } +} + +// A super user bypasses permission checks entirely — worth pinning down, because +// this flag now travels in the JWT rather than being read from the database. +func TestRBAC_SuperUserBypassesChecks(t *testing.T) { + svc, _, db := newTestService(t) + if svc == nil { + return + } + ctx := context.Background() + + user := mustRegister(t, svc, "ivan", "ivan@example.com", "a-good-password") + + evaluator := authz.NewEvaluator(db) + superSubject := authz.Subject{UserID: user.ID, IsSuperUser: true} + + allowed, err := evaluator.IsAuthorized(ctx, superSubject, authz.RequirePermissionPolicy{Permission: PermUsersDelete}) + if err != nil { + t.Fatalf("authorization check failed: %v", err) + } + if !allowed { + t.Error("expected a super user to bypass permission checks") + } +} + +// The synchroniser must be idempotent: running it twice is routine (every deploy) +// and must not duplicate rows or error. +func TestPermissionSync_IsIdempotent(t *testing.T) { + db := testdb.Connect(t) + if db == nil { + return + } + ctx := context.Background() + + sync := authz.NewSynchronizer(db, authz.DefaultRegistry()) + + first, err := sync.Sync(ctx, false) + if err != nil { + t.Fatalf("first sync failed: %v", err) + } + if first.Inserted == 0 { + t.Error("expected the first sync to insert the code-declared permissions") + } + + second, err := sync.Sync(ctx, false) + if err != nil { + t.Fatalf("second sync failed: %v", err) + } + if second.Inserted != 0 { + t.Errorf("second sync inserted %d permissions, want 0 — sync is not idempotent", second.Inserted) + } + + var count int64 + db.Table("permissions").Count(&count) + if int(count) != len(authz.DefaultRegistry().All()) { + t.Errorf("permissions table holds %d rows, want %d", count, len(authz.DefaultRegistry().All())) + } +} + +// A client-generated identifier must survive being written. +// +// This is the property that makes offline-first clients possible: a device +// creates a record while disconnected, mints its own ID, and may already have +// created other records referencing it. If BeforeCreate overwrote that ID on +// arrival, those references would silently point at nothing. +func TestCreate_PreservesClientSuppliedID(t *testing.T) { + _, repo, _ := newTestService(t) + if repo == nil { + return + } + ctx := context.Background() + + // The identifier a disconnected client would have generated for itself. + clientID := id.New() + + user := &User{ + ID: clientID, + Username: "offline-client", + Email: "offline@example.com", + Password: "already-hashed-by-the-service", + IsActive: true, + } + if err := repo.CreateUserWithProfile(ctx, user, &UserProfile{}); err != nil { + t.Fatalf("create failed: %v", err) + } + + if user.ID != clientID { + t.Fatalf("client-supplied ID was overwritten: got %s, want %s", user.ID, clientID) + } + + stored, err := repo.FindByID(ctx, clientID) + if err != nil { + t.Fatalf("could not read back the row by its client-supplied ID: %v", err) + } + if stored.Username != "offline-client" { + t.Errorf("read back the wrong row: %s", stored.Username) + } +} + +// The complementary case: with no ID supplied, one must be generated. +func TestCreate_GeneratesIDWhenAbsent(t *testing.T) { + _, repo, _ := newTestService(t) + if repo == nil { + return + } + ctx := context.Background() + + user := &User{ + Username: "no-id-supplied", + Email: "no-id@example.com", + Password: "already-hashed-by-the-service", + IsActive: true, + } + if err := repo.CreateUserWithProfile(ctx, user, &UserProfile{}); err != nil { + t.Fatalf("create failed: %v", err) + } + + if id.IsZero(user.ID) { + t.Fatal("expected an identifier to be generated") + } + if user.ID.Version() != 7 { + t.Errorf("generated ID is UUID version %d, want 7", user.ID.Version()) + } +} diff --git a/internal/auth/model.go b/internal/auth/model.go index 5fa568b..ee62960 100644 --- a/internal/auth/model.go +++ b/internal/auth/model.go @@ -2,11 +2,16 @@ package auth import ( "time" + + "github.com/google/uuid" + "gorm.io/gorm" + + "github.com/codetheuri/tusk/v2/pkg/id" ) // User handles core authentication data, credentials, and security state. type User struct { - ID uint `json:"id" gorm:"primaryKey"` + ID uuid.UUID `json:"id" gorm:"type:uuid;primaryKey"` Username string `json:"username" gorm:"uniqueIndex;not null"` Email string `json:"email" gorm:"uniqueIndex;not null"` Phone *string `json:"phone,omitempty" gorm:"uniqueIndex"` @@ -24,9 +29,22 @@ type User struct { Profile *UserProfile `json:"profile,omitempty" gorm:"foreignKey:UserID"` } +// BeforeCreate assigns an identifier when the caller has not supplied one. +// +// The zero check is the important part, not the assignment: a client that +// generated its own ID while offline must keep it, because that ID may already be +// referenced by other records it created before it could reach the server. +// Overwriting it here would break those references silently. +func (u *User) BeforeCreate(*gorm.DB) error { + if id.IsZero(u.ID) { + u.ID = id.New() + } + return nil +} + // UserProfile stores personal identity information. type UserProfile struct { - UserID uint `json:"user_id" gorm:"primaryKey"` + UserID uuid.UUID `json:"user_id" gorm:"type:uuid;primaryKey"` FirstName string `json:"first_name"` LastName string `json:"last_name"` Avatar string `json:"avatar"` @@ -37,17 +55,24 @@ type UserProfile struct { // RefreshToken stores hashed refresh tokens for session management and revocation. type RefreshToken struct { - ID uint `json:"id" gorm:"primaryKey"` - UserID uint `json:"user_id" gorm:"not null;index"` + ID uuid.UUID `json:"id" gorm:"type:uuid;primaryKey"` + UserID uuid.UUID `json:"user_id" gorm:"type:uuid;not null;index"` TokenHash string `json:"-" gorm:"uniqueIndex;not null"` ExpiresAt time.Time `json:"expires_at" gorm:"not null"` RevokedAt *time.Time `json:"revoked_at,omitempty"` CreatedAt time.Time `json:"created_at"` } +func (r *RefreshToken) BeforeCreate(*gorm.DB) error { + if id.IsZero(r.ID) { + r.ID = id.New() + } + return nil +} + // Role represents a security role containing permissions. type Role struct { - ID uint `json:"id" gorm:"primaryKey"` + ID uuid.UUID `json:"id" gorm:"type:uuid;primaryKey"` Name string `json:"name" gorm:"uniqueIndex;not null"` Description string `json:"description"` CreatedAt time.Time `json:"created_at"` @@ -55,16 +80,23 @@ type Role struct { Permissions []string `json:"permissions,omitempty" gorm:"-"` } +func (r *Role) BeforeCreate(*gorm.DB) error { + if id.IsZero(r.ID) { + r.ID = id.New() + } + return nil +} + // RolePermission defines the join table linking roles to permissions. type RolePermission struct { - RoleID uint `gorm:"primaryKey"` - PermissionName string `gorm:"primaryKey"` + RoleID uuid.UUID `gorm:"type:uuid;primaryKey"` + PermissionName string `gorm:"primaryKey"` CreatedAt time.Time } // UserRole defines the join table linking users to roles. type UserRole struct { - UserID uint `gorm:"primaryKey"` - RoleID uint `gorm:"primaryKey"` + UserID uuid.UUID `gorm:"type:uuid;primaryKey"` + RoleID uuid.UUID `gorm:"type:uuid;primaryKey"` CreatedAt time.Time } diff --git a/internal/auth/permissions.go b/internal/auth/permissions.go index 864bc6d..34a5c3f 100644 --- a/internal/auth/permissions.go +++ b/internal/auth/permissions.go @@ -1,20 +1,20 @@ package auth -import "github.com/codetheuri/tusk/pkg/authz" +import "github.com/codetheuri/tusk/v2/pkg/authz" // Auth module permission constants to prevent raw string typos across handlers. const ( - PermUsersRead = "users.read" - PermUsersCreate = "users.create" - PermUsersUpdate = "users.update" - PermUsersDelete = "users.delete" - PermRolesRead = "roles.read" - PermRolesCreate = "roles.create" - PermRolesUpdate = "roles.update" - PermRolesDelete = "roles.delete" + PermUsersRead = "users.read" + PermUsersCreate = "users.create" + PermUsersUpdate = "users.update" + PermUsersDelete = "users.delete" + PermRolesRead = "roles.read" + PermRolesCreate = "roles.create" + PermRolesUpdate = "roles.update" + PermRolesDelete = "roles.delete" PermRolePermissionsManage = "roles.permissions.manage" - PermUserRolesManage = "users.roles.manage" - PermPermissionsRead = "permissions.read" + PermUserRolesManage = "users.roles.manage" + PermPermissionsRead = "permissions.read" ) // Permissions exported by the auth module. diff --git a/internal/auth/repository_auth.go b/internal/auth/repository_auth.go index 8d19a19..03218da 100644 --- a/internal/auth/repository_auth.go +++ b/internal/auth/repository_auth.go @@ -1,11 +1,13 @@ package auth import ( + "github.com/google/uuid" + "context" "time" - "github.com/codetheuri/tusk/pkg/logger" - "github.com/codetheuri/tusk/pkg/query" + "github.com/codetheuri/tusk/v2/pkg/logger" + "github.com/codetheuri/tusk/v2/pkg/query" "gorm.io/gorm" ) @@ -50,7 +52,7 @@ func (r *Repository) FindByLogin(ctx context.Context, login string) (*User, erro } // FindByID fetches a user by primary key ID with Profile preloaded. -func (r *Repository) FindByID(ctx context.Context, id uint) (*User, error) { +func (r *Repository) FindByID(ctx context.Context, id uuid.UUID) (*User, error) { var user User if err := r.db.WithContext(ctx).Preload("Profile").First(&user, id).Error; err != nil { return nil, err diff --git a/internal/auth/repository_role.go b/internal/auth/repository_role.go index f55059b..a7f2eef 100644 --- a/internal/auth/repository_role.go +++ b/internal/auth/repository_role.go @@ -1,9 +1,11 @@ package auth import ( + "github.com/google/uuid" + "context" - "github.com/codetheuri/tusk/pkg/authz" + "github.com/codetheuri/tusk/v2/pkg/authz" "gorm.io/gorm" ) @@ -29,7 +31,7 @@ func (r *Repository) ListRoles(ctx context.Context) ([]Role, error) { } // GetRoleByID fetches a role by its ID along with its permissions. -func (r *Repository) GetRoleByID(ctx context.Context, id uint) (*Role, error) { +func (r *Repository) GetRoleByID(ctx context.Context, id uuid.UUID) (*Role, error) { var role Role if err := r.db.WithContext(ctx).First(&role, id).Error; err != nil { return nil, err @@ -48,7 +50,7 @@ func (r *Repository) UpdateRole(ctx context.Context, role *Role) error { } // DeleteRole deletes a role and cascades removal from join tables. -func (r *Repository) DeleteRole(ctx context.Context, id uint) error { +func (r *Repository) DeleteRole(ctx context.Context, id uuid.UUID) error { return r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { if err := tx.Where("role_id = ?", id).Delete(&RolePermission{}).Error; err != nil { return err @@ -70,7 +72,7 @@ func (r *Repository) EnsurePermission(ctx context.Context, name, description str } // AddRolePermission links a permission string to a role. -func (r *Repository) AddRolePermission(ctx context.Context, roleID uint, permName string) error { +func (r *Repository) AddRolePermission(ctx context.Context, roleID uuid.UUID, permName string) error { rp := RolePermission{ RoleID: roleID, PermissionName: permName, @@ -79,12 +81,12 @@ func (r *Repository) AddRolePermission(ctx context.Context, roleID uint, permNam } // RemoveRolePermission removes a permission string from a role. -func (r *Repository) RemoveRolePermission(ctx context.Context, roleID uint, permName string) error { +func (r *Repository) RemoveRolePermission(ctx context.Context, roleID uuid.UUID, permName string) error { return r.db.WithContext(ctx).Where("role_id = ? AND permission_name = ?", roleID, permName).Delete(&RolePermission{}).Error } // AssignUserRole links a user to a specific role. -func (r *Repository) AssignUserRole(ctx context.Context, userID uint, roleID uint) error { +func (r *Repository) AssignUserRole(ctx context.Context, userID uuid.UUID, roleID uuid.UUID) error { ur := UserRole{ UserID: userID, RoleID: roleID, @@ -93,12 +95,12 @@ func (r *Repository) AssignUserRole(ctx context.Context, userID uint, roleID uin } // RemoveUserRole revokes a role from a user. -func (r *Repository) RemoveUserRole(ctx context.Context, userID uint, roleID uint) error { +func (r *Repository) RemoveUserRole(ctx context.Context, userID uuid.UUID, roleID uuid.UUID) error { return r.db.WithContext(ctx).Where("user_id = ? AND role_id = ?", userID, roleID).Delete(&UserRole{}).Error } // GetUserPermissions returns all permission strings granted to a user via assigned roles. -func (r *Repository) GetUserPermissions(ctx context.Context, userID uint) ([]string, error) { +func (r *Repository) GetUserPermissions(ctx context.Context, userID uuid.UUID) ([]string, error) { var perms []string err := r.db.WithContext(ctx). Table("role_permissions"). diff --git a/internal/auth/router.go b/internal/auth/router.go index e5e1867..e9d75cf 100644 --- a/internal/auth/router.go +++ b/internal/auth/router.go @@ -6,9 +6,9 @@ import ( "github.com/danielgtaylor/huma/v2" "gorm.io/gorm" - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/pkg/authz" - "github.com/codetheuri/tusk/pkg/logger" + "github.com/codetheuri/tusk/v2/config" + "github.com/codetheuri/tusk/v2/pkg/authz" + "github.com/codetheuri/tusk/v2/pkg/logger" ) func RegisterRoutes(api huma.API, db *gorm.DB, cfg *config.Config, log logger.Logger) { @@ -38,7 +38,7 @@ func RegisterRoutes(api huma.API, db *gorm.DB, cfg *config.Config, log logger.Lo Path: "/api/v1/auth/login", Summary: "Login user", Description: "Authenticates a user via flexible single field (username, email, or phone) and returns Access + Refresh tokens.", - Tags: []string{"Authentication"} , + Tags: []string{"Authentication"}, }, handler.Login) huma.Register(api, huma.Operation{ diff --git a/internal/auth/service_auth.go b/internal/auth/service_auth.go index dcf4eca..d070575 100644 --- a/internal/auth/service_auth.go +++ b/internal/auth/service_auth.go @@ -1,6 +1,8 @@ package auth import ( + "github.com/google/uuid" + "context" "crypto/rand" "crypto/sha256" @@ -11,10 +13,10 @@ import ( "github.com/golang-jwt/jwt/v5" "golang.org/x/crypto/bcrypt" - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/internal/middleware" - "github.com/codetheuri/tusk/pkg/authz" - "github.com/codetheuri/tusk/pkg/query" + "github.com/codetheuri/tusk/v2/config" + "github.com/codetheuri/tusk/v2/pkg/authz" + "github.com/codetheuri/tusk/v2/pkg/middleware" + "github.com/codetheuri/tusk/v2/pkg/query" ) // Service encapsulates the business logic for Identity, Auth, and RBAC. @@ -159,7 +161,7 @@ func (s *Service) Logout(ctx context.Context, rawRefreshToken string) error { } // GetCurrentUser returns the user model (with preloaded Profile) and permissions for the authenticated user. -func (s *Service) GetCurrentUser(ctx context.Context, userID uint) (*User, []string, error) { +func (s *Service) GetCurrentUser(ctx context.Context, userID uuid.UUID) (*User, []string, error) { user, err := s.repo.FindByID(ctx, userID) if err != nil { return nil, nil, fmt.Errorf("user not found: %w", err) @@ -174,7 +176,7 @@ func (s *Service) GetCurrentUser(ctx context.Context, userID uint) (*User, []str } // UpdateProfile updates the authenticated user's personal identity profile details. -func (s *Service) UpdateProfile(ctx context.Context, userID uint, firstName, lastName, avatar, bio string) (*UserProfile, error) { +func (s *Service) UpdateProfile(ctx context.Context, userID uuid.UUID, firstName, lastName, avatar, bio string) (*UserProfile, error) { user, err := s.repo.FindByID(ctx, userID) if err != nil { return nil, fmt.Errorf("user not found: %w", err) @@ -213,6 +215,9 @@ func (s *Service) generateAccessToken(user *User) (string, error) { expiry := time.Now().Add(s.cfg.AccessTokenTTL) claims := middleware.Claims{ UserID: user.ID, + // Carried in the token so authorization does not need a database round-trip + // on every request. See middleware.Claims for the revocation tradeoff. + IsSuperUser: user.IsSuperUser, RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(expiry), IssuedAt: jwt.NewNumericDate(time.Now()), @@ -222,7 +227,7 @@ func (s *Service) generateAccessToken(user *User) (string, error) { return token.SignedString([]byte(s.cfg.JWTSecret)) } -func (s *Service) issueRefreshToken(ctx context.Context, userID uint) (string, error) { +func (s *Service) issueRefreshToken(ctx context.Context, userID uuid.UUID) (string, error) { bytes := make([]byte, 32) if _, err := rand.Read(bytes); err != nil { return "", err diff --git a/internal/auth/service_role.go b/internal/auth/service_role.go index 193d7f5..0e5faa3 100644 --- a/internal/auth/service_role.go +++ b/internal/auth/service_role.go @@ -1,10 +1,12 @@ package auth import ( + "github.com/google/uuid" + "context" "fmt" - "github.com/codetheuri/tusk/pkg/authz" + "github.com/codetheuri/tusk/v2/pkg/authz" ) type CreateRoleRequest struct { @@ -13,7 +15,7 @@ type CreateRoleRequest struct { } type UpdateRoleRequest struct { - ID uint + ID uuid.UUID Name *string Description *string } @@ -42,7 +44,7 @@ func (s *Service) ListRoles(ctx context.Context) ([]Role, error) { } // GetRoleByID returns a single role by ID. -func (s *Service) GetRoleByID(ctx context.Context, id uint) (*Role, error) { +func (s *Service) GetRoleByID(ctx context.Context, id uuid.UUID) (*Role, error) { return s.repo.GetRoleByID(ctx, id) } @@ -68,12 +70,12 @@ func (s *Service) UpdateRole(ctx context.Context, req *UpdateRoleRequest) (*Role } // DeleteRole removes a role from the system. -func (s *Service) DeleteRole(ctx context.Context, id uint) error { +func (s *Service) DeleteRole(ctx context.Context, id uuid.UUID) error { return s.repo.DeleteRole(ctx, id) } // AddRolePermission attaches a permission string to a role. -func (s *Service) AddRolePermission(ctx context.Context, roleID uint, permName string) error { +func (s *Service) AddRolePermission(ctx context.Context, roleID uuid.UUID, permName string) error { perm, exists := authz.DefaultRegistry().Find(permName) if !exists { return fmt.Errorf("permission '%s' is not a valid system permission", permName) @@ -87,22 +89,22 @@ func (s *Service) AddRolePermission(ctx context.Context, roleID uint, permName s } // RemoveRolePermission detaches a permission string from a role. -func (s *Service) RemoveRolePermission(ctx context.Context, roleID uint, permName string) error { +func (s *Service) RemoveRolePermission(ctx context.Context, roleID uuid.UUID, permName string) error { return s.repo.RemoveRolePermission(ctx, roleID, permName) } // AssignUserRole assigns a role to a user after checking existence. -func (s *Service) AssignUserRole(ctx context.Context, userID uint, roleID uint) error { +func (s *Service) AssignUserRole(ctx context.Context, userID uuid.UUID, roleID uuid.UUID) error { if _, err := s.repo.GetRoleByID(ctx, roleID); err != nil { - return fmt.Errorf("role with ID %d does not exist", roleID) + return fmt.Errorf("role with ID %s does not exist", roleID) } if _, err := s.repo.FindByID(ctx, userID); err != nil { - return fmt.Errorf("user with ID %d does not exist", userID) + return fmt.Errorf("user with ID %s does not exist", userID) } return s.repo.AssignUserRole(ctx, userID, roleID) } // RemoveUserRole revokes a role from a user. -func (s *Service) RemoveUserRole(ctx context.Context, userID uint, roleID uint) error { +func (s *Service) RemoveUserRole(ctx context.Context, userID uuid.UUID, roleID uuid.UUID) error { return s.repo.RemoveUserRole(ctx, userID, roleID) } diff --git a/internal/middleware/jwt.go b/internal/middleware/jwt.go deleted file mode 100644 index 2b8e9b9..0000000 --- a/internal/middleware/jwt.go +++ /dev/null @@ -1,139 +0,0 @@ -package middleware - -import ( - "context" - "net/http" - "strings" - - "github.com/danielgtaylor/huma/v2" - "github.com/golang-jwt/jwt/v5" - "gorm.io/gorm" -) - -type contextKey string - -const ( - ContextKeyUserID contextKey = "user_id" - ContextKeyRole contextKey = "role" - ContextKeyJTI contextKey = "jti" -) - -type Claims struct { - UserID uint `json:"user_id"` - Role string `json:"role"` - jwt.RegisteredClaims -} - -func Authenticate(jwtSecret string) func(http.Handler) http.Handler { - return func(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - authHeader := r.Header.Get("Authorization") - if authHeader == "" { - // To keep it clean, we'd normally call a helper. - // We can just use http.Error or a custom response. - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusUnauthorized) - w.Write([]byte(`{"success":false,"error":{"code":"UNAUTHORIZED","message":"Authorization header is required"}}`)) - return - } - - parts := strings.SplitN(authHeader, " ", 2) - if len(parts) != 2 || strings.ToLower(parts[0]) != "bearer" { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusUnauthorized) - w.Write([]byte(`{"success":false,"error":{"code":"UNAUTHORIZED","message":"Invalid authorization format. Expected: Bearer "}}`)) - return - } - - tokenString := parts[1] - claims := &Claims{} - - token, err := jwt.ParseWithClaims(tokenString, claims, func(token *jwt.Token) (interface{}, error) { - if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok { - return nil, jwt.ErrSignatureInvalid - } - return []byte(jwtSecret), nil - }) - - if err != nil || !token.Valid { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusUnauthorized) - w.Write([]byte(`{"success":false,"error":{"code":"UNAUTHORIZED","message":"Invalid or expired token"}}`)) - return - } - - ctx := context.WithValue(r.Context(), ContextKeyUserID, claims.UserID) - ctx = context.WithValue(ctx, "user_id", claims.UserID) - ctx = context.WithValue(ctx, ContextKeyRole, claims.Role) - ctx = context.WithValue(ctx, "role", claims.Role) - ctx = context.WithValue(ctx, ContextKeyJTI, claims.ID) - ctx = context.WithValue(ctx, "jti", claims.ID) - - next.ServeHTTP(w, r.WithContext(ctx)) - }) - } -} - -func RequireRole(allowedRoles ...string) func(http.Handler) http.Handler { - allowedSet := make(map[string]bool, len(allowedRoles)) - for _, r := range allowedRoles { - allowedSet[r] = true - } - - return func(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - role, ok := r.Context().Value(ContextKeyRole).(string) - if !ok || !allowedSet[role] { - w.Header().Set("Content-Type", "application/json") - w.WriteHeader(http.StatusForbidden) - w.Write([]byte(`{"success":false,"error":{"code":"FORBIDDEN","message":"You do not have permission to perform this action"}}`)) - return - } - next.ServeHTTP(w, r) - }) - } -} - -func GetUserID(ctx context.Context) uint { - id, _ := ctx.Value(ContextKeyUserID).(uint) - return id -} - -// HumaAuthenticate decodes the JWT Authorization header into the Huma request context. -func HumaAuthenticate(api huma.API, jwtSecret string, db *gorm.DB) func(huma.Context, func(huma.Context)) { - return func(ctx huma.Context, next func(huma.Context)) { - authHeader := ctx.Header("Authorization") - if authHeader == "" { - next(ctx) - return - } - - parts := strings.SplitN(authHeader, " ", 2) - if len(parts) == 2 && strings.ToLower(parts[0]) == "bearer" { - tokenString := parts[1] - claims := &Claims{} - token, err := jwt.ParseWithClaims(tokenString, claims, func(token *jwt.Token) (interface{}, error) { - if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok { - return nil, jwt.ErrSignatureInvalid - } - return []byte(jwtSecret), nil - }) - - if err == nil && token.Valid { - reqCtx := ctx.Context() - reqCtx = context.WithValue(reqCtx, ContextKeyUserID, claims.UserID) - reqCtx = context.WithValue(reqCtx, "user_id", claims.UserID) - reqCtx = context.WithValue(reqCtx, ContextKeyRole, claims.Role) - reqCtx = context.WithValue(reqCtx, "role", claims.Role) - - var isSuper bool - db.WithContext(reqCtx).Table("users").Where("id = ?", claims.UserID).Pluck("is_super_user", &isSuper) - reqCtx = context.WithValue(reqCtx, "is_super_user", isSuper) - - ctx = huma.WithContext(ctx, reqCtx) - } - } - - next(ctx) - } -} diff --git a/pkg/app/app.go b/pkg/app/app.go new file mode 100644 index 0000000..dc3b2a8 --- /dev/null +++ b/pkg/app/app.go @@ -0,0 +1,310 @@ +// Package app assembles the pieces every Tusk service needs: a router, the +// middleware chain in the right order, health endpoints, a configured Huma API, +// and graceful shutdown. +// +// It deliberately registers no routes of its own. An application supplies its +// own modules against the returned API: +// +// a, err := app.New(app.Options{ +// Config: cfg, Logger: log, +// Title: "Salio API", +// Tags: []*huma.Tag{{Name: "Customers"}}, +// }) +// customer.RegisterRoutes(a.API(), a.DB(), cfg, log) +// ledger.RegisterRoutes(a.API(), a.DB(), cfg, log) +// a.Run() +// +// Tusk's own auth module is registered the same way, by cmd/api, rather than +// from inside here — so a service that wants different authentication is not +// fighting the framework to get it. +package app + +import ( + "context" + "encoding/json" + "fmt" + "net" + "net/http" + "os" + "os/signal" + "syscall" + "time" + + "github.com/danielgtaylor/huma/v2" + "github.com/danielgtaylor/huma/v2/adapters/humachi" + "github.com/go-chi/chi/v5" + "gorm.io/gorm" + + "github.com/codetheuri/tusk/v2/config" + "github.com/codetheuri/tusk/v2/database" + "github.com/codetheuri/tusk/v2/pkg/logger" + "github.com/codetheuri/tusk/v2/pkg/middleware" + "github.com/codetheuri/tusk/v2/pkg/response" +) + +// TagGroup collects related OpenAPI tags under one heading in the docs sidebar. +type TagGroup struct { + Name string + Tags []string +} + +// Options configures an App. +type Options struct { + // Config and Logger are required. + Config *config.Config + Logger logger.Logger + + // DB is optional. When nil, New opens a connection from Config — which is + // what most services want. Supplying one is for tests, and for a service + // that needs to configure the pool itself. + DB *gorm.DB + + // Title and Version name the API in its documentation. They default to the + // application name and version from Config. + Title string + Version string + + // Description is rendered above the endpoint list. Markdown. + Description string + + Contact *huma.Contact + + // Tags describe the endpoint groups. TagGroups optionally nests them under + // headings in the sidebar. + Tags []*huma.Tag + TagGroups []TagGroup +} + +// App is a configured but not yet running service. +type App struct { + cfg *config.Config + log logger.Logger + db *gorm.DB + router *chi.Mux + api huma.API +} + +// New builds the service. Register routes on the returned App, then call Run. +func New(opts Options) (*App, error) { + if opts.Config == nil { + return nil, fmt.Errorf("app: Options.Config is required") + } + if opts.Logger == nil { + return nil, fmt.Errorf("app: Options.Logger is required") + } + cfg, log := opts.Config, opts.Logger + + db := opts.DB + if db == nil { + var err error + if db, err = database.Connect(cfg, log); err != nil { + return nil, fmt.Errorf("database connection failed: %w", err) + } + } + + r := chi.NewRouter() + + // Middleware order is significant — each layer wraps everything below it. + // RequestID first so every later log line and error carries a correlation ID; + // Recovery above the handlers so a panic becomes a 500 rather than killing the + // process; rate limiting and the body cap before any handler allocates memory + // on behalf of an unauthenticated caller. + rateLimiter := middleware.NewRateLimiter(cfg.RateLimitBurst, cfg.RateLimitRPS, log) + + r.Use(middleware.RequestID()) + r.Use(middleware.Logger(log)) + r.Use(middleware.Recovery(log)) + r.Use(middleware.CORS(cfg.CORSOrigins, log)) + r.Use(middleware.SecurityHeaders) + r.Use(rateLimiter.Limit()) + r.Use(middleware.MaxBodyBytes(cfg.MaxRequestBodyBytes)) + + // Liveness: is the process up? Deliberately does not touch the database — a + // liveness probe that fails on a database blip gets the container killed and + // restarted, which does nothing to fix the database and drops live traffic. + r.Get("/health", healthHandler(cfg)) + r.Get("/live", healthHandler(cfg)) + + // Readiness: can this instance serve traffic right now? This one does check + // the database, because an instance that cannot reach it should be removed + // from the load balancer rather than restarted. + r.Get("/ready", readyHandler(db)) + + response.SetupHuma() + + title := opts.Title + if title == "" { + title = cfg.AppName + } + version := opts.Version + if version == "" { + version = cfg.AppVersion + } + + humaConfig := huma.DefaultConfig(title, version) + humaConfig.Info.Description = opts.Description + humaConfig.Info.Contact = opts.Contact + humaConfig.Tags = opts.Tags + + // Add Bearer JWT Security Scheme to OpenAPI docs + humaConfig.Components = &huma.Components{ + SecuritySchemes: map[string]*huma.SecurityScheme{ + "bearerAuth": { + Type: "http", + Scheme: "bearer", + BearerFormat: "JWT", + Description: "Enter token as: Bearer ", + }, + }, + } + + // Disable $schema from showing up in OpenAPI docs/responses + humaConfig.CreateHooks = []func(huma.Config) huma.Config{ + func(c huma.Config) huma.Config { + c.SchemasPath = "" + return c + }, + } + + // Documentation gating. + // + // Blanking these paths removes the routes entirely rather than hiding them, so + // there is no endpoint left to probe. API endpoints are untouched — only the + // human-facing docs UI and the machine-readable spec disappear. + // + // The spec describes every route, parameter and schema in the system, which is + // a convenient map for anyone looking for a way in. Publishing it is a choice, + // and in production the default answer is no. + if !cfg.DocsEnabled { + humaConfig.DocsPath = "" + humaConfig.OpenAPIPath = "" + } + + api := humachi.New(r, humaConfig) + + if len(opts.TagGroups) > 0 { + openAPI := api.OpenAPI() + if openAPI.Extensions == nil { + openAPI.Extensions = map[string]any{} + } + groups := make([]map[string]any, 0, len(opts.TagGroups)) + for _, g := range opts.TagGroups { + groups = append(groups, map[string]any{"name": g.Name, "tags": g.Tags}) + } + openAPI.Extensions["x-tagGroups"] = groups + } + + // Registered before any route, so every operation inherits it. + api.UseMiddleware(middleware.HumaAuthenticate(api, cfg.JWTSecret)) + + return &App{cfg: cfg, log: log, db: db, router: r, api: api}, nil +} + +// API returns the Huma API to register operations against. +func (a *App) API() huma.API { return a.api } + +// Router returns the underlying chi router, for routes that are not part of the +// API — a server-rendered admin console, or static assets. +func (a *App) Router() *chi.Mux { return a.router } + +// DB returns the database handle the App is using. +func (a *App) DB() *gorm.DB { return a.db } + +// Config returns the configuration the App was built with. +func (a *App) Config() *config.Config { return a.cfg } + +func (a *App) Run() error { + srv := &http.Server{ + Addr: fmt.Sprintf(":%d", a.cfg.ServerPort), + Handler: a.router, + + // ReadHeaderTimeout is the defence against Slowloris: a client that opens a + // connection and dribbles headers forever holds a goroutine hostage. It must + // be set even when ReadTimeout is, because ReadTimeout alone permits a slow + // header phase to consume the entire budget. + ReadHeaderTimeout: 10 * time.Second, + ReadTimeout: a.cfg.ReadTimeout, + WriteTimeout: a.cfg.WriteTimeout, + IdleTimeout: a.cfg.IdleTimeout, + MaxHeaderBytes: 1 << 20, // 1 MiB + } + + ln, err := net.Listen("tcp", srv.Addr) + if err != nil { + return fmt.Errorf("failed to start listener: %w", err) + } + + actualAddr := ln.Addr().(*net.TCPAddr) + a.log.Info(fmt.Sprintf("Server is listening on port %d", actualAddr.Port)) + if a.cfg.DocsEnabled { + a.log.Info(fmt.Sprintf("API documentation at http://localhost:%d/docs", actualAddr.Port)) + } else { + a.log.Info("API documentation is disabled (set DOCS_ENABLED=true to serve /docs and /openapi.json)") + } + + go func() { + if err := srv.Serve(ln); err != nil && err != http.ErrServerClosed { + a.log.Fatal("Server failed to listen or serve", err) + } + }() + + quit := make(chan os.Signal, 1) + signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM) + sig := <-quit + a.log.Warn("Received shutdown signal", "signal", sig.String()) + + ctx, cancel := context.WithTimeout(context.Background(), a.cfg.ShutdownTimeout) + defer cancel() + + a.log.Info("Attempting to shut down gracefully...") + if err := srv.Shutdown(ctx); err != nil { + return fmt.Errorf("server shutdown failed: %w", err) + } + + a.log.Info("Server shut down gracefully.") + return nil +} + +// healthHandler reports process liveness. It deliberately performs no dependency +// checks — see the route registration for why. +func healthHandler(cfg *config.Config) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + writeJSON(w, http.StatusOK, map[string]any{ + "status": "ok", + "app_name": cfg.AppName, + "version": cfg.AppVersion, + "time": time.Now().UTC().Format(time.RFC3339), + }) + } +} + +// readyHandler reports whether this instance can serve traffic, which means +// checking that the database is actually reachable. +func readyHandler(db *gorm.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + ctx, cancel := context.WithTimeout(r.Context(), 2*time.Second) + defer cancel() + + sqlDB, err := db.DB() + if err != nil || sqlDB.PingContext(ctx) != nil { + writeJSON(w, http.StatusServiceUnavailable, map[string]any{ + "status": "unavailable", + "database": "disconnected", + }) + return + } + writeJSON(w, http.StatusOK, map[string]any{ + "status": "ready", + "database": "connected", + }) + } +} + +// writeJSON encodes a response body. Hand-formatting JSON with fmt.Sprintf works +// until a value contains a quote or a backslash, at which point it silently emits +// malformed JSON — encoding/json escapes correctly by construction. +func writeJSON(w http.ResponseWriter, status int, body any) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(body) +} diff --git a/pkg/authz/context.go b/pkg/authz/context.go new file mode 100644 index 0000000..92bcae0 --- /dev/null +++ b/pkg/authz/context.go @@ -0,0 +1,33 @@ +package authz + +import "context" + +// subjectKey is the context key under which the authenticated Subject is stored. +// +// It is an unexported struct type rather than a string on purpose. Context keys +// are compared by type as well as value, so an unexported type cannot be forged +// or accidentally collided with from outside this package — even by code that +// uses the identical name. A plain string key such as "user_id" offers no such +// protection: any dependency writing the same string silently overwrites the +// value, at runtime, with no compile error. +type subjectKey struct{} + +// WithSubject returns a copy of ctx carrying the authenticated subject. +// +// Authentication middleware calls this exactly once per request. Everything +// downstream — guards, handlers, services — reads it back through +// SubjectFromContext rather than reaching for individual claim values. +func WithSubject(ctx context.Context, sub Subject) context.Context { + return context.WithValue(ctx, subjectKey{}, sub) +} + +// SubjectFromContext returns the authenticated subject and whether one is present. +// +// The boolean must be checked. An anonymous request yields the zero Subject, +// which has UserID 0 and IsSuperUser false — indistinguishable from a real +// subject if the caller ignores the flag. Treat false as "unauthenticated", +// never as "a user with no permissions". +func SubjectFromContext(ctx context.Context) (Subject, bool) { + sub, ok := ctx.Value(subjectKey{}).(Subject) + return sub, ok +} diff --git a/pkg/authz/middleware.go b/pkg/authz/middleware.go index d8e6558..da50b60 100644 --- a/pkg/authz/middleware.go +++ b/pkg/authz/middleware.go @@ -7,7 +7,7 @@ import ( "github.com/danielgtaylor/huma/v2" "gorm.io/gorm" - "github.com/codetheuri/tusk/pkg/response" + "github.com/codetheuri/tusk/v2/pkg/response" ) // Security Requirement Helpers for Huma OpenAPI documentation. @@ -20,28 +20,19 @@ var ( ) // SubjectExtractor defines a function signature for extracting Subject from context.Context. +// +// It remains configurable so an application can source identity from somewhere +// other than Tusk's own authentication middleware — a gateway header or an mTLS +// certificate, for example. type SubjectExtractor func(ctx context.Context) (Subject, bool) -// DefaultSubjectExtractor extracts UserID and IsSuperUser from context standard keys. -var DefaultSubjectExtractor SubjectExtractor = func(ctx context.Context) (Subject, bool) { - var sub Subject - - // Extract UserID - if uid, ok := ctx.Value("user_id").(uint); ok { - sub.UserID = uid - } else if uidInt, ok := ctx.Value("user_id").(int); ok { - sub.UserID = uint(uidInt) - } else { - return sub, false - } - - // Extract IsSuperUser flag - if isSuper, ok := ctx.Value("is_super_user").(bool); ok { - sub.IsSuperUser = isSuper - } - - return sub, true -} +// DefaultSubjectExtractor reads the Subject placed in context by authentication +// middleware via WithSubject. +// +// It deliberately does no claim parsing of its own. Identity is resolved once, +// at the edge, and every consumer reads the same value — so there is exactly one +// place where "who is this?" is answered. +var DefaultSubjectExtractor SubjectExtractor = SubjectFromContext // RequirePolicy creates a HTTP middleware that enforces any Policy rule. func RequirePolicy(db *gorm.DB, policy Policy) func(http.Handler) http.Handler { diff --git a/pkg/authz/models.go b/pkg/authz/models.go index 16662af..12eaf34 100644 --- a/pkg/authz/models.go +++ b/pkg/authz/models.go @@ -1,6 +1,13 @@ package authz -import "time" +import ( + "time" + + "github.com/google/uuid" + "gorm.io/gorm" + + "github.com/codetheuri/tusk/v2/pkg/id" +) // PermissionRecord represents the runtime database representation of a code permission. type PermissionRecord struct { @@ -17,21 +24,29 @@ func (PermissionRecord) TableName() string { // Role represents an administrative role grouping multiple permissions. type Role struct { - ID uint `gorm:"primaryKey;autoIncrement" json:"id"` - Name string `gorm:"type:varchar(191);uniqueIndex;not null" json:"name"` - Description string `gorm:"type:text" json:"description"` - Permissions []PermissionRecord `gorm:"many2many:role_permissions;foreignKey:ID;joinForeignKey:role_id;references:Name;joinReferences:permission_name" json:"permissions,omitempty"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` + ID uuid.UUID `gorm:"type:uuid;primaryKey" json:"id"` + Name string `gorm:"type:varchar(191);uniqueIndex;not null" json:"name"` + Description string `gorm:"type:text" json:"description"` + Permissions []PermissionRecord `gorm:"many2many:role_permissions;foreignKey:ID;joinForeignKey:role_id;references:Name;joinReferences:permission_name" json:"permissions,omitempty"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` } // RolePermission links a Role to a PermissionRecord. type RolePermission struct { - RoleID uint `gorm:"primaryKey;index" json:"role_id"` - PermissionName string `gorm:"primaryKey;type:varchar(191);index" json:"permission_name"` + RoleID uuid.UUID `gorm:"type:uuid;primaryKey;index" json:"role_id"` + PermissionName string `gorm:"primaryKey;type:varchar(191);index" json:"permission_name"` CreatedAt time.Time `json:"created_at"` } +// BeforeCreate assigns an identifier when one was not supplied. +func (r *Role) BeforeCreate(*gorm.DB) error { + if id.IsZero(r.ID) { + r.ID = id.New() + } + return nil +} + // TableName explicitly sets table name for RolePermission. func (RolePermission) TableName() string { return "role_permissions" @@ -39,8 +54,8 @@ func (RolePermission) TableName() string { // UserRole links a User ID to a Role ID. type UserRole struct { - UserID uint `gorm:"primaryKey;index" json:"user_id"` - RoleID uint `gorm:"primaryKey;index" json:"role_id"` + UserID uuid.UUID `gorm:"type:uuid;primaryKey;index" json:"user_id"` + RoleID uuid.UUID `gorm:"type:uuid;primaryKey;index" json:"role_id"` CreatedAt time.Time `json:"created_at"` } diff --git a/pkg/authz/policy.go b/pkg/authz/policy.go index 53dfd62..a8ab06d 100644 --- a/pkg/authz/policy.go +++ b/pkg/authz/policy.go @@ -4,12 +4,13 @@ import ( "context" "fmt" + "github.com/google/uuid" "gorm.io/gorm" ) // Subject represents the authenticated user context performing an action. type Subject struct { - UserID uint + UserID uuid.UUID IsSuperUser bool } @@ -40,7 +41,7 @@ func (e *Evaluator) IsAuthorized(ctx context.Context, sub Subject, policy Policy } // UserHasPermission queries whether a user possesses a specific permission through any assigned role. -func (e *Evaluator) UserHasPermission(ctx context.Context, userID uint, permission string) (bool, error) { +func (e *Evaluator) UserHasPermission(ctx context.Context, userID uuid.UUID, permission string) (bool, error) { var count int64 err := e.db.WithContext(ctx).Table("user_roles"). Joins("JOIN role_permissions ON role_permissions.role_id = user_roles.role_id"). diff --git a/pkg/id/id.go b/pkg/id/id.go new file mode 100644 index 0000000..7ff7af2 --- /dev/null +++ b/pkg/id/id.go @@ -0,0 +1,55 @@ +// Package id generates Tusk's primary keys. +// +// It exists to state one decision in one place: Tusk identifiers are UUIDv7. +// Without it, that choice would be re-made implicitly in every model file, and +// changing it would mean finding every call site. +// +// # Why UUID rather than an auto-incrementing integer +// +// An auto-increment key can only be issued by the database, so it cannot be +// created by an offline client, and two databases cannot be merged without +// renumbering. The cost of choosing UUID where an integer would have sufficed is +// eight extra bytes per key; the cost of choosing an integer where UUID was +// needed is rewriting every table and foreign key in the system. The mistakes +// are not symmetric, so Tusk takes the cheap one. +// +// # Why v7 rather than v4 +// +// A primary key lives in a sorted B-tree. Random v4 keys land at arbitrary +// positions, so each insert dirties a different page and the index fragments. +// v7 places a millisecond timestamp in its leading 48 bits, so new keys sort to +// the end of the tree — the write pattern of an auto-increment key, with the +// independence of a UUID. +// +// # What v7 gives away +// +// The timestamp is readable by anyone holding the identifier: an ID reveals when +// its row was created, and any two IDs reveal their relative order. For ordinary +// business records this is unremarkable, since creation time is usually visible +// anyway. Where that correlation is sensitive, expose a separate random +// identifier and keep the v7 key internal. +// +// Note that this is not an enumeration risk. v7 retains 74 random bits, which is +// far beyond guessing; only the timestamp is inferable, never the whole value. +package id + +import "github.com/google/uuid" + +// New returns a new time-ordered identifier. +// +// It panics only if the system's cryptographic random source fails, which is not +// a condition any caller can meaningfully handle — a process that cannot generate +// random bytes cannot safely issue tokens or hash passwords either, so failing +// loudly at the point of breakage is better than propagating an error that every +// call site would ignore. +func New() uuid.UUID { + return uuid.Must(uuid.NewV7()) +} + +// IsZero reports whether an identifier is unset. +// +// Useful in GORM BeforeCreate hooks, which must assign an ID only when the caller +// has not supplied one — an offline client that generated its own ID must keep it. +func IsZero(u uuid.UUID) bool { + return u == uuid.Nil +} diff --git a/pkg/id/id_test.go b/pkg/id/id_test.go new file mode 100644 index 0000000..fe5b3a8 --- /dev/null +++ b/pkg/id/id_test.go @@ -0,0 +1,72 @@ +package id + +import ( + "testing" + "time" + + "github.com/google/uuid" +) + +func TestNew_ProducesVersion7(t *testing.T) { + u := New() + if got := u.Version(); got != 7 { + t.Errorf("UUID version = %d, want 7", got) + } + if got := u.Variant(); got != uuid.RFC4122 { + t.Errorf("UUID variant = %v, want RFC4122", got) + } +} + +// The whole reason for choosing v7 over v4 is that keys sort by creation time, +// which is what preserves B-tree insert locality. If that ordering ever broke, +// the choice would have no remaining benefit over v4 — so it is worth asserting. +func TestNew_IsTimeOrdered(t *testing.T) { + const n = 1000 + ids := make([]uuid.UUID, n) + for i := range ids { + ids[i] = New() + } + + for i := 1; i < n; i++ { + prev, curr := ids[i-1].String(), ids[i].String() + if curr < prev { + t.Fatalf("identifier %d sorts before its predecessor:\n prev: %s\n curr: %s", i, prev, curr) + } + } +} + +func TestNew_IsUnique(t *testing.T) { + const n = 10000 + seen := make(map[uuid.UUID]struct{}, n) + for i := 0; i < n; i++ { + u := New() + if _, dup := seen[u]; dup { + t.Fatalf("duplicate identifier generated after %d draws: %s", i, u) + } + seen[u] = struct{}{} + } +} + +// The embedded timestamp is a documented property, not an accident — the tradeoff +// it implies (an ID reveals its creation time) is called out in the package docs. +func TestNew_EmbedsCurrentTimestamp(t *testing.T) { + before := time.Now().Add(-time.Second) + u := New() + after := time.Now().Add(time.Second) + + sec, nsec := u.Time().UnixTime() + created := time.Unix(sec, nsec) + + if created.Before(before) || created.After(after) { + t.Errorf("embedded timestamp %s is outside the expected window [%s, %s]", created, before, after) + } +} + +func TestIsZero(t *testing.T) { + if !IsZero(uuid.Nil) { + t.Error("IsZero(uuid.Nil) = false, want true") + } + if IsZero(New()) { + t.Error("IsZero(New()) = true, want false") + } +} diff --git a/pkg/logger/logger.go b/pkg/logger/logger.go index e751e72..f367305 100755 --- a/pkg/logger/logger.go +++ b/pkg/logger/logger.go @@ -5,7 +5,7 @@ import ( "log" "os" "strings" - // "sync" + // "sync" ) // Logger interface for logging operations @@ -29,20 +29,20 @@ func NewConsoleLogger() Logger { } } -//format args +// format args func formatArgs(args ...any) string { if len(args) == 0 { return "" } var parts []string - for i := 0; i< len(args); i += 2 { + for i := 0; i < len(args); i += 2 { key := fmt.Sprintf("%v", args[i]) if i+1 < len(args) { value := fmt.Sprintf("%v", args[i+1]) parts = append(parts, fmt.Sprintf("%s=%s", key, value)) - }else { - parts = append(parts, fmt.Sprintf("%s=",key)) - } + } else { + parts = append(parts, fmt.Sprintf("%s=", key)) + } } return " " + strings.Join(parts, ", ") } @@ -72,4 +72,3 @@ func (l *consoleLogger) Fatal(msg string, err error, args ...any) { l.Error(msg, err, args...) os.Exit(1) } - diff --git a/pkg/logger/logger_test.go b/pkg/logger/logger_test.go index a66faa6..a0e83f3 100644 --- a/pkg/logger/logger_test.go +++ b/pkg/logger/logger_test.go @@ -2,8 +2,10 @@ package logger import ( "bytes" + "encoding/json" "errors" "log" + "log/slog" "strings" "testing" ) @@ -35,3 +37,77 @@ func TestConsoleLogger_Error(t *testing.T) { t.Errorf("unexpected error log output: %s", output) } } + +func TestSlogLogger_JSONIsQueryable(t *testing.T) { + var buf bytes.Buffer + l := &slogLogger{l: slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))} + + l.Info("user logged in", "user_id", 42, "role", "admin") + + var record map[string]any + if err := json.Unmarshal(buf.Bytes(), &record); err != nil { + t.Fatalf("output was not valid JSON: %v (%q)", err, buf.String()) + } + if record["msg"] != "user logged in" { + t.Errorf("msg = %v, want %q", record["msg"], "user logged in") + } + // The point of structured logging: user_id is its own field, not text + // embedded in the message that would need a regex to find. + if record["user_id"] != float64(42) { + t.Errorf("user_id = %v, want 42", record["user_id"]) + } + if record["role"] != "admin" { + t.Errorf("role = %v, want admin", record["role"]) + } +} + +func TestSlogLogger_ErrorAttachesErrorField(t *testing.T) { + var buf bytes.Buffer + l := &slogLogger{l: slog.New(slog.NewJSONHandler(&buf, nil))} + + l.Error("save failed", errors.New("connection refused"), "table", "users") + + var record map[string]any + if err := json.Unmarshal(buf.Bytes(), &record); err != nil { + t.Fatalf("output was not valid JSON: %v", err) + } + if record["error"] != "connection refused" { + t.Errorf("error field = %v, want %q", record["error"], "connection refused") + } + if record["table"] != "users" { + t.Errorf("table = %v, want users", record["table"]) + } +} + +func TestSlogLogger_RespectsLevel(t *testing.T) { + var buf bytes.Buffer + l := &slogLogger{l: slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelWarn}))} + + l.Info("this should be filtered out") + if buf.Len() != 0 { + t.Errorf("expected Info to be suppressed at Warn level, got: %s", buf.String()) + } + + l.Warn("this should appear") + if buf.Len() == 0 { + t.Error("expected Warn to be emitted at Warn level") + } +} + +func TestParseLevel(t *testing.T) { + cases := map[string]slog.Level{ + "debug": slog.LevelDebug, + "DEBUG": slog.LevelDebug, + "warn": slog.LevelWarn, + "warning": slog.LevelWarn, + "error": slog.LevelError, + "info": slog.LevelInfo, + "": slog.LevelInfo, + "nonsense": slog.LevelInfo, // must not silence logging + } + for input, want := range cases { + if got := parseLevel(input); got != want { + t.Errorf("parseLevel(%q) = %v, want %v", input, got, want) + } + } +} diff --git a/pkg/logger/slog.go b/pkg/logger/slog.go new file mode 100644 index 0000000..2d9eda2 --- /dev/null +++ b/pkg/logger/slog.go @@ -0,0 +1,102 @@ +package logger + +import ( + "context" + "log/slog" + "os" + "strings" +) + +// slogLogger adapts log/slog to the Logger interface. +// +// slog has been in the standard library since Go 1.21 and is the idiomatic +// answer for structured logging — which is why this is an adapter rather than a +// dependency on a third-party logging package. +// +// The value of structured output is not prettiness: it is that every field is +// separately queryable. "user_id=42" inside a formatted string requires a regex +// to search; a JSON field does not. Once logs are shipped anywhere central, that +// difference decides whether an incident takes minutes or hours. +type slogLogger struct { + l *slog.Logger +} + +// NewJSONLogger writes one JSON object per line at the given level. +// +// This is the production form: machine-parseable, and stable enough that log +// processors can rely on the field names. +func NewJSONLogger(level string) Logger { + return &slogLogger{l: slog.New(slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{ + Level: parseLevel(level), + }))} +} + +// NewTextLogger writes human-readable key=value lines at the given level. +// Suited to a developer's terminal, where a wall of JSON is a hindrance. +func NewTextLogger(level string) Logger { + return &slogLogger{l: slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{ + Level: parseLevel(level), + }))} +} + +// New selects a logger appropriate to the environment: JSON in production, +// readable text elsewhere. +// +// Choosing here rather than at each call site means no code has to know which +// environment it is running in — the only distinction that ever mattered was how +// the bytes are formatted. +func New(production bool, level string) Logger { + if production { + return NewJSONLogger(level) + } + return NewTextLogger(level) +} + +// WithAttrs returns a logger that stamps the given key/value pairs onto every +// record. Use it to bind a request ID once rather than threading it through +// every call. +func (s *slogLogger) WithAttrs(args ...any) Logger { + return &slogLogger{l: s.l.With(args...)} +} + +func (s *slogLogger) Debug(msg string, args ...any) { s.l.Debug(msg, args...) } +func (s *slogLogger) Info(msg string, args ...any) { s.l.Info(msg, args...) } +func (s *slogLogger) Warn(msg string, args ...any) { s.l.Warn(msg, args...) } + +func (s *slogLogger) Error(msg string, err error, args ...any) { + if err != nil { + args = append(args, "error", err.Error()) + } + s.l.Error(msg, args...) +} + +// Fatal logs at error level and terminates the process. +// +// Reserved for unrecoverable startup failures — a process that cannot reach its +// database or load its configuration should stop loudly rather than serve +// requests it cannot fulfil. Never call it from request handling. +func (s *slogLogger) Fatal(msg string, err error, args ...any) { + s.Error(msg, err, args...) + os.Exit(1) +} + +// LogAttrs is available for hot paths where allocation matters; the variadic +// args form allocates a slice per call. +func (s *slogLogger) LogAttrs(ctx context.Context, level slog.Level, msg string, attrs ...slog.Attr) { + s.l.LogAttrs(ctx, level, msg, attrs...) +} + +// parseLevel maps a configured level name onto a slog level, defaulting to Info +// for anything unrecognised. An unreadable LOG_LEVEL should not silence logging. +func parseLevel(level string) slog.Level { + switch strings.ToLower(strings.TrimSpace(level)) { + case "debug": + return slog.LevelDebug + case "warn", "warning": + return slog.LevelWarn + case "error": + return slog.LevelError + default: + return slog.LevelInfo + } +} diff --git a/pkg/mailer/mailer.go b/pkg/mailer/mailer.go index 140f984..9821403 100755 --- a/pkg/mailer/mailer.go +++ b/pkg/mailer/mailer.go @@ -5,8 +5,8 @@ import ( "net/smtp" "strings" - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/pkg/logger" + "github.com/codetheuri/tusk/v2/config" + "github.com/codetheuri/tusk/v2/pkg/logger" ) type MailerService interface { diff --git a/pkg/middleware/bodylimit.go b/pkg/middleware/bodylimit.go new file mode 100644 index 0000000..b168f98 --- /dev/null +++ b/pkg/middleware/bodylimit.go @@ -0,0 +1,23 @@ +package middleware + +import "net/http" + +// MaxBodyBytes caps how much of a request body the server will read. +// +// Without a cap, a single client can stream an unbounded body and the process +// will keep allocating until it is killed — no exploit required, just a slow +// upload that never ends. http.MaxBytesReader stops the read at the limit and +// closes the connection, so the cost of an oversized request is bounded. +// +// The limit applies to the body only. Header size is capped separately by +// http.Server.MaxHeaderBytes. +func MaxBodyBytes(limit int64) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Body != nil { + r.Body = http.MaxBytesReader(w, r.Body, limit) + } + next.ServeHTTP(w, r) + }) + } +} diff --git a/internal/middleware/cors.go b/pkg/middleware/cors.go similarity index 96% rename from internal/middleware/cors.go rename to pkg/middleware/cors.go index 26593d4..d55aaa1 100755 --- a/internal/middleware/cors.go +++ b/pkg/middleware/cors.go @@ -3,7 +3,7 @@ package middleware import ( "net/http" - "github.com/codetheuri/tusk/pkg/logger" + "github.com/codetheuri/tusk/v2/pkg/logger" ) func CORS(allowedOrigins []string, log logger.Logger) func(next http.Handler) http.Handler { @@ -43,4 +43,4 @@ func CORS(allowedOrigins []string, log logger.Logger) func(next http.Handler) ht next.ServeHTTP(w, r) }) } -} \ No newline at end of file +} diff --git a/pkg/middleware/jwt.go b/pkg/middleware/jwt.go new file mode 100644 index 0000000..5a1fc20 --- /dev/null +++ b/pkg/middleware/jwt.go @@ -0,0 +1,223 @@ +// Package middleware provides the global HTTP middleware chain: request IDs, +// logging, panic recovery, CORS, security headers, rate limiting, and +// authentication. +package middleware + +import ( + "context" + "errors" + "fmt" + "net/http" + "strings" + + "github.com/danielgtaylor/huma/v2" + "github.com/golang-jwt/jwt/v5" + "github.com/google/uuid" + + "github.com/codetheuri/tusk/v2/pkg/authz" + "github.com/codetheuri/tusk/v2/pkg/response" + "github.com/codetheuri/tusk/v2/pkg/tenant" +) + +// Claims is the payload carried by a Tusk access token. +// +// IsSuperUser is embedded in the token rather than looked up per request. The +// tradeoff is deliberate: a revoked super-user retains the flag until their +// token expires. That is acceptable for a rarely-changing administrative flag +// and it removes a database round-trip from every authenticated request. Ordinary +// permissions are NOT stored here for exactly the opposite reason — they change +// often, so they are evaluated against the database on each check by pkg/authz. +// +// Role is retained for applications that want a coarse role claim. Tusk's own +// authorization does not consult it; use pkg/authz permissions instead. +// +// TenantID is omitempty and unused by single-tenant applications, which never +// set it and never read it. It belongs in the token rather than in a lookup for +// the same reason the tenant must not come from the request body: it has to be +// something the caller cannot choose. +type Claims struct { + UserID uuid.UUID `json:"user_id"` + TenantID uuid.UUID `json:"tenant_id,omitempty"` + Role string `json:"role,omitempty"` + IsSuperUser bool `json:"is_super_user,omitempty"` + jwt.RegisteredClaims +} + +// ErrNoCredentials indicates the request carried no Authorization header at all. +// It is distinct from a parse or validation failure: an absent credential may be +// legitimate on a public route, whereas a malformed one never is. +var ErrNoCredentials = errors.New("no credentials presented") + +// parseBearerToken extracts and verifies a JWT from an Authorization header. +// +// This is the single place tokens are parsed. Both the net/http and the Huma +// middleware delegate here so the two paths cannot drift apart in how they +// validate signatures or interpret claims. +func parseBearerToken(authHeader, jwtSecret string) (*Claims, error) { + if authHeader == "" { + return nil, ErrNoCredentials + } + + parts := strings.SplitN(authHeader, " ", 2) + if len(parts) != 2 || !strings.EqualFold(parts[0], "bearer") { + return nil, fmt.Errorf("invalid authorization format, expected: Bearer ") + } + + claims := &Claims{} + token, err := jwt.ParseWithClaims(parts[1], claims, func(t *jwt.Token) (any, error) { + // Reject any token not signed with HMAC. Without this check an attacker + // could present an "alg: none" token, or one signed with a public key + // they control, and have it accepted. + if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok { + return nil, jwt.ErrSignatureInvalid + } + return []byte(jwtSecret), nil + }) + if err != nil { + return nil, err + } + if !token.Valid { + return nil, jwt.ErrTokenInvalidClaims + } + + return claims, nil +} + +// subjectFrom builds the authorization Subject carried through the request. +func subjectFrom(claims *Claims) authz.Subject { + return authz.Subject{ + UserID: claims.UserID, + IsSuperUser: claims.IsSuperUser, + } +} + +// Authenticate is net/http middleware that requires a valid token. +// +// Use it to protect a route group mounted outside the Huma API. Requests without +// credentials are rejected — for optional authentication, see HumaAuthenticate. +func Authenticate(jwtSecret string) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + claims, err := parseBearerToken(r.Header.Get("Authorization"), jwtSecret) + if err != nil { + writeUnauthorized(w, err) + return + } + next.ServeHTTP(w, r.WithContext(authz.WithSubject(r.Context(), subjectFrom(claims)))) + }) + } +} + +// HumaAuthenticate resolves identity for every request entering the Huma API. +// +// It distinguishes two cases that the previous implementation conflated: +// +// - No Authorization header: the request proceeds anonymously. Public routes +// work; protected routes are refused by their permission guard with 401. +// - A header that is present but invalid, expired, or malformed: the request +// is refused here with 401. +// +// That distinction matters to clients. Falling through with a bad token produced +// a 403 from the guard, which tells a client "you lack permission" when the truth +// is "your session expired" — so it never knew to refresh the token. +func HumaAuthenticate(api huma.API, jwtSecret string) func(huma.Context, func(huma.Context)) { + return func(ctx huma.Context, next func(huma.Context)) { + claims, err := parseBearerToken(ctx.Header("Authorization"), jwtSecret) + if err != nil { + if errors.Is(err, ErrNoCredentials) { + next(ctx) // anonymous — the guard decides whether that is allowed + return + } + huma.WriteErr(api, ctx, http.StatusUnauthorized, "Invalid or expired token") + return + } + + reqCtx := authz.WithSubject(ctx.Context(), subjectFrom(claims)) + next(huma.WithContext(ctx, reqCtx)) + } +} + +// RequireRole restricts access by the coarse Role claim. +// +// Prefer pkg/authz permissions for anything non-trivial: roles in a token cannot +// be revoked before expiry, and encoding authorization in the credential means +// changing what a user may do requires them to log in again. +func RequireRole(allowedRoles ...string) func(http.Handler) http.Handler { + allowed := make(map[string]bool, len(allowedRoles)) + for _, r := range allowedRoles { + allowed[r] = true + } + + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + role, _ := r.Context().Value(roleKey{}).(string) + if !allowed[role] { + response.WriteJSON(w, http.StatusForbidden, + `{"success":false,"message":"You do not have permission to perform this action"}`) + return + } + next.ServeHTTP(w, r) + }) + } +} + +// roleKey carries the optional Role claim. Unexported so it cannot collide. +type roleKey struct{} + +// WithRole stores the coarse role claim for RequireRole. Applications using +// pkg/authz permissions do not need this. +func WithRole(ctx context.Context, role string) context.Context { + return context.WithValue(ctx, roleKey{}, role) +} + +// GetUserID returns the authenticated user's ID, or the nil UUID when the request +// is anonymous. Callers that need to distinguish the two should use +// authz.SubjectFromContext, which reports presence explicitly. +func GetUserID(ctx context.Context) uuid.UUID { + sub, ok := authz.SubjectFromContext(ctx) + if !ok { + return uuid.Nil + } + return sub.UserID +} + +func writeUnauthorized(w http.ResponseWriter, err error) { + msg := "Invalid or expired token" + if errors.Is(err, ErrNoCredentials) { + msg = "Authorization header is required" + } + response.WriteJSON(w, http.StatusUnauthorized, + fmt.Sprintf(`{"success":false,"message":%q}`, msg)) +} + +// TenantResolver builds a tenant.Resolver that reads the tenant from the token. +// +// A request with no credentials, or with a token carrying no tenant claim, +// resolves to no tenant rather than an error: public routes have to keep +// working, and a request that then touches a tenanted model fails closed at the +// query layer anyway. The failure surfaces where the data is, not at the edge +// where it would also break login. +// +// Wire it in only when the application has tenanted models: +// +// router.Use(tenant.Middleware(middleware.TenantResolver(cfg.JWTSecret))) +func TenantResolver(jwtSecret string) tenant.Resolver { + return func(r *http.Request) (uuid.UUID, error) { + claims, err := parseBearerToken(r.Header.Get("Authorization"), jwtSecret) + if err != nil { + return uuid.Nil, nil + } + return claims.TenantID, nil + } +} + +// HumaTenantResolver is TenantResolver for the Huma API surface. +func HumaTenantResolver(jwtSecret string) tenant.HumaResolver { + return func(ctx huma.Context) (uuid.UUID, error) { + claims, err := parseBearerToken(ctx.Header("Authorization"), jwtSecret) + if err != nil { + return uuid.Nil, nil + } + return claims.TenantID, nil + } +} diff --git a/pkg/middleware/jwt_test.go b/pkg/middleware/jwt_test.go new file mode 100644 index 0000000..9bba206 --- /dev/null +++ b/pkg/middleware/jwt_test.go @@ -0,0 +1,229 @@ +package middleware + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/danielgtaylor/huma/v2" + + "github.com/danielgtaylor/huma/v2/humatest" + "github.com/golang-jwt/jwt/v5" + "github.com/google/uuid" + + "github.com/codetheuri/tusk/v2/pkg/authz" +) + +const testSecret = "test-secret-key-that-is-long-enough-for-hs256" + +// signToken produces a token the middleware should accept. +func signToken(t *testing.T, claims Claims) string { + t.Helper() + if claims.ExpiresAt == nil { + claims.ExpiresAt = jwt.NewNumericDate(time.Now().Add(time.Hour)) + } + signed, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte(testSecret)) + if err != nil { + t.Fatalf("failed to sign test token: %v", err) + } + return signed +} + +func TestParseBearerToken(t *testing.T) { + valid := signToken(t, Claims{UserID: testUserID(7), IsSuperUser: true}) + expired := signToken(t, Claims{ + UserID: testUserID(7), + RegisteredClaims: jwt.RegisteredClaims{ExpiresAt: jwt.NewNumericDate(time.Now().Add(-time.Hour))}, + }) + + // A token signed with a different key must be rejected even though it is + // structurally valid — this is the signature check doing its job. + forged, err := jwt.NewWithClaims(jwt.SigningMethodHS256, Claims{ + UserID: testUserID(99), + RegisteredClaims: jwt.RegisteredClaims{ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour))}, + }).SignedString([]byte("a-completely-different-signing-key-value")) + if err != nil { + t.Fatalf("failed to sign forged token: %v", err) + } + + tests := []struct { + name string + header string + wantErr bool + wantNoCred bool + wantUserID uuid.UUID + }{ + {name: "valid bearer token", header: "Bearer " + valid, wantUserID: testUserID(7)}, + {name: "case-insensitive scheme", header: "bearer " + valid, wantUserID: testUserID(7)}, + {name: "empty header reports no credentials", header: "", wantErr: true, wantNoCred: true}, + {name: "missing scheme", header: valid, wantErr: true}, + {name: "wrong scheme", header: "Basic " + valid, wantErr: true}, + {name: "expired token", header: "Bearer " + expired, wantErr: true}, + {name: "token signed with another key", header: "Bearer " + forged, wantErr: true}, + {name: "garbage token", header: "Bearer not-a-jwt", wantErr: true}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + claims, err := parseBearerToken(tc.header, testSecret) + + if tc.wantErr { + if err == nil { + t.Fatal("expected an error, got nil") + } + // The distinction between "no credentials" and "bad credentials" + // is what lets the Huma middleware answer 401 instead of falling + // through to a misleading 403. + if got := errors.Is(err, ErrNoCredentials); got != tc.wantNoCred { + t.Errorf("errors.Is(err, ErrNoCredentials) = %v, want %v (err: %v)", got, tc.wantNoCred, err) + } + return + } + + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if claims.UserID != tc.wantUserID { + t.Errorf("UserID = %s, want %s", claims.UserID, tc.wantUserID) + } + }) + } +} + +// IsSuperUser now travels in the token instead of being fetched per request. +func TestParseBearerToken_CarriesSuperUserFlag(t *testing.T) { + super := signToken(t, Claims{UserID: testUserID(1), IsSuperUser: true}) + ordinary := signToken(t, Claims{UserID: testUserID(2)}) + + claims, err := parseBearerToken("Bearer "+super, testSecret) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if !claims.IsSuperUser { + t.Error("expected IsSuperUser true for a super-user token") + } + + claims, err = parseBearerToken("Bearer "+ordinary, testSecret) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if claims.IsSuperUser { + t.Error("expected IsSuperUser false when the claim is absent — it must fail closed") + } +} + +func TestAuthenticate_RejectsAnonymous(t *testing.T) { + handlerRan := false + h := Authenticate(testSecret)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + handlerRan = true + })) + + rr := httptest.NewRecorder() + h.ServeHTTP(rr, httptest.NewRequest(http.MethodGet, "/", nil)) + + if rr.Code != http.StatusUnauthorized { + t.Errorf("status = %d, want 401", rr.Code) + } + if handlerRan { + t.Error("handler ran despite missing credentials") + } +} + +func TestAuthenticate_PopulatesSubject(t *testing.T) { + var got authz.Subject + var ok bool + + h := Authenticate(testSecret)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + got, ok = authz.SubjectFromContext(r.Context()) + })) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.Header.Set("Authorization", "Bearer "+signToken(t, Claims{UserID: testUserID(42), IsSuperUser: true})) + h.ServeHTTP(httptest.NewRecorder(), req) + + if !ok { + t.Fatal("expected a Subject in context") + } + if got.UserID != testUserID(42) || !got.IsSuperUser { + t.Errorf("Subject = %+v, want {UserID:%s IsSuperUser:true}", got, testUserID(42)) + } +} + +// This is the H3 regression guard. An absent credential is legitimate on a public +// route; a malformed one never is, and conflating them told clients "forbidden" +// when the truth was "your session expired". +func TestHumaAuthenticate_DistinguishesAbsentFromInvalid(t *testing.T) { + t.Run("no token proceeds anonymously", func(t *testing.T) { + reached := false + api := newProbeAPI(t, func() { reached = true }) + + resp := api.Get("/probe") + if resp.Code != http.StatusOK { + t.Errorf("status = %d, want 200", resp.Code) + } + if !reached { + t.Error("expected the handler to run for an anonymous request") + } + }) + + t.Run("invalid token is refused with 401", func(t *testing.T) { + reached := false + api := newProbeAPI(t, func() { reached = true }) + + resp := api.Get("/probe", "Authorization: Bearer not-a-valid-token") + if resp.Code != http.StatusUnauthorized { + t.Errorf("status = %d, want 401", resp.Code) + } + if reached { + t.Error("handler ran despite an invalid token") + } + }) + + t.Run("expired token is refused with 401", func(t *testing.T) { + api := newProbeAPI(t, func() {}) + + expired := signToken(t, Claims{ + UserID: testUserID(1), + RegisteredClaims: jwt.RegisteredClaims{ExpiresAt: jwt.NewNumericDate(time.Now().Add(-time.Minute))}, + }) + + resp := api.Get("/probe", "Authorization: Bearer "+expired) + if resp.Code != http.StatusUnauthorized { + t.Errorf("status = %d, want 401 (an expired session must be distinguishable from a permission failure)", resp.Code) + } + }) +} + +// probeOutput is the response of the throwaway endpoint used to observe whether +// a request reached its handler. +type probeOutput struct { + Body struct { + OK bool `json:"ok"` + } +} + +// newProbeAPI builds a Huma API with the authentication middleware installed and +// a single unguarded endpoint, so a test can distinguish "refused at the edge" +// from "reached the handler". +func newProbeAPI(t *testing.T, onReach func()) humatest.TestAPI { + t.Helper() + _, api := humatest.New(t) + api.UseMiddleware(HumaAuthenticate(api, testSecret)) + huma.Get(api, "/probe", func(ctx context.Context, _ *struct{}) (*probeOutput, error) { + onReach() + return &probeOutput{}, nil + }) + return api +} + +// testUserID builds a deterministic identifier from a small integer, so tests can +// keep using readable IDs instead of hard-coded UUID literals. The value only has +// to be stable and distinct — it never reaches a database. +func testUserID(n int) uuid.UUID { + var u uuid.UUID + u[15] = byte(n) + return u +} diff --git a/internal/middleware/logger.go b/pkg/middleware/logger.go similarity index 95% rename from internal/middleware/logger.go rename to pkg/middleware/logger.go index 96b8b64..34ba114 100755 --- a/internal/middleware/logger.go +++ b/pkg/middleware/logger.go @@ -4,7 +4,7 @@ import ( "net/http" "time" - "github.com/codetheuri/tusk/pkg/logger" + "github.com/codetheuri/tusk/v2/pkg/logger" ) func Logger(log logger.Logger) func(next http.Handler) http.Handler { diff --git a/internal/middleware/middleware_test.go b/pkg/middleware/middleware_test.go similarity index 96% rename from internal/middleware/middleware_test.go rename to pkg/middleware/middleware_test.go index fcdb471..c8c912b 100644 --- a/internal/middleware/middleware_test.go +++ b/pkg/middleware/middleware_test.go @@ -108,7 +108,7 @@ func TestRecoveryMiddleware(t *testing.T) { func TestJWTMiddleware_Authenticate(t *testing.T) { secret := "super-secret-key" token := jwt.NewWithClaims(jwt.SigningMethodHS256, &Claims{ - UserID: 1, + UserID: testUserID(1), Role: "admin", RegisteredClaims: jwt.RegisteredClaims{ ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)), @@ -118,8 +118,8 @@ func TestJWTMiddleware_Authenticate(t *testing.T) { handler := Authenticate(secret)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { uid := GetUserID(r.Context()) - if uid != 1 { - t.Errorf("expected user_id 1 in context, got %d", uid) + if uid != testUserID(1) { + t.Errorf("expected user_id %s in context, got %s", testUserID(1), uid) } w.WriteHeader(http.StatusOK) })) diff --git a/internal/middleware/ratelimit.go b/pkg/middleware/ratelimit.go similarity index 98% rename from internal/middleware/ratelimit.go rename to pkg/middleware/ratelimit.go index e946535..9ce85f5 100644 --- a/internal/middleware/ratelimit.go +++ b/pkg/middleware/ratelimit.go @@ -5,7 +5,7 @@ import ( "sync" "time" - "github.com/codetheuri/tusk/pkg/logger" + "github.com/codetheuri/tusk/v2/pkg/logger" ) type clientVisitor struct { diff --git a/internal/middleware/ratelimit_test.go b/pkg/middleware/ratelimit_test.go similarity index 100% rename from internal/middleware/ratelimit_test.go rename to pkg/middleware/ratelimit_test.go diff --git a/internal/middleware/recovery.go b/pkg/middleware/recovery.go similarity index 90% rename from internal/middleware/recovery.go rename to pkg/middleware/recovery.go index b9f81a5..883883c 100755 --- a/internal/middleware/recovery.go +++ b/pkg/middleware/recovery.go @@ -5,9 +5,7 @@ import ( "net/http" "runtime/debug" - - "github.com/codetheuri/tusk/pkg/logger" - + "github.com/codetheuri/tusk/v2/pkg/logger" ) // recover from panics and return a 500 error @@ -18,7 +16,7 @@ func Recovery(log logger.Logger) func(next http.Handler) http.Handler { if rcvErr := recover(); rcvErr != nil { //log thr panic var actualErr error - if e, ok := rcvErr.(error); ok { + if e, ok := rcvErr.(error); ok { actualErr = e } else { actualErr = fmt.Errorf("%v", rcvErr) diff --git a/internal/middleware/requestid.go b/pkg/middleware/requestid.go similarity index 56% rename from internal/middleware/requestid.go rename to pkg/middleware/requestid.go index 5cd1e9b..c83d454 100755 --- a/internal/middleware/requestid.go +++ b/pkg/middleware/requestid.go @@ -7,6 +7,15 @@ import ( "github.com/google/uuid" ) +// contextKey is a private named type for context keys owned by this package. +// +// The distinction that matters: a *named* type like this is safe, because +// context lookups compare the key's type as well as its value — no other package +// can construct a middleware.contextKey. An *untyped* string literal such as +// ctx.Value("user_id") is not safe, because any package writing that same literal +// collides silently. Both look similar at the call site; only one of them works. +type contextKey string + const ( RequestIDKey contextKey = "requestID" ) diff --git a/internal/middleware/security_headers.go b/pkg/middleware/security_headers.go similarity index 99% rename from internal/middleware/security_headers.go rename to pkg/middleware/security_headers.go index 02461b0..d9b163d 100755 --- a/internal/middleware/security_headers.go +++ b/pkg/middleware/security_headers.go @@ -9,15 +9,15 @@ func SecurityHeaders(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { // Prevent browsers from guessing (sniffing) MIME types, reducing XSS risk. w.Header().Set("X-Content-Type-Options", "nosniff") - + // Prevent clickjacking attacks by forbidding iframe embedding. w.Header().Set("X-Frame-Options", "DENY") w.Header().Set("Strict-Transport-Security", "max-age=31536000; includeSubDomains; preload") - + // Enable XSS protection filter in browsers. w.Header().Set("X-XSS-Protection", "1; mode=block") next.ServeHTTP(w, r) }) -} \ No newline at end of file +} diff --git a/pkg/query/builder.go b/pkg/query/builder.go index 2eec0c4..b8d723b 100644 --- a/pkg/query/builder.go +++ b/pkg/query/builder.go @@ -42,13 +42,18 @@ func Apply(db *gorm.DB, q Query, cfg Config) *gorm.DB { } // 2. Apply Whitelisted Multi-Column Search + // + // ILIKE, not LIKE: Tusk targets PostgreSQL exclusively (see database.migrationDir), + // where LIKE is case-sensitive. A search for "wanjiku" that silently misses + // "Wanjiku" is not a partial match with a rough edge, it is broken search — + // nobody searching a name or an email types the exact original casing. if q.Search != "" && len(cfg.AllowedSearches) > 0 { var searchConditions []string var searchArgs []interface{} pattern := "%" + q.Search + "%" for _, col := range cfg.AllowedSearches { - searchConditions = append(searchConditions, fmt.Sprintf("%s LIKE ?", col)) + searchConditions = append(searchConditions, fmt.Sprintf("%s ILIKE ?", col)) searchArgs = append(searchArgs, pattern) } diff --git a/pkg/query/query_test.go b/pkg/query/query_test.go new file mode 100644 index 0000000..dfc9363 --- /dev/null +++ b/pkg/query/query_test.go @@ -0,0 +1,146 @@ +package query_test + +import ( + "context" + "testing" + + "github.com/google/uuid" + "gorm.io/gorm" + + "github.com/codetheuri/tusk/v2/pkg/query" + "github.com/codetheuri/tusk/v2/pkg/testdb" +) + +// These run against real PostgreSQL because the defect they cover — LIKE's +// case sensitivity — is specific to the database engine. A fake would not +// reproduce it: it would simply do whatever the test expected. + +type queryTestItem struct { + ID uuid.UUID `gorm:"type:uuid;primaryKey"` + Name string + CreatedAt int64 +} + +func (queryTestItem) TableName() string { return "query_pkg_test_items" } + +func setup(t *testing.T) *gorm.DB { + t.Helper() + db := testdb.Connect(t) + if err := db.AutoMigrate(&queryTestItem{}); err != nil { + t.Fatalf("automigrate: %v", err) + } + t.Cleanup(func() { db.Exec("DROP TABLE IF EXISTS query_pkg_test_items") }) + + seed := []queryTestItem{ + {ID: uuid.New(), Name: "Wanjiku Njoroge", CreatedAt: 1}, + {ID: uuid.New(), Name: "John Otieno", CreatedAt: 2}, + {ID: uuid.New(), Name: "Wanjiru Kamau", CreatedAt: 3}, + } + if err := db.Create(&seed).Error; err != nil { + t.Fatalf("seed: %v", err) + } + return db +} + +var cfg = query.Config{ + DefaultPerPage: 20, + MaxPerPage: 100, + AllowedSearches: []string{"name"}, + AllowedSorts: map[string]string{"name": "name", "created_at": "created_at"}, +} + +// TestApply_SearchIsCaseInsensitive is the regression test for the LIKE/ILIKE +// defect: a lowercase search must find a capitalised name, on PostgreSQL, which +// is the only database Tusk targets. +func TestApply_SearchIsCaseInsensitive(t *testing.T) { + db := setup(t) + + var got []queryTestItem + err := query.Apply(db, query.Query{Search: "wanj"}, cfg).Find(&got).Error + if err != nil { + t.Fatalf("apply: %v", err) + } + if len(got) != 2 { + t.Fatalf("got %d rows, want 2 (Wanjiku, Wanjiru)", len(got)) + } +} + +func TestApply_SearchIgnoresUnlistedColumns(t *testing.T) { + db := setup(t) + + // AllowedSearches lists only "name", so this must not search anything else + // even if the query struct carried a filter naming another column. + restrictedCfg := query.Config{AllowedSearches: []string{"name"}} + var got []queryTestItem + if err := query.Apply(db, query.Query{Search: "otieno"}, restrictedCfg).Find(&got).Error; err != nil { + t.Fatalf("apply: %v", err) + } + if len(got) != 1 || got[0].Name != "John Otieno" { + t.Fatalf("got %+v, want exactly John Otieno", got) + } +} + +func TestPaginate_ReturnsCorrectMetaAcrossPages(t *testing.T) { + db := setup(t) + ctx := context.Background() + + first, meta, err := query.Paginate[queryTestItem](ctx, db, query.Query{Page: 1, PerPage: 2}, cfg) + if err != nil { + t.Fatalf("page 1: %v", err) + } + if len(first) != 2 { + t.Fatalf("page 1 has %d items, want 2", len(first)) + } + if meta.Total != 3 || meta.TotalPages != 2 || !meta.HasNext || meta.HasPrevious { + t.Errorf("page 1 meta = %+v", meta) + } + + second, meta2, err := query.Paginate[queryTestItem](ctx, db, query.Query{Page: 2, PerPage: 2}, cfg) + if err != nil { + t.Fatalf("page 2: %v", err) + } + if len(second) != 1 { + t.Fatalf("page 2 has %d items, want 1", len(second)) + } + if meta2.HasNext || !meta2.HasPrevious { + t.Errorf("page 2 meta = %+v", meta2) + } + + seen := map[uuid.UUID]bool{} + for _, item := range append(first, second...) { + if seen[item.ID] { + t.Errorf("item %s returned on more than one page", item.ID) + } + seen[item.ID] = true + } + if len(seen) != 3 { + t.Errorf("saw %d distinct items across both pages, want 3", len(seen)) + } +} + +func TestPaginate_PerPageIsBoundedByMaxPerPage(t *testing.T) { + db := setup(t) + bounded := cfg + bounded.MaxPerPage = 1 + + items, meta, err := query.Paginate[queryTestItem](context.Background(), db, query.Query{PerPage: 1000}, bounded) + if err != nil { + t.Fatalf("paginate: %v", err) + } + if len(items) != 1 || meta.PerPage != 1 { + t.Errorf("got %d items with per_page=%d, want 1 item, per_page=1", len(items), meta.PerPage) + } +} + +func TestApply_SortsByAllowedColumnOnly(t *testing.T) { + db := setup(t) + + var got []queryTestItem + q := query.Query{Sorts: []query.Sort{{Field: "created_at", Order: query.SortDesc}}} + if err := query.Apply(db, q, cfg).Find(&got).Error; err != nil { + t.Fatalf("apply: %v", err) + } + if len(got) != 3 || got[0].Name != "Wanjiru Kamau" { + t.Fatalf("expected newest-first, got %+v", got) + } +} diff --git a/pkg/session/session.go b/pkg/session/session.go new file mode 100644 index 0000000..d2ef985 --- /dev/null +++ b/pkg/session/session.go @@ -0,0 +1,381 @@ +// Package session provides database-backed cookie sessions for server-rendered +// pages, alongside the JWT authentication Tusk uses for its API. +// +// # Why not just use the JWT +// +// A browser session and an API token want opposite things. A JWT is stateless +// and therefore cannot be revoked before it expires — fine for a mobile client +// holding a short-lived access token, wrong for an operator console where +// "sign out everywhere" and "revoke that person's access now" are requirements. +// Sessions are a database row precisely so that deleting the row ends access. +// +// # What the browser holds +// +// A 256-bit random token, in an HttpOnly cookie. The database stores only its +// SHA-256 digest, so a leaked backup does not hand over live sessions. +// +// SHA-256 rather than bcrypt, deliberately, and for the opposite reason to +// passwords: a password is low-entropy and guessable, so hashing it must be +// slow. This token is 256 bits of randomness — unguessable regardless of hash +// speed — and it is verified on every single request, so a deliberately slow +// hash would be a denial-of-service surface rather than a protection. +// +// # Wiring +// +// mgr, _ := session.New(store, session.Options{ +// Kind: "console", CookieName: "console_session", +// Path: "/console", Secure: cfg.IsProduction(), +// }) +// r.Route("/console", func(r chi.Router) { +// r.Use(mgr.Middleware()) // loads a session if there is one +// r.Get("/login", showLogin) // reads session.FromContext to skip the form +// r.Group(func(r chi.Router) { +// r.Use(mgr.Require("/console/login")) +// r.Get("/dashboard", dashboard) +// }) +// }) +// +// Middleware loads without requiring, and Require refuses. Keeping them separate +// is what stops the classic redirect loop: a login page that redirects on the +// mere presence of a cookie will bounce a visitor holding a stale one between +// /login and /dashboard forever. Here the login page asks whether a *valid* +// session exists, and the middleware has already cleared the cookie if not. +package session + +import ( + "context" + "crypto/rand" + "crypto/sha256" + "encoding/hex" + "fmt" + "net" + "net/http" + "time" + + "github.com/google/uuid" +) + +// Options configures a Manager. +type Options struct { + // Kind names the audience these sessions belong to, for example "console". + // Required: it is what keeps two session audiences in one database apart. + Kind string + + // CookieName defaults to "tusk_session". Give each Kind its own name. + CookieName string + + // Path scopes the cookie. Defaults to "/". Setting it to the console's + // mount point keeps the cookie off every other request. + Path string + + // Domain is normally left empty, which scopes the cookie to the exact host + // that set it — the safer default. + Domain string + + // Secure restricts the cookie to HTTPS. It must be true in production; drive + // it from configuration rather than hard-coding either value. + Secure bool + + // SameSite defaults to http.SameSiteLaxMode, which keeps the cookie off + // cross-site POSTs — the main CSRF vector — while still allowing ordinary + // inbound links to work. + SameSite http.SameSite + + // Lifetime is the absolute maximum age of a session. Defaults to 8 hours. + // It is not extended by activity: a stolen session cannot be kept alive + // forever by using it. + Lifetime time.Duration + + // IdleTimeout ends a session that has gone unused. Zero disables it. + IdleTimeout time.Duration +} + +// Manager issues, validates and ends sessions. +type Manager struct { + store Store + opts Options +} + +// New validates the options and returns a Manager. +func New(store Store, opts Options) (*Manager, error) { + if store == nil { + return nil, fmt.Errorf("session: a Store is required") + } + if opts.Kind == "" { + return nil, fmt.Errorf("session: Options.Kind is required, so that sessions from different audiences cannot be confused") + } + if opts.CookieName == "" { + opts.CookieName = "tusk_session" + } + if opts.Path == "" { + opts.Path = "/" + } + if opts.SameSite == 0 { + opts.SameSite = http.SameSiteLaxMode + } + if opts.Lifetime == 0 { + opts.Lifetime = 8 * time.Hour + } + if opts.IdleTimeout < 0 { + return nil, fmt.Errorf("session: IdleTimeout cannot be negative") + } + return &Manager{store: store, opts: opts}, nil +} + +// contextKey is unexported so no other package can write or overwrite the value. +type contextKey struct{} + +// FromContext returns the session loaded by Middleware, if any. +func FromContext(ctx context.Context) (*Session, bool) { + s, ok := ctx.Value(contextKey{}).(*Session) + return s, ok +} + +// Start creates a session for a subject and sets the cookie. +// +// Call it only after credentials have been verified. It always mints a fresh +// token rather than adopting any the request arrived with, which is what makes +// session fixation impossible: an attacker who plants a cookie value in a +// victim's browser does not learn the one issued at login. +func (m *Manager) Start(w http.ResponseWriter, r *http.Request, subjectID uuid.UUID) (*Session, error) { + token, hash, err := newToken() + if err != nil { + return nil, err + } + + now := time.Now() + sess := &Session{ + TokenHash: hash, + SubjectKind: m.opts.Kind, + SubjectID: subjectID, + IP: clientIP(r), + UserAgent: truncate(r.UserAgent(), 255), + CreatedAt: now, + LastSeenAt: now, + ExpiresAt: now.Add(m.opts.Lifetime), + } + if err := m.store.Create(r.Context(), sess); err != nil { + return nil, fmt.Errorf("session: creating: %w", err) + } + + m.setCookie(w, token, sess.ExpiresAt) + return sess, nil +} + +// Rotate issues a new token for the current session and invalidates the old one. +// +// Call it whenever what the session is allowed to do changes — an elevation to +// administrator, a password change, a switch of acting account. If a token had +// leaked before that point, rotation is what stops it inheriting the new +// privileges. +func (m *Manager) Rotate(w http.ResponseWriter, r *http.Request) (*Session, error) { + current, ok := FromContext(r.Context()) + if !ok { + return nil, ErrNotFound + } + + token, hash, err := newToken() + if err != nil { + return nil, err + } + + now := time.Now() + next := &Session{ + TokenHash: hash, + SubjectKind: current.SubjectKind, + SubjectID: current.SubjectID, + IP: clientIP(r), + UserAgent: truncate(r.UserAgent(), 255), + CreatedAt: now, + LastSeenAt: now, + + // The absolute expiry is inherited, not restarted. Rotation is a + // security measure; letting it extend the lifetime would turn it into a + // way to hold a session open indefinitely. + ExpiresAt: current.ExpiresAt, + } + if err := m.store.Create(r.Context(), next); err != nil { + return nil, fmt.Errorf("session: rotating: %w", err) + } + if err := m.store.Delete(r.Context(), current.ID); err != nil { + return nil, fmt.Errorf("session: removing the rotated session: %w", err) + } + + m.setCookie(w, token, next.ExpiresAt) + return next, nil +} + +// Destroy ends the current session and clears the cookie. +func (m *Manager) Destroy(w http.ResponseWriter, r *http.Request) error { + defer m.clearCookie(w) + + if current, ok := FromContext(r.Context()); ok { + if err := m.store.Delete(r.Context(), current.ID); err != nil { + return fmt.Errorf("session: destroying: %w", err) + } + return nil + } + + // No session in context: the cookie may still name one, if Middleware was + // not mounted on this route. Delete by token so logout is not silently a + // no-op. + cookie, err := r.Cookie(m.opts.CookieName) + if err != nil { + return nil + } + sess, err := m.store.Find(r.Context(), m.opts.Kind, hashToken(cookie.Value), time.Now()) + if err != nil { + return nil // already gone or expired; the cookie is cleared regardless + } + return m.store.Delete(r.Context(), sess.ID) +} + +// DestroyAll ends every session belonging to a subject. +// +// The right response to a password reset or a compromise report: whoever else +// was holding a token is signed out, not merely prevented from signing in again. +func (m *Manager) DestroyAll(ctx context.Context, subjectID uuid.UUID) error { + return m.store.DeleteBySubject(ctx, m.opts.Kind, subjectID) +} + +// DeleteExpired removes rows whose expiry has passed. Run it periodically; +// nothing else deletes them, and expired rows are still returned by nothing but +// still occupy the table. +func (m *Manager) DeleteExpired(ctx context.Context) (int64, error) { + return m.store.DeleteExpired(ctx, time.Now()) +} + +// Middleware loads a valid session into the request context. +// +// It does not reject anything. A request without a session proceeds, which is +// what public pages under the same mount — the login form itself — need. +func (m *Manager) Middleware() func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + cookie, err := r.Cookie(m.opts.CookieName) + if err != nil || cookie.Value == "" { + next.ServeHTTP(w, r) + return + } + + now := time.Now() + sess, err := m.store.Find(r.Context(), m.opts.Kind, hashToken(cookie.Value), now) + if err != nil { + // Clear it. Leaving a dead cookie in place means every later + // request repeats this lookup, and any page that tests for the + // cookie rather than the session will loop. + m.clearCookie(w) + next.ServeHTTP(w, r) + return + } + + if m.opts.IdleTimeout > 0 && now.Sub(sess.LastSeenAt) > m.opts.IdleTimeout { + _ = m.store.Delete(r.Context(), sess.ID) + m.clearCookie(w) + next.ServeHTTP(w, r) + return + } + + m.touch(r.Context(), sess, now) + next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), contextKey{}, sess))) + }) + } +} + +// Require refuses requests that have no session, redirecting to loginPath. +// +// Mount it inside Middleware, on the routes that need protection. +func (m *Manager) Require(loginPath string) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if _, ok := FromContext(r.Context()); !ok { + http.Redirect(w, r, loginPath, http.StatusSeeOther) + return + } + next.ServeHTTP(w, r) + }) + } +} + +// touch records activity, but not on every request. +// +// An idle timeout needs a last-seen time, and updating it per request turns +// every page view into a database write. Writing only once the recorded time is +// meaningfully stale keeps the timeout accurate to within a small fraction of +// itself, which is all it needs to be. +func (m *Manager) touch(ctx context.Context, sess *Session, now time.Time) { + if m.opts.IdleTimeout == 0 { + return + } + if now.Sub(sess.LastSeenAt) < m.opts.IdleTimeout/10 { + return + } + if err := m.store.Touch(ctx, sess.ID, now); err == nil { + sess.LastSeenAt = now + } +} + +func (m *Manager) setCookie(w http.ResponseWriter, token string, expires time.Time) { + http.SetCookie(w, &http.Cookie{ + Name: m.opts.CookieName, + Value: token, + Path: m.opts.Path, + Domain: m.opts.Domain, + Expires: expires, + HttpOnly: true, // unreadable from JavaScript, so XSS cannot exfiltrate it + Secure: m.opts.Secure, + SameSite: m.opts.SameSite, + }) +} + +func (m *Manager) clearCookie(w http.ResponseWriter) { + http.SetCookie(w, &http.Cookie{ + Name: m.opts.CookieName, + Value: "", + Path: m.opts.Path, // must match, or the browser keeps the original + Domain: m.opts.Domain, + MaxAge: -1, + + // The security attributes have to match too. A clearing cookie that + // drops Secure is rejected outright by some browsers on HTTPS. + HttpOnly: true, + Secure: m.opts.Secure, + SameSite: m.opts.SameSite, + }) +} + +// newToken returns a fresh session token and the digest to store for it. +func newToken() (token, hash string, err error) { + buf := make([]byte, 32) + if _, err := rand.Read(buf); err != nil { + return "", "", fmt.Errorf("session: generating a token: %w", err) + } + token = hex.EncodeToString(buf) + return token, hashToken(token), nil +} + +func hashToken(token string) string { + sum := sha256.Sum256([]byte(token)) + return hex.EncodeToString(sum[:]) +} + +// clientIP records where a session was used from, for the benefit of an +// operator reviewing active sessions. +// +// It is recorded, never trusted: behind a proxy RemoteAddr is the proxy, and +// X-Forwarded-For is caller-supplied and trivially forged. Nothing here makes an +// access decision from it. +func clientIP(r *http.Request) string { + host, _, err := net.SplitHostPort(r.RemoteAddr) + if err != nil { + return truncate(r.RemoteAddr, 45) + } + return truncate(host, 45) +} + +func truncate(s string, n int) string { + if len(s) <= n { + return s + } + return s[:n] +} diff --git a/pkg/session/session_test.go b/pkg/session/session_test.go new file mode 100644 index 0000000..52853a4 --- /dev/null +++ b/pkg/session/session_test.go @@ -0,0 +1,411 @@ +package session_test + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/google/uuid" + "gorm.io/gorm" + + "github.com/codetheuri/tusk/v2/pkg/session" + "github.com/codetheuri/tusk/v2/pkg/testdb" +) + +// These run against real PostgreSQL. Sessions are a database row whose whole +// purpose is that deleting it ends access, so the assertions worth making are +// about what a second request finds in the table — not about a fake. + +const loginPath = "/console/login" + +func fixture(t *testing.T, opts session.Options) (*session.Manager, *session.GormStore, *gorm.DB) { + t.Helper() + + db := testdb.Connect(t) + store := session.NewGormStore(db) + + if opts.Kind == "" { + opts.Kind = "console" + } + mgr, err := session.New(store, opts) + if err != nil { + t.Fatalf("new manager: %v", err) + } + return mgr, store, db +} + +// login performs a Start and returns the cookie the browser would hold. +func login(t *testing.T, mgr *session.Manager, subject uuid.UUID) *http.Cookie { + t.Helper() + + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodPost, "/console/login", nil) + req.RemoteAddr = "203.0.113.7:51234" + req.Header.Set("User-Agent", "test-agent") + + if _, err := mgr.Start(rec, req, subject); err != nil { + t.Fatalf("start: %v", err) + } + + cookies := rec.Result().Cookies() + if len(cookies) != 1 { + t.Fatalf("expected one cookie, got %d", len(cookies)) + } + return cookies[0] +} + +// visit sends a request carrying the cookie through Middleware, and reports what +// the handler saw. +func visit(mgr *session.Manager, cookie *http.Cookie) (*httptest.ResponseRecorder, *session.Session, bool) { + var ( + got *session.Session + present bool + ) + h := mgr.Middleware()(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + got, present = session.FromContext(r.Context()) + })) + + req := httptest.NewRequest(http.MethodGet, "/console/dashboard", nil) + if cookie != nil { + req.AddCookie(cookie) + } + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + return rec, got, present +} + +func TestStart_IssuesAnUsableSession(t *testing.T) { + mgr, _, _ := fixture(t, session.Options{}) + subject := uuid.New() + + cookie := login(t, mgr, subject) + + if !cookie.HttpOnly { + t.Error("cookie is not HttpOnly, so JavaScript can read it") + } + if cookie.SameSite != http.SameSiteLaxMode { + t.Errorf("SameSite is %v, want Lax", cookie.SameSite) + } + if len(cookie.Value) != 64 { + t.Errorf("token is %d characters, want 64 (256 bits, hex)", len(cookie.Value)) + } + + _, sess, ok := visit(mgr, cookie) + if !ok { + t.Fatal("the session was not loaded on the next request") + } + if sess.SubjectID != subject { + t.Errorf("session belongs to %s, want %s", sess.SubjectID, subject) + } + if sess.IP != "203.0.113.7" { + t.Errorf("recorded IP %q", sess.IP) + } +} + +// TestStart_StoresOnlyTheDigest is the property that makes a leaked backup +// survivable. +func TestStart_StoresOnlyTheDigest(t *testing.T) { + mgr, _, db := fixture(t, session.Options{}) + + cookie := login(t, mgr, uuid.New()) + + var stored session.Session + if err := db.First(&stored).Error; err != nil { + t.Fatalf("read back: %v", err) + } + if stored.TokenHash == cookie.Value { + t.Fatal("the raw token was stored; a database dump would hand over live sessions") + } + + var matches int64 + db.Model(&session.Session{}).Where("token_hash = ?", cookie.Value).Count(&matches) + if matches != 0 { + t.Error("the raw token is present in the table") + } +} + +func TestMiddleware_IgnoresAnUnknownTokenAndClearsTheCookie(t *testing.T) { + mgr, _, _ := fixture(t, session.Options{}) + + rec, _, ok := visit(mgr, &http.Cookie{Name: "tusk_session", Value: "not-a-real-token"}) + + if ok { + t.Error("an unknown token produced a session") + } + cookies := rec.Result().Cookies() + if len(cookies) != 1 || cookies[0].MaxAge != -1 { + t.Errorf("the dead cookie was not cleared: %#v", cookies) + } +} + +// TestMiddleware_RejectsAnExpiredSession checks the absolute lifetime. The row +// is filtered out by the query, not checked afterwards in Go. +func TestMiddleware_RejectsAnExpiredSession(t *testing.T) { + mgr, store, db := fixture(t, session.Options{}) + + cookie := login(t, mgr, uuid.New()) + + // Age it past its expiry. + if err := db.Model(&session.Session{}). + Where("1 = 1"). + Update("expires_at", time.Now().Add(-time.Minute)).Error; err != nil { + t.Fatalf("age the session: %v", err) + } + + if _, _, ok := visit(mgr, cookie); ok { + t.Error("an expired session was accepted") + } + + // And Find agrees, so nothing downstream can resurrect it. + if _, err := store.Find(context.Background(), "console", cookie.Value, time.Now()); err == nil { + t.Error("Find returned an expired session") + } +} + +func TestMiddleware_RejectsAnIdleSession(t *testing.T) { + mgr, _, db := fixture(t, session.Options{IdleTimeout: time.Hour}) + + cookie := login(t, mgr, uuid.New()) + + if err := db.Model(&session.Session{}). + Where("1 = 1"). + Update("last_seen_at", time.Now().Add(-2*time.Hour)).Error; err != nil { + t.Fatalf("age the session: %v", err) + } + + if _, _, ok := visit(mgr, cookie); ok { + t.Error("a session idle beyond the timeout was accepted") + } + + var remaining int64 + db.Model(&session.Session{}).Count(&remaining) + if remaining != 0 { + t.Error("the idle session was left in the table") + } +} + +// TestKind_SeparatesAudiences: two managers over one table must not see each +// other's sessions, or an operator cookie would authenticate an ordinary user +// route and vice versa. +func TestKind_SeparatesAudiences(t *testing.T) { + db := testdb.Connect(t) + store := session.NewGormStore(db) + + console, err := session.New(store, session.Options{Kind: "console", CookieName: "console_session"}) + if err != nil { + t.Fatalf("console manager: %v", err) + } + portal, err := session.New(store, session.Options{Kind: "portal", CookieName: "portal_session"}) + if err != nil { + t.Fatalf("portal manager: %v", err) + } + + cookie := login(t, console, uuid.New()) + + if _, _, ok := visit(console, cookie); !ok { + t.Fatal("the console session did not load for its own manager") + } + + // Same token value, presented under the portal's cookie name. + if _, _, ok := visit(portal, &http.Cookie{Name: "portal_session", Value: cookie.Value}); ok { + t.Error("a console token authenticated a portal session") + } +} + +// TestRotate_InvalidatesTheOldToken covers privilege change. The old cookie must +// stop working the moment the new one is issued. +func TestRotate_InvalidatesTheOldToken(t *testing.T) { + mgr, _, _ := fixture(t, session.Options{}) + subject := uuid.New() + + old := login(t, mgr, subject) + + var rotated *http.Cookie + h := mgr.Middleware()(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if _, err := mgr.Rotate(w, r); err != nil { + t.Errorf("rotate: %v", err) + } + })) + req := httptest.NewRequest(http.MethodPost, "/console/elevate", nil) + req.AddCookie(old) + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + + for _, c := range rec.Result().Cookies() { + if c.Name == "tusk_session" && c.Value != "" { + rotated = c + } + } + if rotated == nil { + t.Fatal("rotation issued no new cookie") + } + if rotated.Value == old.Value { + t.Fatal("rotation reused the same token") + } + + if _, _, ok := visit(mgr, old); ok { + t.Error("the pre-rotation token still works") + } + _, sess, ok := visit(mgr, rotated) + if !ok { + t.Fatal("the rotated token does not work") + } + if sess.SubjectID != subject { + t.Errorf("rotation changed the subject to %s", sess.SubjectID) + } +} + +// TestRotate_DoesNotExtendTheAbsoluteLifetime: rotation is a security measure, +// not a way to hold a session open indefinitely. +func TestRotate_DoesNotExtendTheAbsoluteLifetime(t *testing.T) { + mgr, _, db := fixture(t, session.Options{}) + + old := login(t, mgr, uuid.New()) + + var before session.Session + if err := db.First(&before).Error; err != nil { + t.Fatalf("read back: %v", err) + } + + h := mgr.Middleware()(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if _, err := mgr.Rotate(w, r); err != nil { + t.Errorf("rotate: %v", err) + } + })) + req := httptest.NewRequest(http.MethodPost, "/console/elevate", nil) + req.AddCookie(old) + h.ServeHTTP(httptest.NewRecorder(), req) + + var after session.Session + if err := db.First(&after).Error; err != nil { + t.Fatalf("read back after rotation: %v", err) + } + if !after.ExpiresAt.Equal(before.ExpiresAt) { + t.Errorf("expiry moved from %v to %v", before.ExpiresAt, after.ExpiresAt) + } +} + +func TestDestroy_EndsTheSession(t *testing.T) { + mgr, _, db := fixture(t, session.Options{}) + + cookie := login(t, mgr, uuid.New()) + + h := mgr.Middleware()(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := mgr.Destroy(w, r); err != nil { + t.Errorf("destroy: %v", err) + } + })) + req := httptest.NewRequest(http.MethodPost, "/console/logout", nil) + req.AddCookie(cookie) + h.ServeHTTP(httptest.NewRecorder(), req) + + var remaining int64 + db.Model(&session.Session{}).Count(&remaining) + if remaining != 0 { + t.Errorf("%d sessions left in the table after logout", remaining) + } + if _, _, ok := visit(mgr, cookie); ok { + t.Error("the token still works after logout") + } +} + +// TestDestroyAll_SignsOutEveryDevice is what a password reset has to do. +func TestDestroyAll_SignsOutEveryDevice(t *testing.T) { + mgr, _, db := fixture(t, session.Options{}) + subject, other := uuid.New(), uuid.New() + + phone := login(t, mgr, subject) + laptop := login(t, mgr, subject) + elsewhere := login(t, mgr, other) + + if err := mgr.DestroyAll(context.Background(), subject); err != nil { + t.Fatalf("destroy all: %v", err) + } + + for name, c := range map[string]*http.Cookie{"phone": phone, "laptop": laptop} { + if _, _, ok := visit(mgr, c); ok { + t.Errorf("the %s session survived", name) + } + } + if _, _, ok := visit(mgr, elsewhere); !ok { + t.Error("another subject's session was destroyed too") + } + + var remaining int64 + db.Model(&session.Session{}).Count(&remaining) + if remaining != 1 { + t.Errorf("%d sessions remain, want 1", remaining) + } +} + +// TestRequire_RedirectsWithoutASession, and lets one through with. +func TestRequire_RedirectsWithoutASession(t *testing.T) { + mgr, _, _ := fixture(t, session.Options{}) + + reached := false + h := mgr.Middleware()(mgr.Require(loginPath)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + reached = true + }))) + + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/console/dashboard", nil)) + + if reached { + t.Error("the protected handler ran without a session") + } + if rec.Code != http.StatusSeeOther { + t.Errorf("status %d, want 303", rec.Code) + } + if loc := rec.Header().Get("Location"); loc != loginPath { + t.Errorf("redirected to %q, want %q", loc, loginPath) + } + + rec = httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/console/dashboard", nil) + req.AddCookie(login(t, mgr, uuid.New())) + h.ServeHTTP(rec, req) + + if !reached { + t.Error("a valid session was refused") + } +} + +func TestDeleteExpired_SweepsOnlyDeadRows(t *testing.T) { + mgr, _, db := fixture(t, session.Options{}) + + liveSubject, deadSubject := uuid.New(), uuid.New() + live := login(t, mgr, liveSubject) + dead := login(t, mgr, deadSubject) + + // Selected by subject rather than by ordering: two sessions created in the + // same microsecond would make an ORDER BY created_at tie-break arbitrary, + // and a test that usually picks the right row is worse than no test. + if err := db.Model(&session.Session{}). + Where("subject_id = ?", deadSubject). + Update("expires_at", time.Now().Add(-time.Hour)).Error; err != nil { + t.Fatalf("age one session: %v", err) + } + + n, err := mgr.DeleteExpired(context.Background()) + if err != nil { + t.Fatalf("sweep: %v", err) + } + if n != 1 { + t.Errorf("swept %d rows, want 1", n) + } + if _, _, ok := visit(mgr, live); !ok { + t.Error("the live session was swept") + } + if _, _, ok := visit(mgr, dead); ok { + t.Error("the expired session survived") + } +} + +func TestNew_RequiresAKind(t *testing.T) { + db := testdb.Connect(t) + if _, err := session.New(session.NewGormStore(db), session.Options{}); err == nil { + t.Error("expected New to require a Kind") + } +} diff --git a/pkg/session/store.go b/pkg/session/store.go new file mode 100644 index 0000000..85ffffc --- /dev/null +++ b/pkg/session/store.go @@ -0,0 +1,132 @@ +package session + +import ( + "context" + "errors" + "time" + + "github.com/google/uuid" + "gorm.io/gorm" + + "github.com/codetheuri/tusk/v2/pkg/id" +) + +// Session is one signed-in browser. +// +// The token the browser holds is deliberately absent. Only its SHA-256 digest is +// stored, so a database dump — a backup on a laptop, a leaked replica, an +// over-broad support query — does not hand over live sessions. This mirrors how +// Tusk stores refresh tokens, and is the same reason password hashes are not +// passwords. +type Session struct { + ID uuid.UUID `gorm:"type:uuid;primaryKey"` + TokenHash string `gorm:"size:64;uniqueIndex;not null"` + + // SubjectKind separates audiences that share a database. An application with + // both an operator console and ordinary user sessions gives them different + // kinds, so a cookie from one can never resolve to a session in the other. + SubjectKind string `gorm:"size:64;not null;index:idx_sessions_subject,priority:1"` + + // SubjectID identifies who is signed in. There is intentionally no foreign + // key: the subject may live in any table — Tusk's own users, or an + // application's separate operators table — and the package has no business + // deciding which. + SubjectID uuid.UUID `gorm:"type:uuid;not null;index:idx_sessions_subject,priority:2"` + + IP string `gorm:"size:45;not null"` + UserAgent string `gorm:"size:255;not null"` + + CreatedAt time.Time `gorm:"not null"` + LastSeenAt time.Time `gorm:"not null"` + ExpiresAt time.Time `gorm:"not null;index"` +} + +func (Session) TableName() string { return "sessions" } + +func (s *Session) BeforeCreate(*gorm.DB) error { + if id.IsZero(s.ID) { + s.ID = id.New() + } + return nil +} + +// ErrNotFound reports that no live session matches. It deliberately does not +// distinguish "no such token" from "expired" — the caller's response is the same +// either way, and telling a client which it was reveals whether a token was ever +// valid. +var ErrNotFound = errors.New("session: not found") + +// Store persists sessions. +// +// An interface because sessions are the one part of an admin panel that +// reasonably lives somewhere other than the main database — Redis, most often, +// when session volume outgrows a table. GormStore is the default. +type Store interface { + Create(ctx context.Context, s *Session) error + + // Find returns the live session for a token digest, or ErrNotFound. It must + // not return expired sessions. + Find(ctx context.Context, kind, tokenHash string, now time.Time) (*Session, error) + + Touch(ctx context.Context, sessionID uuid.UUID, at time.Time) error + Delete(ctx context.Context, sessionID uuid.UUID) error + DeleteBySubject(ctx context.Context, kind string, subjectID uuid.UUID) error + + // DeleteExpired removes rows past their expiry and reports how many. + DeleteExpired(ctx context.Context, now time.Time) (int64, error) +} + +// GormStore stores sessions in the application database. +type GormStore struct { + db *gorm.DB +} + +func NewGormStore(db *gorm.DB) *GormStore { + return &GormStore{db: db} +} + +var _ Store = (*GormStore)(nil) + +func (s *GormStore) Create(ctx context.Context, sess *Session) error { + return s.db.WithContext(ctx).Create(sess).Error +} + +func (s *GormStore) Find(ctx context.Context, kind, tokenHash string, now time.Time) (*Session, error) { + var sess Session + + // Expiry is filtered in the query rather than checked after loading. A row + // read and then rejected in Go is one refactor away from being read and used. + err := s.db.WithContext(ctx). + Where("subject_kind = ? AND token_hash = ? AND expires_at > ?", kind, tokenHash, now). + First(&sess).Error + + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, ErrNotFound + } + if err != nil { + return nil, err + } + return &sess, nil +} + +func (s *GormStore) Touch(ctx context.Context, sessionID uuid.UUID, at time.Time) error { + return s.db.WithContext(ctx). + Model(&Session{}). + Where("id = ?", sessionID). + Update("last_seen_at", at).Error +} + +func (s *GormStore) Delete(ctx context.Context, sessionID uuid.UUID) error { + return s.db.WithContext(ctx).Where("id = ?", sessionID).Delete(&Session{}).Error +} + +func (s *GormStore) DeleteBySubject(ctx context.Context, kind string, subjectID uuid.UUID) error { + return s.db.WithContext(ctx). + Where("subject_kind = ? AND subject_id = ?", kind, subjectID). + Delete(&Session{}).Error +} + +func (s *GormStore) DeleteExpired(ctx context.Context, now time.Time) (int64, error) { + res := s.db.WithContext(ctx).Where("expires_at <= ?", now).Delete(&Session{}) + return res.RowsAffected, res.Error +} diff --git a/pkg/tenant/middleware.go b/pkg/tenant/middleware.go new file mode 100644 index 0000000..35b9460 --- /dev/null +++ b/pkg/tenant/middleware.go @@ -0,0 +1,75 @@ +package tenant + +import ( + "net/http" + + "github.com/danielgtaylor/huma/v2" + "github.com/google/uuid" + + "github.com/codetheuri/tusk/v2/pkg/response" +) + +// Resolver determines the active tenant for a request. +// +// Two outcomes are not failures and must be distinguished: +// +// - (uuid.Nil, nil) — no tenant applies. The request continues unscoped, and +// any tenanted query it goes on to make fails closed with ErrNoTenant. This +// is the correct answer for login, registration and health checks. +// - (id, nil) — this request acts as that tenant. +// +// An error rejects the request outright, for a credential that names a tenant +// the caller may not act as. +// +// A Resolver must read only from sources the caller cannot choose: a verified +// token claim, a subdomain matched against a lookup, a gateway header on a +// trusted network. Never a request body or query parameter — a tenant taken from +// those is not an identity, it is a request to be someone else. +type Resolver func(r *http.Request) (uuid.UUID, error) + +// HumaResolver is Resolver for the Huma API surface. +type HumaResolver func(ctx huma.Context) (uuid.UUID, error) + +// Middleware puts the resolved tenant into the request context. +// +// It is the only supported way for a tenant to enter the system. Handlers read +// it from the context and never accept it as a parameter, so no route can be +// written that takes the caller's word for who they are. +// +// Nothing here activates on its own: an application that registers no middleware +// simply never has a tenant in context, and models that never opted in are +// unaffected either way. +func Middleware(resolve Resolver) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + id, err := resolve(r) + if err != nil { + response.WriteJSON(w, http.StatusUnauthorized, + `{"success":false,"message":"Unable to determine the tenant for this request"}`) + return + } + if id == uuid.Nil { + next.ServeHTTP(w, r) + return + } + next.ServeHTTP(w, r.WithContext(WithTenant(r.Context(), id))) + }) + } +} + +// HumaMiddleware is Middleware for the Huma API surface. +func HumaMiddleware(api huma.API, resolve HumaResolver) func(huma.Context, func(huma.Context)) { + return func(ctx huma.Context, next func(huma.Context)) { + id, err := resolve(ctx) + if err != nil { + huma.WriteErr(api, ctx, http.StatusUnauthorized, + "Unable to determine the tenant for this request") + return + } + if id == uuid.Nil { + next(ctx) + return + } + next(huma.WithContext(ctx, WithTenant(ctx.Context(), id))) + } +} diff --git a/pkg/tenant/middleware_test.go b/pkg/tenant/middleware_test.go new file mode 100644 index 0000000..8075ecc --- /dev/null +++ b/pkg/tenant/middleware_test.go @@ -0,0 +1,99 @@ +package tenant_test + +import ( + "errors" + "net/http" + "net/http/httptest" + "testing" + + "github.com/google/uuid" + + "github.com/codetheuri/tusk/v2/pkg/tenant" +) + +// serve runs one request through the middleware and reports what the handler saw. +func serve(t *testing.T, resolve tenant.Resolver) (status int, got uuid.UUID, present, reached bool) { + t.Helper() + + handler := tenant.Middleware(resolve)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + reached = true + got, present = tenant.FromContext(r.Context()) + w.WriteHeader(http.StatusOK) + })) + + rec := httptest.NewRecorder() + handler.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil)) + return rec.Code, got, present, reached +} + +func TestMiddleware_PutsTheResolvedTenantInContext(t *testing.T) { + want := uuid.New() + + status, got, present, reached := serve(t, func(*http.Request) (uuid.UUID, error) { + return want, nil + }) + + if !reached { + t.Fatal("handler was not reached") + } + if status != http.StatusOK { + t.Errorf("status %d, want 200", status) + } + if !present || got != want { + t.Errorf("handler saw tenant %v (present=%v), want %v", got, present, want) + } +} + +// TestMiddleware_NoTenantIsNotAnError covers login, registration and health +// checks. They must keep working; a request that goes on to touch a tenanted +// model still fails closed at the query layer. +func TestMiddleware_NoTenantIsNotAnError(t *testing.T) { + status, _, present, reached := serve(t, func(*http.Request) (uuid.UUID, error) { + return uuid.Nil, nil + }) + + if !reached { + t.Fatal("handler was not reached; a request without a tenant must still be served") + } + if status != http.StatusOK { + t.Errorf("status %d, want 200", status) + } + if present { + t.Error("a tenant was placed in context when the resolver reported none") + } +} + +func TestMiddleware_ResolverErrorRejectsTheRequest(t *testing.T) { + status, _, _, reached := serve(t, func(*http.Request) (uuid.UUID, error) { + return uuid.New(), errors.New("caller may not act as this tenant") + }) + + if reached { + t.Error("handler ran despite the resolver refusing the request") + } + if status != http.StatusUnauthorized { + t.Errorf("status %d, want 401", status) + } +} + +// TestMiddleware_IgnoresATenantTheCallerSupplies is the property the whole design +// rests on: the tenant comes from the resolver, and a body or query parameter +// naming a different one changes nothing. +func TestMiddleware_IgnoresATenantTheCallerSupplies(t *testing.T) { + real, attacker := uuid.New(), uuid.New() + + var seen uuid.UUID + handler := tenant.Middleware(func(*http.Request) (uuid.UUID, error) { + return real, nil + })(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + seen, _ = tenant.FromContext(r.Context()) + })) + + req := httptest.NewRequest(http.MethodGet, "/?tenant_id="+attacker.String(), nil) + req.Header.Set("X-Tenant-ID", attacker.String()) + handler.ServeHTTP(httptest.NewRecorder(), req) + + if seen != real { + t.Errorf("handler saw %v, want %v — a caller-supplied value reached the context", seen, real) + } +} diff --git a/pkg/tenant/rls.go b/pkg/tenant/rls.go new file mode 100644 index 0000000..ccc4384 --- /dev/null +++ b/pkg/tenant/rls.go @@ -0,0 +1,179 @@ +package tenant + +import ( + "context" + "fmt" + "strings" + + "gorm.io/gorm" +) + +// SettingName is the PostgreSQL run-time parameter carrying the active tenant. +// +// Policies read it with current_setting(SettingName, true) — the second argument +// asks for NULL instead of an error when it is unset. That choice is what makes +// the policy fail closed: an unset parameter compares as NULL, `column = NULL` +// is NULL rather than true, and no rows match. A connection that forgot to +// declare its tenant sees an empty database, not everyone's. +const SettingName = "app.tenant_id" + +// policyName is the name given to the generated policy on every table. Fixed +// rather than derived, so DisableRLS and re-running EnableRLS both know it. +const policyName = "tusk_tenant_isolation" + +// PolicyStatements returns the DDL that puts one table under tenant isolation. +// +// Exported as strings because row-level security belongs in a migration, next to +// the table it protects and under review — not applied at boot by whichever +// process starts first. Paste the output into a migration; EnableRLS is the +// convenience wrapper for tests and local development. +func PolicyStatements(table, column string) []string { + t, c := quoteIdent(table), quoteIdent(column) + predicate := fmt.Sprintf("%s = current_setting('%s', true)::uuid", c, SettingName) + + return []string{ + fmt.Sprintf("ALTER TABLE %s ENABLE ROW LEVEL SECURITY", t), + + // FORCE is not optional. Without it PostgreSQL exempts the table's owner + // from its own policies, and application code very often connects as the + // owner — so the protection would be installed, visible in \d, and doing + // nothing at all. + fmt.Sprintf("ALTER TABLE %s FORCE ROW LEVEL SECURITY", t), + + fmt.Sprintf("DROP POLICY IF EXISTS %s ON %s", quoteIdent(policyName), t), + + // USING filters what a statement may see; WITH CHECK constrains what it + // may write. Both are needed: USING alone would let a tenant INSERT a row + // stamped with someone else's identifier, or UPDATE one of its own rows + // to move it across the boundary. + fmt.Sprintf("CREATE POLICY %s ON %s USING (%s) WITH CHECK (%s)", + quoteIdent(policyName), t, predicate, predicate), + } +} + +// DropPolicyStatements reverses PolicyStatements, for a migration's down step. +func DropPolicyStatements(table string) []string { + t := quoteIdent(table) + return []string{ + fmt.Sprintf("DROP POLICY IF EXISTS %s ON %s", quoteIdent(policyName), t), + fmt.Sprintf("ALTER TABLE %s NO FORCE ROW LEVEL SECURITY", t), + fmt.Sprintf("ALTER TABLE %s DISABLE ROW LEVEL SECURITY", t), + } +} + +// EnableRLS applies the policies for the given models. +// +// Each model must implement Tenanted; passing one that does not is a mistake +// worth reporting rather than skipping, because the caller clearly believed it +// was protected. +func EnableRLS(db *gorm.DB, models ...any) error { + return applyPerModel(db, models, func(table, column string) []string { + return PolicyStatements(table, column) + }) +} + +// DisableRLS removes them again. +func DisableRLS(db *gorm.DB, models ...any) error { + return applyPerModel(db, models, func(table, _ string) []string { + return DropPolicyStatements(table) + }) +} + +func applyPerModel(db *gorm.DB, models []any, build func(table, column string) []string) error { + for _, model := range models { + table, column, err := tableAndTenantColumn(db, model) + if err != nil { + return err + } + for _, stmt := range build(table, column) { + if err := db.Exec(stmt).Error; err != nil { + return fmt.Errorf("tenant: %q: %w", stmt, err) + } + } + } + return nil +} + +// tableAndTenantColumn resolves a model to the table and column the policy needs. +func tableAndTenantColumn(db *gorm.DB, model any) (string, string, error) { + t, ok := model.(Tenanted) + if !ok { + return "", "", fmt.Errorf("tenant: %T does not implement Tenanted, so it has no tenant column to protect", model) + } + + stmt := &gorm.Statement{DB: db} + if err := stmt.Parse(model); err != nil { + return "", "", fmt.Errorf("tenant: cannot parse %T: %w", model, err) + } + + column := t.TenantColumn() + if stmt.Schema.LookUpField(column) == nil { + return "", "", fmt.Errorf("tenant: %s declares tenant column %q, which is not a field on the model", stmt.Schema.Name, column) + } + return stmt.Schema.Table, column, nil +} + +// Transaction runs fn with the tenant declared to PostgreSQL, so row-level +// security applies for the duration. +// +// The declaration has to live inside a transaction. SET LOCAL is scoped to one, +// and reverts on commit or rollback — which is the property that makes this safe +// against connection pooling. A plain SET would persist on the pooled connection +// and hand the next request, belonging to a different tenant, whatever this one +// left behind. +func Transaction(ctx context.Context, db *gorm.DB, fn func(tx *gorm.DB) error) error { + id, err := MustFromContext(ctx) + if err != nil { + return err + } + + return db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + // set_config(name, value, is_local=true) is SET LOCAL in function form. + // SET cannot take a bind parameter, so the literal spelling would mean + // interpolating a value into SQL; this keeps it a parameter. + if err := tx.Exec("SELECT set_config(?, ?, true)", SettingName, id.String()).Error; err != nil { + return fmt.Errorf("tenant: failed to declare %s: %w", SettingName, err) + } + return fn(tx) + }) +} + +// VerifyEnforcement reports whether row-level security actually binds the role +// this connection uses. +// +// Worth calling at startup and logging loudly. PostgreSQL exempts superusers and +// any role holding BYPASSRLS from every policy, without warning — so a system can +// have correct policies, a passing review, and no protection whatsoever. The +// failure is invisible precisely because nothing errors. +func VerifyEnforcement(db *gorm.DB) error { + var role struct { + Name string + Super bool + BypassRL bool + } + + err := db.Raw(` + SELECT rolname AS name, rolsuper AS super, rolbypassrls AS bypass_rl + FROM pg_roles WHERE rolname = current_user + `).Scan(&role).Error + if err != nil { + return fmt.Errorf("tenant: cannot determine whether RLS is enforced: %w", err) + } + + switch { + case role.Super: + return fmt.Errorf("tenant: connected as superuser %q, which bypasses every row-level security policy; "+ + "run the application as a dedicated non-superuser role", role.Name) + case role.BypassRL: + return fmt.Errorf("tenant: role %q holds BYPASSRLS, so row-level security does not apply to it; "+ + "revoke it with ALTER ROLE %s NOBYPASSRLS", role.Name, role.Name) + default: + return nil + } +} + +// quoteIdent quotes a SQL identifier. These come from struct tags rather than +// from users, but a table called "order" still needs quoting to parse. +func quoteIdent(s string) string { + return `"` + strings.ReplaceAll(s, `"`, `""`) + `"` +} diff --git a/pkg/tenant/rls_test.go b/pkg/tenant/rls_test.go new file mode 100644 index 0000000..6d7f1ea --- /dev/null +++ b/pkg/tenant/rls_test.go @@ -0,0 +1,242 @@ +package tenant_test + +import ( + "context" + "database/sql" + "fmt" + "testing" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5" + _ "github.com/jackc/pgx/v5/stdlib" + "gorm.io/driver/postgres" + "gorm.io/gorm" + gormlogger "gorm.io/gorm/logger" + + "github.com/codetheuri/tusk/v2/pkg/id" + "github.com/codetheuri/tusk/v2/pkg/tenant" + "github.com/codetheuri/tusk/v2/pkg/testdb" +) + +// Row-level security is the layer that still holds when the application layer is +// wrong. Testing it therefore means bypassing the application layer on purpose: +// every query below goes through tenant.Unscoped, so nothing but PostgreSQL is +// deciding which rows come back. +// +// The tests connect as a dedicated non-superuser role. This is not fastidiousness +// — PostgreSQL exempts superusers and BYPASSRLS roles from every policy, so the +// same suite run as the usual development superuser would pass without a single +// policy being consulted. + +const ( + rlsRole = "tusk_rls_tester" + rlsPassword = "tusk_rls_tester_pw" +) + +type rlsNote struct { + ID uuid.UUID `gorm:"type:uuid;primaryKey"` + BusinessID uuid.UUID `gorm:"type:uuid;not null;index"` + Title string +} + +func (rlsNote) TableName() string { return "tenant_test_rls_notes" } +func (rlsNote) TenantColumn() string { return "business_id" } + +func (n *rlsNote) BeforeCreate(*gorm.DB) error { + if id.IsZero(n.ID) { + n.ID = id.New() + } + return nil +} + +// rlsFixture builds a table under policy, seeds it as the owner, and returns a +// second connection held by a role the policies actually bind. +func rlsFixture(t *testing.T) (asRole *gorm.DB, alice, bob uuid.UUID) { + t.Helper() + + owner := testdb.Connect(t) + + if err := owner.AutoMigrate(&rlsNote{}); err != nil { + t.Fatalf("automigrate: %v", err) + } + t.Cleanup(func() { owner.Exec("DROP TABLE IF EXISTS " + (rlsNote{}).TableName()) }) + + if err := tenant.EnableRLS(owner, &rlsNote{}); err != nil { + t.Fatalf("enable rls: %v", err) + } + + alice, bob = id.New(), id.New() + seed := []rlsNote{ + {ID: id.New(), BusinessID: alice, Title: "alice-secret"}, + {ID: id.New(), BusinessID: bob, Title: "bob-secret"}, + } + // Explicitly unscoped: the fixture is deliberately building rows for two + // tenants at once, which is precisely what the application layer forbids. + // Going through it here would make the fixture depend on the layer these + // tests exist to be independent of. + if err := tenant.Unscoped(owner).Create(&seed).Error; err != nil { + t.Fatalf("seed: %v", err) + } + + createRole(t, owner) + return connectAsRole(t), alice, bob +} + +func createRole(t *testing.T, owner *gorm.DB) { + t.Helper() + + // Idempotent: a previous run that died before cleanup must not block this one. + stmts := []string{ + fmt.Sprintf(`DO $$ BEGIN + IF NOT EXISTS (SELECT 1 FROM pg_roles WHERE rolname = '%s') THEN + CREATE ROLE %s LOGIN PASSWORD '%s'; + END IF; + END $$`, rlsRole, rlsRole, rlsPassword), + fmt.Sprintf(`GRANT USAGE ON SCHEMA public TO %s`, rlsRole), + fmt.Sprintf(`GRANT SELECT, INSERT, UPDATE, DELETE ON %s TO %s`, (rlsNote{}).TableName(), rlsRole), + } + for _, s := range stmts { + if err := owner.Exec(s).Error; err != nil { + if testdb.Required() { + t.Fatalf("could not prepare the unprivileged test role: %v", err) + } + t.Skipf("cannot create a test role on this database, skipping RLS tests: %v", err) + } + } + + t.Cleanup(func() { + owner.Exec(fmt.Sprintf("DROP OWNED BY %s", rlsRole)) + owner.Exec(fmt.Sprintf("DROP ROLE IF EXISTS %s", rlsRole)) + }) +} + +// connectAsRole opens a second connection as the unprivileged role, reusing the +// host and database the harness was pointed at. +func connectAsRole(t *testing.T) *gorm.DB { + t.Helper() + + // Parsed rather than string-substituted, so TEST_DATABASE_URL works in either + // the keyword/value or the URL form. + cfg, err := pgx.ParseConfig(testdb.DSN()) + if err != nil { + t.Fatalf("parse test DSN: %v", err) + } + sslmode := "disable" + if cfg.TLSConfig != nil { + sslmode = "require" + } + dsn := fmt.Sprintf("host=%s port=%d dbname=%s user=%s password=%s sslmode=%s", + cfg.Host, cfg.Port, cfg.Database, rlsRole, rlsPassword, sslmode) + + sqlDB, err := sql.Open("pgx", dsn) + if err != nil { + t.Fatalf("open as %s: %v", rlsRole, err) + } + t.Cleanup(func() { _ = sqlDB.Close() }) + + db, err := gorm.Open(postgres.New(postgres.Config{Conn: sqlDB}), &gorm.Config{ + Logger: gormlogger.Default.LogMode(gormlogger.Silent), + }) + if err != nil { + t.Fatalf("gorm open as %s: %v", rlsRole, err) + } + if err := tenant.Register(db); err != nil { + t.Fatalf("register callbacks: %v", err) + } + return db +} + +// TestRLS_HidesOtherTenantsEvenWithApplicationScopingBypassed is the whole point +// of the layer: tenant.Unscoped disables everything the Go code does, and the +// rows still do not appear. +func TestRLS_HidesOtherTenantsEvenWithApplicationScopingBypassed(t *testing.T) { + db, alice, _ := rlsFixture(t) + ctx := tenant.WithTenant(context.Background(), alice) + + err := tenant.Transaction(ctx, db, func(tx *gorm.DB) error { + var got []rlsNote + if err := tenant.Unscoped(tx).Find(&got).Error; err != nil { + return err + } + if len(got) != 1 { + t.Errorf("got %d rows, want only alice's 1", len(got)) + } + for _, n := range got { + if n.BusinessID != alice { + t.Errorf("database returned a row owned by %s", n.BusinessID) + } + } + return nil + }) + if err != nil { + t.Fatalf("transaction: %v", err) + } +} + +// TestRLS_WithNoTenantDeclaredReturnsNothing checks the direction that matters. +// An unset parameter must mean "no rows", not "all rows". +func TestRLS_WithNoTenantDeclaredReturnsNothing(t *testing.T) { + db, _, _ := rlsFixture(t) + + var count int64 + if err := db.Raw("SELECT count(*) FROM " + (rlsNote{}).TableName()).Scan(&count).Error; err != nil { + t.Fatalf("count: %v", err) + } + if count != 0 { + t.Errorf("a connection that declared no tenant saw %d rows", count) + } +} + +// TestRLS_RefusesAWriteStampedWithAnotherTenant covers WITH CHECK. USING alone +// would let a tenant insert rows it could then never see — and that another +// tenant would. +func TestRLS_RefusesAWriteStampedWithAnotherTenant(t *testing.T) { + db, alice, bob := rlsFixture(t) + ctx := tenant.WithTenant(context.Background(), alice) + + err := tenant.Transaction(ctx, db, func(tx *gorm.DB) error { + return tenant.Unscoped(tx).Create(&rlsNote{ + ID: id.New(), + BusinessID: bob, + Title: "planted", + }).Error + }) + + if err == nil { + t.Fatal("the database accepted a row belonging to another tenant") + } +} + +// TestRLS_RefusesToMoveARowAcrossTenants is the same guarantee for UPDATE. +func TestRLS_RefusesToMoveARowAcrossTenants(t *testing.T) { + db, alice, bob := rlsFixture(t) + ctx := tenant.WithTenant(context.Background(), alice) + + err := tenant.Transaction(ctx, db, func(tx *gorm.DB) error { + return tenant.Unscoped(tx). + Model(&rlsNote{}). + Where("business_id = ?", alice). + Update("business_id", bob).Error + }) + + if err == nil { + t.Fatal("the database allowed a row to be moved to another tenant") + } +} + +// TestVerifyEnforcement_DistinguishesTheTwoRoles pins the check that stops this +// entire layer from being quietly inert in production. +func TestVerifyEnforcement_DistinguishesTheTwoRoles(t *testing.T) { + asRole, _, _ := rlsFixture(t) + + if err := tenant.VerifyEnforcement(asRole); err != nil { + t.Errorf("expected RLS to bind %s, got: %v", rlsRole, err) + } + + owner := testdb.Connect(t) + if err := tenant.VerifyEnforcement(owner); err == nil { + t.Error("expected the superuser connection to be reported as exempt from RLS") + } else { + t.Logf("correctly reported: %v", err) + } +} diff --git a/pkg/tenant/scope.go b/pkg/tenant/scope.go new file mode 100644 index 0000000..5dacb3e --- /dev/null +++ b/pkg/tenant/scope.go @@ -0,0 +1,285 @@ +package tenant + +import ( + "errors" + "fmt" + "reflect" + + "github.com/google/uuid" + "gorm.io/gorm" + "gorm.io/gorm/clause" + "gorm.io/gorm/schema" +) + +// ErrCrossTenantWrite is returned when a record carrying one tenant's identifier +// is written while another tenant is in context. +// +// This is not the same failure as ErrNoTenant. No tenant at all is usually a +// wiring mistake; a *mismatched* tenant is an attempt, deliberate or not, to +// write into someone else's data, and is worth distinguishing in logs. +var ErrCrossTenantWrite = errors.New("tenant: record belongs to a different tenant") + +// unscopedSetting marks a statement as a deliberate cross-tenant operation. +const unscopedSetting = "tenant:unscoped" + +// Unscoped disables tenant scoping for one chain of calls. +// +// var all []Customer +// tenant.Unscoped(db).Find(&all) // every tenant's customers +// +// Legitimate uses are narrow: migrations, background jobs that fan out across +// tenants, and support tooling. It is a named function rather than a context +// flag precisely so that `grep -rn 'tenant.Unscoped'` lists every place the +// guarantee is suspended, and a reviewer can check each one. +// +// Call it fresh before every statement — never store its return value and issue +// more than one statement against it: +// +// // WRONG: both Creates run against the same *gorm.DB. +// utx := tenant.Unscoped(tx) +// utx.Create(&a) +// utx.Create(&b) // silently reuses a's table and columns +// +// // RIGHT. +// tenant.Unscoped(tx).Create(&a) +// tenant.Unscoped(tx).Create(&b) +// +// This is not specific to Unscoped — it is true of any chained GORM call +// (db.Where(...) has the identical trap) — but it is easy to hit here +// specifically, because Unscoped looks like a modifier you set once per +// transaction rather than per statement. The value it returns has already been +// "obtained" in GORM's terms, so the next chained call mutates its Statement in +// place instead of cloning one. +// +// This is unrelated to GORM's own db.Unscoped(), which disables soft-delete +// filtering. The two are independent and can be combined. +func Unscoped(db *gorm.DB) *gorm.DB { + return db.Set(unscopedSetting, true) +} + +func isUnscoped(db *gorm.DB) bool { + v, ok := db.Statement.Settings.Load(unscopedSetting) + if !ok { + return false + } + on, _ := v.(bool) + return on +} + +// Register installs the tenant callbacks on a GORM connection. +// +// Call it once, immediately after opening the database and before any query +// runs. It is safe in single-tenant applications: every callback returns +// immediately for models that do not implement Tenanted, which is all of them if +// nothing opts in. +func Register(db *gorm.DB) error { + cb := db.Callback() + + return errors.Join( + cb.Create().Before("gorm:create").Register("tenant:create", stampOnCreate), + cb.Query().Before("gorm:query").Register("tenant:query", scopeRead), + cb.Row().Before("gorm:row").Register("tenant:row", scopeRead), + cb.Update().Before("gorm:update").Register("tenant:update", scopeWrite), + cb.Delete().Before("gorm:delete").Register("tenant:delete", scopeWrite), + ) +} + +// tenantColumnOf returns the tenant column declared by the statement's model, or +// "" when the model has not opted in. +// +// It reads Schema.ModelType rather than inspecting Statement.Model or Dest +// directly, because GORM has already resolved those into a schema: a *Customer, +// a *[]Customer and a *[]*Customer all arrive here as the same ModelType. Raw +// SQL leaves Schema nil and is therefore never scoped — see the package doc. +func tenantColumnOf(stmt *gorm.Statement) string { + if stmt == nil || stmt.Schema == nil { + return "" + } + t, ok := reflect.New(stmt.Schema.ModelType).Interface().(Tenanted) + if !ok { + return "" + } + return t.TenantColumn() +} + +// tenantField resolves the declared column to a model field, so a typo in +// TenantColumn() surfaces as a clear error rather than invalid SQL. +func tenantField(stmt *gorm.Statement, col string) (*schema.Field, error) { + f := stmt.Schema.LookUpField(col) + if f == nil { + return nil, fmt.Errorf("tenant: %s declares tenant column %q, which is not a field on the model", stmt.Schema.Name, col) + } + return f, nil +} + +// scopeRead adds the tenant predicate to SELECTs. +func scopeRead(db *gorm.DB) { + col := tenantColumnOf(db.Statement) + if col == "" || isUnscoped(db) { + return + } + + id, ok := FromContext(db.Statement.Context) + if !ok { + _ = db.AddError(ErrNoTenant) + return + } + if _, err := tenantField(db.Statement, col); err != nil { + _ = db.AddError(err) + return + } + + addTenantPredicate(db.Statement, col, id) +} + +// scopeWrite adds the tenant predicate to UPDATEs and DELETEs. +// +// It carries one extra responsibility over scopeRead. GORM refuses an UPDATE or +// DELETE that carries no conditions, so a mistyped chain cannot rewrite a whole +// table. That check runs *after* this callback and is satisfied by any WHERE +// clause — including the one added here. Adding the tenant predicate +// unconditionally would therefore convert "GORM refuses this" into "GORM +// silently rewrites every row this tenant owns", which is worse for being +// plausible. So the same condition is required first, on the same terms GORM +// uses: an explicit WHERE, a primary key on the model, or AllowGlobalUpdate. +func scopeWrite(db *gorm.DB) { + col := tenantColumnOf(db.Statement) + if col == "" || isUnscoped(db) { + return + } + + id, ok := FromContext(db.Statement.Context) + if !ok { + _ = db.AddError(ErrNoTenant) + return + } + if _, err := tenantField(db.Statement, col); err != nil { + _ = db.AddError(err) + return + } + if !db.AllowGlobalUpdate && !hasCondition(db.Statement) { + _ = db.AddError(gorm.ErrMissingWhereClause) + return + } + + addTenantPredicate(db.Statement, col, id) +} + +// hasCondition reports whether the statement already restricts which rows it +// touches, by the same two routes GORM accepts: an explicit WHERE, or a +// populated primary key on the model it was given. +func hasCondition(stmt *gorm.Statement) bool { + if _, ok := stmt.Clauses["WHERE"]; ok { + return true + } + if stmt.Schema == nil || stmt.ReflectValue.Kind() != reflect.Struct { + return false + } + for _, pk := range stmt.Schema.PrimaryFields { + if _, isZero := pk.ValueOf(stmt.Context, stmt.ReflectValue); !isZero { + return true + } + } + return false +} + +// addTenantPredicate appends `AND .= ` to the statement. +func addTenantPredicate(stmt *gorm.Statement, col string, id uuid.UUID) { + // Group the caller's existing conditions before appending, or an OR in them + // swallows the tenant check. `Where("a = ?", 1).Or("b = ?", 2)` with a + // naively appended AND builds: + // + // a = 1 OR b = 2 AND tenant = x + // + // which SQL precedence reads as `a = 1 OR (b = 2 AND tenant = x)` — every + // row matching `a = 1` leaks, regardless of who owns it. Wrapping first + // gives `(a = 1 OR b = 2) AND tenant = x`. GORM's own soft-delete clause + // does exactly this, for exactly this reason. + if c, ok := stmt.Clauses["WHERE"]; ok { + if where, ok := c.Expression.(clause.Where); ok && len(where.Exprs) >= 1 { + for _, expr := range where.Exprs { + if or, ok := expr.(clause.OrConditions); ok && len(or.Exprs) == 1 { + where.Exprs = []clause.Expression{clause.And(where.Exprs...)} + c.Expression = where + stmt.Clauses["WHERE"] = c + break + } + } + } + } + + // clause.CurrentTable qualifies the column, so the predicate stays correct + // once a join puts a second business_id in scope. + stmt.AddClause(clause.Where{Exprs: []clause.Expression{ + clause.Eq{ + Column: clause.Column{Table: clause.CurrentTable, Name: col}, + Value: id, + }, + }}) +} + +// stampOnCreate fills in the tenant column on INSERT, and rejects records that +// arrive already belonging to someone else. +// +// Stamping matters as much as scoping. If callers had to set business_id by hand +// on every create, they would eventually forget one — the same failure the query +// callback exists to prevent, in the other direction. +func stampOnCreate(db *gorm.DB) { + stmt := db.Statement + + col := tenantColumnOf(stmt) + if col == "" || isUnscoped(db) { + return + } + + id, ok := FromContext(stmt.Context) + if !ok { + _ = db.AddError(ErrNoTenant) + return + } + field, err := tenantField(stmt, col) + if err != nil { + _ = db.AddError(err) + return + } + + switch stmt.ReflectValue.Kind() { + case reflect.Slice, reflect.Array: + for i := range stmt.ReflectValue.Len() { + if err := stampOne(stmt, field, stmt.ReflectValue.Index(i), id); err != nil { + _ = db.AddError(err) + return + } + } + case reflect.Struct: + if err := stampOne(stmt, field, stmt.ReflectValue, id); err != nil { + _ = db.AddError(err) + } + } +} + +func stampOne(stmt *gorm.Statement, field *schema.Field, rv reflect.Value, id uuid.UUID) error { + current, isZero := field.ValueOf(stmt.Context, rv) + if isZero { + return field.Set(stmt.Context, rv, id) + } + + // Already set. Honour it only if it agrees with the caller's own tenant; + // anything else is a write into another tenant's data. + switch cur := current.(type) { + case uuid.UUID: + if cur != id { + return fmt.Errorf("%w: record carries %s, context holds %s", ErrCrossTenantWrite, cur, id) + } + case *uuid.UUID: + if cur != nil && *cur != id { + return fmt.Errorf("%w: record carries %s, context holds %s", ErrCrossTenantWrite, *cur, id) + } + default: + // Fail closed. A tenant column of some other type cannot be compared + // here, and guessing would mean allowing an unverified write. + return fmt.Errorf("tenant: column %q on %s must be uuid.UUID or *uuid.UUID, got %T", + field.DBName, stmt.Schema.Name, current) + } + return nil +} diff --git a/pkg/tenant/scope_test.go b/pkg/tenant/scope_test.go new file mode 100644 index 0000000..9f0e19b --- /dev/null +++ b/pkg/tenant/scope_test.go @@ -0,0 +1,385 @@ +package tenant_test + +import ( + "context" + "errors" + "testing" + + "github.com/google/uuid" + "gorm.io/gorm" + + "github.com/codetheuri/tusk/v2/pkg/id" + "github.com/codetheuri/tusk/v2/pkg/tenant" + "github.com/codetheuri/tusk/v2/pkg/testdb" +) + +// These tests run against real PostgreSQL. The thing being verified is the SQL +// GORM ends up sending, and asserting on SQL *text* would pass while the query +// still returned the wrong rows. So every case here asks the only question that +// matters: given two tenants' data in one table, which rows come back? + +// scopedNote opts in to tenancy. +type scopedNote struct { + ID uuid.UUID `gorm:"type:uuid;primaryKey"` + BusinessID uuid.UUID `gorm:"type:uuid;not null;index"` + Title string +} + +func (scopedNote) TableName() string { return "tenant_test_notes" } +func (scopedNote) TenantColumn() string { return "business_id" } + +func (n *scopedNote) BeforeCreate(*gorm.DB) error { + if id.IsZero(n.ID) { + n.ID = id.New() + } + return nil +} + +// plainNote does not. It exists to prove D5: a model that never mentions +// tenancy must behave as though this package were not installed. +type plainNote struct { + ID uuid.UUID `gorm:"type:uuid;primaryKey"` + Title string +} + +func (plainNote) TableName() string { return "tenant_test_plain_notes" } + +func (n *plainNote) BeforeCreate(*gorm.DB) error { + if id.IsZero(n.ID) { + n.ID = id.New() + } + return nil +} + +// fixture returns a database with the callbacks installed and one note owned by +// each of two tenants. +func fixture(t *testing.T) (db *gorm.DB, alice, bob uuid.UUID) { + t.Helper() + + db = testdb.Connect(t) + + // The callbacks are already installed by the harness, exactly as they are by + // database.Connect in production. + if err := db.AutoMigrate(&scopedNote{}, &plainNote{}); err != nil { + t.Fatalf("automigrate: %v", err) + } + + alice, bob = id.New(), id.New() + + // Seeded through Unscoped rather than two separate contexts: the seed is + // setup, and using the code under test to build its own fixture would let a + // bug hide itself. + seed := []scopedNote{ + {ID: id.New(), BusinessID: alice, Title: "alice-one"}, + {ID: id.New(), BusinessID: alice, Title: "alice-two"}, + {ID: id.New(), BusinessID: bob, Title: "bob-one"}, + } + if err := tenant.Unscoped(db).Create(&seed).Error; err != nil { + t.Fatalf("seed: %v", err) + } + + return db, alice, bob +} + +func TestQuery_ReturnsOnlyTheTenantsOwnRows(t *testing.T) { + db, alice, _ := fixture(t) + ctx := tenant.WithTenant(context.Background(), alice) + + var got []scopedNote + if err := db.WithContext(ctx).Find(&got).Error; err != nil { + t.Fatalf("find: %v", err) + } + + if len(got) != 2 { + t.Fatalf("got %d notes, want alice's 2", len(got)) + } + for _, n := range got { + if n.BusinessID != alice { + t.Errorf("leaked a note belonging to %s", n.BusinessID) + } + } +} + +// TestQuery_OrConditionsCannotEscapeTheTenantPredicate covers the precedence +// trap that makes a naive implementation of this package unsafe. +// +// `WHERE title = 'alice-one' OR title = 'bob-one' AND business_id = alice` +// parses as `title = 'alice-one' OR (title = 'bob-one' AND business_id = alice)`. +// Bob's note matches the left branch and comes back to Alice. +func TestQuery_OrConditionsCannotEscapeTheTenantPredicate(t *testing.T) { + db, alice, _ := fixture(t) + ctx := tenant.WithTenant(context.Background(), alice) + + var got []scopedNote + err := db.WithContext(ctx). + Where("title = ?", "alice-one"). + Or("title = ?", "bob-one"). + Find(&got).Error + if err != nil { + t.Fatalf("find: %v", err) + } + + if len(got) != 1 { + t.Fatalf("got %d notes, want only alice-one", len(got)) + } + if got[0].Title != "alice-one" { + t.Errorf("got %q, want alice-one", got[0].Title) + } +} + +// TestQuery_ByPrimaryKeyAcrossTenantsFindsNothing is the "404, not 403" case. +// Bob knows Alice's note ID and asks for it directly. +func TestQuery_ByPrimaryKeyAcrossTenantsFindsNothing(t *testing.T) { + db, alice, bob := fixture(t) + + var target scopedNote + if err := tenant.Unscoped(db).Where("business_id = ?", alice).First(&target).Error; err != nil { + t.Fatalf("locate alice's note: %v", err) + } + + var got scopedNote + err := db.WithContext(tenant.WithTenant(context.Background(), bob)). + First(&got, "id = ?", target.ID).Error + + if !errors.Is(err, gorm.ErrRecordNotFound) { + t.Fatalf("got %v, want ErrRecordNotFound — a known ID must not be a way in", err) + } +} + +// TestQuery_WithNoTenantFailsClosed is the design decision that separates this +// from a filter someone forgot to apply: absent scoping is an error, never a +// silent full-table read. +func TestQuery_WithNoTenantFailsClosed(t *testing.T) { + db, _, _ := fixture(t) + + var got []scopedNote + err := db.WithContext(context.Background()).Find(&got).Error + + if !errors.Is(err, tenant.ErrNoTenant) { + t.Fatalf("got %v, want ErrNoTenant", err) + } + if len(got) != 0 { + t.Fatalf("query returned %d rows despite erroring", len(got)) + } +} + +func TestCreate_StampsTheTenantFromContext(t *testing.T) { + db, alice, _ := fixture(t) + ctx := tenant.WithTenant(context.Background(), alice) + + note := scopedNote{Title: "written-without-a-business-id"} + if err := db.WithContext(ctx).Create(¬e).Error; err != nil { + t.Fatalf("create: %v", err) + } + + if note.BusinessID != alice { + t.Errorf("got business_id %s, want %s", note.BusinessID, alice) + } + + // Confirm it reached the row, not just the struct. + var stored scopedNote + if err := tenant.Unscoped(db).First(&stored, "id = ?", note.ID).Error; err != nil { + t.Fatalf("read back: %v", err) + } + if stored.BusinessID != alice { + t.Errorf("stored business_id %s, want %s", stored.BusinessID, alice) + } +} + +func TestCreate_RejectsARecordBelongingToAnotherTenant(t *testing.T) { + db, alice, bob := fixture(t) + ctx := tenant.WithTenant(context.Background(), bob) + + note := scopedNote{Title: "planted", BusinessID: alice} + err := db.WithContext(ctx).Create(¬e).Error + + if !errors.Is(err, tenant.ErrCrossTenantWrite) { + t.Fatalf("got %v, want ErrCrossTenantWrite", err) + } +} + +func TestCreate_StampsEveryRecordInABatch(t *testing.T) { + db, alice, _ := fixture(t) + ctx := tenant.WithTenant(context.Background(), alice) + + batch := []scopedNote{{Title: "a"}, {Title: "b"}, {Title: "c"}} + if err := db.WithContext(ctx).Create(&batch).Error; err != nil { + t.Fatalf("create batch: %v", err) + } + + for i, n := range batch { + if n.BusinessID != alice { + t.Errorf("batch[%d] got business_id %s, want %s", i, n.BusinessID, alice) + } + } +} + +func TestUpdate_CannotReachAnotherTenantsRow(t *testing.T) { + db, alice, bob := fixture(t) + + var target scopedNote + if err := tenant.Unscoped(db).Where("business_id = ?", alice).First(&target).Error; err != nil { + t.Fatalf("locate alice's note: %v", err) + } + + res := db.WithContext(tenant.WithTenant(context.Background(), bob)). + Model(&scopedNote{}). + Where("id = ?", target.ID). + Update("title", "defaced") + + if res.Error != nil { + t.Fatalf("update: %v", res.Error) + } + if res.RowsAffected != 0 { + t.Errorf("updated %d rows across a tenant boundary", res.RowsAffected) + } + + var after scopedNote + if err := tenant.Unscoped(db).First(&after, "id = ?", target.ID).Error; err != nil { + t.Fatalf("read back: %v", err) + } + if after.Title != target.Title { + t.Errorf("title changed to %q; the row was modified", after.Title) + } +} + +func TestDelete_CannotReachAnotherTenantsRow(t *testing.T) { + db, alice, bob := fixture(t) + + var target scopedNote + if err := tenant.Unscoped(db).Where("business_id = ?", alice).First(&target).Error; err != nil { + t.Fatalf("locate alice's note: %v", err) + } + + res := db.WithContext(tenant.WithTenant(context.Background(), bob)). + Where("id = ?", target.ID). + Delete(&scopedNote{}) + + if res.Error != nil { + t.Fatalf("delete: %v", res.Error) + } + if res.RowsAffected != 0 { + t.Errorf("deleted %d rows across a tenant boundary", res.RowsAffected) + } + + var count int64 + tenant.Unscoped(db).Model(&scopedNote{}).Where("id = ?", target.ID).Count(&count) + if count != 1 { + t.Error("alice's note was deleted by bob") + } +} + +// TestUpdate_StillRefusesAnUnconditionalStatement guards a regression this +// package could easily introduce. +// +// GORM rejects an UPDATE with no conditions, so a mistyped chain cannot rewrite +// a table. Its check runs after the tenant callback and is satisfied by any +// WHERE clause — so adding the tenant predicate unconditionally would turn a +// refusal into "quietly rewrite every row this tenant owns". +func TestUpdate_StillRefusesAnUnconditionalStatement(t *testing.T) { + db, alice, _ := fixture(t) + ctx := tenant.WithTenant(context.Background(), alice) + + err := db.WithContext(ctx).Model(&scopedNote{}).Update("title", "everything").Error + + if !errors.Is(err, gorm.ErrMissingWhereClause) { + t.Fatalf("got %v, want ErrMissingWhereClause", err) + } + + var untouched int64 + tenant.Unscoped(db).Model(&scopedNote{}).Where("title = ?", "everything").Count(&untouched) + if untouched != 0 { + t.Errorf("%d rows were rewritten anyway", untouched) + } +} + +func TestUnscoped_ReachesEveryTenant(t *testing.T) { + db, _, _ := fixture(t) + + var got []scopedNote + if err := tenant.Unscoped(db).Find(&got).Error; err != nil { + t.Fatalf("find: %v", err) + } + if len(got) != 3 { + t.Errorf("got %d notes, want all 3", len(got)) + } +} + +// TestUntenantedModelIsUnaffected is the D5 proof at package level: a model that +// does not implement Tenanted must read, write and delete with no tenant in +// context at all, exactly as it would without this package. +func TestUntenantedModelIsUnaffected(t *testing.T) { + db, _, _ := fixture(t) + ctx := context.Background() + + note := plainNote{Title: "single-tenant"} + if err := db.WithContext(ctx).Create(¬e).Error; err != nil { + t.Fatalf("create: %v", err) + } + + var got []plainNote + if err := db.WithContext(ctx).Find(&got).Error; err != nil { + t.Fatalf("find: %v", err) + } + if len(got) != 1 { + t.Fatalf("got %d notes, want 1", len(got)) + } + + if err := db.WithContext(ctx).Delete(&plainNote{}, "id = ?", note.ID).Error; err != nil { + t.Fatalf("delete: %v", err) + } +} + +// TestUnscoped_MustBeCalledFreshPerStatement pins the trap documented on +// Unscoped: its return value behaves like any chained GORM call (db.Where(...) +// has the identical property) — reusing it across two statements mutates the +// first statement's Table and Schema in place rather than starting a new one. +// This is not a bug in Unscoped; it is a property of GORM this package's doc +// exists to warn callers about, demonstrated here so the warning is never +// allowed to go stale. +// +// Reusing it across two differently-shaped models is severe enough to panic +// inside GORM's reflection rather than merely writing wrong data — which is +// arguably the better outcome, since it fails loudly instead of corrupting a +// row. The panic is caught here specifically to demonstrate that this is what +// misuse produces, not to treat it as acceptable. +func TestUnscoped_MustBeCalledFreshPerStatement(t *testing.T) { + db, _, _ := fixture(t) + + misuseCrashed := func() (crashed bool) { + defer func() { + if recover() != nil { + crashed = true + } + }() + // WRONG usage: one instance reused for two creates of different shapes. + reused := tenant.Unscoped(db) + if err := reused.Create(&scopedNote{Title: "first"}).Error; err != nil { + t.Fatalf("first create: %v", err) + } + _ = reused.Create(&plainNote{Title: "second"}).Error + return false + }() + + if !misuseCrashed { + // It did not panic this time, which can happen depending on the two + // models' field counts. Either way the row must not have landed cleanly + // in plainNote's own table. + var wrongTableCount int64 + tenant.Unscoped(db).Model(&plainNote{}).Where("title = ?", "second").Count(&wrongTableCount) + if wrongTableCount == 1 { + t.Fatal("reusing tenant.Unscoped(db) across two Creates worked correctly — " + + "this test's premise is stale, but treat this failure as real until it is " + + "actually verified fixed upstream") + } + } + + // The fix: a fresh call before each statement. Unaffected by the misuse above. + if err := tenant.Unscoped(db).Create(&plainNote{Title: "correct"}).Error; err != nil { + t.Fatalf("fresh Unscoped create: %v", err) + } + var correct plainNote + if err := tenant.Unscoped(db).First(&correct, "title = ?", "correct").Error; err != nil { + t.Fatalf("the correctly-created row is missing: %v", err) + } +} diff --git a/pkg/tenant/tenant.go b/pkg/tenant/tenant.go new file mode 100644 index 0000000..5ca02fb --- /dev/null +++ b/pkg/tenant/tenant.go @@ -0,0 +1,94 @@ +// Package tenant provides optional multi-tenancy: many customers sharing one +// application and one database, with each customer's rows invisible to the others. +// +// # Tenancy is opt-in, per model +// +// Most applications are single-tenant, and a framework that stamped a tenant +// column onto every table would be unusable for them. Nothing here activates +// unless a model asks for it by implementing Tenanted: +// +// func (Customer) TenantColumn() string { return "business_id" } +// +// A model without that method is untouched — no column, no injected predicate, +// no runtime cost. An application that declares no tenant-scoped models behaves +// exactly as if this package did not exist. +// +// Returning the column *name* rather than assuming "tenant_id" lets each +// application keep its own vocabulary: a shop system says business_id, a +// workspace product says org_id. +// +// # Three layers, and why one is not enough +// +// 1. Context — the tenant is resolved once, at the edge, from a trusted source +// such as a JWT claim. Never from a request body, which the caller controls. +// 2. Query scoping — a GORM callback adds the tenant predicate to every read, +// update and delete of a tenanted model. This is the layer that matters: +// hand-written filters get forgotten, and one forgotten filter is a breach. +// 3. Row-level security — PostgreSQL policies refuse foreign rows regardless of +// what the application asks for. A backstop for the day layer 2 has a bug. +// +// Layer 2 fails closed: a tenanted model queried with no tenant in context +// returns an error rather than every row in the table. Legitimate exceptions — +// migrations, admin tooling, background jobs — must say so explicitly with +// Unscoped, which is greppable in review. +package tenant + +import ( + "context" + "errors" + + "github.com/google/uuid" +) + +// Tenanted marks a model as belonging to a tenant. +// +// The single method is the entire opt-in mechanism. Implement it and every query +// against the model is scoped automatically; omit it and nothing changes. +type Tenanted interface { + // TenantColumn returns the column holding the tenant identifier, + // for example "business_id" or "org_id". + TenantColumn() string +} + +// ErrNoTenant is returned when a tenanted model is accessed with no tenant in +// context and no explicit bypass. +// +// This is deliberately an error and not a silent unscoped query. The failure +// mode of guessing is that one customer sees another's data; the failure mode of +// erroring is a broken request, which is noisy, obvious, and safe. +var ErrNoTenant = errors.New("tenant: no tenant in context (use tenant.Unscoped for deliberate cross-tenant access)") + +// tenantKey is an unexported context key type, so no other package can collide +// with it — including one that uses the same string. +type tenantKey struct{} + +// WithTenant returns a context carrying the tenant identifier. +// +// Call this once per request, from a trusted source. Resolving the tenant from +// anything the caller supplies — a body field, a query parameter, a header they +// control — would let one tenant simply ask for another's data. +func WithTenant(ctx context.Context, id uuid.UUID) context.Context { + return context.WithValue(ctx, tenantKey{}, id) +} + +// FromContext returns the tenant identifier and whether one is present. +// +// The boolean must be checked: the zero UUID is a valid-looking value and would +// silently scope queries to a tenant that does not exist. +func FromContext(ctx context.Context) (uuid.UUID, bool) { + id, ok := ctx.Value(tenantKey{}).(uuid.UUID) + if !ok || id == uuid.Nil { + return uuid.Nil, false + } + return id, true +} + +// MustFromContext returns the tenant identifier, or an error if absent. +// Useful in services that cannot meaningfully proceed without one. +func MustFromContext(ctx context.Context) (uuid.UUID, error) { + id, ok := FromContext(ctx) + if !ok { + return uuid.Nil, ErrNoTenant + } + return id, nil +} diff --git a/pkg/testdb/testdb.go b/pkg/testdb/testdb.go new file mode 100644 index 0000000..6ccabc0 --- /dev/null +++ b/pkg/testdb/testdb.go @@ -0,0 +1,306 @@ +// Package testdb provides a real PostgreSQL database for integration tests. +// +// Unit tests verify a function in isolation; these verify the half of the system +// that only exists once a query is actually issued. That distinction matters more +// than it sounds: GORM builds SQL at runtime from struct tags, so a mismatched +// column type or a broken join compiles cleanly, passes `go vet`, and fails only +// when the query runs. No amount of unit testing catches it. +// +// A service built on Tusk uses the same harness for its own schema, by passing +// its migrator to ConnectWith. The three things this package gets right are each +// a bug someone has already had to find: connecting through pgx because that is +// what production uses, serialising database tests across packages because +// `go test ./...` runs them concurrently, and failing rather than skipping in CI. +// +// SQLite is deliberately not used as a stand-in. It diverges from PostgreSQL on +// precisely the behaviour worth testing — row-level security, partial indexes, +// ON CONFLICT semantics, and type strictness — so a green SQLite suite would be +// evidence of very little. +package testdb + +import ( + "context" + "database/sql" + "fmt" + "os" + "strings" + "sync" + "testing" + "time" + + // pgx, not lib/pq: gorm.io/driver/postgres opens production connections + // through stdlib.OpenDB, so a harness on a different driver would test SQL + // that production never sends. They genuinely differ — the migrator's + // introspection queries reuse a placeholder, which pgx accepts and lib/pq + // rejects outright. + _ "github.com/jackc/pgx/v5/stdlib" + "gorm.io/driver/postgres" + "gorm.io/gorm" + "gorm.io/gorm/logger" + + "github.com/codetheuri/tusk/v2/database" + "github.com/codetheuri/tusk/v2/pkg/tenant" +) + +// defaultDSN points at a local development PostgreSQL. +// +// Note the database name: tests TRUNCATE every table between cases, so they must +// never be pointed at a database holding data anyone cares about. Keeping the +// "_test" suffix in the default is a deliberate guard rail. +const defaultDSN = "host=127.0.0.1 port=5434 user=root password=root dbname=tusk_test sslmode=disable TimeZone=UTC" + +// envDSN overrides the default. CI sets this; developers usually do not need to. +const envDSN = "TEST_DATABASE_URL" + +// envRequire turns an unavailable database from a skip into a failure. +// +// This exists because the alternative is worse than having no tests at all: a CI +// run that skips its integration suite reports success while having verified +// nothing. Locally, skipping is the right behaviour — someone reading the code +// should not need PostgreSQL to run `go test ./...`. In CI it must be fatal. +const envRequire = "REQUIRE_DB_TESTS" + +// testLockKey identifies the advisory lock that serialises database tests. +// +// `go test ./...` runs packages concurrently, and every test here TRUNCATEs the +// whole database. Two packages overlapping produces failures that look like +// application bugs — a row disappearing mid-test, a foreign key violated against +// a record that was definitely just written — and that reproduce only sometimes. +// +// A lock in the harness, rather than `-p 1` in the Makefile, because the harness +// is what everyone actually uses: a plain `go test ./...` typed by hand has to be +// correct too. Packages with no database tests still run in parallel. +// +// The value is arbitrary; it only has to be a constant nothing else uses. +// Advisory locks are scoped to the current database, so separate test databases +// do not contend. +const testLockKey int64 = 0x5455534b // "TUSK" + +// lockTimeout bounds the wait for that lock. A hung test that never reports is +// worse than a failing one, especially in CI where it burns the whole job. +const lockTimeout = 2 * time.Minute + +// The lock is held once per test binary and reference-counted, not taken afresh +// on every Connect. A test that connects twice — one database for the owner and +// one for an unprivileged role, say — would otherwise wait on a lock it is +// already holding and stall until the timeout. +var ( + lockMu sync.Mutex + lockConn *sql.Conn + lockDepth int +) + +// Required reports whether a missing database must fail rather than skip. +// CI sets REQUIRE_DB_TESTS so a green build can never mean "tested nothing". +func Required() bool { + return os.Getenv(envRequire) != "" +} + +// DSN returns the connection string tests should use. +func DSN() string { + if dsn := strings.TrimSpace(os.Getenv(envDSN)); dsn != "" { + return dsn + } + return defaultDSN +} + +// Options selects which schema and which database a test runs against. +type Options struct { + // Migrator supplies the migrations to apply. Nil means Tusk's own — what + // Tusk's tests want, and almost never what an application's do. + Migrator *database.Migrator + + // DefaultDSN overrides the fallback connection string. TEST_DATABASE_URL + // still wins over it, so CI keeps one place to point everything at. + DefaultDSN string +} + +// Connect returns a database migrated with Tusk's own schema. +func Connect(t *testing.T) *gorm.DB { + t.Helper() + return ConnectWith(t, Options{}) +} + +// ConnectWith returns a migrated database, or skips the test when none is +// reachable. +// +// The returned database has every table truncated, so each test starts from a +// known empty state. Cleanup is registered with t.Cleanup rather than left to the +// caller — a forgotten teardown leaks state into whichever test runs next, and +// the resulting failure appears in an unrelated place. +func ConnectWith(t *testing.T, opts Options) *gorm.DB { + t.Helper() + + dsn := DSN() + if os.Getenv(envDSN) == "" && opts.DefaultDSN != "" { + dsn = opts.DefaultDSN + } + + migrator := opts.Migrator + if migrator == nil { + migrator = NewTuskMigrator() + } + + sqlDB, err := sql.Open("pgx", dsn) + if err != nil { + skipOrFail(t, fmt.Errorf("open: %w", err)) + return nil + } + + // sql.Open does not connect. Without an explicit ping, an unreachable + // database surfaces later as a confusing query error instead of a clear + // "there is no database here". + sqlDB.SetConnMaxLifetime(time.Minute) + if err := sqlDB.Ping(); err != nil { + _ = sqlDB.Close() + skipOrFail(t, fmt.Errorf("ping %s: %w", redact(dsn), err)) + return nil + } + + if err := migrator.Up(sqlDB); err != nil { + _ = sqlDB.Close() + t.Fatalf("failed to migrate test database: %v", err) + return nil + } + + gormDB, err := gorm.Open(postgres.New(postgres.Config{Conn: sqlDB}), &gorm.Config{ + // Silent: a failing test should be legible. Every statement echoed to + // stdout buries the actual assertion failure. + Logger: logger.Default.LogMode(logger.Silent), + }) + if err != nil { + _ = sqlDB.Close() + t.Fatalf("failed to open gorm connection: %v", err) + return nil + } + + // Mirrors database.Connect. Tests that never opt a model in to tenancy are + // therefore running the same callbacks production does, which is what makes + // them evidence that tenancy stays inert. + if err := tenant.Register(gormDB); err != nil { + _ = sqlDB.Close() + t.Fatalf("failed to register tenant callbacks: %v", err) + return nil + } + + // Serialise before the first TRUNCATE, release after the last one. + lockDatabase(t, sqlDB) + + Truncate(t, gormDB) + t.Cleanup(func() { + Truncate(t, gormDB) + unlockDatabase() + _ = sqlDB.Close() + }) + + return gormDB +} + +// lockDatabase takes the shared advisory lock, or joins one this binary already +// holds. +// +// The connection is pinned deliberately: advisory locks belong to a session, so +// releasing from a different pooled connection would silently do nothing. +func lockDatabase(t *testing.T, sqlDB *sql.DB) { + t.Helper() + + lockMu.Lock() + defer lockMu.Unlock() + + if lockDepth > 0 { + lockDepth++ + return + } + + ctx, cancel := context.WithTimeout(context.Background(), lockTimeout) + defer cancel() + + conn, err := sqlDB.Conn(ctx) + if err != nil { + t.Fatalf("failed to pin a connection for the test lock: %v", err) + } + if _, err := conn.ExecContext(ctx, "SELECT pg_advisory_lock($1)", testLockKey); err != nil { + _ = conn.Close() + t.Fatalf("timed out after %s waiting for the database test lock; "+ + "another test package may be stuck holding it: %v", lockTimeout, err) + } + + lockConn = conn + lockDepth = 1 +} + +func unlockDatabase() { + lockMu.Lock() + defer lockMu.Unlock() + + if lockDepth == 0 { + return + } + lockDepth-- + if lockDepth > 0 { + return + } + + _, _ = lockConn.ExecContext(context.Background(), "SELECT pg_advisory_unlock($1)", testLockKey) + _ = lockConn.Close() + lockConn = nil +} + +// Truncate empties every application table, resetting identity sequences. +// +// TRUNCATE is used rather than wrapping each test in a rolled-back transaction, +// because the code under test opens transactions of its own — the permission +// synchroniser, for one — and nesting those inside a test transaction changes the +// behaviour being verified. +func Truncate(t *testing.T, db *gorm.DB) { + t.Helper() + + var tables []string + err := db.Raw(` + SELECT tablename FROM pg_tables + WHERE schemaname = 'public' AND tablename <> 'goose_db_version' + `).Scan(&tables).Error + if err != nil { + t.Fatalf("failed to list tables for truncation: %v", err) + } + if len(tables) == 0 { + return + } + + for i, name := range tables { + tables[i] = `"` + name + `"` + } + + // CASCADE because foreign keys otherwise make the order significant; RESTART + // IDENTITY so a test asserting on ID 1 is not defeated by a previous run. + stmt := fmt.Sprintf("TRUNCATE TABLE %s RESTART IDENTITY CASCADE", strings.Join(tables, ", ")) + if err := db.Exec(stmt).Error; err != nil { + t.Fatalf("failed to truncate test tables: %v", err) + } +} + +// skipOrFail skips locally and fails in CI, according to REQUIRE_DB_TESTS. +func skipOrFail(t *testing.T, err error) { + t.Helper() + if os.Getenv(envRequire) != "" { + t.Fatalf("%s is set but no test database is reachable: %v", envRequire, err) + } + t.Skipf("no test database reachable, skipping integration test (set %s to make this fatal): %v", envDSN, err) +} + +// redact removes the password before a DSN reaches a log or a test failure. +func redact(dsn string) string { + fields := strings.Fields(dsn) + for i, f := range fields { + if strings.HasPrefix(f, "password=") { + fields[i] = "password=****" + } + } + return strings.Join(fields, " ") +} + +// NewTuskMigrator returns a migrator over Tusk's own schema, for a test that +// wants Tusk's tables alongside an application's. +func NewTuskMigrator() *database.Migrator { + return database.NewMigrator(database.EmbedMigrations, "migrations/postgres") +} diff --git a/pkg/view/view.go b/pkg/view/view.go new file mode 100644 index 0000000..aa84ba2 --- /dev/null +++ b/pkg/view/view.go @@ -0,0 +1,232 @@ +// Package view renders server-rendered HTML pages alongside a Huma API. +// +// Tusk is API-first, but admin panels are a recurring need, and the usual +// answer — ad-hoc html/template calls scattered through handlers — repeats the +// same three mistakes in every project: +// +// 1. Rendering straight to the ResponseWriter. Go templates write output as +// they execute, so an error halfway through a page has already sent a 200 +// and half the HTML. The client receives a truncated page it believes is +// complete, and the error handler can no longer change the status code. +// Engine renders into a buffer and writes nothing until it succeeds. +// +// 2. Parsing on every request, forever. Convenient in development, wasteful in +// production, and the "cache this later" comment never gets revisited. +// Reload makes it a configuration choice with a correct default. +// +// 3. Discovering a broken template when a user visits the page. Templates are +// parsed once at construction, so a malformed page stops the process from +// starting rather than surfacing as a 500 in production. +// +// # Layouts +// +// Go's html/template has no inheritance; it has named blocks that can reference +// each other within a set. Layout inheritance is therefore expressed by parsing +// the shared files together with one page file, and executing whichever block is +// the outermost one: +// +// base.html {{define "base"}}… {{template "content" .}} …{{end}} +// layout.html {{define "sidebar"}}…{{end}} +// dashboard.html +// {{define "title"}}Dashboard{{end}} +// {{define "content"}}…{{template "sidebar" .}}…{{end}} +// +// Each page gets its own template set, so two pages may define "content" +// differently without colliding. +// +// # Files +// +// The engine reads from an fs.FS, so the same code serves an embed.FS in +// production and os.DirFS in development, where Reload picks up edits without a +// restart. +package view + +import ( + "bytes" + "fmt" + "html/template" + "io" + "io/fs" + "net/http" + "path" + "sort" + "sync" +) + +// Config describes an Engine. +type Config struct { + // FS holds the templates. Required. + FS fs.FS + + // Dir is the directory within FS to read, for example "templates/console". + // Empty means the root of FS. + Dir string + + // Layouts are the shared files parsed into every page set, relative to Dir. + // Order matters only in that later definitions of the same block win. + Layouts []string + + // Entry is the defined template executed for a render, typically the one + // wrapping . Defaults to "base". + Entry string + + // Funcs are made available to every template. Register them here rather than + // per-page; a function missing from one set is a parse error in that set + // alone, which is a confusing way to find out. + Funcs template.FuncMap + + // Reload re-reads and re-parses on every render. Development only: it makes + // template edits visible without a restart, at the cost of doing the work + // every time. Drive it from the environment, never hard-code it true. + Reload bool +} + +// Engine renders pages from a set of templates. +type Engine struct { + cfg Config + entry string + + mu sync.RWMutex + cache map[string]*template.Template +} + +// New parses every page in the configured directory and returns an Engine. +// +// Parsing everything up front is the point: a template that does not compile +// should stop deployment, not wait for a user to open that page. +func New(cfg Config) (*Engine, error) { + if cfg.FS == nil { + return nil, fmt.Errorf("view: Config.FS is required") + } + entry := cfg.Entry + if entry == "" { + entry = "base" + } + + e := &Engine{cfg: cfg, entry: entry, cache: map[string]*template.Template{}} + + pages, err := e.pages() + if err != nil { + return nil, err + } + if len(pages) == 0 { + return nil, fmt.Errorf("view: no page templates found in %q", e.dir()) + } + + for _, page := range pages { + tmpl, err := e.parse(page) + if err != nil { + return nil, err + } + if !cfg.Reload { + e.cache[page] = tmpl + } + } + return e, nil +} + +// Render writes a page to the response. +// +// The status is written only once the page has rendered successfully, so a +// template error can still become a 500 instead of a half-sent 200. +func (e *Engine) Render(w http.ResponseWriter, status int, page string, data any) error { + var buf bytes.Buffer + if err := e.RenderTo(&buf, page, data); err != nil { + return err + } + + w.Header().Set("Content-Type", "text/html; charset=utf-8") + w.WriteHeader(status) + _, err := buf.WriteTo(w) + return err +} + +// RenderTo renders to an arbitrary writer, for tests and for HTML mail. +// +// Callers writing to a network connection should render to a buffer first; this +// method makes no such guarantee of its own. +func (e *Engine) RenderTo(w io.Writer, page string, data any) error { + tmpl, err := e.lookup(page) + if err != nil { + return err + } + if err := tmpl.ExecuteTemplate(w, e.entry, data); err != nil { + return fmt.Errorf("view: rendering %q: %w", page, err) + } + return nil +} + +// lookup returns the parsed set for a page, honouring Reload. +func (e *Engine) lookup(page string) (*template.Template, error) { + if e.cfg.Reload { + return e.parse(page) + } + + e.mu.RLock() + tmpl, ok := e.cache[page] + e.mu.RUnlock() + if !ok { + // Not a missing file — New parsed every page that exists, so reaching + // here means the caller named one that does not. + return nil, fmt.Errorf("view: no such page %q", page) + } + return tmpl, nil +} + +// parse builds one page's template set: the layouts plus the page itself. +func (e *Engine) parse(page string) (*template.Template, error) { + files := make([]string, 0, len(e.cfg.Layouts)+1) + for _, l := range e.cfg.Layouts { + files = append(files, e.join(l)) + } + files = append(files, e.join(page)) + + // Named for the page so template errors identify which one failed. + tmpl, err := template.New(page).Funcs(e.cfg.Funcs).ParseFS(e.cfg.FS, files...) + if err != nil { + return nil, fmt.Errorf("view: parsing %q: %w", page, err) + } + if tmpl.Lookup(e.entry) == nil { + return nil, fmt.Errorf("view: %q does not define template %q; "+ + "either the page or one of its layouts must define it", page, e.entry) + } + return tmpl, nil +} + +// pages lists every .html file in Dir that is not a layout. +func (e *Engine) pages() ([]string, error) { + entries, err := fs.ReadDir(e.cfg.FS, e.dir()) + if err != nil { + return nil, fmt.Errorf("view: reading %q: %w", e.dir(), err) + } + + layouts := make(map[string]bool, len(e.cfg.Layouts)) + for _, l := range e.cfg.Layouts { + layouts[l] = true + } + + var pages []string + for _, entry := range entries { + name := entry.Name() + if entry.IsDir() || path.Ext(name) != ".html" || layouts[name] { + continue + } + pages = append(pages, name) + } + sort.Strings(pages) // deterministic startup errors + return pages, nil +} + +func (e *Engine) dir() string { + if e.cfg.Dir == "" { + return "." + } + return e.cfg.Dir +} + +func (e *Engine) join(name string) string { + if e.cfg.Dir == "" { + return name + } + return path.Join(e.cfg.Dir, name) +} diff --git a/pkg/view/view_test.go b/pkg/view/view_test.go new file mode 100644 index 0000000..6c40db0 --- /dev/null +++ b/pkg/view/view_test.go @@ -0,0 +1,237 @@ +package view_test + +import ( + "html/template" + "net/http" + "net/http/httptest" + "strings" + "testing" + "testing/fstest" + + "github.com/codetheuri/tusk/v2/pkg/view" +) + +// templates returns a small but realistic set: a base wrapping the document, a +// layout contributing a shared block, and two pages defining the same block +// names differently. +func templates() fstest.MapFS { + return fstest.MapFS{ + "tpl/base.html": {Data: []byte( + `{{define "base"}}{{template "title" .}}{{template "content" .}}{{end}}`)}, + "tpl/layout.html": {Data: []byte( + `{{define "sidebar"}}{{end}}`)}, + "tpl/dashboard.html": {Data: []byte( + `{{define "title"}}Dashboard{{end}}{{define "content"}}{{template "sidebar" .}}

{{.Heading}}

{{end}}`)}, + "tpl/login.html": {Data: []byte( + `{{define "title"}}Sign in{{end}}{{define "content"}}
{{end}}`)}, + } +} + +func newEngine(t *testing.T, fsys fstest.MapFS) *view.Engine { + t.Helper() + e, err := view.New(view.Config{ + FS: fsys, + Dir: "tpl", + Layouts: []string{"base.html", "layout.html"}, + }) + if err != nil { + t.Fatalf("new engine: %v", err) + } + return e +} + +func TestRender_ComposesLayoutsWithThePage(t *testing.T) { + e := newEngine(t, templates()) + + rec := httptest.NewRecorder() + err := e.Render(rec, http.StatusOK, "dashboard.html", map[string]any{ + "User": "ada", + "Heading": "Overview", + }) + if err != nil { + t.Fatalf("render: %v", err) + } + + body := rec.Body.String() + for _, want := range []string{"Dashboard", "", "

Overview

"} { + if !strings.Contains(body, want) { + t.Errorf("body missing %q\ngot: %s", want, body) + } + } + if ct := rec.Header().Get("Content-Type"); ct != "text/html; charset=utf-8" { + t.Errorf("Content-Type %q", ct) + } +} + +// TestRender_PagesDoNotCollide is why each page gets its own template set. Both +// pages define "title" and "content"; a single shared set would let whichever +// parsed last win for every page. +func TestRender_PagesDoNotCollide(t *testing.T) { + e := newEngine(t, templates()) + + for page, want := range map[string]string{ + "dashboard.html": "Dashboard", + "login.html": "Sign in", + } { + rec := httptest.NewRecorder() + if err := e.Render(rec, http.StatusOK, page, map[string]any{}); err != nil { + t.Fatalf("render %s: %v", page, err) + } + if !strings.Contains(rec.Body.String(), ""+want+"") { + t.Errorf("%s rendered the wrong title: %s", page, rec.Body.String()) + } + } +} + +// TestRender_WritesNothingWhenTheTemplateFails is the defect this package exists +// to prevent. html/template streams as it executes, so a naive implementation +// has already sent 200 and a partial document before the error appears. +func TestRender_WritesNothingWhenTheTemplateFails(t *testing.T) { + fsys := templates() + // A field the data type does not have. Execution fails only once it reaches + // that action — by which point "before" has already been produced, which is + // exactly the situation a streaming renderer cannot recover from. + fsys["tpl/broken.html"] = &fstest.MapFile{Data: []byte( + `{{define "title"}}Broken{{end}}{{define "content"}}
before
{{.NoSuchField}}{{end}}`)} + + e := newEngine(t, fsys) + + rec := httptest.NewRecorder() + err := e.Render(rec, http.StatusOK, "broken.html", struct{ Heading string }{"ok"}) + + if err == nil { + t.Fatal("expected a render error") + } + if rec.Body.Len() != 0 { + t.Errorf("a partial page reached the client: %q", rec.Body.String()) + } + if rec.Flushed { + t.Error("the response was flushed, so the handler can no longer set a status") + } +} + +// TestNew_RejectsABrokenTemplateAtStartup: a page that does not compile must +// stop the process, not wait for a visitor. +func TestNew_RejectsABrokenTemplateAtStartup(t *testing.T) { + fsys := templates() + fsys["tpl/bad.html"] = &fstest.MapFile{Data: []byte(`{{define "content"}}{{ .Unclosed `)} + + _, err := view.New(view.Config{ + FS: fsys, + Dir: "tpl", + Layouts: []string{"base.html", "layout.html"}, + }) + if err == nil { + t.Fatal("expected New to reject a malformed template") + } + if !strings.Contains(err.Error(), "bad.html") { + t.Errorf("error should name the offending file, got: %v", err) + } +} + +// TestNew_RejectsAPageMissingTheEntryBlock catches the mistake of adding a page +// that forgets {{define "content"}} — it would otherwise render as a blank body. +func TestNew_RejectsAPageMissingTheEntryBlock(t *testing.T) { + fsys := fstest.MapFS{ + "tpl/layout.html": {Data: []byte(`{{define "sidebar"}}x{{end}}`)}, + "tpl/page.html": {Data: []byte(`

no blocks at all

`)}, + } + + _, err := view.New(view.Config{FS: fsys, Dir: "tpl", Layouts: []string{"layout.html"}}) + if err == nil { + t.Fatal("expected New to reject a page that defines no entry template") + } +} + +func TestRender_UnknownPageIsAnError(t *testing.T) { + e := newEngine(t, templates()) + + err := e.Render(httptest.NewRecorder(), http.StatusOK, "nope.html", nil) + if err == nil { + t.Fatal("expected an error for an unknown page") + } +} + +// TestReload_PicksUpEdits covers the development path. Without Reload the first +// parse is cached forever, which is correct in production and maddening locally. +func TestReload_PicksUpEdits(t *testing.T) { + fsys := templates() + + e, err := view.New(view.Config{ + FS: fsys, + Dir: "tpl", + Layouts: []string{"base.html", "layout.html"}, + Reload: true, + }) + if err != nil { + t.Fatalf("new engine: %v", err) + } + + rec := httptest.NewRecorder() + if err := e.Render(rec, http.StatusOK, "login.html", nil); err != nil { + t.Fatalf("render: %v", err) + } + if !strings.Contains(rec.Body.String(), "Sign in") { + t.Fatalf("unexpected first render: %s", rec.Body.String()) + } + + fsys["tpl/login.html"] = &fstest.MapFile{Data: []byte( + `{{define "title"}}Edited{{end}}{{define "content"}}
{{end}}`)} + + rec = httptest.NewRecorder() + if err := e.Render(rec, http.StatusOK, "login.html", nil); err != nil { + t.Fatalf("re-render: %v", err) + } + if !strings.Contains(rec.Body.String(), "Edited") { + t.Errorf("Reload did not pick up the edit: %s", rec.Body.String()) + } +} + +// TestRender_EscapesInterpolatedData confirms html/template's contextual +// escaping survives the layout composition. It is the reason for using +// html/template over text/template, and worth a test that would fail loudly if +// anyone ever "fixed" a rendering problem by switching packages. +func TestRender_EscapesInterpolatedData(t *testing.T) { + e := newEngine(t, templates()) + + rec := httptest.NewRecorder() + err := e.Render(rec, http.StatusOK, "dashboard.html", map[string]any{ + "User": "", + "Heading": "ok", + }) + if err != nil { + t.Fatalf("render: %v", err) + } + if strings.Contains(rec.Body.String(), "") { + t.Error("user data was interpolated without escaping") + } + if !strings.Contains(rec.Body.String(), "<script>") { + t.Errorf("expected escaped output, got: %s", rec.Body.String()) + } +} + +func TestFuncs_AreAvailableToEveryPage(t *testing.T) { + fsys := fstest.MapFS{ + "tpl/base.html": {Data: []byte(`{{define "base"}}{{template "content" .}}{{end}}`)}, + "tpl/a.html": {Data: []byte(`{{define "content"}}{{shout "hi"}}{{end}}`)}, + "tpl/b.html": {Data: []byte(`{{define "content"}}{{shout "yo"}}{{end}}`)}, + } + + e, err := view.New(view.Config{ + FS: fsys, + Dir: "tpl", + Layouts: []string{"base.html"}, + Funcs: template.FuncMap{"shout": strings.ToUpper}, + }) + if err != nil { + t.Fatalf("new engine: %v", err) + } + + rec := httptest.NewRecorder() + if err := e.Render(rec, http.StatusOK, "b.html", nil); err != nil { + t.Fatalf("render: %v", err) + } + if rec.Body.String() != "YO" { + t.Errorf("got %q, want YO", rec.Body.String()) + } +}