diff --git a/src/freshdata/learning/merge.py b/src/freshdata/learning/merge.py index 0501cdb9..50723178 100644 --- a/src/freshdata/learning/merge.py +++ b/src/freshdata/learning/merge.py @@ -15,6 +15,7 @@ from __future__ import annotations import copy +from collections.abc import Mapping from dataclasses import dataclass, field from datetime import datetime, timezone from typing import Any, Literal @@ -294,7 +295,10 @@ def merge_profiles(self_profile: Any, other_profile: Any, *, strategy: Strategy) set(parent_audits[0].protection_candidates) | set(parent_audits[1].protection_candidates) ), - alignment={"merged": True}, + alignment=_merged_alignment( + parent_audits[1] if strategy == "prefer_other" else parent_audits[0], + parent_audits[0] if strategy == "prefer_other" else parent_audits[1], + ), holdout_metrics={}, demotions=[*parent_audits[0].demotions, *parent_audits[1].demotions], notes=notes, @@ -302,6 +306,29 @@ def merge_profiles(self_profile: Any, other_profile: Any, *, strategy: Strategy) return merged +def _source_schema(audit: Any) -> dict[str, str]: + schema = audit.alignment.get("source_schema") if audit is not None else None + if not isinstance(schema, Mapping): + return {} + return {str(c): str(t) for c, t in schema.items()} + + +def _merged_alignment(base_audit: Any, other_audit: Any) -> dict[str, Any]: + """Alignment block for a merged audit. + + Carries the parents' ``source_schema`` (union; the base parent's dtype wins + a disagreement) so replay drift checks compare real pandas dtypes instead + of falling back to the embedded memory's coarse type names. + """ + schema = _source_schema(base_audit) + for column, dtype in _source_schema(other_audit).items(): + schema.setdefault(column, dtype) + alignment: dict[str, Any] = {"merged": True} + if schema: + alignment["source_schema"] = schema + return alignment + + def _resolution_word(strategy: str) -> str: return { "union_min_precision": "dropped both sides", diff --git a/src/freshdata/learning/profile.py b/src/freshdata/learning/profile.py index 5bb93504..396f2320 100644 --- a/src/freshdata/learning/profile.py +++ b/src/freshdata/learning/profile.py @@ -187,27 +187,14 @@ def load(cls, path: str | Path) -> LearningProfile: if "manifest.json" not in raw_members: raise ProfileFormatError(f"{target}: missing required member manifest.json") - manifest = ProfileManifest.from_dict( - json.loads(raw_members["manifest.json"].decode("utf-8")) - ) + manifest = _parse_manifest(raw_members["manifest.json"], target) _check_version(manifest.profile_version, target) missing = [m for m in _REQUIRED_MEMBERS if m not in raw_members] if missing: raise ProfileFormatError(f"{target}: missing required member(s): {', '.join(missing)}") - for name, expected in manifest.member_hashes.items(): - if name not in raw_members: - raise ProfileFormatError(f"{target}: member {name} listed in manifest but absent") - actual = ( - sha256_text(raw_members[name].decode("utf-8")) - if name.endswith(".json") - else _sha256_bytes(raw_members[name]) - ) - if actual != expected: - raise ProfileFormatError( - f"{target}: hash mismatch for {name} (profile corrupt or tampered)" - ) + _verify_member_hashes(manifest, raw_members, target) known = set(_REQUIRED_MEMBERS) | {_VECTORS_MEMBER} for name in sorted(names - known): @@ -218,12 +205,6 @@ def load(cls, path: str | Path) -> LearningProfile: stacklevel=2, ) - rules_payload = json.loads(raw_members["rules.json"].decode("utf-8")) - maps_payload = json.loads(raw_members["value_maps.json"].decode("utf-8")) - memory_payload = json.loads(raw_members["memory.json"].decode("utf-8")) - examples_payload = json.loads(raw_members["examples.json"].decode("utf-8")) - audit_payload = json.loads(raw_members["audit.json"].decode("utf-8")) - vectors = None if _VECTORS_MEMBER in raw_members: try: @@ -239,21 +220,31 @@ def load(cls, path: str | Path) -> LearningProfile: stacklevel=2, ) - memory_data = memory_payload.get("memory") - examples_data = examples_payload.get("examples") - audit_data = audit_payload.get("audit") - return cls( - manifest=manifest, - rules=[ColumnConstraint.from_dict(r) for r in rules_payload.get("rules", [])], - value_maps={ - column: ValueMap.from_dict(vm) - for column, vm in maps_payload.get("value_maps", {}).items() - }, - examples=ExampleBank.from_dict(examples_data) if examples_data else None, - memory=CleaningMemory.from_dict(memory_data) if memory_data else None, - audit_info=ProfileAudit.from_dict(audit_data) if audit_data else None, - vectors=vectors, - ) + try: + payloads = { + name: json.loads(raw_members[name].decode("utf-8")) + for name in _REQUIRED_MEMBERS + if name != "manifest.json" + } + memory_data = payloads["memory.json"].get("memory") + examples_data = payloads["examples.json"].get("examples") + audit_data = payloads["audit.json"].get("audit") + return cls( + manifest=manifest, + rules=[ + ColumnConstraint.from_dict(r) for r in payloads["rules.json"].get("rules", []) + ], + value_maps={ + column: ValueMap.from_dict(vm) + for column, vm in payloads["value_maps.json"].get("value_maps", {}).items() + }, + examples=ExampleBank.from_dict(examples_data) if examples_data else None, + memory=CleaningMemory.from_dict(memory_data) if memory_data else None, + audit_info=ProfileAudit.from_dict(audit_data) if audit_data else None, + vectors=vectors, + ) + except (ValueError, KeyError, TypeError, AttributeError) as exc: + raise ProfileFormatError(f"{target}: malformed profile member: {exc!r}") from exc # -- introspection ------------------------------------------------------- @@ -305,6 +296,50 @@ def _sha256_bytes(payload: bytes) -> str: return hashlib.sha256(payload).hexdigest() +def _parse_manifest(raw: bytes, target: Path) -> ProfileManifest: + """Decode ``manifest.json``; any malformed content is a ProfileFormatError.""" + try: + data = json.loads(raw.decode("utf-8")) + if not isinstance(data, Mapping): + raise TypeError("manifest.json must contain a JSON object") + return ProfileManifest.from_dict(data) + except (ValueError, KeyError, TypeError) as exc: + raise ProfileFormatError(f"{target}: invalid manifest.json: {exc!r}") from exc + + +def _verify_member_hashes( + manifest: ProfileManifest, raw_members: Mapping[str, bytes], target: Path +) -> None: + """Check that every member is hashed in the manifest and every hash matches.""" + # Every member must be covered by a hash: a manifest that omits one + # (partial write, hand edit) would otherwise skip its verification. + must_hash = [m for m in _REQUIRED_MEMBERS if m != "manifest.json"] + if _VECTORS_MEMBER in raw_members: + must_hash.append(_VECTORS_MEMBER) + unhashed = [m for m in must_hash if m not in manifest.member_hashes] + if unhashed: + raise ProfileFormatError( + f"{target}: manifest has no member hash for: {', '.join(unhashed)} " + "(profile incomplete or tampered)" + ) + + for name, expected in manifest.member_hashes.items(): + if name not in raw_members: + raise ProfileFormatError(f"{target}: member {name} listed in manifest but absent") + try: + actual = ( + sha256_text(raw_members[name].decode("utf-8")) + if name.endswith(".json") + else _sha256_bytes(raw_members[name]) + ) + except UnicodeDecodeError as exc: + raise ProfileFormatError(f"{target}: member {name} is not valid UTF-8") from exc + if actual != expected: + raise ProfileFormatError( + f"{target}: hash mismatch for {name} (profile corrupt or tampered)" + ) + + def _check_version(version: str, target: Path) -> None: try: major = int(str(version).split(".", 1)[0]) diff --git a/src/freshdata/learning/replay.py b/src/freshdata/learning/replay.py index 57fc2bb5..44705dd0 100644 --- a/src/freshdata/learning/replay.py +++ b/src/freshdata/learning/replay.py @@ -76,6 +76,17 @@ def resolve_profile(profile: object) -> LearningProfile: ) +#: CleaningMemory signatures store coarse type names; map them to the pandas +#: dtype names ``check_profile_drift`` compares against. +_COARSE_TO_DTYPE = { + "text": "object", + "integer": "int64", + "float": "float64", + "boolean": "bool", + "datetime": "datetime64[ns]", +} + + def _profile_schema(profile: LearningProfile) -> dict[str, str]: """Source schema the profile was learned from (audit first, memory next).""" if profile.audit_info is not None: @@ -85,7 +96,7 @@ def _profile_schema(profile: LearningProfile) -> dict[str, str]: if profile.memory is not None and isinstance(profile.memory.signature, Mapping): columns = profile.memory.signature.get("columns") if isinstance(columns, Mapping) and columns: - return {str(c): str(t) for c, t in columns.items()} + return {str(c): _COARSE_TO_DTYPE.get(str(t), str(t)) for c, t in columns.items()} return {} diff --git a/src/freshdata/memory.py b/src/freshdata/memory.py index 638c6fe6..57c8fae3 100644 --- a/src/freshdata/memory.py +++ b/src/freshdata/memory.py @@ -16,9 +16,12 @@ from __future__ import annotations import json +import math +import os import sqlite3 from dataclasses import dataclass, field from datetime import datetime, timezone +from pathlib import Path from typing import TYPE_CHECKING, Any import pandas as pd @@ -61,6 +64,29 @@ def _signature(df: pd.DataFrame) -> dict[str, Any]: return {"columns": cols, "n_cols": len(cols), "hash": digest} +def _is_missing_cell(value: Any) -> bool: + """True for ``None``, NaN/NaT/``pd.NA`` and blank strings.""" + if value is None: + return True + if isinstance(value, str): + return not value.strip() + try: + return bool(pd.isna(value)) + except (TypeError, ValueError): # non-scalar cells are never "missing" + return False + + +def _json_safe(value: Any) -> Any: + """Recursively replace non-finite floats (NaN/inf) with ``None``.""" + if isinstance(value, float): + return value if math.isfinite(value) else None + if isinstance(value, dict): + return {k: _json_safe(v) for k, v in value.items()} + if isinstance(value, (list, tuple)): + return [_json_safe(v) for v in value] + return value + + def _normalize_decisions(decisions: Any) -> tuple[list[dict], list[dict]]: """Split *decisions* into (accepted, rejected) lists of plain dicts. @@ -91,6 +117,10 @@ def _normalize_decisions(decisions: Any) -> tuple[list[dict], list[dict]]: accepted: list[dict] = [] rejected: list[dict] = [] for r in rows: + # A decisions table read back from CSV holds NaN (or "") for table-level + # steps; replay keys on ``column is None``, so normalise missing cells. + if "column" in r and _is_missing_cell(r["column"]): + r["column"] = None status = str(r.get("status", "accepted")).lower() (rejected if status in ("rejected", "skipped", "declined") else accepted).append(r) return accepted, rejected @@ -197,7 +227,7 @@ def from_dict(cls, payload: dict[str, Any]) -> CleaningMemory: def to_json(self, path: str | None = None) -> str: """Serialize to JSON. With *path*, also write it (``.sqlite``/``.db`` → a tiny key/value SQLite store). Returns the JSON string.""" - text = json.dumps(self.to_dict(), indent=2, default=str) + text = json.dumps(_json_safe(self.to_dict()), indent=2, default=str) if path is not None: if str(path).endswith((".sqlite", ".db")): _sqlite_write(path, self.dataset_id, text) @@ -314,14 +344,28 @@ def learn_cleaning_memory( ) -def load_cleaning_memory(path: str) -> CleaningMemory: - """Load a memory previously saved with :meth:`CleaningMemory.to_json`.""" +def load_cleaning_memory( + path: str | os.PathLike[str], *, dataset_id: str | None = None +) -> CleaningMemory: + """Load a memory previously saved with :meth:`CleaningMemory.to_json`. + + A ``.sqlite``/``.db`` store can hold one memory per ``dataset_id``; pass + *dataset_id* to choose one. Loading a store that holds several memories + without *dataset_id* raises :class:`ValueError` listing the stored ids. + A missing *path* raises :class:`FileNotFoundError` (nothing is created). + """ + if not Path(path).exists(): + raise FileNotFoundError(f"cleaning memory not found: {str(path)!r}") if str(path).endswith((".sqlite", ".db")): - text = _sqlite_read(path) - else: - with open(path, encoding="utf-8") as fh: - text = fh.read() - return CleaningMemory.from_dict(json.loads(text)) + text = _sqlite_read(path, dataset_id) + return CleaningMemory.from_dict(json.loads(text)) + with open(path, encoding="utf-8") as fh: + memory = CleaningMemory.from_dict(json.loads(fh.read())) + if dataset_id is not None and memory.dataset_id != dataset_id: + raise KeyError( + f"no cleaning memory for dataset_id {dataset_id!r} in {str(path)!r} " + f"(it holds {memory.dataset_id!r})") + return memory # -- minimal server-free SQLite key/value store ----------------------------- @@ -341,19 +385,29 @@ def _sqlite_write(path: str, dataset_id: str, text: str) -> None: conn.close() -def _sqlite_read(path: str, dataset_id: str | None = None) -> str: - conn = sqlite3.connect(path) +def _sqlite_read(path: str | os.PathLike[str], dataset_id: str | None = None) -> str: + # Read-only URI: a mistyped path must never create an empty database. + uri = Path(path).resolve().as_uri() + "?mode=ro" + conn = sqlite3.connect(uri, uri=True) try: if dataset_id is not None: row = conn.execute( "SELECT payload FROM cleaning_memory WHERE dataset_id=?", (dataset_id,)).fetchone() - else: - row = conn.execute( - "SELECT payload FROM cleaning_memory LIMIT 1").fetchone() - if not row: - raise KeyError(f"no cleaning memory found in {path!r}") - return str(row[0]) + if not row: + raise KeyError( + f"no cleaning memory for dataset_id {dataset_id!r} in {str(path)!r}") + return str(row[0]) + rows = conn.execute( + "SELECT dataset_id, payload FROM cleaning_memory ORDER BY dataset_id").fetchall() + if not rows: + raise KeyError(f"no cleaning memory found in {str(path)!r}") + if len(rows) > 1: + ids = ", ".join(repr(str(r[0])) for r in rows) + raise ValueError( + f"{str(path)!r} holds {len(rows)} cleaning memories ({ids}); " + "pass dataset_id= to choose one") + return str(rows[0][1]) finally: conn.close() diff --git a/tests/test_memory_learning_persistence.py b/tests/test_memory_learning_persistence.py new file mode 100644 index 00000000..0e3db525 --- /dev/null +++ b/tests/test_memory_learning_persistence.py @@ -0,0 +1,358 @@ +"""Persistence regressions for cleaning memory and learning profiles. + +Covers #306 (explicit SQLite memory selection, no file creation on a missing +path), #308 (merged profiles keep their source schema), #309 (NaN decision +columns and NaN-free memory JSON) and #311 (strict ``.fdprofile`` manifests). +""" + +from __future__ import annotations + +import dataclasses +import hashlib +import io +import json +import math +import zipfile +from pathlib import Path + +import numpy as np +import pandas as pd +import pytest + +import freshdata as fd +from freshdata.learning import ProfileFormatError +from freshdata.learning.profile import LearningProfile, load_profile, save_profile +from freshdata.learning.replay import _profile_schema, check_profile_drift +from freshdata.learning.types import ExampleBank, ExamplePair +from freshdata.memory import CleaningMemory, _normalize_decisions + + +def _strict_json(text: str) -> object: + def _reject(token: str) -> object: + raise AssertionError(f"non-standard JSON token {token}") + + return json.loads(text, parse_constant=_reject) + + +# --------------------------------------------------------------------------- +# #306 load_cleaning_memory on SQLite stores +# --------------------------------------------------------------------------- + + +@pytest.fixture +def two_memory_store(tmp_path: Path) -> Path: + df = pd.DataFrame({"a": [1, 2]}) + store = tmp_path / "memory.sqlite" + fd.learn_cleaning_memory(df, [], "sales", roles={}).to_json(str(store)) + fd.learn_cleaning_memory(df, [], "crm", roles={}).to_json(str(store)) + return store + + +def test_multi_memory_store_requires_dataset_id(two_memory_store: Path) -> None: + with pytest.raises(ValueError, match="dataset_id") as info: + fd.load_cleaning_memory(str(two_memory_store)) + assert "'crm'" in str(info.value) + assert "'sales'" in str(info.value) + + +def test_multi_memory_store_selects_dataset_id(two_memory_store: Path) -> None: + assert fd.load_cleaning_memory(str(two_memory_store), dataset_id="crm").dataset_id == "crm" + assert fd.load_cleaning_memory(two_memory_store, dataset_id="sales").dataset_id == "sales" + + +def test_unknown_dataset_id_raises_key_error(two_memory_store: Path) -> None: + with pytest.raises(KeyError, match="billing"): + fd.load_cleaning_memory(str(two_memory_store), dataset_id="billing") + + +def test_single_memory_store_loads_without_dataset_id(tmp_path: Path) -> None: + store = tmp_path / "one.db" + fd.learn_cleaning_memory(pd.DataFrame({"a": [1]}), [], "only", roles={}).to_json(str(store)) + assert fd.load_cleaning_memory(str(store)).dataset_id == "only" + + +def test_store_path_with_uri_special_characters(tmp_path: Path) -> None: + store = tmp_path / "my store #1 %20.sqlite" + fd.learn_cleaning_memory(pd.DataFrame({"a": [1]}), [], "odd", roles={}).to_json(str(store)) + assert fd.load_cleaning_memory(str(store)).dataset_id == "odd" + + +@pytest.mark.parametrize("name", ["typo.sqlite", "typo.db", "typo.json"]) +def test_missing_path_raises_without_creating_file(tmp_path: Path, name: str) -> None: + target = tmp_path / name + with pytest.raises(FileNotFoundError): + fd.load_cleaning_memory(str(target)) + assert not target.exists() + + +def test_json_memory_dataset_id_mismatch(tmp_path: Path) -> None: + path = tmp_path / "mem.json" + fd.learn_cleaning_memory(pd.DataFrame({"a": [1]}), [], "crm", roles={}).to_json(str(path)) + assert fd.load_cleaning_memory(str(path), dataset_id="crm").dataset_id == "crm" + with pytest.raises(KeyError, match="sales"): + fd.load_cleaning_memory(str(path), dataset_id="sales") + + +# --------------------------------------------------------------------------- +# #308 merged LearningProfile drift +# --------------------------------------------------------------------------- + + +def _email_pair() -> tuple[pd.DataFrame, pd.DataFrame]: + emails = ["a@x.com", "B@Y.COM ", "c@z.org", "d@w.net"] * 5 + messy = pd.DataFrame( + { + "id": range(20), + "email": emails, + "status": ["pend-ing"] * 10 + ["done"] * 10, + "n": range(20), + } + ) + clean = pd.DataFrame( + { + "id": range(20), + "email": [e.strip().lower() for e in emails], + "status": ["pending"] * 10 + ["done"] * 10, + "n": range(20), + } + ) + return messy, clean + + +def test_merged_profile_replays_without_drift_on_training_frame() -> None: + messy, clean = _email_pair() + profile = fd.learn(messy, clean, key="id", min_support=2) + merged = profile.merge(profile) + + _, parent_report = fd.clean(messy, profile=profile, return_report=True, verbose=False) + _, merged_report = fd.clean(messy, profile=merged, return_report=True, verbose=False) + assert parent_report.warnings == [] + assert merged_report.warnings == [] + + gate = check_profile_drift(messy, merged) + assert gate.severity == "none" + assert set(gate.compatible_columns) == {"id", "email", "status", "n"} + assert merged.audit().alignment["source_schema"] == profile.audit().alignment["source_schema"] + + +def test_merged_source_schema_is_union_with_base_precedence() -> None: + messy, clean = _email_pair() + left = fd.learn(messy, clean, key="id", min_support=2) + wider_messy = messy.assign(extra=[1.5] * 20, n=messy["n"].astype("float64")) + wider_clean = clean.assign(extra=[1.5] * 20, n=clean["n"].astype("float64")) + right = fd.learn(wider_messy, wider_clean, key="id", min_support=2) + + union = left.merge(right).audit().alignment["source_schema"] + assert union["extra"] == "float64" + assert union["n"] == "int64" # self wins a dtype disagreement + other_first = left.merge(right, strategy="prefer_other").audit().alignment["source_schema"] + assert other_first["n"] == "float64" + + +def test_memory_signature_fallback_maps_coarse_types() -> None: + messy, clean = _email_pair() + profile = fd.learn(messy, clean, key="id", min_support=2) + memory = CleaningMemory( + dataset_id="sig", + signature={ + "columns": {"id": "integer", "email": "text", "status": "text", "n": "integer"} + }, + ) + bare = dataclasses.replace(profile, audit_info=None, memory=memory) + assert _profile_schema(bare) == { + "id": "int64", + "email": "object", + "status": "object", + "n": "int64", + } + gate = check_profile_drift(messy, bare) + assert gate.severity == "none", gate.reasons + assert "status" in gate.compatible_columns + + +# --------------------------------------------------------------------------- +# #309 NaN decision columns +# --------------------------------------------------------------------------- + + +def test_csv_round_tripped_decisions_replay_table_level_steps() -> None: + df = pd.DataFrame({"a": [1, 1, 2], "b": ["x", "x", "y"]}) + _, report = fd.clean(df, return_report=True, verbose=False) + reviewed = pd.read_csv(io.StringIO(report.to_frame().to_csv(index=False))) + + from_report = fd.learn_cleaning_memory(df, report, "d", roles={}) + from_csv = fd.learn_cleaning_memory(df, reviewed, "d", roles={}) + + def replayed(memory: CleaningMemory) -> list[str]: + _, rep = fd.clean(df, memory=memory, return_report=True, verbose=False) + return [a.step for a in rep.actions if a.memory_influenced and a.step != "memory"] + + assert "drop_duplicates" in replayed(from_report) + assert replayed(from_csv) == replayed(from_report) + assert all(d["column"] is None or isinstance(d["column"], str) for d in from_csv.accepted) + + +def test_missing_decision_columns_normalise_to_none() -> None: + rows = pd.DataFrame( + { + "column": [float("nan"), "", " ", None, "amount"], + "step": ["drop_duplicates", "a", "b", "c", "d"], + } + ) + accepted, rejected = _normalize_decisions(rows) + assert rejected == [] + assert [d["column"] for d in accepted] == [None, None, None, None, "amount"] + listed, _ = _normalize_decisions([{"column": float("nan"), "step": "x"}]) + assert listed[0]["column"] is None + + +def test_memory_json_never_emits_nan_tokens(tmp_path: Path) -> None: + df = pd.DataFrame({"a": [1, 1, 2]}) + reviewed = pd.DataFrame( + {"column": [float("nan")], "step": ["drop_duplicates"], "description": [float("nan")]} + ) + memory = fd.learn_cleaning_memory( + df, reviewed, "d", roles={}, thresholds={"outlier_factor": float("inf")} + ) + payload = _strict_json(memory.to_json()) + assert isinstance(payload, dict) + assert payload["thresholds"]["outlier_factor"] is None + assert payload["accepted"][0]["column"] is None + assert payload["accepted"][0]["description"] is None + + path = tmp_path / "mem.json" + memory.to_json(str(path)) + _strict_json(path.read_text(encoding="utf-8")) + assert fd.load_cleaning_memory(str(path)).accepted[0]["column"] is None + + +# --------------------------------------------------------------------------- +# #311 strict .fdprofile manifests +# --------------------------------------------------------------------------- + + +@pytest.fixture +def good_profile(tmp_path: Path) -> tuple[Path, dict[str, bytes]]: + messy = pd.DataFrame({"id": range(20), "s": ["pend-ing"] * 10 + ["done"] * 10}) + clean = pd.DataFrame({"id": range(20), "s": ["pending"] * 10 + ["done"] * 10}) + path = tmp_path / "good.fdprofile" + save_profile(fd.learn(messy, clean, key="id", min_support=2), path) + with zipfile.ZipFile(path) as archive: + members = {name: archive.read(name) for name in archive.namelist()} + return path, members + + +def _write(target: Path, members: dict[str, bytes], replace: dict[str, bytes]) -> Path: + with zipfile.ZipFile(target, "w") as archive: + for name, payload in members.items(): + archive.writestr(name, replace.get(name, payload)) + return target + + +@pytest.mark.parametrize( + "manifest", + [ + b'{"profile_ver', + b"{}", + b"[]", + b"\xff\xfe", + b'{"profile_version": "1.0", "compartments": 3}', + ], + ids=["truncated", "empty-object", "array", "not-utf8", "bad-field-type"], +) +def test_bad_manifest_raises_profile_format_error( + good_profile: tuple[Path, dict[str, bytes]], tmp_path: Path, manifest: bytes +) -> None: + _, members = good_profile + bad = _write(tmp_path / "bad.fdprofile", members, {"manifest.json": manifest}) + with pytest.raises(ProfileFormatError) as info: + load_profile(bad) + assert info.value.__cause__ is not None + + +def test_missing_member_hashes_fail_closed( + good_profile: tuple[Path, dict[str, bytes]], tmp_path: Path +) -> None: + _, members = good_profile + manifest = dict(json.loads(members["manifest.json"]), member_hashes={}) + bad = _write( + tmp_path / "bad.fdprofile", + members, + {"manifest.json": json.dumps(manifest).encode(), "value_maps.json": b'{"value_maps": {}}'}, + ) + with pytest.raises(ProfileFormatError, match="no member hash"): + load_profile(bad) + + +def test_one_missing_member_hash_fails_closed( + good_profile: tuple[Path, dict[str, bytes]], tmp_path: Path +) -> None: + _, members = good_profile + manifest = json.loads(members["manifest.json"]) + del manifest["member_hashes"]["rules.json"] + bad = _write( + tmp_path / "bad.fdprofile", members, {"manifest.json": json.dumps(manifest).encode()} + ) + with pytest.raises(ProfileFormatError, match="rules.json"): + load_profile(bad) + + +def test_malformed_hashed_member_raises_profile_format_error( + good_profile: tuple[Path, dict[str, bytes]], tmp_path: Path +) -> None: + _, members = good_profile + rules = b"[]" + manifest = json.loads(members["manifest.json"]) + manifest["member_hashes"]["rules.json"] = hashlib.sha256(rules).hexdigest() + bad = _write( + tmp_path / "bad.fdprofile", + members, + {"manifest.json": json.dumps(manifest).encode(), "rules.json": rules}, + ) + with pytest.raises(ProfileFormatError, match="malformed"): + load_profile(bad) + + +def test_saved_profile_round_trips(good_profile: tuple[Path, dict[str, bytes]]) -> None: + path, members = good_profile + loaded = load_profile(path) + assert set(loaded.manifest.member_hashes) == set(members) - {"manifest.json"} + + +def _profile_with_vectors(tmp_path: Path) -> tuple[LearningProfile, Path]: + messy = pd.DataFrame({"id": range(20), "s": ["pend-ing"] * 10 + ["done"] * 10}) + clean = pd.DataFrame({"id": range(20), "s": ["pending"] * 10 + ["done"] * 10}) + profile = fd.learn(messy, clean, key="id", min_support=2) + profile.examples = ExampleBank( + examples=[ExamplePair("s", "odd", "even", "unexplained", 1, False)], + vectors_path="examples_vectors.npz", + embedding_model_id="test-model", + masked=False, + ) + profile.vectors = np.arange(8, dtype="float32").reshape(2, 4) + path = tmp_path / "vectors.fdprofile" + save_profile(profile, path) + return profile, path + + +def test_profile_with_vectors_round_trips(tmp_path: Path) -> None: + profile, path = _profile_with_vectors(tmp_path) + loaded = load_profile(path) + assert "examples_vectors.npz" in loaded.manifest.member_hashes + assert loaded.profile_id == profile.profile_id + assert loaded.vectors is not None + assert np.allclose(loaded.vectors, profile.vectors) + assert not any(math.isnan(v) for v in loaded.vectors.ravel()) + + +def test_unhashed_vectors_member_fails_closed(tmp_path: Path) -> None: + _, path = _profile_with_vectors(tmp_path) + with zipfile.ZipFile(path) as archive: + members = {name: archive.read(name) for name in archive.namelist()} + manifest = json.loads(members["manifest.json"]) + del manifest["member_hashes"]["examples_vectors.npz"] + bad = _write( + tmp_path / "bad.fdprofile", members, {"manifest.json": json.dumps(manifest).encode()} + ) + with pytest.raises(ProfileFormatError, match="examples_vectors.npz"): + load_profile(bad)