Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 23 additions & 1 deletion gtsfm/runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,6 +115,28 @@ def _publish_dask_stats(client: Client, stop: threading.Event, path: Path) -> No
stop.wait(1.0)


def _shutdown_dask_client(client: Client) -> None:
"""Release a Dask cluster without letting teardown mask the pipeline result.

Distributed can time out while joining a worker that it has already
terminated. At that point reconstruction is over and raising the teardown
timeout would incorrectly mark a successful run as failed (or hide the
original pipeline exception). Make one short best-effort close after a
failed graceful shutdown and report cleanup problems as warnings only.
"""

try:
client.shutdown()
return
except Exception as exc: # cleanup must not replace the pipeline outcome
logger.warning("Dask shutdown did not finish cleanly; forcing the client closed: %s", exc)

try:
client.close(timeout=2)
except Exception as exc: # the containing process will release remaining children
logger.warning("Dask client force-close did not finish before process exit: %s", exc)


class GtsfmRunner:
def __init__(self, override_args=None) -> None:
argparser: argparse.ArgumentParser = self.construct_argparser()
Expand Down Expand Up @@ -590,7 +612,7 @@ def run(self) -> None:
if stats_thread is not None:
stats_thread.join(timeout=2)
logger.info("🌟 GTSFM: Shutting down Dask client...")
client.shutdown()
_shutdown_dask_client(client)


if __name__ == "__main__":
Expand Down
48 changes: 48 additions & 0 deletions tests/test_runner_shutdown.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
"""Tests for best-effort Dask runner cleanup."""

from gtsfm.runner import _shutdown_dask_client


class _FakeClient:
def __init__(self, *, shutdown_error: Exception | None = None, close_error: Exception | None = None) -> None:
self.shutdown_error = shutdown_error
self.close_error = close_error
self.shutdown_called = False
self.close_timeout: int | None = None

def shutdown(self) -> None:
self.shutdown_called = True
if self.shutdown_error is not None:
raise self.shutdown_error

def close(self, *, timeout: int) -> None:
self.close_timeout = timeout
if self.close_error is not None:
raise self.close_error


def test_shutdown_dask_client_uses_graceful_shutdown() -> None:
client = _FakeClient()

_shutdown_dask_client(client) # type: ignore[arg-type]

assert client.shutdown_called
assert client.close_timeout is None


def test_shutdown_dask_client_force_closes_after_timeout() -> None:
client = _FakeClient(shutdown_error=TimeoutError())

_shutdown_dask_client(client) # type: ignore[arg-type]

assert client.shutdown_called
assert client.close_timeout == 2


def test_shutdown_dask_client_never_raises_cleanup_error() -> None:
client = _FakeClient(shutdown_error=TimeoutError(), close_error=TimeoutError())

_shutdown_dask_client(client) # type: ignore[arg-type]

assert client.shutdown_called
assert client.close_timeout == 2
23 changes: 23 additions & 0 deletions tests/visualization/test_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -340,6 +340,29 @@ def fake_urlopen(_request: object, **kwargs: object) -> BytesIO:
assert runtime._SSL_CONTEXT.get_ca_certs()


def test_remote_workspace_retries_truncated_status_response(monkeypatch: pytest.MonkeyPatch) -> None:
responses: list[BytesIO] = [BytesIO(), BytesIO(b'{"status": "ready"}')]
attempts = 0

def fake_urlopen(_request: object, **_kwargs: object) -> BytesIO:
nonlocal attempts
response = responses[attempts]
attempts += 1
if attempts == 1:
response.read = lambda: (_ for _ in ()).throw( # type: ignore[method-assign]
runtime.http.client.IncompleteRead(b"", 12)
)
return response

monkeypatch.setattr(runtime.urllib.request, "urlopen", fake_urlopen)
monkeypatch.setattr(runtime.time, "sleep", lambda _seconds: None)

result = runtime.JobManager._remote_json("https://workspace.example/api/jobs/example", "secret")

assert result == {"status": "ready"}
assert attempts == 2


def test_workspace_rejects_unknown_github_sample(tmp_path: Path) -> None:
response = TestClient(create_app(tmp_path)).post("/api/samples/not-a-sample/prepare")

Expand Down
52 changes: 47 additions & 5 deletions visualization/runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@
CONFIG_ROOT = PACKAGE_ROOT / "configs"
_SSL_CONTEXT = ssl.create_default_context(cafile=certifi.where())
_REMOTE_REQUEST_TIMEOUT_SECONDS = 10 * 60
_REMOTE_READ_RETRIES = 3
_RUN_NAME_PATTERN = re.compile(r"[^A-Za-z0-9._-]+")
_OPTIONAL_SUBMODULES = {
"submodule-anysplat": {
Expand Down Expand Up @@ -914,8 +915,21 @@ def _remote_json(url: str, api_key: str, *, payload: dict[str, Any] | None = Non
# Allocating a Modal GPU and loading the CUDA runtime can take more than
# two minutes on the first request. Keep the socket open through that
# cold start instead of launching a second competing verification.
with urllib.request.urlopen(request, timeout=_REMOTE_REQUEST_TIMEOUT_SECONDS, context=_SSL_CONTEXT) as response:
result = json.loads(response.read().decode("utf-8"))
# Status reads are safe to retry when Modal closes a response before
# all Content-Length bytes arrive. Never automatically replay a POST:
# the remote operation may have succeeded even if its response broke.
attempts = _REMOTE_READ_RETRIES if payload is None else 1
for attempt in range(attempts):
try:
with urllib.request.urlopen(
request, timeout=_REMOTE_REQUEST_TIMEOUT_SECONDS, context=_SSL_CONTEXT
) as response:
result = json.loads(response.read().decode("utf-8"))
break
except (http.client.IncompleteRead, http.client.RemoteDisconnected, ConnectionError, TimeoutError):
if attempt + 1 >= attempts:
raise
time.sleep(0.25 * (2**attempt))
if not isinstance(result, dict):
raise ValueError("Remote workspace returned an invalid response")
return result
Expand Down Expand Up @@ -1065,7 +1079,21 @@ def _run_remote(self, job: ManagedJob) -> None:
remote_status = str(created.get("status", "queued"))
while remote_status in {"queued", "running"}:
time.sleep(1)
state = self._remote_json(f"{endpoint}/api/jobs/{remote_id}", job.remote_api_key)
try:
state = self._remote_json(f"{endpoint}/api/jobs/{remote_id}", job.remote_api_key)
except (http.client.IncompleteRead, http.client.RemoteDisconnected, ConnectionError, TimeoutError):
# A Modal proxy/container transition can truncate one
# response. The GPU call is still alive, so retain the
# last known state and reconnect on the next poll.
with self._lock:
if job.status == "cancelled":
break
reconnecting = "Remote status connection was interrupted; reconnecting…"
if not job.log_tail or job.log_tail[-1] != reconnecting:
job.log_tail.append(reconnecting)
del job.log_tail[:-250]
job.updated_at = _utc_now()
continue
remote_status = str(state.get("status", "running"))
try:
live = self._remote_json(f"{endpoint}/api/jobs/{remote_id}/live", job.remote_api_key)
Expand All @@ -1083,7 +1111,14 @@ def _run_remote(self, job: ManagedJob) -> None:
Path(job.live_root) / "live_splats.ply",
)
preview_version = next_version
except (ValueError, urllib.error.URLError, TimeoutError):
except (
ValueError,
urllib.error.URLError,
http.client.IncompleteRead,
http.client.RemoteDisconnected,
ConnectionError,
TimeoutError,
):
pass
with self._lock:
if job.status == "cancelled":
Expand All @@ -1108,7 +1143,14 @@ def _run_remote(self, job: ManagedJob) -> None:
with self._lock:
job.status = remote_status
job.updated_at = _utc_now()
except (KeyError, ValueError, urllib.error.URLError, TimeoutError) as exc:
except (
KeyError,
ValueError,
urllib.error.URLError,
http.client.HTTPException,
ConnectionError,
TimeoutError,
) as exc:
with self._lock:
if job.status != "cancelled":
job.status = "failed"
Expand Down
Loading