diff --git a/AGENTS.md b/AGENTS.md index 93ed96d2f..9a74cc435 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -191,11 +191,14 @@ Each domain has exactly **one** write-owning module (or one tightly-scoped famil | Owner voice candidate (`awareness/owner_candidate.npz`) | `solstone/apps/speakers/owner.py` | | Speaker discovery clusters (`awareness/discovery_clusters.json`, `awareness/discovery_clusters.resolved.json`) | `solstone/apps/speakers/discovery.py` | | Speaker candidate pool (`awareness/speaker_candidates.json`) | `solstone/apps/speakers/candidate_tracker.py` | +| Speaker identify operation ledger (`speakers/identify-operations.jsonl`) | `solstone/think/speaker_identify_operations.py` | | Entity resolution ambiguities (`entities/ambiguities.jsonl`) | `solstone/think/entities/ambiguities.py` | | Entity merge candidates (`entities/review-candidates.jsonl`) | `solstone/think/entities/review_candidates.py` | | Facet review candidates (`facets/review-candidates.jsonl`) | `solstone/think/facet_review_candidates.py` | | Speaker review candidates (`speakers/review-candidates.jsonl`) | `solstone/think/speaker_review_candidates.py` | | Speaker candidate-pair review candidates (`speakers/candidate-pair-review-candidates.jsonl`) | `solstone/think/speaker_candidate_pair_review_candidates.py` | +| Speaker discovery cluster dismissals (`speakers/cluster-dismissals.jsonl`) | `solstone/think/speaker_cluster_dismissals.py` | +| Speaker keep-separate assertions (`speakers/keep-separate.jsonl`) | `solstone/think/speaker_keep_separate.py` | | Facets (`facets/*/facet.json`, `facets/*/relationships/`) | `solstone/think/facets.py` + `solstone/apps/facets/*` (if/when created) | | Observations (`observations.jsonl`) | `solstone/think/entities/observations.py` | | Activities (`facets/*/activities/*.jsonl`) | `solstone/think/activities.py` | diff --git a/scripts/check_journal_io_access.py b/scripts/check_journal_io_access.py index 3896508a0..0bb43b7d7 100644 --- a/scripts/check_journal_io_access.py +++ b/scripts/check_journal_io_access.py @@ -133,6 +133,9 @@ OWNER_FILES: frozenset[str] = frozenset( "solstone/think/entities/voiceprints.py", "solstone/think/facet_review_candidates.py", "solstone/think/speaker_candidate_pair_review_candidates.py", + "solstone/think/speaker_cluster_dismissals.py", + "solstone/think/speaker_identify_operations.py", + "solstone/think/speaker_keep_separate.py", "solstone/think/speaker_review_candidates.py", "solstone/think/facets.py", "solstone/think/identity.py", diff --git a/solstone/apps/speakers/attribution.py b/solstone/apps/speakers/attribution.py index b2891f30c..4f4884fc1 100644 --- a/solstone/apps/speakers/attribution.py +++ b/solstone/apps/speakers/attribution.py @@ -1226,6 +1226,27 @@ def _load_corrections_by_sentence(seg_dir: Path) -> dict[int, dict]: return corrected +def _apply_correction_overlay(label: dict, correction: dict) -> bool: + speaker = correction.get("corrected_speaker") + if speaker is None: + if correction.get("correction_kind") != "identify_undo": + return False + label["speaker"] = None + label["confidence"] = None + label["method"] = None + return True + + label["speaker"] = speaker + label["confidence"] = "high" + if correction.get("original_speaker") == speaker: + label["method"] = "user_confirmed" + elif correction.get("original_speaker") is None: + label["method"] = "user_assigned" + else: + label["method"] = "user_corrected" + return True + + def _speaker_labels_payload( seg_dir: Path, labels: list[dict], @@ -1243,16 +1264,7 @@ def _speaker_labels_payload( sid = label.get("sentence_id") if sid is not None and int(sid) in corrected: corr = corrected[int(sid)] - speaker = corr.get("corrected_speaker") - if speaker is not None: - label["speaker"] = speaker - label["confidence"] = "high" - if corr.get("original_speaker") == speaker: - label["method"] = "user_confirmed" - elif corr.get("original_speaker") is None: - label["method"] = "user_assigned" - else: - label["method"] = "user_corrected" + _apply_correction_overlay(label, corr) logger.info( "Preserved %d user corrections in %s", len(corrected), @@ -1625,6 +1637,101 @@ def apply_label_patches( update_speaker_labels(seg_dir, transform) +def restore_label_rows( + seg_dir: Path, + restorations: list[dict[str, Any]], +) -> dict[str, Any]: + """Compare-restore speaker label rows for identify undo.""" + report: dict[str, Any] = { + "restored_count": 0, + "removed_inserted_count": 0, + "patched_existing_count": 0, + "skipped_count": 0, + "skipped_reasons": {"missing": 0, "changed": 0}, + } + if not restorations: + return report + + def transform(current: dict | None) -> dict | None: + if isinstance(current, dict): + base = dict(current) + labels_value = current.get("labels", []) + labels = [ + dict(label) if isinstance(label, dict) else label + for label in (labels_value if isinstance(labels_value, list) else []) + ] + else: + base = {} + labels = [] + + labels_by_sid: dict[int, dict] = {} + for label in labels: + if not isinstance(label, dict): + continue + sid = _label_sentence_id(label) + if sid is not None: + labels_by_sid[sid] = label + + changed = False + remove_sids: set[int] = set() + for restoration in restorations: + sid = int(restoration["sentence_id"]) + expected = restoration.get("expected_current_label") + current_label = labels_by_sid.get(sid) + if current_label is None: + report["skipped_reasons"]["missing"] += 1 + continue + if current_label != expected: + report["skipped_reasons"]["changed"] += 1 + continue + + prior_state = restoration.get("prior_state") + if prior_state == "absent": + remove_sids.add(sid) + report["removed_inserted_count"] += 1 + elif prior_state == "present": + prior_label = restoration.get("prior_label") + if not isinstance(prior_label, dict): + raise ValueError("present label restoration requires prior_label") + current_label.clear() + current_label.update(prior_label) + report["patched_existing_count"] += 1 + else: + raise ValueError( + f"unknown prior_state for label restore: {prior_state}" + ) + report["restored_count"] += 1 + changed = True + + report["skipped_count"] = sum(report["skipped_reasons"].values()) + if not changed: + return None + + if remove_sids: + labels = [ + label + for label in labels + if not ( + isinstance(label, dict) + and (sid := _label_sentence_id(label)) is not None + and sid in remove_sids + ) + ] + + labels = sorted( + labels, + key=lambda label: ( + _label_sentence_id(label) is None if isinstance(label, dict) else True, + _label_sentence_id(label) if isinstance(label, dict) else 0, + ), + ) + base["labels"] = labels + return base + + update_speaker_labels(seg_dir, transform) + return report + + # --------------------------------------------------------------------------- # Voiceprint accumulation # --------------------------------------------------------------------------- diff --git a/solstone/apps/speakers/candidate_tracker.py b/solstone/apps/speakers/candidate_tracker.py index caa1d100c..511522f58 100644 --- a/solstone/apps/speakers/candidate_tracker.py +++ b/solstone/apps/speakers/candidate_tracker.py @@ -40,7 +40,7 @@ from solstone.think.journal_io import ( hold_lock, read_json, ) -from solstone.think.utils import get_journal +from solstone.think.utils import get_journal, now_ms, segment_start_ts_ms # Synthetic cluster label for producer-proven solo speaker segments. SOLO_CLUSTER_LABEL: int = -1 @@ -128,6 +128,25 @@ class CandidateProfile: ) +@dataclass(frozen=True) +class RetroactiveConfirmPlan: + """Read-only plan for confirming a candidate and backfilling voiceprints.""" + + matched: bool + match_score: float | None + candidate_id: int | None + entity_id: str + candidate_before: dict[str, Any] | None + candidate_after: dict[str, Any] | None + preexisting_voiceprint_keys: tuple[tuple[str, str, str, int], ...] + voiceprints_to_add: tuple[dict[str, Any], ...] + voiceprint_items_to_add: tuple[tuple[np.ndarray, dict[str, Any]], ...] = field( + default_factory=tuple, + repr=False, + compare=False, + ) + + def _routes_helpers(): from solstone.apps.speakers.routes import ( _load_embeddings_file, @@ -137,6 +156,116 @@ def _routes_helpers(): return _load_embeddings_file, _normalize_embedding +def _empty_retroactive_plan( + entity_id: str, + *, + match_score: float | None = None, +) -> RetroactiveConfirmPlan: + return RetroactiveConfirmPlan( + matched=False, + match_score=match_score, + candidate_id=None, + entity_id=entity_id, + candidate_before=None, + candidate_after=None, + preexisting_voiceprint_keys=(), + voiceprints_to_add=(), + voiceprint_items_to_add=(), + ) + + +def _entity_voiceprint_snapshot( + entity_id: str, + normalize_embedding: Callable[[np.ndarray], np.ndarray | None], +) -> tuple[set[tuple[str, str, str, int]], int, np.ndarray | None]: + from solstone.think.entities import load_entity_voiceprints_file + + result = load_entity_voiceprints_file(entity_id) + if result is None: + return set(), 0, None + + embeddings, metadata = result + keys: set[tuple[str, str, str, int]] = set() + for meta in metadata: + try: + keys.add( + ( + str(meta["day"]), + str(meta["segment_key"]), + str(meta["source"]), + int(meta["sentence_id"]), + ) + ) + except (KeyError, TypeError, ValueError): + continue + + existing_norms = [ + normalized + for embedding in embeddings + if (normalized := normalize_embedding(embedding)) is not None + ] + centroid = ( + normalize_embedding(np.mean(existing_norms, axis=0)) if existing_norms else None + ) + return keys, len(embeddings), centroid + + +def _retroactive_owner_context() -> tuple[np.ndarray, float] | None: + from solstone.apps.speakers.attribution import load_owner_centroid + + centroid_data = load_owner_centroid() + if centroid_data is None: + return None + return centroid_data.centroid, centroid_data.threshold + + +def _is_principal_entity(entity_id: str) -> bool: + from solstone.think.entities import get_journal_principal + + principal = get_journal_principal() + return bool(principal and principal.get("id") == entity_id) + + +def _retroactive_segment_noisy(seg_dir: Path, source: str) -> bool: + from solstone.apps.speakers.attribution import ( + NOISY_FLYWHEEL_OVERLAP_MAX, + _read_segment_overlap_fraction, + ) + + jsonl_path = seg_dir / f"{source}.jsonl" + return _read_segment_overlap_fraction(jsonl_path) > NOISY_FLYWHEEL_OVERLAP_MAX + + +def _retroactive_outlier_min_samples() -> int: + from solstone.apps.speakers.attribution import VP_OUTLIER_MIN_SAMPLES + + return VP_OUTLIER_MIN_SAMPLES + + +def _retroactive_outlier_min_similarity() -> float: + from solstone.apps.speakers.attribution import VP_OUTLIER_MIN_SIMILARITY + + return VP_OUTLIER_MIN_SIMILARITY + + +def _retroactive_voiceprint_metadata( + day: str, + stream: str, + segment_key: str, + source: str, + sentence_id: int, +) -> dict[str, Any]: + return { + "day": day, + "segment_key": segment_key, + "source": source, + "stream": stream, + "sentence_id": sentence_id, + "added_at": now_ms(), + "last_seen_ts": segment_start_ts_ms(day, segment_key), + } + + def _source_key(source_segment: dict[str, Any]) -> tuple[str, str, str, str, int]: return ( str(source_segment["day"]), @@ -659,20 +788,91 @@ class CandidateTracker: candidate.confirmed_entity = None self.save() - def retroactive_confirm(self, centroid: np.ndarray, entity_id: str) -> int: - from solstone.apps.speakers.attribution import accumulate_voiceprints + def restore_confirmed_candidate( + self, + cand_id: int, + *, + expected_after: dict[str, Any], + candidate_before: dict[str, Any], + ) -> dict[str, Any]: + """Compare-restore a candidate confirmed by identify undo.""" + report = { + "restored_count": 0, + "skipped_count": 0, + "skipped_reasons": { + "missing": 0, + "already_restored": 0, + "concurrent_change": 0, + }, + } + with hold_lock(self.store_path): + self.load() + candidate = self._candidates.get(int(cand_id)) + if candidate is None: + report["skipped_count"] += 1 + report["skipped_reasons"]["missing"] += 1 + return report + current = candidate.to_json() + if current == candidate_before: + report["skipped_count"] += 1 + report["skipped_reasons"]["already_restored"] += 1 + return report + if current != expected_after: + report["skipped_count"] += 1 + report["skipped_reasons"]["concurrent_change"] += 1 + return report + candidate.status = str(candidate_before.get("status", "pending")) + confirmed_entity = candidate_before.get("confirmed_entity") + candidate.confirmed_entity = ( + str(confirmed_entity) if confirmed_entity is not None else None + ) + self._write() + report["restored_count"] = 1 + return report + def plan_retroactive_confirm( + self, + centroid: np.ndarray, + entity_id: str, + ) -> RetroactiveConfirmPlan: + """Plan retroactive confirmation without mutating tracker or voiceprints.""" load_embeddings_file, normalize_embedding = _routes_helpers() normalized_centroid = normalize_embedding(centroid) if normalized_centroid is None: - return 0 + return _empty_retroactive_plan(entity_id) cand_id, score = self._best_match(normalized_centroid) if cand_id is None or score < MERGE_THRESHOLD: - return 0 + return _empty_retroactive_plan(entity_id, match_score=score) candidate = self._candidates[cand_id] - saved_total = 0 + candidate_before = candidate.to_json() + candidate_after = dict(candidate_before) + candidate_after["status"] = "confirmed" + candidate_after["confirmed_entity"] = entity_id + + preexisting_keys, existing_count, existing_centroid = ( + _entity_voiceprint_snapshot(entity_id, normalize_embedding) + ) + working_keys = set(preexisting_keys) + voiceprint_items: list[tuple[np.ndarray, dict[str, Any]]] = [] + voiceprints_to_add: list[dict[str, Any]] = [] + + owner_context = _retroactive_owner_context() + if owner_context is None or _is_principal_entity(entity_id): + return RetroactiveConfirmPlan( + matched=True, + match_score=score, + candidate_id=cand_id, + entity_id=entity_id, + candidate_before=candidate_before, + candidate_after=candidate_after, + preexisting_voiceprint_keys=tuple(sorted(preexisting_keys)), + voiceprints_to_add=(), + voiceprint_items_to_add=(), + ) + owner_centroid, owner_threshold = owner_context + for source_segment in candidate.source_segments: day = str(source_segment["day"]) segment_key = str(source_segment["segment_key"]) @@ -683,11 +883,14 @@ class CandidateTracker: seg_dir = segment_path(day, segment_key, stream, create=False) if not seg_dir.exists(): continue + if _retroactive_segment_noisy(seg_dir, source): + continue + emb_data = load_embeddings_file(seg_dir / f"{source}.npz") + if emb_data is None: + continue + embeddings, statement_ids, _durations_s = emb_data + sid_to_idx = {int(sid): index for index, sid in enumerate(statement_ids)} if cluster_label == SOLO_CLUSTER_LABEL: - emb_data = load_embeddings_file(seg_dir / f"{source}.npz") - if emb_data is None: - continue - _embeddings, statement_ids, _durations_s = emb_data sentence_ids = sorted(int(sid) for sid in statement_ids) else: integer_labels = _load_integer_speaker_labels(seg_dir, source) @@ -696,27 +899,101 @@ class CandidateTracker: for sid, label in sorted(integer_labels.items()) if int(label) == cluster_label ] - synthetic_labels = [ - { - "sentence_id": sid, - "speaker": entity_id, - "confidence": "high", - "method": "acoustic_cluster", - } - for sid in sentence_ids - ] - if not synthetic_labels: - continue - saved = accumulate_voiceprints( - day, - stream, - segment_key, - synthetic_labels, - source, + + for sid in sentence_ids: + idx = sid_to_idx.get(int(sid)) + if idx is None: + continue + normalized = normalize_embedding(embeddings[idx]) + if normalized is None: + continue + if float(np.dot(normalized, owner_centroid)) >= owner_threshold: + continue + key = (day, segment_key, source, int(sid)) + if key in working_keys: + continue + if ( + existing_count >= _retroactive_outlier_min_samples() + and existing_centroid is not None + and float(np.dot(normalized, existing_centroid)) + < _retroactive_outlier_min_similarity() + ): + continue + metadata = _retroactive_voiceprint_metadata( + day, + stream, + segment_key, + source, + int(sid), + ) + voiceprint_items.append((normalized, metadata)) + voiceprints_to_add.append( + { + "key": { + "day": day, + "segment_key": segment_key, + "source": source, + "sentence_id": int(sid), + }, + "metadata": metadata, + } + ) + working_keys.add(key) + + return RetroactiveConfirmPlan( + matched=True, + match_score=score, + candidate_id=cand_id, + entity_id=entity_id, + candidate_before=candidate_before, + candidate_after=candidate_after, + preexisting_voiceprint_keys=tuple(sorted(preexisting_keys)), + voiceprints_to_add=tuple(voiceprints_to_add), + voiceprint_items_to_add=tuple(voiceprint_items), + ) + + def apply_retroactive_confirm_plan(self, plan: RetroactiveConfirmPlan) -> int: + """Apply a retroactive confirmation plan.""" + if not plan.matched or plan.candidate_id is None: + return 0 + + saved_total = 0 + if plan.voiceprint_items_to_add: + from solstone.think.entities import ( + load_existing_voiceprint_keys, + save_voiceprints_batch, ) - saved_total += sum(saved.values()) + existing_keys = load_existing_voiceprint_keys(plan.entity_id) + items = [ + (embedding, metadata) + for embedding, metadata in plan.voiceprint_items_to_add + if ( + metadata.get("day"), + metadata.get("segment_key"), + metadata.get("source"), + metadata.get("sentence_id"), + ) + not in existing_keys + ] + if items: + try: + saved_total = save_voiceprints_batch(plan.entity_id, items) + except Exception as exc: + logger.warning( + "Failed to apply retroactive voiceprints for %s: %s", + plan.entity_id, + exc, + ) + + candidate = self._candidates.get(int(plan.candidate_id)) + if candidate is None: + return saved_total candidate.status = "confirmed" - candidate.confirmed_entity = entity_id + candidate.confirmed_entity = plan.entity_id self.save() return saved_total + + def retroactive_confirm(self, centroid: np.ndarray, entity_id: str) -> int: + plan = self.plan_retroactive_confirm(centroid, entity_id) + return self.apply_retroactive_confirm_plan(plan) diff --git a/solstone/apps/speakers/tests/test_attribution.py b/solstone/apps/speakers/tests/test_attribution.py index ef7b6cbc5..d22edd625 100644 --- a/solstone/apps/speakers/tests/test_attribution.py +++ b/solstone/apps/speakers/tests/test_attribution.py @@ -1937,6 +1937,327 @@ def test_owner_correction_output_preserves_public_label_overlay(tmp_path): assert data["labels"][0]["method"] == "user_corrected" +def test_restore_label_rows_restores_present_prior_label(tmp_path): + from solstone.apps.speakers.attribution import restore_label_rows + + talents_dir = tmp_path / "talents" + talents_dir.mkdir() + labels_path = talents_dir / "speaker_labels.json" + expected = { + "sentence_id": 1, + "speaker": "target", + "confidence": "high", + "method": "user_identified", + } + prior = { + "sentence_id": 1, + "speaker": "prior", + "confidence": "high", + "method": "acoustic", + } + labels_path.write_text(json.dumps({"labels": [expected]}), encoding="utf-8") + + report = restore_label_rows( + tmp_path, + [ + { + "sentence_id": 1, + "expected_current_label": expected, + "prior_state": "present", + "prior_label": prior, + } + ], + ) + + data = json.loads(labels_path.read_text(encoding="utf-8")) + assert report["restored_count"] == 1 + assert report["patched_existing_count"] == 1 + assert data["labels"] == [prior] + + +def test_restore_label_rows_removes_absent_prior_row(tmp_path): + from solstone.apps.speakers.attribution import restore_label_rows + + talents_dir = tmp_path / "talents" + talents_dir.mkdir() + labels_path = talents_dir / "speaker_labels.json" + inserted = { + "sentence_id": 1, + "speaker": "target", + "confidence": "high", + "method": "user_identified", + } + labels_path.write_text(json.dumps({"labels": [inserted]}), encoding="utf-8") + + report = restore_label_rows( + tmp_path, + [ + { + "sentence_id": 1, + "expected_current_label": inserted, + "prior_state": "absent", + "prior_label": None, + } + ], + ) + + data = json.loads(labels_path.read_text(encoding="utf-8")) + assert report["restored_count"] == 1 + assert report["removed_inserted_count"] == 1 + assert data["labels"] == [] + + +def test_restore_label_rows_skips_changed_current_label(tmp_path): + from solstone.apps.speakers.attribution import restore_label_rows + + talents_dir = tmp_path / "talents" + talents_dir.mkdir() + labels_path = talents_dir / "speaker_labels.json" + expected = { + "sentence_id": 1, + "speaker": "target", + "confidence": "high", + "method": "user_identified", + } + changed = {**expected, "speaker": "other"} + prior = {**expected, "speaker": "prior", "method": "acoustic"} + labels_path.write_text(json.dumps({"labels": [changed]}), encoding="utf-8") + + report = restore_label_rows( + tmp_path, + [ + { + "sentence_id": 1, + "expected_current_label": expected, + "prior_state": "present", + "prior_label": prior, + } + ], + ) + + data = json.loads(labels_path.read_text(encoding="utf-8")) + assert report["restored_count"] == 0 + assert report["skipped_count"] == 1 + assert report["skipped_reasons"] == {"missing": 0, "changed": 1} + assert data["labels"] == [changed] + + +def test_restore_label_rows_rerun_is_noop_skip(tmp_path): + from solstone.apps.speakers.attribution import restore_label_rows + + talents_dir = tmp_path / "talents" + talents_dir.mkdir() + labels_path = talents_dir / "speaker_labels.json" + inserted = { + "sentence_id": 1, + "speaker": "target", + "confidence": "high", + "method": "user_identified", + } + restoration = { + "sentence_id": 1, + "expected_current_label": inserted, + "prior_state": "absent", + "prior_label": None, + } + labels_path.write_text(json.dumps({"labels": [inserted]}), encoding="utf-8") + + first = restore_label_rows(tmp_path, [restoration]) + second = restore_label_rows(tmp_path, [restoration]) + + assert first["restored_count"] == 1 + assert second["restored_count"] == 0 + assert second["skipped_reasons"] == {"missing": 1, "changed": 0} + assert json.loads(labels_path.read_text(encoding="utf-8"))["labels"] == [] + + +def test_identify_tagged_correction_overlays_speaker(tmp_path): + from solstone.apps.speakers.attribution import save_speaker_labels + + talents_dir = tmp_path / "talents" + talents_dir.mkdir() + (talents_dir / "speaker_corrections.json").write_text( + json.dumps( + { + "corrections": [ + { + "sentence_id": 1, + "original_speaker": "alice", + "corrected_speaker": "target", + "original_method": "acoustic", + "timestamp": 1, + "operation_id": "idop_1", + "correction_kind": "identify", + } + ] + } + ), + encoding="utf-8", + ) + + save_speaker_labels( + tmp_path, + [ + { + "sentence_id": 1, + "speaker": "alice", + "confidence": "high", + "method": "acoustic", + } + ], + {}, + ) + + data = json.loads((talents_dir / "speaker_labels.json").read_text()) + assert data["labels"][0]["speaker"] == "target" + assert data["labels"][0]["method"] == "user_corrected" + + +def test_identify_undo_correction_reverts_to_prior_speaker(tmp_path): + from solstone.apps.speakers.attribution import save_speaker_labels + + talents_dir = tmp_path / "talents" + talents_dir.mkdir() + (talents_dir / "speaker_corrections.json").write_text( + json.dumps( + { + "corrections": [ + { + "sentence_id": 1, + "original_speaker": "alice", + "corrected_speaker": "target", + "original_method": "acoustic", + "timestamp": 1, + "operation_id": "idop_1", + "correction_kind": "identify", + }, + { + "sentence_id": 1, + "original_speaker": "target", + "corrected_speaker": "prior", + "original_method": "user_identified", + "timestamp": 2, + "operation_id": "idop_1", + "undo_of_operation_id": "idop_1", + "correction_kind": "identify_undo", + }, + ] + } + ), + encoding="utf-8", + ) + + save_speaker_labels( + tmp_path, + [ + { + "sentence_id": 1, + "speaker": "alice", + "confidence": "high", + "method": "acoustic", + } + ], + {}, + ) + + data = json.loads((talents_dir / "speaker_labels.json").read_text()) + assert data["labels"][0]["speaker"] == "prior" + assert data["labels"][0]["method"] == "user_corrected" + + +def test_identify_undo_correction_with_null_reverts_to_unmatched(tmp_path): + from solstone.apps.speakers.attribution import save_speaker_labels + + talents_dir = tmp_path / "talents" + talents_dir.mkdir() + (talents_dir / "speaker_corrections.json").write_text( + json.dumps( + { + "corrections": [ + { + "sentence_id": 1, + "original_speaker": None, + "corrected_speaker": "target", + "original_method": None, + "timestamp": 1, + "operation_id": "idop_1", + "correction_kind": "identify", + }, + { + "sentence_id": 1, + "original_speaker": "target", + "corrected_speaker": None, + "original_method": "user_identified", + "timestamp": 2, + "operation_id": "idop_1", + "undo_of_operation_id": "idop_1", + "correction_kind": "identify_undo", + }, + ] + } + ), + encoding="utf-8", + ) + + save_speaker_labels( + tmp_path, + [ + { + "sentence_id": 1, + "speaker": "target", + "confidence": "high", + "method": "user_identified", + } + ], + {}, + ) + + data = json.loads((talents_dir / "speaker_labels.json").read_text()) + assert data["labels"][0]["speaker"] is None + assert data["labels"][0]["confidence"] is None + assert data["labels"][0]["method"] is None + + +def test_ordinary_null_correction_keeps_existing_behavior(tmp_path): + from solstone.apps.speakers.attribution import save_speaker_labels + + talents_dir = tmp_path / "talents" + talents_dir.mkdir() + (talents_dir / "speaker_corrections.json").write_text( + json.dumps( + { + "corrections": [ + { + "sentence_id": 1, + "original_speaker": "alice", + "corrected_speaker": None, + "original_method": "acoustic", + "timestamp": 1, + } + ] + } + ), + encoding="utf-8", + ) + + save_speaker_labels( + tmp_path, + [ + { + "sentence_id": 1, + "speaker": "alice", + "confidence": "high", + "method": "acoustic", + } + ], + {}, + ) + + data = json.loads((talents_dir / "speaker_labels.json").read_text()) + assert data["labels"][0]["speaker"] == "alice" + assert data["labels"][0]["method"] == "acoustic" + + # --------------------------------------------------------------------------- # Voiceprint accumulation # --------------------------------------------------------------------------- diff --git a/solstone/apps/speakers/tests/test_candidate_tracker.py b/solstone/apps/speakers/tests/test_candidate_tracker.py index beac32ae7..324cd0a49 100644 --- a/solstone/apps/speakers/tests/test_candidate_tracker.py +++ b/solstone/apps/speakers/tests/test_candidate_tracker.py @@ -34,6 +34,7 @@ from solstone.apps.speakers.encoder_config import ( STABILITY_THRESHOLD, ) from solstone.apps.speakers.owner import OWNER_THRESHOLD +from solstone.apps.speakers.tests.conftest import journal_tree_hash from solstone.think.entities import save_voiceprints_batch STREAM = "test" @@ -126,6 +127,11 @@ def _voiceprint_count(entity_dir: Path) -> int: return len(data["embeddings"]) +def _voiceprint_entries(entity_dir: Path) -> list[dict[str, object]]: + with np.load(entity_dir / "voiceprints.npz", allow_pickle=False) as data: + return [json.loads(str(item)) for item in data["metadata"]] + + def _only_candidate(tracker: CandidateTracker): assert len(tracker._candidates) == 1 return next(iter(tracker._candidates.values())) @@ -1034,6 +1040,60 @@ def test_retroactive_confirm_backfills_with_accumulate_guard(speakers_env, tmp_p assert candidate.confirmed_entity == "alice_test" +def test_plan_retroactive_confirm_writes_nothing(speakers_env, tmp_path): + env = speakers_env() + _setup_owner(env) + env.create_entity("Alice Test") + base = _unit([0.0, 1.0]) + seg_dir = _write_labeled_segment( + env, + "20260101", + "090000_300", + {7: np.stack([base] * 3)}, + ) + tracker = CandidateTracker(tmp_path / "speaker_candidates.json") + tracker.process_segment("20260101", "090000_300", STREAM, "mic_audio", seg_dir) + before = journal_tree_hash(env.journal) + + plan = tracker.plan_retroactive_confirm(base, "alice_test") + + assert plan.matched is True + assert plan.candidate_id == _only_candidate(tracker).cand_id + assert len(plan.voiceprints_to_add) == 3 + assert journal_tree_hash(env.journal) == before + candidate = _only_candidate(tracker) + assert candidate.status == "pending" + assert candidate.confirmed_entity is None + + +def test_plan_retroactive_confirm_keys_match_applied_voiceprints( + speakers_env, + tmp_path, +): + env = speakers_env() + _setup_owner(env) + alice_dir = env.create_entity("Alice Test") + base = _unit([0.0, 1.0]) + seg_dir = _write_labeled_segment( + env, + "20260101", + "090000_300", + {7: np.stack([base] * 3)}, + ) + tracker = CandidateTracker(tmp_path / "speaker_candidates.json") + tracker.process_segment("20260101", "090000_300", STREAM, "mic_audio", seg_dir) + + plan = tracker.plan_retroactive_confirm(base, "alice_test") + saved = tracker.apply_retroactive_confirm_plan(plan) + + assert saved == 3 + planned_metadata = [entry["metadata"] for entry in plan.voiceprints_to_add] + assert _voiceprint_entries(alice_dir) == planned_metadata + candidate = _only_candidate(tracker) + assert candidate.status == "confirmed" + assert candidate.confirmed_entity == "alice_test" + + def test_owner_selection_and_expansion_use_consolidated_survivor(speakers_env): from solstone.apps.speakers.owner import ( _expand_owner_candidate, diff --git a/solstone/think/entities/__init__.py b/solstone/think/entities/__init__.py index e20caa99e..e43042415 100644 --- a/solstone/think/entities/__init__.py +++ b/solstone/think/entities/__init__.py @@ -75,6 +75,7 @@ from solstone.think.entities.history import ( from solstone.think.entities.journal import ( block_journal_entity, create_journal_entity, + delete_created_entity_if_unreferenced, delete_journal_entity, ensure_journal_entity_memory, get_journal_principal, @@ -157,9 +158,11 @@ from solstone.think.entities.saving import ( update_facet_entity_identity, ) from solstone.think.entities.voiceprints import ( + VoiceprintRemovalError, load_entity_voiceprints_file, load_existing_voiceprint_keys, normalize_embedding, + remove_voiceprints_by_key, save_voiceprints_batch, voiceprint_file_path, ) @@ -183,6 +186,7 @@ __all__ = [ "EntityNotFoundError", "EntityWriteError", "EntityAmbiguityError", + "VoiceprintRemovalError", # History "EntityHistoryError", "EntityHistoryRepairRequired", @@ -196,6 +200,7 @@ __all__ = [ # Journal "block_journal_entity", "create_journal_entity", + "delete_created_entity_if_unreferenced", "delete_journal_entity", "ensure_journal_entity_memory", "get_journal_principal", @@ -233,6 +238,7 @@ __all__ = [ "detach_facet_entity", "save_detected_entity", "save_entities", + "remove_voiceprints_by_key", "save_voiceprints_batch", "update_facet_entity_description", "update_facet_entity_identity", diff --git a/solstone/think/entities/history.py b/solstone/think/entities/history.py index 71fb04480..64bb011f4 100644 --- a/solstone/think/entities/history.py +++ b/solstone/think/entities/history.py @@ -6,17 +6,22 @@ Entity trust operations use one per-journal, process-reentrant lock. The fixed global lock order is: - trust -> facet attached-store -> (facet relationship / observations / - voiceprints npz / activities locked_modify / speaker segment) owner locks - -> entity ambiguity store + trust -> speaker identify-ledger -> facet attached-store -> + (facet relationship / observations / voiceprints npz / activities + locked_modify / speaker segment / speaker candidate tracker / + speaker dismissal store / speaker keep-separate store) owner locks -> + entity ambiguity store -The ambiguity-store lock is never acquired before the trust or owner locks. +Mutating speaker identify operations acquire trust -> identify-ledger -> owner +locks (voiceprints / attribution / tracker / sentinel) -> new speaker stores, +never inverted. The ambiguity-store lock is never acquired before the trust or +owner locks. Preview and read-only paths must not acquire these locks at all. The trust lock is backed by ``journal_io.hold_lock`` at ``journal/health/locks/entity-trust``. ``hold_lock`` creates a persistent ``.lock`` sidecar, so journal-tree hash checks that assert byte-neutral preview/read behavior must exclude lock sidecars. Preview and read-only paths -must not acquire this lock at all. +must not acquire the trust lock at all. History layout: diff --git a/solstone/think/entities/journal.py b/solstone/think/entities/journal.py index 8f6e898fc..78d9b73d1 100644 --- a/solstone/think/entities/journal.py +++ b/solstone/think/entities/journal.py @@ -12,12 +12,14 @@ Facet-specific data (description, timestamps) is stored in facet relationships. import json import shutil +from collections import defaultdict from pathlib import Path from typing import Any from solstone.think.entities.core import EntityDict, get_identity_names from solstone.think.entities.history import ( EntityOperationContext, + iter_entity_history, save_entity_identity_with_history, trust_operation_lock, ) @@ -173,6 +175,7 @@ def create_journal_entity( aka: list[str] | None = None, emails: list[str] | None = None, *, + operation: EntityOperationContext | None = None, skip_principal: bool = False, ) -> EntityDict: """Create and persist a new journal-level entity. @@ -212,7 +215,7 @@ def create_journal_entity( ): entity["is_principal"] = True - save_journal_entity(entity) + save_journal_entity(entity, operation=operation) return entity @@ -353,6 +356,420 @@ def delete_journal_entity(entity_id: str) -> dict[str, Any]: return {"success": True, "facets_deleted": facets_deleted} +def delete_created_entity_if_unreferenced( + entity_id: str, + *, + operation_id: str, + expected_identity: dict[str, Any], + expected_history_refs: list[dict[str, Any]], +) -> dict[str, Any]: + """Delete an identify-created entity only when no outside references remain.""" + with trust_operation_lock(): + blocked: dict[str, int] = defaultdict(int) + current = load_journal_entity(entity_id) + if current is None: + return { + "deleted": True, + "blocked_categories": [], + "blocked_counts": {}, + } + if _meaningful_identity(current) != _meaningful_identity(expected_identity): + blocked["concurrent_change"] += 1 + + _check_expected_history( + entity_id, + operation_id=operation_id, + expected_history_refs=expected_history_refs, + blocked=blocked, + ) + _check_entity_dir_contents(entity_id, blocked) + _scan_entity_reference_surfaces( + entity_id, + operation_id=operation_id, + blocked=blocked, + ) + + if blocked: + blocked_counts = dict(sorted(blocked.items())) + return { + "deleted": False, + "blocked_categories": sorted(blocked_counts), + "blocked_counts": blocked_counts, + } + + delete_journal_entity(entity_id) + return { + "deleted": True, + "blocked_categories": [], + "blocked_counts": {}, + } + + +def _meaningful_identity(entity: dict[str, Any]) -> dict[str, Any]: + fields = ("id", "name", "type", "aka", "emails", "is_principal", "blocked") + return {field: entity.get(field) for field in fields if field in entity} + + +def _check_expected_history( + entity_id: str, + *, + operation_id: str, + expected_history_refs: list[dict[str, Any]], + blocked: dict[str, int], +) -> None: + try: + events = list(iter_entity_history(entity_id)) + except Exception: + blocked["unreadable"] += 1 + return + expected_refs = { + (str(ref.get("version_id")), int(ref.get("seq"))) + for ref in expected_history_refs + if ref.get("version_id") is not None and ref.get("seq") is not None + } + current_refs = { + (str(event.get("version_id")), int(event.get("seq"))) + for event in events + if event.get("version_id") is not None and event.get("seq") is not None + } + if len(events) != 1 or current_refs != expected_refs: + blocked["concurrent_change"] += 1 + return + event = events[0] + operation = event.get("operation") + if not isinstance(operation, dict): + blocked["concurrent_change"] += 1 + return + if operation.get("operation_kind") != "speaker_identify": + blocked["concurrent_change"] += 1 + if operation.get("operation_id") != operation_id: + blocked["concurrent_change"] += 1 + + +def _check_entity_dir_contents(entity_id: str, blocked: dict[str, int]) -> None: + entity_dir = Path(get_journal()) / "entities" / entity_id + if not entity_dir.exists(): + return + allowed_files = {entity_dir / "entity.json"} + history_events = entity_dir / "history" / "events" + if history_events.is_dir(): + allowed_files.update(history_events.glob("*.json")) + for path in entity_dir.rglob("*"): + if path.is_dir() or path.name.endswith(".lock"): + continue + if path not in allowed_files: + blocked["unrecognized_file"] += 1 + + +def _scan_entity_reference_surfaces( + entity_id: str, + *, + operation_id: str, + blocked: dict[str, int], +) -> None: + root = Path(get_journal()) + _scan_facet_relationship_refs(root, entity_id, blocked) + _scan_observation_refs(root, entity_id, blocked) + _scan_activity_refs(root, entity_id, blocked) + _scan_segment_speaker_refs(root, entity_id, operation_id, blocked) + _scan_aka_crossrefs(root, entity_id, blocked) + _scan_edge_refs(root, entity_id, blocked) + _scan_jsonl_refs( + root / "entities" / "ambiguities.jsonl", + "ambiguity", + entity_id, + blocked, + predicate=_ambiguity_refs_entity, + ) + _scan_jsonl_refs( + root / "entities" / "review-candidates.jsonl", + "entity_review_candidate", + entity_id, + blocked, + predicate=lambda row, eid: ( + row.get("source_slug") == eid or row.get("target_slug") == eid + ), + ) + _scan_jsonl_refs( + root / "speakers" / "review-candidates.jsonl", + "speaker_review_candidate", + entity_id, + blocked, + predicate=lambda row, eid: ( + row.get("source_id") == eid or row.get("target_id") == eid + ), + ) + _scan_jsonl_refs( + root / "speakers" / "candidate-pair-review-candidates.jsonl", + "candidate_pair", + entity_id, + blocked, + predicate=lambda row, eid: _json_value_present(row, eid), + ) + _scan_speaker_candidate_refs(root, entity_id, blocked) + _scan_keep_separate_refs(entity_id, blocked) + _scan_jsonl_refs( + root / "speakers" / "cluster-dismissals.jsonl", + "dismissal", + entity_id, + blocked, + predicate=lambda row, eid: _json_value_present(row, eid), + ) + _scan_identify_operation_refs(entity_id, operation_id, blocked) + + +def _scan_facet_relationship_refs( + root: Path, + entity_id: str, + blocked: dict[str, int], +) -> None: + facets_dir = root / "facets" + if not facets_dir.is_dir(): + return + for rel_dir in facets_dir.glob(f"*/entities/{entity_id}"): + if rel_dir.exists(): + blocked["facet_relationship"] += 1 + + +def _scan_observation_refs( + root: Path, + entity_id: str, + blocked: dict[str, int], +) -> None: + for path in sorted((root / "facets").glob("*/entities/*/observations.jsonl")): + if path.parent.name == entity_id: + blocked["observation"] += 1 + _scan_jsonl_refs( + path, + "observation", + entity_id, + blocked, + predicate=lambda row, eid: _json_key_value_present( + row, + eid, + {"entity_id", "target_entity_id", "source_entity_id"}, + ), + ) + + +def _scan_activity_refs(root: Path, entity_id: str, blocked: dict[str, int]) -> None: + keys = { + "entity_id", + "active_entities", + "owner_entity_id", + "counterparty_entity_id", + "from_entity_id", + "to_entity_id", + } + for path in sorted((root / "facets").glob("*/activities/*.jsonl")): + _scan_jsonl_refs( + path, + "activity", + entity_id, + blocked, + predicate=lambda row, eid: _json_key_value_present(row, eid, keys), + ) + + +def _scan_segment_speaker_refs( + root: Path, + entity_id: str, + operation_id: str, + blocked: dict[str, int], +) -> None: + chronicle = root / "chronicle" + for labels_path in sorted(chronicle.glob("*/*/*/talents/speaker_labels.json")): + data = _read_json_object(labels_path, blocked) + if data is None: + continue + labels = data.get("labels", []) + if isinstance(labels, list): + for label in labels: + if isinstance(label, dict) and label.get("speaker") == entity_id: + blocked["segment_label"] += 1 + for corr_path in sorted(chronicle.glob("*/*/*/talents/speaker_corrections.json")): + data = _read_json_object(corr_path, blocked) + if data is None: + continue + corrections = data.get("corrections", []) + if not isinstance(corrections, list): + continue + for row in corrections: + if not isinstance(row, dict): + continue + if row.get("operation_id") == operation_id: + continue + if ( + row.get("original_speaker") == entity_id + or row.get("corrected_speaker") == entity_id + ): + blocked["segment_correction"] += 1 + + +def _scan_aka_crossrefs(root: Path, entity_id: str, blocked: dict[str, int]) -> None: + for path in sorted((root / "entities").glob("*/entity.json")): + if path.parent.name == entity_id: + continue + data = _read_json_object(path, blocked) + if data is None: + continue + aka = data.get("aka") + if isinstance(aka, list) and entity_id in aka: + blocked["aka_crossref"] += 1 + + +def _scan_edge_refs(root: Path, entity_id: str, blocked: dict[str, int]) -> None: + if not (root / "indexer" / "journal.sqlite").is_file(): + return + try: + from solstone.think.indexer.edges import count_entity_edges + + count = count_entity_edges(entity_id) + except Exception: + blocked["unreadable"] += 1 + return + if count: + blocked["edge"] += int(count) + + +def _scan_speaker_candidate_refs( + root: Path, + entity_id: str, + blocked: dict[str, int], +) -> None: + data = _read_json_object(root / "awareness" / "speaker_candidates.json", blocked) + if data is None: + return + candidates = data.get("candidates", []) + if not isinstance(candidates, list): + return + for candidate in candidates: + if ( + isinstance(candidate, dict) + and candidate.get("confirmed_entity") == entity_id + ): + blocked["speaker_candidate"] += 1 + + +def _scan_keep_separate_refs(entity_id: str, blocked: dict[str, int]) -> None: + try: + from solstone.think.speaker_keep_separate import fold_assertions + + for assertion in fold_assertions(): + if entity_id in (assertion.entity_id_a, assertion.entity_id_b): + blocked["keep_separate"] += 1 + except Exception: + blocked["unreadable"] += 1 + + +def _scan_identify_operation_refs( + entity_id: str, + operation_id: str, + blocked: dict[str, int], +) -> None: + try: + from solstone.think.speaker_identify_operations import fold_all_operations + + states = fold_all_operations() + except Exception: + blocked["unreadable"] += 1 + return + for state in states: + if state.operation_id == operation_id: + continue + if state.target_entity_id == entity_id: + blocked["identify_operation"] += 1 + continue + if entity_id in state.reviewed_near_match_entity_ids: + blocked["identify_operation"] += 1 + continue + for assertion in state.prepared_plan.get("keep_separate_assertions", []): + if not isinstance(assertion, dict): + continue + if entity_id in ( + assertion.get("entity_id_a"), + assertion.get("entity_id_b"), + assertion.get("planned_target_entity_id"), + assertion.get("reviewed_id"), + ): + blocked["identify_operation"] += 1 + break + + +def _scan_jsonl_refs( + path: Path, + category: str, + entity_id: str, + blocked: dict[str, int], + *, + predicate, +) -> None: + if not path.is_file(): + return + try: + lines = path.read_text(encoding="utf-8").splitlines() + except OSError: + blocked["unreadable"] += 1 + return + for line in lines: + raw = line.strip() + if not raw: + continue + try: + row = json.loads(raw) + except json.JSONDecodeError: + blocked["unreadable"] += 1 + continue + if isinstance(row, dict) and predicate(row, entity_id): + blocked[category] += 1 + + +def _read_json_object(path: Path, blocked: dict[str, int]) -> dict[str, Any] | None: + if not path.is_file(): + return None + try: + data = json.loads(path.read_text(encoding="utf-8")) + except (json.JSONDecodeError, OSError): + blocked["unreadable"] += 1 + return None + return data if isinstance(data, dict) else None + + +def _ambiguity_refs_entity(row: dict[str, Any], entity_id: str) -> bool: + if row.get("resolved_entity_id") == entity_id: + return True + candidates = row.get("ranked_candidates") + return isinstance(candidates, list) and any( + isinstance(candidate, dict) and candidate.get("id") == entity_id + for candidate in candidates + ) + + +def _json_key_value_present(value: Any, entity_id: str, keys: set[str]) -> bool: + if isinstance(value, dict): + for key, child in value.items(): + if key in keys: + if child == entity_id: + return True + if isinstance(child, list) and entity_id in child: + return True + if _json_key_value_present(child, entity_id, keys): + return True + elif isinstance(value, list): + return any(_json_key_value_present(item, entity_id, keys) for item in value) + return False + + +def _json_value_present(value: Any, entity_id: str) -> bool: + if value == entity_id: + return True + if isinstance(value, dict): + return any(_json_value_present(item, entity_id) for item in value.values()) + if isinstance(value, list): + return any(_json_value_present(item, entity_id) for item in value) + return False + + def journal_entity_memory_path(entity_id: str) -> Path: """Return path to journal entity's memory folder. diff --git a/solstone/think/entities/voiceprints.py b/solstone/think/entities/voiceprints.py index 646f3ff2a..957deb9dc 100644 --- a/solstone/think/entities/voiceprints.py +++ b/solstone/think/entities/voiceprints.py @@ -24,6 +24,10 @@ logger = logging.getLogger(__name__) VOICEPRINT_KEYS = ("embeddings", "metadata") +class VoiceprintRemovalError(RuntimeError): + """Raised when a by-key voiceprint removal cannot be applied safely.""" + + def normalize_embedding(emb: np.ndarray) -> np.ndarray | None: """L2-normalize an embedding vector. Returns None if norm is zero.""" import numpy as np @@ -232,6 +236,108 @@ def apply_entity_merge_voiceprint_inverse( return removed +def remove_voiceprints_by_key( + entity_id: str, + removals: list[dict[str, Any]], +) -> dict[str, Any]: + """Remove voiceprint rows by exact key and metadata match. + + Missing rows and metadata mismatches are skipped. Duplicate exact matches + raise because the caller cannot safely attribute one removal to one row. + """ + report: dict[str, Any] = { + "removed_count": 0, + "skipped_count": 0, + "skipped_reasons": {"missing": 0, "metadata_mismatch": 0}, + "file_removed": False, + } + if not removals: + return report + + try: + folder = journal_entity_memory_path(entity_id) + except (RuntimeError, ValueError): + report["skipped_count"] = len(removals) + report["skipped_reasons"]["missing"] = len(removals) + return report + + npz_path = folder / "voiceprints.npz" + if not npz_path.exists(): + report["skipped_count"] = len(removals) + report["skipped_reasons"]["missing"] = len(removals) + return report + + normalized_removals = [ + (_removal_key(removal), removal.get("expected_metadata")) + for removal in removals + ] + + def transform(current: dict[str, np.ndarray]) -> dict[str, np.ndarray] | None: + import numpy as np + + embeddings = current.get("embeddings") + metadata_arr = current.get("metadata") + if embeddings is None or metadata_arr is None: + raise MalformedDataError(npz_path) + + metadata = [json.loads(str(item)) for item in metadata_arr] + remove_indexes: set[int] = set() + for key, expected_metadata in normalized_removals: + exact_matches = [ + index + for index, meta in enumerate(metadata) + if _voiceprint_key(meta) == key and meta == expected_metadata + ] + if len(exact_matches) > 1: + raise VoiceprintRemovalError( + "voiceprint removal locator matched multiple rows" + ) + if len(exact_matches) == 1: + remove_indexes.add(exact_matches[0]) + continue + + key_present = any(_voiceprint_key(meta) == key for meta in metadata) + if key_present: + report["skipped_reasons"]["metadata_mismatch"] += 1 + else: + report["skipped_reasons"]["missing"] += 1 + + if not remove_indexes: + report["skipped_count"] = sum(report["skipped_reasons"].values()) + return None + + keep_indexes = [ + index for index in range(len(metadata)) if index not in remove_indexes + ] + report["removed_count"] = len(remove_indexes) + report["skipped_count"] = sum(report["skipped_reasons"].values()) + if not keep_indexes: + report["file_removed"] = True + return {} + return { + "embeddings": embeddings[keep_indexes], + "metadata": np.asarray( + [json.dumps(metadata[index]) for index in keep_indexes], + dtype=str, + ), + } + + update_npz(npz_path, transform, expected_keys=VOICEPRINT_KEYS) + return report + + +def _removal_key(removal: dict[str, Any]) -> tuple[Any, Any, Any, Any]: + key = removal.get("key") + if not isinstance(key, dict): + raise VoiceprintRemovalError("voiceprint removal key must be an object") + return ( + key.get("day"), + key.get("segment_key"), + key.get("source"), + key.get("sentence_id"), + ) + + def _voiceprint_key(meta: dict[str, Any]) -> tuple[Any, Any, Any, Any]: return ( meta.get("day"), diff --git a/solstone/think/speaker_cluster_dismissals.py b/solstone/think/speaker_cluster_dismissals.py new file mode 100644 index 000000000..058a28dfb --- /dev/null +++ b/solstone/think/speaker_cluster_dismissals.py @@ -0,0 +1,324 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc +"""Append-only speaker discovery-cluster dismissal store. + +Sole write-owner of: + journal/speakers/cluster-dismissals.jsonl +""" + +from __future__ import annotations + +import hashlib +import json +import logging +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from solstone.think.journal_io import append_jsonl, hold_lock +from solstone.think.utils import get_journal + +logger = logging.getLogger(__name__) + +CLUSTER_DISMISSAL_SCHEMA_VERSION = 1 +DISPOSITIONS = {"not_a_person", "quiet"} + + +class ClusterDismissalStoreError(RuntimeError): + """Raised when cluster dismissal storage cannot be trusted.""" + + +@dataclass(frozen=True) +class FoldedClusterDismissal: + """Merged dismissal state folded from overlapping append-only events.""" + + dismissal_id: str + disposition: str + members: tuple[dict[str, Any], ...] + event_ids: tuple[str, ...] + created_at: str + updated_at: str + + @property + def member_count(self) -> int: + return len(self.members) + + @property + def event_count(self) -> int: + return len(self.event_ids) + + +def cluster_dismissals_dir() -> Path: + """Return the speaker dismissal directory, creating it if needed.""" + path = Path(get_journal()) / "speakers" + path.mkdir(parents=True, exist_ok=True) + return path + + +def cluster_dismissals_path() -> Path: + """Return the speaker cluster-dismissal JSONL event-log path.""" + return cluster_dismissals_dir() / "cluster-dismissals.jsonl" + + +def record_cluster_dismissal( + members: list[dict[str, Any]], + disposition: str, +) -> dict[str, Any]: + """Append one discovery-cluster dismissal event.""" + if disposition not in DISPOSITIONS: + raise ValueError(f"unknown cluster dismissal disposition: {disposition}") + canonical_members = _canonical_member_dicts(members) + if not canonical_members: + raise ValueError("cluster dismissal requires at least one member") + ts = utc_now_iso() + event = { + "schema_version": CLUSTER_DISMISSAL_SCHEMA_VERSION, + "event_kind": "dismiss", + "dismiss_event_id": _dismiss_event_id(canonical_members, ts, disposition), + "disposition": disposition, + "members": canonical_members, + "member_count": len(canonical_members), + "ts": ts, + } + append_event(event) + return event + + +def append_event(event: dict[str, Any]) -> None: + """Strict-validate and append one cluster dismissal event.""" + _validate_row(event) + path = cluster_dismissals_path() + with hold_lock(path): + append_jsonl(path, event) + + +def load_events() -> list[dict[str, Any]]: + """Strict-load raw cluster dismissal events.""" + return _load_jsonl_rows(cluster_dismissals_path()) + + +def fold_dismissals() -> list[FoldedClusterDismissal]: + """Fold dismissal events into connected overlap components.""" + events = load_events() + if not events: + return [] + + member_sets = [_member_set(event["members"]) for event in events] + adjacency: list[set[int]] = [set() for _ in events] + for left in range(len(events)): + for right in range(left + 1, len(events)): + if _overlap_ratio_min(member_sets[left], member_sets[right]) >= 0.50: + adjacency[left].add(right) + adjacency[right].add(left) + + folded: list[FoldedClusterDismissal] = [] + seen: set[int] = set() + for start in range(len(events)): + if start in seen: + continue + stack = [start] + component: list[int] = [] + seen.add(start) + while stack: + current = stack.pop() + component.append(current) + for neighbor in sorted(adjacency[current]): + if neighbor in seen: + continue + seen.add(neighbor) + stack.append(neighbor) + folded.append(_fold_component(events, member_sets, component)) + return sorted(folded, key=lambda row: row.dismissal_id) + + +def cluster_dismissal_suppressed(candidate_members: list[dict[str, Any]]) -> bool: + """Return whether a candidate cluster is suppressed by dismissed provenance.""" + candidate_set = _member_set(_canonical_member_dicts(candidate_members)) + if not candidate_set: + return False + for dismissal in fold_dismissals(): + dismissed_set = _member_set(list(dismissal.members)) + if len(candidate_set & dismissed_set) / len(candidate_set) >= 0.50: + return True + return False + + +def list_dismissals() -> list[dict[str, Any]]: + """Return redacted summaries for folded cluster dismissals.""" + return [ + { + "dismissal_id": dismissal.dismissal_id, + "disposition": dismissal.disposition, + "member_count": dismissal.member_count, + "event_count": dismissal.event_count, + "created_at": dismissal.created_at, + "updated_at": dismissal.updated_at, + } + for dismissal in fold_dismissals() + ] + + +def utc_now_iso() -> str: + """Return the current UTC time as an ISO-8601 string ending in Z.""" + return ( + datetime.now(timezone.utc).isoformat(timespec="seconds").replace("+00:00", "Z") + ) + + +def _load_jsonl_rows(path: Path) -> list[dict[str, Any]]: + if not path.exists(): + return [] + rows: list[dict[str, Any]] = [] + try: + with open(path, encoding="utf-8") as handle: + for lineno, line in enumerate(handle, start=1): + raw = line.strip() + if not raw: + continue + try: + row = json.loads(raw) + except json.JSONDecodeError as exc: + message = f"malformed cluster dismissal JSONL at {path}:{lineno}" + logger.error(message) + raise ClusterDismissalStoreError(message) from exc + if not isinstance(row, dict): + message = f"non-object cluster dismissal JSONL at {path}:{lineno}" + logger.error(message) + raise ClusterDismissalStoreError(message) + try: + _validate_row(row) + except ClusterDismissalStoreError as exc: + message = f"invalid cluster dismissal row at {path}:{lineno}: {exc}" + logger.error(message) + raise ClusterDismissalStoreError(message) from exc + rows.append(row) + except OSError as exc: + message = f"failed to read cluster dismissal store {path}: {exc}" + logger.error(message) + raise ClusterDismissalStoreError(message) from exc + return rows + + +def _fold_component( + events: list[dict[str, Any]], + member_sets: list[set[tuple[str, str, str, str, int]]], + component: list[int], +) -> FoldedClusterDismissal: + event_ids = tuple( + sorted(str(events[index]["dismiss_event_id"]) for index in component) + ) + union_members: set[tuple[str, str, str, str, int]] = set() + disposition = "quiet" + timestamps: list[str] = [] + for index in component: + union_members.update(member_sets[index]) + timestamps.append(str(events[index]["ts"])) + if events[index]["disposition"] == "not_a_person": + disposition = "not_a_person" + return FoldedClusterDismissal( + dismissal_id=_folded_dismissal_id(event_ids), + disposition=disposition, + members=tuple(_member_dict(member) for member in sorted(union_members)), + event_ids=event_ids, + created_at=min(timestamps), + updated_at=max(timestamps), + ) + + +def _validate_row(row: dict[str, Any]) -> None: + if row.get("schema_version") != CLUSTER_DISMISSAL_SCHEMA_VERSION: + raise ClusterDismissalStoreError("invalid schema_version") + if row.get("event_kind") != "dismiss": + raise ClusterDismissalStoreError("event_kind must be dismiss") + _required_str(row, "dismiss_event_id") + disposition = _required_str(row, "disposition") + if disposition not in DISPOSITIONS: + raise ClusterDismissalStoreError(f"invalid disposition: {disposition}") + _required_str(row, "ts") + members = row.get("members") + if not isinstance(members, list): + raise ClusterDismissalStoreError("members must be a list") + canonical_members = _canonical_member_dicts(members) + if canonical_members != members: + raise ClusterDismissalStoreError("members must be canonical sorted provenance") + if row.get("member_count") != len(members): + raise ClusterDismissalStoreError("member_count mismatch") + if not members: + raise ClusterDismissalStoreError("dismiss event requires members") + + +def _canonical_member_dicts(members: list[dict[str, Any]]) -> list[dict[str, Any]]: + return [ + _member_dict(item) for item in sorted({_member_tuple(item) for item in members}) + ] + + +def _member_set(members: list[dict[str, Any]]) -> set[tuple[str, str, str, str, int]]: + return {_member_tuple(member) for member in members} + + +def _member_tuple(member: dict[str, Any]) -> tuple[str, str, str, str, int]: + try: + return ( + str(member["day"]), + str(member["stream"]), + str(member["segment_key"]), + str(member["source"]), + int(member["sentence_id"]), + ) + except (KeyError, TypeError, ValueError) as exc: + raise ClusterDismissalStoreError("invalid cluster member provenance") from exc + + +def _member_dict( + member: tuple[str, str, str, str, int] | dict[str, Any], +) -> dict[str, Any]: + if isinstance(member, dict): + member = _member_tuple(member) + day, stream, segment_key, source, sentence_id = member + return { + "day": day, + "stream": stream, + "segment_key": segment_key, + "source": source, + "sentence_id": sentence_id, + } + + +def _overlap_ratio_min( + left: set[tuple[str, str, str, str, int]], + right: set[tuple[str, str, str, str, int]], +) -> float: + denominator = min(len(left), len(right)) + if denominator == 0: + return 0.0 + return len(left & right) / denominator + + +def _dismiss_event_id( + members: list[dict[str, Any]], + ts: str, + disposition: str, +) -> str: + payload = { + "disposition": disposition, + "members": members, + "ts": ts, + } + encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")) + digest = hashlib.sha256(encoded.encode("utf-8")).hexdigest() + return f"cdev_{digest[:24]}" + + +def _folded_dismissal_id(event_ids: tuple[str, ...]) -> str: + encoded = json.dumps(list(event_ids), sort_keys=True, separators=(",", ":")) + digest = hashlib.sha256(encoded.encode("utf-8")).hexdigest() + return f"cdsm_{digest[:24]}" + + +def _required_str(row: dict[str, Any], field: str) -> str: + value = row.get(field) + if not isinstance(value, str) or not value: + raise ClusterDismissalStoreError(f"missing or invalid {field}") + return value diff --git a/solstone/think/speaker_identify_operations.py b/solstone/think/speaker_identify_operations.py new file mode 100644 index 000000000..79ac64151 --- /dev/null +++ b/solstone/think/speaker_identify_operations.py @@ -0,0 +1,546 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc +"""Append-only speaker identify operation ledger. + +Sole write-owner of: + journal/speakers/identify-operations.jsonl +""" + +from __future__ import annotations + +import hashlib +import json +import logging +from dataclasses import dataclass +from dataclasses import field as dataclass_field +from pathlib import Path +from typing import Any + +from solstone.think.entities.history import trust_operation_lock +from solstone.think.journal_io import append_jsonl, hold_lock +from solstone.think.utils import get_journal + +logger = logging.getLogger(__name__) + +IDENTIFY_OPERATION_SCHEMA_VERSION = 1 + +EVENT_KINDS = { + "prepared", + "checkpoint", + "committed", + "repair_required", + "undo_prepared", + "undo_checkpoint", + "undo_committed", + "undo_repair_required", +} +FORWARD_PHASE_ORDER = ( + "entity", + "keep_separate", + "direct_voiceprints", + "corrections", + "labels", + "retro_tracker", + "sentinel", +) +UNDO_PHASE_ORDER = ( + "labels", + "corrections", + "voiceprints", + "tracker", + "sentinel", + "entity", +) + + +class IdentifyOperationError(RuntimeError): + """Raised when the identify operation ledger cannot be read or written.""" + + +@dataclass(frozen=True) +class OperationState: + """Folded state for one identify operation.""" + + operation_id: str + request_id: str + request_fingerprint: str + cluster_member_set: frozenset[tuple[str, str, str, str, int]] + target_entity_id: str | None + target_entity_name: str | None + will_create: bool + entity_type: str | None + reviewed_near_match_entity_ids: tuple[str, ...] + completed_phases: tuple[str, ...] + pending_phases: tuple[str, ...] + terminal_status: str + result: dict[str, Any] | None + undo_report: dict[str, Any] | None + phase_checkpoints: dict[str, dict[str, Any]] + prepared_plan: dict[str, Any] + repair_required: dict[str, Any] | None = None + undo_repair_required: dict[str, Any] | None = None + undo_phase_checkpoints: dict[str, dict[str, Any]] = dataclass_field( + default_factory=dict + ) + + +def identify_operations_dir() -> Path: + """Return the speaker operation-ledger directory, creating it if needed.""" + path = Path(get_journal()) / "speakers" + path.mkdir(parents=True, exist_ok=True) + return path + + +def identify_operations_path() -> Path: + """Return the speaker identify operation ledger path.""" + return identify_operations_dir() / "identify-operations.jsonl" + + +def operation_id_for_request(request_id: str) -> str: + """Return the deterministic public operation id for a caller request id.""" + if not isinstance(request_id, str) or not request_id: + raise ValueError("request_id must be a non-empty string") + digest = hashlib.sha256(request_id.encode("utf-8")).hexdigest() + return f"idop_{digest[:24]}" + + +def request_fingerprint( + *, + cluster_members: list[dict[str, Any]], + target_entity_id: str, + will_create: bool, + entity_type: str, + reviewed_near_match_entity_ids: list[str] | tuple[str, ...] | set[str], +) -> str: + """Hash the immutable identity of an identify request.""" + payload = { + "cluster_members": [ + list(member) + for member in sorted(_member_tuple(item) for item in cluster_members) + ], + "target_entity_id": str(target_entity_id), + "will_create": bool(will_create), + "entity_type": str(entity_type), + "reviewed_near_match_entity_ids": sorted( + {str(item) for item in reviewed_near_match_entity_ids} + ), + } + encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")) + return hashlib.sha256(encoded.encode("utf-8")).hexdigest() + + +def append_event(event: dict[str, Any]) -> None: + """Strict-validate and append one identify ledger event.""" + _validate_row(event) + path = identify_operations_path() + with trust_operation_lock(): + with hold_lock(path): + append_jsonl(path, event) + + +def load_operations() -> list[dict[str, Any]]: + """Strict-load raw identify ledger events.""" + return _load_jsonl_rows(identify_operations_path()) + + +def fold_operation(operation_id: str) -> OperationState | None: + """Fold events for one operation id into current operation state.""" + events = [ + event for event in load_operations() if event["operation_id"] == operation_id + ] + if not events: + return None + return _fold_events(events) + + +def fold_all_operations() -> list[OperationState]: + """Fold all operation ids in the identify ledger.""" + grouped: dict[str, list[dict[str, Any]]] = {} + for event in load_operations(): + grouped.setdefault(event["operation_id"], []).append(event) + return [ + _fold_events(events) + for _operation_id, events in sorted(grouped.items(), key=lambda item: item[0]) + ] + + +def _load_jsonl_rows(path: Path) -> list[dict[str, Any]]: + if not path.exists(): + return [] + + rows: list[dict[str, Any]] = [] + try: + with open(path, encoding="utf-8") as handle: + for lineno, line in enumerate(handle, start=1): + raw = line.strip() + if not raw: + continue + try: + row = json.loads(raw) + except json.JSONDecodeError as exc: + message = f"malformed identify operation JSONL at {path}:{lineno}" + logger.error(message) + raise IdentifyOperationError(message) from exc + if not isinstance(row, dict): + message = f"non-object identify operation JSONL at {path}:{lineno}" + logger.error(message) + raise IdentifyOperationError(message) + try: + _validate_row(row) + except IdentifyOperationError as exc: + message = ( + f"invalid identify operation row at {path}:{lineno}: {exc}" + ) + logger.error(message) + raise IdentifyOperationError(message) from exc + rows.append(row) + except OSError as exc: + message = f"failed to read identify operation ledger {path}: {exc}" + logger.error(message) + raise IdentifyOperationError(message) from exc + return rows + + +def _validate_row(row: dict[str, Any]) -> None: + if row.get("schema_version") != IDENTIFY_OPERATION_SCHEMA_VERSION: + raise IdentifyOperationError("invalid schema_version") + event_kind = _required_str(row, "event_kind") + if event_kind not in EVENT_KINDS: + raise IdentifyOperationError(f"unknown event_kind: {event_kind}") + + _required_str(row, "event_id") + _required_str(row, "operation_id") + _required_str(row, "request_id") + _required_str(row, "ts") + _required_str(row, "caller") + if "actor" not in row: + raise IdentifyOperationError("missing actor") + actor = row["actor"] + if actor is not None and not isinstance(actor, str): + raise IdentifyOperationError("actor must be a string or null") + + if event_kind == "prepared": + _validate_prepared(row) + elif event_kind == "checkpoint": + phase = _required_str(row, "phase") + if phase not in FORWARD_PHASE_ORDER: + raise IdentifyOperationError(f"invalid checkpoint phase: {phase}") + _validate_checkpoint(phase, _required_dict(row, "checkpoint")) + elif event_kind == "committed": + _required_dict(row, "result") + elif event_kind == "repair_required": + _validate_repair(row, undo=False) + elif event_kind == "undo_prepared": + _required_str(row, "undo_started_at") + elif event_kind == "undo_checkpoint": + phase = _required_str(row, "phase") + if phase not in UNDO_PHASE_ORDER: + raise IdentifyOperationError(f"invalid undo checkpoint phase: {phase}") + _required_dict(row, "undo_report_delta") + elif event_kind == "undo_committed": + _required_dict(row, "undo_report") + elif event_kind == "undo_repair_required": + _validate_repair(row, undo=True) + + +def _validate_prepared(row: dict[str, Any]) -> None: + fingerprint = _required_str(row, "request_fingerprint") + if len(fingerprint) != 64: + raise IdentifyOperationError("request_fingerprint must be a sha256 hex digest") + plan = _required_dict(row, "prepared_plan") + if plan.get("plan_schema_version") != 1: + raise IdentifyOperationError("prepared_plan.plan_schema_version must be 1") + if plan.get("operation_id") != row["operation_id"]: + raise IdentifyOperationError("prepared_plan operation_id mismatch") + if plan.get("request_id") != row["request_id"]: + raise IdentifyOperationError("prepared_plan request_id mismatch") + + for field in ( + "planned_at", + "request", + "cluster", + "target", + "entity_identity", + "direct_voiceprints", + "segments", + "retro_confirm", + "sentinel", + "keep_separate_assertions", + ): + if field not in plan: + raise IdentifyOperationError(f"prepared_plan missing {field}") + + request = _required_dict(plan, "request") + for field in ( + "cluster_id", + "name", + "entity_id", + "resolve_only", + "create_new", + "entity_type", + "reviewed_near_match_entity_ids", + ): + if field not in request: + raise IdentifyOperationError(f"prepared_plan.request missing {field}") + + cluster = _required_dict(plan, "cluster") + members = _required_list(cluster, "members") + if int(cluster.get("member_count", -1)) != len(members): + raise IdentifyOperationError("prepared_plan cluster member_count mismatch") + for member in members: + if not isinstance(member, dict): + raise IdentifyOperationError("prepared_plan cluster member is not object") + _member_tuple(member) + + target = _required_dict(plan, "target") + _required_str(target, "entity_id") + _required_str(target, "entity_name") + if not isinstance(target.get("will_create"), bool): + raise IdentifyOperationError("prepared_plan target.will_create must be bool") + + _required_dict(plan, "entity_identity") + _required_dict(plan, "direct_voiceprints") + _required_list(plan, "segments") + _required_dict(plan, "retro_confirm") + _required_dict(plan, "sentinel") + _required_list(plan, "keep_separate_assertions") + + +def _validate_checkpoint(phase: str, checkpoint: dict[str, Any]) -> None: + if checkpoint.get("phase_status") != "complete": + raise IdentifyOperationError("checkpoint.phase_status must be complete") + _required_str(checkpoint, "completed_at") + _required_dict(checkpoint, "counts") + _required_dict(checkpoint, "skipped_reasons") + if phase == "entity": + _required_str(checkpoint, "entity_id") + _required_bool(checkpoint, "entity_created") + _required_str(checkpoint, "identity_after_hash") + _required_list(checkpoint, "history_event_refs") + elif phase == "keep_separate": + _required_list(checkpoint, "pair_keys") + _required_int(checkpoint, "recorded_count") + _required_int(checkpoint, "already_present_count") + elif phase == "direct_voiceprints": + _required_list(checkpoint, "saved_keys") + _required_int(checkpoint, "saved_count") + _required_int(checkpoint, "skipped_existing_count") + elif phase == "corrections": + _required_list(checkpoint, "appended_keys") + _required_int(checkpoint, "appended_count") + _required_int(checkpoint, "skipped_existing_count") + _required_int(checkpoint, "segment_count") + elif phase == "labels": + _required_list(checkpoint, "patched_sentence_keys") + _required_list(checkpoint, "inserted_sentence_keys") + _required_int(checkpoint, "patched_count") + _required_int(checkpoint, "inserted_count") + _required_int(checkpoint, "skipped_already_intended_count") + _required_int(checkpoint, "segment_count") + elif phase == "retro_tracker": + _required_bool(checkpoint, "matched") + if checkpoint.get("candidate_id") is not None and not isinstance( + checkpoint.get("candidate_id"), int + ): + raise IdentifyOperationError("retro checkpoint candidate_id invalid") + _required_list(checkpoint, "saved_keys") + _required_int(checkpoint, "voiceprints_saved_count") + _required_int(checkpoint, "voiceprints_skipped_existing_count") + _required_bool(checkpoint, "tracker_updated") + elif phase == "sentinel": + _required_str(checkpoint, "cluster_key") + _required_bool(checkpoint, "written") + + +def _validate_repair(row: dict[str, Any], *, undo: bool) -> None: + phase = _required_str(row, "phase") + valid_phases = UNDO_PHASE_ORDER if undo else FORWARD_PHASE_ORDER + if phase not in valid_phases: + raise IdentifyOperationError(f"invalid repair phase: {phase}") + _required_str(row, "repair_code") + _required_dict(row, "repair_categories") + _required_dict(row, "undo_report" if undo else "partial_report") + + +def _fold_events(events: list[dict[str, Any]]) -> OperationState: + deduped = _dedupe_events(events) + prepared_events = [event for event in deduped if event["event_kind"] == "prepared"] + if len(prepared_events) != 1: + raise IdentifyOperationError("operation must have exactly one prepared event") + prepared = prepared_events[0] + plan = prepared["prepared_plan"] + + phase_checkpoints: dict[str, dict[str, Any]] = {} + for event in deduped: + if event["event_kind"] != "checkpoint": + continue + phase = event["phase"] + if ( + phase in phase_checkpoints + and phase_checkpoints[phase] != event["checkpoint"] + ): + raise IdentifyOperationError(f"conflicting checkpoint for phase {phase}") + phase_checkpoints[phase] = event["checkpoint"] + + undo_phase_checkpoints: dict[str, dict[str, Any]] = {} + for event in deduped: + if event["event_kind"] != "undo_checkpoint": + continue + phase = event["phase"] + if ( + phase in undo_phase_checkpoints + and undo_phase_checkpoints[phase] != event["undo_report_delta"] + ): + raise IdentifyOperationError( + f"conflicting undo checkpoint for phase {phase}" + ) + undo_phase_checkpoints[phase] = event["undo_report_delta"] + + completed = tuple( + phase for phase in FORWARD_PHASE_ORDER if phase in phase_checkpoints + ) + terminal = _terminal_status(deduped) + pending = _pending_phases(terminal, completed, deduped) + request = plan["request"] + target = plan["target"] + entity_type = ( + str(request.get("entity_type")) if request.get("entity_type") else None + ) + return OperationState( + operation_id=prepared["operation_id"], + request_id=prepared["request_id"], + request_fingerprint=prepared["request_fingerprint"], + cluster_member_set=frozenset( + _member_tuple(member) for member in plan["cluster"]["members"] + ), + target_entity_id=target.get("entity_id"), + target_entity_name=target.get("entity_name"), + will_create=bool(target.get("will_create")), + entity_type=entity_type, + reviewed_near_match_entity_ids=tuple( + str(item) for item in request.get("reviewed_near_match_entity_ids", []) + ), + completed_phases=completed, + pending_phases=pending, + terminal_status=terminal, + result=_last_payload(deduped, "committed", "result"), + undo_report=_last_payload(deduped, "undo_committed", "undo_report"), + phase_checkpoints=phase_checkpoints, + prepared_plan=plan, + repair_required=_last_event(deduped, "repair_required"), + undo_repair_required=_last_event(deduped, "undo_repair_required"), + undo_phase_checkpoints=undo_phase_checkpoints, + ) + + +def _dedupe_events(events: list[dict[str, Any]]) -> list[dict[str, Any]]: + by_id: dict[str, dict[str, Any]] = {} + ordered: list[dict[str, Any]] = [] + for event in events: + event_id = event["event_id"] + existing = by_id.get(event_id) + if existing is None: + by_id[event_id] = event + ordered.append(event) + continue + if existing != event: + raise IdentifyOperationError(f"conflicting duplicate event_id {event_id}") + return ordered + + +def _terminal_status(events: list[dict[str, Any]]) -> str: + kinds = [event["event_kind"] for event in events] + if "undo_repair_required" in kinds: + return "undo_repair_required" + if "undo_committed" in kinds: + return "undone" + if "repair_required" in kinds: + return "repair_required" + if "committed" in kinds: + return "committed" + return "in_progress" + + +def _pending_phases( + terminal: str, + completed: tuple[str, ...], + events: list[dict[str, Any]], +) -> tuple[str, ...]: + if terminal == "in_progress": + completed_set = set(completed) + return tuple( + phase for phase in FORWARD_PHASE_ORDER if phase not in completed_set + ) + if terminal == "repair_required": + repair = _last_event(events, "repair_required") + if repair: + report = repair.get("partial_report", {}) + pending = report.get("pending_phases") + if isinstance(pending, list): + return tuple(str(phase) for phase in pending) + return () + + +def _last_event(events: list[dict[str, Any]], kind: str) -> dict[str, Any] | None: + for event in reversed(events): + if event["event_kind"] == kind: + return event + return None + + +def _last_payload( + events: list[dict[str, Any]], kind: str, field: str +) -> dict[str, Any] | None: + event = _last_event(events, kind) + if event is None: + return None + payload = event.get(field) + return payload if isinstance(payload, dict) else None + + +def _member_tuple(member: dict[str, Any]) -> tuple[str, str, str, str, int]: + try: + return ( + str(member["day"]), + str(member["stream"]), + str(member["segment_key"]), + str(member["source"]), + int(member["sentence_id"]), + ) + except (KeyError, TypeError, ValueError) as exc: + raise IdentifyOperationError("invalid cluster member provenance") from exc + + +def _required_str(row: dict[str, Any], field: str) -> str: + value = row.get(field) + if not isinstance(value, str) or not value: + raise IdentifyOperationError(f"missing or invalid {field}") + return value + + +def _required_dict(row: dict[str, Any], field: str) -> dict[str, Any]: + value = row.get(field) + if not isinstance(value, dict): + raise IdentifyOperationError(f"missing or invalid {field}") + return value + + +def _required_list(row: dict[str, Any], field: str) -> list[Any]: + value = row.get(field) + if not isinstance(value, list): + raise IdentifyOperationError(f"missing or invalid {field}") + return value + + +def _required_int(row: dict[str, Any], field: str) -> int: + value = row.get(field) + if not isinstance(value, int): + raise IdentifyOperationError(f"missing or invalid {field}") + return value + + +def _required_bool(row: dict[str, Any], field: str) -> bool: + value = row.get(field) + if not isinstance(value, bool): + raise IdentifyOperationError(f"missing or invalid {field}") + return value diff --git a/solstone/think/speaker_keep_separate.py b/solstone/think/speaker_keep_separate.py new file mode 100644 index 000000000..86f261d6d --- /dev/null +++ b/solstone/think/speaker_keep_separate.py @@ -0,0 +1,315 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc +"""Append-only speaker keep-separate assertion store. + +Sole write-owner of: + journal/speakers/keep-separate.jsonl +""" + +from __future__ import annotations + +import hashlib +import json +import logging +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from solstone.think.journal_io import append_jsonl, hold_lock +from solstone.think.utils import get_journal + +logger = logging.getLogger(__name__) + +KEEP_SEPARATE_SCHEMA_VERSION = 1 +EVENT_KINDS = {"assert_source", "source_removed"} + + +class KeepSeparateStoreError(RuntimeError): + """Raised when keep-separate assertion storage cannot be trusted.""" + + +@dataclass(frozen=True) +class KeepSeparateAssertion: + """Folded keep-separate assertion for one entity pair.""" + + assertion_id: str + pair_key: str + entity_id_a: str + entity_id_b: str + dismissed_detection_count: int + sources: tuple[dict[str, Any], ...] + created_at: str + updated_at: str + last_recorded_at: str + + @property + def source_count(self) -> int: + return len(self.sources) + + +def keep_separate_dir() -> Path: + """Return the speaker keep-separate directory, creating it if needed.""" + path = Path(get_journal()) / "speakers" + path.mkdir(parents=True, exist_ok=True) + return path + + +def keep_separate_path() -> Path: + """Return the speaker keep-separate JSONL event-log path.""" + return keep_separate_dir() / "keep-separate.jsonl" + + +def pair_key(entity_id_a: str, entity_id_b: str) -> str: + """Return the order-independent entity pair key.""" + return "|".join(sorted([str(entity_id_a), str(entity_id_b)])) + + +def record_keep_separate_assertion( + entity_id_a: str, + entity_id_b: str, + *, + source_kind: str, + operation_id: str | None, + detection_count: int, +) -> dict[str, Any]: + """Append a keep-separate source assertion event.""" + left, right = sorted([str(entity_id_a), str(entity_id_b)]) + event = { + "schema_version": KEEP_SEPARATE_SCHEMA_VERSION, + "event_kind": "assert_source", + "pair_key": pair_key(left, right), + "entity_id_a": left, + "entity_id_b": right, + "source_kind": source_kind, + "operation_id": operation_id, + "detection_count": int(detection_count), + "ts": utc_now_iso(), + } + append_event(event) + return event + + +def remove_operation_sources( + operation_id: str, + pair_keys: list[str] | tuple[str, ...] | set[str], +) -> list[dict[str, Any]]: + """Append tombstones for all current sources from an operation.""" + if not isinstance(operation_id, str) or not operation_id: + raise ValueError("operation_id must be a non-empty string") + wanted = {str(key) for key in pair_keys} + path = keep_separate_path() + appended: list[dict[str, Any]] = [] + with hold_lock(path): + events = _load_jsonl_rows(path) + current_sources = _fold_sources(events) + for key in sorted(wanted): + for source_key in sorted(current_sources.get(key, {})): + source_kind, source_operation_id = source_key + if source_operation_id != operation_id: + continue + event = { + "schema_version": KEEP_SEPARATE_SCHEMA_VERSION, + "event_kind": "source_removed", + "pair_key": key, + "source_kind": source_kind, + "operation_id": operation_id, + "ts": utc_now_iso(), + } + _validate_row(event) + append_jsonl(path, event) + appended.append(event) + current_sources[key].pop(source_key, None) + return appended + + +def append_event(event: dict[str, Any]) -> None: + """Strict-validate and append one keep-separate event.""" + _validate_row(event) + path = keep_separate_path() + with hold_lock(path): + append_jsonl(path, event) + + +def load_events() -> list[dict[str, Any]]: + """Strict-load raw keep-separate events.""" + return _load_jsonl_rows(keep_separate_path()) + + +def fold_assertions() -> list[KeepSeparateAssertion]: + """Fold append-only keep-separate events into active assertions.""" + sources_by_pair = _fold_sources(load_events()) + assertions: list[KeepSeparateAssertion] = [] + for key, sources in sorted(sources_by_pair.items()): + remaining = list(sources.values()) + if not remaining: + continue + detection_count = max(int(source["detection_count"]) for source in remaining) + timestamps = [str(source["recorded_at"]) for source in remaining] + left, right = key.split("|", maxsplit=1) + assertions.append( + KeepSeparateAssertion( + assertion_id=_assertion_id(key), + pair_key=key, + entity_id_a=left, + entity_id_b=right, + dismissed_detection_count=detection_count, + sources=tuple( + sorted( + remaining, + key=lambda source: ( + str(source["source_kind"]), + str(source.get("operation_id") or ""), + ), + ) + ), + created_at=min(timestamps), + updated_at=max(timestamps), + last_recorded_at=max(timestamps), + ) + ) + return assertions + + +def find_assertion(entity_id_a: str, entity_id_b: str) -> KeepSeparateAssertion | None: + """Return the folded assertion for a pair, if present.""" + target = pair_key(entity_id_a, entity_id_b) + for assertion in fold_assertions(): + if assertion.pair_key == target: + return assertion + return None + + +def name_variant_pair_suppressed( + entity_id_a: str, + entity_id_b: str, + current_detection_count: int, +) -> bool: + """Return whether a name-variant pair is suppressed by keep-separate memory.""" + assertion = find_assertion(entity_id_a, entity_id_b) + if assertion is None: + return False + return int(current_detection_count) <= assertion.dismissed_detection_count + + +def list_assertions() -> list[dict[str, Any]]: + """Return redacted summaries for active keep-separate assertions.""" + return [ + { + "assertion_id": assertion.assertion_id, + "pair_key": assertion.pair_key, + "entity_id_a": assertion.entity_id_a, + "entity_id_b": assertion.entity_id_b, + "dismissed_detection_count": assertion.dismissed_detection_count, + "source_count": assertion.source_count, + "created_at": assertion.created_at, + "updated_at": assertion.updated_at, + "last_recorded_at": assertion.last_recorded_at, + } + for assertion in fold_assertions() + ] + + +def utc_now_iso() -> str: + """Return the current UTC time as an ISO-8601 string ending in Z.""" + return ( + datetime.now(timezone.utc).isoformat(timespec="seconds").replace("+00:00", "Z") + ) + + +def _load_jsonl_rows(path: Path) -> list[dict[str, Any]]: + if not path.exists(): + return [] + rows: list[dict[str, Any]] = [] + try: + with open(path, encoding="utf-8") as handle: + for lineno, line in enumerate(handle, start=1): + raw = line.strip() + if not raw: + continue + try: + row = json.loads(raw) + except json.JSONDecodeError as exc: + message = f"malformed keep-separate JSONL at {path}:{lineno}" + logger.error(message) + raise KeepSeparateStoreError(message) from exc + if not isinstance(row, dict): + message = f"non-object keep-separate JSONL at {path}:{lineno}" + logger.error(message) + raise KeepSeparateStoreError(message) + try: + _validate_row(row) + except KeepSeparateStoreError as exc: + message = f"invalid keep-separate row at {path}:{lineno}: {exc}" + logger.error(message) + raise KeepSeparateStoreError(message) from exc + rows.append(row) + except OSError as exc: + message = f"failed to read keep-separate store {path}: {exc}" + logger.error(message) + raise KeepSeparateStoreError(message) from exc + return rows + + +def _fold_sources( + events: list[dict[str, Any]], +) -> dict[str, dict[tuple[str, str | None], dict[str, Any]]]: + sources_by_pair: dict[str, dict[tuple[str, str | None], dict[str, Any]]] = {} + for event in events: + key = str(event["pair_key"]) + sources = sources_by_pair.setdefault(key, {}) + source_key = (str(event["source_kind"]), event.get("operation_id")) + if event["event_kind"] == "source_removed": + sources.pop(source_key, None) + continue + existing = sources.get(source_key) + detection_count = int(event["detection_count"]) + if existing is not None and int(existing["detection_count"]) >= detection_count: + continue + sources[source_key] = { + "source_kind": str(event["source_kind"]), + "operation_id": event.get("operation_id"), + "detection_count": detection_count, + "recorded_at": str(event["ts"]), + } + return sources_by_pair + + +def _validate_row(row: dict[str, Any]) -> None: + if row.get("schema_version") != KEEP_SEPARATE_SCHEMA_VERSION: + raise KeepSeparateStoreError("invalid schema_version") + event_kind = _required_str(row, "event_kind") + if event_kind not in EVENT_KINDS: + raise KeepSeparateStoreError(f"unknown event_kind: {event_kind}") + key = _required_str(row, "pair_key") + _required_str(row, "source_kind") + _required_str(row, "ts") + + if event_kind == "assert_source": + entity_id_a = _required_str(row, "entity_id_a") + entity_id_b = _required_str(row, "entity_id_b") + if key != pair_key(entity_id_a, entity_id_b): + raise KeepSeparateStoreError("pair_key does not match entity ids") + operation_id = row.get("operation_id") + if operation_id is not None and not isinstance(operation_id, str): + raise KeepSeparateStoreError("operation_id must be string or null") + detection_count = row.get("detection_count") + if not isinstance(detection_count, int) or detection_count < 1: + raise KeepSeparateStoreError("detection_count must be a positive int") + return + + operation_id = row.get("operation_id") + if not isinstance(operation_id, str) or not operation_id: + raise KeepSeparateStoreError("source_removed operation_id is required") + + +def _assertion_id(key: str) -> str: + digest = hashlib.sha256(key.encode("utf-8")).hexdigest() + return f"ksep_{digest[:24]}" + + +def _required_str(row: dict[str, Any], field: str) -> str: + value = row.get(field) + if not isinstance(value, str) or not value: + raise KeepSeparateStoreError(f"missing or invalid {field}") + return value diff --git a/tests/test_speaker_cluster_dismissals.py b/tests/test_speaker_cluster_dismissals.py new file mode 100644 index 000000000..92f06b9b6 --- /dev/null +++ b/tests/test_speaker_cluster_dismissals.py @@ -0,0 +1,103 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +from pathlib import Path + +import pytest + +import solstone.think.speaker_cluster_dismissals as store + + +@pytest.fixture +def dismissal_journal(monkeypatch, tmp_path): + monkeypatch.setenv("SOLSTONE_JOURNAL", str(tmp_path)) + import solstone.think.utils as think_utils + + think_utils._journal_path_cache = None + return Path(tmp_path) + + +def _members(start: int, count: int) -> list[dict[str, object]]: + return [ + { + "day": "20260101", + "stream": "test", + "segment_key": "090000_300", + "source": "mic_audio", + "sentence_id": sid, + } + for sid in range(start, start + count) + ] + + +def test_identical_members_are_suppressed(dismissal_journal): + members = _members(1, 3) + store.record_cluster_dismissal(members, "quiet") + + assert store.cluster_dismissal_suppressed(list(reversed(members))) is True + + +def test_grown_rescan_sharing_exactly_half_is_suppressed(dismissal_journal): + store.record_cluster_dismissal(_members(1, 10), "quiet") + candidate = _members(1, 5) + _members(100, 5) + + assert store.cluster_dismissal_suppressed(candidate) is True + + +def test_sharing_less_than_half_is_not_suppressed(dismissal_journal): + store.record_cluster_dismissal(_members(1, 10), "quiet") + candidate = _members(1, 4) + _members(100, 6) + + assert store.cluster_dismissal_suppressed(candidate) is False + + +def test_overlapping_events_fold_to_union_without_duplicates( + dismissal_journal, monkeypatch +): + times = iter(["2026-07-20T12:00:00Z", "2026-07-20T12:00:01Z"]) + monkeypatch.setattr(store, "utc_now_iso", lambda: next(times)) + + store.record_cluster_dismissal(_members(1, 4), "quiet") + store.record_cluster_dismissal(_members(3, 4), "quiet") + + folded = store.fold_dismissals() + assert len(folded) == 1 + assert folded[0].member_count == 6 + assert [member["sentence_id"] for member in folded[0].members] == [1, 2, 3, 4, 5, 6] + assert folded[0].event_count == 2 + assert folded[0].created_at == "2026-07-20T12:00:00Z" + assert folded[0].updated_at == "2026-07-20T12:00:01Z" + + +def test_not_a_person_dominates_quiet_on_merge(dismissal_journal, monkeypatch): + times = iter(["2026-07-20T12:00:00Z", "2026-07-20T12:00:01Z"]) + monkeypatch.setattr(store, "utc_now_iso", lambda: next(times)) + + store.record_cluster_dismissal(_members(1, 4), "quiet") + store.record_cluster_dismissal(_members(3, 4), "not_a_person") + + folded = store.fold_dismissals() + assert len(folded) == 1 + assert folded[0].disposition == "not_a_person" + assert store.list_dismissals()[0]["disposition"] == "not_a_person" + + +def test_second_dismissal_appends_without_rewriting_first_event(dismissal_journal): + store.record_cluster_dismissal(_members(1, 2), "quiet") + first = store.cluster_dismissals_path().read_text(encoding="utf-8") + + store.record_cluster_dismissal(_members(10, 2), "quiet") + second = store.cluster_dismissals_path().read_text(encoding="utf-8") + + assert second.startswith(first) + assert len(first.splitlines()) == 1 + assert len(second.splitlines()) == 2 + + +def test_strict_malformed_row_raises(dismissal_journal): + store.cluster_dismissals_path().write_text("not-json\n", encoding="utf-8") + + with pytest.raises(store.ClusterDismissalStoreError): + store.fold_dismissals() diff --git a/tests/test_speaker_identify_operations.py b/tests/test_speaker_identify_operations.py new file mode 100644 index 000000000..17514f4c3 --- /dev/null +++ b/tests/test_speaker_identify_operations.py @@ -0,0 +1,324 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +import json +from pathlib import Path + +import pytest + +import solstone.think.speaker_identify_operations as ledger + + +@pytest.fixture +def op_journal(monkeypatch, tmp_path): + monkeypatch.setenv("SOLSTONE_JOURNAL", str(tmp_path)) + import solstone.think.utils as think_utils + + think_utils._journal_path_cache = None + return Path(tmp_path) + + +def _members() -> list[dict[str, object]]: + return [ + { + "day": "20260101", + "stream": "test", + "segment_key": "090000_300", + "source": "mic_audio", + "sentence_id": 1, + }, + { + "day": "20260101", + "stream": "test", + "segment_key": "090000_300", + "source": "mic_audio", + "sentence_id": 2, + }, + ] + + +def _prepared_event(request_id: str = "req-1") -> dict[str, object]: + operation_id = ledger.operation_id_for_request(request_id) + members = _members() + fingerprint = ledger.request_fingerprint( + cluster_members=members, + target_entity_id="alice_test", + will_create=True, + entity_type="Person", + reviewed_near_match_entity_ids=["bob_test"], + ) + plan = { + "plan_schema_version": 1, + "operation_id": operation_id, + "request_id": request_id, + "planned_at": "2026-07-20T12:00:00Z", + "request": { + "cluster_id": 7, + "name": "Alice Test", + "entity_id": "alice_test", + "resolve_only": False, + "create_new": True, + "entity_type": "Person", + "reviewed_near_match_entity_ids": ["bob_test"], + }, + "cluster": {"cluster_id": 7, "member_count": len(members), "members": members}, + "target": { + "entity_id": "alice_test", + "entity_name": "Alice Test", + "entity_type": "Person", + "will_create": True, + }, + "entity_identity": { + "prior_identity": None, + "intended_identity": { + "id": "alice_test", + "name": "Alice Test", + "type": "Person", + "created_at": 1, + }, + "expected_history_operation": { + "operation_kind": "speaker_identify", + "operation_id": operation_id, + }, + }, + "direct_voiceprints": {"preexisting_keys": [], "entries_to_add": []}, + "segments": [], + "retro_confirm": { + "matched": False, + "match_score": None, + "candidate_id": None, + "candidate_before": None, + "candidate_after": None, + "preexisting_voiceprint_keys": [], + "voiceprints_to_add": [], + }, + "sentinel": { + "cluster_key": "7", + "prior_entry": None, + "intended_entry": { + "entity_id": "alice_test", + "label": "Alice Test", + "ts": "2026-07-20T12:00:00Z", + }, + }, + "keep_separate_assertions": [], + } + return { + "schema_version": ledger.IDENTIFY_OPERATION_SCHEMA_VERSION, + "event_id": f"{operation_id}:prepared", + "operation_id": operation_id, + "request_id": request_id, + "event_kind": "prepared", + "ts": "2026-07-20T12:00:00Z", + "caller": "test", + "actor": None, + "request_fingerprint": fingerprint, + "prepared_plan": plan, + } + + +def _checkpoint_event( + operation_id: str, request_id: str, phase: str +) -> dict[str, object]: + return { + "schema_version": ledger.IDENTIFY_OPERATION_SCHEMA_VERSION, + "event_id": f"{operation_id}:checkpoint:{phase}", + "operation_id": operation_id, + "request_id": request_id, + "event_kind": "checkpoint", + "ts": "2026-07-20T12:00:01Z", + "caller": "test", + "actor": None, + "phase": phase, + "checkpoint": { + "phase_status": "complete", + "completed_at": "2026-07-20T12:00:01Z", + "counts": {"saved_count": 2}, + "skipped_reasons": {}, + "saved_count": 2, + "skipped_existing_count": 0, + "saved_keys": [ + { + "day": "20260101", + "segment_key": "090000_300", + "source": "mic_audio", + "sentence_id": 1, + } + ], + }, + } + + +def test_operation_id_for_request_is_deterministic(): + assert ledger.operation_id_for_request("abc") == ledger.operation_id_for_request( + "abc" + ) + assert ledger.operation_id_for_request("abc").startswith("idop_") + assert ledger.operation_id_for_request("abc") != ledger.operation_id_for_request( + "abcd" + ) + + +def test_request_fingerprint_changes_for_each_identity_input(): + base = ledger.request_fingerprint( + cluster_members=_members(), + target_entity_id="alice_test", + will_create=True, + entity_type="Person", + reviewed_near_match_entity_ids=["bob_test"], + ) + changed_member = ledger.request_fingerprint( + cluster_members=[{**_members()[0], "sentence_id": 99}], + target_entity_id="alice_test", + will_create=True, + entity_type="Person", + reviewed_near_match_entity_ids=["bob_test"], + ) + changed_target = ledger.request_fingerprint( + cluster_members=_members(), + target_entity_id="carol_test", + will_create=True, + entity_type="Person", + reviewed_near_match_entity_ids=["bob_test"], + ) + changed_create = ledger.request_fingerprint( + cluster_members=_members(), + target_entity_id="alice_test", + will_create=False, + entity_type="Person", + reviewed_near_match_entity_ids=["bob_test"], + ) + changed_type = ledger.request_fingerprint( + cluster_members=_members(), + target_entity_id="alice_test", + will_create=True, + entity_type="Project", + reviewed_near_match_entity_ids=["bob_test"], + ) + changed_reviewed = ledger.request_fingerprint( + cluster_members=_members(), + target_entity_id="alice_test", + will_create=True, + entity_type="Person", + reviewed_near_match_entity_ids=["carol_test"], + ) + + assert len(base) == 64 + assert ( + len( + { + base, + changed_member, + changed_target, + changed_create, + changed_type, + changed_reviewed, + } + ) + == 6 + ) + assert base not in { + changed_member, + changed_target, + changed_create, + changed_type, + changed_reviewed, + } + + +def test_append_and_fold_prepared_resume_state(op_journal): + prepared = _prepared_event() + checkpoint = _checkpoint_event( + str(prepared["operation_id"]), + str(prepared["request_id"]), + "direct_voiceprints", + ) + + ledger.append_event(prepared) + ledger.append_event(checkpoint) + + state = ledger.fold_operation(str(prepared["operation_id"])) + assert state is not None + assert state.operation_id == prepared["operation_id"] + assert state.request_fingerprint == prepared["request_fingerprint"] + assert state.cluster_member_set == { + ("20260101", "test", "090000_300", "mic_audio", 1), + ("20260101", "test", "090000_300", "mic_audio", 2), + } + assert state.target_entity_id == "alice_test" + assert state.target_entity_name == "Alice Test" + assert state.will_create is True + assert state.entity_type == "Person" + assert state.reviewed_near_match_entity_ids == ("bob_test",) + assert state.completed_phases == ("direct_voiceprints",) + assert state.pending_phases == ( + "entity", + "keep_separate", + "corrections", + "labels", + "retro_tracker", + "sentinel", + ) + assert state.terminal_status == "in_progress" + assert state.phase_checkpoints["direct_voiceprints"]["counts"]["saved_count"] == 2 + + +def test_committed_fold_returns_stored_result(op_journal): + prepared = _prepared_event() + operation_id = str(prepared["operation_id"]) + committed = { + "schema_version": ledger.IDENTIFY_OPERATION_SCHEMA_VERSION, + "event_id": f"{operation_id}:committed", + "operation_id": operation_id, + "request_id": prepared["request_id"], + "event_kind": "committed", + "ts": "2026-07-20T12:00:02Z", + "caller": "test", + "actor": None, + "result": {"status": "identified", "operation_id": operation_id}, + } + + ledger.append_event(prepared) + ledger.append_event(committed) + + state = ledger.fold_operation(operation_id) + assert state is not None + assert state.terminal_status == "committed" + assert state.result == {"status": "identified", "operation_id": operation_id} + assert state.pending_phases == () + + +def test_identical_duplicate_event_id_folds_once(op_journal): + prepared = _prepared_event() + path = ledger.identify_operations_path() + path.write_text( + json.dumps(prepared) + "\n" + json.dumps(prepared) + "\n", + encoding="utf-8", + ) + + state = ledger.fold_operation(str(prepared["operation_id"])) + assert state is not None + assert state.terminal_status == "in_progress" + + +def test_non_identical_duplicate_event_id_raises(op_journal): + prepared = _prepared_event() + changed = dict(prepared) + changed["ts"] = "2026-07-20T12:00:09Z" + path = ledger.identify_operations_path() + path.write_text( + json.dumps(prepared) + "\n" + json.dumps(changed) + "\n", + encoding="utf-8", + ) + + with pytest.raises(ledger.IdentifyOperationError): + ledger.fold_operation(str(prepared["operation_id"])) + + +def test_strict_malformed_row_raises(op_journal): + ledger.identify_operations_path().write_text("not-json\n", encoding="utf-8") + + with pytest.raises(ledger.IdentifyOperationError): + ledger.load_operations() diff --git a/tests/test_speaker_keep_separate.py b/tests/test_speaker_keep_separate.py new file mode 100644 index 000000000..3674e82c4 --- /dev/null +++ b/tests/test_speaker_keep_separate.py @@ -0,0 +1,146 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +from pathlib import Path + +import pytest + +import solstone.think.speaker_keep_separate as store + + +@pytest.fixture +def keep_journal(monkeypatch, tmp_path): + monkeypatch.setenv("SOLSTONE_JOURNAL", str(tmp_path)) + import solstone.think.utils as think_utils + + think_utils._journal_path_cache = None + return Path(tmp_path) + + +def test_pair_key_is_order_independent(keep_journal): + assert store.pair_key("alice", "bob") == "alice|bob" + assert store.pair_key("alice", "bob") == store.pair_key("bob", "alice") + + +def test_assertion_folds_present_with_watermark(keep_journal, monkeypatch): + monkeypatch.setattr(store, "utc_now_iso", lambda: "2026-07-20T12:00:00Z") + + store.record_keep_separate_assertion( + "bob", + "alice", + source_kind="explicit_create_near_match", + operation_id="idop_1", + detection_count=3, + ) + + assertion = store.find_assertion("alice", "bob") + assert assertion is not None + assert assertion.pair_key == "alice|bob" + assert assertion.entity_id_a == "alice" + assert assertion.entity_id_b == "bob" + assert assertion.dismissed_detection_count == 3 + assert assertion.source_count == 1 + assert store.list_assertions()[0]["source_count"] == 1 + + +def test_reassert_higher_detection_count_raises_watermark(keep_journal, monkeypatch): + times = iter(["2026-07-20T12:00:00Z", "2026-07-20T12:00:01Z"]) + monkeypatch.setattr(store, "utc_now_iso", lambda: next(times)) + + store.record_keep_separate_assertion( + "alice", + "bob", + source_kind="explicit_create_near_match", + operation_id="idop_1", + detection_count=2, + ) + store.record_keep_separate_assertion( + "bob", + "alice", + source_kind="explicit_create_near_match", + operation_id="idop_1", + detection_count=5, + ) + + assertion = store.find_assertion("alice", "bob") + assert assertion is not None + assert assertion.dismissed_detection_count == 5 + assert assertion.source_count == 1 + assert assertion.updated_at == "2026-07-20T12:00:01Z" + + +def test_tombstone_removes_one_source(keep_journal, monkeypatch): + times = iter( + [ + "2026-07-20T12:00:00Z", + "2026-07-20T12:00:01Z", + "2026-07-20T12:00:02Z", + ] + ) + monkeypatch.setattr(store, "utc_now_iso", lambda: next(times)) + key = store.pair_key("alice", "bob") + + store.record_keep_separate_assertion( + "alice", + "bob", + source_kind="explicit_create_near_match", + operation_id="idop_1", + detection_count=2, + ) + store.record_keep_separate_assertion( + "alice", + "bob", + source_kind="explicit_create_near_match", + operation_id="idop_2", + detection_count=7, + ) + tombstones = store.remove_operation_sources("idop_1", [key]) + + assert len(tombstones) == 1 + assertion = store.find_assertion("alice", "bob") + assert assertion is not None + assert assertion.dismissed_detection_count == 7 + assert assertion.source_count == 1 + + +def test_all_sources_removed_makes_assertion_absent(keep_journal, monkeypatch): + times = iter(["2026-07-20T12:00:00Z", "2026-07-20T12:00:01Z"]) + monkeypatch.setattr(store, "utc_now_iso", lambda: next(times)) + key = store.pair_key("alice", "bob") + + store.record_keep_separate_assertion( + "alice", + "bob", + source_kind="explicit_create_near_match", + operation_id="idop_1", + detection_count=2, + ) + store.remove_operation_sources("idop_1", [key]) + + assert store.find_assertion("alice", "bob") is None + assert store.list_assertions() == [] + + +def test_suppression_predicate_boundaries(keep_journal): + assert store.name_variant_pair_suppressed("alice", "bob", 1) is False + + store.record_keep_separate_assertion( + "alice", + "bob", + source_kind="explicit_create_near_match", + operation_id="idop_1", + detection_count=2, + ) + + assert store.name_variant_pair_suppressed("bob", "alice", 1) is True + assert store.name_variant_pair_suppressed("bob", "alice", 2) is True + assert store.name_variant_pair_suppressed("bob", "alice", 3) is False + + +def test_strict_malformed_row_raises(keep_journal): + store.keep_separate_path().write_text("not-json\n", encoding="utf-8") + + with pytest.raises(store.KeepSeparateStoreError): + store.fold_assertions() diff --git a/tests/test_voiceprint_removal.py b/tests/test_voiceprint_removal.py new file mode 100644 index 000000000..cda698333 --- /dev/null +++ b/tests/test_voiceprint_removal.py @@ -0,0 +1,132 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +import json +from pathlib import Path + +import numpy as np +import pytest + +from solstone.think.entities.voiceprints import ( + VoiceprintRemovalError, + remove_voiceprints_by_key, + save_voiceprints_batch, +) + + +@pytest.fixture +def voiceprint_journal(monkeypatch, tmp_path): + monkeypatch.setenv("SOLSTONE_JOURNAL", str(tmp_path)) + import solstone.think.utils as think_utils + + think_utils._journal_path_cache = None + return Path(tmp_path) + + +def _embedding(value: float) -> np.ndarray: + vector = np.zeros(256, dtype=np.float32) + vector[0] = value + return vector + + +def _metadata(sentence_id: int, *, extra: str | None = None) -> dict[str, object]: + metadata: dict[str, object] = { + "day": "20260101", + "segment_key": "090000_300", + "source": "mic_audio", + "stream": "test", + "sentence_id": sentence_id, + "added_at": 1, + } + if extra is not None: + metadata["extra"] = extra + return metadata + + +def _removal(metadata: dict[str, object]) -> dict[str, object]: + return { + "key": { + "day": metadata["day"], + "segment_key": metadata["segment_key"], + "source": metadata["source"], + "sentence_id": metadata["sentence_id"], + }, + "expected_metadata": metadata, + } + + +def _load_metadata(path: Path) -> list[dict[str, object]]: + with np.load(path, allow_pickle=False) as data: + return [json.loads(str(item)) for item in data["metadata"]] + + +def test_remove_one_voiceprint_by_key_and_metadata(voiceprint_journal): + first = _metadata(1) + second = _metadata(2) + save_voiceprints_batch( + "alice", + [(_embedding(1.0), first), (_embedding(2.0), second)], + ) + path = voiceprint_journal / "entities" / "alice" / "voiceprints.npz" + + report = remove_voiceprints_by_key("alice", [_removal(first)]) + + assert report == { + "removed_count": 1, + "skipped_count": 0, + "skipped_reasons": {"missing": 0, "metadata_mismatch": 0}, + "file_removed": False, + } + assert _load_metadata(path) == [second] + + +def test_missing_key_is_skipped(voiceprint_journal): + existing = _metadata(1) + missing = _metadata(99) + save_voiceprints_batch("alice", [(_embedding(1.0), existing)]) + + report = remove_voiceprints_by_key("alice", [_removal(missing)]) + + assert report["removed_count"] == 0 + assert report["skipped_count"] == 1 + assert report["skipped_reasons"] == {"missing": 1, "metadata_mismatch": 0} + + +def test_metadata_mismatch_is_skipped(voiceprint_journal): + existing = _metadata(1) + expected = _metadata(1, extra="different") + save_voiceprints_batch("alice", [(_embedding(1.0), existing)]) + + report = remove_voiceprints_by_key("alice", [_removal(expected)]) + + assert report["removed_count"] == 0 + assert report["skipped_count"] == 1 + assert report["skipped_reasons"] == {"missing": 0, "metadata_mismatch": 1} + + +def test_remove_all_deletes_file(voiceprint_journal): + existing = _metadata(1) + save_voiceprints_batch("alice", [(_embedding(1.0), existing)]) + path = voiceprint_journal / "entities" / "alice" / "voiceprints.npz" + + report = remove_voiceprints_by_key("alice", [_removal(existing)]) + + assert report["removed_count"] == 1 + assert report["file_removed"] is True + assert not path.exists() + + +def test_duplicate_exact_match_raises(voiceprint_journal): + metadata = _metadata(1) + path = voiceprint_journal / "entities" / "alice" / "voiceprints.npz" + path.parent.mkdir(parents=True) + np.savez_compressed( + path, + embeddings=np.stack([_embedding(1.0), _embedding(2.0)]), + metadata=np.asarray([json.dumps(metadata), json.dumps(metadata)], dtype=str), + ) + + with pytest.raises(VoiceprintRemovalError): + remove_voiceprints_by_key("alice", [_removal(metadata)])