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
+59 -56
View File
@@ -1,7 +1,8 @@
import time
from typing import Optional
from open_webui.internal.db import Base, JSONField, get_db
from sqlalchemy.orm import Session
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
from open_webui.env import DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL
@@ -243,8 +244,9 @@ class UsersTable:
profile_image_url: str = "/user.png",
role: str = "pending",
oauth: Optional[dict] = None,
db: Optional[Session] = None,
) -> Optional[UserModel]:
with get_db() as db:
with get_db_context(db) as db:
user = UserModel(
**{
"id": id,
@@ -267,17 +269,17 @@ class UsersTable:
else:
return None
def get_user_by_id(self, id: str) -> Optional[UserModel]:
def get_user_by_id(self, id: str, db: Optional[Session] = None) -> Optional[UserModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
return UserModel.model_validate(user)
except Exception:
return None
def get_user_by_api_key(self, api_key: str) -> Optional[UserModel]:
def get_user_by_api_key(self, api_key: str, db: Optional[Session] = None) -> Optional[UserModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
user = (
db.query(User)
.join(ApiKey, User.id == ApiKey.user_id)
@@ -288,17 +290,17 @@ class UsersTable:
except Exception:
return None
def get_user_by_email(self, email: str) -> Optional[UserModel]:
def get_user_by_email(self, email: str, db: Optional[Session] = None) -> Optional[UserModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
user = db.query(User).filter_by(email=email).first()
return UserModel.model_validate(user)
except Exception:
return None
def get_user_by_oauth_sub(self, provider: str, sub: str) -> Optional[UserModel]:
def get_user_by_oauth_sub(self, provider: str, sub: str, db: Optional[Session] = None) -> Optional[UserModel]:
try:
with get_db() as db: # type: Session
with get_db_context(db) as db: # type: Session
dialect_name = db.bind.dialect.name
query = db.query(User)
@@ -320,8 +322,9 @@ class UsersTable:
filter: Optional[dict] = None,
skip: Optional[int] = None,
limit: Optional[int] = None,
db: Optional[Session] = None,
) -> dict:
with get_db() as db:
with get_db_context(db) as db:
# Join GroupMember so we can order by group_id when requested
query = db.query(User)
@@ -452,8 +455,8 @@ class UsersTable:
"total": total,
}
def get_users_by_group_id(self, group_id: str) -> list[UserModel]:
with get_db() as db:
def get_users_by_group_id(self, group_id: str, db: Optional[Session] = None) -> list[UserModel]:
with get_db_context(db) as db:
users = (
db.query(User)
.join(GroupMember, User.id == GroupMember.user_id)
@@ -462,30 +465,30 @@ class UsersTable:
)
return [UserModel.model_validate(user) for user in users]
def get_users_by_user_ids(self, user_ids: list[str]) -> list[UserStatusModel]:
with get_db() as db:
def get_users_by_user_ids(self, user_ids: list[str], db: Optional[Session] = None) -> list[UserStatusModel]:
with get_db_context(db) as db:
users = db.query(User).filter(User.id.in_(user_ids)).all()
return [UserModel.model_validate(user) for user in users]
def get_num_users(self) -> Optional[int]:
with get_db() as db:
def get_num_users(self, db: Optional[Session] = None) -> Optional[int]:
with get_db_context(db) as db:
return db.query(User).count()
def has_users(self) -> bool:
with get_db() as db:
def has_users(self, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
return db.query(db.query(User).exists()).scalar()
def get_first_user(self) -> UserModel:
def get_first_user(self, db: Optional[Session] = None) -> UserModel:
try:
with get_db() as db:
with get_db_context(db) as db:
user = db.query(User).order_by(User.created_at).first()
return UserModel.model_validate(user)
except Exception:
return None
def get_user_webhook_url_by_id(self, id: str) -> Optional[str]:
def get_user_webhook_url_by_id(self, id: str, db: Optional[Session] = None) -> Optional[str]:
try:
with get_db() as db:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
if user.settings is None:
@@ -499,8 +502,8 @@ class UsersTable:
except Exception:
return None
def get_num_users_active_today(self) -> Optional[int]:
with get_db() as db:
def get_num_users_active_today(self, db: Optional[Session] = None) -> Optional[int]:
with get_db_context(db) as db:
current_timestamp = int(datetime.datetime.now().timestamp())
today_midnight_timestamp = current_timestamp - (current_timestamp % 86400)
query = db.query(User).filter(
@@ -508,9 +511,9 @@ class UsersTable:
)
return query.count()
def update_user_role_by_id(self, id: str, role: str) -> Optional[UserModel]:
def update_user_role_by_id(self, id: str, role: str, db: Optional[Session] = None) -> Optional[UserModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
db.query(User).filter_by(id=id).update({"role": role})
db.commit()
user = db.query(User).filter_by(id=id).first()
@@ -519,10 +522,10 @@ class UsersTable:
return None
def update_user_status_by_id(
self, id: str, form_data: UserStatus
self, id: str, form_data: UserStatus, db: Optional[Session] = None
) -> Optional[UserModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
db.query(User).filter_by(id=id).update(
{**form_data.model_dump(exclude_none=True)}
)
@@ -534,10 +537,10 @@ class UsersTable:
return None
def update_user_profile_image_url_by_id(
self, id: str, profile_image_url: str
self, id: str, profile_image_url: str, db: Optional[Session] = None
) -> Optional[UserModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
db.query(User).filter_by(id=id).update(
{"profile_image_url": profile_image_url}
)
@@ -549,9 +552,9 @@ class UsersTable:
return None
@throttle(DATABASE_USER_ACTIVE_STATUS_UPDATE_INTERVAL)
def update_last_active_by_id(self, id: str) -> Optional[UserModel]:
def update_last_active_by_id(self, id: str, db: Optional[Session] = None) -> Optional[UserModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
db.query(User).filter_by(id=id).update(
{"last_active_at": int(time.time())}
)
@@ -563,7 +566,7 @@ class UsersTable:
return None
def update_user_oauth_by_id(
self, id: str, provider: str, sub: str
self, id: str, provider: str, sub: str, db: Optional[Session] = None
) -> Optional[UserModel]:
"""
Update or insert an OAuth provider/sub pair into the user's oauth JSON field.
@@ -574,7 +577,7 @@ class UsersTable:
}
"""
try:
with get_db() as db:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=id).first()
if not user:
return None
@@ -594,9 +597,9 @@ class UsersTable:
except Exception:
return None
def update_user_by_id(self, id: str, updated: dict) -> Optional[UserModel]:
def update_user_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[UserModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
db.query(User).filter_by(id=id).update(updated)
db.commit()
@@ -607,9 +610,9 @@ class UsersTable:
print(e)
return None
def update_user_settings_by_id(self, id: str, updated: dict) -> Optional[UserModel]:
def update_user_settings_by_id(self, id: str, updated: dict, db: Optional[Session] = None) -> Optional[UserModel]:
try:
with get_db() as db:
with get_db_context(db) as db:
user_settings = db.query(User).filter_by(id=id).first().settings
if user_settings is None:
@@ -625,15 +628,15 @@ class UsersTable:
except Exception:
return None
def delete_user_by_id(self, id: str) -> bool:
def delete_user_by_id(self, id: str, db: Optional[Session] = None) -> bool:
try:
# Remove User from Groups
Groups.remove_user_from_all_groups(id)
# Delete User Chats
result = Chats.delete_chats_by_user_id(id)
result = Chats.delete_chats_by_user_id(id, db=db)
if result:
with get_db() as db:
with get_db_context(db) as db:
# Delete User
db.query(User).filter_by(id=id).delete()
db.commit()
@@ -644,17 +647,17 @@ class UsersTable:
except Exception:
return False
def get_user_api_key_by_id(self, id: str) -> Optional[str]:
def get_user_api_key_by_id(self, id: str, db: Optional[Session] = None) -> Optional[str]:
try:
with get_db() as db:
with get_db_context(db) as db:
api_key = db.query(ApiKey).filter_by(user_id=id).first()
return api_key.key if api_key else None
except Exception:
return None
def update_user_api_key_by_id(self, id: str, api_key: str) -> bool:
def update_user_api_key_by_id(self, id: str, api_key: str, db: Optional[Session] = None) -> bool:
try:
with get_db() as db:
with get_db_context(db) as db:
db.query(ApiKey).filter_by(user_id=id).delete()
db.commit()
@@ -674,30 +677,30 @@ class UsersTable:
except Exception:
return False
def delete_user_api_key_by_id(self, id: str) -> bool:
def delete_user_api_key_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(ApiKey).filter_by(user_id=id).delete()
db.commit()
return True
except Exception:
return False
def get_valid_user_ids(self, user_ids: list[str]) -> list[str]:
with get_db() as db:
def get_valid_user_ids(self, user_ids: list[str], db: Optional[Session] = None) -> list[str]:
with get_db_context(db) as db:
users = db.query(User).filter(User.id.in_(user_ids)).all()
return [user.id for user in users]
def get_super_admin_user(self) -> Optional[UserModel]:
with get_db() as db:
def get_super_admin_user(self, db: Optional[Session] = None) -> Optional[UserModel]:
with get_db_context(db) as db:
user = db.query(User).filter_by(role="admin").first()
if user:
return UserModel.model_validate(user)
else:
return None
def get_active_user_count(self) -> int:
with get_db() as db:
def get_active_user_count(self, db: Optional[Session] = None) -> int:
with get_db_context(db) as db:
# Consider user active if last_active_at within the last 3 minutes
three_minutes_ago = int(time.time()) - 180
count = (
@@ -705,8 +708,8 @@ class UsersTable:
)
return count
def is_user_active(self, user_id: str) -> bool:
with get_db() as db:
def is_user_active(self, user_id: str, db: Optional[Session] = None) -> bool:
with get_db_context(db) as db:
user = db.query(User).filter_by(id=user_id).first()
if user and user.last_active_at:
# Consider user active if last_active_at within the last 3 minutes