From cc79aa5f7ec97dcd81f801b5e851dcf8e7463a17 Mon Sep 17 00:00:00 2001 From: Nikolay Govorov Date: Fri, 3 Apr 2026 01:23:05 +0100 Subject: Replace grpc to connectrpc --- .github/workflows/build.yml | 5 +- Taskfile.yml | 24 ++++- cmd/mirumd/main.go | 77 ++++++++------ cmd/mirumd/server_admin.go | 57 +++++----- cmd/mirumd/server_grpc.go | 204 ++++++++++++++++++++---------------- cmd/mirumw/client.go | 98 ++++++++++------- cmd/mirumw/main.go | 2 +- go.mod | 5 +- go.sum | 28 +---- proto/admin.proto | 2 +- proto/buf.gen.yaml | 11 ++ proto/mirum.proto | 2 +- 12 files changed, 292 insertions(+), 223 deletions(-) create mode 100644 proto/buf.gen.yaml diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 09c8118..9126654 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -30,10 +30,11 @@ jobs: - name: Install tools run: | echo 'deb [trusted=yes] https://repo.goreleaser.com/apt/ /' | sudo tee /etc/apt/sources.list.d/goreleaser.list - sudo apt update && sudo apt install nfpm protobuf-compiler + sudo apt update && sudo apt install nfpm sh -c "$(curl --location https://taskfile.dev/install.sh)" -- -d -b /usr/local/bin + curl -sSL "https://github.com/bufbuild/buf/releases/download/v1.67.0/buf-$(uname -s)-$(uname -m)" -o /usr/local/bin/buf && chmod +x /usr/local/bin/buf go install google.golang.org/protobuf/cmd/protoc-gen-go@latest - go install google.golang.org/grpc/cmd/protoc-gen-go-grpc@latest + go install connectrpc.com/connect/cmd/protoc-gen-connect-go@latest - name: Import GPG key uses: crazy-max/ghaction-import-gpg@e89d40939c28e39f97cf32126055eeae86ba74ec # v6.3.0 diff --git a/Taskfile.yml b/Taskfile.yml index 4f7b8cf..ab6a14d 100644 --- a/Taskfile.yml +++ b/Taskfile.yml @@ -13,13 +13,14 @@ env: tasks: proto: - desc: Generate gRPC code from proto files + desc: Generate ConnectRPC code from proto files sources: - proto/*.proto generates: - - internal/protocol/pb/*.pb.go + - internal/protocol/pb/**/*.go cmds: - - protoc --go_out=. --go-grpc_out=. proto/mirum.proto proto/admin.proto + - rm -rf internal/protocol/pb + - cd proto && buf generate lint: desc: Run static checks @@ -51,6 +52,23 @@ tasks: - cp {{.BUILD_DIR}}/mirumd-{{.GOOS}}-{{.GOARCH}} {{.BUILD_DIR}}/mirumd - cp {{.BUILD_DIR}}/mirumw-{{.GOOS}}-{{.GOARCH}} {{.BUILD_DIR}}/mirumw + dev: + desc: Generate dev TLS cert and worker keys + cmds: + - mkdir -p dev + - >- + test -f dev/server.crt || + openssl req -x509 -newkey ec -pkeyopt ec_paramgen_curve:prime256v1 + -keyout dev/server.key -out dev/server.crt + -days 365 -nodes -subj "/CN=mirumd" + -addext "subjectAltName=DNS:mirumd,DNS:localhost" + - >- + 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 requires: diff --git a/cmd/mirumd/main.go b/cmd/mirumd/main.go index 91dd8c5..32dc09c 100644 --- a/cmd/mirumd/main.go +++ b/cmd/mirumd/main.go @@ -5,7 +5,9 @@ package main import ( "context" + "crypto/ed25519" "crypto/tls" + "crypto/x509" "encoding/base64" "fmt" "log/slog" @@ -14,15 +16,16 @@ import ( "os" "time" + "connectrpc.com/connect" "dimidiumlabs/mirum/internal/database" "dimidiumlabs/mirum/internal/forges" + "dimidiumlabs/mirum/internal/protocol/pb" + "dimidiumlabs/mirum/internal/protocol/pb/pbconnect" "dimidiumlabs/mirum/internal/supervisor" "github.com/coreos/go-systemd/v22/activation" "github.com/spf13/cobra" - "google.golang.org/grpc" - "google.golang.org/grpc/credentials/insecure" ) func main() { @@ -176,6 +179,7 @@ func daemon(configFile, socketFlag string) { 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) } @@ -191,6 +195,7 @@ func daemon(configFile, socketFlag string) { } if tlsCfg != nil { + grpcLn = tls.NewListener(grpcLn, tlsCfg) webLn = tls.NewListener(webLn, tlsCfg) } @@ -203,13 +208,13 @@ func daemon(configFile, socketFlag string) { } }() go func() { - if err := grpcSrv.Serve(grpcLn); err != nil { + if err := grpcSrv.Serve(grpcLn); err != nil && err != http.ErrServerClosed { slog.Error("grpc server failed", "err", err) os.Exit(1) } }() go func() { - if err := adminSrv.Serve(adminLn); err != nil { + if err := adminSrv.Serve(adminLn); err != nil && err != http.ErrServerClosed { slog.Error("admin server failed", "err", err) os.Exit(1) } @@ -232,9 +237,9 @@ func daemon(configFile, socketFlag string) { srv.Close() - adminSrv.GracefulStop() + adminSrv.Shutdown(context.Background()) wwwSrv.Shutdown(context.Background()) - grpcSrv.GracefulStop() + grpcSrv.Shutdown(context.Background()) } // listeners returns gRPC, web, and admin listeners. @@ -280,37 +285,39 @@ func listeners(cfg *config) (grpcLn, webLn, adminLn net.Listener, err error) { return grpcLn, webLn, adminLn, nil } -func adminClient(socketPath string) pb.AdminClient { +func adminClient(socketPath string) pbconnect.AdminClient { if socketPath == "" { socketPath = "/run/mirumd/admin.sock" } - conn, err := grpc.NewClient("unix://"+socketPath, - grpc.WithTransportCredentials(insecure.NewCredentials()), + return pbconnect.NewAdminClient( + &http.Client{ + Transport: &http.Transport{ + DialContext: func(ctx context.Context, _, _ string) (net.Conn, error) { + return net.Dial("unix", socketPath) + }, + }, + }, + "http://localhost.unix", ) - if err != nil { - fmt.Fprintln(os.Stderr, err) - os.Exit(1) - } - return pb.NewAdminClient(conn) } func userCreate(socketPath, email, password string) { - resp, err := adminClient(socketPath).CreateUser(context.Background(), &pb.CreateUserRequest{ + resp, err := adminClient(socketPath).CreateUser(context.Background(), connect.NewRequest(&pb.CreateUserRequest{ Email: email, Password: password, - }) + })) if err != nil { fmt.Fprintln(os.Stderr, err) os.Exit(1) } - fmt.Println(resp.Id) + fmt.Println(resp.Msg.Id) } func userSetPassword(socketPath, email, password string) { - _, err := adminClient(socketPath).SetPassword(context.Background(), &pb.SetPasswordRequest{ + _, err := adminClient(socketPath).SetPassword(context.Background(), connect.NewRequest(&pb.SetPasswordRequest{ Email: email, Password: password, - }) + })) if err != nil { fmt.Fprintln(os.Stderr, err) os.Exit(1) @@ -319,9 +326,9 @@ func userSetPassword(socketPath, email, password string) { } func userDelete(socketPath, email string) { - _, err := adminClient(socketPath).DeleteUser(context.Background(), &pb.DeleteUserRequest{ + _, err := adminClient(socketPath).DeleteUser(context.Background(), connect.NewRequest(&pb.DeleteUserRequest{ Email: email, - }) + })) if err != nil { fmt.Fprintln(os.Stderr, err) os.Exit(1) @@ -330,25 +337,35 @@ func userDelete(socketPath, email string) { } func workerAdd(socketPath, pubkeyB64 string) { - pubkey, err := base64.StdEncoding.DecodeString(pubkeyB64) + der, 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, - }) + pubkey, err := x509.ParsePKIXPublicKey(der) + if err != nil { + fmt.Fprintln(os.Stderr, "invalid public key:", err) + os.Exit(1) + } + edKey, ok := pubkey.(ed25519.PublicKey) + if !ok { + fmt.Fprintln(os.Stderr, "not an ed25519 key") + os.Exit(1) + } + resp, err := adminClient(socketPath).WorkerAdd(context.Background(), connect.NewRequest(&pb.WorkerAddRequest{ + PublicKey: edKey, + })) if err != nil { fmt.Fprintln(os.Stderr, err) os.Exit(1) } - fmt.Println(resp.Id) + fmt.Println(resp.Msg.Id) } func workerRevoke(socketPath, id string) { - _, err := adminClient(socketPath).WorkerRevoke(context.Background(), &pb.WorkerRevokeRequest{ + _, err := adminClient(socketPath).WorkerRevoke(context.Background(), connect.NewRequest(&pb.WorkerRevokeRequest{ Id: id, - }) + })) if err != nil { fmt.Fprintln(os.Stderr, err) os.Exit(1) @@ -357,12 +374,12 @@ func workerRevoke(socketPath, id string) { } func workerList(socketPath string) { - resp, err := adminClient(socketPath).WorkerList(context.Background(), &pb.WorkerListRequest{}) + resp, err := adminClient(socketPath).WorkerList(context.Background(), connect.NewRequest(&pb.WorkerListRequest{})) if err != nil { fmt.Fprintln(os.Stderr, err) os.Exit(1) } - for _, w := range resp.Workers { + for _, w := range resp.Msg.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 cf72fbc..673e9f6 100644 --- a/cmd/mirumd/server_admin.go +++ b/cmd/mirumd/server_admin.go @@ -7,66 +7,69 @@ import ( "context" "crypto/ed25519" "fmt" + "net/http" - "dimidiumlabs/mirum/internal/protocol/pb" - - "google.golang.org/grpc" + "connectrpc.com/connect" "google.golang.org/protobuf/types/known/timestamppb" + + "dimidiumlabs/mirum/internal/protocol/pb" + "dimidiumlabs/mirum/internal/protocol/pb/pbconnect" ) -func NewAdminServer(srv *server) *grpc.Server { +func NewAdminServer(srv *server) *http.Server { as := &adminService{srv: srv} - s := grpc.NewServer() - pb.RegisterAdminServer(s, as) - return s + path, handler := pbconnect.NewAdminHandler(as) + mux := http.NewServeMux() + mux.Handle(path, handler) + return &http.Server{Handler: mux} } type adminService struct { - pb.UnimplementedAdminServer + pbconnect.UnimplementedAdminHandler srv *server } -func (a *adminService) CreateUser(ctx context.Context, req *pb.CreateUserRequest) (*pb.CreateUserResponse, error) { - id, err := a.srv.db.CreateUser(ctx, req.Email, req.Password, []byte(a.srv.cfg.Pepper)) +func (a *adminService) CreateUser(ctx context.Context, req *connect.Request[pb.CreateUserRequest]) (*connect.Response[pb.CreateUserResponse], error) { + id, err := a.srv.db.CreateUser(ctx, req.Msg.Email, req.Msg.Password, []byte(a.srv.cfg.Pepper)) if err != nil { return nil, err } - return &pb.CreateUserResponse{Id: id}, nil + return connect.NewResponse(&pb.CreateUserResponse{Id: id}), nil } -func (a *adminService) SetPassword(ctx context.Context, req *pb.SetPasswordRequest) (*pb.SetPasswordResponse, error) { - if err := a.srv.db.SetPassword(ctx, req.Email, req.Password, []byte(a.srv.cfg.Pepper)); err != nil { +func (a *adminService) SetPassword(ctx context.Context, req *connect.Request[pb.SetPasswordRequest]) (*connect.Response[pb.SetPasswordResponse], error) { + if err := a.srv.db.SetPassword(ctx, req.Msg.Email, req.Msg.Password, []byte(a.srv.cfg.Pepper)); err != nil { return nil, err } - return &pb.SetPasswordResponse{}, nil + return connect.NewResponse(&pb.SetPasswordResponse{}), nil } -func (a *adminService) DeleteUser(ctx context.Context, req *pb.DeleteUserRequest) (*pb.DeleteUserResponse, error) { - if err := a.srv.db.DeleteUser(ctx, req.Email); err != nil { +func (a *adminService) DeleteUser(ctx context.Context, req *connect.Request[pb.DeleteUserRequest]) (*connect.Response[pb.DeleteUserResponse], error) { + if err := a.srv.db.DeleteUser(ctx, req.Msg.Email); err != nil { return nil, err } - return &pb.DeleteUserResponse{}, nil + return connect.NewResponse(&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)) +func (a *adminService) WorkerAdd(ctx context.Context, req *connect.Request[pb.WorkerAddRequest]) (*connect.Response[pb.WorkerAddResponse], error) { + if len(req.Msg.PublicKey) != ed25519.PublicKeySize { + return nil, fmt.Errorf("invalid public key: expected %d bytes, got %d", ed25519.PublicKeySize, len(req.Msg.PublicKey)) } - id, err := a.srv.db.AddWorker(ctx, req.PublicKey) + id, err := a.srv.db.AddWorker(ctx, req.Msg.PublicKey) if err != nil { return nil, err } - return &pb.WorkerAddResponse{Id: id}, nil + return connect.NewResponse(&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 { +func (a *adminService) WorkerRevoke(ctx context.Context, req *connect.Request[pb.WorkerRevokeRequest]) (*connect.Response[pb.WorkerRevokeResponse], error) { + if err := a.srv.db.RevokeWorker(ctx, req.Msg.Id); err != nil { return nil, err } - return &pb.WorkerRevokeResponse{}, nil + return connect.NewResponse(&pb.WorkerRevokeResponse{}), nil } -func (a *adminService) WorkerList(ctx context.Context, req *pb.WorkerListRequest) (*pb.WorkerListResponse, error) { +func (a *adminService) WorkerList(ctx context.Context, req *connect.Request[pb.WorkerListRequest]) (*connect.Response[pb.WorkerListResponse], error) { workers, err := a.srv.db.ListWorkers(ctx) if err != nil { return nil, err @@ -79,5 +82,5 @@ func (a *adminService) WorkerList(ctx context.Context, req *pb.WorkerListRequest CreatedAt: timestamppb.New(w.CreatedAt), } } - return &pb.WorkerListResponse{Workers: pbWorkers}, nil + return connect.NewResponse(&pb.WorkerListResponse{Workers: pbWorkers}), nil } diff --git a/cmd/mirumd/server_grpc.go b/cmd/mirumd/server_grpc.go index 08841a9..1630728 100644 --- a/cmd/mirumd/server_grpc.go +++ b/cmd/mirumd/server_grpc.go @@ -8,78 +8,83 @@ import ( "crypto/tls" "fmt" "log/slog" + "net" + "net/http" + "strings" "sync" - "sync/atomic" "time" + "connectrpc.com/connect" + "golang.org/x/net/http2" + "golang.org/x/net/http2/h2c" + "google.golang.org/protobuf/types/known/timestamppb" + "dimidiumlabs/mirum/internal/protocol" "dimidiumlabs/mirum/internal/protocol/pb" - - "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" + "dimidiumlabs/mirum/internal/protocol/pb/pbconnect" ) -// connIDKey is the context key for the unique connection identifier. -type connIDKey struct{} +// connKey is the context key for the underlying net.Conn. +type connKey struct{} -func NewGrpcServer(ctx context.Context, srv *server, tlsCfg *tls.Config) *grpc.Server { +// tlsStateKey is the context key for TLS connection state. +type tlsStateKey struct{} + +func NewGrpcServer(ctx context.Context, srv *server, tlsCfg *tls.Config) *http.Server { gsrv := &grpcService{ ctx: ctx, srv: srv, tls: tlsCfg != nil, } - opts := []grpc.ServerOption{ - grpc.UnaryInterceptor(gsrv.unaryInterceptor), - grpc.StreamInterceptor(gsrv.streamInterceptor), - grpc.StatsHandler(&connTracker{gsrv: gsrv}), + opts := []connect.HandlerOption{ + connect.WithInterceptors(gsrv), } - if tlsCfg != nil { - opts = append(opts, grpc.Creds(credentials.NewTLS(tlsCfg))) + path, handler := pbconnect.NewMirumHandler(gsrv, opts...) + mux := http.NewServeMux() + mux.Handle(path, requireGRPC(tlsMiddleware(handler))) + return &http.Server{ + Handler: h2c.NewHandler(mux, &http2.Server{}), + ConnContext: gsrv.connContext, + ConnState: gsrv.connState, + BaseContext: func(_ net.Listener) context.Context { + return ctx + }, } - s := grpc.NewServer(opts...) - pb.RegisterMirumServer(s, gsrv) - return s } -// grpcService is the gRPC transport adapter over server. +// grpcService is the ConnectRPC transport adapter over server. type grpcService struct { - pb.UnimplementedMirumServer + pbconnect.UnimplementedMirumHandler ctx context.Context srv *server tls bool // whether TLS is enabled (for channel binding) - authedConns sync.Map // conn ID (uint64) → true - nextConnID atomic.Uint64 + authedConns sync.Map // net.Conn → true } -func (g *grpcService) Poll(ctx context.Context, req *pb.PollRequest) (*pb.Task, error) { +func (g *grpcService) Poll(ctx context.Context, req *connect.Request[pb.PollRequest]) (*connect.Response[pb.Task], error) { select { case task, ok := <-g.srv.queue: if !ok { - return nil, status.Error(codes.Unavailable, "server is shutting down") + return nil, connect.NewError(connect.CodeUnavailable, fmt.Errorf("server is shutting down")) } slog.Info("task dispatched", "id", task.Id, "repo", task.RepoFullName) - return task, nil + return connect.NewResponse(task), nil case <-ctx.Done(): return nil, ctx.Err() } } -func (g *grpcService) Complete(ctx context.Context, result *pb.TaskResult) (*pb.CompleteResponse, error) { - if err := g.srv.complete(ctx, result.TaskId, result.Success, result.Error); err != nil { +func (g *grpcService) Complete(ctx context.Context, req *connect.Request[pb.TaskResult]) (*connect.Response[pb.CompleteResponse], error) { + if err := g.srv.complete(ctx, req.Msg.TaskId, req.Msg.Success, req.Msg.Error); err != nil { return nil, err } - return &pb.CompleteResponse{}, nil + return connect.NewResponse(&pb.CompleteResponse{}), nil } -func (g *grpcService) Handshake(stream pb.Mirum_HandshakeServer) error { +func (g *grpcService) Handshake(ctx context.Context, stream *connect.BidiStream[pb.HandshakeIn, pb.HandshakeOut]) error { // Step 1: receive worker public key - in, err := stream.Recv() + in, err := stream.Receive() if err != nil { return fmt.Errorf("recv worker challenge: %w", err) } @@ -89,13 +94,13 @@ func (g *grpcService) Handshake(stream pb.Mirum_HandshakeServer) error { } // Look up the worker in the database - worker, err := g.srv.db.LookupWorker(stream.Context(), wc.GetPublicKey()) + 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(stream.Context()) + ekm := g.extractEKM(ctx) hs := protocol.NewServerHandshake() @@ -116,7 +121,7 @@ func (g *grpcService) Handshake(stream pb.Mirum_HandshakeServer) error { } // Step 3: receive worker signature + metadata - in, err = stream.Recv() + in, err = stream.Receive() if err != nil { return fmt.Errorf("recv worker proof: %w", err) } @@ -148,13 +153,13 @@ func (g *grpcService) Handshake(stream pb.Mirum_HandshakeServer) error { ) // Step 4: accept - if id, ok := stream.Context().Value(connIDKey{}).(uint64); ok { - g.authedConns.Store(id, true) + 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 pb.Mirum_HandshakeServer, errMsg *string, warnings []string) error { +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{ @@ -167,95 +172,112 @@ func (g *grpcService) sendResult(stream pb.Mirum_HandshakeServer, errMsg *string }) } -func (g *grpcService) reject(stream pb.Mirum_HandshakeServer, reason string) error { +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, method string) error { - if method == pb.Mirum_Handshake_FullMethodName { +func (g *grpcService) requireAuth(ctx context.Context, procedure string) error { + if procedure == pbconnect.MirumHandshakeProcedure { return nil } - id, ok := ctx.Value(connIDKey{}).(uint64) + conn, ok := ctx.Value(connKey{}).(net.Conn) if !ok { - return status.Error(codes.Unauthenticated, "handshake required") + return connect.NewError(connect.CodeUnauthenticated, fmt.Errorf("handshake required")) } - if _, ok := g.authedConns.Load(id); !ok { - return status.Error(codes.Unauthenticated, "handshake required") + if _, ok := g.authedConns.Load(conn); !ok { + return connect.NewError(connect.CodeUnauthenticated, fmt.Errorf("handshake required")) } return nil } -func (g *grpcService) unaryInterceptor(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) { - if err := g.requireAuth(ctx, info.FullMethod); err != nil { - return nil, err - } - return handler(ctx, req) -} - -func (g *grpcService) streamInterceptor(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error { - if err := g.requireAuth(ss.Context(), info.FullMethod); err != nil { - return err - } - if info.FullMethod == pb.Mirum_Handshake_FullMethodName { - done := make(chan error, 1) - go func() { done <- handler(srv, ss) }() - select { - case err := <-done: - return err - case <-time.After(30 * time.Second): - return fmt.Errorf("handshake timeout") +// 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) } - return handler(srv, ss) } -// connTracker implements stats.Handler to clean up authedPeers on disconnect. -type connTracker struct { - stats.Handler - gsrv *grpcService +// WrapStreamingClient is a no-op (server-side only). +func (g *grpcService) WrapStreamingClient(next connect.StreamingClientFunc) connect.StreamingClientFunc { + return next } -func (t *connTracker) TagConn(ctx context.Context, info *stats.ConnTagInfo) context.Context { - id := t.gsrv.nextConnID.Add(1) - return context.WithValue(ctx, connIDKey{}, id) +// 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) + } } -func (t *connTracker) TagRPC(ctx context.Context, info *stats.RPCTagInfo) context.Context { - return ctx +// 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) } -func (t *connTracker) HandleRPC(ctx context.Context, s stats.RPCStats) {} - -func (t *connTracker) HandleConn(ctx context.Context, s stats.ConnStats) { - if _, ok := s.(*stats.ConnEnd); ok { - if id, ok := ctx.Value(connIDKey{}).(uint64); ok { - t.gsrv.authedConns.Delete(id) - slog.Debug("peer disconnected", "conn", id) - } +// 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 peer context, +// 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 } - p, ok := peer.FromContext(ctx) - if !ok { - return nil - } - ti, ok := p.AuthInfo.(credentials.TLSInfo) - if !ok { + state, ok := ctx.Value(tlsStateKey{}).(*tls.ConnectionState) + if !ok || state == nil { return nil } - ekm, err := ti.State.ExportKeyingMaterial(protocol.EKMLabel, nil, protocol.EKMLength) + 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 { + 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) + } + next.ServeHTTP(w, r) + }) +} diff --git a/cmd/mirumw/client.go b/cmd/mirumw/client.go index 13d4062..999d34b 100644 --- a/cmd/mirumw/client.go +++ b/cmd/mirumw/client.go @@ -9,29 +9,34 @@ import ( "crypto/x509" "fmt" "log/slog" + "net" + "net/http" "os" "runtime/secret" + "sync" "time" + "connectrpc.com/connect" + "google.golang.org/protobuf/types/known/timestamppb" + "dimidiumlabs/mirum/internal/executor" "dimidiumlabs/mirum/internal/protocol" "dimidiumlabs/mirum/internal/protocol/pb" - - "google.golang.org/grpc" - "google.golang.org/grpc/credentials" - "google.golang.org/grpc/peer" - "google.golang.org/protobuf/types/known/timestamppb" + "dimidiumlabs/mirum/internal/protocol/pb/pbconnect" ) type client struct { - cfg *config - conn *grpc.ClientConn - handle pb.MirumClient + cfg *config + http *http.Client + handle pbconnect.MirumClient + tlsMu sync.Mutex + tlsConn *tls.Conn } -func connect(ctx context.Context, cfg *config) (*client, error) { +func dial(ctx context.Context, cfg *config) (*client, error) { tlsCfg := &tls.Config{ MinVersion: tls.VersionTLS13, // required for reliable EKM (channel binding) + NextProtos: []string{"h2"}, // required for gRPC over TLS (ALPN) } if cfg.TLSCA != "" { caCert, err := os.ReadFile(cfg.TLSCA) @@ -45,43 +50,54 @@ func connect(ctx context.Context, cfg *config) (*client, error) { tlsCfg.RootCAs = pool } - conn, err := grpc.NewClient(cfg.Server, - grpc.WithTransportCredentials(credentials.NewTLS(tlsCfg)), - ) - if err != nil { - return nil, err - } + c := &client{cfg: cfg} - c := &client{ - cfg: cfg, - conn: conn, - handle: pb.NewMirumClient(conn), + 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.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 { - if err := c.close(); err != nil { - slog.Warn("handshake: %w", "err", err) - } - + c.close() return nil, err } return c, nil } -func (c *client) close() error { - return c.conn.Close() +func (c *client) close() { + c.tlsMu.Lock() + conn := c.tlsConn + c.tlsMu.Unlock() + if conn != nil { + conn.Close() + } } func (c *client) work(ctx context.Context) error { for ctx.Err() == nil { - task, err := c.handle.Poll(ctx, &pb.PollRequest{}) + resp, err := c.handle.Poll(ctx, connect.NewRequest(&pb.PollRequest{})) if err != nil { return fmt.Errorf("poll: %w", err) } + task := resp.Msg slog.Info("task received", "id", task.Id, "repo", task.RepoFullName) @@ -95,7 +111,7 @@ func (c *client) work(ctx context.Context) error { slog.Info("task passed", "id", task.Id) } - if _, err := c.handle.Complete(ctx, result); err != nil { + if _, err := c.handle.Complete(ctx, connect.NewRequest(result)); err != nil { return fmt.Errorf("complete: %w", err) } } @@ -103,11 +119,7 @@ func (c *client) work(ctx context.Context) error { } func (c *client) handshake(ctx context.Context) error { - var p peer.Peer - stream, err := c.handle.Handshake(ctx, grpc.Peer(&p)) - if err != nil { - return fmt.Errorf("open stream: %w", err) - } + stream := c.handle.Handshake(ctx) // Step 1: send worker public key pubKey, err := protocol.LoadPublicKey(c.cfg.PubKeyFile) @@ -123,7 +135,7 @@ func (c *client) handshake(ctx context.Context) error { } // Step 2: receive server nonce - out, err := stream.Recv() + out, err := stream.Receive() if err != nil { return fmt.Errorf("recv server challenge: %w", err) } @@ -135,11 +147,14 @@ func (c *client) handshake(ctx context.Context) error { // 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") + c.tlsMu.Lock() + conn := c.tlsConn + c.tlsMu.Unlock() + if conn == nil { + return fmt.Errorf("channel binding requested but TLS connection not available") } - ekm, err = ti.State.ExportKeyingMaterial(protocol.EKMLabel, nil, protocol.EKMLength) + state := conn.ConnectionState() + ekm, err = state.ExportKeyingMaterial(protocol.EKMLabel, nil, protocol.EKMLength) if err != nil { return fmt.Errorf("export keying material: %w", err) } @@ -182,7 +197,7 @@ func (c *client) handshake(ctx context.Context) error { } // Step 4: receive result - out, err = stream.Recv() + out, err = stream.Receive() if err != nil { return fmt.Errorf("recv result: %w", err) } @@ -198,6 +213,13 @@ func (c *client) handshake(ctx context.Context) error { 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 diff --git a/cmd/mirumw/main.go b/cmd/mirumw/main.go index ebb1d5a..3ce2596 100644 --- a/cmd/mirumw/main.go +++ b/cmd/mirumw/main.go @@ -32,7 +32,7 @@ func main() { backoff := protocol.NewBackoff() for ctx.Err() == nil { - c, err := connect(ctx, cfg) + c, err := dial(ctx, cfg) if err != nil { slog.Error("connect failed", "err", err) if !backoff.Wait(ctx) { diff --git a/go.mod b/go.mod index 5237b98..62d7700 100644 --- a/go.mod +++ b/go.mod @@ -3,13 +3,14 @@ module dimidiumlabs/mirum go 1.26.1 require ( + connectrpc.com/connect v1.19.1 github.com/coreos/go-systemd/v22 v22.7.0 github.com/jackc/pgx/v5 v5.9.1 github.com/jackc/tern/v2 v2.3.6 github.com/spf13/cobra v1.8.0 go.starlark.net v0.0.0-20260326113308-fadfc96def35 golang.org/x/crypto v0.49.0 - google.golang.org/grpc v1.79.3 + golang.org/x/net v0.52.0 google.golang.org/protobuf v1.36.11 gopkg.in/yaml.v3 v3.0.1 ) @@ -30,9 +31,7 @@ require ( github.com/shopspring/decimal v1.4.0 // indirect github.com/spf13/cast v1.7.0 // indirect github.com/spf13/pflag v1.0.5 // indirect - golang.org/x/net v0.52.0 // indirect golang.org/x/sync v0.20.0 // indirect golang.org/x/sys v0.42.0 // indirect golang.org/x/text v0.35.0 // indirect - google.golang.org/genproto/googleapis/rpc v0.0.0-20260330182312-d5a96adf58d8 // indirect ) diff --git a/go.sum b/go.sum index 12406b8..02f3526 100644 --- a/go.sum +++ b/go.sum @@ -1,3 +1,5 @@ +connectrpc.com/connect v1.19.1 h1:R5M57z05+90EfEvCY1b7hBxDVOUl45PrtXtAV2fOC14= +connectrpc.com/connect v1.19.1/go.mod h1:tN20fjdGlewnSFeZxLKb0xwIZ6ozc3OQs2hTXy4du9w= dario.cat/mergo v1.0.1 h1:Ra4+bf83h2ztPIQYNP99R6m+Y7KfnARDfID+a+vLl4s= dario.cat/mergo v1.0.1/go.mod h1:uNxQE+84aUszobStD9th8a29P2fMDhsBdgRYvZOxGmk= github.com/Masterminds/goutils v1.1.1 h1:5nUrii3FMTL5diU80unEVvNevw1nH4+ZV4DSLVJLSYI= @@ -6,8 +8,6 @@ github.com/Masterminds/semver/v3 v3.3.0 h1:B8LGeaivUe71a5qox1ICM/JLl0NqZSW5CHyL+ github.com/Masterminds/semver/v3 v3.3.0/go.mod h1:4V+yj/TJE1HU9XfppCwVMZq3I84lprf4nC11bSS5beM= github.com/Masterminds/sprig/v3 v3.3.0 h1:mQh0Yrg1XPo6vjYXgtf5OtijNAKJRNcTdOOGZe3tPhs= github.com/Masterminds/sprig/v3 v3.3.0/go.mod h1:Zy1iXRYNqNLUolqCpL4uhk6SHUMAOSCzdgBfDb35Lz0= -github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= -github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= github.com/coreos/go-systemd/v22 v22.7.0 h1:LAEzFkke61DFROc7zNLX/WA2i5J8gYqe0rSj9KI28KA= github.com/coreos/go-systemd/v22 v22.7.0/go.mod h1:xNUYtjHu2EDXbsxz1i41wouACIwT7Ybq9o0BQhMwD0w= github.com/cpuguy83/go-md2man/v2 v2.0.3/go.mod h1:tgQtvFlXSQOSOSIRvRPT7W67SCa46tRHOmNcaadrF8o= @@ -16,12 +16,6 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/frankban/quicktest v1.14.6 h1:7Xjx+VpznH+oBnejlPUj8oUpdxnVs4f8XU8WnHkI4W8= github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0= -github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= -github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= -github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag= -github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE= -github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= -github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= @@ -66,18 +60,6 @@ github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UV github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= -go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= -go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y= -go.opentelemetry.io/otel v1.39.0 h1:8yPrr/S0ND9QEfTfdP9V+SiwT4E0G7Y5MO7p85nis48= -go.opentelemetry.io/otel v1.39.0/go.mod h1:kLlFTywNWrFyEdH0oj2xK0bFYZtHRYUdv1NklR/tgc8= -go.opentelemetry.io/otel/metric v1.39.0 h1:d1UzonvEZriVfpNKEVmHXbdf909uGTOQjA0HF0Ls5Q0= -go.opentelemetry.io/otel/metric v1.39.0/go.mod h1:jrZSWL33sD7bBxg1xjrqyDjnuzTUB0x1nBERXd7Ftcs= -go.opentelemetry.io/otel/sdk v1.39.0 h1:nMLYcjVsvdui1B/4FRkwjzoRVsMK8uL/cj0OyhKzt18= -go.opentelemetry.io/otel/sdk v1.39.0/go.mod h1:vDojkC4/jsTJsE+kh+LXYQlbL8CgrEcwmt1ENZszdJE= -go.opentelemetry.io/otel/sdk/metric v1.39.0 h1:cXMVVFVgsIf2YL6QkRF4Urbr/aMInf+2WKg+sEJTtB8= -go.opentelemetry.io/otel/sdk/metric v1.39.0/go.mod h1:xq9HEVH7qeX69/JnwEfp6fVq5wosJsY1mt4lLfYdVew= -go.opentelemetry.io/otel/trace v1.39.0 h1:2d2vfpEDmCJ5zVYz7ijaJdOF59xLomrvj7bjt6/qCJI= -go.opentelemetry.io/otel/trace v1.39.0/go.mod h1:88w4/PnZSazkGzz/w84VHpQafiU4EtqqlVdxWy+rNOA= go.starlark.net v0.0.0-20260326113308-fadfc96def35 h1:VYAqieSOJNxBDX8KJneTAwvdf4J4zRDE2u+UFXtt9h4= go.starlark.net v0.0.0-20260326113308-fadfc96def35/go.mod h1:Iue6g6iirlfLoVi/DYCi5/x0h/bAOuWF3dULTKpt2Vo= golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4= @@ -90,12 +72,6 @@ golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo= golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/text v0.35.0 h1:JOVx6vVDFokkpaq1AEptVzLTpDe9KGpj5tR4/X+ybL8= golang.org/x/text v0.35.0/go.mod h1:khi/HExzZJ2pGnjenulevKNX1W67CUy0AsXcNubPGCA= -gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk= -gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E= -google.golang.org/genproto/googleapis/rpc v0.0.0-20260330182312-d5a96adf58d8 h1:OHkuo1i98/05rzpm9NBbfEtpJH/k3abEgZUKaAuCI7Y= -google.golang.org/genproto/googleapis/rpc v0.0.0-20260330182312-d5a96adf58d8/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= -google.golang.org/grpc v1.79.3 h1:sybAEdRIEtvcD68Gx7dmnwjZKlyfuc61Dyo9pGXXkKE= -google.golang.org/grpc v1.79.3/go.mod h1:KmT0Kjez+0dde/v2j9vzwoAScgEPx/Bw1CYChhHLrHQ= google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= diff --git a/proto/admin.proto b/proto/admin.proto index 8f628aa..b022bab 100644 --- a/proto/admin.proto +++ b/proto/admin.proto @@ -5,7 +5,7 @@ syntax = "proto3"; package mirum; -option go_package = "internal/protocol/pb"; +option go_package = "dimidiumlabs/mirum/internal/protocol/pb"; import "google/protobuf/timestamp.proto"; diff --git a/proto/buf.gen.yaml b/proto/buf.gen.yaml new file mode 100644 index 0000000..f0e725d --- /dev/null +++ b/proto/buf.gen.yaml @@ -0,0 +1,11 @@ +# Copyright (c) 2026 Nikolay Govorov +# SPDX-License-Identifier: AGPL-3.0-or-later + +version: v2 +plugins: + - local: protoc-gen-go + out: ../internal/protocol/pb + opt: paths=source_relative + - local: protoc-gen-connect-go + out: ../internal/protocol/pb + opt: paths=source_relative diff --git a/proto/mirum.proto b/proto/mirum.proto index 69a856d..ff1f759 100644 --- a/proto/mirum.proto +++ b/proto/mirum.proto @@ -5,7 +5,7 @@ syntax = "proto3"; package mirum; -option go_package = "internal/protocol/pb"; +option go_package = "dimidiumlabs/mirum/internal/protocol/pb"; import "google/protobuf/timestamp.proto"; -- Gilti