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:
mram
2026-09-20 21:33:29 +02:00
parent 5f6a22c2cf
commit e363c9e4e4
14 changed files with 1318 additions and 194 deletions
+192
View File
@@ -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
}
+97
View File
@@ -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)
}
}
}
+9 -1
View File
@@ -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
+35
View File
@@ -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 }
+120
View File
@@ -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
}
+13
View File
@@ -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 = ""
+254
View File
@@ -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[:]
}
+186
View File
@@ -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")
}
}