fix: Add async aupdate_build_config to CustomComponent (#5181)
* Add async aupdate_build_config to CustomComponent * Add test of backward compatibility
This commit is contained in:
parent
a302a946f2
commit
4f5d7d93ad
17 changed files with 141 additions and 151 deletions
|
|
@ -11,8 +11,8 @@ def component():
|
|||
return ChatOllamaComponent()
|
||||
|
||||
|
||||
@patch("httpx.Client.get")
|
||||
def test_get_model_success(mock_get, component):
|
||||
@patch("httpx.AsyncClient.get")
|
||||
async def test_get_model_success(mock_get, component):
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"models": [{"name": "model1"}, {"name": "model2"}]}
|
||||
mock_response.raise_for_status.return_value = None
|
||||
|
|
@ -20,7 +20,7 @@ def test_get_model_success(mock_get, component):
|
|||
|
||||
base_url = "http://localhost:11434"
|
||||
|
||||
model_names = component.get_model(base_url)
|
||||
model_names = await component.get_model(base_url)
|
||||
|
||||
expected_url = urljoin(base_url, "/api/tags")
|
||||
|
||||
|
|
@ -29,8 +29,8 @@ def test_get_model_success(mock_get, component):
|
|||
assert model_names == ["model1", "model2"]
|
||||
|
||||
|
||||
@patch("httpx.Client.get")
|
||||
def test_get_model_failure(mock_get, component):
|
||||
@patch("httpx.AsyncClient.get")
|
||||
async def test_get_model_failure(mock_get, component):
|
||||
# Mock the response for the HTTP GET request to raise an exception
|
||||
mock_get.side_effect = Exception("HTTP request failed")
|
||||
|
||||
|
|
@ -38,10 +38,10 @@ def test_get_model_failure(mock_get, component):
|
|||
|
||||
# Assert that the ValueError is raised when an exception occurs
|
||||
with pytest.raises(ValueError, match="Could not retrieve models"):
|
||||
component.get_model(url)
|
||||
await component.get_model(url)
|
||||
|
||||
|
||||
def test_update_build_config_mirostat_disabled(component):
|
||||
async def test_update_build_config_mirostat_disabled(component):
|
||||
build_config = {
|
||||
"mirostat_eta": {"advanced": False, "value": 0.1},
|
||||
"mirostat_tau": {"advanced": False, "value": 5},
|
||||
|
|
@ -49,7 +49,7 @@ def test_update_build_config_mirostat_disabled(component):
|
|||
field_value = "Disabled"
|
||||
field_name = "mirostat"
|
||||
|
||||
updated_config = component.update_build_config(build_config, field_value, field_name)
|
||||
updated_config = await component.aupdate_build_config(build_config, field_value, field_name)
|
||||
|
||||
assert updated_config["mirostat_eta"]["advanced"] is True
|
||||
assert updated_config["mirostat_tau"]["advanced"] is True
|
||||
|
|
@ -57,7 +57,7 @@ def test_update_build_config_mirostat_disabled(component):
|
|||
assert updated_config["mirostat_tau"]["value"] is None
|
||||
|
||||
|
||||
def test_update_build_config_mirostat_enabled(component):
|
||||
async def test_update_build_config_mirostat_enabled(component):
|
||||
build_config = {
|
||||
"mirostat_eta": {"advanced": False, "value": None},
|
||||
"mirostat_tau": {"advanced": False, "value": None},
|
||||
|
|
@ -65,7 +65,7 @@ def test_update_build_config_mirostat_enabled(component):
|
|||
field_value = "Mirostat 2.0"
|
||||
field_name = "mirostat"
|
||||
|
||||
updated_config = component.update_build_config(build_config, field_value, field_name)
|
||||
updated_config = await component.aupdate_build_config(build_config, field_value, field_name)
|
||||
|
||||
assert updated_config["mirostat_eta"]["advanced"] is False
|
||||
assert updated_config["mirostat_tau"]["advanced"] is False
|
||||
|
|
@ -73,8 +73,8 @@ def test_update_build_config_mirostat_enabled(component):
|
|||
assert updated_config["mirostat_tau"]["value"] == 10
|
||||
|
||||
|
||||
@patch("httpx.Client.get")
|
||||
def test_update_build_config_model_name(mock_get, component):
|
||||
@patch("httpx.AsyncClient.get")
|
||||
async def test_update_build_config_model_name(mock_get, component):
|
||||
# Mock the response for the HTTP GET request
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {"models": [{"name": "model1"}, {"name": "model2"}]}
|
||||
|
|
@ -88,22 +88,22 @@ def test_update_build_config_model_name(mock_get, component):
|
|||
field_value = None
|
||||
field_name = "model_name"
|
||||
|
||||
updated_config = component.update_build_config(build_config, field_value, field_name)
|
||||
updated_config = await component.aupdate_build_config(build_config, field_value, field_name)
|
||||
|
||||
assert updated_config["model_name"]["options"] == ["model1", "model2"]
|
||||
|
||||
|
||||
def test_update_build_config_keep_alive(component):
|
||||
async def test_update_build_config_keep_alive(component):
|
||||
build_config = {"keep_alive": {"value": None, "advanced": False}}
|
||||
field_value = "Keep"
|
||||
field_name = "keep_alive_flag"
|
||||
|
||||
updated_config = component.update_build_config(build_config, field_value, field_name)
|
||||
updated_config = await component.aupdate_build_config(build_config, field_value, field_name)
|
||||
assert updated_config["keep_alive"]["value"] == "-1"
|
||||
assert updated_config["keep_alive"]["advanced"] is True
|
||||
|
||||
field_value = "Immediately"
|
||||
updated_config = component.update_build_config(build_config, field_value, field_name)
|
||||
updated_config = await component.aupdate_build_config(build_config, field_value, field_name)
|
||||
assert updated_config["keep_alive"]["value"] == "0"
|
||||
assert updated_config["keep_alive"]["advanced"] is True
|
||||
|
||||
|
|
|
|||
|
|
@ -1,10 +1,15 @@
|
|||
from typing import Any
|
||||
|
||||
import pytest
|
||||
from langflow.components.agents import AgentComponent
|
||||
from langflow.components.crewai import CrewAIAgentComponent, SequentialTaskComponent
|
||||
from langflow.components.custom_component import CustomComponent
|
||||
from langflow.components.inputs import ChatInput
|
||||
from langflow.components.models import OpenAIModelComponent
|
||||
from langflow.components.outputs import ChatOutput
|
||||
from langflow.schema import dotdict
|
||||
from langflow.template import Output
|
||||
from typing_extensions import override
|
||||
|
||||
|
||||
def test_set_invalid_output():
|
||||
|
|
@ -58,3 +63,21 @@ def test_set_required_inputs_various_components():
|
|||
assert _assert_all_outputs_have_different_required_inputs(chatoutput.outputs)
|
||||
assert _assert_all_outputs_have_different_required_inputs(task.outputs)
|
||||
assert _assert_all_outputs_have_different_required_inputs(agent.outputs)
|
||||
|
||||
|
||||
async def test_update_build_config_backward_compatibility():
|
||||
class TestComponent(CustomComponent):
|
||||
@override
|
||||
def update_build_config(
|
||||
self,
|
||||
build_config: dotdict,
|
||||
field_value: Any,
|
||||
field_name: str | None = None,
|
||||
):
|
||||
build_config["foo"] = "bar"
|
||||
return build_config
|
||||
|
||||
component = TestComponent()
|
||||
build_config = dotdict()
|
||||
build_config = await component.aupdate_build_config(build_config, "", "")
|
||||
assert build_config["foo"] == "bar"
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from datetime import datetime
|
||||
from unittest.mock import patch
|
||||
from uuid import UUID, uuid4
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
from langflow.services.database.models.variable.model import VariableUpdate
|
||||
|
|
@ -9,7 +9,7 @@ from langflow.services.settings.constants import VARIABLES_TO_GET_FROM_ENVIRONME
|
|||
from langflow.services.variable.constants import CREDENTIAL_TYPE, GENERIC_TYPE
|
||||
from langflow.services.variable.service import DatabaseVariableService
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
from sqlmodel import Session, SQLModel
|
||||
from sqlmodel import SQLModel
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
|
||||
|
|
@ -28,16 +28,6 @@ async def session():
|
|||
yield session
|
||||
|
||||
|
||||
def _get_variable(
|
||||
session: Session,
|
||||
service,
|
||||
user_id: UUID | str,
|
||||
name: str,
|
||||
field: str,
|
||||
):
|
||||
return service.get_variable(user_id, name, field, session=session)
|
||||
|
||||
|
||||
async def test_initialize_user_variables__create_and_update(service, session: AsyncSession):
|
||||
user_id = uuid4()
|
||||
field = ""
|
||||
|
|
@ -53,7 +43,7 @@ async def test_initialize_user_variables__create_and_update(service, session: As
|
|||
|
||||
variables = await service.list_variables(user_id, session=session)
|
||||
for name in variables:
|
||||
value = await session.run_sync(_get_variable, service, user_id, name, field)
|
||||
value = await service.get_variable(user_id, name, field, session=session)
|
||||
assert value == env_vars[name]
|
||||
|
||||
assert all(i in variables for i in good_vars)
|
||||
|
|
@ -80,7 +70,7 @@ async def test_get_variable(service, session: AsyncSession):
|
|||
field = ""
|
||||
await service.create_variable(user_id, name, value, session=session)
|
||||
|
||||
result = await session.run_sync(_get_variable, service, user_id, name, field)
|
||||
result = await service.get_variable(user_id, name, field, session=session)
|
||||
|
||||
assert result == value
|
||||
|
||||
|
|
@ -91,7 +81,7 @@ async def test_get_variable__valueerror(service, session: AsyncSession):
|
|||
field = ""
|
||||
|
||||
with pytest.raises(ValueError, match=f"{name} variable not found."):
|
||||
await session.run_sync(_get_variable, service, user_id, name, field)
|
||||
await service.get_variable(user_id, name, field, session=session)
|
||||
|
||||
|
||||
async def test_get_variable__typeerror(service, session: AsyncSession):
|
||||
|
|
@ -103,7 +93,7 @@ async def test_get_variable__typeerror(service, session: AsyncSession):
|
|||
await service.create_variable(user_id, name, value, type_=type_, session=session)
|
||||
|
||||
with pytest.raises(TypeError) as exc:
|
||||
await session.run_sync(_get_variable, service, user_id, name, field)
|
||||
await service.get_variable(user_id, name, field, session=session)
|
||||
|
||||
assert name in str(exc.value)
|
||||
assert "purpose is to prevent the exposure of value" in str(exc.value)
|
||||
|
|
@ -136,9 +126,9 @@ async def test_update_variable(service, session: AsyncSession):
|
|||
field = ""
|
||||
await service.create_variable(user_id, name, old_value, session=session)
|
||||
|
||||
old_recovered = await session.run_sync(_get_variable, service, user_id, name, field)
|
||||
old_recovered = await service.get_variable(user_id, name, field, session=session)
|
||||
result = await service.update_variable(user_id, name, new_value, session=session)
|
||||
new_recovered = await session.run_sync(_get_variable, service, user_id, name, field)
|
||||
new_recovered = await service.get_variable(user_id, name, field, session=session)
|
||||
|
||||
assert old_value == old_recovered
|
||||
assert new_value == new_recovered
|
||||
|
|
@ -197,10 +187,10 @@ async def test_delete_variable(service, session: AsyncSession):
|
|||
field = ""
|
||||
|
||||
await service.create_variable(user_id, name, value, session=session)
|
||||
recovered = await session.run_sync(_get_variable, service, user_id, name, field)
|
||||
recovered = await service.get_variable(user_id, name, field, session=session)
|
||||
await service.delete_variable(user_id, name, session=session)
|
||||
with pytest.raises(ValueError, match=f"{name} variable not found."):
|
||||
await session.run_sync(_get_variable, service, user_id, name, field)
|
||||
await service.get_variable(user_id, name, field, session=session)
|
||||
|
||||
assert recovered == value
|
||||
|
||||
|
|
@ -220,10 +210,10 @@ async def test_delete_variable_by_id(service, session: AsyncSession):
|
|||
field = "field"
|
||||
|
||||
saved = await service.create_variable(user_id, name, value, session=session)
|
||||
recovered = await session.run_sync(_get_variable, service, user_id, name, field)
|
||||
recovered = await service.get_variable(user_id, name, field, session=session)
|
||||
await service.delete_variable_by_id(user_id, saved.id, session=session)
|
||||
with pytest.raises(ValueError, match=f"{name} variable not found."):
|
||||
await session.run_sync(_get_variable, service, user_id, name, field)
|
||||
await service.get_variable(user_id, name, field, session=session)
|
||||
|
||||
assert recovered == value
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue