aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
Diffstat
-rw-r--r--cmd/mirumd/api_auth.go233+233 −0
-rw-r--r--cmd/mirumd/config.go2+2 −0
-rw-r--r--cmd/mirumd/main.go21+19 −2
-rw-r--r--cmd/mirumd/server_admin.go70+37 −33
-rw-r--r--cmd/mirumd/server_grpc.go3+2 −1
-rw-r--r--cmd/mirumd/server_web.go315+224 −91
-rw-r--r--go.mod4+4 −0
-rw-r--r--go.sum10+10 −0
-rw-r--r--internal/database/database.go199+174 −25
-rw-r--r--internal/database/organization.go47+26 −21
-rw-r--r--internal/database/user.go99+72 −27
-rw-r--r--internal/database/worker.go42+29 −13
-rw-r--r--proto/admin.proto17+17 −0
13 files changed, 849 insertions, 213 deletions
diff --git a/cmd/mirumd/api_auth.go b/cmd/mirumd/api_auth.go
new file mode 100644
--- /dev/null
+++ b/cmd/mirumd/api_auth.go
@@ -0,0 +1,233 @@
+// 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
+ }
+
+ caller := CallerFromContext(ctx)
+ if caller == nil {
+ return connect.NewError(connect.CodeUnauthenticated, nil)
+ }
+ if caller.Superuser {
+ return nil
+ }
+
+ switch procedure {
+ // User — self or superuser
+ case pbconnect.AdminUserGetProcedure,
+ pbconnect.AdminUserUpdateProcedure,
+ pbconnect.AdminUserDeleteProcedure:
+ if req == nil {
+ return errDenied
+ }
+
+ if isSelf(*caller, 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, *caller, req, pb.Perm_PERM_ORG_WRITE)
+ case pbconnect.AdminOrgDeleteProcedure:
+ return a.checkOrgPerm(ctx, *caller, req, pb.Perm_PERM_ORG_DELETE)
+
+ // OrgMember — org-scoped
+ case pbconnect.AdminOrgMemberGetProcedure,
+ pbconnect.AdminOrgMemberListProcedure:
+ return a.checkOrgPerm(ctx, *caller, req, pb.Perm_PERM_ORG_MEMBER_READ)
+ case pbconnect.AdminOrgMemberAddProcedure,
+ pbconnect.AdminOrgMemberUpdateProcedure,
+ pbconnect.AdminOrgMemberRemoveProcedure:
+ return a.checkOrgPerm(ctx, *caller, 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, *caller, req, pb.Perm_PERM_WORKER_WRITE)
+
+ // Worker delete — lookup worker's org, then check
+ case pbconnect.AdminWorkerDeleteProcedure:
+ return a.checkWorkerDelete(ctx, *caller, req)
+
+ default:
+ return errDenied
+ }
+}
+
+var errDenied = connect.NewError(connect.CodePermissionDenied, nil)
+
+// checkOrgPerm extracts OrgRef from request and checks caller's role permission.
+func (a *ApiAuthInterceptor) checkOrgPerm(ctx context.Context, caller callerInfo, req connect.AnyRequest, perm pb.Perm) error {
+ ref := orgRefFromRequest(req)
+ if ref == nil {
+ return errDenied
+ }
+ member, err := a.srv.db.GetOrgMember(ctx, caller.UserID, orgRef(ref), database.UserByID(caller.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, caller callerInfo, req connect.AnyRequest) error {
+ m, ok := req.Any().(*pb.WorkerDeleteRequest)
+ if !ok {
+ return errDenied
+ }
+ w, err := a.srv.db.GetWorker(ctx, caller.UserID, 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, caller.UserID, database.OrgByID(*w.OrgID), database.UserByID(caller.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 caller's own user.
+func isSelf(caller callerInfo, 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) == caller.UserID
+ case *pb.UserRef_Email:
+ return v.Email == caller.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/config.go b/cmd/mirumd/config.go
index 9516a1c..fb7a2e0 100644
--- a/cmd/mirumd/config.go
+++ b/cmd/mirumd/config.go
@@ -25,6 +25,8 @@ type config struct {
GrpcTls tlsConfig `yaml:"grpc_tls"`
WebTls *tlsConfig `yaml:"web_tls"` // optional
+ TrustedProxies []string `yaml:"trusted_proxies"` // CIDR list, empty = trust RemoteAddr only
+
GitHubToken string `yaml:"token"`
WebhookSecret string `yaml:"webhook_secret"`
}
diff --git a/cmd/mirumd/main.go b/cmd/mirumd/main.go
index 2357bd6..77f1643 100644
--- a/cmd/mirumd/main.go
+++ b/cmd/mirumd/main.go
@@ -297,9 +297,26 @@ func daemon(configFile, socketFlag string) {
go srv.PurgeSessions(ctx)
- webSrv := NewWebServer(ctx, srv)
+ adminPath, adminHandler := NewAdminHandler(srv)
+
+ webSrv := NewWebServer(ctx, srv, adminPath, adminHandler)
grpcSrv := NewGrpcServer(ctx, srv)
- adminSrv := NewAdminServer(ctx, srv)
+
+ adminMux := http.NewServeMux()
+ adminMux.Handle(adminPath, adminHandler)
+ adminSrv := &http.Server{
+ Handler: adminMux,
+ ConnContext: func(ctx context.Context, _ net.Conn) context.Context {
+ return context.WithValue(ctx, callerKey{}, &callerInfo{
+ UserID: uuid.Nil,
+ Email: "root@localhost",
+ Superuser: true,
+ })
+ },
+ BaseContext: func(_ net.Listener) context.Context {
+ return ctx
+ },
+ }
grpcLn, webLn, adminLn, err := listeners(cfg)
if err != nil {
diff --git a/cmd/mirumd/server_admin.go b/cmd/mirumd/server_admin.go
index b23b4c3..3bf8765 100644
--- a/cmd/mirumd/server_admin.go
+++ b/cmd/mirumd/server_admin.go
@@ -7,7 +7,6 @@ import (
"context"
"errors"
"fmt"
- "net"
"net/http"
"connectrpc.com/connect"
@@ -20,22 +19,14 @@ import (
"dimidiumlabs/mirum/internal/protocol/pb/pbconnect"
)
-func NewAdminServer(ctx context.Context, srv *server) *http.Server {
+// NewAdminHandler creates the ConnectRPC handler with validation and auth.
+func NewAdminHandler(srv *server) (string, http.Handler) {
as := &adminService{srv: srv}
+ auth := &ApiAuthInterceptor{srv: srv}
- path, handler := pbconnect.NewAdminHandler(as,
- connect.WithInterceptors(validate.NewInterceptor()),
+ return pbconnect.NewAdminHandler(as,
+ connect.WithInterceptors(validate.NewInterceptor(), auth),
)
-
- mux := http.NewServeMux()
- mux.Handle(path, handler)
-
- return &http.Server{
- Handler: mux,
- BaseContext: func(_ net.Listener) context.Context {
- return ctx
- },
- }
}
type adminService struct {
@@ -43,6 +34,19 @@ type adminService struct {
srv *server
}
+// anonActorID is a UUID that is not a real user — used for unauthenticated
+// public requests. RLS will show only public data for this actor.
+var anonActorID = uuid.MustParse("ffffffff-ffff-ffff-ffff-ffffffffffff")
+
+// actorID returns the authenticated user's UUID from context,
+// or anonActorID for unauthenticated public requests.
+func actorID(ctx context.Context) uuid.UUID {
+ if c := CallerFromContext(ctx); c != nil {
+ return c.UserID
+ }
+ return anonActorID
+}
+
// --- Error mapping ---
func mapErr(err error) error {
@@ -175,7 +179,7 @@ func workerToProto(w database.Worker) *pb.Worker {
// --- User handlers ---
func (a *adminService) UserCreate(ctx context.Context, req *connect.Request[pb.UserCreateRequest]) (*connect.Response[pb.UserCreateResponse], error) {
- id, err := a.srv.db.UserCreate(ctx, req.Msg.Email, req.Msg.Password, []byte(a.srv.cfg.Pepper))
+ id, err := a.srv.db.UserCreate(ctx, actorID(ctx), req.Msg.Email, req.Msg.Password, []byte(a.srv.cfg.Pepper))
if err != nil {
return nil, mapErr(err)
}
@@ -183,7 +187,7 @@ func (a *adminService) UserCreate(ctx context.Context, req *connect.Request[pb.U
}
func (a *adminService) UserGet(ctx context.Context, req *connect.Request[pb.UserGetRequest]) (*connect.Response[pb.UserGetResponse], error) {
- u, err := a.srv.db.GetUser(ctx, userRef(req.Msg.User))
+ u, err := a.srv.db.GetUser(ctx, actorID(ctx), userRef(req.Msg.User))
if err != nil {
return nil, mapErr(err)
}
@@ -197,7 +201,7 @@ func (a *adminService) UserList(ctx context.Context, req *connect.Request[pb.Use
filter = *req.Msg.Filter
}
- users, total, err := a.srv.db.ListUsers(ctx, cursor, limit, filter)
+ users, total, err := a.srv.db.ListUsers(ctx, actorID(ctx), cursor, limit, filter)
if err != nil {
return nil, mapErr(err)
}
@@ -219,14 +223,14 @@ 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, userRef(req.Msg.User), req.Msg.Email, req.Msg.Password, []byte(a.srv.cfg.Pepper)); err != nil {
+ if err := a.srv.db.UserUpdate(ctx, actorID(ctx), userRef(req.Msg.User), 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, userRef(req.Msg.User)); err != nil {
+ if err := a.srv.db.UserDelete(ctx, actorID(ctx), userRef(req.Msg.User)); err != nil {
return nil, mapErr(err)
}
return connect.NewResponse(&pb.UserDeleteResponse{}), nil
@@ -239,7 +243,7 @@ func (a *adminService) OrgCreate(ctx context.Context, req *connect.Request[pb.Or
if err != nil {
return nil, mapErr(err)
}
- id, err := a.srv.db.CreateOrganization(ctx, req.Msg.Name, slug, req.Msg.Public, userRef(req.Msg.Owner))
+ id, err := a.srv.db.CreateOrganization(ctx, actorID(ctx), req.Msg.Name, slug, req.Msg.Public, userRef(req.Msg.Owner))
if err != nil {
return nil, mapErr(err)
}
@@ -247,7 +251,7 @@ func (a *adminService) OrgCreate(ctx context.Context, req *connect.Request[pb.Or
}
func (a *adminService) OrgGet(ctx context.Context, req *connect.Request[pb.OrgGetRequest]) (*connect.Response[pb.OrgGetResponse], error) {
- o, err := a.srv.db.GetOrg(ctx, orgRef(req.Msg.Org))
+ o, err := a.srv.db.GetOrg(ctx, actorID(ctx), orgRef(req.Msg.Org))
if err != nil {
return nil, mapErr(err)
}
@@ -261,7 +265,7 @@ func (a *adminService) OrgList(ctx context.Context, req *connect.Request[pb.OrgL
filter = *req.Msg.Filter
}
- orgs, total, err := a.srv.db.ListOrganizations(ctx, cursor, limit, filter)
+ orgs, total, err := a.srv.db.ListOrganizations(ctx, actorID(ctx), cursor, limit, filter)
if err != nil {
return nil, mapErr(err)
}
@@ -291,14 +295,14 @@ func (a *adminService) OrgUpdate(ctx context.Context, req *connect.Request[pb.Or
}
slug = &s
}
- if err := a.srv.db.UpdateOrganization(ctx, orgRef(req.Msg.Org), req.Msg.Name, slug, req.Msg.Public); err != nil {
+ if err := a.srv.db.UpdateOrganization(ctx, actorID(ctx), orgRef(req.Msg.Org), 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, orgRef(req.Msg.Org)); err != nil {
+ if err := a.srv.db.DeleteOrganization(ctx, actorID(ctx), orgRef(req.Msg.Org)); err != nil {
return nil, mapErr(err)
}
return connect.NewResponse(&pb.OrgDeleteResponse{}), nil
@@ -311,14 +315,14 @@ func (a *adminService) OrgMemberAdd(ctx context.Context, req *connect.Request[pb
if !ok {
return nil, connect.NewError(connect.CodeInvalidArgument, fmt.Errorf("invalid role"))
}
- if err := a.srv.db.AddOrgMember(ctx, orgRef(req.Msg.Org), userRef(req.Msg.User), role); err != nil {
+ if err := a.srv.db.AddOrgMember(ctx, actorID(ctx), orgRef(req.Msg.Org), userRef(req.Msg.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, orgRef(req.Msg.Org), userRef(req.Msg.User))
+ m, err := a.srv.db.GetOrgMember(ctx, actorID(ctx), orgRef(req.Msg.Org), userRef(req.Msg.User))
if err != nil {
return nil, mapErr(err)
}
@@ -332,7 +336,7 @@ func (a *adminService) OrgMemberList(ctx context.Context, req *connect.Request[p
filter = *req.Msg.Filter
}
- members, total, err := a.srv.db.ListOrgMembers(ctx, orgRef(req.Msg.Org), cursor, limit, filter)
+ members, total, err := a.srv.db.ListOrgMembers(ctx, actorID(ctx), orgRef(req.Msg.Org), cursor, limit, filter)
if err != nil {
return nil, mapErr(err)
}
@@ -358,14 +362,14 @@ func (a *adminService) OrgMemberUpdate(ctx context.Context, req *connect.Request
if !ok {
return nil, connect.NewError(connect.CodeInvalidArgument, fmt.Errorf("invalid role"))
}
- if err := a.srv.db.UpdateOrgMemberRole(ctx, orgRef(req.Msg.Org), userRef(req.Msg.User), role); err != nil {
+ if err := a.srv.db.UpdateOrgMemberRole(ctx, actorID(ctx), orgRef(req.Msg.Org), userRef(req.Msg.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, orgRef(req.Msg.Org), userRef(req.Msg.User)); err != nil {
+ if err := a.srv.db.RemoveOrgMember(ctx, actorID(ctx), orgRef(req.Msg.Org), userRef(req.Msg.User)); err != nil {
return nil, mapErr(err)
}
return connect.NewResponse(&pb.OrgMemberRemoveResponse{}), nil
@@ -379,7 +383,7 @@ func (a *adminService) WorkerCreate(ctx context.Context, req *connect.Request[pb
r := orgRef(req.Msg.Org)
org = &r
}
- id, err := a.srv.db.CreateWorker(ctx, req.Msg.PublicKey, org)
+ id, err := a.srv.db.CreateWorker(ctx, actorID(ctx), req.Msg.PublicKey, org)
if err != nil {
return nil, mapErr(err)
}
@@ -387,7 +391,7 @@ func (a *adminService) WorkerCreate(ctx context.Context, req *connect.Request[pb
}
func (a *adminService) WorkerGet(ctx context.Context, req *connect.Request[pb.WorkerGetRequest]) (*connect.Response[pb.WorkerGetResponse], error) {
- w, err := a.srv.db.GetWorker(ctx, uuid.UUID(req.Msg.Id))
+ w, err := a.srv.db.GetWorker(ctx, actorID(ctx), uuid.UUID(req.Msg.Id))
if err != nil {
return nil, mapErr(err)
}
@@ -401,7 +405,7 @@ func (a *adminService) WorkerList(ctx context.Context, req *connect.Request[pb.W
filter = *req.Msg.Filter
}
- workers, total, err := a.srv.db.ListWorkers(ctx, cursor, limit, filter)
+ workers, total, err := a.srv.db.ListWorkers(ctx, actorID(ctx), cursor, limit, filter)
if err != nil {
return nil, mapErr(err)
}
@@ -423,7 +427,7 @@ 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, uuid.UUID(req.Msg.Id)); err != nil {
+ if err := a.srv.db.DeleteWorker(ctx, actorID(ctx), uuid.UUID(req.Msg.Id)); 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 afe2a6a..5193457 100644
--- a/cmd/mirumd/server_grpc.go
+++ b/cmd/mirumd/server_grpc.go
@@ -17,6 +17,7 @@ import (
"connectrpc.com/connect"
"connectrpc.com/validate"
+ "github.com/google/uuid"
"dimidiumlabs/mirum/internal/protocol"
"dimidiumlabs/mirum/internal/protocol/pb"
@@ -61,7 +62,7 @@ func NewGrpcServer(ctx context.Context, srv *server) *http.Server {
return errors.New("ed25519 certificate required")
}
- if _, err := srv.db.LookupWorker(context.Background(), pubKey); err != nil {
+ if _, err := srv.db.LookupWorker(context.Background(), uuid.Nil, 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 2ecdea4..fc4a6a6 100644
--- a/cmd/mirumd/server_web.go
+++ b/cmd/mirumd/server_web.go
@@ -15,6 +15,14 @@ import (
"io"
"net"
"net/http"
+ "strings"
+ "time"
+
+ "github.com/google/uuid"
+
+ "github.com/go-chi/chi/v5"
+ "github.com/go-chi/chi/v5/middleware"
+ "github.com/go-chi/httprate"
"dimidiumlabs/mirum/internal/database"
"dimidiumlabs/mirum/internal/forges"
@@ -28,106 +36,47 @@ var (
loginTmpl = template.Must(template.ParseFS(templateFS, "templates/layout.html", "templates/login.html"))
)
-func NewWebServer(ctx context.Context, srv *server) *http.Server {
- mux := http.NewServeMux()
-
- mux.HandleFunc("GET /", func(w http.ResponseWriter, r *http.Request) {
- w.Header().Set("Content-Type", "text/html; charset=utf-8")
- var data struct{ Email, CSRF string }
- if c, err := r.Cookie("session"); err == nil {
- if sess, err := srv.db.UserGetSession(r.Context(), c.Value); err == nil {
- data.Email = sess.Email
- data.CSRF = csrfToken(w, r)
- }
- }
- indexTmpl.ExecuteTemplate(w, "layout", data)
- })
-
- mux.HandleFunc("POST /webhook", func(w http.ResponseWriter, r *http.Request) {
- body, err := io.ReadAll(r.Body)
- if err != nil {
- http.Error(w, "read body", http.StatusBadRequest)
- return
- }
-
- ev, err := srv.forge.Webhook(r, body)
- if errors.Is(err, forges.ErrInvalidSignature) {
- http.Error(w, "invalid signature", http.StatusUnauthorized)
- return
- }
- if err != nil {
- http.Error(w, err.Error(), http.StatusBadRequest)
- return
- }
- if ev == nil {
- w.WriteHeader(http.StatusNoContent)
- return
- }
-
- srv.enqueue(ev)
- w.WriteHeader(http.StatusAccepted)
- })
+func NewWebServer(ctx context.Context, srv *server, adminPath string, adminHandler http.Handler) *http.Server {
+ h := &webHandler{srv: srv}
- mux.HandleFunc("GET /auth/login", func(w http.ResponseWriter, r *http.Request) {
- w.Header().Set("Content-Type", "text/html; charset=utf-8")
- loginTmpl.ExecuteTemplate(w, "layout", map[string]string{"CSRF": csrfToken(w, r)})
- })
+ r := chi.NewRouter()
- mux.HandleFunc("POST /auth/login", func(w http.ResponseWriter, r *http.Request) {
- r.Body = http.MaxBytesReader(w, r.Body, 4096)
- if !csrfOK(r) {
- clearCookie(w, "csrf")
- http.Error(w, "invalid csrf token", http.StatusForbidden)
- return
- }
+ r.Use(middleware.CleanPath)
+ r.Use(middleware.StripSlashes)
+ r.Use(middleware.RequestID)
+ r.Use(trustedProxyMiddleware(srv.cfg.TrustedProxies))
+ r.Use(middleware.Logger)
+ r.Use(middleware.Recoverer)
+ r.Use(middleware.Compress(5))
+ r.Use(middleware.Heartbeat("/ping"))
+ r.Use(middleware.Timeout(30 * time.Second))
+ r.Use(middleware.RequestSize(64 << 20)) // 64 MiB global body limit
- email := r.FormValue("email")
- password := r.FormValue("password")
+ // The authorization session sets the user to ctx
+ r.Use(h.SessionMiddleware)
- userID, err := srv.db.UserVerifyPassword(r.Context(), email, password, []byte(srv.cfg.Pepper))
- if err != nil {
- w.WriteHeader(http.StatusUnauthorized)
- loginTmpl.ExecuteTemplate(w, "layout", map[string]string{
- "Error": "Invalid credentials",
- "CSRF": csrfToken(w, r),
- })
- return
- }
+ r.Get("/", h.index)
+ r.Post("/webhook", h.webhook)
- token, err := srv.db.UserCreateSession(r.Context(), userID)
- if err != nil {
- http.Error(w, "internal error", http.StatusInternalServerError)
- return
- }
+ r.Route("/auth", func(r chi.Router) {
+ r.Use(middleware.NoCache)
+ r.Use(httprate.LimitByIP(10, time.Minute))
+ r.Use(middleware.RequestSize(4096))
- http.SetCookie(w, &http.Cookie{
- Name: "session",
- Value: token,
- Path: "/",
- HttpOnly: true,
- Secure: true,
- SameSite: http.SameSiteLaxMode,
- MaxAge: int(database.SessionTTL.Seconds()),
- })
- http.Redirect(w, r, "/", http.StatusSeeOther)
+ r.Get("/login", h.loginPage)
+ r.Post("/login", h.login)
+ r.Post("/logout", h.logout)
})
- mux.HandleFunc("POST /auth/logout", func(w http.ResponseWriter, r *http.Request) {
- if !csrfOK(r) {
- http.Error(w, "invalid csrf token", http.StatusForbidden)
- return
- }
- if c, err := r.Cookie("session"); err == nil {
- srv.db.UserDeleteSession(r.Context(), c.Value)
- }
- clearCookie(w, "session")
- clearCookie(w, "csrf")
- http.Redirect(w, r, "/auth/login", http.StatusSeeOther)
+ r.Route("/api/v1", func(r chi.Router) {
+ r.Use(middleware.NoCache)
+ r.Use(httprate.LimitByIP(300, time.Minute))
+ r.Mount(adminPath, adminHandler)
})
- var tls_config *tls.Config = nil
+ var tlsCfg *tls.Config
if srv.cfg.WebTls != nil {
- tls_config = &tls.Config{
+ tlsCfg = &tls.Config{
MinVersion: tls.VersionTLS13,
GetCertificate: func(_ *tls.ClientHelloInfo) (*tls.Certificate, error) {
cert, err := tls.LoadX509KeyPair(srv.cfg.WebTls.Cert, srv.cfg.WebTls.Key)
@@ -137,14 +86,142 @@ func NewWebServer(ctx context.Context, srv *server) *http.Server {
}
return &http.Server{
- Handler: mux,
- TLSConfig: tls_config,
+ Handler: r,
+ TLSConfig: tlsCfg,
BaseContext: func(_ net.Listener) context.Context {
return ctx
},
}
}
+type callerKey struct{}
+
+type callerInfo struct {
+ UserID uuid.UUID
+ Email string
+ Superuser bool
+}
+
+type webHandler struct {
+ srv *server
+}
+
+// CallerFromContext returns the authenticated caller, or nil.
+func CallerFromContext(ctx context.Context) *callerInfo {
+ if v, ok := ctx.Value(callerKey{}).(*callerInfo); ok {
+ return v
+ }
+ return nil
+}
+
+// SessionMiddleware resolves the session cookie and puts callerInfo 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("session"); err == nil {
+ if sess, err := h.srv.db.UserGetSession(r.Context(), uuid.Nil, c.Value); err == nil {
+ caller := &callerInfo{
+ UserID: sess.UserID,
+ Email: sess.Email,
+ }
+ ctx := context.WithValue(r.Context(), callerKey{}, caller)
+ r = r.WithContext(ctx)
+ }
+ }
+ next.ServeHTTP(w, r)
+ })
+}
+
+func (h *webHandler) index(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "text/html; charset=utf-8")
+ var data struct{ Email, CSRF string }
+ if caller := CallerFromContext(r.Context()); caller != nil {
+ data.Email = caller.Email
+ data.CSRF = csrfToken(w, r)
+ }
+ indexTmpl.ExecuteTemplate(w, "layout", data)
+}
+
+func (h *webHandler) webhook(w http.ResponseWriter, r *http.Request) {
+ body, err := io.ReadAll(r.Body)
+ if err != nil {
+ http.Error(w, "read body", http.StatusBadRequest)
+ return
+ }
+
+ ev, err := h.srv.forge.Webhook(r, body)
+ if errors.Is(err, forges.ErrInvalidSignature) {
+ http.Error(w, "invalid signature", http.StatusUnauthorized)
+ return
+ }
+ if err != nil {
+ http.Error(w, err.Error(), http.StatusBadRequest)
+ return
+ }
+ if ev == nil {
+ w.WriteHeader(http.StatusNoContent)
+ return
+ }
+
+ h.srv.enqueue(ev)
+ w.WriteHeader(http.StatusAccepted)
+}
+
+func (h *webHandler) loginPage(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "text/html; charset=utf-8")
+ loginTmpl.ExecuteTemplate(w, "layout", map[string]string{"CSRF": csrfToken(w, r)})
+}
+
+func (h *webHandler) login(w http.ResponseWriter, r *http.Request) {
+ if !csrfOK(r) {
+ clearCookie(w, "csrf")
+ http.Error(w, "invalid csrf token", http.StatusForbidden)
+ return
+ }
+
+ email := r.FormValue("email")
+ password := r.FormValue("password")
+
+ userID, err := h.srv.db.UserVerifyPassword(r.Context(), uuid.Nil, email, password, []byte(h.srv.cfg.Pepper))
+ if err != nil {
+ w.WriteHeader(http.StatusUnauthorized)
+ loginTmpl.ExecuteTemplate(w, "layout", map[string]string{
+ "Error": "Invalid credentials",
+ "CSRF": csrfToken(w, r),
+ })
+ return
+ }
+
+ token, err := h.srv.db.UserCreateSession(r.Context(), uuid.Nil, userID)
+ if err != nil {
+ http.Error(w, "internal error", http.StatusInternalServerError)
+ return
+ }
+
+ http.SetCookie(w, &http.Cookie{
+ Name: "session",
+ Value: token,
+ Path: "/",
+ HttpOnly: true,
+ Secure: true,
+ SameSite: http.SameSiteLaxMode,
+ MaxAge: int(database.SessionTTL.Seconds()),
+ })
+ http.Redirect(w, r, "/", http.StatusSeeOther)
+}
+
+func (h *webHandler) logout(w http.ResponseWriter, r *http.Request) {
+ if !csrfOK(r) {
+ http.Error(w, "invalid csrf token", http.StatusForbidden)
+ return
+ }
+ if c, err := r.Cookie("session"); err == nil {
+ h.srv.db.UserDeleteSession(r.Context(), uuid.Nil, c.Value)
+ }
+ clearCookie(w, "session")
+ clearCookie(w, "csrf")
+ http.Redirect(w, r, "/auth/login", http.StatusSeeOther)
+}
+
// csrfToken returns the current CSRF token, setting a cookie if absent.
func csrfToken(w http.ResponseWriter, r *http.Request) string {
if c, err := r.Cookie("csrf"); err == nil && c.Value != "" {
@@ -184,3 +261,59 @@ func clearCookie(w http.ResponseWriter, name string) {
MaxAge: -1,
})
}
+
+// trustedProxyMiddleware resolves the real client IP from X-Forwarded-For,
+// walking right-to-left and stopping at the first untrusted hop.
+// Empty cidrs = trust RemoteAddr only (safe default).
+func trustedProxyMiddleware(cidrs []string) func(http.Handler) http.Handler {
+ nets := make([]*net.IPNet, 0, len(cidrs))
+ for _, c := range cidrs {
+ _, n, err := net.ParseCIDR(c)
+ if err != nil {
+ panic("invalid trusted_proxies CIDR: " + c)
+ }
+ nets = append(nets, n)
+ }
+
+ isTrusted := func(ip net.IP) bool {
+ for _, n := range nets {
+ if n.Contains(ip) {
+ return true
+ }
+ }
+ return false
+ }
+
+ return func(next http.Handler) http.Handler {
+ return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if len(nets) == 0 {
+ next.ServeHTTP(w, r)
+ return
+ }
+
+ host, _, _ := net.SplitHostPort(r.RemoteAddr)
+ ip := net.ParseIP(host)
+ if ip == nil || !isTrusted(ip) {
+ // RemoteAddr is not a trusted proxy — use as-is.
+ next.ServeHTTP(w, r)
+ return
+ }
+
+ // Walk X-Forwarded-For right to left.
+ xff := strings.Split(r.Header.Get("X-Forwarded-For"), ",")
+ for i := len(xff) - 1; i >= 0; i-- {
+ candidate := strings.TrimSpace(xff[i])
+ ip = net.ParseIP(candidate)
+ if ip == nil {
+ break // garbage — stop, don't trust anything further left
+ }
+ if !isTrusted(ip) {
+ r.RemoteAddr = candidate + ":0"
+ break
+ }
+ }
+
+ next.ServeHTTP(w, r)
+ })
+ }
+}
diff --git a/go.mod b/go.mod
index b0cef7d..4bb59e7 100644
--- a/go.mod
+++ b/go.mod
@@ -7,6 +7,8 @@ require (
connectrpc.com/connect v1.19.1
connectrpc.com/validate v0.6.0
github.com/coreos/go-systemd/v22 v22.7.0
+ github.com/go-chi/chi/v5 v5.2.5
+ github.com/go-chi/httprate v0.15.0
github.com/google/uuid v1.6.0
github.com/huandu/go-sqlbuilder v1.40.1
github.com/jackc/pgerrcode v0.0.0-20250907135507-afb5586c32a6
@@ -34,11 +36,13 @@ require (
github.com/jackc/pgpassfile v1.0.0 // indirect
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
github.com/jackc/puddle/v2 v2.2.2 // indirect
+ github.com/klauspost/cpuid/v2 v2.2.10 // indirect
github.com/mitchellh/copystructure v1.2.0 // indirect
github.com/mitchellh/reflectwalk v1.0.2 // indirect
github.com/shopspring/decimal v1.4.0 // indirect
github.com/spf13/cast v1.7.0 // indirect
github.com/spf13/pflag v1.0.5 // indirect
+ github.com/zeebo/xxh3 v1.0.2 // indirect
golang.org/x/exp v0.0.0-20250911091902-df9299821621 // indirect
golang.org/x/sync v0.20.0 // indirect
golang.org/x/sys v0.42.0 // indirect
diff --git a/go.sum b/go.sum
index f79ef27..97d9640 100644
--- a/go.sum
+++ b/go.sum
@@ -28,6 +28,10 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8=
github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0=
+github.com/go-chi/chi/v5 v5.2.5 h1:Eg4myHZBjyvJmAFjFvWgrqDTXFyOzjj7YIm3L3mu6Ug=
+github.com/go-chi/chi/v5 v5.2.5/go.mod h1:X7Gx4mteadT3eDOMTsXzmI4/rwUpOwBHLpAfupzFJP0=
+github.com/go-chi/httprate v0.15.0 h1:j54xcWV9KGmPf/X4H32/aTH+wBlrvxL7P+SdnRqxh5g=
+github.com/go-chi/httprate v0.15.0/go.mod h1:rzGHhVrsBn3IMLYDOZQsSU4fJNWcjui4fWKJcCId1R4=
github.com/google/cel-go v0.27.0 h1:e7ih85+4qVrBuqQWTW4FKSqZYokVuc3HnhH5keboFTo=
github.com/google/cel-go v0.27.0/go.mod h1:tTJ11FWqnhw5KKpnWpvW9CJC3Y9GK4EIS0WXnBbebzw=
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
@@ -57,6 +61,8 @@ github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
github.com/jackc/tern/v2 v2.3.6 h1:sqBIZ/CBtfMLz7zdUof0N6cVUBRVGBZ7S+F2OdCp9XU=
github.com/jackc/tern/v2 v2.3.6/go.mod h1:SrtwsdBRKkeTOjuLd6ISNqaLOtaLX+jOTLrpP+lJQe0=
+github.com/klauspost/cpuid/v2 v2.2.10 h1:tBs3QSyvjDyFTq3uoc/9xFpCuOsJQFNPiAhYdw2skhE=
+github.com/klauspost/cpuid/v2 v2.2.10/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
@@ -86,6 +92,10 @@ github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81P
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
+github.com/zeebo/assert v1.3.0 h1:g7C04CbJuIDKNPFHmsk4hwZDO5O+kntRxzaUoNXj+IQ=
+github.com/zeebo/assert v1.3.0/go.mod h1:Pq9JiuJQpG8JLJdtkwrJESF0Foym2/D9XMU5ciN/wJ0=
+github.com/zeebo/xxh3 v1.0.2 h1:xZmwmqxHZA8AI603jOQ0tMqmBr9lPeFwGg6d+xy9DC0=
+github.com/zeebo/xxh3 v1.0.2/go.mod h1:5NWz9Sef7zIDm2JHfFlcQvNekmcEl9ekUZQQKCYaDcA=
go.starlark.net v0.0.0-20260326113308-fadfc96def35 h1:VYAqieSOJNxBDX8KJneTAwvdf4J4zRDE2u+UFXtt9h4=
go.starlark.net v0.0.0-20260326113308-fadfc96def35/go.mod h1:Iue6g6iirlfLoVi/DYCi5/x0h/bAOuWF3dULTKpt2Vo=
go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc=
diff --git a/internal/database/database.go b/internal/database/database.go
index 40940bb..8e36e58 100644
--- a/internal/database/database.go
+++ b/internal/database/database.go
@@ -7,6 +7,8 @@ import (
"context"
"errors"
+ "github.com/google/uuid"
+ "github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/jackc/tern/v2/migrate"
)
@@ -44,6 +46,19 @@ func (db *DB) Close() {
db.Pool.Close()
}
+// beginAs starts a transaction and sets the RLS actor.
+func (db *DB) beginAs(ctx context.Context, actor uuid.UUID) (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)", actor.String()); 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)
@@ -52,67 +67,201 @@ func (db *DB) Migrate(ctx context.Context) error {
}
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.user_id', '00000000-0000-0000-0000-000000000000', true);
+ // This sets the root superuser for the migration's transaction only.
migrator, err := migrate.NewMigrator(ctx, conn.Conn(), "schema_version")
if err != nil {
return errors.Join(ErrMigrate, err)
}
- migrator.AppendMigration("create_users",
- `CREATE TABLE users (
+ 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
- )`,
- `DROP TABLE users`,
- )
+ );
+
+ INSERT INTO users (id, email, password, superuser)
+ VALUES ('00000000-0000-0000-0000-000000000000', 'root@localhost', '', true);
+
+ CREATE FUNCTION app_user_id() RETURNS uuid STABLE AS $$
+ SELECT current_setting('app.user_id', true)::uuid;
+ $$ LANGUAGE sql;
- migrator.AppendMigration("create_sessions",
- `CREATE TABLE sessions (
+ CREATE FUNCTION app_issuper() RETURNS boolean STABLE AS $$
+ SELECT 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`,
- )
+ CREATE INDEX sessions_expires_at ON sessions (expires_at)
+ `, `
+ DROP TABLE sessions
+ `)
- migrator.AppendMigration("create_organizations",
- `CREATE TABLE organizations (
+ 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`,
- )
+ );
+ `, `
+ DROP TABLE organizations
+ `)
- migrator.AppendMigration("create_org_members",
- `CREATE TABLE org_members (
+ 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`,
- )
+ CREATE INDEX org_members_user_id ON org_members (user_id)
+ `, `
+ DROP TABLE org_members
+ `)
- migrator.AppendMigration("create_workers",
- `CREATE TABLE workers (
+ 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`,
- )
+ )
+ `, `
+ 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());
+ CREATE POLICY user_select ON users FOR SELECT USING (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 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());
+ CREATE POLICY select_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_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
index 90a81bd..cca7dff 100644
--- a/internal/database/organization.go
+++ b/internal/database/organization.go
@@ -78,7 +78,13 @@ func resolveOrg(ctx context.Context, tx pgx.Tx, ref OrgRef) (uuid.UUID, error) {
}
// GetOrg returns an org by ref (ID or slug).
-func (db *DB) GetOrg(ctx context.Context, ref OrgRef) (*Organization, error) {
+func (db *DB) GetOrg(ctx context.Context, actor uuid.UUID, 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()
@@ -88,11 +94,10 @@ func (db *DB) GetOrg(ctx context.Context, ref OrgRef) (*Organization, error) {
Build()
var o Organization
- if err := db.Pool.QueryRow(ctx, sql, args...).Scan(&o.ID, &o.Name, &o.Slug, &o.Public, &o.CreatedAt); err != nil {
+ 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
}
@@ -100,8 +105,8 @@ func (db *DB) GetOrg(ctx context.Context, ref OrgRef) (*Organization, error) {
}
// CreateOrganization creates an org and adds the owner as the first member.
-func (db *DB) CreateOrganization(ctx context.Context, name, slug string, public bool, owner UserRef) (uuid.UUID, error) {
- tx, err := db.Pool.Begin(ctx)
+func (db *DB) CreateOrganization(ctx context.Context, actor uuid.UUID, name, slug string, public bool, owner UserRef) (uuid.UUID, error) {
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return uuid.Nil, err
}
@@ -136,12 +141,12 @@ func (db *DB) CreateOrganization(ctx context.Context, name, slug string, public
}
// UpdateOrganization updates an org's name, slug, and/or public flag.
-func (db *DB) UpdateOrganization(ctx context.Context, ref OrgRef, name *string, slug *string, public *bool) error {
+func (db *DB) UpdateOrganization(ctx context.Context, actor uuid.UUID, ref OrgRef, name *string, slug *string, public *bool) error {
if name == nil && slug == nil && public == nil {
return nil
}
- tx, err := db.Pool.Begin(ctx)
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return err
}
@@ -179,8 +184,8 @@ func (db *DB) UpdateOrganization(ctx context.Context, ref OrgRef, name *string,
}
// DeleteOrganization soft-deletes an org and removes all members.
-func (db *DB) DeleteOrganization(ctx context.Context, ref OrgRef) error {
- tx, err := db.Pool.Begin(ctx)
+func (db *DB) DeleteOrganization(ctx context.Context, actor uuid.UUID, ref OrgRef) error {
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return err
}
@@ -205,12 +210,12 @@ func (db *DB) DeleteOrganization(ctx context.Context, ref OrgRef) error {
}
// ListOrganizations returns a page of orgs and the total count.
-func (db *DB) ListOrganizations(ctx context.Context, cursor uuid.UUID, limit int, filter string) ([]Organization, int, error) {
+func (db *DB) ListOrganizations(ctx context.Context, actor uuid.UUID, cursor uuid.UUID, limit int, filter string) ([]Organization, int, error) {
if filter != "" {
return nil, 0, ErrFilterNotImplemented
}
- tx, err := db.Pool.Begin(ctx)
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return nil, 0, err
}
@@ -252,8 +257,8 @@ func (db *DB) ListOrganizations(ctx context.Context, cursor uuid.UUID, limit int
}
// GetOrgMember returns a single member's info.
-func (db *DB) GetOrgMember(ctx context.Context, org OrgRef, user UserRef) (*OrgMember, error) {
- tx, err := db.Pool.Begin(ctx)
+func (db *DB) GetOrgMember(ctx context.Context, actor uuid.UUID, org OrgRef, user UserRef) (*OrgMember, error) {
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return nil, err
}
@@ -284,12 +289,12 @@ func (db *DB) GetOrgMember(ctx context.Context, org OrgRef, user UserRef) (*OrgM
}
// ListOrgMembers returns a page of members for an org.
-func (db *DB) ListOrgMembers(ctx context.Context, org OrgRef, cursor uuid.UUID, limit int, filter string) ([]OrgMember, int, error) {
+func (db *DB) ListOrgMembers(ctx context.Context, actor uuid.UUID, org OrgRef, cursor uuid.UUID, limit int, filter string) ([]OrgMember, int, error) {
if filter != "" {
return nil, 0, ErrFilterNotImplemented
}
- tx, err := db.Pool.Begin(ctx)
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return nil, 0, err
}
@@ -339,8 +344,8 @@ func (db *DB) ListOrgMembers(ctx context.Context, org OrgRef, cursor uuid.UUID,
}
// AddOrgMember adds a user to an org with the given role.
-func (db *DB) AddOrgMember(ctx context.Context, org OrgRef, user UserRef, role string) error {
- tx, err := db.Pool.Begin(ctx)
+func (db *DB) AddOrgMember(ctx context.Context, actor uuid.UUID, org OrgRef, user UserRef, role string) error {
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return err
}
@@ -370,8 +375,8 @@ func (db *DB) AddOrgMember(ctx context.Context, org OrgRef, user UserRef, role s
}
// UpdateOrgMemberRole changes a member's role. Fails if demoting the last owner.
-func (db *DB) UpdateOrgMemberRole(ctx context.Context, org OrgRef, user UserRef, newRole string) error {
- tx, err := db.Pool.Begin(ctx)
+func (db *DB) UpdateOrgMemberRole(ctx context.Context, actor uuid.UUID, org OrgRef, user UserRef, newRole string) error {
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return err
}
@@ -421,8 +426,8 @@ func (db *DB) UpdateOrgMemberRole(ctx context.Context, org OrgRef, user UserRef,
}
// RemoveOrgMember removes a user from an org. Fails if they are the last owner.
-func (db *DB) RemoveOrgMember(ctx context.Context, org OrgRef, user UserRef) error {
- tx, err := db.Pool.Begin(ctx)
+func (db *DB) RemoveOrgMember(ctx context.Context, actor uuid.UUID, org OrgRef, user UserRef) error {
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return err
}
diff --git a/internal/database/user.go b/internal/database/user.go
index 78dae2b..577ac2f 100644
--- a/internal/database/user.go
+++ b/internal/database/user.go
@@ -134,14 +134,20 @@ func hashPassword(password string, pepper []byte) (string, error) {
// 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, email, password string, pepper []byte) (uuid.UUID, error) {
+func (db *DB) UserCreate(ctx context.Context, actor uuid.UUID, email, password string, pepper []byte) (uuid.UUID, error) {
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 := db.Pool.QueryRow(ctx,
+ if err := tx.QueryRow(ctx,
`INSERT INTO users (email, password) VALUES ($1, $2) RETURNING id`,
email, hash,
).Scan(&id); err != nil {
@@ -152,13 +158,19 @@ func (db *DB) UserCreate(ctx context.Context, email, password string, pepper []b
return uuid.Nil, err
}
- return id, nil
+ return id, tx.Commit(ctx)
}
// GetUser returns a user by ref (ID or email).
-func (db *DB) GetUser(ctx context.Context, ref UserRef) (*User, error) {
+func (db *DB) GetUser(ctx context.Context, actor uuid.UUID, 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").
@@ -166,11 +178,10 @@ func (db *DB) GetUser(ctx context.Context, ref UserRef) (*User, error) {
Build()
var u User
- if err := db.Pool.QueryRow(ctx, sql, args...).Scan(&u.ID, &u.Email, &u.CreatedAt); err != nil {
+ 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
}
@@ -178,12 +189,12 @@ func (db *DB) GetUser(ctx context.Context, ref UserRef) (*User, error) {
}
// ListUsers returns a page of users and the total count.
-func (db *DB) ListUsers(ctx context.Context, cursor uuid.UUID, limit int, filter string) ([]User, int, error) {
+func (db *DB) ListUsers(ctx context.Context, actor uuid.UUID, cursor uuid.UUID, limit int, filter string) ([]User, int, error) {
if filter != "" {
return nil, 0, ErrFilterNotImplemented
}
- tx, err := db.Pool.Begin(ctx)
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return nil, 0, err
}
@@ -269,12 +280,12 @@ func resolveUser(ctx context.Context, tx pgx.Tx, ref UserRef) (uuid.UUID, error)
// UserUpdate updates a user's email and/or password.
// Invalidates all sessions when password changes.
-func (db *DB) UserUpdate(ctx context.Context, ref UserRef, email *string, password *string, pepper []byte) error {
+func (db *DB) UserUpdate(ctx context.Context, actor uuid.UUID, ref UserRef, email *string, password *string, pepper []byte) error {
if email == nil && password == nil {
return nil
}
- tx, err := db.Pool.Begin(ctx)
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return err
}
@@ -323,8 +334,8 @@ func (db *DB) UserUpdate(ctx context.Context, ref UserRef, email *string, passwo
// UserDelete soft-deletes a user.
// Fails if the user is the sole owner of any organization.
-func (db *DB) UserDelete(ctx context.Context, ref UserRef) error {
- tx, err := db.Pool.Begin(ctx)
+func (db *DB) UserDelete(ctx context.Context, actor uuid.UUID, ref UserRef) error {
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return err
}
@@ -357,11 +368,17 @@ func (db *DB) UserDelete(ctx context.Context, ref UserRef) error {
}
// UserVerifyPassword checks credentials and returns the user ID.
-func (db *DB) UserVerifyPassword(ctx context.Context, email, password string, pepper []byte) (uuid.UUID, error) {
+func (db *DB) UserVerifyPassword(ctx context.Context, actor uuid.UUID, 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 := db.Pool.QueryRow(ctx,
+ if err := tx.QueryRow(ctx,
`SELECT id, password FROM users WHERE email = $1 AND deleted_at IS NULL`,
email,
).Scan(&id, &hash); err != nil {
@@ -376,12 +393,18 @@ func (db *DB) UserVerifyPassword(ctx context.Context, email, password string, pe
}
// UserGetSession returns session info for a valid, non-expired session.
-func (db *DB) UserGetSession(ctx context.Context, token string) (*Session, error) {
+func (db *DB) UserGetSession(ctx context.Context, actor uuid.UUID, token string) (*Session, error) {
+ tx, err := db.beginAs(ctx, actor)
+ if err != nil {
+ return nil, err
+ }
+ defer tx.Rollback(ctx)
+
var s Session
h := hashToken(token)
var expiresAt time.Time
- if err := db.Pool.QueryRow(ctx,
+ if err := tx.QueryRow(ctx,
`SELECT s.user_id, u.email, 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`,
@@ -391,7 +414,7 @@ func (db *DB) UserGetSession(ctx context.Context, token string) (*Session, error
}
if time.Until(expiresAt) < SessionTTL/2 {
- if _, err := db.Pool.Exec(ctx,
+ if _, err := tx.Exec(ctx,
`UPDATE sessions SET expires_at = now() + $2 WHERE token = $1`,
h, SessionTTL,
); err != nil {
@@ -399,35 +422,57 @@ func (db *DB) UserGetSession(ctx context.Context, token string) (*Session, error
}
}
- return &s, nil
+ return &s, tx.Commit(ctx)
}
// UserCreateSession generates a random token, stores its hash, and returns the token.
-func (db *DB) UserCreateSession(ctx context.Context, userID uuid.UUID) (string, error) {
+func (db *DB) UserCreateSession(ctx context.Context, actor uuid.UUID, 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 := db.Pool.Exec(ctx,
+ if _, err := tx.Exec(ctx,
`INSERT INTO sessions (token, user_id, expires_at) VALUES ($1, $2, now() + $3)`,
hashToken(token), userID, SessionTTL,
); err != nil {
return "", err
}
- return token, nil
+ return token, tx.Commit(ctx)
}
// UserDeleteSession removes a session (logout).
-func (db *DB) UserDeleteSession(ctx context.Context, token string) error {
- _, err := db.Pool.Exec(ctx, `DELETE FROM sessions WHERE token = $1`, hashToken(token))
- return err
+func (db *DB) UserDeleteSession(ctx context.Context, actor uuid.UUID, 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.
+// PurgeExpiredSessions deletes all expired sessions. Runs as root.
func (db *DB) PurgeExpiredSessions(ctx context.Context) error {
- _, err := db.Pool.Exec(ctx, `DELETE FROM sessions WHERE expires_at < now()`)
- return err
+ tx, err := db.beginAs(ctx, uuid.Nil)
+ 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/worker.go b/internal/database/worker.go
index c0ff46b..a92e6e6 100644
--- a/internal/database/worker.go
+++ b/internal/database/worker.go
@@ -26,16 +26,21 @@ type Worker struct {
}
// GetWorker returns a worker by ID.
-func (db *DB) GetWorker(ctx context.Context, id uuid.UUID) (*Worker, error) {
+func (db *DB) GetWorker(ctx context.Context, actor uuid.UUID, 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 := db.Pool.QueryRow(ctx,
+ 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
}
@@ -43,8 +48,8 @@ func (db *DB) GetWorker(ctx context.Context, id uuid.UUID) (*Worker, error) {
}
// CreateWorker registers a new worker with the given public key and optional org.
-func (db *DB) CreateWorker(ctx context.Context, publicKey []byte, org *OrgRef) (uuid.UUID, error) {
- tx, err := db.Pool.Begin(ctx)
+func (db *DB) CreateWorker(ctx context.Context, actor uuid.UUID, publicKey []byte, org *OrgRef) (uuid.UUID, error) {
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return uuid.Nil, err
}
@@ -71,28 +76,33 @@ func (db *DB) CreateWorker(ctx context.Context, publicKey []byte, org *OrgRef) (
}
// DeleteWorker soft-deletes a worker by ID.
-func (db *DB) DeleteWorker(ctx context.Context, id uuid.UUID) error {
- tag, err := db.Pool.Exec(ctx,
+func (db *DB) DeleteWorker(ctx context.Context, actor uuid.UUID, 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 nil
+ return tx.Commit(ctx)
}
// ListWorkers returns a page of workers and the total count.
-func (db *DB) ListWorkers(ctx context.Context, cursor uuid.UUID, limit int, filter string) ([]Worker, int, error) {
+func (db *DB) ListWorkers(ctx context.Context, actor uuid.UUID, cursor uuid.UUID, limit int, filter string) ([]Worker, int, error) {
if filter != "" {
return nil, 0, ErrFilterNotImplemented
}
- tx, err := db.Pool.Begin(ctx)
+ tx, err := db.beginAs(ctx, actor)
if err != nil {
return nil, 0, err
}
@@ -134,9 +144,15 @@ func (db *DB) ListWorkers(ctx context.Context, cursor uuid.UUID, limit int, filt
}
// LookupWorker finds an active worker by its ed25519 public key.
-func (db *DB) LookupWorker(ctx context.Context, publicKey []byte) (*Worker, error) {
+func (db *DB) LookupWorker(ctx context.Context, actor uuid.UUID, 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 := db.Pool.QueryRow(ctx,
+ 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 {
diff --git a/proto/admin.proto b/proto/admin.proto
index e3edb21..7f2231a 100644
--- a/proto/admin.proto
+++ b/proto/admin.proto
@@ -44,6 +44,23 @@ enum Role {
ROLE_MEMBER = 3;
}
+enum Perm {
+ PERM_NONE = 0;
+
+ PERM_USER_READ = 1;
+ PERM_USER_WRITE = 2;
+ PERM_USER_DELETE = 3;
+
+ PERM_ORG_READ = 4;
+ PERM_ORG_WRITE = 5;
+ PERM_ORG_DELETE = 6;
+ PERM_ORG_MEMBER_READ = 7;
+ PERM_ORG_MEMBER_WRITE = 8;
+
+ PERM_WORKER_READ = 9;
+ PERM_WORKER_WRITE = 10;
+}
+
message UserRef {
oneof ref {
option (buf.validate.oneof).required = true;