Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
da02457fd5 | ||
|
|
6a6359a630 | ||
|
|
e4e4348e8d | ||
|
|
b421bb7bfb | ||
|
|
909918657f | ||
|
|
a0435c858e | ||
|
|
fbab0bba33 | ||
|
|
802a64280f | ||
|
|
a88955e35c | ||
|
|
97624470eb | ||
|
|
75f16a0229 | ||
|
|
9481fd8418 | ||
|
|
83812cebf3 | ||
|
|
30a1f55aae | ||
|
|
642cc36a39 | ||
|
|
19e19281de | ||
|
|
08d02d8fa8 | ||
|
|
d1e01f9b78 | ||
|
|
bacb26772a | ||
|
|
e363c9e4e4 |
@@ -60,3 +60,54 @@ jobs:
|
||||
tags: |
|
||||
${{ steps.meta.outputs.image }}:${{ gitea.ref_name }}
|
||||
${{ steps.meta.outputs.image }}:latest
|
||||
|
||||
# On version tags: build the signed Windows binary and attach it (plus
|
||||
# signature and checksum) to a Gitea release for the auto-updater.
|
||||
release:
|
||||
if: gitea.ref_type == 'tag'
|
||||
needs: test
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version: "1.23"
|
||||
|
||||
- name: Build Windows binary
|
||||
run: |
|
||||
GOOS=windows GOARCH=amd64 CGO_ENABLED=0 go build \
|
||||
-ldflags="-s -w -X main.version=${{ gitea.ref_name }}" \
|
||||
-o gpu-turnstile.exe ./cmd/gpu-turnstile
|
||||
|
||||
- name: Sign and checksum
|
||||
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,3 +2,5 @@
|
||||
/gpu-turnstile.exe
|
||||
/compose.yaml
|
||||
/compose.yml
|
||||
/signing/
|
||||
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
FROM golang:1.23 AS build
|
||||
WORKDIR /src
|
||||
COPY go.mod ./
|
||||
COPY go.mod go.sum ./
|
||||
COPY cmd ./cmd
|
||||
COPY internal ./internal
|
||||
ARG VERSION=dev
|
||||
|
||||
@@ -18,7 +18,10 @@ Open WebUI / n8n ────► :8188 ───┘
|
||||
|
||||
- LLM endpoints (`/api/generate`, `/api/chat`, `/api/embed`, `/v1/*`) take the
|
||||
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
|
||||
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
|
||||
@@ -26,23 +29,36 @@ Open WebUI / n8n ────► :8188 ───┘
|
||||
- Everything else (including websockets and all streaming) passes through
|
||||
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
|
||||
|
||||
All configuration is via environment variables; invalid values fail at
|
||||
startup.
|
||||
Configuration comes from environment variables and/or an `.env`-style
|
||||
config file (`KEY=VALUE` lines, `#` comments). File lookup order:
|
||||
`-config <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 |
|
||||
|---|---|---|
|
||||
| `LISTEN_OLLAMA` | `:11434` | Ollama-facing listener |
|
||||
| `LISTEN_COMFY` | `:8188` | ComfyUI-facing listener |
|
||||
| `OLLAMA_URL` | `http://127.0.0.1:11435` | Ollama upstream |
|
||||
| `COMFY_URL` | `http://127.0.0.1:8189` | ComfyUI upstream |
|
||||
| `OLLAMA_URL` | _(empty = disabled)_ | Ollama upstream; set to enable the Ollama consumer |
|
||||
| `COMFY_URL` | _(empty = disabled)_ | ComfyUI upstream; set to enable the ComfyUI consumer |
|
||||
| `UNLOAD_TIMEOUT` | `60s` | Wait for Ollama to unload before an image job |
|
||||
| `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 (400–599, 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) |
|
||||
| `LOGLEVEL` | `warn` | `info` logs every request (colored arrows in text mode), `debug` adds lock transitions. `LOG_LEVEL` works as an alias |
|
||||
| `LOG_FORMAT` | `text` | `json` for structured JSON logs |
|
||||
| `LOG_FILE` | _(empty)_ | Append logs to this file instead of stderr |
|
||||
| `UNLOAD_POLL_INTERVAL` | `500ms` | `/api/ps` poll interval while unloading |
|
||||
| `HISTORY_POLL_INTERVAL` | `1s` | `/history/<id>` poll interval while a job runs |
|
||||
| `PROBE_TIMEOUT` | `5s` | Startup probe of both upstreams |
|
||||
@@ -52,11 +68,15 @@ startup.
|
||||
| `BACKOFF_INITIAL` | `1s` | First retry wait when an upstream refuses a connection |
|
||||
| `BACKOFF_MAX` | `60s` | Cap for the exponential retry backoff |
|
||||
| `PROMPT_CAPTURE_LIMIT` | `65536` | Bytes of the `/prompt` response buffered to find `prompt_id` (pass-through is unaffected) |
|
||||
| `AUTO_UPDATE` | `true` | Poll the Gitea releases API for signed updates |
|
||||
| `UPDATE_INTERVAL` | `6h` | Auto-update check interval |
|
||||
| `UPDATE_REPO` | `https://git.rambossek.at/PUBLIC/gpu-turnstile` | Repository checked for releases |
|
||||
| `UPDATE_ASSET` | `gpu-turnstile.exe` | Release asset to download |
|
||||
|
||||
## Observability
|
||||
|
||||
- `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_image_jobs_total`, `gpu_turnstile_lock_wait_seconds`
|
||||
(histogram, `kind="llm|image"`), `gpu_turnstile_unload_seconds`.
|
||||
@@ -71,9 +91,69 @@ startup.
|
||||
|
||||
```sh
|
||||
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\` (set
|
||||
`LOG_FILE` in the env file — there is no console). 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.
|
||||
|
||||
### 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. 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).
|
||||
`gpu-turnstile --force-update` checks immediately, stages the new binary
|
||||
and restarts the running service (elevating via UAC only if needed).
|
||||
|
||||
### Docker
|
||||
|
||||
```sh
|
||||
docker build -t gpu-turnstile .
|
||||
docker run --rm -p 11434:11434 -p 8188:8188 \
|
||||
@@ -84,9 +164,10 @@ docker run --rm -p 11434:11434 -p 8188:8188 \
|
||||
|
||||
Releases are built by Gitea Actions (`.gitea/workflows/ci.yml`): every push
|
||||
runs `go vet` and `go test -race`, and pushing a semantic-version tag
|
||||
`vX.Y.Z` builds and publishes
|
||||
`git.rambossek.at/<owner>/gpu-turnstile:vX.Y.Z` (and updates `:latest`).
|
||||
No images are built from branches.
|
||||
`vX.Y.Z` publishes the container image
|
||||
(`git.rambossek.at/<owner>/gpu-turnstile:vX.Y.Z` plus `:latest`) and a
|
||||
signed Windows binary attached to a Gitea release. Nothing is built from
|
||||
branches.
|
||||
|
||||
The registry login needs one repository secret (Settings → Actions →
|
||||
Secrets): `REGISTRY_TOKEN` — an access token with `write:package` scope.
|
||||
@@ -99,13 +180,17 @@ go vet ./...
|
||||
go test -race ./...
|
||||
```
|
||||
|
||||
Stdlib only, Go 1.23+. Layout:
|
||||
Go 1.23+; the only external dependency is `golang.org/x/sys` (Windows
|
||||
service integration, unused in the Linux build). Layout:
|
||||
|
||||
```
|
||||
cmd/gpu-turnstile/main.go wiring, config, listeners
|
||||
cmd/gpu-turnstile/main.go wiring, config, listeners, service + updater
|
||||
internal/lock/ two-mode lock (LLM readers / image writer, FIFO)
|
||||
internal/ollama/ ps / unload / warm client
|
||||
internal/comfy/ history poll / free client
|
||||
internal/proxy/ handlers for both listeners
|
||||
internal/metrics/ Prometheus exposition, no dependencies
|
||||
internal/config/ env + .env file configuration
|
||||
internal/update/ signed auto-updater
|
||||
internal/service/ Windows service integration
|
||||
```
|
||||
|
||||
@@ -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
|
||||
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
|
||||
|
||||
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++.
|
||||
On completion (response fully written, including streamed bodies, or client
|
||||
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
|
||||
requests start), waits until n == 0, sets state := `image`. Released after
|
||||
the ComfyUI job finished and models were freed.
|
||||
@@ -78,7 +98,8 @@ ComfyUI listener (`:8188` → `COMFY_URL`):
|
||||
### Image job flow (`POST /prompt`)
|
||||
|
||||
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
|
||||
models), `POST /api/embed {"model":M,"input":"x","keep_alive":0}`. Poll
|
||||
`/api/ps` every `UNLOAD_POLL_INTERVAL` (default 500 ms) until empty or
|
||||
@@ -102,17 +123,28 @@ load time. Off by default.
|
||||
|
||||
## Configuration (env)
|
||||
|
||||
Configuration comes from environment variables and/or an `.env`-style
|
||||
config file (`KEY=VALUE` lines, `#` comments). File lookup order:
|
||||
`-config <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 |
|
||||
|---|---|---|
|
||||
| `LISTEN_OLLAMA` | `:11434` | Ollama-facing listener |
|
||||
| `LISTEN_COMFY` | `:8188` | ComfyUI-facing listener |
|
||||
| `OLLAMA_URL` | `http://127.0.0.1:11435` | upstream |
|
||||
| `COMFY_URL` | `http://127.0.0.1:8189` | upstream |
|
||||
| `OLLAMA_URL` | _(empty = disabled)_ | Ollama upstream; set to enable the Ollama consumer |
|
||||
| `COMFY_URL` | _(empty = disabled)_ | ComfyUI upstream; set to enable the ComfyUI consumer |
|
||||
| `UNLOAD_TIMEOUT` | `60s` | wait for Ollama to unload |
|
||||
| `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 (400–599, 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 |
|
||||
| `LOGLEVEL` | `warn` | `info` logs every request (colored arrows in text mode), `debug` adds lock transitions. `LOG_LEVEL` is accepted as an alias |
|
||||
| `LOG_FORMAT` | `text` | `json` for structured JSON logs |
|
||||
| `LOG_FILE` | `` | append logs to this file instead of stderr (useful as a service) |
|
||||
| `UNLOAD_POLL_INTERVAL` | `500ms` | `/api/ps` poll interval while unloading |
|
||||
| `HISTORY_POLL_INTERVAL` | `1s` | `/history/<id>` poll interval while a job runs |
|
||||
| `PROBE_TIMEOUT` | `5s` | startup probe of both upstreams |
|
||||
@@ -122,15 +154,114 @@ load time. Off by default.
|
||||
| `BACKOFF_INITIAL` | `1s` | first retry wait when an upstream refuses a connection |
|
||||
| `BACKOFF_MAX` | `60s` | cap for the exponential retry backoff |
|
||||
| `PROMPT_CAPTURE_LIMIT` | `65536` | bytes of the `/prompt` response buffered to find `prompt_id` (pass-through is unaffected) |
|
||||
| `AUTO_UPDATE` | `true` | poll the Gitea releases API for signed updates |
|
||||
| `UPDATE_INTERVAL` | `6h` | auto-update check interval |
|
||||
| `UPDATE_REPO` | `https://git.rambossek.at/PUBLIC/gpu-turnstile` | repository to check for releases |
|
||||
| `UPDATE_ASSET` | `gpu-turnstile.exe` | release asset to download |
|
||||
|
||||
Startup fails fast on unparsable values. Both upstreams are probed once at
|
||||
start (`/api/version`, `/system_stats`); failure is logged, not fatal.
|
||||
Startup fails fast on unparsable values and when neither consumer URL is
|
||||
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.
|
||||
|
||||
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`). An existing config in the
|
||||
target location is never overwritten. `--no-copy` registers the current
|
||||
executable location as-is instead.
|
||||
|
||||
### 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. Logs go to the ProgramData directory via
|
||||
`LOG_FILE` since there is no console.
|
||||
- **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
|
||||
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.
|
||||
- **`--force-update`** runs the same check immediately: it downloads,
|
||||
verifies and stages a newer release, 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
|
||||
|
||||
- `GET /healthz` on both listeners: 200 with JSON
|
||||
`{"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:
|
||||
`gpu_turnstile_state{state="…"} 1`, `gpu_turnstile_llm_inflight`,
|
||||
`gpu_turnstile_image_jobs_total`, `gpu_turnstile_lock_wait_seconds`
|
||||
@@ -171,11 +302,15 @@ start (`/api/version`, `/system_stats`); failure is logged, not fatal.
|
||||
|
||||
```
|
||||
gpu-turnstile/
|
||||
cmd/gpu-turnstile/main.go # wiring, config, listeners
|
||||
cmd/gpu-turnstile/main.go # wiring, config, listeners, service + updater
|
||||
internal/lock/lock.go # two-mode lock + tests
|
||||
internal/ollama/client.go # ps / unload / warm
|
||||
internal/comfy/client.go # history poll / free
|
||||
internal/proxy/ # handlers for both listeners
|
||||
internal/metrics/ # Prometheus exposition
|
||||
internal/config/ # env + .env file configuration
|
||||
internal/update/ # signed auto-updater (public key in pubkey.go)
|
||||
internal/service/ # Windows SCM + Linux systemd (notify/watchdog) integration
|
||||
Dockerfile
|
||||
.gitea/workflows/ci.yml
|
||||
README.md
|
||||
@@ -197,11 +332,17 @@ are new.
|
||||
held until `/free` was called.
|
||||
- Streaming test: fake Ollama emits chunks with delays; assert the client
|
||||
receives the first chunk before the last is sent (no buffering).
|
||||
- `internal/config`: env-file parsing, precedence, fail-fast values.
|
||||
- `internal/update`: fake Gitea releases API; staged update happy path,
|
||||
tampered signature rejected, older versions and dev builds skipped.
|
||||
|
||||
## Build and CI
|
||||
|
||||
- Go 1.23+, stdlib only. `CGO_ENABLED=0`, `-ldflags="-s -w"`, version from
|
||||
`git describe` injected via `-X main.version=`.
|
||||
- Go 1.23+, two external dependencies: `golang.org/x/sys` (Windows service
|
||||
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
|
||||
`scratch`), non-root user, `EXPOSE 8188 11434`,
|
||||
`ENTRYPOINT ["/gpu-turnstile"]`.
|
||||
@@ -216,8 +357,12 @@ are new.
|
||||
Login uses the repo secret `REGISTRY_TOKEN` (an access token with
|
||||
`write:package` scope) because the automatic `GITEA_TOKEN` cannot push
|
||||
packages; the username is just `gitea.actor`.
|
||||
- Release: a git tag `vX.Y.Z` produces the versioned image; the Open WebUI
|
||||
compose pins that tag. No images are built from branches.
|
||||
3. on a version tag: also build the Windows binary, sign it with OpenSSL
|
||||
(`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)
|
||||
|
||||
|
||||
+531
-185
@@ -3,259 +3,605 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/fs"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"strconv"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"gpu-turnstile/internal/comfy"
|
||||
"gpu-turnstile/internal/config"
|
||||
"gpu-turnstile/internal/lock"
|
||||
"gpu-turnstile/internal/metrics"
|
||||
"gpu-turnstile/internal/ollama"
|
||||
"gpu-turnstile/internal/proxy"
|
||||
"gpu-turnstile/internal/service"
|
||||
"gpu-turnstile/internal/update"
|
||||
)
|
||||
|
||||
// version is injected at build time via -ldflags "-X main.version=...".
|
||||
var version = "dev"
|
||||
|
||||
type config struct {
|
||||
listenOllama string
|
||||
listenComfy string
|
||||
ollamaURL string
|
||||
comfyURL string
|
||||
unloadTimeout time.Duration
|
||||
jobTimeout time.Duration
|
||||
llmWaitTimeout time.Duration
|
||||
// exitCodeUpdate tells the service recovery configuration to restart the
|
||||
// process: a signed update has been staged and the GPU lock is idle.
|
||||
const exitCodeUpdate = 3
|
||||
|
||||
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
|
||||
// stdoutIsTerminal reports whether stdout is a console (char device), as
|
||||
// opposed to a pipe or file — which is what Docker containers and services
|
||||
// see.
|
||||
func stdoutIsTerminal() bool {
|
||||
fi, err := os.Stdout.Stat()
|
||||
return err == nil && fi.Mode()&os.ModeCharDevice != 0
|
||||
}
|
||||
|
||||
func envDuration(getenv func(string) string, name string, dst *time.Duration) error {
|
||||
v := getenv(name)
|
||||
if v == "" {
|
||||
return nil
|
||||
// parseFlags extracts -config <path> (or -config=<path>), the
|
||||
// --install-service / --remove-service switches, --no-copy, -h/--help,
|
||||
// -v/--version, --force-update and the hidden --elevated-child marker from
|
||||
// args.
|
||||
func parseFlags(args []string) (configPath string, install, remove, noCopy, help, showVersion, forceUpdate, elevatedChild bool, rest []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=")
|
||||
case args[i] == "--install-service" || args[i] == "-install-service":
|
||||
install = true
|
||||
case args[i] == "--remove-service" || args[i] == "-remove-service":
|
||||
remove = true
|
||||
case args[i] == "--no-copy" || args[i] == "-no-copy":
|
||||
noCopy = true
|
||||
case args[i] == "-h" || args[i] == "--help" || args[i] == "-help":
|
||||
help = true
|
||||
case args[i] == "-v" || args[i] == "--version" || args[i] == "-version":
|
||||
showVersion = true
|
||||
case args[i] == "--force-update" || args[i] == "-force-update":
|
||||
forceUpdate = true
|
||||
case args[i] == "--elevated-child":
|
||||
elevatedChild = true
|
||||
default:
|
||||
rest = append(rest, args[i])
|
||||
}
|
||||
}
|
||||
d, err := time.ParseDuration(v)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%s: %w", name, err)
|
||||
}
|
||||
*dst = d
|
||||
return nil
|
||||
return configPath, install, remove, noCopy, help, showVersion, forceUpdate, elevatedChild, rest
|
||||
}
|
||||
|
||||
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,
|
||||
// versionLine is printed at the top of every help and error screen.
|
||||
func versionLine() string { return "gpu-turnstile " + version }
|
||||
|
||||
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,
|
||||
const usageText = `GPU arbitration proxy for Ollama + ComfyUI
|
||||
|
||||
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
|
||||
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() {
|
||||
cfg, err := loadConfig(os.Getenv)
|
||||
if err != nil {
|
||||
fmt.Fprintf(os.Stderr, "gpu-turnstile: %v\n", err)
|
||||
os.Exit(1)
|
||||
configPath, install, remove, noCopy, help, showVersion, forceUpdate, elevatedChild, args := parseFlags(os.Args[1:])
|
||||
if showVersion {
|
||||
fmt.Println(version)
|
||||
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}
|
||||
var handler slog.Handler = slog.NewTextHandler(os.Stderr, opts)
|
||||
if cfg.logJSON {
|
||||
handler = slog.NewJSONHandler(os.Stderr, opts)
|
||||
if exePath, err := os.Executable(); err == nil {
|
||||
update.CleanupOld(exePath)
|
||||
}
|
||||
|
||||
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)
|
||||
slog.SetDefault(log)
|
||||
return log, out, closer
|
||||
}
|
||||
|
||||
// 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 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()
|
||||
u := &update.Updater{Repo: cfg.UpdateRepo, Asset: cfg.UpdateAsset, Version: version, Log: log}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
||||
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)
|
||||
} 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(context.Background(), slog.LevelWarn, "starting gpu-turnstile",
|
||||
log.Log(ctx, slog.LevelWarn, "starting gpu-turnstile",
|
||||
"version", version,
|
||||
"listen_ollama", cfg.listenOllama,
|
||||
"listen_comfy", cfg.listenComfy,
|
||||
"ollama_url", cfg.ollamaURL,
|
||||
"comfy_url", cfg.comfyURL,
|
||||
"unload_timeout", cfg.unloadTimeout,
|
||||
"job_timeout", cfg.jobTimeout,
|
||||
"llm_wait_timeout", cfg.llmWaitTimeout,
|
||||
"unload_poll_interval", cfg.unloadPollInterval,
|
||||
"history_poll_interval", cfg.historyPollInterval,
|
||||
"probe_timeout", cfg.probeTimeout,
|
||||
"free_timeout", cfg.freeTimeout,
|
||||
"warm_timeout", cfg.warmTimeout,
|
||||
"shutdown_timeout", cfg.shutdownTimeout,
|
||||
"backoff_initial", cfg.backoffInitial,
|
||||
"backoff_max", cfg.backoffMax,
|
||||
"prompt_capture_limit", cfg.promptCaptureLimit,
|
||||
"warm_model", cfg.warmModel,
|
||||
"log_level", cfg.logLevel,
|
||||
"log_format", map[bool]string{true: "json", false: "text"}[cfg.logJSON],
|
||||
"listen_ollama", cfg.ListenOllama,
|
||||
"listen_comfy", cfg.ListenComfy,
|
||||
"ollama_url", orDisabled(cfg.OllamaURL),
|
||||
"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)
|
||||
if err != nil {
|
||||
log.Error("invalid configuration", "err", err)
|
||||
os.Exit(1)
|
||||
lk := lock.New(log)
|
||||
// Each consumer is enabled by setting its URL; a disabled consumer gets
|
||||
// no client, no listener and no probe.
|
||||
var ollamaClient *ollama.Client
|
||||
var err error
|
||||
if cfg.OllamaURL != "" {
|
||||
if ollamaClient, err = ollama.New(cfg.OllamaURL, log); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
comfyClient, err := comfy.New(cfg.comfyURL, log)
|
||||
if err != nil {
|
||||
log.Error("invalid configuration", "err", err)
|
||||
os.Exit(1)
|
||||
var comfyClient *comfy.Client
|
||||
if cfg.ComfyURL != "" {
|
||||
if comfyClient, err = comfy.New(cfg.ComfyURL, log); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
srv, err := proxy.New(proxy.Config{
|
||||
OllamaURL: cfg.ollamaURL,
|
||||
ComfyURL: cfg.comfyURL,
|
||||
Lock: lock.New(log),
|
||||
OllamaURL: cfg.OllamaURL,
|
||||
ComfyURL: cfg.ComfyURL,
|
||||
Lock: lk,
|
||||
Ollama: ollamaClient,
|
||||
Comfy: comfyClient,
|
||||
Metrics: metrics.New(),
|
||||
Log: log,
|
||||
LogColor: !cfg.logJSON && os.Getenv("NO_COLOR") == "",
|
||||
LLMWaitTimeout: cfg.llmWaitTimeout,
|
||||
UnloadTimeout: cfg.unloadTimeout,
|
||||
JobTimeout: cfg.jobTimeout,
|
||||
UnloadPollInterval: cfg.unloadPollInterval,
|
||||
HistoryPollInterval: cfg.historyPollInterval,
|
||||
FreeTimeout: cfg.freeTimeout,
|
||||
WarmTimeout: cfg.warmTimeout,
|
||||
BackoffInitial: cfg.backoffInitial,
|
||||
BackoffMax: cfg.backoffMax,
|
||||
PromptCaptureLimit: cfg.promptCaptureLimit,
|
||||
WarmModel: cfg.warmModel,
|
||||
LogColor: !cfg.LogJSON && cfg.LogFile == "" && os.Getenv("NO_COLOR") == "",
|
||||
LogWriter: logOut,
|
||||
LLMWaitTimeout: cfg.LLMWaitTimeout,
|
||||
UnloadTimeout: cfg.UnloadTimeout,
|
||||
JobTimeout: cfg.JobTimeout,
|
||||
UnloadPollInterval: cfg.UnloadPollInterval,
|
||||
HistoryPollInterval: cfg.HistoryPollInterval,
|
||||
FreeTimeout: cfg.FreeTimeout,
|
||||
WarmTimeout: cfg.WarmTimeout,
|
||||
LLMBusyMode: cfg.LLMBusyMode,
|
||||
LLMBusyStatus: cfg.LLMBusyStatus,
|
||||
BusyRetryAfter: cfg.BusyRetryAfter,
|
||||
BackoffInitial: cfg.BackoffInitial,
|
||||
BackoffMax: cfg.BackoffMax,
|
||||
PromptCaptureLimit: cfg.PromptCaptureLimit,
|
||||
WarmModel: cfg.WarmModel,
|
||||
})
|
||||
if err != nil {
|
||||
log.Error("invalid configuration", "err", err)
|
||||
os.Exit(1)
|
||||
return err
|
||||
}
|
||||
|
||||
// Probe both upstreams once; failure is logged, not fatal.
|
||||
probeCtx, probeCancel := context.WithTimeout(context.Background(), cfg.probeTimeout)
|
||||
if err := ollamaClient.Probe(probeCtx); err != nil {
|
||||
log.Warn("ollama probe failed", "url", cfg.ollamaURL, "err", err)
|
||||
// Probe the enabled upstreams once; failure is logged, not fatal.
|
||||
probeCtx, probeCancel := context.WithTimeout(ctx, cfg.ProbeTimeout)
|
||||
if ollamaClient != nil {
|
||||
if err := ollamaClient.Probe(probeCtx); err != nil {
|
||||
log.Warn("ollama probe failed", "url", cfg.OllamaURL, "err", err)
|
||||
}
|
||||
}
|
||||
if err := comfyClient.Probe(probeCtx); err != nil {
|
||||
log.Warn("comfy probe failed", "url", cfg.comfyURL, "err", err)
|
||||
if comfyClient != nil {
|
||||
if err := comfyClient.Probe(probeCtx); err != nil {
|
||||
log.Warn("comfy probe failed", "url", cfg.ComfyURL, "err", err)
|
||||
}
|
||||
}
|
||||
probeCancel()
|
||||
|
||||
ollamaSrv := &http.Server{Addr: cfg.listenOllama, Handler: srv.OllamaHandler()}
|
||||
comfySrv := &http.Server{Addr: cfg.listenComfy, Handler: srv.ComfyHandler()}
|
||||
// Bind the listeners up front so a port conflict fails fast and the
|
||||
// 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)
|
||||
go func() { errCh <- ollamaSrv.ListenAndServe() }()
|
||||
go func() { errCh <- comfySrv.ListenAndServe() }()
|
||||
errCh := make(chan error, len(servers))
|
||||
for i := range servers {
|
||||
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)
|
||||
defer stop()
|
||||
// Tell systemd we are up and start the watchdog pings; both are no-ops
|
||||
// when not running under a notify/watchdog unit.
|
||||
service.NotifyReady()
|
||||
service.StartWatchdog(ctx)
|
||||
|
||||
if cfg.AutoUpdate {
|
||||
go updateLoop(ctx, cfg, log, lk, isService)
|
||||
}
|
||||
|
||||
select {
|
||||
case err := <-errCh:
|
||||
if err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||
log.Error("listener failed", "err", err)
|
||||
os.Exit(1)
|
||||
return err
|
||||
}
|
||||
case <-sigCtx.Done():
|
||||
case <-ctx.Done():
|
||||
log.Info("shutting down")
|
||||
}
|
||||
|
||||
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), cfg.shutdownTimeout)
|
||||
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), cfg.ShutdownTimeout)
|
||||
defer shutdownCancel()
|
||||
ollamaSrv.Shutdown(shutdownCtx)
|
||||
comfySrv.Shutdown(shutdownCtx)
|
||||
for _, s := range servers {
|
||||
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, 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):
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,9 +6,11 @@
|
||||
# ComfyUI --listen 0.0.0.0 --port 8189).
|
||||
services:
|
||||
gpu-turnstile:
|
||||
image: git.rambossek.at/public/gpu-turnstile:v0.1.2
|
||||
image: git.rambossek.at/public/gpu-turnstile:v0.1.6
|
||||
restart: unless-stopped
|
||||
environment:
|
||||
# Each consumer is enabled by setting its URL; leave one unset to
|
||||
# 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
|
||||
@@ -18,6 +20,9 @@ services:
|
||||
# UNLOAD_TIMEOUT: 60s
|
||||
# JOB_TIMEOUT: 15m
|
||||
# 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
|
||||
# LOG_LEVEL: info
|
||||
# LOG_FORMAT: json
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
module gpu-turnstile
|
||||
|
||||
go 1.23
|
||||
|
||||
require golang.org/x/sys v0.29.0
|
||||
|
||||
require github.com/coreos/go-systemd/v22 v22.7.0
|
||||
|
||||
@@ -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=
|
||||
@@ -0,0 +1,225 @@
|
||||
// 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. 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",
|
||||
|
||||
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\"")
|
||||
}
|
||||
if cfg.OllamaURL == "" && cfg.ComfyURL == "" {
|
||||
return cfg, fmt.Errorf("at least one of OLLAMA_URL or COMFY_URL must be set (each URL enables its consumer)")
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"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)
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -74,6 +74,22 @@ func (l *Lock) AcquireLLM(ctx context.Context) error {
|
||||
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.
|
||||
func (l *Lock) ReleaseLLM() {
|
||||
l.mu.Lock()
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
+103
-29
@@ -47,11 +47,22 @@ type Config struct {
|
||||
// LogColor enables ANSI colors in per-request log lines. Ignored when
|
||||
// the log level is above INFO (request lines are not emitted at all).
|
||||
LogColor bool
|
||||
// LogWriter receives the colored per-request lines; nil means stderr.
|
||||
LogWriter io.Writer
|
||||
|
||||
LLMWaitTimeout time.Duration
|
||||
UnloadTimeout time.Duration
|
||||
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
|
||||
@@ -77,34 +88,51 @@ type Config struct {
|
||||
type Server struct {
|
||||
cfg Config
|
||||
log *slog.Logger
|
||||
logWriter io.Writer
|
||||
freeTimeout time.Duration
|
||||
warmTimeout time.Duration
|
||||
captureLimit int64
|
||||
backoffInitial time.Duration
|
||||
backoffMax time.Duration
|
||||
busyMode string
|
||||
busyStatus int
|
||||
busyRetryAfter int
|
||||
|
||||
ollamaProxy *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) {
|
||||
ollamaURL, err := url.Parse(cfg.OllamaURL)
|
||||
if err != nil || ollamaURL.Scheme == "" || ollamaURL.Host == "" {
|
||||
return nil, fmt.Errorf("invalid OLLAMA_URL %q", cfg.OllamaURL)
|
||||
if cfg.OllamaURL == "" && cfg.ComfyURL == "" {
|
||||
return nil, fmt.Errorf("at least one of OllamaURL or ComfyURL is required")
|
||||
}
|
||||
comfyURL, err := url.Parse(cfg.ComfyURL)
|
||||
if err != nil || comfyURL.Scheme == "" || comfyURL.Host == "" {
|
||||
return nil, fmt.Errorf("invalid COMFY_URL %q", cfg.ComfyURL)
|
||||
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)
|
||||
}
|
||||
ollamaURL = u
|
||||
}
|
||||
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)
|
||||
}
|
||||
comfyURL = u
|
||||
}
|
||||
log := cfg.Log
|
||||
if log == nil {
|
||||
log = slog.Default()
|
||||
}
|
||||
if cfg.UnloadPollInterval > 0 {
|
||||
if cfg.UnloadPollInterval > 0 && cfg.Ollama != nil {
|
||||
cfg.Ollama.PollInterval = cfg.UnloadPollInterval
|
||||
}
|
||||
if cfg.HistoryPollInterval > 0 {
|
||||
if cfg.HistoryPollInterval > 0 && cfg.Comfy != nil {
|
||||
cfg.Comfy.PollInterval = cfg.HistoryPollInterval
|
||||
}
|
||||
freeTimeout := cfg.FreeTimeout
|
||||
@@ -127,23 +155,44 @@ func New(cfg Config) (*Server, error) {
|
||||
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,
|
||||
}
|
||||
return &Server{
|
||||
s := &Server{
|
||||
cfg: cfg,
|
||||
log: log,
|
||||
logWriter: cfg.LogWriter,
|
||||
freeTimeout: freeTimeout,
|
||||
warmTimeout: warmTimeout,
|
||||
captureLimit: captureLimit,
|
||||
backoffInitial: backoffInitial,
|
||||
backoffMax: backoffMax,
|
||||
ollamaProxy: newReverseProxy(ollamaURL, retry, log.With("upstream", "ollama")),
|
||||
comfyProxy: newReverseProxy(comfyURL, retry, log.With("upstream", "comfy")),
|
||||
}, 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
|
||||
}
|
||||
|
||||
// retryTransport retries requests whose failure means the upstream never
|
||||
@@ -338,7 +387,11 @@ func (s *Server) reqLine(log *slog.Logger, code, line string, attrs ...any) {
|
||||
fmt.Fprintf(&sb, " %v=%v", attrs[i], attrs[i+1])
|
||||
}
|
||||
sb.WriteByte('\n')
|
||||
os.Stderr.WriteString(sb.String())
|
||||
w := s.logWriter
|
||||
if w == nil {
|
||||
w = os.Stderr
|
||||
}
|
||||
io.WriteString(w, sb.String())
|
||||
}
|
||||
|
||||
// logRequests logs one line per incoming request and one per completed
|
||||
@@ -402,12 +455,27 @@ func (s *Server) OllamaHandler() http.Handler {
|
||||
}
|
||||
|
||||
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)
|
||||
err := s.cfg.Lock.AcquireLLM(wctx)
|
||||
cancel()
|
||||
s.cfg.Metrics.ObserveLockWait("llm", time.Since(start).Seconds())
|
||||
if 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)
|
||||
}
|
||||
return
|
||||
@@ -420,9 +488,13 @@ func (s *Server) OllamaHandler() http.Handler {
|
||||
// ComfyHandler serves the ComfyUI-facing listener.
|
||||
func (s *Server) ComfyHandler() http.Handler {
|
||||
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)
|
||||
return
|
||||
case "/metrics":
|
||||
s.writeMetrics(w)
|
||||
return
|
||||
}
|
||||
if r.Method == http.MethodPost && r.URL.Path == "/prompt" {
|
||||
s.handlePrompt(w, r)
|
||||
@@ -477,19 +549,21 @@ func (s *Server) handlePrompt(w http.ResponseWriter, r *http.Request) {
|
||||
s.cfg.Metrics.ObserveLockWait("image", time.Since(start).Seconds())
|
||||
log.Info("image lock acquired")
|
||||
|
||||
uctx, ucancel := context.WithTimeout(r.Context(), s.cfg.UnloadTimeout)
|
||||
elapsed, uerr := s.cfg.Ollama.UnloadAll(uctx)
|
||||
ucancel()
|
||||
s.cfg.Metrics.ObserveUnload(elapsed.Seconds())
|
||||
switch {
|
||||
case r.Context().Err() != nil:
|
||||
s.cfg.Lock.ReleaseImage()
|
||||
return
|
||||
case uerr != nil:
|
||||
// Degrade, don't fail the user's request on a misbehaving neighbour.
|
||||
log.Warn("ollama unload incomplete; continuing", "err", uerr)
|
||||
default:
|
||||
log.Info("ollama models unloaded", "seconds", elapsed.Seconds())
|
||||
if s.cfg.Ollama != nil {
|
||||
uctx, ucancel := context.WithTimeout(r.Context(), s.cfg.UnloadTimeout)
|
||||
elapsed, uerr := s.cfg.Ollama.UnloadAll(uctx)
|
||||
ucancel()
|
||||
s.cfg.Metrics.ObserveUnload(elapsed.Seconds())
|
||||
switch {
|
||||
case r.Context().Err() != nil:
|
||||
s.cfg.Lock.ReleaseImage()
|
||||
return
|
||||
case uerr != nil:
|
||||
// Degrade, don't fail the user's request on a misbehaving neighbour.
|
||||
log.Warn("ollama unload incomplete; continuing", "err", uerr)
|
||||
default:
|
||||
log.Info("ollama models unloaded", "seconds", elapsed.Seconds())
|
||||
}
|
||||
}
|
||||
|
||||
cw := &captureWriter{ResponseWriter: w, status: http.StatusOK, limit: s.captureLimit}
|
||||
@@ -535,7 +609,7 @@ func (s *Server) finishImageJob(promptID string) {
|
||||
s.cfg.Lock.ReleaseImage()
|
||||
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 {
|
||||
wctx, wcancel := context.WithTimeout(context.Background(), s.warmTimeout)
|
||||
if err := s.cfg.Ollama.Warm(wctx, s.cfg.WarmModel); err != nil {
|
||||
|
||||
@@ -2,6 +2,7 @@ package proxy
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
@@ -345,3 +346,187 @@ func TestPassThroughNoLock(t *testing.T) {
|
||||
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")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
@@ -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) {}
|
||||
@@ -0,0 +1,201 @@
|
||||
//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 (
|
||||
"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.
|
||||
//
|
||||
// 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) error {
|
||||
exe, err := os.Executable()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if abs, absErr := filepath.Abs(exe); absErr == nil {
|
||||
exe = abs
|
||||
}
|
||||
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 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 != "" {
|
||||
// Missing config is not fatal: the service fails fast with a
|
||||
// clear "no consumer URL" error until the user writes one.
|
||||
copyFile(configPath, etcConfig, 0o644) //nolint:errcheck // best effort
|
||||
}
|
||||
} else if configPath != "" {
|
||||
if abs, absErr := filepath.Abs(configPath); absErr == nil {
|
||||
cfg = abs
|
||||
}
|
||||
}
|
||||
if err := os.WriteFile(unitPath, []byte(renderUnit(exe, cfg)), 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 out, err := exec.Command("systemctl", "enable", "--now", Name+".service").CombinedOutput(); err != nil {
|
||||
return fmt.Errorf("systemctl enable --now: %w (%s)", err, out)
|
||||
}
|
||||
return 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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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) 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 }
|
||||
@@ -0,0 +1,398 @@
|
||||
//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 (
|
||||
"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.
|
||||
func Install(configPath string, copyBin bool) 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
|
||||
}
|
||||
}
|
||||
|
||||
installDir, _ := installDirs()
|
||||
if copyBin && !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 err := copyFile(exe, installedExe); err != nil {
|
||||
return fmt.Errorf("copy binary to %s: %w", installedExe, err)
|
||||
}
|
||||
exe = installedExe
|
||||
targetCfg := filepath.Join(installDir, "gpu-turnstile.env")
|
||||
if configPath != "" && !strings.EqualFold(configPath, targetCfg) {
|
||||
if _, statErr := os.Stat(targetCfg); os.IsNotExist(statErr) {
|
||||
copyFile(configPath, targetCfg) //nolint:errcheck // best effort
|
||||
}
|
||||
configPath = targetCfg
|
||||
}
|
||||
}
|
||||
|
||||
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",
|
||||
ServiceStartName: virtualAccount,
|
||||
})
|
||||
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)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
// 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 := s.Control(svc.Stop); err != nil {
|
||||
return false, fmt.Errorf("stop service: %w", err)
|
||||
}
|
||||
deadline := time.Now().Add(30 * time.Second)
|
||||
for {
|
||||
st, err = s.Query()
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("query service: %w", err)
|
||||
}
|
||||
if st.State == svc.Stopped {
|
||||
break
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
return false, fmt.Errorf("service did not stop within 30s")
|
||||
}
|
||||
time.Sleep(300 * time.Millisecond)
|
||||
}
|
||||
if err := s.Start(); err != nil {
|
||||
return false, fmt.Errorf("start service: %w", err)
|
||||
}
|
||||
return true, nil
|
||||
}
|
||||
@@ -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-----
|
||||
`
|
||||
@@ -0,0 +1,254 @@
|
||||
// Package update implements gpu-turnstile's self-updater: it polls the
|
||||
// Gitea releases API, downloads the Windows binary of newer releases, and
|
||||
// verifies its Ed25519 signature (produced by CI with OpenSSL) before
|
||||
// swapping it in next to the running executable.
|
||||
package update
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/ed25519"
|
||||
"crypto/sha256"
|
||||
"crypto/x509"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// maxAssetSize bounds release asset downloads.
|
||||
const maxAssetSize = 512 << 20
|
||||
|
||||
// Updater checks one Gitea repository for newer releases.
|
||||
type Updater struct {
|
||||
Repo string // e.g. https://git.rambossek.at/PUBLIC/gpu-turnstile
|
||||
Asset string // e.g. gpu-turnstile.exe
|
||||
Version string // current version, e.g. v0.1.2 ("dev" disables updates)
|
||||
Log *slog.Logger
|
||||
Client *http.Client
|
||||
}
|
||||
|
||||
type release struct {
|
||||
TagName string `json:"tag_name"`
|
||||
Assets []struct {
|
||||
Name string `json:"name"`
|
||||
BrowserDownloadURL string `json:"browser_download_url"`
|
||||
} `json:"assets"`
|
||||
}
|
||||
|
||||
func (u *Updater) logger() *slog.Logger {
|
||||
if u.Log != nil {
|
||||
return u.Log
|
||||
}
|
||||
return slog.Default()
|
||||
}
|
||||
|
||||
func (u *Updater) httpClient() *http.Client {
|
||||
if u.Client != nil {
|
||||
return u.Client
|
||||
}
|
||||
return &http.Client{Timeout: 5 * time.Minute}
|
||||
}
|
||||
|
||||
// apiURL derives <scheme>://<host>/api/v1/repos/<owner>/<name> from Repo.
|
||||
func (u *Updater) apiURL() (string, error) {
|
||||
repoURL, err := url.Parse(u.Repo)
|
||||
if err != nil || repoURL.Scheme == "" || repoURL.Host == "" {
|
||||
return "", fmt.Errorf("invalid UPDATE_REPO %q", u.Repo)
|
||||
}
|
||||
ownerName := strings.Trim(repoURL.Path, "/")
|
||||
if len(strings.Split(ownerName, "/")) != 2 {
|
||||
return "", fmt.Errorf("UPDATE_REPO %q: expected path /<owner>/<name>", u.Repo)
|
||||
}
|
||||
return fmt.Sprintf("%s://%s/api/v1/repos/%s", repoURL.Scheme, repoURL.Host, ownerName), nil
|
||||
}
|
||||
|
||||
// newerVersion reports whether latest is a higher vX.Y.Z version than
|
||||
// current. Both may carry a leading "v".
|
||||
func newerVersion(current, latest string) (bool, error) {
|
||||
parse := func(s string) ([3]int, error) {
|
||||
var v [3]int
|
||||
parts := strings.Split(strings.TrimPrefix(s, "v"), ".")
|
||||
if len(parts) != 3 {
|
||||
return v, fmt.Errorf("not a vX.Y.Z version: %q", s)
|
||||
}
|
||||
for i, p := range parts {
|
||||
n, err := strconv.Atoi(p)
|
||||
if err != nil {
|
||||
return v, fmt.Errorf("not a vX.Y.Z version: %q", s)
|
||||
}
|
||||
v[i] = n
|
||||
}
|
||||
return v, nil
|
||||
}
|
||||
cur, err := parse(current)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
lat, err := parse(latest)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
for i := 0; i < 3; i++ {
|
||||
if lat[i] != cur[i] {
|
||||
return lat[i] > cur[i], nil
|
||||
}
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func (u *Updater) get(ctx context.Context, url string) ([]byte, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := u.httpClient().Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
io.Copy(io.Discard, resp.Body)
|
||||
return nil, fmt.Errorf("GET %s: %s", url, resp.Status)
|
||||
}
|
||||
return io.ReadAll(io.LimitReader(resp.Body, maxAssetSize))
|
||||
}
|
||||
|
||||
func verifySignature(pubKeyPEM string, data, sig []byte) error {
|
||||
block, _ := pem.Decode([]byte(pubKeyPEM))
|
||||
if block == nil {
|
||||
return fmt.Errorf("invalid embedded public key PEM")
|
||||
}
|
||||
key, err := x509.ParsePKIXPublicKey(block.Bytes)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse public key: %w", err)
|
||||
}
|
||||
pub, ok := key.(ed25519.PublicKey)
|
||||
if !ok {
|
||||
return fmt.Errorf("public key is not Ed25519")
|
||||
}
|
||||
if !ed25519.Verify(pub, data, sig) {
|
||||
return fmt.Errorf("signature verification failed")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// stage swaps data into place at exePath: the running executable is
|
||||
// renamed aside (allowed on Windows) and the new file takes its name.
|
||||
func stage(exePath string, data []byte) error {
|
||||
newPath := exePath + ".new"
|
||||
oldPath := exePath + ".old"
|
||||
os.Remove(oldPath) // leftover from a previous update
|
||||
if err := os.WriteFile(newPath, data, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(exePath, oldPath); err != nil {
|
||||
os.Remove(newPath)
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(newPath, exePath); err != nil {
|
||||
os.Rename(oldPath, exePath) // roll back
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// CleanupOld removes the .old binary left behind by a staged update.
|
||||
// Call once at startup.
|
||||
func CleanupOld(exePath string) {
|
||||
os.Remove(exePath + ".old")
|
||||
os.Remove(exePath + ".new")
|
||||
}
|
||||
|
||||
// Check performs a single update check. staged is true when a newer,
|
||||
// signature-verified binary has been swapped into place at exePath; the
|
||||
// caller should then restart the process. A nil error with staged=false
|
||||
// means "no action" (up to date, disabled, or dev build); a non-nil error
|
||||
// means the check failed and the running binary is untouched.
|
||||
func (u *Updater) Check(ctx context.Context, exePath string) (staged bool, err error) {
|
||||
log := u.logger()
|
||||
if u.Version == "" || u.Version == "dev" {
|
||||
log.Debug("auto-update: dev build, skipping")
|
||||
return false, nil
|
||||
}
|
||||
if publicKeyPEM == "" {
|
||||
log.Debug("auto-update: no public key embedded, skipping")
|
||||
return false, nil
|
||||
}
|
||||
api, err := u.apiURL()
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
body, err := u.get(ctx, api+"/releases/latest")
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("fetch latest release: %w", err)
|
||||
}
|
||||
var rel release
|
||||
if err := json.Unmarshal(body, &rel); err != nil {
|
||||
return false, fmt.Errorf("parse release: %w", err)
|
||||
}
|
||||
newer, err := newerVersion(u.Version, rel.TagName)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if !newer {
|
||||
log.Debug("auto-update: up to date", "version", u.Version, "latest", rel.TagName)
|
||||
return false, nil
|
||||
}
|
||||
|
||||
urls := make(map[string]string, len(rel.Assets))
|
||||
for _, a := range rel.Assets {
|
||||
urls[a.Name] = a.BrowserDownloadURL
|
||||
}
|
||||
assetURL, ok := urls[u.Asset]
|
||||
if !ok {
|
||||
return false, fmt.Errorf("release %s has no asset %q", rel.TagName, u.Asset)
|
||||
}
|
||||
sigURL, ok := urls[u.Asset+".sig"]
|
||||
if !ok {
|
||||
return false, fmt.Errorf("release %s has no signature asset %q", rel.TagName, u.Asset+".sig")
|
||||
}
|
||||
|
||||
data, err := u.get(ctx, assetURL)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("download %s: %w", u.Asset, err)
|
||||
}
|
||||
sig, err := u.get(ctx, sigURL)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("download signature: %w", err)
|
||||
}
|
||||
|
||||
if sumURL, ok := urls[u.Asset+".sha256"]; ok {
|
||||
sumText, err := u.get(ctx, sumURL)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("download checksum: %w", err)
|
||||
}
|
||||
want := strings.Fields(string(sumText))[0]
|
||||
got := hex.EncodeToString(sha256Bytes(data))
|
||||
if !strings.EqualFold(want, got) {
|
||||
return false, fmt.Errorf("sha256 mismatch: got %s, want %s", got, want)
|
||||
}
|
||||
}
|
||||
if err := verifySignature(publicKeyPEM, data, sig); err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
if err := stage(exePath, data); err != nil {
|
||||
return false, fmt.Errorf("stage update: %w", err)
|
||||
}
|
||||
log.Info("auto-update: new version staged", "from", u.Version, "to", rel.TagName)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func sha256Bytes(data []byte) []byte {
|
||||
sum := sha256.Sum256(data)
|
||||
return sum[:]
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
package update
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/ed25519"
|
||||
"crypto/rand"
|
||||
"crypto/x509"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// fakeGitea serves a Gitea-flavored releases API with one release.
|
||||
type fakeGitea struct {
|
||||
srv *httptest.Server
|
||||
pubPEM string
|
||||
asset []byte
|
||||
tag string
|
||||
tamper bool
|
||||
noSig bool
|
||||
}
|
||||
|
||||
func newFakeGitea(t *testing.T, tag string, assetContent []byte) *fakeGitea {
|
||||
t.Helper()
|
||||
pub, priv, err := ed25519.GenerateKey(rand.Reader)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
der, err := x509.MarshalPKIXPublicKey(pub)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
f := &fakeGitea{
|
||||
pubPEM: string(pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: der})),
|
||||
asset: assetContent,
|
||||
tag: tag,
|
||||
}
|
||||
sign := func() []byte { return ed25519.Sign(priv, f.asset) }
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/api/v1/repos/o/r/releases/latest", func(w http.ResponseWriter, r *http.Request) {
|
||||
assets := []map[string]string{
|
||||
{"name": "gpu-turnstile.exe", "browser_download_url": f.srv.URL + "/dl/exe"},
|
||||
{"name": "gpu-turnstile.exe.sha256", "browser_download_url": f.srv.URL + "/dl/sha"},
|
||||
}
|
||||
if !f.noSig {
|
||||
assets = append(assets, map[string]string{"name": "gpu-turnstile.exe.sig", "browser_download_url": f.srv.URL + "/dl/sig"})
|
||||
}
|
||||
json.NewEncoder(w).Encode(map[string]any{"tag_name": f.tag, "assets": assets})
|
||||
})
|
||||
mux.HandleFunc("/dl/exe", func(w http.ResponseWriter, r *http.Request) { w.Write(f.asset) })
|
||||
mux.HandleFunc("/dl/sig", func(w http.ResponseWriter, r *http.Request) {
|
||||
sig := sign()
|
||||
if f.tamper {
|
||||
sig[0] ^= 0xff
|
||||
}
|
||||
w.Write(sig)
|
||||
})
|
||||
mux.HandleFunc("/dl/sha", func(w http.ResponseWriter, r *http.Request) {
|
||||
fmt.Fprintf(w, "%x gpu-turnstile.exe\n", sha256Bytes(f.asset))
|
||||
})
|
||||
f.srv = httptest.NewServer(mux)
|
||||
t.Cleanup(f.srv.Close)
|
||||
return f
|
||||
}
|
||||
|
||||
func (f *fakeGitea) updater(version string) *Updater {
|
||||
return &Updater{Repo: f.srv.URL + "/o/r", Asset: "gpu-turnstile.exe", Version: version}
|
||||
}
|
||||
|
||||
func fakeExe(t *testing.T) string {
|
||||
t.Helper()
|
||||
exe := filepath.Join(t.TempDir(), "gpu-turnstile.exe")
|
||||
if err := os.WriteFile(exe, []byte("old-binary"), 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return exe
|
||||
}
|
||||
|
||||
func withPublicKey(t *testing.T, pem string) {
|
||||
t.Helper()
|
||||
old := publicKeyPEM
|
||||
publicKeyPEM = pem
|
||||
t.Cleanup(func() { publicKeyPEM = old })
|
||||
}
|
||||
|
||||
func TestCheckStagesUpdate(t *testing.T) {
|
||||
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||
withPublicKey(t, f.pubPEM)
|
||||
exe := fakeExe(t)
|
||||
|
||||
staged, err := f.updater("v0.1.2").Check(context.Background(), exe)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !staged {
|
||||
t.Fatal("expected staged update")
|
||||
}
|
||||
content, _ := os.ReadFile(exe)
|
||||
if string(content) != "new-binary" {
|
||||
t.Fatalf("exe content = %q", content)
|
||||
}
|
||||
old, _ := os.ReadFile(exe + ".old")
|
||||
if string(old) != "old-binary" {
|
||||
t.Fatalf(".old content = %q", old)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckRejectsTamperedSignature(t *testing.T) {
|
||||
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||
f.tamper = true
|
||||
withPublicKey(t, f.pubPEM)
|
||||
exe := fakeExe(t)
|
||||
|
||||
staged, err := f.updater("v0.1.2").Check(context.Background(), exe)
|
||||
if err == nil {
|
||||
t.Fatal("expected signature error")
|
||||
}
|
||||
if staged {
|
||||
t.Fatal("must not stage on bad signature")
|
||||
}
|
||||
content, _ := os.ReadFile(exe)
|
||||
if string(content) != "old-binary" {
|
||||
t.Fatal("exe was modified despite bad signature")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckSkipsOlderOrEqual(t *testing.T) {
|
||||
for _, tag := range []string{"v0.1.2", "v0.1.1", "v0.0.9"} {
|
||||
f := newFakeGitea(t, tag, []byte("new-binary"))
|
||||
withPublicKey(t, f.pubPEM)
|
||||
staged, err := f.updater("v0.1.2").Check(context.Background(), fakeExe(t))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if staged {
|
||||
t.Fatalf("tag %s must not stage over v0.1.2", tag)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckSkipsWithoutPublicKey(t *testing.T) {
|
||||
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||
withPublicKey(t, "")
|
||||
staged, err := f.updater("v0.1.2").Check(context.Background(), fakeExe(t))
|
||||
if err != nil || staged {
|
||||
t.Fatalf("staged=%v err=%v, want no action without key", staged, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckSkipsDevBuild(t *testing.T) {
|
||||
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||
withPublicKey(t, f.pubPEM)
|
||||
staged, err := f.updater("dev").Check(context.Background(), fakeExe(t))
|
||||
if err != nil || staged {
|
||||
t.Fatalf("staged=%v err=%v, want no action for dev build", staged, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewerVersion(t *testing.T) {
|
||||
cases := []struct {
|
||||
cur, lat string
|
||||
want bool
|
||||
}{
|
||||
{"v0.1.2", "v0.1.3", true},
|
||||
{"0.1.2", "0.2.0", true},
|
||||
{"v0.1.2", "v1.0.0", true},
|
||||
{"v0.1.2", "v0.1.2", false},
|
||||
{"v1.2.3", "v1.2.10", true},
|
||||
{"v1.2.10", "v1.2.3", false},
|
||||
}
|
||||
for _, c := range cases {
|
||||
got, err := newerVersion(c.cur, c.lat)
|
||||
if err != nil || got != c.want {
|
||||
t.Errorf("newerVersion(%s, %s) = %v, %v; want %v", c.cur, c.lat, got, err, c.want)
|
||||
}
|
||||
}
|
||||
if _, err := newerVersion("v0.1", "v0.1.2"); err == nil {
|
||||
t.Error("expected error for malformed version")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user