From b1757b631c704ffb36931bec789fc3f2775ea996 Mon Sep 17 00:00:00 2001 From: Jonathan Date: Tue, 14 Jul 2026 11:41:51 +0200 Subject: [PATCH] feat(worker): smooth download % from adapter byte-level progress MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The Floor bar used a file-count poller that jumps 0→~99% for small/fast albums (streamrip reports only 0/1). But yt-dlp and slskd already emit real byte-level progress via the on_progress callback the pipeline was discarding (lambda _pct: None). - Thread a throttled on_progress → Job.downloadProgress (write on a >=1% move or every 0.5s; clamped 0..1; write failures swallowed). Called synchronously from adapter.download on the pipeline thread, so it reuses the main conn safely. - The file-count poller stays as the FALLBACK: it skips writing while a `reports` flag is set (an adapter is driving real progress), and the flag is cleared per attempt so a fall-through to streamrip still gets the estimate. - Smooth bars for YouTube + Soulseek now; Qobuz/streamrip keeps the file-count estimate until its per-track callback is hooked later. worker 216 tests / 7-skip (on_progress write/throttle/clamp covered; existing pipeline tests now exercise the fake's on_progress end-to-end). Co-Authored-By: Claude Opus 4.8 (1M context) --- worker/lyra_worker/pipeline.py | 48 +++++++++++++++++++++----- worker/tests/test_download_progress.py | 37 +++++++++++++++++++- 2 files changed, 75 insertions(+), 10 deletions(-) diff --git a/worker/lyra_worker/pipeline.py b/worker/lyra_worker/pipeline.py index a1f4659..8e853fb 100644 --- a/worker/lyra_worker/pipeline.py +++ b/worker/lyra_worker/pipeline.py @@ -1,6 +1,7 @@ import os import shutil import threading +import time from dataclasses import replace from typing import Sequence @@ -26,20 +27,24 @@ def _count_staged_audio(staging: str) -> int: return n -def _poll_download_progress(job_id: str, staging: str, expected: int, stop: "threading.Event") -> None: +def _poll_download_progress(job_id: str, staging: str, expected: int, stop: "threading.Event", + reports: "threading.Event") -> None: """Until `stop`, periodically write Job.downloadProgress = files-in-staging / expected - (capped 0.99). Uses its own short-lived connection (psycopg conns aren't shareable across - threads). A missing DSN or expected<=0 just no-ops.""" + (capped 0.99) — the FALLBACK for adapters (streamrip/Qobuz) that report no byte-level + progress. Skips writing while `reports` is set, i.e. the in-flight adapter is driving + progress itself via on_progress (yt-dlp/slskd). Uses its own short-lived connection + (psycopg conns aren't shareable across threads). A missing DSN or expected<=0 just no-ops.""" dsn = os.environ.get("DATABASE_URL") if not dsn or expected <= 0: return conn = psycopg.connect(dsn) try: while not stop.is_set(): - frac = min(_count_staged_audio(staging) / expected, 0.99) - with conn.cursor() as cur: - cur.execute('UPDATE "Job" SET "downloadProgress" = %s WHERE id = %s', (frac, job_id)) - conn.commit() + if not reports.is_set(): # the adapter isn't reporting real progress → estimate + frac = min(_count_staged_audio(staging) / expected, 0.99) + with conn.cursor() as cur: + cur.execute('UPDATE "Job" SET "downloadProgress" = %s WHERE id = %s', (frac, job_id)) + conn.commit() stop.wait(1.5) except Exception as e: # a progress poller must never affect the job print(f"pipeline: download progress poller error: {e}", flush=True) @@ -62,6 +67,29 @@ def _set_download_progress(conn: psycopg.Connection, job_id: str, frac: float) - conn.commit() +def _make_on_progress(conn: psycopg.Connection, job_id: str, reports: "threading.Event"): + """Build the on_progress(pct) callback threaded into adapter.download. It marks the + adapter as reporting real byte-level progress (so the file-count poller yields) and + throttle-writes Job.downloadProgress. Called synchronously from adapter.download on the + pipeline's own thread, so it safely reuses `conn`.""" + state = {"frac": -1.0, "t": 0.0} + + def on_progress(pct: float) -> None: + reports.set() + frac = 0.0 if pct < 0 else 1.0 if pct > 1 else float(pct) + now = time.monotonic() + # throttle DB writes: on a >=1% move or every 0.5s, not once per received byte + if frac - state["frac"] >= 0.01 or (now - state["t"]) >= 0.5: + state["frac"] = frac + state["t"] = now + try: + _set_download_progress(conn, job_id, frac) + except Exception as e: # a progress write must never fail the download + print(f"pipeline: on_progress write failed: {e}", flush=True) + + return on_progress + + def _request_id(conn: psycopg.Connection, job_id: str) -> str: with conn.cursor() as cur: cur.execute('SELECT "requestId" FROM "Job" WHERE id = %s', (job_id,)) @@ -245,8 +273,9 @@ def run_pipeline( _set_download_progress(conn, job_id, 0.0) # reset for this run expected = target.track_count or (ranked[0].track_count if ranked else 0) _stop = threading.Event() + _reports = threading.Event() # set while the in-flight adapter reports real byte progress _poller = threading.Thread( - target=_poll_download_progress, args=(job_id, staging, expected, _stop), daemon=True + target=_poll_download_progress, args=(job_id, staging, expected, _stop, _reports), daemon=True ) _poller.start() try: @@ -257,7 +286,8 @@ def run_pipeline( continue _mark_chosen(conn, job_id, candidate.source_ref) # reflect the source now in flight shutil.rmtree(staging, ignore_errors=True) # clean slate per attempt - result = adapter.download(candidate, staging, lambda _pct: None) + _reports.clear() # this attempt hasn't reported yet → poller estimates until it does + result = adapter.download(candidate, staging, _make_on_progress(conn, job_id, _reports)) if result.ok: winner = candidate break diff --git a/worker/tests/test_download_progress.py b/worker/tests/test_download_progress.py index 9be3ad5..c578f4e 100644 --- a/worker/tests/test_download_progress.py +++ b/worker/tests/test_download_progress.py @@ -1,6 +1,41 @@ import os +import threading -from lyra_worker.pipeline import _count_staged_audio +from lyra_worker.pipeline import _count_staged_audio, _make_on_progress +from tests.conftest import insert_request + + +def _dp(conn, job_id): + with conn.cursor() as cur: + cur.execute('SELECT "downloadProgress" FROM "Job" WHERE id = %s', (job_id,)) + return cur.fetchone()[0] + + +def test_on_progress_marks_reporting_and_writes(conn): + job_id = insert_request(conn) + reports = threading.Event() + cb = _make_on_progress(conn, job_id, reports) + cb(0.3) + assert reports.is_set() # signals the poller to yield + assert _dp(conn, job_id) == 0.3 + + +def test_on_progress_throttles_tiny_moves(conn): + job_id = insert_request(conn) + cb = _make_on_progress(conn, job_id, threading.Event()) + cb(0.30) + cb(0.305) # <1% and <0.5s later → throttled, not written + assert _dp(conn, job_id) == 0.3 + cb(0.6) # >=1% move → written + assert _dp(conn, job_id) == 0.6 + + +def test_on_progress_clamps_out_of_range(conn): + job_id = insert_request(conn) + _make_on_progress(conn, job_id, threading.Event())(1.5) + assert _dp(conn, job_id) == 1.0 + _make_on_progress(conn, job_id, threading.Event())(-0.2) + assert _dp(conn, job_id) == 0.0 def test_count_staged_audio_counts_nested_audio_files_only(tmp_path):