more checks

This commit is contained in:
William Fu-Hinthorn
2026-01-08 08:18:38 -08:00
parent 46bd949f53
commit df254a7f7f
2 changed files with 80 additions and 20 deletions
+42 -20
View File
@@ -422,7 +422,9 @@ class RunnableCallable(Runnable):
else:
run_manager.on_chain_end(ret)
else:
ret = self.func(*args, **kwargs)
# Still need to set config context for get_config() to work
with set_config_context(config, None) as context:
ret = context.run(self.func, *args, **kwargs)
if self.recurse and isinstance(ret, Runnable):
return ret.invoke(input, config)
return ret
@@ -495,7 +497,13 @@ class RunnableCallable(Runnable):
else:
await run_manager.on_chain_end(ret)
else:
ret = await self.afunc(*args, **kwargs)
# Still need to set config context for get_config() to work
coro = cast(Coroutine[None, None, Any], self.afunc(*args, **kwargs))
if ASYNCIO_ACCEPTS_CONTEXT:
with set_config_context(config, None) as context:
ret = await asyncio.create_task(coro, context=context)
else:
ret = await coro
if self.recurse and isinstance(ret, Runnable):
return await ret.ainvoke(input, config)
return ret
@@ -698,12 +706,14 @@ class RunnableSeq(Runnable):
run_manager.on_chain_end(_process_outputs(self.trace_outputs, input))
return input
else:
for i, step in enumerate(self.steps):
input = (
step.invoke(input, config, **kwargs)
if i == 0
else step.invoke(input, config)
)
# Still need to set config context for get_config() to work
with set_config_context(config, None) as context:
for i, step in enumerate(self.steps):
input = (
context.run(step.invoke, input, config, **kwargs)
if i == 0
else step.invoke(input, config)
)
return input
async def ainvoke(
@@ -764,11 +774,22 @@ class RunnableSeq(Runnable):
)
return input
else:
for i, step in enumerate(self.steps):
if i == 0:
input = await step.ainvoke(input, config, **kwargs)
else:
input = await step.ainvoke(input, config)
# Still need to set config context for get_config() to work
if ASYNCIO_ACCEPTS_CONTEXT:
with set_config_context(config, None) as context:
for i, step in enumerate(self.steps):
if i == 0:
input = await asyncio.create_task(
step.ainvoke(input, config, **kwargs), context=context
)
else:
input = await step.ainvoke(input, config)
else:
for i, step in enumerate(self.steps):
if i == 0:
input = await step.ainvoke(input, config, **kwargs)
else:
input = await step.ainvoke(input, config)
return input
def stream(
@@ -837,13 +858,14 @@ class RunnableSeq(Runnable):
_process_outputs(self.trace_outputs, output)
)
else:
# No tracing - just execute the steps directly
for idx, step in enumerate(self.steps):
if idx == 0:
iterator = step.stream(input, config, **kwargs)
else:
iterator = step.transform(iterator, config)
_consume_iter(iterator)
# No tracing - still need to set config context for get_config() to work
with set_config_context(config, None) as context:
for idx, step in enumerate(self.steps):
if idx == 0:
iterator = step.stream(input, config, **kwargs)
else:
iterator = step.transform(iterator, config)
context.run(_consume_iter, iterator)
yield
async def astream(
@@ -1048,3 +1048,41 @@ async def test_astream_events_traceable_filters_traced_inputs():
if e["event"] == "on_chain_end" and e["name"] == "LangGraph"
)
assert final_event["data"]["output"] == {"value": "b_a_secret_data"}
# =============================================================================
# get_config() Compatibility Tests
# =============================================================================
def test_traceable_enabled_false_allows_get_config():
"""Test that nodes with enabled=False can still call get_config().
This is a regression test for a bug where the trace=False path in
RunnableSeq skipped set_config_context(), breaking get_config() calls.
"""
from langgraph.config import get_config
config_thread_id = None
def my_node(state: SimpleState) -> SimpleState:
nonlocal config_thread_id
config = get_config() # This should work even with enabled=False!
config_thread_id = config.get("configurable", {}).get("thread_id")
return {"value": f"got_config_{state['value']}"}
_set_traceable_config(my_node, enabled=False)
builder = StateGraph(SimpleState)
builder.add_node("my_node", my_node)
builder.add_edge("__start__", "my_node")
graph = builder.compile()
result = graph.invoke(
{"value": "test"}, {"configurable": {"thread_id": "test-thread-123"}}
)
assert result == {"value": "got_config_test"}
assert config_thread_id == "test-thread-123", (
f"Expected thread_id 'test-thread-123', got {config_thread_id}"
)