"""Utility functions for pydantic-settings sources."""

from __future__ import annotations as _annotations

import warnings
from collections import deque
from collections.abc import Mapping, Sequence
from dataclasses import is_dataclass
from enum import Enum
from typing import Any, TypedDict, TypeVar, cast, get_args, get_origin

from pydantic import BaseModel, Json, RootModel, Secret
from pydantic._internal._utils import is_model_class
from pydantic.dataclasses import is_pydantic_dataclass
from pydantic.fields import FieldInfo
from pydantic.types import Strict
from typing_inspection import typing_objects
from typing_inspection.introspection import is_union_origin

from ..exceptions import IncompleteFieldDefinitionWarning, SettingsError
from ..utils import _lenient_issubclass
from .types import EnvNoneType


class InitState(TypedDict, total=False):
    """State shared between settings sources during a single settings resolution."""

    field_info_ids: set[int]
    """The `id()`s of the incomplete `FieldInfo` instances that were already warned about."""


def _warn_if_field_info_incomplete(field_info: FieldInfo, field_name: str, init_state: InitState) -> None:
    """Warn if the field is incomplete, i.e. its annotation contains unresolved forward references.

    An incomplete `FieldInfo` instance is unsafe to inspect — any of its attributes (annotation,
    aliases, metadata, default) may rely on the unresolved annotation, so settings sources may
    silently fail to resolve the field's value. Each instance is only warned about once per
    `init_state`, so that a field accessed by multiple sources during a single settings
    resolution doesn't emit duplicate warnings.
    """
    if getattr(field_info, '_complete', True):
        return
    warned_ids = init_state.setdefault('field_info_ids', set())
    if id(field_info) in warned_ids:
        return
    warned_ids.add(id(field_info))
    warnings.warn(
        f'Field {field_name!r} has an incomplete definition: its annotation contains an unresolved '
        'forward reference, so settings sources may fail to correctly resolve its value. '
        'Call `model_rebuild()` on the model where the field is defined, once all the referenced '
        'types are defined.',
        IncompleteFieldDefinitionWarning,
    )


def _get_env_var_key(key: str, case_sensitive: bool = False) -> str:
    return key if case_sensitive else key.lower()


def _parse_env_none_str(value: str | None, parse_none_str: str | None = None) -> str | None | EnvNoneType:
    return value if not (value == parse_none_str and parse_none_str is not None) else EnvNoneType(value)


def parse_env_vars(
    env_vars: Mapping[str, str | None],
    case_sensitive: bool = False,
    ignore_empty: bool = False,
    parse_none_str: str | None = None,
) -> Mapping[str, str | None]:
    return {
        _get_env_var_key(k, case_sensitive): _parse_env_none_str(v, parse_none_str)
        for k, v in env_vars.items()
        if not (ignore_empty and v == '')
    }


def _substitute_typevars(tp: Any, param_map: dict[Any, Any]) -> Any:
    """Substitute TypeVars in a type annotation with concrete types from param_map."""
    if isinstance(tp, TypeVar) and tp in param_map:
        return param_map[tp]
    args = get_args(tp)
    if not args:
        return tp
    new_args = tuple(_substitute_typevars(arg, param_map) for arg in args)
    if new_args == args:
        return tp
    origin = get_origin(tp)
    if origin is not None:
        try:
            return origin[new_args]
        except TypeError:
            # types.UnionType and similar are not directly subscriptable,
            # reconstruct using | operator
            import functools
            import operator

            return functools.reduce(operator.or_, new_args)
    return tp


def _resolve_type_alias(annotation: Any) -> Any:
    """Resolve a TypeAliasType to its underlying value, substituting type params if parameterized."""
    if typing_objects.is_typealiastype(annotation):
        return annotation.__value__
    origin = get_origin(annotation)
    if typing_objects.is_typealiastype(origin):
        type_params = getattr(origin, '__type_params__', ())
        type_args = get_args(annotation)
        value = origin.__value__
        if type_params and type_args:
            return _substitute_typevars(value, dict(zip(type_params, type_args)))
        return value
    return annotation


def _annotation_is_complex(annotation: Any, metadata: list[Any], init_state: InitState | None = None) -> bool:
    # If the model is a root model, the root annotation should be used to
    # evaluate the complexity.
    annotation = _resolve_type_alias(annotation)
    if annotation is not None and _lenient_issubclass(annotation, RootModel) and annotation is not RootModel:
        annotation = cast('type[RootModel[Any]]', annotation)
        root_field = annotation.model_fields['root']
        if init_state is not None:
            _warn_if_field_info_incomplete(root_field, f'{annotation.__name__}.root', init_state)
        root_annotation = root_field.annotation
        if root_annotation is not None:  # pragma: no branch
            annotation = root_annotation

    if any(isinstance(md, Json) for md in metadata):  # type: ignore[misc]
        return False

    origin = get_origin(annotation)

    # Check if annotation is of the form Annotated[type, metadata].
    if typing_objects.is_annotated(origin):
        # Return result of recursive call on inner type.
        inner, *meta = get_args(annotation)
        return _annotation_is_complex(inner, meta, init_state)

    if _lenient_issubclass(origin, Secret):
        return False

    return (
        _annotation_is_complex_inner(annotation)
        or _annotation_is_complex_inner(origin)
        or hasattr(origin, '__pydantic_core_schema__')
        or hasattr(origin, '__get_pydantic_core_schema__')
    )


def _get_field_metadata(field: FieldInfo) -> list[Any]:
    annotation = _resolve_type_alias(field.annotation)
    metadata = field.metadata
    origin = get_origin(annotation)
    if typing_objects.is_annotated(origin):
        _, *meta = get_args(annotation)
        metadata += meta
    return metadata


def _annotation_is_complex_inner(annotation: type[Any] | None) -> bool:
    if _lenient_issubclass(annotation, (str, bytes)):
        return False

    return _lenient_issubclass(
        annotation, (BaseModel, Mapping, Sequence, tuple, set, frozenset, deque)
    ) or is_dataclass(annotation)


def _union_is_complex(annotation: type[Any] | None, metadata: list[Any], init_state: InitState | None = None) -> bool:
    """Check if a union type contains any complex types."""
    for arg in get_args(annotation):
        if _annotation_is_complex(arg, metadata, init_state):
            return True
        # _annotation_is_complex doesn't handle bare Union types, so when an arg
        # is Annotated[Union[X, Y], ...], stripping Annotated yields a bare Union
        # that _annotation_is_complex can't evaluate.  Recurse into it, but only
        # if the Annotated metadata doesn't suppress complexity (e.g. Json).
        inner = _strip_annotated(arg)
        if inner is not arg:
            _, *inner_meta = get_args(arg)
            if any(isinstance(md, Json) for md in inner_meta):  # type: ignore[misc]
                continue
        if is_union_origin(get_origin(inner)) and _union_is_complex(inner, metadata, init_state):
            return True
    return False


def _union_has_strict_types(annotation: type[Any] | None) -> bool:
    """Check if a union type contains any strict-annotated types."""
    for arg in get_args(annotation):
        if typing_objects.is_annotated(get_origin(arg)):
            _, *meta = get_args(arg)
            if any(isinstance(m, Strict) for m in meta):
                return True
    return False


def _annotation_contains_types(
    annotation: type[Any] | None,
    types: tuple[Any, ...],
    is_include_origin: bool = True,
    is_strip_annotated: bool = False,
    is_instance: bool = False,
    collect: set[Any] | None = None,
) -> bool:
    """Check if a type annotation contains any of the specified types."""
    if is_strip_annotated:
        annotation = _strip_annotated(annotation)
    if is_include_origin is True:
        origin = get_origin(annotation)
        if origin in types:
            if collect is None:
                return True
            collect.add(annotation)
        if is_instance and any(isinstance(origin, type_) for type_ in types):
            if collect is None:
                return True
            collect.add(annotation)
    for type_ in get_args(annotation):
        if (
            _annotation_contains_types(
                type_,
                types,
                is_include_origin=True,
                is_strip_annotated=is_strip_annotated,
                is_instance=is_instance,
                collect=collect,
            )
            and collect is None
        ):
            return True
    if is_instance and any(isinstance(annotation, type_) for type_ in types):
        if collect is None:
            return True
        collect.add(annotation)
    if annotation in types:
        if collect is not None:
            collect.add(annotation)
        return True
    return False


def _strip_annotated(annotation: Any) -> Any:
    if typing_objects.is_annotated(get_origin(annotation)):
        return annotation.__origin__
    else:
        return annotation


def _annotation_enum_val_to_name(annotation: type[Any] | None, value: Any) -> str | None:
    for type_ in (annotation, get_origin(annotation)):
        if _lenient_issubclass(type_, Enum):
            enum_type = cast('type[Enum]', type_)
            if value in enum_type.__members__.values():
                return enum_type(value).name
    for arg in get_args(annotation):
        enum_name = _annotation_enum_val_to_name(arg, value)
        if enum_name is not None:
            return enum_name
    return None


def _annotation_enum_name_to_val(annotation: type[Any] | None, name: Any) -> Any:
    for type_ in (annotation, get_origin(annotation)):
        if _lenient_issubclass(type_, Enum):
            enum_type = cast('type[Enum]', type_)
            if name in enum_type.__members__:
                return enum_type[name]
    for arg in get_args(annotation):
        enum_val = _annotation_enum_name_to_val(arg, name)
        if enum_val is not None:
            return enum_val
    return None


def _literal_has_numeric_enum(annotation: type[Any] | None) -> bool:
    """Check if annotation is a Literal type containing numeric Enum members (IntEnum, (int, Enum), (float, Enum))."""
    if typing_objects.is_literal(get_origin(annotation)):
        return any(isinstance(arg, (int, float)) and isinstance(arg, Enum) for arg in get_args(annotation))
    # Handle Annotated wrapping, e.g. Annotated[Literal[IntEnum.member], Field(...)]
    if typing_objects.is_annotated(get_origin(annotation)):
        inner = get_args(annotation)[0]
        return _literal_has_numeric_enum(inner)
    # Handle Union/Optional wrapping, e.g. Optional[Literal[IntEnum.member]]
    if is_union_origin(get_origin(annotation)):
        return any(_literal_has_numeric_enum(arg) for arg in get_args(annotation))
    return False


def _get_model_fields(model_cls: type[Any]) -> dict[str, FieldInfo]:
    """Get fields from a pydantic model or dataclass."""

    if is_pydantic_dataclass(model_cls) and hasattr(model_cls, '__pydantic_fields__'):
        return model_cls.__pydantic_fields__
    if is_model_class(model_cls):
        return model_cls.model_fields
    raise SettingsError(f'Error: {model_cls.__name__} is not subclass of BaseModel or pydantic.dataclasses.dataclass')


def _get_alias_names(
    field_name: str,
    field_info: FieldInfo,
    alias_path_args: dict[str, int | None] | None = None,
    case_sensitive: bool = True,
    populate_by_name: bool = False,
) -> tuple[tuple[str, ...], bool]:
    """Get alias names for a field, handling alias paths and case sensitivity."""
    from pydantic import AliasChoices, AliasPath

    alias_names: list[str] = []
    is_alias_path_only: bool = True
    if not any((field_info.alias, field_info.validation_alias)):
        alias_names += [field_name]
        is_alias_path_only = False
    else:
        new_alias_paths: list[AliasPath] = []
        for alias in (field_info.alias, field_info.validation_alias):
            if alias is None:
                continue
            elif isinstance(alias, str):
                alias_names.append(alias)
                is_alias_path_only = False
            elif isinstance(alias, AliasChoices):
                for name in alias.choices:
                    if isinstance(name, str):
                        alias_names.append(name)
                        is_alias_path_only = False
                    else:
                        new_alias_paths.append(name)
            else:
                new_alias_paths.append(alias)
        for alias_path in new_alias_paths:
            name = cast(str, alias_path.path[0])
            name = name.lower() if not case_sensitive else name
            if alias_path_args is not None:
                alias_path_args[name] = (
                    alias_path.path[1] if len(alias_path.path) > 1 and isinstance(alias_path.path[1], int) else None
                )
            if not alias_names and is_alias_path_only:
                alias_names.append(name)
        if populate_by_name and field_name not in alias_names:
            alias_names.append(field_name)
            is_alias_path_only = False
    if not case_sensitive:
        alias_names = [alias_name.lower() for alias_name in alias_names]
    return tuple(dict.fromkeys(alias_names)), is_alias_path_only


def _is_function(obj: Any) -> bool:
    """Check if an object is a function."""
    from types import BuiltinFunctionType, FunctionType

    return isinstance(obj, (FunctionType, BuiltinFunctionType))


__all__ = [
    'InitState',
    '_annotation_contains_types',
    '_annotation_enum_name_to_val',
    '_annotation_enum_val_to_name',
    '_annotation_is_complex',
    '_annotation_is_complex_inner',
    '_get_alias_names',
    '_get_env_var_key',
    '_get_model_fields',
    '_is_function',
    '_literal_has_numeric_enum',
    '_parse_env_none_str',
    '_resolve_type_alias',
    '_strip_annotated',
    '_union_has_strict_types',
    '_union_is_complex',
    '_warn_if_field_info_incomplete',
    'parse_env_vars',
]
