mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 13:35:09 +02:00
more checks
This commit is contained in:
@@ -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}"
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user