ref: Auto-fix ruff rules in tests (#4154)

This commit is contained in:
Christophe Bornet 2024-10-16 17:42:36 +02:00 • committed by GitHub
commit 45c8f98692
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
80 changed files with 359 additions and 456 deletions

View file

@ -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

View file

@ -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

View file

@ -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():

View file

@ -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

View file

@ -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]

View file

@ -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",

View file

@ -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():

View file

@ -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():

View file

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

View file

@ -1,5 +1,4 @@
import pytest
from langflow.components.inputs import ChatInput
from langflow.components.outputs import ChatOutput
from langflow.components.prompts import PromptComponent

View file

@ -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)

View file

@ -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]