Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5b752d5637 | ||
|
|
b55d548c5f | ||
|
|
2d5be72436 | ||
|
|
21e40a1774 | ||
|
|
bdc844872d | ||
|
|
299b6dc0bb | ||
|
|
9589e58ce6 | ||
|
|
be0bb36317 | ||
|
|
9b267533c9 | ||
|
|
6f092ddc12 | ||
|
|
0228ccc296 | ||
|
|
9997913929 | ||
|
|
482731ed9e | ||
|
|
e4bdc92ece | ||
|
|
14120bf4a4 | ||
|
|
d7566329ae | ||
|
|
e43ad02fc4 | ||
|
|
232f5b61f2 | ||
|
|
5e7a042cad | ||
|
|
14c2e30478 | ||
|
|
dcdb2cc72e | ||
|
|
44a8e9fbdc | ||
|
|
272a3a467d | ||
|
|
cedf28a809 | ||
|
|
e46e92bee8 | ||
|
|
cf48796183 | ||
|
|
dd3187f951 | ||
|
|
1394a76eea | ||
|
|
da02457fd5 | ||
|
|
6a6359a630 | ||
|
|
e4e4348e8d | ||
|
|
b421bb7bfb | ||
|
|
909918657f | ||
|
|
a0435c858e | ||
|
|
fbab0bba33 | ||
|
|
802a64280f | ||
|
|
a88955e35c | ||
|
|
97624470eb | ||
|
|
75f16a0229 | ||
|
|
9481fd8418 | ||
|
|
83812cebf3 | ||
|
|
30a1f55aae | ||
|
|
642cc36a39 | ||
|
|
19e19281de | ||
|
|
08d02d8fa8 | ||
|
|
d1e01f9b78 | ||
|
|
bacb26772a | ||
|
|
e363c9e4e4 | ||
|
|
5f6a22c2cf | ||
|
|
42e1386811 | ||
|
|
20d5439b5f | ||
|
|
3aaa5d80a9 | ||
|
|
cef845c2e3 | ||
|
|
cf9be55a6c | ||
|
|
0f950f3134 | ||
|
|
73fb8ac3ad |
+60
-3
@@ -37,13 +37,19 @@ jobs:
|
|||||||
exit 1
|
exit 1
|
||||||
fi
|
fi
|
||||||
|
|
||||||
|
- name: Compute lowercase image name
|
||||||
|
id: meta
|
||||||
|
run: |
|
||||||
|
REPO=git.rambossek.at/$(echo "${{ gitea.repository }}" | tr '[:upper:]' '[:lower:]')
|
||||||
|
echo "image=$REPO" >> "$GITHUB_OUTPUT"
|
||||||
|
|
||||||
- uses: docker/setup-buildx-action@v3
|
- uses: docker/setup-buildx-action@v3
|
||||||
|
|
||||||
- uses: docker/login-action@v3
|
- uses: docker/login-action@v3
|
||||||
with:
|
with:
|
||||||
registry: git.rambossek.at
|
registry: git.rambossek.at
|
||||||
username: ${{ gitea.actor }}
|
username: ${{ gitea.actor }}
|
||||||
password: ${{ secrets.GITEA_TOKEN }}
|
password: ${{ secrets.REGISTRY_TOKEN }}
|
||||||
|
|
||||||
- uses: docker/build-push-action@v6
|
- uses: docker/build-push-action@v6
|
||||||
with:
|
with:
|
||||||
@@ -52,5 +58,56 @@ jobs:
|
|||||||
build-args: |
|
build-args: |
|
||||||
VERSION=${{ gitea.ref_name }}
|
VERSION=${{ gitea.ref_name }}
|
||||||
tags: |
|
tags: |
|
||||||
git.rambossek.at/${{ gitea.repository }}:${{ gitea.ref_name }}
|
${{ steps.meta.outputs.image }}:${{ gitea.ref_name }}
|
||||||
git.rambossek.at/${{ gitea.repository }}:latest
|
${{ steps.meta.outputs.image }}:latest
|
||||||
|
|
||||||
|
# On version tags: build the signed Windows binary and attach it (plus
|
||||||
|
# signature and checksum) to a Gitea release for the auto-updater.
|
||||||
|
release:
|
||||||
|
if: gitea.ref_type == 'tag'
|
||||||
|
needs: test
|
||||||
|
runs-on: ubuntu-latest
|
||||||
|
steps:
|
||||||
|
- uses: actions/checkout@v4
|
||||||
|
- uses: actions/setup-go@v5
|
||||||
|
with:
|
||||||
|
go-version: "1.23"
|
||||||
|
|
||||||
|
- name: Build Windows binary
|
||||||
|
run: |
|
||||||
|
GOOS=windows GOARCH=amd64 CGO_ENABLED=0 go build \
|
||||||
|
-ldflags="-s -w -X main.version=${{ gitea.ref_name }}" \
|
||||||
|
-o gpu-turnstile.exe ./cmd/gpu-turnstile
|
||||||
|
|
||||||
|
- name: Sign and checksum
|
||||||
|
env:
|
||||||
|
RELEASE_SIGNING_KEY: ${{ secrets.RELEASE_SIGNING_KEY }}
|
||||||
|
run: |
|
||||||
|
# The key comes via the environment so its value never appears in
|
||||||
|
# the echoed command line of the run log.
|
||||||
|
printf '%s\n' "$RELEASE_SIGNING_KEY" > key.pem
|
||||||
|
chmod 600 key.pem
|
||||||
|
openssl pkeyutl -sign -inkey key.pem -rawin \
|
||||||
|
-in gpu-turnstile.exe -out gpu-turnstile.exe.sig
|
||||||
|
rm -f key.pem
|
||||||
|
sha256sum gpu-turnstile.exe > gpu-turnstile.exe.sha256
|
||||||
|
|
||||||
|
- name: Create release and upload assets
|
||||||
|
env:
|
||||||
|
TOKEN: ${{ secrets.GITEA_TOKEN }}
|
||||||
|
API: https://git.rambossek.at/api/v1/repos/${{ gitea.repository }}
|
||||||
|
TAG: ${{ gitea.ref_name }}
|
||||||
|
run: |
|
||||||
|
set -e
|
||||||
|
ID=$(curl -sf -H "Authorization: token $TOKEN" "$API/releases/tags/$TAG" | jq -r .id || true)
|
||||||
|
if [ -z "$ID" ] || [ "$ID" = "null" ]; then
|
||||||
|
ID=$(curl -sf -X POST -H "Authorization: token $TOKEN" \
|
||||||
|
-H "Content-Type: application/json" \
|
||||||
|
-d "{\"tag_name\":\"$TAG\",\"name\":\"$TAG\"}" \
|
||||||
|
"$API/releases" | jq -r .id)
|
||||||
|
fi
|
||||||
|
for f in gpu-turnstile.exe gpu-turnstile.exe.sig gpu-turnstile.exe.sha256; do
|
||||||
|
curl -sf -X POST -H "Authorization: token $TOKEN" \
|
||||||
|
-F "attachment=@$f" "$API/releases/$ID/assets?name=$f" > /dev/null
|
||||||
|
echo "uploaded $f"
|
||||||
|
done
|
||||||
|
|||||||
@@ -1,2 +1,9 @@
|
|||||||
/gpu-turnstile
|
/gpu-turnstile
|
||||||
/gpu-turnstile.exe
|
/gpu-turnstile.exe
|
||||||
|
/compose.yaml
|
||||||
|
/compose.yml
|
||||||
|
/signing/
|
||||||
|
|
||||||
|
/gpu-turnstile.exe.old
|
||||||
|
/gpu-turnstile.exe.new
|
||||||
|
/gpu-turnstile.exe~
|
||||||
|
|||||||
+1
-1
@@ -1,6 +1,6 @@
|
|||||||
FROM golang:1.23 AS build
|
FROM golang:1.23 AS build
|
||||||
WORKDIR /src
|
WORKDIR /src
|
||||||
COPY go.mod ./
|
COPY go.mod go.sum ./
|
||||||
COPY cmd ./cmd
|
COPY cmd ./cmd
|
||||||
COPY internal ./internal
|
COPY internal ./internal
|
||||||
ARG VERSION=dev
|
ARG VERSION=dev
|
||||||
|
|||||||
@@ -2,9 +2,10 @@
|
|||||||
|
|
||||||
GPU arbitration proxy for Ollama + ComfyUI. One consumer GPU is shared by an
|
GPU arbitration proxy for Ollama + ComfyUI. One consumer GPU is shared by an
|
||||||
LLM server (Ollama) and an image generator (ComfyUI); gpu-turnstile sits in
|
LLM server (Ollama) and an image generator (ComfyUI); gpu-turnstile sits in
|
||||||
front of both and guarantees the GPU is always in exactly one of three states:
|
front of both and guarantees the GPU is always in exactly one of four states:
|
||||||
`idle`, `llm` (N ≥ 1 Ollama requests in flight), or `image` (exactly one
|
`idle`, `llm` (N ≥ 1 Ollama requests in flight), `image` (exactly one
|
||||||
ComfyUI job, Ollama models unloaded). See [SPEC.md](SPEC.md) for the full
|
ComfyUI job, Ollama models unloaded), or `external` (a foreign process such
|
||||||
|
as a game holds the GPU). See [SPEC.md](SPEC.md) for the full
|
||||||
design.
|
design.
|
||||||
|
|
||||||
gpu-turnstile listens on the ports the services normally use; the actual
|
gpu-turnstile listens on the ports the services normally use; the actual
|
||||||
@@ -18,7 +19,10 @@ Open WebUI / n8n ────► :8188 ───┘
|
|||||||
|
|
||||||
- LLM endpoints (`/api/generate`, `/api/chat`, `/api/embed`, `/v1/*`) take the
|
- LLM endpoints (`/api/generate`, `/api/chat`, `/api/embed`, `/v1/*`) take the
|
||||||
LLM lock: concurrent requests allowed, but blocked while an image job is
|
LLM lock: concurrent requests allowed, but blocked while an image job is
|
||||||
active or waiting (image priority).
|
active or waiting (image priority). Blocked requests either hang until the
|
||||||
|
lock is free (`LLM_BUSY_MODE=wait`, default) or fail immediately with 503
|
||||||
|
(or 429) + `Retry-After` (`LLM_BUSY_MODE=reject`) — the latter lets routers
|
||||||
|
like LiteLLM cool down and retry instead of holding a hung connection.
|
||||||
- `POST /prompt` on the ComfyUI listener takes the image lock: new LLM
|
- `POST /prompt` on the ComfyUI listener takes the image lock: new LLM
|
||||||
requests block, in-flight LLMs drain, Ollama models are unloaded, the prompt
|
requests block, in-flight LLMs drain, Ollama models are unloaded, the prompt
|
||||||
is forwarded, and the lock is held until the job finishes and ComfyUI frees
|
is forwarded, and the lock is held until the job finishes and ComfyUI frees
|
||||||
@@ -26,39 +30,217 @@ Open WebUI / n8n ────► :8188 ───┘
|
|||||||
- Everything else (including websockets and all streaming) passes through
|
- Everything else (including websockets and all streaming) passes through
|
||||||
transparently and unbuffered.
|
transparently and unbuffered.
|
||||||
|
|
||||||
|
Each consumer is enabled by setting its URL (`OLLAMA_URL`, `COMFY_URL`) and
|
||||||
|
disabled by leaving it empty — at least one is required. With only Ollama
|
||||||
|
the proxy is a pass-through (no image jobs can arrive); with only ComfyUI
|
||||||
|
the Ollama unload/warm steps are skipped. A third, URL-less consumer —
|
||||||
|
detection of foreign GPU holders such as games — is enabled by `GAME_PROCS`
|
||||||
|
and/or `GPU_FOREIGN_VRAM_MB` (see below).
|
||||||
|
|
||||||
## Configuration
|
## Configuration
|
||||||
|
|
||||||
All configuration is via environment variables; invalid values fail at
|
Configuration comes from environment variables and/or an `.env`-style
|
||||||
startup.
|
config file (`KEY=VALUE` lines, `#` comments). File lookup order:
|
||||||
|
`-config <path>` flag, then `GPU_TURNSTILE_CONFIG`, then
|
||||||
|
`gpu-turnstile.env` next to the executable. Process environment variables
|
||||||
|
override file values. Invalid values fail at startup.
|
||||||
|
|
||||||
| Var | Default | Meaning |
|
| Var | Default | Meaning |
|
||||||
|---|---|---|
|
|---|---|---|
|
||||||
| `LISTEN_OLLAMA` | `:11434` | Ollama-facing listener |
|
| `LISTEN_OLLAMA` | `:11434` | Listener for Ollama-compatible clients |
|
||||||
| `LISTEN_COMFY` | `:8188` | ComfyUI-facing listener |
|
| `LISTEN_COMFY` | `:8188` | Listener for ComfyUI clients |
|
||||||
| `OLLAMA_URL` | `http://127.0.0.1:11435` | Ollama upstream |
|
| `OLLAMA_URL` | _(empty = disabled)_ | Ollama upstream; set to enable the Ollama consumer |
|
||||||
| `COMFY_URL` | `http://127.0.0.1:8189` | ComfyUI upstream |
|
| `COMFY_URL` | _(empty = disabled)_ | ComfyUI upstream; set to enable the ComfyUI consumer |
|
||||||
| `UNLOAD_TIMEOUT` | `60s` | Wait for Ollama to unload before an image job |
|
| `UNLOAD_TIMEOUT` | `60s` | Wait for Ollama to unload before an image job |
|
||||||
| `JOB_TIMEOUT` | `15m` | Wait for a ComfyUI job to finish |
|
| `JOB_TIMEOUT` | `15m` | Wait for a ComfyUI job to finish |
|
||||||
| `LLM_WAIT_TIMEOUT` | `10m` | Max lock wait for an LLM request before 503 |
|
| `LLM_WAIT_TIMEOUT` | `10m` | Max lock wait for an LLM request before 503 (wait mode) |
|
||||||
|
| `LLM_BUSY_MODE` | `wait` | `wait` = hold blocked LLM requests; `reject` = fail them immediately |
|
||||||
|
| `LLM_BUSY_STATUS` | `503` | HTTP status for rejected LLM requests in reject mode (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) |
|
| `WARM_MODEL` | _(empty)_ | Model to reload after an image job (off by default) |
|
||||||
| `LOG_LEVEL` | `info` | `debug` logs every lock transition |
|
| `COMFY_CMD` | _(derived from `COMFY_DIR`; both empty = unmanaged)_ | Supervise ComfyUI: start on demand, stop when idle to free VRAM. Requires `COMFY_URL` |
|
||||||
|
| `COMFY_DIR` | _(empty)_ | Standard venv install root: set alone to supervise ComfyUI with the derived command (`.venv` + `main.py`); also the working directory for `COMFY_CMD` |
|
||||||
|
| `COMFY_IDLE_TIMEOUT` | `5m` | Stop the managed ComfyUI after this long idle |
|
||||||
|
| `COMFY_START_TIMEOUT` | `2m` | Max wait for the managed ComfyUI to come up |
|
||||||
|
| `GAME_PROCS` | _(empty = disabled)_ | Process names (comma-separated); while any runs, the GPU counts as held: requests wait, Ollama unloads, managed ComfyUI stops |
|
||||||
|
| `GPU_FOREIGN_VRAM_MB` | `0` (disabled) | Also treat the GPU as held when a non-ignored process uses more VRAM than this (needs nvidia-smi) |
|
||||||
|
| `GPU_IGNORE_PROCS` | `ollama,ollama app,ollama_llama_server,python,pythonw` | Process names never counted as foreign GPU users |
|
||||||
|
| `GAME_POLL_INTERVAL` | `15s` | How often game/VRAM detection runs (don't go below ~10s — nvidia-smi polls keep the GPU awake) |
|
||||||
|
| `LOGLEVEL` | `warn` | `info` logs every request (colored arrows in text mode), `debug` adds lock transitions. `LOG_LEVEL` works as an alias |
|
||||||
| `LOG_FORMAT` | `text` | `json` for structured JSON logs |
|
| `LOG_FORMAT` | `text` | `json` for structured JSON logs |
|
||||||
|
| `LOG_FILE` | _(empty)_ | Append logs to this file instead of stderr |
|
||||||
|
| `UNLOAD_POLL_INTERVAL` | `500ms` | `/api/ps` poll interval while unloading |
|
||||||
|
| `HISTORY_POLL_INTERVAL` | `1s` | `/history/<id>` poll interval while a job runs |
|
||||||
|
| `PROBE_TIMEOUT` | `5s` | Startup probe of both upstreams (also per-probe health check timeout) |
|
||||||
|
| `HEALTH_INTERVAL` | `30s` | Periodic upstream probe; down/recovered changes are logged |
|
||||||
|
| `FREE_TIMEOUT` | `30s` | `POST /free` call after an image job |
|
||||||
|
| `WARM_TIMEOUT` | `2m` | Warm-model reload after an image job |
|
||||||
|
| `SHUTDOWN_TIMEOUT` | `10s` | Graceful shutdown on SIGINT/SIGTERM |
|
||||||
|
| `BACKOFF_INITIAL` | `1s` | First retry wait when an upstream refuses a connection |
|
||||||
|
| `BACKOFF_MAX` | `60s` | Cap for the exponential retry backoff |
|
||||||
|
| `PROMPT_CAPTURE_LIMIT` | `65536` | Bytes of the `/prompt` response buffered to find `prompt_id` (pass-through is unaffected) |
|
||||||
|
| `AUTO_UPDATE` | `true` | Poll the Gitea releases API for signed updates |
|
||||||
|
| `UPDATE_INTERVAL` | `6h` | Auto-update check interval |
|
||||||
|
| `UPDATE_REPO` | `https://git.rambossek.at/PUBLIC/gpu-turnstile` | Repository checked for releases |
|
||||||
|
| `UPDATE_ASSET` | `gpu-turnstile.exe` | Release asset to download |
|
||||||
|
| `APP_VER` | `stable` | `dev` disables updates, `stable` tracks latest, or pin an exact `vX.Y.Z` |
|
||||||
|
|
||||||
## Observability
|
## Observability
|
||||||
|
|
||||||
- `GET /healthz` (both listeners): `{"state":"idle|llm|image","llm_inflight":N,"image_pending":B}`
|
- `GET /healthz` (both listeners): `{"state":"idle|llm|image|external","llm_inflight":N,"image_pending":B}`
|
||||||
- `GET /metrics` (Ollama listener): Prometheus text format — `gpu_turnstile_state`,
|
- `GET /metrics` (both listeners): Prometheus text format — `gpu_turnstile_state`,
|
||||||
`gpu_turnstile_llm_inflight`, `gpu_turnstile_image_pending`,
|
`gpu_turnstile_llm_inflight`, `gpu_turnstile_image_pending`,
|
||||||
`gpu_turnstile_image_jobs_total`, `gpu_turnstile_lock_wait_seconds`
|
`gpu_turnstile_image_jobs_total`, `gpu_turnstile_lock_wait_seconds`
|
||||||
(histogram, `kind="llm|image"`), `gpu_turnstile_unload_seconds`.
|
(histogram, `kind="llm|image"`), `gpu_turnstile_unload_seconds`.
|
||||||
|
- Logs: startup logs the version and every setting (visible even at the
|
||||||
|
default `warn` level). With `LOGLEVEL=info` or `debug`, every request
|
||||||
|
logs a `-->` incoming line and a `<--` response line with status and
|
||||||
|
duration — ANSI-colored (cyan incoming; green/yellow/red by status
|
||||||
|
class) in text mode, which renders in `docker compose logs` on Windows
|
||||||
|
Terminal. Set `NO_COLOR` to disable colors.
|
||||||
|
|
||||||
|
## Managed ComfyUI (`COMFY_CMD` / `COMFY_DIR`)
|
||||||
|
|
||||||
|
Don't want ComfyUI running 24/7 (it holds VRAM even when idle — and the
|
||||||
|
Desktop app kills its server when you close it)? gpu-turnstile can supervise
|
||||||
|
it: the first request starts it, it stops again after `COMFY_IDLE_TIMEOUT`
|
||||||
|
(default 5m) without work, freeing the GPU for games or the LLM.
|
||||||
|
|
||||||
|
The easy way — point `COMFY_DIR` at a standard venv install (a folder with
|
||||||
|
`.venv` and `main.py`, or `.venv` and `ComfyUI\main.py`) and the launch
|
||||||
|
command is derived from it, including `--port` from `COMFY_URL`:
|
||||||
|
|
||||||
|
```
|
||||||
|
COMFY_URL=http://127.0.0.1:8189
|
||||||
|
COMFY_DIR=C:\ComfyUI
|
||||||
|
```
|
||||||
|
|
||||||
|
For other layouts, spell the command out yourself (run it once manually to
|
||||||
|
confirm it works):
|
||||||
|
|
||||||
|
```
|
||||||
|
COMFY_URL=http://127.0.0.1:8189
|
||||||
|
COMFY_CMD="C:\ComfyUI\.venv\Scripts\python.exe" main.py --port 8189
|
||||||
|
COMFY_DIR=C:\ComfyUI
|
||||||
|
```
|
||||||
|
|
||||||
|
`POST /prompt` waits for the server to answer before taking the GPU lock
|
||||||
|
(LLM traffic keeps flowing while torch loads); crashes are logged and the
|
||||||
|
next request respawns. Shutting gpu-turnstile down stops the child too.
|
||||||
|
|
||||||
|
Running the ComfyUI Desktop app alongside is safe: if something already
|
||||||
|
answers on the port, gpu-turnstile just uses it instead of spawning
|
||||||
|
(and never kills it — it only ever stops its own child). If the managed
|
||||||
|
instance already holds the port when you open the desktop app, the
|
||||||
|
desktop's server is the one that fails to bind.
|
||||||
|
|
||||||
|
One catch when gpu-turnstile runs as a service: the sandboxed service
|
||||||
|
account may not enter your user profile, so a ComfyUI install under
|
||||||
|
`C:\Users\...` (or `/home/...`) fails with "Access is denied".
|
||||||
|
`--install-service` fixes that automatically — it grants
|
||||||
|
`NT SERVICE\gpu-turnstile` recursive access to `COMFY_DIR` on Windows and
|
||||||
|
adds a `BindPaths=` to the systemd unit on Linux. Re-run it after changing
|
||||||
|
`COMFY_DIR`; or grant by hand from an admin shell:
|
||||||
|
`icacls "<COMFY_DIR>" /grant "NT SERVICE\gpu-turnstile:(OI)(CI)M" /T`.
|
||||||
|
|
||||||
|
## Game detection
|
||||||
|
|
||||||
|
Want to game on the same GPU without Ollama/ComfyUI squatting on the VRAM?
|
||||||
|
gpu-turnstile can watch for foreign GPU holders and, while one is active,
|
||||||
|
make LLM/image requests wait (or 503, per `LLM_BUSY_MODE`), unload Ollama's
|
||||||
|
models and stop the managed ComfyUI so the game gets the memory. Two
|
||||||
|
detection paths, each optional, polled every `GAME_POLL_INTERVAL` (15s):
|
||||||
|
|
||||||
|
```
|
||||||
|
GAME_PROCS=cyberpunk2077.exe,bg3.exe # the reliable way on Windows
|
||||||
|
GPU_FOREIGN_VRAM_MB=1024 # catch-all via nvidia-smi
|
||||||
|
```
|
||||||
|
|
||||||
|
`GAME_PROCS` matches running process names (case-insensitive, `.exe`
|
||||||
|
optional). `GPU_FOREIGN_VRAM_MB` asks nvidia-smi which processes hold GPU
|
||||||
|
memory and treats anything not in `GPU_IGNORE_PROCS` above the threshold as
|
||||||
|
foreign — handy as a catch-all, but note that under Windows' WDDM driver
|
||||||
|
graphics-only games may not show up in nvidia-smi's per-process list, so
|
||||||
|
name your games in `GAME_PROCS` there; on Linux both paths work. When the
|
||||||
|
game exits, requests resume automatically.
|
||||||
|
|
||||||
## Build and run
|
## Build and run
|
||||||
|
|
||||||
```sh
|
```sh
|
||||||
go build ./cmd/gpu-turnstile
|
go build ./cmd/gpu-turnstile
|
||||||
./gpu-turnstile
|
OLLAMA_URL=http://127.0.0.1:11435 COMFY_URL=http://127.0.0.1:8189 ./gpu-turnstile
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Running the binary with no arguments in a terminal prints the help screen
|
||||||
|
(same as `-h`/`--help`); without a terminal (services, containers) a bare
|
||||||
|
invocation starts the proxy.
|
||||||
|
|
||||||
|
### Run natively on Windows (current primary deployment)
|
||||||
|
|
||||||
|
Download `gpu-turnstile.exe` from a release and install it as a Windows
|
||||||
|
service — no admin shell needed, a UAC prompt appears automatically and
|
||||||
|
the elevated child does the work (its window waits for Enter so you can
|
||||||
|
read the result):
|
||||||
|
|
||||||
|
```sh
|
||||||
|
gpu-turnstile.exe --install-service # installs into Program Files, auto-start
|
||||||
|
gpu-turnstile.exe --install-service --no-copy # register in place instead
|
||||||
|
gpu-turnstile.exe --remove-service
|
||||||
|
```
|
||||||
|
|
||||||
|
Layout: `C:\Program Files\gpu-turnstile\` holds the exe and
|
||||||
|
`gpu-turnstile.env`, logs go to `C:\ProgramData\gpu-turnstile\`. If you
|
||||||
|
install without a config, the installer writes a sample env file with every
|
||||||
|
setting commented and explained — only `LOG_FILE` is active (a service has
|
||||||
|
no console). Your own `LOG_FILE` setting is always kept. The service always
|
||||||
|
runs
|
||||||
|
as the virtual account `NT SERVICE\gpu-turnstile` (low-privilege,
|
||||||
|
per-service, no password); the installer automatically grants it write
|
||||||
|
access to the install and data directories — nothing else to do.
|
||||||
|
|
||||||
|
Re-running `--install-service` is safe: it stops a running service,
|
||||||
|
replaces the installed binary only if it changed, fixes the registration
|
||||||
|
only where it drifted, and restarts the service only if it was running.
|
||||||
|
|
||||||
|
### Run natively on Linux (systemd)
|
||||||
|
|
||||||
|
The same binary works on Linux. Install it as a systemd service as root:
|
||||||
|
|
||||||
|
```sh
|
||||||
|
gpu-turnstile --install-service # installs into /var/lib/gpu-turnstile, enables + starts
|
||||||
|
gpu-turnstile --install-service --no-copy # register in place instead
|
||||||
|
gpu-turnstile --remove-service
|
||||||
|
```
|
||||||
|
|
||||||
|
The unit (`/etc/systemd/system/gpu-turnstile.service`) is `Type=notify`:
|
||||||
|
`systemctl start` blocks until the listeners are actually bound, a 30 s
|
||||||
|
watchdog restarts the process if it wedges, and logs land in the journal
|
||||||
|
(`journalctl -u gpu-turnstile -f`) unless `LOG_FILE` is set. Install
|
||||||
|
copies the binary to `/var/lib/gpu-turnstile/` and the config to
|
||||||
|
`/etc/gpu-turnstile.env` (edit that one after installing). The service
|
||||||
|
runs sandboxed with `DynamicUser=yes` — a transient low-privilege UID,
|
||||||
|
read-only filesystem except its install dir (so self-update keeps
|
||||||
|
working), no capabilities, syscall-filtered: same least-privilege idea as
|
||||||
|
the Windows virtual account. The notify integration is a no-op in
|
||||||
|
containers and interactive shells.
|
||||||
|
|
||||||
|
**Auto-update is on by default**: the binary checks the repo's latest
|
||||||
|
release on startup and every `UPDATE_INTERVAL`, verifies the Ed25519
|
||||||
|
signature of the download against the public key embedded at build time,
|
||||||
|
and — once the GPU lock is idle — restarts the service onto the new
|
||||||
|
version. `APP_VER` controls the target: `dev` disables updates, `stable`
|
||||||
|
(the default) tracks the latest release, and an exact `vX.Y.Z` pins that
|
||||||
|
release (even as a downgrade or to replace a dev build). Disable entirely
|
||||||
|
with `AUTO_UPDATE=false`. Releases are signed by CI with
|
||||||
|
OpenSSL; the matching public key lives in `internal/update/pubkey.go`
|
||||||
|
(one-time setup: `openssl genpkey -algorithm ed25519 -out private.pem`,
|
||||||
|
`openssl pkey -in private.pem -pubout -out public.pem`; private key goes
|
||||||
|
to the `RELEASE_SIGNING_KEY` repo secret, public key is committed).
|
||||||
|
`gpu-turnstile --force-update` checks immediately, stages the new binary
|
||||||
|
and restarts the running service (elevating via UAC only if needed).
|
||||||
|
|
||||||
|
### Docker
|
||||||
|
|
||||||
```sh
|
```sh
|
||||||
docker build -t gpu-turnstile .
|
docker build -t gpu-turnstile .
|
||||||
docker run --rm -p 11434:11434 -p 8188:8188 \
|
docker run --rm -p 11434:11434 -p 8188:8188 \
|
||||||
@@ -69,9 +251,14 @@ docker run --rm -p 11434:11434 -p 8188:8188 \
|
|||||||
|
|
||||||
Releases are built by Gitea Actions (`.gitea/workflows/ci.yml`): every push
|
Releases are built by Gitea Actions (`.gitea/workflows/ci.yml`): every push
|
||||||
runs `go vet` and `go test -race`, and pushing a semantic-version tag
|
runs `go vet` and `go test -race`, and pushing a semantic-version tag
|
||||||
`vX.Y.Z` builds and publishes
|
`vX.Y.Z` publishes the container image
|
||||||
`git.rambossek.at/<owner>/gpu-turnstile:vX.Y.Z` (and updates `:latest`).
|
(`git.rambossek.at/<owner>/gpu-turnstile:vX.Y.Z` plus `:latest`) and a
|
||||||
No images are built from branches.
|
signed Windows binary attached to a Gitea release. Nothing is built from
|
||||||
|
branches.
|
||||||
|
|
||||||
|
The registry login needs one repository secret (Settings → Actions →
|
||||||
|
Secrets): `REGISTRY_TOKEN` — an access token with `write:package` scope.
|
||||||
|
The automatic `GITEA_TOKEN` cannot push packages.
|
||||||
|
|
||||||
## Development
|
## Development
|
||||||
|
|
||||||
@@ -80,13 +267,17 @@ go vet ./...
|
|||||||
go test -race ./...
|
go test -race ./...
|
||||||
```
|
```
|
||||||
|
|
||||||
Stdlib only, Go 1.23+. Layout:
|
Go 1.23+; the only external dependency is `golang.org/x/sys` (Windows
|
||||||
|
service integration, unused in the Linux build). Layout:
|
||||||
|
|
||||||
```
|
```
|
||||||
cmd/gpu-turnstile/main.go wiring, config, listeners
|
cmd/gpu-turnstile/main.go wiring, config, listeners, service + updater
|
||||||
internal/lock/ two-mode lock (LLM readers / image writer, FIFO)
|
internal/lock/ two-mode lock (LLM readers / image writer, FIFO)
|
||||||
internal/ollama/ ps / unload / warm client
|
internal/ollama/ ps / unload / warm client
|
||||||
internal/comfy/ history poll / free client
|
internal/comfy/ history poll / free client
|
||||||
internal/proxy/ handlers for both listeners
|
internal/proxy/ handlers for both listeners
|
||||||
internal/metrics/ Prometheus exposition, no dependencies
|
internal/metrics/ Prometheus exposition, no dependencies
|
||||||
|
internal/config/ env + .env file configuration
|
||||||
|
internal/update/ signed auto-updater
|
||||||
|
internal/service/ Windows service integration
|
||||||
```
|
```
|
||||||
|
|||||||
@@ -10,11 +10,12 @@ becomes very slow; on Linux it would OOM instead.
|
|||||||
## Goal
|
## Goal
|
||||||
|
|
||||||
A single Go binary that sits in front of **both** services and guarantees that
|
A single Go binary that sits in front of **both** services and guarantees that
|
||||||
at any moment the GPU is in exactly one of three states:
|
at any moment the GPU is in exactly one of four states:
|
||||||
|
|
||||||
- `idle` — nothing in flight
|
- `idle` — nothing in flight
|
||||||
- `llm` — N ≥ 1 Ollama inference requests in flight (concurrency allowed)
|
- `llm` — N ≥ 1 Ollama inference requests in flight (concurrency allowed)
|
||||||
- `image` — exactly one ComfyUI job in flight, Ollama models unloaded
|
- `image` — exactly one ComfyUI job in flight, Ollama models unloaded
|
||||||
|
- `external` — a foreign process (e.g. a game) holds the GPU; new work waits
|
||||||
|
|
||||||
Clients (LiteLLM, Open WebUI, n8n) point at gpu-turnstile instead of at the
|
Clients (LiteLLM, Open WebUI, n8n) point at gpu-turnstile instead of at the
|
||||||
services. gpu-turnstile is transparent for everything that does not touch the
|
services. gpu-turnstile is transparent for everything that does not touch the
|
||||||
@@ -42,6 +43,22 @@ listener is an `httputil.ReverseProxy` to its upstream. Websocket upgrades
|
|||||||
(ComfyUI `/ws`) and streaming bodies (Ollama NDJSON / SSE) must pass through
|
(ComfyUI `/ws`) and streaming bodies (Ollama NDJSON / SSE) must pass through
|
||||||
unbuffered (`FlushInterval = -1`).
|
unbuffered (`FlushInterval = -1`).
|
||||||
|
|
||||||
|
### Modes of operation
|
||||||
|
|
||||||
|
Each GPU consumer is enabled by setting its URL and disabled by leaving it
|
||||||
|
empty — no separate flags. At least one URL must be set; a disabled
|
||||||
|
consumer gets no listener, no startup probe, and no lock participation:
|
||||||
|
|
||||||
|
- **Both set** (default deployment): full arbitration as described below.
|
||||||
|
- **Only `OLLAMA_URL`**: pure pass-through for Ollama; the LLM lock never
|
||||||
|
blocks since no image jobs can arrive.
|
||||||
|
- **Only `COMFY_URL`**: image jobs are tracked and ComfyUI's VRAM is freed
|
||||||
|
afterwards, but the Ollama unload and warm-reload steps are skipped.
|
||||||
|
- **Game detection** is a third, optional consumer without a URL: enabled by
|
||||||
|
`GAME_PROCS` and/or `GPU_FOREIGN_VRAM_MB` it watches for foreign processes
|
||||||
|
holding the GPU (see below) and plugs into the same lock the same way —
|
||||||
|
excluded when both knobs are unset.
|
||||||
|
|
||||||
### Lock semantics
|
### Lock semantics
|
||||||
|
|
||||||
Two-mode lock with image priority (writer-preferring RW lock, where "readers"
|
Two-mode lock with image priority (writer-preferring RW lock, where "readers"
|
||||||
@@ -51,10 +68,20 @@ are LLM requests and the single "writer" is an image job):
|
|||||||
`image` **or while an image job is waiting**. Then state := `llm`, n++.
|
`image` **or while an image job is waiting**. Then state := `llm`, n++.
|
||||||
On completion (response fully written, including streamed bodies, or client
|
On completion (response fully written, including streamed bodies, or client
|
||||||
disconnect) n--; if n == 0 state := `idle`.
|
disconnect) n--; if n == 0 state := `idle`.
|
||||||
|
`LLM_BUSY_MODE` selects what a blocked LLM request sees: `wait` (default)
|
||||||
|
hangs until the lock is free or `LLM_WAIT_TIMEOUT` expires (then 503 +
|
||||||
|
`Retry-After`); `reject` answers immediately with `LLM_BUSY_STATUS`
|
||||||
|
(default 503; 429 works too) + `Retry-After: BUSY_RETRY_AFTER`, which
|
||||||
|
routers like LiteLLM honor for cooldowns/retries.
|
||||||
- **Image job**: `AcquireImage()` marks "image pending" (so no new LLM
|
- **Image job**: `AcquireImage()` marks "image pending" (so no new LLM
|
||||||
requests start), waits until n == 0, sets state := `image`. Released after
|
requests start), waits until n == 0, sets state := `image`. Released after
|
||||||
the ComfyUI job finished and models were freed.
|
the ComfyUI job finished and models were freed.
|
||||||
- Concurrent image jobs queue FIFO behind each other.
|
- Concurrent image jobs queue FIFO behind each other.
|
||||||
|
- **External hold**: game detection calls `SetExternal(holder)` while a
|
||||||
|
foreign process holds the GPU. New LLM and image grants block (same busy
|
||||||
|
handling as above) until `ClearExternal()`; in-flight work is not
|
||||||
|
preempted, it drains. Snapshot reports state `external` once nothing else
|
||||||
|
is in flight.
|
||||||
- All waits are context-aware: a client that disconnects while waiting is
|
- All waits are context-aware: a client that disconnects while waiting is
|
||||||
removed from the queue.
|
removed from the queue.
|
||||||
|
|
||||||
@@ -78,15 +105,18 @@ ComfyUI listener (`:8188` → `COMFY_URL`):
|
|||||||
### Image job flow (`POST /prompt`)
|
### Image job flow (`POST /prompt`)
|
||||||
|
|
||||||
1. `AcquireImage()`.
|
1. `AcquireImage()`.
|
||||||
2. Unload Ollama: `GET /api/ps`; for each model `POST /api/generate
|
2. Unload Ollama (skipped when `OLLAMA_URL` is unset): `GET /api/ps`; for
|
||||||
|
each model `POST /api/generate
|
||||||
{"model":M,"keep_alive":0}`; if that returns non-2xx (embedding-only
|
{"model":M,"keep_alive":0}`; if that returns non-2xx (embedding-only
|
||||||
models), `POST /api/embed {"model":M,"input":"x","keep_alive":0}`. Poll
|
models), `POST /api/embed {"model":M,"input":"x","keep_alive":0}`. Poll
|
||||||
`/api/ps` every 500 ms until empty or `UNLOAD_TIMEOUT`. On timeout: log and
|
`/api/ps` every `UNLOAD_POLL_INTERVAL` (default 500 ms) until empty or
|
||||||
|
`UNLOAD_TIMEOUT`. On timeout: log and
|
||||||
continue (degrade, don't fail the user's request).
|
continue (degrade, don't fail the user's request).
|
||||||
3. Forward the original request body to ComfyUI `/prompt`, return status,
|
3. Forward the original request body to ComfyUI `/prompt`, return status,
|
||||||
headers and body to the caller unchanged, flush.
|
headers and body to the caller unchanged, flush.
|
||||||
4. If the response is 200 and contains `prompt_id`: in a goroutine, poll
|
4. If the response is 200 and contains `prompt_id`: in a goroutine, poll
|
||||||
`GET /history/<prompt_id>` every 1 s until the entry has
|
`GET /history/<prompt_id>` every `HISTORY_POLL_INTERVAL` (default 1 s)
|
||||||
|
until the entry has
|
||||||
`status.completed == true`, `status.status_str == "error"`, or
|
`status.completed == true`, `status.status_str == "error"`, or
|
||||||
`JOB_TIMEOUT`. Then `POST /free {"unload_models":true,"free_memory":true}`.
|
`JOB_TIMEOUT`. Then `POST /free {"unload_models":true,"free_memory":true}`.
|
||||||
Then release the image lock.
|
Then release the image lock.
|
||||||
@@ -98,34 +128,267 @@ state is `idle`, send `POST /api/generate {"model":WARM_MODEL,"keep_alive":-1}`
|
|||||||
with empty prompt to reload the chat model so the next chat doesn't pay the
|
with empty prompt to reload the chat model so the next chat doesn't pay the
|
||||||
load time. Off by default.
|
load time. Off by default.
|
||||||
|
|
||||||
|
### Managed ComfyUI (`COMFY_CMD` / `COMFY_DIR`)
|
||||||
|
|
||||||
|
When `COMFY_CMD` is set, gpu-turnstile runs ComfyUI as a supervised child
|
||||||
|
process instead of expecting an always-on server. Setting only `COMFY_DIR`
|
||||||
|
enables the same management with the launch command derived from the
|
||||||
|
standard venv layout under it (`.venv\Scripts\python.exe` on Windows,
|
||||||
|
`.venv/bin/python` on Linux; `ComfyUI\main.py`, or a flat `main.py` when
|
||||||
|
that is what exists; `--port` from the `COMFY_URL` port). Missing layout
|
||||||
|
files are flagged in the startup log.
|
||||||
|
|
||||||
|
- **Start on demand**: any ComfyUI request spawns it (double quotes in the
|
||||||
|
command line group arguments with spaces; `COMFY_DIR` sets the working
|
||||||
|
directory). `POST /prompt` additionally waits for readiness
|
||||||
|
(`/system_stats`) for up to `COMFY_START_TIMEOUT` *before* taking the GPU
|
||||||
|
lock, so LLM traffic flows while torch loads. Other requests are bridged
|
||||||
|
by the normal retry backoff. Spawn failure or a readiness timeout
|
||||||
|
answers 502.
|
||||||
|
- **Idle stop**: after `COMFY_IDLE_TIMEOUT` without requests or finished
|
||||||
|
jobs — and only while the GPU lock is idle — the process tree is killed,
|
||||||
|
freeing the VRAM ComfyUI holds. The next request restarts it.
|
||||||
|
- **Crash**: an unexpected exit is logged; the next request respawns.
|
||||||
|
gpu-turnstile's own shutdown stops the child too.
|
||||||
|
- **Coexistence**: before spawning, the URL is probed — if another server
|
||||||
|
already answers (e.g. the ComfyUI desktop app), it is used as-is and no
|
||||||
|
child is spawned; the idle watcher and shutdown only ever stop the
|
||||||
|
supervisor's own process, never the external one. The other direction —
|
||||||
|
starting the desktop app while the managed instance holds the port —
|
||||||
|
makes the *desktop* server fail to bind; gpu-turnstile is unaffected.
|
||||||
|
- Its stdout/stderr is forwarded to the log at INFO. The health check
|
||||||
|
skips the intentionally-stopped/starting states; a failed probe while
|
||||||
|
the process is alive and was previously ready is logged as DOWN.
|
||||||
|
- **Permissions**: the service account is sandboxed (Windows virtual
|
||||||
|
account, systemd `DynamicUser`), so a ComfyUI install inside a user
|
||||||
|
profile is off-limits by default. `--install-service` opens it up —
|
||||||
|
a recursive ACL grant for `NT SERVICE\gpu-turnstile` on Windows, a
|
||||||
|
`BindPaths=` in the unit on Linux — reading `COMFY_DIR` from the config
|
||||||
|
it installs. Re-run `--install-service` after changing `COMFY_DIR`, or
|
||||||
|
grant by hand (admin shell):
|
||||||
|
`icacls "<COMFY_DIR>" /grant "NT SERVICE\gpu-turnstile:(OI)(CI)M" /T`.
|
||||||
|
|
||||||
|
## Game detection (foreign GPU holders)
|
||||||
|
|
||||||
|
Games and other foreign GPU users sit outside the URL-based consumer model —
|
||||||
|
nothing proxies through gpu-turnstile for them. Two independent detection
|
||||||
|
paths, polled every `GAME_POLL_INTERVAL` (default 15 s); either one being
|
||||||
|
configured enables the feature:
|
||||||
|
|
||||||
|
- **Process watch list** (`GAME_PROCS`, comma-separated, case-insensitive,
|
||||||
|
`.exe` optional): while any listed process runs, the GPU counts as held.
|
||||||
|
This is the reliable path on Windows.
|
||||||
|
- **Foreign VRAM threshold** (`GPU_FOREIGN_VRAM_MB`): `nvidia-smi
|
||||||
|
--query-compute-apps` lists per-process GPU memory; any process not in
|
||||||
|
`GPU_IGNORE_PROCS` (default: Ollama and python — ComfyUI runs under python)
|
||||||
|
holding more than the threshold counts as a foreign holder. Needs
|
||||||
|
nvidia-smi on the PATH (absent: logged once, path disabled) and works best
|
||||||
|
on Linux — under Windows' WDDM driver, graphics-only games may not appear
|
||||||
|
in the per-process list.
|
||||||
|
|
||||||
|
While a holder is detected, gpu-turnstile:
|
||||||
|
|
||||||
|
1. takes the external hold on the lock (`SetExternal`), so new LLM/image
|
||||||
|
requests wait or are rejected per `LLM_BUSY_MODE`; busy responses name the
|
||||||
|
holder,
|
||||||
|
2. once in-flight work has drained, **frees VRAM for the foreign process**:
|
||||||
|
the managed ComfyUI is stopped (an external server on its port is never
|
||||||
|
touched) and Ollama's resident models are unloaded,
|
||||||
|
3. refuses to spawn the managed ComfyUI; ComfyUI requests that would need a
|
||||||
|
spawn are answered 503 + `Retry-After` (a running desktop instance keeps
|
||||||
|
being proxied),
|
||||||
|
4. logs the transitions at WARN ("GPU held by an external process …" /
|
||||||
|
"released the GPU; resuming") and reports state `external` in `/healthz`
|
||||||
|
and `/metrics`.
|
||||||
|
|
||||||
|
When the holder disappears, the hold is lifted and queued requests proceed.
|
||||||
|
|
||||||
## Configuration (env)
|
## Configuration (env)
|
||||||
|
|
||||||
|
Configuration comes from environment variables and/or an `.env`-style
|
||||||
|
config file (`KEY=VALUE` lines, `#` comments). File lookup order:
|
||||||
|
`-config <path>` flag, then `GPU_TURNSTILE_CONFIG`, then
|
||||||
|
`gpu-turnstile.env` next to the executable. Process environment variables
|
||||||
|
override file values. A missing file is fine; a malformed one is fatal.
|
||||||
|
|
||||||
| Var | Default | Meaning |
|
| Var | Default | Meaning |
|
||||||
|---|---|---|
|
|---|---|---|
|
||||||
| `LISTEN_OLLAMA` | `:11434` | Ollama-facing listener |
|
| `LISTEN_OLLAMA` | `:11434` | listener for Ollama-compatible clients |
|
||||||
| `LISTEN_COMFY` | `:8188` | ComfyUI-facing listener |
|
| `LISTEN_COMFY` | `:8188` | listener for ComfyUI clients |
|
||||||
| `OLLAMA_URL` | `http://127.0.0.1:11435` | upstream |
|
| `OLLAMA_URL` | _(empty = disabled)_ | Ollama upstream; set to enable the Ollama consumer |
|
||||||
| `COMFY_URL` | `http://127.0.0.1:8189` | upstream |
|
| `COMFY_URL` | _(empty = disabled)_ | ComfyUI upstream; set to enable the ComfyUI consumer |
|
||||||
| `UNLOAD_TIMEOUT` | `60s` | wait for Ollama to unload |
|
| `UNLOAD_TIMEOUT` | `60s` | wait for Ollama to unload |
|
||||||
| `JOB_TIMEOUT` | `15m` | wait for ComfyUI job |
|
| `JOB_TIMEOUT` | `15m` | wait for ComfyUI job |
|
||||||
| `LLM_WAIT_TIMEOUT` | `10m` | max time an LLM request waits for the lock before 503 |
|
| `LLM_WAIT_TIMEOUT` | `10m` | max time an LLM request waits for the lock before 503 (wait mode) |
|
||||||
|
| `LLM_BUSY_MODE` | `wait` | `wait` = hold blocked LLM requests; `reject` = fail them immediately |
|
||||||
|
| `LLM_BUSY_STATUS` | `503` | HTTP status for rejected LLM requests in reject mode (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 |
|
| `WARM_MODEL` | `` | optional model to reload after an image job |
|
||||||
| `LOG_LEVEL` | `info` | `debug` logs every lock transition |
|
| `COMFY_CMD` | _(derived from `COMFY_DIR`; both empty = unmanaged)_ | spawn and supervise ComfyUI on demand: first request starts it, idle stop after `COMFY_IDLE_TIMEOUT` frees its VRAM. Requires `COMFY_URL` |
|
||||||
|
| `COMFY_DIR` | `` | standard venv install root: set alone to supervise ComfyUI with the derived launch command (`.venv` + `main.py`, `--port` from `COMFY_URL`); also the working directory for `COMFY_CMD` |
|
||||||
|
| `COMFY_IDLE_TIMEOUT` | `5m` | stop the managed ComfyUI after this long without requests or jobs |
|
||||||
|
| `COMFY_START_TIMEOUT` | `2m` | how long a request waits for the managed ComfyUI to come up |
|
||||||
|
| `GAME_PROCS` | _(empty = disabled)_ | comma-separated process names (case-insensitive, `.exe` optional); while any runs, the GPU counts as held by it: requests wait, Ollama unloads, the managed ComfyUI stops |
|
||||||
|
| `GPU_FOREIGN_VRAM_MB` | `0` (disabled) | also treat the GPU as held when a process not in `GPU_IGNORE_PROCS` uses more VRAM than this; needs nvidia-smi |
|
||||||
|
| `GPU_IGNORE_PROCS` | `ollama,ollama app,ollama_llama_server,python,pythonw` | process names never counted as foreign GPU users |
|
||||||
|
| `GAME_POLL_INTERVAL` | `15s` | how often game/VRAM detection runs (nvidia-smi polls keep the GPU awake; don't go below ~10s) |
|
||||||
|
| `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 (also the per-probe health check timeout) |
|
||||||
|
| `HEALTH_INTERVAL` | `30s` | periodic probe of enabled upstreams; status changes (down/recovered) are logged |
|
||||||
|
| `FREE_TIMEOUT` | `30s` | `POST /free` call after an image job |
|
||||||
|
| `WARM_TIMEOUT` | `2m` | warm-model reload after an image job |
|
||||||
|
| `SHUTDOWN_TIMEOUT` | `10s` | graceful shutdown on SIGINT/SIGTERM |
|
||||||
|
| `BACKOFF_INITIAL` | `1s` | first retry wait when an upstream refuses a connection |
|
||||||
|
| `BACKOFF_MAX` | `60s` | cap for the exponential retry backoff |
|
||||||
|
| `PROMPT_CAPTURE_LIMIT` | `65536` | bytes of the `/prompt` response buffered to find `prompt_id` (pass-through is unaffected) |
|
||||||
|
| `AUTO_UPDATE` | `true` | poll the Gitea releases API for signed updates |
|
||||||
|
| `UPDATE_INTERVAL` | `6h` | auto-update check interval |
|
||||||
|
| `UPDATE_REPO` | `https://git.rambossek.at/PUBLIC/gpu-turnstile` | repository to check for releases |
|
||||||
|
| `UPDATE_ASSET` | `gpu-turnstile.exe` | release asset to download |
|
||||||
|
| `APP_VER` | `stable` | version to run: `dev` disables updates, `stable` tracks the latest release, or an exact `vX.Y.Z` pin (up- or downgraded to) |
|
||||||
|
| `CFG_VER` | _(installer-managed)_ | config format reference written by `--install-service` (always a concrete `vX.Y.Z`; a dev build stamps `v0.0.0`); missing = the file is replaced with a fresh sample (backup `.bak`). New settings are appended (commented out) at install and at every startup after an update changed the version |
|
||||||
|
|
||||||
Startup fails fast on unparsable values. Both upstreams are probed once at
|
Startup fails fast on unparsable values and when neither consumer URL is
|
||||||
start (`/api/version`, `/system_stats`); failure is logged, not fatal.
|
set. Enabled upstreams are probed once at start (`/api/version`,
|
||||||
|
`/system_stats`); failure is logged, not fatal.
|
||||||
|
|
||||||
|
## Native deployment (Windows and Linux)
|
||||||
|
|
||||||
|
The binary runs natively on Windows (the current primary deployment) and on
|
||||||
|
Linux with systemd (the future GPU server), as well as in Docker.
|
||||||
|
|
||||||
|
Service management is the same on both platforms:
|
||||||
|
`gpu-turnstile --install-service [-config path]` installs, registers and
|
||||||
|
starts an auto-start service; `--remove-service` stops and uninstalls it.
|
||||||
|
Both need admin/root; on Windows a non-elevated shell triggers a UAC
|
||||||
|
prompt instead of failing — the command relaunches itself elevated, waits
|
||||||
|
for the child, and mirrors its exit code. The legacy form
|
||||||
|
`gpu-turnstile service install|remove` does the same thing.
|
||||||
|
|
||||||
|
Re-running install on an already-registered service converges instead of
|
||||||
|
failing: a running service is stopped first, the installed binary copy is
|
||||||
|
refreshed only when the content differs, the registration (Windows service
|
||||||
|
config / systemd unit) is updated only where it drifted, and the service is
|
||||||
|
started again only if it was running before.
|
||||||
|
|
||||||
|
By default install creates the canonical layout and copies the binary into
|
||||||
|
it (Windows: `%ProgramFiles%\gpu-turnstile\`, plus
|
||||||
|
`%ProgramData%\gpu-turnstile\` for logs; Linux: `/var/lib/gpu-turnstile/`
|
||||||
|
with the config at `/etc/gpu-turnstile.env`). If there is no config at all,
|
||||||
|
install writes a sample env file covering every setting — each with a
|
||||||
|
comment line, everything commented out — except the `CFG_VER`/`APP_VER`
|
||||||
|
header and `LOG_FILE`, which is active on Windows
|
||||||
|
(`%ProgramData%\gpu-turnstile\gpu-turnstile.log`) since a service has no
|
||||||
|
console; on Linux it stays commented because stderr goes to the journal.
|
||||||
|
The first line is `CFG_VER=vX.Y.Z`, recording the installer version. When a
|
||||||
|
later version's install finds an older `CFG_VER`, it appends every setting
|
||||||
|
the file does not mention (commented or not) at the end and updates
|
||||||
|
`CFG_VER`; a file without `CFG_VER` is invalid and gets replaced by a fresh
|
||||||
|
sample, with the old content kept as `<file>.bak`. `--no-copy` registers
|
||||||
|
the current executable location as-is and leaves the config untouched.
|
||||||
|
|
||||||
|
### Windows
|
||||||
|
|
||||||
|
- `--install-service` creates `%ProgramFiles%\gpu-turnstile\` and
|
||||||
|
`%ProgramData%\gpu-turnstile\`, copies the exe and (if none exists there
|
||||||
|
yet) the `gpu-turnstile.env` into the Program Files directory, and
|
||||||
|
registers that copy as a Windows service; recovery actions restart it
|
||||||
|
5 s after any failure. The install ensures the env file sets `LOG_FILE`
|
||||||
|
to `%ProgramData%\gpu-turnstile\gpu-turnstile.log` since there is no
|
||||||
|
console — an existing `LOG_FILE` setting is kept.
|
||||||
|
- **Account**: the service always runs as the virtual account
|
||||||
|
`NT SERVICE\gpu-turnstile` — a per-service low-privilege identity the
|
||||||
|
SCM manages (no password, automatic logon-as-a-service right, no admin
|
||||||
|
rights, gone when the service is removed). The installer grants it
|
||||||
|
modify access to the install and data directories (self-updates rewrite
|
||||||
|
the exe) and the `LOG_FILE` directory (created if missing), plus read
|
||||||
|
access to the config file when it lives elsewhere. The grants happen
|
||||||
|
after service registration because the virtual account's SID only exists
|
||||||
|
from that point on; if a grant fails the service registration is rolled
|
||||||
|
back.
|
||||||
|
|
||||||
|
### Linux (systemd)
|
||||||
|
|
||||||
|
- `--install-service` copies the binary to `/var/lib/gpu-turnstile/`,
|
||||||
|
copies the config to `/etc/gpu-turnstile.env` if none exists there yet,
|
||||||
|
writes `/etc/systemd/system/gpu-turnstile.service`, then runs `systemctl
|
||||||
|
daemon-reload` and `enable --now`. `--remove-service` removes the unit
|
||||||
|
and the installed binary; the `/etc` config stays. The binary does not
|
||||||
|
go to `/usr/local/sbin` on purpose: replacing a running binary needs
|
||||||
|
write access to its *directory*, and granting the sandboxed service
|
||||||
|
write access to a shared system directory would let a compromised
|
||||||
|
service overwrite other binaries — `/var/lib/gpu-turnstile` is
|
||||||
|
exclusively ours.
|
||||||
|
- **Sandboxing** mirrors the Windows virtual account: the unit runs with
|
||||||
|
`DynamicUser=yes` — a transient per-service UID with no login, no home
|
||||||
|
and no password, managed entirely by systemd. `ProtectSystem=strict`
|
||||||
|
makes the filesystem read-only except `StateDirectory=gpu-turnstile`
|
||||||
|
(the install dir, so self-updates can rewrite the binary), plus
|
||||||
|
`NoNewPrivileges`, `ProtectHome`, `PrivateTmp`, `ProtectKernel*`,
|
||||||
|
`ProtectControlGroups`, `RestrictNamespaces`, `RestrictSUIDSGID`,
|
||||||
|
`RestrictRealtime`, `LockPersonality`, `MemoryDenyWriteExecute`, empty
|
||||||
|
capability sets, `RestrictAddressFamilies=AF_UNIX AF_INET AF_INET6` and
|
||||||
|
`SystemCallFilter=@system-service`. The proxy needs only outbound
|
||||||
|
TCP/UDP and the notify socket, so it loses nothing.
|
||||||
|
- The unit is `Type=notify`: the binary sends `READY=1` via
|
||||||
|
`github.com/coreos/go-systemd` only after the listeners are bound, so
|
||||||
|
`systemctl start` blocks until the proxy accepts connections. A 30 s
|
||||||
|
watchdog (`WatchdogSec=`) is pinged as long as the process runs; three
|
||||||
|
missed pings make systemd restart it. `STOPPING=1` is sent on shutdown.
|
||||||
|
All notify calls are no-ops when `NOTIFY_SOCKET` is unset (containers,
|
||||||
|
interactive shells), and the whole integration is Linux-only — Windows
|
||||||
|
builds carry no-op stubs.
|
||||||
|
- Logs go to the journal (`journalctl -u gpu-turnstile`) or to `LOG_FILE`.
|
||||||
|
- **Auto-update** works the same as on Windows: `Restart=on-failure` with
|
||||||
|
`RestartSec=5s` brings up the staged binary after the updater exits with
|
||||||
|
code 3.
|
||||||
|
- **Auto-update**: on startup and every `UPDATE_INTERVAL`, the binary
|
||||||
|
consults `APP_VER`: `dev` disables updates; `stable` (the default)
|
||||||
|
fetches `UPDATE_REPO`'s latest release and applies it when its tag is a
|
||||||
|
newer `vX.Y.Z` (a `dev` binary cannot be compared and is replaced by the
|
||||||
|
latest release); a `vX.Y.Z` pin fetches that exact tag and stages it on
|
||||||
|
any difference, including downgrades. Applying means downloading
|
||||||
|
`UPDATE_ASSET` plus its `.sig` (and `.sha256` when present) and verifying
|
||||||
|
an Ed25519 signature against the public key embedded in
|
||||||
|
`internal/update/pubkey.go`. A verified binary is swapped in next to the
|
||||||
|
running exe (rename-aside, allowed on Windows), and once the GPU lock is
|
||||||
|
idle the process exits with code 3 so the service recovery restarts it
|
||||||
|
on the new version. Interactive runs only log "restart to apply".
|
||||||
|
Builds without an embedded public key never update.
|
||||||
|
- **`--force-update`** runs the same check immediately, single-shot: one
|
||||||
|
attempt with a 30 s timeout, then exit — "up to date" (exit 0) or the
|
||||||
|
error (exit 1), no retries. When a newer release is found it downloads,
|
||||||
|
verifies and stages it, and if the service is running it restarts it
|
||||||
|
right away (otherwise the new version applies on next start). On Windows
|
||||||
|
it elevates via UAC only when the stage or restart needs permissions the
|
||||||
|
caller does not have.
|
||||||
|
- **Signing setup (one time)**: `openssl genpkey -algorithm ed25519 -out
|
||||||
|
private.pem`; `openssl pkey -in private.pem -pubout -out public.pem`.
|
||||||
|
Private key → repo secret `RELEASE_SIGNING_KEY`; public key → committed into
|
||||||
|
`internal/update/pubkey.go`. CI signs release binaries with
|
||||||
|
`openssl pkeyutl -sign -rawin`.
|
||||||
|
|
||||||
## Observability
|
## Observability
|
||||||
|
|
||||||
- `GET /healthz` on both listeners: 200 with JSON
|
- `GET /healthz` on both listeners: 200 with JSON
|
||||||
`{"state":"idle|llm|image","llm_inflight":N,"image_pending":B}`.
|
`{"state":"idle|llm|image","llm_inflight":N,"image_pending":B}`.
|
||||||
- `GET /metrics` on the Ollama listener: Prometheus text format, no external
|
- `GET /metrics` on both listeners: Prometheus text format, no external
|
||||||
dependency needed:
|
dependency needed:
|
||||||
`gpu_turnstile_state{state="…"} 1`, `gpu_turnstile_llm_inflight`,
|
`gpu_turnstile_state{state="…"} 1`, `gpu_turnstile_llm_inflight`,
|
||||||
`gpu_turnstile_image_jobs_total`, `gpu_turnstile_lock_wait_seconds`
|
`gpu_turnstile_image_jobs_total`, `gpu_turnstile_lock_wait_seconds`
|
||||||
(histogram, label `kind="llm|image"`), `gpu_turnstile_unload_seconds`.
|
(histogram, label `kind="llm|image"`), `gpu_turnstile_unload_seconds`.
|
||||||
- Structured logs (`log/slog`, JSON when `LOG_FORMAT=json`), one line per
|
- Structured logs (`log/slog`, JSON when `LOG_FORMAT=json`), one line per
|
||||||
state transition and per image job phase with `prompt_id`.
|
state transition and per image job phase with `prompt_id`. Startup logs
|
||||||
|
the version and every setting (visible even at the default `warn`
|
||||||
|
level). With `LOGLEVEL=info` or `debug`, every request logs a `-->`
|
||||||
|
incoming line and a `<--` response line with status and duration —
|
||||||
|
ANSI-colored (cyan incoming; green/yellow/red by status class) in text
|
||||||
|
mode, which renders in `docker compose logs` on Windows Terminal. Set
|
||||||
|
`NO_COLOR` to disable colors.
|
||||||
|
|
||||||
## Edge cases to handle
|
## Edge cases to handle
|
||||||
|
|
||||||
@@ -137,6 +400,14 @@ start (`/api/version`, `/system_stats`); failure is logged, not fatal.
|
|||||||
`JOB_TIMEOUT` releases the lock; log at warn.
|
`JOB_TIMEOUT` releases the lock; log at warn.
|
||||||
- Ollama unreachable during unload: continue with the image job; the whole
|
- Ollama unreachable during unload: continue with the image job; the whole
|
||||||
point is not to block users on a misbehaving neighbour.
|
point is not to block users on a misbehaving neighbour.
|
||||||
|
- Upstream unreachable while proxying (connection refused, dial timeout,
|
||||||
|
DNS failure, TLS handshake error): retry with exponential backoff —
|
||||||
|
`BACKOFF_INITIAL`, doubling per attempt, capped at `BACKOFF_MAX` — until
|
||||||
|
the upstream answers or the client disconnects. These are safe to retry:
|
||||||
|
the request never reached the upstream application. 5xx responses are
|
||||||
|
retried the same way, but only when the request body can be replayed
|
||||||
|
(GETs, or bodies with `GetBody`); streamed POSTs are never replayed to
|
||||||
|
avoid duplicate work such as a double-enqueued ComfyUI prompt.
|
||||||
- `POST /prompt` with a body that ComfyUI rejects (400): lock released
|
- `POST /prompt` with a body that ComfyUI rejects (400): lock released
|
||||||
immediately, body passed back.
|
immediately, body passed back.
|
||||||
- Websocket `/ws` connections are long-lived and never take the lock.
|
- Websocket `/ws` connections are long-lived and never take the lock.
|
||||||
@@ -146,11 +417,15 @@ start (`/api/version`, `/system_stats`); failure is logged, not fatal.
|
|||||||
|
|
||||||
```
|
```
|
||||||
gpu-turnstile/
|
gpu-turnstile/
|
||||||
cmd/gpu-turnstile/main.go # wiring, config, listeners
|
cmd/gpu-turnstile/main.go # wiring, config, listeners, service + updater
|
||||||
internal/lock/lock.go # two-mode lock + tests
|
internal/lock/lock.go # two-mode lock + tests
|
||||||
internal/ollama/client.go # ps / unload / warm
|
internal/ollama/client.go # ps / unload / warm
|
||||||
internal/comfy/client.go # history poll / free
|
internal/comfy/client.go # history poll / free
|
||||||
internal/proxy/ # handlers for both listeners
|
internal/proxy/ # handlers for both listeners
|
||||||
|
internal/metrics/ # Prometheus exposition
|
||||||
|
internal/config/ # env + .env file configuration
|
||||||
|
internal/update/ # signed auto-updater (public key in pubkey.go)
|
||||||
|
internal/service/ # Windows SCM + Linux systemd (notify/watchdog) integration
|
||||||
Dockerfile
|
Dockerfile
|
||||||
.gitea/workflows/ci.yml
|
.gitea/workflows/ci.yml
|
||||||
README.md
|
README.md
|
||||||
@@ -172,11 +447,18 @@ are new.
|
|||||||
held until `/free` was called.
|
held until `/free` was called.
|
||||||
- Streaming test: fake Ollama emits chunks with delays; assert the client
|
- Streaming test: fake Ollama emits chunks with delays; assert the client
|
||||||
receives the first chunk before the last is sent (no buffering).
|
receives the first chunk before the last is sent (no buffering).
|
||||||
|
- `internal/config`: env-file parsing, precedence, fail-fast values.
|
||||||
|
- `internal/update`: fake Gitea releases API; staged update happy path,
|
||||||
|
tampered signature rejected, older versions skipped, APP_VER=dev and
|
||||||
|
pinned releases honored.
|
||||||
|
|
||||||
## Build and CI
|
## Build and CI
|
||||||
|
|
||||||
- Go 1.23+, stdlib only. `CGO_ENABLED=0`, `-ldflags="-s -w"`, version from
|
- Go 1.23+, two external dependencies: `golang.org/x/sys` (Windows service
|
||||||
`git describe` injected via `-X main.version=`.
|
integration) and `github.com/coreos/go-systemd` (systemd notify/watchdog,
|
||||||
|
Linux build only). `CGO_ENABLED=0`,
|
||||||
|
`-ldflags="-s -w"`, version from `git describe` injected via
|
||||||
|
`-X main.version=`.
|
||||||
- Dockerfile: multi-stage, final image `gcr.io/distroless/static` (or
|
- Dockerfile: multi-stage, final image `gcr.io/distroless/static` (or
|
||||||
`scratch`), non-root user, `EXPOSE 8188 11434`,
|
`scratch`), non-root user, `EXPOSE 8188 11434`,
|
||||||
`ENTRYPOINT ["/gpu-turnstile"]`.
|
`ENTRYPOINT ["/gpu-turnstile"]`.
|
||||||
@@ -185,16 +467,24 @@ are new.
|
|||||||
available in the runner image
|
available in the runner image
|
||||||
2. on a version tag only (`vX.Y.Z`, enforced): build the image with buildx
|
2. on a version tag only (`vX.Y.Z`, enforced): build the image with buildx
|
||||||
and push it to the Gitea registry
|
and push it to the Gitea registry
|
||||||
`git.rambossek.at/<owner>/gpu-turnstile` tagged `:<tag>` and `:latest`,
|
`git.rambossek.at/<owner>/gpu-turnstile` tagged `:<tag>` and `:latest`
|
||||||
using the workflow token (`${{ secrets.GITEA_TOKEN }}` / `gitea.actor`)
|
(the repository path is lowercased in the workflow; Docker registry
|
||||||
- Release: a git tag `vX.Y.Z` produces the versioned image; the Open WebUI
|
names must be lowercase).
|
||||||
compose pins that tag. No images are built from branches.
|
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`.
|
||||||
|
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)
|
## Deployment (target)
|
||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
gpu-turnstile:
|
gpu-turnstile:
|
||||||
image: git.rambossek.at/<owner>/gpu-turnstile:v0.1.0
|
image: git.rambossek.at/<owner>/gpu-turnstile:v0.1.0 # owner lowercased, e.g. "public"
|
||||||
environment:
|
environment:
|
||||||
OLLAMA_URL: http://<workstation-ip>:11435
|
OLLAMA_URL: http://<workstation-ip>:11435
|
||||||
COMFY_URL: http://<workstation-ip>:8189
|
COMFY_URL: http://<workstation-ip>:8189
|
||||||
|
|||||||
+819
-126
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,36 @@
|
|||||||
|
# Example deployment for gpu-turnstile. Copy to compose.yaml and adjust.
|
||||||
|
#
|
||||||
|
# gpu-turnstile listens on the ports the services normally use; the actual
|
||||||
|
# Ollama and ComfyUI instances run one port higher (11435 / 8189) and must
|
||||||
|
# bind 0.0.0.0 so the container can reach them (OLLAMA_HOST=0.0.0.0:11435,
|
||||||
|
# ComfyUI --listen 0.0.0.0 --port 8189).
|
||||||
|
services:
|
||||||
|
gpu-turnstile:
|
||||||
|
image: git.rambossek.at/public/gpu-turnstile:v0.2.1
|
||||||
|
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
|
||||||
|
# Services on another machine: use its LAN IP instead, e.g.
|
||||||
|
# OLLAMA_URL: http://192.168.1.10:11435
|
||||||
|
# COMFY_URL: http://192.168.1.10:8189
|
||||||
|
# UNLOAD_TIMEOUT: 60s
|
||||||
|
# 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
|
||||||
|
ports:
|
||||||
|
- "11434:11434" # LiteLLM api_base -> http://gpu-turnstile:11434
|
||||||
|
- "8188:8188" # Open WebUI COMFYUI_BASE_URL -> http://gpu-turnstile:8188
|
||||||
|
networks: [internal]
|
||||||
|
|
||||||
|
networks:
|
||||||
|
internal:
|
||||||
|
external: true
|
||||||
@@ -1,3 +1,7 @@
|
|||||||
module gpu-turnstile
|
module gpu-turnstile
|
||||||
|
|
||||||
go 1.23
|
go 1.23
|
||||||
|
|
||||||
|
require golang.org/x/sys v0.29.0
|
||||||
|
|
||||||
|
require github.com/coreos/go-systemd/v22 v22.7.0
|
||||||
|
|||||||
@@ -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,318 @@
|
|||||||
|
// Package config loads gpu-turnstile's configuration from environment
|
||||||
|
// variables and an optional .env-style config file.
|
||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bufio"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Config holds every gpu-turnstile setting.
|
||||||
|
type Config struct {
|
||||||
|
ListenOllama string
|
||||||
|
ListenComfy string
|
||||||
|
OllamaURL string
|
||||||
|
ComfyURL string
|
||||||
|
UnloadTimeout time.Duration
|
||||||
|
JobTimeout time.Duration
|
||||||
|
LLMWaitTimeout time.Duration
|
||||||
|
|
||||||
|
UnloadPollInterval time.Duration
|
||||||
|
HistoryPollInterval time.Duration
|
||||||
|
ProbeTimeout time.Duration
|
||||||
|
HealthInterval time.Duration
|
||||||
|
FreeTimeout time.Duration
|
||||||
|
WarmTimeout time.Duration
|
||||||
|
ShutdownTimeout time.Duration
|
||||||
|
BackoffInitial time.Duration
|
||||||
|
BackoffMax time.Duration
|
||||||
|
PromptCaptureLimit int64
|
||||||
|
|
||||||
|
AutoUpdate bool
|
||||||
|
UpdateInterval time.Duration
|
||||||
|
UpdateRepo string
|
||||||
|
UpdateAsset string
|
||||||
|
|
||||||
|
// AppVersion is the version the user wants to run: "dev" disables
|
||||||
|
// updates, "stable" tracks the latest release, anything else is an
|
||||||
|
// exact vX.Y.Z release to pin. From APP_VER; defaults to "stable".
|
||||||
|
AppVersion string
|
||||||
|
|
||||||
|
// LLMBusyMode is "wait" (hold requests until the lock is free or
|
||||||
|
// LLMWaitTimeout expires) or "reject" (immediately answer with
|
||||||
|
// LLMBusyStatus + Retry-After when an image job is active or pending).
|
||||||
|
LLMBusyMode string
|
||||||
|
LLMBusyStatus int
|
||||||
|
BusyRetryAfter int
|
||||||
|
|
||||||
|
// ComfyCmd spawns and supervises a ComfyUI server on demand. When
|
||||||
|
// ComfyCmd is empty but ComfyDir is set, management is enabled with the
|
||||||
|
// standard venv layout under ComfyDir (.venv + main.py or
|
||||||
|
// ComfyUI/main.py; --port from the COMFY_URL port) — ComfyCmd is the
|
||||||
|
// override for other layouts and doubles as the working directory when
|
||||||
|
// set explicitly. The managed server is stopped after ComfyIdleTimeout
|
||||||
|
// without requests, freeing its VRAM; ComfyStartTimeout bounds how long
|
||||||
|
// a request waits for it to come up.
|
||||||
|
ComfyCmd string
|
||||||
|
ComfyDir string
|
||||||
|
ComfyIdleTimeout time.Duration
|
||||||
|
ComfyStartTimeout time.Duration
|
||||||
|
|
||||||
|
// GameProcs (GAME_PROCS) is a watch list of process names; while any of
|
||||||
|
// them runs, the GPU is treated as held by a foreign process. The
|
||||||
|
// nvidia-smi path (GPUForeignVRAMMB, GPU_FOREIGN_VRAM_MB) does the same
|
||||||
|
// when a process not in GPUIgnoreProcs (GPU_IGNORE_PROCS) holds more than
|
||||||
|
// that many MiB of VRAM. GamePollInterval (GAME_POLL_INTERVAL) is how
|
||||||
|
// often both checks run.
|
||||||
|
GameProcs []string
|
||||||
|
GPUForeignVRAMMB int
|
||||||
|
GPUIgnoreProcs []string
|
||||||
|
GamePollInterval time.Duration
|
||||||
|
|
||||||
|
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,
|
||||||
|
HealthInterval: 30 * time.Second,
|
||||||
|
FreeTimeout: 30 * time.Second,
|
||||||
|
WarmTimeout: 2 * time.Minute,
|
||||||
|
ShutdownTimeout: 10 * time.Second,
|
||||||
|
BackoffInitial: time.Second,
|
||||||
|
BackoffMax: time.Minute,
|
||||||
|
PromptCaptureLimit: 64 * 1024,
|
||||||
|
|
||||||
|
AutoUpdate: true,
|
||||||
|
UpdateInterval: 6 * time.Hour,
|
||||||
|
UpdateRepo: "https://git.rambossek.at/PUBLIC/gpu-turnstile",
|
||||||
|
UpdateAsset: "gpu-turnstile.exe",
|
||||||
|
AppVersion: "stable",
|
||||||
|
|
||||||
|
LLMBusyMode: "wait",
|
||||||
|
LLMBusyStatus: 503,
|
||||||
|
BusyRetryAfter: 30,
|
||||||
|
|
||||||
|
ComfyIdleTimeout: 5 * time.Minute,
|
||||||
|
ComfyStartTimeout: 2 * time.Minute,
|
||||||
|
|
||||||
|
// ComfyUI runs under python; excluding it (and Ollama) by name keeps
|
||||||
|
// our own consumers from tripping the foreign-VRAM check.
|
||||||
|
GPUIgnoreProcs: []string{"ollama", "ollama app", "ollama_llama_server", "python", "pythonw"},
|
||||||
|
GamePollInterval: 15 * time.Second,
|
||||||
|
|
||||||
|
LogLevel: slog.LevelWarn,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ErrNoConsumer is returned by Load when neither OLLAMA_URL nor COMFY_URL
|
||||||
|
// is set. Commands that never talk to an upstream (--force-update) may
|
||||||
|
// ignore it and proceed with the rest of the configuration.
|
||||||
|
var ErrNoConsumer = errors.New("at least one of OLLAMA_URL or COMFY_URL must be set (each URL enables its consumer)")
|
||||||
|
|
||||||
|
// ParseEnvFile parses a .env-style file: KEY=VALUE lines, blank lines and
|
||||||
|
// #-comments are ignored, no quoting. A line without '=' is an error.
|
||||||
|
func ParseEnvFile(r io.Reader) (map[string]string, error) {
|
||||||
|
values := make(map[string]string)
|
||||||
|
scanner := bufio.NewScanner(r)
|
||||||
|
scanner.Buffer(make([]byte, 64*1024), 1024*1024)
|
||||||
|
lineNo := 0
|
||||||
|
for scanner.Scan() {
|
||||||
|
lineNo++
|
||||||
|
line := strings.TrimSpace(scanner.Text())
|
||||||
|
if line == "" || strings.HasPrefix(line, "#") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
key, value, ok := strings.Cut(line, "=")
|
||||||
|
if !ok {
|
||||||
|
return nil, fmt.Errorf("line %d: expected KEY=VALUE", lineNo)
|
||||||
|
}
|
||||||
|
key = strings.TrimSpace(key)
|
||||||
|
if key == "" {
|
||||||
|
return nil, fmt.Errorf("line %d: empty key", lineNo)
|
||||||
|
}
|
||||||
|
values[key] = strings.TrimSpace(value)
|
||||||
|
}
|
||||||
|
return values, scanner.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
// splitList parses a comma-separated setting into trimmed, non-empty items.
|
||||||
|
func splitList(v string) []string {
|
||||||
|
var out []string
|
||||||
|
for _, item := range strings.Split(v, ",") {
|
||||||
|
if item = strings.TrimSpace(item); item != "" {
|
||||||
|
out = append(out, item)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
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},
|
||||||
|
{"COMFY_CMD", &cfg.ComfyCmd},
|
||||||
|
{"COMFY_DIR", &cfg.ComfyDir},
|
||||||
|
{"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},
|
||||||
|
{"HEALTH_INTERVAL", &cfg.HealthInterval},
|
||||||
|
{"FREE_TIMEOUT", &cfg.FreeTimeout},
|
||||||
|
{"WARM_TIMEOUT", &cfg.WarmTimeout},
|
||||||
|
{"SHUTDOWN_TIMEOUT", &cfg.ShutdownTimeout},
|
||||||
|
{"BACKOFF_INITIAL", &cfg.BackoffInitial},
|
||||||
|
{"BACKOFF_MAX", &cfg.BackoffMax},
|
||||||
|
{"COMFY_IDLE_TIMEOUT", &cfg.ComfyIdleTimeout},
|
||||||
|
{"COMFY_START_TIMEOUT", &cfg.ComfyStartTimeout},
|
||||||
|
{"UPDATE_INTERVAL", &cfg.UpdateInterval},
|
||||||
|
{"GAME_POLL_INTERVAL", &cfg.GamePollInterval},
|
||||||
|
} {
|
||||||
|
if err := envDuration(getenv, e.name, e.dst); err != nil {
|
||||||
|
return cfg, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if v := getenv("GAME_PROCS"); v != "" {
|
||||||
|
cfg.GameProcs = splitList(v)
|
||||||
|
}
|
||||||
|
if v := getenv("GPU_IGNORE_PROCS"); v != "" {
|
||||||
|
cfg.GPUIgnoreProcs = splitList(v)
|
||||||
|
}
|
||||||
|
if v := getenv("GPU_FOREIGN_VRAM_MB"); v != "" {
|
||||||
|
n, err := strconv.Atoi(v)
|
||||||
|
if err != nil || n < 0 {
|
||||||
|
return cfg, fmt.Errorf("GPU_FOREIGN_VRAM_MB: must be a non-negative integer (MiB, 0 = disabled)")
|
||||||
|
}
|
||||||
|
cfg.GPUForeignVRAMMB = n
|
||||||
|
}
|
||||||
|
if v := getenv("PROMPT_CAPTURE_LIMIT"); v != "" {
|
||||||
|
n, err := strconv.ParseInt(v, 10, 64)
|
||||||
|
if err != nil || n < 0 {
|
||||||
|
return cfg, fmt.Errorf("PROMPT_CAPTURE_LIMIT: must be a non-negative integer (bytes)")
|
||||||
|
}
|
||||||
|
cfg.PromptCaptureLimit = n
|
||||||
|
}
|
||||||
|
if v := getenv("AUTO_UPDATE"); v != "" {
|
||||||
|
b, err := strconv.ParseBool(v)
|
||||||
|
if err != nil {
|
||||||
|
return cfg, fmt.Errorf("AUTO_UPDATE: must be a boolean (true/false)")
|
||||||
|
}
|
||||||
|
cfg.AutoUpdate = b
|
||||||
|
}
|
||||||
|
if v := getenv("LLM_BUSY_MODE"); v != "" {
|
||||||
|
if v != "wait" && v != "reject" {
|
||||||
|
return cfg, fmt.Errorf("LLM_BUSY_MODE: must be \"wait\" or \"reject\"")
|
||||||
|
}
|
||||||
|
cfg.LLMBusyMode = v
|
||||||
|
}
|
||||||
|
if v := getenv("LLM_BUSY_STATUS"); v != "" {
|
||||||
|
n, err := strconv.Atoi(v)
|
||||||
|
if err != nil || n < 400 || n > 599 {
|
||||||
|
return cfg, fmt.Errorf("LLM_BUSY_STATUS: must be an HTTP status in 400-599")
|
||||||
|
}
|
||||||
|
cfg.LLMBusyStatus = n
|
||||||
|
}
|
||||||
|
if v := getenv("BUSY_RETRY_AFTER"); v != "" {
|
||||||
|
n, err := strconv.Atoi(v)
|
||||||
|
if err != nil || n <= 0 {
|
||||||
|
return cfg, fmt.Errorf("BUSY_RETRY_AFTER: must be a positive integer (seconds)")
|
||||||
|
}
|
||||||
|
cfg.BusyRetryAfter = n
|
||||||
|
}
|
||||||
|
if v := getenv("APP_VER"); v != "" {
|
||||||
|
switch {
|
||||||
|
case v == "dev" || v == "stable":
|
||||||
|
cfg.AppVersion = v
|
||||||
|
default:
|
||||||
|
if _, ok := parseVersion(v); !ok {
|
||||||
|
return cfg, fmt.Errorf("APP_VER: must be \"dev\", \"stable\" or a vX.Y.Z version")
|
||||||
|
}
|
||||||
|
cfg.AppVersion = "v" + strings.TrimPrefix(v, "v")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// LOGLEVEL is the canonical spelling; LOG_LEVEL is kept as an alias.
|
||||||
|
logLevelValue := getenv("LOGLEVEL")
|
||||||
|
if logLevelValue == "" {
|
||||||
|
logLevelValue = getenv("LOG_LEVEL")
|
||||||
|
}
|
||||||
|
if logLevelValue != "" {
|
||||||
|
var level slog.Level
|
||||||
|
if err := level.UnmarshalText([]byte(logLevelValue)); err != nil {
|
||||||
|
return cfg, fmt.Errorf("LOGLEVEL: %w", err)
|
||||||
|
}
|
||||||
|
cfg.LogLevel = level
|
||||||
|
}
|
||||||
|
switch strings.ToLower(getenv("LOG_FORMAT")) {
|
||||||
|
case "", "text":
|
||||||
|
case "json":
|
||||||
|
cfg.LogJSON = true
|
||||||
|
default:
|
||||||
|
return cfg, fmt.Errorf("LOG_FORMAT: must be \"text\" or \"json\"")
|
||||||
|
}
|
||||||
|
if cfg.ComfyCmd != "" && cfg.ComfyURL == "" {
|
||||||
|
return cfg, fmt.Errorf("COMFY_CMD requires COMFY_URL to be set (the proxy needs somewhere to forward)")
|
||||||
|
}
|
||||||
|
if cfg.ComfyCmd == "" && cfg.ComfyDir != "" && cfg.ComfyURL == "" {
|
||||||
|
return cfg, fmt.Errorf("COMFY_DIR without COMFY_CMD requires COMFY_URL to be set (it enables the managed ComfyUI)")
|
||||||
|
}
|
||||||
|
if cfg.OllamaURL == "" && cfg.ComfyURL == "" {
|
||||||
|
return cfg, ErrNoConsumer
|
||||||
|
}
|
||||||
|
return cfg, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,249 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"log/slog"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestDefaults(t *testing.T) {
|
||||||
|
cfg, err := Load(func(k string) string {
|
||||||
|
if k == "OLLAMA_URL" {
|
||||||
|
return "http://127.0.0.1:11435"
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if cfg.ListenOllama != ":11434" || cfg.ListenComfy != ":8188" {
|
||||||
|
t.Fatalf("listen addrs = %s %s", cfg.ListenOllama, cfg.ListenComfy)
|
||||||
|
}
|
||||||
|
if cfg.ComfyURL != "" {
|
||||||
|
t.Fatalf("ComfyURL default = %q, want empty (disabled)", cfg.ComfyURL)
|
||||||
|
}
|
||||||
|
if cfg.UnloadTimeout != time.Minute || cfg.JobTimeout != 15*time.Minute {
|
||||||
|
t.Fatalf("timeouts = %v %v", cfg.UnloadTimeout, cfg.JobTimeout)
|
||||||
|
}
|
||||||
|
if !cfg.AutoUpdate || cfg.UpdateInterval != 6*time.Hour {
|
||||||
|
t.Fatalf("update = %v %v", cfg.AutoUpdate, cfg.UpdateInterval)
|
||||||
|
}
|
||||||
|
if cfg.LogLevel != slog.LevelWarn {
|
||||||
|
t.Fatalf("log level = %v", cfg.LogLevel)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLoadRequiresConsumer(t *testing.T) {
|
||||||
|
_, err := Load(func(string) string { return "" })
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "OLLAMA_URL") {
|
||||||
|
t.Fatalf("err = %v, want missing-consumer error", err)
|
||||||
|
}
|
||||||
|
if !errors.Is(err, ErrNoConsumer) {
|
||||||
|
t.Fatalf("err = %v, want errors.Is(err, ErrNoConsumer)", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAppVersion(t *testing.T) {
|
||||||
|
load := func(appVer string) (Config, error) {
|
||||||
|
return Load(func(k string) string {
|
||||||
|
switch k {
|
||||||
|
case "OLLAMA_URL":
|
||||||
|
return "http://127.0.0.1:11435"
|
||||||
|
case "APP_VER":
|
||||||
|
return appVer
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
})
|
||||||
|
}
|
||||||
|
cfg, err := load("")
|
||||||
|
if err != nil || cfg.AppVersion != "stable" {
|
||||||
|
t.Fatalf("default AppVersion = %q, err %v; want stable", cfg.AppVersion, err)
|
||||||
|
}
|
||||||
|
for _, v := range []string{"dev", "stable"} {
|
||||||
|
if cfg, err := load(v); err != nil || cfg.AppVersion != v {
|
||||||
|
t.Fatalf("APP_VER=%s: got %q, err %v", v, cfg.AppVersion, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if cfg, err := load("1.2.3"); err != nil || cfg.AppVersion != "v1.2.3" {
|
||||||
|
t.Fatalf("APP_VER=1.2.3: got %q, err %v; want normalized v1.2.3", cfg.AppVersion, err)
|
||||||
|
}
|
||||||
|
if _, err := load("nightly"); err == nil {
|
||||||
|
t.Fatal("APP_VER=nightly: want validation error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestComfyCmdRequiresURL(t *testing.T) {
|
||||||
|
_, err := Load(func(k string) string {
|
||||||
|
if k == "COMFY_CMD" {
|
||||||
|
return "python main.py"
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
})
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "COMFY_CMD requires COMFY_URL") {
|
||||||
|
t.Fatalf("err = %v, want COMFY_CMD/COMFY_URL validation error", err)
|
||||||
|
}
|
||||||
|
if _, err := Load(func(k string) string {
|
||||||
|
switch k {
|
||||||
|
case "COMFY_CMD":
|
||||||
|
return "python main.py"
|
||||||
|
case "COMFY_URL":
|
||||||
|
return "http://127.0.0.1:8188"
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("COMFY_CMD with COMFY_URL must load: %v", 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"},
|
||||||
|
{"GPU_FOREIGN_VRAM_MB", "-1"},
|
||||||
|
{"GAME_POLL_INTERVAL", "bogus"},
|
||||||
|
} {
|
||||||
|
_, 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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestGameDetectionSettings(t *testing.T) {
|
||||||
|
cfg, err := Load(func(k string) string {
|
||||||
|
switch k {
|
||||||
|
case "OLLAMA_URL":
|
||||||
|
return "http://127.0.0.1:11435"
|
||||||
|
case "GAME_PROCS":
|
||||||
|
return " cyberpunk2077.exe, hl2.exe ,, "
|
||||||
|
case "GPU_FOREIGN_VRAM_MB":
|
||||||
|
return "1024"
|
||||||
|
case "GPU_IGNORE_PROCS":
|
||||||
|
return "ollama, my-trainer"
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(cfg.GameProcs) != 2 || cfg.GameProcs[0] != "cyberpunk2077.exe" || cfg.GameProcs[1] != "hl2.exe" {
|
||||||
|
t.Fatalf("GameProcs = %v", cfg.GameProcs)
|
||||||
|
}
|
||||||
|
if cfg.GPUForeignVRAMMB != 1024 {
|
||||||
|
t.Fatalf("GPUForeignVRAMMB = %d", cfg.GPUForeignVRAMMB)
|
||||||
|
}
|
||||||
|
if len(cfg.GPUIgnoreProcs) != 2 || cfg.GPUIgnoreProcs[1] != "my-trainer" {
|
||||||
|
t.Fatalf("GPUIgnoreProcs = %v", cfg.GPUIgnoreProcs)
|
||||||
|
}
|
||||||
|
if cfg.GamePollInterval != 15*time.Second {
|
||||||
|
t.Fatalf("GamePollInterval = %v, want 15s default", cfg.GamePollInterval)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Defaults: both detection paths off, ignore list covers our consumers.
|
||||||
|
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 len(cfg.GameProcs) != 0 || cfg.GPUForeignVRAMMB != 0 {
|
||||||
|
t.Fatalf("detection must be off by default: %v %d", cfg.GameProcs, cfg.GPUForeignVRAMMB)
|
||||||
|
}
|
||||||
|
if len(cfg.GPUIgnoreProcs) == 0 {
|
||||||
|
t.Fatal("GPUIgnoreProcs default must not be empty")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestComfyDirOnlyEnablesManaged(t *testing.T) {
|
||||||
|
// COMFY_DIR without COMFY_CMD and without COMFY_URL is a mistake.
|
||||||
|
_, err := Load(func(k string) string {
|
||||||
|
if k == "COMFY_DIR" {
|
||||||
|
return `C:\ComfyUI`
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
})
|
||||||
|
if err == nil || !strings.Contains(err.Error(), "COMFY_DIR") {
|
||||||
|
t.Fatalf("err = %v, want COMFY_DIR/COMFY_URL validation error", err)
|
||||||
|
}
|
||||||
|
// With COMFY_URL it loads — the launch command is derived from the dir.
|
||||||
|
if _, err := Load(func(k string) string {
|
||||||
|
switch k {
|
||||||
|
case "COMFY_DIR":
|
||||||
|
return `C:\ComfyUI`
|
||||||
|
case "COMFY_URL":
|
||||||
|
return "http://127.0.0.1:8189"
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatalf("COMFY_DIR with COMFY_URL must load: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,193 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// sampleEntry is one line pair in the sample env file: a comment and the
|
||||||
|
// (usually commented-out) KEY=VALUE line.
|
||||||
|
type sampleEntry struct {
|
||||||
|
name string
|
||||||
|
value string
|
||||||
|
comment string
|
||||||
|
active bool // rendered uncommented
|
||||||
|
}
|
||||||
|
|
||||||
|
// appVerComment documents APP_VER wherever it is rendered.
|
||||||
|
const appVerComment = `Version to run: "dev" disables updates, "stable" tracks the latest release, or pin an exact release like v0.1.7`
|
||||||
|
|
||||||
|
// sampleEntries lists every setting in sample order. logFile activates the
|
||||||
|
// LOG_FILE line (Windows install); empty keeps it commented like the rest.
|
||||||
|
// CFG_VER and APP_VER are not entries — they head the file, always active.
|
||||||
|
func sampleEntries(logFile string) []sampleEntry {
|
||||||
|
return []sampleEntry{
|
||||||
|
{"LISTEN_OLLAMA", ":11434", "Listen address for Ollama-compatible clients (gpu-turnstile poses as Ollama here)", false},
|
||||||
|
{"LISTEN_COMFY", ":8188", "Listen address for ComfyUI clients (gpu-turnstile poses as ComfyUI here)", false},
|
||||||
|
{"OLLAMA_URL", "http://127.0.0.1:11434", "Ollama upstream URL; setting it enables the Ollama consumer (default: empty = disabled)", false},
|
||||||
|
{"COMFY_URL", "http://127.0.0.1:8188", "ComfyUI upstream URL; setting it enables the ComfyUI consumer (default: empty = disabled)", false},
|
||||||
|
{"WARM_MODEL", "", "Optional model to reload after an image job (default: empty = none)", false},
|
||||||
|
{"COMFY_CMD", `"C:\ComfyUI\.venv\Scripts\python.exe" ComfyUI\main.py --port 8188`, "Spawn and supervise ComfyUI on demand: the first request starts it, it stops after COMFY_IDLE_TIMEOUT to free VRAM (default: derived from COMFY_DIR; both empty = unmanaged)", false},
|
||||||
|
{"COMFY_DIR", `C:\ComfyUI`, "Root of a standard ComfyUI venv install (.venv + main.py): setting it alone supervises ComfyUI with the derived launch command; also the working directory for COMFY_CMD", false},
|
||||||
|
{"COMFY_IDLE_TIMEOUT", "5m", "Stop the managed ComfyUI after this long without requests or jobs (frees VRAM)", false},
|
||||||
|
{"COMFY_START_TIMEOUT", "2m", "How long a request waits for the managed ComfyUI to come up", false},
|
||||||
|
{"GAME_PROCS", "cyberpunk2077.exe,hl2.exe", "While a listed process runs, the GPU counts as held by it: requests wait, Ollama unloads, managed ComfyUI stops (default: empty = disabled)", false},
|
||||||
|
{"GPU_FOREIGN_VRAM_MB", "1024", "Also treat the GPU as held when a process not in GPU_IGNORE_PROCS uses more VRAM than this (needs nvidia-smi; 0/empty = disabled)", false},
|
||||||
|
{"GPU_IGNORE_PROCS", "ollama,ollama app,ollama_llama_server,python,pythonw", "Process names never counted as foreign GPU users (ComfyUI runs under python)", false},
|
||||||
|
{"GAME_POLL_INTERVAL", "15s", "How often game/VRAM detection runs (nvidia-smi polls keep the GPU awake; don't go below ~10s)", false},
|
||||||
|
{"UNLOAD_TIMEOUT", "60s", "How long to wait for Ollama to unload a model", false},
|
||||||
|
{"JOB_TIMEOUT", "15m", "Maximum time to wait for a ComfyUI job", false},
|
||||||
|
{"LLM_WAIT_TIMEOUT", "10m", "Max time an LLM request waits for the GPU before being answered 503 (wait mode)", false},
|
||||||
|
{"LLM_BUSY_MODE", "wait", `How blocked LLM requests are handled: "wait" holds them, "reject" fails them immediately`, false},
|
||||||
|
{"LLM_BUSY_STATUS", "503", "HTTP status for rejected LLM requests in reject mode (400-599, e.g. 429)", false},
|
||||||
|
{"BUSY_RETRY_AFTER", "30", "Seconds sent as Retry-After on busy responses (both modes)", false},
|
||||||
|
{"LOGLEVEL", "warn", `Log verbosity: debug, info, warn, error; "info" logs every request (LOG_LEVEL works too)`, false},
|
||||||
|
{"LOG_FORMAT", "text", `Log format: "text" or "json"`, false},
|
||||||
|
{"LOG_FILE", logFile, "Append logs to this file instead of stderr (a Windows service has no console)", logFile != ""},
|
||||||
|
{"UNLOAD_POLL_INTERVAL", "500ms", "/api/ps poll interval while unloading", false},
|
||||||
|
{"HISTORY_POLL_INTERVAL", "1s", "/history/<id> poll interval while a job runs", false},
|
||||||
|
{"PROBE_TIMEOUT", "5s", "Probe timeout for the startup probe and the periodic health check", false},
|
||||||
|
{"HEALTH_INTERVAL", "30s", "How often enabled upstreams are probed; status changes are logged", false},
|
||||||
|
{"FREE_TIMEOUT", "30s", "Timeout for the POST /free call after an image job", false},
|
||||||
|
{"WARM_TIMEOUT", "2m", "Timeout for the warm-model reload after an image job", false},
|
||||||
|
{"SHUTDOWN_TIMEOUT", "10s", "Graceful shutdown timeout on SIGINT/SIGTERM", false},
|
||||||
|
{"BACKOFF_INITIAL", "1s", "First retry wait when an upstream connection fails", false},
|
||||||
|
{"BACKOFF_MAX", "60s", "Cap for the exponential retry backoff", false},
|
||||||
|
{"PROMPT_CAPTURE_LIMIT", "65536", "Bytes of the /prompt response buffered to find prompt_id (pass-through unaffected)", false},
|
||||||
|
{"AUTO_UPDATE", "true", "Poll the releases API for signed updates", false},
|
||||||
|
{"UPDATE_INTERVAL", "6h", "Auto-update check interval", false},
|
||||||
|
{"UPDATE_REPO", "https://git.rambossek.at/PUBLIC/gpu-turnstile", "Repository to check for releases", false},
|
||||||
|
{"UPDATE_ASSET", "gpu-turnstile.exe", "Release asset to download", false},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SampleEnv renders a sample .env file covering every setting, each with a
|
||||||
|
// comment line. Everything is commented out — so all defaults apply —
|
||||||
|
// except the CFG_VER/APP_VER header and LOG_FILE when logFile is non-empty
|
||||||
|
// (a Windows service has no console). CFG_VER records the version that
|
||||||
|
// wrote the file so installs — and startups after an update — can upgrade
|
||||||
|
// it.
|
||||||
|
func SampleEnv(version, logFile string) string {
|
||||||
|
// CFG_VER is always a concrete vX.Y.Z — never "dev". A dev build
|
||||||
|
// stamps v0.0.0, which sorts older than any release, so the next
|
||||||
|
// release install upgrades the file and stamps a proper version.
|
||||||
|
cfgVer := version
|
||||||
|
if cfgVer == "" || cfgVer == "dev" {
|
||||||
|
cfgVer = "v0.0.0"
|
||||||
|
}
|
||||||
|
var b strings.Builder
|
||||||
|
fmt.Fprintf(&b, "CFG_VER=%s\n", cfgVer)
|
||||||
|
b.WriteString("# Config format reference, written by the installer — do not edit.\n")
|
||||||
|
b.WriteString("# The installer uses it to append newly added settings on updates.\n\n")
|
||||||
|
b.WriteString("# gpu-turnstile configuration\n")
|
||||||
|
b.WriteString("# KEY=VALUE lines; \"#\" starts a comment. Every setting below is at its\n")
|
||||||
|
b.WriteString("# default and commented out — remove the \"#\" to change it.\n")
|
||||||
|
b.WriteString("# At least one of OLLAMA_URL / COMFY_URL must be set for the proxy to start.\n\n")
|
||||||
|
writeEntry(&b, sampleEntry{"APP_VER", "stable", appVerComment, true})
|
||||||
|
for _, e := range sampleEntries(logFile) {
|
||||||
|
writeEntry(&b, e)
|
||||||
|
}
|
||||||
|
return b.String()
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeEntry(b *strings.Builder, e sampleEntry) {
|
||||||
|
fmt.Fprintf(b, "# %s\n", e.comment)
|
||||||
|
if e.active {
|
||||||
|
fmt.Fprintf(b, "%s=%s\n\n", e.name, e.value)
|
||||||
|
} else {
|
||||||
|
fmt.Fprintf(b, "#%s=%s\n\n", e.name, e.value)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SyncSample upgrades an installer-written env file: when its CFG_VER says
|
||||||
|
// it was written by an older gpu-turnstile, every setting the file does not
|
||||||
|
// mention — commented or not — is appended at the end, and CFG_VER is
|
||||||
|
// updated to version. Files without CFG_VER (hand-written or foreign),
|
||||||
|
// up-to-date files and "dev" builds are returned unchanged; changed reports
|
||||||
|
// whether the returned content differs.
|
||||||
|
func SyncSample(data, version, logFile string) (string, bool) {
|
||||||
|
if version == "" || version == "dev" {
|
||||||
|
return data, false
|
||||||
|
}
|
||||||
|
values, err := ParseEnvFile(strings.NewReader(data))
|
||||||
|
if err != nil || values["CFG_VER"] == "" {
|
||||||
|
return data, false // not installer-written; the caller decides
|
||||||
|
}
|
||||||
|
if compareVersions(values["CFG_VER"], version) >= 0 {
|
||||||
|
return data, false // same or newer
|
||||||
|
}
|
||||||
|
|
||||||
|
present := make(map[string]bool)
|
||||||
|
for _, line := range strings.Split(data, "\n") {
|
||||||
|
line = strings.TrimSpace(strings.TrimPrefix(strings.TrimSpace(line), "#"))
|
||||||
|
if name, _, ok := strings.Cut(line, "="); ok {
|
||||||
|
present[strings.TrimSpace(name)] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
lines := strings.Split(data, "\n")
|
||||||
|
for i, l := range lines {
|
||||||
|
if strings.HasPrefix(strings.TrimSpace(l), "CFG_VER=") {
|
||||||
|
lines[i] = "CFG_VER=" + version
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
out := strings.Join(lines, "\n")
|
||||||
|
if !strings.HasSuffix(out, "\n") {
|
||||||
|
out += "\n"
|
||||||
|
}
|
||||||
|
|
||||||
|
var b strings.Builder
|
||||||
|
if !present["APP_VER"] {
|
||||||
|
fmt.Fprintf(&b, "\n# Added by gpu-turnstile %s:\n", version)
|
||||||
|
writeEntry(&b, sampleEntry{"APP_VER", "stable", appVerComment, true})
|
||||||
|
}
|
||||||
|
for _, e := range sampleEntries(logFile) {
|
||||||
|
if !present[e.name] {
|
||||||
|
fmt.Fprintf(&b, "\n# Added by gpu-turnstile %s:\n", version)
|
||||||
|
writeEntry(&b, e)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
out += b.String()
|
||||||
|
return out, out != data
|
||||||
|
}
|
||||||
|
|
||||||
|
// compareVersions orders two vX.Y.Z version strings. Unparseable versions
|
||||||
|
// (including "dev") sort before any release.
|
||||||
|
func compareVersions(a, b string) int {
|
||||||
|
pa, oka := parseVersion(a)
|
||||||
|
pb, okb := parseVersion(b)
|
||||||
|
if oka != okb {
|
||||||
|
if oka {
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
for i := range pa {
|
||||||
|
if pa[i] != pb[i] {
|
||||||
|
if pa[i] > pb[i] {
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
return -1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseVersion(v string) ([3]int, bool) {
|
||||||
|
var out [3]int
|
||||||
|
v = strings.TrimPrefix(strings.TrimSpace(v), "v")
|
||||||
|
parts := strings.Split(v, ".")
|
||||||
|
if len(parts) != 3 {
|
||||||
|
return out, false
|
||||||
|
}
|
||||||
|
for i, p := range parts {
|
||||||
|
n, err := strconv.Atoi(p)
|
||||||
|
if err != nil || n < 0 {
|
||||||
|
return out, false
|
||||||
|
}
|
||||||
|
out[i] = n
|
||||||
|
}
|
||||||
|
return out, true
|
||||||
|
}
|
||||||
@@ -0,0 +1,156 @@
|
|||||||
|
package config
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
const sampleLogPath = `C:\ProgramData\gpu-turnstile\gpu-turnstile.log`
|
||||||
|
|
||||||
|
var allSettingNames = []string{
|
||||||
|
"LISTEN_OLLAMA", "LISTEN_COMFY", "OLLAMA_URL", "COMFY_URL",
|
||||||
|
"WARM_MODEL", "UNLOAD_TIMEOUT", "JOB_TIMEOUT", "LLM_WAIT_TIMEOUT",
|
||||||
|
"LLM_BUSY_MODE", "LLM_BUSY_STATUS", "BUSY_RETRY_AFTER",
|
||||||
|
"LOGLEVEL", "LOG_FORMAT", "LOG_FILE",
|
||||||
|
"UNLOAD_POLL_INTERVAL", "HISTORY_POLL_INTERVAL", "PROBE_TIMEOUT",
|
||||||
|
"FREE_TIMEOUT", "WARM_TIMEOUT", "SHUTDOWN_TIMEOUT",
|
||||||
|
"BACKOFF_INITIAL", "BACKOFF_MAX", "PROMPT_CAPTURE_LIMIT",
|
||||||
|
"AUTO_UPDATE", "UPDATE_INTERVAL", "UPDATE_REPO", "UPDATE_ASSET",
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSampleEnv(t *testing.T) {
|
||||||
|
sample := SampleEnv("v0.1.7", sampleLogPath)
|
||||||
|
|
||||||
|
if !strings.HasPrefix(sample, "CFG_VER=v0.1.7\n") {
|
||||||
|
t.Errorf("first line does not carry CFG_VER: %q", strings.SplitN(sample, "\n", 2)[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
// Every setting known to Load must appear.
|
||||||
|
for _, name := range allSettingNames {
|
||||||
|
if !strings.Contains(sample, name+"=") {
|
||||||
|
t.Errorf("sample is missing %s", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// The sample must parse cleanly; active values are CFG_VER, APP_VER
|
||||||
|
// and LOG_FILE.
|
||||||
|
values, err := ParseEnvFile(strings.NewReader(sample))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("sample does not parse: %v", err)
|
||||||
|
}
|
||||||
|
want := map[string]string{"CFG_VER": "v0.1.7", "APP_VER": "stable", "LOG_FILE": sampleLogPath}
|
||||||
|
if len(values) != len(want) {
|
||||||
|
t.Fatalf("active values = %v, want %v", values, want)
|
||||||
|
}
|
||||||
|
for k, v := range want {
|
||||||
|
if values[k] != v {
|
||||||
|
t.Errorf("%s = %q, want %q", k, values[k], v)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Without a log path LOG_FILE stays commented out.
|
||||||
|
values, err = ParseEnvFile(strings.NewReader(SampleEnv("v0.1.7", "")))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("sample without log path does not parse: %v", err)
|
||||||
|
}
|
||||||
|
if len(values) != 2 || values["LOG_FILE"] != "" {
|
||||||
|
t.Fatalf("active values = %v, want only CFG_VER and APP_VER", values)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSyncSample(t *testing.T) {
|
||||||
|
old := SampleEnv("v0.1.6", sampleLogPath)
|
||||||
|
// Simulate a setting that did not exist yet when the file was written.
|
||||||
|
old = strings.Replace(old, "#UPDATE_ASSET=gpu-turnstile.exe\n", "", 1)
|
||||||
|
|
||||||
|
out, changed := SyncSample(old, "v0.1.7", sampleLogPath)
|
||||||
|
if !changed {
|
||||||
|
t.Fatal("older installer file was not upgraded")
|
||||||
|
}
|
||||||
|
if !strings.Contains(out, "\nCFG_VER=v0.1.7\n") && !strings.HasPrefix(out, "CFG_VER=v0.1.7\n") {
|
||||||
|
t.Error("CFG_VER was not updated to the new version")
|
||||||
|
}
|
||||||
|
if !strings.Contains(out, "#UPDATE_ASSET=gpu-turnstile.exe") {
|
||||||
|
t.Error("missing setting was not appended")
|
||||||
|
}
|
||||||
|
if !strings.Contains(out, "LOG_FILE="+sampleLogPath) {
|
||||||
|
t.Error("existing active LOG_FILE was lost")
|
||||||
|
}
|
||||||
|
if !strings.Contains(out, "Added by gpu-turnstile v0.1.7") {
|
||||||
|
t.Error("appended section is not attributed")
|
||||||
|
}
|
||||||
|
if _, err := ParseEnvFile(strings.NewReader(out)); err != nil {
|
||||||
|
t.Fatalf("upgraded file does not parse: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// A file written by the same or a newer version is left alone.
|
||||||
|
if _, changed := SyncSample(SampleEnv("v0.1.7", sampleLogPath), "v0.1.7", sampleLogPath); changed {
|
||||||
|
t.Error("same-version file was modified")
|
||||||
|
}
|
||||||
|
if _, changed := SyncSample(SampleEnv("v0.2.0", sampleLogPath), "v0.1.7", sampleLogPath); changed {
|
||||||
|
t.Error("newer-version file was modified")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Files without CFG_VER are not installer-written; the installer
|
||||||
|
// replaces them, SyncSample leaves them alone.
|
||||||
|
user := "OLLAMA_URL=http://host:11434\n"
|
||||||
|
if out, changed := SyncSample(user, "v0.1.7", sampleLogPath); changed || out != user {
|
||||||
|
t.Error("file without CFG_VER was modified")
|
||||||
|
}
|
||||||
|
|
||||||
|
// A dev build never upgrades.
|
||||||
|
if _, changed := SyncSample(old, "dev", sampleLogPath); changed {
|
||||||
|
t.Error("dev build modified the file")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSyncSampleAppendsMissingAppVer(t *testing.T) {
|
||||||
|
old := SampleEnv("v0.1.6", sampleLogPath)
|
||||||
|
old = strings.Replace(old, "# "+appVerComment+"\n", "", 1)
|
||||||
|
old = strings.Replace(old, "APP_VER=stable\n", "", 1)
|
||||||
|
|
||||||
|
out, changed := SyncSample(old, "v0.1.7", sampleLogPath)
|
||||||
|
if !changed {
|
||||||
|
t.Fatal("file without APP_VER was not upgraded")
|
||||||
|
}
|
||||||
|
values, err := ParseEnvFile(strings.NewReader(out))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("upgraded file does not parse: %v", err)
|
||||||
|
}
|
||||||
|
if values["APP_VER"] != "stable" {
|
||||||
|
t.Fatalf("APP_VER = %q, want appended default \"stable\"", values["APP_VER"])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSampleEnvDevStampsConcreteVersion(t *testing.T) {
|
||||||
|
// A dev build must never write CFG_VER=dev: it stamps v0.0.0, and the
|
||||||
|
// next release install upgrades the file to a proper version.
|
||||||
|
dev := SampleEnv("dev", sampleLogPath)
|
||||||
|
if !strings.HasPrefix(dev, "CFG_VER=v0.0.0\n") {
|
||||||
|
t.Errorf("dev sample first line: %q", strings.SplitN(dev, "\n", 2)[0])
|
||||||
|
}
|
||||||
|
out, changed := SyncSample(dev, "v0.1.7", sampleLogPath)
|
||||||
|
if !changed || !strings.HasPrefix(out, "CFG_VER=v0.1.7\n") {
|
||||||
|
t.Errorf("dev-stamped file was not upgraded to v0.1.7 (changed=%v)", changed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompareVersions(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
a, b string
|
||||||
|
want int
|
||||||
|
}{
|
||||||
|
{"v0.1.6", "v0.1.7", -1},
|
||||||
|
{"v0.1.7", "v0.1.7", 0},
|
||||||
|
{"v1.0.0", "v0.9.9", 1},
|
||||||
|
{"0.1.7", "v0.1.6", 1},
|
||||||
|
{"dev", "v0.1.7", -1},
|
||||||
|
{"v0.1.7", "dev", 1},
|
||||||
|
{"dev", "dev", 0},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
if got := compareVersions(c.a, c.b); got != c.want {
|
||||||
|
t.Errorf("compareVersions(%q, %q) = %d, want %d", c.a, c.b, got, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,159 @@
|
|||||||
|
// Package game detects processes outside gpu-turnstile's control that hold
|
||||||
|
// the GPU — typically a game — so the proxy can block new GPU work and free
|
||||||
|
// VRAM while they run. Two detection paths: an explicit process watch list
|
||||||
|
// (GAME_PROCS) and a foreign-VRAM threshold via nvidia-smi
|
||||||
|
// (GPU_FOREIGN_VRAM_MB) that catches anything not on the ignore list.
|
||||||
|
package game
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"log/slog"
|
||||||
|
"os/exec"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Process is one running OS process.
|
||||||
|
type Process struct {
|
||||||
|
PID int
|
||||||
|
Name string
|
||||||
|
}
|
||||||
|
|
||||||
|
// computeApp is one process holding GPU memory, as reported by nvidia-smi.
|
||||||
|
type computeApp struct {
|
||||||
|
PID int
|
||||||
|
UsedMB int
|
||||||
|
}
|
||||||
|
|
||||||
|
// Detector checks whether a foreign process holds the GPU. The zero value
|
||||||
|
// (no watch list, no threshold) never detects anything; main only starts the
|
||||||
|
// poll loop when at least one path is configured.
|
||||||
|
type Detector struct {
|
||||||
|
procs map[string]bool // normalized names from GAME_PROCS
|
||||||
|
vramMB int // foreign VRAM threshold; 0 = disabled
|
||||||
|
ignore map[string]bool // normalized names never counted as foreign
|
||||||
|
log *slog.Logger
|
||||||
|
noNvidia bool // nvidia-smi was not found; VRAM path disabled for good
|
||||||
|
}
|
||||||
|
|
||||||
|
// New builds a Detector from the configured watch list, VRAM threshold in
|
||||||
|
// MiB (0 disables the nvidia-smi path) and ignore list. Names are matched
|
||||||
|
// case-insensitively, with or without a trailing ".exe".
|
||||||
|
func New(procs []string, vramMB int, ignore []string, log *slog.Logger) *Detector {
|
||||||
|
if log == nil {
|
||||||
|
log = slog.Default()
|
||||||
|
}
|
||||||
|
return &Detector{
|
||||||
|
procs: nameSet(procs),
|
||||||
|
vramMB: vramMB,
|
||||||
|
ignore: nameSet(ignore),
|
||||||
|
log: log,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// normName lowercases a process name and strips a trailing ".exe" so the
|
||||||
|
// watch and ignore lists match on Windows and Linux spellings alike.
|
||||||
|
func normName(s string) string {
|
||||||
|
return strings.TrimSuffix(strings.ToLower(strings.TrimSpace(s)), ".exe")
|
||||||
|
}
|
||||||
|
|
||||||
|
func nameSet(names []string) map[string]bool {
|
||||||
|
set := make(map[string]bool, len(names))
|
||||||
|
for _, n := range names {
|
||||||
|
if n = normName(n); n != "" {
|
||||||
|
set[n] = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return set
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check looks once for foreign GPU holders and returns a human-readable
|
||||||
|
// description of each (empty when the GPU is free for gpu-turnstile's
|
||||||
|
// consumers). A failing nvidia-smi call is returned as an error only when
|
||||||
|
// the process list found nothing; a missing nvidia-smi binary disables the
|
||||||
|
// VRAM path permanently (logged once).
|
||||||
|
func (d *Detector) Check(ctx context.Context) ([]string, error) {
|
||||||
|
ps, psErr := processes()
|
||||||
|
if d.vramMB <= 0 || d.noNvidia {
|
||||||
|
return d.detect(ps, nil), psErr
|
||||||
|
}
|
||||||
|
apps, err := queryComputeApps(ctx)
|
||||||
|
if errors.Is(err, exec.ErrNotFound) {
|
||||||
|
d.noNvidia = true
|
||||||
|
d.log.Warn("GPU_FOREIGN_VRAM_MB is set but nvidia-smi was not found; VRAM detection disabled")
|
||||||
|
return d.detect(ps, nil), nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return d.detect(ps, nil), err
|
||||||
|
}
|
||||||
|
return d.detect(ps, apps), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// detect is the pure core of Check: given the process table and (optionally)
|
||||||
|
// the nvidia-smi compute-apps list, it returns the foreign holders.
|
||||||
|
func (d *Detector) detect(ps []Process, apps []computeApp) []string {
|
||||||
|
var holders []string
|
||||||
|
for _, p := range ps {
|
||||||
|
if d.procs[normName(p.Name)] {
|
||||||
|
holders = append(holders, fmt.Sprintf("%s (pid %d)", p.Name, p.PID))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if d.vramMB > 0 && apps != nil {
|
||||||
|
names := make(map[int]string, len(ps))
|
||||||
|
for _, p := range ps {
|
||||||
|
names[p.PID] = p.Name
|
||||||
|
}
|
||||||
|
for _, a := range apps {
|
||||||
|
name := names[a.PID]
|
||||||
|
if d.ignore[normName(name)] || a.UsedMB < d.vramMB {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if name == "" {
|
||||||
|
name = "unknown process"
|
||||||
|
}
|
||||||
|
holders = append(holders, fmt.Sprintf("%s (pid %d) using %d MiB VRAM", name, a.PID, a.UsedMB))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return holders
|
||||||
|
}
|
||||||
|
|
||||||
|
// queryComputeApps runs nvidia-smi and parses the per-process VRAM list.
|
||||||
|
// Note: under Windows' WDDM driver, nvidia-smi only sees compute
|
||||||
|
// allocations, so graphics-only games may not appear there — GAME_PROCS is
|
||||||
|
// the reliable path on Windows; on Linux both work.
|
||||||
|
func queryComputeApps(ctx context.Context) ([]computeApp, error) {
|
||||||
|
out, err := exec.CommandContext(ctx, "nvidia-smi",
|
||||||
|
"--query-compute-apps=pid,used_memory", "--format=csv,noheader,nounits").Output()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return parseComputeApps(string(out))
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseComputeApps parses "pid, used_memory" CSV lines (no header, MiB
|
||||||
|
// units). Unsupported rows ("N/A" on WDDM) are skipped.
|
||||||
|
func parseComputeApps(out string) ([]computeApp, error) {
|
||||||
|
var apps []computeApp
|
||||||
|
for _, line := range strings.Split(out, "\n") {
|
||||||
|
line = strings.TrimSpace(line)
|
||||||
|
if line == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
pidStr, memStr, ok := strings.Cut(line, ",")
|
||||||
|
if !ok {
|
||||||
|
return nil, fmt.Errorf("nvidia-smi: unexpected line %q", line)
|
||||||
|
}
|
||||||
|
pid, err := strconv.Atoi(strings.TrimSpace(pidStr))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("nvidia-smi: unexpected pid in %q", line)
|
||||||
|
}
|
||||||
|
mem, err := strconv.Atoi(strings.TrimSpace(memStr))
|
||||||
|
if err != nil {
|
||||||
|
continue // "N/A" and friends: unsupported under WDDM
|
||||||
|
}
|
||||||
|
apps = append(apps, computeApp{PID: pid, UsedMB: mem})
|
||||||
|
}
|
||||||
|
return apps, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,96 @@
|
|||||||
|
package game
|
||||||
|
|
||||||
|
import (
|
||||||
|
"runtime"
|
||||||
|
"slices"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNormName(t *testing.T) {
|
||||||
|
for in, want := range map[string]string{
|
||||||
|
"Cyberpunk2077.exe": "cyberpunk2077",
|
||||||
|
"ollama": "ollama",
|
||||||
|
"OLLAMA APP.EXE": "ollama app",
|
||||||
|
" python.exe ": "python",
|
||||||
|
"hl2": "hl2",
|
||||||
|
} {
|
||||||
|
if got := normName(in); got != want {
|
||||||
|
t.Errorf("normName(%q) = %q, want %q", in, got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseComputeApps(t *testing.T) {
|
||||||
|
apps, err := parseComputeApps("1234, 512\n 42 , 8192 \n\n")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
want := []computeApp{{PID: 1234, UsedMB: 512}, {PID: 42, UsedMB: 8192}}
|
||||||
|
if !slices.Equal(apps, want) {
|
||||||
|
t.Errorf("got %+v, want %+v", apps, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
// WDDM reports "N/A" for memory; those rows are skipped, not fatal.
|
||||||
|
apps, err = parseComputeApps("1234, N/A\n42, 1024\n")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !slices.Equal(apps, []computeApp{{PID: 42, UsedMB: 1024}}) {
|
||||||
|
t.Errorf("got %+v", apps)
|
||||||
|
}
|
||||||
|
|
||||||
|
if _, err := parseComputeApps("garbage\n"); err == nil {
|
||||||
|
t.Error("expected an error for a malformed line")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDetect(t *testing.T) {
|
||||||
|
d := New([]string{"Cyberpunk2077.exe", "hl2"}, 1024,
|
||||||
|
[]string{"ollama", "python", "pythonw"}, nil)
|
||||||
|
ps := []Process{
|
||||||
|
{PID: 10, Name: "ollama.exe"},
|
||||||
|
{PID: 20, Name: "python.exe"},
|
||||||
|
{PID: 30, Name: "cyberpunk2077.exe"},
|
||||||
|
}
|
||||||
|
apps := []computeApp{
|
||||||
|
{PID: 10, UsedMB: 8192}, // ignored: ollama
|
||||||
|
{PID: 20, UsedMB: 4096}, // ignored: python (ComfyUI)
|
||||||
|
{PID: 40, UsedMB: 2048}, // foreign, above threshold
|
||||||
|
{PID: 50, UsedMB: 100}, // foreign but below threshold
|
||||||
|
}
|
||||||
|
holders := d.detect(ps, apps)
|
||||||
|
if len(holders) != 2 {
|
||||||
|
t.Fatalf("got %v, want 2 holders", holders)
|
||||||
|
}
|
||||||
|
if holders[0] != "cyberpunk2077.exe (pid 30)" {
|
||||||
|
t.Errorf("holders[0] = %q", holders[0])
|
||||||
|
}
|
||||||
|
if holders[1] != "unknown process (pid 40) using 2048 MiB VRAM" {
|
||||||
|
t.Errorf("holders[1] = %q", holders[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDetectNothingConfigured(t *testing.T) {
|
||||||
|
d := New(nil, 0, nil, nil)
|
||||||
|
if got := d.detect([]Process{{PID: 1, Name: "game.exe"}}, nil); len(got) != 0 {
|
||||||
|
t.Errorf("got %v, want none", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProcessesLive(t *testing.T) {
|
||||||
|
if runtime.GOOS != "windows" && runtime.GOOS != "linux" {
|
||||||
|
t.Skip("no process listing on this platform")
|
||||||
|
}
|
||||||
|
ps, err := processes()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(ps) == 0 {
|
||||||
|
t.Fatal("no processes listed")
|
||||||
|
}
|
||||||
|
for _, p := range ps {
|
||||||
|
if p.Name == "" {
|
||||||
|
t.Errorf("pid %d has an empty name", p.PID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
//go:build linux
|
||||||
|
|
||||||
|
package game
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// processes lists the running processes from /proc/<pid>/comm.
|
||||||
|
func processes() ([]Process, error) {
|
||||||
|
entries, err := os.ReadDir("/proc")
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var ps []Process
|
||||||
|
for _, e := range entries {
|
||||||
|
pid, err := strconv.Atoi(e.Name())
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
comm, err := os.ReadFile("/proc/" + e.Name() + "/comm")
|
||||||
|
if err != nil {
|
||||||
|
continue // process vanished mid-walk
|
||||||
|
}
|
||||||
|
ps = append(ps, Process{PID: pid, Name: strings.TrimSpace(string(comm))})
|
||||||
|
}
|
||||||
|
return ps, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,7 @@
|
|||||||
|
//go:build !windows && !linux
|
||||||
|
|
||||||
|
package game
|
||||||
|
|
||||||
|
// processes is unsupported on this platform; the process watch list never
|
||||||
|
// matches and the nvidia-smi path reports PIDs without names.
|
||||||
|
func processes() ([]Process, error) { return nil, nil }
|
||||||
@@ -0,0 +1,35 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
package game
|
||||||
|
|
||||||
|
import (
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
)
|
||||||
|
|
||||||
|
// processes lists the running processes via the Toolhelp32 snapshot API.
|
||||||
|
func processes() ([]Process, error) {
|
||||||
|
h, err := windows.CreateToolhelp32Snapshot(windows.TH32CS_SNAPPROCESS, 0)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer windows.CloseHandle(h) //nolint:errcheck // best effort
|
||||||
|
|
||||||
|
var entry windows.ProcessEntry32
|
||||||
|
entry.Size = uint32(unsafe.Sizeof(entry))
|
||||||
|
if err := windows.Process32First(h, &entry); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
var ps []Process
|
||||||
|
for {
|
||||||
|
ps = append(ps, Process{
|
||||||
|
PID: int(entry.ProcessID),
|
||||||
|
Name: windows.UTF16ToString(entry.ExeFile[:]),
|
||||||
|
})
|
||||||
|
if err := windows.Process32Next(h, &entry); err != nil {
|
||||||
|
break // ERROR_NO_MORE_FILES ends the walk
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return ps, nil
|
||||||
|
}
|
||||||
+57
-6
@@ -1,6 +1,8 @@
|
|||||||
// Package lock implements the two-mode GPU arbitration lock: any number of
|
// Package lock implements the two-mode GPU arbitration lock: any number of
|
||||||
// concurrent LLM requests ("readers") or exactly one image job ("writer"),
|
// concurrent LLM requests ("readers") or exactly one image job ("writer"),
|
||||||
// with image jobs taking priority over newly arriving LLM requests.
|
// with image jobs taking priority over newly arriving LLM requests. An
|
||||||
|
// external hold (SetExternal) blocks new grants of both kinds while a
|
||||||
|
// foreign process — e.g. a game — holds the GPU; in-flight work drains.
|
||||||
package lock
|
package lock
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -13,9 +15,10 @@ import (
|
|||||||
type State string
|
type State string
|
||||||
|
|
||||||
const (
|
const (
|
||||||
StateIdle State = "idle"
|
StateIdle State = "idle"
|
||||||
StateLLM State = "llm"
|
StateLLM State = "llm"
|
||||||
StateImage State = "image"
|
StateImage State = "image"
|
||||||
|
StateExternal State = "external"
|
||||||
)
|
)
|
||||||
|
|
||||||
type imageWaiter struct{ id uint64 }
|
type imageWaiter struct{ id uint64 }
|
||||||
@@ -30,6 +33,7 @@ type Lock struct {
|
|||||||
imageActive bool // an image job holds the GPU
|
imageActive bool // an image job holds the GPU
|
||||||
imageQ []imageWaiter
|
imageQ []imageWaiter
|
||||||
nextID uint64
|
nextID uint64
|
||||||
|
external string // non-empty: a foreign process (e.g. a game) holds the GPU
|
||||||
|
|
||||||
log *slog.Logger
|
log *slog.Logger
|
||||||
}
|
}
|
||||||
@@ -52,12 +56,41 @@ func (l *Lock) logTransition(msg string, args ...any) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetExternal records that a process outside gpu-turnstile's control (a
|
||||||
|
// game, another ML job) holds the GPU: new LLM and image grants block until
|
||||||
|
// ClearExternal. In-flight work is not preempted. holder describes the
|
||||||
|
// process for logs and busy responses.
|
||||||
|
func (l *Lock) SetExternal(holder string) {
|
||||||
|
l.mu.Lock()
|
||||||
|
l.external = holder
|
||||||
|
l.broadcast()
|
||||||
|
l.mu.Unlock()
|
||||||
|
l.logTransition("lock transition", "state", StateExternal, "holder", holder)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClearExternal lifts the external hold; waiting LLM and image requests
|
||||||
|
// proceed.
|
||||||
|
func (l *Lock) ClearExternal() {
|
||||||
|
l.mu.Lock()
|
||||||
|
l.external = ""
|
||||||
|
l.broadcast()
|
||||||
|
l.mu.Unlock()
|
||||||
|
l.logTransition("lock transition", "state", StateIdle)
|
||||||
|
}
|
||||||
|
|
||||||
|
// External returns the current external holder, or "" when none.
|
||||||
|
func (l *Lock) External() string {
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
return l.external
|
||||||
|
}
|
||||||
|
|
||||||
// AcquireLLM blocks until no image job is active or pending, then registers
|
// AcquireLLM blocks until no image job is active or pending, then registers
|
||||||
// one in-flight LLM request. Returns ctx.Err() if the context is cancelled
|
// one in-flight LLM request. Returns ctx.Err() if the context is cancelled
|
||||||
// while waiting; no state is changed in that case.
|
// while waiting; no state is changed in that case.
|
||||||
func (l *Lock) AcquireLLM(ctx context.Context) error {
|
func (l *Lock) AcquireLLM(ctx context.Context) error {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
for l.imageActive || len(l.imageQ) > 0 {
|
for l.imageActive || len(l.imageQ) > 0 || l.external != "" {
|
||||||
ch := l.change
|
ch := l.change
|
||||||
l.mu.Unlock()
|
l.mu.Unlock()
|
||||||
select {
|
select {
|
||||||
@@ -74,6 +107,22 @@ func (l *Lock) AcquireLLM(ctx context.Context) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TryAcquireLLM acquires one in-flight LLM slot without waiting and
|
||||||
|
// reports whether it succeeded. It fails when an image job is active or
|
||||||
|
// pending or an external hold is set.
|
||||||
|
func (l *Lock) TryAcquireLLM() bool {
|
||||||
|
l.mu.Lock()
|
||||||
|
if l.imageActive || len(l.imageQ) > 0 || l.external != "" {
|
||||||
|
l.mu.Unlock()
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
l.n++
|
||||||
|
n := l.n
|
||||||
|
l.mu.Unlock()
|
||||||
|
l.logTransition("lock transition", "state", StateLLM, "llm_inflight", n)
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
// ReleaseLLM marks one LLM request as finished.
|
// ReleaseLLM marks one LLM request as finished.
|
||||||
func (l *Lock) ReleaseLLM() {
|
func (l *Lock) ReleaseLLM() {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
@@ -104,7 +153,7 @@ func (l *Lock) AcquireImage(ctx context.Context) error {
|
|||||||
|
|
||||||
for {
|
for {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
if l.imageQ[0].id == w.id && l.n == 0 && !l.imageActive {
|
if l.imageQ[0].id == w.id && l.n == 0 && !l.imageActive && l.external == "" {
|
||||||
l.imageQ = l.imageQ[1:]
|
l.imageQ = l.imageQ[1:]
|
||||||
l.imageActive = true
|
l.imageActive = true
|
||||||
l.mu.Unlock()
|
l.mu.Unlock()
|
||||||
@@ -150,6 +199,8 @@ func (l *Lock) Snapshot() (state State, llmInflight int, imagePending bool) {
|
|||||||
state = StateImage
|
state = StateImage
|
||||||
case l.n > 0:
|
case l.n > 0:
|
||||||
state = StateLLM
|
state = StateLLM
|
||||||
|
case l.external != "":
|
||||||
|
state = StateExternal
|
||||||
default:
|
default:
|
||||||
state = StateIdle
|
state = StateIdle
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -226,3 +226,70 @@ func TestRace(t *testing.T) {
|
|||||||
t.Fatalf("leaked lock state: state=%s n=%d pending=%v", state, n, pending)
|
t.Fatalf("leaked lock state: state=%s n=%d pending=%v", state, n, pending)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestExternalHoldBlocksBoth(t *testing.T) {
|
||||||
|
lk := New(nil)
|
||||||
|
ctx := context.Background()
|
||||||
|
|
||||||
|
lk.SetExternal("game.exe (pid 42)")
|
||||||
|
if lk.TryAcquireLLM() {
|
||||||
|
t.Fatal("TryAcquireLLM succeeded during external hold")
|
||||||
|
}
|
||||||
|
if got := lk.External(); got != "game.exe (pid 42)" {
|
||||||
|
t.Fatalf("External() = %q", got)
|
||||||
|
}
|
||||||
|
if state, _, _ := lk.Snapshot(); state != StateExternal {
|
||||||
|
t.Fatalf("state=%s, want external", state)
|
||||||
|
}
|
||||||
|
|
||||||
|
llmAcquired := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
if err := lk.AcquireLLM(ctx); err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
}
|
||||||
|
close(llmAcquired)
|
||||||
|
}()
|
||||||
|
assertBlocked(t, llmAcquired, "LLM acquire during external hold")
|
||||||
|
|
||||||
|
imageAcquired := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
if err := lk.AcquireImage(ctx); err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
}
|
||||||
|
close(imageAcquired)
|
||||||
|
}()
|
||||||
|
assertBlocked(t, imageAcquired, "image acquire during external hold")
|
||||||
|
|
||||||
|
lk.ClearExternal()
|
||||||
|
// The queued image job wins over the LLM waiter (image priority).
|
||||||
|
waitFor(t, imageAcquired, "image acquire after ClearExternal")
|
||||||
|
assertBlocked(t, llmAcquired, "LLM acquire while image active")
|
||||||
|
lk.ReleaseImage()
|
||||||
|
waitFor(t, llmAcquired, "LLM acquire after image release")
|
||||||
|
lk.ReleaseLLM()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExternalHoldDoesNotPreempt(t *testing.T) {
|
||||||
|
lk := New(nil)
|
||||||
|
ctx := context.Background()
|
||||||
|
if err := lk.AcquireLLM(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
lk.SetExternal("game.exe")
|
||||||
|
// In-flight LLM work keeps the llm state; the hold blocks new grants.
|
||||||
|
if state, _, _ := lk.Snapshot(); state != StateLLM {
|
||||||
|
t.Fatalf("state=%s, want llm while work in flight", state)
|
||||||
|
}
|
||||||
|
if lk.TryAcquireLLM() {
|
||||||
|
t.Fatal("TryAcquireLLM succeeded during external hold")
|
||||||
|
}
|
||||||
|
lk.ReleaseLLM()
|
||||||
|
if state, _, _ := lk.Snapshot(); state != StateExternal {
|
||||||
|
t.Fatalf("state=%s, want external after drain", state)
|
||||||
|
}
|
||||||
|
lk.ClearExternal()
|
||||||
|
if state, _, _ := lk.Snapshot(); state != StateIdle {
|
||||||
|
t.Fatalf("state=%s, want idle", state)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
|
}
|
||||||
@@ -89,7 +89,7 @@ func (m *Metrics) Render(w io.Writer, state string, llmInflight int, imagePendin
|
|||||||
fmt.Fprint(w, `# HELP gpu_turnstile_state Current GPU state (1 for the active state).
|
fmt.Fprint(w, `# HELP gpu_turnstile_state Current GPU state (1 for the active state).
|
||||||
# TYPE gpu_turnstile_state gauge
|
# TYPE gpu_turnstile_state gauge
|
||||||
`)
|
`)
|
||||||
for _, s := range []string{"idle", "llm", "image"} {
|
for _, s := range []string{"idle", "llm", "image", "external"} {
|
||||||
v := 0
|
v := 0
|
||||||
if s == state {
|
if s == state {
|
||||||
v = 1
|
v = 1
|
||||||
|
|||||||
+444
-47
@@ -4,27 +4,35 @@
|
|||||||
package proxy
|
package proxy
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bufio"
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
|
"crypto/tls"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httputil"
|
"net/http/httputil"
|
||||||
"net/url"
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"gpu-turnstile/internal/comfy"
|
"gpu-turnstile/internal/comfy"
|
||||||
"gpu-turnstile/internal/lock"
|
"gpu-turnstile/internal/lock"
|
||||||
"gpu-turnstile/internal/metrics"
|
"gpu-turnstile/internal/metrics"
|
||||||
"gpu-turnstile/internal/ollama"
|
"gpu-turnstile/internal/ollama"
|
||||||
|
"gpu-turnstile/internal/supervise"
|
||||||
)
|
)
|
||||||
|
|
||||||
// captureLimit bounds how much of a /prompt response body is buffered while
|
// defaultCaptureLimit bounds how much of a /prompt response body is
|
||||||
// looking for prompt_id. The body still passes through to the client
|
// buffered while looking for prompt_id. The body still passes through to
|
||||||
// unchanged regardless of size.
|
// the client unchanged regardless of size.
|
||||||
const captureLimit = 64 * 1024
|
const defaultCaptureLimit = 64 * 1024
|
||||||
|
|
||||||
// Config wires a Server.
|
// Config wires a Server.
|
||||||
type Config struct {
|
type Config struct {
|
||||||
@@ -37,45 +45,271 @@ type Config struct {
|
|||||||
Metrics *metrics.Metrics
|
Metrics *metrics.Metrics
|
||||||
Log *slog.Logger
|
Log *slog.Logger
|
||||||
|
|
||||||
|
// ComfySup, when non-nil, is the managed ComfyUI process: any ComfyUI
|
||||||
|
// request starts it on demand, /prompt additionally waits for
|
||||||
|
// readiness before taking the GPU lock. nil = unmanaged upstream.
|
||||||
|
ComfySup *supervise.Process
|
||||||
|
|
||||||
|
// LogColor enables ANSI colors in per-request log lines. Ignored when
|
||||||
|
// the log level is above INFO (request lines are not emitted at all).
|
||||||
|
LogColor bool
|
||||||
|
// LogWriter receives the colored per-request lines; nil means stderr.
|
||||||
|
LogWriter io.Writer
|
||||||
|
|
||||||
LLMWaitTimeout time.Duration
|
LLMWaitTimeout time.Duration
|
||||||
UnloadTimeout time.Duration
|
UnloadTimeout time.Duration
|
||||||
JobTimeout time.Duration
|
JobTimeout time.Duration
|
||||||
WarmModel string
|
|
||||||
|
// LLMBusyMode is "wait" (default) or "reject". In reject mode an LLM
|
||||||
|
// request that arrives while an image job is active or pending is
|
||||||
|
// answered immediately with LLMBusyStatus and a Retry-After header
|
||||||
|
// (BusyRetryAfter seconds) instead of waiting for the lock. In wait
|
||||||
|
// mode the Retry-After header is sent when LLMWaitTimeout expires.
|
||||||
|
LLMBusyMode string
|
||||||
|
LLMBusyStatus int
|
||||||
|
BusyRetryAfter int
|
||||||
|
|
||||||
|
// BackoffInitial and BackoffMax control the exponential retry backoff
|
||||||
|
// when an upstream refuses a connection: the wait doubles from
|
||||||
|
// BackoffInitial up to BackoffMax between attempts. Zero selects the
|
||||||
|
// defaults (1s / 60s).
|
||||||
|
BackoffInitial time.Duration
|
||||||
|
BackoffMax time.Duration
|
||||||
|
|
||||||
|
// UnloadPollInterval and HistoryPollInterval override the clients'
|
||||||
|
// /api/ps and /history poll intervals when > 0.
|
||||||
|
UnloadPollInterval time.Duration
|
||||||
|
HistoryPollInterval time.Duration
|
||||||
|
// FreeTimeout and WarmTimeout bound the /free call and the warm-model
|
||||||
|
// reload; zero selects the defaults.
|
||||||
|
FreeTimeout time.Duration
|
||||||
|
WarmTimeout time.Duration
|
||||||
|
// PromptCaptureLimit overrides defaultCaptureLimit when > 0.
|
||||||
|
PromptCaptureLimit int64
|
||||||
|
|
||||||
|
WarmModel string
|
||||||
}
|
}
|
||||||
|
|
||||||
// Server serves both gpu-turnstile listeners.
|
// Server serves both gpu-turnstile listeners.
|
||||||
type Server struct {
|
type Server struct {
|
||||||
cfg Config
|
cfg Config
|
||||||
log *slog.Logger
|
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
|
ollamaProxy *httputil.ReverseProxy
|
||||||
comfyProxy *httputil.ReverseProxy
|
comfyProxy *httputil.ReverseProxy
|
||||||
}
|
}
|
||||||
|
|
||||||
// New builds a Server, validating the upstream URLs.
|
// New builds a Server, validating the upstream URLs. At least one of
|
||||||
|
// OllamaURL / ComfyURL must be set; an empty URL disables that consumer —
|
||||||
|
// its handler is then never served, its client may be nil, and the image
|
||||||
|
// job flow skips the Ollama unload/warm steps.
|
||||||
func New(cfg Config) (*Server, error) {
|
func New(cfg Config) (*Server, error) {
|
||||||
ollamaURL, err := url.Parse(cfg.OllamaURL)
|
if cfg.OllamaURL == "" && cfg.ComfyURL == "" {
|
||||||
if err != nil || ollamaURL.Scheme == "" || ollamaURL.Host == "" {
|
return nil, fmt.Errorf("at least one of OllamaURL or ComfyURL is required")
|
||||||
return nil, fmt.Errorf("invalid OLLAMA_URL %q", cfg.OllamaURL)
|
|
||||||
}
|
}
|
||||||
comfyURL, err := url.Parse(cfg.ComfyURL)
|
var ollamaURL, comfyURL *url.URL
|
||||||
if err != nil || comfyURL.Scheme == "" || comfyURL.Host == "" {
|
if cfg.OllamaURL != "" {
|
||||||
return nil, fmt.Errorf("invalid COMFY_URL %q", cfg.ComfyURL)
|
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
|
log := cfg.Log
|
||||||
if log == nil {
|
if log == nil {
|
||||||
log = slog.Default()
|
log = slog.Default()
|
||||||
}
|
}
|
||||||
return &Server{
|
if cfg.UnloadPollInterval > 0 && cfg.Ollama != nil {
|
||||||
cfg: cfg,
|
cfg.Ollama.PollInterval = cfg.UnloadPollInterval
|
||||||
log: log,
|
}
|
||||||
ollamaProxy: newReverseProxy(ollamaURL, log.With("upstream", "ollama")),
|
if cfg.HistoryPollInterval > 0 && cfg.Comfy != nil {
|
||||||
comfyProxy: newReverseProxy(comfyURL, log.With("upstream", "comfy")),
|
cfg.Comfy.PollInterval = cfg.HistoryPollInterval
|
||||||
}, nil
|
}
|
||||||
|
freeTimeout := cfg.FreeTimeout
|
||||||
|
if freeTimeout <= 0 {
|
||||||
|
freeTimeout = 30 * time.Second
|
||||||
|
}
|
||||||
|
warmTimeout := cfg.WarmTimeout
|
||||||
|
if warmTimeout <= 0 {
|
||||||
|
warmTimeout = 2 * time.Minute
|
||||||
|
}
|
||||||
|
captureLimit := cfg.PromptCaptureLimit
|
||||||
|
if captureLimit <= 0 {
|
||||||
|
captureLimit = defaultCaptureLimit
|
||||||
|
}
|
||||||
|
backoffInitial := cfg.BackoffInitial
|
||||||
|
if backoffInitial <= 0 {
|
||||||
|
backoffInitial = time.Second
|
||||||
|
}
|
||||||
|
backoffMax := cfg.BackoffMax
|
||||||
|
if backoffMax <= 0 {
|
||||||
|
backoffMax = time.Minute
|
||||||
|
}
|
||||||
|
busyMode := "wait"
|
||||||
|
if cfg.LLMBusyMode == "reject" {
|
||||||
|
busyMode = "reject"
|
||||||
|
}
|
||||||
|
busyStatus := cfg.LLMBusyStatus
|
||||||
|
if busyStatus == 0 {
|
||||||
|
busyStatus = http.StatusServiceUnavailable
|
||||||
|
}
|
||||||
|
busyRetryAfter := cfg.BusyRetryAfter
|
||||||
|
if busyRetryAfter <= 0 {
|
||||||
|
busyRetryAfter = 30
|
||||||
|
}
|
||||||
|
retry := &retryTransport{
|
||||||
|
base: http.DefaultTransport,
|
||||||
|
initial: backoffInitial,
|
||||||
|
max: backoffMax,
|
||||||
|
log: log,
|
||||||
|
}
|
||||||
|
s := &Server{
|
||||||
|
cfg: cfg,
|
||||||
|
log: log,
|
||||||
|
logWriter: cfg.LogWriter,
|
||||||
|
freeTimeout: freeTimeout,
|
||||||
|
warmTimeout: warmTimeout,
|
||||||
|
captureLimit: captureLimit,
|
||||||
|
backoffInitial: backoffInitial,
|
||||||
|
backoffMax: backoffMax,
|
||||||
|
busyMode: busyMode,
|
||||||
|
busyStatus: busyStatus,
|
||||||
|
busyRetryAfter: busyRetryAfter,
|
||||||
|
}
|
||||||
|
if ollamaURL != nil {
|
||||||
|
s.ollamaProxy = newReverseProxy(ollamaURL, retry, log.With("upstream", "ollama"))
|
||||||
|
}
|
||||||
|
if comfyURL != nil {
|
||||||
|
s.comfyProxy = newReverseProxy(comfyURL, retry, log.With("upstream", "comfy"))
|
||||||
|
}
|
||||||
|
return s, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func newReverseProxy(target *url.URL, log *slog.Logger) *httputil.ReverseProxy {
|
// retryTransport retries requests whose failure means the upstream never
|
||||||
|
// saw them — any dial-phase error (connection refused, dial timeout, DNS
|
||||||
|
// failure), TLS handshake errors — plus 5xx responses when the request
|
||||||
|
// body can be replayed (GETs and requests with GetBody set). The wait
|
||||||
|
// doubles from initial up to max between attempts. The loop runs until
|
||||||
|
// the request succeeds, fails in a non-retryable way, or the client's
|
||||||
|
// context is cancelled.
|
||||||
|
type retryTransport struct {
|
||||||
|
base http.RoundTripper
|
||||||
|
initial time.Duration
|
||||||
|
max time.Duration
|
||||||
|
log *slog.Logger
|
||||||
|
}
|
||||||
|
|
||||||
|
// shouldRetry reports whether a RoundTrip error means the request never
|
||||||
|
// reached the upstream application and is therefore safe to send again.
|
||||||
|
func shouldRetry(err error) bool {
|
||||||
|
// Dial-phase failures: refused, timeout, unreachable, DNS (wrapped).
|
||||||
|
var opErr *net.OpError
|
||||||
|
if errors.As(err, &opErr) && opErr.Op == "dial" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
var dnsErr *net.DNSError
|
||||||
|
if errors.As(err, &dnsErr) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
// TLS handshake failures: the HTTP request was never written.
|
||||||
|
var recordErr tls.RecordHeaderError
|
||||||
|
if errors.As(err, &recordErr) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
var certErr *tls.CertificateVerificationError
|
||||||
|
if errors.As(err, &certErr) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
var alertErr tls.AlertError
|
||||||
|
if errors.As(err, &alertErr) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// replayable reports whether the request body can be sent again. Bodies
|
||||||
|
// streamed from the client (GetBody == nil) cannot, so 5xx responses to
|
||||||
|
// POSTs are not retried: the upstream may have partially processed them,
|
||||||
|
// and re-sending could duplicate work (e.g. a second ComfyUI prompt).
|
||||||
|
func replayable(req *http.Request) bool {
|
||||||
|
return req.Body == nil || req.Body == http.NoBody || req.GetBody != nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *retryTransport) logAt(level slog.Level, msg string, args ...any) {
|
||||||
|
if t.log != nil && t.log.Enabled(context.Background(), level) {
|
||||||
|
t.log.Log(context.Background(), level, msg, args...)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *retryTransport) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||||
|
wait := t.initial
|
||||||
|
attempt := 0
|
||||||
|
for {
|
||||||
|
resp, err := t.base.RoundTrip(req)
|
||||||
|
switch {
|
||||||
|
case err != nil && !shouldRetry(err):
|
||||||
|
return nil, err
|
||||||
|
case err == nil && (resp.StatusCode < 500 || !replayable(req)):
|
||||||
|
return resp, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Retryable failure: a transport error, or a 5xx response.
|
||||||
|
var reason string
|
||||||
|
if err != nil {
|
||||||
|
reason = err.Error()
|
||||||
|
} else {
|
||||||
|
reason = resp.Status
|
||||||
|
io.Copy(io.Discard, resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
if req.GetBody != nil {
|
||||||
|
if body, berr := req.GetBody(); berr == nil {
|
||||||
|
req.Body = body
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
attempt++
|
||||||
|
// One WARN per outage episode; subsequent attempts at INFO.
|
||||||
|
level := slog.LevelInfo
|
||||||
|
if attempt == 1 {
|
||||||
|
level = slog.LevelWarn
|
||||||
|
}
|
||||||
|
t.logAt(level, "upstream unavailable; retrying with backoff",
|
||||||
|
"path", req.URL.Path, "reason", reason, "attempt", attempt, "retry_in", wait)
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-req.Context().Done():
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return nil, req.Context().Err()
|
||||||
|
case <-time.After(wait):
|
||||||
|
}
|
||||||
|
wait *= 2
|
||||||
|
if wait > t.max {
|
||||||
|
wait = t.max
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newReverseProxy(target *url.URL, transport http.RoundTripper, log *slog.Logger) *httputil.ReverseProxy {
|
||||||
return &httputil.ReverseProxy{
|
return &httputil.ReverseProxy{
|
||||||
|
Transport: transport,
|
||||||
Rewrite: func(pr *httputil.ProxyRequest) {
|
Rewrite: func(pr *httputil.ProxyRequest) {
|
||||||
pr.SetURL(target)
|
pr.SetURL(target)
|
||||||
pr.SetXForwarded()
|
pr.SetXForwarded()
|
||||||
@@ -106,6 +340,94 @@ func (s *Server) writeMetrics(w http.ResponseWriter) {
|
|||||||
s.cfg.Metrics.Render(w, string(state), n, pending)
|
s.cfg.Metrics.Render(w, string(state), n, pending)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ANSI colors for per-request log lines.
|
||||||
|
const (
|
||||||
|
ansiReset = "\x1b[0m"
|
||||||
|
ansiCyan = "\x1b[36m"
|
||||||
|
ansiGreen = "\x1b[32m"
|
||||||
|
ansiYellow = "\x1b[33m"
|
||||||
|
ansiRed = "\x1b[31m"
|
||||||
|
)
|
||||||
|
|
||||||
|
// statusRecorder remembers the response status while passing everything
|
||||||
|
// through, including streaming flushes and websocket hijacks.
|
||||||
|
type statusRecorder struct {
|
||||||
|
http.ResponseWriter
|
||||||
|
status int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *statusRecorder) WriteHeader(code int) {
|
||||||
|
r.status = code
|
||||||
|
r.ResponseWriter.WriteHeader(code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *statusRecorder) Flush() {
|
||||||
|
if f, ok := r.ResponseWriter.(http.Flusher); ok {
|
||||||
|
f.Flush()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *statusRecorder) Hijack() (net.Conn, *bufio.ReadWriter, error) {
|
||||||
|
h, ok := r.ResponseWriter.(http.Hijacker)
|
||||||
|
if !ok {
|
||||||
|
return nil, nil, errors.New("response writer does not support hijacking")
|
||||||
|
}
|
||||||
|
return h.Hijack()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *statusRecorder) Unwrap() http.ResponseWriter { return r.ResponseWriter }
|
||||||
|
|
||||||
|
// reqLine emits one request log line. With color enabled slog cannot be
|
||||||
|
// used: its text handler escapes the ANSI sequences, so the line is
|
||||||
|
// written to stderr directly in the same key=value shape. Without color
|
||||||
|
// it is a plain slog INFO line.
|
||||||
|
func (s *Server) reqLine(log *slog.Logger, code, line string, attrs ...any) {
|
||||||
|
if !s.cfg.LogColor {
|
||||||
|
log.Info(line, attrs...)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString("time=" + time.Now().Format("2006-01-02T15:04:05.000Z07:00") + " level=INFO ")
|
||||||
|
sb.WriteString(code + line + ansiReset)
|
||||||
|
for i := 0; i+1 < len(attrs); i += 2 {
|
||||||
|
fmt.Fprintf(&sb, " %v=%v", attrs[i], attrs[i+1])
|
||||||
|
}
|
||||||
|
sb.WriteByte('\n')
|
||||||
|
w := s.logWriter
|
||||||
|
if w == nil {
|
||||||
|
w = os.Stderr
|
||||||
|
}
|
||||||
|
io.WriteString(w, sb.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
// logRequests logs one line per incoming request and one per completed
|
||||||
|
// response at INFO level, colored when enabled: cyan "-->" for incoming,
|
||||||
|
// green/yellow/red "<--" for responses by status class. At log levels
|
||||||
|
// above INFO it is a pass-through.
|
||||||
|
func (s *Server) logRequests(listener string, next http.Handler) http.Handler {
|
||||||
|
log := s.log.With("listener", listener)
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if !log.Enabled(r.Context(), slog.LevelInfo) {
|
||||||
|
next.ServeHTTP(w, r)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
start := time.Now()
|
||||||
|
s.reqLine(log, ansiCyan, "--> "+r.Method+" "+r.URL.RequestURI(),
|
||||||
|
"listener", listener, "remote", r.RemoteAddr)
|
||||||
|
rec := &statusRecorder{ResponseWriter: w, status: http.StatusOK}
|
||||||
|
next.ServeHTTP(rec, r)
|
||||||
|
code := ansiGreen
|
||||||
|
switch {
|
||||||
|
case rec.status >= 500:
|
||||||
|
code = ansiRed
|
||||||
|
case rec.status >= 400:
|
||||||
|
code = ansiYellow
|
||||||
|
}
|
||||||
|
s.reqLine(log, code, "<-- "+strconv.Itoa(rec.status)+" "+r.Method+" "+r.URL.RequestURI(),
|
||||||
|
"listener", listener, "ms", time.Since(start).Milliseconds())
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
// llmPaths are the Ollama endpoints that load models into VRAM and therefore
|
// llmPaths are the Ollama endpoints that load models into VRAM and therefore
|
||||||
// take the LLM lock. Everything else passes through unlocked.
|
// take the LLM lock. Everything else passes through unlocked.
|
||||||
var llmPaths = map[string]bool{
|
var llmPaths = map[string]bool{
|
||||||
@@ -122,9 +444,9 @@ func isLLMRequest(r *http.Request) bool {
|
|||||||
return r.Method == http.MethodPost && llmPaths[r.URL.Path]
|
return r.Method == http.MethodPost && llmPaths[r.URL.Path]
|
||||||
}
|
}
|
||||||
|
|
||||||
// OllamaHandler serves the Ollama-facing listener.
|
// OllamaHandler serves the listener for Ollama-compatible clients.
|
||||||
func (s *Server) OllamaHandler() http.Handler {
|
func (s *Server) OllamaHandler() http.Handler {
|
||||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
return s.logRequests("ollama", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
switch r.URL.Path {
|
switch r.URL.Path {
|
||||||
case "/healthz":
|
case "/healthz":
|
||||||
s.writeHealthz(w)
|
s.writeHealthz(w)
|
||||||
@@ -139,42 +461,81 @@ func (s *Server) OllamaHandler() http.Handler {
|
|||||||
}
|
}
|
||||||
|
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
|
if s.busyMode == "reject" {
|
||||||
|
if !s.cfg.Lock.TryAcquireLLM() {
|
||||||
|
s.cfg.Metrics.ObserveLockWait("llm", time.Since(start).Seconds())
|
||||||
|
msg := "GPU busy: image job active or queued"
|
||||||
|
if holder := s.cfg.Lock.External(); holder != "" {
|
||||||
|
msg = "GPU busy: " + holder
|
||||||
|
}
|
||||||
|
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, msg, s.busyStatus)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
s.cfg.Metrics.ObserveLockWait("llm", time.Since(start).Seconds())
|
||||||
|
defer s.cfg.Lock.ReleaseLLM()
|
||||||
|
s.ollamaProxy.ServeHTTP(w, r)
|
||||||
|
return
|
||||||
|
}
|
||||||
wctx, cancel := context.WithTimeout(r.Context(), s.cfg.LLMWaitTimeout)
|
wctx, cancel := context.WithTimeout(r.Context(), s.cfg.LLMWaitTimeout)
|
||||||
err := s.cfg.Lock.AcquireLLM(wctx)
|
err := s.cfg.Lock.AcquireLLM(wctx)
|
||||||
cancel()
|
cancel()
|
||||||
s.cfg.Metrics.ObserveLockWait("llm", time.Since(start).Seconds())
|
s.cfg.Metrics.ObserveLockWait("llm", time.Since(start).Seconds())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if errors.Is(err, context.DeadlineExceeded) && r.Context().Err() == nil {
|
if errors.Is(err, context.DeadlineExceeded) && r.Context().Err() == nil {
|
||||||
|
w.Header().Set("Retry-After", strconv.Itoa(s.busyRetryAfter))
|
||||||
http.Error(w, "GPU busy: timed out waiting for the lock", http.StatusServiceUnavailable)
|
http.Error(w, "GPU busy: timed out waiting for the lock", http.StatusServiceUnavailable)
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
defer s.cfg.Lock.ReleaseLLM()
|
defer s.cfg.Lock.ReleaseLLM()
|
||||||
s.ollamaProxy.ServeHTTP(w, r)
|
s.ollamaProxy.ServeHTTP(w, r)
|
||||||
})
|
}))
|
||||||
}
|
}
|
||||||
|
|
||||||
// ComfyHandler serves the ComfyUI-facing listener.
|
// ComfyHandler serves the listener for ComfyUI clients.
|
||||||
func (s *Server) ComfyHandler() http.Handler {
|
func (s *Server) ComfyHandler() http.Handler {
|
||||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
return s.logRequests("comfy", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
if r.URL.Path == "/healthz" {
|
switch r.URL.Path {
|
||||||
|
case "/healthz":
|
||||||
s.writeHealthz(w)
|
s.writeHealthz(w)
|
||||||
return
|
return
|
||||||
|
case "/metrics":
|
||||||
|
s.writeMetrics(w)
|
||||||
|
return
|
||||||
}
|
}
|
||||||
if r.Method == http.MethodPost && r.URL.Path == "/prompt" {
|
if r.Method == http.MethodPost && r.URL.Path == "/prompt" {
|
||||||
s.handlePrompt(w, r)
|
s.handlePrompt(w, r)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if s.cfg.ComfySup != nil {
|
||||||
|
// Any other ComfyUI request also wakes the managed server; the
|
||||||
|
// retry backoff bridges the time it needs to come up. While a
|
||||||
|
// foreign process holds the GPU we refuse to spawn it — the
|
||||||
|
// request gets the busy answer instead of fighting for VRAM.
|
||||||
|
if holder := s.cfg.Lock.External(); holder != "" && !s.cfg.ComfySup.Running() {
|
||||||
|
w.Header().Set("Retry-After", strconv.Itoa(s.busyRetryAfter))
|
||||||
|
http.Error(w, "GPU busy: "+holder, http.StatusServiceUnavailable)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := s.cfg.ComfySup.EnsureRunning(); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("cannot start ComfyUI: %v", err), http.StatusBadGateway)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
s.comfyProxy.ServeHTTP(w, r)
|
s.comfyProxy.ServeHTTP(w, r)
|
||||||
})
|
}))
|
||||||
}
|
}
|
||||||
|
|
||||||
// captureWriter passes the response through unchanged while recording the
|
// captureWriter passes the response through unchanged while recording the
|
||||||
// status code and the first captureLimit bytes of the body.
|
// status code and the first limit bytes of the body.
|
||||||
type captureWriter struct {
|
type captureWriter struct {
|
||||||
http.ResponseWriter
|
http.ResponseWriter
|
||||||
status int
|
status int
|
||||||
buf bytes.Buffer
|
buf bytes.Buffer
|
||||||
|
limit int64
|
||||||
}
|
}
|
||||||
|
|
||||||
func (w *captureWriter) WriteHeader(code int) {
|
func (w *captureWriter) WriteHeader(code int) {
|
||||||
@@ -183,7 +544,7 @@ func (w *captureWriter) WriteHeader(code int) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (w *captureWriter) Write(p []byte) (int, error) {
|
func (w *captureWriter) Write(p []byte) (int, error) {
|
||||||
if w.buf.Len() < captureLimit {
|
if int64(w.buf.Len()) < w.limit {
|
||||||
w.buf.Write(p)
|
w.buf.Write(p)
|
||||||
}
|
}
|
||||||
return w.ResponseWriter.Write(p)
|
return w.ResponseWriter.Write(p)
|
||||||
@@ -203,6 +564,23 @@ func (w *captureWriter) Unwrap() http.ResponseWriter { return w.ResponseWriter }
|
|||||||
func (s *Server) handlePrompt(w http.ResponseWriter, r *http.Request) {
|
func (s *Server) handlePrompt(w http.ResponseWriter, r *http.Request) {
|
||||||
log := s.log.With("op", "image")
|
log := s.log.With("op", "image")
|
||||||
|
|
||||||
|
// Normally the managed server is brought up *before* taking the GPU
|
||||||
|
// lock: torch can take a minute to load, and LLM traffic should keep
|
||||||
|
// flowing in the meantime. While a foreign process holds the GPU, LLM
|
||||||
|
// traffic is blocked anyway and a fresh ComfyUI would fight it for
|
||||||
|
// VRAM — so the lock comes first in that case.
|
||||||
|
comfyFirst := s.cfg.ComfySup != nil && s.cfg.Lock.External() == ""
|
||||||
|
if comfyFirst {
|
||||||
|
if err := s.cfg.ComfySup.EnsureRunning(); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("cannot start ComfyUI: %v", err), http.StatusBadGateway)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if err := s.cfg.ComfySup.WaitReady(r.Context()); err != nil {
|
||||||
|
http.Error(w, fmt.Sprintf("ComfyUI did not become ready: %v", err), http.StatusBadGateway)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
if err := s.cfg.Lock.AcquireImage(r.Context()); err != nil {
|
if err := s.cfg.Lock.AcquireImage(r.Context()); err != nil {
|
||||||
if errors.Is(err, context.DeadlineExceeded) {
|
if errors.Is(err, context.DeadlineExceeded) {
|
||||||
@@ -213,22 +591,37 @@ func (s *Server) handlePrompt(w http.ResponseWriter, r *http.Request) {
|
|||||||
s.cfg.Metrics.ObserveLockWait("image", time.Since(start).Seconds())
|
s.cfg.Metrics.ObserveLockWait("image", time.Since(start).Seconds())
|
||||||
log.Info("image lock acquired")
|
log.Info("image lock acquired")
|
||||||
|
|
||||||
uctx, ucancel := context.WithTimeout(r.Context(), s.cfg.UnloadTimeout)
|
if s.cfg.ComfySup != nil && !comfyFirst {
|
||||||
elapsed, uerr := s.cfg.Ollama.UnloadAll(uctx)
|
if err := s.cfg.ComfySup.EnsureRunning(); err != nil {
|
||||||
ucancel()
|
s.cfg.Lock.ReleaseImage()
|
||||||
s.cfg.Metrics.ObserveUnload(elapsed.Seconds())
|
http.Error(w, fmt.Sprintf("cannot start ComfyUI: %v", err), http.StatusBadGateway)
|
||||||
switch {
|
return
|
||||||
case r.Context().Err() != nil:
|
}
|
||||||
s.cfg.Lock.ReleaseImage()
|
if err := s.cfg.ComfySup.WaitReady(r.Context()); err != nil {
|
||||||
return
|
s.cfg.Lock.ReleaseImage()
|
||||||
case uerr != nil:
|
http.Error(w, fmt.Sprintf("ComfyUI did not become ready: %v", err), http.StatusBadGateway)
|
||||||
// Degrade, don't fail the user's request on a misbehaving neighbour.
|
return
|
||||||
log.Warn("ollama unload incomplete; continuing", "err", uerr)
|
}
|
||||||
default:
|
|
||||||
log.Info("ollama models unloaded", "seconds", elapsed.Seconds())
|
|
||||||
}
|
}
|
||||||
|
|
||||||
cw := &captureWriter{ResponseWriter: w, status: http.StatusOK}
|
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}
|
||||||
s.comfyProxy.ServeHTTP(cw, r)
|
s.comfyProxy.ServeHTTP(cw, r)
|
||||||
|
|
||||||
var accepted struct {
|
var accepted struct {
|
||||||
@@ -262,7 +655,7 @@ func (s *Server) finishImageJob(promptID string) {
|
|||||||
log.Info("image job completed")
|
log.Info("image job completed")
|
||||||
}
|
}
|
||||||
|
|
||||||
freeCtx, freeCancel := context.WithTimeout(context.Background(), 30*time.Second)
|
freeCtx, freeCancel := context.WithTimeout(context.Background(), s.freeTimeout)
|
||||||
if err := s.cfg.Comfy.Free(freeCtx); err != nil {
|
if err := s.cfg.Comfy.Free(freeCtx); err != nil {
|
||||||
log.Warn("failed to free ComfyUI models", "err", err)
|
log.Warn("failed to free ComfyUI models", "err", err)
|
||||||
}
|
}
|
||||||
@@ -270,10 +663,14 @@ func (s *Server) finishImageJob(promptID string) {
|
|||||||
|
|
||||||
s.cfg.Lock.ReleaseImage()
|
s.cfg.Lock.ReleaseImage()
|
||||||
log.Info("image lock released")
|
log.Info("image lock released")
|
||||||
|
if s.cfg.ComfySup != nil {
|
||||||
|
// The idle clock starts when the job ends, not when it began.
|
||||||
|
s.cfg.ComfySup.NoteActivity()
|
||||||
|
}
|
||||||
|
|
||||||
if s.cfg.WarmModel != "" {
|
if s.cfg.Ollama != nil && s.cfg.WarmModel != "" {
|
||||||
if state, _, _ := s.cfg.Lock.Snapshot(); state == lock.StateIdle {
|
if state, _, _ := s.cfg.Lock.Snapshot(); state == lock.StateIdle {
|
||||||
wctx, wcancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
wctx, wcancel := context.WithTimeout(context.Background(), s.warmTimeout)
|
||||||
if err := s.cfg.Ollama.Warm(wctx, s.cfg.WarmModel); err != nil {
|
if err := s.cfg.Ollama.Warm(wctx, s.cfg.WarmModel); err != nil {
|
||||||
log.Warn("warm model reload failed", "model", s.cfg.WarmModel, "err", err)
|
log.Warn("warm model reload failed", "model", s.cfg.WarmModel, "err", err)
|
||||||
} else {
|
} else {
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package proxy
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bufio"
|
"bufio"
|
||||||
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -45,8 +46,8 @@ type fakes struct {
|
|||||||
rec *recorder
|
rec *recorder
|
||||||
ollama *httptest.Server
|
ollama *httptest.Server
|
||||||
comfy *httptest.Server
|
comfy *httptest.Server
|
||||||
server *httptest.Server // Ollama-facing gpu-turnstile listener
|
server *httptest.Server // gpu-turnstile listener for Ollama-compatible clients
|
||||||
comfySrv *httptest.Server // ComfyUI-facing gpu-turnstile listener
|
comfySrv *httptest.Server // gpu-turnstile listener for ComfyUI clients
|
||||||
freeCh chan struct{}
|
freeCh chan struct{}
|
||||||
chatCh chan struct{}
|
chatCh chan struct{}
|
||||||
historyMu sync.Mutex
|
historyMu sync.Mutex
|
||||||
@@ -345,3 +346,187 @@ func TestPassThroughNoLock(t *testing.T) {
|
|||||||
t.Fatalf("pass-through = %d %s", resp.StatusCode, body)
|
t.Fatalf("pass-through = %d %s", resp.StatusCode, body)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLLMBusyReject(t *testing.T) {
|
||||||
|
f := newFakes(t)
|
||||||
|
|
||||||
|
lk := lock.New(nil)
|
||||||
|
if err := lk.AcquireImage(context.Background()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
ollamaClient, err := ollama.New(f.ollama.URL, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
comfyClient, err := comfy.New(f.comfy.URL, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
srv, err := New(Config{
|
||||||
|
OllamaURL: f.ollama.URL,
|
||||||
|
ComfyURL: f.comfy.URL,
|
||||||
|
Lock: lk,
|
||||||
|
Ollama: ollamaClient,
|
||||||
|
Comfy: comfyClient,
|
||||||
|
Metrics: metrics.New(),
|
||||||
|
LLMWaitTimeout: 2 * time.Second,
|
||||||
|
LLMBusyMode: "reject",
|
||||||
|
BusyRetryAfter: 17,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
front := httptest.NewServer(srv.OllamaHandler())
|
||||||
|
defer front.Close()
|
||||||
|
|
||||||
|
start := time.Now()
|
||||||
|
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{}`))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
resp.Body.Close()
|
||||||
|
if resp.StatusCode != http.StatusServiceUnavailable {
|
||||||
|
t.Fatalf("busy chat status = %d %s", resp.StatusCode, body)
|
||||||
|
}
|
||||||
|
if got := resp.Header.Get("Retry-After"); got != "17" {
|
||||||
|
t.Fatalf("Retry-After = %q", got)
|
||||||
|
}
|
||||||
|
if elapsed := time.Since(start); elapsed > time.Second {
|
||||||
|
t.Fatalf("reject was not immediate: %v", elapsed)
|
||||||
|
}
|
||||||
|
|
||||||
|
// After the image lock is released the next LLM request goes through.
|
||||||
|
lk.ReleaseImage()
|
||||||
|
resp, err = http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{}`))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
resp.Body.Close()
|
||||||
|
if resp.StatusCode != 200 {
|
||||||
|
t.Fatalf("chat after release status = %d", resp.StatusCode)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestLLMBusyWaitTimeoutRetryAfter(t *testing.T) {
|
||||||
|
f := newFakes(t)
|
||||||
|
|
||||||
|
lk := lock.New(nil)
|
||||||
|
if err := lk.AcquireImage(context.Background()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer lk.ReleaseImage()
|
||||||
|
ollamaClient, err := ollama.New(f.ollama.URL, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
comfyClient, err := comfy.New(f.comfy.URL, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
srv, err := New(Config{
|
||||||
|
OllamaURL: f.ollama.URL,
|
||||||
|
ComfyURL: f.comfy.URL,
|
||||||
|
Lock: lk,
|
||||||
|
Ollama: ollamaClient,
|
||||||
|
Comfy: comfyClient,
|
||||||
|
Metrics: metrics.New(),
|
||||||
|
LLMWaitTimeout: 50 * time.Millisecond,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
front := httptest.NewServer(srv.OllamaHandler())
|
||||||
|
defer front.Close()
|
||||||
|
|
||||||
|
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{}`))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
resp.Body.Close()
|
||||||
|
if resp.StatusCode != http.StatusServiceUnavailable {
|
||||||
|
t.Fatalf("timed-out chat status = %d", resp.StatusCode)
|
||||||
|
}
|
||||||
|
if got := resp.Header.Get("Retry-After"); got != "30" {
|
||||||
|
t.Fatalf("Retry-After = %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestComfyOnlyModeSkipsOllama(t *testing.T) {
|
||||||
|
f := newFakes(t)
|
||||||
|
|
||||||
|
comfyClient, err := comfy.New(f.comfy.URL, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
comfyClient.PollInterval = 5 * time.Millisecond
|
||||||
|
srv, err := New(Config{
|
||||||
|
ComfyURL: f.comfy.URL,
|
||||||
|
Lock: lock.New(nil),
|
||||||
|
Comfy: comfyClient,
|
||||||
|
Metrics: metrics.New(),
|
||||||
|
JobTimeout: 2 * time.Second,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
front := httptest.NewServer(srv.ComfyHandler())
|
||||||
|
defer front.Close()
|
||||||
|
|
||||||
|
resp, err := http.Post(front.URL+"/prompt", "application/json", strings.NewReader(`{}`))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
resp.Body.Close()
|
||||||
|
if resp.StatusCode != 200 {
|
||||||
|
t.Fatalf("prompt status = %d", resp.StatusCode)
|
||||||
|
}
|
||||||
|
f.completeJob()
|
||||||
|
select {
|
||||||
|
case <-f.freeCh:
|
||||||
|
case <-time.After(3 * time.Second):
|
||||||
|
t.Fatal("/free never called")
|
||||||
|
}
|
||||||
|
// With Ollama disabled the unload steps must not happen.
|
||||||
|
if i := f.rec.index("unload"); i >= 0 {
|
||||||
|
f.rec.mu.Lock()
|
||||||
|
t.Fatalf("unload called with ollama disabled; events: %v", f.rec.events)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOllamaOnlyMode(t *testing.T) {
|
||||||
|
f := newFakes(t)
|
||||||
|
|
||||||
|
ollamaClient, err := ollama.New(f.ollama.URL, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
srv, err := New(Config{
|
||||||
|
OllamaURL: f.ollama.URL,
|
||||||
|
Lock: lock.New(nil),
|
||||||
|
Ollama: ollamaClient,
|
||||||
|
Metrics: metrics.New(),
|
||||||
|
LLMWaitTimeout: time.Second,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
front := httptest.NewServer(srv.OllamaHandler())
|
||||||
|
defer front.Close()
|
||||||
|
|
||||||
|
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{}`))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
resp.Body.Close()
|
||||||
|
if resp.StatusCode != 200 {
|
||||||
|
t.Fatalf("chat status = %d", resp.StatusCode)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewRequiresConsumer(t *testing.T) {
|
||||||
|
_, err := New(Config{Lock: lock.New(nil), Metrics: metrics.New()})
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("New with no upstream URLs should fail")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,207 @@
|
|||||||
|
package proxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/tls"
|
||||||
|
"io"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"syscall"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// stubTransport fails with ECONNREFUSED for the first fails requests, then
|
||||||
|
// returns a 200 response.
|
||||||
|
type stubTransport struct {
|
||||||
|
fails int
|
||||||
|
calls int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *stubTransport) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||||
|
s.calls++
|
||||||
|
if s.calls <= s.fails {
|
||||||
|
return nil, &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ECONNREFUSED}
|
||||||
|
}
|
||||||
|
return &http.Response{
|
||||||
|
StatusCode: 200,
|
||||||
|
Body: io.NopCloser(strings.NewReader("ok")),
|
||||||
|
Header: make(http.Header),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRetryTransportBackoff(t *testing.T) {
|
||||||
|
st := &stubTransport{fails: 3}
|
||||||
|
rt := &retryTransport{
|
||||||
|
base: st,
|
||||||
|
initial: 10 * time.Millisecond,
|
||||||
|
max: 25 * time.Millisecond,
|
||||||
|
}
|
||||||
|
|
||||||
|
start := time.Now()
|
||||||
|
req, _ := http.NewRequest(http.MethodGet, "http://upstream/api/version", nil)
|
||||||
|
resp, err := rt.RoundTrip(req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
resp.Body.Close()
|
||||||
|
if st.calls != 4 {
|
||||||
|
t.Fatalf("calls = %d, want 4", st.calls)
|
||||||
|
}
|
||||||
|
// Waits: 10ms + 20ms + 25ms (capped) = 55ms minimum.
|
||||||
|
elapsed := time.Since(start)
|
||||||
|
if elapsed < 50*time.Millisecond {
|
||||||
|
t.Fatalf("elapsed = %v, want >= ~55ms of backoff", elapsed)
|
||||||
|
}
|
||||||
|
if elapsed > 5*time.Second {
|
||||||
|
t.Fatalf("elapsed = %v, suspiciously long", elapsed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRetryTransportNonRefusedErrorNotRetried(t *testing.T) {
|
||||||
|
rt := &retryTransport{
|
||||||
|
base: &stubTransport{fails: 0},
|
||||||
|
initial: time.Millisecond,
|
||||||
|
max: time.Millisecond,
|
||||||
|
}
|
||||||
|
req, _ := http.NewRequest(http.MethodGet, "http://upstream/", nil)
|
||||||
|
resp, err := rt.RoundTrip(req)
|
||||||
|
if err != nil || resp.StatusCode != 200 {
|
||||||
|
t.Fatalf("resp=%v err=%v", resp, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type alwaysRefused struct{ calls int }
|
||||||
|
|
||||||
|
func (a *alwaysRefused) RoundTrip(*http.Request) (*http.Response, error) {
|
||||||
|
a.calls++
|
||||||
|
return nil, &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ECONNREFUSED}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRetryTransportContextCancel(t *testing.T) {
|
||||||
|
ar := &alwaysRefused{}
|
||||||
|
rt := &retryTransport{base: ar, initial: time.Second, max: time.Second}
|
||||||
|
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
|
||||||
|
defer cancel()
|
||||||
|
req, _ := http.NewRequestWithContext(ctx, http.MethodGet, "http://upstream/", nil)
|
||||||
|
start := time.Now()
|
||||||
|
_, err := rt.RoundTrip(req)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error after context cancellation")
|
||||||
|
}
|
||||||
|
if time.Since(start) > 2*time.Second {
|
||||||
|
t.Fatal("retry loop did not stop on context cancellation")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRetryViaReverseProxy(t *testing.T) {
|
||||||
|
// Point the proxy at a port nothing listens on; the request should be
|
||||||
|
// retried (not instantly 502) until the client context ends.
|
||||||
|
srv, err := New(Config{
|
||||||
|
OllamaURL: "http://127.0.0.1:1",
|
||||||
|
ComfyURL: "http://127.0.0.1:1",
|
||||||
|
Lock: nil,
|
||||||
|
Metrics: nil,
|
||||||
|
BackoffInitial: 10 * time.Millisecond,
|
||||||
|
BackoffMax: 20 * time.Millisecond,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_ = srv // construction must not panic with minimal config
|
||||||
|
|
||||||
|
rt := &retryTransport{base: http.DefaultTransport, initial: 10 * time.Millisecond, max: 20 * time.Millisecond}
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond)
|
||||||
|
defer cancel()
|
||||||
|
req, _ := http.NewRequestWithContext(ctx, http.MethodGet, "http://127.0.0.1:1/", nil)
|
||||||
|
_, err = rt.RoundTrip(req)
|
||||||
|
if err == nil || !shouldRetry(err) {
|
||||||
|
t.Fatalf("err = %v, want connection refused", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// flakyStatus returns 500 for the first fails requests, then 200.
|
||||||
|
type flakyStatus struct {
|
||||||
|
fails int
|
||||||
|
calls int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *flakyStatus) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||||
|
s.calls++
|
||||||
|
code := 200
|
||||||
|
if s.calls <= s.fails {
|
||||||
|
code = 500
|
||||||
|
}
|
||||||
|
return &http.Response{
|
||||||
|
StatusCode: code,
|
||||||
|
Status: http.StatusText(code),
|
||||||
|
Body: io.NopCloser(strings.NewReader("")),
|
||||||
|
Header: make(http.Header),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRetryTransport5xxGet(t *testing.T) {
|
||||||
|
fs := &flakyStatus{fails: 2}
|
||||||
|
rt := &retryTransport{base: fs, initial: time.Millisecond, max: 2 * time.Millisecond}
|
||||||
|
|
||||||
|
req, _ := http.NewRequest(http.MethodGet, "http://upstream/api/version", nil)
|
||||||
|
resp, err := rt.RoundTrip(req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
resp.Body.Close()
|
||||||
|
if resp.StatusCode != 200 {
|
||||||
|
t.Fatalf("status = %d, want 200", resp.StatusCode)
|
||||||
|
}
|
||||||
|
if fs.calls != 3 {
|
||||||
|
t.Fatalf("calls = %d, want 3", fs.calls)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRetryTransport5xxPostNotRetried(t *testing.T) {
|
||||||
|
fs := &flakyStatus{fails: 10}
|
||||||
|
rt := &retryTransport{base: fs, initial: time.Millisecond, max: time.Millisecond}
|
||||||
|
|
||||||
|
// A streamed body (no GetBody) must not be replayed after a 500.
|
||||||
|
req, _ := http.NewRequest(http.MethodPost, "http://upstream/prompt", io.NopCloser(strings.NewReader("{}")))
|
||||||
|
resp, err := rt.RoundTrip(req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
resp.Body.Close()
|
||||||
|
if resp.StatusCode != 500 {
|
||||||
|
t.Fatalf("status = %d, want 500", resp.StatusCode)
|
||||||
|
}
|
||||||
|
if fs.calls != 1 {
|
||||||
|
t.Fatalf("calls = %d, want 1 (no retry for streamed POST)", fs.calls)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestShouldRetryClassification(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
err error
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"dial refused", &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ECONNREFUSED}, true},
|
||||||
|
{"dial timeout", &net.OpError{Op: "dial", Net: "tcp", Err: timeoutErr{}}, true},
|
||||||
|
{"dns", &net.DNSError{Err: "no such host", IsNotFound: true}, true},
|
||||||
|
{"tls record", tls.RecordHeaderError{Msg: "bad"}, true},
|
||||||
|
{"tls alert", tls.AlertError(42), true},
|
||||||
|
{"read error mid-request", &net.OpError{Op: "read", Net: "tcp", Err: syscall.ECONNRESET}, false},
|
||||||
|
{"plain error", io.EOF, false},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
if got := shouldRetry(c.err); got != c.want {
|
||||||
|
t.Errorf("%s: shouldRetry = %v, want %v", c.name, got, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type timeoutErr struct{}
|
||||||
|
|
||||||
|
func (timeoutErr) Error() string { return "i/o timeout" }
|
||||||
|
func (timeoutErr) Timeout() bool { return true }
|
||||||
|
func (timeoutErr) Temporary() bool { return true }
|
||||||
@@ -0,0 +1,63 @@
|
|||||||
|
// Shared env-file handling for the installers (Windows and Linux).
|
||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"gpu-turnstile/internal/config"
|
||||||
|
)
|
||||||
|
|
||||||
|
// syncEnvFile ensures the installed env file at path is a current
|
||||||
|
// installer-written sample. A missing file is created; a file without a
|
||||||
|
// CFG_VER line (hand-written or from before versioning) is invalid and
|
||||||
|
// replaced by a fresh sample after a .bak backup; a file written by an
|
||||||
|
// older version gets newly added settings appended via config.SyncSample.
|
||||||
|
// logPath activates the LOG_FILE line (Windows); empty leaves it commented.
|
||||||
|
func syncEnvFile(path, version, logPath string) error {
|
||||||
|
data, err := os.ReadFile(path)
|
||||||
|
switch {
|
||||||
|
case os.IsNotExist(err):
|
||||||
|
return writeEnvFile(path, config.SampleEnv(version, logPath))
|
||||||
|
case err != nil:
|
||||||
|
return fmt.Errorf("read %s: %w", path, err)
|
||||||
|
}
|
||||||
|
values, perr := config.ParseEnvFile(strings.NewReader(string(data)))
|
||||||
|
if perr == nil && values["CFG_VER"] != "" {
|
||||||
|
synced, changed := config.SyncSample(string(data), version, logPath)
|
||||||
|
if !changed {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return writeEnvFile(path, synced)
|
||||||
|
}
|
||||||
|
// No readable CFG_VER: the file is invalid — keep a backup and start
|
||||||
|
// from a fresh sample.
|
||||||
|
if err := os.WriteFile(path+".bak", data, 0o644); err != nil {
|
||||||
|
return fmt.Errorf("back up %s: %w", path, err)
|
||||||
|
}
|
||||||
|
return writeEnvFile(path, config.SampleEnv(version, logPath))
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeEnvFile(path, content string) error {
|
||||||
|
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
|
||||||
|
return fmt.Errorf("write %s: %w", path, err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// configuredValue reads one key from the env file the service will load, so
|
||||||
|
// the installers can adapt the sandbox to it (ACL grants, unit directives).
|
||||||
|
// "" when unset or unreadable.
|
||||||
|
func configuredValue(configPath, key 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[key]
|
||||||
|
}
|
||||||
@@ -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,272 @@
|
|||||||
|
//go:build linux
|
||||||
|
|
||||||
|
// Package service integrates gpu-turnstile with systemd on Linux: running
|
||||||
|
// under a unit with readiness notification and watchdog, plus
|
||||||
|
// install/remove helpers that manage a hardened system unit.
|
||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"os/signal"
|
||||||
|
"path/filepath"
|
||||||
|
"syscall"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Name matches the Windows service name; the systemd unit is Name + ".service".
|
||||||
|
const Name = "gpu-turnstile"
|
||||||
|
|
||||||
|
// unitPath is where Install writes the unit file.
|
||||||
|
const unitPath = "/etc/systemd/system/" + Name + ".service"
|
||||||
|
|
||||||
|
// stateDir holds the installed binary (and staged updates); the unit's
|
||||||
|
// StateDirectory= directive makes systemd own it and grant the dynamic
|
||||||
|
// user write access. etcConfig is the config file the unit loads.
|
||||||
|
const (
|
||||||
|
stateDir = "/var/lib/" + Name
|
||||||
|
etcConfig = "/etc/" + Name + ".env"
|
||||||
|
)
|
||||||
|
|
||||||
|
// IsService reports whether the process was started by systemd.
|
||||||
|
func IsService() bool { return os.Getenv("INVOCATION_ID") != "" }
|
||||||
|
|
||||||
|
// Elevated is always true on Linux: there is no UAC equivalent; privilege
|
||||||
|
// errors surface from the failing operation with a "run as root" hint.
|
||||||
|
func Elevated() bool { return true }
|
||||||
|
|
||||||
|
// RelaunchElevated is a Windows-only concept (UAC).
|
||||||
|
func RelaunchElevated([]string) (int, error) {
|
||||||
|
return 0, fmt.Errorf("elevated relaunch is only supported on Windows")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Run executes run with SIGINT/SIGTERM cancellation (which is how systemctl
|
||||||
|
// stop signals the process) and tells systemd when the shutdown begins.
|
||||||
|
func Run(run func(ctx context.Context) error) error {
|
||||||
|
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
|
||||||
|
defer stop()
|
||||||
|
defer NotifyStopping()
|
||||||
|
return run(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
// renderUnit builds the hardened systemd unit: Type=notify so systemctl
|
||||||
|
// start blocks until the listeners are bound, a 30s watchdog, and
|
||||||
|
// restart-on-failure with a 5s delay — which is also what brings up a
|
||||||
|
// staged update after the updater exits with a non-zero code.
|
||||||
|
//
|
||||||
|
// Sandboxing mirrors the Windows virtual account: DynamicUser=yes gives
|
||||||
|
// the service a transient per-service UID with no login and no home, the
|
||||||
|
// filesystem is read-only except StateDirectory (the install dir, so
|
||||||
|
// self-updates can rewrite the binary), and the usual no-privilege-escalation
|
||||||
|
// directives apply. The proxy needs nothing but outbound TCP/UDP and the
|
||||||
|
// notify socket, so it loses nothing. A managed ComfyUI (comfyDir) gets a
|
||||||
|
// BindPaths hole through ProtectHome/ProtectSystem: it reads its venv and
|
||||||
|
// writes output/temp/user data under COMFY_DIR. A venv whose base
|
||||||
|
// interpreter (pyvenv.cfg home) lives outside COMFY_DIR gets an additional
|
||||||
|
// read-only bind.
|
||||||
|
func renderUnit(exePath, configPath, comfyDir string) string {
|
||||||
|
bind := ""
|
||||||
|
if comfyDir != "" {
|
||||||
|
bind = "BindPaths=" + comfyDir + "\n"
|
||||||
|
if home := comfyVenvHome(comfyDir); home != "" {
|
||||||
|
bind += "BindReadOnlyPaths=" + home + "\n"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
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
|
||||||
|
%sProtectSystem=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, bind)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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. When the config sets COMFY_DIR, the unit gets a
|
||||||
|
// BindPaths= for it so the sandboxed service can reach the managed ComfyUI
|
||||||
|
// even under /home. Needs root.
|
||||||
|
//
|
||||||
|
// Re-running install converges an existing unit instead of failing: it is
|
||||||
|
// stopped first if active, the installed binary is replaced only when the
|
||||||
|
// content differs, the unit file is rewritten (followed by daemon-reload)
|
||||||
|
// only when it changed, and the service is started again only if it was
|
||||||
|
// active before.
|
||||||
|
//
|
||||||
|
// The binary lives in the StateDirectory rather than /usr/local/sbin on
|
||||||
|
// purpose: replacing a running binary needs write access to its
|
||||||
|
// *directory*, and granting the sandboxed service user write access to a
|
||||||
|
// shared system directory would let a compromised service overwrite other
|
||||||
|
// binaries. /var/lib/gpu-turnstile is exclusively ours.
|
||||||
|
func Install(configPath string, copyBin bool, version string) error {
|
||||||
|
exe, err := os.Executable()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if abs, absErr := filepath.Abs(exe); absErr == nil {
|
||||||
|
exe = abs
|
||||||
|
}
|
||||||
|
unit := Name + ".service"
|
||||||
|
_, statErr := os.Stat(unitPath)
|
||||||
|
fresh := os.IsNotExist(statErr)
|
||||||
|
wasRunning := exec.Command("systemctl", "is-active", "--quiet", unit).Run() == nil
|
||||||
|
if wasRunning {
|
||||||
|
if out, err := exec.Command("systemctl", "stop", unit).CombinedOutput(); err != nil {
|
||||||
|
return fmt.Errorf("systemctl stop (run as root): %w (%s)", err, out)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg := etcConfig
|
||||||
|
if copyBin {
|
||||||
|
if err := os.MkdirAll(stateDir, 0o755); err != nil {
|
||||||
|
return fmt.Errorf("create %s (run as root): %w", stateDir, err)
|
||||||
|
}
|
||||||
|
installedExe := filepath.Join(stateDir, Name)
|
||||||
|
if exe != installedExe {
|
||||||
|
if same, _ := sameFileContent(exe, installedExe); !same {
|
||||||
|
if err := copyFile(exe, installedExe, 0o755); err != nil {
|
||||||
|
return fmt.Errorf("install binary to %s: %w", installedExe, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
exe = installedExe
|
||||||
|
if _, err := os.Stat(etcConfig); os.IsNotExist(err) && configPath != "" {
|
||||||
|
copyFile(configPath, etcConfig, 0o644) //nolint:errcheck // best effort
|
||||||
|
}
|
||||||
|
// Missing configs get a fully commented sample (LOG_FILE stays
|
||||||
|
// commented — stderr goes to the journal on Linux);
|
||||||
|
// installer-written ones from older versions get new settings
|
||||||
|
// appended; anything without CFG_VER is invalid and gets replaced
|
||||||
|
// (backup kept as .bak).
|
||||||
|
if err := syncEnvFile(etcConfig, version, ""); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
} else if configPath != "" {
|
||||||
|
if abs, absErr := filepath.Abs(configPath); absErr == nil {
|
||||||
|
cfg = abs
|
||||||
|
}
|
||||||
|
}
|
||||||
|
rendered := renderUnit(exe, cfg, configuredValue(cfg, "COMFY_DIR"))
|
||||||
|
if old, _ := os.ReadFile(unitPath); string(old) != rendered {
|
||||||
|
if err := os.WriteFile(unitPath, []byte(rendered), 0o644); err != nil {
|
||||||
|
return fmt.Errorf("write %s (run as root): %w", unitPath, err)
|
||||||
|
}
|
||||||
|
if out, err := exec.Command("systemctl", "daemon-reload").CombinedOutput(); err != nil {
|
||||||
|
return fmt.Errorf("systemctl daemon-reload: %w (%s)", err, out)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if fresh {
|
||||||
|
if out, err := exec.Command("systemctl", "enable", "--now", unit).CombinedOutput(); err != nil {
|
||||||
|
return fmt.Errorf("systemctl enable --now: %w (%s)", err, out)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if exec.Command("systemctl", "is-enabled", "--quiet", unit).Run() != nil {
|
||||||
|
if out, err := exec.Command("systemctl", "enable", unit).CombinedOutput(); err != nil {
|
||||||
|
return fmt.Errorf("systemctl enable: %w (%s)", err, out)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if wasRunning {
|
||||||
|
if out, err := exec.Command("systemctl", "start", unit).CombinedOutput(); err != nil {
|
||||||
|
return fmt.Errorf("systemctl start: %w (%s)", err, out)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// sameFileContent reports whether two files hold identical bytes. A missing
|
||||||
|
// destination is simply "different".
|
||||||
|
func sameFileContent(a, b string) (bool, error) {
|
||||||
|
ba, err := os.ReadFile(a)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
bb, err := os.ReadFile(b)
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
return bytes.Equal(ba, bb), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Remove stops and disables the service and deletes the unit file and the
|
||||||
|
// installed binary. The config file in /etc is left in place (user data).
|
||||||
|
func Remove() error {
|
||||||
|
exec.Command("systemctl", "disable", "--now", Name+".service").Run() // ignore: may not exist
|
||||||
|
if err := os.Remove(unitPath); err != nil && !os.IsNotExist(err) {
|
||||||
|
return fmt.Errorf("remove %s: %w", unitPath, err)
|
||||||
|
}
|
||||||
|
os.RemoveAll(stateDir) // installed binary + staged updates; ignore error
|
||||||
|
if out, err := exec.Command("systemctl", "daemon-reload").CombinedOutput(); err != nil {
|
||||||
|
return fmt.Errorf("systemctl daemon-reload: %w (%s)", err, out)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RestartIfRunning restarts the systemd unit when it is active (used after
|
||||||
|
// a forced update staged a new binary). Reports whether a restart
|
||||||
|
// happened. An inactive or missing unit is not an error. Needs root.
|
||||||
|
func RestartIfRunning() (bool, error) {
|
||||||
|
if err := exec.Command("systemctl", "is-active", "--quiet", Name+".service").Run(); err != nil {
|
||||||
|
return false, nil // inactive or not installed
|
||||||
|
}
|
||||||
|
if out, err := exec.Command("systemctl", "restart", Name+".service").CombinedOutput(); err != nil {
|
||||||
|
return false, fmt.Errorf("systemctl restart (run as root): %w (%s)", err, out)
|
||||||
|
}
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,37 @@
|
|||||||
|
//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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if strings.Contains(unit, "BindPaths") {
|
||||||
|
t.Fatalf("unit without COMFY_DIR must not bind anything:\n%s", unit)
|
||||||
|
}
|
||||||
|
|
||||||
|
unit = renderUnit("/var/lib/gpu-turnstile/gpu-turnstile", "/etc/gpu-turnstile.env", "/home/gpu/ComfyUI")
|
||||||
|
if !strings.Contains(unit, "BindPaths=/home/gpu/ComfyUI\n") {
|
||||||
|
t.Fatalf("unit with COMFY_DIR must bind it:\n%s", 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, string) error { return errUnsupported }
|
||||||
|
|
||||||
|
// Remove is unsupported on non-Windows, non-Linux platforms.
|
||||||
|
func Remove() error { return errUnsupported }
|
||||||
|
|
||||||
|
// Elevated is always true here: there is no UAC concept, and privilege
|
||||||
|
// errors surface from the failing operation with a "run as root" hint.
|
||||||
|
func Elevated() bool { return true }
|
||||||
|
|
||||||
|
// RelaunchElevated is unsupported on non-Windows, non-Linux platforms.
|
||||||
|
func RelaunchElevated([]string) (int, error) { return 0, errUnsupported }
|
||||||
|
|
||||||
|
// RestartIfRunning is a no-op on platforms without service integration.
|
||||||
|
func RestartIfRunning() (bool, error) { return false, nil }
|
||||||
@@ -0,0 +1,549 @@
|
|||||||
|
//go:build windows
|
||||||
|
|
||||||
|
// Package service integrates gpu-turnstile with the Windows Service
|
||||||
|
// Control Manager: running as a service with graceful stop, plus
|
||||||
|
// install/remove helpers. Installed services always run as the virtual
|
||||||
|
// account NT SERVICE\gpu-turnstile — a per-service low-privilege identity
|
||||||
|
// managed by the SCM, with no password and no admin rights.
|
||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"syscall"
|
||||||
|
"time"
|
||||||
|
"unsafe"
|
||||||
|
|
||||||
|
"golang.org/x/sys/windows"
|
||||||
|
"golang.org/x/sys/windows/svc"
|
||||||
|
"golang.org/x/sys/windows/svc/mgr"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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), read access to the
|
||||||
|
// config file if it lives elsewhere, and — when the config sets COMFY_DIR —
|
||||||
|
// recursive modify access to the managed ComfyUI's install tree, which may
|
||||||
|
// live inside a user profile the account otherwise cannot enter. The grants
|
||||||
|
// must come after CreateService: the virtual account's SID only exists once
|
||||||
|
// the service is registered.
|
||||||
|
//
|
||||||
|
// Re-running install on an existing service converges instead of failing:
|
||||||
|
// the service is stopped first if running (so the binary can be replaced),
|
||||||
|
// the installed binary is refreshed only when the content differs, the
|
||||||
|
// registration is updated only where it drifted, and the service is
|
||||||
|
// started again only if it was running before.
|
||||||
|
func Install(configPath string, copyBin bool, version string) error {
|
||||||
|
exe, err := os.Executable()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if abs, absErr := filepath.Abs(exe); absErr == nil {
|
||||||
|
exe = abs
|
||||||
|
}
|
||||||
|
if configPath != "" {
|
||||||
|
if abs, absErr := filepath.Abs(configPath); absErr == nil {
|
||||||
|
configPath = abs
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
m, err := mgr.Connect()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("connect to service manager (run as administrator): %w", err)
|
||||||
|
}
|
||||||
|
defer m.Disconnect()
|
||||||
|
|
||||||
|
// An existing service is converged, not an error. Stop it first so the
|
||||||
|
// binary copy can be replaced, and remember whether to start it again.
|
||||||
|
var s *mgr.Service
|
||||||
|
wasRunning := false
|
||||||
|
if existing, openErr := m.OpenService(Name); openErr == nil {
|
||||||
|
s = existing
|
||||||
|
defer s.Close()
|
||||||
|
if st, qErr := s.Query(); qErr == nil &&
|
||||||
|
(st.State == svc.Running || st.State == svc.StartPending) {
|
||||||
|
wasRunning = true
|
||||||
|
fmt.Println("stopping the running gpu-turnstile service")
|
||||||
|
if err := stopAndWait(s); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
installDir, dataDir := installDirs()
|
||||||
|
if copyBin {
|
||||||
|
targetCfg := filepath.Join(installDir, "gpu-turnstile.env")
|
||||||
|
if !strings.EqualFold(filepath.Dir(exe), installDir) {
|
||||||
|
if err := os.MkdirAll(installDir, 0o755); err != nil {
|
||||||
|
return fmt.Errorf("create %s: %w", installDir, err)
|
||||||
|
}
|
||||||
|
installedExe := filepath.Join(installDir, "gpu-turnstile.exe")
|
||||||
|
if same, _ := sameFileContent(exe, installedExe); !same {
|
||||||
|
fmt.Printf("installing %s\n", installedExe)
|
||||||
|
if err := copyFile(exe, installedExe); err != nil {
|
||||||
|
return fmt.Errorf("copy binary to %s: %w", installedExe, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
exe = installedExe
|
||||||
|
if configPath != "" && !strings.EqualFold(configPath, targetCfg) {
|
||||||
|
if _, statErr := os.Stat(targetCfg); os.IsNotExist(statErr) {
|
||||||
|
copyFile(configPath, targetCfg) //nolint:errcheck // best effort
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// A service has no console: without LOG_FILE the output vanishes, so
|
||||||
|
// the sample pre-wires it to ProgramData. Missing configs get a
|
||||||
|
// fresh sample; installer-written ones from older versions get new
|
||||||
|
// settings appended; anything without CFG_VER is invalid and gets
|
||||||
|
// replaced (backup kept as .bak).
|
||||||
|
logPath := filepath.Join(dataDir, "gpu-turnstile.log")
|
||||||
|
if err := syncEnvFile(targetCfg, version, logPath); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
configPath = targetCfg
|
||||||
|
}
|
||||||
|
binPath := fmt.Sprintf(`"%s" -config "%s"`, exe, configPath)
|
||||||
|
|
||||||
|
if s == nil {
|
||||||
|
s, err = m.CreateService(Name, binPath, mgr.Config{
|
||||||
|
StartType: mgr.StartAutomatic,
|
||||||
|
DisplayName: "gpu-turnstile",
|
||||||
|
Description: "GPU arbitration proxy for Ollama and ComfyUI",
|
||||||
|
ServiceStartName: virtualAccount,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("create service: %w", err)
|
||||||
|
}
|
||||||
|
defer s.Close()
|
||||||
|
|
||||||
|
if err := ensureRecovery(s); err != nil {
|
||||||
|
s.Delete() // roll back so a retry starts clean
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := grantAll(exe, configPath); err != nil {
|
||||||
|
s.Delete() // roll back so a retry starts clean
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
// Best effort: start now instead of waiting for the next boot. A
|
||||||
|
// missing config (no consumer URLs) fails the start; the service stays
|
||||||
|
// registered and can be started once the config exists.
|
||||||
|
fmt.Println("starting the gpu-turnstile service")
|
||||||
|
s.Start()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Existing service: update the registration only where it drifted.
|
||||||
|
cur, err := s.Config()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("query service config: %w", err)
|
||||||
|
}
|
||||||
|
const displayName = "gpu-turnstile"
|
||||||
|
const description = "GPU arbitration proxy for Ollama and ComfyUI"
|
||||||
|
if cur.BinaryPathName != binPath ||
|
||||||
|
cur.StartType != mgr.StartAutomatic ||
|
||||||
|
cur.ServiceStartName != virtualAccount ||
|
||||||
|
cur.DisplayName != displayName ||
|
||||||
|
cur.Description != description {
|
||||||
|
upd := cur
|
||||||
|
upd.BinaryPathName = binPath
|
||||||
|
upd.StartType = mgr.StartAutomatic
|
||||||
|
upd.ServiceStartName = virtualAccount
|
||||||
|
upd.DisplayName = displayName
|
||||||
|
upd.Description = description
|
||||||
|
if err := s.UpdateConfig(upd); err != nil {
|
||||||
|
return fmt.Errorf("update service config: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := ensureRecovery(s); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
// Idempotent: re-assert the virtual account's ACLs (granting an
|
||||||
|
// existing ACE is a no-op).
|
||||||
|
if err := grantAll(exe, configPath); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if wasRunning {
|
||||||
|
fmt.Println("starting the gpu-turnstile service")
|
||||||
|
if err := s.Start(); err != nil {
|
||||||
|
return fmt.Errorf("start service: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ensureRecovery makes sure the service restarts 5s after a failure — the
|
||||||
|
// mechanism that also brings up a staged update. The current actions are
|
||||||
|
// queried first so an up-to-date service is left untouched.
|
||||||
|
func ensureRecovery(s *mgr.Service) error {
|
||||||
|
want := mgr.RecoveryAction{Type: mgr.ServiceRestart, Delay: 5 * time.Second}
|
||||||
|
actions, err := s.RecoveryActions()
|
||||||
|
onFailure, flagErr := s.RecoveryActionsOnNonCrashFailures()
|
||||||
|
if err == nil && flagErr == nil && onFailure && len(actions) == 3 {
|
||||||
|
ok := true
|
||||||
|
for _, a := range actions {
|
||||||
|
if a.Type != want.Type || a.Delay != want.Delay {
|
||||||
|
ok = false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if ok {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := s.SetRecoveryActions([]mgr.RecoveryAction{want, want, want}, 24*60*60); err != nil {
|
||||||
|
return fmt.Errorf("set recovery actions: %w", err)
|
||||||
|
}
|
||||||
|
if err := s.SetRecoveryActionsOnNonCrashFailures(true); err != nil {
|
||||||
|
return fmt.Errorf("set failure actions flag: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// stopAndWait stops the service and waits up to 30s for the stopped state.
|
||||||
|
func stopAndWait(s *mgr.Service) error {
|
||||||
|
if _, err := s.Control(svc.Stop); err != nil {
|
||||||
|
return fmt.Errorf("stop service: %w", err)
|
||||||
|
}
|
||||||
|
deadline := time.Now().Add(30 * time.Second)
|
||||||
|
for {
|
||||||
|
st, err := s.Query()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("query service: %w", err)
|
||||||
|
}
|
||||||
|
if st.State == svc.Stopped {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if time.Now().After(deadline) {
|
||||||
|
return fmt.Errorf("service did not stop within 30s")
|
||||||
|
}
|
||||||
|
time.Sleep(300 * time.Millisecond)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// sameFileContent reports whether two files hold identical bytes. A missing
|
||||||
|
// destination is simply "different".
|
||||||
|
func sameFileContent(a, b string) (bool, error) {
|
||||||
|
ba, err := os.ReadFile(a)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
bb, err := os.ReadFile(b)
|
||||||
|
if os.IsNotExist(err) {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
return bytes.Equal(ba, bb), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// copyFile copies src to dst (0755 on the new file).
|
||||||
|
func copyFile(src, dst string) error {
|
||||||
|
in, err := os.Open(src)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer in.Close()
|
||||||
|
out, err := os.OpenFile(dst, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, 0o755)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if _, err := io.Copy(out, in); err != nil {
|
||||||
|
out.Close()
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return out.Close()
|
||||||
|
}
|
||||||
|
|
||||||
|
// grantAll gives the virtual account every ACL the service needs: modify
|
||||||
|
// on the install and ProgramData directories and the LOG_FILE directory
|
||||||
|
// (if configured elsewhere), read on a config file outside the install
|
||||||
|
// directory.
|
||||||
|
func grantAll(exe, configPath string) error {
|
||||||
|
_, dataDir := installDirs()
|
||||||
|
if err := os.MkdirAll(dataDir, 0o755); err != nil {
|
||||||
|
return fmt.Errorf("create %s: %w", dataDir, err)
|
||||||
|
}
|
||||||
|
if err := grantAccess(dataDir, "(OI)(CI)(M)"); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
exeDir := filepath.Dir(exe)
|
||||||
|
if err := grantAccess(exeDir, "(OI)(CI)(M)"); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if configPath != "" && !strings.HasPrefix(strings.ToLower(configPath), strings.ToLower(exeDir)+`\`) {
|
||||||
|
if err := grantAccess(configPath, "(R)"); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if logFile := configuredValue(configPath, "LOG_FILE"); logFile != "" {
|
||||||
|
dir := filepath.Dir(logFile)
|
||||||
|
if err := os.MkdirAll(dir, 0o755); err == nil {
|
||||||
|
if err := grantAccess(dir, "(OI)(CI)(M)"); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
// A managed ComfyUI whose install lives somewhere the virtual account
|
||||||
|
// may not go (a user profile) needs an explicit grant — recursively,
|
||||||
|
// since ComfyUI also writes output/temp/user data next to its code. A
|
||||||
|
// missing directory is skipped: the startup warning covers it.
|
||||||
|
if comfyDir := configuredValue(configPath, "COMFY_DIR"); comfyDir != "" {
|
||||||
|
if _, err := os.Stat(comfyDir); err == nil {
|
||||||
|
if err := grantAccessTree(comfyDir, "(OI)(CI)(M)"); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
// uv venvs (Comfy-Desktop) redirect to a base interpreter that
|
||||||
|
// can live outside COMFY_DIR; read+execute suffices for it.
|
||||||
|
if home := comfyVenvHome(comfyDir); home != "" {
|
||||||
|
if err := grantAccessTree(home, "(OI)(CI)(RX)"); 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 {
|
||||||
|
return runIcacls(path, perms, false)
|
||||||
|
}
|
||||||
|
|
||||||
|
// grantAccessTree is grantAccess with /T: the ACE is applied to the
|
||||||
|
// existing tree, not just inherited by children created later. Needed when
|
||||||
|
// the tree already exists, e.g. a ComfyUI install in a user profile. On a
|
||||||
|
// large tree (a venv has tens of thousands of files) this takes minutes,
|
||||||
|
// so it says what it is doing instead of looking hung.
|
||||||
|
func grantAccessTree(path, perms string) error {
|
||||||
|
return runIcacls(path, perms, true)
|
||||||
|
}
|
||||||
|
|
||||||
|
func runIcacls(path, perms string, recursive bool) error {
|
||||||
|
args := []string{path, "/grant", virtualAccount + ":" + perms}
|
||||||
|
if recursive {
|
||||||
|
fmt.Printf("granting %s %s access to %s (large trees can take minutes)\n", virtualAccount, perms, path)
|
||||||
|
args = append(args, "/T")
|
||||||
|
}
|
||||||
|
start := time.Now()
|
||||||
|
out, err := exec.Command("icacls", args...).CombinedOutput()
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("grant %s access to %s: %w (%s)", virtualAccount, path, err, strings.TrimSpace(string(out)))
|
||||||
|
}
|
||||||
|
if recursive {
|
||||||
|
fmt.Printf("access granted in %s\n", time.Since(start).Round(time.Second))
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// RestartIfRunning restarts the service when it is installed and running
|
||||||
|
// (used after a forced update staged a new binary). Reports whether a
|
||||||
|
// restart happened. A service that is not installed or not running is not
|
||||||
|
// an error. Needs elevation.
|
||||||
|
func RestartIfRunning() (bool, error) {
|
||||||
|
m, err := mgr.Connect()
|
||||||
|
if err != nil {
|
||||||
|
return false, fmt.Errorf("connect to service manager (run as administrator): %w", err)
|
||||||
|
}
|
||||||
|
defer m.Disconnect()
|
||||||
|
s, err := m.OpenService(Name)
|
||||||
|
if err != nil {
|
||||||
|
return false, nil // not installed
|
||||||
|
}
|
||||||
|
defer s.Close()
|
||||||
|
st, err := s.Query()
|
||||||
|
if err != nil {
|
||||||
|
return false, fmt.Errorf("query service: %w", err)
|
||||||
|
}
|
||||||
|
if st.State != svc.Running && st.State != svc.StartPending {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
if err := stopAndWait(s); err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
if err := s.Start(); err != nil {
|
||||||
|
return false, fmt.Errorf("start service: %w", err)
|
||||||
|
}
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// comfyVenvHome returns the base interpreter directory of the Python venv at
|
||||||
|
// <comfyDir>/.venv when that directory lives outside comfyDir, "" otherwise.
|
||||||
|
// uv-created venvs (Comfy-Desktop) ship a redirector python.exe whose real
|
||||||
|
// interpreter is the pyvenv.cfg "home" tree — typically a sibling of
|
||||||
|
// COMFY_DIR, which a sandbox/ACL covering COMFY_DIR alone does not reach.
|
||||||
|
func comfyVenvHome(comfyDir string) string {
|
||||||
|
data, err := os.ReadFile(filepath.Join(comfyDir, ".venv", "pyvenv.cfg"))
|
||||||
|
if err != nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
for _, line := range strings.Split(string(data), "\n") {
|
||||||
|
k, v, ok := strings.Cut(line, "=")
|
||||||
|
if !ok || strings.TrimSpace(k) != "home" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
home := strings.TrimSpace(v)
|
||||||
|
if home == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
if st, err := os.Stat(home); err != nil || !st.IsDir() {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
rel, err := filepath.Rel(comfyDir, home)
|
||||||
|
if err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) {
|
||||||
|
return home
|
||||||
|
}
|
||||||
|
return "" // inside comfyDir: already covered by the COMFY_DIR grant
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestComfyVenvHome(t *testing.T) {
|
||||||
|
root := t.TempDir()
|
||||||
|
comfy := filepath.Join(root, "ComfyUI")
|
||||||
|
outside := filepath.Join(root, "standalone-env")
|
||||||
|
inside := filepath.Join(comfy, "runtime")
|
||||||
|
for _, d := range []string{filepath.Join(comfy, ".venv"), outside, inside} {
|
||||||
|
if err := os.MkdirAll(d, 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
cfg := filepath.Join(comfy, ".venv", "pyvenv.cfg")
|
||||||
|
|
||||||
|
write := func(home string) {
|
||||||
|
if err := os.WriteFile(cfg, []byte("home = "+home+"\nversion_info = 3.13.0\n"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
write(outside)
|
||||||
|
if got := comfyVenvHome(comfy); got != outside {
|
||||||
|
t.Fatalf("outside home: got %q, want %q", got, outside)
|
||||||
|
}
|
||||||
|
|
||||||
|
write(inside)
|
||||||
|
if got := comfyVenvHome(comfy); got != "" {
|
||||||
|
t.Fatalf("inside home: got %q, want empty", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
write(filepath.Join(root, "does-not-exist"))
|
||||||
|
if got := comfyVenvHome(comfy); got != "" {
|
||||||
|
t.Fatalf("missing home: got %q, want empty", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
if got := comfyVenvHome(filepath.Join(root, "no-venv")); got != "" {
|
||||||
|
t.Fatalf("no pyvenv.cfg: got %q, want empty", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,314 @@
|
|||||||
|
// Package supervise runs an upstream server (ComfyUI) as a managed child
|
||||||
|
// process: started on demand when a request needs it, stopped after an
|
||||||
|
// idle timeout so the GPU memory it holds is freed, and stopped with the
|
||||||
|
// parent. Crashes are logged; the next request respawns it.
|
||||||
|
package supervise
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"os/exec"
|
||||||
|
"path/filepath"
|
||||||
|
"runtime"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ComfyLayout returns the interpreter and script path of a standard ComfyUI
|
||||||
|
// venv install rooted at dir for the given GOOS: .venv\Scripts\python.exe
|
||||||
|
// on Windows, .venv/bin/python elsewhere. The script is main.py — either in
|
||||||
|
// a ComfyUI subdirectory or directly under dir, whichever exists (the
|
||||||
|
// subdirectory form wins ties and is the default when neither exists yet,
|
||||||
|
// so the caller's missing-file warning points at the documented layout).
|
||||||
|
func ComfyLayout(goos, dir string) (python, script string) {
|
||||||
|
if goos == "windows" {
|
||||||
|
python = filepath.Join(dir, ".venv", "Scripts", "python.exe")
|
||||||
|
} else {
|
||||||
|
python = filepath.Join(dir, ".venv", "bin", "python")
|
||||||
|
}
|
||||||
|
script = filepath.Join(dir, "ComfyUI", "main.py")
|
||||||
|
if _, err := os.Stat(script); err != nil {
|
||||||
|
if _, err := os.Stat(filepath.Join(dir, "main.py")); err == nil {
|
||||||
|
script = filepath.Join(dir, "main.py")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return python, script
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultComfyCommand builds the launch command for the standard venv
|
||||||
|
// layout (see ComfyLayout): the script is passed relative to dir so dir
|
||||||
|
// stays the working directory, and --port is taken from comfyURL when the
|
||||||
|
// URL carries one.
|
||||||
|
func DefaultComfyCommand(goos, dir, comfyURL string) string {
|
||||||
|
python, script := ComfyLayout(goos, dir)
|
||||||
|
rel, err := filepath.Rel(dir, script)
|
||||||
|
if err != nil {
|
||||||
|
rel = script
|
||||||
|
}
|
||||||
|
cmd := `"` + python + `" ` + rel
|
||||||
|
if u, err := url.Parse(comfyURL); err == nil && u.Port() != "" {
|
||||||
|
cmd += " --port " + u.Port()
|
||||||
|
}
|
||||||
|
return cmd
|
||||||
|
}
|
||||||
|
|
||||||
|
// Process is one managed child process.
|
||||||
|
type Process struct {
|
||||||
|
name string
|
||||||
|
argv []string
|
||||||
|
dir string
|
||||||
|
probe func(context.Context) error
|
||||||
|
startTimeout time.Duration
|
||||||
|
log *slog.Logger
|
||||||
|
|
||||||
|
mu sync.Mutex
|
||||||
|
cmd *exec.Cmd
|
||||||
|
stopping bool
|
||||||
|
ready bool
|
||||||
|
external bool // someone else serves the port; not our process
|
||||||
|
lastActivity time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
// New parses cmdLine (double quotes group arguments containing spaces) and
|
||||||
|
// prepares a managed process. probe reports whether the server answers
|
||||||
|
// (e.g. the comfy client's Probe); startTimeout bounds WaitReady. dir is
|
||||||
|
// the child's working directory; empty inherits ours.
|
||||||
|
func New(name, cmdLine, dir string, probe func(context.Context) error, startTimeout time.Duration, log *slog.Logger) (*Process, error) {
|
||||||
|
argv, err := splitCommandLine(cmdLine)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("%s command: %w", name, err)
|
||||||
|
}
|
||||||
|
if len(argv) == 0 {
|
||||||
|
return nil, fmt.Errorf("%s command is empty", name)
|
||||||
|
}
|
||||||
|
if log == nil {
|
||||||
|
log = slog.Default()
|
||||||
|
}
|
||||||
|
if startTimeout <= 0 {
|
||||||
|
startTimeout = 2 * time.Minute
|
||||||
|
}
|
||||||
|
return &Process{name: name, argv: argv, dir: dir, probe: probe, startTimeout: startTimeout, log: log}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// splitCommandLine splits a command line on whitespace, treating
|
||||||
|
// double-quoted sections as one argument (quotes removed). Backslashes are
|
||||||
|
// literal — this matches Windows paths.
|
||||||
|
func splitCommandLine(s string) ([]string, error) {
|
||||||
|
var argv []string
|
||||||
|
var cur strings.Builder
|
||||||
|
inQuote := false
|
||||||
|
have := false
|
||||||
|
flush := func() {
|
||||||
|
if have {
|
||||||
|
argv = append(argv, cur.String())
|
||||||
|
cur.Reset()
|
||||||
|
have = false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, r := range s {
|
||||||
|
switch {
|
||||||
|
case r == '"':
|
||||||
|
inQuote = !inQuote
|
||||||
|
have = true
|
||||||
|
case (r == ' ' || r == '\t') && !inQuote:
|
||||||
|
flush()
|
||||||
|
default:
|
||||||
|
cur.WriteRune(r)
|
||||||
|
have = true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if inQuote {
|
||||||
|
return nil, fmt.Errorf("unterminated quote in %q", s)
|
||||||
|
}
|
||||||
|
flush()
|
||||||
|
return argv, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Running reports whether the child process is currently alive.
|
||||||
|
func (p *Process) Running() bool {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
return p.cmd != nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Ready reports whether the server has answered a probe since its last
|
||||||
|
// (re)start. Health checks use it to tell "starting up" from "outage".
|
||||||
|
func (p *Process) Ready() bool {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
return p.ready
|
||||||
|
}
|
||||||
|
|
||||||
|
// MarkReady records that the server answered.
|
||||||
|
func (p *Process) MarkReady() {
|
||||||
|
p.mu.Lock()
|
||||||
|
p.ready = true
|
||||||
|
p.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// NoteActivity resets the idle clock; called for every request served.
|
||||||
|
func (p *Process) NoteActivity() {
|
||||||
|
p.mu.Lock()
|
||||||
|
p.lastActivity = time.Now()
|
||||||
|
p.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
// EnsureRunning starts the child if it is not running. It returns as soon
|
||||||
|
// as the process is spawned; readiness is WaitReady's job (and the proxy's
|
||||||
|
// retry backoff bridges the gap for plain proxied requests). When the URL
|
||||||
|
// already answers — e.g. the ComfyUI desktop app grabbed the port — no
|
||||||
|
// child is spawned: the external server is used as-is, and the idle
|
||||||
|
// watcher never touches it (it only kills its own child).
|
||||||
|
func (p *Process) EnsureRunning() error {
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
p.lastActivity = time.Now()
|
||||||
|
if p.cmd != nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
pctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||||
|
err := p.probe(pctx)
|
||||||
|
cancel()
|
||||||
|
if err == nil {
|
||||||
|
p.ready = true
|
||||||
|
if !p.external {
|
||||||
|
p.external = true
|
||||||
|
p.log.Info(p.name + " is already served externally; not spawning a managed instance")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
p.external = false
|
||||||
|
cmd := exec.Command(p.argv[0], p.argv[1:]...)
|
||||||
|
cmd.Dir = p.dir
|
||||||
|
stdout, err := cmd.StdoutPipe()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
stderr, err := cmd.StderrPipe()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := cmd.Start(); err != nil {
|
||||||
|
return fmt.Errorf("start %s: %w", p.name, err)
|
||||||
|
}
|
||||||
|
p.cmd = cmd
|
||||||
|
p.stopping = false
|
||||||
|
p.ready = false
|
||||||
|
go p.pipeLog(stdout)
|
||||||
|
go p.pipeLog(stderr)
|
||||||
|
go func() {
|
||||||
|
err := cmd.Wait()
|
||||||
|
p.mu.Lock()
|
||||||
|
p.cmd = nil
|
||||||
|
p.ready = false
|
||||||
|
intentional := p.stopping
|
||||||
|
p.mu.Unlock()
|
||||||
|
if intentional {
|
||||||
|
p.log.Info(p.name + " stopped")
|
||||||
|
} else {
|
||||||
|
p.log.Warn(p.name+" exited unexpectedly; the next request restarts it", "err", err)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
p.log.Info(p.name+" starting", "pid", cmd.Process.Pid, "cmd", strings.Join(p.argv, " "))
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// WaitReady blocks until the probe succeeds, ctx ends, or the start
|
||||||
|
// timeout passes.
|
||||||
|
func (p *Process) WaitReady(ctx context.Context) error {
|
||||||
|
ctx, cancel := context.WithTimeout(ctx, p.startTimeout)
|
||||||
|
defer cancel()
|
||||||
|
for {
|
||||||
|
pctx, pcancel := context.WithTimeout(ctx, 5*time.Second)
|
||||||
|
err := p.probe(pctx)
|
||||||
|
pcancel()
|
||||||
|
if err == nil {
|
||||||
|
p.NoteActivity()
|
||||||
|
p.MarkReady()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return fmt.Errorf("%s did not become ready: %w", p.name, ctx.Err())
|
||||||
|
case <-time.After(500 * time.Millisecond):
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Stop kills the child process (the whole tree on Windows) if running.
|
||||||
|
func (p *Process) Stop() {
|
||||||
|
p.mu.Lock()
|
||||||
|
cmd := p.cmd
|
||||||
|
if cmd == nil {
|
||||||
|
p.mu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
|
p.stopping = true
|
||||||
|
p.ready = false
|
||||||
|
p.mu.Unlock()
|
||||||
|
stopTree(cmd)
|
||||||
|
}
|
||||||
|
|
||||||
|
// WatchIdle stops the child after idleTimeout without activity, but only
|
||||||
|
// when gpuIdle reports the GPU lock is free (no active or pending work).
|
||||||
|
// Returns when ctx ends.
|
||||||
|
func (p *Process) WatchIdle(ctx context.Context, idleTimeout time.Duration, gpuIdle func() bool) {
|
||||||
|
ticker := time.NewTicker(5 * time.Second)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
}
|
||||||
|
p.mu.Lock()
|
||||||
|
idleFor := time.Since(p.lastActivity)
|
||||||
|
running := p.cmd != nil
|
||||||
|
p.mu.Unlock()
|
||||||
|
if running && idleFor > idleTimeout && gpuIdle() {
|
||||||
|
p.log.Info(p.name+" idle; stopping to free the GPU", "idle_for", idleFor.Round(time.Second))
|
||||||
|
p.Stop()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// pipeLog forwards one child output stream to the log at INFO, line by
|
||||||
|
// line, prefixed with the process name.
|
||||||
|
func (p *Process) pipeLog(r io.Reader) {
|
||||||
|
buf := make([]byte, 4096)
|
||||||
|
var line string
|
||||||
|
for {
|
||||||
|
n, err := r.Read(buf)
|
||||||
|
line += string(buf[:n])
|
||||||
|
for {
|
||||||
|
i := strings.IndexByte(line, '\n')
|
||||||
|
if i < 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
p.log.Info(p.name + ": " + strings.TrimRight(line[:i], "\r"))
|
||||||
|
line = line[i+1:]
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
if strings.TrimSpace(line) != "" {
|
||||||
|
p.log.Info(p.name + ": " + line)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// stopTree kills cmd's process, including its children on Windows (python
|
||||||
|
// launchers tend to spawn some). The Wait goroutine reaps it.
|
||||||
|
func stopTree(cmd *exec.Cmd) {
|
||||||
|
if runtime.GOOS == "windows" {
|
||||||
|
exec.Command("taskkill", "/T", "/F", "/PID",
|
||||||
|
fmt.Sprint(cmd.Process.Pid)).Run() //nolint:errcheck // best effort
|
||||||
|
return
|
||||||
|
}
|
||||||
|
cmd.Process.Kill() //nolint:errcheck // best effort
|
||||||
|
}
|
||||||
@@ -0,0 +1,241 @@
|
|||||||
|
package supervise
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"log/slog"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestSplitCommandLine(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
in string
|
||||||
|
want []string
|
||||||
|
}{
|
||||||
|
{`python main.py --port 8188`, []string{"python", "main.py", "--port", "8188"}},
|
||||||
|
{`"C:\Program Files\py\python.exe" main.py`, []string{`C:\Program Files\py\python.exe`, "main.py"}},
|
||||||
|
{` spaced out `, []string{"spaced", "out"}},
|
||||||
|
{`a "b c" d`, []string{"a", "b c", "d"}},
|
||||||
|
{"", nil},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
got, err := splitCommandLine(c.in)
|
||||||
|
if err != nil {
|
||||||
|
t.Errorf("splitCommandLine(%q): %v", c.in, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if len(got) != len(c.want) {
|
||||||
|
t.Errorf("splitCommandLine(%q) = %v, want %v", c.in, got, c.want)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
for i := range got {
|
||||||
|
if got[i] != c.want[i] {
|
||||||
|
t.Errorf("splitCommandLine(%q) = %v, want %v", c.in, got, c.want)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, err := splitCommandLine(`"unterminated`); err == nil {
|
||||||
|
t.Error("expected error for unterminated quote")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestHelperProcess is the child executed by the lifecycle tests: it just
|
||||||
|
// sleeps. The env marker is set only around the spawn, so in the normal
|
||||||
|
// test run this returns immediately.
|
||||||
|
func TestHelperProcess(t *testing.T) {
|
||||||
|
if os.Getenv("GO_HELPER_PROCESS") != "1" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
time.Sleep(30 * time.Second)
|
||||||
|
os.Exit(0)
|
||||||
|
}
|
||||||
|
|
||||||
|
func newHelper(t *testing.T, name string) *Process {
|
||||||
|
t.Helper()
|
||||||
|
// The probe always fails: nothing external serves the port, so
|
||||||
|
// EnsureRunning spawns the helper child.
|
||||||
|
p, err := New(name, `"`+os.Args[0]+`" -test.run=TestHelperProcess`, "",
|
||||||
|
func(context.Context) error { return errors.New("nothing there") }, 5*time.Second, slog.Default())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
// startHelper spawns the child with the marker set; exec.Command inherits
|
||||||
|
// the environment at spawn time, so it can be unset right after.
|
||||||
|
func startHelper(t *testing.T, p *Process) {
|
||||||
|
t.Helper()
|
||||||
|
os.Setenv("GO_HELPER_PROCESS", "1")
|
||||||
|
defer os.Unsetenv("GO_HELPER_PROCESS")
|
||||||
|
if err := p.EnsureRunning(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func waitStopped(t *testing.T, p *Process, timeout time.Duration) {
|
||||||
|
t.Helper()
|
||||||
|
deadline := time.Now().Add(timeout)
|
||||||
|
for p.Running() && time.Now().Before(deadline) {
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
}
|
||||||
|
if p.Running() {
|
||||||
|
t.Fatal("process still running")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEnsureRunningAndStop(t *testing.T) {
|
||||||
|
p := newHelper(t, "helper")
|
||||||
|
if p.Running() {
|
||||||
|
t.Fatal("Running before start")
|
||||||
|
}
|
||||||
|
startHelper(t, p)
|
||||||
|
if !p.Running() {
|
||||||
|
t.Fatal("not Running after EnsureRunning")
|
||||||
|
}
|
||||||
|
if err := p.EnsureRunning(); err != nil {
|
||||||
|
t.Fatal("second EnsureRunning must be a no-op")
|
||||||
|
}
|
||||||
|
p.Stop()
|
||||||
|
waitStopped(t, p, 5*time.Second)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestEnsureRunningPrefersExternalServer(t *testing.T) {
|
||||||
|
// The port is already served (e.g. the ComfyUI desktop app): no child
|
||||||
|
// is spawned, the supervisor reports ready, and Stop is a no-op.
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
}))
|
||||||
|
defer srv.Close()
|
||||||
|
probe := func(ctx context.Context) error {
|
||||||
|
req, _ := http.NewRequestWithContext(ctx, http.MethodGet, srv.URL, nil)
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
resp.Body.Close()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
p, err := New("external", `"`+os.Args[0]+`"`, "", probe, 5*time.Second, slog.Default())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := p.EnsureRunning(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if p.Running() {
|
||||||
|
t.Fatal("spawned a child even though the port is already served")
|
||||||
|
}
|
||||||
|
if !p.Ready() {
|
||||||
|
t.Fatal("external server should count as ready")
|
||||||
|
}
|
||||||
|
p.Stop() // must not touch the external server
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWaitReady(t *testing.T) {
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
}))
|
||||||
|
defer srv.Close()
|
||||||
|
probe := func(ctx context.Context) error {
|
||||||
|
req, _ := http.NewRequestWithContext(ctx, http.MethodGet, srv.URL, nil)
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
resp.Body.Close()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
p, err := New("ready", `"`+os.Args[0]+`"`, "", probe, 5*time.Second, slog.Default())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := p.WaitReady(context.Background()); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
failing, err := New("failing", `"`+os.Args[0]+`"`, "",
|
||||||
|
func(context.Context) error { return errors.New("no") }, 500*time.Millisecond, slog.Default())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := failing.WaitReady(context.Background()); err == nil {
|
||||||
|
t.Fatal("expected timeout error from WaitReady")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWatchIdleStops(t *testing.T) {
|
||||||
|
p := newHelper(t, "idle")
|
||||||
|
startHelper(t, p)
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
go p.WatchIdle(ctx, 200*time.Millisecond, func() bool { return true })
|
||||||
|
waitStopped(t, p, 10*time.Second)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWatchIdleRespectsBusyGPU(t *testing.T) {
|
||||||
|
p := newHelper(t, "busy")
|
||||||
|
startHelper(t, p)
|
||||||
|
defer p.Stop()
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
go p.WatchIdle(ctx, 100*time.Millisecond, func() bool { return false })
|
||||||
|
time.Sleep(600 * time.Millisecond)
|
||||||
|
cancel()
|
||||||
|
if !p.Running() {
|
||||||
|
t.Fatal("process was stopped while the GPU was busy")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestComfyLayoutAndDefaultCommand(t *testing.T) {
|
||||||
|
// Nested layout (ComfyUI/main.py under dir) wins.
|
||||||
|
dir := t.TempDir()
|
||||||
|
nested := filepath.Join(dir, "ComfyUI", "main.py")
|
||||||
|
if err := os.MkdirAll(filepath.Dir(nested), 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := os.WriteFile(nested, []byte("x"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
python, script := ComfyLayout("windows", dir)
|
||||||
|
if want := filepath.Join(dir, ".venv", "Scripts", "python.exe"); python != want {
|
||||||
|
t.Errorf("python = %s, want %s", python, want)
|
||||||
|
}
|
||||||
|
if script != nested {
|
||||||
|
t.Errorf("script = %s, want %s", script, nested)
|
||||||
|
}
|
||||||
|
cmd := DefaultComfyCommand("windows", dir, "http://127.0.0.1:8189")
|
||||||
|
want := `"` + filepath.Join(dir, ".venv", "Scripts", "python.exe") + `" ` + filepath.Join("ComfyUI", "main.py") + " --port 8189"
|
||||||
|
if cmd != want {
|
||||||
|
t.Errorf("cmd = %q, want %q", cmd, want)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Flat layout (main.py directly under dir) is found too.
|
||||||
|
flat := t.TempDir()
|
||||||
|
if err := os.WriteFile(filepath.Join(flat, "main.py"), []byte("x"), 0o644); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, script := ComfyLayout("linux", flat); script != filepath.Join(flat, "main.py") {
|
||||||
|
t.Errorf("flat script = %s", script)
|
||||||
|
}
|
||||||
|
cmd = DefaultComfyCommand("linux", flat, "http://comfy.internal")
|
||||||
|
if strings.Contains(cmd, "--port") {
|
||||||
|
t.Errorf("cmd = %q, want no --port for a port-less URL", cmd)
|
||||||
|
}
|
||||||
|
if !strings.HasSuffix(cmd, `" main.py`) {
|
||||||
|
t.Errorf("cmd = %q, want quoted python + relative main.py", cmd)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Neither exists yet: default to the documented nested form so the
|
||||||
|
// startup warning points there.
|
||||||
|
empty := t.TempDir()
|
||||||
|
if _, script := ComfyLayout("windows", empty); script != filepath.Join(empty, "ComfyUI", "main.py") {
|
||||||
|
t.Errorf("missing-layout script = %s", script)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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,282 @@
|
|||||||
|
// Package update implements gpu-turnstile's self-updater: it polls the
|
||||||
|
// Gitea releases API, downloads the Windows binary of newer releases, and
|
||||||
|
// verifies its Ed25519 signature (produced by CI with OpenSSL) before
|
||||||
|
// swapping it in next to the running executable.
|
||||||
|
package update
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/ed25519"
|
||||||
|
"crypto/sha256"
|
||||||
|
"crypto/x509"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"encoding/pem"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// maxAssetSize bounds release asset downloads.
|
||||||
|
const maxAssetSize = 512 << 20
|
||||||
|
|
||||||
|
// Updater checks one Gitea repository for releases.
|
||||||
|
type Updater struct {
|
||||||
|
Repo string // e.g. https://git.rambossek.at/PUBLIC/gpu-turnstile
|
||||||
|
Asset string // e.g. gpu-turnstile.exe
|
||||||
|
Version string // current binary version, e.g. v0.1.2 ("dev" cannot be compared)
|
||||||
|
// Desired is the APP_VER policy: "dev" disables updates, "stable" (or
|
||||||
|
// empty) tracks the latest release, anything else is an exact vX.Y.Z
|
||||||
|
// release tag to pin — staging it even when that means a downgrade.
|
||||||
|
Desired string
|
||||||
|
Log *slog.Logger
|
||||||
|
Client *http.Client
|
||||||
|
}
|
||||||
|
|
||||||
|
type release struct {
|
||||||
|
TagName string `json:"tag_name"`
|
||||||
|
Assets []struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
BrowserDownloadURL string `json:"browser_download_url"`
|
||||||
|
} `json:"assets"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (u *Updater) logger() *slog.Logger {
|
||||||
|
if u.Log != nil {
|
||||||
|
return u.Log
|
||||||
|
}
|
||||||
|
return slog.Default()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (u *Updater) httpClient() *http.Client {
|
||||||
|
if u.Client != nil {
|
||||||
|
return u.Client
|
||||||
|
}
|
||||||
|
return &http.Client{Timeout: 5 * time.Minute}
|
||||||
|
}
|
||||||
|
|
||||||
|
// apiURL derives <scheme>://<host>/api/v1/repos/<owner>/<name> from Repo.
|
||||||
|
func (u *Updater) apiURL() (string, error) {
|
||||||
|
repoURL, err := url.Parse(u.Repo)
|
||||||
|
if err != nil || repoURL.Scheme == "" || repoURL.Host == "" {
|
||||||
|
return "", fmt.Errorf("invalid UPDATE_REPO %q", u.Repo)
|
||||||
|
}
|
||||||
|
ownerName := strings.Trim(repoURL.Path, "/")
|
||||||
|
if len(strings.Split(ownerName, "/")) != 2 {
|
||||||
|
return "", fmt.Errorf("UPDATE_REPO %q: expected path /<owner>/<name>", u.Repo)
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("%s://%s/api/v1/repos/%s", repoURL.Scheme, repoURL.Host, ownerName), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// newerVersion reports whether latest is a higher vX.Y.Z version than
|
||||||
|
// current. Both may carry a leading "v".
|
||||||
|
func newerVersion(current, latest string) (bool, error) {
|
||||||
|
parse := func(s string) ([3]int, error) {
|
||||||
|
var v [3]int
|
||||||
|
parts := strings.Split(strings.TrimPrefix(s, "v"), ".")
|
||||||
|
if len(parts) != 3 {
|
||||||
|
return v, fmt.Errorf("not a vX.Y.Z version: %q", s)
|
||||||
|
}
|
||||||
|
for i, p := range parts {
|
||||||
|
n, err := strconv.Atoi(p)
|
||||||
|
if err != nil {
|
||||||
|
return v, fmt.Errorf("not a vX.Y.Z version: %q", s)
|
||||||
|
}
|
||||||
|
v[i] = n
|
||||||
|
}
|
||||||
|
return v, nil
|
||||||
|
}
|
||||||
|
cur, err := parse(current)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
lat, err := parse(latest)
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
for i := 0; i < 3; i++ {
|
||||||
|
if lat[i] != cur[i] {
|
||||||
|
return lat[i] > cur[i], nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (u *Updater) get(ctx context.Context, url string) ([]byte, error) {
|
||||||
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
resp, err := u.httpClient().Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
io.Copy(io.Discard, resp.Body)
|
||||||
|
return nil, fmt.Errorf("GET %s: %s", url, resp.Status)
|
||||||
|
}
|
||||||
|
return io.ReadAll(io.LimitReader(resp.Body, maxAssetSize))
|
||||||
|
}
|
||||||
|
|
||||||
|
func verifySignature(pubKeyPEM string, data, sig []byte) error {
|
||||||
|
block, _ := pem.Decode([]byte(pubKeyPEM))
|
||||||
|
if block == nil {
|
||||||
|
return fmt.Errorf("invalid embedded public key PEM")
|
||||||
|
}
|
||||||
|
key, err := x509.ParsePKIXPublicKey(block.Bytes)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("parse public key: %w", err)
|
||||||
|
}
|
||||||
|
pub, ok := key.(ed25519.PublicKey)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Errorf("public key is not Ed25519")
|
||||||
|
}
|
||||||
|
if !ed25519.Verify(pub, data, sig) {
|
||||||
|
return fmt.Errorf("signature verification failed")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// stage swaps data into place at exePath: the running executable is
|
||||||
|
// renamed aside (allowed on Windows) and the new file takes its name.
|
||||||
|
func stage(exePath string, data []byte) error {
|
||||||
|
newPath := exePath + ".new"
|
||||||
|
oldPath := exePath + ".old"
|
||||||
|
os.Remove(oldPath) // leftover from a previous update
|
||||||
|
if err := os.WriteFile(newPath, data, 0o755); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := os.Rename(exePath, oldPath); err != nil {
|
||||||
|
os.Remove(newPath)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if err := os.Rename(newPath, exePath); err != nil {
|
||||||
|
os.Rename(oldPath, exePath) // roll back
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// CleanupOld removes the .old binary left behind by a staged update.
|
||||||
|
// Call once at startup.
|
||||||
|
func CleanupOld(exePath string) {
|
||||||
|
os.Remove(exePath + ".old")
|
||||||
|
os.Remove(exePath + ".new")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check performs a single update check. staged is true when a
|
||||||
|
// signature-verified binary has been swapped into place at exePath; the
|
||||||
|
// caller should then restart the process. to is the release tag the check
|
||||||
|
// resolved (the latest release or the pinned tag), set once the release
|
||||||
|
// fetch succeeded — even when staging afterwards fails. A nil error with
|
||||||
|
// staged=false means "no action" (up to date, APP_VER=dev, or no embedded
|
||||||
|
// public key); a non-nil error means the check failed and the running
|
||||||
|
// binary is untouched.
|
||||||
|
func (u *Updater) Check(ctx context.Context, exePath string) (staged bool, to string, err error) {
|
||||||
|
log := u.logger()
|
||||||
|
desired := u.Desired
|
||||||
|
if desired == "" {
|
||||||
|
desired = "stable"
|
||||||
|
}
|
||||||
|
if desired == "dev" {
|
||||||
|
log.Debug("auto-update: APP_VER=dev, skipping")
|
||||||
|
return false, "", nil
|
||||||
|
}
|
||||||
|
if publicKeyPEM == "" {
|
||||||
|
log.Debug("auto-update: no public key embedded, skipping")
|
||||||
|
return false, "", nil
|
||||||
|
}
|
||||||
|
api, err := u.apiURL()
|
||||||
|
if err != nil {
|
||||||
|
return false, "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
pinned := desired != "stable"
|
||||||
|
endpoint := api + "/releases/latest"
|
||||||
|
if pinned {
|
||||||
|
endpoint = api + "/releases/tags/" + desired
|
||||||
|
}
|
||||||
|
body, err := u.get(ctx, endpoint)
|
||||||
|
if err != nil {
|
||||||
|
return false, "", fmt.Errorf("fetch release: %w", err)
|
||||||
|
}
|
||||||
|
var rel release
|
||||||
|
if err := json.Unmarshal(body, &rel); err != nil {
|
||||||
|
return false, "", fmt.Errorf("parse release: %w", err)
|
||||||
|
}
|
||||||
|
to = rel.TagName
|
||||||
|
if pinned {
|
||||||
|
// Pin mode: any difference from the target tag means stage it —
|
||||||
|
// including downgrades and replacing a dev binary.
|
||||||
|
if u.Version == rel.TagName {
|
||||||
|
log.Debug("auto-update: already on pinned version", "version", u.Version)
|
||||||
|
return false, to, nil
|
||||||
|
}
|
||||||
|
} else if u.Version != "" && u.Version != "dev" {
|
||||||
|
// Stable mode: only strictly newer releases count; a dev binary
|
||||||
|
// cannot be compared and is always replaced by the latest release.
|
||||||
|
newer, err := newerVersion(u.Version, rel.TagName)
|
||||||
|
if err != nil {
|
||||||
|
return false, to, err
|
||||||
|
}
|
||||||
|
if !newer {
|
||||||
|
log.Debug("auto-update: up to date", "version", u.Version, "latest", rel.TagName)
|
||||||
|
return false, to, 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, to, fmt.Errorf("release %s has no asset %q", rel.TagName, u.Asset)
|
||||||
|
}
|
||||||
|
sigURL, ok := urls[u.Asset+".sig"]
|
||||||
|
if !ok {
|
||||||
|
return false, to, 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, to, fmt.Errorf("download %s: %w", u.Asset, err)
|
||||||
|
}
|
||||||
|
sig, err := u.get(ctx, sigURL)
|
||||||
|
if err != nil {
|
||||||
|
return false, to, 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, to, fmt.Errorf("download checksum: %w", err)
|
||||||
|
}
|
||||||
|
want := strings.Fields(string(sumText))[0]
|
||||||
|
got := hex.EncodeToString(sha256Bytes(data))
|
||||||
|
if !strings.EqualFold(want, got) {
|
||||||
|
return false, to, fmt.Errorf("sha256 mismatch: got %s, want %s", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := verifySignature(publicKeyPEM, data, sig); err != nil {
|
||||||
|
return false, to, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := stage(exePath, data); err != nil {
|
||||||
|
return false, to, fmt.Errorf("stage update: %w", err)
|
||||||
|
}
|
||||||
|
log.Info("auto-update: new version staged", "from", u.Version, "to", rel.TagName)
|
||||||
|
return true, to, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func sha256Bytes(data []byte) []byte {
|
||||||
|
sum := sha256.Sum256(data)
|
||||||
|
return sum[:]
|
||||||
|
}
|
||||||
@@ -0,0 +1,249 @@
|
|||||||
|
package update
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/ed25519"
|
||||||
|
"crypto/rand"
|
||||||
|
"crypto/x509"
|
||||||
|
"encoding/json"
|
||||||
|
"encoding/pem"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
// fakeGitea serves a Gitea-flavored releases API with one release.
|
||||||
|
type fakeGitea struct {
|
||||||
|
srv *httptest.Server
|
||||||
|
pubPEM string
|
||||||
|
asset []byte
|
||||||
|
tag string
|
||||||
|
tamper bool
|
||||||
|
noSig bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func newFakeGitea(t *testing.T, tag string, assetContent []byte) *fakeGitea {
|
||||||
|
t.Helper()
|
||||||
|
pub, priv, err := ed25519.GenerateKey(rand.Reader)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
der, err := x509.MarshalPKIXPublicKey(pub)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
f := &fakeGitea{
|
||||||
|
pubPEM: string(pem.EncodeToMemory(&pem.Block{Type: "PUBLIC KEY", Bytes: der})),
|
||||||
|
asset: assetContent,
|
||||||
|
tag: tag,
|
||||||
|
}
|
||||||
|
sign := func() []byte { return ed25519.Sign(priv, f.asset) }
|
||||||
|
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
serveRelease := func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
assets := []map[string]string{
|
||||||
|
{"name": "gpu-turnstile.exe", "browser_download_url": f.srv.URL + "/dl/exe"},
|
||||||
|
{"name": "gpu-turnstile.exe.sha256", "browser_download_url": f.srv.URL + "/dl/sha"},
|
||||||
|
}
|
||||||
|
if !f.noSig {
|
||||||
|
assets = append(assets, map[string]string{"name": "gpu-turnstile.exe.sig", "browser_download_url": f.srv.URL + "/dl/sig"})
|
||||||
|
}
|
||||||
|
json.NewEncoder(w).Encode(map[string]any{"tag_name": f.tag, "assets": assets})
|
||||||
|
}
|
||||||
|
mux.HandleFunc("/api/v1/repos/o/r/releases/latest", serveRelease)
|
||||||
|
mux.HandleFunc("/api/v1/repos/o/r/releases/tags/"+f.tag, serveRelease)
|
||||||
|
mux.HandleFunc("/dl/exe", func(w http.ResponseWriter, r *http.Request) { w.Write(f.asset) })
|
||||||
|
mux.HandleFunc("/dl/sig", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
sig := sign()
|
||||||
|
if f.tamper {
|
||||||
|
sig[0] ^= 0xff
|
||||||
|
}
|
||||||
|
w.Write(sig)
|
||||||
|
})
|
||||||
|
mux.HandleFunc("/dl/sha", func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
fmt.Fprintf(w, "%x gpu-turnstile.exe\n", sha256Bytes(f.asset))
|
||||||
|
})
|
||||||
|
f.srv = httptest.NewServer(mux)
|
||||||
|
t.Cleanup(f.srv.Close)
|
||||||
|
return f
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeGitea) updater(version string) *Updater {
|
||||||
|
return &Updater{Repo: f.srv.URL + "/o/r", Asset: "gpu-turnstile.exe", Version: version}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (f *fakeGitea) updaterDesired(version, desired string) *Updater {
|
||||||
|
u := f.updater(version)
|
||||||
|
u.Desired = desired
|
||||||
|
return u
|
||||||
|
}
|
||||||
|
|
||||||
|
func fakeExe(t *testing.T) string {
|
||||||
|
t.Helper()
|
||||||
|
exe := filepath.Join(t.TempDir(), "gpu-turnstile.exe")
|
||||||
|
if err := os.WriteFile(exe, []byte("old-binary"), 0o755); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return exe
|
||||||
|
}
|
||||||
|
|
||||||
|
func withPublicKey(t *testing.T, pem string) {
|
||||||
|
t.Helper()
|
||||||
|
old := publicKeyPEM
|
||||||
|
publicKeyPEM = pem
|
||||||
|
t.Cleanup(func() { publicKeyPEM = old })
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCheckStagesUpdate(t *testing.T) {
|
||||||
|
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||||
|
withPublicKey(t, f.pubPEM)
|
||||||
|
exe := fakeExe(t)
|
||||||
|
|
||||||
|
staged, to, err := f.updater("v0.1.2").Check(context.Background(), exe)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !staged {
|
||||||
|
t.Fatal("expected staged update")
|
||||||
|
}
|
||||||
|
if to != "v9.9.9" {
|
||||||
|
t.Fatalf("to = %q, want v9.9.9", to)
|
||||||
|
}
|
||||||
|
content, _ := os.ReadFile(exe)
|
||||||
|
if string(content) != "new-binary" {
|
||||||
|
t.Fatalf("exe content = %q", content)
|
||||||
|
}
|
||||||
|
old, _ := os.ReadFile(exe + ".old")
|
||||||
|
if string(old) != "old-binary" {
|
||||||
|
t.Fatalf(".old content = %q", old)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCheckRejectsTamperedSignature(t *testing.T) {
|
||||||
|
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||||
|
f.tamper = true
|
||||||
|
withPublicKey(t, f.pubPEM)
|
||||||
|
exe := fakeExe(t)
|
||||||
|
|
||||||
|
staged, _, err := f.updater("v0.1.2").Check(context.Background(), exe)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected signature error")
|
||||||
|
}
|
||||||
|
if staged {
|
||||||
|
t.Fatal("must not stage on bad signature")
|
||||||
|
}
|
||||||
|
content, _ := os.ReadFile(exe)
|
||||||
|
if string(content) != "old-binary" {
|
||||||
|
t.Fatal("exe was modified despite bad signature")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCheckSkipsOlderOrEqual(t *testing.T) {
|
||||||
|
for _, tag := range []string{"v0.1.2", "v0.1.1", "v0.0.9"} {
|
||||||
|
f := newFakeGitea(t, tag, []byte("new-binary"))
|
||||||
|
withPublicKey(t, f.pubPEM)
|
||||||
|
staged, _, err := f.updater("v0.1.2").Check(context.Background(), fakeExe(t))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if staged {
|
||||||
|
t.Fatalf("tag %s must not stage over v0.1.2", tag)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCheckSkipsWithoutPublicKey(t *testing.T) {
|
||||||
|
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||||
|
withPublicKey(t, "")
|
||||||
|
staged, _, err := f.updater("v0.1.2").Check(context.Background(), fakeExe(t))
|
||||||
|
if err != nil || staged {
|
||||||
|
t.Fatalf("staged=%v err=%v, want no action without key", staged, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCheckDevBuildGetsStable(t *testing.T) {
|
||||||
|
// A dev binary cannot be compared; APP_VER=stable replaces it with the
|
||||||
|
// latest release.
|
||||||
|
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||||
|
withPublicKey(t, f.pubPEM)
|
||||||
|
exe := fakeExe(t)
|
||||||
|
staged, _, err := f.updater("dev").Check(context.Background(), exe)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !staged {
|
||||||
|
t.Fatal("expected dev binary to be replaced by the latest release")
|
||||||
|
}
|
||||||
|
content, _ := os.ReadFile(exe)
|
||||||
|
if string(content) != "new-binary" {
|
||||||
|
t.Fatalf("exe content = %q", content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCheckDesiredDevDisables(t *testing.T) {
|
||||||
|
f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
|
||||||
|
withPublicKey(t, f.pubPEM)
|
||||||
|
for _, version := range []string{"dev", "v0.1.2"} {
|
||||||
|
staged, _, err := f.updaterDesired(version, "dev").Check(context.Background(), fakeExe(t))
|
||||||
|
if err != nil || staged {
|
||||||
|
t.Fatalf("version %s: staged=%v err=%v, want no action with APP_VER=dev", version, staged, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCheckPinned(t *testing.T) {
|
||||||
|
// Pin mode stages the exact tag — up or down — and skips when the
|
||||||
|
// binary already matches.
|
||||||
|
for _, version := range []string{"v0.1.2", "v9.9.9", "dev"} {
|
||||||
|
f := newFakeGitea(t, "v0.5.0", []byte("pinned-binary"))
|
||||||
|
withPublicKey(t, f.pubPEM)
|
||||||
|
exe := fakeExe(t)
|
||||||
|
staged, to, err := f.updaterDesired(version, "v0.5.0").Check(context.Background(), exe)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("version %s: %v", version, err)
|
||||||
|
}
|
||||||
|
if !staged {
|
||||||
|
t.Fatalf("version %s: expected pinned v0.5.0 to be staged", version)
|
||||||
|
}
|
||||||
|
if to != "v0.5.0" {
|
||||||
|
t.Fatalf("version %s: to = %q, want v0.5.0", version, to)
|
||||||
|
}
|
||||||
|
content, _ := os.ReadFile(exe)
|
||||||
|
if string(content) != "pinned-binary" {
|
||||||
|
t.Fatalf("version %s: exe content = %q", version, content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
f := newFakeGitea(t, "v0.5.0", []byte("pinned-binary"))
|
||||||
|
withPublicKey(t, f.pubPEM)
|
||||||
|
staged, _, err := f.updaterDesired("v0.5.0", "v0.5.0").Check(context.Background(), fakeExe(t))
|
||||||
|
if err != nil || staged {
|
||||||
|
t.Fatalf("staged=%v err=%v, want no action when already on the pinned version", staged, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewerVersion(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
cur, lat string
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{"v0.1.2", "v0.1.3", true},
|
||||||
|
{"0.1.2", "0.2.0", true},
|
||||||
|
{"v0.1.2", "v1.0.0", true},
|
||||||
|
{"v0.1.2", "v0.1.2", false},
|
||||||
|
{"v1.2.3", "v1.2.10", true},
|
||||||
|
{"v1.2.10", "v1.2.3", false},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
got, err := newerVersion(c.cur, c.lat)
|
||||||
|
if err != nil || got != c.want {
|
||||||
|
t.Errorf("newerVersion(%s, %s) = %v, %v; want %v", c.cur, c.lat, got, err, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, err := newerVersion("v0.1", "v0.1.2"); err == nil {
|
||||||
|
t.Error("expected error for malformed version")
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user