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

This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-09-01 15:24:07 -03:00
commit 5c258d08f0
174 changed files with 14281 additions and 4261 deletions

View file

@ -5,6 +5,9 @@ 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
from langflow.services.database.models.user.user import User, UserCreate
import pytest
from fastapi.testclient import TestClient
from httpx import AsyncClient
@ -43,7 +46,7 @@ async def async_client() -> AsyncGenerator:
# Create client fixture for FastAPI
@pytest.fixture(scope="module")
@pytest.fixture(scope="module", autouse=True)
def client():
from langflow.main import create_app
@ -155,3 +158,53 @@ def session_getter_fixture(client):
@pytest.fixture
def runner():
return CliRunner()
@pytest.fixture
def test_user(client):
user_data = UserCreate(
username="testuser",
password="testpassword",
)
response = client.post("/api/v1/user", json=user_data.dict())
return response.json()
@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()
return user
@pytest.fixture
def logged_in_headers(client, active_user):
login_data = {"username": active_user.username, "password": "testpassword"}
response = client.post("/api/v1/login", data=login_data)
assert response.status_code == 200
tokens = response.json()
a_token = tokens["access_token"]
return {"Authorization": f"Bearer {a_token}"}
@pytest.fixture
def flow(client, json_flow: str, session, active_user):
from langflow.services.database.models.flow.flow import FlowCreate
loaded_json = json.loads(json_flow)
flow_data = FlowCreate(
name="test_flow", data=loaded_json.get("data"), user_id=active_user.id
)
flow = Flow(**flow_data.dict())
session.add(flow)
session.commit()
return flow

View file

@ -1,8 +1,8 @@
from fastapi.testclient import TestClient
def test_zero_shot_agent(client: TestClient):
response = client.get("api/v1/all")
def test_zero_shot_agent(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
agents = json_response["agents"]
@ -113,8 +113,8 @@ def test_zero_shot_agent(client: TestClient):
}
def test_json_agent(client: TestClient):
response = client.get("api/v1/all")
def test_json_agent(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
agents = json_response["agents"]
@ -152,8 +152,8 @@ def test_json_agent(client: TestClient):
}
def test_csv_agent(client: TestClient):
response = client.get("api/v1/all")
def test_csv_agent(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
agents = json_response["agents"]
@ -195,8 +195,8 @@ def test_csv_agent(client: TestClient):
}
def test_initialize_agent(client: TestClient):
response = client.get("api/v1/all")
def test_initialize_agent(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
agents = json_response["agents"]

50
tests/test_api_key.py Normal file
View file

@ -0,0 +1,50 @@
import pytest
from langflow.services.database.models.api_key import ApiKeyCreate
@pytest.fixture
def api_key(client, logged_in_headers, active_user):
api_key = ApiKeyCreate(name="test-api-key")
response = client.post(
"api/v1/api_key", data=api_key.json(), headers=logged_in_headers
)
assert response.status_code == 200, response.text
return response.json()
def test_get_api_keys(client, logged_in_headers, api_key):
response = client.get("api/v1/api_key", headers=logged_in_headers)
assert response.status_code == 200, response.text
data = response.json()
assert "total_count" in data
assert "user_id" in data
assert "api_keys" in data
assert any("test-api-key" in api_key["name"] for api_key in data["api_keys"])
# assert all api keys in data["api_keys"] are masked
assert all("**" in api_key["api_key"] for api_key in data["api_keys"])
# Add more assertions as needed based on the expected data structure and content
def test_create_api_key(client, logged_in_headers):
api_key_name = "test-api-key"
response = client.post(
"api/v1/api_key", json={"name": api_key_name}, headers=logged_in_headers
)
assert response.status_code == 200
data = response.json()
assert "name" in data and data["name"] == api_key_name
assert "api_key" in data
# When creating the API key is returned which is
# the only time the API key is unmasked
assert "**" not in data["api_key"]
def test_delete_api_key(client, logged_in_headers, active_user, api_key):
# Assuming a function to create a test API key, returning the key ID
api_key_id = api_key["id"]
response = client.delete(f"api/v1/api_key/{api_key_id}", headers=logged_in_headers)
assert response.status_code == 200
data = response.json()
assert data["detail"] == "API Key deleted"
# Optionally, add a follow-up check to ensure that the key is actually removed from the database

View file

@ -1,8 +1,8 @@
from fastapi.testclient import TestClient
# def test_chains_settings(client: TestClient):
# response = client.get("api/v1/all")
# def test_chains_settings(client: TestClient, logged_in_headers):
# response = client.get("api/v1/all", headers=logged_in_headers)
# assert response.status_code == 200
# json_response = response.json()
# chains = json_response["chains"]
@ -10,8 +10,8 @@ from fastapi.testclient import TestClient
# Test the ConversationChain object
def test_conversation_chain(client: TestClient):
response = client.get("api/v1/all")
def test_conversation_chain(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
chains = json_response["chains"]
@ -102,8 +102,8 @@ def test_conversation_chain(client: TestClient):
)
def test_llm_chain(client: TestClient):
response = client.get("api/v1/all")
def test_llm_chain(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
chains = json_response["chains"]
@ -173,8 +173,8 @@ def test_llm_chain(client: TestClient):
}
def test_llm_checker_chain(client: TestClient):
response = client.get("api/v1/all")
def test_llm_checker_chain(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
chains = json_response["chains"]
@ -207,8 +207,8 @@ def test_llm_checker_chain(client: TestClient):
assert chain["description"] == ""
def test_llm_math_chain(client: TestClient):
response = client.get("api/v1/all")
def test_llm_math_chain(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
chains = json_response["chains"]
@ -299,8 +299,8 @@ def test_llm_math_chain(client: TestClient):
)
def test_series_character_chain(client: TestClient):
response = client.get("api/v1/all")
def test_series_character_chain(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
chains = json_response["chains"]
@ -367,8 +367,8 @@ def test_series_character_chain(client: TestClient):
)
def test_mid_journey_prompt_chain(client: TestClient):
response = client.get("api/v1/all")
def test_mid_journey_prompt_chain(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
chains = json_response["chains"]
@ -408,8 +408,8 @@ def test_mid_journey_prompt_chain(client: TestClient):
)
def test_time_travel_guide_chain(client: TestClient):
response = client.get("api/v1/all")
def test_time_travel_guide_chain(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
chains = json_response["chains"]

View file

@ -23,8 +23,14 @@ def test_components_path(runner, client, default_settings):
result = runner.invoke(
app,
["--components-path", str(temp_dir), *default_settings],
["run", "--components-path", str(temp_dir), *default_settings],
)
assert result.exit_code == 0, result.stdout
settings_manager = utils.get_settings_manager()
assert str(temp_dir) in settings_manager.settings.COMPONENTS_PATH
def test_superuser(runner, client, session):
result = runner.invoke(app, ["superuser"], input="admin\nadmin\n")
assert result.exit_code == 0, result.stdout
assert "Superuser created successfully." in result.stdout

View file

@ -473,15 +473,16 @@ def test_build_config_no_code():
@pytest.fixture
def component():
def component(client, active_user):
return CustomComponent(
user_id=active_user.id,
field_config={
"fields": {
"llm": {"type": "str"},
"url": {"type": "str"},
"year": {"type": "int"},
}
}
},
)

View file

@ -1,8 +1,9 @@
import json
from langflow.services.database.models.base import orjson_dumps
import orjson
import pytest
from uuid import UUID, uuid4
from sqlalchemy.orm import Session
from sqlmodel import Session
from fastapi.testclient import TestClient
@ -16,7 +17,7 @@ def json_style():
# color: str = Field(index=True)
# emoji: str = Field(index=False)
# flow_id: UUID = Field(default=None, foreign_key="flow.id")
return json.dumps(
return orjson_dumps(
{
"color": "red",
"emoji": "👍",
@ -24,63 +25,69 @@ def json_style():
)
def test_create_flow(client: TestClient, json_flow: str):
flow = json.loads(json_flow)
def test_create_flow(
client: TestClient, json_flow: str, active_user, logged_in_headers
):
flow = orjson.loads(json_flow)
data = flow["data"]
flow = FlowCreate(name="Test Flow", description="description", data=data)
response = client.post("api/v1/flows/", json=flow.dict())
response = client.post("api/v1/flows/", json=flow.dict(), headers=logged_in_headers)
assert response.status_code == 201
assert response.json()["name"] == flow.name
assert response.json()["data"] == flow.data
# flow is optional so we can create a flow without a flow
flow = FlowCreate(name="Test Flow")
response = client.post("api/v1/flows/", json=flow.dict(exclude_unset=True))
response = client.post(
"api/v1/flows/", json=flow.dict(exclude_unset=True), headers=logged_in_headers
)
assert response.status_code == 201
assert response.json()["name"] == flow.name
assert response.json()["data"] == flow.data
def test_read_flows(client: TestClient, json_flow: str):
flow_data = json.loads(json_flow)
def test_read_flows(client: TestClient, json_flow: str, active_user, logged_in_headers):
flow_data = orjson.loads(json_flow)
data = flow_data["data"]
flow = FlowCreate(name="Test Flow", description="description", data=data)
response = client.post("api/v1/flows/", json=flow.dict())
response = client.post("api/v1/flows/", json=flow.dict(), headers=logged_in_headers)
assert response.status_code == 201
assert response.json()["name"] == flow.name
assert response.json()["data"] == flow.data
flow = FlowCreate(name="Test Flow", description="description", data=data)
response = client.post("api/v1/flows/", json=flow.dict())
response = client.post("api/v1/flows/", json=flow.dict(), headers=logged_in_headers)
assert response.status_code == 201
assert response.json()["name"] == flow.name
assert response.json()["data"] == flow.data
response = client.get("api/v1/flows/")
response = client.get("api/v1/flows/", headers=logged_in_headers)
assert response.status_code == 200
assert len(response.json()) > 0
def test_read_flow(client: TestClient, json_flow: str):
flow = json.loads(json_flow)
def test_read_flow(client: TestClient, json_flow: str, active_user, logged_in_headers):
flow = orjson.loads(json_flow)
data = flow["data"]
flow = FlowCreate(name="Test Flow", description="description", data=data)
response = client.post("api/v1/flows/", json=flow.dict())
response = client.post("api/v1/flows/", json=flow.dict(), headers=logged_in_headers)
flow_id = response.json()["id"] # flow_id should be a UUID but is a string
# turn it into a UUID
flow_id = UUID(flow_id)
response = client.get(f"api/v1/flows/{flow_id}")
response = client.get(f"api/v1/flows/{flow_id}", headers=logged_in_headers)
assert response.status_code == 200
assert response.json()["name"] == flow.name
assert response.json()["data"] == flow.data
def test_update_flow(client: TestClient, json_flow: str):
flow = json.loads(json_flow)
def test_update_flow(
client: TestClient, json_flow: str, active_user, logged_in_headers
):
flow = orjson.loads(json_flow)
data = flow["data"]
flow = FlowCreate(name="Test Flow", description="description", data=data)
response = client.post("api/v1/flows/", json=flow.dict())
response = client.post("api/v1/flows/", json=flow.dict(), headers=logged_in_headers)
flow_id = response.json()["id"]
updated_flow = FlowUpdate(
@ -88,7 +95,9 @@ def test_update_flow(client: TestClient, json_flow: str):
description="updated description",
data=data,
)
response = client.patch(f"api/v1/flows/{flow_id}", json=updated_flow.dict())
response = client.patch(
f"api/v1/flows/{flow_id}", json=updated_flow.dict(), headers=logged_in_headers
)
assert response.status_code == 200
assert response.json()["name"] == updated_flow.name
@ -96,19 +105,23 @@ def test_update_flow(client: TestClient, json_flow: str):
# assert response.json()["data"] == updated_flow.data
def test_delete_flow(client: TestClient, json_flow: str):
flow = json.loads(json_flow)
def test_delete_flow(
client: TestClient, json_flow: str, active_user, logged_in_headers
):
flow = orjson.loads(json_flow)
data = flow["data"]
flow = FlowCreate(name="Test Flow", description="description", data=data)
response = client.post("api/v1/flows/", json=flow.dict())
response = client.post("api/v1/flows/", json=flow.dict(), headers=logged_in_headers)
flow_id = response.json()["id"]
response = client.delete(f"api/v1/flows/{flow_id}")
response = client.delete(f"api/v1/flows/{flow_id}", headers=logged_in_headers)
assert response.status_code == 200
assert response.json()["message"] == "Flow deleted successfully"
def test_create_flows(client: TestClient, session: Session, json_flow: str):
flow = json.loads(json_flow)
def test_create_flows(
client: TestClient, session: Session, json_flow: str, logged_in_headers
):
flow = orjson.loads(json_flow)
data = flow["data"]
# Create test data
flow_list = FlowListCreate(
@ -118,7 +131,9 @@ def test_create_flows(client: TestClient, session: Session, json_flow: str):
]
)
# Make request to endpoint
response = client.post("api/v1/flows/batch/", json=flow_list.dict())
response = client.post(
"api/v1/flows/batch/", json=flow_list.dict(), headers=logged_in_headers
)
# Check response status code
assert response.status_code == 201
# Check response data
@ -132,8 +147,10 @@ def test_create_flows(client: TestClient, session: Session, json_flow: str):
assert response_data[1]["data"] == data
def test_upload_file(client: TestClient, session: Session, json_flow: str):
flow = json.loads(json_flow)
def test_upload_file(
client: TestClient, session: Session, json_flow: str, logged_in_headers
):
flow = orjson.loads(json_flow)
data = flow["data"]
# Create test data
flow_list = FlowListCreate(
@ -142,10 +159,11 @@ def test_upload_file(client: TestClient, session: Session, json_flow: str):
FlowCreate(name="Flow 2", description="description", data=data),
]
)
file_contents = json.dumps(flow_list.dict())
file_contents = orjson_dumps(flow_list.dict())
response = client.post(
"api/v1/flows/upload/",
files={"file": ("examples.json", file_contents, "application/json")},
headers=logged_in_headers,
)
# Check response status code
assert response.status_code == 201
@ -160,8 +178,10 @@ def test_upload_file(client: TestClient, session: Session, json_flow: str):
assert response_data[1]["data"] == data
def test_download_file(client: TestClient, session: Session, json_flow):
flow = json.loads(json_flow)
def test_download_file(
client: TestClient, session: Session, json_flow, active_user, logged_in_headers
):
flow = orjson.loads(json_flow)
data = flow["data"]
# Create test data
flow_list = FlowListCreate(
@ -171,11 +191,12 @@ def test_download_file(client: TestClient, session: Session, json_flow):
]
)
for flow in flow_list.flows:
flow.user_id = active_user.id
db_flow = Flow.from_orm(flow)
session.add(db_flow)
session.commit()
# Make request to endpoint
response = client.get("api/v1/flows/download/")
response = client.get("api/v1/flows/download/", headers=logged_in_headers)
# Check response status code
assert response.status_code == 200
# Check response data
@ -189,32 +210,44 @@ def test_download_file(client: TestClient, session: Session, json_flow):
assert response_data[1]["data"] == data
def test_create_flow_with_invalid_data(client: TestClient):
def test_create_flow_with_invalid_data(
client: TestClient, active_user, logged_in_headers
):
flow = {"name": "a" * 256, "data": "Invalid flow data"}
response = client.post("api/v1/flows/", json=flow)
response = client.post("api/v1/flows/", json=flow, headers=logged_in_headers)
assert response.status_code == 422
def test_get_nonexistent_flow(client: TestClient):
def test_get_nonexistent_flow(client: TestClient, active_user, logged_in_headers):
uuid = uuid4()
response = client.get(f"api/v1/flows/{uuid}")
response = client.get(f"api/v1/flows/{uuid}", headers=logged_in_headers)
assert response.status_code == 404
def test_update_flow_idempotency(client: TestClient, json_flow: str):
flow_data = json.loads(json_flow)
def test_update_flow_idempotency(
client: TestClient, json_flow: str, active_user, logged_in_headers
):
flow_data = orjson.loads(json_flow)
data = flow_data["data"]
flow_data = FlowCreate(name="Test Flow", description="description", data=data)
response = client.post("api/v1/flows/", json=flow_data.dict())
response = client.post(
"api/v1/flows/", json=flow_data.dict(), headers=logged_in_headers
)
flow_id = response.json()["id"]
updated_flow = FlowCreate(name="Updated Flow", description="description", data=data)
response1 = client.put(f"api/v1/flows/{flow_id}", json=updated_flow.dict())
response2 = client.put(f"api/v1/flows/{flow_id}", json=updated_flow.dict())
response1 = client.put(
f"api/v1/flows/{flow_id}", json=updated_flow.dict(), headers=logged_in_headers
)
response2 = client.put(
f"api/v1/flows/{flow_id}", json=updated_flow.dict(), headers=logged_in_headers
)
assert response1.json() == response2.json()
def test_update_nonexistent_flow(client: TestClient, json_flow: str):
flow_data = json.loads(json_flow)
def test_update_nonexistent_flow(
client: TestClient, json_flow: str, active_user, logged_in_headers
):
flow_data = orjson.loads(json_flow)
data = flow_data["data"]
uuid = uuid4()
updated_flow = FlowCreate(
@ -222,17 +255,19 @@ def test_update_nonexistent_flow(client: TestClient, json_flow: str):
description="description",
data=data,
)
response = client.patch(f"api/v1/flows/{uuid}", json=updated_flow.dict())
response = client.patch(
f"api/v1/flows/{uuid}", json=updated_flow.dict(), headers=logged_in_headers
)
assert response.status_code == 404
def test_delete_nonexistent_flow(client: TestClient):
def test_delete_nonexistent_flow(client: TestClient, active_user, logged_in_headers):
uuid = uuid4()
response = client.delete(f"api/v1/flows/{uuid}")
response = client.delete(f"api/v1/flows/{uuid}", headers=logged_in_headers)
assert response.status_code == 404
def test_read_empty_flows(client: TestClient):
response = client.get("api/v1/flows/")
def test_read_empty_flows(client: TestClient, active_user, logged_in_headers):
response = client.get("api/v1/flows/", headers=logged_in_headers)
assert response.status_code == 200
assert len(response.json()) == 0

View file

@ -1,3 +1,7 @@
import uuid
from langflow.services.auth.utils import get_password_hash
from langflow.services.database.models.api_key.api_key import ApiKey
from langflow.services.utils import get_settings_manager
import pytest
from fastapi.testclient import TestClient
from langflow.interface.tools.constants import CUSTOM_TOOLS
@ -83,8 +87,141 @@ PROMPT_REQUEST = {
}
def test_get_all(client: TestClient):
response = client.get("api/v1/all")
@pytest.fixture
def created_api_key(session, active_user):
hashed = get_password_hash("random_key")
api_key = ApiKey(
name="test_api_key",
user_id=active_user.id,
api_key="random_key",
hashed_api_key=hashed,
)
session.add(api_key)
session.commit()
session.refresh(api_key)
return api_key
def test_process_flow_invalid_api_key(client, flow, monkeypatch):
# Mock de process_graph_cached
def mock_process_graph_cached(*args, **kwargs):
return {}, "session_id_mock"
settings_manager = get_settings_manager()
settings_manager.auth_settings.AUTO_LOGIN = False
from langflow.api.v1 import endpoints
monkeypatch.setattr(endpoints, "process_graph_cached", mock_process_graph_cached)
headers = {"api-key": "invalid_api_key"}
post_data = {
"inputs": {"key": "value"},
"tweaks": None,
"clear_cache": False,
"session_id": None,
}
response = client.post(f"api/v1/process/{flow.id}", headers=headers, json=post_data)
assert response.status_code == 403
assert response.json() == {"detail": "Invalid or missing API key"}
def test_process_flow_invalid_id(client, monkeypatch, created_api_key):
def mock_process_graph_cached(*args, **kwargs):
return {}, "session_id_mock"
from langflow.api.v1 import endpoints
monkeypatch.setattr(endpoints, "process_graph_cached", mock_process_graph_cached)
api_key = created_api_key.api_key
headers = {"api-key": api_key}
post_data = {
"inputs": {"key": "value"},
"tweaks": None,
"clear_cache": False,
"session_id": None,
}
invalid_id = uuid.uuid4()
response = client.post(
f"api/v1/process/{invalid_id}", headers=headers, json=post_data
)
assert response.status_code == 404
assert f"Flow {invalid_id} not found" in response.json()["detail"]
def test_process_flow_without_autologin(client, flow, monkeypatch, created_api_key):
# Mock de process_graph_cached
from langflow.api.v1 import endpoints
settings_manager = get_settings_manager()
settings_manager.auth_settings.AUTO_LOGIN = False
def mock_process_graph_cached(*args, **kwargs):
return {}, "session_id_mock"
monkeypatch.setattr(endpoints, "process_graph_cached", mock_process_graph_cached)
api_key = created_api_key.api_key
headers = {"api-key": api_key}
# Dummy POST data
post_data = {
"inputs": {"key": "value"},
"tweaks": None,
"clear_cache": False,
"session_id": None,
}
# Make the request to the FastAPI TestClient
response = client.post(f"api/v1/process/{flow.id}", headers=headers, json=post_data)
# Check the response
assert response.status_code == 200, response.json()
assert response.json()["result"] == {}
assert response.json()["session_id"] == "session_id_mock"
def test_process_flow_fails_autologin_off(client, flow, monkeypatch):
# Mock de process_graph_cached
from langflow.api.v1 import endpoints
settings_manager = get_settings_manager()
settings_manager.auth_settings.AUTO_LOGIN = False
def mock_process_graph_cached(*args, **kwargs):
return {}, "session_id_mock"
monkeypatch.setattr(endpoints, "process_graph_cached", mock_process_graph_cached)
headers = {"api-key": "api_key"}
# Dummy POST data
post_data = {
"inputs": {"key": "value"},
"tweaks": None,
"clear_cache": False,
"session_id": None,
}
# Make the request to the FastAPI TestClient
response = client.post(f"api/v1/process/{flow.id}", headers=headers, json=post_data)
# Check the response
assert response.status_code == 403, response.json()
assert response.json() == {"detail": "Invalid or missing API key"}
def test_get_all(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
# We need to test the custom nodes

View file

@ -2,9 +2,9 @@ from fastapi.testclient import TestClient
from langflow.services.utils import get_settings_manager
def test_llms_settings(client: TestClient):
def test_llms_settings(client: TestClient, logged_in_headers):
settings_manager = get_settings_manager()
response = client.get("api/v1/all")
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
llms = json_response["llms"]
@ -103,8 +103,8 @@ def test_llms_settings(client: TestClient):
# }
def test_openai(client: TestClient):
response = client.get("api/v1/all")
def test_openai(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
language_models = json_response["llms"]
@ -369,8 +369,8 @@ def test_openai(client: TestClient):
}
def test_chat_open_ai(client: TestClient):
response = client.get("api/v1/all")
def test_chat_open_ai(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
language_models = json_response["llms"]
@ -542,8 +542,7 @@ def test_chat_open_ai(client: TestClient):
}
assert template["_type"] == "ChatOpenAI"
assert (
model["description"]
== "Wrapper around OpenAI Chat large language models." # noqa E501
model["description"] == "`OpenAI` Chat large language models API." # noqa E501
)
assert set(model["base_classes"]) == {
"BaseLLM",

View file

@ -1,5 +1,4 @@
import json
import pytest
from langchain.chains.base import Chain
from langflow.processing.process import load_flow_from_json

47
tests/test_login.py Normal file
View file

@ -0,0 +1,47 @@
import pytest
from langflow.services.database.models.user import User
from langflow.services.auth.utils import get_password_hash
@pytest.fixture
def test_user():
return User(
username="testuser",
password=get_password_hash(
"testpassword"
), # Assuming password needs to be hashed
is_active=True,
is_superuser=False,
)
def test_login_successful(client, test_user, session):
# Adding the test user to the database
session.add(test_user)
session.commit()
response = client.post(
"api/v1/login", data={"username": "testuser", "password": "testpassword"}
)
assert response.status_code == 200
assert "access_token" in response.json()
def test_login_unsuccessful_wrong_username(client):
response = client.post(
"api/v1/login", data={"username": "wrongusername", "password": "testpassword"}
)
assert response.status_code == 401
assert response.json()["detail"] == "Incorrect username or password"
def test_login_unsuccessful_wrong_password(client, test_user, session):
# Adding the test user to the database
session.add(test_user)
session.commit()
response = client.post(
"api/v1/login", data={"username": "testuser", "password": "wrongpassword"}
)
assert response.status_code == 401
assert response.json()["detail"] == "Incorrect username or password"

View file

@ -2,17 +2,17 @@ from fastapi.testclient import TestClient
from langflow.services.utils import get_settings_manager
def test_prompts_settings(client: TestClient):
def test_prompts_settings(client: TestClient, logged_in_headers):
settings_manager = get_settings_manager()
response = client.get("api/v1/all")
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
prompts = json_response["prompts"]
assert set(prompts.keys()) == set(settings_manager.settings.PROMPTS)
def test_prompt_template(client: TestClient):
response = client.get("api/v1/all")
def test_prompt_template(client: TestClient, logged_in_headers):
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
prompts = json_response["prompts"]

219
tests/test_user.py Normal file
View file

@ -0,0 +1,219 @@
from datetime import datetime
from langflow.services.auth.utils import create_super_user, get_password_hash
from langflow.services.database.models.user.user import User
from langflow.services.utils import get_settings_manager
import pytest
from langflow.services.database.models.user import UserUpdate
@pytest.fixture
def super_user(client, session):
settings_manager = get_settings_manager()
auth_settings = settings_manager.auth_settings
return create_super_user(
db=session,
username=auth_settings.FIRST_SUPERUSER,
password=auth_settings.FIRST_SUPERUSER_PASSWORD,
)
@pytest.fixture
def super_user_headers(client, super_user):
settings_manager = get_settings_manager()
auth_settings = settings_manager.auth_settings
login_data = {
"username": auth_settings.FIRST_SUPERUSER,
"password": auth_settings.FIRST_SUPERUSER_PASSWORD,
}
response = client.post("/api/v1/login", data=login_data)
assert response.status_code == 200
tokens = response.json()
a_token = tokens["access_token"]
return {"Authorization": f"Bearer {a_token}"}
@pytest.fixture
def deactivated_user(session):
user = User(
username="deactivateduser",
password=get_password_hash("testpassword"),
is_active=False,
is_superuser=False,
last_login_at=datetime.now(),
)
session.add(user)
session.commit()
return user
def test_user_waiting_for_approval(client, session):
# Create a user that is not active and has never logged in
user = User(
username="waitingforapproval",
password=get_password_hash("testpassword"),
is_active=False,
last_login_at=None,
)
session.add(user)
session.commit()
login_data = {"username": "waitingforapproval", "password": "testpassword"}
response = client.post("/api/v1/login", data=login_data)
assert response.status_code == 400
assert response.json()["detail"] == "Waiting for approval"
def test_deactivated_user_cannot_login(client, deactivated_user):
login_data = {"username": deactivated_user.username, "password": "testpassword"}
response = client.post("/api/v1/login", data=login_data)
assert response.status_code == 400, response.json()
assert response.json()["detail"] == "Inactive user"
def test_deactivated_user_cannot_access(client, deactivated_user, logged_in_headers):
# Assuming the headers for deactivated_user
response = client.get("/api/v1/users", headers=logged_in_headers)
assert response.status_code == 400, response.json()
assert response.json()["detail"] == "The user doesn't have enough privileges"
def test_data_consistency_after_update(client, active_user, logged_in_headers):
user_id = active_user.id
update_data = UserUpdate(username="newname")
response = client.patch(
f"/api/v1/user/{user_id}", json=update_data.dict(), headers=logged_in_headers
)
assert response.status_code == 200
# Fetch the updated user from the database
response = client.get("/api/v1/user", headers=logged_in_headers)
assert response.json()["username"] == "newname", response.json()
def test_data_consistency_after_delete(client, test_user, super_user_headers):
user_id = test_user["id"]
response = client.delete(f"/api/v1/user/{user_id}", headers=super_user_headers)
assert response.status_code == 200
# Attempt to fetch the deleted user from the database
response = client.get("/api/v1/users", headers=super_user_headers)
assert response.status_code == 200
assert all(user["id"] != user_id for user in response.json()["users"])
def test_inactive_user(client, session):
# Create a user that is not active and has a last_login_at value
user = User(
username="inactiveuser",
password=get_password_hash("testpassword"),
is_active=False,
last_login_at="2023-01-01T00:00:00", # Set to a valid datetime string
)
session.add(user)
session.commit()
login_data = {"username": "inactiveuser", "password": "testpassword"}
response = client.post("/api/v1/login", data=login_data)
assert response.status_code == 400
assert response.json()["detail"] == "Inactive user"
def test_add_user(client, test_user):
assert test_user["username"] == "testuser"
# This is not used in the Frontend at the moment
# def test_read_current_user(client: TestClient, active_user):
# # First we need to login to get the access token
# login_data = {"username": "testuser", "password": "testpassword"}
# response = client.post("/api/v1/login", data=login_data)
# assert response.status_code == 200
# headers = {"Authorization": f"Bearer {response.json()['access_token']}"}
# response = client.get("/api/v1/user", headers=headers)
# assert response.status_code == 200, response.json()
# assert response.json()["username"] == "testuser"
def test_read_all_users(client, super_user_headers):
response = client.get("/api/v1/users", headers=super_user_headers)
assert response.status_code == 200, response.json()
assert isinstance(response.json()["users"], list)
def test_normal_user_cant_read_all_users(client, logged_in_headers):
response = client.get("/api/v1/users", headers=logged_in_headers)
assert response.status_code == 400, response.json()
assert response.json() == {"detail": "The user doesn't have enough privileges"}
def test_patch_user(client, active_user, logged_in_headers):
user_id = active_user.id
update_data = UserUpdate(
username="newname",
)
response = client.patch(
f"/api/v1/user/{user_id}", json=update_data.dict(), headers=logged_in_headers
)
assert response.status_code == 200, response.json()
def test_patch_user_wrong_id(client, active_user, logged_in_headers):
user_id = "wrong_id"
update_data = UserUpdate(
username="newname",
)
response = client.patch(
f"/api/v1/user/{user_id}", json=update_data.dict(), headers=logged_in_headers
)
assert response.status_code == 422, response.json()
assert response.json() == {
"detail": [
{
"loc": ["path", "user_id"],
"msg": "value is not a valid uuid",
"type": "type_error.uuid",
}
]
}
def test_delete_user(client, test_user, super_user_headers):
user_id = test_user["id"]
response = client.delete(f"/api/v1/user/{user_id}", headers=super_user_headers)
assert response.status_code == 200
assert response.json() == {"detail": "User deleted"}
def test_delete_user_wrong_id(client, test_user, super_user_headers):
user_id = "wrong_id"
response = client.delete(f"/api/v1/user/{user_id}", headers=super_user_headers)
assert response.status_code == 422
assert response.json() == {
"detail": [
{
"loc": ["path", "user_id"],
"msg": "value is not a valid uuid",
"type": "type_error.uuid",
}
]
}
def test_normal_user_cant_delete_user(client, test_user, logged_in_headers):
user_id = test_user["id"]
response = client.delete(f"/api/v1/user/{user_id}", headers=logged_in_headers)
assert response.status_code == 400
assert response.json() == {"detail": "The user doesn't have enough privileges"}
# If you still want to test the superuser endpoint
def test_add_super_user_for_testing_purposes_delete_me_before_merge_into_dev(client):
response = client.post("/api/v1/super_user")
assert response.status_code == 200
assert response.json()["username"] == "superuser"

View file

@ -4,9 +4,9 @@ from langflow.services.utils import get_settings_manager
# check that all agents are in settings.agents
# are in json_response["agents"]
def test_vectorstores_settings(client: TestClient):
def test_vectorstores_settings(client: TestClient, logged_in_headers):
settings_manager = get_settings_manager()
response = client.get("api/v1/all")
response = client.get("api/v1/all", headers=logged_in_headers)
assert response.status_code == 200
json_response = response.json()
vectorstores = json_response["vectorstores"]

View file

@ -1,13 +1,16 @@
from fastapi import WebSocketDisconnect
from fastapi.testclient import TestClient
# from langflow.services.chat.manager import ChatManager
import pytest
def test_init_build(client):
def test_init_build(client, active_user, logged_in_headers):
response = client.post(
"api/v1/build/init/test", json={"id": "test", "data": {"key": "value"}}
"api/v1/build/init/test",
json={"id": "test", "data": {"key": "value"}},
headers=logged_in_headers,
)
assert response.status_code == 201
assert response.json() == {"flowId": "test"}
@ -24,10 +27,12 @@ def test_init_build(client):
# assert response.headers["content-type"] == "text/event-stream; charset=utf-8"
def test_websocket_endpoint(client):
def test_websocket_endpoint(client: TestClient, active_user, logged_in_headers):
# Assuming your websocket_endpoint uses chat_manager which caches data from stream_build
access_token = logged_in_headers["Authorization"].split(" ")[1]
with pytest.raises(WebSocketDisconnect):
with client.websocket_connect(
"api/v1/chat/non_existing_client_id"
f"api/v1/chat/non_existing_client_id?token={access_token}"
) as websocket:
websocket.send_json({"type": "test"})
data = websocket.receive_json()