From e363c9e4e47edd0e4461d507fda3878c15d77dd6 Mon Sep 17 00:00:00 2001 From: mram Date: Sun, 20 Sep 2026 21:33:29 +0200 Subject: [PATCH] Native Windows deployment: config file, service, signed auto-update MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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. --- .gitea/workflows/ci.yml | 47 +++ README.md | 53 +++- SPEC.md | 60 +++- cmd/gpu-turnstile/main.go | 441 ++++++++++++++++------------ go.mod | 2 + go.sum | 2 + internal/config/config.go | 192 ++++++++++++ internal/config/config_test.go | 97 ++++++ internal/proxy/proxy.go | 10 +- internal/service/service_other.go | 35 +++ internal/service/service_windows.go | 120 ++++++++ internal/update/pubkey.go | 13 + internal/update/update.go | 254 ++++++++++++++++ internal/update/update_test.go | 186 ++++++++++++ 14 files changed, 1318 insertions(+), 194 deletions(-) create mode 100644 go.sum create mode 100644 internal/config/config.go create mode 100644 internal/config/config_test.go create mode 100644 internal/service/service_other.go create mode 100644 internal/service/service_windows.go create mode 100644 internal/update/pubkey.go create mode 100644 internal/update/update.go create mode 100644 internal/update/update_test.go diff --git a/.gitea/workflows/ci.yml b/.gitea/workflows/ci.yml index 0f3e7d3..1a30987 100644 --- a/.gitea/workflows/ci.yml +++ b/.gitea/workflows/ci.yml @@ -60,3 +60,50 @@ jobs: tags: | ${{ steps.meta.outputs.image }}:${{ gitea.ref_name }} ${{ steps.meta.outputs.image }}:latest + + # On version tags: build the signed Windows binary and attach it (plus + # signature and checksum) to a Gitea release for the auto-updater. + release: + if: gitea.ref_type == 'tag' + needs: test + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version: "1.23" + + - name: Build Windows binary + run: | + GOOS=windows GOARCH=amd64 CGO_ENABLED=0 go build \ + -ldflags="-s -w -X main.version=${{ gitea.ref_name }}" \ + -o gpu-turnstile.exe ./cmd/gpu-turnstile + + - name: Sign and checksum + run: | + printf '%s\n' "${{ secrets.SIGNING_KEY }}" > key.pem + chmod 600 key.pem + openssl pkeyutl -sign -inkey key.pem -rawin \ + -in gpu-turnstile.exe -out gpu-turnstile.exe.sig + rm -f key.pem + sha256sum gpu-turnstile.exe > gpu-turnstile.exe.sha256 + + - name: Create release and upload assets + env: + TOKEN: ${{ secrets.GITEA_TOKEN }} + API: https://git.rambossek.at/api/v1/repos/${{ gitea.repository }} + TAG: ${{ gitea.ref_name }} + run: | + set -e + ID=$(curl -sf -H "Authorization: token $TOKEN" "$API/releases/tags/$TAG" | jq -r .id || true) + if [ -z "$ID" ] || [ "$ID" = "null" ]; then + ID=$(curl -sf -X POST -H "Authorization: token $TOKEN" \ + -H "Content-Type: application/json" \ + -d "{\"tag_name\":\"$TAG\",\"name\":\"$TAG\"}" \ + "$API/releases" | jq -r .id) + fi + for f in gpu-turnstile.exe gpu-turnstile.exe.sig gpu-turnstile.exe.sha256; do + curl -sf -X POST -H "Authorization: token $TOKEN" \ + -F "attachment=@$f" "$API/releases/$ID/assets?name=$f" > /dev/null + echo "uploaded $f" + done diff --git a/README.md b/README.md index d144a4f..f85d0cf 100644 --- a/README.md +++ b/README.md @@ -28,8 +28,11 @@ Open WebUI / n8n ────► :8188 ───┘ ## Configuration -All configuration is via environment variables; invalid values fail at -startup. +Configuration comes from environment variables and/or an `.env`-style +config file (`KEY=VALUE` lines, `#` comments). File lookup order: +`-config ` flag, then `GPU_TURNSTILE_CONFIG`, then +`gpu-turnstile.env` next to the executable. Process environment variables +override file values. Invalid values fail at startup. | Var | Default | Meaning | |---|---|---| @@ -43,6 +46,7 @@ startup. | `WARM_MODEL` | _(empty)_ | Model to reload after an image job (off by default) | | `LOGLEVEL` | `warn` | `info` logs every request (colored arrows in text mode), `debug` adds lock transitions. `LOG_LEVEL` works as an alias | | `LOG_FORMAT` | `text` | `json` for structured JSON logs | +| `LOG_FILE` | _(empty)_ | Append logs to this file instead of stderr | | `UNLOAD_POLL_INTERVAL` | `500ms` | `/api/ps` poll interval while unloading | | `HISTORY_POLL_INTERVAL` | `1s` | `/history/` poll interval while a job runs | | `PROBE_TIMEOUT` | `5s` | Startup probe of both upstreams | @@ -52,6 +56,10 @@ startup. | `BACKOFF_INITIAL` | `1s` | First retry wait when an upstream refuses a connection | | `BACKOFF_MAX` | `60s` | Cap for the exponential retry backoff | | `PROMPT_CAPTURE_LIMIT` | `65536` | Bytes of the `/prompt` response buffered to find `prompt_id` (pass-through is unaffected) | +| `AUTO_UPDATE` | `true` | Poll the Gitea releases API for signed updates | +| `UPDATE_INTERVAL` | `6h` | Auto-update check interval | +| `UPDATE_REPO` | `https://git.rambossek.at/PUBLIC/gpu-turnstile` | Repository checked for releases | +| `UPDATE_ASSET` | `gpu-turnstile.exe` | Release asset to download | ## Observability @@ -74,6 +82,32 @@ go build ./cmd/gpu-turnstile ./gpu-turnstile ``` +### Run natively on Windows (primary deployment) + +Download `gpu-turnstile.exe` from a release, put a `gpu-turnstile.env` +next to it, and run it — or install it as a Windows service from an +elevated shell: + +```sh +gpu-turnstile.exe service install # auto-start service, recovery = restart +gpu-turnstile.exe service remove +``` + +The service uses the config file (services have no convenient +environment); set `LOG_FILE` in it since there is no console. + +**Auto-update is on by default**: the binary checks the repo's latest +release on startup and every `UPDATE_INTERVAL`, verifies the Ed25519 +signature of the download against the public key embedded at build time, +and — once the GPU lock is idle — restarts the service onto the new +version. Disable with `AUTO_UPDATE=false`. Releases are signed by CI with +OpenSSL; the matching public key lives in `internal/update/pubkey.go` +(one-time setup: `openssl genpkey -algorithm ed25519 -out private.pem`, +`openssl pkey -in private.pem -pubout -out public.pem`; private key goes +to the `SIGNING_KEY` repo secret, public key is committed). + +### Docker + ```sh docker build -t gpu-turnstile . docker run --rm -p 11434:11434 -p 8188:8188 \ @@ -84,9 +118,10 @@ docker run --rm -p 11434:11434 -p 8188:8188 \ Releases are built by Gitea Actions (`.gitea/workflows/ci.yml`): every push runs `go vet` and `go test -race`, and pushing a semantic-version tag -`vX.Y.Z` builds and publishes -`git.rambossek.at//gpu-turnstile:vX.Y.Z` (and updates `:latest`). -No images are built from branches. +`vX.Y.Z` publishes the container image +(`git.rambossek.at//gpu-turnstile:vX.Y.Z` plus `:latest`) and a +signed Windows binary attached to a Gitea release. Nothing is built from +branches. The registry login needs one repository secret (Settings → Actions → Secrets): `REGISTRY_TOKEN` — an access token with `write:package` scope. @@ -99,13 +134,17 @@ go vet ./... go test -race ./... ``` -Stdlib only, Go 1.23+. Layout: +Go 1.23+; the only external dependency is `golang.org/x/sys` (Windows +service integration, unused in the Linux build). Layout: ``` -cmd/gpu-turnstile/main.go wiring, config, listeners +cmd/gpu-turnstile/main.go wiring, config, listeners, service + updater internal/lock/ two-mode lock (LLM readers / image writer, FIFO) internal/ollama/ ps / unload / warm client internal/comfy/ history poll / free client internal/proxy/ handlers for both listeners internal/metrics/ Prometheus exposition, no dependencies +internal/config/ env + .env file configuration +internal/update/ signed auto-updater +internal/service/ Windows service integration ``` diff --git a/SPEC.md b/SPEC.md index 5c6bedd..374c7c6 100644 --- a/SPEC.md +++ b/SPEC.md @@ -102,6 +102,12 @@ load time. Off by default. ## Configuration (env) +Configuration comes from environment variables and/or an `.env`-style +config file (`KEY=VALUE` lines, `#` comments). File lookup order: +`-config ` flag, then `GPU_TURNSTILE_CONFIG`, then +`gpu-turnstile.env` next to the executable. Process environment variables +override file values. A missing file is fine; a malformed one is fatal. + | Var | Default | Meaning | |---|---|---| | `LISTEN_OLLAMA` | `:11434` | Ollama-facing listener | @@ -113,6 +119,8 @@ load time. Off by default. | `LLM_WAIT_TIMEOUT` | `10m` | max time an LLM request waits for the lock before 503 | | `WARM_MODEL` | `` | optional model to reload after an image job | | `LOGLEVEL` | `warn` | `info` logs every request (colored arrows in text mode), `debug` adds lock transitions. `LOG_LEVEL` is accepted as an alias | +| `LOG_FORMAT` | `text` | `json` for structured JSON logs | +| `LOG_FILE` | `` | append logs to this file instead of stderr (useful as a service) | | `UNLOAD_POLL_INTERVAL` | `500ms` | `/api/ps` poll interval while unloading | | `HISTORY_POLL_INTERVAL` | `1s` | `/history/` poll interval while a job runs | | `PROBE_TIMEOUT` | `5s` | startup probe of both upstreams | @@ -122,10 +130,39 @@ load time. Off by default. | `BACKOFF_INITIAL` | `1s` | first retry wait when an upstream refuses a connection | | `BACKOFF_MAX` | `60s` | cap for the exponential retry backoff | | `PROMPT_CAPTURE_LIMIT` | `65536` | bytes of the `/prompt` response buffered to find `prompt_id` (pass-through is unaffected) | +| `AUTO_UPDATE` | `true` | poll the Gitea releases API for signed updates | +| `UPDATE_INTERVAL` | `6h` | auto-update check interval | +| `UPDATE_REPO` | `https://git.rambossek.at/PUBLIC/gpu-turnstile` | repository to check for releases | +| `UPDATE_ASSET` | `gpu-turnstile.exe` | release asset to download | Startup fails fast on unparsable values. Both upstreams are probed once at start (`/api/version`, `/system_stats`); failure is logged, not fatal. +## Native Windows deployment + +The binary runs natively on Windows (the primary deployment) as well as in +Docker. + +- `gpu-turnstile.exe service install [-config path]` registers an + auto-start Windows service (needs an elevated shell). Recovery actions + restart it 5 s after any failure. `service remove` uninstalls. +- Use a config file (above) for the service — Windows services have no + convenient environment. Logs go to `LOG_FILE` since there is no console. +- **Auto-update**: on startup and every `UPDATE_INTERVAL`, the binary + checks `UPDATE_REPO`'s latest release; if its tag is a newer `vX.Y.Z`, + it downloads `UPDATE_ASSET` plus its `.sig` (and `.sha256` when present) + and verifies an Ed25519 signature against the public key embedded in + `internal/update/pubkey.go`. A verified binary is swapped in next to the + running exe (rename-aside, allowed on Windows), and once the GPU lock is + idle the process exits with code 3 so the service recovery restarts it + on the new version. Interactive runs only log "restart to apply". + `dev` builds and builds without an embedded public key never update. +- **Signing setup (one time)**: `openssl genpkey -algorithm ed25519 -out + private.pem`; `openssl pkey -in private.pem -pubout -out public.pem`. + Private key → repo secret `SIGNING_KEY`; public key → committed into + `internal/update/pubkey.go`. CI signs release binaries with + `openssl pkeyutl -sign -rawin`. + ## Observability - `GET /healthz` on both listeners: 200 with JSON @@ -171,11 +208,15 @@ start (`/api/version`, `/system_stats`); failure is logged, not fatal. ``` gpu-turnstile/ - cmd/gpu-turnstile/main.go # wiring, config, listeners + cmd/gpu-turnstile/main.go # wiring, config, listeners, service + updater internal/lock/lock.go # two-mode lock + tests internal/ollama/client.go # ps / unload / warm internal/comfy/client.go # history poll / free internal/proxy/ # handlers for both listeners + internal/metrics/ # Prometheus exposition + internal/config/ # env + .env file configuration + internal/update/ # signed auto-updater (public key in pubkey.go) + internal/service/ # Windows service integration Dockerfile .gitea/workflows/ci.yml README.md @@ -197,11 +238,16 @@ are new. held until `/free` was called. - Streaming test: fake Ollama emits chunks with delays; assert the client receives the first chunk before the last is sent (no buffering). +- `internal/config`: env-file parsing, precedence, fail-fast values. +- `internal/update`: fake Gitea releases API; staged update happy path, + tampered signature rejected, older versions and dev builds skipped. ## Build and CI -- Go 1.23+, stdlib only. `CGO_ENABLED=0`, `-ldflags="-s -w"`, version from - `git describe` injected via `-X main.version=`. +- Go 1.23+, `golang.org/x/sys` is the only external dependency (Windows + service integration; not used in the Linux build). `CGO_ENABLED=0`, + `-ldflags="-s -w"`, version from `git describe` injected via + `-X main.version=`. - Dockerfile: multi-stage, final image `gcr.io/distroless/static` (or `scratch`), non-root user, `EXPOSE 8188 11434`, `ENTRYPOINT ["/gpu-turnstile"]`. @@ -216,8 +262,12 @@ are new. Login uses the repo secret `REGISTRY_TOKEN` (an access token with `write:package` scope) because the automatic `GITEA_TOKEN` cannot push packages; the username is just `gitea.actor`. -- Release: a git tag `vX.Y.Z` produces the versioned image; the Open WebUI - compose pins that tag. No images are built from branches. + 3. on a version tag: also build the Windows binary, sign it with OpenSSL + (`SIGNING_KEY` secret), and attach `gpu-turnstile.exe`, `.sig` and + `.sha256` to a Gitea release for the auto-updater. +- Release: a git tag `vX.Y.Z` produces the versioned image and the signed + Windows binary; the Open WebUI compose pins that tag. No images or + binaries are built from branches. ## Deployment (target) diff --git a/cmd/gpu-turnstile/main.go b/cmd/gpu-turnstile/main.go index 7a56ea2..6083496 100644 --- a/cmd/gpu-turnstile/main.go +++ b/cmd/gpu-turnstile/main.go @@ -6,256 +6,335 @@ import ( "context" "errors" "fmt" + "io" "log/slog" "net/http" "os" "os/signal" - "strconv" + "path/filepath" "strings" "syscall" "time" "gpu-turnstile/internal/comfy" + "gpu-turnstile/internal/config" "gpu-turnstile/internal/lock" "gpu-turnstile/internal/metrics" "gpu-turnstile/internal/ollama" "gpu-turnstile/internal/proxy" + "gpu-turnstile/internal/service" + "gpu-turnstile/internal/update" ) // version is injected at build time via -ldflags "-X main.version=...". var version = "dev" -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 - - warmModel string - logLevel slog.Level - logJSON bool -} - -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 -} - -func loadConfig(getenv func(string) string) (config, error) { - cfg := 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, - - logLevel: slog.LevelWarn, - } - 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}, - } { - 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}, - } { - 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 - } - // 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 -} +// exitCodeUpdate tells the service recovery configuration to restart the +// process: a signed update has been staged and the GPU lock is idle. +const exitCodeUpdate = 3 func main() { - cfg, err := loadConfig(os.Getenv) + configPath, args := splitConfigFlag(os.Args[1:]) + if len(args) > 0 && args[0] == "service" { + os.Exit(serviceCommand(configPath, args[1:])) + } + if len(args) > 0 { + fmt.Fprintf(os.Stderr, "usage: gpu-turnstile [-config path] | gpu-turnstile service install|remove [-config path]\n") + os.Exit(2) + } + + if exePath, err := os.Executable(); err == nil { + update.CleanupOld(exePath) + } + + cfg, err := loadMergedConfig(configPath) if err != nil { fmt.Fprintf(os.Stderr, "gpu-turnstile: %v\n", err) os.Exit(1) } + log, logOut, logCloser := newLogger(cfg) + defer logCloser.Close() - opts := &slog.HandlerOptions{Level: cfg.logLevel} - var handler slog.Handler = slog.NewTextHandler(os.Stderr, opts) - if cfg.logJSON { - handler = slog.NewJSONHandler(os.Stderr, opts) + if service.IsService() { + if err := service.Run(func(ctx context.Context) error { return run(ctx, cfg, log, logOut, true) }); err != nil { + log.Error("service failed", "err", err) + os.Exit(1) + } + return + } + ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) + defer stop() + if err := run(ctx, cfg, log, logOut, false); err != nil { + log.Error("listener failed", "err", err) + os.Exit(1) + } +} + +// splitConfigFlag extracts -config (or -config=) from args. +func splitConfigFlag(args []string) (string, []string) { + var configPath string + rest := args[:0] + for i := 0; i < len(args); i++ { + switch { + case args[i] == "-config" && i+1 < len(args): + configPath = args[i+1] + i++ + case strings.HasPrefix(args[i], "-config="): + configPath = strings.TrimPrefix(args[i], "-config=") + default: + rest = append(rest, args[i]) + } + } + return configPath, rest +} + +// defaultConfigPath returns gpu-turnstile.env next to the executable. +func defaultConfigPath() string { + exe, err := os.Executable() + if err != nil { + return "gpu-turnstile.env" + } + return filepath.Join(filepath.Dir(exe), "gpu-turnstile.env") +} + +// resolveConfigPath applies the precedence: -config flag, then +// GPU_TURNSTILE_CONFIG, then the default next to the executable. +func resolveConfigPath(flagValue string) string { + if flagValue != "" { + return flagValue + } + if v := os.Getenv("GPU_TURNSTILE_CONFIG"); v != "" { + return v + } + return defaultConfigPath() +} + +// loadMergedConfig reads the config file (if present) and overlays process +// environment variables on top. A missing file is fine; an unreadable or +// malformed file is fatal. +func loadMergedConfig(flagValue string) (config.Config, error) { + path := resolveConfigPath(flagValue) + values := map[string]string{} + if f, err := os.Open(path); err == nil { + defer f.Close() + parsed, err := config.ParseEnvFile(f) + if err != nil { + return config.Config{}, fmt.Errorf("%s: %w", path, err) + } + values = parsed + } else if !errors.Is(err, os.ErrNotExist) { + return config.Config{}, fmt.Errorf("read config file: %w", err) + } + getenv := func(key string) string { + if v := os.Getenv(key); v != "" { + return v + } + return values[key] + } + return config.Load(getenv) +} + +// newLogger builds the slog logger and returns the output writer (stderr or +// the opened LOG_FILE) plus a closer for it. +func newLogger(cfg config.Config) (*slog.Logger, io.Writer, io.Closer) { + out := io.Writer(os.Stderr) + closer := io.NopCloser(nil) + if cfg.LogFile != "" { + if f, err := os.OpenFile(cfg.LogFile, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o644); err == nil { + out = f + closer = f + } else { + fmt.Fprintf(os.Stderr, "gpu-turnstile: cannot open LOG_FILE %s: %v (logging to stderr)\n", cfg.LogFile, err) + } + } + opts := &slog.HandlerOptions{Level: cfg.LogLevel} + var handler slog.Handler = slog.NewTextHandler(out, opts) + if cfg.LogJSON { + handler = slog.NewJSONHandler(out, opts) } log := slog.New(handler) slog.SetDefault(log) + return log, out, closer +} +func serviceCommand(configPath string, args []string) int { + if len(args) != 1 || (args[0] != "install" && args[0] != "remove") { + fmt.Fprintf(os.Stderr, "usage: gpu-turnstile service install|remove [-config path]\n") + return 2 + } + var err error + if args[0] == "install" { + path := resolveConfigPath(configPath) + if abs, absErr := filepath.Abs(path); absErr == nil { + path = abs + } + err = service.Install(path) + } else { + err = service.Remove() + } + if err != nil { + fmt.Fprintf(os.Stderr, "gpu-turnstile service %s: %v\n", args[0], err) + return 1 + } + fmt.Printf("service %s: %sd\n", service.Name, args[0]) + return 0 +} + +func run(ctx context.Context, cfg config.Config, log *slog.Logger, logOut io.Writer, isService bool) error { // The startup line carries the version and every setting and is emitted // at WARN so it is visible even with the default (quiet) log level. - log.Log(context.Background(), slog.LevelWarn, "starting gpu-turnstile", + log.Log(ctx, slog.LevelWarn, "starting gpu-turnstile", "version", version, - "listen_ollama", cfg.listenOllama, - "listen_comfy", cfg.listenComfy, - "ollama_url", cfg.ollamaURL, - "comfy_url", cfg.comfyURL, - "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, - "prompt_capture_limit", cfg.promptCaptureLimit, - "warm_model", cfg.warmModel, - "log_level", cfg.logLevel, - "log_format", map[bool]string{true: "json", false: "text"}[cfg.logJSON], + "listen_ollama", cfg.ListenOllama, + "listen_comfy", cfg.ListenComfy, + "ollama_url", cfg.OllamaURL, + "comfy_url", cfg.ComfyURL, + "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, + "prompt_capture_limit", cfg.PromptCaptureLimit, + "warm_model", cfg.WarmModel, + "auto_update", cfg.AutoUpdate, + "update_interval", cfg.UpdateInterval, + "update_repo", cfg.UpdateRepo, + "update_asset", cfg.UpdateAsset, + "log_level", cfg.LogLevel, + "log_format", map[bool]string{true: "json", false: "text"}[cfg.LogJSON], + "log_file", cfg.LogFile, ) - ollamaClient, err := ollama.New(cfg.ollamaURL, log) + lk := lock.New(log) + ollamaClient, err := ollama.New(cfg.OllamaURL, log) if err != nil { - log.Error("invalid configuration", "err", err) - os.Exit(1) + return err } - comfyClient, err := comfy.New(cfg.comfyURL, log) + comfyClient, err := comfy.New(cfg.ComfyURL, log) if err != nil { - log.Error("invalid configuration", "err", err) - os.Exit(1) + return err } srv, err := proxy.New(proxy.Config{ - OllamaURL: cfg.ollamaURL, - ComfyURL: cfg.comfyURL, - Lock: lock.New(log), + OllamaURL: cfg.OllamaURL, + ComfyURL: cfg.ComfyURL, + Lock: lk, Ollama: ollamaClient, Comfy: comfyClient, Metrics: metrics.New(), Log: log, - LogColor: !cfg.logJSON && os.Getenv("NO_COLOR") == "", - LLMWaitTimeout: cfg.llmWaitTimeout, - UnloadTimeout: cfg.unloadTimeout, - JobTimeout: cfg.jobTimeout, - UnloadPollInterval: cfg.unloadPollInterval, - HistoryPollInterval: cfg.historyPollInterval, - FreeTimeout: cfg.freeTimeout, - WarmTimeout: cfg.warmTimeout, - BackoffInitial: cfg.backoffInitial, - BackoffMax: cfg.backoffMax, - PromptCaptureLimit: cfg.promptCaptureLimit, - WarmModel: cfg.warmModel, + LogColor: !cfg.LogJSON && cfg.LogFile == "" && os.Getenv("NO_COLOR") == "", + LogWriter: logOut, + LLMWaitTimeout: cfg.LLMWaitTimeout, + UnloadTimeout: cfg.UnloadTimeout, + JobTimeout: cfg.JobTimeout, + UnloadPollInterval: cfg.UnloadPollInterval, + HistoryPollInterval: cfg.HistoryPollInterval, + FreeTimeout: cfg.FreeTimeout, + WarmTimeout: cfg.WarmTimeout, + BackoffInitial: cfg.BackoffInitial, + BackoffMax: cfg.BackoffMax, + PromptCaptureLimit: cfg.PromptCaptureLimit, + WarmModel: cfg.WarmModel, }) if err != nil { - log.Error("invalid configuration", "err", err) - os.Exit(1) + return err } // Probe both upstreams once; failure is logged, not fatal. - probeCtx, probeCancel := context.WithTimeout(context.Background(), cfg.probeTimeout) + probeCtx, probeCancel := context.WithTimeout(ctx, cfg.ProbeTimeout) if err := ollamaClient.Probe(probeCtx); err != nil { - log.Warn("ollama probe failed", "url", cfg.ollamaURL, "err", err) + log.Warn("ollama probe failed", "url", cfg.OllamaURL, "err", err) } if err := comfyClient.Probe(probeCtx); err != nil { - log.Warn("comfy probe failed", "url", cfg.comfyURL, "err", err) + log.Warn("comfy probe failed", "url", cfg.ComfyURL, "err", err) } probeCancel() - ollamaSrv := &http.Server{Addr: cfg.listenOllama, Handler: srv.OllamaHandler()} - comfySrv := &http.Server{Addr: cfg.listenComfy, Handler: srv.ComfyHandler()} + ollamaSrv := &http.Server{Addr: cfg.ListenOllama, Handler: srv.OllamaHandler()} + comfySrv := &http.Server{Addr: cfg.ListenComfy, Handler: srv.ComfyHandler()} errCh := make(chan error, 2) go func() { errCh <- ollamaSrv.ListenAndServe() }() go func() { errCh <- comfySrv.ListenAndServe() }() - sigCtx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) - defer stop() + if cfg.AutoUpdate { + go updateLoop(ctx, cfg, log, lk, isService) + } select { case err := <-errCh: if err != nil && !errors.Is(err, http.ErrServerClosed) { - log.Error("listener failed", "err", err) - os.Exit(1) + return err } - case <-sigCtx.Done(): + case <-ctx.Done(): log.Info("shutting down") } - shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), cfg.shutdownTimeout) + shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), cfg.ShutdownTimeout) defer shutdownCancel() ollamaSrv.Shutdown(shutdownCtx) comfySrv.Shutdown(shutdownCtx) + return nil +} + +// updateLoop checks for signed updates on startup and every UPDATE_INTERVAL. +// In service mode a staged update is applied by exiting with exitCodeUpdate +// once the GPU lock is idle; the service recovery configuration restarts the +// process with the new binary. Interactively it only logs. +func updateLoop(ctx context.Context, cfg config.Config, log *slog.Logger, lk *lock.Lock, isService bool) { + exePath, err := os.Executable() + if err != nil { + log.Warn("auto-update disabled: cannot locate executable", "err", err) + return + } + u := &update.Updater{Repo: cfg.UpdateRepo, Asset: cfg.UpdateAsset, Version: version, Log: log} + for { + staged, err := u.Check(ctx, exePath) + if err != nil && ctx.Err() == nil { + log.Warn("auto-update check failed", "err", err) + } + if staged { + if !isService { + log.Warn("auto-update: new binary staged; restart gpu-turnstile to apply") + return + } + log.Warn("auto-update: staged; restarting once the GPU is idle") + if waitForIdle(ctx, lk, 24*time.Hour) { + log.Warn("auto-update: restarting to apply update") + os.Exit(exitCodeUpdate) + } + return + } + select { + case <-ctx.Done(): + return + case <-time.After(cfg.UpdateInterval): + } + } +} + +// waitForIdle polls the lock until no LLM or image work is active or +// pending, max at most. Returns false on timeout or cancellation. +func waitForIdle(ctx context.Context, lk *lock.Lock, max time.Duration) bool { + deadline := time.Now().Add(max) + for { + if state, _, pending := lk.Snapshot(); state == lock.StateIdle && !pending { + return true + } + if time.Now().After(deadline) { + return false + } + select { + case <-ctx.Done(): + return false + case <-time.After(5 * time.Second): + } + } } diff --git a/go.mod b/go.mod index 9a475e8..05d9acb 100644 --- a/go.mod +++ b/go.mod @@ -1,3 +1,5 @@ module gpu-turnstile go 1.23 + +require golang.org/x/sys v0.29.0 diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..0664caa --- /dev/null +++ b/go.sum @@ -0,0 +1,2 @@ +golang.org/x/sys v0.29.0 h1:TPYlXGxvx1MGTn2GiZDhnjPA9wZzZeGKHHmKhHYvgaU= +golang.org/x/sys v0.29.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= diff --git a/internal/config/config.go b/internal/config/config.go new file mode 100644 index 0000000..add9fdd --- /dev/null +++ b/internal/config/config.go @@ -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 +} diff --git a/internal/config/config_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..2e081ba --- /dev/null +++ b/internal/config/config_test.go @@ -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) + } + } +} diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index 212ec54..f6bccd3 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -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 diff --git a/internal/service/service_other.go b/internal/service/service_other.go new file mode 100644 index 0000000..5821a29 --- /dev/null +++ b/internal/service/service_other.go @@ -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 } diff --git a/internal/service/service_windows.go b/internal/service/service_windows.go new file mode 100644 index 0000000..6760762 --- /dev/null +++ b/internal/service/service_windows.go @@ -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 +} diff --git a/internal/update/pubkey.go b/internal/update/pubkey.go new file mode 100644 index 0000000..9e9597b --- /dev/null +++ b/internal/update/pubkey.go @@ -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 = "" diff --git a/internal/update/update.go b/internal/update/update.go new file mode 100644 index 0000000..b82603e --- /dev/null +++ b/internal/update/update.go @@ -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 :///api/v1/repos// 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 //", 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[:] +} diff --git a/internal/update/update_test.go b/internal/update/update_test.go new file mode 100644 index 0000000..b09e9d8 --- /dev/null +++ b/internal/update/update_test.go @@ -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") + } +}