Compare commits

..
32 Commits
Author SHA1 Message Date
mram e98331bb9e Pin compose example to v0.2.9
ci / test (push) Successful in 14s
ci / docker (push) Successful in 1m8s
ci / release (push) Successful in 16s
2026-09-22 11:23:17 +02:00
mram 54174787d5 Fix control-pipe client race: retry CreateFile on ERROR_PIPE_BUSY
WaitNamedPipe can report an instance that a concurrent client (the
monitor polls status every second) grabs before our CreateFile runs;
the single-attempt Ask then failed with 'no answer from the service'.
Retry the wait+open until a 5s overall deadline.
2026-09-22 11:22:27 +02:00
mram d259b2e96c Add GPU_FOREIGN_UTIL_PCT: game detection via per-process GPU 3D-engine usage
Reads the same PDH counters as Task Manager (\GPU Engine(*)\Utilization
Percentage, locale-independent via PdhAddEnglishCounterW), which cover
graphics work under WDDM — games are caught without an exe list and the
offender is named. Windows-only; a missing or failing counter disables
the path with one log line. dwm (the desktop compositor) joins the
default ignore list.
2026-09-22 10:53:29 +02:00
mram 8fccd333aa Env names in reload diff; monitor shows last VRAM check result; harden control channel against floods
- diffConfig reports the env var names (from new struct tags) instead of
  Go field names, so the reload-env reply names what the user can change
- the monitor's GPU line now shows the game detector's last finding
  (external holders or none) and how long ago the check ran
- control channel: 10s per-connection watchdog (abortive force-close),
  cap of 32 concurrent connections, reload-env rate-limited; command
  read was already capped at 4 KiB
2026-09-22 09:24:27 +02:00
mram 4bd5f34ce7 Pin compose example to v0.2.8
ci / test (push) Successful in 15s
ci / docker (push) Successful in 1m6s
ci / release (push) Successful in 15s
2026-09-22 09:10:14 +02:00
mram b9d3f91403 Monitor: hotkeys (q quit, u update now) with footer line; show GPU VRAM usage 2026-09-22 09:10:14 +02:00
mram 98d2d714ca Add --reload-env: service re-reads and validates its config, restarts when idle if it changed 2026-09-22 09:02:25 +02:00
mram 7a6409831c Pin compose example to v0.2.7
ci / test (push) Successful in 14s
ci / docker (push) Successful in 1m7s
ci / release (push) Successful in 16s
2026-09-22 08:47:47 +02:00
mram c7b8747d85 Fix control pipe accept loop: create instances only after a client connects; CLI force-update never touches the service log file 2026-09-22 08:43:58 +02:00
mram 09d4e81eab Pin compose example to v0.2.6
ci / test (push) Successful in 15s
ci / docker (push) Successful in 1m6s
ci / release (push) Successful in 15s
2026-09-22 08:36:36 +02:00
mram 1264afe40e Monitor restarts itself when the service updated its binary on disk 2026-09-22 08:34:47 +02:00
mram c95d401a36 Pin compose example to v0.2.5
ci / test (push) Successful in 14s
ci / docker (push) Successful in 1m6s
ci / release (push) Successful in 16s
2026-09-22 08:32:06 +02:00
mram f7ea30a494 Fix control pipe races: treat ERROR_PIPE_CONNECTED as success, flush+disconnect before close so replies are never discarded 2026-09-22 08:31:13 +02:00
mram 620be9d57c Add short flags: -i (install), -r (remove), -m (monitor) 2026-09-22 08:23:40 +02:00
mram cffb4151a6 Pin compose example to v0.2.4
ci / test (push) Successful in 14s
ci / docker (push) Successful in 1m9s
ci / release (push) Successful in 15s
2026-09-22 08:19:13 +02:00
mram 16dc7fe462 Add --monitor: live status view (downstream health, GPU lock, queue) via the control channel 2026-09-22 08:17:05 +02:00
mram a74cd49fe5 Pin compose example to v0.2.3
ci / test (push) Successful in 14s
ci / docker (push) Successful in 1m9s
ci / release (push) Successful in 15s
2026-09-21 23:03:45 +02:00
mram cc2a2cad27 Add --update-now: update trigger that only goes through the service's control channel 2026-09-21 23:03:45 +02:00
mram f707d07fd8 Local control channel: unprivileged users can trigger --force-update via the running service 2026-09-21 23:00:38 +02:00
mram ca33a3db82 Pin compose example to v0.2.2
ci / test (push) Successful in 16s
ci / docker (push) Successful in 1m8s
ci / release (push) Successful in 15s
2026-09-21 22:34:56 +02:00
mram 507af1dbff Managed ComfyUI picks up Comfy-Desktop shared models/input/output automatically 2026-09-21 22:34:10 +02:00
mram 45bc9fe27c Strip ANSI escapes from managed ComfyUI output in the log; spawn child with NO_COLOR/TERM=dumb 2026-09-21 22:29:41 +02:00
mram 5b752d5637 Pin compose example to v0.2.1
ci / test (push) Successful in 14s
ci / docker (push) Successful in 1m16s
ci / release (push) Successful in 15s
2026-09-21 22:19:49 +02:00
mram b55d548c5f install-service: grant venv base interpreter (pyvenv.cfg home) outside COMFY_DIR 2026-09-21 22:19:08 +02:00
mram 2d5be72436 --force-update reports from/to versions; elevated parent no longer claims "update applied" when nothing changed 2026-09-21 21:58:50 +02:00
mram 21e40a1774 Pin compose example to v0.2.0
ci / test (push) Successful in 14s
ci / docker (push) Successful in 1m7s
ci / release (push) Successful in 15s
2026-09-21 21:47:48 +02:00
mram bdc844872d Install narrates itself: stop/copy/start steps and the potentially long icacls tree grant 2026-09-21 21:36:45 +02:00
mram 299b6dc0bb Distinguish EACCES from ENOENT in the COMFY_DIR startup check 2026-09-21 21:32:45 +02:00
mram 9589e58ce6 Install opens up COMFY_DIR for the sandboxed service: ACL grant on Windows, BindPaths on Linux 2026-09-21 21:30:35 +02:00
mram be0bb36317 Pin compose example to v0.1.10
ci / test (push) Successful in 28s
ci / docker (push) Successful in 1m15s
ci / release (push) Successful in 17s
2026-09-21 21:20:49 +02:00
mram 9b267533c9 COMFY_DIR alone manages ComfyUI: derive the launch command from the standard venv layout 2026-09-21 21:18:33 +02:00
mram 6f092ddc12 Default GAME_POLL_INTERVAL to 15s: nvidia-smi polls keep the GPU awake 2026-09-21 21:10:21 +02:00
34 changed files with 2457 additions and 251 deletions
+48 -23
View File
@@ -34,8 +34,8 @@ 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 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 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 — 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` detection of foreign GPU holders such as games — is enabled by `GAME_PROCS`,
and/or `GPU_FOREIGN_VRAM_MB` (see below). `GPU_FOREIGN_VRAM_MB` and/or `GPU_FOREIGN_UTIL_PCT` (see below).
## Configuration ## Configuration
@@ -58,14 +58,15 @@ override file values. Invalid values fail at startup.
| `LLM_BUSY_STATUS` | `503` | HTTP status for rejected LLM requests in reject mode (400599, e.g. 429) | | `LLM_BUSY_STATUS` | `503` | HTTP status for rejected LLM requests in reject mode (400599, e.g. 429) |
| `BUSY_RETRY_AFTER` | `30` | Seconds sent as `Retry-After` on busy responses (both modes) | | `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) |
| `COMFY_CMD` | _(empty = unmanaged)_ | Supervise ComfyUI: start on demand, stop when idle to free VRAM. Requires `COMFY_URL` | | `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)_ | Working directory for `COMFY_CMD` | | `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_IDLE_TIMEOUT` | `5m` | Stop the managed ComfyUI after this long idle |
| `COMFY_START_TIMEOUT` | `2m` | Max wait for the managed ComfyUI to come up | | `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 | | `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_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 | | `GPU_FOREIGN_UTIL_PCT` | `0` (disabled) | Also treat the GPU as held when a non-ignored process uses more than this percent of the GPU 3D engine (Windows PDH counters — what Task Manager shows; catches games without an exe list) |
| `GAME_POLL_INTERVAL` | `5s` | How often game/VRAM detection runs | | `GPU_IGNORE_PROCS` | `ollama,ollama app,ollama_llama_server,python,pythonw,dwm` | 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 | | `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 | | `LOG_FILE` | _(empty)_ | Append logs to this file instead of stderr |
@@ -99,18 +100,28 @@ override file values. Invalid values fail at startup.
class) in text mode, which renders in `docker compose logs` on Windows class) in text mode, which renders in `docker compose logs` on Windows
Terminal. Set `NO_COLOR` to disable colors. Terminal. Set `NO_COLOR` to disable colors.
## Managed ComfyUI (`COMFY_CMD`) ## Managed ComfyUI (`COMFY_CMD` / `COMFY_DIR`)
Don't want ComfyUI running 24/7 (it holds VRAM even when idle — and the Don't want ComfyUI running 24/7 (it holds VRAM even when idle — and the
Desktop app kills its server when you close it)? Point `COMFY_CMD` at a Desktop app kills its server when you close it)? gpu-turnstile can supervise
standalone launch command and gpu-turnstile supervises it: the first it: the first request starts it, it stops again after `COMFY_IDLE_TIMEOUT`
request starts it, it stops again after `COMFY_IDLE_TIMEOUT` (default 5m) (default 5m) without work, freeing the GPU for games or the LLM.
without work, freeing the GPU for games or the LLM. Example for a Desktop
install (run it once manually to confirm it works): 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:8188 COMFY_URL=http://127.0.0.1:8189
COMFY_CMD="C:\ComfyUI\.venv\Scripts\python.exe ComfyUI\main.py --port 8188" 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 COMFY_DIR=C:\ComfyUI
``` ```
@@ -124,25 +135,39 @@ answers on the port, gpu-turnstile just uses it instead of spawning
instance already holds the port when you open the desktop app, the instance already holds the port when you open the desktop app, the
desktop's server is the one that fails to bind. 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 ## Game detection
Want to game on the same GPU without Ollama/ComfyUI squatting on the VRAM? 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, 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 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 models and stop the managed ComfyUI so the game gets the memory. Three
detection paths, each optional, polled every `GAME_POLL_INTERVAL` (5s): detection paths, each optional, polled every `GAME_POLL_INTERVAL` (15s):
``` ```
GAME_PROCS=cyberpunk2077.exe,bg3.exe # the reliable way on Windows GPU_FOREIGN_UTIL_PCT=30 # zero-config on Windows: any non-ignored
# process using >30% of the GPU 3D engine
GAME_PROCS=cyberpunk2077.exe,bg3.exe # explicit exe watch list
GPU_FOREIGN_VRAM_MB=1024 # catch-all via nvidia-smi GPU_FOREIGN_VRAM_MB=1024 # catch-all via nvidia-smi
``` ```
`GAME_PROCS` matches running process names (case-insensitive, `.exe` `GPU_FOREIGN_UTIL_PCT` reads the same per-process GPU engine counters as
optional). `GPU_FOREIGN_VRAM_MB` asks nvidia-smi which processes hold GPU Task Manager (PDH), which cover graphics work under Windows' WDDM driver —
memory and treats anything not in `GPU_IGNORE_PROCS` above the threshold as so games show up without naming them, and the offender is named in the log
foreign — handy as a catch-all, but note that under Windows' WDDM driver and monitor. Windows-only. `GAME_PROCS` matches running process names
graphics-only games may not show up in nvidia-smi's per-process list, so (case-insensitive, `.exe` optional). `GPU_FOREIGN_VRAM_MB` asks nvidia-smi
name your games in `GAME_PROCS` there; on Linux both paths work. When the which processes hold GPU memory and treats anything not in
`GPU_IGNORE_PROCS` above the threshold as foreign — handy as a catch-all on
Linux, but under WDDM graphics-only games may not show up in nvidia-smi's
per-process list, so on Windows prefer `GPU_FOREIGN_UTIL_PCT`. When the
game exits, requests resume automatically. game exits, requests resume automatically.
## Build and run ## Build and run
+37 -17
View File
@@ -55,9 +55,9 @@ consumer gets no listener, no startup probe, and no lock participation:
- **Only `COMFY_URL`**: image jobs are tracked and ComfyUI's VRAM is freed - **Only `COMFY_URL`**: image jobs are tracked and ComfyUI's VRAM is freed
afterwards, but the Ollama unload and warm-reload steps are skipped. afterwards, but the Ollama unload and warm-reload steps are skipped.
- **Game detection** is a third, optional consumer without a URL: enabled by - **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 `GAME_PROCS`, `GPU_FOREIGN_VRAM_MB` and/or `GPU_FOREIGN_UTIL_PCT` it
holding the GPU (see below) and plugs into the same lock the same way — watches for foreign processes holding the GPU (see below) and plugs into
excluded when both knobs are unset. the same lock the same way — excluded when all knobs are unset.
### Lock semantics ### Lock semantics
@@ -128,10 +128,15 @@ 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`) ### Managed ComfyUI (`COMFY_CMD` / `COMFY_DIR`)
When `COMFY_CMD` is set, gpu-turnstile runs ComfyUI as a supervised child When `COMFY_CMD` is set, gpu-turnstile runs ComfyUI as a supervised child
process instead of expecting an always-on server: 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 - **Start on demand**: any ComfyUI request spawns it (double quotes in the
command line group arguments with spaces; `COMFY_DIR` sets the working command line group arguments with spaces; `COMFY_DIR` sets the working
@@ -154,24 +159,38 @@ process instead of expecting an always-on server:
- Its stdout/stderr is forwarded to the log at INFO. The health check - Its stdout/stderr is forwarded to the log at INFO. The health check
skips the intentionally-stopped/starting states; a failed probe while skips the intentionally-stopped/starting states; a failed probe while
the process is alive and was previously ready is logged as DOWN. 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) ## Game detection (foreign GPU holders)
Games and other foreign GPU users sit outside the URL-based consumer model — Games and other foreign GPU users sit outside the URL-based consumer model —
nothing proxies through gpu-turnstile for them. Two independent detection nothing proxies through gpu-turnstile for them. Three independent detection
paths, polled every `GAME_POLL_INTERVAL` (default 5 s); either one being paths, polled every `GAME_POLL_INTERVAL` (default 15 s); any one being
configured enables the feature: configured enables the feature:
- **Process watch list** (`GAME_PROCS`, comma-separated, case-insensitive, - **Process watch list** (`GAME_PROCS`, comma-separated, case-insensitive,
`.exe` optional): while any listed process runs, the GPU counts as held. `.exe` optional): while any listed process runs, the GPU counts as held.
This is the reliable path on Windows. - **Foreign 3D-engine utilization** (`GPU_FOREIGN_UTIL_PCT`, Windows only):
per-process GPU engine counters from PDH — the same data Task Manager
shows — cover graphics work under WDDM, so any process not in
`GPU_IGNORE_PROCS` using more than the threshold percent of the 3D engine
counts as a foreign holder, with no exe list needed. Values are summed per
process across engines; a missing/failing PDH counter disables the path
(logged once).
- **Foreign VRAM threshold** (`GPU_FOREIGN_VRAM_MB`): `nvidia-smi - **Foreign VRAM threshold** (`GPU_FOREIGN_VRAM_MB`): `nvidia-smi
--query-compute-apps` lists per-process GPU memory; any process not in --query-compute-apps` lists per-process GPU memory; any process not in
`GPU_IGNORE_PROCS` (default: Ollama and python — ComfyUI runs under python) `GPU_IGNORE_PROCS` (default: Ollama and python — ComfyUI runs under python
holding more than the threshold counts as a foreign holder. Needs — plus dwm, the desktop compositor) holding more than the threshold counts
nvidia-smi on the PATH (absent: logged once, path disabled) and works best as a foreign holder. Needs nvidia-smi on the PATH (absent: logged once,
on Linux — under Windows' WDDM driver, graphics-only games may not appear path disabled) and works best on Linux — under Windows' WDDM driver,
in the per-process list. graphics-only games may not appear in the per-process list.
While a holder is detected, gpu-turnstile: While a holder is detected, gpu-turnstile:
@@ -211,14 +230,15 @@ override file values. A missing file is fine; a malformed one is fatal.
| `LLM_BUSY_STATUS` | `503` | HTTP status for rejected LLM requests in reject mode (400599, e.g. 429) | | `LLM_BUSY_STATUS` | `503` | HTTP status for rejected LLM requests in reject mode (400599, e.g. 429) |
| `BUSY_RETRY_AFTER` | `30` | seconds sent as `Retry-After` on busy responses (both modes) | | `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 |
| `COMFY_CMD` | _(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_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` | `` | working directory for `COMFY_CMD` | | `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_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 | | `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 | | `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_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 | | `GPU_FOREIGN_UTIL_PCT` | `0` (disabled) | also treat the GPU as held when a process not in `GPU_IGNORE_PROCS` uses more than this percent of the GPU 3D engine (Windows PDH counters, as shown by Task Manager; catches games without an exe list) |
| `GAME_POLL_INTERVAL` | `5s` | how often game/VRAM detection runs | | `GPU_IGNORE_PROCS` | `ollama,ollama app,ollama_llama_server,python,pythonw,dwm` | 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 | | `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_FORMAT` | `text` | `json` for structured JSON logs |
| `LOG_FILE` | `` | append logs to this file instead of stderr (useful as a service) | | `LOG_FILE` | `` | append logs to this file instead of stderr (useful as a service) |
+481 -58
View File
@@ -5,6 +5,7 @@ package main
import ( import (
"bufio" "bufio"
"context" "context"
"encoding/json"
"errors" "errors"
"fmt" "fmt"
"io" "io"
@@ -15,12 +16,16 @@ import (
"os" "os"
"os/signal" "os/signal"
"path/filepath" "path/filepath"
"reflect"
"runtime"
"strings" "strings"
"sync"
"syscall" "syscall"
"time" "time"
"gpu-turnstile/internal/comfy" "gpu-turnstile/internal/comfy"
"gpu-turnstile/internal/config" "gpu-turnstile/internal/config"
"gpu-turnstile/internal/control"
"gpu-turnstile/internal/game" "gpu-turnstile/internal/game"
"gpu-turnstile/internal/lock" "gpu-turnstile/internal/lock"
"gpu-turnstile/internal/metrics" "gpu-turnstile/internal/metrics"
@@ -35,9 +40,15 @@ import (
var version = "dev" var version = "dev"
// exitCodeUpdate tells the service recovery configuration to restart the // exitCodeUpdate tells the service recovery configuration to restart the
// process: a signed update has been staged and the GPU lock is idle. // process: a signed update has been staged or a changed config reload was
// requested, and the GPU lock is idle.
const exitCodeUpdate = 3 const exitCodeUpdate = 3
// exitCodeStaged is returned by an elevated --force-update child when it
// staged a new binary, so the non-elevated parent can tell "updated" from
// "up to date" (it cannot see the child's console).
const exitCodeStaged = 4
// stdoutIsTerminal reports whether stdout is a console (char device), as // stdoutIsTerminal reports whether stdout is a console (char device), as
// opposed to a pipe or file — which is what Docker containers and services // opposed to a pipe or file — which is what Docker containers and services
// see. // see.
@@ -50,7 +61,7 @@ func stdoutIsTerminal() bool {
// --install-service / --remove-service switches, --no-copy, -h/--help, // --install-service / --remove-service switches, --no-copy, -h/--help,
// -v/--version, --force-update and the hidden --elevated-child marker from // -v/--version, --force-update and the hidden --elevated-child marker from
// args. // args.
func parseFlags(args []string) (configPath string, install, remove, noCopy, help, showVersion, forceUpdate, elevatedChild bool, rest []string) { func parseFlags(args []string) (configPath string, install, remove, noCopy, help, showVersion, forceUpdate, updateNow, reloadEnv, monitor, elevatedChild bool, rest []string) {
rest = args[:0] rest = args[:0]
for i := 0; i < len(args); i++ { for i := 0; i < len(args); i++ {
switch { switch {
@@ -59,9 +70,9 @@ func parseFlags(args []string) (configPath string, install, remove, noCopy, help
i++ i++
case strings.HasPrefix(args[i], "-config="): case strings.HasPrefix(args[i], "-config="):
configPath = strings.TrimPrefix(args[i], "-config=") configPath = strings.TrimPrefix(args[i], "-config=")
case args[i] == "--install-service" || args[i] == "-install-service": case args[i] == "--install-service" || args[i] == "-install-service" || args[i] == "-i":
install = true install = true
case args[i] == "--remove-service" || args[i] == "-remove-service": case args[i] == "--remove-service" || args[i] == "-remove-service" || args[i] == "-r":
remove = true remove = true
case args[i] == "--no-copy" || args[i] == "-no-copy": case args[i] == "--no-copy" || args[i] == "-no-copy":
noCopy = true noCopy = true
@@ -71,13 +82,19 @@ func parseFlags(args []string) (configPath string, install, remove, noCopy, help
showVersion = true showVersion = true
case args[i] == "--force-update" || args[i] == "-force-update": case args[i] == "--force-update" || args[i] == "-force-update":
forceUpdate = true forceUpdate = true
case args[i] == "--update-now" || args[i] == "-update-now":
updateNow = true
case args[i] == "--reload-env" || args[i] == "-reload-env":
reloadEnv = true
case args[i] == "--monitor" || args[i] == "-monitor" || args[i] == "-m":
monitor = true
case args[i] == "--elevated-child": case args[i] == "--elevated-child":
elevatedChild = true elevatedChild = true
default: default:
rest = append(rest, args[i]) rest = append(rest, args[i])
} }
} }
return configPath, install, remove, noCopy, help, showVersion, forceUpdate, elevatedChild, rest return configPath, install, remove, noCopy, help, showVersion, forceUpdate, updateNow, reloadEnv, monitor, elevatedChild, rest
} }
// versionLine is printed at the top of every help and error screen. // versionLine is printed at the top of every help and error screen.
@@ -87,11 +104,20 @@ const usageText = `GPU arbitration proxy for Ollama + ComfyUI
Usage: Usage:
gpu-turnstile -config <path> run the proxy gpu-turnstile -config <path> run the proxy
gpu-turnstile --install-service [--no-copy] [-config path] install + start as a service gpu-turnstile -i | --install-service [--no-copy] [-config path]
gpu-turnstile --remove-service stop + uninstall the service install + start as a service
gpu-turnstile -r | --remove-service stop + uninstall the service
gpu-turnstile -v | --version print just the version gpu-turnstile -v | --version print just the version
gpu-turnstile --force-update check for a signed update now, gpu-turnstile --force-update check for a signed update now,
apply it and restart the service apply it and restart the service
(no admin needed when the service runs)
gpu-turnstile --update-now like --force-update, but only
through the running service
gpu-turnstile --reload-env make the service re-read and
validate its config file, then
restart onto it if it changed
gpu-turnstile -m | --monitor live status view (downstreams,
GPU lock, queue); Ctrl+C quits
gpu-turnstile -h | --help this help gpu-turnstile -h | --help this help
Options: Options:
@@ -122,12 +148,12 @@ func fatalUsage(format string, args ...any) {
} }
func main() { func main() {
configPath, install, remove, noCopy, help, showVersion, forceUpdate, elevatedChild, args := parseFlags(os.Args[1:]) configPath, install, remove, noCopy, help, showVersion, forceUpdate, updateNow, reloadEnv, monitor, elevatedChild, args := parseFlags(os.Args[1:])
if showVersion { if showVersion {
fmt.Println(version) fmt.Println(version)
return return
} }
bare := configPath == "" && !install && !remove && !forceUpdate && !elevatedChild && len(args) == 0 bare := configPath == "" && !install && !remove && !forceUpdate && !updateNow && !reloadEnv && !monitor && !elevatedChild && len(args) == 0
if help || (bare && stdoutIsTerminal()) { if help || (bare && stdoutIsTerminal()) {
// Bare invocation in a terminal (e.g. double-clicked on Windows) // Bare invocation in a terminal (e.g. double-clicked on Windows)
// shows the help instead of starting a proxy window with no visible // shows the help instead of starting a proxy window with no visible
@@ -150,12 +176,24 @@ func main() {
fatalUsage("error: --install-service and --remove-service are mutually exclusive") fatalUsage("error: --install-service and --remove-service are mutually exclusive")
case forceUpdate && (install || remove): case forceUpdate && (install || remove):
fatalUsage("error: --force-update cannot be combined with --install-service/--remove-service") fatalUsage("error: --force-update cannot be combined with --install-service/--remove-service")
case updateNow && (install || remove || forceUpdate):
fatalUsage("error: --update-now cannot be combined with other commands")
case reloadEnv && (install || remove || forceUpdate || updateNow):
fatalUsage("error: --reload-env cannot be combined with other commands")
case monitor && (install || remove || forceUpdate || updateNow || reloadEnv):
fatalUsage("error: --monitor cannot be combined with other commands")
case install: case install:
os.Exit(serviceCommand(configPath, true, noCopy, elevatedChild)) os.Exit(serviceCommand(configPath, true, noCopy, elevatedChild))
case remove: case remove:
os.Exit(serviceCommand(configPath, false, noCopy, elevatedChild)) os.Exit(serviceCommand(configPath, false, noCopy, elevatedChild))
case forceUpdate: case forceUpdate:
os.Exit(forceUpdateCommand(configPath, elevatedChild)) os.Exit(forceUpdateCommand(configPath, elevatedChild))
case updateNow:
os.Exit(updateNowCommand())
case reloadEnv:
os.Exit(reloadEnvCommand())
case monitor:
os.Exit(monitorCommand())
} }
if len(args) > 0 { if len(args) > 0 {
fatalUsage("error: unknown arguments: %s", strings.Join(args, " ")) fatalUsage("error: unknown arguments: %s", strings.Join(args, " "))
@@ -175,7 +213,7 @@ func main() {
syncEnvFile(resolveConfigPath(configPath), cfg.LogFile, log) syncEnvFile(resolveConfigPath(configPath), cfg.LogFile, log)
if service.IsService() { if service.IsService() {
if err := service.Run(func(ctx context.Context) error { return run(ctx, cfg, log, logOut, true) }); err != nil { if err := service.Run(func(ctx context.Context) error { return run(ctx, cfg, log, logOut, true, configPath) }); err != nil {
log.Error("service failed", "err", err) log.Error("service failed", "err", err)
os.Exit(1) os.Exit(1)
} }
@@ -183,7 +221,7 @@ func main() {
} }
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer stop() defer stop()
if err := run(ctx, cfg, log, logOut, false); err != nil { if err := run(ctx, cfg, log, logOut, false, configPath); err != nil {
log.Error("listener failed", "err", err) log.Error("listener failed", "err", err)
os.Exit(1) os.Exit(1)
} }
@@ -335,23 +373,35 @@ func forceUpdateCommand(configPath string, elevatedChild bool) int {
fmt.Fprintf(os.Stderr, "gpu-turnstile: cannot locate executable: %v\n", err) fmt.Fprintf(os.Stderr, "gpu-turnstile: cannot locate executable: %v\n", err)
return 1 return 1
} }
// One-shot CLI: the updater logs to stderr, never to the service's
// LOG_FILE — that file is ACL'd to the service account, and a CLI run
// has nothing worth persisting there.
cfg.LogFile = ""
log, _, logCloser := newLogger(cfg) log, _, logCloser := newLogger(cfg)
defer logCloser.Close() defer logCloser.Close()
if cfg.AppVersion == "dev" { if cfg.AppVersion == "dev" {
fmt.Printf("%s: APP_VER=dev, updates disabled\n", versionLine()) fmt.Printf("%s: APP_VER=dev, updates disabled\n", versionLine())
return 0 return 0
} }
// A running service can do the privileged work itself (its account owns
// the install dir and it knows when the GPU is idle): ask it over the
// local control channel first, no admin rights needed. Fails fast when
// no service is listening, in which case we do the direct check below.
if reply, err := control.Ask(control.CmdUpdateNow); err == nil {
return printControlReply(reply)
}
u := &update.Updater{Repo: cfg.UpdateRepo, Asset: cfg.UpdateAsset, Version: version, Desired: cfg.AppVersion, Log: log} u := &update.Updater{Repo: cfg.UpdateRepo, Asset: cfg.UpdateAsset, Version: version, Desired: cfg.AppVersion, Log: log}
// Single-shot: one attempt, fail fast when the server is unreachable // Single-shot: one attempt, fail fast when the server is unreachable
// instead of hanging in a TCP connect for minutes. // instead of hanging in a TCP connect for minutes.
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel() defer cancel()
staged, err := u.Check(ctx, exePath) staged, to, err := u.Check(ctx, exePath)
if err != nil && isPermission(err) && !service.Elevated() { if err != nil && isPermission(err) && !service.Elevated() {
code, _ := elevateAndMirror("--force-update") code, _ := elevateAndMirror("--force-update")
if code == 0 { reportElevatedUpdate(code, to)
fmt.Println("update applied (elevated)") if code == 0 || code == exitCodeStaged {
return 0
} }
return code return code
} }
@@ -363,12 +413,13 @@ func forceUpdateCommand(configPath string, elevatedChild bool) int {
fmt.Printf("%s is up to date\n", versionLine()) fmt.Printf("%s is up to date\n", versionLine())
return 0 return 0
} }
fmt.Printf("%s: update staged\n", versionLine()) fmt.Printf("gpu-turnstile: updated from %s to %s\n", version, to)
restarted, err := service.RestartIfRunning() restarted, err := service.RestartIfRunning()
if err != nil && isPermission(err) && !service.Elevated() { if err != nil && isPermission(err) && !service.Elevated() {
code, _ := elevateAndMirror("--force-update") code, _ := elevateAndMirror("--force-update")
if code == 0 { reportElevatedUpdate(code, to)
fmt.Println("update applied (elevated)") if code == 0 || code == exitCodeStaged {
return 0
} }
return code return code
} }
@@ -377,13 +428,69 @@ func forceUpdateCommand(configPath string, elevatedChild bool) int {
return 1 return 1
} }
if restarted { if restarted {
fmt.Println("service restarted on the new version") fmt.Println("service restarted on " + to)
} else { } else {
fmt.Println("no running service; the new version applies on next start") fmt.Println("no running service; " + to + " applies on the next start")
}
if elevatedChild {
// Tell the non-elevated parent (which cannot see this console)
// whether anything was staged, so its mirror message is honest.
return exitCodeStaged
} }
return 0 return 0
} }
// printControlReply prints the service's answer to a control-channel
// request: "OK ..." on stdout (exit 0), "ERR ..." on stderr (exit 1).
func printControlReply(reply string) int {
if msg, ok := strings.CutPrefix(reply, "ERR "); ok {
fmt.Fprintf(os.Stderr, "gpu-turnstile: %s\n", msg)
return 1
}
fmt.Println("gpu-turnstile: " + strings.TrimPrefix(reply, "OK "))
return 0
}
// updateNowCommand only goes through the running service's control channel
// (no direct check, no elevation): the unprivileged update trigger.
func updateNowCommand() int {
reply, err := control.Ask(control.CmdUpdateNow)
if err != nil {
fmt.Fprintf(os.Stderr, "%s\n\ngpu-turnstile: no running service to ask — use --force-update for a direct check\n", versionLine())
return 1
}
return printControlReply(reply)
}
// reloadEnvCommand asks the running service, over the control channel, to
// re-read and validate its config file. The client sends no path — the
// service only ever re-reads its own configured file.
func reloadEnvCommand() int {
reply, err := control.Ask(control.CmdReloadEnv)
if err != nil {
fmt.Fprintf(os.Stderr, "%s\n\ngpu-turnstile: no running service to ask\n", versionLine())
return 1
}
return printControlReply(reply)
}
// reportElevatedUpdate prints the parent's summary of an elevated
// --force-update child: exitCodeStaged means the child staged a new binary,
// 0 means it found nothing to do. to is the tag the parent's own check
// resolved before it hit the permission wall ("" when it never got that
// far).
func reportElevatedUpdate(code int, to string) {
switch {
case code == exitCodeStaged && to != "":
fmt.Printf("gpu-turnstile: updated from %s to %s (elevated)\n", version, to)
case code == exitCodeStaged:
fmt.Println("update applied (elevated)")
case code == 0:
fmt.Printf("%s is up to date\n", versionLine())
}
// Non-zero, non-staged codes: elevateAndMirror already printed the failure.
}
// serviceCommand installs (copyBin = register the canonical-layout copy) // serviceCommand installs (copyBin = register the canonical-layout copy)
// or removes the service and reports the result. On Windows, when the // or removes the service and reports the result. On Windows, when the
// shell is not elevated, the command relaunches itself through a UAC // shell is not elevated, the command relaunches itself through a UAC
@@ -434,7 +541,26 @@ func orDisabled(url string) string {
return url return url
} }
func run(ctx context.Context, cfg config.Config, log *slog.Logger, logOut io.Writer, isService bool) error { // managedComfyCommand resolves how ComfyUI is launched when it is managed:
// COMFY_CMD verbatim, or the standard venv layout under COMFY_DIR. Empty
// when neither is set (unmanaged).
func managedComfyCommand(cfg config.Config) string {
if cfg.ComfyCmd != "" {
return cfg.ComfyCmd
}
if cfg.ComfyDir != "" {
return supervise.DefaultComfyCommand(runtime.GOOS, cfg.ComfyDir, cfg.ComfyURL)
}
return ""
}
func run(ctx context.Context, cfg config.Config, log *slog.Logger, logOut io.Writer, isService bool, configPath string) error {
// ComfyUI can run as a managed child — COMFY_CMD verbatim, or the
// standard venv layout derived from COMFY_DIR alone: started on demand
// by the proxy, stopped after COMFY_IDLE_TIMEOUT idle (and on shutdown)
// so its VRAM is freed.
comfyCmdLine := managedComfyCommand(cfg)
// The startup line carries the version and every setting and is emitted // The startup line carries the version and every setting and is emitted
// at WARN so it is visible even with the default (quiet) log level. // at WARN so it is visible even with the default (quiet) log level.
log.Log(ctx, slog.LevelWarn, "starting gpu-turnstile", log.Log(ctx, slog.LevelWarn, "starting gpu-turnstile",
@@ -460,12 +586,13 @@ func run(ctx context.Context, cfg config.Config, log *slog.Logger, logOut io.Wri
"backoff_max", cfg.BackoffMax, "backoff_max", cfg.BackoffMax,
"prompt_capture_limit", cfg.PromptCaptureLimit, "prompt_capture_limit", cfg.PromptCaptureLimit,
"warm_model", cfg.WarmModel, "warm_model", cfg.WarmModel,
"comfy_cmd", cfg.ComfyCmd, "comfy_cmd", orDisabled(comfyCmdLine),
"comfy_dir", cfg.ComfyDir, "comfy_dir", cfg.ComfyDir,
"comfy_idle_timeout", cfg.ComfyIdleTimeout, "comfy_idle_timeout", cfg.ComfyIdleTimeout,
"comfy_start_timeout", cfg.ComfyStartTimeout, "comfy_start_timeout", cfg.ComfyStartTimeout,
"game_procs", cfg.GameProcs, "game_procs", cfg.GameProcs,
"gpu_foreign_vram_mb", cfg.GPUForeignVRAMMB, "gpu_foreign_vram_mb", cfg.GPUForeignVRAMMB,
"gpu_foreign_util_pct", cfg.GPUForeignUtilPct,
"gpu_ignore_procs", cfg.GPUIgnoreProcs, "gpu_ignore_procs", cfg.GPUIgnoreProcs,
"game_poll_interval", cfg.GamePollInterval, "game_poll_interval", cfg.GamePollInterval,
"auto_update", cfg.AutoUpdate, "auto_update", cfg.AutoUpdate,
@@ -478,6 +605,8 @@ func run(ctx context.Context, cfg config.Config, log *slog.Logger, logOut io.Wri
) )
lk := lock.New(log) lk := lock.New(log)
health := newHealthTracker()
started := time.Now()
// Each consumer is enabled by setting its URL; a disabled consumer gets // Each consumer is enabled by setting its URL; a disabled consumer gets
// no client, no listener and no probe. // no client, no listener and no probe.
var ollamaClient *ollama.Client var ollamaClient *ollama.Client
@@ -494,13 +623,27 @@ func run(ctx context.Context, cfg config.Config, log *slog.Logger, logOut io.Wri
} }
} }
// With COMFY_CMD set, ComfyUI runs as a managed child: started on
// demand by the proxy, stopped after COMFY_IDLE_TIMEOUT idle (and on
// shutdown) so its VRAM is freed.
var comfySup *supervise.Process var comfySup *supervise.Process
if cfg.ComfyCmd != "" { if comfyCmdLine != "" {
if cfg.ComfyCmd == "" {
// Derived from COMFY_DIR: flag a wrong-looking layout early,
// while the operator is still watching the startup log. Inside
// a profile the service account may not enter, os.Stat fails
// with EACCES — that reads as "not found" but means "grant
// access", so say so.
python, script := supervise.ComfyLayout(runtime.GOOS, cfg.ComfyDir)
for _, p := range []string{python, script} {
if _, err := os.Stat(p); err != nil {
if isPermission(err) {
log.Warn("COMFY_DIR: not accessible to the service account; grant access or re-run --install-service", "path", p)
} else {
log.Warn("COMFY_DIR: file not found; ComfyUI requests will fail until it exists", "path", p)
}
}
}
}
var err error var err error
comfySup, err = supervise.New("comfy", cfg.ComfyCmd, cfg.ComfyDir, comfyClient.Probe, cfg.ComfyStartTimeout, log) comfySup, err = supervise.New("comfy", comfyCmdLine, cfg.ComfyDir, comfyClient.Probe, cfg.ComfyStartTimeout, log)
if err != nil { if err != nil {
return err return err
} }
@@ -573,20 +716,25 @@ func run(ctx context.Context, cfg config.Config, log *slog.Logger, logOut io.Wri
} }
probeCtx, probeCancel := context.WithTimeout(ctx, cfg.ProbeTimeout) probeCtx, probeCancel := context.WithTimeout(ctx, cfg.ProbeTimeout)
for name, probe := range probes { for name, probe := range probes {
if err := probe(probeCtx); err != nil && !errors.Is(err, errManagedDown) { err := probe(probeCtx)
health.set(name, err == nil)
if err != nil && !errors.Is(err, errManagedDown) {
log.Warn(name+" probe failed", "err", err) log.Warn(name+" probe failed", "err", err)
} }
} }
probeCancel() probeCancel()
if cfg.HealthInterval > 0 { if cfg.HealthInterval > 0 {
go healthLoop(ctx, cfg.HealthInterval, cfg.ProbeTimeout, log, probes) go healthLoop(ctx, cfg.HealthInterval, cfg.ProbeTimeout, log, probes, health)
} }
// Foreign GPU holders (games, other ML jobs) — enabled by GAME_PROCS // Foreign GPU holders (games, other ML jobs) — enabled by GAME_PROCS,
// and/or GPU_FOREIGN_VRAM_MB — hold the lock externally while they run. // GPU_FOREIGN_VRAM_MB and/or GPU_FOREIGN_UTIL_PCT — hold the lock
if len(cfg.GameProcs) > 0 || cfg.GPUForeignVRAMMB > 0 { // externally while they run. gw collects the VRAM reading and the last
det := game.New(cfg.GameProcs, cfg.GPUForeignVRAMMB, cfg.GPUIgnoreProcs, log) // check result for the status channel.
go gameLoop(ctx, cfg, log, det, lk, ollamaClient, comfySup) gw := &gpuWatch{enabled: len(cfg.GameProcs) > 0 || cfg.GPUForeignVRAMMB > 0 || cfg.GPUForeignUtilPct > 0}
if gw.enabled {
det := game.New(cfg.GameProcs, cfg.GPUForeignVRAMMB, cfg.GPUForeignUtilPct, cfg.GPUIgnoreProcs, log)
go gameLoop(ctx, cfg, log, det, lk, ollamaClient, comfySup, gw)
} }
// Bind the listeners up front so a port conflict fails fast and the // Bind the listeners up front so a port conflict fails fast and the
@@ -624,8 +772,49 @@ func run(ctx context.Context, cfg config.Config, log *slog.Logger, logOut io.Wri
service.NotifyReady() service.NotifyReady()
service.StartWatchdog(ctx) service.StartWatchdog(ctx)
var u *update.Updater
var exePath string
// restartWhenIdle exits with exitCodeUpdate once the GPU lock is idle;
// the service recovery configuration brings the process back. Shared by
// staged updates and config reloads; the once guard makes repeated
// triggers idempotent.
var restartOnce sync.Once
restartWhenIdle := func(reason string) {
log.Warn("restarting once the GPU is idle", "reason", reason)
restartOnce.Do(func() {
go func() {
if waitForIdle(ctx, lk, 24*time.Hour) {
log.Warn("restarting now", "reason", reason)
os.Exit(exitCodeUpdate)
}
}()
})
}
applyStaged := func(to string) {}
if cfg.AutoUpdate { if cfg.AutoUpdate {
go updateLoop(ctx, cfg, log, lk, isService) if p, err := os.Executable(); err != nil {
log.Warn("auto-update disabled: cannot locate executable", "err", err)
} else {
exePath = p
u = &update.Updater{Repo: cfg.UpdateRepo, Asset: cfg.UpdateAsset, Version: version, Desired: cfg.AppVersion, Log: log}
applyStaged = func(to string) {
if !isService {
log.Warn("auto-update: new binary staged; restart gpu-turnstile to apply", "version", to)
return
}
restartWhenIdle("update to " + to)
}
go updateLoop(ctx, cfg.UpdateInterval, log, u, exePath, applyStaged)
}
}
// The control channel (status for --monitor, update-now and reload-env
// triggers) is served whenever running as a service, independent of
// AUTO_UPDATE.
if isService {
serveControl(ctx, log, u, exePath, applyStaged,
statusProvider(cfg, lk, comfySup, health, started, gw),
reloadHandler(cfg, configPath, restartWhenIdle))
} }
select { select {
@@ -649,12 +838,48 @@ func run(ctx context.Context, cfg config.Config, log *slog.Logger, logOut io.Wri
// (idle); health checks skip it instead of logging an outage. // (idle); health checks skip it instead of logging an outage.
var errManagedDown = errors.New("managed upstream intentionally stopped") var errManagedDown = errors.New("managed upstream intentionally stopped")
// gpuWatch records the latest VRAM reading and game-detector finding from
// the game detector's poll loop, for the status channel. Enabled is false
// when game detection is not configured (no polling happens then); Known is
// false until the first successful nvidia-smi reading.
type gpuWatch struct {
mu sync.Mutex
enabled bool
usedMB int
total int
known bool
foreign string // last Check result: external holders, "" when none
at time.Time
}
func (g *gpuWatch) setVRAM(used, total int) {
g.mu.Lock()
g.usedMB, g.total, g.known = used, total, true
g.mu.Unlock()
}
func (g *gpuWatch) setCheck(foreign string) {
g.mu.Lock()
g.foreign, g.at = foreign, time.Now()
g.mu.Unlock()
}
func (g *gpuWatch) get() (used, total int, known bool, foreign string, ageS int64) {
g.mu.Lock()
defer g.mu.Unlock()
ageS = -1
if !g.at.IsZero() {
ageS = int64(time.Since(g.at).Seconds())
}
return g.usedMB, g.total, g.known, g.foreign, ageS
}
// gameLoop polls for foreign GPU holders (a game, another ML job). While one // gameLoop polls for foreign GPU holders (a game, another ML job). While one
// is detected it holds the lock externally so new LLM and image requests // is detected it holds the lock externally so new LLM and image requests
// wait (or are rejected per LLM_BUSY_MODE), and — once in-flight work has // wait (or are rejected per LLM_BUSY_MODE), and — once in-flight work has
// drained — frees VRAM for it: the managed ComfyUI is stopped and Ollama's // drained — frees VRAM for it: the managed ComfyUI is stopped and Ollama's
// resident models are unloaded. // resident models are unloaded.
func gameLoop(ctx context.Context, cfg config.Config, log *slog.Logger, det *game.Detector, lk *lock.Lock, ollamaClient *ollama.Client, comfySup *supervise.Process) { func gameLoop(ctx context.Context, cfg config.Config, log *slog.Logger, det *game.Detector, lk *lock.Lock, ollamaClient *ollama.Client, comfySup *supervise.Process, gw *gpuWatch) {
ticker := time.NewTicker(cfg.GamePollInterval) ticker := time.NewTicker(cfg.GamePollInterval)
defer ticker.Stop() defer ticker.Stop()
held, freed := false, false held, freed := false, false
@@ -668,6 +893,10 @@ func gameLoop(ctx context.Context, cfg config.Config, log *slog.Logger, det *gam
if err != nil && ctx.Err() == nil { if err != nil && ctx.Err() == nil {
log.Warn("game detection failed", "err", err) log.Warn("game detection failed", "err", err)
} }
gw.setCheck(summarizeHolders(holders))
if used, total, verr := game.QueryVRAMMB(ctx); verr == nil {
gw.setVRAM(used, total)
}
switch { switch {
case len(holders) > 0 && !held: case len(holders) > 0 && !held:
held = true held = true
@@ -727,7 +956,8 @@ func freeVRAM(ctx context.Context, unloadTimeout time.Duration, log *slog.Logger
// transitions — "is DOWN" when a previously healthy upstream stops // transitions — "is DOWN" when a previously healthy upstream stops
// answering, "recovered" when it comes back. The first round only // answering, "recovered" when it comes back. The first round only
// establishes the baseline; the startup probe already reported that state. // establishes the baseline; the startup probe already reported that state.
func healthLoop(ctx context.Context, interval, probeTimeout time.Duration, log *slog.Logger, probes map[string]func(context.Context) error) { // Every result goes into the tracker for the status channel.
func healthLoop(ctx context.Context, interval, probeTimeout time.Duration, log *slog.Logger, probes map[string]func(context.Context) error, tracker *healthTracker) {
ticker := time.NewTicker(interval) ticker := time.NewTicker(interval)
defer ticker.Stop() defer ticker.Stop()
up := map[string]bool{} up := map[string]bool{}
@@ -744,9 +974,11 @@ func healthLoop(ctx context.Context, interval, probeTimeout time.Duration, log *
cancel() cancel()
if errors.Is(err, errManagedDown) { if errors.Is(err, errManagedDown) {
managed[name] = true // intentionally stopped; not an outage managed[name] = true // intentionally stopped; not an outage
tracker.set(name, false)
continue continue
} }
now := err == nil now := err == nil
tracker.set(name, now)
if managed[name] { if managed[name] {
// First real probe after an idle stop only re-baselines — // First real probe after an idle stop only re-baselines —
// an on-demand start is not a "recovery". // an on-demand start is not a "recovery".
@@ -767,42 +999,233 @@ func healthLoop(ctx context.Context, interval, probeTimeout time.Duration, log *
} }
} }
// updateLoop checks for signed updates on startup and every UPDATE_INTERVAL. // updateLoop checks for signed updates on startup and every interval;
// In service mode a staged update is applied by exiting with exitCodeUpdate // applyStaged decides what a staged update means (restart when idle as a
// once the GPU lock is idle; the service recovery configuration restarts the // service, log only interactively).
// process with the new binary. Interactively it only logs. func updateLoop(ctx context.Context, interval time.Duration, log *slog.Logger, u *update.Updater, exePath string, applyStaged func(to string)) {
func updateLoop(ctx context.Context, cfg config.Config, log *slog.Logger, lk *lock.Lock, isService bool) {
exePath, err := os.Executable()
if err != nil {
log.Warn("auto-update disabled: cannot locate executable", "err", err)
return
}
u := &update.Updater{Repo: cfg.UpdateRepo, Asset: cfg.UpdateAsset, Version: version, Desired: cfg.AppVersion, Log: log}
for { for {
staged, err := u.Check(ctx, exePath) staged, to, err := u.Check(ctx, exePath)
if err != nil && ctx.Err() == nil { if err != nil && ctx.Err() == nil {
log.Warn("auto-update check failed", "err", err) log.Warn("auto-update check failed", "err", err)
} }
if staged { if staged {
if !isService { applyStaged(to)
log.Warn("auto-update: new binary staged; restart gpu-turnstile to apply")
return
}
log.Warn("auto-update: staged; restarting once the GPU is idle")
if waitForIdle(ctx, lk, 24*time.Hour) {
log.Warn("auto-update: restarting to apply update")
os.Exit(exitCodeUpdate)
}
return return
} }
select { select {
case <-ctx.Done(): case <-ctx.Done():
return return
case <-time.After(cfg.UpdateInterval): case <-time.After(interval):
} }
} }
} }
// serveControl opens the local control channel (named pipe on Windows,
// unix socket on Linux) so unprivileged local users can query status
// (--monitor), trigger an update check (--force-update/--update-now) and
// poke a config reload (--reload-env) without admin rights. The update
// payload is signature-verified regardless of who asks; update triggers
// are rate-limited to one per minute so the channel cannot be used to spam
// restarts. u is nil when AUTO_UPDATE=false.
func serveControl(ctx context.Context, log *slog.Logger, u *update.Updater, exePath string, applyStaged func(to string), status func() string, reload func() string) {
var mu sync.Mutex
var lastTrigger, lastReload time.Time
h := func(cmd string) string {
switch cmd {
case control.CmdStatus:
return "OK " + status()
case control.CmdReloadEnv:
// reload re-reads the config file from disk; a short limiter
// keeps a local flood from turning into disk churn.
mu.Lock()
if wait := 2*time.Second - time.Since(lastReload); wait > 0 {
mu.Unlock()
return fmt.Sprintf("ERR rate limited: retry in %ds", int(wait.Seconds())+1)
}
lastReload = time.Now()
mu.Unlock()
return reload()
case control.CmdUpdateNow:
default:
return "ERR unknown command: " + cmd
}
if u == nil {
return "ERR auto-update is disabled on this instance"
}
mu.Lock()
if wait := time.Minute - time.Since(lastTrigger); wait > 0 {
mu.Unlock()
return fmt.Sprintf("ERR rate limited: retry in %ds", int(wait.Seconds())+1)
}
lastTrigger = time.Now()
mu.Unlock()
cctx, cancel := context.WithTimeout(ctx, 30*time.Second)
defer cancel()
staged, to, err := u.Check(cctx, exePath)
if err != nil {
return "ERR update check failed: " + err.Error()
}
if !staged {
return "OK " + version + " is up to date"
}
applyStaged(to)
return "OK updated from " + version + " to " + to + "; the service restarts once the GPU is idle"
}
if err := control.Serve(ctx, h, log); err != nil {
log.Warn("control channel disabled", "err", err)
}
}
// healthTracker records the latest probe result per upstream for the
// status channel.
type healthTracker struct {
mu sync.Mutex
up map[string]bool
}
func newHealthTracker() *healthTracker {
return &healthTracker{up: map[string]bool{}}
}
func (h *healthTracker) set(name string, up bool) {
h.mu.Lock()
h.up[name] = up
h.mu.Unlock()
}
func (h *healthTracker) get(name string) bool {
h.mu.Lock()
defer h.mu.Unlock()
return h.up[name]
}
// statusDownstream/statusLock/statusSnapshot are the JSON the control
// channel serves on CmdStatus; the monitor mode renders them.
type statusDownstream struct {
Name string `json:"name"`
URL string `json:"url"`
Up bool `json:"up"`
Managed string `json:"managed,omitempty"`
}
type statusLock struct {
State string `json:"state"`
Detail string `json:"detail,omitempty"`
LLMInflight int `json:"llm_inflight"`
LLMWaiting int `json:"llm_waiting"`
ImageQueue int `json:"image_queue"`
External string `json:"external,omitempty"`
SinceS int64 `json:"since_s"`
}
type statusSnapshot struct {
Version string `json:"version"`
UptimeS int64 `json:"uptime_s"`
Downstreams []statusDownstream `json:"downstreams"`
Lock statusLock `json:"lock"`
// GPU carries the latest VRAM reading and detector finding; Enabled is
// false when game detection (and with it nvidia-smi polling) is not
// configured.
GPU statusGPU `json:"gpu"`
// MonitorNote is set client-side (never over the wire) when the
// monitor's own binary differs from the service's version.
MonitorNote string `json:"-"`
}
type statusGPU struct {
// Enabled reports whether game detection is configured (and with it
// VRAM polling); when false the other fields carry no information.
Enabled bool `json:"enabled"`
UsedMB int `json:"used_mb"`
TotalMB int `json:"total_mb"`
Known bool `json:"known"`
// Foreign is the last detector finding (external GPU holders), empty
// when the last check found none.
Foreign string `json:"foreign,omitempty"`
// AgeS is how long ago the last check ran; -1 before the first check.
AgeS int64 `json:"age_s"`
}
// reloadHandler re-reads and validates the service's config file for
// CmdReloadEnv. An invalid config is reported and the service keeps running
// untouched; a valid, changed config triggers a GPU-idle-gated restart onto
// it (same mechanism as staged updates); unchanged is a no-op.
func reloadHandler(current config.Config, configPath string, restartWhenIdle func(reason string)) func() string {
return func() string {
ncfg, err := loadMergedConfig(configPath)
if err != nil {
return "ERR config invalid: " + err.Error()
}
changed := diffConfig(current, ncfg)
if len(changed) == 0 {
return "OK config unchanged"
}
restartWhenIdle("config reload (" + strings.Join(changed, ", ") + ")")
return "OK config valid; restarting once the GPU is idle (changed: " + strings.Join(changed, ", ") + ")"
}
}
// diffConfig lists the env names of settings whose values differ between
// two configs (from each field's env tag, so users recognize them).
func diffConfig(a, b config.Config) []string {
va, vb := reflect.ValueOf(a), reflect.ValueOf(b)
t := va.Type()
var out []string
for i := 0; i < t.NumField(); i++ {
if !reflect.DeepEqual(va.Field(i).Interface(), vb.Field(i).Interface()) {
name := t.Field(i).Tag.Get("env")
if name == "" {
name = t.Field(i).Name
}
out = append(out, name)
}
}
return out
}
// statusProvider assembles the one-line JSON snapshot for CmdStatus.
func statusProvider(cfg config.Config, lk *lock.Lock, comfySup *supervise.Process, health *healthTracker, started time.Time, gw *gpuWatch) func() string {
return func() string {
snap := statusSnapshot{
Version: version,
UptimeS: int64(time.Since(started).Seconds()),
}
used, total, known, foreign, ageS := gw.get()
snap.GPU = statusGPU{
Enabled: gw.enabled,
UsedMB: used, TotalMB: total, Known: known,
Foreign: foreign, AgeS: ageS,
}
if cfg.OllamaURL != "" {
snap.Downstreams = append(snap.Downstreams, statusDownstream{
Name: "ollama", URL: cfg.OllamaURL, Up: health.get("ollama"),
})
}
if cfg.ComfyURL != "" {
d := statusDownstream{Name: "comfy", URL: cfg.ComfyURL, Up: health.get("comfy")}
if comfySup != nil {
d.Managed = comfySup.Status()
}
snap.Downstreams = append(snap.Downstreams, d)
}
st := lk.Status()
snap.Lock = statusLock{
State: string(st.State),
Detail: st.Detail,
LLMInflight: st.LLMInflight,
LLMWaiting: st.LLMWaiting,
ImageQueue: st.ImageQueue,
External: st.External,
SinceS: int64(time.Since(st.Since).Seconds()),
}
b, err := json.Marshal(snap)
if err != nil {
return `{"version":"` + version + `"}`
}
return string(b)
}
}
// waitForIdle polls the lock until no LLM or image work is active or // waitForIdle polls the lock until no LLM or image work is active or
// pending, max at most. Returns false on timeout or cancellation. // pending, max at most. Returns false on timeout or cancellation.
func waitForIdle(ctx context.Context, lk *lock.Lock, max time.Duration) bool { func waitForIdle(ctx context.Context, lk *lock.Lock, max time.Duration) bool {
+303
View File
@@ -0,0 +1,303 @@
package main
import (
"encoding/json"
"fmt"
"os"
"os/exec"
"strings"
"time"
"gpu-turnstile/internal/control"
)
// ANSI colors for the monitor frame.
const (
cReset = "\x1b[0m"
cDim = "\x1b[2m"
cRed = "\x1b[31m"
cGreen = "\x1b[32m"
cYellow = "\x1b[33m"
cCyan = "\x1b[36m"
)
// hotkeysLine is the monitor's footer.
const hotkeysLine = " " + cDim + "q quit · u update now" + cReset + "\x1b[K\n"
// monitorCommand renders a live status view of the running service,
// refreshed every second from the control channel. When the service
// reports a different version and the executable on disk changed (the
// updater replaced it), the monitor restarts itself onto the new binary.
// Hotkeys: q quits, u triggers an update check on the service.
func monitorCommand() int {
if !stdoutIsTerminal() {
fmt.Fprintln(os.Stderr, "gpu-turnstile: --monitor needs an interactive terminal")
return 1
}
enableVirtualTerminal()
restore := enableRawKeys()
defer func() {
if restore != nil {
restore()
}
}()
fmt.Print("\x1b[2J") // clear once; frames then redraw in place
defer fmt.Print(cReset + "\n")
exe, _ := os.Executable()
var exeStamp time.Time
if st, err := os.Stat(exe); err == nil {
exeStamp = st.ModTime()
}
keys := make(chan byte, 8)
go readKeys(keys)
ticker := time.NewTicker(time.Second)
defer ticker.Stop()
var note string
var noteAt time.Time
noteCh := make(chan string, 1)
updatePending := false
poll := func() string {
frame := renderWaiting()
if reply, err := control.Ask(control.CmdStatus); err == nil {
if msg, ok := strings.CutPrefix(reply, "OK "); ok {
var snap statusSnapshot
if json.Unmarshal([]byte(msg), &snap) == nil {
if snap.Version != "" && snap.Version != version {
if exeChanged(exe, exeStamp) {
fmt.Print("\x1b[2J\x1b[H")
fmt.Printf("gpu-turnstile: service updated to %s — restarting the monitor\n", snap.Version)
restartSelf(exe, "--monitor")
return "" // re-execed; this process exits below
}
snap.MonitorNote = fmt.Sprintf("note: the service runs %s, this monitor is %s", snap.Version, version)
}
if note != "" {
snap.MonitorNote = note
}
frame = renderMonitor(snap, termWidth())
}
}
}
return frame
}
for {
frame := poll()
if frame == "" {
return 0 // restartSelf fired
}
fmt.Print("\x1b[H" + frame + "\x1b[J") // home, frame, clear below
select {
case <-ticker.C:
if note != "" && time.Since(noteAt) > 15*time.Second {
note = ""
}
case k, ok := <-keys:
if !ok {
keys = nil
continue
}
switch k {
case 'q', 'Q', 3: // q or Ctrl+C (raw mode delivers it as a byte)
return 0
case 'u', 'U':
if !updatePending {
updatePending = true
note, noteAt = "checking for updates…", time.Now()
go func() {
reply, err := control.Ask(control.CmdUpdateNow)
if err != nil {
noteCh <- "update: no answer from the service"
return
}
msg := strings.TrimPrefix(reply, "OK ")
msg = strings.TrimPrefix(msg, "ERR ")
noteCh <- "update: " + msg
}()
}
}
case n := <-noteCh:
updatePending = false
note, noteAt = n, time.Now()
}
}
}
// readKeys reads single keypresses from stdin (raw mode was enabled by the
// caller) and delivers them until stdin fails.
func readKeys(keys chan<- byte) {
defer close(keys)
buf := make([]byte, 1)
for {
n, err := os.Stdin.Read(buf)
if n > 0 {
keys <- buf[0]
}
if err != nil {
return
}
}
}
// exeChanged reports whether the executable on disk was replaced since the
// recorded stamp (the updater swaps it via rename, which changes ModTime).
func exeChanged(exe string, stamp time.Time) bool {
if exe == "" || stamp.IsZero() {
return false
}
st, err := os.Stat(exe)
return err == nil && !st.ModTime().Equal(stamp)
}
// restartSelf starts a fresh copy of this executable with the given args on
// the same console; the caller exits right after.
func restartSelf(exe string, args ...string) {
cmd := exec.Command(exe, args...)
cmd.Stdin, cmd.Stdout, cmd.Stderr = os.Stdin, os.Stdout, os.Stderr
cmd.Start() //nolint:errcheck // best effort: on failure we just exit
}
func renderWaiting() string {
return cDim + " gpu-turnstile — waiting for a running service…" + cReset + "\x1b[K\n\x1b[K\n" + hotkeysLine
}
// renderMonitor draws one full frame. Each line ends with \x1b[K (clear to
// end of line) so shrinking content leaves no residue.
func renderMonitor(snap statusSnapshot, width int) string {
if width < 40 {
width = 80
}
var b strings.Builder
left := " gpu-turnstile"
if snap.UptimeS > 0 {
left += " " + cDim + "up " + fmtDur(snap.UptimeS) + cReset
}
right := snap.Version
pad := width - printableLen(" gpu-turnstile up "+fmtDur(snap.UptimeS)) - len(right) - 1
if snap.UptimeS == 0 {
pad = width - len(" gpu-turnstile") - len(right) - 1
}
if pad < 1 {
pad = 1
}
b.WriteString(cDim + left + strings.Repeat(" ", pad) + right + cReset + "\x1b[K\n")
b.WriteString(cDim + " " + strings.Repeat("─", width-2) + cReset + "\x1b[K\n")
for _, d := range snap.Downstreams {
b.WriteString(renderDownstream(d) + "\x1b[K\n")
}
b.WriteString("\x1b[K\n")
b.WriteString(renderLock(snap.Lock) + "\x1b[K\n")
if snap.Lock.ImageQueue > 0 {
b.WriteString(fmt.Sprintf(" Queue: %s%d image job(s) waiting%s\x1b[K\n",
cYellow, snap.Lock.ImageQueue, cReset))
}
if snap.GPU.Enabled {
b.WriteString(renderGPU(snap.GPU) + "\x1b[K\n")
}
if snap.MonitorNote != "" {
b.WriteString(" " + cYellow + snap.MonitorNote + cReset + "\x1b[K\n")
}
b.WriteString("\x1b[K\n")
b.WriteString(hotkeysLine)
return b.String()
}
// renderGPU renders the GPU line: VRAM usage (when nvidia-smi answered),
// the game detector's last finding, and how long ago it ran.
func renderGPU(g statusGPU) string {
s := " GPU: "
if g.Known {
s += renderVRAM(g.UsedMB, g.TotalMB)
} else {
s += cDim + "VRAM unknown (nvidia-smi not answering)" + cReset
}
if g.Foreign != "" {
s += " · external: " + cRed + g.Foreign + cReset
} else {
s += cDim + " · no external process" + cReset
}
if g.AgeS >= 0 {
s += cDim + fmt.Sprintf(" (checked %s ago)", fmtDur(g.AgeS)) + cReset
}
return s
}
// renderVRAM renders "4.2 / 16.0 GiB used" (or MiB below 1 GiB).
func renderVRAM(used, total int) string {
format := func(mb int) string {
if mb >= 1024 {
return fmt.Sprintf("%.1f GiB", float64(mb)/1024)
}
return fmt.Sprintf("%d MiB", mb)
}
if total > 0 {
return format(used) + " / " + format(total) + " used"
}
return format(used) + " used"
}
// printableLen counts characters without ANSI escapes (ASCII-only content).
func printableLen(s string) int { return len(s) }
func renderDownstream(d statusDownstream) string {
url := cDim + d.URL + cReset
switch d.Managed {
case "stopped":
return fmt.Sprintf(" %s○%s %-8s %sstopped (managed — starts on demand)%s %s",
cDim, cReset, d.Name, cDim, cReset, url)
case "starting":
return fmt.Sprintf(" %s◌%s %-8s %sstarting…%s %s",
cYellow, cReset, d.Name, cYellow, cReset, url)
}
suffix := ""
if d.Managed == "external" {
suffix = " (external)"
}
if d.Up {
return fmt.Sprintf(" %s●%s %-8s %sUP%s%s %s", cGreen, cReset, d.Name, cGreen, cReset, suffix, url)
}
return fmt.Sprintf(" %s●%s %-8s %sDOWN%s %s", cRed, cReset, d.Name, cRed, cReset, url)
}
func renderLock(l statusLock) string {
dur := cDim + "(" + fmtDur(l.SinceS) + ")" + cReset
switch l.State {
case "idle":
return fmt.Sprintf(" Lock: %sidle%s %s", cGreen, cReset, dur)
case "llm":
s := fmt.Sprintf(" Lock: %sLLM%s — %d in flight", cCyan, cReset, l.LLMInflight)
if l.LLMWaiting > 0 {
s += fmt.Sprintf(", %d waiting", l.LLMWaiting)
}
if l.Detail != "" {
s += " — " + l.Detail
}
return s + " " + dur
case "image":
s := fmt.Sprintf(" Lock: %sIMAGE%s", cYellow, cReset)
if l.Detail != "" {
s += " — " + l.Detail
}
return s + " " + dur
case "external":
return fmt.Sprintf(" Lock: %sEXTERNAL%s — %s %s", cRed, cReset, l.External, dur)
}
return " Lock: unknown"
}
// fmtDur renders seconds as a compact duration ("1m32s", "2h07m").
func fmtDur(s int64) string {
if s < 0 {
s = 0
}
d := time.Duration(s) * time.Second
if d >= time.Hour {
return fmt.Sprintf("%dh%02dm", int(d.Hours()), int(d.Minutes())%60)
}
return d.Round(time.Second).String()
}
+85
View File
@@ -0,0 +1,85 @@
package main
import (
"gpu-turnstile/internal/config"
"reflect"
"strings"
"testing"
)
func TestRenderMonitor(t *testing.T) {
snap := statusSnapshot{
Version: "v0.2.3",
UptimeS: 3725,
Downstreams: []statusDownstream{
{Name: "ollama", URL: "http://127.0.0.1:11435", Up: true},
{Name: "comfy", URL: "http://127.0.0.1:8189", Managed: "stopped"},
},
Lock: statusLock{State: "llm", LLMInflight: 2, LLMWaiting: 1, Detail: "ollama: POST /api/generate", SinceS: 95},
}
frame := renderMonitor(snap, 80)
for _, want := range []string{"v0.2.3", "1h02m", "ollama", "UP", "comfy", "stopped", "LLM", "2 in flight", "1 waiting", "1m35s"} {
if !strings.Contains(frame, want) {
t.Errorf("frame missing %q:\n%s", want, frame)
}
}
snap.Lock = statusLock{State: "external", External: "cyberpunk2077.exe (pid 1234)", SinceS: 3, ImageQueue: 2}
frame = renderMonitor(snap, 40) // narrow: falls back to 80
for _, want := range []string{"EXTERNAL", "cyberpunk2077.exe", "2 image job(s) waiting"} {
if !strings.Contains(frame, want) {
t.Errorf("frame missing %q:\n%s", want, frame)
}
}
snap.GPU = statusGPU{Enabled: true, Known: true, UsedMB: 4300, TotalMB: 16384, Foreign: "cyberpunk2077.exe (pid 1234)", AgeS: 12}
frame = renderMonitor(snap, 80)
for _, want := range []string{"GPU:", "4.2 GiB / 16.0 GiB used", "external:", "checked 12s ago"} {
if !strings.Contains(frame, want) {
t.Errorf("frame missing %q:\n%s", want, frame)
}
}
snap.GPU = statusGPU{Enabled: true, AgeS: 4}
frame = renderMonitor(snap, 80)
for _, want := range []string{"VRAM unknown", "no external process"} {
if !strings.Contains(frame, want) {
t.Errorf("frame missing %q:\n%s", want, frame)
}
}
}
func TestFmtDur(t *testing.T) {
cases := map[int64]string{0: "0s", 5: "5s", 95: "1m35s", 3725: "1h02m", -3: "0s"}
for in, want := range cases {
if got := fmtDur(in); got != want {
t.Errorf("fmtDur(%d) = %q, want %q", in, got, want)
}
}
}
func TestDiffConfig(t *testing.T) {
a := config.Defaults()
b := a
if got := diffConfig(a, b); len(got) != 0 {
t.Fatalf("identical configs: got %v", got)
}
b.LogLevel = -4
b.GameProcs = []string{"game.exe"}
got := diffConfig(a, b)
if len(got) != 2 || got[0] != "GAME_PROCS" || got[1] != "LOGLEVEL" {
t.Fatalf("got %v, want [GAME_PROCS LOGLEVEL]", got)
}
}
// Every Config field must carry an env tag so user-facing output (the
// reload diff) can name the setting the user would actually change.
func TestConfigFieldsHaveEnvTags(t *testing.T) {
typ := reflect.TypeOf(config.Config{})
for i := 0; i < typ.NumField(); i++ {
f := typ.Field(i)
if f.Tag.Get("env") == "" {
t.Errorf("config.Config.%s has no env tag", f.Name)
}
}
}
+39
View File
@@ -0,0 +1,39 @@
//go:build !windows
package main
import (
"os"
"golang.org/x/sys/unix"
)
// enableVirtualTerminal is a no-op: unix terminals speak ANSI natively.
func enableVirtualTerminal() {}
// termWidth is the terminal width in columns, 80 when unknown.
func termWidth() int {
ws, err := unix.IoctlGetWinsize(int(os.Stdout.Fd()), unix.TIOCGWINSZ)
if err != nil || ws.Col == 0 {
return 80
}
return int(ws.Col)
}
// enableRawKeys switches the terminal to per-keypress mode (ICANON and ECHO
// off) and returns the restore function, nil when stdin is not a terminal.
func enableRawKeys() func() {
fd := int(os.Stdin.Fd())
term, err := unix.IoctlGetTermios(fd, unix.TCGETS)
if err != nil {
return nil
}
raw := *term
raw.Lflag &^= unix.ICANON | unix.ECHO
raw.Cc[unix.VMIN] = 1
raw.Cc[unix.VTIME] = 0
if err := unix.IoctlSetTermios(fd, unix.TCSETS, &raw); err != nil {
return nil
}
return func() { unix.IoctlSetTermios(fd, unix.TCSETS, term) } //nolint:errcheck
}
+49
View File
@@ -0,0 +1,49 @@
//go:build windows
package main
import (
"os"
"golang.org/x/sys/windows"
)
// enableVirtualTerminal asks the console to honor ANSI escapes (Windows 10+);
// mintty/Git Bash already does, so errors are ignored.
func enableVirtualTerminal() {
h := windows.Handle(os.Stdout.Fd())
var mode uint32
if err := windows.GetConsoleMode(h, &mode); err != nil {
return
}
windows.SetConsoleMode(h, mode|windows.ENABLE_VIRTUAL_TERMINAL_PROCESSING) //nolint:errcheck
}
// termWidth is the console window width in columns, 80 when unknown.
func termWidth() int {
var info windows.ConsoleScreenBufferInfo
if err := windows.GetConsoleScreenBufferInfo(windows.Handle(os.Stdout.Fd()), &info); err != nil {
return 80
}
return int(info.Window.Right-info.Window.Left) + 1
}
// enableRawKeys puts the console's stdin into per-keypress mode (no line
// buffering, no echo) and returns the restore function. When stdin is not a
// real console (mintty/Git Bash pipes) it returns nil: ptys already deliver
// keystrokes immediately.
func enableRawKeys() func() {
h := windows.Handle(os.Stdin.Fd())
var mode uint32
if err := windows.GetConsoleMode(h, &mode); err != nil {
return nil
}
const (
enableLineInput = 0x0002
enableEchoInput = 0x0004
)
if err := windows.SetConsoleMode(h, mode&^(enableLineInput|enableEchoInput)); err != nil {
return nil
}
return func() { windows.SetConsoleMode(h, mode) } //nolint:errcheck
}
+1 -1
View File
@@ -6,7 +6,7 @@
# ComfyUI --listen 0.0.0.0 --port 8189). # ComfyUI --listen 0.0.0.0 --port 8189).
services: services:
gpu-turnstile: gpu-turnstile:
image: git.rambossek.at/public/gpu-turnstile:v0.1.9 image: git.rambossek.at/public/gpu-turnstile:v0.2.9
restart: unless-stopped restart: unless-stopped
environment: environment:
# Each consumer is enabled by setting its URL; leave one unset to # Each consumer is enabled by setting its URL; leave one unset to
+67 -48
View File
@@ -13,69 +13,77 @@ import (
"time" "time"
) )
// Config holds every gpu-turnstile setting. // Config holds every gpu-turnstile setting. Each field's env tag names the
// environment variable / config-file key that sets it; user-facing output
// (e.g. the reload diff) uses those names, never the Go field names.
type Config struct { type Config struct {
ListenOllama string ListenOllama string `env:"LISTEN_OLLAMA"`
ListenComfy string ListenComfy string `env:"LISTEN_COMFY"`
OllamaURL string OllamaURL string `env:"OLLAMA_URL"`
ComfyURL string ComfyURL string `env:"COMFY_URL"`
UnloadTimeout time.Duration UnloadTimeout time.Duration `env:"UNLOAD_TIMEOUT"`
JobTimeout time.Duration JobTimeout time.Duration `env:"JOB_TIMEOUT"`
LLMWaitTimeout time.Duration LLMWaitTimeout time.Duration `env:"LLM_WAIT_TIMEOUT"`
UnloadPollInterval time.Duration UnloadPollInterval time.Duration `env:"UNLOAD_POLL_INTERVAL"`
HistoryPollInterval time.Duration HistoryPollInterval time.Duration `env:"HISTORY_POLL_INTERVAL"`
ProbeTimeout time.Duration ProbeTimeout time.Duration `env:"PROBE_TIMEOUT"`
HealthInterval time.Duration HealthInterval time.Duration `env:"HEALTH_INTERVAL"`
FreeTimeout time.Duration FreeTimeout time.Duration `env:"FREE_TIMEOUT"`
WarmTimeout time.Duration WarmTimeout time.Duration `env:"WARM_TIMEOUT"`
ShutdownTimeout time.Duration ShutdownTimeout time.Duration `env:"SHUTDOWN_TIMEOUT"`
BackoffInitial time.Duration BackoffInitial time.Duration `env:"BACKOFF_INITIAL"`
BackoffMax time.Duration BackoffMax time.Duration `env:"BACKOFF_MAX"`
PromptCaptureLimit int64 PromptCaptureLimit int64 `env:"PROMPT_CAPTURE_LIMIT"`
AutoUpdate bool AutoUpdate bool `env:"AUTO_UPDATE"`
UpdateInterval time.Duration UpdateInterval time.Duration `env:"UPDATE_INTERVAL"`
UpdateRepo string UpdateRepo string `env:"UPDATE_REPO"`
UpdateAsset string UpdateAsset string `env:"UPDATE_ASSET"`
// AppVersion is the version the user wants to run: "dev" disables // AppVersion is the version the user wants to run: "dev" disables
// updates, "stable" tracks the latest release, anything else is an // updates, "stable" tracks the latest release, anything else is an
// exact vX.Y.Z release to pin. From APP_VER; defaults to "stable". // exact vX.Y.Z release to pin. From APP_VER; defaults to "stable".
AppVersion string AppVersion string `env:"APP_VER"`
// LLMBusyMode is "wait" (hold requests until the lock is free or // LLMBusyMode is "wait" (hold requests until the lock is free or
// LLMWaitTimeout expires) or "reject" (immediately answer with // LLMWaitTimeout expires) or "reject" (immediately answer with
// LLMBusyStatus + Retry-After when an image job is active or pending). // LLMBusyStatus + Retry-After when an image job is active or pending).
LLMBusyMode string LLMBusyMode string `env:"LLM_BUSY_MODE"`
LLMBusyStatus int LLMBusyStatus int `env:"LLM_BUSY_STATUS"`
BusyRetryAfter int BusyRetryAfter int `env:"BUSY_RETRY_AFTER"`
// ComfyCmd spawns and supervises a ComfyUI server on demand (empty = // ComfyCmd spawns and supervises a ComfyUI server on demand. When
// unmanaged, the current behavior). ComfyDir is its working directory. // ComfyCmd is empty but ComfyDir is set, management is enabled with the
// The managed server is stopped after ComfyIdleTimeout without // standard venv layout under ComfyDir (.venv + main.py or
// requests, freeing its VRAM; ComfyStartTimeout bounds how long a // ComfyUI/main.py; --port from the COMFY_URL port) — ComfyCmd is the
// request waits for it to come up. // override for other layouts and doubles as the working directory when
ComfyCmd string // set explicitly. The managed server is stopped after ComfyIdleTimeout
ComfyDir string // without requests, freeing its VRAM; ComfyStartTimeout bounds how long
ComfyIdleTimeout time.Duration // a request waits for it to come up.
ComfyStartTimeout time.Duration ComfyCmd string `env:"COMFY_CMD"`
ComfyDir string `env:"COMFY_DIR"`
ComfyIdleTimeout time.Duration `env:"COMFY_IDLE_TIMEOUT"`
ComfyStartTimeout time.Duration `env:"COMFY_START_TIMEOUT"`
// GameProcs (GAME_PROCS) is a watch list of process names; while any of // 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 // them runs, the GPU is treated as held by a foreign process. The
// nvidia-smi path (GPUForeignVRAMMB, GPU_FOREIGN_VRAM_MB) does the same // nvidia-smi path (GPUForeignVRAMMB, GPU_FOREIGN_VRAM_MB) does the same
// when a process not in GPUIgnoreProcs (GPU_IGNORE_PROCS) holds more than // when a process not in GPUIgnoreProcs (GPU_IGNORE_PROCS) holds more than
// that many MiB of VRAM. GamePollInterval (GAME_POLL_INTERVAL) is how // that many MiB of VRAM, and the PDH path (GPUForeignUtilPct,
// often both checks run. // GPU_FOREIGN_UTIL_PCT) when such a process uses more than that many
GameProcs []string // percent of the GPU 3D engine (Windows only). GamePollInterval
GPUForeignVRAMMB int // (GAME_POLL_INTERVAL) is how often all checks run.
GPUIgnoreProcs []string GameProcs []string `env:"GAME_PROCS"`
GamePollInterval time.Duration GPUForeignVRAMMB int `env:"GPU_FOREIGN_VRAM_MB"`
GPUForeignUtilPct int `env:"GPU_FOREIGN_UTIL_PCT"`
GPUIgnoreProcs []string `env:"GPU_IGNORE_PROCS"`
GamePollInterval time.Duration `env:"GAME_POLL_INTERVAL"`
WarmModel string WarmModel string `env:"WARM_MODEL"`
LogLevel slog.Level LogLevel slog.Level `env:"LOGLEVEL"`
LogJSON bool LogJSON bool `env:"LOG_FORMAT"`
LogFile string LogFile string `env:"LOG_FILE"`
} }
// Defaults returns the configuration used when neither the environment nor // Defaults returns the configuration used when neither the environment nor
@@ -114,9 +122,10 @@ func Defaults() Config {
ComfyStartTimeout: 2 * time.Minute, ComfyStartTimeout: 2 * time.Minute,
// ComfyUI runs under python; excluding it (and Ollama) by name keeps // ComfyUI runs under python; excluding it (and Ollama) by name keeps
// our own consumers from tripping the foreign-VRAM check. // our own consumers from tripping the foreign-VRAM check. dwm is the
GPUIgnoreProcs: []string{"ollama", "ollama app", "ollama_llama_server", "python", "pythonw"}, // desktop compositor — it always shows some 3D-engine usage.
GamePollInterval: 5 * time.Second, GPUIgnoreProcs: []string{"ollama", "ollama app", "ollama_llama_server", "python", "pythonw", "dwm"},
GamePollInterval: 15 * time.Second,
LogLevel: slog.LevelWarn, LogLevel: slog.LevelWarn,
} }
@@ -238,6 +247,13 @@ func Load(getenv func(string) string) (Config, error) {
} }
cfg.GPUForeignVRAMMB = n cfg.GPUForeignVRAMMB = n
} }
if v := getenv("GPU_FOREIGN_UTIL_PCT"); v != "" {
n, err := strconv.Atoi(v)
if err != nil || n < 0 || n > 100 {
return cfg, fmt.Errorf("GPU_FOREIGN_UTIL_PCT: must be an integer in 0-100 (percent of the GPU 3D engine, 0 = disabled)")
}
cfg.GPUForeignUtilPct = n
}
if v := getenv("PROMPT_CAPTURE_LIMIT"); v != "" { if v := getenv("PROMPT_CAPTURE_LIMIT"); v != "" {
n, err := strconv.ParseInt(v, 10, 64) n, err := strconv.ParseInt(v, 10, 64)
if err != nil || n < 0 { if err != nil || n < 0 {
@@ -305,6 +321,9 @@ func Load(getenv func(string) string) (Config, error) {
if cfg.ComfyCmd != "" && cfg.ComfyURL == "" { if cfg.ComfyCmd != "" && cfg.ComfyURL == "" {
return cfg, fmt.Errorf("COMFY_CMD requires COMFY_URL to be set (the proxy needs somewhere to forward)") 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 == "" { if cfg.OllamaURL == "" && cfg.ComfyURL == "" {
return cfg, ErrNoConsumer return cfg, ErrNoConsumer
} }
+27 -2
View File
@@ -201,8 +201,8 @@ func TestGameDetectionSettings(t *testing.T) {
if len(cfg.GPUIgnoreProcs) != 2 || cfg.GPUIgnoreProcs[1] != "my-trainer" { if len(cfg.GPUIgnoreProcs) != 2 || cfg.GPUIgnoreProcs[1] != "my-trainer" {
t.Fatalf("GPUIgnoreProcs = %v", cfg.GPUIgnoreProcs) t.Fatalf("GPUIgnoreProcs = %v", cfg.GPUIgnoreProcs)
} }
if cfg.GamePollInterval != 5*time.Second { if cfg.GamePollInterval != 15*time.Second {
t.Fatalf("GamePollInterval = %v, want 5s default", cfg.GamePollInterval) t.Fatalf("GamePollInterval = %v, want 15s default", cfg.GamePollInterval)
} }
// Defaults: both detection paths off, ignore list covers our consumers. // Defaults: both detection paths off, ignore list covers our consumers.
@@ -222,3 +222,28 @@ func TestGameDetectionSettings(t *testing.T) {
t.Fatal("GPUIgnoreProcs default must not be empty") 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)
}
}
+5 -4
View File
@@ -28,14 +28,15 @@ func sampleEntries(logFile string) []sampleEntry {
{"OLLAMA_URL", "http://127.0.0.1:11434", "Ollama upstream URL; setting it enables the Ollama consumer (default: empty = disabled)", 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}, {"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}, {"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: empty = unmanaged)", 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`, "Working directory for COMFY_CMD (default: empty = inherit)", 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_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}, {"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}, {"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_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}, {"GPU_FOREIGN_UTIL_PCT", "30", "Also treat the GPU as held when a process not in GPU_IGNORE_PROCS uses more than this percent of the GPU 3D engine (Windows Task-Manager counters; catches games without an exe list; 0/empty = disabled)", false},
{"GAME_POLL_INTERVAL", "5s", "How often game/VRAM detection runs", false}, {"GPU_IGNORE_PROCS", "ollama,ollama app,ollama_llama_server,python,pythonw,dwm", "Process names never counted as foreign GPU users (ComfyUI runs under python; dwm is the desktop compositor)", 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}, {"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}, {"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_WAIT_TIMEOUT", "10m", "Max time an LLM request waits for the GPU before being answered 503 (wait mode)", false},
+117
View File
@@ -0,0 +1,117 @@
// Package control exposes a local-only command channel into a running
// gpu-turnstile service: a named pipe on Windows, a unix socket on Linux.
// It lets unprivileged local users ask the service to do privileged work
// that is safe to offer — currently triggering an update check, whose
// payload is signature-verified regardless of who asks. The channel never
// accepts data beyond a one-word command, and the server rate-limits
// triggers, so the worst a local user can cause is a cheap, throttled
// check and a GPU-idle-gated restart onto a signed binary.
//
// Abuse hardening: the command read is capped (4 KiB), each connection is
// force-closed after connTimeout so a stalled client cannot pin a goroutine
// (or a Windows pipe instance) forever, and concurrently served connections
// are capped at maxConns — beyond that, connections are closed on arrival.
// On Windows the pipe's ACL additionally denies network logons, so the
// channel cannot be reached from another machine.
//
// Protocol: the client writes one command line, the server answers with
// one reply line ("OK ..." or "ERR ...") and hangs up.
package control
import (
"bufio"
"errors"
"fmt"
"io"
"strings"
"time"
)
// CmdUpdateNow asks the service to check for, stage and (once the GPU is
// idle) restart onto a signed update immediately.
const CmdUpdateNow = "update-now"
// CmdStatus asks for a one-line JSON status snapshot (monitor mode).
const CmdStatus = "status"
// CmdReloadEnv asks the service to re-read and validate its config file,
// and to restart onto it (once the GPU is idle) when it changed.
const CmdReloadEnv = "reload-env"
// ErrUnavailable means no running service offers the control channel.
var ErrUnavailable = errors.New("control channel unavailable")
// Handler answers one command; the returned string is sent back as one
// line. It must start with "OK " or "ERR ".
type Handler func(cmd string) string
// connTimeout bounds one connection's lifetime: a client that stops
// mid-command or never reads the reply would otherwise pin its goroutine
// (and on Windows one of the pipe instances) indefinitely. A var so tests
// can shrink it.
var connTimeout = 10 * time.Second
// maxConns caps concurrently served connections; beyond it, new
// connections are closed on arrival. Bound on the goroutines a local
// flood can pile up.
const maxConns = 32
var connSem = make(chan struct{}, maxConns)
// serve dispatches connection handling under the concurrency cap. It
// returns false when the cap is reached — the caller must then close the
// connection itself.
func serve(c io.ReadWriteCloser, h Handler) bool {
select {
case connSem <- struct{}{}:
go func() {
defer func() { <-connSem }()
serveConn(c, h)
}()
return true
default:
return false
}
}
// forceCloser is implemented by connections that can be torn down
// abortively, unblocking pending reads and writes (Windows pipe:
// DisconnectNamedPipe; unix socket: a deadline in the past). The
// connection watchdog uses it; normal closes still flush the reply.
type forceCloser interface {
ForceClose() error
}
// serveConn runs the line protocol on one accepted connection.
func serveConn(c io.ReadWriteCloser, h Handler) {
defer c.Close()
if fc, ok := c.(forceCloser); ok {
timer := time.AfterFunc(connTimeout, func() { fc.ForceClose() })
defer timer.Stop()
}
line, err := bufio.NewReader(io.LimitReader(c, 4096)).ReadString('\n')
cmd := strings.TrimSpace(line)
if cmd == "" {
if err != nil {
return
}
fmt.Fprintln(c, "ERR empty command")
return
}
fmt.Fprintln(c, h(cmd))
}
// readReply writes cmd and reads the server's one-line reply.
func readReply(c io.ReadWriteCloser, cmd string) (string, error) {
if _, err := fmt.Fprintln(c, cmd); err != nil {
return "", err
}
// The server hangs up after its reply; a broken-pipe error after the
// last byte still leaves the reply in the buffer.
data, _ := io.ReadAll(io.LimitReader(c, 4096))
line := strings.TrimSpace(string(data))
if line == "" {
return "", ErrUnavailable
}
return line, nil
}
+66
View File
@@ -0,0 +1,66 @@
//go:build linux
package control
import (
"context"
"log/slog"
"net"
"os"
"time"
)
// sockPath lives in the unit's RuntimeDirectory; mode 0666 lets every
// local user ask, nothing can reach it from off the machine.
const sockPath = "/run/gpu-turnstile/control.sock"
// Serve starts the socket listener in the background and returns; only a
// setup failure is reported. Each client connection is answered in its own
// goroutine.
func Serve(ctx context.Context, h Handler, log *slog.Logger) error {
os.Remove(sockPath) // stale socket from a previous run
ln, err := net.Listen("unix", sockPath)
if err != nil {
return err
}
if err := os.Chmod(sockPath, 0o666); err != nil {
ln.Close()
return err
}
go func() {
<-ctx.Done()
ln.Close()
}()
go func() {
for {
c, err := ln.Accept()
if err != nil {
return // shutting down
}
uc := unixConn{c}
if !serve(uc, h) {
uc.ForceClose()
}
}
}()
return nil
}
// unixConn adds an abortive ForceClose to net.Conn: a deadline in the
// past fails pending and future I/O immediately.
type unixConn struct{ net.Conn }
func (c unixConn) ForceClose() error {
c.SetDeadline(time.Now().Add(-time.Second)) //nolint:errcheck // best effort
return c.Conn.Close()
}
// Ask sends one command to the running service and returns its reply.
func Ask(cmd string) (string, error) {
c, err := net.DialTimeout("unix", sockPath, 2*time.Second)
if err != nil {
return "", ErrUnavailable
}
defer c.Close()
return readReply(c, cmd)
}
+18
View File
@@ -0,0 +1,18 @@
//go:build !windows && !linux
package control
import (
"context"
"log/slog"
)
// Serve is a no-op on platforms without a control channel.
func Serve(ctx context.Context, h Handler, log *slog.Logger) error {
return ErrUnavailable
}
// Ask always reports the channel as unavailable.
func Ask(cmd string) (string, error) {
return "", ErrUnavailable
}
+103
View File
@@ -0,0 +1,103 @@
package control
import (
"errors"
"net"
"strings"
"testing"
"time"
)
func TestRoundTrip(t *testing.T) {
server, client := net.Pipe()
go serveConn(server, func(cmd string) string {
if cmd != CmdUpdateNow {
return "ERR unknown command: " + cmd
}
return "OK v0.2.2 is up to date"
})
reply, err := readReply(client, CmdUpdateNow)
if err != nil {
t.Fatal(err)
}
if reply != "OK v0.2.2 is up to date" {
t.Fatalf("reply = %q", reply)
}
server2, client2 := net.Pipe()
go serveConn(server2, func(cmd string) string { return "ERR unknown command: " + cmd })
reply, err = readReply(client2, "bogus")
if err != nil {
t.Fatal(err)
}
if !strings.HasPrefix(reply, "ERR ") {
t.Fatalf("reply = %q, want ERR prefix", reply)
}
}
func TestEmptyReplyIsUnavailable(t *testing.T) {
server, client := net.Pipe()
go serveConn(server, func(cmd string) string {
server.Close() // hang up without answering
return ""
})
if _, err := readReply(client, CmdUpdateNow); !errors.Is(err, ErrUnavailable) {
t.Fatalf("err = %v, want ErrUnavailable", err)
}
}
func TestServeCap(t *testing.T) {
for i := 0; i < maxConns; i++ {
connSem <- struct{}{}
}
defer func() {
for i := 0; i < maxConns; i++ {
<-connSem
}
}()
server, client := net.Pipe()
defer server.Close()
defer client.Close()
if serve(server, func(string) string { return "OK" }) {
t.Fatal("serve accepted a connection beyond the cap")
}
}
// forcePipe records ForceClose calls for the watchdog test.
type forcePipe struct {
net.Conn
forced chan struct{}
}
func (c forcePipe) ForceClose() error {
err := c.Conn.Close()
close(c.forced)
return err
}
func TestConnWatchdog(t *testing.T) {
old := connTimeout
connTimeout = 50 * time.Millisecond
defer func() { connTimeout = old }()
server, client := net.Pipe()
defer client.Close()
fc := forcePipe{Conn: server, forced: make(chan struct{})}
done := make(chan struct{})
go func() {
serveConn(fc, func(string) string { return "OK" })
close(done)
}()
// The client never sends anything; the watchdog must tear the
// connection down instead of blocking forever.
select {
case <-fc.forced:
case <-time.After(5 * time.Second):
t.Fatal("watchdog did not force-close the stalled connection")
}
select {
case <-done:
case <-time.After(5 * time.Second):
t.Fatal("serveConn still blocked after the force close")
}
}
+167
View File
@@ -0,0 +1,167 @@
//go:build windows
package control
import (
"context"
"fmt"
"log/slog"
"os"
"syscall"
"time"
"unsafe"
"golang.org/x/sys/windows"
)
// pipePath is a kernel-local named pipe: no TCP, no firewall prompt.
const pipePath = `\\.\pipe\gpu-turnstile`
// sddlPipe grants full access to Administrators, SYSTEM and the pipe owner,
// and read+write to authenticated users — except network logons, so the
// pipe cannot be reached from another machine over SMB.
const sddlPipe = "D:(D;;GRGW;;;NU)(A;;GA;;;BA)(A;;GA;;;SY)(A;;GA;;;OW)(A;;GRGW;;;AU)"
var (
procConvertSDDL = windows.NewLazySystemDLL("advapi32.dll").
NewProc("ConvertStringSecurityDescriptorToSecurityDescriptorW")
procWaitNamedPipe = windows.NewLazySystemDLL("kernel32.dll").
NewProc("WaitNamedPipeW")
)
func waitNamedPipe(name *uint16, timeout uint32) error {
r, _, err := procWaitNamedPipe.Call(uintptr(unsafe.Pointer(name)), uintptr(timeout))
if r == 0 {
return err
}
return nil
}
func securityAttributesFromSDDL(sddl string) (*windows.SecurityAttributes, error) {
s, err := windows.UTF16PtrFromString(sddl)
if err != nil {
return nil, err
}
var sd *uint16 // SECURITY_DESCRIPTOR*, kept for the process lifetime
r, _, callErr := procConvertSDDL.Call(
uintptr(unsafe.Pointer(s)), 1, /* SDDL_REVISION_1 */
uintptr(unsafe.Pointer(&sd)), 0)
if r == 0 {
return nil, fmt.Errorf("invalid SDDL: %w", callErr)
}
sa := &windows.SecurityAttributes{
Length: uint32(unsafe.Sizeof(windows.SecurityAttributes{})),
SecurityDescriptor: (*windows.SECURITY_DESCRIPTOR)(unsafe.Pointer(sd)),
}
return sa, nil
}
// Serve starts the pipe listener in the background and returns; only a
// setup failure is reported. Each client connection is answered in its own
// goroutine. On shutdown the process exit reaps everything.
func Serve(ctx context.Context, h Handler, log *slog.Logger) error {
sa, err := securityAttributesFromSDDL(sddlPipe)
if err != nil {
return err
}
name, err := windows.UTF16PtrFromString(pipePath)
if err != nil {
return err
}
go func() {
for ctx.Err() == nil {
pipe, err := windows.CreateNamedPipe(name,
windows.PIPE_ACCESS_DUPLEX,
windows.PIPE_TYPE_BYTE|windows.PIPE_READMODE_BYTE|windows.PIPE_WAIT,
16, 4096, 4096, 0, sa)
if err != nil {
log.Warn("control channel stopped", "err", err)
return
}
// Blocks until a client connects — only then is the next
// instance created, so instances are not burned without
// clients. ERROR_PIPE_CONNECTED means the client raced us
// and connected before the call: that is a success. Process
// exit reaps the blocked call on shutdown.
if err := windows.ConnectNamedPipe(pipe, nil); err != nil && err != errnoPipeConnected {
windows.CloseHandle(pipe)
continue
}
conn := &pipeConn{f: os.NewFile(uintptr(pipe), pipePath), h: pipe}
if !serve(conn, h) {
conn.ForceClose()
}
}
}()
return nil
}
// errnoPipeConnected is ConnectNamedPipe's "the client connected before we
// called" result, which means the connection is established.
var errnoPipeConnected = syscall.Errno(535) // ERROR_PIPE_CONNECTED
// errnoPipeBusy is CreateFile's "all pipe instances are busy" result. It
// loses the race against another client that grabbed the instance
// WaitNamedPipe just reported — the caller must wait and retry.
var errnoPipeBusy = syscall.Errno(231) // ERROR_PIPE_BUSY
// pipeConn adapts a pipe handle to io.ReadWriteCloser. Close flushes first
// (FlushFileBuffers blocks until the client has read the reply) and then
// disconnects — closing the bare handle right after writing can discard
// unread reply bytes, which clients see as an empty, failed request.
type pipeConn struct {
f *os.File
h windows.Handle
}
func (c *pipeConn) Read(p []byte) (int, error) { return c.f.Read(p) }
func (c *pipeConn) Write(p []byte) (int, error) { return c.f.Write(p) }
func (c *pipeConn) Close() error {
windows.FlushFileBuffers(c.h) //nolint:errcheck // best effort
windows.DisconnectNamedPipe(c.h) //nolint:errcheck // best effort
return c.f.Close()
}
// ForceClose aborts the connection without flushing: disconnecting
// unblocks pending reads and writes at the cost of possibly discarding an
// unread reply. Used by the connection watchdog; normal closes flush.
func (c *pipeConn) ForceClose() error {
windows.DisconnectNamedPipe(c.h) //nolint:errcheck // best effort
return c.f.Close()
}
// Ask sends one command to the running service and returns its reply.
//
// The server keeps exactly one listening instance per connection, so
// concurrent clients race for it: WaitNamedPipe can report an instance
// that another client grabs before our CreateFile runs (ERROR_PIPE_BUSY).
// Retry on that — with the monitor polling status every second, a single
// attempt loses that race regularly.
func Ask(cmd string) (string, error) {
name, err := windows.UTF16PtrFromString(pipePath)
if err != nil {
return "", err
}
deadline := time.Now().Add(5 * time.Second)
for {
if err := waitNamedPipe(name, 2000); err != nil {
return "", ErrUnavailable
}
handle, err := windows.CreateFile(name,
windows.GENERIC_READ|windows.GENERIC_WRITE, 0, nil,
windows.OPEN_EXISTING, 0, 0)
if err == errnoPipeBusy {
if time.Now().After(deadline) {
return "", ErrUnavailable
}
continue
}
if err != nil {
return "", ErrUnavailable
}
f := os.NewFile(uintptr(handle), pipePath)
defer f.Close()
return readReply(f, cmd)
}
}
+107 -30
View File
@@ -1,8 +1,10 @@
// Package game detects processes outside gpu-turnstile's control that hold // 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 // 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 // VRAM while they run. Three detection paths: an explicit process watch
// (GAME_PROCS) and a foreign-VRAM threshold via nvidia-smi // list (GAME_PROCS), a foreign-VRAM threshold via nvidia-smi
// (GPU_FOREIGN_VRAM_MB) that catches anything not on the ignore list. // (GPU_FOREIGN_VRAM_MB), and a per-process GPU 3D-engine utilization
// threshold via Windows PDH counters (GPU_FOREIGN_UTIL_PCT). The latter two
// catch anything not on the ignore list without naming individual games.
package game package game
import ( import (
@@ -11,6 +13,7 @@ import (
"fmt" "fmt"
"log/slog" "log/slog"
"os/exec" "os/exec"
"slices"
"strconv" "strconv"
"strings" "strings"
) )
@@ -28,28 +31,33 @@ type computeApp struct {
} }
// Detector checks whether a foreign process holds the GPU. The zero value // Detector checks whether a foreign process holds the GPU. The zero value
// (no watch list, no threshold) never detects anything; main only starts the // (no watch list, no thresholds) never detects anything; main only starts
// poll loop when at least one path is configured. // the poll loop when at least one path is configured.
type Detector struct { type Detector struct {
procs map[string]bool // normalized names from GAME_PROCS procs map[string]bool // normalized names from GAME_PROCS
vramMB int // foreign VRAM threshold; 0 = disabled vramMB int // foreign VRAM threshold; 0 = disabled
utilPct int // foreign 3D-engine utilization threshold; 0 = disabled
ignore map[string]bool // normalized names never counted as foreign ignore map[string]bool // normalized names never counted as foreign
log *slog.Logger log *slog.Logger
noNvidia bool // nvidia-smi was not found; VRAM path disabled for good noNvidia bool // nvidia-smi was not found; VRAM path disabled for good
sampler *gpuEngineSampler // open PDH query, opened lazily on first Check
noPDH bool // engine counters unavailable; util path disabled for good
} }
// New builds a Detector from the configured watch list, VRAM threshold in // 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 // MiB, 3D-engine utilization threshold in percent (both 0 = disabled) and
// case-insensitively, with or without a trailing ".exe". // ignore list. Names are matched case-insensitively, with or without a
func New(procs []string, vramMB int, ignore []string, log *slog.Logger) *Detector { // trailing ".exe".
func New(procs []string, vramMB, utilPct int, ignore []string, log *slog.Logger) *Detector {
if log == nil { if log == nil {
log = slog.Default() log = slog.Default()
} }
return &Detector{ return &Detector{
procs: nameSet(procs), procs: nameSet(procs),
vramMB: vramMB, vramMB: vramMB,
ignore: nameSet(ignore), utilPct: utilPct,
log: log, ignore: nameSet(ignore),
log: log,
} }
} }
@@ -73,34 +81,63 @@ func nameSet(names []string) map[string]bool {
// description of each (empty when the GPU is free for gpu-turnstile's // 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 // 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 // the process list found nothing; a missing nvidia-smi binary disables the
// VRAM path permanently (logged once). // VRAM path permanently (logged once), as do missing engine counters.
func (d *Detector) Check(ctx context.Context) ([]string, error) { func (d *Detector) Check(ctx context.Context) ([]string, error) {
ps, psErr := processes() ps, psErr := processes()
if d.vramMB <= 0 || d.noNvidia { utils := d.engineUtil()
return d.detect(ps, nil), psErr var apps []computeApp
if d.vramMB > 0 && !d.noNvidia {
var err error
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")
} else if err != nil {
return d.detect(ps, nil, utils), err
}
} }
apps, err := queryComputeApps(ctx) return d.detect(ps, apps, utils), psErr
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) // engineUtil samples per-process 3D-engine utilization via PDH. The first
// the nvidia-smi compute-apps list, it returns the foreign holders. // call only primes the rate counters and returns nil. A failing open or
func (d *Detector) detect(ps []Process, apps []computeApp) []string { // sample disables the path permanently (logged once).
func (d *Detector) engineUtil() map[int]float64 {
if d.utilPct <= 0 || d.noPDH {
return nil
}
if d.sampler == nil {
s, err := openGPUEngineSampler()
if err != nil {
d.noPDH = true
d.log.Warn("GPU_FOREIGN_UTIL_PCT is set but per-process GPU counters are unavailable; engine detection disabled", "err", err)
return nil
}
d.sampler = s
}
utils, err := d.sampler.sample()
if err != nil {
if errors.Is(err, errNotPrimed) {
return nil
}
d.noPDH = true
d.log.Warn("per-process GPU counters failed; engine detection disabled", "err", err)
return nil
}
return utils
}
// detect is the pure core of Check: given the process table and
// (optionally) the nvidia-smi compute-apps list and the PDH engine
// utilization, it returns the foreign holders.
func (d *Detector) detect(ps []Process, apps []computeApp, utils map[int]float64) []string {
var holders []string var holders []string
for _, p := range ps { for _, p := range ps {
if d.procs[normName(p.Name)] { if d.procs[normName(p.Name)] {
holders = append(holders, fmt.Sprintf("%s (pid %d)", p.Name, p.PID)) holders = append(holders, fmt.Sprintf("%s (pid %d)", p.Name, p.PID))
} }
} }
if d.vramMB > 0 && apps != nil { if (d.vramMB > 0 && apps != nil) || (d.utilPct > 0 && utils != nil) {
names := make(map[int]string, len(ps)) names := make(map[int]string, len(ps))
for _, p := range ps { for _, p := range ps {
names[p.PID] = p.Name names[p.PID] = p.Name
@@ -115,6 +152,23 @@ func (d *Detector) detect(ps []Process, apps []computeApp) []string {
} }
holders = append(holders, fmt.Sprintf("%s (pid %d) using %d MiB VRAM", name, a.PID, a.UsedMB)) holders = append(holders, fmt.Sprintf("%s (pid %d) using %d MiB VRAM", name, a.PID, a.UsedMB))
} }
// Sorted for stable output (map iteration order is random).
pids := make([]int, 0, len(utils))
for pid := range utils {
pids = append(pids, pid)
}
slices.Sort(pids)
for _, pid := range pids {
util := utils[pid]
name := names[pid]
if d.ignore[normName(name)] || util < float64(d.utilPct) {
continue
}
if name == "" {
name = "unknown process"
}
holders = append(holders, fmt.Sprintf("%s (pid %d) using %.0f%% GPU", name, pid, util))
}
} }
return holders return holders
} }
@@ -132,6 +186,29 @@ func queryComputeApps(ctx context.Context) ([]computeApp, error) {
return parseComputeApps(string(out)) return parseComputeApps(string(out))
} }
// QueryVRAMMB returns used and total GPU VRAM in MiB via nvidia-smi.
// Unlike the per-process list this works under WDDM too.
func QueryVRAMMB(ctx context.Context) (used, total int, err error) {
out, err := exec.CommandContext(ctx, "nvidia-smi",
"--query-gpu=memory.used,memory.total", "--format=csv,noheader,nounits").Output()
if err != nil {
return 0, 0, err
}
usedStr, totalStr, ok := strings.Cut(strings.TrimSpace(string(out)), ",")
if !ok {
return 0, 0, fmt.Errorf("nvidia-smi: unexpected output %q", strings.TrimSpace(string(out)))
}
used, err = strconv.Atoi(strings.TrimSpace(usedStr))
if err != nil {
return 0, 0, fmt.Errorf("nvidia-smi: unexpected used memory in %q", strings.TrimSpace(string(out)))
}
total, err = strconv.Atoi(strings.TrimSpace(totalStr))
if err != nil {
return 0, 0, fmt.Errorf("nvidia-smi: unexpected total memory in %q", strings.TrimSpace(string(out)))
}
return used, total, nil
}
// parseComputeApps parses "pid, used_memory" CSV lines (no header, MiB // parseComputeApps parses "pid, used_memory" CSV lines (no header, MiB
// units). Unsupported rows ("N/A" on WDDM) are skipped. // units). Unsupported rows ("N/A" on WDDM) are skipped.
func parseComputeApps(out string) ([]computeApp, error) { func parseComputeApps(out string) ([]computeApp, error) {
+43 -4
View File
@@ -45,7 +45,7 @@ func TestParseComputeApps(t *testing.T) {
} }
func TestDetect(t *testing.T) { func TestDetect(t *testing.T) {
d := New([]string{"Cyberpunk2077.exe", "hl2"}, 1024, d := New([]string{"Cyberpunk2077.exe", "hl2"}, 1024, 0,
[]string{"ollama", "python", "pythonw"}, nil) []string{"ollama", "python", "pythonw"}, nil)
ps := []Process{ ps := []Process{
{PID: 10, Name: "ollama.exe"}, {PID: 10, Name: "ollama.exe"},
@@ -58,7 +58,7 @@ func TestDetect(t *testing.T) {
{PID: 40, UsedMB: 2048}, // foreign, above threshold {PID: 40, UsedMB: 2048}, // foreign, above threshold
{PID: 50, UsedMB: 100}, // foreign but below threshold {PID: 50, UsedMB: 100}, // foreign but below threshold
} }
holders := d.detect(ps, apps) holders := d.detect(ps, apps, nil)
if len(holders) != 2 { if len(holders) != 2 {
t.Fatalf("got %v, want 2 holders", holders) t.Fatalf("got %v, want 2 holders", holders)
} }
@@ -70,13 +70,52 @@ func TestDetect(t *testing.T) {
} }
} }
func TestDetectEngineUtil(t *testing.T) {
d := New(nil, 0, 30, []string{"dwm", "python"}, nil)
ps := []Process{
{PID: 10, Name: "dwm.exe"},
{PID: 20, Name: "game.exe"},
{PID: 40, Name: "browser.exe"},
}
utils := map[int]float64{
10: 45, // ignored: dwm
20: 61, // foreign, above threshold
30: 82, // foreign, unknown name
40: 5, // below threshold
}
holders := d.detect(ps, nil, utils)
want := []string{
"game.exe (pid 20) using 61% GPU",
"unknown process (pid 30) using 82% GPU",
}
if !slices.Equal(holders, want) {
t.Errorf("got %v, want %v", holders, want)
}
}
func TestDetectNothingConfigured(t *testing.T) { func TestDetectNothingConfigured(t *testing.T) {
d := New(nil, 0, nil, nil) d := New(nil, 0, 0, nil, nil)
if got := d.detect([]Process{{PID: 1, Name: "game.exe"}}, nil); len(got) != 0 { if got := d.detect([]Process{{PID: 1, Name: "game.exe"}}, nil, nil); len(got) != 0 {
t.Errorf("got %v, want none", got) t.Errorf("got %v, want none", got)
} }
} }
func TestParseGPUEngineInstance(t *testing.T) {
pid, eng, ok := parseGPUEngineInstance("pid_1234_luid_0x00000000_0x00011A2B_phys_0_eng_0_engtype_3D")
if !ok || pid != 1234 || eng != "3D" {
t.Errorf("got %d, %q, %v", pid, eng, ok)
}
pid, eng, ok = parseGPUEngineInstance("pid_42_luid_0x0_0x0_phys_0_eng_1_engtype_Copy")
if !ok || pid != 42 || eng != "Copy" {
t.Errorf("got %d, %q, %v", pid, eng, ok)
}
for _, bad := range []string{"", "something", "pid_", "pid_x_luid", "pid_-1_luid_0"} {
if _, _, ok := parseGPUEngineInstance(bad); ok {
t.Errorf("%q parsed, want failure", bad)
}
}
}
func TestProcessesLive(t *testing.T) { func TestProcessesLive(t *testing.T) {
if runtime.GOOS != "windows" && runtime.GOOS != "linux" { if runtime.GOOS != "windows" && runtime.GOOS != "linux" {
t.Skip("no process listing on this platform") t.Skip("no process listing on this platform")
+36
View File
@@ -0,0 +1,36 @@
package game
import (
"errors"
"strconv"
"strings"
)
// errNotPrimed marks the first PDH sample after opening a query: rate-based
// counters (like engine utilization) need two collections before they
// return meaningful values.
var errNotPrimed = errors.New("GPU engine counter needs a second sample")
// parseGPUEngineInstance splits a PDH "GPU Engine" instance name —
// "pid_1234_luid_0x00000000_0x00011A2B_phys_0_eng_0_engtype_3D" — into PID
// and engine type ("3D", "Copy", "VideoDecode", ...). engType is empty when
// the name carries no engtype marker.
func parseGPUEngineInstance(name string) (pid int, engType string, ok bool) {
rest, found := strings.CutPrefix(name, "pid_")
if !found {
return 0, "", false
}
digits, rest, found := strings.Cut(rest, "_")
if !found {
return 0, "", false
}
pid, err := strconv.Atoi(digits)
if err != nil || pid < 0 {
return 0, "", false
}
const marker = "engtype_"
if i := strings.LastIndex(rest, marker); i >= 0 {
engType = rest[i+len(marker):]
}
return pid, engType, true
}
+16
View File
@@ -0,0 +1,16 @@
//go:build !windows
package game
import "errors"
// errNoEngineCounters marks platforms without per-process GPU engine
// counters (the PDH path is Windows-only).
var errNoEngineCounters = errors.New("per-process GPU engine counters are only available on Windows")
// gpuEngineSampler is a stub on non-Windows platforms.
type gpuEngineSampler struct{}
func openGPUEngineSampler() (*gpuEngineSampler, error) { return nil, errNoEngineCounters }
func (s *gpuEngineSampler) sample() (map[int]float64, error) { return nil, errNoEngineCounters }
+111
View File
@@ -0,0 +1,111 @@
//go:build windows
package game
import (
"fmt"
"unsafe"
"golang.org/x/sys/windows"
)
// Per-process GPU engine utilization via PDH — the same counters Task
// Manager's "GPU engine" columns read. Unlike nvidia-smi's compute-apps
// this covers graphics work under WDDM, so games show up. The counter is
// added with PdhAddEnglishCounterW, which is independent of the Windows
// display language.
var (
pdhDLL = windows.NewLazySystemDLL("pdh.dll")
procPdhOpenQuery = pdhDLL.NewProc("PdhOpenQueryW")
procPdhAddEnglishCounter = pdhDLL.NewProc("PdhAddEnglishCounterW")
procPdhCollectQueryData = pdhDLL.NewProc("PdhCollectQueryData")
procPdhGetFormattedCounterArray = pdhDLL.NewProc("PdhGetFormattedCounterArrayW")
procPdhCloseQuery = pdhDLL.NewProc("PdhCloseQuery")
)
const (
pdhFmtDouble = 0x00000200 // PDH_FMT_DOUBLE
pdhMoreData = 0x800007D2 // PDH_MORE_DATA
)
// pdhCountervalueItem mirrors PDH_FMT_COUNTERVALUE_ITEM (64-bit, double).
type pdhCountervalueItem struct {
name *uint16
cStatus uint32
_ uint32 // alignment padding
value float64
}
// gpuEngineSampler holds an open PDH query on the wildcard GPU Engine
// utilization counter. Keeping the query open across polls is what makes
// the rate-based values meaningful; the sampler lives as long as the
// process (PdhCloseQuery would only matter on unload).
type gpuEngineSampler struct {
query uintptr // PDH_HQUERY
counter uintptr // PDH_HCOUNTER
primed bool
}
// openGPUEngineSampler opens a query on the per-process GPU engine
// utilization counter (all instances).
func openGPUEngineSampler() (*gpuEngineSampler, error) {
var q uintptr
if r, _, _ := procPdhOpenQuery.Call(0, 0, uintptr(unsafe.Pointer(&q))); r != 0 {
return nil, fmt.Errorf("PdhOpenQuery: status %#x", r)
}
path, err := windows.UTF16PtrFromString(`\GPU Engine(*)\Utilization Percentage`)
if err != nil {
procPdhCloseQuery.Call(q)
return nil, err
}
var c uintptr
if r, _, _ := procPdhAddEnglishCounter.Call(q, uintptr(unsafe.Pointer(path)), 0, uintptr(unsafe.Pointer(&c))); r != 0 {
procPdhCloseQuery.Call(q)
return nil, fmt.Errorf("PdhAddEnglishCounter: status %#x", r)
}
return &gpuEngineSampler{query: q, counter: c}, nil
}
// sample collects the counter once and returns per-PID 3D-engine
// utilization in percent. The first call after open only primes the rate
// calculation and returns errNotPrimed. Processes can drive several 3D
// engines; their values are summed.
func (s *gpuEngineSampler) sample() (map[int]float64, error) {
if r, _, _ := procPdhCollectQueryData.Call(s.query); r != 0 {
return nil, fmt.Errorf("PdhCollectQueryData: status %#x", r)
}
if !s.primed {
s.primed = true
return nil, errNotPrimed
}
var size, count uint32
r, _, _ := procPdhGetFormattedCounterArray.Call(s.counter, pdhFmtDouble,
uintptr(unsafe.Pointer(&size)), uintptr(unsafe.Pointer(&count)), 0)
if r == pdhMoreData && size == 0 {
return nil, nil // no GPU engine instances at all
}
if r != pdhMoreData {
return nil, fmt.Errorf("PdhGetFormattedCounterArray(size): status %#x", r)
}
buf := make([]byte, size)
r, _, _ = procPdhGetFormattedCounterArray.Call(s.counter, pdhFmtDouble,
uintptr(unsafe.Pointer(&size)), uintptr(unsafe.Pointer(&count)),
uintptr(unsafe.Pointer(&buf[0])))
if r != 0 {
return nil, fmt.Errorf("PdhGetFormattedCounterArray: status %#x", r)
}
items := unsafe.Slice((*pdhCountervalueItem)(unsafe.Pointer(&buf[0])), int(count))
out := make(map[int]float64)
for i := range items {
if items[i].cStatus != 0 || items[i].name == nil {
continue
}
pid, engType, ok := parseGPUEngineInstance(windows.UTF16PtrToString(items[i].name))
if !ok || engType != "3D" {
continue // only the 3D engine marks game-like work
}
out[pid] += items[i].value
}
return out, nil
}
+33
View File
@@ -0,0 +1,33 @@
//go:build windows
package game
import (
"errors"
"testing"
"time"
)
// TestGPUEngineSamplerLive opens the real PDH query and takes two samples;
// the first only primes the rate counters. Skipped (not failed) when the
// machine has no GPU counters.
func TestGPUEngineSamplerLive(t *testing.T) {
s, err := openGPUEngineSampler()
if err != nil {
t.Skipf("no GPU engine counters: %v", err)
}
if _, err := s.sample(); !errors.Is(err, errNotPrimed) {
t.Fatalf("first sample: err = %v, want errNotPrimed", err)
}
time.Sleep(200 * time.Millisecond)
utils, err := s.sample()
if err != nil {
t.Fatalf("second sample: %v", err)
}
for pid, util := range utils {
if pid < 0 || util < 0 {
t.Errorf("pid %d: util %.2f", pid, util)
}
}
t.Logf("%d processes with 3D-engine usage", len(utils))
}
+75 -2
View File
@@ -9,6 +9,7 @@ import (
"context" "context"
"log/slog" "log/slog"
"sync" "sync"
"time"
) )
// State is the current GPU occupancy state. // State is the current GPU occupancy state.
@@ -30,10 +31,13 @@ type Lock struct {
change chan struct{} // closed and replaced on every state change change chan struct{} // closed and replaced on every state change
n int // LLM requests in flight n int // LLM requests in flight
llmWaiting int // LLM requests blocked waiting for the GPU
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 external string // non-empty: a foreign process (e.g. a game) holds the GPU
detail string // what the current holder is doing (best effort)
since time.Time // when the current state began
log *slog.Logger log *slog.Logger
} }
@@ -41,7 +45,7 @@ type Lock struct {
// New returns a ready-to-use Lock. log may be nil; if set, every state // New returns a ready-to-use Lock. log may be nil; if set, every state
// transition is logged at debug level. // transition is logged at debug level.
func New(log *slog.Logger) *Lock { func New(log *slog.Logger) *Lock {
return &Lock{change: make(chan struct{}), log: log} return &Lock{change: make(chan struct{}), log: log, since: time.Now()}
} }
// broadcast wakes all waiters. Call with mu held. // broadcast wakes all waiters. Call with mu held.
@@ -63,6 +67,7 @@ func (l *Lock) logTransition(msg string, args ...any) {
func (l *Lock) SetExternal(holder string) { func (l *Lock) SetExternal(holder string) {
l.mu.Lock() l.mu.Lock()
l.external = holder l.external = holder
l.since = time.Now()
l.broadcast() l.broadcast()
l.mu.Unlock() l.mu.Unlock()
l.logTransition("lock transition", "state", StateExternal, "holder", holder) l.logTransition("lock transition", "state", StateExternal, "holder", holder)
@@ -73,6 +78,7 @@ func (l *Lock) SetExternal(holder string) {
func (l *Lock) ClearExternal() { func (l *Lock) ClearExternal() {
l.mu.Lock() l.mu.Lock()
l.external = "" l.external = ""
l.since = time.Now()
l.broadcast() l.broadcast()
l.mu.Unlock() l.mu.Unlock()
l.logTransition("lock transition", "state", StateIdle) l.logTransition("lock transition", "state", StateIdle)
@@ -90,17 +96,31 @@ func (l *Lock) External() string {
// 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()
waiting := false
for l.imageActive || len(l.imageQ) > 0 || l.external != "" { for l.imageActive || len(l.imageQ) > 0 || l.external != "" {
if !waiting {
l.llmWaiting++
waiting = true
}
ch := l.change ch := l.change
l.mu.Unlock() l.mu.Unlock()
select { select {
case <-ctx.Done(): case <-ctx.Done():
l.mu.Lock()
l.llmWaiting--
l.mu.Unlock()
return ctx.Err() return ctx.Err()
case <-ch: case <-ch:
} }
l.mu.Lock() l.mu.Lock()
} }
if waiting {
l.llmWaiting--
}
l.n++ l.n++
if l.n == 1 {
l.since = time.Now()
}
n := l.n n := l.n
l.mu.Unlock() l.mu.Unlock()
l.logTransition("lock transition", "state", StateLLM, "llm_inflight", n) l.logTransition("lock transition", "state", StateLLM, "llm_inflight", n)
@@ -129,6 +149,8 @@ func (l *Lock) ReleaseLLM() {
l.n-- l.n--
n := l.n n := l.n
if l.n == 0 { if l.n == 0 {
l.since = time.Now()
l.detail = ""
l.broadcast() l.broadcast()
} }
l.mu.Unlock() l.mu.Unlock()
@@ -156,6 +178,7 @@ func (l *Lock) AcquireImage(ctx context.Context) error {
if l.imageQ[0].id == w.id && l.n == 0 && !l.imageActive && l.external == "" { 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.since = time.Now()
l.mu.Unlock() l.mu.Unlock()
l.logTransition("lock transition", "state", StateImage) l.logTransition("lock transition", "state", StateImage)
return nil return nil
@@ -184,6 +207,8 @@ func (l *Lock) AcquireImage(ctx context.Context) error {
func (l *Lock) ReleaseImage() { func (l *Lock) ReleaseImage() {
l.mu.Lock() l.mu.Lock()
l.imageActive = false l.imageActive = false
l.since = time.Now()
l.detail = ""
l.broadcast() l.broadcast()
l.mu.Unlock() l.mu.Unlock()
l.logTransition("lock transition", "state", StateIdle) l.logTransition("lock transition", "state", StateIdle)
@@ -206,3 +231,51 @@ func (l *Lock) Snapshot() (state State, llmInflight int, imagePending bool) {
} }
return state, l.n, l.imageActive || len(l.imageQ) > 0 return state, l.n, l.imageActive || len(l.imageQ) > 0
} }
// SetDetail records what the current holder is doing (e.g. the request
// path), for status displays. Best effort: overwritten by each new holder,
// cleared when the GPU goes idle.
func (l *Lock) SetDetail(detail string) {
l.mu.Lock()
l.detail = detail
l.mu.Unlock()
}
// Status is a point-in-time view of the lock for monitoring.
type Status struct {
State State
Detail string
LLMInflight int
LLMWaiting int
ImageActive bool
ImageQueue int
External string
Since time.Time
}
// Status reports the full lock state, including waiters and how long the
// current state has held.
func (l *Lock) Status() Status {
l.mu.Lock()
defer l.mu.Unlock()
s := Status{
Detail: l.detail,
LLMInflight: l.n,
LLMWaiting: l.llmWaiting,
ImageActive: l.imageActive,
ImageQueue: len(l.imageQ),
External: l.external,
Since: l.since,
}
switch {
case l.imageActive:
s.State = StateImage
case l.n > 0:
s.State = StateLLM
case l.external != "":
s.State = StateExternal
default:
s.State = StateIdle
}
return s
}
+2
View File
@@ -491,6 +491,7 @@ func (s *Server) OllamaHandler() http.Handler {
return return
} }
defer s.cfg.Lock.ReleaseLLM() defer s.cfg.Lock.ReleaseLLM()
s.cfg.Lock.SetDetail("ollama: " + r.Method + " " + r.URL.Path)
s.ollamaProxy.ServeHTTP(w, r) s.ollamaProxy.ServeHTTP(w, r)
})) }))
} }
@@ -590,6 +591,7 @@ func (s *Server) handlePrompt(w http.ResponseWriter, r *http.Request) {
} }
s.cfg.Metrics.ObserveLockWait("image", time.Since(start).Seconds()) s.cfg.Metrics.ObserveLockWait("image", time.Since(start).Seconds())
log.Info("image lock acquired") log.Info("image lock acquired")
s.cfg.Lock.SetDetail("comfy: POST /prompt")
if s.cfg.ComfySup != nil && !comfyFirst { if s.cfg.ComfySup != nil && !comfyFirst {
if err := s.cfg.ComfySup.EnsureRunning(); err != nil { if err := s.cfg.ComfySup.EnsureRunning(); err != nil {
+16
View File
@@ -45,3 +45,19 @@ func writeEnvFile(path, content string) error {
} }
return nil 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]
}
+27 -6
View File
@@ -15,6 +15,8 @@ import (
"os/signal" "os/signal"
"path/filepath" "path/filepath"
"syscall" "syscall"
"gpu-turnstile/internal/supervise"
) )
// Name matches the Windows service name; the systemd unit is Name + ".service". // Name matches the Windows service name; the systemd unit is Name + ".service".
@@ -62,8 +64,23 @@ func Run(run func(ctx context.Context) error) error {
// filesystem is read-only except StateDirectory (the install dir, so // filesystem is read-only except StateDirectory (the install dir, so
// self-updates can rewrite the binary), and the usual no-privilege-escalation // self-updates can rewrite the binary), and the usual no-privilege-escalation
// directives apply. The proxy needs nothing but outbound TCP/UDP and the // directives apply. The proxy needs nothing but outbound TCP/UDP and the
// notify socket, so it loses nothing. // notify socket, so it loses nothing. A managed ComfyUI (comfyDir) gets a
func renderUnit(exePath, configPath string) string { // 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, and a Comfy-Desktop shared data dir (models, input,
// output) a read-write one.
func renderUnit(exePath, configPath, comfyDir string) string {
bind := ""
if comfyDir != "" {
bind = "BindPaths=" + comfyDir + "\n"
if home := comfyVenvHome(comfyDir); home != "" {
bind += "BindReadOnlyPaths=" + home + "\n"
}
if shared := supervise.DesktopSharedDir(comfyDir); shared != "" {
bind += "BindPaths=" + shared + "\n"
}
}
return fmt.Sprintf(`[Unit] return fmt.Sprintf(`[Unit]
Description=gpu-turnstile GPU arbitration proxy for Ollama and ComfyUI Description=gpu-turnstile GPU arbitration proxy for Ollama and ComfyUI
After=network-online.target After=network-online.target
@@ -78,7 +95,9 @@ RestartSec=5s
DynamicUser=yes DynamicUser=yes
StateDirectory=%s StateDirectory=%s
ProtectSystem=strict RuntimeDirectory=%s
RuntimeDirectoryMode=0755
%sProtectSystem=strict
ProtectHome=yes ProtectHome=yes
PrivateTmp=yes PrivateTmp=yes
NoNewPrivileges=yes NoNewPrivileges=yes
@@ -100,7 +119,7 @@ SystemCallErrorNumber=EPERM
[Install] [Install]
WantedBy=multi-user.target WantedBy=multi-user.target
`, exePath, configPath, Name) `, exePath, configPath, Name, Name, bind)
} }
// copyFile copies src to dst, creating dst with the given mode. // copyFile copies src to dst, creating dst with the given mode.
@@ -125,7 +144,9 @@ func copyFile(src, dst string, mode os.FileMode) error {
// sure /etc/gpu-turnstile.env exists (copied from the given config file if // sure /etc/gpu-turnstile.env exists (copied from the given config file if
// provided), writes the hardened unit, then enables and starts it. With // provided), writes the hardened unit, then enables and starts it. With
// copyBin=false the current executable location and config path are // copyBin=false the current executable location and config path are
// registered as-is instead. Needs root. // 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 // Re-running install converges an existing unit instead of failing: it is
// stopped first if active, the installed binary is replaced only when the // stopped first if active, the installed binary is replaced only when the
@@ -186,7 +207,7 @@ func Install(configPath string, copyBin bool, version string) error {
cfg = abs cfg = abs
} }
} }
rendered := renderUnit(exe, cfg) rendered := renderUnit(exe, cfg, configuredValue(cfg, "COMFY_DIR"))
if old, _ := os.ReadFile(unitPath); string(old) != rendered { if old, _ := os.ReadFile(unitPath); string(old) != rendered {
if err := os.WriteFile(unitPath, []byte(rendered), 0o644); err != nil { if err := os.WriteFile(unitPath, []byte(rendered), 0o644); err != nil {
return fmt.Errorf("write %s (run as root): %w", unitPath, err) return fmt.Errorf("write %s (run as root): %w", unitPath, err)
+10 -1
View File
@@ -8,7 +8,7 @@ import (
) )
func TestRenderUnit(t *testing.T) { func TestRenderUnit(t *testing.T) {
unit := renderUnit("/var/lib/gpu-turnstile/gpu-turnstile", "/etc/gpu-turnstile.env") unit := renderUnit("/var/lib/gpu-turnstile/gpu-turnstile", "/etc/gpu-turnstile.env", "")
for _, want := range []string{ for _, want := range []string{
"Type=notify", "Type=notify",
"WatchdogSec=30s", "WatchdogSec=30s",
@@ -17,6 +17,7 @@ func TestRenderUnit(t *testing.T) {
"WantedBy=multi-user.target", "WantedBy=multi-user.target",
"DynamicUser=yes", "DynamicUser=yes",
"StateDirectory=gpu-turnstile", "StateDirectory=gpu-turnstile",
"RuntimeDirectory=gpu-turnstile",
"ProtectSystem=strict", "ProtectSystem=strict",
"NoNewPrivileges=yes", "NoNewPrivileges=yes",
"RestrictAddressFamilies=AF_UNIX AF_INET AF_INET6", "RestrictAddressFamilies=AF_UNIX AF_INET AF_INET6",
@@ -26,4 +27,12 @@ func TestRenderUnit(t *testing.T) {
t.Fatalf("unit missing %q:\n%s", want, unit) 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)
}
} }
+60 -22
View File
@@ -25,7 +25,7 @@ import (
"golang.org/x/sys/windows/svc" "golang.org/x/sys/windows/svc"
"golang.org/x/sys/windows/svc/mgr" "golang.org/x/sys/windows/svc/mgr"
"gpu-turnstile/internal/config" "gpu-turnstile/internal/supervise"
) )
// Name is the Windows service name. // Name is the Windows service name.
@@ -106,10 +106,12 @@ func (h *handler) Execute(_ []string, requests <-chan svc.ChangeRequest, status
// the service after 5s on failure — this is also what brings up a staged // 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, // update after the updater exits with a non-zero code. After registering,
// the virtual account is granted modify access to the install and data // the virtual account is granted modify access to the install and data
// directories (self-updates rewrite the exe), and read access to the // directories (self-updates rewrite the exe), read access to the
// config file if it lives elsewhere. The grants must come after // config file if it lives elsewhere, and — when the config sets COMFY_DIR —
// CreateService: the virtual account's SID only exists once the service is // recursive modify access to the managed ComfyUI's install tree, which may
// registered. // 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: // 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 service is stopped first if running (so the binary can be replaced),
@@ -146,6 +148,7 @@ func Install(configPath string, copyBin bool, version string) error {
if st, qErr := s.Query(); qErr == nil && if st, qErr := s.Query(); qErr == nil &&
(st.State == svc.Running || st.State == svc.StartPending) { (st.State == svc.Running || st.State == svc.StartPending) {
wasRunning = true wasRunning = true
fmt.Println("stopping the running gpu-turnstile service")
if err := stopAndWait(s); err != nil { if err := stopAndWait(s); err != nil {
return err return err
} }
@@ -161,6 +164,7 @@ func Install(configPath string, copyBin bool, version string) error {
} }
installedExe := filepath.Join(installDir, "gpu-turnstile.exe") installedExe := filepath.Join(installDir, "gpu-turnstile.exe")
if same, _ := sameFileContent(exe, installedExe); !same { if same, _ := sameFileContent(exe, installedExe); !same {
fmt.Printf("installing %s\n", installedExe)
if err := copyFile(exe, installedExe); err != nil { if err := copyFile(exe, installedExe); err != nil {
return fmt.Errorf("copy binary to %s: %w", installedExe, err) return fmt.Errorf("copy binary to %s: %w", installedExe, err)
} }
@@ -208,6 +212,7 @@ func Install(configPath string, copyBin bool, version string) error {
// Best effort: start now instead of waiting for the next boot. A // Best effort: start now instead of waiting for the next boot. A
// missing config (no consumer URLs) fails the start; the service stays // missing config (no consumer URLs) fails the start; the service stays
// registered and can be started once the config exists. // registered and can be started once the config exists.
fmt.Println("starting the gpu-turnstile service")
s.Start() s.Start()
return nil return nil
} }
@@ -243,6 +248,7 @@ func Install(configPath string, copyBin bool, version string) error {
return err return err
} }
if wasRunning { if wasRunning {
fmt.Println("starting the gpu-turnstile service")
if err := s.Start(); err != nil { if err := s.Start(); err != nil {
return fmt.Errorf("start service: %w", err) return fmt.Errorf("start service: %w", err)
} }
@@ -354,7 +360,7 @@ func grantAll(exe, configPath string) error {
return err return err
} }
} }
if logFile := configuredLogFile(configPath); logFile != "" { if logFile := configuredValue(configPath, "LOG_FILE"); logFile != "" {
dir := filepath.Dir(logFile) dir := filepath.Dir(logFile)
if err := os.MkdirAll(dir, 0o755); err == nil { if err := os.MkdirAll(dir, 0o755); err == nil {
if err := grantAccess(dir, "(OI)(CI)(M)"); err != nil { if err := grantAccess(dir, "(OI)(CI)(M)"); err != nil {
@@ -362,6 +368,31 @@ func grantAll(exe, configPath string) error {
} }
} }
} }
// 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
}
}
// Comfy-Desktop keeps models/input/output in a shared dir next
// to the install; the managed instance writes output there.
if shared := supervise.DesktopSharedDir(comfyDir); shared != "" {
if err := grantAccessTree(shared, "(OI)(CI)(M)"); err != nil {
return err
}
}
}
}
return nil return nil
} }
@@ -466,28 +497,35 @@ func RelaunchElevated(args []string) (int, error) {
// grantAccess gives the virtual account the icacls permission set (e.g. // grantAccess gives the virtual account the icacls permission set (e.g.
// "(OI)(CI)(M)") on path. // "(OI)(CI)(M)") on path.
func grantAccess(path, perms string) error { func grantAccess(path, perms string) error {
out, err := exec.Command("icacls", path, "/grant", virtualAccount+":"+perms).CombinedOutput() 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 { if err != nil {
return fmt.Errorf("grant %s access to %s: %w (%s)", virtualAccount, path, err, strings.TrimSpace(string(out))) 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 return nil
} }
// configuredLogFile reads LOG_FILE from the config file so the installer
// can pre-create and ACL the log directory. "" when unset or unreadable.
func configuredLogFile(configPath string) string {
f, err := os.Open(configPath)
if err != nil {
return ""
}
defer f.Close()
values, err := config.ParseEnvFile(f)
if err != nil {
return ""
}
return values["LOG_FILE"]
}
// RestartIfRunning restarts the service when it is installed and running // RestartIfRunning restarts the service when it is installed and running
// (used after a forced update staged a new binary). Reports whether a // (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 // restart happened. A service that is not installed or not running is not
+38
View File
@@ -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 ""
}
+45
View File
@@ -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)
}
}
+102 -3
View File
@@ -9,13 +9,82 @@ import (
"fmt" "fmt"
"io" "io"
"log/slog" "log/slog"
"net/url"
"os"
"os/exec" "os/exec"
"path/filepath"
"regexp"
"runtime" "runtime"
"strings" "strings"
"sync" "sync"
"time" "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. On a Comfy-Desktop standalone install the shared data
// directory (models, input, output) is added as --*-directory flags so the
// managed instance sees the desktop app's models.
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()
}
if shared := DesktopSharedDir(dir); shared != "" {
for _, sub := range []string{"models", "input", "output"} {
p := filepath.Join(shared, sub)
if st, err := os.Stat(p); err == nil && st.IsDir() {
cmd += ` --` + sub + `-directory "` + p + `"`
}
}
}
return cmd
}
// DesktopSharedDir returns the Comfy-Desktop shared data directory
// (<root>/ComfyUI-Shared) when dir looks like a desktop standalone install
// (<root>/ComfyUI-Installs/<name>/ComfyUI) and the shared models directory
// exists; "" otherwise. The desktop app keeps models, input and output
// there rather than inside the ComfyUI tree.
func DesktopSharedDir(dir string) string {
installs := filepath.Dir(filepath.Dir(dir))
if filepath.Base(installs) != "ComfyUI-Installs" {
return ""
}
shared := filepath.Join(filepath.Dir(installs), "ComfyUI-Shared")
if st, err := os.Stat(filepath.Join(shared, "models")); err == nil && st.IsDir() {
return shared
}
return ""
}
// Process is one managed child process. // Process is one managed child process.
type Process struct { type Process struct {
name string name string
@@ -95,6 +164,24 @@ func (p *Process) Running() bool {
return p.cmd != nil return p.cmd != nil
} }
// Status describes the child for status displays: "external" when something
// else serves the port, "running" once ready, "starting" while the child
// boots, "stopped" otherwise.
func (p *Process) Status() string {
p.mu.Lock()
defer p.mu.Unlock()
switch {
case p.external:
return "external"
case p.cmd == nil:
return "stopped"
case p.ready:
return "running"
default:
return "starting"
}
}
// Ready reports whether the server has answered a probe since its last // Ready reports whether the server has answered a probe since its last
// (re)start. Health checks use it to tell "starting up" from "outage". // (re)start. Health checks use it to tell "starting up" from "outage".
func (p *Process) Ready() bool { func (p *Process) Ready() bool {
@@ -144,6 +231,9 @@ func (p *Process) EnsureRunning() error {
p.external = false p.external = false
cmd := exec.Command(p.argv[0], p.argv[1:]...) cmd := exec.Command(p.argv[0], p.argv[1:]...)
cmd.Dir = p.dir cmd.Dir = p.dir
// Ask the child not to colorize (ComfyUI ignores this and colors
// anyway, so pipeLog also strips escape sequences).
cmd.Env = append(os.Environ(), "NO_COLOR=1", "TERM=dumb")
stdout, err := cmd.StdoutPipe() stdout, err := cmd.StdoutPipe()
if err != nil { if err != nil {
return err return err
@@ -237,7 +327,9 @@ func (p *Process) WatchIdle(ctx context.Context, idleTimeout time.Duration, gpuI
} }
// pipeLog forwards one child output stream to the log at INFO, line by // pipeLog forwards one child output stream to the log at INFO, line by
// line, prefixed with the process name. // line, prefixed with the process name. ANSI escape sequences are
// stripped: ComfyUI colorizes unconditionally, and the escapes only
// render as garbage in a log file.
func (p *Process) pipeLog(r io.Reader) { func (p *Process) pipeLog(r io.Reader) {
buf := make([]byte, 4096) buf := make([]byte, 4096)
var line string var line string
@@ -249,18 +341,25 @@ func (p *Process) pipeLog(r io.Reader) {
if i < 0 { if i < 0 {
break break
} }
p.log.Info(p.name + ": " + strings.TrimRight(line[:i], "\r")) p.log.Info(p.name + ": " + stripANSI(strings.TrimRight(line[:i], "\r")))
line = line[i+1:] line = line[i+1:]
} }
if err != nil { if err != nil {
if strings.TrimSpace(line) != "" { if strings.TrimSpace(line) != "" {
p.log.Info(p.name + ": " + line) p.log.Info(p.name + ": " + stripANSI(line))
} }
return return
} }
} }
} }
// ansiPattern matches CSI escape sequences (colors, cursor moves, …).
var ansiPattern = regexp.MustCompile("\x1b\\[[0-9;?]*[a-zA-Z]")
func stripANSI(s string) string {
return ansiPattern.ReplaceAllString(s, "")
}
// stopTree kills cmd's process, including its children on Windows (python // stopTree kills cmd's process, including its children on Windows (python
// launchers tend to spawn some). The Wait goroutine reaps it. // launchers tend to spawn some). The Wait goroutine reaps it.
func stopTree(cmd *exec.Cmd) { func stopTree(cmd *exec.Cmd) {
+84
View File
@@ -7,6 +7,8 @@ import (
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"os" "os"
"path/filepath"
"strings"
"testing" "testing"
"time" "time"
) )
@@ -190,3 +192,85 @@ func TestWatchIdleRespectsBusyGPU(t *testing.T) {
t.Fatal("process was stopped while the GPU was busy") 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)
}
}
func TestStripANSI(t *testing.T) {
cases := map[string]string{
"\x1b[32m[INFO]\x1b[0m Starting server": "[INFO] Starting server",
"\x1b[1m\x1b[31m[ERROR]\x1b[0m boom": "[ERROR] boom",
"plain line": "plain line",
"\x1b[33mWARN\x1b[0m: \x1b[1mbold\x1b[0m": "WARN: bold",
}
for in, want := range cases {
if got := stripANSI(in); got != want {
t.Errorf("stripANSI(%q) = %q, want %q", in, got, want)
}
}
}
func TestDesktopSharedDir(t *testing.T) {
root := t.TempDir()
comfy := filepath.Join(root, "ComfyUI-Installs", "rtx5080", "ComfyUI")
if err := os.MkdirAll(comfy, 0o755); err != nil {
t.Fatal(err)
}
if got := DesktopSharedDir(comfy); got != "" {
t.Fatalf("no shared dir yet: got %q, want empty", got)
}
shared := filepath.Join(root, "ComfyUI-Shared")
if err := os.MkdirAll(filepath.Join(shared, "models"), 0o755); err != nil {
t.Fatal(err)
}
if got := DesktopSharedDir(comfy); got != shared {
t.Fatalf("got %q, want %q", got, shared)
}
if got := DesktopSharedDir(filepath.Join(root, "plain", "ComfyUI")); got != "" {
t.Fatalf("non-desktop layout: got %q, want empty", got)
}
}
+25 -22
View File
@@ -173,11 +173,13 @@ func CleanupOld(exePath string) {
// Check performs a single update check. staged is true when a // Check performs a single update check. staged is true when a
// signature-verified binary has been swapped into place at exePath; the // signature-verified binary has been swapped into place at exePath; the
// caller should then restart the process. A nil error with staged=false // caller should then restart the process. to is the release tag the check
// means "no action" (up to date, APP_VER=dev, or no embedded public key); // resolved (the latest release or the pinned tag), set once the release
// a non-nil error means the check failed and the running binary is // fetch succeeded — even when staging afterwards fails. A nil error with
// untouched. // staged=false means "no action" (up to date, APP_VER=dev, or no embedded
func (u *Updater) Check(ctx context.Context, exePath string) (staged bool, err error) { // 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() log := u.logger()
desired := u.Desired desired := u.Desired
if desired == "" { if desired == "" {
@@ -185,15 +187,15 @@ func (u *Updater) Check(ctx context.Context, exePath string) (staged bool, err e
} }
if desired == "dev" { if desired == "dev" {
log.Debug("auto-update: APP_VER=dev, skipping") log.Debug("auto-update: APP_VER=dev, skipping")
return false, nil return false, "", nil
} }
if publicKeyPEM == "" { if publicKeyPEM == "" {
log.Debug("auto-update: no public key embedded, skipping") log.Debug("auto-update: no public key embedded, skipping")
return false, nil return false, "", nil
} }
api, err := u.apiURL() api, err := u.apiURL()
if err != nil { if err != nil {
return false, err return false, "", err
} }
pinned := desired != "stable" pinned := desired != "stable"
@@ -203,29 +205,30 @@ func (u *Updater) Check(ctx context.Context, exePath string) (staged bool, err e
} }
body, err := u.get(ctx, endpoint) body, err := u.get(ctx, endpoint)
if err != nil { if err != nil {
return false, fmt.Errorf("fetch release: %w", err) return false, "", fmt.Errorf("fetch release: %w", err)
} }
var rel release var rel release
if err := json.Unmarshal(body, &rel); err != nil { if err := json.Unmarshal(body, &rel); err != nil {
return false, fmt.Errorf("parse release: %w", err) return false, "", fmt.Errorf("parse release: %w", err)
} }
to = rel.TagName
if pinned { if pinned {
// Pin mode: any difference from the target tag means stage it — // Pin mode: any difference from the target tag means stage it —
// including downgrades and replacing a dev binary. // including downgrades and replacing a dev binary.
if u.Version == rel.TagName { if u.Version == rel.TagName {
log.Debug("auto-update: already on pinned version", "version", u.Version) log.Debug("auto-update: already on pinned version", "version", u.Version)
return false, nil return false, to, nil
} }
} else if u.Version != "" && u.Version != "dev" { } else if u.Version != "" && u.Version != "dev" {
// Stable mode: only strictly newer releases count; a dev binary // Stable mode: only strictly newer releases count; a dev binary
// cannot be compared and is always replaced by the latest release. // cannot be compared and is always replaced by the latest release.
newer, err := newerVersion(u.Version, rel.TagName) newer, err := newerVersion(u.Version, rel.TagName)
if err != nil { if err != nil {
return false, err return false, to, err
} }
if !newer { if !newer {
log.Debug("auto-update: up to date", "version", u.Version, "latest", rel.TagName) log.Debug("auto-update: up to date", "version", u.Version, "latest", rel.TagName)
return false, nil return false, to, nil
} }
} }
@@ -235,42 +238,42 @@ func (u *Updater) Check(ctx context.Context, exePath string) (staged bool, err e
} }
assetURL, ok := urls[u.Asset] assetURL, ok := urls[u.Asset]
if !ok { if !ok {
return false, fmt.Errorf("release %s has no asset %q", rel.TagName, u.Asset) return false, to, fmt.Errorf("release %s has no asset %q", rel.TagName, u.Asset)
} }
sigURL, ok := urls[u.Asset+".sig"] sigURL, ok := urls[u.Asset+".sig"]
if !ok { if !ok {
return false, fmt.Errorf("release %s has no signature asset %q", rel.TagName, u.Asset+".sig") return false, to, fmt.Errorf("release %s has no signature asset %q", rel.TagName, u.Asset+".sig")
} }
data, err := u.get(ctx, assetURL) data, err := u.get(ctx, assetURL)
if err != nil { if err != nil {
return false, fmt.Errorf("download %s: %w", u.Asset, err) return false, to, fmt.Errorf("download %s: %w", u.Asset, err)
} }
sig, err := u.get(ctx, sigURL) sig, err := u.get(ctx, sigURL)
if err != nil { if err != nil {
return false, fmt.Errorf("download signature: %w", err) return false, to, fmt.Errorf("download signature: %w", err)
} }
if sumURL, ok := urls[u.Asset+".sha256"]; ok { if sumURL, ok := urls[u.Asset+".sha256"]; ok {
sumText, err := u.get(ctx, sumURL) sumText, err := u.get(ctx, sumURL)
if err != nil { if err != nil {
return false, fmt.Errorf("download checksum: %w", err) return false, to, fmt.Errorf("download checksum: %w", err)
} }
want := strings.Fields(string(sumText))[0] want := strings.Fields(string(sumText))[0]
got := hex.EncodeToString(sha256Bytes(data)) got := hex.EncodeToString(sha256Bytes(data))
if !strings.EqualFold(want, got) { if !strings.EqualFold(want, got) {
return false, fmt.Errorf("sha256 mismatch: got %s, want %s", got, want) return false, to, fmt.Errorf("sha256 mismatch: got %s, want %s", got, want)
} }
} }
if err := verifySignature(publicKeyPEM, data, sig); err != nil { if err := verifySignature(publicKeyPEM, data, sig); err != nil {
return false, err return false, to, err
} }
if err := stage(exePath, data); err != nil { if err := stage(exePath, data); err != nil {
return false, fmt.Errorf("stage update: %w", err) return false, to, fmt.Errorf("stage update: %w", err)
} }
log.Info("auto-update: new version staged", "from", u.Version, "to", rel.TagName) log.Info("auto-update: new version staged", "from", u.Version, "to", rel.TagName)
return true, nil return true, to, nil
} }
func sha256Bytes(data []byte) []byte { func sha256Bytes(data []byte) []byte {
+14 -8
View File
@@ -102,13 +102,16 @@ func TestCheckStagesUpdate(t *testing.T) {
withPublicKey(t, f.pubPEM) withPublicKey(t, f.pubPEM)
exe := fakeExe(t) exe := fakeExe(t)
staged, err := f.updater("v0.1.2").Check(context.Background(), exe) staged, to, err := f.updater("v0.1.2").Check(context.Background(), exe)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
if !staged { if !staged {
t.Fatal("expected staged update") t.Fatal("expected staged update")
} }
if to != "v9.9.9" {
t.Fatalf("to = %q, want v9.9.9", to)
}
content, _ := os.ReadFile(exe) content, _ := os.ReadFile(exe)
if string(content) != "new-binary" { if string(content) != "new-binary" {
t.Fatalf("exe content = %q", content) t.Fatalf("exe content = %q", content)
@@ -125,7 +128,7 @@ func TestCheckRejectsTamperedSignature(t *testing.T) {
withPublicKey(t, f.pubPEM) withPublicKey(t, f.pubPEM)
exe := fakeExe(t) exe := fakeExe(t)
staged, err := f.updater("v0.1.2").Check(context.Background(), exe) staged, _, err := f.updater("v0.1.2").Check(context.Background(), exe)
if err == nil { if err == nil {
t.Fatal("expected signature error") t.Fatal("expected signature error")
} }
@@ -142,7 +145,7 @@ func TestCheckSkipsOlderOrEqual(t *testing.T) {
for _, tag := range []string{"v0.1.2", "v0.1.1", "v0.0.9"} { for _, tag := range []string{"v0.1.2", "v0.1.1", "v0.0.9"} {
f := newFakeGitea(t, tag, []byte("new-binary")) f := newFakeGitea(t, tag, []byte("new-binary"))
withPublicKey(t, f.pubPEM) withPublicKey(t, f.pubPEM)
staged, err := f.updater("v0.1.2").Check(context.Background(), fakeExe(t)) staged, _, err := f.updater("v0.1.2").Check(context.Background(), fakeExe(t))
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -155,7 +158,7 @@ func TestCheckSkipsOlderOrEqual(t *testing.T) {
func TestCheckSkipsWithoutPublicKey(t *testing.T) { func TestCheckSkipsWithoutPublicKey(t *testing.T) {
f := newFakeGitea(t, "v9.9.9", []byte("new-binary")) f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
withPublicKey(t, "") withPublicKey(t, "")
staged, err := f.updater("v0.1.2").Check(context.Background(), fakeExe(t)) staged, _, err := f.updater("v0.1.2").Check(context.Background(), fakeExe(t))
if err != nil || staged { if err != nil || staged {
t.Fatalf("staged=%v err=%v, want no action without key", staged, err) t.Fatalf("staged=%v err=%v, want no action without key", staged, err)
} }
@@ -167,7 +170,7 @@ func TestCheckDevBuildGetsStable(t *testing.T) {
f := newFakeGitea(t, "v9.9.9", []byte("new-binary")) f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
withPublicKey(t, f.pubPEM) withPublicKey(t, f.pubPEM)
exe := fakeExe(t) exe := fakeExe(t)
staged, err := f.updater("dev").Check(context.Background(), exe) staged, _, err := f.updater("dev").Check(context.Background(), exe)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -184,7 +187,7 @@ func TestCheckDesiredDevDisables(t *testing.T) {
f := newFakeGitea(t, "v9.9.9", []byte("new-binary")) f := newFakeGitea(t, "v9.9.9", []byte("new-binary"))
withPublicKey(t, f.pubPEM) withPublicKey(t, f.pubPEM)
for _, version := range []string{"dev", "v0.1.2"} { for _, version := range []string{"dev", "v0.1.2"} {
staged, err := f.updaterDesired(version, "dev").Check(context.Background(), fakeExe(t)) staged, _, err := f.updaterDesired(version, "dev").Check(context.Background(), fakeExe(t))
if err != nil || staged { if err != nil || staged {
t.Fatalf("version %s: staged=%v err=%v, want no action with APP_VER=dev", version, staged, err) t.Fatalf("version %s: staged=%v err=%v, want no action with APP_VER=dev", version, staged, err)
} }
@@ -198,13 +201,16 @@ func TestCheckPinned(t *testing.T) {
f := newFakeGitea(t, "v0.5.0", []byte("pinned-binary")) f := newFakeGitea(t, "v0.5.0", []byte("pinned-binary"))
withPublicKey(t, f.pubPEM) withPublicKey(t, f.pubPEM)
exe := fakeExe(t) exe := fakeExe(t)
staged, err := f.updaterDesired(version, "v0.5.0").Check(context.Background(), exe) staged, to, err := f.updaterDesired(version, "v0.5.0").Check(context.Background(), exe)
if err != nil { if err != nil {
t.Fatalf("version %s: %v", version, err) t.Fatalf("version %s: %v", version, err)
} }
if !staged { if !staged {
t.Fatalf("version %s: expected pinned v0.5.0 to be staged", version) 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) content, _ := os.ReadFile(exe)
if string(content) != "pinned-binary" { if string(content) != "pinned-binary" {
t.Fatalf("version %s: exe content = %q", version, content) t.Fatalf("version %s: exe content = %q", version, content)
@@ -213,7 +219,7 @@ func TestCheckPinned(t *testing.T) {
f := newFakeGitea(t, "v0.5.0", []byte("pinned-binary")) f := newFakeGitea(t, "v0.5.0", []byte("pinned-binary"))
withPublicKey(t, f.pubPEM) withPublicKey(t, f.pubPEM)
staged, err := f.updaterDesired("v0.5.0", "v0.5.0").Check(context.Background(), fakeExe(t)) staged, _, err := f.updaterDesired("v0.5.0", "v0.5.0").Check(context.Background(), fakeExe(t))
if err != nil || staged { if err != nil || staged {
t.Fatalf("staged=%v err=%v, want no action when already on the pinned version", staged, err) t.Fatalf("staged=%v err=%v, want no action when already on the pinned version", staged, err)
} }