mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-09 11:17:53 +02:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4fdf5a2817 |
@@ -9,8 +9,8 @@ import weakref
|
|||||||
from collections import defaultdict, deque
|
from collections import defaultdict, deque
|
||||||
from collections.abc import AsyncIterator, Iterator, Mapping, Sequence
|
from collections.abc import AsyncIterator, Iterator, Mapping, Sequence
|
||||||
from dataclasses import is_dataclass
|
from dataclasses import is_dataclass
|
||||||
from functools import partial
|
from functools import lru_cache, partial
|
||||||
from inspect import isclass
|
from inspect import Parameter, isclass, signature
|
||||||
from typing import Any, Callable, Generic, Union, cast, get_type_hints
|
from typing import Any, Callable, Generic, Union, cast, get_type_hints
|
||||||
from uuid import UUID, uuid5
|
from uuid import UUID, uuid5
|
||||||
|
|
||||||
@@ -3232,6 +3232,15 @@ def _output(
|
|||||||
yield payload
|
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(
|
def _coerce_context(
|
||||||
context_schema: type[ContextT] | None, context: Any
|
context_schema: type[ContextT] | None, context: Any
|
||||||
) -> ContextT | None:
|
) -> ContextT | None:
|
||||||
@@ -3253,10 +3262,17 @@ def _coerce_context(
|
|||||||
if context_schema is None:
|
if context_schema is None:
|
||||||
return context
|
return context
|
||||||
|
|
||||||
schema_is_class = issubclass(context_schema, BaseModel) or is_dataclass(
|
if isinstance(context, dict):
|
||||||
context_schema
|
if is_typeddict(context_schema):
|
||||||
)
|
return cast(ContextT, context)
|
||||||
if isinstance(context, dict) and schema_is_class:
|
|
||||||
return context_schema(**context) # type: ignore[misc]
|
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)
|
return cast(ContextT, context)
|
||||||
|
|||||||
@@ -52,6 +52,14 @@ def test_context_runtime() -> None:
|
|||||||
{"message": "hello world"}, context=Context(api_key="sk_123456")
|
{"message": "hello world"}, context=Context(api_key="sk_123456")
|
||||||
)
|
)
|
||||||
assert result == {"message": "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:
|
def test_override_runtime() -> None:
|
||||||
@@ -134,7 +142,8 @@ def test_context_coercion_dataclass() -> None:
|
|||||||
|
|
||||||
# Test dict coercion with all fields
|
# Test dict coercion with all fields
|
||||||
result = compiled.invoke(
|
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"}
|
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"}
|
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:
|
def test_context_coercion_pydantic() -> None:
|
||||||
"""Test that dict context is coerced to Pydantic model."""
|
"""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
|
# Test dict coercion with all fields
|
||||||
result = compiled.invoke(
|
result = compiled.invoke(
|
||||||
{"message": "test"},
|
{"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']"}
|
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
|
# Test dict passes through for TypedDict
|
||||||
result = compiled.invoke(
|
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"}
|
assert result == {"message": "api_key: sk_test, timeout: 60"}
|
||||||
|
|
||||||
@@ -271,12 +333,6 @@ def test_context_coercion_errors() -> None:
|
|||||||
with pytest.raises(TypeError):
|
with pytest.raises(TypeError):
|
||||||
compiled.invoke({"message": "test"}, context={"timeout": 60})
|
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
|
@pytest.mark.anyio
|
||||||
async def test_context_coercion_async() -> None:
|
async def test_context_coercion_async() -> None:
|
||||||
|
|||||||
Reference in New Issue
Block a user