aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
Diffstat
-rw-r--r--cmd/mirumw/client.go143+143 −0
-rw-r--r--cmd/mirumw/config.go41+41 −0
-rw-r--r--cmd/mirumw/main.go194+23 −171
-rw-r--r--internal/protocol/backoff.go41+41 −0
4 files changed, 248 insertions, 171 deletions
diff --git a/cmd/mirumw/client.go b/cmd/mirumw/client.go
new file mode 100644
--- /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
--- /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
--- /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<<b.attempt)
+
+ // Add jitter: 50%-100% of the computed duration
+ d = d/2 + time.Duration(rand.Int64N(int64(d/2)))
+
+ b.attempt++
+
+ select {
+ case <-time.After(d):
+ return true
+ case <-ctx.Done():
+ return false
+ }
+}