refactor: update component test method and Agent component test for be asynchronous (#5841)
This commit is contained in:
parent
050c12df35
commit
a23be65138
2 changed files with 11 additions and 10 deletions
|
|
@ -1,3 +1,4 @@
|
|||
import asyncio
|
||||
import inspect
|
||||
from typing import Any
|
||||
from unittest.mock import Mock
|
||||
|
|
@ -51,19 +52,19 @@ class ComponentTestBase:
|
|||
msg = f"{self.__class__.__name__} must implement the file_names_mapping fixture"
|
||||
raise NotImplementedError(msg)
|
||||
|
||||
def component_setup(self, component_class: type[Any], default_kwargs: dict[str, Any]) -> Component:
|
||||
async def component_setup(self, component_class: type[Any], default_kwargs: dict[str, Any]) -> Component:
|
||||
mock_vertex = Mock(spec=Vertex)
|
||||
mock_vertex.graph = Mock(spec=Graph)
|
||||
mock_vertex.graph.session_id = str(uuid4())
|
||||
mock_vertex.graph.flow_id = str(uuid4())
|
||||
source_code = inspect.getsource(component_class)
|
||||
source_code = await asyncio.to_thread(inspect.getsource, component_class)
|
||||
component_instance = component_class(_code=source_code, **default_kwargs)
|
||||
component_instance._vertex = mock_vertex
|
||||
return component_instance
|
||||
|
||||
def test_latest_version(self, component_class: type[Any], default_kwargs: dict[str, Any]) -> None:
|
||||
async def test_latest_version(self, component_class: type[Any], default_kwargs: dict[str, Any]) -> None:
|
||||
"""Test that the component works with the latest version."""
|
||||
component_instance = self.component_setup(component_class, default_kwargs)
|
||||
component_instance = await self.component_setup(component_class, default_kwargs)
|
||||
result = component_instance()
|
||||
assert result is not None, "Component returned None for the latest version."
|
||||
|
||||
|
|
|
|||
|
|
@ -22,8 +22,8 @@ class TestAgentComponent(ComponentTestBaseWithoutClient):
|
|||
def file_names_mapping(self):
|
||||
return []
|
||||
|
||||
def component_setup(self, component_class: type[Any], default_kwargs: dict[str, Any]) -> Component:
|
||||
component_instance = super().component_setup(component_class, default_kwargs)
|
||||
async def component_setup(self, component_class: type[Any], default_kwargs: dict[str, Any]) -> Component:
|
||||
component_instance = await super().component_setup(component_class, default_kwargs)
|
||||
# Mock _should_process_output method
|
||||
component_instance._should_process_output = lambda output: False # noqa: ARG005
|
||||
return component_instance
|
||||
|
|
@ -47,7 +47,7 @@ class TestAgentComponent(ComponentTestBaseWithoutClient):
|
|||
}
|
||||
|
||||
async def test_build_config_update(self, component_class, default_kwargs):
|
||||
component = self.component_setup(component_class, default_kwargs)
|
||||
component = await self.component_setup(component_class, default_kwargs)
|
||||
frontend_node = component.to_frontend_node()
|
||||
build_config = frontend_node["data"]["node"]["template"]
|
||||
# Test updating build config for OpenAI
|
||||
|
|
@ -78,9 +78,9 @@ class TestAgentComponent(ComponentTestBaseWithoutClient):
|
|||
assert all(provider in updated_config["agent_llm"]["options"] for provider in MODEL_PROVIDERS_DICT)
|
||||
assert "Anthropic" in updated_config["agent_llm"]["options"]
|
||||
assert updated_config["agent_llm"]["input_types"] == []
|
||||
assert any("sonnet" in option.lower() for option in updated_config["model_name"]["options"]), (
|
||||
f"Options: {updated_config['model_name']['options']}"
|
||||
)
|
||||
assert any(
|
||||
"sonnet" in option.lower() for option in updated_config["model_name"]["options"]
|
||||
), f"Options: {updated_config['model_name']['options']}"
|
||||
|
||||
# Test updating build config for Custom
|
||||
updated_config = await component.update_build_config(build_config, "Custom", "agent_llm")
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue