feat: add dynamic state model creation and update (#3271)
* feat: add initial implementation of dynamic state model creation and output getter in graph state module * feat: implement _reset_all_output_values method to initialize component outputs in custom_component class * feat: add state model management with lazy initialization and dynamic instance getter in custom_component class * feat: Refactor Component class to use public method get_output_by_method Refactor the Component class in the custom_component module to change the visibility of the method `_get_output_by_method` to public by renaming it to `get_output_by_method`. This change improves the accessibility and clarity of the method for external use. * feat: add output setter utility to manage output values in state model properties * feat: implement validation for methods' classes in output getter/setter utilities in state model to ensure proper structure * feat: add state model creation from graph in state_model.py * feat: enhance Graph class with lazy loading for state model creation from graph * feat: add unit tests for state model creation and validation in test_state_model.py * feat: add unit tests for state model creation and validation in test_state_model.py * feat: add functional test for graph state update and validation in test_graph_state_model.py * fix: update _instance_getter function to accept a parameter in component.py for state model instance retrieval * refactor: rename test to clarify purpose in test_state_model.py for functional state update validation * chore: import Finish constant in test_graph_state_model.py for improved clarity and usage in state model tests * refactor: add optional validation in output getter/setter methods for improved method integrity in state model handling * refactor: enhance state model creation with optional validation and error handling for output methods in model.py * refactor: serialize and deserialize GraphStateModel in test_graph_state_model.py * refactor: improve error message and add verbose mode for graph start in test_state_model.py * refactor: remove verbose flag from graph.start in TestCreateStateModel for consistency in test_state_model.py * refactor: disable validation when creating GraphStateModel in state_model.py for improved flexibility * refactor: add validation documentation for method attributes in model.py to enhance code clarity and usability * refactor: expand docstring for build_output_getter in model.py to clarify usage and validation details * refactor: add detailed docstring for build_output_setter in model.py to improve clarity on functionality and usage scenarios * refactor: add comprehensive docstring for create_state_model in model.py to clarify functionality and usage examples * refactor: enhance docstring for create_state_model_from_graph in state_model.py to clarify functionality and provide examples * test: add JSON schema validation in graph state model tests for improved structure and correctness verification * refactor: Improve graph_state_model.json_schema unit test readability and structure.
This commit is contained in:
parent
2ffd723065
commit
c5d9cbae49
7 changed files with 654 additions and 4 deletions
139
src/backend/tests/unit/graph/graph/state/test_state_model.py
Normal file
139
src/backend/tests/unit/graph/graph/state/test_state_model.py
Normal file
|
|
@ -0,0 +1,139 @@
|
|||
import pytest
|
||||
from pydantic import Field
|
||||
|
||||
from langflow.components.inputs import ChatInput
|
||||
from langflow.components.outputs.ChatOutput import ChatOutput
|
||||
from langflow.graph.graph.base import Graph
|
||||
from langflow.graph.graph.constants import Finish
|
||||
from langflow.graph.state.model import create_state_model
|
||||
from langflow.template.field.base import UNDEFINED
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client():
|
||||
pass
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def chat_input_component():
|
||||
return ChatInput()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def chat_output_component():
|
||||
return ChatOutput()
|
||||
|
||||
|
||||
class TestCreateStateModel:
|
||||
# Successfully create a model with valid method return type annotations
|
||||
|
||||
def test_create_model_with_valid_return_type_annotations(self, chat_input_component):
|
||||
StateModel = create_state_model(method_one=chat_input_component.message_response)
|
||||
|
||||
state_instance = StateModel()
|
||||
assert state_instance.method_one is UNDEFINED
|
||||
chat_input_component.set_output_value("message", "test")
|
||||
assert state_instance.method_one == "test"
|
||||
|
||||
def test_create_model_and_assign_values_fails(self, chat_input_component):
|
||||
StateModel = create_state_model(method_one=chat_input_component.message_response)
|
||||
|
||||
state_instance = StateModel()
|
||||
state_instance.method_one = "test"
|
||||
assert state_instance.method_one == "test"
|
||||
|
||||
def test_create_with_multiple_components(self, chat_input_component, chat_output_component):
|
||||
NewStateModel = create_state_model(
|
||||
model_name="NewStateModel",
|
||||
first_method=chat_input_component.message_response,
|
||||
second_method=chat_output_component.message_response,
|
||||
)
|
||||
state_instance = NewStateModel()
|
||||
assert state_instance.first_method is UNDEFINED
|
||||
assert state_instance.second_method is UNDEFINED
|
||||
state_instance.first_method = "test"
|
||||
state_instance.second_method = 123
|
||||
assert state_instance.first_method == "test"
|
||||
assert state_instance.second_method == 123
|
||||
|
||||
def test_create_with_pydantic_field(self, chat_input_component):
|
||||
StateModel = create_state_model(method_one=chat_input_component.message_response, my_attribute=Field(None))
|
||||
|
||||
state_instance = StateModel()
|
||||
state_instance.method_one = "test"
|
||||
state_instance.my_attribute = "test"
|
||||
assert state_instance.method_one == "test"
|
||||
assert state_instance.my_attribute == "test"
|
||||
# my_attribute should be of type Any
|
||||
state_instance.my_attribute = 123
|
||||
assert state_instance.my_attribute == 123
|
||||
|
||||
# Creates a model with fields based on provided keyword arguments
|
||||
def test_create_model_with_fields_from_kwargs(self):
|
||||
StateModel = create_state_model(field_one=(str, "default"), field_two=(int, 123))
|
||||
state_instance = StateModel()
|
||||
assert state_instance.field_one == "default"
|
||||
assert state_instance.field_two == 123
|
||||
|
||||
# Raises ValueError for invalid field type in tuple-based definitions
|
||||
def test_raise_valueerror_for_invalid_field_type_in_tuple(self):
|
||||
with pytest.raises(ValueError, match="Invalid type for field invalid_field"):
|
||||
create_state_model(invalid_field=("not_a_type", "default"))
|
||||
|
||||
# Raises ValueError for unsupported value types in keyword arguments
|
||||
def test_raise_valueerror_for_unsupported_value_types(self):
|
||||
with pytest.raises(ValueError, match="Invalid value type <class 'int'> for field invalid_field"):
|
||||
create_state_model(invalid_field=123)
|
||||
|
||||
# Handles empty keyword arguments gracefully
|
||||
def test_handle_empty_kwargs_gracefully(self):
|
||||
StateModel = create_state_model()
|
||||
state_instance = StateModel()
|
||||
assert state_instance is not None
|
||||
|
||||
# Ensures model name defaults to "State" if not provided
|
||||
def test_default_model_name_to_state(self):
|
||||
StateModel = create_state_model()
|
||||
assert StateModel.__name__ == "State"
|
||||
OtherNameModel = create_state_model(model_name="OtherName")
|
||||
assert OtherNameModel.__name__ == "OtherName"
|
||||
|
||||
# Validates that callable values are properly type-annotated
|
||||
|
||||
def test_create_model_with_invalid_callable(self):
|
||||
class MockComponent:
|
||||
def method_one(self) -> str:
|
||||
return "test"
|
||||
|
||||
def method_two(self) -> int:
|
||||
return 123
|
||||
|
||||
mock_component = MockComponent()
|
||||
with pytest.raises(ValueError, match="get_output_by_method"):
|
||||
create_state_model(method_one=mock_component.method_one, method_two=mock_component.method_two)
|
||||
|
||||
def test_graph_functional_start_state_update(self):
|
||||
chat_input = ChatInput(_id="chat_input")
|
||||
chat_output = ChatOutput(input_value="test", _id="chat_output")
|
||||
chat_output.set(sender_name=chat_input.message_response)
|
||||
ChatStateModel = create_state_model(model_name="ChatState", message=chat_output.message_response)
|
||||
chat_state_model = ChatStateModel()
|
||||
assert chat_state_model.__class__.__name__ == "ChatState"
|
||||
assert chat_state_model.message is UNDEFINED
|
||||
|
||||
graph = Graph(chat_input, chat_output)
|
||||
graph.prepare()
|
||||
# Now iterate through the graph
|
||||
# and check that the graph is running
|
||||
# correctly
|
||||
ids = ["chat_input", "chat_output"]
|
||||
results = []
|
||||
for result in graph.start():
|
||||
results.append(result)
|
||||
|
||||
assert len(results) == 3
|
||||
assert all(result.vertex.id in ids for result in results if hasattr(result, "vertex"))
|
||||
assert results[-1] == Finish()
|
||||
|
||||
assert chat_state_model.__class__.__name__ == "ChatState"
|
||||
assert chat_state_model.message.get_text() == "test"
|
||||
173
src/backend/tests/unit/graph/graph/test_graph_state_model.py
Normal file
173
src/backend/tests/unit/graph/graph/test_graph_state_model.py
Normal file
|
|
@ -0,0 +1,173 @@
|
|||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
||||
from langflow.components.helpers.Memory import MemoryComponent
|
||||
from langflow.components.inputs.ChatInput import ChatInput
|
||||
from langflow.components.models.OpenAIModel import OpenAIModelComponent
|
||||
from langflow.components.outputs.ChatOutput import ChatOutput
|
||||
from langflow.components.prompts.Prompt import PromptComponent
|
||||
from langflow.graph import Graph
|
||||
from langflow.graph.graph.constants import Finish
|
||||
from langflow.graph.graph.state_model import create_state_model_from_graph
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def client():
|
||||
pass
|
||||
|
||||
|
||||
def test_graph_state_model():
|
||||
session_id = "test_session_id"
|
||||
template = """{context}
|
||||
|
||||
User: {user_message}
|
||||
AI: """
|
||||
memory_component = MemoryComponent(_id="chat_memory")
|
||||
memory_component.set(session_id=session_id)
|
||||
chat_input = ChatInput(_id="chat_input")
|
||||
prompt_component = PromptComponent(_id="prompt")
|
||||
prompt_component.set(
|
||||
template=template, user_message=chat_input.message_response, context=memory_component.retrieve_messages_as_text
|
||||
)
|
||||
openai_component = OpenAIModelComponent(_id="openai")
|
||||
openai_component.set(
|
||||
input_value=prompt_component.build_prompt, max_tokens=100, temperature=0.1, api_key="test_api_key"
|
||||
)
|
||||
openai_component.get_output("text_output").value = "Mock response"
|
||||
|
||||
chat_output = ChatOutput(_id="chat_output")
|
||||
chat_output.set(input_value=openai_component.text_response)
|
||||
|
||||
graph = Graph(chat_input, chat_output)
|
||||
|
||||
GraphStateModel = create_state_model_from_graph(graph)
|
||||
assert GraphStateModel.__name__ == "GraphStateModel"
|
||||
assert list(GraphStateModel.model_computed_fields.keys()) == [
|
||||
"chat_input",
|
||||
"chat_output",
|
||||
"openai",
|
||||
"prompt",
|
||||
"chat_memory",
|
||||
]
|
||||
|
||||
|
||||
def test_graph_functional_start_graph_state_update():
|
||||
chat_input = ChatInput(_id="chat_input")
|
||||
chat_input.set(input_value="Test Sender Name")
|
||||
chat_output = ChatOutput(input_value="test", _id="chat_output")
|
||||
chat_output.set(sender_name=chat_input.message_response)
|
||||
|
||||
graph = Graph(chat_input, chat_output)
|
||||
graph.prepare()
|
||||
# Now iterate through the graph
|
||||
# and check that the graph is running
|
||||
# correctly
|
||||
GraphStateModel = create_state_model_from_graph(graph)
|
||||
graph_state_model = GraphStateModel()
|
||||
ids = ["chat_input", "chat_output"]
|
||||
results = []
|
||||
for result in graph.start():
|
||||
results.append(result)
|
||||
|
||||
assert len(results) == 3
|
||||
assert all(result.vertex.id in ids for result in results if hasattr(result, "vertex"))
|
||||
assert results[-1] == Finish()
|
||||
|
||||
assert graph_state_model.__class__.__name__ == "GraphStateModel"
|
||||
assert graph_state_model.chat_input.message.get_text() == "Test Sender Name"
|
||||
assert graph_state_model.chat_output.message.get_text() == "test"
|
||||
|
||||
|
||||
def test_graph_state_model_serialization():
|
||||
chat_input = ChatInput(_id="chat_input")
|
||||
chat_input.set(input_value="Test Sender Name")
|
||||
chat_output = ChatOutput(input_value="test", _id="chat_output")
|
||||
chat_output.set(sender_name=chat_input.message_response)
|
||||
|
||||
graph = Graph(chat_input, chat_output)
|
||||
graph.prepare()
|
||||
# Now iterate through the graph
|
||||
# and check that the graph is running
|
||||
# correctly
|
||||
GraphStateModel = create_state_model_from_graph(graph)
|
||||
graph_state_model = GraphStateModel()
|
||||
ids = ["chat_input", "chat_output"]
|
||||
results = []
|
||||
for result in graph.start():
|
||||
results.append(result)
|
||||
|
||||
assert len(results) == 3
|
||||
assert all(result.vertex.id in ids for result in results if hasattr(result, "vertex"))
|
||||
assert results[-1] == Finish()
|
||||
|
||||
assert graph_state_model.__class__.__name__ == "GraphStateModel"
|
||||
assert graph_state_model.chat_input.message.get_text() == "Test Sender Name"
|
||||
assert graph_state_model.chat_output.message.get_text() == "test"
|
||||
|
||||
serialized_state_model = graph_state_model.model_dump()
|
||||
assert serialized_state_model["chat_input"]["message"]["text"] == "Test Sender Name"
|
||||
|
||||
|
||||
def test_graph_state_model_json_schema():
|
||||
chat_input = ChatInput(_id="chat_input")
|
||||
chat_input.set(input_value="Test Sender Name")
|
||||
chat_output = ChatOutput(input_value="test", _id="chat_output")
|
||||
chat_output.set(sender_name=chat_input.message_response)
|
||||
|
||||
graph = Graph(chat_input, chat_output)
|
||||
graph.prepare()
|
||||
|
||||
GraphStateModel = create_state_model_from_graph(graph)
|
||||
graph_state_model: BaseModel = GraphStateModel()
|
||||
json_schema = graph_state_model.model_json_schema(mode="serialization")
|
||||
|
||||
# Test main schema structure
|
||||
assert json_schema["title"] == "GraphStateModel"
|
||||
assert json_schema["type"] == "object"
|
||||
assert set(json_schema["required"]) == {"chat_input", "chat_output"}
|
||||
|
||||
# Test chat_input and chat_output properties
|
||||
for prop in ["chat_input", "chat_output"]:
|
||||
assert prop in json_schema["properties"]
|
||||
assert json_schema["properties"][prop]["allOf"][0]["$ref"].startswith("#/$defs/")
|
||||
assert json_schema["properties"][prop]["readOnly"] is True
|
||||
|
||||
# Test $defs
|
||||
assert set(json_schema["$defs"].keys()) == {"ChatInputStateModel", "ChatOutputStateModel", "Image", "Message"}
|
||||
|
||||
# Test ChatInputStateModel and ChatOutputStateModel
|
||||
for model in ["ChatInputStateModel", "ChatOutputStateModel"]:
|
||||
assert json_schema["$defs"][model]["type"] == "object"
|
||||
assert json_schema["$defs"][model]["title"] == model
|
||||
assert "message" in json_schema["$defs"][model]["properties"]
|
||||
assert json_schema["$defs"][model]["properties"]["message"]["allOf"][0]["$ref"] == "#/$defs/Message"
|
||||
assert json_schema["$defs"][model]["properties"]["message"]["readOnly"] is True
|
||||
assert json_schema["$defs"][model]["required"] == ["message"]
|
||||
|
||||
# Test Message model
|
||||
message_props = json_schema["$defs"]["Message"]["properties"]
|
||||
assert set(message_props.keys()) == {
|
||||
"text_key",
|
||||
"data",
|
||||
"default_value",
|
||||
"text",
|
||||
"sender",
|
||||
"sender_name",
|
||||
"files",
|
||||
"session_id",
|
||||
"timestamp",
|
||||
"flow_id",
|
||||
}
|
||||
assert message_props["text_key"]["type"] == "string"
|
||||
assert message_props["data"]["type"] == "object"
|
||||
assert "anyOf" in message_props["default_value"]
|
||||
assert "anyOf" in message_props["files"]
|
||||
assert message_props["timestamp"]["type"] == "string"
|
||||
|
||||
# Test Image model
|
||||
image_props = json_schema["$defs"]["Image"]["properties"]
|
||||
assert set(image_props.keys()) == {"path", "url"}
|
||||
for prop in ["path", "url"]:
|
||||
assert "anyOf" in image_props[prop]
|
||||
assert {"type": "string"} in image_props[prop]["anyOf"]
|
||||
assert {"type": "null"} in image_props[prop]["anyOf"]
|
||||
Loading…
Add table
Add a link
Reference in a new issue