From c157e6e45701a1f6101a69e7b62daee37cf9615f Mon Sep 17 00:00:00 2001 From: Nikolay Govorov Date: Tue, 31 Mar 2026 14:55:51 +0100 Subject: auto reconnect grpc client --- cmd/mirumw/client.go | 143 ++++++++++++++++++++++++++ cmd/mirumw/config.go | 41 ++++++++ cmd/mirumw/main.go | 194 +++++------------------------------ internal/protocol/backoff.go | 41 ++++++++ 4 files changed, 248 insertions(+), 171 deletions(-) create mode 100644 cmd/mirumw/client.go create mode 100644 cmd/mirumw/config.go create mode 100644 internal/protocol/backoff.go diff --git a/cmd/mirumw/client.go b/cmd/mirumw/client.go new file mode 100644 index 0000000..779ab11 --- /dev/null +++ b/cmd/mirumw/client.go @@ -0,0 +1,143 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package main + +import ( + "context" + "crypto/tls" + "fmt" + "log/slog" + "net" + "os" + "time" + + "mrdimidium/mirum/internal/protocol" + "mrdimidium/mirum/internal/protocol/pb" + + "google.golang.org/grpc" + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/credentials/insecure" + "google.golang.org/protobuf/types/known/timestamppb" +) + +type client struct { + cfg *config + conn *grpc.ClientConn + handle pb.MirumClient +} + +func connect(ctx context.Context, cfg *config) (*client, error) { + var creds grpc.DialOption + if cfg.Insecure { + host, _, _ := net.SplitHostPort(cfg.Server) + if !isPrivateHost(host) { + return nil, fmt.Errorf("tls_insecure is only allowed for private/loopback addresses, got %s", host) + } + + slog.Warn("TLS disabled, connection is not encrypted") + creds = grpc.WithTransportCredentials(insecure.NewCredentials()) + } else { + creds = grpc.WithTransportCredentials(credentials.NewTLS(&tls.Config{})) + } + + conn, err := grpc.NewClient(cfg.Server, creds) + if err != nil { + return nil, err + } + + c := &client{ + cfg: cfg, + conn: conn, + handle: pb.NewMirumClient(conn), + } + + 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) + } + + return nil, err + } + + return c, nil +} + +func (c *client) close() error { + return c.conn.Close() +} + +func (c *client) handshake(ctx context.Context) error { + stream, err := c.handle.Handshake(ctx) + if err != nil { + return fmt.Errorf("open stream: %w", err) + } + + // Step 1: send worker nonce + workerNonce, err := protocol.GenerateNonce() + if err != nil { + return fmt.Errorf("generate nonce: %w", err) + } + if err := stream.Send(&pb.HandshakeIn{ + Step: &pb.HandshakeIn_WorkerChallenge{ + WorkerChallenge: &pb.WorkerChallenge{Nonce: workerNonce}, + }, + }); err != nil { + return fmt.Errorf("send challenge: %w", err) + } + + // Step 2: receive server challenge, verify server + out, err := stream.Recv() + if err != nil { + return fmt.Errorf("recv server challenge: %w", err) + } + + sc := out.GetServerChallenge() + if sc == nil { + return fmt.Errorf("expected ServerChallenge") + } + if !protocol.VerifyProof([]byte(c.cfg.Secret), workerNonce, sc.GetNonce(), sc.GetProof()) { + return fmt.Errorf("server proof verification failed") + } + + // Step 3: send worker proof + metadata + name := c.cfg.Name + if name == "" { + name, _ = os.Hostname() + } + if err := stream.Send(&pb.HandshakeIn{ + Step: &pb.HandshakeIn_WorkerProof{ + WorkerProof: &pb.WorkerProof{ + Proof: protocol.ComputeProof([]byte(c.cfg.Secret), sc.GetNonce(), workerNonce), + 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.Recv() + 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("server rejected: %s", *sr.Error) + } + + 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/config.go b/cmd/mirumw/config.go new file mode 100644 index 0000000..f8b079e --- /dev/null +++ b/cmd/mirumw/config.go @@ -0,0 +1,41 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package main + +import ( + "fmt" + "os" + + "gopkg.in/yaml.v3" +) + +// Runtime type of this worker binary. Different worker types +// (mirumw-vm, mirumw-docker, etc.) will have different values. +const workerRuntime = "host" + +type config struct { + Name string `yaml:"name"` + Server string `yaml:"server"` + Secret string `yaml:"secret"` + Insecure bool `yaml:"tls_insecure"` +} + +func getConfig(filename string) (*config, error) { + cfg := &config{ + Server: "localhost:2026", + } + + if filename != "" { + data, err := os.ReadFile(filename) + if err != nil { + return nil, fmt.Errorf("couldn't read config: %w", err) + } + + if err := yaml.Unmarshal(data, &cfg); err != nil { + return nil, fmt.Errorf("couldn't parse config: %w", err) + } + } + + return cfg, nil +} diff --git a/cmd/mirumw/main.go b/cmd/mirumw/main.go index 1299b86..0980a9a 100644 --- a/cmd/mirumw/main.go +++ b/cmd/mirumw/main.go @@ -5,200 +5,52 @@ package main import ( "context" - "crypto/tls" "flag" - "fmt" "log/slog" "net" "os" - "time" "mrdimidium/mirum/internal/protocol" - "mrdimidium/mirum/internal/protocol/pb" "mrdimidium/mirum/internal/supervisor" - - "google.golang.org/grpc" - "google.golang.org/grpc/credentials" - "google.golang.org/grpc/credentials/insecure" - "google.golang.org/protobuf/types/known/timestamppb" - "gopkg.in/yaml.v3" ) -// Runtime type of this worker binary. Different worker types -// (mirumw-vm, mirumw-docker, etc.) will have different values. -const workerRuntime = "host" - -type config struct { - Name string `yaml:"name"` - Server string `yaml:"server"` - Secret string `yaml:"secret"` - Insecure bool `yaml:"tls_insecure"` -} - -func GetConfig(filename string) (*config, error) { - cfg := &config{ - Server: "localhost:2026", - } - - if filename != "" { - data, err := os.ReadFile(filename) - if err != nil { - return nil, fmt.Errorf("couldn't read config: %w", err) - } - - if err := yaml.Unmarshal(data, &cfg); err != nil { - return nil, fmt.Errorf("couldn't parse config: %w", err) - } - } - - return cfg, nil -} - -type client struct { - cfg *config - conn *grpc.ClientConn - handle pb.MirumClient -} - -func ClientConnect(cfg *config) (*client, error) { - var creds grpc.DialOption - if cfg.Insecure { - host, _, _ := net.SplitHostPort(cfg.Server) - if !isPrivateHost(host) { - return nil, fmt.Errorf("tls_insecure is only allowed for private/loopback addresses, got %s", host) - } - - slog.Warn("TLS disabled, connection is not encrypted") - creds = grpc.WithTransportCredentials(insecure.NewCredentials()) - } else { - creds = grpc.WithTransportCredentials(credentials.NewTLS(&tls.Config{})) - } - - conn, err := grpc.NewClient(cfg.Server, creds) - if err != nil { - return nil, err - } - - return &client{ - cfg: cfg, - conn: conn, - handle: pb.NewMirumClient(conn), - }, nil -} - -func (c *client) Close() error { - return c.conn.Close() -} - -func (c *client) Handshake(ctx context.Context) error { - stream, err := c.handle.Handshake(ctx) - if err != nil { - return fmt.Errorf("open stream: %w", err) - } - - // Step 1: send worker nonce - workerNonce, err := protocol.GenerateNonce() - if err != nil { - return fmt.Errorf("generate nonce: %w", err) - } - if err := stream.Send(&pb.HandshakeIn{ - Step: &pb.HandshakeIn_WorkerChallenge{ - WorkerChallenge: &pb.WorkerChallenge{Nonce: workerNonce}, - }, - }); err != nil { - return fmt.Errorf("send challenge: %w", err) - } - - // Step 2: receive server challenge, verify server - out, err := stream.Recv() - if err != nil { - return fmt.Errorf("recv server challenge: %w", err) - } - - sc := out.GetServerChallenge() - if sc == nil { - return fmt.Errorf("expected ServerChallenge") - } - if !protocol.VerifyProof([]byte(c.cfg.Secret), workerNonce, sc.GetNonce(), sc.GetProof()) { - return fmt.Errorf("server proof verification failed") - } - - // Step 3: send worker proof + metadata - name := c.cfg.Name - if name == "" { - name, _ = os.Hostname() - } - if err := stream.Send(&pb.HandshakeIn{ - Step: &pb.HandshakeIn_WorkerProof{ - WorkerProof: &pb.WorkerProof{ - Proof: protocol.ComputeProof([]byte(c.cfg.Secret), sc.GetNonce(), workerNonce), - 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.Recv() - 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("server rejected: %s", *sr.Error) - } - - v := sr.GetServerVersion() - slog.Info("handshake ok", "server_version", fmt.Sprintf("%d.%d.%d", v.GetMajor(), v.GetMinor(), v.GetPatch())) - return nil -} - func main() { configFile := flag.String("config", "", "path to config file") flag.Parse() - sup := supervisor.Detect() - ctx := sup.WaitForStop(context.Background()) - - cfg, err := GetConfig(*configFile) + cfg, err := getConfig(*configFile) if err != nil { slog.Error("config", "err", err) os.Exit(1) } - client, err := ClientConnect(cfg) - if err != nil { - slog.Error("failed grpc client", "err", err) - os.Exit(1) - } - defer func() { - if err := client.Close(); err != nil { - slog.Warn("close grpc client", "err", err) + sup := supervisor.Detect() + ctx := sup.WaitForStop(context.Background()) + + sup.Ready() + go sup.StartWatchdog() + + backoff := protocol.NewBackoff() + + for ctx.Err() == nil { + c, err := connect(ctx, cfg) + if err != nil { + slog.Error("connect failed", "err", err) + if !backoff.Wait(ctx) { + break + } + continue } - }() - handshakeCtx, handshakeCancel := context.WithTimeout(ctx, 30*time.Second) - defer handshakeCancel() - if err := client.Handshake(handshakeCtx); err != nil { - slog.Error("handshake failed", "err", err) - os.Exit(1) - } else { slog.Info("connected", "server", cfg.Server) - } + backoff.Reset() - sup.Ready() - go sup.StartWatchdog() + // TODO: c.Work(ctx) — poll tasks, execute, report + <-ctx.Done() + + c.close() + } - <-ctx.Done() slog.Info("shutting down") sup.Stopping() } diff --git a/internal/protocol/backoff.go b/internal/protocol/backoff.go new file mode 100644 index 0000000..8f07892 --- /dev/null +++ b/internal/protocol/backoff.go @@ -0,0 +1,41 @@ +// Copyright (c) 2026 Nikolay Govorov +// SPDX-License-Identifier: AGPL-3.0-or-later + +package protocol + +import ( + "context" + "math/rand/v2" + "time" +) + +type Backoff struct { + attempt int + Min time.Duration + Max time.Duration +} + +func NewBackoff() *Backoff { + return &Backoff{Min: time.Second, Max: 60 * time.Second} +} + +func (b *Backoff) Reset() { + b.attempt = 0 +} + +// Wait sleeps with exponential backoff + jitter. Returns false if ctx is cancelled. +func (b *Backoff) Wait(ctx context.Context) bool { + d := max(b.Max, b.Min<