From 2ab84d1562f562e8f68aa2eac4bdd05c9149bd45 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Fri, 3 May 2024 10:59:00 -0700 Subject: [PATCH] StateGraph: use state schema as input/output schema if already a pydantic model --- langgraph/graph/state.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/langgraph/graph/state.py b/langgraph/graph/state.py index 85b38d585..6c42e139e 100644 --- a/langgraph/graph/state.py +++ b/langgraph/graph/state.py @@ -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: