Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 28 additions & 1 deletion src/freshdata/learning/merge.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -294,14 +295,40 @@ 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,
)
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",
Expand Down
107 changes: 71 additions & 36 deletions src/freshdata/learning/profile.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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:
Expand All @@ -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 -------------------------------------------------------

Expand Down Expand Up @@ -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])
Expand Down
13 changes: 12 additions & 1 deletion src/freshdata/learning/replay.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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 {}


Expand Down
86 changes: 70 additions & 16 deletions src/freshdata/memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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 -----------------------------
Expand All @@ -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()

Expand Down
Loading
Loading