feat: ui build in one single http request (#3020)

* feat: ui build in one single http request

* fix use session_id

* fix frozen

* [autofix.ci] apply automated fixes

* prettier

* add tests

* add tests

* fix mypy

* [autofix.ci] apply automated fixes

---------

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Nicolò Boschi 2024-08-02 15:53:34 +02:00 • committed by GitHub
commit f311a6db54
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 707 additions and 46 deletions

View file

@ -1,6 +1,6 @@
import uuid
import warnings
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING, Optional, Dict
from fastapi import HTTPException
from sqlmodel import Session
@ -122,12 +122,9 @@ def format_elapsed_time(elapsed_time: float) -> str:
return f"{minutes} {minutes_unit}, {seconds} {seconds_unit}"
async def build_graph_from_db(flow_id: str, session: Session, chat_service: "ChatService"):
async def build_graph_from_data(flow_id: str, payload: Dict, **kwargs):
"""Build and cache the graph."""
flow: Optional[Flow] = session.get(Flow, flow_id)
if not flow or not flow.data:
raise ValueError("Invalid flow ID")
graph = Graph.from_payload(flow.data, flow_id, flow_name=flow.name, user_id=str(flow.user_id))
graph = Graph.from_payload(payload, flow_id, **kwargs)
for vertex_id in graph._has_session_id_vertices:
vertex = graph.get_vertex(vertex_id)
if vertex is None:
@ -139,6 +136,19 @@ async def build_graph_from_db(flow_id: str, session: Session, chat_service: "Cha
graph.set_run_id(run_id)
graph.set_run_name()
await graph.initialize_run()
return graph
async def build_graph_from_db_no_cache(flow_id: str, session: Session):
"""Build and cache the graph."""
flow: Optional[Flow] = session.get(Flow, flow_id)
if not flow or not flow.data:
raise ValueError("Invalid flow ID")
return await build_graph_from_data(flow_id, flow.data, flow_name=flow.name, user_id=str(flow.user_id))
async def build_graph_from_db(flow_id: str, session: Session, chat_service: "ChatService"):
graph = await build_graph_from_db_no_cache(flow_id, session)
await chat_service.set_cache(flow_id, graph)
return graph

View file

@ -1,11 +1,17 @@
import asyncio
import json
import time
import traceback
import typing
import uuid
from typing import TYPE_CHECKING, Annotated, Optional
from fastapi import APIRouter, BackgroundTasks, Body, Depends, HTTPException
from fastapi.responses import StreamingResponse
from loguru import logger
from starlette.background import BackgroundTask
from starlette.responses import ContentStream
from starlette.types import Receive
from langflow.api.utils import (
build_and_cache_graph_from_data,
@ -14,6 +20,8 @@ from langflow.api.utils import (
format_exception_message,
get_top_level_vertices,
parse_exception,
build_graph_from_db_no_cache,
build_graph_from_data,
)
from langflow.api.v1.schemas import (
FlowDataRequest,
@ -140,6 +148,296 @@ async def retrieve_vertices_order(
raise HTTPException(status_code=500, detail=str(exc)) from exc
@router.post("/build/{flow_id}/flow")
async def build_flow(
background_tasks: BackgroundTasks,
flow_id: uuid.UUID,
inputs: Annotated[Optional[InputValueRequest], Body(embed=True)] = None,
data: Annotated[Optional[FlowDataRequest], Body(embed=True)] = None,
files: Optional[list[str]] = None,
stop_component_id: Optional[str] = None,
start_component_id: Optional[str] = None,
chat_service: "ChatService" = Depends(get_chat_service),
current_user=Depends(get_current_active_user),
telemetry_service: "TelemetryService" = Depends(get_telemetry_service),
session=Depends(get_session),
):
async def build_graph_and_get_order() -> tuple[list[str], list[str], "Graph"]:
start_time = time.perf_counter()
components_count = None
try:
flow_id_str = str(flow_id)
if not data:
graph = await build_graph_from_db_no_cache(flow_id=flow_id_str, session=session)
else:
graph = await build_graph_from_data(flow_id_str, data.model_dump())
graph.validate_stream()
if stop_component_id or start_component_id:
try:
first_layer = graph.sort_vertices(stop_component_id, start_component_id)
except Exception as exc:
logger.error(exc)
first_layer = graph.sort_vertices()
else:
first_layer = graph.sort_vertices()
for vertex_id in first_layer:
graph.run_manager.add_to_vertices_being_run(vertex_id)
# Now vertices is a list of lists
# We need to get the id of each vertex
# and return the same structure but only with the ids
components_count = len(graph.vertices)
vertices_to_run = list(graph.vertices_to_run.union(get_top_level_vertices(graph, graph.vertices_to_run)))
background_tasks.add_task(
telemetry_service.log_package_playground,
PlaygroundPayload(
playgroundSeconds=int(time.perf_counter() - start_time),
playgroundComponentCount=components_count,
playgroundSuccess=True,
),
)
return first_layer, vertices_to_run, graph
except Exception as exc:
background_tasks.add_task(
telemetry_service.log_package_playground,
PlaygroundPayload(
playgroundSeconds=int(time.perf_counter() - start_time),
playgroundComponentCount=components_count,
playgroundSuccess=False,
playgroundErrorMessage=str(exc),
),
)
if "stream or streaming set to True" in str(exc):
raise HTTPException(status_code=400, detail=str(exc))
logger.error(f"Error checking build status: {exc}")
logger.exception(exc)
raise HTTPException(status_code=500, detail=str(exc)) from exc
async def _build_vertex(vertex_id: str, graph: "Graph") -> VertexBuildResponse:
flow_id_str = str(flow_id)
next_runnable_vertices = []
top_level_vertices = []
start_time = time.perf_counter()
error_message = None
try:
vertex = graph.get_vertex(vertex_id)
try:
lock = chat_service._async_cache_locks[flow_id_str]
(
result_dict,
params,
valid,
artifacts,
vertex,
) = await graph.build_vertex(
chat_service=None,
vertex_id=vertex_id,
user_id=current_user.id,
inputs_dict=inputs.model_dump() if inputs else {},
files=files,
)
next_runnable_vertices = await graph.get_next_runnable_vertices(lock, vertex=vertex, cache=False)
top_level_vertices = graph.get_top_level_vertices(next_runnable_vertices)
result_data_response = ResultDataResponse.model_validate(result_dict, from_attributes=True)
except Exception as exc:
if isinstance(exc, ComponentBuildException):
params = exc.message
tb = exc.formatted_traceback
else:
tb = traceback.format_exc()
logger.exception(f"Error building Component: {exc}")
params = format_exception_message(exc)
message = {"errorMessage": params, "stackTrace": tb}
valid = False
error_message = params
output_label = vertex.outputs[0]["name"] if vertex.outputs else "output"
outputs = {output_label: OutputValue(message=message, type="error")}
result_data_response = ResultDataResponse(results={}, outputs=outputs)
artifacts = {}
background_tasks.add_task(graph.end_all_traces, error=exc)
result_data_response.message = artifacts
# Log the vertex build
if not vertex.will_stream:
background_tasks.add_task(
log_vertex_build,
flow_id=flow_id_str,
vertex_id=vertex_id.split("-")[0],
valid=valid,
params=params,
data=result_data_response,
artifacts=artifacts,
)
timedelta = time.perf_counter() - start_time
duration = format_elapsed_time(timedelta)
result_data_response.duration = duration
result_data_response.timedelta = timedelta
vertex.add_build_time(timedelta)
inactivated_vertices = list(graph.inactivated_vertices)
graph.reset_inactivated_vertices()
graph.reset_activated_vertices()
# graph.stop_vertex tells us if the user asked
# to stop the build of the graph at a certain vertex
# if it is in next_vertices_ids, we need to remove other
# vertices from next_vertices_ids
if graph.stop_vertex and graph.stop_vertex in next_runnable_vertices:
next_runnable_vertices = [graph.stop_vertex]
if not graph.run_manager.vertices_being_run and not next_runnable_vertices:
background_tasks.add_task(graph.end_all_traces)
build_response = VertexBuildResponse(
inactivated_vertices=list(set(inactivated_vertices)),
next_vertices_ids=list(set(next_runnable_vertices)),
top_level_vertices=list(set(top_level_vertices)),
valid=valid,
params=params,
id=vertex.id,
data=result_data_response,
)
background_tasks.add_task(
telemetry_service.log_package_component,
ComponentPayload(
componentName=vertex_id.split("-")[0],
componentSeconds=int(time.perf_counter() - start_time),
componentSuccess=valid,
componentErrorMessage=error_message,
),
)
return build_response
except Exception as exc:
background_tasks.add_task(
telemetry_service.log_package_component,
ComponentPayload(
componentName=vertex_id.split("-")[0],
componentSeconds=int(time.perf_counter() - start_time),
componentSuccess=False,
componentErrorMessage=str(exc),
),
)
logger.error(f"Error building Component: \n\n{exc}")
logger.exception(exc)
message = parse_exception(exc)
raise HTTPException(status_code=500, detail=message) from exc
def send_event(event_type: str, value: dict, queue: asyncio.Queue) -> None:
json_data = {"event": event_type, "data": value}
event_id = uuid.uuid4()
logger.debug(f"sending event {event_id}: {event_type}")
str_data = json.dumps(json_data) + "\n\n"
queue.put_nowait((event_id, str_data.encode("utf-8"), time.time()))
async def build_vertices(
vertex_id: str, graph: "Graph", queue: asyncio.Queue, client_consumed_queue: asyncio.Queue
) -> None:
build_task = asyncio.create_task(await asyncio.to_thread(_build_vertex, vertex_id, graph))
try:
await build_task
except asyncio.CancelledError:
build_task.cancel()
return
vertex_build_response: VertexBuildResponse = build_task.result()
# send built event or error event
send_event("end_vertex", {"build_data": json.loads(vertex_build_response.model_dump_json())}, queue)
await client_consumed_queue.get()
if vertex_build_response.valid:
if vertex_build_response.next_vertices_ids:
tasks = []
for next_vertex_id in vertex_build_response.next_vertices_ids:
task = asyncio.create_task(build_vertices(next_vertex_id, graph, queue, client_consumed_queue))
tasks.append(task)
try:
await asyncio.gather(*tasks)
except asyncio.CancelledError:
for task in tasks:
task.cancel()
return
async def event_generator(queue: asyncio.Queue, client_consumed_queue: asyncio.Queue) -> None:
if not data:
# using another thread since the DB query is I/O bound
vertices_task = asyncio.create_task(await asyncio.to_thread(build_graph_and_get_order))
try:
await vertices_task
except asyncio.CancelledError:
vertices_task.cancel()
return
ids, vertices_to_run, graph = vertices_task.result()
else:
ids, vertices_to_run, graph = await build_graph_and_get_order()
send_event("vertices_sorted", {"ids": ids, "to_run": vertices_to_run}, queue)
await client_consumed_queue.get()
tasks = []
for vertex_id in ids:
task = asyncio.create_task(build_vertices(vertex_id, graph, queue, client_consumed_queue))
tasks.append(task)
try:
await asyncio.gather(*tasks)
except asyncio.CancelledError:
for task in tasks:
task.cancel()
return
send_event("end", {}, queue)
await queue.put((None, None, time.time))
async def consume_and_yield(queue: asyncio.Queue, client_consumed_queue: asyncio.Queue) -> typing.AsyncGenerator:
while True:
event_id, value, put_time = await queue.get()
if value is None:
break
get_time = time.time()
yield value
get_time_yield = time.time()
client_consumed_queue.put_nowait(event_id)
logger.debug(
f"consumed event {str(event_id)} (time in queue, {get_time - put_time:.4f}, client {get_time_yield - get_time:.4f})"
)
asyncio_queue: asyncio.Queue = asyncio.Queue()
asyncio_queue_client_consumed: asyncio.Queue = asyncio.Queue()
main_task = asyncio.create_task(event_generator(asyncio_queue, asyncio_queue_client_consumed))
def on_disconnect():
logger.debug("Client disconnected, closing tasks")
main_task.cancel()
return DisconnectHandlerStreamingResponse(
consume_and_yield(asyncio_queue, asyncio_queue_client_consumed),
media_type="application/x-ndjson",
on_disconnect=on_disconnect,
)
class DisconnectHandlerStreamingResponse(StreamingResponse):
def __init__(
self,
content: ContentStream,
status_code: int = 200,
headers: typing.Mapping[str, str] | None = None,
media_type: str | None = None,
background: BackgroundTask | None = None,
on_disconnect: Optional[typing.Callable] = None,
):
super().__init__(content, status_code, headers, media_type, background)
self.on_disconnect = on_disconnect
async def listen_for_disconnect(self, receive: Receive) -> None:
while True:
message = await receive()
if message["type"] == "http.disconnect":
if self.on_disconnect:
await self.on_disconnect()
break
@router.post("/build/{flow_id}/vertices/{vertex_id}")
async def build_vertex(
flow_id: uuid.UUID,

View file

@ -851,7 +851,7 @@ class Graph:
async def build_vertex(
self,
chat_service: ChatService,
chat_service: Optional[ChatService],
vertex_id: str,
inputs_dict: Optional[Dict[str, str]] = None,
files: Optional[list[str]] = None,
@ -880,14 +880,11 @@ class Graph:
try:
params = ""
if vertex.frozen:
# Check the cache for the vertex
cached_result = await chat_service.get_cache(key=vertex.id)
if isinstance(cached_result, CacheMiss):
await vertex.build(
user_id=user_id, inputs=inputs_dict, fallback_to_env_vars=fallback_to_env_vars, files=files
)
await chat_service.set_cache(key=vertex.id, data=vertex)
if chat_service:
cached_result = await chat_service.get_cache(key=vertex.id)
else:
cached_result = None
if cached_result and not isinstance(cached_result, CacheMiss):
cached_vertex = cached_result["result"]
# Now set update the vertex with the cached vertex
vertex._built = cached_vertex._built
@ -898,12 +895,18 @@ class Graph:
vertex._custom_component = cached_vertex._custom_component
if vertex.result is not None:
vertex.result.used_frozen_result = True
else:
await vertex.build(
user_id=user_id, inputs=inputs_dict, fallback_to_env_vars=fallback_to_env_vars, files=files
)
if chat_service:
await chat_service.set_cache(key=vertex.id, data=vertex)
else:
await vertex.build(
user_id=user_id, inputs=inputs_dict, fallback_to_env_vars=fallback_to_env_vars, files=files
)
await chat_service.set_cache(key=vertex.id, data=vertex)
if chat_service:
await chat_service.set_cache(key=vertex.id, data=vertex)
if vertex.result is not None:
params = f"{vertex._built_object_repr()}{params}"

View file

@ -34,6 +34,7 @@ class TelemetryService(Service):
self.telemetry_queue: asyncio.Queue = asyncio.Queue()
self.client = httpx.AsyncClient(timeout=10.0) # Set a reasonable timeout
self.running = False
self._stopping = False
self.ot = OpenTelemetry(prometheus_enabled=settings_service.settings.prometheus_enabled)
@ -75,11 +76,16 @@ class TelemetryService(Service):
logger.error(f"Unexpected error occurred: {e}")
async def log_package_run(self, payload: RunPayload):
await self.telemetry_queue.put((self.send_telemetry_data, payload, "run"))
await self._queue_event((self.send_telemetry_data, payload, "run"))
async def log_package_shutdown(self):
payload = ShutdownPayload(timeRunning=(datetime.now(timezone.utc) - self._start_time).seconds)
await self.telemetry_queue.put((self.send_telemetry_data, payload, "shutdown"))
await self._queue_event(payload)
async def _queue_event(self, payload):
if self.do_not_track or self._stopping:
return
await self.telemetry_queue.put(payload)
async def log_package_version(self):
python_version = ".".join(platform.python_version().split(".")[:2])
@ -95,13 +101,13 @@ class TelemetryService(Service):
arch=architecture,
autoLogin=self.settings_service.auth_settings.AUTO_LOGIN,
)
await self.telemetry_queue.put((self.send_telemetry_data, payload, None))
await self._queue_event((self.send_telemetry_data, payload, None))
async def log_package_playground(self, payload: PlaygroundPayload):
await self.telemetry_queue.put((self.send_telemetry_data, payload, "playground"))
await self._queue_event((self.send_telemetry_data, payload, "playground"))
async def log_package_component(self, payload: ComponentPayload):
await self.telemetry_queue.put((self.send_telemetry_data, payload, "component"))
await self._queue_event((self.send_telemetry_data, payload, "component"))
async def start(self):
if self.running or self.do_not_track:
@ -123,11 +129,13 @@ class TelemetryService(Service):
logger.error(f"Error flushing logs: {e}")
async def stop(self):
if self.do_not_track:
if self.do_not_track or self._stopping:
return
try:
self.running = False
self._stopping = True
# flush all the remaining events and then stop
await self.flush()
self.running = False
if self.worker_task:
self.worker_task.cancel()
with contextlib.suppress(asyncio.CancelledError):

View file

@ -53,6 +53,7 @@ def pytest_configure(config):
pytest.TWO_OUTPUTS = data_path / "TwoOutputsTest.json"
pytest.VECTOR_STORE_PATH = data_path / "Vector_store.json"
pytest.SIMPLE_API_TEST = data_path / "SimpleAPITest.json"
pytest.MEMORY_CHATBOT_NO_LLM = data_path / "MemoryChatbotNoLLM.json"
pytest.CODE_WITH_SYNTAX_ERROR = """
def get_text():
retun "Hello World"
@ -70,6 +71,7 @@ def get_text():
pytest.CHAT_INPUT,
pytest.TWO_OUTPUTS,
pytest.VECTOR_STORE_PATH,
pytest.MEMORY_CHATBOT_NO_LLM,
]:
assert path.exists(), f"File {path} does not exist. Available files: {list(data_path.iterdir())}"
@ -232,6 +234,12 @@ def json_webhook_test():
return f.read()
@pytest.fixture
def json_memory_chatbot_no_llm():
with open(pytest.MEMORY_CHATBOT_NO_LLM, "r") as f:
return f.read()
@pytest.fixture(name="client", autouse=True)
def client_fixture(session: Session, monkeypatch, request, load_flows_dir):
# Set the database url to a test database

File diff suppressed because one or more lines are too long

View file

@ -0,0 +1,93 @@
import json
from uuid import UUID
from orjson import orjson
from langflow.memory import get_messages
from langflow.services.database.models.flow import FlowCreate, FlowUpdate
def test_build_flow(client, json_memory_chatbot_no_llm, logged_in_headers):
flow_id = _create_flow(client, json_memory_chatbot_no_llm, logged_in_headers)
with client.stream("POST", f"api/v1/build/{flow_id}/flow", json={}, headers=logged_in_headers) as r:
consume_and_assert_stream(r)
check_messages(flow_id)
def test_build_flow_from_request_data(client, json_memory_chatbot_no_llm, logged_in_headers):
flow_id = _create_flow(client, json_memory_chatbot_no_llm, logged_in_headers)
flow_data = client.get("api/v1/flows/" + str(flow_id), headers=logged_in_headers).json()
with client.stream(
"POST", f"api/v1/build/{flow_id}/flow", json={"data": flow_data["data"]}, headers=logged_in_headers
) as r:
consume_and_assert_stream(r)
check_messages(flow_id)
def test_build_flow_with_frozen_path(client, json_memory_chatbot_no_llm, logged_in_headers):
flow_id = _create_flow(client, json_memory_chatbot_no_llm, logged_in_headers)
flow_data = client.get("api/v1/flows/" + str(flow_id), headers=logged_in_headers).json()
flow_data["data"]["nodes"][0]["data"]["node"]["frozen"] = True
response = client.patch(
"api/v1/flows/" + str(flow_id),
json=FlowUpdate(name="Flow", description="description", data=flow_data["data"]).model_dump(),
headers=logged_in_headers,
)
response.raise_for_status()
with client.stream("POST", f"api/v1/build/{flow_id}/flow", json={}, headers=logged_in_headers) as r:
consume_and_assert_stream(r)
check_messages(flow_id)
def check_messages(flow_id):
messages = get_messages(flow_id=UUID(flow_id), order="ASC")
assert len(messages) == 2
assert messages[0].session_id == flow_id
assert messages[0].sender == "User"
assert messages[0].sender_name == "User"
assert messages[0].text == ""
assert messages[1].session_id == flow_id
assert messages[1].sender == "Machine"
assert messages[1].sender_name == "AI"
def consume_and_assert_stream(r):
count = 0
for line in r.iter_lines():
# httpx split by \n, but ndjson sends two \n for each line
if not line:
continue
parsed = json.loads(line)
if count == 0:
assert parsed["event"] == "vertices_sorted"
ids = parsed["data"]["ids"]
ids.sort()
assert ids == ["ChatInput-CIGht", "Memory-amN4Z"]
to_run = parsed["data"]["to_run"]
to_run.sort()
assert to_run == ["ChatInput-CIGht", "ChatOutput-QA7ej", "Memory-amN4Z", "Prompt-iWbCC"]
elif count > 0 and count < 5:
assert parsed["event"] == "end_vertex"
assert parsed["data"]["build_data"] is not None
elif count == 5:
assert parsed["event"] == "end"
else:
raise ValueError(f"Unexpected line: {line}")
count += 1
def _create_flow(client, json_memory_chatbot_no_llm, logged_in_headers):
vector_store = orjson.loads(json_memory_chatbot_no_llm)
data = vector_store["data"]
vector_store = FlowCreate(name="Flow", description="description", data=data, endpoint_name="f")
response = client.post("api/v1/flows/", json=vector_store.model_dump(), headers=logged_in_headers)
response.raise_for_status()
flow_id = response.json()["id"]
return flow_id

View file

@ -16,11 +16,12 @@ const api: AxiosInstance = axios.create({
baseURL: "",
});
const cookies = new Cookies();
function ApiInterceptor() {
const autoLogin = useAuthStore((state) => state.autoLogin);
const setErrorData = useAlertStore((state) => state.setErrorData);
let { accessToken, authenticationErrorCount } = useContext(AuthContext);
const cookies = new Cookies();
const setSaveLoading = useFlowsManagerStore((state) => state.setSaveLoading);
const { mutate: mutationLogout } = useLogout();
const { mutate: mutationRenewAccessToken } = useRefreshAccessToken();
@ -205,4 +206,83 @@ function ApiInterceptor() {
return null;
}
export { ApiInterceptor, api };
export type StreamingRequestParams = {
method: string;
url: string;
onData: (event: object) => Promise<boolean>;
body?: object;
onError?: (statusCode: number) => void;
};
async function performStreamingRequest({
method,
url,
onData,
body,
onError,
}: StreamingRequestParams) {
let headers = {
"Content-Type": "application/json",
// this flag is fundamental to ensure server stops tasks when client disconnects
Connection: "close",
};
const accessToken = cookies.get(LANGFLOW_ACCESS_TOKEN);
if (accessToken) {
headers["Authorization"] = `Bearer ${accessToken}`;
}
const controller = new AbortController();
const params = {
method: method,
headers: headers,
signal: controller.signal,
};
if (body) {
params["body"] = JSON.stringify(body);
}
let current: string[] = [];
let textDecoder = new TextDecoder();
const response = await fetch(url, params);
if (!response.ok) {
if (onError) {
onError(response.status);
} else {
throw new Error("error in streaming request");
}
}
if (response.body === null) {
return;
}
for await (const chunk of response.body) {
const decodedChunk = await textDecoder.decode(chunk);
let all = decodedChunk.split("\n\n");
for (const string of all) {
if (string.endsWith("}")) {
const allString = current.join("") + string;
let data: object;
try {
data = JSON.parse(allString);
current = [];
} catch (e) {
current.push(string);
continue;
}
const shouldContinue = await onData(data);
if (!shouldContinue) {
controller.abort();
return;
}
} else {
current.push(string);
}
}
}
if (current.length > 0) {
const allString = current.join("");
if (allString) {
const data = JSON.parse(current.join(""));
await onData(data);
}
}
}
export { ApiInterceptor, api, performStreamingRequest };

View file

@ -38,7 +38,6 @@ export default function ChatMessage({
const [chatMessage, setChatMessage] = useState(chatMessageString);
const [isStreaming, setIsStreaming] = useState(false);
const eventSource = useRef<EventSource | undefined>(undefined);
const updateFlowPool = useFlowStore((state) => state.updateFlowPool);
const setErrorData = useAlertStore((state) => state.setErrorData);
const chatMessageRef = useRef(chatMessage);

View file

@ -29,7 +29,7 @@ import {
targetHandleType,
} from "../types/flow";
import { FlowStoreType, VertexLayerElementType } from "../types/zustand/flow";
import { buildVertices } from "../utils/buildUtils";
import { buildFlowVerticesWithFallback } from "../utils/buildUtils";
import {
checkChatInput,
checkOldComponents,
@ -607,20 +607,8 @@ const useFlowStore = create<FlowStoreType>((set, get) => ({
);
useFlowStore.getState().updateBuildStatus([vertexBuildData.id], status);
const verticesIds = get().verticesBuild?.verticesIds;
const newFlowBuildStatus = { ...get().flowBuildStatus };
// filter out the vertices that are not status
const verticesToUpdate = verticesIds?.filter(
(id) => newFlowBuildStatus[id]?.status !== BuildStatus.BUILT,
);
if (verticesToUpdate) {
useFlowStore.getState().updateBuildStatus(verticesToUpdate, status);
}
}
await buildVertices({
await buildFlowVerticesWithFallback({
input_value,
files,
flowId: currentFlow!.id,
@ -672,8 +660,8 @@ const useFlowStore = create<FlowStoreType>((set, get) => ({
useFlowStore.getState().updateBuildStatus(idList, BuildStatus.BUILDING);
},
onValidateNodes: validateSubgraph,
nodes: !get().onFlowPage ? get().nodes : undefined,
edges: !get().onFlowPage ? get().edges : undefined,
nodes: get().onFlowPage ? get().nodes : undefined,
edges: get().onFlowPage ? get().edges : undefined,
});
get().setIsBuilding(false);
get().setLockChat(false);
@ -690,7 +678,7 @@ const useFlowStore = create<FlowStoreType>((set, get) => ({
vertices: {
verticesIds: string[];
verticesLayers: VertexLayerElementType[][];
runId: string;
runId?: string;
verticesToRun: string[];
} | null,
) => {

View file

@ -147,7 +147,7 @@ export type FlowStoreType = {
vertices: {
verticesIds: string[];
verticesLayers: VertexLayerElementType[][];
runId: string;
runId?: string;
verticesToRun: string[];
} | null,
) => void;
@ -156,7 +156,7 @@ export type FlowStoreType = {
verticesBuild: {
verticesIds: string[];
verticesLayers: VertexLayerElementType[][];
runId: string;
runId?: string;
verticesToRun: string[];
} | null;
updateBuildStatus: (nodeId: string[], status: BuildStatus) => void;

View file

@ -1,3 +1,5 @@
import { BASE_URL_API } from "@/constants/constants";
import { performStreamingRequest } from "@/controllers/API/api";
import { AxiosError } from "axios";
import { Edge, Node } from "reactflow";
import { BuildStatus } from "../constants/enums";
@ -66,7 +68,7 @@ export async function updateVerticesOrder(
): Promise<{
verticesLayers: VertexLayerElementType[][];
verticesIds: string[];
runId: string;
runId?: string;
verticesToRun: string[];
}> {
return new Promise(async (resolve, reject) => {
@ -115,6 +117,176 @@ export async function updateVerticesOrder(
});
}
export async function buildFlowVerticesWithFallback(
params: BuildVerticesParams,
) {
try {
return await buildFlowVertices(params);
} catch (e: any) {
if (e.message === "endpoint not available") {
return await buildVertices(params);
}
throw e;
}
}
const MIN_VISUAL_BUILD_TIME_MS = 300;
export async function buildFlowVertices({
flowId,
input_value,
files,
startNodeId,
stopNodeId,
onGetOrderSuccess,
onBuildUpdate,
onBuildComplete,
onBuildError,
onBuildStart,
onValidateNodes,
nodes,
edges,
setLockChat,
}: BuildVerticesParams) {
let url = `${BASE_URL_API}build/${flowId}/flow?`;
if (startNodeId) {
url = `${url}&start_component_id=${startNodeId}`;
}
if (stopNodeId) {
url = `${url}&stop_component_id=${stopNodeId}`;
}
const postData = {};
if (typeof input_value !== "undefined") {
postData["inputs"] = { input_value: input_value };
}
if (files) {
postData["files"] = files;
}
if (nodes) {
postData["data"] = {
nodes,
edges,
};
}
const buildResults: Array<boolean> = [];
const verticesStartTimeMs: Map<string, number> = new Map();
const onEvent = async (type, data): Promise<boolean> => {
const onStartVertices = (ids: Array<string>) => {
useFlowStore.getState().updateBuildStatus(ids, BuildStatus.TO_BUILD);
if (onBuildStart)
onBuildStart(ids.map((id) => ({ id: id, reference: id })));
ids.forEach((id) => verticesStartTimeMs.set(id, Date.now()));
};
switch (type) {
case "vertices_sorted": {
const verticesToRun = data.to_run;
const verticesIds = data.ids;
onStartVertices(verticesIds);
let verticesLayers: Array<Array<VertexLayerElementType>> =
verticesIds.map((id: string) => {
return [{ id: id, reference: id }];
});
useFlowStore.getState().updateVerticesBuild({
verticesLayers,
verticesIds,
verticesToRun,
});
if (onValidateNodes) {
try {
onValidateNodes(data.to_run);
if (onGetOrderSuccess) onGetOrderSuccess();
useFlowStore.getState().setIsBuilding(true);
return true;
} catch (e) {
useFlowStore.getState().setIsBuilding(false);
setLockChat && setLockChat(false);
return false;
}
}
return true;
}
case "end_vertex": {
const buildData = data.build_data;
const startTimeMs = verticesStartTimeMs.get(buildData.id);
if (startTimeMs) {
const delta = Date.now() - startTimeMs;
if (delta < MIN_VISUAL_BUILD_TIME_MS) {
// this is a visual trick to make the build process look more natural
await new Promise((resolve) =>
setTimeout(resolve, MIN_VISUAL_BUILD_TIME_MS - delta),
);
}
}
if (onBuildUpdate) {
if (!buildData.valid) {
// lots is a dictionary with the key the output field name and the value the log object
// logs: { [key: string]: { message: any; type: string }[] };
const errorMessages = Object.keys(buildData.data.outputs).map(
(key) => {
const outputs = buildData.data.outputs[key];
if (Array.isArray(outputs)) {
return outputs
.filter((log) => isErrorLogType(log.message))
.map((log) => log.message.errorMessage);
}
if (!isErrorLogType(outputs.message)) {
return [];
}
return [outputs.message.errorMessage];
},
);
onBuildError!("Error Building Component", errorMessages, [
{ id: buildData.id },
]);
onBuildUpdate(buildData, BuildStatus.ERROR, "");
buildResults.push(false);
return false;
} else {
onBuildUpdate(buildData, BuildStatus.BUILT, "");
buildResults.push(true);
}
}
if (buildData.next_vertices_ids) {
onStartVertices(buildData.next_vertices_ids);
}
return true;
}
case "end": {
const allNodesValid = buildResults.every((result) => result);
onBuildComplete!(allNodesValid);
useFlowStore.getState().setIsBuilding(false);
return true;
}
default:
return true;
}
return true;
};
return performStreamingRequest({
method: "POST",
url,
body: postData,
onData: async (event) => {
const type = event["event"];
const data = event["data"];
return await onEvent(type, data);
},
onError: (statusCode) => {
if (statusCode === 404) {
throw new Error("endpoint not available");
}
throw new Error("error in streaming request");
},
});
}
export async function buildVertices({
flowId,
input_value,
@ -252,6 +424,7 @@ export async function buildVertices({
useFlowStore.getState().setIsBuilding(false);
}
}
async function buildVertex({
flowId,
id,