Source code for imednet.models.base

"""Base models for the iMedNet SDK."""

from __future__ import annotations

import os
import types
from collections.abc import Callable
from datetime import datetime
from typing import (
    Any,
    Generic,
    TypeVar,
    Union,
    get_args,
    get_origin,
)

from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from typing_extensions import Self

from imednet.utils.validators import (
    is_missing_value,
    parse_bool,
    parse_datetime,
    parse_dict_or_default,
    parse_int_or_default,
    parse_list_or_default,
    parse_str_or_default,
)

_NORMALIZERS: dict[type, dict[str, Callable[[Any], Any]]] = {}


def _identity(v: Any) -> Any:
    """Return the value as-is."""
    return v


def _optional_str(v: Any) -> Any:
    """Convert value to string, returning None if missing."""
    return None if is_missing_value(v) else parse_str_or_default(v)


def _optional_int(v: Any) -> Any:
    """Convert value to int, returning None if missing."""
    return None if is_missing_value(v) else parse_int_or_default(v)


def _optional_bool(v: Any) -> Any:
    """Convert value to bool, returning None if missing."""
    return None if is_missing_value(v) else parse_bool(v)


def _optional_datetime(v: Any) -> Any:
    """Convert value to datetime, returning None if missing."""
    return None if is_missing_value(v) else parse_datetime(v)


def _extract_single_item(v: Any) -> Any:
    """Extract the first item if a list is provided where an object was expected."""
    if isinstance(v, list) and len(v) > 0 and isinstance(v[0], dict):
        import logging

        logging.getLogger(__name__).warning(
            "Structural shift detected: API returned a list where an object was expected. Coercing by extracting the first item."
        )
        return v[0]
    return v


def _get_normalizer(cls: type[BaseModel], field_name: str) -> Callable[[Any], Any]:
    """Determine the appropriate normalization function for a model field.

    Analyzes type hints to select a parser for strings, integers, booleans,
    datetimes, or nested models.
    """
    if cls in _NORMALIZERS and field_name in _NORMALIZERS[cls]:
        return _NORMALIZERS[cls][field_name]

    field = cls.model_fields[field_name]
    annotation = field.annotation
    origin = get_origin(annotation)
    optional = False

    if origin is Union or origin is getattr(types, "UnionType", type(None)):
        args = [a for a in get_args(annotation) if a is not type(None)]
        if len(args) == 1:
            annotation = args[0]
            origin = get_origin(annotation)
            optional = True

    normalizer = _identity

    if origin is list:
        normalizer = parse_list_or_default
    elif origin is dict:
        normalizer = parse_dict_or_default
    elif annotation is str:
        normalizer = _optional_str if optional else parse_str_or_default
    elif annotation is int:
        normalizer = _optional_int if optional else parse_int_or_default
    elif annotation is bool:
        normalizer = _optional_bool if optional else parse_bool
    elif annotation is datetime:
        normalizer = _optional_datetime if optional else parse_datetime
    elif isinstance(annotation, type) and issubclass(annotation, BaseModel):
        normalizer = _extract_single_item

    if cls not in _NORMALIZERS:
        _NORMALIZERS[cls] = {}
    _NORMALIZERS[cls][field_name] = normalizer
    return normalizer


_drift_reported: set[str] = set()


[docs]class ImednetBaseModel(BaseModel): """Core base model for all iMedNet API responses. Design philosophy: extra='ignore' silently drops new undocumented fields the API introduces. populate_by_name allows models to be instantiated using either pythonic snake_case names or original API camelCase names via Field aliases. str_strip_whitespace trims leading and trailing whitespace from string values. """ model_config = ConfigDict( extra="ignore", populate_by_name=True, str_strip_whitespace=True, )
[docs] @classmethod def from_json(cls, data: Any) -> Self: """Validate data coming from JSON APIs.""" if isinstance(data, list): if len(data) > 0 and isinstance(data[0], dict): import logging logging.getLogger(__name__).warning( f"Structural shift detected: API returned a list where an object ({cls.__name__}) was expected. Coercing by extracting the first item." ) data = data[0] elif len(data) == 0: import logging logging.getLogger(__name__).warning( f"Structural shift detected: API returned an empty list where an object ({cls.__name__}) was expected. Coercing to empty dict." ) data = {} try: return cls.model_validate(data) except Exception as e: import logging msg = f"Drift detected (destructive): {cls.__name__} validation failed: {e}" if msg not in _drift_reported: _drift_reported.add(msg) logging.getLogger("imednet.drift").warning(msg) raise
@model_validator(mode="before") @classmethod def _detect_drift(cls, data: Any) -> Any: """Compare incoming JSON data against the model definition to detect API drift.""" if not isinstance(data, dict): return data import logging logger = logging.getLogger("imednet.drift") defined_fields = set(cls.model_fields.keys()) for name, field in cls.model_fields.items(): # noqa: B007 if field.alias: defined_fields.add(field.alias) if hasattr(cls, "model_computed_fields"): defined_fields.update(cls.model_computed_fields.keys()) incoming_keys = set(data.keys()) unexpected_fields = incoming_keys - defined_fields if unexpected_fields: is_strict = parse_bool(os.getenv("IMEDNET_STRICT_MODE", "false")) msg = f"Drift detected (additive): {cls.__name__} received unexpected fields: {', '.join(sorted(unexpected_fields))}" if is_strict: raise ValueError(msg) if cls.model_config.get("extra") != "ignore": # noqa: SIM102 if msg not in _drift_reported: _drift_reported.add(msg) logger.warning(msg) missing_fields = [] for name, field in cls.model_fields.items(): if field.is_required(): # noqa: SIM102 if name not in incoming_keys and ( not field.alias or field.alias not in incoming_keys ): missing_fields.append(name) if missing_fields: msg = f"Drift detected (destructive): {cls.__name__} missing required fields: {', '.join(sorted(missing_fields))}" if msg not in _drift_reported: _drift_reported.add(msg) logger.warning(msg) return data @field_validator("*", check_fields=False, mode="before") def _normalise(cls, v: Any, info: Any) -> Any: """Normalize common primitive types before validation.""" if not info.field_name: return v # Bolt Optimization: Avoid function call overhead in hot path try: return _NORMALIZERS[cls][info.field_name](v) # type: ignore[index] except KeyError: return _get_normalizer(cls, info.field_name)(v) # type: ignore[arg-type]
[docs]class SortField(ImednetBaseModel): """Sorting information for a field in a paginated response.""" property: str = Field(..., description="Property to sort by") direction: str = Field(..., description="Sort direction (ASC or DESC)")
[docs]class Error(ImednetBaseModel): """Error information in an API response.""" code: str = Field("", description="Error code") message: str = Field("", description="Error message") details: dict[str, Any] = Field(default_factory=dict)
[docs]class Metadata(ImednetBaseModel): """Metadata information in an API response.""" status: str = Field("", description="Response status") method: str = Field("", description="HTTP method") path: str = Field("", description="Request path") timestamp: datetime error: Error = Field(default_factory=lambda: Error(code="", message=""))
T = TypeVar("T")
[docs]class ApiResponse(ImednetBaseModel, Generic[T]): """Generic API response model.""" metadata: Metadata pagination: Pagination | None = None data: T