Add experimental preload endpoint (#1331)
This endpoint allows the Graph to be kept preprocessed in a session.
This commit is contained in:
commit
2d087c6ce2
7 changed files with 171 additions and 46 deletions
|
|
@ -1,3 +1,4 @@
|
||||||
|
import multiprocessing
|
||||||
import platform
|
import platform
|
||||||
import socket
|
import socket
|
||||||
import sys
|
import sys
|
||||||
|
|
@ -9,18 +10,16 @@ from typing import Optional
|
||||||
import httpx
|
import httpx
|
||||||
import typer
|
import typer
|
||||||
from dotenv import load_dotenv
|
from dotenv import load_dotenv
|
||||||
from multiprocess import Process, cpu_count # type: ignore
|
from multiprocess import cpu_count # type: ignore
|
||||||
from rich import box
|
from rich import box
|
||||||
from rich import print as rprint
|
from rich import print as rprint
|
||||||
from rich.console import Console
|
from rich.console import Console
|
||||||
from rich.panel import Panel
|
from rich.panel import Panel
|
||||||
from rich.table import Table
|
from rich.table import Table
|
||||||
from sqlmodel import select
|
|
||||||
|
|
||||||
from langflow.main import setup_app
|
from langflow.main import setup_app
|
||||||
from langflow.services.database.utils import session_getter
|
from langflow.services.deps import get_settings_service
|
||||||
from langflow.services.deps import get_db_service, get_settings_service
|
from langflow.services.utils import initialize_settings_service
|
||||||
from langflow.services.utils import initialize_services, initialize_settings_service
|
|
||||||
from langflow.utils.logger import configure, logger
|
from langflow.utils.logger import configure, logger
|
||||||
|
|
||||||
console = Console()
|
console = Console()
|
||||||
|
|
@ -216,7 +215,8 @@ def run(
|
||||||
|
|
||||||
|
|
||||||
def run_on_mac_or_linux(host, port, log_level, options, app, open_browser=True):
|
def run_on_mac_or_linux(host, port, log_level, options, app, open_browser=True):
|
||||||
webapp_process = Process(target=run_langflow, args=(host, port, log_level, options, app))
|
ctx = multiprocessing.get_context("spawn")
|
||||||
|
webapp_process = ctx.Process(target=run_langflow, args=(host, port, log_level, options))
|
||||||
webapp_process.start()
|
webapp_process.start()
|
||||||
status_code = 0
|
status_code = 0
|
||||||
while status_code != 200:
|
while status_code != 200:
|
||||||
|
|
@ -229,6 +229,7 @@ def run_on_mac_or_linux(host, port, log_level, options, app, open_browser=True):
|
||||||
print_banner(host, port)
|
print_banner(host, port)
|
||||||
if open_browser:
|
if open_browser:
|
||||||
webbrowser.open(f"http://{host}:{port}")
|
webbrowser.open(f"http://{host}:{port}")
|
||||||
|
webapp_process.join()
|
||||||
|
|
||||||
|
|
||||||
def run_on_windows(host, port, log_level, options, app):
|
def run_on_windows(host, port, log_level, options, app):
|
||||||
|
|
@ -298,24 +299,33 @@ def print_banner(host, port):
|
||||||
rprint(panel)
|
rprint(panel)
|
||||||
|
|
||||||
|
|
||||||
def run_langflow(host, port, log_level, options, app):
|
def run_langflow(host, port, log_level, options):
|
||||||
"""
|
"""
|
||||||
Run Langflow server on localhost
|
Run Langflow server on localhost
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
if platform.system() in ["Windows"]:
|
if platform.system() in ["Windows", "Darwin"]:
|
||||||
# Run using uvicorn on MacOS and Windows
|
# Run using uvicorn on MacOS and Windows
|
||||||
# Windows doesn't support gunicorn
|
# Windows doesn't support gunicorn
|
||||||
# MacOS requires an env variable to be set to use gunicorn
|
# MacOS requires an env variable to be set to use gunicorn
|
||||||
|
|
||||||
import uvicorn
|
import uvicorn
|
||||||
|
|
||||||
uvicorn.run(app, host=host, port=port, log_level=log_level)
|
uvicorn.run(
|
||||||
|
"langflow.main:create_app",
|
||||||
|
factory=True,
|
||||||
|
host=host,
|
||||||
|
port=port,
|
||||||
|
log_level=log_level,
|
||||||
|
workers=options["workers"],
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
from langflow.server import LangflowApplication
|
from langflow.server import LangflowApplication
|
||||||
|
|
||||||
LangflowApplication(app, options).run()
|
LangflowApplication(options).run()
|
||||||
except KeyboardInterrupt:
|
except KeyboardInterrupt:
|
||||||
pass
|
logger.info("Shutting down server")
|
||||||
|
sys.exit(0)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.exception(e)
|
logger.exception(e)
|
||||||
sys.exit(1)
|
sys.exit(1)
|
||||||
|
|
|
||||||
|
|
@ -3,12 +3,10 @@ from typing import Annotated, Any, List, Optional, Union
|
||||||
|
|
||||||
import sqlalchemy as sa
|
import sqlalchemy as sa
|
||||||
from fastapi import APIRouter, Body, Depends, HTTPException, UploadFile, status
|
from fastapi import APIRouter, Body, Depends, HTTPException, UploadFile, status
|
||||||
from loguru import logger
|
|
||||||
from sqlmodel import select
|
|
||||||
|
|
||||||
from langflow.api.utils import update_frontend_node_with_template_values
|
from langflow.api.utils import update_frontend_node_with_template_values
|
||||||
from langflow.api.v1.schemas import (
|
from langflow.api.v1.schemas import (
|
||||||
CustomComponentCode,
|
CustomComponentCode,
|
||||||
|
PreloadResponse,
|
||||||
ProcessResponse,
|
ProcessResponse,
|
||||||
TaskResponse,
|
TaskResponse,
|
||||||
TaskStatusResponse,
|
TaskStatusResponse,
|
||||||
|
|
@ -17,12 +15,15 @@ from langflow.api.v1.schemas import (
|
||||||
from langflow.interface.custom.custom_component import CustomComponent
|
from langflow.interface.custom.custom_component import CustomComponent
|
||||||
from langflow.interface.custom.directory_reader import DirectoryReader
|
from langflow.interface.custom.directory_reader import DirectoryReader
|
||||||
from langflow.interface.custom.utils import build_custom_component_template
|
from langflow.interface.custom.utils import build_custom_component_template
|
||||||
from langflow.processing.process import process_graph_cached, process_tweaks
|
from langflow.processing.process import build_graph_and_generate_result, process_graph_cached, process_tweaks
|
||||||
from langflow.services.auth.utils import api_key_security, get_current_active_user
|
from langflow.services.auth.utils import api_key_security, get_current_active_user
|
||||||
from langflow.services.cache.utils import save_uploaded_file
|
from langflow.services.cache.utils import save_uploaded_file
|
||||||
from langflow.services.database.models.flow import Flow
|
from langflow.services.database.models.flow import Flow
|
||||||
from langflow.services.database.models.user.model import User
|
from langflow.services.database.models.user.model import User
|
||||||
from langflow.services.deps import get_session, get_session_service, get_settings_service, get_task_service
|
from langflow.services.deps import get_session, get_session_service, get_settings_service, get_task_service
|
||||||
|
from langflow.services.session.service import SessionService
|
||||||
|
from loguru import logger
|
||||||
|
from sqlmodel import select
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from langflow.worker import process_graph_cached_task
|
from langflow.worker import process_graph_cached_task
|
||||||
|
|
@ -32,9 +33,8 @@ except ImportError:
|
||||||
raise NotImplementedError("Celery is not installed")
|
raise NotImplementedError("Celery is not installed")
|
||||||
|
|
||||||
|
|
||||||
from sqlmodel import Session
|
|
||||||
|
|
||||||
from langflow.services.task.service import TaskService
|
from langflow.services.task.service import TaskService
|
||||||
|
from sqlmodel import Session
|
||||||
|
|
||||||
# build router
|
# build router
|
||||||
router = APIRouter(tags=["Base"])
|
router = APIRouter(tags=["Base"])
|
||||||
|
|
@ -148,6 +148,55 @@ async def process_json(
|
||||||
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||||
|
|
||||||
|
|
||||||
|
# Endpoint to preload a graph
|
||||||
|
@router.post("/process/preload/{flow_id}", response_model=PreloadResponse)
|
||||||
|
async def preload_flow(
|
||||||
|
session: Annotated[Session, Depends(get_session)],
|
||||||
|
flow_id: str,
|
||||||
|
session_id: Optional[str] = None,
|
||||||
|
session_service: SessionService = Depends(get_session_service),
|
||||||
|
api_key_user: User = Depends(api_key_security),
|
||||||
|
clear_session: Annotated[bool, Body(embed=True)] = False, # noqa: F821
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
# Get the flow that matches the flow_id and belongs to the user
|
||||||
|
# flow = session.query(Flow).filter(Flow.id == flow_id).filter(Flow.user_id == api_key_user.id).first()
|
||||||
|
if clear_session:
|
||||||
|
session_service.clear_session(session_id)
|
||||||
|
# Check if the session exists
|
||||||
|
session_data = await session_service.load_session(session_id)
|
||||||
|
# Session data is a tuple of (graph, artifacts)
|
||||||
|
# or (None, None) if the session is empty
|
||||||
|
if isinstance(session_data, tuple):
|
||||||
|
graph, artifacts = session_data
|
||||||
|
is_clear = graph is None and artifacts is None
|
||||||
|
else:
|
||||||
|
is_clear = session_data is None
|
||||||
|
return PreloadResponse(session_id=session_id, is_clear=is_clear)
|
||||||
|
else:
|
||||||
|
if session_id is None:
|
||||||
|
session_id = flow_id
|
||||||
|
flow = session.exec(select(Flow).where(Flow.id == flow_id).where(Flow.user_id == api_key_user.id)).first()
|
||||||
|
if flow is None:
|
||||||
|
raise ValueError(f"Flow {flow_id} not found")
|
||||||
|
|
||||||
|
if flow.data is None:
|
||||||
|
raise ValueError(f"Flow {flow_id} has no data")
|
||||||
|
graph_data = flow.data
|
||||||
|
session_service.clear_session(session_id)
|
||||||
|
# Load the graph using SessionService
|
||||||
|
session_data = await session_service.load_session(session_id, graph_data)
|
||||||
|
graph, artifacts = session_data if session_data else (None, None)
|
||||||
|
if not graph:
|
||||||
|
raise ValueError("Graph not found in the session")
|
||||||
|
_ = await graph.build()
|
||||||
|
session_service.update_session(session_id, (graph, artifacts))
|
||||||
|
return PreloadResponse(session_id=session_id)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.exception(exc)
|
||||||
|
raise HTTPException(status_code=500, detail=str(exc)) from exc
|
||||||
|
|
||||||
|
|
||||||
@router.post(
|
@router.post(
|
||||||
"/predict/{flow_id}",
|
"/predict/{flow_id}",
|
||||||
response_model=ProcessResponse,
|
response_model=ProcessResponse,
|
||||||
|
|
@ -167,36 +216,75 @@ async def process(
|
||||||
task_service: "TaskService" = Depends(get_task_service),
|
task_service: "TaskService" = Depends(get_task_service),
|
||||||
api_key_user: User = Depends(api_key_security),
|
api_key_user: User = Depends(api_key_security),
|
||||||
sync: Annotated[bool, Body(embed=True)] = True, # noqa: F821
|
sync: Annotated[bool, Body(embed=True)] = True, # noqa: F821
|
||||||
|
session_service: SessionService = Depends(get_session_service),
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Endpoint to process an input with a given flow_id.
|
Endpoint to process an input with a given flow_id.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
try:
|
try:
|
||||||
if api_key_user is None:
|
if session_id:
|
||||||
raise HTTPException(
|
session_data = await session_service.load_session(session_id)
|
||||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
graph, artifacts = session_data if session_data else (None, None)
|
||||||
detail="Invalid API Key",
|
task_result: Any = None
|
||||||
|
task_status = None
|
||||||
|
task_id = None
|
||||||
|
if not graph:
|
||||||
|
raise ValueError("Graph not found in the session")
|
||||||
|
result = await build_graph_and_generate_result(
|
||||||
|
graph=graph,
|
||||||
|
inputs=inputs,
|
||||||
|
artifacts=artifacts,
|
||||||
|
session_id=session_id,
|
||||||
|
session_service=session_service,
|
||||||
|
)
|
||||||
|
task_id = str(id(result))
|
||||||
|
if isinstance(result, dict) and "result" in result:
|
||||||
|
task_result = result["result"]
|
||||||
|
session_id = result["session_id"]
|
||||||
|
elif hasattr(result, "result") and hasattr(result, "session_id"):
|
||||||
|
task_result = result.result
|
||||||
|
|
||||||
|
session_id = result.session_id
|
||||||
|
else:
|
||||||
|
task_result = result
|
||||||
|
if task_id:
|
||||||
|
task_response = TaskResponse(id=task_id, href=f"api/v1/task/{task_id}")
|
||||||
|
else:
|
||||||
|
task_response = None
|
||||||
|
return ProcessResponse(
|
||||||
|
result=task_result,
|
||||||
|
status=task_status,
|
||||||
|
task=task_response,
|
||||||
|
session_id=session_id,
|
||||||
|
backend=task_service.backend_name,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Get the flow that matches the flow_id and belongs to the user
|
else:
|
||||||
# flow = session.query(Flow).filter(Flow.id == flow_id).filter(Flow.user_id == api_key_user.id).first()
|
if api_key_user is None:
|
||||||
flow = session.exec(select(Flow).where(Flow.id == flow_id).where(Flow.user_id == api_key_user.id)).first()
|
raise HTTPException(
|
||||||
if flow is None:
|
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||||
raise ValueError(f"Flow {flow_id} not found")
|
detail="Invalid API Key",
|
||||||
|
)
|
||||||
|
|
||||||
if flow.data is None:
|
# Get the flow that matches the flow_id and belongs to the user
|
||||||
raise ValueError(f"Flow {flow_id} has no data")
|
# flow = session.query(Flow).filter(Flow.id == flow_id).filter(Flow.user_id == api_key_user.id).first()
|
||||||
graph_data = flow.data
|
flow = session.exec(select(Flow).where(Flow.id == flow_id).where(Flow.user_id == api_key_user.id)).first()
|
||||||
return await process_graph_data(
|
if flow is None:
|
||||||
graph_data=graph_data,
|
raise ValueError(f"Flow {flow_id} not found")
|
||||||
inputs=inputs,
|
|
||||||
tweaks=tweaks,
|
if flow.data is None:
|
||||||
clear_cache=clear_cache,
|
raise ValueError(f"Flow {flow_id} has no data")
|
||||||
session_id=session_id,
|
graph_data = flow.data
|
||||||
task_service=task_service,
|
return await process_graph_data(
|
||||||
sync=sync,
|
graph_data=graph_data,
|
||||||
)
|
inputs=inputs,
|
||||||
|
tweaks=tweaks,
|
||||||
|
clear_cache=clear_cache,
|
||||||
|
session_id=session_id,
|
||||||
|
task_service=task_service,
|
||||||
|
sync=sync,
|
||||||
|
)
|
||||||
except sa.exc.StatementError as exc:
|
except sa.exc.StatementError as exc:
|
||||||
# StatementError('(builtins.ValueError) badly formed hexadecimal UUID string')
|
# StatementError('(builtins.ValueError) badly formed hexadecimal UUID string')
|
||||||
if "badly formed hexadecimal UUID string" in str(exc):
|
if "badly formed hexadecimal UUID string" in str(exc):
|
||||||
|
|
|
||||||
|
|
@ -64,6 +64,13 @@ class ProcessResponse(BaseModel):
|
||||||
backend: Optional[str] = None
|
backend: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
class PreloadResponse(BaseModel):
|
||||||
|
"""Preload response schema."""
|
||||||
|
|
||||||
|
session_id: Optional[str] = None
|
||||||
|
is_clear: Optional[bool] = None
|
||||||
|
|
||||||
|
|
||||||
# TaskStatusResponse(
|
# TaskStatusResponse(
|
||||||
# status=task.status, result=task.result if task.ready() else None
|
# status=task.status, result=task.result if task.ready() else None
|
||||||
# )
|
# )
|
||||||
|
|
|
||||||
|
|
@ -7,6 +7,7 @@ from fastapi import FastAPI, Request
|
||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
from fastapi.responses import FileResponse
|
from fastapi.responses import FileResponse
|
||||||
from fastapi.staticfiles import StaticFiles
|
from fastapi.staticfiles import StaticFiles
|
||||||
|
|
||||||
from langflow.api import router
|
from langflow.api import router
|
||||||
from langflow.interface.utils import setup_llm_caching
|
from langflow.interface.utils import setup_llm_caching
|
||||||
from langflow.services.plugins.langfuse_plugin import LangfuseInstance
|
from langflow.services.plugins.langfuse_plugin import LangfuseInstance
|
||||||
|
|
@ -102,11 +103,12 @@ def setup_app(static_files_dir: Optional[Path] = None, backend_only: bool = Fals
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
import uvicorn
|
import uvicorn
|
||||||
|
|
||||||
from langflow.__main__ import get_number_of_workers
|
from langflow.__main__ import get_number_of_workers
|
||||||
|
|
||||||
configure()
|
configure()
|
||||||
uvicorn.run(
|
uvicorn.run(
|
||||||
create_app,
|
"langflow.main:create_app",
|
||||||
host="127.0.0.1",
|
host="127.0.0.1",
|
||||||
port=7860,
|
port=7860,
|
||||||
workers=get_number_of_workers(),
|
workers=get_number_of_workers(),
|
||||||
|
|
|
||||||
|
|
@ -7,9 +7,11 @@ from langchain.schema import AgentAction, Document
|
||||||
from langchain.vectorstores.base import VectorStore
|
from langchain.vectorstores.base import VectorStore
|
||||||
from langchain_core.messages import AIMessage
|
from langchain_core.messages import AIMessage
|
||||||
from langchain_core.runnables.base import Runnable
|
from langchain_core.runnables.base import Runnable
|
||||||
|
from langflow.graph.graph.base import Graph
|
||||||
from langflow.interface.custom.custom_component import CustomComponent
|
from langflow.interface.custom.custom_component import CustomComponent
|
||||||
from langflow.interface.run import build_sorted_vertices, get_memory_key, update_memory_keys
|
from langflow.interface.run import build_sorted_vertices, get_memory_key, update_memory_keys
|
||||||
from langflow.services.deps import get_session_service
|
from langflow.services.deps import get_session_service
|
||||||
|
from langflow.services.session.service import SessionService
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
|
|
@ -220,13 +222,27 @@ async def process_graph_cached(
|
||||||
graph, artifacts = session if session else (None, None)
|
graph, artifacts = session if session else (None, None)
|
||||||
if not graph:
|
if not graph:
|
||||||
raise ValueError("Graph not found in the session")
|
raise ValueError("Graph not found in the session")
|
||||||
|
|
||||||
|
result = await build_graph_and_generate_result(graph, inputs, artifacts, session_id, session_service)
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
async def build_graph_and_generate_result(
|
||||||
|
graph: "Graph",
|
||||||
|
session_id: str,
|
||||||
|
inputs: Optional[Union[dict, List[dict]]] = None,
|
||||||
|
artifacts: Optional[Dict[str, Any]] = None,
|
||||||
|
session_service: Optional[SessionService] = None,
|
||||||
|
):
|
||||||
|
"""Build the graph and generate the result"""
|
||||||
built_object = await graph.build()
|
built_object = await graph.build()
|
||||||
processed_inputs = process_inputs(inputs, artifacts or {})
|
processed_inputs = process_inputs(inputs, artifacts or {})
|
||||||
result = await generate_result(built_object, processed_inputs)
|
result = await generate_result(built_object, processed_inputs)
|
||||||
# langchain_object is now updated with the new memory
|
# langchain_object is now updated with the new memory
|
||||||
# we need to update the cache with the updated langchain_object
|
# we need to update the cache with the updated langchain_object
|
||||||
session_service.update_session(session_id, (graph, artifacts))
|
if session_id and session_service:
|
||||||
|
session_service.update_session(session_id, (graph, artifacts))
|
||||||
return Result(result=result, session_id=session_id)
|
return Result(result=result, session_id=session_id)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -2,11 +2,12 @@ from gunicorn.app.base import BaseApplication # type: ignore
|
||||||
|
|
||||||
|
|
||||||
class LangflowApplication(BaseApplication):
|
class LangflowApplication(BaseApplication):
|
||||||
def __init__(self, app, options=None):
|
def __init__(self, options=None):
|
||||||
self.options = options or {}
|
self.options = options or {}
|
||||||
|
from langflow.main import create_app
|
||||||
|
|
||||||
self.options["worker_class"] = "uvicorn.workers.UvicornWorker"
|
self.options["worker_class"] = "uvicorn.workers.UvicornWorker"
|
||||||
self.application = app
|
self.application = create_app()
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
def load_config(self):
|
def load_config(self):
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
from langflow.interface.run import build_sorted_vertices
|
from langflow.interface.run import build_sorted_vertices
|
||||||
from langflow.services.base import Service
|
from langflow.services.base import Service
|
||||||
|
|
@ -14,14 +14,15 @@ class SessionService(Service):
|
||||||
def __init__(self, cache_service):
|
def __init__(self, cache_service):
|
||||||
self.cache_service: "BaseCacheService" = cache_service
|
self.cache_service: "BaseCacheService" = cache_service
|
||||||
|
|
||||||
async def load_session(self, key, data_graph):
|
async def load_session(self, key, data_graph: Optional[dict] = None):
|
||||||
# Check if the data is cached
|
# Check if the data is cached
|
||||||
if key in self.cache_service:
|
if key in self.cache_service:
|
||||||
return self.cache_service.get(key)
|
return self.cache_service.get(key)
|
||||||
|
|
||||||
if key is None:
|
if key is None:
|
||||||
key = self.generate_key(session_id=None, data_graph=data_graph)
|
key = self.generate_key(session_id=None, data_graph=data_graph)
|
||||||
|
if data_graph is None:
|
||||||
|
return (None, None)
|
||||||
# If not cached, build the graph and cache it
|
# If not cached, build the graph and cache it
|
||||||
graph, artifacts = await build_sorted_vertices(data_graph)
|
graph, artifacts = await build_sorted_vertices(data_graph)
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue