From 402b3a13097f00eb51ff1b319bacb38cc472adae Mon Sep 17 00:00:00 2001 From: Nikolay Govorov Date: Thu, 2 Apr 2026 23:36:29 +0100 Subject: Replace hmac handshake to ed25519, support tls binding --- .github/workflows/build.yml | 6 +- Taskfile.yml | 10 ++ cmd/mirumd/config.go | 19 ++- cmd/mirumd/main.go | 107 +++++++++++++- cmd/mirumd/server_admin.go | 37 +++++ cmd/mirumd/server_grpc.go | 72 +++++++--- cmd/mirumw/client.go | 79 +++++++---- cmd/mirumw/config.go | 16 ++- cmd/mirumw/main.go | 14 -- internal/database/database.go | 102 +++++++++++-- internal/protocol/handshake.go | 91 ++++++------ internal/protocol/handshake_test.go | 212 +++++++++++++--------------- internal/protocol/key.go | 70 +++++++++ internal/protocol/key_test.go | 100 +++++++++++++ pkg/mirumd.yaml | 5 +- pkg/mirumw-default.yaml | 11 ++ proto/admin.proto | 32 +++++ proto/mirum.proto | 21 ++- 18 files changed, 740 insertions(+), 264 deletions(-) create mode 100644 internal/protocol/key.go create mode 100644 internal/protocol/key_test.go diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 20aa471..09c8118 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -41,8 +41,8 @@ jobs: passphrase: ${{ secrets.GPG_PASSPHRASE }} gpg_private_key: ${{ secrets.GPG_PRIVATE_KEY }} - - name: Check formatting - run: test -z "$(gofmt -l .)" + - name: Lint + run: task lint - name: Test run: task test @@ -70,8 +70,6 @@ jobs: GPG_PRIVATE_KEY: ${{ secrets.GPG_PRIVATE_KEY }} NFPM_PASSPHRASE: ${{ secrets.GPG_PASSPHRASE }} - - name: Vet - run: go vet ./... - uses: actions/upload-artifact@b7c566a772e6b6bfb58ed0dc250532a479d7789f # v6.0.0 with: diff --git a/Taskfile.yml b/Taskfile.yml index 2373b8b..4f7b8cf 100644 --- a/Taskfile.yml +++ b/Taskfile.yml @@ -8,6 +8,9 @@ vars: DIST_DIR: build/dist ARCHES: amd64 arm64 riscv64 ppc64le +env: + GOEXPERIMENT: runtimesecret + tasks: proto: desc: Generate gRPC code from proto files @@ -18,6 +21,13 @@ tasks: cmds: - protoc --go_out=. --go-grpc_out=. proto/mirum.proto proto/admin.proto + lint: + desc: Run static checks + deps: [proto] + cmds: + - go vet ./... + - gofmt -l . | grep . && exit 1 || true + test: desc: Run all tests deps: [proto] diff --git a/cmd/mirumd/config.go b/cmd/mirumd/config.go index 7eae4e3..13d374c 100644 --- a/cmd/mirumd/config.go +++ b/cmd/mirumd/config.go @@ -11,12 +11,14 @@ import ( ) type config struct { - WwwAddr string `yaml:"www_addr"` - GrpcAddr string `yaml:"grpc_addr"` - AdminSocket string `yaml:"admin_socket"` - DatabaseUri string `yaml:"database_uri"` - WorkerSecret string `yaml:"worker_secret"` - Pepper string `yaml:"pepper"` + WwwAddr string `yaml:"www_addr"` + GrpcAddr string `yaml:"grpc_addr"` + AdminSocket string `yaml:"admin_socket"` + DatabaseUri string `yaml:"database_uri"` + Pepper string `yaml:"pepper"` + + TLSCert string `yaml:"tls_cert"` + TLSKey string `yaml:"tls_key"` GitHubToken string `yaml:"token"` WebhookSecret string `yaml:"webhook_secret"` @@ -26,7 +28,7 @@ func getConfig(filename string) (*config, error) { cfg := &config{ GrpcAddr: ":2026", WwwAddr: ":3000", - AdminSocket: "/run/mirum/admin.sock", + AdminSocket: "/run/mirumd/admin.sock", } data, err := os.ReadFile(filename) @@ -41,9 +43,6 @@ func getConfig(filename string) (*config, error) { if cfg.DatabaseUri == "" { return nil, fmt.Errorf("error: database_uri is required") } - if cfg.WorkerSecret == "" { - return nil, fmt.Errorf("error: worker_secret is required") - } if cfg.WebhookSecret == "" { return nil, fmt.Errorf("error: webhook_secret is required") } diff --git a/cmd/mirumd/main.go b/cmd/mirumd/main.go index 9c413f8..91dd8c5 100644 --- a/cmd/mirumd/main.go +++ b/cmd/mirumd/main.go @@ -5,6 +5,8 @@ package main import ( "context" + "crypto/tls" + "encoding/base64" "fmt" "log/slog" "net" @@ -27,7 +29,7 @@ func main() { var socketPath string root := &cobra.Command{Use: "mirumd", Short: "Mirum CI server"} - root.PersistentFlags().StringVar(&socketPath, "socket", "", "admin socket path (default from config or /run/mirum/admin.sock)") + root.PersistentFlags().StringVar(&socketPath, "socket", "", "admin socket path (default from config or /run/mirumd/admin.sock)") daemonCmd := &cobra.Command{ Use: "daemon", @@ -86,6 +88,42 @@ func main() { deleteUserCmd.Flags().String("email", "", "user email") _ = deleteUserCmd.MarkFlagRequired("email") + workerCmd := &cobra.Command{Use: "worker", Short: "Manage workers"} + root.AddCommand(workerCmd) + + workerAddCmd := &cobra.Command{ + Use: "add", + Short: "Register a worker", + Run: func(cmd *cobra.Command, args []string) { + pubkey, _ := cmd.Flags().GetString("pubkey") + workerAdd(socketPath, pubkey) + }, + } + workerCmd.AddCommand(workerAddCmd) + workerAddCmd.Flags().String("pubkey", "", "base64-encoded ed25519 public key") + _ = workerAddCmd.MarkFlagRequired("pubkey") + + workerRevokeCmd := &cobra.Command{ + Use: "revoke", + Short: "Revoke a worker", + Run: func(cmd *cobra.Command, args []string) { + id, _ := cmd.Flags().GetString("id") + workerRevoke(socketPath, id) + }, + } + workerCmd.AddCommand(workerRevokeCmd) + workerRevokeCmd.Flags().String("id", "", "worker ID") + _ = workerRevokeCmd.MarkFlagRequired("id") + + workerListCmd := &cobra.Command{ + Use: "list", + Short: "List active workers", + Run: func(cmd *cobra.Command, args []string) { + workerList(socketPath) + }, + } + workerCmd.AddCommand(workerListCmd) + if err := root.Execute(); err != nil { os.Exit(1) } @@ -128,8 +166,22 @@ func daemon(configFile, socketFlag string) { sup := supervisor.Detect() ctx := sup.WaitForStop(context.Background()) + var tlsCfg *tls.Config + if cfg.TLSCert != "" && cfg.TLSKey != "" { + cert, err := tls.LoadX509KeyPair(cfg.TLSCert, cfg.TLSKey) + if err != nil { + slog.Error("failed to load TLS certificate", "err", err) + os.Exit(1) + } + tlsCfg = &tls.Config{ + Certificates: []tls.Certificate{cert}, + MinVersion: tls.VersionTLS13, // required for reliable EKM (channel binding) + } + slog.Info("TLS enabled", "cert", cfg.TLSCert) + } + wwwSrv := NewWwwServer(ctx, srv) - grpcSrv := NewGrpcServer(ctx, srv, []byte(cfg.WorkerSecret)) + grpcSrv := NewGrpcServer(ctx, srv, tlsCfg) adminSrv := NewAdminServer(srv) grpcLn, webLn, adminLn, err := listeners(cfg) @@ -138,6 +190,10 @@ func daemon(configFile, socketFlag string) { os.Exit(1) } + if tlsCfg != nil { + webLn = tls.NewListener(webLn, tlsCfg) + } + slog.Info("listening", "grpc", grpcLn.Addr(), "web", webLn.Addr(), "admin", cfg.AdminSocket) go func() { @@ -196,7 +252,7 @@ func listeners(cfg *config) (grpcLn, webLn, adminLn net.Listener, err error) { return nil, nil, nil, err } defer func() { - if err != nil { + if err != nil && grpcLn != nil { _ = grpcLn.Close() } }() @@ -207,7 +263,7 @@ func listeners(cfg *config) (grpcLn, webLn, adminLn net.Listener, err error) { return nil, nil, nil, err } defer func() { - if err != nil { + if err != nil && webLn != nil { _ = webLn.Close() } }() @@ -216,7 +272,7 @@ func listeners(cfg *config) (grpcLn, webLn, adminLn net.Listener, err error) { return nil, nil, nil, err } defer func() { - if err != nil { + if err != nil && adminLn != nil { _ = adminLn.Close() } }() @@ -226,7 +282,7 @@ func listeners(cfg *config) (grpcLn, webLn, adminLn net.Listener, err error) { func adminClient(socketPath string) pb.AdminClient { if socketPath == "" { - socketPath = "/run/mirum/admin.sock" + socketPath = "/run/mirumd/admin.sock" } conn, err := grpc.NewClient("unix://"+socketPath, grpc.WithTransportCredentials(insecure.NewCredentials()), @@ -272,3 +328,42 @@ func userDelete(socketPath, email string) { } fmt.Println("ok") } + +func workerAdd(socketPath, pubkeyB64 string) { + pubkey, err := base64.StdEncoding.DecodeString(pubkeyB64) + if err != nil { + fmt.Fprintln(os.Stderr, "invalid base64:", err) + os.Exit(1) + } + resp, err := adminClient(socketPath).WorkerAdd(context.Background(), &pb.WorkerAddRequest{ + PublicKey: pubkey, + }) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + fmt.Println(resp.Id) +} + +func workerRevoke(socketPath, id string) { + _, err := adminClient(socketPath).WorkerRevoke(context.Background(), &pb.WorkerRevokeRequest{ + Id: id, + }) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + fmt.Println("ok") +} + +func workerList(socketPath string) { + resp, err := adminClient(socketPath).WorkerList(context.Background(), &pb.WorkerListRequest{}) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + for _, w := range resp.Workers { + created := w.CreatedAt.AsTime().Format(time.DateOnly) + fmt.Printf("%s\t%s\t%s\n", w.Id, base64.StdEncoding.EncodeToString(w.PublicKey), created) + } +} diff --git a/cmd/mirumd/server_admin.go b/cmd/mirumd/server_admin.go index 93567d2..cf72fbc 100644 --- a/cmd/mirumd/server_admin.go +++ b/cmd/mirumd/server_admin.go @@ -5,10 +5,13 @@ package main import ( "context" + "crypto/ed25519" + "fmt" "dimidiumlabs/mirum/internal/protocol/pb" "google.golang.org/grpc" + "google.golang.org/protobuf/types/known/timestamppb" ) func NewAdminServer(srv *server) *grpc.Server { @@ -44,3 +47,37 @@ func (a *adminService) DeleteUser(ctx context.Context, req *pb.DeleteUserRequest } return &pb.DeleteUserResponse{}, nil } + +func (a *adminService) WorkerAdd(ctx context.Context, req *pb.WorkerAddRequest) (*pb.WorkerAddResponse, error) { + if len(req.PublicKey) != ed25519.PublicKeySize { + return nil, fmt.Errorf("invalid public key: expected %d bytes, got %d", ed25519.PublicKeySize, len(req.PublicKey)) + } + id, err := a.srv.db.AddWorker(ctx, req.PublicKey) + if err != nil { + return nil, err + } + return &pb.WorkerAddResponse{Id: id}, nil +} + +func (a *adminService) WorkerRevoke(ctx context.Context, req *pb.WorkerRevokeRequest) (*pb.WorkerRevokeResponse, error) { + if err := a.srv.db.RevokeWorker(ctx, req.Id); err != nil { + return nil, err + } + return &pb.WorkerRevokeResponse{}, nil +} + +func (a *adminService) WorkerList(ctx context.Context, req *pb.WorkerListRequest) (*pb.WorkerListResponse, error) { + workers, err := a.srv.db.ListWorkers(ctx) + if err != nil { + return nil, err + } + pbWorkers := make([]*pb.Worker, len(workers)) + for i, w := range workers { + pbWorkers[i] = &pb.Worker{ + Id: w.ID, + PublicKey: w.PublicKey, + CreatedAt: timestamppb.New(w.CreatedAt), + } + } + return &pb.WorkerListResponse{Workers: pbWorkers}, nil +} diff --git a/cmd/mirumd/server_grpc.go b/cmd/mirumd/server_grpc.go index 6490715..08841a9 100644 --- a/cmd/mirumd/server_grpc.go +++ b/cmd/mirumd/server_grpc.go @@ -5,6 +5,7 @@ package main import ( "context" + "crypto/tls" "fmt" "log/slog" "sync" @@ -16,6 +17,8 @@ import ( "google.golang.org/grpc" "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/peer" "google.golang.org/grpc/stats" "google.golang.org/grpc/status" "google.golang.org/protobuf/types/known/timestamppb" @@ -24,17 +27,21 @@ import ( // connIDKey is the context key for the unique connection identifier. type connIDKey struct{} -func NewGrpcServer(ctx context.Context, srv *server, secret []byte) *grpc.Server { +func NewGrpcServer(ctx context.Context, srv *server, tlsCfg *tls.Config) *grpc.Server { gsrv := &grpcService{ - ctx: ctx, - srv: srv, - secret: secret, + ctx: ctx, + srv: srv, + tls: tlsCfg != nil, } - s := grpc.NewServer( + opts := []grpc.ServerOption{ grpc.UnaryInterceptor(gsrv.unaryInterceptor), grpc.StreamInterceptor(gsrv.streamInterceptor), grpc.StatsHandler(&connTracker{gsrv: gsrv}), - ) + } + if tlsCfg != nil { + opts = append(opts, grpc.Creds(credentials.NewTLS(tlsCfg))) + } + s := grpc.NewServer(opts...) pb.RegisterMirumServer(s, gsrv) return s } @@ -45,7 +52,7 @@ type grpcService struct { ctx context.Context srv *server - secret []byte + tls bool // whether TLS is enabled (for channel binding) authedConns sync.Map // conn ID (uint64) → true nextConnID atomic.Uint64 } @@ -71,9 +78,7 @@ func (g *grpcService) Complete(ctx context.Context, result *pb.TaskResult) (*pb. } func (g *grpcService) Handshake(stream pb.Mirum_HandshakeServer) error { - hs := protocol.NewServerHandshake(g.secret) - - // Step 1: receive worker nonce + // Step 1: receive worker public key in, err := stream.Recv() if err != nil { return fmt.Errorf("recv worker challenge: %w", err) @@ -83,23 +88,34 @@ func (g *grpcService) Handshake(stream pb.Mirum_HandshakeServer) error { return fmt.Errorf("expected WorkerChallenge") } - // Step 2: send server nonce + proof - serverNonce, proof, err := hs.Challenge(wc.GetNonce()) + // Look up the worker in the database + worker, err := g.srv.db.LookupWorker(stream.Context(), wc.GetPublicKey()) + if err != nil { + return g.reject(stream, "unknown worker") + } + + // Extract TLS EKM for channel binding (nil if no TLS) + ekm := g.extractEKM(stream.Context()) + + hs := protocol.NewServerHandshake() + + // Step 2: generate challenge nonce + serverNonce, err := hs.Challenge(wc.GetPublicKey(), ekm) if err != nil { return err } if err := stream.Send(&pb.HandshakeOut{ Step: &pb.HandshakeOut_ServerChallenge{ ServerChallenge: &pb.ServerChallenge{ - Nonce: serverNonce, - Proof: proof, + Nonce: serverNonce, + Binded: ekm != nil, }, }, }); err != nil { return fmt.Errorf("send server challenge: %w", err) } - // Step 3: receive worker proof + metadata + // Step 3: receive worker signature + metadata in, err = stream.Recv() if err != nil { return fmt.Errorf("recv worker proof: %w", err) @@ -114,7 +130,7 @@ func (g *grpcService) Handshake(stream pb.Mirum_HandshakeServer) error { return g.reject(stream, "worker_time is required") } - if err := hs.Verify(wp.GetProof(), wt.AsTime()); err != nil { + if err := hs.Verify(wp.GetSignature(), wt.AsTime()); err != nil { return g.reject(stream, err.Error()) } @@ -124,7 +140,7 @@ func (g *grpcService) Handshake(stream pb.Mirum_HandshakeServer) error { } slog.Info("worker connected", - "id", fmt.Sprintf("%x", wp.GetId()), + "worker_id", worker.ID, "name", wp.GetName(), "os", wp.GetOs(), "arch", wp.GetArch(), @@ -221,3 +237,25 @@ func (t *connTracker) HandleConn(ctx context.Context, s stats.ConnStats) { } } } + +// extractEKM returns TLS Exported Keying Material from the peer context, +// or nil if TLS is not enabled. +func (g *grpcService) extractEKM(ctx context.Context) []byte { + if !g.tls { + return nil + } + p, ok := peer.FromContext(ctx) + if !ok { + return nil + } + ti, ok := p.AuthInfo.(credentials.TLSInfo) + if !ok { + return nil + } + ekm, err := ti.State.ExportKeyingMaterial(protocol.EKMLabel, nil, protocol.EKMLength) + if err != nil { + slog.Warn("failed to export keying material", "err", err) + return nil + } + return ekm +} diff --git a/cmd/mirumw/client.go b/cmd/mirumw/client.go index 3f7c0f0..13d4062 100644 --- a/cmd/mirumw/client.go +++ b/cmd/mirumw/client.go @@ -6,10 +6,11 @@ package main import ( "context" "crypto/tls" + "crypto/x509" "fmt" "log/slog" - "net" "os" + "runtime/secret" "time" "dimidiumlabs/mirum/internal/executor" @@ -18,7 +19,7 @@ import ( "google.golang.org/grpc" "google.golang.org/grpc/credentials" - "google.golang.org/grpc/credentials/insecure" + "google.golang.org/grpc/peer" "google.golang.org/protobuf/types/known/timestamppb" ) @@ -29,20 +30,24 @@ type client struct { } func connect(ctx context.Context, cfg *config) (*client, error) { - var creds grpc.DialOption - if cfg.Insecure { - host, _, _ := net.SplitHostPort(cfg.Server) - if !isPrivateHost(host) { - return nil, fmt.Errorf("tls_insecure is only allowed for private/loopback addresses, got %s", host) + tlsCfg := &tls.Config{ + MinVersion: tls.VersionTLS13, // required for reliable EKM (channel binding) + } + if cfg.TLSCA != "" { + caCert, err := os.ReadFile(cfg.TLSCA) + if err != nil { + return nil, fmt.Errorf("read CA cert: %w", err) } - - slog.Warn("TLS disabled, connection is not encrypted") - creds = grpc.WithTransportCredentials(insecure.NewCredentials()) - } else { - creds = grpc.WithTransportCredentials(credentials.NewTLS(&tls.Config{})) + pool := x509.NewCertPool() + if !pool.AppendCertsFromPEM(caCert) { + return nil, fmt.Errorf("failed to parse CA cert") + } + tlsCfg.RootCAs = pool } - conn, err := grpc.NewClient(cfg.Server, creds) + conn, err := grpc.NewClient(cfg.Server, + grpc.WithTransportCredentials(credentials.NewTLS(tlsCfg)), + ) if err != nil { return nil, err } @@ -98,27 +103,26 @@ func (c *client) work(ctx context.Context) error { } func (c *client) handshake(ctx context.Context) error { - stream, err := c.handle.Handshake(ctx) + var p peer.Peer + stream, err := c.handle.Handshake(ctx, grpc.Peer(&p)) if err != nil { return fmt.Errorf("open stream: %w", err) } - hs := protocol.NewClientHandshake([]byte(c.cfg.Secret)) - - // Step 1: send worker nonce - workerNonce, err := hs.Challenge() + // Step 1: send worker public key + pubKey, err := protocol.LoadPublicKey(c.cfg.PubKeyFile) if err != nil { - return fmt.Errorf("generate nonce: %w", err) + return fmt.Errorf("load public key: %w", err) } if err := stream.Send(&pb.HandshakeIn{ Step: &pb.HandshakeIn_WorkerChallenge{ - WorkerChallenge: &pb.WorkerChallenge{Nonce: workerNonce}, + WorkerChallenge: &pb.WorkerChallenge{PublicKey: pubKey}, }, }); err != nil { return fmt.Errorf("send challenge: %w", err) } - // Step 2: receive server challenge, verify and compute proof + // Step 2: receive server nonce out, err := stream.Recv() if err != nil { return fmt.Errorf("recv server challenge: %w", err) @@ -128,12 +132,35 @@ func (c *client) handshake(ctx context.Context) error { return fmt.Errorf("expected ServerChallenge") } - workerProof, err := hs.Verify(sc.GetNonce(), sc.GetProof()) - if err != nil { - return err + // Extract EKM if server requested channel binding + var ekm []byte + if sc.GetBinded() { + ti, ok := p.AuthInfo.(credentials.TLSInfo) + if !ok { + return fmt.Errorf("channel binding requested but TLS info not available") + } + ekm, err = ti.State.ExportKeyingMaterial(protocol.EKMLabel, nil, protocol.EKMLength) + if err != nil { + return fmt.Errorf("export keying material: %w", err) + } + } + + // Step 3: sign nonce (+ optional EKM) inside secret.Do + var signature []byte + var signErr error + secret.Do(func() { + key, err := protocol.LoadPrivateKey(c.cfg.KeyFile) + if err != nil { + signErr = fmt.Errorf("load key: %w", err) + return + } + hs := protocol.NewClientHandshake(key) + signature, signErr = hs.Sign(sc.GetNonce(), ekm) + }) + if signErr != nil { + return signErr } - // Step 3: send worker proof + metadata name := c.cfg.Name if name == "" { name, _ = os.Hostname() @@ -141,7 +168,7 @@ func (c *client) handshake(ctx context.Context) error { if err := stream.Send(&pb.HandshakeIn{ Step: &pb.HandshakeIn_WorkerProof{ WorkerProof: &pb.WorkerProof{ - Proof: workerProof, + Signature: signature, Name: name, Version: protocol.VersionProto(), Os: protocol.DetectOs(), diff --git a/cmd/mirumw/config.go b/cmd/mirumw/config.go index f8b079e..45e0557 100644 --- a/cmd/mirumw/config.go +++ b/cmd/mirumw/config.go @@ -15,10 +15,11 @@ import ( const workerRuntime = "host" type config struct { - Name string `yaml:"name"` - Server string `yaml:"server"` - Secret string `yaml:"secret"` - Insecure bool `yaml:"tls_insecure"` + Name string `yaml:"name"` + Server string `yaml:"server"` + KeyFile string `yaml:"key_file"` + PubKeyFile string `yaml:"pub_key_file"` + TLSCA string `yaml:"tls_ca"` // custom CA cert for self-signed/dev } func getConfig(filename string) (*config, error) { @@ -37,5 +38,12 @@ func getConfig(filename string) (*config, error) { } } + if cfg.KeyFile == "" { + return nil, fmt.Errorf("error: key_file is required") + } + if cfg.PubKeyFile == "" { + return nil, fmt.Errorf("error: pub_key_file is required") + } + return cfg, nil } diff --git a/cmd/mirumw/main.go b/cmd/mirumw/main.go index 40b5f2b..ebb1d5a 100644 --- a/cmd/mirumw/main.go +++ b/cmd/mirumw/main.go @@ -7,7 +7,6 @@ import ( "context" "flag" "log/slog" - "net" "os" "dimidiumlabs/mirum/internal/protocol" @@ -55,16 +54,3 @@ func main() { slog.Info("shutting down") sup.Stopping() } - -func isPrivateHost(host string) bool { - ips, err := net.LookupIP(host) - if err != nil { - return false - } - for _, ip := range ips { - if !ip.IsLoopback() && !ip.IsPrivate() && !ip.IsLinkLocalUnicast() { - return false - } - } - return len(ips) > 0 -} diff --git a/internal/database/database.go b/internal/database/database.go index 6ae3cd2..1e84770 100644 --- a/internal/database/database.go +++ b/internal/database/database.go @@ -15,6 +15,7 @@ import ( "strings" "time" + "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" "github.com/jackc/tern/v2/migrate" "golang.org/x/crypto/argon2" @@ -29,14 +30,17 @@ const ( ) 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") + 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") ) // DB wraps a pgx connection pool. @@ -97,6 +101,16 @@ func (db *DB) Migrate(ctx context.Context) error { `DROP TABLE sessions`, ) + migrator.AppendMigration("create_workers", + `CREATE TABLE workers ( + id UUID PRIMARY KEY DEFAULT uuidv7(), + public_key BYTEA NOT NULL UNIQUE, + created_at TIMESTAMPTZ NOT NULL DEFAULT now(), + revoked_at TIMESTAMPTZ + )`, + `DROP TABLE workers`, + ) + return migrator.Migrate(ctx) } @@ -320,3 +334,75 @@ func hashPassword(password string, pepper []byte) (string, error) { 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, + ) + 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`, + ) + 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() +} diff --git a/internal/protocol/handshake.go b/internal/protocol/handshake.go index e14ac0b..f25cb0a 100644 --- a/internal/protocol/handshake.go +++ b/internal/protocol/handshake.go @@ -4,72 +4,67 @@ package protocol import ( - "crypto/hmac" + "crypto/ed25519" "crypto/rand" - "crypto/sha256" "errors" + "slices" "time" ) var ( - ErrInvalidNonce = errors.New("invalid nonce size") - ErrInvalidProof = errors.New("invalid proof") - ErrClockSkew = errors.New("clock skew too large") - ErrServerRejected = errors.New("server rejected handshake") + ErrInvalidPublicKey = errors.New("invalid public key size") + ErrInvalidNonce = errors.New("invalid nonce size") + ErrInvalidSignature = errors.New("invalid signature") + ErrClockSkew = errors.New("clock skew too large") + ErrServerRejected = errors.New("server rejected handshake") ) const NonceSize = 32 +const EKMLabel = "mirum-handshake" +const EKMLength = 32 + func generateNonce() ([]byte, error) { nonce := make([]byte, NonceSize) _, err := rand.Read(nonce) return nonce, err } -// computeProof returns HMAC-SHA256(secret, first || second). -func computeProof(secret, first, second []byte) []byte { - mac := hmac.New(sha256.New, secret) - mac.Write(first) - mac.Write(second) - return mac.Sum(nil) -} - -func verifyProof(secret, first, second, proof []byte) bool { - expected := computeProof(secret, first, second) - return hmac.Equal(expected, proof) -} - // ServerHandshake holds state for the server side of the handshake protocol. type ServerHandshake struct { - secret []byte - workerNonce []byte + publicKey ed25519.PublicKey serverNonce []byte + ekm []byte // TLS Exported Keying Material, nil if not bound } -func NewServerHandshake(secret []byte) *ServerHandshake { - return &ServerHandshake{secret: secret} +func NewServerHandshake() *ServerHandshake { + return &ServerHandshake{} } -// Challenge validates the worker nonce and returns the server nonce + proof. -func (h *ServerHandshake) Challenge(workerNonce []byte) (serverNonce, proof []byte, err error) { - if len(workerNonce) != NonceSize { - return nil, nil, ErrInvalidNonce +// Challenge validates the worker public key and returns a random nonce. +// If ekm is non-nil, channel binding is enabled and the EKM will be +// included in the signed data during verification. +func (h *ServerHandshake) Challenge(publicKey, ekm []byte) (serverNonce []byte, err error) { + if len(publicKey) != ed25519.PublicKeySize { + return nil, ErrInvalidPublicKey } - h.workerNonce = workerNonce + h.publicKey = publicKey + h.ekm = ekm h.serverNonce, err = generateNonce() if err != nil { - return nil, nil, err + return nil, err } - proof = computeProof(h.secret, workerNonce, h.serverNonce) - return h.serverNonce, proof, nil + return h.serverNonce, nil } -// Verify checks the worker's proof and clock skew. -func (h *ServerHandshake) Verify(proof []byte, workerTime time.Time) error { - if !verifyProof(h.secret, h.serverNonce, h.workerNonce, proof) { - return ErrInvalidProof +// Verify checks the worker's ed25519 signature and clock skew. +// The signed data is nonce (or nonce || ekm if channel binding is enabled). +func (h *ServerHandshake) Verify(signature []byte, workerTime time.Time) error { + signed := slices.Concat(h.serverNonce, h.ekm) + if !ed25519.Verify(h.publicKey, signed, signature) { + return ErrInvalidSignature } skew := time.Since(workerTime).Abs() @@ -82,25 +77,23 @@ func (h *ServerHandshake) Verify(proof []byte, workerTime time.Time) error { // ClientHandshake holds state for the client side of the handshake protocol. type ClientHandshake struct { - secret []byte - workerNonce []byte + privateKey ed25519.PrivateKey } -func NewClientHandshake(secret []byte) *ClientHandshake { - return &ClientHandshake{secret: secret} +func NewClientHandshake(privateKey ed25519.PrivateKey) *ClientHandshake { + return &ClientHandshake{privateKey: privateKey} } -// Challenge generates the worker nonce. -func (h *ClientHandshake) Challenge() (workerNonce []byte, err error) { - h.workerNonce, err = generateNonce() - return h.workerNonce, err +// PublicKey returns the 32-byte ed25519 public key. +func (h *ClientHandshake) PublicKey() []byte { + return h.privateKey.Public().(ed25519.PublicKey) } -// Verify checks the server proof and returns the worker proof. -func (h *ClientHandshake) Verify(serverNonce, serverProof []byte) (proof []byte, err error) { - if !verifyProof(h.secret, h.workerNonce, serverNonce, serverProof) { - return nil, ErrInvalidProof +// Sign signs the server nonce (and optional EKM) with the worker's private key. +// If ekm is non-nil, the signed data is nonce || ekm. +func (h *ClientHandshake) Sign(serverNonce, ekm []byte) ([]byte, error) { + if len(serverNonce) != NonceSize { + return nil, ErrInvalidNonce } - proof = computeProof(h.secret, serverNonce, h.workerNonce) - return proof, nil + return ed25519.Sign(h.privateKey, slices.Concat(serverNonce, ekm)), nil } diff --git a/internal/protocol/handshake_test.go b/internal/protocol/handshake_test.go index 702eebe..d5e336b 100644 --- a/internal/protocol/handshake_test.go +++ b/internal/protocol/handshake_test.go @@ -4,172 +4,158 @@ package protocol import ( - "bytes" - "crypto/hmac" - "crypto/sha256" + "crypto/ed25519" + "crypto/rand" "errors" "testing" "time" ) -func TestGenerateNonce(t *testing.T) { - nonce, err := generateNonce() +func generateTestKey(t *testing.T) ed25519.PrivateKey { + t.Helper() + _, priv, err := ed25519.GenerateKey(rand.Reader) if err != nil { t.Fatal(err) } - if len(nonce) != NonceSize { - t.Fatalf("len = %d, want %d", len(nonce), NonceSize) - } - - nonce2, _ := generateNonce() - if bytes.Equal(nonce, nonce2) { - t.Fatal("two nonces are identical") - } + return priv } -func TestComputeProof(t *testing.T) { - secret := []byte("secret") - first := []byte("first") - second := []byte("second") - - proof := computeProof(secret, first, second) - - // Golden value: HMAC-SHA256("secret", "first" || "second") - mac := hmac.New(sha256.New, secret) - mac.Write(first) - mac.Write(second) - want := mac.Sum(nil) +func TestHandshake_Success(t *testing.T) { + priv := generateTestKey(t) + pub := priv.Public().(ed25519.PublicKey) - if !bytes.Equal(proof, want) { - t.Fatalf("proof mismatch") - } + server := NewServerHandshake() + client := NewClientHandshake(priv) - // Different secret → different proof - other := computeProof([]byte("other"), first, second) - if bytes.Equal(proof, other) { - t.Fatal("different secrets produced same proof") + if got := client.PublicKey(); len(got) != ed25519.PublicKeySize { + t.Fatalf("public key len = %d, want %d", len(got), ed25519.PublicKeySize) } - // Order matters: (first, second) != (second, first) - reversed := computeProof(secret, second, first) - if bytes.Equal(proof, reversed) { - t.Fatal("argument order did not affect proof") + serverNonce, err := server.Challenge(pub, nil) + if err != nil { + t.Fatal(err) } -} -func TestVerifyProof(t *testing.T) { - secret := []byte("secret") - first := []byte("first") - second := []byte("second") - proof := computeProof(secret, first, second) - - tests := []struct { - name string - secret []byte - first []byte - second []byte - proof []byte - want bool - }{ - {"valid", secret, first, second, proof, true}, - {"wrong proof", secret, first, second, []byte("wrong"), false}, - {"wrong secret", []byte("wrong"), first, second, proof, false}, - {"swapped args", secret, second, first, proof, false}, - {"empty proof", secret, first, second, nil, false}, + signature, err := client.Sign(serverNonce, nil) + if err != nil { + t.Fatal(err) } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := verifyProof(tt.secret, tt.first, tt.second, tt.proof) - if got != tt.want { - t.Fatalf("VerifyProof = %v, want %v", got, tt.want) - } - }) + if err := server.Verify(signature, time.Now()); err != nil { + t.Fatal(err) } } -// Full handshake: client and server with matching secrets. -func TestHandshake_Success(t *testing.T) { - secret := []byte("shared-secret") - server := NewServerHandshake(secret) - client := NewClientHandshake(secret) +func TestHandshake_ChannelBinding(t *testing.T) { + priv := generateTestKey(t) + pub := priv.Public().(ed25519.PublicKey) + ekm := make([]byte, 32) + rand.Read(ekm) + + server := NewServerHandshake() + client := NewClientHandshake(priv) - // Step 1: client generates nonce - workerNonce, err := client.Challenge() + serverNonce, err := server.Challenge(pub, ekm) if err != nil { t.Fatal(err) } - // Step 2: server responds with its nonce + proof - serverNonce, serverProof, err := server.Challenge(workerNonce) + signature, err := client.Sign(serverNonce, ekm) if err != nil { t.Fatal(err) } - // Step 3: client verifies server, produces worker proof - workerProof, err := client.Verify(serverNonce, serverProof) - if err != nil { + if err := server.Verify(signature, time.Now()); err != nil { t.Fatal(err) } +} - // Step 4: server verifies worker - if err := server.Verify(workerProof, time.Now()); err != nil { - t.Fatal(err) +func TestHandshake_ChannelBinding_MismatchedEKM(t *testing.T) { + priv := generateTestKey(t) + pub := priv.Public().(ed25519.PublicKey) + + serverEKM := make([]byte, 32) + workerEKM := make([]byte, 32) + rand.Read(serverEKM) + rand.Read(workerEKM) + + server := NewServerHandshake() + client := NewClientHandshake(priv) + + serverNonce, _ := server.Challenge(pub, serverEKM) + signature, _ := client.Sign(serverNonce, workerEKM) + + err := server.Verify(signature, time.Now()) + if !errors.Is(err, ErrInvalidSignature) { + t.Fatalf("err = %v, want ErrInvalidSignature", err) } } -func TestHandshake_WrongSecret(t *testing.T) { - server := NewServerHandshake([]byte("server-secret")) - client := NewClientHandshake([]byte("wrong-secret")) +func TestHandshake_WrongKey(t *testing.T) { + workerKey := generateTestKey(t) + otherKey := generateTestKey(t) - workerNonce, _ := client.Challenge() - _, serverProof, _ := server.Challenge(workerNonce) + server := NewServerHandshake() + client := NewClientHandshake(workerKey) - // Client cannot verify server proof - _, err := client.Verify(server.serverNonce, serverProof) - if !errors.Is(err, ErrInvalidProof) { - t.Fatalf("err = %v, want ErrInvalidProof", err) + serverNonce, _ := server.Challenge(otherKey.Public().(ed25519.PublicKey), nil) + signature, _ := client.Sign(serverNonce, nil) + + err := server.Verify(signature, time.Now()) + if !errors.Is(err, ErrInvalidSignature) { + t.Fatalf("err = %v, want ErrInvalidSignature", err) } } -func TestHandshake_BadNonce(t *testing.T) { - server := NewServerHandshake([]byte("secret")) +func TestHandshake_InvalidPublicKeyLength(t *testing.T) { + server := NewServerHandshake() - _, _, err := server.Challenge([]byte("short")) - if !errors.Is(err, ErrInvalidNonce) { - t.Fatalf("err = %v, want ErrInvalidNonce", err) + _, err := server.Challenge([]byte("short"), nil) + if !errors.Is(err, ErrInvalidPublicKey) { + t.Fatalf("err = %v, want ErrInvalidPublicKey", err) + } +} + +func TestHandshake_TamperedSignature(t *testing.T) { + priv := generateTestKey(t) + pub := priv.Public().(ed25519.PublicKey) + + server := NewServerHandshake() + client := NewClientHandshake(priv) + + serverNonce, _ := server.Challenge(pub, nil) + signature, _ := client.Sign(serverNonce, nil) + + signature[0] ^= 0xff + + err := server.Verify(signature, time.Now()) + if !errors.Is(err, ErrInvalidSignature) { + t.Fatalf("err = %v, want ErrInvalidSignature", err) } } func TestHandshake_ClockSkew(t *testing.T) { - secret := []byte("secret") - server := NewServerHandshake(secret) - client := NewClientHandshake(secret) + priv := generateTestKey(t) + pub := priv.Public().(ed25519.PublicKey) + + server := NewServerHandshake() + client := NewClientHandshake(priv) - workerNonce, _ := client.Challenge() - serverNonce, serverProof, _ := server.Challenge(workerNonce) - workerProof, _ := client.Verify(serverNonce, serverProof) + serverNonce, _ := server.Challenge(pub, nil) + signature, _ := client.Sign(serverNonce, nil) - err := server.Verify(workerProof, time.Now().Add(-2*time.Minute)) + err := server.Verify(signature, time.Now().Add(-2*time.Minute)) if !errors.Is(err, ErrClockSkew) { t.Fatalf("err = %v, want ErrClockSkew", err) } } -func TestHandshake_TamperedProof(t *testing.T) { - secret := []byte("secret") - server := NewServerHandshake(secret) - client := NewClientHandshake(secret) +func TestSign_InvalidNonceLength(t *testing.T) { + priv := generateTestKey(t) + client := NewClientHandshake(priv) - workerNonce, _ := client.Challenge() - serverNonce, serverProof, _ := server.Challenge(workerNonce) - workerProof, _ := client.Verify(serverNonce, serverProof) - - // Flip a byte - workerProof[0] ^= 0xff - - err := server.Verify(workerProof, time.Now()) - if !errors.Is(err, ErrInvalidProof) { - t.Fatalf("err = %v, want ErrInvalidProof", err) + _, err := client.Sign([]byte("short"), nil) + if !errors.Is(err, ErrInvalidNonce) { + t.Fatalf("err = %v, want ErrInvalidNonce", err) } } diff --git a/internal/protocol/key.go b/internal/protocol/key.go new file mode 100644 index 0000000..aab7338 --- /dev/null +++ b/internal/protocol/key.go @@ -0,0 +1,70 @@ +// 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 new file mode 100644 index 0000000..ce0bd98 --- /dev/null +++ b/internal/protocol/key_test.go @@ -0,0 +1,100 @@ +// 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/pkg/mirumd.yaml b/pkg/mirumd.yaml index 63d7d72..411892e 100644 --- a/pkg/mirumd.yaml +++ b/pkg/mirumd.yaml @@ -5,9 +5,10 @@ # See mirumd.socket for details (FileDescriptorName=grpc / web). grpc_addr: :2026 www_addr: :3000 -worker_secret: "" +tls_cert: "" +tls_key: "" webhook_secret: "" token: "" pepper: "" database_uri: "" -admin_socket: /run/mirum/admin.sock +admin_socket: /run/mirumd/admin.sock diff --git a/pkg/mirumw-default.yaml b/pkg/mirumw-default.yaml index 1fc4eee..832e35c 100644 --- a/pkg/mirumw-default.yaml +++ b/pkg/mirumw-default.yaml @@ -2,3 +2,14 @@ # SPDX-License-Identifier: AGPL-3.0-or-later server: localhost:2026 + +# Ed25519 key pair for worker authentication. +# Generate with: +# openssl genpkey -algorithm Ed25519 -out /etc/mirum/mirumw-default.key +# openssl pkey -in /etc/mirum/mirumw-default.key -pubout -out /etc/mirum/mirumw-default.pub +key_file: /etc/mirum/mirumw-default.key +pub_key_file: /etc/mirum/mirumw-default.pub + +# Custom CA certificate for self-signed/dev TLS. +# Leave empty to use system trust store. +tls_ca: "" diff --git a/proto/admin.proto b/proto/admin.proto index 89bee00..8f628aa 100644 --- a/proto/admin.proto +++ b/proto/admin.proto @@ -7,10 +7,16 @@ package mirum; option go_package = "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 WorkerAdd(WorkerAddRequest) returns (WorkerAddResponse); + rpc WorkerRevoke(WorkerRevokeRequest) returns (WorkerRevokeResponse); + rpc WorkerList(WorkerListRequest) returns (WorkerListResponse); } message CreateUserRequest { @@ -34,3 +40,29 @@ message DeleteUserRequest { } message DeleteUserResponse {} + +message WorkerAddRequest { + bytes public_key = 1; +} + +message WorkerAddResponse { + string id = 1; +} + +message WorkerRevokeRequest { + string id = 1; +} + +message WorkerRevokeResponse {} + +message WorkerListRequest {} + +message WorkerListResponse { + repeated Worker workers = 1; +} + +message Worker { + string id = 1; + bytes public_key = 2; + google.protobuf.Timestamp created_at = 3; +} diff --git a/proto/mirum.proto b/proto/mirum.proto index 8f243fb..69a856d 100644 --- a/proto/mirum.proto +++ b/proto/mirum.proto @@ -13,12 +13,12 @@ import "google/protobuf/timestamp.proto"; // GRPC is the only contract between them, so to implement your own worker, // you only need to implement this service. service Mirum { - // Handshake performs mutual authentication via HMAC-SHA256 challenge-response. + // Handshake performs worker authentication via ed25519 challenge-response. // The server will not accept any calls until a handshake is completed. // - // Step 1 (W→S): WorkerChallenge — worker sends random nonce. - // Step 2 (S→W): ServerChallenge — server sends nonce + HMAC(secret, w_nonce || s_nonce). - // Step 3 (W→S): WorkerProof — worker sends HMAC(secret, s_nonce || w_nonce) + metadata. + // Step 1 (W→S): WorkerChallenge — worker sends its ed25519 public key. + // Step 2 (S→W): ServerChallenge — server sends a random nonce. + // Step 3 (W→S): WorkerProof — worker signs the nonce with its private key + metadata. // Step 4 (S→W): ServerResult — server accepts or rejects. rpc Handshake(stream HandshakeIn) returns (stream HandshakeOut); @@ -142,21 +142,20 @@ message HandshakeOut { } } -// Step 1: Worker generates a random nonce and sends it to the server. +// Step 1: Worker sends its ed25519 public key (32 bytes). message WorkerChallenge { - bytes nonce = 1; + bytes public_key = 1; } -// Step 2: Server generates its own nonce and proves it knows the secret. +// Step 2: Server sends a random nonce for the worker to sign. message ServerChallenge { bytes nonce = 1; - bytes proof = 2; // HMAC-SHA256(secret, worker_nonce || server_nonce) + bool binded = 2; // if true, worker must include TLS EKM in signed data } -// Step 3: Worker proves it knows the secret and sends metadata. -// Metadata is sent only after the server is verified. +// Step 3: Worker signs the server nonce and sends metadata. message WorkerProof { - bytes proof = 1; // HMAC-SHA256(secret, server_nonce || worker_nonce) + bytes signature = 1; bytes id = 2; string name = 3; -- Gilti