refac/enh: db session sharing
This commit is contained in:
@@ -62,6 +62,8 @@ from open_webui.utils.auth import (
|
||||
get_password_hash,
|
||||
get_http_authorization_cred,
|
||||
)
|
||||
from open_webui.internal.db import get_session
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.utils.webhook import post_webhook
|
||||
from open_webui.utils.access_control import get_permissions, has_permission
|
||||
from open_webui.utils.groups import apply_default_group_assignment
|
||||
@@ -103,7 +105,10 @@ class SessionUserInfoResponse(SessionUserResponse, UserStatus):
|
||||
|
||||
@router.get("/", response_model=SessionUserInfoResponse)
|
||||
async def get_session_user(
|
||||
request: Request, response: Response, user=Depends(get_current_user)
|
||||
request: Request,
|
||||
response: Response,
|
||||
user=Depends(get_current_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
|
||||
auth_header = request.headers.get("Authorization")
|
||||
@@ -137,7 +142,7 @@ async def get_session_user(
|
||||
)
|
||||
|
||||
user_permissions = get_permissions(
|
||||
user.id, request.app.state.config.USER_PERMISSIONS
|
||||
user.id, request.app.state.config.USER_PERMISSIONS, db=db
|
||||
)
|
||||
|
||||
return {
|
||||
@@ -166,12 +171,15 @@ async def get_session_user(
|
||||
|
||||
@router.post("/update/profile", response_model=UserProfileImageResponse)
|
||||
async def update_profile(
|
||||
form_data: UpdateProfileForm, session_user=Depends(get_verified_user)
|
||||
form_data: UpdateProfileForm,
|
||||
session_user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if session_user:
|
||||
user = Users.update_user_by_id(
|
||||
session_user.id,
|
||||
form_data.model_dump(),
|
||||
db=db,
|
||||
)
|
||||
if user:
|
||||
return user
|
||||
@@ -188,13 +196,17 @@ async def update_profile(
|
||||
|
||||
@router.post("/update/password", response_model=bool)
|
||||
async def update_password(
|
||||
form_data: UpdatePasswordForm, session_user=Depends(get_current_user)
|
||||
form_data: UpdatePasswordForm,
|
||||
session_user=Depends(get_current_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if WEBUI_AUTH_TRUSTED_EMAIL_HEADER:
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.ACTION_PROHIBITED)
|
||||
if session_user:
|
||||
user = Auths.authenticate_user(
|
||||
session_user.email, lambda pw: verify_password(form_data.password, pw)
|
||||
session_user.email,
|
||||
lambda pw: verify_password(form_data.password, pw),
|
||||
db=db,
|
||||
)
|
||||
|
||||
if user:
|
||||
@@ -203,7 +215,7 @@ async def update_password(
|
||||
except Exception as e:
|
||||
raise HTTPException(400, detail=str(e))
|
||||
hashed = get_password_hash(form_data.new_password)
|
||||
return Auths.update_user_password_by_id(user.id, hashed)
|
||||
return Auths.update_user_password_by_id(user.id, hashed, db=db)
|
||||
else:
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.INCORRECT_PASSWORD)
|
||||
else:
|
||||
@@ -214,7 +226,12 @@ async def update_password(
|
||||
# LDAP Authentication
|
||||
############################
|
||||
@router.post("/ldap", response_model=SessionUserResponse)
|
||||
async def ldap_auth(request: Request, response: Response, form_data: LdapForm):
|
||||
async def ldap_auth(
|
||||
request: Request,
|
||||
response: Response,
|
||||
form_data: LdapForm,
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
# Security checks FIRST - before loading any config
|
||||
if not request.app.state.config.ENABLE_LDAP:
|
||||
raise HTTPException(400, detail="LDAP authentication is not enabled")
|
||||
@@ -400,12 +417,12 @@ async def ldap_auth(request: Request, response: Response, form_data: LdapForm):
|
||||
if not connection_user.bind():
|
||||
raise HTTPException(400, "Authentication failed.")
|
||||
|
||||
user = Users.get_user_by_email(email)
|
||||
user = Users.get_user_by_email(email, db=db)
|
||||
if not user:
|
||||
try:
|
||||
role = (
|
||||
"admin"
|
||||
if not Users.has_users()
|
||||
if not Users.has_users(db=db)
|
||||
else request.app.state.config.DEFAULT_USER_ROLE
|
||||
)
|
||||
|
||||
@@ -414,6 +431,7 @@ async def ldap_auth(request: Request, response: Response, form_data: LdapForm):
|
||||
password=str(uuid.uuid4()),
|
||||
name=cn,
|
||||
role=role,
|
||||
db=db,
|
||||
)
|
||||
|
||||
if not user:
|
||||
@@ -424,6 +442,7 @@ async def ldap_auth(request: Request, response: Response, form_data: LdapForm):
|
||||
apply_default_group_assignment(
|
||||
request.app.state.config.DEFAULT_GROUP_ID,
|
||||
user.id,
|
||||
db=db,
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
@@ -434,7 +453,7 @@ async def ldap_auth(request: Request, response: Response, form_data: LdapForm):
|
||||
500, detail="Internal error occurred during LDAP user creation."
|
||||
)
|
||||
|
||||
user = Auths.authenticate_user_by_email(email)
|
||||
user = Auths.authenticate_user_by_email(email, db=db)
|
||||
|
||||
if user:
|
||||
expires_delta = parse_duration(request.app.state.config.JWT_EXPIRES_IN)
|
||||
@@ -464,7 +483,7 @@ async def ldap_auth(request: Request, response: Response, form_data: LdapForm):
|
||||
)
|
||||
|
||||
user_permissions = get_permissions(
|
||||
user.id, request.app.state.config.USER_PERMISSIONS
|
||||
user.id, request.app.state.config.USER_PERMISSIONS, db=db
|
||||
)
|
||||
|
||||
if (
|
||||
@@ -473,9 +492,9 @@ async def ldap_auth(request: Request, response: Response, form_data: LdapForm):
|
||||
and user_groups
|
||||
):
|
||||
if ENABLE_LDAP_GROUP_CREATION:
|
||||
Groups.create_groups_by_group_names(user.id, user_groups)
|
||||
Groups.create_groups_by_group_names(user.id, user_groups, db=db)
|
||||
try:
|
||||
Groups.sync_groups_by_group_names(user.id, user_groups)
|
||||
Groups.sync_groups_by_group_names(user.id, user_groups, db=db)
|
||||
log.info(
|
||||
f"Successfully synced groups for user {user.id}: {user_groups}"
|
||||
)
|
||||
@@ -508,7 +527,12 @@ async def ldap_auth(request: Request, response: Response, form_data: LdapForm):
|
||||
|
||||
|
||||
@router.post("/signin", response_model=SessionUserResponse)
|
||||
async def signin(request: Request, response: Response, form_data: SigninForm):
|
||||
async def signin(
|
||||
request: Request,
|
||||
response: Response,
|
||||
form_data: SigninForm,
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if not ENABLE_PASSWORD_AUTH:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
@@ -529,14 +553,15 @@ async def signin(request: Request, response: Response, form_data: SigninForm):
|
||||
except Exception as e:
|
||||
pass
|
||||
|
||||
if not Users.get_user_by_email(email.lower()):
|
||||
if not Users.get_user_by_email(email.lower(), db=db):
|
||||
await signup(
|
||||
request,
|
||||
response,
|
||||
SignupForm(email=email, password=str(uuid.uuid4()), name=name),
|
||||
db=db,
|
||||
)
|
||||
|
||||
user = Auths.authenticate_user_by_email(email)
|
||||
user = Auths.authenticate_user_by_email(email, db=db)
|
||||
if WEBUI_AUTH_TRUSTED_GROUPS_HEADER and user and user.role != "admin":
|
||||
group_names = request.headers.get(
|
||||
WEBUI_AUTH_TRUSTED_GROUPS_HEADER, ""
|
||||
@@ -544,28 +569,33 @@ async def signin(request: Request, response: Response, form_data: SigninForm):
|
||||
group_names = [name.strip() for name in group_names if name.strip()]
|
||||
|
||||
if group_names:
|
||||
Groups.sync_groups_by_group_names(user.id, group_names)
|
||||
Groups.sync_groups_by_group_names(user.id, group_names, db=db)
|
||||
|
||||
elif WEBUI_AUTH == False:
|
||||
admin_email = "admin@localhost"
|
||||
admin_password = "admin"
|
||||
|
||||
if Users.get_user_by_email(admin_email.lower()):
|
||||
if Users.get_user_by_email(admin_email.lower(), db=db):
|
||||
user = Auths.authenticate_user(
|
||||
admin_email.lower(), lambda pw: verify_password(admin_password, pw)
|
||||
admin_email.lower(),
|
||||
lambda pw: verify_password(admin_password, pw),
|
||||
db=db,
|
||||
)
|
||||
else:
|
||||
if Users.has_users():
|
||||
if Users.has_users(db=db):
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.EXISTING_USERS)
|
||||
|
||||
await signup(
|
||||
request,
|
||||
response,
|
||||
SignupForm(email=admin_email, password=admin_password, name="User"),
|
||||
db=db,
|
||||
)
|
||||
|
||||
user = Auths.authenticate_user(
|
||||
admin_email.lower(), lambda pw: verify_password(admin_password, pw)
|
||||
admin_email.lower(),
|
||||
lambda pw: verify_password(admin_password, pw),
|
||||
db=db,
|
||||
)
|
||||
else:
|
||||
if signin_rate_limiter.is_limited(form_data.email.lower()):
|
||||
@@ -584,7 +614,9 @@ async def signin(request: Request, response: Response, form_data: SigninForm):
|
||||
form_data.password = password_bytes.decode("utf-8", errors="ignore")
|
||||
|
||||
user = Auths.authenticate_user(
|
||||
form_data.email.lower(), lambda pw: verify_password(form_data.password, pw)
|
||||
form_data.email.lower(),
|
||||
lambda pw: verify_password(form_data.password, pw),
|
||||
db=db,
|
||||
)
|
||||
|
||||
if user:
|
||||
@@ -616,7 +648,7 @@ async def signin(request: Request, response: Response, form_data: SigninForm):
|
||||
)
|
||||
|
||||
user_permissions = get_permissions(
|
||||
user.id, request.app.state.config.USER_PERMISSIONS
|
||||
user.id, request.app.state.config.USER_PERMISSIONS, db=db
|
||||
)
|
||||
|
||||
return {
|
||||
@@ -640,8 +672,13 @@ async def signin(request: Request, response: Response, form_data: SigninForm):
|
||||
|
||||
|
||||
@router.post("/signup", response_model=SessionUserResponse)
|
||||
async def signup(request: Request, response: Response, form_data: SignupForm):
|
||||
has_users = Users.has_users()
|
||||
async def signup(
|
||||
request: Request,
|
||||
response: Response,
|
||||
form_data: SignupForm,
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
has_users = Users.has_users(db=db)
|
||||
|
||||
if WEBUI_AUTH:
|
||||
if (
|
||||
@@ -663,7 +700,7 @@ async def signup(request: Request, response: Response, form_data: SignupForm):
|
||||
status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.INVALID_EMAIL_FORMAT
|
||||
)
|
||||
|
||||
if Users.get_user_by_email(form_data.email.lower()):
|
||||
if Users.get_user_by_email(form_data.email.lower(), db=db):
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.EMAIL_TAKEN)
|
||||
|
||||
try:
|
||||
@@ -681,6 +718,7 @@ async def signup(request: Request, response: Response, form_data: SignupForm):
|
||||
form_data.name,
|
||||
form_data.profile_image_url,
|
||||
role,
|
||||
db=db,
|
||||
)
|
||||
|
||||
if user:
|
||||
@@ -723,7 +761,7 @@ async def signup(request: Request, response: Response, form_data: SignupForm):
|
||||
)
|
||||
|
||||
user_permissions = get_permissions(
|
||||
user.id, request.app.state.config.USER_PERMISSIONS
|
||||
user.id, request.app.state.config.USER_PERMISSIONS, db=db
|
||||
)
|
||||
|
||||
if not has_users:
|
||||
@@ -733,6 +771,7 @@ async def signup(request: Request, response: Response, form_data: SignupForm):
|
||||
apply_default_group_assignment(
|
||||
request.app.state.config.DEFAULT_GROUP_ID,
|
||||
user.id,
|
||||
db=db,
|
||||
)
|
||||
|
||||
return {
|
||||
@@ -754,7 +793,9 @@ async def signup(request: Request, response: Response, form_data: SignupForm):
|
||||
|
||||
|
||||
@router.get("/signout")
|
||||
async def signout(request: Request, response: Response):
|
||||
async def signout(
|
||||
request: Request, response: Response, db: Session = Depends(get_session)
|
||||
):
|
||||
|
||||
# get auth token from headers or cookies
|
||||
token = None
|
||||
@@ -776,7 +817,7 @@ async def signout(request: Request, response: Response):
|
||||
if oauth_session_id:
|
||||
response.delete_cookie("oauth_session_id")
|
||||
|
||||
session = OAuthSessions.get_session_by_id(oauth_session_id)
|
||||
session = OAuthSessions.get_session_by_id(oauth_session_id, db=db)
|
||||
oauth_server_metadata_url = (
|
||||
request.app.state.oauth_manager.get_server_metadata_url(session.provider)
|
||||
if session
|
||||
@@ -839,14 +880,17 @@ async def signout(request: Request, response: Response):
|
||||
|
||||
@router.post("/add", response_model=SigninResponse)
|
||||
async def add_user(
|
||||
request: Request, form_data: AddUserForm, user=Depends(get_admin_user)
|
||||
request: Request,
|
||||
form_data: AddUserForm,
|
||||
user=Depends(get_admin_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if not validate_email_format(form_data.email.lower()):
|
||||
raise HTTPException(
|
||||
status.HTTP_400_BAD_REQUEST, detail=ERROR_MESSAGES.INVALID_EMAIL_FORMAT
|
||||
)
|
||||
|
||||
if Users.get_user_by_email(form_data.email.lower()):
|
||||
if Users.get_user_by_email(form_data.email.lower(), db=db):
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.EMAIL_TAKEN)
|
||||
|
||||
try:
|
||||
@@ -862,12 +906,14 @@ async def add_user(
|
||||
form_data.name,
|
||||
form_data.profile_image_url,
|
||||
form_data.role,
|
||||
db=db,
|
||||
)
|
||||
|
||||
if user:
|
||||
apply_default_group_assignment(
|
||||
request.app.state.config.DEFAULT_GROUP_ID,
|
||||
user.id,
|
||||
db=db,
|
||||
)
|
||||
|
||||
token = create_token(data={"id": user.id})
|
||||
@@ -895,7 +941,9 @@ async def add_user(
|
||||
|
||||
|
||||
@router.get("/admin/details")
|
||||
async def get_admin_details(request: Request, user=Depends(get_current_user)):
|
||||
async def get_admin_details(
|
||||
request: Request, user=Depends(get_current_user), db: Session = Depends(get_session)
|
||||
):
|
||||
if request.app.state.config.SHOW_ADMIN_DETAILS:
|
||||
admin_email = request.app.state.config.ADMIN_EMAIL
|
||||
admin_name = None
|
||||
@@ -903,11 +951,11 @@ async def get_admin_details(request: Request, user=Depends(get_current_user)):
|
||||
log.info(f"Admin details - Email: {admin_email}, Name: {admin_name}")
|
||||
|
||||
if admin_email:
|
||||
admin = Users.get_user_by_email(admin_email)
|
||||
admin = Users.get_user_by_email(admin_email, db=db)
|
||||
if admin:
|
||||
admin_name = admin.name
|
||||
else:
|
||||
admin = Users.get_first_user()
|
||||
admin = Users.get_first_user(db=db)
|
||||
if admin:
|
||||
admin_email = admin.email
|
||||
admin_name = admin.name
|
||||
@@ -1149,7 +1197,9 @@ async def update_ldap_config(
|
||||
|
||||
# create api key
|
||||
@router.post("/api_key", response_model=ApiKey)
|
||||
async def generate_api_key(request: Request, user=Depends(get_current_user)):
|
||||
async def generate_api_key(
|
||||
request: Request, user=Depends(get_current_user), db: Session = Depends(get_session)
|
||||
):
|
||||
if not request.app.state.config.ENABLE_API_KEYS or not has_permission(
|
||||
user.id, "features.api_keys", request.app.state.config.USER_PERMISSIONS
|
||||
):
|
||||
@@ -1159,7 +1209,7 @@ async def generate_api_key(request: Request, user=Depends(get_current_user)):
|
||||
)
|
||||
|
||||
api_key = create_api_key()
|
||||
success = Users.update_user_api_key_by_id(user.id, api_key)
|
||||
success = Users.update_user_api_key_by_id(user.id, api_key, db=db)
|
||||
|
||||
if success:
|
||||
return {
|
||||
@@ -1171,14 +1221,18 @@ async def generate_api_key(request: Request, user=Depends(get_current_user)):
|
||||
|
||||
# delete api key
|
||||
@router.delete("/api_key", response_model=bool)
|
||||
async def delete_api_key(user=Depends(get_current_user)):
|
||||
return Users.delete_user_api_key_by_id(user.id)
|
||||
async def delete_api_key(
|
||||
user=Depends(get_current_user), db: Session = Depends(get_session)
|
||||
):
|
||||
return Users.delete_user_api_key_by_id(user.id, db=db)
|
||||
|
||||
|
||||
# get api key
|
||||
@router.get("/api_key", response_model=ApiKey)
|
||||
async def get_api_key(user=Depends(get_current_user)):
|
||||
api_key = Users.get_user_api_key_by_id(user.id)
|
||||
async def get_api_key(
|
||||
user=Depends(get_current_user), db: Session = Depends(get_session)
|
||||
):
|
||||
api_key = Users.get_user_api_key_by_id(user.id, db=db)
|
||||
if api_key:
|
||||
return {
|
||||
"api_key": api_key,
|
||||
|
||||
Reference in New Issue
Block a user