StateGraph: use state schema as input/output schema if already a pydantic model

This commit is contained in:
Nuno Campos
2024-05-03 10:59:00 -07:00
parent 45149d2dc5
commit 2ab84d1562
+15
View File
@@ -3,6 +3,7 @@ from functools import partial
from inspect import signature
from typing import Any, Optional, Sequence, Type, Union, get_type_hints
from langchain_core.pydantic_v1 import BaseModel
from langchain_core.runnables import Runnable, RunnableConfig
from langchain_core.runnables.base import RunnableLike
@@ -171,6 +172,20 @@ class StateGraph(Graph):
class CompiledStateGraph(CompiledGraph):
builder: StateGraph
def get_input_schema(
self, config: Optional[RunnableConfig] = None
) -> type[BaseModel]:
if isinstance(self.builder.schema, BaseModel):
return self.builder.schema
return super().get_input_schema(config)
def get_output_schema(self, config: RunnableConfig | None = None) -> BaseModel:
if isinstance(self.builder.schema, BaseModel):
return self.builder.schema
return super().get_output_schema(config)
def attach_node(self, key: str, node: Optional[Runnable]) -> None:
def _get_state_key(input: dict, config: RunnableConfig, *, key: str) -> Any:
if input is None: