#!/usr/bin/env python3 """Persistent supervised-editing state for interactive extraction jobs. The extraction pipeline produces immutable audio artifacts and a batch manifest. This module layers replayable semantic state on top of that manifest: hits, clusters, constraints, events, suggestions, confidence, and undo snapshots. Supervised edits are cheap, explicit, inspectable, and reproducible. A separate supervised export step renders the mutable state into edited WAV/MIDI/ZIP artifacts without mutating the original batch manifest. """ from __future__ import annotations import copy import json import math import time import uuid from pathlib import Path from typing import Any, Callable STATE_VERSION = "interactive-state-v1" STATE_FILENAME = "supervision_state.json" MAX_UNDO = 30 def now() -> float: return round(time.time(), 6) def state_path(output_dir: str | Path) -> Path: return Path(output_dir) / STATE_FILENAME def manifest_path(output_dir: str | Path) -> Path: return Path(output_dir) / "manifest.json" def load_manifest(output_dir: str | Path) -> dict[str, Any]: path = manifest_path(output_dir) if not path.exists(): raise FileNotFoundError(f"manifest.json not found in {Path(output_dir)}") return json.loads(path.read_text(encoding="utf-8")) def _hit_id(hit: dict[str, Any]) -> str: return f"hit:{int(hit.get('index', 0)):05d}" def _cluster_id(raw: Any) -> str: return f"cluster:{raw}" def _base_label(label: str) -> str: text = str(label or "other") return text.rsplit("_", 1)[0] if "_" in text else text def _new_id(prefix: str) -> str: return f"{prefix}:{uuid.uuid4().hex[:10]}" def _safe_float(value: Any, default: float = 0.0) -> float: try: out = float(value) if math.isfinite(out): return out except Exception: pass return default def _safe_int(value: Any, default: int = 0) -> int: try: return int(value) except Exception: return default def _safe_file_component(value: str) -> str: import re text = str(value or "hit").strip().lower() text = re.sub(r"[^a-z0-9._-]+", "_", text) text = re.sub(r"_+", "_", text).strip("._-") return text or "hit" def _snapshot(state: dict[str, Any]) -> dict[str, Any]: snap = copy.deepcopy(state) snap["undo_stack"] = [] return snap def _push_undo(state: dict[str, Any]) -> None: stack = list(state.get("undo_stack") or []) stack.append(_snapshot(state)) del stack[:-MAX_UNDO] state["undo_stack"] = stack def _event(state: dict[str, Any], event_type: str, payload: dict[str, Any] | None = None, source: str = "system") -> dict[str, Any]: event = { "id": _new_id("event"), "type": event_type, "source": source, "created_at": now(), "payload": payload or {}, } state.setdefault("events", []).append(event) return event def _constraint(state: dict[str, Any], constraint_type: str, payload: dict[str, Any], source: str = "user") -> dict[str, Any]: constraint = { "id": _new_id("constraint"), "type": constraint_type, "source": source, "created_at": now(), **payload, } state.setdefault("constraints", []).append(constraint) _event(state, "constraint.created", {"constraint_id": constraint["id"], "type": constraint_type}, source=source) return constraint def _write_state(output_dir: str | Path, state: dict[str, Any]) -> dict[str, Any]: state["updated_at"] = now() path = state_path(output_dir) path.write_text(json.dumps(state, indent=2, sort_keys=True), encoding="utf-8") return state def _cluster_label_for_hit(hit: dict[str, Any]) -> str: return str(hit.get("cluster_label") or f"{hit.get('label', 'other')}_{hit.get('cluster_id', '0')}") def build_initial_state(job_id: str, manifest: dict[str, Any]) -> dict[str, Any]: hits_by_id: dict[str, dict[str, Any]] = {} clusters: dict[str, dict[str, Any]] = {} raw_hits = list(manifest.get("hits") or []) if not raw_hits: # Older manifests may only contain samples. Keep state valid even then. raw_hits = [] for hit in raw_hits: hid = _hit_id(hit) cid = _cluster_id(hit.get("cluster_id", "unclustered")) cluster_label = _cluster_label_for_hit(hit) hits_by_id[hid] = { "id": hid, "index": int(hit.get("index", len(hits_by_id))), "label": str(hit.get("label") or "other"), "cluster_id": cid, "original_cluster_id": cid, "cluster_label": cluster_label, "onset_sec": _safe_float(hit.get("onset_sec")), "duration_ms": _safe_float(hit.get("duration_ms")), "rms_energy": _safe_float(hit.get("rms_energy")), "spectral_centroid_hz": _safe_float(hit.get("spectral_centroid_hz")), "file": hit.get("file"), "is_representative": bool(hit.get("is_representative")), "source": "detected", "suppressed": False, "favorite": False, "review_status": "unreviewed", "confidence": 0.0, "confidence_reasons": [], "explicit": False, } clusters.setdefault( cid, { "id": cid, "label": cluster_label, "classification": _base_label(cluster_label), "hit_ids": [], "representative_hit_id": None, "locked": False, "user_named": False, "confidence": 0.0, "confidence_reasons": [], "suppressed_count": 0, "original_id": cid, }, )["hit_ids"].append(hid) if bool(hit.get("is_representative")): clusters[cid]["representative_hit_id"] = hid for cid, cluster in clusters.items(): if cluster["representative_hit_id"] is None and cluster["hit_ids"]: cluster["representative_hit_id"] = cluster["hit_ids"][0] state = { "version": STATE_VERSION, "job_id": job_id, "created_at": now(), "updated_at": now(), "manifest_fingerprint": _manifest_fingerprint(manifest), "hits": hits_by_id, "clusters": clusters, "constraints": [], "events": [], "suggestions": [], "undo_stack": [], "counters": {"user_clusters": 0}, } recompute_scores(state) _event( state, "job.state.created", { "hit_count": len(hits_by_id), "cluster_count": len(clusters), "manifest_fingerprint": state["manifest_fingerprint"], }, ) return state def _manifest_fingerprint(manifest: dict[str, Any]) -> str: import hashlib payload = { "params": manifest.get("params"), "hit_count": manifest.get("hit_count"), "cluster_count": manifest.get("cluster_count"), "files": manifest.get("files"), "hits": [ { "index": h.get("index"), "cluster_id": h.get("cluster_id"), "file": h.get("file"), "onset_sec": h.get("onset_sec"), } for h in manifest.get("hits", []) ], } return hashlib.sha256(json.dumps(payload, sort_keys=True).encode("utf-8")).hexdigest() def load_or_create_state(job_id: str, output_dir: str | Path) -> dict[str, Any]: path = state_path(output_dir) if path.exists(): state = json.loads(path.read_text(encoding="utf-8")) if state.get("version") != STATE_VERSION: raise ValueError(f"Unsupported supervision state version: {state.get('version')}") return state manifest = load_manifest(output_dir) state = build_initial_state(job_id, manifest) return _write_state(output_dir, state) def _active_hits(state: dict[str, Any], cluster: dict[str, Any]) -> list[dict[str, Any]]: hits = state.get("hits", {}) return [hits[hid] for hid in cluster.get("hit_ids", []) if hid in hits and not hits[hid].get("suppressed")] def recompute_scores(state: dict[str, Any]) -> None: hits = state.get("hits", {}) clusters = state.get("clusters", {}) energies = sorted(_safe_float(hit.get("rms_energy")) for hit in hits.values()) def energy_rank(value: float) -> float: if not energies: return 0.5 less = sum(1 for item in energies if item <= value) return less / max(1, len(energies)) for cluster in clusters.values(): members = [hits[hid] for hid in cluster.get("hit_ids", []) if hid in hits] active = [hit for hit in members if not hit.get("suppressed")] if not members: confidence = 0.15 reasons = ["empty cluster"] else: labels: dict[str, int] = {} for hit in active: labels[hit.get("label", "other")] = labels.get(hit.get("label", "other"), 0) + 1 majority = max(labels.values()) if labels else 0 purity = majority / max(1, len(active)) size_score = min(1.0, math.log2(len(active) + 1) / 4.0) representative_bonus = 0.12 if cluster.get("representative_hit_id") in cluster.get("hit_ids", []) else 0.0 lock_bonus = 0.12 if cluster.get("locked") else 0.0 confidence = (0.42 * purity) + (0.34 * size_score) + representative_bonus + lock_bonus reasons = [] if len(active) <= 1: reasons.append("singleton cluster") if purity < 0.75: reasons.append("mixed labels") if cluster.get("locked"): reasons.append("user locked") if representative_bonus: reasons.append("has representative") cluster["confidence"] = round(max(0.0, min(1.0, confidence)), 4) cluster["confidence_reasons"] = reasons or ["cohesive cluster"] cluster["suppressed_count"] = sum(1 for hit in members if hit.get("suppressed")) for hit in hits.values(): cluster = clusters.get(hit.get("cluster_id"), {}) active_count = len(_active_hits(state, cluster)) if cluster else 0 label_match = _base_label(str(cluster.get("label", ""))) == str(hit.get("label", "")) energy = energy_rank(_safe_float(hit.get("rms_energy"))) duration_ms = _safe_float(hit.get("duration_ms")) duration_score = 0.65 if duration_ms <= 0 else max(0.0, min(1.0, 1.0 - abs(duration_ms - 180.0) / 700.0)) cluster_conf = _safe_float(cluster.get("confidence"), 0.2) confidence = (0.42 * cluster_conf) + (0.18 * min(1.0, active_count / 4.0)) + (0.18 if label_match else 0.0) + (0.12 * energy) + (0.10 * duration_score) reasons = [] if active_count <= 1: reasons.append("singleton") if not label_match: reasons.append("label differs from cluster") if energy < 0.2: reasons.append("low energy") if hit.get("is_representative"): confidence += 0.08 reasons.append("representative") if hit.get("explicit"): confidence += 0.10 reasons.append("explicit user assignment") if hit.get("suppressed"): confidence = min(confidence, 0.25) reasons.append("suppressed") hit["confidence"] = round(max(0.0, min(1.0, confidence)), 4) hit["confidence_reasons"] = reasons or ["consistent assignment"] hit["cluster_label"] = cluster.get("label", hit.get("cluster_label", "unclustered")) def review_queue(state: dict[str, Any], limit: int = 30) -> list[dict[str, Any]]: rows = [] clusters = state.get("clusters", {}) for hit in state.get("hits", {}).values(): cluster = clusters.get(hit.get("cluster_id"), {}) score = 1.0 - _safe_float(hit.get("confidence"), 0.0) if len(cluster.get("hit_ids", [])) <= 1: score += 0.15 if hit.get("suppressed"): score -= 0.35 if hit.get("review_status") == "accepted": score -= 0.25 rows.append( { "hit_id": hit["id"], "hit_index": hit.get("index"), "label": hit.get("label"), "cluster_id": hit.get("cluster_id"), "cluster_label": cluster.get("label"), "confidence": hit.get("confidence", 0.0), "priority": round(max(0.0, score), 4), "reasons": hit.get("confidence_reasons", []), "suppressed": bool(hit.get("suppressed")), "file": hit.get("file"), } ) rows.sort(key=lambda item: (-item["priority"], item["hit_index"] or 0)) return rows[: max(1, min(int(limit), 200))] def _find_similar_hits(state: dict[str, Any], hit_id: str, *, exclude_cluster: str | None = None, include_suppressed: bool = False, limit: int = 12) -> list[tuple[dict[str, Any], float]]: hits = state.get("hits", {}) src = hits[hit_id] src_centroid = _safe_float(src.get("spectral_centroid_hz")) src_energy = _safe_float(src.get("rms_energy")) scored: list[tuple[dict[str, Any], float]] = [] for candidate in hits.values(): if candidate["id"] == hit_id: continue if exclude_cluster and candidate.get("cluster_id") == exclude_cluster: continue if candidate.get("suppressed") and not include_suppressed: continue label_score = 1.0 if candidate.get("label") == src.get("label") else 0.35 centroid_delta = abs(_safe_float(candidate.get("spectral_centroid_hz")) - src_centroid) centroid_score = max(0.0, 1.0 - centroid_delta / 6000.0) energy_delta = abs(_safe_float(candidate.get("rms_energy")) - src_energy) energy_score = max(0.0, 1.0 - energy_delta / max(src_energy, 1e-4, _safe_float(candidate.get("rms_energy")))) score = (0.48 * label_score) + (0.34 * centroid_score) + (0.18 * energy_score) if score >= 0.62: scored.append((candidate, round(score, 4))) scored.sort(key=lambda item: (-item[1], item[0].get("index", 0))) return scored[:limit] def suggestion_diff(state: dict[str, Any], suggestion: dict[str, Any]) -> dict[str, Any]: """Build an exact before/after preview for a suggestion against current state.""" hits = state.get("hits", {}) clusters = state.get("clusters", {}) stype = suggestion.get("type") hit_ids = [hid for hid in suggestion.get("hit_ids", []) if hid in hits] def cluster_snapshot(cluster_id: str | None) -> dict[str, Any]: cluster = clusters.get(cluster_id or "", {}) members = [hid for hid in cluster.get("hit_ids", []) if hid in hits] active = [hid for hid in members if not hits[hid].get("suppressed")] return { "cluster_id": cluster_id, "label": cluster.get("label", cluster_id), "active_count": len(active), "total_count": len(members), "suppressed_count": sum(1 for hid in members if hits[hid].get("suppressed")), } rows = [] cluster_ids: set[str] = set() for hid in hit_ids: hit = hits[hid] source_cluster_id = hit.get("cluster_id") target_cluster_id = suggestion.get("target_cluster_id") if stype in {"move-hits", "split-hits"} else source_cluster_id cluster_ids.add(str(source_cluster_id)) if target_cluster_id: cluster_ids.add(str(target_cluster_id)) rows.append( { "hit_id": hid, "hit_index": hit.get("index"), "label": hit.get("label"), "from_cluster_id": source_cluster_id, "from_cluster_label": clusters.get(source_cluster_id, {}).get("label"), "to_cluster_id": target_cluster_id, "to_cluster_label": clusters.get(target_cluster_id, {}).get("label") if target_cluster_id else None, "before_suppressed": bool(hit.get("suppressed")), "after_suppressed": bool(hit.get("suppressed")) or stype == "suppress-hits", "confidence": hit.get("confidence"), } ) before = {cid: cluster_snapshot(cid) for cid in sorted(cluster_ids)} after = copy.deepcopy(before) if stype in {"move-hits", "split-hits"}: target = suggestion.get("target_cluster_id") for row in rows: source = row.get("from_cluster_id") if source in after and source != target: after[source]["active_count"] = max(0, after[source]["active_count"] - 1) after[source]["total_count"] = max(0, after[source]["total_count"] - 1) if target in after and source != target: after[target]["active_count"] += 1 after[target]["total_count"] += 1 elif stype == "suppress-hits": for row in rows: source = row.get("from_cluster_id") if source in after and not row.get("before_suppressed"): after[source]["active_count"] = max(0, after[source]["active_count"] - 1) after[source]["suppressed_count"] += 1 return { "type": stype, "affected_hit_count": len(rows), "hits": rows, "clusters_before": before, "clusters_after": after, } def _add_suggestion(state: dict[str, Any], suggestion_type: str, payload: dict[str, Any], confidence: float, reason: str) -> dict[str, Any]: suggestion = { "id": _new_id("suggestion"), "type": suggestion_type, "status": "open", "created_at": now(), "confidence": round(max(0.0, min(1.0, confidence)), 4), "reason": reason, **payload, } suggestion["diff"] = suggestion_diff(state, suggestion) state.setdefault("suggestions", []).append(suggestion) _event(state, "suggestion.created", {"suggestion_id": suggestion["id"], "type": suggestion_type, "reason": reason}) return suggestion def _rebuild_cluster_labels(state: dict[str, Any]) -> None: hits = state.get("hits", {}) for cluster in state.get("clusters", {}).values(): for hid in cluster.get("hit_ids", []): if hid in hits: hits[hid]["cluster_label"] = cluster.get("label", "unclustered") def move_hit(output_dir: str | Path, job_id: str, hit_id: str, target_cluster_id: str, source: str = "user") -> dict[str, Any]: state = load_or_create_state(job_id, output_dir) hits = state.get("hits", {}) clusters = state.get("clusters", {}) if hit_id not in hits: raise KeyError(f"Unknown hit: {hit_id}") if target_cluster_id not in clusters: raise KeyError(f"Unknown cluster: {target_cluster_id}") hit = hits[hit_id] source_cluster_id = hit.get("cluster_id") if source_cluster_id == target_cluster_id: hit["review_status"] = "accepted" recompute_scores(state) return _write_state(output_dir, state) _push_undo(state) if source_cluster_id in clusters: clusters[source_cluster_id]["hit_ids"] = [hid for hid in clusters[source_cluster_id].get("hit_ids", []) if hid != hit_id] clusters[target_cluster_id].setdefault("hit_ids", []) if hit_id not in clusters[target_cluster_id]["hit_ids"]: clusters[target_cluster_id]["hit_ids"].append(hit_id) hit["cluster_id"] = target_cluster_id hit["cluster_label"] = clusters[target_cluster_id].get("label", target_cluster_id) hit["explicit"] = True hit["review_status"] = "accepted" target_rep = clusters[target_cluster_id].get("representative_hit_id") _constraint(state, "force-cluster", {"hit_id": hit_id, "cluster_id": target_cluster_id}, source=source) if target_rep and target_rep != hit_id: _constraint(state, "must-link", {"a": hit_id, "b": target_rep}, source=source) _event(state, "hit.moved", {"hit_id": hit_id, "from_cluster_id": source_cluster_id, "to_cluster_id": target_cluster_id}, source=source) similar = _find_similar_hits(state, hit_id, exclude_cluster=target_cluster_id, limit=10) suggested_ids = [item[0]["id"] for item in similar if item[1] >= 0.72] if suggested_ids: avg = sum(score for _, score in similar if _["id"] in suggested_ids) / len(suggested_ids) _add_suggestion( state, "move-hits", {"hit_ids": suggested_ids, "target_cluster_id": target_cluster_id, "preview_count": len(suggested_ids)}, avg, f"Similar label/spectral/energy profile to {hit_id}", ) _rebuild_cluster_labels(state) recompute_scores(state) return _write_state(output_dir, state) def pull_hit_to_new_cluster(output_dir: str | Path, job_id: str, hit_id: str, label: str | None = None, source: str = "user") -> dict[str, Any]: state = load_or_create_state(job_id, output_dir) hits = state.get("hits", {}) clusters = state.get("clusters", {}) if hit_id not in hits: raise KeyError(f"Unknown hit: {hit_id}") hit = hits[hit_id] source_cluster_id = hit.get("cluster_id") source_rep = clusters.get(source_cluster_id, {}).get("representative_hit_id") _push_undo(state) state.setdefault("counters", {})["user_clusters"] = int(state.get("counters", {}).get("user_clusters", 0)) + 1 base = label or f"{hit.get('label', 'hit')}_user_{state['counters']['user_clusters']}" new_cluster_id = _new_id("cluster:user") if source_cluster_id in clusters: clusters[source_cluster_id]["hit_ids"] = [hid for hid in clusters[source_cluster_id].get("hit_ids", []) if hid != hit_id] clusters[new_cluster_id] = { "id": new_cluster_id, "label": base, "classification": _base_label(base), "hit_ids": [hit_id], "representative_hit_id": hit_id, "locked": False, "user_named": bool(label), "confidence": 0.0, "confidence_reasons": [], "suppressed_count": 0, "original_id": None, } hit["cluster_id"] = new_cluster_id hit["cluster_label"] = base hit["explicit"] = True hit["review_status"] = "accepted" if source_rep and source_rep != hit_id: _constraint(state, "cannot-link", {"a": hit_id, "b": source_rep}, source=source) _constraint(state, "force-cluster", {"hit_id": hit_id, "cluster_id": new_cluster_id}, source=source) _event(state, "hit.pulled_out", {"hit_id": hit_id, "from_cluster_id": source_cluster_id, "to_cluster_id": new_cluster_id}, source=source) similar = _find_similar_hits(state, hit_id, exclude_cluster=new_cluster_id, limit=8) split_ids = [item[0]["id"] for item in similar if item[0].get("cluster_id") == source_cluster_id and item[1] >= 0.70] if split_ids: _add_suggestion( state, "split-hits", {"hit_ids": split_ids, "source_cluster_id": source_cluster_id, "target_cluster_id": new_cluster_id, "preview_count": len(split_ids)}, 0.76, f"Similar to pulled-out hit {hit_id}; preview split from original cluster", ) _rebuild_cluster_labels(state) recompute_scores(state) return _write_state(output_dir, state) def lock_cluster(output_dir: str | Path, job_id: str, cluster_id: str, locked: bool = True, source: str = "user") -> dict[str, Any]: state = load_or_create_state(job_id, output_dir) clusters = state.get("clusters", {}) if cluster_id not in clusters: raise KeyError(f"Unknown cluster: {cluster_id}") _push_undo(state) clusters[cluster_id]["locked"] = bool(locked) _constraint(state, "lock-cluster", {"cluster_id": cluster_id, "locked": bool(locked)}, source=source) _event(state, "cluster.locked" if locked else "cluster.unlocked", {"cluster_id": cluster_id}, source=source) recompute_scores(state) return _write_state(output_dir, state) def suppress_hit(output_dir: str | Path, job_id: str, hit_id: str, reason: str = "bleed", source: str = "user") -> dict[str, Any]: state = load_or_create_state(job_id, output_dir) hits = state.get("hits", {}) if hit_id not in hits: raise KeyError(f"Unknown hit: {hit_id}") _push_undo(state) hit = hits[hit_id] hit["suppressed"] = True hit["review_status"] = "suppressed" hit["explicit"] = True _constraint(state, "suppress-pattern", {"example_hit_id": hit_id, "reason": reason}, source=source) _event(state, "hit.suppressed", {"hit_id": hit_id, "reason": reason}, source=source) similar = _find_similar_hits(state, hit_id, include_suppressed=False, limit=16) suggested_ids = [item[0]["id"] for item in similar if item[1] >= 0.72 and _safe_float(item[0].get("rms_energy")) <= _safe_float(hit.get("rms_energy")) * 1.35] if suggested_ids: _add_suggestion( state, "suppress-hits", {"hit_ids": suggested_ids, "reason_code": reason, "preview_count": len(suggested_ids)}, 0.74, f"Similar low-energy profile to suppressed {reason} example {hit_id}", ) recompute_scores(state) return _write_state(output_dir, state) def restore_hit(output_dir: str | Path, job_id: str, hit_id: str, source: str = "user") -> dict[str, Any]: state = load_or_create_state(job_id, output_dir) hits = state.get("hits", {}) if hit_id not in hits: raise KeyError(f"Unknown hit: {hit_id}") _push_undo(state) hit = hits[hit_id] hit["suppressed"] = False hit["review_status"] = "unreviewed" if hit.get("review_status") == "suppressed" else hit.get("review_status", "unreviewed") hit["explicit"] = True _constraint(state, "restore-hit", {"hit_id": hit_id}, source=source) _event(state, "hit.restored", {"hit_id": hit_id}, source=source) recompute_scores(state) return _write_state(output_dir, state) def force_onset( output_dir: str | Path, job_id: str, onset_sec: float, *, duration_ms: float | None = None, label: str | None = None, target_cluster_id: str | None = None, pre_pad_sec: float = 0.003, source: str = "user", ) -> dict[str, Any]: """Create a user-forced hit from ``stem.wav`` and add it to semantic state.""" import librosa import numpy as np import soundfile as sf from sample_extractor import Hit as AudioHit, classify_hit out = Path(output_dir) stem_path = out / "stem.wav" if not stem_path.exists(): raise FileNotFoundError("stem.wav is required before forcing onsets") state = load_or_create_state(job_id, out) hits = state.setdefault("hits", {}) clusters = state.setdefault("clusters", {}) onset = max(0.0, _safe_float(onset_sec)) audio, sr = sf.read(stem_path, dtype="float32", always_2d=False) if audio.ndim > 1: audio = audio.mean(axis=1) audio = np.asarray(audio, dtype=np.float32) duration = (_safe_float(duration_ms, 0.0) / 1000.0) if duration_ms else 0.0 if duration <= 0: future_onsets = sorted( _safe_float(hit.get("onset_sec")) for hit in hits.values() if not hit.get("suppressed") and _safe_float(hit.get("onset_sec")) > onset + 0.01 ) next_onset = future_onsets[0] if future_onsets else None duration = min(1.5, max(0.08, (next_onset - onset) if next_onset is not None else 0.45)) duration = max(0.02, min(10.0, duration)) start = max(0, int((onset - max(0.0, pre_pad_sec)) * sr)) end = min(len(audio), start + int(duration * sr)) if end <= start: raise ValueError("Forced onset is outside the available stem audio") segment = audio[start:end].copy() fade_len = min(int(0.003 * sr), len(segment) // 4) if fade_len > 0: segment[-fade_len:] *= np.linspace(1, 0, fade_len) rms = float(np.sqrt(np.mean(segment**2))) if len(segment) else 0.0 spectral_centroid = float(librosa.feature.spectral_centroid(y=segment, sr=sr).mean()) if len(segment) >= 32 else 0.0 index = max((_safe_int(hit.get("index"), -1) for hit in hits.values()), default=-1) + 1 tmp_hit = AudioHit(audio=segment, sr=sr, onset_time=onset, duration=len(segment) / sr, index=index, rms_energy=rms, spectral_centroid=spectral_centroid) inferred_label = label or classify_hit(tmp_hit) tmp_hit.label = inferred_label hit_id = _hit_id({"index": index}) safe_label = _safe_file_component(inferred_label or "forced") rel_file = f"review/hits/hit_{index:05d}_{safe_label}_forced.wav" full_path = out / rel_file full_path.parent.mkdir(parents=True, exist_ok=True) sf.write(full_path, segment, sr, subtype="PCM_24") _push_undo(state) if target_cluster_id and target_cluster_id not in clusters: raise KeyError(f"Unknown cluster: {target_cluster_id}") if not target_cluster_id: state.setdefault("counters", {})["user_clusters"] = int(state.get("counters", {}).get("user_clusters", 0)) + 1 target_cluster_id = _new_id("cluster:user") cluster_label = f"{_safe_file_component(inferred_label)}_forced_{state['counters']['user_clusters']}" clusters[target_cluster_id] = { "id": target_cluster_id, "label": cluster_label, "classification": _base_label(cluster_label), "hit_ids": [], "representative_hit_id": hit_id, "locked": False, "user_named": bool(label), "confidence": 0.0, "confidence_reasons": [], "suppressed_count": 0, "original_id": None, } cluster_label = clusters[target_cluster_id].get("label", target_cluster_id) hits[hit_id] = { "id": hit_id, "index": index, "label": str(inferred_label or "other"), "cluster_id": target_cluster_id, "original_cluster_id": None, "cluster_label": cluster_label, "onset_sec": round(onset, 6), "duration_ms": round((len(segment) / sr) * 1000.0, 1), "rms_energy": round(rms, 6), "spectral_centroid_hz": round(spectral_centroid, 1), "file": rel_file, "is_representative": False, "source": "forced", "suppressed": False, "favorite": False, "review_status": "accepted", "confidence": 0.0, "confidence_reasons": [], "explicit": True, } clusters[target_cluster_id].setdefault("hit_ids", []) if hit_id not in clusters[target_cluster_id]["hit_ids"]: clusters[target_cluster_id]["hit_ids"].append(hit_id) if not clusters[target_cluster_id].get("representative_hit_id"): clusters[target_cluster_id]["representative_hit_id"] = hit_id _constraint(state, "force-onset", {"hit_id": hit_id, "onset_sec": round(onset, 6)}, source=source) _constraint(state, "force-cluster", {"hit_id": hit_id, "cluster_id": target_cluster_id}, source=source) _event(state, "hit.force_onset", {"hit_id": hit_id, "onset_sec": round(onset, 6), "cluster_id": target_cluster_id}, source=source) _rebuild_cluster_labels(state) recompute_scores(state) return _write_state(out, state) def set_hit_review_status(output_dir: str | Path, job_id: str, hit_id: str, status: str = "accepted", source: str = "user") -> dict[str, Any]: if status not in {"unreviewed", "accepted", "favorite"}: raise ValueError("status must be unreviewed, accepted, or favorite") state = load_or_create_state(job_id, output_dir) if hit_id not in state.get("hits", {}): raise KeyError(f"Unknown hit: {hit_id}") _push_undo(state) hit = state["hits"][hit_id] hit["review_status"] = status if status == "favorite": hit["favorite"] = True cid = hit.get("cluster_id") if cid in state.get("clusters", {}): state["clusters"][cid]["representative_hit_id"] = hit_id _constraint(state, "pin-representative", {"hit_id": hit_id, "cluster_id": cid}, source=source) _event(state, "hit.reviewed", {"hit_id": hit_id, "status": status}, source=source) recompute_scores(state) return _write_state(output_dir, state) def accept_suggestion(output_dir: str | Path, job_id: str, suggestion_id: str) -> dict[str, Any]: state = load_or_create_state(job_id, output_dir) suggestion = next((s for s in state.get("suggestions", []) if s.get("id") == suggestion_id), None) if not suggestion: raise KeyError(f"Unknown suggestion: {suggestion_id}") if suggestion.get("status") != "open": return state _push_undo(state) stype = suggestion.get("type") if stype in {"move-hits", "split-hits"}: target = suggestion.get("target_cluster_id") for hid in suggestion.get("hit_ids", []): if hid in state.get("hits", {}) and target in state.get("clusters", {}): current = state["hits"][hid].get("cluster_id") if current in state["clusters"]: state["clusters"][current]["hit_ids"] = [x for x in state["clusters"][current].get("hit_ids", []) if x != hid] state["clusters"][target].setdefault("hit_ids", []) if hid not in state["clusters"][target]["hit_ids"]: state["clusters"][target]["hit_ids"].append(hid) state["hits"][hid]["cluster_id"] = target state["hits"][hid]["explicit"] = True _constraint(state, "force-cluster", {"hit_id": hid, "cluster_id": target}, source="accepted-suggestion") elif stype == "suppress-hits": for hid in suggestion.get("hit_ids", []): if hid in state.get("hits", {}): state["hits"][hid]["suppressed"] = True state["hits"][hid]["review_status"] = "suppressed" _constraint(state, "suppress-pattern", {"example_hit_id": hid, "reason": suggestion.get("reason_code", "bleed")}, source="accepted-suggestion") else: raise ValueError(f"Unsupported suggestion type: {stype}") suggestion["status"] = "accepted" suggestion["resolved_at"] = now() _event(state, "suggestion.accepted", {"suggestion_id": suggestion_id, "type": stype}, source="user") _rebuild_cluster_labels(state) recompute_scores(state) return _write_state(output_dir, state) def reject_suggestion(output_dir: str | Path, job_id: str, suggestion_id: str) -> dict[str, Any]: state = load_or_create_state(job_id, output_dir) suggestion = next((s for s in state.get("suggestions", []) if s.get("id") == suggestion_id), None) if not suggestion: raise KeyError(f"Unknown suggestion: {suggestion_id}") _push_undo(state) suggestion["status"] = "rejected" suggestion["resolved_at"] = now() _event(state, "suggestion.rejected", {"suggestion_id": suggestion_id, "type": suggestion.get("type")}, source="user") return _write_state(output_dir, state) def undo_last(output_dir: str | Path, job_id: str) -> dict[str, Any]: state = load_or_create_state(job_id, output_dir) stack = list(state.get("undo_stack") or []) if not stack: return state restored = stack.pop() restored["undo_stack"] = stack _event(restored, "state.undo", {"restored_for_job_id": job_id}, source="user") recompute_scores(restored) return _write_state(output_dir, restored) def explain_cluster(state: dict[str, Any], cluster_id: str) -> dict[str, Any]: clusters = state.get("clusters", {}) hits = state.get("hits", {}) if cluster_id not in clusters: raise KeyError(f"Unknown cluster: {cluster_id}") cluster = clusters[cluster_id] members = [hits[hid] for hid in cluster.get("hit_ids", []) if hid in hits] active = [h for h in members if not h.get("suppressed")] constraints = [c for c in state.get("constraints", []) if c.get("cluster_id") == cluster_id or c.get("hit_id") in cluster.get("hit_ids", []) or c.get("a") in cluster.get("hit_ids", []) or c.get("b") in cluster.get("hit_ids", [])] outliers = sorted(active, key=lambda h: h.get("confidence", 0.0))[:8] labels: dict[str, int] = {} for hit in active: labels[hit.get("label", "other")] = labels.get(hit.get("label", "other"), 0) + 1 return { "cluster_id": cluster_id, "label": cluster.get("label"), "locked": bool(cluster.get("locked")), "confidence": cluster.get("confidence"), "confidence_reasons": cluster.get("confidence_reasons", []), "representative_hit_id": cluster.get("representative_hit_id"), "hit_count": len(members), "active_hit_count": len(active), "suppressed_count": sum(1 for hit in members if hit.get("suppressed")), "label_distribution": labels, "outliers": [{"hit_id": h["id"], "hit_index": h.get("index"), "confidence": h.get("confidence"), "reasons": h.get("confidence_reasons", [])} for h in outliers], "constraints": constraints[-20:], "summary": f"{cluster.get('label')} has {len(active)} active hits, confidence {cluster.get('confidence')}, and {len(constraints)} relevant constraints.", } def public_state(state: dict[str, Any], url_for: Callable[[str], str] | None = None, review_limit: int = 30) -> dict[str, Any]: recompute_scores(state) hits = copy.deepcopy(list(state.get("hits", {}).values())) clusters = copy.deepcopy(list(state.get("clusters", {}).values())) for hit in hits: if url_for and hit.get("file"): hit["url"] = url_for(hit["file"]) clusters.sort(key=lambda c: (-len(c.get("hit_ids", [])), c.get("label", ""))) hits.sort(key=lambda h: h.get("index", 0)) open_suggestions = [copy.deepcopy(s) for s in state.get("suggestions", []) if s.get("status") == "open"] for suggestion in open_suggestions: suggestion["diff"] = suggestion.get("diff") or suggestion_diff(state, suggestion) open_suggestions.sort(key=lambda s: (-_safe_float(s.get("confidence")), s.get("created_at", 0))) latest_export = copy.deepcopy(state.get("latest_export")) if latest_export and url_for and latest_export.get("path"): latest_export["url"] = url_for(latest_export["path"]) return { "version": state.get("version"), "job_id": state.get("job_id"), "created_at": state.get("created_at"), "updated_at": state.get("updated_at"), "summary": { "hit_count": len(hits), "cluster_count": len(clusters), "constraint_count": len(state.get("constraints", [])), "event_count": len(state.get("events", [])), "open_suggestion_count": len(open_suggestions), "suppressed_hit_count": sum(1 for h in hits if h.get("suppressed")), "locked_cluster_count": sum(1 for c in clusters if c.get("locked")), "undo_available": bool(state.get("undo_stack")), "forced_hit_count": sum(1 for h in hits if h.get("source") == "forced"), "latest_export": latest_export, }, "hits": hits, "clusters": clusters, "constraints": state.get("constraints", [])[-100:], "events": state.get("events", [])[-120:], "suggestions": open_suggestions[:50], "review_queue": review_queue(state, review_limit), } def pin_representative(output_dir: str | Path, job_id: str, cluster_id: str, hit_id: str, source: str = "user") -> dict[str, Any]: """Persistently choose a representative hit for a cluster/card.""" state = load_or_create_state(job_id, output_dir) clusters = state.get("clusters", {}) hits = state.get("hits", {}) if cluster_id not in clusters: raise KeyError(f"Unknown cluster: {cluster_id}") if hit_id not in hits: raise KeyError(f"Unknown hit: {hit_id}") if hit_id not in clusters[cluster_id].get("hit_ids", []): raise ValueError(f"Hit {hit_id} is not a member of {cluster_id}") _push_undo(state) for hid in clusters[cluster_id].get("hit_ids", []): if hid in hits: hits[hid]["is_representative"] = (hid == hit_id) clusters[cluster_id]["representative_hit_id"] = hit_id hits[hit_id]["favorite"] = True hits[hit_id]["review_status"] = "favorite" hits[hit_id]["explicit"] = True _constraint(state, "pin-representative", {"hit_id": hit_id, "cluster_id": cluster_id}, source=source) _event(state, "cluster.representative_pinned", {"hit_id": hit_id, "cluster_id": cluster_id}, source=source) recompute_scores(state) return _write_state(output_dir, state) def draw_next_representative(output_dir: str | Path, job_id: str, cluster_id: str, source: str = "user") -> dict[str, Any]: """Cycle a cluster/card to the next available non-suppressed candidate.""" state = load_or_create_state(job_id, output_dir) clusters = state.get("clusters", {}) hits = state.get("hits", {}) if cluster_id not in clusters: raise KeyError(f"Unknown cluster: {cluster_id}") cluster = clusters[cluster_id] active_ids = [hid for hid in cluster.get("hit_ids", []) if hid in hits and not hits[hid].get("suppressed")] if not active_ids: raise ValueError(f"Cluster {cluster_id} has no active hits") current = cluster.get("representative_hit_id") if current in active_ids: next_id = active_ids[(active_ids.index(current) + 1) % len(active_ids)] else: next_id = active_ids[0] return pin_representative(output_dir, job_id, cluster_id, next_id, source=source) def edit_hit_timing( output_dir: str | Path, job_id: str, hit_id: str, *, start_offset_ms: float = 0.0, tail_offset_ms: float = 0.0, source: str = "user", ) -> dict[str, Any]: """Rewrite one hit preview from stem.wav and persist the timing edit. ``start_offset_ms`` trims from the front when positive and extends earlier when negative. ``tail_offset_ms`` extends when positive and trims the tail when negative. The selected hit's file path is replaced so cards and supervised exports immediately use the edited audio. """ import numpy as np import soundfile as sf import librosa out = Path(output_dir) state = load_or_create_state(job_id, out) hits = state.get("hits", {}) if hit_id not in hits: raise KeyError(f"Unknown hit: {hit_id}") hit = hits[hit_id] stem_path = out / "stem.wav" if not stem_path.exists(): raise FileNotFoundError("stem.wav is required for timing edits") audio, sr = sf.read(stem_path, dtype="float32", always_2d=False) if audio.ndim > 1: audio = audio.mean(axis=1) audio = np.asarray(audio, dtype=np.float32) original_onset = _safe_float(hit.get("onset_sec")) original_duration = max(0.02, _safe_float(hit.get("duration_ms"), 100.0) / 1000.0) start_offset = _safe_float(start_offset_ms) / 1000.0 tail_offset = _safe_float(tail_offset_ms) / 1000.0 new_onset = max(0.0, original_onset + start_offset) new_duration = max(0.02, original_duration - start_offset + tail_offset) start = max(0, int(round(new_onset * sr))) end = min(len(audio), start + int(round(new_duration * sr))) if end <= start: raise ValueError("Edited sample range is outside the available stem audio") segment = audio[start:end].copy() fade_len = min(int(0.003 * sr), len(segment) // 4) if fade_len > 0: segment[-fade_len:] *= np.linspace(1, 0, fade_len) rms = float(np.sqrt(np.mean(segment**2))) if len(segment) else 0.0 spectral_centroid = float(librosa.feature.spectral_centroid(y=segment, sr=sr).mean()) if len(segment) >= 32 else 0.0 safe_label = _safe_file_component(hit.get("label") or "edited") rel_file = f"overrides/hits/hit_{_safe_int(hit.get('index')):05d}_{safe_label}_edited.wav" full_path = out / rel_file full_path.parent.mkdir(parents=True, exist_ok=True) sf.write(full_path, segment, sr, subtype="PCM_24") _push_undo(state) hit["onset_sec"] = round(new_onset, 6) hit["duration_ms"] = round((len(segment) / sr) * 1000.0, 1) hit["rms_energy"] = round(rms, 6) hit["spectral_centroid_hz"] = round(spectral_centroid, 1) hit["file"] = rel_file hit["explicit"] = True hit["review_status"] = "accepted" _constraint(state, "edit-hit-timing", {"hit_id": hit_id, "start_offset_ms": round(_safe_float(start_offset_ms), 3), "tail_offset_ms": round(_safe_float(tail_offset_ms), 3)}, source=source) _event(state, "hit.timing_edited", {"hit_id": hit_id, "file": rel_file, "onset_sec": hit["onset_sec"], "duration_ms": hit["duration_ms"]}, source=source) recompute_scores(state) return _write_state(out, state)