fix: add default models to Anthropic and make sure template is updated (#5839)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
parent
64d82d407a
commit
050c12df35
19 changed files with 240 additions and 75 deletions
|
|
@ -1,6 +1,12 @@
|
|||
import inspect
|
||||
from typing import Any
|
||||
from unittest.mock import Mock
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from langflow.custom.custom_component.component import Component
|
||||
from langflow.graph.graph.base import Graph
|
||||
from langflow.graph.vertex.base import Vertex
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from tests.constants import SUPPORTED_VERSIONS
|
||||
|
|
@ -45,9 +51,20 @@ 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:
|
||||
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)
|
||||
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:
|
||||
"""Test that the component works with the latest version."""
|
||||
result = component_class(**default_kwargs)()
|
||||
component_instance = self.component_setup(component_class, default_kwargs)
|
||||
result = component_instance()
|
||||
assert result is not None, "Component returned None for the latest version."
|
||||
|
||||
def test_all_versions_have_a_file_name_defined(self, file_names_mapping: list[VersionComponentMapping]) -> None:
|
||||
|
|
|
|||
|
|
@ -1,9 +1,12 @@
|
|||
import inspect
|
||||
from typing import Any
|
||||
|
||||
from aiofile import async_open
|
||||
from fastapi import status
|
||||
from httpx import AsyncClient
|
||||
from langflow.api.v1.schemas import UpdateCustomComponentRequest
|
||||
from langflow.components.agents.agent import AgentComponent
|
||||
from langflow.custom.utils import build_custom_component_template
|
||||
|
||||
|
||||
async def test_get_version(client: AsyncClient):
|
||||
|
|
@ -46,3 +49,59 @@ async def test_update_component_outputs(client: AsyncClient, logged_in_headers:
|
|||
assert response.status_code == status.HTTP_200_OK
|
||||
output_names = [output["name"] for output in result["outputs"]]
|
||||
assert "tool_output" in output_names
|
||||
|
||||
|
||||
async def test_update_component_model_name_options(client: AsyncClient, logged_in_headers: dict):
|
||||
"""Test that model_name options are updated when selecting a provider."""
|
||||
component = AgentComponent()
|
||||
component_node, cc_instance = build_custom_component_template(
|
||||
component,
|
||||
)
|
||||
|
||||
# Initial template with OpenAI as the provider
|
||||
template = component_node["template"]
|
||||
current_model_names = template["model_name"]["options"]
|
||||
|
||||
# load the code from the file at langflow.components.agents.agent.py asynchronously
|
||||
# we are at str/backend/tests/unit/api/v1/test_endpoints.py
|
||||
# find the file by using the class AgentComponent
|
||||
agent_component_file = inspect.getsourcefile(AgentComponent)
|
||||
async with async_open(agent_component_file, encoding="utf-8") as f:
|
||||
code = await f.read()
|
||||
|
||||
# Create the request to update the component
|
||||
request = UpdateCustomComponentRequest(
|
||||
code=code,
|
||||
frontend_node=component_node,
|
||||
field="agent_llm",
|
||||
field_value="Anthropic",
|
||||
template=template,
|
||||
)
|
||||
|
||||
# Make the request to update the component
|
||||
response = await client.post("api/v1/custom_component/update", json=request.model_dump(), headers=logged_in_headers)
|
||||
result = response.json()
|
||||
|
||||
# Verify the response
|
||||
assert response.status_code == status.HTTP_200_OK, f"Response: {response.json()}"
|
||||
assert "template" in result
|
||||
assert "model_name" in result["template"]
|
||||
assert isinstance(result["template"]["model_name"]["options"], list)
|
||||
assert len(result["template"]["model_name"]["options"]) > 0, (
|
||||
f"Model names: {result['template']['model_name']['options']}"
|
||||
)
|
||||
assert current_model_names != result["template"]["model_name"]["options"], (
|
||||
f"Current model names: {current_model_names}, New model names: {result['template']['model_name']['options']}"
|
||||
)
|
||||
# Now test with Custom provider
|
||||
template["agent_llm"]["value"] = "Custom"
|
||||
request.field_value = "Custom"
|
||||
request.template = template
|
||||
|
||||
response = await client.post("api/v1/custom_component/update", json=request.model_dump(), headers=logged_in_headers)
|
||||
result = response.json()
|
||||
|
||||
# Verify that model_name is not present for Custom provider
|
||||
assert response.status_code == status.HTTP_200_OK
|
||||
assert "template" in result
|
||||
assert "model_name" not in result["template"]
|
||||
|
|
|
|||
|
|
@ -1,8 +1,99 @@
|
|||
import os
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from langflow.base.models.model_input_constants import MODEL_PROVIDERS_DICT
|
||||
from langflow.components.agents.agent import AgentComponent
|
||||
from langflow.components.tools.calculator import CalculatorToolComponent
|
||||
from langflow.custom import Component
|
||||
from langflow.utils.constants import MESSAGE_SENDER_AI, MESSAGE_SENDER_NAME_AI
|
||||
|
||||
from tests.base import ComponentTestBaseWithoutClient
|
||||
from tests.unit.mock_language_model import MockLanguageModel
|
||||
|
||||
|
||||
class TestAgentComponent(ComponentTestBaseWithoutClient):
|
||||
@pytest.fixture
|
||||
def component_class(self):
|
||||
return AgentComponent
|
||||
|
||||
@pytest.fixture
|
||||
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)
|
||||
# Mock _should_process_output method
|
||||
component_instance._should_process_output = lambda output: False # noqa: ARG005
|
||||
return component_instance
|
||||
|
||||
@pytest.fixture
|
||||
def default_kwargs(self):
|
||||
return {
|
||||
"_type": "Agent",
|
||||
"add_current_date_tool": True,
|
||||
"agent_description": "A helpful agent",
|
||||
"agent_llm": MockLanguageModel(),
|
||||
"handle_parsing_errors": True,
|
||||
"input_value": "",
|
||||
"max_iterations": 10,
|
||||
"system_prompt": "You are a helpful assistant.",
|
||||
"tools": [],
|
||||
"verbose": True,
|
||||
"session_id": str(uuid4()),
|
||||
"sender": MESSAGE_SENDER_AI,
|
||||
"sender_name": MESSAGE_SENDER_NAME_AI,
|
||||
}
|
||||
|
||||
async def test_build_config_update(self, component_class, default_kwargs):
|
||||
component = 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
|
||||
component.set(agent_llm="OpenAI")
|
||||
updated_config = await component.update_build_config(build_config, "OpenAI", "agent_llm")
|
||||
assert "agent_llm" in updated_config
|
||||
assert updated_config["agent_llm"]["value"] == "OpenAI"
|
||||
assert isinstance(updated_config["agent_llm"]["options"], list)
|
||||
assert len(updated_config["agent_llm"]["options"]) > 0
|
||||
assert all(provider in updated_config["agent_llm"]["options"] for provider in MODEL_PROVIDERS_DICT)
|
||||
assert "Custom" in updated_config["agent_llm"]["options"]
|
||||
|
||||
# Verify model_name field is populated for OpenAI
|
||||
|
||||
assert "model_name" in updated_config
|
||||
model_name_dict = updated_config["model_name"]
|
||||
assert isinstance(model_name_dict["options"], list)
|
||||
assert len(model_name_dict["options"]) > 0 # OpenAI should have available models
|
||||
assert "gpt-4o" in model_name_dict["options"]
|
||||
|
||||
# Test Anthropic
|
||||
component.set(agent_llm="Anthropic")
|
||||
updated_config = await component.update_build_config(build_config, "Anthropic", "agent_llm")
|
||||
assert "agent_llm" in updated_config
|
||||
assert updated_config["agent_llm"]["value"] == "Anthropic"
|
||||
assert isinstance(updated_config["agent_llm"]["options"], list)
|
||||
assert len(updated_config["agent_llm"]["options"]) > 0
|
||||
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']}"
|
||||
)
|
||||
|
||||
# Test updating build config for Custom
|
||||
updated_config = await component.update_build_config(build_config, "Custom", "agent_llm")
|
||||
assert "agent_llm" in updated_config
|
||||
assert updated_config["agent_llm"]["value"] == "Custom"
|
||||
assert isinstance(updated_config["agent_llm"]["options"], list)
|
||||
assert len(updated_config["agent_llm"]["options"]) > 0
|
||||
assert all(provider in updated_config["agent_llm"]["options"] for provider in MODEL_PROVIDERS_DICT)
|
||||
assert "Custom" in updated_config["agent_llm"]["options"]
|
||||
assert updated_config["agent_llm"]["input_types"] == ["LanguageModel"]
|
||||
|
||||
# Verify model_name field is cleared for Custom
|
||||
assert "model_name" not in updated_config
|
||||
|
||||
|
||||
@pytest.mark.api_key_required
|
||||
|
|
|
|||
|
|
@ -1,17 +1,21 @@
|
|||
from unittest.mock import MagicMock
|
||||
|
||||
from langchain_core.language_models import BaseLanguageModel
|
||||
from pydantic import BaseModel, Field
|
||||
from typing_extensions import override
|
||||
|
||||
|
||||
class MockLanguageModel(BaseLanguageModel):
|
||||
class MockLanguageModel(BaseLanguageModel, BaseModel):
|
||||
"""A mock language model for testing purposes."""
|
||||
|
||||
def __init__(self, response_generator=None):
|
||||
tools: list = Field(default_factory=list)
|
||||
response_generator: callable = Field(default_factory=lambda: lambda msg: f"Response for {msg}")
|
||||
|
||||
def __init__(self, response_generator=None, **kwargs):
|
||||
"""Initialize the mock model with an optional response generator function."""
|
||||
super().__init__()
|
||||
# Use object's __dict__ to bypass pydantic validation
|
||||
object.__setattr__(self, "_response_generator", response_generator or (lambda msg: f"Response for {msg}"))
|
||||
super().__init__(**kwargs)
|
||||
if response_generator:
|
||||
self.response_generator = response_generator
|
||||
|
||||
@override
|
||||
def with_config(self, *args, **kwargs):
|
||||
|
|
@ -30,7 +34,7 @@ class MockLanguageModel(BaseLanguageModel):
|
|||
for msg_list in messages:
|
||||
content = msg_list[-1]["content"] if isinstance(msg_list, list) else msg_list
|
||||
mock_response = MagicMock()
|
||||
mock_response.content = self._response_generator(content)
|
||||
mock_response.content = self.response_generator(content)
|
||||
responses.append(mock_response)
|
||||
return responses
|
||||
|
||||
|
|
@ -61,3 +65,8 @@ class MockLanguageModel(BaseLanguageModel):
|
|||
@override
|
||||
async def apredict_messages(self, *args, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
def bind_tools(self, tools):
|
||||
"""Bind tools to the model for testing."""
|
||||
self.tools = tools
|
||||
return self
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue