From 64bd4d1f13f1343244dbdfe3af8275020f6c3aa3 Mon Sep 17 00:00:00 2001 From: Nick Hollon Date: Fri, 29 May 2026 09:40:22 -0400 Subject: [PATCH] fix(langgraph): merge instead of overwrite in ensure_config for callbacks, tags, metadata, configurable (#7926) --- libs/langgraph/langgraph/_internal/_config.py | 94 +++++++++---- libs/langgraph/tests/test_config_async.py | 99 +++++++++++++- libs/langgraph/tests/test_utils.py | 127 ++++++++++++++++++ 3 files changed, 290 insertions(+), 30 deletions(-) diff --git a/libs/langgraph/langgraph/_internal/_config.py b/libs/langgraph/langgraph/_internal/_config.py index 1bb9128da..45092659a 100644 --- a/libs/langgraph/langgraph/_internal/_config.py +++ b/libs/langgraph/langgraph/_internal/_config.py @@ -79,6 +79,34 @@ def patch_checkpoint_map( return config +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. @@ -113,34 +141,9 @@ def merge_configs(*configs: RunnableConfig | None) -> RunnableConfig: else: base[key] = value elif key == "callbacks": - base_callbacks = base.get("callbacks") - # callbacks can be either None, list[handler] or manager - # so merging two callbacks values has 6 cases - if isinstance(value, list): - if base_callbacks is None: - base["callbacks"] = value.copy() - elif isinstance(base_callbacks, list): - base["callbacks"] = base_callbacks + value - else: - # base_callbacks is a manager - mngr = base_callbacks.copy() - for callback in value: - mngr.add_handler(callback, inherit=True) - base["callbacks"] = mngr - elif isinstance(value, BaseCallbackManager): - # value is a manager - if base_callbacks is None: - base["callbacks"] = value.copy() - elif isinstance(base_callbacks, list): - mngr = value.copy() - for callback in base_callbacks: - mngr.add_handler(callback, inherit=True) - base["callbacks"] = mngr - else: - # base_callbacks is also a manager - base["callbacks"] = base_callbacks.merge(value) - else: - raise NotImplementedError + 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"] @@ -309,7 +312,40 @@ def ensure_config(*configs: RunnableConfig | None) -> RunnableConfig: for k, v in config.items(): if _is_not_empty(v) and k in CONFIG_KEYS: if k == CONF: - empty[k] = cast(dict, v).copy() + # 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": + # Shallow-merge metadata dicts across configs so values + # bound via with_config(...) (e.g. user_id) are preserved + # when later configs supply other metadata keys. + existing = empty.get("metadata") + empty["metadata"] = ( + {**cast(dict, existing), **cast(dict, v)} + if existing + else cast(dict, v).copy() + ) + 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(): diff --git a/libs/langgraph/tests/test_config_async.py b/libs/langgraph/tests/test_config_async.py index a65a5f2af..6ad51c5b2 100644 --- a/libs/langgraph/tests/test_config_async.py +++ b/libs/langgraph/tests/test_config_async.py @@ -1,7 +1,8 @@ import pytest -from langchain_core.callbacks import AsyncCallbackManager +from langchain_core.callbacks import AsyncCallbackManager, BaseCallbackHandler from langgraph._internal._config import get_async_callback_manager_for_config +from langgraph.graph import StateGraph pytestmark = pytest.mark.anyio @@ -17,3 +18,99 @@ def test_new_async_manager_merges_tags_with_config() -> None: config = {"callbacks": None, "tags": ["a"]} manager = get_async_callback_manager_for_config(config, tags=["b"]) assert manager.inheritable_tags == ["a", "b"] + + +class _TrackingCallback(BaseCallbackHandler): + def __init__(self) -> None: + self.called = False + + def on_chain_start(self, *args, **kwargs) -> None: # noqa: ANN002, ANN003 + self.called = True + + +async def test_with_config_callbacks_preserved_in_astream_events() -> None: + """A callback bound via .with_config(...) must survive when + astream_events injects its own internal callback handler. + + Pre-fix: ensure_config overwrites the callbacks key, dropping the + bound handler. Post-fix: the handler list is merged. + """ + builder = StateGraph(dict) + builder.add_node("node", lambda state: state) + builder.add_edge("__start__", "node") + cb = _TrackingCallback() + graph = builder.compile().with_config({"callbacks": [cb]}) + async for _ in graph.astream_events({}, version="v2"): + pass + assert cb.called, "user-bound callback was dropped by ensure_config overwrite" + + +async def test_with_config_configurable_preserved_on_invoke() -> None: + """A configurable key bound via .with_config(...) must survive when + invoke-time config supplies a different configurable key. + + Pre-fix: ensure_config overwrites the entire configurable dict. + Post-fix: the two dicts are shallow-merged per key. + """ + builder = StateGraph(dict) + captured: dict = {} + + def node(state, config): # noqa: ANN001 + captured.update(config.get("configurable") or {}) + return state + + builder.add_node("node", node) + builder.add_edge("__start__", "node") + graph = builder.compile().with_config({"configurable": {"ls_agent_type": "root"}}) + await graph.ainvoke({}, {"configurable": {"thread_id": "T1"}}) + assert captured.get("ls_agent_type") == "root", ( + "bound configurable key was dropped by ensure_config overwrite" + ) + assert captured.get("thread_id") == "T1", "invoke-time key not present" + + +async def test_with_config_metadata_preserved_on_invoke() -> None: + """A metadata key bound via .with_config(...) must survive when + invoke-time config supplies a different metadata key. + + Pre-fix: ensure_config overwrites the entire metadata dict. + Post-fix: the two dicts are shallow-merged per key. + """ + builder = StateGraph(dict) + captured: dict = {} + + def node(state, config): # noqa: ANN001 + captured.update(config.get("metadata") or {}) + return state + + builder.add_node("node", node) + builder.add_edge("__start__", "node") + graph = builder.compile().with_config({"metadata": {"user_id": "U1"}}) + await graph.ainvoke({}, {"metadata": {"correlation_id": "C1"}}) + assert captured.get("user_id") == "U1", ( + "bound metadata key was dropped by ensure_config overwrite" + ) + assert captured.get("correlation_id") == "C1", "invoke-time key not present" + + +async def test_with_config_tags_preserved_on_invoke() -> None: + """Tags bound via .with_config(...) must survive when invoke-time + config supplies its own tags. + + Pre-fix: ensure_config overwrites the entire tags list. + Post-fix: tags are concatenated (matching merge_configs behavior; + no deduplication, no sorting). + """ + builder = StateGraph(dict) + captured: list = [] + + def node(state, config): # noqa: ANN001 + captured.extend(config.get("tags") or []) + return state + + builder.add_node("node", node) + builder.add_edge("__start__", "node") + graph = builder.compile().with_config({"tags": ["bound"]}) + await graph.ainvoke({}, {"tags": ["invoke"]}) + assert "bound" in captured, "bound tag was dropped by ensure_config overwrite" + assert "invoke" in captured, "invoke-time tag not present" diff --git a/libs/langgraph/tests/test_utils.py b/libs/langgraph/tests/test_utils.py index c9e6168ff..65735683b 100644 --- a/libs/langgraph/tests/test_utils.py +++ b/libs/langgraph/tests/test_utils.py @@ -15,12 +15,14 @@ from unittest.mock import MagicMock, patch import langsmith import pytest +from langchain_core.callbacks import BaseCallbackHandler, CallbackManager from langchain_core.runnables import RunnableConfig from langchain_core.tracers import LangChainTracer from typing_extensions import NotRequired, Required, TypedDict from langgraph._internal._config import ( _is_not_empty, + _merge_callbacks, ensure_config, get_callback_manager_for_config, ) @@ -427,3 +429,128 @@ def test_callback_manager_copies_configurable_ids_to_tracing_metadata() -> None: "thread_id": "th-123", "user_id": "uid-1", } + + +class _TrackingCB(BaseCallbackHandler): + """Minimal callback handler used only as a sentinel for merge tests.""" + + def __init__(self, tag: str) -> None: + self.tag = tag + + def __eq__(self, other: object) -> bool: + return isinstance(other, _TrackingCB) and self.tag == other.tag + + def __hash__(self) -> int: + return hash(self.tag) + + +def test_merge_callbacks_none_base_list_new() -> None: + cb = _TrackingCB("a") + merged = _merge_callbacks(None, [cb]) + assert merged == [cb] + + +def test_merge_callbacks_list_base_list_new() -> None: + a, b = _TrackingCB("a"), _TrackingCB("b") + merged = _merge_callbacks([a], [b]) + assert merged == [a, b] + + +def test_merge_callbacks_list_base_manager_new() -> None: + a = _TrackingCB("a") + mgr = CallbackManager(handlers=[_TrackingCB("b")]) + merged = _merge_callbacks([a], mgr) + assert isinstance(merged, CallbackManager) + assert _TrackingCB("a") in merged.handlers + assert _TrackingCB("b") in merged.handlers + + +def test_merge_callbacks_manager_base_list_new() -> None: + mgr = CallbackManager(handlers=[_TrackingCB("a")]) + b = _TrackingCB("b") + merged = _merge_callbacks(mgr, [b]) + assert isinstance(merged, CallbackManager) + assert _TrackingCB("a") in merged.handlers + assert _TrackingCB("b") in merged.handlers + + +def test_merge_callbacks_manager_base_manager_new() -> None: + mgr_a = CallbackManager(handlers=[_TrackingCB("a")]) + mgr_b = CallbackManager(handlers=[_TrackingCB("b")]) + merged = _merge_callbacks(mgr_a, mgr_b) + assert isinstance(merged, CallbackManager) + assert _TrackingCB("a") in merged.handlers + assert _TrackingCB("b") in merged.handlers + + +def test_merge_callbacks_none_base_none_new() -> None: + merged = _merge_callbacks(None, None) + assert merged is None + + +def test_ensure_config_merges_configurable_across_configs() -> None: + a = {"configurable": {"ls_agent_type": "root"}} + b = {"configurable": {"thread_id": "T1"}} + merged = ensure_config(a, b) + assert merged["configurable"]["ls_agent_type"] == "root" + assert merged["configurable"]["thread_id"] == "T1" + + +def test_ensure_config_configurable_later_wins_per_key() -> None: + a = {"configurable": {"shared": "from_a", "only_a": "A"}} + b = {"configurable": {"shared": "from_b", "only_b": "B"}} + merged = ensure_config(a, b) + assert merged["configurable"]["shared"] == "from_b" # later wins per key + assert merged["configurable"]["only_a"] == "A" + assert merged["configurable"]["only_b"] == "B" + + +def test_ensure_config_merges_metadata_across_configs() -> None: + a = {"metadata": {"user_id": "U1"}} + b = {"metadata": {"correlation_id": "C1"}} + merged = ensure_config(a, b) + assert merged["metadata"]["user_id"] == "U1" + assert merged["metadata"]["correlation_id"] == "C1" + + +def test_ensure_config_metadata_later_wins_per_key() -> None: + a = {"metadata": {"shared": "from_a"}} + b = {"metadata": {"shared": "from_b"}} + merged = ensure_config(a, b) + assert merged["metadata"]["shared"] == "from_b" + + +def test_ensure_config_merges_tags_across_configs() -> None: + a = {"tags": ["alpha"]} + b = {"tags": ["beta"]} + merged = ensure_config(a, b) + assert merged["tags"] == ["alpha", "beta"] + + +def test_ensure_config_tags_concat_preserves_order_and_duplicates() -> None: + # Plain concat (matches merge_configs in this file — no dedup, no sort). + a = {"tags": ["shared", "alpha"]} + b = {"tags": ["shared", "beta"]} + merged = ensure_config(a, b) + assert merged["tags"] == ["shared", "alpha", "shared", "beta"] + + +def test_ensure_config_merges_callbacks_across_configs() -> None: + a_cb = _TrackingCB("a") + b_cb = _TrackingCB("b") + merged = ensure_config({"callbacks": [a_cb]}, {"callbacks": [b_cb]}) + assert merged["callbacks"] == [a_cb, b_cb] + + +def test_ensure_config_none_inputs_ignored() -> None: + # mixed with None should not raise + merged = ensure_config(None, {"tags": ["t"]}, None) + assert merged["tags"] == ["t"] + + +def test_ensure_config_empty_inputs() -> None: + # everything empty -> defaults + merged = ensure_config() + assert merged["tags"] == [] + assert merged["configurable"] == {} + assert merged["callbacks"] is None