diff --git a/.mise.toml b/.mise.toml index 04b145a77..8c21999c1 100644 --- a/.mise.toml +++ b/.mise.toml @@ -24,7 +24,7 @@ ripgrep = "latest" dprint = "latest" uv = "0.9.30" ruff = "0.15.0" -ty = "0.0.32" +ty = "0.0.44" "aqua:j178/prek" = "latest" diff --git a/mise.lock b/mise.lock index 04b34ec62..c83b7e0e3 100644 --- a/mise.lock +++ b/mise.lock @@ -131,22 +131,22 @@ checksum = "sha256:093d355ac33c6b8e91e80b8497d5581c61b028c0405e265cf38fd88f9a291 url = "https://github.com/astral-sh/ruff/releases/download/0.15.0/ruff-aarch64-apple-darwin.tar.gz" [[tools.ty]] -version = "0.0.32" +version = "0.0.44" backend = "aqua:astral-sh/ty" [tools.ty."platforms.linux-arm64"] -checksum = "sha256:de848acf867991f495dc346f5d12a7ac470af3ac464fdad7aff22fb8ee931a17" -url = "https://github.com/astral-sh/ty/releases/download/0.0.32/ty-aarch64-unknown-linux-musl.tar.gz" +checksum = "sha256:46a84ec41d2ae8892673f698a2ae31e02278203ededdce0d9229b7c4f5ac8bff" +url = "https://github.com/astral-sh/ty/releases/download/0.0.44/ty-aarch64-unknown-linux-musl.tar.gz" provenance = "github-attestations" [tools.ty."platforms.linux-x64"] -checksum = "sha256:cc58ee952aa551a0cbca495a43325b093dfcb180f0b314bf2c71068d37833b9e" -url = "https://github.com/astral-sh/ty/releases/download/0.0.32/ty-x86_64-unknown-linux-musl.tar.gz" +checksum = "sha256:618549cc5b0fd19b3ca0830a43fee4189623d7bd993bf65315f1f0ba650d944a" +url = "https://github.com/astral-sh/ty/releases/download/0.0.44/ty-x86_64-unknown-linux-musl.tar.gz" provenance = "github-attestations" [tools.ty."platforms.macos-arm64"] -checksum = "sha256:6b03b94d8c2ddcb5db67e6c863ef3c72b83fbcb973e0ff3a4c6daa4352dad009" -url = "https://github.com/astral-sh/ty/releases/download/0.0.32/ty-aarch64-apple-darwin.tar.gz" +checksum = "sha256:e796d5a91886a379d1da7c97772722a205901991d642382f75e2c27c138cfddb" +url = "https://github.com/astral-sh/ty/releases/download/0.0.44/ty-aarch64-apple-darwin.tar.gz" provenance = "github-attestations" [[tools.uv]] diff --git a/pyproject.toml b/pyproject.toml index 58b09c3bd..64b541972 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -361,6 +361,8 @@ safe-synthesizer = "nemo_safe_synthesizer.cli.cli:cli" # Platform-specific imports (e.g. vLLM on macOS) use ty: ignore[unresolved-import] # that appear unused when the package is installed. Suppress the noise. unused-ignore-comment = "ignore" +missing-override-decorator = "error" + [tool.ty.environment] extra-paths = ["typings"] @@ -368,8 +370,6 @@ 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/", - "./tools/*.py" + "./uv-cache", + "./docs/**/*.ipynb", ] diff --git a/script/slurm/nss_top.py b/script/slurm/nss_top.py index 4cda42225..f6630395c 100644 --- a/script/slurm/nss_top.py +++ b/script/slurm/nss_top.py @@ -29,6 +29,7 @@ from textual.containers import Vertical from textual.reactive import reactive from textual.widgets import DataTable, Footer, Header, RichLog, Static +from typing_extensions import override # squeue format fields and matching column headers _SQUEUE_FIELDS = ["%i", "%j", "%T", "%P", "%M", "%l", "%D", "%C", "%R"] @@ -176,6 +177,7 @@ def __init__(self, user: str, log_dir: Path | None, refresh_secs: int) -> None: # ── Layout ──────────────────────────────────────────────────────────────── + @override def compose(self) -> ComposeResult: yield Header(show_clock=True) with Vertical(): diff --git a/src/nemo_safe_synthesizer/cli/artifact_structure.py b/src/nemo_safe_synthesizer/cli/artifact_structure.py index 8d6b96799..2750885f8 100644 --- a/src/nemo_safe_synthesizer/cli/artifact_structure.py +++ b/src/nemo_safe_synthesizer/cli/artifact_structure.py @@ -27,6 +27,8 @@ from pathlib import Path from typing import TYPE_CHECKING, Generic, Self, TypeVar, overload +from typing_extensions import override + from ..observability import get_logger from ..utils import write_json @@ -249,17 +251,21 @@ def path(self) -> Path: """The resolved directory path.""" return self._path + @override def __fspath__(self) -> str: """Support os.fspath() for use with open(), etc.""" return str(self._path) + @override def __str__(self) -> str: """Return the path as a string.""" return str(self._path) + @override def __repr__(self) -> str: return f"BoundDir({self._path!r})" + @override def __eq__(self, other: object) -> bool: """Support comparison with Path objects.""" if isinstance(other, BoundDir): @@ -268,6 +274,7 @@ def __eq__(self, other: object) -> bool: return self._path == other return NotImplemented + @override def __hash__(self) -> int: return hash(self._path) diff --git a/src/nemo_safe_synthesizer/cli/datasets.py b/src/nemo_safe_synthesizer/cli/datasets.py index 7b8662aed..6c6cb8a4a 100644 --- a/src/nemo_safe_synthesizer/cli/datasets.py +++ b/src/nemo_safe_synthesizer/cli/datasets.py @@ -21,6 +21,16 @@ logger = get_logger(__name__) +def _require_dataframe(value: object, *, url: str) -> pd.DataFrame: + if isinstance(value, pd.DataFrame): + return value + raise TypeError(f"Expected dataset reader for {url} to return a pandas DataFrame, got {type(value).__name__}") + + +def _dynamic_callable(value: object) -> Any: + return value + + class DatasetInfo(BaseModel): """Entry in the dataset registry.""" @@ -96,20 +106,22 @@ def fetch(self) -> pd.DataFrame: logger.info(f"Reading dataset from {url}") # Determine the file extension and appropriate reader - match Path(url).suffix.lstrip("."): + reader: Any + extension = Path(url).suffix.lstrip(".") + match extension: case "csv" | "txt": - reader = pd.read_csv + reader = _dynamic_callable(pd.read_csv) default_load_args: dict[str, Any] = {} case "json": - reader = pd.read_json + reader = _dynamic_callable(pd.read_json) default_load_args = {} case "jsonl": - reader = pd.read_json + reader = _dynamic_callable(pd.read_json) default_load_args = {"lines": True} case "parquet": - reader = pd.read_parquet + reader = _dynamic_callable(pd.read_parquet) default_load_args = {} - case extension: + case _: if not extension: extension = f"" raise ValueError(f"Unsupported file extension: {extension}") @@ -118,7 +130,7 @@ def fetch(self) -> pd.DataFrame: final_load_args = {**default_load_args, **(self.load_args or {})} try: - return reader(url, **final_load_args) # ty: ignore[invalid-argument-type] -- reader union includes parquet which has stricter signature + return _require_dataframe(reader(url, **final_load_args), url=url) except Exception as e: logger.error(f"Error reading dataset from {url}: {e}", exc_info=True) raise diff --git a/src/nemo_safe_synthesizer/cli/run.py b/src/nemo_safe_synthesizer/cli/run.py index 0d4cfb7cd..2bd0fe6c3 100644 --- a/src/nemo_safe_synthesizer/cli/run.py +++ b/src/nemo_safe_synthesizer/cli/run.py @@ -9,7 +9,7 @@ import sys from collections.abc import Callable from pathlib import Path -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, ParamSpec, TypeVar import click @@ -35,8 +35,11 @@ from ..sdk.library_builder import SafeSynthesizer from .artifact_structure import Workdir +P = ParamSpec("P") +R = TypeVar("R") -def common_run_options(f: Callable[..., object]) -> Callable[..., object]: + +def common_run_options(f: Callable[P, R]) -> Callable[P, R]: """Decorator to add common options for run commands. Apply this above ``@pydantic_options`` in source order. Python applies diff --git a/src/nemo_safe_synthesizer/cli/utils.py b/src/nemo_safe_synthesizer/cli/utils.py index 4d6908cdf..6084537ca 100644 --- a/src/nemo_safe_synthesizer/cli/utils.py +++ b/src/nemo_safe_synthesizer/cli/utils.py @@ -24,6 +24,7 @@ from pydantic import ValidationError from ..config import SafeSynthesizerParameters +from ..config.parameters import ConfigPatch from ..defaults import DEFAULT_ARTIFACTS_PATH from ..observability import configure_logging_from_workdir, get_logger, initialize_observability from ..utils import merge_dicts @@ -431,7 +432,7 @@ def _initialize_logging_for_cli_from_settings( return run_logger -def merge_overrides(config_path: str | Path | None, overrides: dict) -> SafeSynthesizerParameters: +def merge_overrides(config_path: str | Path | None, overrides: ConfigPatch) -> SafeSynthesizerParameters: """Merge overrides into a SafeSynthesizerParameters object. If config_path is None, use the overrides to create a new SafeSynthesizerParameters object. @@ -446,11 +447,9 @@ def merge_overrides(config_path: str | Path | None, overrides: dict) -> SafeSynt """ try: if config_path is None: - my_config = SafeSynthesizerParameters.model_validate(overrides) + my_config = SafeSynthesizerParameters.from_config_patch(overrides) else: - file_config = SafeSynthesizerParameters.from_yaml(config_path).model_dump(exclude_unset=True) - params = merge_dicts(file_config, overrides) - my_config = SafeSynthesizerParameters.model_validate(params) + my_config = SafeSynthesizerParameters.from_yaml(config_path).with_config_patch(overrides) except ValidationError as e: click.echo(f"{config_path} is invalid:\n{e}") sys.exit(1) diff --git a/src/nemo_safe_synthesizer/config/autoconfig.py b/src/nemo_safe_synthesizer/config/autoconfig.py index cab8c3877..fb7770748 100644 --- a/src/nemo_safe_synthesizer/config/autoconfig.py +++ b/src/nemo_safe_synthesizer/config/autoconfig.py @@ -20,8 +20,7 @@ from ..defaults import DEFAULT_MAX_SEQ_LENGTH, MAX_ROPE_SCALING_FACTOR from ..llm.metadata import ModelMetadata from ..observability import get_logger -from ..utils import merge_dicts -from .parameters import SafeSynthesizerParameters +from .parameters import ConfigPatch, SafeSynthesizerParameters from .types import AUTO_STR if TYPE_CHECKING: @@ -304,14 +303,13 @@ def _build_updated_params( Returns: The validated SafeSynthesizerParameters. """ - new_params = { + new_params: ConfigPatch = { "training": training_params, "data": data_params, "privacy": privacy_params, } - updated_params = merge_dicts(self._config.model_dump(exclude_unset=True), new_params) - logger.debug(f"params to update: {updated_params}") - my_config = SafeSynthesizerParameters.model_validate(updated_params) + logger.debug(f"params to update: {new_params}") + my_config = self._config.with_config_patch(new_params) logger.debug(f"auto-updated config: {my_config.model_dump(exclude_unset=True)}") return my_config diff --git a/src/nemo_safe_synthesizer/config/base.py b/src/nemo_safe_synthesizer/config/base.py index 97cd6458e..d62f2166d 100644 --- a/src/nemo_safe_synthesizer/config/base.py +++ b/src/nemo_safe_synthesizer/config/base.py @@ -9,6 +9,7 @@ from typing import Any from pydantic import BaseModel, ConfigDict +from typing_extensions import override __all__ = ["NSSBaseModel", "LRScheduler", "pydantic_model_config"] @@ -34,6 +35,7 @@ class NSSBaseModel(BaseModel): model_config = pydantic_model_config + @override def dict(self, **kwargs: Any) -> dict[str, Any]: # ty: ignore[invalid-type-form] -- method name shadows builtin dict; backward-compat shim for model_dump """Return a dict representation via ``model_dump`` for backward compatibility.""" return self.model_dump(**kwargs) diff --git a/src/nemo_safe_synthesizer/config/parameters.py b/src/nemo_safe_synthesizer/config/parameters.py index 9829de50b..d39cc7245 100644 --- a/src/nemo_safe_synthesizer/config/parameters.py +++ b/src/nemo_safe_synthesizer/config/parameters.py @@ -4,11 +4,14 @@ from __future__ import annotations import warnings -from typing import Any, Self +from collections.abc import Mapping +from typing import Any, Self, TypeAlias from pydantic import BaseModel, Field, model_validator +from typing_extensions import override from ..configurator.parameters import Parameters +from ..data_processing.records.json_types import JsonValue from ..errors import ParameterError from ..observability import get_logger from ..telemetry import _telemetry_enabled @@ -23,13 +26,16 @@ from .training import TrainingHyperparams from .types import AUTO_STR -__all__ = ["SafeSynthesizerParameters"] +ConfigPatch: TypeAlias = Mapping[str, JsonValue] +_SectionPatch: TypeAlias = dict[str, object] + +__all__ = ["ConfigPatch", "SafeSynthesizerParameters"] logger = get_logger(__name__) -def _collect_set_fields(model: BaseModel) -> dict[str, Any]: +def _collect_set_fields(model: BaseModel) -> _SectionPatch: """Recursively collect a model's explicitly-set fields as a nested dict. Unlike ``model_dump(exclude_unset=True)``, nested models are always @@ -38,7 +44,7 @@ def _collect_set_fields(model: BaseModel) -> dict[str, Any]: ``cfg.generation.validation.foo = True``) are captured. A nested model is included only when it has at least one set field of its own. """ - overrides: dict[str, Any] = {} + overrides: _SectionPatch = {} for name in type(model).model_fields: value = getattr(model, name) if isinstance(value, BaseModel): @@ -192,6 +198,7 @@ def check_timeseries_group_column(self) -> Self: return self @classmethod + @override def from_params(cls, **kwargs) -> "SafeSynthesizerParameters": """Convert singular, flat parameters to nested structure. @@ -237,6 +244,21 @@ def from_params(cls, **kwargs) -> "SafeSynthesizerParameters": extra["emit_telemetry"] = kwargs["emit_telemetry"] return cls(**extra) + @classmethod + def from_config_patch(cls, patch: ConfigPatch) -> Self: + """Validate a sparse top-level config patch as a full configuration.""" + return cls.model_validate(patch) + + def with_config_patch(self, patch: ConfigPatch) -> Self: + """Apply a sparse top-level config patch and revalidate the result. + + Only fields explicitly set on ``self`` are carried into the merge before + applying ``patch``. This preserves file/CLI precedence while keeping + default values implicit for future ``exclude_unset`` dumps. + """ + params = merge_dicts(self.model_dump(exclude_unset=True), patch) + return type(self).model_validate(params) + def with_runtime_overrides(self, runtime: SafeSynthesizerParameters) -> "SafeSynthesizerParameters": """Apply resume-time generation/evaluation/telemetry overrides onto a copy of self. diff --git a/src/nemo_safe_synthesizer/configurator/parameter.py b/src/nemo_safe_synthesizer/configurator/parameter.py index 8d26af89d..7b3f7510a 100644 --- a/src/nemo_safe_synthesizer/configurator/parameter.py +++ b/src/nemo_safe_synthesizer/configurator/parameter.py @@ -13,10 +13,12 @@ import operator from collections.abc import Callable, Sequence from dataclasses import dataclass -from typing import Any, Generic, TypeVar, cast, get_args +from types import NotImplementedType +from typing import Any, Generic, TypeVar, get_args from pydantic import BaseModel, GetCoreSchemaHandler, model_serializer from pydantic_core import core_schema +from typing_extensions import override DataT = TypeVar( "DataT", bound=(int | float | str | bytes | bool | None | Sequence[int | float | str | bytes | bool | BaseModel]) @@ -55,6 +57,7 @@ def ser_model(self) -> "dict[str, DataT] | DataT | Sequence[DataT] | Parameter[D else: return self + @override def __str__(self): return self.__repr__() @@ -86,7 +89,7 @@ def __get_pydantic_core_schema__(cls, source: Any, handler: GetCoreSchemaHandler non_instance_schema = core_schema.no_info_before_validator_function(cls, sequence_t_schema) return core_schema.union_schema([instance_schema, non_instance_schema]) - def _comp_helper(self, other: "Parameter[DataT] | DataT", op: Callable[[Any, Any], bool]) -> bool | None: + def _comp_helper(self, other: object, op: Callable[[Any, Any], bool]) -> bool | NotImplementedType: """Apply a comparison operator ``op`` to ``self.value`` and the unwrapped value of ``other``.""" match other: case Parameter(value=y) if isinstance(self.value, type(y)): @@ -108,6 +111,7 @@ def __gt__(self, other: "Parameter[DataT] | DataT") -> bool | None: def __lt__(self, other: "Parameter[DataT] | DataT") -> bool | None: return self._comp_helper(other, operator.__lt__) + @override def __eq__(self, other: object) -> bool: - result = self._comp_helper(cast("Parameter[DataT] | DataT", other), operator.__eq__) - return cast(bool, result) + result = self._comp_helper(other, operator.__eq__) + return False if result is NotImplemented else result diff --git a/src/nemo_safe_synthesizer/configurator/parameters.py b/src/nemo_safe_synthesizer/configurator/parameters.py index 5e5211666..b6e6fcc79 100644 --- a/src/nemo_safe_synthesizer/configurator/parameters.py +++ b/src/nemo_safe_synthesizer/configurator/parameters.py @@ -27,6 +27,7 @@ from pydantic import ( BaseModel, ) +from typing_extensions import override from ..config.base import ( pydantic_model_config, @@ -55,6 +56,7 @@ def _isparams(self): return True @classmethod + @override def __subclasshook__(cls, c): """Enable ``isinstance()`` checks via duck typing on the ``_isparams`` marker.""" if cls is Parameters: @@ -109,6 +111,7 @@ def _iter_parameters(self, recursive: bool = True) -> Generator[Mapping[str, Any for pg in param_groups: yield from pg._iter_parameters(recursive=True) + @override def __iter__(self) -> Iterator[Mapping[str, Any]]: # ty: ignore[invalid-method-override] -- intentionally overrides pydantic BaseModel.__iter__ with parameter-group semantics """Iterate over all parameters, recursing into nested groups.""" return self._iter_parameters(recursive=True) diff --git a/src/nemo_safe_synthesizer/configurator/pydantic_click_options.py b/src/nemo_safe_synthesizer/configurator/pydantic_click_options.py index d24b7bf9b..3c1ed670c 100644 --- a/src/nemo_safe_synthesizer/configurator/pydantic_click_options.py +++ b/src/nemo_safe_synthesizer/configurator/pydantic_click_options.py @@ -20,13 +20,14 @@ import inspect import types +from collections.abc import Callable from dataclasses import dataclass -from typing import Annotated, Any, Literal, Union, get_args, get_origin +from typing import Annotated, Any, Literal, TypeVar, Union, get_args, get_origin import click from pydantic import BaseModel from pydantic.fields import FieldInfo -from typing_extensions import TypeIs +from typing_extensions import TypeIs, override from ..config.types import AUTO_STR @@ -102,6 +103,7 @@ class FlagParam: ClickParam = LeafParam | FlagParam +F = TypeVar("F", bound=Callable[..., object]) # --------------------------------------------------------------------------- @@ -113,7 +115,7 @@ def _is_basemodel(t: Any) -> TypeIs[type[BaseModel]]: return inspect.isclass(t) and issubclass(t, BaseModel) -def _nullable_model_arg(union_args: tuple) -> type[BaseModel] | None: +def _nullable_model_arg(union_args: tuple[object, ...]) -> type[BaseModel] | None: """Return the BaseModel member of a ``SomeModel | None`` union, or ``None``.""" return next((a for a in union_args if a is not type(None) and _is_basemodel(a)), None) @@ -148,6 +150,7 @@ def __init__(self, base_type: click.ParamType) -> None: self.base_type = base_type self.name = f"{base_type.name}|{AUTO_STR}" + @override def convert( self, value: str, @@ -175,12 +178,12 @@ def convert( return self.base_type.convert(value, param, ctx) -def _has_string_literal(args: set) -> bool: +def _has_string_literal(args: set[object]) -> bool: """Check if any member is a ``Literal`` containing a string value.""" return any(get_origin(a) is Literal and any(isinstance(v, str) for v in get_args(a)) for a in args) -def _is_auto_only_literal_union(args: set) -> bool: +def _is_auto_only_literal_union(args: set[object]) -> bool: """Check that the union's string-valued ``Literal`` members are exactly ``{AUTO_STR}``. Returns ``True`` only if every string-valued Literal member contributes the @@ -198,9 +201,9 @@ def _is_auto_only_literal_union(args: set) -> bool: return string_values == {AUTO_STR} -def _literal_value_types(args: set) -> set: +def _literal_value_types(args: set[object]) -> set[object]: """Replace non-string ``Literal`` annotations with their value types.""" - normalized: set = set() + normalized: set[object] = set() for arg in args: if get_origin(arg) is Literal: normalized.update(type(value) for value in get_args(arg)) @@ -224,7 +227,7 @@ def _click_type(annotation: Any) -> click.ParamType: t = annotation if get_origin(t) is Annotated: t = get_args(t)[0] - args = set(get_args(t)) if get_origin(t) in (Union, types.UnionType) else {t} + args: set[object] = set(get_args(t)) if get_origin(t) in (Union, types.UnionType) else {t} args.discard(type(None)) if _has_string_literal(args): # Auto*Param: Literal["auto"] | -- wrap in AutoParamType so @@ -289,7 +292,7 @@ def _collect_params(cls: type[BaseModel], prefix: str = "") -> list[ClickParam]: # --------------------------------------------------------------------------- -def pydantic_options(model_class: type[BaseModel], field_separator: str = "__"): +def pydantic_options(model_class: type[BaseModel], field_separator: str = "__") -> Callable[[F], F]: """Decorate a Click command with options derived from a Pydantic model. Recurses into nested sub-models, flattening their fields into top-level @@ -309,7 +312,7 @@ def pydantic_options(model_class: type[BaseModel], field_separator: str = "__"): A Click decorator that attaches the generated options to a command. """ - def decorator(f): + def decorator(f: F) -> F: for param in sorted(_collect_params(model_class), key=lambda p: p.name): match param: case FlagParam(field_name=field_name): diff --git a/src/nemo_safe_synthesizer/data_processing/actions/data_actions.py b/src/nemo_safe_synthesizer/data_processing/actions/data_actions.py index a601cd1ca..dc3a0e06a 100644 --- a/src/nemo_safe_synthesizer/data_processing/actions/data_actions.py +++ b/src/nemo_safe_synthesizer/data_processing/actions/data_actions.py @@ -11,7 +11,6 @@ from __future__ import annotations -import json import operator from abc import ABC, abstractmethod from collections import defaultdict @@ -27,7 +26,6 @@ Protocol, TypeVar, Union, - cast, ) import pandas as pd @@ -41,6 +39,7 @@ ValidationInfo, model_validator, ) +from typing_extensions import override from ... import utils from ...observability import get_logger @@ -230,7 +229,7 @@ def set_state(self, state_obj: BaseModel) -> None: def get_state(self, state_obj_type: type[BaseModelT]) -> BaseModelT: """Retrieve and deserialize a previously persisted state object.""" state_obj_json = self._ctx.state[self.hash()] - return state_obj_type.model_validate(json.loads(state_obj_json)) + return state_obj_type.model_validate_json(state_obj_json) DEFAULT_ACTION_CTX = ActionCtx() @@ -252,6 +251,7 @@ class GenerateAction(BaseAction, ABC): phase: ProcessPhase = ProcessPhase.GENERATE + @override def functions(self) -> Functions: """Route ``generate`` to the correct phase slot based on ``self.phase``.""" fns = Functions() @@ -265,6 +265,7 @@ def functions(self) -> Functions: return fns @abstractmethod + @override def generate(self, df: pd.DataFrame) -> pd.DataFrame: """Generate new data based on the existing data in the DataFrame.""" ... @@ -319,6 +320,7 @@ def validate_model(self) -> "BaseAction": return self + @override def generate(self, df: pd.DataFrame) -> pd.DataFrame: df = self._ctx.transforms_util.execute_col_updates(self.col, df, self._expressions) if self.dtype is not None: @@ -338,6 +340,7 @@ class GenRawExpression(GenerateAction): type_: Literal["gen_raw_expression"] = "gen_raw_expression" expressions: list[TransformsUpdate] = [] + @override def generate(self, df: pd.DataFrame) -> pd.DataFrame: return self._ctx.transforms_util.execute_updates(df, self.expressions) @@ -347,6 +350,7 @@ class GenDistribution(GenerateAction): col: str distribution: DistributionT + @override def generate(self, df: pd.DataFrame) -> pd.DataFrame: df[self.col] = self.distribution.sample(num_records=len(df)) return df @@ -361,6 +365,7 @@ def __init__(self, /, **data: Any) -> None: super().__init__(**data) self.data_source = self.data_source.with_ctx(self._ctx) + @override def generate(self, df: pd.DataFrame) -> pd.DataFrame: return self.data_source.generate_data(col=self.col, df=df) @@ -382,15 +387,20 @@ def __init__(self, /, **data: Any) -> None: super().__init__(**data) self.data_source = self.data_source.with_ctx(self._ctx) + @override def preprocess(self, df: pd.DataFrame) -> pd.DataFrame: if self.col in df.columns: - column_index = cast(int, df.columns.get_loc(self.col)) + location = df.columns.get_loc(self.col) + if not isinstance(location, int): + raise ValueError(f"Column {self.col!r} must be unique to replace data source") + column_index = location else: column_index = None self.set_state(self.State(column_index=column_index)) return df.drop(columns=[self.col], errors="ignore") + @override def generate(self, df: pd.DataFrame) -> pd.DataFrame: column_index = self.get_state(self.State).column_index df = self.data_source.generate_data(col=self.col, df=df) @@ -412,6 +422,7 @@ class GenDatetimeDistribution(GenerateAction): col: str distribution: DatetimeDistributionT + @override def generate(self, df: pd.DataFrame) -> pd.DataFrame: df[self.col] = self.distribution.sample(num_records=len(df)) return df @@ -426,6 +437,7 @@ def __init__(self, /, **data: Any) -> None: super().__init__(**data) self._action = GenDataSource(col=self.col, data_source=UniqueIdSource(id_type=self.id_type)).with_ctx(self._ctx) + @override def generate(self, df: pd.DataFrame) -> pd.DataFrame: return self._action.generate(df) @@ -444,6 +456,7 @@ def validate_faker_fn(self) -> "BaseAction": return self + @override def generate(self, df: pd.DataFrame) -> pd.DataFrame: fn = getattr(self._ctx.transforms_util.env._fake, self.faker_fn) df[self.col] = df.apply(lambda _: fn(), axis=1) @@ -452,6 +465,7 @@ def generate(self, df: pd.DataFrame) -> pd.DataFrame: class ValidationAction(BaseAction, ABC): @abstractmethod + @override def _validate_batch(self, batch: pd.DataFrame, df: pd.DataFrame) -> pd.Series: ... @@ -459,6 +473,7 @@ class DropExpression(ValidationAction): type_: Literal["expression_drop"] = "expression_drop" conditions: list[str] = [] + @override def _validate_batch(self, batch: pd.DataFrame, df: pd.DataFrame) -> pd.Series: batch[MetadataColumns.INDEX] = range(len(batch)) batch_copy = batch.copy() @@ -473,6 +488,7 @@ def _validate_batch(self, batch: pd.DataFrame, df: pd.DataFrame) -> pd.Series: class DropDuplicates(ValidationAction): type_: Literal["drop_duplicates"] = "drop_duplicates" + @override def _validate_batch(self, batch: pd.DataFrame, df: pd.DataFrame) -> pd.Series: return ~batch.isin(df).all(axis=1) @@ -483,6 +499,7 @@ class DateConstraint(BaseAction): colB: str operator: Literal["gt", "ge", "lt", "le"] + @override def _validate_batch(self, batch: pd.DataFrame, df: pd.DataFrame) -> pd.Series: """ Filter out all rows where the operator isn't true. The type of the @@ -535,6 +552,7 @@ def _infer_col_dt_format(self, col: pd.Series) -> Optional[str]: return modes.iloc[0] + @override def preprocess(self, df: pd.DataFrame) -> pd.DataFrame: # Retrieve the datetime format, either from the user-config or by inferring it from the column dt_format = self.format @@ -557,6 +575,7 @@ def preprocess(self, df: pd.DataFrame) -> pd.DataFrame: return df + @override def _validate_batch(self, batch: pd.DataFrame, df: pd.DataFrame) -> pd.Series: # Retrieve datetime format, either from the instance or from the state dt_format = self.format or self.get_state(self.State).dt_format @@ -572,6 +591,7 @@ class CategoricalCol(ColAction): type_: Literal["categorical"] = "categorical" values: list[str | int | float] + @override def _validate_batch(self, batch: pd.DataFrame, df: pd.DataFrame) -> pd.Series: return batch[self.name].isin(self.values) diff --git a/src/nemo_safe_synthesizer/data_processing/actions/dates.py b/src/nemo_safe_synthesizer/data_processing/actions/dates.py index 3575391d4..2eba869b5 100644 --- a/src/nemo_safe_synthesizer/data_processing/actions/dates.py +++ b/src/nemo_safe_synthesizer/data_processing/actions/dates.py @@ -415,7 +415,7 @@ def fit_and_transform_dates( names to ``{"format": ..., "min": ...}`` dicts needed by ``transform_dates`` for reversal. """ - date_min_dict = {} + date_min_dict: dict[str, dict[str, str]] = {} object_cols = [col for col, col_type in df.dtypes.items() if col_type == "object"] result_df = df.copy() if not inplace else df for object_col in object_cols: @@ -428,7 +428,7 @@ def fit_and_transform_dates( dates = pd.to_datetime(result_df[object_col], format=inferred_format) min_date = dates.min() result_df[object_col] = (dates - min_date).dt.total_seconds() - date_min_dict[object_col] = { + date_min_dict[str(object_col)] = { "format": inferred_format, "min": str(min_date), } diff --git a/src/nemo_safe_synthesizer/data_processing/actions/distributions.py b/src/nemo_safe_synthesizer/data_processing/actions/distributions.py index 2d9c6e499..cbfc3cfd8 100644 --- a/src/nemo_safe_synthesizer/data_processing/actions/distributions.py +++ b/src/nemo_safe_synthesizer/data_processing/actions/distributions.py @@ -13,11 +13,11 @@ from abc import ABC, abstractmethod from datetime import datetime, timedelta -from functools import partial from typing import Annotated, Any, Literal, Optional, Union import numpy as np from pydantic import BaseModel, Field +from typing_extensions import override class Distribution(BaseModel, ABC): @@ -77,21 +77,20 @@ def _round_datetime(self, dt: datetime, precision: timedelta) -> datetime: rounded_ts = round(dt.timestamp() / precision.total_seconds()) * precision.total_seconds() return datetime.fromtimestamp(rounded_ts) - def _apply_universal_params(self, samples: list[datetime]) -> list[datetime]: - ret: list[str] | list[datetime] = samples - - ops = [] - if self.precision is not None: - ops.append(partial(self._round_datetime, self.precision)) + def _apply_universal_params(self, samples: list[datetime]) -> list[datetime] | list[str]: if self.format is not None: - ops.append(lambda x: x.strftime(self.format)) - + formatted: list[str] = [] + for sample in samples: + if self.precision is not None: + sample = self._round_datetime(sample, self.precision) + formatted.append(sample.strftime(self.format)) + return formatted + + ret: list[datetime] = [] for sample in samples: - n = sample - for op in ops: - n = op(n) - ret.append(n) - + if self.precision is not None: + sample = self._round_datetime(sample, self.precision) + ret.append(sample) return ret @@ -100,6 +99,7 @@ class GaussianDistribution(Distribution): mean: float std_dev: float = Field(gt=0) + @override def sample(self, num_records: int) -> list[float]: return np.random.normal(loc=self.mean, scale=self.std_dev, size=num_records).tolist() @@ -109,6 +109,7 @@ class DatetimeGaussianDistribution(DatetimeDistribution): mean: datetime std_dev: timedelta + @override def sample_datetimes(self, num_records: int) -> list[datetime]: float_samples = GaussianDistribution(mean=self.mean.timestamp(), std_dev=self.std_dev.total_seconds()).sample( num_records=num_records @@ -121,6 +122,7 @@ class UniformDistribution(Distribution): low: float high: float + @override def sample(self, num_records: int) -> list[float]: return np.random.uniform(low=self.low, high=self.high, size=num_records).tolist() @@ -130,6 +132,7 @@ class DatetimeUniformDistribution(DatetimeDistribution): low: datetime high: datetime + @override def sample_datetimes(self, num_records: int) -> list[datetime]: float_samples = UniformDistribution(low=self.low.timestamp(), high=self.high.timestamp()).sample( num_records=num_records diff --git a/src/nemo_safe_synthesizer/data_processing/actions/utils.py b/src/nemo_safe_synthesizer/data_processing/actions/utils.py index 3f7e29320..516d9d7e5 100644 --- a/src/nemo_safe_synthesizer/data_processing/actions/utils.py +++ b/src/nemo_safe_synthesizer/data_processing/actions/utils.py @@ -33,6 +33,7 @@ Field, PrivateAttr, ) +from typing_extensions import override from .dates import parse_date @@ -183,6 +184,7 @@ class UniqueIdSource(DataSource): id_type: Literal["uuid4"] = "uuid4" + @override def generate_data(self, df: pd.DataFrame, col: str = "newcol") -> pd.DataFrame: id_fn: Callable[[Any], Any] = { "uuid4": lambda _: str(uuid.uuid4()), @@ -196,6 +198,7 @@ class ExpressionSource(DataSource): expression: str + @override def generate_data(self, df: pd.DataFrame, col: str = "newcol") -> pd.DataFrame: return self._ctx.transforms_util.execute_col_updates(col, df, [self.expression]) diff --git a/src/nemo_safe_synthesizer/data_processing/assembler.py b/src/nemo_safe_synthesizer/data_processing/assembler.py index 8ec12879b..c70d78d99 100644 --- a/src/nemo_safe_synthesizer/data_processing/assembler.py +++ b/src/nemo_safe_synthesizer/data_processing/assembler.py @@ -20,6 +20,7 @@ from datasets.exceptions import DatasetGenerationError from tqdm.auto import tqdm from transformers import PreTrainedTokenizer +from typing_extensions import override from .. import utils from ..config.parameters import SafeSynthesizerParameters @@ -564,15 +565,18 @@ def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) @property + @override def num_records_train(self) -> int: """Number of records in the training split.""" return len(self.training_dataset) @property + @override def num_records_validation(self) -> int: """Number of records in the validation split.""" return 0 if self.validation_dataset is None else len(self.validation_dataset) + @override def _preprocess_before_splitting(self, tokenized_records: Dataset) -> Dataset: """Tabular data examples do not require any preprocessing before splitting.""" return tokenized_records @@ -687,6 +691,7 @@ def _prepare_dataset_for_training( res = self._run_example_generation(self._fill_context_with_records_generator, processed_dataset) return res + @override def assemble_training_examples(self, data_fraction: float = 1.0) -> TrainingExamples: """Build examples with randomly shuffled records. @@ -917,6 +922,7 @@ def _build_schema_prompt_excluding_pseudo_group( ) self.schema_prompt_ids = tokenizer(self.schema_prompt, add_special_tokens=False)["input_ids"] + @override def _preprocess_before_splitting(self, tokenized_records: Dataset) -> Dataset: """No preprocessing needed - sorting is done after split in _apply_grouped_train_test_split.""" return tokenized_records @@ -989,6 +995,7 @@ def _get_initial_prefill(self) -> dict[str, str]: # Convert lists to joined strings return {group: " " + "\n".join(samples) for group, samples in seen_groups.items()} + @override def _apply_train_test_split(self, dataset: Dataset) -> None: """Override split logic to preserve record order and split along group boundaries.""" self._apply_grouped_train_test_split(dataset) @@ -1036,6 +1043,7 @@ def _apply_grouped_train_test_split(self, dataset: Dataset) -> None: validation_dataset.info.description += "is_val" self.validation_dataset = _maybe_add_row_indices(validation_dataset, self._ROW_INDEX_COLUMN) + @override def _prepare_dataset_for_training( self, dataset: Dataset | None, data_fraction: float, rng: np.random.Generator ) -> Dataset | None: @@ -1079,6 +1087,7 @@ def _prepare_dataset_for_training( return self._run_example_generation(self._fill_context_with_records_generator, processed_dataset) + @override def assemble_training_examples(self, data_fraction: float = 1.0) -> TrainingExamples: """Build examples preserving sequential order within groups. @@ -1152,6 +1161,7 @@ def _flush_example(self, dataset: Dataset, start: int, end: int, stats_target: d stats_target["tokens_per_example"].update(example.num_tokens) return example.to_dict() + @override def _fill_context_with_records_generator(self, dataset: Dataset) -> GeneratorType: """Generate ordered examples, flushing at group and dataset boundaries. @@ -1342,11 +1352,13 @@ def __init__( self.validation_dataset = None @property + @override def num_records_train(self) -> int: """Total number of individual records across all training groups.""" return sum(self.training_dataset["num_records"]) @property + @override def num_records_validation(self) -> int: """Total number of individual records across all validation groups.""" return 0 if self.validation_dataset is None else sum(self.validation_dataset["num_records"]) @@ -1361,6 +1373,7 @@ def num_groups_validation(self) -> int: """Number of groups in the validation split.""" return 0 if self.validation_dataset is None else len(self.validation_dataset) + @override def _preprocess_before_splitting(self, tokenized_records: Dataset) -> Dataset: """Group and order the tokenized records before splitting the dataset.""" if self.order_by is not None: @@ -1538,6 +1551,7 @@ def _prepare_dataset_for_training( return self._run_example_generation(self._fill_context_with_groups_generator, processed_dataset) + @override def assemble_training_examples(self, data_fraction: float = 1.0) -> TrainingExamples: """Build examples with grouped (and optionally ordered) records. diff --git a/src/nemo_safe_synthesizer/data_processing/record_utils.py b/src/nemo_safe_synthesizer/data_processing/record_utils.py index 324d68041..03060777d 100644 --- a/src/nemo_safe_synthesizer/data_processing/record_utils.py +++ b/src/nemo_safe_synthesizer/data_processing/record_utils.py @@ -13,22 +13,28 @@ import json import re import time -from collections.abc import Callable +from collections.abc import Callable, Mapping, Sequence from csv import QUOTE_NONNUMERIC from dataclasses import dataclass, field from datetime import datetime from io import StringIO +from typing import Any import jsonschema import pandas as pd from ..observability import get_logger +from .records.json_types import JsonSchema, JsonValue, is_json_object RECORD_REGEX_PATTERN = r"{.+?}(?:\n|$)" RECORD_REGEX_PATTEN_LOOKAHEAD = r"{.+?}(?=\n|$)" logger = get_logger() +RecordDict = dict[str, Any] +RecordMapping = Mapping[str, Any] +RawRecordMapping = Mapping[Any, Any] + @dataclass class ParsedRecord: @@ -50,7 +56,7 @@ class ParsedRecord: text: str """Original regex-matched JSON string (invariant under reclassification).""" - parsed: dict | None = None + parsed: RecordDict | None = None """Parsed dict when validation succeeded, ``None`` when invalid.""" error: tuple[str, str] | None = None @@ -102,7 +108,7 @@ class ParsedResponse: """Index of the prompt within the batch (set by the processor call).""" @property - def valid_records(self) -> list[dict]: + def valid_records(self) -> list[RecordDict]: """Parsed dicts for records that passed validation.""" return [r.parsed for r in self.records if r.is_valid and r.parsed is not None] @@ -117,7 +123,7 @@ def errors(self) -> list[tuple[str, str]]: return [r.error for r in self.records if r.error is not None] -def is_safe_for_float_conversion(value: str | int | float | None | list | dict) -> bool: +def is_safe_for_float_conversion(value: JsonValue) -> bool: """Check if a value can be safely converted to float64 without overflow. Only ``int`` values can cause overflow; all other types are considered safe. @@ -142,7 +148,7 @@ def is_safe_for_float_conversion(value: str | int | float | None | list | dict) return True -def check_record_for_large_numbers(record: dict) -> str | None: +def check_record_for_large_numbers(record: RecordMapping) -> str | None: """Check if a record contains any numbers that would cause float64 overflow. Args: @@ -161,7 +167,7 @@ def check_record_for_large_numbers(record: dict) -> str | None: return None -def check_if_records_are_ordered(records: list[dict], order_by: str) -> bool: +def check_if_records_are_ordered(records: Sequence[RecordMapping], order_by: str) -> bool: """Check if the records are in ascending order based on the given `order_by` column. Args: @@ -176,6 +182,11 @@ def check_if_records_are_ordered(records: list[dict], order_by: str) -> bool: return order_by_values == sorted_values +def normalize_record_keys(record: RawRecordMapping) -> RecordDict: + """Return a record with string keys, matching JSON object semantics.""" + return {str(key): value for key, value in record.items()} + + def extract_records_from_jsonl_string(jsonl_string: str) -> list[str]: """Extract and return tabular records from the given JSONL string.""" return re.findall(RECORD_REGEX_PATTEN_LOOKAHEAD, jsonl_string) @@ -227,7 +238,7 @@ def _timed(text: str) -> tuple[int, float]: def extract_and_validate_records( jsonl_string: str, - schema: dict, + schema: JsonSchema, encode: Callable[[str], list[int]] | None = None, ) -> ParsedResponse: """Extract and validate records from the given JSONL string. @@ -280,7 +291,7 @@ def _parse_timestamp_to_seconds(value: object, time_format: str) -> int: """ if time_format == "elapsed_seconds": # Value is already in seconds (int for now and float for future) - return int(float(value)) # ty: ignore[invalid-argument-type] -- third-party stub mismatch + return int(float(str(value))) # Parse using strptime format dt = datetime.strptime(str(value), time_format) @@ -297,7 +308,7 @@ def _parse_timestamp_to_seconds(value: object, time_format: str) -> int: return dt.hour * 3600 + dt.minute * 60 + dt.second -def _parse_and_validate_json(matched_json: str, schema: dict) -> tuple[dict | None, tuple[str, str] | None]: +def _parse_and_validate_json(matched_json: str, schema: JsonSchema) -> tuple[RecordDict | None, tuple[str, str] | None]: """Parse JSON string and validate against schema. Args: @@ -310,6 +321,9 @@ def _parse_and_validate_json(matched_json: str, schema: dict) -> tuple[dict | No """ try: matched_dict = json.loads(matched_json) + if not is_json_object(matched_dict): + return None, ("Expected a JSON object", "Invalid JSON") + jsonschema.validate(matched_dict, schema) error_msg = check_record_for_large_numbers(matched_dict) @@ -321,11 +335,11 @@ def _parse_and_validate_json(matched_json: str, schema: dict) -> tuple[dict | No except json.JSONDecodeError as err: return None, (f"Invalid JSON: {err.msg}", "Invalid JSON") except jsonschema.exceptions.ValidationError as err: - return None, (err.message, err.validator) + return None, (err.message, str(err.validator)) def _extract_timestamp_seconds( - record: dict, time_column: str, time_format: str + record: RecordMapping, time_column: str, time_format: str ) -> tuple[int | None, tuple[str, str] | None]: """Extract and parse timestamp from a record. @@ -413,7 +427,7 @@ def _validate_time_interval( def extract_and_validate_timeseries_records( jsonl_string: str, - schema: dict, + schema: JsonSchema, time_column: str, interval_seconds: int | None, time_format: str, @@ -541,7 +555,7 @@ def normalize_dataframe(dataframe: pd.DataFrame) -> pd.DataFrame: ) -def records_to_jsonl(records: pd.DataFrame | list[dict] | dict) -> str: +def records_to_jsonl(records: pd.DataFrame | list[RawRecordMapping] | RawRecordMapping) -> str: """Convert list of records to a JSONL string. Args: @@ -550,9 +564,10 @@ def records_to_jsonl(records: pd.DataFrame | list[dict] | dict) -> str: Returns: The JSONL string. """ - if isinstance(records, pd.DataFrame): - return records.to_json(orient="records", lines=True, force_ascii=False) - elif isinstance(records, (list, dict)): - return pd.DataFrame(records).to_json(orient="records", lines=True, force_ascii=False) - else: - raise ValueError(f"Unsupported type: {type(records)}") + match records: + case pd.DataFrame() as dataframe: + return dataframe.to_json(orient="records", lines=True, force_ascii=False) + case list() | dict(): + return pd.DataFrame(records).to_json(orient="records", lines=True, force_ascii=False) + case _: + raise ValueError(f"Unsupported type: {type(records)}") diff --git a/src/nemo_safe_synthesizer/data_processing/records/base.py b/src/nemo_safe_synthesizer/data_processing/records/base.py index 5cdeaf8fb..a2a73bed4 100644 --- a/src/nemo_safe_synthesizer/data_processing/records/base.py +++ b/src/nemo_safe_synthesizer/data_processing/records/base.py @@ -91,17 +91,17 @@ def tokenize_header(field: str) -> list[str]: def get_type_as_string(value) -> str: """Return the JSON schema type name for a Python scalar value.""" - if isinstance(value, str): - return STRING - elif isinstance(value, bool): - if str(value) in ("True", "False"): + match value: + case str(): + return STRING + case bool(): return BOOL - elif isinstance(value, Number): - return NUMBER - elif value is None: - return NULL - - return NULL + case Number(): + return NUMBER + case None: + return NULL + case _: + return NULL class KVPair: 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/data_processing/records/json_record.py b/src/nemo_safe_synthesizer/data_processing/records/json_record.py index 2d7e3ae63..e928446fb 100644 --- a/src/nemo_safe_synthesizer/data_processing/records/json_record.py +++ b/src/nemo_safe_synthesizer/data_processing/records/json_record.py @@ -12,17 +12,33 @@ from __future__ import annotations +from collections.abc import Iterator from itertools import chain, starmap -from typing import Optional +from typing import Any, Optional + +from typing_extensions import override from . import base +from .json_types import JsonArray, JsonContainer, JsonObject, JsonScalar, JsonValue from .value_path import ( value_path, value_path_to_field_name, ) +__all__ = [ + "JSONRecord", + "JsonArray", + "JsonContainer", + "JsonObject", + "JsonScalar", + "JsonValue", + "convert_flat_dict_to_kv_pairs", + "flatten", + "remove_array_markers", +] + -def flatten(raw, array_marker=base.ARRAY_POS): +def flatten(raw: JsonContainer, array_marker: str = base.ARRAY_POS) -> dict[object, object]: """Recursively flatten a nested dict/list into a single-level dict. Keys are joined with ``NESTING_DELIM``; array indices are encoded as @@ -36,36 +52,39 @@ def flatten(raw, array_marker=base.ARRAY_POS): Returns: A flat dict mapping composite keys to scalar values. """ - if isinstance(raw, list): - # if the whole JSON document is an array, we wrap it in dict - raw = {None: raw} - - def unpack_level(parent_key, parent_val): - if isinstance(parent_val, dict): - for key, value in parent_val.items(): - tmp = str(parent_key) + base.NESTING_DELIM + key - yield tmp, value - elif isinstance(parent_val, list): - i = 0 - for value in parent_val: - if parent_key is None: - tmp = array_marker + str(i) - else: - tmp = str(parent_key) + base.NESTING_DELIM + array_marker + str(i) - - yield tmp, value - i += 1 - else: - yield parent_key, parent_val + match raw: + case list() as values: + # if the whole JSON document is an array, we wrap it in dict + flattened: dict[object, object] = {None: values} + case dict() as values: + flattened = {key: value for key, value in values.items()} + case _: + raise TypeError("flatten expects a JSON object or array") + + def unpack_level(parent_key: object, parent_val: object) -> Iterator[tuple[object, object]]: + match parent_val: + case dict() as values: + for key, value in values.items(): + yield str(parent_key) + base.NESTING_DELIM + str(key), value + case list() as values: + for i, value in enumerate(values): + if parent_key is None: + tmp = array_marker + str(i) + else: + tmp = str(parent_key) + base.NESTING_DELIM + array_marker + str(i) + + yield tmp, value + case scalar: + yield parent_key, scalar while True: - raw = dict(chain.from_iterable(starmap(unpack_level, raw.items()))) - if not any(isinstance(value, dict) for value in raw.values()) and not any( - isinstance(value, list) for value in raw.values() + flattened = dict(chain.from_iterable(starmap(unpack_level, flattened.items()))) + if not any(isinstance(value, dict) for value in flattened.values()) and not any( + isinstance(value, list) for value in flattened.values() ): break - return raw + return flattened def remove_array_markers(data: str) -> tuple[str, int, base.ValuePath]: @@ -76,7 +95,7 @@ def remove_array_markers(data: str) -> tuple[str, int, base.ValuePath]: """ array_count = 0 parts = data.split(base.NESTING_DELIM) - path_items = [] + path_items: list[str | int] = [] for part in parts: if part.startswith(base.ARRAY_POS): array_count += 1 @@ -88,9 +107,9 @@ def remove_array_markers(data: str) -> tuple[str, int, base.ValuePath]: return value_path_to_field_name(path), array_count, path -def convert_flat_dict_to_kv_pairs(data: dict) -> list[base.KVPair]: +def convert_flat_dict_to_kv_pairs(data: dict[Any, Any]) -> list[base.KVPair]: """Convert a flattened dict (from ``flatten``) into a list of ``KVPair`` objects.""" - out = [] + out: list[base.KVPair] = [] for k, v in data.items(): k = str(k) new_key, array_count, value_path = remove_array_markers(k) @@ -106,7 +125,7 @@ class JSONRecord(base.BaseRecord): Provides lookup by JSONPath or ``ValuePath``. """ - def __init__(self, original): + def __init__(self, original: Any): super().__init__(original) self._unpack_json() @@ -118,7 +137,8 @@ def _unpack_json(self) -> None: self.fields.add(pair.field) self.kv_pairs.append(pair) - def unpack(self): + @override + def unpack(self) -> None: self.kv_pairs = [] self.fields = set() self._unpack_json() diff --git a/src/nemo_safe_synthesizer/data_processing/records/json_types.py b/src/nemo_safe_synthesizer/data_processing/records/json_types.py new file mode 100644 index 000000000..61554d2e5 --- /dev/null +++ b/src/nemo_safe_synthesizer/data_processing/records/json_types.py @@ -0,0 +1,36 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Shared recursive JSON type aliases and runtime shape guards.""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import TypeAlias + +from typing_extensions import TypeIs + +JsonScalar: TypeAlias = str | int | float | bool | None +JsonValue: TypeAlias = JsonScalar | dict[str, "JsonValue"] | list["JsonValue"] +JsonObject: TypeAlias = dict[str, JsonValue] +JsonArray: TypeAlias = list[JsonValue] +JsonContainer: TypeAlias = JsonObject | JsonArray +JsonSchema: TypeAlias = Mapping[str, JsonValue] + + +def is_json_value(value: object) -> TypeIs[JsonValue]: + """Return whether ``value`` is representable as JSON.""" + match value: + case str() | int() | float() | bool() | None: + return True + case list() as values: + return all(is_json_value(item) for item in values) + case dict() as values: + return all(isinstance(key, str) and is_json_value(item) for key, item in values.items()) + case _: + return False + + +def is_json_object(value: object) -> TypeIs[JsonObject]: + """Return whether ``value`` is a JSON object with string keys.""" + return isinstance(value, dict) and all(isinstance(key, str) and is_json_value(item) for key, item in value.items()) diff --git a/src/nemo_safe_synthesizer/data_processing/records/value_path.py b/src/nemo_safe_synthesizer/data_processing/records/value_path.py index b63920aa6..06d3dcad2 100644 --- a/src/nemo_safe_synthesizer/data_processing/records/value_path.py +++ b/src/nemo_safe_synthesizer/data_processing/records/value_path.py @@ -73,7 +73,7 @@ def value_path_to_json_path(path: ValuePath) -> str: class InvalidPath(Exception): ... -def unflatten(data: dict[ValuePath, Any]) -> Optional[dict | list]: +def unflatten(data: dict[ValuePath, Any]) -> Optional[dict[Any, Any] | list[Any]]: """Reconstruct a nested dict/list from a flat ``{ValuePath: value}`` mapping. Args: @@ -96,18 +96,20 @@ def unflatten(data: dict[ValuePath, Any]) -> Optional[dict | list]: return result -def _ensure_array_size(result: list, item: int): +def _ensure_array_size(result: list[Any], item: int) -> None: if len(result) <= item: for i in range(len(result), item + 1): result.append(None) -def _ensure_dict_key(result: dict, item: str): +def _ensure_dict_key(result: dict[Any, Any], item: str) -> None: if item not in result: result[item] = None -def _unflatten_path(result: Optional[dict | list], path: ValuePath, value: Any) -> dict | list: +def _unflatten_path( + result: Optional[dict[Any, Any] | list[Any]], path: ValuePath, value: Any +) -> dict[Any, Any] | list[Any]: # Note: result will be a list when working with an array at this level of # the path, and thus the first element of path is an integer. Otherwise # working with an object at this level of the path, result will be a dict @@ -137,7 +139,7 @@ def _unflatten_path(result: Optional[dict | list], path: ValuePath, value: Any) return result -def _unflatten_recursive(result: Any, prev_item: ValuePathItem, items: list[ValuePathItem], value: Any): +def _unflatten_recursive(result: Any, prev_item: ValuePathItem, items: list[ValuePathItem], value: Any) -> None: # Note: result will be a list when working with an array at this level of # the path, and thus the first element of path is an integer. Otherwise # working with an object at this level of the path, result will be a dict diff --git a/src/nemo_safe_synthesizer/evaluation/components/attribute_inference_protection.py b/src/nemo_safe_synthesizer/evaluation/components/attribute_inference_protection.py index 2a9c6e4bd..c95e46b81 100644 --- a/src/nemo_safe_synthesizer/evaluation/components/attribute_inference_protection.py +++ b/src/nemo_safe_synthesizer/evaluation/components/attribute_inference_protection.py @@ -20,6 +20,7 @@ from pydantic import ConfigDict, Field from sentence_transformers import SentenceTransformer, util from sklearn.preprocessing import QuantileTransformer +from typing_extensions import override from ...config.evaluate import QUASI_IDENTIFIER_COUNT from ...config.parameters import SafeSynthesizerParameters @@ -67,6 +68,7 @@ def jinja_context(self) -> dict[str, str]: return d @staticmethod + @override def from_evaluation_datasets( evaluation_datasets: EvaluationDatasets, config: SafeSynthesizerParameters | None = None ) -> AttributeInferenceProtection: @@ -366,16 +368,16 @@ def _aia( # Get all combinations of columns to be the quasi-identifiers # This gets explosive when column count > 500 + training_columns = [str(column) for column in training_df.columns] if len(training_df.columns) < 500: - qi_combos = list(itertools.combinations(training_df.columns, quasi_identifier_count)) + qi_combos = list(itertools.combinations(training_columns, quasi_identifier_count)) else: - columns = list(training_df.columns) qi_combos = [] - for i in range(len(columns) - quasi_identifier_count): - combo = set() + for i in range(len(training_columns) - quasi_identifier_count): + combo = [] for j in range(quasi_identifier_count): - combo.add(columns[i + j]) - qi_combos.append(combo) + combo.append(training_columns[i + j]) + qi_combos.append(tuple(combo)) np.random.seed(5) np.random.shuffle(qi_combos) @@ -415,7 +417,6 @@ def _aia( # As we process the attack dataset, we'll accumulate for each column the number of # correct and incorrect predictions - training_columns = [str(column) for column in training_df.columns] correct = {predict_column: 0 for predict_column in training_columns} incorrect = {predict_column: 0 for predict_column in training_columns} @@ -578,7 +579,8 @@ def _aia( for i in range(len(entropy)): entropy_wts.append(0) else: - arr = (entropy - min(entropy)) / (max(entropy) - min(entropy)) + entropy_arr = np.asarray(entropy, dtype=float) + arr = (entropy_arr - min(entropy)) / (max(entropy) - min(entropy)) entropy_wts = arr / arr.sum() i = 0 diff --git a/src/nemo_safe_synthesizer/evaluation/components/column_distribution.py b/src/nemo_safe_synthesizer/evaluation/components/column_distribution.py index 9543f5902..9e181cda3 100644 --- a/src/nemo_safe_synthesizer/evaluation/components/column_distribution.py +++ b/src/nemo_safe_synthesizer/evaluation/components/column_distribution.py @@ -8,6 +8,7 @@ import pandas as pd from plotly.graph_objects import Figure from pydantic import BaseModel, Field +from typing_extensions import override from ...artifacts.analyzers.field_features import FieldType from ...config.parameters import SafeSynthesizerParameters @@ -121,6 +122,7 @@ def jinja_context(self) -> dict: return d @staticmethod + @override def from_evaluation_datasets( evaluation_datasets: EvaluationDatasets, config: SafeSynthesizerParameters | None = None ) -> ColumnDistribution: diff --git a/src/nemo_safe_synthesizer/evaluation/components/correlation.py b/src/nemo_safe_synthesizer/evaluation/components/correlation.py index 8763d49fe..10dbd11d1 100644 --- a/src/nemo_safe_synthesizer/evaluation/components/correlation.py +++ b/src/nemo_safe_synthesizer/evaluation/components/correlation.py @@ -8,6 +8,7 @@ import numpy as np import pandas as pd from pydantic import ConfigDict, Field +from typing_extensions import override from ...config.parameters import SafeSynthesizerParameters from ...evaluation.components.component import Component @@ -63,6 +64,7 @@ def jinja_context(self) -> dict: return d @staticmethod + @override def from_evaluation_datasets( evaluation_datasets: EvaluationDatasets, config: SafeSynthesizerParameters | None = None ) -> Correlation: diff --git a/src/nemo_safe_synthesizer/evaluation/components/data_privacy_score.py b/src/nemo_safe_synthesizer/evaluation/components/data_privacy_score.py index ac1364b64..787274d98 100644 --- a/src/nemo_safe_synthesizer/evaluation/components/data_privacy_score.py +++ b/src/nemo_safe_synthesizer/evaluation/components/data_privacy_score.py @@ -4,6 +4,7 @@ from __future__ import annotations from pydantic import Field +from typing_extensions import override from ...observability import get_logger from ..data_model.evaluation_score import ( @@ -22,6 +23,7 @@ class DataPrivacyScore(CompositeScore): name: str = Field(default="Data Privacy Score") @staticmethod + @override def from_components(components: list[Component] | Component, name: str = "Data Privacy Score") -> DataPrivacyScore: """Compute the Data Privacy Score from privacy sub-metric components.""" if isinstance(components, Component): diff --git a/src/nemo_safe_synthesizer/evaluation/components/dataset_statistics.py b/src/nemo_safe_synthesizer/evaluation/components/dataset_statistics.py index 8d5407722..64245c6cd 100644 --- a/src/nemo_safe_synthesizer/evaluation/components/dataset_statistics.py +++ b/src/nemo_safe_synthesizer/evaluation/components/dataset_statistics.py @@ -6,6 +6,7 @@ from functools import cached_property from pydantic import Field +from typing_extensions import override from ...config.parameters import SafeSynthesizerParameters from ...evaluation.components.component import Component @@ -54,6 +55,7 @@ def jinja_context(self) -> dict: return stats @staticmethod + @override def from_evaluation_datasets( evaluation_datasets: EvaluationDatasets, config: SafeSynthesizerParameters | None = None ) -> DatasetStatistics: diff --git a/src/nemo_safe_synthesizer/evaluation/components/deep_structure.py b/src/nemo_safe_synthesizer/evaluation/components/deep_structure.py index 380302c5e..03ba8de70 100644 --- a/src/nemo_safe_synthesizer/evaluation/components/deep_structure.py +++ b/src/nemo_safe_synthesizer/evaluation/components/deep_structure.py @@ -9,6 +9,7 @@ import pandas as pd from category_encoders.count import CountEncoder from pydantic import ConfigDict, Field +from typing_extensions import override from ...artifacts.analyzers.field_features import ( FieldType, @@ -56,6 +57,7 @@ def jinja_context(self) -> dict: return d @staticmethod + @override def from_evaluation_datasets( evaluation_datasets: EvaluationDatasets, config: SafeSynthesizerParameters | None = None ) -> DeepStructure: diff --git a/src/nemo_safe_synthesizer/evaluation/components/membership_inference_protection.py b/src/nemo_safe_synthesizer/evaluation/components/membership_inference_protection.py index b307615f3..2a8171659 100644 --- a/src/nemo_safe_synthesizer/evaluation/components/membership_inference_protection.py +++ b/src/nemo_safe_synthesizer/evaluation/components/membership_inference_protection.py @@ -15,6 +15,7 @@ from sentence_transformers import SentenceTransformer, util from sklearn.metrics import accuracy_score, precision_score from sklearn.preprocessing import QuantileTransformer +from typing_extensions import override from ...config.evaluate import DEFAULT_RECORD_COUNT from ...config.parameters import SafeSynthesizerParameters @@ -68,6 +69,7 @@ def jinja_context(self) -> dict: return d @staticmethod + @override def from_evaluation_datasets( evaluation_datasets: EvaluationDatasets, config: SafeSynthesizerParameters | None = None ) -> MembershipInferenceProtection: @@ -248,8 +250,8 @@ def _compute_mia( ) -> tuple[ float, list[str], - dict[str, list[int]], - dict[str, list[int]], + dict[float, int], + dict[float, int], ]: """Core membership inference attack implementation for a single run. @@ -286,8 +288,8 @@ def _compute_mia( pd.concat([training_df_attack, test_df_norm]).reset_index(drop=True).sample(frac=1, random_state=run) ) - attack_synth_dist_text = [[0] for i in range(len(attack_df))] - attack_synth_indices_text = [[0] for i in range(len(attack_df))] + attack_synth_dist_text: list[list[float]] = [[0.0] for i in range(len(attack_df))] + attack_synth_indices_text: list[list[int]] = [[0] for i in range(len(attack_df))] # Get the NN dist for text for the entire attack dataset @@ -362,8 +364,8 @@ def _compute_mia( score = 0 attack_summary = [] - tp_cnts = {} - fp_cnts = {} + tp_cnts: dict[float, int] = {} + fp_cnts: dict[float, int] = {} # Using the above text and tabular distances we now compute an overall distance score for # every record in the attack dataset. We then conduct 36 individual mia attacks on this one big @@ -519,8 +521,8 @@ def mia( scores = [] attack_sum_values = [] - tps_values = {} - fps_values = {} + tps_values: dict[float, int] = {} + fps_values: dict[float, int] = {} for i in [0.1, 0.2, 0.3, 0.4]: tps_values[i] = 0 fps_values[i] = 0 diff --git a/src/nemo_safe_synthesizer/evaluation/components/multi_modal_figures.py b/src/nemo_safe_synthesizer/evaluation/components/multi_modal_figures.py index 4b3ed1fc9..d04f2ce95 100644 --- a/src/nemo_safe_synthesizer/evaluation/components/multi_modal_figures.py +++ b/src/nemo_safe_synthesizer/evaluation/components/multi_modal_figures.py @@ -676,7 +676,7 @@ def generate_text_structure_similarity_figures( "average_words_per_sentence", "average_characters_per_word", ] - figures = [] + figures: list[go.Figure] = [] for key in statistics_keys: if training_statistics.per_record_statistics.empty or synthetic_statistics.per_record_statistics.empty: break @@ -684,7 +684,8 @@ def generate_text_structure_similarity_figures( training_statistics.per_record_statistics[key], synthetic_statistics.per_record_statistics[key], ) - figures.append(figure) + if figure is not None: + figures.append(figure) if not figures: return None diff --git a/src/nemo_safe_synthesizer/evaluation/components/pii_replay.py b/src/nemo_safe_synthesizer/evaluation/components/pii_replay.py index cf857dbe2..213ebf9b4 100644 --- a/src/nemo_safe_synthesizer/evaluation/components/pii_replay.py +++ b/src/nemo_safe_synthesizer/evaluation/components/pii_replay.py @@ -7,6 +7,7 @@ from functools import cached_property from pydantic import BaseModel, Field +from typing_extensions import override from ...config.parameters import SafeSynthesizerParameters from ...evaluation.components.component import Component @@ -72,6 +73,7 @@ def jinja_context(self) -> dict: return d @staticmethod + @override def from_evaluation_datasets( evaluation_datasets: EvaluationDatasets, config: SafeSynthesizerParameters | None = None ) -> PIIReplay: diff --git a/src/nemo_safe_synthesizer/evaluation/components/sqs_score.py b/src/nemo_safe_synthesizer/evaluation/components/sqs_score.py index c523f74db..714bf1bdf 100644 --- a/src/nemo_safe_synthesizer/evaluation/components/sqs_score.py +++ b/src/nemo_safe_synthesizer/evaluation/components/sqs_score.py @@ -4,6 +4,7 @@ from __future__ import annotations from pydantic import Field +from typing_extensions import override from ...artifacts.analyzers.field_features import ( FieldType, @@ -29,6 +30,7 @@ class SQSScore(CompositeScore): name: str = Field(default="Synthetic Quality Score") @staticmethod + @override def from_components(components: list[Component] | Component, name: str = "Synthetic Quality Score") -> SQSScore: """Compute the SQS from a list of quality sub-metric components.""" if isinstance(components, Component): diff --git a/src/nemo_safe_synthesizer/evaluation/components/text_semantic_similarity.py b/src/nemo_safe_synthesizer/evaluation/components/text_semantic_similarity.py index e1f6e04cc..686673353 100644 --- a/src/nemo_safe_synthesizer/evaluation/components/text_semantic_similarity.py +++ b/src/nemo_safe_synthesizer/evaluation/components/text_semantic_similarity.py @@ -20,6 +20,7 @@ stop_after_attempt, wait_exponential, ) +from typing_extensions import override from ...artifacts.analyzers.field_features import FieldType from ...config.evaluate import DEFAULT_RECORD_COUNT @@ -95,6 +96,7 @@ def jinja_context(self) -> dict: return ctx @staticmethod + @override def from_evaluation_datasets( evaluation_datasets: EvaluationDatasets, config: SafeSynthesizerParameters | None = None ) -> TextSemanticSimilarity: diff --git a/src/nemo_safe_synthesizer/evaluation/components/text_structure_similarity.py b/src/nemo_safe_synthesizer/evaluation/components/text_structure_similarity.py index 9654d4e07..4b82764fe 100644 --- a/src/nemo_safe_synthesizer/evaluation/components/text_structure_similarity.py +++ b/src/nemo_safe_synthesizer/evaluation/components/text_structure_similarity.py @@ -9,6 +9,7 @@ import numpy as np import pandas as pd from pydantic import BaseModel, ConfigDict, Field +from typing_extensions import override from ...artifacts.analyzers.field_features import FieldType from ...config.parameters import SafeSynthesizerParameters @@ -93,6 +94,7 @@ def jinja_context(self) -> dict: return d @staticmethod + @override def from_evaluation_datasets( evaluation_datasets: EvaluationDatasets, config: SafeSynthesizerParameters | None = None ) -> TextStructureSimilarity: diff --git a/src/nemo_safe_synthesizer/evaluation/nearest_neighbors.py b/src/nemo_safe_synthesizer/evaluation/nearest_neighbors.py index 0b7dd3fea..c70aec3bf 100644 --- a/src/nemo_safe_synthesizer/evaluation/nearest_neighbors.py +++ b/src/nemo_safe_synthesizer/evaluation/nearest_neighbors.py @@ -31,6 +31,7 @@ os.environ.setdefault("OMP_NUM_THREADS", "8") import numpy as np +import numpy.typing as npt import torch from sklearn.neighbors import NearestNeighbors @@ -61,7 +62,7 @@ def __init__(self, n_neighbors: int = 5): self.n_neighbors = n_neighbors self._torch_device = self._detect_device() self.use_gpu = self._torch_device is not None - self._index = None + self._index: NearestNeighbors | None = None self._data_t: torch.Tensor | None = None @classmethod @@ -87,7 +88,7 @@ def _detect_device(cls) -> torch.device | None: cls._device_checked = True return cls._torch_device - def fit(self, data: np.ndarray) -> NearestNeighborSearch: + def fit(self, data: npt.NDArray[np.float32]) -> NearestNeighborSearch: """Build the search index from data. Args: @@ -110,7 +111,11 @@ def fit(self, data: np.ndarray) -> NearestNeighborSearch: return self - def kneighbors(self, queries: np.ndarray, n_neighbors: int | None = None) -> tuple[np.ndarray, np.ndarray]: + def kneighbors( + self, + queries: npt.NDArray[np.float32], + n_neighbors: int | None = None, + ) -> tuple[npt.NDArray[np.float32], npt.NDArray[np.int64]]: """Find k nearest neighbors for query points. Args: diff --git a/src/nemo_safe_synthesizer/evaluation/reports/multimodal/multimodal_report.py b/src/nemo_safe_synthesizer/evaluation/reports/multimodal/multimodal_report.py index b0d7358eb..90410d1b7 100644 --- a/src/nemo_safe_synthesizer/evaluation/reports/multimodal/multimodal_report.py +++ b/src/nemo_safe_synthesizer/evaluation/reports/multimodal/multimodal_report.py @@ -21,6 +21,7 @@ ColumnDistribution, ColumnDistributionPlotRow, ) +from ....evaluation.components.component import Component from ....evaluation.components.correlation import ( Correlation, ) @@ -150,7 +151,7 @@ def from_dataframes( mandatory_columns=MultimodalReport._get_config_value("mandatory_columns", [], config), ) - components = [] + components: list[Component] = [] attribute_inference_protection = AttributeInferenceProtection( score=EvaluationScore(grade=PrivacyGrade.UNAVAILABLE) diff --git a/src/nemo_safe_synthesizer/generation/backend.py b/src/nemo_safe_synthesizer/generation/backend.py index 97c64f906..ce44a8334 100644 --- a/src/nemo_safe_synthesizer/generation/backend.py +++ b/src/nemo_safe_synthesizer/generation/backend.py @@ -8,6 +8,8 @@ import abc from collections.abc import Callable +from typing_extensions import override + from .. import utils from ..cli.artifact_structure import Workdir from ..config import SafeSynthesizerParameters @@ -56,6 +58,7 @@ class GeneratorBackend(metaclass=abc.ABCMeta): """Working directory containing model artifacts.""" @classmethod + @override def __subclasshook__(cls, subclass): return ( hasattr(subclass, "prepare_args") diff --git a/src/nemo_safe_synthesizer/generation/processors.py b/src/nemo_safe_synthesizer/generation/processors.py index 1811287aa..901365efc 100644 --- a/src/nemo_safe_synthesizer/generation/processors.py +++ b/src/nemo_safe_synthesizer/generation/processors.py @@ -8,6 +8,8 @@ from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Any +from typing_extensions import override + from ..config import SafeSynthesizerParameters from ..config.generate import ValidationParameters from ..data_processing.record_utils import ( @@ -126,6 +128,7 @@ def _encode(self): class TabularDataProcessor(Processor): """Processor for standard (non-grouped, non-time-series) tabular data.""" + @override def _process_text_generation(self, text: str) -> ParsedResponse: """Extract and validate records from a flat JSONL string. @@ -185,6 +188,7 @@ def __init__( self.interval_seconds = interval_seconds self.time_format: str = time_format + @override def _process_text_generation(self, text: str) -> ParsedResponse: """Extract, validate, and check temporal ordering of time-series records. @@ -240,6 +244,7 @@ def __init__( self.bos_token = bos_token self.eos_token = eos_token + @override def _process_text_generation(self, text: str) -> ParsedResponse: """Process the output from the fine-tuned model. diff --git a/src/nemo_safe_synthesizer/generation/regex_manager.py b/src/nemo_safe_synthesizer/generation/regex_manager.py index 990f34437..0f83207f8 100644 --- a/src/nemo_safe_synthesizer/generation/regex_manager.py +++ b/src/nemo_safe_synthesizer/generation/regex_manager.py @@ -42,7 +42,7 @@ # Helper method not exported by outlines_core and outlines doesn't have it anymore # past outlines==0.11.8 -def _get_num_items_pattern(min_items: int | None, max_items: int | None, **kwargs) -> str | None: +def _get_num_items_pattern(min_items: int | None, max_items: int | None, **kwargs: object) -> str | None: """Return a regex quantifier ``{min,max}`` for array/object items.""" min_items = int(min_items or 0) if max_items is None: @@ -60,7 +60,7 @@ def _build_object_key_prefix(name: str, whitespace_pattern: str) -> str: return f'{whitespace_pattern}"{re.escape(key_inner)}"{whitespace_pattern}:{whitespace_pattern}' -def _properties_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs) -> str: +def _properties_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs: object) -> str: """Build a regex matching a JSON object with known property names.""" regex = "" regex += r"\{" @@ -115,7 +115,7 @@ def _properties_regex(instance: dict[str, Any], whitespace_pattern: str, **kwarg return regex -def _enum_regex(instance: dict[str, Any], **kwargs) -> str: +def _enum_regex(instance: dict[str, Any], **kwargs: object) -> str: """Build a regex matching any value in the schema's ``enum`` array.""" choices = [] for choice in instance["enum"]: @@ -141,7 +141,7 @@ def _enum_regex(instance: dict[str, Any], **kwargs) -> str: return f"({'|'.join(choices)})" -def _string_type_regex(instance: dict[str, Any], **kwargs) -> str: +def _string_type_regex(instance: dict[str, Any], **kwargs: object) -> str: if "maxLength" in instance or "minLength" in instance: max_items = instance.get("maxLength", "") min_items = instance.get("minLength", "") @@ -177,7 +177,7 @@ def _string_type_regex(instance: dict[str, Any], **kwargs) -> str: return JSON_STRING -def _type_array_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs) -> str: +def _type_array_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs: object) -> str: num_repeats = _get_num_items_pattern(instance.get("minItems"), instance.get("maxItems")) if num_repeats is None: return rf"\[{whitespace_pattern}\]" @@ -202,7 +202,7 @@ def _type_array_regex(instance: dict[str, Any], whitespace_pattern: str, **kwarg return rf"\[{whitespace_pattern}({'|'.join(regexes)})(,{whitespace_pattern}({'|'.join(regexes)})){num_repeats}){allow_empty}{whitespace_pattern}\]" -def _type_object_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs) -> str: +def _type_object_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs: object) -> str: # pattern for json object with values defined by instance["additionalProperties"] # enforces value type constraints recursively, "minProperties", and "maxProperties" # doesn't enforce "required", "dependencies", "propertyNames" "any/all/on Of" @@ -223,7 +223,7 @@ def _type_object_regex(instance: dict[str, Any], whitespace_pattern: str, **kwar return r"\{" + whitespace_pattern + multiple_key_value_pattern + whitespace_pattern + r"\}" -def _type_int_regex(instance: dict[str, Any], **kwargs) -> str: +def _type_int_regex(instance: dict[str, Any], **kwargs: object) -> str: if "minimum" in instance and "maximum" in instance: min_int = int(instance["minimum"]) max_int = int(instance["maximum"]) @@ -245,7 +245,7 @@ def _type_int_regex(instance: dict[str, Any], **kwargs) -> str: return regex -def _type_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs) -> str: +def _type_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs: object) -> str: """Dispatch to the appropriate regex builder based on the ``type`` keyword. The ``type`` keyword may be a string naming a single basic type or an @@ -253,14 +253,6 @@ def _type_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs) -> matches any of the listed types. """ instance_type = instance["type"] - dispatch = { - "string": _string_type_regex, - "integer": _type_int_regex, - "array": _type_array_regex, - "object": _type_object_regex, - } - - kwargs = {"instance": instance, "whitespace_pattern": whitespace_pattern} match instance_type: # match first on the list type as the other will attempt to hash it to index into # the dict @@ -276,9 +268,14 @@ def _type_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs) -> # https://cswr.github.io/JsonSchema/spec/multiple_types/ regexes = [_build_regex({"type": t}, whitespace_pattern) for t in instance_type if t != "object"] return rf"({'|'.join(regexes)})" - case x if x in dispatch: - return dispatch[x](**kwargs) # ty: ignore[invalid-argument-type] -- kwargs values are correct per-key; ty can't narrow the union through dict unpacking - + case "string": + return _string_type_regex(instance) + case "integer": + return _type_int_regex(instance) + case "array": + return _type_array_regex(instance, whitespace_pattern) + case "object": + return _type_object_regex(instance, whitespace_pattern) case "number": return NUMBER case "boolean": @@ -289,7 +286,7 @@ def _type_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs) -> raise NotImplementedError(f"Unsupported type={instance_type}") -def _build_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs) -> str: +def _build_regex(instance: dict[str, Any], whitespace_pattern: str, **kwargs: object) -> str: """Convert a JSON schema fragment into a regex string. Supports the subset of JSON Schema needed for TabFT schemas -- diff --git a/src/nemo_safe_synthesizer/generation/results.py b/src/nemo_safe_synthesizer/generation/results.py index 5e6a7a40d..6809646e7 100644 --- a/src/nemo_safe_synthesizer/generation/results.py +++ b/src/nemo_safe_synthesizer/generation/results.py @@ -14,6 +14,7 @@ from ..data_processing.actions.utils import ( MetadataColumns, ) +from ..data_processing.record_utils import normalize_record_keys from ..data_processing.stats import RunningStatistics from ..defaults import ( EPS, @@ -294,7 +295,7 @@ def _apply_data_actions_fn(self, batch: Batch) -> None: rid = record.parsed[record_id_key] if new_record := id_to_valid_records.get(rid): del new_record[record_id_key] - record.parsed = new_record + record.parsed = normalize_record_keys(new_record) elif new_record := id_to_rejected_records.get(rid): del new_record[record_id_key] record.invalidate(rejected_record_to_error(new_record)) diff --git a/src/nemo_safe_synthesizer/generation/timeseries_backend.py b/src/nemo_safe_synthesizer/generation/timeseries_backend.py index d7f4de6b3..e6f43f97f 100644 --- a/src/nemo_safe_synthesizer/generation/timeseries_backend.py +++ b/src/nemo_safe_synthesizer/generation/timeseries_backend.py @@ -13,6 +13,7 @@ from pathlib import Path import pandas as pd +from typing_extensions import override from vllm.sampling_params import SamplingParams from .. import utils @@ -237,6 +238,7 @@ def __init__(self, config: SafeSynthesizerParameters, model_metadata: ModelMetad self._group_prefills: dict[str, str] = initial_prefill_value self._groups: list[str] = list(self._group_prefills.keys()) + @override def _get_prompt_token_count(self) -> int: """Return the longest active prompt length for ``SamplingParams``. @@ -890,6 +892,7 @@ def _retain_single_valid_response(self, batch: Batch) -> list[dict]: return final_records + @override def generate( self, data_actions_fn: utils.DataActionsFn | None = None, diff --git a/src/nemo_safe_synthesizer/generation/vllm_backend.py b/src/nemo_safe_synthesizer/generation/vllm_backend.py index aca79ec1a..04b557a4f 100644 --- a/src/nemo_safe_synthesizer/generation/vllm_backend.py +++ b/src/nemo_safe_synthesizer/generation/vllm_backend.py @@ -12,10 +12,11 @@ import time from functools import partial from pathlib import Path -from typing import Any, cast +from typing import Any, TypeGuard, cast import torch from transformers import PreTrainedTokenizerBase +from typing_extensions import override from vllm import LLM as vLLM from vllm import RequestOutput from vllm.config import StructuredOutputsConfig @@ -48,7 +49,7 @@ from ..llm.metadata import ModelMetadata from ..llm.utils import ModelRef, cleanup_memory, get_max_vram from ..observability import get_logger, heartbeat -from ..utils import all_equal_type, load_json +from ..utils import load_json logger = get_logger(__name__) @@ -128,6 +129,20 @@ def _secure_outlines_cache_dir() -> None: _secure_outlines_cache_dir() +def _is_nested_int_list(value: object) -> TypeGuard[list[list[int]]]: + """Return whether ``value`` is a non-empty list of int lists.""" + return ( + isinstance(value, list) + and bool(value) + and all(isinstance(row, list) and all(isinstance(item, int) for item in row) for row in value) + ) + + +def _is_flat_int_list(value: object) -> TypeGuard[list[int]]: + """Return whether ``value`` is a list of ints.""" + return isinstance(value, list) and all(isinstance(item, int) for item in value) + + def _is_redis_available() -> bool: """Return True if the ``redis`` package is importable.""" try: @@ -243,6 +258,7 @@ def __init__( self.lora_req = LoRARequest("lora", 1, str(adapter_path)) if adapter_path else None self._torn_down = False + @override def teardown(self) -> None: """Release GPU memory and distributed resources. Idempotent -- safe to call multiple times.""" if self._torn_down: @@ -270,6 +286,7 @@ def __del__(self) -> None: except Exception: logger.debug("VllmBackend teardown failed during garbage collection", exc_info=True) + @override def initialize(self, **kwargs) -> None: """Initialize and load the model into memory. @@ -465,6 +482,7 @@ def _transform_kwargs_to_sampling_params( return sampling_params + @override def prepare_params(self, **kwargs) -> None: """Parse parameters and configure the generation method. @@ -538,11 +556,10 @@ def _generate( case torch.Tensor(): logger.debug("vllm generate: prompt_token_ids (torch.Tensor)") result = self._gen_method(prompt_token_ids=input_ids.tolist()) - case [[*_inner], *_] if all_equal_type(input_ids, int): # ty: ignore[invalid-argument-type] - assert isinstance(input_ids, list) - logger.debug(f"vllm generate: prompt_token_ids ({len(input_ids)} prompts)") - result = self._gen_method(prompt_token_ids=input_ids) - case [*ids] if all_equal_type(ids, int, flatten_iter=False): + case list() as ids if _is_nested_int_list(ids): + logger.debug(f"vllm generate: prompt_token_ids ({len(ids)} prompts)") + result = self._gen_method(prompt_token_ids=ids) + case list() as ids if _is_flat_int_list(ids): logger.debug("vllm generate: prompt_token_ids (single flat list)") result = self._gen_method(prompt_token_ids=[ids]) case None: @@ -637,6 +654,7 @@ def _log_batch_timing_and_progress( }, ) + @override def generate( self, data_actions_fn: utils.DataActionsFn | None = None, diff --git a/src/nemo_safe_synthesizer/llm/utils.py b/src/nemo_safe_synthesizer/llm/utils.py index 5ac6ed0ba..a628a0dd1 100644 --- a/src/nemo_safe_synthesizer/llm/utils.py +++ b/src/nemo_safe_synthesizer/llm/utils.py @@ -15,7 +15,9 @@ from dataclasses import dataclass from fnmatch import fnmatchcase from pathlib import Path -from typing import TYPE_CHECKING, Any, ClassVar, Literal, Self, cast +from typing import TYPE_CHECKING, Any, ClassVar, Literal, Self, TypeAlias, cast + +from typing_extensions import TypeIs from ..observability import get_logger @@ -28,6 +30,17 @@ logger = get_logger(__name__) +AutoMapValue: TypeAlias = str | list[object] +WeightMap: TypeAlias = dict[str, str] + + +def _is_weight_map(value: object) -> TypeIs[WeightMap]: + return ( + isinstance(value, dict) + and bool(value) + and all(isinstance(key, str) and isinstance(shard_name, str) for key, shard_name in value.items()) + ) + @dataclass(frozen=True, slots=True) class ModelRef: @@ -301,25 +314,27 @@ def _remote_code_components(cls, model_dir: Path) -> list[tuple[str, Path | None except (OSError, json.JSONDecodeError): return [] - auto_map = data.get("auto_map") - if not isinstance(auto_map, dict): - return [] - - components: list[tuple[str, Path | None]] = [] - for value in auto_map.values(): - for class_ref in cls._auto_map_class_refs(value): - component = cls._remote_code_component(class_ref) - if component is not None: - components.append(component) - return components + match data.get("auto_map"): + case dict() as auto_map: + components: list[tuple[str, Path | None]] = [] + for value in auto_map.values(): + match value: + case (str() | list()) as auto_map_value: + for class_ref in cls._auto_map_class_refs(auto_map_value): + component = cls._remote_code_component(class_ref) + if component is not None: + components.append(component) + return components + case _: + return [] @staticmethod - def _auto_map_class_refs(value: object) -> list[str]: - if isinstance(value, str): - return [value] - if isinstance(value, list): - return [item for item in value if isinstance(item, str)] - return [] + def _auto_map_class_refs(value: AutoMapValue) -> list[str]: + match value: + case str() as class_ref: + return [class_ref] + case list() as class_refs: + return [item for item in class_refs if isinstance(item, str)] @staticmethod def _remote_code_component(class_ref: str) -> tuple[str, Path | None] | None: @@ -352,14 +367,11 @@ def _index_references_existing_shards(model_dir: Path, index_path: Path) -> bool except (OSError, json.JSONDecodeError): return False - weight_map = data.get("weight_map") - if not isinstance(weight_map, dict) or not weight_map: - return False - - shard_names = {name for name in weight_map.values() if isinstance(name, str)} - if not shard_names: - return False - return all((model_dir / name).is_file() for name in shard_names) + match data.get("weight_map"): + case weight_map if _is_weight_map(weight_map): + return all((model_dir / shard_name).is_file() for shard_name in set(weight_map.values())) + case _: + return False def partial_cached_snapshot(self) -> Path | None: """Return the local HF snapshot for this repo/revision, even if it is partial.""" diff --git a/src/nemo_safe_synthesizer/observability.py b/src/nemo_safe_synthesizer/observability.py index c4d5b1c92..0d6a6c415 100644 --- a/src/nemo_safe_synthesizer/observability.py +++ b/src/nemo_safe_synthesizer/observability.py @@ -48,6 +48,7 @@ import colorama from structlog.types import BindableLogger, Processor +from typing_extensions import override if TYPE_CHECKING: from .cli.artifact_structure import Workdir @@ -301,6 +302,7 @@ def _render_table_data_for_console(logger: logging.Logger, method_name: str, eve class DiscardSensitiveMessages(logging.Filter): """Discards messages marked as sensitive via the `sensitive` flag.""" + @override def filter(self, record: logging.LogRecord) -> bool: return not getattr(record, "sensitive", False) @@ -312,6 +314,7 @@ def __init__(self, include_categories: set[LogCategory] | None = None): super().__init__() self.include_categories = include_categories + @override def filter(self, record: logging.LogRecord) -> bool: if self.include_categories is None: return True @@ -603,6 +606,7 @@ def __init__(self, logger: logging.Logger, category: LogCategory): super().__init__(logger, {}) self.category = category + @override def process(self, msg: str, kwargs: MutableMapping[str, Any]) -> tuple[str, MutableMapping[str, Any]]: # Set category via contextvar to avoid it appearing in ExtraAdder output _current_log_category.set(self.category.value) @@ -656,9 +660,11 @@ def __init__(self, base_logger: logging.Logger): ) @property + @override def name(self) -> str: return self._logger.name + @override def isEnabledFor(self, level: int) -> bool: return self._logger.isEnabledFor(level) 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 2bba807ee..e8887bd29 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 @@ -572,9 +598,10 @@ class EntityExtractorGliner(EntityExtractor): _entity_cache: dict[tuple, list] @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,24 +613,24 @@ 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 _detect_entities_chunked( @@ -627,11 +654,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 [] + predict_entities = getattr(model, "predict_entities") entities_key = tuple(sorted(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 @@ -646,7 +675,7 @@ def _detect_entities_chunked( temp_entities = self._entity_cache.get((hash(chunk), entities_key)) if temp_entities is None: n_cache_miss += 1 - temp_entities = self._model.predict_entities( + temp_entities = predict_entities( chunk, entity_labels, threshold=self._ner_threshold, @@ -686,17 +715,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 @@ -713,11 +745,13 @@ 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 - entities_key = tuple(sorted(entity_labels)) + batch_predict_entities = getattr(model, "batch_predict_entities") + entities_key = tuple(sorted(entities)) for text in texts: text = str(text) start = 0 @@ -728,9 +762,9 @@ def batch_update_cache(self, texts: list[str], entity_labels: Optional[set[str]] for batch_n, batch in enumerate( [chunks[t : t + self._batch_size] for t in range(0, len(chunks), self._batch_size)] ): - entities_lists = self._model.batch_predict_entities( + entities_lists = batch_predict_entities( batch, - entity_labels, + entities, threshold=self._ner_threshold, flat_ner=False, ) @@ -754,6 +788,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 = [] @@ -761,6 +796,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 = [] @@ -769,13 +805,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..a0cbe129f 100644 --- a/src/nemo_safe_synthesizer/pii_replacer/ner/model.py +++ b/src/nemo_safe_synthesizer/pii_replacer/ner/model.py @@ -6,18 +6,194 @@ from __future__ import annotations import re +from collections.abc import Mapping +from numbers import Real +from typing import Any, Literal, TypeAlias, TypedDict, 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) -> TypeIs[list[str]]: + return isinstance(rows, list) and bool(rows) and all(isinstance(row, str) for row in rows) + + +def _is_record_rows(rows: object) -> TypeIs[list[NERInputRecord]]: + return isinstance(rows, list) and bool(rows) and all(isinstance(row, dict) 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 +211,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..611728da5 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,14 @@ class _ProcPayload: """ +def _is_prediction(value: object) -> TypeIs[Prediction]: + return isinstance(value, NERPrediction) or isinstance(value, dict) + + +def _is_prediction_list(value: object) -> TypeIs[PredictionList]: + return isinstance(value, list) and all(_is_prediction(item) for item in value) + + 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 +73,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 +166,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 +212,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 +229,29 @@ 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: + for prediction in payload.out_data: + if _is_prediction(prediction): + single_preds.append(prediction) 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/src/nemo_safe_synthesizer/privacy/dp_transformers/dp_utils.py b/src/nemo_safe_synthesizer/privacy/dp_transformers/dp_utils.py index 6e7ab505c..9c2335bf7 100644 --- a/src/nemo_safe_synthesizer/privacy/dp_transformers/dp_utils.py +++ b/src/nemo_safe_synthesizer/privacy/dp_transformers/dp_utils.py @@ -27,7 +27,6 @@ import torch from accelerate.optimizer import AcceleratedOptimizer from datasets import Dataset -from opacus.accountants import RDPAccountant from peft import PeftModel from torch import nn from torch.utils.data import DataLoader @@ -45,6 +44,7 @@ utils, ) from transformers.trainer import TRAINING_ARGS_NAME +from typing_extensions import override from ...observability import get_logger from . import linear # imported for side effects # noqa @@ -89,6 +89,7 @@ def __init__( self.noise_multiplier = noise_multiplier self.sampling_probability = sampling_probability + @override def on_substep_end( self, args: training_args.TrainingArguments, @@ -126,6 +127,7 @@ def on_substep_end( self._on_substep_end_was_called = True + @override def on_step_end( self, args: training_args.TrainingArguments, @@ -166,12 +168,13 @@ def on_step_end( if not self.accountant.use_prv: # Use RDPAccountant, which uses `.step()` to increment number of # steps, required for accurate epsilon calculation. - acct = cast(RDPAccountant, self.accountant.accountant) + acct = self.accountant.rdp_accountant() acct.step( noise_multiplier=self.noise_multiplier, sample_rate=self.sampling_probability, ) + @override def on_save( self, args: training_args.TrainingArguments, @@ -194,6 +197,7 @@ def on_save( """ return self._check_max_epsilon_exceeded(state, control) + @override def on_evaluate( self, args: training_args.TrainingArguments, @@ -252,6 +256,7 @@ class DataCollatorForPrivateCausalLanguageModeling(DataCollatorForLanguageModeli def __init__(self, tokenizer: PreTrainedTokenizer): super().__init__(tokenizer=tokenizer, mlm=False) + @override def __call__( self, features, @@ -288,6 +293,7 @@ class DataCollatorForPrivateTokenClassification(DataCollatorForTokenClassificati def __init__(self, tokenizer: PreTrainedTokenizer): super().__init__(tokenizer=tokenizer) + @override def __call__( self, features, @@ -503,6 +509,7 @@ def num_steps(self) -> int: ) return _num_steps + @override def create_optimizer(self, model: nn.Module | None = None) -> torch.optim.Optimizer: """Create the base optimizer then wrap it with Opacus DPOptimizer.""" _ = model # Signature matches transformers v5; base method uses self.model. @@ -516,10 +523,12 @@ class DPOptimizer(opacus.optimizers.DPOptimizer): optimizer so learning rate scheduling and other param_group updates work. """ + @override @property def param_groups(self) -> list: return self.original_optimizer.param_groups + @override @param_groups.setter def param_groups(self, param_groups: list) -> None: self.original_optimizer.param_groups = param_groups @@ -542,6 +551,7 @@ def param_groups(self, param_groups: list) -> None: return self.optimizer + @override def training_step( self, model: nn.Module, @@ -583,6 +593,7 @@ def training_step( return loss.detach() / self.args.gradient_accumulation_steps + @override def _get_train_sampler(self, train_dataset: Dataset | None = None) -> torch.utils.data.Sampler | None: # ty: ignore[invalid-method-override] -- HF Trainer stub imprecision """Return the entity-level (or record-level) sampler for training.""" ds = train_dataset if train_dataset is not None else self.train_dataset @@ -607,6 +618,7 @@ def _get_train_sampler(self, train_dataset: Dataset | None = None) -> torch.util batch_size=self.args.per_device_train_batch_size, ) + @override def get_train_dataloader(self) -> DataLoader: """Returns a torch DataLoader that uses an entity-level sampler.""" train_dataset = self.train_dataset @@ -621,6 +633,7 @@ def get_train_dataloader(self) -> DataLoader: pin_memory=self.args.dataloader_pin_memory, ) + @override def _save(self, output_dir: str | None = None, state_dict: dict[str, Any] | None = None) -> None: """Save the PEFT adapter (unwrap GradSampleModule) and tokenizer. diff --git a/src/nemo_safe_synthesizer/privacy/dp_transformers/linear.py b/src/nemo_safe_synthesizer/privacy/dp_transformers/linear.py index 3c5473412..11c28f234 100644 --- a/src/nemo_safe_synthesizer/privacy/dp_transformers/linear.py +++ b/src/nemo_safe_synthesizer/privacy/dp_transformers/linear.py @@ -22,6 +22,19 @@ from opt_einsum import contract +def _contract_tensor(expression: str, *operands: torch.Tensor) -> torch.Tensor: + result = contract(expression, *operands) + if not isinstance(result, torch.Tensor): + raise TypeError("expected opt_einsum.contract to return a torch.Tensor for torch operands") + return result + + +def _linear_parameter(parameter: torch.Tensor) -> nn.Parameter: + if not isinstance(parameter, nn.Parameter): + raise TypeError("expected nn.Linear parameter") + return parameter + + @register_grad_sampler(nn.Linear) def compute_linear_grad_sample( layer: nn.Linear, activations: list[torch.Tensor], backprops: torch.Tensor @@ -41,10 +54,10 @@ def compute_linear_grad_sample( per-sample gradient tensor of shape ``(batch, ...)``. """ activation = activations[0] - ret = {} - if layer.weight.requires_grad: - gs = contract("n...i,n...j->nij", backprops.float(), activation.float()) - ret[layer.weight] = gs + ret: dict[nn.Parameter, torch.Tensor] = {} + weight = _linear_parameter(layer.weight) + if weight.requires_grad: + ret[weight] = _contract_tensor("n...i,n...j->nij", backprops.float(), activation.float()) if layer.bias is not None and layer.bias.requires_grad: - ret[layer.bias] = contract("n...k->nk", backprops.float()) + ret[layer.bias] = _contract_tensor("n...k->nk", backprops.float()) return ret diff --git a/src/nemo_safe_synthesizer/privacy/dp_transformers/privacy_args.py b/src/nemo_safe_synthesizer/privacy/dp_transformers/privacy_args.py index de8b8b1e4..fd7eba109 100644 --- a/src/nemo_safe_synthesizer/privacy/dp_transformers/privacy_args.py +++ b/src/nemo_safe_synthesizer/privacy/dp_transformers/privacy_args.py @@ -19,7 +19,7 @@ import threading import warnings from dataclasses import dataclass, field -from typing import Literal, cast +from typing import Literal import numpy as np from opacus.accountants import RDPAccountant @@ -145,11 +145,21 @@ def compute_epsilon(self, steps: int) -> float: # HF Trainer runs an extra optimizer step for an incomplete # gradient-accumulation batch at the end of an epoch. steps = min(steps, self.max_compositions) - acct = cast(PRVAccountant, self.accountant) - return acct.compute_epsilon(steps)[2] + return self.prv_accountant().compute_epsilon(steps)[2] else: - acct = cast(RDPAccountant, self.accountant) - return acct.get_epsilon(self.delta) + return self.rdp_accountant().get_epsilon(self.delta) + + def prv_accountant(self) -> PRVAccountant: + """Return the PRV accountant when this instance is in PRV mode.""" + if not isinstance(self.accountant, PRVAccountant): + raise TypeError("SafeSynthesizerAccountant is not using a PRV accountant") + return self.accountant + + def rdp_accountant(self) -> RDPAccountant: + """Return the RDP accountant when this instance is in RDP mode.""" + if not isinstance(self.accountant, RDPAccountant): + raise TypeError("SafeSynthesizerAccountant is not using an RDP accountant") + return self.accountant @dataclass diff --git a/src/nemo_safe_synthesizer/privacy/dp_transformers/sampler.py b/src/nemo_safe_synthesizer/privacy/dp_transformers/sampler.py index ee445260a..bf839bded 100644 --- a/src/nemo_safe_synthesizer/privacy/dp_transformers/sampler.py +++ b/src/nemo_safe_synthesizer/privacy/dp_transformers/sampler.py @@ -21,6 +21,7 @@ import torch from opacus.utils.uniform_sampler import UniformWithReplacementSampler from torch.utils.data.sampler import BatchSampler, RandomSampler, Sampler +from typing_extensions import override from ...observability import get_logger @@ -50,6 +51,7 @@ def __len__(self) -> int: raise TypeError("entity_sampler must implement __len__") return len(self.entity_sampler) + @override def __iter__(self) -> Iterator[list[int]]: """Iterate over batches of dataset indices, one sample per entity per batch. @@ -133,6 +135,7 @@ def __init__(self, *args, **kwargs): self.empty_batches = 0 super().__init__(*args, **kwargs) + @override def __len__(self) -> int: """Return the number of batches that will be yielded (non-empty only). @@ -142,6 +145,7 @@ def __len__(self) -> int: """ return self.steps - self.empty_batches + @override def __iter__(self) -> Iterator[list[int]]: """Iterate over batches, each drawn uniformly with replacement; skip empty batches. diff --git a/src/nemo_safe_synthesizer/sdk/config_builder.py b/src/nemo_safe_synthesizer/sdk/config_builder.py index fed86bbdb..d6f124f29 100644 --- a/src/nemo_safe_synthesizer/sdk/config_builder.py +++ b/src/nemo_safe_synthesizer/sdk/config_builder.py @@ -6,7 +6,7 @@ from __future__ import annotations from collections.abc import Mapping -from typing import Any, Self, TypeAlias, TypeVar, cast +from typing import Self, TypeAlias, TypeVar, overload import pandas as pd from pydantic import BaseModel @@ -27,9 +27,6 @@ logger = get_logger(__name__) -KT = TypeVar("KT") -VT = TypeVar("VT") - NSSParameters = ( DataParameters | EvaluationParameters @@ -43,7 +40,8 @@ ParamT = TypeVar("ParamT", bound=BaseModel) DataSource = pd.DataFrame | str -ParamDict: TypeAlias = dict[str, str | int | float | bool | None | list[Any] | Mapping[KT, VT]] +RawConfig: TypeAlias = Mapping[str, object] +ParamDict: TypeAlias = RawConfig class ConfigBuilder(object): @@ -56,8 +54,8 @@ class ConfigBuilder(object): ``SafeSynthesizerParameters``. Each ``with_*`` method accepts an optional typed config object or - a plain dict, plus ``**kwargs`` overrides. ``kwargs`` always take - precedence over fields in the config/dict. All ``with_*`` methods + a raw mapping, plus ``**kwargs`` overrides. ``kwargs`` always take + precedence over fields in the config/mapping. All ``with_*`` methods return ``Self`` so subclasses preserve their concrete type through fluent chains. @@ -101,14 +99,23 @@ def __init__(self, config: SafeSynthesizerParameters | None = None) -> None: "_time_series_config", ] - def _resolve_config(self, values: ParamDict | NSSParameters | None, cls: type[ParamT], **kwargs) -> ParamT: + @overload + def _resolve_config(self, values: ParamT, cls: type[ParamT], **kwargs: object) -> ParamT: ... + + @overload + def _resolve_config(self, values: RawConfig, cls: type[ParamT], **kwargs: object) -> ParamT: ... + + @overload + def _resolve_config(self, values: None, cls: type[ParamT], **kwargs: object) -> ParamT: ... + + def _resolve_config(self, values: object, cls: type[ParamT], **kwargs: object) -> ParamT: """Resolve configuration from various input types. Precedence: ``kwargs`` override ``values``; ``values`` override model defaults. Args: - values: Existing config, a raw dict, or ``None`` for + values: Existing config, a raw mapping, or ``None`` for defaults-only. cls: The Pydantic model class to validate against. **kwargs: Field-level overrides applied on top. @@ -116,14 +123,19 @@ def _resolve_config(self, values: ParamDict | NSSParameters | None, cls: type[Pa Returns: A validated config instance of type ``cls``. """ - overrides = kwargs match values: - case BaseModel() as model: - return cast(ParamT, model.model_copy(update=overrides)) - case dict() as d: - return cls.model_validate(d).model_copy(update=overrides) case None: - return cls(**overrides) + return cls.model_validate(kwargs) + case BaseModel() as model: + if not isinstance(model, cls): + raise TypeError(f"Expected {cls.__name__}, got {type(model).__name__}") + raw_values = model.model_dump() + raw_values.update(kwargs) + return cls.model_validate(raw_values) + case Mapping() as mapping: + raw_values = dict(mapping) + raw_values.update(kwargs) + return cls.model_validate(raw_values) case _: raise TypeError(f"Unsupported config type: {type(values)}") @@ -139,11 +151,11 @@ def with_data_source(self, df_source: DataSource) -> Self: self._data_source = df_source return self - def with_data(self, config: DataParameters | ParamDict | None = None, **kwargs) -> Self: + def with_data(self, config: DataParameters | RawConfig | None = None, **kwargs: object) -> Self: """Configure data processing settings. Args: - config: Data configuration object or dict. + config: Data configuration object or raw mapping. **kwargs: Field-level overrides (e.g. ``holdout_size``). Returns: @@ -152,11 +164,11 @@ def with_data(self, config: DataParameters | ParamDict | None = None, **kwargs) self._data_config = self._resolve_config(values=config, cls=DataParameters, **kwargs) return self - def with_train(self, config: TrainingHyperparams | ParamDict | None = None, **kwargs) -> Self: + def with_train(self, config: TrainingHyperparams | RawConfig | None = None, **kwargs: object) -> Self: """Configure training hyperparameters. Args: - config: Training configuration object or dict. + config: Training configuration object or raw mapping. **kwargs: Field-level overrides (e.g. ``learning_rate``). Returns: @@ -167,11 +179,11 @@ def with_train(self, config: TrainingHyperparams | ParamDict | None = None, **kw ) return self - def with_generate(self, config: GenerateParameters | ParamDict | None = None, **kwargs) -> Self: + def with_generate(self, config: GenerateParameters | RawConfig | None = None, **kwargs: object) -> Self: """Configure generation settings. Args: - config: Generation configuration object or dict. + config: Generation configuration object or raw mapping. **kwargs: Field-level overrides (e.g. ``num_records``). Returns: @@ -182,11 +194,11 @@ def with_generate(self, config: GenerateParameters | ParamDict | None = None, ** ) return self - def with_time_series(self, config: TimeSeriesParameters | ParamDict | None = None, **kwargs) -> Self: + def with_time_series(self, config: TimeSeriesParameters | RawConfig | None = None, **kwargs: object) -> Self: """Configure time-series synthesis settings. Args: - config: Time-series configuration object or dict. + config: Time-series configuration object or raw mapping. **kwargs: Field-level overrides (e.g. ``time_column``). Returns: @@ -196,12 +208,12 @@ def with_time_series(self, config: TimeSeriesParameters | ParamDict | None = Non return self def with_differential_privacy( - self, config: DifferentialPrivacyHyperparams | ParamDict | None = None, **kwargs + self, config: DifferentialPrivacyHyperparams | RawConfig | None = None, **kwargs: object ) -> Self: """Configure differential privacy settings. Args: - config: DP configuration object or dict. + config: DP configuration object or raw mapping. **kwargs: Field-level overrides (e.g. ``epsilon``). Returns: @@ -211,7 +223,7 @@ def with_differential_privacy( return self def with_replace_pii( - self, config: PiiReplacerConfig | ParamDict | None = None, *, enable: bool = True, **kwargs + self, config: PiiReplacerConfig | RawConfig | None = None, *, enable: bool = True, **kwargs: object ) -> Self: """Configure PII replacement settings. @@ -231,7 +243,7 @@ def with_replace_pii( in ``from_params``. Args: - config: PII replacement configuration object or dict. + config: PII replacement configuration object or raw mapping. enable: When ``False``, disables PII replacement entirely and clears any previously set config. **kwargs: Field-level overrides (e.g. ``classify``). @@ -241,7 +253,7 @@ def with_replace_pii( Raises: ValueError: If ``config`` is not a ``PiiReplacerConfig``, - dict, or ``None``. + raw mapping, or ``None``. Example:: @@ -251,25 +263,26 @@ def with_replace_pii( self._replace_pii_config = None return self - cfg = None match config: - case PiiReplacerConfig() as m: - cfg = m.model_copy(update=kwargs, deep=True) - case dict() as d: - cfg = PiiReplacerConfig.model_validate(d).model_copy(update=kwargs, deep=True) + case PiiReplacerConfig() | Mapping() as values: + cfg = self._resolve_config(values=values, cls=PiiReplacerConfig, **kwargs) case None: - cfg = PiiReplacerConfig.get_default_config().model_copy(update=kwargs, deep=True) + cfg = self._resolve_config( + values=PiiReplacerConfig.get_default_config(), + cls=PiiReplacerConfig, + **kwargs, + ) case _: - raise ValueError(f"Config must be a PiiReplacerConfig, dict, or None, got {config!r}") + raise ValueError(f"Config must be a PiiReplacerConfig, raw mapping, or None, got {config!r}") self._replace_pii_config = cfg return self - def with_evaluate(self, config: EvaluationParameters | ParamDict | None = None, **kwargs) -> Self: + def with_evaluate(self, config: EvaluationParameters | RawConfig | None = None, **kwargs: object) -> Self: """Configure evaluation settings. Args: - config: Evaluation configuration object or dict. + config: Evaluation configuration object or raw mapping. **kwargs: Field-level overrides (e.g. ``enabled``). Returns: diff --git a/src/nemo_safe_synthesizer/sdk/library_builder.py b/src/nemo_safe_synthesizer/sdk/library_builder.py index 678936fd5..98bf006aa 100644 --- a/src/nemo_safe_synthesizer/sdk/library_builder.py +++ b/src/nemo_safe_synthesizer/sdk/library_builder.py @@ -44,6 +44,7 @@ sanitize_model_for_telemetry, ) from ..training.huggingface_backend import HuggingFaceBackend +from ..utils import is_dataframe from .config_builder import ConfigBuilder logger = get_logger(__name__) @@ -110,9 +111,10 @@ def _build_telemetry_event(ss: SafeSynthesizer, status: TaskStatusEnum) -> NSSTr records_bucket = "undefined" columns_bucket = "undefined" - if isinstance(ss._data_source, pd.DataFrame): - records_bucket = bucket_records(len(ss._data_source)) - columns_bucket = bucket_columns(len(ss._data_source.columns)) + data_source = ss._data_source + if is_dataframe(data_source): + records_bucket = bucket_records(len(data_source)) + columns_bucket = bucket_columns(len(data_source.columns)) gpu = get_device_name() diff --git a/src/nemo_safe_synthesizer/telemetry.py b/src/nemo_safe_synthesizer/telemetry.py index 2702c75e7..2930227b1 100644 --- a/src/nemo_safe_synthesizer/telemetry.py +++ b/src/nemo_safe_synthesizer/telemetry.py @@ -23,10 +23,10 @@ from datetime import datetime, timezone from enum import Enum from pathlib import Path, PureWindowsPath -from typing import TYPE_CHECKING, Any, ClassVar, cast +from typing import TYPE_CHECKING, Any, ClassVar from urllib.parse import urlsplit, urlunsplit -from pydantic import BaseModel, Field +from pydantic import BaseModel, ConfigDict, Field from .observability import get_logger @@ -78,7 +78,7 @@ def _redact_endpoint(endpoint: str) -> str: except ValueError: return "" query = "" if parsed.query else "" - return cast(str, urlunsplit((parsed.scheme, parsed.netloc, parsed.path, query, parsed.fragment))) + return urlunsplit((parsed.scheme, parsed.netloc, parsed.path, query, parsed.fragment)) def _deployment_type() -> DeploymentTypeEnum: @@ -169,7 +169,7 @@ class NSSTrainingAndGenerationEvent(BaseModel): nemo_source: NemoSourceEnum = Field( default=NemoSourceEnum.SAFE_SYNTHESIZER, - alias="nemoSource", + serialization_alias="nemoSource", description="The NeMo product that created the event.", ) task: str = Field( @@ -183,72 +183,72 @@ class NSSTrainingAndGenerationEvent(BaseModel): ) deployment_type: DeploymentTypeEnum = Field( default_factory=_deployment_type, - alias="deploymentType", + serialization_alias="deploymentType", description="How Safe Synthesizer was invoked (cli, sdk, nmp).", ) # Timing job_duration_sec: float = Field( default=-1.0, - alias="jobDurationSec", + serialization_alias="jobDurationSec", description="Wall-clock duration of the job in seconds. -1.0 if not available.", ) # Generation metrics num_records_generated: int = Field( default=-1, - alias="numRecordsGenerated", + serialization_alias="numRecordsGenerated", description="Number of valid synthetic records produced. -1 if not available.", ) num_tokens_generated: int = Field( default=-1, - alias="numTokensGenerated", + serialization_alias="numTokensGenerated", description="Number of tokens generated by the model. -1 if not available.", ) # Feature flags replace_pii_enabled: bool = Field( default=False, - alias="replacePiiEnabled", + serialization_alias="replacePiiEnabled", description="Whether PII replacement was enabled for this run.", ) differential_privacy_enabled: bool = Field( default=False, - alias="differentialPrivacyEnabled", + serialization_alias="differentialPrivacyEnabled", description="Whether differential privacy training was enabled for this run.", ) time_series_enabled: bool = Field( default=False, - alias="timeSeriesEnabled", + serialization_alias="timeSeriesEnabled", description="Whether time-series mode was enabled for this run.", ) group_by_enabled: bool = Field( default=False, - alias="groupByEnabled", + serialization_alias="groupByEnabled", description="Whether group-by was set on the input data for this run.", ) # Input characteristics (bucketed to avoid transmitting exact counts) input_records_bucket: str = Field( default="undefined", - alias="inputRecordsBucket", + serialization_alias="inputRecordsBucket", description="Bucketed count of input training records (e.g. '101-1000'). Use bucket_records().", ) input_columns_bucket: str = Field( default="undefined", - alias="inputColumnsBucket", + serialization_alias="inputColumnsBucket", description="Bucketed count of input columns (e.g. '6-10'). Use bucket_columns().", ) # Evaluation scores (-1.0 when evaluation was skipped or unavailable) synthetic_quality_score: float = Field( default=-1.0, - alias="syntheticQualityScore", + serialization_alias="syntheticQualityScore", description="Top-level Synthetic Quality Score from the evaluation report. -1.0 if not available.", ) data_privacy_score: float = Field( default=-1.0, - alias="dataPrivacyScore", + serialization_alias="dataPrivacyScore", description="Top-level Data Privacy Score from the evaluation report. -1.0 if not available.", ) @@ -262,7 +262,7 @@ class NSSTrainingAndGenerationEvent(BaseModel): description="GPU device name (e.g. 'NVIDIA A100 80GB PCIe'). 'undefined' if not on GPU.", ) - model_config = {"populate_by_name": True} + model_config = ConfigDict(populate_by_name=True) @dataclass diff --git a/src/nemo_safe_synthesizer/training/backend.py b/src/nemo_safe_synthesizer/training/backend.py index 8ae53a7a6..79be210c2 100644 --- a/src/nemo_safe_synthesizer/training/backend.py +++ b/src/nemo_safe_synthesizer/training/backend.py @@ -6,6 +6,7 @@ from __future__ import annotations import abc +from collections.abc import Callable from dataclasses import dataclass from pathlib import Path @@ -19,6 +20,7 @@ TrainerCallback, TrainingArguments, ) +from typing_extensions import override from ..cli.artifact_structure import Workdir from ..config import SafeSynthesizerParameters @@ -117,8 +119,8 @@ class TrainingBackend(metaclass=abc.ABCMeta): load_params: dict """Raw parameters used when calling ``from_pretrained``.""" - trainer_type: type[OpacusDPTrainer | Trainer] - """Trainer class to instantiate -- standard ``Trainer`` or ``OpacusDPTrainer`` for DP.""" + trainer_type: Callable[..., Trainer] + """Trainer factory to instantiate standard or DP-aware HuggingFace trainers.""" trainer: OpacusDPTrainer | Trainer """Instantiated trainer, created during ``prepare_params``.""" @@ -187,6 +189,7 @@ def __init__( self.workdir = workdir @classmethod + @override def __subclasshook__(cls, subclass): if cls is not TrainingBackend: return NotImplemented diff --git a/src/nemo_safe_synthesizer/training/callbacks.py b/src/nemo_safe_synthesizer/training/callbacks.py index dc83d4624..5da058a40 100644 --- a/src/nemo_safe_synthesizer/training/callbacks.py +++ b/src/nemo_safe_synthesizer/training/callbacks.py @@ -15,6 +15,7 @@ TrainerState, TrainingArguments, ) +from typing_extensions import override if TYPE_CHECKING: from torch.utils.data import DataLoader @@ -101,6 +102,7 @@ def __init__( "repetition_penalty": kws.get("repetition_penalty", DEFAULT_SAMPLING_PARAMETERS["repetition_penalty"]), } + @override def on_evaluate( self, args: TrainingArguments, @@ -205,6 +207,7 @@ def __init__(self): self.training_bar = None self.prediction_bar = None + @override def on_train_begin(self, args: TrainingArguments, state: TrainerState, control: TrainerControl, **kwargs) -> None: if state.is_world_process_zero: self.training_bar = tqdm( @@ -214,11 +217,13 @@ def on_train_begin(self, args: TrainingArguments, state: TrainerState, control: ) self.current_step = 0 + @override def on_step_end(self, args: TrainingArguments, state: TrainerState, control: TrainerControl, **kwargs) -> None: if state.is_world_process_zero and self.training_bar is not None: self.training_bar.update(state.global_step - self.current_step) self.current_step = state.global_step + @override def on_prediction_step( self, args: TrainingArguments, @@ -242,12 +247,14 @@ def on_prediction_step( ) self.prediction_bar.update(1) + @override def on_evaluate(self, args: TrainingArguments, state: TrainerState, control: TrainerControl, **kwargs) -> None: if state.is_world_process_zero: if self.prediction_bar is not None: self.prediction_bar.close() self.prediction_bar = None + @override def on_predict( self, args: TrainingArguments, @@ -261,6 +268,7 @@ def on_predict( self.prediction_bar.close() self.prediction_bar = None + @override def on_log( self, args: TrainingArguments, @@ -284,6 +292,7 @@ def on_log( if "loss" in logs: self.training_bar.set_description(f"Training in progress [loss = {logs['loss']: .4f}]") + @override def on_train_end(self, args: TrainingArguments, state: TrainerState, control: TrainerControl, **kwargs) -> None: if state.is_world_process_zero and self.training_bar is not None: self.training_bar.close() @@ -310,6 +319,7 @@ def __init__(self, log_interval: float = 60.0): self._last_log_ts = time.monotonic() self._last_log_global_step = 0 + @override def on_train_begin( self, args: TrainingArguments, @@ -333,6 +343,7 @@ def _checked_log_if(self, cond: bool, state: TrainerState, control: TrainerContr return control return None + @override def on_epoch_end( self, args: TrainingArguments, @@ -344,6 +355,7 @@ def on_epoch_end( # when the last log was emitted. return self._checked_log_if(True, state, control) + @override def on_substep_end( self, args: TrainingArguments, @@ -357,6 +369,7 @@ def on_substep_end( # We leave it to `on_log` below to actually reset the last log timestamp. return self._checked_log_if(time.monotonic() - self._last_log_ts >= self._log_interval, state, control) + @override def on_log( self, args: TrainingArguments, diff --git a/src/nemo_safe_synthesizer/training/huggingface_backend.py b/src/nemo_safe_synthesizer/training/huggingface_backend.py index 264331785..e4be63210 100644 --- a/src/nemo_safe_synthesizer/training/huggingface_backend.py +++ b/src/nemo_safe_synthesizer/training/huggingface_backend.py @@ -11,9 +11,8 @@ import time from collections.abc import Callable from contextlib import redirect_stdout -from functools import partial from pathlib import Path -from typing import TYPE_CHECKING, Any, cast +from typing import TYPE_CHECKING, Any import numpy as np import pandas as pd @@ -36,6 +35,7 @@ ) from transformers.trainer_pt_utils import get_model_param_count from transformers.utils.quantization_config import QuantizationConfigMixin +from typing_extensions import override from .. import utils from ..cli.artifact_structure import BoundDir @@ -84,6 +84,11 @@ DEFAULT_ROPE_THETA = 10000.0 +# Training arguments fixed by Safe Synthesizer at runtime. +# +# Training duration is controlled by ``num_input_records_to_sample`` and the +# assembled ``data_fraction``, not by epochs. These values keep the HuggingFace +# Trainer behavior stable across CLI and SDK entry points. FIXED_RUNTIME_TRAINING_ARGS = { # the training time is set by the number of training records "num_train_epochs": 1, @@ -95,12 +100,27 @@ "bf16": True, "ddp_find_unused_parameters": False, } -"""Training arguments fixed by Safe Synthesizer at runtime. -Training duration is controlled by ``num_input_records_to_sample`` and the -assembled ``data_fraction``, not by epochs. These values keep the HuggingFace -Trainer behavior stable across CLI and SDK entry points. -""" + +def _standard_trainer_factory(**kwargs: Any) -> Trainer: + return Trainer(**kwargs) + + +def _opacus_trainer_factory( + *, + privacy_args: PrivacyArguments, + true_dataset_size: int, + data_fraction: float, +) -> Callable[..., Trainer]: + def factory(**kwargs: Any) -> Trainer: + return OpacusDPTrainer( + privacy_args=privacy_args, + true_dataset_size=true_dataset_size, + data_fraction=data_fraction, + **kwargs, + ) + + return factory class HuggingFaceBackend(TrainingBackend): @@ -119,7 +139,7 @@ class HuggingFaceBackend(TrainingBackend): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - self.trainer_type: type[Trainer] | partial[OpacusDPTrainer] = Trainer + self.trainer_type: Callable[..., Trainer] = _standard_trainer_factory self.model_loader_type = AutoModelForCausalLM self.training_output_dir = Path(self.workdir.train.cache) self.model_ref = ModelRef.parse(self.params.training.pretrained_model) @@ -306,6 +326,7 @@ def _apply_rope_scaling(self, framework_params: dict, **kwargs: Any) -> None: logger.warning(msg) @traced_runtime("prepare_config") + @override def prepare_config(self, add_max_memory: bool = True, **kwargs: Any) -> None: """Set common model arguments for initializing a model. @@ -371,6 +392,7 @@ def _prepare_quantize_base(self, **quantize_params: dict) -> None: logger.info(f"using loftq with {scheme.effective_bits} bits") self.quant_params["loftq_config"] = LoftQConfig(loftq_bits=scheme.effective_bits) + @override def maybe_quantize(self, **quant_params: dict) -> None: """Apply LoRA wrapping (and optional k-bit quantization) to the model.""" self._prepare_quantize_base(**quant_params) @@ -392,6 +414,7 @@ def maybe_quantize(self, **quant_params: dict) -> None: f"Using PEFT - {parameter_count:.2f} million parameters are trainable", ) + @override def load_model(self, **model_args: Any) -> None: """Load an ``AutoModelForCausalLM`` instance with specified arguments. @@ -494,8 +517,7 @@ def _configure_dp_training(self, training_args: dict) -> DataCollatorForPrivateT per_sample_max_grad_norm=privacy.per_sample_max_grad_norm, ) - self.trainer_type = partial( # ty: ignore[invalid-assignment] -- partial is assignable at runtime - OpacusDPTrainer, + self.trainer_type = _opacus_trainer_factory( privacy_args=privacy_args, true_dataset_size=self.true_dataset_size, data_fraction=self.data_fraction, @@ -537,8 +559,7 @@ def _create_trainer( Returns: The configured Trainer instance. """ - factory = cast(Callable[..., Trainer], self.trainer_type) - trainer = factory( + trainer = self.trainer_type( model=self.model, processing_class=self.tokenizer, args=training_args, @@ -601,6 +622,7 @@ def _add_inference_eval_callback(self, trainer: Trainer, training_args: dict) -> ) @traced_runtime("prepare_params") + @override def prepare_params(self, **training_args: Any) -> None: """Prepare training parameters and create the trainer. @@ -700,6 +722,7 @@ def _log_dataset_statistics(self, assembler: TrainingExampleAssembler) -> None: } logger.user.info("", extra=extra) + @override def prepare_training_data(self) -> None: """Validate, preprocess, and tokenize the training dataset. @@ -778,6 +801,7 @@ def _propagate_max_tokens_per_example(self) -> None: self.model_metadata.max_tokens_per_example = int(math.ceil(tokens_per_example.max)) @utils.time_function + @override def train(self, **training_args: Any) -> None: """Run the full training pipeline and populate ``results``. @@ -807,6 +831,7 @@ def train(self, **training_args: Any) -> None: elapsed_time=training_time_sec, ) + @override def save_model(self) -> None: """Save the fine-tuning adapter and related artifacts under ``self.workdir``. @@ -837,6 +862,7 @@ def save_model(self) -> None: indent=4, ) + @override def teardown(self) -> None: """Release GPU memory, distributed resources, and trainer state. Idempotent -- safe to call multiple times.""" if getattr(self, "_torn_down", False): @@ -871,6 +897,7 @@ def teardown(self) -> None: except Exception: logger.warning("destroy_process_group failed during teardown", exc_info=True) + @override def __str__(self): f = f"HuggingFaceBackend(pretrained_model={self.params.training.pretrained_model}, params={self.params})" return f diff --git a/src/nemo_safe_synthesizer/utils.py b/src/nemo_safe_synthesizer/utils.py index 93ac5bac8..707dc74b2 100644 --- a/src/nemo_safe_synthesizer/utils.py +++ b/src/nemo_safe_synthesizer/utils.py @@ -13,13 +13,14 @@ import json import os import time -from collections.abc import Callable, Generator, Iterable +from collections.abc import Callable, Generator, Iterable, Mapping from pathlib import Path -from typing import TYPE_CHECKING, Any, Protocol +from typing import TYPE_CHECKING, Any, ParamSpec, Protocol, TypeGuard, TypeVar import numpy as np import pandas as pd from pandas import DataFrame +from typing_extensions import TypeIs from .data_processing.stats import Statistics from .observability import get_logger @@ -36,6 +37,13 @@ _HF_OFFLINE_ENV_VARS = ("HF_HUB_OFFLINE", "TRANSFORMERS_OFFLINE") +P = ParamSpec("P") +R = TypeVar("R") + + +def _is_statistics_list(stats: Statistics | list[Statistics]) -> TypeIs[list[Statistics]]: + return isinstance(stats, list) + def env_flag_is_true(name: str, *, default: bool = False) -> bool: """Return whether ``name`` is set to a truthy env value. @@ -125,11 +133,11 @@ def log_stats( title: Optional table title. """ headers = headers or [] - stats = stats if isinstance(stats, list) else [stats] + stats_list = stats if _is_statistics_list(stats) else [stats] # Build structured data - processor will render as table for console structured_stats = {} - for header, stat in zip(headers, stats): + for header, stat in zip(headers, stats_list): key = header.lower().replace(" ", "_") structured_stats[key] = { "min": round_number_if_float(stat.min), @@ -153,11 +161,11 @@ def log_stats( ) -def log_training_example_stats(stats_dict: dict[str, Statistics], **kwargs) -> None: +def log_training_example_stats(stats_dict: dict[str, Statistics]) -> None: """Log training example statistics from the given dictionary.""" stats = list(stats_dict.values()) headers = list([name.replace("_", " ").capitalize() for name in stats_dict.keys()]) - log_stats(title="Training Example Statistics", stats=stats, headers=headers, **kwargs) + log_stats(title="Training Example Statistics", stats=stats, headers=headers) def round_number_if_float(number: int | float, precision: int = 3) -> int | float: @@ -180,7 +188,7 @@ def smart_read_table(df_or_path: str | Path | pd.DataFrame) -> pd.DataFrame: Raises: ValueError: If the file extension is not supported. """ - if isinstance(df_or_path, pd.DataFrame): + if is_dataframe(df_or_path): return df_or_path path = str(df_or_path) @@ -200,11 +208,11 @@ def smart_read_table(df_or_path: str | Path | pd.DataFrame) -> pd.DataFrame: return df -def time_function(func: Callable[..., Any]) -> Callable[..., Any]: +def time_function(func: Callable[P, R]) -> Callable[P, R]: """Decorator to log the time taken by a function to execute.""" @functools.wraps(func) - def time_closure(*args: Any, **kwargs: Any) -> Any: + def time_closure(*args: P.args, **kwargs: P.kwargs) -> R: start = time.perf_counter() result = func(*args, **kwargs) time_elapsed = time.perf_counter() - start @@ -255,24 +263,24 @@ def debug_fmt(df: pd.DataFrame) -> str: return df.head(5).to_json(orient="records", date_format="iso") -def merge_dicts(base: dict, new: dict) -> dict: +def merge_dicts(base: Mapping[str, Any], new: Mapping[str, Any]) -> dict[str, Any]: """Deep-merge two dicts, preferring values from ``new`` on conflict.""" - result = base.copy() + result = dict(base) for k, new_v in new.items(): base_v = result.get(k) - if isinstance(base_v, dict) and isinstance(new_v, dict): + if isinstance(base_v, Mapping) and isinstance(new_v, Mapping): result[k] = merge_dicts(base_v, new_v) else: result[k] = new_v return result -def is_iterable(x: object) -> bool: +def is_iterable(x: object) -> TypeGuard[Iterable[object]]: """Check whether ``x`` has both ``__iter__`` and ``__getitem__``.""" return hasattr(x, "__iter__") and hasattr(x, "__getitem__") -def flatten(iter: Iterable) -> Generator: +def flatten(iter: Iterable[object]) -> Generator[object]: """Flatten a possibly nested iterable. Strings are yielded as-is (not broken into characters). Dicts are @@ -311,8 +319,13 @@ def typecheck(x: object) -> bool: return True +def is_dataframe(x: object) -> TypeIs[pd.DataFrame]: + """Return whether ``x`` is a pandas ``DataFrame``.""" + return isinstance(x, pd.DataFrame) + + def write_json( - data: dict, + data: Mapping[str, object], path: str | os.PathLike[str], encoding: str | None = None, indent: int | None = None, @@ -324,7 +337,7 @@ def write_json( json.dump(data, file, indent=indent) -def load_json(path: str | Path, encoding: str | None = None) -> dict: +def load_json(path: str | Path, encoding: str | None = None) -> dict[str, Any]: """Load JSON file and return the content as a dict.""" with Path(path).open(encoding=encoding) as file: return json.load(file) diff --git a/tests/config/test_autoconfig.py b/tests/config/test_autoconfig.py index 189b2627f..9d7c1d91f 100644 --- a/tests/config/test_autoconfig.py +++ b/tests/config/test_autoconfig.py @@ -72,9 +72,9 @@ class AutoConfigTestCase: def get_config(self) -> SafeSynthesizerParameters: """Get the config, calling it if it's a factory function.""" - if callable(self.config): - return self.config() # ty: ignore[call-top-callable] -- dynamic callable - return self.config + if isinstance(self.config, SafeSynthesizerParameters): + return self.config + return self.config() AUTO_NO_DP = AutoConfigTestCase( diff --git a/tests/config/test_parameters.py b/tests/config/test_parameters.py index 1dda3cd2e..f19e489b7 100644 --- a/tests/config/test_parameters.py +++ b/tests/config/test_parameters.py @@ -133,6 +133,27 @@ def test_from_params_none_disables_pii(): assert SafeSynthesizerParameters.from_params(replace_pii=None).replace_pii is None +def test_from_config_patch_validates_sparse_config(): + config = SafeSynthesizerParameters.from_config_patch({"replace_pii": None}) + + assert config.replace_pii is None + + +def test_with_config_patch_merges_sparse_patch_and_keeps_defaults_implicit(): + config = SafeSynthesizerParameters.model_validate({"generation": {"num_records": 77}}) + + merged = config.with_config_patch({"generation": {"temperature": 0.7}, "training": {"batch_size": 4}}) + + assert merged.generation.num_records == 77 + assert merged.generation.temperature == 0.7 + assert merged.generation.use_structured_generation is False + assert merged.training.batch_size == 4 + assert merged.model_dump(exclude_unset=True) == { + "generation": {"num_records": 77, "temperature": 0.7}, + "training": {"batch_size": 4}, + } + + def _resolve(obj: object, path: str) -> object: """Resolve a dotted attribute ``path`` (e.g. ``generation.validation.foo``).""" for part in path.split("."): diff --git a/tests/configurator/test_pydantic_click_options.py b/tests/configurator/test_pydantic_click_options.py index 10b8162e8..694e38508 100644 --- a/tests/configurator/test_pydantic_click_options.py +++ b/tests/configurator/test_pydantic_click_options.py @@ -355,6 +355,7 @@ def cmd(**kwargs): pass flag = next(p for p in cmd.params if p.name == "no_nested") + assert isinstance(flag, click.Option) assert flag.is_flag 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/data_processing/test_budget.py b/tests/data_processing/test_budget.py index 4ca5a51df..828e62b49 100644 --- a/tests/data_processing/test_budget.py +++ b/tests/data_processing/test_budget.py @@ -11,6 +11,7 @@ import pandas as pd import pytest from transformers import BatchEncoding, PreTrainedTokenizerBase +from typing_extensions import override from nemo_safe_synthesizer.data_processing.budget import ( compute_max_new_tokens, @@ -26,11 +27,13 @@ def __init__(self) -> None: self.texts: list[str] = [] self.encoded_text = "" + @override def __call__(self, texts: list[str], *, add_special_tokens: bool) -> dict[str, list[list[int]]]: assert add_special_tokens is False self.texts = texts return {"input_ids": [[ord(char) for char in text] for text in texts]} + @override def encode(self, text: str, *, add_special_tokens: bool) -> list[int]: assert add_special_tokens is False self.encoded_text = text @@ -79,10 +82,12 @@ class _BatchEncodingTokenizer(PreTrainedTokenizerBase): def __init__(self) -> None: self.encode_calls = 0 + @override def __call__(self, texts: list[str], *, add_special_tokens: bool) -> BatchEncoding: assert add_special_tokens is False return BatchEncoding({"input_ids": [[ord(char) for char in text] for text in texts]}) + @override def encode(self, text: str, *, add_special_tokens: bool) -> list[int]: assert add_special_tokens is False self.encode_calls += 1 diff --git a/tests/data_processing/test_data_actions.py b/tests/data_processing/test_data_actions.py new file mode 100644 index 000000000..d82c03907 --- /dev/null +++ b/tests/data_processing/test_data_actions.py @@ -0,0 +1,48 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import json + +import pandas as pd + +from nemo_safe_synthesizer.data_processing.actions.data_actions import DatetimeCol, ReplaceDataSource +from nemo_safe_synthesizer.data_processing.actions.utils import ActionCtx, UniqueIdSource + + +def test_replace_datasource_state_stays_json_and_restores_column_index(): + ctx = ActionCtx() + action = ReplaceDataSource(col="replacement", data_source=UniqueIdSource()).with_ctx(ctx) + source = pd.DataFrame( + { + "left": [1, 2], + "replacement": ["old-a", "old-b"], + "right": [3, 4], + } + ) + + preprocessed = action.preprocess(source) + + assert list(preprocessed.columns) == ["left", "right"] + assert ctx.state[action.hash()] == '{"column_index":1}' + + generated = action.generate(preprocessed) + + assert list(generated.columns) == ["left", "replacement", "right"] + assert generated["replacement"].notna().all() + + +def test_datetime_col_state_stays_json_and_validates_with_inferred_format(): + ctx = ActionCtx() + action = DatetimeCol(name="started_at").with_ctx(ctx) + source = pd.DataFrame({"started_at": ["2024-01-20", "2024-01-21"]}) + + action.preprocess(source) + + state = json.loads(ctx.state[action.hash()]) + assert state == {"dt_format": "%Y-%m-%d"} + + batch = pd.DataFrame({"started_at": ["2024-01-22", "not a date"]}) + valid, rejected = action.validate_batch(batch, pd.DataFrame()) + + assert valid["started_at"].tolist() == ["2024-01-22"] + assert rejected["started_at"].tolist() == ["not a date"] diff --git a/tests/data_processing/test_distributions.py b/tests/data_processing/test_distributions.py new file mode 100644 index 000000000..088054dcf --- /dev/null +++ b/tests/data_processing/test_distributions.py @@ -0,0 +1,21 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from datetime import datetime, timedelta + +from typing_extensions import override + +from nemo_safe_synthesizer.data_processing.actions.distributions import DatetimeDistribution + + +class FixedDatetimeDistribution(DatetimeDistribution): + @override + def sample_datetimes(self, num_records: int) -> list[datetime]: + start = datetime(2024, 1, 1, 12, 20) + return [start + timedelta(minutes=20 * offset) for offset in range(num_records)] + + +def test_datetime_distribution_applies_precision_and_format_without_mutating_samples(): + distribution = FixedDatetimeDistribution(precision=timedelta(hours=1), format="%H:%M") + + assert distribution.sample(2) == ["12:00", "13:00"] diff --git a/tests/data_processing/test_records.py b/tests/data_processing/test_records.py index d1670a36f..8a1ff1a22 100644 --- a/tests/data_processing/test_records.py +++ b/tests/data_processing/test_records.py @@ -18,6 +18,8 @@ is_safe_for_float_conversion, normalize_dataframe, ) +from nemo_safe_synthesizer.data_processing.records.json_record import flatten +from nemo_safe_synthesizer.data_processing.records.json_types import is_json_object, is_json_value def _mock_encode(text: str) -> list[int]: @@ -25,6 +27,24 @@ def _mock_encode(text: str) -> list[int]: return list(range(len(text))) +def test_flatten_handles_top_level_array_and_nested_records(): + assert flatten([{"name": "Alice"}, {"name": "Bob", "scores": [1, 2]}]) == { + "_nssarray_0*#N#*name": "Alice", + "_nssarray_1*#N#*name": "Bob", + "_nssarray_1*#N#*scores*#N#*_nssarray_0": 1, + "_nssarray_1*#N#*scores*#N#*_nssarray_1": 2, + } + + +def test_json_type_guards_accept_recursive_json_objects(): + value = {"name": "Alice", "scores": [1, 2, None], "profile": {"active": True}} + + assert is_json_value(value) is True + assert is_json_object(value) is True + assert is_json_object({1: "not-json-object"}) is False + assert is_json_value({"bad": object()}) is False + + def test_is_safe_for_float_conversion(): # Test with safe values assert is_safe_for_float_conversion(100) is True diff --git a/tests/evaluation/components/benchmark_nearest_neighbor.py b/tests/evaluation/components/benchmark_nearest_neighbor.py index d7d21c54c..a9193db61 100644 --- a/tests/evaluation/components/benchmark_nearest_neighbor.py +++ b/tests/evaluation/components/benchmark_nearest_neighbor.py @@ -20,6 +20,8 @@ import logging import os +from typing_extensions import override + # Must be set BEFORE importing numpy/sklearn (which load OpenBLAS) # Cap threads to avoid OpenBLAS crashes on high-core machines (>128 cores) os.environ.setdefault("OPENBLAS_NUM_THREADS", "8") @@ -96,6 +98,7 @@ class BenchmarkResult: n_queries: int k: int + @override def __str__(self) -> str: return ( f"{self.name}: fit={self.fit_time_sec:.3f}s, " diff --git a/tests/generation/test_generation.py b/tests/generation/test_generation.py index 952634075..aaa04634b 100644 --- a/tests/generation/test_generation.py +++ b/tests/generation/test_generation.py @@ -13,6 +13,7 @@ DateConstraint, data_actions_fn, ) +from nemo_safe_synthesizer.data_processing.record_utils import normalize_record_keys from nemo_safe_synthesizer.errors import GenerationError from nemo_safe_synthesizer.generation.batch import Batch from nemo_safe_synthesizer.generation.processors import ParsedRecord, ParsedResponse @@ -260,11 +261,15 @@ def test_apply_data_actions(fixture_mock_processor, caplog): batch = Batch(fixture_mock_processor) batch._responses = [ ParsedResponse( - records=[ParsedRecord(text=str(r), parsed=r) for r in data.iloc[:3].to_dict("records")], + records=[ + ParsedRecord(text=str(r), parsed=normalize_record_keys(r)) for r in data.iloc[:3].to_dict("records") + ], prompt_number=1, ), ParsedResponse( - records=[ParsedRecord(text=str(r), parsed=r) for r in data.iloc[3:].to_dict("records")], + records=[ + ParsedRecord(text=str(r), parsed=normalize_record_keys(r)) for r in data.iloc[3:].to_dict("records") + ], prompt_number=2, ), ] 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..9c55b8905 --- /dev/null +++ b/tests/pii_replacer/test_ner_model.py @@ -0,0 +1,117 @@ +# 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.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) diff --git a/tests/preflight/test_plugin_registration.py b/tests/preflight/test_plugin_registration.py index 10ef7518e..626b11db3 100644 --- a/tests/preflight/test_plugin_registration.py +++ b/tests/preflight/test_plugin_registration.py @@ -10,6 +10,7 @@ import pandas as pd import pytest +from typing_extensions import override from nemo_safe_synthesizer.config.parameters import SafeSynthesizerParameters from nemo_safe_synthesizer.config.preflight import PreflightParameters @@ -56,6 +57,7 @@ class MyCheck(ConfigCheck): name = "myplugin.registered" label = "Registered plugin" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: collector.warning("myplugin_fired", "hello") @@ -74,6 +76,7 @@ class MyCheck(ConfigCheck): name = "myplugin.fires" label = "Plugin fires" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: collector.warning("plugin_signal", "plugin ran") @@ -101,6 +104,7 @@ class BadPlugin(ConfigCheck): name = "gpu.rogue" label = "Rogue plugin" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: return @@ -128,6 +132,7 @@ class First(ConfigCheck): name = "myplugin.first" label = "First" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: return @@ -135,6 +140,7 @@ class Duplicate(ConfigCheck): name = "myplugin.first" label = "Also first" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: return @@ -142,6 +148,7 @@ class Third(ConfigCheck): name = "myplugin.third" label = "Third" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: return @@ -217,6 +224,7 @@ class CoreCollision(ConfigCheck): name = "myplugin.collides" label = "Collides" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: return @@ -259,6 +267,7 @@ class PrereqPlugin(ConfigCheck): name = "myplugin.prereq" label = "Prereq" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: collector.warning("prereq_ran", "should not fire when disabled") @@ -267,6 +276,7 @@ class DependentPlugin(ConfigCheck): label = "Dependent" requires = ("myplugin.prereq",) + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: collector.warning("dependent_ran", "should not fire when prereq disabled") @@ -316,6 +326,7 @@ class SoloCheck(ConfigCheck): name = "myplugin.solo" label = "Solo" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: collector.warning("solo_ran", "solo") @@ -346,6 +357,7 @@ class Crasher(ConfigCheck): name = "myplugin.crasher" label = "Crasher" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: raise RuntimeError("boom") @@ -353,6 +365,7 @@ class Follower(ConfigCheck): name = "myplugin.follower" label = "Follower" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: collector.warning("followed", "ran after crash") @@ -388,6 +401,7 @@ class P1(ConfigCheck): name = "myplugin.stage_cfg" label = "Plugin cfg" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: return @@ -395,6 +409,7 @@ class P2(DataFrameCheck): name = "myplugin.stage_df" label = "Plugin df" + @override def check(self, ctx: DataFrameView, collector: IssueCollector) -> None: return diff --git a/tests/preflight/test_preflight.py b/tests/preflight/test_preflight.py index 143157c7e..a3a5db648 100644 --- a/tests/preflight/test_preflight.py +++ b/tests/preflight/test_preflight.py @@ -14,6 +14,7 @@ import pytest from rich.console import Console from transformers import PretrainedConfig, PreTrainedTokenizerBase +from typing_extensions import override from nemo_safe_synthesizer.config.data import DataParameters from nemo_safe_synthesizer.config.parameters import SafeSynthesizerParameters @@ -58,10 +59,12 @@ def _issue_by_code(issues: list[PreflightIssue], code: str) -> PreflightIssue: class _PseudoColumnSensitiveTokenizer(PreTrainedTokenizerBase): """Tokenizer that makes pseudo-column leakage visible in budget tests.""" + @override def encode(self, text: str, *, add_special_tokens: bool) -> list[int]: assert add_special_tokens is False return [] + @override def __call__(self, texts: list[str], *, add_special_tokens: bool) -> dict[str, list[list[int]]]: assert add_special_tokens is False return {"input_ids": [[0] * (100 if PSEUDO_GROUP_COLUMN in text else 1) for text in texts]} @@ -1265,6 +1268,7 @@ class AlwaysError(ConfigCheck): name = "plugintest.base" label = "Base" + @override def check(self, ctx, collector): collector.error("test_err", "forced error") @@ -1273,6 +1277,7 @@ class AlwaysPass(ConfigCheck): label = "Dependent" requires = ("plugintest.base",) + @override def check(self, ctx, collector): return @@ -1280,6 +1285,7 @@ class Independent(ConfigCheck): name = "plugintest.indep" label = "Independent" + @override def check(self, ctx, collector): return @@ -1297,6 +1303,7 @@ class WarnOnly(ConfigCheck): name = "plugintest.warn_base" label = "Base" + @override def check(self, ctx, collector): collector.warning("test_warn", "just a warning") @@ -1305,6 +1312,7 @@ class AlwaysPass(ConfigCheck): label = "Dependent" requires = ("plugintest.warn_base",) + @override def check(self, ctx, collector): return @@ -1321,6 +1329,7 @@ class AdvisoryError(AdvisoryCheck): name = "plugintest.advisory_base" label = "Advisory base" + @override def check(self, ctx, collector): collector.error("advisory_err", "advisory error") @@ -1329,6 +1338,7 @@ class Dependent(ConfigCheck): label = "Dependent" requires = ("plugintest.advisory_base",) + @override def check(self, ctx, collector): return diff --git a/tests/preflight/test_registry_validation.py b/tests/preflight/test_registry_validation.py index c81f16090..fa405d0af 100644 --- a/tests/preflight/test_registry_validation.py +++ b/tests/preflight/test_registry_validation.py @@ -6,6 +6,7 @@ from __future__ import annotations import pytest +from typing_extensions import override from nemo_safe_synthesizer.preflight import ( AdvisoryCheck, @@ -25,6 +26,7 @@ class _NoopConfig(ConfigCheck): name = "plugintest.noop_cfg" label = "Noop config" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: return @@ -33,6 +35,7 @@ class _NoopConfigB(ConfigCheck): name = "plugintest.noop_cfg_b" label = "Noop config B" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: return @@ -41,6 +44,7 @@ class _NoopDataFrame(DataFrameCheck): name = "plugintest.noop_df" label = "Noop df" + @override def check(self, ctx: DataFrameView, collector: IssueCollector) -> None: return @@ -49,6 +53,7 @@ class _NoopMetadata(MetadataCheck): name = "plugintest.noop_meta" label = "Noop meta" + @override def check(self, ctx: MetadataView, collector: IssueCollector) -> None: return @@ -57,6 +62,7 @@ class _NoopAdvisory(AdvisoryCheck): name = "plugintest.noop_adv" label = "Noop advisory" + @override def check(self, ctx: DataFrameView, collector: IssueCollector) -> None: return @@ -84,6 +90,7 @@ class NeedsMissing(ConfigCheck): label = "Needs missing" requires = ("plugintest.does_not_exist",) + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: return @@ -97,6 +104,7 @@ class LeaderConfig(ConfigCheck): name = "plugintest.leader" label = "Leader" + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: return @@ -105,6 +113,7 @@ class NeedsLeader(ConfigCheck): label = "Needs leader" requires = ("plugintest.leader",) + @override def check(self, ctx: ConfigView, collector: IssueCollector) -> None: return diff --git a/tests/sdk/test_builder.py b/tests/sdk/test_builder.py index ae6cf7183..9108cbdd3 100644 --- a/tests/sdk/test_builder.py +++ b/tests/sdk/test_builder.py @@ -12,7 +12,7 @@ DEFAULT_PII_TRANSFORM_CONFIG, PiiReplacerConfig, ) -from nemo_safe_synthesizer.sdk.library_builder import SafeSynthesizer, _emit_nss_telemetry +from nemo_safe_synthesizer.sdk.library_builder import SafeSynthesizer, _build_telemetry_event, _emit_nss_telemetry from nemo_safe_synthesizer.telemetry import DeploymentTypeEnum, TaskStatusEnum _SMALL_DF = pd.DataFrame({"a": [1, 2, 3]}) @@ -401,6 +401,16 @@ def _builder_for_telemetry() -> SafeSynthesizer: class TestTelemetryEmission: + def test_build_telemetry_event_buckets_dataframe_input(self, monkeypatch): + monkeypatch.setattr("nemo_safe_synthesizer.sdk.library_builder.get_device_name", lambda: "undefined") + builder = _builder_for_telemetry() + builder._data_source = pd.DataFrame({f"col_{idx}": range(250) for idx in range(6)}) + + event = _build_telemetry_event(builder, TaskStatusEnum.COMPLETED) + + assert event.input_records_bucket == "201-1000" + assert event.input_columns_bucket == "6-10" + def test_run_emits_completed_after_save_results(self, monkeypatch, tmp_path: Path): builder = SafeSynthesizer(config=SafeSynthesizerParameters(), save_path=tmp_path) emitted = [] diff --git a/tests/sdk/test_config_builder.py b/tests/sdk/test_config_builder.py new file mode 100644 index 000000000..acc114666 --- /dev/null +++ b/tests/sdk/test_config_builder.py @@ -0,0 +1,43 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from typing import Any, cast + +import pytest +from pydantic import ValidationError + +from nemo_safe_synthesizer.config import GenerateParameters +from nemo_safe_synthesizer.config.replace_pii import PiiReplacerConfig +from nemo_safe_synthesizer.sdk.config_builder import ConfigBuilder + + +def test_with_generate_validates_raw_config_with_kwargs(): + with pytest.raises(ValidationError, match="patience"): + ConfigBuilder().with_generate(config={"num_records": 10}, patience=0) + + +def test_with_generate_validates_typed_config_with_kwargs(): + with pytest.raises(ValidationError, match="patience"): + ConfigBuilder().with_generate(config=GenerateParameters(num_records=10), patience=0) + + +def test_with_generate_rejects_wrong_typed_config_object(): + wrong_config = cast(Any, PiiReplacerConfig.get_default_config()) + + with pytest.raises(TypeError, match="Expected GenerateParameters"): + ConfigBuilder().with_generate(config=wrong_config) + + +def test_with_replace_pii_validates_default_config_with_kwargs(): + with pytest.raises(ValidationError, match="Invalid locale"): + ConfigBuilder().with_replace_pii(globals={"locales": ["not-a-locale"]}) + + +def test_with_replace_pii_resolves_raw_config_with_kwargs(): + builder = ConfigBuilder().with_replace_pii( + config=PiiReplacerConfig.get_default_config().model_dump(), + globals={"classify": {"enable_classify": False}}, + ) + + assert builder._replace_pii_config is not None + assert builder._replace_pii_config.globals.classify.enable_classify is False diff --git a/tools/diff-lockfile.py b/tools/diff-lockfile.py index f5f3ed10a..c69a09673 100644 --- a/tools/diff-lockfile.py +++ b/tools/diff-lockfile.py @@ -36,7 +36,7 @@ import subprocess import tomllib from enum import Enum -from typing import Annotated, Optional +from typing import Annotated, Any, Optional import git import typer @@ -116,12 +116,17 @@ def get_lockfile_content(repo: git.Repo, ref: str, path: str) -> str: return blob.data_stream.read().decode() -def _extract_source(raw: dict) -> str: +def _extract_source(raw: dict[str, Any]) -> str: """Return a human-readable source string from a ``[[package]]`` entry.""" - src = raw.get("source", {}) - if isinstance(src, dict): - return src.get("registry", src.get("git", src.get("path", ""))) - return str(src) if src else "" + match raw.get("source", {}): + case {"registry": value} | {"git": value} | {"path": value}: + return str(value) if value is not None else "" + case dict(): + return "" + case src if src: + return str(src) + case _: + return "" def parse_packages(content: str) -> dict[str, Package]: diff --git a/tools/hf_network_guard_proxy.py b/tools/hf_network_guard_proxy.py index 673c389fe..df9791667 100644 --- a/tools/hf_network_guard_proxy.py +++ b/tools/hf_network_guard_proxy.py @@ -8,6 +8,7 @@ # "pydantic", # "proxy.py>=2.4.10", # "structlog", +# "typing-extensions>=4.15.0", # ] # /// # pyright: reportMissingImports=false @@ -74,7 +75,7 @@ import sys import threading import time -from contextlib import nullcontext +from contextlib import AbstractContextManager, nullcontext from dataclasses import dataclass, field from enum import IntEnum from multiprocessing import Manager @@ -90,6 +91,7 @@ from proxy.http.parser import HttpParser from proxy.http.proxy import HttpProxyBasePlugin from pydantic import BaseModel, ConfigDict, Field +from typing_extensions import override DEFAULT_HOST = "127.0.0.1" DEFAULT_PORT = 8765 @@ -170,7 +172,7 @@ def summary(self) -> ProxySummary: class RunningServer: """Background proxy.py server instance.""" - server: Proxy + server: AbstractContextManager[object] host: str port: int state: ProxyState @@ -201,6 +203,7 @@ def configure(cls, *, state: ProxyState, allow_passthrough_hf: bool) -> None: cls.state = state cls.allow_passthrough_hf = allow_passthrough_hf + @override def before_upstream_connection(self, request: HttpParser) -> HttpParser | None: """Record the destination and decide whether proxy.py may connect upstream.""" proxy_request = _proxy_request_from_proxy_py(request) diff --git a/tools/patch_dependabot.py b/tools/patch_dependabot.py index 22db97894..62355aa7a 100644 --- a/tools/patch_dependabot.py +++ b/tools/patch_dependabot.py @@ -7,6 +7,7 @@ # "requests>=2.33.0", # "tomlkit>=0.13.0", # "packaging>=24", +# "typing-extensions>=4.15.0", # ] # /// # What is this? @@ -30,26 +31,31 @@ import os import subprocess from pathlib import Path +from typing import Any, TypeAlias import requests import tomlkit import tomlkit.items +from packaging.markers import Marker from packaging.requirements import Requirement from packaging.specifiers import SpecifierSet from packaging.utils import canonicalize_name from packaging.version import Version from tomlkit.container import OutOfOrderTableProxy +from typing_extensions import TypeIs TOO_HARD_TO_UPGRADE = frozenset({canonicalize_name(n) for n in ("torch", "transformers", "vllm")}) REPOSITORY_NAME = os.environ.get("REPOSITORY_NAME") or "NVIDIA-NeMo/Safe-Synthesizer" +TomlTable: TypeAlias = tomlkit.TOMLDocument | tomlkit.items.Table | OutOfOrderTableProxy + # --------------------------------------------------------------------------- # PEP 508 / specifier helpers # --------------------------------------------------------------------------- -def _format_requirement(name: str, spec: SpecifierSet, marker) -> str: +def _format_requirement(name: str, spec: SpecifierSet, marker: Marker | None) -> str: s = str(spec).strip() out = f"{name}{s}" if s else str(name) if marker is not None: @@ -121,50 +127,72 @@ def _bump_direct_deps_in_doc( """ changed: list[str] = [] - proj = doc.get("project") - if _is_table(proj): - main_deps = proj.get("dependencies") # type: ignore[union-attr] - if isinstance(main_deps, tomlkit.items.Array): - if _bump_array_item(main_deps, cname, floor, display_name): - changed.append("project.dependencies") - - opt = proj.get("optional-dependencies") # type: ignore[union-attr] - if _is_table(opt): - for extra_name, extra_arr in opt.items(): # type: ignore[union-attr] - if isinstance(extra_arr, tomlkit.items.Array): - if _bump_array_item(extra_arr, cname, floor, display_name): - changed.append(f"project.optional-dependencies.{extra_name}") - - dep_groups = doc.get("dependency-groups") - if _is_table(dep_groups): - for gname, garr in dep_groups.items(): # type: ignore[union-attr] - if isinstance(garr, tomlkit.items.Array): - if _bump_array_item(garr, cname, floor, display_name): - changed.append(f"dependency-groups.{gname}") + match _get_table(doc, "project"): + case None: + pass + case proj: + match _get_array(proj, "dependencies"): + case tomlkit.items.Array() as main_deps: + if _bump_array_item(main_deps, cname, floor, display_name): + changed.append("project.dependencies") + + match _get_table(proj, "optional-dependencies"): + case None: + pass + case opt: + for extra_name, extra_arr in opt.items(): + match extra_arr: + case tomlkit.items.Array() as extra_dependencies: + if _bump_array_item(extra_dependencies, cname, floor, display_name): + changed.append(f"project.optional-dependencies.{extra_name}") + + match _get_table(doc, "dependency-groups"): + case None: + pass + case dep_groups: + for gname, garr in dep_groups.items(): + match garr: + case tomlkit.items.Array() as group_dependencies: + if _bump_array_item(group_dependencies, cname, floor, display_name): + changed.append(f"dependency-groups.{gname}") return changed def _collect_direct_dep_names(doc: tomlkit.TOMLDocument) -> frozenset[str]: names: set[str] = set() - proj = doc.get("project") - if _is_table(proj): - for line in proj.get("dependencies") or []: # type: ignore[union-attr] - n = _safe_req_name(str(line)) if isinstance(line, str) else None - if n: - names.add(n) - for lines in (proj.get("optional-dependencies") or {}).values(): # type: ignore[union-attr] - for line in lines or []: + match _get_table(doc, "project"): + case None: + pass + case proj: + for line in _get_array(proj, "dependencies") or []: n = _safe_req_name(str(line)) if isinstance(line, str) else None if n: names.add(n) - dep_groups = doc.get("dependency-groups") - if _is_table(dep_groups): - for items in dep_groups.values(): # type: ignore[union-attr] - for item in items or []: - n = _safe_req_name(str(item)) if isinstance(item, str) else None - if n: - names.add(n) + + match _get_table(proj, "optional-dependencies"): + case None: + pass + case optional_dependencies: + for lines in optional_dependencies.values(): + match lines: + case tomlkit.items.Array() as dependency_lines: + for line in dependency_lines: + n = _safe_req_name(str(line)) if isinstance(line, str) else None + if n: + names.add(n) + + match _get_table(doc, "dependency-groups"): + case None: + pass + case dep_groups: + for items in dep_groups.values(): + match items: + case tomlkit.items.Array() as dependency_items: + for item in dependency_items: + n = _safe_req_name(str(item)) if isinstance(item, str) else None + if n: + names.add(n) return frozenset(names) @@ -205,16 +233,19 @@ def _write_constraints_txt(path: Path, constraint_lines: list[str]) -> None: def _replace_constraint_dependencies_array(doc: tomlkit.TOMLDocument, lines: list[str]) -> None: - if "tool" not in doc or not _is_table(doc["tool"]): - raise SystemExit("pyproject.toml: missing or invalid [tool] table") - uv = doc["tool"]["uv"] - if not _is_table(uv): - raise SystemExit("pyproject.toml: missing or invalid [tool.uv] table") - arr = tomlkit.array() - arr.multiline(True) - for line in sorted(lines): - arr.append(line) - uv["constraint-dependencies"] = arr + match _get_table(doc, "tool"): + case None: + raise SystemExit("pyproject.toml: missing or invalid [tool] table") + case tool: + match _get_table(tool, "uv"): + case None: + raise SystemExit("pyproject.toml: missing or invalid [tool.uv] table") + case uv: + arr = tomlkit.array() + arr.multiline(True) + for line in sorted(lines): + arr.append(line) + uv["constraint-dependencies"] = arr # --------------------------------------------------------------------------- @@ -222,7 +253,7 @@ def _replace_constraint_dependencies_array(doc: tomlkit.TOMLDocument, lines: lis # --------------------------------------------------------------------------- -def _collect_max_floors(deps: list[dict], too_hard: frozenset[str]) -> dict[str, str]: +def _collect_max_floors(deps: list[dict[str, Any]], too_hard: frozenset[str]) -> dict[str, str]: """Map canonical name -> highest advisory floor version across all alerts.""" out: dict[str, str] = {} for dep in deps: @@ -240,7 +271,7 @@ def _collect_max_floors(deps: list[dict], too_hard: frozenset[str]) -> dict[str, return out -def _advisory_spelling_by_canonical(deps: list[dict], cnames: frozenset[str]) -> dict[str, str]: +def _advisory_spelling_by_canonical(deps: list[dict[str, Any]], cnames: frozenset[str]) -> dict[str, str]: out: dict[str, str] = {} for dep in deps: p = dep["dependency"]["package"]["name"] @@ -257,8 +288,18 @@ def _advisory_spelling_by_canonical(deps: list[dict], cnames: frozenset[str]) -> # --------------------------------------------------------------------------- -def _is_table(x: object) -> bool: - return isinstance(x, (tomlkit.items.Table, OutOfOrderTableProxy)) +def _is_table(x: object) -> TypeIs[TomlTable]: + return isinstance(x, (tomlkit.TOMLDocument, tomlkit.items.Table, OutOfOrderTableProxy)) + + +def _get_table(table: TomlTable, key: str) -> TomlTable | None: + value = table.get(key) + return value if _is_table(value) else None + + +def _get_array(table: TomlTable, key: str) -> tomlkit.items.Array | None: + value = table.get(key) + return value if isinstance(value, tomlkit.items.Array) else None # --------------------------------------------------------------------------- @@ -278,9 +319,9 @@ def main() -> None: if not github_token: raise SystemExit("GITHUB_TOKEN is not set, required to fetch dependabot alerts") headers = {"Authorization": f"Bearer {github_token}"} - all_deps: list[dict] = [] + all_deps: list[dict[str, Any]] = [] next_url: str | None = url - params: dict | None = {"per_page": 100} + params: dict[str, int] | None = {"per_page": 100} while next_url: response = requests.get(next_url, headers=headers, params=params) response.raise_for_status() @@ -292,7 +333,7 @@ def main() -> None: json.dump(all_deps, f) with open(dependabot_file, encoding="utf-8") as f: - deps = json.load(f) + deps: list[dict[str, Any]] = json.load(f) doc = tomlkit.parse(pyproject_path.read_text(encoding="utf-8")) direct_dep_names = _collect_direct_dep_names(doc) @@ -322,14 +363,23 @@ def main() -> None: # ------------------------------------------------------------------ # Pass 2: transitive deps → constraint-dependencies # ------------------------------------------------------------------ - uv_section = doc["tool"]["uv"] if _is_table(doc.get("tool")) else None - raw = uv_section.get("constraint-dependencies") if _is_table(uv_section) else None - if raw is None: - current: list[str] = [] - elif not isinstance(raw, tomlkit.items.Array): - raise SystemExit("pyproject: [tool.uv] constraint-dependencies is not an array") - else: - current = [str(x) for x in raw if str(x).strip()] + match _get_table(doc, "tool"): + case None: + raw_constraints = None + case tool: + match _get_table(tool, "uv"): + case None: + raw_constraints = None + case uv_section: + raw_constraints = uv_section.get("constraint-dependencies") + + match raw_constraints: + case None: + current: list[str] = [] + case tomlkit.items.Array() as constraints: + current = [str(item) for item in constraints if str(item).strip()] + case _: + raise SystemExit("pyproject: [tool.uv] constraint-dependencies is not an array") updated = list(current) for cname, floor in sorted(floors.items()): diff --git a/typings/opacus/accountants/__init__.pyi b/typings/opacus/accountants/__init__.pyi index 493e97a5f..2f2a581a1 100644 --- a/typings/opacus/accountants/__init__.pyi +++ b/typings/opacus/accountants/__init__.pyi @@ -2,10 +2,12 @@ # SPDX-License-Identifier: Apache-2.0 from typing import Any +from typing_extensions import override class Accountant: def get_epsilon(self, delta: float, **kwargs: Any) -> float: ... class RDPAccountant(Accountant): def step(self, *, noise_multiplier: float, sample_rate: float) -> None: ... + @override def get_epsilon(self, delta: float, **kwargs: Any) -> float: ... diff --git a/typings/opacus/optimizers/__init__.pyi b/typings/opacus/optimizers/__init__.pyi index 1a8bfb180..6dbdb2407 100644 --- a/typings/opacus/optimizers/__init__.pyi +++ b/typings/opacus/optimizers/__init__.pyi @@ -4,6 +4,7 @@ from typing import Any from torch.optim import Optimizer +from typing_extensions import override class DPOptimizer(Optimizer): original_optimizer: Optimizer @@ -16,7 +17,9 @@ class DPOptimizer(Optimizer): expected_batch_size: int, **kwargs: Any, ) -> None: ... + @override def step(self, closure: Any = None) -> Any: ... + @override def zero_grad(self, set_to_none: bool = True) -> None: ... @property def param_groups(self) -> list[dict[str, Any]]: ... diff --git a/typings/opacus/utils/uniform_sampler.pyi b/typings/opacus/utils/uniform_sampler.pyi index d5a44612d..e0142fbe4 100644 --- a/typings/opacus/utils/uniform_sampler.pyi +++ b/typings/opacus/utils/uniform_sampler.pyi @@ -5,6 +5,7 @@ from collections.abc import Iterator import torch from torch.utils.data import Sampler +from typing_extensions import override class UniformWithReplacementSampler(Sampler[list[int]]): num_samples: int @@ -13,4 +14,5 @@ class UniformWithReplacementSampler(Sampler[list[int]]): generator: torch.Generator | None def __init__(self, *, num_samples: int, sample_rate: float, generator: object = None) -> None: ... def __len__(self) -> int: ... + @override def __iter__(self) -> Iterator[list[int]]: ...