ref: Make create_or_update_starter_projects async (#5165)

Make create_or_update_starter_projects async
This commit is contained in:
Christophe Bornet 2024-12-10 08:51:35 +01:00 • committed by GitHub
commit fe6ec1690b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 59 additions and 60 deletions

View file

@ -5,7 +5,7 @@ from fastapi import status
from httpx import AsyncClient
from langflow.graph.schema import RunOutputs
from langflow.initial_setup.setup import load_starter_projects
from langflow.load import run_flow_from_json
from langflow.load.load import arun_flow_from_json
@pytest.mark.api_key_required
@ -78,9 +78,9 @@ 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():
async def test_run_flow_from_json_object():
"""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)
project = next(project for _, project in await load_starter_projects() if "Basic Prompting" in project["name"])
results = await arun_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,5 +1,3 @@
import asyncio
import pytest
from langflow.services.deps import get_settings_service
@ -72,7 +70,7 @@ async def test_create_starter_projects():
await initialize_services(fix_migration=False)
settings_service = get_settings_service()
types_dict = await get_and_cache_all_types_dict(settings_service)
await asyncio.to_thread(create_or_update_starter_projects, types_dict)
await create_or_update_starter_projects(types_dict)
assert "test_performance.db" in settings_service.settings.database_url

View file

@ -256,8 +256,8 @@ def test_update_source_handle():
assert updated_edge["data"]["sourceHandle"]["id"] == "last_node"
def test_serialize_graph():
starter_projects = load_starter_projects()
async def test_serialize_graph():
starter_projects = await load_starter_projects()
data = starter_projects[0][1]["data"]
graph = Graph.from_payload(data)
assert isinstance(graph, Graph)

View file

@ -1,4 +1,3 @@
import asyncio
import json
from typing import NamedTuple
from uuid import UUID, uuid4
@ -605,7 +604,7 @@ async def test_delete_nonexistent_flow(client: AsyncClient, logged_in_headers):
@pytest.mark.usefixtures("active_user")
async def test_read_only_starter_projects(client: AsyncClient, logged_in_headers):
response = await client.get("api/v1/flows/basic_examples/", headers=logged_in_headers)
starter_projects = await asyncio.to_thread(load_starter_projects)
starter_projects = await load_starter_projects()
assert response.status_code == 200
assert len(response.json()) == len(starter_projects)

View file

@ -1,6 +1,4 @@
import asyncio
from datetime import datetime
from pathlib import Path
import anyio
import pytest
@ -18,15 +16,15 @@ from sqlalchemy.orm import selectinload
from sqlmodel import select
def test_load_starter_projects():
projects = load_starter_projects()
async def test_load_starter_projects():
projects = await load_starter_projects()
assert isinstance(projects, list)
assert all(isinstance(project[1], dict) for project in projects)
assert all(isinstance(project[0], Path) for project in projects)
assert all(isinstance(project[0], anyio.Path) for project in projects)
def test_get_project_data():
projects = load_starter_projects()
async def test_get_project_data():
projects = await load_starter_projects()
for _, project in projects:
(
project_name,
@ -56,7 +54,7 @@ def test_get_project_data():
async def test_create_or_update_starter_projects():
async with async_session_scope() as session:
# Get the number of projects returned by load_starter_projects
num_projects = len(await asyncio.to_thread(load_starter_projects))
num_projects = len(await load_starter_projects())
# Get the number of projects in the database
stmt = select(Folder).options(selectinload(Folder.flows)).where(Folder.name == STARTER_FOLDER_NAME)

View file

@ -22,7 +22,7 @@ from langflow.load import load_flow_from_json
async def test_load_flow_from_json_object():
"""Test loading a flow from a json file and applying tweaks."""
result = await asyncio.to_thread(load_starter_projects)
result = await load_starter_projects()
project = result[0][1]
loaded = await asyncio.to_thread(load_flow_from_json, project)
assert loaded is not None