mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-19 22:25:44 +02:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4fdf5a2817 |
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user