feat: Add AsyncSession support for non-blocking db operations (#4408)

* Add AsyncSession support for non-blocking db operations

* Use sqlalchemy extras
This commit is contained in:
Christophe Bornet 2024-11-06 12:21:53 +01:00 • committed by GitHub
commit e853a13d57
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 641 additions and 100 deletions

View file

@ -1,3 +1,4 @@
import asyncio
import inspect
import platform
import socket
@ -27,7 +28,7 @@ from langflow.services.database.models.folder.utils import (
create_default_folder_if_it_doesnt_exist,
)
from langflow.services.database.utils import session_getter
from langflow.services.deps import get_db_service, get_settings_service, session_scope
from langflow.services.deps import async_session_scope, get_db_service, get_settings_service
from langflow.services.settings.constants import DEFAULT_SUPERUSER
from langflow.services.utils import initialize_services
from langflow.utils.version import fetch_latest_version, get_version_info
@ -486,28 +487,35 @@ def api_key(
if not auth_settings.AUTO_LOGIN:
typer.echo("Auto login is disabled. API keys cannot be created through the CLI.")
return
with session_scope() as session:
from langflow.services.database.models.user.model import User
superuser = session.exec(select(User).where(User.username == DEFAULT_SUPERUSER)).first()
if not superuser:
typer.echo("Default superuser not found. This command requires a superuser and AUTO_LOGIN to be enabled.")
return
from langflow.services.database.models.api_key import ApiKey, ApiKeyCreate
from langflow.services.database.models.api_key.crud import (
create_api_key,
delete_api_key,
)
async def aapi_key():
async with async_session_scope() as session:
from langflow.services.database.models.user.model import User
api_key = session.exec(select(ApiKey).where(ApiKey.user_id == superuser.id)).first()
if api_key:
delete_api_key(session, api_key.id)
superuser = (await session.exec(select(User).where(User.username == DEFAULT_SUPERUSER))).first()
if not superuser:
typer.echo(
"Default superuser not found. This command requires a superuser and AUTO_LOGIN to be enabled."
)
return None
from langflow.services.database.models.api_key import ApiKey, ApiKeyCreate
from langflow.services.database.models.api_key.crud import (
create_api_key,
delete_api_key,
)
api_key_create = ApiKeyCreate(name="CLI")
unmasked_api_key = create_api_key(session, api_key_create, user_id=superuser.id)
session.commit()
# Create a banner to display the API key and tell the user it won't be shown again
api_key_banner(unmasked_api_key)
api_key = (await session.exec(select(ApiKey).where(ApiKey.user_id == superuser.id))).first()
if api_key:
await delete_api_key(session, api_key.id)
api_key_create = ApiKeyCreate(name="CLI")
unmasked_api_key = await create_api_key(session, api_key_create, user_id=superuser.id)
await session.commit()
return unmasked_api_key
unmasked_api_key = asyncio.run(aapi_key())
# Create a banner to display the API key and tell the user it won't be shown again
api_key_banner(unmasked_api_key)
def api_key_banner(unmasked_api_key) -> None:

View file

@ -9,6 +9,7 @@ from fastapi_pagination import Params
from loguru import logger
from sqlalchemy import delete
from sqlmodel import Session
from sqlmodel.ext.asyncio.session import AsyncSession
from langflow.graph.graph.base import Graph
from langflow.services.auth.utils import get_current_active_user
@ -16,7 +17,7 @@ from langflow.services.database.models import User
from langflow.services.database.models.flow import Flow
from langflow.services.database.models.transactions.model import TransactionTable
from langflow.services.database.models.vertex_builds.model import VertexBuildTable
from langflow.services.deps import get_session
from langflow.services.deps import get_async_session, get_session
from langflow.services.store.utils import get_lf_version_from_pypi
if TYPE_CHECKING:
@ -31,6 +32,7 @@ MIN_PAGE_SIZE = 1
CurrentActiveUser = Annotated[User, Depends(get_current_active_user)]
DbSession = Annotated[Session, Depends(get_session)]
AsyncDbSession = Annotated[AsyncSession, Depends(get_async_session)]
def has_api_terms(word: str):

View file

@ -3,7 +3,7 @@ from uuid import UUID
from fastapi import APIRouter, Depends, HTTPException, Response
from langflow.api.utils import CurrentActiveUser, DbSession
from langflow.api.utils import AsyncDbSession, CurrentActiveUser, DbSession
from langflow.api.v1.schemas import ApiKeyCreateRequest, ApiKeysResponse
from langflow.services.auth import utils as auth_utils
@ -20,12 +20,12 @@ router = APIRouter(tags=["APIKey"], prefix="/api_key")
@router.get("/")
async def get_api_keys_route(
db: DbSession,
db: AsyncDbSession,
current_user: CurrentActiveUser,
) -> ApiKeysResponse:
try:
user_id = current_user.id
keys = get_api_keys(db, user_id)
keys = await get_api_keys(db, user_id)
return ApiKeysResponse(total_count=len(keys), user_id=user_id, api_keys=keys)
except Exception as exc:
@ -36,11 +36,11 @@ async def get_api_keys_route(
async def create_api_key_route(
req: ApiKeyCreate,
current_user: CurrentActiveUser,
db: DbSession,
db: AsyncDbSession,
) -> UnmaskedApiKeyRead:
try:
user_id = current_user.id
return create_api_key(db, req, user_id=user_id)
return await create_api_key(db, req, user_id=user_id)
except Exception as e:
raise HTTPException(status_code=400, detail=str(e)) from e
@ -48,10 +48,10 @@ async def create_api_key_route(
@router.delete("/{api_key_id}", dependencies=[Depends(auth_utils.get_current_active_user)])
async def delete_api_key_route(
api_key_id: UUID,
db: DbSession,
db: AsyncDbSession,
):
try:
delete_api_key(db, api_key_id)
await delete_api_key(db, api_key_id)
except Exception as e:
raise HTTPException(status_code=400, detail=str(e)) from e
return {"detail": "API Key deleted"}

View file

@ -5,6 +5,7 @@ from typing import TYPE_CHECKING
from uuid import UUID
from sqlmodel import Session, select
from sqlmodel.ext.asyncio.session import AsyncSession
from langflow.services.database.models.api_key import ApiKey, ApiKeyCreate, ApiKeyRead, UnmaskedApiKeyRead
@ -12,13 +13,13 @@ if TYPE_CHECKING:
from sqlmodel.sql.expression import SelectOfScalar
def get_api_keys(session: Session, user_id: UUID) -> list[ApiKeyRead]:
async def get_api_keys(session: AsyncSession, user_id: UUID) -> list[ApiKeyRead]:
query: SelectOfScalar = select(ApiKey).where(ApiKey.user_id == user_id)
api_keys = session.exec(query).all()
api_keys = (await session.exec(query)).all()
return [ApiKeyRead.model_validate(api_key) for api_key in api_keys]
def create_api_key(session: Session, api_key_create: ApiKeyCreate, user_id: UUID) -> UnmaskedApiKeyRead:
async def create_api_key(session: AsyncSession, api_key_create: ApiKeyCreate, user_id: UUID) -> UnmaskedApiKeyRead:
# Generate a random API key with 32 bytes of randomness
generated_api_key = f"sk-{secrets.token_urlsafe(32)}"
@ -30,20 +31,20 @@ def create_api_key(session: Session, api_key_create: ApiKeyCreate, user_id: UUID
)
session.add(api_key)
session.commit()
session.refresh(api_key)
await session.commit()
await session.refresh(api_key)
unmasked = UnmaskedApiKeyRead.model_validate(api_key, from_attributes=True)
unmasked.api_key = generated_api_key
return unmasked
def delete_api_key(session: Session, api_key_id: UUID) -> None:
api_key = session.get(ApiKey, api_key_id)
async def delete_api_key(session: AsyncSession, api_key_id: UUID) -> None:
api_key = await session.get(ApiKey, api_key_id)
if api_key is None:
msg = "API Key not found"
raise ValueError(msg)
session.delete(api_key)
session.commit()
await session.delete(api_key)
await session.commit()
def check_key(session: Session, api_key: str) -> ApiKey | None:

View file

@ -1,8 +1,9 @@
from __future__ import annotations
import asyncio
import sqlite3
import time
from contextlib import contextmanager
from contextlib import asynccontextmanager, contextmanager
from datetime import datetime, timezone
from pathlib import Path
from typing import TYPE_CHECKING
@ -14,7 +15,9 @@ from loguru import logger
from sqlalchemy import event, inspect
from sqlalchemy.engine import Engine
from sqlalchemy.exc import OperationalError
from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine
from sqlmodel import Session, SQLModel, create_engine, select, text
from sqlmodel.ext.asyncio.session import AsyncSession
from langflow.services.base import Service
from langflow.services.database import models
@ -39,12 +42,17 @@ class DatabaseService(Service):
msg = "No database URL provided"
raise ValueError(msg)
self.database_url: str = settings_service.settings.database_url
self._sanitize_database_url()
# This file is in langflow.services.database.manager.py
# the ini is in langflow
langflow_dir = Path(__file__).parent.parent.parent
self.script_location = langflow_dir / "alembic"
self.alembic_cfg_path = langflow_dir / "alembic.ini"
# register the event listener for sqlite as part of this class.
# Using decorator will make the method not able to use self
event.listen(Engine, "connect", self.on_connection)
self.engine = self._create_engine()
self.async_engine = self._create_async_engine()
alembic_log_file = self.settings_service.settings.alembic_log_file
# Check if the provided path is absolute, cross-platform.
@ -56,10 +64,47 @@ class DatabaseService(Service):
self.alembic_log_path = Path(langflow_dir) / alembic_log_file
def reload_engine(self) -> None:
self._sanitize_database_url()
self.engine = self._create_engine()
self.async_engine = self._create_async_engine()
def _sanitize_database_url(self):
if self.database_url.startswith("postgres://"):
self.database_url = self.database_url.replace("postgres://", "postgresql://")
logger.warning(
"Fixed postgres dialect in database URL. Replacing postgres:// with postgresql://. "
"To avoid this warning, update the database URL."
)
def _create_engine(self) -> Engine:
"""Create the engine for the database."""
return create_engine(
self.database_url,
connect_args=self._get_connect_args(),
pool_size=self.settings_service.settings.pool_size,
max_overflow=self.settings_service.settings.max_overflow,
)
def _create_async_engine(self) -> AsyncEngine:
"""Create the engine for the database."""
url_components = self.database_url.split("://", maxsplit=1)
if url_components[0].startswith("sqlite"):
database_url = "sqlite+aiosqlite://"
kwargs = {}
else:
kwargs = {
"pool_size": self.settings_service.settings.pool_size,
"max_overflow": self.settings_service.settings.max_overflow,
}
database_url = "postgresql+psycopg://" if url_components[0].startswith("postgresql") else url_components[0]
database_url += url_components[1]
return create_async_engine(
database_url,
connect_args=self._get_connect_args(),
**kwargs,
)
def _get_connect_args(self):
if self.settings_service.settings.database_url and self.settings_service.settings.database_url.startswith(
"sqlite"
):
@ -69,33 +114,12 @@ class DatabaseService(Service):
}
else:
connect_args = {}
try:
# register the event listener for sqlite as part of this class.
# Using decorator will make the method not able to use self
event.listen(Engine, "connect", self.on_connection)
return create_engine(
self.database_url,
connect_args=connect_args,
pool_size=self.settings_service.settings.pool_size,
max_overflow=self.settings_service.settings.max_overflow,
)
except sa.exc.NoSuchModuleError as exc:
if "postgres" in str(exc) and not self.database_url.startswith("postgresql"):
# https://stackoverflow.com/questions/62688256/sqlalchemy-exc-nosuchmoduleerror-cant-load-plugin-sqlalchemy-dialectspostgre
self.database_url = self.database_url.replace("postgres://", "postgresql://")
logger.warning(
"Fixed postgres dialect in database URL. Replacing postgres:// with postgresql://. "
"To avoid this warning, update the database URL."
)
return self._create_engine()
msg = "Error creating database engine"
raise RuntimeError(msg) from exc
return connect_args
def on_connection(self, dbapi_connection, _connection_record) -> None:
from sqlite3 import Connection as sqliteConnection
if isinstance(dbapi_connection, sqliteConnection):
if isinstance(
dbapi_connection, sqlite3.Connection | sa.dialects.sqlite.aiosqlite.AsyncAdapt_aiosqlite_connection
):
pragmas: dict = self.settings_service.settings.sqlite_pragmas or {}
pragmas_list = []
for key, val in pragmas.items():
@ -117,6 +141,11 @@ class DatabaseService(Service):
with Session(self.engine) as session:
yield session
@asynccontextmanager
async def with_async_session(self):
async with AsyncSession(self.async_engine) as session:
yield session
def migrate_flows_if_auto_login(self) -> None:
# if auto_login is enabled, we need to migrate the flows
# to the default superuser if they don't have a user id
@ -334,3 +363,4 @@ class DatabaseService(Service):
async def teardown(self) -> None:
await asyncio.to_thread(self._teardown)
await self.async_engine.dispose()

View file

@ -1,6 +1,6 @@
from __future__ import annotations
from contextlib import contextmanager
from contextlib import asynccontextmanager, contextmanager
from typing import TYPE_CHECKING
from loguru import logger
@ -8,9 +8,10 @@ from loguru import logger
from langflow.services.schema import ServiceType
if TYPE_CHECKING:
from collections.abc import Generator
from collections.abc import AsyncGenerator, Generator
from sqlmodel import Session
from sqlmodel.ext.asyncio.session import AsyncSession
from langflow.services.cache.service import AsyncBaseCacheService, CacheService
from langflow.services.chat.service import ChatService
@ -162,6 +163,17 @@ def get_session() -> Generator[Session, None, None]:
yield session
async def get_async_session() -> AsyncGenerator[AsyncSession, None]:
"""Retrieves an async session from the database service.
Yields:
Session: An async session object.
"""
async with get_db_service().with_async_session() as session:
yield session
@contextmanager
def session_scope() -> Generator[Session, None, None]:
"""Context manager for managing a session scope.
@ -188,6 +200,32 @@ def session_scope() -> Generator[Session, None, None]:
raise
@asynccontextmanager
async def async_session_scope() -> AsyncGenerator[AsyncSession, None]:
"""Context manager for managing an async session scope.
This context manager is used to manage an async session scope for database operations.
It ensures that the session is properly committed if no exceptions occur,
and rolled back if an exception is raised.
Yields:
session: The async session object.
Raises:
Exception: If an error occurs during the session scope.
"""
db_service = get_db_service()
async with db_service.with_async_session() as session:
try:
yield session
await session.commit()
except Exception:
logger.exception("An error occurred during the session scope.")
await session.rollback()
raise
def get_cache_service() -> CacheService | AsyncBaseCacheService:
"""Retrieves the cache service from the service manager.