fix(langgraph): merge instead of overwrite in ensure_config for callbacks, tags, metadata, configurable (#7926)

This commit is contained in:
Nick Hollon
2026-05-29 09:40:22 -04:00
committed by GitHub
parent f25a0f4f4c
commit 64bd4d1f13
3 changed files with 290 additions and 30 deletions
+65 -29
View File
@@ -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():
+98 -1
View File
@@ -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"
+127
View File
@@ -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