Compare commits
187
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d6a2277e6a | ||
|
|
a59fdee094 | ||
|
|
3fe2018589 | ||
|
|
00bf55c96f | ||
|
|
1e4fdd12e0 | ||
|
|
3b6e15dad5 | ||
|
|
de51120d6a | ||
|
|
3e5aa90e3f | ||
|
|
8cc7c88692 | ||
|
|
e7e97e4b4e | ||
|
|
97bfbfecd1 | ||
|
|
0c8a40774f | ||
|
|
f45c885171 | ||
|
|
da8358916c | ||
|
|
0268f43e7f | ||
|
|
6d42d2255e | ||
|
|
1103898d7b | ||
|
|
28220ff2dd | ||
|
|
fa9bfd41be | ||
|
|
0e47053823 | ||
|
|
bea72cadbe | ||
|
|
6624b3fb19 | ||
|
|
fc603bcc92 | ||
|
|
2300d4d5f7 | ||
|
|
c7890d2a5a | ||
|
|
633e991edc | ||
|
|
64b3785b0b | ||
|
|
4b3c62d783 | ||
|
|
103a5dbbe4 | ||
|
|
86f6c9ec8e | ||
|
|
e1a76f9aca | ||
|
|
57263edad7 | ||
|
|
4c0a0a26b7 | ||
|
|
634936eb10 | ||
|
|
422dbd1313 | ||
|
|
d23f22d56b | ||
|
|
eaa87ff041 | ||
|
|
0ee9bc35d0 | ||
|
|
a1cd15964c | ||
|
|
895e62df7c | ||
|
|
49c175301a | ||
|
|
b92d35ac5d | ||
|
|
b6bd125f23 | ||
|
|
53331a6f5a | ||
|
|
1d445855d2 | ||
|
|
9dc01d4b8b | ||
|
|
e93b6366e6 | ||
|
|
89cee3f824 | ||
|
|
078bad8e0f | ||
|
|
032c5aba86 | ||
|
|
a8371baaa8 | ||
|
|
5a615d2c66 | ||
|
|
12480b81b1 | ||
|
|
ef1cb57944 | ||
|
|
6447788fdf | ||
|
|
afa361913c | ||
|
|
0d3c81b596 | ||
|
|
98c9e68aae | ||
|
|
772e3e82b4 | ||
|
|
228019c9a8 | ||
|
|
e38b5f5741 | ||
|
|
f60f5f1f59 | ||
|
|
7640f42dce | ||
|
|
3eeed0cfe9 | ||
|
|
b946fded12 | ||
|
|
7718a7a10b | ||
|
|
0687e13b57 | ||
|
|
89d585a36e | ||
|
|
3f5114923b | ||
|
|
43f1d46fe6 | ||
|
|
3223a20f88 | ||
|
|
2d37627ab2 | ||
|
|
54a88f694b | ||
|
|
3441201be0 | ||
|
|
de5d848189 | ||
|
|
aa7c2f5801 | ||
|
|
b7f14ce4ad | ||
|
|
7a928da1a0 | ||
|
|
0961a5e512 | ||
|
|
0619df8dff | ||
|
|
a91b27b518 | ||
|
|
605236566c | ||
|
|
918c400f29 | ||
|
|
76d16ac1bd | ||
|
|
f22606f3ad | ||
|
|
f8492dd676 | ||
|
|
f6afa43d07 | ||
|
|
835a889326 | ||
|
|
2f3888549b | ||
|
|
e9e4d95d06 | ||
|
|
2048d2793a | ||
|
|
52ce462396 | ||
|
|
e300739d78 | ||
|
|
7efae54cf8 | ||
|
|
0257bfcfec | ||
|
|
933c1a2cf1 | ||
|
|
c59a787a10 | ||
|
|
ed6d7a9652 | ||
|
|
3a1013527f | ||
|
|
243220fc3a | ||
|
|
d57b4babc5 | ||
|
|
c198a8349a | ||
|
|
c4f3ca457d | ||
|
|
e1e3a477a7 | ||
|
|
e20add344c | ||
|
|
0a07202634 | ||
|
|
3cf1007f28 | ||
|
|
1044483edf | ||
|
|
a04c972b71 | ||
|
|
8c1afe6d9d | ||
|
|
809314a459 | ||
|
|
2dcfd0bb55 | ||
|
|
93616a9c2a | ||
|
|
4072ec3db4 | ||
|
|
0548386cb3 | ||
|
|
45ec840ead | ||
|
|
fcc6e4a843 | ||
|
|
99a98eaefe | ||
|
|
f72439d6cf | ||
|
|
a355ad4ab6 | ||
|
|
daefad8769 | ||
|
|
afe013a6bc | ||
|
|
f0764532e2 | ||
|
|
d11d608c2d | ||
|
|
109199e024 | ||
|
|
9d133870e5 | ||
|
|
0d9a392e8d | ||
|
|
e6284a5d5f | ||
|
|
69567e2e56 | ||
|
|
8e98e7a1be | ||
|
|
b0785c4e6d | ||
|
|
5baf82bfb9 | ||
|
|
f5d33aad8c | ||
|
|
539309ea84 | ||
|
|
e450a37d4b | ||
|
|
77b66abd94 | ||
|
|
1a39061849 | ||
|
|
dd1aa3654d | ||
|
|
8430f9843c | ||
|
|
3299d5986b | ||
|
|
a53ddc35a8 | ||
|
|
62eddfad41 | ||
|
|
5462eeacaa | ||
|
|
9aa01d4e72 | ||
|
|
17d181e427 | ||
|
|
1762c57355 | ||
|
|
e3eda1af2c | ||
|
|
40b3c4f460 | ||
|
|
eed841a8ca | ||
|
|
49ab178db7 | ||
|
|
3482399197 | ||
|
|
df4d016a7d | ||
|
|
128b07c923 | ||
|
|
fd6d21401b | ||
|
|
c9adcb2647 | ||
|
|
e6d3103ba8 | ||
|
|
1e6de9155b | ||
|
|
a41dc8bcac | ||
|
|
a8c5ce5c31 | ||
|
|
e0e19f3dc9 | ||
|
|
b73f31b237 | ||
|
|
6bcd3429ac | ||
|
|
81b6bbc4d3 | ||
|
|
def15b8423 | ||
|
|
4df7d4e97e | ||
|
|
42b63488e9 | ||
|
|
ee3e87c0a7 | ||
|
|
b8f1d7f19d | ||
|
|
155b9345b9 | ||
|
|
37c5df6f3a | ||
|
|
3be001f3fd | ||
|
|
28c7bacefb | ||
|
|
fbb258d2e2 | ||
|
|
6837ba25ac | ||
|
|
366b55d9d1 | ||
|
|
be007e9d77 | ||
|
|
030bc47515 | ||
|
|
31d4db65e8 | ||
|
|
fd30a6c4ad | ||
|
|
dcaed7cbf4 | ||
|
|
9615cd5294 | ||
|
|
e4c1ef0de6 | ||
|
|
5229a9504c | ||
|
|
214a859344 | ||
|
|
b72436a4e5 | ||
|
|
c2955dbe92 | ||
|
|
a6f008ec38 |
@@ -40,6 +40,7 @@ sudo apt-get install -y \
|
||||
libwebkit2gtk-4.1-dev libgtk-3-dev libpango1.0-dev libcairo2-dev \
|
||||
libsoup-3.0-dev libgdk-pixbuf-2.0-dev \
|
||||
libayatana-appindicator3-dev librsvg2-dev libssl-dev libxdo-dev \
|
||||
gstreamer1.0-plugins-good \
|
||||
libasound2-dev build-essential curl wget file
|
||||
```
|
||||
|
||||
|
||||
+1
-1
@@ -4,7 +4,7 @@
|
||||
|
||||
| Version | Supported |
|
||||
|---------|-----------|
|
||||
| 0.3.x (latest release + `main` previews) | ✅ Current — all fixes land here |
|
||||
| 0.5.x (latest release + `main` previews) | ✅ Current — all fixes land here |
|
||||
| 0.2.7 | ⚠️ Legacy stable — security fixes only, upgrade recommended |
|
||||
| < 0.2.7 | ❌ No longer supported |
|
||||
|
||||
|
||||
@@ -51,6 +51,13 @@ jobs:
|
||||
- os: ubuntu-latest
|
||||
platform: linux-x86_64
|
||||
experimental: false
|
||||
- os: ubuntu-24.04-arm
|
||||
platform: linux-aarch64
|
||||
# Apple Silicon under Asahi Linux. Experimental: the Vulkan
|
||||
# (Honeykrisp GPU) build path is new and the hosted arm64
|
||||
# runner has no GPU — it validates that the binary builds;
|
||||
# on-host Vulkan acceleration is exercised by users.
|
||||
experimental: true
|
||||
- os: windows-latest
|
||||
platform: windows-x86_64
|
||||
experimental: false
|
||||
@@ -80,11 +87,18 @@ jobs:
|
||||
# Linux-only: upstream `buildcpu.sh` enables `-DGGML_BLAS=ON` which
|
||||
# requires a system BLAS implementation at cmake configure time.
|
||||
- name: Linux system deps (BLAS for ggml-blas backend)
|
||||
if: matrix.platform == 'linux-x86_64'
|
||||
if: startsWith(matrix.platform, 'linux')
|
||||
run: |
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y libopenblas-dev pkg-config
|
||||
|
||||
# linux-aarch64: let the build script's Vulkan path (Honeykrisp GPU
|
||||
# on Asahi) engage instead of silently falling back to CPU.
|
||||
- name: Vulkan dev deps (linux-aarch64 GPU backend)
|
||||
if: matrix.platform == 'linux-aarch64'
|
||||
run: |
|
||||
sudo apt-get install -y glslc libvulkan-dev spirv-headers
|
||||
|
||||
- name: Build omnivoice-tts
|
||||
shell: bash
|
||||
# Pass values through env (quoted) rather than ${{ }} interpolation
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
# Installer smoke — runs scripts/install.sh / scripts/install.ps1 end-to-end
|
||||
# on all three desktop platforms so the one-liner installers can't rot.
|
||||
#
|
||||
# Gated by `paths` because a cold run downloads multi-GB wheels (torch) and
|
||||
# takes ~15-30 min per OS; it only needs to fire when an installer or this
|
||||
# workflow changes. The heavy Tauri bundles stay in release.yml (tag push).
|
||||
|
||||
name: Install smoke
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
paths:
|
||||
- "scripts/install.sh"
|
||||
- "scripts/install.ps1"
|
||||
- ".github/workflows/install-smoke.yml"
|
||||
push:
|
||||
branches: [main]
|
||||
paths:
|
||||
- "scripts/install.sh"
|
||||
- "scripts/install.ps1"
|
||||
- ".github/workflows/install-smoke.yml"
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
env:
|
||||
FORCE_JAVASCRIPT_ACTIONS_TO_NODE24: true
|
||||
|
||||
jobs:
|
||||
install:
|
||||
name: Install (${{ matrix.os }})
|
||||
runs-on: ${{ matrix.os }}
|
||||
timeout-minutes: 60
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
os: [ubuntu-22.04, macos-latest, windows-latest]
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
# Running `sh scripts/install.sh` from the repo root exercises the
|
||||
# repo-root resolution (script dir is scripts/, project root one level
|
||||
# up) — the exact bug that made a local run clone a duplicate repo.
|
||||
# Binary mode is the default: prebuilt release asset, checksum verified.
|
||||
- name: Run installer — binary (macOS/Linux)
|
||||
if: runner.os != 'Windows'
|
||||
run: sh scripts/install.sh
|
||||
|
||||
- name: Verify install — binary (macOS/Linux)
|
||||
if: runner.os != 'Windows'
|
||||
run: |
|
||||
if [ "$(uname)" = "Darwin" ]; then
|
||||
test -d "/Applications/VoiceStudio.app" || { echo "::error::VoiceStudio.app missing from /Applications"; exit 1; }
|
||||
echo "✓ VoiceStudio.app installed in /Applications"
|
||||
else
|
||||
test -x "$HOME/.local/bin/VoiceStudio" || { echo "::error::AppImage missing from ~/.local/bin"; exit 1; }
|
||||
"$HOME/.local/bin/VoiceStudio" --appimage-help >/dev/null 2>&1 || true
|
||||
echo "✓ AppImage installed and executable"
|
||||
fi
|
||||
|
||||
# Source mode stays covered end-to-end behind --source.
|
||||
- name: Run installer — source (macOS/Linux)
|
||||
if: runner.os != 'Windows'
|
||||
run: sh scripts/install.sh --source
|
||||
|
||||
- name: Verify install — source (macOS/Linux)
|
||||
if: runner.os != 'Windows'
|
||||
working-directory: ${{ github.workspace }}
|
||||
run: |
|
||||
test -d .venv || { echo "::error::.venv missing"; exit 1; }
|
||||
test -f frontend/dist/index.html || { echo "::error::frontend build missing"; exit 1; }
|
||||
echo "✓ venv + frontend bundle present"
|
||||
|
||||
# Binary mode is the default; CI runs msiexec silently.
|
||||
- name: Run installer — binary (Windows)
|
||||
if: runner.os == 'Windows'
|
||||
env:
|
||||
CI: true
|
||||
shell: pwsh
|
||||
run: '& { $ErrorActionPreference = "Stop"; & "${{ github.workspace }}\scripts\install.ps1" }'
|
||||
|
||||
- name: Verify install — binary (Windows)
|
||||
if: runner.os == 'Windows'
|
||||
shell: pwsh
|
||||
run: |
|
||||
$paths = @(
|
||||
"HKLM:\Software\Microsoft\Windows\CurrentVersion\Uninstall\*",
|
||||
"HKLM:\Software\WOW6432Node\Microsoft\Windows\CurrentVersion\Uninstall\*",
|
||||
"HKCU:\Software\Microsoft\Windows\CurrentVersion\Uninstall\*"
|
||||
)
|
||||
$key = Get-ItemProperty $paths -ErrorAction SilentlyContinue |
|
||||
Where-Object { $_.DisplayName -match "VoiceStudio|OmniVoice" } |
|
||||
Select-Object -First 1
|
||||
if (-not $key) {
|
||||
Get-ItemProperty $paths -ErrorAction SilentlyContinue |
|
||||
Where-Object DisplayName | ForEach-Object { Write-Host " installed: $($_.DisplayName)" }
|
||||
Write-Host "::error::MSI product not registered"; exit 1
|
||||
}
|
||||
Write-Host "✓ MSI product registered: $($key.DisplayName)"
|
||||
|
||||
# Source mode stays covered end-to-end behind -Source.
|
||||
- name: Run installer — source (Windows)
|
||||
if: runner.os == 'Windows'
|
||||
env:
|
||||
VOICESTUDIO_INSTALL_MODE: source
|
||||
shell: pwsh
|
||||
run: '& { $ErrorActionPreference = "Stop"; & "${{ github.workspace }}\scripts\install.ps1" }'
|
||||
|
||||
- name: Verify install — source (Windows)
|
||||
if: runner.os == 'Windows'
|
||||
shell: pwsh
|
||||
run: |
|
||||
if (-not (Test-Path .venv)) { Write-Host "::error::.venv missing"; exit 1 }
|
||||
if (-not (Test-Path frontend\dist\index.html)) { Write-Host "::error::frontend build missing"; exit 1 }
|
||||
Write-Host "✓ venv + frontend bundle present"
|
||||
|
||||
- name: Upload install log on failure
|
||||
if: failure()
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: install-log-${{ matrix.os }}
|
||||
path: |
|
||||
/Users/runner/Library/Application Support/OmniVoice/*.log
|
||||
/home/runner/.local/share/VoiceStudio/*.log
|
||||
${{ runner.temp }}/VoiceStudio/**/*.log
|
||||
if-no-files-found: ignore
|
||||
@@ -159,6 +159,13 @@ tests/probe/reports/
|
||||
# reports). Working notes for whoever is driving a change, not a repo artifact.
|
||||
/remote/
|
||||
|
||||
# OmniVoice GGUF runtime build artifacts (scripts/build-omnivoice-tts.sh).
|
||||
# Only 0-byte placeholders of omnivoice-tts-* are tracked; real binaries,
|
||||
# the checksums manifest and the copied libggml shared libs ship via CI.
|
||||
bin/libggml*
|
||||
bin/checksums.sha256
|
||||
bin/omnivoice-tts-linux-aarch64
|
||||
|
||||
# Dubbing-demo intermediates. The .mp4/.srt/manifest.json in this directory ARE
|
||||
# committed (they ship with the app); the per-language source WAVs are just the
|
||||
# inputs scripts/render_dub_demo_audio.py hands to scripts/build_dub_demo.sh.
|
||||
|
||||
+86
-3
@@ -10,20 +10,65 @@ the frozen-backend fallback mirror it for their toolchains.
|
||||
|
||||
**Highlights**
|
||||
|
||||
- The backend now answers within a second of launch and narrates its startup step by step
|
||||
- Reporting a bug from an outdated build now offers the latest release first
|
||||
- The backend is only announced ready once it can actually serve, and crash-loop restarts now pace themselves
|
||||
### Changed
|
||||
|
||||
### Added
|
||||
|
||||
### Docs
|
||||
|
||||
### Fixed
|
||||
|
||||
## [0.5.1] — 2026-08-28
|
||||
|
||||
**Highlights**
|
||||
|
||||
- OmniVoice generation on Apple Silicon now runs in a crash-isolated child, so fatal MPS memory exits no longer take down the local backend (#1697, #1698) — thanks @ndntran14!
|
||||
- Model-load GPU exhaustion now returns a sanitized, actionable dubbing error, and readiness correctly attributes the shared model status to TTS (#1695)
|
||||
- Source-mode development now restarts an isolated backend crash without tearing down the UI, while repeated crash loops still stop loudly with diagnostics (#1690)
|
||||
- Dubbing playback now keeps an audible companion source when a WebView can render the preview picture but cannot decode its audio (#1692)
|
||||
- Model Catalogue engine rows now use the available desktop width and keep identity, runtime state, and actions from crowding one another (#1689)
|
||||
- VoiceStudio now acts as a local speech platform: other apps can trigger its native dictation or connect through versioned HTTP, WebSocket, JSON-RPC, CLI, and MCP transports (#1646)
|
||||
- A timed-out in-process dub transcription no longer starts a second WhisperX/CTranslate2 call over the abandoned native worker, preventing the overlapping access that preceded Windows `0xC0000005` exits (#1669)
|
||||
- Windows debugger termination code `0x40010004` is no longer misreported as a backend crash or charged against automatic restart recovery (#1663)
|
||||
- Studio now keeps one generation reservation across page changes, preventing a remount from stacking native jobs until the backend reports capacity busy or is killed under memory pressure (#1670)
|
||||
- Uploaded dubbing videos are normalized to browser-safe H.264/AAC before preview, preventing valid VP9, AV1, or Opus media from failing with “no supported sources” (#1644)
|
||||
- Dubbing now separates spoken and target languages, preserves translations through segment cleanup, and lets failed translations be retried or skipped without restarting the batch (#1654) — thanks @Number16BusShelter!
|
||||
- Importing replacement SRT subtitles now keeps each cue bound to the best-overlapping source speaker and clone instead of resetting every line to a random default voice (#1660) — thanks @invio-a11y!
|
||||
- Uploading a Dub preview no longer blocks every backend request while ffmpeg extracts its audio (#1667) — thanks @tfreyd!
|
||||
- Docker quick starts now require the administrator key needed through container NAT instead of starting a UI whose protected actions return 403 (#1651) — thanks @wd357dui!
|
||||
- WSL2 AMD containers now use the `/dev/dxg` ROCDXG bridge with actionable GPU diagnostics instead of silently falling back to CPU (#1655) — thanks @wd357dui!
|
||||
- Ad-hoc voice-clone references now stay alive until cancelled or timed-out GPU work actually stops reading them, so prompt caching can finish instead of failing on a deleted temp file (#1668) — thanks @tfreyd!
|
||||
- Dictation now stays bound to the app where it started and recovers locally from silent recognizer output (#1175)
|
||||
- The backend now answers within a second of launch and narrates its startup step by step (#1550)
|
||||
- Reporting a bug from an outdated build now offers the latest release first (#1547)
|
||||
- The backend is only announced ready once it can actually serve, and crash-loop restarts now pace themselves (#1548)
|
||||
- Invisible watermarking no longer stalls — or silently skips — the first take of a session (#1615)
|
||||
- Dub subtitles can be retimed, inserted, and merged in either direction from the segment table (#1612) — thanks @invio-a11y!
|
||||
|
||||
### Changed
|
||||
- Model Catalogue now uses one breathable workspace canvas with simpler pane and engine-family navigation instead of nested cards and scroll regions (#1685)
|
||||
- Linux source launchers now catch missing libxdo and GStreamer audio plugins before they can cause a linker error or an aborted, blank WebKit renderer (#1680, #1682)
|
||||
- Dictation now carries one native output session from shortcut-down through final delivery, restores text, HTML, image, or file-list clipboards only when untouched, keeps Wayland copy-safe unless current-focus insertion is explicitly enabled, and retries silent Sherpa speech only through an already-installed local ASR model (#1175)
|
||||
- The backend binds its port immediately and reports startup progress live — `/health` answers 503-with-step and a new `/startup/progress` endpoint lists every step while PyTorch, API routes, and database migrations load in the background, so "starting at step X" is never mistakable for "dead"; the desktop splash narrates each step (#1550)
|
||||
|
||||
### Added
|
||||
- A bundled Rust loopback sidecar exposes dictation start/stop/toggle, focused-output sessions, discovery, and JSON-RPC; the backend adds versioned streaming events and a dependency-free CLI bridge for Herdr, coding agents, editors, desktop apps, and TUIs (#1646)
|
||||
- Headless NVIDIA and ROCm machines can now join as worker-only Docker Compose services with no published UI and durable protocol-v2 enrollment; update both machines together before reconnecting (#1638) — thanks @jkrogers9862!
|
||||
- Linux ARM64 (Asahi Apple Silicon) support for the OmniVoice GGUF engine — a `linux-aarch64` binary built with GGML Vulkan where the toolchain allows it, so Apple GPUs accelerate generation through the open-source Honeykrisp driver instead of falling back to CPU-only (#1641)
|
||||
- One-command install on every desktop OS: `curl -fsSL https://voicestudio.sh/install | sh` (macOS/Linux/WSL) or `irm https://voicestudio.sh/install | iex` (Windows) — the URL serves the right script per platform, and Windows gains a source installer (`scripts/install.ps1`) with a 3-OS CI smoke (#1626)
|
||||
- Per-line subtitle management in the dub table: a line's end time is editable alongside its start (typing a time and dragging its timeline edge now take the same path), lines merge with the previous row as well as the next (`Ctrl/Cmd+Shift+M`), and a new line can be inserted into the gap after any row (#1612) — thanks @invio-a11y!
|
||||
- CI now enforces performance regression budgets on the hot paths — operation-count tests pin streaming TTS to one synthesis per sentence and cached dub re-mixes to zero re-synthesis; fast-path guards cover zero re-decoding and ⌈N/W⌉ native batch calls when enabled (#1594)
|
||||
- Default-engine dubbing now synthesizes several segments per forward pass instead of one call per line — the width follows the host's device headroom (1 on CPU and low-VRAM cards, up to 8), `OMNIVOICE_DUB_BATCH_WIDTH` overrides it, and engines without native batching keep the single-segment path (#1594)
|
||||
- `/ws/tts` now reports real time-to-first-audio, and its RTF measures synthesis alone so a slow client can't inflate it (#1594)
|
||||
- The locally cached AudioSeal watermark generator warms on a background thread ~35s after boot (`OMNIVOICE_PRELOAD_WATERMARK=0` opts out; explicitly setting `=1` may download it), so the first synthesis no longer serializes the audioseal import + model load inline — measured at ~42s on a cold filesystem, 3s short of a 90s client timeout (#1576) — thanks @paoloantinori!
|
||||
- Voices you've cloned stay "warm" across restarts — encoded references now persist to disk (~10 KB each), so the first generation of a session skips the re-encode and any transcription pass; `OMNIVOICE_PROMPT_DISK_CACHE=0` opts out (#1565)
|
||||
- Optional FlashInfer acceleration for the default engine on CUDA (`OMNIVOICE_FLASHINFER=1`, ~2.2x measured) — needs the optional `flashinfer-python` package; missing package or kernel failure logs why and falls back to the standard path (#1565)
|
||||
- The bug reporter notices when you're on an outdated build and offers the latest release before filing — with a "File anyway" escape hatch — and stamps a `Build status` line into every report so up-to-date reports are tellable from stale ones (#1547)
|
||||
- Settings → Performance & Device gains a compute-device override (Auto / CUDA / ROCm / XPU / MPS / CPU, or `OMNIVOICE_DEVICE`) — pin the device when auto-detect picks wrong; only devices your machine actually has are offered (#1557)
|
||||
- Opt-in 24-layer PocketTTS checkpoints via `OMNIVOICE_POCKETTTS_24L` — better prosody for it/de/es/pt at roughly 2x render time (still faster than real-time); the fast 6-layer model stays the default (#1613) — thanks @paoloantinori!
|
||||
|
||||
### Docs
|
||||
- Supported-version and install guidance now identifies 0.5.1 as the stable desktop and container release (#1687)
|
||||
- The Docker Hub overview now shows the current engine-switching demo, Model Catalogue, and gallery voice workflow (#1593)
|
||||
- The Docker Hub overview and install guide now show the v0.5 tags and the built-in API-key/share-PIN security model instead of obsolete v0.4 and no-authentication guidance (#1592)
|
||||
- The READMEs now lead with download buttons and a three-step first-clone walkthrough, and a new benchmarks page anchors measured per-engine/per-device numbers on the in-repo harness (#1555)
|
||||
@@ -31,6 +76,42 @@ the frozen-backend fallback mirror it for their toolchains.
|
||||
- The OmniVoice guide now covers combining style attributes with a reference clip (consistent instruct stabilizes cloning; the reference wins conflicts), inline pronunciation control (pinyin / CMU phonemes), and corrects the claim that the default engine can't do voice design — it can, from attributes (#1565)
|
||||
|
||||
### Fixed
|
||||
- Workspaces now measure their responsive width when the post-bootstrap shell actually mounts, so native UI scaling reflows Projects and History instead of crushing the Dubbing demo into unreadable columns (#1683)
|
||||
- Dubbing keeps the source-language selector visible after a local file is chosen, so ASR can be pinned before transcription starts (#1678) — thanks @Lonki-lomki-cloud!
|
||||
- First-run media-engine downloads become available to TTS immediately without a restart, and missing media-process failures now point to repair controls (#1677) — thanks @farhataligpt-dev!
|
||||
- Source installs on AMD GPUs honour `OMNIVOICE_TORCH_VARIANT=rocm`: `bun run desktop` now swaps in the ROCm torch wheel after `uv sync` and launches the backend without re-syncing, instead of silently reverting to the CPU-only CUDA build on every start (#1665) — thanks @uberclokr!
|
||||
- `bun run desktop` on a fresh clone no longer fails with "resource path `../../frontend/dist` doesn't exist" — the dev launcher creates the placeholder Tauri resource directory before compiling (#1664) — thanks @uberclokr!
|
||||
- macOS no longer loses TTS after the first request when Python lacks `os.waitid`; subprocess ownership now uses a safe `waitpid` fallback without risking reused process groups (#1656) — thanks @paoloantinori!
|
||||
- Desktop startup, Retry, reset, uninstall, shutdown, and crash recovery now share one backend lifecycle owner; quitting interrupts first-run installers and gracefully drains then force-cleans the full backend process tree, so overlaps cannot duplicate or orphan it (#1635) — thanks @Xohaibxobi!
|
||||
- Large Stories and Audiobook projects now persist in IndexedDB instead of overflowing the `omnivoice.app` localStorage envelope, with quota-safe migration and orderly exit/reload flushing (#1636) — thanks @leodzai!
|
||||
- OmniVoice and its crash-isolated subprocess now route to AMD ROCm GPUs instead of warning and falling back to CPU (#1629) — thanks @j4r3kb!
|
||||
- Dictation now cancels pending startup work, capture resources, sockets, and timers when the capture widget closes, preventing late work against a destroyed webview (#1645)
|
||||
- Streaming generation failures now show recognized recovery guidance and appear in Diagnostics instead of only returning a generic error (#1607)
|
||||
- The worker-capacity transport test no longer races its own setup: the 1-slot limit now goes through the enrollment handshake instead of mutating client config after connect, where the server's stream-open ConfigUpdate (carrying the registered capacity of 2) could overwrite it and fake an over-accept; failed CI twice on 2026-08-21 (#1630)
|
||||
- Moving words across a speaker boundary in a dub — merging two lines and splitting them again — no longer dubs the second half in the first speaker's voice; each half now keeps the speaker, voice, direction, gain, and language of whoever actually says it (#1612) — thanks @invio-a11y!
|
||||
- Dictation on a WebView that refuses a 16 kHz audio context (WKWebView) now low-passes before downsampling, so frequencies above 8 kHz stop folding into the speech the recognizer is fed (#1610)
|
||||
- A microphone context that cannot be resumed now reports a mic error instead of leaving the dictation pill on "Listening" while capturing nothing (#1610)
|
||||
- Dictation no longer retains a whole session's audio for silent-model recovery — an open mic grew that buffer by ~115 MB an hour; the recent two minutes are kept instead (#1610)
|
||||
- The clipboard-delivery status is now translated in all 21 languages, so Wayland users — where clipboard delivery is the default — no longer see an English string (#1610)
|
||||
- A native sherpa-onnx load failure of any exception type now degrades to "engine unavailable" instead of taking the dictation WebSocket down (#1610)
|
||||
- Dictation now ships Whisper Tiny as its one cross-platform default, avoiding Parakeet's measured empty decoding on Windows while keeping Parakeet selectable behind runtime fallback (#1175)
|
||||
- Re-mixing a dub no longer decodes, rewrites, and re-reads every cached segment — same-rate cached audio is reused directly (and rejected if truncated), switching timing modes can't reuse slot-truncated audio as natural-rate, and RVC respects natural-rate modes (#1594)
|
||||
- PocketTTS French works again — pocket-tts only ships a 24-layer French model and rejected the name the sidecar asked for, so every French request failed at model load; French now always loads `french_24l` (#1613) — thanks @paoloantinori!
|
||||
- Installing IndexTTS 2.5 no longer fails claiming an interrupted download — the weights repo ships `config.yaml` and VoiceStudio demanded a `config_v2_5.yaml` that exists in no upstream release; both names are accepted, so a hand-renamed checkout keeps working (#1611) — thanks @zuiaiyutu!
|
||||
- IndexTTS 2.5 no longer has long-text generation killed at 60 seconds — the sidecar now proves it is alive every 5 seconds while `infer()` runs, and its deadline rises to 900s (`OMNIVOICE_INDEXTTS_RECV_TIMEOUT_S`) (#1611) — thanks @zuiaiyutu!
|
||||
- The OpenAI-compatible `/v1/audio/speech` route now reuses the shared cached engine for explicit `model` ids instead of constructing a fresh engine — and its sidecar/model load, a ~28s floor per call for subprocess engines — on every request, with the same single-engine-resident discipline `/generate` applies (#1614) — thanks @paoloantinori!
|
||||
- The setup wizard's RAM check no longer blocks 8 GB machines whose OS reports ~7.8 GB usable — the thresholds now tolerate reserved memory, and `OMNIVOICE_RAM_PREFLIGHT=0` turns a genuine block into a warning for those who accept the OOM risk (#1618)
|
||||
- Invisible watermarking now runs eagerly instead of through `torch.compile` — AudioSeal's lazy compile sent the first embed of every session into Inductor's C++ codegen, which failed outright on macOS hosts whose toolchain couldn't serve it and shipped the audio unmarked after a 30-40s wait; first embed drops from 9.70s to 0.26s (#1615) — thanks @paoloantinori!
|
||||
- The macOS Accessibility blocker now rechecks while visible and closes as soon as the grant is enabled instead of keeping a stale permission prompt on screen (#1609)
|
||||
- The dubbing editor's video and transcript columns can now be resized by pointer or keyboard, and the chosen split persists across launches (#1571) — thanks @invio-a11y!
|
||||
- CPU-only synthesis now gets a bounded ten-minute execution budget, and a render that exhausts it is reported as a compute timeout instead of misleading "generation capacity is busy" queue pressure (#1588) — thanks @ChienNguyen1111!
|
||||
- Rapid Launchpad ↔ Dub navigation now replaces the workspace DOM owner cleanly, so late media/waveform cleanup cannot trigger React's `insertBefore` crash (#1590) — thanks @nicolas-jacques!
|
||||
- Watermark embedding failures now log the full traceback instead of just the exception message, so a silently-unmarked-audio incident (audio passes through unmarked by design) is diagnosable from the log alone (#1576) — thanks @paoloantinori!
|
||||
- Dubbing now recovers rapid two-speaker exchanges when diarization collapses them, defaults new projects to lip sync without overwriting saved timing choices, and keeps the editor usable on narrow screens (#1584) — thanks @victordonat0!
|
||||
- `OMNIVOICE_ASR_BACKEND=omnivoice` now selects the PyTorch-native Whisper path, so the documented ROCm escape hatch no longer fails as an unknown engine (#1582) — thanks @patmansk!
|
||||
- Network Sharing from Windows MSI/portable installs now serves the bundled web interface to LAN devices instead of redirecting them to their own `localhost` (#1589) — thanks @TWIISTED-STUDIOS!
|
||||
- Exported dubbed videos now mark the dubbed language as the default audio stream while keeping Original available as an explicit choice (#1575) — thanks @invio-a11y!
|
||||
- Cloning references can no longer exhaust system memory: transcript-free clips up to 75 seconds are searched in five bounded passages, longer clips ask to be trimmed, and supplied transcripts remain capped at 20 seconds to preserve alignment (#1578) — thanks @ACKAPOB!
|
||||
- Stored artifact subpaths now resolve after moving a data directory between Windows, macOS, Linux, and Docker, while traversal and symlink escapes remain blocked (#1559) — thanks @Eman-Yousaf!
|
||||
- A remote browser hitting an API-key-configured server's admin 403 now gets the API-key login form instead of endless console 403s, while desktop and PIN-only/no-key servers keep the plain loopback error so guests are never offered a login no key can satisfy (#1568) — thanks @paoloantinori!
|
||||
- The crash-isolated ASR sidecar and its download preflight now agree on which model to load — setting the shared faster-whisper model variable applies to both variants instead of the sidecar quietly using a different one (#1556)
|
||||
@@ -38,6 +119,8 @@ the frozen-backend fallback mirror it for their toolchains.
|
||||
- Supervisor restarts after repeat crashes now back off (immediate, then 5s, then 15s) instead of respawning back-to-back, so a tight crash loop can't burn the whole restart budget in seconds (#1548)
|
||||
- The Linux desktop cleanup regression test now isolates build artifacts, so an existing developer build can no longer change its result (#1566)
|
||||
|
||||
- Renaming, deleting, or revoking consent on a voice (and starring/clearing history, recording exports) now live-updates every open tab again — the sync routes' WebSocket events were silently dropped, which could look like "all my voices are gone" (#1561) — thanks @paoloantinori!
|
||||
|
||||
### CI
|
||||
- Project agents now share pinned Vite and FastAPI skills from skills.sh (#1594)
|
||||
- Weekly full-history secret scans no longer mistake the Ed25519 private-key type name for committed key material (#1591)
|
||||
|
||||
@@ -60,7 +60,7 @@
|
||||
| macOS 13.3+ | DMG, Apple Silicon | [Install on macOS](docs/install/macos.md) |
|
||||
| Windows 10/11 | MSI, x64 | [Install on Windows](docs/install/windows.md) |
|
||||
| Linux | AppImage, x86_64 with glibc 2.39+ | [Install on Linux](docs/install/linux.md) |
|
||||
| Docker | CUDA, ROCm, or CPU | [Run with Docker](docs/install/docker.md) |
|
||||
| Docker | CUDA, ROCm, or CPU; worker-only GPU profiles | [Run with Docker](docs/install/docker.md) |
|
||||
|
||||
Download packages from the [latest release](https://github.com/debpalash/VoiceStudio/releases/latest). First launch creates a managed Python environment and downloads the default model. Later launches reuse both.
|
||||
|
||||
@@ -103,7 +103,7 @@ Use `bun run dev` for the browser UI. See [Contributing](.github/CONTRIBUTING.md
|
||||
| **Voice Design** | Create a voice from age, accent, pitch, style, and delivery instructions |
|
||||
| **Video Dubbing** | Transcribe, translate, preserve speakers, synthesize, and export video |
|
||||
| **Stories and audiobooks** | Multi-voice scripts · EPUB/PDF import · chapter rendering · `.m4b` export |
|
||||
| **Dictation Widget** | System-wide shortcut, live transcription, optional local-LLM cleanup |
|
||||
| **[Dictation Widget](docs/features/dictation.md)** | System-wide shortcut, live transcription, optional local-LLM cleanup |
|
||||
| **Vocal Isolation** | Demucs speech/background separation |
|
||||
| **Speaker Diarization** | Pyannote and WhisperX speaker assignment |
|
||||
| **Batch Queue** | Queue large sets of audio and video jobs with per-job progress |
|
||||
@@ -254,7 +254,7 @@ FastAPI backend
|
||||
|
||||
<a id="api"></a>
|
||||
|
||||
## OpenAI-compatible API
|
||||
## Local speech platform and OpenAI-compatible API
|
||||
|
||||
Point an OpenAI-compatible audio client at the local backend:
|
||||
|
||||
@@ -267,6 +267,8 @@ Point an OpenAI-compatible audio client at the local backend:
|
||||
|---|---|
|
||||
| `POST /v1/audio/speech` | TTS to `mp3`, `opus`, `aac`, `flac`, `wav`, or `pcm`; select a profile with `voice` and an engine with `model` |
|
||||
| `POST /v1/audio/transcriptions` | STT to `json`, `text`, `verbose_json`, `srt`, or `vtt` |
|
||||
| `WS /v1/audio/transcriptions/stream` | Live PCM/WebM transcription with partial, utterance, and session-final events |
|
||||
| `GET /.well-known/voicestudio-speech` | Discover HTTP, WebSocket, MCP, and native dictation-control transports |
|
||||
| `GET /v1/audio/voices` | List local voice profiles and engines |
|
||||
|
||||
```python
|
||||
@@ -283,7 +285,12 @@ with client.audio.speech.with_streaming_response.create(
|
||||
response.stream_to_file("speech.wav")
|
||||
```
|
||||
|
||||
The full API reference is in **Settings → OpenAPI Reference**. For LAN, Tailscale, or proxy access, read [API authentication](docs/api-auth.md) before exposing the backend.
|
||||
The bundled Rust control sidecar also lets Herdr, coding agents, VS Code,
|
||||
desktop apps, and TUIs trigger the existing system-wide dictation flow or reuse
|
||||
its safe native insertion. See the [speech platform guide](docs/speech-platform.md).
|
||||
The full API reference is in **Settings → OpenAPI Reference**. For LAN,
|
||||
Tailscale, or proxy access, read [API authentication](docs/api-auth.md) before
|
||||
exposing the backend.
|
||||
|
||||
### Agent skills
|
||||
|
||||
@@ -312,7 +319,7 @@ The [notebook](notebooks/OmniVoice_Studio_Colab.ipynb) runs the app and web UI o
|
||||
| Fix setup | [Troubleshooting](docs/install/troubleshooting.md) · [model downloads](docs/downloading-models.md) · [Hugging Face token](docs/setup/huggingface-token.md) |
|
||||
| Choose an engine | [Engine guides](docs/engines/README.md) · [benchmarks](docs/benchmarks.md) · [expressive speech](docs/expressive-speech.md) |
|
||||
| Tune hardware | [Performance](docs/performance.md) · [remote workers](docs/remote-workers.md) |
|
||||
| Build integrations | [API auth](docs/api-auth.md) · [MCP](docs/mcp.md) · [examples](examples/README.md) |
|
||||
| Build integrations | [Speech platform](docs/speech-platform.md) · [API auth](docs/api-auth.md) · [MCP](docs/mcp.md) · [examples](examples/README.md) |
|
||||
| Build VoiceStudio | [Contributing](.github/CONTRIBUTING.md) · [engine acceptance](docs/engine-acceptance.md) |
|
||||
| Track changes | [Changelog](CHANGELOG.md) · [roadmap](docs/ROADMAP.md) · [latest release](https://github.com/debpalash/VoiceStudio/releases/latest) |
|
||||
| Remove everything | [Uninstall guide](docs/install/uninstall.md) |
|
||||
|
||||
@@ -54,6 +54,15 @@ def _server_mode() -> bool:
|
||||
return os.environ.get("OMNIVOICE_SERVER_MODE", "").strip().lower() in _TRUTHY
|
||||
|
||||
|
||||
def validate_server_admin_key() -> None:
|
||||
"""Reject an explicitly blank key before a server-mode app starts."""
|
||||
raw_key = os.environ.get("OMNIVOICE_API_KEY")
|
||||
if _server_mode() and raw_key is not None and not raw_key.strip():
|
||||
raise RuntimeError(
|
||||
"OMNIVOICE_API_KEY is blank; configure a non-whitespace administrator key"
|
||||
)
|
||||
|
||||
|
||||
def _configured_pin(request) -> str | None:
|
||||
"""The active share PIN (``app.state.network_share.pin``) or None. Read via
|
||||
getattr so a bare Request stub (or a request that hit before lifespan set
|
||||
|
||||
@@ -330,7 +330,7 @@ async def _render_archetype_wav(a: dict, out_path: Path) -> None:
|
||||
# GPU pool and brick the backend (#730 class). Budget comes from the shared
|
||||
# length-scaled helper (#1190) instead of the flat 300s default.
|
||||
from services.model_manager import generate_timeout_s
|
||||
_budget = generate_timeout_s(text)
|
||||
_budget = generate_timeout_s(text, engine=model)
|
||||
audio_tensor = await run_on_gpu_pool_guarded(
|
||||
lambda: _infer(_PREVIEW_SEED), what="Archetype preview generate",
|
||||
timeout=_budget)
|
||||
@@ -357,15 +357,11 @@ async def _render_archetype_wav(a: dict, out_path: Path) -> None:
|
||||
# Runs on the dedicated watermark pool (#1190): AudioSeal embedding is CPU
|
||||
# work that holds no VRAM, so it must not occupy a GPU worker ahead of the
|
||||
# next generate on 1-worker hosts.
|
||||
from services.watermark import mark_synthetic
|
||||
from services.model_manager import get_watermark_pool
|
||||
import functools
|
||||
audio_tensor = await run_on_gpu_pool_guarded(
|
||||
functools.partial(mark_synthetic, audio_tensor, model.sampling_rate,
|
||||
context="archetypes.render"),
|
||||
what="Archetype watermark",
|
||||
from services.watermark import mark_synthetic_async
|
||||
audio_tensor = await mark_synthetic_async(
|
||||
audio_tensor, model.sampling_rate,
|
||||
context="archetypes.render",
|
||||
timeout=generate_timeout_s(""),
|
||||
executor=get_watermark_pool(),
|
||||
)
|
||||
|
||||
out_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
@@ -361,7 +361,7 @@ LONGFORM_NUM_STEP = 32
|
||||
LONGFORM_GUIDANCE_SCALE = 2.0
|
||||
|
||||
|
||||
def _seed_segment_rng(base_seed, text: str, nonce: int = 0) -> None:
|
||||
def _seed_segment_rng(base_seed, text: str, nonce: int = 0) -> int | None:
|
||||
"""Apply a profile's pinned seed to this synth call (#1139).
|
||||
|
||||
``_resolve_voice`` has always fetched the profile ``seed`` — but only the
|
||||
@@ -380,11 +380,13 @@ def _seed_segment_rng(base_seed, text: str, nonce: int = 0) -> None:
|
||||
must cover /generate and here together, not one path.
|
||||
"""
|
||||
if base_seed is None:
|
||||
return
|
||||
return None
|
||||
import torch
|
||||
|
||||
from services.audiobook import segment_seed
|
||||
torch.manual_seed(segment_seed(base_seed, text, nonce))
|
||||
seed = segment_seed(base_seed, text, nonce)
|
||||
torch.manual_seed(seed)
|
||||
return seed
|
||||
|
||||
|
||||
def _base_seed(opts: ExpressiveOptions, voice: dict):
|
||||
@@ -508,16 +510,21 @@ def _build_synth(
|
||||
"get_model": get_model, "language": language, "opts": opts}
|
||||
|
||||
backend = cls()
|
||||
extra = _generic_extra_kwargs(opts)
|
||||
native_proxy = bool(getattr(cls, "supports_native_omnivoice_controls", False))
|
||||
extra = (_omnivoice_sampling_kwargs(opts) if native_proxy
|
||||
else _generic_extra_kwargs(opts))
|
||||
next_nonce = _make_occ_counter(opts)
|
||||
|
||||
def synth(text, voice_id, speed=None):
|
||||
v = resolve(voice_id)
|
||||
_seed_segment_rng(_base_seed(opts, v), text, next_nonce())
|
||||
seed = _seed_segment_rng(_base_seed(opts, v), text, next_nonce())
|
||||
call_extra = dict(extra)
|
||||
if native_proxy and seed is not None:
|
||||
call_extra["seed"] = seed
|
||||
return backend.generate(
|
||||
text, language=language, ref_audio=v["ref_audio"],
|
||||
ref_text=v["ref_text"], instruct=v["instruct"], duration=None,
|
||||
speed=float(speed) if speed else 1.0, **extra,
|
||||
speed=float(speed) if speed else 1.0, **call_extra,
|
||||
)
|
||||
return {"mode": "generic", "resolve": resolve, "engine_id": engine_id,
|
||||
"synth": synth, "sample_rate": backend.sample_rate}
|
||||
|
||||
+188
-13
@@ -103,6 +103,75 @@ def _set_progress(job, stage, percent=0, **extra):
|
||||
job["progress"] = {"stage": stage, "percent": percent, **extra}
|
||||
|
||||
|
||||
#: Override for the native dub batch width. Set to 1 to disable batching.
|
||||
BATCH_WIDTH_ENV = "OMNIVOICE_DUB_BATCH_WIDTH"
|
||||
|
||||
#: Hard ceiling on the override — a batch this wide is already amortizing
|
||||
#: almost all of the per-call setup, and beyond it the failure mode is an OOM
|
||||
#: that costs more than the saving.
|
||||
_MAX_BATCH_WIDTH = 16
|
||||
|
||||
|
||||
def _native_batch_width(backend) -> int:
|
||||
"""How many segments to render in one native batch on THIS host.
|
||||
|
||||
A native batch widens the forward pass, so the width cannot be a constant.
|
||||
The default engine declares ``min_vram_gb = 6.0`` for a SINGLE job; an
|
||||
unconditional 8-wide batch would OOM the 4-8 GB CUDA cards and the MPS
|
||||
Macs where the per-segment path succeeds today — turning a throughput
|
||||
optimization into a regression on exactly the hardware that already
|
||||
struggles (#1616 is a 4 GB card reporting capacity failures). Default
|
||||
behaviour must not get riskier on a host, so the width is derived from
|
||||
measured headroom and falls back to 1 (no batching) when unknown.
|
||||
|
||||
CPU hosts get 1: batching there buys no kernel amortization and only
|
||||
multiplies peak RAM.
|
||||
"""
|
||||
override = os.environ.get(BATCH_WIDTH_ENV, "").strip()
|
||||
if override:
|
||||
try:
|
||||
return max(1, min(_MAX_BATCH_WIDTH, int(override)))
|
||||
except (TypeError, ValueError):
|
||||
logger.warning(
|
||||
"%s=%r is not an integer — deriving the batch width from the host instead.",
|
||||
BATCH_WIDTH_ENV, override,
|
||||
)
|
||||
try:
|
||||
from core.device_caps import detect_host_caps
|
||||
caps = detect_host_caps()
|
||||
except Exception: # noqa: BLE001 — an unprobeable host takes the safe path
|
||||
return 1
|
||||
if caps.family == "cpu" or not caps.vram_gb:
|
||||
return 1
|
||||
headroom = caps.vram_gb - float(getattr(backend, "min_vram_gb", 0.0) or 0.0)
|
||||
if headroom < 2.0:
|
||||
return 1
|
||||
if headroom < 6.0:
|
||||
return 2
|
||||
if headroom < 12.0:
|
||||
return 4
|
||||
return 8
|
||||
|
||||
|
||||
def _batch_timeout_s(texts: list[str], backend) -> float:
|
||||
"""Execution budget for one native batch.
|
||||
|
||||
Not the sum of the per-item budgets: ``generate_timeout_s`` returns a
|
||||
floor (300s GPU / 600s CPU) plus per-length overage, so summing it across
|
||||
eight items yields a ~2400s budget — and a wedged batch would hold a
|
||||
GPU-pool worker for forty minutes before the reset this file depends on
|
||||
(#730). One floor covers wedge detection for the whole call; only the
|
||||
length-driven overage is genuinely additive.
|
||||
"""
|
||||
from services.model_manager import generate_timeout_s
|
||||
|
||||
floor = generate_timeout_s("", engine=backend)
|
||||
overage = sum(
|
||||
max(0.0, generate_timeout_s(text, engine=backend) - floor) for text in texts
|
||||
)
|
||||
return floor + overage
|
||||
|
||||
|
||||
async def _run_batch_pipeline(job_id: str, job: dict):
|
||||
"""Full batch dub pipeline: extract → transcribe → translate → generate → mix → export."""
|
||||
import subprocess
|
||||
@@ -279,6 +348,111 @@ async def _run_batch_pipeline(job_id: str, job: dict):
|
||||
full_audio = torch.zeros(1, total_samples)
|
||||
total_segs = len(translated_segments)
|
||||
|
||||
# Native engines can amortize encoder/decoder setup across a small
|
||||
# batch. Keep the adapter seam optional: engines without a real batch
|
||||
# implementation inherit TTSBackend.generate_batch(), which preserves
|
||||
# the established one-segment behavior below.
|
||||
from services.tts_backend import TTSBackend
|
||||
batched_audio: dict[int, torch.Tensor] = {}
|
||||
has_native_batch = type(backend).generate_batch is not TTSBackend.generate_batch
|
||||
if has_native_batch:
|
||||
from services.text_normalization import normalize_for_tts
|
||||
|
||||
batch_ref_audio = None
|
||||
batch_ref_text = None
|
||||
if job.get("voice_id"):
|
||||
from core.db import db_conn
|
||||
from core.config import VOICES_DIR as _VD
|
||||
with db_conn() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT * FROM voice_profiles WHERE id=?",
|
||||
(job["voice_id"],),
|
||||
).fetchone()
|
||||
if row:
|
||||
if row["is_locked"] and row["locked_audio_path"]:
|
||||
batch_ref_audio = os.path.join(_VD, row["locked_audio_path"])
|
||||
elif row["ref_audio_path"]:
|
||||
batch_ref_audio = os.path.join(_VD, row["ref_audio_path"])
|
||||
batch_ref_text = row["ref_text"]
|
||||
|
||||
batch_width = _native_batch_width(backend)
|
||||
|
||||
async def _prefetch_batch(first_index: int) -> None:
|
||||
"""Render the batch beginning at ``first_index`` into
|
||||
``batched_audio``.
|
||||
|
||||
Rendered on demand rather than prerendering the whole track:
|
||||
the tensors are popped as they are placed, so peak host memory
|
||||
is one batch instead of every segment of the language — and
|
||||
the progress bar tracks placement instead of running to the
|
||||
end and restarting at segment 1.
|
||||
"""
|
||||
if job["status"] == "cancelled":
|
||||
return
|
||||
batch_rows = []
|
||||
index = first_index
|
||||
while index < total_segs and len(batch_rows) < batch_width:
|
||||
seg = translated_segments[index]
|
||||
if (seg.get("end", 0) - seg.get("start", 0) > 0.05
|
||||
and seg.get("text", "").strip()):
|
||||
batch_rows.append((index, seg))
|
||||
index += 1
|
||||
if len(batch_rows) < 2:
|
||||
return # nothing to amortize — the per-segment path is equal
|
||||
batch_indices = [index for index, _ in batch_rows]
|
||||
batch_texts = [
|
||||
normalize_for_tts(row.get("text", "").strip(), target_lang)
|
||||
for _, row in batch_rows
|
||||
]
|
||||
batch_durations = [
|
||||
row.get("end", 0) - row.get("start", 0)
|
||||
for _, row in batch_rows
|
||||
]
|
||||
|
||||
def _render_native_batch():
|
||||
generated = backend.generate_batch(
|
||||
batch_texts,
|
||||
language=target_lang,
|
||||
ref_audio=batch_ref_audio,
|
||||
ref_text=batch_ref_text,
|
||||
duration=batch_durations,
|
||||
num_step=16,
|
||||
guidance_scale=2.0,
|
||||
speed=1.0,
|
||||
denoise=True,
|
||||
postprocess_output=True,
|
||||
)
|
||||
if len(generated) != len(batch_indices):
|
||||
raise RuntimeError(
|
||||
f"native batch returned {len(generated)} outputs for "
|
||||
f"{len(batch_indices)} segments"
|
||||
)
|
||||
rendered = []
|
||||
for audio_out in generated:
|
||||
if not getattr(backend, "applies_own_mastering", False):
|
||||
audio_out = apply_mastering(audio_out, sample_rate=sr)
|
||||
rendered.append(normalize_audio(audio_out, target_dBFS=-2.0))
|
||||
return rendered
|
||||
|
||||
try:
|
||||
rendered = await run_on_gpu_pool_guarded(
|
||||
_render_native_batch,
|
||||
what="Batch generate",
|
||||
timeout=_batch_timeout_s(batch_texts, backend),
|
||||
)
|
||||
batched_audio.update(zip(batch_indices, rendered))
|
||||
except TimeoutError:
|
||||
# Do not immediately queue the same expensive work again:
|
||||
# the timed-out pool task may still be holding the device.
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Native TTS batch failed for segments %s-%s; falling back per segment: %s",
|
||||
batch_indices[0] + 1,
|
||||
batch_indices[-1] + 1,
|
||||
e,
|
||||
)
|
||||
|
||||
for i, seg in enumerate(translated_segments):
|
||||
if job["status"] == "cancelled":
|
||||
return
|
||||
@@ -356,10 +530,15 @@ async def _run_batch_pipeline(job_id: str, job: dict):
|
||||
# Budget is the shared length-scaled one (#1190): a long segment
|
||||
# on CPU-class hardware no longer dies on the flat 300s.
|
||||
from services.model_manager import generate_timeout_s
|
||||
audio_tensor = await run_on_gpu_pool_guarded(
|
||||
_gen, what="Batch generate",
|
||||
timeout=generate_timeout_s(seg_text),
|
||||
)
|
||||
if has_native_batch and i not in batched_audio:
|
||||
await _prefetch_batch(i)
|
||||
if i in batched_audio:
|
||||
audio_tensor = batched_audio.pop(i)
|
||||
else:
|
||||
audio_tensor = await run_on_gpu_pool_guarded(
|
||||
_gen, what="Batch generate",
|
||||
timeout=generate_timeout_s(seg_text, engine=backend),
|
||||
)
|
||||
|
||||
# Fit to slot
|
||||
target_samples_seg = int(seg_duration * sr)
|
||||
@@ -413,19 +592,15 @@ async def _run_batch_pipeline(job_id: str, job: dict):
|
||||
# unmarked while the interactive dub pipeline marked every segment.
|
||||
# One whole-track embed (chunked internally, #1045) is equivalent to
|
||||
# dub_generate's per-segment marks: the 16-bit message repeats
|
||||
# throughout. Runs in the GPU pool like generate's finalize; never
|
||||
# raises (degrades to unmarked on failure, same as every producer).
|
||||
# throughout. Never raises (degrades to unmarked on failure, same as
|
||||
# every producer).
|
||||
# Dispatched to the dedicated watermark pool, not the GPU pool (#1190):
|
||||
# AudioSeal embedding is CPU work that holds no VRAM, and a whole-track
|
||||
# embed is long enough that occupying a GPU worker with it stalled the
|
||||
# next language's segments on 1-worker hosts.
|
||||
from services.watermark import mark_synthetic
|
||||
from services.model_manager import get_watermark_pool
|
||||
import functools
|
||||
full_audio = await loop.run_in_executor(
|
||||
get_watermark_pool(),
|
||||
functools.partial(mark_synthetic, full_audio, sr,
|
||||
context="batch.dub_track"),
|
||||
from services.watermark import mark_synthetic_async
|
||||
full_audio = await mark_synthetic_async(
|
||||
full_audio, sr, context="batch.dub_track",
|
||||
)
|
||||
|
||||
# Same assembly pattern as dub_generate.py:390 — `full_audio` is a
|
||||
|
||||
@@ -27,6 +27,10 @@ Protocol:
|
||||
"detail": "..."} — error ("detail"
|
||||
kept for legacy)
|
||||
|
||||
Sherpa ``final`` frames additionally carry
|
||||
``"final_kind": "utterance"|"summary"``. Utterances are mid-session
|
||||
commits; the summary is the authoritative whole-session result at EOF.
|
||||
|
||||
Every ``final`` text is normalised by services.text_polish (leading
|
||||
capital for Latin scripts, terminal punctuation, single-spaced) so the
|
||||
pasted result reads like typed text. Partials are raw.
|
||||
@@ -34,10 +38,14 @@ Protocol:
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import tempfile
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, WebSocket, WebSocketDisconnect
|
||||
|
||||
@@ -47,6 +55,9 @@ from services.text_polish import polish_text
|
||||
router = APIRouter()
|
||||
logger = logging.getLogger("omnivoice.capture_ws")
|
||||
|
||||
SPEECH_PROTOCOL = "voicestudio.speech.v1"
|
||||
PLATFORM_STREAM_PATH = "/v1/audio/transcriptions/stream"
|
||||
|
||||
# How often (seconds) to run transcription on the accumulated buffer.
|
||||
# Shorter = more responsive but more GPU load.
|
||||
PARTIAL_INTERVAL_S = float(os.environ.get("OMNIVOICE_STREAM_INTERVAL", "2.0"))
|
||||
@@ -70,17 +81,79 @@ _AEC_NEAR = 0x00 # microphone frame (clean it, then buffer for ASR)
|
||||
_AEC_FAR = 0x01 # playback reference frame (feed the echo model only)
|
||||
|
||||
|
||||
def _requested_pcm_sample_rate(query_params) -> int | None:
|
||||
"""Return a bounded PCM rate for ``?pcm=1``/``?aec=1`` sessions."""
|
||||
raw_pcm = query_params.get("pcm") in ("1", "true", "on")
|
||||
aec = query_params.get("aec") in ("1", "true", "on")
|
||||
if not raw_pcm and not aec:
|
||||
return None
|
||||
# Client-supplied ``?sr=`` values outside the range real capture devices use
|
||||
# are replaced with 16 kHz. The rate sizes server-side state — RecoveryTail
|
||||
# multiplies it by RECOVERY_TAIL_SECONDS to compute its byte ceiling — so an
|
||||
# absurd rate must never be believed: it would re-open the unbounded-memory
|
||||
# path the recovery-tail cap closed.
|
||||
SR_MIN, SR_MAX = 8000, 96000
|
||||
|
||||
|
||||
def _is_end_control(text: str | None) -> bool:
|
||||
"""Accept the versioned JSON control frame and the legacy ``EOF`` frame."""
|
||||
if text == "EOF":
|
||||
return True
|
||||
if not text:
|
||||
return False
|
||||
try:
|
||||
message = json.loads(text)
|
||||
except (TypeError, json.JSONDecodeError):
|
||||
return False
|
||||
return isinstance(message, dict) and message.get("type") == "input_audio.end"
|
||||
|
||||
|
||||
class _PlatformWebSocket:
|
||||
"""Add v1 session metadata without changing the legacy WebSocket contract."""
|
||||
|
||||
def __init__(self, websocket: WebSocket):
|
||||
self._websocket = websocket
|
||||
self.session_id = uuid.uuid4().hex
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
return getattr(self._websocket, name)
|
||||
|
||||
async def send_json(self, data: Any, mode: str = "text") -> None:
|
||||
if isinstance(data, dict):
|
||||
data = dict(data)
|
||||
data.setdefault("protocol", SPEECH_PROTOCOL)
|
||||
data.setdefault("session_id", self.session_id)
|
||||
if data.get("type") == "final":
|
||||
data.setdefault("final_kind", "summary")
|
||||
await self._websocket.send_json(data, mode=mode)
|
||||
|
||||
|
||||
def _bounded_sample_rate(query_params) -> int:
|
||||
try:
|
||||
sample_rate = int(query_params.get("sr", "16000"))
|
||||
except (TypeError, ValueError):
|
||||
return 16000
|
||||
return sample_rate if 8000 <= sample_rate <= 96000 else 16000
|
||||
return sample_rate if SR_MIN <= sample_rate <= SR_MAX else 16000
|
||||
|
||||
|
||||
def _requested_pcm_sample_rate(query_params) -> int | None:
|
||||
"""Return the bounded rate when the client transport is raw PCM.
|
||||
|
||||
Sherpa clients omit ``pcm=1`` because the selected model already defines
|
||||
that transport. If the model is demoted or its runtime is unavailable, the
|
||||
legacy recognizer fallback must still decode those same bytes as PCM.
|
||||
"""
|
||||
raw_pcm = query_params.get("pcm") in ("1", "true", "on")
|
||||
aec = query_params.get("aec") in ("1", "true", "on")
|
||||
sherpa_pcm = False
|
||||
requested_model = query_params.get("model")
|
||||
if requested_model:
|
||||
try:
|
||||
from services.sherpa_dictation import is_sherpa_model
|
||||
sherpa_pcm = is_sherpa_model(requested_model)
|
||||
except Exception: # noqa: BLE001
|
||||
# A broken sherpa install must not decide the framing question —
|
||||
# sherpa_pcm stays False and the session negotiates the
|
||||
# MediaRecorder path; availability is re-probed (and reported)
|
||||
# when the model is actually selected.
|
||||
sherpa_pcm = False
|
||||
if not raw_pcm and not aec and not sherpa_pcm:
|
||||
return None
|
||||
return _bounded_sample_rate(query_params)
|
||||
|
||||
|
||||
def _demux_aec_frame(data: bytes) -> tuple[str, bytes]:
|
||||
@@ -137,21 +210,47 @@ def _select_sherpa_spec(websocket: WebSocket):
|
||||
from services import sherpa_dictation as sd
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def _usable_spec(model_id):
|
||||
spec = sd.get_spec(model_id)
|
||||
if spec is not None and sd.is_demoted(spec.id):
|
||||
logger.warning(
|
||||
"dictation model %s is demoted — using the capture ASR fallback",
|
||||
spec.id,
|
||||
)
|
||||
return None
|
||||
return spec
|
||||
|
||||
requested = websocket.query_params.get("model")
|
||||
if requested:
|
||||
return sd.get_spec(requested) # explicit selection (may be None if bad)
|
||||
return _usable_spec(requested) # explicit selection (may be unavailable)
|
||||
# Fall back to the persisted dictation pref.
|
||||
try:
|
||||
from services.asr_backend import dictation_model_id
|
||||
mid = dictation_model_id()
|
||||
except Exception:
|
||||
mid = None
|
||||
return sd.get_spec(mid) if mid else None
|
||||
return _usable_spec(mid) if mid else None
|
||||
|
||||
|
||||
@router.websocket(PLATFORM_STREAM_PATH)
|
||||
@router.websocket("/ws/transcribe")
|
||||
async def ws_transcribe(websocket: WebSocket):
|
||||
"""Stream audio in, get partial + final transcription out."""
|
||||
is_platform_stream = websocket.url.path == PLATFORM_STREAM_PATH
|
||||
if is_platform_stream:
|
||||
websocket = _PlatformWebSocket(websocket)
|
||||
# A browser can reach localhost regardless of the page's own origin.
|
||||
# Reject ambient cross-site WebSocket handshakes before the loopback-host
|
||||
# shortcut or accept(), while keeping native clients (no Origin header)
|
||||
# and configured/same-origin browser UIs working (#1646 review).
|
||||
origin = websocket.headers.get("origin")
|
||||
if origin:
|
||||
from core.csrf import origin_allowed
|
||||
|
||||
if not origin_allowed(websocket):
|
||||
await websocket.close(code=1008, reason="browser origin not allowed")
|
||||
return
|
||||
# Loopback origin guard — refuse anything not from 127.0.0.1, ::1, or
|
||||
# localhost. Privileged HTTP routers use Depends(require_admin) at router
|
||||
# level; WebSocket dependency injection differs across FastAPI versions, so we
|
||||
@@ -166,6 +265,16 @@ async def ws_transcribe(websocket: WebSocket):
|
||||
return
|
||||
|
||||
await websocket.accept()
|
||||
if is_platform_stream:
|
||||
await websocket.send_json({
|
||||
"type": "session.started",
|
||||
"input_format": (
|
||||
"audio/pcm;encoding=s16le;channels=1"
|
||||
if _requested_pcm_sample_rate(websocket.query_params) is not None
|
||||
else "audio/webm;codecs=opus"
|
||||
),
|
||||
"sample_rate": _bounded_sample_rate(websocket.query_params),
|
||||
})
|
||||
|
||||
# Live-dictation engine selection. When a sherpa-onnx model is selected
|
||||
# (via ?model= or the dictation.model_id pref) AND sherpa is installed,
|
||||
@@ -288,7 +397,7 @@ async def ws_transcribe(websocket: WebSocket):
|
||||
total_bytes += len(data)
|
||||
last_audio_time = time.monotonic()
|
||||
continue
|
||||
if msg.get("text") == "EOF":
|
||||
if _is_end_control(msg.get("text")):
|
||||
# Client signals end-of-audio but stays connected for `final`.
|
||||
running = False
|
||||
break
|
||||
@@ -422,6 +531,64 @@ SHERPA_OFFLINE_SILENCE_S = float(os.environ.get("OMNIVOICE_SHERPA_OFFLINE_SILENC
|
||||
SHERPA_OFFLINE_RMS_FLOOR = float(os.environ.get("OMNIVOICE_SHERPA_OFFLINE_RMS", "0.01"))
|
||||
|
||||
|
||||
#: Seconds of audio retained for silent-model recovery. Recovery only needs
|
||||
#: enough speech to prove the model is broken and to re-transcribe what was
|
||||
#: said; retaining the whole session grew ~115 MB/hour at 16 kHz on an open
|
||||
#: mic, unbounded, and only ever got read when the fallback fired.
|
||||
RECOVERY_TAIL_DEFAULT_SECONDS = 120.0
|
||||
RECOVERY_TAIL_MAX_SECONDS = 300.0
|
||||
|
||||
|
||||
def _bounded_recovery_tail_seconds(value: str | None) -> float:
|
||||
"""Parse the recovery tail override without allowing unbounded buffers."""
|
||||
try:
|
||||
seconds = float(value) if value is not None else RECOVERY_TAIL_DEFAULT_SECONDS
|
||||
except (TypeError, ValueError):
|
||||
return RECOVERY_TAIL_DEFAULT_SECONDS
|
||||
if not math.isfinite(seconds) or seconds <= 0:
|
||||
return RECOVERY_TAIL_DEFAULT_SECONDS
|
||||
return min(seconds, RECOVERY_TAIL_MAX_SECONDS)
|
||||
|
||||
|
||||
RECOVERY_TAIL_SECONDS = _bounded_recovery_tail_seconds(
|
||||
os.environ.get("OMNIVOICE_DICTATION_RECOVERY_TAIL_S")
|
||||
)
|
||||
|
||||
|
||||
class RecoveryTail:
|
||||
"""The most recent ``RECOVERY_TAIL_SECONDS`` of session audio.
|
||||
|
||||
Keeps the *tail* rather than the head: a long dictation's useful speech is
|
||||
what the user just said, and the silent-model check cares about how much
|
||||
audio the session carried overall — which ``total_bytes`` still reports
|
||||
truthfully after trimming.
|
||||
"""
|
||||
|
||||
__slots__ = ("_buf", "_max", "total_bytes")
|
||||
|
||||
def __init__(self, sample_rate: int, seconds: float = RECOVERY_TAIL_SECONDS):
|
||||
# int16 mono → 2 bytes/sample. Floor of one frame so a nonsense rate
|
||||
# or seconds value can't produce a zero-length buffer.
|
||||
self._max = max(2, int(seconds * max(1, sample_rate)) * 2)
|
||||
self._buf = bytearray()
|
||||
self.total_bytes = 0
|
||||
|
||||
def extend(self, pcm: bytes) -> None:
|
||||
self._buf.extend(pcm)
|
||||
self.total_bytes += len(pcm)
|
||||
excess = len(self._buf) - self._max
|
||||
if excess > 0:
|
||||
# int16 mono: trim whole samples only. A split frame can carry an
|
||||
# odd byte count, and an odd trim would leave the tail starting
|
||||
# mid-sample — every later sample byte-shifted, and the recovery
|
||||
# transcription fed noise.
|
||||
excess += excess % 2
|
||||
del self._buf[:excess]
|
||||
|
||||
def tail(self) -> bytes:
|
||||
return bytes(self._buf)
|
||||
|
||||
|
||||
def is_model_silent(text: str, heard_speech: bool, pcm_bytes: int) -> bool:
|
||||
"""True when the dictation model produced NO text despite real speech.
|
||||
|
||||
@@ -448,19 +615,74 @@ def _pcm16_to_f32(pcm: bytes):
|
||||
return np.frombuffer(pcm, dtype=np.int16).astype(np.float32) / 32768.0
|
||||
|
||||
|
||||
async def _sherpa_session(websocket: WebSocket):
|
||||
"""Shared WS receive setup for the sherpa handlers.
|
||||
def _pcm16_rms(pcm: bytes) -> float:
|
||||
samples = _pcm16_to_f32(pcm)
|
||||
if not len(samples):
|
||||
return 0.0
|
||||
return float((samples * samples).mean() ** 0.5)
|
||||
|
||||
Returns ``(get_frame, state)`` where ``get_frame`` is an async callable
|
||||
that yields the next near-end (mic) PCM bytes, ``b""`` for a keepalive/ref
|
||||
frame, or ``None`` on EOF/disconnect. ``state`` carries sample rate, AEC,
|
||||
and the disconnect flag for the caller's finaliser.
|
||||
"""
|
||||
pcm_sr = 16000
|
||||
|
||||
async def _recover_silent_sherpa(
|
||||
spec, pcm: bytes, pcm_sr: int,
|
||||
) -> tuple[str, list[dict]]:
|
||||
"""Retry a token-silent Sherpa session through an installed local ASR."""
|
||||
logger.warning(
|
||||
"dictation model %s decoded NOTHING from %.1fs of speech-level audio "
|
||||
"— falling back to the capture ASR engine for this session",
|
||||
spec.id, len(pcm) / float(max(1, pcm_sr) * 2),
|
||||
)
|
||||
try:
|
||||
pcm_sr = int(websocket.query_params.get("sr", "16000"))
|
||||
except (TypeError, ValueError):
|
||||
pcm_sr = 16000
|
||||
from services.asr_backend import asr_model_missing_error
|
||||
fallback_missing = await asyncio.to_thread(
|
||||
asr_model_missing_error,
|
||||
purpose="dictation",
|
||||
skip_sherpa=True,
|
||||
require_installed=True,
|
||||
)
|
||||
if fallback_missing is not None:
|
||||
logger.warning(
|
||||
"dictation silent-model fallback is not installed (%s); "
|
||||
"skipping recovery to avoid an automatic download",
|
||||
fallback_missing.get("missing_repo_id", "unknown"),
|
||||
)
|
||||
return "", []
|
||||
|
||||
result = await _transcribe_buffer_full(
|
||||
[pcm], pcm_sr=pcm_sr, skip_sherpa=True,
|
||||
)
|
||||
text = polish_text(_result_text(result))
|
||||
if not text:
|
||||
return "", []
|
||||
# The RMS gate can fire on fan/keyboard noise. Only another recognizer
|
||||
# producing words proves the audio held speech and makes persistent
|
||||
# demotion safe.
|
||||
try:
|
||||
from services.sherpa_dictation import demote_model
|
||||
if await asyncio.to_thread(demote_model, spec.id):
|
||||
logger.error(
|
||||
"dictation model %s demoted on this machine — it will no longer be "
|
||||
"auto-selected. Pick it again in Settings to give it another chance.",
|
||||
spec.id,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("silent-model demotion failed")
|
||||
segments = (result or {}).get("segments") or [
|
||||
{"start": 0.0, "end": None, "text": text}
|
||||
]
|
||||
return text, segments
|
||||
except Exception:
|
||||
logger.exception("dictation silent-model fallback failed")
|
||||
return "", []
|
||||
|
||||
|
||||
async def _sherpa_session(websocket: WebSocket):
|
||||
"""Shared WS setup for the sherpa handlers.
|
||||
|
||||
Returns ``(pcm_sr, aec)``: the bounded PCM sample rate for the session
|
||||
and the echo canceller when ``?aec=1`` requested one (``None`` otherwise
|
||||
or when AEC setup fails).
|
||||
"""
|
||||
pcm_sr = _bounded_sample_rate(websocket.query_params)
|
||||
aec = None
|
||||
if websocket.query_params.get("aec") in ("1", "true", "on"):
|
||||
try:
|
||||
@@ -498,7 +720,7 @@ async def _recv_pcm_frame(websocket: WebSocket, aec):
|
||||
return "skip", b""
|
||||
return "near", aec.process_near_end(payload)
|
||||
return "near", data
|
||||
if msg.get("text") == "EOF":
|
||||
if _is_end_control(msg.get("text")):
|
||||
return "eof", b""
|
||||
return "skip", b""
|
||||
|
||||
@@ -569,6 +791,8 @@ async def _run_sherpa_streaming(websocket: WebSocket, spec):
|
||||
|
||||
last_partial = ""
|
||||
committed: list[str] = [] # finalized utterances this session
|
||||
session_pcm = RecoveryTail(pcm_sr) # bounded audio for silent-model recovery
|
||||
heard_speech = False
|
||||
client_disconnected = False
|
||||
|
||||
async def _send(payload) -> bool:
|
||||
@@ -610,6 +834,9 @@ async def _run_sherpa_streaming(websocket: WebSocket, spec):
|
||||
break
|
||||
if kind == "skip":
|
||||
continue
|
||||
session_pcm.extend(pcm)
|
||||
if not heard_speech and _pcm16_rms(pcm) >= SHERPA_OFFLINE_RMS_FLOOR:
|
||||
heard_speech = True
|
||||
text, endpoint = await asyncio.to_thread(_decode_after_feed, pcm)
|
||||
if endpoint:
|
||||
# Commit this utterance (polished — it gets pasted); reset
|
||||
@@ -618,6 +845,7 @@ async def _run_sherpa_streaming(websocket: WebSocket, spec):
|
||||
if text:
|
||||
committed.append(text)
|
||||
await _send({"type": "final", "text": text,
|
||||
"final_kind": "utterance",
|
||||
"segments": [{"start": 0.0, "end": None, "text": text}],
|
||||
"language": "auto", "engine": backend.id})
|
||||
rec.reset(stream)
|
||||
@@ -644,7 +872,28 @@ async def _run_sherpa_streaming(websocket: WebSocket, spec):
|
||||
# Pieces are already polished; the join is too (polish is idempotent).
|
||||
full = " ".join(t for t in committed if t).strip()
|
||||
segments = [{"start": 0.0, "end": None, "text": t} for t in committed if t]
|
||||
|
||||
model_silent = is_model_silent(full, heard_speech, session_pcm.total_bytes)
|
||||
if model_silent:
|
||||
recovered, recovered_segments = await _recover_silent_sherpa(
|
||||
spec, session_pcm.tail(), pcm_sr,
|
||||
)
|
||||
if recovered:
|
||||
full = recovered
|
||||
segments = recovered_segments
|
||||
|
||||
if not client_disconnected:
|
||||
payload = {"type": "final", "text": full, "final_kind": "summary",
|
||||
"segments": segments,
|
||||
"language": "auto", "engine": backend.id}
|
||||
if model_silent:
|
||||
payload["engine"] = "capture-asr-fallback" if full else backend.id
|
||||
payload["model_silent"] = spec.id
|
||||
payload["warning"] = (
|
||||
f"The selected dictation model ({spec.id}) produced no text from your "
|
||||
"speech. Switched to the fallback engine for this session — pick a "
|
||||
"different model in Settings → Dictation."
|
||||
)
|
||||
if full:
|
||||
# Hard-bounded refinement (~4s): never delays this summary `final`
|
||||
# beyond OMNIVOICE_REFINE_TIMEOUT_S even with a dead LLM endpoint.
|
||||
@@ -653,14 +902,9 @@ async def _run_sherpa_streaming(websocket: WebSocket, spec):
|
||||
refined = await maybe_refine_async(full)
|
||||
except Exception:
|
||||
refined = None
|
||||
payload = {"type": "final", "text": full, "segments": segments,
|
||||
"language": "auto", "engine": backend.id}
|
||||
if refined and refined != full:
|
||||
payload["refined_text"] = refined
|
||||
await _send(payload)
|
||||
else:
|
||||
await _send({"type": "final", "text": "", "segments": [],
|
||||
"language": "auto", "engine": backend.id})
|
||||
await _send(payload)
|
||||
try:
|
||||
await websocket.close()
|
||||
except Exception:
|
||||
@@ -697,7 +941,7 @@ async def _run_sherpa_offline(websocket: WebSocket, spec):
|
||||
# whisper/zipformer transcribe the same bytes). Keep the whole session's
|
||||
# audio and whether any of it was speech-level, so the finaliser can tell
|
||||
# "user said nothing" (fine) from "model produced nothing" (broken).
|
||||
session_pcm = bytearray()
|
||||
session_pcm = RecoveryTail(pcm_sr)
|
||||
heard_speech = False
|
||||
running = True
|
||||
client_disconnected = False
|
||||
@@ -716,12 +960,6 @@ async def _run_sherpa_offline(websocket: WebSocket, spec):
|
||||
client_disconnected = True
|
||||
return False
|
||||
|
||||
def _rms(pcm: bytes) -> float:
|
||||
samples = _pcm16_to_f32(pcm)
|
||||
if not len(samples):
|
||||
return 0.0
|
||||
return float((samples * samples).mean() ** 0.5)
|
||||
|
||||
def _decode_window(pcm: bytes) -> str:
|
||||
samples = _pcm16_to_f32(pcm)
|
||||
if not len(samples):
|
||||
@@ -740,7 +978,7 @@ async def _run_sherpa_offline(websocket: WebSocket, spec):
|
||||
continue
|
||||
buf.extend(pcm)
|
||||
session_pcm.extend(pcm)
|
||||
if not heard_speech and _rms(pcm) >= SHERPA_OFFLINE_RMS_FLOOR:
|
||||
if not heard_speech and _pcm16_rms(pcm) >= SHERPA_OFFLINE_RMS_FLOOR:
|
||||
heard_speech = True
|
||||
last_audio = time.monotonic()
|
||||
except WebSocketDisconnect:
|
||||
@@ -766,6 +1004,7 @@ async def _run_sherpa_offline(websocket: WebSocket, spec):
|
||||
if text:
|
||||
committed.append(text)
|
||||
await _send({"type": "final", "text": text,
|
||||
"final_kind": "utterance",
|
||||
"segments": [{"start": 0.0, "end": None, "text": text}],
|
||||
"language": "auto", "engine": backend.id})
|
||||
|
||||
@@ -777,8 +1016,8 @@ async def _run_sherpa_offline(websocket: WebSocket, spec):
|
||||
continue
|
||||
snapshot = bytes(buf)
|
||||
if len(snapshot) > sil_bytes and \
|
||||
_rms(snapshot[-sil_bytes:]) < SHERPA_OFFLINE_RMS_FLOOR:
|
||||
if _rms(snapshot[:-sil_bytes]) >= SHERPA_OFFLINE_RMS_FLOOR:
|
||||
_pcm16_rms(snapshot[-sil_bytes:]) < SHERPA_OFFLINE_RMS_FLOOR:
|
||||
if _pcm16_rms(snapshot[:-sil_bytes]) >= SHERPA_OFFLINE_RMS_FLOOR:
|
||||
await _commit(snapshot)
|
||||
else:
|
||||
# Pure silence — drop it (keep the gate window for
|
||||
@@ -824,39 +1063,18 @@ async def _run_sherpa_offline(websocket: WebSocket, spec):
|
||||
# quiet user — hand the session to the capture ASR backend so the user
|
||||
# still gets their words, and say which model let them down. Bounded to
|
||||
# this session; the pref is left alone so the user stays in control.
|
||||
model_silent = is_model_silent(full, heard_speech, len(session_pcm))
|
||||
model_silent = is_model_silent(full, heard_speech, session_pcm.total_bytes)
|
||||
if model_silent:
|
||||
logger.warning(
|
||||
"dictation model %s decoded NOTHING from %.1fs of speech-level audio "
|
||||
"— falling back to the capture ASR engine for this session",
|
||||
spec.id, len(session_pcm) / float(max(1, pcm_sr) * 2),
|
||||
recovered, recovered_segments = await _recover_silent_sherpa(
|
||||
spec, session_pcm.tail(), pcm_sr,
|
||||
)
|
||||
# Demote it so the NEXT session doesn't repeat this round trip. The
|
||||
# curated default can be broken on a platform we never tested (the
|
||||
# NeMo-TDT decoder is, on Windows), and observing it beats guessing.
|
||||
try:
|
||||
from services.sherpa_dictation import demote_model
|
||||
if demote_model(spec.id):
|
||||
logger.error(
|
||||
"dictation model %s demoted on this machine — it will no longer be "
|
||||
"auto-selected. Pick it again in Settings to give it another chance.",
|
||||
spec.id,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("silent-model demotion failed")
|
||||
try:
|
||||
result = await _transcribe_buffer_full([bytes(session_pcm)], pcm_sr=pcm_sr)
|
||||
fb_text = polish_text((result or {}).get("text", "") or "")
|
||||
if fb_text:
|
||||
full = fb_text
|
||||
segments = (result or {}).get("segments") or [
|
||||
{"start": 0.0, "end": None, "text": fb_text}
|
||||
]
|
||||
except Exception:
|
||||
logger.exception("dictation silent-model fallback failed")
|
||||
if recovered:
|
||||
full = recovered
|
||||
segments = recovered_segments
|
||||
|
||||
if not client_disconnected:
|
||||
payload = {"type": "final", "text": full, "segments": segments,
|
||||
payload = {"type": "final", "text": full, "final_kind": "summary",
|
||||
"segments": segments,
|
||||
"language": "auto", "engine": backend.id}
|
||||
if model_silent:
|
||||
# The client surfaces this so a silently-broken model can't look
|
||||
@@ -884,6 +1102,35 @@ async def _run_sherpa_offline(websocket: WebSocket, spec):
|
||||
pass
|
||||
|
||||
|
||||
def _result_text(result: dict | None) -> str:
|
||||
"""Normalize text from every ASR backend result shape.
|
||||
|
||||
Some backends return a top-level ``text`` value, while WhisperX, Faster
|
||||
Whisper, Moonshine, and OpenAI-compatible ASR expose only ``segments`` and
|
||||
``chunks``. Dictation partials and finals must interpret both contracts the
|
||||
same way.
|
||||
"""
|
||||
if not isinstance(result, dict):
|
||||
return ""
|
||||
|
||||
text = result.get("text")
|
||||
if isinstance(text, str) and text.strip():
|
||||
return text.strip()
|
||||
|
||||
for key in ("segments", "chunks"):
|
||||
items = result.get(key)
|
||||
if not isinstance(items, (list, tuple)):
|
||||
continue
|
||||
text = " ".join(
|
||||
str(item.get("text", "")).strip()
|
||||
for item in items
|
||||
if isinstance(item, dict) and item.get("text")
|
||||
).strip()
|
||||
if text:
|
||||
return text
|
||||
return ""
|
||||
|
||||
|
||||
async def _transcribe_buffer(chunks: list[bytes], *, pcm_sr: int | None = None) -> str:
|
||||
"""Quick partial transcription of the current audio buffer."""
|
||||
|
||||
@@ -898,7 +1145,7 @@ async def _transcribe_buffer(chunks: list[bytes], *, pcm_sr: int | None = None)
|
||||
def _run():
|
||||
backend = get_capture_asr_backend()
|
||||
result = backend.transcribe(tmp, word_timestamps=False)
|
||||
return result.get("text", "")
|
||||
return _result_text(result)
|
||||
|
||||
# Bound dictation transcribes (#730): a wedged whisperx/CTranslate2 call
|
||||
# must not hold its GPU-pool worker forever and starve TTS / other ASR
|
||||
@@ -912,7 +1159,9 @@ async def _transcribe_buffer(chunks: list[bytes], *, pcm_sr: int | None = None)
|
||||
pass
|
||||
|
||||
|
||||
async def _transcribe_buffer_full(chunks: list[bytes], *, pcm_sr: int | None = None) -> dict:
|
||||
async def _transcribe_buffer_full(
|
||||
chunks: list[bytes], *, pcm_sr: int | None = None, skip_sherpa: bool = False,
|
||||
) -> dict:
|
||||
"""Full transcription with timing info for the final result."""
|
||||
tmp = _pcm16_to_wav(b"".join(chunks), pcm_sr) if pcm_sr else _chunks_to_wav(chunks)
|
||||
if tmp is None:
|
||||
@@ -924,15 +1173,13 @@ async def _transcribe_buffer_full(chunks: list[bytes], *, pcm_sr: int | None = N
|
||||
from services.asr_backend import get_capture_asr_backend, run_transcribe_guarded
|
||||
|
||||
def _run():
|
||||
backend = get_capture_asr_backend()
|
||||
backend = get_capture_asr_backend(skip_sherpa=skip_sherpa)
|
||||
t0 = time.perf_counter()
|
||||
result = backend.transcribe(tmp, word_timestamps=False)
|
||||
elapsed = round(time.perf_counter() - t0, 2)
|
||||
|
||||
segments = result.get("segments", [])
|
||||
full_text = result.get("text", "")
|
||||
if not full_text and segments:
|
||||
full_text = " ".join(s.get("text", "") for s in segments).strip()
|
||||
full_text = _result_text(result)
|
||||
|
||||
# Wave 1.1: strip Whisper hallucination loops from the final
|
||||
# text (the string that gets auto-pasted). Segments keep the
|
||||
|
||||
+326
-34
@@ -134,6 +134,82 @@ _save_job = dub_pipeline.save_job
|
||||
# paste (or a mis-aimed binary) burn CPU in the parser.
|
||||
_MAX_SUBTITLE_PASTE_CHARS = 2_000_000
|
||||
|
||||
_SRT_REPLACED_FIELDS = {
|
||||
"id",
|
||||
"start",
|
||||
"end",
|
||||
"text",
|
||||
"text_original",
|
||||
"translations",
|
||||
"translate_error",
|
||||
"translate_degraded",
|
||||
}
|
||||
|
||||
|
||||
def _best_overlapping_segment(cue: dict, existing: list[dict]) -> dict | None:
|
||||
"""Return the prior segment with the strongest temporal overlap."""
|
||||
cue_start = float(cue.get("start") or 0.0)
|
||||
cue_end = float(cue.get("end") or cue_start)
|
||||
cue_mid = (cue_start + cue_end) / 2.0
|
||||
best = None
|
||||
best_key = None
|
||||
for index, segment in enumerate(existing):
|
||||
start = float(segment.get("start") or 0.0)
|
||||
end = float(segment.get("end") or start)
|
||||
overlap = min(cue_end, end) - max(cue_start, start)
|
||||
if overlap <= 0:
|
||||
continue
|
||||
midpoint_distance = abs(cue_mid - ((start + end) / 2.0))
|
||||
key = (overlap, -midpoint_distance, -index)
|
||||
if best_key is None or key > best_key:
|
||||
best = segment
|
||||
best_key = key
|
||||
return best
|
||||
|
||||
|
||||
def _carry_srt_voice_metadata(
|
||||
cues: list[dict],
|
||||
existing: list[dict],
|
||||
segment_clones: dict | None,
|
||||
speaker_clones: dict | None = None,
|
||||
) -> tuple[list[dict], dict]:
|
||||
"""Replace subtitle content while retaining the source cast assignment."""
|
||||
source_clones = dict(segment_clones or {})
|
||||
source_speaker_clones = dict(speaker_clones or {})
|
||||
# Replacement cues get new positional ids. Starting from the old map would
|
||||
# let an unmatched cue whose new id happens to equal an old id inherit an
|
||||
# unrelated reference. Only explicitly overlap-matched references survive.
|
||||
clones = {}
|
||||
merged_segments = []
|
||||
for new_id, cue in enumerate(cues):
|
||||
prior = _best_overlapping_segment(cue, existing)
|
||||
metadata = {
|
||||
key: value
|
||||
for key, value in (prior or {}).items()
|
||||
if key not in _SRT_REPLACED_FIELDS
|
||||
}
|
||||
merged = {
|
||||
**metadata,
|
||||
"id": new_id,
|
||||
"start": cue.get("start", 0.0),
|
||||
"end": cue.get("end", 0.0),
|
||||
"text": cue.get("text", ""),
|
||||
"text_original": cue.get("text", ""),
|
||||
}
|
||||
if not merged.get("speaker_id"):
|
||||
merged["speaker_id"] = cue.get("speaker_id") or "Speaker 1"
|
||||
if prior is not None:
|
||||
prior_id = str(prior.get("id", ""))
|
||||
clone = source_clones.get(prior_id)
|
||||
if clone is None:
|
||||
clone = source_speaker_clones.get(prior.get("speaker_id"))
|
||||
if clone is not None:
|
||||
clones[str(new_id)] = clone
|
||||
if merged.get("profile_id") == f"auto-seg:{prior_id}":
|
||||
merged["profile_id"] = f"auto-seg:{new_id}"
|
||||
merged_segments.append(merged)
|
||||
return merged_segments, clones
|
||||
|
||||
|
||||
@router.post("/dub/parse-subtitle-text")
|
||||
def dub_parse_subtitle_text(req: ParseSubtitleTextRequest):
|
||||
@@ -234,7 +310,32 @@ async def dub_import_srt(job_id: str, file: UploadFile = File(...)):
|
||||
else:
|
||||
segments = result.segments
|
||||
|
||||
prior_segments = [
|
||||
segment for segment in (job.get("segments") or []) if isinstance(segment, dict)
|
||||
]
|
||||
segments, segment_clones = _carry_srt_voice_metadata(
|
||||
segments,
|
||||
prior_segments,
|
||||
job.get("segment_clones"),
|
||||
job.get("speaker_clones"),
|
||||
)
|
||||
job["segments"] = segments
|
||||
job["segment_clones"] = segment_clones
|
||||
# A pooled speaker clone is keyed only by a display label. Replacement
|
||||
# cues can reuse that label without overlapping the original speaker, so
|
||||
# retain matched pooled references as segment-specific clones above and
|
||||
# drop the global map before rebuilding the cast.
|
||||
job["speaker_clones"] = {}
|
||||
if segment_clones:
|
||||
from services.speaker_clone import build_cast_sources
|
||||
|
||||
job["cast_sources"] = build_cast_sources(
|
||||
segments,
|
||||
None,
|
||||
segment_clones,
|
||||
)
|
||||
else:
|
||||
job.pop("cast_sources", None)
|
||||
# `source_lang` stays whatever the user (or the upload step) set; we
|
||||
# don't try to language-detect off the cue text — that's noisy and the
|
||||
# user usually knows what their .srt is.
|
||||
@@ -351,12 +452,13 @@ async def preview_upload(video: UploadFile = File(...)):
|
||||
safe_name = f"{uuid.uuid4().hex[:12]}"
|
||||
vid_path = os.path.join(PREVIEW_DIR, f"{safe_name}{ext}")
|
||||
wav_path = os.path.join(PREVIEW_DIR, f"{safe_name}.wav")
|
||||
|
||||
with open(vid_path, "wb") as f:
|
||||
f.write(await video.read())
|
||||
|
||||
has_audio = False
|
||||
if ext not in [".wav", ".mp3", ".m4a", ".aac"]:
|
||||
payload = await video.read()
|
||||
|
||||
def _write_and_extract() -> bool:
|
||||
with open(vid_path, "wb") as f:
|
||||
f.write(payload)
|
||||
if ext in {".wav", ".mp3", ".m4a", ".aac"}:
|
||||
return False
|
||||
try:
|
||||
ffmpeg_cmd = [
|
||||
find_ffmpeg(), "-y", "-i", vid_path,
|
||||
@@ -368,10 +470,16 @@ async def preview_upload(video: UploadFile = File(...)):
|
||||
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
|
||||
timeout=300,
|
||||
)
|
||||
has_audio = True
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning("FFmpeg extraction failed: %s", log_safe(e))
|
||||
pass
|
||||
return False
|
||||
|
||||
# File writes and ffmpeg are blocking operations. Keep them on the bounded
|
||||
# CPU pool so a large preview cannot stall unrelated API requests (#1667).
|
||||
has_audio = await asyncio.get_running_loop().run_in_executor(
|
||||
_cpu_pool, _write_and_extract
|
||||
)
|
||||
|
||||
return {
|
||||
"url": f"/preview/{safe_name}{ext}",
|
||||
@@ -410,12 +518,38 @@ _ingest_gen = dub_pipeline.ingest_pipeline
|
||||
#: container so a mislabelled video can't slip past the video-skipping branch.
|
||||
_AUDIO_EXTS = {".wav", ".mp3", ".m4a", ".aac", ".flac", ".ogg", ".opus", ".wma"}
|
||||
|
||||
# Source-language choices exposed by the first-party dub UI. Keeping this an
|
||||
# allow-list rejects language names and private-use BCP-47 tags before they are
|
||||
# persisted as ASR overrides. Values are normalized to lowercase below.
|
||||
_DUB_SOURCE_LANG_CODES = frozenset({
|
||||
"af", "sq", "am", "ar", "hy", "az", "eu", "be", "bn", "bs", "bg",
|
||||
"my", "ca", "cmn-hans", "cmn-hant", "hr", "cs", "da", "nl", "en",
|
||||
"et", "fi", "fr", "gl", "ka", "de", "el", "gu", "ht", "ha", "haw",
|
||||
"he", "hi", "hu", "is", "id", "it", "ja", "jw", "kn", "kk", "km",
|
||||
"ko", "ku", "ky", "lo", "la", "lv", "lt", "mk", "ms", "ml", "mt",
|
||||
"mi", "mr", "mn", "ne", "no", "ps", "fa", "pl", "pt", "pa", "ro",
|
||||
"ru", "sm", "gd", "sr", "sn", "sd", "si", "sk", "sl", "so", "es",
|
||||
"su", "sw", "sv", "tg", "ta", "te", "th", "tr", "uk", "ur", "uz",
|
||||
"vi", "cy", "xh", "yi", "yo", "zu",
|
||||
})
|
||||
|
||||
|
||||
def _source_lang_override(value: str | None) -> str | None:
|
||||
"""Normalize a user-selected source language; auto/und means detect."""
|
||||
code = (value or "").strip().lower()
|
||||
if code in {"", "auto", "und"}:
|
||||
return None
|
||||
if code not in _DUB_SOURCE_LANG_CODES:
|
||||
raise HTTPException(status_code=400, detail="Invalid source language code")
|
||||
return code
|
||||
|
||||
|
||||
@router.post("/dub/upload")
|
||||
async def dub_upload(
|
||||
video: UploadFile = File(...),
|
||||
job_id: Optional[str] = Form(None),
|
||||
input_type: str = Form("video"),
|
||||
source_lang: Optional[str] = Form(None),
|
||||
):
|
||||
"""Accept a media upload, write to disk, queue background prep task.
|
||||
|
||||
@@ -445,6 +579,7 @@ async def dub_upload(
|
||||
detail=f"Audio-only dubbing needs an audio file ({', '.join(sorted(_AUDIO_EXTS))}); got '{ext or 'no extension'}'.",
|
||||
)
|
||||
|
||||
source_lang_override = _source_lang_override(source_lang)
|
||||
os.makedirs(job_dir, exist_ok=True)
|
||||
|
||||
video_path = os.path.join(job_dir, f"original{ext}")
|
||||
@@ -456,7 +591,13 @@ async def dub_upload(
|
||||
await task_manager.add_task(
|
||||
task_id, "prep",
|
||||
_ingest_gen, job_id, job_dir,
|
||||
{"kind": "file", "path": video_path, "input_type": input_type}, filename,
|
||||
{
|
||||
"kind": "file",
|
||||
"path": video_path,
|
||||
"input_type": input_type,
|
||||
"source_lang": source_lang_override,
|
||||
},
|
||||
filename,
|
||||
)
|
||||
return JSONResponse(
|
||||
status_code=202,
|
||||
@@ -478,6 +619,7 @@ async def dub_ingest_url(req: DubIngestUrlRequest, request: Request):
|
||||
status_code=400,
|
||||
detail="URL must start with http:// or https://. Paste a full video link (e.g. https://youtube.com/watch?v=…) or drop a local file instead.",
|
||||
)
|
||||
source_lang_override = _source_lang_override(req.source_lang)
|
||||
|
||||
try:
|
||||
import yt_dlp # noqa: F401
|
||||
@@ -513,6 +655,7 @@ async def dub_ingest_url(req: DubIngestUrlRequest, request: Request):
|
||||
"fetch_subs": bool(req.fetch_subs),
|
||||
"sub_langs": req.sub_langs or None,
|
||||
"cookie_file": cookie_path,
|
||||
"source_lang": source_lang_override,
|
||||
}
|
||||
try:
|
||||
await task_manager.add_task(
|
||||
@@ -577,6 +720,118 @@ def _clamp_num_speakers(value) -> Optional[int]:
|
||||
return value if 1 <= value <= 20 else None
|
||||
|
||||
|
||||
def _recover_from_phrase_embeddings(
|
||||
diar_pipe,
|
||||
diarized_segments: list[dict],
|
||||
*,
|
||||
phrases: list[dict],
|
||||
requested_speakers: int | None,
|
||||
audio_target: str,
|
||||
segments: list[dict],
|
||||
words: list,
|
||||
):
|
||||
"""Recover rapid turns when pyannote collapses a two-speaker exchange.
|
||||
|
||||
Uses ASR phrase boundaries and the embedding/audio components already
|
||||
loaded by speaker-diarization-3.1. Weak or imbalanced clusters are rejected
|
||||
so ordinary single-speaker recordings remain untouched. Returns
|
||||
``(segments, separation)`` or ``None``.
|
||||
"""
|
||||
present = {
|
||||
str(seg.get("speaker_id")) for seg in diarized_segments
|
||||
if seg.get("speaker_id")
|
||||
}
|
||||
if len(present) > 1:
|
||||
return None
|
||||
usable_phrases = [
|
||||
phrase for phrase in phrases
|
||||
if phrase.get("text")
|
||||
and float(phrase.get("end", 0.0)) - float(phrase.get("start", 0.0)) >= 0.75
|
||||
]
|
||||
if len(usable_phrases) < 4:
|
||||
return None
|
||||
requested = int(requested_speakers) if requested_speakers else 2
|
||||
if requested != 2:
|
||||
return None
|
||||
embedding = getattr(diar_pipe, "_embedding", None)
|
||||
audio = getattr(diar_pipe, "_audio", None)
|
||||
if embedding is None or audio is None:
|
||||
return None
|
||||
try:
|
||||
import numpy as np
|
||||
from pyannote.core import Segment as _PyannoteSegment
|
||||
from sklearn.cluster import AgglomerativeClustering
|
||||
|
||||
vectors = []
|
||||
durations = []
|
||||
for phrase in usable_phrases:
|
||||
start, end = float(phrase["start"]), float(phrase["end"])
|
||||
duration = end - start
|
||||
waveform, _ = audio.crop(
|
||||
audio_target, _PyannoteSegment(start, end),
|
||||
duration=duration, mode="pad",
|
||||
)
|
||||
vector = np.asarray(embedding(waveform[None])).reshape(-1)
|
||||
if not np.isfinite(vector).all():
|
||||
return None
|
||||
vectors.append(vector)
|
||||
durations.append(duration)
|
||||
matrix = np.vstack(vectors)
|
||||
labels = np.asarray(AgglomerativeClustering(
|
||||
n_clusters=2, metric="cosine", linkage="average",
|
||||
).fit_predict(matrix))
|
||||
if len(set(labels.tolist())) != 2:
|
||||
return None
|
||||
|
||||
counts = [int(np.sum(labels == cluster)) for cluster in (0, 1)]
|
||||
cluster_durations = [
|
||||
float(sum(duration for duration, label in zip(durations, labels) if label == cluster))
|
||||
for cluster in (0, 1)
|
||||
]
|
||||
if min(counts) < 2 or min(cluster_durations) < 1.5:
|
||||
return None
|
||||
|
||||
normalized = matrix / np.maximum(np.linalg.norm(matrix, axis=1, keepdims=True), 1e-8)
|
||||
similarities = normalized @ normalized.T
|
||||
within, cross = [], []
|
||||
for left in range(len(labels)):
|
||||
for right in range(left + 1, len(labels)):
|
||||
target = within if labels[left] == labels[right] else cross
|
||||
target.append(float(similarities[left, right]))
|
||||
if not within or not cross:
|
||||
return None
|
||||
separation = float(np.mean(within) - np.mean(cross))
|
||||
min_separation = 0.12 if requested_speakers == 2 else 0.18
|
||||
if separation < min_separation:
|
||||
logger.info(
|
||||
"phrase-embedding speaker recovery rejected (separation=%.3f < %.3f)",
|
||||
separation, min_separation,
|
||||
)
|
||||
return None
|
||||
|
||||
speaker_map = {}
|
||||
turns = []
|
||||
for phrase, label in zip(usable_phrases, labels.tolist()):
|
||||
if label not in speaker_map:
|
||||
speaker_map[label] = f"Speaker {len(speaker_map) + 1}"
|
||||
turns.append({
|
||||
"start": float(phrase["start"]),
|
||||
"end": float(phrase["end"]),
|
||||
"speaker": speaker_map[label],
|
||||
})
|
||||
# Assignment mutates segment dictionaries. Work on copies so a recovery
|
||||
# rejected by the final two-speaker check cannot leak partial labels
|
||||
# into the ordinary pyannote result.
|
||||
assigned = assign_speakers_from_turns([dict(item) for item in segments], turns)
|
||||
recovered = resplit_segments_by_turns(assigned, words, turns)
|
||||
if len({item.get("speaker_id") for item in recovered if item.get("speaker_id")}) < 2:
|
||||
return None
|
||||
return recovered, separation
|
||||
except Exception:
|
||||
logger.exception("phrase-embedding speaker recovery failed")
|
||||
return None
|
||||
|
||||
|
||||
@router.get("/dub/transcribe-stream/{job_id}")
|
||||
async def dub_transcribe_stream(
|
||||
job_id: str,
|
||||
@@ -912,6 +1167,12 @@ async def dub_transcribe_stream(
|
||||
# Words (global-timeline) retained so diarization can re-split a segment
|
||||
# that spans two speakers' turns at the word boundary (#486).
|
||||
all_words: list = []
|
||||
# Preserve the ASR backend's natural phrase boundaries before
|
||||
# segment_transcript merges short neighboring phrases. Pyannote 3.1
|
||||
# occasionally collapses rapid exchanges into one dominant speaker; in
|
||||
# that narrow case these phrase spans give its own WeSpeaker embedding
|
||||
# model clean candidate utterances for a conservative recovery pass.
|
||||
asr_phrase_segments: list[dict] = []
|
||||
detected_lang = None
|
||||
next_seg_id = 0
|
||||
chunk_errors: list[str] = []
|
||||
@@ -981,21 +1242,13 @@ async def dub_transcribe_stream(
|
||||
"error_code": failure["code"],
|
||||
}
|
||||
|
||||
# Retry a failed/timed-out chunk once on a fresh pool before giving
|
||||
# up. Otherwise a transient wedge on the FIRST chunk (whisperx often
|
||||
# cold-loads its model there, the #730 hang) drops that whole window
|
||||
# and the transcript is "missing the beginning, only middle+end".
|
||||
# The retry reuses the same audio window, so a recovered chunk fills
|
||||
# the hole instead of leaving silent gaps.
|
||||
# Retry an ordinary completed failure once. A timed-out native call
|
||||
# is different: its thread is still executing and must not overlap
|
||||
# a retry against the same backend (#1669).
|
||||
part = None
|
||||
timed_out = False
|
||||
for _attempt in range(1, _CHUNK_TRANSCRIBE_ATTEMPTS + 1):
|
||||
# A wedged chunk gets the SAME guarded-timeout + pool-reset
|
||||
# semantics as the whole-file paths (#730/#851):
|
||||
# run_transcribe_guarded bounds the call, abandons the poisoned
|
||||
# pool so the retry (and any concurrent TTS work) gets a fresh
|
||||
# worker, and raises the actionable ASRTimeoutError. Run it as
|
||||
# a task and poll so we can keep yielding pings — the
|
||||
# EventSource connection drops without them.
|
||||
# Run as a task and poll so pings keep the EventSource alive.
|
||||
task = asyncio.ensure_future(run_transcribe_guarded(
|
||||
_gpu_pool, _transcribe_chunk,
|
||||
what=f"Dub chunk {i + 1}/{chunks_n}",
|
||||
@@ -1010,9 +1263,12 @@ async def dub_transcribe_stream(
|
||||
try:
|
||||
part = task.result()
|
||||
except ASRTimeoutError:
|
||||
# The guard already reset the pool; keep the actionable
|
||||
# message (it names the durable fixes, and — after repeated
|
||||
# timeouts — the crash-isolated engine escape hatch).
|
||||
# Python cannot kill an in-process native transcribe. Do
|
||||
# not swap pools and retry over the still-running call:
|
||||
# concurrent whisperx/CTranslate2 access caused the native
|
||||
# Windows access violation in #1669. Stop this transcript;
|
||||
# the worker remains honestly occupied until it exits.
|
||||
timed_out = True
|
||||
logger.error(
|
||||
"Transcribe chunk %d/%d timed out after %.0fs (attempt %d/%d, job=%s)",
|
||||
i + 1, chunks_n, transcribe_timeout_s, _attempt,
|
||||
@@ -1031,23 +1287,36 @@ async def dub_transcribe_stream(
|
||||
# error-part; the timeout path already reset the pool).
|
||||
if part is not None and not part.get("error"):
|
||||
break
|
||||
if _attempt < _CHUNK_TRANSCRIBE_ATTEMPTS:
|
||||
if timed_out:
|
||||
break
|
||||
if not timed_out and _attempt < _CHUNK_TRANSCRIBE_ATTEMPTS:
|
||||
logger.warning(
|
||||
"Retrying transcribe chunk %d/%d after failure/timeout (next attempt %d/%d, job=%s)",
|
||||
i + 1, chunks_n, _attempt + 1, _CHUNK_TRANSCRIBE_ATTEMPTS, log_safe(job_id),
|
||||
)
|
||||
# A completed exception did not wedge the worker. Resetting
|
||||
# the pool here leaked a healthy executor on every ordinary
|
||||
# decode failure; run_transcribe_guarded already resets the
|
||||
# pool on the only case that needs it: a real timeout.
|
||||
# A completed exception did not leave native work behind,
|
||||
# so retrying this same audio window is safe.
|
||||
if part.get("error"):
|
||||
chunk_errors.append(part["error"])
|
||||
if part.get("error_code"):
|
||||
chunk_error_codes.append(part["error_code"])
|
||||
logger.warning("Chunk %d/%d error: %s", i + 1, chunks_n, log_safe(part["error"]))
|
||||
if timed_out:
|
||||
break
|
||||
if detected_lang is None and part.get("language"):
|
||||
detected_lang = part["language"]
|
||||
asr_speaker_turns.extend(part.get("speaker_turns") or [])
|
||||
for _phrase in part.get("chunks", []) or []:
|
||||
_pts = _phrase.get("timestamp") or (None, None)
|
||||
_ptext = (_phrase.get("text") or "").strip()
|
||||
try:
|
||||
_ps, _pe = float(_pts[0]), float(_pts[1])
|
||||
except (TypeError, ValueError, IndexError):
|
||||
continue
|
||||
if _ptext and _pe > _ps:
|
||||
asr_phrase_segments.append({
|
||||
"start": _ps, "end": _pe, "text": _ptext,
|
||||
})
|
||||
chunk_segs = segment_transcript(part, duration=t1, scene_cuts=scene_cuts)
|
||||
# Same word source segment_transcript used (already global-timeline),
|
||||
# kept for the post-diarization speaker re-split (#486).
|
||||
@@ -1313,7 +1582,25 @@ async def dub_transcribe_stream(
|
||||
assigned = assign_speakers_from_diarization(all_segments, diar)
|
||||
# #486: split any segment that spans two speakers' turns at the
|
||||
# word boundary (single-speaker segments pass through unchanged).
|
||||
return resplit_segments_by_diarization(assigned, all_words, diar), None, "pyannote"
|
||||
resplit = resplit_segments_by_diarization(assigned, all_words, diar)
|
||||
recovered = _recover_from_phrase_embeddings(
|
||||
diar_pipe,
|
||||
resplit,
|
||||
phrases=asr_phrase_segments,
|
||||
requested_speakers=num_speakers,
|
||||
audio_target=asr_audio_target,
|
||||
segments=all_segments,
|
||||
words=all_words,
|
||||
)
|
||||
if recovered is not None:
|
||||
recovered_segments, separation = recovered
|
||||
logger.info(
|
||||
"Recovered rapid two-speaker exchange from ASR phrase embeddings "
|
||||
"(phrases=%d, separation=%.3f).",
|
||||
len(asr_phrase_segments), separation,
|
||||
)
|
||||
return recovered_segments, None, "phrase_embeddings"
|
||||
return resplit, None, "pyannote"
|
||||
except Exception as e:
|
||||
logger.exception("Diarization failed")
|
||||
# Inline ASR turns beat the silence-gap heuristic as a crash
|
||||
@@ -1522,7 +1809,9 @@ async def dub_transcribe_stream(
|
||||
except Exception as e:
|
||||
logger.warning("speaker_clone extraction skipped: %s", e)
|
||||
|
||||
job["source_lang"] = ((detected_lang or "en").split("_")[0][:2] or "en").lower()
|
||||
job["source_lang"] = job.get("source_lang_override") or (
|
||||
(detected_lang or "en").split("_")[0][:2] or "en"
|
||||
).lower()
|
||||
job["full_transcript"] = " ".join(s.get("text", "") for s in final_segs)
|
||||
_save_job(job_id, job)
|
||||
|
||||
@@ -1719,7 +2008,9 @@ async def dub_transcribe(job_id: str, num_speakers: Optional[int] = None):
|
||||
except Exception as e:
|
||||
logger.warning("Failed to unload ASR backend: %s", e)
|
||||
|
||||
job["source_lang"] = (detected_lang or "en").split("_")[0][:2].lower()
|
||||
job["source_lang"] = job.get("source_lang_override") or (
|
||||
(detected_lang or "en").split("_")[0][:2] or "en"
|
||||
).lower()
|
||||
|
||||
scene_cuts = job.get("scene_cuts") or []
|
||||
segments = segment_transcript(result, duration=job.get("duration", 0.0), scene_cuts=scene_cuts)
|
||||
@@ -1775,7 +2066,8 @@ async def dub_transcribe(job_id: str, num_speakers: Optional[int] = None):
|
||||
# Bound the whole-file transcribe (#730): a wedged whisperx/CTranslate2
|
||||
# call would otherwise hold its GPU-pool worker forever and starve
|
||||
# every other request into a "can't reach backend". run_transcribe_guarded
|
||||
# also resets the pool on timeout so capacity is restored.
|
||||
# leaves an unkillable native worker accounted for on timeout so a
|
||||
# retry cannot overlap it (#1669).
|
||||
segments_result = await run_transcribe_guarded(_gpu_pool, _transcribe, what="Dub")
|
||||
except asyncio.CancelledError:
|
||||
job["aborted"] = True
|
||||
|
||||
@@ -572,7 +572,7 @@ def _build_audio_export_cmd(
|
||||
async def dub_download(
|
||||
job_id: str,
|
||||
preserve_bg: bool = Query(True, description="Mix background noise into dubbed tracks"),
|
||||
default_track: str = Query("original"),
|
||||
default_track: str = Query("", description="Default audio track; omitted selects the first dubbed track"),
|
||||
include_tracks: str = Query("", description="Comma-separated list of tracks to include (e.g. 'original,de,es'). Empty = include all."),
|
||||
save_authorization: str = Header("", alias="X-VoiceStudio-Path-Authorization"),
|
||||
burn_subs: bool = Query(False, description="Burn subtitles into the video stream (forces re-encode). Uses dual-subtitle layout when dual=1."),
|
||||
@@ -607,6 +607,18 @@ async def dub_download(
|
||||
for key, value in filtered_tracks.items()
|
||||
}
|
||||
|
||||
# A dub export should play the dub without requiring player-specific track
|
||||
# selection. Keep ``original`` as an explicit opt-in, but when callers omit
|
||||
# the preference choose the first generated dub consistently (#1575).
|
||||
if (
|
||||
filtered_tracks
|
||||
and not (default_track == "original" and include_original)
|
||||
and default_track not in filtered_tracks
|
||||
):
|
||||
default_track = next(iter(filtered_tracks))
|
||||
elif not filtered_tracks and include_original:
|
||||
default_track = "original"
|
||||
|
||||
if not filtered_tracks and not include_original:
|
||||
raise HTTPException(status_code=400, detail="No tracks selected for export")
|
||||
|
||||
@@ -631,12 +643,17 @@ async def dub_download(
|
||||
fmt = (out_format or "m4a").lower()
|
||||
if fmt not in _AUDIO_FORMAT_CODECS:
|
||||
fmt = "m4a"
|
||||
# lang_code is already constrained to an existing track key, but
|
||||
# allowlist-sanitize it before it reaches the output path so a path
|
||||
# component can never carry separators/traversal (same pattern as
|
||||
# safe_name below).
|
||||
safe_lang = "".join(c for c in lang_code if c.isalnum() or c in "-_") or "track"
|
||||
out_path = os.path.join(exports_dir, f"dubbed_audio_{safe_lang}_{stamp}.{fmt}")
|
||||
# Keep route/job data out of the filesystem and logging trust boundary.
|
||||
# The selected format reaches the path only through literal branches.
|
||||
if fmt == "wav":
|
||||
output_name = f"dubbed_audio_{stamp}.wav"
|
||||
elif fmt == "mp3":
|
||||
output_name = f"dubbed_audio_{stamp}.mp3"
|
||||
elif fmt == "flac":
|
||||
output_name = f"dubbed_audio_{stamp}.flac"
|
||||
else:
|
||||
output_name = f"dubbed_audio_{stamp}.m4a"
|
||||
out_path = os.path.join(exports_dir, output_name)
|
||||
bg = _optional_dub_artifact(job.get("no_vocals_path"), job_id) if preserve_bg else None
|
||||
cmd = _build_audio_export_cmd(ffmpeg, track_info["path"], bg, out_path, fmt)
|
||||
try:
|
||||
@@ -654,15 +671,28 @@ async def dub_download(
|
||||
)
|
||||
if not os.path.exists(out_path) or os.path.getsize(out_path) == 0:
|
||||
raise HTTPException(status_code=500, detail="ffmpeg audio export produced no output file")
|
||||
logger.info("Dub audio export wrote %s (%d bytes)", out_path, os.path.getsize(out_path))
|
||||
logger.info("Dub audio export completed (%d bytes)", os.path.getsize(out_path))
|
||||
|
||||
base_name = os.path.splitext(job.get("filename", "output"))[0]
|
||||
safe_name = "".join(c for c in base_name if c.isalnum() or c in "-_ ").strip() or "output"
|
||||
dl_name = f"dubbed_{safe_name}_{safe_lang}_{stamp}.{fmt}"
|
||||
# Response metadata must not become a second path-like sink for job or
|
||||
# request data. Keep the user-selected format through explicit literal
|
||||
# branches; source names and language keys never enter the label.
|
||||
if fmt == "wav":
|
||||
dl_name = f"dubbed_audio_{stamp}.wav"
|
||||
elif fmt == "mp3":
|
||||
dl_name = f"dubbed_audio_{stamp}.mp3"
|
||||
elif fmt == "flac":
|
||||
dl_name = f"dubbed_audio_{stamp}.flac"
|
||||
else:
|
||||
dl_name = f"dubbed_audio_{stamp}.m4a"
|
||||
media_type = _MEDIA_TYPES.get(f".{fmt}", "audio/mp4")
|
||||
save_path = _consume_native_save(save_authorization)
|
||||
if save_path:
|
||||
return _native_save(out_path, save_path, dl_name, media_type=media_type)
|
||||
# Keep the request-derived download label out of the filesystem
|
||||
# trust boundary. It is response metadata, not a source or
|
||||
# destination path (CodeQL, #1575).
|
||||
result = _native_save(out_path, save_path, "dubbed_audio", media_type=media_type)
|
||||
result["display_name"] = dl_name
|
||||
return result
|
||||
return FileResponse(
|
||||
out_path, media_type=media_type,
|
||||
headers={"Content-Disposition": content_disposition(dl_name)},
|
||||
@@ -887,7 +917,10 @@ async def dub_download(
|
||||
if default_track == "original" and include_original:
|
||||
cmd += ["-disposition:a:0", "default"]
|
||||
else:
|
||||
target_idx = 0
|
||||
# A stale/missing language preference still means "play a dub", not
|
||||
# "silently fall back to the source". The first processed dub is the
|
||||
# deterministic fallback; ``original`` above remains explicit.
|
||||
target_idx = tracks_to_process[0]["stream_idx"] if tracks_to_process else 0
|
||||
for t in tracks_to_process:
|
||||
if t['lang_code'] == default_track:
|
||||
target_idx = t["stream_idx"]
|
||||
@@ -1566,7 +1599,10 @@ async def dub_download_audio(
|
||||
return _native_save(wav_path, save_path, dl_name, media_type="audio/wav")
|
||||
return FileResponse(
|
||||
wav_path, media_type="audio/wav",
|
||||
headers={"Content-Disposition": content_disposition(dl_name)},
|
||||
headers={
|
||||
"Cache-Control": "no-store",
|
||||
"Content-Disposition": content_disposition(dl_name),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import os
|
||||
import re
|
||||
import json
|
||||
import struct
|
||||
import logging
|
||||
import time
|
||||
import asyncio
|
||||
@@ -80,6 +81,62 @@ def _prepare_oom_retry(error: Exception, *, execution_target: str) -> bool:
|
||||
return True
|
||||
|
||||
|
||||
def _cached_payload_intact(path: str, info) -> bool:
|
||||
"""Cheap truth check on a cached WAV whose header we are about to trust.
|
||||
|
||||
The natural-rate fast path hands the mixer a PATH instead of decoded
|
||||
audio, so a cache whose header reads fine but whose payload is truncated
|
||||
would only fail later, during assembly — after the timing plan (Smart Fit,
|
||||
video stretch) had been computed from the header's frame count. The plan
|
||||
would then describe audio that no longer exists and the segment would be
|
||||
replaced by slot-length silence, leaving the persisted video plan and the
|
||||
rendered track disagreeing.
|
||||
|
||||
Comparing the declared frame count against the physical ``data`` chunk
|
||||
catches that without decoding: a truncated file cannot hold the samples
|
||||
its header claims. Anything failing here falls through to the decoding path, which
|
||||
already degrades to a warning plus silence. Formats with no fixed
|
||||
bits-per-sample (compressed caches) are left to the decoder as before.
|
||||
"""
|
||||
try:
|
||||
bits = int(getattr(info, "bits_per_sample", 0) or 0)
|
||||
frames = int(getattr(info, "num_frames", 0) or 0)
|
||||
channels = int(getattr(info, "num_channels", 0) or 0)
|
||||
if bits <= 0 or frames <= 0 or channels <= 0:
|
||||
# Undecidable metadata fails CLOSED (review on #1620): these caches
|
||||
# are PCM WAVs this module wrote itself, so anything else is
|
||||
# unexpected — and the decode path this falls through to handles
|
||||
# every format the fast path would have.
|
||||
return False
|
||||
payload = frames * channels * (bits // 8)
|
||||
if payload <= 0:
|
||||
return False
|
||||
|
||||
# A WAV may carry JUNK/LIST metadata before data, so its header is not
|
||||
# necessarily 44 bytes. Locate the data chunk instead of counting
|
||||
# metadata as audio; otherwise an extended header can mask truncation.
|
||||
file_size = os.path.getsize(path)
|
||||
with open(path, "rb") as wav:
|
||||
header = wav.read(12)
|
||||
if len(header) != 12 or header[:4] != b"RIFF" or header[8:12] != b"WAVE":
|
||||
return False
|
||||
offset = 12
|
||||
while offset + 8 <= file_size:
|
||||
wav.seek(offset)
|
||||
chunk_id = wav.read(4)
|
||||
chunk_size_raw = wav.read(4)
|
||||
if len(chunk_id) != 4 or len(chunk_size_raw) != 4:
|
||||
return False
|
||||
chunk_size = struct.unpack("<I", chunk_size_raw)[0]
|
||||
data_offset = offset + 8
|
||||
if chunk_id == b"data":
|
||||
return chunk_size >= payload and file_size >= data_offset + payload
|
||||
offset = data_offset + chunk_size + (chunk_size % 2)
|
||||
return False
|
||||
except Exception: # noqa: BLE001 — an unstattable cache is the decoder's problem
|
||||
return False
|
||||
|
||||
|
||||
def _underrun_min_rate() -> float:
|
||||
"""Floor for the underrun fill (audio slowed toward its slot, never below
|
||||
this rate). Default 0.85 stays natural-sounding; OMNIVOICE_UNDERRUN_MIN_RATE=1.0
|
||||
@@ -446,6 +503,18 @@ async def dub_generate(job_id: str, req: DubRequest):
|
||||
backend = await resolve_generation_backend(require_cloning=True)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
except Exception as e:
|
||||
from core.failure import is_gpu_oom
|
||||
|
||||
if not is_gpu_oom(e):
|
||||
raise
|
||||
from core.public_errors import public_exception_response
|
||||
|
||||
payload = public_exception_response(
|
||||
e,
|
||||
fallback="The TTS model could not be loaded.",
|
||||
)
|
||||
raise HTTPException(status_code=503, detail=payload["detail"]) from e
|
||||
|
||||
async def _stream(task_id):
|
||||
total = len(req.segments)
|
||||
@@ -593,11 +662,11 @@ async def dub_generate(job_id: str, req: DubRequest):
|
||||
voice_match = (req.voice_match or "per_line").lower()
|
||||
_consistent_ref_memo: dict = {}
|
||||
remote_audio: dict[int, str] = {}
|
||||
# Strategy-transition guard: smart_fit re-mixes the *natural-rate*
|
||||
# per-segment WAVs from disk. If the previous run used strict_slot,
|
||||
# the on-disk WAVs are slot-squeezed ("slotted") — reusing them would
|
||||
# double-compress. Force one full regen; afterwards seg_wav_kind is
|
||||
# "natural" and partial regen / fit-only re-mix (regen_only=[]) work.
|
||||
# Strategy-transition guard: concise, stretch_video and smart_fit all
|
||||
# re-mix *natural-rate* per-segment WAVs. If the previous run used
|
||||
# strict_slot, the on-disk WAVs are slot-squeezed ("slotted") — the
|
||||
# missing tails cannot be recovered by a re-mix. Force one full regen;
|
||||
# afterwards partial regen / fit-only re-mix (regen_only=[]) is safe.
|
||||
# Jobs predating this field have unknown kind → also regen once.
|
||||
# P1.3: the kind is per-track now (each language renders under its own
|
||||
# strategy); the flat job["seg_wav_kind"] is only consulted for jobs
|
||||
@@ -608,7 +677,7 @@ async def dub_generate(job_id: str, req: DubRequest):
|
||||
_wav_kind = (
|
||||
_kind_map.get(lang_code) if isinstance(_kind_map, dict) else job.get("seg_wav_kind")
|
||||
)
|
||||
if strategy == "smart_fit" and regen_only is not None and _wav_kind != "natural":
|
||||
if strategy != "strict_slot" and regen_only is not None and _wav_kind != "natural":
|
||||
regen_only = None
|
||||
# Manifest: stable segment id per current index. Per-segment WAVs are
|
||||
# named by stable id (dub_seg_path) so regen reuses the right audio after
|
||||
@@ -759,15 +828,38 @@ async def dub_generate(job_id: str, req: DubRequest):
|
||||
if os.path.exists(seg_wav_path):
|
||||
try:
|
||||
_t_cache_0 = time.perf_counter()
|
||||
# Natural-rate caches are already the exact assembly
|
||||
# input. Keep the durable path in the manifest so the
|
||||
# mixer decodes it once; the old path decoded here,
|
||||
# wrote an identical mix_<id> scratch WAV, then decoded
|
||||
# that copy again. Header-only inspection preserves
|
||||
# the resample fallback for caches made by an engine
|
||||
# with a different sample rate.
|
||||
if strategy != "strict_slot":
|
||||
try:
|
||||
cached_info = torchaudio.info(seg_wav_path)
|
||||
except Exception:
|
||||
cached_info = None
|
||||
if (
|
||||
cached_info is not None
|
||||
and int(cached_info.sample_rate) == int(backend.sample_rate)
|
||||
and _cached_payload_intact(seg_wav_path, cached_info)
|
||||
):
|
||||
all_segment_wavs.append(
|
||||
(seg.start, seg.end, seg_wav_path, backend.sample_rate)
|
||||
)
|
||||
sync_scores.append(getattr(seg, 'sync_ratio', None) or 1.0)
|
||||
_t_cache += time.perf_counter() - _t_cache_0
|
||||
continue
|
||||
|
||||
cached_wav, cached_sr = torchaudio.load(seg_wav_path)
|
||||
if cached_sr != backend.sample_rate:
|
||||
import torchaudio.functional as AF
|
||||
cached_wav = AF.resample(cached_wav, cached_sr, backend.sample_rate)
|
||||
# Pad/trim to slot — except smart_fit, whose mix
|
||||
# loop needs the natural-rate length to compute the
|
||||
# audio/video split (the seg_wav_kind guard above
|
||||
# guarantees these cached WAVs are natural-rate).
|
||||
if strategy != "smart_fit":
|
||||
# strict_slot persists slot-sized buffers. Every other
|
||||
# strategy consumes natural-rate audio and lets the mix
|
||||
# loop fit it to the current timeline.
|
||||
if strategy == "strict_slot":
|
||||
target_samples = int(seg_duration * backend.sample_rate)
|
||||
current_samples = cached_wav.shape[-1]
|
||||
if target_samples > current_samples:
|
||||
@@ -1091,7 +1183,7 @@ async def dub_generate(job_id: str, req: DubRequest):
|
||||
_num_step, req.guidance_scale, seg_speed, seg_profile, seg_effect_preset,
|
||||
),
|
||||
what="Dub generate",
|
||||
timeout=generate_timeout_s(seg.text),
|
||||
timeout=generate_timeout_s(seg.text, engine=backend),
|
||||
)
|
||||
_t_tts += time.perf_counter() - _t_tts_0
|
||||
|
||||
@@ -1164,12 +1256,15 @@ async def dub_generate(job_id: str, req: DubRequest):
|
||||
if rvc_sr == backend.sample_rate:
|
||||
audio_tensor = rvc_wav
|
||||
|
||||
target_samples = int(seg_duration * backend.sample_rate)
|
||||
current_samples = audio_tensor.shape[-1]
|
||||
if target_samples > current_samples:
|
||||
audio_tensor = torch.nn.functional.pad(audio_tensor, (0, target_samples - current_samples))
|
||||
elif current_samples > target_samples:
|
||||
audio_tensor = audio_tensor[..., :target_samples]
|
||||
if strategy == "strict_slot":
|
||||
target_samples = int(seg_duration * backend.sample_rate)
|
||||
current_samples = audio_tensor.shape[-1]
|
||||
if target_samples > current_samples:
|
||||
audio_tensor = torch.nn.functional.pad(
|
||||
audio_tensor, (0, target_samples - current_samples)
|
||||
)
|
||||
elif current_samples > target_samples:
|
||||
audio_tensor = audio_tensor[..., :target_samples]
|
||||
except Exception as e:
|
||||
yield f"data: {json.dumps({'type': 'warning', 'segment': i, 'message': f'RVC skipped: {str(e)[:120]}'})}\n\n"
|
||||
|
||||
@@ -1208,7 +1303,15 @@ async def dub_generate(job_id: str, req: DubRequest):
|
||||
pass
|
||||
_release_audio_tensors()
|
||||
except Exception as e:
|
||||
yield f"data: {json.dumps({'type': 'error', 'segment': i, 'error': str(e)})}\n\n"
|
||||
# A task-stream error bypasses the global exception handler.
|
||||
# Never publish engine exception text here: allocator errors
|
||||
# carry process tables and arbitrary failures can carry paths,
|
||||
# tokens, or source text. The shared helper enriches recognized
|
||||
# classes using VoiceStudio-owned constants only.
|
||||
from core.public_errors import stream_generation_failure
|
||||
|
||||
error_detail = stream_generation_failure(e)["detail"]
|
||||
yield f"data: {json.dumps({'type': 'error', 'segment': i, 'error': error_detail})}\n\n"
|
||||
sr = backend.sample_rate
|
||||
all_segment_wavs.append(_store_mix_wav(seg.start, seg.end, torch.zeros(1, max(0, int(seg_duration * sr))), sr, f"mix_{seg_id}"))
|
||||
sync_scores.append(1.0)
|
||||
@@ -1356,7 +1459,21 @@ async def dub_generate(job_id: str, req: DubRequest):
|
||||
seg_gain = getattr(seg_ref, "gain", None) if seg_ref is not None else None
|
||||
seg_gain = seg_gain if seg_gain is not None else 1.0
|
||||
seg_gain = max(0.0, min(2.0, seg_gain))
|
||||
wav = _load_entry_wav((start, end, wav_path, sr), sr)
|
||||
try:
|
||||
wav = _load_entry_wav((start, end, wav_path, sr), sr)
|
||||
except Exception as e:
|
||||
# A WAV header can be readable while its payload is
|
||||
# truncated. Direct cache reuse deliberately defers the
|
||||
# decode to assembly, so preserve the old recovery contract
|
||||
# here: warn and fill this slot with silence instead of
|
||||
# aborting the entire dub.
|
||||
warning = {
|
||||
"type": "warning",
|
||||
"segment": i,
|
||||
"message": f"cached seg lost, padding silence: {str(e)[:120]}",
|
||||
}
|
||||
yield f"data: {json.dumps(warning)}\n\n"
|
||||
wav = torch.zeros(1, max(0, int((end - start) * sr)))
|
||||
adjusted = wav * seg_gain
|
||||
if adjusted.ndim == 2 and adjusted.shape[0] > 1:
|
||||
adjusted = adjusted.mean(dim=0, keepdim=True)
|
||||
@@ -1798,7 +1915,7 @@ async def preview_segment(job_id: str, req: SegmentPreviewRequest):
|
||||
from services.model_manager import generate_timeout_s
|
||||
audio_tensor = await run_on_gpu_pool_guarded(
|
||||
_gen, what="Dub preview generate",
|
||||
timeout=generate_timeout_s(req.text),
|
||||
timeout=generate_timeout_s(req.text, engine=backend),
|
||||
)
|
||||
|
||||
sr = backend.sample_rate
|
||||
|
||||
@@ -8,6 +8,7 @@ import asyncio
|
||||
import tempfile
|
||||
import contextlib
|
||||
import logging
|
||||
import threading
|
||||
import traceback
|
||||
from typing import Optional
|
||||
from fastapi import APIRouter, File, Form, UploadFile, HTTPException
|
||||
@@ -32,6 +33,83 @@ router = APIRouter()
|
||||
logger = logging.getLogger("omnivoice.generate")
|
||||
|
||||
|
||||
class _TempReferenceLease:
|
||||
"""Delete a request-owned reference once every abandoned reader drains."""
|
||||
|
||||
def __init__(self, path: str):
|
||||
self.path = path
|
||||
self._lock = threading.Lock()
|
||||
self._active = 0
|
||||
self._request_done = False
|
||||
self._deleted = False
|
||||
|
||||
def acquire(self):
|
||||
with self._lock:
|
||||
if self._request_done:
|
||||
raise RuntimeError("reference lease acquired after request cleanup")
|
||||
self._active += 1
|
||||
once_lock = threading.Lock()
|
||||
released = False
|
||||
|
||||
def release() -> None:
|
||||
nonlocal released
|
||||
with once_lock:
|
||||
if released:
|
||||
return
|
||||
released = True
|
||||
self._release()
|
||||
|
||||
return release
|
||||
|
||||
def _release(self) -> None:
|
||||
delete = False
|
||||
with self._lock:
|
||||
self._active -= 1
|
||||
if self._active < 0:
|
||||
raise RuntimeError("reference lease released too many times")
|
||||
if self._request_done and self._active == 0 and not self._deleted:
|
||||
self._deleted = True
|
||||
delete = True
|
||||
if delete:
|
||||
with contextlib.suppress(OSError):
|
||||
os.remove(self.path)
|
||||
|
||||
def finish_request(self) -> None:
|
||||
delete = False
|
||||
with self._lock:
|
||||
self._request_done = True
|
||||
if self._active == 0 and not self._deleted:
|
||||
self._deleted = True
|
||||
delete = True
|
||||
if delete:
|
||||
with contextlib.suppress(OSError):
|
||||
os.remove(self.path)
|
||||
|
||||
|
||||
async def _run_with_reference_lease(lease, factory):
|
||||
"""Hold an ad-hoc reference through one local GPU-pool dispatch."""
|
||||
if lease is None:
|
||||
return await factory(None)
|
||||
release = lease.acquire()
|
||||
abandoned = False
|
||||
try:
|
||||
return await factory(release)
|
||||
except GpuPoolBusyError:
|
||||
# Busy means no job started; release now. The callback may already have
|
||||
# done so, and the lease token is deliberately idempotent.
|
||||
release()
|
||||
abandoned = True
|
||||
raise
|
||||
except (asyncio.CancelledError, GpuJobTimeoutError):
|
||||
# The guard owns release now: immediately for a queued cancellation,
|
||||
# or from the worker finalizer after an in-flight job drains.
|
||||
abandoned = True
|
||||
raise
|
||||
finally:
|
||||
if not abandoned:
|
||||
release()
|
||||
|
||||
|
||||
def _profile_instruct(row):
|
||||
"""Validator-safe instruct for a stored profile row.
|
||||
|
||||
@@ -381,6 +459,31 @@ def _is_timeout_failure(e) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _is_media_process_launch_failure(exc: BaseException) -> bool:
|
||||
"""Identify an ffmpeg/ffprobe launch ENOENT without guessing from a file name."""
|
||||
if not isinstance(exc, FileNotFoundError):
|
||||
return False
|
||||
|
||||
# A regular missing reference/model file may itself be named "ffmpeg".
|
||||
# Require the innermost raise site to be Python's process launcher so that
|
||||
# basename collisions keep the normal missing-file diagnosis (#1677).
|
||||
traceback_cursor = exc.__traceback__
|
||||
if traceback_cursor is None:
|
||||
return False
|
||||
while traceback_cursor.tb_next is not None:
|
||||
traceback_cursor = traceback_cursor.tb_next
|
||||
origin_module = traceback_cursor.tb_frame.f_globals.get("__name__", "")
|
||||
if origin_module != "subprocess" and not origin_module.startswith("asyncio."):
|
||||
return False
|
||||
|
||||
filename = getattr(exc, "filename", None)
|
||||
if not filename:
|
||||
return "[winerror 2]" in str(exc).lower()
|
||||
return os.path.basename(str(filename)).lower() in {
|
||||
"ffmpeg", "ffmpeg.exe", "ffprobe", "ffprobe.exe",
|
||||
}
|
||||
|
||||
|
||||
def _oom_friendly_reraise(e):
|
||||
"""Best-effort cache flush + the user-facing OOM hint shared by both
|
||||
inference paths."""
|
||||
@@ -405,6 +508,21 @@ def _oom_friendly_reraise(e):
|
||||
# that lost its +x bit) is NOT an OOM — don't send the user to the Flush
|
||||
# button; tell them what's actually wrong.
|
||||
es = str(e)
|
||||
# #1677: Windows CreateProcess reports a missing executable as a bare
|
||||
# ``FileNotFoundError: [WinError 2] ...`` with no filename, while POSIX
|
||||
# includes the missing ffmpeg/ffprobe name. The bundled-media downloader
|
||||
# now republishes PATH as soon as it finishes, but a failed/blocked
|
||||
# download still needs an actionable recovery rather than the unknown-
|
||||
# error dead end. Keep missing reference/model files on their own path.
|
||||
for _exc in _exception_chain(e):
|
||||
if _is_media_process_launch_failure(_exc):
|
||||
raise RuntimeError(
|
||||
"A required media program couldn't be launched. Open "
|
||||
"Settings → Audio tools and use "
|
||||
"Download/Repair for the media engine, then retry. If Audio "
|
||||
"tools is already ready, repair the selected TTS engine and "
|
||||
f"restart VoiceStudio. Underlying error: {_safe_exc_text(_exc)}"
|
||||
) from e
|
||||
if isinstance(e, PermissionError) or "Permission denied" in es or "Errno 13" in es:
|
||||
raise RuntimeError(
|
||||
f"A required engine binary couldn't be executed (permission denied). "
|
||||
@@ -596,7 +714,7 @@ def _oom_friendly_reraise(e):
|
||||
) from e
|
||||
|
||||
|
||||
def _generate_timeout_s(text: str) -> float:
|
||||
def _generate_timeout_s(text: str, *, execution_device=None) -> float:
|
||||
"""Wall-clock budget for one generate, scaled to the request.
|
||||
|
||||
Thin alias for the canonical helper, which moved to
|
||||
@@ -605,7 +723,7 @@ def _generate_timeout_s(text: str) -> float:
|
||||
as they did, silently keeping the flat 300s).
|
||||
"""
|
||||
from services.model_manager import generate_timeout_s
|
||||
return generate_timeout_s(text)
|
||||
return generate_timeout_s(text, execution_device=execution_device)
|
||||
|
||||
|
||||
def _run_inference(
|
||||
@@ -695,15 +813,17 @@ def _run_backend_inference(
|
||||
backend, text, language, ref_audio_path, ref_text, instruct, duration,
|
||||
num_step, guidance_scale, speed, denoise, postprocess_output,
|
||||
used_seed, effect_preset="broadcast",
|
||||
max_chunk_chars=None, crossfade_ms=None, *, dropped_sink=None,
|
||||
max_chunk_chars=None, crossfade_ms=None, *, t_shift=None,
|
||||
layer_penalty_factor=None, position_temperature=None,
|
||||
class_temperature=None, dropped_sink=None,
|
||||
):
|
||||
"""Engine-aware twin of :func:`_run_inference` (issue #312).
|
||||
|
||||
Runs the request through a pluggable ``TTSBackend`` adapter instead of the
|
||||
VoiceStudio model directly. The adapter protocol is narrower than the
|
||||
VoiceStudio-native surface — engine-specific extras (``t_shift``,
|
||||
``layer_penalty_factor``, …) only exist on the native path, which is why
|
||||
VoiceStudio itself still goes through ``_run_inference``.
|
||||
VoiceStudio model directly. A crash-isolated OmniVoice proxy advertises
|
||||
``supports_native_omnivoice_controls`` and receives the same advanced
|
||||
controls and per-call seed as the native path; other adapters keep the
|
||||
narrower protocol unchanged.
|
||||
"""
|
||||
import torch
|
||||
try:
|
||||
@@ -718,6 +838,18 @@ def _run_backend_inference(
|
||||
instruct=instruct, num_step=num_step, guidance_scale=guidance_scale,
|
||||
speed=speed, denoise=denoise, postprocess_output=postprocess_output,
|
||||
)
|
||||
native_proxy = bool(
|
||||
getattr(backend, "supports_native_omnivoice_controls", False)
|
||||
)
|
||||
if native_proxy:
|
||||
gen_kwargs.update({
|
||||
key: value for key, value in {
|
||||
"t_shift": t_shift,
|
||||
"layer_penalty_factor": layer_penalty_factor,
|
||||
"position_temperature": position_temperature,
|
||||
"class_temperature": class_temperature,
|
||||
}.items() if value is not None
|
||||
})
|
||||
sr = backend.sample_rate
|
||||
|
||||
# Inline [pause Nms] markers (issue #276) work for every engine — the
|
||||
@@ -727,10 +859,17 @@ def _run_backend_inference(
|
||||
has_pause = len(segments) > 1 or (segments and segments[0][1] > 0)
|
||||
|
||||
if has_pause:
|
||||
first_span = True
|
||||
|
||||
def _gen_span(span_text):
|
||||
nonlocal first_span
|
||||
# Per-span duration is left to the engine; an explicit overall
|
||||
# `duration` can't be meaningfully split across spans.
|
||||
return backend.generate(span_text, duration=None, **gen_kwargs)
|
||||
span_kwargs = dict(gen_kwargs)
|
||||
if native_proxy and first_span and used_seed is not None:
|
||||
span_kwargs["seed"] = used_seed
|
||||
first_span = False
|
||||
return backend.generate(span_text, duration=None, **span_kwargs)
|
||||
audio_out = _render_with_pauses(_gen_span, segments, sr)
|
||||
else:
|
||||
# Wave 1.2: sentence-boundary chunking for long text (see
|
||||
@@ -747,12 +886,19 @@ def _run_backend_inference(
|
||||
for i, chunk_text in enumerate(text_chunks):
|
||||
if used_seed is not None:
|
||||
torch.manual_seed(used_seed + i)
|
||||
parts.append(backend.generate(chunk_text, duration=None, **gen_kwargs))
|
||||
chunk_kwargs = dict(gen_kwargs)
|
||||
if native_proxy and used_seed is not None:
|
||||
chunk_kwargs["seed"] = used_seed + i
|
||||
parts.append(backend.generate(
|
||||
chunk_text, duration=None, **chunk_kwargs
|
||||
))
|
||||
_note_generate_progress()
|
||||
audio_out = concatenate_audio_chunks(parts, sr, _xfade_ms,
|
||||
texts=text_chunks,
|
||||
sink=dropped_sink)
|
||||
else:
|
||||
if native_proxy and used_seed is not None:
|
||||
gen_kwargs["seed"] = used_seed
|
||||
audio_out = backend.generate(text, duration=duration, **gen_kwargs)
|
||||
|
||||
return _apply_effect_chain(
|
||||
@@ -870,7 +1016,6 @@ async def _finalize_generation(
|
||||
Returns ``(watermarked_tensor, meta)`` where ``meta`` carries
|
||||
``id`` / ``filename`` / ``duration`` / ``gen_time``.
|
||||
"""
|
||||
loop = asyncio.get_running_loop()
|
||||
# Invisible AudioSeal provenance watermark on the final audio. Embedding
|
||||
# was previously only wired into the dub pipeline (dub_generate.py), so
|
||||
# plain TTS came out unmarked despite the setting being on — and the same
|
||||
@@ -882,12 +1027,9 @@ async def _finalize_generation(
|
||||
# AudioSeal embedding is CPU work that holds no VRAM, so occupying a GPU
|
||||
# worker with it only delays the next generate on 1-worker hosts.
|
||||
if not already_marked:
|
||||
from services.watermark import mark_synthetic
|
||||
from services.model_manager import get_watermark_pool
|
||||
audio_tensor = await loop.run_in_executor(
|
||||
get_watermark_pool(),
|
||||
functools.partial(mark_synthetic, audio_tensor, sample_rate,
|
||||
context="generate.finalize"),
|
||||
from services.watermark import mark_synthetic_async
|
||||
audio_tensor = await mark_synthetic_async(
|
||||
audio_tensor, sample_rate, context="generate.finalize",
|
||||
)
|
||||
gen_time = round(time.time() - start_time, 2)
|
||||
|
||||
@@ -1198,6 +1340,10 @@ async def generate_speech(
|
||||
_backend = None
|
||||
_engine_min_vram_gb = getattr(backend_cls, "min_vram_gb", 0.0)
|
||||
_routing_notice = None
|
||||
# Remote renders deliberately skip this host's capability gate. Keep the
|
||||
# local fallback call's timeout device-neutral so the closure is valid
|
||||
# without pretending the control plane describes the remote worker.
|
||||
_routing = {"effective_device": None}
|
||||
|
||||
if not _remote:
|
||||
# Single-active-engine memory discipline: hand back any OTHER resident
|
||||
@@ -1297,6 +1443,7 @@ async def generate_speech(
|
||||
|
||||
ref_audio_path = None
|
||||
cleanup_ref = False
|
||||
ref_lease = None
|
||||
used_seed = seed
|
||||
resolved_profile_id = None
|
||||
history_mode = None # profile.kind when a profile drives; else inferred at insert
|
||||
@@ -1384,6 +1531,7 @@ async def generate_speech(
|
||||
f.write(await ref_audio.read())
|
||||
ref_audio_path = f.name
|
||||
cleanup_ref = True
|
||||
ref_lease = _TempReferenceLease(ref_audio_path)
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
@@ -1400,13 +1548,19 @@ async def generate_speech(
|
||||
# built-in ASR fallback), so a timeout degrades to None rather than
|
||||
# failing the whole generate.
|
||||
try:
|
||||
ref_text = await run_on_gpu_pool_guarded(
|
||||
functools.partial(transcribe_reference, ref_audio_path),
|
||||
what="Reference transcribe",
|
||||
# Floor budget (#1190): a reference clip is seconds of audio,
|
||||
# so the length-scaled bonus never applies — but the timeout is
|
||||
# explicit here too, so no dispatch relies on a hidden default.
|
||||
timeout=_generate_timeout_s(""),
|
||||
ref_text = await _run_with_reference_lease(
|
||||
ref_lease,
|
||||
lambda release: run_on_gpu_pool_guarded(
|
||||
functools.partial(transcribe_reference, ref_audio_path),
|
||||
what="Reference transcribe",
|
||||
# Floor budget (#1190): a reference clip is seconds of audio,
|
||||
# so the length-scaled bonus never applies — but the timeout is
|
||||
# explicit here too, so no dispatch relies on a hidden default.
|
||||
timeout=_generate_timeout_s(
|
||||
"", execution_device=_routing["effective_device"]
|
||||
),
|
||||
on_abandon=release,
|
||||
)
|
||||
)
|
||||
# TimeoutError covers both the execution bound and pool saturation:
|
||||
# this path is best-effort either way.
|
||||
@@ -1523,7 +1677,7 @@ async def generate_speech(
|
||||
local=gpu_gateway.LocalCall(
|
||||
_remote_only_local_call(_target_label),
|
||||
what="TTS generate",
|
||||
timeout=_generate_timeout_s(text),
|
||||
timeout=_generate_timeout_s(text, execution_device=_routing["effective_device"]),
|
||||
min_vram_gb=_engine_min_vram_gb,
|
||||
),
|
||||
remote=_remote_call,
|
||||
@@ -1662,19 +1816,30 @@ async def generate_speech(
|
||||
"target_label": e.worker_label or _target_label,
|
||||
"hint": e.hint,
|
||||
})
|
||||
except Exception:
|
||||
except Exception as exc:
|
||||
# Mid-job remote failure is NOT quietly redone here: the client
|
||||
# treats a retryable error as "surface it", so the user decides
|
||||
# whether to spend the same minutes again on this machine.
|
||||
logger.error("Remote generation failed", exc_info=True)
|
||||
from core.public_errors import stream_failure
|
||||
yield _line({"type": "error", **stream_failure("generation_failed")})
|
||||
# whether to spend the same minutes again on this machine. Like
|
||||
# the local streaming path, this in-band frame stands in for the
|
||||
# global 500 handler, so it journals the scrubbed failure and
|
||||
# names a recognized cause instead of the bare generic string
|
||||
# (#1607).
|
||||
logger.error(
|
||||
"Remote generation failed (class=%s)",
|
||||
type(exc).__name__,
|
||||
)
|
||||
from core.public_errors import stream_generation_failure
|
||||
from core import error_journal
|
||||
|
||||
error_journal.record(
|
||||
exc, route="/generate", trace=traceback.format_exc()
|
||||
)
|
||||
yield _line({"type": "error", **stream_generation_failure(exc)})
|
||||
finally:
|
||||
if not render.done():
|
||||
render.cancel()
|
||||
if cleanup_ref and ref_audio_path:
|
||||
with contextlib.suppress(OSError):
|
||||
os.remove(ref_audio_path)
|
||||
if cleanup_ref and ref_lease is not None:
|
||||
ref_lease.finish_request()
|
||||
|
||||
return StreamingResponse(
|
||||
_remote_stream_events(),
|
||||
@@ -1716,6 +1881,17 @@ async def generate_speech(
|
||||
instruct=instruct, num_step=num_step,
|
||||
guidance_scale=guidance_scale, speed=speed,
|
||||
denoise=denoise, postprocess_output=postprocess_output,
|
||||
**({
|
||||
key: value for key, value in {
|
||||
"t_shift": t_shift,
|
||||
"layer_penalty_factor": layer_penalty_factor,
|
||||
"position_temperature": position_temperature,
|
||||
"class_temperature": class_temperature,
|
||||
"seed": used_seed + i if used_seed is not None else None,
|
||||
}.items() if value is not None
|
||||
} if getattr(
|
||||
_backend, "supports_native_omnivoice_controls", False
|
||||
) else {}),
|
||||
)
|
||||
sr = _backend.sample_rate
|
||||
skip = getattr(_backend, "applies_own_mastering", False)
|
||||
@@ -1779,33 +1955,45 @@ async def generate_speech(
|
||||
if _has_pause or len(_text_chunks) <= 1:
|
||||
# Single-shot pipeline, unchanged — streamed as one chunk.
|
||||
if _backend is not None:
|
||||
audio_tensor = await run_on_gpu_pool_guarded(
|
||||
functools.partial(
|
||||
_run_backend_inference,
|
||||
_backend, text, language, ref_audio_path, ref_text,
|
||||
instruct, duration, num_step, guidance_scale, speed,
|
||||
denoise, postprocess_output, used_seed, effect_preset,
|
||||
max_chunk_chars, crossfade_ms, dropped_sink=_dropped_sink,
|
||||
),
|
||||
what="TTS generate",
|
||||
min_vram_gb=_engine_min_vram_gb,
|
||||
timeout=_generate_timeout_s(text),
|
||||
audio_tensor = await _run_with_reference_lease(
|
||||
ref_lease,
|
||||
lambda release: run_on_gpu_pool_guarded(
|
||||
functools.partial(
|
||||
_run_backend_inference,
|
||||
_backend, text, language, ref_audio_path, ref_text,
|
||||
instruct, duration, num_step, guidance_scale, speed,
|
||||
denoise, postprocess_output, used_seed, effect_preset,
|
||||
max_chunk_chars, crossfade_ms, t_shift=t_shift,
|
||||
layer_penalty_factor=layer_penalty_factor,
|
||||
position_temperature=position_temperature,
|
||||
class_temperature=class_temperature,
|
||||
dropped_sink=_dropped_sink,
|
||||
),
|
||||
what="TTS generate",
|
||||
min_vram_gb=_engine_min_vram_gb,
|
||||
timeout=_generate_timeout_s(text, execution_device=_routing["effective_device"]),
|
||||
on_abandon=release,
|
||||
)
|
||||
)
|
||||
sample_rate = _backend.sample_rate
|
||||
else:
|
||||
audio_tensor = await run_on_gpu_pool_guarded(
|
||||
functools.partial(
|
||||
_run_inference,
|
||||
_model, text, language, ref_audio_path, ref_text,
|
||||
instruct, duration, num_step, guidance_scale, speed,
|
||||
t_shift, denoise, postprocess_output,
|
||||
layer_penalty_factor, position_temperature,
|
||||
class_temperature, used_seed, effect_preset,
|
||||
max_chunk_chars, crossfade_ms, dropped_sink=_dropped_sink,
|
||||
),
|
||||
what="TTS generate",
|
||||
min_vram_gb=_engine_min_vram_gb,
|
||||
timeout=_generate_timeout_s(text),
|
||||
audio_tensor = await _run_with_reference_lease(
|
||||
ref_lease,
|
||||
lambda release: run_on_gpu_pool_guarded(
|
||||
functools.partial(
|
||||
_run_inference,
|
||||
_model, text, language, ref_audio_path, ref_text,
|
||||
instruct, duration, num_step, guidance_scale, speed,
|
||||
t_shift, denoise, postprocess_output,
|
||||
layer_penalty_factor, position_temperature,
|
||||
class_temperature, used_seed, effect_preset,
|
||||
max_chunk_chars, crossfade_ms, dropped_sink=_dropped_sink,
|
||||
),
|
||||
what="TTS generate",
|
||||
min_vram_gb=_engine_min_vram_gb,
|
||||
timeout=_generate_timeout_s(text, execution_device=_routing["effective_device"]),
|
||||
on_abandon=release,
|
||||
)
|
||||
)
|
||||
sample_rate = _model.sampling_rate
|
||||
yield _line({
|
||||
@@ -1822,12 +2010,10 @@ async def generate_speech(
|
||||
# (#1190): AudioSeal embedding is CPU work that owns no
|
||||
# VRAM, and on a 1-worker host it used to serialize
|
||||
# directly ahead of the next generate.
|
||||
from services.watermark import mark_synthetic
|
||||
from services.model_manager import get_watermark_pool
|
||||
_preview = await asyncio.get_running_loop().run_in_executor(
|
||||
get_watermark_pool(),
|
||||
functools.partial(mark_synthetic, audio_tensor, sample_rate,
|
||||
context="generate.stream_preview"),
|
||||
from services.watermark import mark_synthetic_async
|
||||
_preview = await mark_synthetic_async(
|
||||
audio_tensor, sample_rate,
|
||||
context="generate.stream_preview",
|
||||
)
|
||||
yield _line({"type": "chunk", "seq": 0, "pcm": _pcm16_b64(_preview)})
|
||||
else:
|
||||
@@ -1836,25 +2022,27 @@ async def generate_speech(
|
||||
for i, chunk_text in enumerate(_text_chunks):
|
||||
# Bounded per chunk + pool-reset on hang (#730 class);
|
||||
# a timeout surfaces as an "error" event below.
|
||||
raw, preview, sample_rate = await run_on_gpu_pool_guarded(
|
||||
functools.partial(_render_stream_chunk, i, chunk_text),
|
||||
what="TTS generate",
|
||||
min_vram_gb=_engine_min_vram_gb,
|
||||
# Budget scaled to THIS chunk (#1190) — the flat
|
||||
# 300s here is what made long streamed renders fail
|
||||
# even after the v0.3.22 scaled budget shipped.
|
||||
timeout=_generate_timeout_s(chunk_text),
|
||||
raw, preview, sample_rate = await _run_with_reference_lease(
|
||||
ref_lease,
|
||||
lambda release: run_on_gpu_pool_guarded(
|
||||
functools.partial(_render_stream_chunk, i, chunk_text),
|
||||
what="TTS generate",
|
||||
min_vram_gb=_engine_min_vram_gb,
|
||||
# Budget scaled to THIS chunk (#1190) — the flat
|
||||
# 300s here is what made long streamed renders fail
|
||||
# even after the v0.3.22 scaled budget shipped.
|
||||
timeout=_generate_timeout_s(chunk_text, execution_device=_routing["effective_device"]),
|
||||
on_abandon=release,
|
||||
)
|
||||
)
|
||||
parts.append(raw)
|
||||
# Provenance-mark the streamed copy off the GPU pool
|
||||
# (#1169 mark, #1190 placement): CPU-only AudioSeal
|
||||
# work must not occupy a GPU worker between chunks.
|
||||
from services.watermark import mark_synthetic
|
||||
from services.model_manager import get_watermark_pool
|
||||
preview = await asyncio.get_running_loop().run_in_executor(
|
||||
get_watermark_pool(),
|
||||
functools.partial(mark_synthetic, preview, sample_rate,
|
||||
context="generate.stream_preview"),
|
||||
from services.watermark import mark_synthetic_async
|
||||
preview = await mark_synthetic_async(
|
||||
preview, sample_rate,
|
||||
context="generate.stream_preview",
|
||||
)
|
||||
if i == 0:
|
||||
# After the first render so lazy-loading engines
|
||||
@@ -1869,7 +2057,7 @@ async def generate_speech(
|
||||
audio_tensor = await run_on_gpu_pool_guarded(
|
||||
functools.partial(_assemble_stream_chunks, parts, sample_rate),
|
||||
what="TTS assemble",
|
||||
timeout=_generate_timeout_s(text),
|
||||
timeout=_generate_timeout_s(text, execution_device=_routing["effective_device"]),
|
||||
)
|
||||
|
||||
_, meta = await _finalize_generation(
|
||||
@@ -1898,7 +2086,7 @@ async def generate_speech(
|
||||
# Client went away mid-stream — same semantics as aborting a
|
||||
# classic /generate mid-render: nothing is saved.
|
||||
raise
|
||||
except (GpuJobTimeoutError, GpuPoolBusyError) as e:
|
||||
except GpuPoolBusyError as e:
|
||||
# In-band error frame carries the machine-readable retryable
|
||||
# marker (#1190) — an NDJSON consumer can back off instead of
|
||||
# guessing from the prose.
|
||||
@@ -1907,20 +2095,44 @@ async def generate_speech(
|
||||
failure = stream_failure("generation_busy")
|
||||
failure["retry_after"] = getattr(e, "retry_after", 30)
|
||||
yield _line({"type": "error", **failure})
|
||||
except GpuJobTimeoutError:
|
||||
# The worker started and spent its full execution budget. That
|
||||
# is compute time, not queue pressure (#1588).
|
||||
logger.error("Streaming generation exceeded its compute budget")
|
||||
from core.public_errors import stream_failure
|
||||
failure = stream_failure("generation_timeout")
|
||||
failure["retry_after"] = 30
|
||||
yield _line({"type": "error", **failure})
|
||||
except ValueError:
|
||||
logger.error("Streaming generation request rejected")
|
||||
from core.public_errors import stream_failure
|
||||
yield _line({"type": "error", **stream_failure("invalid_request")})
|
||||
except Exception:
|
||||
logger.error("Streaming generation failed unexpectedly")
|
||||
from core.public_errors import stream_failure
|
||||
yield _line({"type": "error", **stream_failure("generation_failed")})
|
||||
except Exception as exc:
|
||||
# A streaming request answers 200 and carries its failure as an
|
||||
# in-band error frame, so it never reaches the global 500
|
||||
# handler — which is where a classic /generate failure gets its
|
||||
# scrubbed journal entry (Diagnostics / recent errors) AND its
|
||||
# classified, actionable message. Both have to be reproduced
|
||||
# here or a streaming generation failure is invisible in the
|
||||
# diagnostic bundle and opaque to the user (#1607). The raw
|
||||
# exception is NOT logged: it can carry a reference-clip path or
|
||||
# a provider secret, and only the journal scrubs before storing.
|
||||
logger.error(
|
||||
"Streaming generation failed unexpectedly (class=%s)",
|
||||
type(exc).__name__,
|
||||
)
|
||||
from core.public_errors import stream_generation_failure
|
||||
from core import error_journal
|
||||
|
||||
error_journal.record(
|
||||
exc, route="/generate", trace=traceback.format_exc()
|
||||
)
|
||||
yield _line({"type": "error", **stream_generation_failure(exc)})
|
||||
finally:
|
||||
# Ownership of the temp reference clip moves to this generator
|
||||
# in stream mode (the route returns before rendering starts).
|
||||
if cleanup_ref and ref_audio_path:
|
||||
with contextlib.suppress(OSError):
|
||||
os.remove(ref_audio_path)
|
||||
if cleanup_ref and ref_lease is not None:
|
||||
ref_lease.finish_request()
|
||||
|
||||
# Routing notice (#21): known before the stream starts, so it rides the
|
||||
# same headers the classic path uses — and now also carries "your
|
||||
@@ -1958,7 +2170,11 @@ async def generate_speech(
|
||||
_backend, text, language, ref_audio_path, ref_text, instruct,
|
||||
duration, num_step, guidance_scale, speed, denoise,
|
||||
postprocess_output, used_seed, effect_preset,
|
||||
max_chunk_chars, crossfade_ms, dropped_sink=_dropped_text,
|
||||
max_chunk_chars, crossfade_ms, t_shift=t_shift,
|
||||
layer_penalty_factor=layer_penalty_factor,
|
||||
position_temperature=position_temperature,
|
||||
class_temperature=class_temperature,
|
||||
dropped_sink=_dropped_text,
|
||||
)
|
||||
else:
|
||||
_local_render = functools.partial(
|
||||
@@ -1969,14 +2185,18 @@ async def generate_speech(
|
||||
class_temperature, used_seed, effect_preset,
|
||||
max_chunk_chars, crossfade_ms, dropped_sink=_dropped_text,
|
||||
)
|
||||
audio_tensor = await gpu_gateway.run(
|
||||
_REMOTE_OP,
|
||||
local=gpu_gateway.LocalCall(
|
||||
_local_render, what="TTS generate",
|
||||
timeout=_generate_timeout_s(text),
|
||||
min_vram_gb=_engine_min_vram_gb,
|
||||
),
|
||||
decision=_decision,
|
||||
audio_tensor = await _run_with_reference_lease(
|
||||
ref_lease,
|
||||
lambda release: gpu_gateway.run(
|
||||
_REMOTE_OP,
|
||||
local=gpu_gateway.LocalCall(
|
||||
_local_render, what="TTS generate",
|
||||
timeout=_generate_timeout_s(text, execution_device=_routing["effective_device"]),
|
||||
min_vram_gb=_engine_min_vram_gb,
|
||||
on_abandon=release,
|
||||
),
|
||||
decision=_decision,
|
||||
)
|
||||
)
|
||||
# Read after generation: engines with lazy model loading report
|
||||
# their real rate only once weights are up.
|
||||
@@ -2108,9 +2328,8 @@ async def generate_speech(
|
||||
),
|
||||
)
|
||||
finally:
|
||||
if cleanup_ref and ref_audio_path:
|
||||
with contextlib.suppress(OSError):
|
||||
os.remove(ref_audio_path)
|
||||
if cleanup_ref and ref_lease is not None:
|
||||
ref_lease.finish_request()
|
||||
|
||||
def _safe_output_path(name):
|
||||
if not name:
|
||||
|
||||
@@ -160,7 +160,9 @@ _OPENAI_VOICE_ALIASES = {
|
||||
|
||||
def _resolve_engine(model_id: str):
|
||||
"""Map an OpenAI model name to a VoiceStudio backend."""
|
||||
from services.tts_backend import get_backend_class, get_active_tts_backend
|
||||
from services.tts_backend import (
|
||||
get_backend_class, get_active_tts_backend, get_engine_instance_for,
|
||||
)
|
||||
|
||||
# Accept OpenAI model names as pass-through to the active engine.
|
||||
if model_id in ("tts-1", "tts-1-hd"):
|
||||
@@ -177,8 +179,18 @@ def _resolve_engine(model_id: str):
|
||||
)
|
||||
from services.tts_backend import OmniVoiceBackend
|
||||
if cls is OmniVoiceBackend:
|
||||
# OmniVoice only ever runs as the shared active engine — the
|
||||
# explicit-omnivoice request is the active-engine request.
|
||||
return get_active_tts_backend()
|
||||
return cls()
|
||||
# Cached singleton, not a fresh cls(): SubprocessBackend engines would
|
||||
# spawn a sidecar process and reload their model on EVERY request, and
|
||||
# register a new atexit hook each time (get_engine_instance's contract).
|
||||
# No router-local cache on top of it: the shared cache is keyed by
|
||||
# CLASS precisely so id rebinds/evictions can't serve a stale instance,
|
||||
# and cross-engine memory discipline is create_speech's
|
||||
# evict_other_tts_engines call (the same seam /generate uses) — not a
|
||||
# bespoke unload here.
|
||||
return get_engine_instance_for(model_id)
|
||||
except ValueError:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
@@ -388,6 +400,15 @@ async def create_speech(req: SpeechRequest):
|
||||
# VRAM eviction runs in get_model()'s warm-return path now, covering every
|
||||
# native TTS generate (this route, WS TTS, dub, batch, audiobook).
|
||||
|
||||
# Single-active-engine memory discipline (MM2-01), the same call /generate
|
||||
# makes before its load: hand back every OTHER resident TTS engine's model
|
||||
# before this one warms up, so switching `model` ids across requests —
|
||||
# explicit id → explicit id, or explicit id → the tts-1/omnivoice aliases —
|
||||
# can't stack multi-GB engines/sidecars. No-op when nothing else is
|
||||
# resident; opt out with OMNIVOICE_SINGLE_ENGINE_RESIDENT=0.
|
||||
from services.engine_memory import evict_other_tts_engines
|
||||
await evict_other_tts_engines(backend.id)
|
||||
|
||||
# ── #1033/#1037/#1014: warm the engine under the LOAD budget before the
|
||||
# generate clock starts. The T4 verification (#1014) measured a fresh
|
||||
# install's first /v1/audio/speech burning its whole 300s generate budget
|
||||
@@ -454,7 +475,7 @@ async def create_speech(req: SpeechRequest):
|
||||
from services.model_manager import generate_timeout_s
|
||||
wav, sr = await run_on_gpu_pool_guarded(
|
||||
lambda: _run_tts(backend, text, kw), what="OpenAI TTS generate",
|
||||
timeout=generate_timeout_s(text))
|
||||
timeout=generate_timeout_s(text, engine=backend))
|
||||
except Exception as e:
|
||||
# #1172/#1173: typed failures get their real status + actionable
|
||||
# message (400 bad input / 503 broken engine binary) instead of a
|
||||
|
||||
@@ -76,6 +76,7 @@ _cancelled: set[str] = set()
|
||||
_active_installs: set[str] = set()
|
||||
_active_installs_lock = threading.Lock()
|
||||
_install_tasks: set[asyncio.Task] = set()
|
||||
_install_tasks_by_repo: dict[str, asyncio.Task] = {}
|
||||
|
||||
|
||||
def _download_max_workers() -> int:
|
||||
@@ -420,16 +421,11 @@ async def install_model(req: InstallModelRequest):
|
||||
f"Retry in {remaining}s or check your network."
|
||||
),
|
||||
)
|
||||
with _active_installs_lock:
|
||||
if req.repo_id in _active_installs:
|
||||
return {"status": "already_running", "repo_id": req.repo_id}
|
||||
_active_installs.add(req.repo_id)
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
def _do():
|
||||
token = hf_progress.current_repo_id.set(req.repo_id)
|
||||
target_token = hf_progress.current_target.set("local")
|
||||
_cancelled.discard(req.repo_id) # clear any stale cancel from a prior run
|
||||
hf_progress.emit({
|
||||
"repo_id": req.repo_id,
|
||||
"filename": req.repo_id,
|
||||
@@ -688,17 +684,53 @@ async def install_model(req: InstallModelRequest):
|
||||
with _active_installs_lock:
|
||||
_active_installs.discard(req.repo_id)
|
||||
|
||||
try:
|
||||
task = loop.create_task(asyncio.to_thread(_do))
|
||||
_install_tasks.add(task)
|
||||
task.add_done_callback(_install_tasks.discard)
|
||||
except Exception:
|
||||
with _active_installs_lock:
|
||||
with _active_installs_lock:
|
||||
if req.repo_id in _active_installs:
|
||||
return {"status": "already_running", "repo_id": req.repo_id}
|
||||
_active_installs.add(req.repo_id)
|
||||
# Admission and task publication are one atomic generation boundary:
|
||||
# cancellation can never observe an admitted install without its task.
|
||||
_cancelled.discard(req.repo_id)
|
||||
try:
|
||||
task = loop.create_task(asyncio.to_thread(_do))
|
||||
_install_tasks.add(task)
|
||||
_install_tasks_by_repo[req.repo_id] = task
|
||||
except Exception:
|
||||
_active_installs.discard(req.repo_id)
|
||||
raise
|
||||
raise
|
||||
|
||||
def install_finished(completed: asyncio.Task) -> None:
|
||||
with _active_installs_lock:
|
||||
_install_tasks.discard(completed)
|
||||
if _install_tasks_by_repo.get(req.repo_id) is completed:
|
||||
_install_tasks_by_repo.pop(req.repo_id, None)
|
||||
|
||||
task.add_done_callback(install_finished)
|
||||
return {"status": "install_started", "repo_id": req.repo_id}
|
||||
|
||||
|
||||
async def cancel_install_and_wait(repo_id: str) -> None:
|
||||
"""Request cancellation and retain authority until its thread exits."""
|
||||
from worker.async_utils import drain_task # noqa: PLC0415
|
||||
|
||||
with _active_installs_lock:
|
||||
_cancelled.add(repo_id)
|
||||
_install_cooldowns.pop(repo_id, None)
|
||||
task = _install_tasks_by_repo.get(repo_id)
|
||||
if task is None:
|
||||
return
|
||||
try:
|
||||
# asyncio.to_thread cannot stop snapshot_download mid-file. Cancelling
|
||||
# its wrapper would only detach the thread, so wait until the blocking
|
||||
# call observes the flag or naturally returns.
|
||||
await drain_task(task)
|
||||
finally:
|
||||
with _active_installs_lock:
|
||||
current = _install_tasks_by_repo.get(repo_id)
|
||||
if current is None or current is task:
|
||||
_cancelled.discard(repo_id)
|
||||
|
||||
|
||||
@router.post("/models/install/cancel")
|
||||
async def cancel_install(req: InstallModelRequest):
|
||||
"""Request cancellation of an in-flight install (FDL-11).
|
||||
|
||||
@@ -62,6 +62,11 @@ def setup_status():
|
||||
_MIN_NVIDIA_DRIVER = 555
|
||||
_RAM_FAIL_GB = 8
|
||||
_RAM_WARN_GB = 12
|
||||
# Installed DIMMs never fully reach the OS: firmware, integrated graphics and
|
||||
# kernel reservations shave off up to ~7% (an "8 GB" Windows laptop reports
|
||||
# ~7.8 GB usable). Thresholds are compared with this allowance applied so the
|
||||
# machines a threshold is meant to admit aren't blocked by that gap (#1618).
|
||||
_RAM_RESERVED_ALLOWANCE = 0.93
|
||||
|
||||
|
||||
def _run_cmd(args: list[str], timeout: float = 2.0) -> tuple[int, str]:
|
||||
@@ -352,17 +357,28 @@ def preflight():
|
||||
|
||||
# ── RAM
|
||||
ram = _ram_gb()
|
||||
# Escape hatch (#1618): a preflight should inform, not brick setup —
|
||||
# OMNIVOICE_RAM_PREFLIGHT=0 downgrades the hard block to a warning for
|
||||
# users who accept the OOM risk. Same opt-out shape as
|
||||
# OMNIVOICE_ASR_VRAM_PREFLIGHT.
|
||||
ram_gate = os.environ.get(
|
||||
"OMNIVOICE_RAM_PREFLIGHT", "1"
|
||||
).strip().lower() not in ("0", "false", "no")
|
||||
if ram == 0:
|
||||
ram_status, ram_detail, ram_fix = (
|
||||
"warn", "Could not detect system RAM.",
|
||||
"Install psutil in the backend environment or ignore this warning.",
|
||||
)
|
||||
elif ram < _RAM_FAIL_GB:
|
||||
elif ram < _RAM_FAIL_GB * _RAM_RESERVED_ALLOWANCE:
|
||||
ram_status, ram_detail, ram_fix = (
|
||||
"fail", f"{ram:.1f} GB total (need ≥ {_RAM_FAIL_GB} GB)",
|
||||
"The app will OOM on first dub. Close other apps or upgrade RAM.",
|
||||
"fail" if ram_gate else "warn",
|
||||
f"{ram:.1f} GB total (need ≥ {_RAM_FAIL_GB} GB)",
|
||||
"The app will OOM on first dub. Close other apps or upgrade RAM."
|
||||
if ram_gate else
|
||||
"RAM check disabled via OMNIVOICE_RAM_PREFLIGHT=0 — dubbing may "
|
||||
"OOM on this machine.",
|
||||
)
|
||||
elif ram < _RAM_WARN_GB:
|
||||
elif ram < _RAM_WARN_GB * _RAM_RESERVED_ALLOWANCE:
|
||||
ram_status, ram_detail, ram_fix = (
|
||||
"warn", f"{ram:.1f} GB total ({_RAM_WARN_GB}+ GB recommended)",
|
||||
"Long videos may hit swap. Keep other apps closed during dubbing.",
|
||||
|
||||
@@ -0,0 +1,160 @@
|
||||
"""Discovery contract for VoiceStudio's local speech platform.
|
||||
|
||||
Interfaces should discover this document instead of hard-coding whichever
|
||||
dictation route the desktop happens to use. Endpoint URLs are relative so the
|
||||
same response works on loopback, a tailnet GPU host, and a reverse proxy.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Literal
|
||||
|
||||
from fastapi import APIRouter
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from core.version import APP_VERSION
|
||||
|
||||
router = APIRouter(tags=["Speech Platform"])
|
||||
|
||||
SPEECH_PROTOCOL = "voicestudio.speech.v1"
|
||||
STREAM_PATH = "/v1/audio/transcriptions/stream"
|
||||
|
||||
|
||||
class EndpointCapability(BaseModel):
|
||||
path: str
|
||||
transport: Literal["http", "websocket", "mcp-streamable-http", "mcp-stdio"]
|
||||
method: str | None = None
|
||||
protocol: str | None = None
|
||||
|
||||
|
||||
class StreamInputCapability(BaseModel):
|
||||
framing: Literal["binary"] = "binary"
|
||||
formats: list[str]
|
||||
default_format: str
|
||||
sample_rate_query: str = "sr"
|
||||
end_control: dict[str, str]
|
||||
|
||||
|
||||
class StreamOutputCapability(BaseModel):
|
||||
framing: Literal["json"] = "json"
|
||||
events: list[str]
|
||||
final_kinds: list[str]
|
||||
|
||||
|
||||
class SpeechFeatureCapabilities(BaseModel):
|
||||
batch_transcription: bool = True
|
||||
streaming_transcription: bool = True
|
||||
partial_transcripts: bool = True
|
||||
utterance_finals: bool = True
|
||||
session_summary: bool = True
|
||||
word_timestamps: bool = True
|
||||
local_refinement: bool = True
|
||||
acoustic_echo_cancellation: bool = True
|
||||
native_dictation_control: bool = False
|
||||
|
||||
|
||||
class SpeechAuthCapabilities(BaseModel):
|
||||
loopback: Literal["none"] = "none"
|
||||
remote: Literal["bearer"] = "bearer"
|
||||
header: str = "Authorization: Bearer <OMNIVOICE_API_KEY>"
|
||||
browser_session_endpoint: str = "/api/auth/session"
|
||||
websocket_ticket_endpoint: str = "/api/auth/ws-ticket"
|
||||
websocket_ticket_query_parameter: Literal["ws_ticket"] = "ws_ticket"
|
||||
|
||||
|
||||
class SpeechCapabilities(BaseModel):
|
||||
schema_: Literal["voicestudio.speech-capabilities"] = Field(
|
||||
default="voicestudio.speech-capabilities",
|
||||
serialization_alias="schema",
|
||||
)
|
||||
protocol: Literal["voicestudio.speech.v1"] = SPEECH_PROTOCOL
|
||||
protocol_version: Literal["1.0"] = "1.0"
|
||||
service: str = "VoiceStudio"
|
||||
service_version: str = APP_VERSION
|
||||
local_first: bool = True
|
||||
endpoints: dict[str, EndpointCapability]
|
||||
stream_input: StreamInputCapability
|
||||
stream_output: StreamOutputCapability
|
||||
features: SpeechFeatureCapabilities
|
||||
authentication: SpeechAuthCapabilities
|
||||
|
||||
|
||||
def speech_capabilities() -> SpeechCapabilities:
|
||||
"""Return the stable, side-effect-free integration contract."""
|
||||
endpoints = {
|
||||
"capabilities": EndpointCapability(
|
||||
path="/.well-known/voicestudio-speech",
|
||||
transport="http",
|
||||
method="GET",
|
||||
),
|
||||
"batch_transcription": EndpointCapability(
|
||||
path="/v1/audio/transcriptions",
|
||||
transport="http",
|
||||
method="POST",
|
||||
protocol="openai.audio.transcriptions",
|
||||
),
|
||||
"streaming_transcription": EndpointCapability(
|
||||
path=STREAM_PATH,
|
||||
transport="websocket",
|
||||
protocol=SPEECH_PROTOCOL,
|
||||
),
|
||||
"mcp": EndpointCapability(
|
||||
path="/mcp",
|
||||
transport="mcp-streamable-http",
|
||||
method="POST",
|
||||
protocol="mcp",
|
||||
),
|
||||
"mcp_stdio": EndpointCapability(
|
||||
path="python -m backend.mcp_shim",
|
||||
transport="mcp-stdio",
|
||||
protocol="mcp",
|
||||
),
|
||||
}
|
||||
native_control = False
|
||||
try:
|
||||
control_port = int(os.environ.get("VOICESTUDIO_SPEECH_CONTROL_PORT", ""))
|
||||
except (TypeError, ValueError):
|
||||
control_port = 0
|
||||
if 0 < control_port <= 65535:
|
||||
native_control = True
|
||||
endpoints["native_dictation_control"] = EndpointCapability(
|
||||
path=f"http://127.0.0.1:{control_port}/v1/capabilities",
|
||||
transport="http",
|
||||
method="GET",
|
||||
protocol=SPEECH_PROTOCOL,
|
||||
)
|
||||
|
||||
return SpeechCapabilities(
|
||||
endpoints=endpoints,
|
||||
stream_input=StreamInputCapability(
|
||||
formats=[
|
||||
"audio/pcm;encoding=s16le;channels=1",
|
||||
"audio/webm;codecs=opus",
|
||||
],
|
||||
default_format="audio/webm;codecs=opus",
|
||||
end_control={"type": "input_audio.end"},
|
||||
),
|
||||
stream_output=StreamOutputCapability(
|
||||
events=["session.started", "status", "partial", "final", "error"],
|
||||
final_kinds=["utterance", "summary"],
|
||||
),
|
||||
features=SpeechFeatureCapabilities(
|
||||
native_dictation_control=native_control,
|
||||
),
|
||||
authentication=SpeechAuthCapabilities(),
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/.well-known/voicestudio-speech",
|
||||
response_model=SpeechCapabilities,
|
||||
response_model_by_alias=True,
|
||||
)
|
||||
@router.get(
|
||||
"/v1/audio/capabilities",
|
||||
response_model=SpeechCapabilities,
|
||||
response_model_by_alias=True,
|
||||
)
|
||||
async def get_speech_capabilities() -> SpeechCapabilities:
|
||||
"""Advertise batch, streaming, and agent-facing speech transports."""
|
||||
return speech_capabilities()
|
||||
@@ -10,7 +10,8 @@ as they're generated. This unlocks:
|
||||
Protocol:
|
||||
→ Client sends JSON: {"text": "...", "voice": "profile_id", ...}
|
||||
← Server sends binary audio chunks (PCM16 @ 24kHz mono) as generated
|
||||
← Server sends JSON: {"type": "done", "duration_s": 4.2, "gen_time_s": 1.1}
|
||||
← Server sends JSON: {"type": "done", "duration_s": 4.2,
|
||||
"gen_time_s": 1.1, "ttfa_ms": 180.0, "rtf": 0.262}
|
||||
← Server sends JSON: {"type": "error", "detail": "..."}
|
||||
|
||||
The chunked delivery targets <100ms time-to-first-audio (TTFA) on warm models.
|
||||
@@ -33,6 +34,30 @@ logger = logging.getLogger("omnivoice.tts_stream")
|
||||
# Smaller chunks = lower latency but more WebSocket overhead.
|
||||
CHUNK_SAMPLES = int(os.environ.get("OMNIVOICE_STREAM_CHUNK", "4800"))
|
||||
|
||||
# Module seam for deterministic latency-contract tests. Keep every timing
|
||||
# sample on the same monotonic clock.
|
||||
_perf_counter = time.perf_counter
|
||||
|
||||
|
||||
async def _resolve_stream_backend(engine_id: str | None):
|
||||
"""Resolve the live-stream engine without bypassing host isolation."""
|
||||
from services.tts_backend import (
|
||||
OmniVoiceBackend,
|
||||
active_backend_id,
|
||||
get_active_tts_backend,
|
||||
get_backend_class,
|
||||
)
|
||||
|
||||
if engine_id:
|
||||
return get_backend_class(engine_id)()
|
||||
|
||||
cls = get_backend_class(active_backend_id())
|
||||
if cls is OmniVoiceBackend:
|
||||
from services.model_manager import get_model
|
||||
|
||||
return get_active_tts_backend(model=await get_model())
|
||||
return get_active_tts_backend()
|
||||
|
||||
|
||||
class StreamTTSRequest(BaseModel):
|
||||
"""Client request for streaming TTS."""
|
||||
@@ -85,7 +110,7 @@ async def ws_tts(websocket: WebSocket):
|
||||
})
|
||||
continue
|
||||
|
||||
t0 = time.perf_counter()
|
||||
t0 = _perf_counter()
|
||||
text = data["text"]
|
||||
|
||||
# Remote GPU: this socket stays on this machine, and says so.
|
||||
@@ -127,10 +152,6 @@ async def ws_tts(websocket: WebSocket):
|
||||
|
||||
try:
|
||||
# Resolve engine
|
||||
from services.tts_backend import (
|
||||
get_active_tts_backend,
|
||||
get_backend_class,
|
||||
)
|
||||
engine_id = data.get("engine")
|
||||
# #1224: leave a breadcrumb when memory is already tight before
|
||||
# a heavy load. /generate has done this since the 16 GB-Mac
|
||||
@@ -146,13 +167,7 @@ async def ws_tts(websocket: WebSocket):
|
||||
log_if_low(f"TTS stream load ({engine_id or 'active engine'})")
|
||||
except Exception:
|
||||
pass
|
||||
if engine_id:
|
||||
cls = get_backend_class(engine_id)
|
||||
backend = cls()
|
||||
else:
|
||||
from services.model_manager import get_model
|
||||
model = await get_model()
|
||||
backend = get_active_tts_backend(model=model)
|
||||
backend = await _resolve_stream_backend(engine_id)
|
||||
|
||||
# ── Routing gate (#21 — no silent CPU fallback). WebSockets have
|
||||
# no response headers, so this uses frames: an error frame +
|
||||
@@ -258,6 +273,11 @@ async def ws_tts(websocket: WebSocket):
|
||||
from services.model_manager import run_on_gpu_pool_guarded
|
||||
|
||||
def _generate(sentence_text):
|
||||
# Timed INSIDE the pool worker: the guarded dispatch below
|
||||
# can queue behind other jobs, and queue wait is not
|
||||
# synthesis (review on #1620) — under contention it would
|
||||
# inflate rtf without the engine slowing at all.
|
||||
_synth_t0 = _perf_counter()
|
||||
from services.audio_dsp import apply_mastering, normalize_audio
|
||||
from services.watermark import mark_synthetic
|
||||
wav = backend.generate(sentence_text, **kw)
|
||||
@@ -279,12 +299,19 @@ async def ws_tts(websocket: WebSocket):
|
||||
# watermark._iter_chunks), which is inherent to marking
|
||||
# ultra-short clips, not a coverage gap.
|
||||
wav = mark_synthetic(wav, sr_actual, context="tts_stream.sentence")
|
||||
return wav, sr_actual
|
||||
return wav, sr_actual, _perf_counter() - _synth_t0
|
||||
|
||||
import torch
|
||||
total_samples = 0
|
||||
sr = backend.sample_rate
|
||||
started = False
|
||||
first_audio_at: float | None = None
|
||||
# Synthesis time only. The wall clock below also carries socket
|
||||
# delivery and the per-chunk event-loop yields, so deriving RTF
|
||||
# from it reports "how slow was the client" as if it were engine
|
||||
# throughput — on a slow consumer that inflates RTF without the
|
||||
# engine having changed at all.
|
||||
synth_time = 0.0
|
||||
|
||||
for sentence in sentences:
|
||||
# Bounded + pool-reset on hang so a wedged generate can't
|
||||
@@ -294,11 +321,12 @@ async def ws_tts(websocket: WebSocket):
|
||||
# Length-scaled budget per sentence (#1190) — the flat 300s
|
||||
# default is gone from every dispatch.
|
||||
from services.model_manager import generate_timeout_s
|
||||
wav_tensor, sr = await run_on_gpu_pool_guarded(
|
||||
wav_tensor, sr, sentence_synth_s = await run_on_gpu_pool_guarded(
|
||||
functools.partial(_generate, sentence),
|
||||
what="TTS generate",
|
||||
timeout=generate_timeout_s(sentence),
|
||||
timeout=generate_timeout_s(sentence, engine=backend),
|
||||
)
|
||||
synth_time += sentence_synth_s
|
||||
|
||||
if not started:
|
||||
# Send metadata after the first generation so
|
||||
@@ -325,25 +353,49 @@ async def ws_tts(websocket: WebSocket):
|
||||
end = min(sent_samples + CHUNK_SAMPLES, n_samples)
|
||||
chunk = pcm_bytes[sent_samples * 2: end * 2]
|
||||
await websocket.send_bytes(chunk)
|
||||
if first_audio_at is None:
|
||||
# TTFA ends when the first audio bytes have been
|
||||
# handed to the socket. The previous log used the
|
||||
# whole-render duration and called it TTFA.
|
||||
first_audio_at = _perf_counter()
|
||||
sent_samples = end
|
||||
# Yield to event loop between chunks for responsiveness
|
||||
await asyncio.sleep(0)
|
||||
total_samples += n_samples
|
||||
|
||||
gen_time = round(time.perf_counter() - t0, 3)
|
||||
finished_at = _perf_counter()
|
||||
wall_time_raw = max(0.0, finished_at - t0)
|
||||
synth_time_raw = max(0.0, synth_time)
|
||||
gen_time = round(wall_time_raw, 3)
|
||||
duration = round(total_samples / sr, 3)
|
||||
ttfa_ms = (
|
||||
round(max(0.0, first_audio_at - t0) * 1000.0, 1)
|
||||
if first_audio_at is not None
|
||||
else None
|
||||
)
|
||||
# RTF is a render metric: synthesis seconds per audio second.
|
||||
rtf = (
|
||||
round(synth_time_raw / (total_samples / sr), 3)
|
||||
if total_samples > 0
|
||||
else None
|
||||
)
|
||||
|
||||
await websocket.send_json({
|
||||
"type": "done",
|
||||
"duration_s": duration,
|
||||
"gen_time_s": gen_time,
|
||||
"ttfa_ms": ttfa_ms,
|
||||
"rtf": rtf,
|
||||
"samples": total_samples,
|
||||
"sample_rate": sr,
|
||||
"engine": backend.id,
|
||||
})
|
||||
logger.info(
|
||||
"TTS stream: %.1fs audio in %.1fs (TTFA=%.0fms)",
|
||||
duration, gen_time, gen_time * 1000,
|
||||
"TTS stream: %.1fs audio in %.1fs (TTFA=%s, RTF=%s)",
|
||||
duration,
|
||||
gen_time,
|
||||
f"{ttfa_ms:.0f}ms" if ttfa_ms is not None else "n/a",
|
||||
f"{rtf:.3f}" if rtf is not None else "n/a",
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
|
||||
+266
-55
@@ -23,13 +23,16 @@ appears and is replaced by the GPU gateway.
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from api.dependencies import require_admin
|
||||
from worker import registry, routing, service
|
||||
from worker.async_utils import drain_task, to_thread_and_defer_cancellation
|
||||
|
||||
logger = logging.getLogger("omnivoice.worker")
|
||||
|
||||
@@ -158,6 +161,19 @@ def agent_status() -> dict:
|
||||
return worker_agent.agent.status()
|
||||
|
||||
|
||||
@router.get("/agent/readiness", include_in_schema=False, response_model=None)
|
||||
def agent_readiness() -> JSONResponse:
|
||||
"""Container readiness: 200 only after this process registered as a worker."""
|
||||
from worker import agent as worker_agent # noqa: PLC0415
|
||||
|
||||
readiness = worker_agent.agent.readiness()
|
||||
return JSONResponse(
|
||||
status_code=200 if readiness["ready"] else 503,
|
||||
content=readiness,
|
||||
headers={} if readiness["ready"] else {"Retry-After": "2"},
|
||||
)
|
||||
|
||||
|
||||
def _refuse_when_env_pinned(worker_agent) -> None:
|
||||
"""OMNIVOICE_WORKER_MODE wins over the setting everywhere else.
|
||||
|
||||
@@ -176,6 +192,63 @@ def _refuse_when_env_pinned(worker_agent) -> None:
|
||||
)
|
||||
|
||||
|
||||
async def _finish_cleanup(awaitable):
|
||||
"""Run rollback to completion even if its HTTP task was cancelled."""
|
||||
task = asyncio.create_task(awaitable)
|
||||
await drain_task(task)
|
||||
return task.result()
|
||||
|
||||
|
||||
async def _set_worker_mode(worker_agent, enabled: bool) -> None:
|
||||
_result, cancelled = await to_thread_and_defer_cancellation(
|
||||
worker_agent.set_worker_mode_enabled, enabled
|
||||
)
|
||||
if cancelled:
|
||||
raise asyncio.CancelledError
|
||||
|
||||
|
||||
async def _restore_agent_transaction(
|
||||
worker_agent, previous: dict, *, was_running: bool
|
||||
) -> None:
|
||||
"""Restore durable enrollment/settings and the exact prior live state."""
|
||||
try:
|
||||
await _finish_cleanup(worker_agent.agent.stop())
|
||||
await _finish_cleanup(worker_agent.restore_enrollment(previous))
|
||||
if was_running and not worker_agent.agent.running:
|
||||
await _finish_cleanup(worker_agent.agent.start())
|
||||
elif not was_running and worker_agent.agent.running:
|
||||
await _finish_cleanup(worker_agent.agent.stop())
|
||||
except worker_agent.EnrollmentRollbackError:
|
||||
raise
|
||||
except BaseException as exc:
|
||||
message = (
|
||||
"The previous worker state could not be restored safely. "
|
||||
"Worker mode remains stopped; fix its enrollment/settings storage, then retry."
|
||||
)
|
||||
with contextlib.suppress(BaseException):
|
||||
await _finish_cleanup(worker_agent.agent.stop())
|
||||
worker_agent.agent.last_error = message
|
||||
raise worker_agent.EnrollmentRollbackError(message) from exc
|
||||
|
||||
|
||||
def _raise_agent_transaction_failure(
|
||||
worker_agent, operation: BaseException, rollback: BaseException | None
|
||||
) -> None:
|
||||
if isinstance(operation, asyncio.CancelledError):
|
||||
if rollback is not None:
|
||||
logger.error(
|
||||
"Worker rollback failed during request cancellation",
|
||||
exc_info=(type(rollback), rollback, rollback.__traceback__),
|
||||
)
|
||||
raise operation
|
||||
if rollback is not None:
|
||||
raise HTTPException(status_code=409, detail=str(rollback)) from rollback
|
||||
if isinstance(operation, Exception):
|
||||
worker_agent.agent.last_error = str(operation)
|
||||
raise HTTPException(status_code=409, detail=str(operation)) from operation
|
||||
raise operation
|
||||
|
||||
|
||||
@router.post("/agent/join")
|
||||
async def join_control_plane(request: JoinRequest) -> dict:
|
||||
"""Redeem a join code and start working for that control plane.
|
||||
@@ -200,26 +273,37 @@ async def join_control_plane(request: JoinRequest) -> dict:
|
||||
# says it joined and never lends anything (CodeRabbit).
|
||||
_refuse_when_env_pinned(worker_agent)
|
||||
async with worker_agent.agent.lifecycle:
|
||||
# A rejoin replaces a working enrollment. Keep enough to put it back:
|
||||
# pinning the new certificate overwrites the old one on disk, so a
|
||||
# failed rejoin would otherwise leave the machine unable to reconnect
|
||||
# to the control plane it was already serving.
|
||||
previous = worker_agent.snapshot_enrollment()
|
||||
await worker_agent.agent.stop()
|
||||
try:
|
||||
previous, cancelled = await to_thread_and_defer_cancellation(
|
||||
worker_agent.snapshot_enrollment
|
||||
)
|
||||
except worker_agent.EnrollmentStateError as exc:
|
||||
worker_agent.agent.last_error = str(exc)
|
||||
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
||||
if cancelled:
|
||||
raise asyncio.CancelledError
|
||||
was_running = worker_agent.agent.running
|
||||
|
||||
# A rejoin stops a working agent before the replacement is accepted.
|
||||
# Stop, acceptance and the durable setting are one transaction: every
|
||||
# failure, including cancellation, restores both trust and live state.
|
||||
try:
|
||||
await worker_agent.agent.stop()
|
||||
await worker_agent.agent.start(token_text=token)
|
||||
# Success is the control plane ACCEPTING this worker, not the
|
||||
# connection being scheduled — see wait_until_registered.
|
||||
await worker_agent.agent.wait_until_registered()
|
||||
except Exception as exc:
|
||||
worker_agent.agent.last_error = str(exc)
|
||||
await worker_agent.agent.stop()
|
||||
await worker_agent.restore_enrollment(previous)
|
||||
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
||||
await _set_worker_mode(worker_agent, True)
|
||||
except BaseException as exc:
|
||||
rollback_exc = None
|
||||
try:
|
||||
await _restore_agent_transaction(
|
||||
worker_agent, previous, was_running=was_running
|
||||
)
|
||||
except BaseException as rollback_error:
|
||||
rollback_exc = rollback_error
|
||||
_raise_agent_transaction_failure(worker_agent, exc, rollback_exc)
|
||||
worker_agent.agent.last_error = ""
|
||||
# Persisted only after the join actually worked: a machine that failed
|
||||
# to enrol must not come back up trying again forever.
|
||||
worker_agent.set_worker_mode_enabled(True)
|
||||
return worker_agent.agent.status()
|
||||
|
||||
|
||||
@@ -235,19 +319,35 @@ async def set_agent_enabled(request: EnableRequest) -> dict:
|
||||
|
||||
_refuse_when_env_pinned(worker_agent)
|
||||
async with worker_agent.agent.lifecycle:
|
||||
if request.enabled:
|
||||
try:
|
||||
try:
|
||||
previous, cancelled = await to_thread_and_defer_cancellation(
|
||||
worker_agent.snapshot_enrollment
|
||||
)
|
||||
except worker_agent.EnrollmentStateError as exc:
|
||||
worker_agent.agent.last_error = str(exc)
|
||||
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
||||
if cancelled:
|
||||
raise asyncio.CancelledError
|
||||
was_running = worker_agent.agent.running
|
||||
|
||||
try:
|
||||
if request.enabled:
|
||||
await worker_agent.agent.start()
|
||||
await worker_agent.agent.wait_until_registered()
|
||||
except Exception as exc:
|
||||
worker_agent.agent.last_error = str(exc)
|
||||
await _set_worker_mode(worker_agent, True)
|
||||
else:
|
||||
await worker_agent.agent.stop()
|
||||
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
||||
worker_agent.agent.last_error = ""
|
||||
worker_agent.set_worker_mode_enabled(True)
|
||||
else:
|
||||
await worker_agent.agent.stop()
|
||||
worker_agent.set_worker_mode_enabled(False)
|
||||
await _set_worker_mode(worker_agent, False)
|
||||
except BaseException as exc:
|
||||
rollback_exc = None
|
||||
try:
|
||||
await _restore_agent_transaction(
|
||||
worker_agent, previous, was_running=was_running
|
||||
)
|
||||
except BaseException as rollback_error:
|
||||
rollback_exc = rollback_error
|
||||
_raise_agent_transaction_failure(worker_agent, exc, rollback_exc)
|
||||
worker_agent.agent.last_error = ""
|
||||
return worker_agent.agent.status()
|
||||
|
||||
|
||||
@@ -263,9 +363,14 @@ def create_enrollment(request: EnrollRequest) -> dict:
|
||||
status_code=409,
|
||||
detail="Remote workers are turned off. Enable them in Settings → System → Remote workers first.",
|
||||
)
|
||||
token = service.control_plane.create_enrollment(
|
||||
endpoint=request.endpoint, label=request.label, ttl_seconds=request.ttl_seconds
|
||||
)
|
||||
try:
|
||||
token = service.control_plane.create_enrollment(
|
||||
endpoint=request.endpoint,
|
||||
label=request.label,
|
||||
ttl_seconds=request.ttl_seconds,
|
||||
)
|
||||
except service.EndpointCertificateError as exc:
|
||||
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
||||
return {
|
||||
"token": token.encode(),
|
||||
"endpoint": token.endpoint,
|
||||
@@ -275,23 +380,55 @@ def create_enrollment(request: EnrollRequest) -> dict:
|
||||
}
|
||||
|
||||
|
||||
def _persist_worker_update(
|
||||
worker_id: str, request: WorkerUpdate
|
||||
):
|
||||
"""Write policy on a worker thread; live publication stays loop-owned."""
|
||||
return registry.update_policy(
|
||||
worker_id,
|
||||
name=request.name,
|
||||
enabled=request.enabled,
|
||||
priority=request.priority,
|
||||
)
|
||||
|
||||
|
||||
@router.patch("/{worker_id}")
|
||||
def update_worker(worker_id: str, request: WorkerUpdate) -> dict:
|
||||
worker = registry.get(worker_id)
|
||||
if worker is None:
|
||||
async def update_worker(worker_id: str, request: WorkerUpdate) -> dict:
|
||||
pool = service.control_plane.pool if service.control_plane.running else None
|
||||
live = None
|
||||
was_pending = False
|
||||
if pool is not None:
|
||||
# Quiesce dispatch before releasing authority for the SQLite write.
|
||||
# The publication after the await restores the exact prior state, so a
|
||||
# concurrent registration handoff remains quiesced for its own reason.
|
||||
with registry.authority_guard():
|
||||
live = pool.get(worker_id)
|
||||
if live is not None:
|
||||
was_pending = live.registration_pending
|
||||
live.registration_pending = True
|
||||
updated = None
|
||||
cancelled = False
|
||||
try:
|
||||
updated, cancelled = await to_thread_and_defer_cancellation(
|
||||
_persist_worker_update, worker_id, request
|
||||
)
|
||||
finally:
|
||||
if pool is not None:
|
||||
with registry.authority_guard():
|
||||
if updated is not None:
|
||||
# Pool state, including the cached record the scheduler
|
||||
# reads, belongs to the app's event loop.
|
||||
pool.refresh_record(updated)
|
||||
current = pool.get(worker_id)
|
||||
if current is live:
|
||||
current.registration_pending = was_pending
|
||||
if updated is None:
|
||||
if cancelled:
|
||||
raise asyncio.CancelledError
|
||||
raise HTTPException(status_code=404, detail="No such worker.")
|
||||
if request.name is not None:
|
||||
registry.rename(worker_id, request.name)
|
||||
if request.enabled is not None:
|
||||
registry.set_enabled(worker_id, request.enabled)
|
||||
if request.priority is not None:
|
||||
registry.set_priority(worker_id, request.priority)
|
||||
updated = registry.get(worker_id)
|
||||
# Keep the live copy in step, so the scheduler and its logs do not go on
|
||||
# using the name or priority this worker had when it connected.
|
||||
if updated is not None and service.control_plane.running:
|
||||
service.control_plane.pool.refresh_record(updated)
|
||||
return updated.to_dict() if updated else {}
|
||||
if cancelled:
|
||||
raise asyncio.CancelledError
|
||||
return updated.to_dict()
|
||||
|
||||
|
||||
@router.post("/{worker_id}/consent")
|
||||
@@ -305,7 +442,7 @@ def grant_consent(worker_id: str) -> dict:
|
||||
|
||||
|
||||
@router.post("/{worker_id}/resume")
|
||||
def clear_breaker(worker_id: str) -> dict:
|
||||
async def clear_breaker(worker_id: str) -> dict:
|
||||
"""Clear a paused worker's circuit breakers.
|
||||
|
||||
The user fixed the machine and knows it — a breaker with no manual clear is
|
||||
@@ -320,18 +457,53 @@ def clear_breaker(worker_id: str) -> dict:
|
||||
|
||||
|
||||
@router.delete("/{worker_id}")
|
||||
def revoke_worker(worker_id: str) -> dict:
|
||||
async def revoke_worker(worker_id: str) -> dict:
|
||||
"""Remove a worker — which means revoke its key, not hide the row.
|
||||
|
||||
Its in-flight work is released so it can be retried elsewhere rather than
|
||||
waiting out a lease on a machine that will never answer again.
|
||||
"""
|
||||
if registry.get(worker_id) is None:
|
||||
pool = service.control_plane.pool if service.control_plane.running else None
|
||||
live = None
|
||||
was_pending = False
|
||||
if pool is not None:
|
||||
with registry.authority_guard():
|
||||
live = pool.get(worker_id)
|
||||
if live is not None:
|
||||
was_pending = live.registration_pending
|
||||
live.registration_pending = True
|
||||
try:
|
||||
revoked, cancelled = await to_thread_and_defer_cancellation(
|
||||
registry.revoke, worker_id
|
||||
)
|
||||
except BaseException:
|
||||
if pool is not None:
|
||||
with registry.authority_guard():
|
||||
current = pool.get(worker_id)
|
||||
if current is live:
|
||||
current.registration_pending = was_pending
|
||||
raise
|
||||
if not revoked:
|
||||
if pool is not None:
|
||||
with registry.authority_guard():
|
||||
current = pool.get(worker_id)
|
||||
if current is live:
|
||||
current.registration_pending = was_pending
|
||||
if cancelled:
|
||||
raise asyncio.CancelledError
|
||||
raise HTTPException(status_code=404, detail="No such worker.")
|
||||
registry.revoke(worker_id)
|
||||
if service.control_plane.running:
|
||||
service.control_plane.scheduler.on_disconnected(worker_id)
|
||||
service.control_plane.pool.breakers.forget_worker(worker_id)
|
||||
|
||||
# The tombstone committed before any egress/session mutation. Everything
|
||||
# below is loop-owned and published under the same scheduler authority read
|
||||
# used by next_assignment(), so no task can bind in the handoff window.
|
||||
with registry.authority_guard():
|
||||
if service.control_plane.running:
|
||||
if service.control_plane.servicer is not None:
|
||||
service.control_plane.servicer.revoke_worker_sessions(worker_id)
|
||||
service.control_plane.scheduler.on_disconnected(worker_id)
|
||||
service.control_plane.pool.breakers.forget_worker(worker_id)
|
||||
if cancelled:
|
||||
raise asyncio.CancelledError
|
||||
return {"ok": True, "revoked": worker_id}
|
||||
|
||||
|
||||
@@ -375,7 +547,9 @@ async def submit_task(request: Request, body: SubmitTaskRequest) -> dict:
|
||||
|
||||
scheduler = service.control_plane.scheduler
|
||||
try:
|
||||
task = scheduler.submit(
|
||||
submit = getattr(scheduler, "submit_async", None)
|
||||
submit = submit if callable(submit) else scheduler.submit
|
||||
submitted = submit(
|
||||
operation=body.operation,
|
||||
engine=body.engine,
|
||||
model_id=body.model_id,
|
||||
@@ -384,6 +558,7 @@ async def submit_task(request: Request, body: SubmitTaskRequest) -> dict:
|
||||
deadline_seconds=body.deadline_seconds,
|
||||
pinned_worker_id=routing.decide().worker_id or None,
|
||||
)
|
||||
task = await submitted if asyncio.iscoroutine(submitted) else submitted
|
||||
except QueueFull as exc:
|
||||
raise HTTPException(status_code=429, detail=str(exc)) from exc
|
||||
|
||||
@@ -505,8 +680,33 @@ async def set_inbound_enabled(request: InboundEnableRequest) -> dict:
|
||||
"machine. Change that environment setting and restart VoiceStudio."
|
||||
),
|
||||
)
|
||||
|
||||
requested_bind = (
|
||||
inbound_service.normalise_bind_host(request.bind)
|
||||
if request.bind
|
||||
else inbound_service.bind_host()
|
||||
)
|
||||
requested_port = request.port or inbound_service.bind_port()
|
||||
if (
|
||||
request.enabled
|
||||
and inbound_service.node.running
|
||||
and (
|
||||
requested_bind != inbound_service.bind_host()
|
||||
or requested_port != inbound_service.node.port
|
||||
)
|
||||
):
|
||||
# start() is intentionally idempotent while a listener owns its
|
||||
# socket. Persisting a new endpoint here would make the UI report a
|
||||
# narrower/different bind while the original socket stayed live.
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail=(
|
||||
"Turn off Accept connections before changing its bind address "
|
||||
"or port."
|
||||
),
|
||||
)
|
||||
if request.bind:
|
||||
inbound_service.set_bind_host(request.bind)
|
||||
inbound_service.set_bind_host(requested_bind)
|
||||
if request.port:
|
||||
inbound_service.set_bind_port(request.port)
|
||||
inbound_service.set_enabled(request.enabled)
|
||||
@@ -535,6 +735,7 @@ def issue_inbound_key(request: IssueKeyRequest) -> dict:
|
||||
is stored, so it cannot be shown again, only replaced.
|
||||
"""
|
||||
from worker.inbound import service as inbound_service # noqa: PLC0415
|
||||
from worker.inbound.keys import KeyLimitExceeded # noqa: PLC0415
|
||||
|
||||
if not inbound_service.node.running:
|
||||
raise HTTPException(
|
||||
@@ -544,7 +745,10 @@ def issue_inbound_key(request: IssueKeyRequest) -> dict:
|
||||
"Settings → System → Remote workers → Accept connections first."
|
||||
),
|
||||
)
|
||||
issued = inbound_service.node.keys.issue(request.label)
|
||||
try:
|
||||
issued = inbound_service.node.keys.issue(request.label)
|
||||
except KeyLimitExceeded as exc:
|
||||
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
||||
return {
|
||||
"key_id": issued.key.key_id,
|
||||
"label": issued.key.label,
|
||||
@@ -555,12 +759,12 @@ def issue_inbound_key(request: IssueKeyRequest) -> dict:
|
||||
|
||||
|
||||
@router.delete("/inbound/keys/{key_id}")
|
||||
def revoke_inbound_key(key_id: str) -> dict:
|
||||
async def revoke_inbound_key(key_id: str) -> dict:
|
||||
"""Revoke one panel. Everyone else stays connected — the whole reason keys
|
||||
are per panel rather than one shared node key."""
|
||||
from worker.inbound import service as inbound_service # noqa: PLC0415
|
||||
|
||||
if not inbound_service.node.keys.revoke(key_id):
|
||||
if not await inbound_service.node.revoke_key(key_id):
|
||||
raise HTTPException(status_code=404, detail="No such key.")
|
||||
return inbound_service.node.snapshot()
|
||||
|
||||
@@ -579,6 +783,7 @@ async def add_inbound_connection(request: ConnectRequest) -> dict:
|
||||
"""Paste a connection string from a GPU machine and dial it."""
|
||||
from worker.inbound import service as inbound_service # noqa: PLC0415
|
||||
from worker.inbound.connection_string import InvalidConnectionString # noqa: PLC0415
|
||||
from worker.inbound.connector import InboundConnectionError # noqa: PLC0415
|
||||
|
||||
if not service.control_plane.running:
|
||||
raise HTTPException(
|
||||
@@ -597,12 +802,18 @@ async def add_inbound_connection(request: ConnectRequest) -> dict:
|
||||
# surfaces as "cannot connect", which is what a firewall, a wrong port
|
||||
# and a dead node all say too.
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
except InboundConnectionError as exc:
|
||||
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
||||
return {"endpoint": connection.endpoint, "connections": inbound_service.outbound.snapshot()}
|
||||
|
||||
|
||||
@router.delete("/inbound/connections/{endpoint}")
|
||||
async def remove_inbound_connection(endpoint: str) -> dict:
|
||||
from worker.inbound import service as inbound_service # noqa: PLC0415
|
||||
from worker.inbound.connector import InboundConnectionError # noqa: PLC0415
|
||||
|
||||
await inbound_service.outbound.remove(endpoint)
|
||||
try:
|
||||
await inbound_service.outbound.remove(endpoint)
|
||||
except InboundConnectionError as exc:
|
||||
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
||||
return {"connections": inbound_service.outbound.snapshot()}
|
||||
|
||||
+10
-10
@@ -159,17 +159,16 @@ models:
|
||||
- repo_id: "csukuangfj/sherpa-onnx-nemo-parakeet-tdt-0.6b-v3-int8"
|
||||
label: "Parakeet TDT v3 (sherpa-onnx — dictation, 25 EU langs)"
|
||||
role: ASR
|
||||
size_gb: 0.18
|
||||
size_gb: 0.67
|
||||
engine: sherpa-onnx
|
||||
dictation_id: sherpa-parakeet-tdt-v3
|
||||
tag: offline
|
||||
curated_on: [all]
|
||||
note: "Recommended live-dictation default. CPU, int8 ONNX. Requires sherpa-onnx."
|
||||
note: "Multilingual European-language dictation. CPU, int8 ONNX. Requires sherpa-onnx."
|
||||
|
||||
- repo_id: "csukuangfj/sherpa-onnx-nemo-parakeet-tdt-0.6b-v2-int8"
|
||||
label: "Parakeet TDT v2 (sherpa-onnx — dictation, English)"
|
||||
role: ASR
|
||||
size_gb: 0.17
|
||||
size_gb: 0.66
|
||||
engine: sherpa-onnx
|
||||
dictation_id: sherpa-parakeet-tdt-v2
|
||||
tag: offline
|
||||
@@ -178,7 +177,7 @@ models:
|
||||
- repo_id: "csukuangfj/sherpa-onnx-streaming-zipformer-bilingual-zh-en-2023-02-20"
|
||||
label: "Zipformer Bilingual (sherpa-onnx — streaming, zh+en)"
|
||||
role: ASR
|
||||
size_gb: 0.13
|
||||
size_gb: 0.2
|
||||
engine: sherpa-onnx
|
||||
dictation_id: sherpa-zipformer-bilingual-zh-en
|
||||
tag: streaming
|
||||
@@ -187,7 +186,7 @@ models:
|
||||
- repo_id: "csukuangfj/sherpa-onnx-streaming-paraformer-bilingual-zh-en"
|
||||
label: "Paraformer Bilingual (sherpa-onnx — streaming, zh+en)"
|
||||
role: ASR
|
||||
size_gb: 0.115
|
||||
size_gb: 0.24
|
||||
engine: sherpa-onnx
|
||||
dictation_id: sherpa-paraformer-bilingual-zh-en
|
||||
tag: streaming
|
||||
@@ -196,7 +195,7 @@ models:
|
||||
- repo_id: "csukuangfj/sherpa-onnx-streaming-zipformer-en-20M-2023-02-17"
|
||||
label: "Zipformer Streaming EN 20M (sherpa-onnx — streaming, English)"
|
||||
role: ASR
|
||||
size_gb: 0.128
|
||||
size_gb: 0.044
|
||||
engine: sherpa-onnx
|
||||
dictation_id: sherpa-zipformer-en-20m
|
||||
tag: streaming
|
||||
@@ -205,7 +204,7 @@ models:
|
||||
- repo_id: "csukuangfj/sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23"
|
||||
label: "Zipformer Streaming ZH 14M (sherpa-onnx — streaming, Chinese)"
|
||||
role: ASR
|
||||
size_gb: 0.074
|
||||
size_gb: 0.025
|
||||
engine: sherpa-onnx
|
||||
dictation_id: sherpa-zipformer-zh-14m
|
||||
tag: streaming
|
||||
@@ -214,11 +213,12 @@ models:
|
||||
- repo_id: "csukuangfj/sherpa-onnx-whisper-tiny"
|
||||
label: "Whisper Tiny (sherpa-onnx — dictation, 90+ langs)"
|
||||
role: ASR
|
||||
size_gb: 0.116
|
||||
size_gb: 0.104
|
||||
engine: sherpa-onnx
|
||||
dictation_id: sherpa-whisper-tiny
|
||||
tag: offline
|
||||
note: "Multilingual offline dictation (auto-detect). CPU, int8 ONNX. Requires sherpa-onnx."
|
||||
curated_on: [all]
|
||||
note: "Recommended cross-platform dictation default (auto-detect). CPU, int8 ONNX. Requires sherpa-onnx."
|
||||
|
||||
# ── Diarisation ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -0,0 +1,609 @@
|
||||
"""Nested subprocess ownership for desktop-managed backend operations.
|
||||
|
||||
The desktop owns the backend with an OS process group/Job. Engine and
|
||||
installer operations also need an independently terminable subtree: killing
|
||||
only their direct child on a timeout leaves uv/git/model workers holding pipes
|
||||
and mutating files. A small direct-child supervisor bridges both lifetimes.
|
||||
|
||||
On POSIX the supervisor is the unreaped leader of a nested process group. A
|
||||
control-pipe EOF (including kernel EOF when the backend dies) kills that group;
|
||||
the parent also drains the group before reaping its stable leader. On Windows
|
||||
the supervisor assigns the operation, while suspended, to a nested
|
||||
kill-on-close Job. The outer desktop Job still contains both levels.
|
||||
|
||||
Standalone/server launches use the same nested owner, preserving their
|
||||
independently terminable subtree without relying on ``taskkill`` or discovery.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import signal
|
||||
import struct
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
|
||||
_RESULT = struct.Struct("!i")
|
||||
_DESKTOP_MARKER = "OMNIVOICE_DESKTOP_CONTAINED"
|
||||
_DRAIN_FD_ENV = "OMNIVOICE_DESKTOP_DRAIN_FD"
|
||||
|
||||
|
||||
def backend_drain_fd(*, required: bool = False) -> Optional[int]:
|
||||
"""Validated Rust-owned drain writer inherited by the desktop backend."""
|
||||
if os.name != "posix" or os.environ.get(_DESKTOP_MARKER) != "1":
|
||||
return None
|
||||
try:
|
||||
fd = int(os.environ[_DRAIN_FD_ENV])
|
||||
os.fstat(fd)
|
||||
except (KeyError, ValueError, OSError) as exc:
|
||||
if required:
|
||||
raise RuntimeError(
|
||||
"desktop backend is missing its live nested-operation drain descriptor"
|
||||
) from exc
|
||||
return None
|
||||
return fd
|
||||
|
||||
|
||||
def secure_backend_drain_fd() -> None:
|
||||
"""Restore CLOEXEC after Rust's one intentional backend inheritance."""
|
||||
fd = backend_drain_fd(required=True)
|
||||
if fd is not None:
|
||||
os.set_inheritable(fd, False)
|
||||
|
||||
|
||||
class OwnedPopen:
|
||||
"""Popen-compatible handle for a desktop-owned nested operation."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
proc: subprocess.Popen,
|
||||
control_fd: int,
|
||||
result_fd: int,
|
||||
) -> None:
|
||||
self._proc = proc
|
||||
self._control_fd: Optional[int] = control_fd
|
||||
self._result_fd: Optional[int] = result_fd
|
||||
self._returncode: Optional[int] = None
|
||||
self._lock = threading.RLock()
|
||||
|
||||
# Popen callers use these directly (protocol pipes and log drains).
|
||||
self.stdin = proc.stdin
|
||||
self.stdout = proc.stdout
|
||||
self.stderr = proc.stderr
|
||||
|
||||
@property
|
||||
def pid(self) -> int:
|
||||
return self._proc.pid
|
||||
|
||||
@property
|
||||
def args(self) -> Any:
|
||||
return self._proc.args
|
||||
|
||||
@property
|
||||
def returncode(self) -> Optional[int]:
|
||||
return self._returncode
|
||||
|
||||
def _close_control(self) -> None:
|
||||
fd, self._control_fd = self._control_fd, None
|
||||
if fd is not None:
|
||||
try:
|
||||
os.close(fd)
|
||||
except OSError:
|
||||
# Cleanup is idempotent; another teardown path already closed it.
|
||||
pass
|
||||
|
||||
def _read_result(self, fallback: int) -> int:
|
||||
fd, self._result_fd = self._result_fd, None
|
||||
if fd is None:
|
||||
return fallback
|
||||
try:
|
||||
payload = b""
|
||||
while len(payload) < _RESULT.size:
|
||||
chunk = os.read(fd, _RESULT.size - len(payload))
|
||||
if not chunk:
|
||||
break
|
||||
payload += chunk
|
||||
return _RESULT.unpack(payload)[0] if len(payload) == _RESULT.size else fallback
|
||||
except OSError:
|
||||
return fallback
|
||||
finally:
|
||||
try:
|
||||
os.close(fd)
|
||||
except OSError:
|
||||
# The descriptor may have been closed by cancellation cleanup.
|
||||
pass
|
||||
|
||||
def _posix_exited_unreaped(self) -> bool:
|
||||
flags = os.WEXITED | os.WNOHANG | os.WNOWAIT
|
||||
info = os.waitid(os.P_PID, self.pid, flags)
|
||||
return info is not None and info.si_pid != 0
|
||||
|
||||
def _posix_exited_reaping(self) -> Optional[int]:
|
||||
"""macOS fallback for :meth:`_posix_exited_unreaped` (#1656).
|
||||
|
||||
CPython on macOS does not expose ``os.waitid`` (HAVE_WAITID is not set
|
||||
in its build), so the WNOWAIT probe is unavailable there. This
|
||||
fallback *reaps* the wrapper with ``waitpid(WNOHANG)``: it returns
|
||||
the wrapper's exit code once it has exited, None while it is still
|
||||
running, and raises ``ChildProcessError`` when another owner already
|
||||
reaped it (the same refusal the waitid probe gives).
|
||||
|
||||
Reaping earlier than the WNOWAIT dance loses the pre-reap group kill
|
||||
in :meth:`poll`; that is safe because the supervisor's control-pipe
|
||||
EOF already terminates the whole nested group (#1635 design).
|
||||
"""
|
||||
pid, status = os.waitpid(self.pid, os.WNOHANG)
|
||||
if pid != self.pid:
|
||||
return None
|
||||
rc = os.waitstatus_to_exitcode(status)
|
||||
# Publish on the underlying Popen so its own wait()/poll() no-op.
|
||||
self._proc.returncode = rc
|
||||
return rc
|
||||
|
||||
def _posix_exit_state_reaping(self) -> Optional[int]:
|
||||
""":meth:`_posix_exited_reaping` plus one concession: if the leader
|
||||
was already reaped through *this* Popen (``_proc.returncode`` known),
|
||||
report that code rather than refusing — reaping by our own handle is
|
||||
not the foreign reaper the ECHILD refusal exists for."""
|
||||
try:
|
||||
return self._posix_exited_reaping()
|
||||
except ChildProcessError:
|
||||
return self._proc.returncode
|
||||
|
||||
def _signal_owned_group(self, sig: int) -> None:
|
||||
# The numeric group is safe only while its direct-child leader remains
|
||||
# ours and unreaped. ECHILD therefore refuses rather than guessing.
|
||||
try:
|
||||
os.waitid(os.P_PID, self.pid, os.WEXITED | os.WNOHANG | os.WNOWAIT)
|
||||
except ChildProcessError:
|
||||
return
|
||||
except AttributeError:
|
||||
# macOS CPython has no os.waitid (#1656). waitpid still proves
|
||||
# that this exact numeric pid is our live child: ECHILD refuses a
|
||||
# foreign-reaped/reused pid, while pid == self.pid records an exit
|
||||
# without ever signalling the now-unowned process-group number.
|
||||
try:
|
||||
pid, status = os.waitpid(self.pid, os.WNOHANG)
|
||||
except ChildProcessError:
|
||||
return
|
||||
if pid == self.pid:
|
||||
self._proc.returncode = os.waitstatus_to_exitcode(status)
|
||||
return
|
||||
try:
|
||||
os.killpg(self.pid, sig)
|
||||
except ProcessLookupError:
|
||||
# The owned group exited between the waitid probe and the signal.
|
||||
pass
|
||||
|
||||
def poll(self) -> Optional[int]:
|
||||
with self._lock:
|
||||
if self._returncode is not None:
|
||||
return self._returncode
|
||||
if os.name == "posix":
|
||||
try:
|
||||
if hasattr(os, "waitid"):
|
||||
if not self._posix_exited_unreaped():
|
||||
return None
|
||||
self._signal_owned_group(signal.SIGKILL)
|
||||
wrapper_rc = self._proc.wait()
|
||||
else:
|
||||
# macOS CPython: no os.waitid (#1656) — the reaping
|
||||
# probe already terminated/killed nothing; the group
|
||||
# is torn down by the control-pipe EOF in _close_control.
|
||||
wrapper_rc = self._posix_exit_state_reaping()
|
||||
if wrapper_rc is None:
|
||||
return None
|
||||
except ChildProcessError:
|
||||
# Never signal a potentially reused group after another
|
||||
# owner reaped the stable leader.
|
||||
return None
|
||||
else:
|
||||
wrapper_rc = self._proc.poll()
|
||||
if wrapper_rc is None:
|
||||
return None
|
||||
self._close_control()
|
||||
self._returncode = self._read_result(wrapper_rc)
|
||||
return self._returncode
|
||||
|
||||
def wait(self, timeout: Optional[float] = None) -> int:
|
||||
deadline = None if timeout is None else time.monotonic() + timeout
|
||||
while True:
|
||||
rc = self.poll()
|
||||
if rc is not None:
|
||||
return rc
|
||||
if deadline is not None and time.monotonic() >= deadline:
|
||||
raise subprocess.TimeoutExpired(self.args, timeout)
|
||||
time.sleep(0.01)
|
||||
|
||||
def terminate(self) -> None:
|
||||
with self._lock:
|
||||
if self._returncode is not None:
|
||||
return
|
||||
self._close_control()
|
||||
if os.name == "posix":
|
||||
self._signal_owned_group(signal.SIGTERM)
|
||||
else:
|
||||
# Closing the control pipe asks the supervisor to terminate
|
||||
# its nested Job. The stable wrapper handle is a fallback.
|
||||
try:
|
||||
self._proc.terminate()
|
||||
except OSError:
|
||||
# The wrapper exited after the return-code check.
|
||||
pass
|
||||
|
||||
def kill(self) -> None:
|
||||
with self._lock:
|
||||
if self._returncode is not None:
|
||||
return
|
||||
self._close_control()
|
||||
if os.name == "posix":
|
||||
self._signal_owned_group(signal.SIGKILL)
|
||||
else:
|
||||
try:
|
||||
self._proc.kill()
|
||||
except OSError:
|
||||
# The wrapper exited after the return-code check.
|
||||
pass
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
return getattr(self._proc, name)
|
||||
|
||||
def __del__(self) -> None:
|
||||
self._close_control()
|
||||
fd, self._result_fd = self._result_fd, None
|
||||
if fd is not None:
|
||||
try:
|
||||
os.close(fd)
|
||||
except OSError:
|
||||
# Finalization may race explicit wait or cancellation cleanup.
|
||||
pass
|
||||
|
||||
|
||||
def spawn_owned(argv: list[str], **kwargs: Any) -> "subprocess.Popen | OwnedPopen":
|
||||
"""Spawn an operation with a stable, independently terminable owner."""
|
||||
|
||||
drain_fd = backend_drain_fd(required=True) if os.name == "posix" else None
|
||||
control_read, control_write = os.pipe()
|
||||
result_read, result_write = os.pipe()
|
||||
control_token = control_read
|
||||
result_token = result_write
|
||||
if os.name == "nt":
|
||||
import msvcrt
|
||||
|
||||
control_token = msvcrt.get_osfhandle(control_read)
|
||||
result_token = msvcrt.get_osfhandle(result_write)
|
||||
wrapper_argv = _supervisor_argv(
|
||||
control_token,
|
||||
result_token,
|
||||
argv,
|
||||
)
|
||||
wrapper_kwargs = dict(kwargs)
|
||||
if os.name == "posix":
|
||||
wrapper_kwargs["start_new_session"] = True
|
||||
pass_fds = [control_read, result_write]
|
||||
if drain_fd is not None:
|
||||
pass_fds.append(drain_fd)
|
||||
if wrapper_kwargs.get("env") is not None:
|
||||
wrapper_env = dict(wrapper_kwargs["env"])
|
||||
wrapper_env[_DESKTOP_MARKER] = "1"
|
||||
wrapper_env[_DRAIN_FD_ENV] = str(drain_fd)
|
||||
wrapper_kwargs["env"] = wrapper_env
|
||||
wrapper_kwargs["pass_fds"] = tuple(pass_fds)
|
||||
else:
|
||||
# Python's Windows fd inheritance requires inheritable CRT handles.
|
||||
# All unrelated descriptors are non-inheritable by default (PEP 446).
|
||||
os.set_handle_inheritable(control_token, True)
|
||||
os.set_handle_inheritable(result_token, True)
|
||||
wrapper_kwargs["close_fds"] = False
|
||||
try:
|
||||
proc = subprocess.Popen(wrapper_argv, **wrapper_kwargs)
|
||||
except BaseException:
|
||||
# The finally block exclusively owns the child-side endpoints. Closing
|
||||
# them here as well risks closing a reused descriptor in another thread.
|
||||
for fd in (control_write, result_read):
|
||||
try:
|
||||
os.close(fd)
|
||||
except OSError:
|
||||
# A partial spawn may already have closed a parent-side endpoint.
|
||||
pass
|
||||
raise
|
||||
finally:
|
||||
for fd in (control_read, result_write):
|
||||
try:
|
||||
os.close(fd)
|
||||
except OSError:
|
||||
# Popen may have consumed an inherited child-side endpoint.
|
||||
pass
|
||||
return OwnedPopen(proc, control_write, result_read)
|
||||
|
||||
|
||||
def _supervisor_argv(
|
||||
control_token: int,
|
||||
result_token: int,
|
||||
argv: list[str],
|
||||
) -> list[str]:
|
||||
prefix = [sys.executable]
|
||||
if not getattr(sys, "frozen", False):
|
||||
prefix.append(str(Path(__file__).resolve().parents[1] / "main.py"))
|
||||
return [
|
||||
*prefix,
|
||||
"--supervise",
|
||||
str(control_token),
|
||||
str(result_token),
|
||||
"--",
|
||||
*map(str, argv),
|
||||
]
|
||||
|
||||
|
||||
def _write_result(fd: int, returncode: int) -> None:
|
||||
try:
|
||||
os.write(fd, _RESULT.pack(int(returncode)))
|
||||
except OSError:
|
||||
# The caller may have cancelled and closed its result reader.
|
||||
pass
|
||||
finally:
|
||||
try:
|
||||
os.close(fd)
|
||||
except OSError:
|
||||
# Writing or cancellation may already have closed the descriptor.
|
||||
pass
|
||||
|
||||
|
||||
def _operation_env() -> dict[str, str]:
|
||||
env = os.environ.copy()
|
||||
# The operation intentionally does not own the Rust drain writer. Avoid
|
||||
# exposing a stale numeric token which nested code could mistake as valid.
|
||||
env.pop(_DRAIN_FD_ENV, None)
|
||||
env.pop(_DESKTOP_MARKER, None)
|
||||
return env
|
||||
|
||||
|
||||
def _supervise_posix(control_fd: int, result_fd: int, argv: list[str]) -> int:
|
||||
def cancel_on_eof() -> None:
|
||||
try:
|
||||
while os.read(control_fd, 1):
|
||||
pass
|
||||
except OSError:
|
||||
# Closing the control descriptor is itself a cancellation signal.
|
||||
pass
|
||||
os.killpg(os.getpgrp(), signal.SIGKILL)
|
||||
|
||||
threading.Thread(target=cancel_on_eof, daemon=True).start()
|
||||
try:
|
||||
child = subprocess.Popen(argv, close_fds=True, env=_operation_env())
|
||||
rc = child.wait()
|
||||
except OSError:
|
||||
rc = 127
|
||||
_write_result(result_fd, rc)
|
||||
# Drain children which outlived the operation before the stable group
|
||||
# leader exits. SIGKILL intentionally includes this supervisor.
|
||||
os.killpg(os.getpgrp(), signal.SIGKILL)
|
||||
return rc # unreachable
|
||||
|
||||
|
||||
def _windows_job() -> tuple[Any, Any, Any]:
|
||||
import ctypes
|
||||
import ctypes.wintypes as wintypes
|
||||
|
||||
kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
|
||||
kernel32.CloseHandle.argtypes = (wintypes.HANDLE,)
|
||||
kernel32.CloseHandle.restype = wintypes.BOOL
|
||||
kernel32.TerminateJobObject.argtypes = (wintypes.HANDLE, wintypes.UINT)
|
||||
kernel32.TerminateJobObject.restype = wintypes.BOOL
|
||||
kernel32.ReadFile.argtypes = (
|
||||
wintypes.HANDLE,
|
||||
ctypes.c_void_p,
|
||||
wintypes.DWORD,
|
||||
ctypes.POINTER(wintypes.DWORD),
|
||||
ctypes.c_void_p,
|
||||
)
|
||||
kernel32.ReadFile.restype = wintypes.BOOL
|
||||
kernel32.WriteFile.argtypes = (
|
||||
wintypes.HANDLE,
|
||||
ctypes.c_void_p,
|
||||
wintypes.DWORD,
|
||||
ctypes.POINTER(wintypes.DWORD),
|
||||
ctypes.c_void_p,
|
||||
)
|
||||
kernel32.WriteFile.restype = wintypes.BOOL
|
||||
create = kernel32.CreateJobObjectW
|
||||
create.argtypes = (ctypes.c_void_p, wintypes.LPCWSTR)
|
||||
create.restype = wintypes.HANDLE
|
||||
job = create(None, None)
|
||||
if not job:
|
||||
raise OSError(ctypes.get_last_error(), "CreateJobObjectW")
|
||||
|
||||
class BasicLimits(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("PerProcessUserTimeLimit", ctypes.c_longlong),
|
||||
("PerJobUserTimeLimit", ctypes.c_longlong),
|
||||
("LimitFlags", wintypes.DWORD),
|
||||
("MinimumWorkingSetSize", ctypes.c_size_t),
|
||||
("MaximumWorkingSetSize", ctypes.c_size_t),
|
||||
("ActiveProcessLimit", wintypes.DWORD),
|
||||
("Affinity", ctypes.c_size_t),
|
||||
("PriorityClass", wintypes.DWORD),
|
||||
("SchedulingClass", wintypes.DWORD),
|
||||
]
|
||||
|
||||
class IoCounters(ctypes.Structure):
|
||||
_fields_ = [(name, ctypes.c_ulonglong) for name in (
|
||||
"ReadOperationCount", "WriteOperationCount", "OtherOperationCount",
|
||||
"ReadTransferCount", "WriteTransferCount", "OtherTransferCount",
|
||||
)]
|
||||
|
||||
class ExtendedLimits(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("BasicLimitInformation", BasicLimits),
|
||||
("IoInfo", IoCounters),
|
||||
("ProcessMemoryLimit", ctypes.c_size_t),
|
||||
("JobMemoryLimit", ctypes.c_size_t),
|
||||
("PeakProcessMemoryUsed", ctypes.c_size_t),
|
||||
("PeakJobMemoryUsed", ctypes.c_size_t),
|
||||
]
|
||||
|
||||
info = ExtendedLimits()
|
||||
info.BasicLimitInformation.LimitFlags = 0x00002000 # KILL_ON_JOB_CLOSE
|
||||
set_info = kernel32.SetInformationJobObject
|
||||
set_info.argtypes = (wintypes.HANDLE, ctypes.c_int, ctypes.c_void_p, wintypes.DWORD)
|
||||
set_info.restype = wintypes.BOOL
|
||||
if not set_info(job, 9, ctypes.byref(info), ctypes.sizeof(info)):
|
||||
error = ctypes.get_last_error()
|
||||
kernel32.CloseHandle(job)
|
||||
raise OSError(error, "SetInformationJobObject")
|
||||
return job, kernel32, wintypes
|
||||
|
||||
|
||||
def _resume_windows_process(kernel32: Any, wintypes: Any, pid: int) -> None:
|
||||
import ctypes
|
||||
|
||||
class ThreadEntry(ctypes.Structure):
|
||||
_fields_ = [
|
||||
("dwSize", wintypes.DWORD),
|
||||
("cntUsage", wintypes.DWORD),
|
||||
("th32ThreadID", wintypes.DWORD),
|
||||
("th32OwnerProcessID", wintypes.DWORD),
|
||||
("tpBasePri", wintypes.LONG),
|
||||
("tpDeltaPri", wintypes.LONG),
|
||||
("dwFlags", wintypes.DWORD),
|
||||
]
|
||||
|
||||
kernel32.CreateToolhelp32Snapshot.argtypes = (wintypes.DWORD, wintypes.DWORD)
|
||||
kernel32.CreateToolhelp32Snapshot.restype = wintypes.HANDLE
|
||||
kernel32.Thread32First.argtypes = (wintypes.HANDLE, ctypes.POINTER(ThreadEntry))
|
||||
kernel32.Thread32First.restype = wintypes.BOOL
|
||||
kernel32.Thread32Next.argtypes = (wintypes.HANDLE, ctypes.POINTER(ThreadEntry))
|
||||
kernel32.Thread32Next.restype = wintypes.BOOL
|
||||
kernel32.OpenThread.argtypes = (wintypes.DWORD, wintypes.BOOL, wintypes.DWORD)
|
||||
kernel32.OpenThread.restype = wintypes.HANDLE
|
||||
kernel32.ResumeThread.argtypes = (wintypes.HANDLE,)
|
||||
kernel32.ResumeThread.restype = wintypes.DWORD
|
||||
|
||||
snapshot = kernel32.CreateToolhelp32Snapshot(0x00000004, 0)
|
||||
invalid = ctypes.c_void_p(-1).value
|
||||
if snapshot == invalid:
|
||||
raise OSError(ctypes.get_last_error(), "CreateToolhelp32Snapshot")
|
||||
try:
|
||||
entry = ThreadEntry(dwSize=ctypes.sizeof(ThreadEntry))
|
||||
found = kernel32.Thread32First(snapshot, ctypes.byref(entry))
|
||||
while found:
|
||||
if entry.th32OwnerProcessID == pid:
|
||||
thread = kernel32.OpenThread(0x0002, False, entry.th32ThreadID)
|
||||
if not thread:
|
||||
raise OSError(ctypes.get_last_error(), "OpenThread")
|
||||
try:
|
||||
if kernel32.ResumeThread(thread) == 0xFFFFFFFF:
|
||||
raise OSError(ctypes.get_last_error(), "ResumeThread")
|
||||
return
|
||||
finally:
|
||||
kernel32.CloseHandle(thread)
|
||||
found = kernel32.Thread32Next(snapshot, ctypes.byref(entry))
|
||||
finally:
|
||||
kernel32.CloseHandle(snapshot)
|
||||
raise OSError("suspended operation thread was not found")
|
||||
|
||||
|
||||
def _supervise_windows(control_fd: int, result_fd: int, argv: list[str]) -> int:
|
||||
import ctypes
|
||||
|
||||
job, kernel32, wintypes = _windows_job()
|
||||
cancelled = threading.Event()
|
||||
job_lock = threading.Lock()
|
||||
job_open = True
|
||||
|
||||
def terminate_job() -> None:
|
||||
with job_lock:
|
||||
if job_open:
|
||||
kernel32.TerminateJobObject(job, 1)
|
||||
|
||||
def cancel_on_eof() -> None:
|
||||
byte = ctypes.create_string_buffer(1)
|
||||
count = wintypes.DWORD()
|
||||
while kernel32.ReadFile(
|
||||
wintypes.HANDLE(control_fd), byte, 1, ctypes.byref(count), None
|
||||
) and count.value:
|
||||
pass
|
||||
kernel32.CloseHandle(wintypes.HANDLE(control_fd))
|
||||
cancelled.set()
|
||||
terminate_job()
|
||||
|
||||
threading.Thread(target=cancel_on_eof, daemon=True).start()
|
||||
child: Optional[subprocess.Popen] = None
|
||||
rc = 127
|
||||
try:
|
||||
child = subprocess.Popen(
|
||||
argv,
|
||||
close_fds=True,
|
||||
env=_operation_env(),
|
||||
creationflags=0x08000000 | 0x00000004, # NO_WINDOW | SUSPENDED
|
||||
)
|
||||
assign = kernel32.AssignProcessToJobObject
|
||||
assign.argtypes = (wintypes.HANDLE, wintypes.HANDLE)
|
||||
assign.restype = wintypes.BOOL
|
||||
if not assign(job, wintypes.HANDLE(child._handle)):
|
||||
raise OSError(ctypes.get_last_error(), "AssignProcessToJobObject")
|
||||
if cancelled.is_set():
|
||||
terminate_job()
|
||||
else:
|
||||
_resume_windows_process(kernel32, wintypes, child.pid)
|
||||
rc = child.wait()
|
||||
# A successful direct child may leave helpers behind; terminate the
|
||||
# nested stable Job before reporting completion.
|
||||
terminate_job()
|
||||
except OSError:
|
||||
terminate_job()
|
||||
if child is not None:
|
||||
try:
|
||||
# Assignment itself may have failed, leaving this suspended
|
||||
# process outside the nested Job. Terminate it through its
|
||||
# stable process handle before waiting; never strand an
|
||||
# unassigned operation or rely on the outer desktop Job.
|
||||
child.kill()
|
||||
except OSError:
|
||||
# The suspended child may have exited during Job teardown.
|
||||
pass
|
||||
try:
|
||||
child.wait(timeout=5)
|
||||
except (OSError, subprocess.TimeoutExpired):
|
||||
# The outer desktop Job remains the terminal containment fallback.
|
||||
pass
|
||||
finally:
|
||||
payload = _RESULT.pack(int(rc))
|
||||
payload_buffer = ctypes.create_string_buffer(payload)
|
||||
written = wintypes.DWORD()
|
||||
kernel32.WriteFile(
|
||||
wintypes.HANDLE(result_fd),
|
||||
payload_buffer,
|
||||
len(payload),
|
||||
ctypes.byref(written),
|
||||
None,
|
||||
)
|
||||
kernel32.CloseHandle(wintypes.HANDLE(result_fd))
|
||||
with job_lock:
|
||||
job_open = False
|
||||
kernel32.CloseHandle(job)
|
||||
return rc
|
||||
|
||||
|
||||
def supervisor_main(args: list[str]) -> int:
|
||||
if len(args) < 5 or args[0] != "--supervise" or args[3] != "--":
|
||||
return 2
|
||||
control_fd = int(args[1])
|
||||
result_fd = int(args[2])
|
||||
argv = args[4:]
|
||||
secure_backend_drain_fd()
|
||||
if os.name == "posix":
|
||||
return _supervise_posix(control_fd, result_fd, argv)
|
||||
return _supervise_windows(control_fd, result_fd, argv)
|
||||
|
||||
|
||||
def _main() -> int:
|
||||
return supervisor_main(sys.argv[1:])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(_main())
|
||||
@@ -177,6 +177,23 @@ def gfx_for_hsa_override(value: str) -> str | None:
|
||||
#: The ROCm kernel driver interface. Its absence, or its presence without
|
||||
#: permission, are the two commonest reasons a ROCm host silently runs on CPU.
|
||||
_KFD_DEVICE = "/dev/kfd"
|
||||
_DXG_DEVICE = "/dev/dxg"
|
||||
_DXG_RUNTIME_PATHS = (
|
||||
"/usr/lib/libdxcore.so",
|
||||
"/usr/lib/librocdxg.so",
|
||||
"/usr/share/rocdxg/dids.conf",
|
||||
)
|
||||
|
||||
|
||||
def _rocm_requires_dxg_detection(version: object) -> bool:
|
||||
"""Whether WSL's ROCDXG bridge still needs its explicit opt-in."""
|
||||
try:
|
||||
parts = str(version).split(".")
|
||||
return (int(parts[0]), int(parts[1])) < (7, 13)
|
||||
except (IndexError, TypeError, ValueError):
|
||||
# Unknown versions get the conservative advice. The variable is
|
||||
# harmless on newer runtimes and necessary on every older one.
|
||||
return True
|
||||
|
||||
|
||||
def why_no_gpu(torch) -> tuple[str, ...]:
|
||||
@@ -230,6 +247,40 @@ def why_no_gpu(torch) -> tuple[str, ...]:
|
||||
# /dev/kfd only exists on Linux; on any other platform its absence
|
||||
# says nothing, so don't invent a reason.
|
||||
if sys.platform.startswith("linux"):
|
||||
if not os.path.exists(_KFD_DEVICE) and os.path.exists(_DXG_DEVICE):
|
||||
if not os.access(_DXG_DEVICE, os.R_OK | os.W_OK):
|
||||
return (
|
||||
f"ROCm {hip} is installed and {_DXG_DEVICE} exists, "
|
||||
"but this process cannot open it — pass "
|
||||
"--device /dev/dxg to the WSL container",
|
||||
)
|
||||
dxg_detection = os.environ.get("HSA_ENABLE_DXG_DETECTION", "").strip()
|
||||
if dxg_detection == "0":
|
||||
return (
|
||||
f"ROCm {hip} is installed and {_DXG_DEVICE} is reachable, "
|
||||
"but HSA_ENABLE_DXG_DETECTION=0 explicitly disables the "
|
||||
"WSL GPU bridge; remove it or set it to 1",
|
||||
)
|
||||
if _rocm_requires_dxg_detection(hip) and dxg_detection != "1":
|
||||
return (
|
||||
f"ROCm {hip} is installed and {_DXG_DEVICE} is "
|
||||
"reachable, but this pre-7.13 runtime requires "
|
||||
"HSA_ENABLE_DXG_DETECTION=1 inside WSL containers",
|
||||
)
|
||||
missing = [
|
||||
path for path in _DXG_RUNTIME_PATHS if not os.path.exists(path)
|
||||
]
|
||||
if missing:
|
||||
return (
|
||||
f"ROCm {hip} can reach {_DXG_DEVICE}, but the WSL "
|
||||
"ROCDXG runtime mounts are incomplete; missing: "
|
||||
f"{', '.join(missing)}",
|
||||
)
|
||||
return (
|
||||
f"ROCm {hip} and the WSL ROCDXG bridge are reachable, "
|
||||
"but no GPU was enumerated — verify the AMD Windows "
|
||||
"driver, librocdxg/ROCm compatibility, and host `rocminfo`",
|
||||
)
|
||||
if not os.path.exists(_KFD_DEVICE):
|
||||
return (
|
||||
f"ROCm {hip} is installed but {_KFD_DEVICE} is not "
|
||||
|
||||
@@ -23,9 +23,17 @@ logger = logging.getLogger("omnivoice.events")
|
||||
_listeners: list[asyncio.Queue] = []
|
||||
_lock = asyncio.Lock()
|
||||
|
||||
# The loop that serves /ws/events, captured on first use. Sync FastAPI
|
||||
# endpoints (rename/delete profile, revoke consent) run in threadpool workers
|
||||
# where `asyncio.get_running_loop()` raises, which used to silently drop their
|
||||
# events — the UI then never refetched the voice list (#1158 class).
|
||||
_serving_loop: asyncio.AbstractEventLoop | None = None
|
||||
|
||||
|
||||
async def subscribe() -> asyncio.Queue:
|
||||
"""Register a new listener. Returns a Queue that receives event dicts."""
|
||||
global _serving_loop
|
||||
_serving_loop = asyncio.get_running_loop()
|
||||
q: asyncio.Queue = asyncio.Queue(maxsize=64)
|
||||
async with _lock:
|
||||
_listeners.append(q)
|
||||
@@ -57,11 +65,29 @@ def emit(kind: str, payload: dict[str, Any] | None = None) -> None:
|
||||
}
|
||||
event_str = json.dumps(event)
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
loop.create_task(_broadcast(event_str))
|
||||
caller_loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
# No event loop running (unlikely in FastAPI context but safe)
|
||||
caller_loop = None
|
||||
target_loop = _serving_loop or caller_loop
|
||||
if target_loop is None:
|
||||
# No serving loop yet — nobody to notify; dropping is correct.
|
||||
logger.debug("No event loop — event dropped: %s", kind)
|
||||
return
|
||||
try:
|
||||
if caller_loop is target_loop:
|
||||
target_loop.create_task(_broadcast(event_str))
|
||||
else:
|
||||
# Sync endpoints and async producers on a foreign loop must both
|
||||
# hand off: the lock and listener queues belong to serving_loop.
|
||||
target_loop.call_soon_threadsafe(_schedule_broadcast, event_str)
|
||||
except RuntimeError:
|
||||
# The serving loop closed between capture and use (app shutdown).
|
||||
logger.debug("Event loop closed — event dropped: %s", kind)
|
||||
|
||||
|
||||
def _schedule_broadcast(event_str: str) -> None:
|
||||
"""Run `_broadcast` on the serving loop; called via call_soon_threadsafe."""
|
||||
asyncio.get_running_loop().create_task(_broadcast(event_str))
|
||||
|
||||
|
||||
async def _broadcast(event_str: str) -> None:
|
||||
@@ -73,11 +99,11 @@ async def _broadcast(event_str: str) -> None:
|
||||
q.put_nowait(event_str)
|
||||
except asyncio.QueueFull:
|
||||
# Slow consumer — drop oldest, then push. Not a race (#1163):
|
||||
# every queue op runs on the single event loop, and there is
|
||||
# no await between the QueueFull and this get_nowait/put_nowait
|
||||
# pair — no consumer can interleave, so get_nowait cannot raise
|
||||
# QueueEmpty here. emit() from a foreign thread drops the event
|
||||
# before ever touching a queue (see the RuntimeError branch).
|
||||
# every queue op runs on the single event loop (a foreign
|
||||
# thread's emit() hands off via call_soon_threadsafe first),
|
||||
# and there is no await between the QueueFull and this
|
||||
# get_nowait/put_nowait pair — no consumer can interleave, so
|
||||
# get_nowait cannot raise QueueEmpty here.
|
||||
try:
|
||||
q.get_nowait()
|
||||
q.put_nowait(event_str)
|
||||
|
||||
@@ -52,6 +52,7 @@ _REDACTED_VALUE = "***REDACTED***"
|
||||
# One-line "what to do" per docs-taxonomy key. Keys mirror error_docs_map's
|
||||
# taxonomy; the docs URL itself stays owned by error_docs_map.
|
||||
_HINTS: dict[str, str] = {
|
||||
"GPU_OOM": "Close other GPU-heavy apps or unload models, then retry. You can also choose CPU in Settings → Performance & Device or select a smaller TTS engine.",
|
||||
"WORKER_AT_CAPACITY": "Wait for a running job on that worker to finish, or choose another available worker and retry.",
|
||||
"MODEL_NOT_INSTALLED": "Install or enable this engine on the worker machine, then refresh its capabilities and retry.",
|
||||
"MODEL_NOT_DOWNLOADED": "Open Models, install this model on the selected worker, then retry when the download completes.",
|
||||
@@ -290,6 +291,9 @@ def append_hf_mirror_hint(text: str) -> str:
|
||||
# must NOT be added: its bare "timed out" trigger would stamp a "video server"
|
||||
# hint on a model-load timeout that leaks through the 500 handler.
|
||||
_CONTEXT_FREE_HINT_CLASSES = frozenset({
|
||||
# Device allocator signatures are specific enough to attach the shared
|
||||
# recovery without exposing CUDA's process table or filesystem paths.
|
||||
"GPU_OOM",
|
||||
"SOCKS_PROXY_SUPPORT_MISSING",
|
||||
"SSL_HANDSHAKE_FAILURE",
|
||||
# Its trigger is an exact OpenSSL string, so it cannot be confused with
|
||||
@@ -323,6 +327,38 @@ def append_hint(text: str) -> str:
|
||||
return f"{text} — {hint}" if hint else text
|
||||
|
||||
|
||||
_GPU_OOM_SIGNATURES = (
|
||||
"cuda out of memory",
|
||||
"cuda error: out of memory",
|
||||
"cuda_error_out_of_memory",
|
||||
"mps backend out of memory",
|
||||
"hip out of memory",
|
||||
"out of memory on device",
|
||||
)
|
||||
|
||||
|
||||
def is_gpu_oom(error: BaseException | str) -> bool:
|
||||
"""Recognize device OOMs through wrappers without importing torch."""
|
||||
pending: list[BaseException] = [error] if isinstance(error, BaseException) else []
|
||||
seen: set[int] = set()
|
||||
while pending:
|
||||
current = pending.pop()
|
||||
if id(current) in seen:
|
||||
continue
|
||||
seen.add(id(current))
|
||||
if type(current).__name__ == "OutOfMemoryError":
|
||||
return True
|
||||
if any(signature in str(current).lower() for signature in _GPU_OOM_SIGNATURES):
|
||||
return True
|
||||
if current.__cause__ is not None:
|
||||
pending.append(current.__cause__)
|
||||
if current.__context__ is not None:
|
||||
pending.append(current.__context__)
|
||||
if isinstance(error, str):
|
||||
return any(signature in error.lower() for signature in _GPU_OOM_SIGNATURES)
|
||||
return False
|
||||
|
||||
|
||||
def classify(reason: str) -> str:
|
||||
"""Map a failure reason to a docs-taxonomy key, or "" when unknown.
|
||||
|
||||
@@ -330,6 +366,8 @@ def classify(reason: str) -> str:
|
||||
backend log / diagnostic names the same class the UI deeplink will use.
|
||||
"""
|
||||
low = (reason or "").lower()
|
||||
if is_gpu_oom(low):
|
||||
return "GPU_OOM"
|
||||
if "pkg_resources" in low:
|
||||
return "PKG_RESOURCES_MISSING"
|
||||
if "quarantine" in low or "is damaged" in low or "gatekeeper" in low:
|
||||
|
||||
@@ -28,6 +28,15 @@ def stream_failure(code: str) -> dict[str, object]:
|
||||
"detail": "Generation capacity is busy. Try again shortly.",
|
||||
"retryable": True,
|
||||
},
|
||||
"generation_timeout": {
|
||||
"code": "generation_timeout",
|
||||
"detail": (
|
||||
"Generation exceeded the compute-time limit. The backend is "
|
||||
"still running; try a shorter passage or raise the generation "
|
||||
"timeout."
|
||||
),
|
||||
"retryable": True,
|
||||
},
|
||||
"invalid_request": {
|
||||
"code": "invalid_request",
|
||||
"detail": "The generation request could not be processed.",
|
||||
@@ -65,6 +74,48 @@ def stream_failure(code: str) -> dict[str, object]:
|
||||
return dict(failures.get(code, failures["generation_failed"]))
|
||||
|
||||
|
||||
def stream_generation_failure(error: BaseException | object) -> dict[str, object]:
|
||||
"""``generation_failed`` stream metadata, enriched with the actual cause.
|
||||
|
||||
The bare "Generation failed. Check the selected engine and try again." is
|
||||
the floor for an *unrecognized* failure. When the private exception DOES
|
||||
classify to a known failure class — a corrupt model cache, an unreachable
|
||||
Hugging Face mirror, a missing ffmpeg/ffprobe, a Windows paging-file limit,
|
||||
a SOCKS/TLS proxy problem, … — the stable VoiceStudio-owned remediation for
|
||||
that class is appended so the user can self-diagnose instead of guessing
|
||||
which engine or which failure. This is the same enrichment the classic
|
||||
(non-streaming) ``/generate`` 500 already gets via
|
||||
:func:`public_exception_response`; the in-band streaming error frame
|
||||
replaces the global 500 handler for a streaming request and used to bypass
|
||||
it entirely (#1607).
|
||||
|
||||
Only VoiceStudio-owned constants are copied — never a substring of
|
||||
``error`` (Constitution I). Never raises: a diagnosis failure must not
|
||||
replace the failure being diagnosed.
|
||||
"""
|
||||
payload = stream_failure("generation_failed")
|
||||
try:
|
||||
enriched = public_exception_response(error, fallback=str(payload["detail"]))
|
||||
except Exception:
|
||||
return payload
|
||||
hint = enriched.get("hint")
|
||||
if hint:
|
||||
payload["detail"] = enriched["detail"]
|
||||
payload["hint"] = hint
|
||||
topic = enriched.get("docs_topic")
|
||||
if topic:
|
||||
payload["docs_topic"] = topic
|
||||
try:
|
||||
from core import error_docs_map
|
||||
|
||||
url = error_docs_map.ERROR_DOCS.get(topic, "")
|
||||
except Exception:
|
||||
url = ""
|
||||
if url:
|
||||
payload["docs_url"] = url
|
||||
return payload
|
||||
|
||||
|
||||
def public_failure(
|
||||
logger: logging.Logger,
|
||||
log_message: str,
|
||||
|
||||
@@ -24,7 +24,7 @@ from pathlib import Path
|
||||
# tests/test_app_version.py::test_all_version_files_in_lockstep and bumped by
|
||||
# release.yml's version-bump job, so it stays equal to
|
||||
# pyproject/tauri.conf/Cargo/package.json.
|
||||
_FALLBACK_VERSION = "0.5.0"
|
||||
_FALLBACK_VERSION = "0.5.1"
|
||||
|
||||
|
||||
def _fallback_version() -> str:
|
||||
|
||||
@@ -28,6 +28,7 @@ packages. The parent only ever spawns it as a subprocess.
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
from typing import TYPE_CHECKING
|
||||
@@ -164,6 +165,23 @@ class IndexTTS2Backend(SubprocessBackend):
|
||||
from engines.indextts.bootstrap import resolve_indextts_venv
|
||||
return resolve_indextts_venv()
|
||||
|
||||
@property
|
||||
def recv_timeout_s(self) -> float:
|
||||
# IndexTTS was the only sidecar left on the 60s class default while
|
||||
# pockettts and omnivoice-subprocess both raised theirs. infer() is one
|
||||
# blocking upstream call, so a long passage legitimately outruns 60s and
|
||||
# the parent's watchdog killed a healthy synthesis (#1611). main.py also
|
||||
# heartbeats during infer(), which is what actually proves liveness —
|
||||
# this deadline is the ceiling for a sidecar that has gone genuinely
|
||||
# silent. OMNIVOICE_INDEXTTS_RECV_TIMEOUT_S tunes it.
|
||||
try:
|
||||
v = float(os.environ.get("OMNIVOICE_INDEXTTS_RECV_TIMEOUT_S", "900"))
|
||||
except (ValueError, TypeError):
|
||||
return 900.0
|
||||
if not math.isfinite(v): # reject inf/nan so the deadline can't be disabled
|
||||
return 900.0
|
||||
return max(30.0, v)
|
||||
|
||||
@classmethod
|
||||
def sidecar_script(cls):
|
||||
from engines.indextts.bootstrap import INDEXTTS_SIDECAR_SCRIPT
|
||||
|
||||
@@ -63,11 +63,13 @@ Restrictions:
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import contextlib
|
||||
import json
|
||||
import os
|
||||
import struct
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
import traceback
|
||||
|
||||
|
||||
@@ -117,11 +119,59 @@ EMOTION_KWARGS_ALLOWLIST = frozenset({
|
||||
# ── wire protocol ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
#: Seconds between keep-alive progress frames during a long blocking call.
|
||||
_HEARTBEAT_S = 5.0
|
||||
|
||||
#: Serializes _send across threads (the heartbeat below + the main loop) so
|
||||
#: concurrent length+body writes can't interleave and corrupt the framing.
|
||||
_send_lock = threading.Lock()
|
||||
|
||||
|
||||
def _send(stream, obj: dict) -> None:
|
||||
body = json.dumps(obj, separators=(",", ":")).encode("utf-8")
|
||||
stream.write(struct.pack("!I", len(body)))
|
||||
stream.write(body)
|
||||
stream.flush()
|
||||
with _send_lock:
|
||||
stream.write(struct.pack("!I", len(body)))
|
||||
stream.write(body)
|
||||
stream.flush()
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _heartbeat(stdout, stage: str):
|
||||
"""Emit a progress frame every ~5s for the duration of the block.
|
||||
|
||||
IndexTTS spends the whole of a cold load and the whole of ``infer()``
|
||||
inside one blocking upstream call, saying nothing on the wire. The parent
|
||||
reads that silence two ways, and BOTH kill a perfectly healthy synthesis
|
||||
of a long passage (#1611):
|
||||
|
||||
* ``SubprocessBackend.generate`` re-arms its recv watchdog on every
|
||||
frame, so with no frames it hard-kills the sidecar at recv_timeout_s;
|
||||
* each frame also reports activity to the GPU pool's execution clock
|
||||
(#1367), so with no frames the outer generate budget expires and
|
||||
blames the hardware.
|
||||
|
||||
Raising the deadline alone therefore does not fix long-text generation —
|
||||
the sidecar has to prove it is alive. Percent climbs 1..99 because the
|
||||
upstream call exposes no real progress; it is a liveness signal, not a
|
||||
measurement.
|
||||
"""
|
||||
stop = threading.Event()
|
||||
|
||||
def _beat() -> None:
|
||||
pct = 1
|
||||
while not stop.wait(_HEARTBEAT_S):
|
||||
pct = min(pct + 1, 99)
|
||||
try:
|
||||
_send(stdout, {"op": "progress", "stage": stage, "percent": pct})
|
||||
except Exception:
|
||||
return # pipe gone — the main loop will surface it
|
||||
hb = threading.Thread(target=_beat, name=f"indextts-{stage}-heartbeat", daemon=True)
|
||||
hb.start()
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
stop.set()
|
||||
hb.join(timeout=_HEARTBEAT_S + 1)
|
||||
|
||||
|
||||
def _recv(stream):
|
||||
@@ -160,14 +210,40 @@ def _torch_bf16_supported() -> bool:
|
||||
return False
|
||||
|
||||
|
||||
#: Model-config filenames to look for, most-preferred first, per version.
|
||||
#: IndexTeam/IndexTTS-2.5 ships ``config.yaml``; VoiceStudio used to demand
|
||||
#: ``config_v2_5.yaml``, a name that exists in no upstream revision, so the
|
||||
#: install failed until the user hand-renamed the file (#1611). Both names are
|
||||
#: accepted now — the hand-renamed installs must keep working untouched — and
|
||||
#: the renamed one wins, because a user who created it did so deliberately.
|
||||
_CFG_NAMES = {
|
||||
"2.5": ("config_v2_5.yaml", "config.yaml"),
|
||||
"2": ("config.yaml",),
|
||||
}
|
||||
|
||||
|
||||
def _resolve_cfg_path(model_dir: str, *, version: str) -> str:
|
||||
"""First accepted config that exists in ``model_dir``.
|
||||
|
||||
Falls back to the last candidate when none exist, so the failure surfaces
|
||||
as upstream's own "no such file" naming a real expected path rather than
|
||||
a name no upstream release has ever shipped.
|
||||
"""
|
||||
names = _CFG_NAMES.get(version, _CFG_NAMES["2"])
|
||||
for name in names:
|
||||
candidate = os.path.join(model_dir, name)
|
||||
if os.path.isfile(candidate):
|
||||
return candidate
|
||||
return os.path.join(model_dir, names[-1])
|
||||
|
||||
|
||||
def _model_init_kwargs(
|
||||
repo_dir: str, *, version: str, reduced_precision: bool,
|
||||
) -> dict:
|
||||
"""Build version-specific constructor arguments for IndexTTS 2.5 or 2."""
|
||||
model_dir = os.path.join(repo_dir, "checkpoints")
|
||||
cfg_name = "config_v2_5.yaml" if version == "2.5" else "config.yaml"
|
||||
kwargs = {
|
||||
"cfg_path": os.path.join(model_dir, cfg_name),
|
||||
"cfg_path": _resolve_cfg_path(model_dir, version=version),
|
||||
"model_dir": model_dir,
|
||||
"use_cuda_kernel": False,
|
||||
"use_deepspeed": False,
|
||||
@@ -216,7 +292,8 @@ def _load_model(stdout) -> object:
|
||||
model_kw = _model_init_kwargs(
|
||||
repo_dir, version=_model_version, reduced_precision=reduced_precision,
|
||||
)
|
||||
_model = IndexTTS2(**model_kw)
|
||||
with _heartbeat(stdout, "loading_model"):
|
||||
_model = IndexTTS2(**model_kw)
|
||||
|
||||
_send(stdout, {"op": "progress", "stage": "loading_model", "percent": 100})
|
||||
return _model
|
||||
@@ -276,7 +353,10 @@ def _handle_synthesize(msg: dict, stdout) -> None:
|
||||
tmp_path = tmp.name
|
||||
try:
|
||||
infer_kw["output_path"] = tmp_path
|
||||
model.infer(**infer_kw)
|
||||
# A long passage keeps infer() busy for minutes with nothing on the
|
||||
# wire; without this the parent kills the sidecar mid-synthesis (#1611).
|
||||
with _heartbeat(stdout, "synthesizing"):
|
||||
model.infer(**infer_kw)
|
||||
pcm_b64, sr, n_samples = _wav_to_pcm_b64(tmp_path)
|
||||
finally:
|
||||
try:
|
||||
|
||||
@@ -98,6 +98,7 @@ def _platform_slug() -> str:
|
||||
darwin-x86_64
|
||||
windows-x86_64
|
||||
linux-x86_64
|
||||
linux-aarch64
|
||||
"""
|
||||
system = platform.system().lower()
|
||||
machine = platform.machine().lower()
|
||||
@@ -107,6 +108,8 @@ def _platform_slug() -> str:
|
||||
return "darwin-x86_64"
|
||||
if system == "windows":
|
||||
return "windows-x86_64"
|
||||
if system == "linux" and machine in ("arm64", "aarch64"):
|
||||
return "linux-aarch64"
|
||||
# Linux + everything else falls into the linux slug.
|
||||
return "linux-x86_64"
|
||||
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
"""omnivoice-subprocess: the resident OmniVoice TTS engine in a crash-isolated
|
||||
sidecar process (#730/#1190).
|
||||
|
||||
The default ``omnivoice`` engine runs in-process on the GPU ``ThreadPoolExecutor``.
|
||||
The ``omnivoice`` engine runs in-process on CUDA, ROCm, and CPU. On MPS it is
|
||||
resolved to :class:`OmniVoiceMPSSubprocessBackend` so a fatal native allocator
|
||||
exit cannot take down the local API process.
|
||||
When a generate or load there exceeds its execution budget the pool is "reset"
|
||||
but the abandoned worker *thread* cannot be killed (Python cannot interrupt a
|
||||
native torch/MPS call), so it holds the MPS device until it finishes on its
|
||||
@@ -13,16 +15,11 @@ timeout the parent's watchdog calls ``proc.kill()``, reclaiming the child's
|
||||
VRAM/device, and the next request transparently respawns a fresh sidecar. That
|
||||
is the one thing the in-process engine structurally cannot do.
|
||||
|
||||
OPT-IN (Settings -> Engines, or ``OMNIVOICE_TTS_BACKEND=omnivoice-subprocess``);
|
||||
the in-process ``omnivoice`` stays the default so existing users see no change.
|
||||
The explicit ``omnivoice-subprocess`` id remains available on every host for
|
||||
operators who want the same containment elsewhere.
|
||||
|
||||
Tradeoff vs the in-process engine: identical model and quality, a little extra
|
||||
per-call overhead (one stdio round-trip), and it does not carry the native
|
||||
advanced-parameter surface (``t_shift`` / ``layer_penalty_factor`` /
|
||||
``position_temperature`` / ``class_temperature``) or parent-side seed
|
||||
determinism, because the generic ``backend.generate`` path does not forward
|
||||
those. Acceptable for unattended / reaction-triggered use where reliability
|
||||
matters more than those controls.
|
||||
Tradeoff vs the in-process engine: identical model, controls, seed behavior,
|
||||
and quality, with a little extra per-call overhead (one stdio round-trip).
|
||||
|
||||
Unlike IndexTTS / dots.tts / Supertonic-3, this sidecar runs under the PARENT
|
||||
interpreter (``venv_python() -> sys.executable``): the goal here is crash
|
||||
@@ -51,7 +48,7 @@ class OmniVoiceSubprocessBackend(SubprocessBackend):
|
||||
id = "omnivoice-subprocess"
|
||||
display_name = "OmniVoice (subprocess-isolated, killable on timeout)"
|
||||
_DEFAULT_SAMPLE_RATE = 24000
|
||||
gpu_compat = ("cuda", "mps", "cpu")
|
||||
gpu_compat = ("cuda", "rocm", "mps", "cpu")
|
||||
# Match OmniVoiceBackend: the measured floor below which a render that
|
||||
# should take seconds runs for minutes (the #1226/#1222 4 GB reports).
|
||||
min_vram_gb = 6.0
|
||||
@@ -102,4 +99,34 @@ class OmniVoiceSubprocessBackend(SubprocessBackend):
|
||||
return ["multi"]
|
||||
|
||||
|
||||
__all__ = ["OmniVoiceSubprocessBackend"]
|
||||
class OmniVoiceMPSSubprocessBackend(OmniVoiceSubprocessBackend):
|
||||
"""Effective ``omnivoice`` implementation on MPS.
|
||||
|
||||
Native torch/MPS allocator failures can terminate the process without a
|
||||
catchable Python exception. Keeping the same engine id and model surface in
|
||||
a child makes that failure recoverable while Settings, APIs, and saved
|
||||
projects continue to refer to ``omnivoice``.
|
||||
"""
|
||||
|
||||
id = "omnivoice"
|
||||
display_name = "VoiceStudio (k2-fsa/OmniVoice, 600+ languages)"
|
||||
supports_native_omnivoice_controls = True
|
||||
|
||||
def generate(self, text: str, **kw):
|
||||
from services.model_manager import make_room_before_generate
|
||||
|
||||
make_room_before_generate()
|
||||
try:
|
||||
return super().generate(text, **kw)
|
||||
except RuntimeError as exc:
|
||||
if "sidecar closed pipe mid-generate" not in str(exc):
|
||||
raise
|
||||
raise RuntimeError(
|
||||
"The isolated OmniVoice engine stopped during generation, "
|
||||
"usually because macOS reclaimed it under memory pressure. "
|
||||
"The VoiceStudio backend is still running. Close memory-heavy "
|
||||
"apps or select a smaller TTS engine, then retry."
|
||||
) from exc
|
||||
|
||||
|
||||
__all__ = ["OmniVoiceMPSSubprocessBackend", "OmniVoiceSubprocessBackend"]
|
||||
|
||||
@@ -50,6 +50,8 @@ OMNIVOICE_SAMPLE_RATE = 24000
|
||||
_GEN_KW_ALLOWLIST = (
|
||||
"language", "instruct", "duration", "num_step", "guidance_scale",
|
||||
"speed", "denoise", "postprocess_output", "preprocess_prompt",
|
||||
"t_shift", "layer_penalty_factor", "position_temperature",
|
||||
"class_temperature", "audio_chunk_duration", "audio_chunk_threshold",
|
||||
)
|
||||
|
||||
_model = None
|
||||
@@ -183,6 +185,12 @@ def _handle_synthesize(msg: dict, stdout) -> None:
|
||||
ref_text = msg.get("ref_text") or None
|
||||
gen_kw = {k: msg[k] for k in _GEN_KW_ALLOWLIST if k in msg}
|
||||
|
||||
seed = msg.get("seed")
|
||||
if seed is not None:
|
||||
import torch
|
||||
|
||||
torch.manual_seed(int(seed))
|
||||
|
||||
audios = model.generate(
|
||||
text=text, ref_audio=ref_audio, ref_text=ref_text, **gen_kw
|
||||
)
|
||||
|
||||
@@ -151,6 +151,40 @@ def _pocket_language(raw) -> str:
|
||||
)
|
||||
|
||||
|
||||
_TRUTHY = {"1", "true", "yes", "on"}
|
||||
|
||||
|
||||
def _has_24l_config(language: str) -> bool:
|
||||
"""Whether the installed pocket-tts ships a 24-layer checkpoint for
|
||||
``language`` (it/de/es/pt/fr in 2.1.0; english has none)."""
|
||||
try:
|
||||
from pocket_tts.models.tts_model import CONFIGS_DIR # type: ignore[import-not-found] # noqa: PLC0415
|
||||
except Exception as exc: # noqa: BLE001 — absence of the package is not fatal here
|
||||
# Log it, though: if a future pocket-tts moves CONFIGS_DIR, the 24L
|
||||
# opt-in would otherwise go silently inert.
|
||||
print(f"pockettts sidecar: 24l config probe failed: {exc!r}", file=sys.stderr)
|
||||
return False
|
||||
from pathlib import Path # noqa: PLC0415
|
||||
|
||||
return (Path(CONFIGS_DIR) / f"{language}_24l.yaml").is_file()
|
||||
|
||||
|
||||
def _model_config_name(language: str) -> str:
|
||||
"""Pocket-tts config name to load: the 6-layer default, or the 24-layer
|
||||
checkpoint when OMNIVOICE_POCKETTTS_24L is set and one exists for the
|
||||
language. Opt-in only — defaults keep the fast model; the 24-layer variant
|
||||
trades roughly 4x transformer compute for better prosody.
|
||||
|
||||
French is the exception: pocket-tts 2.1.0 only ships a 24-layer French
|
||||
model and load_model(language="french") raises, so French always maps to
|
||||
french_24l regardless of the env var."""
|
||||
if language == "french":
|
||||
return "french_24l"
|
||||
if os.environ.get("OMNIVOICE_POCKETTTS_24L", "").strip().lower() not in _TRUTHY:
|
||||
return language
|
||||
return f"{language}_24l" if _has_24l_config(language) else language
|
||||
|
||||
|
||||
def _load_model(stdout, language: str):
|
||||
"""Cold-construct the PocketTTS model for ``language`` (cached per language).
|
||||
Emits progress frames for the parent watchdog. Raises on failure (e.g.
|
||||
@@ -178,7 +212,7 @@ def _load_model(stdout, language: str):
|
||||
try:
|
||||
from pocket_tts import TTSModel # type: ignore[import-not-found] # noqa: PLC0415
|
||||
|
||||
model = TTSModel.load_model(language=language)
|
||||
model = TTSModel.load_model(language=_model_config_name(language))
|
||||
_MODELS[language] = model
|
||||
finally:
|
||||
stop.set()
|
||||
|
||||
+107
-17
@@ -9,6 +9,24 @@ _backend_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
if _backend_dir not in sys.path:
|
||||
sys.path.insert(0, _backend_dir)
|
||||
|
||||
# PyInstaller re-executes this entry module when the frozen backend binary is
|
||||
# launched. Nested operation supervisors therefore dispatch here, before math,
|
||||
# logging, FastAPI, torch, or any application initialization. Source launches
|
||||
# use this same entry contract so frozen/source behavior cannot drift.
|
||||
if __name__ == "__main__" and len(sys.argv) > 1 and sys.argv[1] == "--supervise":
|
||||
from core.contained_subprocess import supervisor_main
|
||||
|
||||
raise SystemExit(supervisor_main(sys.argv[1:]))
|
||||
|
||||
# Rust clears CLOEXEC only for the backend exec. Re-arm PEP 446 immediately:
|
||||
# nested supervisors receive this descriptor solely through explicit pass_fds,
|
||||
# so a third-party close_fds=False child cannot hold the desktop drain barrier.
|
||||
from core.contained_subprocess import secure_backend_drain_fd # noqa: E402
|
||||
|
||||
secure_backend_drain_fd()
|
||||
|
||||
import math # noqa: E402
|
||||
|
||||
# Windows: run every child process (ffmpeg, engine sidecars, yt-dlp, demucs, …)
|
||||
# WITHOUT popping a console window. The backend itself is spawned console-less by
|
||||
# the Tauri shell, so on Windows each console subprocess it launches would
|
||||
@@ -369,19 +387,36 @@ def _env_flag(name: str, default: bool = False) -> bool:
|
||||
_EAGER = _env_flag("OMNIVOICE_EAGER_INIT", default=("pytest" in sys.modules))
|
||||
|
||||
|
||||
def _env_float(name: str, default: float) -> float:
|
||||
"""Parse a float env override, rejecting negative and non-finite values.
|
||||
|
||||
Shared by the preload-delay / timeout knobs: NaN would silently never
|
||||
fire, a negative would fire during startup I/O, so both fall back to the
|
||||
default instead (the bug class CodeRabbit flagged on the watermark knob
|
||||
in PR #1577 — latent in the older copies too, closed here for all)."""
|
||||
raw = os.environ.get(name, "")
|
||||
try:
|
||||
value = float(raw) if raw.strip() else default
|
||||
except ValueError:
|
||||
return default
|
||||
return value if math.isfinite(value) and value >= 0 else default
|
||||
|
||||
|
||||
def _capture_preload_delay_s() -> float:
|
||||
"""Seconds after boot before the dictation (capture ASR) model warms.
|
||||
|
||||
Late enough that it never competes with startup I/O or the TTS preload;
|
||||
overridable via OMNIVOICE_CAPTURE_PRELOAD_DELAY (mostly for tests)."""
|
||||
raw = os.environ.get("OMNIVOICE_CAPTURE_PRELOAD_DELAY", "")
|
||||
try:
|
||||
v = float(raw)
|
||||
if v >= 0:
|
||||
return v
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
return 30.0
|
||||
return _env_float("OMNIVOICE_CAPTURE_PRELOAD_DELAY", 30.0)
|
||||
|
||||
def _watermark_preload_delay_s() -> float:
|
||||
"""Seconds after boot before the AudioSeal generator warm-up fires.
|
||||
|
||||
Own knob, NOT ``_capture_preload_delay_s`` + offset: a capture-specific
|
||||
env override must not retime the watermark warm too, and the two cold
|
||||
imports shouldn't fire on the same tick (CodeRabbit, PR #1577). Default
|
||||
35s sits ~5s past the capture-ASR warm for the same reason."""
|
||||
return _env_float("OMNIVOICE_PRELOAD_WATERMARK_DELAY", 35.0)
|
||||
|
||||
|
||||
def _capture_preload_ram_ok(min_free_bytes: int = 4 * 1024**3) -> bool:
|
||||
@@ -398,14 +433,7 @@ def _capture_preload_ram_ok(min_free_bytes: int = 4 * 1024**3) -> bool:
|
||||
def _mcp_start_timeout_s() -> float:
|
||||
"""Seconds to wait for the MCP session manager to start before giving up
|
||||
and serving without it (#632). Overridable via OMNIVOICE_MCP_START_TIMEOUT_S."""
|
||||
raw = os.environ.get("OMNIVOICE_MCP_START_TIMEOUT_S", "")
|
||||
try:
|
||||
v = float(raw)
|
||||
if v > 0:
|
||||
return v
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
return 30.0
|
||||
return max(_env_float("OMNIVOICE_MCP_START_TIMEOUT_S", 30.0), 0.001)
|
||||
|
||||
|
||||
async def _serve_mcp(session_manager, ready: "asyncio.Event", stop: "asyncio.Event") -> None:
|
||||
@@ -638,6 +666,7 @@ def _phase_a_build_inner() -> None:
|
||||
events,
|
||||
capture,
|
||||
capture_ws,
|
||||
speech_platform,
|
||||
dictation,
|
||||
openai_compat,
|
||||
tts_stream,
|
||||
@@ -657,7 +686,7 @@ def _phase_a_build_inner() -> None:
|
||||
system, profiles, exports, generation, dub_core, dub_generate,
|
||||
dub_export, dub_translate, projects, glossary, engines, tools,
|
||||
stories, setup, gallery, archetypes, describe_voice, community,
|
||||
batch, watermark, events, capture, capture_ws, dictation,
|
||||
batch, watermark, events, capture, capture_ws, speech_platform, dictation,
|
||||
openai_compat, tts_stream, marketplace, personas, sonitranslate,
|
||||
audiobook, longform_jobs, pronunciation, settings_router,
|
||||
media_tools_router, auth_router, _mcp_bindings_router, workers_router,
|
||||
@@ -852,6 +881,8 @@ async def _phase_b(app: FastAPI) -> None:
|
||||
# #1174: arm model loads for THIS run — an in-process relaunch may carry a
|
||||
# stale shutting-down flag from a previous lifespan.
|
||||
model_loads_reset_shutdown()
|
||||
from services.model_manager import begin_watermark_pool_lifecycle
|
||||
begin_watermark_pool_lifecycle()
|
||||
app.state.idle_task = asyncio.create_task(idle_worker())
|
||||
app.state.worker_task = asyncio.create_task(task_manager.worker())
|
||||
# Warm the TTS model in the background so first /generate is instant.
|
||||
@@ -905,6 +936,50 @@ async def _phase_b(app: FastAPI) -> None:
|
||||
else:
|
||||
logger.info("Capture ASR preload disabled; dictation ASR will load on first use.")
|
||||
|
||||
# Watermark: warm the AudioSeal generator in the background so the first
|
||||
# mark_synthetic doesn't serialize the audioseal import + model load
|
||||
# inside the first synthesis (measured ~42 s inline on a cold filesystem,
|
||||
# 2026-08-17 macOS report — 3 s short of the client's 90 s timeout).
|
||||
# Small model on CPU; deferred a few seconds past the capture-ASR warm so
|
||||
# the two cold imports don't contend for the same disk, and no RAM guard
|
||||
# is needed. Runs on the watermark pool — where the model is used — not
|
||||
# the shared default executor.
|
||||
if _env_flag("OMNIVOICE_PRELOAD_WATERMARK", default=True):
|
||||
async def _preload_watermark():
|
||||
await asyncio.sleep(_watermark_preload_delay_s())
|
||||
loop = asyncio.get_running_loop()
|
||||
from services import watermark as _watermark
|
||||
|
||||
# Gate BEFORE touching get_watermark_pool(): the pool is lazy so
|
||||
# hosts with watermarking disabled never spawn its thread, and
|
||||
# creating it unconditionally would break that invariant. The
|
||||
# race with a first embed is benign — pool creation is itself
|
||||
# lock-guarded.
|
||||
if not _watermark.will_mark():
|
||||
logger.debug("Watermark preload skipped (disabled or audioseal absent)")
|
||||
return
|
||||
from services.model_manager import get_watermark_pool
|
||||
|
||||
# Default startup may warm an existing local checkpoint but may
|
||||
# not fetch one. Only an explicit user opt-in permits a download.
|
||||
raw_preload = os.environ.get("OMNIVOICE_PRELOAD_WATERMARK", "")
|
||||
allow_download = raw_preload.strip().lower() in {"1", "true", "yes", "on"}
|
||||
|
||||
try:
|
||||
await loop.run_in_executor(
|
||||
get_watermark_pool(),
|
||||
lambda: _watermark.prefetch_generator(
|
||||
allow_download=allow_download
|
||||
),
|
||||
)
|
||||
except Exception:
|
||||
# prefetch_generator swallows its own errors; this guards the
|
||||
# setup half (imports, pool construction) so a broken warm-up
|
||||
# is visible now, not as an unretrieved exception at shutdown.
|
||||
logger.warning("Watermark preload task failed", exc_info=True)
|
||||
|
||||
app.state.watermark_preload_task = asyncio.create_task(_preload_watermark())
|
||||
|
||||
# ── MCP session manager (Wave 2.2) ────────────────────────────────────
|
||||
# Run it in its OWN task owning the full enter→exit lifecycle (anyio
|
||||
# task-affinity, see _serve_mcp); only wait, with a timeout, for ready —
|
||||
@@ -934,6 +1009,9 @@ async def _phase_b(app: FastAPI) -> None:
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
from api.dependencies import validate_server_admin_key
|
||||
|
||||
validate_server_admin_key()
|
||||
# Startup watchdog (#632): a silent hang during startup (e.g. a model-load /
|
||||
# MCP deadlock on some platforms) means "Application startup complete" never
|
||||
# logs and the app sits forever with no error. If startup hasn't finished
|
||||
@@ -1080,8 +1158,20 @@ async def lifespan(app: FastAPI):
|
||||
getattr(app.state, "worker_task", None),
|
||||
getattr(app.state, "preload_task", None),
|
||||
getattr(app.state, "capture_preload_task", None),
|
||||
getattr(app.state, "watermark_preload_task", None),
|
||||
timeout=20.0,
|
||||
)
|
||||
# The watermark warm-up runs on its dedicated 1-worker pool. Cancellation
|
||||
# detaches the asyncio future but cannot kill a thread inside AudioSeal,
|
||||
# so drain it fully before lifespan teardown reports completion.
|
||||
try:
|
||||
from services.model_manager import shutdown_watermark_pool as _wm_drain
|
||||
|
||||
_wm_drain()
|
||||
except Exception:
|
||||
# Best-effort drain: a failure here must not abort the remaining
|
||||
# shutdown steps (model unload, MCP teardown) below.
|
||||
logger.warning("Watermark pool drain failed at shutdown", exc_info=True)
|
||||
# Unload the model and free GPU memory
|
||||
try:
|
||||
import services.model_manager as mm
|
||||
|
||||
@@ -190,6 +190,7 @@ class ParseSubtitleTextRequest(BaseModel):
|
||||
class DubIngestUrlRequest(BaseModel):
|
||||
url: str
|
||||
job_id: Optional[str] = None
|
||||
source_lang: Optional[str] = None
|
||||
# When true and the URL is a caption-bearing host (YouTube, Vimeo, TED…),
|
||||
# ask yt-dlp to also download the original-language + any additional
|
||||
# sub_langs as VTT. The UI uses this to seed a transcript without running
|
||||
|
||||
@@ -33,7 +33,9 @@ WS_TICKET_PREFIX = "ovs_ws_ticket_"
|
||||
_TOKEN_BYTES = 32
|
||||
_ENCODED_TOKEN_LENGTH = 43
|
||||
_TOKEN_BODY_RE = re.compile(rf"^[A-Za-z0-9_-]{{{_ENCODED_TOKEN_LENGTH}}}$")
|
||||
_ALLOWED_WS_PATHS = frozenset({"/ws/events", "/ws/transcribe"})
|
||||
_ALLOWED_WS_PATHS = frozenset(
|
||||
{"/ws/events", "/ws/transcribe", "/v1/audio/transcriptions/stream"}
|
||||
)
|
||||
_ADMIN_CAPABILITIES = frozenset({"consume", "admin"})
|
||||
_KEY_GENERATION_INFO = b"omnivoice-admin-key-generation-v1"
|
||||
|
||||
|
||||
+103
-43
@@ -88,10 +88,9 @@ def reset_pool_after_wedge(executor, *, what: str = "ASR") -> bool:
|
||||
|
||||
|
||||
# ── Consecutive-timeout streak → recommend the crash-isolated engine ────────
|
||||
# A pool reset restores *capacity*, but the wedged CTranslate2/whisperx thread
|
||||
# keeps its VRAM until the process exits. When guarded transcribes keep timing
|
||||
# out back-to-back in one session, resets clearly aren't recovering the
|
||||
# underlying hang — the durable fix is the crash-isolated sidecar engine
|
||||
# A timed-out CTranslate2/whisperx thread keeps its worker and VRAM until the
|
||||
# native call exits. When guarded transcribes keep timing out back-to-back in
|
||||
# one session, the durable fix is the crash-isolated sidecar engine
|
||||
# (services.subprocess_asr, #393), whose child process CAN be hard-killed to
|
||||
# reclaim the hung call and its VRAM. We only *recommend* it (log + error
|
||||
# message); we never switch engines automatically (owner rule: no silent
|
||||
@@ -146,23 +145,18 @@ def _isolated_engine_hint(streak: int) -> str:
|
||||
|
||||
async def run_transcribe_guarded(executor, fn, *, what: str = "ASR",
|
||||
timeout: float = ASR_TRANSCRIBE_TIMEOUT_S,
|
||||
timeout_env: str = "OMNIVOICE_ASR_TRANSCRIBE_TIMEOUT_S"):
|
||||
timeout_env: str = "OMNIVOICE_ASR_TRANSCRIBE_TIMEOUT_S",
|
||||
reset_on_timeout: bool = False):
|
||||
"""Run a blocking transcribe ``fn`` in ``executor`` with a hard wall-clock
|
||||
bound. On timeout, raise :class:`ASRTimeoutError` with guidance instead of
|
||||
letting the request hang forever.
|
||||
|
||||
``run_in_executor`` cannot cancel the underlying thread, so a wedged
|
||||
transcribe (a CTranslate2 / whisperx / VAD hang seen on some Windows + CUDA
|
||||
setups, #730) keeps occupying its GPU-pool worker. With a 1–2 worker pool
|
||||
that starves every *other* request — including TTS generate — and the next
|
||||
thing the user does surfaces as "Can't reach the local backend" even though
|
||||
the process is alive. So on timeout we also ``reset()`` the pool when it
|
||||
supports it (``_ResilientGpuPool``): the wedged thread is abandoned and the
|
||||
next submit gets a fresh worker, restoring capacity without an app restart.
|
||||
The orphaned thread still holds its VRAM until the process exits, which is
|
||||
why the message still recommends a smaller ASR model / Flush as the durable
|
||||
fix. Executors without ``reset`` (a plain ThreadPoolExecutor in tests) just
|
||||
get the bound + actionable error.
|
||||
``run_in_executor`` cannot cancel the underlying thread, so a timed-out
|
||||
in-process CTranslate2/whisperx call still owns its model and device. The
|
||||
default deliberately leaves that worker accounted for: swapping in a fresh
|
||||
pool and immediately retrying the same backend overlaps two native calls,
|
||||
which produced the Windows access violation in #1669. A caller backed by a
|
||||
genuinely killable process may opt into ``reset_on_timeout``.
|
||||
"""
|
||||
loop = asyncio.get_running_loop()
|
||||
# Same SystemExit containment as the TTS pool (#1133 class): an ASR
|
||||
@@ -171,16 +165,16 @@ async def run_transcribe_guarded(executor, fn, *, what: str = "ASR",
|
||||
try:
|
||||
result = await asyncio.wait_for(fut, timeout=timeout)
|
||||
except asyncio.TimeoutError:
|
||||
# Free the poisoned pool so a hung transcribe can't keep starving TTS /
|
||||
# other ASR work (the "can't reach backend" symptom, #730).
|
||||
reset_pool_after_wedge(executor, what=what)
|
||||
if reset_on_timeout:
|
||||
reset_pool_after_wedge(executor, what=what)
|
||||
streak = _note_transcribe_timeout()
|
||||
msg = (
|
||||
f"{what} transcription exceeded {timeout:.0f}s and was abandoned — "
|
||||
"the backend is running, but the ASR model is too heavy for the "
|
||||
"available compute. Most often the GPU is VRAM-starved: the resident "
|
||||
"TTS model and a large ASR model (large-v3) contend for memory. "
|
||||
"Capacity was restored automatically, but for a durable fix Flush the "
|
||||
"The native call cannot be killed safely, so its capacity remains "
|
||||
"reserved until it exits. For a durable fix Flush the "
|
||||
"TTS model to free VRAM, pick a smaller ASR model in "
|
||||
f"Model Catalogue → Models, or set ASR to CPU. (Raise {timeout_env} "
|
||||
"for very long transcribes.)"
|
||||
@@ -2564,7 +2558,10 @@ def _auto_detect() -> str:
|
||||
def active_backend_id() -> str:
|
||||
explicit = os.environ.get("OMNIVOICE_ASR_BACKEND")
|
||||
if explicit:
|
||||
return explicit
|
||||
# #1582's public spelling predates the registry name. Keep it as a
|
||||
# compatibility alias for the PyTorch-native Whisper implementation
|
||||
# that can use ROCm/HIP; every ASR consumer resolves through here.
|
||||
return "pytorch-whisper" if explicit == "omnivoice" else explicit
|
||||
from core import prefs
|
||||
picked = prefs.get("asr_backend")
|
||||
if picked:
|
||||
@@ -2992,7 +2989,7 @@ def _capture_prefers_parakeet() -> bool:
|
||||
return _parakeet_mlx_installed()
|
||||
|
||||
|
||||
def get_capture_asr_backend() -> ASRBackend:
|
||||
def get_capture_asr_backend(*, skip_sherpa: bool = False) -> ASRBackend:
|
||||
"""Pick the fastest ASR engine for capture / dictation.
|
||||
|
||||
Selection order:
|
||||
@@ -3017,6 +3014,9 @@ def get_capture_asr_backend() -> ASRBackend:
|
||||
|
||||
Returns a cached singleton so the model stays warm between calls; the
|
||||
singleton is rebuilt if the selected sherpa model changes.
|
||||
|
||||
``skip_sherpa`` is used only to validate a token-silent Sherpa result with
|
||||
the installed capture fallback before persisting model demotion.
|
||||
"""
|
||||
global _capture_backend, _capture_backend_key
|
||||
|
||||
@@ -3025,7 +3025,7 @@ def get_capture_asr_backend() -> ASRBackend:
|
||||
# call get_sherpa_dictation_backend concurrently) can't both build a model.
|
||||
with _capture_backend_lock:
|
||||
# 0. Honor an explicit sherpa dictation model selection.
|
||||
sherpa_id = dictation_model_id()
|
||||
sherpa_id = None if skip_sherpa else dictation_model_id()
|
||||
if sherpa_id:
|
||||
ok, _ = SherpaDictationBackend.is_available()
|
||||
if ok:
|
||||
@@ -3192,7 +3192,10 @@ def _capture_whisper_repo() -> str | None:
|
||||
return os.environ.get("OMNIVOICE_PYTORCH_ASR_MODEL", _PYTORCH_ASR_DEFAULT)
|
||||
|
||||
|
||||
def _recommended_asr_model(purpose: str, missing_repo: str | None) -> dict | None:
|
||||
def _recommended_asr_model(
|
||||
purpose: str, missing_repo: str | None, *, prefer_sherpa: bool = True,
|
||||
excluded_sherpa_model_id: str | None = None,
|
||||
) -> dict | None:
|
||||
"""The catalog entry to offer in the download CTA.
|
||||
|
||||
Offline: the missing repo itself when it's in the catalog (guarantees
|
||||
@@ -3212,20 +3215,38 @@ def _recommended_asr_model(purpose: str, missing_repo: str | None) -> dict | Non
|
||||
|
||||
by_id = {m["repo_id"]: m for m in KNOWN_MODELS}
|
||||
exact = by_id.get(missing_repo) if missing_repo else None
|
||||
want_sherpa = False
|
||||
if purpose == "dictation":
|
||||
if exact is not None and exact.get("engine") == "sherpa-onnx":
|
||||
|
||||
def _eligible(m: dict, *, sherpa: bool) -> bool:
|
||||
if (m.get("engine") == "sherpa-onnx") != sherpa:
|
||||
return False
|
||||
if sherpa and m.get("dictation_id") == excluded_sherpa_model_id:
|
||||
return False
|
||||
return _model_supported(m)
|
||||
|
||||
if purpose != "dictation":
|
||||
if exact is not None and _model_supported(exact):
|
||||
return _shape(exact)
|
||||
prefer_sherpa = False
|
||||
|
||||
if purpose == "dictation" and prefer_sherpa:
|
||||
ok, _ = SherpaDictationBackend.is_available()
|
||||
want_sherpa = ok
|
||||
if not want_sherpa and exact is not None and _model_supported(exact):
|
||||
if ok:
|
||||
if exact is not None and _eligible(exact, sherpa=True):
|
||||
return _shape(exact)
|
||||
for m in KNOWN_MODELS:
|
||||
if (m.get("role") == "ASR" and _eligible(m, sherpa=True)
|
||||
and _model_curated(m)):
|
||||
return _shape(m)
|
||||
|
||||
# No usable Sherpa recommendation remains (runtime unavailable, explicit
|
||||
# fallback probe, or the sole curated entry is the demoted model). Offer
|
||||
# the exact capture fallback so download → retry cannot loop.
|
||||
if exact is not None and _eligible(exact, sherpa=False):
|
||||
return _shape(exact)
|
||||
for m in KNOWN_MODELS:
|
||||
if m.get("role") != "ASR":
|
||||
continue
|
||||
if (m.get("engine") == "sherpa-onnx") != want_sherpa:
|
||||
continue
|
||||
if _model_curated(m) and _model_supported(m):
|
||||
if _eligible(m, sherpa=False) and _model_curated(m):
|
||||
return _shape(m)
|
||||
return None
|
||||
|
||||
@@ -3256,7 +3277,9 @@ def _repo_installed(repo: str) -> bool:
|
||||
|
||||
def asr_model_missing_error(*, purpose: str = "transcribe",
|
||||
sherpa_model_id: str | None = None,
|
||||
backend_id: str | None = None) -> dict | None:
|
||||
backend_id: str | None = None,
|
||||
skip_sherpa: bool = False,
|
||||
require_installed: bool = False) -> dict | None:
|
||||
"""None when the active ASR selection can transcribe without downloading
|
||||
anything; otherwise the typed ``{"error": "asr_model_missing", ...}``
|
||||
payload for a 409 / SSE / WS error with a download CTA.
|
||||
@@ -3268,6 +3291,11 @@ def asr_model_missing_error(*, purpose: str = "transcribe",
|
||||
``?model=`` override. Installed state comes from the same HF-cache helpers
|
||||
the model store uses (see :func:`_repo_installed`), so the answer matches
|
||||
the Model Catalogue → Models install badges.
|
||||
``skip_sherpa`` probes only the non-Sherpa capture fallback; silent-model
|
||||
recovery uses it before deciding whether persistent demotion is warranted.
|
||||
``require_installed`` makes unknown/custom selections fail closed for that
|
||||
recovery path so it can never turn the normal fail-open policy into an
|
||||
implicit model download.
|
||||
|
||||
FAIL-OPEN rule: a repo the model catalog doesn't know (a custom
|
||||
``ASR_MODEL_*`` pin, pytorch-whisper's default repo, an unrecognized
|
||||
@@ -3277,27 +3305,55 @@ def asr_model_missing_error(*, purpose: str = "transcribe",
|
||||
a broken preflight must degrade to the old behaviour, not block ASR.
|
||||
"""
|
||||
try:
|
||||
prefer_sherpa_recommendation = not skip_sherpa
|
||||
excluded_sherpa_model_id = None
|
||||
if purpose == "dictation":
|
||||
sid = sherpa_model_id or dictation_model_id()
|
||||
sid = None if skip_sherpa else (sherpa_model_id or dictation_model_id())
|
||||
if sid:
|
||||
ok, _ = SherpaDictationBackend.is_available()
|
||||
if ok:
|
||||
from services import sherpa_dictation as _sd
|
||||
spec = _sd.get_spec(sid)
|
||||
# A recognizer observed returning silence must follow the
|
||||
# same capture fallback as execution, even when the
|
||||
# frontend keeps sending its persisted `?model=` value.
|
||||
if spec is not None:
|
||||
if _sd.is_installed(spec):
|
||||
return None
|
||||
return {
|
||||
"error": ASR_MODEL_MISSING,
|
||||
"missing_repo_id": spec.repo_id,
|
||||
"recommended": _recommended_asr_model(purpose, spec.repo_id),
|
||||
}
|
||||
if _sd.is_demoted(spec.id):
|
||||
excluded_sherpa_model_id = spec.id
|
||||
else:
|
||||
if _sd.is_installed(spec):
|
||||
return None
|
||||
return {
|
||||
"error": ASR_MODEL_MISSING,
|
||||
"missing_repo_id": spec.repo_id,
|
||||
"recommended": _recommended_asr_model(
|
||||
purpose, spec.repo_id,
|
||||
),
|
||||
}
|
||||
repo = _capture_whisper_repo()
|
||||
else:
|
||||
repo = _offline_asr_repo(backend_id)
|
||||
if repo is None:
|
||||
if require_installed:
|
||||
return {
|
||||
"error": ASR_MODEL_MISSING,
|
||||
"missing_repo_id": "unresolved-capture-fallback",
|
||||
"recommended": None,
|
||||
}
|
||||
return None # explicit opt-in engine — can't (and shouldn't) preflight
|
||||
from api.routers.setup.models import get_model_catalog
|
||||
if require_installed:
|
||||
if _repo_installed(repo):
|
||||
return None
|
||||
return {
|
||||
"error": ASR_MODEL_MISSING,
|
||||
"missing_repo_id": repo,
|
||||
"recommended": _recommended_asr_model(
|
||||
purpose, repo,
|
||||
prefer_sherpa=prefer_sherpa_recommendation,
|
||||
excluded_sherpa_model_id=excluded_sherpa_model_id,
|
||||
),
|
||||
}
|
||||
if get_model_catalog().get(repo) is None:
|
||||
return None # not installable from the CTA — fail open (see docstring)
|
||||
if _repo_installed(repo):
|
||||
@@ -3305,7 +3361,11 @@ def asr_model_missing_error(*, purpose: str = "transcribe",
|
||||
return {
|
||||
"error": ASR_MODEL_MISSING,
|
||||
"missing_repo_id": repo,
|
||||
"recommended": _recommended_asr_model(purpose, repo),
|
||||
"recommended": _recommended_asr_model(
|
||||
purpose, repo,
|
||||
prefer_sherpa=prefer_sherpa_recommendation,
|
||||
excluded_sherpa_model_id=excluded_sherpa_model_id,
|
||||
),
|
||||
}
|
||||
except Exception: # noqa: BLE001 — preflight is best-effort, never a blocker
|
||||
logger.warning("ASR install preflight failed — proceeding without it",
|
||||
|
||||
@@ -742,6 +742,74 @@ def _ensure_browser_playable_mp4(video_path: str) -> str:
|
||||
return video_path
|
||||
|
||||
|
||||
async def _ensure_browser_playable_mp4_for_job(job_id: str, video_path: str) -> str:
|
||||
"""Normalize an upload through the job's cancellable process registry."""
|
||||
is_mp4 = video_path.lower().endswith(".mp4")
|
||||
vcodec, acodec = await asyncio.to_thread(_probe_codecs, video_path)
|
||||
if is_mp4 and vcodec in _BROWSER_VIDEO_CODECS and acodec in _BROWSER_AUDIO_CODECS:
|
||||
return video_path
|
||||
|
||||
target = os.path.splitext(video_path)[0] + ".mp4"
|
||||
if target == video_path:
|
||||
target = os.path.splitext(video_path)[0] + ".browser.mp4"
|
||||
run_proc = run_proc_factory(job_id)
|
||||
ffmpeg_bin = find_ffmpeg()
|
||||
|
||||
async def attempt(cmd: list[str]) -> int:
|
||||
try:
|
||||
proc, _stdout, _stderr = await run_proc(cmd, timeout=1800.0)
|
||||
return proc.returncode
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Browser-media normalization process failed for %s: %s",
|
||||
log_safe(video_path),
|
||||
log_safe(exc),
|
||||
)
|
||||
return 1
|
||||
|
||||
rc = 1
|
||||
if not is_mp4:
|
||||
rc = await attempt(
|
||||
[
|
||||
ffmpeg_bin, "-y", "-i", video_path,
|
||||
"-c:v", "copy", "-c:a", "copy",
|
||||
"-movflags", "+faststart", target,
|
||||
]
|
||||
)
|
||||
if rc == 0 and os.path.exists(target):
|
||||
target_vcodec, target_acodec = await asyncio.to_thread(_probe_codecs, target)
|
||||
if (
|
||||
target_vcodec not in _BROWSER_VIDEO_CODECS
|
||||
or target_acodec not in _BROWSER_AUDIO_CODECS
|
||||
):
|
||||
rc = 1
|
||||
else:
|
||||
rc = 1
|
||||
if rc != 0:
|
||||
rc = await attempt(
|
||||
[
|
||||
ffmpeg_bin, "-y", "-i", video_path,
|
||||
"-c:v", "libx264", "-preset", "veryfast", "-crf", "23",
|
||||
"-pix_fmt", "yuv420p", "-c:a", "aac", "-b:a", "192k",
|
||||
"-movflags", "+faststart", target,
|
||||
]
|
||||
)
|
||||
if rc == 0 and os.path.exists(target) and target != video_path:
|
||||
try:
|
||||
os.remove(video_path)
|
||||
except OSError:
|
||||
pass # Best effort: the normalized target is already complete.
|
||||
return target
|
||||
logger.warning(
|
||||
"Could not transcode %s to browser-playable mp4 — the in-app "
|
||||
"video player may render this file as a black box.",
|
||||
log_safe(video_path),
|
||||
)
|
||||
return video_path
|
||||
|
||||
|
||||
# Bounded retry for transient download failures (#579/#598). yt-dlp's own
|
||||
# `retries`/`fragment_retries` cover per-fragment HTTP flakes, but a broken
|
||||
# pipe ([Errno 32]) raised while the write side of a pipe closes mid-stream
|
||||
@@ -1257,6 +1325,13 @@ async def ingest_pipeline(
|
||||
except Exception:
|
||||
dur = 0.0
|
||||
|
||||
# URL downloads already pass through this guard in yt_download_sync.
|
||||
# Uploaded videos did not, so a valid VP9/AV1/Opus upload could be
|
||||
# processed successfully but remain undecodable by the in-app WebView.
|
||||
# Codec probing/transcoding is blocking; keep it off the event loop.
|
||||
if source.get("kind") != "url" and input_type != "audio":
|
||||
video_path = await _ensure_browser_playable_mp4_for_job(job_id, video_path)
|
||||
|
||||
# Content-hash cache: reuse artifacts from previous matching jobs.
|
||||
content_hash = await asyncio.to_thread(compute_file_hash, audio_path)
|
||||
cached = find_cached_job(content_hash, job_id)
|
||||
@@ -1295,6 +1370,7 @@ async def ingest_pipeline(
|
||||
"scene_cuts": scene_cuts,
|
||||
"youtube_subs": youtube_subs_by_lang or None,
|
||||
"input_type": input_type,
|
||||
"source_lang_override": source.get("source_lang"),
|
||||
}
|
||||
if not put_and_save_job(
|
||||
job_id, full_job, filename=filename, duration=dur, content_hash=content_hash,
|
||||
@@ -1323,6 +1399,7 @@ async def ingest_pipeline(
|
||||
"scene_cuts": [],
|
||||
"youtube_subs": youtube_subs_by_lang or None,
|
||||
"input_type": input_type,
|
||||
"source_lang_override": source.get("source_lang"),
|
||||
}
|
||||
if not put_and_save_job(
|
||||
job_id, partial, filename=filename, duration=dur, content_hash=content_hash,
|
||||
|
||||
@@ -177,6 +177,9 @@ class LocalCall:
|
||||
queue_timeout: Optional[float] = None
|
||||
# The engine's declared VRAM floor; only shapes the timeout message.
|
||||
min_vram_gb: float = 0.0
|
||||
# Called once a local worker abandoned by its waiter can no longer touch
|
||||
# request-owned inputs. Normal completion does not call it (#1668).
|
||||
on_abandon: Optional[Callable[[], None]] = None
|
||||
# Some remote-first callers cannot construct the local callable without
|
||||
# loading the very model they are trying to offload. Prepare it only when
|
||||
# routing/fallback actually selects this machine.
|
||||
@@ -465,6 +468,7 @@ async def _run_local(call: LocalCall, *, admit: bool = False, executor=None) ->
|
||||
queue_timeout=call.queue_timeout,
|
||||
min_vram_gb=call.min_vram_gb,
|
||||
executor=executor,
|
||||
on_abandon=call.on_abandon,
|
||||
)
|
||||
|
||||
|
||||
@@ -499,7 +503,9 @@ async def _run_remote(
|
||||
deadline = _default_deadline(call.operation, params.get("text"))
|
||||
|
||||
try:
|
||||
task = scheduler.submit(
|
||||
submit = getattr(scheduler, "submit_async", None)
|
||||
submit = submit if callable(submit) else scheduler.submit
|
||||
submitted = submit(
|
||||
operation=call.operation,
|
||||
engine=call.engine,
|
||||
model_id=call.model_id,
|
||||
@@ -508,6 +514,7 @@ async def _run_remote(
|
||||
deadline_seconds=deadline,
|
||||
pinned_worker_id=decision.worker_id,
|
||||
)
|
||||
task = await submitted if asyncio.iscoroutine(submitted) else submitted
|
||||
except QueueFull as exc:
|
||||
raise _NotDispatched(str(exc)) from exc
|
||||
|
||||
|
||||
@@ -154,6 +154,17 @@ def bundled_dir() -> str:
|
||||
return os.path.join(media_tools_dir(), f"ffbin-{_FFBIN_COMMIT[:12]}", _platform_key())
|
||||
|
||||
|
||||
def _publish_bundled_on_path() -> None:
|
||||
"""Make a newly validated bundle visible to bare-name subprocess calls."""
|
||||
directory = os.path.abspath(bundled_dir())
|
||||
current = os.environ.get("PATH", "")
|
||||
entries = current.split(os.pathsep) if current else []
|
||||
if os.path.normcase(directory) in {os.path.normcase(entry) for entry in entries if entry}:
|
||||
return
|
||||
os.environ["PATH"] = os.pathsep.join([directory, *entries])
|
||||
logger.info("Published the acquired media-tool directory on PATH")
|
||||
|
||||
|
||||
def _exe(name: str) -> str:
|
||||
return f"{name}.exe" if sys.platform == "win32" else name
|
||||
|
||||
@@ -231,12 +242,21 @@ def acquire_bundled(wait: bool = False) -> dict:
|
||||
_ops["acquire"].update(state="running", progress=0.0, error=None)
|
||||
|
||||
if all(bundled_tool_path(t) and _binary_runs(bundled_tool_path(t)) for t in TOOLS):
|
||||
# The bundle may have arrived after startup's one-time PATH publish
|
||||
# (first-run acquisition is asynchronous). Make it visible to pydub
|
||||
# and other dependencies that launch ffmpeg/ffprobe by bare name now,
|
||||
# without requiring a backend restart (#1677).
|
||||
_publish_bundled_on_path()
|
||||
_set_op("acquire", state="done", progress=1.0)
|
||||
return _op_snapshot()["acquire"]
|
||||
|
||||
def _worker():
|
||||
try:
|
||||
_do_acquire()
|
||||
# Startup cannot publish binaries which do not exist yet. The
|
||||
# background worker must complete that second half atomically with
|
||||
# installation so the very next synthesis can use the tools.
|
||||
_publish_bundled_on_path()
|
||||
_set_op("acquire", state="done", progress=1.0, error=None)
|
||||
logger.info("media-tools: bundled ffmpeg/ffprobe installed at %s", bundled_dir())
|
||||
except Exception as e:
|
||||
|
||||
@@ -4,8 +4,9 @@ import sys
|
||||
import time
|
||||
import asyncio
|
||||
import logging
|
||||
import queue
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor, Executor
|
||||
from concurrent.futures import Executor, Future, ThreadPoolExecutor
|
||||
|
||||
from utils.containment import contain_system_exit
|
||||
|
||||
@@ -397,7 +398,13 @@ def __getattr__(name: str):
|
||||
# (generation.py, tts_stream.py) were the last unguarded dispatch — and the
|
||||
# residual on-main reports all fail on generate:start (audio). This is the same
|
||||
# guard generalised so every GPU dispatch shares one recovery path.
|
||||
_GENERATE_TIMEOUT_EXPLICIT = "OMNIVOICE_GENERATE_TIMEOUT_S" in os.environ
|
||||
GPU_JOB_TIMEOUT_S = float(os.environ.get("OMNIVOICE_GENERATE_TIMEOUT_S", "300.0"))
|
||||
_CONFIGURED_GPU_JOB_TIMEOUT_S = GPU_JOB_TIMEOUT_S
|
||||
# CPU synthesis is healthy but substantially slower than accelerated inference.
|
||||
# Keep a separate, bounded floor so a short render on CPU is not abandoned at
|
||||
# the GPU-oriented five-minute deadline (#1588).
|
||||
CPU_JOB_TIMEOUT_S = float(os.environ.get("OMNIVOICE_CPU_GENERATE_TIMEOUT_S", "600.0"))
|
||||
|
||||
# Queue-wait budget — a SEPARATE, deliberately generous clock (#1190/#1202).
|
||||
# The execution bound above must never be spent waiting in line: a job queued
|
||||
@@ -500,7 +507,9 @@ class GpuPoolBusyError(TimeoutError):
|
||||
self.retry_after = max(1, int(round(retry_after)))
|
||||
|
||||
|
||||
def generate_timeout_s(text: "str | None") -> float:
|
||||
def generate_timeout_s(
|
||||
text: "str | None", *, engine: object = None, execution_device: "str | None" = None,
|
||||
) -> float:
|
||||
"""THE wall-clock execution budget for one synthesis job, scaled to input.
|
||||
|
||||
Single source of truth for every TTS dispatch (#1190/#1202). The
|
||||
@@ -516,10 +525,33 @@ def generate_timeout_s(text: "str | None") -> float:
|
||||
CPU-class hardware, still bounded (a wedged job is caught in minutes, not
|
||||
hours).
|
||||
"""
|
||||
return max(
|
||||
GPU_JOB_TIMEOUT_S,
|
||||
GPU_JOB_TIMEOUT_S + (max(0, len(text or "") - 1200) / 40.0),
|
||||
)
|
||||
base = GPU_JOB_TIMEOUT_S
|
||||
try:
|
||||
from core.device_caps import detect_host_caps
|
||||
family = execution_device or detect_host_caps().family
|
||||
if execution_device is None and engine is not None:
|
||||
from services.engine_routing import resolve_routing
|
||||
compat = getattr(engine, "gpu_compat", None)
|
||||
if compat is None:
|
||||
compat = getattr(type(engine), "gpu_compat", (family, "cpu"))
|
||||
if tuple(compat) == ("cpu",):
|
||||
family = "cpu"
|
||||
else:
|
||||
family = resolve_routing(
|
||||
compat, detect_host_caps(),
|
||||
float(getattr(engine, "min_vram_gb", 0.0) or 0.0),
|
||||
)["effective_device"]
|
||||
universal_override = (
|
||||
_GENERATE_TIMEOUT_EXPLICIT
|
||||
or GPU_JOB_TIMEOUT_S != _CONFIGURED_GPU_JOB_TIMEOUT_S
|
||||
)
|
||||
if family == "cpu" and not universal_override:
|
||||
base = CPU_JOB_TIMEOUT_S
|
||||
except Exception:
|
||||
# Device probing is advisory here; the configured universal bound is
|
||||
# still safe when a platform probe is unavailable during startup.
|
||||
pass
|
||||
return base + (max(0, len(text or "") - 1200) / 40.0)
|
||||
|
||||
|
||||
def _retry_after_estimate(stats: dict) -> float:
|
||||
@@ -599,7 +631,8 @@ async def run_on_gpu_pool_guarded(fn, *, what: str = "GPU job",
|
||||
timeout: "float | None" = None,
|
||||
executor=None,
|
||||
queue_timeout: "float | None" = None,
|
||||
min_vram_gb: float = 0.0):
|
||||
min_vram_gb: float = 0.0,
|
||||
on_abandon=None):
|
||||
"""Run blocking ``fn`` on the GPU pool, bounding **execution** — not the
|
||||
wait for a free worker.
|
||||
|
||||
@@ -627,6 +660,12 @@ async def run_on_gpu_pool_guarded(fn, *, what: str = "GPU job",
|
||||
at 0 — the default, and correct for every non-TTS job on this pool
|
||||
(reference transcribe, watermarking, dub steps) — the under-provisioned-GPU
|
||||
wording is never used, because nothing measured says it applies (#1226).
|
||||
|
||||
``on_abandon`` is called once, after a job whose caller stopped waiting can
|
||||
no longer access its inputs. A queued job that is cancelled before it
|
||||
starts calls it immediately; a running thread calls it from ``_job``'s
|
||||
finalizer. Normal completion never calls it. This lets request-owned temp
|
||||
files outlive abandoned workers without delaying ordinary requests (#1668).
|
||||
"""
|
||||
loop = asyncio.get_running_loop()
|
||||
ex = executor if executor is not None else _get_gpu_pool()
|
||||
@@ -641,6 +680,24 @@ async def run_on_gpu_pool_guarded(fn, *, what: str = "GPU job",
|
||||
# job's model-load heartbeats (#1367). A dict, not a nonlocal: the closure
|
||||
# runs on a pool thread while the waiter reads from the event loop.
|
||||
_ident_box: dict = {}
|
||||
_abandon_lock = threading.Lock()
|
||||
_abandon_state = {
|
||||
"requested": False,
|
||||
"finished": False,
|
||||
"callback_called": False,
|
||||
}
|
||||
|
||||
def _fire_abandon_callback() -> None:
|
||||
if on_abandon is None:
|
||||
return
|
||||
with _abandon_lock:
|
||||
if _abandon_state["callback_called"]:
|
||||
return
|
||||
_abandon_state["callback_called"] = True
|
||||
try:
|
||||
on_abandon()
|
||||
except Exception: # noqa: BLE001 — cleanup cannot hide the pool result
|
||||
logger.exception("%s abandon cleanup failed", _log_safe(what))
|
||||
|
||||
def _job():
|
||||
# First thing the worker does: tell the awaiting coroutine the
|
||||
@@ -657,8 +714,26 @@ async def run_on_gpu_pool_guarded(fn, *, what: str = "GPU job",
|
||||
# Idents are reused by the OS; a stale heartbeat under this ident
|
||||
# must not vouch for some future job on the same thread.
|
||||
_MODEL_LOAD_ACTIVITY.pop(threading.get_ident(), None)
|
||||
with _abandon_lock:
|
||||
_abandon_state["finished"] = True
|
||||
abandoned = _abandon_state["requested"]
|
||||
if abandoned:
|
||||
_fire_abandon_callback()
|
||||
|
||||
concurrent_fut = ex.submit(_job)
|
||||
fut = asyncio.wrap_future(concurrent_fut, loop=loop)
|
||||
|
||||
def _abandon() -> None:
|
||||
# Keep the concurrent future so we can distinguish a job cancelled out
|
||||
# of the queue from a thread that Python cannot stop once it has begun.
|
||||
cancelled_before_start = concurrent_fut.cancel()
|
||||
with _abandon_lock:
|
||||
_abandon_state["requested"] = True
|
||||
finished = _abandon_state["finished"]
|
||||
fut.cancel()
|
||||
if cancelled_before_start or finished:
|
||||
_fire_abandon_callback()
|
||||
|
||||
fut = loop.run_in_executor(ex, _job)
|
||||
waiter = asyncio.ensure_future(started.wait())
|
||||
try:
|
||||
# Phase 1 — queue wait. Watch the future too, so a job that fails or is
|
||||
@@ -671,6 +746,7 @@ async def run_on_gpu_pool_guarded(fn, *, what: str = "GPU job",
|
||||
# Caller went away (client disconnect). We stop awaiting the job, so
|
||||
# make sure its eventual result/exception is consumed rather than
|
||||
# logged as "Future exception was never retrieved".
|
||||
_abandon()
|
||||
fut.add_done_callback(_swallow_abandoned)
|
||||
raise
|
||||
finally:
|
||||
@@ -680,7 +756,7 @@ async def run_on_gpu_pool_guarded(fn, *, what: str = "GPU job",
|
||||
# Never picked up: cancel it out of the queue (a not-yet-started
|
||||
# concurrent future cancels cleanly) and report saturation, NOT a
|
||||
# too-heavy job.
|
||||
fut.cancel()
|
||||
_abandon()
|
||||
fut.add_done_callback(_swallow_abandoned)
|
||||
stats = gpu_pool_stats(ex)
|
||||
logger.warning(
|
||||
@@ -751,14 +827,14 @@ async def run_on_gpu_pool_guarded(fn, *, what: str = "GPU job",
|
||||
# Caller went away mid-execution. The old wait_for cancelled the
|
||||
# wrapper itself; asyncio.wait does not, so do both halves here or the
|
||||
# eventual result is logged as "Future exception was never retrieved".
|
||||
fut.cancel()
|
||||
_abandon()
|
||||
fut.add_done_callback(_swallow_abandoned)
|
||||
raise
|
||||
except asyncio.TimeoutError as timeout_exc:
|
||||
# Parity with the old wait_for semantics: cancel the asyncio wrapper;
|
||||
# the worker thread keeps going regardless. Consume whatever it
|
||||
# eventually produces.
|
||||
fut.cancel()
|
||||
_abandon()
|
||||
fut.add_done_callback(_swallow_abandoned)
|
||||
# Capture the stacks BEFORE reset(): reset() replaces the executor, and
|
||||
# once the wedged thread is no longer a pool worker we can no longer
|
||||
@@ -1057,21 +1133,166 @@ def _timeout_guidance(
|
||||
# doubling the effective queue depth of a streamed multi-chunk render.
|
||||
# Giving it its own tiny pool removes that head-of-line blocking with no VRAM
|
||||
# risk, because the work was never on the device to begin with.
|
||||
_watermark_pool_singleton: "ThreadPoolExecutor | None" = None
|
||||
_watermark_pool_lock = threading.Lock()
|
||||
_WATERMARK_STOP = object()
|
||||
|
||||
|
||||
def get_watermark_pool() -> ThreadPoolExecutor:
|
||||
"""Dedicated 1-worker pool for provenance marking. Built lazily so hosts
|
||||
with watermarking disabled never spawn the thread."""
|
||||
global _watermark_pool_singleton
|
||||
if _watermark_pool_singleton is None:
|
||||
with _watermark_pool_lock:
|
||||
if _watermark_pool_singleton is None:
|
||||
_watermark_pool_singleton = ThreadPoolExecutor(
|
||||
max_workers=1, thread_name_prefix="watermark",
|
||||
class _WatermarkExecutor(Executor):
|
||||
"""Single daemon worker with a bounded shutdown contract.
|
||||
|
||||
``ThreadPoolExecutor`` uses non-daemon workers that Python joins at exit,
|
||||
so ``wait=False`` still delays process exit while ``wait=True`` can hang
|
||||
lifespan teardown forever. AudioSeal loading is not cooperatively
|
||||
cancellable; a daemon worker plus a bounded join is the only thread-based
|
||||
contract that both preserves in-process model warm-up and guarantees exit.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._items: queue.Queue = queue.Queue()
|
||||
self._lock = threading.Lock()
|
||||
self._shutdown = False
|
||||
self._thread: threading.Thread | None = None
|
||||
|
||||
def submit(self, fn, /, *args, **kwargs) -> Future:
|
||||
future: Future = Future()
|
||||
with self._lock:
|
||||
if self._shutdown:
|
||||
raise RuntimeError("cannot schedule new futures after shutdown")
|
||||
if self._thread is None:
|
||||
self._thread = threading.Thread(
|
||||
target=self._run,
|
||||
name="watermark_0",
|
||||
daemon=True,
|
||||
)
|
||||
return _watermark_pool_singleton
|
||||
self._thread.start()
|
||||
self._items.put((future, fn, args, kwargs))
|
||||
return future
|
||||
|
||||
def _run(self) -> None:
|
||||
while True:
|
||||
item = self._items.get()
|
||||
if item is _WATERMARK_STOP:
|
||||
return
|
||||
future, fn, args, kwargs = item
|
||||
if not future.set_running_or_notify_cancel():
|
||||
continue
|
||||
try:
|
||||
future.set_result(fn(*args, **kwargs))
|
||||
except (Exception, SystemExit, KeyboardInterrupt) as exc:
|
||||
future.set_exception(exc)
|
||||
|
||||
def is_stopped(self) -> bool:
|
||||
"""Whether shutdown has completed and this executor can be replaced."""
|
||||
with self._lock:
|
||||
return self._shutdown and (
|
||||
self._thread is None or not self._thread.is_alive()
|
||||
)
|
||||
|
||||
def is_shutdown(self) -> bool:
|
||||
with self._lock:
|
||||
return self._shutdown
|
||||
|
||||
def shutdown(
|
||||
self,
|
||||
wait: bool = True,
|
||||
*,
|
||||
cancel_futures: bool = False,
|
||||
timeout: float | None = None,
|
||||
) -> bool:
|
||||
with self._lock:
|
||||
self._shutdown = True
|
||||
thread = self._thread
|
||||
if cancel_futures:
|
||||
while True:
|
||||
try:
|
||||
item = self._items.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
if item is not _WATERMARK_STOP:
|
||||
item[0].cancel()
|
||||
self._items.put(_WATERMARK_STOP)
|
||||
if wait and thread is not None:
|
||||
thread.join(timeout=timeout)
|
||||
return thread is None or not thread.is_alive()
|
||||
|
||||
|
||||
_watermark_pool_singleton: "_WatermarkExecutor | None" = None
|
||||
_watermark_pool_lock = threading.Lock()
|
||||
_watermark_pool_accepting = True
|
||||
|
||||
|
||||
def begin_watermark_pool_lifecycle() -> None:
|
||||
"""Open watermark submissions for a newly-started app lifespan."""
|
||||
global _watermark_pool_accepting, _watermark_pool_singleton
|
||||
with _watermark_pool_lock:
|
||||
if (
|
||||
_watermark_pool_singleton is not None
|
||||
and _watermark_pool_singleton.is_stopped()
|
||||
):
|
||||
_watermark_pool_singleton = None
|
||||
_watermark_pool_accepting = (
|
||||
_watermark_pool_singleton is None
|
||||
or not _watermark_pool_singleton.is_shutdown()
|
||||
)
|
||||
|
||||
|
||||
def get_watermark_pool() -> _WatermarkExecutor:
|
||||
"""Dedicated 1-worker pool for provenance marking. Built lazily so hosts
|
||||
with watermarking disabled never spawn the thread.
|
||||
|
||||
The executor is captured and returned UNDER the lock: reading the global
|
||||
again after an unlocked null-check could race shutdown_watermark_pool's
|
||||
reset and hand out None (CodeRabbit, PR #1577)."""
|
||||
global _watermark_pool_accepting, _watermark_pool_singleton
|
||||
with _watermark_pool_lock:
|
||||
if not _watermark_pool_accepting:
|
||||
if (
|
||||
_watermark_pool_singleton is not None
|
||||
and _watermark_pool_singleton.is_stopped()
|
||||
):
|
||||
_watermark_pool_singleton = None
|
||||
_watermark_pool_accepting = True
|
||||
else:
|
||||
raise RuntimeError("watermark executor is shutting down")
|
||||
if (
|
||||
_watermark_pool_singleton is not None
|
||||
and _watermark_pool_singleton.is_stopped()
|
||||
):
|
||||
_watermark_pool_singleton = None
|
||||
if _watermark_pool_singleton is None:
|
||||
_watermark_pool_singleton = _WatermarkExecutor()
|
||||
return _watermark_pool_singleton
|
||||
|
||||
|
||||
def shutdown_watermark_pool(*, timeout: float = 20.0) -> None:
|
||||
"""Drain the watermark pool at app shutdown (PR #1577).
|
||||
|
||||
Refuse queued work and wait for the active operation: Python cannot kill
|
||||
a thread inside AudioSeal loading, so returning early would let model
|
||||
initialization continue during interpreter teardown. The draining pool
|
||||
remains published until its worker stops, preventing concurrent producers
|
||||
from creating a replacement that escapes this shutdown. A process that
|
||||
keeps running after lifespan shutdown (the test suite does exactly this)
|
||||
gets a fresh pool once the old worker has actually stopped."""
|
||||
global _watermark_pool_accepting, _watermark_pool_singleton
|
||||
with _watermark_pool_lock:
|
||||
_watermark_pool_accepting = False
|
||||
pool = _watermark_pool_singleton
|
||||
if pool is not None:
|
||||
stopped = pool.shutdown(
|
||||
wait=True,
|
||||
cancel_futures=True,
|
||||
timeout=max(0.0, float(timeout)),
|
||||
)
|
||||
if stopped:
|
||||
with _watermark_pool_lock:
|
||||
if _watermark_pool_singleton is pool:
|
||||
_watermark_pool_singleton = None
|
||||
else:
|
||||
logger.warning(
|
||||
"Watermark worker exceeded the %.1fs shutdown deadline; "
|
||||
"abandoning its daemon thread",
|
||||
timeout,
|
||||
)
|
||||
|
||||
|
||||
model = None # type: ignore
|
||||
@@ -2628,6 +2849,21 @@ async def preload_model():
|
||||
if model is not None:
|
||||
return # already loaded
|
||||
|
||||
# On MPS the configured ``omnivoice`` id resolves to a crash-isolated
|
||||
# sidecar. Warming the native singleton here would put the same fatal MPS
|
||||
# allocator risk back into the API process before the isolated engine is
|
||||
# ever asked to synthesize.
|
||||
try:
|
||||
from core.device_caps import detect_host_caps
|
||||
|
||||
if detect_host_caps().family == "mps":
|
||||
logger.info(
|
||||
"Native TTS preload skipped: OmniVoice uses crash isolation on this host."
|
||||
)
|
||||
return
|
||||
except Exception: # noqa: BLE001 -- preload selection must stay best-effort
|
||||
logger.debug("effective TTS preload selection failed", exc_info=True)
|
||||
|
||||
# A machine lending its GPU has no local user to warm the model FOR. This
|
||||
# preload exists to make the first /generate feel instant for the person
|
||||
# sitting in front of the app; on a headless node there is nobody sitting
|
||||
@@ -2717,7 +2953,10 @@ async def preload_model():
|
||||
"The TTS model could not be loaded. Settings → Logs → Backend "
|
||||
"has the full error."
|
||||
)
|
||||
_set_loading("failed", detail, error=detail)
|
||||
# `sub_stage` is a public API enum and the frontend keys failure state
|
||||
# off `error`. Keep the human-readable word "failed" in the detail,
|
||||
# not in the state machine (#1695).
|
||||
_set_loading("error", detail, error=detail)
|
||||
|
||||
def get_model_status():
|
||||
is_loaded = model is not None
|
||||
|
||||
@@ -79,6 +79,7 @@ class Segment:
|
||||
text: str
|
||||
speaker_id: str = "Speaker 1"
|
||||
id: str = field(default_factory=lambda: str(uuid.uuid4())[:8])
|
||||
extra: dict = field(default_factory=dict)
|
||||
|
||||
@property
|
||||
def duration(self) -> float:
|
||||
@@ -90,6 +91,7 @@ class Segment:
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
**self.extra,
|
||||
"id": self.id,
|
||||
"start": round(self.start, 2),
|
||||
"end": round(self.end, 2),
|
||||
@@ -98,6 +100,46 @@ class Segment:
|
||||
}
|
||||
|
||||
|
||||
def _merge_segment_extra(target: Segment, incoming: Segment, *, prepend: bool) -> None:
|
||||
"""Preserve editor metadata when cleanup folds ``incoming`` into ``target``."""
|
||||
for key, value in incoming.extra.items():
|
||||
target.extra.setdefault(key, value)
|
||||
|
||||
def joined(left: object, right: object) -> str:
|
||||
return _clean(f"{left or ''} {right or ''}")
|
||||
|
||||
target_original = target.extra.get("text_original")
|
||||
incoming_original = incoming.extra.get("text_original")
|
||||
if target_original is not None or incoming_original is not None:
|
||||
target.extra["text_original"] = (
|
||||
joined(incoming_original, target_original)
|
||||
if prepend
|
||||
else joined(target_original, incoming_original)
|
||||
)
|
||||
|
||||
raw_target_translations = target.extra.get("translations")
|
||||
raw_incoming_translations = incoming.extra.get("translations")
|
||||
target_translations = raw_target_translations if isinstance(raw_target_translations, dict) else {}
|
||||
incoming_translations = (
|
||||
raw_incoming_translations if isinstance(raw_incoming_translations, dict) else {}
|
||||
)
|
||||
if target_translations or incoming_translations:
|
||||
merged = {}
|
||||
languages = {
|
||||
*target_translations.keys(),
|
||||
*incoming_translations.keys(),
|
||||
}
|
||||
for language in languages:
|
||||
target_text = target_translations.get(language)
|
||||
incoming_text = incoming_translations.get(language)
|
||||
merged[language] = (
|
||||
joined(incoming_text, target_text)
|
||||
if prepend
|
||||
else joined(target_text, incoming_text)
|
||||
)
|
||||
target.extra["translations"] = merged
|
||||
|
||||
|
||||
def _clean(text: str) -> str:
|
||||
return _WS.sub(" ", (text or "").strip())
|
||||
|
||||
@@ -317,12 +359,14 @@ def _merge_short(segments: List[Segment]) -> List[Segment]:
|
||||
i += 1
|
||||
continue
|
||||
if target is prev:
|
||||
_merge_segment_extra(prev, s, prepend=False)
|
||||
prev.text = _clean(prev.text + " " + s.text)
|
||||
prev.end = max(prev.end, s.end)
|
||||
segments.pop(i)
|
||||
did_merge = True
|
||||
continue
|
||||
if target is nxt:
|
||||
_merge_segment_extra(nxt, s, prepend=True)
|
||||
nxt.text = _clean(s.text + " " + nxt.text)
|
||||
nxt.start = min(nxt.start, s.start)
|
||||
segments.pop(i)
|
||||
@@ -360,6 +404,7 @@ def _stitch_adjacent_shorts(segments: List[Segment]) -> List[Segment]:
|
||||
and b.duration <= STITCH_DUR
|
||||
and combined_dur <= MAX_DUR
|
||||
):
|
||||
_merge_segment_extra(a, b, prepend=False)
|
||||
a.text = _clean(a.text + " " + b.text)
|
||||
a.end = b.end
|
||||
segments.pop(i + 1)
|
||||
@@ -386,6 +431,11 @@ def clean_up_segments(segments: List[dict]) -> List[dict]:
|
||||
text=_clean(str(s.get("text", ""))),
|
||||
speaker_id=str(s.get("speaker_id") or "Speaker 1"),
|
||||
id=str(s.get("id") or uuid.uuid4().hex[:8]),
|
||||
extra={
|
||||
key: value
|
||||
for key, value in s.items()
|
||||
if key not in {"id", "start", "end", "text", "speaker_id"}
|
||||
},
|
||||
))
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
|
||||
@@ -254,6 +254,31 @@ def get_text(key: str, default: Optional[str] = None) -> Optional[str]:
|
||||
return default
|
||||
|
||||
|
||||
def get_text_state(key: str) -> tuple[bool, str]:
|
||||
"""Return ``(is_present, value)`` without hiding storage failures.
|
||||
|
||||
Rollback snapshots must distinguish a missing row from an unreadable
|
||||
database. ``get_text`` deliberately collapses those cases for ordinary
|
||||
preference reads, so transactional callers use this strict variant.
|
||||
"""
|
||||
if key == _TOKEN_KEY or key.startswith(_SECRET_PREFIX):
|
||||
raise ValueError(
|
||||
"get_text_state refuses to read an encrypted secret row; "
|
||||
"use get_hf_token()/get_secret() for secrets"
|
||||
)
|
||||
from core.db import db_conn
|
||||
|
||||
with db_conn() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT value FROM settings WHERE key = ?", (key,)
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return False, ""
|
||||
if row[0] is None:
|
||||
return True, ""
|
||||
return True, str(row[0])
|
||||
|
||||
|
||||
def set_text(key: str, value: str) -> None:
|
||||
"""Persist a non-encrypted text value into the settings table.
|
||||
|
||||
@@ -274,6 +299,19 @@ def set_text(key: str, value: str) -> None:
|
||||
)
|
||||
|
||||
|
||||
def clear_text(key: str) -> None:
|
||||
"""Remove a non-encrypted text setting, preserving a missing-row default."""
|
||||
if key == _TOKEN_KEY or key.startswith(_SECRET_PREFIX):
|
||||
raise ValueError(
|
||||
"clear_text refuses to delete an encrypted secret row; "
|
||||
"use clear_hf_token()/clear_secret() for secrets"
|
||||
)
|
||||
from core.db import db_conn
|
||||
|
||||
with db_conn() as conn:
|
||||
conn.execute("DELETE FROM settings WHERE key = ?", (key,))
|
||||
|
||||
|
||||
# ── Phase 4 Plan 04-01 (GGUF-04): per-engine quant override ────────────────
|
||||
#
|
||||
# Settings row "gguf_quant_override" holds either:
|
||||
|
||||
@@ -110,7 +110,7 @@ class SherpaModelSpec:
|
||||
# the same HF tree API on 2026-08-07 — not estimated. Every one of the seven
|
||||
# was wrong before, and in both directions, which is worse than uniformly
|
||||
# optimistic: the two Parakeets under-reported by ~3.8x (0.18 -> 0.67 GB),
|
||||
# so the recommended default quietly downloaded four times what the picker
|
||||
# so installing v3 quietly downloaded four times what the picker
|
||||
# promised on a metered or small-disk machine; but the two low-RAM
|
||||
# zipformers OVER-reported by ~3x (0.128 -> 0.044), making the fallback
|
||||
# models look bulkier than the heavyweights they exist to rescue users
|
||||
@@ -129,7 +129,6 @@ _MODELS: dict[str, SherpaModelSpec] = {
|
||||
kind="offline-transducer",
|
||||
size_gb=0.67,
|
||||
languages="25 European languages",
|
||||
recommended=True,
|
||||
heavy=True,
|
||||
model_type="nemo_transducer",
|
||||
files={
|
||||
@@ -223,6 +222,7 @@ _MODELS: dict[str, SherpaModelSpec] = {
|
||||
kind="offline-whisper",
|
||||
size_gb=0.104,
|
||||
languages="90+ languages (auto-detect)",
|
||||
recommended=True,
|
||||
files={
|
||||
"encoder": "tiny-encoder.int8.onnx",
|
||||
"decoder": "tiny-decoder.int8.onnx",
|
||||
@@ -231,7 +231,7 @@ _MODELS: dict[str, SherpaModelSpec] = {
|
||||
),
|
||||
}
|
||||
|
||||
DEFAULT_MODEL_ID = "sherpa-parakeet-tdt-v3"
|
||||
DEFAULT_MODEL_ID = "sherpa-whisper-tiny"
|
||||
|
||||
# repo_id → model id, so the model-store list (keyed by repo_id) can be
|
||||
# enriched with the dictation metadata, and so capture can map either key.
|
||||
@@ -261,6 +261,16 @@ def sherpa_available() -> tuple[bool, str]:
|
||||
return True, "ready"
|
||||
except ImportError as e:
|
||||
return False, f"sherpa-onnx not installed: {e}. Install with: uv add sherpa-onnx"
|
||||
except Exception as e: # noqa: BLE001 — an availability probe must fail closed
|
||||
# Native wheel failures surface as OSError/RuntimeError rather than
|
||||
# ImportError (missing DLL/dylib/so, loader or runtime init failure) —
|
||||
# but the set is open-ended: an extension module is free to raise
|
||||
# anything at init. This is an availability question, so ANY failure to
|
||||
# import means "not available", never an exception escaping to the
|
||||
# caller. SherpaDictationBackend.is_available() calls this directly and
|
||||
# capture_ws.ws_transcribe calls that without a guard, so an unexpected
|
||||
# type here took the WebSocket down instead of falling back (#1610).
|
||||
return False, f"sherpa-onnx unavailable ({type(e).__name__}): {e}"
|
||||
|
||||
|
||||
def _resolve_model_dir(spec: SherpaModelSpec, *, download: bool = True) -> str:
|
||||
@@ -397,13 +407,10 @@ def build_online_recognizer(spec: SherpaModelSpec, *, download: bool = True):
|
||||
# transcribe the same bytes. It is a defect inside sherpa-onnx that the app
|
||||
# cannot fix by configuration.
|
||||
#
|
||||
# The curated default therefore cannot be trusted to WORK just because it is
|
||||
# installed — and which platforms are affected is not knowable up front, so
|
||||
# hard-coding a different default per OS would only be a guess. Instead the app
|
||||
# learns from what it observes: when a session hears real speech and the model
|
||||
# returns nothing, that model is demoted on THIS machine and stops being
|
||||
# selected. Self-correcting wherever the breakage actually is, and a no-op
|
||||
# everywhere it isn't.
|
||||
# Installation alone therefore cannot prove that a recognizer works. When a
|
||||
# session hears real speech and the model returns nothing, that model is
|
||||
# demoted on this machine and stops being selected. This self-corrects wherever
|
||||
# the decoder defect appears and is a no-op everywhere it does not.
|
||||
|
||||
#: prefs key holding the list of model ids demoted on this machine.
|
||||
PREF_SILENT_MODELS = "dictation.silent_models"
|
||||
|
||||
@@ -60,6 +60,7 @@ from pathlib import Path
|
||||
from typing import Callable, Optional
|
||||
|
||||
from core.config import DATA_DIR
|
||||
from core.contained_subprocess import OwnedPopen, spawn_owned
|
||||
|
||||
logger = logging.getLogger("omnivoice.sidecar_install")
|
||||
|
||||
@@ -108,7 +109,11 @@ class SidecarSpec:
|
||||
weights_repo_id: Optional[str] = None # HF repo downloaded into <checkout>/<weights_subdir>
|
||||
weights_revision: Optional[str] = None # reviewed HF commit
|
||||
weights_subdir: str = "checkpoints"
|
||||
weights_config_name: str = "config.yaml" # required model config inside weights_subdir
|
||||
# Model-config filenames accepted inside weights_subdir. A tuple, not a
|
||||
# single name: IndexTTS 2.5's weights repo ships config.yaml, but installs
|
||||
# predating #1611 were only usable after hand-renaming it to
|
||||
# config_v2_5.yaml, and those must keep working without a reinstall.
|
||||
weights_config_names: tuple[str, ...] = ("config.yaml",)
|
||||
docs_path: str = "docs/engines" # where the manual-install fallback lives
|
||||
required_bytes: int = 12 * _GIB # conservative source+venv+weights estimate for preflight
|
||||
# Called after a successful install/uninstall so the engine's memoised
|
||||
@@ -146,7 +151,7 @@ SPECS: dict[str, SidecarSpec] = {
|
||||
weights_repo_id="IndexTeam/IndexTTS-2.5",
|
||||
weights_revision="d0aa86e75bb6f3437f3831e95056fa72842d89ef",
|
||||
weights_subdir="checkpoints",
|
||||
weights_config_name="config_v2_5.yaml",
|
||||
weights_config_names=("config.yaml", "config_v2_5.yaml"),
|
||||
docs_path="docs/engines/indextts.md",
|
||||
# ~0.1 GB source + up to ~6 GB venv (torch + transformers<5) +
|
||||
# ~6 GB weights. Deliberately conservative; the preflight subtracts
|
||||
@@ -879,11 +884,11 @@ def _weights_present(spec: SidecarSpec) -> bool:
|
||||
actual = marker[:2] if len(marker) >= 2 else marker + [""]
|
||||
if actual != expected:
|
||||
return False
|
||||
return _weights_floor_ok(wdir, config_name=spec.weights_config_name)
|
||||
return _weights_floor_ok(wdir, config_names=spec.weights_config_names)
|
||||
|
||||
|
||||
def _weights_floor_ok(wdir: Path, *, config_name: str = "config.yaml") -> bool:
|
||||
if not (wdir / config_name).is_file():
|
||||
def _weights_floor_ok(wdir: Path, *, config_names: tuple[str, ...] = ("config.yaml",)) -> bool:
|
||||
if not any((wdir / name).is_file() for name in config_names):
|
||||
return False
|
||||
floor = 5 * 1024 * 1024
|
||||
try:
|
||||
@@ -969,7 +974,7 @@ def _step_fetch_weights(spec: SidecarSpec, job: dict) -> None:
|
||||
hf_progress.unregister_listener(listener_id)
|
||||
hf_progress.current_repo_id.reset(repo_token)
|
||||
|
||||
if not _weights_floor_ok(wdir, config_name=spec.weights_config_name):
|
||||
if not _weights_floor_ok(wdir, config_names=spec.weights_config_names):
|
||||
raise _StepError(
|
||||
"Weight download finished but no plausible weight files were found — "
|
||||
"the download was likely interrupted.",
|
||||
@@ -994,6 +999,11 @@ def _step_persist(spec: SidecarSpec, job: dict) -> None:
|
||||
# ── Subprocess runner with live log capture ────────────────────────────────
|
||||
|
||||
|
||||
def _install_containment_kwargs() -> dict:
|
||||
"""Nested process-group/Job ownership is supplied by ``spawn_owned``."""
|
||||
return {}
|
||||
|
||||
|
||||
def _run_logged(job: dict, argv: list[str], *, timeout: float,
|
||||
env: "dict[str, str] | None" = None) -> int:
|
||||
"""Run *argv*, streaming combined stdout+stderr lines into the job log.
|
||||
@@ -1008,13 +1018,11 @@ def _run_logged(job: dict, argv: list[str], *, timeout: float,
|
||||
killed child — a blocking ``for line in proc.stdout`` on this thread
|
||||
would hang past the timeout waiting for pipe EOF.
|
||||
"""
|
||||
popen_kwargs: dict = {}
|
||||
if os.name == "posix":
|
||||
# New session → we can kill the whole process group on timeout
|
||||
# instead of only the direct child.
|
||||
popen_kwargs["start_new_session"] = True
|
||||
# ``spawn_owned`` creates the local timeout group/Job before the operation
|
||||
# starts and links it to backend death through its control pipe.
|
||||
popen_kwargs = _install_containment_kwargs()
|
||||
try:
|
||||
proc = subprocess.Popen(
|
||||
proc = spawn_owned(
|
||||
argv,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
@@ -1049,29 +1057,18 @@ def _run_logged(job: dict, argv: list[str], *, timeout: float,
|
||||
|
||||
|
||||
def _kill_tree(proc: "subprocess.Popen") -> None:
|
||||
"""Kill the child and its whole process tree, on every platform.
|
||||
|
||||
POSIX: the child was started in its own session, so SIGKILL the group.
|
||||
Windows: ``proc.kill()`` only terminates the direct child — a git/uv
|
||||
helper it spawned would keep running (and writing into the checkout)
|
||||
past our timeout — so use ``taskkill /T`` to fell the tree.
|
||||
"""
|
||||
if os.name == "posix":
|
||||
import signal
|
||||
"""Kill an operation through its stable nested group/Job owner."""
|
||||
if isinstance(proc, OwnedPopen):
|
||||
# The retained supervisor/process-group or nested Job is the stable
|
||||
# per-operation owner. Do not fall back to a direct PID kill.
|
||||
proc.kill()
|
||||
try:
|
||||
os.killpg(proc.pid, signal.SIGKILL)
|
||||
return
|
||||
except (ProcessLookupError, PermissionError, OSError):
|
||||
pass # group already gone / not ours — fall through to plain kill
|
||||
else: # Windows
|
||||
try:
|
||||
subprocess.run(
|
||||
["taskkill", "/F", "/T", "/PID", str(proc.pid)],
|
||||
capture_output=True, timeout=15,
|
||||
)
|
||||
return
|
||||
except (OSError, subprocess.SubprocessError):
|
||||
pass # taskkill unavailable/failed — fall through to plain kill
|
||||
proc.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
pass
|
||||
return
|
||||
# A test double or a legacy caller without the nested owner can only be
|
||||
# stopped through its stable direct-process handle.
|
||||
try:
|
||||
proc.kill()
|
||||
except OSError:
|
||||
|
||||
@@ -32,9 +32,9 @@ Threat-model summary (see Plan 02-01 frontmatter):
|
||||
AUTH-05 installed (``HFTokenRedactor``) on the root logger.
|
||||
T-02-04 — compromised sidecar emitting unexpected ops: parent allowlist
|
||||
``PARENT_INBOUND_OPS`` rejects everything else.
|
||||
T-02-05 — Tauri group-kill scope: ``start_new_session=True`` on Unix
|
||||
and ``CREATE_NEW_PROCESS_GROUP`` on Windows isolate the
|
||||
sidecar's process group.
|
||||
T-02-05 — nested containment: a retained supervisor process group/Job owns
|
||||
each engine operation and is linked to backend death by a control
|
||||
pipe, while still permitting independent timeout teardown.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -56,6 +56,7 @@ from typing import Optional
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from core.contained_subprocess import spawn_owned
|
||||
from services.tts_backend import TTSBackend
|
||||
|
||||
logger = logging.getLogger("omnivoice.subprocess_backend")
|
||||
@@ -470,13 +471,6 @@ class SubprocessBackend(TTSBackend):
|
||||
"env": env,
|
||||
"bufsize": 0, # unbuffered binary pipes
|
||||
}
|
||||
# Process-group isolation so the Tauri lib.rs group-kill in shutdown
|
||||
# doesn't escape into other children. See T-02-05.
|
||||
if sys.platform == "win32":
|
||||
kwargs["creationflags"] = subprocess.CREATE_NEW_PROCESS_GROUP
|
||||
else:
|
||||
kwargs["start_new_session"] = True
|
||||
|
||||
# `venv_python()` resolves the engine's interpreter, and on a cold
|
||||
# first run that is not cheap: it spawns each candidate to import the
|
||||
# engine (bounded, but tens of seconds on a slow disk), and if none is
|
||||
@@ -509,7 +503,7 @@ class SubprocessBackend(TTSBackend):
|
||||
self.id, Path(python_path).name, Path(script_path).name,
|
||||
)
|
||||
try:
|
||||
self._proc = subprocess.Popen([python_path, script_path], **kwargs)
|
||||
self._proc = spawn_owned([python_path, script_path], **kwargs)
|
||||
except OSError as exc:
|
||||
raise InvalidBinaryError(
|
||||
python_path,
|
||||
|
||||
@@ -341,6 +341,44 @@ class TTSBackend(ABC):
|
||||
Engines that don't support this will ignore the parameter.
|
||||
"""
|
||||
|
||||
def generate_batch(
|
||||
self,
|
||||
texts: list[str],
|
||||
*,
|
||||
ref_audio=None,
|
||||
ref_text=None,
|
||||
instruct=None,
|
||||
language=None,
|
||||
duration=None,
|
||||
speed=1.0,
|
||||
**extras,
|
||||
) -> list[torch.Tensor]:
|
||||
"""Synthesize several utterances, preserving the single-item contract.
|
||||
|
||||
Engines with a native batch forward pass override this method. The
|
||||
default keeps every existing adapter correct while giving callers one
|
||||
stable seam and per-item keyword handling.
|
||||
"""
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
def _item(value, index):
|
||||
return value[index] if isinstance(value, list) else value
|
||||
|
||||
return [
|
||||
self.generate(
|
||||
text,
|
||||
ref_audio=_item(ref_audio, index),
|
||||
ref_text=_item(ref_text, index),
|
||||
instruct=_item(instruct, index),
|
||||
language=_item(language, index),
|
||||
duration=_item(duration, index),
|
||||
speed=_item(speed, index),
|
||||
**extras,
|
||||
)
|
||||
for index, text in enumerate(texts)
|
||||
]
|
||||
|
||||
# ── Lifecycle (Phase 2 will enforce per-engine overrides) ──────────────
|
||||
#
|
||||
# Today every backend lazily loads its weights on first `generate()` and
|
||||
@@ -631,7 +669,7 @@ class OmniVoiceBackend(TTSBackend):
|
||||
|
||||
id = "omnivoice"
|
||||
display_name = "VoiceStudio (k2-fsa/OmniVoice, 600+ languages)"
|
||||
gpu_compat = ("cuda", "mps", "cpu")
|
||||
gpu_compat = ("cuda", "rocm", "mps", "cpu")
|
||||
# Derived from the pool's own per-job budget (_GPU_VRAM_PER_JOB_GB = 5.0 in
|
||||
# model_manager, itself measured from the ~1.6 GB forward + autoregressive
|
||||
# decode and the co-loaded WhisperX on the clone path), plus room for the
|
||||
@@ -717,6 +755,73 @@ class OmniVoiceBackend(TTSBackend):
|
||||
)
|
||||
return audios[0]
|
||||
|
||||
def generate_batch(self, texts: list[str], **kw) -> list[torch.Tensor]:
|
||||
"""Use OmniVoice's native variable-length batch generation.
|
||||
|
||||
Batch callers pass per-item language, duration, speed and reference
|
||||
lists. Reusable clone prompts are prepared once and handed to the
|
||||
model together; an incomplete prompt batch falls back to the proven
|
||||
single-item path instead of changing synthesis semantics.
|
||||
"""
|
||||
self._ensure_loaded()
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
def _items(value):
|
||||
if isinstance(value, list):
|
||||
return value
|
||||
return [value] * len(texts)
|
||||
|
||||
def _item_kwargs(index):
|
||||
return {
|
||||
key: value[index] if isinstance(value, list) else value
|
||||
for key, value in kw.items()
|
||||
}
|
||||
|
||||
ref_audios = _items(kw.get("ref_audio"))
|
||||
ref_texts = _items(kw.get("ref_text"))
|
||||
cache_ref = bool(kw.get("cache_ref", True))
|
||||
preprocess_prompt = bool(kw.get("preprocess_prompt", True))
|
||||
prompts = []
|
||||
if any(ref_audios):
|
||||
for ref_audio, ref_text in zip(ref_audios, ref_texts):
|
||||
if not ref_audio:
|
||||
prompts = []
|
||||
break
|
||||
prompt = _get_clone_prompt(
|
||||
self._model,
|
||||
ref_audio,
|
||||
ref_text,
|
||||
preprocess_prompt,
|
||||
store=cache_ref,
|
||||
)
|
||||
if prompt is None:
|
||||
prompts = []
|
||||
break
|
||||
prompts.append(prompt)
|
||||
|
||||
if any(ref_audios) and len(prompts) != len(texts):
|
||||
return [self.generate(text, **_item_kwargs(i))
|
||||
for i, text in enumerate(texts)]
|
||||
|
||||
gen_kw = dict(
|
||||
language=kw.get("language"),
|
||||
instruct=kw.get("instruct"),
|
||||
duration=kw.get("duration"),
|
||||
speed=kw.get("speed", 1.0),
|
||||
denoise=kw.get("denoise", True),
|
||||
postprocess_output=kw.get("postprocess_output", True),
|
||||
num_step=kw.get("num_step", 16),
|
||||
guidance_scale=kw.get("guidance_scale", 2.0),
|
||||
preprocess_prompt=preprocess_prompt,
|
||||
)
|
||||
if prompts:
|
||||
gen_kw["voice_clone_prompt"] = prompts
|
||||
else:
|
||||
gen_kw["ref_audio"] = None
|
||||
gen_kw["ref_text"] = None
|
||||
return self._model.generate(text=texts, **gen_kw)
|
||||
|
||||
def unload(self) -> None:
|
||||
"""Release the OmniVoice model (MM2-02). OmniVoice shares the singleton
|
||||
owned by ``model_manager``, so dropping our local ref isn't enough — we
|
||||
@@ -2293,6 +2398,7 @@ def list_backends() -> list[dict]:
|
||||
|
||||
out: list[dict] = []
|
||||
for bid, cls in _REGISTRY.items():
|
||||
cls = _effective_backend_class(bid, cls, caps.family)
|
||||
try:
|
||||
ok, msg = cls.is_available()
|
||||
except Exception:
|
||||
@@ -2366,10 +2472,29 @@ def list_backends() -> list[dict]:
|
||||
return out
|
||||
|
||||
|
||||
def _effective_backend_class(
|
||||
backend_id: str,
|
||||
backend_cls: type[TTSBackend],
|
||||
host_family: str | None = None,
|
||||
) -> type[TTSBackend]:
|
||||
"""Resolve host-specific containment without changing the configured id."""
|
||||
if backend_id != "omnivoice":
|
||||
return backend_cls
|
||||
if host_family is None:
|
||||
from core.device_caps import detect_host_caps
|
||||
|
||||
host_family = detect_host_caps().family
|
||||
if host_family != "mps":
|
||||
return backend_cls
|
||||
from engines.omnivoice_subprocess import OmniVoiceMPSSubprocessBackend
|
||||
|
||||
return OmniVoiceMPSSubprocessBackend
|
||||
|
||||
|
||||
def get_backend_class(backend_id: str) -> type[TTSBackend]:
|
||||
if backend_id not in _REGISTRY:
|
||||
raise ValueError(f"Unknown TTS backend: {backend_id!r}. Known: {list(_REGISTRY)}")
|
||||
return _REGISTRY[backend_id]
|
||||
return _effective_backend_class(backend_id, _REGISTRY[backend_id])
|
||||
|
||||
|
||||
def cloning_capable_engine_ids() -> list[str]:
|
||||
|
||||
+266
-41
@@ -20,12 +20,17 @@ Usage:
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
import torch
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from core.prefs import resolve
|
||||
|
||||
logger = logging.getLogger("omnivoice.watermark")
|
||||
@@ -37,6 +42,20 @@ _detector = None
|
||||
_audioseal_available: Optional[bool] = None
|
||||
# Monotonic stamp of the last embed/detect, for the idle release below.
|
||||
_last_used = 0.0
|
||||
# Per-model locks for the lazy builds below: the startup prefetch thread
|
||||
# races the first embed, and both must share ONE build (a double load doubles
|
||||
# the cold-start cost the prefetch exists to hide). One lock PER MODEL — a
|
||||
# single shared lock made the ~42s generator prefetch block unrelated detector
|
||||
# loads and the idle reaper behind it. release_idle_models acquires both, in
|
||||
# this fixed order (nothing else nests them, so no cycle is possible).
|
||||
_generator_lock = threading.Lock()
|
||||
_detector_lock = threading.Lock()
|
||||
|
||||
# True when the generator exists ONLY because the startup prefetch built it
|
||||
# and no embed/detect has used it since. The idle reaper grants one extra
|
||||
# idle window before dropping such a model, so a first synthesis at minute
|
||||
# 20 still finds it warm (code-review finding 2 on the prefetch PR).
|
||||
_prefetched_unused = False
|
||||
|
||||
# 16-bit message: "OM" in ASCII = 0x4F 0x4D = 0100_1111 0100_1101
|
||||
# This is our signature — every VoiceStudio-generated audio carries it.
|
||||
@@ -50,6 +69,86 @@ OMNI_MESSAGE = [0, 1, 0, 0, 1, 1, 1, 1, 0, 1, 0, 0, 1, 1, 0, 1]
|
||||
_CHUNK_SECONDS = 30
|
||||
|
||||
|
||||
# AudioSeal vendors moshi's ``@torch_compile_lazy`` on SEANetEncoder.forward,
|
||||
# so the first EMBED — not the model load, which prefetch already warms —
|
||||
# calls torch.compile and drops into Inductor's C++ codegen. On hosts whose
|
||||
# C++ toolchain can't serve Inductor that compile raises CppCompileError, the
|
||||
# embed fail-opens, and audio ships unmarked: a macOS arm64 deployment lost
|
||||
# provenance marking on 10/10 takes while paying 30-40 s for the first failed
|
||||
# compile and 5-8 s for each later one (#1615).
|
||||
#
|
||||
# The compile is pure cost even where it succeeds. Measured on an M3 (5 s of
|
||||
# 24 kHz audio, three consecutive embeds): compiled 9.70 / 0.26 / 0.23 s vs
|
||||
# eager 0.30 / 0.28 / 0.27 s — a ~10 s first-embed tax to save ~0.03 s per
|
||||
# later embed, on CPU work that is already bounded by the 30 s chunk loop.
|
||||
# So watermarking runs eager on every platform.
|
||||
def _moshi_compile_module():
|
||||
"""AudioSeal's vendored moshi compile switch module, or None.
|
||||
|
||||
Resolved per call rather than at import: ``_check_available()`` is what
|
||||
guarantees audioseal is importable, and it runs later than this module.
|
||||
"""
|
||||
try:
|
||||
from audioseal.libs.moshi.utils import compile as moshi_compile
|
||||
except Exception: # noqa: BLE001 — any import shape change degrades, not crashes
|
||||
return None
|
||||
return moshi_compile
|
||||
|
||||
|
||||
_eager_lock = threading.Lock()
|
||||
#: Depth of nested/concurrent eager scopes, and the switch value to put back
|
||||
#: when the last one exits. One dict rather than two module scalars: the
|
||||
#: fields are only meaningful together, and only under _eager_lock.
|
||||
_eager_state: dict = {"depth": 0, "saved": None}
|
||||
_eager_guard_warned = False
|
||||
|
||||
|
||||
def _warn_missing_eager_guard() -> None:
|
||||
global _eager_guard_warned
|
||||
_eager_guard_warned = True
|
||||
logger.info(
|
||||
"audioseal's no_compile switch is unavailable — watermarking may run "
|
||||
"through torch.compile and pay (or fail) an Inductor C++ compile (#1615)."
|
||||
)
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _eager_audioseal():
|
||||
"""Run the AudioSeal model eagerly, restoring the switch on the way out.
|
||||
|
||||
Upstream's own ``no_compile()`` saves and restores ``_compile_disabled``
|
||||
per call, which is not safe when two watermark calls overlap: the first to
|
||||
exit restores False while the second is still mid-embed, handing it back
|
||||
the compile this whole fix exists to avoid. So the flag is reference
|
||||
counted here — it goes True on the outermost entry and only comes back on
|
||||
the outermost exit — rather than serializing embeds behind a lock, which
|
||||
would cost real throughput on concurrent generations.
|
||||
|
||||
Degrades to a plain call if a future audioseal drops the helper
|
||||
(``tests/test_watermark_no_torch_compile_1615.py`` fails loudly on that
|
||||
upgrade rather than letting the compile creep back in).
|
||||
"""
|
||||
moshi = _moshi_compile_module()
|
||||
if moshi is None:
|
||||
if not _eager_guard_warned:
|
||||
_warn_missing_eager_guard()
|
||||
yield
|
||||
return
|
||||
with _eager_lock:
|
||||
if _eager_state["depth"] == 0:
|
||||
_eager_state["saved"] = moshi._compile_disabled
|
||||
_eager_state["depth"] += 1
|
||||
moshi._compile_disabled = True
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
with _eager_lock:
|
||||
_eager_state["depth"] -= 1
|
||||
if _eager_state["depth"] == 0:
|
||||
moshi._compile_disabled = _eager_state["saved"]
|
||||
_eager_state["saved"] = None
|
||||
|
||||
|
||||
def _iter_chunks(audio: torch.Tensor, sample_rate: int):
|
||||
"""Yield ≤ ~_CHUNK_SECONDS slices of (batch, channels, samples) audio
|
||||
along the time axis. A sub-second tail is folded into the previous chunk
|
||||
@@ -77,28 +176,82 @@ def _check_available() -> bool:
|
||||
return _audioseal_available
|
||||
|
||||
|
||||
def _get_generator():
|
||||
"""Lazy-load the AudioSeal generator model."""
|
||||
global _generator, _last_used
|
||||
_last_used = time.monotonic()
|
||||
if _generator is None:
|
||||
from audioseal import AudioSeal
|
||||
_generator = AudioSeal.load_generator("audioseal_wm_16bits")
|
||||
_generator.eval()
|
||||
logger.info("AudioSeal generator loaded (16-bit message mode)")
|
||||
return _generator
|
||||
def _get_generator(mark_prefetched: bool = False):
|
||||
"""Lazy-load the AudioSeal generator model.
|
||||
|
||||
Owns the idle-reaper grace in ONE critical section: the startup prefetch
|
||||
claims it (``mark_prefetched=True``) only when THIS call builds the model,
|
||||
and every other call (a real embed) consumes it — no call-site blocks, no
|
||||
window between two lock scopes where the claim could land on an
|
||||
already-used model.
|
||||
"""
|
||||
global _generator, _last_used, _prefetched_unused
|
||||
with _generator_lock:
|
||||
_last_used = time.monotonic()
|
||||
if _generator is None:
|
||||
from audioseal import AudioSeal
|
||||
_generator = AudioSeal.load_generator("audioseal_wm_16bits")
|
||||
_generator.eval()
|
||||
logger.info("AudioSeal generator loaded (16-bit message mode)")
|
||||
_prefetched_unused = mark_prefetched
|
||||
elif not mark_prefetched:
|
||||
_prefetched_unused = False
|
||||
return _generator
|
||||
|
||||
|
||||
def _get_detector():
|
||||
"""Lazy-load the AudioSeal detector model."""
|
||||
global _detector, _last_used
|
||||
_last_used = time.monotonic()
|
||||
if _detector is None:
|
||||
from audioseal import AudioSeal
|
||||
_detector = AudioSeal.load_detector("audioseal_detector_16bits")
|
||||
_detector.eval()
|
||||
logger.info("AudioSeal detector loaded (16-bit message mode)")
|
||||
return _detector
|
||||
with _detector_lock:
|
||||
_last_used = time.monotonic()
|
||||
if _detector is None:
|
||||
from audioseal import AudioSeal
|
||||
_detector = AudioSeal.load_detector("audioseal_detector_16bits")
|
||||
_detector.eval()
|
||||
logger.info("AudioSeal detector loaded (16-bit message mode)")
|
||||
return _detector
|
||||
|
||||
|
||||
def _generator_checkpoint_cached() -> bool:
|
||||
"""Return whether AudioSeal can warm without contacting Hugging Face.
|
||||
|
||||
AudioSeal 0.2 stores the checkpoint in ``<cache>/audioseal`` even though
|
||||
it uses huggingface_hub to fetch it. Keep startup local-first: an ordinary
|
||||
boot may consume that file, but must never turn prefetch into a download.
|
||||
"""
|
||||
cache_root = os.environ.get("AUDIOSEAL_CACHE_DIR") or os.environ.get(
|
||||
"XDG_CACHE_HOME"
|
||||
)
|
||||
root = Path(cache_root).expanduser() if cache_root else Path.home() / ".cache"
|
||||
return (root / "audioseal" / "generator_base.pth").is_file()
|
||||
|
||||
|
||||
def prefetch_generator(*, allow_download: bool = False) -> None:
|
||||
"""Warm the AudioSeal generator eagerly (startup background thread).
|
||||
|
||||
The first ``mark_synthetic`` otherwise pays the audioseal import plus the
|
||||
generator load inline — measured at ~42 s on a cold filesystem (2026-08-17
|
||||
macOS deployment), serialized inside the first synthesis and 3 s short of
|
||||
a 90 s client timeout. Warming here overlaps that span with the TTS model
|
||||
load. No-op when watermarking is off or audioseal is absent; a failure
|
||||
logs and leaves the lazy path to retry on first embed. Default startup is
|
||||
also cache-only; a download is allowed only when the user explicitly set
|
||||
``OMNIVOICE_PRELOAD_WATERMARK=1``.
|
||||
"""
|
||||
try:
|
||||
if not will_mark():
|
||||
logger.debug("Watermark prefetch skipped (disabled or audioseal absent)")
|
||||
return
|
||||
if not allow_download and not _generator_checkpoint_cached():
|
||||
logger.info("Watermark prefetch skipped: AudioSeal checkpoint is not cached")
|
||||
return
|
||||
_get_generator(mark_prefetched=True)
|
||||
logger.info("AudioSeal generator prefetched in the background")
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Watermark prefetch failed; the first embed will retry inline",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
|
||||
def release_idle_models(idle_seconds: float, *, now: Optional[float] = None) -> bool:
|
||||
@@ -114,14 +267,28 @@ def release_idle_models(idle_seconds: float, *, now: Optional[float] = None) ->
|
||||
Returns True if anything was released. Never raises: this runs from the
|
||||
idle reaper, which must survive it.
|
||||
"""
|
||||
global _generator, _detector
|
||||
if _generator is None and _detector is None:
|
||||
return False
|
||||
stamp = time.monotonic() if now is None else float(now)
|
||||
if stamp - _last_used < idle_seconds:
|
||||
return False
|
||||
_generator = None
|
||||
_detector = None
|
||||
global _generator, _detector, _prefetched_unused
|
||||
with _generator_lock, _detector_lock:
|
||||
if _generator is None and _detector is None:
|
||||
return False
|
||||
stamp = time.monotonic() if now is None else float(now)
|
||||
if stamp - _last_used < idle_seconds:
|
||||
return False
|
||||
if _prefetched_unused:
|
||||
# The startup prefetch built the generator and nothing has used
|
||||
# it yet. Drop the grace (one extra idle window only) instead of
|
||||
# the model, so a first synthesis shortly after boot still finds
|
||||
# it warm — the exact scenario the prefetch exists for.
|
||||
_prefetched_unused = False
|
||||
logger.info(
|
||||
"Idle watermark models are prefetch-warmed but unused; "
|
||||
"granting one more idle window before releasing."
|
||||
)
|
||||
return False
|
||||
# Under the locks so a release racing the prefetch or a first embed
|
||||
# can't wipe a model the lazy path just built.
|
||||
_generator = None
|
||||
_detector = None
|
||||
logger.info("Idle timeout reached. Released the AudioSeal watermark models.")
|
||||
return True
|
||||
|
||||
@@ -200,6 +367,62 @@ def mark_synthetic(
|
||||
return marked
|
||||
|
||||
|
||||
async def mark_synthetic_async(
|
||||
waveform: torch.Tensor,
|
||||
sample_rate: int,
|
||||
*,
|
||||
context: str,
|
||||
force: bool = False,
|
||||
timeout: float | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Dispatch marking without letting a draining pool lose finished audio."""
|
||||
import asyncio
|
||||
import functools
|
||||
|
||||
from services.model_manager import (
|
||||
GpuJobTimeoutError,
|
||||
GpuPoolBusyError,
|
||||
get_watermark_pool,
|
||||
run_on_gpu_pool_guarded,
|
||||
)
|
||||
|
||||
try:
|
||||
pool = get_watermark_pool()
|
||||
except RuntimeError:
|
||||
logger.warning("Watermark skipped while the prior worker is shutting down")
|
||||
return waveform
|
||||
|
||||
job = functools.partial(
|
||||
mark_synthetic, waveform, sample_rate, context=context, force=force
|
||||
)
|
||||
try:
|
||||
if timeout is not None:
|
||||
return await run_on_gpu_pool_guarded(
|
||||
job, what="Audio watermark", timeout=timeout, executor=pool
|
||||
)
|
||||
return await asyncio.get_running_loop().run_in_executor(pool, job)
|
||||
except (GpuJobTimeoutError, GpuPoolBusyError):
|
||||
# Watermarking is provenance best-effort: a typed execution overrun or
|
||||
# queue saturation must not discard synthesis that already completed.
|
||||
logger.warning("Watermark skipped after its bounded dispatch expired")
|
||||
return waveform
|
||||
except asyncio.CancelledError:
|
||||
# A queued future is cancelled during pool teardown. Caller-driven
|
||||
# cancellation while the pool is live must retain normal semantics.
|
||||
if not pool.is_shutdown():
|
||||
raise
|
||||
logger.warning("Watermark skipped while the pool is shutting down")
|
||||
return waveform
|
||||
except RuntimeError:
|
||||
# Shutdown may begin after admission but before Executor.submit().
|
||||
# Preserve unrelated worker failures; only lifecycle rejection is
|
||||
# fail-open because finished synthesis must not be lost to teardown.
|
||||
if not pool.is_shutdown():
|
||||
raise
|
||||
logger.warning("Watermark skipped while the pool is shutting down")
|
||||
return waveform
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def embed_watermark(
|
||||
waveform: torch.Tensor,
|
||||
@@ -243,13 +466,14 @@ def embed_watermark(
|
||||
|
||||
# AudioSeal operates at 16kHz internally; it handles resampling, but
|
||||
# we need to inform it of the source rate for correct embedding.
|
||||
watermarked = torch.cat(
|
||||
[
|
||||
generator(seg, sample_rate=sample_rate, message=msg)
|
||||
for seg in _iter_chunks(audio, sample_rate)
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
with _eager_audioseal():
|
||||
watermarked = torch.cat(
|
||||
[
|
||||
generator(seg, sample_rate=sample_rate, message=msg)
|
||||
for seg in _iter_chunks(audio, sample_rate)
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
|
||||
# Restore original shape
|
||||
if len(original_shape) == 2:
|
||||
@@ -260,7 +484,7 @@ def embed_watermark(
|
||||
return watermarked
|
||||
|
||||
except Exception as e:
|
||||
logger.warning("Watermark embedding failed (passing through original): %s", e)
|
||||
logger.warning("Watermark embedding failed (passing through original): %s", e, exc_info=True)
|
||||
return waveform
|
||||
|
||||
|
||||
@@ -307,12 +531,13 @@ def detect_watermark(
|
||||
# embedding does, and a splice where only part of the file is
|
||||
# VoiceStudio audio still registers (a whole-file average would dilute it).
|
||||
best_conf, decoded_msg = -1.0, None
|
||||
for seg in _iter_chunks(audio, sample_rate):
|
||||
result = detector.detect_watermark(seg, sample_rate=sample_rate, message_threshold=0.5)
|
||||
seg_conf = float(result[0]) if isinstance(result, tuple) else 0.0
|
||||
if seg_conf > best_conf:
|
||||
best_conf = seg_conf
|
||||
decoded_msg = result[1] if isinstance(result, tuple) and len(result) > 1 else None
|
||||
with _eager_audioseal():
|
||||
for seg in _iter_chunks(audio, sample_rate):
|
||||
result = detector.detect_watermark(seg, sample_rate=sample_rate, message_threshold=0.5)
|
||||
seg_conf = float(result[0]) if isinstance(result, tuple) else 0.0
|
||||
if seg_conf > best_conf:
|
||||
best_conf = seg_conf
|
||||
decoded_msg = result[1] if isinstance(result, tuple) and len(result) > 1 else None
|
||||
confidence = max(best_conf, 0.0)
|
||||
|
||||
# Decode message bits
|
||||
@@ -337,7 +562,7 @@ def detect_watermark(
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.warning("Watermark detection failed: %s", e)
|
||||
logger.warning("Watermark detection failed: %s", e, exc_info=True)
|
||||
return {
|
||||
"is_watermarked": False,
|
||||
"confidence": 0.0,
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Dependency-free client for VoiceStudio's local speech platform."""
|
||||
@@ -0,0 +1,278 @@
|
||||
"""CLI/module bridge for terminals, editor extensions, and agent hooks.
|
||||
|
||||
The desktop app must be running for native dictation control. Batch
|
||||
transcription can also target a standalone or remote VoiceStudio backend.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import ipaddress
|
||||
import json
|
||||
import mimetypes
|
||||
import os
|
||||
from pathlib import Path
|
||||
import secrets
|
||||
import sys
|
||||
from typing import Any
|
||||
from urllib import error, request
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
DEFAULT_CONTROL_URL = "http://127.0.0.1:3902"
|
||||
DEFAULT_ENGINE_URL = "http://127.0.0.1:3900"
|
||||
|
||||
|
||||
class SpeechClientError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class _RejectCredentialRedirect(request.HTTPRedirectHandler):
|
||||
def redirect_request(self, req, fp, code, msg, headers, newurl): # noqa: ARG002
|
||||
raise SpeechClientError("VoiceStudio refused a credentialed redirect")
|
||||
|
||||
|
||||
def _join_url(base_url: str, path: str) -> str:
|
||||
return f"{base_url.rstrip('/')}/{path.lstrip('/')}"
|
||||
|
||||
|
||||
def _decode_error(exc: error.HTTPError) -> str:
|
||||
try:
|
||||
body = exc.read().decode("utf-8", errors="replace")
|
||||
except Exception:
|
||||
body = ""
|
||||
try:
|
||||
detail = json.loads(body)
|
||||
except (TypeError, json.JSONDecodeError):
|
||||
detail = body.strip()
|
||||
return f"HTTP {exc.code}: {detail or exc.reason}"
|
||||
|
||||
|
||||
def _is_loopback_host(host: str | None) -> bool:
|
||||
if not host:
|
||||
return False
|
||||
if host.lower() == "localhost":
|
||||
return True
|
||||
try:
|
||||
return ipaddress.ip_address(host).is_loopback
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def _open(req: request.Request, timeout: float = 300.0) -> tuple[bytes, str]:
|
||||
target = urlsplit(req.full_url)
|
||||
scheme = target.scheme.lower()
|
||||
if scheme not in {"http", "https"}:
|
||||
raise SpeechClientError("VoiceStudio URLs must use http:// or https://")
|
||||
credentialed = bool(req.get_header("Authorization"))
|
||||
if credentialed and scheme != "https" and not _is_loopback_host(target.hostname):
|
||||
raise SpeechClientError("Remote VoiceStudio credentials require https://")
|
||||
try:
|
||||
opener = (
|
||||
request.build_opener(_RejectCredentialRedirect())
|
||||
if credentialed
|
||||
else request.build_opener()
|
||||
)
|
||||
with opener.open(req, timeout=timeout) as response: # noqa: S310
|
||||
return response.read(), response.headers.get("Content-Type", "")
|
||||
except error.HTTPError as exc:
|
||||
raise SpeechClientError(_decode_error(exc)) from exc
|
||||
except error.URLError as exc:
|
||||
raise SpeechClientError(f"VoiceStudio is unavailable: {exc.reason}") from exc
|
||||
|
||||
|
||||
def _json_request(method: str, url: str, payload: Any | None = None) -> Any:
|
||||
data = None if payload is None else json.dumps(payload).encode("utf-8")
|
||||
headers = {"Accept": "application/json"}
|
||||
if data is not None:
|
||||
headers["Content-Type"] = "application/json"
|
||||
body, _ = _open(request.Request(url, data=data, headers=headers, method=method), timeout=10.0)
|
||||
try:
|
||||
return json.loads(body)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise SpeechClientError("VoiceStudio returned invalid JSON") from exc
|
||||
|
||||
|
||||
def _encode_multipart(
|
||||
*,
|
||||
filename: str,
|
||||
audio: bytes,
|
||||
fields: dict[str, str],
|
||||
boundary: str | None = None,
|
||||
) -> tuple[bytes, str]:
|
||||
boundary = boundary or f"voicestudio-{secrets.token_hex(16)}"
|
||||
marker = boundary.encode("ascii")
|
||||
parts: list[bytes] = []
|
||||
for name, value in fields.items():
|
||||
parts.extend(
|
||||
[
|
||||
b"--" + marker + b"\r\n",
|
||||
f'Content-Disposition: form-data; name="{name}"\r\n\r\n'.encode(),
|
||||
value.encode("utf-8"),
|
||||
b"\r\n",
|
||||
]
|
||||
)
|
||||
safe_filename = Path(filename).name.replace('"', "") or "audio.wav"
|
||||
content_type = mimetypes.guess_type(safe_filename)[0] or "application/octet-stream"
|
||||
if Path(safe_filename).suffix.lower() in {".wav", ".wave"}:
|
||||
content_type = "audio/wav"
|
||||
parts.extend(
|
||||
[
|
||||
b"--" + marker + b"\r\n",
|
||||
(
|
||||
'Content-Disposition: form-data; name="file"; '
|
||||
f'filename="{safe_filename}"\r\n'
|
||||
).encode(),
|
||||
f"Content-Type: {content_type}\r\n\r\n".encode(),
|
||||
audio,
|
||||
b"\r\n--" + marker + b"--\r\n",
|
||||
]
|
||||
)
|
||||
return b"".join(parts), f"multipart/form-data; boundary={boundary}"
|
||||
|
||||
|
||||
def _control(args: argparse.Namespace, action: str) -> int:
|
||||
method = "GET" if action in {"status", "capabilities"} else "POST"
|
||||
path = {
|
||||
"status": "/v1/status",
|
||||
"capabilities": "/v1/capabilities",
|
||||
"start": "/v1/dictation/start",
|
||||
"stop": "/v1/dictation/stop",
|
||||
"toggle": "/v1/dictation/toggle",
|
||||
}[action]
|
||||
result = _json_request(method, _join_url(args.control_url, path))
|
||||
print(json.dumps(result, ensure_ascii=False, indent=2))
|
||||
return 0
|
||||
|
||||
|
||||
def _read_audio(path: str, stdin_filename: str) -> tuple[bytes, str]:
|
||||
if path == "-":
|
||||
return sys.stdin.buffer.read(), stdin_filename
|
||||
audio_path = Path(path)
|
||||
try:
|
||||
return audio_path.read_bytes(), audio_path.name
|
||||
except OSError as exc:
|
||||
display_name = path.replace("\\", "/").rsplit("/", 1)[-1] or "audio input"
|
||||
reason = exc.strerror or type(exc).__name__
|
||||
raise SpeechClientError(f"could not read '{display_name}': {reason}") from exc
|
||||
|
||||
|
||||
def _response_text(body: bytes, content_type: str) -> str:
|
||||
decoded = body.decode("utf-8", errors="replace")
|
||||
if "json" not in content_type.lower():
|
||||
return decoded
|
||||
try:
|
||||
payload = json.loads(decoded)
|
||||
except json.JSONDecodeError:
|
||||
return decoded
|
||||
if isinstance(payload, dict) and isinstance(payload.get("text"), str):
|
||||
return payload["text"]
|
||||
return decoded
|
||||
|
||||
|
||||
def _transcribe(args: argparse.Namespace) -> int:
|
||||
audio, filename = _read_audio(args.audio, args.stdin_filename)
|
||||
fields = {
|
||||
"model": args.model,
|
||||
"response_format": args.response_format,
|
||||
}
|
||||
if args.language:
|
||||
fields["language"] = args.language
|
||||
body, content_type = _encode_multipart(filename=filename, audio=audio, fields=fields)
|
||||
headers = {"Content-Type": content_type, "Accept": "application/json, text/plain"}
|
||||
api_key = os.environ.get("OMNIVOICE_API_KEY", "").strip()
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
|
||||
output_session_id = None
|
||||
if args.insert:
|
||||
session = _json_request(
|
||||
"POST", _join_url(args.control_url, "/v1/output/sessions")
|
||||
)
|
||||
output_session_id = session["session_id"]
|
||||
|
||||
session_needs_cleanup = output_session_id is not None
|
||||
try:
|
||||
response_body, response_type = _open(
|
||||
request.Request(
|
||||
_join_url(args.engine_url, "/v1/audio/transcriptions"),
|
||||
data=body,
|
||||
headers=headers,
|
||||
method="POST",
|
||||
)
|
||||
)
|
||||
if output_session_id is not None:
|
||||
_json_request(
|
||||
"POST",
|
||||
_join_url(
|
||||
args.control_url,
|
||||
f"/v1/output/sessions/{output_session_id}/insert",
|
||||
),
|
||||
{"text": _response_text(response_body, response_type)},
|
||||
)
|
||||
session_needs_cleanup = False
|
||||
finally:
|
||||
if session_needs_cleanup:
|
||||
try:
|
||||
_json_request(
|
||||
"DELETE",
|
||||
_join_url(args.control_url, f"/v1/output/sessions/{output_session_id}"),
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
# Best-effort cleanup must not replace the original failure or
|
||||
# KeyboardInterrupt that brought control into this finally.
|
||||
pass
|
||||
|
||||
sys.stdout.buffer.write(response_body)
|
||||
if response_body and not response_body.endswith(b"\n"):
|
||||
sys.stdout.buffer.write(b"\n")
|
||||
return 0
|
||||
|
||||
|
||||
def _parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(
|
||||
prog="voicestudio-speech",
|
||||
description="Control and consume VoiceStudio's local speech platform.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--control-url",
|
||||
default=os.environ.get("VOICESTUDIO_SPEECH_URL", DEFAULT_CONTROL_URL),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--engine-url",
|
||||
default=os.environ.get("VOICESTUDIO_URL", DEFAULT_ENGINE_URL),
|
||||
)
|
||||
subparsers = parser.add_subparsers(dest="command", required=True)
|
||||
for command in ("status", "capabilities", "start", "stop", "toggle"):
|
||||
subparsers.add_parser(command)
|
||||
|
||||
transcribe = subparsers.add_parser("transcribe")
|
||||
transcribe.add_argument("audio", help="audio file, or - for stdin")
|
||||
transcribe.add_argument("--stdin-filename", default="audio.wav")
|
||||
transcribe.add_argument("--model", default="whisper-1")
|
||||
transcribe.add_argument("--language")
|
||||
transcribe.add_argument(
|
||||
"--format",
|
||||
dest="response_format",
|
||||
choices=("json", "text", "verbose_json", "srt", "vtt"),
|
||||
default="text",
|
||||
)
|
||||
transcribe.add_argument(
|
||||
"--insert",
|
||||
action="store_true",
|
||||
help="insert the result into the app focused when this command starts",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
args = _parser().parse_args(argv)
|
||||
try:
|
||||
if args.command == "transcribe":
|
||||
return _transcribe(args)
|
||||
return _control(args, args.command)
|
||||
except (SpeechClientError, KeyError) as exc:
|
||||
print(f"voicestudio-speech: {exc}", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -80,12 +80,10 @@ def test_timeout_error_is_a_timeouterror_subclass():
|
||||
assert issubclass(ASRTimeoutError, TimeoutError)
|
||||
|
||||
|
||||
def test_timeout_resets_a_resilient_pool_to_restore_capacity():
|
||||
# #730: a wedged transcribe holds its GPU-pool worker forever; with a 1-2
|
||||
# worker pool that starves TTS generate and surfaces as "can't reach
|
||||
# backend". On timeout, run_transcribe_guarded must reset() a pool that
|
||||
# supports it (the real _ResilientGpuPool) so the next submit gets a fresh
|
||||
# worker — capacity restored without an app restart.
|
||||
def test_timeout_does_not_overlap_an_in_process_native_worker():
|
||||
# #1669: reset() cannot kill the old native thread. A fresh pool let the
|
||||
# retry enter the same whisperx/CTranslate2 model concurrently and the
|
||||
# process died with 0xC0000005. Keep the old worker accounted for instead.
|
||||
class _FakePool(ThreadPoolExecutor):
|
||||
def __init__(self):
|
||||
super().__init__(max_workers=1)
|
||||
@@ -105,7 +103,7 @@ def test_timeout_resets_a_resilient_pool_to_restore_capacity():
|
||||
await run_transcribe_guarded(pool, _hang, what="Dub", timeout=0.2)
|
||||
|
||||
asyncio.run(_go())
|
||||
assert pool.reset_calls == 1
|
||||
assert pool.reset_calls == 0
|
||||
pool.shutdown(wait=False)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,259 @@
|
||||
"""Stable nested operation ownership (model-free, cross-platform seams)."""
|
||||
import ctypes
|
||||
import builtins
|
||||
import os
|
||||
import runpy
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import types
|
||||
from ctypes import wintypes
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from core import contained_subprocess as owned
|
||||
|
||||
|
||||
class _Call:
|
||||
def __init__(self, fn):
|
||||
self.fn = fn
|
||||
|
||||
def __call__(self, *args):
|
||||
return self.fn(*args)
|
||||
|
||||
|
||||
def test_supervisor_argv_uses_entry_module_for_source_and_frozen_binary(monkeypatch):
|
||||
monkeypatch.delattr(owned.sys, "frozen", raising=False)
|
||||
source = owned._supervisor_argv(3, 4, ["operation"])
|
||||
assert source[:2] == [sys.executable, str(Path(owned.__file__).parents[1] / "main.py")]
|
||||
assert source[2:] == ["--supervise", "3", "4", "--", "operation"]
|
||||
|
||||
monkeypatch.setattr(owned.sys, "frozen", True, raising=False)
|
||||
frozen = owned._supervisor_argv(3, 4, ["operation"])
|
||||
assert frozen == [sys.executable, "--supervise", "3", "4", "--", "operation"]
|
||||
|
||||
|
||||
def test_source_main_dispatches_supervisor_before_heavy_imports(monkeypatch):
|
||||
calls = []
|
||||
fake = types.ModuleType("core.contained_subprocess")
|
||||
fake.supervisor_main = lambda args: calls.append(args) or 23
|
||||
monkeypatch.setitem(sys.modules, "core.contained_subprocess", fake)
|
||||
main_path = Path(owned.__file__).parents[1] / "main.py"
|
||||
monkeypatch.setattr(
|
||||
sys,
|
||||
"argv",
|
||||
[str(main_path), "--supervise", "3", "4", "--", "operation"],
|
||||
)
|
||||
original_import = builtins.__import__
|
||||
|
||||
def guard_heavy_import(name, *args, **kwargs):
|
||||
if name == "math":
|
||||
raise AssertionError("supervisor dispatch reached application imports")
|
||||
return original_import(name, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", guard_heavy_import)
|
||||
with pytest.raises(SystemExit, match="23"):
|
||||
runpy.run_path(str(main_path), run_name="__main__")
|
||||
assert calls == [["--supervise", "3", "4", "--", "operation"]]
|
||||
|
||||
|
||||
@pytest.mark.skipif(os.name != "posix", reason="Unix drain pipe contract")
|
||||
def test_drain_fd_is_explicitly_inherited_by_wrapper_but_not_operation(monkeypatch):
|
||||
drain_read, drain_write = os.pipe()
|
||||
monkeypatch.setenv("OMNIVOICE_DESKTOP_CONTAINED", "1")
|
||||
monkeypatch.setenv("OMNIVOICE_DESKTOP_DRAIN_FD", str(drain_write))
|
||||
owned.secure_backend_drain_fd()
|
||||
assert not os.get_inheritable(drain_write)
|
||||
implicit_probe = subprocess.check_output(
|
||||
[
|
||||
sys.executable,
|
||||
"-c",
|
||||
"import os; "
|
||||
"fd=int(os.environ['OMNIVOICE_DESKTOP_DRAIN_FD']); "
|
||||
"\ntry: os.fstat(fd); print('leaked')"
|
||||
"\nexcept OSError: print('closed')",
|
||||
],
|
||||
close_fds=False,
|
||||
text=True,
|
||||
)
|
||||
assert implicit_probe.strip() == "closed"
|
||||
script = (
|
||||
"import os,time; token=os.environ.get('OMNIVOICE_DESKTOP_DRAIN_FD'); "
|
||||
"marker=os.environ.get('OMNIVOICE_DESKTOP_CONTAINED'); "
|
||||
"\nif token is None and marker is None: state='stripped'"
|
||||
"\nelse:"
|
||||
"\n try: os.fstat(int(token)); state='leaked'"
|
||||
"\n except OSError: state='closed'"
|
||||
"\nprint(state, flush=True); time.sleep(60)"
|
||||
)
|
||||
proc = owned.spawn_owned(
|
||||
[sys.executable, "-c", script],
|
||||
stdout=subprocess.PIPE,
|
||||
text=True,
|
||||
)
|
||||
try:
|
||||
assert proc.stdout.readline().strip() == "stripped"
|
||||
os.close(drain_write)
|
||||
drain_write = -1
|
||||
os.set_blocking(drain_read, False)
|
||||
with pytest.raises(BlockingIOError):
|
||||
os.read(drain_read, 1) # wrapper still holds the only writer
|
||||
proc.kill()
|
||||
proc.wait(timeout=5)
|
||||
deadline = time.monotonic() + 2
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
if os.read(drain_read, 1) == b"":
|
||||
break
|
||||
except BlockingIOError:
|
||||
time.sleep(0.01)
|
||||
else:
|
||||
pytest.fail("wrapper exit did not close the desktop drain writer")
|
||||
finally:
|
||||
if drain_write >= 0:
|
||||
os.close(drain_write)
|
||||
os.close(drain_read)
|
||||
if proc.poll() is None:
|
||||
proc.kill()
|
||||
proc.wait(timeout=5)
|
||||
|
||||
|
||||
def test_invalid_or_missing_desktop_drain_fd_fails_safe(monkeypatch):
|
||||
monkeypatch.setenv("OMNIVOICE_DESKTOP_CONTAINED", "1")
|
||||
monkeypatch.setenv("OMNIVOICE_DESKTOP_DRAIN_FD", "not-an-fd")
|
||||
with pytest.raises(RuntimeError, match="missing its live.*drain descriptor"):
|
||||
owned.spawn_owned([sys.executable, "-c", "print('unsafe')"])
|
||||
|
||||
monkeypatch.delenv("OMNIVOICE_DESKTOP_DRAIN_FD")
|
||||
with pytest.raises(RuntimeError, match="missing its live.*drain descriptor"):
|
||||
owned.secure_backend_drain_fd()
|
||||
|
||||
monkeypatch.delenv("OMNIVOICE_DESKTOP_CONTAINED")
|
||||
assert owned.backend_drain_fd(required=True) is None
|
||||
proc = owned.spawn_owned(
|
||||
[sys.executable, "-c", "print('standalone')"],
|
||||
stdout=subprocess.PIPE,
|
||||
text=True,
|
||||
)
|
||||
assert proc.stdout.readline().strip() == "standalone"
|
||||
assert proc.wait(timeout=5) == 0
|
||||
|
||||
|
||||
def test_windows_operation_is_in_kill_on_close_job_before_resume(monkeypatch):
|
||||
"""The child gets no instruction before stable nested Job assignment."""
|
||||
events = []
|
||||
job_closed = threading.Event()
|
||||
job = 99
|
||||
|
||||
def close_handle(handle):
|
||||
value = getattr(handle, "value", handle)
|
||||
events.append(("close", value))
|
||||
if value == job:
|
||||
job_closed.set()
|
||||
return True
|
||||
|
||||
kernel = type("Kernel", (), {})()
|
||||
kernel.AssignProcessToJobObject = _Call(
|
||||
lambda assigned_job, process: events.append(("assign", assigned_job, process)) or True
|
||||
)
|
||||
kernel.TerminateJobObject = _Call(
|
||||
lambda assigned_job, code: events.append(("terminate", assigned_job, code)) or True
|
||||
)
|
||||
kernel.WriteFile = _Call(
|
||||
lambda handle, payload, size, written, overlap: events.append(("write", size)) or True
|
||||
)
|
||||
kernel.CloseHandle = _Call(close_handle)
|
||||
|
||||
def read_control(*_args):
|
||||
job_closed.wait(2)
|
||||
return False
|
||||
|
||||
kernel.ReadFile = _Call(read_control)
|
||||
monkeypatch.setattr(owned, "_windows_job", lambda: (job, kernel, wintypes))
|
||||
monkeypatch.setattr(
|
||||
owned,
|
||||
"_resume_windows_process",
|
||||
lambda _kernel, _types, pid: events.append(("resume", pid)),
|
||||
)
|
||||
|
||||
class Child:
|
||||
_handle = 77
|
||||
pid = 123
|
||||
|
||||
def wait(self, timeout=None):
|
||||
events.append(("wait", timeout))
|
||||
return 0
|
||||
|
||||
monkeypatch.setattr(
|
||||
owned.subprocess,
|
||||
"Popen",
|
||||
lambda *args, **kwargs: events.append(("spawn", kwargs["creationflags"])) or Child(),
|
||||
)
|
||||
|
||||
assert owned._supervise_windows(11, 12, ["operation.exe"]) == 0
|
||||
assert job_closed.wait(1)
|
||||
|
||||
names = [event[0] for event in events]
|
||||
assert names.index("assign") < names.index("resume") < names.index("wait")
|
||||
assert names.index("wait") < names.index("terminate") < names.index("write")
|
||||
|
||||
|
||||
def test_windows_assignment_failure_kills_suspended_unowned_child(monkeypatch):
|
||||
"""A child outside the nested Job must be killed through its stable handle."""
|
||||
events = []
|
||||
job_closed = threading.Event()
|
||||
job = 99
|
||||
|
||||
def close_handle(handle):
|
||||
value = getattr(handle, "value", handle)
|
||||
events.append(("close", value))
|
||||
if value == job:
|
||||
job_closed.set()
|
||||
return True
|
||||
|
||||
kernel = type("Kernel", (), {})()
|
||||
kernel.AssignProcessToJobObject = _Call(
|
||||
lambda assigned_job, process: events.append(("assign", assigned_job, process))
|
||||
or False
|
||||
)
|
||||
kernel.TerminateJobObject = _Call(
|
||||
lambda assigned_job, code: events.append(("terminate", assigned_job, code)) or True
|
||||
)
|
||||
kernel.WriteFile = _Call(
|
||||
lambda handle, payload, size, written, overlap: events.append(("write", size)) or True
|
||||
)
|
||||
kernel.CloseHandle = _Call(close_handle)
|
||||
|
||||
def read_control(*_args):
|
||||
job_closed.wait(2)
|
||||
return False
|
||||
|
||||
kernel.ReadFile = _Call(read_control)
|
||||
monkeypatch.setattr(owned, "_windows_job", lambda: (job, kernel, wintypes))
|
||||
monkeypatch.setattr(ctypes, "get_last_error", lambda: 5, raising=False)
|
||||
|
||||
class Child:
|
||||
_handle = 77
|
||||
pid = 123
|
||||
|
||||
def kill(self):
|
||||
events.append(("kill",))
|
||||
|
||||
def wait(self, timeout=None):
|
||||
events.append(("wait", timeout))
|
||||
return 1
|
||||
|
||||
monkeypatch.setattr(
|
||||
owned.subprocess,
|
||||
"Popen",
|
||||
lambda *args, **kwargs: events.append(("spawn", kwargs["creationflags"])) or Child(),
|
||||
)
|
||||
|
||||
assert owned._supervise_windows(11, 12, ["operation.exe"]) == 127
|
||||
assert job_closed.wait(1)
|
||||
|
||||
names = [event[0] for event in events]
|
||||
assert names.index("assign") < names.index("terminate") < names.index("kill")
|
||||
assert names.index("kill") < names.index("wait") < names.index("write")
|
||||
@@ -0,0 +1,126 @@
|
||||
"""macOS fallback for the os.waitid probe (#1656).
|
||||
|
||||
CPython on macOS does not expose os.waitid, so OwnedPopen's WNOWAIT dance
|
||||
crashed with AttributeError on every poll after the first spawn. These tests
|
||||
simulate that platform (monkeypatch os.waitid away) and pin the fallback:
|
||||
poll/wait/kill must work, exit codes must be real, and an already-reaped
|
||||
leader must be refused (ChildProcessError path), never signalled blind.
|
||||
"""
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from core import contained_subprocess as owned
|
||||
|
||||
|
||||
def _make_owned(argv):
|
||||
cr, cw = os.pipe()
|
||||
rr, rw = os.pipe()
|
||||
proc = subprocess.Popen(argv, start_new_session=True)
|
||||
os.close(cw)
|
||||
os.close(rw) # result writer gone: _read_result falls back to wrapper rc
|
||||
return owned.OwnedPopen(proc, cr, rr), proc
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def no_waitid(monkeypatch):
|
||||
monkeypatch.delattr(os, "waitid", raising=False)
|
||||
|
||||
|
||||
def test_poll_running_then_exited_without_waitid(no_waitid):
|
||||
h, _ = _make_owned([sys.executable, "-c", "import time; time.sleep(1.5)"])
|
||||
try:
|
||||
assert h.poll() is None, "running child must poll None"
|
||||
h._proc.wait()
|
||||
deadline = time.monotonic() + 5
|
||||
rc = None
|
||||
while rc is None and time.monotonic() < deadline:
|
||||
rc = h.poll()
|
||||
time.sleep(0.05)
|
||||
assert rc == 0
|
||||
assert h.poll() == 0
|
||||
finally:
|
||||
h._close_control()
|
||||
if h._result_fd is not None:
|
||||
os.close(h._result_fd)
|
||||
|
||||
|
||||
def test_poll_reports_real_exit_code_without_waitid(no_waitid):
|
||||
h, _ = _make_owned([sys.executable, "-c", "raise SystemExit(3)"])
|
||||
try:
|
||||
deadline = time.monotonic() + 5
|
||||
while h.poll() is None and time.monotonic() < deadline:
|
||||
time.sleep(0.05)
|
||||
assert h.poll() == 3
|
||||
finally:
|
||||
h._close_control()
|
||||
if h._result_fd is not None:
|
||||
os.close(h._result_fd)
|
||||
|
||||
|
||||
def test_wait_returns_after_kill_without_waitid(no_waitid):
|
||||
h, _ = _make_owned([sys.executable, "-c", "import time; time.sleep(30)"])
|
||||
try:
|
||||
h.kill()
|
||||
rc = h.wait(timeout=5)
|
||||
assert rc != 0
|
||||
finally:
|
||||
h._close_control()
|
||||
if h._result_fd is not None:
|
||||
os.close(h._result_fd)
|
||||
|
||||
|
||||
def test_reaped_by_own_popen_reports_code_without_waitid(no_waitid):
|
||||
h, proc = _make_owned([sys.executable, "-c", "pass"])
|
||||
try:
|
||||
proc.wait() # reaped through OUR handle: known code, not a refusal
|
||||
assert h.poll() == 0
|
||||
finally:
|
||||
h._close_control()
|
||||
if h._result_fd is not None:
|
||||
os.close(h._result_fd)
|
||||
|
||||
|
||||
def test_foreign_reaped_leader_is_refused_without_waitid(no_waitid):
|
||||
h, proc = _make_owned([sys.executable, "-c", "pass"])
|
||||
try:
|
||||
# Reap OUTSIDE this handle: Popen never learns the code, so poll must
|
||||
# refuse (None) rather than guess or signal a maybe-reused group.
|
||||
while True:
|
||||
pid, _ = os.waitpid(proc.pid, os.WNOHANG)
|
||||
if pid == proc.pid:
|
||||
break
|
||||
time.sleep(0.05)
|
||||
assert h.poll() is None
|
||||
finally:
|
||||
h._close_control()
|
||||
if h._result_fd is not None:
|
||||
os.close(h._result_fd)
|
||||
|
||||
|
||||
def test_kill_after_pid_reuse_does_not_signal_without_waitid(no_waitid, monkeypatch):
|
||||
"""A foreign-reaped leader's reused numeric pid must not authorize killpg."""
|
||||
import signal as _signal
|
||||
|
||||
h, proc = _make_owned([sys.executable, "-c", "pass"])
|
||||
try:
|
||||
while True:
|
||||
pid, _ = os.waitpid(proc.pid, os.WNOHANG)
|
||||
if pid == proc.pid:
|
||||
break
|
||||
time.sleep(0.05)
|
||||
# Model the numeric pid being reused: kill(pid, 0) would succeed even
|
||||
# though waitpid still reports that the original child is no longer
|
||||
# ours. The old guard therefore reached killpg and fails this test.
|
||||
monkeypatch.setattr(os, "kill", lambda _pid, _sig: None)
|
||||
signalled = []
|
||||
monkeypatch.setattr(os, "killpg", lambda pid, sig: signalled.append((pid, sig)))
|
||||
h._signal_owned_group(_signal.SIGKILL)
|
||||
assert signalled == []
|
||||
finally:
|
||||
h._close_control()
|
||||
if h._result_fd is not None:
|
||||
os.close(h._result_fd)
|
||||
@@ -1,16 +1,15 @@
|
||||
"""A dictation model that decodes nothing gets demoted, not re-selected forever.
|
||||
|
||||
`sherpa-parakeet-tdt-v3` is the curated default, and on Windows it installs
|
||||
cleanly, loads without error, and returns an empty token list for clear speech
|
||||
On Windows, `sherpa-parakeet-tdt-v3` installs cleanly, loads without error,
|
||||
and returns an empty token list for clear speech
|
||||
(both quantisations, both decoding methods, sherpa-onnx 1.13.3 and 1.13.4)
|
||||
while whisper and zipformer transcribe the same bytes. The defect is inside
|
||||
sherpa-onnx's NeMo-TDT decoder — unfixable from here by configuration.
|
||||
|
||||
Hard-coding a different default per OS would be a guess: we have evidence for
|
||||
one platform only. So the app observes instead. When a session hears real
|
||||
speech and the model returns nothing, that model is demoted ON THIS MACHINE and
|
||||
stops being auto-selected, which self-corrects wherever the breakage actually
|
||||
is and is a no-op everywhere it isn't.
|
||||
Whisper Tiny is now the cross-platform default, while Parakeet remains
|
||||
selectable. Runtime demotion still protects users who select a recognizer that
|
||||
loads successfully but decodes nothing: it is demoted on this machine and the
|
||||
next session follows the capture fallback.
|
||||
|
||||
These tests pin the demotion round trip and, critically, that the user can
|
||||
always take back control by re-picking the model.
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
"""A dictation model that decodes NOTHING must fall back, not fail silently.
|
||||
|
||||
Found on Windows with the curated default `sherpa-parakeet-tdt-v3`: the model
|
||||
Found on Windows with `sherpa-parakeet-tdt-v3`: the model
|
||||
downloads, loads with zero errors, and is correctly detected as a TDT model
|
||||
(`num_durations: 5`) — then returns an empty token list for clear speech.
|
||||
Measured against the same 18.9s WAV, on the same machine, same sherpa-onnx:
|
||||
|
||||
sherpa-whisper-tiny -> "Alright, here we are. I hope that's all..."
|
||||
sherpa-zipformer-en-20m -> "ANTS BOTH IN WHAT DISGUISED THIS THAT..."
|
||||
parakeet-tdt-v3 (int8) -> '' <-- the curated default
|
||||
parakeet-tdt-v3 (int8) -> ''
|
||||
parakeet-tdt-v3 (fp32) -> ''
|
||||
parakeet-tdt-v2 (int8) -> ''
|
||||
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import io
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
from fastapi import UploadFile
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_preview_ffmpeg_does_not_block_event_loop(monkeypatch, tmp_path):
|
||||
from api.routers import dub_core
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
started = asyncio.Event()
|
||||
release = threading.Event()
|
||||
|
||||
def slow_ffmpeg(*_args, **_kwargs):
|
||||
loop.call_soon_threadsafe(started.set)
|
||||
assert release.wait(timeout=2)
|
||||
|
||||
monkeypatch.setattr(dub_core, "PREVIEW_DIR", str(tmp_path))
|
||||
monkeypatch.setattr(dub_core, "find_ffmpeg", lambda: "ffmpeg")
|
||||
monkeypatch.setattr(dub_core.subprocess, "run", slow_ffmpeg)
|
||||
upload = UploadFile(filename="preview.mp4", file=io.BytesIO(b"video"))
|
||||
|
||||
before = loop.time()
|
||||
task = asyncio.create_task(dub_core.preview_upload(upload))
|
||||
try:
|
||||
await asyncio.wait_for(started.wait(), timeout=1)
|
||||
assert loop.time() - before < 0.5
|
||||
finally:
|
||||
release.set()
|
||||
|
||||
result = await task
|
||||
assert result["audioUrl"].endswith(".wav")
|
||||
@@ -0,0 +1,71 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_abandoned_reader_keeps_adhoc_reference_until_worker_finishes(tmp_path):
|
||||
from api.routers.generation import (
|
||||
_TempReferenceLease,
|
||||
_run_with_reference_lease,
|
||||
)
|
||||
from services.model_manager import run_on_gpu_pool_guarded
|
||||
|
||||
reference = tmp_path / "reference.wav"
|
||||
reference.write_bytes(b"voice")
|
||||
lease = _TempReferenceLease(str(reference))
|
||||
started = threading.Event()
|
||||
release_worker = threading.Event()
|
||||
worker_read = threading.Event()
|
||||
|
||||
def read_reference():
|
||||
started.set()
|
||||
assert release_worker.wait(timeout=2)
|
||||
assert reference.read_bytes() == b"voice"
|
||||
worker_read.set()
|
||||
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
task = asyncio.create_task(
|
||||
_run_with_reference_lease(
|
||||
lease,
|
||||
lambda on_abandon: run_on_gpu_pool_guarded(
|
||||
read_reference,
|
||||
executor=executor,
|
||||
timeout=1,
|
||||
on_abandon=on_abandon,
|
||||
),
|
||||
)
|
||||
)
|
||||
assert await asyncio.to_thread(started.wait, 1)
|
||||
task.cancel()
|
||||
cancelled = await asyncio.gather(task, return_exceptions=True)
|
||||
assert isinstance(cancelled[0], asyncio.CancelledError)
|
||||
|
||||
lease.finish_request()
|
||||
assert reference.exists()
|
||||
release_worker.set()
|
||||
assert await asyncio.to_thread(worker_read.wait, 1)
|
||||
|
||||
for _ in range(100):
|
||||
if not reference.exists():
|
||||
break
|
||||
await asyncio.sleep(0.01)
|
||||
assert not reference.exists()
|
||||
|
||||
|
||||
def test_normal_request_deletes_adhoc_reference_immediately(tmp_path):
|
||||
from api.routers.generation import _TempReferenceLease
|
||||
|
||||
reference = tmp_path / "reference.wav"
|
||||
reference.write_bytes(b"voice")
|
||||
lease = _TempReferenceLease(str(reference))
|
||||
|
||||
release = lease.acquire()
|
||||
release()
|
||||
lease.finish_request()
|
||||
|
||||
assert not reference.exists()
|
||||
@@ -0,0 +1,70 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import threading
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_abandon_callback_waits_for_running_worker_to_finish():
|
||||
from services.model_manager import run_on_gpu_pool_guarded
|
||||
|
||||
started = threading.Event()
|
||||
release = threading.Event()
|
||||
cleaned = threading.Event()
|
||||
|
||||
def job():
|
||||
started.set()
|
||||
assert release.wait(timeout=2)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
task = asyncio.create_task(
|
||||
run_on_gpu_pool_guarded(
|
||||
job,
|
||||
executor=executor,
|
||||
timeout=1,
|
||||
on_abandon=cleaned.set,
|
||||
)
|
||||
)
|
||||
assert await asyncio.to_thread(started.wait, 1)
|
||||
task.cancel()
|
||||
cancelled = await asyncio.gather(task, return_exceptions=True)
|
||||
assert isinstance(cancelled[0], asyncio.CancelledError)
|
||||
|
||||
assert not cleaned.is_set()
|
||||
release.set()
|
||||
assert await asyncio.to_thread(cleaned.wait, 1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_queued_cancellation_releases_without_running_job():
|
||||
from services.model_manager import GpuPoolBusyError, run_on_gpu_pool_guarded
|
||||
|
||||
hog_started = threading.Event()
|
||||
release_hog = threading.Event()
|
||||
cleaned = threading.Event()
|
||||
queued_job_ran = threading.Event()
|
||||
|
||||
def hog():
|
||||
hog_started.set()
|
||||
assert release_hog.wait(timeout=2)
|
||||
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
hog_future = executor.submit(hog)
|
||||
assert hog_started.wait(timeout=1)
|
||||
try:
|
||||
with pytest.raises(GpuPoolBusyError):
|
||||
await run_on_gpu_pool_guarded(
|
||||
queued_job_ran.set,
|
||||
executor=executor,
|
||||
timeout=1,
|
||||
queue_timeout=0.05,
|
||||
on_abandon=cleaned.set,
|
||||
)
|
||||
assert cleaned.is_set()
|
||||
assert not queued_job_ran.is_set()
|
||||
finally:
|
||||
release_hog.set()
|
||||
hog_future.result(timeout=1)
|
||||
@@ -17,18 +17,31 @@ import json
|
||||
import math
|
||||
import array
|
||||
import base64
|
||||
import io
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from services.subprocess_backend import SubprocessBackend, RECV_TIMEOUT_S
|
||||
from services.tts_backend import get_backend_class
|
||||
from engines.omnivoice_subprocess import OmniVoiceSubprocessBackend
|
||||
from services.subprocess_backend import (
|
||||
RECV_TIMEOUT_S,
|
||||
SubprocessBackend,
|
||||
)
|
||||
from services.tts_backend import OmniVoiceBackend, get_backend_class, list_backends
|
||||
from engines.omnivoice_subprocess import (
|
||||
OmniVoiceMPSSubprocessBackend,
|
||||
OmniVoiceSubprocessBackend,
|
||||
)
|
||||
|
||||
|
||||
# ── stub sidecar (model-free) ──────────────────────────────────────────────
|
||||
|
||||
STUB_SIDECAR = r'''
|
||||
import sys, json, struct, time, math, array, base64
|
||||
import sys, os, json, struct, time, math, array, base64, subprocess
|
||||
|
||||
def _send(o):
|
||||
b = json.dumps(o, separators=(",", ":")).encode()
|
||||
@@ -60,9 +73,20 @@ while True:
|
||||
sys.exit(0)
|
||||
elif op == "synthesize":
|
||||
t = m.get("text", "")
|
||||
if t == "CRASH":
|
||||
os._exit(137)
|
||||
if t == "HANG":
|
||||
while True: # wedge forever; the parent must hard-kill us
|
||||
time.sleep(1)
|
||||
if t == "HANG_CHILD":
|
||||
subprocess.Popen([
|
||||
sys.executable,
|
||||
"-c",
|
||||
"import os,time; time.sleep(1); "
|
||||
"open(os.environ['OMNIVOICE_TIMEOUT_MARKER'], 'w').write('bad')",
|
||||
])
|
||||
while True:
|
||||
time.sleep(1)
|
||||
# Emit progress frames before the audio when asked, to exercise the
|
||||
# parent's progress-consuming recv loop (the cold-load fix).
|
||||
if t.startswith("PROG:"):
|
||||
@@ -98,6 +122,80 @@ def test_registry_resolves_to_subprocess_backend():
|
||||
assert get_backend_class("omnivoice-subprocess") is OmniVoiceSubprocessBackend
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("family", "expected_name"),
|
||||
[("mps", "OmniVoiceMPSSubprocessBackend"), ("cuda", "OmniVoiceBackend"),
|
||||
("cpu", "OmniVoiceBackend")],
|
||||
)
|
||||
def test_omnivoice_is_crash_isolated_only_on_mps(monkeypatch, family, expected_name):
|
||||
from core.device_caps import HostCaps
|
||||
|
||||
available = (family, "cpu") if family != "cpu" else ("cpu",)
|
||||
monkeypatch.setattr(
|
||||
"core.device_caps.detect_host_caps",
|
||||
lambda: HostCaps(family=family, available_families=available),
|
||||
)
|
||||
|
||||
resolved = get_backend_class("omnivoice")
|
||||
assert resolved.__name__ == expected_name
|
||||
if family != "mps":
|
||||
assert resolved is OmniVoiceBackend
|
||||
|
||||
|
||||
def test_engine_catalogue_reports_effective_mps_isolation(monkeypatch):
|
||||
from core.device_caps import HostCaps
|
||||
from services import tts_backend
|
||||
|
||||
monkeypatch.setattr(tts_backend, "_REGISTRY", {"omnivoice": OmniVoiceBackend})
|
||||
monkeypatch.setattr(
|
||||
"core.device_caps.detect_host_caps",
|
||||
lambda: HostCaps(family="mps", available_families=("mps", "cpu")),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"engines.omnivoice_subprocess.OmniVoiceSubprocessBackend.is_available",
|
||||
classmethod(lambda cls: (True, "ready")),
|
||||
)
|
||||
|
||||
row = next(item for item in list_backends() if item["id"] == "omnivoice")
|
||||
assert row["isolation_mode"] == "subprocess"
|
||||
|
||||
|
||||
def test_mps_startup_does_not_preload_native_model(monkeypatch):
|
||||
from core.device_caps import HostCaps
|
||||
from services import model_manager
|
||||
|
||||
monkeypatch.setattr(
|
||||
"core.device_caps.detect_host_caps",
|
||||
lambda: HostCaps(family="mps", available_families=("mps", "cpu")),
|
||||
)
|
||||
monkeypatch.setenv("OMNIVOICE_TTS_BACKEND", "omnivoice")
|
||||
monkeypatch.setattr(model_manager, "model", None)
|
||||
|
||||
async def fail_load():
|
||||
raise AssertionError("native OmniVoice must not load in the API process on MPS")
|
||||
|
||||
monkeypatch.setattr(model_manager, "_load_model_with_timeout", fail_load)
|
||||
asyncio.run(model_manager.preload_model())
|
||||
|
||||
|
||||
def test_streaming_mps_path_does_not_load_native_model(monkeypatch):
|
||||
from api.routers.tts_stream import _resolve_stream_backend
|
||||
from services import model_manager, tts_backend
|
||||
|
||||
sentinel = object()
|
||||
monkeypatch.setattr(tts_backend, "active_backend_id", lambda: "omnivoice")
|
||||
monkeypatch.setattr(
|
||||
tts_backend, "get_backend_class", lambda _id: OmniVoiceMPSSubprocessBackend,
|
||||
)
|
||||
monkeypatch.setattr(tts_backend, "get_active_tts_backend", lambda: sentinel)
|
||||
|
||||
async def fail_load():
|
||||
raise AssertionError("streaming must not load native OmniVoice on MPS")
|
||||
|
||||
monkeypatch.setattr(model_manager, "get_model", fail_load)
|
||||
assert asyncio.run(_resolve_stream_backend(None)) is sentinel
|
||||
|
||||
|
||||
def test_is_marked_subprocess_isolated():
|
||||
# list_backends() detects isolation via this duck-typed marker, not issubclass.
|
||||
assert getattr(OmniVoiceSubprocessBackend, "_is_subprocess_isolated", False) is True
|
||||
@@ -136,6 +234,40 @@ def test_base_default_recv_timeout_is_60s():
|
||||
assert _PlainBackend().recv_timeout_s == 60.0
|
||||
|
||||
|
||||
def test_sidecar_spawn_delegates_all_containment_to_nested_owner(monkeypatch, tmp_path):
|
||||
from services import subprocess_backend as backend_module
|
||||
|
||||
captured = {}
|
||||
|
||||
class StubProcess:
|
||||
stderr = io.BytesIO()
|
||||
|
||||
@staticmethod
|
||||
def poll():
|
||||
return None
|
||||
|
||||
def fake_spawn(argv, **kwargs):
|
||||
captured.update(kwargs)
|
||||
return StubProcess()
|
||||
|
||||
monkeypatch.setattr(_PlainBackend, "venv_python", classmethod(lambda cls: Path(sys.executable)))
|
||||
monkeypatch.setattr(
|
||||
_PlainBackend,
|
||||
"sidecar_script",
|
||||
classmethod(lambda cls: tmp_path / "stub.py"),
|
||||
)
|
||||
monkeypatch.setattr(backend_module, "spawn_owned", fake_spawn)
|
||||
monkeypatch.setattr(backend_module, "_ensure_reaper_running", lambda: None)
|
||||
backend = _PlainBackend()
|
||||
monkeypatch.setattr(backend, "_recv_with_timeout", lambda _timeout: {"op": "ready"})
|
||||
|
||||
try:
|
||||
backend._spawn()
|
||||
assert not ({"start_new_session", "creationflags", "preexec_fn"} & captured.keys())
|
||||
finally:
|
||||
backend._proc = None
|
||||
|
||||
|
||||
def test_omnivoice_subprocess_recv_timeout_overrides_default():
|
||||
b = OmniVoiceSubprocessBackend()
|
||||
assert b.recv_timeout_s == 300.0 # aligns with the generate budget
|
||||
@@ -205,6 +337,48 @@ def test_wedged_sidecar_is_hard_killed_and_recovers(stub_sidecar, monkeypatch):
|
||||
b.shutdown()
|
||||
|
||||
|
||||
def test_mps_proxy_survives_fatal_child_exit_and_recovers(stub_sidecar, monkeypatch):
|
||||
_use_stub(monkeypatch, stub_sidecar)
|
||||
monkeypatch.setattr(
|
||||
"services.model_manager.make_room_before_generate", lambda: None,
|
||||
)
|
||||
b = OmniVoiceMPSSubprocessBackend()
|
||||
try:
|
||||
with pytest.raises(RuntimeError, match="backend is still running"):
|
||||
b.generate("CRASH")
|
||||
assert b._proc is not None and b._proc.poll() is not None
|
||||
assert b.generate("ok").shape[1] == 24000
|
||||
finally:
|
||||
b.shutdown()
|
||||
|
||||
|
||||
def test_desktop_timeout_kills_engine_subtree_before_late_mutation(
|
||||
stub_sidecar, monkeypatch, tmp_path
|
||||
):
|
||||
marker = tmp_path / "late-engine-mutation"
|
||||
monkeypatch.setenv("OMNIVOICE_DESKTOP_CONTAINED", "1")
|
||||
drain_read, drain_write = os.pipe()
|
||||
monkeypatch.setenv("OMNIVOICE_DESKTOP_DRAIN_FD", str(drain_write))
|
||||
monkeypatch.setenv("OMNIVOICE_TIMEOUT_MARKER", str(marker))
|
||||
_use_stub(monkeypatch, stub_sidecar)
|
||||
monkeypatch.setattr(
|
||||
OmniVoiceSubprocessBackend,
|
||||
"recv_timeout_s",
|
||||
property(lambda self: 0.3),
|
||||
)
|
||||
b = OmniVoiceSubprocessBackend()
|
||||
try:
|
||||
with pytest.raises(RuntimeError):
|
||||
b.generate("HANG_CHILD")
|
||||
time.sleep(1.2)
|
||||
assert not marker.exists()
|
||||
assert b.generate("ok").shape[1] == 24000
|
||||
finally:
|
||||
b.shutdown()
|
||||
os.close(drain_write)
|
||||
os.close(drain_read)
|
||||
|
||||
|
||||
def test_generate_does_not_deadlock_when_called_on_gpu_pool_worker(stub_sidecar, monkeypatch):
|
||||
# Regression: /v1/audio/speech and /generate dispatch backend.generate() via
|
||||
# run_on_gpu_pool_guarded, i.e. ON a gpu-pool worker. generate() must NOT
|
||||
@@ -222,3 +396,94 @@ def test_generate_does_not_deadlock_when_called_on_gpu_pool_worker(stub_sidecar,
|
||||
assert tensor.shape[1] == 24000
|
||||
finally:
|
||||
b.shutdown()
|
||||
|
||||
|
||||
def test_sidecar_forwards_native_controls_and_applies_seed(monkeypatch):
|
||||
import torch
|
||||
from engines.omnivoice_subprocess import main as sidecar
|
||||
|
||||
calls = []
|
||||
seeds = []
|
||||
frames = []
|
||||
|
||||
class FakeModel:
|
||||
sampling_rate = 24000
|
||||
|
||||
def generate(self, **kwargs):
|
||||
calls.append(kwargs)
|
||||
return [torch.zeros(1, 16)]
|
||||
|
||||
monkeypatch.setattr(sidecar, "_load_model", lambda _stdout: FakeModel())
|
||||
monkeypatch.setattr(sidecar, "_send", lambda _stdout, frame: frames.append(frame))
|
||||
real_manual_seed = torch.manual_seed
|
||||
monkeypatch.setattr(
|
||||
torch, "manual_seed", lambda seed: (seeds.append(seed), real_manual_seed(seed))[1],
|
||||
)
|
||||
|
||||
sidecar._handle_synthesize({
|
||||
"text": "hello",
|
||||
"seed": 123,
|
||||
"t_shift": 0.4,
|
||||
"layer_penalty_factor": 0.2,
|
||||
"position_temperature": 0.7,
|
||||
"class_temperature": 0.8,
|
||||
"audio_chunk_duration": 10,
|
||||
"audio_chunk_threshold": 0.6,
|
||||
}, object())
|
||||
|
||||
assert seeds == [123]
|
||||
assert calls == [{
|
||||
"text": "hello",
|
||||
"ref_audio": None,
|
||||
"ref_text": None,
|
||||
"t_shift": 0.4,
|
||||
"layer_penalty_factor": 0.2,
|
||||
"position_temperature": 0.7,
|
||||
"class_temperature": 0.8,
|
||||
"audio_chunk_duration": 10,
|
||||
"audio_chunk_threshold": 0.6,
|
||||
}]
|
||||
assert frames[-1]["op"] == "audio"
|
||||
|
||||
|
||||
def test_generation_proxy_forwards_native_controls_and_seed():
|
||||
import torch
|
||||
from api.routers.generation import _run_backend_inference
|
||||
|
||||
calls = []
|
||||
|
||||
class Proxy:
|
||||
id = "omnivoice"
|
||||
display_name = "OmniVoice"
|
||||
sample_rate = 24000
|
||||
applies_own_mastering = True
|
||||
supports_native_omnivoice_controls = True
|
||||
|
||||
def generate(self, text, **kwargs):
|
||||
calls.append((text, kwargs))
|
||||
return torch.zeros(1, 240)
|
||||
|
||||
_run_backend_inference(
|
||||
Proxy(), "hello", "en", None, None, None, None,
|
||||
16, 2.0, 1.0, False, False, 321,
|
||||
t_shift=0.4, layer_penalty_factor=0.2,
|
||||
position_temperature=0.7, class_temperature=0.8,
|
||||
)
|
||||
|
||||
assert calls == [("hello", {
|
||||
"duration": None,
|
||||
"language": "en",
|
||||
"ref_audio": None,
|
||||
"ref_text": None,
|
||||
"instruct": None,
|
||||
"num_step": 16,
|
||||
"guidance_scale": 2.0,
|
||||
"speed": 1.0,
|
||||
"denoise": False,
|
||||
"postprocess_output": False,
|
||||
"t_shift": 0.4,
|
||||
"layer_penalty_factor": 0.2,
|
||||
"position_temperature": 0.7,
|
||||
"class_temperature": 0.8,
|
||||
"seed": 321,
|
||||
})]
|
||||
|
||||
@@ -0,0 +1,76 @@
|
||||
"""#1618 — RAM preflight must not hard-block the machines it means to admit.
|
||||
|
||||
An "8 GB" machine reports ~7.8 GB usable (firmware/iGPU/kernel reservations),
|
||||
so comparing reported RAM against the marketing-size threshold blocked exactly
|
||||
the boundary hardware the ≥8 GB rule intends to allow. The check now applies
|
||||
``_RAM_RESERVED_ALLOWANCE`` to both thresholds, and
|
||||
``OMNIVOICE_RAM_PREFLIGHT=0`` downgrades a genuine fail to a warning.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from api.routers.setup import wizard
|
||||
|
||||
|
||||
def _ram_check(monkeypatch, ram_gb: float, env: str | None = None) -> dict:
|
||||
# Keep the preflight hermetic: stub the probes that hit the network or
|
||||
# auto-acquire media tools, so each RAM assertion stays fast and offline.
|
||||
monkeypatch.setattr(wizard, "_network_check", lambda: {
|
||||
"id": "network", "label": "Network", "status": "pass",
|
||||
"detail": "stubbed", "fix": None, "mirror_reachable": True,
|
||||
})
|
||||
import services.media_tools as media_tools
|
||||
monkeypatch.setattr(media_tools, "summary", lambda auto_acquire=True: None)
|
||||
monkeypatch.setattr(wizard, "_ram_gb", lambda: ram_gb)
|
||||
if env is None:
|
||||
monkeypatch.delenv("OMNIVOICE_RAM_PREFLIGHT", raising=False)
|
||||
else:
|
||||
monkeypatch.setenv("OMNIVOICE_RAM_PREFLIGHT", env)
|
||||
resp = wizard.preflight()
|
||||
checks = resp["checks"] if isinstance(resp, dict) else resp.checks
|
||||
for c in checks:
|
||||
c = c if isinstance(c, dict) else c.model_dump()
|
||||
if c["id"] == "ram":
|
||||
return c
|
||||
raise AssertionError("no ram check in preflight response")
|
||||
|
||||
|
||||
def test_8gb_installed_reporting_7_84_usable_is_not_blocked(monkeypatch):
|
||||
"""The #1618 report: 7.84 GB usable on an 8 GB laptop was a hard fail."""
|
||||
check = _ram_check(monkeypatch, 7.84)
|
||||
assert check["status"] != "fail"
|
||||
|
||||
|
||||
def test_boundary_at_allowance_passes_the_fail_gate(monkeypatch):
|
||||
check = _ram_check(
|
||||
monkeypatch, wizard._RAM_FAIL_GB * wizard._RAM_RESERVED_ALLOWANCE
|
||||
)
|
||||
assert check["status"] != "fail"
|
||||
|
||||
|
||||
def test_genuinely_low_ram_still_fails(monkeypatch):
|
||||
check = _ram_check(monkeypatch, 6.0)
|
||||
assert check["status"] == "fail"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("env", ["0", "false", "no"])
|
||||
def test_escape_hatch_downgrades_fail_to_warn(monkeypatch, env):
|
||||
check = _ram_check(monkeypatch, 6.0, env=env)
|
||||
assert check["status"] == "warn"
|
||||
assert "OMNIVOICE_RAM_PREFLIGHT" in (check["fix"] or "")
|
||||
|
||||
|
||||
def test_escape_hatch_not_triggered_by_other_values(monkeypatch):
|
||||
check = _ram_check(monkeypatch, 6.0, env="1")
|
||||
assert check["status"] == "fail"
|
||||
|
||||
|
||||
def test_12gb_installed_reporting_11_8_usable_passes_clean(monkeypatch):
|
||||
"""Same reservation gap at the warn threshold: 12 GB installed ≈ 11.8."""
|
||||
check = _ram_check(monkeypatch, 11.8)
|
||||
assert check["status"] == "pass"
|
||||
|
||||
|
||||
def test_warn_band_between_thresholds(monkeypatch):
|
||||
check = _ram_check(monkeypatch, 9.0)
|
||||
assert check["status"] == "warn"
|
||||
@@ -76,6 +76,35 @@ class TestUnloadOnABC:
|
||||
)
|
||||
|
||||
|
||||
def test_omnivoice_native_batch_preserves_per_item_controls():
|
||||
"""The adapter forwards variable-length batch controls to OmniVoice."""
|
||||
import torch
|
||||
|
||||
tts = _load_tts_backend_module()
|
||||
calls = []
|
||||
|
||||
class _Model:
|
||||
sampling_rate = 24000
|
||||
|
||||
def generate(self, **kwargs):
|
||||
calls.append(kwargs)
|
||||
return [torch.zeros(1, 12000), torch.zeros(1, 24000)]
|
||||
|
||||
backend = tts.OmniVoiceBackend(model=_Model())
|
||||
outputs = backend.generate_batch(
|
||||
["short", "long"],
|
||||
language=["en", "es"],
|
||||
duration=[0.5, 1.0],
|
||||
speed=[1.0, 0.8],
|
||||
)
|
||||
|
||||
assert [output.shape[-1] for output in outputs] == [12000, 24000]
|
||||
assert calls[0]["text"] == ["short", "long"]
|
||||
assert calls[0]["language"] == ["en", "es"]
|
||||
assert calls[0]["duration"] == [0.5, 1.0]
|
||||
assert calls[0]["speed"] == [1.0, 0.8]
|
||||
|
||||
|
||||
class TestUnloadDefaultBehavior:
|
||||
"""The default no-op must actually be safe to call."""
|
||||
|
||||
@@ -154,4 +183,4 @@ class TestExistingSubclassesInherit:
|
||||
assert callable(getattr(cls, "unload", None)), (
|
||||
f"{cls.__name__} has no callable unload() — even via the "
|
||||
"ABC inheritance. Did someone shadow it?"
|
||||
)
|
||||
)
|
||||
|
||||
+857
-86
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,61 @@
|
||||
"""Cancellation helpers for work that cannot be stopped mid-call."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Callable
|
||||
from typing import Any, TypeVar
|
||||
|
||||
_Result = TypeVar("_Result")
|
||||
|
||||
|
||||
async def drain_task(task: asyncio.Task[Any]) -> None:
|
||||
"""Wait for ``task`` even if the waiter is cancelled again."""
|
||||
while not task.done():
|
||||
try:
|
||||
await asyncio.shield(task)
|
||||
except asyncio.CancelledError:
|
||||
continue
|
||||
except BaseException:
|
||||
break
|
||||
if task.done():
|
||||
try:
|
||||
task.result()
|
||||
except BaseException:
|
||||
pass
|
||||
|
||||
|
||||
async def to_thread_and_drain_on_cancel(
|
||||
function: Callable[..., _Result], /, *args: Any
|
||||
) -> _Result:
|
||||
"""Run a blocking call without detaching it when its waiter is cancelled."""
|
||||
thread_task = asyncio.create_task(asyncio.to_thread(function, *args))
|
||||
try:
|
||||
return await asyncio.shield(thread_task)
|
||||
except asyncio.CancelledError:
|
||||
await drain_task(thread_task)
|
||||
raise
|
||||
|
||||
|
||||
async def to_thread_and_defer_cancellation(
|
||||
function: Callable[..., _Result], /, *args: Any
|
||||
) -> tuple[_Result, bool]:
|
||||
"""Finish a durable call and report cancellation after its result is known.
|
||||
|
||||
Authority writes need their event-loop publication even when the HTTP
|
||||
caller disappears while SQLite is committing. Returning the cancellation
|
||||
flag lets the caller publish that result first, then propagate cancellation.
|
||||
"""
|
||||
thread_task = asyncio.create_task(asyncio.to_thread(function, *args))
|
||||
try:
|
||||
return await asyncio.shield(thread_task), False
|
||||
except asyncio.CancelledError:
|
||||
await drain_task(thread_task)
|
||||
return thread_task.result(), True
|
||||
|
||||
|
||||
__all__ = [
|
||||
"drain_task",
|
||||
"to_thread_and_defer_cancellation",
|
||||
"to_thread_and_drain_on_cancel",
|
||||
]
|
||||
@@ -59,10 +59,18 @@ _VRAM_PER_JOB_BYTES = 5 * 1024**3
|
||||
# cpu — oversubscription just thrashes
|
||||
_ALWAYS_SERIAL = frozenset({"mps", "mlx", "cpu", ""})
|
||||
|
||||
# Absolute ceiling regardless of how much memory a card reports. Beyond this
|
||||
# the bottleneck stops being VRAM and starts being scheduler overhead and
|
||||
# host-side I/O contention.
|
||||
_MAX_DERIVED = 4
|
||||
# Absolute protocol ceiling regardless of how much memory a peer reports.
|
||||
# Beyond this the bottleneck stops being VRAM and starts being scheduler
|
||||
# overhead and host-side I/O contention. It is public because every wire
|
||||
# boundary must clamp to the same number; a UINT32_MAX heartbeat must not grow
|
||||
# a scheduler queue that local derivation would never create.
|
||||
MAX_CONCURRENT_TASKS = 4
|
||||
|
||||
|
||||
def clamp_concurrency(value: int, *, allow_zero: bool = False) -> int:
|
||||
"""Bound an advertised concurrency value to the server's safe range."""
|
||||
minimum = 0 if allow_zero else 1
|
||||
return max(minimum, min(MAX_CONCURRENT_TASKS, int(value)))
|
||||
|
||||
# Bounds on how long a parked slot is held. The caller passes the timed-out
|
||||
# job's own execution budget — the longest its thread can still legitimately be
|
||||
@@ -104,7 +112,7 @@ def derive_concurrency(
|
||||
budget = max(min_model_bytes, _VRAM_PER_JOB_BYTES)
|
||||
if budget <= 0:
|
||||
return 1
|
||||
return max(1, min(_MAX_DERIVED, int(free_memory_bytes // budget)))
|
||||
return clamp_concurrency(int(free_memory_bytes // budget))
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -149,6 +157,9 @@ class WorkerCapacity:
|
||||
resident_models: set[str] = field(default_factory=set)
|
||||
slots: dict[str, ModelSlot] = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.max_concurrent_tasks = clamp_concurrency(self.max_concurrent_tasks)
|
||||
|
||||
@staticmethod
|
||||
def slot_key(engine: str, model_id: str) -> str:
|
||||
return f"{engine}:{model_id}"
|
||||
@@ -198,6 +209,21 @@ class WorkerCapacity:
|
||||
)
|
||||
slot.active += 1
|
||||
|
||||
def reserve_unknown(self) -> None:
|
||||
"""Consume worker-wide capacity for claimed work we cannot classify.
|
||||
|
||||
Reconciliation will tell the peer to cancel a terminal or unknown
|
||||
attempt, but until that cancellation lands it is still using the GPU.
|
||||
"""
|
||||
self.active_tasks += 1
|
||||
|
||||
def release_unknown(self) -> bool:
|
||||
"""Release one exact reconciled claim with no model-slot identity."""
|
||||
if self.active_tasks <= 0:
|
||||
return False
|
||||
self.active_tasks -= 1
|
||||
return True
|
||||
|
||||
def release(
|
||||
self,
|
||||
engine: str,
|
||||
@@ -267,8 +293,12 @@ class WorkerCapacity:
|
||||
) -> None:
|
||||
"""Adopt a heartbeat snapshot. The worker is the source of truth for
|
||||
what it is actually running."""
|
||||
self.active_tasks = max(0, active_tasks)
|
||||
reported_ceiling = self.active_tasks + max(0, available_slots)
|
||||
self.active_tasks = clamp_concurrency(active_tasks, allow_zero=True)
|
||||
bounded_available = clamp_concurrency(available_slots, allow_zero=True)
|
||||
bounded_available = min(
|
||||
bounded_available, MAX_CONCURRENT_TASKS - self.active_tasks
|
||||
)
|
||||
reported_ceiling = self.active_tasks + bounded_available
|
||||
if reported_ceiling > 0:
|
||||
# Adopted, not merely grown. The worker computes this as its own
|
||||
# ``max_concurrent_tasks``, so a ceiling we refuse to lower is one
|
||||
@@ -305,4 +335,10 @@ class WorkerCapacity:
|
||||
}
|
||||
|
||||
|
||||
__all__ = ["ModelSlot", "WorkerCapacity", "derive_concurrency"]
|
||||
__all__ = [
|
||||
"MAX_CONCURRENT_TASKS",
|
||||
"ModelSlot",
|
||||
"WorkerCapacity",
|
||||
"clamp_concurrency",
|
||||
"derive_concurrency",
|
||||
]
|
||||
|
||||
@@ -35,6 +35,9 @@ from typing import Optional
|
||||
# plane may run in a process that never loads torch). test_worker_deadlines.py
|
||||
# asserts the two agree, so a change there cannot silently drift from here.
|
||||
_GENERATE_TIMEOUT_S = float(os.environ.get("OMNIVOICE_GENERATE_TIMEOUT_S", "300.0"))
|
||||
_CPU_GENERATE_TIMEOUT_S = float(
|
||||
os.environ.get("OMNIVOICE_CPU_GENERATE_TIMEOUT_S", "600.0")
|
||||
)
|
||||
_MODEL_LOAD_EXTRA_S = float(os.environ.get("OMNIVOICE_MODEL_LOAD_TIMEOUT_S", "1800.0"))
|
||||
_HEARTBEAT_GRACE_S = float(os.environ.get("OMNIVOICE_MODEL_LOAD_HEARTBEAT_GRACE_S", "30.0"))
|
||||
|
||||
@@ -123,20 +126,40 @@ class Deadlines:
|
||||
}
|
||||
|
||||
|
||||
def _base_execution_seconds(text: Optional[str]) -> float:
|
||||
def _base_execution_seconds(
|
||||
text: Optional[str], *, execution_device: Optional[str] = None
|
||||
) -> float:
|
||||
"""Delegate to model_manager's budget; fall back to its formula.
|
||||
|
||||
The lazy import keeps this module usable in a process that has no torch —
|
||||
the control plane schedules work it never executes.
|
||||
"""
|
||||
target_device = str(execution_device or "cpu").lower()
|
||||
if target_device not in {"cpu", "cuda", "mps", "mlx", "directml", "rocm", "xpu"}:
|
||||
target_device = "cpu"
|
||||
try:
|
||||
from services import model_manager # noqa: PLC0415 — intentionally lazy
|
||||
|
||||
return float(model_manager.generate_timeout_s(text))
|
||||
return float(
|
||||
model_manager.generate_timeout_s(
|
||||
text, execution_device=target_device
|
||||
)
|
||||
)
|
||||
except Exception:
|
||||
base = _GENERATE_TIMEOUT_S
|
||||
try:
|
||||
if (
|
||||
target_device == "cpu"
|
||||
and "OMNIVOICE_GENERATE_TIMEOUT_S" not in os.environ
|
||||
):
|
||||
base = _CPU_GENERATE_TIMEOUT_S
|
||||
except Exception:
|
||||
# Capability detection is optional in the torch-free control
|
||||
# plane; retain the configured universal bounded fallback.
|
||||
pass
|
||||
return max(
|
||||
_GENERATE_TIMEOUT_S,
|
||||
_GENERATE_TIMEOUT_S + max(0, len(text or "") - _FREE_CHARS) / _CHARS_PER_SECOND,
|
||||
base,
|
||||
base + max(0, len(text or "") - _FREE_CHARS) / _CHARS_PER_SECOND,
|
||||
)
|
||||
|
||||
|
||||
@@ -147,6 +170,7 @@ def for_task(
|
||||
model_resident: bool = False,
|
||||
model_downloaded: bool = True,
|
||||
input_seconds: float = 0.0,
|
||||
execution_device: Optional[str] = None,
|
||||
) -> Deadlines:
|
||||
"""Compute the deadlines for one attempt.
|
||||
|
||||
@@ -158,7 +182,9 @@ def for_task(
|
||||
op = Operation.coerce(operation)
|
||||
multiplier, grace = _PROFILE[op]
|
||||
|
||||
execution = _base_execution_seconds(text) * multiplier
|
||||
execution = _base_execution_seconds(
|
||||
text, execution_device=execution_device
|
||||
) * multiplier
|
||||
# Media-length operations scale on duration, not characters.
|
||||
if input_seconds > 0:
|
||||
execution = max(execution, input_seconds * multiplier)
|
||||
|
||||
+520
-91
@@ -17,15 +17,19 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import errno
|
||||
import hashlib
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import io
|
||||
import zipfile
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
import zipfile
|
||||
from typing import Any, Awaitable, Callable, Optional
|
||||
|
||||
from worker.async_utils import drain_task
|
||||
from worker.errors import ErrorClass, WorkerError
|
||||
|
||||
logger = logging.getLogger("omnivoice.worker")
|
||||
@@ -44,6 +48,18 @@ INPUT_ERRORS_PARAM = "input_errors"
|
||||
# leak on the worker that unpurged artifacts were on the control plane.
|
||||
INPUT_CACHE_LIMIT_BYTES = 2 * 1024 * 1024 * 1024
|
||||
_FALLBACK_INPUT_FETCH_SECONDS = 600.0
|
||||
_STALE_INPUT_PARTIAL_SECONDS = 60 * 60.0
|
||||
|
||||
# Pruning runs in worker threads and every executor instance shares the same
|
||||
# on-disk cache, so active-path leases are process-wide and thread-safe.
|
||||
_INPUT_CACHE_LEASE_LOCK = threading.Lock()
|
||||
_INPUT_CACHE_LEASES: dict[str, int] = {}
|
||||
_INPUT_CACHE_FETCH_LEASES: dict[str, int] = {}
|
||||
_INPUT_CACHE_MUTATIONS: set[str] = set()
|
||||
# Concurrent fetches keep distinct partial files but serialize the instant a
|
||||
# verified generation is published at its content address.
|
||||
_INPUT_CACHE_FETCH_LOCKS: dict[str, asyncio.Lock] = {}
|
||||
_INPUT_CACHE_FETCH_USERS: dict[str, int] = {}
|
||||
|
||||
# on_progress(fraction: float, stage: str)
|
||||
# on_model_loading(fraction: float, detail: str)
|
||||
@@ -88,6 +104,7 @@ class TaskExecutor:
|
||||
self._on_model_loading = on_model_loading
|
||||
self._fetch_input = fetch_input
|
||||
self._input_dir = input_dir
|
||||
self._blocking_tasks: set[asyncio.Task] = set()
|
||||
|
||||
async def execute(
|
||||
self,
|
||||
@@ -110,33 +127,35 @@ class TaskExecutor:
|
||||
"""
|
||||
operation = (assignment.operation or "tts").lower()
|
||||
params = _parse_params(assignment.params_json)
|
||||
params = await self._materialize_inputs(
|
||||
params, leased_inputs = await self._materialize_inputs(
|
||||
assignment, params, fetch_input or self._fetch_input
|
||||
)
|
||||
|
||||
handler = {
|
||||
"tts": self._run_tts,
|
||||
"clone": self._run_tts,
|
||||
"audiobook": self._run_audiobook,
|
||||
"dub_segments": self._run_dub_segments,
|
||||
}.get(operation)
|
||||
if handler is None:
|
||||
raise TaskFailure(
|
||||
WorkerError(
|
||||
error_class=ErrorClass.CAPABILITY,
|
||||
code="OPERATION_UNSUPPORTED",
|
||||
message=f"This worker cannot run '{operation}' tasks.",
|
||||
hint="Run this task locally, or use a worker that supports it.",
|
||||
try:
|
||||
handler = {
|
||||
"tts": self._run_tts,
|
||||
"clone": self._run_tts,
|
||||
"audiobook": self._run_audiobook,
|
||||
"dub_segments": self._run_dub_segments,
|
||||
}.get(operation)
|
||||
if handler is None:
|
||||
raise TaskFailure(
|
||||
WorkerError(
|
||||
error_class=ErrorClass.CAPABILITY,
|
||||
code="OPERATION_UNSUPPORTED",
|
||||
message=f"This worker cannot run '{operation}' tasks.",
|
||||
hint="Run this task locally, or use a worker that supports it.",
|
||||
)
|
||||
)
|
||||
return await handler(
|
||||
assignment,
|
||||
params,
|
||||
_Reporters(
|
||||
on_progress or self._on_progress,
|
||||
on_model_loading or self._on_model_loading,
|
||||
),
|
||||
)
|
||||
return await handler(
|
||||
assignment,
|
||||
params,
|
||||
_Reporters(
|
||||
on_progress or self._on_progress,
|
||||
on_model_loading or self._on_model_loading,
|
||||
),
|
||||
)
|
||||
finally:
|
||||
self._release_inputs_after_active_work(leased_inputs)
|
||||
|
||||
async def _run_dub_segments(self, assignment, params: dict, report: "_Reporters") -> dict:
|
||||
"""Render every requested dub line under one lease and return one bundle."""
|
||||
@@ -150,8 +169,9 @@ class TaskExecutor:
|
||||
))
|
||||
load_budget, run_budget = _budgets(assignment)
|
||||
await report.loading(0.0, f"preparing {assignment.engine}")
|
||||
backend = await self._bounded(
|
||||
asyncio.to_thread(self._load_backend, assignment.engine),
|
||||
backend = await self._bounded_thread(
|
||||
self._load_backend,
|
||||
assignment.engine,
|
||||
timeout=load_budget, code="MODEL_LOAD_TIMEOUT", what=f"Loading '{assignment.engine}'",
|
||||
)
|
||||
await report.loading(1.0, "model ready")
|
||||
@@ -159,11 +179,15 @@ class TaskExecutor:
|
||||
for index, row in enumerate(rows):
|
||||
row = dict(row)
|
||||
row["ref_audio"] = refs[index] if index < len(refs) else None
|
||||
audio = await self._bounded(
|
||||
asyncio.to_thread(self._synthesize_dub_segment, backend, row),
|
||||
audio = await self._bounded_thread(
|
||||
self._synthesize_dub_segment,
|
||||
backend,
|
||||
row,
|
||||
timeout=run_budget, code="EXECUTION_TIMEOUT", what=f"Dubbing segment {index + 1}",
|
||||
)
|
||||
payload, _meta = await asyncio.to_thread(self._encode, audio, row, backend)
|
||||
payload, _meta = await self._thread_call(
|
||||
self._encode, audio, row, backend
|
||||
)
|
||||
rendered.append((int(row.get("index", index)), payload))
|
||||
await report.progress((index + 1) / len(rows), f"segment {index + 1} of {len(rows)}")
|
||||
|
||||
@@ -184,9 +208,11 @@ class TaskExecutor:
|
||||
from services.text_normalization import normalize_for_tts
|
||||
|
||||
text = normalize_for_tts(row.get("text") or "", row.get("language"))
|
||||
seed = None
|
||||
if row.get("seed") is not None:
|
||||
import torch
|
||||
torch.manual_seed(int(row["seed"]))
|
||||
seed = int(row["seed"])
|
||||
torch.manual_seed(seed)
|
||||
kwargs = {
|
||||
"language": row.get("language") if row.get("language") != "Auto" else None,
|
||||
"ref_audio": row.get("ref_audio"), "ref_text": row.get("ref_text"),
|
||||
@@ -197,6 +223,11 @@ class TaskExecutor:
|
||||
"speed": float(row.get("speed") or 1.0), "denoise": True,
|
||||
"postprocess_output": True,
|
||||
}
|
||||
if (
|
||||
getattr(backend, "supports_native_omnivoice_controls", False)
|
||||
and seed is not None
|
||||
):
|
||||
kwargs["seed"] = seed
|
||||
audio = backend.generate(text=text, **kwargs)
|
||||
preset = row.get("effect_preset") or "broadcast"
|
||||
if preset != "raw":
|
||||
@@ -225,8 +256,9 @@ class TaskExecutor:
|
||||
load_budget, run_budget = _budgets(assignment)
|
||||
|
||||
await report.loading(0.0, f"preparing {assignment.engine}")
|
||||
backend = await self._bounded(
|
||||
asyncio.to_thread(self._load_backend, assignment.engine),
|
||||
backend = await self._bounded_thread(
|
||||
self._load_backend,
|
||||
assignment.engine,
|
||||
timeout=load_budget,
|
||||
code="MODEL_LOAD_TIMEOUT",
|
||||
what=f"Loading '{assignment.engine}'",
|
||||
@@ -234,16 +266,22 @@ class TaskExecutor:
|
||||
await report.loading(1.0, "model ready")
|
||||
|
||||
await report.progress(0.05, "synthesising")
|
||||
audio = await self._bounded(
|
||||
asyncio.to_thread(self._synthesize, backend, text, params),
|
||||
audio = await self._bounded_thread(
|
||||
self._synthesize,
|
||||
backend,
|
||||
text,
|
||||
params,
|
||||
timeout=run_budget,
|
||||
code="EXECUTION_TIMEOUT",
|
||||
what="Synthesis",
|
||||
)
|
||||
await report.progress(0.9, "encoding")
|
||||
|
||||
payload, meta = await self._bounded(
|
||||
asyncio.to_thread(self._encode, audio, params, backend),
|
||||
payload, meta = await self._bounded_thread(
|
||||
self._encode,
|
||||
audio,
|
||||
params,
|
||||
backend,
|
||||
timeout=run_budget,
|
||||
code="EXECUTION_TIMEOUT",
|
||||
what="Encoding",
|
||||
@@ -264,19 +302,26 @@ class TaskExecutor:
|
||||
))
|
||||
load_budget, run_budget = _budgets(assignment)
|
||||
await report.loading(0.0, f"preparing {assignment.engine}")
|
||||
backend = await self._bounded(
|
||||
asyncio.to_thread(self._load_backend, assignment.engine),
|
||||
backend = await self._bounded_thread(
|
||||
self._load_backend,
|
||||
assignment.engine,
|
||||
timeout=load_budget, code="MODEL_LOAD_TIMEOUT",
|
||||
what=f"Loading '{assignment.engine}'",
|
||||
)
|
||||
await report.loading(1.0, "model ready")
|
||||
await report.progress(0.05, "synthesising chapter")
|
||||
audio = await self._bounded(
|
||||
asyncio.to_thread(self._synthesize_audiobook, backend, spans, voices, params),
|
||||
audio = await self._bounded_thread(
|
||||
self._synthesize_audiobook,
|
||||
backend,
|
||||
spans,
|
||||
voices,
|
||||
params,
|
||||
timeout=run_budget, code="EXECUTION_TIMEOUT", what="Audiobook chapter",
|
||||
)
|
||||
await report.progress(0.9, "encoding")
|
||||
payload, meta = await asyncio.to_thread(self._encode, audio, params, backend)
|
||||
payload, meta = await self._thread_call(
|
||||
self._encode, audio, params, backend
|
||||
)
|
||||
await report.progress(1.0, "done")
|
||||
return {"meta": meta, "payload": payload}
|
||||
|
||||
@@ -294,7 +339,10 @@ class TaskExecutor:
|
||||
key: value for key, value in opts.to_manifest().items()
|
||||
if value is not None and key not in ("seed", "vary_repeats")
|
||||
}
|
||||
if isinstance(backend, OmniVoiceBackend):
|
||||
native_proxy = bool(
|
||||
getattr(backend, "supports_native_omnivoice_controls", False)
|
||||
)
|
||||
if isinstance(backend, OmniVoiceBackend) or native_proxy:
|
||||
extra.setdefault("num_step", 32)
|
||||
extra.setdefault("guidance_scale", 2.0)
|
||||
for key in ("emo_vector", "emo_text", "emo_alpha"):
|
||||
@@ -303,11 +351,13 @@ class TaskExecutor:
|
||||
def synth(text, index, speed=None):
|
||||
voice = voices[int(index)]
|
||||
base_seed = opts.seed if opts.seed is not None else voice.get("seed")
|
||||
seed = None
|
||||
if base_seed is not None:
|
||||
import torch
|
||||
nonce = occurrence["value"] if opts.vary_repeats else 0
|
||||
occurrence["value"] += 1
|
||||
torch.manual_seed(segment_seed(base_seed, text, nonce))
|
||||
seed = segment_seed(base_seed, text, nonce)
|
||||
torch.manual_seed(seed)
|
||||
kwargs = {
|
||||
"language": language,
|
||||
"ref_audio": voice.get("ref_audio"),
|
||||
@@ -316,6 +366,8 @@ class TaskExecutor:
|
||||
"speed": float(speed) if speed else 1.0,
|
||||
**extra,
|
||||
}
|
||||
if native_proxy and seed is not None:
|
||||
kwargs["seed"] = seed
|
||||
return backend.generate(text, **kwargs)
|
||||
|
||||
spans = [Span(voice_id=str(i), text=row.get("text", ""),
|
||||
@@ -329,7 +381,9 @@ class TaskExecutor:
|
||||
|
||||
# ── Inputs ────────────────────────────────────────────────────────────
|
||||
|
||||
async def _materialize_inputs(self, assignment, params: dict, fetch) -> dict:
|
||||
async def _materialize_inputs(
|
||||
self, assignment, params: dict, fetch
|
||||
) -> tuple[dict, list[str]]:
|
||||
"""Turn declared inputs into local files, then point the params at them.
|
||||
|
||||
The control plane sends artifact ids, never paths — its own paths mean
|
||||
@@ -352,7 +406,7 @@ class TaskExecutor:
|
||||
|
||||
refs = [ref for ref in (getattr(assignment, "inputs", None) or []) if ref.artifact_id]
|
||||
if not refs:
|
||||
return params
|
||||
return params, []
|
||||
if fetch is None:
|
||||
raise TaskFailure(
|
||||
WorkerError(
|
||||
@@ -365,16 +419,24 @@ class TaskExecutor:
|
||||
|
||||
_, run_budget = _budgets(assignment)
|
||||
local: dict[str, str] = {}
|
||||
for ref in refs:
|
||||
local[ref.artifact_id] = await self._bounded(
|
||||
self._fetch_one(ref, fetch),
|
||||
timeout=min(run_budget, _FALLBACK_INPUT_FETCH_SECONDS),
|
||||
code="INPUT_FETCH_TIMEOUT",
|
||||
what=f"Fetching '{ref.filename or ref.artifact_id}'",
|
||||
)
|
||||
return _rewrite_params(params, local)
|
||||
leased: list[str] = []
|
||||
try:
|
||||
for ref in refs:
|
||||
path = await self._bounded(
|
||||
self._fetch_one(ref, fetch, retain=True),
|
||||
timeout=min(run_budget, _FALLBACK_INPUT_FETCH_SECONDS),
|
||||
code="INPUT_FETCH_TIMEOUT",
|
||||
what=f"Fetching '{ref.filename or ref.artifact_id}'",
|
||||
)
|
||||
local[ref.artifact_id] = path
|
||||
leased.append(path)
|
||||
return _rewrite_params(params, local), leased
|
||||
except BaseException:
|
||||
for path in leased:
|
||||
_release_input_cache_path(path)
|
||||
raise
|
||||
|
||||
async def _fetch_one(self, ref, fetch) -> str:
|
||||
async def _fetch_one(self, ref, fetch, *, retain: bool = False) -> str:
|
||||
"""The local copy of one input, downloaded only if we lack it.
|
||||
|
||||
Content-addressed: the name is the hash the control plane computed, so
|
||||
@@ -382,37 +444,127 @@ class TaskExecutor:
|
||||
worker — costs no transfer at all.
|
||||
"""
|
||||
directory = self._input_dir or default_input_dir()
|
||||
os.makedirs(directory, exist_ok=True)
|
||||
await self._thread_call(_durable_makedirs, directory)
|
||||
destination = os.path.join(directory, _cache_name(ref))
|
||||
if _already_held(destination, ref):
|
||||
_touch(destination)
|
||||
return destination
|
||||
return await self._fetch_one_owned(
|
||||
ref,
|
||||
fetch,
|
||||
directory=directory,
|
||||
destination=destination,
|
||||
retain=retain,
|
||||
)
|
||||
|
||||
partial = f"{destination}.{uuid.uuid4().hex}.part"
|
||||
async def _fetch_one_owned(
|
||||
self,
|
||||
ref,
|
||||
fetch,
|
||||
*,
|
||||
directory: str,
|
||||
destination: str,
|
||||
retain: bool,
|
||||
) -> str:
|
||||
"""Validate, fetch, and safely publish one content address."""
|
||||
await _acquire_input_cache_path(destination, fetching=True)
|
||||
leased_result = destination
|
||||
succeeded = False
|
||||
try:
|
||||
await fetch(ref, partial)
|
||||
except TaskFailure:
|
||||
raise
|
||||
except Exception as exc:
|
||||
_discard(partial)
|
||||
raise TaskFailure(
|
||||
WorkerError(
|
||||
# Transient on purpose: an id we cannot resolve now is far
|
||||
# more often a dropped stream than a permanently missing
|
||||
# file, and one wasted retry beats failing real work.
|
||||
error_class=ErrorClass.TRANSIENT,
|
||||
code="INPUT_FETCH_FAILED",
|
||||
message=f"Could not fetch '{ref.filename or ref.artifact_id}': {exc}",
|
||||
hint="The control plane may have restarted; the task will be retried.",
|
||||
)
|
||||
) from exc
|
||||
# Cache hits still hash the advertised content identity. Filename
|
||||
# plus size is not proof after disk corruption or external edits.
|
||||
if await self._thread_call(_already_held, destination, ref):
|
||||
await self._thread_call(_touch, destination)
|
||||
succeeded = True
|
||||
return destination
|
||||
|
||||
# Off the loop: hashing a source video on the event loop thread would
|
||||
# stall every heartbeat this worker owes the control plane.
|
||||
await asyncio.to_thread(_verify, partial, ref)
|
||||
os.replace(partial, destination)
|
||||
await asyncio.to_thread(_prune_input_cache, directory)
|
||||
return destination
|
||||
partial = f"{destination}.{uuid.uuid4().hex}.part"
|
||||
_lease_input_cache_path(partial)
|
||||
finalized = False
|
||||
try:
|
||||
try:
|
||||
await fetch(ref, partial)
|
||||
except TaskFailure:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise _input_fetch_failure(ref, exc) from exc
|
||||
|
||||
# Hashing and durability barriers can both block on a large
|
||||
# source or slow disk. Keep them off the loop and drain before
|
||||
# cleanup so Windows never unlinks a file still in use.
|
||||
await self._thread_call(_verify, partial, ref)
|
||||
gate_key = _cache_path_key(destination)
|
||||
gate = _INPUT_CACHE_FETCH_LOCKS.setdefault(
|
||||
gate_key, asyncio.Lock()
|
||||
)
|
||||
_INPUT_CACHE_FETCH_USERS[gate_key] = (
|
||||
_INPUT_CACHE_FETCH_USERS.get(gate_key, 0) + 1
|
||||
)
|
||||
try:
|
||||
async with gate:
|
||||
# A concurrent fetch may have published these exact
|
||||
# bytes while this one was downloading its own partial.
|
||||
if await self._thread_call(
|
||||
_already_held, destination, ref
|
||||
):
|
||||
await self._thread_call(_discard, partial)
|
||||
finalized = True
|
||||
elif _claim_input_cache_mutation(destination):
|
||||
try:
|
||||
await self._thread_call(
|
||||
_durable_replace, partial, destination
|
||||
)
|
||||
finally:
|
||||
_finish_input_cache_mutation(destination)
|
||||
finalized = True
|
||||
else:
|
||||
# Another execution is actively reading the
|
||||
# canonical generation. Never unlink or replace
|
||||
# bytes underneath it; publish this verified fetch
|
||||
# under a leased sibling path and let a later
|
||||
# unshared fetch repair canonical.
|
||||
stem, suffix = os.path.splitext(destination)
|
||||
alternate = (
|
||||
f"{stem}.{uuid.uuid4().hex}.generation{suffix}"
|
||||
)
|
||||
await _acquire_input_cache_path(
|
||||
alternate, fetching=True
|
||||
)
|
||||
try:
|
||||
await self._thread_call(
|
||||
_durable_replace, partial, alternate
|
||||
)
|
||||
except BaseException:
|
||||
_release_input_cache_path(
|
||||
alternate, fetching=True
|
||||
)
|
||||
await self._thread_call(_discard, alternate)
|
||||
raise
|
||||
finalized = True
|
||||
_release_input_cache_path(
|
||||
destination, fetching=True
|
||||
)
|
||||
leased_result = alternate
|
||||
except OSError as exc:
|
||||
raise _input_fetch_failure(ref, exc) from exc
|
||||
finally:
|
||||
remaining = _INPUT_CACHE_FETCH_USERS[gate_key] - 1
|
||||
if remaining:
|
||||
_INPUT_CACHE_FETCH_USERS[gate_key] = remaining
|
||||
else:
|
||||
_INPUT_CACHE_FETCH_USERS.pop(gate_key, None)
|
||||
if _INPUT_CACHE_FETCH_LOCKS.get(gate_key) is gate:
|
||||
_INPUT_CACHE_FETCH_LOCKS.pop(gate_key, None)
|
||||
finally:
|
||||
if not finalized:
|
||||
await self._thread_call(_discard, partial)
|
||||
_release_input_cache_path(partial)
|
||||
|
||||
await self._thread_call(_prune_input_cache, directory)
|
||||
succeeded = True
|
||||
return leased_result
|
||||
finally:
|
||||
if retain and succeeded:
|
||||
_promote_input_cache_lease(leased_result)
|
||||
else:
|
||||
_release_input_cache_path(leased_result, fetching=True)
|
||||
|
||||
# ── Engine plumbing ───────────────────────────────────────────────────
|
||||
|
||||
@@ -508,6 +660,10 @@ class TaskExecutor:
|
||||
duration, num_step, guidance_scale, speed, denoise,
|
||||
postprocess_output, used_seed, effect_preset,
|
||||
max_chunk_chars, crossfade_ms,
|
||||
t_shift=params.get("t_shift"),
|
||||
layer_penalty_factor=params.get("layer_penalty_factor"),
|
||||
position_temperature=params.get("position_temperature"),
|
||||
class_temperature=params.get("class_temperature"),
|
||||
)
|
||||
except Exception as exc:
|
||||
from worker import errors as worker_errors # noqa: PLC0415
|
||||
@@ -551,6 +707,92 @@ class TaskExecutor:
|
||||
|
||||
# ── Bounding ──────────────────────────────────────────────────────────
|
||||
|
||||
def _release_inputs_after_active_work(self, paths: list[str]) -> None:
|
||||
"""Keep files leased while a timed-out engine thread still owns them."""
|
||||
pending = [task for task in self._blocking_tasks if not task.done()]
|
||||
if not pending:
|
||||
for path in paths:
|
||||
_release_input_cache_path(path)
|
||||
return
|
||||
remaining = {"count": len(pending)}
|
||||
|
||||
def finished(_task: asyncio.Task) -> None:
|
||||
remaining["count"] -= 1
|
||||
if remaining["count"] == 0:
|
||||
for path in paths:
|
||||
_release_input_cache_path(path)
|
||||
|
||||
for task in pending:
|
||||
task.add_done_callback(finished)
|
||||
|
||||
def _start_thread(self, function, /, *args) -> asyncio.Task:
|
||||
task = asyncio.create_task(asyncio.to_thread(function, *args))
|
||||
self._blocking_tasks.add(task)
|
||||
|
||||
def finished(completed: asyncio.Task) -> None:
|
||||
self._blocking_tasks.discard(completed)
|
||||
if not completed.cancelled():
|
||||
# Timed-out calls intentionally finish in the background. Read
|
||||
# their exception so asyncio never reports an unowned task.
|
||||
completed.exception()
|
||||
|
||||
task.add_done_callback(finished)
|
||||
return task
|
||||
|
||||
async def _thread_call(self, function, /, *args):
|
||||
task = self._start_thread(function, *args)
|
||||
try:
|
||||
return await asyncio.shield(task)
|
||||
except asyncio.CancelledError:
|
||||
await drain_task(task)
|
||||
raise
|
||||
|
||||
async def drain_active_work(self) -> None:
|
||||
"""Wait until every blocking engine call has relinquished the process."""
|
||||
cancelled = bool(
|
||||
(current := asyncio.current_task()) is not None and current.cancelling()
|
||||
)
|
||||
while self._blocking_tasks:
|
||||
for task in list(self._blocking_tasks):
|
||||
try:
|
||||
await asyncio.shield(task)
|
||||
except asyncio.CancelledError:
|
||||
# Cancellation cannot make a Python GPU thread stop. Hold
|
||||
# authority until it really exits, then propagate the
|
||||
# cancellation so callers never publish a false free slot.
|
||||
cancelled = True
|
||||
await drain_task(task)
|
||||
except BaseException:
|
||||
# The owner reports/classifies the engine exception. This
|
||||
# barrier only establishes that the thread has finished.
|
||||
pass
|
||||
if cancelled:
|
||||
raise asyncio.CancelledError
|
||||
|
||||
async def _bounded_thread(
|
||||
self, function, /, *args, timeout: float, code: str, what: str
|
||||
):
|
||||
"""Bound a blocking call without losing ownership of its live thread."""
|
||||
task = self._start_thread(function, *args)
|
||||
try:
|
||||
done, _pending = await asyncio.wait({task}, timeout=timeout)
|
||||
except asyncio.CancelledError:
|
||||
await drain_task(task)
|
||||
raise
|
||||
if done:
|
||||
return task.result()
|
||||
# A GPU call cannot be killed. Return the timeout so the scheduler can
|
||||
# park its slot, but retain the task above so terminal authority loss
|
||||
# can drain it before claiming this worker has stopped.
|
||||
raise TaskFailure(
|
||||
WorkerError(
|
||||
error_class=ErrorClass.TIMEOUT,
|
||||
code=code,
|
||||
message=f"{what} exceeded the {timeout:g}s budget for this task.",
|
||||
hint="Try a shorter input, or a worker with more headroom.",
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _bounded(coro, *, timeout: float, code: str, what: str):
|
||||
"""Run ``coro`` under the server's budget for this phase.
|
||||
@@ -636,6 +878,158 @@ def default_input_dir() -> str:
|
||||
return os.path.join(tempfile.gettempdir(), "omnivoice-worker-inputs")
|
||||
|
||||
|
||||
def _cache_path_key(path: str) -> str:
|
||||
return os.path.normcase(os.path.abspath(path))
|
||||
|
||||
|
||||
def _lease_input_cache_path(path: str, *, fetching: bool = False) -> bool:
|
||||
key = _cache_path_key(path)
|
||||
with _INPUT_CACHE_LEASE_LOCK:
|
||||
if key in _INPUT_CACHE_MUTATIONS:
|
||||
return False
|
||||
_INPUT_CACHE_LEASES[key] = _INPUT_CACHE_LEASES.get(key, 0) + 1
|
||||
if fetching:
|
||||
_INPUT_CACHE_FETCH_LEASES[key] = (
|
||||
_INPUT_CACHE_FETCH_LEASES.get(key, 0) + 1
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
async def _acquire_input_cache_path(
|
||||
path: str, *, fetching: bool = False
|
||||
) -> None:
|
||||
while not _lease_input_cache_path(path, fetching=fetching):
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
|
||||
def _release_input_cache_path(path: str, *, fetching: bool = False) -> None:
|
||||
key = _cache_path_key(path)
|
||||
with _INPUT_CACHE_LEASE_LOCK:
|
||||
if fetching:
|
||||
fetch_remaining = _INPUT_CACHE_FETCH_LEASES.get(key, 0) - 1
|
||||
if fetch_remaining > 0:
|
||||
_INPUT_CACHE_FETCH_LEASES[key] = fetch_remaining
|
||||
else:
|
||||
_INPUT_CACHE_FETCH_LEASES.pop(key, None)
|
||||
remaining = _INPUT_CACHE_LEASES.get(key, 0) - 1
|
||||
if remaining > 0:
|
||||
_INPUT_CACHE_LEASES[key] = remaining
|
||||
else:
|
||||
_INPUT_CACHE_LEASES.pop(key, None)
|
||||
|
||||
|
||||
def _promote_input_cache_lease(path: str) -> None:
|
||||
"""Turn a fetcher's lease into the active execution lease it returns."""
|
||||
key = _cache_path_key(path)
|
||||
with _INPUT_CACHE_LEASE_LOCK:
|
||||
remaining = _INPUT_CACHE_FETCH_LEASES.get(key, 0) - 1
|
||||
if remaining > 0:
|
||||
_INPUT_CACHE_FETCH_LEASES[key] = remaining
|
||||
else:
|
||||
_INPUT_CACHE_FETCH_LEASES.pop(key, None)
|
||||
|
||||
|
||||
def _leased_input_cache_paths() -> set[str]:
|
||||
with _INPUT_CACHE_LEASE_LOCK:
|
||||
return set(_INPUT_CACHE_LEASES)
|
||||
|
||||
|
||||
def _claim_input_cache_mutation(path: str) -> bool:
|
||||
key = _cache_path_key(path)
|
||||
with _INPUT_CACHE_LEASE_LOCK:
|
||||
if key in _INPUT_CACHE_MUTATIONS:
|
||||
return False
|
||||
active_leases = _INPUT_CACHE_LEASES.get(
|
||||
key, 0
|
||||
) - _INPUT_CACHE_FETCH_LEASES.get(key, 0)
|
||||
# Fetchers can safely converge under the publication gate. A lease
|
||||
# already promoted to an execution may have this exact pathname open.
|
||||
if active_leases > 0:
|
||||
return False
|
||||
_INPUT_CACHE_MUTATIONS.add(key)
|
||||
return True
|
||||
|
||||
|
||||
def _finish_input_cache_mutation(path: str) -> None:
|
||||
key = _cache_path_key(path)
|
||||
with _INPUT_CACHE_LEASE_LOCK:
|
||||
_INPUT_CACHE_MUTATIONS.discard(key)
|
||||
|
||||
|
||||
def _fsync_file(path: str) -> None:
|
||||
with open(path, "r+b") as handle:
|
||||
os.fsync(handle.fileno())
|
||||
|
||||
|
||||
def _fsync_parent_directory(directory: str) -> None:
|
||||
directory_flag = getattr(os, "O_DIRECTORY", None)
|
||||
if directory_flag is None:
|
||||
return
|
||||
unsupported = {
|
||||
errno.EINVAL,
|
||||
getattr(errno, "ENOTSUP", errno.EINVAL),
|
||||
getattr(errno, "EOPNOTSUPP", errno.EINVAL),
|
||||
}
|
||||
try:
|
||||
descriptor = os.open(directory, os.O_RDONLY | directory_flag)
|
||||
except OSError as exc:
|
||||
if exc.errno in unsupported:
|
||||
return
|
||||
raise
|
||||
try:
|
||||
os.fsync(descriptor)
|
||||
except OSError as exc:
|
||||
if exc.errno not in unsupported:
|
||||
raise
|
||||
finally:
|
||||
os.close(descriptor)
|
||||
|
||||
|
||||
def _durable_makedirs(directory: str) -> None:
|
||||
target = os.path.abspath(directory)
|
||||
missing: list[str] = []
|
||||
current = target
|
||||
while not os.path.isdir(current):
|
||||
if os.path.exists(current):
|
||||
if os.path.isdir(current):
|
||||
break
|
||||
raise NotADirectoryError(current)
|
||||
missing.append(current)
|
||||
parent = os.path.dirname(current)
|
||||
if parent == current:
|
||||
break
|
||||
current = parent
|
||||
for path in reversed(missing):
|
||||
try:
|
||||
os.mkdir(path)
|
||||
except FileExistsError:
|
||||
if not os.path.isdir(path):
|
||||
raise
|
||||
_fsync_parent_directory(os.path.dirname(path) or ".")
|
||||
if not missing:
|
||||
_fsync_parent_directory(os.path.dirname(target) or ".")
|
||||
|
||||
|
||||
def _durable_replace(source: str, destination: str) -> None:
|
||||
_fsync_file(source)
|
||||
os.replace(source, destination)
|
||||
_fsync_parent_directory(os.path.dirname(destination) or ".")
|
||||
|
||||
|
||||
def _input_fetch_failure(ref, error: BaseException) -> TaskFailure:
|
||||
return TaskFailure(
|
||||
WorkerError(
|
||||
# Transient on purpose: an id we cannot resolve now is far more
|
||||
# often a dropped stream/disk barrier than a permanently missing
|
||||
# file, and one wasted retry beats failing real work.
|
||||
error_class=ErrorClass.TRANSIENT,
|
||||
code="INPUT_FETCH_FAILED",
|
||||
message=f"Could not fetch '{ref.filename or ref.artifact_id}': {error}",
|
||||
hint="The control plane may have restarted; the task will be retried.",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _cache_name(ref) -> str:
|
||||
"""A safe, content-addressed local name for one input.
|
||||
|
||||
@@ -656,13 +1050,25 @@ def _cache_name(ref) -> str:
|
||||
def _already_held(path: str, ref) -> bool:
|
||||
"""Do we already have this exact input?
|
||||
|
||||
Size alone: the name is the content hash and the only writer is an atomic
|
||||
rename, so a file of the right size at this name cannot be different bytes.
|
||||
The filename is content-addressed, but disks and external edits can still
|
||||
change bytes at that name. Re-hash the advertised identity before reuse.
|
||||
"""
|
||||
try:
|
||||
expected = int(getattr(ref, "size_bytes", 0) or 0)
|
||||
return os.path.isfile(path) and (not expected or os.path.getsize(path) == expected)
|
||||
except OSError: # pragma: no cover
|
||||
if not os.path.isfile(path):
|
||||
return False
|
||||
if expected and os.path.getsize(path) != expected:
|
||||
return False
|
||||
expected_hash = (getattr(ref, "sha256", "") or "").strip().lower()
|
||||
if expected_hash:
|
||||
digest = hashlib.sha256()
|
||||
with open(path, "rb") as handle:
|
||||
for block in iter(lambda: handle.read(1024 * 1024), b""):
|
||||
digest.update(block)
|
||||
if digest.hexdigest() != expected_hash:
|
||||
return False
|
||||
return True
|
||||
except OSError:
|
||||
return False
|
||||
|
||||
|
||||
@@ -723,22 +1129,45 @@ def _verify(path: str, ref) -> None:
|
||||
)
|
||||
|
||||
|
||||
def _prune_input_cache(directory: str, limit_bytes: int = INPUT_CACHE_LIMIT_BYTES) -> None:
|
||||
"""Keep the input cache under its ceiling, oldest first."""
|
||||
def _prune_input_cache(
|
||||
directory: str,
|
||||
limit_bytes: int = INPUT_CACHE_LIMIT_BYTES,
|
||||
now: Optional[float] = None,
|
||||
) -> None:
|
||||
"""Keep the cache bounded without deleting inputs a task is still using."""
|
||||
try:
|
||||
entries = []
|
||||
total = 0
|
||||
stamp = time.time() if now is None else now
|
||||
leased = _leased_input_cache_paths()
|
||||
for name in os.listdir(directory):
|
||||
path = os.path.join(directory, name)
|
||||
if name.endswith(".part") or not os.path.isfile(path):
|
||||
if not os.path.isfile(path):
|
||||
continue
|
||||
stat = os.stat(path)
|
||||
entries.append((stat.st_mtime, stat.st_size, path))
|
||||
key = _cache_path_key(path)
|
||||
is_partial = name.endswith(".part")
|
||||
if (
|
||||
is_partial
|
||||
and key not in leased
|
||||
and stamp - stat.st_mtime >= _STALE_INPUT_PARTIAL_SECONDS
|
||||
):
|
||||
os.remove(path)
|
||||
continue
|
||||
total += stat.st_size
|
||||
# Active finals and partial transfers count toward the ceiling but
|
||||
# cannot be evicted. Young unleased .part files may belong to a
|
||||
# process that has not yet rebuilt its in-memory lease after fork;
|
||||
# the age sweep will remove them if they are crash leftovers.
|
||||
if key not in leased and not is_partial:
|
||||
entries.append((stat.st_mtime, stat.st_size, path))
|
||||
for _mtime, size, path in sorted(entries):
|
||||
if total <= limit_bytes:
|
||||
break
|
||||
os.remove(path)
|
||||
try:
|
||||
os.remove(path)
|
||||
except FileNotFoundError:
|
||||
continue
|
||||
total -= size
|
||||
except OSError: # pragma: no cover — a full cache is not a failed task
|
||||
logger.debug("Could not prune the worker input cache", exc_info=True)
|
||||
|
||||
@@ -25,6 +25,7 @@ in the dialog that shows it once.
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import errno
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
@@ -301,10 +302,31 @@ def save_worker_key(path: str, keypair: WorkerKeypair) -> None:
|
||||
tmp = f"{path}.tmp"
|
||||
fd = os.open(tmp, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
|
||||
try:
|
||||
os.write(fd, keypair.private_bytes())
|
||||
finally:
|
||||
remaining = memoryview(keypair.private_bytes())
|
||||
while remaining:
|
||||
written = os.write(fd, remaining)
|
||||
if written <= 0:
|
||||
raise OSError("could not finish writing the worker identity key")
|
||||
remaining = remaining[written:]
|
||||
os.fsync(fd)
|
||||
except Exception:
|
||||
os.close(fd)
|
||||
os.replace(tmp, path)
|
||||
try:
|
||||
os.unlink(tmp)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
raise
|
||||
else:
|
||||
os.close(fd)
|
||||
try:
|
||||
os.replace(tmp, path)
|
||||
except Exception:
|
||||
try:
|
||||
os.unlink(tmp)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
raise
|
||||
_fsync_parent_directory(directory)
|
||||
try:
|
||||
os.chmod(path, 0o600)
|
||||
except OSError:
|
||||
@@ -313,6 +335,30 @@ def save_worker_key(path: str, keypair: WorkerKeypair) -> None:
|
||||
pass
|
||||
|
||||
|
||||
def _fsync_parent_directory(directory: str) -> None:
|
||||
directory_flag = getattr(os, "O_DIRECTORY", None)
|
||||
if directory_flag is None:
|
||||
return
|
||||
unsupported = {
|
||||
errno.EINVAL,
|
||||
getattr(errno, "ENOTSUP", errno.EINVAL),
|
||||
getattr(errno, "EOPNOTSUPP", errno.EINVAL),
|
||||
}
|
||||
try:
|
||||
descriptor = os.open(directory, os.O_RDONLY | directory_flag)
|
||||
except OSError as exc:
|
||||
if exc.errno in unsupported:
|
||||
return
|
||||
raise
|
||||
try:
|
||||
os.fsync(descriptor)
|
||||
except OSError as exc:
|
||||
if exc.errno not in unsupported:
|
||||
raise
|
||||
finally:
|
||||
os.close(descriptor)
|
||||
|
||||
|
||||
def load_worker_key(path: str) -> Optional[WorkerKeypair]:
|
||||
try:
|
||||
with open(path, "rb") as fh:
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -22,20 +22,106 @@ import logging
|
||||
import os
|
||||
import socket
|
||||
import ssl
|
||||
from typing import Optional
|
||||
from typing import BinaryIO, Optional
|
||||
|
||||
import grpc
|
||||
|
||||
from worker import identity, registry, tls
|
||||
from worker.async_utils import to_thread_and_drain_on_cancel
|
||||
from worker.inbound.connection_string import Connection
|
||||
from worker.inbound.listener import KEY_METADATA_KEY
|
||||
from worker.protocol.gen import worker_v1_pb2 as pb
|
||||
from worker.protocol.gen import worker_v1_pb2_grpc as pb_grpc
|
||||
from worker.transport.client import MAX_MESSAGE_BYTES, backoff_delay
|
||||
from worker.transport.client import (
|
||||
MAX_MESSAGE_BYTES,
|
||||
TerminalRegistrationError,
|
||||
backoff_delay,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_PUSH_CHUNK_BYTES = 1024 * 1024
|
||||
# Register remains provisional until the node confirms that it durably saved
|
||||
# the panel-assigned identity. Match the control plane's provisional-session
|
||||
# lifetime so a peer that stops after Register cannot strand this connector.
|
||||
_REGISTRATION_CONFIRMATION_TIMEOUT_SECONDS = 30.0
|
||||
_REMOTE_SHUTDOWN_TIMEOUT_SECONDS = 30.0
|
||||
|
||||
_FileVersion = tuple[int, int, int, int, int]
|
||||
|
||||
|
||||
def _write_all(handle: BinaryIO, payload: bytes) -> None:
|
||||
remaining = memoryview(payload)
|
||||
while remaining:
|
||||
written = handle.write(remaining)
|
||||
if not written:
|
||||
raise OSError("result write made no progress")
|
||||
remaining = remaining[written:]
|
||||
|
||||
|
||||
def _remove_quietly(path: str) -> None:
|
||||
with contextlib.suppress(OSError):
|
||||
os.remove(path)
|
||||
|
||||
|
||||
class InboundConnectionError(RuntimeError):
|
||||
"""A pasted inbound connection could not be validated or activated."""
|
||||
|
||||
|
||||
class InboundConnectionRollbackError(InboundConnectionError):
|
||||
"""A failed connection change could not restore its prior generation."""
|
||||
|
||||
|
||||
class RemoteShutdownUnavailable(InboundConnectionError):
|
||||
"""The node may retain work, but no live stream can revoke it safely."""
|
||||
|
||||
|
||||
def _file_version(stat: os.stat_result) -> _FileVersion:
|
||||
"""Fields that identify both a staged path and the bytes hashed from it."""
|
||||
return (
|
||||
int(stat.st_dev),
|
||||
int(stat.st_ino),
|
||||
int(stat.st_size),
|
||||
int(stat.st_mtime_ns),
|
||||
int(stat.st_ctime_ns),
|
||||
)
|
||||
|
||||
|
||||
def _hash_staged_input(path: str) -> tuple[int, str, _FileVersion]:
|
||||
"""Hash one stable generation without ever allocating the whole file."""
|
||||
digest = hashlib.sha256()
|
||||
received = 0
|
||||
with open(path, "rb") as handle:
|
||||
before = _file_version(os.fstat(handle.fileno()))
|
||||
while True:
|
||||
block = handle.read(_PUSH_CHUNK_BYTES)
|
||||
if not block:
|
||||
break
|
||||
received += len(block)
|
||||
digest.update(block)
|
||||
after = _file_version(os.fstat(handle.fileno()))
|
||||
if before != after or received != before[2]:
|
||||
raise RuntimeError("the staged task input changed while it was being hashed")
|
||||
return received, digest.hexdigest(), before
|
||||
|
||||
|
||||
def _validate_staged_input(path: str, expected: _FileVersion) -> None:
|
||||
"""Reject a replacement or in-place edit between hashing and streaming."""
|
||||
try:
|
||||
current = _file_version(os.stat(path))
|
||||
except OSError as exc:
|
||||
raise RuntimeError("the staged task input is no longer available") from exc
|
||||
if current != expected:
|
||||
raise RuntimeError("the staged task input changed before it could be sent")
|
||||
|
||||
|
||||
def _validate_open_staged_input(
|
||||
handle: BinaryIO, path: str, expected: _FileVersion
|
||||
) -> None:
|
||||
"""The open generation and its path must still be the bytes we hashed."""
|
||||
if _file_version(os.fstat(handle.fileno())) != expected:
|
||||
raise RuntimeError("the staged task input changed before it could be sent")
|
||||
_validate_staged_input(path, expected)
|
||||
|
||||
|
||||
def _fetch_pinned_certificate(
|
||||
@@ -70,9 +156,15 @@ class NodeConnection:
|
||||
self._connection = connection
|
||||
self._label = label or connection.host
|
||||
self._outbox: asyncio.Queue[pb.ServerMessage] = asyncio.Queue()
|
||||
self._active_session = None
|
||||
self._stub: Optional[pb_grpc.NodeServiceStub] = None
|
||||
self._worker_id = ""
|
||||
self._stop = asyncio.Event()
|
||||
self._session_closed = asyncio.Event()
|
||||
self._session_closed.set()
|
||||
self._shutdown_confirmed = asyncio.Event()
|
||||
self._registration_ready = asyncio.Event()
|
||||
self._remote_protocol_retained = False
|
||||
self._last_error = ""
|
||||
|
||||
@property
|
||||
@@ -105,6 +197,10 @@ class NodeConnection:
|
||||
attempt = 0
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except TerminalRegistrationError as exc:
|
||||
self._remote_protocol_retained = False
|
||||
self._last_error = str(exc)
|
||||
raise
|
||||
except Exception:
|
||||
attempt += 1
|
||||
self._last_error = "Connection failed; check the backend log for details."
|
||||
@@ -114,8 +210,143 @@ class NodeConnection:
|
||||
await asyncio.wait_for(self._stop.wait(), timeout=delay)
|
||||
|
||||
async def stop(self) -> None:
|
||||
if self._stop.is_set():
|
||||
return
|
||||
if self._shutdown_confirmed.is_set() and not self._remote_protocol_retained:
|
||||
self._stop.set()
|
||||
return
|
||||
if self._active_session is None:
|
||||
if self._remote_protocol_retained:
|
||||
raise RemoteShutdownUnavailable(
|
||||
"That GPU machine is offline and may still be running work. "
|
||||
"Reconnect it, then remove the connection again."
|
||||
)
|
||||
self._stop.set()
|
||||
return
|
||||
# EOF is indistinguishable from a network blip and deliberately keeps
|
||||
# node execution alive for reconnect. Send an explicit terminal frame
|
||||
# and wait for the node to drain before removal reports success.
|
||||
self._shutdown_confirmed.clear()
|
||||
await self._outbox.put(
|
||||
pb.ServerMessage(
|
||||
shutdown=pb.Shutdown(reason="This GPU-machine connection was removed.")
|
||||
)
|
||||
)
|
||||
confirmed = asyncio.create_task(self._shutdown_confirmed.wait())
|
||||
disconnected = asyncio.create_task(self._session_closed.wait())
|
||||
try:
|
||||
done, _pending = await asyncio.wait(
|
||||
{confirmed, disconnected},
|
||||
timeout=_REMOTE_SHUTDOWN_TIMEOUT_SECONDS,
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
if confirmed not in done and not self._shutdown_confirmed.is_set():
|
||||
raise RemoteShutdownUnavailable(
|
||||
"The GPU machine disconnected before it confirmed shutdown. "
|
||||
"Reconnect it, then remove the connection again."
|
||||
)
|
||||
finally:
|
||||
confirmed.cancel()
|
||||
disconnected.cancel()
|
||||
await asyncio.gather(confirmed, disconnected, return_exceptions=True)
|
||||
self._remote_protocol_retained = False
|
||||
self._stop.set()
|
||||
|
||||
async def close(self) -> None:
|
||||
"""End this process without revoking reconnectable remote work."""
|
||||
self._stop.set()
|
||||
|
||||
def confirm_remote_shutdown(self, session) -> None:
|
||||
if self._active_session is not session:
|
||||
return
|
||||
self._remote_protocol_retained = False
|
||||
self._shutdown_confirmed.set()
|
||||
|
||||
def confirm_registration(self, session) -> None:
|
||||
"""Publish readiness only after the shared servicer activated the session."""
|
||||
if self._active_session is session:
|
||||
self._registration_ready.set()
|
||||
|
||||
async def wait_until_registered(
|
||||
self, task: asyncio.Task, *, timeout: float = 30.0
|
||||
) -> None:
|
||||
"""Wait for activation or surface a terminal/background dial failure."""
|
||||
ready = asyncio.create_task(self._registration_ready.wait())
|
||||
try:
|
||||
done, _pending = await asyncio.wait(
|
||||
{ready, task}, timeout=timeout, return_when=asyncio.FIRST_COMPLETED
|
||||
)
|
||||
if ready in done:
|
||||
return
|
||||
if task in done:
|
||||
if task.cancelled():
|
||||
raise InboundConnectionError(
|
||||
"The GPU-machine connection stopped before it became ready."
|
||||
)
|
||||
exc = task.exception()
|
||||
if exc is not None:
|
||||
raise InboundConnectionError(str(exc)) from exc
|
||||
raise InboundConnectionError(
|
||||
"That GPU machine did not finish connecting in time."
|
||||
)
|
||||
finally:
|
||||
ready.cancel()
|
||||
await asyncio.gather(ready, return_exceptions=True)
|
||||
|
||||
async def probe(self) -> None:
|
||||
"""Authenticate a replacement paste without publishing a worker session."""
|
||||
try:
|
||||
certificate_pem = await asyncio.to_thread(
|
||||
_fetch_pinned_certificate, self._connection
|
||||
)
|
||||
async with self._channel(certificate_pem) as channel:
|
||||
stub = pb_grpc.NodeServiceStub(channel)
|
||||
metadata = ((KEY_METADATA_KEY, self._connection.secret),)
|
||||
stream = stub.Attach(self._outbound(), metadata=metadata)
|
||||
try:
|
||||
first = await asyncio.wait_for(
|
||||
stream.read(),
|
||||
timeout=_REGISTRATION_CONFIRMATION_TIMEOUT_SECONDS,
|
||||
)
|
||||
finally:
|
||||
stream.cancel()
|
||||
except InboundConnectionError:
|
||||
raise
|
||||
except grpc.aio.AioRpcError as exc:
|
||||
detail = exc.details() or "The GPU machine rejected this connection."
|
||||
raise InboundConnectionError(detail) from exc
|
||||
except asyncio.TimeoutError as exc:
|
||||
raise InboundConnectionError(
|
||||
"That GPU machine did not answer in time."
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
raise InboundConnectionError(str(exc)) from exc
|
||||
|
||||
if first == grpc.aio.EOF or first.WhichOneof("payload") != "register":
|
||||
raise InboundConnectionError(
|
||||
"That machine answered, but not as a VoiceStudio GPU node."
|
||||
)
|
||||
request = first.register
|
||||
validate = getattr(self._servicer, "validate_inbound_request", None)
|
||||
refusal = validate(request) if callable(validate) else None
|
||||
if refusal is not None and refusal.error.code:
|
||||
raise InboundConnectionError(
|
||||
f"{refusal.error.code}: {refusal.error.message}"
|
||||
)
|
||||
public_key = bytes(request.public_key)
|
||||
if len(public_key) != 32:
|
||||
raise InboundConnectionError("That machine sent no usable identity.")
|
||||
key_id = identity.key_id_for(public_key)
|
||||
if registry.is_revoked(key_id):
|
||||
raise InboundConnectionError(
|
||||
"This GPU machine was removed from this app. Add it again to use it."
|
||||
)
|
||||
known = registry.get_by_key_id(key_id)
|
||||
if known is not None and not self._proves_key_possession(request, known):
|
||||
raise InboundConnectionError(
|
||||
"That machine could not prove its saved identity."
|
||||
)
|
||||
|
||||
async def _connect_once(self) -> None:
|
||||
# A fresh outbox per attempt. The queue used to be built once and
|
||||
# reused, so anything a dying session left behind became the NEXT
|
||||
@@ -123,7 +354,10 @@ class NodeConnection:
|
||||
# registration it requires first, aborted the call, and the pair span
|
||||
# at full speed: on hardware this reached session epoch 2445 inside a
|
||||
# second, with the log reading "Locally aborted" over and over.
|
||||
self._worker_id = ""
|
||||
self._stub = None
|
||||
self._outbox = asyncio.Queue()
|
||||
self._active_session = None
|
||||
certificate_pem = await asyncio.to_thread(
|
||||
_fetch_pinned_certificate, self._connection
|
||||
)
|
||||
@@ -140,33 +374,98 @@ class NodeConnection:
|
||||
"That machine answered, but not as a VoiceStudio GPU node."
|
||||
)
|
||||
|
||||
response = self._register(first.register)
|
||||
response = await self._register(first.register)
|
||||
if response.error.code:
|
||||
# A refusal here is a decision, not a blip: the node is a
|
||||
# different machine than the one this key was trusted for, or
|
||||
# its version cannot work with ours. Reconnecting cannot fix
|
||||
# either, so surface it rather than looping.
|
||||
raise RuntimeError(f"{response.error.code}: {response.error.message}")
|
||||
# either. Deliver the verdict before surfacing it locally so
|
||||
# the node can retire work retained across the dead stream;
|
||||
# closing first strands that executor with nobody left able to
|
||||
# cancel it.
|
||||
await self._outbox.put(pb.ServerMessage(registered=response))
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
stream.read(),
|
||||
timeout=_REGISTRATION_CONFIRMATION_TIMEOUT_SECONDS,
|
||||
)
|
||||
except (asyncio.TimeoutError, grpc.aio.AioRpcError):
|
||||
pass
|
||||
raise TerminalRegistrationError(
|
||||
f"{response.error.code}: {response.error.message}"
|
||||
)
|
||||
|
||||
self._worker_id = response.worker_id
|
||||
self._stub = stub
|
||||
self._last_error = ""
|
||||
await self._outbox.put(pb.ServerMessage(registered=response))
|
||||
|
||||
session = self._servicer.session_for(self._worker_id)
|
||||
if session is None:
|
||||
raise RuntimeError("the session went away before the stream opened")
|
||||
|
||||
pump = asyncio.create_task(self._pump_outbound(session))
|
||||
try:
|
||||
await self._servicer.run_inbound_stream(session, _Frames(stream), self)
|
||||
await self._complete_registration(stream, response, stub)
|
||||
finally:
|
||||
pump.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError, Exception):
|
||||
await pump
|
||||
self._stub = None
|
||||
# Idempotent after activation; essential before it. A user can
|
||||
# remove this connection while the node is still persisting
|
||||
# identity, and cancellation must release the old worker's
|
||||
# scheduling gate immediately rather than wait for expiry.
|
||||
self._servicer.discard_unopened_session(
|
||||
response.worker_id, session_token=response.session_token
|
||||
)
|
||||
|
||||
def _register(self, request: pb.RegisterRequest) -> pb.RegisterResponse:
|
||||
async def _complete_registration(self, stream, response, stub) -> None:
|
||||
"""Validate durable acceptance, then run the exact issued session."""
|
||||
try:
|
||||
confirmation = await asyncio.wait_for(
|
||||
stream.read(), timeout=_REGISTRATION_CONFIRMATION_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError as exc:
|
||||
raise RuntimeError(
|
||||
"That GPU machine did not confirm registration in time."
|
||||
) from exc
|
||||
except grpc.aio.AioRpcError as exc:
|
||||
detail = exc.details() or ""
|
||||
error_code = detail.partition(":")[0].strip()
|
||||
if exc.code() == grpc.StatusCode.FAILED_PRECONDITION and error_code in {
|
||||
"AUTH_FAILED",
|
||||
"LOCAL_STATE",
|
||||
"UPGRADE_REQUIRED",
|
||||
}:
|
||||
raise TerminalRegistrationError(detail) from exc
|
||||
raise
|
||||
if confirmation == grpc.aio.EOF:
|
||||
raise RuntimeError(
|
||||
"That GPU machine disconnected before confirming registration."
|
||||
)
|
||||
if confirmation.WhichOneof("payload") != "heartbeat":
|
||||
raise RuntimeError(
|
||||
"That GPU machine sent an invalid registration confirmation."
|
||||
)
|
||||
|
||||
session = self._servicer.session_for(
|
||||
response.worker_id, session_token=response.session_token
|
||||
)
|
||||
if session is None:
|
||||
raise RuntimeError("the session went away before the stream opened")
|
||||
|
||||
self._worker_id = response.worker_id
|
||||
self._stub = stub
|
||||
self._last_error = ""
|
||||
self._active_session = session
|
||||
self._remote_protocol_retained = True
|
||||
self._shutdown_confirmed.clear()
|
||||
self._session_closed.clear()
|
||||
|
||||
pump = asyncio.create_task(self._pump_outbound(session))
|
||||
try:
|
||||
await self._servicer.run_inbound_stream(
|
||||
session, _Frames(stream, first=confirmation), self
|
||||
)
|
||||
finally:
|
||||
pump.cancel()
|
||||
with contextlib.suppress(asyncio.CancelledError, Exception):
|
||||
await pump
|
||||
self._stub = None
|
||||
self._worker_id = ""
|
||||
if self._active_session is session:
|
||||
self._active_session = None
|
||||
self._session_closed.set()
|
||||
|
||||
async def _register(self, request: pb.RegisterRequest) -> pb.RegisterResponse:
|
||||
"""Trust on first sight, then require the same key forever after.
|
||||
|
||||
Pasting the connection string is the consent — the user went to the
|
||||
@@ -175,14 +474,28 @@ class NodeConnection:
|
||||
is a licence for a different machine to answer at that address later,
|
||||
which is why the key is bound on first contact.
|
||||
"""
|
||||
refusal = self._servicer.validate_inbound_request(request)
|
||||
if refusal is not None:
|
||||
return refusal
|
||||
worker, refusal = await to_thread_and_drain_on_cancel(
|
||||
self._authenticate_registration, request
|
||||
)
|
||||
if refusal is not None:
|
||||
return refusal
|
||||
return await self._servicer.establish_session(
|
||||
worker, request, address=self._connection.endpoint
|
||||
)
|
||||
|
||||
def _authenticate_registration(self, request: pb.RegisterRequest):
|
||||
"""Resolve inbound identity without running SQLite on the app loop."""
|
||||
public_key = bytes(request.public_key)
|
||||
if len(public_key) != 32:
|
||||
return self._servicer._refuse(
|
||||
return None, self._servicer._refuse(
|
||||
"AUTH_FAILED", "That machine sent no usable identity."
|
||||
)
|
||||
key_id = identity.key_id_for(public_key)
|
||||
if registry.is_revoked(key_id):
|
||||
return self._servicer._refuse(
|
||||
return None, self._servicer._refuse(
|
||||
"AUTH_FAILED",
|
||||
"This GPU machine was removed from this app. Add it again to use it.",
|
||||
)
|
||||
@@ -219,14 +532,12 @@ class NodeConnection:
|
||||
)
|
||||
worker = known
|
||||
if worker is None:
|
||||
return self._servicer._refuse(
|
||||
return None, self._servicer._refuse(
|
||||
"AUTH_FAILED",
|
||||
"That machine could not prove it is the one this key was added for.",
|
||||
)
|
||||
|
||||
return self._servicer.register_inbound(
|
||||
worker, request, address=self._connection.endpoint
|
||||
)
|
||||
return worker, None
|
||||
|
||||
@staticmethod
|
||||
def _proves_key_possession(request: pb.RegisterRequest, known) -> bool:
|
||||
@@ -249,9 +560,34 @@ class NodeConnection:
|
||||
public_key, message, bytes(request.challenge_signature)
|
||||
)
|
||||
|
||||
async def _outbound(self):
|
||||
def _outbound(self):
|
||||
# grpc closes request iterators itself when the peer ends a stream. An
|
||||
# async generator can still be suspended in ``Queue.get`` at that
|
||||
# point, making its concurrent ``aclose`` fail and leak teardown into
|
||||
# the next channel. A plain async iterator has no generator-finalizer
|
||||
# race and keeps the same one-frame-at-a-time backpressure.
|
||||
return _OutboundFrames(self)
|
||||
|
||||
def fence_session_egress(self, session) -> None:
|
||||
"""Drop frames copied before a replacement generation activated."""
|
||||
if self._active_session is not session:
|
||||
return
|
||||
while True:
|
||||
yield await self._outbox.get()
|
||||
try:
|
||||
self._outbox.get_nowait()
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
|
||||
def revoke_session(self, session) -> None:
|
||||
"""Synchronously fence frames already copied into the request queue."""
|
||||
if self._active_session is not session:
|
||||
return
|
||||
while True:
|
||||
try:
|
||||
self._outbox.get_nowait()
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
self._outbox.put_nowait(None)
|
||||
|
||||
async def _pump_outbound(self, session) -> None:
|
||||
"""Move the servicer's per-session outbox onto the dialled stream.
|
||||
@@ -260,8 +596,18 @@ class NodeConnection:
|
||||
to cross into the request generator instead, because this side is the
|
||||
caller.
|
||||
"""
|
||||
while True:
|
||||
await self._outbox.put(await session.outbox.get())
|
||||
task = asyncio.current_task()
|
||||
if task is not None:
|
||||
session.egress_tasks.add(task)
|
||||
try:
|
||||
while not session.revoked and not getattr(session, "egress_fenced", False):
|
||||
message = await session.outbox.get()
|
||||
if session.revoked or getattr(session, "egress_fenced", False):
|
||||
return
|
||||
await self._outbox.put(message)
|
||||
finally:
|
||||
if task is not None:
|
||||
session.egress_tasks.discard(task)
|
||||
|
||||
# ── Artifacts ─────────────────────────────────────────────────────────
|
||||
|
||||
@@ -277,32 +623,52 @@ class NodeConnection:
|
||||
if stub is None:
|
||||
raise RuntimeError("that GPU machine is not connected")
|
||||
|
||||
size = os.path.getsize(path)
|
||||
digest = hashlib.sha256()
|
||||
with open(path, "rb") as handle:
|
||||
digest.update(handle.read())
|
||||
size, digest, version = await to_thread_and_drain_on_cancel(
|
||||
_hash_staged_input, path
|
||||
)
|
||||
# Hashing and the gRPC request are separate operations. Re-resolve the
|
||||
# path immediately before handing the iterator to gRPC so a replaced
|
||||
# staging file is never described by the old generation's digest.
|
||||
await to_thread_and_drain_on_cancel(_validate_staged_input, path, version)
|
||||
|
||||
declared = pb.ArtifactRef()
|
||||
declared.CopyFrom(ref)
|
||||
declared.size_bytes = size
|
||||
declared.sha256 = digest.hexdigest()
|
||||
declared.sha256 = digest
|
||||
if not declared.filename:
|
||||
declared.filename = os.path.basename(path)
|
||||
|
||||
async def chunks():
|
||||
offset = 0
|
||||
with open(path, "rb") as handle:
|
||||
while True:
|
||||
data = handle.read(_PUSH_CHUNK_BYTES)
|
||||
handle = await to_thread_and_drain_on_cancel(open, path, "rb")
|
||||
try:
|
||||
await to_thread_and_drain_on_cancel(
|
||||
_validate_open_staged_input, handle, path, version
|
||||
)
|
||||
while offset < size:
|
||||
data = await to_thread_and_drain_on_cancel(
|
||||
handle.read, min(_PUSH_CHUNK_BYTES, size - offset)
|
||||
)
|
||||
if not data:
|
||||
break
|
||||
raise RuntimeError(
|
||||
"the staged task input changed before it could be sent"
|
||||
)
|
||||
offset += len(data)
|
||||
last = offset == size
|
||||
if last:
|
||||
# Do not publish the terminal frame until both the open
|
||||
# generation and its path still match what was hashed.
|
||||
await to_thread_and_drain_on_cancel(
|
||||
_validate_open_staged_input, handle, path, version
|
||||
)
|
||||
yield pb.ArtifactChunk(
|
||||
ref=declared,
|
||||
offset=offset - len(data),
|
||||
data=data,
|
||||
last=offset >= size,
|
||||
last=last,
|
||||
)
|
||||
finally:
|
||||
await to_thread_and_drain_on_cancel(handle.close)
|
||||
|
||||
ack = await stub.PushInput(
|
||||
chunks(), metadata=((KEY_METADATA_KEY, self._connection.secret),)
|
||||
@@ -327,7 +693,11 @@ class NodeConnection:
|
||||
offset = 0
|
||||
complete = False
|
||||
try:
|
||||
with open(destination, "wb") as handle:
|
||||
handle = None
|
||||
try:
|
||||
handle = await to_thread_and_drain_on_cancel(
|
||||
open, destination, "wb"
|
||||
)
|
||||
async for chunk in stub.FetchResult(
|
||||
request, metadata=((KEY_METADATA_KEY, self._connection.secret),)
|
||||
):
|
||||
@@ -339,27 +709,31 @@ class NodeConnection:
|
||||
raise RuntimeError(
|
||||
"the result is larger than the control plane accepts"
|
||||
)
|
||||
handle.write(chunk.data)
|
||||
digest.update(chunk.data)
|
||||
offset += len(chunk.data)
|
||||
data = bytes(chunk.data)
|
||||
await to_thread_and_drain_on_cancel(_write_all, handle, data)
|
||||
digest.update(data)
|
||||
offset += len(data)
|
||||
if chunk.last:
|
||||
complete = True
|
||||
break
|
||||
finally:
|
||||
if handle is not None:
|
||||
await to_thread_and_drain_on_cancel(handle.close)
|
||||
except asyncio.CancelledError:
|
||||
await to_thread_and_drain_on_cancel(_remove_quietly, destination)
|
||||
raise
|
||||
except Exception:
|
||||
with contextlib.suppress(OSError):
|
||||
os.remove(destination)
|
||||
await to_thread_and_drain_on_cancel(_remove_quietly, destination)
|
||||
raise
|
||||
|
||||
# A truncated file that is renamed into place and called done is the
|
||||
# exact failure the upload path was hardened against; the pull
|
||||
# direction gets the same treatment.
|
||||
if not complete:
|
||||
with contextlib.suppress(OSError):
|
||||
os.remove(destination)
|
||||
await to_thread_and_drain_on_cancel(_remove_quietly, destination)
|
||||
raise RuntimeError("the result ended before its final chunk")
|
||||
if ref.sha256 and digest.hexdigest() != ref.sha256:
|
||||
with contextlib.suppress(OSError):
|
||||
os.remove(destination)
|
||||
await to_thread_and_drain_on_cancel(_remove_quietly, destination)
|
||||
raise RuntimeError(
|
||||
"the result did not match the checksum that machine declared"
|
||||
)
|
||||
@@ -368,14 +742,46 @@ class NodeConnection:
|
||||
class _Frames:
|
||||
"""Adapts a gRPC client stream to the ``async for`` the read loop expects."""
|
||||
|
||||
def __init__(self, stream) -> None:
|
||||
def __init__(self, stream, *, first=None) -> None:
|
||||
self._stream = stream
|
||||
self._first = first
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
if self._first is not None:
|
||||
message = self._first
|
||||
self._first = None
|
||||
return message
|
||||
message = await self._stream.read()
|
||||
if message == grpc.aio.EOF:
|
||||
raise StopAsyncIteration
|
||||
return message
|
||||
|
||||
|
||||
class _OutboundFrames:
|
||||
"""Cancellation-safe request iterator for the inverted Attach stream."""
|
||||
|
||||
def __init__(self, connection: NodeConnection) -> None:
|
||||
self._connection = connection
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
connection = self._connection
|
||||
while True:
|
||||
message = await connection._outbox.get()
|
||||
if message is None:
|
||||
raise StopAsyncIteration
|
||||
session = connection._active_session
|
||||
if session is not None:
|
||||
if session.revoked:
|
||||
raise StopAsyncIteration
|
||||
if (
|
||||
getattr(session, "egress_fenced", False)
|
||||
and message.WhichOneof("payload") != "shutdown"
|
||||
):
|
||||
continue
|
||||
return message
|
||||
|
||||
+192
-17
@@ -15,6 +15,7 @@ settings store.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import errno
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
@@ -45,6 +46,20 @@ _MAX_FAILURES = 5
|
||||
_LOCKOUT_SECONDS = 60.0
|
||||
_FAILURE_WINDOW_SECONDS = 300.0
|
||||
|
||||
# ``Attach`` is the only RPC that records presence. Persisting on every
|
||||
# reconnect lets an authenticated peer turn harmless telemetry into an fsync
|
||||
# storm on the gRPC event loop, so coalesce it to a useful reporting cadence.
|
||||
_LAST_SEEN_PERSIST_INTERVAL_SECONDS = 60.0
|
||||
|
||||
# Authentication deliberately scans every stored hash in constant time. Keep
|
||||
# that work and the JSON credential file bounded even if an administrator
|
||||
# repeatedly issues replacements.
|
||||
MAX_PANEL_KEYS = 256
|
||||
|
||||
|
||||
class KeyLimitExceeded(RuntimeError):
|
||||
"""No additional panel credential can be retained safely."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class PanelKey:
|
||||
@@ -89,6 +104,31 @@ def _peer_host(peer: str) -> str:
|
||||
return peer
|
||||
|
||||
|
||||
def _fsync_parent_directory(directory: str) -> None:
|
||||
"""Make a preceding directory-entry replacement durable when supported."""
|
||||
directory_flag = getattr(os, "O_DIRECTORY", None)
|
||||
if directory_flag is None:
|
||||
return
|
||||
unsupported = {
|
||||
errno.EINVAL,
|
||||
getattr(errno, "ENOTSUP", errno.EINVAL),
|
||||
getattr(errno, "EOPNOTSUPP", errno.EINVAL),
|
||||
}
|
||||
try:
|
||||
descriptor = os.open(directory, os.O_RDONLY | directory_flag)
|
||||
except OSError as exc:
|
||||
if exc.errno in unsupported:
|
||||
return
|
||||
raise
|
||||
try:
|
||||
os.fsync(descriptor)
|
||||
except OSError as exc:
|
||||
if exc.errno not in unsupported:
|
||||
raise
|
||||
finally:
|
||||
os.close(descriptor)
|
||||
|
||||
|
||||
@dataclass
|
||||
class IssuedKey:
|
||||
"""The one and only time the plaintext exists outside the caller's hands."""
|
||||
@@ -113,6 +153,10 @@ class KeyStore:
|
||||
self._connection_secrets: dict[str, str] = {}
|
||||
self._connection_fingerprints: dict[str, str] = {}
|
||||
self._failures: dict[str, _Failures] = {}
|
||||
# A failed persistence attempt must remain denied in this process but
|
||||
# still be retryable. Keeping this separate from PanelKey.revoked lets
|
||||
# the next DELETE attempt write the durable transition again.
|
||||
self._pending_revocations: set[str] = set()
|
||||
self._load()
|
||||
|
||||
# ── Persistence ───────────────────────────────────────────────────────
|
||||
@@ -170,10 +214,31 @@ class KeyStore:
|
||||
# `identity.save_worker_key` uses for the Ed25519 private key.
|
||||
fd = os.open(tmp, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
|
||||
try:
|
||||
os.write(fd, payload)
|
||||
finally:
|
||||
remaining = memoryview(payload)
|
||||
while remaining:
|
||||
written = os.write(fd, remaining)
|
||||
if written <= 0:
|
||||
raise OSError("could not finish writing the inbound key file")
|
||||
remaining = remaining[written:]
|
||||
os.fsync(fd)
|
||||
except Exception:
|
||||
os.close(fd)
|
||||
os.replace(tmp, self._path)
|
||||
try:
|
||||
os.unlink(tmp)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
raise
|
||||
else:
|
||||
os.close(fd)
|
||||
try:
|
||||
os.replace(tmp, self._path)
|
||||
except Exception:
|
||||
try:
|
||||
os.unlink(tmp)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
raise
|
||||
_fsync_parent_directory(directory)
|
||||
try:
|
||||
os.chmod(self._path, 0o600)
|
||||
except OSError:
|
||||
@@ -196,8 +261,37 @@ class KeyStore:
|
||||
created_at=now,
|
||||
)
|
||||
with self._lock:
|
||||
previous = self._keys.get(key.key_id)
|
||||
pruned: dict[str, PanelKey] = {}
|
||||
if previous is None and len(self._keys) >= MAX_PANEL_KEYS:
|
||||
revoked = sorted(
|
||||
(
|
||||
stored
|
||||
for stored in self._keys.values()
|
||||
if stored.revoked
|
||||
and stored.key_id not in self._pending_revocations
|
||||
),
|
||||
key=lambda stored: stored.created_at,
|
||||
)
|
||||
while len(self._keys) >= MAX_PANEL_KEYS and revoked:
|
||||
stale = revoked.pop(0)
|
||||
pruned[stale.key_id] = self._keys.pop(stale.key_id)
|
||||
if len(self._keys) >= MAX_PANEL_KEYS:
|
||||
self._keys.update(pruned)
|
||||
raise KeyLimitExceeded(
|
||||
"This GPU machine already has as many panel keys as it accepts. "
|
||||
"Revoke an unused key, then try again."
|
||||
)
|
||||
self._keys[key.key_id] = key
|
||||
self._save_locked()
|
||||
try:
|
||||
self._save_locked()
|
||||
except Exception:
|
||||
if previous is None:
|
||||
self._keys.pop(key.key_id, None)
|
||||
else:
|
||||
self._keys[key.key_id] = previous
|
||||
self._keys.update(pruned)
|
||||
raise
|
||||
return IssuedKey(key=key, secret=secret)
|
||||
|
||||
def revoke(self, key_id: str) -> bool:
|
||||
@@ -206,8 +300,14 @@ class KeyStore:
|
||||
key = self._keys.get(key_id)
|
||||
if key is None or key.revoked:
|
||||
return False
|
||||
self._pending_revocations.add(key_id)
|
||||
key.revoked = True
|
||||
self._save_locked()
|
||||
try:
|
||||
self._save_locked()
|
||||
except Exception:
|
||||
key.revoked = False
|
||||
raise
|
||||
self._pending_revocations.discard(key_id)
|
||||
return True
|
||||
|
||||
def remember_worker_id(self, key_id: str, worker_id: str) -> None:
|
||||
@@ -216,15 +316,46 @@ class KeyStore:
|
||||
return
|
||||
with self._lock:
|
||||
key = self._keys.get(key_id)
|
||||
if key is None or key.worker_id == worker_id:
|
||||
if (
|
||||
key is None
|
||||
or key.revoked
|
||||
or key_id in self._pending_revocations
|
||||
):
|
||||
raise PermissionError("the panel key was revoked during registration")
|
||||
if key.worker_id == worker_id:
|
||||
return
|
||||
previous_worker_id = key.worker_id
|
||||
key.worker_id = worker_id
|
||||
self._save_locked()
|
||||
try:
|
||||
self._save_locked()
|
||||
except Exception:
|
||||
# A callback retry must attempt the durable write again. If
|
||||
# the failed value remains in memory, the equality fast path
|
||||
# above accepts it as saved and the node reconnects with an id
|
||||
# that disappears on process restart.
|
||||
key.worker_id = previous_worker_id
|
||||
raise
|
||||
|
||||
def worker_id_for(self, key_id: str) -> str:
|
||||
with self._lock:
|
||||
key = self._keys.get(key_id)
|
||||
return key.worker_id if key is not None else ""
|
||||
return (
|
||||
key.worker_id
|
||||
if key is not None
|
||||
and not key.revoked
|
||||
and key_id not in self._pending_revocations
|
||||
else ""
|
||||
)
|
||||
|
||||
def is_active(self, key_id: str) -> bool:
|
||||
"""Whether this key still has authority to use an existing session."""
|
||||
with self._lock:
|
||||
key = self._keys.get(key_id)
|
||||
return (
|
||||
key is not None
|
||||
and not key.revoked
|
||||
and key_id not in self._pending_revocations
|
||||
)
|
||||
|
||||
def list_keys(self) -> list[dict]:
|
||||
with self._lock:
|
||||
@@ -232,7 +363,10 @@ class KeyStore:
|
||||
|
||||
def any_active(self) -> bool:
|
||||
with self._lock:
|
||||
return any(not k.revoked for k in self._keys.values())
|
||||
return any(
|
||||
not key.revoked and key.key_id not in self._pending_revocations
|
||||
for key in self._keys.values()
|
||||
)
|
||||
|
||||
# ── Panel-side connection credentials ───────────────────────────────
|
||||
|
||||
@@ -241,10 +375,23 @@ class KeyStore:
|
||||
) -> None:
|
||||
"""Persist a pasted node secret outside the UI-readable settings store."""
|
||||
with self._lock:
|
||||
previous_secret = self._connection_secrets.get(endpoint)
|
||||
previous_fingerprint = self._connection_fingerprints.get(endpoint)
|
||||
self._connection_secrets[endpoint] = secret
|
||||
if fingerprint:
|
||||
self._connection_fingerprints[endpoint] = fingerprint
|
||||
self._save_locked()
|
||||
try:
|
||||
self._save_locked()
|
||||
except Exception:
|
||||
if previous_secret is None:
|
||||
self._connection_secrets.pop(endpoint, None)
|
||||
else:
|
||||
self._connection_secrets[endpoint] = previous_secret
|
||||
if previous_fingerprint is None:
|
||||
self._connection_fingerprints.pop(endpoint, None)
|
||||
else:
|
||||
self._connection_fingerprints[endpoint] = previous_fingerprint
|
||||
raise
|
||||
|
||||
def connection_secret(self, endpoint: str) -> str:
|
||||
with self._lock:
|
||||
@@ -256,9 +403,19 @@ class KeyStore:
|
||||
|
||||
def forget_connection_secret(self, endpoint: str) -> None:
|
||||
with self._lock:
|
||||
if self._connection_secrets.pop(endpoint, None) is not None:
|
||||
self._connection_fingerprints.pop(endpoint, None)
|
||||
previous_secret = self._connection_secrets.get(endpoint)
|
||||
if previous_secret is None:
|
||||
return
|
||||
previous_fingerprint = self._connection_fingerprints.get(endpoint)
|
||||
self._connection_secrets.pop(endpoint, None)
|
||||
self._connection_fingerprints.pop(endpoint, None)
|
||||
try:
|
||||
self._save_locked()
|
||||
except Exception:
|
||||
self._connection_secrets[endpoint] = previous_secret
|
||||
if previous_fingerprint is not None:
|
||||
self._connection_fingerprints[endpoint] = previous_fingerprint
|
||||
raise
|
||||
|
||||
# ── Authentication ────────────────────────────────────────────────────
|
||||
|
||||
@@ -268,7 +425,9 @@ class KeyStore:
|
||||
record = self._failures.get(peer)
|
||||
return record is not None and record.locked_until > self._now()
|
||||
|
||||
def authenticate(self, secret: str, *, peer: str = "") -> Optional[PanelKey]:
|
||||
def authenticate(
|
||||
self, secret: str, *, peer: str = "", record_seen: bool = True
|
||||
) -> Optional[PanelKey]:
|
||||
"""Return the matching live key, or None.
|
||||
|
||||
Compares against every stored key in constant time and does not stop at
|
||||
@@ -286,7 +445,11 @@ class KeyStore:
|
||||
candidate = hash_secret(secret) if secret else ""
|
||||
matched: Optional[PanelKey] = None
|
||||
for key in self._keys.values():
|
||||
if key.revoked or not candidate:
|
||||
if (
|
||||
key.revoked
|
||||
or key.key_id in self._pending_revocations
|
||||
or not candidate
|
||||
):
|
||||
continue
|
||||
if constant_time_equals(key.secret_hash, candidate):
|
||||
matched = key
|
||||
@@ -296,9 +459,21 @@ class KeyStore:
|
||||
return None
|
||||
|
||||
self._failures.pop(peer_host, None)
|
||||
matched.last_seen_at = now
|
||||
matched.last_seen_peer = peer
|
||||
self._save_locked()
|
||||
if record_seen and (
|
||||
matched.last_seen_at <= 0.0
|
||||
or now - matched.last_seen_at
|
||||
>= _LAST_SEEN_PERSIST_INTERVAL_SECONDS
|
||||
):
|
||||
previous_at = matched.last_seen_at
|
||||
previous_peer = matched.last_seen_peer
|
||||
matched.last_seen_at = now
|
||||
matched.last_seen_peer = peer
|
||||
try:
|
||||
self._save_locked()
|
||||
except Exception:
|
||||
matched.last_seen_at = previous_at
|
||||
matched.last_seen_peer = previous_peer
|
||||
raise
|
||||
return matched
|
||||
|
||||
def _record_failure_locked(self, peer: str, now: float) -> None:
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -14,11 +14,13 @@ second box — which is why neither implies the other.
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import ipaddress
|
||||
import logging
|
||||
import os
|
||||
from typing import Optional
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
from worker.async_utils import drain_task, to_thread_and_drain_on_cancel
|
||||
from worker.inbound.artifacts import ArtifactStore, KeyedArtifactTransport
|
||||
from worker.inbound.connection_log import ConnectionLog
|
||||
from worker.inbound.connection_string import (
|
||||
@@ -27,6 +29,7 @@ from worker.inbound.connection_string import (
|
||||
format_connection,
|
||||
parse_connection,
|
||||
)
|
||||
from worker.inbound.connector import InboundConnectionRollbackError
|
||||
from worker.inbound.keys import KeyStore
|
||||
from worker.inbound.listener import DEFAULT_BIND, DEFAULT_PORT, NodeListener
|
||||
|
||||
@@ -38,6 +41,41 @@ _PORT_KEY = "inbound_node_port"
|
||||
_SAVED_KEY = "inbound_saved_nodes"
|
||||
|
||||
|
||||
async def _finish_rollback(rollback, *, description: str) -> None:
|
||||
"""Finish a lifecycle rollback even if its caller is cancelled again."""
|
||||
task = asyncio.create_task(rollback, name="inbound-connection-rollback")
|
||||
try:
|
||||
await asyncio.shield(task)
|
||||
except BaseException:
|
||||
await drain_task(task)
|
||||
if task.cancelled():
|
||||
raise InboundConnectionRollbackError(
|
||||
f"Could not {description}: rollback was cancelled."
|
||||
)
|
||||
return task.result()
|
||||
|
||||
|
||||
def _normalise_listener_host(value: str) -> str:
|
||||
"""Return the bare identity gRPC and X.509 expect for an IP literal."""
|
||||
candidate = (value or "").strip()
|
||||
inner = (
|
||||
candidate[1:-1]
|
||||
if len(candidate) >= 2
|
||||
and candidate.startswith("[")
|
||||
and candidate.endswith("]")
|
||||
else candidate
|
||||
)
|
||||
try:
|
||||
return str(ipaddress.ip_address(inner))
|
||||
except ValueError:
|
||||
return candidate
|
||||
|
||||
|
||||
def normalise_bind_host(value: str) -> str:
|
||||
"""Canonicalise a requested listener host before comparing or saving it."""
|
||||
return _normalise_listener_host(value)
|
||||
|
||||
|
||||
def _setting(name: str, default: str = "") -> str:
|
||||
try:
|
||||
from services import settings_store # noqa: PLC0415
|
||||
@@ -86,13 +124,15 @@ def bind_host() -> str:
|
||||
point. Reaching this node from another machine should be a decision
|
||||
somebody made, not a side effect of turning the feature on.
|
||||
"""
|
||||
return (
|
||||
os.environ.get("OMNIVOICE_INBOUND_BIND") or _setting(_BIND_KEY) or DEFAULT_BIND
|
||||
return _normalise_listener_host(
|
||||
os.environ.get("OMNIVOICE_INBOUND_BIND")
|
||||
or _setting(_BIND_KEY)
|
||||
or DEFAULT_BIND
|
||||
)
|
||||
|
||||
|
||||
def set_bind_host(value: str) -> None:
|
||||
_set_setting(_BIND_KEY, (value or "").strip() or DEFAULT_BIND)
|
||||
_set_setting(_BIND_KEY, _normalise_listener_host(value) or DEFAULT_BIND)
|
||||
|
||||
|
||||
def bind_port() -> int:
|
||||
@@ -148,7 +188,9 @@ def is_exposed(host: Optional[str] = None) -> bool:
|
||||
connection string then admits clients beyond this machine. Transport
|
||||
remains pinned TLS (docs/adr/inbound-node-mode.md).
|
||||
"""
|
||||
return (host if host is not None else bind_host()) not in (
|
||||
return _normalise_listener_host(
|
||||
host if host is not None else bind_host()
|
||||
).lower() not in (
|
||||
"127.0.0.1",
|
||||
"localhost",
|
||||
"::1",
|
||||
@@ -189,6 +231,7 @@ class InboundNode:
|
||||
self._keys: Optional[KeyStore] = None
|
||||
self._log = ConnectionLog()
|
||||
self._idle_sweep: Optional[asyncio.Task] = None
|
||||
self._lifecycle_lock = asyncio.Lock()
|
||||
self.startup_error: Optional[str] = None
|
||||
|
||||
@property
|
||||
@@ -209,18 +252,12 @@ class InboundNode:
|
||||
def port(self) -> int:
|
||||
return self._listener.port if self._listener else 0
|
||||
|
||||
def _client_factory(self, artifacts: KeyedArtifactTransport, key_id: str):
|
||||
# Imported here so a machine that never accepts connections does not
|
||||
# pay for the executor or grpc at startup.
|
||||
def _prepare_client(self, key_id: str) -> dict:
|
||||
"""Probe keys, host and accelerators away from the listener loop."""
|
||||
from worker import capabilities # noqa: PLC0415
|
||||
from worker.agent import _paths as agent_paths # noqa: PLC0415
|
||||
from worker.executor import TaskExecutor # noqa: PLC0415
|
||||
from worker.identity import load_or_create_worker_key # noqa: PLC0415
|
||||
from worker.transport.client import ( # noqa: PLC0415
|
||||
WorkerClient,
|
||||
WorkerConfig,
|
||||
describe_host,
|
||||
)
|
||||
from worker.transport.client import describe_host # noqa: PLC0415
|
||||
|
||||
locations = agent_paths()
|
||||
os.makedirs(locations["root"], exist_ok=True)
|
||||
@@ -228,6 +265,29 @@ class InboundNode:
|
||||
discovered = capabilities.discover(include_unavailable=True)
|
||||
host = describe_host()
|
||||
host["gpus"] = capabilities.describe_gpus()
|
||||
return {
|
||||
"keypair": keypair,
|
||||
"discovered": discovered,
|
||||
"host": host,
|
||||
"worker_id": self.keys.worker_id_for(key_id),
|
||||
"max_concurrent_tasks": capabilities.max_concurrent_tasks(discovered),
|
||||
}
|
||||
|
||||
async def _client_factory(
|
||||
self, artifacts: KeyedArtifactTransport, key_id: str
|
||||
):
|
||||
# Imported here so a machine that never accepts connections does not
|
||||
# pay for the executor or grpc at startup.
|
||||
from worker import capabilities # noqa: PLC0415
|
||||
from worker.executor import TaskExecutor # noqa: PLC0415
|
||||
from worker.transport.client import ( # noqa: PLC0415
|
||||
WorkerClient,
|
||||
WorkerConfig,
|
||||
)
|
||||
|
||||
prepared = await to_thread_and_drain_on_cancel(
|
||||
self._prepare_client, key_id
|
||||
)
|
||||
|
||||
executor = TaskExecutor()
|
||||
return WorkerClient(
|
||||
@@ -235,23 +295,28 @@ class InboundNode:
|
||||
endpoint="",
|
||||
cert_fingerprint="",
|
||||
certificate_pem=b"",
|
||||
keypair=keypair,
|
||||
keypair=prepared["keypair"],
|
||||
# Per panel key, not per node: each panel keeps its own
|
||||
# registry, so the same machine is a different worker id to
|
||||
# each of them, and the node signs its challenge over that id.
|
||||
worker_id=self.keys.worker_id_for(key_id),
|
||||
worker_id=prepared["worker_id"],
|
||||
enrollment_token="",
|
||||
max_concurrent_tasks=capabilities.max_concurrent_tasks(discovered),
|
||||
capabilities=discovered,
|
||||
host=host,
|
||||
max_concurrent_tasks=prepared["max_concurrent_tasks"],
|
||||
capabilities=prepared["discovered"],
|
||||
host=prepared["host"],
|
||||
),
|
||||
execute=executor.execute,
|
||||
capability_probe=lambda: capabilities.discover(include_unavailable=True),
|
||||
on_registered=lambda wid: self.keys.remember_worker_id(key_id, wid),
|
||||
artifacts=artifacts,
|
||||
drain_active_work=executor.drain_active_work,
|
||||
)
|
||||
|
||||
async def start(self) -> None:
|
||||
async with self._lifecycle_lock:
|
||||
await self._start()
|
||||
|
||||
async def _start(self) -> None:
|
||||
if self._listener is not None:
|
||||
return
|
||||
self.startup_error = None
|
||||
@@ -264,9 +329,18 @@ class InboundNode:
|
||||
)
|
||||
try:
|
||||
await listener.start(host=bind_host(), port=bind_port())
|
||||
except asyncio.CancelledError:
|
||||
# NodeListener cleans a partially bound server before returning.
|
||||
# If that cleanup itself failed it retains the handle; publish it
|
||||
# here so a later stop can retry rather than losing a live socket.
|
||||
if listener.running:
|
||||
self._listener = listener
|
||||
raise
|
||||
except Exception as exc:
|
||||
# A node that cannot listen must say so in the UI rather than look
|
||||
# enabled and quietly accept nothing.
|
||||
if listener.running:
|
||||
self._listener = listener
|
||||
self.startup_error = str(exc)
|
||||
logger.error("Could not start the inbound listener: %s", exc)
|
||||
return
|
||||
@@ -282,12 +356,48 @@ class InboundNode:
|
||||
)
|
||||
|
||||
async def stop(self) -> None:
|
||||
sweep, self._idle_sweep = self._idle_sweep, None
|
||||
if sweep is not None:
|
||||
sweep.cancel()
|
||||
listener, self._listener = self._listener, None
|
||||
if listener is not None:
|
||||
await listener.stop()
|
||||
async with self._lifecycle_lock:
|
||||
await self._stop()
|
||||
|
||||
async def _stop(self) -> None:
|
||||
sweep = self._idle_sweep
|
||||
listener = self._listener
|
||||
|
||||
async def shutdown() -> None:
|
||||
if sweep is not None:
|
||||
sweep.cancel()
|
||||
await asyncio.gather(sweep, return_exceptions=True)
|
||||
if listener is not None:
|
||||
await listener.stop()
|
||||
|
||||
stopping = asyncio.create_task(shutdown(), name="inbound-node-stop")
|
||||
try:
|
||||
await asyncio.shield(stopping)
|
||||
except asyncio.CancelledError:
|
||||
await drain_task(stopping)
|
||||
if stopping.cancelled():
|
||||
raise
|
||||
failure = stopping.exception()
|
||||
if failure is not None:
|
||||
raise failure
|
||||
if self._idle_sweep is sweep:
|
||||
self._idle_sweep = None
|
||||
if self._listener is listener:
|
||||
self._listener = None
|
||||
raise
|
||||
except BaseException:
|
||||
await drain_task(stopping)
|
||||
raise
|
||||
if self._idle_sweep is sweep:
|
||||
self._idle_sweep = None
|
||||
if self._listener is listener:
|
||||
self._listener = None
|
||||
|
||||
async def revoke_key(self, key_id: str) -> bool:
|
||||
"""Durably revoke one panel and withdraw all of its live sessions."""
|
||||
if self._listener is not None:
|
||||
return await self._listener.revoke_key_and_wait(key_id)
|
||||
return self.keys.revoke(key_id)
|
||||
|
||||
def connection_string(self, secret: str, *, host: Optional[str] = None) -> str:
|
||||
"""The one artifact a user copies to another machine.
|
||||
@@ -330,6 +440,8 @@ class OutboundNodes:
|
||||
self._connections: dict[str, object] = {}
|
||||
self._tasks: dict[str, asyncio.Task] = {}
|
||||
self._credentials = credentials
|
||||
self._lifecycle_lock = asyncio.Lock()
|
||||
self._servicer = None
|
||||
|
||||
@property
|
||||
def credentials(self) -> KeyStore:
|
||||
@@ -367,67 +479,337 @@ class OutboundNodes:
|
||||
async def add(self, text: str, servicer) -> Connection:
|
||||
"""Parse, save and dial. Raises InvalidConnectionString on a bad paste."""
|
||||
connection = parse_connection(text)
|
||||
entries = self.saved()
|
||||
# Keyed by endpoint: re-pasting a rotated key for the same machine
|
||||
# replaces it rather than leaving a dead entry that retries forever.
|
||||
entries = [e for e in entries if _endpoint_of(e) != connection.endpoint]
|
||||
entries.append(connection.endpoint)
|
||||
self.credentials.remember_connection_secret(
|
||||
connection.endpoint, connection.secret, connection.fingerprint
|
||||
)
|
||||
self._save(entries)
|
||||
async with self._lifecycle_lock:
|
||||
return await self._add(connection, servicer)
|
||||
|
||||
# Tear down any live session to this machine BEFORE dialling. Without
|
||||
# this, re-pasting for an already-connected machine saved the new key
|
||||
# and then short-circuited on the existing connection — so a wrong key
|
||||
# reported success, kept working on the old session, and only failed
|
||||
# after a restart, by which time nothing pointed at the paste that
|
||||
# caused it. Verified on hardware.
|
||||
await self._drop(connection.endpoint)
|
||||
await self._dial(connection, servicer)
|
||||
return connection
|
||||
|
||||
async def _drop(self, endpoint: str) -> None:
|
||||
connection = self._connections.pop(endpoint, None)
|
||||
task = self._tasks.pop(endpoint, None)
|
||||
if connection is not None:
|
||||
await connection.stop()
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
|
||||
async def remove(self, endpoint: str) -> bool:
|
||||
entries = [e for e in self.saved() if _endpoint_of(e) != endpoint]
|
||||
self._save(entries)
|
||||
self.credentials.forget_connection_secret(endpoint)
|
||||
existed = endpoint in self._connections
|
||||
await self._drop(endpoint)
|
||||
return existed
|
||||
|
||||
async def start_all(self, servicer) -> None:
|
||||
for entry in self.saved():
|
||||
try:
|
||||
await self._dial(self._connection_for(entry), servicer)
|
||||
except InvalidConnectionString as exc:
|
||||
logger.warning(
|
||||
"Ignoring a saved connection that no longer parses: %s", exc
|
||||
)
|
||||
|
||||
async def _dial(self, connection: Connection, servicer) -> None:
|
||||
async def _add(self, connection: Connection, servicer) -> Connection:
|
||||
from worker.inbound.connector import NodeConnection # noqa: PLC0415
|
||||
|
||||
if connection.endpoint in self._connections:
|
||||
original_entries = self.saved()
|
||||
old_secret = self.credentials.connection_secret(connection.endpoint)
|
||||
old_fingerprint = self.credentials.connection_fingerprint(connection.endpoint)
|
||||
existing = self._connections.get(connection.endpoint)
|
||||
existing_task = self._tasks.get(connection.endpoint)
|
||||
|
||||
# Pasting the already-running key is idempotent. Probing it would be a
|
||||
# duplicate Attach and could disturb state deliberately retained by
|
||||
# that same key.
|
||||
if (
|
||||
existing is not None
|
||||
and old_secret == connection.secret
|
||||
and old_fingerprint == connection.fingerprint
|
||||
):
|
||||
if existing_task is not None and not existing_task.done():
|
||||
return connection
|
||||
# Terminal registration failures leave their diagnostic connector
|
||||
# in the snapshot. Re-pasting after an upgrade/repair must really
|
||||
# redial, not mistake that dead object for a healthy connection.
|
||||
if self._connections.get(connection.endpoint) is existing:
|
||||
self._connections.pop(connection.endpoint, None)
|
||||
if self._tasks.get(connection.endpoint) is existing_task:
|
||||
self._tasks.pop(connection.endpoint, None)
|
||||
try:
|
||||
await self._dial(connection, servicer, wait_until_ready=True)
|
||||
except BaseException as operation:
|
||||
try:
|
||||
await _finish_rollback(
|
||||
self._restore_failed_redial(
|
||||
connection.endpoint, existing, existing_task
|
||||
),
|
||||
description="restore the previous inbound connector",
|
||||
)
|
||||
except InboundConnectionRollbackError as rollback:
|
||||
raise rollback from operation
|
||||
raise
|
||||
return connection
|
||||
|
||||
if existing is not None:
|
||||
# Authenticate and apply identity/version policy before touching
|
||||
# the only working connector or its durable credential.
|
||||
await NodeConnection(servicer, connection).probe()
|
||||
|
||||
entries = [
|
||||
entry
|
||||
for entry in original_entries
|
||||
if _endpoint_of(entry) != connection.endpoint
|
||||
]
|
||||
# Keyed by endpoint: re-pasting a rotated key for the same machine
|
||||
# replaces it rather than leaving a dead entry that retries forever.
|
||||
entries.append(connection.endpoint)
|
||||
try:
|
||||
if existing is not None:
|
||||
# An offline connector can own work retained on the node. Its
|
||||
# shutdown guard must run before replacement state is persisted.
|
||||
await self._drop(connection.endpoint)
|
||||
self.credentials.remember_connection_secret(
|
||||
connection.endpoint, connection.secret, connection.fingerprint
|
||||
)
|
||||
self._save(entries)
|
||||
await self._dial(
|
||||
connection, servicer, wait_until_ready=existing is not None
|
||||
)
|
||||
except BaseException as operation:
|
||||
try:
|
||||
await _finish_rollback(
|
||||
self._rollback_add(
|
||||
connection,
|
||||
servicer,
|
||||
existing,
|
||||
original_entries,
|
||||
old_secret,
|
||||
old_fingerprint,
|
||||
),
|
||||
description="restore the previous inbound connection",
|
||||
)
|
||||
except InboundConnectionRollbackError as rollback:
|
||||
raise rollback from operation
|
||||
raise
|
||||
return connection
|
||||
|
||||
async def _restore_failed_redial(
|
||||
self, endpoint: str, existing, existing_task: Optional[asyncio.Task]
|
||||
) -> None:
|
||||
failure = None
|
||||
candidate = self._connections.get(endpoint)
|
||||
candidate_task = self._tasks.get(endpoint)
|
||||
if candidate is not None and candidate is not existing:
|
||||
try:
|
||||
await self._close_candidate(endpoint, candidate, candidate_task)
|
||||
except BaseException as exc:
|
||||
failure = exc
|
||||
if endpoint not in self._connections:
|
||||
self._connections[endpoint] = existing
|
||||
if endpoint not in self._tasks and existing_task is not None:
|
||||
self._tasks[endpoint] = existing_task
|
||||
if failure is not None:
|
||||
raise InboundConnectionRollbackError(
|
||||
"The previous GPU-machine connector could not be restored safely."
|
||||
) from failure
|
||||
|
||||
async def _close_candidate(self, endpoint: str, candidate, candidate_task) -> None:
|
||||
try:
|
||||
close = getattr(candidate, "close", None)
|
||||
if callable(close):
|
||||
await close()
|
||||
finally:
|
||||
if candidate_task is not None:
|
||||
candidate_task.cancel()
|
||||
await asyncio.gather(candidate_task, return_exceptions=True)
|
||||
if self._connections.get(endpoint) is candidate:
|
||||
self._connections.pop(endpoint, None)
|
||||
if self._tasks.get(endpoint) is candidate_task:
|
||||
self._tasks.pop(endpoint, None)
|
||||
|
||||
async def _rollback_add(
|
||||
self,
|
||||
connection: Connection,
|
||||
servicer,
|
||||
existing,
|
||||
original_entries: list[str],
|
||||
old_secret: str,
|
||||
old_fingerprint: str,
|
||||
) -> None:
|
||||
"""Restore both live and durable generations after a failed replacement."""
|
||||
endpoint = connection.endpoint
|
||||
failures = []
|
||||
candidate = self._connections.get(endpoint)
|
||||
candidate_task = self._tasks.get(endpoint)
|
||||
if candidate is not None and candidate is not existing:
|
||||
try:
|
||||
await self._close_candidate(endpoint, candidate, candidate_task)
|
||||
except BaseException as exc:
|
||||
failures.append(exc)
|
||||
try:
|
||||
if old_secret:
|
||||
self.credentials.remember_connection_secret(
|
||||
endpoint, old_secret, old_fingerprint
|
||||
)
|
||||
else:
|
||||
self.credentials.forget_connection_secret(endpoint)
|
||||
except BaseException as exc:
|
||||
failures.append(exc)
|
||||
try:
|
||||
self._save(original_entries)
|
||||
except BaseException as exc:
|
||||
failures.append(exc)
|
||||
if (
|
||||
not failures
|
||||
and existing is not None
|
||||
and old_secret
|
||||
and endpoint not in self._connections
|
||||
):
|
||||
try:
|
||||
await self._dial(
|
||||
Connection(
|
||||
host=connection.host,
|
||||
port=connection.port,
|
||||
secret=old_secret,
|
||||
fingerprint=old_fingerprint,
|
||||
),
|
||||
servicer,
|
||||
)
|
||||
except BaseException as exc:
|
||||
failures.append(exc)
|
||||
if failures:
|
||||
raise InboundConnectionRollbackError(
|
||||
"The previous GPU-machine connection could not be restored safely. "
|
||||
"It remains stopped; fix its connection/settings storage, then "
|
||||
"paste the original connection again."
|
||||
) from failures[0]
|
||||
|
||||
async def _drop(self, endpoint: str) -> None:
|
||||
connection = self._connections.get(endpoint)
|
||||
task = self._tasks.get(endpoint)
|
||||
if connection is not None:
|
||||
await connection.stop()
|
||||
if self._connections.get(endpoint) is connection:
|
||||
self._connections.pop(endpoint, None)
|
||||
if self._tasks.get(endpoint) is task:
|
||||
self._tasks.pop(endpoint, None)
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
|
||||
async def remove(self, endpoint: str) -> bool:
|
||||
async with self._lifecycle_lock:
|
||||
return await self._remove(endpoint)
|
||||
|
||||
async def _remove(self, endpoint: str) -> bool:
|
||||
existed = endpoint in self._connections
|
||||
original_entries = self.saved()
|
||||
secret = self.credentials.connection_secret(endpoint)
|
||||
fingerprint = self.credentials.connection_fingerprint(endpoint)
|
||||
previous_connection = None
|
||||
if secret and fingerprint:
|
||||
try:
|
||||
previous_connection = self._connection_for(endpoint)
|
||||
except InvalidConnectionString:
|
||||
pass
|
||||
# A disconnected node deliberately retains work for reconnect. Do not
|
||||
# erase the only connector/key capable of delivering terminal shutdown.
|
||||
try:
|
||||
await self._drop(endpoint)
|
||||
entries = [e for e in original_entries if _endpoint_of(e) != endpoint]
|
||||
# Remove the protected credential first. If that durable write
|
||||
# fails, restore both durable generations and the live connector.
|
||||
self.credentials.forget_connection_secret(endpoint)
|
||||
self._save(entries)
|
||||
except BaseException as operation:
|
||||
try:
|
||||
await _finish_rollback(
|
||||
self._rollback_remove(
|
||||
endpoint,
|
||||
original_entries,
|
||||
secret,
|
||||
fingerprint,
|
||||
previous_connection,
|
||||
),
|
||||
description="restore removed inbound connection state",
|
||||
)
|
||||
except InboundConnectionRollbackError as rollback:
|
||||
raise rollback from operation
|
||||
raise
|
||||
return existed
|
||||
|
||||
async def _rollback_remove(
|
||||
self,
|
||||
endpoint: str,
|
||||
original_entries: list[str],
|
||||
secret: str,
|
||||
fingerprint: str,
|
||||
previous_connection: Optional[Connection],
|
||||
) -> None:
|
||||
failures = []
|
||||
try:
|
||||
if secret:
|
||||
self.credentials.remember_connection_secret(
|
||||
endpoint, secret, fingerprint
|
||||
)
|
||||
except BaseException as exc:
|
||||
failures.append(exc)
|
||||
try:
|
||||
self._save(original_entries)
|
||||
except BaseException as exc:
|
||||
failures.append(exc)
|
||||
if (
|
||||
not failures
|
||||
and previous_connection is not None
|
||||
and self._servicer is not None
|
||||
and endpoint not in self._connections
|
||||
):
|
||||
try:
|
||||
await self._dial(previous_connection, self._servicer)
|
||||
except BaseException as exc:
|
||||
failures.append(exc)
|
||||
if failures:
|
||||
raise InboundConnectionRollbackError(
|
||||
"The removed GPU-machine connection could not be restored safely. "
|
||||
"It remains stopped; fix its connection/settings storage, then retry."
|
||||
) from failures[0]
|
||||
|
||||
async def start_all(self, servicer) -> None:
|
||||
async with self._lifecycle_lock:
|
||||
for entry in self.saved():
|
||||
try:
|
||||
await self._dial(self._connection_for(entry), servicer)
|
||||
except InvalidConnectionString as exc:
|
||||
logger.warning(
|
||||
"Ignoring a saved connection that no longer parses: %s", exc
|
||||
)
|
||||
|
||||
async def _dial(
|
||||
self, connection: Connection, servicer, *, wait_until_ready: bool = False
|
||||
) -> None:
|
||||
from worker.inbound.connector import NodeConnection # noqa: PLC0415
|
||||
|
||||
self._servicer = servicer
|
||||
existing = self._connections.get(connection.endpoint)
|
||||
if existing is not None:
|
||||
if wait_until_ready and getattr(existing, "_connection", None) != connection:
|
||||
from worker.inbound.connector import ( # noqa: PLC0415
|
||||
InboundConnectionError,
|
||||
)
|
||||
|
||||
raise InboundConnectionError(
|
||||
"A different connection to that GPU machine is already active."
|
||||
)
|
||||
return
|
||||
node = NodeConnection(servicer, connection)
|
||||
self._connections[connection.endpoint] = node
|
||||
self._tasks[connection.endpoint] = asyncio.create_task(
|
||||
task = asyncio.create_task(
|
||||
node.run_forever(), name=f"inbound-node-{connection.endpoint}"
|
||||
)
|
||||
task.add_done_callback(self._observe_connection_result)
|
||||
self._tasks[connection.endpoint] = task
|
||||
if wait_until_ready:
|
||||
await node.wait_until_registered(task)
|
||||
|
||||
@staticmethod
|
||||
def _observe_connection_result(task: asyncio.Task) -> None:
|
||||
"""Retrieve terminal dial errors; NodeConnection retains the UI detail."""
|
||||
if task.cancelled():
|
||||
return
|
||||
try:
|
||||
task.exception()
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
async def stop(self) -> None:
|
||||
async with self._lifecycle_lock:
|
||||
await self._stop_all()
|
||||
|
||||
async def _stop_all(self) -> None:
|
||||
for connection in list(self._connections.values()):
|
||||
await connection.stop()
|
||||
for task in list(self._tasks.values()):
|
||||
close = getattr(connection, "close", None)
|
||||
if callable(close):
|
||||
await close()
|
||||
else:
|
||||
await connection.stop()
|
||||
tasks = list(self._tasks.values())
|
||||
for task in tasks:
|
||||
task.cancel()
|
||||
if tasks:
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
self._connections.clear()
|
||||
self._tasks.clear()
|
||||
|
||||
|
||||
+38
-5
@@ -17,7 +17,7 @@ from dataclasses import dataclass, field
|
||||
from typing import Iterator, Optional
|
||||
|
||||
from worker.breaker import BreakerRegistry
|
||||
from worker.capacity import WorkerCapacity, derive_concurrency
|
||||
from worker.capacity import WorkerCapacity, clamp_concurrency, derive_concurrency
|
||||
from worker.clock import resolve
|
||||
from worker.identity import Session
|
||||
from worker.registry import RemoteWorker
|
||||
@@ -33,6 +33,9 @@ _HEARTBEAT_MISS_SECONDS = 90.0
|
||||
# ping is a ~25-second view: current enough to notice a link degrading, long
|
||||
# enough that one slow answer cannot move it.
|
||||
_LATENCY_WINDOW = 5
|
||||
_KNOWN_EXECUTION_DEVICES = frozenset(
|
||||
{"cpu", "cuda", "mps", "mlx", "directml", "rocm", "xpu"}
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -55,6 +58,9 @@ class ConnectedWorker:
|
||||
# The address this worker connected FROM, as the control plane saw it.
|
||||
address: str = ""
|
||||
draining: bool = False
|
||||
# Registration handoff temporarily stops new assignments without
|
||||
# conflating that transport state with a user-requested drain/shutdown.
|
||||
registration_pending: bool = False
|
||||
# Attempt ids this worker claims to be running. Rebuilt on every reconnect
|
||||
# from its own report, never inferred.
|
||||
in_flight: set[str] = field(default_factory=set)
|
||||
@@ -79,7 +85,11 @@ class ConnectedWorker:
|
||||
"""
|
||||
if self.stale():
|
||||
return "offline"
|
||||
if self.draining or self.capacity.available_slots <= 0:
|
||||
if (
|
||||
self.draining
|
||||
or self.registration_pending
|
||||
or self.capacity.available_slots <= 0
|
||||
):
|
||||
return "busy"
|
||||
return "ready"
|
||||
|
||||
@@ -100,6 +110,21 @@ class ConnectedWorker:
|
||||
return bool(cap.get("supported")) and bool(cap.get("installed", True))
|
||||
return False
|
||||
|
||||
def execution_device(self, engine: str, model_id: str, operation: str) -> str:
|
||||
"""Device used by the exact capability selected for this task."""
|
||||
for cap in self.record.capabilities:
|
||||
if cap.get("engine") != engine:
|
||||
continue
|
||||
if model_id and cap.get("model_id") not in (model_id, "", None):
|
||||
continue
|
||||
if operation and operation not in (cap.get("operations") or [operation]):
|
||||
continue
|
||||
if cap.get("cpu_fallback"):
|
||||
return "cpu"
|
||||
backend = str(cap.get("backend") or "").lower()
|
||||
return backend if backend in _KNOWN_EXECUTION_DEVICES else "cpu"
|
||||
return "cpu"
|
||||
|
||||
def is_warm(self, engine: str, model_id: str) -> bool:
|
||||
return self.capacity.is_resident(engine, model_id)
|
||||
|
||||
@@ -157,7 +182,7 @@ class WorkerPool:
|
||||
epoch=epoch,
|
||||
capacity=WorkerCapacity(
|
||||
worker_id=record.id,
|
||||
max_concurrent_tasks=max(1, max_concurrent_tasks),
|
||||
max_concurrent_tasks=clamp_concurrency(max_concurrent_tasks),
|
||||
backend=backend,
|
||||
),
|
||||
connected_at=stamp,
|
||||
@@ -213,6 +238,10 @@ class WorkerPool:
|
||||
def disconnect(self, worker_id: str) -> Optional[ConnectedWorker]:
|
||||
return self._connected.pop(worker_id, None)
|
||||
|
||||
def restore_connection(self, worker: ConnectedWorker) -> None:
|
||||
"""Restore an exact live snapshot after replacement activation fails."""
|
||||
self._connected[worker.worker_id] = worker
|
||||
|
||||
def get(self, worker_id: str) -> Optional[ConnectedWorker]:
|
||||
return self._connected.get(worker_id)
|
||||
|
||||
@@ -282,10 +311,14 @@ class WorkerPool:
|
||||
worker.capacity.slots[key] = ModelSlot(
|
||||
engine=cap.get("engine", ""),
|
||||
model_id=cap.get("model_id", ""),
|
||||
derived_concurrency=max(0, declared),
|
||||
derived_concurrency=clamp_concurrency(
|
||||
declared, allow_zero=True
|
||||
),
|
||||
)
|
||||
else:
|
||||
slot.derived_concurrency = max(0, declared)
|
||||
slot.derived_concurrency = clamp_concurrency(
|
||||
declared, allow_zero=True
|
||||
)
|
||||
|
||||
def stale_workers(self, *, now: Optional[float] = None) -> list[ConnectedWorker]:
|
||||
return [w for w in self if w.stale(now=now)]
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
# Generated by the protocol buffer compiler. DO NOT EDIT!
|
||||
# NO CHECKED-IN PROTOBUF GENCODE
|
||||
# source: worker_v1.proto
|
||||
# Protobuf Python Version: 6.33.5
|
||||
# Protobuf Python Version: 7.35.1
|
||||
"""Generated protocol buffer code."""
|
||||
from google.protobuf import descriptor as _descriptor
|
||||
from google.protobuf import descriptor_pool as _descriptor_pool
|
||||
@@ -11,9 +11,9 @@ from google.protobuf import symbol_database as _symbol_database
|
||||
from google.protobuf.internal import builder as _builder
|
||||
_runtime_version.ValidateProtobufRuntimeVersion(
|
||||
_runtime_version.Domain.PUBLIC,
|
||||
6,
|
||||
33,
|
||||
5,
|
||||
7,
|
||||
35,
|
||||
1,
|
||||
'',
|
||||
'worker_v1.proto'
|
||||
)
|
||||
|
||||
@@ -5,7 +5,7 @@ import warnings
|
||||
|
||||
from . import worker_v1_pb2 as worker__v1__pb2
|
||||
|
||||
GRPC_GENERATED_VERSION = '1.81.1'
|
||||
GRPC_GENERATED_VERSION = '1.83.0'
|
||||
GRPC_VERSION = grpc.__version__
|
||||
_version_not_supported = False
|
||||
|
||||
|
||||
@@ -7,7 +7,9 @@
|
||||
// Rules of the road (goal_v2.md A5):
|
||||
// * Additive-only within v1. Never renumber, never reuse a field number.
|
||||
// * Version negotiation happens at Register; the server may refuse with
|
||||
// UPGRADE_REQUIRED. Supported skew window is N-2.
|
||||
// UPGRADE_REQUIRED. Semantic protocol v2 is intentionally incompatible
|
||||
// with v1 because enrollment became a durable two-phase handshake; never
|
||||
// infer a release-based skew window across that boundary.
|
||||
// * The Control stream carries SMALL messages only. Artifacts (reference
|
||||
// audio in, rendered audio/video out) move through UploadResult /
|
||||
// DownloadArtifact. A large payload on the control stream head-of-line
|
||||
|
||||
+254
-53
@@ -19,12 +19,15 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
from core.db import db_conn
|
||||
from worker.capacity import clamp_concurrency
|
||||
from worker.clock import resolve
|
||||
from worker import identity
|
||||
|
||||
@@ -33,6 +36,14 @@ logger = logging.getLogger("omnivoice.worker")
|
||||
# Default scheduling preference. Higher wins; equal priorities fall through to
|
||||
# least-busy, which is the actual default behaviour for a homogeneous setup.
|
||||
_DEFAULT_PRIORITY = 50
|
||||
_AUTHORITY_LOCK = threading.RLock()
|
||||
|
||||
|
||||
@contextmanager
|
||||
def authority_guard():
|
||||
"""Serialize durable authority changes with live-session publication."""
|
||||
with _AUTHORITY_LOCK:
|
||||
yield
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -180,6 +191,54 @@ def purge_expired_enrollments(*, now: Optional[float] = None) -> int:
|
||||
# ── Workers ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _new_worker(
|
||||
*,
|
||||
name: str,
|
||||
public_key: bytes,
|
||||
endpoint: str,
|
||||
host: Optional[dict],
|
||||
capabilities: Optional[list[dict]],
|
||||
max_concurrent_tasks: int,
|
||||
consent_granted: bool,
|
||||
stamp: float,
|
||||
) -> RemoteWorker:
|
||||
key_id = identity.key_id_for(public_key)
|
||||
return RemoteWorker(
|
||||
id=uuid.uuid4().hex[:12],
|
||||
name=name or key_id,
|
||||
key_id=key_id,
|
||||
public_key=public_key,
|
||||
endpoint=endpoint,
|
||||
host=host or {},
|
||||
capabilities=capabilities or [],
|
||||
max_concurrent_tasks=clamp_concurrency(max_concurrent_tasks),
|
||||
consent_granted_at=stamp if consent_granted else None,
|
||||
created_at=stamp,
|
||||
)
|
||||
|
||||
|
||||
def _insert_worker(conn, worker: RemoteWorker) -> None:
|
||||
conn.execute(
|
||||
"INSERT INTO remote_workers "
|
||||
"(id, name, key_id, public_key, enabled, revoked, priority, endpoint, host_json, "
|
||||
" capabilities_json, max_concurrent_tasks, session_epoch, consent_granted_at, created_at) "
|
||||
"VALUES (?, ?, ?, ?, 1, 0, ?, ?, ?, ?, ?, 0, ?, ?)",
|
||||
(
|
||||
worker.id,
|
||||
worker.name,
|
||||
worker.key_id,
|
||||
worker.public_key,
|
||||
worker.priority,
|
||||
worker.endpoint,
|
||||
json.dumps(worker.host),
|
||||
json.dumps(worker.capabilities),
|
||||
worker.max_concurrent_tasks,
|
||||
worker.consent_granted_at,
|
||||
worker.created_at,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def enroll_worker(
|
||||
*,
|
||||
name: str,
|
||||
@@ -198,49 +257,128 @@ def enroll_worker(
|
||||
your own desktop is not agreeing to use someone else's.
|
||||
"""
|
||||
stamp = resolve(now)
|
||||
key_id = identity.key_id_for(public_key)
|
||||
existing = get_by_key_id(key_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
worker = RemoteWorker(
|
||||
id=uuid.uuid4().hex[:12],
|
||||
name=name or key_id,
|
||||
key_id=key_id,
|
||||
worker = _new_worker(
|
||||
name=name,
|
||||
public_key=public_key,
|
||||
endpoint=endpoint,
|
||||
host=host or {},
|
||||
capabilities=capabilities or [],
|
||||
max_concurrent_tasks=max(1, int(max_concurrent_tasks)),
|
||||
consent_granted_at=stamp if consent_granted else None,
|
||||
created_at=stamp,
|
||||
host=host,
|
||||
capabilities=capabilities,
|
||||
max_concurrent_tasks=max_concurrent_tasks,
|
||||
consent_granted=consent_granted,
|
||||
stamp=stamp,
|
||||
)
|
||||
with db_conn() as conn:
|
||||
conn.execute(
|
||||
"INSERT INTO remote_workers "
|
||||
"(id, name, key_id, public_key, enabled, revoked, priority, endpoint, host_json, "
|
||||
" capabilities_json, max_concurrent_tasks, session_epoch, consent_granted_at, created_at) "
|
||||
"VALUES (?, ?, ?, ?, 1, 0, ?, ?, ?, ?, ?, 0, ?, ?)",
|
||||
(
|
||||
worker.id,
|
||||
worker.name,
|
||||
worker.key_id,
|
||||
worker.public_key,
|
||||
worker.priority,
|
||||
worker.endpoint,
|
||||
json.dumps(worker.host),
|
||||
json.dumps(worker.capabilities),
|
||||
worker.max_concurrent_tasks,
|
||||
worker.consent_granted_at,
|
||||
worker.created_at,
|
||||
),
|
||||
)
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT * FROM remote_workers WHERE key_id = ?", (worker.key_id,)
|
||||
).fetchone()
|
||||
if row is not None:
|
||||
return _row_to_worker(row)
|
||||
_insert_worker(conn, worker)
|
||||
logger.info("Enrolled remote worker %s (%s)", worker.name, worker.key_id)
|
||||
return worker
|
||||
|
||||
|
||||
def enroll_with_token(
|
||||
token: identity.EnrollmentToken,
|
||||
*,
|
||||
name: str,
|
||||
public_key: bytes,
|
||||
endpoint: str = "",
|
||||
host: Optional[dict] = None,
|
||||
capabilities: Optional[list[dict]] = None,
|
||||
max_concurrent_tasks: int = 1,
|
||||
consent_granted: bool = True,
|
||||
now: Optional[float] = None,
|
||||
) -> Optional[RemoteWorker]:
|
||||
"""Consume a token and bind its worker identity in one transaction."""
|
||||
stamp = resolve(now)
|
||||
worker = _new_worker(
|
||||
name=name,
|
||||
public_key=public_key,
|
||||
endpoint=endpoint,
|
||||
host=host,
|
||||
capabilities=capabilities,
|
||||
max_concurrent_tasks=max_concurrent_tasks,
|
||||
consent_granted=consent_granted,
|
||||
stamp=stamp,
|
||||
)
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
enrollment = conn.execute(
|
||||
"SELECT secret_hash, expires_at, used_at FROM remote_worker_enrollments "
|
||||
"WHERE token_id = ?",
|
||||
(token.token_id,),
|
||||
).fetchone()
|
||||
if (
|
||||
enrollment is None
|
||||
or enrollment["used_at"] is not None
|
||||
or stamp > float(enrollment["expires_at"])
|
||||
or not identity.constant_time_equals(
|
||||
enrollment["secret_hash"], token.secret_hash
|
||||
)
|
||||
):
|
||||
return None
|
||||
existing_row = conn.execute(
|
||||
"SELECT * FROM remote_workers WHERE key_id = ?", (worker.key_id,)
|
||||
).fetchone()
|
||||
if existing_row is not None and bool(existing_row["revoked"]):
|
||||
return None
|
||||
enrolled = _row_to_worker(existing_row) if existing_row is not None else worker
|
||||
consumed = conn.execute(
|
||||
"UPDATE remote_worker_enrollments SET used_at = ?, used_by_worker = ? "
|
||||
"WHERE token_id = ? AND used_at IS NULL",
|
||||
(stamp, enrolled.id, token.token_id),
|
||||
)
|
||||
if consumed.rowcount != 1:
|
||||
return None
|
||||
if existing_row is None:
|
||||
_insert_worker(conn, worker)
|
||||
if existing_row is None:
|
||||
logger.info("Enrolled remote worker %s (%s)", worker.name, worker.key_id)
|
||||
return enrolled
|
||||
|
||||
|
||||
def recover_enrollment_with_token(
|
||||
token: identity.EnrollmentToken, *, public_key: bytes
|
||||
) -> Optional[RemoteWorker]:
|
||||
"""Resolve a spent token only to the exact key it originally enrolled.
|
||||
|
||||
This is only the durable lookup half of recovery. The transport must also
|
||||
verify a fresh signature from this key before issuing another session; a
|
||||
spent token and an observed public key are not proof of private-key
|
||||
possession.
|
||||
"""
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
enrollment = conn.execute(
|
||||
"SELECT secret_hash, used_at, used_by_worker "
|
||||
"FROM remote_worker_enrollments WHERE token_id = ?",
|
||||
(token.token_id,),
|
||||
).fetchone()
|
||||
if (
|
||||
enrollment is None
|
||||
or enrollment["used_at"] is None
|
||||
or not enrollment["used_by_worker"]
|
||||
or not identity.constant_time_equals(
|
||||
enrollment["secret_hash"], token.secret_hash
|
||||
)
|
||||
):
|
||||
return None
|
||||
row = conn.execute(
|
||||
"SELECT * FROM remote_workers WHERE id = ?",
|
||||
(enrollment["used_by_worker"],),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
worker = _row_to_worker(row)
|
||||
if worker.revoked or worker.public_key != public_key:
|
||||
return None
|
||||
return worker
|
||||
|
||||
|
||||
def get(worker_id: str) -> Optional[RemoteWorker]:
|
||||
with db_conn() as conn:
|
||||
row = conn.execute("SELECT * FROM remote_workers WHERE id = ?", (worker_id,)).fetchone()
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT * FROM remote_workers WHERE id = ?", (worker_id,)
|
||||
).fetchone()
|
||||
return _row_to_worker(row) if row else None
|
||||
|
||||
|
||||
@@ -268,7 +406,7 @@ def begin_session(worker_id: str, *, now: Optional[float] = None) -> int:
|
||||
otherwise delivers two accepts for one assignment.
|
||||
"""
|
||||
stamp = resolve(now)
|
||||
with db_conn() as conn:
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
conn.execute(
|
||||
"UPDATE remote_workers SET session_epoch = session_epoch + 1, last_seen_at = ? WHERE id = ?",
|
||||
(stamp, worker_id),
|
||||
@@ -292,29 +430,79 @@ def update_capabilities(
|
||||
capabilities: list[dict],
|
||||
host: Optional[dict] = None,
|
||||
max_concurrent_tasks: Optional[int] = None,
|
||||
_conn=None,
|
||||
) -> None:
|
||||
sets = ["capabilities_json = ?"]
|
||||
params: list = [json.dumps(capabilities)]
|
||||
if host is not None:
|
||||
sets.append("host_json = ?")
|
||||
params.append(json.dumps(host))
|
||||
if max_concurrent_tasks is not None:
|
||||
sets.append("max_concurrent_tasks = ?")
|
||||
params.append(max(1, int(max_concurrent_tasks)))
|
||||
params.append(worker_id)
|
||||
with db_conn() as conn:
|
||||
conn.execute(f"UPDATE remote_workers SET {', '.join(sets)} WHERE id = ?", params)
|
||||
params = (
|
||||
json.dumps(capabilities),
|
||||
json.dumps(host) if host is not None else None,
|
||||
clamp_concurrency(max_concurrent_tasks)
|
||||
if max_concurrent_tasks is not None
|
||||
else None,
|
||||
worker_id,
|
||||
)
|
||||
sql = """
|
||||
UPDATE remote_workers
|
||||
SET capabilities_json = ?,
|
||||
host_json = COALESCE(?, host_json),
|
||||
max_concurrent_tasks = COALESCE(?, max_concurrent_tasks)
|
||||
WHERE id = ?
|
||||
"""
|
||||
if _conn is not None:
|
||||
# The caller owns the surrounding SQLite transaction. Acquiring the
|
||||
# authority lock inside it can deadlock against revoke, which takes
|
||||
# that lock before waiting for the same database write lock.
|
||||
_conn.execute(sql, params)
|
||||
else:
|
||||
with _AUTHORITY_LOCK:
|
||||
with db_conn() as conn:
|
||||
conn.execute(sql, params)
|
||||
|
||||
|
||||
def update_policy(
|
||||
worker_id: str,
|
||||
*,
|
||||
name: Optional[str] = None,
|
||||
enabled: Optional[bool] = None,
|
||||
priority: Optional[int] = None,
|
||||
) -> Optional[RemoteWorker]:
|
||||
"""Atomically update user-controlled policy and return the committed row."""
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT * FROM remote_workers WHERE id = ?", (worker_id,)
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
if name is not None or enabled is not None or priority is not None:
|
||||
conn.execute(
|
||||
"""
|
||||
UPDATE remote_workers
|
||||
SET name = COALESCE(?, name),
|
||||
enabled = COALESCE(?, enabled),
|
||||
priority = COALESCE(?, priority)
|
||||
WHERE id = ?
|
||||
""",
|
||||
(
|
||||
name,
|
||||
1 if enabled is True else 0 if enabled is False else None,
|
||||
max(0, min(100, int(priority))) if priority is not None else None,
|
||||
worker_id,
|
||||
),
|
||||
)
|
||||
row = conn.execute(
|
||||
"SELECT * FROM remote_workers WHERE id = ?", (worker_id,)
|
||||
).fetchone()
|
||||
return _row_to_worker(row) if row is not None else None
|
||||
|
||||
|
||||
def set_enabled(worker_id: str, enabled: bool) -> None:
|
||||
with db_conn() as conn:
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
conn.execute(
|
||||
"UPDATE remote_workers SET enabled = ? WHERE id = ?", (1 if enabled else 0, worker_id)
|
||||
)
|
||||
|
||||
|
||||
def set_priority(worker_id: str, priority: int) -> None:
|
||||
with db_conn() as conn:
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
conn.execute(
|
||||
"UPDATE remote_workers SET priority = ? WHERE id = ?",
|
||||
(max(0, min(100, int(priority))), worker_id),
|
||||
@@ -322,7 +510,7 @@ def set_priority(worker_id: str, priority: int) -> None:
|
||||
|
||||
|
||||
def rename(worker_id: str, name: str) -> None:
|
||||
with db_conn() as conn:
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
conn.execute("UPDATE remote_workers SET name = ? WHERE id = ?", (name, worker_id))
|
||||
|
||||
|
||||
@@ -334,7 +522,7 @@ def revoke(worker_id: str, *, now: Optional[float] = None) -> bool:
|
||||
who could simply enroll again.
|
||||
"""
|
||||
stamp = resolve(now)
|
||||
with db_conn() as conn:
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
cur = conn.execute(
|
||||
"UPDATE remote_workers SET revoked = 1, revoked_at = ?, enabled = 0 WHERE id = ?",
|
||||
(stamp, worker_id),
|
||||
@@ -345,15 +533,23 @@ def revoke(worker_id: str, *, now: Optional[float] = None) -> bool:
|
||||
|
||||
|
||||
def is_revoked(key_id: str) -> bool:
|
||||
with db_conn() as conn:
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT revoked FROM remote_workers WHERE key_id = ?", (key_id,)
|
||||
).fetchone()
|
||||
return bool(row and row["revoked"])
|
||||
|
||||
|
||||
def is_enabled(worker_id: str) -> bool:
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT enabled FROM remote_workers WHERE id = ?", (worker_id,)
|
||||
).fetchone()
|
||||
return bool(row and row["enabled"])
|
||||
|
||||
|
||||
def grant_consent(worker_id: str, *, now: Optional[float] = None) -> None:
|
||||
with db_conn() as conn:
|
||||
with _AUTHORITY_LOCK, db_conn() as conn:
|
||||
conn.execute(
|
||||
"UPDATE remote_workers SET consent_granted_at = ? WHERE id = ?",
|
||||
(resolve(now), worker_id),
|
||||
@@ -397,16 +593,20 @@ def authenticate(
|
||||
|
||||
__all__ = [
|
||||
"RemoteWorker",
|
||||
"authority_guard",
|
||||
"authenticate",
|
||||
"begin_session",
|
||||
"create_enrollment",
|
||||
"enroll_with_token",
|
||||
"enroll_worker",
|
||||
"get",
|
||||
"get_by_key_id",
|
||||
"grant_consent",
|
||||
"is_revoked",
|
||||
"is_enabled",
|
||||
"list_workers",
|
||||
"purge_expired_enrollments",
|
||||
"recover_enrollment_with_token",
|
||||
"redeem_enrollment",
|
||||
"rename",
|
||||
"revoke",
|
||||
@@ -414,4 +614,5 @@ __all__ = [
|
||||
"set_priority",
|
||||
"touch",
|
||||
"update_capabilities",
|
||||
"update_policy",
|
||||
]
|
||||
|
||||
+332
-24
@@ -27,14 +27,20 @@ be minutes.
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import enum
|
||||
import logging
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, fields
|
||||
from functools import partial
|
||||
from typing import Callable, Optional
|
||||
|
||||
from worker import deadlines as deadline_policy
|
||||
from worker import task_store
|
||||
from worker import registry, task_store
|
||||
from worker.async_utils import (
|
||||
to_thread_and_defer_cancellation,
|
||||
to_thread_and_drain_on_cancel,
|
||||
)
|
||||
from worker.breaker import Attribution
|
||||
from worker.capacity import WorkerCapacity
|
||||
from worker.clock import resolve
|
||||
@@ -100,6 +106,29 @@ class SchedulerStopped(RuntimeError):
|
||||
"""
|
||||
|
||||
|
||||
def _adopt_task_state(target: Task, source: Task) -> None:
|
||||
"""Commit a reconciled copy while preserving public task/attempt identities."""
|
||||
current_attempts = {attempt.attempt_id: attempt for attempt in target.attempts}
|
||||
adopted_attempts = []
|
||||
for source_attempt in source.attempts:
|
||||
target_attempt = current_attempts.get(source_attempt.attempt_id)
|
||||
if target_attempt is None:
|
||||
target_attempt = copy.deepcopy(source_attempt)
|
||||
else:
|
||||
for item in fields(Attempt):
|
||||
setattr(
|
||||
target_attempt,
|
||||
item.name,
|
||||
copy.deepcopy(getattr(source_attempt, item.name)),
|
||||
)
|
||||
adopted_attempts.append(target_attempt)
|
||||
|
||||
for item in fields(Task):
|
||||
if item.name != "attempts":
|
||||
setattr(target, item.name, copy.deepcopy(getattr(source, item.name)))
|
||||
target.attempts[:] = adopted_attempts
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Assignment:
|
||||
"""One task bound to one worker, ready to send."""
|
||||
@@ -110,6 +139,18 @@ class Assignment:
|
||||
deadlines: deadline_policy.Deadlines
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ReconciliationGeneration:
|
||||
"""Immutable task copies persisted before their live generation is adopted."""
|
||||
|
||||
worker_id: str
|
||||
stamp: float
|
||||
originals: tuple[Task, ...]
|
||||
snapshots: tuple[Task, ...]
|
||||
candidates: tuple[Task, ...]
|
||||
zombies: tuple[str, ...] = ()
|
||||
|
||||
|
||||
class Scheduler:
|
||||
"""Owns the queue, the selection pipeline, and the deadline sweeper."""
|
||||
|
||||
@@ -137,6 +178,10 @@ class Scheduler:
|
||||
# task_id → when its progress was last written through. Cleared with
|
||||
# the task, so it cannot outlive what it describes.
|
||||
self._progress_saved_at: dict[str, float] = {}
|
||||
# Async producers serialize admission while durable input staging is
|
||||
# off-loop. This keeps queue/idempotency decisions atomic without
|
||||
# holding the event loop hostage to a multi-gigabyte copy and hash.
|
||||
self._submission_lock = asyncio.Lock()
|
||||
|
||||
# ── Persistence seam ──────────────────────────────────────────────────
|
||||
|
||||
@@ -281,12 +326,99 @@ class Scheduler:
|
||||
)
|
||||
if deadline_seconds:
|
||||
task.deadline_at = stamp + deadline_seconds
|
||||
self._tasks[task.task_id] = task
|
||||
if self._persist:
|
||||
task_store.create(task, now=stamp)
|
||||
persisted = task_store.create(task, now=stamp)
|
||||
if persisted.task_id != task.task_id:
|
||||
self._tasks.setdefault(persisted.task_id, persisted)
|
||||
return persisted
|
||||
task = persisted
|
||||
# Publish to the dispatcher only after durable admission succeeds. A
|
||||
# failed DB/input-stage write must not send work the API rejected.
|
||||
self._tasks[task.task_id] = task
|
||||
self._emit("queued", task)
|
||||
return task
|
||||
|
||||
async def submit_async(
|
||||
self,
|
||||
*,
|
||||
operation: str,
|
||||
engine: str,
|
||||
model_id: str,
|
||||
params: Optional[dict] = None,
|
||||
priority: PriorityClass = PriorityClass.INTERACTIVE,
|
||||
idempotency_key: Optional[str] = None,
|
||||
max_attempts: int = 3,
|
||||
deadline_seconds: Optional[float] = None,
|
||||
pinned_worker_id: Optional[str] = None,
|
||||
now: Optional[float] = None,
|
||||
) -> Task:
|
||||
"""Admit durably without hashing/copying task inputs on the app loop."""
|
||||
async with self._submission_lock:
|
||||
stamp = resolve(now)
|
||||
if idempotency_key:
|
||||
for existing in self._tasks.values():
|
||||
if existing.idempotency_key == idempotency_key:
|
||||
return existing
|
||||
if self._persist:
|
||||
stored, cancelled = await to_thread_and_defer_cancellation(
|
||||
task_store.get_by_idempotency_key, idempotency_key
|
||||
)
|
||||
if stored is not None:
|
||||
self._tasks.setdefault(stored.task_id, stored)
|
||||
if cancelled:
|
||||
raise asyncio.CancelledError
|
||||
return stored
|
||||
if cancelled:
|
||||
raise asyncio.CancelledError
|
||||
|
||||
if self.queue_depth >= self.max_queue_depth:
|
||||
raise QueueFull(
|
||||
f"The remote task queue is full ({self.max_queue_depth} waiting). "
|
||||
"Wait for current work to finish, or add another worker."
|
||||
)
|
||||
|
||||
task = Task(
|
||||
task_id=uuid.uuid4().hex[:16],
|
||||
operation=operation,
|
||||
engine=engine,
|
||||
model_id=model_id,
|
||||
params=params or {},
|
||||
priority=priority,
|
||||
idempotency_key=idempotency_key,
|
||||
max_attempts=max_attempts,
|
||||
created_at=stamp,
|
||||
pinned_worker_id=pinned_worker_id,
|
||||
)
|
||||
if deadline_seconds:
|
||||
task.deadline_at = stamp + deadline_seconds
|
||||
if self._persist:
|
||||
persisted, cancelled = await to_thread_and_defer_cancellation(
|
||||
partial(task_store.create, task, now=stamp)
|
||||
)
|
||||
if persisted.task_id != task.task_id:
|
||||
self._tasks.setdefault(persisted.task_id, persisted)
|
||||
if cancelled:
|
||||
raise asyncio.CancelledError
|
||||
return persisted
|
||||
task = persisted
|
||||
if cancelled:
|
||||
# The durable copy cannot be interrupted midway. Settle it
|
||||
# before publication so cancellation can never dispatch a
|
||||
# task whose caller did not finish submitting it.
|
||||
task.cancel(reason="submission was cancelled", now=resolve())
|
||||
await to_thread_and_drain_on_cancel(
|
||||
partial(task_store.save, task, now=resolve())
|
||||
)
|
||||
self._tasks[task.task_id] = task
|
||||
self._emit("cancelled", task)
|
||||
raise asyncio.CancelledError
|
||||
|
||||
# The dispatcher only learns the task after input bytes and its DB
|
||||
# row are durable. A failed stage cannot escape as runnable work.
|
||||
self._tasks[task.task_id] = task
|
||||
self._emit("queued", task)
|
||||
return task
|
||||
|
||||
def adopt(self, task: Task) -> None:
|
||||
"""Take ownership of a task loaded from disk after a restart."""
|
||||
self._tasks[task.task_id] = task
|
||||
@@ -378,7 +510,11 @@ class Scheduler:
|
||||
for worker in self.pool:
|
||||
if task.pinned_worker_id and worker.worker_id != task.pinned_worker_id:
|
||||
continue
|
||||
if not worker.record.schedulable or worker.draining:
|
||||
if (
|
||||
not worker.record.schedulable
|
||||
or worker.draining
|
||||
or worker.registration_pending
|
||||
):
|
||||
continue
|
||||
if worker.stale(now=stamp):
|
||||
continue
|
||||
@@ -469,6 +605,13 @@ class Scheduler:
|
||||
empty or every waiting task is blocked on capacity. A task with no
|
||||
capable worker at all is failed here rather than left to age out.
|
||||
"""
|
||||
# Selection and binding are one authority read. A revoke/disable route
|
||||
# holds the same lock until its durable row and live pool record agree,
|
||||
# so another thread cannot bind work in that publication window.
|
||||
with registry.authority_guard():
|
||||
return self._next_assignment(now=now)
|
||||
|
||||
def _next_assignment(self, *, now: Optional[float] = None) -> Optional[Assignment]:
|
||||
stamp = resolve(now)
|
||||
for task in self._queued_order():
|
||||
try:
|
||||
@@ -506,6 +649,9 @@ class Scheduler:
|
||||
model_resident=worker.is_warm(task.engine, task.model_id),
|
||||
model_downloaded=True,
|
||||
input_seconds=float(task.params.get("input_seconds") or 0.0),
|
||||
execution_device=worker.execution_device(
|
||||
task.engine, task.model_id, task.operation
|
||||
),
|
||||
)
|
||||
attempt.renew_lease(budget.accept_seconds, now=now)
|
||||
self._save(task, now=now)
|
||||
@@ -697,9 +843,25 @@ class Scheduler:
|
||||
if epoch is not None and attempt.session_epoch != epoch:
|
||||
return False, task
|
||||
|
||||
committed, attempt = task.commit_result(
|
||||
attempt_id, result_ref=result_ref, session_epoch=epoch, now=stamp
|
||||
)
|
||||
if self._persist:
|
||||
# Prepare and persist the terminal generation before publishing it
|
||||
# to the live graph. A failed database write must leave the
|
||||
# redelivery path looking exactly as it did before this frame:
|
||||
# still in flight, still consuming capacity, and not yet credited
|
||||
# as a breaker success.
|
||||
candidate = copy.deepcopy(task)
|
||||
committed, _ = candidate.commit_result(
|
||||
attempt_id, result_ref=result_ref, session_epoch=epoch, now=stamp
|
||||
)
|
||||
if committed:
|
||||
task_store.commit_result(candidate, result_json=result, now=stamp)
|
||||
else:
|
||||
task_store.save(candidate, now=stamp)
|
||||
_adopt_task_state(task, candidate)
|
||||
else:
|
||||
committed, attempt = task.commit_result(
|
||||
attempt_id, result_ref=result_ref, session_epoch=epoch, now=stamp
|
||||
)
|
||||
self._release_slot(task, attempt, now=stamp)
|
||||
worker = self.pool.get(attempt.worker_id)
|
||||
if worker is not None and committed:
|
||||
@@ -708,10 +870,6 @@ class Scheduler:
|
||||
WorkerCapacity.slot_key(task.engine, task.model_id),
|
||||
now=stamp,
|
||||
)
|
||||
if committed and self._persist:
|
||||
task_store.commit_result(task, result_json=result, now=stamp)
|
||||
elif self._persist:
|
||||
self._save(task, now=stamp)
|
||||
self._emit("completed" if committed else "duplicate", task)
|
||||
return committed, task
|
||||
|
||||
@@ -776,28 +934,85 @@ class Scheduler:
|
||||
This is where the duplicate-execution bug would live if a disconnect
|
||||
were treated as a failure: the worker may be seconds from delivering.
|
||||
"""
|
||||
generation = self.prepare_disconnected(worker_id, now=now)
|
||||
try:
|
||||
# A disconnect is one generation, just like reconnect. If any
|
||||
# write fails, neither the database nor the live task graph may
|
||||
# contain a half-marked set of attempts.
|
||||
self.persist_reconciliation(generation)
|
||||
return self.apply_disconnected(generation)
|
||||
finally:
|
||||
# A failed persistence write must not leave a dead stream's pool
|
||||
# entry schedulable.
|
||||
self.pool.disconnect(worker_id)
|
||||
|
||||
def prepare_disconnected(
|
||||
self,
|
||||
worker_id: str,
|
||||
*,
|
||||
now: Optional[float] = None,
|
||||
include_task_ids: Optional[set[str]] = None,
|
||||
) -> _ReconciliationGeneration:
|
||||
"""Build, but do not publish, one disconnected task generation."""
|
||||
stamp = resolve(now)
|
||||
affected: list[Task] = []
|
||||
for task in self._tasks.values():
|
||||
included = include_task_ids or set()
|
||||
originals = tuple(
|
||||
task
|
||||
for task in self._tasks.values()
|
||||
if task.task_id in included
|
||||
or (
|
||||
task.active_attempt is not None
|
||||
and task.active_attempt.worker_id == worker_id
|
||||
)
|
||||
)
|
||||
snapshots = tuple(copy.deepcopy(task) for task in originals)
|
||||
candidates = tuple(copy.deepcopy(task) for task in snapshots)
|
||||
for task in candidates:
|
||||
attempt = task.active_attempt
|
||||
if attempt is None or attempt.worker_id != worker_id:
|
||||
continue
|
||||
grace = deadline_policy.default_grace_seconds(task.operation)
|
||||
task.mark_disconnected(attempt.attempt_id, grace_seconds=grace, now=stamp)
|
||||
affected.append(task)
|
||||
self._save(task, now=stamp)
|
||||
self._emit("worker_lost", task)
|
||||
self.pool.disconnect(worker_id)
|
||||
return affected
|
||||
task.mark_disconnected(
|
||||
attempt.attempt_id, grace_seconds=grace, now=stamp
|
||||
)
|
||||
return _ReconciliationGeneration(
|
||||
worker_id=worker_id,
|
||||
stamp=stamp,
|
||||
originals=originals,
|
||||
snapshots=snapshots,
|
||||
candidates=candidates,
|
||||
)
|
||||
|
||||
def on_reconnected(
|
||||
self, worker_id: str, *, in_flight: set[str], now: Optional[float] = None
|
||||
self,
|
||||
worker_id: str,
|
||||
*,
|
||||
in_flight: set[str],
|
||||
now: Optional[float] = None,
|
||||
before_persist: Optional[Callable[[object], None]] = None,
|
||||
) -> list[str]:
|
||||
"""Reconcile against what the worker says it is running.
|
||||
|
||||
Returns attempt ids the worker should cancel — work we have already
|
||||
written off, which it must stop burning a GPU on.
|
||||
"""
|
||||
generation = self.prepare_reconnected(
|
||||
worker_id, in_flight=in_flight, now=now
|
||||
)
|
||||
self.persist_reconciliation(
|
||||
generation, before_persist=before_persist
|
||||
)
|
||||
return self.apply_reconnected(generation)
|
||||
|
||||
def prepare_reconnected(
|
||||
self,
|
||||
worker_id: str,
|
||||
*,
|
||||
in_flight: set[str],
|
||||
now: Optional[float] = None,
|
||||
include_task_ids: Optional[set[str]] = None,
|
||||
) -> _ReconciliationGeneration:
|
||||
"""Build, but do not publish, one reconnect reconciliation."""
|
||||
stamp = resolve(now)
|
||||
known = {
|
||||
attempt.attempt_id: attempt
|
||||
@@ -810,7 +1025,20 @@ class Scheduler:
|
||||
for attempt_id in in_flight
|
||||
if attempt_id not in known or known[attempt_id].state.terminal
|
||||
]
|
||||
for task in list(self._tasks.values()):
|
||||
included = include_task_ids or set()
|
||||
originals = tuple(
|
||||
task
|
||||
for task in self._tasks.values()
|
||||
if task.task_id in included
|
||||
or (
|
||||
not task.state.terminal
|
||||
and task.active_attempt is not None
|
||||
and task.active_attempt.worker_id == worker_id
|
||||
)
|
||||
)
|
||||
snapshots = tuple(copy.deepcopy(task) for task in originals)
|
||||
candidates = tuple(copy.deepcopy(task) for task in snapshots)
|
||||
for task in candidates:
|
||||
if task.state.terminal:
|
||||
continue
|
||||
reconcile(
|
||||
@@ -823,8 +1051,84 @@ class Scheduler:
|
||||
resume_lease_seconds=self._budget_for(task).progress_lease_seconds,
|
||||
now=stamp,
|
||||
)
|
||||
self._save(task, now=stamp)
|
||||
return zombies
|
||||
return _ReconciliationGeneration(
|
||||
worker_id=worker_id,
|
||||
stamp=stamp,
|
||||
originals=originals,
|
||||
snapshots=snapshots,
|
||||
candidates=candidates,
|
||||
zombies=tuple(zombies),
|
||||
)
|
||||
|
||||
def persist_reconciliation(
|
||||
self,
|
||||
generation: _ReconciliationGeneration,
|
||||
*,
|
||||
before_persist: Optional[Callable[[object], None]] = None,
|
||||
) -> None:
|
||||
"""Persist a prepared generation without touching the live task graph."""
|
||||
# Nothing in the live graph changes until the complete generation is
|
||||
# durable. A failed write therefore cannot queue duplicate execution
|
||||
# while the old worker is still finishing the original attempt.
|
||||
if self._persist or before_persist is not None:
|
||||
task_store.save_many(
|
||||
generation.candidates if self._persist else [],
|
||||
now=generation.stamp,
|
||||
before_save=before_persist,
|
||||
)
|
||||
|
||||
def reconciliation_is_current(
|
||||
self, generation: _ReconciliationGeneration
|
||||
) -> bool:
|
||||
"""Did the live graph remain unchanged while persistence was off-loop?"""
|
||||
return all(
|
||||
self._tasks.get(original.task_id) is original
|
||||
and original == snapshot
|
||||
for original, snapshot in zip(
|
||||
generation.originals, generation.snapshots
|
||||
)
|
||||
)
|
||||
|
||||
def apply_reconnected(
|
||||
self, generation: _ReconciliationGeneration
|
||||
) -> list[str]:
|
||||
"""Publish a durable reconnect generation on the scheduler's loop."""
|
||||
if not self.reconciliation_is_current(generation):
|
||||
raise RuntimeError("reconnect reconciliation generation changed")
|
||||
changed: list[Task] = []
|
||||
for original, candidate in zip(
|
||||
generation.originals, generation.candidates
|
||||
):
|
||||
if original != candidate:
|
||||
changed.append(original)
|
||||
_adopt_task_state(original, candidate)
|
||||
for task in changed:
|
||||
self._emit(
|
||||
"failed"
|
||||
if task.state.terminal
|
||||
else "requeued"
|
||||
if task.state is TaskState.QUEUED
|
||||
else "resumed",
|
||||
task,
|
||||
)
|
||||
return list(generation.zombies)
|
||||
|
||||
def apply_disconnected(
|
||||
self, generation: _ReconciliationGeneration
|
||||
) -> list[Task]:
|
||||
"""Publish a durable disconnect generation on the scheduler's loop."""
|
||||
if not self.reconciliation_is_current(generation):
|
||||
raise RuntimeError("disconnect reconciliation generation changed")
|
||||
changed: list[Task] = []
|
||||
for original, candidate in zip(
|
||||
generation.originals, generation.candidates
|
||||
):
|
||||
if original != candidate:
|
||||
changed.append(original)
|
||||
_adopt_task_state(original, candidate)
|
||||
for task in changed:
|
||||
self._emit("worker_lost", task)
|
||||
return list(generation.originals)
|
||||
|
||||
def cancel(self, task_id: str, *, reason: str = "cancelled", now: Optional[float] = None) -> bool:
|
||||
task = self._tasks.get(task_id)
|
||||
@@ -997,6 +1301,10 @@ class Scheduler:
|
||||
text=task.params.get("text"),
|
||||
model_resident=bool(worker and worker.is_warm(task.engine, task.model_id)),
|
||||
input_seconds=float(task.params.get("input_seconds") or 0.0),
|
||||
execution_device=(
|
||||
worker.execution_device(task.engine, task.model_id, task.operation)
|
||||
if worker else None
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
||||
+124
-19
@@ -15,6 +15,7 @@ import asyncio
|
||||
import logging
|
||||
import os
|
||||
from typing import Optional
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
from worker.clock import resolve
|
||||
|
||||
@@ -24,10 +25,52 @@ logger = logging.getLogger("omnivoice.worker")
|
||||
_SWEEP_INTERVAL_SECONDS = 5.0
|
||||
# How often the dispatcher looks for queued work it can place.
|
||||
_DISPATCH_INTERVAL_SECONDS = 1.0
|
||||
# Finished rows and their artifacts remain inspectable for a week, then leave
|
||||
# in bounded batches. The loop keeps a long-lived control plane from growing
|
||||
# forever while the startup pass covers apps that are rarely left open.
|
||||
_ARTIFACT_GC_INTERVAL_SECONDS = 60 * 60.0
|
||||
_ARTIFACT_GC_BATCH_SIZE = 250
|
||||
|
||||
DEFAULT_PORT = 7443
|
||||
|
||||
|
||||
class EndpointCertificateError(ValueError):
|
||||
"""An advertised endpoint is not usable with the live TLS certificate."""
|
||||
|
||||
|
||||
def _endpoint_host(endpoint: str) -> str:
|
||||
"""Extract the hostname from the gRPC ``host:port`` enrollment target."""
|
||||
target = (endpoint or "").strip()
|
||||
try:
|
||||
parsed = urlsplit(f"//{target}")
|
||||
host = parsed.hostname
|
||||
port = parsed.port
|
||||
except ValueError as exc:
|
||||
raise EndpointCertificateError(
|
||||
f"Worker endpoint must be host:port; got {endpoint!r}."
|
||||
) from exc
|
||||
if (
|
||||
not host
|
||||
or port is None
|
||||
or parsed.username is not None
|
||||
or parsed.password is not None
|
||||
or parsed.path
|
||||
or parsed.query
|
||||
or parsed.fragment
|
||||
):
|
||||
raise EndpointCertificateError(
|
||||
f"Worker endpoint must be host:port; got {endpoint!r}."
|
||||
)
|
||||
return host
|
||||
|
||||
|
||||
def _format_endpoint(host: str, port: int) -> str:
|
||||
clean_host = (host or "").strip().strip("[]")
|
||||
if ":" in clean_host:
|
||||
clean_host = f"[{clean_host}]"
|
||||
return f"{clean_host}:{port}"
|
||||
|
||||
|
||||
def remote_workers_enabled() -> bool:
|
||||
"""Opt-in gate. Off unless the user turned it on.
|
||||
|
||||
@@ -93,6 +136,7 @@ class ControlPlane:
|
||||
self._server = None
|
||||
self._tasks: list[asyncio.Task] = []
|
||||
self._started = False
|
||||
self._lifecycle_lock = asyncio.Lock()
|
||||
self.startup_error: Optional[str] = None
|
||||
# The port we actually bound, which is not necessarily the configured
|
||||
# one — an enrollment token carries this, so advertising the config
|
||||
@@ -108,6 +152,10 @@ class ControlPlane:
|
||||
return self.credentials.fingerprint if self.credentials else ""
|
||||
|
||||
async def start(self, *, port: Optional[int] = None) -> None:
|
||||
async with self._lifecycle_lock:
|
||||
await self._start(port=port)
|
||||
|
||||
async def _start(self, *, port: Optional[int] = None) -> None:
|
||||
if self._started:
|
||||
return
|
||||
# Imported here, not at module scope: a user who never enables remote
|
||||
@@ -117,25 +165,39 @@ class ControlPlane:
|
||||
from worker.scheduler import Scheduler # noqa: PLC0415
|
||||
from worker.transport.server import WorkerServicer, serve # noqa: PLC0415
|
||||
|
||||
locations = paths()
|
||||
os.makedirs(locations["root"], exist_ok=True)
|
||||
self.credentials = tls.load_or_create(
|
||||
locations["certificate"], locations["private_key"]
|
||||
)
|
||||
self.pool = WorkerPool()
|
||||
self.scheduler = Scheduler(self.pool)
|
||||
# Recover anything that was in flight when the app last quit. The
|
||||
# workers holding those tasks may still be rendering.
|
||||
self.scheduler.restore()
|
||||
|
||||
self.servicer = WorkerServicer(
|
||||
self.scheduler,
|
||||
self.pool,
|
||||
artifact_dir=locations["artifacts"],
|
||||
cert_fingerprint=self.credentials.fingerprint,
|
||||
)
|
||||
self._port = port or control_port()
|
||||
try:
|
||||
locations = paths()
|
||||
os.makedirs(locations["root"], exist_ok=True)
|
||||
wanted_hostnames = tls.default_hostnames()
|
||||
advertised_host = _endpoint_host(self.default_endpoint())
|
||||
if advertised_host not in wanted_hostnames:
|
||||
wanted_hostnames.append(advertised_host)
|
||||
self.credentials = tls.load_or_create(
|
||||
locations["certificate"],
|
||||
locations["private_key"],
|
||||
hostnames=wanted_hostnames,
|
||||
)
|
||||
self.pool = WorkerPool()
|
||||
self.scheduler = Scheduler(self.pool)
|
||||
try:
|
||||
await self._artifact_gc_once(
|
||||
locations["artifacts"], retry_upload_parts=False
|
||||
)
|
||||
except Exception:
|
||||
# Retention is best-effort. A locked file or temporarily busy
|
||||
# database must not make remote compute unavailable.
|
||||
logger.exception("Remote worker startup artifact sweep failed")
|
||||
# Recover anything that was in flight when the app last quit. The
|
||||
# workers holding those tasks may still be rendering.
|
||||
self.scheduler.restore()
|
||||
|
||||
self.servicer = WorkerServicer(
|
||||
self.scheduler,
|
||||
self.pool,
|
||||
artifact_dir=locations["artifacts"],
|
||||
cert_fingerprint=self.credentials.fingerprint,
|
||||
)
|
||||
self._server = await serve(
|
||||
self.servicer,
|
||||
port=self._port,
|
||||
@@ -148,12 +210,20 @@ class ControlPlane:
|
||||
self._tasks = [
|
||||
asyncio.create_task(self._sweep_loop(), name="worker-sweep"),
|
||||
asyncio.create_task(self._dispatch_loop(), name="worker-dispatch"),
|
||||
asyncio.create_task(
|
||||
self._artifact_gc_loop(locations["artifacts"]),
|
||||
name="worker-artifact-gc",
|
||||
),
|
||||
]
|
||||
self._started = True
|
||||
self.startup_error = None
|
||||
logger.info("Remote worker control plane started on port %d", self._port)
|
||||
|
||||
async def stop(self) -> None:
|
||||
async with self._lifecycle_lock:
|
||||
await self._stop()
|
||||
|
||||
async def _stop(self) -> None:
|
||||
for task in self._tasks:
|
||||
task.cancel()
|
||||
if self._tasks:
|
||||
@@ -199,6 +269,29 @@ class ControlPlane:
|
||||
except Exception:
|
||||
logger.exception("Worker sweep failed")
|
||||
|
||||
async def _artifact_gc_once(
|
||||
self, artifact_root: str, *, retry_upload_parts: bool = True
|
||||
) -> int:
|
||||
"""Purge one bounded retention batch without blocking the event loop."""
|
||||
from worker import task_store # noqa: PLC0415
|
||||
|
||||
removed = await asyncio.to_thread(
|
||||
task_store.purge_finished,
|
||||
root=artifact_root,
|
||||
limit=_ARTIFACT_GC_BATCH_SIZE,
|
||||
)
|
||||
if retry_upload_parts and self.servicer is not None:
|
||||
self.servicer.sweep_orphaned_upload_parts()
|
||||
return removed
|
||||
|
||||
async def _artifact_gc_loop(self, artifact_root: str) -> None:
|
||||
while True:
|
||||
await asyncio.sleep(_ARTIFACT_GC_INTERVAL_SECONDS)
|
||||
try:
|
||||
await self._artifact_gc_once(artifact_root)
|
||||
except Exception:
|
||||
logger.exception("Remote worker artifact sweep failed")
|
||||
|
||||
async def _dispatch_loop(self) -> None:
|
||||
"""Place queued work on eligible workers.
|
||||
|
||||
@@ -233,10 +326,21 @@ class ControlPlane:
|
||||
|
||||
def create_enrollment(self, *, endpoint: str = "", label: str = "", ttl_seconds: int = 900):
|
||||
"""Mint a join token carrying this control plane's fingerprint."""
|
||||
from worker import tls # noqa: PLC0415
|
||||
from worker import registry # noqa: PLC0415
|
||||
|
||||
advertised_endpoint = endpoint.strip() or self.default_endpoint()
|
||||
advertised_host = _endpoint_host(advertised_endpoint)
|
||||
if self.credentials is not None and not tls.covers(
|
||||
self.credentials, advertised_host
|
||||
):
|
||||
raise EndpointCertificateError(
|
||||
f"The running certificate does not cover {advertised_host!r}. "
|
||||
"Set OMNIVOICE_WORKER_ENDPOINT_HOST to that hostname and restart "
|
||||
"VoiceStudio before creating this enrollment."
|
||||
)
|
||||
return registry.create_enrollment(
|
||||
endpoint=endpoint or self.default_endpoint(),
|
||||
endpoint=advertised_endpoint,
|
||||
cert_fingerprint=self.fingerprint,
|
||||
label=label,
|
||||
ttl_seconds=ttl_seconds,
|
||||
@@ -262,7 +366,7 @@ class ControlPlane:
|
||||
or tls.primary_ip()
|
||||
or "127.0.0.1"
|
||||
)
|
||||
return f"{host}:{self._port or control_port()}"
|
||||
return _format_endpoint(host, self._port or control_port())
|
||||
|
||||
def snapshot(self, *, now: Optional[float] = None) -> dict:
|
||||
"""Everything the workers UI needs in one call."""
|
||||
@@ -380,6 +484,7 @@ async def stop() -> None:
|
||||
__all__ = [
|
||||
"ControlPlane",
|
||||
"DEFAULT_PORT",
|
||||
"EndpointCertificateError",
|
||||
"control_plane",
|
||||
"control_port",
|
||||
"paths",
|
||||
|
||||
+305
-66
@@ -17,6 +17,7 @@ flips the task to completed, so the ack can only follow a durable fact.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import errno
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
@@ -25,7 +26,8 @@ import os
|
||||
import re
|
||||
import shutil
|
||||
import time
|
||||
from typing import Iterable, Iterator, Optional
|
||||
import uuid
|
||||
from typing import Callable, Iterable, Iterator, Optional
|
||||
|
||||
from core.db import db_conn
|
||||
from core.path_security import UnsafePath, resolve_within, safe_filename
|
||||
@@ -133,6 +135,7 @@ INPUTS_PARAM_KEY = "inputs"
|
||||
|
||||
_HASH_CHUNK_BYTES = 1024 * 1024
|
||||
_SAFE_EXTENSION = re.compile(r"^\.[A-Za-z0-9]{1,8}$")
|
||||
_CONTENT_ARTIFACT = re.compile(r"^([0-9a-f]{64})(?:\.[A-Za-z0-9]{1,8})?$")
|
||||
|
||||
|
||||
class InputStagingError(RuntimeError):
|
||||
@@ -143,6 +146,60 @@ class InputStagingError(RuntimeError):
|
||||
"""
|
||||
|
||||
|
||||
def _fsync_parent_directory(directory: str) -> None:
|
||||
"""Persist directory entry changes where the platform supports it."""
|
||||
directory_flag = getattr(os, "O_DIRECTORY", None)
|
||||
if directory_flag is None:
|
||||
return
|
||||
unsupported = {
|
||||
errno.EINVAL,
|
||||
getattr(errno, "ENOTSUP", errno.EINVAL),
|
||||
getattr(errno, "EOPNOTSUPP", errno.EINVAL),
|
||||
}
|
||||
try:
|
||||
descriptor = os.open(directory, os.O_RDONLY | directory_flag)
|
||||
except OSError as exc:
|
||||
if exc.errno in unsupported:
|
||||
return
|
||||
raise
|
||||
try:
|
||||
os.fsync(descriptor)
|
||||
except OSError as exc:
|
||||
if exc.errno not in unsupported:
|
||||
raise
|
||||
finally:
|
||||
os.close(descriptor)
|
||||
|
||||
|
||||
def _fsync_file(path: str) -> None:
|
||||
with open(path, "r+b") as handle:
|
||||
os.fsync(handle.fileno())
|
||||
|
||||
|
||||
def _durable_makedirs(directory: str) -> None:
|
||||
"""Create a directory hierarchy and persist each parent entry."""
|
||||
target = os.path.abspath(directory)
|
||||
missing: list[str] = []
|
||||
current = target
|
||||
while not os.path.isdir(current):
|
||||
if os.path.exists(current):
|
||||
raise NotADirectoryError(current)
|
||||
missing.append(current)
|
||||
parent = os.path.dirname(current)
|
||||
if parent == current:
|
||||
break
|
||||
current = parent
|
||||
for path in reversed(missing):
|
||||
try:
|
||||
os.mkdir(path)
|
||||
except FileExistsError:
|
||||
if not os.path.isdir(path):
|
||||
raise
|
||||
_fsync_parent_directory(os.path.dirname(path) or ".")
|
||||
if not missing:
|
||||
_fsync_parent_directory(os.path.dirname(target) or ".")
|
||||
|
||||
|
||||
def artifact_root(*, create_dir: bool = True) -> str:
|
||||
"""The directory the control plane serves artifacts from.
|
||||
|
||||
@@ -154,7 +211,7 @@ def artifact_root(*, create_dir: bool = True) -> str:
|
||||
|
||||
root = paths()["artifacts"]
|
||||
if create_dir:
|
||||
os.makedirs(os.path.join(root, INPUTS_DIRNAME), exist_ok=True)
|
||||
_durable_makedirs(os.path.join(root, INPUTS_DIRNAME))
|
||||
return root
|
||||
|
||||
|
||||
@@ -183,7 +240,40 @@ def _digest(path: str) -> tuple[str, int]:
|
||||
return digest.hexdigest(), size
|
||||
|
||||
|
||||
def stage_input(source: str, *, root: Optional[str] = None, now: Optional[float] = None) -> dict:
|
||||
def _staged_entry_matches(path: str, entry: dict) -> bool:
|
||||
"""Verify staged bytes against both metadata and their content address."""
|
||||
if not os.path.isfile(path):
|
||||
return False
|
||||
artifact_id = str(entry.get("artifact_id") or "")
|
||||
portable_name = artifact_id.replace("\\", "/").rsplit("/", 1)[-1]
|
||||
named = _CONTENT_ARTIFACT.fullmatch(portable_name)
|
||||
if named is None:
|
||||
return False
|
||||
try:
|
||||
actual_digest, actual_size = _digest(path)
|
||||
except OSError:
|
||||
return False
|
||||
recorded_digest = str(entry.get("sha256") or "").strip().lower()
|
||||
recorded_size = entry.get("size_bytes")
|
||||
if recorded_digest and actual_digest != recorded_digest:
|
||||
return False
|
||||
if recorded_size is not None:
|
||||
try:
|
||||
if actual_size != int(recorded_size):
|
||||
return False
|
||||
except (TypeError, ValueError):
|
||||
return False
|
||||
if actual_digest != named.group(1):
|
||||
return False
|
||||
# Backfill metadata on a legacy row once its content address proves it.
|
||||
entry["sha256"] = actual_digest
|
||||
entry["size_bytes"] = actual_size
|
||||
return True
|
||||
|
||||
|
||||
def stage_input(
|
||||
source: str, *, root: Optional[str] = None, now: Optional[float] = None
|
||||
) -> dict:
|
||||
"""Copy one input into the artifact store, keyed by its content hash.
|
||||
|
||||
Returns the record that ends up on the task row. ``source`` is kept in it
|
||||
@@ -195,27 +285,58 @@ def stage_input(source: str, *, root: Optional[str] = None, now: Optional[float]
|
||||
try:
|
||||
digest, size = _digest(source)
|
||||
except OSError as exc:
|
||||
raise InputStagingError(f"Could not read the task input {source!r}: {exc}") from exc
|
||||
raise InputStagingError(
|
||||
f"Could not read the task input {source!r}: {exc}"
|
||||
) from exc
|
||||
|
||||
artifact_id = os.path.join(INPUTS_DIRNAME, f"{digest}{_extension(source)}")
|
||||
try:
|
||||
destination = resolve_within(base, artifact_id)
|
||||
except UnsafePath as exc: # pragma: no cover — the id is ours, hex only
|
||||
raise InputStagingError(f"Refusing to stage {source!r} outside the artifact store") from exc
|
||||
raise InputStagingError(
|
||||
f"Refusing to stage {source!r} outside the artifact store"
|
||||
) from exc
|
||||
|
||||
partial = destination.with_name(
|
||||
f".{destination.name}.{uuid.uuid4().hex}.part"
|
||||
)
|
||||
try:
|
||||
destination.parent.mkdir(parents=True, exist_ok=True)
|
||||
# Same size at a content-addressed name means the same bytes: the only
|
||||
# writer is the rename below, so a truncated file cannot exist here.
|
||||
if not (destination.is_file() and destination.stat().st_size == size):
|
||||
partial = destination.with_name(destination.name + ".part")
|
||||
_durable_makedirs(str(destination.parent))
|
||||
expected = {
|
||||
"artifact_id": artifact_id,
|
||||
"sha256": digest,
|
||||
"size_bytes": size,
|
||||
}
|
||||
if not _staged_entry_matches(str(destination), expected):
|
||||
shutil.copyfile(source, partial)
|
||||
copied_digest, copied_size = _digest(str(partial))
|
||||
if copied_digest != digest or copied_size != size:
|
||||
raise InputStagingError(
|
||||
f"The task input {source!r} changed while it was being staged."
|
||||
)
|
||||
_fsync_file(str(partial))
|
||||
os.replace(partial, destination)
|
||||
_fsync_parent_directory(str(destination.parent))
|
||||
# Freshness, not decoration: the purge dates an unreferenced input by
|
||||
# its mtime, so re-using a staged voice has to renew it.
|
||||
os.utime(destination, (stamp, stamp))
|
||||
_fsync_file(str(destination))
|
||||
_fsync_parent_directory(str(destination.parent))
|
||||
except OSError as exc:
|
||||
raise InputStagingError(f"Could not stage the task input {source!r}: {exc}") from exc
|
||||
raise InputStagingError(
|
||||
f"Could not stage the task input {source!r}: {exc}"
|
||||
) from exc
|
||||
finally:
|
||||
try:
|
||||
os.remove(partial)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except OSError:
|
||||
logger.debug(
|
||||
"Could not remove the staged-input partial %s",
|
||||
partial,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
filename = os.path.basename(str(source)) or f"{digest}{_extension(source)}"
|
||||
return {
|
||||
@@ -254,8 +375,13 @@ def ensure_staged(
|
||||
"""
|
||||
params = task.params if isinstance(task.params, dict) else {}
|
||||
recorded = params.get(INPUTS_PARAM_KEY)
|
||||
entries: list[dict] = [e for e in recorded if isinstance(e, dict)] if isinstance(recorded, list) else []
|
||||
if root:
|
||||
entries: list[dict] = (
|
||||
[e for e in recorded if isinstance(e, dict)]
|
||||
if isinstance(recorded, list)
|
||||
else []
|
||||
)
|
||||
base = root or artifact_root()
|
||||
if entries:
|
||||
# A task may have been staged when it was submitted under the default
|
||||
# store, then dispatched by a servicer configured with another store.
|
||||
# Recorded metadata is not proof that this servicer can serve it.
|
||||
@@ -263,7 +389,10 @@ def ensure_staged(
|
||||
for entry in entries:
|
||||
artifact_id = str(entry.get("artifact_id") or "")
|
||||
try:
|
||||
available = bool(artifact_id and resolve_within(root, artifact_id).is_file())
|
||||
path = resolve_within(base, artifact_id)
|
||||
available = bool(
|
||||
artifact_id and _staged_entry_matches(str(path), entry)
|
||||
)
|
||||
except UnsafePath:
|
||||
available = False
|
||||
if available:
|
||||
@@ -271,7 +400,7 @@ def ensure_staged(
|
||||
continue
|
||||
source = str(entry.get("source") or "")
|
||||
if source and os.path.isfile(source):
|
||||
replacement = stage_input(source, root=root, now=now)
|
||||
replacement = stage_input(source, root=base, now=now)
|
||||
replacement.update(key=entry.get("key"), index=entry.get("index"))
|
||||
refreshed.append(replacement)
|
||||
else:
|
||||
@@ -289,7 +418,7 @@ def ensure_staged(
|
||||
# id here. Only what exists on this disk is an input.
|
||||
if not os.path.isfile(value):
|
||||
continue
|
||||
entry = stage_input(value, root=root, now=now)
|
||||
entry = stage_input(value, root=base, now=now)
|
||||
entry["key"] = key
|
||||
entry["index"] = index
|
||||
entries.append(entry)
|
||||
@@ -345,6 +474,63 @@ def _referenced_artifacts(conn) -> set[str]:
|
||||
return referenced
|
||||
|
||||
|
||||
def _purge_result_directories(
|
||||
task_ids: Iterable[str], *, root: Optional[str] = None
|
||||
) -> tuple[list[str], int]:
|
||||
"""Delete result directories before their owning rows become unreachable."""
|
||||
task_ids = list(task_ids)
|
||||
try:
|
||||
base = root or artifact_root(create_dir=False)
|
||||
except Exception: # pragma: no cover — no data dir at all
|
||||
logger.debug("No artifact root to purge", exc_info=True)
|
||||
return [], 0
|
||||
if not os.path.isdir(base):
|
||||
return task_ids, 0
|
||||
|
||||
cleaned: list[str] = []
|
||||
removed = 0
|
||||
for task_id in task_ids:
|
||||
try:
|
||||
path = resolve_within(base, safe_filename(task_id))
|
||||
except UnsafePath:
|
||||
continue
|
||||
if not os.path.exists(path):
|
||||
try:
|
||||
_fsync_parent_directory(base)
|
||||
except OSError:
|
||||
logger.debug(
|
||||
"Could not persist task artifact cleanup at %s",
|
||||
base,
|
||||
exc_info=True,
|
||||
)
|
||||
continue
|
||||
cleaned.append(task_id)
|
||||
continue
|
||||
if not os.path.isdir(path):
|
||||
logger.warning("Refusing to purge non-directory task artifact %s", path)
|
||||
continue
|
||||
try:
|
||||
shutil.rmtree(path)
|
||||
except OSError:
|
||||
logger.debug("Could not purge task artifacts at %s", path, exc_info=True)
|
||||
continue
|
||||
try:
|
||||
_fsync_parent_directory(base)
|
||||
except OSError:
|
||||
# The bytes are gone from this process's view, but the directory
|
||||
# deletion is not a crash-durable fact yet. Keep the DB row as the
|
||||
# retry index until a later sweep can establish that barrier.
|
||||
logger.debug(
|
||||
"Could not persist task artifact cleanup at %s",
|
||||
base,
|
||||
exc_info=True,
|
||||
)
|
||||
continue
|
||||
cleaned.append(task_id)
|
||||
removed += 1
|
||||
return cleaned, removed
|
||||
|
||||
|
||||
def purge_artifacts(
|
||||
task_ids: Iterable[str], referenced: set[str], *, cutoff: float, root: Optional[str] = None
|
||||
) -> int:
|
||||
@@ -356,7 +542,7 @@ def purge_artifacts(
|
||||
rows were judged by. Nothing here raises — a purge that fails is a disk
|
||||
that stays fuller than we wanted, not a failed request.
|
||||
"""
|
||||
removed = 0
|
||||
_cleaned, removed = _purge_result_directories(task_ids, root=root)
|
||||
try:
|
||||
base = root or artifact_root(create_dir=False)
|
||||
except Exception: # pragma: no cover — no data dir at all
|
||||
@@ -365,15 +551,6 @@ def purge_artifacts(
|
||||
if not os.path.isdir(base):
|
||||
return 0
|
||||
|
||||
for task_id in task_ids:
|
||||
try:
|
||||
path = resolve_within(base, safe_filename(task_id))
|
||||
except UnsafePath:
|
||||
continue
|
||||
if os.path.isdir(path):
|
||||
shutil.rmtree(path, ignore_errors=True)
|
||||
removed += 1
|
||||
|
||||
inputs_dir = os.path.join(base, INPUTS_DIRNAME)
|
||||
try:
|
||||
names = os.listdir(inputs_dir)
|
||||
@@ -484,6 +661,25 @@ def _upsert_attempts(conn, task: Task) -> None:
|
||||
)
|
||||
|
||||
|
||||
def _save_with_conn(conn, task: Task, *, stamp: float) -> None:
|
||||
conn.execute(
|
||||
"UPDATE remote_tasks SET state=?, excluded_json=?, error_json=?, result_ref=?, "
|
||||
"updated_at=?, deadline_at=?, finished_at=?, pinned_worker_id=? WHERE id=?",
|
||||
(
|
||||
task.state.value,
|
||||
json.dumps(sorted(task.excluded_workers)),
|
||||
_dump_error(task.error),
|
||||
task.result_ref,
|
||||
stamp,
|
||||
task.deadline_at,
|
||||
task.finished_at,
|
||||
task.pinned_worker_id,
|
||||
task.task_id,
|
||||
),
|
||||
)
|
||||
_upsert_attempts(conn, task)
|
||||
|
||||
|
||||
def save(task: Task, *, now: Optional[float] = None) -> None:
|
||||
"""Write the whole task + attempt graph.
|
||||
|
||||
@@ -492,22 +688,22 @@ def save(task: Task, *, now: Optional[float] = None) -> None:
|
||||
"""
|
||||
stamp = resolve(now)
|
||||
with db_conn() as conn:
|
||||
conn.execute(
|
||||
"UPDATE remote_tasks SET state=?, excluded_json=?, error_json=?, result_ref=?, "
|
||||
"updated_at=?, deadline_at=?, finished_at=?, pinned_worker_id=? WHERE id=?",
|
||||
(
|
||||
task.state.value,
|
||||
json.dumps(sorted(task.excluded_workers)),
|
||||
_dump_error(task.error),
|
||||
task.result_ref,
|
||||
stamp,
|
||||
task.deadline_at,
|
||||
task.finished_at,
|
||||
task.pinned_worker_id,
|
||||
task.task_id,
|
||||
),
|
||||
)
|
||||
_upsert_attempts(conn, task)
|
||||
_save_with_conn(conn, task, stamp=stamp)
|
||||
|
||||
|
||||
def save_many(
|
||||
tasks: Iterable[Task],
|
||||
*,
|
||||
now: Optional[float] = None,
|
||||
before_save: Optional[Callable[[object], None]] = None,
|
||||
) -> None:
|
||||
"""Persist one reconciliation generation atomically."""
|
||||
stamp = resolve(now)
|
||||
with db_conn() as conn:
|
||||
if before_save is not None:
|
||||
before_save(conn)
|
||||
for task in tasks:
|
||||
_save_with_conn(conn, task, stamp=stamp)
|
||||
|
||||
|
||||
def commit_result(
|
||||
@@ -584,23 +780,30 @@ def load_unfinished() -> list[Task]:
|
||||
still be rendering, and reconciliation decides each one's fate once the
|
||||
workers reconnect.
|
||||
"""
|
||||
live = ", ".join(f"'{s.value}'" for s in TaskState if not s.terminal)
|
||||
live = json.dumps([s.value for s in TaskState if not s.terminal])
|
||||
with db_conn() as conn:
|
||||
rows = conn.execute(
|
||||
f"SELECT * FROM remote_tasks WHERE state IN ({live}) ORDER BY priority ASC, created_at ASC"
|
||||
"""
|
||||
SELECT * FROM remote_tasks
|
||||
WHERE state IN (SELECT value FROM json_each(?))
|
||||
ORDER BY priority ASC, created_at ASC
|
||||
""",
|
||||
(live,),
|
||||
).fetchall()
|
||||
return [_row_to_task(r, _attempts_for(conn, r["id"])) for r in rows]
|
||||
|
||||
|
||||
def list_tasks(*, states: Optional[Iterable[TaskState]] = None, limit: int = 100) -> list[Task]:
|
||||
sql = "SELECT * FROM remote_tasks"
|
||||
params: list = []
|
||||
if states:
|
||||
placeholders = ", ".join("?" for _ in states)
|
||||
sql += f" WHERE state IN ({placeholders})"
|
||||
params.extend(s.value for s in states)
|
||||
sql += " ORDER BY created_at DESC LIMIT ?"
|
||||
params.append(limit)
|
||||
sql = """
|
||||
SELECT * FROM remote_tasks
|
||||
WHERE state IN (SELECT value FROM json_each(?))
|
||||
ORDER BY created_at DESC LIMIT ?
|
||||
"""
|
||||
params = (json.dumps([s.value for s in states]), limit)
|
||||
else:
|
||||
sql = "SELECT * FROM remote_tasks ORDER BY created_at DESC LIMIT ?"
|
||||
params = (limit,)
|
||||
with db_conn() as conn:
|
||||
rows = conn.execute(sql, params).fetchall()
|
||||
return [_row_to_task(r, _attempts_for(conn, r["id"])) for r in rows]
|
||||
@@ -611,6 +814,7 @@ def purge_finished(
|
||||
older_than_seconds: float = 7 * 24 * 3600,
|
||||
now: Optional[float] = None,
|
||||
root: Optional[str] = None,
|
||||
limit: Optional[int] = None,
|
||||
) -> int:
|
||||
"""Drop old finished tasks — rows *and* the bytes they own.
|
||||
|
||||
@@ -619,31 +823,66 @@ def purge_finished(
|
||||
audio. Neither was ever deleted, so the feature grew the user's disk for
|
||||
as long as they used it.
|
||||
"""
|
||||
if limit is not None and limit <= 0:
|
||||
return 0
|
||||
cutoff = resolve(now) - older_than_seconds
|
||||
terminal = ", ".join(f"'{s.value}'" for s in TaskState if s.terminal)
|
||||
terminal = json.dumps([s.value for s in TaskState if s.terminal])
|
||||
with db_conn() as conn:
|
||||
doomed = [
|
||||
row["id"]
|
||||
for row in conn.execute(
|
||||
f"SELECT id FROM remote_tasks WHERE state IN ({terminal}) AND finished_at < ?",
|
||||
(cutoff,),
|
||||
"""
|
||||
SELECT id FROM remote_tasks
|
||||
WHERE state IN (SELECT value FROM json_each(?))
|
||||
AND finished_at < ?
|
||||
ORDER BY finished_at ASC, id ASC
|
||||
LIMIT ?
|
||||
""",
|
||||
(terminal, cutoff, -1 if limit is None else int(limit)),
|
||||
).fetchall()
|
||||
]
|
||||
conn.execute(
|
||||
f"DELETE FROM remote_task_attempts WHERE task_id IN "
|
||||
f"(SELECT id FROM remote_tasks WHERE state IN ({terminal}) AND finished_at < ?)",
|
||||
(cutoff,),
|
||||
)
|
||||
cur = conn.execute(
|
||||
f"DELETE FROM remote_tasks WHERE state IN ({terminal}) AND finished_at < ?", (cutoff,)
|
||||
)
|
||||
removed = cur.rowcount
|
||||
# Results are attempt-scoped. Delete them before their task rows, so a
|
||||
# crash or transient Windows lock cannot erase the only index from which a
|
||||
# future sweep could find those bytes.
|
||||
cleaned, _artifacts_removed = _purge_result_directories(doomed, root=root)
|
||||
with db_conn() as conn:
|
||||
eligible: list[str] = []
|
||||
if cleaned:
|
||||
eligible = [
|
||||
row["id"]
|
||||
for row in conn.execute(
|
||||
"""
|
||||
SELECT id FROM remote_tasks
|
||||
WHERE id IN (SELECT value FROM json_each(?))
|
||||
AND state IN (SELECT value FROM json_each(?))
|
||||
AND finished_at < ?
|
||||
""",
|
||||
(json.dumps(cleaned), terminal, cutoff),
|
||||
).fetchall()
|
||||
]
|
||||
removed = 0
|
||||
if eligible:
|
||||
conn.execute(
|
||||
"""
|
||||
DELETE FROM remote_task_attempts
|
||||
WHERE task_id IN (SELECT value FROM json_each(?))
|
||||
""",
|
||||
(json.dumps(eligible),),
|
||||
)
|
||||
cur = conn.execute(
|
||||
"""
|
||||
DELETE FROM remote_tasks
|
||||
WHERE id IN (SELECT value FROM json_each(?))
|
||||
""",
|
||||
(json.dumps(eligible),),
|
||||
)
|
||||
removed = cur.rowcount
|
||||
# Read the survivors inside the same transaction that deleted the
|
||||
# rows: an input is only unreferenced relative to what is left.
|
||||
referenced = _referenced_artifacts(conn)
|
||||
# Filesystem work outside the transaction — a slow rmtree must not hold
|
||||
# SQLite's write lock against the dispatch loop.
|
||||
purge_artifacts(doomed, referenced, cutoff=cutoff, root=root)
|
||||
# Shared content-addressed inputs remain discoverable without their old
|
||||
# task row, so they can be swept after the transaction.
|
||||
purge_artifacts((), referenced, cutoff=cutoff, root=root)
|
||||
return removed
|
||||
|
||||
|
||||
|
||||
+296
-34
@@ -20,15 +20,21 @@ There is deliberately no way to disable verification.
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime as _dt
|
||||
import errno
|
||||
import ipaddress
|
||||
import logging
|
||||
import os
|
||||
import socket
|
||||
import ssl
|
||||
import tempfile
|
||||
import threading
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from cryptography import x509
|
||||
from cryptography.exceptions import InvalidSignature, UnsupportedAlgorithm
|
||||
from cryptography.hazmat.primitives import hashes, serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import ec
|
||||
from cryptography.x509.oid import NameOID
|
||||
@@ -39,6 +45,8 @@ logger = logging.getLogger("omnivoice.worker")
|
||||
|
||||
_CERT_VALID_DAYS = 825 # the CA/Browser Forum maximum; long enough to be quiet
|
||||
_RENEW_WITHIN_DAYS = 30
|
||||
_PENDING_PAIR_MAGIC = b"omnivoice-tls-pair-v1\n"
|
||||
_PROCESS_CREDENTIAL_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def unverified_client_context() -> ssl.SSLContext:
|
||||
@@ -72,7 +80,7 @@ def _san_entries(hostnames: list[str]) -> list[x509.GeneralName]:
|
||||
"""
|
||||
entries: list[x509.GeneralName] = []
|
||||
for host in hostnames:
|
||||
host = (host or "").strip()
|
||||
host = _normalise_host(host)
|
||||
if not host:
|
||||
continue
|
||||
try:
|
||||
@@ -118,9 +126,23 @@ def covers(credentials: "ServerCredentials", host: str) -> bool:
|
||||
).value
|
||||
except Exception:
|
||||
return False
|
||||
names = set(san.get_values_for_type(x509.DNSName))
|
||||
names |= {str(ip) for ip in san.get_values_for_type(x509.IPAddress)}
|
||||
return host in names
|
||||
names = {_normalise_host(name) for name in san.get_values_for_type(x509.DNSName)}
|
||||
names |= {
|
||||
_normalise_host(str(ip)) for ip in san.get_values_for_type(x509.IPAddress)
|
||||
}
|
||||
return _normalise_host(host) in names
|
||||
|
||||
|
||||
def _normalise_host(host: str) -> str:
|
||||
"""Canonicalise a SAN identity for comparison, not for DNS resolution."""
|
||||
candidate = (host or "").strip().strip("[]")
|
||||
try:
|
||||
return str(ipaddress.ip_address(candidate))
|
||||
except ValueError:
|
||||
# DNS names are case-insensitive, and a trailing root dot does not
|
||||
# name a different host. cryptography intentionally does no matching
|
||||
# for us because the certificate is its own trust root.
|
||||
return candidate.rstrip(".").lower()
|
||||
|
||||
|
||||
def default_hostnames() -> list[str]:
|
||||
@@ -206,25 +228,50 @@ def load_or_create(
|
||||
desktop control plane presents as "my workers all went offline" with
|
||||
nothing in the UI explaining why.
|
||||
"""
|
||||
existing = _load(cert_path, key_path)
|
||||
wanted = hostnames or default_hostnames()
|
||||
# Every explicitly requested identity must be present. This matters for an
|
||||
# inbound listener bound to a user-entered LAN address: keeping a stable
|
||||
# certificate that does not name that address makes mandatory hostname
|
||||
# verification fail even though its fingerprint is correct.
|
||||
missing = [
|
||||
host for host in wanted if existing is not None and not covers(existing, host)
|
||||
]
|
||||
if existing is not None and not _expiring_soon(existing) and not missing:
|
||||
return existing
|
||||
if existing is not None:
|
||||
reason = (
|
||||
"expiring" if _expiring_soon(existing) else "missing a requested hostname"
|
||||
)
|
||||
logger.info("Control-plane certificate is %s — regenerating.", reason)
|
||||
credentials = generate_self_signed(hostnames=wanted)
|
||||
_save(cert_path, key_path, credentials)
|
||||
return credentials
|
||||
with _credential_lock(cert_path):
|
||||
existing = _load(cert_path, key_path)
|
||||
pending_present, pending = _load_pending_pair(cert_path, key_path)
|
||||
if existing is not None and pending_present:
|
||||
# A valid final pair is either the old generation (the transaction
|
||||
# never began replacing it) or the fully committed new one. Both
|
||||
# are safer than rotating a pin merely because cleanup was cut
|
||||
# short. The lock makes it safe to remove the stale journal here.
|
||||
_remove_pending_pair(cert_path, key_path)
|
||||
elif existing is None and pending is not None:
|
||||
logger.warning(
|
||||
"Recovering an interrupted control-plane TLS credential update."
|
||||
)
|
||||
_install_pair(cert_path, key_path, pending)
|
||||
_remove_pending_pair(cert_path, key_path)
|
||||
existing = pending
|
||||
|
||||
wanted = hostnames or default_hostnames()
|
||||
# Every explicitly requested identity must be present. This matters for
|
||||
# a listener bound to a user-entered address: keeping a stable
|
||||
# certificate that does not name it makes mandatory hostname
|
||||
# verification fail even though its fingerprint is correct.
|
||||
missing = [
|
||||
host
|
||||
for host in wanted
|
||||
if existing is not None and not covers(existing, host)
|
||||
]
|
||||
if existing is not None and not _expiring_soon(existing) and not missing:
|
||||
return existing
|
||||
if existing is not None:
|
||||
reason = (
|
||||
"expiring"
|
||||
if _expiring_soon(existing)
|
||||
else "missing a requested hostname"
|
||||
)
|
||||
logger.info("Control-plane certificate is %s — regenerating.", reason)
|
||||
elif pending_present and pending is None:
|
||||
logger.warning(
|
||||
"Ignoring a corrupt interrupted TLS credential update and "
|
||||
"generating a new pair."
|
||||
)
|
||||
credentials = generate_self_signed(hostnames=wanted)
|
||||
_save(cert_path, key_path, credentials)
|
||||
return credentials
|
||||
|
||||
|
||||
def _load(cert_path: str, key_path: str) -> Optional[ServerCredentials]:
|
||||
@@ -233,11 +280,40 @@ def _load(cert_path: str, key_path: str) -> Optional[ServerCredentials]:
|
||||
cert_pem = fh.read()
|
||||
with open(key_path, "rb") as fh:
|
||||
key_pem = fh.read()
|
||||
except (FileNotFoundError, PermissionError):
|
||||
except FileNotFoundError:
|
||||
return None
|
||||
return _credentials_from_pem(cert_pem, key_pem)
|
||||
|
||||
|
||||
def _credentials_from_pem(
|
||||
cert_pem: bytes, key_pem: bytes
|
||||
) -> Optional[ServerCredentials]:
|
||||
"""Parse a credential pair only when the private key belongs to the cert."""
|
||||
try:
|
||||
certificate = x509.load_pem_x509_certificate(cert_pem)
|
||||
except ValueError:
|
||||
private_key = serialization.load_pem_private_key(key_pem, password=None)
|
||||
certificate_public_key = certificate.public_key().public_bytes(
|
||||
serialization.Encoding.DER,
|
||||
serialization.PublicFormat.SubjectPublicKeyInfo,
|
||||
)
|
||||
private_public_key = private_key.public_key().public_bytes(
|
||||
serialization.Encoding.DER,
|
||||
serialization.PublicFormat.SubjectPublicKeyInfo,
|
||||
)
|
||||
certificate.public_key().verify(
|
||||
certificate.signature,
|
||||
certificate.tbs_certificate_bytes,
|
||||
ec.ECDSA(certificate.signature_hash_algorithm),
|
||||
)
|
||||
except (
|
||||
AttributeError,
|
||||
InvalidSignature,
|
||||
TypeError,
|
||||
UnsupportedAlgorithm,
|
||||
ValueError,
|
||||
):
|
||||
return None
|
||||
if certificate_public_key != private_public_key:
|
||||
return None
|
||||
return ServerCredentials(
|
||||
certificate_pem=cert_pem,
|
||||
@@ -258,19 +334,205 @@ def _expiring_soon(
|
||||
|
||||
|
||||
def _save(cert_path: str, key_path: str, credentials: ServerCredentials) -> None:
|
||||
os.makedirs(os.path.dirname(os.path.abspath(cert_path)), exist_ok=True)
|
||||
with open(cert_path, "wb") as fh:
|
||||
fh.write(credentials.certificate_pem)
|
||||
# The private key gets the same 0600 treatment as worker keys.
|
||||
fd = os.open(key_path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
|
||||
try:
|
||||
os.write(fd, credentials.private_key_pem)
|
||||
finally:
|
||||
os.close(fd)
|
||||
"""Durably replace a pair, leaving enough state to finish after a crash.
|
||||
|
||||
No filesystem atomically renames two independent paths. A mode-0600
|
||||
journal is therefore made durable first; the loader uses it only when the
|
||||
final paths are torn or corrupt. Each final path is itself an atomic
|
||||
sibling rename, and the journal is removed only after both directory
|
||||
entries and their contents are durable.
|
||||
"""
|
||||
pending_path = _pending_pair_path(cert_path, key_path)
|
||||
_atomic_write(pending_path, _encode_pending_pair(credentials))
|
||||
_install_pair(cert_path, key_path, credentials)
|
||||
_remove_pending_pair(cert_path, key_path)
|
||||
|
||||
|
||||
def _install_pair(
|
||||
cert_path: str, key_path: str, credentials: ServerCredentials
|
||||
) -> None:
|
||||
# Key first, certificate second: the public certificate is the commit
|
||||
# marker. A crash between them is detected by _load and recovered from the
|
||||
# already-durable pending pair.
|
||||
_atomic_write(key_path, credentials.private_key_pem)
|
||||
try:
|
||||
os.chmod(key_path, 0o600)
|
||||
except OSError:
|
||||
# Windows and some network filesystems do not honour POSIX modes; the
|
||||
# key remains in the app's per-user data directory there.
|
||||
pass
|
||||
_atomic_write(cert_path, credentials.certificate_pem)
|
||||
installed = _load(cert_path, key_path)
|
||||
if installed is None or installed.fingerprint != credentials.fingerprint:
|
||||
raise OSError("TLS credential pair failed verification after persistence")
|
||||
|
||||
|
||||
def _atomic_write(path: str, payload: bytes) -> None:
|
||||
"""Fsync and atomically replace one mode-restricted credential file."""
|
||||
directory = os.path.dirname(os.path.abspath(path))
|
||||
_durable_makedirs(directory)
|
||||
fd, temporary = tempfile.mkstemp(
|
||||
prefix=f".{os.path.basename(path)}.", suffix=".tmp", dir=directory
|
||||
)
|
||||
try:
|
||||
with os.fdopen(fd, "wb") as fh:
|
||||
fh.write(payload)
|
||||
fh.flush()
|
||||
os.fsync(fh.fileno())
|
||||
os.replace(temporary, path)
|
||||
_fsync_parent_directory(directory)
|
||||
except Exception:
|
||||
try:
|
||||
os.unlink(temporary)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
raise
|
||||
|
||||
|
||||
def _pending_pair_path(cert_path: str, key_path: str) -> str:
|
||||
cert_name = os.path.basename(cert_path)
|
||||
key_name = os.path.basename(key_path)
|
||||
directory = os.path.dirname(os.path.abspath(cert_path))
|
||||
return os.path.join(directory, f".{cert_name}.{key_name}.pending")
|
||||
|
||||
|
||||
def _encode_pending_pair(credentials: ServerCredentials) -> bytes:
|
||||
certificate = credentials.certificate_pem
|
||||
return (
|
||||
_PENDING_PAIR_MAGIC
|
||||
+ str(len(certificate)).encode("ascii")
|
||||
+ b"\n"
|
||||
+ certificate
|
||||
+ credentials.private_key_pem
|
||||
)
|
||||
|
||||
|
||||
def _load_pending_pair(
|
||||
cert_path: str, key_path: str
|
||||
) -> tuple[bool, Optional[ServerCredentials]]:
|
||||
path = _pending_pair_path(cert_path, key_path)
|
||||
try:
|
||||
with open(path, "rb") as fh:
|
||||
raw = fh.read()
|
||||
except FileNotFoundError:
|
||||
return False, None
|
||||
if not raw.startswith(_PENDING_PAIR_MAGIC):
|
||||
return True, None
|
||||
size_line, separator, payload = raw[len(_PENDING_PAIR_MAGIC) :].partition(b"\n")
|
||||
if not separator:
|
||||
return True, None
|
||||
try:
|
||||
certificate_size = int(size_line)
|
||||
except ValueError:
|
||||
return True, None
|
||||
if certificate_size <= 0 or certificate_size >= len(payload):
|
||||
return True, None
|
||||
return True, _credentials_from_pem(
|
||||
payload[:certificate_size], payload[certificate_size:]
|
||||
)
|
||||
|
||||
|
||||
def _remove_pending_pair(cert_path: str, key_path: str) -> None:
|
||||
path = _pending_pair_path(cert_path, key_path)
|
||||
try:
|
||||
os.unlink(path)
|
||||
except FileNotFoundError:
|
||||
return
|
||||
_fsync_parent_directory(os.path.dirname(path) or ".")
|
||||
|
||||
|
||||
def _fsync_parent_directory(directory: str) -> None:
|
||||
"""Make a preceding directory-entry change durable where supported."""
|
||||
directory_flag = getattr(os, "O_DIRECTORY", None)
|
||||
if directory_flag is None:
|
||||
return
|
||||
unsupported = {
|
||||
errno.EINVAL,
|
||||
getattr(errno, "ENOTSUP", errno.EINVAL),
|
||||
getattr(errno, "EOPNOTSUPP", errno.EINVAL),
|
||||
}
|
||||
try:
|
||||
descriptor = os.open(directory, os.O_RDONLY | directory_flag)
|
||||
except OSError as exc:
|
||||
if exc.errno in unsupported:
|
||||
return
|
||||
raise
|
||||
try:
|
||||
os.fsync(descriptor)
|
||||
except OSError as exc:
|
||||
if exc.errno not in unsupported:
|
||||
raise
|
||||
finally:
|
||||
os.close(descriptor)
|
||||
|
||||
|
||||
def _durable_makedirs(directory: str) -> None:
|
||||
"""Persist every newly-created credential-directory entry."""
|
||||
target = os.path.abspath(directory)
|
||||
missing: list[str] = []
|
||||
current = target
|
||||
while not os.path.isdir(current):
|
||||
if os.path.exists(current):
|
||||
if os.path.isdir(current):
|
||||
break
|
||||
raise NotADirectoryError(current)
|
||||
missing.append(current)
|
||||
parent = os.path.dirname(current)
|
||||
if parent == current:
|
||||
break
|
||||
current = parent
|
||||
|
||||
for path in reversed(missing):
|
||||
try:
|
||||
os.mkdir(path)
|
||||
except FileExistsError:
|
||||
if not os.path.isdir(path):
|
||||
raise
|
||||
_fsync_parent_directory(os.path.dirname(path) or ".")
|
||||
|
||||
if not missing:
|
||||
# Retry the exact barrier that may have failed after a prior mkdir.
|
||||
_fsync_parent_directory(os.path.dirname(target) or ".")
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _credential_lock(cert_path: str) -> Iterator[None]:
|
||||
"""Serialise first-start/renewal across threads and app processes."""
|
||||
directory = os.path.dirname(os.path.abspath(cert_path))
|
||||
_durable_makedirs(directory)
|
||||
lock_path = os.path.join(directory, f".{os.path.basename(cert_path)}.lock")
|
||||
with _PROCESS_CREDENTIAL_LOCK:
|
||||
descriptor = os.open(lock_path, os.O_RDWR | os.O_CREAT, 0o600)
|
||||
is_windows = os.name == "nt"
|
||||
acquired = False
|
||||
try:
|
||||
if is_windows:
|
||||
import msvcrt # noqa: PLC0415
|
||||
|
||||
if os.fstat(descriptor).st_size == 0:
|
||||
os.write(descriptor, b"\0")
|
||||
os.fsync(descriptor)
|
||||
os.lseek(descriptor, 0, os.SEEK_SET)
|
||||
msvcrt.locking(descriptor, msvcrt.LK_LOCK, 1)
|
||||
else:
|
||||
import fcntl # noqa: PLC0415
|
||||
|
||||
fcntl.flock(descriptor, fcntl.LOCK_EX)
|
||||
acquired = True
|
||||
yield
|
||||
finally:
|
||||
try:
|
||||
if acquired and is_windows:
|
||||
import msvcrt # noqa: PLC0415
|
||||
|
||||
os.lseek(descriptor, 0, os.SEEK_SET)
|
||||
msvcrt.locking(descriptor, msvcrt.LK_UNLCK, 1)
|
||||
elif acquired:
|
||||
import fcntl # noqa: PLC0415
|
||||
|
||||
fcntl.flock(descriptor, fcntl.LOCK_UN)
|
||||
finally:
|
||||
os.close(descriptor)
|
||||
|
||||
|
||||
def pin_matches(certificate_der: bytes, expected_fingerprint: str) -> bool:
|
||||
|
||||
@@ -48,8 +48,10 @@ from typing import Awaitable, Callable, Optional, Protocol
|
||||
|
||||
import grpc
|
||||
|
||||
from worker.async_utils import drain_task, to_thread_and_drain_on_cancel
|
||||
from worker import errors as worker_errors
|
||||
from worker import identity
|
||||
from worker.capacity import clamp_concurrency
|
||||
from worker.errors import ErrorClass, WorkerError
|
||||
from worker.identity import EnrollmentToken, WorkerKeypair
|
||||
from worker.protocol.gen import worker_v1_pb2 as pb
|
||||
@@ -103,6 +105,8 @@ class ArtifactTransport(Protocol):
|
||||
|
||||
async def stage_in(self, ref: pb.ArtifactRef, destination: str) -> None: ...
|
||||
|
||||
def result_acked(self, artifacts: list[pb.ArtifactRef]) -> None: ...
|
||||
|
||||
# Used when an assignment carries no lease (an older control plane, or a test).
|
||||
# Mirrors deadlines.py's _HEARTBEAT_GRACE_S * 4.
|
||||
_DEFAULT_PROGRESS_LEASE_SECONDS = 120.0
|
||||
@@ -115,6 +119,29 @@ _MIN_KEEPALIVE_INTERVAL_SECONDS = 0.05
|
||||
_EXECUTOR_KWARGS = frozenset({"on_progress", "on_model_loading", "fetch_input"})
|
||||
|
||||
|
||||
def _write_all(handle, payload: bytes) -> None:
|
||||
"""Write a complete chunk, including through short-writing file wrappers."""
|
||||
remaining = memoryview(payload)
|
||||
while remaining:
|
||||
written = handle.write(remaining)
|
||||
if written is None or written <= 0:
|
||||
raise OSError("input destination made no write progress")
|
||||
remaining = remaining[written:]
|
||||
|
||||
|
||||
def _close_and_remove(handle, destination: str) -> None:
|
||||
"""Finish file cleanup as one blocking operation after cancellation."""
|
||||
if handle is not None:
|
||||
with contextlib.suppress(OSError):
|
||||
handle.close()
|
||||
with contextlib.suppress(OSError):
|
||||
os.remove(destination)
|
||||
|
||||
|
||||
class TerminalRegistrationError(RuntimeError):
|
||||
"""A registration failure that reconnecting cannot repair."""
|
||||
|
||||
|
||||
def keepalive_interval(lease_seconds: float) -> float:
|
||||
"""How often a running task must renew its progress lease.
|
||||
|
||||
@@ -267,28 +294,44 @@ class WorkerClient:
|
||||
cancel: Optional[Callable[[str], Awaitable[None]]] = None,
|
||||
capability_probe: Optional[Callable[[], list[dict]]] = None,
|
||||
on_registered: Optional[Callable[[str], None]] = None,
|
||||
on_activated: Optional[Callable[[str], None]] = None,
|
||||
artifacts: Optional["ArtifactTransport"] = None,
|
||||
drain_active_work: Optional[Callable[[], Awaitable[None]]] = None,
|
||||
) -> None:
|
||||
self.config = config
|
||||
self.config.max_concurrent_tasks = clamp_concurrency(
|
||||
self.config.max_concurrent_tasks
|
||||
)
|
||||
self._execute = execute
|
||||
self._cancel = cancel
|
||||
self._capability_probe = capability_probe
|
||||
# Discovery imports engine adapters and inspects model storage. Share
|
||||
# one off-loop probe across reconnect/task/prewarm/idle refresh races;
|
||||
# a cancelled waiter drains it before returning so no detached probe
|
||||
# can mutate global engine state after authority is gone.
|
||||
self._capability_probe_task: Optional[asyncio.Task] = None
|
||||
# Outbound mode moves artifacts with RPCs this side initiates
|
||||
# (UploadResult / DownloadArtifact), which is only possible because
|
||||
# this side dialled. In inbound mode the node cannot call the panel at
|
||||
# all, so both directions are driven from the panel and this hook
|
||||
# swaps in the staging that makes that work. None means outbound.
|
||||
self._artifacts = artifacts
|
||||
self._drain_active_work = drain_active_work
|
||||
# Lets the agent persist the server-assigned id. Without it a restarted
|
||||
# worker signs its challenge with an empty worker_id, the signature
|
||||
# never matches, and reconnecting needs a fresh enrollment token —
|
||||
# which would make key-based identity pointless.
|
||||
self._on_registered = on_registered
|
||||
# Register only reserves a provisional server generation. Readiness is
|
||||
# published separately, after ConfigUpdate proves Control activated it.
|
||||
self._on_activated = on_activated
|
||||
self._activation_confirmed = False
|
||||
self._reporter_kwargs = _accepted_reporter_kwargs(execute)
|
||||
self._outbox = _Outbox()
|
||||
self._pending: dict[str, PendingResult] = {}
|
||||
self._running: dict[str, asyncio.Task] = {}
|
||||
self._keepalives: dict[str, asyncio.Task] = {}
|
||||
self._maintenance: set[asyncio.Task] = set()
|
||||
self._epoch = 0
|
||||
self._session_token = ""
|
||||
# Negotiated by ConfigUpdate; None means "use the executor's own
|
||||
@@ -298,6 +341,12 @@ class WorkerClient:
|
||||
# same channel rather than the control stream.
|
||||
self._stub = None
|
||||
self._stop = asyncio.Event()
|
||||
# Drain is a graceful reconnect, not terminal shutdown. This event
|
||||
# half-closes only the current stream after every active result is ACKed
|
||||
# while ``_stop`` remains reserved for cancelling the agent itself.
|
||||
self._reconnect_requested = asyncio.Event()
|
||||
self._draining = False
|
||||
self._accepting_assignments = True
|
||||
|
||||
# ── Connection ────────────────────────────────────────────────────────
|
||||
|
||||
@@ -318,6 +367,8 @@ class WorkerClient:
|
||||
|
||||
async def run_forever(self) -> None:
|
||||
"""Connect, serve, and reconnect until stopped."""
|
||||
if not self._stop.is_set():
|
||||
self._accepting_assignments = True
|
||||
attempt = 0
|
||||
while not self._stop.is_set():
|
||||
try:
|
||||
@@ -325,6 +376,13 @@ class WorkerClient:
|
||||
attempt = 0
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except TerminalRegistrationError:
|
||||
# The control plane has made a durable decision that this
|
||||
# identity may not reconnect. Work deliberately survives an
|
||||
# ordinary network drop, but must not survive revocation and
|
||||
# keep using the GPU with no authority able to cancel it.
|
||||
await self._cancel_active_work()
|
||||
raise
|
||||
except Exception as exc:
|
||||
attempt += 1
|
||||
delay = backoff_delay(attempt)
|
||||
@@ -338,6 +396,52 @@ class WorkerClient:
|
||||
|
||||
async def stop(self) -> None:
|
||||
self._stop.set()
|
||||
self._reconnect_requested.set()
|
||||
await self._cancel_active_work()
|
||||
|
||||
async def _cancel_active_work(self) -> None:
|
||||
"""Cancel retained tasks without permanently disabling reconnect."""
|
||||
# Close admission before taking any snapshots. Attach/Control may still
|
||||
# have a frame ready while revocation drains an uninterruptible GPU
|
||||
# call; accepting that frame here lets it escape the snapshot entirely.
|
||||
self._accepting_assignments = False
|
||||
maintenance = list(self._maintenance)
|
||||
for task in maintenance:
|
||||
task.cancel()
|
||||
running = list(self._running.items())
|
||||
keepalives = list(self._keepalives.values())
|
||||
for key, task in running:
|
||||
self._stop_keepalive(key)
|
||||
task.cancel()
|
||||
for keepalive in keepalives:
|
||||
keepalive.cancel()
|
||||
cancel_callbacks = []
|
||||
if self._cancel is not None:
|
||||
cancel_callbacks = [
|
||||
asyncio.create_task(self._cancel(key.split("/")[0]))
|
||||
for key, _task in running
|
||||
]
|
||||
# Cancel every maintenance, task, and keepalive wrapper before awaiting
|
||||
# any uninterruptible one. A prewarm stuck in a model-load thread must
|
||||
# not delay revocation of active user renders.
|
||||
draining = [
|
||||
*maintenance,
|
||||
*(task for _key, task in running),
|
||||
*keepalives,
|
||||
*cancel_callbacks,
|
||||
]
|
||||
if draining:
|
||||
await asyncio.gather(
|
||||
*draining, return_exceptions=True
|
||||
)
|
||||
self._maintenance.clear()
|
||||
for key, task in running:
|
||||
if self._running.get(key) is task:
|
||||
self._running.pop(key, None)
|
||||
if self._drain_active_work is not None:
|
||||
await self._drain_active_work()
|
||||
self._keepalives.clear()
|
||||
self._pending.clear()
|
||||
|
||||
async def _connect_once(self) -> None:
|
||||
async with self._channel() as channel:
|
||||
@@ -359,6 +463,8 @@ class WorkerClient:
|
||||
try:
|
||||
async for message in stream:
|
||||
await self._on_server_message(message)
|
||||
if self._stop.is_set():
|
||||
break
|
||||
finally:
|
||||
heartbeat.cancel()
|
||||
# The channel closes with this block, so a stub kept past it
|
||||
@@ -375,7 +481,27 @@ class WorkerClient:
|
||||
# before; inbound calls them from the Attach handler. Neither mode gets its
|
||||
# own copy of registration, zombie reconciliation or redelivery.
|
||||
|
||||
def build_register_request(self) -> pb.RegisterRequest:
|
||||
async def _probe_capabilities(self) -> list[dict]:
|
||||
if self._capability_probe is None:
|
||||
return list(self.config.capabilities or [])
|
||||
task = self._capability_probe_task
|
||||
if task is None or task.done():
|
||||
task = asyncio.create_task(
|
||||
to_thread_and_drain_on_cancel(self._capability_probe),
|
||||
name="worker-capability-probe",
|
||||
)
|
||||
self._capability_probe_task = task
|
||||
try:
|
||||
capabilities = await asyncio.shield(task)
|
||||
except asyncio.CancelledError:
|
||||
await drain_task(task)
|
||||
raise
|
||||
finally:
|
||||
if task.done() and self._capability_probe_task is task:
|
||||
self._capability_probe_task = None
|
||||
return list(capabilities or [])
|
||||
|
||||
async def build_register_request(self) -> pb.RegisterRequest:
|
||||
"""This worker's self-description. Identical in both modes."""
|
||||
challenge = identity.new_challenge()
|
||||
nonce = identity.new_challenge()
|
||||
@@ -387,9 +513,8 @@ class WorkerClient:
|
||||
nonce=nonce,
|
||||
)
|
||||
)
|
||||
capabilities = (
|
||||
self._capability_probe() if self._capability_probe else self.config.capabilities
|
||||
)
|
||||
capabilities = await self._probe_capabilities()
|
||||
self.config.capabilities = capabilities
|
||||
return pb.RegisterRequest(
|
||||
envelope=pb.Envelope(sequence=self._epoch),
|
||||
protocol_version_min=PROTOCOL_VERSION,
|
||||
@@ -403,7 +528,9 @@ class WorkerClient:
|
||||
key_id=self.config.keypair.key_id,
|
||||
host=codec.host_to_pb(self.config.host or describe_host()),
|
||||
capabilities=[codec.capability_to_pb(c) for c in capabilities],
|
||||
max_concurrent_tasks=self.config.max_concurrent_tasks,
|
||||
max_concurrent_tasks=clamp_concurrency(
|
||||
self.config.max_concurrent_tasks
|
||||
),
|
||||
in_flight=[
|
||||
codec.task_ref(t.split("/")[0], t.split("/")[1], self._epoch)
|
||||
for t in self._running
|
||||
@@ -415,27 +542,80 @@ class WorkerClient:
|
||||
async def accept_registration(self, response: pb.RegisterResponse) -> None:
|
||||
"""Adopt the control plane's answer and recover in-flight state."""
|
||||
if response.error.code:
|
||||
raise RuntimeError(f"{response.error.code}: {response.error.message}")
|
||||
raise TerminalRegistrationError(
|
||||
f"{response.error.code}: {response.error.message}"
|
||||
)
|
||||
# An enrollment token is already spent when this response arrives.
|
||||
# Commit the reconnect identity before adopting the live session; if
|
||||
# local durable state cannot be written, retrying the spent token can
|
||||
# never repair the worker and must reach the caller immediately.
|
||||
if self._on_registered is not None:
|
||||
try:
|
||||
await to_thread_and_drain_on_cancel(
|
||||
self._on_registered, response.worker_id
|
||||
)
|
||||
except Exception as exc:
|
||||
raise TerminalRegistrationError(
|
||||
"LOCAL_STATE: accepted enrollment could not be persisted"
|
||||
) from exc
|
||||
|
||||
self._draining = False
|
||||
self._reconnect_requested.clear()
|
||||
if not self._stop.is_set():
|
||||
self._accepting_assignments = True
|
||||
self._activation_confirmed = False
|
||||
self._epoch = response.session_epoch
|
||||
self._session_token = response.session_token
|
||||
self.config.worker_id = response.worker_id
|
||||
# The token is spent; every later connection proves key possession.
|
||||
self.config.enrollment_token = ""
|
||||
if self._on_registered is not None:
|
||||
try:
|
||||
self._on_registered(response.worker_id)
|
||||
except Exception:
|
||||
logger.warning("Could not persist the worker id", exc_info=True)
|
||||
|
||||
authoritative = {ref.attempt_id for ref in response.authoritative_in_flight}
|
||||
await self._cancel_zombies(authoritative)
|
||||
await self._redeliver_pending()
|
||||
|
||||
def confirm_activation(self) -> None:
|
||||
"""Publish readiness once Control proves the provisional session live."""
|
||||
if self._activation_confirmed:
|
||||
return
|
||||
if self._on_activated is not None:
|
||||
try:
|
||||
self._on_activated(self.config.worker_id)
|
||||
except Exception as exc:
|
||||
raise TerminalRegistrationError(
|
||||
"LOCAL_STATE: activated enrollment could not be published"
|
||||
) from exc
|
||||
self._activation_confirmed = True
|
||||
|
||||
async def next_outbound(self) -> pb.WorkerMessage:
|
||||
"""The next frame this worker wants to send."""
|
||||
return await self._outbox.get()
|
||||
|
||||
def prepare_inbound_session(self) -> None:
|
||||
"""Start a fresh stream while retaining running and pending work.
|
||||
|
||||
An inbound listener creates the protocol owner once per panel key, not
|
||||
once per transport generation. Frames queued for the dead stream are
|
||||
stale, but `_running` and `_pending` are precisely the state the next
|
||||
Register must reconcile and redeliver.
|
||||
"""
|
||||
self._outbox = _Outbox()
|
||||
self._session_token = ""
|
||||
self._stub = None
|
||||
self._draining = False
|
||||
self._reconnect_requested.clear()
|
||||
if not self._stop.is_set():
|
||||
self._accepting_assignments = True
|
||||
|
||||
@property
|
||||
def reconnect_requested(self) -> bool:
|
||||
return self._reconnect_requested.is_set()
|
||||
|
||||
@property
|
||||
def outbound_pending(self) -> bool:
|
||||
"""Whether a terminal/control frame still needs transport delivery."""
|
||||
return not self._outbox.empty()
|
||||
|
||||
def start_heartbeat(self, response: pb.RegisterResponse) -> asyncio.Task:
|
||||
"""Begin the heartbeat this session's liveness depends on.
|
||||
|
||||
@@ -458,14 +638,27 @@ class WorkerClient:
|
||||
await self._on_server_message(message)
|
||||
|
||||
async def _register(self, stub) -> pb.RegisterResponse:
|
||||
return await stub.Register(self.build_register_request())
|
||||
return await stub.Register(await self.build_register_request())
|
||||
|
||||
# ── Outbound ──────────────────────────────────────────────────────────
|
||||
|
||||
async def _outbound(self):
|
||||
while True:
|
||||
message = await self._outbox.get()
|
||||
yield message
|
||||
while not self._reconnect_requested.is_set():
|
||||
message = asyncio.create_task(self._outbox.get())
|
||||
reconnect = asyncio.create_task(self._reconnect_requested.wait())
|
||||
try:
|
||||
done, _pending = await asyncio.wait(
|
||||
{message, reconnect}, return_when=asyncio.FIRST_COMPLETED
|
||||
)
|
||||
if message not in done:
|
||||
return
|
||||
yield message.result()
|
||||
self._maybe_finish_drain()
|
||||
finally:
|
||||
for task in (message, reconnect):
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
await asyncio.gather(message, reconnect, return_exceptions=True)
|
||||
|
||||
async def _send(self, message: pb.WorkerMessage, *, bulk: bool = False) -> None:
|
||||
"""Enqueue a frame. ``bulk`` is for anything that can be large.
|
||||
@@ -479,17 +672,19 @@ class WorkerClient:
|
||||
async def _heartbeat_loop(self, interval: float) -> None:
|
||||
while True:
|
||||
await asyncio.sleep(interval)
|
||||
await self._send(
|
||||
pb.WorkerMessage(
|
||||
heartbeat=pb.Heartbeat(
|
||||
active_tasks=len(self._running),
|
||||
available_slots=max(
|
||||
0, self.config.max_concurrent_tasks - len(self._running)
|
||||
),
|
||||
resident_models=self._resident_models(),
|
||||
)
|
||||
)
|
||||
await self._send(self.heartbeat_message())
|
||||
|
||||
def heartbeat_message(self) -> pb.WorkerMessage:
|
||||
"""Build the worker's current liveness/capacity frame."""
|
||||
return pb.WorkerMessage(
|
||||
heartbeat=pb.Heartbeat(
|
||||
active_tasks=len(self._running),
|
||||
available_slots=max(
|
||||
0, self.config.max_concurrent_tasks - len(self._running)
|
||||
),
|
||||
resident_models=self._resident_models(),
|
||||
)
|
||||
)
|
||||
|
||||
def _resident_models(self) -> list[str]:
|
||||
return [
|
||||
@@ -503,7 +698,7 @@ class WorkerClient:
|
||||
if self._capability_probe is None:
|
||||
return
|
||||
try:
|
||||
capabilities = self._capability_probe()
|
||||
capabilities = await self._probe_capabilities()
|
||||
self.config.capabilities = capabilities
|
||||
await self._send(pb.WorkerMessage(capabilities=pb.CapabilityUpdate(
|
||||
capabilities=[codec.capability_to_pb(c) for c in capabilities]
|
||||
@@ -531,9 +726,14 @@ class WorkerClient:
|
||||
# keepalive for an attempt the server has disowned is exactly the
|
||||
# frame that resurrects a cancelled task.
|
||||
self._stop_keepalive(key)
|
||||
task = self._running.pop(key, None)
|
||||
task = self._running.get(key)
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
# CancelAck releases server capacity. Do not send it while an
|
||||
# uninterruptible synthesis/load thread still owns the GPU, or the
|
||||
# replacement assignment can overlap and corrupt or OOM the worker.
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
self._running.pop(key, None)
|
||||
if self._cancel is not None:
|
||||
await self._cancel(key.split("/")[0])
|
||||
|
||||
@@ -550,24 +750,76 @@ class WorkerClient:
|
||||
)
|
||||
elif kind == "result_ack":
|
||||
# Only now is it safe to forget the result.
|
||||
self._pending.pop(self._key(message.result_ack.ref), None)
|
||||
pending = self._pending.pop(self._key(message.result_ack.ref), None)
|
||||
if pending is not None and self._artifacts is not None:
|
||||
result_acked_async = getattr(
|
||||
self._artifacts, "result_acked_async", None
|
||||
)
|
||||
if callable(result_acked_async):
|
||||
await result_acked_async(pending.artifacts)
|
||||
else:
|
||||
result_acked = getattr(self._artifacts, "result_acked", None)
|
||||
if callable(result_acked):
|
||||
result_acked(pending.artifacts)
|
||||
self._maybe_finish_drain()
|
||||
elif kind == "config":
|
||||
if message.config.max_concurrent_tasks:
|
||||
self.config.max_concurrent_tasks = message.config.max_concurrent_tasks
|
||||
self.config.max_concurrent_tasks = clamp_concurrency(
|
||||
message.config.max_concurrent_tasks
|
||||
)
|
||||
if message.config.inline_result_threshold_bytes:
|
||||
# Negotiated, so the two sides cannot drift: the control plane
|
||||
# is the one that knows how much it is willing to take on the
|
||||
# control stream, and it may lower this at any time.
|
||||
self._inline_threshold = int(message.config.inline_result_threshold_bytes)
|
||||
self.confirm_activation()
|
||||
elif kind == "ping":
|
||||
# Answer immediately; the server times the round trip.
|
||||
await self._send(pb.WorkerMessage(pong=pb.Pong(nonce=message.ping.nonce)))
|
||||
elif kind == "drain":
|
||||
self._stop.set()
|
||||
self._accepting_assignments = False
|
||||
self._draining = True
|
||||
if message.drain.reconnect_to:
|
||||
self.config.endpoint = message.drain.reconnect_to
|
||||
self._maybe_finish_drain()
|
||||
elif kind == "shutdown":
|
||||
await self._cancel_active_work()
|
||||
self._stop.set()
|
||||
if self._artifacts is not None:
|
||||
await self._send(
|
||||
pb.WorkerMessage(
|
||||
goodbye=pb.WorkerGoodbye(
|
||||
reason="The control-plane connection was removed."
|
||||
)
|
||||
)
|
||||
)
|
||||
# Publish the terminal acknowledgement before asking the inbound
|
||||
# Attach loop to leave. Reversing these lets that loop observe
|
||||
# reconnect_requested and close the stream with Goodbye still
|
||||
# queued, so the control plane cannot prove remote work drained.
|
||||
self._reconnect_requested.set()
|
||||
elif kind == "prewarm":
|
||||
asyncio.create_task(self._on_prewarm(message.prewarm), name="worker-prewarm")
|
||||
if not self._accepting_assignments:
|
||||
return
|
||||
task = asyncio.create_task(
|
||||
self._on_prewarm(message.prewarm), name="worker-prewarm"
|
||||
)
|
||||
self._maintenance.add(task)
|
||||
task.add_done_callback(self._maintenance_finished)
|
||||
|
||||
def _maintenance_finished(self, task: asyncio.Task) -> None:
|
||||
self._maintenance.discard(task)
|
||||
self._maybe_finish_drain()
|
||||
|
||||
def _maybe_finish_drain(self) -> None:
|
||||
if (
|
||||
self._draining
|
||||
and not self._running
|
||||
and not self._pending
|
||||
and not self._maintenance
|
||||
and self._outbox.empty()
|
||||
):
|
||||
self._reconnect_requested.set()
|
||||
|
||||
async def _on_prewarm(self, request: pb.PrewarmRequest) -> None:
|
||||
"""Load/download a catalog model, then report the resulting capability."""
|
||||
@@ -591,15 +843,20 @@ class WorkerClient:
|
||||
await self._install_catalog_repo(repo_ids[0])
|
||||
from worker.executor import TaskExecutor # noqa: PLC0415
|
||||
|
||||
await asyncio.to_thread(TaskExecutor._load_backend, engine)
|
||||
await to_thread_and_drain_on_cancel(TaskExecutor._load_backend, engine)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception:
|
||||
logger.warning("Prewarm failed for %s", engine or request.model_id, exc_info=True)
|
||||
finally:
|
||||
await self.refresh_capabilities()
|
||||
await self.refresh_capabilities()
|
||||
|
||||
async def _install_catalog_repo(self, repo_id: str) -> None:
|
||||
"""Run the existing setup installer and pipe its hf_progress upstream."""
|
||||
from api.routers.setup.download import InstallModelRequest, install_model # noqa: PLC0415
|
||||
from api.routers.setup.download import ( # noqa: PLC0415
|
||||
InstallModelRequest,
|
||||
cancel_install_and_wait,
|
||||
install_model,
|
||||
)
|
||||
from utils import download_aggregator, hf_progress # noqa: PLC0415
|
||||
|
||||
hf_progress.install()
|
||||
@@ -627,12 +884,20 @@ class WorkerClient:
|
||||
loop.call_soon_threadsafe(_finish)
|
||||
|
||||
listener_id = hf_progress.register_listener(listener)
|
||||
started_install = False
|
||||
install_completed = False
|
||||
try:
|
||||
await install_model(InstallModelRequest(repo_id=repo_id, target="local"))
|
||||
response = await install_model(
|
||||
InstallModelRequest(repo_id=repo_id, target="local")
|
||||
)
|
||||
started_install = response.get("status") == "install_started"
|
||||
event = await asyncio.wait_for(terminal, timeout=_FALLBACK_MODEL_LOAD_SECONDS)
|
||||
if event.get("phase") != "install_done":
|
||||
raise RuntimeError(event.get("error") or "model install did not complete")
|
||||
install_completed = True
|
||||
finally:
|
||||
if started_install and not install_completed:
|
||||
await cancel_install_and_wait(repo_id)
|
||||
hf_progress.unregister_listener(listener_id)
|
||||
|
||||
@staticmethod
|
||||
@@ -641,6 +906,20 @@ class WorkerClient:
|
||||
|
||||
async def _on_assignment(self, assignment: pb.TaskAssignment) -> None:
|
||||
key = self._key(assignment.ref)
|
||||
if not self._accepting_assignments or self._stop.is_set():
|
||||
await self._send(
|
||||
pb.WorkerMessage(
|
||||
rejected=pb.TaskRejected(
|
||||
ref=assignment.ref,
|
||||
error=pb.Error(
|
||||
error_class=pb.ERROR_CLASS_TRANSIENT,
|
||||
code="WORKER_STOPPING",
|
||||
message="The worker is relinquishing this control plane.",
|
||||
),
|
||||
)
|
||||
)
|
||||
)
|
||||
return
|
||||
if len(self._running) >= self.config.max_concurrent_tasks:
|
||||
# Declining because we are full is normal and penalty-free; the
|
||||
# scheduler's view of our capacity is only ever advisory.
|
||||
@@ -672,9 +951,11 @@ class WorkerClient:
|
||||
# release the slot too, or the reserved task keeps running work
|
||||
# the scheduler never saw accepted — and double-executes after
|
||||
# reassignment. The stream-death case lands here as well.
|
||||
task = self._running.pop(key, None)
|
||||
task = self._running.get(key)
|
||||
if task is not None:
|
||||
task.cancel()
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
self._running.pop(key, None)
|
||||
raise
|
||||
|
||||
async def _run(self, assignment: pb.TaskAssignment) -> None:
|
||||
@@ -745,14 +1026,32 @@ class WorkerClient:
|
||||
failure: WorkerError = (
|
||||
exc.error if isinstance(exc, TaskFailure) else worker_errors.from_exception(exc)
|
||||
)
|
||||
if (
|
||||
failure.error_class is ErrorClass.TIMEOUT
|
||||
and self._drain_active_work is not None
|
||||
):
|
||||
# A timeout only ends the coroutine's wait; Python cannot
|
||||
# cancel the GPU thread underneath it. Do not send FAILED or
|
||||
# release this admission slot until the executor proves that
|
||||
# work relinquished the device, otherwise a capacity-1 worker
|
||||
# can accept a replacement on top of the timed-out render.
|
||||
await self._drain_active_work()
|
||||
await self._fail(assignment.ref, failure)
|
||||
finally:
|
||||
# Also covers the abnormal exits — cancellation, a crash between
|
||||
# the two _stop_keepalive calls above — so the timer can never
|
||||
# outlive the task that owns it.
|
||||
self._stop_keepalive(key)
|
||||
self._running.pop(key, None)
|
||||
await self.refresh_capabilities()
|
||||
# Discovery may inspect the just-used backend. Keep the admission
|
||||
# reservation until that off-loop probe has drained, otherwise a
|
||||
# heartbeat can advertise the slot while this generation still
|
||||
# owns task-finalization work. The inner finally also releases it
|
||||
# when shutdown cancels the refresh waiter.
|
||||
try:
|
||||
await self.refresh_capabilities()
|
||||
finally:
|
||||
self._running.pop(key, None)
|
||||
self._maybe_finish_drain()
|
||||
|
||||
async def _fail(self, ref: pb.TaskRef, error: WorkerError) -> None:
|
||||
await self._send(
|
||||
@@ -855,8 +1154,12 @@ class WorkerClient:
|
||||
await self._report_upload(ref, 0.0)
|
||||
|
||||
offset = 0
|
||||
metadata = ((SESSION_METADATA_KEY, self._session_token),)
|
||||
for _ in range(_MAX_UPLOAD_RESUMES):
|
||||
ack = await stub.UploadResult(self._result_chunks(ref, artifact, payload, offset))
|
||||
ack = await stub.UploadResult(
|
||||
self._result_chunks(ref, artifact, payload, offset),
|
||||
metadata=metadata,
|
||||
)
|
||||
if ack.committed:
|
||||
break
|
||||
resumed = int(ack.bytes_received)
|
||||
@@ -967,25 +1270,29 @@ class WorkerClient:
|
||||
request.session_token = self._session_token
|
||||
offset = 0
|
||||
complete = False
|
||||
handle = None
|
||||
try:
|
||||
with open(destination, "wb") as handle:
|
||||
async for chunk in self._stub.DownloadArtifact(request):
|
||||
if int(chunk.offset) != offset:
|
||||
raise RuntimeError(
|
||||
f"input offset {chunk.offset} did not match {offset} bytes received"
|
||||
)
|
||||
handle.write(chunk.data)
|
||||
offset += len(chunk.data)
|
||||
if chunk.last:
|
||||
complete = True
|
||||
break
|
||||
handle = await to_thread_and_drain_on_cancel(open, destination, "wb")
|
||||
async for chunk in self._stub.DownloadArtifact(request):
|
||||
if int(chunk.offset) != offset:
|
||||
raise RuntimeError(
|
||||
f"input offset {chunk.offset} did not match {offset} bytes received"
|
||||
)
|
||||
await to_thread_and_drain_on_cancel(_write_all, handle, chunk.data)
|
||||
offset += len(chunk.data)
|
||||
if chunk.last:
|
||||
complete = True
|
||||
break
|
||||
await to_thread_and_drain_on_cancel(handle.close)
|
||||
handle = None
|
||||
except asyncio.CancelledError:
|
||||
await to_thread_and_drain_on_cancel(_close_and_remove, handle, destination)
|
||||
raise
|
||||
except Exception:
|
||||
with contextlib.suppress(OSError):
|
||||
os.remove(destination)
|
||||
await to_thread_and_drain_on_cancel(_close_and_remove, handle, destination)
|
||||
raise
|
||||
if not complete:
|
||||
with contextlib.suppress(OSError):
|
||||
os.remove(destination)
|
||||
await to_thread_and_drain_on_cancel(_close_and_remove, None, destination)
|
||||
raise RuntimeError("input download ended before its final chunk")
|
||||
|
||||
def _executor_kwargs(self, assignment: pb.TaskAssignment) -> dict[str, Callable]:
|
||||
@@ -1120,6 +1427,7 @@ __all__ = [
|
||||
"MAX_MESSAGE_BYTES",
|
||||
"UPLOAD_STAGE",
|
||||
"PendingResult",
|
||||
"TerminalRegistrationError",
|
||||
"WorkerClient",
|
||||
"WorkerConfig",
|
||||
"backoff_delay",
|
||||
|
||||
@@ -10,7 +10,7 @@ import logging
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
from worker.capacity import derive_concurrency
|
||||
from worker.capacity import clamp_concurrency, derive_concurrency
|
||||
from worker.deadlines import Deadlines
|
||||
from worker.errors import ErrorClass, WorkerError
|
||||
from worker.lifecycle import Attempt, PriorityClass, Task
|
||||
@@ -225,7 +225,7 @@ def capability_to_pb(cap: dict) -> pb.ModelCapability:
|
||||
resident=bool(cap.get("resident")),
|
||||
min_memory_bytes=int(cap.get("min_memory_bytes") or 0),
|
||||
precision=str(cap.get("precision") or ""),
|
||||
derived_concurrency=max(0, declared),
|
||||
derived_concurrency=clamp_concurrency(declared, allow_zero=True),
|
||||
cpu_fallback=bool(cap.get("cpu_fallback")),
|
||||
repo_ids=list(cap.get("repo_ids") or []),
|
||||
display_name=str(cap.get("display_name") or ""),
|
||||
@@ -243,7 +243,9 @@ def capability_from_pb(message: pb.ModelCapability) -> dict:
|
||||
"resident": message.resident,
|
||||
"min_memory_bytes": message.min_memory_bytes,
|
||||
"precision": message.precision,
|
||||
"derived_concurrency": message.derived_concurrency,
|
||||
"derived_concurrency": clamp_concurrency(
|
||||
message.derived_concurrency, allow_zero=True
|
||||
),
|
||||
"cpu_fallback": message.cpu_fallback,
|
||||
"repo_ids": list(message.repo_ids),
|
||||
"display_name": message.display_name,
|
||||
|
||||
+2349
-186
File diff suppressed because it is too large
Load Diff
+2
-1
@@ -7,7 +7,8 @@ VoiceStudio installer. Today:
|
||||
|------|------------|---------|
|
||||
| `omnivoice-tts-darwin-arm64` | `ServeurpersoCom/omnivoice.cpp` @ pinned SHA | GGUF inference runtime — Apple Silicon |
|
||||
| `omnivoice-tts-darwin-x86_64` | same | Intel Mac |
|
||||
| `omnivoice-tts-linux-x86_64` | same | Linux |
|
||||
| `omnivoice-tts-linux-x86_64` | same | Linux (x86_64) |
|
||||
| `omnivoice-tts-linux-aarch64` | same | Linux ARM64 — Apple Silicon under Asahi; built with GGML Vulkan where the toolchain supports it, so the Honeykrisp driver can accelerate generation |
|
||||
| `omnivoice-tts-windows-x86_64.exe` | same | Windows |
|
||||
| `checksums.sha256` | computed by `scripts/build-omnivoice-tts.sh` | SHA-256 manifest — verified by `VoiceStudioGGUFBackend.is_available()` |
|
||||
|
||||
|
||||
@@ -5,9 +5,11 @@
|
||||
# docker compose -f deploy/docker-compose.yml --profile cpu up # CPU mode
|
||||
# docker compose -f deploy/docker-compose.yml --profile gpu up # NVIDIA GPU
|
||||
# docker compose -f deploy/docker-compose.yml --profile rocm up # AMD GPU (ROCm)
|
||||
# docker compose -f deploy/docker-compose.yml --profile worker-gpu up # NVIDIA worker
|
||||
# docker compose -f deploy/docker-compose.yml --profile worker-rocm up # AMD worker
|
||||
#
|
||||
# All services bind to port 3900, so they MUST be opt-in via profiles —
|
||||
# otherwise `compose up` would race them and one would fail to bind.
|
||||
# The Studio services publish the UI on loopback and the authenticated TLS
|
||||
# worker control plane on 7443. Worker-only services publish no host port.
|
||||
#
|
||||
# First run downloads ~4 GB of models. Progress is shown in logs.
|
||||
# Open http://localhost:3900 once the health check passes.
|
||||
@@ -15,9 +17,9 @@
|
||||
# SECURITY: The port is bound to 127.0.0.1 by default — only this
|
||||
# machine can reach the API. To expose VoiceStudio on your LAN (or
|
||||
# through a reverse proxy / tunnel), change the port mapping to
|
||||
# "0.0.0.0:3900:3900" or "3900:3900". VoiceStudio itself ships no
|
||||
# authentication — if you expose it, put it behind a reverse proxy
|
||||
# with auth (Caddy basic_auth, nginx + htpasswd, Tailscale, etc.).
|
||||
# "0.0.0.0:3900:3900" or "3900:3900". Export a long random
|
||||
# OMNIVOICE_API_KEY before starting a Studio profile; use HTTPS or an
|
||||
# encrypted private overlay whenever traffic leaves a fully trusted LAN.
|
||||
# ──────────────────────────────────────────────────────────────
|
||||
|
||||
services:
|
||||
@@ -33,6 +35,7 @@ services:
|
||||
profiles: ["cpu"]
|
||||
ports:
|
||||
- "127.0.0.1:3900:3900"
|
||||
- "${OMNIVOICE_WORKER_PUBLISH_HOST:-127.0.0.1}:${OMNIVOICE_WORKER_PORT:-7443}:${OMNIVOICE_WORKER_PORT:-7443}"
|
||||
volumes:
|
||||
- omnivoice-data:/app/omnivoice_data
|
||||
environment:
|
||||
@@ -53,6 +56,11 @@ services:
|
||||
# discoverable. If you front the container with your own auth proxy on
|
||||
# loopback, set this to 0 to re-enable the strict gate.
|
||||
- OMNIVOICE_SERVER_MODE=1
|
||||
# Required for server-mode settings, diagnostics, and other admin
|
||||
# mutations because Docker NAT cannot prove a browser is loopback.
|
||||
- OMNIVOICE_API_KEY=${OMNIVOICE_API_KEY:?export a long random OMNIVOICE_API_KEY before starting a Studio profile}
|
||||
- OMNIVOICE_WORKER_PORT=${OMNIVOICE_WORKER_PORT:-7443}
|
||||
- OMNIVOICE_WORKER_ENDPOINT_HOST=${OMNIVOICE_WORKER_ENDPOINT_HOST:-}
|
||||
healthcheck:
|
||||
test: ["CMD", "curl", "-sf", "http://localhost:3900/health"]
|
||||
interval: 30s
|
||||
@@ -71,6 +79,7 @@ services:
|
||||
profiles: ["gpu"]
|
||||
ports:
|
||||
- "127.0.0.1:3900:3900"
|
||||
- "${OMNIVOICE_WORKER_PUBLISH_HOST:-127.0.0.1}:${OMNIVOICE_WORKER_PORT:-7443}:${OMNIVOICE_WORKER_PORT:-7443}"
|
||||
volumes:
|
||||
- omnivoice-data:/app/omnivoice_data
|
||||
environment:
|
||||
@@ -86,6 +95,9 @@ services:
|
||||
# See the CPU service above — relaxes the loopback origin gate for the
|
||||
# headless Docker deployment (issue #261). Set to 0 to re-enable it.
|
||||
- OMNIVOICE_SERVER_MODE=1
|
||||
- OMNIVOICE_API_KEY=${OMNIVOICE_API_KEY:?export a long random OMNIVOICE_API_KEY before starting a Studio profile}
|
||||
- OMNIVOICE_WORKER_PORT=${OMNIVOICE_WORKER_PORT:-7443}
|
||||
- OMNIVOICE_WORKER_ENDPOINT_HOST=${OMNIVOICE_WORKER_ENDPOINT_HOST:-}
|
||||
healthcheck:
|
||||
test: ["CMD", "curl", "-sf", "http://localhost:3900/health"]
|
||||
interval: 30s
|
||||
@@ -121,6 +133,7 @@ services:
|
||||
profiles: ["rocm"]
|
||||
ports:
|
||||
- "127.0.0.1:3900:3900"
|
||||
- "${OMNIVOICE_WORKER_PUBLISH_HOST:-127.0.0.1}:${OMNIVOICE_WORKER_PORT:-7443}:${OMNIVOICE_WORKER_PORT:-7443}"
|
||||
devices:
|
||||
- /dev/kfd
|
||||
- /dev/dri
|
||||
@@ -136,6 +149,9 @@ services:
|
||||
# loopback origin gate for the headless Docker deployment.
|
||||
- OMNIVOICE_BIND_HOST=0.0.0.0
|
||||
- OMNIVOICE_SERVER_MODE=1
|
||||
- OMNIVOICE_API_KEY=${OMNIVOICE_API_KEY:?export a long random OMNIVOICE_API_KEY before starting a Studio profile}
|
||||
- OMNIVOICE_WORKER_PORT=${OMNIVOICE_WORKER_PORT:-7443}
|
||||
- OMNIVOICE_WORKER_ENDPOINT_HOST=${OMNIVOICE_WORKER_ENDPOINT_HOST:-}
|
||||
# RDNA3 consumer cards (RX 7900 XTX/XT and friends, gfx1100): if the
|
||||
# GPU is not detected, uncomment the override below. The backend
|
||||
# auto-sets it for known consumer GFX IDs, so try without it first.
|
||||
@@ -148,5 +164,79 @@ services:
|
||||
start_period: 180s
|
||||
restart: unless-stopped
|
||||
|
||||
# ── Worker-only NVIDIA GPU mode
|
||||
# Generate a join code on the control plane, then start with:
|
||||
# OMNIVOICE_WORKER_TOKEN='ovw_…' docker compose \
|
||||
# -f deploy/docker-compose.yml --profile worker-gpu up -d
|
||||
# No port is published: uvicorn only hosts the application lifespan that
|
||||
# owns the outbound worker agent. The browser UI is not needed or exposed.
|
||||
omnivoice-worker-gpu:
|
||||
image: ghcr.io/debpalash/omnivoice-studio:latest
|
||||
build:
|
||||
context: ..
|
||||
dockerfile: deploy/Dockerfile
|
||||
container_name: omnivoice-worker-gpu
|
||||
profiles: ["worker-gpu"]
|
||||
entrypoint: ["python3", "-m", "uvicorn"]
|
||||
command: ["backend.main:app", "--host", "127.0.0.1", "--port", "3900"]
|
||||
volumes:
|
||||
- omnivoice-worker-data:/app/omnivoice_data
|
||||
environment:
|
||||
- HF_HOME=/app/omnivoice_data/huggingface
|
||||
- HF_TOKEN=${HF_TOKEN:-}
|
||||
- OMNIVOICE_DATA_DIR=/app/omnivoice_data
|
||||
- PYTHONPATH=/app/backend
|
||||
- PYTHONUNBUFFERED=1
|
||||
- OMNIVOICE_SERVER_MODE=0
|
||||
- OMNIVOICE_WORKER_MODE=1
|
||||
- OMNIVOICE_WORKER_TOKEN=${OMNIVOICE_WORKER_TOKEN:-}
|
||||
- OMNIVOICE_WORKER_ENDPOINT=${OMNIVOICE_WORKER_ENDPOINT:-}
|
||||
deploy:
|
||||
resources:
|
||||
reservations:
|
||||
devices:
|
||||
- driver: nvidia
|
||||
count: 1
|
||||
capabilities: [gpu]
|
||||
healthcheck:
|
||||
test: ["CMD", "curl", "-fsS", "http://127.0.0.1:3900/workers/agent/readiness"]
|
||||
interval: 10s
|
||||
timeout: 5s
|
||||
retries: 6
|
||||
start_period: 180s
|
||||
restart: unless-stopped
|
||||
|
||||
# ── Worker-only AMD GPU (ROCm) mode
|
||||
omnivoice-worker-rocm:
|
||||
image: ghcr.io/debpalash/omnivoice-studio:rocm
|
||||
container_name: omnivoice-worker-rocm
|
||||
profiles: ["worker-rocm"]
|
||||
entrypoint: ["python3", "-m", "uvicorn"]
|
||||
command: ["backend.main:app", "--host", "127.0.0.1", "--port", "3900"]
|
||||
devices:
|
||||
- /dev/kfd
|
||||
- /dev/dri
|
||||
volumes:
|
||||
- omnivoice-worker-rocm-data:/app/omnivoice_data
|
||||
environment:
|
||||
- HF_HOME=/app/omnivoice_data/huggingface
|
||||
- HF_TOKEN=${HF_TOKEN:-}
|
||||
- OMNIVOICE_DATA_DIR=/app/omnivoice_data
|
||||
- PYTHONPATH=/app/backend
|
||||
- PYTHONUNBUFFERED=1
|
||||
- OMNIVOICE_SERVER_MODE=0
|
||||
- OMNIVOICE_WORKER_MODE=1
|
||||
- OMNIVOICE_WORKER_TOKEN=${OMNIVOICE_WORKER_TOKEN:-}
|
||||
- OMNIVOICE_WORKER_ENDPOINT=${OMNIVOICE_WORKER_ENDPOINT:-}
|
||||
healthcheck:
|
||||
test: ["CMD", "curl", "-fsS", "http://127.0.0.1:3900/workers/agent/readiness"]
|
||||
interval: 10s
|
||||
timeout: 5s
|
||||
retries: 6
|
||||
start_period: 180s
|
||||
restart: unless-stopped
|
||||
|
||||
volumes:
|
||||
omnivoice-data:
|
||||
omnivoice-worker-data:
|
||||
omnivoice-worker-rocm-data:
|
||||
|
||||
@@ -43,21 +43,30 @@ the entire pipeline runs on CPU, just slower. Pull size: ~5 GB compressed
|
||||
## Quick start (CPU)
|
||||
|
||||
```bash
|
||||
export OMNIVOICE_API_KEY="$(python3 -c 'import secrets; print(secrets.token_urlsafe(32))')"
|
||||
|
||||
docker run -d --name omnivoice \
|
||||
-p 127.0.0.1:3900:3900 \
|
||||
-e OMNIVOICE_API_KEY="$OMNIVOICE_API_KEY" \
|
||||
-v omnivoice-data:/app/omnivoice_data \
|
||||
-v ~/.cache/huggingface:/root/.cache/huggingface \
|
||||
palashdeb/omnivoice-studio:latest
|
||||
```
|
||||
|
||||
Open <http://localhost:3900>. The first run downloads a few GB of model weights —
|
||||
follow `docker logs -f omnivoice` to watch progress.
|
||||
follow `docker logs -f omnivoice` to watch progress. When the UI asks for an
|
||||
API key, paste the generated value; settings and diagnostic actions require
|
||||
this administrator session because Docker NAT hides the browser's true
|
||||
loopback origin.
|
||||
|
||||
## Quick start (NVIDIA GPU)
|
||||
|
||||
```bash
|
||||
export OMNIVOICE_API_KEY="$(python3 -c 'import secrets; print(secrets.token_urlsafe(32))')"
|
||||
|
||||
docker run -d --name omnivoice --gpus all \
|
||||
-p 127.0.0.1:3900:3900 \
|
||||
-e OMNIVOICE_API_KEY="$OMNIVOICE_API_KEY" \
|
||||
-v omnivoice-data:/app/omnivoice_data \
|
||||
-v ~/.cache/huggingface:/root/.cache/huggingface \
|
||||
palashdeb/omnivoice-studio:latest
|
||||
@@ -74,9 +83,12 @@ CUDA-only and runs on CPU on AMD hardware). No toolkit needed — pass the GPU
|
||||
through as device nodes; the host only needs the `amdgpu` kernel driver:
|
||||
|
||||
```bash
|
||||
export OMNIVOICE_API_KEY="$(python3 -c 'import secrets; print(secrets.token_urlsafe(32))')"
|
||||
|
||||
docker run -d --name omnivoice \
|
||||
--device /dev/kfd --device /dev/dri \
|
||||
-p 127.0.0.1:3900:3900 \
|
||||
-e OMNIVOICE_API_KEY="$OMNIVOICE_API_KEY" \
|
||||
-v omnivoice-data:/app/omnivoice_data \
|
||||
-v ~/.cache/huggingface:/root/.cache/huggingface \
|
||||
palashdeb/omnivoice-studio:rocm
|
||||
@@ -87,8 +99,9 @@ Podman users: same two `--device` flags (Quadlet: `AddDevice=/dev/kfd` +
|
||||
`-e HSA_OVERRIDE_GFX_VERSION=11.0.0` if the GPU isn't detected — details in
|
||||
the [Docker install guide](https://github.com/debpalash/VoiceStudio/blob/main/docs/install/docker.md).
|
||||
|
||||
There's also a Compose file in the repo with `cpu` / `gpu` / `rocm` profiles
|
||||
— see the [Docker install guide](https://github.com/debpalash/VoiceStudio/blob/main/docs/install/docker.md).
|
||||
There's also a Compose file in the repo with `cpu` / `gpu` / `rocm` profiles,
|
||||
plus `worker-gpu` / `worker-rocm` profiles that lend a headless GPU without
|
||||
publishing the web UI — see the [Docker install guide](https://github.com/debpalash/VoiceStudio/blob/main/docs/install/docker.md).
|
||||
|
||||
---
|
||||
|
||||
@@ -98,12 +111,12 @@ There's also a Compose file in the repo with `cpu` / `gpu` / `rocm` profiles
|
||||
|-----|--------------|
|
||||
| `:latest` | **Rolling preview** — latest commit on `main`, at or ahead of the last release. This is the preview channel; pin `:stable` for production. |
|
||||
| `:stable` | Most recent versioned release (updated on every `v*` git tag) |
|
||||
| `:0.5.0` | Exact release version |
|
||||
| `:0.5.1` | Exact release version |
|
||||
| `:0.5` | Latest patch within the `0.5` minor |
|
||||
| `:main` | Alias of the same rolling `main` build as `:latest` |
|
||||
| `:sha-xxxxxxx` | A specific commit (produced by manual workflow dispatch) |
|
||||
| `:rocm` | **AMD GPU (ROCm) build** of the rolling preview — the ROCm analogue of `:latest` |
|
||||
| `:stable-rocm`, `:0.5.0-rocm`, `:0.5-rocm`, `:sha-xxxxxxx-rocm` | ROCm builds of the corresponding tags above |
|
||||
| `:stable-rocm`, `:0.5.1-rocm`, `:0.5-rocm`, `:sha-xxxxxxx-rocm` | ROCm builds of the corresponding tags above |
|
||||
|
||||
Preview builds always come from `main` and never version-sort below `:stable`,
|
||||
so upgrades flow naturally. The same images and tags
|
||||
|
||||
+8
-8
@@ -17,7 +17,7 @@ Phase 4 · The two bets ▓▓▓▓▓▓▓▓▓▓ 6 / 6 ✅
|
||||
Phase 5 · Productisation ░░░░░░░░░░ 0 / 5 🚫 demand-driven
|
||||
|
||||
Design track ▓▓▓▓▓▓▓▓▓░ ongoing · 14 primitives + ~67 migrated inline styles · DubTab/Header/Sidebar/CloneDesignTab drained
|
||||
Performance track ░░░░░░░░░░ not started
|
||||
Performance track ▓▓▓░░░░░░░ underway · profiling, preload, isolated engines + cache-remix I/O
|
||||
Feature-magic track ░░░░░░░░░░ not started
|
||||
Quality track ▓▓░░░░░░░░ 12 smoke tests, 10 error messages rewritten
|
||||
```
|
||||
@@ -193,16 +193,16 @@ None on the critical path to world-class. All are answers to real demand.
|
||||
| Design-system primitives (14) | ✅ | Full inventory above. |
|
||||
| Migrate remaining inline styles | 🟡 | **Four biggest offenders drained 2026-04-20** — DubTab **93 → 2**, Header **24 → 1**, Sidebar **21 → 8**, CloneDesignTab **33 → 0**. All remaining are genuinely dynamic (per-row `--row-accent` CSS custom props in Sidebar, per-bar `height/animationDelay` in WaveBars, `opacity` computed from index in skeleton rows, `fontSize` by prop). New class systems: `.dub-*` (DubTab), `.hq-col-*/.hq-stats__*/.hq-logo-*` (Header), `.sidebar-tile--*/.sidebar__scroll/.history-*--*` (Sidebar), `.clone-*/.label-row--*` (CloneDesignTab). Drag-hover on `.file-drag` and `.dub-idle-drop` now toggles `.is-dragging` instead of mutating styles via DOM. Remaining 119 across the tail (Launchpad, KeyboardCheatsheet, DubSegmentRow, WaveformTimeline, etc.) — less concentrated, lower-leverage. |
|
||||
|
||||
### ⚡ Performance track _(⏳ not started)_
|
||||
### ⚡ Performance track _(🟡 underway)_
|
||||
|
||||
| Item | Status | Current measurement |
|
||||
|------|:---:|------|
|
||||
| Batched TTS (8–16 segments per forward pass) | ⏳ | 1 segment per call today. |
|
||||
| Kill per-segment disk round-trip | ⏳ | `dub_generate.py:132-133` saves + re-reads per segment. |
|
||||
| Cold start ≤1.5 s to first audible sample | ⏳ | Currently 4+ s on Apple Silicon. |
|
||||
| Batched TTS (host-derived width per forward pass) | 🟡 | The batch queue feeds OmniVoice's native variable-length forward pass, with the width derived from device headroom (1 on CPU/low-VRAM hosts, up to 8) and overridable via `OMNIVOICE_DUB_BATCH_WIDTH`; adapters without native batching retain the single-segment fallback. |
|
||||
| Kill per-segment disk round-trip | 🟡 | Long-video assembly stays disk-backed to bound RAM. Unchanged same-rate natural segments now skip the redundant decode → scratch encode → decode cycle; fresh segments still persist once and reload for assembly. |
|
||||
| Cold start ≤1.5 s to first audible sample | 🟡 | Installed models preload in the background and `scripts/bench_pipeline.py` measures cold/warm synthesis; target is not yet verified. |
|
||||
| Speculative regeneration on hover | ⏳ | — |
|
||||
| Crash-sandbox engines (subprocess isolation) | ⏳ | Single CUDA OOM still kills server. |
|
||||
| Interaction budgets (<50 ms UI, <200 ms preview, <4 s first seg) | ⏳ | Not measured. |
|
||||
| Crash-sandbox engines (subprocess isolation) | 🟡 | Killable sidecar engines and opt-in `omnivoice-subprocess` are live; the default OmniVoice engine now routes through the crash sandbox on MPS, while CUDA/ROCm/CPU retain the lower-overhead in-process path. |
|
||||
| Interaction budgets (<50 ms UI, <200 ms preview, <4 s first seg) | 🟡 | `/ws/tts` reports real TTFA, total generation time and RTF; frontend responsiveness instrumentation exists, but no cross-surface budget gate yet. |
|
||||
| Dedicated dev-week per quarter | ⏳ | Cadence not yet booked. |
|
||||
|
||||
### ✨ Feature-magic track _(⏳ not started)_
|
||||
@@ -220,7 +220,7 @@ None on the critical path to world-class. All are answers to real demand.
|
||||
| Item | Status | Notes |
|
||||
|------|:---:|------|
|
||||
| Every bug ships a regression test | ⏳ | Rule written, not yet enforced in CI. |
|
||||
| Perf regression budget (≤5 % on fixture clip) | ⏳ | No fixture clip yet. |
|
||||
| Perf regression budget (≤5 % on fixture clip) | ✅ | Shipped 2026-08-20 as hardware-independent **operation-count budgets** (`tests/test_perf_operation_budgets.py`) — stricter than 5 %, and CI-stable where wall-clock on varying runners is not: one generate per sentence on `/ws/tts`, zero TTS calls on cached dub re-mixes; zero decode/rewrite and ⌈N/W⌉ `generate_batch` guards activate with their respective fast paths. See docs/performance.md §Performance budgets. |
|
||||
| Accessibility (keyboard-first, WCAG AA, ARIA live regions) | 🟡 | Focus rings token defined; full audit pending. |
|
||||
| Privacy (zero telemetry by default, per-feature opt-in) | ✅ | Enforced in Settings → Privacy tab. |
|
||||
| Docs updated per phase | 🟡 | STRUCTURE.md, ROADMAP.md, ui/README.md current (research/ + design/ retired 2026-07-12). |
|
||||
|
||||
+5
-1
@@ -60,10 +60,14 @@ VoiceStudio/
|
||||
│ └── frontend/ Node-based frontend tests
|
||||
│
|
||||
├── scripts/ ⟵ dev / build / release shell + python scripts
|
||||
│ ├── install.sh universal installer
|
||||
│ ├── install.sh universal installer (macOS/Linux/WSL)
|
||||
│ ├── install.ps1 universal installer (Windows)
|
||||
│ ├── run.sh universal launcher
|
||||
│ ├── smoke-test.sh end-to-end validation
|
||||
│ └── desktop-prod.sh production desktop build
|
||||
|
||||
├── infra/ ⟵ edge/deploy workers (not the Docker deploy path)
|
||||
│ └── install-redirect/ voicestudio.sh/install — UA-sniffing installer worker
|
||||
│
|
||||
├── deploy/ ⟵ Docker deployment configs
|
||||
│ ├── Dockerfile single-stage CUDA image
|
||||
|
||||
@@ -77,7 +77,7 @@ Both spike URLs are **real, live, and the intended artifacts**. The "VoiceStudio
|
||||
|
||||
# GGUF variant — adds a bundled binary, not a Python dep
|
||||
# Built once per platform in CI and stored alongside the Tauri installer:
|
||||
# bin/omnivoice-tts-{darwin-arm64,darwin-x86_64,windows-x86_64,linux-x86_64}
|
||||
# bin/omnivoice-tts-{darwin-arm64,darwin-x86_64,windows-x86_64,linux-x86_64,linux-aarch64}
|
||||
# Quants pulled at first use via huggingface_hub
|
||||
```
|
||||
|
||||
@@ -149,7 +149,8 @@ User opens app / first run │ Settings → Engines → Default
|
||||
┌─────────────────────────────────────────────────────────────────────────┐
|
||||
│ SubprocessBackend (Phase 2 primitive) │
|
||||
│ │
|
||||
│ spawn: bin/omnivoice-tts-{darwin-arm64|x86_64|win|linux} │
|
||||
│ spawn: bin/omnivoice-tts-{darwin-arm64|x86_64|win|linux| │
|
||||
│ linux-aarch64} │
|
||||
│ --model $HF_HUB_CACHE/.../omnivoice-base-{Q}.gguf │
|
||||
│ --codec $HF_HUB_CACHE/.../omnivoice-tokenizer-{Q}.gguf │
|
||||
│ --lang <user lang> │
|
||||
@@ -229,7 +230,8 @@ bin/ # bundled in installer (per platform)
|
||||
├── omnivoice-tts-darwin-arm64
|
||||
├── omnivoice-tts-darwin-x86_64
|
||||
├── omnivoice-tts-windows-x86_64.exe
|
||||
└── omnivoice-tts-linux-x86_64
|
||||
├── omnivoice-tts-linux-x86_64
|
||||
└── omnivoice-tts-linux-aarch64
|
||||
|
||||
.planning/decisions/
|
||||
├── SPIKE-01-gguf.md # ADR (this research → human review → ratified)
|
||||
|
||||
+16
-7
@@ -1,10 +1,14 @@
|
||||
# Agentic voice: VoiceStudio as a TTS/STT provider
|
||||
|
||||
VoiceStudio exposes an **OpenAI-compatible API**, so any agent framework that
|
||||
VoiceStudio exposes a **local speech platform**—OpenAI-compatible batch audio,
|
||||
a versioned transcription WebSocket, native dictation control, and MCP—so any agent framework that
|
||||
speaks to OpenAI's audio endpoints can use your local VoiceStudio for speech —
|
||||
in your own cloned voice, with nothing leaving your machine. You bring the
|
||||
agent runtime; VoiceStudio is the voice.
|
||||
|
||||
For dictating directly into Claude Code, Codex, Pi, Antigravity CLI, Herdr, or
|
||||
another focused prompt, use the [Rust control sidecar](speech-platform.md).
|
||||
|
||||
This is "agentic v1": VoiceStudio is a provider, not the orchestrator. You wire
|
||||
your own agent (a support line, a desk assistant, a Discord persona) and point
|
||||
its TTS/STT at VoiceStudio.
|
||||
@@ -16,13 +20,17 @@ its TTS/STT at VoiceStudio.
|
||||
|
||||
## The endpoints
|
||||
|
||||
VoiceStudio serves these on `http://localhost:3900/v1` (or your
|
||||
[remote backend URL](remote-gpu.md)):
|
||||
VoiceStudio's service root is `http://localhost:3900` (or your
|
||||
[remote backend URL](remote-gpu.md)). OpenAI-compatible clients use
|
||||
`http://localhost:3900/v1` as their base URL, while discovery stays at the
|
||||
service root: `http://localhost:3900/.well-known/voicestudio-speech`.
|
||||
|
||||
| OpenAI route | VoiceStudio support |
|
||||
|---|---|
|
||||
| `POST /v1/audio/speech` | TTS. `model` = engine id, `voice` = a voice-profile id (your clone) or preset, `response_format` incl. `pcm` and `wav`, `speed`. Default output is 24 kHz. |
|
||||
| `POST /v1/audio/transcriptions` | STT (Whisper-family). |
|
||||
| `WS /v1/audio/transcriptions/stream` | Live partial/final STT from PCM or WebM. |
|
||||
| `GET /.well-known/voicestudio-speech` | Machine-readable transport discovery. |
|
||||
| `GET /v1/audio/voices` | list available voices (VoiceStudio extension). |
|
||||
|
||||
A contract test (`tests/test_agentic_provider_contract.py`) pins this request
|
||||
@@ -74,10 +82,11 @@ single local agent, pipecat is lighter.
|
||||
|
||||
## Remote backend
|
||||
|
||||
Running VoiceStudio on a [remote GPU box](remote-gpu.md)? Use that backend's URL
|
||||
as `base_url` and pass its `OMNIVOICE_API_KEY` as the `api_key` — the same
|
||||
bearer the rest of the app uses. Keep it on your tailnet, not the open
|
||||
internet.
|
||||
Running VoiceStudio on a [remote GPU box](remote-gpu.md)? Append `/v1` to that
|
||||
backend's service-root URL for the OpenAI client's `base_url`, and pass its
|
||||
`OMNIVOICE_API_KEY` as the `api_key` — the same bearer the rest of the app uses.
|
||||
Keep the unmodified service root for `/.well-known/voicestudio-speech`
|
||||
discovery, and keep the backend on your tailnet, not the open internet.
|
||||
|
||||
## Use your own voice responsibly
|
||||
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
# Exporting dubbed video
|
||||
|
||||
Video exports can contain both the source audio and one or more dubbed tracks.
|
||||
VoiceStudio marks the selected dubbed language as the default so ordinary video
|
||||
players and messaging apps play the dub immediately. Choose **Original** in the
|
||||
Default Track control when the source audio should play first instead.
|
||||
@@ -215,6 +215,14 @@ launching the backend (or in **Settings → Credentials**):
|
||||
pro endpoint).
|
||||
- **Microsoft Translator:** `MICROSOFT_API_KEY` (optionally `MICROSOFT_BASE_URL`).
|
||||
|
||||
## Editing workspace
|
||||
|
||||
After transcription, drag the divider between the video/timeline and transcript
|
||||
columns to give either side more room. The divider also works from the keyboard:
|
||||
focus it, use Left/Right Arrow in 5% steps, or Home/End for the minimum/maximum.
|
||||
VoiceStudio remembers the split on this device. Narrow workspaces stack the two
|
||||
panels instead so neither editor becomes unusably small.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
- **"The 'google' translation engine needs the optional deep_translator Python
|
||||
|
||||
@@ -77,6 +77,14 @@ missing, preventing slow disks or antivirus scans from hiding a valid venv.
|
||||
Set `OMNIVOICE_INDEXTTS_IMPORT_PROBE_TIMEOUT_S` to raise the default 60-second
|
||||
probe limit.
|
||||
|
||||
### Long-text generation
|
||||
|
||||
A long passage can keep `infer()` busy for several minutes. The sidecar emits a
|
||||
keep-alive frame every 5 seconds while it works, so the parent can tell a slow
|
||||
synthesis from a wedged one, and waits up to 900 seconds for a sidecar that has
|
||||
gone genuinely silent. Set `OMNIVOICE_INDEXTTS_RECV_TIMEOUT_S` (minimum 30) to
|
||||
tune that ceiling.
|
||||
|
||||
IndexTTS 2.5 requires a language token. VoiceStudio maps locale codes and
|
||||
language names to the five supported languages and detects Chinese, Japanese,
|
||||
or Arabic script for Auto requests. Ambiguous Latin text defaults to English.
|
||||
@@ -95,9 +103,14 @@ confirm that the configured directory contains:
|
||||
```text
|
||||
pyproject.toml
|
||||
indextts/infer_v2_5.py
|
||||
checkpoints/config_v2_5.yaml
|
||||
checkpoints/config.yaml
|
||||
```
|
||||
|
||||
`IndexTeam/IndexTTS-2.5` ships the model config as `config.yaml`. Earlier
|
||||
installs only worked after hand-renaming it to `config_v2_5.yaml`; both names
|
||||
are accepted, so a renamed checkout keeps working as-is and needs no
|
||||
reinstall.
|
||||
|
||||
### `uv` not found
|
||||
|
||||
Install `uv` from <https://docs.astral.sh/uv/> or configure the bundled binary
|
||||
|
||||
@@ -20,8 +20,9 @@ instead — same model family, no NeMo dependency:
|
||||
|
||||
- **Apple Silicon:** [parakeet-mlx](parakeet-mlx.md) (installed by default on
|
||||
mac-ARM source installs).
|
||||
- **Any platform, CPU:** [sherpa-onnx-asr](sherpa-onnx-asr.md) — its default
|
||||
dictation model is an int8 ONNX export of Parakeet TDT v3.
|
||||
- **Any platform, CPU:** [sherpa-onnx-asr](sherpa-onnx-asr.md) — selectable
|
||||
int8 ONNX exports of Parakeet TDT v2/v3; Whisper Tiny remains the
|
||||
cross-platform dictation default.
|
||||
|
||||
## Selecting it
|
||||
|
||||
|
||||
@@ -46,6 +46,15 @@ failing at spawn time; build one with
|
||||
`scripts/build-omnivoice-tts.sh --platform <slug>` or use the default
|
||||
in-process engine.
|
||||
|
||||
**Linux ARM64 (Asahi Apple Silicon):** the `linux-aarch64` binary prefers
|
||||
GGML's Vulkan backend when built on a host with `glslc` and the Khronos
|
||||
SPIRV headers installed (Arch: `pacman -S shaderc spirv-headers`; Debian:
|
||||
`apt install glslc libvulkan-dev spirv-headers`), so Apple GPUs accelerate
|
||||
generation through the open-source Honeykrisp driver. Without those deps the
|
||||
build falls back to CPU. Expect roughly 2–4x slower generation than macOS
|
||||
Metal while upstream Mesa and llama.cpp Vulkan optimizations mature; still
|
||||
well ahead of CPU-only.
|
||||
|
||||
## Integrity and self-healing
|
||||
|
||||
Before reporting ready, the engine:
|
||||
|
||||
@@ -6,12 +6,12 @@ wedged generation can be hard-killed and its VRAM/device reclaimed.
|
||||
|
||||
## Why this engine exists
|
||||
|
||||
The default `omnivoice` engine runs in-process on the GPU worker pool. On
|
||||
VRAM-tight machines (Apple Silicon MPS especially) a heavy generation or model
|
||||
load can exceed its execution budget. When that happens the worker is
|
||||
"abandoned" but **cannot be killed** (Python cannot interrupt a native torch /
|
||||
MPS call), so it keeps holding the GPU device until it finishes on its own, and
|
||||
every later synth queues behind it and hangs (#730 / #1190).
|
||||
An in-process `omnivoice` engine runs on the GPU worker pool. On VRAM-tight
|
||||
machines a heavy generation or model load can exceed its execution budget.
|
||||
When that happens the worker is "abandoned" but **cannot be killed** (Python
|
||||
cannot interrupt a native torch call), so it keeps holding the GPU device until
|
||||
it finishes on its own, and every later synth queues behind it and hangs
|
||||
(#730 / #1190).
|
||||
|
||||
`omnivoice-subprocess` runs the model in a child process spawned via the same
|
||||
`SubprocessBackend` primitive used by IndexTTS, Supertonic-3, and dots.tts. A
|
||||
@@ -26,21 +26,23 @@ structurally cannot do.
|
||||
must recover on its own instead of hanging until a manual restart.
|
||||
- **VRAM-starved MPS hosts** that hit the abandoned-worker cascade.
|
||||
|
||||
For interactive single-shot use on a machine with comfortable VRAM, the default
|
||||
in-process `omnivoice` engine is faster (no stdio round-trip) and remains the
|
||||
default.
|
||||
On Apple Silicon, the default `omnivoice` id automatically uses this isolated
|
||||
implementation. CUDA, ROCm, and CPU keep the in-process implementation and its
|
||||
lower call overhead.
|
||||
|
||||
## Selecting it
|
||||
|
||||
- **Settings -> Engines**, or
|
||||
- **Model Catalogue → Engines**, or
|
||||
- `OMNIVOICE_TTS_BACKEND=omnivoice-subprocess`
|
||||
|
||||
It is **opt-in**; the in-process engine stays the default, so existing setups
|
||||
see no change.
|
||||
The explicit engine is opt-in on CUDA, ROCm, and CPU. Apple Silicon gets the
|
||||
same isolation automatically while keeping the default `omnivoice` id in APIs,
|
||||
Settings, and saved projects.
|
||||
|
||||
## Platform support
|
||||
|
||||
- **CUDA, MPS, and CPU** (same as the in-process VoiceStudio engine).
|
||||
- **CUDA, AMD ROCm on Linux, MPS, and CPU** (same as the in-process
|
||||
VoiceStudio engine).
|
||||
- **No extra install.** Unlike IndexTTS / dots.tts / Supertonic-3, this sidecar
|
||||
runs under VoiceStudio's own interpreter, because the goal here is crash
|
||||
isolation, not dependency isolation. If the default `omnivoice` engine works
|
||||
@@ -53,11 +55,9 @@ see no change.
|
||||
- A wedged generation is **killed and recovered** at the recv-timeout deadline
|
||||
(`OMNIVOICE_SIDECAR_RECV_TIMEOUT_S`, default 300s, aligned with the generate
|
||||
budget) instead of hanging indefinitely.
|
||||
- It does **not** carry the native advanced-parameter surface
|
||||
(`t_shift` / `layer_penalty_factor` / `position_temperature` /
|
||||
`class_temperature`) or parent-side seed determinism, because the generic
|
||||
engine path does not forward those. For plain voice-clone and design
|
||||
synthesis this is a non-issue.
|
||||
- The default Apple Silicon proxy preserves native advanced parameters,
|
||||
deterministic seeds, and longform quality settings across the process
|
||||
boundary.
|
||||
- The recv-timeout deadline is per call and assumes the route's text chunking:
|
||||
`/generate` and `/v1/audio/speech` split long text into pieces of at most
|
||||
`max_chunk_chars` before calling the engine, so each call stays short. A
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user