From b2cb3e55b69f8c36e8125991df4b212acd63e847 Mon Sep 17 00:00:00 2001 From: codetheuri Date: Wed, 9 Sep 2026 16:15:23 +0300 Subject: [PATCH 1/4] refactor!: UUID-native, Postgres-only, and reusable packages out of internal/ MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Hardens Tusk to production grade so a downstream service can be built on it. Developed against Salio, whose backend rewrite is the first consumer. BREAKING CHANGES * Primary keys are UUIDv7, generated in Go by the new pkg/id, never by the database. An offline-first client has to create a row and know its ID before the server has seen it, which no server-side default can serve. Migration 00001 was rewritten rather than layered with a conversion; an existing database must be recreated. Tokens issued before this carry a numeric user_id and will fail to parse, so all sessions are invalidated. * PostgreSQL is the only supported driver. MySQL and SQLite are dropped and LoadConfig refuses them 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. Row-level security is Postgres-only regardless. * internal/middleware -> pkg/middleware, internal/app -> pkg/app, and internal/platform/database/gorm.go -> database/connect.go (NewGoRMDB -> Connect). Go's internal/ rule made these unreachable from a separate module, so a downstream service got pkg/* and config only. internal/auth stays internal deliberately: an application with different authentication should write its own against pkg/authz rather than fight the framework's. * pkg/app.New takes an Options struct and registers no routes of its own. Callers register their modules against App.API(). Also in this change * pkg/tenant: opt-in per-model multi-tenancy in three independent layers — context propagation, GORM callbacks that inject the predicate so a developer cannot forget it, and Postgres RLS so the database refuses cross-tenant rows regardless of application bugs. Opt-in because most applications are single-tenant and tenancy that cannot be switched off would make Tusk useless for them. tenant.VerifyEnforcement reports the trap that superusers and BYPASSRLS roles are exempt from every policy silently. * pkg/testdb: a shared PostgreSQL harness. Database tests across packages serialise on an advisory lock, because go test runs packages concurrently and each truncates the schema — which produced failures that looked like application bugs. * database.Migrator: a reusable migration runner, so a downstream service runs its own migrations through the same code path as the CLI. Adds Reset, which the round-trip migration test needed: it had been reverting a single step and, with only one migration in the tree, appeared to verify the whole rollback path. * lib/pq removed. Config, the migrate CLI and the tests used it while GORM used pgx, so the schema was built by one driver and used by another. * pkg/query: LIKE -> ILIKE. Searching for "wanj" did not match "Wanjiku" on PostgreSQL. Fixes the same latent bug in internal/auth's user search, and adds the package's first tests. Co-Authored-By: Claude Opus 5 --- .env.example | 74 ++- .github/workflows/go.yml | 35 +- Makefile | 25 +- README.md | 1 + cmd/api/main.go | 51 ++- cmd/migrate/main.go | 59 +-- cmd/tusk/main.go | 4 +- config/config.go | 206 ++++++--- config/config_test.go | 223 ++++++++- .../database/gorm.go => database/connect.go | 26 +- database/migrate.go | 173 ++++++- .../00001_identity_auth_rbac.mysql.sql | 86 ---- .../postgres/00001_identity_auth_rbac.sql | 107 +++++ database/migrations_test.go | 129 ++++++ docs/architecture.md | 21 +- docs/database-and-migrations.md | 2 +- go.mod | 8 +- go.sum | 8 - internal/app/app.go | 177 -------- internal/auth/dto.go | 20 +- internal/auth/handler_auth.go | 7 +- internal/auth/integration_test.go | 425 ++++++++++++++++++ internal/auth/model.go | 50 ++- internal/auth/permissions.go | 20 +- internal/auth/repository_auth.go | 4 +- internal/auth/repository_role.go | 16 +- internal/auth/router.go | 2 +- internal/auth/service_auth.go | 13 +- internal/auth/service_role.go | 20 +- internal/middleware/jwt.go | 139 ------ pkg/app/app.go | 310 +++++++++++++ pkg/authz/context.go | 33 ++ pkg/authz/middleware.go | 31 +- pkg/authz/models.go | 37 +- pkg/authz/policy.go | 5 +- pkg/id/id.go | 55 +++ pkg/id/id_test.go | 72 +++ pkg/logger/logger.go | 13 +- pkg/logger/logger_test.go | 76 ++++ pkg/logger/slog.go | 102 +++++ pkg/middleware/bodylimit.go | 23 + {internal => pkg}/middleware/cors.go | 2 +- pkg/middleware/jwt.go | 223 +++++++++ pkg/middleware/jwt_test.go | 229 ++++++++++ {internal => pkg}/middleware/logger.go | 0 .../middleware/middleware_test.go | 6 +- {internal => pkg}/middleware/ratelimit.go | 0 .../middleware/ratelimit_test.go | 0 {internal => pkg}/middleware/recovery.go | 4 +- {internal => pkg}/middleware/requestid.go | 9 + .../middleware/security_headers.go | 6 +- pkg/query/builder.go | 7 +- pkg/query/query_test.go | 146 ++++++ pkg/tenant/middleware.go | 75 ++++ pkg/tenant/middleware_test.go | 99 ++++ pkg/tenant/rls.go | 179 ++++++++ pkg/tenant/rls_test.go | 242 ++++++++++ pkg/tenant/scope.go | 285 ++++++++++++ pkg/tenant/scope_test.go | 385 ++++++++++++++++ pkg/tenant/tenant.go | 94 ++++ pkg/testdb/testdb.go | 306 +++++++++++++ 61 files changed, 4520 insertions(+), 665 deletions(-) rename internal/platform/database/gorm.go => database/connect.go (71%) delete mode 100644 database/migrations/00001_identity_auth_rbac.mysql.sql create mode 100644 database/migrations/postgres/00001_identity_auth_rbac.sql create mode 100644 database/migrations_test.go mode change 100644 => 100755 go.sum delete mode 100644 internal/app/app.go create mode 100644 internal/auth/integration_test.go delete mode 100644 internal/middleware/jwt.go create mode 100644 pkg/app/app.go create mode 100644 pkg/authz/context.go create mode 100644 pkg/id/id.go create mode 100644 pkg/id/id_test.go create mode 100644 pkg/logger/slog.go create mode 100644 pkg/middleware/bodylimit.go rename {internal => pkg}/middleware/cors.go (99%) create mode 100644 pkg/middleware/jwt.go create mode 100644 pkg/middleware/jwt_test.go rename {internal => pkg}/middleware/logger.go (100%) rename {internal => pkg}/middleware/middleware_test.go (96%) rename {internal => pkg}/middleware/ratelimit.go (100%) rename {internal => pkg}/middleware/ratelimit_test.go (100%) rename {internal => pkg}/middleware/recovery.go (95%) rename {internal => pkg}/middleware/requestid.go (56%) rename {internal => pkg}/middleware/security_headers.go (99%) create mode 100644 pkg/query/query_test.go create mode 100644 pkg/tenant/middleware.go create mode 100644 pkg/tenant/middleware_test.go create mode 100644 pkg/tenant/rls.go create mode 100644 pkg/tenant/rls_test.go create mode 100644 pkg/tenant/scope.go create mode 100644 pkg/tenant/scope_test.go create mode 100644 pkg/tenant/tenant.go create mode 100644 pkg/testdb/testdb.go 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/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..9312721 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/internal/auth" + "github.com/codetheuri/tusk/pkg/app" "github.com/codetheuri/tusk/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..8a9cf01 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" + // PostgreSQL driver, registered for its side effects. + _ "github.com/jackc/pgx/v5/stdlib" + "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" ) 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..07a4ecc 100644 --- a/cmd/tusk/main.go +++ b/cmd/tusk/main.go @@ -9,8 +9,8 @@ import ( "github.com/codetheuri/tusk/config" // Import modules to trigger explicit permission registration in init() + "github.com/codetheuri/tusk/database" _ "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" ) @@ -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 71% rename from internal/platform/database/gorm.go rename to database/connect.go index a9c719f..6464d82 100755 --- a/internal/platform/database/gorm.go +++ b/database/connect.go @@ -1,5 +1,11 @@ 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" @@ -7,27 +13,22 @@ import ( "github.com/codetheuri/tusk/config" "github.com/codetheuri/tusk/pkg/logger" + "github.com/codetheuri/tusk/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_test.go b/database/migrations_test.go new file mode 100644 index 0000000..113ac63 --- /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/database" + "github.com/codetheuri/tusk/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/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/go.mod b/go.mod index e567b15..4e29830 100644 --- a/go.mod +++ b/go.mod @@ -4,8 +4,8 @@ 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..c5db869 100644 --- a/internal/auth/dto.go +++ b/internal/auth/dto.go @@ -1,6 +1,8 @@ package auth import ( + "github.com/google/uuid" + "github.com/codetheuri/tusk/pkg/authz" "github.com/codetheuri/tusk/pkg/query" "github.com/codetheuri/tusk/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..1b5c16a 100644 --- a/internal/auth/handler_auth.go +++ b/internal/auth/handler_auth.go @@ -6,6 +6,7 @@ import ( "github.com/danielgtaylor/huma/v2" + "github.com/codetheuri/tusk/pkg/authz" "github.com/codetheuri/tusk/pkg/logger" "github.com/codetheuri/tusk/pkg/query" "github.com/codetheuri/tusk/pkg/response" @@ -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/integration_test.go b/internal/auth/integration_test.go new file mode 100644 index 0000000..8fba6ac --- /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/config" + "github.com/codetheuri/tusk/pkg/authz" + "github.com/codetheuri/tusk/pkg/id" + "github.com/codetheuri/tusk/pkg/logger" + "github.com/codetheuri/tusk/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..93f8888 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/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..6af98c2 100644 --- a/internal/auth/permissions.go +++ b/internal/auth/permissions.go @@ -4,17 +4,17 @@ import "github.com/codetheuri/tusk/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..5f46bd5 100644 --- a/internal/auth/repository_auth.go +++ b/internal/auth/repository_auth.go @@ -1,6 +1,8 @@ package auth import ( + "github.com/google/uuid" + "context" "time" @@ -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..4e592c1 100644 --- a/internal/auth/repository_role.go +++ b/internal/auth/repository_role.go @@ -1,6 +1,8 @@ package auth import ( + "github.com/google/uuid" + "context" "github.com/codetheuri/tusk/pkg/authz" @@ -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..e926569 100644 --- a/internal/auth/router.go +++ b/internal/auth/router.go @@ -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..6651211 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" @@ -12,8 +14,8 @@ import ( "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/middleware" "github.com/codetheuri/tusk/pkg/query" ) @@ -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..957895d 100644 --- a/internal/auth/service_role.go +++ b/internal/auth/service_role.go @@ -1,6 +1,8 @@ package auth import ( + "github.com/google/uuid" + "context" "fmt" @@ -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..c66e4a1 --- /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/config" + "github.com/codetheuri/tusk/database" + "github.com/codetheuri/tusk/pkg/logger" + "github.com/codetheuri/tusk/pkg/middleware" + "github.com/codetheuri/tusk/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..feadf2f 100644 --- a/pkg/authz/middleware.go +++ b/pkg/authz/middleware.go @@ -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..f0e32f7 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/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/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 99% rename from internal/middleware/cors.go rename to pkg/middleware/cors.go index 26593d4..3b00fad 100755 --- a/internal/middleware/cors.go +++ b/pkg/middleware/cors.go @@ -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..999e61a --- /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/pkg/authz" + "github.com/codetheuri/tusk/pkg/response" + "github.com/codetheuri/tusk/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..f6f513d --- /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/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 100% rename from internal/middleware/logger.go rename to pkg/middleware/logger.go 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 100% rename from internal/middleware/ratelimit.go rename to pkg/middleware/ratelimit.go 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 95% rename from internal/middleware/recovery.go rename to pkg/middleware/recovery.go index b9f81a5..ef56879 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" - ) // 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..0b6eda4 --- /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/pkg/query" + "github.com/codetheuri/tusk/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/tenant/middleware.go b/pkg/tenant/middleware.go new file mode 100644 index 0000000..04c72e7 --- /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/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..8f99762 --- /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/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..6fa9ed9 --- /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/pkg/id" + "github.com/codetheuri/tusk/pkg/tenant" + "github.com/codetheuri/tusk/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..14ae2ea --- /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/pkg/id" + "github.com/codetheuri/tusk/pkg/tenant" + "github.com/codetheuri/tusk/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..fea7a1b --- /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/database" + "github.com/codetheuri/tusk/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") +} From e44c240910c474d6ec999183003da8ddbeb355f1 Mon Sep 17 00:00:00 2001 From: codetheuri Date: Wed, 9 Sep 2026 16:15:23 +0300 Subject: [PATCH 2/4] feat: server-rendered pages and database-backed sessions MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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 mistakes in every project. Both packages are additive; nothing else imports them. pkg/view renders into a buffer and writes nothing until the page succeeds, so a template error halfway through can still become a 500 instead of a truncated 200 the client believes is complete. Every page is parsed at construction, so a malformed template stops the process starting rather than surfacing when a user opens that page. Reload makes per-request parsing a configuration choice with a correct default instead of a "cache this later" comment nobody revisits. pkg/session stores only the SHA-256 digest of the token the browser holds, so a leaked backup does not hand over live sessions. SHA-256 rather than bcrypt, for the opposite reason to passwords: the token is 256 bits of randomness, so it is unguessable regardless of hash speed, and it is verified on every request, where a deliberately slow hash would be a denial-of-service surface. Middleware loads a session without requiring one; Require refuses. Keeping them separate is what stops 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 database, and Session.SubjectID carries no foreign key on purpose — the subject may live in Tusk's own users table or in an application's separate operators table, and this package has no business deciding which. Co-Authored-By: Claude Opus 5 --- .../migrations/postgres/00002_sessions.sql | 42 ++ pkg/session/session.go | 381 ++++++++++++++++ pkg/session/session_test.go | 411 ++++++++++++++++++ pkg/session/store.go | 132 ++++++ pkg/view/view.go | 232 ++++++++++ pkg/view/view_test.go | 237 ++++++++++ 6 files changed, 1435 insertions(+) create mode 100644 database/migrations/postgres/00002_sessions.sql create mode 100644 pkg/session/session.go create mode 100644 pkg/session/session_test.go create mode 100644 pkg/session/store.go create mode 100644 pkg/view/view.go create mode 100644 pkg/view/view_test.go 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/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..fc7aaef --- /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/pkg/session" + "github.com/codetheuri/tusk/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..c5b632c --- /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/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/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..d61ea94 --- /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/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()) + } +} From 7c8ddf10f2b9ffaf95c22a1d7280a3b66950507b Mon Sep 17 00:00:00 2001 From: codetheuri Date: Wed, 9 Sep 2026 16:15:23 +0300 Subject: [PATCH 3/4] docs: roadmap and agent guide roadmap.md records what changed in each hardening pass and why, including the behaviour breaks and the two findings worth carrying forward: that a superuser is exempt from every RLS policy silently, and that the test harness had been exercising a different driver than production. CLAUDE.md gives an agent working in this repo the conventions and the layering rules up front, rather than having them inferred from whichever file is open. Co-Authored-By: Claude Opus 5 --- CLAUDE.md | 118 +++++++++++++++++++++++++++++ docs/roadmap.md | 193 ++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 311 insertions(+) create mode 100644 CLAUDE.md create mode 100644 docs/roadmap.md 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/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. From d570ffc3c8ec4e2e78b40d9b61f99abe138052d8 Mon Sep 17 00:00:00 2001 From: codetheuri Date: Wed, 9 Sep 2026 16:15:23 +0300 Subject: [PATCH 4/4] build!: move to the /v2 module path MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit v1.0.0 is published, and the preceding commits are breaking: UUID primary keys, internal/ packages promoted to pkg/, NewGoRMDB renamed to database.Connect, and MySQL and SQLite dropped. Go resolves a module's major version from its path, so github.com/codetheuri/tusk can only ever serve v1.x — tagging a breaking release there would hand a broken build to anyone on v1.0.0 who ran `go get -u`, with no warning and no way to pin around it. The suffix is what makes the two versions independent: v1.0.0 keeps resolving for anyone still on it, and v2 is opted into by changing an import path. Import paths only; no behaviour changes in this commit. Co-Authored-By: Claude Opus 5 --- cmd/api/main.go | 8 ++++---- cmd/migrate/main.go | 4 ++-- cmd/tusk/main.go | 10 +++++----- database/connect.go | 6 +++--- database/migrations_test.go | 4 ++-- docs/authorization.md | 2 +- go.mod | 2 +- internal/auth/dto.go | 6 +++--- internal/auth/handler_auth.go | 8 ++++---- internal/auth/handler_role.go | 2 +- internal/auth/integration_test.go | 10 +++++----- internal/auth/model.go | 2 +- internal/auth/permissions.go | 2 +- internal/auth/repository_auth.go | 4 ++-- internal/auth/repository_role.go | 2 +- internal/auth/router.go | 6 +++--- internal/auth/service_auth.go | 8 ++++---- internal/auth/service_role.go | 2 +- pkg/app/app.go | 10 +++++----- pkg/authz/middleware.go | 2 +- pkg/authz/models.go | 2 +- pkg/mailer/mailer.go | 4 ++-- pkg/middleware/cors.go | 2 +- pkg/middleware/jwt.go | 6 +++--- pkg/middleware/jwt_test.go | 2 +- pkg/middleware/logger.go | 2 +- pkg/middleware/ratelimit.go | 2 +- pkg/middleware/recovery.go | 2 +- pkg/query/query_test.go | 4 ++-- pkg/session/session_test.go | 4 ++-- pkg/session/store.go | 2 +- pkg/tenant/middleware.go | 2 +- pkg/tenant/middleware_test.go | 2 +- pkg/tenant/rls_test.go | 6 +++--- pkg/tenant/scope_test.go | 6 +++--- pkg/testdb/testdb.go | 4 ++-- pkg/view/view_test.go | 2 +- 37 files changed, 77 insertions(+), 77 deletions(-) diff --git a/cmd/api/main.go b/cmd/api/main.go index 9312721..4d5869b 100755 --- a/cmd/api/main.go +++ b/cmd/api/main.go @@ -4,10 +4,10 @@ package main import ( "github.com/danielgtaylor/huma/v2" - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/internal/auth" - "github.com/codetheuri/tusk/pkg/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() { diff --git a/cmd/migrate/main.go b/cmd/migrate/main.go index 8a9cf01..53fc1d1 100755 --- a/cmd/migrate/main.go +++ b/cmd/migrate/main.go @@ -13,8 +13,8 @@ import ( // PostgreSQL driver, registered for its side effects. _ "github.com/jackc/pgx/v5/stdlib" - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/database" + "github.com/codetheuri/tusk/v2/config" + "github.com/codetheuri/tusk/v2/database" ) func main() { diff --git a/cmd/tusk/main.go b/cmd/tusk/main.go index 07a4ecc..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/database" - _ "github.com/codetheuri/tusk/internal/auth" - "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() { diff --git a/database/connect.go b/database/connect.go index 6464d82..19e76c9 100755 --- a/database/connect.go +++ b/database/connect.go @@ -11,9 +11,9 @@ import ( "fmt" "time" - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/pkg/logger" - "github.com/codetheuri/tusk/pkg/tenant" + "github.com/codetheuri/tusk/v2/config" + "github.com/codetheuri/tusk/v2/pkg/logger" + "github.com/codetheuri/tusk/v2/pkg/tenant" "gorm.io/driver/postgres" "gorm.io/gorm" diff --git a/database/migrations_test.go b/database/migrations_test.go index 113ac63..ce5c955 100644 --- a/database/migrations_test.go +++ b/database/migrations_test.go @@ -7,8 +7,8 @@ import ( _ "github.com/jackc/pgx/v5/stdlib" - "github.com/codetheuri/tusk/database" - "github.com/codetheuri/tusk/pkg/testdb" + "github.com/codetheuri/tusk/v2/database" + "github.com/codetheuri/tusk/v2/pkg/testdb" ) // migrationTestDB is created and dropped by this test. 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/go.mod b/go.mod index 4e29830..e00306b 100644 --- a/go.mod +++ b/go.mod @@ -1,4 +1,4 @@ -module github.com/codetheuri/tusk +module github.com/codetheuri/tusk/v2 go 1.25.7 diff --git a/internal/auth/dto.go b/internal/auth/dto.go index c5db869..e79798e 100644 --- a/internal/auth/dto.go +++ b/internal/auth/dto.go @@ -3,9 +3,9 @@ package auth import ( "github.com/google/uuid" - "github.com/codetheuri/tusk/pkg/authz" - "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/query" + "github.com/codetheuri/tusk/v2/pkg/response" ) // ------------------------------------------------------------- diff --git a/internal/auth/handler_auth.go b/internal/auth/handler_auth.go index 1b5c16a..bb5fe6d 100644 --- a/internal/auth/handler_auth.go +++ b/internal/auth/handler_auth.go @@ -6,10 +6,10 @@ import ( "github.com/danielgtaylor/huma/v2" - "github.com/codetheuri/tusk/pkg/authz" - "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 { 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 index 8fba6ac..cb8eb5b 100644 --- a/internal/auth/integration_test.go +++ b/internal/auth/integration_test.go @@ -7,11 +7,11 @@ import ( "gorm.io/gorm" - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/pkg/authz" - "github.com/codetheuri/tusk/pkg/id" - "github.com/codetheuri/tusk/pkg/logger" - "github.com/codetheuri/tusk/pkg/testdb" + "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. diff --git a/internal/auth/model.go b/internal/auth/model.go index 93f8888..ee62960 100644 --- a/internal/auth/model.go +++ b/internal/auth/model.go @@ -6,7 +6,7 @@ import ( "github.com/google/uuid" "gorm.io/gorm" - "github.com/codetheuri/tusk/pkg/id" + "github.com/codetheuri/tusk/v2/pkg/id" ) // User handles core authentication data, credentials, and security state. diff --git a/internal/auth/permissions.go b/internal/auth/permissions.go index 6af98c2..34a5c3f 100644 --- a/internal/auth/permissions.go +++ b/internal/auth/permissions.go @@ -1,6 +1,6 @@ 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 ( diff --git a/internal/auth/repository_auth.go b/internal/auth/repository_auth.go index 5f46bd5..03218da 100644 --- a/internal/auth/repository_auth.go +++ b/internal/auth/repository_auth.go @@ -6,8 +6,8 @@ import ( "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" ) diff --git a/internal/auth/repository_role.go b/internal/auth/repository_role.go index 4e592c1..a7f2eef 100644 --- a/internal/auth/repository_role.go +++ b/internal/auth/repository_role.go @@ -5,7 +5,7 @@ import ( "context" - "github.com/codetheuri/tusk/pkg/authz" + "github.com/codetheuri/tusk/v2/pkg/authz" "gorm.io/gorm" ) diff --git a/internal/auth/router.go b/internal/auth/router.go index e926569..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) { diff --git a/internal/auth/service_auth.go b/internal/auth/service_auth.go index 6651211..d070575 100644 --- a/internal/auth/service_auth.go +++ b/internal/auth/service_auth.go @@ -13,10 +13,10 @@ import ( "github.com/golang-jwt/jwt/v5" "golang.org/x/crypto/bcrypt" - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/pkg/authz" - "github.com/codetheuri/tusk/pkg/middleware" - "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. diff --git a/internal/auth/service_role.go b/internal/auth/service_role.go index 957895d..0e5faa3 100644 --- a/internal/auth/service_role.go +++ b/internal/auth/service_role.go @@ -6,7 +6,7 @@ import ( "context" "fmt" - "github.com/codetheuri/tusk/pkg/authz" + "github.com/codetheuri/tusk/v2/pkg/authz" ) type CreateRoleRequest struct { diff --git a/pkg/app/app.go b/pkg/app/app.go index c66e4a1..dc3b2a8 100644 --- a/pkg/app/app.go +++ b/pkg/app/app.go @@ -35,11 +35,11 @@ import ( "github.com/go-chi/chi/v5" "gorm.io/gorm" - "github.com/codetheuri/tusk/config" - "github.com/codetheuri/tusk/database" - "github.com/codetheuri/tusk/pkg/logger" - "github.com/codetheuri/tusk/pkg/middleware" - "github.com/codetheuri/tusk/pkg/response" + "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. diff --git a/pkg/authz/middleware.go b/pkg/authz/middleware.go index feadf2f..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. diff --git a/pkg/authz/models.go b/pkg/authz/models.go index f0e32f7..12eaf34 100644 --- a/pkg/authz/models.go +++ b/pkg/authz/models.go @@ -6,7 +6,7 @@ import ( "github.com/google/uuid" "gorm.io/gorm" - "github.com/codetheuri/tusk/pkg/id" + "github.com/codetheuri/tusk/v2/pkg/id" ) // PermissionRecord represents the runtime database representation of a code permission. 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/cors.go b/pkg/middleware/cors.go index 3b00fad..d55aaa1 100755 --- a/pkg/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 { diff --git a/pkg/middleware/jwt.go b/pkg/middleware/jwt.go index 999e61a..5a1fc20 100644 --- a/pkg/middleware/jwt.go +++ b/pkg/middleware/jwt.go @@ -14,9 +14,9 @@ import ( "github.com/golang-jwt/jwt/v5" "github.com/google/uuid" - "github.com/codetheuri/tusk/pkg/authz" - "github.com/codetheuri/tusk/pkg/response" - "github.com/codetheuri/tusk/pkg/tenant" + "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. diff --git a/pkg/middleware/jwt_test.go b/pkg/middleware/jwt_test.go index f6f513d..9bba206 100644 --- a/pkg/middleware/jwt_test.go +++ b/pkg/middleware/jwt_test.go @@ -14,7 +14,7 @@ import ( "github.com/golang-jwt/jwt/v5" "github.com/google/uuid" - "github.com/codetheuri/tusk/pkg/authz" + "github.com/codetheuri/tusk/v2/pkg/authz" ) const testSecret = "test-secret-key-that-is-long-enough-for-hs256" diff --git a/pkg/middleware/logger.go b/pkg/middleware/logger.go index 96b8b64..34ba114 100755 --- a/pkg/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/pkg/middleware/ratelimit.go b/pkg/middleware/ratelimit.go index e946535..9ce85f5 100644 --- a/pkg/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/pkg/middleware/recovery.go b/pkg/middleware/recovery.go index ef56879..883883c 100755 --- a/pkg/middleware/recovery.go +++ b/pkg/middleware/recovery.go @@ -5,7 +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 diff --git a/pkg/query/query_test.go b/pkg/query/query_test.go index 0b6eda4..dfc9363 100644 --- a/pkg/query/query_test.go +++ b/pkg/query/query_test.go @@ -7,8 +7,8 @@ import ( "github.com/google/uuid" "gorm.io/gorm" - "github.com/codetheuri/tusk/pkg/query" - "github.com/codetheuri/tusk/pkg/testdb" + "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 diff --git a/pkg/session/session_test.go b/pkg/session/session_test.go index fc7aaef..52853a4 100644 --- a/pkg/session/session_test.go +++ b/pkg/session/session_test.go @@ -10,8 +10,8 @@ import ( "github.com/google/uuid" "gorm.io/gorm" - "github.com/codetheuri/tusk/pkg/session" - "github.com/codetheuri/tusk/pkg/testdb" + "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 diff --git a/pkg/session/store.go b/pkg/session/store.go index c5b632c..85ffffc 100644 --- a/pkg/session/store.go +++ b/pkg/session/store.go @@ -8,7 +8,7 @@ import ( "github.com/google/uuid" "gorm.io/gorm" - "github.com/codetheuri/tusk/pkg/id" + "github.com/codetheuri/tusk/v2/pkg/id" ) // Session is one signed-in browser. diff --git a/pkg/tenant/middleware.go b/pkg/tenant/middleware.go index 04c72e7..35b9460 100644 --- a/pkg/tenant/middleware.go +++ b/pkg/tenant/middleware.go @@ -6,7 +6,7 @@ import ( "github.com/danielgtaylor/huma/v2" "github.com/google/uuid" - "github.com/codetheuri/tusk/pkg/response" + "github.com/codetheuri/tusk/v2/pkg/response" ) // Resolver determines the active tenant for a request. diff --git a/pkg/tenant/middleware_test.go b/pkg/tenant/middleware_test.go index 8f99762..8075ecc 100644 --- a/pkg/tenant/middleware_test.go +++ b/pkg/tenant/middleware_test.go @@ -8,7 +8,7 @@ import ( "github.com/google/uuid" - "github.com/codetheuri/tusk/pkg/tenant" + "github.com/codetheuri/tusk/v2/pkg/tenant" ) // serve runs one request through the middleware and reports what the handler saw. diff --git a/pkg/tenant/rls_test.go b/pkg/tenant/rls_test.go index 6fa9ed9..6d7f1ea 100644 --- a/pkg/tenant/rls_test.go +++ b/pkg/tenant/rls_test.go @@ -13,9 +13,9 @@ import ( "gorm.io/gorm" gormlogger "gorm.io/gorm/logger" - "github.com/codetheuri/tusk/pkg/id" - "github.com/codetheuri/tusk/pkg/tenant" - "github.com/codetheuri/tusk/pkg/testdb" + "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 diff --git a/pkg/tenant/scope_test.go b/pkg/tenant/scope_test.go index 14ae2ea..9f0e19b 100644 --- a/pkg/tenant/scope_test.go +++ b/pkg/tenant/scope_test.go @@ -8,9 +8,9 @@ import ( "github.com/google/uuid" "gorm.io/gorm" - "github.com/codetheuri/tusk/pkg/id" - "github.com/codetheuri/tusk/pkg/tenant" - "github.com/codetheuri/tusk/pkg/testdb" + "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 diff --git a/pkg/testdb/testdb.go b/pkg/testdb/testdb.go index fea7a1b..6ccabc0 100644 --- a/pkg/testdb/testdb.go +++ b/pkg/testdb/testdb.go @@ -38,8 +38,8 @@ import ( "gorm.io/gorm" "gorm.io/gorm/logger" - "github.com/codetheuri/tusk/database" - "github.com/codetheuri/tusk/pkg/tenant" + "github.com/codetheuri/tusk/v2/database" + "github.com/codetheuri/tusk/v2/pkg/tenant" ) // defaultDSN points at a local development PostgreSQL. diff --git a/pkg/view/view_test.go b/pkg/view/view_test.go index d61ea94..6c40db0 100644 --- a/pkg/view/view_test.go +++ b/pkg/view/view_test.go @@ -8,7 +8,7 @@ import ( "testing" "testing/fstest" - "github.com/codetheuri/tusk/pkg/view" + "github.com/codetheuri/tusk/v2/pkg/view" ) // templates returns a small but realistic set: a base wrapping the document, a