From 88817b26c17256d567d9206012ec252f7f3ffd55 Mon Sep 17 00:00:00 2001 From: Nikolay Govorov Date: Mon, 6 Apr 2026 08:16:47 +0100 Subject: Auth & database refactoring --- cmd/mirumd/actor.go | 252 +++++ cmd/mirumd/api_auth.go | 236 ----- cmd/mirumd/api_cli.go | 11 +- cmd/mirumd/database.go | 1499 +++++++++++++++++++++++++++++ cmd/mirumd/id.go | 302 ++++++ cmd/mirumd/id_test.go | 76 ++ cmd/mirumd/main.go | 5 +- cmd/mirumd/server.go | 5 +- cmd/mirumd/server_admin.go | 241 +++-- cmd/mirumd/server_grpc.go | 3 +- cmd/mirumd/server_web.go | 21 +- internal/database/actor.go | 117 --- internal/database/database.go | 306 ------ internal/database/organization.go | 477 --------- internal/database/user.go | 492 ---------- internal/database/validate.go | 60 -- internal/database/worker.go | 165 ---- 17 files changed, 2317 insertions(+), 1951 deletions(-) create mode 100644 cmd/mirumd/actor.go delete mode 100644 cmd/mirumd/api_auth.go create mode 100644 cmd/mirumd/database.go create mode 100644 cmd/mirumd/id.go create mode 100644 cmd/mirumd/id_test.go delete mode 100644 internal/database/actor.go delete mode 100644 internal/database/database.go delete mode 100644 internal/database/organization.go delete mode 100644 internal/database/user.go delete mode 100644 internal/database/validate.go delete mode 100644 internal/database/worker.go diff --git a/cmd/mirumd/actor.go b/cmd/mirumd/actor.go new file mode 100644 index 0000000..bdd49c2 --- /dev/null +++ b/cmd/mirumd/actor.go @@ -0,0 +1,252 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package main + +import ( + "context" + "errors" + "slices" + + "dimidiumlabs/mirum/internal/protocol/pb" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5" +) + +var ( + ErrPermissionDenied = errors.New("database: permission denied") + ErrUnauthenticated = errors.New("database: authentication required") +) + +var anonPermissions = []pb.Perm{ + pb.Perm_PERM_ORG_READ, +} + +var userGlobalPermissions = []pb.Perm{ + pb.Perm_PERM_ORG_READ, + pb.Perm_PERM_ORG_WRITE, + pb.Perm_PERM_USER_READ, +} + +// rolePermissions is the single source of truth for role → perm bundles. +// RLS checks only tenancy (membership); action authz lives here. +var rolePermissions = map[string][]pb.Perm{ + "owner": { + pb.Perm_PERM_ORG_READ, + pb.Perm_PERM_ORG_WRITE, + pb.Perm_PERM_ORG_DELETE, + pb.Perm_PERM_ORG_MEMBER_READ, + pb.Perm_PERM_ORG_MEMBER_WRITE, + pb.Perm_PERM_WORKER_READ, + pb.Perm_PERM_WORKER_WRITE, + }, + "admin": { + pb.Perm_PERM_ORG_READ, + pb.Perm_PERM_ORG_WRITE, + pb.Perm_PERM_ORG_MEMBER_READ, + pb.Perm_PERM_ORG_MEMBER_WRITE, + pb.Perm_PERM_WORKER_READ, + pb.Perm_PERM_WORKER_WRITE, + }, + "member": { + pb.Perm_PERM_ORG_READ, + pb.Perm_PERM_ORG_MEMBER_READ, + pb.Perm_PERM_WORKER_READ, + }, +} + +// Actor is the principal making a database request. It carries identity, +// display metadata, and coarse capability. Zero value is invalid: dbID +// panics, so a missing initialisation cannot silently grant privileges. +// +// Synthetic actors (System/Operator/Anon) live only as Go constants — +// they are not rows in the users table, so they cannot be logged in as +// even if somebody writes a password into the DB. +// +// Authorization is divided into two planes: +// - Tenancy: an actor can only see a subset of resources to which +// they have access (public or through organization membership). +// Any select statement will return only records accessible to the actor. +// - RBAC: what the actor can do with records (create/read/write) is implemented here. +// Any rights we grant here are a strict subset of the Tenancy rights. +// The list of perms can be either explicit (for tokens) or implied (for user roles). +type Actor struct { + kind actorKind + id uuid.UUID + email string + superuser bool +} + +type actorKind uint8 + +const ( + actorInvalid actorKind = iota + actorUser + actorOperator + actorSystem + actorAnon +) + +// ActorKind is the exported form of actorKind for audit sinks and logging. +type ActorKind uint8 + +const ( + KindInvalid ActorKind = iota + KindUser + KindOperator + KindSystem + KindAnon +) + +var ( + anonUUID = uuid.MustParse("ffffffff-ffff-ffff-ffff-ffffffffffff") + systemUUID = uuid.MustParse("00000000-0000-0000-0000-000000000001") + operatorUUID = uuid.MustParse("00000000-0000-0000-0000-000000000002") +) + +// UserActor identifies an authenticated user from a session or token. +func UserActor(id UserID, email string, superuser bool) Actor { + if id.IsZero() { + panic("database: UserActor with nil UUID") + } + if email == "" { + panic("database: UserActor with empty email") + } + return Actor{kind: actorUser, id: id.UUID(), email: email, superuser: superuser} +} + +// OperatorActor is the principal for externally invoked privileged +// operations (admin socket). Distinguishable from System in audit logs. +func OperatorActor() Actor { + return Actor{kind: actorOperator, id: operatorUUID, email: "operator@mirum.local", superuser: true} +} + +// SystemActor is the principal for internal machinery (mTLS handshake, +// session bootstrap, background jobs). Not an operator action. +func SystemActor() Actor { + return Actor{kind: actorSystem, id: systemUUID, email: "system@mirum.local", superuser: true} +} + +// AnonActor is the principal for unauthenticated public requests. +func AnonActor() Actor { + return Actor{kind: actorAnon, id: anonUUID, email: "anonymous@mirum.local"} +} + +func (a Actor) Kind() ActorKind { + switch a.kind { + case actorUser: + return KindUser + case actorOperator: + return KindOperator + case actorSystem: + return KindSystem + case actorAnon: + return KindAnon + } + return KindInvalid +} + +func (a Actor) UserID() UserID { return UserID(a.id) } +func (a Actor) Email() string { return a.email } +func (a Actor) IsSuperuser() bool { return a.superuser } + +// dbID returns the UUID to write into app.user_id. Panics on zero value. +func (a Actor) dbID() uuid.UUID { + if a.id == uuid.Nil { + panic("database: zero-value Actor; use UserActor/SystemActor/OperatorActor/AnonActor") + } + return a.id +} + +// kindString returns the string written into app.actor_kind. +// It must match the values tested by app_issuper() in the SQL migration. +func (a Actor) kindString() string { + switch a.kind { + case actorUser: + return "user" + case actorOperator: + return "operator" + case actorSystem: + return "system" + case actorAnon: + return "anon" + } + panic("database: zero-value Actor; use UserActor/SystemActor/OperatorActor/AnonActor") +} + +// checkGlobal checks a global-scope perm (no specific org). Pure, no DB. +func checkGlobal(actor Actor, perm pb.Perm) error { + switch actor.kind { + case actorOperator, actorSystem: + return nil + case actorAnon: + if slices.Contains(anonPermissions, perm) { + return nil + } + return ErrUnauthenticated + case actorUser: + if actor.superuser { + return nil + } + if slices.Contains(userGlobalPermissions, perm) { + return nil + } + return ErrPermissionDenied + default: + return ErrPermissionDenied + } +} + +// checkPerm checks an org-scoped perm within an existing transaction. +func checkPerm(ctx context.Context, tx pgx.Tx, actor Actor, orgID OrgID, perm pb.Perm) error { + switch actor.kind { + case actorOperator, actorSystem: + return nil + case actorAnon: + return ErrUnauthenticated + case actorUser: + if actor.superuser { + return nil + } + var role string + err := tx.QueryRow(ctx, + `SELECT role FROM org_members WHERE org_id = $1 AND user_id = $2`, + orgID, actor.id, + ).Scan(&role) + if err != nil { + return ErrPermissionDenied + } + if !slices.Contains(rolePermissions[role], perm) { + return ErrPermissionDenied + } + return nil + default: + return ErrPermissionDenied + } +} + +// checkSystem checks that actor is the internal system principal. +func checkSystem(actor Actor) error { + if actor.kind == actorSystem { + return nil + } + return ErrPermissionDenied +} + +// checkSelf checks that actor is the target user or superuser. +func checkSelf(actor Actor, targetID UserID) error { + switch actor.kind { + case actorOperator, actorSystem: + return nil + case actorAnon: + return ErrUnauthenticated + case actorUser: + if actor.superuser || targetID == actor.UserID() { + return nil + } + return ErrPermissionDenied + default: + return ErrPermissionDenied + } +} diff --git a/cmd/mirumd/api_auth.go b/cmd/mirumd/api_auth.go deleted file mode 100644 index 38abde4..0000000 --- a/cmd/mirumd/api_auth.go +++ /dev/null @@ -1,236 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package main - -import ( - "context" - "slices" - - "connectrpc.com/connect" - "github.com/google/uuid" - - "dimidiumlabs/mirum/internal/database" - "dimidiumlabs/mirum/internal/protocol/pb" - "dimidiumlabs/mirum/internal/protocol/pb/pbconnect" -) - -// Role → permissions mapping. -var rolePermissions = map[pb.Role][]pb.Perm{ - pb.Role_ROLE_OWNER: { - pb.Perm_PERM_ORG_READ, pb.Perm_PERM_ORG_WRITE, pb.Perm_PERM_ORG_DELETE, - pb.Perm_PERM_ORG_MEMBER_READ, pb.Perm_PERM_ORG_MEMBER_WRITE, - pb.Perm_PERM_WORKER_READ, pb.Perm_PERM_WORKER_WRITE, - }, - pb.Role_ROLE_ADMIN: { - pb.Perm_PERM_ORG_READ, pb.Perm_PERM_ORG_WRITE, - pb.Perm_PERM_ORG_MEMBER_READ, pb.Perm_PERM_ORG_MEMBER_WRITE, - pb.Perm_PERM_WORKER_READ, pb.Perm_PERM_WORKER_WRITE, - }, - pb.Role_ROLE_MEMBER: { - pb.Perm_PERM_ORG_READ, - pb.Perm_PERM_ORG_MEMBER_READ, - pb.Perm_PERM_WORKER_READ, - }, -} - -// ApiAuthInterceptor enforces authorization on admin RPCs. -// Authentication is handled upstream: web sessionMiddleware for TCP, -// ConnContext for unix socket. -type ApiAuthInterceptor struct { - srv *server -} - -func (a *ApiAuthInterceptor) authorize(ctx context.Context, procedure string, req connect.AnyRequest) error { - // Public routes — no auth required, visibility enforced in handler/DB. - switch procedure { - case pbconnect.AdminOrgListProcedure, - pbconnect.AdminOrgGetProcedure: - return nil - } - - actor := ActorFromContext(ctx) - if actor.Kind() == database.KindAnon { - return errUnauthenticated - } - if actor.IsSuperuser() { - return nil - } - - switch procedure { - // User — self or superuser - case pbconnect.AdminUserGetProcedure, - pbconnect.AdminUserUpdateProcedure, - pbconnect.AdminUserDeleteProcedure: - if req == nil { - return errDenied - } - - if isSelf(actor, req) { - return nil - } - return errDenied - - // User — superuser only - case pbconnect.AdminUserCreateProcedure, - pbconnect.AdminUserListProcedure: - return errDenied - - // Org — any authenticated - case pbconnect.AdminOrgCreateProcedure: - return nil - case pbconnect.AdminOrgUpdateProcedure: - return a.checkOrgPerm(ctx, actor, req, pb.Perm_PERM_ORG_WRITE) - case pbconnect.AdminOrgDeleteProcedure: - return a.checkOrgPerm(ctx, actor, req, pb.Perm_PERM_ORG_DELETE) - - // OrgMember — org-scoped - case pbconnect.AdminOrgMemberGetProcedure, - pbconnect.AdminOrgMemberListProcedure: - return a.checkOrgPerm(ctx, actor, req, pb.Perm_PERM_ORG_MEMBER_READ) - case pbconnect.AdminOrgMemberAddProcedure, - pbconnect.AdminOrgMemberUpdateProcedure, - pbconnect.AdminOrgMemberRemoveProcedure: - return a.checkOrgPerm(ctx, actor, req, pb.Perm_PERM_ORG_MEMBER_WRITE) - - // Worker — any authenticated can read - case pbconnect.AdminWorkerGetProcedure, - pbconnect.AdminWorkerListProcedure: - return nil - - // Worker create — with org: org perm, without: superuser - case pbconnect.AdminWorkerCreateProcedure: - if orgRefFromRequest(req) == nil { - return errDenied - } - return a.checkOrgPerm(ctx, actor, req, pb.Perm_PERM_WORKER_WRITE) - - // Worker delete — lookup worker's org, then check - case pbconnect.AdminWorkerDeleteProcedure: - return a.checkWorkerDelete(ctx, actor, req) - - default: - return errDenied - } -} - -var ( - errDenied = newAPIError(connect.CodePermissionDenied, pb.ErrorReason_ERROR_REASON_PERMISSION_DENIED, nil) - errUnauthenticated = newAPIError(connect.CodeUnauthenticated, pb.ErrorReason_ERROR_REASON_UNAUTHENTICATED, nil) -) - -// checkOrgPerm extracts OrgRef from request and checks the actor's role permission. -func (a *ApiAuthInterceptor) checkOrgPerm(ctx context.Context, actor database.Actor, req connect.AnyRequest, perm pb.Perm) error { - ref := orgRefFromRequest(req) - if ref == nil { - return errDenied - } - member, err := a.srv.db.GetOrgMember(ctx, actor, orgRef(ref), database.UserByID(actor.UserID())) - if err != nil { - return errDenied - } - if !slices.Contains(rolePermissions[roleToProto[member.Role]], perm) { - return errDenied - } - return nil -} - -// checkWorkerDelete looks up the worker's org and checks permission. -func (a *ApiAuthInterceptor) checkWorkerDelete(ctx context.Context, actor database.Actor, req connect.AnyRequest) error { - m, ok := req.Any().(*pb.WorkerDeleteRequest) - if !ok { - return errDenied - } - w, err := a.srv.db.GetWorker(ctx, actor, uuid.UUID(m.Id)) - if err != nil { - return errDenied - } - if w.OrgID == nil { - return errDenied // global worker — superuser only - } - member, err := a.srv.db.GetOrgMember(ctx, actor, database.OrgByID(*w.OrgID), database.UserByID(actor.UserID())) - if err != nil { - return errDenied - } - if !slices.Contains(rolePermissions[roleToProto[member.Role]], pb.Perm_PERM_WORKER_WRITE) { - return errDenied - } - return nil -} - -// isSelf checks if the request targets the actor's own user. -func isSelf(actor database.Actor, req connect.AnyRequest) bool { - var ref *pb.UserRef - switch m := req.Any().(type) { - case *pb.UserGetRequest: - ref = m.User - case *pb.UserUpdateRequest: - ref = m.User - case *pb.UserDeleteRequest: - ref = m.User - } - if ref == nil { - return false - } - switch v := ref.GetRef().(type) { - case *pb.UserRef_Id: - return uuid.UUID(v.Id) == actor.UserID() - case *pb.UserRef_Email: - return v.Email == actor.Email() - default: - return false - } -} - -// orgRefFromRequest extracts the OrgRef from requests that carry one. -func orgRefFromRequest(req connect.AnyRequest) *pb.OrgRef { - if req == nil { - return nil - } - switch m := req.Any().(type) { - case *pb.OrgGetRequest: - return m.Org - case *pb.OrgUpdateRequest: - return m.Org - case *pb.OrgDeleteRequest: - return m.Org - case *pb.OrgMemberAddRequest: - return m.Org - case *pb.OrgMemberGetRequest: - return m.Org - case *pb.OrgMemberListRequest: - return m.Org - case *pb.OrgMemberUpdateRequest: - return m.Org - case *pb.OrgMemberRemoveRequest: - return m.Org - case *pb.WorkerCreateRequest: - return m.Org - default: - return nil - } -} - -func (a *ApiAuthInterceptor) WrapUnary(next connect.UnaryFunc) connect.UnaryFunc { - return func(ctx context.Context, req connect.AnyRequest) (connect.AnyResponse, error) { - if err := a.authorize(ctx, req.Spec().Procedure, req); err != nil { - return nil, err - } - return next(ctx, req) - } -} - -func (a *ApiAuthInterceptor) WrapStreamingClient(next connect.StreamingClientFunc) connect.StreamingClientFunc { - return func(ctx context.Context, spec connect.Spec) connect.StreamingClientConn { - panic("admin service does not make outbound streaming calls") - } -} - -func (a *ApiAuthInterceptor) WrapStreamingHandler(next connect.StreamingHandlerFunc) connect.StreamingHandlerFunc { - return func(ctx context.Context, conn connect.StreamingHandlerConn) error { - if err := a.authorize(ctx, conn.Spec().Procedure, nil); err != nil { - return err - } - return next(ctx, conn) - } -} diff --git a/cmd/mirumd/api_cli.go b/cmd/mirumd/api_cli.go index 6769798..6d53511 100644 --- a/cmd/mirumd/api_cli.go +++ b/cmd/mirumd/api_cli.go @@ -19,7 +19,6 @@ import ( "strings" "unicode" - "github.com/google/uuid" "github.com/spf13/cobra" "github.com/spf13/pflag" "google.golang.org/protobuf/proto" @@ -245,11 +244,11 @@ func registerMessageField(cmd *cobra.Command, fd protoreflect.FieldDescriptor, f func parseBytesFlag(fieldName, v string) ([]byte, error) { switch { case fieldName == "id" || strings.HasSuffix(fieldName, "_id"): - u, err := uuid.Parse(v) + id, err := ParseAnyID(v) if err != nil { - return nil, fmt.Errorf("invalid uuid: %w", err) + return nil, fmt.Errorf("invalid id: %w", err) } - return u[:], nil + return id[:], nil case strings.Contains(fieldName, "key"): der, err := base64.StdEncoding.DecodeString(v) if err != nil { @@ -419,9 +418,7 @@ func formatScalar(fd protoreflect.FieldDescriptor, v protoreflect.Value) string func formatBytes(b []byte) string { if len(b) == 16 { - if u, err := uuid.FromBytes(b); err == nil { - return u.String() - } + return FormatAnyID(b) } return base64.StdEncoding.EncodeToString(b) } diff --git a/cmd/mirumd/database.go b/cmd/mirumd/database.go new file mode 100644 index 0000000..6c03f7a --- /dev/null +++ b/cmd/mirumd/database.go @@ -0,0 +1,1499 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package main + +import ( + "context" + "crypto/hmac" + "crypto/rand" + "crypto/sha256" + "crypto/subtle" + "encoding/base64" + "errors" + "fmt" + "log/slog" + "net/mail" + "regexp" + "strings" + "time" + + "dimidiumlabs/mirum/internal/config" + "dimidiumlabs/mirum/internal/protocol/pb" + + sb "github.com/huandu/go-sqlbuilder" + "github.com/jackc/pgerrcode" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgconn" + "github.com/jackc/pgx/v5/pgxpool" + "github.com/jackc/tern/v2/migrate" + "golang.org/x/crypto/argon2" +) + +var ( + ErrAcquire = errors.New("database: failed to acquire connection") + ErrAlreadyMember = errors.New("database: already a member") + ErrEmailTaken = errors.New("database: email already taken") + ErrInvalidCreds = errors.New("database: invalid credentials") + ErrInvalidEmail = errors.New("database: invalid email") + ErrInvalidRole = errors.New("database: invalid role") + ErrInvalidSlug = errors.New("database: invalid slug") + ErrLastOwner = errors.New("database: last owner") + ErrMigrate = errors.New("database: failed to create migrator") + ErrNotImplemented = errors.New("database: filter not implemented") + ErrNotMember = errors.New("database: not a member") + ErrOpen = errors.New("database: failed to open") + ErrOrgNotFound = errors.New("database: organization not found") + ErrPing = errors.New("database: failed to ping") + ErrReservedEmail = errors.New("database: email uses a reserved domain") + ErrSlugTaken = errors.New("database: slug already taken") + ErrSoleOwner = errors.New("database: sole owner of an organization") + ErrUserNotFound = errors.New("database: user not found") + ErrWorkerNotFound = errors.New("database: worker not found") +) + +// reservedEmailSuffix is the domain carved out for synthetic actors +// (system/operator/anon). Real users cannot register with this suffix. +const reservedEmailSuffix = "@mirum.local" + +const ( + saltLen = 16 + + argonTime = 3 + argonMemory = 64 * 1024 // 64 MB + argonKeyLen = 32 + argonThreads = 2 +) + +var slugRe = regexp.MustCompile(`^[a-zA-Z0-9]+(?:-[a-zA-Z0-9]+)*$`) + +// UserRef identifies a user by ID or email. +type UserRef struct { + id UserID + email string +} + +func UserByID(id UserID) UserRef { return UserRef{id: id} } + +func UserByEmail(email string) UserRef { return UserRef{email: email} } + +func (r UserRef) where() (string, any) { + if !r.id.IsZero() { + return "id", r.id + } + return "email", r.email +} + +// OrgRef identifies an organization by ID or slug. +type OrgRef struct { + id OrgID + slug string +} + +func OrgByID(id OrgID) OrgRef { return OrgRef{id: id} } + +func OrgBySlug(slug string) OrgRef { return OrgRef{slug: slug} } + +func (r OrgRef) IsZero() bool { return r.id.IsZero() && r.slug == "" } + +func (r OrgRef) where() (string, any) { + if !r.id.IsZero() { + return "id", r.id + } + return "slug", r.slug +} + +// DB wraps a pgx connection pool. +type DB struct { + Pool *pgxpool.Pool +} + +// User holds info about a user. +type User struct { + ID UserID + Email string + CreatedAt time.Time +} + +// Organization holds info about an organization. +type Organization struct { + ID OrgID + Name string + Slug string + Public bool + CreatedAt time.Time +} + +// OrgMember pairs a user with their role in an organization. +type OrgMember struct { + User User + Role string + JoinedAt time.Time +} + +// Worker holds info about a registered worker. +type Worker struct { + ID WorkerID + OrgID *OrgID + PublicKey []byte + CreatedAt time.Time +} + +// DatabaseOpen connects to PostgreSQL and returns a DB. +func DatabaseOpen(ctx context.Context, dsn string) (*DB, error) { + pool, err := pgxpool.New(ctx, dsn) + if err != nil { + return nil, errors.Join(ErrOpen, err) + } + + if err := pool.Ping(ctx); err != nil { + pool.Close() + return nil, errors.Join(ErrPing, err) + } + + return &DB{Pool: pool}, nil +} + +// Close closes the connection pool. +func (db *DB) Close() { + db.Pool.Close() +} + +// apicall starts a transaction and sets the RLS actor. Both app.user_id +// and app.actor_kind are populated: app_issuper() checks actor_kind for +// System/Operator principals, and app.user_id for real user superusers. +func (db *DB) apicall(ctx context.Context, actor Actor, access, validate, doit func(pgx.Tx) error) error { + tx, err := db.Pool.Begin(ctx) + if err != nil { + return err + } + defer tx.Rollback(ctx) + + if _, err := tx.Exec(ctx, + `SELECT set_config('app.user_id', $1, true), + set_config('app.actor_kind', $2, true)`, + actor.dbID().String(), actor.kindString(), + ); err != nil { + return err + } + + if err := access(tx); err != nil { + slog.Debug("access denied", "err", err) + return err + } + if validate != nil { + if err := validate(tx); err != nil { + slog.Debug("validation failed", "err", err) + return err + } + } + if err := doit(tx); err != nil { + slog.Debug("exec failed", "err", err) + return err + } + + return tx.Commit(ctx) +} + +// Migrate applies all pending migrations. +func (db *DB) Migrate(ctx context.Context) error { + conn, err := db.Pool.Acquire(ctx) + if err != nil { + return errors.Join(ErrAcquire, err) + } + defer conn.Release() + + migrator, err := migrate.NewMigrator(ctx, conn.Conn(), "schema_version") + if err != nil { + return errors.Join(ErrMigrate, err) + } + + migrator.AppendMigration("create_users", ` + CREATE TABLE users ( + id UUID PRIMARY KEY DEFAULT uuidv7(), + email TEXT NOT NULL UNIQUE, + password TEXT NOT NULL, + superuser BOOLEAN NOT NULL DEFAULT false, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + deleted_at TIMESTAMPTZ + ); + + CREATE FUNCTION app_user_id() RETURNS uuid STABLE AS $$ + SELECT current_setting('app.user_id', true)::uuid; + $$ LANGUAGE sql; + + -- app_issuper has two independent branches: + -- (a) runtime setting app.actor_kind is 'system' or 'operator' — + -- set only by apicall from Go for synthetic principals and by + -- DML migrations; cannot be injected via login since there is + -- no matching users row to authenticate against. + -- (b) the current app.user_id resolves to a users row with + -- superuser = true — real support-agent style superusers. + CREATE FUNCTION app_issuper() RETURNS boolean STABLE AS $$ + SELECT + current_setting('app.actor_kind', true) IN ('system', 'operator') + OR EXISTS ( + SELECT 1 FROM users + WHERE id = current_setting('app.user_id', true)::uuid + AND superuser = true + ); + $$ LANGUAGE sql; + `, ` + DROP FUNCTION app_issuper; + DROP FUNCTION app_user_id; + DROP TABLE users; + `) + + migrator.AppendMigration("create_sessions", ` + CREATE TABLE sessions ( + token TEXT PRIMARY KEY, + user_id UUID NOT NULL REFERENCES users(id), + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + expires_at TIMESTAMPTZ NOT NULL + ); + CREATE INDEX sessions_expires_at ON sessions (expires_at); + `, ` + DROP TABLE sessions; + `) + + migrator.AppendMigration("create_organizations", ` + CREATE TABLE organizations ( + id UUID PRIMARY KEY DEFAULT uuidv7(), + name TEXT NOT NULL, + slug TEXT NOT NULL UNIQUE, + public BOOLEAN NOT NULL DEFAULT false, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + deleted_at TIMESTAMPTZ + ); + `, ` + DROP TABLE organizations; + `) + + migrator.AppendMigration("create_org_members", ` + CREATE TABLE org_members ( + org_id UUID NOT NULL REFERENCES organizations(id), + user_id UUID NOT NULL REFERENCES users(id), + role TEXT NOT NULL DEFAULT 'member' CHECK (role IN ('owner', 'admin', 'member')), + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + PRIMARY KEY (org_id, user_id) + ); + CREATE INDEX org_members_user_id ON org_members (user_id); + + CREATE FUNCTION is_member(org uuid) RETURNS boolean STABLE AS $$ + SELECT EXISTS ( + SELECT 1 FROM org_members + WHERE org_id = org AND user_id = app_user_id() + ); + $$ LANGUAGE sql; + + CREATE FUNCTION is_authenticated() RETURNS boolean STABLE AS $$ + SELECT current_setting('app.actor_kind', true) NOT IN ('', 'anon'); + $$ LANGUAGE sql; + `, ` + DROP FUNCTION is_authenticated; + DROP FUNCTION is_member; + DROP TABLE org_members; + `) + + migrator.AppendMigration("create_workers", ` + CREATE TABLE workers ( + id UUID PRIMARY KEY DEFAULT uuidv7(), + org_id UUID REFERENCES organizations(id), + public_key BYTEA NOT NULL UNIQUE, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + revoked_at TIMESTAMPTZ + ); + `, ` + DROP TABLE workers + `) + + migrator.AppendMigration("rls_users", ` + ALTER TABLE users ENABLE ROW LEVEL SECURITY; + ALTER TABLE users FORCE ROW LEVEL SECURITY; + + CREATE POLICY superuser ON users FOR ALL USING (app_issuper()); + CREATE POLICY self ON users FOR ALL USING (id = app_user_id()); + CREATE POLICY shared_org ON users FOR SELECT USING (EXISTS ( + SELECT 1 FROM org_members target + JOIN org_members mine ON mine.org_id = target.org_id + WHERE target.user_id = users.id + AND mine.user_id = app_user_id() + )); + CREATE POLICY write_auth ON users AS RESTRICTIVE FOR INSERT WITH CHECK (is_authenticated()); + CREATE POLICY update_auth ON users AS RESTRICTIVE FOR UPDATE USING (is_authenticated()); + CREATE POLICY delete_auth ON users AS RESTRICTIVE FOR DELETE USING (is_authenticated()); + `, ` + DROP POLICY write_auth ON users; + DROP POLICY update_auth ON users; + DROP POLICY delete_auth ON users; + DROP POLICY superuser ON users; + DROP POLICY self ON users; + DROP POLICY shared_org ON users; + + ALTER TABLE users DISABLE ROW LEVEL SECURITY; + `) + + migrator.AppendMigration("rls_sessions", ` + ALTER TABLE sessions ENABLE ROW LEVEL SECURITY; + ALTER TABLE sessions FORCE ROW LEVEL SECURITY; + + CREATE POLICY superuser ON sessions FOR ALL USING (app_issuper()); + CREATE POLICY own_sessions ON sessions FOR ALL USING (user_id = app_user_id()); + CREATE POLICY write_auth ON sessions AS RESTRICTIVE FOR INSERT WITH CHECK (is_authenticated()); + CREATE POLICY update_auth ON sessions AS RESTRICTIVE FOR UPDATE USING (is_authenticated()); + CREATE POLICY delete_auth ON sessions AS RESTRICTIVE FOR DELETE USING (is_authenticated()); + `, ` + DROP POLICY write_auth ON sessions; + DROP POLICY update_auth ON sessions; + DROP POLICY delete_auth ON sessions; + DROP POLICY superuser ON sessions; + DROP POLICY own_sessions ON sessions; + + ALTER TABLE sessions DISABLE ROW LEVEL SECURITY; + `) + + migrator.AppendMigration("rls_organizations", ` + ALTER TABLE organizations ENABLE ROW LEVEL SECURITY; + ALTER TABLE organizations FORCE ROW LEVEL SECURITY; + + CREATE POLICY superuser ON organizations FOR ALL USING (app_issuper()); + CREATE POLICY public_org ON organizations FOR ALL USING (public); + CREATE POLICY member_org ON organizations FOR ALL USING (is_member(id)); + CREATE POLICY write_auth ON organizations AS RESTRICTIVE FOR INSERT WITH CHECK (is_authenticated()); + CREATE POLICY update_auth ON organizations AS RESTRICTIVE FOR UPDATE USING (is_authenticated()); + CREATE POLICY delete_auth ON organizations AS RESTRICTIVE FOR DELETE USING (is_authenticated()); + `, ` + DROP POLICY write_auth ON organizations; + DROP POLICY update_auth ON organizations; + DROP POLICY delete_auth ON organizations; + DROP POLICY superuser ON organizations; + DROP POLICY public_org ON organizations; + DROP POLICY member_org ON organizations; + + ALTER TABLE organizations DISABLE ROW LEVEL SECURITY; + `) + + migrator.AppendMigration("rls_org_members", ` + ALTER TABLE org_members ENABLE ROW LEVEL SECURITY; + ALTER TABLE org_members FORCE ROW LEVEL SECURITY; + + CREATE POLICY superuser ON org_members FOR ALL USING (app_issuper()); + + -- Self path (own membership rows) uses a pure column predicate + -- so is_member's inner query can resolve without recursion. + CREATE POLICY self_member ON org_members FOR ALL USING (user_id = app_user_id()); + CREATE POLICY org_member ON org_members FOR ALL USING (is_member(org_id)); + CREATE POLICY write_auth ON org_members AS RESTRICTIVE FOR INSERT WITH CHECK (is_authenticated()); + CREATE POLICY update_auth ON org_members AS RESTRICTIVE FOR UPDATE USING (is_authenticated()); + CREATE POLICY delete_auth ON org_members AS RESTRICTIVE FOR DELETE USING (is_authenticated()); + `, ` + DROP POLICY write_auth ON org_members; + DROP POLICY update_auth ON org_members; + DROP POLICY delete_auth ON org_members; + DROP POLICY superuser ON org_members; + DROP POLICY self_member ON org_members; + DROP POLICY org_member ON org_members; + + ALTER TABLE org_members DISABLE ROW LEVEL SECURITY; + `) + + migrator.AppendMigration("rls_workers", ` + ALTER TABLE workers ENABLE ROW LEVEL SECURITY; + ALTER TABLE workers FORCE ROW LEVEL SECURITY; + + CREATE POLICY superuser ON workers FOR ALL USING (app_issuper()); + CREATE POLICY org_worker ON workers FOR ALL USING ( + org_id IS NOT NULL AND is_member(org_id) + ); + CREATE POLICY write_auth ON workers AS RESTRICTIVE FOR INSERT WITH CHECK (is_authenticated()); + CREATE POLICY update_auth ON workers AS RESTRICTIVE FOR UPDATE USING (is_authenticated()); + CREATE POLICY delete_auth ON workers AS RESTRICTIVE FOR DELETE USING (is_authenticated()); + `, ` + DROP POLICY write_auth ON workers; + DROP POLICY update_auth ON workers; + DROP POLICY delete_auth ON workers; + DROP POLICY superuser ON workers; + DROP POLICY org_worker ON workers; + + ALTER TABLE workers DISABLE ROW LEVEL SECURITY; + `) + + return migrator.Migrate(ctx) +} + +// UserCreate hashes the password with argon2id and inserts a new user. +// The pepper is a server-side secret not stored in the database. +func (db *DB) UserCreate(ctx context.Context, actor Actor, email, password string, pepper []byte) (UserID, error) { + var id UserID + err := db.apicall( + ctx, actor, + func(tx pgx.Tx) error { + if !actor.IsSuperuser() { + return ErrPermissionDenied + } + + return nil + }, + func(tx pgx.Tx) error { + if strings.HasSuffix(strings.ToLower(email), reservedEmailSuffix) { + return ErrReservedEmail + } + + return nil + }, + func(tx pgx.Tx) error { + hash, err := hashPassword(password, pepper) + if err != nil { + return err + } + + if err := tx.QueryRow(ctx, + `INSERT INTO users (email, password) VALUES ($1, $2) RETURNING id`, + email, hash, + ).Scan(&id); err != nil { + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) && pgErr.Code == pgerrcode.UniqueViolation { + return ErrEmailTaken + } + + return err + } + return nil + }, + ) + return id, err +} + +// UserGet returns a user by ref (ID or email). +func (db *DB) UserGet(ctx context.Context, actor Actor, ref UserRef) (*User, error) { + var u User + err := db.apicall( + ctx, actor, + func(tx pgx.Tx) error { return checkGlobal(actor, pb.Perm_PERM_USER_READ) }, + nil, + func(tx pgx.Tx) error { + col, val := ref.where() + + q := sb.PostgreSQL.NewSelectBuilder() + sql, args := q.Select("id", "email", "created_at"). + From("users"). + Where(q.Equal(col, val), q.IsNull("deleted_at")). + Build() + + if err := tx.QueryRow(ctx, sql, args...).Scan(&u.ID, &u.Email, &u.CreatedAt); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return ErrUserNotFound + } + return err + } + + return nil + }, + ) + return &u, err +} + +// UserList returns a page of users and the total count. +func (db *DB) UserList(ctx context.Context, actor Actor, cursor UserID, limit int, filter string) ([]User, int, error) { + var users []User + var total int + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { return checkGlobal(actor, pb.Perm_PERM_USER_READ) }, + func(tx pgx.Tx) error { + if filter != "" { + return ErrNotImplemented + } + return nil + }, + func(tx pgx.Tx) error { + if err := tx.QueryRow(ctx, + `SELECT count(*) FROM users WHERE deleted_at IS NULL`, + ).Scan(&total); err != nil { + return err + } + + q := sb.PostgreSQL.NewSelectBuilder() + q.Select("id", "email", "created_at"). + From("users"). + Where(q.IsNull("deleted_at")). + OrderBy("id"). + Limit(limit) + if !cursor.IsZero() { + q.Where(q.GreaterThan("id", cursor)) + } + + sql, args := q.Build() + rows, err := tx.Query(ctx, sql, args...) + if err != nil { + return err + } + defer rows.Close() + + for rows.Next() { + var u User + if err := rows.Scan(&u.ID, &u.Email, &u.CreatedAt); err != nil { + return err + } + users = append(users, u) + } + return rows.Err() + }, + ) + return users, total, err +} + +// UserUpdate updates a user's email and/or password. +// Invalidates all sessions when password changes. +func (db *DB) UserUpdate(ctx context.Context, actor Actor, ref UserRef, email *string, password *string, pepper []byte) error { + if email == nil && password == nil { + return nil + } + + var id UserID + return db.apicall(ctx, actor, + func(tx pgx.Tx) error { + var err error + id, err = resolveUser(ctx, tx, ref) + if err != nil { + return err + } + return checkSelf(actor, id) + }, + func(tx pgx.Tx) error { + if email != nil && strings.HasSuffix(strings.ToLower(*email), reservedEmailSuffix) { + return ErrReservedEmail + } + return nil + }, + func(tx pgx.Tx) error { + ub := sb.PostgreSQL.NewUpdateBuilder() + ub.Update("users") + + if email != nil { + ub.SetMore(ub.Assign("email", *email)) + } + if password != nil { + hash, err := hashPassword(*password, pepper) + if err != nil { + return err + } + ub.SetMore(ub.Assign("password", hash)) + } + ub.Where(ub.Equal("id", id)) + + sql, args := ub.Build() + if _, err := tx.Exec(ctx, sql, args...); err != nil { + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) && pgErr.Code == pgerrcode.UniqueViolation { + return ErrEmailTaken + } + return err + } + + if password != nil { + if _, err := tx.Exec(ctx, `DELETE FROM sessions WHERE user_id = $1`, id); err != nil { + return err + } + } + return nil + }, + ) +} + +// UserDelete soft-deletes a user. +// Fails if the user is the sole owner of any organization. +func (db *DB) UserDelete(ctx context.Context, actor Actor, ref UserRef) error { + var id UserID + return db.apicall(ctx, actor, + func(tx pgx.Tx) error { + var err error + id, err = resolveUser(ctx, tx, ref) + if err != nil { + return err + } + return checkSelf(actor, id) + }, + func(tx pgx.Tx) error { + return checkNotSoleOwner(ctx, tx, id) + }, + func(tx pgx.Tx) error { + if _, err := tx.Exec(ctx, `DELETE FROM sessions WHERE user_id = $1`, id); err != nil { + return err + } + if _, err := tx.Exec(ctx, `DELETE FROM org_members WHERE user_id = $1`, id); err != nil { + return err + } + if _, err := tx.Exec(ctx, + `UPDATE users SET email = id::text, password = '', deleted_at = now() WHERE id = $1`, id, + ); err != nil { + return err + } + return nil + }, + ) +} + +// UserVerifyPassword checks credentials and returns the user ID. +func (db *DB) UserVerifyPassword(ctx context.Context, actor Actor, email, password string, pepper []byte) (UserID, error) { + var id UserID + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { return checkSystem(actor) }, + nil, + func(tx pgx.Tx) error { + var hash string + if err := tx.QueryRow(ctx, + `SELECT id, password FROM users WHERE email = $1 AND deleted_at IS NULL`, + email, + ).Scan(&id, &hash); err != nil { + return ErrInvalidCreds + } + if !verifyHash(password, hash, pepper) { + return ErrInvalidCreds + } + return nil + }, + ) + return id, err +} + +// UserSessionGet resolves a session token into the Actor it authenticates. +// Returns an invalid zero Actor on any error; callers must check err. +func (db *DB) UserSessionGet(ctx context.Context, actor Actor, token string) (Actor, error) { + var ( + userID UserID + email string + superuser bool + ) + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { return checkSystem(actor) }, + nil, + func(tx pgx.Tx) error { + var expiresAt time.Time + h := hashToken(token) + + if err := tx.QueryRow(ctx, + `SELECT s.user_id, u.email, u.superuser, s.expires_at + FROM sessions s JOIN users u ON u.id = s.user_id + WHERE s.token = $1 AND s.expires_at > now() AND u.deleted_at IS NULL`, + h, + ).Scan(&userID, &email, &superuser, &expiresAt); err != nil { + return err + } + + if time.Until(expiresAt) < config.SessionTTL/2 { + if _, err := tx.Exec(ctx, + `UPDATE sessions SET expires_at = now() + $2 WHERE token = $1`, + h, config.SessionTTL, + ); err != nil { + return err + } + } + return nil + }, + ) + if err != nil { + return Actor{}, err + } + return UserActor(userID, email, superuser), nil +} + +// UserSessionCreate generates a random token, stores its hash, and returns the token. +func (db *DB) UserSessionCreate(ctx context.Context, actor Actor, userID UserID) (string, error) { + var token string + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { return checkSystem(actor) }, + nil, + func(tx pgx.Tx) error { + buf := make([]byte, 32) + if _, err := rand.Read(buf); err != nil { + return err + } + token = base64.RawURLEncoding.EncodeToString(buf) + + if _, err := tx.Exec(ctx, + `INSERT INTO sessions (token, user_id, expires_at) VALUES ($1, $2, now() + $3)`, + hashToken(token), userID, config.SessionTTL, + ); err != nil { + return err + } + return nil + }, + ) + return token, err +} + +// UserSessionDelete removes a session (logout). +func (db *DB) UserSessionDelete(ctx context.Context, actor Actor, token string) error { + return db.apicall(ctx, actor, + func(tx pgx.Tx) error { return checkSystem(actor) }, + nil, + func(tx pgx.Tx) error { + _, err := tx.Exec(ctx, `DELETE FROM sessions WHERE token = $1`, hashToken(token)) + return err + }, + ) +} + +// UserSessionPurgeExpired deletes all expired sessions. Runs as SystemActor. +func (db *DB) UserSessionPurgeExpired(ctx context.Context) error { + return db.apicall(ctx, SystemActor(), + func(tx pgx.Tx) error { return checkSystem(SystemActor()) }, + nil, + func(tx pgx.Tx) error { + _, err := tx.Exec(ctx, `DELETE FROM sessions WHERE expires_at < now()`) + return err + }, + ) +} + +// OrgGet returns an org by ref (ID or slug). +func (db *DB) OrgGet(ctx context.Context, actor Actor, ref OrgRef) (*Organization, error) { + var o Organization + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { return checkGlobal(actor, pb.Perm_PERM_ORG_READ) }, + nil, + func(tx pgx.Tx) error { + col, val := ref.where() + q := sb.PostgreSQL.NewSelectBuilder() + + sql, args := q.Select("id", "name", "slug", "public", "created_at"). + From("organizations"). + Where(q.Equal(col, val), q.IsNull("deleted_at")). + Build() + + if err := tx.QueryRow(ctx, sql, args...).Scan(&o.ID, &o.Name, &o.Slug, &o.Public, &o.CreatedAt); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return ErrOrgNotFound + } + return err + } + return nil + }, + ) + return &o, err +} + +// OrgCreate creates an org and adds the owner as the first member. +func (db *DB) OrgCreate(ctx context.Context, actor Actor, name, slug string, public bool, owner UserRef) (OrgID, error) { + var orgID OrgID + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { return checkGlobal(actor, pb.Perm_PERM_ORG_WRITE) }, + nil, + func(tx pgx.Tx) error { + userID, err := resolveUser(ctx, tx, owner) + if err != nil { + return err + } + + if err := tx.QueryRow(ctx, + `INSERT INTO organizations (name, slug, public) VALUES ($1, $2, $3) RETURNING id`, + name, slug, public, + ).Scan(&orgID); err != nil { + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) && pgErr.Code == pgerrcode.UniqueViolation { + return ErrSlugTaken + } + return err + } + + _, err = tx.Exec(ctx, + `INSERT INTO org_members (org_id, user_id, role) VALUES ($1, $2, 'owner')`, + orgID, userID, + ) + return err + }, + ) + return orgID, err +} + +// OrgUpdate updates an org's name, slug, and/or public flag. +func (db *DB) OrgUpdate(ctx context.Context, actor Actor, ref OrgRef, name *string, slug *string, public *bool) error { + if name == nil && slug == nil && public == nil { + return nil + } + + var id OrgID + return db.apicall(ctx, actor, + func(tx pgx.Tx) error { + var err error + id, err = resolveOrg(ctx, tx, ref) + if err != nil { + return err + } + return checkPerm(ctx, tx, actor, id, pb.Perm_PERM_ORG_WRITE) + }, + nil, + func(tx pgx.Tx) error { + ub := sb.PostgreSQL.NewUpdateBuilder() + ub.Update("organizations") + if name != nil { + ub.SetMore(ub.Assign("name", *name)) + } + if slug != nil { + ub.SetMore(ub.Assign("slug", *slug)) + } + if public != nil { + ub.SetMore(ub.Assign("public", *public)) + } + ub.Where(ub.Equal("id", id)) + + sql, args := ub.Build() + tag, err := tx.Exec(ctx, sql, args...) + if err != nil { + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) && pgErr.Code == pgerrcode.UniqueViolation { + return ErrSlugTaken + } + return err + } + if tag.RowsAffected() == 0 { + return ErrOrgNotFound + } + return nil + }, + ) +} + +// OrgDelete soft-deletes an org and removes all members. +func (db *DB) OrgDelete(ctx context.Context, actor Actor, ref OrgRef) error { + var id OrgID + return db.apicall(ctx, actor, + func(tx pgx.Tx) error { + var err error + id, err = resolveOrg(ctx, tx, ref) + if err != nil { + return err + } + return checkPerm(ctx, tx, actor, id, pb.Perm_PERM_ORG_DELETE) + }, + nil, + func(tx pgx.Tx) error { + if _, err := tx.Exec(ctx, `DELETE FROM org_members WHERE org_id = $1`, id); err != nil { + return err + } + tag, err := tx.Exec(ctx, + `UPDATE organizations SET slug = id::text, deleted_at = now() WHERE id = $1`, id, + ) + if err != nil { + return err + } + if tag.RowsAffected() == 0 { + return ErrOrgNotFound + } + return nil + }, + ) +} + +// OrgList returns a page of orgs and the total count. +func (db *DB) OrgList(ctx context.Context, actor Actor, cursor OrgID, limit int, filter string) ([]Organization, int, error) { + var orgs []Organization + var total int + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { return checkGlobal(actor, pb.Perm_PERM_ORG_READ) }, + func(tx pgx.Tx) error { + if filter != "" { + return ErrNotImplemented + } + return nil + }, + func(tx pgx.Tx) error { + if err := tx.QueryRow(ctx, + `SELECT count(*) FROM organizations WHERE deleted_at IS NULL`, + ).Scan(&total); err != nil { + return err + } + + q := sb.PostgreSQL.NewSelectBuilder() + q.Select("id", "name", "slug", "public", "created_at"). + From("organizations"). + Where(q.IsNull("deleted_at")). + OrderBy("id"). + Limit(limit) + if !cursor.IsZero() { + q.Where(q.GreaterThan("id", cursor)) + } + + sql, args := q.Build() + rows, err := tx.Query(ctx, sql, args...) + if err != nil { + return err + } + defer rows.Close() + + for rows.Next() { + var o Organization + if err := rows.Scan(&o.ID, &o.Name, &o.Slug, &o.Public, &o.CreatedAt); err != nil { + return err + } + orgs = append(orgs, o) + } + return rows.Err() + }, + ) + return orgs, total, err +} + +// OrgMemberGet returns a single member's info. +func (db *DB) OrgMemberGet(ctx context.Context, actor Actor, org OrgRef, user UserRef) (*OrgMember, error) { + var m OrgMember + var orgID OrgID + var userID UserID + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { + var err error + orgID, err = resolveOrg(ctx, tx, org) + if err != nil { + return err + } + return checkPerm(ctx, tx, actor, orgID, pb.Perm_PERM_ORG_MEMBER_READ) + }, + nil, + func(tx pgx.Tx) error { + var err error + userID, err = resolveUser(ctx, tx, user) + if err != nil { + return err + } + + if err := tx.QueryRow(ctx, + `SELECT u.id, u.email, u.created_at, om.role, om.created_at + FROM org_members om + JOIN users u ON u.id = om.user_id + WHERE om.org_id = $1 AND om.user_id = $2`, orgID, userID, + ).Scan(&m.User.ID, &m.User.Email, &m.User.CreatedAt, &m.Role, &m.JoinedAt); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return ErrNotMember + } + return err + } + return nil + }, + ) + return &m, err +} + +// OrgMembersList returns a page of members for an org. +func (db *DB) OrgMembersList(ctx context.Context, actor Actor, org OrgRef, cursor UserID, limit int, filter string) ([]OrgMember, int, error) { + var members []OrgMember + var total int + var orgID OrgID + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { + var err error + orgID, err = resolveOrg(ctx, tx, org) + if err != nil { + return err + } + return checkPerm(ctx, tx, actor, orgID, pb.Perm_PERM_ORG_MEMBER_READ) + }, + func(tx pgx.Tx) error { + if filter != "" { + return ErrNotImplemented + } + return nil + }, + func(tx pgx.Tx) error { + if err := tx.QueryRow(ctx, + `SELECT count(*) FROM org_members WHERE org_id = $1`, orgID, + ).Scan(&total); err != nil { + return err + } + + q := sb.PostgreSQL.NewSelectBuilder() + q.Select("u.id", "u.email", "u.created_at", "m.role", "m.created_at"). + From("org_members m"). + Join("users u", "u.id = m.user_id"). + Where(q.Equal("m.org_id", orgID), q.IsNull("u.deleted_at")). + OrderBy("u.id"). + Limit(limit) + if !cursor.IsZero() { + q.Where(q.GreaterThan("u.id", cursor)) + } + + sql, args := q.Build() + rows, err := tx.Query(ctx, sql, args...) + if err != nil { + return err + } + defer rows.Close() + + for rows.Next() { + var m OrgMember + if err := rows.Scan(&m.User.ID, &m.User.Email, &m.User.CreatedAt, &m.Role, &m.JoinedAt); err != nil { + return err + } + members = append(members, m) + } + return rows.Err() + }, + ) + return members, total, err +} + +// OrgMemberAdd adds a user to an org with the given role. +func (db *DB) OrgMemberAdd(ctx context.Context, actor Actor, org OrgRef, user UserRef, role string) error { + var orgID OrgID + return db.apicall(ctx, actor, + func(tx pgx.Tx) error { + var err error + orgID, err = resolveOrg(ctx, tx, org) + if err != nil { + return err + } + return checkPerm(ctx, tx, actor, orgID, pb.Perm_PERM_ORG_MEMBER_WRITE) + }, + nil, + func(tx pgx.Tx) error { + userID, err := resolveUser(ctx, tx, user) + if err != nil { + return err + } + if _, err := tx.Exec(ctx, + `INSERT INTO org_members (org_id, user_id, role) VALUES ($1, $2, $3)`, + orgID, userID, role, + ); err != nil { + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) && pgErr.Code == pgerrcode.UniqueViolation { + return ErrAlreadyMember + } + return err + } + return nil + }, + ) +} + +// OrgMemberUpdateRole changes a member's role. Fails if demoting the last owner. +func (db *DB) OrgMemberUpdateRole(ctx context.Context, actor Actor, org OrgRef, user UserRef, newRole string) error { + var orgID OrgID + var userID UserID + return db.apicall(ctx, actor, + func(tx pgx.Tx) error { + var err error + orgID, err = resolveOrg(ctx, tx, org) + if err != nil { + return err + } + return checkPerm(ctx, tx, actor, orgID, pb.Perm_PERM_ORG_MEMBER_WRITE) + }, + nil, + func(tx pgx.Tx) error { + var err error + userID, err = resolveUser(ctx, tx, user) + if err != nil { + return err + } + + var currentRole string + if err := tx.QueryRow(ctx, + `SELECT role FROM org_members WHERE org_id = $1 AND user_id = $2 FOR UPDATE`, + orgID, userID, + ).Scan(¤tRole); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return ErrNotMember + } + return err + } + + if currentRole == "owner" && newRole != "owner" { + var ownerCount int + if err := tx.QueryRow(ctx, + `SELECT count(*) FROM org_members WHERE org_id = $1 AND role = 'owner'`, + orgID, + ).Scan(&ownerCount); err != nil { + return err + } + if ownerCount <= 1 { + return ErrLastOwner + } + } + + _, err = tx.Exec(ctx, + `UPDATE org_members SET role = $1 WHERE org_id = $2 AND user_id = $3`, + newRole, orgID, userID, + ) + return err + }, + ) +} + +// OrgMemberRemove removes a user from an org. Fails if they are the last owner. +func (db *DB) OrgMemberRemove(ctx context.Context, actor Actor, org OrgRef, user UserRef) error { + var orgID OrgID + return db.apicall(ctx, actor, + func(tx pgx.Tx) error { + var err error + orgID, err = resolveOrg(ctx, tx, org) + if err != nil { + return err + } + return checkPerm(ctx, tx, actor, orgID, pb.Perm_PERM_ORG_MEMBER_WRITE) + }, + nil, + func(tx pgx.Tx) error { + userID, err := resolveUser(ctx, tx, user) + if err != nil { + return err + } + + var role string + if err := tx.QueryRow(ctx, + `SELECT role FROM org_members WHERE org_id = $1 AND user_id = $2 FOR UPDATE`, + orgID, userID, + ).Scan(&role); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return ErrNotMember + } + return err + } + + if role == "owner" { + var ownerCount int + if err := tx.QueryRow(ctx, + `SELECT count(*) FROM org_members WHERE org_id = $1 AND role = 'owner'`, + orgID, + ).Scan(&ownerCount); err != nil { + return err + } + if ownerCount <= 1 { + return ErrLastOwner + } + } + + _, err = tx.Exec(ctx, + `DELETE FROM org_members WHERE org_id = $1 AND user_id = $2`, + orgID, userID, + ) + return err + }, + ) +} + +// WorkerGet returns a worker by ID. +func (db *DB) WorkerGet(ctx context.Context, actor Actor, id WorkerID) (*Worker, error) { + var w Worker + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { + var orgID *OrgID + if err := tx.QueryRow(ctx, + `SELECT org_id FROM workers WHERE id = $1 AND revoked_at IS NULL`, id, + ).Scan(&orgID); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return ErrWorkerNotFound + } + return err + } + if orgID != nil { + return checkPerm(ctx, tx, actor, *orgID, pb.Perm_PERM_WORKER_READ) + } + if !actor.IsSuperuser() { + return ErrPermissionDenied + } + return nil + }, + nil, + func(tx pgx.Tx) error { + return tx.QueryRow(ctx, + `SELECT id, public_key, org_id, created_at FROM workers WHERE id = $1 AND revoked_at IS NULL`, id, + ).Scan(&w.ID, &w.PublicKey, &w.OrgID, &w.CreatedAt) + }, + ) + return &w, err +} + +// WorkerCreate registers a new worker with the given public key and optional org. +func (db *DB) WorkerCreate(ctx context.Context, actor Actor, publicKey []byte, org *OrgRef) (WorkerID, error) { + var workerID WorkerID + var orgID *OrgID + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { + if org != nil { + id, err := resolveOrg(ctx, tx, *org) + if err != nil { + return err + } + orgID = &id + return checkPerm(ctx, tx, actor, id, pb.Perm_PERM_WORKER_WRITE) + } + if !actor.IsSuperuser() { + return ErrPermissionDenied + } + return nil + }, + nil, + func(tx pgx.Tx) error { + return tx.QueryRow(ctx, + `INSERT INTO workers (public_key, org_id) VALUES ($1, $2) RETURNING id`, + publicKey, orgID, + ).Scan(&workerID) + }, + ) + return workerID, err +} + +// WorkerDelete soft-deletes a worker by ID. +func (db *DB) WorkerDelete(ctx context.Context, actor Actor, id WorkerID) error { + return db.apicall(ctx, actor, + func(tx pgx.Tx) error { + var orgID *OrgID + if err := tx.QueryRow(ctx, + `SELECT org_id FROM workers WHERE id = $1 AND revoked_at IS NULL`, id, + ).Scan(&orgID); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return ErrWorkerNotFound + } + return err + } + if orgID != nil { + return checkPerm(ctx, tx, actor, *orgID, pb.Perm_PERM_WORKER_WRITE) + } + if !actor.IsSuperuser() { + return ErrPermissionDenied + } + return nil + }, + nil, + func(tx pgx.Tx) error { + tag, err := tx.Exec(ctx, + `UPDATE workers SET revoked_at = now() WHERE id = $1 AND revoked_at IS NULL`, id, + ) + if err != nil { + return err + } + if tag.RowsAffected() == 0 { + return ErrWorkerNotFound + } + return nil + }, + ) +} + +// WorkerList returns a page of workers and the total count. +func (db *DB) WorkerList(ctx context.Context, actor Actor, cursor WorkerID, limit int, filter string) ([]Worker, int, error) { + var workers []Worker + var total int + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { + if actor.kind == actorAnon { + return ErrUnauthenticated + } + return nil + }, + func(tx pgx.Tx) error { + if filter != "" { + return ErrNotImplemented + } + return nil + }, + func(tx pgx.Tx) error { + if err := tx.QueryRow(ctx, + `SELECT count(*) FROM workers WHERE revoked_at IS NULL`, + ).Scan(&total); err != nil { + return err + } + + q := sb.PostgreSQL.NewSelectBuilder() + q.Select("id", "public_key", "org_id", "created_at"). + From("workers"). + Where(q.IsNull("revoked_at")). + OrderBy("id"). + Limit(limit) + if !cursor.IsZero() { + q.Where(q.GreaterThan("id", cursor)) + } + + sql, args := q.Build() + rows, err := tx.Query(ctx, sql, args...) + if err != nil { + return err + } + defer rows.Close() + + for rows.Next() { + var w Worker + if err := rows.Scan(&w.ID, &w.PublicKey, &w.OrgID, &w.CreatedAt); err != nil { + return err + } + workers = append(workers, w) + } + return rows.Err() + }, + ) + return workers, total, err +} + +// WorkerLookup finds an active worker by its ed25519 public key. +func (db *DB) WorkerLookup(ctx context.Context, actor Actor, publicKey []byte) (*Worker, error) { + var w Worker + err := db.apicall(ctx, actor, + func(tx pgx.Tx) error { return checkSystem(actor) }, + nil, + func(tx pgx.Tx) error { + if err := tx.QueryRow(ctx, + `SELECT id, public_key, org_id, created_at FROM workers WHERE public_key = $1 AND revoked_at IS NULL`, + publicKey, + ).Scan(&w.ID, &w.PublicKey, &w.OrgID, &w.CreatedAt); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return ErrWorkerNotFound + } + return err + } + return nil + }, + ) + return &w, err +} + +// hashToken returns the hex-encoded SHA-256 of a session token. +func hashToken(token string) string { + h := sha256.Sum256([]byte(token)) + return fmt.Sprintf("%x", h) +} + +// verifyHash parses a PHC-format argon2id string and compares. +// Format: $argon2id$v=19$m=65536,t=3,p=2$$ +func verifyHash(password, encoded string, pepper []byte) bool { + // $argon2id$v=19$m=65536,t=3,p=2$salt$key → 6 parts + parts := strings.Split(encoded, "$") + if len(parts) != 6 || parts[1] != "argon2id" { + return false + } + + var memory, time uint32 + var threads uint8 + if _, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &time, &threads); err != nil { + return false + } + + salt, err := base64.RawStdEncoding.DecodeString(parts[4]) + if err != nil { + return false + } + expectedKey, err := base64.RawStdEncoding.DecodeString(parts[5]) + if err != nil { + return false + } + + mac := hmac.New(sha256.New, pepper) + mac.Write([]byte(password)) + peppered := mac.Sum(nil) + + key := argon2.IDKey(peppered, salt, time, memory, threads, uint32(len(expectedKey))) + + return subtle.ConstantTimeCompare(key, expectedKey) == 1 +} + +// hashPassword produces a PHC-format string: +// $argon2id$v=19$m=65536,t=3,p=2$$ +func hashPassword(password string, pepper []byte) (string, error) { + salt := make([]byte, saltLen) + if _, err := rand.Read(salt); err != nil { + return "", err + } + + // Apply pepper: HMAC-SHA256(pepper, password) + mac := hmac.New(sha256.New, pepper) + mac.Write([]byte(password)) + peppered := mac.Sum(nil) + + key := argon2.IDKey(peppered, salt, argonTime, argonMemory, argonThreads, argonKeyLen) + + return fmt.Sprintf("$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s", + argon2.Version, + argonMemory, argonTime, argonThreads, + base64.RawStdEncoding.EncodeToString(salt), + base64.RawStdEncoding.EncodeToString(key), + ), nil +} + +// checkNotSoleOwner returns ErrSoleOwner if the user is the only owner of any org. +func checkNotSoleOwner(ctx context.Context, tx pgx.Tx, userID UserID) error { + var slug string + err := tx.QueryRow(ctx, + `SELECT o.slug FROM org_members m + JOIN organizations o ON o.id = m.org_id + WHERE m.role = 'owner' AND o.deleted_at IS NULL + GROUP BY o.id, o.slug + HAVING count(*) = 1 AND bool_or(m.user_id = $1) + LIMIT 1`, userID, + ).Scan(&slug) + if err == nil { + return ErrSoleOwner + } + if errors.Is(err, pgx.ErrNoRows) { + return nil + } + return err +} + +// resolveUser locks and returns the user ID within a transaction. +func resolveUser(ctx context.Context, tx pgx.Tx, ref UserRef) (UserID, error) { + col, val := ref.where() + q := sb.PostgreSQL.NewSelectBuilder() + + sql, args := q.Select("id").From("users"). + Where(q.Equal(col, val), q.IsNull("deleted_at")). + ForUpdate(). + Build() + + var id UserID + if err := tx.QueryRow(ctx, sql, args...).Scan(&id); err != nil { + var zero UserID + if errors.Is(err, pgx.ErrNoRows) { + return zero, ErrUserNotFound + } + + return zero, err + } + + return id, nil +} + +// resolveOrg locks and returns the org ID within a transaction. +func resolveOrg(ctx context.Context, tx pgx.Tx, ref OrgRef) (OrgID, error) { + col, val := ref.where() + q := sb.PostgreSQL.NewSelectBuilder() + + sql, args := q.Select("id").From("organizations"). + Where(q.Equal(col, val), q.IsNull("deleted_at")). + ForUpdate(). + Build() + + var id OrgID + if err := tx.QueryRow(ctx, sql, args...).Scan(&id); err != nil { + var zero OrgID + if errors.Is(err, pgx.ErrNoRows) { + return zero, ErrOrgNotFound + } + + return zero, err + } + + return id, nil +} + +// ValidateEmail checks that the value is a valid email address. +func ValidateEmail(value string) error { + if _, err := mail.ParseAddress(value); err != nil { + return ErrInvalidEmail + } + return nil +} + +// ValidateSlug checks format and returns the normalized (lowercased) slug. +func ValidateSlug(value string) (string, error) { + if len(value) < 2 || len(value) > 64 || !slugRe.MatchString(value) { + return "", ErrInvalidSlug + } + return strings.ToLower(value), nil +} + +// ValidateRole checks that the value is a valid role string. +// Derives valid roles from rolePermissions — single source of truth. +func ValidateRole(value string) error { + if _, ok := rolePermissions[value]; !ok { + return ErrInvalidRole + } + return nil +} diff --git a/cmd/mirumd/id.go b/cmd/mirumd/id.go new file mode 100644 index 0000000..09ec141 --- /dev/null +++ b/cmd/mirumd/id.go @@ -0,0 +1,302 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package main + +import ( + "database/sql/driver" + "errors" + "fmt" + "log/slog" + "strings" + + "github.com/google/uuid" +) + +var ( + ErrBadID = errors.New("database: bad id") + ErrBadIDPrefix = errors.New("database: wrong id prefix") +) + +type ( + OrgID = ID[OrgKind] + UserID = ID[UserKind] + WorkerID = ID[WorkerKind] +) + +// IDKind is a phantom-type tag that distinguishes otherwise-identical +// 16-byte IDs at the Go type level. Each tag is a zero-sized struct +// that carries only its 3-letter prefix. +type IDKind interface { + UserKind | OrgKind | WorkerKind + Prefix() string +} + +type ( + UserKind struct{} + OrgKind struct{} + WorkerKind struct{} +) + +func (OrgKind) Prefix() string { return "org" } +func (UserKind) Prefix() string { return "usr" } +func (WorkerKind) Prefix() string { return "wrk" } + +// ID[K] is a typed UUIDv7. Instantiations with different K are distinct +// types, so passing a UserID where an OrgID is expected is a compile +// error. Cross-kind conversion requires an explicit cast, visible in +// review. +type ID[K IDKind] uuid.UUID + +func NewID[K IDKind]() ID[K] { + return ID[K](uuid.Must(uuid.NewV7())) +} + +func IDFromBytes[K IDKind](b []byte) (ID[K], error) { + var zero ID[K] + if len(b) != 16 { + return zero, fmt.Errorf("%w: got %d bytes", ErrBadID, len(b)) + } + u := uuid.UUID(b) + if err := validateV7(u); err != nil { + return zero, err + } + return ID[K](u), nil +} + +func validateV7(u uuid.UUID) error { + if v := u.Variant(); v != uuid.RFC4122 { + return fmt.Errorf("%w: variant %s", ErrBadID, v) + } + if v := u.Version(); v != 7 { + return fmt.Errorf("%w: version %d, want 7", ErrBadID, v) + } + return nil +} + +func (id ID[K]) UUID() uuid.UUID { return uuid.UUID(id) } +func (id ID[K]) Bytes() []byte { return id[:] } + +func (id ID[K]) IsZero() bool { + var zero ID[K] + return id == zero +} + +// String returns the prefixed form for logs, JSON, errors and anywhere +// the entity type isn't obvious from context. For URL path segments +// where the route already names the type, use Bare(). +func (id ID[K]) String() string { + var k K + return k.Prefix() + "_" + encodeBase58(uuid.UUID(id)) +} + +// Bare returns the base58 form without a type prefix — ≤ 22 chars. +// Use this for URL path segments where the route already identifies +// the entity ("/org/:id"); use String() everywhere else. +func (id ID[K]) Bare() string { + return encodeBase58(uuid.UUID(id)) +} + +func (id ID[K]) LogValue() slog.Value { + return slog.StringValue(id.String()) +} + +func (id ID[K]) MarshalText() ([]byte, error) { + return []byte(id.String()), nil +} + +func (id *ID[K]) UnmarshalText(b []byte) error { + parsed, err := ParseID[K](string(b)) + if err != nil { + return err + } + *id = parsed + return nil +} + +func (id *ID[K]) Scan(src any) error { + var u uuid.UUID + if err := u.Scan(src); err != nil { + return err + } + *id = ID[K](u) + return nil +} + +func (id ID[K]) Value() (driver.Value, error) { + return uuid.UUID(id).Value() +} + +// ParseID accepts the prefixed base58 ("_"), bare base58 +// (≤ 22 chars, as returned by Bare()), or a canonical UUID string. +// Bare forms stay supported so existing CLI flags, manual SQL lookups +// and URL path parameters keep working without a flag day. +func ParseID[K IDKind](s string) (ID[K], error) { + var k K + u, err := parseID(k.Prefix(), s) + return ID[K](u), err +} + +func MustParseID[K IDKind](s string) ID[K] { + id, err := ParseID[K](s) + if err != nil { + panic(err) + } + return id +} + +// ParseAnyID parses a prefixed, bare base58, or canonical UUID string +// into raw 16 bytes. Unlike ParseID[K], it does not require a known +// entity kind — any 3-letter prefix is stripped silently. Intended for +// CLI boundary code that dispatches on proto field names. +func ParseAnyID(s string) ([16]byte, error) { + if s == "" { + return [16]byte{}, fmt.Errorf("%w: empty", ErrBadID) + } + // Strip any typed prefix. + if len(s) > 4 && s[3] == '_' { + s = s[4:] + } + if len(s) <= 22 { + return decodeBase58(s) + } + u, err := uuid.Parse(s) + if err != nil { + return [16]byte{}, fmt.Errorf("%w: %w", ErrBadID, err) + } + return u, nil +} + +// FormatAnyID formats raw 16-byte ID as bare base58. Returns base64 for +// non-16-byte inputs as a fallback. +func FormatAnyID(b []byte) string { + if len(b) == 16 { + return encodeBase58(uuid.UUID(b)) + } + return fmt.Sprintf("%x", b) +} + +func parseID(prefix, s string) (uuid.UUID, error) { + if s == "" { + return uuid.Nil, fmt.Errorf("%w: empty", ErrBadID) + } + + if rest, ok := strings.CutPrefix(s, prefix+"_"); ok { + return decodeBase58(rest) + } + + // Typed prefix with the wrong value — reject so a cross-type paste + // doesn't silently fall through to the bare path. + if len(s) > 4 && s[3] == '_' { + return uuid.Nil, fmt.Errorf("%w: want %q, got %q", ErrBadIDPrefix, prefix, s[:3]) + } + + // Bare form: base58 (≤ 22 chars, from Bare()) or canonical UUID + // (32–36 chars, from manual SQL / legacy CLI). + if len(s) <= 22 { + return decodeBase58(s) + } + + u, err := uuid.Parse(s) + if err != nil { + return uuid.Nil, fmt.Errorf("%w: %w", ErrBadID, err) + } + + return u, nil +} + +// --- base58 (Bitcoin alphabet) --- +// +// 16 bytes fit in ≤ 22 base58 chars: log₅₈(2¹²⁸) ≈ 21.86. +// UUIDv7 values in the post-1970 range always have a non-zero +// leading byte, so encoded length is effectively a constant 21-22. + +const b58Alphabet = "123456789ABCDEFGHJKLMNPQRSTUVWXYZabcdefghijkmnopqrstuvwxyz" + +var b58Index [256]byte + +func init() { + for i := range b58Index { + b58Index[i] = 0xff + } + for i := 0; i < len(b58Alphabet); i++ { + b58Index[b58Alphabet[i]] = byte(i) + } +} + +func encodeBase58(src uuid.UUID) string { + zeros := 0 + for zeros < 16 && src[zeros] == 0 { + zeros++ + } + + buf := src // array copy; long-division is in place + start := zeros + out := make([]byte, 0, 22) + + for start < 16 { + rem := 0 + for i := start; i < 16; i++ { + v := rem*256 + int(buf[i]) + buf[i] = byte(v / 58) + rem = v % 58 + } + out = append(out, b58Alphabet[rem]) + for start < 16 && buf[start] == 0 { + start++ + } + } + + for i := 0; i < zeros; i++ { + out = append(out, b58Alphabet[0]) + } + + // Reverse into big-endian order. + for i, j := 0, len(out)-1; i < j; i, j = i+1, j-1 { + out[i], out[j] = out[j], out[i] + } + + return string(out) +} + +// decodeBase58 is strict about length — any input that round-trips to a +// value of a different byte length is rejected; we never want a +// 15-byte or 17-byte payload masquerading as a UUID. +func decodeBase58(s string) (uuid.UUID, error) { + if s == "" { + return uuid.Nil, fmt.Errorf("%w: empty", ErrBadID) + } + zeros := 0 + for zeros < len(s) && s[zeros] == b58Alphabet[0] { + zeros++ + } + + var out uuid.UUID + for i := zeros; i < len(s); i++ { + v := b58Index[s[i]] + if v == 0xff { + return uuid.Nil, fmt.Errorf("%w: bad base58 char %q", ErrBadID, s[i]) + } + carry := int(v) + for j := 15; j >= 0; j-- { + acc := int(out[j])*58 + carry + out[j] = byte(acc) + carry = acc >> 8 + } + if carry != 0 { + return uuid.Nil, fmt.Errorf("%w: overflow", ErrBadID) + } + } + + // The '1' prefix count in the string must equal the leading-zero + // byte count in the result — any mismatch means the input decoded + // to a different byte length than a UUID. + actualZeros := 0 + for actualZeros < 16 && out[actualZeros] == 0 { + actualZeros++ + } + if actualZeros != zeros { + return uuid.Nil, fmt.Errorf("%w: length mismatch", ErrBadID) + } + return out, nil +} diff --git a/cmd/mirumd/id_test.go b/cmd/mirumd/id_test.go new file mode 100644 index 0000000..0d06fc1 --- /dev/null +++ b/cmd/mirumd/id_test.go @@ -0,0 +1,76 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package main + +import ( + "testing" + + "github.com/google/uuid" +) + +func TestIDRoundTrip(t *testing.T) { + for range 50 { + id := NewID[UserKind]() + got, err := ParseID[UserKind](id.String()) + if err != nil { + t.Fatalf("parse prefixed: %v", err) + } + if got != id { + t.Fatalf("prefixed: got %v, want %v", got, id) + } + got, err = ParseID[UserKind](id.Bare()) + if err != nil { + t.Fatalf("parse bare: %v", err) + } + if got != id { + t.Fatalf("bare: got %v, want %v", got, id) + } + } +} + +func TestIDRoundTripCanonicalUUID(t *testing.T) { + id := NewID[OrgKind]() + got, err := ParseID[OrgKind](id.UUID().String()) + if err != nil { + t.Fatalf("parse canonical: %v", err) + } + if got != id { + t.Fatalf("canonical: got %v, want %v", got, id) + } +} + +func TestIDCrossPrefixRejected(t *testing.T) { + id := NewID[UserKind]() + _, err := ParseID[OrgKind](id.String()) + if err == nil { + t.Fatal("expected error parsing usr_ as org") + } +} + +func TestBase58EdgeCases(t *testing.T) { + cases := [][16]byte{ + {}, // all zeros + {0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1}, // minimal + {0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, + 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff}, // max + } + for _, raw := range cases { + enc := encodeBase58(uuid.UUID(raw)) + dec, err := decodeBase58(enc) + if err != nil { + t.Fatalf("decode(%q): %v", enc, err) + } + if dec != uuid.UUID(raw) { + t.Fatalf("roundtrip: got %x, want %x", dec, raw) + } + } +} + +func TestIDFromBytesRejectsV4(t *testing.T) { + v4 := uuid.Must(uuid.NewRandom()) // v4 + _, err := IDFromBytes[UserKind](v4[:]) + if err == nil { + t.Fatal("expected v4 rejection") + } +} diff --git a/cmd/mirumd/main.go b/cmd/mirumd/main.go index 9898929..73aa3f9 100644 --- a/cmd/mirumd/main.go +++ b/cmd/mirumd/main.go @@ -13,7 +13,6 @@ import ( "os" "dimidiumlabs/mirum/internal/config" - "dimidiumlabs/mirum/internal/database" "dimidiumlabs/mirum/internal/forges" "dimidiumlabs/mirum/internal/protocol/pb" "dimidiumlabs/mirum/internal/protocol/pb/pbconnect" @@ -82,7 +81,7 @@ func daemon(configFile, socketFlag string) error { ctx, cancel := context.WithCancel(sup.WaitForStop(context.Background())) defer cancel() - db, err := database.Open(ctx, cfg.DatabaseUri) + db, err := DatabaseOpen(ctx, cfg.DatabaseUri) if err != nil { slog.Error("couldn't open database", "err", err) return err @@ -115,7 +114,7 @@ func daemon(configFile, socketFlag string) error { adminSrv := hardenServer(&http.Server{ Handler: adminMux, ConnContext: func(ctx context.Context, _ net.Conn) context.Context { - return context.WithValue(ctx, actorKey{}, database.OperatorActor()) + return context.WithValue(ctx, actorKey{}, OperatorActor()) }, BaseContext: func(_ net.Listener) context.Context { return ctx diff --git a/cmd/mirumd/server.go b/cmd/mirumd/server.go index 3a86052..11585d9 100644 --- a/cmd/mirumd/server.go +++ b/cmd/mirumd/server.go @@ -12,7 +12,6 @@ import ( "time" "dimidiumlabs/mirum/internal/config" - "dimidiumlabs/mirum/internal/database" "dimidiumlabs/mirum/internal/forges" "dimidiumlabs/mirum/internal/protocol/pb" ) @@ -20,7 +19,7 @@ import ( // server holds the shared application state. type server struct { cfg *appConfig - db *database.DB + db *DB forge forges.Forge queue chan *pb.Task @@ -43,7 +42,7 @@ func (s *server) PurgeSessions(ctx context.Context) { for { select { case <-ticker.C: - if err := s.db.PurgeExpiredSessions(ctx); err != nil { + if err := s.db.UserSessionPurgeExpired(ctx); err != nil { slog.Error("purge sessions", "err", err) } case <-ctx.Done(): diff --git a/cmd/mirumd/server_admin.go b/cmd/mirumd/server_admin.go index f758c0f..7e1dae1 100644 --- a/cmd/mirumd/server_admin.go +++ b/cmd/mirumd/server_admin.go @@ -11,21 +11,18 @@ import ( "connectrpc.com/connect" "connectrpc.com/validate" - "github.com/google/uuid" "google.golang.org/protobuf/types/known/timestamppb" - "dimidiumlabs/mirum/internal/database" "dimidiumlabs/mirum/internal/protocol/pb" "dimidiumlabs/mirum/internal/protocol/pb/pbconnect" ) -// NewAdminHandler creates the ConnectRPC handler with validation and auth. +// NewAdminHandler creates the ConnectRPC handler with validation. +// Authorization is handled inside DB methods, not by an interceptor. func NewAdminHandler(srv *server) (string, http.Handler) { as := &adminService{srv: srv} - auth := &ApiAuthInterceptor{srv: srv} - return pbconnect.NewAdminHandler(as, - connect.WithInterceptors(validate.NewInterceptor(), auth), + connect.WithInterceptors(validate.NewInterceptor()), ) } @@ -52,19 +49,21 @@ var errSpecs = []struct { code connect.Code reason pb.ErrorReason }{ - {database.ErrUserNotFound, connect.CodeNotFound, pb.ErrorReason_ERROR_REASON_USER_NOT_FOUND}, - {database.ErrOrgNotFound, connect.CodeNotFound, pb.ErrorReason_ERROR_REASON_ORG_NOT_FOUND}, - {database.ErrWorkerNotFound, connect.CodeNotFound, pb.ErrorReason_ERROR_REASON_WORKER_NOT_FOUND}, - {database.ErrNotMember, connect.CodeNotFound, pb.ErrorReason_ERROR_REASON_MEMBER_NOT_FOUND}, - {database.ErrEmailTaken, connect.CodeAlreadyExists, pb.ErrorReason_ERROR_REASON_EMAIL_TAKEN}, - {database.ErrSlugTaken, connect.CodeAlreadyExists, pb.ErrorReason_ERROR_REASON_SLUG_TAKEN}, - {database.ErrAlreadyMember, connect.CodeAlreadyExists, pb.ErrorReason_ERROR_REASON_ALREADY_MEMBER}, - {database.ErrLastOwner, connect.CodeFailedPrecondition, pb.ErrorReason_ERROR_REASON_LAST_OWNER}, - {database.ErrSoleOwner, connect.CodeFailedPrecondition, pb.ErrorReason_ERROR_REASON_SOLE_OWNER}, - {database.ErrInvalidSlug, connect.CodeInvalidArgument, pb.ErrorReason_ERROR_REASON_INVALID_SLUG}, - {database.ErrInvalidRole, connect.CodeInvalidArgument, pb.ErrorReason_ERROR_REASON_INVALID_ROLE}, - {database.ErrReservedEmail, connect.CodeInvalidArgument, pb.ErrorReason_ERROR_REASON_RESERVED_EMAIL}, - {database.ErrFilterNotImplemented, connect.CodeUnimplemented, pb.ErrorReason_ERROR_REASON_UNIMPLEMENTED}, + {ErrUserNotFound, connect.CodeNotFound, pb.ErrorReason_ERROR_REASON_USER_NOT_FOUND}, + {ErrOrgNotFound, connect.CodeNotFound, pb.ErrorReason_ERROR_REASON_ORG_NOT_FOUND}, + {ErrWorkerNotFound, connect.CodeNotFound, pb.ErrorReason_ERROR_REASON_WORKER_NOT_FOUND}, + {ErrNotMember, connect.CodeNotFound, pb.ErrorReason_ERROR_REASON_MEMBER_NOT_FOUND}, + {ErrEmailTaken, connect.CodeAlreadyExists, pb.ErrorReason_ERROR_REASON_EMAIL_TAKEN}, + {ErrSlugTaken, connect.CodeAlreadyExists, pb.ErrorReason_ERROR_REASON_SLUG_TAKEN}, + {ErrAlreadyMember, connect.CodeAlreadyExists, pb.ErrorReason_ERROR_REASON_ALREADY_MEMBER}, + {ErrLastOwner, connect.CodeFailedPrecondition, pb.ErrorReason_ERROR_REASON_LAST_OWNER}, + {ErrSoleOwner, connect.CodeFailedPrecondition, pb.ErrorReason_ERROR_REASON_SOLE_OWNER}, + {ErrInvalidSlug, connect.CodeInvalidArgument, pb.ErrorReason_ERROR_REASON_INVALID_SLUG}, + {ErrInvalidRole, connect.CodeInvalidArgument, pb.ErrorReason_ERROR_REASON_INVALID_ROLE}, + {ErrReservedEmail, connect.CodeInvalidArgument, pb.ErrorReason_ERROR_REASON_RESERVED_EMAIL}, + {ErrPermissionDenied, connect.CodePermissionDenied, pb.ErrorReason_ERROR_REASON_PERMISSION_DENIED}, + {ErrUnauthenticated, connect.CodeUnauthenticated, pb.ErrorReason_ERROR_REASON_UNAUTHENTICATED}, + {ErrNotImplemented, connect.CodeUnimplemented, pb.ErrorReason_ERROR_REASON_UNIMPLEMENTED}, } func mapErr(err error) error { @@ -82,25 +81,33 @@ func mapErr(err error) error { // --- Ref converters --- -func userRef(r *pb.UserRef) database.UserRef { +func userRef(r *pb.UserRef) (UserRef, error) { switch v := r.GetRef().(type) { case *pb.UserRef_Id: - return database.UserByID(uuid.UUID(v.Id)) + id, err := IDFromBytes[UserKind](v.Id) + if err != nil { + return UserRef{}, err + } + return UserByID(id), nil case *pb.UserRef_Email: - return database.UserByEmail(v.Email) + return UserByEmail(v.Email), nil default: - return database.UserByID(uuid.Nil) + return UserRef{}, nil } } -func orgRef(r *pb.OrgRef) database.OrgRef { +func orgRef(r *pb.OrgRef) (OrgRef, error) { switch v := r.GetRef().(type) { case *pb.OrgRef_Id: - return database.OrgByID(uuid.UUID(v.Id)) + id, err := IDFromBytes[OrgKind](v.Id) + if err != nil { + return OrgRef{}, err + } + return OrgByID(id), nil case *pb.OrgRef_Slug: - return database.OrgBySlug(v.Slug) + return OrgBySlug(v.Slug), nil default: - return database.OrgByID(uuid.Nil) + return OrgRef{}, nil } } @@ -125,7 +132,7 @@ const ( maxPageSize = 200 ) -func pageParams(p *pb.PageRequest) (cursor uuid.UUID, limit int) { +func pageParams[K IDKind](p *pb.PageRequest) (cursor ID[K], limit int, err error) { limit = defaultPageSize if p != nil { if p.PageSize > 0 && int(p.PageSize) < maxPageSize { @@ -134,48 +141,51 @@ func pageParams(p *pb.PageRequest) (cursor uuid.UUID, limit int) { limit = maxPageSize } if len(p.Cursor) == 16 { - cursor = uuid.UUID(p.Cursor) + cursor, err = IDFromBytes[K](p.Cursor) + if err != nil { + return + } } } return } -func pageResponse(items int, limit int, lastID uuid.UUID, total int) *pb.PageResponse { +func pageResponse[K IDKind](items int, limit int, lastID ID[K], total int) *pb.PageResponse { resp := &pb.PageResponse{TotalCount: int32(total)} if items == limit { - resp.NextCursor = lastID[:] + resp.NextCursor = lastID.Bytes() } return resp } // --- Proto converters --- -func userToProto(u database.User) *pb.User { +func userToProto(u User) *pb.User { return &pb.User{ - Id: u.ID[:], Email: u.Email, CreatedAt: timestamppb.New(u.CreatedAt), + Id: u.ID.Bytes(), Email: u.Email, CreatedAt: timestamppb.New(u.CreatedAt), } } -func orgToProto(o database.Organization) *pb.Org { +func orgToProto(o Organization) *pb.Org { return &pb.Org{ - Id: o.ID[:], Name: o.Name, Slug: o.Slug, + Id: o.ID.Bytes(), Name: o.Name, Slug: o.Slug, Public: o.Public, CreatedAt: timestamppb.New(o.CreatedAt), } } -func memberToProto(m database.OrgMember) *pb.OrgMemberInfo { +func memberToProto(m OrgMember) *pb.OrgMemberInfo { return &pb.OrgMemberInfo{ User: userToProto(m.User), Role: roleToProto[m.Role], JoinedAt: timestamppb.New(m.JoinedAt), } } -func workerToProto(w database.Worker) *pb.Worker { +func workerToProto(w Worker) *pb.Worker { pw := &pb.Worker{ - Id: w.ID[:], PublicKey: w.PublicKey, CreatedAt: timestamppb.New(w.CreatedAt), + Id: w.ID.Bytes(), PublicKey: w.PublicKey, CreatedAt: timestamppb.New(w.CreatedAt), } if w.OrgID != nil { - pw.OrgId = w.OrgID[:] + pw.OrgId = w.OrgID.Bytes() } return pw } @@ -187,11 +197,15 @@ func (a *adminService) UserCreate(ctx context.Context, req *connect.Request[pb.U if err != nil { return nil, mapErr(err) } - return connect.NewResponse(&pb.UserCreateResponse{Id: id[:]}), nil + return connect.NewResponse(&pb.UserCreateResponse{Id: id.Bytes()}), nil } func (a *adminService) UserGet(ctx context.Context, req *connect.Request[pb.UserGetRequest]) (*connect.Response[pb.UserGetResponse], error) { - u, err := a.srv.db.GetUser(ctx, ActorFromContext(ctx), userRef(req.Msg.User)) + ref, err := userRef(req.Msg.User) + if err != nil { + return nil, mapErr(err) + } + u, err := a.srv.db.UserGet(ctx, ActorFromContext(ctx), ref) if err != nil { return nil, mapErr(err) } @@ -199,13 +213,16 @@ func (a *adminService) UserGet(ctx context.Context, req *connect.Request[pb.User } func (a *adminService) UserList(ctx context.Context, req *connect.Request[pb.UserListRequest]) (*connect.Response[pb.UserListResponse], error) { - cursor, limit := pageParams(req.Msg.Page) + cursor, limit, err := pageParams[UserKind](req.Msg.Page) + if err != nil { + return nil, mapErr(err) + } filter := "" if req.Msg.Filter != nil { filter = *req.Msg.Filter } - users, total, err := a.srv.db.ListUsers(ctx, ActorFromContext(ctx), cursor, limit, filter) + users, total, err := a.srv.db.UserList(ctx, ActorFromContext(ctx), cursor, limit, filter) if err != nil { return nil, mapErr(err) } @@ -215,7 +232,7 @@ func (a *adminService) UserList(ctx context.Context, req *connect.Request[pb.Use out[i] = userToProto(users[i]) } - var lastID uuid.UUID + var lastID UserID if len(users) > 0 { lastID = users[len(users)-1].ID } @@ -227,14 +244,22 @@ func (a *adminService) UserList(ctx context.Context, req *connect.Request[pb.Use } func (a *adminService) UserUpdate(ctx context.Context, req *connect.Request[pb.UserUpdateRequest]) (*connect.Response[pb.UserUpdateResponse], error) { - if err := a.srv.db.UserUpdate(ctx, ActorFromContext(ctx), userRef(req.Msg.User), req.Msg.Email, req.Msg.Password, []byte(a.srv.cfg.Pepper)); err != nil { + ref, err := userRef(req.Msg.User) + if err != nil { + return nil, mapErr(err) + } + if err := a.srv.db.UserUpdate(ctx, ActorFromContext(ctx), ref, req.Msg.Email, req.Msg.Password, []byte(a.srv.cfg.Pepper)); err != nil { return nil, mapErr(err) } return connect.NewResponse(&pb.UserUpdateResponse{}), nil } func (a *adminService) UserDelete(ctx context.Context, req *connect.Request[pb.UserDeleteRequest]) (*connect.Response[pb.UserDeleteResponse], error) { - if err := a.srv.db.UserDelete(ctx, ActorFromContext(ctx), userRef(req.Msg.User)); err != nil { + ref, err := userRef(req.Msg.User) + if err != nil { + return nil, mapErr(err) + } + if err := a.srv.db.UserDelete(ctx, ActorFromContext(ctx), ref); err != nil { return nil, mapErr(err) } return connect.NewResponse(&pb.UserDeleteResponse{}), nil @@ -243,19 +268,27 @@ func (a *adminService) UserDelete(ctx context.Context, req *connect.Request[pb.U // --- Org handlers --- func (a *adminService) OrgCreate(ctx context.Context, req *connect.Request[pb.OrgCreateRequest]) (*connect.Response[pb.OrgCreateResponse], error) { - slug, err := database.ValidateSlug(req.Msg.Slug) + slug, err := ValidateSlug(req.Msg.Slug) if err != nil { return nil, mapErr(err) } - id, err := a.srv.db.CreateOrganization(ctx, ActorFromContext(ctx), req.Msg.Name, slug, req.Msg.Public, userRef(req.Msg.Owner)) + owner, err := userRef(req.Msg.Owner) if err != nil { return nil, mapErr(err) } - return connect.NewResponse(&pb.OrgCreateResponse{Id: id[:]}), nil + id, err := a.srv.db.OrgCreate(ctx, ActorFromContext(ctx), req.Msg.Name, slug, req.Msg.Public, owner) + if err != nil { + return nil, mapErr(err) + } + return connect.NewResponse(&pb.OrgCreateResponse{Id: id.Bytes()}), nil } func (a *adminService) OrgGet(ctx context.Context, req *connect.Request[pb.OrgGetRequest]) (*connect.Response[pb.OrgGetResponse], error) { - o, err := a.srv.db.GetOrg(ctx, ActorFromContext(ctx), orgRef(req.Msg.Org)) + ref, err := orgRef(req.Msg.Org) + if err != nil { + return nil, mapErr(err) + } + o, err := a.srv.db.OrgGet(ctx, ActorFromContext(ctx), ref) if err != nil { return nil, mapErr(err) } @@ -263,13 +296,16 @@ func (a *adminService) OrgGet(ctx context.Context, req *connect.Request[pb.OrgGe } func (a *adminService) OrgList(ctx context.Context, req *connect.Request[pb.OrgListRequest]) (*connect.Response[pb.OrgListResponse], error) { - cursor, limit := pageParams(req.Msg.Page) + cursor, limit, err := pageParams[OrgKind](req.Msg.Page) + if err != nil { + return nil, mapErr(err) + } filter := "" if req.Msg.Filter != nil { filter = *req.Msg.Filter } - orgs, total, err := a.srv.db.ListOrganizations(ctx, ActorFromContext(ctx), cursor, limit, filter) + orgs, total, err := a.srv.db.OrgList(ctx, ActorFromContext(ctx), cursor, limit, filter) if err != nil { return nil, mapErr(err) } @@ -279,7 +315,7 @@ func (a *adminService) OrgList(ctx context.Context, req *connect.Request[pb.OrgL out[i] = orgToProto(orgs[i]) } - var lastID uuid.UUID + var lastID OrgID if len(orgs) > 0 { lastID = orgs[len(orgs)-1].ID } @@ -293,20 +329,28 @@ func (a *adminService) OrgList(ctx context.Context, req *connect.Request[pb.OrgL func (a *adminService) OrgUpdate(ctx context.Context, req *connect.Request[pb.OrgUpdateRequest]) (*connect.Response[pb.OrgUpdateResponse], error) { var slug *string if req.Msg.Slug != nil { - s, err := database.ValidateSlug(*req.Msg.Slug) + s, err := ValidateSlug(*req.Msg.Slug) if err != nil { return nil, mapErr(err) } slug = &s } - if err := a.srv.db.UpdateOrganization(ctx, ActorFromContext(ctx), orgRef(req.Msg.Org), req.Msg.Name, slug, req.Msg.Public); err != nil { + ref, err := orgRef(req.Msg.Org) + if err != nil { + return nil, mapErr(err) + } + if err := a.srv.db.OrgUpdate(ctx, ActorFromContext(ctx), ref, req.Msg.Name, slug, req.Msg.Public); err != nil { return nil, mapErr(err) } return connect.NewResponse(&pb.OrgUpdateResponse{}), nil } func (a *adminService) OrgDelete(ctx context.Context, req *connect.Request[pb.OrgDeleteRequest]) (*connect.Response[pb.OrgDeleteResponse], error) { - if err := a.srv.db.DeleteOrganization(ctx, ActorFromContext(ctx), orgRef(req.Msg.Org)); err != nil { + ref, err := orgRef(req.Msg.Org) + if err != nil { + return nil, mapErr(err) + } + if err := a.srv.db.OrgDelete(ctx, ActorFromContext(ctx), ref); err != nil { return nil, mapErr(err) } return connect.NewResponse(&pb.OrgDeleteResponse{}), nil @@ -319,14 +363,30 @@ func (a *adminService) OrgMemberAdd(ctx context.Context, req *connect.Request[pb if !ok { return nil, newAPIError(connect.CodeInvalidArgument, pb.ErrorReason_ERROR_REASON_INVALID_ROLE, nil) } - if err := a.srv.db.AddOrgMember(ctx, ActorFromContext(ctx), orgRef(req.Msg.Org), userRef(req.Msg.User), role); err != nil { + org, err := orgRef(req.Msg.Org) + if err != nil { + return nil, mapErr(err) + } + user, err := userRef(req.Msg.User) + if err != nil { + return nil, mapErr(err) + } + if err := a.srv.db.OrgMemberAdd(ctx, ActorFromContext(ctx), org, user, role); err != nil { return nil, mapErr(err) } return connect.NewResponse(&pb.OrgMemberAddResponse{}), nil } func (a *adminService) OrgMemberGet(ctx context.Context, req *connect.Request[pb.OrgMemberGetRequest]) (*connect.Response[pb.OrgMemberGetResponse], error) { - m, err := a.srv.db.GetOrgMember(ctx, ActorFromContext(ctx), orgRef(req.Msg.Org), userRef(req.Msg.User)) + org, err := orgRef(req.Msg.Org) + if err != nil { + return nil, mapErr(err) + } + user, err := userRef(req.Msg.User) + if err != nil { + return nil, mapErr(err) + } + m, err := a.srv.db.OrgMemberGet(ctx, ActorFromContext(ctx), org, user) if err != nil { return nil, mapErr(err) } @@ -334,13 +394,20 @@ func (a *adminService) OrgMemberGet(ctx context.Context, req *connect.Request[pb } func (a *adminService) OrgMemberList(ctx context.Context, req *connect.Request[pb.OrgMemberListRequest]) (*connect.Response[pb.OrgMemberListResponse], error) { - cursor, limit := pageParams(req.Msg.Page) + cursor, limit, err := pageParams[UserKind](req.Msg.Page) + if err != nil { + return nil, mapErr(err) + } filter := "" if req.Msg.Filter != nil { filter = *req.Msg.Filter } + org, err := orgRef(req.Msg.Org) + if err != nil { + return nil, mapErr(err) + } - members, total, err := a.srv.db.ListOrgMembers(ctx, ActorFromContext(ctx), orgRef(req.Msg.Org), cursor, limit, filter) + members, total, err := a.srv.db.OrgMembersList(ctx, ActorFromContext(ctx), org, cursor, limit, filter) if err != nil { return nil, mapErr(err) } @@ -350,7 +417,7 @@ func (a *adminService) OrgMemberList(ctx context.Context, req *connect.Request[p out[i] = memberToProto(members[i]) } - var lastID uuid.UUID + var lastID UserID if len(members) > 0 { lastID = members[len(members)-1].User.ID } @@ -366,14 +433,30 @@ func (a *adminService) OrgMemberUpdate(ctx context.Context, req *connect.Request if !ok { return nil, newAPIError(connect.CodeInvalidArgument, pb.ErrorReason_ERROR_REASON_INVALID_ROLE, nil) } - if err := a.srv.db.UpdateOrgMemberRole(ctx, ActorFromContext(ctx), orgRef(req.Msg.Org), userRef(req.Msg.User), role); err != nil { + org, err := orgRef(req.Msg.Org) + if err != nil { + return nil, mapErr(err) + } + user, err := userRef(req.Msg.User) + if err != nil { + return nil, mapErr(err) + } + if err := a.srv.db.OrgMemberUpdateRole(ctx, ActorFromContext(ctx), org, user, role); err != nil { return nil, mapErr(err) } return connect.NewResponse(&pb.OrgMemberUpdateResponse{}), nil } func (a *adminService) OrgMemberRemove(ctx context.Context, req *connect.Request[pb.OrgMemberRemoveRequest]) (*connect.Response[pb.OrgMemberRemoveResponse], error) { - if err := a.srv.db.RemoveOrgMember(ctx, ActorFromContext(ctx), orgRef(req.Msg.Org), userRef(req.Msg.User)); err != nil { + org, err := orgRef(req.Msg.Org) + if err != nil { + return nil, mapErr(err) + } + user, err := userRef(req.Msg.User) + if err != nil { + return nil, mapErr(err) + } + if err := a.srv.db.OrgMemberRemove(ctx, ActorFromContext(ctx), org, user); err != nil { return nil, mapErr(err) } return connect.NewResponse(&pb.OrgMemberRemoveResponse{}), nil @@ -382,20 +465,27 @@ func (a *adminService) OrgMemberRemove(ctx context.Context, req *connect.Request // --- Worker handlers --- func (a *adminService) WorkerCreate(ctx context.Context, req *connect.Request[pb.WorkerCreateRequest]) (*connect.Response[pb.WorkerCreateResponse], error) { - var org *database.OrgRef + var org *OrgRef if req.Msg.Org != nil { - r := orgRef(req.Msg.Org) + r, err := orgRef(req.Msg.Org) + if err != nil { + return nil, mapErr(err) + } org = &r } - id, err := a.srv.db.CreateWorker(ctx, ActorFromContext(ctx), req.Msg.PublicKey, org) + id, err := a.srv.db.WorkerCreate(ctx, ActorFromContext(ctx), req.Msg.PublicKey, org) if err != nil { return nil, mapErr(err) } - return connect.NewResponse(&pb.WorkerCreateResponse{Id: id[:]}), nil + return connect.NewResponse(&pb.WorkerCreateResponse{Id: id.Bytes()}), nil } func (a *adminService) WorkerGet(ctx context.Context, req *connect.Request[pb.WorkerGetRequest]) (*connect.Response[pb.WorkerGetResponse], error) { - w, err := a.srv.db.GetWorker(ctx, ActorFromContext(ctx), uuid.UUID(req.Msg.Id)) + wid, err := IDFromBytes[WorkerKind](req.Msg.Id) + if err != nil { + return nil, mapErr(err) + } + w, err := a.srv.db.WorkerGet(ctx, ActorFromContext(ctx), wid) if err != nil { return nil, mapErr(err) } @@ -403,13 +493,16 @@ func (a *adminService) WorkerGet(ctx context.Context, req *connect.Request[pb.Wo } func (a *adminService) WorkerList(ctx context.Context, req *connect.Request[pb.WorkerListRequest]) (*connect.Response[pb.WorkerListResponse], error) { - cursor, limit := pageParams(req.Msg.Page) + cursor, limit, err := pageParams[WorkerKind](req.Msg.Page) + if err != nil { + return nil, mapErr(err) + } filter := "" if req.Msg.Filter != nil { filter = *req.Msg.Filter } - workers, total, err := a.srv.db.ListWorkers(ctx, ActorFromContext(ctx), cursor, limit, filter) + workers, total, err := a.srv.db.WorkerList(ctx, ActorFromContext(ctx), cursor, limit, filter) if err != nil { return nil, mapErr(err) } @@ -419,7 +512,7 @@ func (a *adminService) WorkerList(ctx context.Context, req *connect.Request[pb.W out[i] = workerToProto(workers[i]) } - var lastID uuid.UUID + var lastID WorkerID if len(workers) > 0 { lastID = workers[len(workers)-1].ID } @@ -431,7 +524,11 @@ func (a *adminService) WorkerList(ctx context.Context, req *connect.Request[pb.W } func (a *adminService) WorkerDelete(ctx context.Context, req *connect.Request[pb.WorkerDeleteRequest]) (*connect.Response[pb.WorkerDeleteResponse], error) { - if err := a.srv.db.DeleteWorker(ctx, ActorFromContext(ctx), uuid.UUID(req.Msg.Id)); err != nil { + wid, err := IDFromBytes[WorkerKind](req.Msg.Id) + if err != nil { + return nil, mapErr(err) + } + if err := a.srv.db.WorkerDelete(ctx, ActorFromContext(ctx), wid); err != nil { return nil, mapErr(err) } return connect.NewResponse(&pb.WorkerDeleteResponse{}), nil diff --git a/cmd/mirumd/server_grpc.go b/cmd/mirumd/server_grpc.go index 580154e..a41604d 100644 --- a/cmd/mirumd/server_grpc.go +++ b/cmd/mirumd/server_grpc.go @@ -19,7 +19,6 @@ import ( "connectrpc.com/validate" "dimidiumlabs/mirum/internal/config" - "dimidiumlabs/mirum/internal/database" "dimidiumlabs/mirum/internal/protocol" "dimidiumlabs/mirum/internal/protocol/pb" "dimidiumlabs/mirum/internal/protocol/pb/pbconnect" @@ -62,7 +61,7 @@ func NewGrpcServer(ctx context.Context, srv *server) *http.Server { return errors.New("ed25519 certificate required") } - if _, err := srv.db.LookupWorker(context.Background(), database.SystemActor(), pubKey); err != nil { + if _, err := srv.db.WorkerLookup(context.Background(), SystemActor(), pubKey); err != nil { return fmt.Errorf("unknown worker: %w", err) } diff --git a/cmd/mirumd/server_web.go b/cmd/mirumd/server_web.go index 877d700..13c5154 100644 --- a/cmd/mirumd/server_web.go +++ b/cmd/mirumd/server_web.go @@ -22,7 +22,6 @@ import ( "github.com/go-chi/httprate" "dimidiumlabs/mirum/internal/config" - "dimidiumlabs/mirum/internal/database" "dimidiumlabs/mirum/internal/forges" "dimidiumlabs/mirum/internal/protocol/pb" ) @@ -126,18 +125,18 @@ type webHandler struct { } // ActorFromContext returns the authenticated actor, or AnonActor if none. -func ActorFromContext(ctx context.Context) database.Actor { - if v, ok := ctx.Value(actorKey{}).(database.Actor); ok { +func ActorFromContext(ctx context.Context) Actor { + if v, ok := ctx.Value(actorKey{}).(Actor); ok { return v } - return database.AnonActor() + return AnonActor() } // SessionMiddleware resolves the session cookie and puts the Actor in context. func (h *webHandler) SessionMiddleware(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if c, err := r.Cookie(sessionCookie); err == nil { - if actor, err := h.srv.db.UserGetSession(r.Context(), database.SystemActor(), c.Value); err == nil { + if actor, err := h.srv.db.UserSessionGet(r.Context(), SystemActor(), c.Value); err == nil { ctx := context.WithValue(r.Context(), actorKey{}, actor) r = r.WithContext(ctx) } @@ -146,12 +145,12 @@ func (h *webHandler) SessionMiddleware(next http.Handler) http.Handler { }) } -type authedHandler func(w http.ResponseWriter, r *http.Request, actor database.Actor) +type authedHandler func(w http.ResponseWriter, r *http.Request, actor Actor) func authonly(next authedHandler) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { actor := ActorFromContext(r.Context()) - if actor.Kind() == database.KindAnon { + if actor.Kind() == KindAnon { http.Redirect(w, r, "/auth/login", http.StatusSeeOther) return } @@ -159,7 +158,7 @@ func authonly(next authedHandler) http.HandlerFunc { } } -func (h *webHandler) index(w http.ResponseWriter, r *http.Request, actor database.Actor) { +func (h *webHandler) index(w http.ResponseWriter, r *http.Request, actor Actor) { h.assets.renderPage(w, "dashboard", http.StatusOK, map[string]any{ "user": map[string]string{"email": actor.Email()}, "csrf": csrfToken(w, r), @@ -246,13 +245,13 @@ func (h *webHandler) login(w http.ResponseWriter, r *http.Request) { email := r.FormValue("email") password := r.FormValue("password") - userID, err := h.srv.db.UserVerifyPassword(r.Context(), database.SystemActor(), email, password, []byte(h.srv.cfg.Pepper)) + userID, err := h.srv.db.UserVerifyPassword(r.Context(), SystemActor(), email, password, []byte(h.srv.cfg.Pepper)) if err != nil { h.renderLogin(w, r, http.StatusUnauthorized, pb.ErrorReason_ERROR_REASON_INVALID_CREDENTIALS) return } - token, err := h.srv.db.UserCreateSession(r.Context(), database.SystemActor(), userID) + token, err := h.srv.db.UserSessionCreate(r.Context(), SystemActor(), userID) if err != nil { slog.Error("create session failed", "err", err) h.renderLogin(w, r, http.StatusInternalServerError, pb.ErrorReason_ERROR_REASON_INTERNAL) @@ -279,7 +278,7 @@ func (h *webHandler) logout(w http.ResponseWriter, r *http.Request) { return } if c, err := r.Cookie(sessionCookie); err == nil { - h.srv.db.UserDeleteSession(r.Context(), database.SystemActor(), c.Value) + h.srv.db.UserSessionDelete(r.Context(), SystemActor(), c.Value) } clearCookie(w, sessionCookie) clearCookie(w, csrfCookie) diff --git a/internal/database/actor.go b/internal/database/actor.go deleted file mode 100644 index 30825dc..0000000 --- a/internal/database/actor.go +++ /dev/null @@ -1,117 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package database - -import "github.com/google/uuid" - -// Actor is the principal making a database request. It carries identity, -// display metadata, and coarse capability. Zero value is invalid: dbID -// panics, so a missing initialisation cannot silently grant privileges. -// -// Synthetic actors (System/Operator/Anon) live only as Go constants — -// they are not rows in the users table, so they cannot be logged in as -// even if somebody writes a password into the DB. -type Actor struct { - kind actorKind - id uuid.UUID - email string - superuser bool -} - -type actorKind uint8 - -const ( - actorInvalid actorKind = iota - actorUser - actorOperator - actorSystem - actorAnon -) - -// ActorKind is the exported form of actorKind for audit sinks and logging. -type ActorKind uint8 - -const ( - KindInvalid ActorKind = iota - KindUser - KindOperator - KindSystem - KindAnon -) - -var ( - systemUUID = uuid.MustParse("00000000-0000-0000-0000-000000000001") - operatorUUID = uuid.MustParse("00000000-0000-0000-0000-000000000002") - anonUUID = uuid.MustParse("ffffffff-ffff-ffff-ffff-ffffffffffff") -) - -// UserActor identifies an authenticated user from a session or token. -func UserActor(id uuid.UUID, email string, superuser bool) Actor { - if id == uuid.Nil { - panic("database: UserActor with nil UUID") - } - if email == "" { - panic("database: UserActor with empty email") - } - return Actor{kind: actorUser, id: id, email: email, superuser: superuser} -} - -// OperatorActor is the principal for externally invoked privileged -// operations (admin socket). Distinguishable from System in audit logs. -func OperatorActor() Actor { - return Actor{kind: actorOperator, id: operatorUUID, email: "operator@mirum.local", superuser: true} -} - -// SystemActor is the principal for internal machinery (mTLS handshake, -// session bootstrap, background jobs). Not an operator action. -func SystemActor() Actor { - return Actor{kind: actorSystem, id: systemUUID, email: "system@mirum.local", superuser: true} -} - -// AnonActor is the principal for unauthenticated public requests. -func AnonActor() Actor { - return Actor{kind: actorAnon, id: anonUUID, email: "anonymous@mirum.local"} -} - -func (a Actor) Kind() ActorKind { - switch a.kind { - case actorUser: - return KindUser - case actorOperator: - return KindOperator - case actorSystem: - return KindSystem - case actorAnon: - return KindAnon - } - return KindInvalid -} - -func (a Actor) UserID() uuid.UUID { return a.id } -func (a Actor) Email() string { return a.email } -func (a Actor) IsSuperuser() bool { return a.superuser } - -// dbID returns the UUID to write into app.user_id. Panics on zero value. -func (a Actor) dbID() uuid.UUID { - if a.id == uuid.Nil { - panic("database: zero-value Actor; use UserActor/SystemActor/OperatorActor/AnonActor") - } - return a.id -} - -// kindString returns the string written into app.actor_kind. -// It must match the values tested by app_issuper() in the SQL migration. -func (a Actor) kindString() string { - switch a.kind { - case actorUser: - return "user" - case actorOperator: - return "operator" - case actorSystem: - return "system" - case actorAnon: - return "anon" - } - panic("database: zero-value Actor; use UserActor/SystemActor/OperatorActor/AnonActor") -} diff --git a/internal/database/database.go b/internal/database/database.go deleted file mode 100644 index be01a16..0000000 --- a/internal/database/database.go +++ /dev/null @@ -1,306 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package database - -import ( - "context" - "errors" - - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgxpool" - "github.com/jackc/tern/v2/migrate" -) - -var ( - ErrOpen = errors.New("database: failed to open") - ErrPing = errors.New("database: failed to ping") - ErrAcquire = errors.New("database: failed to acquire connection") - ErrMigrate = errors.New("database: failed to create migrator") - ErrFilterNotImplemented = errors.New("database: filter not implemented") -) - -// DB wraps a pgx connection pool. -type DB struct { - Pool *pgxpool.Pool -} - -// Open connects to PostgreSQL and returns a DB. -func Open(ctx context.Context, dsn string) (*DB, error) { - pool, err := pgxpool.New(ctx, dsn) - if err != nil { - return nil, errors.Join(ErrOpen, err) - } - - if err := pool.Ping(ctx); err != nil { - pool.Close() - return nil, errors.Join(ErrPing, err) - } - - return &DB{Pool: pool}, nil -} - -// Close closes the connection pool. -func (db *DB) Close() { - db.Pool.Close() -} - -// beginAs starts a transaction and sets the RLS actor. Both app.user_id -// and app.actor_kind are populated: app_issuper() checks actor_kind for -// System/Operator principals, and app.user_id for real user superusers. -func (db *DB) beginAs(ctx context.Context, actor Actor) (pgx.Tx, error) { - tx, err := db.Pool.Begin(ctx) - if err != nil { - return nil, err - } - if _, err := tx.Exec(ctx, - `SELECT set_config('app.user_id', $1, true), - set_config('app.actor_kind', $2, true)`, - actor.dbID().String(), actor.kindString(), - ); err != nil { - tx.Rollback(ctx) - return nil, err - } - return tx, nil -} - -// Migrate applies all pending migrations. -func (db *DB) Migrate(ctx context.Context) error { - conn, err := db.Pool.Acquire(ctx) - if err != nil { - return errors.Join(ErrAcquire, err) - } - defer conn.Release() - - // NOTE: RLS is active on all tables. Migrations with DML (INSERT/UPDATE/DELETE) - // must prefix the SQL with: - // SELECT set_config('app.actor_kind', 'system', true); - // This flips app_issuper() via the actor_kind branch for the migration's - // transaction only. No synthetic row in users is needed. - migrator, err := migrate.NewMigrator(ctx, conn.Conn(), "schema_version") - if err != nil { - return errors.Join(ErrMigrate, err) - } - - migrator.AppendMigration("create_users", ` - CREATE TABLE users ( - id UUID PRIMARY KEY DEFAULT uuidv7(), - email TEXT NOT NULL UNIQUE, - password TEXT NOT NULL, - superuser BOOLEAN NOT NULL DEFAULT false, - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), - deleted_at TIMESTAMPTZ - ); - - CREATE FUNCTION app_user_id() RETURNS uuid STABLE AS $$ - SELECT current_setting('app.user_id', true)::uuid; - $$ LANGUAGE sql; - - -- app_issuper has two independent branches: - -- (a) runtime setting app.actor_kind is 'system' or 'operator' — - -- set only by beginAs from Go for synthetic principals and by - -- DML migrations; cannot be injected via login since there is - -- no matching users row to authenticate against. - -- (b) the current app.user_id resolves to a users row with - -- superuser = true — real support-agent style superusers. - CREATE FUNCTION app_issuper() RETURNS boolean STABLE AS $$ - SELECT - current_setting('app.actor_kind', true) IN ('system', 'operator') - OR EXISTS ( - SELECT 1 FROM users - WHERE id = current_setting('app.user_id', true)::uuid - AND superuser = true - ); - $$ LANGUAGE sql; - `, ` - DROP FUNCTION app_issuper; - DROP FUNCTION app_user_id; - DROP TABLE users - `) - - migrator.AppendMigration("create_sessions", ` - CREATE TABLE sessions ( - token TEXT PRIMARY KEY, - user_id UUID NOT NULL REFERENCES users(id), - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), - expires_at TIMESTAMPTZ NOT NULL - ); - CREATE INDEX sessions_expires_at ON sessions (expires_at) - `, ` - DROP TABLE sessions - `) - - migrator.AppendMigration("create_organizations", ` - CREATE TABLE organizations ( - id UUID PRIMARY KEY DEFAULT uuidv7(), - name TEXT NOT NULL, - slug TEXT NOT NULL UNIQUE, - public BOOLEAN NOT NULL DEFAULT false, - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), - deleted_at TIMESTAMPTZ - ); - `, ` - DROP TABLE organizations - `) - - migrator.AppendMigration("create_org_members", ` - CREATE TABLE org_members ( - org_id UUID NOT NULL REFERENCES organizations(id), - user_id UUID NOT NULL REFERENCES users(id), - role TEXT NOT NULL DEFAULT 'member' CHECK (role IN ('owner', 'admin', 'member')), - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), - PRIMARY KEY (org_id, user_id) - ); - CREATE INDEX org_members_user_id ON org_members (user_id) - `, ` - DROP TABLE org_members - `) - - migrator.AppendMigration("create_workers", ` - CREATE TABLE workers ( - id UUID PRIMARY KEY DEFAULT uuidv7(), - org_id UUID REFERENCES organizations(id), - public_key BYTEA NOT NULL UNIQUE, - created_at TIMESTAMPTZ NOT NULL DEFAULT now(), - revoked_at TIMESTAMPTZ - ) - `, ` - DROP TABLE workers - `) - - migrator.AppendMigration("rls_helper_functions", ` - CREATE FUNCTION has_org_role(org uuid, roles text[]) RETURNS boolean STABLE AS $$ - SELECT EXISTS ( - SELECT 1 FROM org_members - WHERE org_id = org AND user_id = app_user_id() AND role = ANY(roles) - ); - $$ LANGUAGE sql - `, ` - DROP FUNCTION has_org_role - `) - - migrator.AppendMigration("rls_users", ` - ALTER TABLE users ENABLE ROW LEVEL SECURITY; - ALTER TABLE users FORCE ROW LEVEL SECURITY; - - CREATE POLICY superuser ON users FOR ALL USING (app_issuper()); - - -- SELECT is split into base case (own row) and extended case (users - -- who share an org membership with the caller). The base case is a - -- pure predicate — it lets the extended subquery resolve - -- app_user_id()'s own memberships without recursion. - CREATE POLICY user_select_self ON users FOR SELECT - USING (id = app_user_id()); - CREATE POLICY user_select_shared_org ON users FOR SELECT - USING (EXISTS ( - SELECT 1 FROM org_members target - JOIN org_members mine ON mine.org_id = target.org_id - WHERE target.user_id = users.id - AND mine.user_id = app_user_id() - )); - - CREATE POLICY user_update ON users FOR UPDATE USING (id = app_user_id()); - CREATE POLICY user_delete ON users FOR DELETE USING (id = app_user_id()); - CREATE POLICY user_insert ON users FOR INSERT WITH CHECK (app_issuper()); - `, ` - DROP POLICY superuser ON users; - DROP POLICY user_select_self ON users; - DROP POLICY user_select_shared_org ON users; - DROP POLICY user_insert ON users; - DROP POLICY user_update ON users; - DROP POLICY user_delete ON users; - - ALTER TABLE users DISABLE ROW LEVEL SECURITY; - `) - - migrator.AppendMigration("rls_sessions", ` - ALTER TABLE sessions ENABLE ROW LEVEL SECURITY; - ALTER TABLE sessions FORCE ROW LEVEL SECURITY; - - CREATE POLICY superuser ON sessions FOR ALL USING (app_issuper()); - CREATE POLICY own_sessions ON sessions FOR ALL USING (user_id = app_user_id()); - `, ` - DROP POLICY superuser ON sessions; - DROP POLICY own_sessions ON sessions; - - ALTER TABLE sessions DISABLE ROW LEVEL SECURITY; - `) - - migrator.AppendMigration("rls_organizations", ` - ALTER TABLE organizations ENABLE ROW LEVEL SECURITY; - ALTER TABLE organizations FORCE ROW LEVEL SECURITY; - - CREATE POLICY superuser ON organizations FOR ALL USING (app_issuper()); - CREATE POLICY select_org ON organizations FOR SELECT USING (public OR has_org_role(id, '{owner,admin,member}')); - CREATE POLICY update_org ON organizations FOR UPDATE USING (has_org_role(id, '{owner,admin}')); - CREATE POLICY delete_org ON organizations FOR DELETE USING (has_org_role(id, '{owner}')); - CREATE POLICY insert_org ON organizations FOR INSERT WITH CHECK (true); - `, ` - DROP POLICY superuser ON organizations; - DROP POLICY select_org ON organizations; - DROP POLICY update_org ON organizations; - DROP POLICY delete_org ON organizations; - DROP POLICY insert_org ON organizations; - - ALTER TABLE organizations DISABLE ROW LEVEL SECURITY; - `) - - migrator.AppendMigration("rls_org_members", ` - ALTER TABLE org_members ENABLE ROW LEVEL SECURITY; - ALTER TABLE org_members FORCE ROW LEVEL SECURITY; - - CREATE POLICY superuser ON org_members FOR ALL USING (app_issuper()); - - -- SELECT is split into two permissive policies. The base case - -- (own membership rows) is a pure column predicate with no function - -- call, which is enough for has_org_role's inner query to succeed — - -- has_org_role always filters by user_id = app_user_id(). Without - -- this split, the admin-scoped policy calls has_org_role which - -- queries org_members which calls has_org_role → stack overflow. - CREATE POLICY select_own_member ON org_members FOR SELECT - USING (user_id = app_user_id()); - CREATE POLICY select_org_member ON org_members FOR SELECT - USING (has_org_role(org_id, '{owner,admin,member}')); - - CREATE POLICY insert_member ON org_members FOR INSERT WITH CHECK ( - has_org_role(org_id, '{owner,admin}') - OR ( - user_id = app_user_id() - AND role = 'owner' - AND NOT EXISTS (SELECT 1 FROM org_members existing WHERE existing.org_id = org_members.org_id) - ) - ); - CREATE POLICY update_member ON org_members FOR UPDATE USING (has_org_role(org_id, '{owner,admin}')); - CREATE POLICY delete_member ON org_members FOR DELETE USING (has_org_role(org_id, '{owner,admin}')); - `, ` - DROP POLICY superuser ON org_members; - DROP POLICY select_own_member ON org_members; - DROP POLICY select_org_member ON org_members; - DROP POLICY insert_member ON org_members; - DROP POLICY update_member ON org_members; - DROP POLICY delete_member ON org_members; - - ALTER TABLE org_members DISABLE ROW LEVEL SECURITY; - `) - - migrator.AppendMigration("rls_workers", ` - ALTER TABLE workers ENABLE ROW LEVEL SECURITY; - ALTER TABLE workers FORCE ROW LEVEL SECURITY; - - CREATE POLICY superuser ON workers FOR ALL USING (app_issuper()); - CREATE POLICY select_worker ON workers FOR SELECT USING (org_id IS NOT NULL AND has_org_role(org_id, '{owner,admin,member}')); - CREATE POLICY insert_worker ON workers FOR INSERT WITH CHECK (org_id IS NULL OR has_org_role(org_id, '{owner,admin}')); - CREATE POLICY update_worker ON workers FOR UPDATE USING (org_id IS NOT NULL AND has_org_role(org_id, '{owner,admin}')); - CREATE POLICY delete_worker ON workers FOR DELETE USING (org_id IS NOT NULL AND has_org_role(org_id, '{owner,admin}')); - `, ` - DROP POLICY superuser ON workers; - DROP POLICY select_worker ON workers; - DROP POLICY insert_worker ON workers; - DROP POLICY update_worker ON workers; - DROP POLICY delete_worker ON workers; - - ALTER TABLE workers DISABLE ROW LEVEL SECURITY; - `) - - return migrator.Migrate(ctx) -} diff --git a/internal/database/organization.go b/internal/database/organization.go deleted file mode 100644 index e5e8a3b..0000000 --- a/internal/database/organization.go +++ /dev/null @@ -1,477 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package database - -import ( - "context" - "errors" - "time" - - "github.com/google/uuid" - sb "github.com/huandu/go-sqlbuilder" - "github.com/jackc/pgerrcode" - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgconn" -) - -var ( - ErrOrgNotFound = errors.New("database: organization not found") - ErrSlugTaken = errors.New("database: slug already taken") - ErrAlreadyMember = errors.New("database: already a member") - ErrNotMember = errors.New("database: not a member") - ErrLastOwner = errors.New("database: last owner") -) - -// OrgRef identifies an organization by ID or slug. -type OrgRef struct { - id uuid.UUID - slug string -} - -func OrgByID(id uuid.UUID) OrgRef { return OrgRef{id: id} } -func OrgBySlug(slug string) OrgRef { return OrgRef{slug: slug} } - -func (r OrgRef) where() (string, any) { - if r.id != uuid.Nil { - return "id", r.id - } - return "slug", r.slug -} - -// Organization holds info about an organization. -type Organization struct { - ID uuid.UUID - Name string - Slug string - Public bool - CreatedAt time.Time -} - -// OrgMember pairs a user with their role in an organization. -type OrgMember struct { - User User - Role string - JoinedAt time.Time -} - -// resolveOrg locks and returns the org ID within a transaction. -func resolveOrg(ctx context.Context, tx pgx.Tx, ref OrgRef) (uuid.UUID, error) { - col, val := ref.where() - q := sb.PostgreSQL.NewSelectBuilder() - - sql, args := q.Select("id").From("organizations"). - Where(q.Equal(col, val), q.IsNull("deleted_at")). - ForUpdate(). - Build() - - var id uuid.UUID - if err := tx.QueryRow(ctx, sql, args...).Scan(&id); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return uuid.Nil, ErrOrgNotFound - } - - return uuid.Nil, err - } - - return id, nil -} - -// GetOrg returns an org by ref (ID or slug). -func (db *DB) GetOrg(ctx context.Context, actor Actor, ref OrgRef) (*Organization, error) { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return nil, err - } - defer tx.Rollback(ctx) - - col, val := ref.where() - q := sb.PostgreSQL.NewSelectBuilder() - - sql, args := q.Select("id", "name", "slug", "public", "created_at"). - From("organizations"). - Where(q.Equal(col, val), q.IsNull("deleted_at")). - Build() - - var o Organization - if err := tx.QueryRow(ctx, sql, args...).Scan(&o.ID, &o.Name, &o.Slug, &o.Public, &o.CreatedAt); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return nil, ErrOrgNotFound - } - return nil, err - } - - return &o, nil -} - -// CreateOrganization creates an org and adds the owner as the first member. -func (db *DB) CreateOrganization(ctx context.Context, actor Actor, name, slug string, public bool, owner UserRef) (uuid.UUID, error) { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return uuid.Nil, err - } - defer tx.Rollback(ctx) - - userID, err := resolveUser(ctx, tx, owner) - if err != nil { - return uuid.Nil, err - } - - var orgID uuid.UUID - if err := tx.QueryRow(ctx, - `INSERT INTO organizations (name, slug, public) VALUES ($1, $2, $3) RETURNING id`, - name, slug, public, - ).Scan(&orgID); err != nil { - var pgErr *pgconn.PgError - if errors.As(err, &pgErr) && pgErr.Code == pgerrcode.UniqueViolation { - return uuid.Nil, ErrSlugTaken - } - - return uuid.Nil, err - } - - if _, err := tx.Exec(ctx, - `INSERT INTO org_members (org_id, user_id, role) VALUES ($1, $2, 'owner')`, - orgID, userID, - ); err != nil { - return uuid.Nil, err - } - - return orgID, tx.Commit(ctx) -} - -// UpdateOrganization updates an org's name, slug, and/or public flag. -func (db *DB) UpdateOrganization(ctx context.Context, actor Actor, ref OrgRef, name *string, slug *string, public *bool) error { - if name == nil && slug == nil && public == nil { - return nil - } - - tx, err := db.beginAs(ctx, actor) - if err != nil { - return err - } - defer tx.Rollback(ctx) - - id, err := resolveOrg(ctx, tx, ref) - if err != nil { - return err - } - - ub := sb.PostgreSQL.NewUpdateBuilder() - ub.Update("organizations") - if name != nil { - ub.SetMore(ub.Assign("name", *name)) - } - if slug != nil { - ub.SetMore(ub.Assign("slug", *slug)) - } - if public != nil { - ub.SetMore(ub.Assign("public", *public)) - } - ub.Where(ub.Equal("id", id)) - - sql, args := ub.Build() - if _, err := tx.Exec(ctx, sql, args...); err != nil { - var pgErr *pgconn.PgError - if errors.As(err, &pgErr) && pgErr.Code == pgerrcode.UniqueViolation { - return ErrSlugTaken - } - - return err - } - - return tx.Commit(ctx) -} - -// DeleteOrganization soft-deletes an org and removes all members. -func (db *DB) DeleteOrganization(ctx context.Context, actor Actor, ref OrgRef) error { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return err - } - defer tx.Rollback(ctx) - - id, err := resolveOrg(ctx, tx, ref) - if err != nil { - return err - } - - if _, err := tx.Exec(ctx, `DELETE FROM org_members WHERE org_id = $1`, id); err != nil { - return err - } - - if _, err := tx.Exec(ctx, - `UPDATE organizations SET slug = id::text, deleted_at = now() WHERE id = $1`, id, - ); err != nil { - return err - } - - return tx.Commit(ctx) -} - -// ListOrganizations returns a page of orgs and the total count. -func (db *DB) ListOrganizations(ctx context.Context, actor Actor, cursor uuid.UUID, limit int, filter string) ([]Organization, int, error) { - if filter != "" { - return nil, 0, ErrFilterNotImplemented - } - - tx, err := db.beginAs(ctx, actor) - if err != nil { - return nil, 0, err - } - defer tx.Rollback(ctx) - - var total int - if err := tx.QueryRow(ctx, - `SELECT count(*) FROM organizations WHERE deleted_at IS NULL`, - ).Scan(&total); err != nil { - return nil, 0, err - } - - q := sb.PostgreSQL.NewSelectBuilder() - q.Select("id", "name", "slug", "public", "created_at"). - From("organizations"). - Where(q.IsNull("deleted_at")). - OrderBy("id"). - Limit(limit) - if cursor != uuid.Nil { - q.Where(q.GreaterThan("id", cursor)) - } - - sql, args := q.Build() - rows, err := tx.Query(ctx, sql, args...) - if err != nil { - return nil, 0, err - } - defer rows.Close() - - var orgs []Organization - for rows.Next() { - var o Organization - if err := rows.Scan(&o.ID, &o.Name, &o.Slug, &o.Public, &o.CreatedAt); err != nil { - return nil, 0, err - } - orgs = append(orgs, o) - } - return orgs, total, rows.Err() -} - -// GetOrgMember returns a single member's info. -func (db *DB) GetOrgMember(ctx context.Context, actor Actor, org OrgRef, user UserRef) (*OrgMember, error) { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return nil, err - } - defer tx.Rollback(ctx) - - orgID, err := resolveOrg(ctx, tx, org) - if err != nil { - return nil, err - } - userID, err := resolveUser(ctx, tx, user) - if err != nil { - return nil, err - } - - var m OrgMember - if err := tx.QueryRow(ctx, - `SELECT u.id, u.email, u.created_at, om.role, om.created_at - FROM org_members om - JOIN users u ON u.id = om.user_id - WHERE om.org_id = $1 AND om.user_id = $2`, orgID, userID, - ).Scan(&m.User.ID, &m.User.Email, &m.User.CreatedAt, &m.Role, &m.JoinedAt); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return nil, ErrNotMember - } - return nil, err - } - return &m, nil -} - -// ListOrgMembers returns a page of members for an org. -func (db *DB) ListOrgMembers(ctx context.Context, actor Actor, org OrgRef, cursor uuid.UUID, limit int, filter string) ([]OrgMember, int, error) { - if filter != "" { - return nil, 0, ErrFilterNotImplemented - } - - tx, err := db.beginAs(ctx, actor) - if err != nil { - return nil, 0, err - } - defer tx.Rollback(ctx) - - orgID, err := resolveOrg(ctx, tx, org) - if err != nil { - return nil, 0, err - } - - var total int - if err := tx.QueryRow(ctx, - `SELECT count(*) FROM org_members WHERE org_id = $1`, orgID, - ).Scan(&total); err != nil { - return nil, 0, err - } - - q := sb.PostgreSQL.NewSelectBuilder() - q.Select("u.id", "u.email", "u.created_at", "m.role", "m.created_at"). - From("org_members m"). - Join("users u", "u.id = m.user_id"). - Where(q.Equal("m.org_id", orgID), q.IsNull("u.deleted_at")). - OrderBy("u.id"). - Limit(limit) - if cursor != uuid.Nil { - q.Where(q.GreaterThan("u.id", cursor)) - } - - sql, args := q.Build() - rows, err := tx.Query(ctx, sql, args...) - if err != nil { - return nil, 0, err - } - defer rows.Close() - - var members []OrgMember - for rows.Next() { - var m OrgMember - if err := rows.Scan(&m.User.ID, &m.User.Email, &m.User.CreatedAt, &m.Role, &m.JoinedAt); err != nil { - return nil, 0, err - } - - members = append(members, m) - } - - return members, total, rows.Err() -} - -// AddOrgMember adds a user to an org with the given role. -func (db *DB) AddOrgMember(ctx context.Context, actor Actor, org OrgRef, user UserRef, role string) error { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return err - } - defer tx.Rollback(ctx) - - orgID, err := resolveOrg(ctx, tx, org) - if err != nil { - return err - } - userID, err := resolveUser(ctx, tx, user) - if err != nil { - return err - } - - if _, err := tx.Exec(ctx, - `INSERT INTO org_members (org_id, user_id, role) VALUES ($1, $2, $3)`, - orgID, userID, role, - ); err != nil { - var pgErr *pgconn.PgError - if errors.As(err, &pgErr) && pgErr.Code == pgerrcode.UniqueViolation { - return ErrAlreadyMember - } - return err - } - - return tx.Commit(ctx) -} - -// UpdateOrgMemberRole changes a member's role. Fails if demoting the last owner. -func (db *DB) UpdateOrgMemberRole(ctx context.Context, actor Actor, org OrgRef, user UserRef, newRole string) error { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return err - } - defer tx.Rollback(ctx) - - orgID, err := resolveOrg(ctx, tx, org) - if err != nil { - return err - } - userID, err := resolveUser(ctx, tx, user) - if err != nil { - return err - } - - var currentRole string - if err := tx.QueryRow(ctx, - `SELECT role FROM org_members WHERE org_id = $1 AND user_id = $2 FOR UPDATE`, - orgID, userID, - ).Scan(¤tRole); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return ErrNotMember - } - return err - } - - if currentRole == "owner" && newRole != "owner" { - var ownerCount int - if err := tx.QueryRow(ctx, - `SELECT count(*) FROM org_members WHERE org_id = $1 AND role = 'owner'`, - orgID, - ).Scan(&ownerCount); err != nil { - return err - } - if ownerCount <= 1 { - return ErrLastOwner - } - } - - if _, err := tx.Exec(ctx, - `UPDATE org_members SET role = $1 WHERE org_id = $2 AND user_id = $3`, - newRole, orgID, userID, - ); err != nil { - return err - } - - return tx.Commit(ctx) -} - -// RemoveOrgMember removes a user from an org. Fails if they are the last owner. -func (db *DB) RemoveOrgMember(ctx context.Context, actor Actor, org OrgRef, user UserRef) error { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return err - } - defer tx.Rollback(ctx) - - orgID, err := resolveOrg(ctx, tx, org) - if err != nil { - return err - } - userID, err := resolveUser(ctx, tx, user) - if err != nil { - return err - } - - var role string - if err := tx.QueryRow(ctx, - `SELECT role FROM org_members WHERE org_id = $1 AND user_id = $2 FOR UPDATE`, - orgID, userID, - ).Scan(&role); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return ErrNotMember - } - return err - } - - if role == "owner" { - var ownerCount int - if err := tx.QueryRow(ctx, - `SELECT count(*) FROM org_members WHERE org_id = $1 AND role = 'owner'`, - orgID, - ).Scan(&ownerCount); err != nil { - return err - } - if ownerCount <= 1 { - return ErrLastOwner - } - } - - if _, err := tx.Exec(ctx, - `DELETE FROM org_members WHERE org_id = $1 AND user_id = $2`, - orgID, userID, - ); err != nil { - return err - } - - return tx.Commit(ctx) -} diff --git a/internal/database/user.go b/internal/database/user.go deleted file mode 100644 index a6d250c..0000000 --- a/internal/database/user.go +++ /dev/null @@ -1,492 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package database - -import ( - "context" - "crypto/hmac" - "crypto/rand" - "crypto/sha256" - "crypto/subtle" - "encoding/base64" - "errors" - "fmt" - "strings" - "time" - - "github.com/google/uuid" - sb "github.com/huandu/go-sqlbuilder" - "github.com/jackc/pgerrcode" - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgconn" - "golang.org/x/crypto/argon2" - - "dimidiumlabs/mirum/internal/config" -) - -var ( - ErrSoleOwner = errors.New("database: sole owner of an organization") - ErrEmailTaken = errors.New("database: email already taken") - ErrInvalidCreds = errors.New("database: invalid credentials") - ErrUserNotFound = errors.New("database: user not found") - ErrReservedEmail = errors.New("database: email uses a reserved domain") -) - -// reservedEmailSuffix is the domain carved out for synthetic actors -// (system/operator/anon). Real users cannot register with this suffix. -const reservedEmailSuffix = "@mirum.local" - -// UserRef identifies a user by ID or email. -type UserRef struct { - id uuid.UUID - email string -} - -func UserByID(id uuid.UUID) UserRef { return UserRef{id: id} } -func UserByEmail(email string) UserRef { return UserRef{email: email} } - -func (r UserRef) where() (string, any) { - if r.id != uuid.Nil { - return "id", r.id - } - return "email", r.email -} - -const ( - saltLen = 16 - - argonTime = 3 - argonMemory = 64 * 1024 // 64 MB - argonKeyLen = 32 - argonThreads = 2 -) - -// User holds info about a user. -type User struct { - ID uuid.UUID - Email string - CreatedAt time.Time -} - -// hashToken returns the hex-encoded SHA-256 of a session token. -func hashToken(token string) string { - h := sha256.Sum256([]byte(token)) - return fmt.Sprintf("%x", h) -} - -// verifyHash parses a PHC-format argon2id string and compares. -// Format: $argon2id$v=19$m=65536,t=3,p=2$$ -func verifyHash(password, encoded string, pepper []byte) bool { - // $argon2id$v=19$m=65536,t=3,p=2$salt$key → 6 parts - parts := strings.Split(encoded, "$") - if len(parts) != 6 || parts[1] != "argon2id" { - return false - } - - var memory, time uint32 - var threads uint8 - if _, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &time, &threads); err != nil { - return false - } - - salt, err := base64.RawStdEncoding.DecodeString(parts[4]) - if err != nil { - return false - } - expectedKey, err := base64.RawStdEncoding.DecodeString(parts[5]) - if err != nil { - return false - } - - mac := hmac.New(sha256.New, pepper) - mac.Write([]byte(password)) - peppered := mac.Sum(nil) - - key := argon2.IDKey(peppered, salt, time, memory, threads, uint32(len(expectedKey))) - - return subtle.ConstantTimeCompare(key, expectedKey) == 1 -} - -// hashPassword produces a PHC-format string: -// $argon2id$v=19$m=65536,t=3,p=2$$ -func hashPassword(password string, pepper []byte) (string, error) { - salt := make([]byte, saltLen) - if _, err := rand.Read(salt); err != nil { - return "", err - } - - // Apply pepper: HMAC-SHA256(pepper, password) - mac := hmac.New(sha256.New, pepper) - mac.Write([]byte(password)) - peppered := mac.Sum(nil) - - key := argon2.IDKey(peppered, salt, argonTime, argonMemory, argonThreads, argonKeyLen) - - return fmt.Sprintf("$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s", - argon2.Version, - argonMemory, argonTime, argonThreads, - base64.RawStdEncoding.EncodeToString(salt), - base64.RawStdEncoding.EncodeToString(key), - ), nil -} - -// UserCreate hashes the password with argon2id and inserts a new user. -// The pepper is a server-side secret not stored in the database. -func (db *DB) UserCreate(ctx context.Context, actor Actor, email, password string, pepper []byte) (uuid.UUID, error) { - if strings.HasSuffix(strings.ToLower(email), reservedEmailSuffix) { - return uuid.Nil, ErrReservedEmail - } - - hash, err := hashPassword(password, pepper) - if err != nil { - return uuid.Nil, err - } - - tx, err := db.beginAs(ctx, actor) - if err != nil { - return uuid.Nil, err - } - defer tx.Rollback(ctx) - - var id uuid.UUID - if err := tx.QueryRow(ctx, - `INSERT INTO users (email, password) VALUES ($1, $2) RETURNING id`, - email, hash, - ).Scan(&id); err != nil { - var pgErr *pgconn.PgError - if errors.As(err, &pgErr) && pgErr.Code == pgerrcode.UniqueViolation { - return uuid.Nil, ErrEmailTaken - } - return uuid.Nil, err - } - - return id, tx.Commit(ctx) -} - -// GetUser returns a user by ref (ID or email). -func (db *DB) GetUser(ctx context.Context, actor Actor, ref UserRef) (*User, error) { - col, val := ref.where() - - tx, err := db.beginAs(ctx, actor) - if err != nil { - return nil, err - } - defer tx.Rollback(ctx) - - q := sb.PostgreSQL.NewSelectBuilder() - sql, args := q.Select("id", "email", "created_at"). - From("users"). - Where(q.Equal(col, val), q.IsNull("deleted_at")). - Build() - - var u User - if err := tx.QueryRow(ctx, sql, args...).Scan(&u.ID, &u.Email, &u.CreatedAt); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return nil, ErrUserNotFound - } - return nil, err - } - - return &u, nil -} - -// ListUsers returns a page of users and the total count. -func (db *DB) ListUsers(ctx context.Context, actor Actor, cursor uuid.UUID, limit int, filter string) ([]User, int, error) { - if filter != "" { - return nil, 0, ErrFilterNotImplemented - } - - tx, err := db.beginAs(ctx, actor) - if err != nil { - return nil, 0, err - } - defer tx.Rollback(ctx) - - var total int - if err := tx.QueryRow(ctx, - `SELECT count(*) FROM users WHERE deleted_at IS NULL`, - ).Scan(&total); err != nil { - return nil, 0, err - } - - q := sb.PostgreSQL.NewSelectBuilder() - q.Select("id", "email", "created_at"). - From("users"). - Where(q.IsNull("deleted_at")). - OrderBy("id"). - Limit(limit) - if cursor != uuid.Nil { - q.Where(q.GreaterThan("id", cursor)) - } - - sql, args := q.Build() - rows, err := tx.Query(ctx, sql, args...) - if err != nil { - return nil, 0, err - } - defer rows.Close() - - var users []User - for rows.Next() { - var u User - if err := rows.Scan(&u.ID, &u.Email, &u.CreatedAt); err != nil { - return nil, 0, err - } - users = append(users, u) - } - - return users, total, rows.Err() -} - -// checkNotSoleOwner returns ErrSoleOwner if the user is the only owner of any org. -func checkNotSoleOwner(ctx context.Context, tx pgx.Tx, userID uuid.UUID) error { - var slug string - err := tx.QueryRow(ctx, - `SELECT o.slug FROM org_members m - JOIN organizations o ON o.id = m.org_id - WHERE m.role = 'owner' AND o.deleted_at IS NULL - GROUP BY o.id, o.slug - HAVING count(*) = 1 AND bool_or(m.user_id = $1) - LIMIT 1`, userID, - ).Scan(&slug) - if err == nil { - return ErrSoleOwner - } - if errors.Is(err, pgx.ErrNoRows) { - return nil - } - return err -} - -// resolveUser locks and returns the user ID within a transaction. -func resolveUser(ctx context.Context, tx pgx.Tx, ref UserRef) (uuid.UUID, error) { - col, val := ref.where() - q := sb.PostgreSQL.NewSelectBuilder() - - sql, args := q.Select("id").From("users"). - Where(q.Equal(col, val), q.IsNull("deleted_at")). - ForUpdate(). - Build() - - var id uuid.UUID - if err := tx.QueryRow(ctx, sql, args...).Scan(&id); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return uuid.Nil, ErrUserNotFound - } - - return uuid.Nil, err - } - - return id, nil -} - -// UserUpdate updates a user's email and/or password. -// Invalidates all sessions when password changes. -func (db *DB) UserUpdate(ctx context.Context, actor Actor, ref UserRef, email *string, password *string, pepper []byte) error { - if email == nil && password == nil { - return nil - } - if email != nil && strings.HasSuffix(strings.ToLower(*email), reservedEmailSuffix) { - return ErrReservedEmail - } - - tx, err := db.beginAs(ctx, actor) - if err != nil { - return err - } - defer tx.Rollback(ctx) - - id, err := resolveUser(ctx, tx, ref) - if err != nil { - return err - } - - ub := sb.PostgreSQL.NewUpdateBuilder() - ub.Update("users") - - if email != nil { - ub.SetMore(ub.Assign("email", *email)) - } - - if password != nil { - hash, err := hashPassword(*password, pepper) - if err != nil { - return err - } - - ub.SetMore(ub.Assign("password", hash)) - } - - ub.Where(ub.Equal("id", id)) - - sql, args := ub.Build() - if _, err := tx.Exec(ctx, sql, args...); err != nil { - var pgErr *pgconn.PgError - if errors.As(err, &pgErr) && pgErr.Code == pgerrcode.UniqueViolation { - return ErrEmailTaken - } - return err - } - - if password != nil { - if _, err := tx.Exec(ctx, `DELETE FROM sessions WHERE user_id = $1`, id); err != nil { - return err - } - } - - return tx.Commit(ctx) -} - -// UserDelete soft-deletes a user. -// Fails if the user is the sole owner of any organization. -func (db *DB) UserDelete(ctx context.Context, actor Actor, ref UserRef) error { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return err - } - defer tx.Rollback(ctx) - - id, err := resolveUser(ctx, tx, ref) - if err != nil { - return err - } - - if err := checkNotSoleOwner(ctx, tx, id); err != nil { - return err - } - - if _, err = tx.Exec(ctx, `DELETE FROM sessions WHERE user_id = $1`, id); err != nil { - return err - } - - if _, err = tx.Exec(ctx, `DELETE FROM org_members WHERE user_id = $1`, id); err != nil { - return err - } - - if _, err := tx.Exec(ctx, - `UPDATE users SET email = id::text, password = '', deleted_at = now() WHERE id = $1`, id, - ); err != nil { - return err - } - - return tx.Commit(ctx) -} - -// UserVerifyPassword checks credentials and returns the user ID. -func (db *DB) UserVerifyPassword(ctx context.Context, actor Actor, email, password string, pepper []byte) (uuid.UUID, error) { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return uuid.Nil, err - } - defer tx.Rollback(ctx) - - var id uuid.UUID - var hash string - - if err := tx.QueryRow(ctx, - `SELECT id, password FROM users WHERE email = $1 AND deleted_at IS NULL`, - email, - ).Scan(&id, &hash); err != nil { - return uuid.Nil, ErrInvalidCreds - } - - if !verifyHash(password, hash, pepper) { - return uuid.Nil, ErrInvalidCreds - } - - return id, nil -} - -// UserGetSession resolves a session token into the Actor it authenticates. -// Returns an invalid zero Actor on any error; callers must check err. -func (db *DB) UserGetSession(ctx context.Context, actor Actor, token string) (Actor, error) { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return Actor{}, err - } - defer tx.Rollback(ctx) - - var ( - userID uuid.UUID - email string - superuser bool - expiresAt time.Time - ) - h := hashToken(token) - - if err := tx.QueryRow(ctx, - `SELECT s.user_id, u.email, u.superuser, s.expires_at - FROM sessions s JOIN users u ON u.id = s.user_id - WHERE s.token = $1 AND s.expires_at > now() AND u.deleted_at IS NULL`, - h, - ).Scan(&userID, &email, &superuser, &expiresAt); err != nil { - return Actor{}, err - } - - if time.Until(expiresAt) < config.SessionTTL/2 { - if _, err := tx.Exec(ctx, - `UPDATE sessions SET expires_at = now() + $2 WHERE token = $1`, - h, config.SessionTTL, - ); err != nil { - return Actor{}, err - } - } - - if err := tx.Commit(ctx); err != nil { - return Actor{}, err - } - return UserActor(userID, email, superuser), nil -} - -// UserCreateSession generates a random token, stores its hash, and returns the token. -func (db *DB) UserCreateSession(ctx context.Context, actor Actor, userID uuid.UUID) (string, error) { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return "", err - } - defer tx.Rollback(ctx) - - buf := make([]byte, 32) - if _, err := rand.Read(buf); err != nil { - return "", err - } - token := base64.RawURLEncoding.EncodeToString(buf) - - if _, err := tx.Exec(ctx, - `INSERT INTO sessions (token, user_id, expires_at) VALUES ($1, $2, now() + $3)`, - hashToken(token), userID, config.SessionTTL, - ); err != nil { - return "", err - } - - return token, tx.Commit(ctx) -} - -// UserDeleteSession removes a session (logout). -func (db *DB) UserDeleteSession(ctx context.Context, actor Actor, token string) error { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return err - } - defer tx.Rollback(ctx) - - if _, err := tx.Exec(ctx, `DELETE FROM sessions WHERE token = $1`, hashToken(token)); err != nil { - return err - } - return tx.Commit(ctx) -} - -// PurgeExpiredSessions deletes all expired sessions. Runs as SystemActor. -func (db *DB) PurgeExpiredSessions(ctx context.Context) error { - tx, err := db.beginAs(ctx, SystemActor()) - if err != nil { - return err - } - defer tx.Rollback(ctx) - - if _, err := tx.Exec(ctx, `DELETE FROM sessions WHERE expires_at < now()`); err != nil { - return err - } - return tx.Commit(ctx) -} diff --git a/internal/database/validate.go b/internal/database/validate.go deleted file mode 100644 index 9dc2b44..0000000 --- a/internal/database/validate.go +++ /dev/null @@ -1,60 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package database - -import ( - "errors" - "net/mail" - "regexp" - "strings" -) - -var slugRe = regexp.MustCompile(`^[a-zA-Z0-9]+(?:-[a-zA-Z0-9]+)*$`) - -var validRoles = map[string]bool{ - "owner": true, - "admin": true, - "member": true, -} - -var ( - ErrInvalidEmail = errors.New("database: invalid email") - ErrInvalidSlug = errors.New("database: invalid slug") - ErrInvalidRole = errors.New("database: invalid role") -) - -// ValidateEmail checks that the value is a valid email address. -func ValidateEmail(value string) error { - if _, err := mail.ParseAddress(value); err != nil { - return ErrInvalidEmail - } - return nil -} - -// ValidateSlug checks format and returns the normalized (lowercased) slug. -func ValidateSlug(value string) (string, error) { - if len(value) < 2 || len(value) > 64 || !slugRe.MatchString(value) { - return "", ErrInvalidSlug - } - return strings.ToLower(value), nil -} - -// ValidateRole checks that the value is a valid role string. -func ValidateRole(value string) error { - if !validRoles[value] { - return ErrInvalidRole - } - return nil -} - -// ClampPageSize clamps a page_size to [1, max]. If v is 0 (unset), returns defaultSize. -func ClampPageSize(v int32, defaultSize, max int32) int32 { - if v <= 0 { - return defaultSize - } - if v > max { - return max - } - return v -} diff --git a/internal/database/worker.go b/internal/database/worker.go deleted file mode 100644 index 3680901..0000000 --- a/internal/database/worker.go +++ /dev/null @@ -1,165 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package database - -import ( - "context" - "errors" - "time" - - "github.com/google/uuid" - sb "github.com/huandu/go-sqlbuilder" - "github.com/jackc/pgx/v5" -) - -var ( - ErrWorkerNotFound = errors.New("database: worker not found") -) - -// Worker holds info about a registered worker. -type Worker struct { - ID uuid.UUID - OrgID *uuid.UUID - PublicKey []byte - CreatedAt time.Time -} - -// GetWorker returns a worker by ID. -func (db *DB) GetWorker(ctx context.Context, actor Actor, id uuid.UUID) (*Worker, error) { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return nil, err - } - defer tx.Rollback(ctx) - - var w Worker - if err := tx.QueryRow(ctx, - `SELECT id, public_key, org_id, created_at FROM workers WHERE id = $1 AND revoked_at IS NULL`, - id, - ).Scan(&w.ID, &w.PublicKey, &w.OrgID, &w.CreatedAt); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return nil, ErrWorkerNotFound - } - return nil, err - } - - return &w, nil -} - -// CreateWorker registers a new worker with the given public key and optional org. -func (db *DB) CreateWorker(ctx context.Context, actor Actor, publicKey []byte, org *OrgRef) (uuid.UUID, error) { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return uuid.Nil, err - } - defer tx.Rollback(ctx) - - var orgID *uuid.UUID - if org != nil { - id, err := resolveOrg(ctx, tx, *org) - if err != nil { - return uuid.Nil, err - } - orgID = &id - } - - var workerID uuid.UUID - if err := tx.QueryRow(ctx, - `INSERT INTO workers (public_key, org_id) VALUES ($1, $2) RETURNING id`, - publicKey, orgID, - ).Scan(&workerID); err != nil { - return uuid.Nil, err - } - - return workerID, tx.Commit(ctx) -} - -// DeleteWorker soft-deletes a worker by ID. -func (db *DB) DeleteWorker(ctx context.Context, actor Actor, id uuid.UUID) error { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return err - } - defer tx.Rollback(ctx) - - tag, err := tx.Exec(ctx, - `UPDATE workers SET revoked_at = now() WHERE id = $1 AND revoked_at IS NULL`, id, - ) - if err != nil { - return err - } - if tag.RowsAffected() == 0 { - return ErrWorkerNotFound - } - - return tx.Commit(ctx) -} - -// ListWorkers returns a page of workers and the total count. -func (db *DB) ListWorkers(ctx context.Context, actor Actor, cursor uuid.UUID, limit int, filter string) ([]Worker, int, error) { - if filter != "" { - return nil, 0, ErrFilterNotImplemented - } - - tx, err := db.beginAs(ctx, actor) - if err != nil { - return nil, 0, err - } - defer tx.Rollback(ctx) - - var total int - if err := tx.QueryRow(ctx, - `SELECT count(*) FROM workers WHERE revoked_at IS NULL`, - ).Scan(&total); err != nil { - return nil, 0, err - } - - q := sb.PostgreSQL.NewSelectBuilder() - q.Select("id", "public_key", "org_id", "created_at"). - From("workers"). - Where(q.IsNull("revoked_at")). - OrderBy("id"). - Limit(limit) - if cursor != uuid.Nil { - q.Where(q.GreaterThan("id", cursor)) - } - - sql, args := q.Build() - rows, err := tx.Query(ctx, sql, args...) - if err != nil { - return nil, 0, err - } - defer rows.Close() - - var workers []Worker - for rows.Next() { - var w Worker - if err := rows.Scan(&w.ID, &w.PublicKey, &w.OrgID, &w.CreatedAt); err != nil { - return nil, 0, err - } - workers = append(workers, w) - } - return workers, total, rows.Err() -} - -// LookupWorker finds an active worker by its ed25519 public key. -func (db *DB) LookupWorker(ctx context.Context, actor Actor, publicKey []byte) (*Worker, error) { - tx, err := db.beginAs(ctx, actor) - if err != nil { - return nil, err - } - defer tx.Rollback(ctx) - - var w Worker - if err := tx.QueryRow(ctx, - `SELECT id, public_key, org_id, created_at FROM workers WHERE public_key = $1 AND revoked_at IS NULL`, - publicKey, - ).Scan(&w.ID, &w.PublicKey, &w.OrgID, &w.CreatedAt); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return nil, ErrWorkerNotFound - } - return nil, err - } - return &w, nil -} -- Gilti