5 Commits
Author SHA1 Message Date
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
19 changed files with 1615 additions and 198 deletions
+47
View File
@@ -60,3 +60,50 @@ 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
run: |
printf '%s\n' "${{ secrets.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/
+61 -9
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
@@ -28,8 +31,11 @@ Open WebUI / n8n ────► :8188 ───┘
## 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 |
|---|---|---| |---|---|---|
@@ -39,10 +45,14 @@ startup.
| `COMFY_URL` | `http://127.0.0.1:8189` | ComfyUI upstream | | `COMFY_URL` | `http://127.0.0.1:8189` | ComfyUI upstream |
| `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) |
| `LOGLEVEL` | `warn` | `info` logs every request (colored arrows in text mode), `debug` adds lock transitions. `LOG_LEVEL` works as an alias | | `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 |
@@ -52,6 +62,10 @@ startup.
| `BACKOFF_INITIAL` | `1s` | First retry wait when an upstream refuses a connection | | `BACKOFF_INITIAL` | `1s` | First retry wait when an upstream refuses a connection |
| `BACKOFF_MAX` | `60s` | Cap for the exponential retry backoff | | `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 |
## Observability ## Observability
@@ -74,6 +88,39 @@ go build ./cmd/gpu-turnstile
./gpu-turnstile ./gpu-turnstile
``` ```
### Run natively on Windows (primary deployment)
Download `gpu-turnstile.exe` from a release, put a `gpu-turnstile.env`
next to it, and run it — or install it as a Windows service from an
elevated shell:
```sh
gpu-turnstile.exe service install # auto-start service, recovery = restart
gpu-turnstile.exe service remove
```
The service uses the config file (services have no convenient
environment); set `LOG_FILE` in it since there is no console.
Suggested layout: `C:\Program Files\gpu-turnstile\` for the exe and
`gpu-turnstile.env`, logs under `C:\ProgramData\gpu-turnstile\` via
`LOG_FILE`. The service runs as `LocalSystem` by default, which can write
the install directory for self-updates. For least privilege, run it as the
virtual account `NT SERVICE\gpu-turnstile` and grant write access to just
those two directories.
**Auto-update is on by default**: the binary checks the repo's latest
release on startup and every `UPDATE_INTERVAL`, verifies the Ed25519
signature of the download against the public key embedded at build time,
and — once the GPU lock is idle — restarts the service onto the new
version. Disable with `AUTO_UPDATE=false`. Releases are signed by CI with
OpenSSL; the matching public key lives in `internal/update/pubkey.go`
(one-time setup: `openssl genpkey -algorithm ed25519 -out private.pem`,
`openssl pkey -in private.pem -pubout -out public.pem`; private key goes
to the `RELEASE_SIGNING_KEY` repo secret, public key is committed).
### 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 \
@@ -84,9 +131,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.
@@ -99,13 +147,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
``` ```
+73 -6
View File
@@ -51,6 +51,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.
@@ -102,6 +107,12 @@ 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 |
@@ -110,9 +121,14 @@ load time. Off by default.
| `COMFY_URL` | `http://127.0.0.1:8189` | upstream | | `COMFY_URL` | `http://127.0.0.1:8189` | upstream |
| `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 |
| `LOGLEVEL` | `warn` | `info` logs every request (colored arrows in text mode), `debug` adds lock transitions. `LOG_LEVEL` is accepted as an alias | | `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 |
@@ -122,10 +138,48 @@ load time. Off by default.
| `BACKOFF_INITIAL` | `1s` | first retry wait when an upstream refuses a connection | | `BACKOFF_INITIAL` | `1s` | first retry wait when an upstream refuses a connection |
| `BACKOFF_MAX` | `60s` | cap for the exponential retry backoff | | `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 |
Startup fails fast on unparsable values. Both upstreams are probed once at Startup fails fast on unparsable values. Both upstreams are probed once at
start (`/api/version`, `/system_stats`); failure is logged, not fatal. start (`/api/version`, `/system_stats`); failure is logged, not fatal.
## Native Windows deployment
The binary runs natively on Windows (the primary deployment) as well as in
Docker.
- `gpu-turnstile.exe service install [-config path]` registers an
auto-start Windows service (needs an elevated shell). Recovery actions
restart it 5 s after any failure. `service remove` uninstalls.
- **Layout**: install to `C:\Program Files\gpu-turnstile\` (exe plus
`gpu-turnstile.env`); logs belong in `C:\ProgramData\gpu-turnstile\` via
`LOG_FILE`. The service must be able to write its install directory for
self-updates — Program Files is writable by LocalSystem and admins, which
is why running as the default `LocalSystem` account is the simple choice.
- **Account**: the default `LocalSystem` works out of the box. For least
privilege, create the service with the virtual account
`NT SERVICE\gpu-turnstile` and grant it write access to the install and
log directories only (no network logon, no user profile).
- Use a config file (above) for the service — Windows services have no
convenient environment. Logs go to `LOG_FILE` since there is no console.
- **Auto-update**: on startup and every `UPDATE_INTERVAL`, the binary
checks `UPDATE_REPO`'s latest release; if its tag is a newer `vX.Y.Z`,
it downloads `UPDATE_ASSET` plus its `.sig` (and `.sha256` when present)
and verifies an Ed25519 signature against the public key embedded in
`internal/update/pubkey.go`. A verified binary is swapped in next to the
running exe (rename-aside, allowed on Windows), and once the GPU lock is
idle the process exits with code 3 so the service recovery restarts it
on the new version. Interactive runs only log "restart to apply".
`dev` builds and builds without an embedded public key never update.
- **Signing setup (one time)**: `openssl genpkey -algorithm ed25519 -out
private.pem`; `openssl pkey -in private.pem -pubout -out public.pem`.
Private key → repo secret `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
@@ -171,11 +225,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 service integration
Dockerfile Dockerfile
.gitea/workflows/ci.yml .gitea/workflows/ci.yml
README.md README.md
@@ -197,11 +255,16 @@ 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 and dev builds skipped.
## Build and CI ## Build and CI
- Go 1.23+, stdlib only. `CGO_ENABLED=0`, `-ldflags="-s -w"`, version from - Go 1.23+, `golang.org/x/sys` is the only external dependency (Windows
`git describe` injected via `-X main.version=`. service integration; not used in the Linux build). `CGO_ENABLED=0`,
`-ldflags="-s -w"`, version from `git describe` injected via
`-X main.version=`.
- Dockerfile: multi-stage, final image `gcr.io/distroless/static` (or - 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"]`.
@@ -216,8 +279,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)
+266 -181
View File
@@ -6,256 +6,341 @@ import (
"context" "context"
"errors" "errors"
"fmt" "fmt"
"io"
"log/slog" "log/slog"
"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
historyPollInterval time.Duration
probeTimeout time.Duration
freeTimeout time.Duration
warmTimeout time.Duration
shutdownTimeout time.Duration
backoffInitial time.Duration
backoffMax time.Duration
promptCaptureLimit int64
warmModel string
logLevel slog.Level
logJSON bool
}
func envDuration(getenv func(string) string, name string, dst *time.Duration) error {
v := getenv(name)
if v == "" {
return nil
}
d, err := time.ParseDuration(v)
if err != nil {
return fmt.Errorf("%s: %w", name, err)
}
*dst = d
return nil
}
func loadConfig(getenv func(string) string) (config, error) {
cfg := config{
listenOllama: ":11434",
listenComfy: ":8188",
ollamaURL: "http://127.0.0.1:11435",
comfyURL: "http://127.0.0.1:8189",
unloadTimeout: time.Minute,
jobTimeout: 15 * time.Minute,
llmWaitTimeout: 10 * time.Minute,
unloadPollInterval: 500 * time.Millisecond,
historyPollInterval: time.Second,
probeTimeout: 5 * time.Second,
freeTimeout: 30 * time.Second,
warmTimeout: 2 * time.Minute,
shutdownTimeout: 10 * time.Second,
backoffInitial: time.Second,
backoffMax: time.Minute,
promptCaptureLimit: 64 * 1024,
logLevel: slog.LevelWarn,
}
for _, e := range []struct {
name string
dst *string
}{
{"LISTEN_OLLAMA", &cfg.listenOllama},
{"LISTEN_COMFY", &cfg.listenComfy},
{"OLLAMA_URL", &cfg.ollamaURL},
{"COMFY_URL", &cfg.comfyURL},
{"WARM_MODEL", &cfg.warmModel},
} {
if v := getenv(e.name); v != "" {
*e.dst = v
}
}
for _, e := range []struct {
name string
dst *time.Duration
}{
{"UNLOAD_TIMEOUT", &cfg.unloadTimeout},
{"JOB_TIMEOUT", &cfg.jobTimeout},
{"LLM_WAIT_TIMEOUT", &cfg.llmWaitTimeout},
{"UNLOAD_POLL_INTERVAL", &cfg.unloadPollInterval},
{"HISTORY_POLL_INTERVAL", &cfg.historyPollInterval},
{"PROBE_TIMEOUT", &cfg.probeTimeout},
{"FREE_TIMEOUT", &cfg.freeTimeout},
{"WARM_TIMEOUT", &cfg.warmTimeout},
{"SHUTDOWN_TIMEOUT", &cfg.shutdownTimeout},
{"BACKOFF_INITIAL", &cfg.backoffInitial},
{"BACKOFF_MAX", &cfg.backoffMax},
} {
if err := envDuration(getenv, e.name, e.dst); err != nil {
return cfg, err
}
}
if v := getenv("PROMPT_CAPTURE_LIMIT"); v != "" {
n, err := strconv.ParseInt(v, 10, 64)
if err != nil || n < 0 {
return cfg, fmt.Errorf("PROMPT_CAPTURE_LIMIT: must be a non-negative integer (bytes)")
}
cfg.promptCaptureLimit = n
}
// LOGLEVEL is the canonical spelling; LOG_LEVEL is kept as an alias.
logLevelValue := getenv("LOGLEVEL")
if logLevelValue == "" {
logLevelValue = getenv("LOG_LEVEL")
}
if logLevelValue != "" {
var level slog.Level
if err := level.UnmarshalText([]byte(logLevelValue)); err != nil {
return cfg, fmt.Errorf("LOGLEVEL: %w", err)
}
cfg.logLevel = level
}
switch strings.ToLower(getenv("LOG_FORMAT")) {
case "", "text":
case "json":
cfg.logJSON = true
default:
return cfg, fmt.Errorf("LOG_FORMAT: must be \"text\" or \"json\"")
}
return cfg, nil
}
func main() { func main() {
cfg, err := loadConfig(os.Getenv) configPath, args := splitConfigFlag(os.Args[1:])
if len(args) > 0 && args[0] == "service" {
os.Exit(serviceCommand(configPath, args[1:]))
}
if len(args) > 0 {
fmt.Fprintf(os.Stderr, "usage: gpu-turnstile [-config path] | gpu-turnstile service install|remove [-config path]\n")
os.Exit(2)
}
if exePath, err := os.Executable(); err == nil {
update.CleanupOld(exePath)
}
cfg, err := loadMergedConfig(configPath)
if err != nil { if err != nil {
fmt.Fprintf(os.Stderr, "gpu-turnstile: %v\n", err) fmt.Fprintf(os.Stderr, "gpu-turnstile: %v\n", err)
os.Exit(1) os.Exit(1)
} }
log, logOut, logCloser := newLogger(cfg)
defer logCloser.Close()
opts := &slog.HandlerOptions{Level: cfg.logLevel} if service.IsService() {
var handler slog.Handler = slog.NewTextHandler(os.Stderr, opts) if err := service.Run(func(ctx context.Context) error { return run(ctx, cfg, log, logOut, true) }); err != nil {
if cfg.logJSON { log.Error("service failed", "err", err)
handler = slog.NewJSONHandler(os.Stderr, opts) os.Exit(1)
}
return
}
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer stop()
if err := run(ctx, cfg, log, logOut, false); err != nil {
log.Error("listener failed", "err", err)
os.Exit(1)
}
}
// splitConfigFlag extracts -config <path> (or -config=<path>) from args.
func splitConfigFlag(args []string) (string, []string) {
var configPath string
rest := args[:0]
for i := 0; i < len(args); i++ {
switch {
case args[i] == "-config" && i+1 < len(args):
configPath = args[i+1]
i++
case strings.HasPrefix(args[i], "-config="):
configPath = strings.TrimPrefix(args[i], "-config=")
default:
rest = append(rest, args[i])
}
}
return configPath, rest
}
// defaultConfigPath returns gpu-turnstile.env next to the executable.
func defaultConfigPath() string {
exe, err := os.Executable()
if err != nil {
return "gpu-turnstile.env"
}
return filepath.Join(filepath.Dir(exe), "gpu-turnstile.env")
}
// resolveConfigPath applies the precedence: -config flag, then
// GPU_TURNSTILE_CONFIG, then the default next to the executable.
func resolveConfigPath(flagValue string) string {
if flagValue != "" {
return flagValue
}
if v := os.Getenv("GPU_TURNSTILE_CONFIG"); v != "" {
return v
}
return defaultConfigPath()
}
// loadMergedConfig reads the config file (if present) and overlays process
// environment variables on top. A missing file is fine; an unreadable or
// malformed file is fatal.
func loadMergedConfig(flagValue string) (config.Config, error) {
path := resolveConfigPath(flagValue)
values := map[string]string{}
if f, err := os.Open(path); err == nil {
defer f.Close()
parsed, err := config.ParseEnvFile(f)
if err != nil {
return config.Config{}, fmt.Errorf("%s: %w", path, err)
}
values = parsed
} else if !errors.Is(err, os.ErrNotExist) {
return config.Config{}, fmt.Errorf("read config file: %w", err)
}
getenv := func(key string) string {
if v := os.Getenv(key); v != "" {
return v
}
return values[key]
}
return config.Load(getenv)
}
// newLogger builds the slog logger and returns the output writer (stderr or
// the opened LOG_FILE) plus a closer for it.
func newLogger(cfg config.Config) (*slog.Logger, io.Writer, io.Closer) {
out := io.Writer(os.Stderr)
closer := io.NopCloser(nil)
if cfg.LogFile != "" {
if f, err := os.OpenFile(cfg.LogFile, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o644); err == nil {
out = f
closer = f
} else {
fmt.Fprintf(os.Stderr, "gpu-turnstile: cannot open LOG_FILE %s: %v (logging to stderr)\n", cfg.LogFile, err)
}
}
opts := &slog.HandlerOptions{Level: cfg.LogLevel}
var handler slog.Handler = slog.NewTextHandler(out, opts)
if cfg.LogJSON {
handler = slog.NewJSONHandler(out, opts)
} }
log := slog.New(handler) log := slog.New(handler)
slog.SetDefault(log) slog.SetDefault(log)
return log, out, closer
}
func serviceCommand(configPath string, args []string) int {
if len(args) != 1 || (args[0] != "install" && args[0] != "remove") {
fmt.Fprintf(os.Stderr, "usage: gpu-turnstile service install|remove [-config path]\n")
return 2
}
var err error
if args[0] == "install" {
path := resolveConfigPath(configPath)
if abs, absErr := filepath.Abs(path); absErr == nil {
path = abs
}
err = service.Install(path)
} else {
err = service.Remove()
}
if err != nil {
fmt.Fprintf(os.Stderr, "gpu-turnstile service %s: %v\n", args[0], err)
return 1
}
fmt.Printf("service %s: %sd\n", service.Name, args[0])
return 0
}
func run(ctx context.Context, cfg config.Config, log *slog.Logger, logOut io.Writer, isService bool) error {
// The startup line carries the version and every setting and is emitted // 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. // at WARN so it is visible even with the default (quiet) log level.
log.Log(context.Background(), slog.LevelWarn, "starting gpu-turnstile", log.Log(ctx, slog.LevelWarn, "starting gpu-turnstile",
"version", version, "version", version,
"listen_ollama", cfg.listenOllama, "listen_ollama", cfg.ListenOllama,
"listen_comfy", cfg.listenComfy, "listen_comfy", cfg.ListenComfy,
"ollama_url", cfg.ollamaURL, "ollama_url", cfg.OllamaURL,
"comfy_url", cfg.comfyURL, "comfy_url", cfg.ComfyURL,
"unload_timeout", cfg.unloadTimeout, "unload_timeout", cfg.UnloadTimeout,
"job_timeout", cfg.jobTimeout, "job_timeout", cfg.JobTimeout,
"llm_wait_timeout", cfg.llmWaitTimeout, "llm_wait_timeout", cfg.LLMWaitTimeout,
"unload_poll_interval", cfg.unloadPollInterval, "llm_busy_mode", cfg.LLMBusyMode,
"history_poll_interval", cfg.historyPollInterval, "llm_busy_status", cfg.LLMBusyStatus,
"probe_timeout", cfg.probeTimeout, "busy_retry_after", cfg.BusyRetryAfter,
"free_timeout", cfg.freeTimeout, "unload_poll_interval", cfg.UnloadPollInterval,
"warm_timeout", cfg.warmTimeout, "history_poll_interval", cfg.HistoryPollInterval,
"shutdown_timeout", cfg.shutdownTimeout, "probe_timeout", cfg.ProbeTimeout,
"backoff_initial", cfg.backoffInitial, "free_timeout", cfg.FreeTimeout,
"backoff_max", cfg.backoffMax, "warm_timeout", cfg.WarmTimeout,
"prompt_capture_limit", cfg.promptCaptureLimit, "shutdown_timeout", cfg.ShutdownTimeout,
"warm_model", cfg.warmModel, "backoff_initial", cfg.BackoffInitial,
"log_level", cfg.logLevel, "backoff_max", cfg.BackoffMax,
"log_format", map[bool]string{true: "json", false: "text"}[cfg.logJSON], "prompt_capture_limit", cfg.PromptCaptureLimit,
"warm_model", cfg.WarmModel,
"auto_update", cfg.AutoUpdate,
"update_interval", cfg.UpdateInterval,
"update_repo", cfg.UpdateRepo,
"update_asset", cfg.UpdateAsset,
"log_level", cfg.LogLevel,
"log_format", map[bool]string{true: "json", false: "text"}[cfg.LogJSON],
"log_file", cfg.LogFile,
) )
ollamaClient, err := ollama.New(cfg.ollamaURL, log) lk := lock.New(log)
ollamaClient, err := ollama.New(cfg.OllamaURL, log)
if err != nil { if err != nil {
log.Error("invalid configuration", "err", err) return err
os.Exit(1)
} }
comfyClient, err := comfy.New(cfg.comfyURL, log) comfyClient, err := comfy.New(cfg.ComfyURL, log)
if err != nil { if err != nil {
log.Error("invalid configuration", "err", err) return 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,
LogColor: !cfg.logJSON && os.Getenv("NO_COLOR") == "", LogColor: !cfg.LogJSON && cfg.LogFile == "" && os.Getenv("NO_COLOR") == "",
LLMWaitTimeout: cfg.llmWaitTimeout, LogWriter: logOut,
UnloadTimeout: cfg.unloadTimeout, LLMWaitTimeout: cfg.LLMWaitTimeout,
JobTimeout: cfg.jobTimeout, UnloadTimeout: cfg.UnloadTimeout,
UnloadPollInterval: cfg.unloadPollInterval, JobTimeout: cfg.JobTimeout,
HistoryPollInterval: cfg.historyPollInterval, UnloadPollInterval: cfg.UnloadPollInterval,
FreeTimeout: cfg.freeTimeout, HistoryPollInterval: cfg.HistoryPollInterval,
WarmTimeout: cfg.warmTimeout, FreeTimeout: cfg.FreeTimeout,
BackoffInitial: cfg.backoffInitial, WarmTimeout: cfg.WarmTimeout,
BackoffMax: cfg.backoffMax, LLMBusyMode: cfg.LLMBusyMode,
PromptCaptureLimit: cfg.promptCaptureLimit, LLMBusyStatus: cfg.LLMBusyStatus,
WarmModel: cfg.warmModel, 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 both upstreams once; failure is logged, not fatal.
probeCtx, probeCancel := context.WithTimeout(context.Background(), cfg.probeTimeout) probeCtx, probeCancel := context.WithTimeout(ctx, cfg.ProbeTimeout)
if err := ollamaClient.Probe(probeCtx); err != nil { if err := ollamaClient.Probe(probeCtx); err != nil {
log.Warn("ollama probe failed", "url", cfg.ollamaURL, "err", err) log.Warn("ollama probe failed", "url", cfg.OllamaURL, "err", err)
} }
if err := comfyClient.Probe(probeCtx); err != nil { 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()} ollamaSrv := &http.Server{Addr: cfg.ListenOllama, Handler: srv.OllamaHandler()}
comfySrv := &http.Server{Addr: cfg.listenComfy, Handler: srv.ComfyHandler()} comfySrv := &http.Server{Addr: cfg.ListenComfy, Handler: srv.ComfyHandler()}
errCh := make(chan error, 2) errCh := make(chan error, 2)
go func() { errCh <- ollamaSrv.ListenAndServe() }() go func() { errCh <- ollamaSrv.ListenAndServe() }()
go func() { errCh <- comfySrv.ListenAndServe() }() go func() { errCh <- comfySrv.ListenAndServe() }()
sigCtx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) if cfg.AutoUpdate {
defer stop() 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) ollamaSrv.Shutdown(shutdownCtx)
comfySrv.Shutdown(shutdownCtx) comfySrv.Shutdown(shutdownCtx)
return nil
}
// updateLoop checks for signed updates on startup and every UPDATE_INTERVAL.
// In service mode a staged update is applied by exiting with exitCodeUpdate
// once the GPU lock is idle; the service recovery configuration restarts the
// process with the new binary. Interactively it only logs.
func updateLoop(ctx context.Context, cfg config.Config, log *slog.Logger, lk *lock.Lock, isService bool) {
exePath, err := os.Executable()
if err != nil {
log.Warn("auto-update disabled: cannot locate executable", "err", err)
return
}
u := &update.Updater{Repo: cfg.UpdateRepo, Asset: cfg.UpdateAsset, Version: version, Log: log}
for {
staged, err := u.Check(ctx, exePath)
if err != nil && ctx.Err() == nil {
log.Warn("auto-update check failed", "err", err)
}
if staged {
if !isService {
log.Warn("auto-update: new binary staged; restart gpu-turnstile to apply")
return
}
log.Warn("auto-update: staged; restarting once the GPU is idle")
if waitForIdle(ctx, lk, 24*time.Hour) {
log.Warn("auto-update: restarting to apply update")
os.Exit(exitCodeUpdate)
}
return
}
select {
case <-ctx.Done():
return
case <-time.After(cfg.UpdateInterval):
}
}
}
// waitForIdle polls the lock until no LLM or image work is active or
// pending, max at most. Returns false on timeout or cancellation.
func waitForIdle(ctx context.Context, lk *lock.Lock, max time.Duration) bool {
deadline := time.Now().Add(max)
for {
if state, _, pending := lk.Snapshot(); state == lock.StateIdle && !pending {
return true
}
if time.Now().After(deadline) {
return false
}
select {
case <-ctx.Done():
return false
case <-time.After(5 * time.Second):
}
}
} }
+4 -1
View File
@@ -6,7 +6,7 @@
# ComfyUI --listen 0.0.0.0 --port 8189). # ComfyUI --listen 0.0.0.0 --port 8189).
services: services:
gpu-turnstile: gpu-turnstile:
image: git.rambossek.at/public/gpu-turnstile:v0.1.2 image: git.rambossek.at/public/gpu-turnstile:v0.1.3
restart: unless-stopped restart: unless-stopped
environment: environment:
# Services on the Docker host itself: # Services on the Docker host itself:
@@ -18,6 +18,9 @@ services:
# 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
+2
View File
@@ -1,3 +1,5 @@
module gpu-turnstile module gpu-turnstile
go 1.23 go 1.23
require golang.org/x/sys v0.29.0
+2
View File
@@ -0,0 +1,2 @@
golang.org/x/sys v0.29.0 h1:TPYlXGxvx1MGTn2GiZDhnjPA9wZzZeGKHHmKhHYvgaU=
golang.org/x/sys v0.29.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
+223
View File
@@ -0,0 +1,223 @@
// Package config loads gpu-turnstile's configuration from environment
// variables and an optional .env-style config file.
package config
import (
"bufio"
"fmt"
"io"
"log/slog"
"strconv"
"strings"
"time"
)
// Config holds every gpu-turnstile setting.
type Config struct {
ListenOllama string
ListenComfy string
OllamaURL string
ComfyURL string
UnloadTimeout time.Duration
JobTimeout time.Duration
LLMWaitTimeout time.Duration
UnloadPollInterval time.Duration
HistoryPollInterval time.Duration
ProbeTimeout time.Duration
FreeTimeout time.Duration
WarmTimeout time.Duration
ShutdownTimeout time.Duration
BackoffInitial time.Duration
BackoffMax time.Duration
PromptCaptureLimit int64
AutoUpdate bool
UpdateInterval time.Duration
UpdateRepo string
UpdateAsset string
// 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.
func Defaults() Config {
return Config{
ListenOllama: ":11434",
ListenComfy: ":8188",
OllamaURL: "http://127.0.0.1:11435",
ComfyURL: "http://127.0.0.1:8189",
UnloadTimeout: time.Minute,
JobTimeout: 15 * time.Minute,
LLMWaitTimeout: 10 * time.Minute,
UnloadPollInterval: 500 * time.Millisecond,
HistoryPollInterval: time.Second,
ProbeTimeout: 5 * time.Second,
FreeTimeout: 30 * time.Second,
WarmTimeout: 2 * time.Minute,
ShutdownTimeout: 10 * time.Second,
BackoffInitial: time.Second,
BackoffMax: time.Minute,
PromptCaptureLimit: 64 * 1024,
AutoUpdate: true,
UpdateInterval: 6 * time.Hour,
UpdateRepo: "https://git.rambossek.at/PUBLIC/gpu-turnstile",
UpdateAsset: "gpu-turnstile.exe",
LLMBusyMode: "wait",
LLMBusyStatus: 503,
BusyRetryAfter: 30,
LogLevel: slog.LevelWarn,
}
}
// ParseEnvFile parses a .env-style file: KEY=VALUE lines, blank lines and
// #-comments are ignored, no quoting. A line without '=' is an error.
func ParseEnvFile(r io.Reader) (map[string]string, error) {
values := make(map[string]string)
scanner := bufio.NewScanner(r)
scanner.Buffer(make([]byte, 64*1024), 1024*1024)
lineNo := 0
for scanner.Scan() {
lineNo++
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
key, value, ok := strings.Cut(line, "=")
if !ok {
return nil, fmt.Errorf("line %d: expected KEY=VALUE", lineNo)
}
key = strings.TrimSpace(key)
if key == "" {
return nil, fmt.Errorf("line %d: empty key", lineNo)
}
values[key] = strings.TrimSpace(value)
}
return values, scanner.Err()
}
func envDuration(getenv func(string) string, name string, dst *time.Duration) error {
v := getenv(name)
if v == "" {
return nil
}
d, err := time.ParseDuration(v)
if err != nil {
return fmt.Errorf("%s: %w", name, err)
}
*dst = d
return nil
}
// Load overlays values from getenv onto the Defaults. Unknown keys are
// ignored. Invalid values are fatal.
func Load(getenv func(string) string) (Config, error) {
cfg := Defaults()
for _, e := range []struct {
name string
dst *string
}{
{"LISTEN_OLLAMA", &cfg.ListenOllama},
{"LISTEN_COMFY", &cfg.ListenComfy},
{"OLLAMA_URL", &cfg.OllamaURL},
{"COMFY_URL", &cfg.ComfyURL},
{"WARM_MODEL", &cfg.WarmModel},
{"UPDATE_REPO", &cfg.UpdateRepo},
{"UPDATE_ASSET", &cfg.UpdateAsset},
{"LOG_FILE", &cfg.LogFile},
} {
if v := getenv(e.name); v != "" {
*e.dst = v
}
}
for _, e := range []struct {
name string
dst *time.Duration
}{
{"UNLOAD_TIMEOUT", &cfg.UnloadTimeout},
{"JOB_TIMEOUT", &cfg.JobTimeout},
{"LLM_WAIT_TIMEOUT", &cfg.LLMWaitTimeout},
{"UNLOAD_POLL_INTERVAL", &cfg.UnloadPollInterval},
{"HISTORY_POLL_INTERVAL", &cfg.HistoryPollInterval},
{"PROBE_TIMEOUT", &cfg.ProbeTimeout},
{"FREE_TIMEOUT", &cfg.FreeTimeout},
{"WARM_TIMEOUT", &cfg.WarmTimeout},
{"SHUTDOWN_TIMEOUT", &cfg.ShutdownTimeout},
{"BACKOFF_INITIAL", &cfg.BackoffInitial},
{"BACKOFF_MAX", &cfg.BackoffMax},
{"UPDATE_INTERVAL", &cfg.UpdateInterval},
} {
if err := envDuration(getenv, e.name, e.dst); err != nil {
return cfg, err
}
}
if v := getenv("PROMPT_CAPTURE_LIMIT"); v != "" {
n, err := strconv.ParseInt(v, 10, 64)
if err != nil || n < 0 {
return cfg, fmt.Errorf("PROMPT_CAPTURE_LIMIT: must be a non-negative integer (bytes)")
}
cfg.PromptCaptureLimit = n
}
if v := getenv("AUTO_UPDATE"); v != "" {
b, err := strconv.ParseBool(v)
if err != nil {
return cfg, fmt.Errorf("AUTO_UPDATE: must be a boolean (true/false)")
}
cfg.AutoUpdate = b
}
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
}
// LOGLEVEL is the canonical spelling; LOG_LEVEL is kept as an alias.
logLevelValue := getenv("LOGLEVEL")
if logLevelValue == "" {
logLevelValue = getenv("LOG_LEVEL")
}
if logLevelValue != "" {
var level slog.Level
if err := level.UnmarshalText([]byte(logLevelValue)); err != nil {
return cfg, fmt.Errorf("LOGLEVEL: %w", err)
}
cfg.LogLevel = level
}
switch strings.ToLower(getenv("LOG_FORMAT")) {
case "", "text":
case "json":
cfg.LogJSON = true
default:
return cfg, fmt.Errorf("LOG_FORMAT: must be \"text\" or \"json\"")
}
return cfg, nil
}
+100
View File
@@ -0,0 +1,100 @@
package config
import (
"log/slog"
"strings"
"testing"
"time"
)
func TestDefaults(t *testing.T) {
cfg, err := Load(func(string) string { return "" })
if err != nil {
t.Fatal(err)
}
if cfg.ListenOllama != ":11434" || cfg.ListenComfy != ":8188" {
t.Fatalf("listen addrs = %s %s", cfg.ListenOllama, cfg.ListenComfy)
}
if cfg.UnloadTimeout != time.Minute || cfg.JobTimeout != 15*time.Minute {
t.Fatalf("timeouts = %v %v", cfg.UnloadTimeout, cfg.JobTimeout)
}
if !cfg.AutoUpdate || cfg.UpdateInterval != 6*time.Hour {
t.Fatalf("update = %v %v", cfg.AutoUpdate, cfg.UpdateInterval)
}
if cfg.LogLevel != slog.LevelWarn {
t.Fatalf("log level = %v", cfg.LogLevel)
}
}
func TestParseEnvFile(t *testing.T) {
input := `# comment
OLLAMA_URL=http://host:11435
LOGLEVEL=debug
SPACED = value with spaces
`
values, err := ParseEnvFile(strings.NewReader(input))
if err != nil {
t.Fatal(err)
}
if values["OLLAMA_URL"] != "http://host:11435" {
t.Fatalf("OLLAMA_URL = %q", values["OLLAMA_URL"])
}
if values["LOGLEVEL"] != "debug" {
t.Fatalf("LOGLEVEL = %q", values["LOGLEVEL"])
}
if values["SPACED"] != "value with spaces" {
t.Fatalf("SPACED = %q", values["SPACED"])
}
}
func TestParseEnvFileMalformed(t *testing.T) {
_, err := ParseEnvFile(strings.NewReader("OK=1\nNOT_A_PAIR\n"))
if err == nil || !strings.Contains(err.Error(), "line 2") {
t.Fatalf("err = %v, want line 2 error", err)
}
}
func TestEnvOverridesFile(t *testing.T) {
file := map[string]string{"OLLAMA_URL": "http://file:1", "UNLOAD_TIMEOUT": "42s"}
env := map[string]string{"OLLAMA_URL": "http://env:2"}
getenv := func(k string) string {
if v := env[k]; v != "" {
return v
}
return file[k]
}
cfg, err := Load(getenv)
if err != nil {
t.Fatal(err)
}
if cfg.OllamaURL != "http://env:2" {
t.Fatalf("OllamaURL = %q, want env value", cfg.OllamaURL)
}
if cfg.UnloadTimeout != 42*time.Second {
t.Fatalf("UnloadTimeout = %v, want file value", cfg.UnloadTimeout)
}
}
func TestLoadErrors(t *testing.T) {
for _, tc := range []struct{ key, value string }{
{"UNLOAD_TIMEOUT", "bogus"},
{"PROMPT_CAPTURE_LIMIT", "-5"},
{"AUTO_UPDATE", "maybe"},
{"LOGLEVEL", "shouty"},
{"LOG_FORMAT", "yaml"},
{"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
}
return ""
})
if err == nil {
t.Errorf("%s=%s: expected error", tc.key, tc.value)
}
}
}
+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()
}
+51 -1
View File
@@ -47,11 +47,22 @@ type Config struct {
// LogColor enables ANSI colors in per-request log lines. Ignored when // LogColor enables ANSI colors in per-request log lines. Ignored when
// the log level is above INFO (request lines are not emitted at all). // the log level is above INFO (request lines are not emitted at all).
LogColor bool 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 // BackoffInitial and BackoffMax control the exponential retry backoff
// when an upstream refuses a connection: the wait doubles from // when an upstream refuses a connection: the wait doubles from
// BackoffInitial up to BackoffMax between attempts. Zero selects the // BackoffInitial up to BackoffMax between attempts. Zero selects the
@@ -77,11 +88,15 @@ 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 backoffInitial time.Duration
backoffMax time.Duration backoffMax time.Duration
busyMode string
busyStatus int
busyRetryAfter int
ollamaProxy *httputil.ReverseProxy ollamaProxy *httputil.ReverseProxy
comfyProxy *httputil.ReverseProxy comfyProxy *httputil.ReverseProxy
@@ -127,6 +142,18 @@ func New(cfg Config) (*Server, error) {
if backoffMax <= 0 { if backoffMax <= 0 {
backoffMax = time.Minute 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{ retry := &retryTransport{
base: http.DefaultTransport, base: http.DefaultTransport,
initial: backoffInitial, initial: backoffInitial,
@@ -136,11 +163,15 @@ func New(cfg Config) (*Server, error) {
return &Server{ return &Server{
cfg: cfg, cfg: cfg,
log: log, log: log,
logWriter: cfg.LogWriter,
freeTimeout: freeTimeout, freeTimeout: freeTimeout,
warmTimeout: warmTimeout, warmTimeout: warmTimeout,
captureLimit: captureLimit, captureLimit: captureLimit,
backoffInitial: backoffInitial, backoffInitial: backoffInitial,
backoffMax: backoffMax, backoffMax: backoffMax,
busyMode: busyMode,
busyStatus: busyStatus,
busyRetryAfter: busyRetryAfter,
ollamaProxy: newReverseProxy(ollamaURL, retry, log.With("upstream", "ollama")), ollamaProxy: newReverseProxy(ollamaURL, retry, log.With("upstream", "ollama")),
comfyProxy: newReverseProxy(comfyURL, retry, log.With("upstream", "comfy")), comfyProxy: newReverseProxy(comfyURL, retry, log.With("upstream", "comfy")),
}, nil }, nil
@@ -338,7 +369,11 @@ func (s *Server) reqLine(log *slog.Logger, code, line string, attrs ...any) {
fmt.Fprintf(&sb, " %v=%v", attrs[i], attrs[i+1]) fmt.Fprintf(&sb, " %v=%v", attrs[i], attrs[i+1])
} }
sb.WriteByte('\n') sb.WriteByte('\n')
os.Stderr.WriteString(sb.String()) w := s.logWriter
if w == nil {
w = os.Stderr
}
io.WriteString(w, sb.String())
} }
// logRequests logs one line per incoming request and one per completed // logRequests logs one line per incoming request and one per completed
@@ -402,12 +437,27 @@ 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
+106
View File
@@ -2,6 +2,7 @@ package proxy
import ( import (
"bufio" "bufio"
"context"
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
@@ -345,3 +346,108 @@ 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)
}
}
+35
View File
@@ -0,0 +1,35 @@
//go:build !windows
// Package service provides the non-Windows stubs for the Windows service
// integration. Run falls back to plain signal handling; install/remove
// are unsupported.
package service
import (
"context"
"errors"
"os/signal"
"syscall"
)
// Name matches the Windows service name.
const Name = "gpu-turnstile"
var errUnsupported = errors.New("service management is only supported on Windows")
// IsService is always false on non-Windows platforms.
func IsService() bool { return false }
// Run executes run with SIGINT/SIGTERM cancellation, mirroring the
// interactive behavior.
func Run(run func(ctx context.Context) error) error {
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer stop()
return run(ctx)
}
// Install is unsupported on non-Windows platforms.
func Install(string) error { return errUnsupported }
// Remove is unsupported on non-Windows platforms.
func Remove() error { return errUnsupported }
+120
View File
@@ -0,0 +1,120 @@
//go:build windows
// Package service integrates gpu-turnstile with the Windows Service
// Control Manager: running as a service with graceful stop, plus
// install/remove helpers.
package service
import (
"context"
"fmt"
"os"
"time"
"golang.org/x/sys/windows/svc"
"golang.org/x/sys/windows/svc/mgr"
)
// Name is the Windows service name.
const Name = "gpu-turnstile"
// IsService reports whether the process is running as a Windows service.
func IsService() bool {
isSvc, err := svc.IsWindowsService()
return err == nil && isSvc
}
// Run executes run as a Windows service. SCM Stop and Shutdown cancel the
// context passed to run, triggering the same graceful shutdown as SIGTERM
// in interactive mode.
func Run(run func(ctx context.Context) error) error {
return svc.Run(Name, &handler{run: run})
}
type handler struct {
run func(ctx context.Context) error
}
func (h *handler) Execute(_ []string, requests <-chan svc.ChangeRequest, status chan<- svc.Status) (bool, uint32) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
status <- svc.Status{State: svc.StartPending}
errCh := make(chan error, 1)
go func() { errCh <- h.run(ctx) }()
status <- svc.Status{State: svc.Running, Accepts: svc.AcceptStop | svc.AcceptShutdown}
for {
select {
case err := <-errCh:
status <- svc.Status{State: svc.Stopped}
if err != nil {
return true, 1
}
return false, 0
case c := <-requests:
switch c.Cmd {
case svc.Interrogate:
status <- c.CurrentStatus
case svc.Stop, svc.Shutdown:
status <- svc.Status{State: svc.StopPending}
cancel()
}
}
}
}
// Install registers gpu-turnstile as an auto-start Windows service whose
// binPath loads the given config file. Recovery actions restart the
// service after 5s on failure — this is also what brings up a staged
// update after the updater exits with a non-zero code.
func Install(configPath string) error {
exe, err := os.Executable()
if err != nil {
return err
}
m, err := mgr.Connect()
if err != nil {
return fmt.Errorf("connect to service manager (run as administrator): %w", err)
}
defer m.Disconnect()
binPath := fmt.Sprintf(`"%s" -config "%s"`, exe, configPath)
s, err := m.CreateService(Name, binPath, mgr.Config{
StartType: mgr.StartAutomatic,
DisplayName: "gpu-turnstile",
Description: "GPU arbitration proxy for Ollama and ComfyUI",
})
if err != nil {
return fmt.Errorf("create service: %w", err)
}
defer s.Close()
restart := mgr.RecoveryAction{Type: mgr.ServiceRestart, Delay: 5 * time.Second}
if err := s.SetRecoveryActions([]mgr.RecoveryAction{restart, restart, restart}, 24*60*60); err != nil {
return fmt.Errorf("set recovery actions: %w", err)
}
if err := s.SetRecoveryActionsOnNonCrashFailures(true); err != nil {
return fmt.Errorf("set failure actions flag: %w", err)
}
return nil
}
// Remove stops (if running) and unregisters the service.
func Remove() error {
m, err := mgr.Connect()
if err != nil {
return fmt.Errorf("connect to service manager (run as administrator): %w", err)
}
defer m.Disconnect()
s, err := m.OpenService(Name)
if err != nil {
return fmt.Errorf("open service: %w", err)
}
defer s.Close()
s.Control(svc.Stop) // ignore error: may already be stopped
if err := s.Delete(); err != nil {
return fmt.Errorf("delete service: %w", err)
}
return nil
}
+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-----
MCowBQYDK2VwAyEA+gJbSvgeYX58woPQGbSC8x8Zw4OTDiiQ7/19seZKfSQ=
-----END PUBLIC KEY-----
`
+254
View File
@@ -0,0 +1,254 @@
// Package update implements gpu-turnstile's self-updater: it polls the
// Gitea releases API, downloads the Windows binary of newer releases, and
// verifies its Ed25519 signature (produced by CI with OpenSSL) before
// swapping it in next to the running executable.
package update
import (
"context"
"crypto/ed25519"
"crypto/sha256"
"crypto/x509"
"encoding/hex"
"encoding/json"
"encoding/pem"
"fmt"
"io"
"log/slog"
"net/http"
"net/url"
"os"
"strconv"
"strings"
"time"
)
// maxAssetSize bounds release asset downloads.
const maxAssetSize = 512 << 20
// Updater checks one Gitea repository for newer releases.
type Updater struct {
Repo string // e.g. https://git.rambossek.at/PUBLIC/gpu-turnstile
Asset string // e.g. gpu-turnstile.exe
Version string // current version, e.g. v0.1.2 ("dev" disables updates)
Log *slog.Logger
Client *http.Client
}
type release struct {
TagName string `json:"tag_name"`
Assets []struct {
Name string `json:"name"`
BrowserDownloadURL string `json:"browser_download_url"`
} `json:"assets"`
}
func (u *Updater) logger() *slog.Logger {
if u.Log != nil {
return u.Log
}
return slog.Default()
}
func (u *Updater) httpClient() *http.Client {
if u.Client != nil {
return u.Client
}
return &http.Client{Timeout: 5 * time.Minute}
}
// apiURL derives <scheme>://<host>/api/v1/repos/<owner>/<name> from Repo.
func (u *Updater) apiURL() (string, error) {
repoURL, err := url.Parse(u.Repo)
if err != nil || repoURL.Scheme == "" || repoURL.Host == "" {
return "", fmt.Errorf("invalid UPDATE_REPO %q", u.Repo)
}
ownerName := strings.Trim(repoURL.Path, "/")
if len(strings.Split(ownerName, "/")) != 2 {
return "", fmt.Errorf("UPDATE_REPO %q: expected path /<owner>/<name>", u.Repo)
}
return fmt.Sprintf("%s://%s/api/v1/repos/%s", repoURL.Scheme, repoURL.Host, ownerName), nil
}
// newerVersion reports whether latest is a higher vX.Y.Z version than
// current. Both may carry a leading "v".
func newerVersion(current, latest string) (bool, error) {
parse := func(s string) ([3]int, error) {
var v [3]int
parts := strings.Split(strings.TrimPrefix(s, "v"), ".")
if len(parts) != 3 {
return v, fmt.Errorf("not a vX.Y.Z version: %q", s)
}
for i, p := range parts {
n, err := strconv.Atoi(p)
if err != nil {
return v, fmt.Errorf("not a vX.Y.Z version: %q", s)
}
v[i] = n
}
return v, nil
}
cur, err := parse(current)
if err != nil {
return false, err
}
lat, err := parse(latest)
if err != nil {
return false, err
}
for i := 0; i < 3; i++ {
if lat[i] != cur[i] {
return lat[i] > cur[i], nil
}
}
return false, nil
}
func (u *Updater) get(ctx context.Context, url string) ([]byte, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, err
}
resp, err := u.httpClient().Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
io.Copy(io.Discard, resp.Body)
return nil, fmt.Errorf("GET %s: %s", url, resp.Status)
}
return io.ReadAll(io.LimitReader(resp.Body, maxAssetSize))
}
func verifySignature(pubKeyPEM string, data, sig []byte) error {
block, _ := pem.Decode([]byte(pubKeyPEM))
if block == nil {
return fmt.Errorf("invalid embedded public key PEM")
}
key, err := x509.ParsePKIXPublicKey(block.Bytes)
if err != nil {
return fmt.Errorf("parse public key: %w", err)
}
pub, ok := key.(ed25519.PublicKey)
if !ok {
return fmt.Errorf("public key is not Ed25519")
}
if !ed25519.Verify(pub, data, sig) {
return fmt.Errorf("signature verification failed")
}
return nil
}
// stage swaps data into place at exePath: the running executable is
// renamed aside (allowed on Windows) and the new file takes its name.
func stage(exePath string, data []byte) error {
newPath := exePath + ".new"
oldPath := exePath + ".old"
os.Remove(oldPath) // leftover from a previous update
if err := os.WriteFile(newPath, data, 0o755); err != nil {
return err
}
if err := os.Rename(exePath, oldPath); err != nil {
os.Remove(newPath)
return err
}
if err := os.Rename(newPath, exePath); err != nil {
os.Rename(oldPath, exePath) // roll back
return err
}
return nil
}
// CleanupOld removes the .old binary left behind by a staged update.
// Call once at startup.
func CleanupOld(exePath string) {
os.Remove(exePath + ".old")
os.Remove(exePath + ".new")
}
// Check performs a single update check. staged is true when a newer,
// signature-verified binary has been swapped into place at exePath; the
// caller should then restart the process. A nil error with staged=false
// means "no action" (up to date, disabled, or dev build); a non-nil error
// means the check failed and the running binary is untouched.
func (u *Updater) Check(ctx context.Context, exePath string) (staged bool, err error) {
log := u.logger()
if u.Version == "" || u.Version == "dev" {
log.Debug("auto-update: dev build, skipping")
return false, nil
}
if publicKeyPEM == "" {
log.Debug("auto-update: no public key embedded, skipping")
return false, nil
}
api, err := u.apiURL()
if err != nil {
return false, err
}
body, err := u.get(ctx, api+"/releases/latest")
if err != nil {
return false, fmt.Errorf("fetch latest release: %w", err)
}
var rel release
if err := json.Unmarshal(body, &rel); err != nil {
return false, fmt.Errorf("parse release: %w", err)
}
newer, err := newerVersion(u.Version, rel.TagName)
if err != nil {
return false, err
}
if !newer {
log.Debug("auto-update: up to date", "version", u.Version, "latest", rel.TagName)
return false, nil
}
urls := make(map[string]string, len(rel.Assets))
for _, a := range rel.Assets {
urls[a.Name] = a.BrowserDownloadURL
}
assetURL, ok := urls[u.Asset]
if !ok {
return false, fmt.Errorf("release %s has no asset %q", rel.TagName, u.Asset)
}
sigURL, ok := urls[u.Asset+".sig"]
if !ok {
return false, fmt.Errorf("release %s has no signature asset %q", rel.TagName, u.Asset+".sig")
}
data, err := u.get(ctx, assetURL)
if err != nil {
return false, fmt.Errorf("download %s: %w", u.Asset, err)
}
sig, err := u.get(ctx, sigURL)
if err != nil {
return false, fmt.Errorf("download signature: %w", err)
}
if sumURL, ok := urls[u.Asset+".sha256"]; ok {
sumText, err := u.get(ctx, sumURL)
if err != nil {
return false, fmt.Errorf("download checksum: %w", err)
}
want := strings.Fields(string(sumText))[0]
got := hex.EncodeToString(sha256Bytes(data))
if !strings.EqualFold(want, got) {
return false, fmt.Errorf("sha256 mismatch: got %s, want %s", got, want)
}
}
if err := verifySignature(publicKeyPEM, data, sig); err != nil {
return false, err
}
if err := stage(exePath, data); err != nil {
return false, fmt.Errorf("stage update: %w", err)
}
log.Info("auto-update: new version staged", "from", u.Version, "to", rel.TagName)
return true, nil
}
func sha256Bytes(data []byte) []byte {
sum := sha256.Sum256(data)
return sum[:]
}
+186
View File
@@ -0,0 +1,186 @@
package update
import (
"context"
"crypto/ed25519"
"crypto/rand"
"crypto/x509"
"encoding/json"
"encoding/pem"
"fmt"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"testing"
)
// fakeGitea serves a Gitea-flavored releases API with one release.
type fakeGitea struct {
srv *httptest.Server
pubPEM string
asset []byte
tag string
tamper bool
noSig bool
}
func newFakeGitea(t *testing.T, tag string, assetContent []byte) *fakeGitea {
t.Helper()
pub, priv, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
t.Fatal(err)
}
der, err := x509.MarshalPKIXPublicKey(pub)
if err != nil {
t.Fatal(err)
}
f := &fakeGitea{
pubPEM: string(pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: der})),
asset: assetContent,
tag: tag,
}
sign := func() []byte { return ed25519.Sign(priv, f.asset) }
mux := http.NewServeMux()
mux.HandleFunc("/api/v1/repos/o/r/releases/latest", func(w http.ResponseWriter, r *http.Request) {
assets := []map[string]string{
{"name": "gpu-turnstile.exe", "browser_download_url": f.srv.URL + "/dl/exe"},
{"name": "gpu-turnstile.exe.sha256", "browser_download_url": f.srv.URL + "/dl/sha"},
}
if !f.noSig {
assets = append(assets, map[string]string{"name": "gpu-turnstile.exe.sig", "browser_download_url": f.srv.URL + "/dl/sig"})
}
json.NewEncoder(w).Encode(map[string]any{"tag_name": f.tag, "assets": assets})
})
mux.HandleFunc("/dl/exe", func(w http.ResponseWriter, r *http.Request) { w.Write(f.asset) })
mux.HandleFunc("/dl/sig", func(w http.ResponseWriter, r *http.Request) {
sig := sign()
if f.tamper {
sig[0] ^= 0xff
}
w.Write(sig)
})
mux.HandleFunc("/dl/sha", func(w http.ResponseWriter, r *http.Request) {
fmt.Fprintf(w, "%x gpu-turnstile.exe\n", sha256Bytes(f.asset))
})
f.srv = httptest.NewServer(mux)
t.Cleanup(f.srv.Close)
return f
}
func (f *fakeGitea) updater(version string) *Updater {
return &Updater{Repo: f.srv.URL + "/o/r", Asset: "gpu-turnstile.exe", Version: version}
}
func fakeExe(t *testing.T) string {
t.Helper()
exe := filepath.Join(t.TempDir(), "gpu-turnstile.exe")
if err := os.WriteFile(exe, []byte("old-binary"), 0o755); err != nil {
t.Fatal(err)
}
return exe
}
func withPublicKey(t *testing.T, pem string) {
t.Helper()
old := publicKeyPEM
publicKeyPEM = pem
t.Cleanup(func() { publicKeyPEM = old })
}
func TestCheckStagesUpdate(t *testing.T) {
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
withPublicKey(t, f.pubPEM)
exe := fakeExe(t)
staged, err := f.updater("v0.1.2").Check(context.Background(), exe)
if err != nil {
t.Fatal(err)
}
if !staged {
t.Fatal("expected staged update")
}
content, _ := os.ReadFile(exe)
if string(content) != "new-binary" {
t.Fatalf("exe content = %q", content)
}
old, _ := os.ReadFile(exe + ".old")
if string(old) != "old-binary" {
t.Fatalf(".old content = %q", old)
}
}
func TestCheckRejectsTamperedSignature(t *testing.T) {
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
f.tamper = true
withPublicKey(t, f.pubPEM)
exe := fakeExe(t)
staged, err := f.updater("v0.1.2").Check(context.Background(), exe)
if err == nil {
t.Fatal("expected signature error")
}
if staged {
t.Fatal("must not stage on bad signature")
}
content, _ := os.ReadFile(exe)
if string(content) != "old-binary" {
t.Fatal("exe was modified despite bad signature")
}
}
func TestCheckSkipsOlderOrEqual(t *testing.T) {
for _, tag := range []string{"v0.1.2", "v0.1.1", "v0.0.9"} {
f := newFakeGitea(t, tag, []byte("new-binary"))
withPublicKey(t, f.pubPEM)
staged, err := f.updater("v0.1.2").Check(context.Background(), fakeExe(t))
if err != nil {
t.Fatal(err)
}
if staged {
t.Fatalf("tag %s must not stage over v0.1.2", tag)
}
}
}
func TestCheckSkipsWithoutPublicKey(t *testing.T) {
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
withPublicKey(t, "")
staged, err := f.updater("v0.1.2").Check(context.Background(), fakeExe(t))
if err != nil || staged {
t.Fatalf("staged=%v err=%v, want no action without key", staged, err)
}
}
func TestCheckSkipsDevBuild(t *testing.T) {
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
withPublicKey(t, f.pubPEM)
staged, err := f.updater("dev").Check(context.Background(), fakeExe(t))
if err != nil || staged {
t.Fatalf("staged=%v err=%v, want no action for dev build", staged, err)
}
}
func TestNewerVersion(t *testing.T) {
cases := []struct {
cur, lat string
want bool
}{
{"v0.1.2", "v0.1.3", true},
{"0.1.2", "0.2.0", true},
{"v0.1.2", "v1.0.0", true},
{"v0.1.2", "v0.1.2", false},
{"v1.2.3", "v1.2.10", true},
{"v1.2.10", "v1.2.3", false},
}
for _, c := range cases {
got, err := newerVersion(c.cur, c.lat)
if err != nil || got != c.want {
t.Errorf("newerVersion(%s, %s) = %v, %v; want %v", c.cur, c.lat, got, err, c.want)
}
}
if _, err := newerVersion("v0.1", "v0.1.2"); err == nil {
t.Error("expected error for malformed version")
}
}