Fix CUDA downloads on unsupported platforms (#770)

* Fix CUDA downloads on unsupported platforms

* fix: align CUDA status nullability

* fix: require CUDA download support flag
This commit is contained in:
Andrew Barnes
2026-07-20 19:22:57 -07:00
committed by Jamie Pine
parent 3f4631c865
commit 12ce7a4c35
7 changed files with 78 additions and 6 deletions
+4
View File
@@ -26,6 +26,10 @@ async def download_cuda_backend():
"""Download the CUDA backend binary."""
from ..services import cuda
unsupported_reason = cuda.get_cuda_download_unsupported_reason()
if unsupported_reason:
raise HTTPException(status_code=409, detail=unsupported_reason)
if cuda.get_cuda_binary_path() is not None:
raise HTTPException(status_code=409, detail="CUDA backend already downloaded")
+32 -1
View File
@@ -21,9 +21,9 @@ import tarfile
from pathlib import Path
from typing import Optional
from .. import __version__
from ..config import get_data_dir
from ..utils.progress import get_progress_manager
from .. import __version__
logger = logging.getLogger(__name__)
@@ -31,6 +31,8 @@ GITHUB_RELEASES_URL = "https://github.com/jamiepine/voicebox/releases/download"
PROGRESS_KEY = "cuda-backend"
CUDA_DOWNLOAD_UNSUPPORTED_REASON = "Downloadable CUDA backend releases are currently only published for Windows."
# The current expected CUDA libs version. Bump this when we change the
# CUDA toolkit version or torch's CUDA dependency changes (e.g. cu126 -> cu128).
CUDA_LIBS_VERSION = "cu128-v1"
@@ -63,6 +65,25 @@ def get_cuda_exe_name() -> str:
return "voicebox-server-cuda"
def is_cuda_download_supported() -> bool:
"""Return whether this platform has a matching CUDA release asset."""
return sys.platform == "win32"
def get_cuda_download_unsupported_reason() -> str | None:
"""Explain why this platform cannot use the release-download flow."""
if is_cuda_download_supported():
return None
return CUDA_DOWNLOAD_UNSUPPORTED_REASON
def ensure_cuda_download_supported() -> None:
"""Raise if downloading would fetch an asset built for another platform."""
reason = get_cuda_download_unsupported_reason()
if reason:
raise RuntimeError(reason)
def get_cuda_binary_path() -> Optional[Path]:
"""Return path to the CUDA executable if it exists inside the onedir."""
p = get_cuda_dir() / get_cuda_exe_name()
@@ -103,12 +124,15 @@ def get_cuda_status() -> dict:
cuda_path = get_cuda_binary_path()
progress = progress_manager.get_progress(PROGRESS_KEY)
cuda_libs_version = get_installed_cuda_libs_version()
unsupported_reason = get_cuda_download_unsupported_reason()
return {
"available": cuda_path is not None,
"active": is_cuda_active(),
"binary_path": str(cuda_path) if cuda_path else None,
"cuda_libs_version": cuda_libs_version,
"download_supported": unsupported_reason is None,
"unsupported_reason": unsupported_reason,
"downloading": progress is not None and progress.get("status") == "downloading",
"download_progress": progress,
}
@@ -257,6 +281,8 @@ async def download_cuda_binary(version: Optional[str] = None):
async def _download_cuda_binary_locked(version: Optional[str] = None):
"""Inner implementation of download_cuda_binary, called under _download_lock."""
ensure_cuda_download_supported()
import httpx
if version is None:
@@ -387,6 +413,11 @@ async def check_and_update_cuda_binary():
if not cuda_path:
return # No CUDA binary installed, nothing to update
unsupported_reason = get_cuda_download_unsupported_reason()
if unsupported_reason:
logger.info("Skipping CUDA backend auto-update: %s", unsupported_reason)
return
need_server = _needs_server_download()
need_libs = _needs_cuda_libs_download()
+32
View File
@@ -0,0 +1,32 @@
import sys as py_sys
import types
import pytest
from backend.services import cuda
def test_cuda_status_reports_unsupported_linux_download(monkeypatch, tmp_path):
monkeypatch.setattr(cuda.sys, "platform", "linux")
monkeypatch.setattr(cuda, "get_data_dir", lambda: tmp_path)
status = cuda.get_cuda_status()
assert status["available"] is False
assert status["download_supported"] is False
assert status["unsupported_reason"] == cuda.CUDA_DOWNLOAD_UNSUPPORTED_REASON
@pytest.mark.asyncio
async def test_cuda_download_rejects_linux_before_network(monkeypatch, tmp_path):
monkeypatch.setattr(cuda.sys, "platform", "linux")
monkeypatch.setattr(cuda, "get_data_dir", lambda: tmp_path)
class UnexpectedClient:
def __init__(self, *args, **kwargs):
raise AssertionError("unsupported platforms should not start a release download")
monkeypatch.setitem(py_sys.modules, "httpx", types.SimpleNamespace(AsyncClient=UnexpectedClient))
with pytest.raises(RuntimeError, match="currently only published for Windows"):
await cuda._download_cuda_binary_locked("v0.5.0")