Merge remote-tracking branch 'origin/dev' into celery

This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-09-22 14:50:19 -03:00
commit b3febf25dd
31 changed files with 582 additions and 263 deletions

View file

@ -2,13 +2,14 @@ from contextlib import contextmanager
import json
from pathlib import Path
from typing import AsyncGenerator, TYPE_CHECKING
from langflow.api.v1.flows import get_session
from langflow.graph.graph.base import Graph
from langflow.services.auth.utils import get_password_hash
from langflow.services.database.models.flow.flow import Flow, FlowCreate
from langflow.services.database.models.user.user import User, UserCreate
import orjson
from langflow.services.database.utils import session_getter
from langflow.services.getters import get_db_manager
import pytest
from fastapi.testclient import TestClient
from httpx import AsyncClient
@ -16,6 +17,9 @@ from sqlmodel import SQLModel, Session, create_engine
from sqlmodel.pool import StaticPool
from typer.testing import CliRunner
# we need to import tmpdir
import tempfile
if TYPE_CHECKING:
from langflow.services.database.manager import DatabaseService
@ -61,22 +65,6 @@ def session_fixture():
yield session
@pytest.fixture(name="client")
def client_fixture(session: Session, monkeypatch):
def get_session_override():
return session
monkeypatch.setenv("LANGFLOW_AUTO_LOGIN", False)
from langflow.main import create_app
app = create_app()
app.dependency_overrides[get_session] = get_session_override
with TestClient(app) as client:
yield client
app.dependency_overrides.clear()
class Config:
broker_url = "redis://localhost:6379/0"
result_backend = "redis://localhost:6379/0"
@ -182,27 +170,33 @@ def json_vector_store():
return f.read()
# @contextmanager
# def session_getter():
# try:
# session = Session(engine)
# yield session
# except Exception as e:
# print("Session rollback because of exception:", e)
# session.rollback()
# raise
# finally:
# session.close()
@pytest.fixture(name="client", autouse=True)
def client_fixture(session: Session, monkeypatch):
# Set the database url to a test database
db_dir = tempfile.mkdtemp()
db_path = Path(db_dir) / "test.db"
monkeypatch.setenv("LANGFLOW_DATABASE_URL", f"sqlite:///{db_path}")
monkeypatch.setenv("LANGFLOW_AUTO_LOGIN", False)
def get_session_override():
return session
from langflow.main import create_app
app = create_app()
# app.dependency_overrides[get_session] = get_session_override
with TestClient(app) as client:
yield client
# app.dependency_overrides.clear()
monkeypatch.undo()
# clear the temp db
db_path.unlink()
# create a fixture for session_getter above
@pytest.fixture(name="session_getter")
def session_getter_fixture(client):
engine = create_engine(
"sqlite://", connect_args={"check_same_thread": False}, poolclass=StaticPool
)
SQLModel.metadata.create_all(engine)
@contextmanager
def blank_session_getter(db_service: "DatabaseService"):
with Session(db_service.engine) as session:
@ -228,17 +222,18 @@ def test_user(client):
@pytest.fixture(scope="function")
def active_user(client, session):
user = User(
username="activeuser",
password=get_password_hash(
"testpassword"
), # Assuming password needs to be hashed
is_active=True,
is_superuser=False,
)
session.add(user)
session.commit()
def active_user(client):
db_manager = get_db_manager()
with session_getter(db_manager) as session:
user = User(
username="activeuser",
password=get_password_hash("testpassword"),
is_active=True,
is_superuser=False,
)
session.add(user)
session.commit()
session.refresh(user)
return user
@ -253,7 +248,7 @@ def logged_in_headers(client, active_user):
@pytest.fixture
def flow(client, json_flow: str, session, active_user):
def flow(client, json_flow: str, active_user):
from langflow.services.database.models.flow.flow import FlowCreate
loaded_json = json.loads(json_flow)
@ -261,8 +256,10 @@ def flow(client, json_flow: str, session, active_user):
name="test_flow", data=loaded_json.get("data"), user_id=active_user.id
)
flow = Flow(**flow_data.dict())
session.add(flow)
session.commit()
with session_getter(get_db_manager()) as session:
session.add(flow)
session.commit()
session.refresh(flow)
return flow