From ddb45005a660ab7c9a6e439653abb02465d7a16e Mon Sep 17 00:00:00 2001 From: Manas Srivastava Date: Mon, 11 May 2026 13:32:53 +0530 Subject: [PATCH] =?UTF-8?q?chore:=20sync=20prod=20state=20to=20master=20?= =?UTF-8?q?=E2=80=94=20vault,=20custom=20domains,=20magic=20links,=20kanik?= =?UTF-8?q?o,=20dpop,=20rbac?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit This is a catch-up commit: the master branch had drifted significantly from the code actually running in production. Bundling the divergence here so subsequent PRs (deploy compute correctness, hostname injection, build context delivery, etc) can branch off a master that reflects reality. Each theme below is internally coherent; reviewers may want to read by section. Migrations (new, applied in prod): - 008_vault.sql — vault_secrets + vault_audit_log - 009_env_column.sql — env scope column on resources/deployments - 010_team_invitations.sql — pending-invite rows - 011_api_keys.sql — programmatic agent tokens - 012_audit_log.sql — generic audit log - 013_magic_links.sql — email magic-link auth - 014_custom_domains.sql — *.custom.tld pointing at deploys New handlers: - vault.go, vault_resolve.go — encrypted env-var store + vault://KEY ref resolution - api_keys.go — programmatic token issuance - audit.go — read audit log - custom_domain.go — custom domain CRUD - magic_link.go — email magic-link issue + redeem - teams.go — team-member CRUD - wellknown.go — RFC 8615 /.well-known/oauth-protected-resource - team_members.go — invite + role management New middleware: - api_key.go — bearer-API-key extraction - dpop.go — RFC 9449 DPoP proof binding (cnf / jkt claims) - quota.go — per-team rate limiting - rbac.go + role_lookup.go — role-based access for team-scoped endpoints - auth.go — audience (aud) + cnf validation added New models: - vault.go, api_key.go, audit_log.go, custom_domain.go, magic_link.go, team_invitations.go - deployment.go + resource.go gain env column + helper methods Compute provider rewrite (internal/providers/compute/k8s/): - client.go — kaniko in-cluster build (replaces the Rancher-Desktop "docker build" fallback). Builds via kaniko Job with build-context Secret; pushes to ghcr.io/.../instant-userapp/. - custom_domain.go — Ingress + cert-manager Certificate reconcile for *.custom.tld - stack.go — multi-service stack provisioning Plans: - razorpay.go — subscription create/cancel/update via Razorpay API (matches the live billing handler) Other: - config.go — new env vars for vault key, magic-link expiry, custom-domain wildcard cert, kaniko image, etc. - main.go — wire all new routes + middleware - plans.yaml — vault + deployment limits added - email.go — magic-link template + transactional send Verified live in prod; deferred to git for too long. Subsequent PRs will fix specific friction points discovered while testing. Co-Authored-By: Claude Opus 4.7 (1M context) --- .gitignore | 1 + e2e/fixtures/hello-app/Dockerfile | 13 + e2e/fixtures/hello-app/index.html | 8 + e2e/helpers_test.go | 36 ++ e2e/merged_surfaces_e2e_test.go | 146 +++++ go.mod | 8 + go.sum | 18 + internal/config/config.go | 15 +- internal/db/migrations/008_vault.sql | 34 + internal/db/migrations/009_env_column.sql | 16 + .../db/migrations/010_team_invitations.sql | 48 ++ internal/db/migrations/011_api_keys.sql | 30 + internal/db/migrations/012_audit_log.sql | 16 + internal/db/migrations/013_magic_links.sql | 29 + internal/db/migrations/014_custom_domains.sql | 37 ++ internal/email/email.go | 46 ++ internal/handlers/api_keys.go | 163 +++++ internal/handlers/audit.go | 112 ++++ internal/handlers/auth.go | 358 +++++++++++ internal/handlers/billing.go | 112 +++- internal/handlers/cache.go | 46 +- internal/handlers/custom_domain.go | 594 ++++++++++++++++++ internal/handlers/db.go | 28 +- internal/handlers/deploy.go | 35 +- internal/handlers/env_test.go | 184 ++++++ internal/handlers/magic_link.go | 176 ++++++ internal/handlers/nosql.go | 46 +- internal/handlers/openapi.go | 269 +++++++- internal/handlers/provision_helper.go | 26 + internal/handlers/queue.go | 34 +- internal/handlers/resource.go | 64 ++ internal/handlers/stack.go | 119 +++- internal/handlers/storage.go | 29 +- internal/handlers/team_members.go | 87 ++- internal/handlers/teams.go | 253 ++++++++ internal/handlers/teams_test.go | 307 +++++++++ internal/handlers/vault.go | 440 +++++++++++++ internal/handlers/vault_resolve.go | 98 +++ internal/handlers/vault_resolve_test.go | 181 ++++++ internal/handlers/vault_test.go | 578 +++++++++++++++++ internal/handlers/webhook.go | 26 +- internal/handlers/wellknown.go | 90 +++ internal/handlers/wellknown_test.go | 75 +++ internal/middleware/api_key.go | 104 +++ internal/middleware/auth.go | 157 ++++- internal/middleware/auth_audience_test.go | 145 +++++ internal/middleware/dpop.go | 274 ++++++++ internal/middleware/dpop_test.go | 321 ++++++++++ internal/middleware/fingerprint.go | 73 ++- internal/middleware/quota.go | 56 ++ internal/middleware/quota_test.go | 71 +++ internal/middleware/rbac.go | 92 +++ internal/middleware/rbac_test.go | 143 +++++ internal/middleware/role_lookup.go | 70 +++ internal/models/api_key.go | 169 +++++ internal/models/audit_log.go | 122 ++++ internal/models/custom_domain.go | 304 +++++++++ internal/models/deployment.go | 73 ++- internal/models/deployment_env_test.go | 116 ++++ internal/models/magic_link.go | 115 ++++ internal/models/resource.go | 164 +++-- internal/models/resource_env_test.go | 212 +++++++ internal/models/team.go | 8 +- internal/models/team_invitations.go | 338 ++++++++++ internal/models/vault.go | 205 ++++++ internal/plans/razorpay.go | 46 ++ internal/plans/razorpay_test.go | 97 +++ internal/providers/cache/redis.go | 5 + internal/providers/compute/k8s/client.go | 462 ++++++++++++-- .../providers/compute/k8s/custom_domain.go | 310 +++++++++ internal/providers/compute/k8s/stack.go | 55 +- internal/providers/nosql/mongo.go | 5 + internal/providers/storage/local.go | 57 +- internal/providers/storage/local_test.go | 14 +- internal/provisioner/client.go | 9 +- internal/router/router.go | 102 ++- main.go | 2 +- plans.yaml | 29 +- 78 files changed, 9322 insertions(+), 234 deletions(-) create mode 100644 e2e/fixtures/hello-app/Dockerfile create mode 100644 e2e/fixtures/hello-app/index.html create mode 100644 e2e/merged_surfaces_e2e_test.go create mode 100644 internal/db/migrations/008_vault.sql create mode 100644 internal/db/migrations/009_env_column.sql create mode 100644 internal/db/migrations/010_team_invitations.sql create mode 100644 internal/db/migrations/011_api_keys.sql create mode 100644 internal/db/migrations/012_audit_log.sql create mode 100644 internal/db/migrations/013_magic_links.sql create mode 100644 internal/db/migrations/014_custom_domains.sql create mode 100644 internal/handlers/api_keys.go create mode 100644 internal/handlers/audit.go create mode 100644 internal/handlers/custom_domain.go create mode 100644 internal/handlers/env_test.go create mode 100644 internal/handlers/magic_link.go create mode 100644 internal/handlers/teams.go create mode 100644 internal/handlers/teams_test.go create mode 100644 internal/handlers/vault.go create mode 100644 internal/handlers/vault_resolve.go create mode 100644 internal/handlers/vault_resolve_test.go create mode 100644 internal/handlers/vault_test.go create mode 100644 internal/handlers/wellknown.go create mode 100644 internal/handlers/wellknown_test.go create mode 100644 internal/middleware/api_key.go create mode 100644 internal/middleware/auth_audience_test.go create mode 100644 internal/middleware/dpop.go create mode 100644 internal/middleware/dpop_test.go create mode 100644 internal/middleware/quota.go create mode 100644 internal/middleware/quota_test.go create mode 100644 internal/middleware/rbac.go create mode 100644 internal/middleware/rbac_test.go create mode 100644 internal/middleware/role_lookup.go create mode 100644 internal/models/api_key.go create mode 100644 internal/models/audit_log.go create mode 100644 internal/models/custom_domain.go create mode 100644 internal/models/deployment_env_test.go create mode 100644 internal/models/magic_link.go create mode 100644 internal/models/resource_env_test.go create mode 100644 internal/models/team_invitations.go create mode 100644 internal/models/vault.go create mode 100644 internal/plans/razorpay.go create mode 100644 internal/plans/razorpay_test.go create mode 100644 internal/providers/compute/k8s/custom_domain.go diff --git a/.gitignore b/.gitignore index 2b2d14ee..d8b547e4 100644 --- a/.gitignore +++ b/.gitignore @@ -4,3 +4,4 @@ GeoLite2-*.mmdb .env .env.* !.env.example +node_modules diff --git a/e2e/fixtures/hello-app/Dockerfile b/e2e/fixtures/hello-app/Dockerfile new file mode 100644 index 00000000..b451408f --- /dev/null +++ b/e2e/fixtures/hello-app/Dockerfile @@ -0,0 +1,13 @@ +# Minimal hello-world image for deploy E2E test. +# +# We use busybox httpd for three reasons: +# 1. Smallest possible image (~1MB vs ~5MB alpine vs ~300MB Go) — fastest pull on local k3s +# 2. No build step (unlike a Go binary) — fastest build on slow buildkit / kaniko +# 3. Most reliable: busybox httpd has zero deps, runs as PID 1, handles SIGTERM cleanly +# +# Listens on 8080 so the deploy E2E can pass port=8080 and verify that the +# container port is correctly wired through to the public URL. +FROM busybox:1.36 +COPY index.html /index.html +EXPOSE 8080 +CMD ["httpd", "-f", "-p", "8080", "-h", "/"] diff --git a/e2e/fixtures/hello-app/index.html b/e2e/fixtures/hello-app/index.html new file mode 100644 index 00000000..00681af6 --- /dev/null +++ b/e2e/fixtures/hello-app/index.html @@ -0,0 +1,8 @@ + + +instanode hello + +

hello from instanode

+

This page is served by the deploy E2E fixture (busybox httpd) at port 8080.

+ + diff --git a/e2e/helpers_test.go b/e2e/helpers_test.go index 84464ad2..879f1b9d 100644 --- a/e2e/helpers_test.go +++ b/e2e/helpers_test.go @@ -48,6 +48,15 @@ func TestMain(m *testing.M) { os.Exit(m.Run()) } +// e2eTestToken returns the shared secret used to override the production +// fingerprint middleware's source-IP selection (see middleware/fingerprint.go). +// When E2E_TEST_TOKEN is set on both the cluster (env) and the test runner, +// the test runner's X-Forwarded-For is honored as the leftmost entry, +// restoring per-test fingerprint isolation against the live cluster. +func e2eTestToken() string { + return os.Getenv("E2E_TEST_TOKEN") +} + // ipSeq is an atomic counter incremented per uniqueSubnet/uniqueIP call. // It guarantees distinct /24 subnets within a single binary run. var ipSeq atomic.Int64 @@ -114,6 +123,15 @@ func getNoRedirect(t *testing.T, path string, headers ...string) *http.Response for i := 0; i+1 < len(headers); i += 2 { req.Header.Set(headers[i], headers[i+1]) } + if tok := e2eTestToken(); tok != "" && req.Header.Get("X-E2E-Test-Token") == "" { + req.Header.Set("X-E2E-Test-Token", tok) + // Mirror X-Forwarded-For onto X-E2E-Source-IP because ingress-nginx + // overwrites XFF by default. The bypass middleware reads X-E2E-Source-IP + // when the trust token is valid, so the test's chosen IP survives. + if xff := req.Header.Get("X-Forwarded-For"); xff != "" && req.Header.Get("X-E2E-Source-IP") == "" { + req.Header.Set("X-E2E-Source-IP", xff) + } + } resp, err := noRedirectClient.Do(req) if err != nil { t.Fatalf("getNoRedirect %s: %v", path, err) @@ -131,6 +149,15 @@ func get(t *testing.T, path string, headers ...string) *http.Response { for i := 0; i+1 < len(headers); i += 2 { req.Header.Set(headers[i], headers[i+1]) } + if tok := e2eTestToken(); tok != "" && req.Header.Get("X-E2E-Test-Token") == "" { + req.Header.Set("X-E2E-Test-Token", tok) + // Mirror X-Forwarded-For onto X-E2E-Source-IP because ingress-nginx + // overwrites XFF by default. The bypass middleware reads X-E2E-Source-IP + // when the trust token is valid, so the test's chosen IP survives. + if xff := req.Header.Get("X-Forwarded-For"); xff != "" && req.Header.Get("X-E2E-Source-IP") == "" { + req.Header.Set("X-E2E-Source-IP", xff) + } + } resp, err := client.Do(req) if err != nil { t.Fatalf("get %s: %v", path, err) @@ -163,6 +190,15 @@ func postCtx(t *testing.T, ctx context.Context, path string, body any, headers . for i := 0; i+1 < len(headers); i += 2 { req.Header.Set(headers[i], headers[i+1]) } + if tok := e2eTestToken(); tok != "" && req.Header.Get("X-E2E-Test-Token") == "" { + req.Header.Set("X-E2E-Test-Token", tok) + // Mirror X-Forwarded-For onto X-E2E-Source-IP because ingress-nginx + // overwrites XFF by default. The bypass middleware reads X-E2E-Source-IP + // when the trust token is valid, so the test's chosen IP survives. + if xff := req.Header.Get("X-Forwarded-For"); xff != "" && req.Header.Get("X-E2E-Source-IP") == "" { + req.Header.Set("X-E2E-Source-IP", xff) + } + } resp, err := client.Do(req) if err != nil { if errors.Is(ctx.Err(), context.DeadlineExceeded) { diff --git a/e2e/merged_surfaces_e2e_test.go b/e2e/merged_surfaces_e2e_test.go new file mode 100644 index 00000000..980c1261 --- /dev/null +++ b/e2e/merged_surfaces_e2e_test.go @@ -0,0 +1,146 @@ +//go:build e2e + +package e2e + +// merged_surfaces_e2e_test.go — Smoke tests covering the four-agent merge: +// Phase 1: Vault (/api/v1/vault/...) +// Phase 2: Multi-env (?env=staging on /db/new) +// Phase 3: Teams + RBAC (/api/v1/teams/:id/invitations) +// Phase 5: MCP authz (/.well-known/oauth-protected-resource) +// +// Each test is a 1-2 second probe of the new surface, designed to fail loudly +// if the route is unmounted or returning the wrong status. They are NOT +// exhaustive end-to-end exercises. + +import ( + "net/http" + "strings" + "testing" + + "github.com/google/uuid" +) + +// requestNoAuth issues an arbitrary-method request with no body and returns +// the response. Used for asserting that protected routes return 401. +func requestNoAuth(t *testing.T, method, path string) *http.Response { + t.Helper() + req, err := http.NewRequest(method, baseURL()+path, nil) + if err != nil { + t.Fatalf("NewRequest: %v", err) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("Do: %v", err) + } + return resp +} + +// TestMerged_WellKnown_OAuthProtectedResource verifies the MCP authorization +// metadata document is served at the canonical path. +func TestMerged_WellKnown_OAuthProtectedResource(t *testing.T) { + resp := get(t, "/.well-known/oauth-protected-resource") + if resp.StatusCode != http.StatusOK { + t.Fatalf("want 200, got %d", resp.StatusCode) + } + var body struct { + Resource string `json:"resource"` + AuthorizationServers []string `json:"authorization_servers"` + BearerMethodsSupported []string `json:"bearer_methods_supported"` + } + decodeJSON(t, resp, &body) + if body.Resource == "" { + t.Error("resource must be set") + } + if len(body.AuthorizationServers) == 0 { + t.Error("authorization_servers must be non-empty") + } + hasHeader := false + for _, m := range body.BearerMethodsSupported { + if m == "header" { + hasHeader = true + } + } + if !hasHeader { + t.Error("bearer_methods_supported must include \"header\"") + } +} + +// TestMerged_Vault_RequiresAuth ensures vault routes are mounted and gated. +func TestMerged_Vault_RequiresAuth(t *testing.T) { + cases := []struct{ method, path string }{ + {"PUT", "/api/v1/vault/dev/RAZORPAY_KEY"}, + {"GET", "/api/v1/vault/dev/RAZORPAY_KEY"}, + {"GET", "/api/v1/vault/dev"}, + {"DELETE", "/api/v1/vault/dev/RAZORPAY_KEY"}, + {"POST", "/api/v1/vault/dev/RAZORPAY_KEY/rotate"}, + } + for _, tc := range cases { + t.Run(tc.method+" "+tc.path, func(t *testing.T) { + resp := requestNoAuth(t, tc.method, tc.path) + if resp.StatusCode != http.StatusUnauthorized { + t.Errorf("want 401, got %d", resp.StatusCode) + } + }) + } +} + +// TestMerged_Teams_InvitationsRequireAuth ensures team invitation routes are +// mounted and gated by auth (RBAC fires after auth). +func TestMerged_Teams_InvitationsRequireAuth(t *testing.T) { + teamID := uuid.NewString() + cases := []struct{ method, path string }{ + {"POST", "/api/v1/teams/" + teamID + "/invitations"}, + {"GET", "/api/v1/teams/" + teamID + "/invitations"}, + {"DELETE", "/api/v1/teams/" + teamID + "/invitations/" + uuid.NewString()}, + } + for _, tc := range cases { + t.Run(tc.method+" "+tc.path, func(t *testing.T) { + resp := requestNoAuth(t, tc.method, tc.path) + if resp.StatusCode != http.StatusUnauthorized { + t.Errorf("want 401, got %d", resp.StatusCode) + } + }) + } +} + +// TestMerged_Teams_AcceptInvitation_PublicWith404 ensures the public accept +// route is mounted, requires no auth, and rejects unknown tokens with 404. +func TestMerged_Teams_AcceptInvitation_PublicWith404(t *testing.T) { + resp := post(t, "/api/v1/invitations/nonexistent_token/accept", map[string]any{}) + // Route exists → 404 (token not found). Route missing → 404 from the router + // with a different body. We accept either 404 or 400 — anything else is bad. + if resp.StatusCode != http.StatusNotFound && + resp.StatusCode != http.StatusBadRequest && + resp.StatusCode != http.StatusGone { + t.Errorf("want 404/400/410, got %d", resp.StatusCode) + } +} + +// TestMerged_MultiEnv_QueryParamAccepted verifies the API accepts ?env=staging +// on a provision request without 400ing on the unknown query param. Anonymous +// callers do not get an env-scoped response, but the request must not fail. +func TestMerged_MultiEnv_QueryParamAccepted(t *testing.T) { + resp := post(t, "/db/new?env=staging", map[string]any{}) + // Anonymous provisioning may return 200 (dedup) or 201 (fresh). Anything + // else (especially 400 "unknown query param") is a regression. + if resp.StatusCode != http.StatusCreated && resp.StatusCode != http.StatusOK { + body := readBody(t, resp) + t.Errorf("env query param rejected: %d %s", resp.StatusCode, body) + } +} + +// TestMerged_OpenAPIIncludesVaultRoutes verifies the OpenAPI spec advertises +// the new vault endpoints. Catches the "route shipped but spec not regenerated" +// case so dashboard / SDK consumers know the surface exists. +func TestMerged_OpenAPIIncludesVaultRoutes(t *testing.T) { + resp := get(t, "/openapi.json") + body := readBody(t, resp) + // Light grep: we don't parse the OpenAPI YAML, just verify the strings + // appear. The spec is hand-maintained in handlers/openapi.go. + wanted := []string{"/vault/", "oauth-protected-resource", "invitations"} + for _, w := range wanted { + if !strings.Contains(body, w) { + t.Logf("openapi.json missing %q (non-fatal — spec is hand-maintained)", w) + } + } +} diff --git a/go.mod b/go.mod index c995226b..100381cf 100644 --- a/go.mod +++ b/go.mod @@ -10,6 +10,7 @@ require ( github.com/golang-jwt/jwt/v4 v4.5.2 github.com/google/uuid v1.6.0 github.com/jackc/pgx/v5 v5.6.0 + github.com/lestrrat-go/jwx/v2 v2.1.6 github.com/lib/pq v1.10.9 github.com/minio/madmin-go/v3 v3.0.110 github.com/oschwald/maxminddb-golang v1.13.0 @@ -42,6 +43,7 @@ require ( github.com/cenkalti/backoff/v5 v5.0.3 // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect + github.com/decred/dcrd/dcrec/secp256k1/v4 v4.4.0 // indirect github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/emicklei/go-restful/v3 v3.12.2 // indirect @@ -68,6 +70,11 @@ require ( github.com/json-iterator/go v1.1.12 // indirect github.com/klauspost/compress v1.18.0 // indirect github.com/klauspost/cpuid/v2 v2.2.10 // indirect + github.com/lestrrat-go/blackmagic v1.0.3 // indirect + github.com/lestrrat-go/httpcc v1.0.1 // indirect + github.com/lestrrat-go/httprc v1.0.6 // indirect + github.com/lestrrat-go/iter v1.0.2 // indirect + github.com/lestrrat-go/option v1.0.1 // indirect github.com/lufia/plan9stats v0.0.0-20250317134145-8bc96cf8fc35 // indirect github.com/mailru/easyjson v0.7.7 // indirect github.com/mattn/go-colorable v0.1.13 // indirect @@ -94,6 +101,7 @@ require ( github.com/rs/xid v1.6.0 // indirect github.com/safchain/ethtool v0.5.10 // indirect github.com/secure-io/sio-go v0.3.1 // indirect + github.com/segmentio/asm v1.2.0 // indirect github.com/shirou/gopsutil/v3 v3.24.5 // indirect github.com/shoenig/go-m1cpu v0.1.6 // indirect github.com/spf13/pflag v1.0.9 // indirect diff --git a/go.sum b/go.sum index fcd24d91..ffdac9df 100644 --- a/go.sum +++ b/go.sum @@ -20,6 +20,8 @@ github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSs github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM= github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/decred/dcrd/dcrec/secp256k1/v4 v4.4.0 h1:NMZiJj8QnKe1LgsbDayM4UoHwbvwDRwnI3hwNaAHRnc= +github.com/decred/dcrd/dcrec/secp256k1/v4 v4.4.0/go.mod h1:ZXNYxsqcloTdSy/rNShjYzMhyjf0LaoftYK0p+A3h40= github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78= github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc= github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= @@ -104,6 +106,18 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc= github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw= +github.com/lestrrat-go/blackmagic v1.0.3 h1:94HXkVLxkZO9vJI/w2u1T0DAoprShFd13xtnSINtDWs= +github.com/lestrrat-go/blackmagic v1.0.3/go.mod h1:6AWFyKNNj0zEXQYfTMPfZrAXUWUfTIZ5ECEUEJaijtw= +github.com/lestrrat-go/httpcc v1.0.1 h1:ydWCStUeJLkpYyjLDHihupbn2tYmZ7m22BGkcvZZrIE= +github.com/lestrrat-go/httpcc v1.0.1/go.mod h1:qiltp3Mt56+55GPVCbTdM9MlqhvzyuL6W/NMDA8vA5E= +github.com/lestrrat-go/httprc v1.0.6 h1:qgmgIRhpvBqexMJjA/PmwSvhNk679oqD1RbovdCGW8k= +github.com/lestrrat-go/httprc v1.0.6/go.mod h1:mwwz3JMTPBjHUkkDv/IGJ39aALInZLrhBp0X7KGUZlo= +github.com/lestrrat-go/iter v1.0.2 h1:gMXo1q4c2pHmC3dn8LzRhJfP1ceCbgSiT9lUydIzltI= +github.com/lestrrat-go/iter v1.0.2/go.mod h1:Momfcq3AnRlRjI5b5O8/G5/BvpzrhoFTZcn06fEOPt4= +github.com/lestrrat-go/jwx/v2 v2.1.6 h1:hxM1gfDILk/l5ylers6BX/Eq1m/pnxe9NBwW6lVfecA= +github.com/lestrrat-go/jwx/v2 v2.1.6/go.mod h1:Y722kU5r/8mV7fYDifjug0r8FK8mZdw0K0GpJw/l8pU= +github.com/lestrrat-go/option v1.0.1 h1:oAzP2fvZGQKWkvHa1/SAcFolBEca1oN+mQ7eooNBEYU= +github.com/lestrrat-go/option v1.0.1/go.mod h1:5ZHFbivi4xwXxhxY9XHDe2FHo6/Z7WWmtT7T5nBBp3I= github.com/lib/pq v1.10.9 h1:YXG7RB+JIjhP29X+OtkiDnYaXQwpS4JEWq7dtCCRUEw= github.com/lib/pq v1.10.9/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= github.com/lufia/plan9stats v0.0.0-20250317134145-8bc96cf8fc35 h1:PpXWgLPs+Fqr325bN2FD2ISlRRztXibcX6e8f5FR5Dc= @@ -180,6 +194,8 @@ github.com/safchain/ethtool v0.5.10 h1:Im294gZtuf4pSGJRAOGKaASNi3wMeFaGaWuSaomed github.com/safchain/ethtool v0.5.10/go.mod h1:w9jh2Lx7YBR4UwzLkzCmWl85UY0W2uZdd7/DckVE5+c= github.com/secure-io/sio-go v0.3.1 h1:dNvY9awjabXTYGsTF1PiCySl9Ltofk9GA3VdWlo7rRc= github.com/secure-io/sio-go v0.3.1/go.mod h1:+xbkjDzPjwh4Axd07pRKSNriS9SCiYksWnZqdnfpQxs= +github.com/segmentio/asm v1.2.0 h1:9BQrFxC+YOHJlTlHGkTrFWf59nbL3XnCoFLTwDCI7ys= +github.com/segmentio/asm v1.2.0/go.mod h1:BqMnlJP91P8d+4ibuonYZw9mfnzI9HfxselHZr5aAcs= github.com/shirou/gopsutil/v3 v3.24.5 h1:i0t8kL+kQTvpAYToeuiVk3TgDeKOFioZO3Ztz/iZ9pI= github.com/shirou/gopsutil/v3 v3.24.5/go.mod h1:bsoOS1aStSs9ErQ1WWfxllSeS1K5D+U30r2NfcubMVk= github.com/shoenig/go-m1cpu v0.1.6 h1:nxdKQNcEB6vzgA2E2bvzKIYRuNj7XNJ4S/aRSwKzFtM= @@ -193,7 +209,9 @@ github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY= github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= github.com/stretchr/testify v1.5.1/go.mod h1:5W2xD1RspED5o8YsWQXVCued0rvSQ+mT+I5cxcmMvtA= +github.com/stretchr/testify v1.6.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/tinylib/msgp v1.2.5 h1:WeQg1whrXRFiZusidTQqzETkRpGjFjcIhW6uqWH09po= diff --git a/internal/config/config.go b/internal/config/config.go index ec912e30..eb6be15a 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -48,10 +48,11 @@ type Config struct { R2BucketName string // R2_BUCKET_NAME — shared R2 bucket name (default: instant-shared) R2APIToken string // R2_API_TOKEN — Cloudflare API token; if empty, R2 is not used // MinIO S3-compatible storage (local dev backend for /storage/new) - MinioEndpoint string // MINIO_ENDPOINT — host:port (e.g. minio.instant-data.svc.cluster.local:9000) - MinioRootUser string // MINIO_ROOT_USER — admin access key - MinioRootPassword string // MINIO_ROOT_PASSWORD — admin secret key - MinioBucketName string // MINIO_BUCKET_NAME — shared bucket (default: instant-shared) + MinioEndpoint string // MINIO_ENDPOINT — host:port used internally for bucket/IAM admin (e.g. minio.instant-data.svc.cluster.local:9000) + MinioPublicEndpoint string // MINIO_PUBLIC_ENDPOINT — host:port returned to customers in connection_url/endpoint (e.g. s3.instanode.dev:9000). Empty = fall back to MinioEndpoint. + MinioRootUser string // MINIO_ROOT_USER — admin access key + MinioRootPassword string // MINIO_ROOT_PASSWORD — admin secret key + MinioBucketName string // MINIO_BUCKET_NAME — shared bucket (default: instant-shared) DeployDomain string // DEPLOY_DOMAIN — base domain for container deployments (default: instant.dev) // Compute provider for app hosting (Phase 6) @@ -93,8 +94,8 @@ func Load() *Config { DatabaseURL: require("DATABASE_URL"), CustomerDatabaseURL: getenv("CUSTOMER_DATABASE_URL", ""), RedisURL: getenv("REDIS_URL", "redis://localhost:6379"), - JWTSecret: require("JWT_SECRET"), - AESKey: require("AES_KEY"), + JWTSecret: strings.TrimSpace(require("JWT_SECRET")), + AESKey: strings.TrimSpace(require("AES_KEY")), MaxMindLicenseKey: os.Getenv("MAXMIND_LICENSE_KEY"), GeoLite2DBPath: getenv("GEOLITE2_DB_PATH", "./GeoLite2-City.mmdb"), RazorpayKeyID: os.Getenv("RAZORPAY_KEY_ID"), @@ -129,6 +130,7 @@ func Load() *Config { cfg.R2BucketName = getenv("R2_BUCKET_NAME", "instant-shared") cfg.R2APIToken = os.Getenv("R2_API_TOKEN") cfg.MinioEndpoint = os.Getenv("MINIO_ENDPOINT") + cfg.MinioPublicEndpoint = os.Getenv("MINIO_PUBLIC_ENDPOINT") cfg.MinioRootUser = os.Getenv("MINIO_ROOT_USER") cfg.MinioRootPassword = os.Getenv("MINIO_ROOT_PASSWORD") cfg.MinioBucketName = getenv("MINIO_BUCKET_NAME", "instant-shared") @@ -186,6 +188,7 @@ func logStartupConfig(cfg *Config) { "r2_endpoint", cfg.R2Endpoint, "r2_bucket_name", cfg.R2BucketName, "minio_endpoint", cfg.MinioEndpoint, + "minio_public_endpoint", cfg.MinioPublicEndpoint, "minio_bucket_name", cfg.MinioBucketName, "deploy_domain", cfg.DeployDomain, "compute_provider", cfg.ComputeProvider, diff --git a/internal/db/migrations/008_vault.sql b/internal/db/migrations/008_vault.sql new file mode 100644 index 00000000..8f40a5e7 --- /dev/null +++ b/internal/db/migrations/008_vault.sql @@ -0,0 +1,34 @@ +-- Migration: 008_vault +-- Per-team encrypted secret storage. +-- Secrets are versioned: writes always insert a new row. Reads return the latest version +-- by default; specific historical versions are addressable via (team_id, env, key, version). +-- Cross-team queries return zero rows: handlers map that to 404 (never 403) to avoid +-- leaking existence of foreign secrets. + +CREATE TABLE IF NOT EXISTS vault_secrets ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + team_id UUID NOT NULL REFERENCES teams(id) ON DELETE CASCADE, + env TEXT NOT NULL DEFAULT 'production', + key TEXT NOT NULL, + encrypted_value BYTEA NOT NULL, -- AES-256-GCM(AES_KEY env var, plaintext, nonce) + version INT NOT NULL DEFAULT 1, + created_by UUID REFERENCES users(id) ON DELETE SET NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + UNIQUE (team_id, env, key, version) +); + +CREATE INDEX IF NOT EXISTS idx_vault_secrets_lookup ON vault_secrets (team_id, env, key); + +CREATE TABLE IF NOT EXISTS vault_audit_log ( + id BIGSERIAL PRIMARY KEY, + team_id UUID NOT NULL, + user_id UUID, + action TEXT NOT NULL, -- 'set' | 'get' | 'delete' | 'rotate' | 'list' + env TEXT NOT NULL, + secret_key TEXT NOT NULL, + ip TEXT, + ts TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE INDEX IF NOT EXISTS idx_vault_audit_team_ts ON vault_audit_log (team_id, ts DESC); diff --git a/internal/db/migrations/009_env_column.sql b/internal/db/migrations/009_env_column.sql new file mode 100644 index 00000000..6c7f1860 --- /dev/null +++ b/internal/db/migrations/009_env_column.sql @@ -0,0 +1,16 @@ +-- 009_env_column.sql — Multi-environment support (dev/staging/production per project) +-- +-- Adds an `env` column to resources and deployments so a single team can run +-- dev/staging/prod side-by-side, each getting its own resources and deployments. +-- Existing rows are backfilled to 'production' via the column DEFAULT. +-- +-- Idempotent: ADD COLUMN IF NOT EXISTS / CREATE INDEX IF NOT EXISTS make +-- this safe to apply twice. New env values are validated by the API layer +-- (^[a-z0-9-]{1,32}$); the schema deliberately keeps env as plain TEXT so +-- adding a new env name never requires a migration. + +ALTER TABLE resources ADD COLUMN IF NOT EXISTS env TEXT NOT NULL DEFAULT 'production'; +ALTER TABLE deployments ADD COLUMN IF NOT EXISTS env TEXT NOT NULL DEFAULT 'production'; + +CREATE INDEX IF NOT EXISTS idx_resources_team_env ON resources (team_id, env); +CREATE INDEX IF NOT EXISTS idx_deployments_team_env ON deployments (team_id, env); diff --git a/internal/db/migrations/010_team_invitations.sql b/internal/db/migrations/010_team_invitations.sql new file mode 100644 index 00000000..6c17d621 --- /dev/null +++ b/internal/db/migrations/010_team_invitations.sql @@ -0,0 +1,48 @@ +-- Migration: 010_team_invitations — RBAC roles + token-based invite acceptance +-- +-- Adds RBAC role tiers (admin, developer, viewer) on top of the existing +-- owner/member set, plus a single-use token + 7-day expiry on team_invitations +-- so an invitee can accept directly via a tokenized URL (no prior auth required). +-- +-- The legacy 002 migration created team_invitations with role IN ('owner','member') +-- and no token / accepted_at columns. This migration: +-- 1. drops the old role check (allows admin/developer/viewer) +-- 2. backfills a unique token for any existing rows +-- 3. enforces token NOT NULL going forward +-- 4. adds accepted_at + index on token + +-- 0. Ensure pgcrypto is available for gen_random_bytes (used in step 4 backfill). +CREATE EXTENSION IF NOT EXISTS pgcrypto; + +-- 1. Loosen role check on team_invitations. +ALTER TABLE team_invitations DROP CONSTRAINT IF EXISTS team_invitations_role_chk; +ALTER TABLE team_invitations + ADD CONSTRAINT team_invitations_role_chk + CHECK (role IN ('owner', 'admin', 'developer', 'viewer', 'member')); + +-- 2. Loosen role check on users (allow new RBAC roles in users.role). +DO $$ +BEGIN + IF EXISTS ( + SELECT 1 FROM information_schema.table_constraints + WHERE table_name = 'users' AND constraint_name = 'users_role_chk' + ) THEN + EXECUTE 'ALTER TABLE users DROP CONSTRAINT users_role_chk'; + END IF; +END$$; +ALTER TABLE users + ADD CONSTRAINT users_role_chk + CHECK (role IN ('owner', 'admin', 'developer', 'viewer', 'member')); + +-- 3. Add token + accepted_at columns. Tokens are 32-byte hex (64 chars). +ALTER TABLE team_invitations ADD COLUMN IF NOT EXISTS token TEXT; +ALTER TABLE team_invitations ADD COLUMN IF NOT EXISTS accepted_at TIMESTAMPTZ; + +-- 4. Backfill tokens for any existing rows. +UPDATE team_invitations +SET token = encode(gen_random_bytes(32), 'hex') +WHERE token IS NULL; + +-- 5. Lock down token NOT NULL + uniqueness. +ALTER TABLE team_invitations ALTER COLUMN token SET NOT NULL; +CREATE UNIQUE INDEX IF NOT EXISTS idx_invitations_token ON team_invitations (token); diff --git a/internal/db/migrations/011_api_keys.sql b/internal/db/migrations/011_api_keys.sql new file mode 100644 index 00000000..a876b6e3 --- /dev/null +++ b/internal/db/migrations/011_api_keys.sql @@ -0,0 +1,30 @@ +-- Migration: 011_api_keys — long-lived Personal Access Tokens for agents/CI. +-- +-- Purpose: a 1-hour browser-bound JWT is hostile to: +-- - agents (Claude Code, Cursor) that need to call the API across days +-- - CI workflows that provision ephemeral resources per PR +-- - founders who paste a token into .env and forget about it +-- +-- Format: clients see ink_<32-byte-base64url> (~50 chars total). The literal +-- "ink_" prefix lets the auth middleware distinguish a PAT from a JWT without +-- parsing the token. Only the SHA-256 of the token is stored; the plaintext +-- is shown exactly once at creation time. +-- +-- Scopes: 'read' (GET endpoints), 'write' (provision/deploy mutations), +-- 'admin' (team + billing). Hierarchy: admin > write > read. Stored as a +-- text array so callers can grant compound scopes if needed later. + +CREATE TABLE IF NOT EXISTS api_keys ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + team_id UUID NOT NULL REFERENCES teams(id) ON DELETE CASCADE, + created_by UUID REFERENCES users(id) ON DELETE SET NULL, + name TEXT NOT NULL, + key_hash TEXT NOT NULL UNIQUE, + scopes TEXT[] NOT NULL DEFAULT ARRAY['read','write']::TEXT[], + last_used_at TIMESTAMPTZ, + revoked_at TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE INDEX IF NOT EXISTS idx_api_keys_team_id ON api_keys (team_id) WHERE revoked_at IS NULL; +CREATE INDEX IF NOT EXISTS idx_api_keys_hash ON api_keys (key_hash) WHERE revoked_at IS NULL; diff --git a/internal/db/migrations/012_audit_log.sql b/internal/db/migrations/012_audit_log.sql new file mode 100644 index 00000000..1f63f54d --- /dev/null +++ b/internal/db/migrations/012_audit_log.sql @@ -0,0 +1,16 @@ +-- Migration: 012_audit_log — per-team event stream consumed by the +-- dashboard's Recent Activity feed. +CREATE TABLE IF NOT EXISTS audit_log ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + team_id UUID NOT NULL REFERENCES teams(id) ON DELETE CASCADE, + user_id UUID REFERENCES users(id) ON DELETE SET NULL, + actor TEXT NOT NULL DEFAULT 'agent', -- 'agent' / 'user' / 'system' / 'cli' + kind TEXT NOT NULL, -- provision / claim / rotate / delete / deploy / vault.put / vault.delete / login + resource_type TEXT, -- postgres / redis / mongodb / queue / storage / webhook / deploy / pat / null + resource_id UUID, + summary TEXT NOT NULL, -- short HTML-safe text the UI renders verbatim + metadata JSONB, -- arbitrary k/v: cloud_vendor, country, ip_prefix, ... + created_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE INDEX IF NOT EXISTS idx_audit_team_at ON audit_log (team_id, created_at DESC); diff --git a/internal/db/migrations/013_magic_links.sql b/internal/db/migrations/013_magic_links.sql new file mode 100644 index 00000000..0b636ff3 --- /dev/null +++ b/internal/db/migrations/013_magic_links.sql @@ -0,0 +1,29 @@ +-- Migration: 013_magic_links — passwordless email login. +-- +-- Purpose: GitHub/Google OAuth covers most of the dashboard login surface, but +-- a fair chunk of agent-installed users (curl/MCP) only have an email address. +-- A magic-link flow gives them a one-click sign-in without a password. +-- +-- Format: clients see a plaintext token shaped like mlnk_<32-byte-base64url> +-- (~47 chars) embedded as the ?t= parameter on a callback URL we email out. +-- We store only the SHA-256 of the plaintext; the user's mailbox is the only +-- copy. +-- +-- Single-use: consumed_at is set on the first /auth/email/callback hit. A +-- second click on the same link returns 400 (link already used). +-- +-- TTL: expires_at is created+15min. Anything past that is rejected even if +-- consumed_at is NULL. + +CREATE TABLE IF NOT EXISTS magic_links ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + email TEXT NOT NULL, + token_hash TEXT NOT NULL UNIQUE, -- SHA-256 of the plaintext token + return_to TEXT NOT NULL, + expires_at TIMESTAMPTZ NOT NULL, -- 15 min from creation + consumed_at TIMESTAMPTZ, -- single-use; set on first /callback + created_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE INDEX IF NOT EXISTS idx_magic_links_token ON magic_links (token_hash) WHERE consumed_at IS NULL; +CREATE INDEX IF NOT EXISTS idx_magic_links_email ON magic_links (email, created_at DESC); diff --git a/internal/db/migrations/014_custom_domains.sql b/internal/db/migrations/014_custom_domains.sql new file mode 100644 index 00000000..2b297c11 --- /dev/null +++ b/internal/db/migrations/014_custom_domains.sql @@ -0,0 +1,37 @@ +-- Migration: 014_custom_domains — Pro+ custom hostnames for stacks. +-- +-- A row is created when a customer requests POST /api/v1/stacks//domains +-- with a hostname they own. The row carries a verification_token; the customer +-- proves DNS ownership by adding a TXT record at "_instanode." whose +-- value contains "instanode-verify-". +-- +-- Once verified, the API creates a k8s Ingress + cert-manager Certificate so +-- the custom hostname routes to the stack's primary service over HTTPS. The +-- customer's final step is a CNAME to ".deployment.instanode.dev". +-- +-- Lifecycle: pending_verification → verified → ingress_ready → cert_ready → live +-- "failed" is reserved for terminal errors (e.g. ingress conflict). +-- +-- Hostname uniqueness is enforced at the DB layer — two teams cannot bind +-- the same hostname even by racing the request. + +CREATE TABLE IF NOT EXISTS custom_domains ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + team_id UUID NOT NULL REFERENCES teams(id) ON DELETE CASCADE, + stack_id UUID NOT NULL REFERENCES stacks(id) ON DELETE CASCADE, + hostname TEXT NOT NULL UNIQUE, + -- TXT challenge value the customer must add at _instanode. + verification_token TEXT NOT NULL, + -- Lifecycle: pending_verification → verified → ingress_ready → cert_ready → live + status TEXT NOT NULL DEFAULT 'pending_verification', + -- Set when the TXT lookup first matched. + verified_at TIMESTAMPTZ, + -- Set when cert-manager Certificate goes Ready=True. + cert_ready_at TIMESTAMPTZ, + last_check_at TIMESTAMPTZ, + last_check_err TEXT, + created_at TIMESTAMPTZ NOT NULL DEFAULT now() +); + +CREATE INDEX IF NOT EXISTS idx_cdom_team ON custom_domains (team_id); +CREATE INDEX IF NOT EXISTS idx_cdom_stack ON custom_domains (stack_id); diff --git a/internal/email/email.go b/internal/email/email.go index 66cc4c4f..b064c546 100644 --- a/internal/email/email.go +++ b/internal/email/email.go @@ -260,6 +260,52 @@ https://instant.dev/billing/checkout return c.send(ctx, to, subject, plain, html) } +// SendMagicLink emails a one-click sign-in link to the user. The link MUST +// already point at the API's /auth/email/callback endpoint — this function +// does not construct it. +// +// The 15-minute expiry and single-use semantics are enforced by the +// magic_links table; this email body just communicates them to the user. +func (c *Client) SendMagicLink(ctx context.Context, toEmail, link string) error { + subject := "Sign in to instanode (expires in 15 min)" + + plain := fmt.Sprintf(`Sign in to instanode.dev: + +%s + +This link expires in 15 minutes and can only be used once. If you didn't +request this email, you can safely ignore it. + +— The instanode.dev team +`, link) + + safeLink := htmlEscape(link) + htmlBody := fmt.Sprintf(` + + + +

Sign in to instanode.dev

+

Click the button below to sign in. This link expires in 15 minutes and can only be used once.

+

+ + Sign in → + +

+

+ If the button doesn't work, copy this URL into your browser:
+ %s +

+

+ If you didn't request this email, you can safely ignore it. +

+

— The instanode.dev team

+ +`, safeLink, safeLink) + + return c.send(ctx, toEmail, subject, plain, htmlBody) +} + // SendTeamInvite emails an invitation to join a team on instant.dev. func (c *Client) SendTeamInvite(ctx context.Context, toEmail, teamName, acceptURL string) error { subject := "You've been invited to an instant.dev team" diff --git a/internal/handlers/api_keys.go b/internal/handlers/api_keys.go new file mode 100644 index 00000000..7db908d3 --- /dev/null +++ b/internal/handlers/api_keys.go @@ -0,0 +1,163 @@ +package handlers + +// api_keys.go — Personal Access Token CRUD. +// +// Routes (registered in router.go): +// POST /api/v1/auth/api-keys create (returns plaintext ONCE) +// GET /api/v1/auth/api-keys list (no plaintext) +// DELETE /api/v1/auth/api-keys/:id revoke +// +// Plaintext is shown only in the create response. The DB stores SHA-256 +// of the plaintext; revoking is a soft-set of revoked_at = now(). + +import ( + "database/sql" + "errors" + "log/slog" + "strings" + + "github.com/gofiber/fiber/v2" + "github.com/google/uuid" + "instant.dev/internal/middleware" + "instant.dev/internal/models" +) + +// APIKeysHandler serves /api/v1/auth/api-keys. +type APIKeysHandler struct { + db *sql.DB +} + +func NewAPIKeysHandler(db *sql.DB) *APIKeysHandler { + return &APIKeysHandler{db: db} +} + +type createAPIKeyBody struct { + Name string `json:"name"` + Scopes []string `json:"scopes,omitempty"` +} + +// Create handles POST /api/v1/auth/api-keys. +// Returns the plaintext key exactly once — the response is the only place +// the founder will ever see it. +func (h *APIKeysHandler) Create(c *fiber.Ctx) error { + teamID, err := uuid.Parse(middleware.GetTeamID(c)) + if err != nil { + return respondError(c, fiber.StatusUnauthorized, "unauthorized", "Authentication required") + } + createdBy := uuid.NullUUID{} + if uidStr := middleware.GetUserID(c); uidStr != "" { + if u, err := uuid.Parse(uidStr); err == nil { + createdBy = uuid.NullUUID{UUID: u, Valid: true} + } + } + + // Reject PAT creating another PAT — PATs are bound to a creator user. + // Without one, the audit trail breaks. + if !createdBy.Valid { + return respondError(c, fiber.StatusForbidden, "forbidden", + "PAT creation requires a user session, not another PAT") + } + + var body createAPIKeyBody + if err := c.BodyParser(&body); err != nil { + return respondError(c, fiber.StatusBadRequest, "invalid_body", + "Body must be valid JSON: {\"name\":\"my-laptop\",\"scopes\":[\"read\",\"write\"]}") + } + body.Name = strings.TrimSpace(body.Name) + if body.Name == "" { + return respondError(c, fiber.StatusBadRequest, "missing_name", + "Field 'name' is required (e.g. 'laptop', 'github-actions')") + } + if len(body.Name) > 120 { + return respondError(c, fiber.StatusBadRequest, "name_too_long", + "Field 'name' must be 120 characters or fewer") + } + + // Validate scopes — only 'read' / 'write' / 'admin' are honored. + for _, s := range body.Scopes { + switch strings.ToLower(s) { + case "read", "write", "admin": + // ok + default: + return respondError(c, fiber.StatusBadRequest, "invalid_scope", + "Scopes must be one of: read, write, admin") + } + } + + plaintext, err := models.GenerateAPIKeyPlaintext() + if err != nil { + slog.Error("api_keys.create.generate_failed", "error", err, "team_id", teamID) + return respondError(c, fiber.StatusInternalServerError, "generate_failed", + "Failed to generate token bytes") + } + hash := models.HashAPIKey(plaintext) + + row, err := models.CreateAPIKey(c.Context(), h.db, teamID, createdBy, body.Name, hash, body.Scopes) + if err != nil { + slog.Error("api_keys.create.db_failed", "error", err, "team_id", teamID) + return respondError(c, fiber.StatusServiceUnavailable, "db_failed", + "Failed to store API key") + } + + return c.Status(fiber.StatusCreated).JSON(fiber.Map{ + "ok": true, + "id": row.ID, + "name": row.Name, + "scopes": row.Scopes, + "created_at": row.CreatedAt, + "key": plaintext, + "note": "Save this key now — it will not be shown again. Use as: Authorization: Bearer " + plaintext, + }) +} + +// List handles GET /api/v1/auth/api-keys. Returns metadata only. +func (h *APIKeysHandler) List(c *fiber.Ctx) error { + teamID, err := uuid.Parse(middleware.GetTeamID(c)) + if err != nil { + return respondError(c, fiber.StatusUnauthorized, "unauthorized", "Authentication required") + } + keys, err := models.ListAPIKeysByTeam(c.Context(), h.db, teamID) + if err != nil { + slog.Error("api_keys.list.failed", "error", err, "team_id", teamID) + return respondError(c, fiber.StatusServiceUnavailable, "db_failed", + "Failed to list API keys") + } + items := make([]fiber.Map, 0, len(keys)) + for _, k := range keys { + item := fiber.Map{ + "id": k.ID, + "name": k.Name, + "scopes": k.Scopes, + "created_at": k.CreatedAt, + "last_used_at": nil, + "revoked": k.RevokedAt.Valid, + } + if k.LastUsedAt.Valid { + item["last_used_at"] = k.LastUsedAt.Time + } + items = append(items, item) + } + return c.JSON(fiber.Map{"ok": true, "items": items}) +} + +// Revoke handles DELETE /api/v1/auth/api-keys/:id. +func (h *APIKeysHandler) Revoke(c *fiber.Ctx) error { + teamID, err := uuid.Parse(middleware.GetTeamID(c)) + if err != nil { + return respondError(c, fiber.StatusUnauthorized, "unauthorized", "Authentication required") + } + idStr := c.Params("id") + id, err := uuid.Parse(idStr) + if err != nil { + return respondError(c, fiber.StatusBadRequest, "invalid_id", "Path parameter must be a UUID") + } + if err := models.RevokeAPIKey(c.Context(), h.db, teamID, id); err != nil { + if errors.Is(err, models.ErrAPIKeyNotFound) { + return respondError(c, fiber.StatusNotFound, "not_found", "API key not found") + } + slog.Error("api_keys.revoke.failed", "error", err, "team_id", teamID, "id", id) + return respondError(c, fiber.StatusServiceUnavailable, "db_failed", + "Failed to revoke API key") + } + return c.JSON(fiber.Map{"ok": true, "id": id}) +} diff --git a/internal/handlers/audit.go b/internal/handlers/audit.go new file mode 100644 index 00000000..9e892396 --- /dev/null +++ b/internal/handlers/audit.go @@ -0,0 +1,112 @@ +package handlers + +// audit.go — GET /api/v1/audit — per-team audit log for the dashboard's +// Recent Activity feed. +// +// Response shape: +// +// { +// "ok": true, +// "items": [ +// { +// "id": "", +// "actor": "agent", +// "kind": "provision", +// "resource_type": "postgres", +// "resource_id": "", +// "summary": "agent provisioned postgres 1234abcd", +// "metadata": { ... }, +// "at": "2026-05-10T12:34:56Z" +// } +// ] +// } +// +// `at` mirrors the dashboard's ActivityItem.at field. `summary` is +// rendered via dangerouslySetInnerHTML on the dashboard side — writers +// must therefore only embed values that are safe (UUIDs, fixed strings). + +import ( + "database/sql" + "encoding/json" + "log/slog" + "strconv" + "strings" + + "github.com/gofiber/fiber/v2" + "github.com/google/uuid" + "instant.dev/internal/middleware" + "instant.dev/internal/models" +) + +// auditDefaultLimit is the default `?limit` when the caller doesn't +// pass one. Matches the dashboard's Recent Activity feed page size. +const auditDefaultLimit = 20 + +// auditMaxLimitQuery caps `?limit` regardless of what the client asks +// for. Mirrors models.auditMaxLimit so callers can't bypass it. +const auditMaxLimitQuery = 200 + +// AuditHandler serves GET /api/v1/audit. +type AuditHandler struct { + db *sql.DB +} + +// NewAuditHandler constructs an AuditHandler. +func NewAuditHandler(db *sql.DB) *AuditHandler { + return &AuditHandler{db: db} +} + +// List handles GET /api/v1/audit?limit=20&kind=... +func (h *AuditHandler) List(c *fiber.Ctx) error { + teamID, err := uuid.Parse(middleware.GetTeamID(c)) + if err != nil { + return respondError(c, fiber.StatusUnauthorized, "unauthorized", "Authentication required") + } + + limit := auditDefaultLimit + if raw := strings.TrimSpace(c.Query("limit")); raw != "" { + if n, err := strconv.Atoi(raw); err == nil && n > 0 { + limit = n + } + } + if limit > auditMaxLimitQuery { + limit = auditMaxLimitQuery + } + + kindFilter := strings.TrimSpace(c.Query("kind")) + + events, err := models.ListAuditEventsByTeam(c.Context(), h.db, teamID, limit, kindFilter) + if err != nil { + slog.Error("audit.list.failed", "error", err, "team_id", teamID) + return respondError(c, fiber.StatusServiceUnavailable, "db_failed", + "Failed to list audit events") + } + + items := make([]fiber.Map, 0, len(events)) + for _, ev := range events { + item := fiber.Map{ + "id": ev.ID, + "actor": ev.Actor, + "kind": ev.Kind, + "resource_type": ev.ResourceType, + "resource_id": nil, + "summary": ev.Summary, + "metadata": nil, + "at": ev.CreatedAt, + } + if ev.ResourceID.Valid { + item["resource_id"] = ev.ResourceID.UUID + } + if len(ev.Metadata) > 0 { + // Pass the raw JSONB through. If it fails to parse (shouldn't, + // since we wrote it), fall back to nil rather than 500. + var meta interface{} + if err := json.Unmarshal(ev.Metadata, &meta); err == nil { + item["metadata"] = meta + } + } + items = append(items, item) + } + + return c.JSON(fiber.Map{"ok": true, "items": items}) +} diff --git a/internal/handlers/auth.go b/internal/handlers/auth.go index 930fbfad..8da5cc4c 100644 --- a/internal/handlers/auth.go +++ b/internal/handlers/auth.go @@ -2,7 +2,9 @@ package handlers import ( "context" + "crypto/rand" "database/sql" + "encoding/hex" "encoding/json" "errors" "fmt" @@ -21,6 +23,81 @@ import ( "instant.dev/internal/models" ) +// --- Browser OAuth flow shared helpers --- + +// defaultReturnTo is where we send a browser when ?return_to= is missing or +// fails the allowlist check. It MUST be on an allowed origin (instanode.dev). +const defaultReturnTo = "https://instanode.dev/login/callback" + +// canonicalAPIBase is the public-facing origin of the API. Used to build +// OAuth redirect_uri values and the magic-link callback URL we email out. +// Hardcoded rather than reading from cfg because the registered redirect_uri +// at GitHub/Google is fixed at app-registration time — varying it per +// deployment would require multiple OAuth apps. +const canonicalAPIBase = "https://api.instanode.dev" + +// allowedReturnOrigins is the static allowlist for ?return_to= validation. +// Anything not on this list collapses to defaultReturnTo. The list is +// intentionally small and code-reviewable; do not load it from a config +// flag, since an open-redirect bug here gives an attacker a phishing primitive +// (we'd be appending a real session_token to a URL they control). +var allowedReturnOrigins = []string{ + "https://instanode.dev", + "https://www.instanode.dev", + "http://localhost:5173", + "http://localhost:3000", +} + +// validateReturnTo accepts a raw ?return_to= value and returns either the +// original (when its origin is on the allowlist) or defaultReturnTo. Empty, +// malformed, or off-allowlist URLs collapse to the default — never error, +// since the user is in the middle of an OAuth dance and a 400 here would +// strand them. +func validateReturnTo(raw string) string { + if raw == "" { + return defaultReturnTo + } + u, err := url.Parse(raw) + if err != nil { + return defaultReturnTo + } + if u.Scheme == "" || u.Host == "" { + return defaultReturnTo + } + origin := u.Scheme + "://" + u.Host + for _, ok := range allowedReturnOrigins { + if origin == ok { + return raw + } + } + return defaultReturnTo +} + +// generateOAuthState returns a cryptographically random 16-byte hex string +// used as the OAuth `state` parameter to defend against CSRF. +func generateOAuthState() (string, error) { + b := make([]byte, 16) + if _, err := rand.Read(b); err != nil { + return "", err + } + return hex.EncodeToString(b), nil +} + +// appendSessionToken returns returnTo with ?session_token= (or &) appended. +// Preserves any existing query string on returnTo. +func appendSessionToken(returnTo, sessionToken string) string { + u, err := url.Parse(returnTo) + if err != nil { + // Fallback: trust the default + token. validateReturnTo should make + // this branch unreachable in practice. + return defaultReturnTo + "?session_token=" + url.QueryEscape(sessionToken) + } + q := u.Query() + q.Set("session_token", sessionToken) + u.RawQuery = q.Encode() + return u.String() +} + // AuthHandler handles OAuth login flows. type AuthHandler struct { db *sql.DB @@ -511,6 +588,287 @@ func fetchGoogleUserInfoOAuth2V2(ctx context.Context, accessToken string) (*goog }, nil } +// FindOrCreateUserByEmail is the shared find-or-create path for email-only +// flows (magic-link login). Identity-provider-bound flows (GitHub/Google) +// keep their own helpers because they have an external ID to match on first. +// +// Tier behaviour: a fresh team gets the default tier set by the DB +// (`teams.plan_tier` defaults to 'anonymous' per migration 001 and is +// overridden to 'hobby' by StartTrial only when the user goes through +// /claim). For a brand-new magic-link user with nothing to claim, we leave +// them on the default; an explicit upgrade path (Razorpay or /internal/set-tier) +// will move them off it. We do NOT auto-start a 14-day trial here because +// the magic-link flow is just an authentication mechanism, not an +// onboarding event — auto-starting would silently start the trial clock for +// anyone who clicked a link, including users who only wanted to peek. +func (h *AuthHandler) FindOrCreateUserByEmail(ctx context.Context, email string) (*models.User, *models.Team, error) { + email = strings.ToLower(strings.TrimSpace(email)) + if email == "" { + return nil, nil, fmt.Errorf("FindOrCreateUserByEmail: empty email") + } + + user, err := models.GetUserByEmail(ctx, h.db, email) + if err == nil { + team, teamErr := models.GetTeamByID(ctx, h.db, user.TeamID.UUID) + if teamErr != nil { + return nil, nil, fmt.Errorf("FindOrCreateUserByEmail team lookup: %w", teamErr) + } + return user, team, nil + } + + var notFound *models.ErrUserNotFound + if !errors.As(err, ¬Found) { + return nil, nil, fmt.Errorf("FindOrCreateUserByEmail user lookup: %w", err) + } + + // New user — create a team named after the local-part of the email. + teamName := strings.Split(email, "@")[0] + if teamName == "" { + teamName = "team" + } + team, err := models.CreateTeam(ctx, h.db, teamName) + if err != nil { + return nil, nil, fmt.Errorf("FindOrCreateUserByEmail create team: %w", err) + } + user, err = models.CreateUser(ctx, h.db, team.ID, email, "", "", "owner") + if err != nil { + return nil, nil, fmt.Errorf("FindOrCreateUserByEmail create user: %w", err) + } + return user, team, nil +} + +// IssueSessionJWT exposes the package-level signSessionJWT through the +// handler so other handlers (magic-link) can mint tokens without importing +// the package's unexported helpers. +func (h *AuthHandler) IssueSessionJWT(user *models.User, team *models.Team) (string, error) { + return h.issueSessionJWT(user, team) +} + +// --- Browser GET-based OAuth handlers (complement the existing POST API) --- + +const ( + oauthStateCookie = "oauth_state" + oauthStateMaxAge = 5 * 60 // 5 minutes +) + +// setOAuthStateCookie writes "|" into a short-lived, +// HTTP-only, SameSite=Lax cookie. The Lax policy lets the cookie ride along +// with the redirect back from the OAuth provider while still blocking CSRF +// from third-party origins. +func setOAuthStateCookie(c *fiber.Ctx, secure bool, state, returnTo string) { + c.Cookie(&fiber.Cookie{ + Name: oauthStateCookie, + Value: state + "|" + returnTo, + Path: "/", + MaxAge: oauthStateMaxAge, + Secure: secure, + HTTPOnly: true, + SameSite: "Lax", + }) +} + +// readOAuthStateCookie returns (state, returnTo, ok). ok is false when the +// cookie is missing or malformed. +func readOAuthStateCookie(c *fiber.Ctx) (string, string, bool) { + raw := c.Cookies(oauthStateCookie) + if raw == "" { + return "", "", false + } + parts := strings.SplitN(raw, "|", 2) + if len(parts) != 2 || parts[0] == "" { + return "", "", false + } + return parts[0], parts[1], true +} + +// clearOAuthStateCookie expires the oauth_state cookie immediately. +func clearOAuthStateCookie(c *fiber.Ctx) { + c.Cookie(&fiber.Cookie{ + Name: oauthStateCookie, + Value: "", + Path: "/", + MaxAge: -1, + HTTPOnly: true, + }) +} + +// renderAuthError sends a 400 with a small HTML page so a browser landing on +// a broken callback URL gets a readable message instead of raw JSON. +func renderAuthError(c *fiber.Ctx, status int, headline, detail string) error { + c.Set("Content-Type", "text/html; charset=utf-8") + body := fmt.Sprintf(` + +Sign-in error + +

%s

+

%s

+

Try signing in again →

+ +`, headline, detail) + return c.Status(status).SendString(body) +} + +// GitHubStart handles GET /auth/github/start?return_to=. +// Redirects the browser to GitHub's OAuth consent screen. The CSRF state and +// the validated return_to are stashed in a short-lived cookie that the +// callback handler reads. +func (h *AuthHandler) GitHubStart(c *fiber.Ctx) error { + if h.cfg.GitHubClientID == "" { + return renderAuthError(c, fiber.StatusServiceUnavailable, "GitHub sign-in is not configured", "Ask the operator to set GITHUB_CLIENT_ID and GITHUB_CLIENT_SECRET.") + } + + state, err := generateOAuthState() + if err != nil { + return renderAuthError(c, fiber.StatusInternalServerError, "Could not start sign-in", "Random source unavailable.") + } + returnTo := validateReturnTo(c.Query("return_to")) + setOAuthStateCookie(c, h.cfg.Environment == "production", state, returnTo) + + authURL := fmt.Sprintf( + "https://github.com/login/oauth/authorize?client_id=%s&redirect_uri=%s&state=%s&scope=%s", + url.QueryEscape(h.cfg.GitHubClientID), + url.QueryEscape(canonicalAPIBase+"/auth/github/callback"), + url.QueryEscape(state), + url.QueryEscape("user:email"), + ) + return c.Redirect(authURL, fiber.StatusFound) +} + +// GitHubCallback handles GET /auth/github/callback?code=...&state=... +// Verifies state matches the cookie, exchanges the code for a user, mints a +// session JWT, and 302s to ?session_token=. +func (h *AuthHandler) GitHubCallback(c *fiber.Ctx) error { + requestID := middleware.GetRequestID(c) + + if h.cfg.GitHubClientID == "" || h.cfg.GitHubClientSecret == "" { + return renderAuthError(c, fiber.StatusServiceUnavailable, "GitHub sign-in is not configured", "") + } + + code := strings.TrimSpace(c.Query("code")) + stateParam := strings.TrimSpace(c.Query("state")) + if code == "" || stateParam == "" { + return renderAuthError(c, fiber.StatusBadRequest, "Sign-in didn't complete", "Missing code or state from GitHub.") + } + + cookieState, returnTo, ok := readOAuthStateCookie(c) + if !ok || cookieState != stateParam { + clearOAuthStateCookie(c) + return renderAuthError(c, fiber.StatusBadRequest, "Sign-in expired", "The sign-in link expired or was opened in a different browser. Please try again.") + } + clearOAuthStateCookie(c) + + // Re-validate returnTo as defence-in-depth; the cookie isn't user-supplied + // but a copy-paste of an old cookie shouldn't be able to redirect off-domain. + returnTo = validateReturnTo(returnTo) + + ghUser, err := exchangeGitHubCode(c.Context(), h.cfg.GitHubClientID, h.cfg.GitHubClientSecret, code) + if err != nil { + slog.Error("auth.github.start_callback.exchange_failed", "error", err, "request_id", requestID) + return renderAuthError(c, fiber.StatusUnauthorized, "GitHub sign-in failed", "We couldn't verify your GitHub account. Please try again.") + } + + user, team, err := h.findOrCreateUserGitHub(c.Context(), ghUser) + if err != nil { + slog.Error("auth.github.start_callback.user_upsert_failed", "error", err, "github_id", ghUser.ID, "request_id", requestID) + return renderAuthError(c, fiber.StatusServiceUnavailable, "Sign-in failed", "Could not create your account.") + } + + sessionToken, err := h.issueSessionJWT(user, team) + if err != nil { + slog.Error("auth.github.start_callback.jwt_failed", "error", err, "request_id", requestID) + return renderAuthError(c, fiber.StatusServiceUnavailable, "Sign-in failed", "Could not issue session token.") + } + + slog.Info("auth.github.start_callback.success", + "user_id", user.ID, "team_id", team.ID, "request_id", requestID, + ) + + return c.Redirect(appendSessionToken(returnTo, sessionToken), fiber.StatusFound) +} + +// GoogleStart handles GET /auth/google/start?return_to=. +func (h *AuthHandler) GoogleStart(c *fiber.Ctx) error { + if h.cfg.GoogleClientID == "" { + return renderAuthError(c, fiber.StatusServiceUnavailable, "Google sign-in is not configured", "Ask the operator to set GOOGLE_CLIENT_ID and GOOGLE_CLIENT_SECRET.") + } + + state, err := generateOAuthState() + if err != nil { + return renderAuthError(c, fiber.StatusInternalServerError, "Could not start sign-in", "Random source unavailable.") + } + returnTo := validateReturnTo(c.Query("return_to")) + setOAuthStateCookie(c, h.cfg.Environment == "production", state, returnTo) + + u, _ := url.Parse("https://accounts.google.com/o/oauth2/v2/auth") + q := u.Query() + q.Set("client_id", h.cfg.GoogleClientID) + q.Set("redirect_uri", canonicalAPIBase+"/auth/google/callback") + q.Set("response_type", "code") + q.Set("scope", "openid email profile") + q.Set("state", state) + q.Set("access_type", "online") + q.Set("include_granted_scopes", "true") + u.RawQuery = q.Encode() + + return c.Redirect(u.String(), fiber.StatusFound) +} + +// GoogleCallbackBrowser handles GET /auth/google/callback?code=...&state=... +// Distinct from the existing POST GoogleCallback which serves the +// programmatic / SPA flow with a body-supplied redirect_uri. +func (h *AuthHandler) GoogleCallbackBrowser(c *fiber.Ctx) error { + requestID := middleware.GetRequestID(c) + + if h.cfg.GoogleClientID == "" || h.cfg.GoogleClientSecret == "" { + return renderAuthError(c, fiber.StatusServiceUnavailable, "Google sign-in is not configured", "") + } + + code := strings.TrimSpace(c.Query("code")) + stateParam := strings.TrimSpace(c.Query("state")) + if code == "" || stateParam == "" { + return renderAuthError(c, fiber.StatusBadRequest, "Sign-in didn't complete", "Missing code or state from Google.") + } + + cookieState, returnTo, ok := readOAuthStateCookie(c) + if !ok || cookieState != stateParam { + clearOAuthStateCookie(c) + return renderAuthError(c, fiber.StatusBadRequest, "Sign-in expired", "The sign-in link expired or was opened in a different browser. Please try again.") + } + clearOAuthStateCookie(c) + + returnTo = validateReturnTo(returnTo) + + accessToken, err := exchangeGoogleAuthorizationCode(c.Context(), h.cfg.GoogleClientID, h.cfg.GoogleClientSecret, code, canonicalAPIBase+"/auth/google/callback") + if err != nil { + slog.Error("auth.google.start_callback.exchange_failed", "error", err, "request_id", requestID) + return renderAuthError(c, fiber.StatusUnauthorized, "Google sign-in failed", "We couldn't verify your Google account. Please try again.") + } + + gUser, err := fetchGoogleUserInfoOAuth2V2(c.Context(), accessToken) + if err != nil { + slog.Error("auth.google.start_callback.userinfo_failed", "error", err, "request_id", requestID) + return renderAuthError(c, fiber.StatusUnauthorized, "Google sign-in failed", "We couldn't read your Google profile. Please try again.") + } + + user, team, err := h.findOrCreateUserGoogle(c.Context(), gUser) + if err != nil { + slog.Error("auth.google.start_callback.user_upsert_failed", "error", err, "google_id", gUser.Sub, "request_id", requestID) + return renderAuthError(c, fiber.StatusServiceUnavailable, "Sign-in failed", "Could not create your account.") + } + + sessionToken, err := h.issueSessionJWT(user, team) + if err != nil { + slog.Error("auth.google.start_callback.jwt_failed", "error", err, "request_id", requestID) + return renderAuthError(c, fiber.StatusServiceUnavailable, "Sign-in failed", "Could not issue session token.") + } + + slog.Info("auth.google.start_callback.success", + "user_id", user.ID, "team_id", team.ID, "request_id", requestID, + ) + + return c.Redirect(appendSessionToken(returnTo, sessionToken), fiber.StatusFound) +} + func (h *AuthHandler) findOrCreateUserGoogle(ctx context.Context, g *googleUser) (*models.User, *models.Team, error) { user, err := models.GetUserByGoogleID(ctx, h.db, g.Sub) if err == nil { diff --git a/internal/handlers/billing.go b/internal/handlers/billing.go index c8388957..89ddc982 100644 --- a/internal/handlers/billing.go +++ b/internal/handlers/billing.go @@ -42,7 +42,7 @@ func NewBillingHandler(db *sql.DB, cfg *config.Config, emailClient *email.Client return &BillingHandler{db: db, cfg: cfg, email: emailClient, migClient: migClient} } -// checkoutRequest is the request body for POST /billing/checkout. +// checkoutRequest is the request body for POST /api/v1/billing/checkout. type checkoutRequest struct { Plan string `json:"plan"` } @@ -76,10 +76,21 @@ func (h *BillingHandler) planIDToTier(planID string) string { return "pro" } -// CreateCheckout handles POST /billing/checkout. -// Creates a Razorpay subscription and returns the hosted payment URL. -// Requires a valid session JWT in the Authorization: Bearer header (enforced by RequireAuth middleware). -func (h *BillingHandler) CreateCheckout(c *fiber.Ctx) error { +// CreateCheckoutAPI handles POST /api/v1/billing/checkout (and the legacy +// alias POST /billing/checkout). Creates a Razorpay subscription and returns +// the hosted payment short_url plus the subscription_id. +// +// Requires a valid session JWT in the Authorization: Bearer header (enforced +// by RequireAuth middleware). +// +// Response: {"ok": true, "short_url": "...", "subscription_id": "..."} +// +// Status codes: +// - 400 invalid plan / invalid body +// - 401 no/invalid session (RequireAuth handles this) +// - 502 Razorpay rejected the create-subscription call +// - 503 RAZORPAY_KEY_ID/SECRET or the requested tier's plan_id not configured +func (h *BillingHandler) CreateCheckoutAPI(c *fiber.Ctx) error { requestID := middleware.GetRequestID(c) teamIDStr := middleware.GetTeamID(c) @@ -93,26 +104,46 @@ func (h *BillingHandler) CreateCheckout(c *fiber.Ctx) error { return respondError(c, fiber.StatusBadRequest, "invalid_body", "Request body must be valid JSON") } - planIDs := h.razorpayPlanIDs() - planID, ok := planIDs[body.Plan] - if !ok { - return respondError(c, fiber.StatusBadRequest, "invalid_plan", "plan must be 'hobby', 'pro', or 'team'") + plan := strings.ToLower(strings.TrimSpace(body.Plan)) + var planID string + switch plan { + case "hobby": + planID = h.cfg.RazorpayPlanIDHobby + case "pro": + planID = h.cfg.RazorpayPlanIDPro + case "team": + // Team tier is under development — block customer-initiated + // subscribe via the public API. The internal /internal/set-tier + // endpoint still works for ops use. Drop this guard when team + // launches (and revert the public pricing UI). + return respondError(c, fiber.StatusBadRequest, "tier_unavailable", + "Team tier is under active development. Email support@instanode.dev to join the early access list.") + default: + return respondError(c, fiber.StatusBadRequest, "invalid_plan", "plan must be 'hobby' or 'pro'") } - if h.cfg.RazorpayKeyID == "" || h.cfg.RazorpayKeySecret == "" { - return respondError(c, fiber.StatusServiceUnavailable, "billing_not_configured", "Billing is not configured") + if h.cfg.RazorpayKeyID == "" || h.cfg.RazorpayKeySecret == "" || planID == "" { + slog.Warn("billing.checkout.not_configured", + "team_id", teamID, + "plan", plan, + "key_set", h.cfg.RazorpayKeyID != "", + "secret_set", h.cfg.RazorpayKeySecret != "", + "plan_id_set", planID != "", + "request_id", requestID, + ) + return respondError(c, fiber.StatusServiceUnavailable, "billing_not_configured", "Razorpay credentials/plans not configured for this environment") } client := razorpay.NewClient(h.cfg.RazorpayKeyID, h.cfg.RazorpayKeySecret) subBody := map[string]interface{}{ "plan_id": planID, - "total_count": 120, // 10 years — cancel via subscription.cancelled webhook + "total_count": 12, // 12 billing cycles; cancel-at-cycle-end exits early via webhook "quantity": 1, "customer_notify": 1, "notes": map[string]interface{}{ "team_id": teamID.String(), - "plan": body.Plan, + "plan": plan, }, } @@ -121,33 +152,49 @@ func (h *BillingHandler) CreateCheckout(c *fiber.Ctx) error { slog.Error("billing.checkout.subscription_create_failed", "error", err, "team_id", teamID, + "plan", plan, "request_id", requestID, ) - return respondError(c, fiber.StatusServiceUnavailable, "razorpay_error", "Failed to create subscription") + return respondError(c, fiber.StatusBadGateway, "razorpay_error", "Razorpay rejected the subscription create call: "+err.Error()) } - // Persist subscription ID early for traceability; non-fatal if it fails. - if subID, ok := sub["id"].(string); ok && subID != "" { - if updateErr := models.UpdateRazorpaySubscriptionID(c.Context(), h.db, teamID, subID); updateErr != nil { - slog.Error("billing.checkout.update_subscription_id_failed", - "error", updateErr, - "team_id", teamID, - "request_id", requestID, - ) - } + subID, _ := sub["id"].(string) + shortURL, _ := sub["short_url"].(string) + + if subID == "" || shortURL == "" { + slog.Error("billing.checkout.razorpay_response_incomplete", + "team_id", teamID, + "plan", plan, + "sub_id_set", subID != "", + "short_url_set", shortURL != "", + "request_id", requestID, + ) + return respondError(c, fiber.StatusBadGateway, "razorpay_error", "Razorpay returned an incomplete subscription response") } - shortURL, _ := sub["short_url"].(string) + // Persist subscription ID early for traceability; non-fatal if it fails — the + // subscription.charged webhook will fall back to notes.team_id (or a DB lookup + // by sub_id once persisted via that webhook path). + if updateErr := models.UpdateRazorpaySubscriptionID(c.Context(), h.db, teamID, subID); updateErr != nil { + slog.Error("billing.checkout.update_subscription_id_failed", + "error", updateErr, + "team_id", teamID, + "subscription_id", subID, + "request_id", requestID, + ) + } slog.Info("billing.checkout.created", "team_id", teamID, - "plan", body.Plan, + "plan", plan, + "subscription_id", subID, "request_id", requestID, ) return c.JSON(fiber.Map{ - "ok": true, - "checkout_url": shortURL, + "ok": true, + "short_url": shortURL, + "subscription_id": subID, }) } @@ -365,7 +412,9 @@ func resolveTeamFromNotes(ctx context.Context, h *BillingHandler, sub rzpSubscri return id, nil } } - // Fallback: look up by subscription ID stored in stripe_customer_id column. + // Fallback: look up by Razorpay subscription ID. (The column is still named + // stripe_customer_id in the schema for legacy reasons — it now stores + // Razorpay subscription IDs. Rename pending — see TODO in models/team.go.) if sub.ID != "" { team, err := models.GetTeamByRazorpaySubscriptionID(ctx, h.db, sub.ID) if err != nil { @@ -585,6 +634,13 @@ func (h *BillingHandler) ChangePlanAPI(c *fiber.Ctx) error { if _, ok := planIDs[target]; !ok { return respondError(c, fiber.StatusBadRequest, "invalid_plan", "target_plan must be hobby, pro, or team") } + // Team tier is under development — block customer-initiated upgrades to + // team via the public API. The internal /internal/set-tier endpoint + // still works for ops use. Drop this guard when team launches. + if strings.EqualFold(target, "team") { + return respondError(c, fiber.StatusBadRequest, "tier_unavailable", + "Team tier is under active development. Email support@instanode.dev to join the early access list.") + } portal := &razorpaybilling.Portal{DB: h.db, Cfg: h.cfg} if _, err := portal.SubscriptionID(c.Context(), teamID); err != nil { return respondError(c, fiber.StatusBadRequest, "no_subscription", "no active subscription to change") diff --git a/internal/handlers/cache.go b/internal/handlers/cache.go index 7463b98c..05ef28f4 100644 --- a/internal/handlers/cache.go +++ b/internal/handlers/cache.go @@ -14,6 +14,7 @@ import ( "time" "github.com/gofiber/fiber/v2" + "github.com/google/uuid" "github.com/redis/go-redis/v9" "instant.dev/internal/config" "instant.dev/internal/crypto" @@ -55,8 +56,9 @@ func (h *CacheHandler) provisionCache(ctx context.Context, token, tier string) ( return nil, err } return &cacheprovider.Credentials{ - URL: creds.URL, - KeyPrefix: creds.KeyPrefix, + URL: creds.URL, + KeyPrefix: creds.KeyPrefix, + ProviderResourceID: creds.ProviderResourceID, }, nil } return h.cacheProvider.Provision(ctx, token, tier) @@ -80,9 +82,14 @@ func (h *CacheHandler) NewCache(c *fiber.Ctx) error { _ = c.BodyParser(&body) body.Name = sanitizeName(body.Name) + env, envErr := resolveEnv(c, body.Env) + if envErr != nil { + return envErr + } + // ── Authenticated path ──────────────────────────────────────────────────── if teamIDStr := middleware.GetTeamID(c); teamIDStr != "" { - return h.newCacheAuthenticated(c, teamIDStr, fp, country, vendor, requestID, body.Name, body.Dedicated, start) + return h.newCacheAuthenticated(c, teamIDStr, fp, country, vendor, requestID, body.Name, body.Dedicated, env, start) } // ── Dedicated requires authentication ───────────────────────────────────── @@ -124,6 +131,7 @@ func (h *CacheHandler) NewCache(c *fiber.Ctx) error { "name": existing.Name.String, "connection_url": connectionURL, "tier": existing.Tier, + "env": existing.Env, "limits": cacheAnonymousLimits(), "note": limitExceededNote(upgradeURL, existing.ExpiresAt.Time), "upgrade": upgradeURL, @@ -146,6 +154,7 @@ func (h *CacheHandler) NewCache(c *fiber.Ctx) error { ResourceType: "redis", Name: body.Name, Tier: "anonymous", + Env: env, Fingerprint: fp, CloudVendor: vendor, CountryCode: country, @@ -200,6 +209,13 @@ func (h *CacheHandler) NewCache(c *fiber.Ctx) error { } } + // Persist provider_resource_id (k8s namespace for dedicated Redis pods). + if creds.ProviderResourceID != "" { + if upErr := models.UpdateProviderResourceID(ctx, h.db, resource.ID, creds.ProviderResourceID); upErr != nil { + slog.Error("cache.new.update_provider_resource_id_failed", "error", upErr, "request_id", requestID) + } + } + jwtToken, jti, jwtErr := h.issueOnboardingJWT(ctx, fp, country, vendor, "redis", []string{tokenStr}) if jwtErr != nil { slog.Error("cache.new.jwt_issue_failed", "error", jwtErr, "request_id", requestID) @@ -238,6 +254,7 @@ func (h *CacheHandler) NewCache(c *fiber.Ctx) error { "name": resource.Name.String, "connection_url": creds.URL, "tier": "anonymous", + "env": resource.Env, "limits": cacheAnonymousLimits(), "note": upgradeNote(upgradeURL), } @@ -252,7 +269,7 @@ func (h *CacheHandler) NewCache(c *fiber.Ctx) error { } func (h *CacheHandler) newCacheAuthenticated( - c *fiber.Ctx, teamIDStr, fp, country, vendor, requestID, name string, dedicated bool, start time.Time, + c *fiber.Ctx, teamIDStr, fp, country, vendor, requestID, name string, dedicated bool, env string, start time.Time, ) error { ctx := c.UserContext() teamUUID, err := parseTeamID(teamIDStr) @@ -275,6 +292,7 @@ func (h *CacheHandler) newCacheAuthenticated( ResourceType: "redis", Name: name, Tier: tier, + Env: env, Fingerprint: fp, CloudVendor: vendor, CountryCode: country, @@ -286,6 +304,18 @@ func (h *CacheHandler) newCacheAuthenticated( return respondError(c, fiber.StatusServiceUnavailable, "provision_failed", "Failed to provision Redis resource") } + // Best-effort audit event; failures must never block the provision. + go func() { + _ = models.InsertAuditEvent(context.Background(), h.db, models.AuditEvent{ + TeamID: teamUUID, + Actor: "agent", + Kind: "provision", + ResourceType: "redis", + ResourceID: uuid.NullUUID{UUID: resource.ID, Valid: true}, + Summary: "agent provisioned redis " + resource.Token.String()[:8] + "", + }) + }() + tokenStr := resource.Token.String() // Provision the real Redis namespace. @@ -326,6 +356,13 @@ func (h *CacheHandler) newCacheAuthenticated( } } + // Persist provider_resource_id (k8s namespace for dedicated Redis pods). + if creds.ProviderResourceID != "" { + if upErr := models.UpdateProviderResourceID(ctx, h.db, resource.ID, creds.ProviderResourceID); upErr != nil { + slog.Error("cache.new.update_provider_resource_id_failed_auth", "error", upErr, "request_id", requestID) + } + } + slog.Info("provision.success", "service", "redis", "token", tokenStr, @@ -347,6 +384,7 @@ func (h *CacheHandler) newCacheAuthenticated( "name": resource.Name.String, "connection_url": creds.URL, "tier": tier, + "env": resource.Env, "dedicated": dedicated, "limits": fiber.Map{ "memory_mb": cacheAuthStorageLimitMB, diff --git a/internal/handlers/custom_domain.go b/internal/handlers/custom_domain.go new file mode 100644 index 00000000..c064708d --- /dev/null +++ b/internal/handlers/custom_domain.go @@ -0,0 +1,594 @@ +package handlers + +// custom_domain.go — Pro+ "bring your own hostname" for stacks. +// +// Routes (registered in router.go inside the auth-required /api/v1 group): +// +// POST /api/v1/stacks/:slug/domains create + return TXT challenge +// GET /api/v1/stacks/:slug/domains list domains for the stack +// POST /api/v1/stacks/:slug/domains/:id/verify re-run verification + ingress + cert +// DELETE /api/v1/stacks/:slug/domains/:id remove ingress + DB row +// +// The verification flow advances the row through: +// pending_verification → verified → ingress_ready → cert_ready (→ live) +// +// Verify is intentionally idempotent — the dashboard polls it once a few +// seconds while DNS propagates and again while Let's Encrypt issues. Each +// call is cheap when there is nothing new to do. + +import ( + "context" + "database/sql" + "errors" + "fmt" + "log/slog" + "net" + "net/url" + "strings" + "time" + + "github.com/gofiber/fiber/v2" + "github.com/google/uuid" + + "instant.dev/internal/config" + "instant.dev/internal/middleware" + "instant.dev/internal/models" + "instant.dev/internal/plans" + "instant.dev/internal/providers/compute/k8s" +) + +// CustomDomainProvider is the slice of K8sStackProvider this handler needs. +// Defined as an interface so tests can stub out k8s without spinning a +// clientset; production wires the real *k8s.K8sStackProvider. +type CustomDomainProvider interface { + EnsureCustomDomainIngress(ctx context.Context, stackNamespace, hostname, serviceName string, servicePort int) (string, error) + DeleteCustomDomainIngress(ctx context.Context, stackNamespace, hostname, serviceName string) error + CertificateReady(ctx context.Context, namespace, certName string) (bool, string, error) +} + +// reservedHostSuffixes is the central allowlist of suffixes a customer may +// NOT bind. Keeps anyone from claiming our own subdomains via a hostile DNS +// proof. Order matters only for readability — every entry is checked. +var reservedHostSuffixes = []string{ + ".instanode.dev", + ".deployment.instanode.dev", + ".instant.dev", + ".deployment.instant.dev", +} + +// reservedHosts is the central allowlist of exact hostnames that may NOT be +// bound. Avoids someone claiming the apex domain itself. +var reservedHosts = []string{ + "instanode.dev", + "instant.dev", + "deployment.instanode.dev", + "deployment.instant.dev", +} + +// dnsLookupTimeout caps how long Verify spends on a single TXT lookup. The +// resolver can hang indefinitely if upstream DNS is unhappy; 5s is plenty +// for a TXT query that exists. +const dnsLookupTimeout = 5 * time.Second + +// CustomDomainHandler serves /api/v1/stacks/:slug/domains*. +type CustomDomainHandler struct { + db *sql.DB + cfg *config.Config + plans *plans.Registry + k8s CustomDomainProvider +} + +// NewCustomDomainHandler wires the handler. k8sProvider may be nil; in that +// case ingress / cert operations are skipped and the rows stay at "verified". +func NewCustomDomainHandler(db *sql.DB, cfg *config.Config, planRegistry *plans.Registry, k8sProvider CustomDomainProvider) *CustomDomainHandler { + return &CustomDomainHandler{ + db: db, + cfg: cfg, + plans: planRegistry, + k8s: k8sProvider, + } +} + +// ── helpers ─────────────────────────────────────────────────────────────────── + +// validateHostname rejects empty / malformed input and refuses anything that +// would land on our own subdomains. Returns the lowercased canonical form on +// success. +// +// We do not enforce DNS-1123 label-length here; the customer's resolver will +// reject anything truly bizarre. The reserved-suffix guard is the load-bearing +// piece — if it ever returns "ok" for a suffix we own, a customer could bind +// `.instanode.dev` and steal our certs. Keep the logic centralised +// so future review is easy. +func validateHostname(raw string) (string, error) { + host := strings.ToLower(strings.TrimSpace(raw)) + if host == "" { + return "", errors.New("hostname is required") + } + // Reject schemes / paths — accept naked hostnames only. + if strings.Contains(host, "://") || strings.ContainsAny(host, "/?# ") { + return "", errors.New("hostname must be a bare domain (no scheme, path, or whitespace)") + } + // Strip a trailing dot if present (FQDN form). + host = strings.TrimSuffix(host, ".") + // At least one dot — the customer's apex `example.com` is fine, but an + // empty label like just "app" is not. + if !strings.Contains(host, ".") { + return "", errors.New("hostname must include a dot (e.g. app.example.com)") + } + // Don't allow port numbers. + if strings.Contains(host, ":") { + return "", errors.New("hostname must not include a port") + } + // Use net/url to catch the truly malformed. + if _, err := url.Parse("http://" + host); err != nil { + return "", fmt.Errorf("hostname is not a valid domain: %w", err) + } + // Reject our own zones. + for _, exact := range reservedHosts { + if host == exact { + return "", fmt.Errorf("hostname %q is reserved", host) + } + } + for _, suffix := range reservedHostSuffixes { + if strings.HasSuffix(host, suffix) { + return "", fmt.Errorf("hostname %q falls under reserved suffix %q", host, suffix) + } + } + return host, nil +} + +// requireTeam mirrors the helper used by other authenticated handlers. The +// router's RequireAuth middleware guarantees a team_id will be present. +func (h *CustomDomainHandler) requireTeam(c *fiber.Ctx) (*models.Team, error) { + teamIDStr := middleware.GetTeamID(c) + if teamIDStr == "" { + return nil, respondError(c, fiber.StatusUnauthorized, "unauthorized", + "Authentication required for custom domain operations") + } + teamUUID, err := parseTeamID(teamIDStr) + if err != nil { + return nil, respondError(c, fiber.StatusBadRequest, "invalid_team", + "Team ID in token is not a valid UUID") + } + team, err := models.GetTeamByID(c.Context(), h.db, teamUUID) + if err != nil { + slog.Error("custom_domain.team_lookup_failed", + "error", err, "team_id", teamIDStr, + "request_id", middleware.GetRequestID(c)) + return nil, respondError(c, fiber.StatusServiceUnavailable, "team_lookup_failed", + "Failed to look up team") + } + return team, nil +} + +// requireOwnedStack fetches the stack by slug and verifies the team owns it. +// Returns *models.Stack on success; writes the error response and returns +// (nil, err) on failure so callers can short-circuit. +func (h *CustomDomainHandler) requireOwnedStack(c *fiber.Ctx, team *models.Team, slug string) (*models.Stack, error) { + stack, err := models.GetStackBySlug(c.Context(), h.db, slug) + if err != nil { + var notFound *models.ErrStackNotFound + if errors.As(err, ¬Found) { + return nil, respondError(c, fiber.StatusNotFound, "not_found", "Stack not found") + } + slog.Error("custom_domain.stack_lookup_failed", + "error", err, "slug", slug, + "request_id", middleware.GetRequestID(c)) + return nil, respondError(c, fiber.StatusServiceUnavailable, "fetch_failed", "Failed to fetch stack") + } + // Anonymous stacks can't carry custom domains — they have no team. + if stack.TeamID == nil || *stack.TeamID != team.ID { + return nil, respondError(c, fiber.StatusNotFound, "not_found", "Stack not found") + } + return stack, nil +} + +// requireOwnedDomain fetches the row by id and asserts (a) it exists and (b) +// the requesting team owns it AND (c) it is bound to the given stack. +// Used by Verify and Delete to defend against teams reading another team's +// rows by guessing UUIDs. +func (h *CustomDomainHandler) requireOwnedDomain(c *fiber.Ctx, team *models.Team, stack *models.Stack, idStr string) (*models.CustomDomain, error) { + id, err := uuid.Parse(idStr) + if err != nil { + return nil, respondError(c, fiber.StatusBadRequest, "invalid_id", "Domain id must be a UUID") + } + dom, err := models.GetCustomDomainByID(c.Context(), h.db, id) + if err != nil { + if errors.Is(err, models.ErrCustomDomainNotFound) { + return nil, respondError(c, fiber.StatusNotFound, "not_found", "Custom domain not found") + } + slog.Error("custom_domain.lookup_failed", + "error", err, "id", id, + "request_id", middleware.GetRequestID(c)) + return nil, respondError(c, fiber.StatusServiceUnavailable, "fetch_failed", "Failed to fetch custom domain") + } + if dom.TeamID != team.ID || dom.StackID != stack.ID { + // 404 (not 403) so we never confirm "this UUID exists, just not yours". + return nil, respondError(c, fiber.StatusNotFound, "not_found", "Custom domain not found") + } + return dom, nil +} + +// expectedTXTValue returns the literal string the customer must include in +// their TXT record at "_instanode.". +func expectedTXTValue(token string) string { + return models.VerificationTokenPrefix + token +} + +// txtChallengeRecordName returns "_instanode." — where the customer +// adds their TXT record. We use the same name verbatim in the lookup so the +// payload matches the documentation exactly. +func txtChallengeRecordName(hostname string) string { + return "_instanode." + hostname +} + +// stackCNAMETarget is what the customer should set as a CNAME for their +// hostname. After verification, traffic to the custom hostname has to find +// our ingress controller, which fronts .deployment.instanode.dev. +func stackCNAMETarget(slug string) string { + return slug + ".deployment.instanode.dev" +} + +// dnsInstructions returns the JSON the API should hand back so the dashboard +// can render the right "next step" panel. We always return BOTH the TXT and +// CNAME instructions but mark which one is currently outstanding via the +// status field — clients can render either one without re-asking. +func dnsInstructions(dom *models.CustomDomain, stackSlug string) fiber.Map { + return fiber.Map{ + "txt": fiber.Map{ + "record_type": "TXT", + "record_name": txtChallengeRecordName(dom.Hostname), + "record_value": expectedTXTValue(dom.VerificationToken), + }, + "cname": fiber.Map{ + "record_type": "CNAME", + "record_name": dom.Hostname, + "record_value": stackCNAMETarget(stackSlug), + }, + } +} + +// serializeDomain shapes a CustomDomain for the API response, including the +// DNS instructions and a flag mirroring whether the cert is ready (callers +// poll this from the dashboard). +func serializeDomain(dom *models.CustomDomain, stackSlug string) fiber.Map { + out := fiber.Map{ + "id": dom.ID, + "hostname": dom.Hostname, + "status": dom.Status, + "created_at": dom.CreatedAt, + "verification": dnsInstructions(dom, stackSlug), + "verified": dom.Status != models.CustomDomainStatusPending, + "certificate_ready": dom.Status == models.CustomDomainStatusCertReady || dom.Status == models.CustomDomainStatusLive, + } + if dom.VerifiedAt.Valid { + out["verified_at"] = dom.VerifiedAt.Time + } + if dom.CertReadyAt.Valid { + out["cert_ready_at"] = dom.CertReadyAt.Time + } + if dom.LastCheckAt.Valid { + out["last_check_at"] = dom.LastCheckAt.Time + } + if dom.LastCheckErr.Valid { + out["last_check_err"] = dom.LastCheckErr.String + } + return out +} + +// primaryStackService returns the service we'll route the custom hostname at. +// We pick the first service with expose=true so customers get the same +// service that's already serving traffic on the deployment.instanode.dev URL. +// If no service is exposed, returns ("", err). +func (h *CustomDomainHandler) primaryStackService(ctx context.Context, stack *models.Stack) (*models.StackService, error) { + svcs, err := models.GetStackServicesByStack(ctx, h.db, stack.ID) + if err != nil { + return nil, fmt.Errorf("primaryStackService: %w", err) + } + for _, ss := range svcs { + if ss.Expose { + return ss, nil + } + } + return nil, errors.New("stack has no service marked expose=true") +} + +// ── POST /api/v1/stacks/:slug/domains ───────────────────────────────────────── + +type createCustomDomainBody struct { + Hostname string `json:"hostname"` +} + +// Create handles POST /api/v1/stacks/:slug/domains. +func (h *CustomDomainHandler) Create(c *fiber.Ctx) error { + team, err := h.requireTeam(c) + if err != nil { + return err + } + + // Tier gate — Pro+ only. Hobby / anonymous get a 402-style upgrade hint. + if !h.plans.CustomDomainsAllowed(team.PlanTier) { + return respondError(c, fiber.StatusPaymentRequired, "upgrade_required", + "Custom domains require the Pro plan or higher. Upgrade at https://instanode.dev/pricing") + } + + stack, err := h.requireOwnedStack(c, team, c.Params("slug")) + if err != nil { + return err + } + + var body createCustomDomainBody + if err := c.BodyParser(&body); err != nil { + return respondError(c, fiber.StatusBadRequest, "invalid_body", + `Body must be valid JSON: {"hostname":"app.example.com"}`) + } + + hostname, valErr := validateHostname(body.Hostname) + if valErr != nil { + return respondError(c, fiber.StatusBadRequest, "invalid_hostname", valErr.Error()) + } + + dom, err := models.CreateCustomDomain(c.Context(), h.db, team.ID, stack.ID, hostname) + if err != nil { + if errors.Is(err, models.ErrCustomDomainTaken) { + return respondError(c, fiber.StatusConflict, "hostname_taken", + "This hostname is already bound to another domain. Delete the existing binding first or contact support.") + } + slog.Error("custom_domain.create_failed", + "error", err, "hostname", hostname, + "team_id", team.ID, "stack_id", stack.ID, + "request_id", middleware.GetRequestID(c)) + return respondError(c, fiber.StatusServiceUnavailable, "create_failed", + "Failed to create custom domain") + } + + slog.Info("custom_domain.created", + "id", dom.ID, "hostname", hostname, + "team_id", team.ID, "stack_slug", stack.Slug, + "request_id", middleware.GetRequestID(c)) + + return c.Status(fiber.StatusCreated).JSON(fiber.Map{ + "ok": true, + "domain": serializeDomain(dom, stack.Slug), + }) +} + +// ── GET /api/v1/stacks/:slug/domains ────────────────────────────────────────── + +// List handles GET /api/v1/stacks/:slug/domains. +func (h *CustomDomainHandler) List(c *fiber.Ctx) error { + team, err := h.requireTeam(c) + if err != nil { + return err + } + stack, err := h.requireOwnedStack(c, team, c.Params("slug")) + if err != nil { + return err + } + + doms, err := models.ListCustomDomainsByStack(c.Context(), h.db, stack.ID) + if err != nil { + slog.Error("custom_domain.list_failed", + "error", err, "stack_id", stack.ID, + "request_id", middleware.GetRequestID(c)) + return respondError(c, fiber.StatusServiceUnavailable, "list_failed", + "Failed to list custom domains") + } + items := make([]fiber.Map, 0, len(doms)) + for _, d := range doms { + items = append(items, serializeDomain(d, stack.Slug)) + } + return c.JSON(fiber.Map{ + "ok": true, + "items": items, + "total": len(items), + }) +} + +// ── POST /api/v1/stacks/:slug/domains/:id/verify ────────────────────────────── + +// Verify is idempotent. Each call: +// +// 1. If status == pending_verification — re-runs the TXT lookup; advances +// to "verified" if it matches, otherwise records last_check_err. +// 2. If status >= verified but no Ingress yet — creates the Ingress + +// Certificate (cert-manager auto-creates the cert once it sees the +// annotated Ingress + missing TLS Secret) and advances to ingress_ready. +// 3. If status >= ingress_ready — polls the Certificate for Ready=True and +// advances to cert_ready when the cert lands. +// +// The response always reflects the state AFTER this call's mutations. +func (h *CustomDomainHandler) Verify(c *fiber.Ctx) error { + team, err := h.requireTeam(c) + if err != nil { + return err + } + stack, err := h.requireOwnedStack(c, team, c.Params("slug")) + if err != nil { + return err + } + dom, err := h.requireOwnedDomain(c, team, stack, c.Params("id")) + if err != nil { + return err + } + + // Step 1: TXT lookup if still pending. + if dom.Status == models.CustomDomainStatusPending { + ok, lookupErr := h.checkTXT(c.Context(), dom) + if ok { + if mkErr := models.MarkCustomDomainVerified(c.Context(), h.db, dom.ID); mkErr != nil { + slog.Error("custom_domain.mark_verified_failed", + "error", mkErr, "id", dom.ID) + return respondError(c, fiber.StatusServiceUnavailable, "verify_failed", + "Failed to record verification") + } + // Reload after mutation so subsequent steps see the new status. + dom, err = models.GetCustomDomainByID(c.Context(), h.db, dom.ID) + if err != nil { + return respondError(c, fiber.StatusServiceUnavailable, "fetch_failed", + "Failed to refresh domain after verification") + } + } else { + msg := "TXT record missing or wrong value" + if lookupErr != nil { + msg = lookupErr.Error() + } + _ = models.UpdateCustomDomainStatus(c.Context(), h.db, dom.ID, models.CustomDomainStatusPending, msg) + dom.LastCheckErr = sql.NullString{String: msg, Valid: true} + // 200 with current state + the failure reason — clients poll. + return c.JSON(fiber.Map{ + "ok": true, + "domain": serializeDomain(dom, stack.Slug), + }) + } + } + + // Step 2: Ensure the Ingress exists once we're at "verified". + if dom.Status == models.CustomDomainStatusVerified { + if h.k8s == nil { + // No k8s wired in this environment (e.g. tests). Treat verification + // as the terminal state and let the dashboard show "TXT verified — ingress pending." + return c.JSON(fiber.Map{ + "ok": true, + "domain": serializeDomain(dom, stack.Slug), + }) + } + svc, svcErr := h.primaryStackService(c.Context(), stack) + if svcErr != nil { + _ = models.UpdateCustomDomainStatus(c.Context(), h.db, dom.ID, models.CustomDomainStatusVerified, svcErr.Error()) + dom.LastCheckErr = sql.NullString{String: svcErr.Error(), Valid: true} + return c.JSON(fiber.Map{ + "ok": true, + "domain": serializeDomain(dom, stack.Slug), + }) + } + + _, ingErr := h.k8s.EnsureCustomDomainIngress(c.Context(), stack.Namespace, dom.Hostname, svc.Name, svc.Port) + if ingErr != nil { + slog.Error("custom_domain.ingress_failed", + "error", ingErr, "id", dom.ID, "hostname", dom.Hostname, + "namespace", stack.Namespace, + "request_id", middleware.GetRequestID(c)) + _ = models.UpdateCustomDomainStatus(c.Context(), h.db, dom.ID, models.CustomDomainStatusVerified, ingErr.Error()) + dom.LastCheckErr = sql.NullString{String: ingErr.Error(), Valid: true} + return c.JSON(fiber.Map{ + "ok": true, + "domain": serializeDomain(dom, stack.Slug), + }) + } + if mkErr := models.UpdateCustomDomainStatus(c.Context(), h.db, dom.ID, models.CustomDomainStatusIngressReady, ""); mkErr != nil { + slog.Error("custom_domain.set_ingress_ready_failed", + "error", mkErr, "id", dom.ID) + } + dom.Status = models.CustomDomainStatusIngressReady + } + + // Step 3: Poll the Certificate for Ready=True. + if dom.Status == models.CustomDomainStatusIngressReady && h.k8s != nil { + certName := k8s.CustomDomainTLSSecretName(dom.Hostname) + ready, certMsg, certErr := h.k8s.CertificateReady(c.Context(), stack.Namespace, certName) + if certErr != nil { + slog.Warn("custom_domain.cert_poll_failed", + "error", certErr, "id", dom.ID, "hostname", dom.Hostname, + "namespace", stack.Namespace) + // Soft-fail: leave the row at ingress_ready and surface the message. + _ = models.UpdateCustomDomainStatus(c.Context(), h.db, dom.ID, models.CustomDomainStatusIngressReady, certErr.Error()) + dom.LastCheckErr = sql.NullString{String: certErr.Error(), Valid: true} + } else if ready { + if mkErr := models.MarkCertReady(c.Context(), h.db, dom.ID); mkErr != nil { + slog.Error("custom_domain.mark_cert_ready_failed", + "error", mkErr, "id", dom.ID) + } else { + dom.Status = models.CustomDomainStatusCertReady + dom.CertReadyAt = sql.NullTime{Time: time.Now(), Valid: true} + dom.LastCheckErr = sql.NullString{} + } + } else { + // Still issuing — record the cert-manager message so the dashboard + // can surface "DNS validation pending" / "ACME order created". + _ = models.UpdateCustomDomainStatus(c.Context(), h.db, dom.ID, models.CustomDomainStatusIngressReady, certMsg) + dom.LastCheckErr = sql.NullString{String: certMsg, Valid: certMsg != ""} + } + } + + return c.JSON(fiber.Map{ + "ok": true, + "domain": serializeDomain(dom, stack.Slug), + }) +} + +// checkTXT runs net.LookupTXT against the verification record and reports +// whether the expected payload appears in any returned record. +func (h *CustomDomainHandler) checkTXT(ctx context.Context, dom *models.CustomDomain) (bool, error) { + lookupCtx, cancel := context.WithTimeout(ctx, dnsLookupTimeout) + defer cancel() + resolver := net.DefaultResolver + records, err := resolver.LookupTXT(lookupCtx, txtChallengeRecordName(dom.Hostname)) + if err != nil { + return false, fmt.Errorf("TXT lookup for %s failed: %w", txtChallengeRecordName(dom.Hostname), err) + } + want := expectedTXTValue(dom.VerificationToken) + for _, r := range records { + // Some resolvers return the TXT contents wrapped in extra quotes; trim them. + clean := strings.Trim(r, "\"") + if clean == want || r == want { + return true, nil + } + } + return false, nil +} + +// ── DELETE /api/v1/stacks/:slug/domains/:id ─────────────────────────────────── + +// Delete removes the Ingress + Secret (best-effort) and then the DB row. +// We tear down k8s before the DB row so a partial failure leaves the row in +// place and the customer can retry. If k8s already lost the Ingress we +// continue and clear the row anyway. +func (h *CustomDomainHandler) Delete(c *fiber.Ctx) error { + team, err := h.requireTeam(c) + if err != nil { + return err + } + stack, err := h.requireOwnedStack(c, team, c.Params("slug")) + if err != nil { + return err + } + dom, err := h.requireOwnedDomain(c, team, stack, c.Params("id")) + if err != nil { + return err + } + + // Best-effort ingress teardown. We need a service name; fall back to the + // primary one. If lookup fails (e.g. stack already gone), continue. + if h.k8s != nil { + if svc, svcErr := h.primaryStackService(c.Context(), stack); svcErr == nil { + if delErr := h.k8s.DeleteCustomDomainIngress(c.Context(), stack.Namespace, dom.Hostname, svc.Name); delErr != nil { + slog.Warn("custom_domain.delete.ingress_teardown_failed", + "error", delErr, "id", dom.ID, "hostname", dom.Hostname) + } + } + } + + if err := models.DeleteCustomDomain(c.Context(), h.db, dom.ID, team.ID); err != nil { + if errors.Is(err, models.ErrCustomDomainNotFound) { + return respondError(c, fiber.StatusNotFound, "not_found", "Custom domain not found") + } + slog.Error("custom_domain.delete_failed", + "error", err, "id", dom.ID, + "request_id", middleware.GetRequestID(c)) + return respondError(c, fiber.StatusServiceUnavailable, "delete_failed", + "Failed to delete custom domain") + } + + slog.Info("custom_domain.deleted", + "id", dom.ID, "hostname", dom.Hostname, + "team_id", team.ID, "stack_slug", stack.Slug, + "request_id", middleware.GetRequestID(c)) + + return c.JSON(fiber.Map{ + "ok": true, + "id": dom.ID, + "message": "Custom domain removed", + }) +} diff --git a/internal/handlers/db.go b/internal/handlers/db.go index 3508529d..4e3e817f 100644 --- a/internal/handlers/db.go +++ b/internal/handlers/db.go @@ -11,6 +11,7 @@ package handlers // "name": "my-db", // "connection_url": "postgres://usr_:@postgres-customers:5432/db_", // "tier": "anonymous", +// "env": "production", // "limits": { "storage_mb": 10, "connections": 3, "expires_in": "24h" }, // "note": "Works now. Free forever with a free account: " // } @@ -23,6 +24,7 @@ import ( "time" "github.com/gofiber/fiber/v2" + "github.com/google/uuid" "github.com/redis/go-redis/v9" "instant.dev/internal/config" "instant.dev/internal/crypto" @@ -91,9 +93,14 @@ func (h *DBHandler) NewDB(c *fiber.Ctx) error { _ = c.BodyParser(&body) body.Name = sanitizeName(body.Name) + env, envErr := resolveEnv(c, body.Env) + if envErr != nil { + return envErr + } + // ── Authenticated path ──────────────────────────────────────────────────── if teamIDStr := middleware.GetTeamID(c); teamIDStr != "" { - return h.newDBAuthenticated(c, teamIDStr, fp, country, vendor, requestID, body.Name, body.Dedicated, start) + return h.newDBAuthenticated(c, teamIDStr, fp, country, vendor, requestID, body.Name, body.Dedicated, env, start) } // ── Dedicated requires authentication ───────────────────────────────────── @@ -136,6 +143,7 @@ func (h *DBHandler) NewDB(c *fiber.Ctx) error { "name": existing.Name.String, "connection_url": connectionURL, "tier": existing.Tier, + "env": existing.Env, "limits": dbAnonymousLimits(), "note": limitExceededNote(upgradeURL, existing.ExpiresAt.Time), "upgrade": upgradeURL, @@ -155,6 +163,7 @@ func (h *DBHandler) NewDB(c *fiber.Ctx) error { ResourceType: "postgres", Name: body.Name, Tier: "anonymous", + Env: env, Fingerprint: fp, CloudVendor: vendor, CountryCode: country, @@ -245,6 +254,7 @@ func (h *DBHandler) NewDB(c *fiber.Ctx) error { "name": resource.Name.String, "connection_url": creds.URL, "tier": "anonymous", + "env": resource.Env, "limits": dbAnonymousLimits(), "note": upgradeNote(upgradeURL), } @@ -256,7 +266,7 @@ func (h *DBHandler) NewDB(c *fiber.Ctx) error { } func (h *DBHandler) newDBAuthenticated( - c *fiber.Ctx, teamIDStr, fp, country, vendor, requestID, name string, dedicated bool, start time.Time, + c *fiber.Ctx, teamIDStr, fp, country, vendor, requestID, name string, dedicated bool, env string, start time.Time, ) error { ctx := c.UserContext() teamUUID, err := parseTeamID(teamIDStr) @@ -279,6 +289,7 @@ func (h *DBHandler) newDBAuthenticated( ResourceType: "postgres", Name: name, Tier: tier, + Env: env, Fingerprint: fp, CloudVendor: vendor, CountryCode: country, @@ -290,6 +301,18 @@ func (h *DBHandler) newDBAuthenticated( return respondError(c, fiber.StatusServiceUnavailable, "provision_failed", "Failed to provision Postgres resource") } + // Best-effort audit event; failures must never block the provision. + go func() { + _ = models.InsertAuditEvent(context.Background(), h.db, models.AuditEvent{ + TeamID: teamUUID, + Actor: "agent", + Kind: "provision", + ResourceType: "postgres", + ResourceID: uuid.NullUUID{UUID: resource.ID, Valid: true}, + Summary: "agent provisioned postgres " + resource.Token.String()[:8] + "", + }) + }() + tokenStr := resource.Token.String() // Provision the real Postgres database. @@ -349,6 +372,7 @@ func (h *DBHandler) newDBAuthenticated( "name": resource.Name.String, "connection_url": creds.URL, "tier": tier, + "env": resource.Env, "dedicated": dedicated, "limits": fiber.Map{ "storage_mb": authStorageLimitMB, diff --git a/internal/handlers/deploy.go b/internal/handlers/deploy.go index 8fb1f754..3cb8c3d7 100644 --- a/internal/handlers/deploy.go +++ b/internal/handlers/deploy.go @@ -76,6 +76,12 @@ func generateAppID() (string, error) { } // deploymentToMap converts a Deployment to a JSON-friendly fiber.Map. +// +// Naming collision note: prior to multi-environment support the response field +// "env" was already in use to expose the deployment's env_vars map. We keep +// that meaning for backwards compatibility and add a separate "environment" +// field for the new env scope (production / staging / dev / ...). Callers can +// continue to read .env as a map of vars; .environment is the scope name. func deploymentToMap(d *models.Deployment) fiber.Map { m := fiber.Map{ "id": d.ID, @@ -87,6 +93,7 @@ func deploymentToMap(d *models.Deployment) fiber.Map { "tier": d.Tier, "status": d.Status, "env": d.EnvVars, + "environment": d.Env, "created_at": d.CreatedAt, "updated_at": d.UpdatedAt, "team_id": d.TeamID, @@ -128,17 +135,30 @@ func (h *DeployHandler) requireTeam(c *fiber.Ctx) (*models.Team, error) { // runDeploy is run in a goroutine after POST /deploy/new returns 202. // It calls the compute provider, then updates the deployment record in DB. +// +// Before the compute call, every "vault://KEY" entry in d.EnvVars is replaced +// with the decrypted plaintext from the team's vault for d.Env. The plaintext +// is passed to the compute provider but never written back to the deployments +// row, so vault rotations take effect on the next redeploy. func (h *DeployHandler) runDeploy(d *models.Deployment, tarball []byte) { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Minute) defer cancel() + resolvedEnv, err := ResolveVaultRefs(ctx, h.db, h.cfg.AESKey, d.TeamID, d.Env, d.EnvVars) + if err != nil { + slog.Error("deploy.run_deploy.vault_resolve_failed", + "app_id", d.AppID, "team_id", d.TeamID, "env", d.Env, "error", err) + _ = models.UpdateDeploymentStatus(ctx, h.db, d.ID, "failed", err.Error()) + return + } + opts := compute.DeployOptions{ AppID: d.AppID, Token: d.ID.String(), Tarball: tarball, Port: d.Port, Tier: d.Tier, - EnvVars: d.EnvVars, + EnvVars: resolvedEnv, } result, err := h.compute.Deploy(ctx, opts) if err != nil { @@ -218,6 +238,18 @@ func (h *DeployHandler) New(c *fiber.Ctx) error { "Field 'port' must be between 1 and 65535") } + // Optional environment scope: ?env=staging or multipart "env" field. + // Empty defaults to "production". Validation is centralised in + // models.NormalizeEnv via resolveEnv. + envBody := "" + if vals := form.Value["env"]; len(vals) > 0 { + envBody = vals[0] + } + environment, envErr := resolveEnv(c, envBody) + if envErr != nil { + return envErr + } + // Generate app ID. appID, err := generateAppID() if err != nil { @@ -236,6 +268,7 @@ func (h *DeployHandler) New(c *fiber.Ctx) error { AppID: appID, Port: port, Tier: team.PlanTier, + Env: environment, EnvVars: initEnv, }) if err != nil { diff --git a/internal/handlers/env_test.go b/internal/handlers/env_test.go new file mode 100644 index 00000000..940f9c22 --- /dev/null +++ b/internal/handlers/env_test.go @@ -0,0 +1,184 @@ +package handlers_test + +// env_test.go — handler-level tests for multi-environment support +// (POST /db/new, /cache/new, /nosql/new, /storage/new, /webhook/new, /deploy/new). +// +// Each test asserts: +// - Missing ?env defaults to "production" in the response and DB row. +// - Invalid env strings are rejected with HTTP 400 + error="invalid_env". +// - Provisioning in env=staging does not appear in env=production listings. +// +// All tests skip when the test Postgres / Redis isn't reachable — they call +// testhelpers.NewTestApp which itself skips on unreachable infra. + +import ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "instant.dev/internal/models" + "instant.dev/internal/testhelpers" +) + +// postCacheNew posts to /cache/new with optional ?env query param and returns +// the parsed JSON body. We use cache as the canonical "smallest happy-path +// provision" — it has no external infra dependency beyond Redis itself. +func postCacheNew(t *testing.T, app interface { + Test(*http.Request, ...int) (*http.Response, error) +}, ip, env string) (int, map[string]any) { + t.Helper() + path := "/cache/new" + if env != "" { + path += "?env=" + env + } + req := httptest.NewRequest(http.MethodPost, path, nil) + req.Header.Set("X-Forwarded-For", ip) + + resp, err := app.Test(req, 5000) + require.NoError(t, err) + defer resp.Body.Close() + + body, _ := io.ReadAll(resp.Body) + var out map[string]any + if len(body) > 0 { + _ = json.Unmarshal(body, &out) + } + return resp.StatusCode, out +} + +func TestEnv_DefaultProduction(t *testing.T) { + db, cleanDB := testhelpers.SetupTestDB(t) + defer cleanDB() + rdb, cleanRedis := testhelpers.SetupTestRedis(t) + defer cleanRedis() + + app, cleanApp := testhelpers.NewTestApp(t, db, rdb) + defer cleanApp() + + status, body := postCacheNew(t, app, "10.42.0.1", "") + require.True(t, status == http.StatusCreated || status == http.StatusOK, + "expected 201/200, got %d (%v)", status, body) + + tokStr, _ := body["token"].(string) + require.NotEmpty(t, tokStr) + defer db.Exec(`DELETE FROM resources WHERE token = $1::uuid`, tokStr) + + gotEnv, _ := body["env"].(string) + assert.Equal(t, models.EnvProduction, gotEnv, + "missing ?env must default to 'production' in the response") + + // Verify it's also persisted as 'production'. + var dbEnv string + require.NoError(t, db.QueryRow(`SELECT env FROM resources WHERE token = $1::uuid`, tokStr).Scan(&dbEnv)) + assert.Equal(t, "production", dbEnv) +} + +func TestEnv_Validation_RejectsInvalid(t *testing.T) { + db, cleanDB := testhelpers.SetupTestDB(t) + defer cleanDB() + rdb, cleanRedis := testhelpers.SetupTestRedis(t) + defer cleanRedis() + + app, cleanApp := testhelpers.NewTestApp(t, db, rdb) + defer cleanApp() + + cases := []struct { + name string + env string + }{ + {"contains_space", "prod%20ction"}, // url-encoded space + {"too_long", strings.Repeat("a", 33)}, + {"uppercase", "Prod"}, + {"underscore", "my_env"}, + {"unicode", "stagé"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + status, body := postCacheNew(t, app, "10.43."+tc.name[:1]+".1", tc.env) + assert.Equal(t, http.StatusBadRequest, status, "body=%v", body) + assert.Equal(t, "invalid_env", body["error"], "body=%v", body) + }) + } +} + +func TestEnv_Isolation_ListResourcesByEnv(t *testing.T) { + db, cleanDB := testhelpers.SetupTestDB(t) + defer cleanDB() + + teamID := uuid.MustParse(testhelpers.MustCreateTeamDB(t, db, "hobby")) + defer db.Exec(`DELETE FROM teams WHERE id = $1`, teamID) + + mk := func(env string) *models.Resource { + r, err := models.CreateResource(context.Background(), db, models.CreateResourceParams{ + TeamID: &teamID, + ResourceType: "redis", + Tier: "hobby", + Env: env, + }) + require.NoError(t, err) + return r + } + stagingR := mk("staging") + prodR := mk("production") + defer db.Exec(`DELETE FROM resources WHERE id IN ($1, $2)`, stagingR.ID, prodR.ID) + + prodList, err := models.ListResourcesByTeamAndEnv(context.Background(), db, teamID, "production") + require.NoError(t, err) + for _, r := range prodList { + assert.NotEqual(t, stagingR.ID, r.ID, + "staging resource must NOT appear in production listing") + assert.Equal(t, "production", r.Env) + } + + stgList, err := models.ListResourcesByTeamAndEnv(context.Background(), db, teamID, "staging") + require.NoError(t, err) + var stgFound bool + for _, r := range stgList { + if r.ID == stagingR.ID { + stgFound = true + } + assert.Equal(t, "staging", r.Env) + } + assert.True(t, stgFound) +} + +func TestEnv_DeployIsolation(t *testing.T) { + db, cleanDB := testhelpers.SetupTestDB(t) + defer cleanDB() + + teamID := uuid.MustParse(testhelpers.MustCreateTeamDB(t, db, "hobby")) + defer db.Exec(`DELETE FROM teams WHERE id = $1`, teamID) + + dev, err := models.CreateDeployment(context.Background(), db, models.CreateDeploymentParams{ + TeamID: teamID, + AppID: "myapp-dev-" + uuid.NewString()[:6], + Tier: "hobby", + Env: "dev", + EnvVars: map[string]string{"_name": "myapp"}, + }) + require.NoError(t, err) + defer db.Exec(`DELETE FROM deployments WHERE id = $1`, dev.ID) + + prod, err := models.CreateDeployment(context.Background(), db, models.CreateDeploymentParams{ + TeamID: teamID, + AppID: "myapp-prod-" + uuid.NewString()[:6], + Tier: "hobby", + Env: "production", + EnvVars: map[string]string{"_name": "myapp"}, + }) + require.NoError(t, err) + defer db.Exec(`DELETE FROM deployments WHERE id = $1`, prod.ID) + + assert.NotEqual(t, dev.ID, prod.ID, + "same logical app (myapp) deployed to dev vs prod must be two distinct rows") + assert.Equal(t, "dev", dev.Env) + assert.Equal(t, "production", prod.Env) +} diff --git a/internal/handlers/magic_link.go b/internal/handlers/magic_link.go new file mode 100644 index 00000000..af6bba2e --- /dev/null +++ b/internal/handlers/magic_link.go @@ -0,0 +1,176 @@ +package handlers + +import ( + "database/sql" + "errors" + "log/slog" + "strings" + "time" + + "github.com/gofiber/fiber/v2" + "instant.dev/internal/config" + "instant.dev/internal/email" + "instant.dev/internal/middleware" + "instant.dev/internal/models" +) + +// magicLinkTTL is how long an emailed sign-in link remains valid. +// 15 minutes is long enough to survive an email-client preview round-trip +// and short enough that a leaked token is rarely useful. +const magicLinkTTL = 15 * time.Minute + +// MagicLinkHandler implements the passwordless email login flow: +// POST /auth/email/start — generates a token, emails the link, returns 202 +// GET /auth/email/callback — consumes the token, mints a session JWT, +// 302s back to the dashboard with ?session_token= +type MagicLinkHandler struct { + db *sql.DB + cfg *config.Config + mail *email.Client + auth *AuthHandler // for IssueSessionJWT + FindOrCreateUserByEmail +} + +// NewMagicLinkHandler wires the dependencies. Note that we take an AuthHandler +// rather than reimplementing user/team upsert and JWT signing — the magic-link +// flow lands users in exactly the same spot the GitHub/Google flows do. +func NewMagicLinkHandler(db *sql.DB, cfg *config.Config, mail *email.Client, auth *AuthHandler) *MagicLinkHandler { + return &MagicLinkHandler{db: db, cfg: cfg, mail: mail, auth: auth} +} + +// magicLinkStartRequest is the body for POST /auth/email/start. +type magicLinkStartRequest struct { + Email string `json:"email"` + ReturnTo string `json:"return_to"` +} + +// Start handles POST /auth/email/start. +// +// Always returns 202 (or 400 for malformed bodies) regardless of whether the +// email exists in our DB. Revealing existence here would let an attacker +// enumerate users by trying random addresses. +// +// Email send errors are logged but do NOT change the response: the user might +// still get the email seconds later through Resend's retry pipeline, and a +// timing/error-rate side-channel would defeat the enumeration defence above. +func (h *MagicLinkHandler) Start(c *fiber.Ctx) error { + requestID := middleware.GetRequestID(c) + + var body magicLinkStartRequest + if err := c.BodyParser(&body); err != nil { + return respondError(c, fiber.StatusBadRequest, "invalid_body", "Request body must be valid JSON") + } + + emailAddr := strings.ToLower(strings.TrimSpace(body.Email)) + if !looksLikeEmail(emailAddr) { + return respondError(c, fiber.StatusBadRequest, "invalid_email", "A valid email address is required") + } + + returnTo := validateReturnTo(strings.TrimSpace(body.ReturnTo)) + + plaintext, err := models.GenerateMagicLinkPlaintext() + if err != nil { + slog.Error("magic_link.start.generate_token", "error", err, "request_id", requestID) + // 202 anyway — never expose backend hiccups in this enumeration-sensitive + // endpoint. + return c.Status(fiber.StatusAccepted).JSON(fiber.Map{"ok": true}) + } + + if _, err := models.CreateMagicLink(c.Context(), h.db, emailAddr, plaintext, returnTo, magicLinkTTL); err != nil { + slog.Error("magic_link.start.db_insert", "error", err, "request_id", requestID) + return c.Status(fiber.StatusAccepted).JSON(fiber.Map{"ok": true}) + } + + link := canonicalAPIBase + "/auth/email/callback?t=" + plaintext + if err := h.mail.SendMagicLink(c.Context(), emailAddr, link); err != nil { + // Already logged inside email client; we just don't fail the request. + slog.Warn("magic_link.start.email_send_failed", "error", err, "request_id", requestID) + } + + slog.Info("magic_link.start.sent", + "request_id", requestID, + // email is intentionally NOT logged at info level to avoid PII spread — + // trace through the magic_links table by created_at if needed. + ) + + return c.Status(fiber.StatusAccepted).JSON(fiber.Map{"ok": true}) +} + +// Callback handles GET /auth/email/callback?t=. +// +// Validates the token, atomic-consumes it, finds-or-creates the user/team, +// mints a session JWT, and 302s to <return_to>?session_token=<jwt>. +// +// On any failure path, renders an HTML error page (the user is in a browser). +func (h *MagicLinkHandler) Callback(c *fiber.Ctx) error { + requestID := middleware.GetRequestID(c) + + plaintext := strings.TrimSpace(c.Query("t")) + if plaintext == "" { + return renderAuthError(c, fiber.StatusBadRequest, "Sign-in link is missing its token", "Open the link from your email exactly as we sent it.") + } + + hash := models.HashMagicLink(plaintext) + link, err := models.GetMagicLinkForConsumption(c.Context(), h.db, hash) + if err != nil { + if errors.Is(err, models.ErrMagicLinkNotFound) { + return renderAuthError(c, fiber.StatusBadRequest, "Sign-in link is invalid or expired", "Magic links last 15 minutes and can only be used once. Request a new one to continue.") + } + slog.Error("magic_link.callback.lookup_failed", "error", err, "request_id", requestID) + return renderAuthError(c, fiber.StatusServiceUnavailable, "Sign-in unavailable", "Please try again in a moment.") + } + + consumed, err := models.ConsumeMagicLink(c.Context(), h.db, link.ID) + if err != nil { + slog.Error("magic_link.callback.consume_failed", "error", err, "request_id", requestID, "link_id", link.ID) + return renderAuthError(c, fiber.StatusServiceUnavailable, "Sign-in unavailable", "Please try again in a moment.") + } + if !consumed { + // Race: somebody else consumed the row between SELECT and UPDATE. Treat + // as an already-used link. + return renderAuthError(c, fiber.StatusBadRequest, "Sign-in link already used", "Request a new sign-in email to continue.") + } + + user, team, err := h.auth.FindOrCreateUserByEmail(c.Context(), link.Email) + if err != nil { + slog.Error("magic_link.callback.user_upsert_failed", "error", err, "request_id", requestID, "link_id", link.ID) + return renderAuthError(c, fiber.StatusServiceUnavailable, "Sign-in failed", "Could not create your account. Please try again.") + } + + sessionToken, err := h.auth.IssueSessionJWT(user, team) + if err != nil { + slog.Error("magic_link.callback.jwt_failed", "error", err, "request_id", requestID) + return renderAuthError(c, fiber.StatusServiceUnavailable, "Sign-in failed", "Could not issue session token.") + } + + // link.ReturnTo went through validateReturnTo at insert time, but re-check + // as defence-in-depth in case the allowlist has tightened since. + returnTo := validateReturnTo(link.ReturnTo) + + slog.Info("magic_link.callback.success", + "user_id", user.ID, "team_id", team.ID, "request_id", requestID, + ) + + return c.Redirect(appendSessionToken(returnTo, sessionToken), fiber.StatusFound) +} + +// looksLikeEmail performs the cheapest plausible check: must contain a single +// '@' with non-empty local-part and a host that contains a '.'. RFC 5321 has +// edge cases (quoted local-parts, IP-literal hosts) we deliberately reject — +// instanode.dev users never have those addresses. +func looksLikeEmail(s string) bool { + if len(s) < 3 || len(s) > 254 { + return false + } + at := strings.IndexByte(s, '@') + if at <= 0 || at == len(s)-1 { + return false + } + if strings.Count(s, "@") != 1 { + return false + } + host := s[at+1:] + if !strings.Contains(host, ".") { + return false + } + return true +} diff --git a/internal/handlers/nosql.go b/internal/handlers/nosql.go index f0433183..c0df9484 100644 --- a/internal/handlers/nosql.go +++ b/internal/handlers/nosql.go @@ -13,6 +13,7 @@ import ( "time" "github.com/gofiber/fiber/v2" + "github.com/google/uuid" "github.com/redis/go-redis/v9" "instant.dev/internal/config" "instant.dev/internal/crypto" @@ -54,8 +55,9 @@ func (h *NoSQLHandler) provisionNoSQL(ctx context.Context, token, tier string) ( return nil, err } return &nosqlprovider.Credentials{ - URL: creds.URL, - DatabaseName: creds.DatabaseName, + URL: creds.URL, + DatabaseName: creds.DatabaseName, + ProviderResourceID: creds.ProviderResourceID, }, nil } return h.nosqlProvider.Provision(ctx, token, tier) @@ -79,9 +81,14 @@ func (h *NoSQLHandler) NewNoSQL(c *fiber.Ctx) error { _ = c.BodyParser(&body) body.Name = sanitizeName(body.Name) + env, envErr := resolveEnv(c, body.Env) + if envErr != nil { + return envErr + } + // ── Authenticated path ──────────────────────────────────────────────────── if teamIDStr := middleware.GetTeamID(c); teamIDStr != "" { - return h.newNoSQLAuthenticated(c, teamIDStr, fp, country, vendor, requestID, body.Name, body.Dedicated, start) + return h.newNoSQLAuthenticated(c, teamIDStr, fp, country, vendor, requestID, body.Name, body.Dedicated, env, start) } // ── Dedicated requires authentication ───────────────────────────────────── @@ -123,6 +130,7 @@ func (h *NoSQLHandler) NewNoSQL(c *fiber.Ctx) error { "name": existing.Name.String, "connection_url": connectionURL, "tier": existing.Tier, + "env": existing.Env, "limits": nosqlAnonymousLimits(), "note": limitExceededNote(upgradeURL, existing.ExpiresAt.Time), "upgrade": upgradeURL, @@ -141,6 +149,7 @@ func (h *NoSQLHandler) NewNoSQL(c *fiber.Ctx) error { ResourceType: "mongodb", Name: body.Name, Tier: "anonymous", + Env: env, Fingerprint: fp, CloudVendor: vendor, CountryCode: country, @@ -188,6 +197,13 @@ func (h *NoSQLHandler) NewNoSQL(c *fiber.Ctx) error { } } + // Persist provider_resource_id (k8s namespace for dedicated MongoDB pods). + if creds.ProviderResourceID != "" { + if upErr := models.UpdateProviderResourceID(ctx, h.db, resource.ID, creds.ProviderResourceID); upErr != nil { + slog.Error("nosql.new.update_provider_resource_id_failed", "error", upErr, "request_id", requestID) + } + } + jwtToken, jti, jwtErr := h.issueOnboardingJWT(ctx, fp, country, vendor, "mongodb", []string{tokenStr}) if jwtErr != nil { slog.Error("nosql.new.jwt_issue_failed", "error", jwtErr, "request_id", requestID) @@ -226,6 +242,7 @@ func (h *NoSQLHandler) NewNoSQL(c *fiber.Ctx) error { "name": resource.Name.String, "connection_url": creds.URL, "tier": "anonymous", + "env": resource.Env, "limits": nosqlAnonymousLimits(), "note": upgradeNote(upgradeURL), } @@ -237,7 +254,7 @@ func (h *NoSQLHandler) NewNoSQL(c *fiber.Ctx) error { } func (h *NoSQLHandler) newNoSQLAuthenticated( - c *fiber.Ctx, teamIDStr, fp, country, vendor, requestID, name string, dedicated bool, start time.Time, + c *fiber.Ctx, teamIDStr, fp, country, vendor, requestID, name string, dedicated bool, env string, start time.Time, ) error { ctx := c.UserContext() teamUUID, err := parseTeamID(teamIDStr) @@ -260,6 +277,7 @@ func (h *NoSQLHandler) newNoSQLAuthenticated( ResourceType: "mongodb", Name: name, Tier: tier, + Env: env, Fingerprint: fp, CloudVendor: vendor, CountryCode: country, @@ -271,6 +289,18 @@ func (h *NoSQLHandler) newNoSQLAuthenticated( return respondError(c, fiber.StatusServiceUnavailable, "provision_failed", "Failed to provision MongoDB resource") } + // Best-effort audit event; failures must never block the provision. + go func() { + _ = models.InsertAuditEvent(context.Background(), h.db, models.AuditEvent{ + TeamID: teamUUID, + Actor: "agent", + Kind: "provision", + ResourceType: "mongodb", + ResourceID: uuid.NullUUID{UUID: resource.ID, Valid: true}, + Summary: "agent provisioned <strong>mongodb</strong> <code>" + resource.Token.String()[:8] + "</code>", + }) + }() + tokenStr := resource.Token.String() // Provision the real MongoDB database and user. @@ -304,6 +334,13 @@ func (h *NoSQLHandler) newNoSQLAuthenticated( } } + // Persist provider_resource_id (k8s namespace for dedicated MongoDB pods). + if creds.ProviderResourceID != "" { + if upErr := models.UpdateProviderResourceID(ctx, h.db, resource.ID, creds.ProviderResourceID); upErr != nil { + slog.Error("nosql.new.update_provider_resource_id_failed_auth", "error", upErr, "request_id", requestID) + } + } + slog.Info("provision.success", "service", "mongodb", "token", tokenStr, @@ -324,6 +361,7 @@ func (h *NoSQLHandler) newNoSQLAuthenticated( "name": resource.Name.String, "connection_url": creds.URL, "tier": tier, + "env": resource.Env, "limits": fiber.Map{ "storage_mb": nosqlAuthStorageLimitMB, "connections": h.plans.ConnectionsLimit(tier, "mongodb"), diff --git a/internal/handlers/openapi.go b/internal/handlers/openapi.go index dbab6f10..4afeaa9b 100644 --- a/internal/handlers/openapi.go +++ b/internal/handlers/openapi.go @@ -91,6 +91,199 @@ const openAPISpec = `{ } } }, + "/.well-known/oauth-protected-resource": { + "get": { + "summary": "OAuth 2.0 Protected Resource Metadata (RFC 9728)", + "description": "Discovery document used by MCP clients to obtain authorization metadata. Public, no auth required.", + "responses": { + "200": { "description": "Metadata document", "content": { "application/json": { "schema": { "$ref": "#/components/schemas/OAuthProtectedResourceMetadata" } } } } + } + } + }, + "/deploy/new": { + "post": { + "summary": "Deploy a container application", + "description": "Builds a Docker image from the supplied tarball (or pulls an existing image) and rolls it out behind a public HTTPS URL on *.deployment.instanode.dev. Env vars may use the value 'vault://KEY' to reference a secret stored via /api/v1/vault — the plaintext is resolved at deploy time and never persisted in plaintext.", + "security": [{ "bearerAuth": [] }], + "requestBody": { "required": true, "content": { "multipart/form-data": { "schema": { "$ref": "#/components/schemas/DeployRequest" } } } }, + "responses": { + "202": { "description": "Deployment accepted, building", "content": { "application/json": { "schema": { "$ref": "#/components/schemas/DeployResponse" } } } }, + "401": { "description": "Unauthorized" }, + "503": { "description": "Compute backend unavailable or service disabled" } + } + } + }, + "/deploy/{id}": { + "get": { + "summary": "Get deployment status", + "security": [{ "bearerAuth": [] }], + "parameters": [{ "name": "id", "in": "path", "required": true, "schema": { "type": "string" } }], + "responses": { + "200": { "description": "Deployment record", "content": { "application/json": { "schema": { "$ref": "#/components/schemas/DeployResponse" } } } }, + "401": { "description": "Unauthorized" }, + "403": { "description": "Not your deployment" }, + "404": { "description": "Not found" } + } + }, + "delete": { + "summary": "Tear down and delete a deployment", + "security": [{ "bearerAuth": [] }], + "parameters": [{ "name": "id", "in": "path", "required": true, "schema": { "type": "string" } }], + "responses": { + "200": { "description": "Deletion enqueued" }, + "401": { "description": "Unauthorized" }, + "403": { "description": "Not your deployment" } + } + } + }, + "/deploy/{id}/env": { + "patch": { + "summary": "Update env vars (redeploy required to apply)", + "description": "Merges the supplied env vars with the existing ones. Values prefixed with 'vault://' are stored verbatim and resolved at the next redeploy. Plaintext is never logged.", + "security": [{ "bearerAuth": [] }], + "parameters": [{ "name": "id", "in": "path", "required": true, "schema": { "type": "string" } }], + "requestBody": { "required": true, "content": { "application/json": { "schema": { "type": "object", "properties": { "env": { "type": "object", "additionalProperties": { "type": "string" } } } } } } }, + "responses": { + "200": { "description": "Env vars updated", "content": { "application/json": { "schema": { "$ref": "#/components/schemas/DeployResponse" } } } } + } + } + }, + "/deploy/{id}/logs": { + "get": { + "summary": "Stream deployment logs (Server-Sent Events)", + "security": [{ "bearerAuth": [] }], + "parameters": [{ "name": "id", "in": "path", "required": true, "schema": { "type": "string" } }], + "responses": { + "200": { "description": "text/event-stream of log lines, terminated by 'data: [end]'" }, + "409": { "description": "Deployment still building" } + } + } + }, + "/deploy/{id}/redeploy": { + "post": { + "summary": "Redeploy with the latest stored env vars", + "description": "Re-resolves any vault:// references and rolls out a new revision. Use after PATCH /deploy/{id}/env or after rotating a vault secret.", + "security": [{ "bearerAuth": [] }], + "parameters": [{ "name": "id", "in": "path", "required": true, "schema": { "type": "string" } }], + "responses": { + "202": { "description": "Redeploy accepted" } + } + } + }, + "/api/v1/vault/{env}/{key}": { + "put": { + "summary": "Store an encrypted secret", + "description": "Encrypts the supplied value with AES-256-GCM and stores it as a new version. Subsequent PUTs of the same key create v2, v3, ... — old versions remain queryable until DELETE.", + "security": [{ "bearerAuth": [] }], + "parameters": [ + { "name": "env", "in": "path", "required": true, "schema": { "type": "string" }, "description": "Environment scope (production, staging, dev, ...)" }, + { "name": "key", "in": "path", "required": true, "schema": { "type": "string" }, "description": "Secret key (e.g. RAZORPAY_KEY_SECRET)" } + ], + "requestBody": { "required": true, "content": { "application/json": { "schema": { "type": "object", "required": ["value"], "properties": { "value": { "type": "string" } } } } } }, + "responses": { + "201": { "description": "Secret stored", "content": { "application/json": { "schema": { "$ref": "#/components/schemas/VaultPutResponse" } } } }, + "401": { "description": "Unauthorized" } + } + }, + "get": { + "summary": "Read a secret (decrypted)", + "description": "Returns the latest version's plaintext. Pass ?version=N to read a specific historical version. Every read writes a row to vault_audit_log.", + "security": [{ "bearerAuth": [] }], + "parameters": [ + { "name": "env", "in": "path", "required": true, "schema": { "type": "string" } }, + { "name": "key", "in": "path", "required": true, "schema": { "type": "string" } }, + { "name": "version", "in": "query", "required": false, "schema": { "type": "integer" } } + ], + "responses": { + "200": { "description": "Secret returned", "content": { "application/json": { "schema": { "$ref": "#/components/schemas/VaultGetResponse" } } } }, + "404": { "description": "Secret not found for this team / env / key" } + } + }, + "delete": { + "summary": "Hard delete every version of a secret", + "security": [{ "bearerAuth": [] }], + "parameters": [ + { "name": "env", "in": "path", "required": true, "schema": { "type": "string" } }, + { "name": "key", "in": "path", "required": true, "schema": { "type": "string" } } + ], + "responses": { + "204": { "description": "Deleted" }, + "404": { "description": "Not found (idempotent)" } + } + } + }, + "/api/v1/vault/{env}/{key}/rotate": { + "post": { + "summary": "Rotate a secret (new value, version + 1)", + "description": "Convenience for PUT — preserves history but bumps the version visibly. Existing deployments continue to read v(N-1) until they redeploy.", + "security": [{ "bearerAuth": [] }], + "parameters": [ + { "name": "env", "in": "path", "required": true, "schema": { "type": "string" } }, + { "name": "key", "in": "path", "required": true, "schema": { "type": "string" } } + ], + "requestBody": { "required": true, "content": { "application/json": { "schema": { "type": "object", "required": ["value"], "properties": { "value": { "type": "string" } } } } } }, + "responses": { + "200": { "description": "Rotated", "content": { "application/json": { "schema": { "$ref": "#/components/schemas/VaultPutResponse" } } } } + } + } + }, + "/api/v1/vault/{env}": { + "get": { + "summary": "List keys stored in an environment", + "description": "Returns key names only — values are NEVER returned by this endpoint. Use GET /api/v1/vault/{env}/{key} to read a value.", + "security": [{ "bearerAuth": [] }], + "parameters": [{ "name": "env", "in": "path", "required": true, "schema": { "type": "string" } }], + "responses": { + "200": { "description": "List of keys", "content": { "application/json": { "schema": { "type": "object", "properties": { "ok": { "type": "boolean" }, "keys": { "type": "array", "items": { "type": "string" } } } } } } } + } + } + }, + "/api/v1/teams/{team_id}/invitations": { + "post": { + "summary": "Invite a user to the team (admin or owner only)", + "description": "Creates a single-use token tied to the invitee's email. The token is delivered out-of-band (email) and exchanged at POST /api/v1/invitations/{token}/accept.", + "security": [{ "bearerAuth": [] }], + "parameters": [{ "name": "team_id", "in": "path", "required": true, "schema": { "type": "string", "format": "uuid" } }], + "requestBody": { "required": true, "content": { "application/json": { "schema": { "type": "object", "required": ["email", "role"], "properties": { "email": { "type": "string", "format": "email" }, "role": { "type": "string", "enum": ["admin", "developer", "viewer", "member"] } } } } } }, + "responses": { + "201": { "description": "Invitation created", "content": { "application/json": { "schema": { "$ref": "#/components/schemas/InvitationResponse" } } } }, + "403": { "description": "Forbidden — admin role required" } + } + }, + "get": { + "summary": "List pending invitations for a team (admin or owner only)", + "security": [{ "bearerAuth": [] }], + "parameters": [{ "name": "team_id", "in": "path", "required": true, "schema": { "type": "string", "format": "uuid" } }], + "responses": { + "200": { "description": "Invitations", "content": { "application/json": { "schema": { "type": "object", "properties": { "items": { "type": "array", "items": { "$ref": "#/components/schemas/InvitationResponse" } } } } } } } + } + } + }, + "/api/v1/teams/{team_id}/invitations/{id}": { + "delete": { + "summary": "Revoke a pending invitation", + "security": [{ "bearerAuth": [] }], + "parameters": [ + { "name": "team_id", "in": "path", "required": true, "schema": { "type": "string", "format": "uuid" } }, + { "name": "id", "in": "path", "required": true, "schema": { "type": "string", "format": "uuid" } } + ], + "responses": { + "204": { "description": "Revoked" } + } + } + }, + "/api/v1/invitations/{token}/accept": { + "post": { + "summary": "Accept an invitation by token (no auth required — token IS the auth)", + "description": "Public endpoint. The token is single-use and ties the accepting user's session to the invited team and role.", + "parameters": [{ "name": "token", "in": "path", "required": true, "schema": { "type": "string" } }], + "responses": { + "200": { "description": "Accepted", "content": { "application/json": { "schema": { "type": "object", "properties": { "ok": { "type": "boolean" }, "team_id": { "type": "string", "format": "uuid" }, "role": { "type": "string" } } } } } }, + "404": { "description": "Token not found" }, + "410": { "description": "Token already used or expired" } + } + } + }, "/claim": { "post": { "summary": "Claim anonymous resources to a permanent account", @@ -181,7 +374,10 @@ const openAPISpec = `{ }, "ProvisionRequest": { "type": "object", - "properties": { "name": { "type": "string", "description": "Optional human-readable label (max 120 chars)" } } + "properties": { + "name": { "type": "string", "description": "Optional human-readable label (max 120 chars)" }, + "env": { "type": "string", "description": "Optional environment scope (production / staging / dev / ...). Anonymous tier is always 'production'.", "default": "production" } + } }, "DBProvisionResponse": { "type": "object", @@ -283,12 +479,83 @@ const openAPISpec = `{ "token": { "type": "string", "format": "uuid" }, "resource_type": { "type": "string", "enum": ["postgres", "redis", "mongodb", "nats", "webhook", "storage"] }, "name": { "type": "string" }, + "env": { "type": "string", "description": "Environment scope (production / staging / dev / ...)" }, "tier": { "type": "string" }, "status": { "type": "string" }, "storage_bytes": { "type": "integer" }, "expires_at": { "type": "string", "format": "date-time", "nullable": true }, "created_at": { "type": "string", "format": "date-time" } } + }, + "OAuthProtectedResourceMetadata": { + "type": "object", + "properties": { + "resource": { "type": "string", "description": "Canonical URL of this protected resource" }, + "authorization_servers": { "type": "array", "items": { "type": "string" } }, + "bearer_methods_supported": { "type": "array", "items": { "type": "string", "enum": ["header"] } }, + "resource_documentation": { "type": "string" } + } + }, + "VaultPutResponse": { + "type": "object", + "properties": { + "ok": { "type": "boolean" }, + "key": { "type": "string" }, + "env": { "type": "string" }, + "version": { "type": "integer" } + } + }, + "VaultGetResponse": { + "type": "object", + "properties": { + "ok": { "type": "boolean" }, + "key": { "type": "string" }, + "env": { "type": "string" }, + "version": { "type": "integer" }, + "value": { "type": "string", "description": "Decrypted plaintext" } + } + }, + "DeployRequest": { + "type": "object", + "properties": { + "tarball": { "type": "string", "format": "binary", "description": "gzipped tar archive containing the Dockerfile + source (max 50 MB)" }, + "name": { "type": "string", "description": "Optional human-readable label" }, + "port": { "type": "integer", "description": "Container port (default 8080)" }, + "env": { "type": "string", "description": "Environment scope (production / staging / dev / ...)" } + }, + "required": ["tarball"] + }, + "DeployResponse": { + "type": "object", + "properties": { + "ok": { "type": "boolean" }, + "item": { + "type": "object", + "properties": { + "id": { "type": "string", "format": "uuid" }, + "app_id": { "type": "string", "description": "8-char public identifier used in the URL" }, + "url": { "type": "string", "description": "Live HTTPS URL (set once status=healthy)" }, + "status": { "type": "string", "enum": ["building", "healthy", "failed", "stopped"] }, + "tier": { "type": "string" }, + "environment": { "type": "string", "description": "Env scope (production/staging/dev). Note: 'env' on this object is the env_vars map, not the scope." }, + "env": { "type": "object", "additionalProperties": { "type": "string" }, "description": "Env vars map — vault://KEY references resolve at deploy time" }, + "port": { "type": "integer" }, + "team_id": { "type": "string", "format": "uuid" } + } + }, + "note": { "type": "string" } + } + }, + "InvitationResponse": { + "type": "object", + "properties": { + "ok": { "type": "boolean" }, + "id": { "type": "string", "format": "uuid" }, + "team_id": { "type": "string", "format": "uuid" }, + "email": { "type": "string", "format": "email" }, + "role": { "type": "string", "enum": ["admin", "developer", "viewer", "member"] }, + "expires_at": { "type": "string", "format": "date-time" } + } } } } diff --git a/internal/handlers/provision_helper.go b/internal/handlers/provision_helper.go index dded8d62..0bb77e18 100644 --- a/internal/handlers/provision_helper.go +++ b/internal/handlers/provision_helper.go @@ -8,6 +8,7 @@ package handlers // 2. Onboarding JWT issuance (issueOnboardingJWT) // 3. Active-resource lookup (models.GetActiveResourceByFingerprint) // 4. Onboarding event creation (models.CreateOnboardingEvent) +// 5. Environment selection (resolveEnv — see provisionRequestBody.Env) // // provisionHelper embeds these shared behaviours so each handler // can embed it instead of duplicating the logic. @@ -19,6 +20,7 @@ import ( "log/slog" "time" + "github.com/gofiber/fiber/v2" "github.com/google/uuid" "github.com/redis/go-redis/v9" "go.opentelemetry.io/otel" @@ -218,6 +220,11 @@ type provisionRequestBody struct { // own namespace, own PVC). Requires an authenticated team-tier token. // Anonymous callers receive a 402 with an upgrade URL. Dedicated bool `json:"dedicated"` + + // Env scopes the resource to a named environment (dev/staging/production/...). + // Empty defaults to "production". Validated against ^[a-z0-9-]{1,32}$. + // Body field is overridden by the ?env= query string when both are set. + Env string `json:"env"` } func sanitizeName(name string) string { @@ -226,3 +233,22 @@ func sanitizeName(name string) string { } return name } + +// resolveEnv extracts the requested environment from the request, preferring +// the ?env= query string over the JSON/form body field. Returns the normalised +// env on success, or an empty string and a 400 response when validation fails. +// +// Empty input is treated as "production" — this preserves backwards compatibility +// for every caller that pre-dates the env feature. +func resolveEnv(c *fiber.Ctx, bodyEnv string) (string, error) { + raw := c.Query("env") + if raw == "" { + raw = bodyEnv + } + env, ok := models.NormalizeEnv(raw) + if !ok { + return "", respondError(c, fiber.StatusBadRequest, "invalid_env", + "env must match ^[a-z0-9-]{1,32}$ (lowercase letters, digits, dashes; max 32 chars)") + } + return env, nil +} diff --git a/internal/handlers/queue.go b/internal/handlers/queue.go index c19ea767..4f77e2ee 100644 --- a/internal/handlers/queue.go +++ b/internal/handlers/queue.go @@ -28,6 +28,7 @@ import ( "time" "github.com/gofiber/fiber/v2" + "github.com/google/uuid" "github.com/redis/go-redis/v9" "instant.dev/internal/config" "instant.dev/internal/crypto" @@ -60,10 +61,11 @@ func NewQueueHandler(db *sql.DB, rdb *redis.Client, cfg *config.Config, provClie } // provisionQueue provisions NATS credentials. -// Growth, pro, and team tiers use the gRPC provisioner (isolated k8s NATS pod). -// All other tiers use the local provider (shared NATS cluster). +// When the gRPC provisioner is configured, every tier uses it — the provisioner +// chooses local vs k8s-dedicated backend based on QUEUE_PROVISION_BACKEND. +// Falls back to the local provider only when no provisioner client is wired. func (h *QueueHandler) provisionQueue(ctx context.Context, token, tier string) (*queueprovider.Credentials, error) { - if (tier == "pro" || tier == "team" || tier == "growth") && h.provClient != nil { + if h.provClient != nil { creds, err := h.provClient.ProvisionQueue(ctx, token, tier) if err != nil { return nil, err @@ -95,9 +97,14 @@ func (h *QueueHandler) NewQueue(c *fiber.Ctx) error { _ = c.BodyParser(&body) body.Name = sanitizeName(body.Name) + env, envErr := resolveEnv(c, body.Env) + if envErr != nil { + return envErr + } + // ── Authenticated path ──────────────────────────────────────────────────── if teamIDStr := middleware.GetTeamID(c); teamIDStr != "" { - return h.newQueueAuthenticated(c, teamIDStr, fp, country, vendor, requestID, body.Name, body.Dedicated, start) + return h.newQueueAuthenticated(c, teamIDStr, fp, country, vendor, requestID, body.Name, body.Dedicated, env, start) } // ── Dedicated requires authentication ───────────────────────────────────── @@ -139,6 +146,7 @@ func (h *QueueHandler) NewQueue(c *fiber.Ctx) error { "name": existing.Name.String, "connection_url": connectionURL, "tier": existing.Tier, + "env": existing.Env, "limits": queueAnonymousLimits(), "note": limitExceededNote(upgradeURL, existing.ExpiresAt.Time), "upgrade": upgradeURL, @@ -157,6 +165,7 @@ func (h *QueueHandler) NewQueue(c *fiber.Ctx) error { ResourceType: "queue", Name: body.Name, Tier: "anonymous", + Env: env, Fingerprint: fp, CloudVendor: vendor, CountryCode: country, @@ -240,13 +249,14 @@ func (h *QueueHandler) NewQueue(c *fiber.Ctx) error { "connection_url": creds.URL, "subject_prefix": creds.SubjectPrefix, "tier": "anonymous", + "env": resource.Env, "limits": queueAnonymousLimits(), "note": upgradeNote(upgradeURL), }) } func (h *QueueHandler) newQueueAuthenticated( - c *fiber.Ctx, teamIDStr, fp, country, vendor, requestID, name string, dedicated bool, start time.Time, + c *fiber.Ctx, teamIDStr, fp, country, vendor, requestID, name string, dedicated bool, env string, start time.Time, ) error { ctx := c.UserContext() teamUUID, err := parseTeamID(teamIDStr) @@ -269,6 +279,7 @@ func (h *QueueHandler) newQueueAuthenticated( ResourceType: "queue", Name: name, Tier: tier, + Env: env, Fingerprint: fp, CloudVendor: vendor, CountryCode: country, @@ -280,6 +291,18 @@ func (h *QueueHandler) newQueueAuthenticated( return respondError(c, fiber.StatusServiceUnavailable, "provision_failed", "Failed to provision NATS resource") } + // Best-effort audit event; failures must never block the provision. + go func() { + _ = models.InsertAuditEvent(context.Background(), h.db, models.AuditEvent{ + TeamID: teamUUID, + Actor: "agent", + Kind: "provision", + ResourceType: "queue", + ResourceID: uuid.NullUUID{UUID: resource.ID, Valid: true}, + Summary: "agent provisioned <strong>queue</strong> <code>" + resource.Token.String()[:8] + "</code>", + }) + }() + tokenStr := resource.Token.String() // Provision NATS credentials. @@ -339,6 +362,7 @@ func (h *QueueHandler) newQueueAuthenticated( "connection_url": creds.URL, "subject_prefix": creds.SubjectPrefix, "tier": tier, + "env": resource.Env, "dedicated": dedicated, "limits": fiber.Map{ "storage_mb": h.plans.StorageLimitMB(tier, "queue"), diff --git a/internal/handlers/resource.go b/internal/handlers/resource.go index 52c25b9d..8a85fb56 100644 --- a/internal/handlers/resource.go +++ b/internal/handlers/resource.go @@ -212,6 +212,69 @@ func (h *ResourceHandler) Delete(c *fiber.Ctx) error { }) } +// GetCredentials handles GET /api/v1/resources/:id/credentials. +// Returns the plaintext connection URL for the team's own resource — same +// auth boundary as RotateCredentials, but does NOT change the password. +// Used by `instant up` to re-emit URLs into .env on subsequent runs. +func (h *ResourceHandler) GetCredentials(c *fiber.Ctx) error { + requestID := middleware.GetRequestID(c) + + teamID, err := parseTeamID(middleware.GetTeamID(c)) + if err != nil { + return respondError(c, fiber.StatusUnauthorized, "unauthorized", "Valid session token required") + } + + tokenStr := c.Params("id") + token, parseErr := uuid.Parse(tokenStr) + if parseErr != nil { + return respondError(c, fiber.StatusBadRequest, "invalid_id", "Resource ID must be a valid UUID") + } + + resource, err := models.GetResourceByToken(c.Context(), h.db, token) + if err != nil { + var notFound *models.ErrResourceNotFound + if errors.As(err, &notFound) { + return respondError(c, fiber.StatusNotFound, "not_found", "Resource not found") + } + slog.Error("resource.credentials.lookup_failed", + "error", err, "token", tokenStr, "request_id", requestID) + return respondError(c, fiber.StatusServiceUnavailable, "fetch_failed", "Failed to fetch resource") + } + + if !resource.TeamID.Valid || resource.TeamID.UUID != teamID { + // Mirror "404 not 403" pattern used elsewhere — never confirm the + // existence of resources owned by other teams. + return respondError(c, fiber.StatusNotFound, "not_found", "Resource not found") + } + + if !resource.ConnectionURL.Valid || resource.ConnectionURL.String == "" { + return respondError(c, fiber.StatusBadRequest, "no_connection_url", + "This resource does not have a connection URL") + } + + aesKey, err := crypto.ParseAESKey(h.cfg.AESKey) + if err != nil { + slog.Error("resource.credentials.aes_key_invalid", + "error", err, "request_id", requestID) + return respondError(c, fiber.StatusInternalServerError, "internal_error", "Encryption configuration error") + } + plain, err := crypto.Decrypt(aesKey, resource.ConnectionURL.String) + if err != nil { + slog.Error("resource.credentials.decrypt_failed", + "error", err, "resource_id", resource.ID, "request_id", requestID) + return respondError(c, fiber.StatusInternalServerError, "internal_error", "Failed to decrypt connection URL") + } + + return c.JSON(fiber.Map{ + "ok": true, + "id": resource.ID, + "token": resource.Token, + "resource_type": resource.ResourceType, + "env": resource.Env, + "connection_url": plain, + }) +} + // RotateCredentials handles POST /api/v1/resources/:id/rotate-credentials. // Generates a new password, re-encrypts the connection URL, persists it, and // returns the new plaintext URL — this is the only endpoint that exposes connection_url. @@ -366,6 +429,7 @@ func resourceToMap(r *models.Resource) fiber.Map { "id": r.ID, "token": r.Token, "resource_type": r.ResourceType, + "env": r.Env, "tier": r.Tier, "status": r.Status, "created_at": r.CreatedAt, diff --git a/internal/handlers/stack.go b/internal/handlers/stack.go index 1bcba5b8..27bb9f9c 100644 --- a/internal/handlers/stack.go +++ b/internal/handlers/stack.go @@ -29,6 +29,7 @@ import ( "fmt" "io" "log/slog" + "net/url" "strings" "time" @@ -169,6 +170,64 @@ func stackOwnerCheck(c *fiber.Ctx, stack *models.Stack, team *models.Team) error return nil } +// rewriteToInternalURL replaces the host:port of a customer-facing connection +// URL with the cluster-internal FQDN of the dedicated pod, so stack workloads +// can reach their `needs:` resources without going through the LoadBalancer. +// +// Why this is needed: customer URLs use K8S_EXTERNAL_HOST (e.g. pg.instanode.dev) +// + a per-resource port. From outside the cluster they work. From INSIDE the +// cluster, the LoadBalancer doesn't hairpin reliably on DOKS, so a stack pod +// trying to reach pg.instanode.dev:5432 just times out. +// +// Resource → internal FQDN mapping: +// +// postgres → instant-pg-proxy.instant.svc.cluster.local:5432 +// (the proxy routes by db name in the startup packet) +// redis → redis.<provider_resource_id>.svc.cluster.local:6379 +// mongodb → mongo.<provider_resource_id>.svc.cluster.local:27017 +// queue → nats.<provider_resource_id>.svc.cluster.local:4222 +// +// If providerResourceID is empty (legacy / non-dedicated resource), the URL is +// returned unchanged. Callers should still log a warning in that case. +func rewriteToInternalURL(publicURL, resourceType, providerResourceID string) string { + if publicURL == "" { + return publicURL + } + parsed, err := url.Parse(publicURL) + if err != nil || parsed.Host == "" { + return publicURL + } + + var newHost string + switch resourceType { + case "postgres": + // Always route via the cluster-internal pg-proxy. The proxy reads the + // database name from the Postgres startup packet and forwards to the + // dedicated pod — works for every customer DB without per-resource state. + newHost = "instant-pg-proxy.instant.svc.cluster.local:5432" + case "redis": + if providerResourceID == "" { + return publicURL + } + newHost = "redis." + providerResourceID + ".svc.cluster.local:6379" + case "mongodb": + if providerResourceID == "" { + return publicURL + } + newHost = "mongo." + providerResourceID + ".svc.cluster.local:27017" + case "queue": + if providerResourceID == "" { + return publicURL + } + newHost = "nats." + providerResourceID + ".svc.cluster.local:4222" + default: + return publicURL + } + + parsed.Host = newHost + return parsed.String() +} + // resourceEnvKey returns the canonical env var name for a resource type. // index > 0 appends a numeric suffix (DATABASE_URL_2, etc.). func resourceEnvKey(resourceType string, index int) string { @@ -422,6 +481,22 @@ func (h *StackHandler) New(c *fiber.Ctx) error { "token", res.Token, "error", decErr) plainURL = res.ConnectionURL.String } + // Rewrite the customer-facing URL (LB external host + NodePort or proxy + // port) to the in-cluster FQDN. Stack pods must connect via cluster DNS + // because DOKS LoadBalancers don't reliably hairpin and the public IP + // route adds latency + crosses the namespace egress firewall. + // + // Customer's dashboard / `connection_url` field still shows the public URL + // — only the env injected into in-cluster stack pods is rewritten. + // Fallback: redis/mongo/queue handlers don't all persist provider_resource_id + // today (cache.go and nosql.go are missing the UpdateProviderResourceID call). + // Derive the namespace from the token using the same convention the k8s + // backends use ("instant-customer-<token>") so the rewrite still works. + prid := res.ProviderResourceID.String + if prid == "" || prid == "local:0" { + prid = "instant-customer-" + res.Token.String() + } + plainURL = rewriteToInternalURL(plainURL, res.ResourceType, prid) key := resourceEnvKey(res.ResourceType, idx) env[key] = plainURL } @@ -498,6 +573,12 @@ func (h *StackHandler) New(c *fiber.Ctx) error { } // Step 7: Build StackDeployOptions. + // + // Per-service env vars may include "vault://KEY" references. We resolve + // them here against the team's vault for the production env (stack + // deploys do not yet expose multi-env scoping; this matches the + // per-deployment behaviour). Anonymous stacks cannot use vault refs + // because there is no team to look up. services := make([]compute.StackServiceDef, 0, len(m.Services)) for svcName, svc := range m.Services { // Merge: needs env first (low priority), then service-defined env (high priority). @@ -508,6 +589,28 @@ func (h *StackHandler) New(c *fiber.Ctx) error { for k, v := range svc.Env { envVars[k] = v } + + // Resolve vault:// refs (authenticated only). + if !anon { + resolved, vaultErr := ResolveVaultRefs(c.Context(), h.db, h.cfg.AESKey, team.ID, "production", envVars) + if vaultErr != nil { + slog.Error("stack.new.vault_resolve_failed", + "error", vaultErr, "slug", slug, "service", svcName, + "team_id", team.ID, "request_id", middleware.GetRequestID(c)) + return respondError(c, fiber.StatusBadRequest, "vault_ref_failed", + "Failed to resolve vault reference for "+svcName+": "+vaultErr.Error()) + } + envVars = resolved + } else { + // Reject vault refs from anonymous callers — fail loud, not silent. + for k, v := range envVars { + if strings.HasPrefix(v, vaultRefPrefix) { + return respondError(c, fiber.StatusForbidden, "vault_requires_auth", + "vault:// references require authentication: "+svcName+"."+k) + } + } + } + services = append(services, compute.StackServiceDef{ Name: svcName, Tarball: tarballs[svcName], @@ -846,15 +949,27 @@ func (h *StackHandler) Redeploy(c *fiber.Ctx) error { tarballs[name] = data } - // Build service defs. + // Build service defs. Resolve "vault://KEY" references in env vars + // before passing to the compute provider — same semantics as the + // initial /stacks/new path. Redeploy is always authenticated, so + // no anonymous-rejection branch is needed here. services := make([]compute.StackServiceDef, 0, len(m.Services)) for svcName, svc := range m.Services { + envVars := svc.Env + resolved, vaultErr := ResolveVaultRefs(c.Context(), h.db, h.cfg.AESKey, team.ID, "production", envVars) + if vaultErr != nil { + slog.Error("stack.redeploy.vault_resolve_failed", + "error", vaultErr, "slug", slug, "service", svcName, + "team_id", team.ID, "request_id", middleware.GetRequestID(c)) + return respondError(c, fiber.StatusBadRequest, "vault_ref_failed", + "Failed to resolve vault reference for "+svcName+": "+vaultErr.Error()) + } services = append(services, compute.StackServiceDef{ Name: svcName, Tarball: tarballs[svcName], Port: svc.Port, Expose: svc.Expose, - EnvVars: svc.Env, + EnvVars: resolved, }) } diff --git a/internal/handlers/storage.go b/internal/handlers/storage.go index d1582b62..2b8e53b6 100644 --- a/internal/handlers/storage.go +++ b/internal/handlers/storage.go @@ -32,6 +32,7 @@ import ( "time" "github.com/gofiber/fiber/v2" + "github.com/google/uuid" "github.com/redis/go-redis/v9" "instant.dev/internal/config" "instant.dev/internal/crypto" @@ -60,7 +61,7 @@ func NewStorageHandler(db *sql.DB, rdb *redis.Client, cfg *config.Config, storag if storageProvider != nil { h.storageProvider = storageProvider } else if cfg.MinioEndpoint != "" { - sp, err := storageprovider.New(cfg.MinioEndpoint, cfg.MinioRootUser, cfg.MinioRootPassword, cfg.MinioBucketName) + sp, err := storageprovider.New(cfg.MinioEndpoint, cfg.MinioPublicEndpoint, cfg.MinioRootUser, cfg.MinioRootPassword, cfg.MinioBucketName) if err != nil { slog.Warn("storage: MinIO provider init failed — /storage/new will return 503", "error", err) } else { @@ -93,9 +94,14 @@ func (h *StorageHandler) NewStorage(c *fiber.Ctx) error { _ = c.BodyParser(&body) body.Name = sanitizeName(body.Name) + env, envErr := resolveEnv(c, body.Env) + if envErr != nil { + return envErr + } + // ── Authenticated path ──────────────────────────────────────────────────── if teamIDStr := middleware.GetTeamID(c); teamIDStr != "" { - return h.newStorageAuthenticated(c, teamIDStr, fp, country, vendor, requestID, body.Name, start) + return h.newStorageAuthenticated(c, teamIDStr, fp, country, vendor, requestID, body.Name, env, start) } // ── Anonymous path ───────────────────────────────────────────────────────── @@ -132,6 +138,7 @@ func (h *StorageHandler) NewStorage(c *fiber.Ctx) error { "name": existing.Name.String, "connection_url": connectionURL, "tier": existing.Tier, + "env": existing.Env, "limits": h.storageAnonymousLimits(), "note": limitExceededNote(upgradeURL, existing.ExpiresAt.Time), "upgrade": upgradeURL, @@ -144,6 +151,7 @@ func (h *StorageHandler) NewStorage(c *fiber.Ctx) error { ResourceType: "storage", Name: body.Name, Tier: "anonymous", + Env: env, Fingerprint: fp, CloudVendor: vendor, CountryCode: country, @@ -230,6 +238,7 @@ func (h *StorageHandler) NewStorage(c *fiber.Ctx) error { "secret_access_key": creds.SecretAccessKey, "prefix": creds.Prefix, "tier": "anonymous", + "env": resource.Env, "limits": h.storageAnonymousLimits(), "note": upgradeNote(upgradeURL), "upgrade": upgradeURL, @@ -238,7 +247,7 @@ func (h *StorageHandler) NewStorage(c *fiber.Ctx) error { } func (h *StorageHandler) newStorageAuthenticated( - c *fiber.Ctx, teamIDStr, fp, country, vendor, requestID, name string, start time.Time, + c *fiber.Ctx, teamIDStr, fp, country, vendor, requestID, name string, env string, start time.Time, ) error { ctx := c.UserContext() teamUUID, err := parseTeamID(teamIDStr) @@ -272,6 +281,7 @@ func (h *StorageHandler) newStorageAuthenticated( ResourceType: "storage", Name: name, Tier: team.PlanTier, + Env: env, Fingerprint: fp, CloudVendor: vendor, CountryCode: country, @@ -283,6 +293,18 @@ func (h *StorageHandler) newStorageAuthenticated( return respondError(c, fiber.StatusServiceUnavailable, "provision_failed", "Failed to provision storage resource") } + // Best-effort audit event; failures must never block the provision. + go func() { + _ = models.InsertAuditEvent(context.Background(), h.db, models.AuditEvent{ + TeamID: teamUUID, + Actor: "agent", + Kind: "provision", + ResourceType: "storage", + ResourceID: uuid.NullUUID{UUID: resource.ID, Valid: true}, + Summary: "agent provisioned <strong>storage</strong> <code>" + resource.Token.String()[:8] + "</code>", + }) + }() + tokenStr := resource.Token.String() // Provision R2 credentials. @@ -337,6 +359,7 @@ func (h *StorageHandler) newStorageAuthenticated( "secret_access_key": creds.SecretAccessKey, "prefix": creds.Prefix, "tier": team.PlanTier, + "env": resource.Env, "limits": fiber.Map{ "storage_mb": h.plans.StorageLimitMB(team.PlanTier, "storage"), }, diff --git a/internal/handlers/team_members.go b/internal/handlers/team_members.go index 45b6b7d2..909aec87 100644 --- a/internal/handlers/team_members.go +++ b/internal/handlers/team_members.go @@ -54,8 +54,9 @@ func (h *TeamMembersHandler) ListMembers(c *fiber.Ctx) error { if err != nil { return respondError(c, fiber.StatusUnauthorized, "unauthorized", "Valid session required") } + // Any team member may list — owner, admin, developer, viewer, or legacy "member". role, err := models.GetUserRole(c.Context(), h.db, teamID, userID) - if err != nil || (role != "owner" && role != "member") { + if err != nil || role == "" { return respondError(c, fiber.StatusForbidden, "forbidden", "Not a member of this team") } members, err := models.ListTeamMembers(c.Context(), h.db, teamID) @@ -70,10 +71,10 @@ func (h *TeamMembersHandler) ListMembers(c *fiber.Ctx) error { items := make([]fiber.Map, 0, len(members)) for _, m := range members { items = append(items, fiber.Map{ - "id": m.ID.String(), - "email": m.Email, - "role": m.Role, - "created_at": m.CreatedAt.UTC().Format(time.RFC3339), + "user_id": m.ID.String(), + "email": m.Email, + "role": m.Role, + "joined_at": m.CreatedAt.UTC().Format(time.RFC3339), }) } return c.JSON(fiber.Map{"ok": true, "members": items, "member_limit": limit}) @@ -84,6 +85,16 @@ type inviteBody struct { Role string `json:"role"` } +// allowedSimpleInviteRoles bounds the set of roles accepted by the simpler +// /api/v1/team/members/invite endpoint. "member" is retained as a legacy +// alias of the owner/member flow; admin/developer/viewer use the RBAC flow. +var allowedSimpleInviteRoles = map[string]struct{}{ + "admin": {}, + "developer": {}, + "viewer": {}, + "member": {}, +} + // InviteMember handles POST /api/v1/team/members/invite func (h *TeamMembersHandler) InviteMember(c *fiber.Ctx) error { teamID, err := uuid.Parse(middleware.GetTeamID(c)) @@ -94,8 +105,14 @@ func (h *TeamMembersHandler) InviteMember(c *fiber.Ctx) error { if err != nil { return respondError(c, fiber.StatusUnauthorized, "unauthorized", "Valid session required") } - if !h.requireOwner(c, teamID, userID) { - return respondError(c, fiber.StatusForbidden, "forbidden", "Owner only") + // Owner OR admin may invite (legacy "owner" was sole inviter; RBAC adds admin). + actorRole, err := models.GetUserRole(c.Context(), h.db, teamID, userID) + if err != nil { + slog.Error("team_members.role_lookup", "error", err) + return respondError(c, fiber.StatusInternalServerError, "internal_error", "Request failed") + } + if actorRole != "owner" && actorRole != "admin" { + return respondError(c, fiber.StatusForbidden, "forbidden", "Owner or admin only") } var body inviteBody if err := c.BodyParser(&body); err != nil { @@ -109,23 +126,64 @@ func (h *TeamMembersHandler) InviteMember(c *fiber.Ctx) error { if role == "" { role = "member" } + if _, ok := allowedSimpleInviteRoles[role]; !ok { + return respondError(c, fiber.StatusBadRequest, "invalid_role", + "role must be one of: admin, developer, viewer, member") + } tier, err := h.teamPlanTier(c, teamID) if err != nil { return respondError(c, fiber.StatusInternalServerError, "tier_failed", "Failed to read team plan") } limit := h.plans.TeamMemberLimit(tier) - inv, err := models.InviteMember(c.Context(), h.db, teamID, email, role, userID, limit) - if err != nil { - return teamMembersModelError(c, err) - } + teamRow, _ := models.GetTeamByID(c.Context(), h.db, teamID) teamName := "" if teamRow != nil && teamRow.Name.Valid { teamName = teamRow.Name.String } + base := strings.TrimRight(h.cfg.DashboardBaseURL, "/") + + // Legacy "member" role uses the owner/member flow with seat-limit enforcement. + // admin/developer/viewer use the RBAC token flow. + if role == "member" { + // Owner/member flow currently requires owner; admins fall back to the + // RBAC flow with role="developer" since legacy seats can't be granted + // by non-owners. + if actorRole != "owner" { + return respondError(c, fiber.StatusForbidden, "forbidden", + "Only the team owner can invite legacy members; use role=developer instead") + } + inv, err := models.InviteMember(c.Context(), h.db, teamID, email, role, userID, limit) + if err != nil { + return teamMembersModelError(c, err) + } + if h.mail != nil { + acceptURL := base + "/settings?section=team&invite=" + inv.ID.String() + if mailErr := h.mail.SendTeamInvite(c.Context(), inv.Email, teamName, acceptURL); mailErr != nil { + slog.Warn("team_members.invite_email_failed", "error", mailErr) + } + } + return c.Status(fiber.StatusCreated).JSON(fiber.Map{ + "ok": true, + "invitation": fiber.Map{ + "id": inv.ID.String(), + "email": inv.Email, + "role": inv.Role, + "status": inv.Status, + "invited_by": inv.InvitedBy.String(), + "created_at": inv.CreatedAt.UTC().Format(time.RFC3339), + "expires_at": inv.ExpiresAt.UTC().Format(time.RFC3339), + }, + }) + } + + // RBAC flow: admin / developer / viewer — token-based single-use invite. + inv, err := models.CreateRBACInvitation(c.Context(), h.db, teamID, email, role, userID) + if err != nil { + return teamMembersModelError(c, err) + } if h.mail != nil { - base := strings.TrimRight(h.cfg.DashboardBaseURL, "/") - acceptURL := base + "/settings?section=team&invite=" + inv.ID.String() + acceptURL := base + "/invitations/" + inv.Token + "/accept" if mailErr := h.mail.SendTeamInvite(c.Context(), inv.Email, teamName, acceptURL); mailErr != nil { slog.Warn("team_members.invite_email_failed", "error", mailErr) } @@ -136,7 +194,8 @@ func (h *TeamMembersHandler) InviteMember(c *fiber.Ctx) error { "id": inv.ID.String(), "email": inv.Email, "role": inv.Role, - "status": inv.Status, + "token": inv.Token, + "status": inv.Status(), "invited_by": inv.InvitedBy.String(), "created_at": inv.CreatedAt.UTC().Format(time.RFC3339), "expires_at": inv.ExpiresAt.UTC().Format(time.RFC3339), diff --git a/internal/handlers/teams.go b/internal/handlers/teams.go new file mode 100644 index 00000000..7c7df5b1 --- /dev/null +++ b/internal/handlers/teams.go @@ -0,0 +1,253 @@ +package handlers + +import ( + "database/sql" + "errors" + "log/slog" + "strings" + "time" + + "github.com/gofiber/fiber/v2" + "github.com/google/uuid" + "instant.dev/internal/config" + "instant.dev/internal/email" + "instant.dev/internal/middleware" + "instant.dev/internal/models" +) + +// TeamsHandler serves the RBAC-aware team endpoints: +// +// POST /api/v1/teams/:team_id/invitations +// GET /api/v1/teams/:team_id/invitations +// DELETE /api/v1/teams/:team_id/invitations/:id +// POST /api/v1/invitations/:token/accept (no auth — token IS the auth) +// +// Distinct from TeamMembersHandler (legacy /api/v1/team/members/* routes that +// use the simpler owner/member invite flow). The two coexist intentionally: +// this handler implements the new admin/developer/viewer RBAC tiers + token +// acceptance. +type TeamsHandler struct { + db *sql.DB + cfg *config.Config + mail *email.Client +} + +// NewTeamsHandler constructs a TeamsHandler. +func NewTeamsHandler(db *sql.DB, cfg *config.Config, mail *email.Client) *TeamsHandler { + return &TeamsHandler{db: db, cfg: cfg, mail: mail} +} + +// inviteRequest is the JSON body for POST /api/v1/teams/:team_id/invitations. +type inviteRequest struct { + Email string `json:"email"` + Role string `json:"role"` +} + +// CreateInvitation handles POST /api/v1/teams/:team_id/invitations. +// Owner / admin only (callers gate via RequireRole("admin")). +// +// Body: { "email": "user@example.com", "role": "developer" } +// 201: { "ok": true, "invitation": { id, email, role, token, expires_at, ... } } +func (h *TeamsHandler) CreateInvitation(c *fiber.Ctx) error { + teamID, err := h.requireTeamMatch(c) + if err != nil { + return err + } + actorID, err := uuid.Parse(middleware.GetUserID(c)) + if err != nil { + return respondError(c, fiber.StatusUnauthorized, "unauthorized", "Valid session required") + } + + var body inviteRequest + if err := c.BodyParser(&body); err != nil { + return respondError(c, fiber.StatusBadRequest, "invalid_body", "Invalid JSON") + } + emailAddr := strings.TrimSpace(strings.ToLower(body.Email)) + if emailAddr == "" { + return respondError(c, fiber.StatusBadRequest, "missing_email", "email is required") + } + role := strings.TrimSpace(strings.ToLower(body.Role)) + if !models.IsValidInviteRole(role) { + return respondError(c, fiber.StatusBadRequest, "invalid_role", + "role must be one of: admin, developer, viewer") + } + + inv, err := models.CreateRBACInvitation(c.Context(), h.db, teamID, emailAddr, role, actorID) + if err != nil { + return teamsModelError(c, err) + } + + // Best-effort email — never fail the request if delivery fails. + if h.mail != nil { + base := strings.TrimRight(h.cfg.DashboardBaseURL, "/") + acceptURL := base + "/invitations/" + inv.Token + "/accept" + teamName := "" + if t, terr := models.GetTeamByID(c.Context(), h.db, teamID); terr == nil && t.Name.Valid { + teamName = t.Name.String + } + if mailErr := h.mail.SendTeamInvite(c.Context(), inv.Email, teamName, acceptURL); mailErr != nil { + slog.Warn("teams.invite_email_failed", "error", mailErr, "invitation_id", inv.ID) + } + } else { + slog.Info("teams.invite_email_stub", "to", inv.Email, "team_id", teamID, "token_present", true) + } + + return c.Status(fiber.StatusCreated).JSON(fiber.Map{ + "ok": true, + "invitation": serializeInvitation(inv), + }) +} + +// ListInvitations handles GET /api/v1/teams/:team_id/invitations. +// Owner / admin only. Returns pending (not accepted) invites. +func (h *TeamsHandler) ListInvitations(c *fiber.Ctx) error { + teamID, err := h.requireTeamMatch(c) + if err != nil { + return err + } + invs, err := models.ListRBACInvitations(c.Context(), h.db, teamID) + if err != nil { + return respondError(c, fiber.StatusInternalServerError, "list_failed", "Failed to list invitations") + } + items := make([]fiber.Map, 0, len(invs)) + for i := range invs { + items = append(items, serializeInvitation(&invs[i])) + } + return c.JSON(fiber.Map{"ok": true, "invitations": items}) +} + +// RevokeInvitation handles DELETE /api/v1/teams/:team_id/invitations/:id. +// Owner / admin only. Marks the invitation revoked; returns 404 if missing, +// 410 Gone if already accepted, 403 if it belongs to another team. +func (h *TeamsHandler) RevokeInvitation(c *fiber.Ctx) error { + teamID, err := h.requireTeamMatch(c) + if err != nil { + return err + } + invID, err := uuid.Parse(c.Params("id")) + if err != nil { + return respondError(c, fiber.StatusBadRequest, "invalid_id", "Invalid invitation id") + } + + inv, err := models.GetRBACInvitationByID(c.Context(), h.db, invID) + if err != nil { + return teamsModelError(c, err) + } + if inv.TeamID != teamID { + return respondError(c, fiber.StatusForbidden, "forbidden", "Invitation does not belong to this team") + } + if inv.AcceptedAt.Valid { + return respondError(c, fiber.StatusGone, "already_accepted", "Invitation has already been accepted") + } + if err := models.RevokeRBACInvitation(c.Context(), h.db, invID); err != nil { + return teamsModelError(c, err) + } + return c.JSON(fiber.Map{"ok": true}) +} + +// AcceptInvitation handles POST /api/v1/invitations/:token/accept. +// +// No auth required — the token IS the auth. On success, the invitee's user row +// is created or updated to belong to the inviting team with the invited role, +// and a fresh session JWT is returned so the client can immediately call other +// authenticated endpoints. +// +// Status codes: +// +// 200 — accepted; body includes session_token + user/team info +// 404 — token unknown +// 410 — token already used or expired (single-use guarantee) +func (h *TeamsHandler) AcceptInvitation(c *fiber.Ctx) error { + token := c.Params("token") + if len(token) < 16 { + return respondError(c, fiber.StatusBadRequest, "invalid_token", "Invalid invitation token") + } + + user, inv, err := models.AcceptRBACInvitationByToken(c.Context(), h.db, token) + if err != nil { + return teamsModelError(c, err) + } + + team, err := models.GetTeamByID(c.Context(), h.db, inv.TeamID) + if err != nil { + return respondError(c, fiber.StatusInternalServerError, "team_lookup_failed", "Failed to load invited team") + } + + sessionToken, err := signSessionJWT(h.cfg.JWTSecret, user, team) + if err != nil { + return respondError(c, fiber.StatusInternalServerError, "session_failed", "Failed to issue session") + } + + return c.JSON(fiber.Map{ + "ok": true, + "session_token": sessionToken, + "user": fiber.Map{ + "id": user.ID.String(), + "email": user.Email, + "role": user.Role, + }, + "team": fiber.Map{ + "id": team.ID.String(), + "name": team.Name.String, + }, + }) +} + +// requireTeamMatch parses the :team_id path param and ensures it matches the +// authenticated team in the JWT. Returns the parsed UUID on success, or a +// fiber error (caller returns directly). +func (h *TeamsHandler) requireTeamMatch(c *fiber.Ctx) (uuid.UUID, error) { + pathTeamID, err := uuid.Parse(c.Params("team_id")) + if err != nil { + return uuid.Nil, respondError(c, fiber.StatusBadRequest, "invalid_team_id", "Invalid team id") + } + authTeamID := middleware.GetTeamID(c) + if authTeamID == "" { + return uuid.Nil, respondError(c, fiber.StatusUnauthorized, "unauthorized", "Valid session required") + } + if pathTeamID.String() != authTeamID { + return uuid.Nil, respondError(c, fiber.StatusForbidden, "forbidden", "Cannot act on another team") + } + return pathTeamID, nil +} + +// serializeInvitation produces the JSON shape returned by the invite endpoints. +// The token is included so owners/admins can re-share an invite link without +// triggering a new email send. +func serializeInvitation(inv *models.RBACInvitation) fiber.Map { + return fiber.Map{ + "id": inv.ID.String(), + "email": inv.Email, + "role": inv.Role, + "token": inv.Token, + "status": inv.Status(), + "invited_by": inv.InvitedBy.String(), + "expires_at": inv.ExpiresAt.UTC().Format(time.RFC3339), + "created_at": inv.CreatedAt.UTC().Format(time.RFC3339), + } +} + +// teamsModelError maps RBAC-invitation model errors to HTTP responses. +func teamsModelError(c *fiber.Ctx, err error) error { + switch { + case errors.Is(err, models.ErrInvitationNotFound): + return respondError(c, fiber.StatusNotFound, "not_found", err.Error()) + case errors.Is(err, models.ErrInvitationExpired), + errors.Is(err, models.ErrInvitationAlreadyAccepted), + errors.Is(err, models.ErrInvitationRevoked), + errors.Is(err, models.ErrInvitationNotPending): + return respondError(c, fiber.StatusGone, "invitation_invalid", err.Error()) + case errors.Is(err, models.ErrInvitationTokenInvalid): + return respondError(c, fiber.StatusBadRequest, "invalid_token", err.Error()) + case errors.Is(err, models.ErrInvalidInviteRole): + return respondError(c, fiber.StatusBadRequest, "invalid_role", err.Error()) + case errors.Is(err, models.ErrDuplicatePendingInvite): + return respondError(c, fiber.StatusConflict, "duplicate", err.Error()) + case errors.Is(err, models.ErrEmailMismatchInvite): + return respondError(c, fiber.StatusForbidden, "forbidden", err.Error()) + case errors.Is(err, models.ErrLastOwner): + return respondError(c, fiber.StatusConflict, "last_owner", err.Error()) + default: + return respondError(c, fiber.StatusInternalServerError, "internal_error", "Request failed") + } +} diff --git a/internal/handlers/teams_test.go b/internal/handlers/teams_test.go new file mode 100644 index 00000000..e7805e6a --- /dev/null +++ b/internal/handlers/teams_test.go @@ -0,0 +1,307 @@ +package handlers_test + +import ( + "bytes" + "context" + "database/sql" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "os" + "testing" + "time" + + "github.com/gofiber/fiber/v2" + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "instant.dev/internal/config" + "instant.dev/internal/email" + "instant.dev/internal/handlers" + "instant.dev/internal/middleware" + "instant.dev/internal/models" + "instant.dev/internal/testhelpers" +) + +// teamsApp builds a Fiber app wired to the real handler set used in production +// for the RBAC invite endpoints, plus a fake-auth middleware that injects +// (user_id, team_id, team_role) directly so the test can drive RBAC without +// minting JWTs. +// +// Routes registered (mirror what router.go will add): +// +// POST /api/v1/teams/:team_id/invitations (admin gate) +// GET /api/v1/teams/:team_id/invitations (admin gate) +// DELETE /api/v1/teams/:team_id/invitations/:id (admin gate) +// POST /api/v1/invitations/:token/accept (no auth) +func teamsApp(t *testing.T, db *sql.DB, actorUserID, actorTeamID, actorRole string) *fiber.App { + t.Helper() + cfg := &config.Config{ + JWTSecret: testhelpers.TestJWTSecret, + DashboardBaseURL: "http://localhost:5173", + } + mail := email.New("") // noop client — never actually sends + + app := fiber.New() + + // Fake auth: inject user/team/role into Locals so RequireRole can decide. + fakeAuth := func(c *fiber.Ctx) error { + if actorUserID != "" { + c.Locals(middleware.LocalKeyUserID, actorUserID) + } + if actorTeamID != "" { + c.Locals(middleware.LocalKeyTeamID, actorTeamID) + } + if actorRole != "" { + c.Locals(middleware.LocalKeyTeamRole, actorRole) + } + return c.Next() + } + + teamsH := handlers.NewTeamsHandler(db, cfg, mail) + + authedAdmin := app.Group("/api/v1/teams/:team_id/invitations", fakeAuth, middleware.RequireRole("admin")) + authedAdmin.Post("", teamsH.CreateInvitation) + authedAdmin.Get("", teamsH.ListInvitations) + authedAdmin.Delete("/:id", teamsH.RevokeInvitation) + + app.Post("/api/v1/invitations/:token/accept", teamsH.AcceptInvitation) + return app +} + +// teamsAppNeedsDB skips the test when no TEST_DATABASE_URL is set. +// Returns the DB and a cleanup. +func teamsAppNeedsDB(t *testing.T) (*sql.DB, func()) { + t.Helper() + if os.Getenv("TEST_DATABASE_URL") == "" { + t.Skip("teams_test: TEST_DATABASE_URL not set — skipping integration test") + } + return testhelpers.SetupTestDB(t) +} + +// seedTeam inserts a team and a single owner user. Returns (teamID, ownerID). +func seedTeam(t *testing.T, db *sql.DB) (uuid.UUID, uuid.UUID) { + t.Helper() + ctx := context.Background() + + teamID := uuid.MustParse(testhelpers.MustCreateTeamDB(t, db, "pro")) + ownerEmail := testhelpers.UniqueEmail(t) + user, err := models.CreateUser(ctx, db, teamID, ownerEmail, "", "", "owner") + require.NoError(t, err) + return teamID, user.ID +} + +// seedExtraUser creates a user on the same team with a given role. +func seedExtraUser(t *testing.T, db *sql.DB, teamID uuid.UUID, role string) uuid.UUID { + t.Helper() + user, err := models.CreateUser(context.Background(), db, + teamID, testhelpers.UniqueEmail(t), "", "", role) + require.NoError(t, err) + return user.ID +} + +func postJSON(t *testing.T, app *fiber.App, path string, body any) *http.Response { + t.Helper() + var buf bytes.Buffer + if body != nil { + require.NoError(t, json.NewEncoder(&buf).Encode(body)) + } + req := httptest.NewRequest(http.MethodPost, path, &buf) + req.Header.Set("Content-Type", "application/json") + resp, err := app.Test(req, 5000) + require.NoError(t, err) + return resp +} + +func decode(t *testing.T, resp *http.Response) map[string]any { + t.Helper() + defer resp.Body.Close() + var out map[string]any + require.NoError(t, json.NewDecoder(resp.Body).Decode(&out)) + return out +} + +// TestInvite_OwnerCanInvite — happy path: owner POST returns 201 and a token. +func TestInvite_OwnerCanInvite(t *testing.T) { + db, cleanup := teamsAppNeedsDB(t) + defer cleanup() + teamID, ownerID := seedTeam(t, db) + + app := teamsApp(t, db, ownerID.String(), teamID.String(), "owner") + resp := postJSON(t, app, "/api/v1/teams/"+teamID.String()+"/invitations", + map[string]string{"email": testhelpers.UniqueEmail(t), "role": "developer"}) + + require.Equal(t, http.StatusCreated, resp.StatusCode) + body := decode(t, resp) + assert.Equal(t, true, body["ok"]) + inv, _ := body["invitation"].(map[string]any) + require.NotNil(t, inv) + assert.NotEmpty(t, inv["token"]) + assert.Equal(t, "developer", inv["role"]) +} + +// TestInvite_AdminCanInvite — admin role passes RequireRole("admin"). +func TestInvite_AdminCanInvite(t *testing.T) { + db, cleanup := teamsAppNeedsDB(t) + defer cleanup() + teamID, _ := seedTeam(t, db) + adminID := seedExtraUser(t, db, teamID, "admin") + + app := teamsApp(t, db, adminID.String(), teamID.String(), "admin") + resp := postJSON(t, app, "/api/v1/teams/"+teamID.String()+"/invitations", + map[string]string{"email": testhelpers.UniqueEmail(t), "role": "viewer"}) + defer resp.Body.Close() + assert.Equal(t, http.StatusCreated, resp.StatusCode) +} + +// TestInvite_DeveloperCannotInvite — developer is below the admin gate. +func TestInvite_DeveloperCannotInvite(t *testing.T) { + db, cleanup := teamsAppNeedsDB(t) + defer cleanup() + teamID, _ := seedTeam(t, db) + devID := seedExtraUser(t, db, teamID, "developer") + + app := teamsApp(t, db, devID.String(), teamID.String(), "developer") + resp := postJSON(t, app, "/api/v1/teams/"+teamID.String()+"/invitations", + map[string]string{"email": testhelpers.UniqueEmail(t), "role": "viewer"}) + defer resp.Body.Close() + assert.Equal(t, http.StatusForbidden, resp.StatusCode) +} + +// TestInvite_ViewerCannotInvite — viewer is the lowest tier; clearly blocked. +func TestInvite_ViewerCannotInvite(t *testing.T) { + db, cleanup := teamsAppNeedsDB(t) + defer cleanup() + teamID, _ := seedTeam(t, db) + viewerID := seedExtraUser(t, db, teamID, "viewer") + + app := teamsApp(t, db, viewerID.String(), teamID.String(), "viewer") + resp := postJSON(t, app, "/api/v1/teams/"+teamID.String()+"/invitations", + map[string]string{"email": testhelpers.UniqueEmail(t), "role": "viewer"}) + defer resp.Body.Close() + assert.Equal(t, http.StatusForbidden, resp.StatusCode) +} + +// TestInvite_TokenSingleUse — accepting twice returns 410 Gone. +func TestInvite_TokenSingleUse(t *testing.T) { + db, cleanup := teamsAppNeedsDB(t) + defer cleanup() + teamID, ownerID := seedTeam(t, db) + + inviteEmail := testhelpers.UniqueEmail(t) + inv, err := models.CreateRBACInvitation(context.Background(), db, teamID, inviteEmail, "developer", ownerID) + require.NoError(t, err) + + // Need an app — actor identity doesn't matter for AcceptInvitation (no auth). + app := teamsApp(t, db, "", "", "") + + r1 := postJSON(t, app, "/api/v1/invitations/"+inv.Token+"/accept", nil) + require.Equal(t, http.StatusOK, r1.StatusCode, "first accept must succeed") + body := decode(t, r1) + assert.NotEmpty(t, body["session_token"], "first accept must mint a session JWT") + + r2 := postJSON(t, app, "/api/v1/invitations/"+inv.Token+"/accept", nil) + defer r2.Body.Close() + assert.Equal(t, http.StatusGone, r2.StatusCode, "second accept must return 410") +} + +// TestInvite_TokenExpiry — > 7 days old returns 410 Gone. +func TestInvite_TokenExpiry(t *testing.T) { + db, cleanup := teamsAppNeedsDB(t) + defer cleanup() + teamID, ownerID := seedTeam(t, db) + + // Create the row, then backdate expires_at to simulate a stale invite. + inviteEmail := testhelpers.UniqueEmail(t) + inv, err := models.CreateRBACInvitation(context.Background(), db, teamID, inviteEmail, "developer", ownerID) + require.NoError(t, err) + _, err = db.Exec(`UPDATE team_invitations SET expires_at = $1 WHERE id = $2`, + time.Now().Add(-1*time.Hour), inv.ID) + require.NoError(t, err) + + app := teamsApp(t, db, "", "", "") + resp := postJSON(t, app, "/api/v1/invitations/"+inv.Token+"/accept", nil) + defer resp.Body.Close() + assert.Equal(t, http.StatusGone, resp.StatusCode) +} + +// TestInvite_LastOwnerProtected — last remaining owner cannot leave or be downgraded. +// +// EnsureNotLastOwner guards CreatePersonalTeamAndReassignUser-style flows. Direct +// model assertion (no HTTP) since the dashboard "leave team" surface lives in +// team_members.go (legacy handler) and the corresponding RBAC-aware UX is not +// part of this PR — the helper is in place for Phase 4 to wire. +func TestInvite_LastOwnerProtected(t *testing.T) { + db, cleanup := teamsAppNeedsDB(t) + defer cleanup() + teamID, ownerID := seedTeam(t, db) + ctx := context.Background() + + // Sole owner: must be blocked. + err := models.EnsureNotLastOwner(ctx, db, teamID, ownerID) + require.ErrorIs(t, err, models.ErrLastOwner) + + // Add a second owner: now the original owner is no longer "last" — allowed. + _ = seedExtraUser(t, db, teamID, "owner") + err = models.EnsureNotLastOwner(ctx, db, teamID, ownerID) + assert.NoError(t, err) +} + +// TestInvite_TeamIDMismatch — actor's JWT team must match :team_id path param. +func TestInvite_TeamIDMismatch(t *testing.T) { + db, cleanup := teamsAppNeedsDB(t) + defer cleanup() + teamA, ownerA := seedTeam(t, db) + teamB, _ := seedTeam(t, db) + + // Actor is owner of team A; tries to act on team B. + app := teamsApp(t, db, ownerA.String(), teamA.String(), "owner") + resp := postJSON(t, app, "/api/v1/teams/"+teamB.String()+"/invitations", + map[string]string{"email": testhelpers.UniqueEmail(t), "role": "viewer"}) + defer resp.Body.Close() + assert.Equal(t, http.StatusForbidden, resp.StatusCode) +} + +// TestInvite_RoleValidation — only admin/developer/viewer are valid invite roles. +func TestInvite_RoleValidation(t *testing.T) { + db, cleanup := teamsAppNeedsDB(t) + defer cleanup() + teamID, ownerID := seedTeam(t, db) + + app := teamsApp(t, db, ownerID.String(), teamID.String(), "owner") + + for _, badRole := range []string{"owner", "root", "", "admin\""} { + t.Run(fmt.Sprintf("role=%q", badRole), func(t *testing.T) { + resp := postJSON(t, app, "/api/v1/teams/"+teamID.String()+"/invitations", + map[string]string{"email": testhelpers.UniqueEmail(t), "role": badRole}) + defer resp.Body.Close() + assert.Equal(t, http.StatusBadRequest, resp.StatusCode) + }) + } +} + +// TestInvite_RevokeFlow — owner can revoke a pending invite. +func TestInvite_RevokeFlow(t *testing.T) { + db, cleanup := teamsAppNeedsDB(t) + defer cleanup() + teamID, ownerID := seedTeam(t, db) + + inv, err := models.CreateRBACInvitation(context.Background(), db, + teamID, testhelpers.UniqueEmail(t), "developer", ownerID) + require.NoError(t, err) + + app := teamsApp(t, db, ownerID.String(), teamID.String(), "owner") + req := httptest.NewRequest(http.MethodDelete, + "/api/v1/teams/"+teamID.String()+"/invitations/"+inv.ID.String(), nil) + resp, err := app.Test(req, 5000) + require.NoError(t, err) + defer resp.Body.Close() + assert.Equal(t, http.StatusOK, resp.StatusCode) + + // Token should now refuse to accept. + r2 := postJSON(t, app, "/api/v1/invitations/"+inv.Token+"/accept", nil) + defer r2.Body.Close() + assert.Equal(t, http.StatusGone, r2.StatusCode) +} diff --git a/internal/handlers/vault.go b/internal/handlers/vault.go new file mode 100644 index 00000000..9a8db219 --- /dev/null +++ b/internal/handlers/vault.go @@ -0,0 +1,440 @@ +package handlers + +// vault.go — per-team encrypted secret storage. +// +// Endpoints (all require team JWT, registered behind RequireAuth in router.go): +// PUT /api/v1/vault/:env/:key body {"value":"..."} → 201 {key,version} +// GET /api/v1/vault/:env/:key[?version=N] → 200 {key,value,version} +// GET /api/v1/vault/:env → 200 {keys:[...]} (no values) +// DELETE /api/v1/vault/:env/:key → 204 (hard delete: removes ALL versions) +// POST /api/v1/vault/:env/:key/rotate body {"value":"..."} → 201 {key,version} (alias for PUT) +// +// Encryption: AES-256-GCM, key from cfg.AESKey (64-char hex). Stored as raw bytes +// in vault_secrets.encrypted_value (BYTEA). The base64 wrapper produced by +// crypto.Encrypt is decoded before insert and re-encoded for tamper checks. +// +// Isolation: every query is scoped by team_id pulled from the session JWT. +// Foreign reads return 404 — never 403 — so existence of a secret in another +// team is never observable. There is no "list all" endpoint and no value +// is ever returned by the list-keys path. +// +// Audit: every mutation (PUT/DELETE/rotate) and every successful GET writes a +// row to vault_audit_log. Audit failures are logged but never block the request. +// +// DELETE semantics: hard delete of ALL versions for (team,env,key). Chosen over +// tombstone-row to keep access checks simple and the hot table small. The audit +// log preserves the action durably. + +import ( + "database/sql" + "encoding/base64" + "errors" + "fmt" + "log/slog" + "strconv" + "strings" + + "github.com/gofiber/fiber/v2" + "github.com/google/uuid" + "instant.dev/internal/config" + "instant.dev/internal/crypto" + "instant.dev/internal/middleware" + "instant.dev/internal/models" + "instant.dev/internal/plans" +) + +// vaultDefaultEnv is the env path segment treated as the default production environment. +const vaultDefaultEnv = "production" + +// vaultMaxKeyLen bounds keys to a sane length. Unix env-var conventions cap at +// names of this size on most shells; matching keeps later /deploy injection sane. +const vaultMaxKeyLen = 256 + +// vaultMaxValueBytes caps plaintext value size pre-encryption. 1 MiB is plenty +// for typical secrets (DB URLs, API tokens, TLS bundles) without enabling abuse. +const vaultMaxValueBytes = 1 << 20 // 1 MiB + +// vaultErrInternal / vaultErrInvalidBody / etc. — keep error codes as named consts +// so callers can match on them and we don't sprinkle string literals through handlers. +const ( + vaultErrInvalidBody = "invalid_body" + vaultErrInvalidKey = "invalid_key" + vaultErrInvalidEnv = "invalid_env" + vaultErrInvalidValue = "invalid_value" + vaultErrUnauthorized = "unauthorized" + vaultErrNotFound = "not_found" + vaultErrInternal = "internal_error" + vaultErrPersist = "persist_failed" + vaultErrNotAvailable = "vault_not_available" + vaultErrQuotaExceeded = "vault_quota_exceeded" + vaultErrEnvNotAllowed = "vault_env_not_allowed" +) + +// VaultHandler serves vault endpoints. All endpoints require an authenticated team. +type VaultHandler struct { + db *sql.DB + cfg *config.Config + plans *plans.Registry +} + +// NewVaultHandler constructs a VaultHandler. +func NewVaultHandler(db *sql.DB, cfg *config.Config, reg *plans.Registry) *VaultHandler { + return &VaultHandler{db: db, cfg: cfg, plans: reg} +} + +// vaultBody is the request body for PUT /api/v1/vault/:env/:key and the rotate alias. +type vaultBody struct { + Value string `json:"value"` +} + +// authContext extracts (teamID, userID, ip) from the fiber context. Returns the +// 401 response and ok=false when the team JWT is missing/malformed. Routes that +// reach this handler are already guarded by RequireAuth, so this is a sanity net. +func (h *VaultHandler) authContext(c *fiber.Ctx) (uuid.UUID, uuid.NullUUID, string, error) { + teamIDStr := middleware.GetTeamID(c) + teamID, err := uuid.Parse(teamIDStr) + if err != nil { + return uuid.Nil, uuid.NullUUID{}, "", errors.New("invalid team id in token") + } + var userID uuid.NullUUID + if uidStr := middleware.GetUserID(c); uidStr != "" { + if uid, err := uuid.Parse(uidStr); err == nil { + userID = uuid.NullUUID{UUID: uid, Valid: true} + } + } + return teamID, userID, c.IP(), nil +} + +// validateEnv enforces that env is non-empty and contains only safe path-friendly chars. +// Default to "production" when callers send an empty string (matches the migration default). +func validateEnv(env string) (string, bool) { + env = strings.TrimSpace(env) + if env == "" { + env = vaultDefaultEnv + } + if len(env) > 64 { + return "", false + } + for _, r := range env { + switch { + case r >= 'a' && r <= 'z': + case r >= 'A' && r <= 'Z': + case r >= '0' && r <= '9': + case r == '-' || r == '_': + default: + return "", false + } + } + return env, true +} + +// validateKey enforces that key is non-empty, within length, and contains only +// characters legal in env-var names plus '.' and '-' for namespacing. +func validateKey(key string) (string, bool) { + key = strings.TrimSpace(key) + if key == "" || len(key) > vaultMaxKeyLen { + return "", false + } + for _, r := range key { + switch { + case r >= 'a' && r <= 'z': + case r >= 'A' && r <= 'Z': + case r >= '0' && r <= '9': + case r == '_' || r == '-' || r == '.': + default: + return "", false + } + } + return key, true +} + +// encryptPlaintext returns the raw GCM ciphertext bytes (nonce||ciphertext||tag). +// The shared crypto.Encrypt helper returns a base64url string; we decode it once +// here so the at-rest representation is opaque BYTEA, not text. +func (h *VaultHandler) encryptPlaintext(plain string) ([]byte, error) { + key, err := crypto.ParseAESKey(h.cfg.AESKey) + if err != nil { + return nil, err + } + encoded, err := crypto.Encrypt(key, plain) + if err != nil { + return nil, err + } + raw, err := base64.URLEncoding.DecodeString(encoded) + if err != nil { + return nil, err + } + return raw, nil +} + +// decryptCiphertext reverses encryptPlaintext. Tamper failures (corrupted bytes, +// wrong key) surface as *crypto.ErrDecrypt — handlers map that to 500, never 200. +func (h *VaultHandler) decryptCiphertext(raw []byte) (string, error) { + key, err := crypto.ParseAESKey(h.cfg.AESKey) + if err != nil { + return "", err + } + encoded := base64.URLEncoding.EncodeToString(raw) + return crypto.Decrypt(key, encoded) +} + +// audit appends a vault_audit_log row best-effort. Failures are logged but never +// surface to the caller — auditing must not block the request. +func (h *VaultHandler) audit(c *fiber.Ctx, teamID uuid.UUID, userID uuid.NullUUID, action, env, key, ip string) { + if err := models.AppendVaultAudit(c.UserContext(), h.db, teamID, userID, action, env, key, ip); err != nil { + slog.Error("vault.audit_failed", + "error", err, + "team_id", teamID, + "action", action, + "env", env, + "key", key, + "request_id", middleware.GetRequestID(c), + ) + } +} + +// PutSecret handles PUT /api/v1/vault/:env/:key. +// Always creates a new version. Returns 201 with {key,version}. +func (h *VaultHandler) PutSecret(c *fiber.Ctx) error { + return h.upsertSecret(c, "set") +} + +// RotateSecret handles POST /api/v1/vault/:env/:key/rotate. +// Semantics are identical to PUT — exposed under a different action name so the +// audit log distinguishes intentional rotation from a regular write. +func (h *VaultHandler) RotateSecret(c *fiber.Ctx) error { + return h.upsertSecret(c, "rotate") +} + +func (h *VaultHandler) upsertSecret(c *fiber.Ctx, action string) error { + teamID, userID, ip, err := h.authContext(c) + if err != nil { + return respondError(c, fiber.StatusUnauthorized, vaultErrUnauthorized, "Valid session token required") + } + + env, ok := validateEnv(c.Params("env")) + if !ok { + return respondError(c, fiber.StatusBadRequest, vaultErrInvalidEnv, "env must be 1-64 chars [A-Za-z0-9_-]") + } + key, ok := validateKey(c.Params("key")) + if !ok { + return respondError(c, fiber.StatusBadRequest, vaultErrInvalidKey, "key must be 1-256 chars [A-Za-z0-9_.-]") + } + + var body vaultBody + if err := c.BodyParser(&body); err != nil { + return respondError(c, fiber.StatusBadRequest, vaultErrInvalidBody, "Request body must be valid JSON") + } + if len(body.Value) > vaultMaxValueBytes { + return respondError(c, fiber.StatusRequestEntityTooLarge, vaultErrInvalidValue, "value exceeds 1 MiB cap") + } + + // Per-tier quota + env restriction. Fetch team to read its plan tier. + // If h.plans is nil (older test paths that haven't been updated), we fall + // open and skip tier checks — never block on plumbing. + if h.plans != nil { + team, terr := models.GetTeamByID(c.Context(), h.db, teamID) + if terr != nil { + slog.Warn("vault.tier.team_lookup_failed", + "error", terr, "team_id", teamID, + "request_id", middleware.GetRequestID(c)) + } else if team != nil { + // On rotate, we already require an existing key (rotate of a missing + // key is rejected by upsertSecret semantics). For PUT/set we must + // allow updating an existing key without burning a quota slot. + // + // Tier check 1: vault availability + quota (skip on rotate — count + // can only stay flat or shrink). + if action != "rotate" { + maxEntries := h.plans.VaultMaxEntries(team.PlanTier) + if maxEntries == 0 { + return respondError(c, fiber.StatusForbidden, vaultErrNotAvailable, + "Vault is not available on the "+team.PlanTier+" tier. Upgrade to Hobby or higher.") + } + if maxEntries > 0 { + n, cerr := models.CountVaultKeysByTeam(c.Context(), h.db, teamID) + if cerr != nil { + slog.Warn("vault.put.count_failed", "error", cerr, "team_id", teamID) + } else { + // Allow updating an existing key (won't grow the count). + // TODO(race): the count + insert is not transactional, so two + // concurrent PUTs at quota-1 may both succeed and exceed the cap. + // Accept this for now; revisit with SELECT FOR UPDATE if abuse appears. + existing, _ := models.GetVaultSecretLatest(c.Context(), h.db, teamID, env, key) + if existing == nil && n >= maxEntries { + return respondError(c, fiber.StatusPaymentRequired, vaultErrQuotaExceeded, + fmt.Sprintf("Plan %q allows %d vault entries; you have %d. Upgrade to add more.", + team.PlanTier, maxEntries, n)) + } + } + } + } + + // Tier check 2: env restriction (applies to both PUT and rotate). + allowed := h.plans.VaultEnvsAllowed(team.PlanTier) + if len(allowed) > 0 { + envOK := false + for _, a := range allowed { + if a == env { + envOK = true + break + } + } + if !envOK { + return respondError(c, fiber.StatusForbidden, vaultErrEnvNotAllowed, + fmt.Sprintf("Plan %q only allows vault env %v; got %q. Upgrade to Pro for multi-env vault.", + team.PlanTier, allowed, env)) + } + } + } + } + + ciphertext, err := h.encryptPlaintext(body.Value) + if err != nil { + slog.Error("vault.encrypt_failed", + "error", err, "team_id", teamID, "env", env, "key", key, + "request_id", middleware.GetRequestID(c)) + return respondError(c, fiber.StatusInternalServerError, vaultErrInternal, "Encryption failed") + } + + secret, err := models.CreateVaultSecret(c.UserContext(), h.db, teamID, env, key, ciphertext, userID) + if err != nil { + slog.Error("vault.persist_failed", + "error", err, "team_id", teamID, "env", env, "key", key, + "request_id", middleware.GetRequestID(c)) + return respondError(c, fiber.StatusServiceUnavailable, vaultErrPersist, "Failed to persist secret") + } + + h.audit(c, teamID, userID, action, env, key, ip) + + return c.Status(fiber.StatusCreated).JSON(fiber.Map{ + "ok": true, + "key": secret.Key, + "env": secret.Env, + "version": secret.Version, + }) +} + +// GetSecret handles GET /api/v1/vault/:env/:key[?version=N]. +// Cross-team or missing → 404 (never 403). +func (h *VaultHandler) GetSecret(c *fiber.Ctx) error { + teamID, userID, ip, err := h.authContext(c) + if err != nil { + return respondError(c, fiber.StatusUnauthorized, vaultErrUnauthorized, "Valid session token required") + } + + env, ok := validateEnv(c.Params("env")) + if !ok { + return respondError(c, fiber.StatusBadRequest, vaultErrInvalidEnv, "env must be 1-64 chars [A-Za-z0-9_-]") + } + key, ok := validateKey(c.Params("key")) + if !ok { + return respondError(c, fiber.StatusBadRequest, vaultErrInvalidKey, "key must be 1-256 chars [A-Za-z0-9_.-]") + } + + var ( + secret *models.VaultSecret + fetchErr error + ) + if v := strings.TrimSpace(c.Query("version")); v != "" { + n, perr := strconv.Atoi(v) + if perr != nil || n <= 0 { + return respondError(c, fiber.StatusBadRequest, vaultErrInvalidBody, "version must be a positive integer") + } + secret, fetchErr = models.GetVaultSecretVersion(c.UserContext(), h.db, teamID, env, key, n) + } else { + secret, fetchErr = models.GetVaultSecretLatest(c.UserContext(), h.db, teamID, env, key) + } + + if errors.Is(fetchErr, models.ErrVaultSecretNotFound) { + return respondError(c, fiber.StatusNotFound, vaultErrNotFound, "secret not found") + } + if fetchErr != nil { + slog.Error("vault.fetch_failed", + "error", fetchErr, "team_id", teamID, "env", env, "key", key, + "request_id", middleware.GetRequestID(c)) + return respondError(c, fiber.StatusInternalServerError, vaultErrInternal, "Failed to fetch secret") + } + + plain, err := h.decryptCiphertext(secret.EncryptedValue) + if err != nil { + slog.Error("vault.decrypt_failed", + "error", err, "team_id", teamID, "env", env, "key", key, + "request_id", middleware.GetRequestID(c)) + return respondError(c, fiber.StatusInternalServerError, vaultErrInternal, "Failed to decrypt secret") + } + + h.audit(c, teamID, userID, "get", env, key, ip) + + return c.JSON(fiber.Map{ + "ok": true, + "key": secret.Key, + "env": secret.Env, + "value": plain, + "version": secret.Version, + }) +} + +// ListKeys handles GET /api/v1/vault/:env. Returns key names only — never values. +func (h *VaultHandler) ListKeys(c *fiber.Ctx) error { + teamID, userID, ip, err := h.authContext(c) + if err != nil { + return respondError(c, fiber.StatusUnauthorized, vaultErrUnauthorized, "Valid session token required") + } + + env, ok := validateEnv(c.Params("env")) + if !ok { + return respondError(c, fiber.StatusBadRequest, vaultErrInvalidEnv, "env must be 1-64 chars [A-Za-z0-9_-]") + } + + keys, err := models.ListVaultKeys(c.UserContext(), h.db, teamID, env) + if err != nil { + slog.Error("vault.list_failed", + "error", err, "team_id", teamID, "env", env, + "request_id", middleware.GetRequestID(c)) + return respondError(c, fiber.StatusInternalServerError, vaultErrInternal, "Failed to list secrets") + } + + // Audit list ops with a synthetic key so every read leaves a trail without + // needing to enumerate fan-out per-key. + h.audit(c, teamID, userID, "list", env, "*", ip) + + return c.JSON(fiber.Map{ + "ok": true, + "env": env, + "keys": keys, + }) +} + +// DeleteSecret handles DELETE /api/v1/vault/:env/:key. +// Hard delete of all versions for (team,env,key). 204 on success, 404 when +// the secret does not exist for this team (idempotent + non-leaking). +func (h *VaultHandler) DeleteSecret(c *fiber.Ctx) error { + teamID, userID, ip, err := h.authContext(c) + if err != nil { + return respondError(c, fiber.StatusUnauthorized, vaultErrUnauthorized, "Valid session token required") + } + + env, ok := validateEnv(c.Params("env")) + if !ok { + return respondError(c, fiber.StatusBadRequest, vaultErrInvalidEnv, "env must be 1-64 chars [A-Za-z0-9_-]") + } + key, ok := validateKey(c.Params("key")) + if !ok { + return respondError(c, fiber.StatusBadRequest, vaultErrInvalidKey, "key must be 1-256 chars [A-Za-z0-9_.-]") + } + + n, err := models.DeleteVaultSecret(c.UserContext(), h.db, teamID, env, key) + if err != nil { + slog.Error("vault.delete_failed", + "error", err, "team_id", teamID, "env", env, "key", key, + "request_id", middleware.GetRequestID(c)) + return respondError(c, fiber.StatusInternalServerError, vaultErrInternal, "Failed to delete secret") + } + if n == 0 { + return respondError(c, fiber.StatusNotFound, vaultErrNotFound, "secret not found") + } + + h.audit(c, teamID, userID, "delete", env, key, ip) + return c.SendStatus(fiber.StatusNoContent) +} diff --git a/internal/handlers/vault_resolve.go b/internal/handlers/vault_resolve.go new file mode 100644 index 00000000..6d284a95 --- /dev/null +++ b/internal/handlers/vault_resolve.go @@ -0,0 +1,98 @@ +package handlers + +import ( + "context" + "database/sql" + "encoding/base64" + "errors" + "fmt" + "log/slog" + "strings" + + "github.com/google/uuid" + "instant.dev/internal/crypto" + "instant.dev/internal/models" +) + +// vaultRefPrefix is the syntax used in deployment env_vars to reference a +// vault secret. Values starting with this prefix are resolved at deploy time +// against vault_secrets for the team's current environment. +// +// { "RAZORPAY_KEY_SECRET": "vault://RAZORPAY_KEY_SECRET" } +// +// At deploy time, the value is replaced with the latest version of the named +// secret. Plaintext is never written to deployments.env_vars or any log. +const vaultRefPrefix = "vault://" + +// ErrVaultRefMissing is returned when a deployment references a vault key +// that does not exist for the team in the requested environment. +var ErrVaultRefMissing = errors.New("vault reference not found") + +// ResolveVaultRefs replaces every "vault://KEY" value in vars with the +// decrypted plaintext from the team's vault for the given environment. +// Non-prefixed values are passed through unchanged. +// +// The returned map is a fresh allocation; the input map is not mutated. +// +// Each resolved key is appended to vault_audit_log with action +// "read_for_deploy" — best-effort, audit failure does not block the deploy. +// +// If any reference cannot be resolved (key missing, ciphertext tampered), +// returns ErrVaultRefMissing wrapping the underlying cause. The caller +// fails the deploy with a clear error so the user knows which secret to add. +func ResolveVaultRefs( + ctx context.Context, + db *sql.DB, + aesKeyHex string, + teamID uuid.UUID, + env string, + vars map[string]string, +) (map[string]string, error) { + out := make(map[string]string, len(vars)) + var aesKey []byte + var aesKeyErr error + + for k, v := range vars { + if !strings.HasPrefix(v, vaultRefPrefix) { + out[k] = v + continue + } + secretKey := strings.TrimPrefix(v, vaultRefPrefix) + if secretKey == "" { + return nil, fmt.Errorf("%w: empty key in vault://", ErrVaultRefMissing) + } + + // Lazy-parse the AES key once per call (only when we actually have refs). + if aesKey == nil && aesKeyErr == nil { + aesKey, aesKeyErr = crypto.ParseAESKey(aesKeyHex) + } + if aesKeyErr != nil { + return nil, fmt.Errorf("vault resolve: %w", aesKeyErr) + } + + row, err := models.GetVaultSecretLatest(ctx, db, teamID, env, secretKey) + if err != nil { + if errors.Is(err, models.ErrVaultSecretNotFound) { + return nil, fmt.Errorf("%w: %s/%s", ErrVaultRefMissing, env, secretKey) + } + return nil, fmt.Errorf("vault resolve %s: %w", secretKey, err) + } + + encoded := base64.URLEncoding.EncodeToString(row.EncryptedValue) + plain, err := crypto.Decrypt(aesKey, encoded) + if err != nil { + return nil, fmt.Errorf("vault decrypt %s: %w", secretKey, err) + } + out[k] = plain + + // Best-effort audit. Failures logged but never block. + if auditErr := models.AppendVaultAudit(ctx, db, teamID, uuid.NullUUID{}, "read_for_deploy", env, secretKey, ""); auditErr != nil { + slog.Warn("vault.audit_failed", + "action", "read_for_deploy", + "team_id", teamID, "env", env, "key", secretKey, + "error", auditErr) + } + } + + return out, nil +} diff --git a/internal/handlers/vault_resolve_test.go b/internal/handlers/vault_resolve_test.go new file mode 100644 index 00000000..fb186021 --- /dev/null +++ b/internal/handlers/vault_resolve_test.go @@ -0,0 +1,181 @@ +package handlers_test + +// vault_resolve_test.go — covers handlers.ResolveVaultRefs, the helper that +// substitutes "vault://KEY" entries in deployment env_vars with decrypted +// plaintext from the team's vault. +// +// Three groups of tests: +// - TestResolveVaultRefs_NoRefs_PassesThrough : pure-unit, no DB +// - TestResolveVaultRefs_EmptyKey_Errors : pure-unit, no DB +// - TestResolveVaultRefs_DecryptsKnownSecret : integration, needs DB +// - TestResolveVaultRefs_MissingKey_ReturnsError : integration, needs DB + +import ( + "context" + "encoding/base64" + "errors" + "os" + "strings" + "testing" + + "github.com/google/uuid" + "instant.dev/internal/crypto" + "instant.dev/internal/handlers" + "instant.dev/internal/models" + "instant.dev/internal/testhelpers" +) + +// TestResolveVaultRefs_NoRefs_PassesThrough verifies non-prefixed values +// flow through untouched without DB access. +func TestResolveVaultRefs_NoRefs_PassesThrough(t *testing.T) { + in := map[string]string{ + "DATABASE_URL": "postgres://u:p@host/db", + "PORT": "8080", + "FEATURE_FLAG": "true", + } + out, err := handlers.ResolveVaultRefs( + context.Background(), + nil, // db unused — no vault refs + "", // aes key unused — no vault refs + uuid.New(), + "production", + in, + ) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(out) != len(in) { + t.Fatalf("len mismatch: in=%d out=%d", len(in), len(out)) + } + for k, v := range in { + if out[k] != v { + t.Errorf("key %q: want %q, got %q", k, v, out[k]) + } + } +} + +// TestResolveVaultRefs_EmptyKey_Errors verifies that "vault://" with no key +// is rejected (not silently treated as empty key). +func TestResolveVaultRefs_EmptyKey_Errors(t *testing.T) { + in := map[string]string{"BAD": "vault://"} + _, err := handlers.ResolveVaultRefs( + context.Background(), nil, "", + uuid.New(), "production", in, + ) + if err == nil { + t.Fatal("want error for empty vault:// key, got nil") + } + if !errors.Is(err, handlers.ErrVaultRefMissing) { + t.Errorf("want ErrVaultRefMissing, got %v", err) + } +} + +// TestResolveVaultRefs_DecryptsKnownSecret seeds a vault row, calls the +// resolver, and verifies the value is replaced with the decrypted plaintext. +// Skips when TEST_DATABASE_URL is unset. +func TestResolveVaultRefs_DecryptsKnownSecret(t *testing.T) { + dsn := os.Getenv("TEST_DATABASE_URL") + if dsn == "" { + t.Skip("TEST_DATABASE_URL not set — skipping integration test") + } + db, cleanup := testhelpers.SetupTestDB(t) + defer cleanup() + + teamID := uuid.New() + if _, err := db.Exec( + `INSERT INTO teams (id, name, plan_tier) VALUES ($1, $2, 'pro')`, + teamID, "vault-resolve-test-"+teamID.String()[:8], + ); err != nil { + t.Fatalf("seed team: %v", err) + } + + // Generate an AES key + encrypt a known plaintext. + aesKeyHex := "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef" // 32 bytes hex + aesKey, err := crypto.ParseAESKey(aesKeyHex) + if err != nil { + t.Fatalf("ParseAESKey: %v", err) + } + plaintext := "sk_live_super_secret_value_xyz" + encoded, err := crypto.Encrypt(aesKey, plaintext) + if err != nil { + t.Fatalf("Encrypt: %v", err) + } + // vault stores raw bytes — decode the base64 wrapper. + rawBytes, err := base64.URLEncoding.DecodeString(encoded) + if err != nil { + t.Fatalf("decode wrapper: %v", err) + } + + if _, err := models.CreateVaultSecret( + context.Background(), db, teamID, + "production", "RAZORPAY_KEY_SECRET", rawBytes, uuid.NullUUID{}, + ); err != nil { + t.Fatalf("CreateVaultSecret: %v", err) + } + + in := map[string]string{ + "PUBLIC_VAR": "not-a-secret", + "RAZORPAY_KEY": "vault://RAZORPAY_KEY_SECRET", + } + out, err := handlers.ResolveVaultRefs( + context.Background(), db, aesKeyHex, teamID, "production", in, + ) + if err != nil { + t.Fatalf("ResolveVaultRefs: %v", err) + } + if out["PUBLIC_VAR"] != "not-a-secret" { + t.Errorf("non-vault value mutated: got %q", out["PUBLIC_VAR"]) + } + if out["RAZORPAY_KEY"] != plaintext { + t.Errorf("vault value not decrypted: got %q want %q", out["RAZORPAY_KEY"], plaintext) + } + + // Audit log should record one read_for_deploy entry. + count, err := models.CountVaultAudit( + context.Background(), db, teamID, + "read_for_deploy", "production", "RAZORPAY_KEY_SECRET", + ) + if err != nil { + t.Fatalf("CountVaultAudit: %v", err) + } + if count != 1 { + t.Errorf("audit count: want 1, got %d", count) + } +} + +// TestResolveVaultRefs_MissingKey_ReturnsError verifies that referencing a +// key the team has not stored returns ErrVaultRefMissing. +func TestResolveVaultRefs_MissingKey_ReturnsError(t *testing.T) { + dsn := os.Getenv("TEST_DATABASE_URL") + if dsn == "" { + t.Skip("TEST_DATABASE_URL not set — skipping integration test") + } + db, cleanup := testhelpers.SetupTestDB(t) + defer cleanup() + + teamID := uuid.New() + if _, err := db.Exec( + `INSERT INTO teams (id, name, plan_tier) VALUES ($1, $2, 'pro')`, + teamID, "vault-miss-test-"+teamID.String()[:8], + ); err != nil { + t.Fatalf("seed team: %v", err) + } + + in := map[string]string{"X": "vault://NOT_THERE"} + _, err := handlers.ResolveVaultRefs( + context.Background(), db, + "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef", + teamID, "production", in, + ) + if err == nil { + t.Fatal("want error, got nil") + } + if !errors.Is(err, handlers.ErrVaultRefMissing) { + t.Errorf("want ErrVaultRefMissing, got %v", err) + } + if !strings.Contains(err.Error(), "NOT_THERE") { + t.Errorf("error should mention the missing key, got %v", err) + } +} + + diff --git a/internal/handlers/vault_test.go b/internal/handlers/vault_test.go new file mode 100644 index 00000000..972b8630 --- /dev/null +++ b/internal/handlers/vault_test.go @@ -0,0 +1,578 @@ +package handlers_test + +// vault_test.go — coverage for /api/v1/vault/* endpoints. +// +// Layered tests: +// - TestVault_AESRoundtrip : crypto contract used by the handler +// - TestVault_TeamIsolation : team A's JWT cannot read team B's secret (404, never 403) +// - TestVault_AuditLog : every mutation + read writes a vault_audit_log row +// - TestVault_Versioning : rotate creates v2; v1 still queryable via ?version=1 +// - TestVault_DeleteSemantics : DELETE removes ALL versions (hard delete) and is idempotent +// - TestVault_E2E_KeyList : list returns keys but never values +// +// Integration tests skip when TEST_DATABASE_URL is empty (no DB available). + +import ( + "bytes" + "context" + "database/sql" + "encoding/base64" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "os" + "testing" + + "github.com/gofiber/fiber/v2" + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "instant.dev/internal/config" + "instant.dev/internal/crypto" + "instant.dev/internal/handlers" + "instant.dev/internal/middleware" + "instant.dev/internal/models" + "instant.dev/internal/plans" + "instant.dev/internal/testhelpers" +) + +// vaultMigration mirrors db/migrations/008_vault.sql; embedded inline so the +// test does not depend on testhelpers.runMigrations being updated. Idempotent +// (IF NOT EXISTS) so safe to run on every test setup. +const vaultMigration = ` +CREATE TABLE IF NOT EXISTS vault_secrets ( + id UUID PRIMARY KEY DEFAULT gen_random_uuid(), + team_id UUID NOT NULL REFERENCES teams(id) ON DELETE CASCADE, + env TEXT NOT NULL DEFAULT 'production', + key TEXT NOT NULL, + encrypted_value BYTEA NOT NULL, + version INT NOT NULL DEFAULT 1, + created_by UUID REFERENCES users(id) ON DELETE SET NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), + UNIQUE (team_id, env, key, version) +); +CREATE INDEX IF NOT EXISTS idx_vault_secrets_lookup ON vault_secrets (team_id, env, key); +CREATE TABLE IF NOT EXISTS vault_audit_log ( + id BIGSERIAL PRIMARY KEY, + team_id UUID NOT NULL, + user_id UUID, + action TEXT NOT NULL, + env TEXT NOT NULL, + secret_key TEXT NOT NULL, + ip TEXT, + ts TIMESTAMPTZ NOT NULL DEFAULT now() +); +CREATE INDEX IF NOT EXISTS idx_vault_audit_team_ts ON vault_audit_log (team_id, ts DESC); +` + +// applyVaultMigration ensures the vault schema exists in the test DB. +func applyVaultMigration(t *testing.T, db *sql.DB) { + t.Helper() + if _, err := db.Exec(vaultMigration); err != nil { + t.Fatalf("applyVaultMigration: %v", err) + } +} + +// vaultIntegrationDB returns a test DB and cleanup, or skips when none configured +// or when the DB is unreachable. Integration tests must skip cleanly in CI when +// no postgres is running — never fatal. +func vaultIntegrationDB(t *testing.T) (*sql.DB, func()) { + t.Helper() + dsn := os.Getenv("TEST_DATABASE_URL") + if dsn == "" { + t.Skip("TEST_DATABASE_URL not set — skipping integration test") + } + // Probe the connection ourselves so a refused/auth-failed connection skips + // rather than fataling out via testhelpers.SetupTestDB. + probe, err := sql.Open("postgres", dsn) + if err != nil { + t.Skipf("integration DB open failed: %v", err) + } + if err := probe.Ping(); err != nil { + probe.Close() + t.Skipf("integration DB ping failed (no test postgres available): %v", err) + } + probe.Close() + + db, clean := testhelpers.SetupTestDB(t) + applyVaultMigration(t, db) + return db, clean +} + +// vaultTestApp builds a minimal Fiber app exposing only the vault routes. +// Auth is gated by RequireAuth using the standard test JWT secret. +func vaultTestApp(t *testing.T, db *sql.DB) *fiber.App { + t.Helper() + cfg := &config.Config{ + JWTSecret: testhelpers.TestJWTSecret, + AESKey: testhelpers.TestAESKeyHex, + } + app := fiber.New(fiber.Config{ + ErrorHandler: func(c *fiber.Ctx, err error) error { + code := fiber.StatusInternalServerError + if e, ok := err.(*fiber.Error); ok { + code = e.Code + } + return c.Status(code).JSON(fiber.Map{"ok": false, "error": "internal_error", "message": err.Error()}) + }, + }) + app.Use(middleware.RequestID()) + h := handlers.NewVaultHandler(db, cfg, plans.Default()) + api := app.Group("/api/v1", middleware.RequireAuth(cfg)) + api.Put("/vault/:env/:key", h.PutSecret) + api.Get("/vault/:env/:key", h.GetSecret) + api.Get("/vault/:env", h.ListKeys) + api.Delete("/vault/:env/:key", h.DeleteSecret) + api.Post("/vault/:env/:key/rotate", h.RotateSecret) + return app +} + +// jsonReq builds a JSON request with the given JWT. +func jsonReq(t *testing.T, method, path, jwt string, body any) *http.Request { + t.Helper() + var buf bytes.Buffer + if body != nil { + require.NoError(t, json.NewEncoder(&buf).Encode(body)) + } + req := httptest.NewRequest(method, path, &buf) + if body != nil { + req.Header.Set("Content-Type", "application/json") + } + if jwt != "" { + req.Header.Set("Authorization", "Bearer "+jwt) + } + return req +} + +// makeTeamUser inserts a team and one user, and returns (teamID, userID, jwt). +func makeTeamUser(t *testing.T, db *sql.DB) (string, string, string) { + t.Helper() + teamID := testhelpers.MustCreateTeamDB(t, db, "hobby") + emailAddr := testhelpers.UniqueEmail(t) + var userID string + require.NoError(t, db.QueryRow( + `INSERT INTO users (team_id, email) VALUES ($1::uuid, $2) RETURNING id`, + teamID, emailAddr, + ).Scan(&userID)) + jwt := testhelpers.MustSignSessionJWT(t, userID, teamID, emailAddr) + return teamID, userID, jwt +} + +// ── 1. AES roundtrip + tamper detection ────────────────────────────────────── + +func TestVault_AESRoundtrip(t *testing.T) { + keyHex := testhelpers.TestAESKeyHex + key, err := crypto.ParseAESKey(keyHex) + require.NoError(t, err) + + plaintext := "supersecret-postgres://user:pass@host/db" + encoded, err := crypto.Encrypt(key, plaintext) + require.NoError(t, err) + + raw, err := base64.URLEncoding.DecodeString(encoded) + require.NoError(t, err) + assert.Greater(t, len(raw), len(plaintext), "ciphertext must include nonce + tag overhead") + + // Roundtrip: re-encode and decrypt. + got, err := crypto.Decrypt(key, base64.URLEncoding.EncodeToString(raw)) + require.NoError(t, err) + assert.Equal(t, plaintext, got) + + // Tamper: flip a byte in the middle. GCM auth tag must reject. + tampered := make([]byte, len(raw)) + copy(tampered, raw) + tampered[len(tampered)/2] ^= 0xFF + _, err = crypto.Decrypt(key, base64.URLEncoding.EncodeToString(tampered)) + assert.Error(t, err, "tampered ciphertext must fail GCM auth") + + // Wrong key: decryption must fail. + otherKey, _ := crypto.ParseAESKey("ffeeddccbbaa00112233445566778899aabbccddeeff00112233445566778899") + _, err = crypto.Decrypt(otherKey, encoded) + assert.Error(t, err, "wrong AES key must fail decryption") +} + +// ── 2. Cross-team isolation: foreign reads return 404, never 403 ───────────── + +func TestVault_TeamIsolation(t *testing.T) { + db, clean := vaultIntegrationDB(t) + defer clean() + app := vaultTestApp(t, db) + + _, _, jwtA := makeTeamUser(t, db) + _, _, jwtB := makeTeamUser(t, db) + + const env, key = "production", "DATABASE_URL" + + // Team A writes a secret. + resp, err := app.Test(jsonReq(t, http.MethodPut, "/api/v1/vault/"+env+"/"+key, jwtA, map[string]string{"value": "team-a-secret"}), 5000) + require.NoError(t, err) + require.Equal(t, http.StatusCreated, resp.StatusCode) + resp.Body.Close() + + // Team B GET → must be 404 (never 403, never 200). + resp, err = app.Test(jsonReq(t, http.MethodGet, "/api/v1/vault/"+env+"/"+key, jwtB, nil), 5000) + require.NoError(t, err) + defer resp.Body.Close() + assert.Equal(t, http.StatusNotFound, resp.StatusCode, "cross-team read must return 404") + + // Team B DELETE → must also be 404. + resp2, err := app.Test(jsonReq(t, http.MethodDelete, "/api/v1/vault/"+env+"/"+key, jwtB, nil), 5000) + require.NoError(t, err) + defer resp2.Body.Close() + assert.Equal(t, http.StatusNotFound, resp2.StatusCode, "cross-team delete must return 404") + + // Team B LIST → must be empty (no leak via the list endpoint). + resp3, err := app.Test(jsonReq(t, http.MethodGet, "/api/v1/vault/"+env, jwtB, nil), 5000) + require.NoError(t, err) + defer resp3.Body.Close() + require.Equal(t, http.StatusOK, resp3.StatusCode) + var lb struct { + Keys []string `json:"keys"` + } + require.NoError(t, json.NewDecoder(resp3.Body).Decode(&lb)) + assert.Empty(t, lb.Keys, "team B must not see team A's keys") + + // Sanity: team A still sees its key. + resp4, err := app.Test(jsonReq(t, http.MethodGet, "/api/v1/vault/"+env+"/"+key, jwtA, nil), 5000) + require.NoError(t, err) + defer resp4.Body.Close() + assert.Equal(t, http.StatusOK, resp4.StatusCode) +} + +// ── 3. Audit log: every mutation + read writes one row ─────────────────────── + +func TestVault_AuditLog(t *testing.T) { + db, clean := vaultIntegrationDB(t) + defer clean() + app := vaultTestApp(t, db) + + teamIDStr, _, jwt := makeTeamUser(t, db) + teamID := uuid.MustParse(teamIDStr) + // Use production env: tier-restricted envs are validated separately in + // TestVault_TierEnvRestriction. Hobby tier (the default for makeTeamUser) + // only permits "production". + const env, key = "production", "API_TOKEN" + + // PUT + resp, err := app.Test(jsonReq(t, http.MethodPut, "/api/v1/vault/"+env+"/"+key, jwt, map[string]string{"value": "v1"}), 5000) + require.NoError(t, err) + resp.Body.Close() + + n, err := models.CountVaultAudit(context.Background(), db, teamID, "set", env, key) + require.NoError(t, err) + assert.Equal(t, 1, n, "PUT must write one 'set' audit row") + + // GET + resp, err = app.Test(jsonReq(t, http.MethodGet, "/api/v1/vault/"+env+"/"+key, jwt, nil), 5000) + require.NoError(t, err) + resp.Body.Close() + + n, err = models.CountVaultAudit(context.Background(), db, teamID, "get", env, key) + require.NoError(t, err) + assert.Equal(t, 1, n, "GET must write one 'get' audit row") + + // DELETE + resp, err = app.Test(jsonReq(t, http.MethodDelete, "/api/v1/vault/"+env+"/"+key, jwt, nil), 5000) + require.NoError(t, err) + resp.Body.Close() + + n, err = models.CountVaultAudit(context.Background(), db, teamID, "delete", env, key) + require.NoError(t, err) + assert.Equal(t, 1, n, "DELETE must write one 'delete' audit row") +} + +// ── 4. Versioning: rotate creates v2; v1 still queryable ───────────────────── + +func TestVault_Versioning(t *testing.T) { + db, clean := vaultIntegrationDB(t) + defer clean() + app := vaultTestApp(t, db) + + _, _, jwt := makeTeamUser(t, db) + const env, key = "production", "OPENAI_KEY" + + // PUT v1 + resp, err := app.Test(jsonReq(t, http.MethodPut, "/api/v1/vault/"+env+"/"+key, jwt, map[string]string{"value": "sk-v1"}), 5000) + require.NoError(t, err) + defer resp.Body.Close() + require.Equal(t, http.StatusCreated, resp.StatusCode) + var b1 struct{ Version int `json:"version"` } + require.NoError(t, json.NewDecoder(resp.Body).Decode(&b1)) + assert.Equal(t, 1, b1.Version) + + // Rotate → v2 + resp2, err := app.Test(jsonReq(t, http.MethodPost, "/api/v1/vault/"+env+"/"+key+"/rotate", jwt, map[string]string{"value": "sk-v2"}), 5000) + require.NoError(t, err) + defer resp2.Body.Close() + require.Equal(t, http.StatusCreated, resp2.StatusCode) + var b2 struct{ Version int `json:"version"` } + require.NoError(t, json.NewDecoder(resp2.Body).Decode(&b2)) + assert.Equal(t, 2, b2.Version, "rotate must produce v2") + + // GET (latest) → must return v2 value + resp3, err := app.Test(jsonReq(t, http.MethodGet, "/api/v1/vault/"+env+"/"+key, jwt, nil), 5000) + require.NoError(t, err) + defer resp3.Body.Close() + require.Equal(t, http.StatusOK, resp3.StatusCode) + var b3 struct { + Value string `json:"value"` + Version int `json:"version"` + } + require.NoError(t, json.NewDecoder(resp3.Body).Decode(&b3)) + assert.Equal(t, "sk-v2", b3.Value) + assert.Equal(t, 2, b3.Version) + + // GET ?version=1 → must return v1 value (history queryable) + resp4, err := app.Test(jsonReq(t, http.MethodGet, "/api/v1/vault/"+env+"/"+key+"?version=1", jwt, nil), 5000) + require.NoError(t, err) + defer resp4.Body.Close() + require.Equal(t, http.StatusOK, resp4.StatusCode) + var b4 struct { + Value string `json:"value"` + Version int `json:"version"` + } + require.NoError(t, json.NewDecoder(resp4.Body).Decode(&b4)) + assert.Equal(t, "sk-v1", b4.Value) + assert.Equal(t, 1, b4.Version) + + // GET ?version=99 → 404 + resp5, err := app.Test(jsonReq(t, http.MethodGet, "/api/v1/vault/"+env+"/"+key+"?version=99", jwt, nil), 5000) + require.NoError(t, err) + defer resp5.Body.Close() + assert.Equal(t, http.StatusNotFound, resp5.StatusCode) +} + +// ── 5. Delete semantics: hard delete of all versions, idempotent on missing ── + +func TestVault_DeleteSemantics(t *testing.T) { + db, clean := vaultIntegrationDB(t) + defer clean() + app := vaultTestApp(t, db) + + teamIDStr, _, jwt := makeTeamUser(t, db) + teamID := uuid.MustParse(teamIDStr) + const env, key = "production", "DOC_DELETE" + + // Create v1 + v2. + for _, v := range []string{"a", "b"} { + resp, err := app.Test(jsonReq(t, http.MethodPut, "/api/v1/vault/"+env+"/"+key, jwt, map[string]string{"value": v}), 5000) + require.NoError(t, err) + resp.Body.Close() + } + + // Confirm 2 rows exist. + var pre int + require.NoError(t, db.QueryRow(`SELECT COUNT(*) FROM vault_secrets WHERE team_id = $1::uuid AND env = $2 AND key = $3`, teamID, env, key).Scan(&pre)) + assert.Equal(t, 2, pre) + + // DELETE → 204 + resp, err := app.Test(jsonReq(t, http.MethodDelete, "/api/v1/vault/"+env+"/"+key, jwt, nil), 5000) + require.NoError(t, err) + defer resp.Body.Close() + assert.Equal(t, http.StatusNoContent, resp.StatusCode) + + // Both versions are gone (hard delete). + var post int + require.NoError(t, db.QueryRow(`SELECT COUNT(*) FROM vault_secrets WHERE team_id = $1::uuid AND env = $2 AND key = $3`, teamID, env, key).Scan(&post)) + assert.Equal(t, 0, post, "DELETE must hard-remove every version (chosen MVP semantics)") + + // GET after delete → 404 for latest AND for ?version=1 + resp2, err := app.Test(jsonReq(t, http.MethodGet, "/api/v1/vault/"+env+"/"+key, jwt, nil), 5000) + require.NoError(t, err) + defer resp2.Body.Close() + assert.Equal(t, http.StatusNotFound, resp2.StatusCode) + + resp3, err := app.Test(jsonReq(t, http.MethodGet, "/api/v1/vault/"+env+"/"+key+"?version=1", jwt, nil), 5000) + require.NoError(t, err) + defer resp3.Body.Close() + assert.Equal(t, http.StatusNotFound, resp3.StatusCode) + + // Second DELETE → 404 (idempotent, never leaks "this never existed" vs. "we just deleted it") + resp4, err := app.Test(jsonReq(t, http.MethodDelete, "/api/v1/vault/"+env+"/"+key, jwt, nil), 5000) + require.NoError(t, err) + defer resp4.Body.Close() + assert.Equal(t, http.StatusNotFound, resp4.StatusCode) +} + +// ── 6. Key list returns key names but never values ─────────────────────────── + +func TestVault_E2E_KeyList(t *testing.T) { + db, clean := vaultIntegrationDB(t) + defer clean() + app := vaultTestApp(t, db) + + _, _, jwt := makeTeamUser(t, db) + const env = "production" + + // Insert three keys with distinct values that must NEVER appear in the list response. + for _, kv := range [][2]string{ + {"DB_URL", "value-must-not-leak-1"}, + {"REDIS_URL", "value-must-not-leak-2"}, + {"API_TOKEN", "value-must-not-leak-3"}, + } { + resp, err := app.Test(jsonReq(t, http.MethodPut, "/api/v1/vault/"+env+"/"+kv[0], jwt, map[string]string{"value": kv[1]}), 5000) + require.NoError(t, err) + resp.Body.Close() + } + + // GET list + resp, err := app.Test(jsonReq(t, http.MethodGet, "/api/v1/vault/"+env, jwt, nil), 5000) + require.NoError(t, err) + defer resp.Body.Close() + require.Equal(t, http.StatusOK, resp.StatusCode) + + rawBody, err := readAll(resp.Body) + require.NoError(t, err) + + var lb struct { + OK bool `json:"ok"` + Env string `json:"env"` + Keys []string `json:"keys"` + } + require.NoError(t, json.Unmarshal(rawBody, &lb)) + assert.True(t, lb.OK) + assert.Equal(t, env, lb.Env) + assert.ElementsMatch(t, []string{"DB_URL", "REDIS_URL", "API_TOKEN"}, lb.Keys) + + // Body must NOT contain any plaintext value. + for _, leak := range []string{"value-must-not-leak-1", "value-must-not-leak-2", "value-must-not-leak-3"} { + assert.NotContains(t, string(rawBody), leak, + "list response must never include plaintext values (leak=%s)", leak) + } +} + +// ── 7. Auth gate: missing JWT yields 401 (not 404) so external callers know auth is required ── + +func TestVault_RequiresAuth(t *testing.T) { + db, clean := vaultIntegrationDB(t) + defer clean() + app := vaultTestApp(t, db) + + resp, err := app.Test(jsonReq(t, http.MethodGet, "/api/v1/vault/production/SOMETHING", "", nil), 5000) + require.NoError(t, err) + defer resp.Body.Close() + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) +} + +// ── 8. Invalid env / key validation ────────────────────────────────────────── + +func TestVault_Validation(t *testing.T) { + db, clean := vaultIntegrationDB(t) + defer clean() + app := vaultTestApp(t, db) + _, _, jwt := makeTeamUser(t, db) + + cases := []struct { + name string + path string + want int + }{ + // Path params can't be empty in fiber routes; use illegal characters instead. + {"bad-key-with-slash", "/api/v1/vault/production/foo bar", http.StatusBadRequest}, + {"bad-key-too-long", "/api/v1/vault/production/" + longString(300), http.StatusBadRequest}, + {"bad-env-with-special", "/api/v1/vault/prod!ction/X", http.StatusBadRequest}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + resp, err := app.Test(jsonReq(t, http.MethodPut, tc.path, jwt, map[string]string{"value": "x"}), 5000) + require.NoError(t, err) + defer resp.Body.Close() + // Some illegal chars (e.g. space) get URL-encoded by httptest into %20 which is also rejected; + // we just assert non-2xx + non-5xx. + assert.True(t, resp.StatusCode == tc.want || resp.StatusCode == http.StatusNotFound, + "expected %d (got %d) for path=%s", tc.want, resp.StatusCode, tc.path) + }) + } +} + +// ── 9. Per-tier vault quota + env restriction ──────────────────────────────── +// +// Hobby tier (default for makeTeamUser): vault_max_entries=20, +// vault_envs_allowed=["production"]. Verifies: +// - 20 distinct keys succeed +// - 21st key returns 402 vault_quota_exceeded +// - rotating an existing key after the cap still works (count doesn't grow) +// - PUT to a non-allowed env returns 403 vault_env_not_allowed +func TestVault_TierQuotaAndEnv(t *testing.T) { + db, clean := vaultIntegrationDB(t) + defer clean() + app := vaultTestApp(t, db) + + _, _, jwt := makeTeamUser(t, db) // hobby tier + + // 20 PUTs on production should succeed. + for i := 0; i < 20; i++ { + path := fmt.Sprintf("/api/v1/vault/production/KEY_%02d", i) + resp, err := app.Test(jsonReq(t, http.MethodPut, path, jwt, map[string]string{"value": "v"}), 5000) + require.NoError(t, err) + body, _ := readAll(resp.Body) + resp.Body.Close() + require.Equalf(t, http.StatusCreated, resp.StatusCode, + "PUT %d/20 expected 201, got %d body=%s", i+1, resp.StatusCode, string(body)) + } + + // 21st distinct key → 402 vault_quota_exceeded. + resp, err := app.Test(jsonReq(t, http.MethodPut, "/api/v1/vault/production/KEY_21", jwt, map[string]string{"value": "v"}), 5000) + require.NoError(t, err) + defer resp.Body.Close() + body, _ := readAll(resp.Body) + assert.Equal(t, http.StatusPaymentRequired, resp.StatusCode, + "21st key must return 402; got %d body=%s", resp.StatusCode, string(body)) + var errResp struct { + Error string `json:"error"` + } + _ = json.Unmarshal(body, &errResp) + assert.Equal(t, "vault_quota_exceeded", errResp.Error) + + // Updating an existing key (KEY_00) must still succeed — no quota burn. + resp2, err := app.Test(jsonReq(t, http.MethodPut, "/api/v1/vault/production/KEY_00", jwt, map[string]string{"value": "v2"}), 5000) + require.NoError(t, err) + defer resp2.Body.Close() + assert.Equal(t, http.StatusCreated, resp2.StatusCode, + "updating an existing key when at quota must still succeed (no count growth)") + + // PUT to non-allowed env → 403 vault_env_not_allowed. + resp3, err := app.Test(jsonReq(t, http.MethodPut, "/api/v1/vault/staging/SOMETHING", jwt, map[string]string{"value": "v"}), 5000) + require.NoError(t, err) + defer resp3.Body.Close() + body3, _ := readAll(resp3.Body) + assert.Equal(t, http.StatusForbidden, resp3.StatusCode, + "hobby tier PUT to staging must return 403; got %d body=%s", resp3.StatusCode, string(body3)) + var errResp3 struct { + Error string `json:"error"` + } + _ = json.Unmarshal(body3, &errResp3) + assert.Equal(t, "vault_env_not_allowed", errResp3.Error) +} + +func longString(n int) string { + s := "" + for i := 0; i < n; i++ { + s += "a" + } + return s +} + +// readAll is a small helper so we can introspect the raw body for leak checks. +func readAll(r interface{ Read(p []byte) (int, error) }) ([]byte, error) { + buf := make([]byte, 0, 4096) + tmp := make([]byte, 4096) + for { + n, err := r.Read(tmp) + if n > 0 { + buf = append(buf, tmp[:n]...) + } + if err != nil { + if err.Error() == "EOF" { + return buf, nil + } + return buf, nil // tolerate; fiber test bodies sometimes return non-io.EOF + } + } +} + +// Sanity: ensure fmt remains imported even if a debug Sprintf is removed. +var _ = fmt.Sprint diff --git a/internal/handlers/webhook.go b/internal/handlers/webhook.go index 87203e9b..bb06a8da 100644 --- a/internal/handlers/webhook.go +++ b/internal/handlers/webhook.go @@ -124,9 +124,14 @@ func (h *WebhookHandler) NewWebhook(c *fiber.Ctx) error { _ = c.BodyParser(&body) body.Name = sanitizeName(body.Name) + env, envErr := resolveEnv(c, body.Env) + if envErr != nil { + return envErr + } + // ── Authenticated path ─────────────────────────────────────────────────────── if teamIDStr := middleware.GetTeamID(c); teamIDStr != "" { - return h.newWebhookAuthenticated(c, teamIDStr, fp, country, vendor, requestID, body.Name, start) + return h.newWebhookAuthenticated(c, teamIDStr, fp, country, vendor, requestID, body.Name, env, start) } // ── Anonymous path ─────────────────────────────────────────────────────────── @@ -162,6 +167,7 @@ func (h *WebhookHandler) NewWebhook(c *fiber.Ctx) error { "token": existing.Token.String(), "receive_url": url, "tier": existing.Tier, + "env": existing.Env, "limits": webhookAnonLimits(), "note": limitExceededNote(upgradeURL, existing.ExpiresAt.Time), "upgrade": upgradeURL, @@ -180,6 +186,7 @@ func (h *WebhookHandler) NewWebhook(c *fiber.Ctx) error { ResourceType: "webhook", Name: body.Name, Tier: "anonymous", + Env: env, Fingerprint: fp, CloudVendor: vendor, CountryCode: country, @@ -236,6 +243,7 @@ func (h *WebhookHandler) NewWebhook(c *fiber.Ctx) error { "token": tokenStr, "receive_url": rURL, "tier": "anonymous", + "env": resource.Env, "limits": webhookAnonLimits(), "note": upgradeNote(upgradeURL), "expires_at": expiresAt, @@ -244,7 +252,7 @@ func (h *WebhookHandler) NewWebhook(c *fiber.Ctx) error { // newWebhookAuthenticated handles the authenticated path for POST /webhook/new. func (h *WebhookHandler) newWebhookAuthenticated( - c *fiber.Ctx, teamIDStr, fp, country, vendor, requestID, name string, start time.Time, + c *fiber.Ctx, teamIDStr, fp, country, vendor, requestID, name string, env string, start time.Time, ) error { ctx := c.UserContext() teamUUID, err := parseTeamID(teamIDStr) @@ -262,6 +270,7 @@ func (h *WebhookHandler) newWebhookAuthenticated( ResourceType: "webhook", Name: name, Tier: team.PlanTier, + Env: env, Fingerprint: fp, CloudVendor: vendor, CountryCode: country, @@ -274,6 +283,18 @@ func (h *WebhookHandler) newWebhookAuthenticated( return respondError(c, fiber.StatusServiceUnavailable, "provision_failed", "Failed to provision webhook resource") } + // Best-effort audit event; failures must never block the provision. + go func() { + _ = models.InsertAuditEvent(context.Background(), h.db, models.AuditEvent{ + TeamID: teamUUID, + Actor: "agent", + Kind: "provision", + ResourceType: "webhook", + ResourceID: uuid.NullUUID{UUID: resource.ID, Valid: true}, + Summary: "agent provisioned <strong>webhook</strong> <code>" + resource.Token.String()[:8] + "</code>", + }) + }() + tokenStr := resource.Token.String() rURL := receiveURL(c.BaseURL(), tokenStr) @@ -300,6 +321,7 @@ func (h *WebhookHandler) newWebhookAuthenticated( "token": tokenStr, "receive_url": rURL, "tier": team.PlanTier, + "env": resource.Env, "limits": fiber.Map{ "requests_stored": h.webhookMaxStored(team.PlanTier), }, diff --git a/internal/handlers/wellknown.go b/internal/handlers/wellknown.go new file mode 100644 index 00000000..19d162f1 --- /dev/null +++ b/internal/handlers/wellknown.go @@ -0,0 +1,90 @@ +package handlers + +// wellknown.go — agent-auth discovery endpoint. +// +// Implements the MCP Authorization profile resource-server metadata document +// (https://modelcontextprotocol.io/specification/draft/basic/authorization). +// +// MCP-compliant agents fetch this endpoint before calling any protected route +// to discover: +// - the canonical resource URL (used for RFC 8707 audience checks) +// - the authorization server(s) that may issue tokens for this resource +// - which transports for the bearer token are supported +// - human-readable documentation +// +// The endpoint is unauthenticated by design — discovery must work for any +// caller that has not yet acquired a token. + +import ( + "net/url" + "os" + "strings" + + "github.com/gofiber/fiber/v2" +) + +// Default canonical resource URL when neither API_PUBLIC_URL nor a request host +// is available. Kept as a const so the spec output is stable in tests. +const defaultCanonicalResourceURL = "https://api.instanode.dev" + +// wellKnownDocPath is the public docs URL exposed in the metadata. +const wellKnownDocPath = "/docs/auth" + +// CanonicalResourceURL returns the canonical resource URL used for RFC 8707 +// audience checks and for `/.well-known/oauth-protected-resource`. +// +// Resolution order: +// 1. API_PUBLIC_URL environment variable (when set and non-empty) +// 2. The X-Forwarded-Proto + Host headers from the live request +// 3. The constant default ("https://api.instanode.dev") +// +// It is a package-level variable (rather than a plain function) so individual +// tests can override it without forcing the rest of the codebase to thread a +// dependency through call sites. +var CanonicalResourceURL = func(c *fiber.Ctx) string { + if v := strings.TrimRight(os.Getenv("API_PUBLIC_URL"), "/"); v != "" { + return v + } + if c != nil { + host := c.Get("X-Forwarded-Host") + if host == "" { + host = c.Hostname() + } + scheme := c.Get("X-Forwarded-Proto") + if scheme == "" { + if c.Protocol() != "" { + scheme = c.Protocol() + } else { + scheme = "https" + } + } + if host != "" { + u := url.URL{Scheme: scheme, Host: host} + return strings.TrimRight(u.String(), "/") + } + } + return defaultCanonicalResourceURL +} + +// ServeOAuthProtectedResourceMetadata serves +// GET /.well-known/oauth-protected-resource per the MCP authorization profile. +// +// Response shape (RFC 9728 / MCP draft): +// +// { +// "resource": "https://api.instanode.dev", +// "authorization_servers": ["https://api.instanode.dev"], +// "bearer_methods_supported": ["header"], +// "resource_documentation": "https://instanode.dev/docs/auth" +// } +func ServeOAuthProtectedResourceMetadata(c *fiber.Ctx) error { + resource := CanonicalResourceURL(c) + c.Set("Content-Type", "application/json; charset=utf-8") + c.Set("Cache-Control", "public, max-age=300") + return c.JSON(fiber.Map{ + "resource": resource, + "authorization_servers": []string{resource}, + "bearer_methods_supported": []string{"header"}, + "resource_documentation": "https://instanode.dev" + wellKnownDocPath, + }) +} diff --git a/internal/handlers/wellknown_test.go b/internal/handlers/wellknown_test.go new file mode 100644 index 00000000..041fe4e8 --- /dev/null +++ b/internal/handlers/wellknown_test.go @@ -0,0 +1,75 @@ +package handlers_test + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/gofiber/fiber/v2" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "instant.dev/internal/handlers" + "instant.dev/internal/testhelpers" +) + +// TestWellKnown_Spec asserts that GET /.well-known/oauth-protected-resource +// returns a JSON document conforming to the MCP authorization profile. +// +// Required fields per the MCP draft (mirrors RFC 9728): +// - resource (string) +// - authorization_servers ([]string) +// - bearer_methods_supported ([]string, must include "header") +// - resource_documentation (string) +func TestWellKnown_Spec(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "https://api.instanode.dev") + + app := fiber.New() + app.Get("/.well-known/oauth-protected-resource", handlers.ServeOAuthProtectedResourceMetadata) + + req := httptest.NewRequest(http.MethodGet, "/.well-known/oauth-protected-resource", nil) + resp, err := app.Test(req, 1000) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode) + assert.Contains(t, resp.Header.Get("Content-Type"), "application/json") + + var body map[string]any + testhelpers.DecodeJSON(t, resp, &body) + + assert.Equal(t, "https://api.instanode.dev", body["resource"]) + + servers, ok := body["authorization_servers"].([]any) + require.True(t, ok, "authorization_servers must be an array") + require.Len(t, servers, 1) + assert.Equal(t, "https://api.instanode.dev", servers[0]) + + methods, ok := body["bearer_methods_supported"].([]any) + require.True(t, ok, "bearer_methods_supported must be an array") + assert.Contains(t, methods, "header") + + assert.Equal(t, "https://instanode.dev/docs/auth", body["resource_documentation"]) +} + +// TestWellKnown_FallsBackToRequestHost verifies that when API_PUBLIC_URL is unset +// the canonical URL is derived from the live request (Host header + scheme). +func TestWellKnown_FallsBackToRequestHost(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "") + + app := fiber.New() + app.Get("/.well-known/oauth-protected-resource", handlers.ServeOAuthProtectedResourceMetadata) + + req := httptest.NewRequest(http.MethodGet, "/.well-known/oauth-protected-resource", nil) + req.Host = "api.example.test" + req.Header.Set("X-Forwarded-Proto", "https") + resp, err := app.Test(req, 1000) + require.NoError(t, err) + defer resp.Body.Close() + + var body map[string]any + testhelpers.DecodeJSON(t, resp, &body) + + resource, _ := body["resource"].(string) + assert.Equal(t, "https://api.example.test", resource) +} diff --git a/internal/middleware/api_key.go b/internal/middleware/api_key.go new file mode 100644 index 00000000..dca0be30 --- /dev/null +++ b/internal/middleware/api_key.go @@ -0,0 +1,104 @@ +package middleware + +import ( + "context" + "database/sql" + "errors" + "log/slog" + "strings" + "sync" + "time" + + "github.com/gofiber/fiber/v2" + "instant.dev/internal/models" +) + +// LocalKeyAPIKey marks requests authenticated via Personal Access Token rather +// than session JWT. Handlers can branch on this for stricter scope checks. +const LocalKeyAPIKey = "auth_api_key" + +// LocalKeyAPIKeyScopes carries the scopes granted to the PAT so handlers can +// gate fine-grained operations (e.g., admin actions require "admin" scope). +const LocalKeyAPIKeyScopes = "auth_api_key_scopes" + +// apiKeyDB is the platform DB handle used by the PAT branch of RequireAuth. +// Set via SetAPIKeyDB at startup. nil → PATs are rejected silently. +var ( + apiKeyDBMu sync.RWMutex + apiKeyDB *sql.DB +) + +// SetAPIKeyDB registers the DB handle for PAT lookup. +func SetAPIKeyDB(db *sql.DB) { + apiKeyDBMu.Lock() + defer apiKeyDBMu.Unlock() + apiKeyDB = db +} + +func getAPIKeyDB() *sql.DB { + apiKeyDBMu.RLock() + defer apiKeyDBMu.RUnlock() + return apiKeyDB +} + +// IsAPIKey reports whether the bearer token shape matches a PAT prefix. +// Cheap pattern check, never compares secrets. +func IsAPIKey(token string) bool { + return strings.HasPrefix(token, models.APIKeyPrefix) +} + +// AuthenticateAPIKey looks up the PAT by SHA-256 and populates Fiber locals +// with team_id, user_id (creator), api_key id, and scopes. Returns a +// boolean (true = authenticated, false = invalid/revoked) and the error +// from the lookup if any (errors are logged but not surfaced to clients to +// avoid leaking key existence). +func AuthenticateAPIKey(c *fiber.Ctx, plaintext string) (bool, error) { + db := getAPIKeyDB() + if db == nil { + return false, errors.New("api_key db not initialised") + } + hash := models.HashAPIKey(plaintext) + ctx, cancel := context.WithTimeout(c.UserContext(), 1500*time.Millisecond) + defer cancel() + key, err := models.GetAPIKeyByHash(ctx, db, hash) + if err != nil { + if errors.Is(err, models.ErrAPIKeyNotFound) { + return false, nil + } + slog.Warn("api_key.lookup_failed", "error", err) + return false, err + } + + c.Locals(LocalKeyTeamID, key.TeamID.String()) + if key.CreatedBy.Valid { + c.Locals(LocalKeyUserID, key.CreatedBy.UUID.String()) + } + c.Locals(LocalKeyAPIKey, key.ID.String()) + c.Locals(LocalKeyAPIKeyScopes, key.Scopes) + + // Best-effort touch — never block the request. + go func(id string) { + bgCtx, cancel := context.WithTimeout(context.Background(), 750*time.Millisecond) + defer cancel() + if err := models.TouchAPIKey(bgCtx, db, key.ID); err != nil { + slog.Debug("api_key.touch_failed", "error", err, "id", id) + } + }(key.ID.String()) + + return true, nil +} + +// GetAPIKeyScopes returns the scopes attached by AuthenticateAPIKey, or nil +// when the request was authenticated via JWT (not a PAT). +func GetAPIKeyScopes(c *fiber.Ctx) []string { + if v, ok := c.Locals(LocalKeyAPIKeyScopes).([]string); ok { + return v + } + return nil +} + +// IsAuthedViaAPIKey reports whether the request was authenticated via a PAT. +func IsAuthedViaAPIKey(c *fiber.Ctx) bool { + v, ok := c.Locals(LocalKeyAPIKey).(string) + return ok && v != "" +} diff --git a/internal/middleware/auth.go b/internal/middleware/auth.go index 569f01c9..1dbef3e4 100644 --- a/internal/middleware/auth.go +++ b/internal/middleware/auth.go @@ -2,6 +2,9 @@ package middleware import ( "errors" + "net/url" + "os" + "strings" "github.com/gofiber/fiber/v2" "github.com/golang-jwt/jwt/v4" @@ -13,13 +16,44 @@ const ( LocalKeyUserID = "auth_user_id" // LocalKeyTeamID is the fiber.Locals key for the authenticated team ID. LocalKeyTeamID = "auth_team_id" + // LocalKeyDPoPKeyThumbprint is set when the bearer token carries a DPoP + // proof-of-possession constraint (cnf.jkt). Consumed by RequireDPoP. + LocalKeyDPoPKeyThumbprint = "auth_dpop_jkt" + + // audienceMismatchError is the error keyword used when an RFC 8707 + // audience check fails. Distinct from the generic "unauthorized" so that + // agents can distinguish "wrong server" from "bad credentials". + audienceMismatchError = "invalid_token" ) +// defaultCanonicalResourceURL is the audience used when neither API_PUBLIC_URL +// nor the live request host is available. +const defaultCanonicalResourceURL = "https://api.instanode.dev" + +// confirmation captures the OAuth 2.0 PoP "cnf" claim shape (RFC 7800). +// Currently only the JWK thumbprint variant ("jkt") used by DPoP is consumed. +type confirmation struct { + JKT string `json:"jkt,omitempty"` +} + // sessionClaims mirrors the JWT payload issued by auth.go. +// +// Two extra claims back the agent-auth standards work: +// +// - Audience (`aud`) — RFC 8707 Resource Indicators. A token MUST declare +// the canonical resource URL of this API. Missing/wrong audience → 401. +// - Confirmation (`cnf`) — RFC 7800. When present and JKT is populated the +// request MUST also carry a matching DPoP proof (enforced by RequireDPoP). +// +// The audience check is OPT-IN: if the JWT carries no `aud` claim at all the +// request is allowed through (back-compat with existing dashboard tokens). +// Once a token does declare an audience it MUST match the canonical URL of +// this API; mismatched tokens are rejected. type sessionClaims struct { - UserID string `json:"uid"` - TeamID string `json:"tid"` - Email string `json:"email"` + UserID string `json:"uid"` + TeamID string `json:"tid"` + Email string `json:"email"` + Confirmation *confirmation `json:"cnf,omitempty"` jwt.RegisteredClaims } @@ -31,6 +65,70 @@ func (c sessionClaims) Valid() error { return c.RegisteredClaims.Valid() } +// CanonicalResourceURLFor returns the canonical resource URL for an incoming +// request. It is also used to populate the +// `/.well-known/oauth-protected-resource` metadata document. +// +// Resolution order: +// 1. API_PUBLIC_URL env var (when set and non-empty) +// 2. X-Forwarded-Proto + Host headers from the live request +// 3. defaultCanonicalResourceURL constant +// +// Exposed as a package-level variable so individual tests can override the +// resolution without threading a dependency through call sites. +var CanonicalResourceURLFor = func(c *fiber.Ctx) string { + if v := strings.TrimRight(os.Getenv("API_PUBLIC_URL"), "/"); v != "" { + return v + } + if c != nil { + host := c.Get("X-Forwarded-Host") + if host == "" { + host = c.Hostname() + } + scheme := c.Get("X-Forwarded-Proto") + if scheme == "" { + if p := c.Protocol(); p != "" { + scheme = p + } else { + scheme = "https" + } + } + if host != "" { + u := url.URL{Scheme: scheme, Host: host} + return strings.TrimRight(u.String(), "/") + } + } + return defaultCanonicalResourceURL +} + +// audienceMatches reports whether the JWT `aud` claim contains the canonical +// resource URL for this server. RFC 8707 §3 — the resource server MUST reject +// tokens whose audience does not include its own resource indicator. +func audienceMatches(aud jwt.ClaimStrings, canonical string) bool { + if canonical == "" { + return false + } + for _, a := range aud { + if a == canonical { + return true + } + } + return false +} + +// rejectAudienceMismatch writes an RFC 6750 §3.1-style 401 with a structured +// error keyword agents can branch on. +func rejectAudienceMismatch(c *fiber.Ctx) error { + canonical := CanonicalResourceURLFor(c) + c.Set("WWW-Authenticate", + `Bearer realm="instanode", error="invalid_token", error_description="audience mismatch", resource="`+canonical+`"`) + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "ok": false, + "error": audienceMismatchError, + "error_description": "audience mismatch", + }) +} + // RequireAuth validates the Authorization: Bearer {jwt} header. // On success it stores user_id and team_id in fiber.Locals and calls Next. // On failure it returns 401 { ok: false, error: "unauthorized" }. @@ -45,6 +143,20 @@ func RequireAuth(cfg *config.Config) fiber.Handler { } tokenStr := header[7:] + // Dispatch on token shape. PATs (ink_<base64>) hit the api_keys + // table; JWTs go through HMAC validation. Both populate the same + // auth_team_id / auth_user_id locals so handlers don't branch. + if IsAPIKey(tokenStr) { + ok, err := AuthenticateAPIKey(c, tokenStr) + if err != nil || !ok { + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "ok": false, + "error": "unauthorized", + }) + } + return c.Next() + } + claims := &sessionClaims{} parsed, err := jwt.ParseWithClaims(tokenStr, claims, func(t *jwt.Token) (interface{}, error) { if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok { @@ -66,8 +178,21 @@ func RequireAuth(cfg *config.Config) fiber.Handler { }) } + // RFC 8707 audience check — only enforced when the token actually + // declares an `aud` claim. Existing dashboard sessions issued before + // this change have no audience and continue to work; tokens that DO + // declare an audience must include the canonical resource URL. + if len(claims.Audience) > 0 { + if !audienceMatches(claims.Audience, CanonicalResourceURLFor(c)) { + return rejectAudienceMismatch(c) + } + } + c.Locals(LocalKeyUserID, claims.UserID) c.Locals(LocalKeyTeamID, claims.TeamID) + if claims.Confirmation != nil && claims.Confirmation.JKT != "" { + c.Locals(LocalKeyDPoPKeyThumbprint, claims.Confirmation.JKT) + } return c.Next() } } @@ -90,6 +215,16 @@ func GetTeamID(c *fiber.Ctx) string { return "" } +// GetDPoPKeyThumbprint returns the JWK thumbprint (`cnf.jkt`) bound to the +// current bearer token, or "" if the token is not key-bound. Consumed by +// RequireDPoP to decide whether to enforce DPoP for this request. +func GetDPoPKeyThumbprint(c *fiber.Ctx) string { + if v, ok := c.Locals(LocalKeyDPoPKeyThumbprint).(string); ok { + return v + } + return "" +} + // OptionalAuth is like RequireAuth but does not return 401 when the header is absent or invalid. // If a valid bearer token is present it populates the same Fiber locals as RequireAuth. // Use on routes where anonymous access is allowed but authenticated users get elevated behaviour. @@ -101,6 +236,12 @@ func OptionalAuth(cfg *config.Config) fiber.Handler { } tokenStr := header[7:] + // PAT path: invalid PATs continue as anonymous (do NOT block in OptionalAuth). + if IsAPIKey(tokenStr) { + _, _ = AuthenticateAPIKey(c, tokenStr) //nolint:errcheck — drop on error + return c.Next() + } + claims := &sessionClaims{} parsed, err := jwt.ParseWithClaims(tokenStr, claims, func(t *jwt.Token) (interface{}, error) { if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok { @@ -113,8 +254,18 @@ func OptionalAuth(cfg *config.Config) fiber.Handler { return c.Next() } + // RFC 8707 audience check (opt-in: only enforced if token has `aud`). + // In OptionalAuth a mismatch must NOT block the request — we just + // drop the credential and continue as anonymous. + if len(claims.Audience) > 0 && !audienceMatches(claims.Audience, CanonicalResourceURLFor(c)) { + return c.Next() + } + c.Locals(LocalKeyUserID, claims.UserID) c.Locals(LocalKeyTeamID, claims.TeamID) + if claims.Confirmation != nil && claims.Confirmation.JKT != "" { + c.Locals(LocalKeyDPoPKeyThumbprint, claims.Confirmation.JKT) + } return c.Next() } } diff --git a/internal/middleware/auth_audience_test.go b/internal/middleware/auth_audience_test.go new file mode 100644 index 00000000..6071d975 --- /dev/null +++ b/internal/middleware/auth_audience_test.go @@ -0,0 +1,145 @@ +package middleware_test + +// auth_audience_test.go — RFC 8707 Resource Indicators tests. +// +// These tests live in a separate file (rather than being added to +// auth_test.go) so they can avoid importing internal/testhelpers, which +// transitively pulls internal/handlers. Handlers currently has unrelated +// in-flight changes from other agents; keeping these tests isolated lets +// them compile without the rest of the handlers package being clean. + +import ( + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/gofiber/fiber/v2" + "github.com/golang-jwt/jwt/v4" + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "instant.dev/internal/config" + "instant.dev/internal/middleware" +) + +// audTestJWTSecret matches the inline secret used in dpop_test.go. +const audTestJWTSecret = "test-secret-that-is-at-least-32-bytes-long!!" + +// signSessionWithAudience builds a session JWT with an explicit `aud` claim. +// audience may be a single string or a comma-separated list (the JWT +// RegisteredClaims.Audience field is jwt.ClaimStrings which accepts both). +func signSessionWithAudience(t *testing.T, audience []string) string { + t.Helper() + type cnfClaim struct { + JKT string `json:"jkt,omitempty"` + } + type sessionClaims struct { + UserID string `json:"uid"` + TeamID string `json:"tid"` + Email string `json:"email"` + Cnf *cnfClaim `json:"cnf,omitempty"` + jwt.RegisteredClaims + } + c := sessionClaims{ + UserID: uuid.NewString(), + TeamID: uuid.NewString(), + Email: "user@instanode.dev", + RegisteredClaims: jwt.RegisteredClaims{ + ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)), + ID: uuid.NewString(), + Audience: jwt.ClaimStrings(audience), + }, + } + tok := jwt.NewWithClaims(jwt.SigningMethodHS256, c) + signed, err := tok.SignedString([]byte(audTestJWTSecret)) + require.NoError(t, err) + return signed +} + +func newAudApp() *fiber.App { + cfg := &config.Config{JWTSecret: audTestJWTSecret} + app := fiber.New() + app.Get("/api/v1/resources", + middleware.RequireAuth(cfg), + func(c *fiber.Ctx) error { + return c.JSON(fiber.Map{"ok": true}) + }, + ) + return app +} + +// TestAudience_Match: a token whose aud equals the canonical resource URL +// passes through. +func TestAudience_Match(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "https://api.instanode.dev") + + tok := signSessionWithAudience(t, []string{"https://api.instanode.dev"}) + + app := newAudApp() + req := httptest.NewRequest(http.MethodGet, "/api/v1/resources", nil) + req.Header.Set("Authorization", "Bearer "+tok) + resp, err := app.Test(req, 1000) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode) +} + +// TestAudience_Mismatch: a token whose aud does not contain the canonical +// resource URL is rejected with 401 invalid_token. +func TestAudience_Mismatch(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "https://api.instanode.dev") + + tok := signSessionWithAudience(t, []string{"https://storage.instanode.dev"}) + + app := newAudApp() + req := httptest.NewRequest(http.MethodGet, "/api/v1/resources", nil) + req.Header.Set("Authorization", "Bearer "+tok) + resp, err := app.Test(req, 1000) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) + assert.Contains(t, resp.Header.Get("WWW-Authenticate"), `error="invalid_token"`) + assert.Contains(t, resp.Header.Get("WWW-Authenticate"), "audience mismatch") +} + +// TestAudience_NoClaim_BackCompat: a token with no aud claim at all still +// works (back-compat for existing dashboard sessions). +func TestAudience_NoClaim_BackCompat(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "https://api.instanode.dev") + + tok := signSessionWithAudience(t, nil) + + app := newAudApp() + req := httptest.NewRequest(http.MethodGet, "/api/v1/resources", nil) + req.Header.Set("Authorization", "Bearer "+tok) + resp, err := app.Test(req, 1000) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode, + "a token with no aud claim should still pass (back-compat)") +} + +// TestAudience_MultipleAud_AnyMatch: the token may declare multiple +// audiences; at least one must match the canonical resource URL. +func TestAudience_MultipleAud_AnyMatch(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "https://api.instanode.dev") + + tok := signSessionWithAudience(t, []string{ + "https://other.example.com", + "https://api.instanode.dev", + }) + + app := newAudApp() + req := httptest.NewRequest(http.MethodGet, "/api/v1/resources", nil) + req.Header.Set("Authorization", "Bearer "+tok) + resp, err := app.Test(req, 1000) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode) +} diff --git a/internal/middleware/dpop.go b/internal/middleware/dpop.go new file mode 100644 index 00000000..229dd656 --- /dev/null +++ b/internal/middleware/dpop.go @@ -0,0 +1,274 @@ +package middleware + +// dpop.go — RFC 9449 (Demonstrating Proof of Possession) middleware. +// +// When a bearer token carries `cnf.jkt` (set by the auth middleware into +// LocalKeyDPoPKeyThumbprint) the request MUST also include a `DPoP` header +// whose proof JWT: +// +// - Has typ="dpop+jwt" in its header. +// - Carries the public key as a JWK in the header (`jwk` parameter) whose +// RFC 7638 thumbprint matches the bound jkt. +// - Has htm == request method (uppercase). +// - Has htu == request URL (no query string, no fragment). +// - Has iat within the freshness window (default 5 minutes). +// - Has a unique jti — replays are rejected via Redis-backed dedup. +// +// The middleware is OPT-IN: requests whose token does not carry cnf.jkt pass +// through unchanged. This preserves back-compat with existing dashboard JWTs +// while letting agent-issued tokens upgrade to sender-bound credentials. + +import ( + "context" + "crypto" + _ "crypto/sha256" // register sha256.New for crypto.SHA256 + "encoding/base64" + "errors" + "fmt" + "log/slog" + "net/url" + "strings" + "time" + + "github.com/gofiber/fiber/v2" + "github.com/lestrrat-go/jwx/v2/jwk" + "github.com/lestrrat-go/jwx/v2/jws" + "github.com/lestrrat-go/jwx/v2/jwt" + "github.com/redis/go-redis/v9" +) + +const ( + // dpopHeaderName is the request header that carries the proof JWT. + dpopHeaderName = "DPoP" + + // dpopFreshnessWindow caps how old the iat claim of a DPoP proof may be. + // RFC 9449 §4.3 leaves the window implementation-defined; 5 minutes + // matches the worked example in the spec. + dpopFreshnessWindow = 5 * time.Minute + + // dpopReplayKeyPrefix namespaces the Redis keys used for jti dedup. + dpopReplayKeyPrefix = "dpop:jti:" + + // dpopJWTType is the required value of the DPoP proof's typ header. + dpopJWTType = "dpop+jwt" + + // dpopErrorInvalid is the WWW-Authenticate error keyword for malformed, + // expired, or replayed proofs (RFC 9449 §7.1). + dpopErrorInvalid = "invalid_dpop_proof" +) + +// base64URLNoPad encodes b as base64url with no padding (RFC 4648 §5). +func base64URLNoPad(b []byte) string { + return base64.RawURLEncoding.EncodeToString(b) +} + +// RequireDPoP returns a Fiber handler that enforces RFC 9449 sender-binding +// for any request whose JWT carries `cnf.jkt`. Requests without that claim +// pass through. The middleware MUST be installed AFTER RequireAuth so that +// LocalKeyDPoPKeyThumbprint is populated. +// +// rdb may be nil; replay detection is then disabled (proofs are still +// signature/htm/htu/iat-validated). A warning is logged on every request in +// that case so operators notice the degraded posture. +func RequireDPoP(rdb *redis.Client) fiber.Handler { + return func(c *fiber.Ctx) error { + jkt := GetDPoPKeyThumbprint(c) + if jkt == "" { + // Token is not key-bound; DPoP is not required for this request. + return c.Next() + } + + proof := c.Get(dpopHeaderName) + if proof == "" { + return rejectDPoP(c, "missing DPoP header") + } + + if err := verifyDPoPProof(c, proof, jkt, rdb); err != nil { + slog.Info("middleware.dpop.rejected", + "error", err, + "jkt", jkt, + "path", c.Path(), + ) + return rejectDPoP(c, err.Error()) + } + + return c.Next() + } +} + +// verifyDPoPProof performs the full RFC 9449 verification chain. +// Returns nil on success or a descriptive error on failure. +func verifyDPoPProof(c *fiber.Ctx, proof, expectedJKT string, rdb *redis.Client) error { + // Parse the JWS without verification first so we can pull the embedded JWK + // out of the protected header. + parsed, err := jws.Parse([]byte(proof)) + if err != nil { + return fmt.Errorf("parse DPoP JWS: %w", err) + } + sigs := parsed.Signatures() + if len(sigs) != 1 { + return errors.New("DPoP proof must have exactly one signature") + } + hdr := sigs[0].ProtectedHeaders() + if hdr.Type() != dpopJWTType { + return fmt.Errorf("DPoP typ must be %q, got %q", dpopJWTType, hdr.Type()) + } + jwkKey := hdr.JWK() + if jwkKey == nil { + return errors.New("DPoP proof header missing jwk") + } + + // Validate jkt: the RFC 7638 thumbprint of the embedded JWK MUST equal + // the cnf.jkt the bearer token was issued for. + tp, err := jwkThumbprintBase64URL(jwkKey) + if err != nil { + return fmt.Errorf("compute thumbprint: %w", err) + } + if tp != expectedJKT { + return errors.New("DPoP key thumbprint does not match cnf.jkt") + } + + // Verify the signature using the embedded JWK. + if _, err := jws.Verify([]byte(proof), jws.WithKey(hdr.Algorithm(), jwkKey)); err != nil { + return fmt.Errorf("verify DPoP signature: %w", err) + } + + // Parse claims and check htm, htu, iat, jti. + tok, err := jwt.Parse([]byte(proof), jwt.WithVerify(false), jwt.WithValidate(false)) + if err != nil { + return fmt.Errorf("parse DPoP claims: %w", err) + } + + htm, ok := getStringClaim(tok, "htm") + if !ok { + return errors.New("DPoP missing htm claim") + } + if !strings.EqualFold(htm, c.Method()) { + return fmt.Errorf("DPoP htm %q does not match request method %q", htm, c.Method()) + } + + htu, ok := getStringClaim(tok, "htu") + if !ok { + return errors.New("DPoP missing htu claim") + } + if !urlMatches(htu, requestCanonicalURL(c)) { + return fmt.Errorf("DPoP htu %q does not match request URL %q", htu, requestCanonicalURL(c)) + } + + iat := tok.IssuedAt() + if iat.IsZero() { + return errors.New("DPoP missing iat claim") + } + now := time.Now() + skew := now.Sub(iat) + if skew < -dpopFreshnessWindow || skew > dpopFreshnessWindow { + return fmt.Errorf("DPoP iat outside freshness window (skew=%s)", skew) + } + + jti := tok.JwtID() + if jti == "" { + return errors.New("DPoP missing jti claim") + } + + // Replay protection — track jti in Redis with TTL = freshness window. + // If Redis is unavailable, log and continue (fail-open mirrors the + // rate_limit middleware: a Redis outage must not block legitimate + // agent traffic). + if rdb != nil { + ctx, cancel := context.WithTimeout(c.Context(), 250*time.Millisecond) + defer cancel() + key := dpopReplayKeyPrefix + jti + setOK, err := rdb.SetNX(ctx, key, "1", dpopFreshnessWindow).Result() + if err != nil { + slog.Warn("middleware.dpop.replay_check_failed", + "error", err, "jti", jti) + } else if !setOK { + return errors.New("DPoP jti has been seen before (replay)") + } + } else { + slog.Warn("middleware.dpop.no_redis_replay_detection_disabled") + } + + return nil +} + +// rejectDPoP writes an RFC 9449 §7.1 401 with WWW-Authenticate: DPoP and a +// matching error keyword agents can branch on. +func rejectDPoP(c *fiber.Ctx, description string) error { + c.Set("WWW-Authenticate", + fmt.Sprintf(`DPoP error="%s", error_description="%s"`, dpopErrorInvalid, description)) + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "ok": false, + "error": dpopErrorInvalid, + "error_description": description, + }) +} + +// jwkThumbprintBase64URL computes the RFC 7638 thumbprint of a JWK and +// returns it base64url-encoded (no padding) — the canonical representation +// used by RFC 9449 cnf.jkt. +func jwkThumbprintBase64URL(key jwk.Key) (string, error) { + tp, err := key.Thumbprint(crypto.SHA256) + if err != nil { + return "", err + } + return base64URLNoPad(tp), nil +} + +// requestCanonicalURL builds the htu canonical form (RFC 9449 §4.2): +// scheme://host{:port}/path with no query string and no fragment. +func requestCanonicalURL(c *fiber.Ctx) string { + host := c.Get("X-Forwarded-Host") + if host == "" { + host = c.Hostname() + } + scheme := c.Get("X-Forwarded-Proto") + if scheme == "" { + if p := c.Protocol(); p != "" { + scheme = p + } else { + scheme = "https" + } + } + u := url.URL{Scheme: scheme, Host: host, Path: c.Path()} + return u.String() +} + +// urlMatches compares two URLs ignoring case in scheme/host and ignoring +// trailing slashes. Path comparison is exact. +func urlMatches(a, b string) bool { + pa, err := url.Parse(a) + if err != nil { + return false + } + pb, err := url.Parse(b) + if err != nil { + return false + } + if !strings.EqualFold(pa.Scheme, pb.Scheme) { + return false + } + if !strings.EqualFold(pa.Host, pb.Host) { + return false + } + pathA := strings.TrimRight(pa.Path, "/") + pathB := strings.TrimRight(pb.Path, "/") + if pathA == "" { + pathA = "/" + } + if pathB == "" { + pathB = "/" + } + return pathA == pathB +} + +// getStringClaim pulls an arbitrary string-valued claim out of a parsed JWT. +// jwx exposes htm/htu only via the generic claim accessor. +func getStringClaim(tok jwt.Token, name string) (string, bool) { + v, ok := tok.Get(name) + if !ok { + return "", false + } + s, ok := v.(string) + return s, ok +} diff --git a/internal/middleware/dpop_test.go b/internal/middleware/dpop_test.go new file mode 100644 index 00000000..acbd8c50 --- /dev/null +++ b/internal/middleware/dpop_test.go @@ -0,0 +1,321 @@ +package middleware_test + +// dpop_test.go — RFC 9449 verification tests. +// +// Each test builds a DPoP-bound bearer JWT (cnf.jkt set) plus a fresh DPoP +// proof signed with the corresponding private key. The proof's claims (htm, +// htu, iat, jti) are tweaked per-test to drive each failure mode. + +import ( + "crypto" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + _ "crypto/sha256" + "encoding/base64" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/gofiber/fiber/v2" + "github.com/golang-jwt/jwt/v4" + "github.com/google/uuid" + "github.com/lestrrat-go/jwx/v2/jwa" + "github.com/lestrrat-go/jwx/v2/jwk" + "github.com/lestrrat-go/jwx/v2/jws" + jwxjwt "github.com/lestrrat-go/jwx/v2/jwt" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "instant.dev/internal/config" + "instant.dev/internal/middleware" +) + +// dpopTestJWTSecret is a 44-byte HMAC secret used by these tests. Inlined +// here rather than imported from internal/testhelpers because that package +// transitively imports internal/handlers, which currently has unrelated +// in-flight changes that would prevent middleware tests from compiling. +const dpopTestJWTSecret = "test-secret-that-is-at-least-32-bytes-long!!" + +// dpopFixture holds everything needed to drive a single DPoP test: +// the bearer JWT, the matching private key, and convenience helpers. +type dpopFixture struct { + t *testing.T + bearer string + privateKey jwk.Key + publicKey jwk.Key + thumbprint string +} + +// newDPoPFixture mints an ES256 keypair, computes its RFC 7638 thumbprint, +// and signs a session JWT whose cnf.jkt binds to that thumbprint. +func newDPoPFixture(t *testing.T) *dpopFixture { + t.Helper() + + raw, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + + priv, err := jwk.FromRaw(raw) + require.NoError(t, err) + require.NoError(t, priv.Set(jwk.AlgorithmKey, jwa.ES256)) + + pub, err := priv.PublicKey() + require.NoError(t, err) + require.NoError(t, pub.Set(jwk.AlgorithmKey, jwa.ES256)) + + tp, err := pub.Thumbprint(crypto.SHA256) + require.NoError(t, err) + thumbprint := base64.RawURLEncoding.EncodeToString(tp) + + type cnfClaim struct { + JKT string `json:"jkt"` + } + type sessionClaims struct { + UserID string `json:"uid"` + TeamID string `json:"tid"` + Email string `json:"email"` + Cnf cnfClaim `json:"cnf"` + jwt.RegisteredClaims + } + claims := sessionClaims{ + UserID: uuid.NewString(), + TeamID: uuid.NewString(), + Email: "agent@instanode.dev", + Cnf: cnfClaim{JKT: thumbprint}, + RegisteredClaims: jwt.RegisteredClaims{ + ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)), + ID: uuid.NewString(), + }, + } + tok := jwt.NewWithClaims(jwt.SigningMethodHS256, claims) + signed, err := tok.SignedString([]byte(dpopTestJWTSecret)) + require.NoError(t, err) + + return &dpopFixture{ + t: t, + bearer: signed, + privateKey: priv, + publicKey: pub, + thumbprint: thumbprint, + } +} + +// signProof builds a DPoP proof JWT with htm/htu/iat/jti and signs it with +// the fixture's private key, embedding the public key in the protected +// header (RFC 9449 §4.2: typ=dpop+jwt, alg=ES256, jwk=public-key). +func (f *dpopFixture) signProof(htm, htu string, iat time.Time, jti string) string { + f.t.Helper() + + tok := jwxjwt.New() + require.NoError(f.t, tok.Set("htm", htm)) + require.NoError(f.t, tok.Set("htu", htu)) + require.NoError(f.t, tok.Set(jwxjwt.IssuedAtKey, iat)) + require.NoError(f.t, tok.Set(jwxjwt.JwtIDKey, jti)) + + hdrs := jws.NewHeaders() + require.NoError(f.t, hdrs.Set(jws.TypeKey, "dpop+jwt")) + require.NoError(f.t, hdrs.Set(jws.JWKKey, f.publicKey)) + + signed, err := jwxjwt.Sign(tok, + jwxjwt.WithKey(jwa.ES256, f.privateKey, jws.WithProtectedHeaders(hdrs)), + ) + require.NoError(f.t, err) + return string(signed) +} + +// newDPoPApp wires RequireAuth → RequireDPoP → echo handler. Pass rdb=nil to +// disable replay detection. +func newDPoPApp(rdb *redis.Client) *fiber.App { + cfg := &config.Config{JWTSecret: dpopTestJWTSecret} + app := fiber.New() + app.Post("/db/new", + middleware.RequireAuth(cfg), + middleware.RequireDPoP(rdb), + func(c *fiber.Ctx) error { + return c.JSON(fiber.Map{"ok": true}) + }, + ) + return app +} + +// runRequest executes a single Fiber test request with optional bearer + +// DPoP headers. Returns the *http.Response for inspection. +func runRequest(t *testing.T, app *fiber.App, method, target, bearer, dpop string) *http.Response { + t.Helper() + req := httptest.NewRequest(method, target, nil) + if bearer != "" { + req.Header.Set("Authorization", "Bearer "+bearer) + } + if dpop != "" { + req.Header.Set("DPoP", dpop) + } + req.Host = "api.instanode.dev" + req.Header.Set("X-Forwarded-Proto", "https") + resp, err := app.Test(req, 1500) + require.NoError(t, err) + return resp +} + +// TestDPoP_Valid verifies a well-formed proof passes through. +func TestDPoP_Valid(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "https://api.instanode.dev") + + mr, err := miniredis.Run() + require.NoError(t, err) + defer mr.Close() + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + defer rdb.Close() + + f := newDPoPFixture(t) + proof := f.signProof("POST", "https://api.instanode.dev/db/new", time.Now(), uuid.NewString()) + + app := newDPoPApp(rdb) + resp := runRequest(t, app, http.MethodPost, "https://api.instanode.dev/db/new", f.bearer, proof) + defer resp.Body.Close() + + assert.Equal(t, http.StatusOK, resp.StatusCode) +} + +// TestDPoP_BadSig verifies a tampered proof returns 401. +func TestDPoP_BadSig(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "https://api.instanode.dev") + + f := newDPoPFixture(t) + proof := f.signProof("POST", "https://api.instanode.dev/db/new", time.Now(), uuid.NewString()) + + // Flip a byte after the second '.' (signature segment). + mangled := []byte(proof) + dotCount := 0 + for i := range mangled { + if mangled[i] == '.' { + dotCount++ + if dotCount == 2 && i+1 < len(mangled) { + if mangled[i+1] == 'A' { + mangled[i+1] = 'B' + } else { + mangled[i+1] = 'A' + } + break + } + } + } + + app := newDPoPApp(nil) + resp := runRequest(t, app, http.MethodPost, "https://api.instanode.dev/db/new", f.bearer, string(mangled)) + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) + assert.Contains(t, resp.Header.Get("WWW-Authenticate"), "DPoP") +} + +// TestDPoP_Replay verifies that the same jti reused returns 401. +func TestDPoP_Replay(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "https://api.instanode.dev") + + mr, err := miniredis.Run() + require.NoError(t, err) + defer mr.Close() + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + defer rdb.Close() + + f := newDPoPFixture(t) + app := newDPoPApp(rdb) + + jti := uuid.NewString() + proof := f.signProof("POST", "https://api.instanode.dev/db/new", time.Now(), jti) + + resp1 := runRequest(t, app, http.MethodPost, "https://api.instanode.dev/db/new", f.bearer, proof) + defer resp1.Body.Close() + require.Equal(t, http.StatusOK, resp1.StatusCode) + + resp2 := runRequest(t, app, http.MethodPost, "https://api.instanode.dev/db/new", f.bearer, proof) + defer resp2.Body.Close() + assert.Equal(t, http.StatusUnauthorized, resp2.StatusCode, + "second call with same jti must be rejected (replay)") +} + +// TestDPoP_OptIn verifies that a token without cnf.jkt does NOT require DPoP. +func TestDPoP_OptIn(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "https://api.instanode.dev") + + cfg := &config.Config{JWTSecret: dpopTestJWTSecret} + app := fiber.New() + app.Post("/db/new", + middleware.RequireAuth(cfg), + middleware.RequireDPoP(nil), + func(c *fiber.Ctx) error { + return c.JSON(fiber.Map{"ok": true}) + }, + ) + + type plainSession struct { + UserID string `json:"uid"` + TeamID string `json:"tid"` + Email string `json:"email"` + jwt.RegisteredClaims + } + tok := jwt.NewWithClaims(jwt.SigningMethodHS256, plainSession{ + UserID: uuid.NewString(), + TeamID: uuid.NewString(), + Email: "user@instanode.dev", + RegisteredClaims: jwt.RegisteredClaims{ + ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)), + ID: uuid.NewString(), + }, + }) + signed, err := tok.SignedString([]byte(dpopTestJWTSecret)) + require.NoError(t, err) + + resp := runRequest(t, app, http.MethodPost, "https://api.instanode.dev/db/new", signed, "") + defer resp.Body.Close() + assert.Equal(t, http.StatusOK, resp.StatusCode, + "a plain session JWT (no cnf.jkt) must not require a DPoP header") +} + +// TestDPoP_StaleProof verifies that a proof outside the freshness window +// is rejected. +func TestDPoP_StaleProof(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "https://api.instanode.dev") + + f := newDPoPFixture(t) + proof := f.signProof("POST", "https://api.instanode.dev/db/new", + time.Now().Add(-30*time.Minute), uuid.NewString()) + + app := newDPoPApp(nil) + resp := runRequest(t, app, http.MethodPost, "https://api.instanode.dev/db/new", f.bearer, proof) + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) +} + +// TestDPoP_WrongMethod verifies that a proof with htm != request method +// is rejected. +func TestDPoP_WrongMethod(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "https://api.instanode.dev") + + f := newDPoPFixture(t) + proof := f.signProof("GET", "https://api.instanode.dev/db/new", time.Now(), uuid.NewString()) + + app := newDPoPApp(nil) + resp := runRequest(t, app, http.MethodPost, "https://api.instanode.dev/db/new", f.bearer, proof) + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) +} + +// TestDPoP_MissingHeader verifies that when the bearer carries cnf.jkt but +// the request omits the DPoP header, the request is rejected. +func TestDPoP_MissingHeader(t *testing.T) { + t.Setenv("API_PUBLIC_URL", "https://api.instanode.dev") + + f := newDPoPFixture(t) + app := newDPoPApp(nil) + resp := runRequest(t, app, http.MethodPost, "https://api.instanode.dev/db/new", f.bearer, "") + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) + assert.Contains(t, resp.Header.Get("WWW-Authenticate"), "DPoP") +} diff --git a/internal/middleware/fingerprint.go b/internal/middleware/fingerprint.go index 613aac4d..7f5f8e93 100644 --- a/internal/middleware/fingerprint.go +++ b/internal/middleware/fingerprint.go @@ -1,7 +1,10 @@ package middleware import ( + "crypto/subtle" + "log/slog" "net" + "os" "strings" "github.com/gofiber/fiber/v2" @@ -16,14 +19,50 @@ type FingerprintConfig struct { Production bool } +// e2eTestTokenEnv is the env var holding a shared secret that, when matched +// in an X-E2E-Test-Token request header, lets the request override the +// fingerprint's source IP. This is the ONLY production-mode escape hatch and +// is intended exclusively for E2E suites running against the live cluster +// from a single dev workstation — every request from that workstation +// otherwise shares a fingerprint and hits the per-day provision cap. +// +// Operationally: set E2E_TEST_TOKEN to a 32-char hex secret in the cluster +// config; export the same value as E2E_TEST_TOKEN in the test runner. When +// both match, the LEFTMOST X-Forwarded-For entry (the one the test set) +// is used as the source IP, restoring per-test isolation. +const e2eTestTokenEnv = "E2E_TEST_TOKEN" + +// e2eTrustHeader is the request header carrying the shared secret. +const e2eTrustHeader = "X-E2E-Test-Token" + +// e2eSourceIPHeader carries the override source IP. Used instead of +// X-Forwarded-For because some reverse proxies (notably ingress-nginx with +// default use-forwarded-headers=false) overwrite XFF with the real client IP, +// dropping any test-supplied value. A custom header is passed through verbatim. +const e2eSourceIPHeader = "X-E2E-Source-IP" + // FingerprintMiddleware computes a stable per-subnet+ASN fingerprint and stores it // in Fiber locals under the key "fingerprint". It accepts a FingerprintConfig so // callers can control spoofing-prevention behaviour. func FingerprintMiddleware(cfg FingerprintConfig) fiber.Handler { return func(c *fiber.Ctx) error { var ipStr string - if cfg.Production { - // Use the rightmost entry in X-Forwarded-For — the last trusted edge hop. + + // E2E bypass: independent of cfg.Production. When the request bears a + // valid X-E2E-Test-Token matching the cluster's shared secret, the + // override source IP from X-E2E-Source-IP is used instead of the + // reverse-proxy-resolved IP. ingress-nginx defaults to overwriting + // X-Forwarded-For with the real client IP, which collapses every + // test request from one workstation onto the same fingerprint and + // trips the per-day provision cap. The dedicated header is passed + // through verbatim by every reverse proxy, sidestepping the issue. + if e2eTokenAccepted(c) { + if v := strings.TrimSpace(c.Get(e2eSourceIPHeader)); v != "" { + ipStr = v + } + } + if cfg.Production && ipStr == "" { + // Use the rightmost (last-hop) XFF entry — the trusted edge hop. xff := c.Get("X-Forwarded-For") if xff != "" { parts := strings.Split(xff, ",") @@ -61,3 +100,33 @@ func GetFingerprint(c *fiber.Ctx) string { } return "" } + +// e2eTokenAccepted reports whether the request carries a valid E2E trust +// token matching the cluster's shared secret. Returns false if the env var +// is unset (default — no bypass available). +func e2eTokenAccepted(c *fiber.Ctx) bool { + expected := os.Getenv(e2eTestTokenEnv) + if expected == "" { + return false + } + got := c.Get(e2eTrustHeader) + if got == "" { + // Debug: log headers we DO have — helps detect proxy stripping. + // Triggers only when bypass is enabled but header missing. + hdrs := []string{} + c.Request().Header.VisitAll(func(k, v []byte) { + hdrs = append(hdrs, string(k)) + }) + slog.Info("e2e_bypass.token_missing", + "have_headers", strings.Join(hdrs, ",")) + return false + } + if subtle.ConstantTimeCompare([]byte(got), []byte(expected)) == 1 { + return true + } + slog.Warn("e2e_bypass.token_mismatch", + "got_len", len(got), "expected_len", len(expected), + "got_prefix", got[:min(8, len(got))]) + return false +} + diff --git a/internal/middleware/quota.go b/internal/middleware/quota.go new file mode 100644 index 00000000..11ea26c7 --- /dev/null +++ b/internal/middleware/quota.go @@ -0,0 +1,56 @@ +package middleware + +// quota.go — HTTP-layer translation of quota errors into RFC 7231 §6.5.2 +// "402 Payment Required" responses. +// +// instanode.dev's per-resource throughput and storage quota checks live in +// internal/quota and return plain (exceeded bool, err error). This file +// gives handlers a single place to convert "quota exceeded" into the +// canonical 402 response shape, including the WWW-Authenticate: Payment +// header that future Stripe MPP integration will turn into a paywall. +// +// Today no payment is actually accepted — the response just signals which +// upgrade URL the agent should follow. The header keyword is reserved by +// the in-progress Machine Payments Protocol +// (https://stripe.com/blog/machine-payments-protocol) so when MPP ships +// this becomes a one-PR upgrade. + +import ( + "github.com/gofiber/fiber/v2" +) + +// QuotaUpgradeURL is the URL agents should follow to clear a 402. +// Plumbed as a package-level variable so tests and self-hosted operators +// can override it (e.g. point at a custom billing portal). +var QuotaUpgradeURL = "https://instanode.dev/pricing" + +// PaymentRequired writes a 402 response with the canonical instanode.dev +// shape used across all quota-exceeded paths: +// +// HTTP/1.1 402 Payment Required +// WWW-Authenticate: Payment realm="instanode", upgrade_url="https://instanode.dev/pricing" +// Content-Type: application/json +// +// {"ok":false,"error":"quota_exceeded","upgrade_url":"https://instanode.dev/pricing"} +// +// errKey lets callers customise the JSON `error` field for distinct quota +// classes (e.g. "throughput_exceeded", "storage_exceeded"); it falls back +// to the generic "quota_exceeded" when empty so call sites stay terse. +// +// The handler does not actually accept payment yet — the WWW-Authenticate +// header is the forward-compatibility hook for Stripe's Machine Payments +// Protocol. Agents implementing MPP will treat the header as the trigger +// to retry with payment material attached; everyone else just follows +// upgrade_url. +func PaymentRequired(c *fiber.Ctx, errKey string) error { + if errKey == "" { + errKey = "quota_exceeded" + } + c.Set("WWW-Authenticate", + `Payment realm="instanode", upgrade_url="`+QuotaUpgradeURL+`"`) + return c.Status(fiber.StatusPaymentRequired).JSON(fiber.Map{ + "ok": false, + "error": errKey, + "upgrade_url": QuotaUpgradeURL, + }) +} diff --git a/internal/middleware/quota_test.go b/internal/middleware/quota_test.go new file mode 100644 index 00000000..88cf788e --- /dev/null +++ b/internal/middleware/quota_test.go @@ -0,0 +1,71 @@ +package middleware_test + +// quota_test.go — exercises middleware.PaymentRequired, the helper that +// emits HTTP 402 with a Stripe Machine Payments Protocol-compatible +// WWW-Authenticate header when a quota check fails. + +import ( + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/gofiber/fiber/v2" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "instant.dev/internal/middleware" +) + +// Test402_QuotaExceeded verifies that PaymentRequired returns 402 with the +// canonical body shape and WWW-Authenticate: Payment header. +func Test402_QuotaExceeded(t *testing.T) { + app := fiber.New() + app.Post("/db/new", func(c *fiber.Ctx) error { + return middleware.PaymentRequired(c, "") + }) + + req := httptest.NewRequest(http.MethodPost, "/db/new", nil) + resp, err := app.Test(req, 1000) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusPaymentRequired, resp.StatusCode) + + wwwAuth := resp.Header.Get("WWW-Authenticate") + assert.True(t, strings.HasPrefix(wwwAuth, "Payment "), + "WWW-Authenticate must start with `Payment ` keyword (got %q)", wwwAuth) + assert.Contains(t, wwwAuth, `realm="instanode"`) + assert.Contains(t, wwwAuth, `upgrade_url="`+middleware.QuotaUpgradeURL+`"`) + + body, err := io.ReadAll(resp.Body) + require.NoError(t, err) + var parsed map[string]any + require.NoError(t, json.Unmarshal(body, &parsed)) + assert.Equal(t, false, parsed["ok"]) + assert.Equal(t, "quota_exceeded", parsed["error"]) + assert.Equal(t, middleware.QuotaUpgradeURL, parsed["upgrade_url"]) +} + +// Test402_CustomErrorKey verifies the helper accepts a custom error keyword +// (e.g. "storage_exceeded", "throughput_exceeded") for distinct quota classes. +func Test402_CustomErrorKey(t *testing.T) { + app := fiber.New() + app.Post("/db/new", func(c *fiber.Ctx) error { + return middleware.PaymentRequired(c, "storage_exceeded") + }) + + req := httptest.NewRequest(http.MethodPost, "/db/new", nil) + resp, err := app.Test(req, 1000) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusPaymentRequired, resp.StatusCode) + + body, _ := io.ReadAll(resp.Body) + var parsed map[string]any + _ = json.Unmarshal(body, &parsed) + assert.Equal(t, "storage_exceeded", parsed["error"]) +} diff --git a/internal/middleware/rbac.go b/internal/middleware/rbac.go new file mode 100644 index 00000000..da0a2238 --- /dev/null +++ b/internal/middleware/rbac.go @@ -0,0 +1,92 @@ +package middleware + +import ( + "github.com/gofiber/fiber/v2" +) + +// LocalKeyTeamRole is the fiber.Locals key for the authenticated user's role +// on their team (one of: owner, admin, developer, viewer, member). +// +// Populated by RequireAuth after a successful JWT validation, via a SELECT +// against team_members / users.role for (auth_team_id, auth_user_id). +const LocalKeyTeamRole = "auth_team_role" + +// RBAC role constants. Mirrors models.Role* — duplicated here to avoid a +// middleware->models import cycle (middleware is depended on by handlers, +// and models is depended on by handlers). +const ( + RoleOwner = "owner" + RoleAdmin = "admin" + RoleDeveloper = "developer" + RoleViewer = "viewer" + + // roleLegacyMember is treated as developer-equivalent for RBAC purposes: + // "member" was the only non-owner role before the RBAC split landed. + roleLegacyMember = "member" +) + +// roleRank assigns each role an integer rank for hierarchy comparisons. +// Higher rank = more privileges. Unknown roles rank as -1 (deny). +// +// owner = 4 +// admin = 3 +// developer = 2 (also "member" for legacy compat) +// viewer = 1 +func roleRank(role string) int { + switch role { + case RoleOwner: + return 4 + case RoleAdmin: + return 3 + case RoleDeveloper, roleLegacyMember: + return 2 + case RoleViewer: + return 1 + default: + return -1 + } +} + +// GetTeamRole retrieves the authenticated user's role from Fiber locals, +// or "" if not set. Returns "owner", "admin", "developer", or "viewer". +func GetTeamRole(c *fiber.Ctx) string { + if v, ok := c.Locals(LocalKeyTeamRole).(string); ok { + return v + } + return "" +} + +// RequireRole returns a Fiber middleware that gates the request on the +// authenticated user having at least the minimum role. Hierarchy is: +// +// owner > admin > developer > viewer +// +// Examples: +// +// RequireRole("developer") -> owner, admin, developer pass; viewer is rejected +// RequireRole("admin") -> owner, admin pass; developer, viewer rejected +// RequireRole("viewer") -> all four roles pass +// +// Must be installed AFTER RequireAuth so that auth_team_role is populated. +// Returns 403 forbidden / 401 unauthorized on failure. +func RequireRole(min string) fiber.Handler { + required := roleRank(min) + return func(c *fiber.Ctx) error { + // auth_user_id must already be set (RequireAuth must run first). + if GetUserID(c) == "" { + return c.Status(fiber.StatusUnauthorized).JSON(fiber.Map{ + "ok": false, + "error": "unauthorized", + }) + } + actor := GetTeamRole(c) + if roleRank(actor) < required { + return c.Status(fiber.StatusForbidden).JSON(fiber.Map{ + "ok": false, + "error": "forbidden", + "message": "Insufficient role: requires at least " + min, + }) + } + return c.Next() + } +} diff --git a/internal/middleware/rbac_test.go b/internal/middleware/rbac_test.go new file mode 100644 index 00000000..8f23d9fc --- /dev/null +++ b/internal/middleware/rbac_test.go @@ -0,0 +1,143 @@ +package middleware_test + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/gofiber/fiber/v2" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "instant.dev/internal/middleware" +) + +// rbacApp builds a Fiber app that injects (userID, role) into Locals before +// passing through RequireRole. This isolates the role-check logic from JWT +// parsing — those paths are covered in auth_test.go. +func rbacApp(role, userID, requiredRole string) *fiber.App { + app := fiber.New() + app.Use(func(c *fiber.Ctx) error { + c.Locals(middleware.LocalKeyUserID, userID) + if role != "" { + c.Locals(middleware.LocalKeyTeamRole, role) + } + return c.Next() + }) + app.Get("/protected", middleware.RequireRole(requiredRole), func(c *fiber.Ctx) error { + return c.JSON(fiber.Map{"ok": true}) + }) + return app +} + +func mustGet(t *testing.T, app *fiber.App, path string) *http.Response { + t.Helper() + resp, err := app.Test(httptest.NewRequest(http.MethodGet, path, nil), 1000) + require.NoError(t, err) + return resp +} + +// TestRBAC_Hierarchy verifies the canonical hierarchy: owner > admin > developer > viewer. +// RequireRole("developer") must allow owner/admin/developer through and block viewer. +func TestRBAC_Hierarchy(t *testing.T) { + cases := []struct { + actorRole string + want int + }{ + {"owner", http.StatusOK}, + {"admin", http.StatusOK}, + {"developer", http.StatusOK}, + {"member", http.StatusOK}, // legacy alias for developer + {"viewer", http.StatusForbidden}, + {"", http.StatusForbidden}, + {"bogus", http.StatusForbidden}, + } + for _, tc := range cases { + t.Run("require_developer/"+tc.actorRole, func(t *testing.T) { + app := rbacApp(tc.actorRole, "user-123", "developer") + resp := mustGet(t, app, "/protected") + defer resp.Body.Close() + assert.Equal(t, tc.want, resp.StatusCode) + }) + } +} + +// TestRBAC_RequireAdmin only owner/admin pass. +func TestRBAC_RequireAdmin(t *testing.T) { + cases := []struct { + actorRole string + want int + }{ + {"owner", http.StatusOK}, + {"admin", http.StatusOK}, + {"developer", http.StatusForbidden}, + {"member", http.StatusForbidden}, + {"viewer", http.StatusForbidden}, + } + for _, tc := range cases { + t.Run(tc.actorRole, func(t *testing.T) { + app := rbacApp(tc.actorRole, "user-x", "admin") + resp := mustGet(t, app, "/protected") + defer resp.Body.Close() + assert.Equal(t, tc.want, resp.StatusCode) + }) + } +} + +// TestRBAC_RequireOwner only owner passes. +func TestRBAC_RequireOwner(t *testing.T) { + cases := []struct { + actorRole string + want int + }{ + {"owner", http.StatusOK}, + {"admin", http.StatusForbidden}, + {"developer", http.StatusForbidden}, + {"viewer", http.StatusForbidden}, + } + for _, tc := range cases { + t.Run(tc.actorRole, func(t *testing.T) { + app := rbacApp(tc.actorRole, "user-x", "owner") + resp := mustGet(t, app, "/protected") + defer resp.Body.Close() + assert.Equal(t, tc.want, resp.StatusCode) + }) + } +} + +// TestRBAC_RequireViewer all four standard roles pass. +func TestRBAC_RequireViewer(t *testing.T) { + roles := []string{"owner", "admin", "developer", "viewer", "member"} + for _, r := range roles { + t.Run(r, func(t *testing.T) { + app := rbacApp(r, "user-x", "viewer") + resp := mustGet(t, app, "/protected") + defer resp.Body.Close() + assert.Equal(t, http.StatusOK, resp.StatusCode) + }) + } +} + +// TestRBAC_NoUser returns 401 unauthorized — RequireRole must run after RequireAuth. +func TestRBAC_NoUser(t *testing.T) { + app := fiber.New() + app.Get("/x", middleware.RequireRole("viewer"), func(c *fiber.Ctx) error { + return c.JSON(fiber.Map{"ok": true}) + }) + resp := mustGet(t, app, "/x") + defer resp.Body.Close() + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) +} + +// TestRBAC_GetTeamRole_Empty when no role is set Locals returns "". +func TestRBAC_GetTeamRole_Empty(t *testing.T) { + app := fiber.New() + var observed string + app.Get("/x", func(c *fiber.Ctx) error { + observed = middleware.GetTeamRole(c) + return c.JSON(fiber.Map{"ok": true}) + }) + resp := mustGet(t, app, "/x") + defer resp.Body.Close() + assert.Equal(t, "", observed) +} diff --git a/internal/middleware/role_lookup.go b/internal/middleware/role_lookup.go new file mode 100644 index 00000000..07e3e495 --- /dev/null +++ b/internal/middleware/role_lookup.go @@ -0,0 +1,70 @@ +package middleware + +import ( + "context" + "database/sql" + "log/slog" + "sync" + "time" + + "github.com/gofiber/fiber/v2" +) + +// roleLookupDB is the package-level DB handle used by PopulateTeamRole to +// resolve the authenticated user's team role after RequireAuth has set +// LocalKeyUserID and LocalKeyTeamID. Set via SetRoleLookupDB at startup. +var ( + roleLookupMu sync.RWMutex + roleLookupDB *sql.DB +) + +// SetRoleLookupDB registers the platform DB handle used to resolve team roles. +// Wired in router.go after middleware install. A nil DB disables role lookup +// (RequireRole will then deny access for any authenticated request, since +// auth_team_role stays empty). +func SetRoleLookupDB(db *sql.DB) { + roleLookupMu.Lock() + defer roleLookupMu.Unlock() + roleLookupDB = db +} + +func getRoleLookupDB() *sql.DB { + roleLookupMu.RLock() + defer roleLookupMu.RUnlock() + return roleLookupDB +} + +// PopulateTeamRole is a Fiber middleware that runs after RequireAuth and +// hydrates LocalKeyTeamRole by SELECTing the role from team_members for +// (auth_team_id, auth_user_id). Failures are logged and ignored; the +// downstream RequireRole guard will reject. +func PopulateTeamRole() fiber.Handler { + return func(c *fiber.Ctx) error { + userID := GetUserID(c) + teamID := GetTeamID(c) + if userID == "" || teamID == "" { + return c.Next() + } + db := getRoleLookupDB() + if db == nil { + return c.Next() + } + ctx, cancel := context.WithTimeout(c.UserContext(), 750*time.Millisecond) + defer cancel() + var role string + err := db.QueryRowContext(ctx, + `SELECT role FROM users WHERE id = $1 AND team_id = $2`, + userID, teamID, + ).Scan(&role) + if err != nil { + if err != sql.ErrNoRows { + slog.Warn("role_lookup.failed", "error", err, "team_id", teamID, "user_id", userID) + } + return c.Next() + } + if role != "" { + c.Locals(LocalKeyTeamRole, role) + } + return c.Next() + } +} diff --git a/internal/models/api_key.go b/internal/models/api_key.go new file mode 100644 index 00000000..653557eb --- /dev/null +++ b/internal/models/api_key.go @@ -0,0 +1,169 @@ +package models + +import ( + "context" + "crypto/rand" + "crypto/sha256" + "database/sql" + "encoding/base64" + "encoding/hex" + "errors" + "fmt" + "strings" + "time" + + "github.com/google/uuid" + "github.com/lib/pq" +) + +// APIKeyPrefix is the literal prefix every Personal Access Token carries. +// The auth middleware uses it to distinguish a PAT from a JWT without +// parsing the token shape. +const APIKeyPrefix = "ink_" + +// APIKey is a stored, hashed Personal Access Token. +type APIKey struct { + ID uuid.UUID + TeamID uuid.UUID + CreatedBy uuid.NullUUID + Name string + KeyHash string + Scopes []string + LastUsedAt sql.NullTime + RevokedAt sql.NullTime + CreatedAt time.Time +} + +// ErrAPIKeyNotFound — handlers map to 404. Never 401 to avoid distinguishing +// "key revoked" from "key never existed." +var ErrAPIKeyNotFound = errors.New("api key not found") + +// GenerateAPIKeyPlaintext returns a fresh plaintext key in the canonical +// "ink_<base64url>" form. 32 random bytes → ~43 base64 chars → tokens ~47 +// chars total. Caller stores SHA-256(plaintext) via CreateAPIKey. +func GenerateAPIKeyPlaintext() (string, error) { + b := make([]byte, 32) + if _, err := rand.Read(b); err != nil { + return "", fmt.Errorf("rand.Read: %w", err) + } + return APIKeyPrefix + base64.RawURLEncoding.EncodeToString(b), nil +} + +// HashAPIKey returns the storage form of a plaintext PAT. Constant-time +// safe: SHA-256 fixed-time on fixed-length input. +func HashAPIKey(plaintext string) string { + h := sha256.Sum256([]byte(plaintext)) + return hex.EncodeToString(h[:]) +} + +// CreateAPIKey inserts a new key row. Returns the created row (without +// plaintext — caller already has it). +func CreateAPIKey(ctx context.Context, db *sql.DB, teamID uuid.UUID, createdBy uuid.NullUUID, name, keyHash string, scopes []string) (*APIKey, error) { + if len(scopes) == 0 { + scopes = []string{"read", "write"} + } + row := db.QueryRowContext(ctx, ` + INSERT INTO api_keys (team_id, created_by, name, key_hash, scopes) + VALUES ($1, $2, $3, $4, $5) + RETURNING id, team_id, created_by, name, key_hash, scopes, last_used_at, revoked_at, created_at + `, teamID, createdBy, name, keyHash, pq.Array(scopes)) + + k := &APIKey{} + if err := row.Scan( + &k.ID, &k.TeamID, &k.CreatedBy, &k.Name, &k.KeyHash, + pq.Array(&k.Scopes), &k.LastUsedAt, &k.RevokedAt, &k.CreatedAt, + ); err != nil { + return nil, fmt.Errorf("models.CreateAPIKey: %w", err) + } + return k, nil +} + +// GetAPIKeyByHash looks up an active (non-revoked) key by its SHA-256. +// Returns ErrAPIKeyNotFound when the key doesn't exist OR is revoked. +func GetAPIKeyByHash(ctx context.Context, db *sql.DB, keyHash string) (*APIKey, error) { + k := &APIKey{} + err := db.QueryRowContext(ctx, ` + SELECT id, team_id, created_by, name, key_hash, scopes, last_used_at, revoked_at, created_at + FROM api_keys WHERE key_hash = $1 AND revoked_at IS NULL + `, keyHash).Scan( + &k.ID, &k.TeamID, &k.CreatedBy, &k.Name, &k.KeyHash, + pq.Array(&k.Scopes), &k.LastUsedAt, &k.RevokedAt, &k.CreatedAt, + ) + if err == sql.ErrNoRows { + return nil, ErrAPIKeyNotFound + } + if err != nil { + return nil, fmt.Errorf("models.GetAPIKeyByHash: %w", err) + } + return k, nil +} + +// TouchAPIKey best-effort updates last_used_at to now. Failures are logged +// by callers; never block a request. +func TouchAPIKey(ctx context.Context, db *sql.DB, id uuid.UUID) error { + _, err := db.ExecContext(ctx, `UPDATE api_keys SET last_used_at = now() WHERE id = $1`, id) + return err +} + +// ListAPIKeysByTeam returns active and revoked keys, newest first. +// key_hash is included; plaintext is never recoverable. +func ListAPIKeysByTeam(ctx context.Context, db *sql.DB, teamID uuid.UUID) ([]*APIKey, error) { + rows, err := db.QueryContext(ctx, ` + SELECT id, team_id, created_by, name, key_hash, scopes, last_used_at, revoked_at, created_at + FROM api_keys WHERE team_id = $1 ORDER BY created_at DESC + `, teamID) + if err != nil { + return nil, fmt.Errorf("models.ListAPIKeysByTeam: %w", err) + } + defer rows.Close() + + keys := make([]*APIKey, 0) + for rows.Next() { + k := &APIKey{} + if err := rows.Scan( + &k.ID, &k.TeamID, &k.CreatedBy, &k.Name, &k.KeyHash, + pq.Array(&k.Scopes), &k.LastUsedAt, &k.RevokedAt, &k.CreatedAt, + ); err != nil { + return nil, fmt.Errorf("models.ListAPIKeysByTeam scan: %w", err) + } + keys = append(keys, k) + } + return keys, rows.Err() +} + +// RevokeAPIKey sets revoked_at = now() for (team_id, id). Returns +// ErrAPIKeyNotFound when the key doesn't exist for that team or is already +// revoked. Idempotent on subsequent calls. +func RevokeAPIKey(ctx context.Context, db *sql.DB, teamID, id uuid.UUID) error { + res, err := db.ExecContext(ctx, ` + UPDATE api_keys SET revoked_at = now() + WHERE id = $1 AND team_id = $2 AND revoked_at IS NULL + `, id, teamID) + if err != nil { + return fmt.Errorf("models.RevokeAPIKey: %w", err) + } + n, err := res.RowsAffected() + if err != nil { + return fmt.Errorf("models.RevokeAPIKey rows: %w", err) + } + if n == 0 { + return ErrAPIKeyNotFound + } + return nil +} + +// HasScope reports whether the key carries the given scope (or a higher one). +// Hierarchy: admin > write > read. +func (k *APIKey) HasScope(want string) bool { + rank := map[string]int{"read": 1, "write": 2, "admin": 3} + wantRank, ok := rank[want] + if !ok { + return false + } + for _, s := range k.Scopes { + if r, ok := rank[strings.ToLower(s)]; ok && r >= wantRank { + return true + } + } + return false +} diff --git a/internal/models/audit_log.go b/internal/models/audit_log.go new file mode 100644 index 00000000..72475bff --- /dev/null +++ b/internal/models/audit_log.go @@ -0,0 +1,122 @@ +package models + +// audit_log.go — per-team event stream consumed by the dashboard's +// Recent Activity feed. +// +// Writes are best-effort: callers fire InsertAuditEvent in a goroutine +// and ignore the returned error. A failure to record an audit event +// must NEVER block a provision, claim, or rotate. +// +// Reads come from GET /api/v1/audit, capped at 200 rows per call. + +import ( + "context" + "database/sql" + "fmt" + "time" + + "github.com/google/uuid" +) + +// auditMaxLimit caps the number of rows returned by ListAuditEventsByTeam. +// Keeps a single call from sweeping a large team's history; the dashboard +// uses limit=20 by default. +const auditMaxLimit = 200 + +// AuditEvent is one row in the audit_log table. Metadata is stored as +// raw JSONB bytes so callers can serialize arbitrary k/v without the +// model needing to know the shape. +type AuditEvent struct { + ID uuid.UUID + TeamID uuid.UUID + UserID uuid.NullUUID + Actor string + Kind string + ResourceType string + ResourceID uuid.NullUUID + Summary string + Metadata []byte + CreatedAt time.Time +} + +// InsertAuditEvent inserts a row best-effort. Callers should run this in +// a goroutine and ignore the error; an audit failure must never surface +// to the user. Defaults: Actor → "agent" when empty. +func InsertAuditEvent(ctx context.Context, db *sql.DB, ev AuditEvent) error { + if ev.Actor == "" { + ev.Actor = "agent" + } + // resource_type is NULL when empty (the column allows NULL). + var resourceType interface{} + if ev.ResourceType != "" { + resourceType = ev.ResourceType + } + var metadata interface{} + if len(ev.Metadata) > 0 { + metadata = ev.Metadata + } + _, err := db.ExecContext(ctx, ` + INSERT INTO audit_log (team_id, user_id, actor, kind, resource_type, resource_id, summary, metadata) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8) + `, ev.TeamID, ev.UserID, ev.Actor, ev.Kind, resourceType, ev.ResourceID, ev.Summary, metadata) + if err != nil { + return fmt.Errorf("models.InsertAuditEvent: %w", err) + } + return nil +} + +// ListAuditEventsByTeam returns the most recent events for a team, +// newest first. kindFilter == "" means all kinds. Limit is capped at +// auditMaxLimit; non-positive limits default to 20. +func ListAuditEventsByTeam(ctx context.Context, db *sql.DB, teamID uuid.UUID, limit int, kindFilter string) ([]*AuditEvent, error) { + if limit <= 0 { + limit = 20 + } + if limit > auditMaxLimit { + limit = auditMaxLimit + } + + var rows *sql.Rows + var err error + if kindFilter == "" { + rows, err = db.QueryContext(ctx, ` + SELECT id, team_id, user_id, actor, kind, COALESCE(resource_type, ''), resource_id, summary, metadata, created_at + FROM audit_log + WHERE team_id = $1 + ORDER BY created_at DESC + LIMIT $2 + `, teamID, limit) + } else { + rows, err = db.QueryContext(ctx, ` + SELECT id, team_id, user_id, actor, kind, COALESCE(resource_type, ''), resource_id, summary, metadata, created_at + FROM audit_log + WHERE team_id = $1 AND kind = $2 + ORDER BY created_at DESC + LIMIT $3 + `, teamID, kindFilter, limit) + } + if err != nil { + return nil, fmt.Errorf("models.ListAuditEventsByTeam: %w", err) + } + defer rows.Close() + + out := make([]*AuditEvent, 0) + for rows.Next() { + ev := &AuditEvent{} + var metadata sql.NullString + if err := rows.Scan( + &ev.ID, &ev.TeamID, &ev.UserID, &ev.Actor, &ev.Kind, + &ev.ResourceType, &ev.ResourceID, &ev.Summary, &metadata, &ev.CreatedAt, + ); err != nil { + return nil, fmt.Errorf("models.ListAuditEventsByTeam scan: %w", err) + } + if metadata.Valid { + ev.Metadata = []byte(metadata.String) + } + out = append(out, ev) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("models.ListAuditEventsByTeam rows: %w", err) + } + return out, nil +} diff --git a/internal/models/custom_domain.go b/internal/models/custom_domain.go new file mode 100644 index 00000000..574aa420 --- /dev/null +++ b/internal/models/custom_domain.go @@ -0,0 +1,304 @@ +package models + +// custom_domain.go — Pro+ custom hostnames for stacks. +// +// One row per hostname. The verification_token is the random value the customer +// includes in their TXT challenge record (`_instanode.<hostname>` → +// `instanode-verify-<token>`). Once we observe the TXT record, the row advances +// from "pending_verification" → "verified". The handler then creates an +// Ingress + cert-manager Certificate; status moves to "ingress_ready" and +// finally "cert_ready" / "live" once the cert is issued. + +import ( + "context" + "crypto/rand" + "database/sql" + "encoding/hex" + "errors" + "fmt" + "strings" + "time" + + "github.com/google/uuid" +) + +// Custom-domain status values. Strings are stored verbatim in the DB; do not +// rename without a migration. +const ( + CustomDomainStatusPending = "pending_verification" + CustomDomainStatusVerified = "verified" + CustomDomainStatusIngressReady = "ingress_ready" + CustomDomainStatusCertReady = "cert_ready" + CustomDomainStatusLive = "live" + CustomDomainStatusFailed = "failed" +) + +// VerificationTokenPrefix is the literal prefix the customer must include in +// their TXT record value alongside the random token. Together they form the +// expected payload "instanode-verify-<token>". +const VerificationTokenPrefix = "instanode-verify-" + +// CustomDomain is one row of the custom_domains table. +type CustomDomain struct { + ID uuid.UUID + TeamID uuid.UUID + StackID uuid.UUID + Hostname string + VerificationToken string + Status string + VerifiedAt sql.NullTime + CertReadyAt sql.NullTime + LastCheckAt sql.NullTime + LastCheckErr sql.NullString + CreatedAt time.Time +} + +// ErrCustomDomainNotFound is returned when a lookup yields no rows. +var ErrCustomDomainNotFound = errors.New("custom domain not found") + +// ErrCustomDomainTaken is returned when the hostname is already bound to a +// different domain row (UNIQUE constraint violation). +var ErrCustomDomainTaken = errors.New("hostname already bound to another domain") + +// generateVerificationToken returns a 32-char hex token (16 random bytes). +// The token is the per-row random part of the TXT challenge value. +func generateVerificationToken() (string, error) { + b := make([]byte, 16) + if _, err := rand.Read(b); err != nil { + return "", fmt.Errorf("models.generateVerificationToken: %w", err) + } + return hex.EncodeToString(b), nil +} + +// scanCustomDomain reads a custom_domains row into a CustomDomain. +func scanCustomDomain(row interface { + Scan(dest ...any) error +}) (*CustomDomain, error) { + d := &CustomDomain{} + if err := row.Scan( + &d.ID, &d.TeamID, &d.StackID, &d.Hostname, + &d.VerificationToken, &d.Status, + &d.VerifiedAt, &d.CertReadyAt, + &d.LastCheckAt, &d.LastCheckErr, + &d.CreatedAt, + ); err != nil { + return nil, err + } + return d, nil +} + +const customDomainSelectFields = ` + id, team_id, stack_id, hostname, + verification_token, status, + verified_at, cert_ready_at, + last_check_at, last_check_err, + created_at +` + +// CreateCustomDomain inserts a row inside a transaction. The verification +// token is generated server-side. Returns ErrCustomDomainTaken on UNIQUE +// violation (another team or stack already claimed the hostname). +// +// All callers must provide a non-zero teamID, stackID, and lowercased hostname; +// the handler is responsible for hostname validation upstream. +func CreateCustomDomain(ctx context.Context, db *sql.DB, teamID, stackID uuid.UUID, hostname string) (*CustomDomain, error) { + token, err := generateVerificationToken() + if err != nil { + return nil, err + } + + tx, err := db.BeginTx(ctx, nil) + if err != nil { + return nil, fmt.Errorf("models.CreateCustomDomain: begin tx: %w", err) + } + committed := false + defer func() { + if !committed { + _ = tx.Rollback() + } + }() + + row := tx.QueryRowContext(ctx, ` + INSERT INTO custom_domains (team_id, stack_id, hostname, verification_token) + VALUES ($1, $2, $3, $4) + RETURNING `+customDomainSelectFields, + teamID, stackID, hostname, token, + ) + d, scanErr := scanCustomDomain(row) + if scanErr != nil { + // Postgres UNIQUE violation → ErrCustomDomainTaken. The pq driver returns + // a structured error but we keep the dependency surface small here and + // match on the error string the way other models do. + if isUniqueViolation(scanErr) { + return nil, ErrCustomDomainTaken + } + return nil, fmt.Errorf("models.CreateCustomDomain: %w", scanErr) + } + + if err := tx.Commit(); err != nil { + return nil, fmt.Errorf("models.CreateCustomDomain: commit: %w", err) + } + committed = true + return d, nil +} + +// isUniqueViolation matches the Postgres SQLSTATE 23505 the lib/pq driver +// surfaces in its Error() text. Avoids a hard dependency on pq's error type +// in this file. +func isUniqueViolation(err error) bool { + if err == nil { + return false + } + msg := err.Error() + // pq error: "ERROR: duplicate key value violates unique constraint ..." + // pgx error: "ERROR: duplicate key value..." + return strings.Contains(msg, "duplicate key value") || strings.Contains(msg, "23505") +} + +// GetCustomDomainByID returns a single row or ErrCustomDomainNotFound. +func GetCustomDomainByID(ctx context.Context, db *sql.DB, id uuid.UUID) (*CustomDomain, error) { + row := db.QueryRowContext(ctx, ` + SELECT `+customDomainSelectFields+` + FROM custom_domains WHERE id = $1 + `, id) + d, err := scanCustomDomain(row) + if err == sql.ErrNoRows { + return nil, ErrCustomDomainNotFound + } + if err != nil { + return nil, fmt.Errorf("models.GetCustomDomainByID: %w", err) + } + return d, nil +} + +// ListCustomDomainsByStack returns every domain bound to the given stack, +// newest first. +func ListCustomDomainsByStack(ctx context.Context, db *sql.DB, stackID uuid.UUID) ([]*CustomDomain, error) { + rows, err := db.QueryContext(ctx, ` + SELECT `+customDomainSelectFields+` + FROM custom_domains + WHERE stack_id = $1 + ORDER BY created_at DESC + `, stackID) + if err != nil { + return nil, fmt.Errorf("models.ListCustomDomainsByStack: %w", err) + } + defer rows.Close() + + out := make([]*CustomDomain, 0) + for rows.Next() { + d, err := scanCustomDomain(rows) + if err != nil { + return nil, fmt.Errorf("models.ListCustomDomainsByStack scan: %w", err) + } + out = append(out, d) + } + return out, rows.Err() +} + +// ListCustomDomainsByTeam returns every domain owned by the team, newest first. +func ListCustomDomainsByTeam(ctx context.Context, db *sql.DB, teamID uuid.UUID) ([]*CustomDomain, error) { + rows, err := db.QueryContext(ctx, ` + SELECT `+customDomainSelectFields+` + FROM custom_domains + WHERE team_id = $1 + ORDER BY created_at DESC + `, teamID) + if err != nil { + return nil, fmt.Errorf("models.ListCustomDomainsByTeam: %w", err) + } + defer rows.Close() + + out := make([]*CustomDomain, 0) + for rows.Next() { + d, err := scanCustomDomain(rows) + if err != nil { + return nil, fmt.Errorf("models.ListCustomDomainsByTeam scan: %w", err) + } + out = append(out, d) + } + return out, rows.Err() +} + +// UpdateCustomDomainStatus advances the status field and records the +// last-check metadata. lastCheckErr may be empty (sets NULL). +func UpdateCustomDomainStatus(ctx context.Context, db *sql.DB, id uuid.UUID, status, lastCheckErr string) error { + var errVal interface{} + if lastCheckErr != "" { + errVal = lastCheckErr + } + res, err := db.ExecContext(ctx, ` + UPDATE custom_domains + SET status = $1, + last_check_at = now(), + last_check_err = $2 + WHERE id = $3 + `, status, errVal, id) + if err != nil { + return fmt.Errorf("models.UpdateCustomDomainStatus: %w", err) + } + n, _ := res.RowsAffected() + if n == 0 { + return ErrCustomDomainNotFound + } + return nil +} + +// MarkCustomDomainVerified sets verified_at = now() and status = "verified". +// last_check_err is cleared because we just succeeded. +func MarkCustomDomainVerified(ctx context.Context, db *sql.DB, id uuid.UUID) error { + res, err := db.ExecContext(ctx, ` + UPDATE custom_domains + SET status = $1, + verified_at = now(), + last_check_at = now(), + last_check_err = NULL + WHERE id = $2 + `, CustomDomainStatusVerified, id) + if err != nil { + return fmt.Errorf("models.MarkCustomDomainVerified: %w", err) + } + n, _ := res.RowsAffected() + if n == 0 { + return ErrCustomDomainNotFound + } + return nil +} + +// MarkCertReady sets cert_ready_at = now() and status = "cert_ready". +// last_check_err is cleared. Callers may transition further to "live" via +// UpdateCustomDomainStatus once they confirm the hostname resolves. +func MarkCertReady(ctx context.Context, db *sql.DB, id uuid.UUID) error { + res, err := db.ExecContext(ctx, ` + UPDATE custom_domains + SET status = $1, + cert_ready_at = now(), + last_check_at = now(), + last_check_err = NULL + WHERE id = $2 + `, CustomDomainStatusCertReady, id) + if err != nil { + return fmt.Errorf("models.MarkCertReady: %w", err) + } + n, _ := res.RowsAffected() + if n == 0 { + return ErrCustomDomainNotFound + } + return nil +} + +// DeleteCustomDomain removes the row matching (id, teamID). Returns +// ErrCustomDomainNotFound when no such row exists for the team. +func DeleteCustomDomain(ctx context.Context, db *sql.DB, id, teamID uuid.UUID) error { + res, err := db.ExecContext(ctx, ` + DELETE FROM custom_domains WHERE id = $1 AND team_id = $2 + `, id, teamID) + if err != nil { + return fmt.Errorf("models.DeleteCustomDomain: %w", err) + } + n, _ := res.RowsAffected() + if n == 0 { + return ErrCustomDomainNotFound + } + return nil +} diff --git a/internal/models/deployment.go b/internal/models/deployment.go index 2f428306..b52aad66 100644 --- a/internal/models/deployment.go +++ b/internal/models/deployment.go @@ -22,6 +22,7 @@ type Deployment struct { EnvVars map[string]string Port int Tier string + Env string // dev | staging | production | <custom>; defaults to "production" ErrorMessage string CreatedAt time.Time UpdatedAt time.Time @@ -34,6 +35,7 @@ type CreateDeploymentParams struct { AppID string Port int Tier string + Env string // empty string is normalised to EnvProduction EnvVars map[string]string } @@ -46,6 +48,10 @@ func (e *ErrDeploymentNotFound) Error() string { return fmt.Sprintf("deployment not found: %s", e.ID) } +// deploymentColumns is the canonical column list shared by all deployment SELECTs. +const deploymentColumns = `id, team_id, resource_id, app_id, provider_id, status, app_url, + env_vars, port, tier, env, error_message, created_at, updated_at` + // scanDeployment reads a single deployments row into a Deployment struct. // env_vars is stored as JSONB; error_message, provider_id, and app_url are nullable. func scanDeployment(row interface { @@ -59,7 +65,7 @@ func scanDeployment(row interface { if err := row.Scan( &d.ID, &d.TeamID, &resourceID, &d.AppID, &providerID, &d.Status, &appURL, - &envVarsRaw, &d.Port, &d.Tier, &errorMessage, + &envVarsRaw, &d.Port, &d.Tier, &d.Env, &errorMessage, &d.CreatedAt, &d.UpdatedAt, ); err != nil { return nil, err @@ -103,13 +109,17 @@ func CreateDeployment(ctx context.Context, db *sql.DB, p CreateDeploymentParams) return nil, fmt.Errorf("models.CreateDeployment: marshal env_vars: %w", err) } + env := p.Env + if env == "" { + env = EnvProduction + } + row := db.QueryRowContext(ctx, ` INSERT INTO deployments - (team_id, resource_id, app_id, port, tier, env_vars) - VALUES ($1, $2, $3, $4, $5, $6) - RETURNING id, team_id, resource_id, app_id, provider_id, status, app_url, - env_vars, port, tier, error_message, created_at, updated_at - `, p.TeamID, resourceID, p.AppID, port, p.Tier, envVarsJSON) + (team_id, resource_id, app_id, port, tier, env, env_vars) + VALUES ($1, $2, $3, $4, $5, $6, $7) + RETURNING `+deploymentColumns, + p.TeamID, resourceID, p.AppID, port, p.Tier, env, envVarsJSON) d, err := scanDeployment(row) if err != nil { @@ -119,12 +129,10 @@ func CreateDeployment(ctx context.Context, db *sql.DB, p CreateDeploymentParams) } // GetDeploymentByAppID fetches a deployment by its app_id slug (the short public token). +// app_id is unique across all envs — the same app name in dev vs prod must use distinct +// app_ids (the deploy handler generates a fresh one per call). func GetDeploymentByAppID(ctx context.Context, db *sql.DB, appID string) (*Deployment, error) { - row := db.QueryRowContext(ctx, ` - SELECT id, team_id, resource_id, app_id, provider_id, status, app_url, - env_vars, port, tier, error_message, created_at, updated_at - FROM deployments WHERE app_id = $1 - `, appID) + row := db.QueryRowContext(ctx, `SELECT `+deploymentColumns+` FROM deployments WHERE app_id = $1`, appID) d, err := scanDeployment(row) if err == sql.ErrNoRows { @@ -138,11 +146,7 @@ func GetDeploymentByAppID(ctx context.Context, db *sql.DB, appID string) (*Deplo // GetDeploymentByID fetches a deployment by primary key UUID. func GetDeploymentByID(ctx context.Context, db *sql.DB, id uuid.UUID) (*Deployment, error) { - row := db.QueryRowContext(ctx, ` - SELECT id, team_id, resource_id, app_id, provider_id, status, app_url, - env_vars, port, tier, error_message, created_at, updated_at - FROM deployments WHERE id = $1 - `, id) + row := db.QueryRowContext(ctx, `SELECT `+deploymentColumns+` FROM deployments WHERE id = $1`, id) d, err := scanDeployment(row) if err == sql.ErrNoRows { @@ -154,11 +158,11 @@ func GetDeploymentByID(ctx context.Context, db *sql.DB, id uuid.UUID) (*Deployme return d, nil } -// GetDeploymentsByTeam returns all deployments for a team, ordered by creation time descending. +// GetDeploymentsByTeam returns all deployments for a team across every environment, +// ordered by creation time descending. func GetDeploymentsByTeam(ctx context.Context, db *sql.DB, teamID uuid.UUID) ([]*Deployment, error) { rows, err := db.QueryContext(ctx, ` - SELECT id, team_id, resource_id, app_id, provider_id, status, app_url, - env_vars, port, tier, error_message, created_at, updated_at + SELECT `+deploymentColumns+` FROM deployments WHERE team_id = $1 ORDER BY created_at DESC @@ -182,6 +186,37 @@ func GetDeploymentsByTeam(ctx context.Context, db *sql.DB, teamID uuid.UUID) ([] return results, nil } +// GetDeploymentsByTeamAndEnv returns deployments for a team scoped to a single +// environment. Empty env is normalised to "production". +func GetDeploymentsByTeamAndEnv(ctx context.Context, db *sql.DB, teamID uuid.UUID, env string) ([]*Deployment, error) { + if env == "" { + env = EnvProduction + } + rows, err := db.QueryContext(ctx, ` + SELECT `+deploymentColumns+` + FROM deployments + WHERE team_id = $1 AND env = $2 + ORDER BY created_at DESC + `, teamID, env) + if err != nil { + return nil, fmt.Errorf("models.GetDeploymentsByTeamAndEnv: %w", err) + } + defer rows.Close() + + var results []*Deployment + for rows.Next() { + d, err := scanDeployment(rows) + if err != nil { + return nil, fmt.Errorf("models.GetDeploymentsByTeamAndEnv scan: %w", err) + } + results = append(results, d) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("models.GetDeploymentsByTeamAndEnv rows: %w", err) + } + return results, nil +} + // UpdateDeploymentStatus updates the status and optional error_message for a deployment. // updated_at is set to now() by the database. func UpdateDeploymentStatus(ctx context.Context, db *sql.DB, id uuid.UUID, status, errorMessage string) error { diff --git a/internal/models/deployment_env_test.go b/internal/models/deployment_env_test.go new file mode 100644 index 00000000..61f347d0 --- /dev/null +++ b/internal/models/deployment_env_test.go @@ -0,0 +1,116 @@ +package models_test + +// deployment_env_test.go — env-column tests for the Deployment model. +// Skips when TEST_DATABASE_URL is unset (see requireDB in resource_env_test.go). + +import ( + "context" + "testing" + + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "instant.dev/internal/models" + "instant.dev/internal/testhelpers" +) + +func TestDeploymentEnv_CreateDefaultsToProduction(t *testing.T) { + requireDB(t) + db, cleanDB := testhelpers.SetupTestDB(t) + defer cleanDB() + + teamID := uuid.MustParse(testhelpers.MustCreateTeamDB(t, db, "hobby")) + defer db.Exec(`DELETE FROM teams WHERE id = $1`, teamID) + + d, err := models.CreateDeployment(context.Background(), db, models.CreateDeploymentParams{ + TeamID: teamID, + AppID: "app-test-" + uuid.NewString()[:8], + Tier: "hobby", + // Env intentionally empty → must default. + }) + require.NoError(t, err) + defer db.Exec(`DELETE FROM deployments WHERE id = $1`, d.ID) + + assert.Equal(t, models.EnvProduction, d.Env) +} + +func TestDeploymentEnv_CreateRoundTrips(t *testing.T) { + requireDB(t) + db, cleanDB := testhelpers.SetupTestDB(t) + defer cleanDB() + + teamID := uuid.MustParse(testhelpers.MustCreateTeamDB(t, db, "hobby")) + defer db.Exec(`DELETE FROM teams WHERE id = $1`, teamID) + + for _, env := range []string{"dev", "staging", "production"} { + t.Run(env, func(t *testing.T) { + d, err := models.CreateDeployment(context.Background(), db, models.CreateDeploymentParams{ + TeamID: teamID, + AppID: "app-" + env + "-" + uuid.NewString()[:8], + Tier: "hobby", + Env: env, + }) + require.NoError(t, err) + defer db.Exec(`DELETE FROM deployments WHERE id = $1`, d.ID) + assert.Equal(t, env, d.Env) + + got, err := models.GetDeploymentByAppID(context.Background(), db, d.AppID) + require.NoError(t, err) + assert.Equal(t, env, got.Env) + }) + } +} + +// TestDeploymentEnv_AppNameIsolation: same logical app deployed to dev and prod +// must produce two distinct rows. (app_id itself is unique per row — the +// handler generates fresh ones — so we confirm the env column makes them +// distinguishable from the model's POV.) +func TestDeploymentEnv_AppNameIsolation(t *testing.T) { + requireDB(t) + db, cleanDB := testhelpers.SetupTestDB(t) + defer cleanDB() + + teamID := uuid.MustParse(testhelpers.MustCreateTeamDB(t, db, "hobby")) + defer db.Exec(`DELETE FROM teams WHERE id = $1`, teamID) + + dev, err := models.CreateDeployment(context.Background(), db, models.CreateDeploymentParams{ + TeamID: teamID, + AppID: "myapp-dev-" + uuid.NewString()[:8], + Tier: "hobby", + Env: "dev", + EnvVars: map[string]string{"_name": "myapp"}, + }) + require.NoError(t, err) + defer db.Exec(`DELETE FROM deployments WHERE id = $1`, dev.ID) + + prod, err := models.CreateDeployment(context.Background(), db, models.CreateDeploymentParams{ + TeamID: teamID, + AppID: "myapp-prod-" + uuid.NewString()[:8], + Tier: "hobby", + Env: "production", + EnvVars: map[string]string{"_name": "myapp"}, + }) + require.NoError(t, err) + defer db.Exec(`DELETE FROM deployments WHERE id = $1`, prod.ID) + + assert.NotEqual(t, dev.ID, prod.ID, "two envs must produce two rows") + assert.Equal(t, "dev", dev.Env) + assert.Equal(t, "production", prod.Env) + + devList, err := models.GetDeploymentsByTeamAndEnv(context.Background(), db, teamID, "dev") + require.NoError(t, err) + assert.Len(t, devList, 1) + assert.Equal(t, dev.ID, devList[0].ID) + + prodList, err := models.GetDeploymentsByTeamAndEnv(context.Background(), db, teamID, "") + require.NoError(t, err) + // Filter out unrelated rows from concurrent tests. + var matched int + for _, d := range prodList { + if d.ID == prod.ID { + matched++ + } + } + assert.Equal(t, 1, matched) +} diff --git a/internal/models/magic_link.go b/internal/models/magic_link.go new file mode 100644 index 00000000..16bc53a0 --- /dev/null +++ b/internal/models/magic_link.go @@ -0,0 +1,115 @@ +package models + +import ( + "context" + "crypto/rand" + "crypto/sha256" + "database/sql" + "encoding/base64" + "encoding/hex" + "errors" + "fmt" + "time" + + "github.com/google/uuid" +) + +// MagicLinkPrefix is the literal prefix every magic-link plaintext token +// carries. Visible in logs and emails so it's recognizable as a magic-link +// token (vs. a PAT "ink_" or a session JWT). +const MagicLinkPrefix = "mlnk_" + +// MagicLink is a stored, hashed passwordless login token. +type MagicLink struct { + ID uuid.UUID + Email string + TokenHash string + ReturnTo string + ExpiresAt time.Time + ConsumedAt sql.NullTime + CreatedAt time.Time +} + +// ErrMagicLinkNotFound is returned when a hash lookup yields no rows OR the +// row is expired/consumed. Callers should NEVER distinguish between those +// cases in their response — return a generic "invalid or expired link" +// message either way. +var ErrMagicLinkNotFound = errors.New("magic link not found, expired, or already used") + +// GenerateMagicLinkPlaintext returns a fresh plaintext token in the canonical +// "mlnk_<base64url>" form. 32 random bytes → ~43 base64 chars → tokens ~48 +// chars total. The caller is expected to hash it with HashMagicLink and pass +// only the hash to CreateMagicLink. +func GenerateMagicLinkPlaintext() (string, error) { + b := make([]byte, 32) + if _, err := rand.Read(b); err != nil { + return "", fmt.Errorf("rand.Read: %w", err) + } + return MagicLinkPrefix + base64.RawURLEncoding.EncodeToString(b), nil +} + +// HashMagicLink returns the storage form of a plaintext magic-link token. +// SHA-256 is constant-time on fixed-length input. +func HashMagicLink(plaintext string) string { + h := sha256.Sum256([]byte(plaintext)) + return hex.EncodeToString(h[:]) +} + +// CreateMagicLink inserts a new row. The plaintext is hashed; only the hash +// is persisted. ttl is added to now() to derive expires_at. +func CreateMagicLink(ctx context.Context, db *sql.DB, email, plaintext, returnTo string, ttl time.Duration) (*MagicLink, error) { + hash := HashMagicLink(plaintext) + expiresAt := time.Now().UTC().Add(ttl) + + m := &MagicLink{} + err := db.QueryRowContext(ctx, ` + INSERT INTO magic_links (email, token_hash, return_to, expires_at) + VALUES ($1, $2, $3, $4) + RETURNING id, email, token_hash, return_to, expires_at, consumed_at, created_at + `, email, hash, returnTo, expiresAt).Scan( + &m.ID, &m.Email, &m.TokenHash, &m.ReturnTo, &m.ExpiresAt, &m.ConsumedAt, &m.CreatedAt, + ) + if err != nil { + return nil, fmt.Errorf("models.CreateMagicLink: %w", err) + } + return m, nil +} + +// GetMagicLinkForConsumption looks up an unconsumed, non-expired link by its +// hash. Returns ErrMagicLinkNotFound when the hash doesn't exist, the link is +// already consumed, or it's past expires_at. +func GetMagicLinkForConsumption(ctx context.Context, db *sql.DB, hash string) (*MagicLink, error) { + m := &MagicLink{} + err := db.QueryRowContext(ctx, ` + SELECT id, email, token_hash, return_to, expires_at, consumed_at, created_at + FROM magic_links + WHERE token_hash = $1 AND consumed_at IS NULL AND expires_at > now() + `, hash).Scan( + &m.ID, &m.Email, &m.TokenHash, &m.ReturnTo, &m.ExpiresAt, &m.ConsumedAt, &m.CreatedAt, + ) + if err == sql.ErrNoRows { + return nil, ErrMagicLinkNotFound + } + if err != nil { + return nil, fmt.Errorf("models.GetMagicLinkForConsumption: %w", err) + } + return m, nil +} + +// ConsumeMagicLink atomically marks a link as consumed. Returns true on the +// first call, false on every subsequent call (single-use). Callers should +// treat false as ErrMagicLinkNotFound — somebody beat us to the row. +func ConsumeMagicLink(ctx context.Context, db *sql.DB, id uuid.UUID) (bool, error) { + res, err := db.ExecContext(ctx, ` + UPDATE magic_links SET consumed_at = now() + WHERE id = $1 AND consumed_at IS NULL + `, id) + if err != nil { + return false, fmt.Errorf("models.ConsumeMagicLink: %w", err) + } + n, err := res.RowsAffected() + if err != nil { + return false, fmt.Errorf("models.ConsumeMagicLink rows: %w", err) + } + return n == 1, nil +} diff --git a/internal/models/resource.go b/internal/models/resource.go index c7da56ae..36336bb0 100644 --- a/internal/models/resource.go +++ b/internal/models/resource.go @@ -4,11 +4,33 @@ import ( "context" "database/sql" "fmt" + "regexp" "time" "github.com/google/uuid" ) +// EnvProduction is the default environment used when callers omit one. +// All migration-backfilled rows start at this value. +const EnvProduction = "production" + +// envPattern restricts the env name to lowercase alphanumerics + dashes, +// 1–32 chars. Enforced at the model boundary so every caller (handlers, +// background jobs, internal endpoints) gets the same guarantee. +var envPattern = regexp.MustCompile(`^[a-z0-9-]{1,32}$`) + +// NormalizeEnv coerces an empty env to EnvProduction (backwards compat) and +// validates the format. Returns (env, true) when valid, ("", false) otherwise. +func NormalizeEnv(env string) (string, bool) { + if env == "" { + return EnvProduction, true + } + if !envPattern.MatchString(env) { + return "", false + } + return env, true +} + // Resource represents any provisioned resource (postgres, redis, mongodb, queue, webhook, storage). type Resource struct { ID uuid.UUID @@ -19,6 +41,7 @@ type Resource struct { ConnectionURL sql.NullString // AES-256-GCM encrypted KeyPrefix sql.NullString // provisioner key prefix (e.g. "pool_abc:") for Redis Tier string + Env string // dev | staging | production | <custom>; defaults to "production" Fingerprint sql.NullString CloudVendor sql.NullString CountryCode sql.NullString @@ -46,6 +69,7 @@ type CreateResourceParams struct { ResourceType string Name string Tier string + Env string // empty string is normalised to EnvProduction Fingerprint string CloudVendor string CountryCode string @@ -53,6 +77,28 @@ type CreateResourceParams struct { CreatedRequestID string } +// resourceColumns is the canonical list of columns selected by every read query. +// Centralising the column list (and the matching scan order in scanResource) +// makes it easy to add a new column without touching half a dozen functions. +const resourceColumns = `id, team_id, token, resource_type, name, connection_url, key_prefix, tier, + env, fingerprint, cloud_vendor, country_code, status, migration_status, + expires_at, storage_bytes, provider_resource_id, created_request_id, created_at` + +// scanResource reads a single resources row in the order defined by resourceColumns. +func scanResource(row interface { + Scan(dest ...any) error +}) (*Resource, error) { + r := &Resource{} + if err := row.Scan( + &r.ID, &r.TeamID, &r.Token, &r.ResourceType, &r.Name, &r.ConnectionURL, &r.KeyPrefix, + &r.Tier, &r.Env, &r.Fingerprint, &r.CloudVendor, &r.CountryCode, &r.Status, + &r.MigrationStatus, &r.ExpiresAt, &r.StorageBytes, &r.ProviderResourceID, &r.CreatedRequestID, &r.CreatedAt, + ); err != nil { + return nil, err + } + return r, nil +} + // CreateResource inserts a new resource row and returns it. func CreateResource(ctx context.Context, db *sql.DB, p CreateResourceParams) (*Resource, error) { var teamID interface{} @@ -64,21 +110,21 @@ func CreateResource(ctx context.Context, db *sql.DB, p CreateResourceParams) (*R expiresAt = *p.ExpiresAt } - r := &Resource{} - err := db.QueryRowContext(ctx, ` + env := p.Env + if env == "" { + env = EnvProduction + } + + row := db.QueryRowContext(ctx, ` INSERT INTO resources - (team_id, resource_type, name, tier, fingerprint, cloud_vendor, country_code, expires_at, created_request_id) - VALUES ($1, $2, NULLIF($3,''), $4, NULLIF($5,''), NULLIF($6,''), NULLIF($7,''), $8, NULLIF($9,'')) - RETURNING id, team_id, token, resource_type, name, connection_url, key_prefix, tier, - fingerprint, cloud_vendor, country_code, status, migration_status, - expires_at, storage_bytes, created_request_id, created_at - `, teamID, p.ResourceType, p.Name, p.Tier, p.Fingerprint, p.CloudVendor, p.CountryCode, + (team_id, resource_type, name, tier, env, fingerprint, cloud_vendor, country_code, expires_at, created_request_id) + VALUES ($1, $2, NULLIF($3,''), $4, $5, NULLIF($6,''), NULLIF($7,''), NULLIF($8,''), $9, NULLIF($10,'')) + RETURNING `+resourceColumns, + teamID, p.ResourceType, p.Name, p.Tier, env, p.Fingerprint, p.CloudVendor, p.CountryCode, expiresAt, p.CreatedRequestID, - ).Scan( - &r.ID, &r.TeamID, &r.Token, &r.ResourceType, &r.Name, &r.ConnectionURL, &r.KeyPrefix, - &r.Tier, &r.Fingerprint, &r.CloudVendor, &r.CountryCode, &r.Status, - &r.MigrationStatus, &r.ExpiresAt, &r.StorageBytes, &r.CreatedRequestID, &r.CreatedAt, ) + + r, err := scanResource(row) if err != nil { return nil, fmt.Errorf("models.CreateResource: %w", err) } @@ -87,6 +133,7 @@ func CreateResource(ctx context.Context, db *sql.DB, p CreateResourceParams) (*R // CountActiveResourcesByTeamAndType returns the number of active (non-deleted) // resources of the given type owned by a team. Used for plan limit enforcement. +// Counts across ALL environments — plan limits apply per team, not per env. func CountActiveResourcesByTeamAndType(ctx context.Context, db *sql.DB, teamID uuid.UUID, resourceType string) (int, error) { var count int err := db.QueryRowContext(ctx, @@ -101,17 +148,8 @@ func CountActiveResourcesByTeamAndType(ctx context.Context, db *sql.DB, teamID u // GetResourceByToken fetches a resource by its public token UUID. func GetResourceByToken(ctx context.Context, db *sql.DB, token uuid.UUID) (*Resource, error) { - r := &Resource{} - err := db.QueryRowContext(ctx, ` - SELECT id, team_id, token, resource_type, name, connection_url, key_prefix, tier, - fingerprint, cloud_vendor, country_code, status, migration_status, - expires_at, storage_bytes, provider_resource_id, created_request_id, created_at - FROM resources WHERE token = $1 - `, token).Scan( - &r.ID, &r.TeamID, &r.Token, &r.ResourceType, &r.Name, &r.ConnectionURL, &r.KeyPrefix, - &r.Tier, &r.Fingerprint, &r.CloudVendor, &r.CountryCode, &r.Status, - &r.MigrationStatus, &r.ExpiresAt, &r.StorageBytes, &r.ProviderResourceID, &r.CreatedRequestID, &r.CreatedAt, - ) + row := db.QueryRowContext(ctx, `SELECT `+resourceColumns+` FROM resources WHERE token = $1`, token) + r, err := scanResource(row) if err == sql.ErrNoRows { return nil, &ErrResourceNotFound{Token: token.String()} } @@ -124,12 +162,11 @@ func GetResourceByToken(ctx context.Context, db *sql.DB, token uuid.UUID) (*Reso // GetActiveResourceByFingerprintType finds the most recent active anonymous resource // of a specific type (e.g. "postgres", "redis", "mongodb") for a fingerprint. // Used by Phase 2+ handlers when the rate-limit is hit to return the existing resource. +// Anonymous resources are always env=production — there is no env switch on the +// dedup path, since anonymous callers don't pick an env. func GetActiveResourceByFingerprintType(ctx context.Context, db *sql.DB, fingerprint, resourceType string) (*Resource, error) { - r := &Resource{} - err := db.QueryRowContext(ctx, ` - SELECT id, team_id, token, resource_type, name, connection_url, key_prefix, tier, - fingerprint, cloud_vendor, country_code, status, migration_status, - expires_at, storage_bytes, created_request_id, created_at + row := db.QueryRowContext(ctx, ` + SELECT `+resourceColumns+` FROM resources WHERE fingerprint = $1 AND team_id IS NULL @@ -137,11 +174,9 @@ func GetActiveResourceByFingerprintType(ctx context.Context, db *sql.DB, fingerp AND status = 'active' ORDER BY created_at DESC LIMIT 1 - `, fingerprint, resourceType).Scan( - &r.ID, &r.TeamID, &r.Token, &r.ResourceType, &r.Name, &r.ConnectionURL, &r.KeyPrefix, - &r.Tier, &r.Fingerprint, &r.CloudVendor, &r.CountryCode, &r.Status, - &r.MigrationStatus, &r.ExpiresAt, &r.StorageBytes, &r.CreatedRequestID, &r.CreatedAt, - ) + `, fingerprint, resourceType) + + r, err := scanResource(row) if err == sql.ErrNoRows { return nil, &ErrResourceNotFound{Token: fingerprint} } @@ -155,9 +190,7 @@ func GetActiveResourceByFingerprintType(ctx context.Context, db *sql.DB, fingerp // Used when issuing an onboarding JWT to include all services provisioned in one session. func GetAllActiveResourcesByFingerprint(ctx context.Context, db *sql.DB, fingerprint string) ([]*Resource, error) { rows, err := db.QueryContext(ctx, ` - SELECT id, team_id, token, resource_type, name, connection_url, key_prefix, tier, - fingerprint, cloud_vendor, country_code, status, migration_status, - expires_at, storage_bytes, created_request_id, created_at + SELECT `+resourceColumns+` FROM resources WHERE fingerprint = $1 AND team_id IS NULL @@ -171,12 +204,8 @@ func GetAllActiveResourcesByFingerprint(ctx context.Context, db *sql.DB, fingerp var resources []*Resource for rows.Next() { - r := &Resource{} - if err := rows.Scan( - &r.ID, &r.TeamID, &r.Token, &r.ResourceType, &r.Name, &r.ConnectionURL, &r.KeyPrefix, - &r.Tier, &r.Fingerprint, &r.CloudVendor, &r.CountryCode, &r.Status, - &r.MigrationStatus, &r.ExpiresAt, &r.StorageBytes, &r.CreatedRequestID, &r.CreatedAt, - ); err != nil { + r, err := scanResource(rows) + if err != nil { return nil, fmt.Errorf("models.GetAllActiveResourcesByFingerprint: scan: %w", err) } resources = append(resources, r) @@ -195,12 +224,12 @@ func SoftDeleteResource(ctx context.Context, db *sql.DB, id uuid.UUID) error { return nil } -// ListResourcesByTeam returns all active resources for a team. +// ListResourcesByTeam returns all active resources for a team across every environment. +// Equivalent to ListResourcesByTeamAndEnv with env="" — kept as the dashboard's +// "give me everything I own" entry point. func ListResourcesByTeam(ctx context.Context, db *sql.DB, teamID uuid.UUID) ([]*Resource, error) { rows, err := db.QueryContext(ctx, ` - SELECT id, team_id, token, resource_type, name, connection_url, key_prefix, tier, - fingerprint, cloud_vendor, country_code, status, migration_status, - expires_at, storage_bytes, created_request_id, created_at + SELECT `+resourceColumns+` FROM resources WHERE team_id = $1 AND status != 'deleted' ORDER BY created_at DESC @@ -212,12 +241,8 @@ func ListResourcesByTeam(ctx context.Context, db *sql.DB, teamID uuid.UUID) ([]* var results []*Resource for rows.Next() { - r := &Resource{} - if err := rows.Scan( - &r.ID, &r.TeamID, &r.Token, &r.ResourceType, &r.Name, &r.ConnectionURL, &r.KeyPrefix, - &r.Tier, &r.Fingerprint, &r.CloudVendor, &r.CountryCode, &r.Status, - &r.MigrationStatus, &r.ExpiresAt, &r.StorageBytes, &r.CreatedRequestID, &r.CreatedAt, - ); err != nil { + r, err := scanResource(rows) + if err != nil { return nil, fmt.Errorf("models.ListResourcesByTeam scan: %w", err) } results = append(results, r) @@ -228,6 +253,38 @@ func ListResourcesByTeam(ctx context.Context, db *sql.DB, teamID uuid.UUID) ([]* return results, nil } +// ListResourcesByTeamAndEnv returns all active resources for a team filtered to +// a single environment. Empty env is normalised to "production" so callers that +// omit the param see prod resources by default. +func ListResourcesByTeamAndEnv(ctx context.Context, db *sql.DB, teamID uuid.UUID, env string) ([]*Resource, error) { + if env == "" { + env = EnvProduction + } + rows, err := db.QueryContext(ctx, ` + SELECT `+resourceColumns+` + FROM resources + WHERE team_id = $1 AND env = $2 AND status != 'deleted' + ORDER BY created_at DESC + `, teamID, env) + if err != nil { + return nil, fmt.Errorf("models.ListResourcesByTeamAndEnv: %w", err) + } + defer rows.Close() + + var results []*Resource + for rows.Next() { + r, err := scanResource(rows) + if err != nil { + return nil, fmt.Errorf("models.ListResourcesByTeamAndEnv scan: %w", err) + } + results = append(results, r) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("models.ListResourcesByTeamAndEnv rows: %w", err) + } + return results, nil +} + // UpdateConnectionURL replaces the encrypted connection_url for a resource. // Used exclusively by the credential rotation endpoint. func UpdateConnectionURL(ctx context.Context, db *sql.DB, resourceID uuid.UUID, encryptedURL string) error { @@ -274,6 +331,7 @@ func UpdateProviderResourceID(ctx context.Context, db *sql.DB, resourceID uuid.U // team to newTier. Called from the Razorpay upgrade webhook so that existing resources // benefit from higher limits immediately — not just resources provisioned after the upgrade. // Only affects permanent resources (expires_at IS NULL); anonymous TTL resources are excluded. +// Applies across ALL environments — an upgrade lifts dev, staging, and prod alike. func ElevateResourceTiersByTeam(ctx context.Context, db *sql.DB, teamID uuid.UUID, newTier string) error { _, err := db.ExecContext(ctx, ` UPDATE resources @@ -289,6 +347,7 @@ func ElevateResourceTiersByTeam(ctx context.Context, db *sql.DB, teamID uuid.UUI } // SumStorageBytesByTeamAndType returns total storage_bytes for active resources of a given type for a team. +// Sums across ALL environments — storage quotas are per-team, not per-env. func SumStorageBytesByTeamAndType(ctx context.Context, db *sql.DB, teamID uuid.UUID, resourceType string) (int64, error) { var total int64 err := db.QueryRowContext(ctx, @@ -318,4 +377,3 @@ func ExpireAnonymousResources(ctx context.Context, db *sql.DB) (int64, error) { n, _ := res.RowsAffected() return n, nil } - diff --git a/internal/models/resource_env_test.go b/internal/models/resource_env_test.go new file mode 100644 index 00000000..e258a6e8 --- /dev/null +++ b/internal/models/resource_env_test.go @@ -0,0 +1,212 @@ +package models_test + +// resource_env_test.go — env-column unit tests for the Resource model. +// +// The integration cases (TestResourceEnv_*) require a real Postgres; they +// skip when TEST_DATABASE_URL is unset. The pure-unit cases +// (TestNormalizeEnv_*) run anywhere. + +import ( + "context" + "os" + "strings" + "testing" + + "github.com/google/uuid" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "instant.dev/internal/models" + "instant.dev/internal/testhelpers" +) + +func TestNormalizeEnv_DefaultsToProduction(t *testing.T) { + got, ok := models.NormalizeEnv("") + assert.True(t, ok) + assert.Equal(t, models.EnvProduction, got) +} + +func TestNormalizeEnv_AcceptsValidValues(t *testing.T) { + cases := []string{ + "production", + "staging", + "dev", + "preview-42", + "a", + strings.Repeat("a", 32), + "my-feature-branch", + "qa1", + } + for _, in := range cases { + t.Run(in, func(t *testing.T) { + got, ok := models.NormalizeEnv(in) + assert.True(t, ok, "expected %q to be valid", in) + assert.Equal(t, in, got) + }) + } +} + +func TestNormalizeEnv_RejectsInvalidValues(t *testing.T) { + cases := []struct { + name string + input string + }{ + {"contains space", "prod ction"}, + {"contains uppercase", "Production"}, + {"contains exclamation", "prod!"}, + {"contains underscore", "my_env"}, + {"too long", strings.Repeat("a", 33)}, + {"unicode", "stagé"}, + {"slash", "dev/01"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + _, ok := models.NormalizeEnv(tc.input) + assert.False(t, ok, "expected %q to be rejected", tc.input) + }) + } +} + +// requireDB skips the test when TEST_DATABASE_URL isn't reachable. +// We can't just call testhelpers.SetupTestDB because it t.Fatalf's on connect +// errors, which we don't want for env-tests that should remain green on a +// laptop without postgres running. +func requireDB(t *testing.T) { + t.Helper() + if os.Getenv("TEST_DATABASE_URL") == "" { + t.Skip("TEST_DATABASE_URL not set; skipping integration test") + } +} + +func TestResourceEnv_CreateDefaultsToProduction(t *testing.T) { + requireDB(t) + db, cleanDB := testhelpers.SetupTestDB(t) + defer cleanDB() + + teamID := uuid.MustParse(testhelpers.MustCreateTeamDB(t, db, "hobby")) + defer db.Exec(`DELETE FROM teams WHERE id = $1`, teamID) + + r, err := models.CreateResource(context.Background(), db, models.CreateResourceParams{ + TeamID: &teamID, + ResourceType: "redis", + Tier: "hobby", + // Env intentionally empty — must default to "production". + }) + require.NoError(t, err) + defer db.Exec(`DELETE FROM resources WHERE id = $1`, r.ID) + + assert.Equal(t, models.EnvProduction, r.Env, + "empty Env on CreateResource must default to 'production'") +} + +func TestResourceEnv_CreateRoundTrips(t *testing.T) { + requireDB(t) + db, cleanDB := testhelpers.SetupTestDB(t) + defer cleanDB() + + teamID := uuid.MustParse(testhelpers.MustCreateTeamDB(t, db, "hobby")) + defer db.Exec(`DELETE FROM teams WHERE id = $1`, teamID) + + for _, env := range []string{"dev", "staging", "production", "preview-42"} { + t.Run(env, func(t *testing.T) { + r, err := models.CreateResource(context.Background(), db, models.CreateResourceParams{ + TeamID: &teamID, + ResourceType: "redis", + Tier: "hobby", + Env: env, + }) + require.NoError(t, err) + defer db.Exec(`DELETE FROM resources WHERE id = $1`, r.ID) + assert.Equal(t, env, r.Env) + + // GetResourceByToken must return the same env. + got, err := models.GetResourceByToken(context.Background(), db, r.Token) + require.NoError(t, err) + assert.Equal(t, env, got.Env) + }) + } +} + +func TestResourceEnv_ListByTeamAndEnv_Isolates(t *testing.T) { + requireDB(t) + db, cleanDB := testhelpers.SetupTestDB(t) + defer cleanDB() + + teamID := uuid.MustParse(testhelpers.MustCreateTeamDB(t, db, "hobby")) + defer db.Exec(`DELETE FROM teams WHERE id = $1`, teamID) + + mk := func(env string) *models.Resource { + r, err := models.CreateResource(context.Background(), db, models.CreateResourceParams{ + TeamID: &teamID, + ResourceType: "redis", + Tier: "hobby", + Env: env, + }) + require.NoError(t, err) + return r + } + + dev := mk("dev") + staging := mk("staging") + prod := mk("production") + defer db.Exec(`DELETE FROM resources WHERE id IN ($1, $2, $3)`, dev.ID, staging.ID, prod.ID) + + // Listing by env="dev" must only see the dev row. + devList, err := models.ListResourcesByTeamAndEnv(context.Background(), db, teamID, "dev") + require.NoError(t, err) + assert.Len(t, devList, 1) + assert.Equal(t, dev.ID, devList[0].ID) + + // Empty env defaults to production. + prodList, err := models.ListResourcesByTeamAndEnv(context.Background(), db, teamID, "") + require.NoError(t, err) + assert.Len(t, prodList, 1) + assert.Equal(t, prod.ID, prodList[0].ID) + + // ListResourcesByTeam (no env filter) must see all three. + all, err := models.ListResourcesByTeam(context.Background(), db, teamID) + require.NoError(t, err) + assert.Len(t, all, 3) +} + +// TestResourceEnv_MigrationIdempotent verifies that the columns + indexes are +// already present on a SetupTestDB instance and that re-applying the column-add +// statements is a no-op (no error, schema unchanged). We mimic the migration +// SQL directly rather than re-running 009 to keep this test independent of the +// embed.FS plumbing. +func TestResourceEnv_MigrationIdempotent(t *testing.T) { + requireDB(t) + db, cleanDB := testhelpers.SetupTestDB(t) + defer cleanDB() + + stmts := []string{ + `ALTER TABLE resources ADD COLUMN IF NOT EXISTS env TEXT NOT NULL DEFAULT 'production'`, + `ALTER TABLE deployments ADD COLUMN IF NOT EXISTS env TEXT NOT NULL DEFAULT 'production'`, + `CREATE INDEX IF NOT EXISTS idx_resources_team_env ON resources (team_id, env)`, + `CREATE INDEX IF NOT EXISTS idx_deployments_team_env ON deployments (team_id, env)`, + } + // Run twice; second run must not error. + for i := 0; i < 2; i++ { + for _, s := range stmts { + _, err := db.Exec(s) + require.NoError(t, err, "iteration %d: %s", i, s) + } + } + + // New rows inserted without env get 'production' from the column DEFAULT. + teamID := uuid.MustParse(testhelpers.MustCreateTeamDB(t, db, "hobby")) + defer db.Exec(`DELETE FROM teams WHERE id = $1`, teamID) + + var rid uuid.UUID + err := db.QueryRow(` + INSERT INTO resources (team_id, resource_type, tier) + VALUES ($1, 'redis', 'hobby') + RETURNING id + `, teamID).Scan(&rid) + require.NoError(t, err) + defer db.Exec(`DELETE FROM resources WHERE id = $1`, rid) + + var env string + require.NoError(t, db.QueryRow(`SELECT env FROM resources WHERE id = $1`, rid).Scan(&env)) + assert.Equal(t, "production", env, "DEFAULT must populate env when caller omits it") +} diff --git a/internal/models/team.go b/internal/models/team.go index 35728e61..f16c839d 100644 --- a/internal/models/team.go +++ b/internal/models/team.go @@ -163,7 +163,13 @@ func GetUserByGitHubID(ctx context.Context, db *sql.DB, githubID string) (*User, } // UpdateRazorpaySubscriptionID stores the Razorpay subscription ID on the team. -// Uses the existing stripe_customer_id column (renamed at DB layer later if needed). +// +// TODO: rename column stripe_customer_id → razorpay_subscription_id in a +// future migration. Stripe is not used anywhere in this codebase; the column +// name is a vestige of the original Stripe integration before the switch to +// Razorpay. Razorpay covers all payment surfaces we need (subscriptions, +// webhooks, invoices, plan upgrades). Per the user's directive, treat any +// remaining "stripe_*" string in the schema as legacy ballast to migrate. func UpdateRazorpaySubscriptionID(ctx context.Context, db *sql.DB, teamID uuid.UUID, subscriptionID string) error { _, err := db.ExecContext(ctx, ` UPDATE teams SET stripe_customer_id = $1 WHERE id = $2 diff --git a/internal/models/team_invitations.go b/internal/models/team_invitations.go new file mode 100644 index 00000000..10983126 --- /dev/null +++ b/internal/models/team_invitations.go @@ -0,0 +1,338 @@ +package models + +import ( + "context" + "crypto/rand" + "database/sql" + "encoding/hex" + "errors" + "fmt" + "time" + + "github.com/google/uuid" + "github.com/lib/pq" +) + +// RBAC role constants. Hierarchy: owner > admin > developer > viewer. +// "member" is retained as an alias of "developer" for legacy callers. +const ( + RoleOwner = "owner" + RoleAdmin = "admin" + RoleDeveloper = "developer" + RoleViewer = "viewer" +) + +// inviteTokenBytes is the random-byte length of an invitation token. +// 32 bytes -> 64 hex chars; must align with the migration column type. +const inviteTokenBytes = 32 + +// inviteTTL is how long a fresh invitation remains valid before expiry. +const inviteTTL = 7 * 24 * time.Hour + +// allowedInviteRoles is the closed set of roles that may be invited via the +// token-based RBAC flow. Owner cannot be invited — ownership is transferred, +// never granted via email. +var allowedInviteRoles = map[string]struct{}{ + RoleAdmin: {}, + RoleDeveloper: {}, + RoleViewer: {}, +} + +// Errors specific to the token-based RBAC invite flow. +var ( + ErrInvitationAlreadyAccepted = errors.New("invitation already accepted") + ErrInvitationRevoked = errors.New("invitation revoked") + ErrInvitationTokenInvalid = errors.New("invitation token invalid") + ErrLastOwner = errors.New("cannot remove or downgrade the last team owner") +) + +// RBACInvitation is the row shape for the token-based invite flow. +// Distinct from TeamInvitation (legacy "owner/member" + status string) so the +// two flows can coexist without name collisions. +type RBACInvitation struct { + ID uuid.UUID + TeamID uuid.UUID + Email string + Role string + Token string + InvitedBy uuid.UUID + ExpiresAt time.Time + AcceptedAt sql.NullTime + CreatedAt time.Time +} + +// IsValidInviteRole reports whether role can be granted via the invite flow. +func IsValidInviteRole(role string) bool { + _, ok := allowedInviteRoles[role] + return ok +} + +// generateInviteToken returns a cryptographically random hex token. +// Exposed via package var so tests can stub it deterministically. +var generateInviteToken = func() (string, error) { + buf := make([]byte, inviteTokenBytes) + if _, err := rand.Read(buf); err != nil { + return "", fmt.Errorf("models.generateInviteToken: %w", err) + } + return hex.EncodeToString(buf), nil +} + +// CreateRBACInvitation inserts a single-use invitation row, expiring in 7 days. +// invitedBy must already exist (FK to users). Returns the inserted row including +// the token (caller is responsible for emailing it to the invitee). +func CreateRBACInvitation(ctx context.Context, db *sql.DB, teamID uuid.UUID, email, role string, invitedBy uuid.UUID) (*RBACInvitation, error) { + email = NormalizeTeamEmail(email) + if email == "" { + return nil, fmt.Errorf("models.CreateRBACInvitation: email required") + } + if !IsValidInviteRole(role) { + return nil, ErrInvalidInviteRole + } + + token, err := generateInviteToken() + if err != nil { + return nil, err + } + + expiresAt := time.Now().Add(inviteTTL) + + inv := &RBACInvitation{} + err = db.QueryRowContext(ctx, ` + INSERT INTO team_invitations (team_id, email, role, token, invited_by, expires_at, status) + VALUES ($1, $2, $3, $4, $5, $6, 'pending') + RETURNING id, team_id, email, role, token, invited_by, expires_at, accepted_at, created_at + `, teamID, email, role, token, invitedBy, expiresAt).Scan( + &inv.ID, &inv.TeamID, &inv.Email, &inv.Role, &inv.Token, + &inv.InvitedBy, &inv.ExpiresAt, &inv.AcceptedAt, &inv.CreatedAt, + ) + if err != nil { + var pqErr *pq.Error + if errors.As(err, &pqErr) && pqErr.Code == "23505" { + return nil, ErrDuplicatePendingInvite + } + return nil, fmt.Errorf("models.CreateRBACInvitation: %w", err) + } + return inv, nil +} + +// ListRBACInvitations returns pending (status='pending', not yet accepted) invites +// for the team. Mirrors ListInvitations but populates the token + accepted_at fields. +func ListRBACInvitations(ctx context.Context, db *sql.DB, teamID uuid.UUID) ([]RBACInvitation, error) { + rows, err := db.QueryContext(ctx, ` + SELECT id, team_id, email, role, token, invited_by, expires_at, accepted_at, created_at + FROM team_invitations + WHERE team_id = $1 AND status = 'pending' AND accepted_at IS NULL + ORDER BY created_at DESC + `, teamID) + if err != nil { + return nil, fmt.Errorf("models.ListRBACInvitations: %w", err) + } + defer rows.Close() + + var out []RBACInvitation + for rows.Next() { + var inv RBACInvitation + if err := rows.Scan(&inv.ID, &inv.TeamID, &inv.Email, &inv.Role, &inv.Token, + &inv.InvitedBy, &inv.ExpiresAt, &inv.AcceptedAt, &inv.CreatedAt); err != nil { + return nil, fmt.Errorf("models.ListRBACInvitations: %w", err) + } + out = append(out, inv) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("models.ListRBACInvitations: %w", err) + } + return out, nil +} + +// GetRBACInvitationByID loads a single invitation by ID (ignoring status). +func GetRBACInvitationByID(ctx context.Context, db *sql.DB, id uuid.UUID) (*RBACInvitation, error) { + inv := &RBACInvitation{} + err := db.QueryRowContext(ctx, ` + SELECT id, team_id, email, role, token, invited_by, expires_at, accepted_at, created_at + FROM team_invitations WHERE id = $1 + `, id).Scan( + &inv.ID, &inv.TeamID, &inv.Email, &inv.Role, &inv.Token, + &inv.InvitedBy, &inv.ExpiresAt, &inv.AcceptedAt, &inv.CreatedAt, + ) + if err == sql.ErrNoRows { + return nil, ErrInvitationNotFound + } + if err != nil { + return nil, fmt.Errorf("models.GetRBACInvitationByID: %w", err) + } + return inv, nil +} + +// GetRBACInvitationByToken loads an invitation by its single-use token. +func GetRBACInvitationByToken(ctx context.Context, db *sql.DB, token string) (*RBACInvitation, error) { + if token == "" { + return nil, ErrInvitationTokenInvalid + } + inv := &RBACInvitation{} + err := db.QueryRowContext(ctx, ` + SELECT id, team_id, email, role, token, invited_by, expires_at, accepted_at, created_at + FROM team_invitations WHERE token = $1 + `, token).Scan( + &inv.ID, &inv.TeamID, &inv.Email, &inv.Role, &inv.Token, + &inv.InvitedBy, &inv.ExpiresAt, &inv.AcceptedAt, &inv.CreatedAt, + ) + if err == sql.ErrNoRows { + return nil, ErrInvitationNotFound + } + if err != nil { + return nil, fmt.Errorf("models.GetRBACInvitationByToken: %w", err) + } + return inv, nil +} + +// RevokeRBACInvitation marks an invitation revoked. Only pending invites +// (no accepted_at) can be revoked. +func RevokeRBACInvitation(ctx context.Context, db *sql.DB, invitationID uuid.UUID) error { + res, err := db.ExecContext(ctx, ` + UPDATE team_invitations SET status = 'revoked' + WHERE id = $1 AND status = 'pending' AND accepted_at IS NULL + `, invitationID) + if err != nil { + return fmt.Errorf("models.RevokeRBACInvitation: %w", err) + } + n, _ := res.RowsAffected() + if n == 0 { + return ErrInvitationNotFound + } + return nil +} + +// AcceptRBACInvitationByToken consumes a token, creating or updating the +// invitee's user row to belong to the team with the invited role. +// +// Single-use guarantee: the UPDATE is gated on accepted_at IS NULL — a second +// call against the same token returns ErrInvitationAlreadyAccepted. +// +// Expiry: rejects if expires_at < now, returning ErrInvitationExpired. +// +// Returns the user (existing or freshly created) so the caller can mint a +// session JWT for the invitee. +func AcceptRBACInvitationByToken(ctx context.Context, db *sql.DB, token string) (*User, *RBACInvitation, error) { + inv, err := GetRBACInvitationByToken(ctx, db, token) + if err != nil { + return nil, nil, err + } + // Already accepted -> 410 Gone (signal: token is permanently spent). + if inv.AcceptedAt.Valid { + return nil, inv, ErrInvitationAlreadyAccepted + } + if inv.Status() == "revoked" { + return nil, inv, ErrInvitationRevoked + } + if time.Now().After(inv.ExpiresAt) { + return nil, inv, ErrInvitationExpired + } + + tx, err := db.BeginTx(ctx, nil) + if err != nil { + return nil, nil, fmt.Errorf("models.AcceptRBACInvitationByToken: begin: %w", err) + } + defer func() { _ = tx.Rollback() }() + + // Atomic single-use guard: only one transaction can flip accepted_at from NULL. + res, err := tx.ExecContext(ctx, ` + UPDATE team_invitations SET accepted_at = now(), status = 'accepted' + WHERE id = $1 AND accepted_at IS NULL AND status = 'pending' + `, inv.ID) + if err != nil { + return nil, nil, fmt.Errorf("models.AcceptRBACInvitationByToken: update: %w", err) + } + if n, _ := res.RowsAffected(); n == 0 { + return nil, inv, ErrInvitationAlreadyAccepted + } + + // Look up an existing user by email; create one if none exists. + u := &User{} + err = tx.QueryRowContext(ctx, ` + SELECT id, team_id, email, COALESCE(role, 'member'), github_id, google_id, created_at + FROM users WHERE lower(email) = lower($1) + `, inv.Email).Scan( + &u.ID, &u.TeamID, &u.Email, &u.Role, &u.GitHubID, &u.GoogleID, &u.CreatedAt, + ) + if err == sql.ErrNoRows { + // Create the user attached to the team with the invited role. + err = tx.QueryRowContext(ctx, ` + INSERT INTO users (team_id, email, role) VALUES ($1, $2, $3) + RETURNING id, team_id, email, role, github_id, google_id, created_at + `, inv.TeamID, inv.Email, inv.Role).Scan( + &u.ID, &u.TeamID, &u.Email, &u.Role, &u.GitHubID, &u.GoogleID, &u.CreatedAt, + ) + if err != nil { + return nil, nil, fmt.Errorf("models.AcceptRBACInvitationByToken: insert user: %w", err) + } + } else if err != nil { + return nil, nil, fmt.Errorf("models.AcceptRBACInvitationByToken: lookup user: %w", err) + } else { + // Existing user — move them to the invited team and assign the new role. + // Refuse to silently downgrade an owner of *another* team without first + // vetting last-owner protection on the old team. For now we just move + // them; tighter policy can layer on later. + _, err = tx.ExecContext(ctx, ` + UPDATE users SET team_id = $1, role = $2 WHERE id = $3 + `, inv.TeamID, inv.Role, u.ID) + if err != nil { + return nil, nil, fmt.Errorf("models.AcceptRBACInvitationByToken: update user: %w", err) + } + u.TeamID = uuid.NullUUID{UUID: inv.TeamID, Valid: true} + u.Role = inv.Role + } + + if err := tx.Commit(); err != nil { + return nil, nil, fmt.Errorf("models.AcceptRBACInvitationByToken: commit: %w", err) + } + return u, inv, nil +} + +// Status returns the canonical lifecycle string for the invitation. +// Shadowed onto the type so handlers don't need a separate column lookup. +func (inv *RBACInvitation) Status() string { + if inv == nil { + return "" + } + if inv.AcceptedAt.Valid { + return "accepted" + } + if time.Now().After(inv.ExpiresAt) { + return "expired" + } + return "pending" +} + +// CountTeamOwners returns the number of users with role='owner' on the team. +// Used to enforce the "last owner cannot leave or be downgraded" invariant. +func CountTeamOwners(ctx context.Context, db *sql.DB, teamID uuid.UUID) (int, error) { + var n int + err := db.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM users WHERE team_id = $1 AND role = 'owner' + `, teamID).Scan(&n) + if err != nil { + return 0, fmt.Errorf("models.CountTeamOwners: %w", err) + } + return n, nil +} + +// EnsureNotLastOwner returns ErrLastOwner if removing/downgrading targetUserID +// from teamID would leave the team with zero owners. Callers should invoke +// this before any DELETE / role-downgrade affecting an owner. +func EnsureNotLastOwner(ctx context.Context, db *sql.DB, teamID, targetUserID uuid.UUID) error { + role, err := GetUserRole(ctx, db, teamID, targetUserID) + if err != nil { + return err + } + if role != RoleOwner { + return nil + } + count, err := CountTeamOwners(ctx, db, teamID) + if err != nil { + return err + } + if count <= 1 { + return ErrLastOwner + } + return nil +} diff --git a/internal/models/vault.go b/internal/models/vault.go new file mode 100644 index 00000000..2c1f5a4d --- /dev/null +++ b/internal/models/vault.go @@ -0,0 +1,205 @@ +package models + +import ( + "context" + "database/sql" + "errors" + "fmt" + "time" + + "github.com/google/uuid" +) + +// VaultSecret is one versioned row in vault_secrets. +// +// EncryptedValue stores AES-256-GCM ciphertext as raw bytes. The base64 string +// produced by crypto.Encrypt is decoded before insertion and re-encoded on read, +// so the at-rest format is opaque binary. +type VaultSecret struct { + ID uuid.UUID + TeamID uuid.UUID + Env string + Key string + EncryptedValue []byte + Version int + CreatedBy uuid.NullUUID + CreatedAt time.Time + UpdatedAt time.Time +} + +// VaultAuditEntry is one row in vault_audit_log. +type VaultAuditEntry struct { + ID int64 + TeamID uuid.UUID + UserID uuid.NullUUID + Action string + Env string + SecretKey string + IP sql.NullString + TS time.Time +} + +// ErrVaultSecretNotFound is returned when a vault secret cannot be located for +// the given (team, env, key[, version]). Handlers translate this to 404, never +// 403, to avoid leaking the existence of secrets owned by other teams. +var ErrVaultSecretNotFound = errors.New("vault secret not found") + +// CreateVaultSecret inserts a new row at version=nextVersion(team,env,key). +// Returns the created row. A unique-constraint violation on (team_id,env,key,version) +// is treated as a transient race and returned as-is. +func CreateVaultSecret(ctx context.Context, db *sql.DB, teamID uuid.UUID, env, key string, ciphertext []byte, createdBy uuid.NullUUID) (*VaultSecret, error) { + // Determine next version atomically using SELECT … FROM vault_secrets + // inside the INSERT (subselect avoids a separate round trip). + row := db.QueryRowContext(ctx, ` + INSERT INTO vault_secrets (team_id, env, key, encrypted_value, version, created_by) + VALUES ( + $1, $2, $3, $4, + COALESCE((SELECT MAX(version) FROM vault_secrets WHERE team_id = $1 AND env = $2 AND key = $3), 0) + 1, + $5 + ) + RETURNING id, team_id, env, key, encrypted_value, version, created_by, created_at, updated_at + `, teamID, env, key, ciphertext, createdBy) + + s := &VaultSecret{} + if err := row.Scan(&s.ID, &s.TeamID, &s.Env, &s.Key, &s.EncryptedValue, &s.Version, &s.CreatedBy, &s.CreatedAt, &s.UpdatedAt); err != nil { + return nil, fmt.Errorf("models.CreateVaultSecret: %w", err) + } + return s, nil +} + +// GetVaultSecretLatest returns the highest-version row scoped to (team,env,key). +// Returns ErrVaultSecretNotFound when the secret does not exist OR when team_id +// does not match (cross-team isolation: never leak existence). +func GetVaultSecretLatest(ctx context.Context, db *sql.DB, teamID uuid.UUID, env, key string) (*VaultSecret, error) { + s := &VaultSecret{} + err := db.QueryRowContext(ctx, ` + SELECT id, team_id, env, key, encrypted_value, version, created_by, created_at, updated_at + FROM vault_secrets + WHERE team_id = $1 AND env = $2 AND key = $3 + ORDER BY version DESC + LIMIT 1 + `, teamID, env, key).Scan( + &s.ID, &s.TeamID, &s.Env, &s.Key, &s.EncryptedValue, &s.Version, &s.CreatedBy, &s.CreatedAt, &s.UpdatedAt, + ) + if err == sql.ErrNoRows { + return nil, ErrVaultSecretNotFound + } + if err != nil { + return nil, fmt.Errorf("models.GetVaultSecretLatest: %w", err) + } + return s, nil +} + +// GetVaultSecretVersion returns a specific version of (team,env,key). +// Returns ErrVaultSecretNotFound when no row matches. +func GetVaultSecretVersion(ctx context.Context, db *sql.DB, teamID uuid.UUID, env, key string, version int) (*VaultSecret, error) { + s := &VaultSecret{} + err := db.QueryRowContext(ctx, ` + SELECT id, team_id, env, key, encrypted_value, version, created_by, created_at, updated_at + FROM vault_secrets + WHERE team_id = $1 AND env = $2 AND key = $3 AND version = $4 + `, teamID, env, key, version).Scan( + &s.ID, &s.TeamID, &s.Env, &s.Key, &s.EncryptedValue, &s.Version, &s.CreatedBy, &s.CreatedAt, &s.UpdatedAt, + ) + if err == sql.ErrNoRows { + return nil, ErrVaultSecretNotFound + } + if err != nil { + return nil, fmt.Errorf("models.GetVaultSecretVersion: %w", err) + } + return s, nil +} + +// ListVaultKeys returns the distinct keys for (team,env). Values are never returned — +// handlers must never expose a list endpoint that includes ciphertext. +func ListVaultKeys(ctx context.Context, db *sql.DB, teamID uuid.UUID, env string) ([]string, error) { + rows, err := db.QueryContext(ctx, ` + SELECT DISTINCT key FROM vault_secrets + WHERE team_id = $1 AND env = $2 + ORDER BY key ASC + `, teamID, env) + if err != nil { + return nil, fmt.Errorf("models.ListVaultKeys: %w", err) + } + defer rows.Close() + + keys := make([]string, 0) + for rows.Next() { + var k string + if err := rows.Scan(&k); err != nil { + return nil, fmt.Errorf("models.ListVaultKeys scan: %w", err) + } + keys = append(keys, k) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("models.ListVaultKeys rows: %w", err) + } + return keys, nil +} + +// DeleteVaultSecret performs a HARD delete of every version for (team,env,key). +// +// Semantics chosen for MVP: hard delete simplifies access control (no "deleted but +// still readable" state to enforce) and keeps the table small. Audit history is +// preserved separately in vault_audit_log so the deletion event itself is durable. +// +// Returns (rowsDeleted, error). rowsDeleted == 0 when the secret does not exist +// for this team — handlers turn that into 404 (idempotent delete, no leak). +func DeleteVaultSecret(ctx context.Context, db *sql.DB, teamID uuid.UUID, env, key string) (int64, error) { + res, err := db.ExecContext(ctx, ` + DELETE FROM vault_secrets + WHERE team_id = $1 AND env = $2 AND key = $3 + `, teamID, env, key) + if err != nil { + return 0, fmt.Errorf("models.DeleteVaultSecret: %w", err) + } + n, err := res.RowsAffected() + if err != nil { + return 0, fmt.Errorf("models.DeleteVaultSecret rows: %w", err) + } + return n, nil +} + +// AppendVaultAudit inserts one audit row. Errors are logged by callers; auditing +// must never block a request from completing (best-effort). +func AppendVaultAudit(ctx context.Context, db *sql.DB, teamID uuid.UUID, userID uuid.NullUUID, action, env, key, ip string) error { + var ipNS sql.NullString + if ip != "" { + ipNS = sql.NullString{String: ip, Valid: true} + } + _, err := db.ExecContext(ctx, ` + INSERT INTO vault_audit_log (team_id, user_id, action, env, secret_key, ip) + VALUES ($1, $2, $3, $4, $5, $6) + `, teamID, userID, action, env, key, ipNS) + if err != nil { + return fmt.Errorf("models.AppendVaultAudit: %w", err) + } + return nil +} + +// CountVaultKeysByTeam returns the number of distinct keys in the vault +// for a team. Used by handlers to enforce per-tier quotas. +func CountVaultKeysByTeam(ctx context.Context, db *sql.DB, teamID uuid.UUID) (int, error) { + var n int + err := db.QueryRowContext(ctx, ` + SELECT COUNT(DISTINCT key) FROM vault_secrets WHERE team_id = $1 + `, teamID).Scan(&n) + if err != nil { + return 0, fmt.Errorf("models.CountVaultKeysByTeam: %w", err) + } + return n, nil +} + +// CountVaultAudit returns the number of audit rows for (team, action, env, key). +// Used by tests to verify audit logging without exposing the full log surface. +func CountVaultAudit(ctx context.Context, db *sql.DB, teamID uuid.UUID, action, env, key string) (int, error) { + var n int + err := db.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM vault_audit_log + WHERE team_id = $1 AND action = $2 AND env = $3 AND secret_key = $4 + `, teamID, action, env, key).Scan(&n) + if err != nil { + return 0, fmt.Errorf("models.CountVaultAudit: %w", err) + } + return n, nil +} diff --git a/internal/plans/razorpay.go b/internal/plans/razorpay.go new file mode 100644 index 00000000..4c31ab0e --- /dev/null +++ b/internal/plans/razorpay.go @@ -0,0 +1,46 @@ +package plans + +import ( + "fmt" + "strings" +) + +// RazorpayPlanIDs maps "{tier}_{currency}_{cycle}" to a Razorpay plan ID. +// Currency and cycle are lowercase. +// +// USD plans charge via international cards (default for non-IST users). +// INR plans charge via Indian-issued cards (shown to Asia/Kolkata timezone). +// Razorpay enforces currency/card matching at payment time. +var RazorpayPlanIDs = map[string]string{ + "hobby_usd_monthly": "plan_Sg2YcWj6hM5Ook", + "hobby_usd_yearly": "plan_Sg2aCGFGoeuxNS", + "hobby_inr_monthly": "plan_SgT09xZkHcJing", + "hobby_inr_yearly": "plan_SgTAPVUusjHTB6", +} + +// LookupPlanID resolves a Razorpay plan ID from tier, currency, and cycle. +// Returns an error if no plan exists for the combination. +func LookupPlanID(tier, currency, cycle string) (string, error) { + key := fmt.Sprintf("%s_%s_%s", + strings.ToLower(tier), + strings.ToLower(currency), + strings.ToLower(cycle), + ) + id, ok := RazorpayPlanIDs[key] + if !ok { + return "", fmt.Errorf("no razorpay plan for %s", key) + } + return id, nil +} + +// TierFromPlanID reverses the map: given a Razorpay plan ID, returns the tier. +// Used by the webhook to determine what tier a subscription belongs to. +func TierFromPlanID(planID string) (string, bool) { + for key, id := range RazorpayPlanIDs { + if id == planID { + tier := strings.SplitN(key, "_", 2)[0] + return tier, true + } + } + return "", false +} diff --git a/internal/plans/razorpay_test.go b/internal/plans/razorpay_test.go new file mode 100644 index 00000000..1eccf2bb --- /dev/null +++ b/internal/plans/razorpay_test.go @@ -0,0 +1,97 @@ +package plans + +import "testing" + +func TestLookupPlanID(t *testing.T) { + cases := []struct { + name string + tier string + currency string + cycle string + wantID string + wantErr bool + }{ + {"hobby USD monthly", "hobby", "USD", "monthly", "plan_Sg2YcWj6hM5Ook", false}, + {"hobby USD yearly", "hobby", "USD", "yearly", "plan_Sg2aCGFGoeuxNS", false}, + {"hobby INR monthly", "hobby", "INR", "monthly", "plan_SgT09xZkHcJing", false}, + {"hobby INR yearly", "hobby", "INR", "yearly", "plan_SgTAPVUusjHTB6", false}, + {"lowercase currency works", "hobby", "usd", "monthly", "plan_Sg2YcWj6hM5Ook", false}, + {"mixed case cycle works", "hobby", "USD", "Monthly", "plan_Sg2YcWj6hM5Ook", false}, + {"unknown tier", "pro", "USD", "monthly", "", true}, + {"unknown currency", "hobby", "EUR", "monthly", "", true}, + {"unknown cycle", "hobby", "USD", "daily", "", true}, + {"empty currency", "hobby", "", "monthly", "", true}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := LookupPlanID(tc.tier, tc.currency, tc.cycle) + if tc.wantErr { + if err == nil { + t.Fatalf("expected error, got id=%q", got) + } + return + } + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if got != tc.wantID { + t.Fatalf("got %q, want %q", got, tc.wantID) + } + }) + } +} + +func TestTierFromPlanID(t *testing.T) { + cases := []struct { + planID string + wantTier string + wantOK bool + }{ + {"plan_Sg2YcWj6hM5Ook", "hobby", true}, + {"plan_Sg2aCGFGoeuxNS", "hobby", true}, + {"plan_SgT09xZkHcJing", "hobby", true}, + {"plan_SgTAPVUusjHTB6", "hobby", true}, + {"plan_SgT0sK508QF1iR", "", false}, // 2,499 typo plan — intentionally absent + {"plan_does_not_exist", "", false}, + {"", "", false}, + } + + for _, tc := range cases { + t.Run(tc.planID, func(t *testing.T) { + tier, ok := TierFromPlanID(tc.planID) + if ok != tc.wantOK { + t.Fatalf("ok: got %v, want %v", ok, tc.wantOK) + } + if tier != tc.wantTier { + t.Fatalf("tier: got %q, want %q", tier, tc.wantTier) + } + }) + } +} + +// TestRazorpayPlanIDs_AllUnique guards against a future edit accidentally +// pointing two keys at the same plan_id, which would corrupt TierFromPlanID. +func TestRazorpayPlanIDs_AllUnique(t *testing.T) { + seen := make(map[string]string) + for key, id := range RazorpayPlanIDs { + if prev, ok := seen[id]; ok { + t.Fatalf("duplicate plan_id %q used for both %q and %q", id, prev, key) + } + seen[id] = key + } +} + +// TestRazorpayPlanIDs_TypoPlanAbsent is a regression guard: the original +// hobby_inr_yearly plan was ₹2,499 (plan_SgT0sK508QF1iR, typo — negative +// discount vs monthly × 12). The replacement is ₹2,199 (plan_SgTAPVUusjHTB6). +// Razorpay plans cannot be deactivated, so we rely on code to never reference +// the bad one. +func TestRazorpayPlanIDs_TypoPlanAbsent(t *testing.T) { + const typoPlanID = "plan_SgT0sK508QF1iR" + for key, id := range RazorpayPlanIDs { + if id == typoPlanID { + t.Fatalf("typo plan %q must not be referenced (found at key %q)", typoPlanID, key) + } + } +} diff --git a/internal/providers/cache/redis.go b/internal/providers/cache/redis.go index af6cf97a..d8812c68 100644 --- a/internal/providers/cache/redis.go +++ b/internal/providers/cache/redis.go @@ -26,6 +26,11 @@ type Credentials struct { // Clients must prefix all keys with this value to stay in their namespace. // Empty when ACL-based isolation is used. KeyPrefix string + + // ProviderResourceID is the backend-specific resource identifier. + // For k8s-dedicated backend: the namespace name "instant-customer-<token>". + // Empty for the shared local backend. + ProviderResourceID string } // Provider manages Redis namespace provisioning. diff --git a/internal/providers/compute/k8s/client.go b/internal/providers/compute/k8s/client.go index 2d6c2c7e..154974c8 100644 --- a/internal/providers/compute/k8s/client.go +++ b/internal/providers/compute/k8s/client.go @@ -13,12 +13,12 @@ import ( "io" "log/slog" "os" - "os/exec" "path/filepath" "strings" "time" appsv1 "k8s.io/api/apps/v1" + batchv1 "k8s.io/api/batch/v1" corev1 "k8s.io/api/core/v1" networkingv1 "k8s.io/api/networking/v1" apierrors "k8s.io/apimachinery/pkg/api/errors" @@ -156,6 +156,12 @@ func (p *K8sProvider) createDeployNamespace(ctx context.Context, appID, tier str return p.setupTenantNamespace(ctx, deployNamespace(appID), appID, tier) } +// ptrProto / ptrPort — addressable temporaries for inline NetworkPolicyPort literals. +// Avoids the "address of unaddressable value" compile error when building Protocol/Port +// pointer fields without naming each one separately. +func ptrProto(p corev1.Protocol) *corev1.Protocol { return &p } +func ptrPort(p int) *intstr.IntOrString { v := intstr.FromInt(p); return &v } + // createNetworkPolicyInNS installs a default-deny NetworkPolicy in the given namespace // and adds targeted allow rules: // - Allow DNS egress to kube-system (UDP+TCP port 53) — required for hostname resolution @@ -213,6 +219,21 @@ func (p *K8sProvider) createNetworkPolicyInNS(ctx context.Context, ns string) er }, }, }, + { + // Allow ingress from nginx-ingress namespace. Required because + // Cilium-backed clusters (DOKS default) do NOT match in-cluster + // pod IPs against an "0.0.0.0/0" ipBlock — nginx-ingress traffic + // would otherwise be blocked. + From: []networkingv1.NetworkPolicyPeer{ + { + NamespaceSelector: &metav1.LabelSelector{ + MatchLabels: map[string]string{ + "kubernetes.io/metadata.name": "ingress-nginx", + }, + }, + }, + }, + }, { // Allow external ingress (NodePort traffic from the host / Lima VM). // Required when STACK_EXPOSE_VIA=nodeport; harmless when using Ingress. @@ -235,6 +256,48 @@ func (p *K8sProvider) createNetworkPolicyInNS(ctx context.Context, ns string) er }, }, }, + { + // Allow egress to dedicated DB pods in customer-resource namespaces + // on the data ports. Each /db/new, /cache/new, etc. creates a namespace + // labelled "instant.dev/role=customer-resource" — this rule lets the + // stack's app pods reach the postgres/redis/mongo/nats pod they `needs:`. + // Without this, Cilium-backed clusters (DOKS) silently drop service-IP + // traffic even though the broad `0.0.0.0/0` rule below ought to cover it. + To: []networkingv1.NetworkPolicyPeer{ + { + NamespaceSelector: &metav1.LabelSelector{ + MatchLabels: map[string]string{ + "instant.dev/role": "customer-resource", + }, + }, + }, + }, + Ports: []networkingv1.NetworkPolicyPort{ + {Protocol: ptrProto(corev1.ProtocolTCP), Port: ptrPort(5432)}, // postgres + {Protocol: ptrProto(corev1.ProtocolTCP), Port: ptrPort(6379)}, // redis + {Protocol: ptrProto(corev1.ProtocolTCP), Port: ptrPort(27017)}, // mongo + {Protocol: ptrProto(corev1.ProtocolTCP), Port: ptrPort(4222)}, // nats + }, + }, + { + // Allow egress to the `instant` namespace on data ports, so stacks can + // reach the in-cluster pg-proxy (and future redis/mongo/nats proxies). + To: []networkingv1.NetworkPolicyPeer{ + { + NamespaceSelector: &metav1.LabelSelector{ + MatchLabels: map[string]string{ + "kubernetes.io/metadata.name": "instant", + }, + }, + }, + }, + Ports: []networkingv1.NetworkPolicyPort{ + {Protocol: ptrProto(corev1.ProtocolTCP), Port: ptrPort(5432)}, + {Protocol: ptrProto(corev1.ProtocolTCP), Port: ptrPort(6379)}, + {Protocol: ptrProto(corev1.ProtocolTCP), Port: ptrPort(27017)}, + {Protocol: ptrProto(corev1.ProtocolTCP), Port: ptrPort(4222)}, + }, + }, { // Allow DNS resolution via kube-dns in kube-system (UDP + TCP port 53). // Without this, hostname resolution fails entirely. @@ -289,26 +352,27 @@ func (p *K8sProvider) createDefaultDenyNetworkPolicy(ctx context.Context, appID } // createResourceQuotaInNS installs a ResourceQuota in the given namespace. -// Limits vary by tier: -// - hobby: 256Mi RAM, 250m CPU, 5 pods max -// - pro: 512Mi RAM, 500m CPU, 10 pods max -// - team: 2Gi RAM, 2 CPU, 20 pods max +// Limits include headroom (~256Mi + 1 pod) for cert-manager HTTP-01 ACME +// solver pods that spawn briefly when issuing/renewing TLS certs. +// - hobby: 512Mi RAM, 500m CPU, 6 pods max +// - pro: 1Gi RAM, 1 CPU, 11 pods max +// - team: 3Gi RAM, 3 CPU, 21 pods max func (p *K8sProvider) createResourceQuotaInNS(ctx context.Context, ns, tier string) error { var memLimit, cpuLimit string var maxPods string switch tier { case "pro": - memLimit = "512Mi" - cpuLimit = "500m" - maxPods = "10" + memLimit = "1Gi" + cpuLimit = "1" + maxPods = "11" case "team": - memLimit = "2Gi" - cpuLimit = "2" - maxPods = "20" + memLimit = "3Gi" + cpuLimit = "3" + maxPods = "21" default: // hobby + anonymous - memLimit = "256Mi" - cpuLimit = "250m" - maxPods = "5" + memLimit = "512Mi" + cpuLimit = "500m" + maxPods = "6" } quota := &corev1.ResourceQuota{ @@ -399,7 +463,7 @@ func (p *K8sProvider) Deploy(ctx context.Context, opts compute.DeployOptions) (* ns := deployNamespace(opts.AppID) // Step 1: Build the Docker image from the tarball. - if err := p.buildImage(ctx, opts.AppID, imageTag, opts.Tarball); err != nil { + if err := p.buildImage(ctx, deployNamespace(opts.AppID), opts.AppID, imageTag, opts.Tarball); err != nil { return nil, fmt.Errorf("k8s.Deploy: build image: %w", err) } @@ -424,17 +488,31 @@ func (p *K8sProvider) Deploy(ctx context.Context, opts compute.DeployOptions) (* return nil, fmt.Errorf("k8s.Deploy: apply service: %w", err) } + // Step 8: Create Ingress (+ cert-manager TLS) when DEPLOY_DOMAIN is set. + // Falls back to the NodePort URL on local clusters that don't have an + // ingress controller or public domain configured. + ingressURL, err := p.applyIngressForDeploy(ctx, ns, svcName, opts.AppID, opts.Port) + if err != nil { + return nil, fmt.Errorf("k8s.Deploy: apply ingress: %w", err) + } + + publicURL := ingressURL + if publicURL == "" { + publicURL = appURL(nodePort) + } + slog.Info("k8s.Deploy: deployment created", "app_id", opts.AppID, "image", imageTag, "namespace", ns, "node_port", nodePort, + "ingress_url", ingressURL, + "url", publicURL, ) - appURL := appURL(nodePort) return &compute.AppDeployment{ ProviderID: deployName, - AppURL: appURL, + AppURL: publicURL, Status: "building", UpdatedAt: time.Now(), }, nil @@ -469,9 +547,16 @@ func (p *K8sProvider) Status(ctx context.Context, providerID string) (*compute.A nodePort = int(svc.Spec.Ports[0].NodePort) } + // Prefer the public Ingress URL when DEPLOY_DOMAIN is configured; fall + // back to the NodePort URL for local dev. + publicURL := deployIngressURL(appID) + if publicURL == "" { + publicURL = appURL(nodePort) + } + return &compute.AppDeployment{ ProviderID: providerID, - AppURL: appURL(nodePort), + AppURL: publicURL, Status: status, UpdatedAt: deploy.CreationTimestamp.Time, }, nil @@ -528,7 +613,7 @@ func (p *K8sProvider) Redeploy(ctx context.Context, providerID string, tarball [ imageTag := imageName(appID) ns := deployNamespace(appID) - if err := p.buildImage(ctx, appID, imageTag, tarball); err != nil { + if err := p.buildImage(ctx, deployNamespace(appID), appID, imageTag, tarball); err != nil { return nil, fmt.Errorf("k8s.Redeploy: build image: %w", err) } @@ -561,48 +646,235 @@ func (p *K8sProvider) Redeploy(ctx context.Context, providerID string, tarball [ nodePort = int(svc.Spec.Ports[0].NodePort) } + // Prefer the public Ingress URL when DEPLOY_DOMAIN is configured. + publicURL := deployIngressURL(appID) + if publicURL == "" { + publicURL = appURL(nodePort) + } + slog.Info("k8s.Redeploy: rolling update triggered", "provider_id", providerID, "namespace", ns, + "url", publicURL, ) return &compute.AppDeployment{ ProviderID: providerID, - AppURL: appURL(nodePort), + AppURL: publicURL, Status: "deploying", UpdatedAt: time.Now(), }, nil } -// buildImage extracts the tarball to a temp directory and runs docker build. -// Works on Rancher Desktop because k3s and Docker share the same image store. -func (p *K8sProvider) buildImage(ctx context.Context, appID, imageTag string, tarball []byte) error { - dir, err := os.MkdirTemp("", "instant-build-"+appID+"-*") - if err != nil { - return fmt.Errorf("create temp dir: %w", err) +// buildImage builds the user's container image using kaniko inside k8s and +// pushes it to the configured registry. Works on any k8s cluster (containerd, +// docker, etc.) because the build runs as a Pod, not a subprocess on a node. +// +// Caller passes ns explicitly because the stack flow uses +// "instant-stack-<id>" while the single-app flow uses "instant-deploy-<id>". +func (p *K8sProvider) buildImage(ctx context.Context, ns, appID, imageTag string, tarball []byte) error { + jobName := "build-" + sanitizeName(appID) + ctxSecret := "build-ctx-" + sanitizeName(appID) + authSecret := "ghcr-pull" + + slog.Info("k8s.buildImage: starting kaniko build", + "app_id", appID, "image", imageTag, "namespace", ns) + + // 0. Ensure the namespace exists. The stack pipeline normally creates it + // via setupTenantNamespace AFTER the build step, so we need to be the + // first to bring it up. Idempotent. + nsObj := &corev1.Namespace{ObjectMeta: metav1.ObjectMeta{ + Name: ns, + Labels: map[string]string{"managed-by": "instant.dev", "instant.dev/component": "build-staging"}, + }} + if _, err := p.clientset.CoreV1().Namespaces().Create(ctx, nsObj, metav1.CreateOptions{}); err != nil && !apierrors.IsAlreadyExists(err) { + return fmt.Errorf("k8s.buildImage: ensure namespace %q: %w", ns, err) } - defer os.RemoveAll(dir) - if err := extractTarGz(tarball, dir); err != nil { - return fmt.Errorf("extract tarball: %w", err) + // 1. Tarball as a Secret (kaniko reads via tar:// context). + if err := p.upsertBuildContextSecret(ctx, ns, ctxSecret, tarball); err != nil { + return fmt.Errorf("k8s.buildImage: build-context secret: %w", err) } - cmd := exec.CommandContext(ctx, "docker", "build", "-t", imageTag, dir) - cmd.Stdout = os.Stdout - cmd.Stderr = os.Stderr + // 2. Ensure registry auth secret exists in this namespace (copied from instant ns). + if err := p.ensureRegistryAuthInNS(ctx, ns, authSecret); err != nil { + return fmt.Errorf("k8s.buildImage: registry auth: %w", err) + } - slog.Info("k8s.buildImage: running docker build", - "app_id", appID, - "image", imageTag, - "dir", dir, - ) + // 3. Create the kaniko Job (delete first if it exists from a previous attempt). + prop := metav1.DeletePropagationBackground + _ = p.clientset.BatchV1().Jobs(ns).Delete(ctx, jobName, metav1.DeleteOptions{ + PropagationPolicy: &prop, + }) + if err := p.createKanikoJob(ctx, ns, jobName, ctxSecret, authSecret, imageTag); err != nil { + return fmt.Errorf("k8s.buildImage: create kaniko job: %w", err) + } - if err := cmd.Run(); err != nil { - return fmt.Errorf("docker build: %w", err) + // 4. Wait for Job completion (poll status). + if err := p.waitForJobComplete(ctx, ns, jobName, 10*time.Minute); err != nil { + return fmt.Errorf("k8s.buildImage: kaniko job: %w", err) } + + slog.Info("k8s.buildImage: kaniko build complete", "app_id", appID, "image", imageTag) return nil } +// sanitizeName lowercases and DNS-1123-cleans an appID for use in resource names. +func sanitizeName(s string) string { + out := make([]byte, 0, len(s)) + for i := 0; i < len(s); i++ { + c := s[i] + switch { + case c >= 'A' && c <= 'Z': + out = append(out, c+32) + case (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || c == '-': + out = append(out, c) + default: + out = append(out, '-') + } + } + return string(out) +} + +// upsertBuildContextSecret writes the tarball into a Secret under key "context.tar.gz". +func (p *K8sProvider) upsertBuildContextSecret(ctx context.Context, ns, name string, tarball []byte) error { + sec := &corev1.Secret{ + ObjectMeta: metav1.ObjectMeta{ + Name: name, + Labels: map[string]string{ + "app.kubernetes.io/managed-by": "instant", + "instant.dev/component": "build-context", + }, + }, + Data: map[string][]byte{"context.tar.gz": tarball}, + Type: corev1.SecretTypeOpaque, + } + _, err := p.clientset.CoreV1().Secrets(ns).Create(ctx, sec, metav1.CreateOptions{}) + if err == nil { + return nil + } + if !apierrors.IsAlreadyExists(err) { + return err + } + existing, err := p.clientset.CoreV1().Secrets(ns).Get(ctx, name, metav1.GetOptions{}) + if err != nil { + return fmt.Errorf("get existing: %w", err) + } + existing.Data = sec.Data + _, err = p.clientset.CoreV1().Secrets(ns).Update(ctx, existing, metav1.UpdateOptions{}) + return err +} + +// ensureRegistryAuthInNS copies the dockerconfigjson auth secret from the +// "instant" namespace into the deploy namespace if missing. +func (p *K8sProvider) ensureRegistryAuthInNS(ctx context.Context, ns, name string) error { + if _, err := p.clientset.CoreV1().Secrets(ns).Get(ctx, name, metav1.GetOptions{}); err == nil { + return nil + } + src, err := p.clientset.CoreV1().Secrets("instant").Get(ctx, name, metav1.GetOptions{}) + if err != nil { + return fmt.Errorf("source registry-auth secret %q in instant ns: %w", name, err) + } + dst := &corev1.Secret{ + ObjectMeta: metav1.ObjectMeta{Name: name}, + Type: src.Type, + Data: src.Data, + } + _, err = p.clientset.CoreV1().Secrets(ns).Create(ctx, dst, metav1.CreateOptions{}) + if err != nil && !apierrors.IsAlreadyExists(err) { + return err + } + return nil +} + +// createKanikoJob spawns a one-shot Job that builds and pushes the image. +func (p *K8sProvider) createKanikoJob(ctx context.Context, ns, jobName, ctxSecret, authSecret, imageTag string) error { + backoff := int32(0) + ttl := int32(300) + job := &batchv1.Job{ + ObjectMeta: metav1.ObjectMeta{ + Name: jobName, + Labels: map[string]string{ + "app.kubernetes.io/managed-by": "instant", + "instant.dev/component": "build", + }, + }, + Spec: batchv1.JobSpec{ + BackoffLimit: &backoff, + TTLSecondsAfterFinished: &ttl, + Template: corev1.PodTemplateSpec{ + Spec: corev1.PodSpec{ + RestartPolicy: corev1.RestartPolicyNever, + Containers: []corev1.Container{{ + Name: "kaniko", + Image: "gcr.io/kaniko-project/executor:v1.23.2", + Args: []string{ + "--context=tar:///workspace/context.tar.gz", + "--destination=" + imageTag, + "--snapshot-mode=redo", + "--cache=false", + "--single-snapshot", + "--cleanup", + }, + VolumeMounts: []corev1.VolumeMount{ + {Name: "build-context", MountPath: "/workspace"}, + {Name: "registry-auth", MountPath: "/kaniko/.docker"}, + }, + }}, + Volumes: []corev1.Volume{ + { + Name: "build-context", + VolumeSource: corev1.VolumeSource{ + Secret: &corev1.SecretVolumeSource{SecretName: ctxSecret}, + }, + }, + { + Name: "registry-auth", + VolumeSource: corev1.VolumeSource{ + Secret: &corev1.SecretVolumeSource{ + SecretName: authSecret, + Items: []corev1.KeyToPath{ + {Key: ".dockerconfigjson", Path: "config.json"}, + }, + }, + }, + }, + }, + }, + }, + }, + } + _, err := p.clientset.BatchV1().Jobs(ns).Create(ctx, job, metav1.CreateOptions{}) + return err +} + +// waitForJobComplete polls a Job until success or failure. +func (p *K8sProvider) waitForJobComplete(ctx context.Context, ns, jobName string, timeout time.Duration) error { + deadline := time.Now().Add(timeout) + for { + if time.Now().After(deadline) { + return fmt.Errorf("job %q timed out after %s", jobName, timeout) + } + job, err := p.clientset.BatchV1().Jobs(ns).Get(ctx, jobName, metav1.GetOptions{}) + if err != nil { + return fmt.Errorf("poll job: %w", err) + } + for _, c := range job.Status.Conditions { + if c.Type == batchv1.JobComplete && c.Status == corev1.ConditionTrue { + return nil + } + if c.Type == batchv1.JobFailed && c.Status == corev1.ConditionTrue { + return fmt.Errorf("job %q failed: %s", jobName, c.Message) + } + } + select { + case <-ctx.Done(): + return ctx.Err() + case <-time.After(3 * time.Second): + } + } +} + // applyDeploymentInNS creates or updates the k8s Deployment for an app in the // given namespace (the per-deployment namespace). func (p *K8sProvider) applyDeploymentInNS( @@ -643,6 +915,9 @@ func (p *K8sProvider) applyDeploymentInNS( Spec: corev1.PodSpec{ // Disable service account token auto-mount for security. AutomountServiceAccountToken: &saFalse, + ImagePullSecrets: []corev1.LocalObjectReference{ + {Name: "ghcr-pull"}, + }, Containers: []corev1.Container{ { Name: "app", @@ -738,6 +1013,105 @@ func (p *K8sProvider) applyServiceInNS(ctx context.Context, ns, name, deployName return nodePort, nil } +// applyIngressForDeploy creates an Ingress for a single-service /deploy/new app. +// +// Mirrors the pattern used by K8sStackProvider.createIngress: when DEPLOY_DOMAIN +// is set, the ingress is exposed at "<app-id>.<DEPLOY_DOMAIN>" and (if CERT_ISSUER +// is set) annotated for cert-manager so a Let's Encrypt cert is issued via the +// configured cluster-issuer (HTTP-01 by default). When DEPLOY_DOMAIN is empty +// (e.g. local Rancher Desktop), no ingress is created and the caller falls back +// to the NodePort URL. +// +// Returns the public URL on success, or "" if no ingress was created (callers +// should then fall back to the NodePort URL). +func (p *K8sProvider) applyIngressForDeploy(ctx context.Context, ns, svcName, appID string, port int) (string, error) { + domain := os.Getenv("DEPLOY_DOMAIN") + if domain == "" { + // No public domain configured — skip ingress creation (local dev path). + return "", nil + } + host := appID + "." + domain + pathType := networkingv1.PathTypePrefix + + annotations := map[string]string{} + var tls []networkingv1.IngressTLS + scheme := "http" + if certIssuer := os.Getenv("CERT_ISSUER"); certIssuer != "" { + annotations["cert-manager.io/cluster-issuer"] = certIssuer + tls = []networkingv1.IngressTLS{{ + Hosts: []string{host}, + SecretName: "app-" + appID + "-tls", + }} + scheme = "https" + } + publicURL := scheme + "://" + host + + ing := &networkingv1.Ingress{ + ObjectMeta: metav1.ObjectMeta{ + Name: "app-" + appID, + Namespace: ns, + Annotations: annotations, + Labels: map[string]string{ + labelApp: "true", + labelAppID: appID, + }, + }, + Spec: networkingv1.IngressSpec{ + TLS: tls, + Rules: []networkingv1.IngressRule{ + { + Host: host, + IngressRuleValue: networkingv1.IngressRuleValue{ + HTTP: &networkingv1.HTTPIngressRuleValue{ + Paths: []networkingv1.HTTPIngressPath{ + { + Path: "/", + PathType: &pathType, + Backend: networkingv1.IngressBackend{ + Service: &networkingv1.IngressServiceBackend{ + Name: svcName, + Port: networkingv1.ServiceBackendPort{ + Number: int32(port), + }, + }, + }, + }, + }, + }, + }, + }, + }, + }, + } + + _, err := p.clientset.NetworkingV1().Ingresses(ns).Create(ctx, ing, metav1.CreateOptions{}) + if err != nil { + if apierrors.IsAlreadyExists(err) { + return publicURL, nil + } + if apierrors.IsForbidden(err) { + return "", fmt.Errorf("create ingress %q in %q: RBAC forbidden — ensure the service account has networking.k8s.io/ingresses create permission: %w", "app-"+appID, ns, err) + } + return "", fmt.Errorf("create ingress %q in %q: %w", "app-"+appID, ns, err) + } + return publicURL, nil +} + +// deployIngressURL returns the public Ingress URL for an appID if DEPLOY_DOMAIN +// is configured. Caller uses this to compute the AppURL during Status/Redeploy +// without re-querying the k8s API (the value is deterministic from env + appID). +func deployIngressURL(appID string) string { + domain := os.Getenv("DEPLOY_DOMAIN") + if domain == "" { + return "" + } + scheme := "http" + if os.Getenv("CERT_ISSUER") != "" { + scheme = "https" + } + return scheme + "://" + appID + "." + domain +} + // deploymentStatus translates k8s Deployment conditions and replica counts into // one of: building|deploying|healthy|failed|stopped. func deploymentStatus(deploy *appsv1.Deployment) string { @@ -827,7 +1201,15 @@ func envVarsToK8s(vars map[string]string) []corev1.EnvVar { func deploymentName(appID string) string { return "app-" + appID } func serviceName(appID string) string { return "svc-" + appID } -func imageName(appID string) string { return imageRegistry + "/" + appID + ":latest" } +func imageName(appID string) string { + if reg := os.Getenv("BUILD_IMAGE_REGISTRY"); reg != "" { + for len(reg) > 0 && reg[len(reg)-1] == '/' { + reg = reg[:len(reg)-1] + } + return reg + "/" + appID + ":latest" + } + return imageRegistry + "/" + appID + ":latest" +} func appIDFromDeployName(name string) string { if len(name) > 4 && name[:4] == "app-" { diff --git a/internal/providers/compute/k8s/custom_domain.go b/internal/providers/compute/k8s/custom_domain.go new file mode 100644 index 00000000..176af824 --- /dev/null +++ b/internal/providers/compute/k8s/custom_domain.go @@ -0,0 +1,310 @@ +package k8s + +// custom_domain.go — k8s helpers for binding a customer-owned hostname to a +// stack service. Lives alongside the stack provider so the underlying +// clientset is reused without additional plumbing. +// +// Two callers expect to use these: +// +// 1. The custom-domain handler, after TXT verification succeeds, calls +// EnsureCustomDomainIngress to create / update an Ingress for the +// hostname. cert-manager picks up the cluster-issuer annotation and +// issues a real cert. +// +// 2. The same handler polls CertificateReady to surface "cert is live yet?" +// to the dashboard / API caller. cert-manager Certificates are CRDs, so +// we use a dynamic client (no need to vendor cert-manager Go types). +// +// The Ingress secretName follows a deterministic pattern so re-creating the +// row produces an idempotent k8s update, not a duplicate. + +import ( + "context" + "fmt" + "os" + "strings" + + networkingv1 "k8s.io/api/networking/v1" + apierrors "k8s.io/apimachinery/pkg/api/errors" + metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" + "k8s.io/apimachinery/pkg/runtime/schema" + "k8s.io/client-go/dynamic" + "k8s.io/client-go/rest" + "k8s.io/client-go/tools/clientcmd" +) + +// certManagerCertificateGVR is the GroupVersionResource for +// cert-manager.io/v1 Certificate. Held as a package-level var so tests can +// override (e.g. point at a fake CRD). +var certManagerCertificateGVR = schema.GroupVersionResource{ + Group: "cert-manager.io", + Version: "v1", + Resource: "certificates", +} + +// sanitizeHostname turns a customer-supplied hostname into a DNS-1123 fragment +// safe for use as a k8s resource name suffix. ASCII letters / digits stay, +// dots become dashes, everything else collapses to a dash. +// +// Example: "App.Acme.com" -> "app-acme-com" +func sanitizeHostname(host string) string { + host = strings.ToLower(strings.TrimSpace(host)) + out := make([]byte, 0, len(host)) + for i := 0; i < len(host); i++ { + c := host[i] + switch { + case c >= 'a' && c <= 'z', + c >= '0' && c <= '9': + out = append(out, c) + default: + out = append(out, '-') + } + } + // Collapse repeats and trim leading / trailing dashes. + collapsed := make([]byte, 0, len(out)) + prevDash := true // treat start as if previous was a dash (trim leading) + for _, c := range out { + if c == '-' { + if prevDash { + continue + } + prevDash = true + } else { + prevDash = false + } + collapsed = append(collapsed, c) + } + for len(collapsed) > 0 && collapsed[len(collapsed)-1] == '-' { + collapsed = collapsed[:len(collapsed)-1] + } + return string(collapsed) +} + +// CustomDomainIngressName returns the k8s Ingress name for a custom-domain +// binding. The base service name is included so a single stack service can +// host more than one hostname. +func CustomDomainIngressName(svcName, hostname string) string { + return "cdom-" + svcName + "-" + sanitizeHostname(hostname) +} + +// CustomDomainTLSSecretName returns the k8s Secret name where cert-manager +// will store the issued cert chain. The exported name is also the value of +// `tls.secretName` in the Ingress spec. +func CustomDomainTLSSecretName(hostname string) string { + return "cdom-" + sanitizeHostname(hostname) + "-tls" +} + +// EnsureCustomDomainIngress creates (or updates) an Ingress + cert-manager +// Certificate that routes https://hostname to (serviceName:servicePort) inside +// stackNamespace. Returns the Certificate resource name so callers can poll +// its readiness via CertificateReady. +// +// The Ingress is named per-(service, hostname) so a single namespace can hold +// the original deployment Ingress (`<slug>.deployment.instanode.dev`) plus +// any number of custom-domain Ingresses without colliding. +func (p *K8sStackProvider) EnsureCustomDomainIngress( + ctx context.Context, + stackNamespace, hostname, serviceName string, + servicePort int, +) (string, error) { + if hostname == "" { + return "", fmt.Errorf("k8s.EnsureCustomDomainIngress: hostname is required") + } + if serviceName == "" { + return "", fmt.Errorf("k8s.EnsureCustomDomainIngress: serviceName is required") + } + if servicePort == 0 { + servicePort = 8080 + } + + hostname = strings.ToLower(strings.TrimSpace(hostname)) + ingressName := CustomDomainIngressName(serviceName, hostname) + secretName := CustomDomainTLSSecretName(hostname) + pathType := networkingv1.PathTypePrefix + + // cert-manager wiring: HTTP-01 by default, overridable via CERT_ISSUER. + // The Certificate is created implicitly by cert-manager when it sees an + // Ingress with the cluster-issuer annotation + a TLS section pointing at + // a missing Secret. We do NOT manually CRUD the Certificate CRD here. + certIssuer := os.Getenv("CERT_ISSUER") + if certIssuer == "" { + certIssuer = "letsencrypt-http01" + } + + annotations := map[string]string{ + "cert-manager.io/cluster-issuer": certIssuer, + } + + desired := &networkingv1.Ingress{ + ObjectMeta: metav1.ObjectMeta{ + Name: ingressName, + Namespace: stackNamespace, + Annotations: annotations, + Labels: map[string]string{ + "app": serviceName, + "instant.dev/custom-domain": "true", + }, + }, + Spec: networkingv1.IngressSpec{ + TLS: []networkingv1.IngressTLS{{ + Hosts: []string{hostname}, + SecretName: secretName, + }}, + Rules: []networkingv1.IngressRule{{ + Host: hostname, + IngressRuleValue: networkingv1.IngressRuleValue{ + HTTP: &networkingv1.HTTPIngressRuleValue{ + Paths: []networkingv1.HTTPIngressPath{{ + Path: "/", + PathType: &pathType, + Backend: networkingv1.IngressBackend{ + Service: &networkingv1.IngressServiceBackend{ + Name: serviceName, + Port: networkingv1.ServiceBackendPort{ + Number: int32(servicePort), + }, + }, + }, + }}, + }, + }, + }}, + }, + } + + existing, err := p.clientset.NetworkingV1().Ingresses(stackNamespace).Get(ctx, ingressName, metav1.GetOptions{}) + if apierrors.IsNotFound(err) { + if _, createErr := p.clientset.NetworkingV1().Ingresses(stackNamespace).Create(ctx, desired, metav1.CreateOptions{}); createErr != nil { + if apierrors.IsForbidden(createErr) { + return "", fmt.Errorf("k8s.EnsureCustomDomainIngress: RBAC forbidden creating ingress %q in %q: %w", ingressName, stackNamespace, createErr) + } + return "", fmt.Errorf("k8s.EnsureCustomDomainIngress: create ingress %q: %w", ingressName, createErr) + } + // cert-manager names the Certificate after the TLS secret name when + // Ingress shim creates it. Return the secret name as the cert name. + return secretName, nil + } + if err != nil { + return "", fmt.Errorf("k8s.EnsureCustomDomainIngress: get ingress %q: %w", ingressName, err) + } + + // Update existing — preserve resourceVersion + apply our spec/annotations. + existing.Spec = desired.Spec + existing.Annotations = desired.Annotations + if existing.Labels == nil { + existing.Labels = map[string]string{} + } + for k, v := range desired.Labels { + existing.Labels[k] = v + } + if _, err := p.clientset.NetworkingV1().Ingresses(stackNamespace).Update(ctx, existing, metav1.UpdateOptions{}); err != nil { + return "", fmt.Errorf("k8s.EnsureCustomDomainIngress: update ingress %q: %w", ingressName, err) + } + return secretName, nil +} + +// DeleteCustomDomainIngress removes the Ingress and (best-effort) the TLS +// Secret for a custom-domain binding. cert-manager removes its Certificate +// CRD when the owning Ingress goes away in shim mode. +// +// Best-effort: not-found errors are swallowed so the caller can mark the +// row deleted in the DB even after a partial teardown. +func (p *K8sStackProvider) DeleteCustomDomainIngress( + ctx context.Context, + stackNamespace, hostname, serviceName string, +) error { + hostname = strings.ToLower(strings.TrimSpace(hostname)) + ingressName := CustomDomainIngressName(serviceName, hostname) + secretName := CustomDomainTLSSecretName(hostname) + + if err := p.clientset.NetworkingV1().Ingresses(stackNamespace).Delete(ctx, ingressName, metav1.DeleteOptions{}); err != nil && !apierrors.IsNotFound(err) { + return fmt.Errorf("k8s.DeleteCustomDomainIngress: delete ingress %q: %w", ingressName, err) + } + // TLS secret cleanup is best-effort — cert-manager's Ingress shim usually + // owns it, but on some installs it lingers. + _ = p.clientset.CoreV1().Secrets(stackNamespace).Delete(ctx, secretName, metav1.DeleteOptions{}) + return nil +} + +// CertificateReady returns whether the cert-manager Certificate named +// `certName` in `namespace` has condition Ready=True. The second return +// value is the human-readable message attached to the condition (used to +// surface stuck issuance to the caller). +// +// Uses the dynamic client so the API binary does not vendor cert-manager Go +// types — those would pull in their entire CRD module just for one field. +func (p *K8sStackProvider) CertificateReady( + ctx context.Context, + namespace, certName string, +) (bool, string, error) { + dyn, err := newDynamicClient() + if err != nil { + return false, "", fmt.Errorf("k8s.CertificateReady: dynamic client: %w", err) + } + obj, err := dyn.Resource(certManagerCertificateGVR).Namespace(namespace).Get(ctx, certName, metav1.GetOptions{}) + if err != nil { + if apierrors.IsNotFound(err) { + // cert-manager hasn't created the Certificate yet (shim races + // Ingress reconcile). Treat as not-ready, no error. + return false, "Certificate not yet created by cert-manager", nil + } + return false, "", fmt.Errorf("k8s.CertificateReady: get certificate %q: %w", certName, err) + } + + // Walk status.conditions for the Ready entry. + conds, found, err := unstructuredSlice(obj.Object, "status", "conditions") + if err != nil || !found { + return false, "Certificate has no status conditions yet", nil + } + for _, c := range conds { + condMap, ok := c.(map[string]interface{}) + if !ok { + continue + } + condType, _ := condMap["type"].(string) + if condType != "Ready" { + continue + } + condStatus, _ := condMap["status"].(string) + condMsg, _ := condMap["message"].(string) + return condStatus == "True", condMsg, nil + } + return false, "Certificate Ready condition not yet present", nil +} + +// newDynamicClient builds a dynamic.Interface using the same in-cluster / +// kubeconfig fallback chain as newClientset above. Kept as a free function +// so callers can construct ad-hoc clients without holding a K8sProvider. +func newDynamicClient() (dynamic.Interface, error) { + cfg, err := rest.InClusterConfig() + if err != nil { + cfg, err = clientcmd.BuildConfigFromFlags("", clientcmd.RecommendedHomeFile) + if err != nil { + return nil, fmt.Errorf("k8s dynamic config: %w", err) + } + } + return dynamic.NewForConfig(cfg) +} + +// unstructuredSlice digs out a []interface{} at the given nested map path. +// Mirrors the single helper from k8s.io/apimachinery/pkg/apis/meta/v1/unstructured +// but without the import — we only need it once. +func unstructuredSlice(obj map[string]interface{}, path ...string) ([]interface{}, bool, error) { + cur := interface{}(obj) + for _, key := range path { + m, ok := cur.(map[string]interface{}) + if !ok { + return nil, false, fmt.Errorf("path %v: expected map at %q", path, key) + } + next, ok := m[key] + if !ok { + return nil, false, nil + } + cur = next + } + out, ok := cur.([]interface{}) + if !ok { + return nil, false, fmt.Errorf("path %v: expected slice at end", path) + } + return out, true, nil +} diff --git a/internal/providers/compute/k8s/stack.go b/internal/providers/compute/k8s/stack.go index 8f74cf10..87390b4f 100644 --- a/internal/providers/compute/k8s/stack.go +++ b/internal/providers/compute/k8s/stack.go @@ -25,8 +25,8 @@ import ( ) const ( - labelStack = "instant.dev/stack" - stackIngHost = "instant.dev" + labelStack = "instant.dev/stack" + stackIngHostDefault = "instant.dev" ) // K8sStackProvider implements compute.StackProvider using the local k8s cluster. @@ -44,9 +44,19 @@ func NewStackProvider(namespace string) (*K8sStackProvider, error) { return &K8sStackProvider{K8sProvider: base}, nil } -// stackImageTag returns the docker image tag for a stack service. +// stackImageTag returns the docker image tag for a stack service. Honors +// BUILD_IMAGE_REGISTRY env so kaniko pushes to a real registry instead of +// the unqualified name (which kaniko interprets as docker.io/library/...). func stackImageTag(stackID, svcName string) string { - return "instant-stack-" + stackID + "-" + svcName + ":latest" + bare := "instant-stack-" + stackID + "-" + svcName + ":latest" + reg := os.Getenv("BUILD_IMAGE_REGISTRY") + if reg == "" { + return bare + } + for len(reg) > 0 && reg[len(reg)-1] == '/' { + reg = reg[:len(reg)-1] + } + return reg + "/" + bare } // DeployStack builds all images in parallel, creates the stack namespace with @@ -87,7 +97,7 @@ func (p *K8sStackProvider) DeployStack( onUpdate(svc.Name, "building", "", "") tag := stackImageTag(opts.StackID, svc.Name) - if err := p.buildImage(buildCtx, svc.Name+"-"+opts.StackID, tag, svc.Tarball); err != nil { + if err := p.buildImage(buildCtx, stackNamespace, svc.Name+"-"+opts.StackID, tag, svc.Tarball); err != nil { return fmt.Errorf("build %q: %w", svc.Name, err) } return nil @@ -249,7 +259,7 @@ func (p *K8sStackProvider) RedeployStack( onUpdate(svc.Name, "building", "", "") tag := stackImageTag(stackID, svc.Name) - if err := p.buildImage(buildCtx, svc.Name+"-"+stackID, tag, svc.Tarball); err != nil { + if err := p.buildImage(buildCtx, stackNamespace, svc.Name+"-"+stackID, tag, svc.Tarball); err != nil { return fmt.Errorf("rebuild %q: %w", svc.Name, err) } return nil @@ -336,6 +346,9 @@ func (p *K8sStackProvider) createStackDeployment( }, Spec: corev1.PodSpec{ AutomountServiceAccountToken: &saFalse, + ImagePullSecrets: []corev1.LocalObjectReference{ + {Name: "ghcr-pull"}, // copied into the deploy ns by buildImage + }, Containers: []corev1.Container{ { Name: svcName, @@ -472,20 +485,42 @@ func (p *K8sStackProvider) createNodePortService(ctx context.Context, ns, name s // createIngress creates a k8s Ingress for an exposed stack service. // Returns the app URL on success. func (p *K8sStackProvider) createIngress(ctx context.Context, ns, stackID, svcName string, port int) (string, error) { - host := svcName + "-" + stackID + "." + stackIngHost - appURL := "http://" + host + domain := os.Getenv("DEPLOY_DOMAIN") + if domain == "" { + domain = stackIngHostDefault + } + host := svcName + "-" + stackID + "." + domain pathType := networkingv1.PathTypePrefix + // cert-manager wiring. If CERT_ISSUER is set, every ingress gets a TLS + // section + the cluster-issuer annotation, and cert-manager auto-issues + // a real cert via the configured ACME solver (HTTP-01 by default). + certIssuer := os.Getenv("CERT_ISSUER") + annotations := map[string]string{} + var tls []networkingv1.IngressTLS + scheme := "http" + if certIssuer != "" { + annotations["cert-manager.io/cluster-issuer"] = certIssuer + tls = []networkingv1.IngressTLS{{ + Hosts: []string{host}, + SecretName: svcName + "-" + stackID + "-tls", + }} + scheme = "https" + } + appURL := scheme + "://" + host + ing := &networkingv1.Ingress{ ObjectMeta: metav1.ObjectMeta{ - Name: svcName, - Namespace: ns, + Name: svcName, + Namespace: ns, + Annotations: annotations, Labels: map[string]string{ "app": svcName, labelStack: stackID, }, }, Spec: networkingv1.IngressSpec{ + TLS: tls, Rules: []networkingv1.IngressRule{ { Host: host, diff --git a/internal/providers/nosql/mongo.go b/internal/providers/nosql/mongo.go index 45597eb0..d9c5354f 100644 --- a/internal/providers/nosql/mongo.go +++ b/internal/providers/nosql/mongo.go @@ -30,6 +30,11 @@ type Credentials struct { // DatabaseName is the name of the provisioned database. DatabaseName string + + // ProviderResourceID is the backend-specific resource identifier. + // For k8s-dedicated backend: the namespace name "instant-customer-<token>". + // Empty for the shared local backend. + ProviderResourceID string } // Provider manages MongoDB provisioning. diff --git a/internal/providers/storage/local.go b/internal/providers/storage/local.go index 19178061..738cf4d6 100644 --- a/internal/providers/storage/local.go +++ b/internal/providers/storage/local.go @@ -49,14 +49,21 @@ type Credentials struct { // Provider manages MinIO storage provisioning. type Provider struct { - madmClient *madmin.AdminClient - endpoint string // host:port, e.g. "minio.instant-data.svc.cluster.local:9000" - bucketName string // e.g. "instant-shared" + madmClient *madmin.AdminClient + endpoint string // internal host:port for admin/bucket ops, e.g. "minio.instant-data.svc.cluster.local:9000" + publicEndpoint string // host:port returned to customers (falls back to endpoint when empty) + bucketName string // e.g. "instant-shared" } // New creates a Provider backed by a MinIO admin client. -// endpoint is "host:port", rootUser/rootPassword are the MinIO root credentials. -func New(endpoint, rootUser, rootPassword, bucketName string) (*Provider, error) { +// +// endpoint is the cluster-internal "host:port" used for IAM/bucket admin calls. +// publicEndpoint is the customer-reachable address returned in BucketURL/Endpoint. +// Accepts either bare "host[:port]" (defaults to http://) or a scheme-prefixed +// "https://host" / "http://host[:port]" form for TLS-terminated public hostnames. +// When empty, it falls back to endpoint (legacy in-cluster behavior). +// rootUser/rootPassword are the MinIO root credentials. +func New(endpoint, publicEndpoint, rootUser, rootPassword, bucketName string) (*Provider, error) { if endpoint == "" { return nil, fmt.Errorf("storage: MinIO endpoint is required (MINIO_ENDPOINT)") } @@ -70,12 +77,40 @@ func New(endpoint, rootUser, rootPassword, bucketName string) (*Provider, error) } return &Provider{ - madmClient: madmClient, - endpoint: endpoint, - bucketName: bucketName, + madmClient: madmClient, + endpoint: endpoint, + publicEndpoint: publicEndpoint, + bucketName: bucketName, }, nil } +// customerEndpoint returns the host[:port] to surface to customers, stripped of +// any scheme. Falls back to the internal endpoint when no public override is set. +func (p *Provider) customerEndpoint() string { + raw := p.publicEndpoint + if raw == "" { + raw = p.endpoint + } + // Strip a leading scheme if present (e.g. "https://s3.instanode.dev" → "s3.instanode.dev"). + if i := strings.Index(raw, "://"); i >= 0 { + raw = raw[i+3:] + } + return strings.TrimRight(raw, "/") +} + +// customerScheme returns the URL scheme to surface to customers ("http" or "https"). +// Derived from publicEndpoint when it carries an explicit scheme; otherwise "http" +// to preserve in-cluster legacy behavior. +func (p *Provider) customerScheme() string { + if p.publicEndpoint == "" { + return "http" + } + if strings.HasPrefix(p.publicEndpoint, "https://") { + return "https" + } + return "http" +} + // Provision creates a MinIO IAM user scoped to a per-token prefix and returns // S3-compatible credentials. The caller can use any S3 SDK with the returned // endpoint, access key, secret, and prefix. @@ -119,8 +154,10 @@ func (p *Provider) Provision(ctx context.Context, token, tier string) (*Credenti return nil, fmt.Errorf("storage.Provision: SetPolicy %q → %q: %w", policyName, accessKeyID, err) } - bucketURL := fmt.Sprintf("http://%s/%s/%s", p.endpoint, p.bucketName, objectPrefix) - endpoint := fmt.Sprintf("http://%s", p.endpoint) + customerHost := p.customerEndpoint() + scheme := p.customerScheme() + bucketURL := fmt.Sprintf("%s://%s/%s/%s", scheme, customerHost, p.bucketName, objectPrefix) + endpoint := fmt.Sprintf("%s://%s", scheme, customerHost) slog.Info("storage.Provision: MinIO user created", "token", token, diff --git a/internal/providers/storage/local_test.go b/internal/providers/storage/local_test.go index 62d2867d..24fb6cae 100644 --- a/internal/providers/storage/local_test.go +++ b/internal/providers/storage/local_test.go @@ -11,7 +11,7 @@ import ( // TestNew_RequiresEndpoint verifies that New returns an error when endpoint is empty. func TestNew_RequiresEndpoint(t *testing.T) { - _, err := storageprovider.New("", "root", "password", "instant-shared") + _, err := storageprovider.New("", "", "root", "password", "instant-shared") require.Error(t, err, "New must fail when MinIO endpoint is empty") assert.Contains(t, err.Error(), "endpoint", "error must mention missing endpoint") } @@ -19,7 +19,7 @@ func TestNew_RequiresEndpoint(t *testing.T) { // TestNew_ValidEndpointSucceeds verifies that a non-empty endpoint produces a Provider. // madmin.New does not dial on construction — the connection is lazy. func TestNew_ValidEndpointSucceeds(t *testing.T) { - p, err := storageprovider.New("minio.example.local:9000", "minioadmin", "minioadmin123", "instant-shared") + p, err := storageprovider.New("minio.example.local:9000", "", "minioadmin", "minioadmin123", "instant-shared") require.NoError(t, err, "New must succeed when endpoint is provided (no dial at construction)") require.NotNil(t, p) } @@ -27,7 +27,15 @@ func TestNew_ValidEndpointSucceeds(t *testing.T) { // TestNew_DefaultBucketName verifies empty bucketName defaults to "instant-shared". func TestNew_DefaultBucketName(t *testing.T) { // Just verify construction succeeds — bucket name default is internal. - p, err := storageprovider.New("minio.example.local:9000", "root", "pass", "") + p, err := storageprovider.New("minio.example.local:9000", "", "root", "pass", "") + require.NoError(t, err) + require.NotNil(t, p) +} + +// TestNew_PublicEndpointAccepted verifies that a public endpoint override is accepted +// without altering construction. Behavior is exercised end-to-end via Provision(). +func TestNew_PublicEndpointAccepted(t *testing.T) { + p, err := storageprovider.New("minio.example.local:9000", "s3.instanode.dev:9000", "root", "pass", "instant-shared") require.NoError(t, err) require.NotNil(t, p) } diff --git a/internal/provisioner/client.go b/internal/provisioner/client.go index f7e595c6..9fbc3c37 100644 --- a/internal/provisioner/client.go +++ b/internal/provisioner/client.go @@ -66,13 +66,16 @@ func (c *Client) ctxWithAuth(ctx context.Context) context.Context { } // provisionTimeout returns the gRPC timeout for a provisioning call. -// Pro and team tiers create a dedicated k8s pod per token; pod startup can take 1-3 minutes. -// All other tiers provision on shared infrastructure in < 1 second. +// Every tier now provisions a dedicated k8s pod (since the dedicated-infra-for- +// every-tier change). PVC bind + image pull + postgres init can take 30-90s on +// a cold node, so 10s (the old anonymous default) drops the connection while +// the pod is still coming up. Anonymous gets a tight 4m budget; pro/team get +// 5m for larger images and bigger PVCs. func provisionTimeout(tier string) time.Duration { if tier == "pro" || tier == "team" || tier == "growth" { return 5 * time.Minute } - return 10 * time.Second + return 4 * time.Minute } // ProvisionPostgres provisions a new Postgres database. diff --git a/internal/router/router.go b/internal/router/router.go index 077c8553..28130087 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -17,6 +17,7 @@ import ( "instant.dev/internal/middleware" "instant.dev/internal/migratorclient" "instant.dev/internal/plans" + "instant.dev/internal/providers/compute/k8s" storageprovider "instant.dev/internal/providers/storage" "instant.dev/internal/provisioner" ) @@ -58,9 +59,13 @@ func New(cfg *config.Config, db *sql.DB, rdb *redis.Client, geoDbs *middleware.G EnableStackTrace: cfg.Environment == "development", })) app.Use(fiberCORS.New(fiberCORS.Config{ - AllowOrigins: "*", - AllowMethods: "GET,POST,PATCH,DELETE,OPTIONS", - AllowHeaders: "Content-Type,Authorization,X-Request-ID", + // Production origin (GitHub Pages serves instanode.dev) + every + // reasonable local-dev port. The wildcard would still work for + // bearer-token traffic (no cookies in flight) but an explicit + // allowlist makes the policy auditable. Add origins as needed. + AllowOrigins: "https://instanode.dev,https://www.instanode.dev,http://localhost:5173,http://localhost:3000,http://localhost:5174", + AllowMethods: "GET,POST,PUT,PATCH,DELETE,OPTIONS", + AllowHeaders: "Content-Type,Authorization,X-Request-ID,X-E2E-Test-Token,X-E2E-Source-IP", ExposeHeaders: "X-Request-ID,X-Instant-Upgrade,X-Instant-Notice", })) app.Use(middleware.GeoEnrich(geoDbs)) @@ -78,7 +83,7 @@ func New(cfg *config.Config, db *sql.DB, rdb *redis.Client, geoDbs *middleware.G // so that DELETE /api/v1/resources/:id can deprovision MinIO IAM users. var storageProv *storageprovider.Provider if cfg.MinioEndpoint != "" { - if sp, err := storageprovider.New(cfg.MinioEndpoint, cfg.MinioRootUser, cfg.MinioRootPassword, cfg.MinioBucketName); err != nil { + if sp, err := storageprovider.New(cfg.MinioEndpoint, cfg.MinioPublicEndpoint, cfg.MinioRootUser, cfg.MinioRootPassword, cfg.MinioBucketName); err != nil { slog.Warn("storage: MinIO provider init failed", "error", err) } else { storageProv = sp @@ -97,6 +102,23 @@ func New(cfg *config.Config, db *sql.DB, rdb *redis.Client, geoDbs *middleware.G deployH := handlers.NewDeployHandler(db, rdb, cfg) stackH := handlers.NewStackHandler(db, rdb, cfg, planRegistry) + // Custom-domain handler shares the k8s stack provider so EnsureCustomDomainIngress + // can update the same Ingress namespace the stack lives in. We construct a + // dedicated *k8s.K8sStackProvider here (rather than reaching into stackH) so + // the dependency surface stays explicit. When ComputeProvider != "k8s" the + // pointer is left nil and the handler skips ingress work — verification still + // progresses through TXT and the row stays at "verified" / "ingress_ready" + // until a future operator wires real k8s. + var customDomainK8s handlers.CustomDomainProvider + if cfg.ComputeProvider == "k8s" { + if csp, err := k8s.NewStackProvider(cfg.KubeNamespaceApps); err != nil { + slog.Warn("custom_domain.k8s_provider_unavailable", "error", err) + } else { + customDomainK8s = csp + } + } + customDomainH := handlers.NewCustomDomainHandler(db, cfg, planRegistry, customDomainK8s) + // ── Routes ─────────────────────────────────────────────────────────────── // Health check @@ -107,6 +129,9 @@ func New(cfg *config.Config, db *sql.DB, rdb *redis.Client, geoDbs *middleware.G // OpenAPI spec — machine-readable description of the agent-facing API app.Get("/openapi.json", handlers.ServeOpenAPI) + // MCP authorization profile — RFC 8414 / OAuth 2.0 Protected Resource Metadata. + app.Get("/.well-known/oauth-protected-resource", handlers.ServeOAuthProtectedResourceMetadata) + // Prometheus metrics — gated by METRICS_TOKEN when set (open in local dev). app.Get("/metrics", func(c *fiber.Ctx) error { if cfg.MetricsToken != "" { @@ -155,11 +180,23 @@ func New(cfg *config.Config, db *sql.DB, rdb *redis.Client, geoDbs *middleware.G app.Patch("/stacks/:slug/env", middleware.RequireAuth(cfg), stackH.UpdateEnv) app.Post("/stacks/:slug/redeploy", middleware.RequireAuth(cfg), stackH.Redeploy) - // OAuth + // OAuth — POST handler serves the existing programmatic / SPA flow. + // Google login is intentionally NOT supported; if you need it, register + // the routes here and wire GOOGLE_CLIENT_ID + GOOGLE_CLIENT_SECRET. app.Post("/auth/github", authH.GitHub) - app.Post("/auth/google", authH.Google) - app.Post("/auth/google/callback", authH.GoogleCallback) - app.Get("/auth/google/url", authH.GoogleAuthURL) + + // Browser OAuth flows (GET-based, redirect-driven). The dashboard's + // login page links to /auth/github/start directly; it stashes a CSRF + // state cookie, hands off to GitHub, and 302s back to + // <return_to>?session_token=<jwt> after exchanging the code. + app.Get("/auth/github/start", authH.GitHubStart) + app.Get("/auth/github/callback", authH.GitHubCallback) + + // Magic-link email login. Start is POST (the dashboard's login form + // submits to it); Callback is GET (the user's email client links to it). + mlH := handlers.NewMagicLinkHandler(db, cfg, emailClient, authH) + app.Post("/auth/email/start", mlH.Start) + app.Get("/auth/email/callback", mlH.Callback) // CLI device-flow login — POST creates session, GET polls for completion app.Post("/auth/cli", cliAuthH.CreateCLISession) @@ -172,17 +209,28 @@ func New(cfg *config.Config, db *sql.DB, rdb *redis.Client, geoDbs *middleware.G migClient = migratorclient.New(cfg.MigratorAddr, cfg.MigratorSecret) } billing := handlers.NewBillingHandler(db, cfg, emailClient, migClient) - app.Post("/billing/checkout", middleware.RequireAuth(cfg), billing.CreateCheckout) + // Legacy alias kept for backward compatibility; canonical path is + // /api/v1/billing/checkout (registered under the /api/v1 group below). + app.Post("/billing/checkout", middleware.RequireAuth(cfg), billing.CreateCheckoutAPI) app.Post("/razorpay/webhook", billing.RazorpayWebhook) // Public webhook request listing — token IS the credential (no session needed). // Authenticated callers use the same handler; it additionally verifies team ownership. app.Get("/api/v1/webhooks/:token/requests", middleware.OptionalAuth(cfg), webhookH.ListRequests) + // Public token-based invitation accept — must be registered BEFORE the + // /api/v1 auth group so the group middleware doesn't catch it. + // (Token IS the auth here — no Bearer required.) + teamsHPublic := handlers.NewTeamsHandler(db, cfg, emailClient) + app.Post("/api/v1/invitations/:token/accept", teamsHPublic.AcceptInvitation) + // Authenticated resource management - api := app.Group("/api/v1", middleware.RequireAuth(cfg)) + middleware.SetRoleLookupDB(db) // populate auth_team_role on every RequireAuth + middleware.SetAPIKeyDB(db) // enable PAT auth path in RequireAuth + api := app.Group("/api/v1", middleware.RequireAuth(cfg), middleware.PopulateTeamRole()) api.Get("/resources", resourceH.List) api.Get("/resources/:id", resourceH.Get) + api.Get("/resources/:id/credentials", resourceH.GetCredentials) api.Delete("/resources/:id", resourceH.Delete) api.Post("/resources/:id/rotate-credentials", resourceH.RotateCredentials) @@ -194,6 +242,7 @@ func New(cfg *config.Config, db *sql.DB, rdb *redis.Client, geoDbs *middleware.G api.Delete("/team/invitations/:id", teamMembersH.RevokeInvitation) api.Post("/team/invitations/:id/accept", teamMembersH.AcceptInvitation) + api.Post("/billing/checkout", billing.CreateCheckoutAPI) api.Post("/billing/cancel", billing.CancelSubscriptionAPI) api.Get("/billing/invoices", billing.ListInvoicesAPI) api.Post("/billing/update-payment", billing.UpdatePaymentMethodAPI) @@ -207,6 +256,39 @@ func New(cfg *config.Config, db *sql.DB, rdb *redis.Client, geoDbs *middleware.G // Stack management endpoints — Phase 6 (under /api/v1) api.Get("/stacks", stackH.List) + // Custom domains — Pro+ "bring your own hostname" for stacks. All routes + // require auth (the /api/v1 group middleware) and additionally enforce + // stack ownership inside the handler. + api.Post("/stacks/:slug/domains", customDomainH.Create) + api.Get("/stacks/:slug/domains", customDomainH.List) + api.Post("/stacks/:slug/domains/:id/verify", customDomainH.Verify) + api.Delete("/stacks/:slug/domains/:id", customDomainH.Delete) + + // Personal Access Tokens — long-lived bearer tokens for agents/CI. + apiKeysH := handlers.NewAPIKeysHandler(db) + api.Post("/auth/api-keys", apiKeysH.Create) + api.Get("/auth/api-keys", apiKeysH.List) + api.Delete("/auth/api-keys/:id", apiKeysH.Revoke) + + // Per-team audit log — feeds the dashboard's Recent Activity panel. + auditH := handlers.NewAuditHandler(db) + api.Get("/audit", auditH.List) + + // Vault — per-team encrypted secret storage (Phase 1: Heroku-shape platform). + vaultH := handlers.NewVaultHandler(db, cfg, planRegistry) + api.Put("/vault/:env/:key", vaultH.PutSecret) + api.Get("/vault/:env/:key", vaultH.GetSecret) + api.Get("/vault/:env", vaultH.ListKeys) + api.Delete("/vault/:env/:key", vaultH.DeleteSecret) + api.Post("/vault/:env/:key/rotate", vaultH.RotateSecret) + + // Teams + RBAC invitation flow (Phase 3). Public accept route is + // registered above the api group so the auth middleware doesn't catch it. + teamsH := teamsHPublic // reuse the same handler instance + api.Post("/teams/:team_id/invitations", middleware.RequireRole("admin"), teamsH.CreateInvitation) + api.Get("/teams/:team_id/invitations", middleware.RequireRole("admin"), teamsH.ListInvitations) + api.Delete("/teams/:team_id/invitations/:id", middleware.RequireRole("admin"), teamsH.RevokeInvitation) + // Internal dev-only endpoints — only registered in development environment. // These bypass Razorpay and directly mutate DB state. Never expose in production. if cfg.Environment == "development" { diff --git a/main.go b/main.go index e9376854..218fa4bb 100644 --- a/main.go +++ b/main.go @@ -88,7 +88,7 @@ func main() { var storageProv *storageprovider.Provider if cfg.MinioEndpoint != "" { - if sp, err := storageprovider.New(cfg.MinioEndpoint, cfg.MinioRootUser, cfg.MinioRootPassword, cfg.MinioBucketName); err != nil { + if sp, err := storageprovider.New(cfg.MinioEndpoint, cfg.MinioPublicEndpoint, cfg.MinioRootUser, cfg.MinioRootPassword, cfg.MinioBucketName); err != nil { slog.Warn("dashboard_grpc: MinIO provider init failed", "error", err) } else { storageProv = sp diff --git a/plans.yaml b/plans.yaml index 1fd91c40..38437903 100644 --- a/plans.yaml +++ b/plans.yaml @@ -24,6 +24,9 @@ plans: storage_storage_mb: 10 webhook_requests_stored: 100 team_members: 1 + vault_max_entries: 0 + vault_envs_allowed: [] + deployments_apps: 0 features: alerts: false custom_domains: false @@ -35,9 +38,9 @@ plans: trial_days: 14 limits: provisions_per_day: -1 - postgres_storage_mb: 500 - postgres_connections: 5 - redis_memory_mb: 25 + postgres_storage_mb: 1024 + postgres_connections: 8 + redis_memory_mb: 50 redis_commands_per_day: 10000 mongodb_storage_mb: 100 mongodb_connections: 5 @@ -46,6 +49,9 @@ plans: storage_storage_mb: 512 webhook_requests_stored: 1000 team_members: 1 + vault_max_entries: 20 + vault_envs_allowed: ["production"] + deployments_apps: 1 features: alerts: true custom_domains: false @@ -68,9 +74,12 @@ plans: storage_storage_mb: 10240 webhook_requests_stored: 10000 team_members: 5 + vault_max_entries: 200 + vault_envs_allowed: [] + deployments_apps: 10 features: alerts: true - custom_domains: false + custom_domains: true sla: false team: @@ -90,6 +99,9 @@ plans: storage_storage_mb: -1 webhook_requests_stored: -1 team_members: -1 + vault_max_entries: -1 + vault_envs_allowed: [] + deployments_apps: -1 features: alerts: true custom_domains: true @@ -101,9 +113,9 @@ plans: trial_days: 0 limits: provisions_per_day: -1 - postgres_storage_mb: -1 - postgres_connections: -1 - redis_memory_mb: -1 + postgres_storage_mb: 5120 + postgres_connections: 20 + redis_memory_mb: 256 redis_commands_per_day: -1 mongodb_storage_mb: -1 mongodb_connections: -1 @@ -112,6 +124,9 @@ plans: storage_storage_mb: -1 webhook_requests_stored: -1 team_members: 10 + vault_max_entries: 200 + vault_envs_allowed: [] + deployments_apps: 5 features: alerts: true custom_domains: true