aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorNikolay Govorov <me@govorov.online>2026-04-02 23:36:29 +0100
committerNikolay Govorov <me@govorov.online>2026-04-02 23:36:29 +0100
commit402b3a13097f00eb51ff1b319bacb38cc472adae (patch)
treeb66ef76e031d328ed63eb02a110b56230d3db9ea
parentd1d50c5003b30dcada50af9ae898f7c07ec822fc (diff)
downloadtar
tar.gz
tar.bz2
tar.lz
tar.xz
tar.zst
zip
Replace hmac handshake to ed25519, support tls binding
Diffstat
-rw-r--r--.github/workflows/build.yml6+2 −4
-rw-r--r--Taskfile.yml10+10 −0
-rw-r--r--cmd/mirumd/config.go19+9 −10
-rw-r--r--cmd/mirumd/main.go107+101 −6
-rw-r--r--cmd/mirumd/server_admin.go37+37 −0
-rw-r--r--cmd/mirumd/server_grpc.go72+55 −17
-rw-r--r--cmd/mirumw/client.go79+53 −26
-rw-r--r--cmd/mirumw/config.go16+12 −4
-rw-r--r--cmd/mirumw/main.go14+0 −14
-rw-r--r--internal/database/database.go102+94 −8
-rw-r--r--internal/protocol/handshake.go91+42 −49
-rw-r--r--internal/protocol/handshake_test.go212+99 −113
-rw-r--r--internal/protocol/key.go70+70 −0
-rw-r--r--internal/protocol/key_test.go100+100 −0
-rw-r--r--pkg/mirumd.yaml5+3 −2
-rw-r--r--pkg/mirumw-default.yaml11+11 −0
-rw-r--r--proto/admin.proto32+32 −0
-rw-r--r--proto/mirum.proto21+10 −11
18 files changed, 740 insertions, 264 deletions
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
--- /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
--- /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;