refac/enh: db session sharing

This commit is contained in:
Timothy Jaeryang Baek
2025-12-28 22:00:44 +04:00
parent d4de26bd05
commit 2041ab483e
20 changed files with 600 additions and 562 deletions
+51 -48
View File
@@ -1,10 +1,12 @@
import json
import logging
import time
from typing import Optional
import uuid
from open_webui.internal.db import Base, get_db
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from open_webui.models.files import (
File,
@@ -157,9 +159,9 @@ class KnowledgeFileListResponse(BaseModel):
class KnowledgeTable:
def insert_new_knowledge(
self, user_id: str, form_data: KnowledgeForm
self, user_id: str, form_data: KnowledgeForm, db: Optional[Session] = None
) -> Optional[KnowledgeModel]:
with get_db() as db:
with get_db_context(db) as db:
knowledge = KnowledgeModel(
**{
**form_data.model_dump(),
@@ -183,15 +185,15 @@ class KnowledgeTable:
return None
def get_knowledge_bases(
self, skip: int = 0, limit: int = 30
self, skip: int = 0, limit: int = 30, db: Optional[Session] = None
) -> list[KnowledgeUserModel]:
with get_db() as db:
with get_db_context(db) as db:
all_knowledge = (
db.query(Knowledge).order_by(Knowledge.updated_at.desc()).all()
)
user_ids = list(set(knowledge.user_id for knowledge in all_knowledge))
users = Users.get_users_by_user_ids(user_ids) if user_ids else []
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users_dict = {user.id: user for user in users}
knowledge_bases = []
@@ -208,10 +210,10 @@ class KnowledgeTable:
return knowledge_bases
def search_knowledge_bases(
self, user_id: str, filter: dict, skip: int = 0, limit: int = 30
self, user_id: str, filter: dict, skip: int = 0, limit: int = 30, db: Optional[Session] = None
) -> KnowledgeListResponse:
try:
with get_db() as db:
with get_db_context(db) as db:
query = db.query(Knowledge, User).outerjoin(
User, User.id == Knowledge.user_id
)
@@ -267,14 +269,14 @@ class KnowledgeTable:
return KnowledgeListResponse(items=[], total=0)
def search_knowledge_files(
self, filter: dict, skip: int = 0, limit: int = 30
self, filter: dict, skip: int = 0, limit: int = 30, db: Optional[Session] = None
) -> KnowledgeFileListResponse:
"""
Scalable version: search files across all knowledge bases the user has
READ access to, without loading all KBs or using large IN() lists.
"""
try:
with get_db() as db:
with get_db_context(db) as db:
# Base query: join Knowledge → KnowledgeFile → File
query = (
db.query(File, User)
@@ -327,20 +329,20 @@ class KnowledgeTable:
print("search_knowledge_files error:", e)
return KnowledgeFileListResponse(items=[], total=0)
def check_access_by_user_id(self, id, user_id, permission="write") -> bool:
knowledge = self.get_knowledge_by_id(id)
def check_access_by_user_id(self, id, user_id, permission="write", db: Optional[Session] = None) -> bool:
knowledge = self.get_knowledge_by_id(id, db=db)
if not knowledge:
return False
if knowledge.user_id == user_id:
return True
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id)}
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
return has_access(user_id, permission, knowledge.access_control, user_group_ids)
def get_knowledge_bases_by_user_id(
self, user_id: str, permission: str = "write"
self, user_id: str, permission: str = "write", db: Optional[Session] = None
) -> list[KnowledgeUserModel]:
knowledge_bases = self.get_knowledge_bases()
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id)}
knowledge_bases = self.get_knowledge_bases(db=db)
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
return [
knowledge_base
for knowledge_base in knowledge_bases
@@ -350,32 +352,32 @@ class KnowledgeTable:
)
]
def get_knowledge_by_id(self, id: str) -> Optional[KnowledgeModel]:
def get_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> Optional[KnowledgeModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
knowledge = db.query(Knowledge).filter_by(id=id).first()
return KnowledgeModel.model_validate(knowledge) if knowledge else None
except Exception:
return None
def get_knowledge_by_id_and_user_id(
self, id: str, user_id: str
self, id: str, user_id: str, db: Optional[Session] = None
) -> Optional[KnowledgeModel]:
knowledge = self.get_knowledge_by_id(id)
knowledge = self.get_knowledge_by_id(id, db=db)
if not knowledge:
return None
if knowledge.user_id == user_id:
return knowledge
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id)}
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user_id, db=db)}
if has_access(user_id, "write", knowledge.access_control, user_group_ids):
return knowledge
return None
def get_knowledges_by_file_id(self, file_id: str) -> list[KnowledgeModel]:
def get_knowledges_by_file_id(self, file_id: str, db: Optional[Session] = None) -> list[KnowledgeModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
knowledges = (
db.query(Knowledge)
.join(KnowledgeFile, Knowledge.id == KnowledgeFile.knowledge_id)
@@ -395,9 +397,10 @@ class KnowledgeTable:
filter: dict,
skip: int = 0,
limit: int = 30,
db: Optional[Session] = None,
) -> KnowledgeFileListResponse:
try:
with get_db() as db:
with get_db_context(db) as db:
query = (
db.query(File, User)
.join(KnowledgeFile, File.id == KnowledgeFile.file_id)
@@ -470,9 +473,9 @@ class KnowledgeTable:
print(e)
return KnowledgeFileListResponse(items=[], total=0)
def get_files_by_id(self, knowledge_id: str) -> list[FileModel]:
def get_files_by_id(self, knowledge_id: str, db: Optional[Session] = None) -> list[FileModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
files = (
db.query(File)
.join(KnowledgeFile, File.id == KnowledgeFile.file_id)
@@ -483,18 +486,18 @@ class KnowledgeTable:
except Exception:
return []
def get_file_metadatas_by_id(self, knowledge_id: str) -> list[FileMetadataResponse]:
def get_file_metadatas_by_id(self, knowledge_id: str, db: Optional[Session] = None) -> list[FileMetadataResponse]:
try:
with get_db() as db:
files = self.get_files_by_id(knowledge_id)
with get_db_context(db) as db:
files = self.get_files_by_id(knowledge_id, db=db)
return [FileMetadataResponse(**file.model_dump()) for file in files]
except Exception:
return []
def add_file_to_knowledge_by_id(
self, knowledge_id: str, file_id: str, user_id: str
self, knowledge_id: str, file_id: str, user_id: str, db: Optional[Session] = None
) -> Optional[KnowledgeFileModel]:
with get_db() as db:
with get_db_context(db) as db:
knowledge_file = KnowledgeFileModel(
**{
"id": str(uuid.uuid4()),
@@ -518,9 +521,9 @@ class KnowledgeTable:
except Exception:
return None
def remove_file_from_knowledge_by_id(self, knowledge_id: str, file_id: str) -> bool:
def remove_file_from_knowledge_by_id(self, knowledge_id: str, file_id: str, db: Optional[Session] = None) -> bool:
try:
with get_db() as db:
with get_db_context(db) as db:
db.query(KnowledgeFile).filter_by(
knowledge_id=knowledge_id, file_id=file_id
).delete()
@@ -529,9 +532,9 @@ class KnowledgeTable:
except Exception:
return False
def reset_knowledge_by_id(self, id: str) -> Optional[KnowledgeModel]:
def reset_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> Optional[KnowledgeModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
# Delete all knowledge_file entries for this knowledge_id
db.query(KnowledgeFile).filter_by(knowledge_id=id).delete()
db.commit()
@@ -544,17 +547,17 @@ class KnowledgeTable:
)
db.commit()
return self.get_knowledge_by_id(id=id)
return self.get_knowledge_by_id(id=id, db=db)
except Exception as e:
log.exception(e)
return None
def update_knowledge_by_id(
self, id: str, form_data: KnowledgeForm, overwrite: bool = False
self, id: str, form_data: KnowledgeForm, overwrite: bool = False, db: Optional[Session] = None
) -> Optional[KnowledgeModel]:
try:
with get_db() as db:
knowledge = self.get_knowledge_by_id(id=id)
with get_db_context(db) as db:
knowledge = self.get_knowledge_by_id(id=id, db=db)
db.query(Knowledge).filter_by(id=id).update(
{
**form_data.model_dump(),
@@ -562,17 +565,17 @@ class KnowledgeTable:
}
)
db.commit()
return self.get_knowledge_by_id(id=id)
return self.get_knowledge_by_id(id=id, db=db)
except Exception as e:
log.exception(e)
return None
def update_knowledge_data_by_id(
self, id: str, data: dict
self, id: str, data: dict, db: Optional[Session] = None
) -> Optional[KnowledgeModel]:
try:
with get_db() as db:
knowledge = self.get_knowledge_by_id(id=id)
with get_db_context(db) as db:
knowledge = self.get_knowledge_by_id(id=id, db=db)
db.query(Knowledge).filter_by(id=id).update(
{
"data": data,
@@ -580,22 +583,22 @@ class KnowledgeTable:
}
)
db.commit()
return self.get_knowledge_by_id(id=id)
return self.get_knowledge_by_id(id=id, db=db)
except Exception as e:
log.exception(e)
return None
def delete_knowledge_by_id(self, id: str) -> bool:
def delete_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> bool:
try:
with get_db() as db:
with get_db_context(db) as db:
db.query(Knowledge).filter_by(id=id).delete()
db.commit()
return True
except Exception:
return False
def delete_all_knowledge(self) -> bool:
with get_db() as db:
def delete_all_knowledge(self, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
try:
db.query(Knowledge).delete()
db.commit()