diff --git a/libs/langgraph/langgraph/pregel/_call.py b/libs/langgraph/langgraph/pregel/_call.py index 0cd007042..8a07a9387 100644 --- a/libs/langgraph/langgraph/pregel/_call.py +++ b/libs/langgraph/langgraph/pregel/_call.py @@ -192,6 +192,19 @@ def get_runnable_for_entrypoint(func: Callable[..., Any]) -> Runnable: return CACHE.setdefault(key, run) +def _compose_trace_inputs( + base_fn: Callable[[Any], Any], + process_fn: Callable[[Any], Any], +) -> Callable[[Any], Any]: + """Compose trace_inputs with process_inputs from traceable config.""" + + def composed(inputs: Any) -> Any: + exploded = base_fn(inputs) + return process_fn(exploded) + + return composed + + def get_runnable_for_task(func: Callable[..., Any]) -> Runnable: key = (func, True) if key in CACHE: @@ -206,10 +219,27 @@ def get_runnable_for_task(func: Callable[..., Any]) -> Runnable: else: name = str(func) - if is_async_callable(func): + # Check for traceable config early to handle __unwrapped__ + raw_config = getattr(func, "__traceable_config__", None) + traceable_config = None + if raw_config and isinstance(raw_config, dict): + traceable_config = { + "process_inputs": raw_config.get("process_inputs"), + "process_outputs": raw_config.get("process_outputs"), + "enabled": raw_config.get("enabled"), + "__unwrapped__": raw_config.get("__unwrapped__"), + } + + # Use unwrapped function if available to avoid double-tracing + # when @traceable and @task are used together + func_to_run = func + if traceable_config and traceable_config.get("__unwrapped__"): + func_to_run = traceable_config["__unwrapped__"] + + if is_async_callable(func_to_run): run = RunnableCallable( None, - func, + func_to_run, explode_args=True, name=name, trace=False, @@ -217,20 +247,43 @@ def get_runnable_for_task(func: Callable[..., Any]) -> Runnable: ) else: run = RunnableCallable( - func, - functools.wraps(func)(functools.partial(run_in_executor, None, func)), + func_to_run, + functools.wraps(func_to_run)( + functools.partial(run_in_executor, None, func_to_run) + ), explode_args=True, name=name, trace=False, recurse=False, ) + + # Build trace_inputs - compose with process_inputs if provided + # Use original func for signature to match expected args + base_trace_inputs = functools.partial( + _explode_args_trace_inputs, inspect.signature(func) + ) + if traceable_config and traceable_config.get("process_inputs"): + trace_inputs = _compose_trace_inputs( + base_trace_inputs, traceable_config["process_inputs"] + ) + else: + trace_inputs = base_trace_inputs + + # Build seq_kwargs + seq_kwargs: dict[str, Any] = { + "name": name, + "trace_inputs": trace_inputs, + } + if traceable_config: + if traceable_config.get("enabled") is False: + seq_kwargs["trace"] = False + if traceable_config.get("process_outputs"): + seq_kwargs["trace_outputs"] = traceable_config["process_outputs"] + seq = RunnableSeq( run, ChannelWrite([ChannelWriteEntry(RETURN)]), - name=name, - trace_inputs=functools.partial( - _explode_args_trace_inputs, inspect.signature(func) - ), + **seq_kwargs, ) if not _lookup_module_and_qualname(func): return seq diff --git a/libs/langgraph/tests/test_traceable_integration.py b/libs/langgraph/tests/test_traceable_integration.py index 0e93dde69..c6207ccbd 100644 --- a/libs/langgraph/tests/test_traceable_integration.py +++ b/libs/langgraph/tests/test_traceable_integration.py @@ -15,6 +15,8 @@ from typing import Any import pytest from typing_extensions import TypedDict +from langgraph._internal._runnable import ASYNCIO_ACCEPTS_CONTEXT +from langgraph.func import entrypoint, task from langgraph.graph import StateGraph from tests.fake_tracer import FakeTracer @@ -380,6 +382,10 @@ def test_traceable_config_enabled_none_honors_external_context(): assert len(node_runs) == 1 +@pytest.mark.skipif( + not ASYNCIO_ACCEPTS_CONTEXT, + reason="Requires Python 3.11+ for async context propagation", +) async def test_traceable_context_propagates_to_children(): """Traceable node should become the parent context for downstream nodes.""" tracer = FakeTracer() @@ -421,3 +427,409 @@ async def test_traceable_context_propagates_to_children(): assert child_run.dotted_order.startswith(traceable_run.dotted_order) assert child_run.parent_run_id == subgraph_run.id assert subgraph_run.parent_run_id == traceable_run.id + + +# ============================================================================= +# Functional API Tests (@entrypoint and @task) +# ============================================================================= + + +def test_entrypoint_traceable_config_process_inputs(): + """Test that @entrypoint honors __traceable_config__ with process_inputs.""" + tracer = FakeTracer() + + def filter_inputs(inputs: Any) -> dict[str, Any]: + return {"input": "[ENTRYPOINT_INPUT_REDACTED]"} + + def my_entrypoint(value: str) -> str: + return f"processed_{value}" + + _set_traceable_config(my_entrypoint, process_inputs=filter_inputs) + + workflow = entrypoint()(my_entrypoint) + + result = workflow.invoke("secret_data", {"callbacks": [tracer]}) + + assert result == "processed_secret_data" + + runs = tracer.flattened_runs() + entrypoint_runs = [r for r in runs if r.name == "my_entrypoint"] + + assert len(entrypoint_runs) == 1, ( + f"Expected 1 my_entrypoint run, got {len(entrypoint_runs)}. " + f"All runs: {[r.name for r in runs]}" + ) + + # Verify inputs were filtered + assert entrypoint_runs[0].inputs == {"input": "[ENTRYPOINT_INPUT_REDACTED]"}, ( + f"Expected inputs to be filtered, got {entrypoint_runs[0].inputs}" + ) + + +def test_entrypoint_traceable_config_process_outputs(): + """Test that @entrypoint honors __traceable_config__ with process_outputs.""" + tracer = FakeTracer() + + def filter_outputs(outputs: Any) -> Any: + return {"output": "[ENTRYPOINT_OUTPUT_REDACTED]"} + + def my_entrypoint(value: str) -> str: + return f"secret_result_{value}" + + _set_traceable_config(my_entrypoint, process_outputs=filter_outputs) + + workflow = entrypoint()(my_entrypoint) + + result = workflow.invoke("input", {"callbacks": [tracer]}) + + # Actual result should be unfiltered + assert result == "secret_result_input" + + runs = tracer.flattened_runs() + entrypoint_runs = [r for r in runs if r.name == "my_entrypoint"] + + assert len(entrypoint_runs) == 1 + + # Verify outputs were filtered in trace + assert entrypoint_runs[0].outputs == {"output": "[ENTRYPOINT_OUTPUT_REDACTED]"}, ( + f"Expected outputs to be filtered, got {entrypoint_runs[0].outputs}" + ) + + +def test_entrypoint_traceable_config_enabled_false(): + """Test that @entrypoint with enabled=False skips trace creation.""" + tracer = FakeTracer() + + def hidden_entrypoint(value: str) -> str: + return f"hidden_{value}" + + _set_traceable_config(hidden_entrypoint, enabled=False) + + workflow = entrypoint()(hidden_entrypoint) + + result = workflow.invoke("test", {"callbacks": [tracer]}) + + assert result == "hidden_test" + + runs = tracer.flattened_runs() + run_names = [r.name for r in runs] + + # The entrypoint node should NOT be traced (enabled=False) + assert "hidden_entrypoint" not in run_names, ( + f"hidden_entrypoint should not be traced but found in {run_names}" + ) + + +async def test_entrypoint_traceable_config_async(): + """Test that async @entrypoint honors __traceable_config__.""" + tracer = FakeTracer() + + def filter_inputs(inputs: Any) -> dict[str, Any]: + return {"input": "[ASYNC_ENTRYPOINT_REDACTED]"} + + async def async_entrypoint(value: str) -> str: + return f"async_{value}" + + _set_traceable_config(async_entrypoint, process_inputs=filter_inputs) + + workflow = entrypoint()(async_entrypoint) + + result = await workflow.ainvoke("secret", {"callbacks": [tracer]}) + + assert result == "async_secret" + + runs = tracer.flattened_runs() + entrypoint_runs = [r for r in runs if r.name == "async_entrypoint"] + + assert len(entrypoint_runs) == 1 + assert entrypoint_runs[0].inputs == {"input": "[ASYNC_ENTRYPOINT_REDACTED]"} + + +def test_task_traceable_config_process_inputs(): + """Test that @task honors __traceable_config__ with process_inputs.""" + tracer = FakeTracer() + + def filter_inputs(inputs: Any) -> dict[str, Any]: + return {"args": "[TASK_INPUT_REDACTED]"} + + def my_task_func(value: str) -> str: + return f"task_{value}" + + _set_traceable_config(my_task_func, process_inputs=filter_inputs) + + my_task = task(my_task_func) + + @entrypoint() + def workflow(value: str) -> str: + future = my_task(value) + return future.result() + + result = workflow.invoke("secret_data", {"callbacks": [tracer]}) + + assert result == "task_secret_data" + + runs = tracer.flattened_runs() + task_runs = [r for r in runs if r.name == "my_task_func"] + + assert len(task_runs) == 1, ( + f"Expected 1 my_task_func run, got {len(task_runs)}. " + f"All runs: {[r.name for r in runs]}" + ) + + # Verify inputs were filtered + assert task_runs[0].inputs == {"args": "[TASK_INPUT_REDACTED]"}, ( + f"Expected inputs to be filtered, got {task_runs[0].inputs}" + ) + + +def test_task_traceable_config_process_outputs(): + """Test that @task honors __traceable_config__ with process_outputs.""" + tracer = FakeTracer() + + def filter_outputs(outputs: Any) -> Any: + return {"output": "[TASK_OUTPUT_REDACTED]"} + + def my_task_func(value: str) -> str: + return f"secret_result_{value}" + + _set_traceable_config(my_task_func, process_outputs=filter_outputs) + + my_task = task(my_task_func) + + @entrypoint() + def workflow(value: str) -> str: + future = my_task(value) + return future.result() + + result = workflow.invoke("input", {"callbacks": [tracer]}) + + # Actual result should be unfiltered + assert result == "secret_result_input" + + runs = tracer.flattened_runs() + task_runs = [r for r in runs if r.name == "my_task_func"] + + assert len(task_runs) == 1 + + # Verify outputs were filtered in trace + assert task_runs[0].outputs == {"output": "[TASK_OUTPUT_REDACTED]"}, ( + f"Expected outputs to be filtered, got {task_runs[0].outputs}" + ) + + +def test_task_traceable_config_enabled_false(): + """Test that @task with enabled=False skips trace creation.""" + tracer = FakeTracer() + + def hidden_task_func(value: str) -> str: + return f"hidden_{value}" + + _set_traceable_config(hidden_task_func, enabled=False) + + hidden_task = task(hidden_task_func) + + def visible_task_func(value: str) -> str: + return f"visible_{value}" + + visible_task = task(visible_task_func) + + @entrypoint() + def workflow(value: str) -> str: + hidden_result = hidden_task(value).result() + visible_result = visible_task(hidden_result).result() + return visible_result + + result = workflow.invoke("test", {"callbacks": [tracer]}) + + assert result == "visible_hidden_test" + + runs = tracer.flattened_runs() + run_names = [r.name for r in runs] + + # visible_task_func should be traced + assert "visible_task_func" in run_names, ( + f"Expected visible_task_func in {run_names}" + ) + + # hidden_task_func should NOT be traced (enabled=False) + assert "hidden_task_func" not in run_names, ( + f"hidden_task_func should not be traced but found in {run_names}" + ) + + +@pytest.mark.skipif( + not ASYNCIO_ACCEPTS_CONTEXT, + reason="Requires Python 3.11+ for async context propagation", +) +async def test_task_traceable_config_async(): + """Test that async @task honors __traceable_config__.""" + tracer = FakeTracer() + + def filter_inputs(inputs: Any) -> dict[str, Any]: + return {"args": "[ASYNC_TASK_REDACTED]"} + + async def async_task_func(value: str) -> str: + return f"async_task_{value}" + + _set_traceable_config(async_task_func, process_inputs=filter_inputs) + + async_task = task(async_task_func) + + @entrypoint() + async def workflow(value: str) -> str: + result = await async_task(value) + return result + + result = await workflow.ainvoke("secret", {"callbacks": [tracer]}) + + assert result == "async_task_secret" + + runs = tracer.flattened_runs() + task_runs = [r for r in runs if r.name == "async_task_func"] + + assert len(task_runs) == 1 + assert task_runs[0].inputs == {"args": "[ASYNC_TASK_REDACTED]"} + + +def test_entrypoint_with_traceable_tasks_context_propagation(): + """Test that trace context propagates from entrypoint through tasks.""" + tracer = FakeTracer() + + def filter_inputs(inputs: Any) -> dict[str, Any]: + return {"input": "[MASKED]"} + + def my_task_func(value: str) -> str: + return f"task_{value}" + + my_task = task(my_task_func) + + def my_entrypoint(value: str) -> str: + future = my_task(value) + return future.result() + + _set_traceable_config(my_entrypoint, process_inputs=filter_inputs) + + workflow = entrypoint()(my_entrypoint) + + result = workflow.invoke("test", {"callbacks": [tracer]}) + + assert result == "task_test" + + runs = tracer.flattened_runs() + entrypoint_runs = [r for r in runs if r.name == "my_entrypoint"] + task_runs = [r for r in runs if r.name == "my_task_func"] + + assert len(entrypoint_runs) == 1 + assert len(task_runs) == 1 + + entrypoint_run = entrypoint_runs[0] + task_run = task_runs[0] + + # Verify same trace_id (in same trace) + assert task_run.trace_id == entrypoint_run.trace_id + + # Verify task's dotted_order starts with entrypoint's (is a descendant) + assert task_run.dotted_order.startswith(entrypoint_run.dotted_order) + + +def test_task_traceable_config_process_inputs_error_handling(): + """Test that errors in process_inputs don't leak PII for @task.""" + tracer = FakeTracer() + + def bad_filter(inputs: Any) -> Any: + raise ValueError("filter crashed!") + + def my_task_func(value: str) -> str: + return f"task_{value}" + + _set_traceable_config(my_task_func, process_inputs=bad_filter) + + my_task = task(my_task_func) + + @entrypoint() + def workflow(value: str) -> str: + future = my_task(value) + return future.result() + + # Workflow should still execute successfully + result = workflow.invoke("secret_pii_data", {"callbacks": [tracer]}) + assert result == "task_secret_pii_data" + + # Trace should show error placeholder, NOT the raw PII data + runs = tracer.flattened_runs() + task_runs = [r for r in runs if r.name == "my_task_func"] + + assert len(task_runs) == 1 + # Should NOT contain "secret_pii_data" + assert task_runs[0].inputs == {"error": ""} + + +def test_task_traceable_config_process_outputs_error_handling(): + """Test that errors in process_outputs don't leak PII for @task.""" + tracer = FakeTracer() + + def bad_filter(outputs: Any) -> Any: + raise RuntimeError("output filter crashed!") + + def my_task_func(value: str) -> str: + return "secret_output_pii" + + _set_traceable_config(my_task_func, process_outputs=bad_filter) + + my_task = task(my_task_func) + + @entrypoint() + def workflow(value: str) -> str: + future = my_task(value) + return future.result() + + # Workflow should still execute successfully + result = workflow.invoke("input", {"callbacks": [tracer]}) + assert result == "secret_output_pii" + + # Trace should show error placeholder, NOT the raw PII data + runs = tracer.flattened_runs() + task_runs = [r for r in runs if r.name == "my_task_func"] + + assert len(task_runs) == 1 + # Should NOT contain "secret_output_pii" + assert task_runs[0].outputs == {"error": ""} + + +def test_task_with_real_traceable_decorator(): + """Test @task with actual @ls.traceable decorator from langsmith.""" + import langsmith as ls + + tracer = FakeTracer() + + @ls.traceable(process_inputs=lambda inputs: {"value": "[LANGSMITH_REDACTED]"}) + def traceable_task_func(value: str) -> str: + return f"result_{value}" + + # Skip if langsmith version doesn't expose __traceable_config__ + if not hasattr(traceable_task_func, "__traceable_config__"): + pytest.skip("langsmith version doesn't expose __traceable_config__") + + my_task = task(traceable_task_func) + + @entrypoint() + def workflow(value: str) -> str: + future = my_task(value) + return future.result() + + result = workflow.invoke("secret_data", {"callbacks": [tracer]}) + + assert result == "result_secret_data" + + runs = tracer.flattened_runs() + task_runs = [r for r in runs if r.name == "traceable_task_func"] + + assert len(task_runs) == 1, ( + f"Expected 1 traceable_task_func run, got {len(task_runs)}. " + f"All runs: {[r.name for r in runs]}" + ) + + # Verify inputs were filtered by the real @traceable decorator's process_inputs + assert task_runs[0].inputs == {"value": "[LANGSMITH_REDACTED]"}, ( + f"Expected inputs to be filtered, got {task_runs[0].inputs}" + )