Native Windows deployment: config file, service, signed auto-update
- internal/config: .env-style config file (gpu-turnstile.env next to the exe, -config flag or GPU_TURNSTILE_CONFIG); process env overrides file. - internal/service: Windows service via golang.org/x/sys/windows/svc — graceful SCM stop, 'service install/remove' commands, restart-on-failure recovery (also applies staged updates). First external dependency, Windows-only; Linux/Docker build unaffected (go.mod stays at 1.23). - internal/update: polls the Gitea releases API, verifies the Ed25519 signature of the downloaded binary against an embedded public key (openssl-signed by CI), swaps it in next to the running exe, and once the GPU lock is idle exits with code 3 so service recovery restarts onto the new version. Dev builds and empty pubkey never update. - CI: tag builds additionally produce gpu-turnstile.exe + .sig + .sha256 attached to a Gitea release. - LOG_FILE env var so the service has somewhere to log.
This commit is contained in:
@@ -0,0 +1,192 @@
|
||||
// Package config loads gpu-turnstile's configuration from environment
|
||||
// variables and an optional .env-style config file.
|
||||
package config
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Config holds every gpu-turnstile setting.
|
||||
type Config struct {
|
||||
ListenOllama string
|
||||
ListenComfy string
|
||||
OllamaURL string
|
||||
ComfyURL string
|
||||
UnloadTimeout time.Duration
|
||||
JobTimeout time.Duration
|
||||
LLMWaitTimeout time.Duration
|
||||
|
||||
UnloadPollInterval time.Duration
|
||||
HistoryPollInterval time.Duration
|
||||
ProbeTimeout time.Duration
|
||||
FreeTimeout time.Duration
|
||||
WarmTimeout time.Duration
|
||||
ShutdownTimeout time.Duration
|
||||
BackoffInitial time.Duration
|
||||
BackoffMax time.Duration
|
||||
PromptCaptureLimit int64
|
||||
|
||||
AutoUpdate bool
|
||||
UpdateInterval time.Duration
|
||||
UpdateRepo string
|
||||
UpdateAsset string
|
||||
|
||||
WarmModel string
|
||||
LogLevel slog.Level
|
||||
LogJSON bool
|
||||
LogFile string
|
||||
}
|
||||
|
||||
// Defaults returns the configuration used when neither the environment nor
|
||||
// a config file sets a value.
|
||||
func Defaults() Config {
|
||||
return Config{
|
||||
ListenOllama: ":11434",
|
||||
ListenComfy: ":8188",
|
||||
OllamaURL: "http://127.0.0.1:11435",
|
||||
ComfyURL: "http://127.0.0.1:8189",
|
||||
UnloadTimeout: time.Minute,
|
||||
JobTimeout: 15 * time.Minute,
|
||||
LLMWaitTimeout: 10 * time.Minute,
|
||||
|
||||
UnloadPollInterval: 500 * time.Millisecond,
|
||||
HistoryPollInterval: time.Second,
|
||||
ProbeTimeout: 5 * time.Second,
|
||||
FreeTimeout: 30 * time.Second,
|
||||
WarmTimeout: 2 * time.Minute,
|
||||
ShutdownTimeout: 10 * time.Second,
|
||||
BackoffInitial: time.Second,
|
||||
BackoffMax: time.Minute,
|
||||
PromptCaptureLimit: 64 * 1024,
|
||||
|
||||
AutoUpdate: true,
|
||||
UpdateInterval: 6 * time.Hour,
|
||||
UpdateRepo: "https://git.rambossek.at/PUBLIC/gpu-turnstile",
|
||||
UpdateAsset: "gpu-turnstile.exe",
|
||||
|
||||
LogLevel: slog.LevelWarn,
|
||||
}
|
||||
}
|
||||
|
||||
// ParseEnvFile parses a .env-style file: KEY=VALUE lines, blank lines and
|
||||
// #-comments are ignored, no quoting. A line without '=' is an error.
|
||||
func ParseEnvFile(r io.Reader) (map[string]string, error) {
|
||||
values := make(map[string]string)
|
||||
scanner := bufio.NewScanner(r)
|
||||
scanner.Buffer(make([]byte, 64*1024), 1024*1024)
|
||||
lineNo := 0
|
||||
for scanner.Scan() {
|
||||
lineNo++
|
||||
line := strings.TrimSpace(scanner.Text())
|
||||
if line == "" || strings.HasPrefix(line, "#") {
|
||||
continue
|
||||
}
|
||||
key, value, ok := strings.Cut(line, "=")
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("line %d: expected KEY=VALUE", lineNo)
|
||||
}
|
||||
key = strings.TrimSpace(key)
|
||||
if key == "" {
|
||||
return nil, fmt.Errorf("line %d: empty key", lineNo)
|
||||
}
|
||||
values[key] = strings.TrimSpace(value)
|
||||
}
|
||||
return values, scanner.Err()
|
||||
}
|
||||
|
||||
func envDuration(getenv func(string) string, name string, dst *time.Duration) error {
|
||||
v := getenv(name)
|
||||
if v == "" {
|
||||
return nil
|
||||
}
|
||||
d, err := time.ParseDuration(v)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s: %w", name, err)
|
||||
}
|
||||
*dst = d
|
||||
return nil
|
||||
}
|
||||
|
||||
// Load overlays values from getenv onto the Defaults. Unknown keys are
|
||||
// ignored. Invalid values are fatal.
|
||||
func Load(getenv func(string) string) (Config, error) {
|
||||
cfg := Defaults()
|
||||
for _, e := range []struct {
|
||||
name string
|
||||
dst *string
|
||||
}{
|
||||
{"LISTEN_OLLAMA", &cfg.ListenOllama},
|
||||
{"LISTEN_COMFY", &cfg.ListenComfy},
|
||||
{"OLLAMA_URL", &cfg.OllamaURL},
|
||||
{"COMFY_URL", &cfg.ComfyURL},
|
||||
{"WARM_MODEL", &cfg.WarmModel},
|
||||
{"UPDATE_REPO", &cfg.UpdateRepo},
|
||||
{"UPDATE_ASSET", &cfg.UpdateAsset},
|
||||
{"LOG_FILE", &cfg.LogFile},
|
||||
} {
|
||||
if v := getenv(e.name); v != "" {
|
||||
*e.dst = v
|
||||
}
|
||||
}
|
||||
for _, e := range []struct {
|
||||
name string
|
||||
dst *time.Duration
|
||||
}{
|
||||
{"UNLOAD_TIMEOUT", &cfg.UnloadTimeout},
|
||||
{"JOB_TIMEOUT", &cfg.JobTimeout},
|
||||
{"LLM_WAIT_TIMEOUT", &cfg.LLMWaitTimeout},
|
||||
{"UNLOAD_POLL_INTERVAL", &cfg.UnloadPollInterval},
|
||||
{"HISTORY_POLL_INTERVAL", &cfg.HistoryPollInterval},
|
||||
{"PROBE_TIMEOUT", &cfg.ProbeTimeout},
|
||||
{"FREE_TIMEOUT", &cfg.FreeTimeout},
|
||||
{"WARM_TIMEOUT", &cfg.WarmTimeout},
|
||||
{"SHUTDOWN_TIMEOUT", &cfg.ShutdownTimeout},
|
||||
{"BACKOFF_INITIAL", &cfg.BackoffInitial},
|
||||
{"BACKOFF_MAX", &cfg.BackoffMax},
|
||||
{"UPDATE_INTERVAL", &cfg.UpdateInterval},
|
||||
} {
|
||||
if err := envDuration(getenv, e.name, e.dst); err != nil {
|
||||
return cfg, err
|
||||
}
|
||||
}
|
||||
if v := getenv("PROMPT_CAPTURE_LIMIT"); v != "" {
|
||||
n, err := strconv.ParseInt(v, 10, 64)
|
||||
if err != nil || n < 0 {
|
||||
return cfg, fmt.Errorf("PROMPT_CAPTURE_LIMIT: must be a non-negative integer (bytes)")
|
||||
}
|
||||
cfg.PromptCaptureLimit = n
|
||||
}
|
||||
if v := getenv("AUTO_UPDATE"); v != "" {
|
||||
b, err := strconv.ParseBool(v)
|
||||
if err != nil {
|
||||
return cfg, fmt.Errorf("AUTO_UPDATE: must be a boolean (true/false)")
|
||||
}
|
||||
cfg.AutoUpdate = b
|
||||
}
|
||||
// LOGLEVEL is the canonical spelling; LOG_LEVEL is kept as an alias.
|
||||
logLevelValue := getenv("LOGLEVEL")
|
||||
if logLevelValue == "" {
|
||||
logLevelValue = getenv("LOG_LEVEL")
|
||||
}
|
||||
if logLevelValue != "" {
|
||||
var level slog.Level
|
||||
if err := level.UnmarshalText([]byte(logLevelValue)); err != nil {
|
||||
return cfg, fmt.Errorf("LOGLEVEL: %w", err)
|
||||
}
|
||||
cfg.LogLevel = level
|
||||
}
|
||||
switch strings.ToLower(getenv("LOG_FORMAT")) {
|
||||
case "", "text":
|
||||
case "json":
|
||||
cfg.LogJSON = true
|
||||
default:
|
||||
return cfg, fmt.Errorf("LOG_FORMAT: must be \"text\" or \"json\"")
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestDefaults(t *testing.T) {
|
||||
cfg, err := Load(func(string) string { return "" })
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if cfg.ListenOllama != ":11434" || cfg.ListenComfy != ":8188" {
|
||||
t.Fatalf("listen addrs = %s %s", cfg.ListenOllama, cfg.ListenComfy)
|
||||
}
|
||||
if cfg.UnloadTimeout != time.Minute || cfg.JobTimeout != 15*time.Minute {
|
||||
t.Fatalf("timeouts = %v %v", cfg.UnloadTimeout, cfg.JobTimeout)
|
||||
}
|
||||
if !cfg.AutoUpdate || cfg.UpdateInterval != 6*time.Hour {
|
||||
t.Fatalf("update = %v %v", cfg.AutoUpdate, cfg.UpdateInterval)
|
||||
}
|
||||
if cfg.LogLevel != slog.LevelWarn {
|
||||
t.Fatalf("log level = %v", cfg.LogLevel)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseEnvFile(t *testing.T) {
|
||||
input := `# comment
|
||||
OLLAMA_URL=http://host:11435
|
||||
|
||||
LOGLEVEL=debug
|
||||
SPACED = value with spaces
|
||||
`
|
||||
values, err := ParseEnvFile(strings.NewReader(input))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if values["OLLAMA_URL"] != "http://host:11435" {
|
||||
t.Fatalf("OLLAMA_URL = %q", values["OLLAMA_URL"])
|
||||
}
|
||||
if values["LOGLEVEL"] != "debug" {
|
||||
t.Fatalf("LOGLEVEL = %q", values["LOGLEVEL"])
|
||||
}
|
||||
if values["SPACED"] != "value with spaces" {
|
||||
t.Fatalf("SPACED = %q", values["SPACED"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseEnvFileMalformed(t *testing.T) {
|
||||
_, err := ParseEnvFile(strings.NewReader("OK=1\nNOT_A_PAIR\n"))
|
||||
if err == nil || !strings.Contains(err.Error(), "line 2") {
|
||||
t.Fatalf("err = %v, want line 2 error", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnvOverridesFile(t *testing.T) {
|
||||
file := map[string]string{"OLLAMA_URL": "http://file:1", "UNLOAD_TIMEOUT": "42s"}
|
||||
env := map[string]string{"OLLAMA_URL": "http://env:2"}
|
||||
getenv := func(k string) string {
|
||||
if v := env[k]; v != "" {
|
||||
return v
|
||||
}
|
||||
return file[k]
|
||||
}
|
||||
cfg, err := Load(getenv)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if cfg.OllamaURL != "http://env:2" {
|
||||
t.Fatalf("OllamaURL = %q, want env value", cfg.OllamaURL)
|
||||
}
|
||||
if cfg.UnloadTimeout != 42*time.Second {
|
||||
t.Fatalf("UnloadTimeout = %v, want file value", cfg.UnloadTimeout)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadErrors(t *testing.T) {
|
||||
for _, tc := range []struct{ key, value string }{
|
||||
{"UNLOAD_TIMEOUT", "bogus"},
|
||||
{"PROMPT_CAPTURE_LIMIT", "-5"},
|
||||
{"AUTO_UPDATE", "maybe"},
|
||||
{"LOGLEVEL", "shouty"},
|
||||
{"LOG_FORMAT", "yaml"},
|
||||
} {
|
||||
_, err := Load(func(k string) string {
|
||||
if k == tc.key {
|
||||
return tc.value
|
||||
}
|
||||
return ""
|
||||
})
|
||||
if err == nil {
|
||||
t.Errorf("%s=%s: expected error", tc.key, tc.value)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -47,6 +47,8 @@ type Config struct {
|
||||
// LogColor enables ANSI colors in per-request log lines. Ignored when
|
||||
// the log level is above INFO (request lines are not emitted at all).
|
||||
LogColor bool
|
||||
// LogWriter receives the colored per-request lines; nil means stderr.
|
||||
LogWriter io.Writer
|
||||
|
||||
LLMWaitTimeout time.Duration
|
||||
UnloadTimeout time.Duration
|
||||
@@ -77,6 +79,7 @@ type Config struct {
|
||||
type Server struct {
|
||||
cfg Config
|
||||
log *slog.Logger
|
||||
logWriter io.Writer
|
||||
freeTimeout time.Duration
|
||||
warmTimeout time.Duration
|
||||
captureLimit int64
|
||||
@@ -136,6 +139,7 @@ func New(cfg Config) (*Server, error) {
|
||||
return &Server{
|
||||
cfg: cfg,
|
||||
log: log,
|
||||
logWriter: cfg.LogWriter,
|
||||
freeTimeout: freeTimeout,
|
||||
warmTimeout: warmTimeout,
|
||||
captureLimit: captureLimit,
|
||||
@@ -338,7 +342,11 @@ func (s *Server) reqLine(log *slog.Logger, code, line string, attrs ...any) {
|
||||
fmt.Fprintf(&sb, " %v=%v", attrs[i], attrs[i+1])
|
||||
}
|
||||
sb.WriteByte('\n')
|
||||
os.Stderr.WriteString(sb.String())
|
||||
w := s.logWriter
|
||||
if w == nil {
|
||||
w = os.Stderr
|
||||
}
|
||||
io.WriteString(w, sb.String())
|
||||
}
|
||||
|
||||
// logRequests logs one line per incoming request and one per completed
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
//go:build !windows
|
||||
|
||||
// Package service provides the non-Windows stubs for the Windows service
|
||||
// integration. Run falls back to plain signal handling; install/remove
|
||||
// are unsupported.
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
// Name matches the Windows service name.
|
||||
const Name = "gpu-turnstile"
|
||||
|
||||
var errUnsupported = errors.New("service management is only supported on Windows")
|
||||
|
||||
// IsService is always false on non-Windows platforms.
|
||||
func IsService() bool { return false }
|
||||
|
||||
// Run executes run with SIGINT/SIGTERM cancellation, mirroring the
|
||||
// interactive behavior.
|
||||
func Run(run func(ctx context.Context) error) error {
|
||||
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
|
||||
defer stop()
|
||||
return run(ctx)
|
||||
}
|
||||
|
||||
// Install is unsupported on non-Windows platforms.
|
||||
func Install(string) error { return errUnsupported }
|
||||
|
||||
// Remove is unsupported on non-Windows platforms.
|
||||
func Remove() error { return errUnsupported }
|
||||
@@ -0,0 +1,120 @@
|
||||
//go:build windows
|
||||
|
||||
// Package service integrates gpu-turnstile with the Windows Service
|
||||
// Control Manager: running as a service with graceful stop, plus
|
||||
// install/remove helpers.
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
"golang.org/x/sys/windows/svc"
|
||||
"golang.org/x/sys/windows/svc/mgr"
|
||||
)
|
||||
|
||||
// Name is the Windows service name.
|
||||
const Name = "gpu-turnstile"
|
||||
|
||||
// IsService reports whether the process is running as a Windows service.
|
||||
func IsService() bool {
|
||||
isSvc, err := svc.IsWindowsService()
|
||||
return err == nil && isSvc
|
||||
}
|
||||
|
||||
// Run executes run as a Windows service. SCM Stop and Shutdown cancel the
|
||||
// context passed to run, triggering the same graceful shutdown as SIGTERM
|
||||
// in interactive mode.
|
||||
func Run(run func(ctx context.Context) error) error {
|
||||
return svc.Run(Name, &handler{run: run})
|
||||
}
|
||||
|
||||
type handler struct {
|
||||
run func(ctx context.Context) error
|
||||
}
|
||||
|
||||
func (h *handler) Execute(_ []string, requests <-chan svc.ChangeRequest, status chan<- svc.Status) (bool, uint32) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
status <- svc.Status{State: svc.StartPending}
|
||||
errCh := make(chan error, 1)
|
||||
go func() { errCh <- h.run(ctx) }()
|
||||
status <- svc.Status{State: svc.Running, Accepts: svc.AcceptStop | svc.AcceptShutdown}
|
||||
|
||||
for {
|
||||
select {
|
||||
case err := <-errCh:
|
||||
status <- svc.Status{State: svc.Stopped}
|
||||
if err != nil {
|
||||
return true, 1
|
||||
}
|
||||
return false, 0
|
||||
case c := <-requests:
|
||||
switch c.Cmd {
|
||||
case svc.Interrogate:
|
||||
status <- c.CurrentStatus
|
||||
case svc.Stop, svc.Shutdown:
|
||||
status <- svc.Status{State: svc.StopPending}
|
||||
cancel()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Install registers gpu-turnstile as an auto-start Windows service whose
|
||||
// binPath loads the given config file. Recovery actions restart the
|
||||
// service after 5s on failure — this is also what brings up a staged
|
||||
// update after the updater exits with a non-zero code.
|
||||
func Install(configPath string) error {
|
||||
exe, err := os.Executable()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
m, err := mgr.Connect()
|
||||
if err != nil {
|
||||
return fmt.Errorf("connect to service manager (run as administrator): %w", err)
|
||||
}
|
||||
defer m.Disconnect()
|
||||
|
||||
binPath := fmt.Sprintf(`"%s" -config "%s"`, exe, configPath)
|
||||
s, err := m.CreateService(Name, binPath, mgr.Config{
|
||||
StartType: mgr.StartAutomatic,
|
||||
DisplayName: "gpu-turnstile",
|
||||
Description: "GPU arbitration proxy for Ollama and ComfyUI",
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("create service: %w", err)
|
||||
}
|
||||
defer s.Close()
|
||||
|
||||
restart := mgr.RecoveryAction{Type: mgr.ServiceRestart, Delay: 5 * time.Second}
|
||||
if err := s.SetRecoveryActions([]mgr.RecoveryAction{restart, restart, restart}, 24*60*60); err != nil {
|
||||
return fmt.Errorf("set recovery actions: %w", err)
|
||||
}
|
||||
if err := s.SetRecoveryActionsOnNonCrashFailures(true); err != nil {
|
||||
return fmt.Errorf("set failure actions flag: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Remove stops (if running) and unregisters the service.
|
||||
func Remove() error {
|
||||
m, err := mgr.Connect()
|
||||
if err != nil {
|
||||
return fmt.Errorf("connect to service manager (run as administrator): %w", err)
|
||||
}
|
||||
defer m.Disconnect()
|
||||
s, err := m.OpenService(Name)
|
||||
if err != nil {
|
||||
return fmt.Errorf("open service: %w", err)
|
||||
}
|
||||
defer s.Close()
|
||||
s.Control(svc.Stop) // ignore error: may already be stopped
|
||||
if err := s.Delete(); err != nil {
|
||||
return fmt.Errorf("delete service: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
package update
|
||||
|
||||
// publicKeyPEM is the PEM-encoded Ed25519 public key that matches the
|
||||
// SIGNING_KEY secret used by CI to sign release binaries. Generate a
|
||||
// keypair once with:
|
||||
//
|
||||
// openssl genpkey -algorithm ed25519 -out private.pem
|
||||
// openssl pkey -in private.pem -pubout -out public.pem
|
||||
//
|
||||
// Paste the contents of public.pem here and commit; store private.pem as
|
||||
// the SIGNING_KEY repository secret. When empty, the updater refuses to
|
||||
// update (e.g. development builds).
|
||||
var publicKeyPEM = ""
|
||||
@@ -0,0 +1,254 @@
|
||||
// Package update implements gpu-turnstile's self-updater: it polls the
|
||||
// Gitea releases API, downloads the Windows binary of newer releases, and
|
||||
// verifies its Ed25519 signature (produced by CI with OpenSSL) before
|
||||
// swapping it in next to the running executable.
|
||||
package update
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/ed25519"
|
||||
"crypto/sha256"
|
||||
"crypto/x509"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// maxAssetSize bounds release asset downloads.
|
||||
const maxAssetSize = 512 << 20
|
||||
|
||||
// Updater checks one Gitea repository for newer releases.
|
||||
type Updater struct {
|
||||
Repo string // e.g. https://git.rambossek.at/PUBLIC/gpu-turnstile
|
||||
Asset string // e.g. gpu-turnstile.exe
|
||||
Version string // current version, e.g. v0.1.2 ("dev" disables updates)
|
||||
Log *slog.Logger
|
||||
Client *http.Client
|
||||
}
|
||||
|
||||
type release struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Assets []struct {
|
||||
Name string `json:"name"`
|
||||
BrowserDownloadURL string `json:"browser_download_url"`
|
||||
} `json:"assets"`
|
||||
}
|
||||
|
||||
func (u *Updater) logger() *slog.Logger {
|
||||
if u.Log != nil {
|
||||
return u.Log
|
||||
}
|
||||
return slog.Default()
|
||||
}
|
||||
|
||||
func (u *Updater) httpClient() *http.Client {
|
||||
if u.Client != nil {
|
||||
return u.Client
|
||||
}
|
||||
return &http.Client{Timeout: 5 * time.Minute}
|
||||
}
|
||||
|
||||
// apiURL derives <scheme>://<host>/api/v1/repos/<owner>/<name> from Repo.
|
||||
func (u *Updater) apiURL() (string, error) {
|
||||
repoURL, err := url.Parse(u.Repo)
|
||||
if err != nil || repoURL.Scheme == "" || repoURL.Host == "" {
|
||||
return "", fmt.Errorf("invalid UPDATE_REPO %q", u.Repo)
|
||||
}
|
||||
ownerName := strings.Trim(repoURL.Path, "/")
|
||||
if len(strings.Split(ownerName, "/")) != 2 {
|
||||
return "", fmt.Errorf("UPDATE_REPO %q: expected path /<owner>/<name>", u.Repo)
|
||||
}
|
||||
return fmt.Sprintf("%s://%s/api/v1/repos/%s", repoURL.Scheme, repoURL.Host, ownerName), nil
|
||||
}
|
||||
|
||||
// newerVersion reports whether latest is a higher vX.Y.Z version than
|
||||
// current. Both may carry a leading "v".
|
||||
func newerVersion(current, latest string) (bool, error) {
|
||||
parse := func(s string) ([3]int, error) {
|
||||
var v [3]int
|
||||
parts := strings.Split(strings.TrimPrefix(s, "v"), ".")
|
||||
if len(parts) != 3 {
|
||||
return v, fmt.Errorf("not a vX.Y.Z version: %q", s)
|
||||
}
|
||||
for i, p := range parts {
|
||||
n, err := strconv.Atoi(p)
|
||||
if err != nil {
|
||||
return v, fmt.Errorf("not a vX.Y.Z version: %q", s)
|
||||
}
|
||||
v[i] = n
|
||||
}
|
||||
return v, nil
|
||||
}
|
||||
cur, err := parse(current)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
lat, err := parse(latest)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
for i := 0; i < 3; i++ {
|
||||
if lat[i] != cur[i] {
|
||||
return lat[i] > cur[i], nil
|
||||
}
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func (u *Updater) get(ctx context.Context, url string) ([]byte, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := u.httpClient().Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
io.Copy(io.Discard, resp.Body)
|
||||
return nil, fmt.Errorf("GET %s: %s", url, resp.Status)
|
||||
}
|
||||
return io.ReadAll(io.LimitReader(resp.Body, maxAssetSize))
|
||||
}
|
||||
|
||||
func verifySignature(pubKeyPEM string, data, sig []byte) error {
|
||||
block, _ := pem.Decode([]byte(pubKeyPEM))
|
||||
if block == nil {
|
||||
return fmt.Errorf("invalid embedded public key PEM")
|
||||
}
|
||||
key, err := x509.ParsePKIXPublicKey(block.Bytes)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse public key: %w", err)
|
||||
}
|
||||
pub, ok := key.(ed25519.PublicKey)
|
||||
if !ok {
|
||||
return fmt.Errorf("public key is not Ed25519")
|
||||
}
|
||||
if !ed25519.Verify(pub, data, sig) {
|
||||
return fmt.Errorf("signature verification failed")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// stage swaps data into place at exePath: the running executable is
|
||||
// renamed aside (allowed on Windows) and the new file takes its name.
|
||||
func stage(exePath string, data []byte) error {
|
||||
newPath := exePath + ".new"
|
||||
oldPath := exePath + ".old"
|
||||
os.Remove(oldPath) // leftover from a previous update
|
||||
if err := os.WriteFile(newPath, data, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(exePath, oldPath); err != nil {
|
||||
os.Remove(newPath)
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(newPath, exePath); err != nil {
|
||||
os.Rename(oldPath, exePath) // roll back
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// CleanupOld removes the .old binary left behind by a staged update.
|
||||
// Call once at startup.
|
||||
func CleanupOld(exePath string) {
|
||||
os.Remove(exePath + ".old")
|
||||
os.Remove(exePath + ".new")
|
||||
}
|
||||
|
||||
// Check performs a single update check. staged is true when a newer,
|
||||
// signature-verified binary has been swapped into place at exePath; the
|
||||
// caller should then restart the process. A nil error with staged=false
|
||||
// means "no action" (up to date, disabled, or dev build); a non-nil error
|
||||
// means the check failed and the running binary is untouched.
|
||||
func (u *Updater) Check(ctx context.Context, exePath string) (staged bool, err error) {
|
||||
log := u.logger()
|
||||
if u.Version == "" || u.Version == "dev" {
|
||||
log.Debug("auto-update: dev build, skipping")
|
||||
return false, nil
|
||||
}
|
||||
if publicKeyPEM == "" {
|
||||
log.Debug("auto-update: no public key embedded, skipping")
|
||||
return false, nil
|
||||
}
|
||||
api, err := u.apiURL()
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
body, err := u.get(ctx, api+"/releases/latest")
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("fetch latest release: %w", err)
|
||||
}
|
||||
var rel release
|
||||
if err := json.Unmarshal(body, &rel); err != nil {
|
||||
return false, fmt.Errorf("parse release: %w", err)
|
||||
}
|
||||
newer, err := newerVersion(u.Version, rel.TagName)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if !newer {
|
||||
log.Debug("auto-update: up to date", "version", u.Version, "latest", rel.TagName)
|
||||
return false, nil
|
||||
}
|
||||
|
||||
urls := make(map[string]string, len(rel.Assets))
|
||||
for _, a := range rel.Assets {
|
||||
urls[a.Name] = a.BrowserDownloadURL
|
||||
}
|
||||
assetURL, ok := urls[u.Asset]
|
||||
if !ok {
|
||||
return false, fmt.Errorf("release %s has no asset %q", rel.TagName, u.Asset)
|
||||
}
|
||||
sigURL, ok := urls[u.Asset+".sig"]
|
||||
if !ok {
|
||||
return false, fmt.Errorf("release %s has no signature asset %q", rel.TagName, u.Asset+".sig")
|
||||
}
|
||||
|
||||
data, err := u.get(ctx, assetURL)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("download %s: %w", u.Asset, err)
|
||||
}
|
||||
sig, err := u.get(ctx, sigURL)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("download signature: %w", err)
|
||||
}
|
||||
|
||||
if sumURL, ok := urls[u.Asset+".sha256"]; ok {
|
||||
sumText, err := u.get(ctx, sumURL)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("download checksum: %w", err)
|
||||
}
|
||||
want := strings.Fields(string(sumText))[0]
|
||||
got := hex.EncodeToString(sha256Bytes(data))
|
||||
if !strings.EqualFold(want, got) {
|
||||
return false, fmt.Errorf("sha256 mismatch: got %s, want %s", got, want)
|
||||
}
|
||||
}
|
||||
if err := verifySignature(publicKeyPEM, data, sig); err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
if err := stage(exePath, data); err != nil {
|
||||
return false, fmt.Errorf("stage update: %w", err)
|
||||
}
|
||||
log.Info("auto-update: new version staged", "from", u.Version, "to", rel.TagName)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func sha256Bytes(data []byte) []byte {
|
||||
sum := sha256.Sum256(data)
|
||||
return sum[:]
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
package update
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/ed25519"
|
||||
"crypto/rand"
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// fakeGitea serves a Gitea-flavored releases API with one release.
|
||||
type fakeGitea struct {
|
||||
srv *httptest.Server
|
||||
pubPEM string
|
||||
asset []byte
|
||||
tag string
|
||||
tamper bool
|
||||
noSig bool
|
||||
}
|
||||
|
||||
func newFakeGitea(t *testing.T, tag string, assetContent []byte) *fakeGitea {
|
||||
t.Helper()
|
||||
pub, priv, err := ed25519.GenerateKey(rand.Reader)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
der, err := x509.MarshalPKIXPublicKey(pub)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
f := &fakeGitea{
|
||||
pubPEM: string(pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: der})),
|
||||
asset: assetContent,
|
||||
tag: tag,
|
||||
}
|
||||
sign := func() []byte { return ed25519.Sign(priv, f.asset) }
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/api/v1/repos/o/r/releases/latest", func(w http.ResponseWriter, r *http.Request) {
|
||||
assets := []map[string]string{
|
||||
{"name": "gpu-turnstile.exe", "browser_download_url": f.srv.URL + "/dl/exe"},
|
||||
{"name": "gpu-turnstile.exe.sha256", "browser_download_url": f.srv.URL + "/dl/sha"},
|
||||
}
|
||||
if !f.noSig {
|
||||
assets = append(assets, map[string]string{"name": "gpu-turnstile.exe.sig", "browser_download_url": f.srv.URL + "/dl/sig"})
|
||||
}
|
||||
json.NewEncoder(w).Encode(map[string]any{"tag_name": f.tag, "assets": assets})
|
||||
})
|
||||
mux.HandleFunc("/dl/exe", func(w http.ResponseWriter, r *http.Request) { w.Write(f.asset) })
|
||||
mux.HandleFunc("/dl/sig", func(w http.ResponseWriter, r *http.Request) {
|
||||
sig := sign()
|
||||
if f.tamper {
|
||||
sig[0] ^= 0xff
|
||||
}
|
||||
w.Write(sig)
|
||||
})
|
||||
mux.HandleFunc("/dl/sha", func(w http.ResponseWriter, r *http.Request) {
|
||||
fmt.Fprintf(w, "%x gpu-turnstile.exe\n", sha256Bytes(f.asset))
|
||||
})
|
||||
f.srv = httptest.NewServer(mux)
|
||||
t.Cleanup(f.srv.Close)
|
||||
return f
|
||||
}
|
||||
|
||||
func (f *fakeGitea) updater(version string) *Updater {
|
||||
return &Updater{Repo: f.srv.URL + "/o/r", Asset: "gpu-turnstile.exe", Version: version}
|
||||
}
|
||||
|
||||
func fakeExe(t *testing.T) string {
|
||||
t.Helper()
|
||||
exe := filepath.Join(t.TempDir(), "gpu-turnstile.exe")
|
||||
if err := os.WriteFile(exe, []byte("old-binary"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return exe
|
||||
}
|
||||
|
||||
func withPublicKey(t *testing.T, pem string) {
|
||||
t.Helper()
|
||||
old := publicKeyPEM
|
||||
publicKeyPEM = pem
|
||||
t.Cleanup(func() { publicKeyPEM = old })
|
||||
}
|
||||
|
||||
func TestCheckStagesUpdate(t *testing.T) {
|
||||
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||
withPublicKey(t, f.pubPEM)
|
||||
exe := fakeExe(t)
|
||||
|
||||
staged, err := f.updater("v0.1.2").Check(context.Background(), exe)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !staged {
|
||||
t.Fatal("expected staged update")
|
||||
}
|
||||
content, _ := os.ReadFile(exe)
|
||||
if string(content) != "new-binary" {
|
||||
t.Fatalf("exe content = %q", content)
|
||||
}
|
||||
old, _ := os.ReadFile(exe + ".old")
|
||||
if string(old) != "old-binary" {
|
||||
t.Fatalf(".old content = %q", old)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckRejectsTamperedSignature(t *testing.T) {
|
||||
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||
f.tamper = true
|
||||
withPublicKey(t, f.pubPEM)
|
||||
exe := fakeExe(t)
|
||||
|
||||
staged, err := f.updater("v0.1.2").Check(context.Background(), exe)
|
||||
if err == nil {
|
||||
t.Fatal("expected signature error")
|
||||
}
|
||||
if staged {
|
||||
t.Fatal("must not stage on bad signature")
|
||||
}
|
||||
content, _ := os.ReadFile(exe)
|
||||
if string(content) != "old-binary" {
|
||||
t.Fatal("exe was modified despite bad signature")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckSkipsOlderOrEqual(t *testing.T) {
|
||||
for _, tag := range []string{"v0.1.2", "v0.1.1", "v0.0.9"} {
|
||||
f := newFakeGitea(t, tag, []byte("new-binary"))
|
||||
withPublicKey(t, f.pubPEM)
|
||||
staged, err := f.updater("v0.1.2").Check(context.Background(), fakeExe(t))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if staged {
|
||||
t.Fatalf("tag %s must not stage over v0.1.2", tag)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckSkipsWithoutPublicKey(t *testing.T) {
|
||||
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||
withPublicKey(t, "")
|
||||
staged, err := f.updater("v0.1.2").Check(context.Background(), fakeExe(t))
|
||||
if err != nil || staged {
|
||||
t.Fatalf("staged=%v err=%v, want no action without key", staged, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckSkipsDevBuild(t *testing.T) {
|
||||
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||
withPublicKey(t, f.pubPEM)
|
||||
staged, err := f.updater("dev").Check(context.Background(), fakeExe(t))
|
||||
if err != nil || staged {
|
||||
t.Fatalf("staged=%v err=%v, want no action for dev build", staged, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewerVersion(t *testing.T) {
|
||||
cases := []struct {
|
||||
cur, lat string
|
||||
want bool
|
||||
}{
|
||||
{"v0.1.2", "v0.1.3", true},
|
||||
{"0.1.2", "0.2.0", true},
|
||||
{"v0.1.2", "v1.0.0", true},
|
||||
{"v0.1.2", "v0.1.2", false},
|
||||
{"v1.2.3", "v1.2.10", true},
|
||||
{"v1.2.10", "v1.2.3", false},
|
||||
}
|
||||
for _, c := range cases {
|
||||
got, err := newerVersion(c.cur, c.lat)
|
||||
if err != nil || got != c.want {
|
||||
t.Errorf("newerVersion(%s, %s) = %v, %v; want %v", c.cur, c.lat, got, err, c.want)
|
||||
}
|
||||
}
|
||||
if _, err := newerVersion("v0.1", "v0.1.2"); err == nil {
|
||||
t.Error("expected error for malformed version")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user