Refactor API schemas and update dependencies

This commit is contained in:
Gabriel Luiz Freitas Almeida 2024-02-27 11:36:58 -03:00
commit b7e52f62be
3 changed files with 28 additions and 22 deletions

View file

@ -4,12 +4,12 @@ from pathlib import Path
from typing import Any, Dict, List, Optional, Union from typing import Any, Dict, List, Optional, Union
from uuid import UUID from uuid import UUID
from langflow.api.utils import serialize_field from pydantic import BaseModel, Field, field_validator
from langflow.services.database.models.api_key.model import ApiKeyRead from langflow.services.database.models.api_key.model import ApiKeyRead
from langflow.services.database.models.base import orjson_dumps from langflow.services.database.models.base import orjson_dumps
from langflow.services.database.models.flow import FlowCreate, FlowRead from langflow.services.database.models.flow import FlowCreate, FlowRead
from langflow.services.database.models.user import UserRead from langflow.services.database.models.user import UserRead
from pydantic import BaseModel, Field, field_serializer, field_validator
class BuildStatus(Enum): class BuildStatus(Enum):
@ -161,7 +161,9 @@ class StreamData(BaseModel):
data: dict data: dict
def __str__(self) -> str: def __str__(self) -> str:
return f"event: {self.event}\ndata: {orjson_dumps(self.data, indent_2=False)}\n\n" return (
f"event: {self.event}\ndata: {orjson_dumps(self.data, indent_2=False)}\n\n"
)
class CustomComponentCode(BaseModel): class CustomComponentCode(BaseModel):
@ -220,18 +222,12 @@ class VerticesOrderResponse(BaseModel):
ids: List[List[str]] ids: List[List[str]]
class ResultData(BaseModel): class ResultDataResponse(BaseModel):
results: Optional[Any] = Field(default_factory=dict) results: Optional[Any] = Field(default_factory=dict)
artifacts: Optional[Any] = Field(default_factory=dict) artifacts: Optional[Any] = Field(default_factory=dict)
timedelta: Optional[float] = None timedelta: Optional[float] = None
duration: Optional[str] = None duration: Optional[str] = None
@field_serializer("results")
def serialize_results(self, value):
if isinstance(value, dict):
return {key: serialize_field(val) for key, val in value.items()}
return serialize_field(value)
class VertexBuildResponse(BaseModel): class VertexBuildResponse(BaseModel):
id: Optional[str] = None id: Optional[str] = None
@ -239,7 +235,7 @@ class VertexBuildResponse(BaseModel):
valid: bool valid: bool
params: Optional[str] params: Optional[str]
"""JSON string of the params.""" """JSON string of the params."""
data: ResultData data: ResultDataResponse
"""Mapping of vertex ids to result dict containing the param name and result value.""" """Mapping of vertex ids to result dict containing the param name and result value."""
timestamp: Optional[datetime] = Field(default_factory=datetime.utcnow) timestamp: Optional[datetime] = Field(default_factory=datetime.utcnow)
"""Timestamp of the build.""" """Timestamp of the build."""

View file

@ -1,12 +1,13 @@
from typing import TYPE_CHECKING, Any, Dict, Optional, Type, Union from typing import TYPE_CHECKING, Any, Dict, Optional, Type, Union
import duckdb import duckdb
from langflow.services.deps import get_monitor_service
from loguru import logger from loguru import logger
from pydantic import BaseModel from pydantic import BaseModel
from langflow.services.deps import get_monitor_service
if TYPE_CHECKING: if TYPE_CHECKING:
from langflow.api.v1.schemas import ResultData from langflow.api.v1.schemas import ResultDataResponse
INDEX_KEY = "index" INDEX_KEY = "index"
@ -45,7 +46,9 @@ def model_to_sql_column_definitions(model: Type[BaseModel]) -> dict:
return columns return columns
def drop_and_create_table_if_schema_mismatch(db_path: str, table_name: str, model: Type[BaseModel]): def drop_and_create_table_if_schema_mismatch(
db_path: str, table_name: str, model: Type[BaseModel]
):
with duckdb.connect(db_path) as conn: with duckdb.connect(db_path) as conn:
# Get the current schema from the database # Get the current schema from the database
try: try:
@ -66,8 +69,12 @@ def drop_and_create_table_if_schema_mismatch(db_path: str, table_name: str, mode
conn.execute(f"CREATE SEQUENCE seq_{table_name} START 1;") conn.execute(f"CREATE SEQUENCE seq_{table_name} START 1;")
except duckdb.CatalogException: except duckdb.CatalogException:
pass pass
desired_schema[INDEX_KEY] = f"INTEGER PRIMARY KEY DEFAULT NEXTVAL('seq_{table_name}')" desired_schema[INDEX_KEY] = (
columns_sql = ", ".join(f"{name} {data_type}" for name, data_type in desired_schema.items()) f"INTEGER PRIMARY KEY DEFAULT NEXTVAL('seq_{table_name}')"
)
columns_sql = ", ".join(
f"{name} {data_type}" for name, data_type in desired_schema.items()
)
create_table_sql = f"CREATE TABLE {table_name} ({columns_sql})" create_table_sql = f"CREATE TABLE {table_name} ({columns_sql})"
conn.execute(create_table_sql) conn.execute(create_table_sql)
@ -138,7 +145,7 @@ async def log_vertex_build(
vertex_id: str, vertex_id: str,
valid: bool, valid: bool,
params: Any, params: Any,
data: "ResultData", data: "ResultDataResponse",
artifacts: Optional[dict] = None, artifacts: Optional[dict] = None,
): ):
try: try:

View file

@ -2,14 +2,15 @@ import time
from typing import Callable from typing import Callable
import socketio import socketio
from sqlmodel import select
from langflow.api.utils import format_elapsed_time from langflow.api.utils import format_elapsed_time
from langflow.api.v1.schemas import ResultData, VertexBuildResponse from langflow.api.v1.schemas import ResultDataResponse, VertexBuildResponse
from langflow.graph.graph.base import Graph from langflow.graph.graph.base import Graph
from langflow.graph.vertex.base import StatelessVertex from langflow.graph.vertex.base import StatelessVertex
from langflow.services.database.models.flow.model import Flow from langflow.services.database.models.flow.model import Flow
from langflow.services.deps import get_session from langflow.services.deps import get_session
from langflow.services.monitor.utils import log_vertex_build from langflow.services.monitor.utils import log_vertex_build
from sqlmodel import select
def set_socketio_server(socketio_server): def set_socketio_server(socketio_server):
@ -73,7 +74,7 @@ async def build_vertex(
artifacts = vertex.artifacts artifacts = vertex.artifacts
timedelta = time.perf_counter() - start_time timedelta = time.perf_counter() - start_time
duration = format_elapsed_time(timedelta) duration = format_elapsed_time(timedelta)
result_dict = ResultData( result_dict = ResultDataResponse(
results=result_dict, results=result_dict,
artifacts=artifacts, artifacts=artifacts,
duration=duration, duration=duration,
@ -82,7 +83,7 @@ async def build_vertex(
except Exception as exc: except Exception as exc:
params = str(exc) params = str(exc)
valid = False valid = False
result_dict = ResultData(results={}) result_dict = ResultDataResponse(results={})
artifacts = {} artifacts = {}
set_cache(flow_id, graph) set_cache(flow_id, graph)
await log_vertex_build( await log_vertex_build(
@ -95,7 +96,9 @@ async def build_vertex(
) )
# Emit the vertex build response # Emit the vertex build response
response = VertexBuildResponse(valid=valid, params=params, id=vertex.id, data=result_dict) response = VertexBuildResponse(
valid=valid, params=params, id=vertex.id, data=result_dict
)
await sio.emit("vertex_build", data=response.model_dump(), to=sid) await sio.emit("vertex_build", data=response.model_dump(), to=sid)
except Exception as exc: except Exception as exc: