aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
Diffstat
-rw-r--r--.github/workflows/build.yml5+3 −2
-rw-r--r--Taskfile.yml24+21 −3
-rw-r--r--cmd/mirumd/main.go77+47 −30
-rw-r--r--cmd/mirumd/server_admin.go57+30 −27
-rw-r--r--cmd/mirumd/server_grpc.go204+113 −91
-rw-r--r--cmd/mirumw/client.go98+60 −38
-rw-r--r--cmd/mirumw/main.go2+1 −1
-rw-r--r--go.mod5+2 −3
-rw-r--r--go.sum28+2 −26
-rw-r--r--proto/admin.proto2+1 −1
-rw-r--r--proto/buf.gen.yaml11+11 −0
-rw-r--r--proto/mirum.proto2+1 −1
12 files changed, 292 insertions, 223 deletions
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
--- /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";