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:
Gabriel Luiz Freitas Almeida 2024-08-12 21:53:57 -03:00 • committed by GitHub
commit c5d9cbae49
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 654 additions and 4 deletions

View 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"

View 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"]