From 56a0be15a8858897050f64930d4a364723834e63 Mon Sep 17 00:00:00 2001 From: domrichardson <100129001+domrichardson@users.noreply.github.com> Date: Wed, 24 Jun 2026 14:33:17 +0100 Subject: [PATCH] feat: update agent button on server page --- cmd/main.go | 2 +- internal/grpc/client.go | 7 +- internal/grpc/pb/keymanager.pb.go | 11 ++- internal/sync/sync.go | 113 ++++++++++++++++++++++++++++-- 4 files changed, 122 insertions(+), 11 deletions(-) diff --git a/cmd/main.go b/cmd/main.go index 865f601..f49ff27 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -33,7 +33,7 @@ func main() { defer stop() log.Printf("keymanager-agent %s starting (server=%s, poll=%s)", Version, cfg.ServerURL, cfg.PollInterval) - if err := agentsync.Run(ctx, cfg); err != nil { + if err := agentsync.Run(ctx, cfg, Version); err != nil { log.Fatalf("agent error: %v", err) } } diff --git a/internal/grpc/client.go b/internal/grpc/client.go index c3d8199..7a44926 100644 --- a/internal/grpc/client.go +++ b/internal/grpc/client.go @@ -73,13 +73,14 @@ func (c *Client) Register(serverID, preRegToken, hostname, ipAddress, osInfo str return resp.AgentToken, nil } -func (c *Client) SyncKeys(serverID, agentToken string) ([]string, error) { +func (c *Client) SyncKeys(serverID, agentToken, version string) ([]string, error) { ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) defer cancel() resp, err := c.client.SyncKeys(ctx, &pb.SyncRequest{ - ServerId: serverID, - AgentToken: agentToken, + ServerId: serverID, + AgentToken: agentToken, + AgentVersion: version, }) if err != nil { return nil, err diff --git a/internal/grpc/pb/keymanager.pb.go b/internal/grpc/pb/keymanager.pb.go index 514c8e9..da52687 100644 --- a/internal/grpc/pb/keymanager.pb.go +++ b/internal/grpc/pb/keymanager.pb.go @@ -23,8 +23,9 @@ type RegisterResponse struct { } type SyncRequest struct { - ServerId string `json:"server_id"` - AgentToken string `json:"agent_token"` + ServerId string `json:"server_id"` + AgentToken string `json:"agent_token"` + AgentVersion string `json:"agent_version,omitempty"` } type SyncResponse struct { @@ -49,12 +50,18 @@ type ServerCommand struct { CommandId string `json:"command_id"` GenerateKey *GenerateKeyCmd `json:"generate_key,omitempty"` DeleteKey *DeleteKeyCmd `json:"delete_key,omitempty"` + UpdateAgent *UpdateAgentCmd `json:"update_agent,omitempty"` } type DeleteKeyCmd struct { Label string `json:"label"` } +type UpdateAgentCmd struct { + Version string `json:"version"` + GiteaBaseURL string `json:"gitea_base_url"` +} + type GenerateKeyCmd struct { Label string `json:"label"` KeyType string `json:"key_type,omitempty"` diff --git a/internal/sync/sync.go b/internal/sync/sync.go index c815245..8985891 100644 --- a/internal/sync/sync.go +++ b/internal/sync/sync.go @@ -2,10 +2,15 @@ package agentsync import ( "context" + "crypto/sha256" + "encoding/hex" "fmt" + "io" "log" "net" + "net/http" "os" + "os/exec" "runtime" "strings" "time" @@ -16,7 +21,7 @@ import ( "github.com/mrhid6/keymanager/agent/internal/keys" ) -func Run(ctx context.Context, cfg *config.Config) error { +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) @@ -60,7 +65,7 @@ func Run(ctx context.Context, cfg *config.Config) error { defer ticker.Stop() // Run immediately on startup - if err := poll(client, cfg); err != nil { + if err := poll(client, cfg, version); err != nil { log.Printf("poll error: %v", err) } @@ -69,15 +74,15 @@ func Run(ctx context.Context, cfg *config.Config) error { case <-ctx.Done(): return nil case <-ticker.C: - if err := poll(client, cfg); err != nil { + if err := poll(client, cfg, version); err != nil { log.Printf("poll error: %v", err) } } } } -func poll(client *grpcclient.Client, cfg *config.Config) error { - desired, err := client.SyncKeys(cfg.ServerID, cfg.AgentToken) +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) } @@ -165,6 +170,9 @@ func connectAndHandleStream(ctx context.Context, cfg *config.Config) error { if cmd.DeleteKey != nil { go handleDeleteKey(cmd) } + if cmd.UpdateAgent != nil { + go handleUpdateAgent(cmd) + } } } @@ -184,6 +192,101 @@ func handleDeleteKey(cmd *pb.ServerCommand) { 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