mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-26 17:42:24 +02:00
Add generics tests
Signed-off-by: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com>
This commit is contained in:
@@ -24,12 +24,14 @@ from typing import (
|
||||
Any,
|
||||
Dict,
|
||||
Generator,
|
||||
Generic,
|
||||
Iterator,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Sequence,
|
||||
Tuple,
|
||||
TypeVar,
|
||||
Union,
|
||||
get_type_hints,
|
||||
)
|
||||
@@ -37,11 +39,7 @@ from typing import (
|
||||
import httpx
|
||||
import pytest
|
||||
from langchain_core.language_models import GenericFakeChatModel
|
||||
from langchain_core.runnables import (
|
||||
RunnableConfig,
|
||||
RunnableLambda,
|
||||
RunnablePassthrough,
|
||||
)
|
||||
from langchain_core.runnables import RunnableConfig, RunnableLambda, RunnablePassthrough
|
||||
from langchain_core.runnables.graph import Edge
|
||||
from langsmith import traceable
|
||||
from pytest_mock import MockerFixture
|
||||
@@ -3109,10 +3107,10 @@ def test_nested_pydantic_models(version: str) -> None:
|
||||
# Import necessary modules
|
||||
|
||||
if version == "v1":
|
||||
from pydantic.v1 import ( # type: ignore
|
||||
from pydantic.v1 import (
|
||||
BaseModel,
|
||||
ByteSize,
|
||||
Field,
|
||||
Field, # type: ignore
|
||||
SecretStr,
|
||||
confloat,
|
||||
conint,
|
||||
@@ -3120,10 +3118,10 @@ def test_nested_pydantic_models(version: str) -> None:
|
||||
constr,
|
||||
)
|
||||
else:
|
||||
from pydantic import ( # type: ignore
|
||||
from pydantic import (
|
||||
BaseModel,
|
||||
ByteSize,
|
||||
Field,
|
||||
Field, # type: ignore
|
||||
SecretStr,
|
||||
confloat,
|
||||
conint,
|
||||
@@ -3423,6 +3421,43 @@ def test_pydantic_v1_state_root_validator():
|
||||
assert res["text"] == "Hello, Validated John!"
|
||||
|
||||
|
||||
def test_pydantic_generics():
|
||||
from pydantic import BaseModel
|
||||
|
||||
class A(BaseModel):
|
||||
a: str
|
||||
|
||||
class B(BaseModel):
|
||||
b: str
|
||||
|
||||
AorB = TypeVar("AorB", A, B)
|
||||
|
||||
class C(BaseModel, Generic[AorB]):
|
||||
c: AorB
|
||||
|
||||
class State(BaseModel):
|
||||
text: str
|
||||
count: int
|
||||
c: C[A]
|
||||
|
||||
input_state = {"text": "1", "count": 0, "c": {"c": {"a": "1"}}}
|
||||
expected_input = State.model_validate(input_state)
|
||||
|
||||
def process_node(state: State):
|
||||
assert state == expected_input
|
||||
new_text = ", the type of c is " + str(type(state.c.c))
|
||||
|
||||
return {"text": state.text + new_text, "count": state.count + 1}
|
||||
|
||||
builder = StateGraph(State)
|
||||
builder.add_node("process", process_node)
|
||||
builder.add_edge(START, "process")
|
||||
g = builder.compile()
|
||||
|
||||
g.invoke(input_state)
|
||||
g.invoke(expected_input)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
|
||||
def test_in_one_fan_out_state_graph_waiting_edge_plus_regular(
|
||||
request: pytest.FixtureRequest, checkpointer_name: str
|
||||
|
||||
Reference in New Issue
Block a user