From 0b3e983f6b8a80d88e92980cbfc0937cf381ea7f Mon Sep 17 00:00:00 2001 From: Nikolay Govorov Date: Sat, 4 Apr 2026 04:53:12 +0100 Subject: Use mTLS with ed25519 keys for handshake --- Taskfile.yml | 5 +- cmd/mirumd/config.go | 16 +- cmd/mirumd/main.go | 72 +++---- cmd/mirumd/server_admin.go | 13 +- cmd/mirumd/server_grpc.go | 284 ++++++---------------------- cmd/mirumd/server_web.go | 111 ++++++----- cmd/mirumw/client.go | 194 ++++--------------- cmd/mirumw/config.go | 12 +- cmd/mirumw/main.go | 1 + go.mod | 1 - go.sum | 2 - internal/protocol/handshake.go | 162 +++++++--------- internal/protocol/handshake_test.go | 180 ++++++------------ proto/mirum.proto | 79 +------- 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; -- Gilti