ref: Auto-fix ruff rules in tests (#4154)
This commit is contained in:
parent
51b3909d60
commit
45c8f98692
80 changed files with 359 additions and 456 deletions
|
|
@ -1,18 +1,14 @@
|
|||
import os
|
||||
|
||||
from astrapy.db import AstraDB
|
||||
import pytest
|
||||
|
||||
from astrapy.db import AstraDB
|
||||
from langchain_core.documents import Document
|
||||
from langflow.components.embeddings import OpenAIEmbeddingsComponent
|
||||
from langflow.components.vectorstores import AstraVectorStoreComponent
|
||||
from tests.api_keys import get_astradb_application_token, get_astradb_api_endpoint, get_openai_api_key
|
||||
from tests.integration.components.mock_components import TextToData
|
||||
from tests.integration.utils import ComponentInputHandle
|
||||
from langchain_core.documents import Document
|
||||
|
||||
|
||||
from langflow.schema.data import Data
|
||||
from tests.integration.utils import run_single_component
|
||||
from tests.api_keys import get_astradb_api_endpoint, get_astradb_application_token, get_openai_api_key
|
||||
from tests.integration.components.mock_components import TextToData
|
||||
from tests.integration.utils import ComponentInputHandle, run_single_component
|
||||
|
||||
BASIC_COLLECTION = "test_basic"
|
||||
SEARCH_COLLECTION = "test_search"
|
||||
|
|
@ -30,7 +26,7 @@ ALL_COLLECTIONS = [
|
|||
]
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
@pytest.fixture
|
||||
def astradb_client(request):
|
||||
client = AstraDB(api_endpoint=get_astradb_api_endpoint(), token=get_astradb_application_token())
|
||||
yield client
|
||||
|
|
@ -139,7 +135,7 @@ def test_astra_vectorize():
|
|||
|
||||
@pytest.mark.api_key_required
|
||||
def test_astra_vectorize_with_provider_api_key():
|
||||
"""tests vectorize using an openai api key"""
|
||||
"""Tests vectorize using an openai api key."""
|
||||
from langchain_astradb import AstraDBVectorStore, CollectionVectorServiceOptions
|
||||
|
||||
application_token = get_astradb_application_token()
|
||||
|
|
@ -196,7 +192,7 @@ def test_astra_vectorize_with_provider_api_key():
|
|||
|
||||
@pytest.mark.api_key_required
|
||||
def test_astra_vectorize_passes_authentication():
|
||||
"""tests vectorize using the authentication parameter"""
|
||||
"""Tests vectorize using the authentication parameter."""
|
||||
from langchain_astradb import AstraDBVectorStore, CollectionVectorServiceOptions
|
||||
|
||||
store = None
|
||||
|
|
|
|||
|
|
@ -1,10 +1,9 @@
|
|||
import pytest
|
||||
|
||||
from langflow.components.helpers.ParseJSONData import ParseJSONDataComponent
|
||||
from langflow.components.inputs import ChatInput
|
||||
from langflow.schema import Data
|
||||
from tests.integration.components.mock_components import TextToData
|
||||
from tests.integration.utils import run_single_component, ComponentInputHandle
|
||||
from tests.integration.utils import ComponentInputHandle, run_single_component
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -1,10 +1,9 @@
|
|||
import pytest
|
||||
from langflow.components.inputs import ChatInput
|
||||
from langflow.memory import get_messages
|
||||
from langflow.schema.message import Message
|
||||
from tests.integration.utils import run_single_component
|
||||
|
||||
from langflow.components.inputs import ChatInput
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default():
|
||||
|
|
|
|||
|
|
@ -1,14 +1,12 @@
|
|||
import pytest
|
||||
from langflow.components.inputs import TextInputComponent
|
||||
from langflow.schema.message import Message
|
||||
from tests.integration.utils import run_single_component
|
||||
|
||||
from langflow.components.inputs import TextInputComponent
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_input():
|
||||
outputs = await run_single_component(TextInputComponent, run_input="sample text", input_type="text")
|
||||
print(outputs)
|
||||
assert isinstance(outputs["text"], Message)
|
||||
assert outputs["text"].text == "sample text"
|
||||
assert outputs["text"].sender is None
|
||||
|
|
|
|||
|
|
@ -1,8 +1,7 @@
|
|||
import json
|
||||
from typing import List
|
||||
|
||||
from langflow.custom import Component
|
||||
from langflow.inputs import StrInput, BoolInput
|
||||
from langflow.inputs import BoolInput, StrInput
|
||||
from langflow.schema import Data
|
||||
from langflow.template import Output
|
||||
|
||||
|
|
@ -21,5 +20,5 @@ class TextToData(Component):
|
|||
return Data(data=json.loads(text))
|
||||
return Data(text=text)
|
||||
|
||||
def create_data(self) -> List[Data]:
|
||||
def create_data(self) -> list[Data]:
|
||||
return [self._to_data(t) for t in self.text_data]
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import os
|
||||
import pytest
|
||||
|
||||
import pytest
|
||||
from langflow.components.models.OpenAIModel import OpenAIModelComponent
|
||||
from langflow.components.output_parsers.OutputParser import OutputParserComponent
|
||||
from langflow.components.prompts.Prompt import PromptComponent
|
||||
|
|
@ -23,7 +23,7 @@ async def test_csv_output_parser_openai():
|
|||
prompt_handler = ComponentInputHandle(
|
||||
clazz=PromptComponent,
|
||||
inputs={
|
||||
"template": "List the first five positive integers.\n\n{format_instructions}",
|
||||
"template": "List the first five positive integers.\n\n{format_instructions}", # noqa: RUF027
|
||||
"format_instructions": format_instructions,
|
||||
},
|
||||
output_name="prompt",
|
||||
|
|
|
|||
|
|
@ -1,10 +1,9 @@
|
|||
import pytest
|
||||
from langflow.components.outputs import ChatOutput
|
||||
from langflow.memory import get_messages
|
||||
from langflow.schema.message import Message
|
||||
from tests.integration.utils import run_single_component
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_string():
|
||||
|
|
|
|||
|
|
@ -1,9 +1,8 @@
|
|||
import pytest
|
||||
from langflow.components.outputs import TextOutputComponent
|
||||
from langflow.schema.message import Message
|
||||
from tests.integration.utils import run_single_component
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test():
|
||||
|
|
|
|||
|
|
@ -1,13 +1,11 @@
|
|||
import pytest
|
||||
from langflow.components.prompts import PromptComponent
|
||||
from langflow.schema.message import Message
|
||||
from tests.integration.utils import run_single_component
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test():
|
||||
outputs = await run_single_component(PromptComponent, inputs={"template": "test {var1}", "var1": "from the var"})
|
||||
print(outputs)
|
||||
assert isinstance(outputs["prompt"], Message)
|
||||
assert outputs["prompt"].text == "test from the var"
|
||||
|
|
|
|||
|
|
@ -1,5 +1,4 @@
|
|||
import pytest
|
||||
|
||||
from langflow.components.inputs import ChatInput
|
||||
from langflow.components.outputs import ChatOutput
|
||||
from langflow.components.prompts import PromptComponent
|
||||
|
|
|
|||
|
|
@ -3,7 +3,6 @@ from uuid import uuid4
|
|||
import pytest
|
||||
from fastapi import status
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from langflow.graph.schema import RunOutputs
|
||||
from langflow.initial_setup.setup import load_starter_projects
|
||||
from langflow.load import run_flow_from_json
|
||||
|
|
@ -80,9 +79,8 @@ async def test_run_with_inputs_and_outputs(client, starter_project, created_api_
|
|||
@pytest.mark.noclient
|
||||
@pytest.mark.api_key_required
|
||||
def test_run_flow_from_json_object():
|
||||
"""Test loading a flow from a json file and applying tweaks"""
|
||||
_, projects = zip(*load_starter_projects())
|
||||
project = [project for project in projects if "Basic Prompting" in project["name"]][0]
|
||||
"""Test loading a flow from a json file and applying tweaks."""
|
||||
project = next(project for _, project in load_starter_projects() if "Basic Prompting" in project["name"])
|
||||
results = run_flow_from_json(project, input_value="test", fallback_to_env_vars=True)
|
||||
assert results is not None
|
||||
assert all(isinstance(result, RunOutputs) for result in results)
|
||||
|
|
|
|||
|
|
@ -1,21 +1,19 @@
|
|||
import dataclasses
|
||||
import os
|
||||
import uuid
|
||||
from typing import Optional, Any
|
||||
from typing import Any
|
||||
|
||||
import requests
|
||||
from astrapy.admin import parse_api_endpoint
|
||||
|
||||
from langflow.api.v1.schemas import InputValueRequest
|
||||
from langflow.custom import Component
|
||||
from langflow.field_typing import Embeddings
|
||||
from langflow.graph import Graph
|
||||
from langflow.processing.process import run_graph_internal
|
||||
import requests
|
||||
|
||||
|
||||
def check_env_vars(*vars):
|
||||
"""
|
||||
Check if all specified environment variables are set.
|
||||
"""Check if all specified environment variables are set.
|
||||
|
||||
Args:
|
||||
*vars (str): The environment variables to check.
|
||||
|
|
@ -27,8 +25,7 @@ def check_env_vars(*vars):
|
|||
|
||||
|
||||
def valid_nvidia_vectorize_region(api_endpoint: str) -> bool:
|
||||
"""
|
||||
Check if the specified region is valid.
|
||||
"""Check if the specified region is valid.
|
||||
|
||||
Args:
|
||||
region (str): The region to check.
|
||||
|
|
@ -38,8 +35,9 @@ def valid_nvidia_vectorize_region(api_endpoint: str) -> bool:
|
|||
"""
|
||||
parsed_endpoint = parse_api_endpoint(api_endpoint)
|
||||
if not parsed_endpoint:
|
||||
raise ValueError("Invalid ASTRA_DB_API_ENDPOINT")
|
||||
return parsed_endpoint.region in ["us-east-2"]
|
||||
msg = "Invalid ASTRA_DB_API_ENDPOINT"
|
||||
raise ValueError(msg)
|
||||
return parsed_endpoint.region == "us-east-2"
|
||||
|
||||
|
||||
class MockEmbeddings(Embeddings):
|
||||
|
|
@ -70,15 +68,15 @@ class JSONFlow:
|
|||
if node["data"]["type"] == component_type:
|
||||
result.append(node["id"])
|
||||
if not result:
|
||||
raise ValueError(
|
||||
f"Component of type {component_type} not found, available types: {', '.join(set(node['data']['type'] for node in self.json['data']['nodes']))}"
|
||||
)
|
||||
msg = f"Component of type {component_type} not found, available types: {', '.join({node['data']['type'] for node in self.json['data']['nodes']})}"
|
||||
raise ValueError(msg)
|
||||
return result
|
||||
|
||||
def get_component_by_type(self, component_type):
|
||||
components = self.get_components_by_type(component_type)
|
||||
if len(components) > 1:
|
||||
raise ValueError(f"Multiple components of type {component_type} found")
|
||||
msg = f"Multiple components of type {component_type} found"
|
||||
raise ValueError(msg)
|
||||
return components[0]
|
||||
|
||||
def set_value(self, component_id, key, value):
|
||||
|
|
@ -86,13 +84,15 @@ class JSONFlow:
|
|||
for node in self.json["data"]["nodes"]:
|
||||
if node["id"] == component_id:
|
||||
if key not in node["data"]["node"]["template"]:
|
||||
raise ValueError(f"Component {component_id} does not have input {key}")
|
||||
msg = f"Component {component_id} does not have input {key}"
|
||||
raise ValueError(msg)
|
||||
node["data"]["node"]["template"][key]["value"] = value
|
||||
node["data"]["node"]["template"][key]["load_from_db"] = False
|
||||
done = True
|
||||
break
|
||||
if not done:
|
||||
raise ValueError(f"Component {component_id} not found")
|
||||
msg = f"Component {component_id} not found"
|
||||
raise ValueError(msg)
|
||||
|
||||
|
||||
def download_flow_from_github(name: str, version: str) -> JSONFlow:
|
||||
|
|
@ -105,18 +105,15 @@ def download_flow_from_github(name: str, version: str) -> JSONFlow:
|
|||
|
||||
|
||||
async def run_json_flow(
|
||||
json_flow: JSONFlow, run_input: Optional[Any] = None, session_id: Optional[str] = None
|
||||
json_flow: JSONFlow, run_input: Any | None = None, session_id: str | None = None
|
||||
) -> dict[str, Any]:
|
||||
graph = Graph.from_payload(json_flow.json)
|
||||
return await run_flow(graph, run_input, session_id)
|
||||
|
||||
|
||||
async def run_flow(graph: Graph, run_input: Optional[Any] = None, session_id: Optional[str] = None) -> dict[str, Any]:
|
||||
async def run_flow(graph: Graph, run_input: Any | None = None, session_id: str | None = None) -> dict[str, Any]:
|
||||
graph.prepare()
|
||||
if run_input:
|
||||
graph_run_inputs = [InputValueRequest(input_value=run_input, type="chat")]
|
||||
else:
|
||||
graph_run_inputs = []
|
||||
graph_run_inputs = [InputValueRequest(input_value=run_input, type="chat")] if run_input else []
|
||||
|
||||
flow_id = str(uuid.uuid4())
|
||||
|
||||
|
|
@ -137,23 +134,24 @@ class ComponentInputHandle:
|
|||
|
||||
async def run_single_component(
|
||||
clazz: type,
|
||||
inputs: dict = None,
|
||||
run_input: Optional[Any] = None,
|
||||
session_id: Optional[str] = None,
|
||||
input_type: Optional[str] = "chat",
|
||||
inputs: dict | None = None,
|
||||
run_input: Any | None = None,
|
||||
session_id: str | None = None,
|
||||
input_type: str | None = "chat",
|
||||
) -> dict[str, Any]:
|
||||
user_id = str(uuid.uuid4())
|
||||
flow_id = str(uuid.uuid4())
|
||||
graph = Graph(user_id=user_id, flow_id=flow_id)
|
||||
|
||||
def _add_component(clazz: type, inputs: Optional[dict] = None) -> str:
|
||||
def _add_component(clazz: type, inputs: dict | None = None) -> str:
|
||||
raw_inputs = {}
|
||||
if inputs:
|
||||
for key, value in inputs.items():
|
||||
if not isinstance(value, ComponentInputHandle):
|
||||
raw_inputs[key] = value
|
||||
if isinstance(value, Component):
|
||||
raise ValueError("Component inputs must be wrapped in ComponentInputHandle")
|
||||
msg = "Component inputs must be wrapped in ComponentInputHandle"
|
||||
raise ValueError(msg)
|
||||
component = clazz(**raw_inputs, _user_id=user_id)
|
||||
component_id = graph.add_component(component)
|
||||
if inputs:
|
||||
|
|
@ -165,10 +163,7 @@ async def run_single_component(
|
|||
|
||||
component_id = _add_component(clazz, inputs)
|
||||
graph.prepare()
|
||||
if run_input:
|
||||
graph_run_inputs = [InputValueRequest(input_value=run_input, type=input_type)]
|
||||
else:
|
||||
graph_run_inputs = []
|
||||
graph_run_inputs = [InputValueRequest(input_value=run_input, type=input_type)] if run_input else []
|
||||
|
||||
_, _ = await run_graph_internal(
|
||||
graph, flow_id, session_id=session_id, inputs=graph_run_inputs, outputs=[component_id]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue