Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 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
102 changes: 81 additions & 21 deletions server.py
Original file line number Diff line number Diff line change
Expand Up @@ -1687,13 +1687,47 @@ 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"}
have = {k.strip().lower()
for k in (existing.get("stems_all") or existing.get("stems") or {})}
if wanted <= have:
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated
all_urls = existing.get("stems_all") or existing.get("stems") or {}
return {
"job_id": job_id,
"stems": {s: all_urls[s.strip().lower()]
for s in stem_list if s.strip().lower() in all_urls},
"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 +1738,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 +1825,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 +1928,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
158 changes: 158 additions & 0 deletions tests/test_stem_set_cache.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,158 @@
"""The cache must honour the requested STEM SET, not just the audio and the model. (#10)

`job_id` is `(audio_hash, model)` and deliberately does not include the stem set. That's
fine for the on-disk cache — `_check_cache` requires every requested stem to be present —
but `_enqueue_job` short-circuited on the in-memory jobs table and returned a completed
job's stems *regardless of what had been asked for*:

POST /separate?model=bs_roformer_sw -> drums, bass, vocals, other
POST /separate?model=bs_roformer_sw&stems=...,guitar,piano -> drums, bass, vocals, other

...in 0 ms, with no error. The caller asked for guitar and piano, got neither, and the
response *looked* like a fast success. Silent and fast is the worst combination: nothing
tells you the answer is stale rather than authoritative.

Extracted via AST, like test_cache_cleanup, to avoid server.py's torch/whisperx import
chain — which is exactly why this can run in CI.
"""
import ast
import collections
import threading
import time
from pathlib import Path

import pytest

SERVER_PY = Path(__file__).parent.parent / "server.py"


def _load_enqueue_job(jobs, max_concurrent=2):
"""Extract _enqueue_job with a namespace standing in for server.py's module globals."""
tree = ast.parse(SERVER_PY.read_text(encoding="utf-8"))
node = next(n for n in ast.iter_child_nodes(tree)
if isinstance(n, ast.FunctionDef) and n.name == "_enqueue_job")
mod = ast.Module(body=[node], type_ignores=[])
ast.copy_location(mod, node)

started = []

class _FakeThread:
def __init__(self, target=None, args=(), daemon=None):
self._args = args

def start(self):
started.append(self._args) # record; never actually separate anything

ns = {
"jobs": jobs,
"jobs_lock": threading.Lock(),
"active_lock": threading.Lock(),
"active_count": 0,
"MAX_CONCURRENT": max_concurrent,
"threading": type("T", (), {"Thread": _FakeThread}),
"time": time,
"_run_roformer": lambda *a: None,
"_run_demucs": lambda *a: None,
"_is_roformer_model": lambda m: "roformer" in m,
}
exec(compile(ast.unparse(mod), "<test>", "exec"), ns)
return ns["_enqueue_job"], started


JOB_ID = "deadbeef-bs-roformer-sw"
FOUR = {"drums": "/download/x/drums.flac", "bass": "/download/x/bass.flac",
"vocals": "/download/x/vocals.flac", "other": "/download/x/other.flac"}
SIX = dict(FOUR, guitar="/download/x/guitar.flac", piano="/download/x/piano.flac")
SIX_NAMES = ["drums", "bass", "vocals", "other", "guitar", "piano"]


def _completed(stems_all, stem_list=None):
return collections.OrderedDict({JOB_ID: {
"job_id": JOB_ID, "status": "complete", "progress": 100,
"stems": dict(stems_all), "stems_all": dict(stems_all),
"stem_list": list(stem_list or stems_all), "missing": [],
"error": None, "model": "bs_roformer_sw", "created_at": time.time(),
}})


def _in_flight(stem_list):
return collections.OrderedDict({JOB_ID: {
"job_id": JOB_ID, "status": "processing", "progress": 40,
"stems": {}, "stems_all": {}, "stem_list": list(stem_list), "missing": [],
"error": None, "model": "bs_roformer_sw", "created_at": time.time(),
}})


def test_superset_request_is_not_served_from_a_smaller_completed_job():
"""THE bug: a 6-stem request answered with the cached 4, instantly, with no error."""
jobs = _completed(FOUR)
enqueue, started = _load_enqueue_job(jobs)
result = enqueue(JOB_ID, "/tmp/a.ogg", SIX_NAMES, "bs_roformer_sw")

assert result.get("cached") is not True, (
"serving the cached 4-stem result for a 6-stem request is silent data loss — the "
"caller asked for guitar and piano and got neither, with no error"
)
assert started, "it must actually re-separate rather than return a short result"


def test_exact_match_is_served_from_cache():
jobs = _completed(SIX)
enqueue, started = _load_enqueue_job(jobs)
result = enqueue(JOB_ID, "/tmp/a.ogg", SIX_NAMES, "bs_roformer_sw")
assert result["cached"] is True
assert set(result["stems"]) == set(SIX)
assert not started


def test_subset_request_is_served_from_a_larger_completed_job():
"""Fewer stems than were computed is a legitimate hit — and returns only what was asked
for, not everything we happen to have lying around."""
jobs = _completed(SIX)
enqueue, started = _load_enqueue_job(jobs)
result = enqueue(JOB_ID, "/tmp/a.ogg", ["vocals", "drums"], "bs_roformer_sw")
assert result["cached"] is True
assert set(result["stems"]) == {"vocals", "drums"}
assert not started


def test_case_and_whitespace_do_not_defeat_the_coverage_check():
jobs = _completed(SIX)
enqueue, _ = _load_enqueue_job(jobs)
assert enqueue(JOB_ID, "/tmp/a.ogg", [" Vocals ", "DRUMS"], "bs_roformer_sw")["cached"]


def test_a_job_from_before_this_fix_still_serves_its_stems():
# Old jobs carry only `stems` (no `stems_all`). Coverage must fall back to it, or every
# pre-existing entry would be needlessly re-separated.
jobs = collections.OrderedDict({JOB_ID: {
"job_id": JOB_ID, "status": "complete", "progress": 100,
"stems": dict(SIX), "error": None, "model": "bs_roformer_sw",
"created_at": time.time(),
}})
enqueue, started = _load_enqueue_job(jobs)
result = enqueue(JOB_ID, "/tmp/a.ogg", ["vocals", "guitar"], "bs_roformer_sw")
assert result["cached"] is True
assert not started


def test_in_flight_job_with_a_smaller_set_is_not_silently_joined():
"""Riding along on a running 4-stem job completes without guitar/piano — the same silent
loss, merely delayed."""
jobs = _in_flight(["drums", "bass", "vocals", "other"])
enqueue, _ = _load_enqueue_job(jobs)
result = enqueue(JOB_ID, "/tmp/a.ogg", SIX_NAMES, "bs_roformer_sw")
assert "error" in result
assert result.get("status") != "processing"


def test_in_flight_job_that_covers_us_is_joined():
jobs = _in_flight(SIX_NAMES)
enqueue, started = _load_enqueue_job(jobs)
result = enqueue(JOB_ID, "/tmp/a.ogg", ["vocals", "guitar"], "bs_roformer_sw")
assert result["status"] == "processing"
assert not started, "must attach to the running job, not start a second separation"


if __name__ == "__main__":
raise SystemExit(pytest.main([__file__, "-v"]))
Loading