31 Commits
Author SHA1 Message Date
mram dcdb2cc72e Pin compose example to v0.1.7
ci / test (push) Successful in 13s
ci / docker (push) Successful in 1m7s
ci / release (push) Successful in 16s
2026-09-21 10:04:29 +02:00
mram 44a8e9fbdc CFG_VER/APP_VER in the env file: invalid configs replaced, updates follow APP_VER (dev/stable/pin) 2026-09-21 09:28:29 +02:00
mram 272a3a467d force-update: single-shot with a 30s fail-fast timeout 2026-09-21 09:09:27 +02:00
mram cedf28a809 force-update does not require consumer URLs (ErrNoConsumer sentinel) 2026-09-21 09:05:46 +02:00
mram e46e92bee8 Sample env carries a version marker; install appends settings missing since the writing version 2026-09-21 08:54:47 +02:00
mram cf48796183 Install writes a fully commented sample env file when no config exists 2026-09-21 08:48:51 +02:00
mram dd3187f951 Install sets a default LOG_FILE in the env file on Windows (no console as a service) 2026-09-21 08:37:25 +02:00
mram 1394a76eea Make service install converge: stop running service, refresh binary/config only on change, restart only if it was running 2026-09-21 08:30:55 +02:00
mram da02457fd5 Pin compose example to v0.1.6
ci / test (push) Successful in 14s
ci / docker (push) Successful in 1m7s
ci / release (push) Successful in 14s
2026-09-21 08:22:42 +02:00
mram 6a6359a630 Add --force-update: immediate signed-update check, stage + service restart 2026-09-21 08:22:11 +02:00
mram e4e4348e8d Start the Windows service right after install (parity with enable --now) 2026-09-21 08:12:13 +02:00
mram b421bb7bfb Add --version and print the version atop every help and error screen 2026-09-21 08:10:46 +02:00
mram 909918657f Fix install/remove success message (installd -> installed) 2026-09-21 08:08:24 +02:00
mram a0435c858e Bare run in a terminal prints the help screen; add -h/--help
A zero-argument invocation now shows the usage text when stdout is a
console (double-clicked exe, interactive shell). Without a terminal —
Docker entrypoint, services, pipes — a bare invocation still starts the
proxy, so the container image and service behavior are unchanged.
2026-09-21 08:06:35 +02:00
mram fbab0bba33 Relaunch through UAC when (un)installing the service unprivileged
--install-service/--remove-service on Windows no longer fail with
'Access is denied' from a normal shell: the process re-runs itself via
ShellExecuteEx 'runas', waits for the elevated child and mirrors its
exit code. The child gets --elevated-child and pauses for a keypress so
its console output stays readable. Declining the prompt reports
'UAC prompt declined'.
2026-09-21 08:00:24 +02:00
mram 802a64280f Self-install into canonical layout on both platforms, --no-copy to opt out
Windows: --install-service creates %ProgramFiles%\gpu-turnstile and
%ProgramData%\gpu-turnstile, copies the exe and (if absent) the env
file in, and registers the copy. Linux: binary goes to
/var/lib/gpu-turnstile (not /usr/local/sbin: replacing a running binary
needs directory write, which must not be granted on a shared system dir
to a sandboxed service). --no-copy registers the current location
as-is on both platforms.
2026-09-21 07:52:33 +02:00
mram a88955e35c Sandbox the systemd unit: DynamicUser, read-only FS, no capabilities
The Linux install now mirrors the Windows virtual-account hardening: the
unit runs with DynamicUser=yes (transient per-service UID, no login),
ProtectSystem=strict with only StateDirectory writable (the install dir,
so self-update can rewrite the binary), NoNewPrivileges, empty
capability sets, restricted address families and a @system-service
syscall filter. Install copies the binary to /var/lib/gpu-turnstile and
the config to /etc/gpu-turnstile.env; Remove cleans up the unit and
binary but keeps the config.
2026-09-21 00:11:25 +02:00
mram 97624470eb Install the Windows service as the NT SERVICE virtual account only
--install-service now registers the service under
NT SERVICE\gpu-turnstile (low-privilege, per-service, no password) and
grants it modify access to the install dir (for self-updates) and the
LOG_FILE dir, plus read access to an external config file. Grants run
after CreateService because the virtual account's SID does not exist
before registration; a failed grant rolls back the registration.
2026-09-20 23:53:54 +02:00
mram 75f16a0229 Copy go.sum into the Docker build; pin compose example to v0.1.5
ci / test (push) Successful in 13s
ci / docker (push) Successful in 1m7s
ci / release (push) Successful in 14s
The image build broke with the first Linux-imported dependency
(go-systemd): the Dockerfile copied go.mod only, and the missing go.sum
never mattered while golang.org/x/sys was Windows-only.
2026-09-20 22:49:55 +02:00
mram 9481fd8418 Embed rotated release public key; pin compose example to v0.1.4
ci / test (push) Successful in 13s
ci / docker (push) Failing after 50s
ci / release (push) Successful in 27s
2026-09-20 22:37:28 +02:00
mram 83812cebf3 Add Linux systemd support and --install-service/--remove-service flags
The service package now has a Linux implementation alongside the Windows
one: systemd unit install/remove (/etc/systemd/system), readiness
notification (READY=1 via go-systemd) sent only after the listeners are
bound, a 30s watchdog, and STOPPING on shutdown. Listeners are pre-bound
so port conflicts fail fast and the readiness signal is truthful. The
notify calls are no-ops without NOTIFY_SOCKET (containers, shells) and
on non-Linux builds. --install-service/--remove-service work on both
platforms; the 'service install|remove' subcommand remains as an alias.
2026-09-20 22:35:59 +02:00
mram 30a1f55aae Make consumers optional: a URL enables its mode, empty disables it
OLLAMA_URL and COMFY_URL no longer have defaults; each consumer
(listener, client, startup probe, lock participation) is enabled by
setting its URL and disabled by leaving it empty. At least one must be
set. Ollama-only mode is a pure pass-through; ComfyUI-only mode skips
the unload and warm-reload steps. /metrics is now served on both
listeners. This is the extension pattern for future consumers such as
local game detection.
2026-09-20 22:29:16 +02:00
mram 642cc36a39 Pass signing key via env so its value is never echoed in CI logs 2026-09-20 22:21:25 +02:00
mram 19e19281de Pin compose example to v0.1.3
ci / test (push) Successful in 12s
ci / docker (push) Successful in 1m8s
ci / release (push) Failing after 25s
2026-09-20 22:16:09 +02:00
mram 08d02d8fa8 Embed release public key; rename signing secret to RELEASE_SIGNING_KEY 2026-09-20 22:15:19 +02:00
mram d1e01f9b78 Ignore signing/ directory (Ed25519 release keys) 2026-09-20 22:13:10 +02:00
mram bacb26772a Add LLM busy modes: wait (hang) or reject with Retry-After
LLM_BUSY_MODE=reject answers blocked LLM requests immediately with
LLM_BUSY_STATUS (default 503, 429 works) and Retry-After, so routers
like LiteLLM can cool down and retry instead of holding a hung
connection. The default wait mode now also sends Retry-After when
LLM_WAIT_TIMEOUT expires. Document the service account (LocalSystem
default, NT SERVICE virtual-account hardening) and the Program
Files / ProgramData install layout.
2026-09-20 21:43:03 +02:00
mram e363c9e4e4 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.
2026-09-20 21:33:29 +02:00
mram 5f6a22c2cf Pin compose example to v0.1.2
ci / test (push) Successful in 47s
ci / docker (push) Successful in 1m5s
2026-09-20 20:07:38 +02:00
mram 42e1386811 Extend retry backoff; rework logging (LOGLEVEL, request lines, colors)
Backoff now covers all dial-phase errors (refused, timeout, DNS), TLS
handshake failures, and 5xx responses with replayable bodies; streamed
POSTs are never replayed to avoid duplicate work. First retry of an
episode logs at WARN, subsequent attempts at INFO.

LOGLEVEL (LOG_LEVEL kept as alias) now defaults to warn: startup logs
version plus every setting; INFO adds one line per incoming request and
per response with status/duration, ANSI-colored in text mode (bypasses
slog's escaping so colors render in docker compose logs); NO_COLOR or
LOG_FORMAT=json disables colors.
2026-09-20 20:00:57 +02:00
mram 20d5439b5f Retry upstream connection-refused with exponential backoff
A refused dial (service down/restarting) is retried with a wait that
doubles from BACKOFF_INITIAL (1s) up to BACKOFF_MAX (60s) until the
upstream answers or the client disconnects. Handles the Windows WSA
errno (10061) as well as POSIX ECONNREFUSED. compose.yaml.example now
uses host.docker.internal like the working local deployment.
2026-09-20 19:45:48 +02:00
29 changed files with 3964 additions and 232 deletions
+51
View File
@@ -60,3 +60,54 @@ jobs:
tags: | tags: |
${{ steps.meta.outputs.image }}:${{ gitea.ref_name }} ${{ steps.meta.outputs.image }}:${{ gitea.ref_name }}
${{ steps.meta.outputs.image }}:latest ${{ 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
env:
RELEASE_SIGNING_KEY: ${{ secrets.RELEASE_SIGNING_KEY }}
run: |
# The key comes via the environment so its value never appears in
# the echoed command line of the run log.
printf '%s\n' "$RELEASE_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
+2
View File
@@ -2,3 +2,5 @@
/gpu-turnstile.exe /gpu-turnstile.exe
/compose.yaml /compose.yaml
/compose.yml /compose.yml
/signing/
+1 -1
View File
@@ -1,6 +1,6 @@
FROM golang:1.23 AS build FROM golang:1.23 AS build
WORKDIR /src WORKDIR /src
COPY go.mod ./ COPY go.mod go.sum ./
COPY cmd ./cmd COPY cmd ./cmd
COPY internal ./internal COPY internal ./internal
ARG VERSION=dev ARG VERSION=dev
+118 -14
View File
@@ -18,7 +18,10 @@ Open WebUI / n8n ────► :8188 ───┘
- LLM endpoints (`/api/generate`, `/api/chat`, `/api/embed`, `/v1/*`) take the - LLM endpoints (`/api/generate`, `/api/chat`, `/api/embed`, `/v1/*`) take the
LLM lock: concurrent requests allowed, but blocked while an image job is LLM lock: concurrent requests allowed, but blocked while an image job is
active or waiting (image priority). active or waiting (image priority). Blocked requests either hang until the
lock is free (`LLM_BUSY_MODE=wait`, default) or fail immediately with 503
(or 429) + `Retry-After` (`LLM_BUSY_MODE=reject`) — the latter lets routers
like LiteLLM cool down and retry instead of holding a hung connection.
- `POST /prompt` on the ComfyUI listener takes the image lock: new LLM - `POST /prompt` on the ComfyUI listener takes the image lock: new LLM
requests block, in-flight LLMs drain, Ollama models are unloaded, the prompt requests block, in-flight LLMs drain, Ollama models are unloaded, the prompt
is forwarded, and the lock is held until the job finishes and ComfyUI frees is forwarded, and the lock is held until the job finishes and ComfyUI frees
@@ -26,46 +29,142 @@ Open WebUI / n8n ────► :8188 ───┘
- Everything else (including websockets and all streaming) passes through - Everything else (including websockets and all streaming) passes through
transparently and unbuffered. transparently and unbuffered.
Each consumer is enabled by setting its URL (`OLLAMA_URL`, `COMFY_URL`) and
disabled by leaving it empty — at least one is required. With only Ollama
the proxy is a pass-through (no image jobs can arrive); with only ComfyUI
the Ollama unload/warm steps are skipped. Future consumers (e.g. local game
detection) plug into the same lock the same way.
## Configuration ## Configuration
All configuration is via environment variables; invalid values fail at Configuration comes from environment variables and/or an `.env`-style
startup. config file (`KEY=VALUE` lines, `#` comments). File lookup order:
`-config <path>` 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 | | Var | Default | Meaning |
|---|---|---| |---|---|---|
| `LISTEN_OLLAMA` | `:11434` | Ollama-facing listener | | `LISTEN_OLLAMA` | `:11434` | Ollama-facing listener |
| `LISTEN_COMFY` | `:8188` | ComfyUI-facing listener | | `LISTEN_COMFY` | `:8188` | ComfyUI-facing listener |
| `OLLAMA_URL` | `http://127.0.0.1:11435` | Ollama upstream | | `OLLAMA_URL` | _(empty = disabled)_ | Ollama upstream; set to enable the Ollama consumer |
| `COMFY_URL` | `http://127.0.0.1:8189` | ComfyUI upstream | | `COMFY_URL` | _(empty = disabled)_ | ComfyUI upstream; set to enable the ComfyUI consumer |
| `UNLOAD_TIMEOUT` | `60s` | Wait for Ollama to unload before an image job | | `UNLOAD_TIMEOUT` | `60s` | Wait for Ollama to unload before an image job |
| `JOB_TIMEOUT` | `15m` | Wait for a ComfyUI job to finish | | `JOB_TIMEOUT` | `15m` | Wait for a ComfyUI job to finish |
| `LLM_WAIT_TIMEOUT` | `10m` | Max lock wait for an LLM request before 503 | | `LLM_WAIT_TIMEOUT` | `10m` | Max lock wait for an LLM request before 503 (wait mode) |
| `LLM_BUSY_MODE` | `wait` | `wait` = hold blocked LLM requests; `reject` = fail them immediately |
| `LLM_BUSY_STATUS` | `503` | HTTP status for rejected LLM requests in reject mode (400599, e.g. 429) |
| `BUSY_RETRY_AFTER` | `30` | Seconds sent as `Retry-After` on busy responses (both modes) |
| `WARM_MODEL` | _(empty)_ | Model to reload after an image job (off by default) | | `WARM_MODEL` | _(empty)_ | Model to reload after an image job (off by default) |
| `LOG_LEVEL` | `info` | `debug` logs every lock transition | | `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_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 | | `UNLOAD_POLL_INTERVAL` | `500ms` | `/api/ps` poll interval while unloading |
| `HISTORY_POLL_INTERVAL` | `1s` | `/history/<id>` poll interval while a job runs | | `HISTORY_POLL_INTERVAL` | `1s` | `/history/<id>` poll interval while a job runs |
| `PROBE_TIMEOUT` | `5s` | Startup probe of both upstreams | | `PROBE_TIMEOUT` | `5s` | Startup probe of both upstreams |
| `FREE_TIMEOUT` | `30s` | `POST /free` call after an image job | | `FREE_TIMEOUT` | `30s` | `POST /free` call after an image job |
| `WARM_TIMEOUT` | `2m` | Warm-model reload after an image job | | `WARM_TIMEOUT` | `2m` | Warm-model reload after an image job |
| `SHUTDOWN_TIMEOUT` | `10s` | Graceful shutdown on SIGINT/SIGTERM | | `SHUTDOWN_TIMEOUT` | `10s` | Graceful shutdown on SIGINT/SIGTERM |
| `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) | | `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 |
| `APP_VER` | `stable` | `dev` disables updates, `stable` tracks latest, or pin an exact `vX.Y.Z` |
## Observability ## Observability
- `GET /healthz` (both listeners): `{"state":"idle|llm|image","llm_inflight":N,"image_pending":B}` - `GET /healthz` (both listeners): `{"state":"idle|llm|image","llm_inflight":N,"image_pending":B}`
- `GET /metrics` (Ollama listener): Prometheus text format — `gpu_turnstile_state`, - `GET /metrics` (both listeners): Prometheus text format — `gpu_turnstile_state`,
`gpu_turnstile_llm_inflight`, `gpu_turnstile_image_pending`, `gpu_turnstile_llm_inflight`, `gpu_turnstile_image_pending`,
`gpu_turnstile_image_jobs_total`, `gpu_turnstile_lock_wait_seconds` `gpu_turnstile_image_jobs_total`, `gpu_turnstile_lock_wait_seconds`
(histogram, `kind="llm|image"`), `gpu_turnstile_unload_seconds`. (histogram, `kind="llm|image"`), `gpu_turnstile_unload_seconds`.
- Logs: startup logs the version and every setting (visible even at the
default `warn` level). With `LOGLEVEL=info` or `debug`, every request
logs a `-->` incoming line and a `<--` response line with status and
duration — ANSI-colored (cyan incoming; green/yellow/red by status
class) in text mode, which renders in `docker compose logs` on Windows
Terminal. Set `NO_COLOR` to disable colors.
## Build and run ## Build and run
```sh ```sh
go build ./cmd/gpu-turnstile go build ./cmd/gpu-turnstile
./gpu-turnstile OLLAMA_URL=http://127.0.0.1:11435 COMFY_URL=http://127.0.0.1:8189 ./gpu-turnstile
``` ```
Running the binary with no arguments in a terminal prints the help screen
(same as `-h`/`--help`); without a terminal (services, containers) a bare
invocation starts the proxy.
### Run natively on Windows (current primary deployment)
Download `gpu-turnstile.exe` from a release and install it as a Windows
service — no admin shell needed, a UAC prompt appears automatically and
the elevated child does the work (its window waits for Enter so you can
read the result):
```sh
gpu-turnstile.exe --install-service # installs into Program Files, auto-start
gpu-turnstile.exe --install-service --no-copy # register in place instead
gpu-turnstile.exe --remove-service
```
Layout: `C:\Program Files\gpu-turnstile\` holds the exe and
`gpu-turnstile.env`, logs go to `C:\ProgramData\gpu-turnstile\`. If you
install without a config, the installer writes a sample env file with every
setting commented and explained — only `LOG_FILE` is active (a service has
no console). Your own `LOG_FILE` setting is always kept. The service always
runs
as the virtual account `NT SERVICE\gpu-turnstile` (low-privilege,
per-service, no password); the installer automatically grants it write
access to the install and data directories — nothing else to do.
Re-running `--install-service` is safe: it stops a running service,
replaces the installed binary only if it changed, fixes the registration
only where it drifted, and restarts the service only if it was running.
### Run natively on Linux (systemd)
The same binary works on Linux. Install it as a systemd service as root:
```sh
gpu-turnstile --install-service # installs into /var/lib/gpu-turnstile, enables + starts
gpu-turnstile --install-service --no-copy # register in place instead
gpu-turnstile --remove-service
```
The unit (`/etc/systemd/system/gpu-turnstile.service`) is `Type=notify`:
`systemctl start` blocks until the listeners are actually bound, a 30 s
watchdog restarts the process if it wedges, and logs land in the journal
(`journalctl -u gpu-turnstile -f`) unless `LOG_FILE` is set. Install
copies the binary to `/var/lib/gpu-turnstile/` and the config to
`/etc/gpu-turnstile.env` (edit that one after installing). The service
runs sandboxed with `DynamicUser=yes` — a transient low-privilege UID,
read-only filesystem except its install dir (so self-update keeps
working), no capabilities, syscall-filtered: same least-privilege idea as
the Windows virtual account. The notify integration is a no-op in
containers and interactive shells.
**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. `APP_VER` controls the target: `dev` disables updates, `stable`
(the default) tracks the latest release, and an exact `vX.Y.Z` pins that
release (even as a downgrade or to replace a dev build). Disable entirely
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 `RELEASE_SIGNING_KEY` repo secret, public key is committed).
`gpu-turnstile --force-update` checks immediately, stages the new binary
and restarts the running service (elevating via UAC only if needed).
### Docker
```sh ```sh
docker build -t gpu-turnstile . docker build -t gpu-turnstile .
docker run --rm -p 11434:11434 -p 8188:8188 \ docker run --rm -p 11434:11434 -p 8188:8188 \
@@ -76,9 +175,10 @@ docker run --rm -p 11434:11434 -p 8188:8188 \
Releases are built by Gitea Actions (`.gitea/workflows/ci.yml`): every push 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 runs `go vet` and `go test -race`, and pushing a semantic-version tag
`vX.Y.Z` builds and publishes `vX.Y.Z` publishes the container image
`git.rambossek.at/<owner>/gpu-turnstile:vX.Y.Z` (and updates `:latest`). (`git.rambossek.at/<owner>/gpu-turnstile:vX.Y.Z` plus `:latest`) and a
No images are built from branches. signed Windows binary attached to a Gitea release. Nothing is built from
branches.
The registry login needs one repository secret (Settings → Actions → The registry login needs one repository secret (Settings → Actions →
Secrets): `REGISTRY_TOKEN` — an access token with `write:package` scope. Secrets): `REGISTRY_TOKEN` — an access token with `write:package` scope.
@@ -91,13 +191,17 @@ go vet ./...
go test -race ./... 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/lock/ two-mode lock (LLM readers / image writer, FIFO)
internal/ollama/ ps / unload / warm client internal/ollama/ ps / unload / warm client
internal/comfy/ history poll / free client internal/comfy/ history poll / free client
internal/proxy/ handlers for both listeners internal/proxy/ handlers for both listeners
internal/metrics/ Prometheus exposition, no dependencies internal/metrics/ Prometheus exposition, no dependencies
internal/config/ env + .env file configuration
internal/update/ signed auto-updater
internal/service/ Windows service integration
``` ```
+200 -14
View File
@@ -42,6 +42,21 @@ listener is an `httputil.ReverseProxy` to its upstream. Websocket upgrades
(ComfyUI `/ws`) and streaming bodies (Ollama NDJSON / SSE) must pass through (ComfyUI `/ws`) and streaming bodies (Ollama NDJSON / SSE) must pass through
unbuffered (`FlushInterval = -1`). unbuffered (`FlushInterval = -1`).
### Modes of operation
Each GPU consumer is enabled by setting its URL and disabled by leaving it
empty — no separate flags. At least one URL must be set; a disabled
consumer gets no listener, no startup probe, and no lock participation:
- **Both set** (default deployment): full arbitration as described below.
- **Only `OLLAMA_URL`**: pure pass-through for Ollama; the LLM lock never
blocks since no image jobs can arrive.
- **Only `COMFY_URL`**: image jobs are tracked and ComfyUI's VRAM is freed
afterwards, but the Ollama unload and warm-reload steps are skipped.
- Future consumers (e.g. detecting a local game holding VRAM) plug into the
same lock the same way: enabled by their config knob, excluded when
absent.
### Lock semantics ### Lock semantics
Two-mode lock with image priority (writer-preferring RW lock, where "readers" Two-mode lock with image priority (writer-preferring RW lock, where "readers"
@@ -51,6 +66,11 @@ are LLM requests and the single "writer" is an image job):
`image` **or while an image job is waiting**. Then state := `llm`, n++. `image` **or while an image job is waiting**. Then state := `llm`, n++.
On completion (response fully written, including streamed bodies, or client On completion (response fully written, including streamed bodies, or client
disconnect) n--; if n == 0 state := `idle`. disconnect) n--; if n == 0 state := `idle`.
`LLM_BUSY_MODE` selects what a blocked LLM request sees: `wait` (default)
hangs until the lock is free or `LLM_WAIT_TIMEOUT` expires (then 503 +
`Retry-After`); `reject` answers immediately with `LLM_BUSY_STATUS`
(default 503; 429 works too) + `Retry-After: BUSY_RETRY_AFTER`, which
routers like LiteLLM honor for cooldowns/retries.
- **Image job**: `AcquireImage()` marks "image pending" (so no new LLM - **Image job**: `AcquireImage()` marks "image pending" (so no new LLM
requests start), waits until n == 0, sets state := `image`. Released after requests start), waits until n == 0, sets state := `image`. Released after
the ComfyUI job finished and models were freed. the ComfyUI job finished and models were freed.
@@ -78,7 +98,8 @@ ComfyUI listener (`:8188` → `COMFY_URL`):
### Image job flow (`POST /prompt`) ### Image job flow (`POST /prompt`)
1. `AcquireImage()`. 1. `AcquireImage()`.
2. Unload Ollama: `GET /api/ps`; for each model `POST /api/generate 2. Unload Ollama (skipped when `OLLAMA_URL` is unset): `GET /api/ps`; for
each model `POST /api/generate
{"model":M,"keep_alive":0}`; if that returns non-2xx (embedding-only {"model":M,"keep_alive":0}`; if that returns non-2xx (embedding-only
models), `POST /api/embed {"model":M,"input":"x","keep_alive":0}`. Poll models), `POST /api/embed {"model":M,"input":"x","keep_alive":0}`. Poll
`/api/ps` every `UNLOAD_POLL_INTERVAL` (default 500 ms) until empty or `/api/ps` every `UNLOAD_POLL_INTERVAL` (default 500 ms) until empty or
@@ -102,39 +123,181 @@ load time. Off by default.
## Configuration (env) ## Configuration (env)
Configuration comes from environment variables and/or an `.env`-style
config file (`KEY=VALUE` lines, `#` comments). File lookup order:
`-config <path>` 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 | | Var | Default | Meaning |
|---|---|---| |---|---|---|
| `LISTEN_OLLAMA` | `:11434` | Ollama-facing listener | | `LISTEN_OLLAMA` | `:11434` | Ollama-facing listener |
| `LISTEN_COMFY` | `:8188` | ComfyUI-facing listener | | `LISTEN_COMFY` | `:8188` | ComfyUI-facing listener |
| `OLLAMA_URL` | `http://127.0.0.1:11435` | upstream | | `OLLAMA_URL` | _(empty = disabled)_ | Ollama upstream; set to enable the Ollama consumer |
| `COMFY_URL` | `http://127.0.0.1:8189` | upstream | | `COMFY_URL` | _(empty = disabled)_ | ComfyUI upstream; set to enable the ComfyUI consumer |
| `UNLOAD_TIMEOUT` | `60s` | wait for Ollama to unload | | `UNLOAD_TIMEOUT` | `60s` | wait for Ollama to unload |
| `JOB_TIMEOUT` | `15m` | wait for ComfyUI job | | `JOB_TIMEOUT` | `15m` | wait for ComfyUI job |
| `LLM_WAIT_TIMEOUT` | `10m` | max time an LLM request waits for the lock before 503 | | `LLM_WAIT_TIMEOUT` | `10m` | max time an LLM request waits for the lock before 503 (wait mode) |
| `LLM_BUSY_MODE` | `wait` | `wait` = hold blocked LLM requests; `reject` = fail them immediately |
| `LLM_BUSY_STATUS` | `503` | HTTP status for rejected LLM requests in reject mode (400599, e.g. 429) |
| `BUSY_RETRY_AFTER` | `30` | seconds sent as `Retry-After` on busy responses (both modes) |
| `WARM_MODEL` | `` | optional model to reload after an image job | | `WARM_MODEL` | `` | optional model to reload after an image job |
| `LOG_LEVEL` | `info` | `debug` logs every lock transition | | `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 | | `UNLOAD_POLL_INTERVAL` | `500ms` | `/api/ps` poll interval while unloading |
| `HISTORY_POLL_INTERVAL` | `1s` | `/history/<id>` poll interval while a job runs | | `HISTORY_POLL_INTERVAL` | `1s` | `/history/<id>` poll interval while a job runs |
| `PROBE_TIMEOUT` | `5s` | startup probe of both upstreams | | `PROBE_TIMEOUT` | `5s` | startup probe of both upstreams |
| `FREE_TIMEOUT` | `30s` | `POST /free` call after an image job | | `FREE_TIMEOUT` | `30s` | `POST /free` call after an image job |
| `WARM_TIMEOUT` | `2m` | warm-model reload after an image job | | `WARM_TIMEOUT` | `2m` | warm-model reload after an image job |
| `SHUTDOWN_TIMEOUT` | `10s` | graceful shutdown on SIGINT/SIGTERM | | `SHUTDOWN_TIMEOUT` | `10s` | graceful shutdown on SIGINT/SIGTERM |
| `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) | | `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 |
| `APP_VER` | `stable` | version to run: `dev` disables updates, `stable` tracks the latest release, or an exact `vX.Y.Z` pin (up- or downgraded to) |
| `CFG_VER` | _(installer-managed)_ | config format reference written by `--install-service`; missing = the file is replaced with a fresh sample (backup `.bak`) |
Startup fails fast on unparsable values. Both upstreams are probed once at Startup fails fast on unparsable values and when neither consumer URL is
start (`/api/version`, `/system_stats`); failure is logged, not fatal. set. Enabled upstreams are probed once at start (`/api/version`,
`/system_stats`); failure is logged, not fatal.
## Native deployment (Windows and Linux)
The binary runs natively on Windows (the current primary deployment) and on
Linux with systemd (the future GPU server), as well as in Docker.
Service management is the same on both platforms:
`gpu-turnstile --install-service [-config path]` installs, registers and
starts an auto-start service; `--remove-service` stops and uninstalls it.
Both need admin/root; on Windows a non-elevated shell triggers a UAC
prompt instead of failing — the command relaunches itself elevated, waits
for the child, and mirrors its exit code. The legacy form
`gpu-turnstile service install|remove` does the same thing.
Re-running install on an already-registered service converges instead of
failing: a running service is stopped first, the installed binary copy is
refreshed only when the content differs, the registration (Windows service
config / systemd unit) is updated only where it drifted, and the service is
started again only if it was running before.
By default install creates the canonical layout and copies the binary into
it (Windows: `%ProgramFiles%\gpu-turnstile\`, plus
`%ProgramData%\gpu-turnstile\` for logs; Linux: `/var/lib/gpu-turnstile/`
with the config at `/etc/gpu-turnstile.env`). If there is no config at all,
install writes a sample env file covering every setting — each with a
comment line, everything commented out — except the `CFG_VER`/`APP_VER`
header and `LOG_FILE`, which is active on Windows
(`%ProgramData%\gpu-turnstile\gpu-turnstile.log`) since a service has no
console; on Linux it stays commented because stderr goes to the journal.
The first line is `CFG_VER=vX.Y.Z`, recording the installer version. When a
later version's install finds an older `CFG_VER`, it appends every setting
the file does not mention (commented or not) at the end and updates
`CFG_VER`; a file without `CFG_VER` is invalid and gets replaced by a fresh
sample, with the old content kept as `<file>.bak`. `--no-copy` registers
the current executable location as-is and leaves the config untouched.
### Windows
- `--install-service` creates `%ProgramFiles%\gpu-turnstile\` and
`%ProgramData%\gpu-turnstile\`, copies the exe and (if none exists there
yet) the `gpu-turnstile.env` into the Program Files directory, and
registers that copy as a Windows service; recovery actions restart it
5 s after any failure. The install ensures the env file sets `LOG_FILE`
to `%ProgramData%\gpu-turnstile\gpu-turnstile.log` since there is no
console — an existing `LOG_FILE` setting is kept.
- **Account**: the service always runs as the virtual account
`NT SERVICE\gpu-turnstile` — a per-service low-privilege identity the
SCM manages (no password, automatic logon-as-a-service right, no admin
rights, gone when the service is removed). The installer grants it
modify access to the install and data directories (self-updates rewrite
the exe) and the `LOG_FILE` directory (created if missing), plus read
access to the config file when it lives elsewhere. The grants happen
after service registration because the virtual account's SID only exists
from that point on; if a grant fails the service registration is rolled
back.
### Linux (systemd)
- `--install-service` copies the binary to `/var/lib/gpu-turnstile/`,
copies the config to `/etc/gpu-turnstile.env` if none exists there yet,
writes `/etc/systemd/system/gpu-turnstile.service`, then runs `systemctl
daemon-reload` and `enable --now`. `--remove-service` removes the unit
and the installed binary; the `/etc` config stays. The binary does not
go to `/usr/local/sbin` on purpose: replacing a running binary needs
write access to its *directory*, and granting the sandboxed service
write access to a shared system directory would let a compromised
service overwrite other binaries — `/var/lib/gpu-turnstile` is
exclusively ours.
- **Sandboxing** mirrors the Windows virtual account: the unit runs with
`DynamicUser=yes` — a transient per-service UID with no login, no home
and no password, managed entirely by systemd. `ProtectSystem=strict`
makes the filesystem read-only except `StateDirectory=gpu-turnstile`
(the install dir, so self-updates can rewrite the binary), plus
`NoNewPrivileges`, `ProtectHome`, `PrivateTmp`, `ProtectKernel*`,
`ProtectControlGroups`, `RestrictNamespaces`, `RestrictSUIDSGID`,
`RestrictRealtime`, `LockPersonality`, `MemoryDenyWriteExecute`, empty
capability sets, `RestrictAddressFamilies=AF_UNIX AF_INET AF_INET6` and
`SystemCallFilter=@system-service`. The proxy needs only outbound
TCP/UDP and the notify socket, so it loses nothing.
- The unit is `Type=notify`: the binary sends `READY=1` via
`github.com/coreos/go-systemd` only after the listeners are bound, so
`systemctl start` blocks until the proxy accepts connections. A 30 s
watchdog (`WatchdogSec=`) is pinged as long as the process runs; three
missed pings make systemd restart it. `STOPPING=1` is sent on shutdown.
All notify calls are no-ops when `NOTIFY_SOCKET` is unset (containers,
interactive shells), and the whole integration is Linux-only — Windows
builds carry no-op stubs.
- Logs go to the journal (`journalctl -u gpu-turnstile`) or to `LOG_FILE`.
- **Auto-update** works the same as on Windows: `Restart=on-failure` with
`RestartSec=5s` brings up the staged binary after the updater exits with
code 3.
- **Auto-update**: on startup and every `UPDATE_INTERVAL`, the binary
consults `APP_VER`: `dev` disables updates; `stable` (the default)
fetches `UPDATE_REPO`'s latest release and applies it when its tag is a
newer `vX.Y.Z` (a `dev` binary cannot be compared and is replaced by the
latest release); a `vX.Y.Z` pin fetches that exact tag and stages it on
any difference, including downgrades. Applying means downloading
`UPDATE_ASSET` plus its `.sig` (and `.sha256` when present) and verifying
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".
Builds without an embedded public key never update.
- **`--force-update`** runs the same check immediately, single-shot: one
attempt with a 30 s timeout, then exit — "up to date" (exit 0) or the
error (exit 1), no retries. When a newer release is found it downloads,
verifies and stages it, and if the service is running it restarts it
right away (otherwise the new version applies on next start). On Windows
it elevates via UAC only when the stage or restart needs permissions the
caller does not have.
- **Signing setup (one time)**: `openssl genpkey -algorithm ed25519 -out
private.pem`; `openssl pkey -in private.pem -pubout -out public.pem`.
Private key → repo secret `RELEASE_SIGNING_KEY`; public key → committed into
`internal/update/pubkey.go`. CI signs release binaries with
`openssl pkeyutl -sign -rawin`.
## Observability ## Observability
- `GET /healthz` on both listeners: 200 with JSON - `GET /healthz` on both listeners: 200 with JSON
`{"state":"idle|llm|image","llm_inflight":N,"image_pending":B}`. `{"state":"idle|llm|image","llm_inflight":N,"image_pending":B}`.
- `GET /metrics` on the Ollama listener: Prometheus text format, no external - `GET /metrics` on both listeners: Prometheus text format, no external
dependency needed: dependency needed:
`gpu_turnstile_state{state="…"} 1`, `gpu_turnstile_llm_inflight`, `gpu_turnstile_state{state="…"} 1`, `gpu_turnstile_llm_inflight`,
`gpu_turnstile_image_jobs_total`, `gpu_turnstile_lock_wait_seconds` `gpu_turnstile_image_jobs_total`, `gpu_turnstile_lock_wait_seconds`
(histogram, label `kind="llm|image"`), `gpu_turnstile_unload_seconds`. (histogram, label `kind="llm|image"`), `gpu_turnstile_unload_seconds`.
- Structured logs (`log/slog`, JSON when `LOG_FORMAT=json`), one line per - Structured logs (`log/slog`, JSON when `LOG_FORMAT=json`), one line per
state transition and per image job phase with `prompt_id`. state transition and per image job phase with `prompt_id`. Startup logs
the version and every setting (visible even at the default `warn`
level). With `LOGLEVEL=info` or `debug`, every request logs a `-->`
incoming line and a `<--` response line with status and duration —
ANSI-colored (cyan incoming; green/yellow/red by status class) in text
mode, which renders in `docker compose logs` on Windows Terminal. Set
`NO_COLOR` to disable colors.
## Edge cases to handle ## Edge cases to handle
@@ -146,6 +309,14 @@ start (`/api/version`, `/system_stats`); failure is logged, not fatal.
`JOB_TIMEOUT` releases the lock; log at warn. `JOB_TIMEOUT` releases the lock; log at warn.
- Ollama unreachable during unload: continue with the image job; the whole - Ollama unreachable during unload: continue with the image job; the whole
point is not to block users on a misbehaving neighbour. point is not to block users on a misbehaving neighbour.
- Upstream unreachable while proxying (connection refused, dial timeout,
DNS failure, TLS handshake error): retry with exponential backoff —
`BACKOFF_INITIAL`, doubling per attempt, capped at `BACKOFF_MAX` — until
the upstream answers or the client disconnects. These are safe to retry:
the request never reached the upstream application. 5xx responses are
retried the same way, but only when the request body can be replayed
(GETs, or bodies with `GetBody`); streamed POSTs are never replayed to
avoid duplicate work such as a double-enqueued ComfyUI prompt.
- `POST /prompt` with a body that ComfyUI rejects (400): lock released - `POST /prompt` with a body that ComfyUI rejects (400): lock released
immediately, body passed back. immediately, body passed back.
- Websocket `/ws` connections are long-lived and never take the lock. - Websocket `/ws` connections are long-lived and never take the lock.
@@ -155,11 +326,15 @@ start (`/api/version`, `/system_stats`); failure is logged, not fatal.
``` ```
gpu-turnstile/ 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/lock/lock.go # two-mode lock + tests
internal/ollama/client.go # ps / unload / warm internal/ollama/client.go # ps / unload / warm
internal/comfy/client.go # history poll / free internal/comfy/client.go # history poll / free
internal/proxy/ # handlers for both listeners 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 SCM + Linux systemd (notify/watchdog) integration
Dockerfile Dockerfile
.gitea/workflows/ci.yml .gitea/workflows/ci.yml
README.md README.md
@@ -181,11 +356,18 @@ are new.
held until `/free` was called. held until `/free` was called.
- Streaming test: fake Ollama emits chunks with delays; assert the client - Streaming test: fake Ollama emits chunks with delays; assert the client
receives the first chunk before the last is sent (no buffering). 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 skipped, APP_VER=dev and
pinned releases honored.
## Build and CI ## Build and CI
- Go 1.23+, stdlib only. `CGO_ENABLED=0`, `-ldflags="-s -w"`, version from - Go 1.23+, two external dependencies: `golang.org/x/sys` (Windows service
`git describe` injected via `-X main.version=`. integration) and `github.com/coreos/go-systemd` (systemd notify/watchdog,
Linux build only). `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 - Dockerfile: multi-stage, final image `gcr.io/distroless/static` (or
`scratch`), non-root user, `EXPOSE 8188 11434`, `scratch`), non-root user, `EXPOSE 8188 11434`,
`ENTRYPOINT ["/gpu-turnstile"]`. `ENTRYPOINT ["/gpu-turnstile"]`.
@@ -200,8 +382,12 @@ are new.
Login uses the repo secret `REGISTRY_TOKEN` (an access token with Login uses the repo secret `REGISTRY_TOKEN` (an access token with
`write:package` scope) because the automatic `GITEA_TOKEN` cannot push `write:package` scope) because the automatic `GITEA_TOKEN` cannot push
packages; the username is just `gitea.actor`. packages; the username is just `gitea.actor`.
- Release: a git tag `vX.Y.Z` produces the versioned image; the Open WebUI 3. on a version tag: also build the Windows binary, sign it with OpenSSL
compose pins that tag. No images are built from branches. (`RELEASE_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) ## Deployment (target)
+543 -157
View File
@@ -3,228 +3,614 @@
package main package main
import ( import (
"bufio"
"context" "context"
"errors" "errors"
"fmt" "fmt"
"io"
"io/fs"
"log/slog" "log/slog"
"net"
"net/http" "net/http"
"os" "os"
"os/signal" "os/signal"
"strconv" "path/filepath"
"strings" "strings"
"syscall" "syscall"
"time" "time"
"gpu-turnstile/internal/comfy" "gpu-turnstile/internal/comfy"
"gpu-turnstile/internal/config"
"gpu-turnstile/internal/lock" "gpu-turnstile/internal/lock"
"gpu-turnstile/internal/metrics" "gpu-turnstile/internal/metrics"
"gpu-turnstile/internal/ollama" "gpu-turnstile/internal/ollama"
"gpu-turnstile/internal/proxy" "gpu-turnstile/internal/proxy"
"gpu-turnstile/internal/service"
"gpu-turnstile/internal/update"
) )
// version is injected at build time via -ldflags "-X main.version=...". // version is injected at build time via -ldflags "-X main.version=...".
var version = "dev" var version = "dev"
type config struct { // exitCodeUpdate tells the service recovery configuration to restart the
listenOllama string // process: a signed update has been staged and the GPU lock is idle.
listenComfy string const exitCodeUpdate = 3
ollamaURL string
comfyURL string
unloadTimeout time.Duration
jobTimeout time.Duration
llmWaitTimeout time.Duration
unloadPollInterval time.Duration // stdoutIsTerminal reports whether stdout is a console (char device), as
historyPollInterval time.Duration // opposed to a pipe or file — which is what Docker containers and services
probeTimeout time.Duration // see.
freeTimeout time.Duration func stdoutIsTerminal() bool {
warmTimeout time.Duration fi, err := os.Stdout.Stat()
shutdownTimeout time.Duration return err == nil && fi.Mode()&os.ModeCharDevice != 0
promptCaptureLimit int64
warmModel string
logLevel slog.Level
logJSON bool
} }
func envDuration(getenv func(string) string, name string, dst *time.Duration) error { // parseFlags extracts -config <path> (or -config=<path>), the
v := getenv(name) // --install-service / --remove-service switches, --no-copy, -h/--help,
if v == "" { // -v/--version, --force-update and the hidden --elevated-child marker from
return nil // args.
} func parseFlags(args []string) (configPath string, install, remove, noCopy, help, showVersion, forceUpdate, elevatedChild bool, rest []string) {
d, err := time.ParseDuration(v) rest = args[:0]
if err != nil { for i := 0; i < len(args); i++ {
return fmt.Errorf("%s: %w", name, err) switch {
} case args[i] == "-config" && i+1 < len(args):
*dst = d configPath = args[i+1]
return nil i++
} case strings.HasPrefix(args[i], "-config="):
configPath = strings.TrimPrefix(args[i], "-config=")
func loadConfig(getenv func(string) string) (config, error) { case args[i] == "--install-service" || args[i] == "-install-service":
cfg := config{ install = true
listenOllama: ":11434", case args[i] == "--remove-service" || args[i] == "-remove-service":
listenComfy: ":8188", remove = true
ollamaURL: "http://127.0.0.1:11435", case args[i] == "--no-copy" || args[i] == "-no-copy":
comfyURL: "http://127.0.0.1:8189", noCopy = true
unloadTimeout: time.Minute, case args[i] == "-h" || args[i] == "--help" || args[i] == "-help":
jobTimeout: 15 * time.Minute, help = true
llmWaitTimeout: 10 * time.Minute, case args[i] == "-v" || args[i] == "--version" || args[i] == "-version":
showVersion = true
unloadPollInterval: 500 * time.Millisecond, case args[i] == "--force-update" || args[i] == "-force-update":
historyPollInterval: time.Second, forceUpdate = true
probeTimeout: 5 * time.Second, case args[i] == "--elevated-child":
freeTimeout: 30 * time.Second, elevatedChild = true
warmTimeout: 2 * time.Minute,
shutdownTimeout: 10 * time.Second,
promptCaptureLimit: 64 * 1024,
logLevel: slog.LevelInfo,
}
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},
} {
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("LOG_LEVEL"); v != "" {
var level slog.Level
if err := level.UnmarshalText([]byte(v)); err != nil {
return cfg, fmt.Errorf("LOG_LEVEL: %w", err)
}
cfg.logLevel = level
}
switch strings.ToLower(getenv("LOG_FORMAT")) {
case "", "text":
case "json":
cfg.logJSON = true
default: default:
return cfg, fmt.Errorf("LOG_FORMAT: must be \"text\" or \"json\"") rest = append(rest, args[i])
} }
return cfg, nil }
return configPath, install, remove, noCopy, help, showVersion, forceUpdate, elevatedChild, rest
}
// versionLine is printed at the top of every help and error screen.
func versionLine() string { return "gpu-turnstile " + version }
const usageText = `GPU arbitration proxy for Ollama + ComfyUI
Usage:
gpu-turnstile -config <path> run the proxy
gpu-turnstile --install-service [--no-copy] [-config path] install + start as a service
gpu-turnstile --remove-service stop + uninstall the service
gpu-turnstile -v | --version print just the version
gpu-turnstile --force-update check for a signed update now,
apply it and restart the service
gpu-turnstile -h | --help this help
Options:
-config <path> config file (default: gpu-turnstile.env next to the exe)
--install-service copies the binary into the canonical location
(%ProgramFiles%\gpu-turnstile or /var/lib/gpu-turnstile)
unless --no-copy; on Windows a UAC prompt appears when
the shell is not elevated
--no-copy with --install-service: register the current location as-is
All runtime settings are environment variables or KEY=VALUE lines in the
config file (OLLAMA_URL, COMFY_URL, LOGLEVEL, ...); see README.md.
`
// printHelp prints the version header plus the full help text.
func printHelp() {
fmt.Printf("%s — %s", versionLine(), usageText)
}
// fatalUsage prints the version header, an error message and the one-line
// usage summary, then exits with code 2.
func fatalUsage(format string, args ...any) {
fmt.Fprintf(os.Stderr, "%s\n\n", versionLine())
fmt.Fprintf(os.Stderr, format+"\n\n", args...)
fmt.Fprintln(os.Stderr, "usage: gpu-turnstile [-config path] [--install-service [--no-copy] | --remove-service]")
fmt.Fprintln(os.Stderr, " gpu-turnstile service install|remove [-config path]")
os.Exit(2)
} }
func main() { func main() {
cfg, err := loadConfig(os.Getenv) configPath, install, remove, noCopy, help, showVersion, forceUpdate, elevatedChild, args := parseFlags(os.Args[1:])
if err != nil { if showVersion {
fmt.Fprintf(os.Stderr, "gpu-turnstile: %v\n", err) fmt.Println(version)
os.Exit(1) return
}
bare := configPath == "" && !install && !remove && !forceUpdate && !elevatedChild && len(args) == 0
if help || (bare && stdoutIsTerminal()) {
// Bare invocation in a terminal (e.g. double-clicked on Windows)
// shows the help instead of starting a proxy window with no visible
// explanation. Without a terminal — Docker containers, services,
// pipes — a bare invocation starts the proxy as before.
printHelp()
return
}
if len(args) > 0 && args[0] == "service" {
// Legacy subcommand form: gpu-turnstile service install|remove.
if len(args) != 2 || (args[1] != "install" && args[1] != "remove") {
fatalUsage("error: expected 'service install' or 'service remove'")
}
install = args[1] == "install"
remove = !install
args = nil
}
switch {
case install && remove:
fatalUsage("error: --install-service and --remove-service are mutually exclusive")
case forceUpdate && (install || remove):
fatalUsage("error: --force-update cannot be combined with --install-service/--remove-service")
case install:
os.Exit(serviceCommand(configPath, true, noCopy, elevatedChild))
case remove:
os.Exit(serviceCommand(configPath, false, noCopy, elevatedChild))
case forceUpdate:
os.Exit(forceUpdateCommand(configPath, elevatedChild))
}
if len(args) > 0 {
fatalUsage("error: unknown arguments: %s", strings.Join(args, " "))
} }
opts := &slog.HandlerOptions{Level: cfg.logLevel} if exePath, err := os.Executable(); err == nil {
var handler slog.Handler = slog.NewTextHandler(os.Stderr, opts) update.CleanupOld(exePath)
if cfg.logJSON { }
handler = slog.NewJSONHandler(os.Stderr, opts)
cfg, err := loadMergedConfig(configPath)
if err != nil {
fmt.Fprintf(os.Stderr, "%s\n\ngpu-turnstile: %v\n", versionLine(), err)
os.Exit(1)
}
log, logOut, logCloser := newLogger(cfg)
defer logCloser.Close()
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)
}
}
// 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) log := slog.New(handler)
slog.SetDefault(log) slog.SetDefault(log)
return log, out, closer
}
log.Info("starting gpu-turnstile", // waitForEnter keeps an elevated child's console window open until the
// user has read the output.
func waitForEnter() {
fmt.Print("\nPress Enter to close this window...")
bufio.NewReader(os.Stdin).ReadString('\n')
}
// elevateAndMirror relaunches the current command elevated (UAC) and
// mirrors the child's exit code. verb is used in messages.
func elevateAndMirror(verb string) (int, bool) {
args := append(append([]string{}, os.Args[1:]...), "--elevated-child")
code, err := service.RelaunchElevated(args)
if errors.Is(err, service.ErrUserCancelled) {
fmt.Fprintln(os.Stderr, "gpu-turnstile: UAC prompt declined")
return 1, true
}
if err != nil {
fmt.Fprintf(os.Stderr, "gpu-turnstile: could not elevate: %v\n", err)
return 1, true
}
if code != 0 {
fmt.Fprintf(os.Stderr, "gpu-turnstile %s failed in the elevated process (exit %d)\n", verb, code)
return code, true
}
return 0, true
}
// isPermission reports whether err is a permission problem (Windows
// ERROR_ACCESS_DENIED, POSIX EACCES/EPERM, possibly wrapped).
func isPermission(err error) bool {
return errors.Is(err, fs.ErrPermission) || strings.Contains(strings.ToLower(err.Error()), "access is denied")
}
// forceUpdateCommand checks for a signed update immediately, stages it if
// newer, and restarts the service when it is running so the new binary
// takes effect. Staging into a system directory and restarting a service
// need admin rights; instead of prompting unconditionally, permission
// failures trigger the UAC relaunch so a dev copy in a user-writable
// directory updates without a prompt.
func forceUpdateCommand(configPath string, elevatedChild bool) int {
if elevatedChild {
defer waitForEnter()
}
cfg, err := loadMergedConfig(configPath)
if errors.Is(err, config.ErrNoConsumer) {
err = nil // update settings do not depend on a consumer URL
}
if err != nil {
fmt.Fprintf(os.Stderr, "%s\n\ngpu-turnstile: %v\n", versionLine(), err)
return 1
}
exePath, err := os.Executable()
if err != nil {
fmt.Fprintf(os.Stderr, "gpu-turnstile: cannot locate executable: %v\n", err)
return 1
}
log, _, logCloser := newLogger(cfg)
defer logCloser.Close()
if cfg.AppVersion == "dev" {
fmt.Printf("%s: APP_VER=dev, updates disabled\n", versionLine())
return 0
}
u := &update.Updater{Repo: cfg.UpdateRepo, Asset: cfg.UpdateAsset, Version: version, Desired: cfg.AppVersion, Log: log}
// Single-shot: one attempt, fail fast when the server is unreachable
// instead of hanging in a TCP connect for minutes.
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
staged, err := u.Check(ctx, exePath)
if err != nil && isPermission(err) && !service.Elevated() {
code, _ := elevateAndMirror("--force-update")
if code == 0 {
fmt.Println("update applied (elevated)")
}
return code
}
if err != nil {
fmt.Fprintf(os.Stderr, "%s\n\ngpu-turnstile: update check failed: %v\n", versionLine(), err)
return 1
}
if !staged {
fmt.Printf("%s is up to date\n", versionLine())
return 0
}
fmt.Printf("%s: update staged\n", versionLine())
restarted, err := service.RestartIfRunning()
if err != nil && isPermission(err) && !service.Elevated() {
code, _ := elevateAndMirror("--force-update")
if code == 0 {
fmt.Println("update applied (elevated)")
}
return code
}
if err != nil {
fmt.Fprintf(os.Stderr, "gpu-turnstile: update staged but service restart failed: %v\n", err)
return 1
}
if restarted {
fmt.Println("service restarted on the new version")
} else {
fmt.Println("no running service; the new version applies on next start")
}
return 0
}
// serviceCommand installs (copyBin = register the canonical-layout copy)
// or removes the service and reports the result. On Windows, when the
// shell is not elevated, the command relaunches itself through a UAC
// prompt and mirrors the elevated child's exit code. An elevated child
// waits for a keypress so its console window does not flash closed before
// the output can be read.
func serviceCommand(configPath string, install, noCopy, elevatedChild bool) int {
verb, doneVerb := "remove", "removed"
if install {
verb, doneVerb = "install", "installed"
}
if elevatedChild {
defer waitForEnter()
}
if !service.Elevated() {
code, done := elevateAndMirror(verb)
if done && code != 0 {
return code
}
if done {
fmt.Printf("service %s: %s (elevated)\n", service.Name, doneVerb)
return 0
}
}
var err error
if install {
path := resolveConfigPath(configPath)
if abs, absErr := filepath.Abs(path); absErr == nil {
path = abs
}
err = service.Install(path, !noCopy, version)
} else {
err = service.Remove()
}
if err != nil {
fmt.Fprintf(os.Stderr, "gpu-turnstile service %s: %v\n", verb, err)
return 1
}
fmt.Printf("service %s: %s\n", service.Name, doneVerb)
return 0
}
// orDisabled renders an empty URL as "disabled" for the startup dump.
func orDisabled(url string) string {
if url == "" {
return "disabled"
}
return url
}
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(ctx, slog.LevelWarn, "starting gpu-turnstile",
"version", version, "version", version,
"listen_ollama", cfg.listenOllama, "listen_ollama", cfg.ListenOllama,
"listen_comfy", cfg.listenComfy, "listen_comfy", cfg.ListenComfy,
"ollama_url", cfg.ollamaURL, "ollama_url", orDisabled(cfg.OllamaURL),
"comfy_url", cfg.comfyURL, "comfy_url", orDisabled(cfg.ComfyURL),
"unload_timeout", cfg.UnloadTimeout,
"job_timeout", cfg.JobTimeout,
"llm_wait_timeout", cfg.LLMWaitTimeout,
"llm_busy_mode", cfg.LLMBusyMode,
"llm_busy_status", cfg.LLMBusyStatus,
"busy_retry_after", cfg.BusyRetryAfter,
"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)
if err != nil { // Each consumer is enabled by setting its URL; a disabled consumer gets
log.Error("invalid configuration", "err", err) // no client, no listener and no probe.
os.Exit(1) var ollamaClient *ollama.Client
var err error
if cfg.OllamaURL != "" {
if ollamaClient, err = ollama.New(cfg.OllamaURL, log); err != nil {
return err
}
}
var comfyClient *comfy.Client
if cfg.ComfyURL != "" {
if comfyClient, err = comfy.New(cfg.ComfyURL, log); err != nil {
return err
} }
comfyClient, err := comfy.New(cfg.comfyURL, log)
if err != nil {
log.Error("invalid configuration", "err", err)
os.Exit(1)
} }
srv, err := proxy.New(proxy.Config{ srv, err := proxy.New(proxy.Config{
OllamaURL: cfg.ollamaURL, OllamaURL: cfg.OllamaURL,
ComfyURL: cfg.comfyURL, ComfyURL: cfg.ComfyURL,
Lock: lock.New(log), Lock: lk,
Ollama: ollamaClient, Ollama: ollamaClient,
Comfy: comfyClient, Comfy: comfyClient,
Metrics: metrics.New(), Metrics: metrics.New(),
Log: log, Log: log,
LLMWaitTimeout: cfg.llmWaitTimeout, LogColor: !cfg.LogJSON && cfg.LogFile == "" && os.Getenv("NO_COLOR") == "",
UnloadTimeout: cfg.unloadTimeout, LogWriter: logOut,
JobTimeout: cfg.jobTimeout, LLMWaitTimeout: cfg.LLMWaitTimeout,
UnloadPollInterval: cfg.unloadPollInterval, UnloadTimeout: cfg.UnloadTimeout,
HistoryPollInterval: cfg.historyPollInterval, JobTimeout: cfg.JobTimeout,
FreeTimeout: cfg.freeTimeout, UnloadPollInterval: cfg.UnloadPollInterval,
WarmTimeout: cfg.warmTimeout, HistoryPollInterval: cfg.HistoryPollInterval,
PromptCaptureLimit: cfg.promptCaptureLimit, FreeTimeout: cfg.FreeTimeout,
WarmModel: cfg.warmModel, WarmTimeout: cfg.WarmTimeout,
LLMBusyMode: cfg.LLMBusyMode,
LLMBusyStatus: cfg.LLMBusyStatus,
BusyRetryAfter: cfg.BusyRetryAfter,
BackoffInitial: cfg.BackoffInitial,
BackoffMax: cfg.BackoffMax,
PromptCaptureLimit: cfg.PromptCaptureLimit,
WarmModel: cfg.WarmModel,
}) })
if err != nil { if err != nil {
log.Error("invalid configuration", "err", err) return err
os.Exit(1)
} }
// Probe both upstreams once; failure is logged, not fatal. // Probe the enabled upstreams once; failure is logged, not fatal.
probeCtx, probeCancel := context.WithTimeout(context.Background(), cfg.probeTimeout) probeCtx, probeCancel := context.WithTimeout(ctx, cfg.ProbeTimeout)
if ollamaClient != nil {
if err := ollamaClient.Probe(probeCtx); err != nil { 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 comfyClient != nil {
if err := comfyClient.Probe(probeCtx); err != nil { 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() probeCancel()
ollamaSrv := &http.Server{Addr: cfg.listenOllama, Handler: srv.OllamaHandler()} // Bind the listeners up front so a port conflict fails fast and the
comfySrv := &http.Server{Addr: cfg.listenComfy, Handler: srv.ComfyHandler()} // readiness notification below really means "accepting connections".
var servers []*http.Server
var listeners []net.Listener
bind := func(addr string, handler http.Handler, consumer string) error {
ln, err := net.Listen("tcp", addr)
if err != nil {
return fmt.Errorf("listen %s on %s: %w", consumer, addr, err)
}
servers = append(servers, &http.Server{Addr: addr, Handler: handler})
listeners = append(listeners, ln)
log.Warn("listening", "consumer", consumer, "addr", addr)
return nil
}
if ollamaClient != nil {
if err := bind(cfg.ListenOllama, srv.OllamaHandler(), "ollama"); err != nil {
return err
}
}
if comfyClient != nil {
if err := bind(cfg.ListenComfy, srv.ComfyHandler(), "comfy"); err != nil {
return err
}
}
errCh := make(chan error, 2) errCh := make(chan error, len(servers))
go func() { errCh <- ollamaSrv.ListenAndServe() }() for i := range servers {
go func() { errCh <- comfySrv.ListenAndServe() }() go func(s *http.Server, ln net.Listener) { errCh <- s.Serve(ln) }(servers[i], listeners[i])
}
sigCtx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) // Tell systemd we are up and start the watchdog pings; both are no-ops
defer stop() // when not running under a notify/watchdog unit.
service.NotifyReady()
service.StartWatchdog(ctx)
if cfg.AutoUpdate {
go updateLoop(ctx, cfg, log, lk, isService)
}
select { select {
case err := <-errCh: case err := <-errCh:
if err != nil && !errors.Is(err, http.ErrServerClosed) { if err != nil && !errors.Is(err, http.ErrServerClosed) {
log.Error("listener failed", "err", err) return err
os.Exit(1)
} }
case <-sigCtx.Done(): case <-ctx.Done():
log.Info("shutting down") log.Info("shutting down")
} }
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), cfg.shutdownTimeout) shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), cfg.ShutdownTimeout)
defer shutdownCancel() defer shutdownCancel()
ollamaSrv.Shutdown(shutdownCtx) for _, s := range servers {
comfySrv.Shutdown(shutdownCtx) s.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, Desired: cfg.AppVersion, 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):
}
}
} }
+15 -4
View File
@@ -1,17 +1,28 @@
# Example deployment for gpu-turnstile. Copy to compose.yaml and adjust. # Example deployment for gpu-turnstile. Copy to compose.yaml and adjust.
# #
# gpu-turnstile listens on the ports the services normally use; the actual # gpu-turnstile listens on the ports the services normally use; the actual
# Ollama and ComfyUI instances run one port higher on the workstation. # Ollama and ComfyUI instances run one port higher (11435 / 8189) and must
# bind 0.0.0.0 so the container can reach them (OLLAMA_HOST=0.0.0.0:11435,
# ComfyUI --listen 0.0.0.0 --port 8189).
services: services:
gpu-turnstile: gpu-turnstile:
image: git.rambossek.at/public/gpu-turnstile:v0.1.1 image: git.rambossek.at/public/gpu-turnstile:v0.1.7
restart: unless-stopped restart: unless-stopped
environment: environment:
OLLAMA_URL: http://<workstation-ip>:11435 # Each consumer is enabled by setting its URL; leave one unset to
COMFY_URL: http://<workstation-ip>:8189 # disable that side (no listener, no probe, no lock participation).
# Services on the Docker host itself:
OLLAMA_URL: http://host.docker.internal:11435
COMFY_URL: http://host.docker.internal:8189
# Services on another machine: use its LAN IP instead, e.g.
# OLLAMA_URL: http://192.168.1.10:11435
# COMFY_URL: http://192.168.1.10:8189
# UNLOAD_TIMEOUT: 60s # UNLOAD_TIMEOUT: 60s
# JOB_TIMEOUT: 15m # JOB_TIMEOUT: 15m
# LLM_WAIT_TIMEOUT: 10m # LLM_WAIT_TIMEOUT: 10m
# LLM_BUSY_MODE: reject # wait (default) hangs; reject fails fast
# LLM_BUSY_STATUS: 429 # status in reject mode (default 503)
# BUSY_RETRY_AFTER: 30 # Retry-After seconds on busy responses
# WARM_MODEL: qwen3:14b # WARM_MODEL: qwen3:14b
# LOG_LEVEL: info # LOG_LEVEL: info
# LOG_FORMAT: json # LOG_FORMAT: json
+4
View File
@@ -1,3 +1,7 @@
module gpu-turnstile module gpu-turnstile
go 1.23 go 1.23
require golang.org/x/sys v0.29.0
require github.com/coreos/go-systemd/v22 v22.7.0
+4
View File
@@ -0,0 +1,4 @@
github.com/coreos/go-systemd/v22 v22.7.0 h1:LAEzFkke61DFROc7zNLX/WA2i5J8gYqe0rSj9KI28KA=
github.com/coreos/go-systemd/v22 v22.7.0/go.mod h1:xNUYtjHu2EDXbsxz1i41wouACIwT7Ybq9o0BQhMwD0w=
golang.org/x/sys v0.29.0 h1:TPYlXGxvx1MGTn2GiZDhnjPA9wZzZeGKHHmKhHYvgaU=
golang.org/x/sys v0.29.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
+248
View File
@@ -0,0 +1,248 @@
// Package config loads gpu-turnstile's configuration from environment
// variables and an optional .env-style config file.
package config
import (
"bufio"
"errors"
"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
// AppVersion is the version the user wants to run: "dev" disables
// updates, "stable" tracks the latest release, anything else is an
// exact vX.Y.Z release to pin. From APP_VER; defaults to "stable".
AppVersion string
// LLMBusyMode is "wait" (hold requests until the lock is free or
// LLMWaitTimeout expires) or "reject" (immediately answer with
// LLMBusyStatus + Retry-After when an image job is active or pending).
LLMBusyMode string
LLMBusyStatus int
BusyRetryAfter int
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. The upstream URLs default to empty: a
// consumer is enabled by setting its URL, disabled by leaving it empty.
func Defaults() Config {
return Config{
ListenOllama: ":11434",
ListenComfy: ":8188",
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",
AppVersion: "stable",
LLMBusyMode: "wait",
LLMBusyStatus: 503,
BusyRetryAfter: 30,
LogLevel: slog.LevelWarn,
}
}
// ErrNoConsumer is returned by Load when neither OLLAMA_URL nor COMFY_URL
// is set. Commands that never talk to an upstream (--force-update) may
// ignore it and proceed with the rest of the configuration.
var ErrNoConsumer = errors.New("at least one of OLLAMA_URL or COMFY_URL must be set (each URL enables its consumer)")
// 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
}
if v := getenv("LLM_BUSY_MODE"); v != "" {
if v != "wait" && v != "reject" {
return cfg, fmt.Errorf("LLM_BUSY_MODE: must be \"wait\" or \"reject\"")
}
cfg.LLMBusyMode = v
}
if v := getenv("LLM_BUSY_STATUS"); v != "" {
n, err := strconv.Atoi(v)
if err != nil || n < 400 || n > 599 {
return cfg, fmt.Errorf("LLM_BUSY_STATUS: must be an HTTP status in 400-599")
}
cfg.LLMBusyStatus = n
}
if v := getenv("BUSY_RETRY_AFTER"); v != "" {
n, err := strconv.Atoi(v)
if err != nil || n <= 0 {
return cfg, fmt.Errorf("BUSY_RETRY_AFTER: must be a positive integer (seconds)")
}
cfg.BusyRetryAfter = n
}
if v := getenv("APP_VER"); v != "" {
switch {
case v == "dev" || v == "stable":
cfg.AppVersion = v
default:
if _, ok := parseVersion(v); !ok {
return cfg, fmt.Errorf("APP_VER: must be \"dev\", \"stable\" or a vX.Y.Z version")
}
cfg.AppVersion = "v" + strings.TrimPrefix(v, "v")
}
}
// 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\"")
}
if cfg.OllamaURL == "" && cfg.ComfyURL == "" {
return cfg, ErrNoConsumer
}
return cfg, nil
}
+151
View File
@@ -0,0 +1,151 @@
package config
import (
"errors"
"log/slog"
"strings"
"testing"
"time"
)
func TestDefaults(t *testing.T) {
cfg, err := Load(func(k string) string {
if k == "OLLAMA_URL" {
return "http://127.0.0.1:11435"
}
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.ComfyURL != "" {
t.Fatalf("ComfyURL default = %q, want empty (disabled)", cfg.ComfyURL)
}
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 TestLoadRequiresConsumer(t *testing.T) {
_, err := Load(func(string) string { return "" })
if err == nil || !strings.Contains(err.Error(), "OLLAMA_URL") {
t.Fatalf("err = %v, want missing-consumer error", err)
}
if !errors.Is(err, ErrNoConsumer) {
t.Fatalf("err = %v, want errors.Is(err, ErrNoConsumer)", err)
}
}
func TestAppVersion(t *testing.T) {
load := func(appVer string) (Config, error) {
return Load(func(k string) string {
switch k {
case "OLLAMA_URL":
return "http://127.0.0.1:11435"
case "APP_VER":
return appVer
}
return ""
})
}
cfg, err := load("")
if err != nil || cfg.AppVersion != "stable" {
t.Fatalf("default AppVersion = %q, err %v; want stable", cfg.AppVersion, err)
}
for _, v := range []string{"dev", "stable"} {
if cfg, err := load(v); err != nil || cfg.AppVersion != v {
t.Fatalf("APP_VER=%s: got %q, err %v", v, cfg.AppVersion, err)
}
}
if cfg, err := load("1.2.3"); err != nil || cfg.AppVersion != "v1.2.3" {
t.Fatalf("APP_VER=1.2.3: got %q, err %v; want normalized v1.2.3", cfg.AppVersion, err)
}
if _, err := load("nightly"); err == nil {
t.Fatal("APP_VER=nightly: want validation error")
}
}
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"},
{"LLM_BUSY_MODE", "bogus"},
{"LLM_BUSY_STATUS", "200"},
{"BUSY_RETRY_AFTER", "0"},
} {
_, err := Load(func(k string) string {
if k == tc.key {
return tc.value
}
if k == "OLLAMA_URL" {
return "http://127.0.0.1:11435"
}
return ""
})
if err == nil {
t.Errorf("%s=%s: expected error", tc.key, tc.value)
}
}
}
+176
View File
@@ -0,0 +1,176 @@
package config
import (
"fmt"
"strconv"
"strings"
)
// sampleEntry is one line pair in the sample env file: a comment and the
// (usually commented-out) KEY=VALUE line.
type sampleEntry struct {
name string
value string
comment string
active bool // rendered uncommented
}
// appVerComment documents APP_VER wherever it is rendered.
const appVerComment = `Version to run: "dev" disables updates, "stable" tracks the latest release, or pin an exact release like v0.1.7`
// sampleEntries lists every setting in sample order. logFile activates the
// LOG_FILE line (Windows install); empty keeps it commented like the rest.
// CFG_VER and APP_VER are not entries — they head the file, always active.
func sampleEntries(logFile string) []sampleEntry {
return []sampleEntry{
{"LISTEN_OLLAMA", ":11434", "Ollama-facing listener address", false},
{"LISTEN_COMFY", ":8188", "ComfyUI-facing listener address", false},
{"OLLAMA_URL", "http://127.0.0.1:11434", "Ollama upstream URL; setting it enables the Ollama consumer (default: empty = disabled)", false},
{"COMFY_URL", "http://127.0.0.1:8188", "ComfyUI upstream URL; setting it enables the ComfyUI consumer (default: empty = disabled)", false},
{"WARM_MODEL", "", "Optional model to reload after an image job (default: empty = none)", false},
{"UNLOAD_TIMEOUT", "60s", "How long to wait for Ollama to unload a model", false},
{"JOB_TIMEOUT", "15m", "Maximum time to wait for a ComfyUI job", false},
{"LLM_WAIT_TIMEOUT", "10m", "Max time an LLM request waits for the GPU before being answered 503 (wait mode)", false},
{"LLM_BUSY_MODE", "wait", `How blocked LLM requests are handled: "wait" holds them, "reject" fails them immediately`, false},
{"LLM_BUSY_STATUS", "503", "HTTP status for rejected LLM requests in reject mode (400-599, e.g. 429)", false},
{"BUSY_RETRY_AFTER", "30", "Seconds sent as Retry-After on busy responses (both modes)", false},
{"LOGLEVEL", "warn", `Log verbosity: debug, info, warn, error; "info" logs every request (LOG_LEVEL works too)`, false},
{"LOG_FORMAT", "text", `Log format: "text" or "json"`, false},
{"LOG_FILE", logFile, "Append logs to this file instead of stderr (a Windows service has no console)", logFile != ""},
{"UNLOAD_POLL_INTERVAL", "500ms", "/api/ps poll interval while unloading", false},
{"HISTORY_POLL_INTERVAL", "1s", "/history/<id> poll interval while a job runs", false},
{"PROBE_TIMEOUT", "5s", "Startup probe timeout for the enabled upstreams", false},
{"FREE_TIMEOUT", "30s", "Timeout for the POST /free call after an image job", false},
{"WARM_TIMEOUT", "2m", "Timeout for the warm-model reload after an image job", false},
{"SHUTDOWN_TIMEOUT", "10s", "Graceful shutdown timeout on SIGINT/SIGTERM", false},
{"BACKOFF_INITIAL", "1s", "First retry wait when an upstream connection fails", false},
{"BACKOFF_MAX", "60s", "Cap for the exponential retry backoff", false},
{"PROMPT_CAPTURE_LIMIT", "65536", "Bytes of the /prompt response buffered to find prompt_id (pass-through unaffected)", false},
{"AUTO_UPDATE", "true", "Poll the releases API for signed updates", false},
{"UPDATE_INTERVAL", "6h", "Auto-update check interval", false},
{"UPDATE_REPO", "https://git.rambossek.at/PUBLIC/gpu-turnstile", "Repository to check for releases", false},
{"UPDATE_ASSET", "gpu-turnstile.exe", "Release asset to download", false},
}
}
// SampleEnv renders a sample .env file covering every setting, each with a
// comment line. Everything is commented out — so all defaults apply —
// except the CFG_VER/APP_VER header and LOG_FILE when logFile is non-empty
// (a Windows service has no console). CFG_VER records the version that
// wrote the file so later installs can upgrade it.
func SampleEnv(version, logFile string) string {
var b strings.Builder
fmt.Fprintf(&b, "CFG_VER=%s\n", version)
b.WriteString("# Config format reference, written by the installer — do not edit.\n")
b.WriteString("# The installer uses it to append newly added settings on updates.\n\n")
b.WriteString("# gpu-turnstile configuration\n")
b.WriteString("# KEY=VALUE lines; \"#\" starts a comment. Every setting below is at its\n")
b.WriteString("# default and commented out — remove the \"#\" to change it.\n")
b.WriteString("# At least one of OLLAMA_URL / COMFY_URL must be set for the proxy to start.\n\n")
writeEntry(&b, sampleEntry{"APP_VER", "stable", appVerComment, true})
for _, e := range sampleEntries(logFile) {
writeEntry(&b, e)
}
return b.String()
}
func writeEntry(b *strings.Builder, e sampleEntry) {
fmt.Fprintf(b, "# %s\n", e.comment)
if e.active {
fmt.Fprintf(b, "%s=%s\n\n", e.name, e.value)
} else {
fmt.Fprintf(b, "#%s=%s\n\n", e.name, e.value)
}
}
// SyncSample upgrades an installer-written env file: when its CFG_VER says
// it was written by an older gpu-turnstile, every setting the file does not
// mention — commented or not — is appended at the end, and CFG_VER is
// updated to version. Files without CFG_VER (hand-written or foreign),
// up-to-date files and "dev" builds are returned unchanged; changed reports
// whether the returned content differs.
func SyncSample(data, version, logFile string) (string, bool) {
if version == "" || version == "dev" {
return data, false
}
values, err := ParseEnvFile(strings.NewReader(data))
if err != nil || values["CFG_VER"] == "" {
return data, false // not installer-written; the caller decides
}
if compareVersions(values["CFG_VER"], version) >= 0 {
return data, false // same or newer
}
present := make(map[string]bool)
for _, line := range strings.Split(data, "\n") {
line = strings.TrimSpace(strings.TrimPrefix(strings.TrimSpace(line), "#"))
if name, _, ok := strings.Cut(line, "="); ok {
present[strings.TrimSpace(name)] = true
}
}
lines := strings.Split(data, "\n")
for i, l := range lines {
if strings.HasPrefix(strings.TrimSpace(l), "CFG_VER=") {
lines[i] = "CFG_VER=" + version
break
}
}
out := strings.Join(lines, "\n")
if !strings.HasSuffix(out, "\n") {
out += "\n"
}
var b strings.Builder
if !present["APP_VER"] {
fmt.Fprintf(&b, "\n# Added by gpu-turnstile %s:\n", version)
writeEntry(&b, sampleEntry{"APP_VER", "stable", appVerComment, true})
}
for _, e := range sampleEntries(logFile) {
if !present[e.name] {
fmt.Fprintf(&b, "\n# Added by gpu-turnstile %s:\n", version)
writeEntry(&b, e)
}
}
out += b.String()
return out, out != data
}
// compareVersions orders two vX.Y.Z version strings. Unparseable versions
// (including "dev") sort before any release.
func compareVersions(a, b string) int {
pa, oka := parseVersion(a)
pb, okb := parseVersion(b)
if oka != okb {
if oka {
return 1
}
return -1
}
for i := range pa {
if pa[i] != pb[i] {
if pa[i] > pb[i] {
return 1
}
return -1
}
}
return 0
}
func parseVersion(v string) ([3]int, bool) {
var out [3]int
v = strings.TrimPrefix(strings.TrimSpace(v), "v")
parts := strings.Split(v, ".")
if len(parts) != 3 {
return out, false
}
for i, p := range parts {
n, err := strconv.Atoi(p)
if err != nil || n < 0 {
return out, false
}
out[i] = n
}
return out, true
}
+143
View File
@@ -0,0 +1,143 @@
package config
import (
"strings"
"testing"
)
const sampleLogPath = `C:\ProgramData\gpu-turnstile\gpu-turnstile.log`
var allSettingNames = []string{
"LISTEN_OLLAMA", "LISTEN_COMFY", "OLLAMA_URL", "COMFY_URL",
"WARM_MODEL", "UNLOAD_TIMEOUT", "JOB_TIMEOUT", "LLM_WAIT_TIMEOUT",
"LLM_BUSY_MODE", "LLM_BUSY_STATUS", "BUSY_RETRY_AFTER",
"LOGLEVEL", "LOG_FORMAT", "LOG_FILE",
"UNLOAD_POLL_INTERVAL", "HISTORY_POLL_INTERVAL", "PROBE_TIMEOUT",
"FREE_TIMEOUT", "WARM_TIMEOUT", "SHUTDOWN_TIMEOUT",
"BACKOFF_INITIAL", "BACKOFF_MAX", "PROMPT_CAPTURE_LIMIT",
"AUTO_UPDATE", "UPDATE_INTERVAL", "UPDATE_REPO", "UPDATE_ASSET",
}
func TestSampleEnv(t *testing.T) {
sample := SampleEnv("v0.1.7", sampleLogPath)
if !strings.HasPrefix(sample, "CFG_VER=v0.1.7\n") {
t.Errorf("first line does not carry CFG_VER: %q", strings.SplitN(sample, "\n", 2)[0])
}
// Every setting known to Load must appear.
for _, name := range allSettingNames {
if !strings.Contains(sample, name+"=") {
t.Errorf("sample is missing %s", name)
}
}
// The sample must parse cleanly; active values are CFG_VER, APP_VER
// and LOG_FILE.
values, err := ParseEnvFile(strings.NewReader(sample))
if err != nil {
t.Fatalf("sample does not parse: %v", err)
}
want := map[string]string{"CFG_VER": "v0.1.7", "APP_VER": "stable", "LOG_FILE": sampleLogPath}
if len(values) != len(want) {
t.Fatalf("active values = %v, want %v", values, want)
}
for k, v := range want {
if values[k] != v {
t.Errorf("%s = %q, want %q", k, values[k], v)
}
}
// Without a log path LOG_FILE stays commented out.
values, err = ParseEnvFile(strings.NewReader(SampleEnv("v0.1.7", "")))
if err != nil {
t.Fatalf("sample without log path does not parse: %v", err)
}
if len(values) != 2 || values["LOG_FILE"] != "" {
t.Fatalf("active values = %v, want only CFG_VER and APP_VER", values)
}
}
func TestSyncSample(t *testing.T) {
old := SampleEnv("v0.1.6", sampleLogPath)
// Simulate a setting that did not exist yet when the file was written.
old = strings.Replace(old, "#UPDATE_ASSET=gpu-turnstile.exe\n", "", 1)
out, changed := SyncSample(old, "v0.1.7", sampleLogPath)
if !changed {
t.Fatal("older installer file was not upgraded")
}
if !strings.Contains(out, "\nCFG_VER=v0.1.7\n") && !strings.HasPrefix(out, "CFG_VER=v0.1.7\n") {
t.Error("CFG_VER was not updated to the new version")
}
if !strings.Contains(out, "#UPDATE_ASSET=gpu-turnstile.exe") {
t.Error("missing setting was not appended")
}
if !strings.Contains(out, "LOG_FILE="+sampleLogPath) {
t.Error("existing active LOG_FILE was lost")
}
if !strings.Contains(out, "Added by gpu-turnstile v0.1.7") {
t.Error("appended section is not attributed")
}
if _, err := ParseEnvFile(strings.NewReader(out)); err != nil {
t.Fatalf("upgraded file does not parse: %v", err)
}
// A file written by the same or a newer version is left alone.
if _, changed := SyncSample(SampleEnv("v0.1.7", sampleLogPath), "v0.1.7", sampleLogPath); changed {
t.Error("same-version file was modified")
}
if _, changed := SyncSample(SampleEnv("v0.2.0", sampleLogPath), "v0.1.7", sampleLogPath); changed {
t.Error("newer-version file was modified")
}
// Files without CFG_VER are not installer-written; the installer
// replaces them, SyncSample leaves them alone.
user := "OLLAMA_URL=http://host:11434\n"
if out, changed := SyncSample(user, "v0.1.7", sampleLogPath); changed || out != user {
t.Error("file without CFG_VER was modified")
}
// A dev build never upgrades.
if _, changed := SyncSample(old, "dev", sampleLogPath); changed {
t.Error("dev build modified the file")
}
}
func TestSyncSampleAppendsMissingAppVer(t *testing.T) {
old := SampleEnv("v0.1.6", sampleLogPath)
old = strings.Replace(old, "# "+appVerComment+"\n", "", 1)
old = strings.Replace(old, "APP_VER=stable\n", "", 1)
out, changed := SyncSample(old, "v0.1.7", sampleLogPath)
if !changed {
t.Fatal("file without APP_VER was not upgraded")
}
values, err := ParseEnvFile(strings.NewReader(out))
if err != nil {
t.Fatalf("upgraded file does not parse: %v", err)
}
if values["APP_VER"] != "stable" {
t.Fatalf("APP_VER = %q, want appended default \"stable\"", values["APP_VER"])
}
}
func TestCompareVersions(t *testing.T) {
cases := []struct {
a, b string
want int
}{
{"v0.1.6", "v0.1.7", -1},
{"v0.1.7", "v0.1.7", 0},
{"v1.0.0", "v0.9.9", 1},
{"0.1.7", "v0.1.6", 1},
{"dev", "v0.1.7", -1},
{"v0.1.7", "dev", 1},
{"dev", "dev", 0},
}
for _, c := range cases {
if got := compareVersions(c.a, c.b); got != c.want {
t.Errorf("compareVersions(%q, %q) = %d, want %d", c.a, c.b, got, c.want)
}
}
}
+16
View File
@@ -74,6 +74,22 @@ func (l *Lock) AcquireLLM(ctx context.Context) error {
return nil return nil
} }
// TryAcquireLLM acquires one in-flight LLM slot without waiting and
// reports whether it succeeded. It fails when an image job is active or
// pending.
func (l *Lock) TryAcquireLLM() bool {
l.mu.Lock()
if l.imageActive || len(l.imageQ) > 0 {
l.mu.Unlock()
return false
}
l.n++
n := l.n
l.mu.Unlock()
l.logTransition("lock transition", "state", StateLLM, "llm_inflight", n)
return true
}
// ReleaseLLM marks one LLM request as finished. // ReleaseLLM marks one LLM request as finished.
func (l *Lock) ReleaseLLM() { func (l *Lock) ReleaseLLM() {
l.mu.Lock() l.mu.Lock()
+51
View File
@@ -0,0 +1,51 @@
package lock
import (
"context"
"testing"
"time"
)
func TestTryAcquireLLM(t *testing.T) {
lk := New(nil)
if !lk.TryAcquireLLM() {
t.Fatal("TryAcquireLLM on idle lock should succeed")
}
if _, n, _ := lk.Snapshot(); n != 1 {
t.Fatalf("n = %d, want 1", n)
}
// While an image job is pending, TryAcquireLLM must fail.
imageWaiting := make(chan struct{})
go func() {
lk.AcquireImage(context.Background())
close(imageWaiting)
}()
deadline := time.Now().Add(2 * time.Second)
for {
lk.mu.Lock()
queued := len(lk.imageQ)
lk.mu.Unlock()
if queued == 1 {
break
}
if time.Now().After(deadline) {
t.Fatal("image waiter never queued")
}
time.Sleep(time.Millisecond)
}
if lk.TryAcquireLLM() {
t.Fatal("TryAcquireLLM with image pending should fail")
}
lk.ReleaseLLM()
<-imageWaiting
if lk.TryAcquireLLM() {
t.Fatal("TryAcquireLLM with image active should fail")
}
lk.ReleaseImage()
if !lk.TryAcquireLLM() {
t.Fatal("TryAcquireLLM after image release should succeed")
}
lk.ReleaseLLM()
}
+319 -18
View File
@@ -4,15 +4,22 @@
package proxy package proxy
import ( import (
"bufio"
"bytes" "bytes"
"context" "context"
"crypto/tls"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"io"
"log/slog" "log/slog"
"net"
"net/http" "net/http"
"net/http/httputil" "net/http/httputil"
"net/url" "net/url"
"os"
"strconv"
"strings"
"time" "time"
"gpu-turnstile/internal/comfy" "gpu-turnstile/internal/comfy"
@@ -37,10 +44,32 @@ type Config struct {
Metrics *metrics.Metrics Metrics *metrics.Metrics
Log *slog.Logger Log *slog.Logger
// 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 LLMWaitTimeout time.Duration
UnloadTimeout time.Duration UnloadTimeout time.Duration
JobTimeout time.Duration JobTimeout time.Duration
// LLMBusyMode is "wait" (default) or "reject". In reject mode an LLM
// request that arrives while an image job is active or pending is
// answered immediately with LLMBusyStatus and a Retry-After header
// (BusyRetryAfter seconds) instead of waiting for the lock. In wait
// mode the Retry-After header is sent when LLMWaitTimeout expires.
LLMBusyMode string
LLMBusyStatus int
BusyRetryAfter int
// BackoffInitial and BackoffMax control the exponential retry backoff
// when an upstream refuses a connection: the wait doubles from
// BackoffInitial up to BackoffMax between attempts. Zero selects the
// defaults (1s / 60s).
BackoffInitial time.Duration
BackoffMax time.Duration
// UnloadPollInterval and HistoryPollInterval override the clients' // UnloadPollInterval and HistoryPollInterval override the clients'
// /api/ps and /history poll intervals when > 0. // /api/ps and /history poll intervals when > 0.
UnloadPollInterval time.Duration UnloadPollInterval time.Duration
@@ -59,32 +88,51 @@ type Config struct {
type Server struct { type Server struct {
cfg Config cfg Config
log *slog.Logger log *slog.Logger
logWriter io.Writer
freeTimeout time.Duration freeTimeout time.Duration
warmTimeout time.Duration warmTimeout time.Duration
captureLimit int64 captureLimit int64
backoffInitial time.Duration
backoffMax time.Duration
busyMode string
busyStatus int
busyRetryAfter int
ollamaProxy *httputil.ReverseProxy ollamaProxy *httputil.ReverseProxy
comfyProxy *httputil.ReverseProxy comfyProxy *httputil.ReverseProxy
} }
// New builds a Server, validating the upstream URLs. // New builds a Server, validating the upstream URLs. At least one of
// OllamaURL / ComfyURL must be set; an empty URL disables that consumer —
// its handler is then never served, its client may be nil, and the image
// job flow skips the Ollama unload/warm steps.
func New(cfg Config) (*Server, error) { func New(cfg Config) (*Server, error) {
ollamaURL, err := url.Parse(cfg.OllamaURL) if cfg.OllamaURL == "" && cfg.ComfyURL == "" {
if err != nil || ollamaURL.Scheme == "" || ollamaURL.Host == "" { return nil, fmt.Errorf("at least one of OllamaURL or ComfyURL is required")
}
var ollamaURL, comfyURL *url.URL
if cfg.OllamaURL != "" {
u, err := url.Parse(cfg.OllamaURL)
if err != nil || u.Scheme == "" || u.Host == "" {
return nil, fmt.Errorf("invalid OLLAMA_URL %q", cfg.OllamaURL) return nil, fmt.Errorf("invalid OLLAMA_URL %q", cfg.OllamaURL)
} }
comfyURL, err := url.Parse(cfg.ComfyURL) ollamaURL = u
if err != nil || comfyURL.Scheme == "" || comfyURL.Host == "" { }
if cfg.ComfyURL != "" {
u, err := url.Parse(cfg.ComfyURL)
if err != nil || u.Scheme == "" || u.Host == "" {
return nil, fmt.Errorf("invalid COMFY_URL %q", cfg.ComfyURL) return nil, fmt.Errorf("invalid COMFY_URL %q", cfg.ComfyURL)
} }
comfyURL = u
}
log := cfg.Log log := cfg.Log
if log == nil { if log == nil {
log = slog.Default() log = slog.Default()
} }
if cfg.UnloadPollInterval > 0 { if cfg.UnloadPollInterval > 0 && cfg.Ollama != nil {
cfg.Ollama.PollInterval = cfg.UnloadPollInterval cfg.Ollama.PollInterval = cfg.UnloadPollInterval
} }
if cfg.HistoryPollInterval > 0 { if cfg.HistoryPollInterval > 0 && cfg.Comfy != nil {
cfg.Comfy.PollInterval = cfg.HistoryPollInterval cfg.Comfy.PollInterval = cfg.HistoryPollInterval
} }
freeTimeout := cfg.FreeTimeout freeTimeout := cfg.FreeTimeout
@@ -99,19 +147,163 @@ func New(cfg Config) (*Server, error) {
if captureLimit <= 0 { if captureLimit <= 0 {
captureLimit = defaultCaptureLimit captureLimit = defaultCaptureLimit
} }
return &Server{ backoffInitial := cfg.BackoffInitial
if backoffInitial <= 0 {
backoffInitial = time.Second
}
backoffMax := cfg.BackoffMax
if backoffMax <= 0 {
backoffMax = time.Minute
}
busyMode := "wait"
if cfg.LLMBusyMode == "reject" {
busyMode = "reject"
}
busyStatus := cfg.LLMBusyStatus
if busyStatus == 0 {
busyStatus = http.StatusServiceUnavailable
}
busyRetryAfter := cfg.BusyRetryAfter
if busyRetryAfter <= 0 {
busyRetryAfter = 30
}
retry := &retryTransport{
base: http.DefaultTransport,
initial: backoffInitial,
max: backoffMax,
log: log,
}
s := &Server{
cfg: cfg, cfg: cfg,
log: log, log: log,
logWriter: cfg.LogWriter,
freeTimeout: freeTimeout, freeTimeout: freeTimeout,
warmTimeout: warmTimeout, warmTimeout: warmTimeout,
captureLimit: captureLimit, captureLimit: captureLimit,
ollamaProxy: newReverseProxy(ollamaURL, log.With("upstream", "ollama")), backoffInitial: backoffInitial,
comfyProxy: newReverseProxy(comfyURL, log.With("upstream", "comfy")), backoffMax: backoffMax,
}, nil busyMode: busyMode,
busyStatus: busyStatus,
busyRetryAfter: busyRetryAfter,
}
if ollamaURL != nil {
s.ollamaProxy = newReverseProxy(ollamaURL, retry, log.With("upstream", "ollama"))
}
if comfyURL != nil {
s.comfyProxy = newReverseProxy(comfyURL, retry, log.With("upstream", "comfy"))
}
return s, nil
} }
func newReverseProxy(target *url.URL, log *slog.Logger) *httputil.ReverseProxy { // retryTransport retries requests whose failure means the upstream never
// saw them — any dial-phase error (connection refused, dial timeout, DNS
// failure), TLS handshake errors — plus 5xx responses when the request
// body can be replayed (GETs and requests with GetBody set). The wait
// doubles from initial up to max between attempts. The loop runs until
// the request succeeds, fails in a non-retryable way, or the client's
// context is cancelled.
type retryTransport struct {
base http.RoundTripper
initial time.Duration
max time.Duration
log *slog.Logger
}
// shouldRetry reports whether a RoundTrip error means the request never
// reached the upstream application and is therefore safe to send again.
func shouldRetry(err error) bool {
// Dial-phase failures: refused, timeout, unreachable, DNS (wrapped).
var opErr *net.OpError
if errors.As(err, &opErr) && opErr.Op == "dial" {
return true
}
var dnsErr *net.DNSError
if errors.As(err, &dnsErr) {
return true
}
// TLS handshake failures: the HTTP request was never written.
var recordErr tls.RecordHeaderError
if errors.As(err, &recordErr) {
return true
}
var certErr *tls.CertificateVerificationError
if errors.As(err, &certErr) {
return true
}
var alertErr tls.AlertError
if errors.As(err, &alertErr) {
return true
}
return false
}
// replayable reports whether the request body can be sent again. Bodies
// streamed from the client (GetBody == nil) cannot, so 5xx responses to
// POSTs are not retried: the upstream may have partially processed them,
// and re-sending could duplicate work (e.g. a second ComfyUI prompt).
func replayable(req *http.Request) bool {
return req.Body == nil || req.Body == http.NoBody || req.GetBody != nil
}
func (t *retryTransport) logAt(level slog.Level, msg string, args ...any) {
if t.log != nil && t.log.Enabled(context.Background(), level) {
t.log.Log(context.Background(), level, msg, args...)
}
}
func (t *retryTransport) RoundTrip(req *http.Request) (*http.Response, error) {
wait := t.initial
attempt := 0
for {
resp, err := t.base.RoundTrip(req)
switch {
case err != nil && !shouldRetry(err):
return nil, err
case err == nil && (resp.StatusCode < 500 || !replayable(req)):
return resp, nil
}
// Retryable failure: a transport error, or a 5xx response.
var reason string
if err != nil {
reason = err.Error()
} else {
reason = resp.Status
io.Copy(io.Discard, resp.Body)
resp.Body.Close()
if req.GetBody != nil {
if body, berr := req.GetBody(); berr == nil {
req.Body = body
}
}
}
attempt++
// One WARN per outage episode; subsequent attempts at INFO.
level := slog.LevelInfo
if attempt == 1 {
level = slog.LevelWarn
}
t.logAt(level, "upstream unavailable; retrying with backoff",
"path", req.URL.Path, "reason", reason, "attempt", attempt, "retry_in", wait)
select {
case <-req.Context().Done():
if err != nil {
return nil, err
}
return nil, req.Context().Err()
case <-time.After(wait):
}
wait *= 2
if wait > t.max {
wait = t.max
}
}
}
func newReverseProxy(target *url.URL, transport http.RoundTripper, log *slog.Logger) *httputil.ReverseProxy {
return &httputil.ReverseProxy{ return &httputil.ReverseProxy{
Transport: transport,
Rewrite: func(pr *httputil.ProxyRequest) { Rewrite: func(pr *httputil.ProxyRequest) {
pr.SetURL(target) pr.SetURL(target)
pr.SetXForwarded() pr.SetXForwarded()
@@ -142,6 +334,94 @@ func (s *Server) writeMetrics(w http.ResponseWriter) {
s.cfg.Metrics.Render(w, string(state), n, pending) s.cfg.Metrics.Render(w, string(state), n, pending)
} }
// ANSI colors for per-request log lines.
const (
ansiReset = "\x1b[0m"
ansiCyan = "\x1b[36m"
ansiGreen = "\x1b[32m"
ansiYellow = "\x1b[33m"
ansiRed = "\x1b[31m"
)
// statusRecorder remembers the response status while passing everything
// through, including streaming flushes and websocket hijacks.
type statusRecorder struct {
http.ResponseWriter
status int
}
func (r *statusRecorder) WriteHeader(code int) {
r.status = code
r.ResponseWriter.WriteHeader(code)
}
func (r *statusRecorder) Flush() {
if f, ok := r.ResponseWriter.(http.Flusher); ok {
f.Flush()
}
}
func (r *statusRecorder) Hijack() (net.Conn, *bufio.ReadWriter, error) {
h, ok := r.ResponseWriter.(http.Hijacker)
if !ok {
return nil, nil, errors.New("response writer does not support hijacking")
}
return h.Hijack()
}
func (r *statusRecorder) Unwrap() http.ResponseWriter { return r.ResponseWriter }
// reqLine emits one request log line. With color enabled slog cannot be
// used: its text handler escapes the ANSI sequences, so the line is
// written to stderr directly in the same key=value shape. Without color
// it is a plain slog INFO line.
func (s *Server) reqLine(log *slog.Logger, code, line string, attrs ...any) {
if !s.cfg.LogColor {
log.Info(line, attrs...)
return
}
var sb strings.Builder
sb.WriteString("time=" + time.Now().Format("2006-01-02T15:04:05.000Z07:00") + " level=INFO ")
sb.WriteString(code + line + ansiReset)
for i := 0; i+1 < len(attrs); i += 2 {
fmt.Fprintf(&sb, " %v=%v", attrs[i], attrs[i+1])
}
sb.WriteByte('\n')
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
// response at INFO level, colored when enabled: cyan "-->" for incoming,
// green/yellow/red "<--" for responses by status class. At log levels
// above INFO it is a pass-through.
func (s *Server) logRequests(listener string, next http.Handler) http.Handler {
log := s.log.With("listener", listener)
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if !log.Enabled(r.Context(), slog.LevelInfo) {
next.ServeHTTP(w, r)
return
}
start := time.Now()
s.reqLine(log, ansiCyan, "--> "+r.Method+" "+r.URL.RequestURI(),
"listener", listener, "remote", r.RemoteAddr)
rec := &statusRecorder{ResponseWriter: w, status: http.StatusOK}
next.ServeHTTP(rec, r)
code := ansiGreen
switch {
case rec.status >= 500:
code = ansiRed
case rec.status >= 400:
code = ansiYellow
}
s.reqLine(log, code, "<-- "+strconv.Itoa(rec.status)+" "+r.Method+" "+r.URL.RequestURI(),
"listener", listener, "ms", time.Since(start).Milliseconds())
})
}
// llmPaths are the Ollama endpoints that load models into VRAM and therefore // llmPaths are the Ollama endpoints that load models into VRAM and therefore
// take the LLM lock. Everything else passes through unlocked. // take the LLM lock. Everything else passes through unlocked.
var llmPaths = map[string]bool{ var llmPaths = map[string]bool{
@@ -160,7 +440,7 @@ func isLLMRequest(r *http.Request) bool {
// OllamaHandler serves the Ollama-facing listener. // OllamaHandler serves the Ollama-facing listener.
func (s *Server) OllamaHandler() http.Handler { func (s *Server) OllamaHandler() http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { return s.logRequests("ollama", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path { switch r.URL.Path {
case "/healthz": case "/healthz":
s.writeHealthz(w) s.writeHealthz(w)
@@ -175,34 +455,53 @@ func (s *Server) OllamaHandler() http.Handler {
} }
start := time.Now() start := time.Now()
if s.busyMode == "reject" {
if !s.cfg.Lock.TryAcquireLLM() {
s.cfg.Metrics.ObserveLockWait("llm", time.Since(start).Seconds())
s.log.Info("llm request rejected; GPU busy",
"path", r.URL.Path, "status", s.busyStatus)
w.Header().Set("Retry-After", strconv.Itoa(s.busyRetryAfter))
http.Error(w, "GPU busy: image job active or queued", s.busyStatus)
return
}
s.cfg.Metrics.ObserveLockWait("llm", time.Since(start).Seconds())
defer s.cfg.Lock.ReleaseLLM()
s.ollamaProxy.ServeHTTP(w, r)
return
}
wctx, cancel := context.WithTimeout(r.Context(), s.cfg.LLMWaitTimeout) wctx, cancel := context.WithTimeout(r.Context(), s.cfg.LLMWaitTimeout)
err := s.cfg.Lock.AcquireLLM(wctx) err := s.cfg.Lock.AcquireLLM(wctx)
cancel() cancel()
s.cfg.Metrics.ObserveLockWait("llm", time.Since(start).Seconds()) s.cfg.Metrics.ObserveLockWait("llm", time.Since(start).Seconds())
if err != nil { if err != nil {
if errors.Is(err, context.DeadlineExceeded) && r.Context().Err() == nil { if errors.Is(err, context.DeadlineExceeded) && r.Context().Err() == nil {
w.Header().Set("Retry-After", strconv.Itoa(s.busyRetryAfter))
http.Error(w, "GPU busy: timed out waiting for the lock", http.StatusServiceUnavailable) http.Error(w, "GPU busy: timed out waiting for the lock", http.StatusServiceUnavailable)
} }
return return
} }
defer s.cfg.Lock.ReleaseLLM() defer s.cfg.Lock.ReleaseLLM()
s.ollamaProxy.ServeHTTP(w, r) s.ollamaProxy.ServeHTTP(w, r)
}) }))
} }
// ComfyHandler serves the ComfyUI-facing listener. // ComfyHandler serves the ComfyUI-facing listener.
func (s *Server) ComfyHandler() http.Handler { func (s *Server) ComfyHandler() http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { return s.logRequests("comfy", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/healthz" { switch r.URL.Path {
case "/healthz":
s.writeHealthz(w) s.writeHealthz(w)
return return
case "/metrics":
s.writeMetrics(w)
return
} }
if r.Method == http.MethodPost && r.URL.Path == "/prompt" { if r.Method == http.MethodPost && r.URL.Path == "/prompt" {
s.handlePrompt(w, r) s.handlePrompt(w, r)
return return
} }
s.comfyProxy.ServeHTTP(w, r) s.comfyProxy.ServeHTTP(w, r)
}) }))
} }
// captureWriter passes the response through unchanged while recording the // captureWriter passes the response through unchanged while recording the
@@ -250,6 +549,7 @@ func (s *Server) handlePrompt(w http.ResponseWriter, r *http.Request) {
s.cfg.Metrics.ObserveLockWait("image", time.Since(start).Seconds()) s.cfg.Metrics.ObserveLockWait("image", time.Since(start).Seconds())
log.Info("image lock acquired") log.Info("image lock acquired")
if s.cfg.Ollama != nil {
uctx, ucancel := context.WithTimeout(r.Context(), s.cfg.UnloadTimeout) uctx, ucancel := context.WithTimeout(r.Context(), s.cfg.UnloadTimeout)
elapsed, uerr := s.cfg.Ollama.UnloadAll(uctx) elapsed, uerr := s.cfg.Ollama.UnloadAll(uctx)
ucancel() ucancel()
@@ -264,6 +564,7 @@ func (s *Server) handlePrompt(w http.ResponseWriter, r *http.Request) {
default: default:
log.Info("ollama models unloaded", "seconds", elapsed.Seconds()) log.Info("ollama models unloaded", "seconds", elapsed.Seconds())
} }
}
cw := &captureWriter{ResponseWriter: w, status: http.StatusOK, limit: s.captureLimit} cw := &captureWriter{ResponseWriter: w, status: http.StatusOK, limit: s.captureLimit}
s.comfyProxy.ServeHTTP(cw, r) s.comfyProxy.ServeHTTP(cw, r)
@@ -308,7 +609,7 @@ func (s *Server) finishImageJob(promptID string) {
s.cfg.Lock.ReleaseImage() s.cfg.Lock.ReleaseImage()
log.Info("image lock released") log.Info("image lock released")
if s.cfg.WarmModel != "" { if s.cfg.Ollama != nil && s.cfg.WarmModel != "" {
if state, _, _ := s.cfg.Lock.Snapshot(); state == lock.StateIdle { if state, _, _ := s.cfg.Lock.Snapshot(); state == lock.StateIdle {
wctx, wcancel := context.WithTimeout(context.Background(), s.warmTimeout) wctx, wcancel := context.WithTimeout(context.Background(), s.warmTimeout)
if err := s.cfg.Ollama.Warm(wctx, s.cfg.WarmModel); err != nil { if err := s.cfg.Ollama.Warm(wctx, s.cfg.WarmModel); err != nil {
+185
View File
@@ -2,6 +2,7 @@ package proxy
import ( import (
"bufio" "bufio"
"context"
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
@@ -345,3 +346,187 @@ func TestPassThroughNoLock(t *testing.T) {
t.Fatalf("pass-through = %d %s", resp.StatusCode, body) t.Fatalf("pass-through = %d %s", resp.StatusCode, body)
} }
} }
func TestLLMBusyReject(t *testing.T) {
f := newFakes(t)
lk := lock.New(nil)
if err := lk.AcquireImage(context.Background()); err != nil {
t.Fatal(err)
}
ollamaClient, err := ollama.New(f.ollama.URL, nil)
if err != nil {
t.Fatal(err)
}
comfyClient, err := comfy.New(f.comfy.URL, nil)
if err != nil {
t.Fatal(err)
}
srv, err := New(Config{
OllamaURL: f.ollama.URL,
ComfyURL: f.comfy.URL,
Lock: lk,
Ollama: ollamaClient,
Comfy: comfyClient,
Metrics: metrics.New(),
LLMWaitTimeout: 2 * time.Second,
LLMBusyMode: "reject",
BusyRetryAfter: 17,
})
if err != nil {
t.Fatal(err)
}
front := httptest.NewServer(srv.OllamaHandler())
defer front.Close()
start := time.Now()
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{}`))
if err != nil {
t.Fatal(err)
}
body, _ := io.ReadAll(resp.Body)
resp.Body.Close()
if resp.StatusCode != http.StatusServiceUnavailable {
t.Fatalf("busy chat status = %d %s", resp.StatusCode, body)
}
if got := resp.Header.Get("Retry-After"); got != "17" {
t.Fatalf("Retry-After = %q", got)
}
if elapsed := time.Since(start); elapsed > time.Second {
t.Fatalf("reject was not immediate: %v", elapsed)
}
// After the image lock is released the next LLM request goes through.
lk.ReleaseImage()
resp, err = http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{}`))
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if resp.StatusCode != 200 {
t.Fatalf("chat after release status = %d", resp.StatusCode)
}
}
func TestLLMBusyWaitTimeoutRetryAfter(t *testing.T) {
f := newFakes(t)
lk := lock.New(nil)
if err := lk.AcquireImage(context.Background()); err != nil {
t.Fatal(err)
}
defer lk.ReleaseImage()
ollamaClient, err := ollama.New(f.ollama.URL, nil)
if err != nil {
t.Fatal(err)
}
comfyClient, err := comfy.New(f.comfy.URL, nil)
if err != nil {
t.Fatal(err)
}
srv, err := New(Config{
OllamaURL: f.ollama.URL,
ComfyURL: f.comfy.URL,
Lock: lk,
Ollama: ollamaClient,
Comfy: comfyClient,
Metrics: metrics.New(),
LLMWaitTimeout: 50 * time.Millisecond,
})
if err != nil {
t.Fatal(err)
}
front := httptest.NewServer(srv.OllamaHandler())
defer front.Close()
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{}`))
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if resp.StatusCode != http.StatusServiceUnavailable {
t.Fatalf("timed-out chat status = %d", resp.StatusCode)
}
if got := resp.Header.Get("Retry-After"); got != "30" {
t.Fatalf("Retry-After = %q", got)
}
}
func TestComfyOnlyModeSkipsOllama(t *testing.T) {
f := newFakes(t)
comfyClient, err := comfy.New(f.comfy.URL, nil)
if err != nil {
t.Fatal(err)
}
comfyClient.PollInterval = 5 * time.Millisecond
srv, err := New(Config{
ComfyURL: f.comfy.URL,
Lock: lock.New(nil),
Comfy: comfyClient,
Metrics: metrics.New(),
JobTimeout: 2 * time.Second,
})
if err != nil {
t.Fatal(err)
}
front := httptest.NewServer(srv.ComfyHandler())
defer front.Close()
resp, err := http.Post(front.URL+"/prompt", "application/json", strings.NewReader(`{}`))
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if resp.StatusCode != 200 {
t.Fatalf("prompt status = %d", resp.StatusCode)
}
f.completeJob()
select {
case <-f.freeCh:
case <-time.After(3 * time.Second):
t.Fatal("/free never called")
}
// With Ollama disabled the unload steps must not happen.
if i := f.rec.index("unload"); i >= 0 {
f.rec.mu.Lock()
t.Fatalf("unload called with ollama disabled; events: %v", f.rec.events)
}
}
func TestOllamaOnlyMode(t *testing.T) {
f := newFakes(t)
ollamaClient, err := ollama.New(f.ollama.URL, nil)
if err != nil {
t.Fatal(err)
}
srv, err := New(Config{
OllamaURL: f.ollama.URL,
Lock: lock.New(nil),
Ollama: ollamaClient,
Metrics: metrics.New(),
LLMWaitTimeout: time.Second,
})
if err != nil {
t.Fatal(err)
}
front := httptest.NewServer(srv.OllamaHandler())
defer front.Close()
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{}`))
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if resp.StatusCode != 200 {
t.Fatalf("chat status = %d", resp.StatusCode)
}
}
func TestNewRequiresConsumer(t *testing.T) {
_, err := New(Config{Lock: lock.New(nil), Metrics: metrics.New()})
if err == nil {
t.Fatal("New with no upstream URLs should fail")
}
}
+207
View File
@@ -0,0 +1,207 @@
package proxy
import (
"context"
"crypto/tls"
"io"
"net"
"net/http"
"strings"
"syscall"
"testing"
"time"
)
// stubTransport fails with ECONNREFUSED for the first fails requests, then
// returns a 200 response.
type stubTransport struct {
fails int
calls int
}
func (s *stubTransport) RoundTrip(req *http.Request) (*http.Response, error) {
s.calls++
if s.calls <= s.fails {
return nil, &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ECONNREFUSED}
}
return &http.Response{
StatusCode: 200,
Body: io.NopCloser(strings.NewReader("ok")),
Header: make(http.Header),
}, nil
}
func TestRetryTransportBackoff(t *testing.T) {
st := &stubTransport{fails: 3}
rt := &retryTransport{
base: st,
initial: 10 * time.Millisecond,
max: 25 * time.Millisecond,
}
start := time.Now()
req, _ := http.NewRequest(http.MethodGet, "http://upstream/api/version", nil)
resp, err := rt.RoundTrip(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if st.calls != 4 {
t.Fatalf("calls = %d, want 4", st.calls)
}
// Waits: 10ms + 20ms + 25ms (capped) = 55ms minimum.
elapsed := time.Since(start)
if elapsed < 50*time.Millisecond {
t.Fatalf("elapsed = %v, want >= ~55ms of backoff", elapsed)
}
if elapsed > 5*time.Second {
t.Fatalf("elapsed = %v, suspiciously long", elapsed)
}
}
func TestRetryTransportNonRefusedErrorNotRetried(t *testing.T) {
rt := &retryTransport{
base: &stubTransport{fails: 0},
initial: time.Millisecond,
max: time.Millisecond,
}
req, _ := http.NewRequest(http.MethodGet, "http://upstream/", nil)
resp, err := rt.RoundTrip(req)
if err != nil || resp.StatusCode != 200 {
t.Fatalf("resp=%v err=%v", resp, err)
}
}
type alwaysRefused struct{ calls int }
func (a *alwaysRefused) RoundTrip(*http.Request) (*http.Response, error) {
a.calls++
return nil, &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ECONNREFUSED}
}
func TestRetryTransportContextCancel(t *testing.T) {
ar := &alwaysRefused{}
rt := &retryTransport{base: ar, initial: time.Second, max: time.Second}
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
defer cancel()
req, _ := http.NewRequestWithContext(ctx, http.MethodGet, "http://upstream/", nil)
start := time.Now()
_, err := rt.RoundTrip(req)
if err == nil {
t.Fatal("expected error after context cancellation")
}
if time.Since(start) > 2*time.Second {
t.Fatal("retry loop did not stop on context cancellation")
}
}
func TestRetryViaReverseProxy(t *testing.T) {
// Point the proxy at a port nothing listens on; the request should be
// retried (not instantly 502) until the client context ends.
srv, err := New(Config{
OllamaURL: "http://127.0.0.1:1",
ComfyURL: "http://127.0.0.1:1",
Lock: nil,
Metrics: nil,
BackoffInitial: 10 * time.Millisecond,
BackoffMax: 20 * time.Millisecond,
})
if err != nil {
t.Fatal(err)
}
_ = srv // construction must not panic with minimal config
rt := &retryTransport{base: http.DefaultTransport, initial: 10 * time.Millisecond, max: 20 * time.Millisecond}
ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond)
defer cancel()
req, _ := http.NewRequestWithContext(ctx, http.MethodGet, "http://127.0.0.1:1/", nil)
_, err = rt.RoundTrip(req)
if err == nil || !shouldRetry(err) {
t.Fatalf("err = %v, want connection refused", err)
}
}
// flakyStatus returns 500 for the first fails requests, then 200.
type flakyStatus struct {
fails int
calls int
}
func (s *flakyStatus) RoundTrip(req *http.Request) (*http.Response, error) {
s.calls++
code := 200
if s.calls <= s.fails {
code = 500
}
return &http.Response{
StatusCode: code,
Status: http.StatusText(code),
Body: io.NopCloser(strings.NewReader("")),
Header: make(http.Header),
}, nil
}
func TestRetryTransport5xxGet(t *testing.T) {
fs := &flakyStatus{fails: 2}
rt := &retryTransport{base: fs, initial: time.Millisecond, max: 2 * time.Millisecond}
req, _ := http.NewRequest(http.MethodGet, "http://upstream/api/version", nil)
resp, err := rt.RoundTrip(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if resp.StatusCode != 200 {
t.Fatalf("status = %d, want 200", resp.StatusCode)
}
if fs.calls != 3 {
t.Fatalf("calls = %d, want 3", fs.calls)
}
}
func TestRetryTransport5xxPostNotRetried(t *testing.T) {
fs := &flakyStatus{fails: 10}
rt := &retryTransport{base: fs, initial: time.Millisecond, max: time.Millisecond}
// A streamed body (no GetBody) must not be replayed after a 500.
req, _ := http.NewRequest(http.MethodPost, "http://upstream/prompt", io.NopCloser(strings.NewReader("{}")))
resp, err := rt.RoundTrip(req)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if resp.StatusCode != 500 {
t.Fatalf("status = %d, want 500", resp.StatusCode)
}
if fs.calls != 1 {
t.Fatalf("calls = %d, want 1 (no retry for streamed POST)", fs.calls)
}
}
func TestShouldRetryClassification(t *testing.T) {
cases := []struct {
name string
err error
want bool
}{
{"dial refused", &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ECONNREFUSED}, true},
{"dial timeout", &net.OpError{Op: "dial", Net: "tcp", Err: timeoutErr{}}, true},
{"dns", &net.DNSError{Err: "no such host", IsNotFound: true}, true},
{"tls record", tls.RecordHeaderError{Msg: "bad"}, true},
{"tls alert", tls.AlertError(42), true},
{"read error mid-request", &net.OpError{Op: "read", Net: "tcp", Err: syscall.ECONNRESET}, false},
{"plain error", io.EOF, false},
}
for _, c := range cases {
if got := shouldRetry(c.err); got != c.want {
t.Errorf("%s: shouldRetry = %v, want %v", c.name, got, c.want)
}
}
}
type timeoutErr struct{}
func (timeoutErr) Error() string { return "i/o timeout" }
func (timeoutErr) Timeout() bool { return true }
func (timeoutErr) Temporary() bool { return true }
+47
View File
@@ -0,0 +1,47 @@
// Shared env-file handling for the installers (Windows and Linux).
package service
import (
"fmt"
"os"
"strings"
"gpu-turnstile/internal/config"
)
// syncEnvFile ensures the installed env file at path is a current
// installer-written sample. A missing file is created; a file without a
// CFG_VER line (hand-written or from before versioning) is invalid and
// replaced by a fresh sample after a .bak backup; a file written by an
// older version gets newly added settings appended via config.SyncSample.
// logPath activates the LOG_FILE line (Windows); empty leaves it commented.
func syncEnvFile(path, version, logPath string) error {
data, err := os.ReadFile(path)
switch {
case os.IsNotExist(err):
return writeEnvFile(path, config.SampleEnv(version, logPath))
case err != nil:
return fmt.Errorf("read %s: %w", path, err)
}
values, perr := config.ParseEnvFile(strings.NewReader(string(data)))
if perr == nil && values["CFG_VER"] != "" {
synced, changed := config.SyncSample(string(data), version, logPath)
if !changed {
return nil
}
return writeEnvFile(path, synced)
}
// No readable CFG_VER: the file is invalid — keep a backup and start
// from a fresh sample.
if err := os.WriteFile(path+".bak", data, 0o644); err != nil {
return fmt.Errorf("back up %s: %w", path, err)
}
return writeEnvFile(path, config.SampleEnv(version, logPath))
}
func writeEnvFile(path, content string) error {
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
return fmt.Errorf("write %s: %w", path, err)
}
return nil
}
+8
View File
@@ -0,0 +1,8 @@
package service
import "errors"
// ErrUserCancelled is returned by RelaunchElevated when the user declines
// the UAC prompt. Windows-only in practice; defined here so cross-platform
// callers can compare against it.
var ErrUserCancelled = errors.New("UAC prompt declined")
+46
View File
@@ -0,0 +1,46 @@
//go:build linux
package service
import (
"context"
"os"
"strconv"
"time"
"github.com/coreos/go-systemd/v22/daemon"
)
// NotifyReady tells systemd the service is up (Type=notify). It is a no-op
// when NOTIFY_SOCKET is unset, e.g. in a container or interactive shell.
func NotifyReady() {
daemon.SdNotify(false, daemon.SdNotifyReady)
}
// NotifyStopping tells systemd the service is shutting down.
func NotifyStopping() {
daemon.SdNotify(false, daemon.SdNotifyStopping)
}
// StartWatchdog pings the systemd watchdog every half of WATCHDOG_USEC
// until ctx is cancelled. It is a no-op unless systemd started the process
// with a watchdog configured (WatchdogSec= in the unit).
func StartWatchdog(ctx context.Context) {
usec, err := strconv.Atoi(os.Getenv("WATCHDOG_USEC"))
if err != nil || usec <= 0 {
return
}
interval := time.Duration(usec) * time.Microsecond / 2
go func() {
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
daemon.SdNotify(false, daemon.SdNotifyWatchdog)
}
}
}()
}
+14
View File
@@ -0,0 +1,14 @@
//go:build !linux
package service
import "context"
// NotifyReady is a no-op outside Linux (no systemd notify socket).
func NotifyReady() {}
// NotifyStopping is a no-op outside Linux.
func NotifyStopping() {}
// StartWatchdog is a no-op outside Linux.
func StartWatchdog(context.Context) {}
+259
View File
@@ -0,0 +1,259 @@
//go:build linux
// Package service integrates gpu-turnstile with systemd on Linux: running
// under a unit with readiness notification and watchdog, plus
// install/remove helpers that manage a hardened system unit.
package service
import (
"bytes"
"context"
"fmt"
"io"
"os"
"os/exec"
"os/signal"
"path/filepath"
"syscall"
)
// Name matches the Windows service name; the systemd unit is Name + ".service".
const Name = "gpu-turnstile"
// unitPath is where Install writes the unit file.
const unitPath = "/etc/systemd/system/" + Name + ".service"
// stateDir holds the installed binary (and staged updates); the unit's
// StateDirectory= directive makes systemd own it and grant the dynamic
// user write access. etcConfig is the config file the unit loads.
const (
stateDir = "/var/lib/" + Name
etcConfig = "/etc/" + Name + ".env"
)
// IsService reports whether the process was started by systemd.
func IsService() bool { return os.Getenv("INVOCATION_ID") != "" }
// Elevated is always true on Linux: there is no UAC equivalent; privilege
// errors surface from the failing operation with a "run as root" hint.
func Elevated() bool { return true }
// RelaunchElevated is a Windows-only concept (UAC).
func RelaunchElevated([]string) (int, error) {
return 0, fmt.Errorf("elevated relaunch is only supported on Windows")
}
// Run executes run with SIGINT/SIGTERM cancellation (which is how systemctl
// stop signals the process) and tells systemd when the shutdown begins.
func Run(run func(ctx context.Context) error) error {
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer stop()
defer NotifyStopping()
return run(ctx)
}
// renderUnit builds the hardened systemd unit: Type=notify so systemctl
// start blocks until the listeners are bound, a 30s watchdog, and
// restart-on-failure with a 5s delay — which is also what brings up a
// staged update after the updater exits with a non-zero code.
//
// Sandboxing mirrors the Windows virtual account: DynamicUser=yes gives
// the service a transient per-service UID with no login and no home, the
// filesystem is read-only except StateDirectory (the install dir, so
// self-updates can rewrite the binary), and the usual no-privilege-escalation
// directives apply. The proxy needs nothing but outbound TCP/UDP and the
// notify socket, so it loses nothing.
func renderUnit(exePath, configPath string) string {
return fmt.Sprintf(`[Unit]
Description=gpu-turnstile GPU arbitration proxy for Ollama and ComfyUI
After=network-online.target
Wants=network-online.target
[Service]
Type=notify
WatchdogSec=30s
ExecStart=%q -config %q
Restart=on-failure
RestartSec=5s
DynamicUser=yes
StateDirectory=%s
ProtectSystem=strict
ProtectHome=yes
PrivateTmp=yes
NoNewPrivileges=yes
ProtectKernelTunables=yes
ProtectKernelModules=yes
ProtectKernelLogs=yes
ProtectControlGroups=yes
ProtectClock=yes
RestrictNamespaces=yes
RestrictSUIDSGID=yes
RestrictRealtime=yes
LockPersonality=yes
MemoryDenyWriteExecute=yes
CapabilityBoundingSet=
AmbientCapabilities=
RestrictAddressFamilies=AF_UNIX AF_INET AF_INET6
SystemCallFilter=@system-service
SystemCallErrorNumber=EPERM
[Install]
WantedBy=multi-user.target
`, exePath, configPath, Name)
}
// copyFile copies src to dst, creating dst with the given mode.
func copyFile(src, dst string, mode os.FileMode) error {
in, err := os.Open(src)
if err != nil {
return err
}
defer in.Close()
out, err := os.OpenFile(dst, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, mode)
if err != nil {
return err
}
defer out.Close()
if _, err := io.Copy(out, in); err != nil {
return err
}
return out.Close()
}
// Install copies the current executable into /var/lib/gpu-turnstile, makes
// sure /etc/gpu-turnstile.env exists (copied from the given config file if
// provided), writes the hardened unit, then enables and starts it. With
// copyBin=false the current executable location and config path are
// registered as-is instead. Needs root.
//
// Re-running install converges an existing unit instead of failing: it is
// stopped first if active, the installed binary is replaced only when the
// content differs, the unit file is rewritten (followed by daemon-reload)
// only when it changed, and the service is started again only if it was
// active before.
//
// The binary lives in the StateDirectory rather than /usr/local/sbin on
// purpose: replacing a running binary needs write access to its
// *directory*, and granting the sandboxed service user write access to a
// shared system directory would let a compromised service overwrite other
// binaries. /var/lib/gpu-turnstile is exclusively ours.
func Install(configPath string, copyBin bool, version string) error {
exe, err := os.Executable()
if err != nil {
return err
}
if abs, absErr := filepath.Abs(exe); absErr == nil {
exe = abs
}
unit := Name + ".service"
_, statErr := os.Stat(unitPath)
fresh := os.IsNotExist(statErr)
wasRunning := exec.Command("systemctl", "is-active", "--quiet", unit).Run() == nil
if wasRunning {
if out, err := exec.Command("systemctl", "stop", unit).CombinedOutput(); err != nil {
return fmt.Errorf("systemctl stop (run as root): %w (%s)", err, out)
}
}
cfg := etcConfig
if copyBin {
if err := os.MkdirAll(stateDir, 0o755); err != nil {
return fmt.Errorf("create %s (run as root): %w", stateDir, err)
}
installedExe := filepath.Join(stateDir, Name)
if exe != installedExe {
if same, _ := sameFileContent(exe, installedExe); !same {
if err := copyFile(exe, installedExe, 0o755); err != nil {
return fmt.Errorf("install binary to %s: %w", installedExe, err)
}
}
}
exe = installedExe
if _, err := os.Stat(etcConfig); os.IsNotExist(err) && configPath != "" {
copyFile(configPath, etcConfig, 0o644) //nolint:errcheck // best effort
}
// Missing configs get a fully commented sample (LOG_FILE stays
// commented — stderr goes to the journal on Linux);
// installer-written ones from older versions get new settings
// appended; anything without CFG_VER is invalid and gets replaced
// (backup kept as .bak).
if err := syncEnvFile(etcConfig, version, ""); err != nil {
return err
}
} else if configPath != "" {
if abs, absErr := filepath.Abs(configPath); absErr == nil {
cfg = abs
}
}
rendered := renderUnit(exe, cfg)
if old, _ := os.ReadFile(unitPath); string(old) != rendered {
if err := os.WriteFile(unitPath, []byte(rendered), 0o644); err != nil {
return fmt.Errorf("write %s (run as root): %w", unitPath, err)
}
if out, err := exec.Command("systemctl", "daemon-reload").CombinedOutput(); err != nil {
return fmt.Errorf("systemctl daemon-reload: %w (%s)", err, out)
}
}
if fresh {
if out, err := exec.Command("systemctl", "enable", "--now", unit).CombinedOutput(); err != nil {
return fmt.Errorf("systemctl enable --now: %w (%s)", err, out)
}
return nil
}
if exec.Command("systemctl", "is-enabled", "--quiet", unit).Run() != nil {
if out, err := exec.Command("systemctl", "enable", unit).CombinedOutput(); err != nil {
return fmt.Errorf("systemctl enable: %w (%s)", err, out)
}
}
if wasRunning {
if out, err := exec.Command("systemctl", "start", unit).CombinedOutput(); err != nil {
return fmt.Errorf("systemctl start: %w (%s)", err, out)
}
}
return nil
}
// sameFileContent reports whether two files hold identical bytes. A missing
// destination is simply "different".
func sameFileContent(a, b string) (bool, error) {
ba, err := os.ReadFile(a)
if err != nil {
return false, err
}
bb, err := os.ReadFile(b)
if os.IsNotExist(err) {
return false, nil
}
if err != nil {
return false, err
}
return bytes.Equal(ba, bb), nil
}
// Remove stops and disables the service and deletes the unit file and the
// installed binary. The config file in /etc is left in place (user data).
func Remove() error {
exec.Command("systemctl", "disable", "--now", Name+".service").Run() // ignore: may not exist
if err := os.Remove(unitPath); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("remove %s: %w", unitPath, err)
}
os.RemoveAll(stateDir) // installed binary + staged updates; ignore error
if out, err := exec.Command("systemctl", "daemon-reload").CombinedOutput(); err != nil {
return fmt.Errorf("systemctl daemon-reload: %w (%s)", err, out)
}
return nil
}
// RestartIfRunning restarts the systemd unit when it is active (used after
// a forced update staged a new binary). Reports whether a restart
// happened. An inactive or missing unit is not an error. Needs root.
func RestartIfRunning() (bool, error) {
if err := exec.Command("systemctl", "is-active", "--quiet", Name+".service").Run(); err != nil {
return false, nil // inactive or not installed
}
if out, err := exec.Command("systemctl", "restart", Name+".service").CombinedOutput(); err != nil {
return false, fmt.Errorf("systemctl restart (run as root): %w (%s)", err, out)
}
return true, nil
}
+29
View File
@@ -0,0 +1,29 @@
//go:build linux
package service
import (
"strings"
"testing"
)
func TestRenderUnit(t *testing.T) {
unit := renderUnit("/var/lib/gpu-turnstile/gpu-turnstile", "/etc/gpu-turnstile.env")
for _, want := range []string{
"Type=notify",
"WatchdogSec=30s",
`ExecStart="/var/lib/gpu-turnstile/gpu-turnstile" -config "/etc/gpu-turnstile.env"`,
"Restart=on-failure",
"WantedBy=multi-user.target",
"DynamicUser=yes",
"StateDirectory=gpu-turnstile",
"ProtectSystem=strict",
"NoNewPrivileges=yes",
"RestrictAddressFamilies=AF_UNIX AF_INET AF_INET6",
"SystemCallFilter=@system-service",
} {
if !strings.Contains(unit, want) {
t.Fatalf("unit missing %q:\n%s", want, unit)
}
}
}
+45
View File
@@ -0,0 +1,45 @@
//go:build !windows && !linux
// Package service provides the stubs for platforms without service
// integration (Windows uses the SCM, Linux uses systemd). 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 and Linux (systemd)")
// 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, non-Linux platforms.
func Install(string, bool, string) error { return errUnsupported }
// Remove is unsupported on non-Windows, non-Linux platforms.
func Remove() error { return errUnsupported }
// Elevated is always true here: there is no UAC concept, and privilege
// errors surface from the failing operation with a "run as root" hint.
func Elevated() bool { return true }
// RelaunchElevated is unsupported on non-Windows, non-Linux platforms.
func RelaunchElevated([]string) (int, error) { return 0, errUnsupported }
// RestartIfRunning is a no-op on platforms without service integration.
func RestartIfRunning() (bool, error) { return false, nil }
+520
View File
@@ -0,0 +1,520 @@
//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. Installed services always run as the virtual
// account NT SERVICE\gpu-turnstile — a per-service low-privilege identity
// managed by the SCM, with no password and no admin rights.
package service
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"os"
"os/exec"
"path/filepath"
"strings"
"syscall"
"time"
"unsafe"
"golang.org/x/sys/windows"
"golang.org/x/sys/windows/svc"
"golang.org/x/sys/windows/svc/mgr"
"gpu-turnstile/internal/config"
)
// Name is the Windows service name.
const Name = "gpu-turnstile"
// virtualAccount is the per-service identity the service runs as. The SCM
// manages it: no password, automatic "log on as a service" right, gone
// when the service is removed.
const virtualAccount = `NT SERVICE\` + Name
// installDirs returns the canonical install (Program Files) and data
// (ProgramData) directories.
func installDirs() (install, data string) {
pf := os.Getenv("ProgramFiles")
if pf == "" {
pf = `C:\Program Files`
}
pd := os.Getenv("ProgramData")
if pd == "" {
pd = `C:\ProgramData`
}
return filepath.Join(pf, Name), filepath.Join(pd, Name)
}
// 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 running
// as the NT SERVICE\gpu-turnstile virtual account, whose binPath loads the
// given config file. With copyBin it first creates the canonical layout —
// the binary is copied into %ProgramFiles%\gpu-turnstile and the config
// next to it (an existing config there is kept), %ProgramData%\gpu-turnstile
// is created for logs — and registers that copy; with copyBin=false the
// current executable location is registered as-is. 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. After registering,
// the virtual account is granted modify access to the install and data
// directories (self-updates rewrite the exe), and read access to the
// config file if it lives elsewhere. The grants must come after
// CreateService: the virtual account's SID only exists once the service is
// registered.
//
// Re-running install on an existing service converges instead of failing:
// the service is stopped first if running (so the binary can be replaced),
// the installed binary is refreshed only when the content differs, the
// registration is updated only where it drifted, and the service is
// started again only if it was running before.
func Install(configPath string, copyBin bool, version string) error {
exe, err := os.Executable()
if err != nil {
return err
}
if abs, absErr := filepath.Abs(exe); absErr == nil {
exe = abs
}
if configPath != "" {
if abs, absErr := filepath.Abs(configPath); absErr == nil {
configPath = abs
}
}
m, err := mgr.Connect()
if err != nil {
return fmt.Errorf("connect to service manager (run as administrator): %w", err)
}
defer m.Disconnect()
// An existing service is converged, not an error. Stop it first so the
// binary copy can be replaced, and remember whether to start it again.
var s *mgr.Service
wasRunning := false
if existing, openErr := m.OpenService(Name); openErr == nil {
s = existing
defer s.Close()
if st, qErr := s.Query(); qErr == nil &&
(st.State == svc.Running || st.State == svc.StartPending) {
wasRunning = true
if err := stopAndWait(s); err != nil {
return err
}
}
}
installDir, dataDir := installDirs()
if copyBin {
targetCfg := filepath.Join(installDir, "gpu-turnstile.env")
if !strings.EqualFold(filepath.Dir(exe), installDir) {
if err := os.MkdirAll(installDir, 0o755); err != nil {
return fmt.Errorf("create %s: %w", installDir, err)
}
installedExe := filepath.Join(installDir, "gpu-turnstile.exe")
if same, _ := sameFileContent(exe, installedExe); !same {
if err := copyFile(exe, installedExe); err != nil {
return fmt.Errorf("copy binary to %s: %w", installedExe, err)
}
}
exe = installedExe
if configPath != "" && !strings.EqualFold(configPath, targetCfg) {
if _, statErr := os.Stat(targetCfg); os.IsNotExist(statErr) {
copyFile(configPath, targetCfg) //nolint:errcheck // best effort
}
}
}
// A service has no console: without LOG_FILE the output vanishes, so
// the sample pre-wires it to ProgramData. Missing configs get a
// fresh sample; installer-written ones from older versions get new
// settings appended; anything without CFG_VER is invalid and gets
// replaced (backup kept as .bak).
logPath := filepath.Join(dataDir, "gpu-turnstile.log")
if err := syncEnvFile(targetCfg, version, logPath); err != nil {
return err
}
configPath = targetCfg
}
binPath := fmt.Sprintf(`"%s" -config "%s"`, exe, configPath)
if s == nil {
s, err = m.CreateService(Name, binPath, mgr.Config{
StartType: mgr.StartAutomatic,
DisplayName: "gpu-turnstile",
Description: "GPU arbitration proxy for Ollama and ComfyUI",
ServiceStartName: virtualAccount,
})
if err != nil {
return fmt.Errorf("create service: %w", err)
}
defer s.Close()
if err := ensureRecovery(s); err != nil {
s.Delete() // roll back so a retry starts clean
return err
}
if err := grantAll(exe, configPath); err != nil {
s.Delete() // roll back so a retry starts clean
return err
}
// Best effort: start now instead of waiting for the next boot. A
// missing config (no consumer URLs) fails the start; the service stays
// registered and can be started once the config exists.
s.Start()
return nil
}
// Existing service: update the registration only where it drifted.
cur, err := s.Config()
if err != nil {
return fmt.Errorf("query service config: %w", err)
}
const displayName = "gpu-turnstile"
const description = "GPU arbitration proxy for Ollama and ComfyUI"
if cur.BinaryPathName != binPath ||
cur.StartType != mgr.StartAutomatic ||
cur.ServiceStartName != virtualAccount ||
cur.DisplayName != displayName ||
cur.Description != description {
upd := cur
upd.BinaryPathName = binPath
upd.StartType = mgr.StartAutomatic
upd.ServiceStartName = virtualAccount
upd.DisplayName = displayName
upd.Description = description
if err := s.UpdateConfig(upd); err != nil {
return fmt.Errorf("update service config: %w", err)
}
}
if err := ensureRecovery(s); err != nil {
return err
}
// Idempotent: re-assert the virtual account's ACLs (granting an
// existing ACE is a no-op).
if err := grantAll(exe, configPath); err != nil {
return err
}
if wasRunning {
if err := s.Start(); err != nil {
return fmt.Errorf("start service: %w", err)
}
}
return nil
}
// ensureRecovery makes sure the service restarts 5s after a failure — the
// mechanism that also brings up a staged update. The current actions are
// queried first so an up-to-date service is left untouched.
func ensureRecovery(s *mgr.Service) error {
want := mgr.RecoveryAction{Type: mgr.ServiceRestart, Delay: 5 * time.Second}
actions, err := s.RecoveryActions()
onFailure, flagErr := s.RecoveryActionsOnNonCrashFailures()
if err == nil && flagErr == nil && onFailure && len(actions) == 3 {
ok := true
for _, a := range actions {
if a.Type != want.Type || a.Delay != want.Delay {
ok = false
}
}
if ok {
return nil
}
}
if err := s.SetRecoveryActions([]mgr.RecoveryAction{want, want, want}, 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
}
// stopAndWait stops the service and waits up to 30s for the stopped state.
func stopAndWait(s *mgr.Service) error {
if _, err := s.Control(svc.Stop); err != nil {
return fmt.Errorf("stop service: %w", err)
}
deadline := time.Now().Add(30 * time.Second)
for {
st, err := s.Query()
if err != nil {
return fmt.Errorf("query service: %w", err)
}
if st.State == svc.Stopped {
return nil
}
if time.Now().After(deadline) {
return fmt.Errorf("service did not stop within 30s")
}
time.Sleep(300 * time.Millisecond)
}
}
// sameFileContent reports whether two files hold identical bytes. A missing
// destination is simply "different".
func sameFileContent(a, b string) (bool, error) {
ba, err := os.ReadFile(a)
if err != nil {
return false, err
}
bb, err := os.ReadFile(b)
if os.IsNotExist(err) {
return false, nil
}
if err != nil {
return false, err
}
return bytes.Equal(ba, bb), nil
}
// copyFile copies src to dst (0755 on the new file).
func copyFile(src, dst string) error {
in, err := os.Open(src)
if err != nil {
return err
}
defer in.Close()
out, err := os.OpenFile(dst, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o755)
if err != nil {
return err
}
if _, err := io.Copy(out, in); err != nil {
out.Close()
return err
}
return out.Close()
}
// grantAll gives the virtual account every ACL the service needs: modify
// on the install and ProgramData directories and the LOG_FILE directory
// (if configured elsewhere), read on a config file outside the install
// directory.
func grantAll(exe, configPath string) error {
_, dataDir := installDirs()
if err := os.MkdirAll(dataDir, 0o755); err != nil {
return fmt.Errorf("create %s: %w", dataDir, err)
}
if err := grantAccess(dataDir, "(OI)(CI)(M)"); err != nil {
return err
}
exeDir := filepath.Dir(exe)
if err := grantAccess(exeDir, "(OI)(CI)(M)"); err != nil {
return err
}
if configPath != "" && !strings.HasPrefix(strings.ToLower(configPath), strings.ToLower(exeDir)+`\`) {
if err := grantAccess(configPath, "(R)"); err != nil {
return err
}
}
if logFile := configuredLogFile(configPath); logFile != "" {
dir := filepath.Dir(logFile)
if err := os.MkdirAll(dir, 0o755); err == nil {
if err := grantAccess(dir, "(OI)(CI)(M)"); err != nil {
return err
}
}
}
return nil
}
// Remove stops (if running) and unregisters the service. The virtual
// account ceases to exist with it; the ACL grants on the install and log
// directories are left in place (harmless without the account).
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
}
var procShellExecuteExW = windows.NewLazySystemDLL("shell32.dll").NewProc("ShellExecuteExW")
const seeMaskNoCloseProcess = 0x40
// shellExecuteInfo mirrors SHELLEXECUTEINFOW (64-bit layout).
type shellExecuteInfo struct {
cbSize uint32
fMask uint32
hwnd uintptr
lpVerb *uint16
lpFile *uint16
lpParameters *uint16
lpDirectory *uint16
nShow int32
_ int32
hInstApp uintptr
lpIDList unsafe.Pointer
lpClass *uint16
hkeyClass uintptr
dwHotKey uint32
_ uint32
hIcon uintptr
hProcess windows.Handle
}
// Elevated reports whether the current process token is UAC-elevated.
func Elevated() bool {
var token windows.Token
if err := windows.OpenProcessToken(windows.CurrentProcess(), windows.TOKEN_QUERY, &token); err != nil {
return false
}
defer token.Close()
return token.IsElevated()
}
// RelaunchElevated re-runs the current executable elevated via the UAC
// "runas" verb with the given arguments, waits for the child, and returns
// its exit code. The child gets a fresh console window for its output.
func RelaunchElevated(args []string) (int, error) {
exe, err := os.Executable()
if err != nil {
return 0, err
}
quoted := make([]string, len(args))
for i, a := range args {
quoted[i] = syscall.EscapeArg(a)
}
cwd, _ := os.Getwd()
verb, _ := windows.UTF16PtrFromString("runas")
exeP, _ := windows.UTF16PtrFromString(exe)
params, _ := windows.UTF16PtrFromString(strings.Join(quoted, " "))
dir, _ := windows.UTF16PtrFromString(cwd)
info := shellExecuteInfo{
fMask: seeMaskNoCloseProcess,
lpVerb: verb,
lpFile: exeP,
lpParameters: params,
lpDirectory: dir,
nShow: windows.SW_NORMAL,
}
info.cbSize = uint32(unsafe.Sizeof(info))
r, _, callErr := procShellExecuteExW.Call(uintptr(unsafe.Pointer(&info)))
if r == 0 {
if errors.Is(callErr, syscall.Errno(1223)) { // ERROR_CANCELLED
return 0, ErrUserCancelled
}
return 0, fmt.Errorf("ShellExecuteEx: %w", callErr)
}
defer windows.CloseHandle(windows.Handle(info.hProcess))
windows.WaitForSingleObject(windows.Handle(info.hProcess), windows.INFINITE)
var code uint32
if err := windows.GetExitCodeProcess(windows.Handle(info.hProcess), &code); err != nil {
return 0, err
}
return int(code), nil
}
// grantAccess gives the virtual account the icacls permission set (e.g.
// "(OI)(CI)(M)") on path.
func grantAccess(path, perms string) error {
out, err := exec.Command("icacls", path, "/grant", virtualAccount+":"+perms).CombinedOutput()
if err != nil {
return fmt.Errorf("grant %s access to %s: %w (%s)", virtualAccount, path, err, strings.TrimSpace(string(out)))
}
return nil
}
// configuredLogFile reads LOG_FILE from the config file so the installer
// can pre-create and ACL the log directory. "" when unset or unreadable.
func configuredLogFile(configPath string) string {
f, err := os.Open(configPath)
if err != nil {
return ""
}
defer f.Close()
values, err := config.ParseEnvFile(f)
if err != nil {
return ""
}
return values["LOG_FILE"]
}
// RestartIfRunning restarts the service when it is installed and running
// (used after a forced update staged a new binary). Reports whether a
// restart happened. A service that is not installed or not running is not
// an error. Needs elevation.
func RestartIfRunning() (bool, error) {
m, err := mgr.Connect()
if err != nil {
return false, fmt.Errorf("connect to service manager (run as administrator): %w", err)
}
defer m.Disconnect()
s, err := m.OpenService(Name)
if err != nil {
return false, nil // not installed
}
defer s.Close()
st, err := s.Query()
if err != nil {
return false, fmt.Errorf("query service: %w", err)
}
if st.State != svc.Running && st.State != svc.StartPending {
return false, nil
}
if err := stopAndWait(s); err != nil {
return false, err
}
if err := s.Start(); err != nil {
return false, fmt.Errorf("start service: %w", err)
}
return true, nil
}
+16
View File
@@ -0,0 +1,16 @@
package update
// publicKeyPEM is the PEM-encoded Ed25519 public key that matches the
// RELEASE_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 RELEASE_SIGNING_KEY repository secret. When empty, the updater refuses
// to update (e.g. development builds).
var publicKeyPEM = `-----BEGIN PUBLIC KEY-----
MCowBQYDK2VwAyEAqTAJ0CCeAQI7MhFlgc5xNmF/CfvLVUAY3ZAoeAS0tT8=
-----END PUBLIC KEY-----
`
+279
View File
@@ -0,0 +1,279 @@
// 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 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 binary version, e.g. v0.1.2 ("dev" cannot be compared)
// Desired is the APP_VER policy: "dev" disables updates, "stable" (or
// empty) tracks the latest release, anything else is an exact vX.Y.Z
// release tag to pin — staging it even when that means a downgrade.
Desired string
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
// 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, APP_VER=dev, or no embedded public key);
// 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()
desired := u.Desired
if desired == "" {
desired = "stable"
}
if desired == "dev" {
log.Debug("auto-update: APP_VER=dev, 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
}
pinned := desired != "stable"
endpoint := api + "/releases/latest"
if pinned {
endpoint = api + "/releases/tags/" + desired
}
body, err := u.get(ctx, endpoint)
if err != nil {
return false, fmt.Errorf("fetch release: %w", err)
}
var rel release
if err := json.Unmarshal(body, &rel); err != nil {
return false, fmt.Errorf("parse release: %w", err)
}
if pinned {
// Pin mode: any difference from the target tag means stage it —
// including downgrades and replacing a dev binary.
if u.Version == rel.TagName {
log.Debug("auto-update: already on pinned version", "version", u.Version)
return false, nil
}
} else if u.Version != "" && u.Version != "dev" {
// Stable mode: only strictly newer releases count; a dev binary
// cannot be compared and is always replaced by the latest release.
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[:]
}
+243
View File
@@ -0,0 +1,243 @@
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()
serveRelease := 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("/api/v1/repos/o/r/releases/latest", serveRelease)
mux.HandleFunc("/api/v1/repos/o/r/releases/tags/"+f.tag, serveRelease)
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 (f *fakeGitea) updaterDesired(version, desired string) *Updater {
u := f.updater(version)
u.Desired = desired
return u
}
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 TestCheckDevBuildGetsStable(t *testing.T) {
// A dev binary cannot be compared; APP_VER=stable replaces it with the
// latest release.
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
withPublicKey(t, f.pubPEM)
exe := fakeExe(t)
staged, err := f.updater("dev").Check(context.Background(), exe)
if err != nil {
t.Fatal(err)
}
if !staged {
t.Fatal("expected dev binary to be replaced by the latest release")
}
content, _ := os.ReadFile(exe)
if string(content) != "new-binary" {
t.Fatalf("exe content = %q", content)
}
}
func TestCheckDesiredDevDisables(t *testing.T) {
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
withPublicKey(t, f.pubPEM)
for _, version := range []string{"dev", "v0.1.2"} {
staged, err := f.updaterDesired(version, "dev").Check(context.Background(), fakeExe(t))
if err != nil || staged {
t.Fatalf("version %s: staged=%v err=%v, want no action with APP_VER=dev", version, staged, err)
}
}
}
func TestCheckPinned(t *testing.T) {
// Pin mode stages the exact tag — up or down — and skips when the
// binary already matches.
for _, version := range []string{"v0.1.2", "v9.9.9", "dev"} {
f := newFakeGitea(t, "v0.5.0", []byte("pinned-binary"))
withPublicKey(t, f.pubPEM)
exe := fakeExe(t)
staged, err := f.updaterDesired(version, "v0.5.0").Check(context.Background(), exe)
if err != nil {
t.Fatalf("version %s: %v", version, err)
}
if !staged {
t.Fatalf("version %s: expected pinned v0.5.0 to be staged", version)
}
content, _ := os.ReadFile(exe)
if string(content) != "pinned-binary" {
t.Fatalf("version %s: exe content = %q", version, content)
}
}
f := newFakeGitea(t, "v0.5.0", []byte("pinned-binary"))
withPublicKey(t, f.pubPEM)
staged, err := f.updaterDesired("v0.5.0", "v0.5.0").Check(context.Background(), fakeExe(t))
if err != nil || staged {
t.Fatalf("staged=%v err=%v, want no action when already on the pinned version", 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")
}
}