refac: async db
This commit is contained in:
@@ -4,8 +4,9 @@ import uuid
|
||||
from typing import Optional
|
||||
from functools import lru_cache
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, get_db, get_db_context
|
||||
from sqlalchemy import select, delete, update, or_, func, cast
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, get_async_db_context
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.users import User, UserModel, Users, UserResponse
|
||||
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
|
||||
@@ -13,7 +14,6 @@ from open_webui.models.access_grants import AccessGrantModel, AccessGrants
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from sqlalchemy import BigInteger, Column, Text, JSON
|
||||
from sqlalchemy import or_, func, cast
|
||||
|
||||
####################
|
||||
# Note DB Schema
|
||||
@@ -88,18 +88,18 @@ class NoteListResponse(BaseModel):
|
||||
|
||||
|
||||
class NoteTable:
|
||||
def _get_access_grants(self, note_id: str, db: Optional[Session] = None) -> list[AccessGrantModel]:
|
||||
return AccessGrants.get_grants_by_resource('note', note_id, db=db)
|
||||
async def _get_access_grants(self, note_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
|
||||
return await AccessGrants.get_grants_by_resource('note', note_id, db=db)
|
||||
|
||||
def _to_note_model(
|
||||
async def _to_note_model(
|
||||
self,
|
||||
note: Note,
|
||||
access_grants: Optional[list[AccessGrantModel]] = None,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> NoteModel:
|
||||
note_data = NoteModel.model_validate(note).model_dump(exclude={'access_grants'})
|
||||
note_data['access_grants'] = (
|
||||
access_grants if access_grants is not None else self._get_access_grants(note_data['id'], db=db)
|
||||
access_grants if access_grants is not None else await self._get_access_grants(note_data['id'], db=db)
|
||||
)
|
||||
return NoteModel.model_validate(note_data)
|
||||
|
||||
@@ -113,8 +113,8 @@ class NoteTable:
|
||||
permission=permission,
|
||||
)
|
||||
|
||||
def insert_new_note(self, user_id: str, form_data: NoteForm, db: Optional[Session] = None) -> Optional[NoteModel]:
|
||||
with get_db_context(db) as db:
|
||||
async def insert_new_note(self, user_id: str, form_data: NoteForm, db: Optional[AsyncSession] = None) -> Optional[NoteModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
note = NoteModel(
|
||||
**{
|
||||
'id': str(uuid.uuid4()),
|
||||
@@ -129,38 +129,39 @@ class NoteTable:
|
||||
new_note = Note(**note.model_dump(exclude={'access_grants'}))
|
||||
|
||||
db.add(new_note)
|
||||
db.commit()
|
||||
AccessGrants.set_access_grants('note', note.id, form_data.access_grants, db=db)
|
||||
return self._to_note_model(new_note, db=db)
|
||||
await db.commit()
|
||||
await AccessGrants.set_access_grants('note', note.id, form_data.access_grants, db=db)
|
||||
return await self._to_note_model(new_note, db=db)
|
||||
|
||||
def get_notes(self, skip: int = 0, limit: int = 50, db: Optional[Session] = None) -> list[NoteModel]:
|
||||
with get_db_context(db) as db:
|
||||
query = db.query(Note).order_by(Note.updated_at.desc())
|
||||
async def get_notes(self, skip: int = 0, limit: int = 50, db: Optional[AsyncSession] = None) -> list[NoteModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(Note).order_by(Note.updated_at.desc())
|
||||
if skip is not None:
|
||||
query = query.offset(skip)
|
||||
stmt = stmt.offset(skip)
|
||||
if limit is not None:
|
||||
query = query.limit(limit)
|
||||
notes = query.all()
|
||||
stmt = stmt.limit(limit)
|
||||
result = await db.execute(stmt)
|
||||
notes = result.scalars().all()
|
||||
note_ids = [note.id for note in notes]
|
||||
grants_map = AccessGrants.get_grants_by_resources('note', note_ids, db=db)
|
||||
return [self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db)
|
||||
return [await self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes]
|
||||
|
||||
def search_notes(
|
||||
async def search_notes(
|
||||
self,
|
||||
user_id: str,
|
||||
filter: dict = {},
|
||||
skip: int = 0,
|
||||
limit: int = 30,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> NoteListResponse:
|
||||
with get_db_context(db) as db:
|
||||
query = db.query(Note, User).outerjoin(User, User.id == Note.user_id)
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(Note, User).outerjoin(User, User.id == Note.user_id)
|
||||
if filter:
|
||||
query_key = filter.get('query')
|
||||
if query_key:
|
||||
# Normalize search by removing hyphens and spaces (e.g., "todo" matches "to-do" and "to do")
|
||||
normalized_query = query_key.replace('-', '').replace(' ', '')
|
||||
query = query.filter(
|
||||
stmt = stmt.filter(
|
||||
or_(
|
||||
func.replace(func.replace(Note.title, '-', ''), ' ', '').ilike(f'%{normalized_query}%'),
|
||||
func.replace(
|
||||
@@ -173,9 +174,9 @@ class NoteTable:
|
||||
|
||||
view_option = filter.get('view_option')
|
||||
if view_option == 'created':
|
||||
query = query.filter(Note.user_id == user_id)
|
||||
stmt = stmt.filter(Note.user_id == user_id)
|
||||
elif view_option == 'shared':
|
||||
query = query.filter(Note.user_id != user_id)
|
||||
stmt = stmt.filter(Note.user_id != user_id)
|
||||
|
||||
# Apply access control filtering
|
||||
if 'permission' in filter:
|
||||
@@ -183,9 +184,9 @@ class NoteTable:
|
||||
else:
|
||||
permission = 'write'
|
||||
|
||||
query = self._has_permission(
|
||||
stmt = self._has_permission(
|
||||
db,
|
||||
query,
|
||||
stmt,
|
||||
filter,
|
||||
permission=permission,
|
||||
)
|
||||
@@ -195,87 +196,95 @@ class NoteTable:
|
||||
|
||||
if order_by == 'name':
|
||||
if direction == 'asc':
|
||||
query = query.order_by(Note.title.asc())
|
||||
stmt = stmt.order_by(Note.title.asc())
|
||||
else:
|
||||
query = query.order_by(Note.title.desc())
|
||||
stmt = stmt.order_by(Note.title.desc())
|
||||
elif order_by == 'created_at':
|
||||
if direction == 'asc':
|
||||
query = query.order_by(Note.created_at.asc())
|
||||
stmt = stmt.order_by(Note.created_at.asc())
|
||||
else:
|
||||
query = query.order_by(Note.created_at.desc())
|
||||
stmt = stmt.order_by(Note.created_at.desc())
|
||||
elif order_by == 'updated_at':
|
||||
if direction == 'asc':
|
||||
query = query.order_by(Note.updated_at.asc())
|
||||
stmt = stmt.order_by(Note.updated_at.asc())
|
||||
else:
|
||||
query = query.order_by(Note.updated_at.desc())
|
||||
stmt = stmt.order_by(Note.updated_at.desc())
|
||||
else:
|
||||
query = query.order_by(Note.updated_at.desc())
|
||||
stmt = stmt.order_by(Note.updated_at.desc())
|
||||
|
||||
else:
|
||||
query = query.order_by(Note.updated_at.desc())
|
||||
stmt = stmt.order_by(Note.updated_at.desc())
|
||||
|
||||
# Count BEFORE pagination
|
||||
total = query.count()
|
||||
count_result = await db.execute(
|
||||
select(func.count()).select_from(stmt.subquery())
|
||||
)
|
||||
total = count_result.scalar()
|
||||
|
||||
if skip:
|
||||
query = query.offset(skip)
|
||||
stmt = stmt.offset(skip)
|
||||
if limit:
|
||||
query = query.limit(limit)
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
items = query.all()
|
||||
result = await db.execute(stmt)
|
||||
items = result.all()
|
||||
|
||||
note_ids = [note.id for note, _ in items]
|
||||
grants_map = AccessGrants.get_grants_by_resources('note', note_ids, db=db)
|
||||
grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db)
|
||||
|
||||
notes = []
|
||||
for note, user in items:
|
||||
notes.append(
|
||||
NoteUserResponse(
|
||||
**self._to_note_model(
|
||||
**(await self._to_note_model(
|
||||
note,
|
||||
access_grants=grants_map.get(note.id, []),
|
||||
db=db,
|
||||
).model_dump(),
|
||||
)).model_dump(),
|
||||
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
|
||||
)
|
||||
)
|
||||
|
||||
return NoteListResponse(items=notes, total=total)
|
||||
|
||||
def get_notes_by_user_id(
|
||||
async def get_notes_by_user_id(
|
||||
self,
|
||||
user_id: str,
|
||||
permission: str = 'read',
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[NoteModel]:
|
||||
with get_db_context(db) as db:
|
||||
user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id, db=db)]
|
||||
async with get_async_db_context(db) as db:
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = [group.id for group in user_groups]
|
||||
|
||||
query = db.query(Note).order_by(Note.updated_at.desc())
|
||||
query = self._has_permission(db, query, {'user_id': user_id, 'group_ids': user_group_ids}, permission)
|
||||
stmt = select(Note).order_by(Note.updated_at.desc())
|
||||
stmt = self._has_permission(db, stmt, {'user_id': user_id, 'group_ids': user_group_ids}, permission)
|
||||
|
||||
if skip is not None:
|
||||
query = query.offset(skip)
|
||||
stmt = stmt.offset(skip)
|
||||
if limit is not None:
|
||||
query = query.limit(limit)
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
notes = query.all()
|
||||
result = await db.execute(stmt)
|
||||
notes = result.scalars().all()
|
||||
note_ids = [note.id for note in notes]
|
||||
grants_map = AccessGrants.get_grants_by_resources('note', note_ids, db=db)
|
||||
return [self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db)
|
||||
return [await self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes]
|
||||
|
||||
def get_note_by_id(self, id: str, db: Optional[Session] = None) -> Optional[NoteModel]:
|
||||
with get_db_context(db) as db:
|
||||
note = db.query(Note).filter(Note.id == id).first()
|
||||
return self._to_note_model(note, db=db) if note else None
|
||||
async def get_note_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[NoteModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Note).filter(Note.id == id))
|
||||
note = result.scalars().first()
|
||||
return await self._to_note_model(note, db=db) if note else None
|
||||
|
||||
def update_note_by_id(
|
||||
self, id: str, form_data: NoteUpdateForm, db: Optional[Session] = None
|
||||
async def update_note_by_id(
|
||||
self, id: str, form_data: NoteUpdateForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[NoteModel]:
|
||||
with get_db_context(db) as db:
|
||||
note = db.query(Note).filter(Note.id == id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Note).filter(Note.id == id))
|
||||
note = result.scalars().first()
|
||||
if not note:
|
||||
return None
|
||||
|
||||
@@ -289,19 +298,19 @@ class NoteTable:
|
||||
note.meta = {**note.meta, **form_data['meta']}
|
||||
|
||||
if 'access_grants' in form_data:
|
||||
AccessGrants.set_access_grants('note', id, form_data['access_grants'], db=db)
|
||||
await AccessGrants.set_access_grants('note', id, form_data['access_grants'], db=db)
|
||||
|
||||
note.updated_at = int(time.time_ns())
|
||||
|
||||
db.commit()
|
||||
return self._to_note_model(note, db=db) if note else None
|
||||
await db.commit()
|
||||
return await self._to_note_model(note, db=db) if note else None
|
||||
|
||||
def delete_note_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
async def delete_note_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
AccessGrants.revoke_all_access('note', id, db=db)
|
||||
db.query(Note).filter(Note.id == id).delete()
|
||||
db.commit()
|
||||
async with get_async_db_context(db) as db:
|
||||
await AccessGrants.revoke_all_access('note', id, db=db)
|
||||
await db.execute(delete(Note).filter(Note.id == id))
|
||||
await db.commit()
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
Reference in New Issue
Block a user