diff options
Diffstat
| -rw-r--r-- | cmd/mirumd/api_auth.go | 233 | +233 −0 |
| -rw-r--r-- | cmd/mirumd/config.go | 2 | +2 −0 |
| -rw-r--r-- | cmd/mirumd/main.go | 21 | +19 −2 |
| -rw-r--r-- | cmd/mirumd/server_admin.go | 70 | +37 −33 |
| -rw-r--r-- | cmd/mirumd/server_grpc.go | 3 | +2 −1 |
| -rw-r--r-- | cmd/mirumd/server_web.go | 315 | +224 −91 |
| -rw-r--r-- | go.mod | 4 | +4 −0 |
| -rw-r--r-- | go.sum | 10 | +10 −0 |
| -rw-r--r-- | internal/database/database.go | 199 | +174 −25 |
| -rw-r--r-- | internal/database/organization.go | 47 | +26 −21 |
| -rw-r--r-- | internal/database/user.go | 99 | +72 −27 |
| -rw-r--r-- | internal/database/worker.go | 42 | +29 −13 |
| -rw-r--r-- | proto/admin.proto | 17 | +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; |
