File size: 3,998 Bytes
03d531b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
#!/usr/bin/env python3
"""Smoke-test manifest-backed interactive supervision endpoints."""

from __future__ import annotations

import io
import json
import sys
import time
from pathlib import Path
from urllib.parse import quote

import soundfile as sf
from fastapi.testclient import TestClient

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))

from app import app  # noqa: E402
from synth_generator import generate_test_song  # noqa: E402


def wait_for_job(client: TestClient, job_id: str) -> dict:
    for _ in range(80):
        payload = client.get(f"/api/jobs/{job_id}").json()
        if payload["status"] in {"complete", "error"}:
            return payload
        time.sleep(0.15)
    raise TimeoutError(job_id)


def post_json(client: TestClient, path: str, body: dict | None = None) -> dict:
    response = client.post(path, json=body or {})
    response.raise_for_status()
    return response.json()


def main() -> int:
    song = generate_test_song(pattern_name="funk", bars=1, bpm=124, add_bass=False)
    buf = io.BytesIO()
    sf.write(buf, song.drums_only, song.sr, format="WAV")
    buf.seek(0)

    client = TestClient(app)
    response = client.post(
        "/api/jobs",
        files={"file": ("interactive.wav", buf, "audio/wav")},
        data={"params": json.dumps({"stem": "all", "clustering_mode": "online_preview", "target_min": 3, "target_max": 10})},
    )
    response.raise_for_status()
    job_id = response.json()["id"]
    job = wait_for_job(client, job_id)
    assert job["status"] == "complete", job.get("error")

    state = client.get(f"/api/jobs/{job_id}/state").json()
    assert state["summary"]["hit_count"] > 0
    assert state["summary"]["cluster_count"] > 0
    assert state["review_queue"], "expected uncertainty review queue"

    hit_id = state["hits"][0]["id"]
    cluster_id = state["clusters"][0]["id"]
    q_hit = quote(hit_id, safe="")
    q_cluster = quote(cluster_id, safe="")

    state = post_json(client, f"/api/jobs/{job_id}/clusters/{q_cluster}/lock", {"locked": True})
    assert state["summary"]["locked_cluster_count"] >= 1

    state = post_json(client, f"/api/jobs/{job_id}/hits/{q_hit}/review", {"status": "favorite"})
    assert state["summary"]["constraint_count"] >= 1

    explanation = client.get(f"/api/jobs/{job_id}/explain/cluster/{q_cluster}")
    explanation.raise_for_status()
    assert explanation.json()["cluster_id"] == cluster_id

    state = post_json(client, f"/api/jobs/{job_id}/hits/{q_hit}/pull-out", {})
    assert state["summary"]["cluster_count"] >= 1
    assert state["summary"]["undo_available"] is True
    assert any(c["type"] in {"cannot-link", "force-cluster"} for c in state["constraints"])

    state = post_json(client, f"/api/jobs/{job_id}/undo", {})
    assert state["summary"]["hit_count"] > 0

    if len(state["clusters"]) > 1:
        target = next(c for c in state["clusters"] if c["id"] != state["hits"][0]["cluster_id"])
        state = post_json(
            client,
            f"/api/jobs/{job_id}/hits/{q_hit}/move",
            {"target_cluster_id": target["id"]},
        )
        assert any(c["type"] == "force-cluster" for c in state["constraints"])

    if len(state["hits"]) > 1:
        suppress_hit = quote(state["hits"][1]["id"], safe="")
        state = post_json(client, f"/api/jobs/{job_id}/hits/{suppress_hit}/suppress", {"reason": "bleed"})
        assert state["summary"]["suppressed_hit_count"] >= 1

    suggestions = client.get(f"/api/jobs/{job_id}/suggestions")
    suggestions.raise_for_status()

    print(json.dumps({
        "status": "ok",
        "job_id": job_id,
        "hit_count": state["summary"]["hit_count"],
        "cluster_count": state["summary"]["cluster_count"],
        "constraints": state["summary"]["constraint_count"],
        "events": state["summary"]["event_count"],
        "suggestions": state["summary"]["open_suggestion_count"],
    }, indent=2))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())