9 Commits
Author SHA1 Message Date
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
16 changed files with 679 additions and 73 deletions
+153 -35
View File
@@ -17,11 +17,13 @@ import (
"path/filepath" "path/filepath"
"runtime" "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"
@@ -39,6 +41,11 @@ var version = "dev"
// process: a signed update has been staged and the GPU lock is idle. // process: a signed update has been staged 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.
@@ -51,7 +58,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, 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 {
@@ -72,13 +79,15 @@ 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] == "--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, 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.
@@ -93,6 +102,9 @@ Usage:
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 -h | --help this help gpu-turnstile -h | --help this help
Options: Options:
@@ -123,12 +135,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, 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 && !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
@@ -151,12 +163,16 @@ 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 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())
} }
if len(args) > 0 { if len(args) > 0 {
fatalUsage("error: unknown arguments: %s", strings.Join(args, " ")) fatalUsage("error: unknown arguments: %s", strings.Join(args, " "))
@@ -342,17 +358,25 @@ func forceUpdateCommand(configPath string, elevatedChild bool) int {
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
} }
@@ -364,12 +388,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
} }
@@ -378,13 +403,57 @@ 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)
}
// 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
@@ -659,7 +728,35 @@ func run(ctx context.Context, cfg config.Config, log *slog.Logger, logOut io.Wri
service.StartWatchdog(ctx) service.StartWatchdog(ctx)
if cfg.AutoUpdate { if cfg.AutoUpdate {
go updateLoop(ctx, cfg, log, lk, isService) exePath, err := os.Executable()
if err != nil {
log.Warn("auto-update disabled: cannot locate executable", "err", err)
} else {
u := &update.Updater{Repo: cfg.UpdateRepo, Asset: cfg.UpdateAsset, Version: version, Desired: cfg.AppVersion, Log: log}
// applyStaged is shared by the hourly loop and the control
// channel; the once guard keeps a second trigger from
// double-waiting on the GPU lock.
var once sync.Once
applyStaged := func(to string) {
if !isService {
log.Warn("auto-update: new binary staged; restart gpu-turnstile to apply", "version", to)
return
}
log.Warn("auto-update: staged; restarting once the GPU is idle", "version", to)
once.Do(func() {
go func() {
if waitForIdle(ctx, lk, 24*time.Hour) {
log.Warn("auto-update: restarting to apply update")
os.Exit(exitCodeUpdate)
}
}()
})
}
go updateLoop(ctx, cfg.UpdateInterval, log, u, exePath, applyStaged)
if isService {
serveControl(ctx, log, u, exePath, applyStaged)
}
}
} }
select { select {
@@ -801,42 +898,63 @@ 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 trigger an update
// check via --force-update without admin rights. The check payload is
// signature-verified regardless of who asks; triggers are rate-limited to
// one per minute so the channel cannot be used to spam restarts.
func serveControl(ctx context.Context, log *slog.Logger, u *update.Updater, exePath string, applyStaged func(to string)) {
var mu sync.Mutex
var lastTrigger time.Time
h := func(cmd string) string {
if cmd != control.CmdUpdateNow {
return "ERR unknown command: " + cmd
}
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)
}
}
// 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 {
+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.2.0 image: git.rambossek.at/public/gpu-turnstile:v0.2.3
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
+61
View File
@@ -0,0 +1,61 @@
// 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.
//
// 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"
)
// CmdUpdateNow asks the service to check for, stage and (once the GPU is
// idle) restart onto a signed update immediately.
const CmdUpdateNow = "update-now"
// 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
// serveConn runs the line protocol on one accepted connection.
func serveConn(c io.ReadWriteCloser, h Handler) {
defer c.Close()
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
}
+54
View File
@@ -0,0 +1,54 @@
//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
}
go serveConn(c, h)
}
}()
return nil
}
// 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
}
+46
View File
@@ -0,0 +1,46 @@
package control
import (
"errors"
"net"
"strings"
"testing"
)
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)
}
}
+112
View File
@@ -0,0 +1,112 @@
//go:build windows
package control
import (
"context"
"fmt"
"log/slog"
"os"
"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
}
go func() {
// Blocks until a client connects; on process exit the
// handle goes away with everything else.
if err := windows.ConnectNamedPipe(pipe, nil); err != nil {
windows.CloseHandle(pipe)
return
}
f := os.NewFile(uintptr(pipe), pipePath)
serveConn(f, h) // closes f, and with it the pipe handle
}()
}
}()
return nil
}
// Ask sends one command to the running service and returns its reply.
func Ask(cmd string) (string, error) {
name, err := windows.UTF16PtrFromString(pipePath)
if err != nil {
return "", err
}
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 != nil {
return "", ErrUnavailable
}
f := os.NewFile(uintptr(handle), pipePath)
defer f.Close()
return readReply(f, cmd)
}
+15 -2
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".
@@ -64,11 +66,20 @@ func Run(run func(ctx context.Context) error) error {
// 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. A managed ComfyUI (comfyDir) gets a // notify socket, so it loses nothing. A managed ComfyUI (comfyDir) gets a
// BindPaths hole through ProtectHome/ProtectSystem: it reads its venv and // BindPaths hole through ProtectHome/ProtectSystem: it reads its venv and
// writes output/temp/user data under COMFY_DIR. // 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 { func renderUnit(exePath, configPath, comfyDir string) string {
bind := "" bind := ""
if comfyDir != "" { if comfyDir != "" {
bind = "BindPaths=" + comfyDir + "\n" 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
@@ -84,6 +95,8 @@ RestartSec=5s
DynamicUser=yes DynamicUser=yes
StateDirectory=%s StateDirectory=%s
RuntimeDirectory=%s
RuntimeDirectoryMode=0755
%sProtectSystem=strict %sProtectSystem=strict
ProtectHome=yes ProtectHome=yes
PrivateTmp=yes PrivateTmp=yes
@@ -106,7 +119,7 @@ SystemCallErrorNumber=EPERM
[Install] [Install]
WantedBy=multi-user.target WantedBy=multi-user.target
`, exePath, configPath, Name, bind) `, 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.
+1
View File
@@ -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",
+17 -1
View File
@@ -24,6 +24,8 @@ import (
"golang.org/x/sys/windows" "golang.org/x/sys/windows"
"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/supervise"
) )
// Name is the Windows service name. // Name is the Windows service name.
@@ -375,6 +377,20 @@ func grantAll(exe, configPath string) error {
if err := grantAccessTree(comfyDir, "(OI)(CI)(M)"); err != nil { if err := grantAccessTree(comfyDir, "(OI)(CI)(M)"); err != nil {
return err 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
@@ -496,7 +512,7 @@ func grantAccessTree(path, perms string) error {
func runIcacls(path, perms string, recursive bool) error { func runIcacls(path, perms string, recursive bool) error {
args := []string{path, "/grant", virtualAccount + ":" + perms} args := []string{path, "/grant", virtualAccount + ":" + perms}
if recursive { if recursive {
fmt.Printf("granting %s modify access to %s (large trees can take minutes)\n", virtualAccount, path) fmt.Printf("granting %s %s access to %s (large trees can take minutes)\n", virtualAccount, perms, path)
args = append(args, "/T") args = append(args, "/T")
} }
start := time.Now() start := time.Now()
+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)
}
}
+44 -4
View File
@@ -13,6 +13,7 @@ import (
"os" "os"
"os/exec" "os/exec"
"path/filepath" "path/filepath"
"regexp"
"runtime" "runtime"
"strings" "strings"
"sync" "sync"
@@ -43,7 +44,9 @@ func ComfyLayout(goos, dir string) (python, script string) {
// DefaultComfyCommand builds the launch command for the standard venv // DefaultComfyCommand builds the launch command for the standard venv
// layout (see ComfyLayout): the script is passed relative to dir so dir // layout (see ComfyLayout): the script is passed relative to dir so dir
// stays the working directory, and --port is taken from comfyURL when the // stays the working directory, and --port is taken from comfyURL when the
// URL carries one. // 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 { func DefaultComfyCommand(goos, dir, comfyURL string) string {
python, script := ComfyLayout(goos, dir) python, script := ComfyLayout(goos, dir)
rel, err := filepath.Rel(dir, script) rel, err := filepath.Rel(dir, script)
@@ -54,9 +57,34 @@ func DefaultComfyCommand(goos, dir, comfyURL string) string {
if u, err := url.Parse(comfyURL); err == nil && u.Port() != "" { if u, err := url.Parse(comfyURL); err == nil && u.Port() != "" {
cmd += " --port " + 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 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
@@ -185,6 +213,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
@@ -278,7 +309,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
@@ -290,18 +323,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) {
+35
View File
@@ -239,3 +239,38 @@ func TestComfyLayoutAndDefaultCommand(t *testing.T) {
t.Errorf("missing-layout script = %s", script) 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)
} }