package agentsync import ( "context" "crypto/sha256" "encoding/hex" "fmt" "io" "log" "net" "net/http" "os" "os/exec" "runtime" "strings" "time" "github.com/mrhid6/keymanager/agent/internal/config" grpcclient "github.com/mrhid6/keymanager/agent/internal/grpc" "github.com/mrhid6/keymanager/agent/internal/grpc/pb" "github.com/mrhid6/keymanager/agent/internal/keys" ) func Run(ctx context.Context, cfg *config.Config, version string) error { client, err := grpcclient.New(cfg.ServerURL, cfg.TLS) if err != nil { return fmt.Errorf("dial grpc: %w", err) } defer client.Close() // Register if we have a pre-reg token if cfg.PreRegToken != "" { log.Println("registering with server...") hostname, _ := os.Hostname() ipAddress := localIP() osInfo := fmt.Sprintf("%s %s", runtime.GOOS, runtime.GOARCH) agentToken, err := client.Register(cfg.ServerID, cfg.PreRegToken, hostname, ipAddress, osInfo) if err != nil { return fmt.Errorf("registration failed: %w", err) } cfg.AgentToken = agentToken cfg.PreRegToken = "" if err := config.Save(cfg); err != nil { return fmt.Errorf("save config: %w", err) } log.Println("registration successful") client.Close() client, err = grpcclient.New(cfg.ServerURL, cfg.TLS) if err != nil { return fmt.Errorf("reconnect: %w", err) } } if cfg.AgentToken == "" { return fmt.Errorf("no agent token available — registration required") } // Start the command stream alongside the poll loop. go runCommandStream(ctx, cfg) ticker := time.NewTicker(cfg.PollInterval) defer ticker.Stop() // Run immediately on startup if err := poll(client, cfg, version); err != nil { log.Printf("poll error: %v", err) } for { select { case <-ctx.Done(): return nil case <-ticker.C: if err := poll(client, cfg, version); err != nil { log.Printf("poll error: %v", err) } } } } func poll(client *grpcclient.Client, cfg *config.Config, version string) error { desired, err := client.SyncKeys(cfg.ServerID, cfg.AgentToken, version) if err != nil { return fmt.Errorf("SyncKeys: %w", err) } current, err := keys.ReadAuthorizedKeys() if err != nil { return fmt.Errorf("read authorized_keys: %w", err) } if !keys.StateChanged(current, desired) { log.Println("authorized_keys unchanged, skipping write") return nil } if err := keys.WriteAuthorizedKeys(desired); err != nil { return fmt.Errorf("write authorized_keys: %w", err) } log.Printf("authorized_keys updated (%d keys)", len(desired)) return nil } // runCommandStream maintains a persistent bidirectional stream with the server // for instant command delivery. Reconnects with exponential backoff on failure. func runCommandStream(ctx context.Context, cfg *config.Config) { backoff := time.Second const maxBackoff = 2 * time.Minute for { select { case <-ctx.Done(): return default: } if err := connectAndHandleStream(ctx, cfg); err != nil { if ctx.Err() != nil { return } log.Printf("command stream error: %v, reconnecting in %s", err, backoff) select { case <-ctx.Done(): return case <-time.After(backoff): } if backoff < maxBackoff { backoff *= 2 } } else { backoff = time.Second } } } func connectAndHandleStream(ctx context.Context, cfg *config.Config) error { client, err := grpcclient.New(cfg.ServerURL, cfg.TLS) if err != nil { return fmt.Errorf("dial: %w", err) } defer client.Close() stream, err := client.CommandStream(ctx) if err != nil { return fmt.Errorf("open stream: %w", err) } if err := stream.Send(&pb.AgentMessage{ ServerId: cfg.ServerID, AgentToken: cfg.AgentToken, Ready: &pb.AgentReady{}, }); err != nil { return fmt.Errorf("send auth: %w", err) } log.Println("command stream connected") for { cmd, err := stream.Recv() if err != nil { return fmt.Errorf("recv: %w", err) } if cmd.GenerateKey != nil { go handleGenerateKey(cfg, cmd) } if cmd.DeleteKey != nil { go handleDeleteKey(cmd) } if cmd.UpdateAgent != nil { go handleUpdateAgent(cmd) } } } func handleDeleteKey(cmd *pb.ServerCommand) { label := cmd.DeleteKey.Label keyPath := fmt.Sprintf("/root/.ssh/keymanager_%s", strings.ReplaceAll(label, " ", "_")) if err := keys.RemoveSSHIdentity(keyPath); err != nil { log.Printf("remove ssh identity failed (cmd=%s): %v", cmd.CommandId, err) } for _, path := range []string{keyPath, keyPath + ".pub"} { if err := os.Remove(path); err != nil && !os.IsNotExist(err) { log.Printf("delete key file %s (cmd=%s): %v", path, cmd.CommandId, err) } } log.Printf("deleted local key files for %q (cmd=%s)", label, cmd.CommandId) } func handleUpdateAgent(cmd *pb.ServerCommand) { u := cmd.UpdateAgent arch := runtime.GOARCH // "amd64" or "arm64" tag := "agent%2Fv" + u.Version binaryURL := fmt.Sprintf("%s/mrhid6/keymanager/releases/download/%s/keymanager-agent-linux-%s", u.GiteaBaseURL, tag, arch) checksumURL := fmt.Sprintf("%s/mrhid6/keymanager/releases/download/%s/checksums.txt", u.GiteaBaseURL, tag) log.Printf("updating agent to v%s from %s (cmd=%s)", u.Version, u.GiteaBaseURL, cmd.CommandId) // Download binary tmpBin := "/tmp/keymanager-agent-update" if err := downloadFile(binaryURL, tmpBin); err != nil { log.Printf("update download failed (cmd=%s): %v", cmd.CommandId, err) return } // Download and verify checksum checksumData, err := httpGetBytes(checksumURL) if err != nil { log.Printf("update checksum fetch failed (cmd=%s): %v", cmd.CommandId, err) return } if err := verifyChecksum(tmpBin, fmt.Sprintf("keymanager-agent-linux-%s", arch), checksumData); err != nil { log.Printf("update checksum mismatch (cmd=%s): %v", cmd.CommandId, err) os.Remove(tmpBin) return } if err := os.Chmod(tmpBin, 0755); err != nil { log.Printf("update chmod failed (cmd=%s): %v", cmd.CommandId, err) return } if err := os.Rename(tmpBin, "/usr/local/bin/keymanager-agent"); err != nil { log.Printf("update replace binary failed (cmd=%s): %v", cmd.CommandId, err) return } log.Printf("agent binary replaced, restarting service (cmd=%s)", cmd.CommandId) exec.Command("systemctl", "restart", "keymanager-agent").Run() } func downloadFile(url, dest string) error { resp, err := http.Get(url) //nolint:gosec if err != nil { return err } defer resp.Body.Close() if resp.StatusCode != http.StatusOK { return fmt.Errorf("HTTP %d from %s", resp.StatusCode, url) } f, err := os.Create(dest) if err != nil { return err } defer f.Close() _, err = io.Copy(f, resp.Body) return err } func httpGetBytes(url string) ([]byte, error) { resp, err := http.Get(url) //nolint:gosec if err != nil { return nil, err } defer resp.Body.Close() if resp.StatusCode != http.StatusOK { return nil, fmt.Errorf("HTTP %d from %s", resp.StatusCode, url) } return io.ReadAll(resp.Body) } func verifyChecksum(filePath, filename string, checksumData []byte) error { f, err := os.Open(filePath) if err != nil { return err } defer f.Close() h := sha256.New() if _, err := io.Copy(h, f); err != nil { return err } actual := hex.EncodeToString(h.Sum(nil)) for _, line := range strings.Split(string(checksumData), "\n") { fields := strings.Fields(line) if len(fields) == 2 && fields[1] == filename { if fields[0] != actual { return fmt.Errorf("expected %s got %s", fields[0], actual) } return nil } } return fmt.Errorf("no checksum entry found for %s", filename) } func handleGenerateKey(cfg *config.Config, cmd *pb.ServerCommand) { g := cmd.GenerateKey label := g.Label keyPath := fmt.Sprintf("/root/.ssh/keymanager_%s", strings.ReplaceAll(label, " ", "_")) opts := keys.KeyGenOptions{ KeyType: g.KeyType, KeySize: g.KeySize, Passphrase: g.Passphrase, Comment: g.Comment, } pubKey, err := keys.GenerateKeyPair(keyPath, opts) if err != nil { log.Printf("key generation failed (cmd=%s): %v", cmd.CommandId, err) return } privKeyData, err := os.ReadFile(keyPath) if err != nil { log.Printf("read private key failed (cmd=%s): %v", cmd.CommandId, err) return } client, err := grpcclient.New(cfg.ServerURL, cfg.TLS) if err != nil { log.Printf("dial for key upload failed (cmd=%s): %v", cmd.CommandId, err) return } defer client.Close() keyID, err := client.UploadGeneratedKey(cfg.ServerID, cfg.AgentToken, pubKey, string(privKeyData), label) if err != nil { log.Printf("key upload failed (cmd=%s): %v", cmd.CommandId, err) return } if err := keys.AddSSHIdentity(keyPath); err != nil { log.Printf("add ssh identity failed (cmd=%s): %v", cmd.CommandId, err) } log.Printf("generated and uploaded key %q (key_id=%s, cmd=%s)", label, keyID, cmd.CommandId) } func localIP() string { addrs, err := net.InterfaceAddrs() if err != nil { return "" } for _, addr := range addrs { if ipNet, ok := addr.(*net.IPNet); ok && !ipNet.IP.IsLoopback() { if ipNet.IP.To4() != nil { return ipNet.IP.String() } } } return "" } // GenerateAndUpload generates an SSH keypair and uploads the public key to the server. func GenerateAndUpload(cfg *config.Config, label string) error { client, err := grpcclient.New(cfg.ServerURL, cfg.TLS) if err != nil { return err } defer client.Close() keyPath := fmt.Sprintf("/root/.ssh/keymanager_%s", strings.ReplaceAll(label, " ", "_")) pubKey, err := keys.GenerateKeyPair(keyPath, keys.KeyGenOptions{Comment: label}) if err != nil { return err } privKeyData, err := os.ReadFile(keyPath) if err != nil { return fmt.Errorf("read private key: %w", err) } keyID, err := client.UploadGeneratedKey(cfg.ServerID, cfg.AgentToken, pubKey, string(privKeyData), label) if err != nil { return err } if err := keys.AddSSHIdentity(keyPath); err != nil { log.Printf("add ssh identity: %v", err) } log.Printf("uploaded generated key %s (key_id=%s)", label, keyID) return nil }