from __future__ import annotations

from collections import ChainMap
from collections.abc import Mapping, Sequence
from os import getenv
from typing import Any, cast

from langchain_core.callbacks import (
    AsyncCallbackManager,
    BaseCallbackManager,
    CallbackManager,
    Callbacks,
)
from langchain_core.runnables import RunnableConfig
from langchain_core.runnables.config import (
    CONFIG_KEYS,
    COPIABLE_KEYS,
    var_child_runnable_config,
)
from langgraph.checkpoint.base import CheckpointMetadata

from langgraph._internal._constants import (
    CONF,
    CONFIG_KEY_CHECKPOINT_ID,
    CONFIG_KEY_CHECKPOINT_MAP,
    CONFIG_KEY_CHECKPOINT_NS,
    NS_END,
    NS_SEP,
)

DEFAULT_RECURSION_LIMIT = int(getenv("LANGGRAPH_DEFAULT_RECURSION_LIMIT", "10007"))
DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT = int(
    getenv("LANGGRAPH_DELTA_MAX_SUPERSTEPS_SINCE_SNAPSHOT", "5000")
)


def recast_checkpoint_ns(ns: str) -> str:
    """Remove task IDs from checkpoint namespace.

    Args:
        ns: The checkpoint namespace with task IDs.

    Returns:
        str: The checkpoint namespace without task IDs.
    """
    return NS_SEP.join(
        part.split(NS_END)[0] for part in ns.split(NS_SEP) if not part.isdigit()
    )


def patch_configurable(
    config: RunnableConfig | None, patch: dict[str, Any]
) -> RunnableConfig:
    if config is None:
        return {CONF: patch}
    elif CONF not in config:
        return {**config, CONF: patch}
    else:
        return {**config, CONF: {**config[CONF], **patch}}


def patch_checkpoint_map(
    config: RunnableConfig | None, metadata: CheckpointMetadata | None
) -> RunnableConfig:
    if config is None:
        return config
    elif parents := (metadata.get("parents") if metadata else None):
        conf = config[CONF]
        return patch_configurable(
            config,
            {
                CONFIG_KEY_CHECKPOINT_MAP: {
                    **parents,
                    conf[CONFIG_KEY_CHECKPOINT_NS]: conf[CONFIG_KEY_CHECKPOINT_ID],
                },
            },
        )
    else:
        return config


def _copy_mapping_value(value: Any) -> Any:
    return dict(value) if isinstance(value, Mapping) else value


def _merge_metadata(
    base: Mapping[str, Any] | None, new: Mapping[str, Any] | None
) -> dict[str, Any]:
    """Merge metadata without mutating inputs.

    Top-level keys merge with newer values winning. `lc_versions` is the only
    mapping-valued key that merges one level deeper, so independent LangChain
    packages can contribute package versions without changing generic metadata
    semantics. Mapping values are copied one level, so deeper nested objects
    remain shared. `None` inputs are treated as empty metadata.

    Mirrors `langchain_core.runnables.config._merge_metadata_dicts` so configs
    that pass through LangGraph's merge keep the same `lc_versions` semantics;
    keep the two in sync. Unlike lc-core, this copies mapping values one level
    (intentionally more defensive — do not "simplify" back to a shared ref).
    """
    merged = {key: _copy_mapping_value(value) for key, value in (base or {}).items()}
    for key, value in (new or {}).items():
        if (
            key == "lc_versions"
            and isinstance(merged.get(key), Mapping)
            and isinstance(value, Mapping)
        ):
            merged[key] = {
                **cast(Mapping[str, Any], merged[key]),
                **value,
            }
        else:
            merged[key] = _copy_mapping_value(value)
    return merged


def _merge_callbacks(base: Callbacks, new: Callbacks) -> Callbacks:
    """Merge two callbacks values (None / list / BaseCallbackManager).

    Six cases total (3 base types x 2 non-None new types).
    """
    if new is None:
        return base
    if base is None:
        return new.copy() if isinstance(new, (list, BaseCallbackManager)) else new
    if isinstance(new, list):
        if isinstance(base, list):
            return base + new
        if isinstance(base, BaseCallbackManager):
            mngr = base.copy()
            for cb in new:
                mngr.add_handler(cb, inherit=True)
            return mngr
    elif isinstance(new, BaseCallbackManager):
        if isinstance(base, list):
            mngr = new.copy()
            for cb in base:
                mngr.add_handler(cb, inherit=True)
            return mngr
        if isinstance(base, BaseCallbackManager):
            return base.merge(new)
    raise NotImplementedError(f"Unsupported callback types: {type(base)}, {type(new)}")


def merge_configs(*configs: RunnableConfig | None) -> RunnableConfig:
    """Merge multiple configs into one.

    Args:
        *configs: The configs to merge.

    Returns:
        RunnableConfig: The merged config.
    """
    base: RunnableConfig = {}
    # Even though the keys aren't literals, this is correct
    # because both dicts are the same type
    for config in configs:
        if config is None:
            continue
        for key, value in config.items():
            if not value:
                continue
            if key == "metadata":
                base[key] = _merge_metadata(
                    cast(Mapping[str, Any] | None, base.get(key)),
                    cast(Mapping[str, Any], value),
                )
            elif key == "tags":
                if base_value := base.get(key):
                    base[key] = [*base_value, *value]  # type: ignore
                else:
                    base[key] = value  # type: ignore[literal-required]
            elif key == CONF:
                if base_value := base.get(key):
                    base[key] = {**base_value, **value}  # type: ignore[dict-item]
                else:
                    base[key] = value
            elif key == "callbacks":
                base["callbacks"] = _merge_callbacks(
                    base.get("callbacks"), cast(Callbacks, value)
                )
            elif key == "recursion_limit":
                if config["recursion_limit"] != DEFAULT_RECURSION_LIMIT:
                    base["recursion_limit"] = config["recursion_limit"]
            else:
                base[key] = config[key]  # type: ignore[literal-required]
    if CONF not in base:
        base[CONF] = {}
    return base


def patch_config(
    config: RunnableConfig | None,
    *,
    callbacks: Callbacks = None,
    recursion_limit: int | None = None,
    max_concurrency: int | None = None,
    run_name: str | None = None,
    configurable: dict[str, Any] | None = None,
) -> RunnableConfig:
    """Patch a config with new values.

    Args:
        config: The config to patch.
        callbacks: The callbacks to set.
        recursion_limit: The recursion limit to set.
        max_concurrency: The max number of concurrent steps to run, which also applies to parallelized steps.
        run_name: The run name to set.
        configurable: The configurable to set.

    Returns:
        RunnableConfig: The patched config.
    """
    config = config.copy() if config is not None else {}
    if callbacks is not None:
        # If we're replacing callbacks, we need to unset run_name
        # As that should apply only to the same run as the original callbacks
        config["callbacks"] = callbacks
        if "run_name" in config:
            del config["run_name"]
        if "run_id" in config:
            del config["run_id"]
    if recursion_limit is not None:
        config["recursion_limit"] = recursion_limit
    if max_concurrency is not None:
        config["max_concurrency"] = max_concurrency
    if run_name is not None:
        config["run_name"] = run_name
    if configurable is not None:
        config[CONF] = {**config.get(CONF, {}), **configurable}
    return config


def get_callback_manager_for_config(
    config: RunnableConfig, tags: Sequence[str] | None = None
) -> CallbackManager:
    """Get a callback manager for a config.

    Args:
        config: The config.

    Returns:
        CallbackManager: The callback manager.
    """
    from langchain_core.callbacks.manager import CallbackManager

    # merge tags
    all_tags = config.get("tags")
    if all_tags is not None and tags is not None:
        all_tags = [*all_tags, *tags]
    elif tags is not None:
        all_tags = list(tags)
    # use existing callbacks if they exist
    if (callbacks := config.get("callbacks")) and isinstance(
        callbacks, CallbackManager
    ):
        if all_tags:
            callbacks.add_tags(all_tags)
        if metadata := config.get("metadata"):
            callbacks.add_metadata(metadata)
        manager = callbacks
    else:
        # otherwise create a new manager
        manager = CallbackManager.configure(
            inheritable_callbacks=config.get("callbacks"),
            inheritable_tags=all_tags,
            inheritable_metadata=config.get("metadata"),
            langsmith_inheritable_metadata=_get_tracing_metadata_defaults(config),
        )
    return manager


def get_async_callback_manager_for_config(
    config: RunnableConfig,
    tags: Sequence[str] | None = None,
) -> AsyncCallbackManager:
    """Get an async callback manager for a config.

    Args:
        config: The config.

    Returns:
        AsyncCallbackManager: The async callback manager.
    """
    from langchain_core.callbacks.manager import AsyncCallbackManager

    # merge tags
    all_tags = config.get("tags")
    if all_tags is not None and tags is not None:
        all_tags = [*all_tags, *tags]
    elif tags is not None:
        all_tags = list(tags)
    # use existing callbacks if they exist
    if (callbacks := config.get("callbacks")) and isinstance(
        callbacks, AsyncCallbackManager
    ):
        if all_tags:
            callbacks.add_tags(all_tags)
        if metadata := config.get("metadata"):
            callbacks.add_metadata(metadata)
        manager = callbacks
    else:
        # otherwise create a new manager
        manager = AsyncCallbackManager.configure(
            inheritable_callbacks=config.get("callbacks"),
            inheritable_tags=all_tags,
            inheritable_metadata=config.get("metadata"),
            langsmith_inheritable_metadata=_get_tracing_metadata_defaults(config),
        )
    return manager


def _is_not_empty(value: Any) -> bool:
    if isinstance(value, (list, tuple, dict)):
        return len(value) > 0
    else:
        return value is not None


def ensure_config(*configs: RunnableConfig | None) -> RunnableConfig:
    """Return a config with all keys, merging any provided configs.

    Args:
        *configs: Configs to merge before ensuring defaults.

    Returns:
        RunnableConfig: The merged and ensured config.
    """
    empty = RunnableConfig(
        tags=[],
        metadata=ChainMap(),
        callbacks=None,
        recursion_limit=DEFAULT_RECURSION_LIMIT,
        configurable={},
    )
    if var_config := var_child_runnable_config.get():
        empty.update(
            {
                k: v.copy() if k in COPIABLE_KEYS else v  # type: ignore[attr-defined]
                for k, v in var_config.items()
                if _is_not_empty(v)
            },
        )
    for config in configs:
        if config is None:
            continue
        for k, v in config.items():
            if _is_not_empty(v) and k in CONFIG_KEYS:
                if k == CONF:
                    # Shallow-merge configurable dicts across configs so values
                    # bound via with_config(...) (e.g. ls_agent_type) are
                    # preserved when later configs (e.g. invoke-time) only
                    # specify a subset of keys like thread_id.
                    existing = empty.get(k)
                    empty[k] = (
                        {**cast(dict, existing), **cast(dict, v)}
                        if existing
                        else cast(dict, v).copy()
                    )
                elif k == "callbacks":
                    empty["callbacks"] = _merge_callbacks(
                        empty.get("callbacks"), cast(Callbacks, v)
                    )
                elif k == "metadata":
                    # Matches merge_configs: top-level metadata keys merge, and
                    # only `lc_versions` merges one level deeper.
                    empty["metadata"] = _merge_metadata(
                        cast(Mapping[str, Any] | None, empty.get("metadata")),
                        cast(Mapping[str, Any], v),
                    )
                elif k == "tags":
                    # Concatenate tags across configs so values bound via
                    # with_config(...) are preserved when later configs
                    # supply additional tags. Matches merge_configs.
                    existing_tags: list[str] | None = empty.get("tags")
                    empty["tags"] = (
                        [*existing_tags, *cast(list, v)]
                        if existing_tags
                        else list(cast(list, v))
                    )
                else:
                    empty[k] = v  # type: ignore[literal-required]
        for k, v in config.items():
            if _is_not_empty(v) and k not in CONFIG_KEYS:
                empty[CONF][k] = v

    configurable = empty.get("configurable")
    metadata = empty.get("metadata")
    if configurable and metadata is not None:
        for key in _PROPAGATE_TO_METADATA:
            if key in metadata:
                continue
            value = configurable.get(key)
            if value:
                metadata[key] = value
    return empty


_OMIT = ("key", "token", "secret", "password", "auth")


def _exclude_as_metadata(key: str, value: Any) -> bool:
    key_lower = key.casefold()
    return (
        key.startswith("__")
        or not isinstance(value, (str, int, float, bool))
        or any(substr in key_lower for substr in _OMIT)
    )


def _get_tracing_metadata_defaults(
    config: RunnableConfig,
) -> dict[str, Any] | None:
    """Get tracer-only metadata defaults from configurable values."""
    configurable = config.get("configurable")
    if not configurable:
        return None
    metadata: dict[str, Any] = {}
    for key, value in configurable.items():
        if _exclude_as_metadata(key, value):
            continue
        metadata[key] = value
    return metadata or None


_PROPAGATE_TO_METADATA = frozenset(
    (
        "thread_id",
        "checkpoint_id",
        "checkpoint_ns",
        "task_id",
        "run_id",
        "assistant_id",
        "graph_id",
    )
)


def filter_to_user_tags(tags: Sequence[str] | None) -> list[str] | None:
    """Drop langgraph's internal `seq:step:*` bookkeeping tags.

    `seq:step:N` tags are added internally to mark sequence steps; everything
    else (user-supplied tags and any other framework tags) is kept. Returns the
    surviving tags, or `None` if none remain. Shared by the `messages` and
    `tasks` stream handlers so both surface the same tag set on their metadata.
    """
    if not tags:
        return None
    filtered = [t for t in tags if not t.startswith("seq:step")]
    return filtered or None
