From a82b57af5f68ce74c18a1d3ba1d839222d1edfd5 Mon Sep 17 00:00:00 2001 From: Nikolay Govorov Date: Fri, 3 Apr 2026 17:18:23 +0100 Subject: Add org entity --- cmd/mirumd/main.go | 266 ++++++++++++++++++++- cmd/mirumd/server_admin.go | 106 ++++++++- cmd/mirumd/server_web.go | 8 +- cmd/mirumw/client.go | 1 + internal/database/database.go | 354 +++------------------------- internal/database/organization.go | 314 ++++++++++++++++++++++++ internal/database/user.go | 285 ++++++++++++++++++++++ internal/database/worker.go | 91 +++++++ internal/protocol/handshake.go | 71 +++++- internal/protocol/handshake_test.go | 90 +++++++ internal/protocol/key.go | 70 ------ internal/protocol/key_test.go | 100 -------- internal/protocol/platform.go | 45 ++++ internal/protocol/platform_test.go | 45 ++++ internal/protocol/version.go | 54 ----- internal/protocol/version_test.go | 45 ---- proto/admin.proto | 113 ++++++++- 17 files changed, 1432 insertions(+), 626 deletions(-) create mode 100644 internal/database/organization.go create mode 100644 internal/database/user.go create mode 100644 internal/database/worker.go delete mode 100644 internal/protocol/key.go delete mode 100644 internal/protocol/key_test.go create mode 100644 internal/protocol/platform_test.go delete mode 100644 internal/protocol/version.go delete mode 100644 internal/protocol/version_test.go diff --git a/cmd/mirumd/main.go b/cmd/mirumd/main.go index 32dc09c..b7210e6 100644 --- a/cmd/mirumd/main.go +++ b/cmd/mirumd/main.go @@ -14,12 +14,15 @@ import ( "net" "net/http" "os" + "os/user" + "strconv" "time" - "connectrpc.com/connect" "dimidiumlabs/mirum/internal/database" "dimidiumlabs/mirum/internal/forges" + "connectrpc.com/connect" + "dimidiumlabs/mirum/internal/protocol/pb" "dimidiumlabs/mirum/internal/protocol/pb/pbconnect" "dimidiumlabs/mirum/internal/supervisor" @@ -127,6 +130,132 @@ func main() { } workerCmd.AddCommand(workerListCmd) + orgCmd := &cobra.Command{Use: "org", Short: "Manage organizations"} + root.AddCommand(orgCmd) + + orgCreateCmd := &cobra.Command{ + Use: "create", + Short: "Create an organization", + Run: func(cmd *cobra.Command, args []string) { + name, _ := cmd.Flags().GetString("name") + slug, _ := cmd.Flags().GetString("slug") + public, _ := cmd.Flags().GetBool("public") + owner, _ := cmd.Flags().GetString("owner") + orgCreate(socketPath, name, slug, public, owner) + }, + } + orgCmd.AddCommand(orgCreateCmd) + orgCreateCmd.Flags().String("name", "", "display name") + orgCreateCmd.Flags().String("slug", "", "URL slug") + orgCreateCmd.Flags().Bool("public", false, "public visibility") + orgCreateCmd.Flags().String("owner", "", "owner email") + _ = orgCreateCmd.MarkFlagRequired("name") + _ = orgCreateCmd.MarkFlagRequired("slug") + _ = orgCreateCmd.MarkFlagRequired("owner") + + orgDeleteCmd := &cobra.Command{ + Use: "delete", + Short: "Delete an organization", + Run: func(cmd *cobra.Command, args []string) { + slug, _ := cmd.Flags().GetString("slug") + orgDelete(socketPath, slug) + }, + } + orgCmd.AddCommand(orgDeleteCmd) + orgDeleteCmd.Flags().String("slug", "", "org slug") + _ = orgDeleteCmd.MarkFlagRequired("slug") + + orgRenameCmd := &cobra.Command{ + Use: "rename", + Short: "Rename an organization", + Run: func(cmd *cobra.Command, args []string) { + slug, _ := cmd.Flags().GetString("slug") + newName, _ := cmd.Flags().GetString("name") + newSlug, _ := cmd.Flags().GetString("new-slug") + orgRename(socketPath, slug, newName, newSlug) + }, + } + orgCmd.AddCommand(orgRenameCmd) + orgRenameCmd.Flags().String("slug", "", "current slug") + orgRenameCmd.Flags().String("name", "", "new display name") + orgRenameCmd.Flags().String("new-slug", "", "new slug") + _ = orgRenameCmd.MarkFlagRequired("slug") + _ = orgRenameCmd.MarkFlagRequired("name") + _ = orgRenameCmd.MarkFlagRequired("new-slug") + + orgListCmd := &cobra.Command{ + Use: "list", + Short: "List organizations", + Run: func(cmd *cobra.Command, args []string) { + email, _ := cmd.Flags().GetString("user") + orgList(socketPath, email) + }, + } + orgCmd.AddCommand(orgListCmd) + orgListCmd.Flags().String("user", "", "filter by user email (optional)") + + orgMemberAddCmd := &cobra.Command{ + Use: "add-member", + Short: "Add a member to an organization", + Run: func(cmd *cobra.Command, args []string) { + org, _ := cmd.Flags().GetString("org") + email, _ := cmd.Flags().GetString("email") + role, _ := cmd.Flags().GetString("role") + orgMemberAdd(socketPath, org, email, role) + }, + } + orgCmd.AddCommand(orgMemberAddCmd) + orgMemberAddCmd.Flags().String("org", "", "org slug") + orgMemberAddCmd.Flags().String("email", "", "user email") + orgMemberAddCmd.Flags().String("role", "member", "role (owner, admin, member)") + _ = orgMemberAddCmd.MarkFlagRequired("org") + _ = orgMemberAddCmd.MarkFlagRequired("email") + + orgMemberRemoveCmd := &cobra.Command{ + Use: "remove-member", + Short: "Remove a member from an organization", + Run: func(cmd *cobra.Command, args []string) { + org, _ := cmd.Flags().GetString("org") + email, _ := cmd.Flags().GetString("email") + orgMemberRemove(socketPath, org, email) + }, + } + orgCmd.AddCommand(orgMemberRemoveCmd) + orgMemberRemoveCmd.Flags().String("org", "", "org slug") + orgMemberRemoveCmd.Flags().String("email", "", "user email") + _ = orgMemberRemoveCmd.MarkFlagRequired("org") + _ = orgMemberRemoveCmd.MarkFlagRequired("email") + + orgSetRoleCmd := &cobra.Command{ + Use: "set-role", + Short: "Change a member's role", + Run: func(cmd *cobra.Command, args []string) { + org, _ := cmd.Flags().GetString("org") + email, _ := cmd.Flags().GetString("email") + role, _ := cmd.Flags().GetString("role") + orgMemberSetRole(socketPath, org, email, role) + }, + } + orgCmd.AddCommand(orgSetRoleCmd) + orgSetRoleCmd.Flags().String("org", "", "org slug") + orgSetRoleCmd.Flags().String("email", "", "user email") + orgSetRoleCmd.Flags().String("role", "", "new role (owner, admin, member)") + _ = orgSetRoleCmd.MarkFlagRequired("org") + _ = orgSetRoleCmd.MarkFlagRequired("email") + _ = orgSetRoleCmd.MarkFlagRequired("role") + + orgMemberListCmd := &cobra.Command{ + Use: "list-members", + Short: "List members of an organization", + Run: func(cmd *cobra.Command, args []string) { + org, _ := cmd.Flags().GetString("org") + orgMemberList(socketPath, org) + }, + } + orgCmd.AddCommand(orgMemberListCmd) + orgMemberListCmd.Flags().String("org", "", "org slug") + _ = orgMemberListCmd.MarkFlagRequired("org") + if err := root.Execute(); err != nil { os.Exit(1) } @@ -237,9 +366,9 @@ func daemon(configFile, socketFlag string) { srv.Close() - adminSrv.Shutdown(context.Background()) wwwSrv.Shutdown(context.Background()) grpcSrv.Shutdown(context.Background()) + adminSrv.Shutdown(context.Background()) } // listeners returns gRPC, web, and admin listeners. @@ -273,6 +402,7 @@ func listeners(cfg *config) (grpcLn, webLn, adminLn net.Listener, err error) { } }() + _ = os.Remove(cfg.AdminSocket) if adminLn, err = net.Listen("unix", cfg.AdminSocket); err != nil { return nil, nil, nil, err } @@ -282,6 +412,24 @@ func listeners(cfg *config) (grpcLn, webLn, adminLn net.Listener, err error) { } }() + grp, err := user.LookupGroup("workerd") + if err != nil { + return nil, nil, nil, fmt.Errorf("lookup group workerd: %w", err) + } + + gid, err := strconv.Atoi(grp.Gid) + if err != nil { + return nil, nil, nil, fmt.Errorf("parse gid: %w", err) + } + + if err = os.Chown(cfg.AdminSocket, 0, gid); err != nil { + return nil, nil, nil, fmt.Errorf("chown admin socket: %w", err) + } + + if err = os.Chmod(cfg.AdminSocket, 0660); err != nil { + return nil, nil, nil, fmt.Errorf("chmod admin socket: %w", err) + } + return grpcLn, webLn, adminLn, nil } @@ -302,7 +450,7 @@ func adminClient(socketPath string) pbconnect.AdminClient { } func userCreate(socketPath, email, password string) { - resp, err := adminClient(socketPath).CreateUser(context.Background(), connect.NewRequest(&pb.CreateUserRequest{ + resp, err := adminClient(socketPath).UserCreate(context.Background(), connect.NewRequest(&pb.UserCreateRequest{ Email: email, Password: password, })) @@ -314,7 +462,7 @@ func userCreate(socketPath, email, password string) { } func userSetPassword(socketPath, email, password string) { - _, err := adminClient(socketPath).SetPassword(context.Background(), connect.NewRequest(&pb.SetPasswordRequest{ + _, err := adminClient(socketPath).UserSetPassword(context.Background(), connect.NewRequest(&pb.UserSetPasswordRequest{ Email: email, Password: password, })) @@ -326,7 +474,7 @@ func userSetPassword(socketPath, email, password string) { } func userDelete(socketPath, email string) { - _, err := adminClient(socketPath).DeleteUser(context.Background(), connect.NewRequest(&pb.DeleteUserRequest{ + _, err := adminClient(socketPath).UserDelete(context.Background(), connect.NewRequest(&pb.UserDeleteRequest{ Email: email, })) if err != nil { @@ -384,3 +532,111 @@ func workerList(socketPath string) { fmt.Printf("%s\t%s\t%s\n", w.Id, base64.StdEncoding.EncodeToString(w.PublicKey), created) } } + +func orgCreate(socketPath, name, slug string, public bool, ownerEmail string) { + resp, err := adminClient(socketPath).OrgCreate(context.Background(), connect.NewRequest(&pb.OrgCreateRequest{ + Name: name, + Slug: slug, + Public: public, + OwnerEmail: ownerEmail, + })) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + fmt.Println(resp.Msg.Id) +} + +func orgDelete(socketPath, slug string) { + _, err := adminClient(socketPath).OrgDelete(context.Background(), connect.NewRequest(&pb.OrgDeleteRequest{ + Slug: slug, + })) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + fmt.Println("ok") +} + +func orgRename(socketPath, slug, newName, newSlug string) { + _, err := adminClient(socketPath).OrgRename(context.Background(), connect.NewRequest(&pb.OrgRenameRequest{ + Slug: slug, + NewName: newName, + NewSlug: newSlug, + })) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + fmt.Println("ok") +} + +func orgList(socketPath, email string) { + req := &pb.OrgListRequest{} + if email != "" { + req.UserEmail = &email + } + resp, err := adminClient(socketPath).OrgList(context.Background(), connect.NewRequest(req)) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + for _, o := range resp.Msg.Organizations { + visibility := "private" + if o.Public { + visibility = "public" + } + fmt.Printf("%s\t%s\t%s\t%s\n", o.Slug, o.Name, visibility, o.CreatedAt.AsTime().Format(time.DateOnly)) + } +} + +func orgMemberAdd(socketPath, orgSlug, email, role string) { + _, err := adminClient(socketPath).OrgMemberAdd(context.Background(), connect.NewRequest(&pb.OrgMemberAddRequest{ + OrgSlug: orgSlug, + Email: email, + Role: role, + })) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + fmt.Println("ok") +} + +func orgMemberRemove(socketPath, orgSlug, email string) { + _, err := adminClient(socketPath).OrgMemberRemove(context.Background(), connect.NewRequest(&pb.OrgMemberRemoveRequest{ + OrgSlug: orgSlug, + Email: email, + })) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + fmt.Println("ok") +} + +func orgMemberSetRole(socketPath, orgSlug, email, role string) { + _, err := adminClient(socketPath).OrgMemberSetRole(context.Background(), connect.NewRequest(&pb.OrgMemberSetRoleRequest{ + OrgSlug: orgSlug, + Email: email, + Role: role, + })) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + fmt.Println("ok") +} + +func orgMemberList(socketPath, orgSlug string) { + resp, err := adminClient(socketPath).OrgMemberList(context.Background(), connect.NewRequest(&pb.OrgMemberListRequest{ + OrgSlug: orgSlug, + })) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + for _, m := range resp.Msg.Members { + fmt.Printf("%s\t%s\t%s\n", m.Email, m.Role, m.JoinedAt.AsTime().Format(time.DateOnly)) + } +} diff --git a/cmd/mirumd/server_admin.go b/cmd/mirumd/server_admin.go index 673e9f6..2962887 100644 --- a/cmd/mirumd/server_admin.go +++ b/cmd/mirumd/server_admin.go @@ -12,6 +12,7 @@ import ( "connectrpc.com/connect" "google.golang.org/protobuf/types/known/timestamppb" + "dimidiumlabs/mirum/internal/database" "dimidiumlabs/mirum/internal/protocol/pb" "dimidiumlabs/mirum/internal/protocol/pb/pbconnect" ) @@ -29,26 +30,26 @@ type adminService struct { srv *server } -func (a *adminService) CreateUser(ctx context.Context, req *connect.Request[pb.CreateUserRequest]) (*connect.Response[pb.CreateUserResponse], error) { - id, err := a.srv.db.CreateUser(ctx, req.Msg.Email, req.Msg.Password, []byte(a.srv.cfg.Pepper)) +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)) if err != nil { return nil, err } - return connect.NewResponse(&pb.CreateUserResponse{Id: id}), nil + return connect.NewResponse(&pb.UserCreateResponse{Id: id}), nil } -func (a *adminService) SetPassword(ctx context.Context, req *connect.Request[pb.SetPasswordRequest]) (*connect.Response[pb.SetPasswordResponse], error) { - if err := a.srv.db.SetPassword(ctx, req.Msg.Email, req.Msg.Password, []byte(a.srv.cfg.Pepper)); err != nil { +func (a *adminService) UserSetPassword(ctx context.Context, req *connect.Request[pb.UserSetPasswordRequest]) (*connect.Response[pb.UserSetPasswordResponse], error) { + if err := a.srv.db.UserSetPassword(ctx, req.Msg.Email, req.Msg.Password, []byte(a.srv.cfg.Pepper)); err != nil { return nil, err } - return connect.NewResponse(&pb.SetPasswordResponse{}), nil + return connect.NewResponse(&pb.UserSetPasswordResponse{}), nil } -func (a *adminService) DeleteUser(ctx context.Context, req *connect.Request[pb.DeleteUserRequest]) (*connect.Response[pb.DeleteUserResponse], error) { - if err := a.srv.db.DeleteUser(ctx, req.Msg.Email); err != 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, req.Msg.Email); err != nil { return nil, err } - return connect.NewResponse(&pb.DeleteUserResponse{}), nil + return connect.NewResponse(&pb.UserDeleteResponse{}), nil } func (a *adminService) WorkerAdd(ctx context.Context, req *connect.Request[pb.WorkerAddRequest]) (*connect.Response[pb.WorkerAddResponse], error) { @@ -63,7 +64,7 @@ func (a *adminService) WorkerAdd(ctx context.Context, req *connect.Request[pb.Wo } func (a *adminService) WorkerRevoke(ctx context.Context, req *connect.Request[pb.WorkerRevokeRequest]) (*connect.Response[pb.WorkerRevokeResponse], error) { - if err := a.srv.db.RevokeWorker(ctx, req.Msg.Id); err != nil { + if err := a.srv.db.WorkerRevoke(ctx, req.Msg.Id); err != nil { return nil, err } return connect.NewResponse(&pb.WorkerRevokeResponse{}), nil @@ -84,3 +85,88 @@ func (a *adminService) WorkerList(ctx context.Context, req *connect.Request[pb.W } return connect.NewResponse(&pb.WorkerListResponse{Workers: pbWorkers}), nil } + +func (a *adminService) OrgCreate(ctx context.Context, req *connect.Request[pb.OrgCreateRequest]) (*connect.Response[pb.OrgCreateResponse], error) { + id, err := a.srv.db.CreateOrganization(ctx, req.Msg.Name, req.Msg.Slug, req.Msg.Public, req.Msg.OwnerEmail) + if err != nil { + return nil, err + } + return connect.NewResponse(&pb.OrgCreateResponse{Id: id}), 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, req.Msg.Slug); err != nil { + return nil, err + } + return connect.NewResponse(&pb.OrgDeleteResponse{}), nil +} + +func (a *adminService) OrgRename(ctx context.Context, req *connect.Request[pb.OrgRenameRequest]) (*connect.Response[pb.OrgRenameResponse], error) { + if err := a.srv.db.RenameOrganization(ctx, req.Msg.Slug, req.Msg.NewName, req.Msg.NewSlug); err != nil { + return nil, err + } + return connect.NewResponse(&pb.OrgRenameResponse{}), nil +} + +func (a *adminService) OrgList(ctx context.Context, req *connect.Request[pb.OrgListRequest]) (*connect.Response[pb.OrgListResponse], error) { + var ( + orgs []database.Organization + err error + ) + if req.Msg.UserEmail != nil { + orgs, err = a.srv.db.ListUserOrganizations(ctx, *req.Msg.UserEmail) + } else { + orgs, err = a.srv.db.ListAllOrganizations(ctx) + } + if err != nil { + return nil, err + } + pbOrgs := make([]*pb.Org, len(orgs)) + for i, o := range orgs { + pbOrgs[i] = &pb.Org{ + Id: o.ID, + Name: o.Name, + Slug: o.Slug, + Public: o.Public, + CreatedAt: timestamppb.New(o.CreatedAt), + } + } + return connect.NewResponse(&pb.OrgListResponse{Organizations: pbOrgs}), nil +} + +func (a *adminService) OrgMemberAdd(ctx context.Context, req *connect.Request[pb.OrgMemberAddRequest]) (*connect.Response[pb.OrgMemberAddResponse], error) { + if err := a.srv.db.AddOrgMember(ctx, req.Msg.OrgSlug, req.Msg.Email, req.Msg.Role); err != nil { + return nil, err + } + return connect.NewResponse(&pb.OrgMemberAddResponse{}), 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, req.Msg.OrgSlug, req.Msg.Email); err != nil { + return nil, err + } + return connect.NewResponse(&pb.OrgMemberRemoveResponse{}), nil +} + +func (a *adminService) OrgMemberSetRole(ctx context.Context, req *connect.Request[pb.OrgMemberSetRoleRequest]) (*connect.Response[pb.OrgMemberSetRoleResponse], error) { + if err := a.srv.db.ChangeOrgMemberRole(ctx, req.Msg.OrgSlug, req.Msg.Email, req.Msg.Role); err != nil { + return nil, err + } + return connect.NewResponse(&pb.OrgMemberSetRoleResponse{}), nil +} + +func (a *adminService) OrgMemberList(ctx context.Context, req *connect.Request[pb.OrgMemberListRequest]) (*connect.Response[pb.OrgMemberListResponse], error) { + members, err := a.srv.db.ListOrgMembers(ctx, req.Msg.OrgSlug) + if err != nil { + return nil, err + } + pbMembers := make([]*pb.OrgMemberInfo, len(members)) + for i, m := range members { + pbMembers[i] = &pb.OrgMemberInfo{ + Email: m.User.Email, + Role: m.Role, + JoinedAt: timestamppb.New(m.JoinedAt), + } + } + return connect.NewResponse(&pb.OrgMemberListResponse{Members: pbMembers}), nil +} diff --git a/cmd/mirumd/server_web.go b/cmd/mirumd/server_web.go index 5c922c6..608cf61 100644 --- a/cmd/mirumd/server_web.go +++ b/cmd/mirumd/server_web.go @@ -83,7 +83,7 @@ func wwwRoutes(srv *server) http.Handler { 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.GetSession(r.Context(), c.Value); err == nil { + if sess, err := srv.db.UserGetSession(r.Context(), c.Value); err == nil { data.Email = sess.Email data.CSRF = csrfToken(w, r) } @@ -132,7 +132,7 @@ func wwwRoutes(srv *server) http.Handler { email := r.FormValue("email") password := r.FormValue("password") - userID, err := srv.db.VerifyPassword(r.Context(), email, password, []byte(srv.cfg.Pepper)) + 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{ @@ -142,7 +142,7 @@ func wwwRoutes(srv *server) http.Handler { return } - token, err := srv.db.CreateSession(r.Context(), userID) + token, err := srv.db.UserCreateSession(r.Context(), userID) if err != nil { http.Error(w, "internal error", http.StatusInternalServerError) return @@ -166,7 +166,7 @@ func wwwRoutes(srv *server) http.Handler { return } if c, err := r.Cookie("session"); err == nil { - srv.db.DeleteSession(r.Context(), c.Value) + srv.db.UserDeleteSession(r.Context(), c.Value) } clearCookie(w, "session") clearCookie(w, "csrf") diff --git a/cmd/mirumw/client.go b/cmd/mirumw/client.go index 999d34b..3496cac 100644 --- a/cmd/mirumw/client.go +++ b/cmd/mirumw/client.go @@ -222,5 +222,6 @@ func (c *client) handshake(ctx context.Context) error { v := sr.GetServerVersion() slog.Info("handshake ok", "server_version", fmt.Sprintf("%d.%d.%d", v.GetMajor(), v.GetMinor(), v.GetPatch())) + return nil } diff --git a/internal/database/database.go b/internal/database/database.go index 1e84770..a6cbaed 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -5,42 +5,17 @@ package database import ( "context" - "crypto/hmac" - "crypto/rand" - "crypto/sha256" - "crypto/subtle" - "encoding/base64" "errors" - "fmt" - "strings" - "time" - "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" "github.com/jackc/tern/v2/migrate" - "golang.org/x/crypto/argon2" -) - -const ( - argonMemory = 64 * 1024 // 64 MB - argonTime = 3 - argonThreads = 2 - argonKeyLen = 32 - saltLen = 16 ) var ( - errOpen = errors.New("database: failed to open") - errPing = errors.New("database: failed to ping") - errAcquire = errors.New("database: failed to acquire connection") - errMigrate = errors.New("database: failed to create migrator") - errCreateUser = errors.New("database: failed to create user") - errSetPassword = errors.New("database: failed to set password") - errDeleteUser = errors.New("database: failed to delete user") - errInvalidCreds = errors.New("invalid credentials") - errAddWorker = errors.New("database: failed to add worker") - errRevokeWorker = errors.New("database: failed to revoke worker") - ErrWorkerNotFound = errors.New("database: worker not found") + ErrOpen = errors.New("database: failed to open") + ErrPing = errors.New("database: failed to ping") + ErrAcquire = errors.New("database: failed to acquire connection") + ErrMigrate = errors.New("database: failed to create migrator") ) // DB wraps a pgx connection pool. @@ -52,12 +27,14 @@ type DB struct { func Open(ctx context.Context, dsn string) (*DB, error) { pool, err := pgxpool.New(ctx, dsn) if err != nil { - return nil, errors.Join(errOpen, err) + return nil, errors.Join(ErrOpen, err) } + if err := pool.Ping(ctx); err != nil { pool.Close() - return nil, errors.Join(errPing, err) + return nil, errors.Join(ErrPing, err) } + return &DB{Pool: pool}, nil } @@ -70,13 +47,13 @@ func (db *DB) Close() { func (db *DB) Migrate(ctx context.Context) error { conn, err := db.Pool.Acquire(ctx) if err != nil { - return errors.Join(errAcquire, err) + return errors.Join(ErrAcquire, err) } defer conn.Release() migrator, err := migrate.NewMigrator(ctx, conn.Conn(), "schema_version") if err != nil { - return errors.Join(errMigrate, err) + return errors.Join(ErrMigrate, err) } migrator.AppendMigration("create_users", @@ -111,298 +88,29 @@ func (db *DB) Migrate(ctx context.Context) error { `DROP TABLE workers`, ) - return migrator.Migrate(ctx) -} - -const SessionTTL = 14 * 24 * time.Hour // 14 days - -// hashToken returns the hex-encoded SHA-256 of a session token. -func hashToken(token string) string { - h := sha256.Sum256([]byte(token)) - return fmt.Sprintf("%x", h) -} - -// CreateSession generates a random token, stores its hash, and returns the token. -func (db *DB) CreateSession(ctx context.Context, userID string) (string, error) { - b := make([]byte, 32) - if _, err := rand.Read(b); err != nil { - return "", err - } - token := base64.RawURLEncoding.EncodeToString(b) - - _, err := db.Pool.Exec(ctx, - `INSERT INTO sessions (token, user_id, expires_at) VALUES ($1, $2, now() + $3)`, - hashToken(token), userID, SessionTTL, - ) - if err != nil { - return "", err - } - return token, nil -} - -// Session holds info about an authenticated session. -type Session struct { - UserID string - Email string -} - -// GetSession returns session info for a valid, non-expired session. -// It extends the session expiry only when less than half the TTL remains, -// avoiding a write on every request. -func (db *DB) GetSession(ctx context.Context, token string) (*Session, error) { - h := hashToken(token) - var s Session - var expiresAt time.Time - err := db.Pool.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`, - h, - ).Scan(&s.UserID, &s.Email, &expiresAt) - if err != nil { - return nil, err - } - - if time.Until(expiresAt) < SessionTTL/2 { - db.Pool.Exec(ctx, - `UPDATE sessions SET expires_at = now() + $2 WHERE token = $1`, - h, SessionTTL, - ) - } - - return &s, nil -} - -// DeleteSession removes a session (logout). -func (db *DB) DeleteSession(ctx context.Context, token string) error { - _, err := db.Pool.Exec(ctx, `DELETE FROM sessions WHERE token = $1`, hashToken(token)) - return err -} - -// PurgeExpiredSessions deletes all expired sessions. -func (db *DB) PurgeExpiredSessions(ctx context.Context) error { - _, err := db.Pool.Exec(ctx, `DELETE FROM sessions WHERE expires_at < now()`) - return err -} - -// CreateUser 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) CreateUser(ctx context.Context, email, password string, pepper []byte) (string, error) { - hash, err := hashPassword(password, pepper) - if err != nil { - return "", err - } - - var id string - err = db.Pool.QueryRow(ctx, - `INSERT INTO users (email, password) VALUES ($1, $2) RETURNING id`, - email, hash, - ).Scan(&id) - if err != nil { - return "", errors.Join(errCreateUser, err) - } - - return id, nil -} - -// SetPassword updates the password and invalidates all existing sessions. -func (db *DB) SetPassword(ctx context.Context, email, password string, pepper []byte) error { - hash, err := hashPassword(password, pepper) - if err != nil { - return err - } - - tx, err := db.Pool.Begin(ctx) - if err != nil { - return errors.Join(errSetPassword, err) - } - defer tx.Rollback(ctx) - - var id string - err = tx.QueryRow(ctx, - `UPDATE users SET password = $1 WHERE email = $2 AND deleted_at IS NULL RETURNING id`, - hash, email, - ).Scan(&id) - if err != nil { - return errors.Join(errSetPassword, err) - } - - _, err = tx.Exec(ctx, `DELETE FROM sessions WHERE user_id = $1`, id) - if err != nil { - return errors.Join(errSetPassword, err) - } - - return tx.Commit(ctx) -} - -// DeleteUser clears all fields but keeps the row to preserve the id. -// Also deletes all sessions for that user in a single transaction. -func (db *DB) DeleteUser(ctx context.Context, email string) error { - tx, err := db.Pool.Begin(ctx) - if err != nil { - return errors.Join(errDeleteUser, err) - } - defer tx.Rollback(ctx) - - var id string - err = tx.QueryRow(ctx, - `UPDATE users SET email = id::text, password = '', deleted_at = now() WHERE email = $1 RETURNING id`, - email, - ).Scan(&id) - if err != nil { - return errors.Join(errDeleteUser, err) - } - - _, err = tx.Exec(ctx, `DELETE FROM sessions WHERE user_id = $1`, id) - if err != nil { - return errors.Join(errDeleteUser, err) - } - - return tx.Commit(ctx) -} - -// VerifyPassword checks credentials and returns the user ID. -func (db *DB) VerifyPassword(ctx context.Context, email, password string, pepper []byte) (string, error) { - var id, hash string - err := db.Pool.QueryRow(ctx, - `SELECT id, password FROM users WHERE email = $1 AND deleted_at IS NULL`, - email, - ).Scan(&id, &hash) - if err != nil { - return "", errInvalidCreds - } - - if !verifyHash(password, hash, pepper) { - return "", errInvalidCreds - } - - return id, nil -} - -// verifyHash parses a PHC-format argon2id string and compares. -// Format: $argon2id$v=19$m=65536,t=3,p=2$$ -func verifyHash(password, encoded string, pepper []byte) bool { - // $argon2id$v=19$m=65536,t=3,p=2$salt$key → 6 parts - parts := strings.Split(encoded, "$") - if len(parts) != 6 || parts[1] != "argon2id" { - return false - } - - var memory, time uint32 - var threads uint8 - if _, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &time, &threads); err != nil { - return false - } - - salt, err := base64.RawStdEncoding.DecodeString(parts[4]) - if err != nil { - return false - } - expectedKey, err := base64.RawStdEncoding.DecodeString(parts[5]) - if err != nil { - return false - } - - mac := hmac.New(sha256.New, pepper) - mac.Write([]byte(password)) - peppered := mac.Sum(nil) - - key := argon2.IDKey(peppered, salt, time, memory, threads, uint32(len(expectedKey))) - - return subtle.ConstantTimeCompare(key, expectedKey) == 1 -} - -// hashPassword produces a PHC-format string: -// $argon2id$v=19$m=65536,t=3,p=2$$ -func hashPassword(password string, pepper []byte) (string, error) { - salt := make([]byte, saltLen) - if _, err := rand.Read(salt); err != nil { - return "", err - } - - // Apply pepper: HMAC-SHA256(pepper, password) - mac := hmac.New(sha256.New, pepper) - mac.Write([]byte(password)) - peppered := mac.Sum(nil) - - key := argon2.IDKey(peppered, salt, argonTime, argonMemory, argonThreads, argonKeyLen) - - return fmt.Sprintf("$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s", - argon2.Version, - argonMemory, argonTime, argonThreads, - base64.RawStdEncoding.EncodeToString(salt), - base64.RawStdEncoding.EncodeToString(key), - ), nil -} - -// Worker holds info about a registered worker. -type Worker struct { - ID string - PublicKey []byte - CreatedAt time.Time -} - -// LookupWorker finds an active worker by its ed25519 public key. -func (db *DB) LookupWorker(ctx context.Context, publicKey []byte) (*Worker, error) { - var w Worker - err := db.Pool.QueryRow(ctx, - `SELECT id, public_key, created_at FROM workers WHERE public_key = $1 AND revoked_at IS NULL`, - publicKey, - ).Scan(&w.ID, &w.PublicKey, &w.CreatedAt) - if err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return nil, ErrWorkerNotFound - } - return nil, err - } - return &w, nil -} - -// AddWorker registers a new worker with the given public key. -func (db *DB) AddWorker(ctx context.Context, publicKey []byte) (string, error) { - var id string - err := db.Pool.QueryRow(ctx, - `INSERT INTO workers (public_key) VALUES ($1) RETURNING id`, - publicKey, - ).Scan(&id) - if err != nil { - return "", errors.Join(errAddWorker, err) - } - return id, nil -} - -// RevokeWorker soft-deletes a worker by ID. -func (db *DB) RevokeWorker(ctx context.Context, id string) error { - tag, err := db.Pool.Exec(ctx, - `UPDATE workers SET revoked_at = now() WHERE id = $1 AND revoked_at IS NULL`, - id, + 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`, ) - if err != nil { - return errors.Join(errRevokeWorker, err) - } - if tag.RowsAffected() == 0 { - return ErrWorkerNotFound - } - return nil -} -// ListWorkers returns all active (non-revoked) workers. -func (db *DB) ListWorkers(ctx context.Context) ([]Worker, error) { - rows, err := db.Pool.Query(ctx, - `SELECT id, public_key, created_at FROM workers WHERE revoked_at IS NULL ORDER BY created_at`, + 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`, ) - if err != nil { - return nil, err - } - defer rows.Close() - var workers []Worker - for rows.Next() { - var w Worker - if err := rows.Scan(&w.ID, &w.PublicKey, &w.CreatedAt); err != nil { - return nil, err - } - workers = append(workers, w) - } - return workers, rows.Err() + return migrator.Migrate(ctx) } diff --git a/internal/database/organization.go b/internal/database/organization.go new file mode 100644 index 0000000..53f60e7 --- /dev/null +++ b/internal/database/organization.go @@ -0,0 +1,314 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package database + +import ( + "context" + "errors" + "strings" + "time" + + "github.com/jackc/pgx/v5" +) + +var ( + ErrOrgNotFound = errors.New("database: organization not found") + ErrAlreadyMember = errors.New("database: user is already a member") + ErrNotMember = errors.New("database: user is not a member") + ErrLastOwner = errors.New("database: cannot remove or demote the last owner") + ErrSlugTaken = errors.New("database: slug already taken") +) + +// Organization holds info about an organization. +type Organization struct { + ID string + Name string + Slug string + Public bool + CreatedAt time.Time +} + +// OrgMember pairs a user with their role in an organization. +type OrgMember struct { + User User + Role string + JoinedAt time.Time +} + +// CreateOrganization creates an org and adds the owner as the first member. +func (db *DB) CreateOrganization(ctx context.Context, name, slug string, public bool, ownerEmail string) (string, error) { + tx, err := db.Pool.Begin(ctx) + if err != nil { + return "", err + } + defer tx.Rollback(ctx) + + var userID string + if err := tx.QueryRow(ctx, + `SELECT id FROM users WHERE email = $1 AND deleted_at IS NULL`, ownerEmail, + ).Scan(&userID); err != nil { + return "", ErrUserNotFound + } + + var orgID string + if err := tx.QueryRow(ctx, + `INSERT INTO organizations (name, slug, public) VALUES ($1, $2, $3) RETURNING id`, name, slug, public, + ).Scan(&orgID); err != nil { + return "", err + } + + if _, err := tx.Exec(ctx, + `INSERT INTO org_members (org_id, user_id, role) VALUES ($1, $2, 'owner')`, orgID, userID, + ); err != nil { + return "", err + } + + return orgID, tx.Commit(ctx) +} + +// DeleteOrganization soft-deletes an org and removes all members. +func (db *DB) DeleteOrganization(ctx context.Context, slug string) error { + tx, err := db.Pool.Begin(ctx) + if err != nil { + return err + } + defer tx.Rollback(ctx) + + var orgID string + if err := tx.QueryRow(ctx, + `UPDATE organizations SET slug = id::text, deleted_at = now() + WHERE slug = $1 AND deleted_at IS NULL RETURNING id`, slug, + ).Scan(&orgID); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return ErrOrgNotFound + } + return err + } + + if _, err := tx.Exec(ctx, `DELETE FROM org_members WHERE org_id = $1`, orgID); err != nil { + return err + } + + return tx.Commit(ctx) +} + +// RenameOrganization updates the name and/or slug. +func (db *DB) RenameOrganization(ctx context.Context, currentSlug, newName, newSlug string) error { + tag, err := db.Pool.Exec(ctx, + `UPDATE organizations SET name = $1, slug = $2 + WHERE slug = $3 AND deleted_at IS NULL`, + newName, newSlug, currentSlug, + ) + if err != nil { + if strings.Contains(err.Error(), "unique") { + return ErrSlugTaken + } + return err + } + if tag.RowsAffected() == 0 { + return ErrOrgNotFound + } + return nil +} + +// GetOrganizationBySlug returns an org by its slug. +func (db *DB) GetOrganizationBySlug(ctx context.Context, slug string) (*Organization, error) { + var o Organization + err := db.Pool.QueryRow(ctx, + `SELECT id, name, slug, public, created_at FROM organizations WHERE slug = $1 AND deleted_at IS NULL`, + slug, + ).Scan(&o.ID, &o.Name, &o.Slug, &o.Public, &o.CreatedAt) + if err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return nil, ErrOrgNotFound + } + return nil, err + } + return &o, nil +} + +// ListAllOrganizations returns all active orgs. +func (db *DB) ListAllOrganizations(ctx context.Context) ([]Organization, error) { + rows, err := db.Pool.Query(ctx, + `SELECT id, name, slug, public, created_at FROM organizations WHERE deleted_at IS NULL ORDER BY name`, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + var orgs []Organization + for rows.Next() { + var o Organization + if err := rows.Scan(&o.ID, &o.Name, &o.Slug, &o.Public, &o.CreatedAt); err != nil { + return nil, err + } + orgs = append(orgs, o) + } + return orgs, rows.Err() +} + +// ListUserOrganizations returns all orgs a user belongs to. +func (db *DB) ListUserOrganizations(ctx context.Context, email string) ([]Organization, error) { + rows, err := db.Pool.Query(ctx, + `SELECT o.id, o.name, o.slug, o.public, o.created_at + FROM organizations o + JOIN org_members m ON m.org_id = o.id + JOIN users u ON u.id = m.user_id + WHERE u.email = $1 AND u.deleted_at IS NULL AND o.deleted_at IS NULL + ORDER BY o.name`, + email, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + var orgs []Organization + for rows.Next() { + var o Organization + if err := rows.Scan(&o.ID, &o.Name, &o.Slug, &o.Public, &o.CreatedAt); err != nil { + return nil, err + } + orgs = append(orgs, o) + } + return orgs, rows.Err() +} + +// ListOrgMembers returns all members of an org. +func (db *DB) ListOrgMembers(ctx context.Context, orgSlug string) ([]OrgMember, error) { + rows, err := db.Pool.Query(ctx, + `SELECT u.id, u.email, u.created_at, m.role, m.created_at + FROM org_members m + JOIN users u ON u.id = m.user_id + JOIN organizations o ON o.id = m.org_id + WHERE o.slug = $1 AND o.deleted_at IS NULL AND u.deleted_at IS NULL + ORDER BY m.created_at`, + orgSlug, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + var members []OrgMember + for rows.Next() { + var m OrgMember + if err := rows.Scan(&m.User.ID, &m.User.Email, &m.User.CreatedAt, &m.Role, &m.JoinedAt); err != nil { + return nil, err + } + members = append(members, m) + } + return members, rows.Err() +} + +// AddOrgMember adds a user to an org with the given role. +func (db *DB) AddOrgMember(ctx context.Context, orgSlug, email, role string) error { + tag, err := db.Pool.Exec(ctx, + `INSERT INTO org_members (org_id, user_id, role) + SELECT o.id, u.id, $3 + FROM organizations o, users u + WHERE o.slug = $1 AND u.email = $2 AND o.deleted_at IS NULL AND u.deleted_at IS NULL`, + orgSlug, email, role, + ) + if err != nil { + if strings.Contains(err.Error(), "duplicate key") { + return ErrAlreadyMember + } + return err + } + if tag.RowsAffected() == 0 { + return ErrOrgNotFound + } + return nil +} + +// RemoveOrgMember removes a user from an org. Fails if they are the last owner. +func (db *DB) RemoveOrgMember(ctx context.Context, orgSlug, email string) error { + tx, err := db.Pool.Begin(ctx) + if err != nil { + return err + } + defer tx.Rollback(ctx) + + var orgID, userID, role string + if err := tx.QueryRow(ctx, + `SELECT o.id, u.id, m.role + FROM org_members m + JOIN organizations o ON o.id = m.org_id + JOIN users u ON u.id = m.user_id + WHERE o.slug = $1 AND u.email = $2 AND o.deleted_at IS NULL`, + orgSlug, email, + ).Scan(&orgID, &userID, &role); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return ErrNotMember + } + return err + } + + if role == "owner" { + var ownerCount int + if err := tx.QueryRow(ctx, + `SELECT count(*) FROM org_members WHERE org_id = $1 AND role = 'owner'`, orgID, + ).Scan(&ownerCount); err != nil { + return err + } + if ownerCount <= 1 { + return ErrLastOwner + } + } + + if _, err := tx.Exec(ctx, + `DELETE FROM org_members WHERE org_id = $1 AND user_id = $2`, orgID, userID, + ); err != nil { + return err + } + + return tx.Commit(ctx) +} + +// ChangeOrgMemberRole changes a member's role. Fails if demoting the last owner. +func (db *DB) ChangeOrgMemberRole(ctx context.Context, orgSlug, email, newRole string) error { + tx, err := db.Pool.Begin(ctx) + if err != nil { + return err + } + defer tx.Rollback(ctx) + + var orgID, userID, currentRole string + if err := tx.QueryRow(ctx, + `SELECT o.id, u.id, m.role + FROM org_members m + JOIN organizations o ON o.id = m.org_id + JOIN users u ON u.id = m.user_id + WHERE o.slug = $1 AND u.email = $2 AND o.deleted_at IS NULL`, + orgSlug, email, + ).Scan(&orgID, &userID, ¤tRole); err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return ErrNotMember + } + return err + } + + if currentRole == "owner" && newRole != "owner" { + var ownerCount int + if err := tx.QueryRow(ctx, + `SELECT count(*) FROM org_members WHERE org_id = $1 AND role = 'owner'`, orgID, + ).Scan(&ownerCount); err != nil { + return err + } + if ownerCount <= 1 { + return ErrLastOwner + } + } + + if _, err := tx.Exec(ctx, + `UPDATE org_members SET role = $1 WHERE org_id = $2 AND user_id = $3`, + newRole, orgID, userID, + ); err != nil { + return err + } + + return tx.Commit(ctx) +} diff --git a/internal/database/user.go b/internal/database/user.go new file mode 100644 index 0000000..6cf5f6f --- /dev/null +++ b/internal/database/user.go @@ -0,0 +1,285 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package database + +import ( + "context" + "crypto/hmac" + "crypto/rand" + "crypto/sha256" + "crypto/subtle" + "encoding/base64" + "errors" + "fmt" + "strings" + "time" + + "github.com/jackc/pgx/v5" + "golang.org/x/crypto/argon2" +) + +var ( + ErrCreateUser = errors.New("database: failed to create user") + ErrSetPassword = errors.New("database: failed to set password") + ErrDeleteUser = errors.New("database: failed to delete user") + ErrInvalidCreds = errors.New("invalid credentials") + ErrSoleOwner = errors.New("database: user is the sole owner of an organization") + ErrUserNotFound = errors.New("database: user not found") +) + +const ( + saltLen = 16 + + argonTime = 3 + argonMemory = 64 * 1024 // 64 MB + argonKeyLen = 32 + argonThreads = 2 + + SessionTTL = 14 * 24 * time.Hour // 14 days +) + +// User holds info about a user. +type User struct { + ID string + Email string + CreatedAt time.Time +} + +// Session holds info about an authenticated session. +type Session struct { + UserID string + Email string +} + +// hashToken returns the hex-encoded SHA-256 of a session token. +func hashToken(token string) string { + h := sha256.Sum256([]byte(token)) + return fmt.Sprintf("%x", h) +} + +// verifyHash parses a PHC-format argon2id string and compares. +// Format: $argon2id$v=19$m=65536,t=3,p=2$$ +func verifyHash(password, encoded string, pepper []byte) bool { + // $argon2id$v=19$m=65536,t=3,p=2$salt$key → 6 parts + parts := strings.Split(encoded, "$") + if len(parts) != 6 || parts[1] != "argon2id" { + return false + } + + var memory, time uint32 + var threads uint8 + if _, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &time, &threads); err != nil { + return false + } + + salt, err := base64.RawStdEncoding.DecodeString(parts[4]) + if err != nil { + return false + } + expectedKey, err := base64.RawStdEncoding.DecodeString(parts[5]) + if err != nil { + return false + } + + mac := hmac.New(sha256.New, pepper) + mac.Write([]byte(password)) + peppered := mac.Sum(nil) + + key := argon2.IDKey(peppered, salt, time, memory, threads, uint32(len(expectedKey))) + + return subtle.ConstantTimeCompare(key, expectedKey) == 1 +} + +// hashPassword produces a PHC-format string: +// $argon2id$v=19$m=65536,t=3,p=2$$ +func hashPassword(password string, pepper []byte) (string, error) { + salt := make([]byte, saltLen) + if _, err := rand.Read(salt); err != nil { + return "", err + } + + // Apply pepper: HMAC-SHA256(pepper, password) + mac := hmac.New(sha256.New, pepper) + mac.Write([]byte(password)) + peppered := mac.Sum(nil) + + key := argon2.IDKey(peppered, salt, argonTime, argonMemory, argonThreads, argonKeyLen) + + return fmt.Sprintf("$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s", + argon2.Version, + argonMemory, argonTime, argonThreads, + base64.RawStdEncoding.EncodeToString(salt), + base64.RawStdEncoding.EncodeToString(key), + ), nil +} + +// UserCreate hashes the password with argon2id and inserts a new user. +// The pepper is a server-side secret not stored in the database. +func (db *DB) UserCreate(ctx context.Context, email, password string, pepper []byte) (string, error) { + hash, err := hashPassword(password, pepper) + if err != nil { + return "", err + } + + var id string + err = db.Pool.QueryRow(ctx, + `INSERT INTO users (email, password) VALUES ($1, $2) RETURNING id`, + email, hash, + ).Scan(&id) + if err != nil { + return "", errors.Join(ErrCreateUser, err) + } + + return id, nil +} + +// UserDelete clears all fields but keeps the row to preserve the id. +// Also deletes all sessions for that user in a single transaction. +func (db *DB) UserDelete(ctx context.Context, email string) error { + tx, err := db.Pool.Begin(ctx) + if err != nil { + return errors.Join(ErrDeleteUser, err) + } + defer tx.Rollback(ctx) + + var id string + err = tx.QueryRow(ctx, + `UPDATE users SET email = id::text, password = '', deleted_at = now() WHERE email = $1 RETURNING id`, + email, + ).Scan(&id) + if err != nil { + return errors.Join(ErrDeleteUser, err) + } + + var soloOwnedSlug string + err = tx.QueryRow(ctx, + `SELECT o.slug FROM org_members m + JOIN organizations o ON o.id = m.org_id + WHERE m.role = 'owner' AND o.deleted_at IS NULL + GROUP BY o.id, o.slug + HAVING count(*) = 1 AND bool_or(m.user_id = $1) + LIMIT 1`, id, + ).Scan(&soloOwnedSlug) + if err == nil { + return ErrSoleOwner + } + if !errors.Is(err, pgx.ErrNoRows) { + return errors.Join(ErrDeleteUser, err) + } + + if _, err = tx.Exec(ctx, `DELETE FROM org_members WHERE user_id = $1`, id); err != nil { + return errors.Join(ErrDeleteUser, err) + } + + if _, err = tx.Exec(ctx, `DELETE FROM sessions WHERE user_id = $1`, id); err != nil { + return errors.Join(ErrDeleteUser, err) + } + + return tx.Commit(ctx) +} + +// UserSetPassword updates the password and invalidates all existing sessions. +func (db *DB) UserSetPassword(ctx context.Context, email, password string, pepper []byte) error { + hash, err := hashPassword(password, pepper) + if err != nil { + return err + } + + tx, err := db.Pool.Begin(ctx) + if err != nil { + return errors.Join(ErrSetPassword, err) + } + defer tx.Rollback(ctx) + + var id string + err = tx.QueryRow(ctx, + `UPDATE users SET password = $1 WHERE email = $2 AND deleted_at IS NULL RETURNING id`, + hash, email, + ).Scan(&id) + if err != nil { + return errors.Join(ErrSetPassword, err) + } + + _, err = tx.Exec(ctx, `DELETE FROM sessions WHERE user_id = $1`, id) + if err != nil { + return errors.Join(ErrSetPassword, err) + } + + return tx.Commit(ctx) +} + +// UserVerifyPassword checks credentials and returns the user ID. +func (db *DB) UserVerifyPassword(ctx context.Context, email, password string, pepper []byte) (string, error) { + var id, hash string + err := db.Pool.QueryRow(ctx, + `SELECT id, password FROM users WHERE email = $1 AND deleted_at IS NULL`, + email, + ).Scan(&id, &hash) + if err != nil { + return "", ErrInvalidCreds + } + + if !verifyHash(password, hash, pepper) { + return "", ErrInvalidCreds + } + + return id, nil +} + +// UserGetSession returns session info for a valid, non-expired session. +// It extends the session expiry only when less than half the TTL remains, +// avoiding a write on every request. +func (db *DB) UserGetSession(ctx context.Context, token string) (*Session, error) { + h := hashToken(token) + var s Session + var expiresAt time.Time + err := db.Pool.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`, + h, + ).Scan(&s.UserID, &s.Email, &expiresAt) + if err != nil { + return nil, err + } + + if time.Until(expiresAt) < SessionTTL/2 { + db.Pool.Exec(ctx, + `UPDATE sessions SET expires_at = now() + $2 WHERE token = $1`, + h, SessionTTL, + ) + } + + return &s, nil +} + +// UserCreateSession generates a random token, stores its hash, and returns the token. +func (db *DB) UserCreateSession(ctx context.Context, userID string) (string, error) { + b := make([]byte, 32) + if _, err := rand.Read(b); err != nil { + return "", err + } + token := base64.RawURLEncoding.EncodeToString(b) + + _, err := db.Pool.Exec(ctx, + `INSERT INTO sessions (token, user_id, expires_at) VALUES ($1, $2, now() + $3)`, + hashToken(token), userID, SessionTTL, + ) + if err != nil { + return "", err + } + return token, nil +} + +// 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 +} + +// PurgeExpiredSessions deletes all expired sessions. +func (db *DB) PurgeExpiredSessions(ctx context.Context) error { + _, err := db.Pool.Exec(ctx, `DELETE FROM sessions WHERE expires_at < now()`) + return err +} diff --git a/internal/database/worker.go b/internal/database/worker.go new file mode 100644 index 0000000..bdd4b0b --- /dev/null +++ b/internal/database/worker.go @@ -0,0 +1,91 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package database + +import ( + "context" + "errors" + "time" + + "github.com/jackc/pgx/v5" +) + +var ( + ErrAddWorker = errors.New("database: failed to add worker") + ErrRevokeWorker = errors.New("database: failed to revoke worker") + ErrWorkerNotFound = errors.New("database: worker not found") +) + +// Worker holds info about a registered worker. +type Worker struct { + ID string + PublicKey []byte + CreatedAt time.Time +} + +// ListWorkers returns all active (non-revoked) workers. +func (db *DB) ListWorkers(ctx context.Context) ([]Worker, error) { + rows, err := db.Pool.Query(ctx, + `SELECT id, public_key, created_at FROM workers WHERE revoked_at IS NULL ORDER BY created_at`, + ) + if err != nil { + return nil, err + } + defer rows.Close() + + var workers []Worker + for rows.Next() { + var w Worker + if err := rows.Scan(&w.ID, &w.PublicKey, &w.CreatedAt); err != nil { + return nil, err + } + workers = append(workers, w) + } + + return workers, rows.Err() +} + +// LookupWorker finds an active worker by its ed25519 public key. +func (db *DB) LookupWorker(ctx context.Context, publicKey []byte) (*Worker, error) { + var w Worker + err := db.Pool.QueryRow(ctx, + `SELECT id, public_key, created_at FROM workers WHERE public_key = $1 AND revoked_at IS NULL`, + publicKey, + ).Scan(&w.ID, &w.PublicKey, &w.CreatedAt) + if err != nil { + if errors.Is(err, pgx.ErrNoRows) { + return nil, ErrWorkerNotFound + } + return nil, err + } + return &w, nil +} + +// AddWorker registers a new worker with the given public key. +func (db *DB) AddWorker(ctx context.Context, publicKey []byte) (string, error) { + var id string + err := db.Pool.QueryRow(ctx, + `INSERT INTO workers (public_key) VALUES ($1) RETURNING id`, + publicKey, + ).Scan(&id) + if err != nil { + return "", errors.Join(ErrAddWorker, err) + } + return id, nil +} + +// WorkerRevoke soft-deletes a worker by ID. +func (db *DB) WorkerRevoke(ctx context.Context, id string) error { + tag, err := db.Pool.Exec(ctx, + `UPDATE workers SET revoked_at = now() WHERE id = $1 AND revoked_at IS NULL`, + id, + ) + if err != nil { + return errors.Join(ErrRevokeWorker, err) + } + if tag.RowsAffected() == 0 { + return ErrWorkerNotFound + } + return nil +} diff --git a/internal/protocol/handshake.go b/internal/protocol/handshake.go index f25cb0a..7f54a24 100644 --- a/internal/protocol/handshake.go +++ b/internal/protocol/handshake.go @@ -6,23 +6,34 @@ package protocol import ( "crypto/ed25519" "crypto/rand" + "crypto/x509" + "encoding/pem" "errors" + "fmt" + "os" "slices" "time" ) var ( - ErrInvalidPublicKey = errors.New("invalid public key size") - ErrInvalidNonce = errors.New("invalid nonce size") - ErrInvalidSignature = errors.New("invalid signature") + ErrKeyNotPEM = errors.New("key file does not contain a PEM block") + ErrKeyNotPKCS8 = errors.New("key file does not contain a PKCS8 private key") + ErrKeyNotPKIX = errors.New("key file does not contain a PKIX public key") + ErrKeyNotEd25519 = errors.New("key file does not contain an ed25519 key") + ErrClockSkew = errors.New("clock skew too large") + ErrInvalidNonce = errors.New("invalid nonce size") ErrServerRejected = errors.New("server rejected handshake") + ErrInvalidPublicKey = errors.New("invalid public key size") + ErrInvalidSignature = errors.New("invalid signature") ) const NonceSize = 32 -const EKMLabel = "mirum-handshake" -const EKMLength = 32 +const ( + EKMLabel = "mirum-handshake" + EKMLength = 32 +) func generateNonce() ([]byte, error) { nonce := make([]byte, NonceSize) @@ -30,6 +41,56 @@ func generateNonce() ([]byte, error) { return nonce, err } +// LoadPrivateKey reads a PEM-encoded PKCS8 ed25519 private key from path. +func LoadPrivateKey(path string) (ed25519.PrivateKey, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, fmt.Errorf("read key: %w", err) + } + + block, _ := pem.Decode(data) + if block == nil { + return nil, ErrKeyNotPEM + } + + key, err := x509.ParsePKCS8PrivateKey(block.Bytes) + if err != nil { + return nil, fmt.Errorf("%w: %w", ErrKeyNotPKCS8, err) + } + + edKey, ok := key.(ed25519.PrivateKey) + if !ok { + return nil, ErrKeyNotEd25519 + } + + return edKey, nil +} + +// LoadPublicKey reads a PEM-encoded PKIX ed25519 public key from path. +func LoadPublicKey(path string) (ed25519.PublicKey, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, fmt.Errorf("read key: %w", err) + } + + block, _ := pem.Decode(data) + if block == nil { + return nil, ErrKeyNotPEM + } + + key, err := x509.ParsePKIXPublicKey(block.Bytes) + if err != nil { + return nil, fmt.Errorf("%w: %w", ErrKeyNotPKIX, err) + } + + edKey, ok := key.(ed25519.PublicKey) + if !ok { + return nil, ErrKeyNotEd25519 + } + + return edKey, nil +} + // ServerHandshake holds state for the server side of the handshake protocol. type ServerHandshake struct { publicKey ed25519.PublicKey diff --git a/internal/protocol/handshake_test.go b/internal/protocol/handshake_test.go index d5e336b..14f254a 100644 --- a/internal/protocol/handshake_test.go +++ b/internal/protocol/handshake_test.go @@ -6,11 +6,101 @@ package protocol import ( "crypto/ed25519" "crypto/rand" + "crypto/x509" + "encoding/pem" "errors" + "os" + "path/filepath" "testing" "time" ) +func writeTestKeyPair(t *testing.T) (privPath, pubPath string, pub ed25519.PublicKey) { + t.Helper() + pub, priv, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + + dir := t.TempDir() + + privDER, err := x509.MarshalPKCS8PrivateKey(priv) + if err != nil { + t.Fatal(err) + } + privPath = filepath.Join(dir, "test.key") + if err := os.WriteFile(privPath, pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: privDER}), 0o600); err != nil { + t.Fatal(err) + } + + pubDER, err := x509.MarshalPKIXPublicKey(pub) + if err != nil { + t.Fatal(err) + } + pubPath = filepath.Join(dir, "test.pub") + if err := os.WriteFile(pubPath, pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: pubDER}), 0o644); err != nil { + t.Fatal(err) + } + + return privPath, pubPath, pub +} + +func TestLoadPrivateKey(t *testing.T) { + privPath, _, wantPub := writeTestKeyPair(t) + + key, err := LoadPrivateKey(privPath) + if err != nil { + t.Fatal(err) + } + + gotPub := key.Public().(ed25519.PublicKey) + if !gotPub.Equal(wantPub) { + t.Fatal("public key mismatch") + } +} + +func TestLoadPublicKey(t *testing.T) { + _, pubPath, wantPub := writeTestKeyPair(t) + + key, err := LoadPublicKey(pubPath) + if err != nil { + t.Fatal(err) + } + + if !key.Equal(wantPub) { + t.Fatal("public key mismatch") + } +} + +func TestLoadPrivateKey_NotFound(t *testing.T) { + _, err := LoadPrivateKey("/nonexistent/path") + if err == nil { + t.Fatal("expected error") + } +} + +func TestLoadPrivateKey_NotPEM(t *testing.T) { + path := filepath.Join(t.TempDir(), "bad.key") + os.WriteFile(path, []byte("not pem"), 0o600) + + _, err := LoadPrivateKey(path) + if err == nil { + t.Fatal("expected error") + } +} + +func TestLoadPrivateKey_WrongKeyType(t *testing.T) { + // Write a PEM block with garbage DER + data := pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: []byte("garbage")}) + path := filepath.Join(t.TempDir(), "bad.key") + os.WriteFile(path, data, 0o600) + + _, err := LoadPrivateKey(path) + if err == nil { + t.Fatal("expected error") + } +} + func generateTestKey(t *testing.T) ed25519.PrivateKey { t.Helper() _, priv, err := ed25519.GenerateKey(rand.Reader) diff --git a/internal/protocol/key.go b/internal/protocol/key.go deleted file mode 100644 index aab7338..0000000 --- a/internal/protocol/key.go +++ /dev/null @@ -1,70 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package protocol - -import ( - "crypto/ed25519" - "crypto/x509" - "encoding/pem" - "errors" - "fmt" - "os" -) - -var ( - errKeyNotPEM = errors.New("key file does not contain a PEM block") - errKeyNotPKCS8 = errors.New("key file does not contain a PKCS8 private key") - errKeyNotPKIX = errors.New("key file does not contain a PKIX public key") - errKeyNotEd25519 = errors.New("key file does not contain an ed25519 key") -) - -// LoadPrivateKey reads a PEM-encoded PKCS8 ed25519 private key from path. -func LoadPrivateKey(path string) (ed25519.PrivateKey, error) { - data, err := os.ReadFile(path) - if err != nil { - return nil, fmt.Errorf("read key: %w", err) - } - - block, _ := pem.Decode(data) - if block == nil { - return nil, errKeyNotPEM - } - - key, err := x509.ParsePKCS8PrivateKey(block.Bytes) - if err != nil { - return nil, fmt.Errorf("%w: %w", errKeyNotPKCS8, err) - } - - edKey, ok := key.(ed25519.PrivateKey) - if !ok { - return nil, errKeyNotEd25519 - } - - return edKey, nil -} - -// LoadPublicKey reads a PEM-encoded PKIX ed25519 public key from path. -func LoadPublicKey(path string) (ed25519.PublicKey, error) { - data, err := os.ReadFile(path) - if err != nil { - return nil, fmt.Errorf("read key: %w", err) - } - - block, _ := pem.Decode(data) - if block == nil { - return nil, errKeyNotPEM - } - - key, err := x509.ParsePKIXPublicKey(block.Bytes) - if err != nil { - return nil, fmt.Errorf("%w: %w", errKeyNotPKIX, err) - } - - edKey, ok := key.(ed25519.PublicKey) - if !ok { - return nil, errKeyNotEd25519 - } - - return edKey, nil -} diff --git a/internal/protocol/key_test.go b/internal/protocol/key_test.go deleted file mode 100644 index ce0bd98..0000000 --- a/internal/protocol/key_test.go +++ /dev/null @@ -1,100 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package protocol - -import ( - "crypto/ed25519" - "crypto/rand" - "crypto/x509" - "encoding/pem" - "os" - "path/filepath" - "testing" -) - -func writeTestKeyPair(t *testing.T) (privPath, pubPath string, pub ed25519.PublicKey) { - t.Helper() - pub, priv, err := ed25519.GenerateKey(rand.Reader) - if err != nil { - t.Fatal(err) - } - - dir := t.TempDir() - - privDER, err := x509.MarshalPKCS8PrivateKey(priv) - if err != nil { - t.Fatal(err) - } - privPath = filepath.Join(dir, "test.key") - if err := os.WriteFile(privPath, pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: privDER}), 0600); err != nil { - t.Fatal(err) - } - - pubDER, err := x509.MarshalPKIXPublicKey(pub) - if err != nil { - t.Fatal(err) - } - pubPath = filepath.Join(dir, "test.pub") - if err := os.WriteFile(pubPath, pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: pubDER}), 0644); err != nil { - t.Fatal(err) - } - - return privPath, pubPath, pub -} - -func TestLoadPrivateKey(t *testing.T) { - privPath, _, wantPub := writeTestKeyPair(t) - - key, err := LoadPrivateKey(privPath) - if err != nil { - t.Fatal(err) - } - - gotPub := key.Public().(ed25519.PublicKey) - if !gotPub.Equal(wantPub) { - t.Fatal("public key mismatch") - } -} - -func TestLoadPublicKey(t *testing.T) { - _, pubPath, wantPub := writeTestKeyPair(t) - - key, err := LoadPublicKey(pubPath) - if err != nil { - t.Fatal(err) - } - - if !key.Equal(wantPub) { - t.Fatal("public key mismatch") - } -} - -func TestLoadPrivateKey_NotFound(t *testing.T) { - _, err := LoadPrivateKey("/nonexistent/path") - if err == nil { - t.Fatal("expected error") - } -} - -func TestLoadPrivateKey_NotPEM(t *testing.T) { - path := filepath.Join(t.TempDir(), "bad.key") - os.WriteFile(path, []byte("not pem"), 0600) - - _, err := LoadPrivateKey(path) - if err == nil { - t.Fatal("expected error") - } -} - -func TestLoadPrivateKey_WrongKeyType(t *testing.T) { - // Write a PEM block with garbage DER - data := pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: []byte("garbage")}) - path := filepath.Join(t.TempDir(), "bad.key") - os.WriteFile(path, data, 0600) - - _, err := LoadPrivateKey(path) - if err == nil { - t.Fatal("expected error") - } -} diff --git a/internal/protocol/platform.go b/internal/protocol/platform.go index beac035..638a664 100644 --- a/internal/protocol/platform.go +++ b/internal/protocol/platform.go @@ -4,11 +4,24 @@ package protocol import ( + "errors" + "fmt" "runtime" + "strconv" + "strings" + mirum "dimidiumlabs/mirum" "dimidiumlabs/mirum/internal/protocol/pb" ) +var ErrInvalidVersion = errors.New("invalid version string") + +var ( + Major uint32 + Minor uint32 + Patch uint32 +) + var osMap = map[string]pb.Os{ "linux": pb.Os_OS_LINUX, "darwin": pb.Os_OS_DARWIN, @@ -57,3 +70,35 @@ func DetectArch() pb.Arch { } return pb.Arch_ARCH_UNSPECIFIED } + +func init() { + var err error + Major, Minor, Patch, err = ParseVersion(mirum.Version) + if err != nil { + panic(fmt.Sprintf("version: %v", err)) + } +} + +// ParseVersion parses a "major.minor.patch" string. +func ParseVersion(s string) (major, minor, patch uint32, err error) { + s = strings.TrimSpace(s) + parts := strings.SplitN(s, ".", 3) + if len(parts) != 3 { + return 0, 0, 0, ErrInvalidVersion + } + n0, err0 := strconv.ParseUint(parts[0], 10, 32) + n1, err1 := strconv.ParseUint(parts[1], 10, 32) + n2, err2 := strconv.ParseUint(parts[2], 10, 32) + if err0 != nil || err1 != nil || err2 != nil { + return 0, 0, 0, ErrInvalidVersion + } + return uint32(n0), uint32(n1), uint32(n2), nil +} + +func VersionString() string { + return fmt.Sprintf("%d.%d.%d", Major, Minor, Patch) +} + +func VersionProto() *pb.Version { + return &pb.Version{Major: Major, Minor: Minor, Patch: Patch} +} diff --git a/internal/protocol/platform_test.go b/internal/protocol/platform_test.go new file mode 100644 index 0000000..276f335 --- /dev/null +++ b/internal/protocol/platform_test.go @@ -0,0 +1,45 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package protocol + +import ( + "errors" + "testing" +) + +func TestParseVersion(t *testing.T) { + tests := []struct { + name string + in string + major, minor, patch uint32 + wantErr bool + }{ + {"valid", "0.1.0", 0, 1, 0, false}, + {"large", "10.20.30", 10, 20, 30, false}, + {"whitespace", " 1.2.3 ", 1, 2, 3, false}, + {"empty", "", 0, 0, 0, true}, + {"two parts", "1.2", 0, 0, 0, true}, + {"one part", "1", 0, 0, 0, true}, + {"letters", "a.b.c", 0, 0, 0, true}, + {"negative", "-1.0.0", 0, 0, 0, true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + major, minor, patch, err := ParseVersion(tt.in) + if tt.wantErr { + if !errors.Is(err, ErrInvalidVersion) { + t.Fatalf("ParseVersion(%q) err = %v, want ErrInvalidVersion", tt.in, err) + } + return + } + if err != nil { + t.Fatalf("ParseVersion(%q) unexpected error: %v", tt.in, err) + } + if major != tt.major || minor != tt.minor || patch != tt.patch { + t.Fatalf("ParseVersion(%q) = %d.%d.%d, want %d.%d.%d", + tt.in, major, minor, patch, tt.major, tt.minor, tt.patch) + } + }) + } +} diff --git a/internal/protocol/version.go b/internal/protocol/version.go deleted file mode 100644 index 0cd40a9..0000000 --- a/internal/protocol/version.go +++ /dev/null @@ -1,54 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package protocol - -import ( - "errors" - "fmt" - "strconv" - "strings" - - mirum "dimidiumlabs/mirum" - "dimidiumlabs/mirum/internal/protocol/pb" -) - -var ErrInvalidVersion = errors.New("invalid version string") - -var ( - Major uint32 - Minor uint32 - Patch uint32 -) - -func init() { - var err error - Major, Minor, Patch, err = ParseVersion(mirum.Version) - if err != nil { - panic(fmt.Sprintf("version: %v", err)) - } -} - -// ParseVersion parses a "major.minor.patch" string. -func ParseVersion(s string) (major, minor, patch uint32, err error) { - s = strings.TrimSpace(s) - parts := strings.SplitN(s, ".", 3) - if len(parts) != 3 { - return 0, 0, 0, ErrInvalidVersion - } - n0, err0 := strconv.ParseUint(parts[0], 10, 32) - n1, err1 := strconv.ParseUint(parts[1], 10, 32) - n2, err2 := strconv.ParseUint(parts[2], 10, 32) - if err0 != nil || err1 != nil || err2 != nil { - return 0, 0, 0, ErrInvalidVersion - } - return uint32(n0), uint32(n1), uint32(n2), nil -} - -func VersionString() string { - return fmt.Sprintf("%d.%d.%d", Major, Minor, Patch) -} - -func VersionProto() *pb.Version { - return &pb.Version{Major: Major, Minor: Minor, Patch: Patch} -} diff --git a/internal/protocol/version_test.go b/internal/protocol/version_test.go deleted file mode 100644 index 276f335..0000000 --- a/internal/protocol/version_test.go +++ /dev/null @@ -1,45 +0,0 @@ -// Copyright (c) 2026 Nikolay Govorov -// SPDX-License-Identifier: AGPL-3.0-or-later - -package protocol - -import ( - "errors" - "testing" -) - -func TestParseVersion(t *testing.T) { - tests := []struct { - name string - in string - major, minor, patch uint32 - wantErr bool - }{ - {"valid", "0.1.0", 0, 1, 0, false}, - {"large", "10.20.30", 10, 20, 30, false}, - {"whitespace", " 1.2.3 ", 1, 2, 3, false}, - {"empty", "", 0, 0, 0, true}, - {"two parts", "1.2", 0, 0, 0, true}, - {"one part", "1", 0, 0, 0, true}, - {"letters", "a.b.c", 0, 0, 0, true}, - {"negative", "-1.0.0", 0, 0, 0, true}, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - major, minor, patch, err := ParseVersion(tt.in) - if tt.wantErr { - if !errors.Is(err, ErrInvalidVersion) { - t.Fatalf("ParseVersion(%q) err = %v, want ErrInvalidVersion", tt.in, err) - } - return - } - if err != nil { - t.Fatalf("ParseVersion(%q) unexpected error: %v", tt.in, err) - } - if major != tt.major || minor != tt.minor || patch != tt.patch { - t.Fatalf("ParseVersion(%q) = %d.%d.%d, want %d.%d.%d", - tt.in, major, minor, patch, tt.major, tt.minor, tt.patch) - } - }) - } -} diff --git a/proto/admin.proto b/proto/admin.proto index b022bab..a8f4b74 100644 --- a/proto/admin.proto +++ b/proto/admin.proto @@ -10,36 +10,128 @@ option go_package = "dimidiumlabs/mirum/internal/protocol/pb"; import "google/protobuf/timestamp.proto"; service Admin { - rpc CreateUser(CreateUserRequest) returns (CreateUserResponse); - rpc SetPassword(SetPasswordRequest) returns (SetPasswordResponse); - rpc DeleteUser(DeleteUserRequest) returns (DeleteUserResponse); + rpc UserCreate(UserCreateRequest) returns (UserCreateResponse); + rpc UserDelete(UserDeleteRequest) returns (UserDeleteResponse); + rpc UserSetPassword(UserSetPasswordRequest) returns (UserSetPasswordResponse); + rpc OrgList(OrgListRequest) returns (OrgListResponse); + rpc OrgCreate(OrgCreateRequest) returns (OrgCreateResponse); + rpc OrgDelete(OrgDeleteRequest) returns (OrgDeleteResponse); + rpc OrgRename(OrgRenameRequest) returns (OrgRenameResponse); + + rpc OrgMemberList(OrgMemberListRequest) returns (OrgMemberListResponse); + rpc OrgMemberAdd(OrgMemberAddRequest) returns (OrgMemberAddResponse); + rpc OrgMemberRemove(OrgMemberRemoveRequest) returns (OrgMemberRemoveResponse); + rpc OrgMemberSetRole(OrgMemberSetRoleRequest) returns (OrgMemberSetRoleResponse); + + rpc WorkerList(WorkerListRequest) returns (WorkerListResponse); rpc WorkerAdd(WorkerAddRequest) returns (WorkerAddResponse); rpc WorkerRevoke(WorkerRevokeRequest) returns (WorkerRevokeResponse); - rpc WorkerList(WorkerListRequest) returns (WorkerListResponse); } -message CreateUserRequest { +message UserCreateRequest { string email = 1; string password = 2; } -message CreateUserResponse { +message UserCreateResponse { string id = 1; } -message SetPasswordRequest { +message UserSetPasswordRequest { string email = 1; string password = 2; } -message SetPasswordResponse {} +message UserSetPasswordResponse {} -message DeleteUserRequest { +message UserDeleteRequest { string email = 1; } -message DeleteUserResponse {} +message UserDeleteResponse {} + +// Organization management + +message OrgCreateRequest { + string name = 1; + string slug = 2; + bool public = 3; + string owner_email = 4; +} + +message OrgCreateResponse { + string id = 1; +} + +message OrgDeleteRequest { + string slug = 1; +} + +message OrgDeleteResponse {} + +message OrgRenameRequest { + string slug = 1; + string new_name = 2; + string new_slug = 3; +} + +message OrgRenameResponse {} + +message OrgListRequest { + optional string user_email = 1; +} + +message OrgListResponse { + repeated Org organizations = 1; +} + +message Org { + string id = 1; + string name = 2; + string slug = 3; + bool public = 4; + google.protobuf.Timestamp created_at = 5; +} + +// Organization members + +message OrgMemberAddRequest { + string org_slug = 1; + string email = 2; + string role = 3; +} + +message OrgMemberAddResponse {} + +message OrgMemberRemoveRequest { + string org_slug = 1; + string email = 2; +} + +message OrgMemberRemoveResponse {} + +message OrgMemberSetRoleRequest { + string org_slug = 1; + string email = 2; + string role = 3; +} + +message OrgMemberSetRoleResponse {} + +message OrgMemberListRequest { + string org_slug = 1; +} + +message OrgMemberListResponse { + repeated OrgMemberInfo members = 1; +} + +message OrgMemberInfo { + string email = 1; + string role = 2; + google.protobuf.Timestamp joined_at = 3; +} message WorkerAddRequest { bytes public_key = 1; @@ -66,3 +158,4 @@ message Worker { bytes public_key = 2; google.protobuf.Timestamp created_at = 3; } + -- Gilti