fix: Fix async usage in app startup (#4285)

Fix async usage in app startup
This commit is contained in:
Christophe Bornet 2024-10-27 16:00:28 +01:00 • committed by GitHub
commit c8bdcf36b0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 24 additions and 26 deletions

View file

@ -3,7 +3,6 @@ import json
import shutil import shutil
import time import time
from collections import defaultdict from collections import defaultdict
from collections.abc import Awaitable
from copy import deepcopy from copy import deepcopy
from datetime import datetime, timezone from datetime import datetime, timezone
from pathlib import Path from pathlib import Path
@ -600,12 +599,7 @@ def find_existing_flow(session, flow_id, flow_endpoint_name):
return None return None
async def create_or_update_starter_projects(get_all_components_coro: Awaitable[dict]) -> None: def create_or_update_starter_projects(all_types_dict: dict) -> None:
try:
all_types_dict = await get_all_components_coro
except Exception:
logger.exception("Error loading components")
raise
with session_scope() as session: with session_scope() as session:
new_folder = create_starter_folder(session) new_folder = create_starter_folder(session)
starter_projects = load_starter_projects() starter_projects = load_starter_projects()

View file

@ -8,7 +8,6 @@ from http import HTTPStatus
from pathlib import Path from pathlib import Path
from urllib.parse import urlencode from urllib.parse import urlencode
import nest_asyncio
from fastapi import FastAPI, HTTPException, Request, Response, status from fastapi import FastAPI, HTTPException, Request, Response, status
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse, JSONResponse from fastapi.responses import FileResponse, JSONResponse
@ -87,28 +86,25 @@ class JavaScriptMIMETypeMiddleware(BaseHTTPMiddleware):
return response return response
telemetry_service_tasks = set()
def get_lifespan(*, fix_migration=False, version=None): def get_lifespan(*, fix_migration=False, version=None):
def _initialize():
initialize_services(fix_migration=fix_migration)
setup_llm_caching()
initialize_super_user_if_needed()
@asynccontextmanager @asynccontextmanager
async def lifespan(_app: FastAPI): async def lifespan(_app: FastAPI):
nest_asyncio.apply()
# Startup message # Startup message
if version: if version:
rprint(f"[bold green]Starting Langflow v{version}...[/bold green]") rprint(f"[bold green]Starting Langflow v{version}...[/bold green]")
else: else:
rprint("[bold green]Starting Langflow...[/bold green]") rprint("[bold green]Starting Langflow...[/bold green]")
try: try:
initialize_services(fix_migration=fix_migration) await asyncio.to_thread(_initialize)
setup_llm_caching() all_types_dict = await get_and_cache_all_types_dict(get_settings_service())
initialize_super_user_if_needed() await asyncio.to_thread(create_or_update_starter_projects, all_types_dict)
task = asyncio.create_task(get_and_cache_all_types_dict(get_settings_service())) get_telemetry_service().start()
await create_or_update_starter_projects(task) await asyncio.to_thread(load_flows_from_directory)
telemetry_service_task = asyncio.create_task(get_telemetry_service().start())
telemetry_service_tasks.add(telemetry_service_task)
telemetry_service_task.add_done_callback(telemetry_service_tasks.discard)
load_flows_from_directory()
yield yield
except Exception as exc: except Exception as exc:
if "langflow migration --fix" not in str(exc): if "langflow migration --fix" not in str(exc):

View file

@ -1,7 +1,6 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import contextlib
import os import os
import platform import platform
from datetime import datetime, timezone from datetime import datetime, timezone
@ -112,7 +111,7 @@ class TelemetryService(Service):
async def log_package_component(self, payload: ComponentPayload) -> None: async def log_package_component(self, payload: ComponentPayload) -> None:
await self._queue_event((self.send_telemetry_data, payload, "component")) await self._queue_event((self.send_telemetry_data, payload, "component"))
async def start(self) -> None: def start(self) -> None:
if self.running or self.do_not_track: if self.running or self.do_not_track:
return return
try: try:
@ -131,6 +130,15 @@ class TelemetryService(Service):
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
logger.exception("Error flushing logs") logger.exception("Error flushing logs")
async def _cancel_task(self, task: asyncio.Task, cancel_msg: str) -> None:
task.cancel(cancel_msg)
try:
await task
except asyncio.CancelledError:
current_task = asyncio.current_task()
if current_task and current_task.cancelling() > 0:
raise
async def stop(self) -> None: async def stop(self) -> None:
if self.do_not_track or self._stopping: if self.do_not_track or self._stopping:
return return
@ -140,9 +148,9 @@ class TelemetryService(Service):
await self.flush() await self.flush()
self.running = False self.running = False
if self.worker_task: if self.worker_task:
self.worker_task.cancel() await self._cancel_task(self.worker_task, "Cancel telemetry worker task")
with contextlib.suppress(asyncio.CancelledError): if self.log_package_version_task:
await self.worker_task await self._cancel_task(self.log_package_version_task, "Cancel telemetry log package version task")
await self.client.aclose() await self.client.aclose()
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
logger.exception("Error stopping tracing service") logger.exception("Error stopping tracing service")