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:
parent
07d8f2e04b
commit
e853a13d57
8 changed files with 641 additions and 100 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue