Compare commits

...
Author SHA1 Message Date
William Fu-Hinthorn 4fdf5a2817 chore: Exclude extra args 2025-08-11 07:31:43 -07:00
2 changed files with 88 additions and 16 deletions
+23 -7
View File
@@ -9,8 +9,8 @@ import weakref
from collections import defaultdict, deque
from collections.abc import AsyncIterator, Iterator, Mapping, Sequence
from dataclasses import is_dataclass
from functools import partial
from inspect import isclass
from functools import lru_cache, partial
from inspect import Parameter, isclass, signature
from typing import Any, Callable, Generic, Union, cast, get_type_hints
from uuid import UUID, uuid5
@@ -3232,6 +3232,15 @@ def _output(
yield payload
@lru_cache(maxsize=128)
def _get_schema_fields(context_schema: type) -> set[str] | None:
# Must be dataclass right now.
sig = signature(context_schema)
if any(p.kind == Parameter.VAR_KEYWORD for p in sig.parameters.values()):
return None
return set(get_type_hints(context_schema).keys())
def _coerce_context(
context_schema: type[ContextT] | None, context: Any
) -> ContextT | None:
@@ -3253,10 +3262,17 @@ def _coerce_context(
if context_schema is None:
return context
schema_is_class = issubclass(context_schema, BaseModel) or is_dataclass(
context_schema
)
if isinstance(context, dict) and schema_is_class:
return context_schema(**context) # type: ignore[misc]
if isinstance(context, dict):
if is_typeddict(context_schema):
return cast(ContextT, context)
if is_dataclass(context_schema):
fields = _get_schema_fields(context_schema)
if fields is None:
return context_schema(**context)
return context_schema(**{k: v for k, v in context.items() if k in fields})
if issubclass(context_schema, BaseModel):
return context_schema(**context)
return cast(ContextT, context)
+65 -9
View File
@@ -52,6 +52,14 @@ def test_context_runtime() -> None:
{"message": "hello world"}, context=Context(api_key="sk_123456")
)
assert result == {"message": "api key: sk_123456"}
result = compiled.invoke(
{"message": "hello world"}, context={"api_key": "sk_1234567"}
)
assert result == {"message": "api key: sk_1234567"}
result = compiled.invoke(
{"message": "hello world"}, context={"api_key": "sk_1234568", "foo": "bar"}
)
assert result == {"message": "api key: sk_1234568"}
def test_override_runtime() -> None:
@@ -134,7 +142,8 @@ def test_context_coercion_dataclass() -> None:
# Test dict coercion with all fields
result = compiled.invoke(
{"message": "test"}, context={"api_key": "sk_test", "timeout": 60}
{"message": "test"},
context={"api_key": "sk_test", "timeout": 60, "extra_arg": "extra_value"},
)
assert result == {"message": "api_key: sk_test, timeout: 60"}
@@ -149,6 +158,53 @@ def test_context_coercion_dataclass() -> None:
assert result == {"message": "api_key: sk_test3, timeout: 90"}
def test_context_coercion_dataclass_custom_init() -> None:
"""Test that dict context is coerced to dataclass."""
@dataclass(init=False)
class Context:
api_key: str
timeout: int = 30
def __init__(self, api_key: str, timeout: int = 30, **kwargs: Any):
self.api_key = api_key
self.timeout = timeout
setattr(self, "extra_arg", kwargs.get("extra_arg"))
class State(TypedDict):
message: str
def node_with_context(state: State, runtime: Runtime[Context]) -> dict[str, Any]:
return {
"message": f"api_key: {runtime.context.api_key}, timeout: {runtime.context.timeout}, extra_arg: {runtime.context.extra_arg}"
}
graph = StateGraph(state_schema=State, context_schema=Context)
graph.add_node("node", node_with_context)
graph.add_edge(START, "node")
graph.add_edge("node", END)
compiled = graph.compile()
# Test dict coercion with all fields
result = compiled.invoke(
{"message": "test"},
context={"api_key": "sk_test", "timeout": 60, "extra_arg": "extra_value"},
)
assert result == {
"message": "api_key: sk_test, timeout: 60, extra_arg: extra_value"
}
# Test dict coercion with default field
result = compiled.invoke({"message": "test"}, context={"api_key": "sk_test2"})
assert result == {"message": "api_key: sk_test2, timeout: 30, extra_arg: None"}
# Test with actual dataclass instance (should still work)
result = compiled.invoke(
{"message": "test"}, context=Context(api_key="sk_test3", timeout=90)
)
assert result == {"message": "api_key: sk_test3, timeout: 90, extra_arg: None"}
def test_context_coercion_pydantic() -> None:
"""Test that dict context is coerced to Pydantic model."""
@@ -174,7 +230,12 @@ def test_context_coercion_pydantic() -> None:
# Test dict coercion with all fields
result = compiled.invoke(
{"message": "test"},
context={"api_key": "sk_test", "timeout": 60, "tags": ["prod", "v2"]},
context={
"api_key": "sk_test",
"timeout": 60,
"tags": ["prod", "v2"],
"extra_arg": "extra_value",
},
)
assert result == {"message": "api_key: sk_test, timeout: 60, tags: ['prod', 'v2']"}
@@ -214,7 +275,8 @@ def test_context_coercion_typeddict() -> None:
# Test dict passes through for TypedDict
result = compiled.invoke(
{"message": "test"}, context={"api_key": "sk_test", "timeout": 60}
{"message": "test"},
context={"api_key": "sk_test", "timeout": 60, "extra_arg": "extra_value"},
)
assert result == {"message": "api_key: sk_test, timeout: 60"}
@@ -271,12 +333,6 @@ def test_context_coercion_errors() -> None:
with pytest.raises(TypeError):
compiled.invoke({"message": "test"}, context={"timeout": 60})
# Test invalid dict keys
with pytest.raises(TypeError):
compiled.invoke(
{"message": "test"}, context={"api_key": "test", "invalid_field": "value"}
)
@pytest.mark.anyio
async def test_context_coercion_async() -> None: