aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorNikolay Govorov <me@govorov.online>2026-04-04 04:53:12 +0100
committerNikolay Govorov <me@govorov.online>2026-04-04 04:53:12 +0100
commit0b3e983f6b8a80d88e92980cbfc0937cf381ea7f (patch)
tree398fa28de76d9b6e495f35c3fd8a3c10be284eb5
parentc3c9802b8a2a16077793498b995b6a632f4641c5 (diff)
downloadtar
tar.gz
tar.bz2
tar.lz
tar.xz
tar.zst
zip
Use mTLS with ed25519 keys for handshake
Diffstat
-rw-r--r--Taskfile.yml5+1 −4
-rw-r--r--cmd/mirumd/config.go16+12 −4
-rw-r--r--cmd/mirumd/main.go72+31 −41
-rw-r--r--cmd/mirumd/server_admin.go13+11 −2
-rw-r--r--cmd/mirumd/server_grpc.go284+62 −222
-rw-r--r--cmd/mirumd/server_web.go111+60 −51
-rw-r--r--cmd/mirumw/client.go194+42 −152
-rw-r--r--cmd/mirumw/config.go12+4 −8
-rw-r--r--cmd/mirumw/main.go1+1 −0
-rw-r--r--go.mod1+0 −1
-rw-r--r--go.sum2+0 −2
-rw-r--r--internal/protocol/handshake.go162+64 −98
-rw-r--r--internal/protocol/handshake_test.go180+54 −126
-rw-r--r--proto/mirum.proto79+8 −71
14 files changed, 350 insertions, 782 deletions
diff --git a/Taskfile.yml b/Taskfile.yml
index ab6a14d..31d8e3f 100644
--- a/Taskfile.yml
+++ b/Taskfile.yml
@@ -53,7 +53,7 @@ tasks:
- cp {{.BUILD_DIR}}/mirumw-{{.GOOS}}-{{.GOARCH}} {{.BUILD_DIR}}/mirumw
dev:
- desc: Generate dev TLS cert and worker keys
+ desc: Generate dev TLS cert and worker key
cmds:
- mkdir -p dev
- >-
@@ -65,9 +65,6 @@ tasks:
- >-
test -f dev/worker.key ||
openssl genpkey -algorithm ed25519 -out dev/worker.key
- - >-
- test -f dev/worker.pub ||
- openssl pkey -in dev/worker.key -pubout -out dev/worker.pub
package:
desc: Build deb/rpm packages for all architectures
diff --git a/cmd/mirumd/config.go b/cmd/mirumd/config.go
index 13d374c..9516a1c 100644
--- a/cmd/mirumd/config.go
+++ b/cmd/mirumd/config.go
@@ -10,15 +10,20 @@ import (
"gopkg.in/yaml.v3"
)
+type tlsConfig struct {
+ Cert string `yaml:"cert"`
+ Key string `yaml:"key"`
+}
+
type config struct {
- WwwAddr string `yaml:"www_addr"`
+ WebAddr string `yaml:"web_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"`
+ GrpcTls tlsConfig `yaml:"grpc_tls"`
+ WebTls *tlsConfig `yaml:"web_tls"` // optional
GitHubToken string `yaml:"token"`
WebhookSecret string `yaml:"webhook_secret"`
@@ -27,7 +32,7 @@ type config struct {
func getConfig(filename string) (*config, error) {
cfg := &config{
GrpcAddr: ":2026",
- WwwAddr: ":3000",
+ WebAddr: ":3000",
AdminSocket: "/run/mirumd/admin.sock",
}
@@ -49,6 +54,9 @@ func getConfig(filename string) (*config, error) {
if cfg.Pepper == "" {
return nil, fmt.Errorf("error: pepper is required")
}
+ if cfg.GrpcTls.Cert == "" || cfg.GrpcTls.Key == "" {
+ return nil, fmt.Errorf("error: grpc_tls.cert and grpc_tls.key are required")
+ }
return cfg, nil
}
diff --git a/cmd/mirumd/main.go b/cmd/mirumd/main.go
index b0b8368..2357bd6 100644
--- a/cmd/mirumd/main.go
+++ b/cmd/mirumd/main.go
@@ -6,7 +6,6 @@ package main
import (
"context"
"crypto/ed25519"
- "crypto/tls"
"crypto/x509"
"encoding/base64"
"fmt"
@@ -18,12 +17,11 @@ import (
"strconv"
"time"
- "dimidiumlabs/mirum/internal/database"
- "dimidiumlabs/mirum/internal/forges"
-
"connectrpc.com/connect"
"github.com/google/uuid"
+ "dimidiumlabs/mirum/internal/database"
+ "dimidiumlabs/mirum/internal/forges"
"dimidiumlabs/mirum/internal/protocol/pb"
"dimidiumlabs/mirum/internal/protocol/pb/pbconnect"
"dimidiumlabs/mirum/internal/supervisor"
@@ -273,14 +271,17 @@ func daemon(configFile, socketFlag string) {
slog.Info("config loaded", "configfile", configFile)
- db, err := database.Open(context.Background(), cfg.DatabaseUri)
+ sup := supervisor.Detect()
+ ctx := sup.WaitForStop(context.Background())
+
+ db, err := database.Open(ctx, cfg.DatabaseUri)
if err != nil {
slog.Error("couldn't open database: %w", "err", err)
os.Exit(1)
}
defer db.Close()
- if err := db.Migrate(context.Background()); err != nil {
+ if err := db.Migrate(ctx); err != nil {
slog.Error("migration failed: %w", "err", err)
os.Exit(1)
}
@@ -288,33 +289,17 @@ func daemon(configFile, socketFlag string) {
slog.Info("database ready")
srv := &server{
- cfg: cfg,
db: db,
+ cfg: cfg,
forge: &forges.GitHub{Secret: cfg.WebhookSecret, Token: cfg.GitHubToken},
queue: make(chan *pb.Task, 100),
}
- 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)
- NextProtos: []string{"h2"}, // required for gRPC over TLS (ALPN)
- }
- slog.Info("TLS enabled", "cert", cfg.TLSCert)
- }
+ go srv.PurgeSessions(ctx)
- wwwSrv := NewWwwServer(ctx, srv)
- grpcSrv := NewGrpcServer(ctx, srv, tlsCfg)
- adminSrv := NewAdminServer(srv)
+ webSrv := NewWebServer(ctx, srv)
+ grpcSrv := NewGrpcServer(ctx, srv)
+ adminSrv := NewAdminServer(ctx, srv)
grpcLn, webLn, adminLn, err := listeners(cfg)
if err != nil {
@@ -322,21 +307,22 @@ func daemon(configFile, socketFlag string) {
os.Exit(1)
}
- if tlsCfg != nil {
- grpcLn = tls.NewListener(grpcLn, tlsCfg)
- webLn = tls.NewListener(webLn, tlsCfg)
- }
-
slog.Info("listening", "grpc", grpcLn.Addr(), "web", webLn.Addr(), "admin", cfg.AdminSocket)
go func() {
- if err := wwwSrv.Serve(webLn); err != nil && err != http.ErrServerClosed {
+ var err error
+ if webSrv.TLSConfig != nil {
+ err = webSrv.ServeTLS(webLn, "", "")
+ } else {
+ err = webSrv.Serve(webLn)
+ }
+ if err != nil && err != http.ErrServerClosed {
slog.Error("web server failed", "err", err)
os.Exit(1)
}
}()
go func() {
- if err := grpcSrv.Serve(grpcLn); err != nil && err != http.ErrServerClosed {
+ if err := grpcSrv.ServeTLS(grpcLn, "", ""); err != nil && err != http.ErrServerClosed {
slog.Error("grpc server failed", "err", err)
os.Exit(1)
}
@@ -348,8 +334,6 @@ func daemon(configFile, socketFlag string) {
}
}()
- go srv.PurgeSessions(ctx)
-
sup.Ready()
go sup.StartWatchdog(ctx)
@@ -365,9 +349,15 @@ func daemon(configFile, socketFlag string) {
srv.Close()
- wwwSrv.Shutdown(context.Background())
- grpcSrv.Shutdown(context.Background())
- adminSrv.Shutdown(context.Background())
+ if err := webSrv.Shutdown(context.Background()); err != nil {
+ slog.Error("web server shutdown", "err", err)
+ }
+ if err := grpcSrv.Shutdown(context.Background()); err != nil {
+ slog.Error("grpc server shutdown", "err", err)
+ }
+ if err := adminSrv.Shutdown(context.Background()); err != nil {
+ slog.Error("admin server shutdown", "err", err)
+ }
}
// listeners returns gRPC, web, and admin listeners.
@@ -392,7 +382,7 @@ func listeners(cfg *config) (grpcLn, webLn, adminLn net.Listener, err error) {
if lns := named["web"]; len(lns) > 0 {
webLn = lns[0]
- } else if webLn, err = net.Listen("tcp", cfg.WwwAddr); err != nil {
+ } else if webLn, err = net.Listen("tcp", cfg.WebAddr); err != nil {
return nil, nil, nil, err
}
defer func() {
@@ -425,7 +415,7 @@ func listeners(cfg *config) (grpcLn, webLn, adminLn net.Listener, err error) {
return nil, nil, nil, fmt.Errorf("chown admin socket: %w", err)
}
- if err = os.Chmod(cfg.AdminSocket, 0660); err != nil {
+ if err = os.Chmod(cfg.AdminSocket, 0o660); err != nil {
return nil, nil, nil, fmt.Errorf("chmod admin socket: %w", err)
}
diff --git a/cmd/mirumd/server_admin.go b/cmd/mirumd/server_admin.go
index 94ab144..b23b4c3 100644
--- a/cmd/mirumd/server_admin.go
+++ b/cmd/mirumd/server_admin.go
@@ -7,6 +7,7 @@ import (
"context"
"errors"
"fmt"
+ "net"
"net/http"
"connectrpc.com/connect"
@@ -19,14 +20,22 @@ import (
"dimidiumlabs/mirum/internal/protocol/pb/pbconnect"
)
-func NewAdminServer(srv *server) *http.Server {
+func NewAdminServer(ctx context.Context, srv *server) *http.Server {
as := &adminService{srv: srv}
+
path, handler := pbconnect.NewAdminHandler(as,
connect.WithInterceptors(validate.NewInterceptor()),
)
+
mux := http.NewServeMux()
mux.Handle(path, handler)
- return &http.Server{Handler: mux}
+
+ return &http.Server{
+ Handler: mux,
+ BaseContext: func(_ net.Listener) context.Context {
+ return ctx
+ },
+ }
}
type adminService struct {
diff --git a/cmd/mirumd/server_grpc.go b/cmd/mirumd/server_grpc.go
index 1630728..afe2a6a 100644
--- a/cmd/mirumd/server_grpc.go
+++ b/cmd/mirumd/server_grpc.go
@@ -5,61 +5,83 @@ package main
import (
"context"
+ "crypto/ed25519"
"crypto/tls"
+ "crypto/x509"
+ "errors"
"fmt"
"log/slog"
"net"
"net/http"
- "strings"
- "sync"
"time"
"connectrpc.com/connect"
- "golang.org/x/net/http2"
- "golang.org/x/net/http2/h2c"
- "google.golang.org/protobuf/types/known/timestamppb"
+ "connectrpc.com/validate"
"dimidiumlabs/mirum/internal/protocol"
"dimidiumlabs/mirum/internal/protocol/pb"
"dimidiumlabs/mirum/internal/protocol/pb/pbconnect"
)
-// connKey is the context key for the underlying net.Conn.
-type connKey struct{}
+func NewGrpcServer(ctx context.Context, srv *server) *http.Server {
+ gsrv := &grpcService{srv: srv}
-// tlsStateKey is the context key for TLS connection state.
-type tlsStateKey struct{}
+ path, handler := pbconnect.NewMirumHandler(gsrv,
+ connect.WithInterceptors(validate.NewInterceptor()),
+ )
-func NewGrpcServer(ctx context.Context, srv *server, tlsCfg *tls.Config) *http.Server {
- gsrv := &grpcService{
- ctx: ctx,
- srv: srv,
- tls: tlsCfg != nil,
- }
- opts := []connect.HandlerOption{
- connect.WithInterceptors(gsrv),
- }
- path, handler := pbconnect.NewMirumHandler(gsrv, opts...)
mux := http.NewServeMux()
- mux.Handle(path, requireGRPC(tlsMiddleware(handler)))
+ mux.Handle(path, workerLog(handler))
+
return &http.Server{
- Handler: h2c.NewHandler(mux, &http2.Server{}),
- ConnContext: gsrv.connContext,
- ConnState: gsrv.connState,
+ Handler: mux,
BaseContext: func(_ net.Listener) context.Context {
return ctx
},
+ TLSConfig: &tls.Config{
+ NextProtos: []string{"h2"},
+ MinVersion: tls.VersionTLS13,
+ ClientAuth: tls.RequireAnyClientCert,
+ GetCertificate: func(_ *tls.ClientHelloInfo) (*tls.Certificate, error) {
+ cert, err := tls.LoadX509KeyPair(srv.cfg.GrpcTls.Cert, srv.cfg.GrpcTls.Key)
+ return &cert, err
+ },
+ VerifyPeerCertificate: func(rawCerts [][]byte, _ [][]*x509.Certificate) error {
+ if len(rawCerts) == 0 {
+ return errors.New("client certificate required")
+ }
+
+ c, err := x509.ParseCertificate(rawCerts[0])
+ if err != nil {
+ return fmt.Errorf("parse client cert: %w", err)
+ }
+
+ pubKey, ok := c.PublicKey.(ed25519.PublicKey)
+ if !ok {
+ return errors.New("ed25519 certificate required")
+ }
+
+ if _, err := srv.db.LookupWorker(context.Background(), pubKey); err != nil {
+ return fmt.Errorf("unknown worker: %w", err)
+ }
+
+ // Clock skew: NotBefore is set to time.Now() when the cert was generated.
+ // Checked here (once per TLS handshake), not in the interceptor,
+ // because HTTP/2 reuses the connection and NotBefore would go stale.
+ if skew := time.Since(c.NotBefore).Abs(); skew > time.Minute {
+ return fmt.Errorf("%w: %s", protocol.ErrClockSkew, skew.Truncate(time.Second))
+ }
+
+ return nil
+ },
+ },
}
}
// grpcService is the ConnectRPC transport adapter over server.
type grpcService struct {
pbconnect.UnimplementedMirumHandler
-
- ctx context.Context
- srv *server
- tls bool // whether TLS is enabled (for channel binding)
- authedConns sync.Map // net.Conn → true
+ srv *server
}
func (g *grpcService) Poll(ctx context.Context, req *connect.Request[pb.PollRequest]) (*connect.Response[pb.Task], error) {
@@ -82,202 +104,20 @@ func (g *grpcService) Complete(ctx context.Context, req *connect.Request[pb.Task
return connect.NewResponse(&pb.CompleteResponse{}), nil
}
-func (g *grpcService) Handshake(ctx context.Context, stream *connect.BidiStream[pb.HandshakeIn, pb.HandshakeOut]) error {
- // Step 1: receive worker public key
- in, err := stream.Receive()
- if err != nil {
- return fmt.Errorf("recv worker challenge: %w", err)
- }
- wc := in.GetWorkerChallenge()
- if wc == nil {
- return fmt.Errorf("expected WorkerChallenge")
- }
-
- // Look up the worker in the database
- worker, err := g.srv.db.LookupWorker(ctx, wc.GetPublicKey())
- if err != nil {
- return g.reject(stream, "unknown worker")
- }
-
- // Extract TLS EKM for channel binding (nil if no TLS)
- ekm := g.extractEKM(ctx)
-
- 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,
- Binded: ekm != nil,
- },
- },
- }); err != nil {
- return fmt.Errorf("send server challenge: %w", err)
- }
-
- // Step 3: receive worker signature + metadata
- in, err = stream.Receive()
- if err != nil {
- return fmt.Errorf("recv worker proof: %w", err)
- }
- wp := in.GetWorkerProof()
- if wp == nil {
- return fmt.Errorf("expected WorkerProof")
- }
-
- wt := wp.GetWorkerTime()
- if wt == nil {
- return g.reject(stream, "worker_time is required")
- }
-
- if err := hs.Verify(wp.GetSignature(), wt.AsTime()); err != nil {
- return g.reject(stream, err.Error())
- }
-
- var warnings []string
- if skew := time.Since(wt.AsTime()).Abs(); skew > 10*time.Second {
- warnings = append(warnings, fmt.Sprintf("clock drift: %s", skew.Truncate(time.Second)))
- }
-
- slog.Info("worker connected",
- "worker_id", worker.ID,
- "name", wp.GetName(),
- "os", wp.GetOs(),
- "arch", wp.GetArch(),
- "runtime", wp.GetRuntime(),
- )
-
- // Step 4: accept
- if conn, ok := ctx.Value(connKey{}).(net.Conn); ok {
- g.authedConns.Store(conn, true)
- }
- return g.sendResult(stream, nil, warnings)
-}
-
-func (g *grpcService) sendResult(stream *connect.BidiStream[pb.HandshakeIn, pb.HandshakeOut], errMsg *string, warnings []string) error {
- return stream.Send(&pb.HandshakeOut{
- Step: &pb.HandshakeOut_ServerResult{
- ServerResult: &pb.ServerResult{
- Error: errMsg,
- ServerVersion: protocol.VersionProto(),
- ServerTime: timestamppb.Now(),
- Warnings: warnings,
- },
- },
- })
-}
-
-func (g *grpcService) reject(stream *connect.BidiStream[pb.HandshakeIn, pb.HandshakeOut], reason string) error {
- if err := g.sendResult(stream, &reason, nil); err != nil {
- return err
- }
- return fmt.Errorf("%s", reason)
-}
-
-func (g *grpcService) requireAuth(ctx context.Context, procedure string) error {
- if procedure == pbconnect.MirumHandshakeProcedure {
- return nil
- }
- conn, ok := ctx.Value(connKey{}).(net.Conn)
- if !ok {
- return connect.NewError(connect.CodeUnauthenticated, fmt.Errorf("handshake required"))
- }
- if _, ok := g.authedConns.Load(conn); !ok {
- return connect.NewError(connect.CodeUnauthenticated, fmt.Errorf("handshake required"))
- }
- return nil
-}
-
-// WrapUnary implements connect.Interceptor for unary RPCs (Poll, Complete).
-func (g *grpcService) WrapUnary(next connect.UnaryFunc) connect.UnaryFunc {
- return func(ctx context.Context, req connect.AnyRequest) (connect.AnyResponse, error) {
- if err := g.requireAuth(ctx, req.Spec().Procedure); err != nil {
- return nil, err
- }
- return next(ctx, req)
- }
-}
-
-// WrapStreamingClient is a no-op (server-side only).
-func (g *grpcService) WrapStreamingClient(next connect.StreamingClientFunc) connect.StreamingClientFunc {
- return next
-}
-
-// WrapStreamingHandler implements connect.Interceptor for streaming RPCs (Handshake).
-func (g *grpcService) WrapStreamingHandler(next connect.StreamingHandlerFunc) connect.StreamingHandlerFunc {
- return func(ctx context.Context, conn connect.StreamingHandlerConn) error {
- if err := g.requireAuth(ctx, conn.Spec().Procedure); err != nil {
- return err
- }
- if conn.Spec().Procedure == pbconnect.MirumHandshakeProcedure {
- done := make(chan error, 1)
- go func() { done <- next(ctx, conn) }()
- select {
- case err := <-done:
- return err
- case <-time.After(30 * time.Second):
- return fmt.Errorf("handshake timeout")
- }
- }
- return next(ctx, conn)
- }
-}
-
-// connContext stores the net.Conn in the context for connection tracking.
-func (g *grpcService) connContext(ctx context.Context, c net.Conn) context.Context {
- return context.WithValue(ctx, connKey{}, c)
-}
-
-// connState cleans up authedConns when a connection closes.
-func (g *grpcService) connState(c net.Conn, state http.ConnState) {
- if state == http.StateClosed {
- g.authedConns.Delete(c)
- slog.Debug("peer disconnected")
- }
-}
-
-// extractEKM returns TLS Exported Keying Material from the context,
-// or nil if TLS is not enabled.
-func (g *grpcService) extractEKM(ctx context.Context) []byte {
- if !g.tls {
- return nil
- }
- state, ok := ctx.Value(tlsStateKey{}).(*tls.ConnectionState)
- if !ok || state == nil {
- return nil
- }
- ekm, err := state.ExportKeyingMaterial(protocol.EKMLabel, nil, protocol.EKMLength)
- if err != nil {
- slog.Warn("failed to export keying material", "err", err)
- return nil
- }
- return ekm
-}
-
-// requireGRPC rejects requests that do not use the gRPC wire protocol.
-func requireGRPC(next http.Handler) http.Handler {
+// workerLog logs worker metadata from the mTLS client certificate
+// and sets the server version response header.
+func workerLog(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- ct := r.Header.Get("Content-Type")
- if ct != "application/grpc" && !strings.HasPrefix(ct, "application/grpc+") {
- http.Error(w, "only gRPC protocol is supported", http.StatusUnsupportedMediaType)
- return
- }
- next.ServeHTTP(w, r)
- })
-}
-
-// tlsMiddleware injects the TLS connection state into the request context.
-func tlsMiddleware(next http.Handler) http.Handler {
- return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- if r.TLS != nil {
- ctx := context.WithValue(r.Context(), tlsStateKey{}, r.TLS)
- r = r.WithContext(ctx)
+ if r.TLS != nil && len(r.TLS.PeerCertificates) > 0 {
+ if meta := protocol.ParseWorkerMeta(r.TLS.PeerCertificates[0]); meta != nil {
+ slog.Info("worker request",
+ "name", meta.Name,
+ "version", meta.Version,
+ "path", r.URL.Path,
+ )
+ }
}
+ w.Header().Set("X-Server-Version", protocol.VersionString())
next.ServeHTTP(w, r)
})
}
diff --git a/cmd/mirumd/server_web.go b/cmd/mirumd/server_web.go
index 608cf61..2ecdea4 100644
--- a/cmd/mirumd/server_web.go
+++ b/cmd/mirumd/server_web.go
@@ -7,6 +7,7 @@ import (
"context"
"crypto/rand"
"crypto/subtle"
+ "crypto/tls"
"embed"
"encoding/base64"
"errors"
@@ -27,56 +28,7 @@ var (
loginTmpl = template.Must(template.ParseFS(templateFS, "templates/layout.html", "templates/login.html"))
)
-// csrfToken returns the current CSRF token, setting a cookie if absent.
-func csrfToken(w http.ResponseWriter, r *http.Request) string {
- if c, err := r.Cookie("csrf"); err == nil && c.Value != "" {
- return c.Value
- }
- b := make([]byte, 32)
- rand.Read(b)
- token := base64.RawURLEncoding.EncodeToString(b)
- http.SetCookie(w, &http.Cookie{
- Name: "csrf",
- Value: token,
- Path: "/",
- HttpOnly: true,
- Secure: true,
- SameSite: http.SameSiteStrictMode,
- })
- return token
-}
-
-// csrfOK checks that the form field matches the cookie (double-submit).
-func csrfOK(r *http.Request) bool {
- cookie, err := r.Cookie("csrf")
- if err != nil || cookie.Value == "" {
- return false
- }
- field := r.FormValue("csrf")
- return subtle.ConstantTimeCompare([]byte(cookie.Value), []byte(field)) == 1
-}
-
-func clearCookie(w http.ResponseWriter, name string) {
- http.SetCookie(w, &http.Cookie{
- Name: name,
- Value: "",
- Path: "/",
- HttpOnly: true,
- Secure: true,
- MaxAge: -1,
- })
-}
-
-func NewWwwServer(ctx context.Context, srv *server) *http.Server {
- return &http.Server{
- Handler: wwwRoutes(srv),
- BaseContext: func(_ net.Listener) context.Context {
- return ctx
- },
- }
-}
-
-func wwwRoutes(srv *server) http.Handler {
+func NewWebServer(ctx context.Context, srv *server) *http.Server {
mux := http.NewServeMux()
mux.HandleFunc("GET /", func(w http.ResponseWriter, r *http.Request) {
@@ -173,5 +125,62 @@ func wwwRoutes(srv *server) http.Handler {
http.Redirect(w, r, "/auth/login", http.StatusSeeOther)
})
- return mux
+ var tls_config *tls.Config = nil
+ if srv.cfg.WebTls != nil {
+ tls_config = &tls.Config{
+ MinVersion: tls.VersionTLS13,
+ GetCertificate: func(_ *tls.ClientHelloInfo) (*tls.Certificate, error) {
+ cert, err := tls.LoadX509KeyPair(srv.cfg.WebTls.Cert, srv.cfg.WebTls.Key)
+ return &cert, err
+ },
+ }
+ }
+
+ return &http.Server{
+ Handler: mux,
+ TLSConfig: tls_config,
+ BaseContext: func(_ net.Listener) context.Context {
+ return ctx
+ },
+ }
+}
+
+// csrfToken returns the current CSRF token, setting a cookie if absent.
+func csrfToken(w http.ResponseWriter, r *http.Request) string {
+ if c, err := r.Cookie("csrf"); err == nil && c.Value != "" {
+ return c.Value
+ }
+ b := make([]byte, 32)
+ rand.Read(b)
+ token := base64.RawURLEncoding.EncodeToString(b)
+ http.SetCookie(w, &http.Cookie{
+ Name: "csrf",
+ Value: token,
+ Path: "/",
+ HttpOnly: true,
+ Secure: true,
+ SameSite: http.SameSiteStrictMode,
+ })
+ return token
+}
+
+// csrfOK checks that the form field matches the cookie (double-submit).
+func csrfOK(r *http.Request) bool {
+ cookie, err := r.Cookie("csrf")
+ if err != nil || cookie.Value == "" {
+ return false
+ }
+ field := r.FormValue("csrf")
+ return subtle.ConstantTimeCompare([]byte(cookie.Value), []byte(field)) == 1
+}
+
+func clearCookie(w http.ResponseWriter, name string) {
+ http.SetCookie(w, &http.Cookie{
+ Name: name,
+ Value: "",
+ Path: "/",
+ HttpOnly: true,
+ Secure: true,
+ MaxAge: -1,
+ })
}
diff --git a/cmd/mirumw/client.go b/cmd/mirumw/client.go
index 3496cac..7f3886d 100644
--- a/cmd/mirumw/client.go
+++ b/cmd/mirumw/client.go
@@ -9,15 +9,11 @@ import (
"crypto/x509"
"fmt"
"log/slog"
- "net"
"net/http"
"os"
- "runtime/secret"
- "sync"
- "time"
+ "runtime"
"connectrpc.com/connect"
- "google.golang.org/protobuf/types/known/timestamppb"
"dimidiumlabs/mirum/internal/executor"
"dimidiumlabs/mirum/internal/protocol"
@@ -26,70 +22,67 @@ import (
)
type client struct {
- cfg *config
- http *http.Client
- handle pbconnect.MirumClient
- tlsMu sync.Mutex
- tlsConn *tls.Conn
+ cfg *config
+ http *http.Client
+ handle pbconnect.MirumClient
}
func dial(ctx context.Context, cfg *config) (*client, error) {
+ name := cfg.Name
+ if name == "" {
+ name, _ = os.Hostname()
+ }
+
+ meta := &protocol.WorkerMeta{
+ Os: runtime.GOOS,
+ Arch: runtime.GOARCH,
+ Name: name,
+ Runtime: workerRuntime,
+ Version: protocol.VersionString(),
+ }
+
tlsCfg := &tls.Config{
- MinVersion: tls.VersionTLS13, // required for reliable EKM (channel binding)
- NextProtos: []string{"h2"}, // required for gRPC over TLS (ALPN)
+ NextProtos: []string{"h2"},
+ MinVersion: tls.VersionTLS13,
+ GetClientCertificate: func(_ *tls.CertificateRequestInfo) (*tls.Certificate, error) {
+ key, err := protocol.LoadPrivateKey(cfg.KeyFile)
+ if err != nil {
+ return nil, fmt.Errorf("load key: %w", err)
+ }
+
+ cert, err := protocol.SelfSignedCert(key, meta)
+ return &cert, err
+ },
}
+
if cfg.TLSCA != "" {
caCert, err := os.ReadFile(cfg.TLSCA)
if err != nil {
return nil, fmt.Errorf("read CA cert: %w", err)
}
+
pool := x509.NewCertPool()
if !pool.AppendCertsFromPEM(caCert) {
return nil, fmt.Errorf("failed to parse CA cert")
}
+
tlsCfg.RootCAs = pool
}
- c := &client{cfg: cfg}
-
- c.http = &http.Client{
- Transport: &http.Transport{
- DialTLSContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
- conn, err := tls.Dial(network, addr, tlsCfg)
- if err != nil {
- return nil, err
- }
- c.tlsMu.Lock()
- c.tlsConn = conn
- c.tlsMu.Unlock()
- return conn, nil
- },
- ForceAttemptHTTP2: true,
+ c := &client{
+ cfg: cfg,
+ http: &http.Client{
+ Transport: &http.Transport{TLSClientConfig: tlsCfg, ForceAttemptHTTP2: true},
},
}
c.handle = pbconnect.NewMirumClient(c.http, "https://"+cfg.Server, connect.WithGRPC())
slog.Info("dialing", "server", cfg.Server)
- hsCtx, hsCancel := context.WithTimeout(ctx, 30*time.Second)
- defer hsCancel()
-
- if err := c.handshake(hsCtx); err != nil {
- c.close()
- return nil, err
- }
-
return c, nil
}
-func (c *client) close() {
- c.tlsMu.Lock()
- conn := c.tlsConn
- c.tlsMu.Unlock()
- if conn != nil {
- conn.Close()
- }
-}
+func (c *client) close() {}
func (c *client) work(ctx context.Context) error {
for ctx.Err() == nil {
@@ -97,8 +90,12 @@ func (c *client) work(ctx context.Context) error {
if err != nil {
return fmt.Errorf("poll: %w", err)
}
- task := resp.Msg
+ for _, w := range resp.Header().Values("X-Warning") {
+ slog.Warn("server warning", "msg", w)
+ }
+
+ task := resp.Msg
slog.Info("task received", "id", task.Id, "repo", task.RepoFullName)
execErr := executor.Run(task.CloneUrl, task.Branch)
@@ -115,113 +112,6 @@ func (c *client) work(ctx context.Context) error {
return fmt.Errorf("complete: %w", err)
}
}
- return ctx.Err()
-}
-
-func (c *client) handshake(ctx context.Context) error {
- stream := c.handle.Handshake(ctx)
-
- // Step 1: send worker public key
- pubKey, err := protocol.LoadPublicKey(c.cfg.PubKeyFile)
- if err != nil {
- return fmt.Errorf("load public key: %w", err)
- }
- if err := stream.Send(&pb.HandshakeIn{
- Step: &pb.HandshakeIn_WorkerChallenge{
- WorkerChallenge: &pb.WorkerChallenge{PublicKey: pubKey},
- },
- }); err != nil {
- return fmt.Errorf("send challenge: %w", err)
- }
- // Step 2: receive server nonce
- out, err := stream.Receive()
- if err != nil {
- return fmt.Errorf("recv server challenge: %w", err)
- }
- sc := out.GetServerChallenge()
- if sc == nil {
- return fmt.Errorf("expected ServerChallenge")
- }
-
- // Extract EKM if server requested channel binding
- var ekm []byte
- if sc.GetBinded() {
- c.tlsMu.Lock()
- conn := c.tlsConn
- c.tlsMu.Unlock()
- if conn == nil {
- return fmt.Errorf("channel binding requested but TLS connection not available")
- }
- state := conn.ConnectionState()
- ekm, err = 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
- }
-
- name := c.cfg.Name
- if name == "" {
- name, _ = os.Hostname()
- }
- if err := stream.Send(&pb.HandshakeIn{
- Step: &pb.HandshakeIn_WorkerProof{
- WorkerProof: &pb.WorkerProof{
- Signature: signature,
- Name: name,
- Version: protocol.VersionProto(),
- Os: protocol.DetectOs(),
- Arch: protocol.DetectArch(),
- Runtime: workerRuntime,
- WorkerTime: timestamppb.Now(),
- },
- },
- }); err != nil {
- return fmt.Errorf("send proof: %w", err)
- }
-
- // Step 4: receive result
- out, err = stream.Receive()
- if err != nil {
- return fmt.Errorf("recv result: %w", err)
- }
- sr := out.GetServerResult()
- if sr == nil {
- return fmt.Errorf("expected ServerResult")
- }
- if sr.Error != nil {
- return fmt.Errorf("%w: %s", protocol.ErrServerRejected, *sr.Error)
- }
-
- for _, w := range sr.GetWarnings() {
- slog.Warn("server warning", "msg", w)
- }
-
- if err := stream.CloseRequest(); err != nil {
- return fmt.Errorf("close stream: %w", err)
- }
- if err := stream.CloseResponse(); err != nil {
- return fmt.Errorf("close stream: %w", err)
- }
-
- v := sr.GetServerVersion()
- slog.Info("handshake ok", "server_version", fmt.Sprintf("%d.%d.%d", v.GetMajor(), v.GetMinor(), v.GetPatch()))
-
- return nil
+ return ctx.Err()
}
diff --git a/cmd/mirumw/config.go b/cmd/mirumw/config.go
index 45e0557..50e82d3 100644
--- a/cmd/mirumw/config.go
+++ b/cmd/mirumw/config.go
@@ -15,11 +15,10 @@ import (
const workerRuntime = "host"
type config struct {
- 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
+ Name string `yaml:"name"`
+ Server string `yaml:"server"`
+ KeyFile string `yaml:"key_file"`
+ TLSCA string `yaml:"tls_ca"` // custom CA cert for self-signed/dev
}
func getConfig(filename string) (*config, error) {
@@ -41,9 +40,6 @@ 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 3ce2596..762b772 100644
--- a/cmd/mirumw/main.go
+++ b/cmd/mirumw/main.go
@@ -38,6 +38,7 @@ func main() {
if !backoff.Wait(ctx) {
break
}
+
continue
}
diff --git a/go.mod b/go.mod
index 889be98..b0cef7d 100644
--- a/go.mod
+++ b/go.mod
@@ -15,7 +15,6 @@ require (
github.com/spf13/cobra v1.8.0
go.starlark.net v0.0.0-20260326113308-fadfc96def35
golang.org/x/crypto v0.49.0
- golang.org/x/net v0.52.0
google.golang.org/protobuf v1.36.11
gopkg.in/yaml.v3 v3.0.1
)
diff --git a/go.sum b/go.sum
index 931ac40..f79ef27 100644
--- a/go.sum
+++ b/go.sum
@@ -94,8 +94,6 @@ golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4=
golang.org/x/crypto v0.49.0/go.mod h1:ErX4dUh2UM+CFYiXZRTcMpEcN8b/1gxEuv3nODoYtCA=
golang.org/x/exp v0.0.0-20250911091902-df9299821621 h1:2id6c1/gto0kaHYyrixvknJ8tUK/Qs5IsmBtrc+FtgU=
golang.org/x/exp v0.0.0-20250911091902-df9299821621/go.mod h1:TwQYMMnGpvZyc+JpB/UAuTNIsVJifOlSkrZkhcvpVUk=
-golang.org/x/net v0.52.0 h1:He/TN1l0e4mmR3QqHMT2Xab3Aj3L9qjbhRm78/6jrW0=
-golang.org/x/net v0.52.0/go.mod h1:R1MAz7uMZxVMualyPXb+VaqGSa3LIaUqk0eEt3w36Sw=
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo=
diff --git a/internal/protocol/handshake.go b/internal/protocol/handshake.go
index 7f54a24..5d020a2 100644
--- a/internal/protocol/handshake.go
+++ b/internal/protocol/handshake.go
@@ -6,12 +6,14 @@ package protocol
import (
"crypto/ed25519"
"crypto/rand"
+ "crypto/tls"
"crypto/x509"
"encoding/pem"
"errors"
"fmt"
+ "math/big"
+ "net/url"
"os"
- "slices"
"time"
)
@@ -21,24 +23,50 @@ var (
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")
+ ErrClockSkew = errors.New("clock skew too large")
)
-const NonceSize = 32
+// WorkerMeta describes the worker for embedding in a self-signed X.509
+// certificate as a URI SAN (mirum:worker?name=...&version=...&...).
+type WorkerMeta struct {
+ Name string
+ Version string
+ Os string
+ Arch string
+ Runtime string
+}
-const (
- EKMLabel = "mirum-handshake"
- EKMLength = 32
-)
+// URI encodes worker metadata as a mirum: URI.
+func (m *WorkerMeta) URI() *url.URL {
+ return &url.URL{
+ Scheme: "mirum",
+ Opaque: "worker",
+ RawQuery: url.Values{
+ "name": {m.Name},
+ "version": {m.Version},
+ "os": {m.Os},
+ "arch": {m.Arch},
+ "runtime": {m.Runtime},
+ }.Encode(),
+ }
+}
-func generateNonce() ([]byte, error) {
- nonce := make([]byte, NonceSize)
- _, err := rand.Read(nonce)
- return nonce, err
+// ParseWorkerMeta extracts WorkerMeta from a certificate's URI SANs.
+// Returns nil if no mirum:worker URI is found.
+func ParseWorkerMeta(cert *x509.Certificate) *WorkerMeta {
+ for _, u := range cert.URIs {
+ if u.Scheme == "mirum" && u.Opaque == "worker" {
+ q := u.Query()
+ return &WorkerMeta{
+ Name: q.Get("name"),
+ Version: q.Get("version"),
+ Os: q.Get("os"),
+ Arch: q.Get("arch"),
+ Runtime: q.Get("runtime"),
+ }
+ }
+ }
+ return nil
}
// LoadPrivateKey reads a PEM-encoded PKCS8 ed25519 private key from path.
@@ -66,95 +94,33 @@ func LoadPrivateKey(path string) (ed25519.PrivateKey, error) {
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)
+// SelfSignedCert generates a self-signed X.509 certificate from an ed25519
+// private key with worker metadata encoded as a URI SAN. The server extracts
+// the public key for authentication, metadata from the URI, and uses NotBefore
+// for clock skew detection.
+func SelfSignedCert(key ed25519.PrivateKey, meta *WorkerMeta) (tls.Certificate, error) {
+ serial, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 128))
if err != nil {
- return nil, fmt.Errorf("read key: %w", err)
+ return tls.Certificate{}, fmt.Errorf("generate serial: %w", err)
}
- block, _ := pem.Decode(data)
- if block == nil {
- return nil, ErrKeyNotPEM
+ now := time.Now()
+ tmpl := &x509.Certificate{
+ SerialNumber: serial,
+ NotBefore: now,
+ NotAfter: now.Add(24 * time.Hour),
+ KeyUsage: x509.KeyUsageDigitalSignature,
+ ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth},
+ URIs: []*url.URL{meta.URI()},
}
- key, err := x509.ParsePKIXPublicKey(block.Bytes)
+ certDER, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, key.Public(), key)
if err != nil {
- return nil, fmt.Errorf("%w: %w", ErrKeyNotPKIX, err)
+ return tls.Certificate{}, fmt.Errorf("create certificate: %w", 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
- serverNonce []byte
- ekm []byte // TLS Exported Keying Material, nil if not bound
-}
-
-func NewServerHandshake() *ServerHandshake {
- return &ServerHandshake{}
-}
-
-// 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.publicKey = publicKey
- h.ekm = ekm
- h.serverNonce, err = generateNonce()
- if err != nil {
- return nil, err
- }
-
- return h.serverNonce, nil
-}
-
-// 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()
- if skew > time.Minute {
- return ErrClockSkew
- }
-
- return nil
-}
-
-// ClientHandshake holds state for the client side of the handshake protocol.
-type ClientHandshake struct {
- privateKey ed25519.PrivateKey
-}
-
-func NewClientHandshake(privateKey ed25519.PrivateKey) *ClientHandshake {
- return &ClientHandshake{privateKey: privateKey}
-}
-
-// PublicKey returns the 32-byte ed25519 public key.
-func (h *ClientHandshake) PublicKey() []byte {
- return h.privateKey.Public().(ed25519.PublicKey)
-}
-
-// 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
- }
- return ed25519.Sign(h.privateKey, slices.Concat(serverNonce, ekm)), nil
+ return tls.Certificate{
+ Certificate: [][]byte{certDER},
+ PrivateKey: key,
+ }, nil
}
diff --git a/internal/protocol/handshake_test.go b/internal/protocol/handshake_test.go
index 14f254a..d5365b8 100644
--- a/internal/protocol/handshake_test.go
+++ b/internal/protocol/handshake_test.go
@@ -8,7 +8,6 @@ import (
"crypto/rand"
"crypto/x509"
"encoding/pem"
- "errors"
"os"
"path/filepath"
"testing"
@@ -59,19 +58,6 @@ func TestLoadPrivateKey(t *testing.T) {
}
}
-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 {
@@ -90,7 +76,6 @@ func TestLoadPrivateKey_NotPEM(t *testing.T) {
}
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)
@@ -101,151 +86,94 @@ func TestLoadPrivateKey_WrongKeyType(t *testing.T) {
}
}
-func generateTestKey(t *testing.T) ed25519.PrivateKey {
- t.Helper()
+func TestSelfSignedCert(t *testing.T) {
_, priv, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
t.Fatal(err)
}
- return priv
-}
-func TestHandshake_Success(t *testing.T) {
- priv := generateTestKey(t)
- pub := priv.Public().(ed25519.PublicKey)
-
- server := NewServerHandshake()
- client := NewClientHandshake(priv)
-
- if got := client.PublicKey(); len(got) != ed25519.PublicKeySize {
- t.Fatalf("public key len = %d, want %d", len(got), ed25519.PublicKeySize)
+ meta := &WorkerMeta{
+ Name: "test-worker",
+ Version: "1.2.3",
+ Os: "linux",
+ Arch: "amd64",
+ Runtime: "host",
}
- serverNonce, err := server.Challenge(pub, nil)
+ cert, err := SelfSignedCert(priv, meta)
if err != nil {
t.Fatal(err)
}
- signature, err := client.Sign(serverNonce, nil)
- if err != nil {
- t.Fatal(err)
- }
-
- if err := server.Verify(signature, time.Now()); err != nil {
- t.Fatal(err)
+ if len(cert.Certificate) != 1 {
+ t.Fatalf("expected 1 cert, got %d", len(cert.Certificate))
}
-}
-
-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)
- serverNonce, err := server.Challenge(pub, ekm)
+ parsed, err := x509.ParseCertificate(cert.Certificate[0])
if err != nil {
t.Fatal(err)
}
- signature, err := client.Sign(serverNonce, ekm)
- if err != nil {
- t.Fatal(err)
+ // Public key matches
+ pubKey, ok := parsed.PublicKey.(ed25519.PublicKey)
+ if !ok {
+ t.Fatal("certificate does not contain an ed25519 public key")
}
-
- if err := server.Verify(signature, time.Now()); err != nil {
- t.Fatal(err)
+ if !pubKey.Equal(priv.Public().(ed25519.PublicKey)) {
+ t.Fatal("public key mismatch")
}
-}
-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)
+ // Key usage
+ if parsed.KeyUsage&x509.KeyUsageDigitalSignature == 0 {
+ t.Fatal("missing DigitalSignature key usage")
}
-}
-
-func TestHandshake_WrongKey(t *testing.T) {
- workerKey := generateTestKey(t)
- otherKey := generateTestKey(t)
-
- server := NewServerHandshake()
- client := NewClientHandshake(workerKey)
-
- 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)
+ if len(parsed.ExtKeyUsage) != 1 || parsed.ExtKeyUsage[0] != x509.ExtKeyUsageClientAuth {
+ t.Fatal("missing ClientAuth extended key usage")
}
-}
-
-func TestHandshake_InvalidPublicKeyLength(t *testing.T) {
- server := NewServerHandshake()
- _, err := server.Challenge([]byte("short"), nil)
- if !errors.Is(err, ErrInvalidPublicKey) {
- t.Fatalf("err = %v, want ErrInvalidPublicKey", err)
+ // NotBefore is recent (used for clock skew)
+ if time.Since(parsed.NotBefore).Abs() > 5*time.Second {
+ t.Fatalf("NotBefore too far from now: %v", parsed.NotBefore)
}
}
-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
+func TestSelfSignedCert_WorkerMeta(t *testing.T) {
+ _, priv, err := ed25519.GenerateKey(rand.Reader)
+ if err != nil {
+ t.Fatal(err)
+ }
- err := server.Verify(signature, time.Now())
- if !errors.Is(err, ErrInvalidSignature) {
- t.Fatalf("err = %v, want ErrInvalidSignature", err)
+ want := &WorkerMeta{
+ Name: "my-worker",
+ Version: "0.5.1",
+ Os: "darwin",
+ Arch: "arm64",
+ Runtime: "docker",
}
-}
-func TestHandshake_ClockSkew(t *testing.T) {
- priv := generateTestKey(t)
- pub := priv.Public().(ed25519.PublicKey)
+ cert, err := SelfSignedCert(priv, want)
+ if err != nil {
+ t.Fatal(err)
+ }
- server := NewServerHandshake()
- client := NewClientHandshake(priv)
+ parsed, err := x509.ParseCertificate(cert.Certificate[0])
+ if err != nil {
+ t.Fatal(err)
+ }
- serverNonce, _ := server.Challenge(pub, nil)
- signature, _ := client.Sign(serverNonce, nil)
+ got := ParseWorkerMeta(parsed)
+ if got == nil {
+ t.Fatal("ParseWorkerMeta returned nil")
+ }
- err := server.Verify(signature, time.Now().Add(-2*time.Minute))
- if !errors.Is(err, ErrClockSkew) {
- t.Fatalf("err = %v, want ErrClockSkew", err)
+ if *got != *want {
+ t.Fatalf("meta mismatch:\n got: %+v\nwant: %+v", got, want)
}
}
-func TestSign_InvalidNonceLength(t *testing.T) {
- priv := generateTestKey(t)
- client := NewClientHandshake(priv)
-
- _, err := client.Sign([]byte("short"), nil)
- if !errors.Is(err, ErrInvalidNonce) {
- t.Fatalf("err = %v, want ErrInvalidNonce", err)
+func TestParseWorkerMeta_NoCert(t *testing.T) {
+ cert := &x509.Certificate{}
+ if meta := ParseWorkerMeta(cert); meta != nil {
+ t.Fatalf("expected nil, got %+v", meta)
}
}
diff --git a/proto/mirum.proto b/proto/mirum.proto
index ff1f759..dd85e26 100644
--- a/proto/mirum.proto
+++ b/proto/mirum.proto
@@ -7,21 +7,17 @@ package mirum;
option go_package = "dimidiumlabs/mirum/internal/protocol/pb";
-import "google/protobuf/timestamp.proto";
-
// Describes the contract between the worker and the server.
// GRPC is the only contract between them, so to implement your own worker,
// you only need to implement this service.
+//
+// Authentication is handled via mTLS: the worker presents a self-signed
+// X.509 certificate containing its ed25519 public key. The server verifies
+// the key against its database during the TLS handshake.
+//
+// Worker metadata (name, version, os, arch) and clock skew detection
+// are embedded in the certificate (URI SAN and NotBefore).
service Mirum {
- // 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 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);
-
// When a worker has free resources, it requests a task from the server.
// The call will block if the server currently has no tasks.
rpc Poll(PollRequest) returns (Task);
@@ -126,66 +122,7 @@ enum Arch {
ARCH_OPENRISC = 32;
}
-// Worker → Server
-message HandshakeIn {
- oneof step {
- WorkerChallenge worker_challenge = 1; // step 1
- WorkerProof worker_proof = 2; // step 3
- }
-}
-
-// Server → Worker
-message HandshakeOut {
- oneof step {
- ServerChallenge server_challenge = 1; // step 2
- ServerResult server_result = 2; // step 4
- }
-}
-
-// Step 1: Worker sends its ed25519 public key (32 bytes).
-message WorkerChallenge {
- bytes public_key = 1;
-}
-
-// Step 2: Server sends a random nonce for the worker to sign.
-message ServerChallenge {
- bytes nonce = 1;
- bool binded = 2; // if true, worker must include TLS EKM in signed data
-}
-
-// Step 3: Worker signs the server nonce and sends metadata.
-message WorkerProof {
- bytes signature = 1;
-
- bytes id = 2;
- string name = 3;
- Version version = 4;
-
- Os os = 5;
- Arch arch = 6;
- string runtime = 7;
-
- google.protobuf.Timestamp worker_time = 8;
-}
-
-// Step 4: Server accepts or rejects the worker.
-message ServerResult {
- // A successful handshake doesn't mean the worker can run:
- // the versions may be incompatible, the clocks may be out of sync,
- // server limits may have been exceeded, or something else entirely.
- optional string error = 1;
-
- Version server_version = 2;
- google.protobuf.Timestamp server_time = 3;
-
- // Non-fatal warnings for the worker to log (clock drift, upcoming
- // deprecations, expiring secrets, known vulnerabilities, etc.).
- repeated string warnings = 4;
-}
-
-message PollRequest {
-
-}
+message PollRequest {}
message Task {
string id = 1;