Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 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
131 changes: 106 additions & 25 deletions server.py
Original file line number Diff line number Diff line change
Expand Up @@ -1613,10 +1613,23 @@ def _check_cache(job_id, stem_list, model):

stems_found = {}
for stem_name in stem_list:
for ext in (".mp3", ".wav", ".flac"):
p = cache_path / f"{stem_name}{ext}"
if p.exists():
stems_found[stem_name] = f"/download/{job_id}/{stem_name}{ext}"
# Probe the LOWERCASE filename first (what the workers now write), then the caller's
# own casing (what pre-fix cache entries were written with). Probing only the
# caller's spelling would miss a cache entry that exists — a mixed-case `stems=`
# request would silently re-separate, and after a restart the in-memory jobs table
# is empty, so this is the ONLY path that can find it.
lower = stem_name.strip().lower()
for candidate in (lower, stem_name):
hit = None
for ext in (".mp3", ".wav", ".flac"):
p = cache_path / f"{candidate}{ext}"
if p.exists():
# The key echoes what the caller asked for; the URL points at the file
# that actually exists on disk.
hit = f"/download/{job_id}/{candidate}{ext}"
break
if hit:
stems_found[stem_name] = hit
break

if len(stems_found) == len(stem_list):
Expand Down Expand Up @@ -1687,13 +1700,55 @@ def _enqueue_job(job_id, audio_path, stem_list, model):
"""Create a job and start processing in background."""
global active_count

wanted = {s.strip().lower() for s in stem_list}

with jobs_lock:
# If job already exists and is processing/complete, return it
# If a job for this audio+model already exists, reuse it — but ONLY if it actually
# covers what this caller asked for.
#
# It used to return the existing job's stems regardless. Since job_id is
# (audio, model) and NOT the stem set, a 4-stem separation followed by a 6-stem
# request on the same audio silently returned the 4 — in 0 ms, looking like a fast
# success. The caller asked for guitar and piano and simply didn't get them, with no
# error and nothing to indicate the answer was stale rather than authoritative.
existing = jobs.get(job_id)
if existing and existing["status"] in ("processing", "complete"):
if existing["status"] == "complete":
return {"job_id": job_id, "stems": existing["stems"], "cached": True}
return {"job_id": job_id, "status": "processing"}
# Normalize the KEYS before both the coverage check and the lookup.
#
# A job from before this fix stores `stems` keyed by the ORIGINAL caller's
# casing (`{"Vocals": ...}`). Testing coverage case-insensitively while
# looking values up by the lowercased name would pass the check and then
# find nothing — returning `cached: True` with an empty stems dict. A
# confident, instant, empty answer is worse than the bug this PR fixes.
all_urls = {k.strip().lower(): v
for k, v in (existing.get("stems_all")
or existing.get("stems") or {}).items()}
if wanted <= set(all_urls):
return {
"job_id": job_id,
# Keys echo what the CALLER asked for; values come from the
# normalized map.
"stems": {s: all_urls[s.strip().lower()] for s in stem_list},
"cached": True,
}
Comment thread
topkoa marked this conversation as resolved.
Outdated
# Not covered: the previous run produced a smaller set (an older, narrower
# request, or a job from before this fix). Fall through and re-separate, so
# the caller gets what it actually asked for.
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
else:
# In flight. If its stem set covers ours we can ride along; if not, the
# result would be short, so don't attach to it — but we also can't run two
# jobs under one id, so tell the caller plainly instead of handing back a
# job that will complete without their stems.
in_flight = {s.strip().lower() for s in (existing.get("stem_list") or [])}
if not in_flight or wanted <= in_flight:
return {"job_id": job_id, "status": "processing"}
Comment thread
topkoa marked this conversation as resolved.
Outdated
return {
"error": "A separation for this audio is already running with a smaller "
"stem set. Retry once it finishes and the extra stems will be "
"computed.",
"job_id": job_id,
}

with active_lock:
if active_count >= MAX_CONCURRENT:
Expand All @@ -1704,6 +1759,12 @@ def _enqueue_job(job_id, audio_path, stem_list, model):
"status": "processing",
"progress": 0,
"stems": {},
# What THIS run was asked for, and (once complete) everything the model actually
# produced. job_id is (audio, model) and does not include the stem set, so a later
# request for a different set has to be able to tell whether this job covers it.
"stem_list": list(stem_list),
"stems_all": {},
"missing": [],
"error": None,
"model": model,
"created_at": time.time(),
Expand Down Expand Up @@ -1785,19 +1846,25 @@ def _run_demucs(job_id, audio_path, stem_list, model):
cache_path = _cache_entry_path(job_id)
cache_path.mkdir(parents=True, exist_ok=True)

stems_result = {}
for stem_name in stem_list:
src = out_track_dir / f"{stem_name}.wav"
if not src.exists():
continue

wav_dest = cache_path / f"{stem_name}.wav"
# Cache EVERY stem the model produced, not just this caller's subset — see the note
# in _run_roformer. demucs writes its full stem set to out_track_dir whether we take
# them or not; dropping the extras means the next request for a superset re-runs a
# separation we already did.
stems_all = {}
for src in sorted(out_track_dir.glob("*.wav")):
label = src.stem.strip().lower()
wav_dest = cache_path / f"{label}.wav"
shutil.copy2(src, wav_dest)
stems_result[stem_name] = f"/download/{job_id}/{stem_name}.wav"
stems_all[label] = f"/download/{job_id}/{label}.wav"

stems_result = {s: stems_all[s.strip().lower()]
for s in stem_list if s.strip().lower() in stems_all}
Comment thread
topkoa marked this conversation as resolved.

if stems_result:
if stems_all:
_remember_cache_entry(job_id)
_update_job(job_id, status="complete", progress=100, stems=stems_result)
missing = [s for s in stem_list if s.strip().lower() not in stems_all]
_update_job(job_id, status="complete", progress=100, stems=stems_result,
stems_all=stems_all, missing=missing)

except subprocess.TimeoutExpired:
proc.kill()
Expand Down Expand Up @@ -1882,18 +1949,32 @@ def _run_roformer(job_id, audio_path, stem_list, model):
cache_path = _cache_entry_path(job_id)
cache_path.mkdir(parents=True, exist_ok=True)

stems_result = {}
for stem_name in stem_list:
src = produced.get(stem_name.strip().lower())
if not src or not src.exists():
# Cache EVERY stem the model produced, not just the ones this caller asked for.
#
# The separation computes them all regardless — bs_roformer_sw always emits six —
# so discarding the unrequested ones means the next request for a larger stem set
# pays the full inference again (~2 min on a GPU) for stems we had already made and
# deleted. Caching all of them turns that into a cache hit, and it is what lets the
# cache honour a request for a SUPERSET of what was originally asked for.
stems_all = {}
for label, src in produced.items():
if not src.exists():
continue
dest = cache_path / f"{stem_name}.flac"
dest = cache_path / f"{label}.flac"
shutil.copy2(src, dest)
stems_result[stem_name] = f"/download/{job_id}/{stem_name}.flac"
stems_all[label] = f"/download/{job_id}/{label}.flac"

# ...but return only what this caller asked for.
stems_result = {s: stems_all[s.strip().lower()]
for s in stem_list if s.strip().lower() in stems_all}

if stems_result:
if stems_all:
_remember_cache_entry(job_id)
_update_job(job_id, status="complete", progress=100, stems=stems_result)
# `missing` is explicit: a caller that asked for `guitar` from a 4-stem model should
# be told it isn't coming, not silently handed a short dict and left to wonder.
missing = [s for s in stem_list if s.strip().lower() not in stems_all]
_update_job(job_id, status="complete", progress=100, stems=stems_result,
stems_all=stems_all, missing=missing)

except subprocess.TimeoutExpired:
if proc:
Expand Down
19 changes: 13 additions & 6 deletions tests/test_cache_cleanup.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@

def _extract_function(func_name: str):
"""Extract a function from server.py source code using AST."""
source = SERVER_PY.read_text()
source = SERVER_PY.read_text(encoding="utf-8")
tree = ast.parse(source)
for node in ast.iter_child_nodes(tree):
if isinstance(node, ast.FunctionDef) and node.name == func_name:
Expand Down Expand Up @@ -100,10 +100,17 @@ def _parse_requirement_pin(line: str):


def _read_requirements():
"""Read requirements.txt and return dict of package -> min_version."""
"""Read requirements.txt and return dict of package -> min_version.

encoding="utf-8" explicitly. Without it, open() uses the platform default — UTF-8 on
Linux (so CI is green) and **cp1252 on Windows**, where any non-ASCII character in the
file (an em-dash in a comment is enough) raises UnicodeDecodeError. These tests were
passing in CI while being broken for every Windows contributor, which is the worst way
for a test to fail: invisibly, and only for other people.
"""
req_path = Path(__file__).parent.parent / "requirements.txt"
pins = {}
with open(req_path) as f:
with open(req_path, encoding="utf-8") as f:
for line in f:
parsed = _parse_requirement_pin(line)
if parsed:
Expand Down Expand Up @@ -158,7 +165,7 @@ def test_after_includes_network_online(self):
"""After= must reference both network.target AND network-online.target
to prevent the service from starting before the network stack is
fully ready (Wants= alone doesn't enforce ordering)."""
content = self.SERVICE_PATH.read_text()
content = self.SERVICE_PATH.read_text(encoding="utf-8")
assert "After=network.target network-online.target" in content, (
"After= should be 'network.target network-online.target', "
"not just 'network.target'"
Expand All @@ -176,7 +183,7 @@ def test_cleanup_removes_stale_jobs(self):
"""Extract _cache_cleanup_loop from server.py and verify that
shutil.rmtree is immediately followed by jobs.pop with the
same entry name, wrapped in the jobs_lock."""
source = SERVER_PY.read_text()
source = SERVER_PY.read_text(encoding="utf-8")
tree = ast.parse(source)

# Find _cache_cleanup_loop function
Expand Down Expand Up @@ -220,7 +227,7 @@ def test_sleep_at_end_of_loop_not_beginning(self):
"""Extract _cache_cleanup_loop and verify time.sleep is at the
END of the while body, not the beginning. This ensures the
first sweep runs immediately on startup."""
source = SERVER_PY.read_text()
source = SERVER_PY.read_text(encoding="utf-8")
tree = ast.parse(source)

func_node = None
Expand Down
Loading
Loading