diff --git a/pyproject.toml b/pyproject.toml index 6b6c02a1f..cc72bac22 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -352,13 +352,13 @@ safe-synthesizer = "nemo_safe_synthesizer.cli.cli:cli" # that appear unused when the package is installed. Suppress the noise. unused-ignore-comment = "ignore" + [tool.ty.environment] extra-paths = ["typings"] [tool.ty.src] # Below is a list of excluded directories from ty typechecks. exclude = [ - ".uv_cache", # Cache Dir in CI - "./docs/**/*.ipynb", - "./src/nemo_safe_synthesizer/pii_replacer/", + "./uv-cache", + "./docs/**/*.ipynb", ] diff --git a/src/nemo_safe_synthesizer/data_processing/records/fragment.py b/src/nemo_safe_synthesizer/data_processing/records/fragment.py index 655fa8603..d799f3346 100644 --- a/src/nemo_safe_synthesizer/data_processing/records/fragment.py +++ b/src/nemo_safe_synthesizer/data_processing/records/fragment.py @@ -14,9 +14,9 @@ import time import uuid from collections import defaultdict -from dataclasses import dataclass +from dataclasses import dataclass, field from datetime import datetime -from typing import Any +from typing import Any, NotRequired, TypeAlias, TypedDict from ...pii_replacer.ner.entity import Score from ...pii_replacer.ner.predictor import NERPrediction @@ -32,6 +32,62 @@ class MetadataError(Exception): E2F = "fields_by_entity" +class NERRawPredictionPayload(TypedDict): + """Raw ``NERPrediction.as_dict`` payload consumed by metadata helpers.""" + + text: str + start: int + end: int + label: str + source: str + score: float | None + field: NotRequired[str | None] + value_path: NotRequired[tuple[str | int, ...] | list[str | int] | None] + substring_match: NotRequired[bool | None] + + +class NERFieldLabelPayload(TypedDict): + """Per-field label entry in the NER model metadata payload.""" + + start: int + end: int + label: str + score: float | None + source: str + text: str + + +NERMetadataFieldsPayload: TypeAlias = dict[str, dict[str, dict[str, list[NERFieldLabelPayload]]]] + + +class NEREntityMapPayload(TypedDict): + """Score-bucketed entity summary in the NER model metadata payload.""" + + score_high: list[str] + score_med: list[str] + score_low: list[str] + fields_by_entity: dict[str, list[str]] + + +class NERMetadataPayload(TypedDict): + """API-facing NER metadata payload.""" + + record_id: str + fields: NERMetadataFieldsPayload + entities: NEREntityMapPayload + received_at: str + + +NERRecordPayload: TypeAlias = dict[str, Any] + + +class NERApiResponseRow(TypedDict): + """One API response row for dict-based NER model predictions.""" + + data: NERRecordPayload + model_metadata: NERMetadataPayload + + @dataclass class Metadata: """Merged record metadata aggregated from one or more ``MetadataFragment`` objects. @@ -43,18 +99,23 @@ class Metadata: record_id: str - fields: dict + fields: NERMetadataFieldsPayload """Nested dict of per-field, per-fragment metadata.""" - entities: dict + entities: NEREntityMapPayload """Entity map produced by ``predictions_to_dict``.""" received_at: str """ISO-8601 timestamp of the earliest fragment.""" - def as_dict(self): + def as_dict(self) -> NERMetadataPayload: """Serialize to a plain dictionary.""" - return self.__dict__ + return { + "record_id": self.record_id, + "fields": self.fields, + "entities": self.entities, + "received_at": self.received_at, + } @dataclass @@ -75,6 +136,7 @@ class MetadataFragment: fragment_ts: str fragment_epoch: float fragment_name: str + fields: dict[str, dict[str, list[NERFieldLabelPayload]]] = field(init=False) def __post_init__(self): self.fields = defaultdict(lambda: defaultdict(list)) @@ -84,7 +146,12 @@ def fragment_datetime(self) -> datetime: """Fragment creation time as a ``datetime`` object.""" return datetime.fromtimestamp(self.fragment_epoch) - def add_field_data(self, field_name: str, metadata_type: str, field_data: dict | list): + def add_field_data( + self, + field_name: str, + metadata_type: str, + field_data: NERFieldLabelPayload | list[NERFieldLabelPayload], + ) -> None: """Append metadata entries for a field. Args: @@ -97,12 +164,10 @@ def add_field_data(self, field_name: str, metadata_type: str, field_data: dict | """ if isinstance(field_data, list): self.fields[field_name][metadata_type].extend(field_data) - elif isinstance(field_data, dict): - self.fields[field_name][metadata_type].append(field_data) else: - raise TypeError("field_data must be a dict or list, got ", type(field_data)) + self.fields[field_name][metadata_type].append(field_data) - def as_dict(self): + def as_dict(self) -> dict[str, Any]: """Serialize to a plain dictionary.""" return self.__dict__ @@ -126,15 +191,29 @@ def merge_fragments(*fragments, ts: str | None = None) -> Metadata: else: record_id = fragments[0].record_id - # todo(dn): there might be a better way to build up this object - merged_fragment = defaultdict(lambda: defaultdict(lambda: defaultdict(list))) + merged_fragment: dict[str, dict[str, dict[str, list[NERFieldLabelPayload]]]] = defaultdict( + lambda: defaultdict(lambda: defaultdict(list)) + ) ts = ts or min([f.fragment_datetime for f in fragments]).isoformat() + "Z" fragment: MetadataFragment for fragment in fragments: for field_name, field_data in fragment.fields.items(): for meta_type, meta_data in field_data.items(): merged_fragment[field_name][fragment.fragment_name][meta_type].extend(meta_data) - return Metadata(record_id=record_id, fields=merged_fragment, received_at=ts, entities={}) + fields: NERMetadataFieldsPayload = { + field_name: { + fragment_name: {meta_type: list(meta_data) for meta_type, meta_data in fragment_data.items()} + for fragment_name, fragment_data in field_data.items() + } + for field_name, field_data in merged_fragment.items() + } + empty_entities: NEREntityMapPayload = { + SCORE_HIGH: [], + SCORE_MED: [], + SCORE_LOW: [], + E2F: {}, + } + return Metadata(record_id=record_id, fields=fields, received_at=ts, entities=empty_entities) def fragment_for_record(record_id: str, fragment_name: str) -> MetadataFragment: @@ -154,7 +233,7 @@ def predictions_to_dict( *, high_score: float = Score.HIGH, med_score: float = Score.MED, -) -> tuple[dict, dict]: +) -> tuple[dict[str, list[NERFieldLabelPayload]], NEREntityMapPayload]: """Aggregate NER predictions into per-field results and an entity map. Groups predictions by field and builds a score-bucketed entity map:: @@ -174,13 +253,11 @@ def predictions_to_dict( Returns: A tuple of (predictions_by_field, entity_map). """ - entity_map: dict[str, Any] = { - SCORE_HIGH: set(), - SCORE_MED: set(), - SCORE_LOW: set(), - E2F: defaultdict(set), - } - predictions_by_key = defaultdict(list) + high_entities: set[str] = set() + medium_entities: set[str] = set() + low_entities: set[str] = set() + fields_by_entity: dict[str, set[str]] = defaultdict(set) + predictions_by_key: dict[str, list[NERFieldLabelPayload]] = defaultdict(list) for prediction in predictions: if prediction.field is None: continue @@ -198,20 +275,22 @@ def predictions_to_dict( # no score is emitted. Predictions here could be # hit or miss so we throw it into medium if prediction.score is None: - entity_map[SCORE_MED].add(prediction.label) + medium_entities.add(prediction.label) elif prediction.score >= high_score: - entity_map[SCORE_HIGH].add(prediction.label) + high_entities.add(prediction.label) elif prediction.score >= med_score: - entity_map[SCORE_MED].add(prediction.label) + medium_entities.add(prediction.label) else: - entity_map[SCORE_LOW].add(prediction.label) - entity_map[E2F][prediction.label].add(prediction.field) + low_entities.add(prediction.label) + fields_by_entity[prediction.label].add(prediction.field) for _, preds in predictions_by_key.items(): preds.sort(key=lambda p: p["start"]) - for level in (SCORE_HIGH, SCORE_MED, SCORE_LOW): - entity_map[level] = list(entity_map[level]) - for entity, _set in entity_map[E2F].items(): - entity_map[E2F][entity] = list(_set) + entity_map: NEREntityMapPayload = { + SCORE_HIGH: list(high_entities), + SCORE_MED: list(medium_entities), + SCORE_LOW: list(low_entities), + E2F: {entity: list(fields) for entity, fields in fields_by_entity.items()}, + } return predictions_by_key, entity_map @@ -219,7 +298,7 @@ def fragment_from_ner_predictions( fragment_name: str, predictions: list[NERPrediction], record_id: str, -) -> tuple[MetadataFragment, dict]: +) -> tuple[MetadataFragment, NEREntityMapPayload]: """Build a ``MetadataFragment`` and entity map from NER predictions. Args: @@ -244,9 +323,9 @@ def fragment_from_ner_predictions( return fragment, ent_map -def build_ner_metadata(preds: list[dict]) -> Metadata: - """Construct a ``Metadata`` object from raw prediction dicts.""" - ner_preds = [NERPrediction.from_dict(p) for p in preds] +def build_ner_metadata(preds: list[NERRawPredictionPayload]) -> NERMetadataPayload: + """Construct an API-facing metadata payload from raw prediction dicts.""" + ner_preds = [NERPrediction.from_dict(dict(p)) for p in preds] fragment, ent_map = fragment_from_ner_predictions( "ner", ner_preds, @@ -257,7 +336,11 @@ def build_ner_metadata(preds: list[dict]) -> Metadata: return meta.as_dict() -def create_ner_api_response(records: list[dict], predictions: list[list[dict]], pure_dict: bool = False) -> list[dict]: +def create_ner_api_response( + records: list[NERRecordPayload], + predictions: list[list[NERRawPredictionPayload]], + pure_dict: bool = False, +) -> list[NERApiResponseRow]: """Build an API-compatible list of ``{data, model_metadata}`` dicts. Args: @@ -268,10 +351,20 @@ def create_ner_api_response(records: list[dict], predictions: list[list[dict]], Returns: List of dicts, each containing ``data`` and ``model_metadata`` keys. """ - out = [ + out: list[NERApiResponseRow] = [ {"data": record, "model_metadata": build_ner_metadata(prediction)} for record, prediction in zip(records, predictions) ] if pure_dict: - return json.loads(json.dumps(out)) + data_rows = json.loads(json.dumps([row["data"] for row in out])) + if not isinstance(data_rows, list): + raise TypeError("expected JSON round-trip to preserve response row list") + rows: list[NERApiResponseRow] = [] + for data, row in zip(data_rows, out): + if not isinstance(data, dict): + raise TypeError("expected JSON round-trip to preserve response data dictionaries") + record: NERRecordPayload = {str(key): value for key, value in data.items()} + response_row: NERApiResponseRow = {"data": record, "model_metadata": row["model_metadata"]} + rows.append(response_row) + return rows return out diff --git a/src/nemo_safe_synthesizer/pii_replacer/data_editor/detect.py b/src/nemo_safe_synthesizer/pii_replacer/data_editor/detect.py index 87342781e..600506f10 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/data_editor/detect.py +++ b/src/nemo_safe_synthesizer/pii_replacer/data_editor/detect.py @@ -11,7 +11,7 @@ from itertools import chain, islice from time import monotonic from timeit import default_timer as timer -from typing import Optional +from typing import Any, Optional import json_repair import pandas as pd @@ -19,6 +19,7 @@ from gliner import GLiNER from openai import OpenAI from pydantic import ConfigDict, TypeAdapter, ValidationError +from typing_extensions import override from ...observability import get_logger from ...utils import hf_offline_enabled @@ -260,6 +261,9 @@ def classify_columns( if not formatted_prompt: return {} + if client is None: + raise RuntimeError("InferenceAPI classifier not initialized.") + llm_start = timer() response = client.chat.completions.create( model=DefaultLLMConfig.config_id(), @@ -271,6 +275,10 @@ def classify_columns( max_tokens=DefaultLLMConfig.MAX_OUTPUT_TOKENS, ) entities_str = response.choices[0].message.content + if entities_str is None: + on_validation_error() + return {} + llm_elapsed = timer() - llm_start logger.info( f"LLM column classification took {llm_elapsed} seconds.", @@ -285,17 +293,18 @@ def classify_columns( return {col: ent if ent in entities else UNKNOWN_ENTITY for col, ent in col_entities.items()} -def sample_columns(df: pd.DataFrame, num_samples: int, random_state: Optional[int] = None) -> dict[str, pd.Series]: +def sample_columns( + df: pd.DataFrame, num_samples: int | None, random_state: Optional[int] = None +) -> dict[str, list[str]]: """Sample up to ``num_samples`` unique values per non-empty column for classification prompts.""" nonempty_columns = df.dropna(axis="columns", how="all").columns - col_samples = {} + col_samples: dict[str, list[str]] = {} for col in nonempty_columns: filtered = df[col][df[col].apply(lambda x: len(str(x)) < MAX_COL_STR_LEN)].dropna() if filtered.empty: continue - col_samples[col] = ( - filtered.sample(frac=1, random_state=random_state).value_counts().index[:num_samples].astype(str) - ) + samples = filtered.sample(frac=1, random_state=random_state).value_counts().index[:num_samples].astype(str) + col_samples[str(col)] = list(samples) return col_samples @@ -353,6 +362,7 @@ def detect_types(self, df: pd.DataFrame, entities: Optional[set[str]]) -> dict[s class ColumnClassifierNoop(ColumnClassifier): """No-op classifier that assigns ``UNKNOWN_ENTITY`` to every column.""" + @override def detect_types(self, df: pd.DataFrame, entities: Optional[set[str]] = None) -> dict[str, Optional[str]]: return {col: UNKNOWN_ENTITY for col in df.columns} @@ -386,11 +396,15 @@ def __init__(self): self._llm = None self._num_samples = None - def detect_types(self, df: pd.DataFrame, entities: set[str]) -> dict[str, Optional[str]]: + @override + def detect_types(self, df: pd.DataFrame, entities: Optional[set[str]] = None) -> dict[str, Optional[str]]: """Sample column data and call the inference API to classify columns into entity types.""" if self._llm is None: raise Exception("InferenceAPI classifier not initialized. Use get_classifier() method.") + if entities is None: + entities = DEFAULT_ENTITIES + return classify_columns( df=df, entities=entities, @@ -488,9 +502,11 @@ def batch_update_cache(self, texts: list[str], entities: Optional[set[str]] = No class EntityExtractorNoop(EntityExtractor): """No-op extractor that returns no entities.""" + @override def extract_entity_values(self, text: str, entities: Optional[set[str]]) -> list[dict[str, str]]: return [] + @override def extract_ner_predictions(self, text: str, entities: Optional[set[str]]) -> list[NERPrediction]: return [] @@ -514,41 +530,51 @@ class EntityExtractorRegexp(EntityExtractor): _entity_types: set[str] - def pipeline_from_entities(self, entities: set[str]) -> Callable[[], Pipeline]: + def pipeline_from_entities(self, entities: Optional[set[str]]) -> Callable[[], Pipeline]: """Build a pipeline factory for the given entity set (or ``_entity_types`` if empty).""" if not entities: entities = self._entity_types predictor_filter = LabelSetPredictorFilter(entities) factory = NERFactory(regex_only=True) ner = factory.create(predictor_filter=predictor_filter) - return ner.pipeline_factory - - def _detect_entities(self, text: str, entities: set[str]) -> list[dict]: + if isinstance(ner, ner_mp.NERParallel): + return ner.pipeline_factory + pipeline = ner.pipeline + if pipeline is None: + raise RuntimeError("NER pipeline is not configured") + return lambda: pipeline + + def _detect_entities(self, text: str, entities: Optional[set[str]]) -> list[NERPrediction]: # Ensure text is a string - Jinja templates may pass non-string types (e.g., float/NaN) text = str(text) pipeline_factory = self.pipeline_from_entities(entities) predictor = ner_mp.NERParallel(pipeline_factory=pipeline_factory, num_proc=2) results = predictor.predict({"text": text}) - return results + if not isinstance(results, list): + return [] + return [result for result in results if isinstance(result, NERPrediction)] + @override def extract_entity_values(self, text: str, entities: Optional[set[str]] = None) -> list[dict[str, str]]: # _detect_entities already converts text to string detected = self._detect_entities(text, entities) return [{"entity": e.label, "value": e.text} for e in detected] + @override def extract_ner_predictions(self, text: str, entities: Optional[set[str]] = None) -> list[NERPrediction]: # _detect_entities already converts text to string return self._detect_entities(text, entities) @classmethod + @override def get_entity_extractor( cls, - clsfy_cfg: ClassifyConfig, + clsfy_config: ClassifyConfig, ) -> EntityExtractor: """Return a regex extractor with entity types from ``clsfy_cfg`` (or ``DEFAULT_ENTITIES``).""" entity_types = DEFAULT_ENTITIES - if clsfy_cfg.ner_entities: - entity_types = clsfy_cfg.ner_entities + if clsfy_config.ner_entities: + entity_types = clsfy_config.ner_entities self = cls() self._entity_types = entity_types return self @@ -569,12 +595,13 @@ class EntityExtractorGliner(EntityExtractor): _batch_mode_enabled: bool # Map (text sha hash, entity set) -> list of detected entities in GLiNER format - _entity_cache: dict[tuple, list] + _entity_cache: dict[tuple[int, tuple[str, ...]], list[dict[str, Any]]] @classmethod + @override def get_entity_extractor( cls, - clsfy_cfg: ClassifyConfig, + clsfy_config: ClassifyConfig, ) -> EntityExtractorGliner: """Load GLiNER model and return extractor configured from ``clsfy_cfg``.""" extractor = cls() @@ -586,28 +613,32 @@ def get_entity_extractor( ) extractor._model = GLiNER.from_pretrained( - clsfy_cfg.gliner_model, + clsfy_config.gliner_model, map_location=map_location, local_files_only=hf_offline_enabled(), ) entity_types = DEFAULT_ENTITIES - if clsfy_cfg.ner_entities: - entity_types = clsfy_cfg.ner_entities + if clsfy_config.ner_entities: + entity_types = clsfy_config.ner_entities extractor._entity_types = entity_types extractor._ner_threshold = 0.3 - if clsfy_cfg.ner_threshold is not None: - extractor._ner_threshold = clsfy_cfg.ner_threshold - extractor._batch_mode_enabled = clsfy_cfg.gliner_batch_mode_enabled - extractor._chunk_length = clsfy_cfg.gliner_batch_mode_chunk_length + if clsfy_config.ner_threshold is not None: + extractor._ner_threshold = clsfy_config.ner_threshold + extractor._batch_mode_enabled = clsfy_config.gliner_batch_mode_enabled + extractor._chunk_length = clsfy_config.gliner_batch_mode_chunk_length extractor._chunk_overlap = 128 if extractor._chunk_length <= extractor._chunk_overlap: extractor._chunk_overlap = 0 extractor._entity_cache = {} - extractor._batch_size = clsfy_cfg.gliner_batch_mode_batch_size + extractor._batch_size = clsfy_config.gliner_batch_mode_batch_size return extractor - def _predict_entities(self, text: str, entity_labels: list[str]) -> list[dict]: - predict_entities = getattr(self._model, "predict_entities", None) + def _predict_entities(self, text: str, entity_labels: list[str]) -> list[dict[str, Any]]: + model = self._model + if model is None: + return [] + + predict_entities = getattr(model, "predict_entities", None) if predict_entities is not None: return predict_entities( text, @@ -616,14 +647,18 @@ def _predict_entities(self, text: str, entity_labels: list[str]) -> list[dict]: flat_ner=False, ) - return self._model.batch_predict_entities( - [text], - entity_labels, - threshold=self._ner_threshold, - flat_ner=False, - )[0] + batch_predict_entities = getattr(model, "batch_predict_entities", None) + if batch_predict_entities is not None: + return batch_predict_entities( + [text], + entity_labels, + threshold=self._ner_threshold, + flat_ner=False, + )[0] + + raise AttributeError("GLiNER model has neither predict_entities nor batch_predict_entities") - def _batch_predict_entities(self, texts: list[str], entity_labels: list[str]) -> list[list[dict]]: + def _batch_predict_entities(self, texts: list[str], entity_labels: list[str]) -> list[list[dict[str, Any]]]: batch_predict_entities = getattr(self._model, "batch_predict_entities", None) if batch_predict_entities is not None: return batch_predict_entities( @@ -651,7 +686,7 @@ def _detect_entities_chunked( self, text: str, entity_labels: Optional[set[str]], - ) -> list[dict]: + ) -> list[dict[str, Any]]: """Detect entities from text using GLiNER with chunking; update ``column_report``. Chunks text to stay within model context; merges overlapping chunks. Returns @@ -668,12 +703,13 @@ def _detect_entities_chunked( if entity_labels is None: entity_labels = self._entity_types - if entity_labels is None: + model = self._model + if entity_labels is None or model is None: return [] gliner_entity_labels = sorted(entity_labels) entities_key = tuple(gliner_entity_labels) start = 0 - entities = [] + entities: list[dict[str, Any]] = [] # Occasionally text interpreted as type other than string by jinja text = str(text) nchunks = 0 @@ -723,17 +759,20 @@ def _detect_entities_chunked( return entities + @override def extract_entity_values(self, text: str, entities: Optional[set[str]] = None) -> list[dict[str, str]]: detected = self._detect_entities_chunked(text, entities) return [{"entity": e["label"], "value": e["text"]} for e in detected] + @override def extract_ner_predictions(self, text: str, entities: Optional[set[str]]) -> list[NERPrediction]: return [ NERPrediction(e["text"], e["start"], e["end"], e["label"], "GLiNER", 9.0) for e in self._detect_entities_chunked(text, entities) ] - def batch_update_cache(self, texts: list[str], entity_labels: Optional[set[str]] = None): + @override + def batch_update_cache(self, texts: list[str], entities: Optional[set[str]] = None): if not self._batch_mode_enabled: return @@ -750,11 +789,12 @@ def batch_update_cache(self, texts: list[str], entity_labels: Optional[set[str]] last_log = monotonic() chunks = [] - if entity_labels is None: - entity_labels = self._entity_types - if entity_labels is None: + if entities is None: + entities = self._entity_types + model = self._model + if entities is None or model is None: return - gliner_entity_labels = sorted(entity_labels) + gliner_entity_labels = sorted(entities) entities_key = tuple(gliner_entity_labels) for text in texts: text = str(text) @@ -778,8 +818,8 @@ def batch_update_cache(self, texts: list[str], entity_labels: Optional[set[str]] }, ) last_log = monotonic() - for idx, entities in enumerate(entities_lists): - self._entity_cache[(hash(batch[idx]), entities_key)] = entities + for idx, batch_entities in enumerate(entities_lists): + self._entity_cache[(hash(batch[idx]), entities_key)] = batch_entities class EntityExtractorMulti(EntityExtractor): @@ -787,6 +827,7 @@ class EntityExtractorMulti(EntityExtractor): extractors: list[EntityExtractor] + @override def extract_entity_values(self, text: str, entities: Optional[set[str]] = None) -> list[dict[str, str]]: """Return combined entity/value dicts from all sub-extractors.""" retval = [] @@ -794,6 +835,7 @@ def extract_entity_values(self, text: str, entities: Optional[set[str]] = None) retval += extractor.extract_entity_values(text, entities) return retval + @override def extract_ner_predictions(self, text: str, entities: Optional[set[str]] = None) -> list[NERPrediction]: """Return merged NER predictions from all sub-extractors.""" predictions = [] @@ -802,13 +844,15 @@ def extract_ner_predictions(self, text: str, entities: Optional[set[str]] = None return predictions @classmethod - def get_entity_extractor(cls, clsfy_cfg: ClassifyConfig) -> EntityExtractorMulti: + @override + def get_entity_extractor(cls, clsfy_config: ClassifyConfig) -> EntityExtractorMulti: """Return an empty composite; add extractors with ``add_entity_extractor``.""" self = cls() self.extractors = [] return self - def batch_update_cache(self, texts, entities: Optional[set[str]] = None): + @override + def batch_update_cache(self, texts: list[str], entities: Optional[set[str]] = None): for extractor in self.extractors: extractor.batch_update_cache(texts, entities) diff --git a/src/nemo_safe_synthesizer/pii_replacer/data_editor/edit.py b/src/nemo_safe_synthesizer/pii_replacer/data_editor/edit.py index 3b966d07c..2bd6f2318 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/data_editor/edit.py +++ b/src/nemo_safe_synthesizer/pii_replacer/data_editor/edit.py @@ -32,6 +32,25 @@ Rule = dict[str, Any] +def _as_str_list(value: Any) -> list[str]: + if value is None: + return [] + if isinstance(value, list): + return [str(item) for item in value] + if isinstance(value, pd.Index): + return [str(item) for item in value] + return [str(value)] + + +def _locked_columns(env: Environment) -> set[str]: + return set(_as_str_list(env.globals_config.get("lock_columns"))) + + +def _as_positions(value: Any) -> list[int]: + values = value if isinstance(value, list) else [value] + return [position for position in values if isinstance(position, int)] + + class TransformFnAccounting: """Tracks which transform functions or filters are applied to each column for reporting. @@ -50,22 +69,26 @@ def __init__(self, included_fns: list[str]): self.included_fns = set(included_fns) self.column_fns = defaultdict(set) - def update(self, column_names: str | Iterable[str], fns: str | set[str]) -> None: + def update(self, column_names: str | Iterable[str], fns: str | Iterable[str] | None) -> None: """Record that the given functions/filters were applied to the given columns. Args: column_names: Column name(s) to record; a single string or iterable of strings. fns: Name(s) of functions or filters applied; intersected with ``included_fns``. """ - if isinstance(fns, str): - fns = set([fns]) - fns &= self.included_fns - if not fns: - fns = {"jinja"} + if fns is None: + fn_set: set[str] = set() + elif isinstance(fns, str): + fn_set = {fns} + else: + fn_set = set(fns) + fn_set &= self.included_fns + if not fn_set: + fn_set = {"jinja"} if isinstance(column_names, str): column_names = [column_names] for column_name in column_names: - self.column_fns[column_name] |= fns + self.column_fns[column_name] |= fn_set @dataclass @@ -163,7 +186,7 @@ class Step: """ _env: Environment - _vars: dict[str, str | dict | list] + _vars: dict[str, Any] def do_make_template(self, template_str: str) -> Template: """Build a Jinja template from the string (may raise ``TemplateError``).""" @@ -176,7 +199,6 @@ def make_template(self, template_str: str) -> Template: except TemplateError as e: raise Exception( f"Error building jinja template '{template_str}': {e}", - error_id="param", ) def template_to_fnames(self, template_str: str) -> set[str]: @@ -196,7 +218,7 @@ def _render_column( **kwargs, ) -> str: """Render the full column as a single string from the template with ``column`` and ``vars``.""" - self._env.maybe_seed(column) + self._env.maybe_seed(column.to_json()) for k, v in kwargs.items(): setattr(column, k, v) return template.render(column=column, vars=self._vars, **kwargs) @@ -250,9 +272,10 @@ def _render_cell( if progress is not None: progress.log_throttled() - this = row[column["name"]] - self._env.maybe_seed(this) - self._env.entity_extractor.current_column = column["name"] + column_name = str(column["name"]) + this = row[column_name] + self._env.maybe_seed(str(this)) + self._env.entity_extractor.current_column = column_name if foreach: foreach_str = foreach.render( row=row, @@ -263,8 +286,6 @@ def _render_cell( **kwargs, ) - foreach_itr = None - try: foreach_itr = ast.literal_eval(foreach_str) except (ValueError, TypeError, SyntaxError): @@ -273,29 +294,28 @@ def _render_cell( except (json.JSONDecodeError, TypeError): foreach_itr = None - try: - iter(foreach_itr) - except TypeError: + if not isinstance(foreach_itr, Iterable): return f"[Error] '{foreach_str}' is not iterable built-in python type or JSON blob." + foreach_items = list(foreach_itr) else: - foreach_itr = [None] + foreach_items = [None] try: - cell = this - for foreach_item in foreach_itr: + cell: Any = this + for foreach_item in foreach_items: cell = template.render( row=row, index=row.name, column=column, this=cell, - items=foreach_itr, + items=foreach_items, item=foreach_item, vars=self._vars, **kwargs, ) if fnreport: - fnreport.update(column["name"], fn_names) + fnreport.update(column_name, fn_names) if progress is not None: progress.status.row_n += 1 progress.log_throttled() @@ -305,7 +325,7 @@ def _render_cell( except Exception as e: if fallback_template: if fnreport: - fnreport.update(column["name"], fallback_fn_names) + fnreport.update(column_name, fallback_fn_names) cell = self._render_cell( row, fallback_template, @@ -328,25 +348,26 @@ def _render_row(self, template: Template, row: pd.Series, **kwargs) -> str: def _add_columns(self, df: pd.DataFrame, rules: list[Rule]) -> None: """Insert new columns into ``df`` per rules (``name`` and optional ``position``).""" for col in rules: - name = col["name"] + name = str(col["name"]) position = col.get("position") if position is None: position = len(df.columns) - df.insert(position, name, None) + if isinstance(position, int): + df.insert(position, name, None) def _rename_columns(self, df: pd.DataFrame, rules: list[Rule]) -> None: """Rename columns per rules (name → value); skip columns in ``lock_columns``.""" - locked = set(self._env._env.globals["globals"].get("lock_columns") or []) + locked = _locked_columns(self._env) column_names = { - col["name"]: col["value"] + str(col["name"]): str(col["value"]) for col in rules - if "name" in col and col["name"] not in locked and col["value"] not in locked + if "name" in col and str(col["name"]) not in locked and str(col["value"]) not in locked } df.rename(columns=column_names, inplace=True) def _drop_rows(self, df: pd.DataFrame, rules: list[Rule]) -> None: """Drop rows for which each rule's condition template renders to ``True``.""" - conditions = [self.make_template(rule["condition"]) for rule in rules] + conditions = [self.make_template(str(rule["condition"])) for rule in rules] for condition in conditions: row_filter = df.apply(lambda row: self._render_row(condition, row), axis=1) df.drop(index=df.index[row_filter == "True"], inplace=True) @@ -367,34 +388,27 @@ def _parse_drop_columns_rule( position = rule.get("position") condition = rule.get("condition") rule_coltypes = rule.get("type") + condition_template: Template | None = None if column_name: - if isinstance(column_name, list): - columns = column_name - else: - columns = [column_name] + columns = _as_str_list(column_name) elif rule_entities: - if not isinstance(rule_entities, list): - rule_entities = [rule_entities] - columns = [col for col in df.columns if entities.get(col) in rule_entities] + entity_names = set(_as_str_list(rule_entities)) + columns = [str(col) for col in df.columns if entities.get(str(col)) in entity_names] elif rule_coltypes: - if not isinstance(rule_coltypes, list): - rule_coltypes = [rule_coltypes] - columns = [col for col in df.columns if column_types.get(col) in rule_coltypes] + type_names = set(_as_str_list(rule_coltypes)) + columns = [str(col) for col in df.columns if column_types.get(str(col)) in type_names] elif position is not None: - if isinstance(position, list): - position_list = position - else: - position_list = [position] + position_list = _as_positions(position) columns = [df.columns[pos] for pos in position_list] - elif condition: - columns = df.columns - condition = self.make_template(condition) + columns = [str(column) for column in columns] + elif isinstance(condition, str): + columns = [str(column) for column in df.columns] + condition_template = self.make_template(condition) else: raise Exception( f"column drop rule must contain one of name, entity, position, or condition. {rule}", - error_id="param", ) - return columns, condition + return columns, condition_template def _drop_columns( self, @@ -405,11 +419,11 @@ def _drop_columns( fnreport: TransformFnAccounting | None = None, ) -> None: """Drop columns per rules (by name/entity/type/position/condition); respect ``lock_columns``; update ``fnreport``.""" - locked = set(self._env._env.globals["globals"].get("lock_columns") or []) + locked = _locked_columns(self._env) try: for rule in rules: columns, condition_tmpl = self._parse_drop_columns_rule(rule, df, entities, column_types) - to_drop = [] + to_drop: Iterable[str] if condition_tmpl: column_properties = {} for position, column_name in enumerate(columns): @@ -424,8 +438,9 @@ def _drop_columns( lambda column: self._render_column(condition_tmpl, column, **column_properties[column.name]), axis=0, ) - colfilter[locked & set(columns)] = "False" - to_drop = df.columns[colfilter == "True"] + locked_columns = list(locked & set(columns)) + colfilter.loc[locked_columns] = "False" + to_drop = [str(column) for column in df.columns[colfilter == "True"]] df.drop(columns=to_drop, inplace=True) else: to_drop = list(set(columns) - locked) @@ -435,7 +450,6 @@ def _drop_columns( except KeyError as keyerr: raise Exception( f"Attempting to drop nonexistent column: {keyerr}", - error_id="param", ) def update_ner_cache(self, texts: pd.Series, entities: set[str] | None = None) -> None: @@ -454,34 +468,28 @@ def _parse_update_rows_rule( rule_entities = rule.get("entity") condition = rule.get("condition") rule_coltypes = rule.get("type") + condition_template: Template | None = None if column_name: - if isinstance(column_name, list): - columns = column_name - else: - columns = [column_name] + columns = _as_str_list(column_name) elif rule_entities: - if not isinstance(rule_entities, list): - rule_entities = [rule_entities] - columns = [col for col in df.columns if entities.get(col) in rule_entities] + entity_names = set(_as_str_list(rule_entities)) + columns = [str(col) for col in df.columns if entities.get(str(col)) in entity_names] elif rule_coltypes: - if not isinstance(rule_coltypes, list): - rule_coltypes = [rule_coltypes] - columns = [col for col in df.columns if column_types.get(col) in rule_coltypes] - elif condition: - columns = df.columns - condition = self.make_template(condition) + type_names = set(_as_str_list(rule_coltypes)) + columns = [str(col) for col in df.columns if column_types.get(str(col)) in type_names] + elif isinstance(condition, str): + columns = [str(column) for column in df.columns] + condition_template = self.make_template(condition) else: raise Exception( f"row update rule must contain one of name, entity, or condition. {rule}", - error_id="param", ) for column in columns: if column not in df.columns: raise Exception( f"The column '{column}' was not found. If you are adding a column and wish to access it, be sure to place the column.add rule in a step prior to the step accessing the column.", - error_id="param", ) - return columns, condition + return columns, condition_template def _update_rows( self, @@ -490,7 +498,7 @@ def _update_rows( entities: dict[str, str | None], column_types: dict[str, str | None], progress: ProgressLog, - fnreport: TransformFnAccounting, + fnreport: TransformFnAccounting | None = None, ) -> None: """Apply row-update rules to DataFrame cells; skip locked columns; update progress and fnreport. @@ -511,25 +519,26 @@ def _update_rows( Returns: None. The DataFrame is modified in place. """ - locked = set(self._env._env.globals["globals"].get("lock_columns") or []) + locked = _locked_columns(self._env) progress.status.update_rule_n_total = len(rules) for rule_n, rule in enumerate(rules): progress.status.update_rule_n = rule_n - progress.status.update_rule_description = rule.get("description") + progress.status.update_rule_description = str(rule.get("description") or "") columns, condition = self._parse_update_rows_rule(rule, df, entities, column_types) columns = [col for col in columns if col not in locked] - foreach = rule.get("foreach") - if foreach: - foreach = self.make_template(foreach) + foreach_value = rule.get("foreach") + foreach = self.make_template(foreach_value) if isinstance(foreach_value, str) else None - fns = self.template_to_fnames(rule["value"]) - value = self.make_template(rule["value"]) + value_text = str(rule["value"]) + fns = self.template_to_fnames(value_text) + value = self.make_template(value_text) fallback_value = rule.get("fallback_value") - fallback_fns = None - if fallback_value: + fallback_fns: set[str] | None = None + fallback_template: Template | None = None + if isinstance(fallback_value, str): fallback_fns = self.template_to_fnames(fallback_value) - fallback_value = self.make_template(fallback_value) + fallback_template = self.make_template(fallback_value) progress.status.column_n_total = len(columns) for position, column_name in enumerate(columns): column_properties = { @@ -561,7 +570,7 @@ def _update_rows( args=( value, column_properties, - fallback_value, + fallback_template, foreach, progress, fns, @@ -578,7 +587,7 @@ def execute( df: pd.DataFrame, entities: dict[str, str | None], column_types: dict[str, str | None], - step_config: dict[str, dict], + step_config: dict[str, Any], env: Environment, progress: ProgressLog, fnreport: TransformFnAccounting | None, @@ -647,7 +656,6 @@ def instantiate_vars(var_name: str, var_value: dict | list | str, step: Step, df # If it's valid jinja syntax but some other error occurred, assume user error. raise Exception( f"Error building jinja template for var '{var_name}': '{var_value}': {e}", - error_id="param", ) try: @@ -693,14 +701,15 @@ def _config_globals(self, entity_extractor: EntityExtractor | None) -> None: entity_extractor=entity_extractor, ) - def __init__(self, config: dict[str, dict], entity_extractor: EntityExtractor | None) -> None: + def __init__(self, config: dict[str, Any], entity_extractor: EntityExtractor | None = None) -> None: self.config = config self._config_globals(entity_extractor) @classmethod def load_yaml(cls, yaml_str: str) -> Editor: """Build an ``Editor`` from a YAML string (e.g. ``yaml.safe_load(yaml_str)``).""" - return cls(yaml.safe_load(yaml_str)) + config = yaml.safe_load(yaml_str) + return cls(config if isinstance(config, dict) else {}) def _process_df( self, diff --git a/src/nemo_safe_synthesizer/pii_replacer/data_editor/environment.py b/src/nemo_safe_synthesizer/pii_replacer/data_editor/environment.py index ec51569ee..2a89b518c 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/data_editor/environment.py +++ b/src/nemo_safe_synthesizer/pii_replacer/data_editor/environment.py @@ -9,7 +9,7 @@ from collections.abc import Callable from datetime import date, datetime, timedelta from functools import partial -from typing import Any, Optional +from typing import Any, Optional, cast import dateutil.parser import jinja2 @@ -279,6 +279,7 @@ class Environment: _env: SandboxedEnvironment _fake: Faker entity_extractor: EntityExtractor + globals_config: dict[str, Any] def __init__( self, @@ -291,13 +292,16 @@ def __init__( self.entity_extractor = entity_extractor else: self.entity_extractor = EntityExtractorNoop() + self.globals_config = globals_config or {} self._env = SandboxedEnvironment(loader=jinja2.BaseLoader()) self._fake = Faker(locale=locales, seed=seed) - self._env.globals["fake"] = self._fake - self._env.globals["globals"] = globals_config - self._env.globals["random"] = random - self._env.globals["re"] = re - self._env.globals["timedelta"] = timedelta + # Jinja templates intentionally expose arbitrary helper objects in globals. + template_globals = cast("dict[str, Any]", self._env.globals) + template_globals["fake"] = self._fake + template_globals["globals"] = self.globals_config + template_globals["random"] = random + template_globals["re"] = re + template_globals["timedelta"] = timedelta self._env.filters["hash"] = partial(sha256, str(seed)) self._env.filters["isna"] = pd.isna self._env.filters["fake"] = lambda faker_type: getattr(self._fake, faker_type)() @@ -327,7 +331,7 @@ def __init__( ) ) self._env.filters["fake_entities"] = lambda text, entities=None, on_error=None, extended=False: ( - entity_extractor.extract_and_replace_entities( + self.entity_extractor.extract_and_replace_entities( partial( fake_entities_fn, str(seed), diff --git a/src/nemo_safe_synthesizer/pii_replacer/data_editor/transform_test_utils.py b/src/nemo_safe_synthesizer/pii_replacer/data_editor/transform_test_utils.py index c28a2431a..21e2db13a 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/data_editor/transform_test_utils.py +++ b/src/nemo_safe_synthesizer/pii_replacer/data_editor/transform_test_utils.py @@ -6,10 +6,12 @@ from unittest.mock import MagicMock import pandas as pd +from typing_extensions import override from ..ner.ner import NERPrediction from .detect import ( UNKNOWN_ENTITY, + ClassifyConfig, ColumnClassifier, EntityExtractorGliner, IAPIClassifierConfig, @@ -29,20 +31,22 @@ class EntityExtractorMock(EntityExtractorGliner): extract_ner_predictions_called: bool @classmethod - def get_entity_extractor(cls, *args, **kwargs) -> EntityExtractorMock | None: + @override + def get_entity_extractor(cls, clsfy_config: ClassifyConfig) -> EntityExtractorMock: self = cls() self.extract_ner_predictions_called = False self._chunk_length = 512 self._chunk_overlap = 128 self._entity_cache = {} self._batch_size = 1000 - self._entity_types = ["test", "entity", "types"] + self._entity_types = {"test", "entity", "types"} self._model = MagicMock() self._ner_threshold = 0.9 self._batch_mode_enabled = False return self - def _detect_entities_chunked(self, text: str, entities: set[str] | None = None) -> list[dict]: + @override + def _detect_entities_chunked(self, text: str, entity_labels: set[str] | None = None) -> list[dict]: """Return fixed entity dicts for the known test string (first_name, ssn, nofake). Input is assumed to match the string used in transform tests; indices are @@ -54,6 +58,7 @@ def _detect_entities_chunked(self, text: str, entities: set[str] | None = None) {"label": "nofake", "text": "Unfake-able", "start": 45, "end": 56}, ] + @override def extract_ner_predictions(self, text: str, entities: set[str] | None) -> list[NERPrediction]: """Return fixed NER predictions and set ``extract_ner_predictions_called`` to ``True``.""" self.extract_ner_predictions_called = True @@ -86,7 +91,8 @@ def get_deployed_llm_classifier( classifier._num_samples = num_samples return classifier - def detect_types(self, df: pd.DataFrame, all_entities: set[str]) -> dict[str, str | None]: + @override + def detect_types(self, df: pd.DataFrame, entities: set[str] | None = None) -> dict[str, str | None]: """Return a hardcoded column-to-entity map for known test columns. Only columns present in the internal mapping and in ``all_entities`` @@ -100,7 +106,7 @@ def detect_types(self, df: pd.DataFrame, all_entities: set[str]) -> dict[str, st Returns: Map of column name to entity name (or ``UNKNOWN_ENTITY``). """ - entities: dict[str, str] = { + column_entities: dict[str, str] = { "AddressLine1": "street_address", "AddressLine1a": "street_address", "AddressLine1b": "street_address", @@ -119,6 +125,6 @@ def detect_types(self, df: pd.DataFrame, all_entities: set[str]) -> dict[str, st # columns = sample_columns(df, self._num_samples) # Always see something that isn't really there.. - all_entities = all_entities | {"ghost_entity"} - entities = {col: entities[col] for col in entities if entities[col] in all_entities} - return {col: entities.get(col, UNKNOWN_ENTITY) for col in df.columns} + allowed_entities = (entities or set()) | {"ghost_entity"} + configured_entities = {col: value for col, value in column_entities.items() if value in allowed_entities} + return {col: configured_entities.get(col, UNKNOWN_ENTITY) for col in df.columns} diff --git a/src/nemo_safe_synthesizer/pii_replacer/nemo_pii.py b/src/nemo_safe_synthesizer/pii_replacer/nemo_pii.py index 9710bcc16..4233c616c 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/nemo_pii.py +++ b/src/nemo_safe_synthesizer/pii_replacer/nemo_pii.py @@ -213,7 +213,8 @@ def _build_column_statistics( detected_values[entity_name] = entity_report.values elif field.entity is not None: # Non-text column with detected entity - detected_counts[field.entity] = field.entity_count + if field.entity_count is not None: + detected_counts[field.entity] = field.entity_count detected_values[field.entity] = set(field.entity_values) # Get transform functions for this column diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/custom.py b/src/nemo_safe_synthesizer/pii_replacer/ner/custom.py index 05e06c2e6..e1d2cb8ef 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/custom.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/custom.py @@ -45,7 +45,7 @@ score_map = {"low": Score.LOW, "med": Score.MED, "high": Score.HIGH} -PatternListType = list[re.Pattern] +PatternListType = list[re.Pattern[str]] def _str_to_pattern(data: str | list[str]) -> PatternListType: @@ -72,7 +72,7 @@ class CustomRegexPattern: header_skip: Optional[str | list[str]] = field(default_factory=list) # set after init - regex_compiled: re.Pattern = None + regex_compiled: re.Pattern[str] | None = None header_match_compiled: PatternListType = field(default_factory=list) header_skip_compiled: PatternListType = field(default_factory=list) @@ -80,9 +80,7 @@ def __post_init__(self): if self.score not in score_map: raise CustomPredictorError("score must be one of low, med, high") - self.regex_compiled = self.regex - if not isinstance(self.regex, re.Pattern): - self.regex_compiled = re.compile(str(self.regex)) + self.regex_compiled = re.compile(str(self.regex)) self._load_header_patterns() @@ -94,6 +92,8 @@ def _load_header_patterns(self): self.header_skip_compiled = _str_to_pattern(self.header_skip) def get_synthesizer_pattern(self): + if self.regex_compiled is None: + raise CustomPredictorError("regex pattern was not compiled") _score = score_map[self.score] return RegexPattern( pattern=self.regex_compiled, diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/datetime.py b/src/nemo_safe_synthesizer/pii_replacer/ner/datetime.py index 1cfe9f551..64d233bcf 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/datetime.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/datetime.py @@ -5,6 +5,8 @@ import itertools import math +import re +from collections.abc import Sequence from dataclasses import dataclass from datetime import datetime from typing import Optional @@ -12,6 +14,7 @@ from dateparser import parse from dateparser.date import get_date_from_timestamp from dateparser.search import search_dates +from typing_extensions import override from ...data_processing.records.json_record import JSONRecord from .entity import ( @@ -98,23 +101,28 @@ @dataclass class BaseContext(PredictorContext): - header_contexts: list - header_regexes: list = None - header_tokens: list = None + header_contexts: Sequence[str | re.Pattern[str]] + header_regexes: re.Pattern[str] | None = None + header_tokens: re.Pattern[str] | None = None def __post_init__(self): self.header_regexes, self.header_tokens = split_header_contexts(self.header_contexts) + def get_entity_label(self, match: tuple[str, datetime]) -> str: + raise NotImplementedError + @dataclass class DateTimeContext(BaseContext): - def get_entity_label(self, match: tuple): + @override + def get_entity_label(self, match: tuple[str, datetime]) -> str: return Entity.DATETIME.tag @dataclass class BirthDateContext(BaseContext): - def get_entity_label(self, _): + @override + def get_entity_label(self, match: tuple[str, datetime]) -> str: return Entity.BIRTH_DATE.tag @@ -169,14 +177,16 @@ class DateTime(Predictor): """Date date/time matcher.""" default_name: str = "datetime" + _context: BaseContext - def __init__(self, name: str = None): + def __init__(self, name: str | None = None): if name is None: name = self.default_name super().__init__(name) self._context = DateTimeContext(LABELS) - def evaluate(self, in_record: JSONRecord) -> list[NERPrediction]: + @override + def evaluate(self, in_data: JSONRecord) -> list[NERPrediction]: """ Given a single record determine if any entities are represented. @@ -187,8 +197,8 @@ def evaluate(self, in_record: JSONRecord) -> list[NERPrediction]: Returns: A list of entity predictions sorted by score. Top score is first entry in list. """ - result_set_by_field = [[] for _ in in_record.kv_pairs] - for field_matches, record_field in zip(result_set_by_field, in_record.kv_pairs): + result_set_by_field: list[list[NERPrediction]] = [[] for _ in in_data.kv_pairs] + for field_matches, record_field in zip(result_set_by_field, in_data.kv_pairs): # NOTE(jm): Changed to require header context no matter what, too many # FPs when looking in unstructured text if not self.header_has_context( diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/factory.py b/src/nemo_safe_synthesizer/pii_replacer/ner/factory.py index e7168a648..35c80f370 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/factory.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/factory.py @@ -10,6 +10,8 @@ from enum import StrEnum from typing import Optional +from typing_extensions import override + from ...data_processing.records.base import normalize_labels from ...observability import get_logger from .metadata import FieldLabelCondition @@ -53,6 +55,7 @@ class LabelSetPredictorFilter(PredictorFilter): def __init__(self, included_labels: set[str]): self._included_labels = normalize_labels(included_labels) + @override def should_include(self, predictor: Predictor) -> bool: if isinstance(predictor, RegexPredictor): # For Regex predictor, we filter based on ``entity.tag`` (that's what is visible to the user) @@ -101,7 +104,7 @@ def __init__( parallel: bool = True, use_nlp: bool = False, regex_only: bool = False, - ner_max_runtime_seconds: int = None, + ner_max_runtime_seconds: int | None = None, ): if custom_predictors is None: custom_predictors = [] @@ -111,6 +114,7 @@ def __init__( self._ner_pipeline_type = NERPipelineType.from_flags(use_nlp=use_nlp, regex_only=regex_only) self._ner_max_runtime_seconds = ner_max_runtime_seconds + @override def create( self, *, @@ -186,6 +190,7 @@ class StaticNERFactory(NERFactoryBase): def __init__(self, ner: NER | NERParallel): self._ner = ner + @override def create( self, *, diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/fasttext.py b/src/nemo_safe_synthesizer/pii_replacer/ner/fasttext.py index 5c5334882..126dab7be 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/fasttext.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/fasttext.py @@ -4,6 +4,7 @@ from __future__ import annotations from collections.abc import Callable, Iterable +from typing import Any from .models import ( ModelManifest, @@ -15,9 +16,29 @@ # TODO: Figure out import situation and resolve noqa: F821 exceptions through # ner/ directory. -spacy = None -dot = None -np = None +spacy: Any +dot: Any +np: Any +Doc: Any +Span: Any +try: + import numpy as _np + import spacy as _spacy # ty: ignore[unresolved-import] + from numpy import dot as _dot + from spacy.tokens import Doc as _Doc # ty: ignore[unresolved-import] + from spacy.tokens import Span as _Span # ty: ignore[unresolved-import] +except ImportError: + spacy = None + dot = None + np = None + Doc = None + Span = None +else: + spacy = _spacy + dot = _dot + np = _np + Doc = _Doc + Span = _Span manifest = ModelManifest( @@ -128,6 +149,8 @@ class FTEntityMatcher: @classmethod def factory(cls, model: ModelManifest = manifest) -> FTEntityMatcher: model_objects = get_cache_manager().resolve(manifest) + if model_objects is None: + raise RuntimeError(f"Could not resolve manifest {manifest}") return cls(**model_objects) def __init__(self, *, pos_neg_terms: dict, ft_word_vecs: dict, ft_ngram_vecs: dict): @@ -200,7 +223,7 @@ def get_ft_vec(self, word): Get FastText vector for a word. If it's OOV, gather the vectors for it's ngrams and average them """ - vec = [0] * 100 + vec: Any = [0] * 100 # If FT already knows about this word, go ahead and get it's vector if word in self.ft_word_vecs: vec = self.ft_word_vecs[word] @@ -217,7 +240,7 @@ def get_ft_vec(self, word): # for this OOV word for nh in ngram_hashes: vec += self.ft_ngram_vecs[nh] - vec = self.norm(vec / len(ngram_hashes)) + vec = self.norm(np.array(vec) / len(ngram_hashes)) return vec diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/helpers.py b/src/nemo_safe_synthesizer/pii_replacer/ner/helpers.py index 14f89d0b1..99b45546c 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/helpers.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/helpers.py @@ -3,9 +3,17 @@ from __future__ import annotations +from typing import Any + from .entity import Entity -spacy = None +spacy: Any +try: + import spacy as _spacy # ty: ignore[unresolved-import] +except ImportError: + spacy = None +else: + spacy = _spacy def entities_to_html(text: str, entities: list): @@ -15,6 +23,9 @@ def entities_to_html(text: str, entities: list): Returns: an HTML string """ + if spacy is None: + raise RuntimeError("spacy is not installed") + palette = [ "#7aecec", "#bfeeb7", diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/metadata.py b/src/nemo_safe_synthesizer/pii_replacer/ner/metadata.py index dfb10017e..1235f6247 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/metadata.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/metadata.py @@ -3,14 +3,14 @@ from __future__ import annotations -from dataclasses import asdict, dataclass +from dataclasses import dataclass from dataclasses import field as Field from enum import StrEnum from math import ceil -from typing import Optional +from typing import Any, Optional, TypedDict from ...data_processing.records.json_record import JSONRecord -from .ner import NER, PipelineResult +from .ner import NER, PipelineResult, Timings from .ner_mp import NERParallel @@ -19,6 +19,53 @@ class FieldAttribute(StrEnum): CATEGORICAL = "categorical" +class EntityMetadataPayload(TypedDict): + label: str + count: int + f_ratio: float + approx_cardinality: int + sources: list[str] + field_label_f_ratio: float + + +class TypeMetadataPayload(TypedDict): + type: str + count: int + + +class FieldMetadataPayload(TypedDict): + field: str + count: int + approx_cardinality: int + missing: int + pct_missing: float + pct_total_unique: float + s_score: float + entities: list[EntityMetadataPayload] + types: list[TypeMetadataPayload] + field_labels: list[str] + field_attributes: list[FieldAttribute] + + +class EntitySummaryPayload(TypedDict): + label: str + fields: list[str] + count: int + approx_distinct_count: int + sources: list[str] + + +class FieldsMetadataPayload(TypedDict): + fields: list[FieldMetadataPayload] + entities: list[EntitySummaryPayload] + + +class DatasetMetadataPayload(TypedDict): + project_record_count: int + total_field_count: int + data: FieldsMetadataPayload + + @dataclass(frozen=True) class EntityMetadata: label: str @@ -43,6 +90,16 @@ class EntityMetadata: This field is used to determine if an entity should be applied as a field_label in transformation pipelines.""" + def dict(self) -> EntityMetadataPayload: + return { + "label": self.label, + "count": self.count, + "f_ratio": self.f_ratio, + "approx_cardinality": self.approx_cardinality, + "sources": self.sources, + "field_label_f_ratio": self.field_label_f_ratio, + } + @dataclass(frozen=True) class TypeMetadata: @@ -55,6 +112,12 @@ class TypeMetadata: count: int """Number of times this type appeared in the values of a field.""" + def dict(self) -> TypeMetadataPayload: + return { + "type": self.type, + "count": self.count, + } + @dataclass(frozen=True) class FieldMetadata: @@ -101,8 +164,20 @@ class FieldMetadata: field_attributes: list[FieldAttribute] = Field(default_factory=list) """Attributes detected for this field.""" - def dict(self): - return asdict(self) + def dict(self) -> FieldMetadataPayload: + return { + "field": self.field, + "count": self.count, + "approx_cardinality": self.approx_cardinality, + "missing": self.missing, + "pct_missing": self.pct_missing, + "pct_total_unique": self.pct_total_unique, + "s_score": self.s_score, + "entities": [entity.dict() for entity in self.entities], + "types": [type_metadata.dict() for type_metadata in self.types], + "field_labels": self.field_labels, + "field_attributes": self.field_attributes, + } @dataclass(frozen=True) @@ -129,8 +204,14 @@ class EntitySummary: to the entity summary. """ - def dict(self) -> dict: - return asdict(self) + def dict(self) -> EntitySummaryPayload: + return { + "label": self.label, + "fields": self.fields, + "count": self.count, + "approx_distinct_count": self.approx_distinct_count, + "sources": self.sources, + } @dataclass(frozen=True) @@ -159,8 +240,15 @@ def add_field(self, field_metadata: FieldMetadata): def add_entity(self, entity_summary: EntitySummary): self.data.entities.append(entity_summary) - def to_dict(self): - return asdict(self) + def to_dict(self) -> DatasetMetadataPayload: + return { + "project_record_count": self.project_record_count, + "total_field_count": self.total_field_count, + "data": { + "fields": [field_metadata.dict() for field_metadata in self.data.fields], + "entities": [entity_summary.dict() for entity_summary in self.data.entities], + }, + } @dataclass(frozen=True) @@ -174,6 +262,46 @@ def explain(self, label: str) -> str: return f"At least {self.min_f_ratio * 100}% of all records were labeled with {label}" +class _DatasetMetadataTracker: + def __init__(self, field_label_condition: FieldLabelCondition | None = None): + self.field_label_condition = field_label_condition or FieldLabelCondition() + self._field_names: list[str] = [] + self._record_count = 0 + + def add_field_names(self, field_names: list[str]) -> None: + for field_name in field_names: + if field_name not in self._field_names: + self._field_names.append(field_name) + + def update_fields(self, records: list[JSONRecord]) -> None: + self._record_count += len(records) + + def update_entities(self, records: list[JSONRecord], record_labels: PipelineResult) -> None: + return None + + def get_snapshot(self) -> DatasetMetadata: + fields = [ + FieldMetadata( + field=field_name, + count=0, + approx_cardinality=0, + missing=self._record_count, + pct_missing=100.0 if self._record_count else 0.0, + pct_total_unique=0.0, + s_score=0.0, + ) + for field_name in self._field_names + ] + return DatasetMetadata( + project_record_count=self._record_count, + total_field_count=len(self._field_names), + data=FieldsMetadata(fields=fields), + ) + + def get_entity_detail(self, entity_label: str) -> dict[str, Any]: + return {} + + class MetadataService: """ Service that provides functionality to label records and also track model_metadata across whole dataset. @@ -184,10 +312,10 @@ class MetadataService: def __init__( self, ner: NER | NERParallel, - field_label_condition: FieldLabelCondition = None, + field_label_condition: FieldLabelCondition | None = None, ): self.ner = ner - self.dataset_metadata_tracker = _DatasetMetadataTracker(field_label_condition=field_label_condition) # noqa: F821 + self.dataset_metadata_tracker = _DatasetMetadataTracker(field_label_condition=field_label_condition) def add_field_names(self, field_names: list[str]): """ @@ -209,20 +337,31 @@ def predict( min_score: float = 0.0, timings_only: bool = False, include_labels: Optional[set[str]] = None, - ) -> PipelineResult: + ) -> PipelineResult | dict[str, Any]: # potential improvements here # - if a field is already classified as something on a field level -> do we skip doing NER on that field? + if timings_only: + timings = self.ner.predict( + records, + dict_result=True, + min_score=min_score, + timings_only=True, + include_labels=include_labels, + ) + if not isinstance(timings, Timings): + raise RuntimeError("NER timings result was not returned") + return timings.to_dict() + record_labels = self.ner.predict( records, dict_result=True, min_score=min_score, - timings_only=timings_only, + timings_only=False, include_labels=include_labels, ) - - if timings_only: - return record_labels.to_dict() + if isinstance(record_labels, Timings): + raise RuntimeError("NER predictions were not returned") # Update model_metadata based on records that were classified self.dataset_metadata_tracker.update_fields(records) diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/model.py b/src/nemo_safe_synthesizer/pii_replacer/ner/model.py index 749ecf34a..428514362 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/model.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/model.py @@ -6,18 +6,198 @@ from __future__ import annotations import re +from collections.abc import Mapping +from numbers import Real +from typing import Any, Literal, TypeAlias, TypedDict, TypeGuard, overload -from ...data_processing.records.fragment import create_ner_api_response +from typing_extensions import TypeIs + +from ...data_processing.records.fragment import ( + NERApiResponseRow, + NERRawPredictionPayload, + create_ner_api_response, +) +from ...pii_replacer.ner import ner, pipeline, regex from ...pii_replacer.ner.entity import Score -from ...pii_replacer.ner.ner import ner -from ...pii_replacer.ner.pipeline import pipeline -from ...pii_replacer.ner.regex import regex INPUT_ERR = "Input data must be a string, dict, or a list of either" _source_validator = re.compile(r"[A-Za-z_]{2,15}$") +class NERPredictorTimingPayload(TypedDict): + total_time_ms: float + total_time_ms_avg: float + + +class NERTimingsPayload(TypedDict): + records: int + total_predictions: int + total_time_ms: float + total_time_ms_avg: float + time_per_prediction_ms: float + predictors: dict[str, NERPredictorTimingPayload] + + +NERPredictionRows: TypeAlias = list[list[NERRawPredictionPayload]] +NERModelPredictionResponse: TypeAlias = NERPredictionRows | list[NERApiResponseRow] +NERInputRecord: TypeAlias = dict[str, Any] +NERInputRows: TypeAlias = list[str] | list[NERInputRecord] +RawPredictionRows: TypeAlias = list[list[object]] + + +def _is_string_rows(rows: object) -> TypeGuard[list[str]]: + return isinstance(rows, list) and bool(rows) and all(isinstance(row, str) for row in rows) + + +def _is_record_rows(rows: object) -> TypeGuard[list[NERInputRecord]]: + return ( + isinstance(rows, list) + and bool(rows) + and all(isinstance(row, dict) and all(isinstance(key, str) for key in row) for row in rows) + ) + + +def _is_raw_prediction_rows(value: object) -> TypeIs[RawPredictionRows]: + return isinstance(value, list) and all(isinstance(row, list) for row in value) + + +def _required_str(value: object, field_name: str) -> str: + if not isinstance(value, str): + raise TypeError(f"NER prediction field {field_name!r} must be a string") + return value + + +def _required_int(value: object, field_name: str) -> int: + if not isinstance(value, int): + raise TypeError(f"NER prediction field {field_name!r} must be an integer") + return value + + +def _optional_float(value: object, field_name: str) -> float | None: + if value is None: + return None + if isinstance(value, bool) or not isinstance(value, Real): + raise TypeError(f"NER prediction field {field_name!r} must be a float or None") + return float(value) + + +def _required_float(value: object, field_name: str) -> float: + if isinstance(value, bool) or not isinstance(value, Real): + raise TypeError(f"NER timing field {field_name!r} must be a float") + return float(value) + + +def _optional_str(value: object, field_name: str) -> str | None: + if value is None: + return None + if not isinstance(value, str): + raise TypeError(f"NER prediction field {field_name!r} must be a string or None") + return value + + +def _optional_value_path(value: object, field_name: str) -> tuple[str | int, ...] | list[str | int] | None: + if value is None: + return None + if not isinstance(value, (tuple, list)): + raise TypeError(f"NER prediction field {field_name!r} must be a string/integer path or None") + path_parts: list[str | int] = [] + for part in value: + if not isinstance(part, (str, int)): + raise TypeError(f"NER prediction field {field_name!r} must be a string/integer path or None") + path_parts.append(part) + if isinstance(value, tuple): + return tuple(path_parts) + return path_parts + + +def _optional_bool(value: object, field_name: str) -> bool | None: + if value is None: + return None + if not isinstance(value, bool): + raise TypeError(f"NER prediction field {field_name!r} must be a boolean or None") + return value + + +def _raw_prediction_payload(value: object) -> NERRawPredictionPayload: + match value: + case Mapping() as prediction: + payload: NERRawPredictionPayload = { + "text": _required_str(prediction.get("text"), "text"), + "start": _required_int(prediction.get("start"), "start"), + "end": _required_int(prediction.get("end"), "end"), + "label": _required_str(prediction.get("label"), "label"), + "source": _required_str(prediction.get("source"), "source"), + "score": _optional_float(prediction.get("score"), "score"), + } + match prediction: + case {"field": field}: + payload["field"] = _optional_str(field, "field") + match prediction: + case {"value_path": value_path}: + payload["value_path"] = _optional_value_path(value_path, "value_path") + match prediction: + case {"substring_match": substring_match}: + payload["substring_match"] = _optional_bool(substring_match, "substring_match") + return payload + case _: + raise TypeError("NER prediction rows must contain dictionaries") + + +def _prediction_rows(value: object) -> NERPredictionRows: + match value: + case list() as rows if _is_raw_prediction_rows(rows): + return [[_raw_prediction_payload(prediction) for prediction in row] for row in rows] + case _: + raise TypeError("NER predictions must be a list of prediction rows") + + +def _timings_payload_from_mapping(value: object) -> NERTimingsPayload: + match value: + case Mapping() as timing_data: + match timing_data.get("predictors"): + case Mapping() as predictors: + predictor_timings: dict[str, NERPredictorTimingPayload] = {} + for predictor, timing in predictors.items(): + match predictor, timing: + case str() as predictor_name, Mapping() as predictor_timing: + predictor_timings[predictor_name] = { + "total_time_ms": _required_float( + predictor_timing.get("total_time_ms"), "total_time_ms" + ), + "total_time_ms_avg": _required_float( + predictor_timing.get("total_time_ms_avg"), "total_time_ms_avg" + ), + } + case _, Mapping(): + raise TypeError("NER timing predictor names must be strings") + case str(), _: + raise TypeError("NER predictor timings must be dictionaries") + case _: + raise TypeError("NER timing predictor names must be strings") + case _: + raise TypeError("NER timings field 'predictors' must be a dictionary") + case _: + raise TypeError("NER timings must be a dictionary") + + return { + "records": _required_int(timing_data.get("records"), "records"), + "total_predictions": _required_int(timing_data.get("total_predictions"), "total_predictions"), + "total_time_ms": _required_float(timing_data.get("total_time_ms"), "total_time_ms"), + "total_time_ms_avg": _required_float(timing_data.get("total_time_ms_avg"), "total_time_ms_avg"), + "time_per_prediction_ms": _required_float(timing_data.get("time_per_prediction_ms"), "time_per_prediction_ms"), + "predictors": predictor_timings, + } + + +def _timings_payload(value: object) -> NERTimingsPayload: + match getattr(value, "to_dict", None): + case to_dict if callable(to_dict): + return _timings_payload_from_mapping(to_dict()) + case _: + raise TypeError("timings_only=True must return NER timings") + + def _parse_custom_source(source: str) -> tuple[str, str]: """Return a namespace, name str tuple""" parts = source.split("/") @@ -35,49 +215,112 @@ class Model: several NER techniques into a simple interface """ - def __init__(self, *args, exclude: list[str] | None = None): + def __init__(self, *args: str, exclude: list[str] | None = None): if args and exclude: raise ValueError("Cannot include and exclude predictors") - _pipeline = pipeline.from_source_string_list(include=args, exclude=exclude) + include = list(args) if args else None + if include is not None: + _pipeline = pipeline.from_source_string_list(include=include) + elif exclude is not None: + _pipeline = pipeline.from_source_string_list(exclude=exclude) + else: + _pipeline = pipeline.from_source_string_list() self._ner = ner.NER(pipeline=_pipeline) @property def predictors(self) -> list[str]: - return [pred.source for pred in self._ner.pipeline.predictors] + active_pipeline = self._ner.pipeline + if active_pipeline is None: + raise RuntimeError("NER pipeline is not configured") + return [pred.source for pred in active_pipeline.predictors] - def predict(self, input_data: str | dict | list[str] | list[dict], *, timings_only=False) -> list[dict] | dict: - if isinstance(input_data, (str, dict)): - input_data = [input_data] + @overload + def predict( + self, + input_data: str | dict[str, Any] | list[str] | list[dict[str, Any]], + *, + timings_only: Literal[True], + ) -> NERTimingsPayload: ... - if not isinstance(input_data, list): - raise ValueError(INPUT_ERR) + @overload + def predict( + self, + input_data: str, + *, + timings_only: Literal[False] = False, + ) -> NERPredictionRows: ... - if not isinstance(input_data[0], (str, dict)): - raise ValueError(INPUT_ERR) + @overload + def predict( + self, + input_data: list[str], + *, + timings_only: Literal[False] = False, + ) -> list[list[NERRawPredictionPayload]]: ... - _target_type = type(input_data[0]) + @overload + def predict( + self, + input_data: dict[str, Any] | list[dict[str, Any]], + *, + timings_only: Literal[False] = False, + ) -> list[NERApiResponseRow]: ... - for _target in input_data: - if not isinstance(_target, _target_type): + @overload + def predict( + self, + input_data: str | dict[str, Any] | list[str] | list[dict[str, Any]], + *, + timings_only: bool = False, + ) -> NERModelPredictionResponse | NERTimingsPayload: ... + + def predict( + self, + input_data: str | dict[str, Any] | list[str] | list[dict[str, Any]], + *, + timings_only: bool = False, + ) -> NERModelPredictionResponse | NERTimingsPayload: + match input_data: + case str() as text: + input_rows: NERInputRows = [text] + case dict() as record: + input_rows = [record] + case list() as rows if _is_string_rows(rows): + input_rows = rows + case list() as rows if _is_record_rows(rows): + input_rows = rows + case _: raise ValueError(INPUT_ERR) - predictions = self._ner.predict(input_data, timings_only=timings_only, dict_result=True) + predictions = self._ner.predict(input_rows, timings_only=timings_only, dict_result=True) if timings_only: - return predictions.to_dict() - if _target_type is str: - return predictions + return _timings_payload(predictions) - return create_ner_api_response(input_data, predictions, pure_dict=True) + prediction_rows = _prediction_rows(predictions) + match input_rows: + case list() as rows if _is_string_rows(rows): + return prediction_rows + case list() as rows if _is_record_rows(rows): + return create_ner_api_response( + rows, + prediction_rows, + pure_dict=True, + ) + case _: + raise ValueError(INPUT_ERR) - def add_regex(self, source: str, pattern: str | re.Pattern, score: Score | None = None): + def add_regex(self, source: str, pattern: str | re.Pattern, score: float | None = None): namespace, name = _parse_custom_source(source) if score is None: score = Score.HIGH - pattern = regex.Pattern(pattern=pattern, raw_score=score.value) - predictor = regex.RegexPredictor(name=name, namespace=namespace, patterns=[pattern]) - self._ner.pipeline.add_predictors(predictor) + regex_pattern = regex.Pattern(pattern=pattern, raw_score=score) + predictor = regex.RegexPredictor(name=name, namespace=namespace, patterns=[regex_pattern]) + active_pipeline = self._ner.pipeline + if active_pipeline is None: + raise RuntimeError("NER pipeline is not configured") + active_pipeline.add_predictors(predictor) def create_empty() -> Model: diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/models.py b/src/nemo_safe_synthesizer/pii_replacer/ner/models.py index 9ae09211d..a2d46f2bd 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/models.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/models.py @@ -9,7 +9,7 @@ from dataclasses import dataclass from enum import Enum from pathlib import Path -from typing import Any, Optional +from typing import Any, ClassVar, Optional from ...observability import get_logger @@ -107,7 +107,7 @@ def from_system(cls) -> StorageConfig: return cls(bucket=DEFAULT_BUCKET, cache_dir=cache_dir) -def get_cache_manager(storage_config: StorageConfig = None) -> CacheManager: +def get_cache_manager(storage_config: StorageConfig | None = None) -> CacheManager: """Returns a singleton instance of ``CacheManager``.""" return CacheManager.get_instance(storage_config) @@ -122,7 +122,7 @@ class CacheManager: storage_config: A storage config. """ - __instance = None + __instance: ClassVar[CacheManager | None] = None """Used to hold a singleton of ``CacheManager``""" _cache: dict[str, dict[str, Any]] @@ -139,13 +139,15 @@ def reset(): CacheManager.__instance = None @classmethod - def get_instance(cls, storage_config: StorageConfig = None) -> CacheManager: + def get_instance(cls, storage_config: StorageConfig | None = None) -> CacheManager: """Returns a singleton instance of ``CacheManager``.""" if not CacheManager.__instance: CacheManager(storage_config) + if CacheManager.__instance is None: + raise RuntimeError("CacheManager singleton was not initialized") return CacheManager.__instance - def __init__(self, storage_config: StorageConfig = None): + def __init__(self, storage_config: StorageConfig | None = None): if CacheManager.__instance: raise Exception("Cannot instantiate a singleton.") else: diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/ner.py b/src/nemo_safe_synthesizer/pii_replacer/ner/ner.py index a5706b793..d73b6e1a8 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/ner.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/ner.py @@ -10,7 +10,7 @@ from collections import defaultdict from dataclasses import dataclass from dataclasses import field as dataclasses_field -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING, Any, Literal, Optional, overload from ...data_processing.records.json_record import JSONRecord from ...data_processing.records.value_path import ( @@ -55,15 +55,17 @@ class NERPrediction: """ @property - def as_dict(self): + def as_dict(self) -> dict[str, Any]: return self.__dict__ @classmethod - def from_dict(cls, source: dict): + def from_dict(cls, source: dict[str, Any]) -> NERPrediction: return cls(**source) @property - def json_path(self): + def json_path(self) -> str | None: + if self.value_path is None: + return None return value_path_to_json_path(self.value_path) def get_dedupe_key_by_label(self, record: JSONRecord) -> tuple: @@ -95,8 +97,11 @@ def get_dedupe_key_by_label(self, record: JSONRecord) -> tuple: ) -"""Represents the prediction results of a pipeline""" -PipelineResult = list[list[NERPrediction | dict]] +"""Represents the prediction results of a pipeline.""" +PredictionDict = dict[str, Any] +Prediction = NERPrediction | PredictionDict +PredictionList = list[Prediction] +PipelineResult = PredictionList | list[PredictionList] PRECISION = 4 @@ -123,7 +128,7 @@ class Timings: predictors: dict = dataclasses_field(default_factory=lambda: defaultdict(PredictorTimings)) def to_dict(self): - ret = { + ret: dict[str, Any] = { "records": self.records, "total_predictions": self.total_predictions, "total_time_ms": round(self.total_time * 1000, PRECISION), @@ -156,22 +161,46 @@ def set_avg(self, num_cpu: int = 1): class NER: """Entity Recognition Pipeline""" - pipeline: Pipeline + pipeline: Pipeline | None - def __init__(self, predictor_cache_size=DEFAULT_CACHE_SIZE, pipeline: Pipeline = None): + def __init__(self, predictor_cache_size: int = DEFAULT_CACHE_SIZE, pipeline: Pipeline | None = None): self.predictor_cache_size = predictor_cache_size self.pipeline = pipeline + @overload def predict( self, in_data: InData, *, - pipeline: Pipeline = None, + pipeline: Pipeline | None = None, + dict_result: bool = False, + min_score: float = 0.0, + timings_only: Literal[True], + include_labels: Optional[set[str]] = None, + ) -> Timings: ... + + @overload + def predict( + self, + in_data: InData, + *, + pipeline: Pipeline | None = None, + dict_result: bool = False, + min_score: float = 0.0, + timings_only: Literal[False] = False, + include_labels: Optional[set[str]] = None, + ) -> PipelineResult: ... + + def predict( + self, + in_data: InData, + *, + pipeline: Pipeline | None = None, dict_result: bool = False, min_score: float = 0.0, timings_only: bool = False, include_labels: Optional[set[str]] = None, - ) -> PipelineResult: + ) -> PipelineResult | Timings: """ Predict entities from a string, list or dictionary object. @@ -198,23 +227,25 @@ def predict( # the instance configured pipeline if not pipeline: pipeline = self.pipeline + if pipeline is None: + raise NERError("Pipeline is empty") processed_in_data = input_to_json_records(in_data) # type: List[JSONRecord] # create an empty array for each item # that we are predicting - slots = [[] for _ in processed_in_data] + slots: list[PredictionList] = [[] for _ in processed_in_data] # for each predictor we need to add all # predictions for each item to the right slot - result_set = [] + result_set: list[list[NERPrediction]] = [] timings = Timings(records=len(processed_in_data)) input_record: JSONRecord for input_record in processed_in_data: timings.total_predictions += len(input_record.kv_pairs) - record_predictions = [] + record_predictions: list[NERPrediction] = [] for predictor in pipeline.iter_predictors(): start = time.perf_counter() predictions = predictor.evaluate(input_record) @@ -239,7 +270,7 @@ def predict( return timings for slot, preds in zip(slots, result_set): - results = preds + results: PredictionList = list(preds) if dict_result: results = [p.as_dict for p in preds] slot.extend(results) diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/ner_mp.py b/src/nemo_safe_synthesizer/pii_replacer/ner/ner_mp.py index da0f76d7d..557349daf 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/ner_mp.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/ner_mp.py @@ -13,12 +13,12 @@ from typing import Any, Optional import joblib.externals.loky as loky +from typing_extensions import TypeIs -from ...data_processing.records.json_record import JSONRecord from ...observability import get_logger from . import pipeline -from .ner import NER, PipelineResult, Timings -from .utils import InData +from .ner import NER, NERPrediction, PipelineResult, Prediction, PredictionList, Timings +from .utils import InData, input_to_json_records logger = get_logger(__name__) @@ -30,10 +30,10 @@ class _ProcPayload: seq: int in_data: InData - out_data: dict | Timings | PipelineResult = None + out_data: Timings | PipelineResult | None = None -_ner_predictor = None # type: ner.NER +_ner_predictor: NER | None = None """This global var should only ever be init'd to an NER instance by a process that is part of the process pool. By making it a global it becomes more stable to @@ -42,6 +42,25 @@ class _ProcPayload: """ +def _is_prediction(value: object) -> TypeIs[Prediction]: + return isinstance(value, NERPrediction) or (isinstance(value, dict) and all(isinstance(key, str) for key in value)) + + +def _is_prediction_list(value: object) -> TypeIs[PredictionList]: + return isinstance(value, list) and all(_is_prediction(item) for item in value) + + +def _extend_single_predictions(target: PredictionList, out_data: object) -> None: + """Append worker output for a single input record to ``target``.""" + if _is_prediction_list(out_data): + target.extend(out_data) + return + if isinstance(out_data, list): + for row in out_data: + if _is_prediction_list(row): + target.extend(row) + + def _set_ner_predictor(pipeline_factory: Callable[[], pipeline.Pipeline]): """This is the init routine when forking the process pool and should be passed into the pool constructor @@ -65,6 +84,8 @@ def _predict(payload: _ProcPayload, **kwargs) -> _ProcPayload: this is ever called """ try: + if _ner_predictor is None: + raise RuntimeError("NER worker was not initialized") payload.out_data = _ner_predictor.predict(payload.in_data, **kwargs) return payload except Exception as e: @@ -156,29 +177,29 @@ def predict(self, in_data: InData, **kwargs) -> Timings | PipelineResult: list_input = isinstance(in_data, list) if list_input: - if len(in_data) > 0 and isinstance(in_data[0], JSONRecord): - # Send pure dicts to NER, as it's much faster to pickle/unpickle - # pure dicts that JSONRecords (and that's what multiprocessing is doing) - in_data = [record.original for record in in_data] + # Send pure dicts to NER, as it's much faster to pickle/unpickle + # pure dicts that JSONRecords (and that's what multiprocessing is doing) + records = [record.original for record in input_to_json_records(in_data)] - record_chunks = iter_record_chunks(iter(in_data), CHUNK_SIZE) + chunks: list[tuple[InData, int]] = [ + (chunk, len(chunk)) for chunk in iter_record_chunks(iter(records), CHUNK_SIZE) + ] else: # We need to handle the case where a non-list is sent in # as our prediction object and we need to makesure the payload # we send to the worker is the raw input, not a list - record_chunks = [in_data] + chunks = [(in_data, 1)] result_data_tracker = _ResultData() total_chunks = 0 submitted_chunks = [] - for i, p in enumerate(record_chunks): + for i, (data, chunk_size) in enumerate(chunks): result_data_tracker.lock.acquire() logger.info(f"Submitting chunk number {i + 1} to NER workers.") - data = list(p) if list_input else p payload = _ProcPayload(seq=i, in_data=data) - submitted_chunks.append(_ChunkInfo(seq=i, chunk_size=len(data) if list_input else 1)) + submitted_chunks.append(_ChunkInfo(seq=i, chunk_size=chunk_size)) self.pool.submit(_predict, payload, **kwargs).add_done_callback(result_data_tracker.handle_results) @@ -202,16 +223,15 @@ def predict(self, in_data: InData, **kwargs) -> Timings | PipelineResult: self._initialize_pool() # in case it's reused break - result_data_tracker.progress_callback.flush() result_data = result_data_tracker.results logger.info("NER prediction completed.") if timings_only: - result_timings = iter(result_data) - timings = next(result_timings).out_data - for other_timings in result_timings: - timings.join(other_timings.out_data) + result_timings = [result.out_data for result in result_data if isinstance(result.out_data, Timings)] + timings = result_timings[0] if result_timings else Timings() + for other_timings in result_timings[1:]: + timings.join(other_timings) timings.set_avg(num_cpu=self.num_proc) return timings @@ -220,17 +240,27 @@ def predict(self, in_data: InData, **kwargs) -> Timings | PipelineResult: # Restore the predictions to the order they would # have been if predicting on a single worker - preds = [] + if list_input: + batch_preds: list[PredictionList] = [] + else: + single_preds: PredictionList = [] for seq in sorted(list(all_chunks.keys())): if (payload := completed_chunks.get(seq, None)) is not None: - preds.extend(payload.out_data) + if isinstance(payload.out_data, list): + if list_input: + for row in payload.out_data: + if _is_prediction_list(row): + batch_preds.append(row) + else: + _extend_single_predictions(single_preds, payload.out_data) else: # Add an empty spot for each record in the chunk that wasn't completed logger.warning(f"NER for chunk number {seq + 1} did not complete.") - preds.extend([[]] * all_chunks[seq].chunk_size) + if list_input: + batch_preds.extend([[]] * all_chunks[seq].chunk_size) - return preds + return batch_preds if list_input else single_preds def __exit__(self, exc_type, exc_value, traceback): logger.info("Shutting down NER worker pool.") diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/nlp.py b/src/nemo_safe_synthesizer/pii_replacer/ner/nlp.py index d6861c316..39bb5d35b 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/nlp.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/nlp.py @@ -10,6 +10,8 @@ from time import perf_counter from typing import TYPE_CHECKING, Any, Optional +from typing_extensions import override + from ...data_processing.records.base import KVPair from ...data_processing.records.json_record import JSONRecord from ...data_processing.records.value_path import ValuePath @@ -26,10 +28,25 @@ from .predictor import Predictor from .utils import is_string_a_number -spacy = None -Doc = None -Span = None -srsly = None +spacy: Any +Doc: Any +Span: Any +srsly: Any +try: + import spacy as _spacy # ty: ignore[unresolved-import] + import srsly as _srsly # ty: ignore[unresolved-import] + from spacy.tokens import Doc as _Doc # ty: ignore[unresolved-import] + from spacy.tokens import Span as _Span # ty: ignore[unresolved-import] +except ImportError: + spacy = None + Doc = None + Span = None + srsly = None +else: + spacy = _spacy + Doc = _Doc + Span = _Span + srsly = _srsly if TYPE_CHECKING: @@ -75,7 +92,7 @@ SPACY_DELIM = " is " -def _is_valid_spacy_entity(ent: Span): +def _is_valid_spacy_entity(ent: Any): """Removes entities predicted by Spacy ML models that are not contained in ENTITY_VALID_CHARACTERS, or that are greater than ENTITY_MAX_CHARS in length @@ -210,45 +227,40 @@ def _get_spacy_ent_score(_, ent: Entity) -> float: elif ent == Entity.PERSON_NAME: return Score.LOW else: - return None + return 0.0 class SpacyPredictor(Predictor): nlp: Any - timings: dict[str, Number] + timings: dict[str, float] default_name: str = "spacy" def __init__( self, - name: str = None, - model: str = None, + name: str | None = None, + model: Any | None = None, namespace: Optional[str] = None, ): if spacy is None: raise RuntimeError("spacy is not installed and must be for spacy predictors") self.timings = {} - - if name is None and model is None: - # We load our default Spacy model - name = "spacy" - start_time = perf_counter() - model_bytes = get_cache_manager().resolve(spacy_manifest)["model_data"] # noqa: F821 - nlp = spacy.blank("en") - ner_pipe = nlp.create_pipe("ner") - nlp.add_pipe(ner_pipe) - nlp.from_bytes(model_bytes) - self.timings[self.default_name] = perf_counter() - start_time - self.nlp = nlp - else: - raise ValueError("Spacy predictor no longer works") + if model is None: + raise ValueError("Spacy predictor requires a model loaded with from_manifest()") + name = name or self.default_name + self.nlp = model super().__init__(name=name, namespace=namespace) @classmethod def from_manifest(cls, manifest: ModelManifest) -> "SpacyPredictor": + if spacy is None: + raise RuntimeError("spacy is not installed and must be for spacy predictors") start_time = perf_counter() - model_bytes = get_cache_manager().resolve(manifest)["model_data"] + model_data = get_cache_manager().resolve(manifest) + if model_data is None: + raise RuntimeError(f"Could not resolve manifest {manifest}") + model_bytes = model_data["model_data"] nlp = spacy.blank("en") ner_pipe = nlp.create_pipe("ner") nlp.add_pipe(ner_pipe) @@ -263,6 +275,7 @@ def _predict(self, input_text: str) -> Doc: doc.set_extension(const.NER_SCORE, default=None, force=True) return self.nlp(doc.text) + @override def evaluate(self, in_data: JSONRecord) -> list[NERPrediction]: fields = _flatten_fields(in_data.kv_pairs) diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/person_name.py b/src/nemo_safe_synthesizer/pii_replacer/ner/person_name.py index d917585d9..d37026327 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/person_name.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/person_name.py @@ -11,8 +11,10 @@ import re from collections.abc import Iterable from dataclasses import dataclass, field +from typing import Any from flashtext import KeywordProcessor +from typing_extensions import override from ...data_processing.records.base import KVPair, tokenize_header from ...data_processing.records.json_record import JSONRecord @@ -45,8 +47,8 @@ MAX_STR_LEN = 64 -def build_name_only_headers(others: Iterable[str]) -> list[re.Pattern]: - out = [] +def build_name_only_headers(others: Iterable[str]) -> list[str]: + out: list[str] = [] for other in others: out.append(r"{}.?{}".format("name", other)) out.append(r"{}.?{}".format(other, "name")) @@ -59,13 +61,13 @@ class WordList: names MUST match the ref names of the files from the manifest """ - word_list: KeywordProcessor = field(default_factory=frozenset) + word_list: KeywordProcessor | frozenset[str] = field(default_factory=frozenset) """The master list of actual names""" - headers: re.Pattern | None = None # NOTE: init'd as a FrozenSet then converted + headers: re.Pattern[str] | frozenset[str] = field(default_factory=frozenset) """The list of partial header names that can trigger the prediction flow""" - headers_neg: KeywordProcessor = None # NOTE: init'd as a FrozenSet then converted + headers_neg: KeywordProcessor | frozenset[str] = field(default_factory=frozenset) """A list of header tokens that should not be present to trigger prediction flow""" headers_pairs: frozenset[str] = field(default_factory=frozenset) @@ -83,16 +85,22 @@ def __post_init__(self): # matching header values from them, and then re-set our master # header list tmp = build_name_only_headers(self.headers_pairs) + if not isinstance(self.headers, frozenset): + raise TypeError("headers must be loaded as a frozenset") _header_strings = list(self.headers | frozenset(tmp)) _header_regex = re.compile("|".join(_header_strings), re.IGNORECASE) self.headers = _header_regex _neg_headers = KeywordProcessor() + if not isinstance(self.headers_neg, frozenset): + raise TypeError("negative headers must be loaded as a frozenset") _neg_headers.add_keywords_from_list(list(self.headers_neg)) self.headers_neg = _neg_headers # self.word_list = re.compile("|".join([w + "$" for w in list(self.word_list)]), re.IGNORECASE) _word_list = KeywordProcessor() + if not isinstance(self.word_list, frozenset): + raise TypeError("word list must be loaded as a frozenset") _word_list.add_keywords_from_list(list(self.word_list)) self.word_list = _word_list @@ -100,14 +108,32 @@ def __post_init__(self): # _parts.add_keywords_from_list(list(self.parts)) # self.parts = _parts + @property + def word_processor(self) -> KeywordProcessor: + if not isinstance(self.word_list, KeywordProcessor): + raise TypeError("word list has not been initialized") + return self.word_list + + @property + def header_pattern(self) -> re.Pattern[str]: + if not isinstance(self.headers, re.Pattern): + raise TypeError("headers have not been initialized") + return self.headers + + @property + def negative_header_processor(self) -> KeywordProcessor: + if not isinstance(self.headers_neg, KeywordProcessor): + raise TypeError("negative headers have not been initialized") + return self.headers_neg + @classmethod - def init_from_manifest(cls, manifest: ModelManifest = None): + def init_from_manifest(cls, manifest: ModelManifest | None = None) -> WordList: if manifest is None: manifest = DEFAULT_MANIFEST cache_data = get_cache_manager().resolve(manifest, skip_pickle=True) if cache_data is None: raise RuntimeError("Model cached returned None data for word list") - kwargs = {} + kwargs: dict[str, Any] = {} # NOTE: this functionality depends on the dataclass attrs # being named the same as the keys in the ``ObjectRef`` # instances since those keys are what is returned @@ -126,11 +152,12 @@ def __init__(self): super().__init__(self.default_name) self.word_list = WordList.init_from_manifest() - def create_prediction(self, record_field: KVPair): + def create_prediction(self, record_field: KVPair) -> NERPrediction: + value = str(record_field.value) return NERPrediction( - text=record_field.value, + text=value, start=0, - end=len(record_field.value), + end=len(value), field=record_field.field, value_path=record_field.value_path, score=Score.HIGH, @@ -149,12 +176,12 @@ def check_exact_name_header_data(self, field_value: str) -> bool: if len(token) == 1: continue - in_main_word_list = self.word_list.word_list.extract_keywords(token) + in_main_word_list = self.word_list.word_processor.extract_keywords(token) # if the token is not in any of these lists, fail if ( not in_main_word_list - and not self.word_list.headers.match(token) + and not self.word_list.header_pattern.match(token) # and token not in self.word_list.headers and token not in self.word_list.parts ): @@ -171,13 +198,14 @@ def _is_neg_header_in_value(self, value) -> bool: value_tokens = tokenize_header(str(value)) for _token in value_tokens: # if _token in self.word_list.headers_neg: - if self.word_list.headers_neg.match(_token): + if self.word_list.negative_header_processor.extract_keywords(_token): return True return False - def evaluate(self, in_record: JSONRecord) -> list[NERPrediction]: - record_fields = in_record.kv_pairs - result_set_by_field = [set() for _ in record_fields] + @override + def evaluate(self, in_data: JSONRecord) -> list[NERPrediction]: + record_fields = in_data.kv_pairs + result_set_by_field: list[set[NERPrediction]] = [set() for _ in record_fields] record_field: KVPair for field_matches, record_field in zip(result_set_by_field, record_fields): @@ -189,16 +217,16 @@ def evaluate(self, in_record: JSONRecord) -> list[NERPrediction]: # check if any negative header fields exist for header_token in record_field.field_tokens: - if self.word_list.headers_neg.extract_keywords(header_token): + if self.word_list.negative_header_processor.extract_keywords(header_token): continue # tokenize the value and see if any of tokens exist # in the negative header list # if record_field.value and self._is_neg_header_in_value(record_field.value): - if self.word_list.headers_neg.extract_keywords(record_field.value): + if self.word_list.negative_header_processor.extract_keywords(record_field.value): continue - if not self.header_has_context(record_field, self.KEY, regex_patterns=self.word_list.headers): + if not self.header_has_context(record_field, self.KEY, regex_patterns=self.word_list.header_pattern): # specialy handling with the field name is exactly "name", we need # to check if every token in the value exists in one of our specific # name or modifier lists @@ -213,7 +241,7 @@ def evaluate(self, in_record: JSONRecord) -> list[NERPrediction]: # for token in re.finditer(TOKEN_REGEX, record_field.value): # token_str = token.group(0).lower() - if not self.word_list.word_list.extract_keywords(record_field.value): + if not self.word_list.word_processor.extract_keywords(record_field.value): continue # if the token is in the word list, we consider diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/pipeline.py b/src/nemo_safe_synthesizer/pii_replacer/ner/pipeline.py index 5e0feff52..b049d5ae3 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/pipeline.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/pipeline.py @@ -7,6 +7,9 @@ from enum import Enum from numbers import Number from pathlib import Path +from typing import Protocol + +from typing_extensions import override from ...observability import get_logger from . import person_name @@ -24,6 +27,10 @@ """User defined custom regex predictors and patterns""" +class PredictorFactory(Protocol): + def __call__(self) -> Predictor: ... + + class PredictionSource(Enum): """This enum stores default source tags for NLP and other more complex predictors that have associated "models" that need @@ -36,6 +43,7 @@ class PredictionSource(Enum): BIRTH_DATE = BirthDateTime @property + @override def name(self): return f"{Predictor.default_namespace}/{self.value.default_name}" # pylint: disable=no-member @@ -50,26 +58,26 @@ class Pipeline: predictors: list[Predictor] load_timings: dict[str, Number] - def __init__(self, predictors: list[Predictor] = None): - self.predictors = predictors or [] + def __init__(self, predictors: Sequence[Predictor] | None = None): + self.predictors = list(predictors) if predictors else [] self.load_timings = {} - def _next_udf_name(self): + def _next_udf_name(self) -> str: return f"user_defined_predictor_{len(self.predictors) + 1}" - def add_predictors(self, predictors: Predictor | list[Predictor]): + def add_predictors(self, predictors: Predictor | Sequence[Predictor]) -> Pipeline: if isinstance(predictors, Predictor): self.predictors.append(predictors) - if isinstance(predictors, list): + else: self.predictors.extend(predictors) return self - def add_pattern(self, pattern: Pattern, name: str = None): + def add_pattern(self, pattern: Pattern, name: str | None = None): name = name or self._next_udf_name() predictor = RegexPredictor.from_pattern(pattern, name=name) self.add_predictors(predictor) - def add_regex(self, regex: str, name: str = None, namespace: str = None): + def add_regex(self, regex: str, name: str | None = None, namespace: str | None = None): name = name or self._next_udf_name() predictor = RegexPredictor.from_regex(regex, name=name) self.add_predictors(predictor) @@ -81,7 +89,7 @@ def merge(self, other_pipeline: Pipeline) -> Pipeline: def iter_predictors(self) -> Iterator[Predictor]: return iter(self.predictors) - def get_predictor(self, source: str) -> Predictor: + def get_predictor(self, source: str) -> Predictor | None: """Returns the first predictor by source name in the pipeline Args: @@ -104,10 +112,10 @@ def add_predictors_from_yaml(self, file_path: str = CUSTOM_CONFIG): logger.info("Custom Predictors: Not Found, skipping") return logger.info("Custom Predictors: loading from %s", file_path) - self.add_predictors(get_predictors_from_yaml(_path)) + self.add_predictors(get_predictors_from_yaml(str(_path))) @classmethod - def from_class_refs(cls, predictors: Sequence[type[Predictor]]) -> Pipeline: + def from_class_refs(cls, predictors: Sequence[PredictorFactory]) -> Pipeline: klasses = [p() for p in predictors] return cls(klasses) @@ -151,7 +159,7 @@ def create_default_ner(full: bool = False) -> NER: return NER(pipeline=pipe) -def from_source_string_list(*, include: list[str] = None, exclude: list[str] = None) -> Pipeline: +def from_source_string_list(*, include: list[str] | None = None, exclude: list[str] | None = None) -> Pipeline: if include and exclude: raise ValueError("cannot include and exclude") diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/predictor.py b/src/nemo_safe_synthesizer/pii_replacer/ner/predictor.py index 24b3d04fd..fdd67478a 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/predictor.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/predictor.py @@ -4,6 +4,7 @@ from __future__ import annotations from abc import ABC, abstractmethod +from collections.abc import Sequence from dataclasses import dataclass from re import Pattern from typing import Optional @@ -74,7 +75,7 @@ class ContextSpan: to search for any matches from the ``pattern_list`` objects. """ - pattern_list: list[str | Pattern] + pattern_list: Sequence[str | Pattern[str]] span: int = DEFAULT_CONTEXT_SPAN_SIZE def is_match(self, data: str, start: int, end: int) -> bool: @@ -100,7 +101,7 @@ def is_match(self, data: str, start: int, end: int) -> bool: return False -def is_context_matched(data: str, start: int, end: int, spans: list[ContextSpan]) -> bool: +def is_context_matched(data: str, start: int, end: int, spans: Sequence[ContextSpan]) -> bool: for span in spans: if span.is_match(data, start, end): return True @@ -135,7 +136,7 @@ class Predictor(ABC): BOTH = 3 default_namespace: str = "safe-synthesizer" - default_name: str = None + default_name: str | None = None """Subclasses can set a default name to use that can be directly accessed as a class attr if need be. @@ -144,7 +145,7 @@ class Predictor(ABC): def __init__( self, name: str, - namespace: str = None, + namespace: str | None = None, predictor_context: Optional[PredictorContext] = None, ): if namespace is None: @@ -165,15 +166,15 @@ def header_has_context( self, field_pair: KVPair, header_context_source: int, - token_patterns: Pattern = None, - regex_patterns: Pattern = None, + token_patterns: Pattern[str] | None = None, + regex_patterns: Pattern[str] | None = None, ) -> bool: """Checks to see if the field has a label match.""" _field = field_pair if header_context_source == self.BOTH: - search_string = (_field.field + " " + _field.value if _field.field else _field.value).casefold() + search_string = (f"{_field.field} {str(_field.value)}" if _field.field else str(_field.value)).casefold() elif header_context_source == self.VALUE: - search_string = _field.value.casefold() + search_string = str(_field.value).casefold() else: if _field.field is None: return False diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regex.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regex.py index 51f291c90..90990eed6 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regex.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regex.py @@ -6,9 +6,12 @@ import itertools import re from collections import defaultdict +from collections.abc import Sequence from dataclasses import dataclass, field from typing import Optional +from typing_extensions import override + from ...data_processing.records.base import KVPair from ...data_processing.records.json_record import JSONRecord from .entity import Entity, Score @@ -17,8 +20,8 @@ def split_header_contexts( - contexts: list[str | re.Pattern], -) -> tuple[re.Pattern | None, re.Pattern | None]: + contexts: Sequence[str | re.Pattern[str]], +) -> tuple[re.Pattern[str] | None, re.Pattern[str] | None]: """Split a list of strings and re.Patterns into two distcit regexes. Returns (regexes, tokens) @@ -56,7 +59,7 @@ class Pattern: `NERError` if `pattern` is not a string or regex Pattern """ - pattern: str | re.Pattern + pattern: str | re.Pattern[str] context_score: Optional[float] = Score.HIGH """This is the optimal score that you want to assign when context exists @@ -72,21 +75,21 @@ class Pattern: """If set, do not emit a match if only the raw regex matches without any context """ - header_contexts: Optional[list[str | re.Pattern]] = field(default_factory=list) + header_contexts: Sequence[str | re.Pattern[str]] = field(default_factory=list) """A list of strings or regexes that should be used to check the name of the field / header for a match. If there are any matches here, then the ``context_score`` value will be used as the matched score """ - header_regexes: Optional[re.Pattern] = field(init=False, default=None) - header_tokens: Optional[re.Pattern] = field(init=False, default=None) + header_regexes: re.Pattern[str] | None = field(init=False, default=None) + header_tokens: re.Pattern[str] | None = field(init=False, default=None) - neg_header_contexts: Optional[list[str | re.Pattern]] = field(default_factory=list) + neg_header_contexts: Sequence[str | re.Pattern[str]] = field(default_factory=list) """A list of strings or regexes that can be used to disqualify a field from being analyzed. If used, any matches were will short-circuit processing for a given key/value pair.""" - neg_header_regexes: Optional[re.Pattern] = field(init=False, default=None) - neg_header_tokens: Optional[re.Pattern] = field(init=False, default=None) + neg_header_regexes: re.Pattern[str] | None = field(init=False, default=None) + neg_header_tokens: re.Pattern[str] | None = field(init=False, default=None) header_context_source: int = Predictor.KEY """If doing header context searching, this dictates where to search for the context. We default @@ -94,7 +97,7 @@ class Pattern: of the field name and value """ - span_contexts: Optional[ContextSpan | list[ContextSpan]] = field(default_factory=list) + span_contexts: ContextSpan | Sequence[ContextSpan] | None = field(default_factory=list) """A list of ``ContextSpan`` instances that will be used, if provided, to search surrounding text of a string match for other discrete strings or matching regular expressions. See the ``ContextSpan`` usage for more details. @@ -117,6 +120,14 @@ def __post_init__(self): self.neg_header_regexes, self.neg_header_tokens = split_header_contexts(self.neg_header_contexts) +def _span_contexts(spans: ContextSpan | Sequence[ContextSpan] | None) -> Sequence[ContextSpan]: + if spans is None: + return () + if isinstance(spans, ContextSpan): + return (spans,) + return spans + + class RegexPredictor(Predictor): """Base class that represents a single entity. @@ -127,14 +138,14 @@ class RegexPredictor(Predictor): def __init__( self, name: Optional[str] = None, - patterns: list[Pattern] = None, + patterns: Sequence[Pattern] | None = None, entity: Optional[Entity] = None, namespace: Optional[str] = None, ): if patterns is None: patterns = [] - self.patterns = patterns + self.patterns = list(patterns) self.entity = entity # NOTE: If a name is not provided, then we will use the @@ -145,7 +156,7 @@ def __init__( super().__init__(name, namespace=namespace) - def validate_match(self, matched_text: str, original_text: str): + def validate_match(self, matched_text: str, original_text: str) -> bool: """ A base method for regex rules to implement. @@ -165,7 +176,8 @@ def filter_by_range_by_score(self, field_matches: set[NERPrediction]) -> list[NE return [max(ps, key=lambda p: p.score) for _, ps in by_range] - def evaluate(self, in_record: JSONRecord, res_by_field=False) -> list[NERPrediction]: + @override + def evaluate(self, in_data: JSONRecord) -> list[NERPrediction]: """ Given a single record determine if any entities are represented. @@ -177,8 +189,8 @@ def evaluate(self, in_record: JSONRecord, res_by_field=False) -> list[NERPredict A list of entity predictions sorted by score. Top score is first entry in list. """ - record_fields = in_record.kv_pairs - result_set_by_field = [set() for _ in record_fields] + record_fields = in_data.kv_pairs + result_set_by_field: list[set[NERPrediction]] = [set() for _ in record_fields] record_field: KVPair for field_matches, record_field in zip(result_set_by_field, record_fields): @@ -234,7 +246,7 @@ def evaluate(self, in_record: JSONRecord, res_by_field=False) -> list[NERPredict record_field.value, start_pos, end_pos, - pattern.span_contexts, + _span_contexts(pattern.span_contexts), ): _score = pattern.context_score elif pattern.ignore_raw_score: @@ -257,20 +269,17 @@ def evaluate(self, in_record: JSONRecord, res_by_field=False) -> list[NERPredict filtered_results = map(self.filter_by_range_by_score, result_set_by_field) - if res_by_field: - return [list(res_set) for res_set in result_set_by_field] - results_flat = itertools.chain.from_iterable(filtered_results) results = sorted(results_flat, key=lambda i: i.score, reverse=True) return list(results) @classmethod - def from_pattern(cls, pattern: Pattern, name: str = None, namespace: str = None): + def from_pattern(cls, pattern: Pattern, name: str | None = None, namespace: str | None = None): return cls(patterns=[pattern], name=name, namespace=namespace) @classmethod - def from_regex(cls, regex: str, name: str = None, namespace: str = None): + def from_regex(cls, regex: str, name: str | None = None, namespace: str | None = None): pattern = Pattern(pattern=regex, raw_score=Score.MAX) return RegexPredictor.from_pattern(pattern, name=name, namespace=namespace) @@ -363,7 +372,11 @@ def get_predictors(self) -> list[RegexPredictor]: return out_predictors -def phrase_predictors_from_entity_ruler(name: str, er_patterns: list[dict], entity_map: dict) -> list[RegexPredictor]: +def phrase_predictors_from_entity_ruler( + name: str, + er_patterns: Sequence[dict[str, object]], + entity_map: dict[str, str | Entity], +) -> list[RegexPredictor]: """Given a list of Spacy EntityRuler patterns, create a phrase matcher predictor. """ @@ -372,11 +385,14 @@ def phrase_predictors_from_entity_ruler(name: str, er_patterns: list[dict], enti builder = PhraseMatcherBuilder(name) for er_pattern in er_patterns: - _label = er_label_map.get(er_pattern["label"], None) + label_key = er_pattern.get("label") + if not isinstance(label_key, str): + continue + _label = er_label_map.get(label_key, None) # _label = er_pattern["label"] if not _label: continue - _pattern = er_pattern["pattern"] + _pattern = er_pattern.get("pattern") if isinstance(_pattern, str): builder.add_phrase(_label, _pattern, case=True) @@ -386,12 +402,19 @@ def phrase_predictors_from_entity_ruler(name: str, er_patterns: list[dict], enti # to determine if it should have a whitespace added # before it - # seed the base string - _pattern = iter(_pattern) - _str = next(_pattern)["LOWER"] - + token_values: list[str] = [] for part in _pattern: - part = part["LOWER"] + match part: + case {"LOWER": str(lower)}: + token_values.append(lower) + case _: + continue + if not token_values: + continue + + # seed the base string + _str = token_values[0] + for part in token_values[1:]: if not part.isalnum() and len(part) == 1: _str += part else: diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/age.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/age.py index fcf13b4db..430e540ac 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/age.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/age.py @@ -10,6 +10,8 @@ import re +from typing_extensions import override + from ..entity import Entity from ..regex import Pattern, RegexPredictor @@ -39,9 +41,10 @@ def __init__(self): super().__init__(entity=entity, patterns=[age, desc]) - def validate_match(self, matched_str: str, _): + @override + def validate_match(self, matched_text: str, original_text: str) -> bool: try: - age = float(matched_str) + age = float(matched_text) return 0 <= age <= 120 * 12 except ValueError: pass diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/credit_card.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/credit_card.py index 7fc5647b9..450e532d1 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/credit_card.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/credit_card.py @@ -6,6 +6,7 @@ import re from stdnum import luhn +from typing_extensions import override from ..entity import Entity from ..predictor import ContextSpan @@ -150,5 +151,6 @@ def __init__(self): entity = Entity.CREDIT_CARD_NUMBER super().__init__(entity=entity, patterns=PATTERNS) - def validate_match(self, matched_text: str, _): + @override + def validate_match(self, matched_text: str, original_text: str) -> bool: return _is_luhn_checksum_valid(matched_text) diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/domain_name.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/domain_name.py index 43c3d2582..0f98abd8c 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/domain_name.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/domain_name.py @@ -6,6 +6,7 @@ import re import tldextract +from typing_extensions import override from ..entity import Entity from ..predictor import ContextSpan @@ -70,11 +71,12 @@ class DomainName(RegexPredictor): def __init__(self): entity = Entity.DOMAIN_NAME - self.tld_extract = tldextract.TLDExtract(suffix_list_urls=None) + self.tld_extract = tldextract.TLDExtract(suffix_list_urls=()) super().__init__(name="domain_name", entity=entity, patterns=[MATCHER, ISOLATED_MATCHER]) - def validate_match(self, in_text: str, _) -> bool: - result = self.tld_extract(in_text) + @override + def validate_match(self, matched_text: str, original_text: str) -> bool: + result = self.tld_extract(matched_text) return result.fqdn != "" diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/email.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/email.py index fe18b6681..026998209 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/email.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/email.py @@ -4,6 +4,7 @@ from __future__ import annotations import tldextract +from typing_extensions import override from ..entity import Entity, Score from ..regex import Pattern, RegexPredictor @@ -23,9 +24,10 @@ def __init__(self): raw_score=Score.HIGH, ) - self.tld_extract = tldextract.TLDExtract(suffix_list_urls=None) + self.tld_extract = tldextract.TLDExtract(suffix_list_urls=()) super().__init__(entity=entity, patterns=[match]) - def validate_match(self, in_text: str, _) -> bool: - result = self.tld_extract(in_text) + @override + def validate_match(self, matched_text: str, original_text: str) -> bool: + result = self.tld_extract(matched_text) return result.fqdn != "" diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/generic_key.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/generic_key.py index 83a60ad22..5ac209c92 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/generic_key.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/generic_key.py @@ -16,14 +16,14 @@ r"secret", ] -LABELS = [ +LABELS: list[str | re.Pattern[str]] = [ "token", ] for reg in REG: LABELS.append(re.compile("^" + reg)) -SPANNER_PATTERNS = ["token"] +SPANNER_PATTERNS: list[str | re.Pattern[str]] = ["token"] for reg in REG: SPANNER_PATTERNS.append(re.compile(reg)) diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/github.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/github.py index 42bcd7bae..28e9cc109 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/github.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/github.py @@ -5,6 +5,8 @@ import re +from typing_extensions import override + from ..entity import Entity, Score from ..predictor import ContextSpan from ..regex import Pattern, RegexPredictor @@ -41,11 +43,12 @@ def __init__(self): patterns=[_match_1], ) - def validate_match(self, matched_text, original_value): - if re.search(COMMIT_URL.format(matched_text), original_value): + @override + def validate_match(self, matched_text: str, original_text: str) -> bool: + if re.search(COMMIT_URL.format(matched_text), original_text): return False - if re.search(f"(?:id|h)={matched_text}", original_value): + if re.search(f"(?:id|h)={matched_text}", original_text): return False - if "commit" in original_value: + if "commit" in original_text: return False return True diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/iban.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/iban.py index 7221b25b3..5723f3595 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/iban.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/iban.py @@ -16,9 +16,11 @@ import string +from typing_extensions import override + # Import 're2' regex engine if installed, if not- import 'regex' try: - import re2 as re + import re2 as re # ty: ignore[unresolved-import] except ImportError: import regex as re @@ -222,17 +224,11 @@ def __init__(self): super().__init__(entity=Entity.IBAN_CODE, patterns=patterns) - def validate_match(self, in_text: str, _) -> bool: - pattern_text = in_text.replace(" ", "") + @override + def validate_match(self, matched_text: str, original_text: str) -> bool: + pattern_text = matched_text.replace(" ", "") is_valid_checksum = IBAN.__generate_iban_check_digits(pattern_text) == pattern_text[2:4] - # score = EntityRecognizer.MIN_SCORE - result = False - if is_valid_checksum: - if IBAN.__is_valid_format(pattern_text): - result = True - elif IBAN.__is_valid_format(pattern_text.upper()): - result = None - return result + return bool(is_valid_checksum and IBAN.__is_valid_format(pattern_text)) @staticmethod def __number_iban(iban): diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/imei.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/imei.py index 7797fbdae..97ac15864 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/imei.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/imei.py @@ -4,6 +4,7 @@ from __future__ import annotations from stdnum import imei +from typing_extensions import override from ..entity import Entity from ..predictor import ContextSpan @@ -38,9 +39,10 @@ def __init__(self): patterns=[unlikely_match, likely_match], ) - def validate_match(self, in_text: str, _) -> bool: + @override + def validate_match(self, matched_text: str, original_text: str) -> bool: try: - check = imei.validate(in_text) + check = imei.validate(matched_text) except Exception: return False return True if check else False diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/ip_address.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/ip_address.py index 60b84107d..cf6af3ab1 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/ip_address.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/ip_address.py @@ -5,6 +5,8 @@ import ipaddress +from typing_extensions import override + from ..entity import Entity, Score from ..regex import Pattern, RegexPredictor @@ -32,7 +34,8 @@ def __init__(self): patterns=[possible_match, likely_match], ) - def validate_match(self, matched_text: str, _): + @override + def validate_match(self, matched_text: str, original_text: str) -> bool: try: ipaddress.ip_address(matched_text) except ValueError: diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/jwt.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/jwt.py index c88f4c3e8..113d4f9bc 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/jwt.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/jwt.py @@ -6,6 +6,8 @@ import base64 import json +from typing_extensions import override + from ..entity import Entity, Score from ..regex import Pattern, RegexPredictor @@ -24,7 +26,8 @@ def __init__(self): super().__init__(entity=entity, patterns=[_match_1]) # https://github.com/Yelp/detect-secrets/blob/master/detect_secrets/plugins/jwt.py - def validate_match(self, matched_text: str, _): + @override + def validate_match(self, matched_text: str, original_text: str) -> bool: parts = matched_text.split(".") for idx, part in enumerate(parts): try: diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/swift.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/swift.py index 521a44949..67fc70280 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/swift.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/swift.py @@ -5,6 +5,8 @@ import re +from typing_extensions import override + from ..entity import Entity from ..regex import Pattern, RegexPredictor from .iban import regex_per_country @@ -55,7 +57,8 @@ def __init__(self): super().__init__(entity=Entity.SWIFT_CODE, patterns=patterns) - def validate_match(self, in_text: str, _) -> bool: - country_code = in_text[4:6] + @override + def validate_match(self, matched_text: str, original_text: str) -> bool: + country_code = matched_text[4:6] # Keys of the more extensive dict in iban regex are country codes used for swift return country_code.upper() in regex_per_country diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/url.py b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/url.py index 5c1c3c605..8b2b907ae 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/url.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/regexes/url.py @@ -4,6 +4,7 @@ from __future__ import annotations import tldextract +from typing_extensions import override from ..entity import Entity, Score from ..predictor import ContextSpan @@ -28,9 +29,10 @@ def __init__(self): header_contexts=URL_LABELS, span_contexts=SPANNER, ) - self.tld_extract = tldextract.TLDExtract(suffix_list_urls=None) + self.tld_extract = tldextract.TLDExtract(suffix_list_urls=()) super().__init__(entity=Entity.URL, patterns=[match]) - def validate_match(self, in_text: str, orig: str) -> bool: - result = self.tld_extract(in_text) + @override + def validate_match(self, matched_text: str, original_text: str) -> bool: + result = self.tld_extract(matched_text) return result.fqdn != "" diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/report/metadata.py b/src/nemo_safe_synthesizer/pii_replacer/ner/report/metadata.py index 7cdc19275..937e67016 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/report/metadata.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/report/metadata.py @@ -16,7 +16,7 @@ def convert_to_report(metadata: DatasetMetadata) -> DatasetMetadataReport: """Converts internal model_metadata object to the report.""" fields_report = [_convert_field(field) for field in metadata.data.fields] - entity_summary_report = [EntitySummaryReport(**e.model_dump()) for e in metadata.data.entities] + entity_summary_report = [EntitySummaryReport(**e.dict()) for e in metadata.data.entities] return DatasetMetadataReport( record_count=metadata.project_record_count, @@ -34,7 +34,7 @@ def _convert_field(field: FieldMetadata) -> FieldMetadataReport: approx_distinct_count=field.approx_cardinality, missing_count=field.missing, labels=field.field_labels, - attributes=field.field_attributes, + attributes=[attribute.value for attribute in field.field_attributes], entities=entities, types=[TypeReport(type=tm.type, count=tm.count) for tm in field.types], ) diff --git a/src/nemo_safe_synthesizer/pii_replacer/ner/utils.py b/src/nemo_safe_synthesizer/pii_replacer/ner/utils.py index d00d5ab8b..45c308cdf 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/utils.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/utils.py @@ -3,27 +3,35 @@ from __future__ import annotations -from ...data_processing.records.json_record import JSONRecord +from collections.abc import Sequence +from typing import TypeAlias -# Valid input data includes a str, tuple or dict -InData = str | list | dict | JSONRecord +from ...data_processing.records.json_record import JsonObject, JSONRecord + +JsonRecordInput: TypeAlias = str | JsonObject | JSONRecord +InData: TypeAlias = JsonRecordInput | Sequence[JsonRecordInput] def input_to_json_records(in_data: InData) -> list[JSONRecord]: """Try and convert python objects to a list of Fields""" - if isinstance(in_data, JSONRecord): - return [in_data] - if isinstance(in_data, (str, dict)): - return [JSONRecord(in_data)] - if isinstance(in_data, list): - out = [] - for record in in_data: - if isinstance(record, JSONRecord): - out.append(record) - else: - out.append(JSONRecord(record)) - return out - raise TypeError("Input data not supported.") + match in_data: + case JSONRecord() as record: + return [record] + case (str() | dict()) as record: + return [JSONRecord(record)] + case list() as records: + out: list[JSONRecord] = [] + for record in records: + match record: + case JSONRecord() as json_record: + out.append(json_record) + case (str() | dict()) as raw_record: + out.append(JSONRecord(raw_record)) + case _: + raise TypeError("Input data not supported.") + return out + case _: + raise TypeError("Input data not supported.") def is_string_a_number(value) -> bool: diff --git a/tests/data_processing/records/test_fragment.py b/tests/data_processing/records/test_fragment.py new file mode 100644 index 000000000..8e2dbc5d5 --- /dev/null +++ b/tests/data_processing/records/test_fragment.py @@ -0,0 +1,65 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from nemo_safe_synthesizer.data_processing.records.fragment import ( + E2F, + SCORE_HIGH, + SCORE_LOW, + SCORE_MED, + NERRawPredictionPayload, + build_ner_metadata, + create_ner_api_response, +) + + +def _prediction() -> NERRawPredictionPayload: + return { + "text": "alice@example.com", + "start": 0, + "end": 17, + "label": "email", + "source": "safe-synthesizer/email", + "score": 0.99, + "field": "contact", + "value_path": ("contact",), + "substring_match": None, + } + + +def test_build_ner_metadata_returns_payload_dict(): + metadata = build_ner_metadata([_prediction()]) + + assert isinstance(metadata, dict) + assert metadata["record_id"] + assert metadata["received_at"].endswith("Z") + assert metadata["fields"]["contact"]["ner"]["labels"] == [ + { + "start": 0, + "end": 17, + "label": "email", + "score": 0.99, + "source": "safe-synthesizer/email", + "text": "alice@example.com", + } + ] + assert metadata["entities"][SCORE_HIGH] == ["email"] + assert metadata["entities"][SCORE_MED] == [] + assert metadata["entities"][SCORE_LOW] == [] + assert metadata["entities"][E2F] == {"email": ["contact"]} + + +def test_create_ner_api_response_pure_dict_preserves_shape_and_normalizes_metadata(): + response = create_ner_api_response( + [{"contact": "alice@example.com"}], + [[_prediction()]], + pure_dict=True, + ) + + assert response[0]["data"] == {"contact": "alice@example.com"} + + metadata = response[0]["model_metadata"] + assert type(metadata["fields"]) is dict + assert type(metadata["fields"]["contact"]) is dict + assert type(metadata["fields"]["contact"]["ner"]) is dict + assert type(metadata["entities"][E2F]) is dict + assert metadata["fields"]["contact"]["ner"]["labels"][0]["label"] == "email" diff --git a/tests/pii_replacer/test_metadata.py b/tests/pii_replacer/test_metadata.py new file mode 100644 index 000000000..54b1063b7 --- /dev/null +++ b/tests/pii_replacer/test_metadata.py @@ -0,0 +1,86 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from nemo_safe_synthesizer.pii_replacer.ner.metadata import ( + DatasetMetadata, + EntityMetadata, + EntitySummary, + FieldAttribute, + FieldMetadata, + TypeMetadata, +) + + +def test_metadata_dict_payloads_preserve_dataclass_shape(): + entity = EntityMetadata( + label="email", + count=2, + f_ratio=0.5, + approx_cardinality=2, + sources=["regex"], + field_label_f_ratio=1.0, + ) + field = FieldMetadata( + field="email_address", + count=2, + approx_cardinality=2, + missing=0, + pct_missing=0.0, + pct_total_unique=100.0, + s_score=1.0, + entities=[entity], + types=[TypeMetadata(type="str", count=2)], + field_labels=["email"], + field_attributes=[FieldAttribute.ID], + ) + entity_summary = EntitySummary( + label="email", + fields=["email_address"], + count=2, + approx_distinct_count=2, + sources=["regex"], + ) + metadata = DatasetMetadata(project_record_count=2, total_field_count=1) + metadata.add_field(field) + metadata.add_entity(entity_summary) + + expected_field = { + "field": "email_address", + "count": 2, + "approx_cardinality": 2, + "missing": 0, + "pct_missing": 0.0, + "pct_total_unique": 100.0, + "s_score": 1.0, + "entities": [ + { + "label": "email", + "count": 2, + "f_ratio": 0.5, + "approx_cardinality": 2, + "sources": ["regex"], + "field_label_f_ratio": 1.0, + } + ], + "types": [{"type": "str", "count": 2}], + "field_labels": ["email"], + "field_attributes": [FieldAttribute.ID], + } + expected_entity_summary = { + "label": "email", + "fields": ["email_address"], + "count": 2, + "approx_distinct_count": 2, + "sources": ["regex"], + } + + assert field.dict() == expected_field + assert entity_summary.dict() == expected_entity_summary + assert metadata.to_dict() == { + "project_record_count": 2, + "total_field_count": 1, + "data": { + "fields": [expected_field], + "entities": [expected_entity_summary], + }, + } diff --git a/tests/pii_replacer/test_ner_model.py b/tests/pii_replacer/test_ner_model.py new file mode 100644 index 000000000..d73c4e75f --- /dev/null +++ b/tests/pii_replacer/test_ner_model.py @@ -0,0 +1,128 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from typing import Any + +import pytest + +from nemo_safe_synthesizer.data_processing.records.fragment import NERRawPredictionPayload +from nemo_safe_synthesizer.data_processing.records.json_record import JSONRecord +from nemo_safe_synthesizer.pii_replacer.ner.model import Model +from nemo_safe_synthesizer.pii_replacer.ner.ner import NERPrediction +from nemo_safe_synthesizer.pii_replacer.ner.ner_mp import _extend_single_predictions +from nemo_safe_synthesizer.pii_replacer.ner.utils import input_to_json_records + + +def _prediction() -> NERRawPredictionPayload: + return { + "text": "alice@example.com", + "start": 0, + "end": 17, + "label": "email", + "source": "safe-synthesizer/email", + "score": 0.99, + "field": "contact", + "value_path": ("contact",), + "substring_match": None, + } + + +class _Timings: + def to_dict(self) -> dict[str, Any]: + return { + "records": 1, + "total_predictions": 1, + "total_time_ms": 2.0, + "total_time_ms_avg": 2.0, + "time_per_prediction_ms": 2.0, + "predictors": {}, + } + + +class _NERStub: + def __init__(self, result: list[list[NERRawPredictionPayload]]) -> None: + self.result = result + self.calls: list[tuple[object, dict[str, object]]] = [] + + def predict(self, input_data: object, **kwargs: object) -> object: + self.calls.append((input_data, kwargs)) + if kwargs["timings_only"]: + return _Timings() + return self.result + + +def _model_with_stub(result: list[list[NERRawPredictionPayload]]) -> tuple[Model, _NERStub]: + model = Model.__new__(Model) + stub = _NERStub(result) + model._ner = stub + return model, stub + + +def test_predict_dict_input_returns_api_response_list(): + model, stub = _model_with_stub([[_prediction()]]) + + response = model.predict({"contact": "alice@example.com"}) + + assert response[0]["data"] == {"contact": "alice@example.com"} + assert response[0]["model_metadata"]["fields"]["contact"]["ner"]["labels"][0]["label"] == "email" + assert stub.calls == [ + ( + [{"contact": "alice@example.com"}], + {"timings_only": False, "dict_result": True}, + ) + ] + + +def test_predict_string_input_returns_prediction_rows(): + prediction = _prediction() + model, stub = _model_with_stub([[prediction]]) + + response = model.predict("alice@example.com") + + assert response == [[prediction]] + assert stub.calls == [ + ( + ["alice@example.com"], + {"timings_only": False, "dict_result": True}, + ) + ] + + +def test_predict_timings_only_returns_timing_payload_for_dict_input(): + model, _stub = _model_with_stub([[_prediction()]]) + + response = model.predict({"contact": "alice@example.com"}, timings_only=True) + + assert response == { + "records": 1, + "total_predictions": 1, + "total_time_ms": 2.0, + "total_time_ms_avg": 2.0, + "time_per_prediction_ms": 2.0, + "predictors": {}, + } + + +def test_input_to_json_records_preserves_json_records_and_wraps_raw_records(): + existing = JSONRecord({"contact": "alice@example.com"}) + + records = input_to_json_records([existing, {"name": "Alice"}, "raw text"]) + + assert records[0] is existing + assert [record.original for record in records[1:]] == [{"name": "Alice"}, "raw text"] + + +def test_input_to_json_records_rejects_unsupported_list_items(): + bad_input: Any = [1] + + with pytest.raises(TypeError, match="Input data not supported"): + input_to_json_records(bad_input) + + +def test_extend_single_predictions_accepts_row_oriented_worker_output(): + prediction = NERPrediction("alice@example.com", 0, 17, "email", "regex", 0.99) + target: list[NERPrediction | dict[str, Any]] = [] + + _extend_single_predictions(target, [[prediction]]) + + assert target == [prediction]