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
+28 -27
View File
@@ -3,7 +3,8 @@ import time
import uuid
from typing import Optional
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.users import User
from pydantic import BaseModel, ConfigDict
@@ -121,9 +122,9 @@ class FeedbackListResponse(BaseModel):
class FeedbackTable:
def insert_new_feedback(
self, user_id: str, form_data: FeedbackForm
self, user_id: str, form_data: FeedbackForm, db: Optional[Session] = None
) -> Optional[FeedbackModel]:
with get_db() as db:
with get_db_context(db) as db:
id = str(uuid.uuid4())
feedback = FeedbackModel(
**{
@@ -148,9 +149,9 @@ class FeedbackTable:
log.exception(f"Error creating a new feedback: {e}")
return None
def get_feedback_by_id(self, id: str) -> Optional[FeedbackModel]:
def get_feedback_by_id(self, id: str, db: Optional[Session] = None) -> Optional[FeedbackModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id).first()
if not feedback:
return None
@@ -159,10 +160,10 @@ class FeedbackTable:
return None
def get_feedback_by_id_and_user_id(
self, id: str, user_id: str
self, id: str, user_id: str, db: Optional[Session] = None
) -> Optional[FeedbackModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first()
if not feedback:
return None
@@ -171,9 +172,9 @@ class FeedbackTable:
return None
def get_feedback_items(
self, filter: dict = {}, skip: int = 0, limit: int = 30
self, filter: dict = {}, skip: int = 0, limit: int = 30, db: Optional[Session] = None
) -> FeedbackListResponse:
with get_db() as db:
with get_db_context(db) as db:
query = db.query(Feedback, User).join(User, Feedback.user_id == User.id)
if filter:
@@ -234,8 +235,8 @@ class FeedbackTable:
return FeedbackListResponse(items=feedbacks, total=total)
def get_all_feedbacks(self) -> list[FeedbackModel]:
with get_db() as db:
def get_all_feedbacks(self, db: Optional[Session] = None) -> list[FeedbackModel]:
with get_db_context(db) as db:
return [
FeedbackModel.model_validate(feedback)
for feedback in db.query(Feedback)
@@ -243,8 +244,8 @@ class FeedbackTable:
.all()
]
def get_feedbacks_by_type(self, type: str) -> list[FeedbackModel]:
with get_db() as db:
def get_feedbacks_by_type(self, type: str, db: Optional[Session] = None) -> list[FeedbackModel]:
with get_db_context(db) as db:
return [
FeedbackModel.model_validate(feedback)
for feedback in db.query(Feedback)
@@ -253,8 +254,8 @@ class FeedbackTable:
.all()
]
def get_feedbacks_by_user_id(self, user_id: str) -> list[FeedbackModel]:
with get_db() as db:
def get_feedbacks_by_user_id(self, user_id: str, db: Optional[Session] = None) -> list[FeedbackModel]:
with get_db_context(db) as db:
return [
FeedbackModel.model_validate(feedback)
for feedback in db.query(Feedback)
@@ -264,9 +265,9 @@ class FeedbackTable:
]
def update_feedback_by_id(
self, id: str, form_data: FeedbackForm
self, id: str, form_data: FeedbackForm, db: Optional[Session] = None
) -> Optional[FeedbackModel]:
with get_db() as db:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id).first()
if not feedback:
return None
@@ -284,9 +285,9 @@ class FeedbackTable:
return FeedbackModel.model_validate(feedback)
def update_feedback_by_id_and_user_id(
self, id: str, user_id: str, form_data: FeedbackForm
self, id: str, user_id: str, form_data: FeedbackForm, db: Optional[Session] = None
) -> Optional[FeedbackModel]:
with get_db() as db:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first()
if not feedback:
return None
@@ -303,8 +304,8 @@ class FeedbackTable:
db.commit()
return FeedbackModel.model_validate(feedback)
def delete_feedback_by_id(self, id: str) -> bool:
with get_db() as db:
def delete_feedback_by_id(self, id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id).first()
if not feedback:
return False
@@ -312,8 +313,8 @@ class FeedbackTable:
db.commit()
return True
def delete_feedback_by_id_and_user_id(self, id: str, user_id: str) -> bool:
with get_db() as db:
def delete_feedback_by_id_and_user_id(self, id: str, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first()
if not feedback:
return False
@@ -321,8 +322,8 @@ class FeedbackTable:
db.commit()
return True
def delete_feedbacks_by_user_id(self, user_id: str) -> bool:
with get_db() as db:
def delete_feedbacks_by_user_id(self, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
feedbacks = db.query(Feedback).filter_by(user_id=user_id).all()
if not feedbacks:
return False
@@ -331,8 +332,8 @@ class FeedbackTable:
db.commit()
return True
def delete_all_feedbacks(self) -> bool:
with get_db() as db:
def delete_all_feedbacks(self, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
feedbacks = db.query(Feedback).all()
if not feedbacks:
return False