refac/enh: db session sharing
This commit is contained in:
@@ -30,6 +30,8 @@ from fastapi.responses import FileResponse, StreamingResponse
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.access_control import has_access, has_permission
|
||||
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL, STATIC_DIR
|
||||
from open_webui.internal.db import get_session
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
@@ -59,6 +61,7 @@ async def get_models(
|
||||
direction: Optional[str] = None,
|
||||
page: Optional[int] = 1,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
|
||||
limit = PAGE_ITEM_COUNT
|
||||
@@ -79,13 +82,13 @@ async def get_models(
|
||||
filter["direction"] = direction
|
||||
|
||||
if not user.role == "admin" or not BYPASS_ADMIN_ACCESS_CONTROL:
|
||||
groups = Groups.get_groups_by_member_id(user.id)
|
||||
groups = Groups.get_groups_by_member_id(user.id, db=db)
|
||||
if groups:
|
||||
filter["group_ids"] = [group.id for group in groups]
|
||||
|
||||
filter["user_id"] = user.id
|
||||
|
||||
return Models.search_models(user.id, filter=filter, skip=skip, limit=limit)
|
||||
return Models.search_models(user.id, filter=filter, skip=skip, limit=limit, db=db)
|
||||
|
||||
|
||||
###########################
|
||||
@@ -94,8 +97,8 @@ async def get_models(
|
||||
|
||||
|
||||
@router.get("/base", response_model=list[ModelResponse])
|
||||
async def get_base_models(user=Depends(get_admin_user)):
|
||||
return Models.get_base_models()
|
||||
async def get_base_models(user=Depends(get_admin_user), db: Session = Depends(get_session)):
|
||||
return Models.get_base_models(db=db)
|
||||
|
||||
|
||||
###########################
|
||||
@@ -104,11 +107,11 @@ async def get_base_models(user=Depends(get_admin_user)):
|
||||
|
||||
|
||||
@router.get("/tags", response_model=list[str])
|
||||
async def get_model_tags(user=Depends(get_verified_user)):
|
||||
async def get_model_tags(user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
if user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL:
|
||||
models = Models.get_models()
|
||||
models = Models.get_models(db=db)
|
||||
else:
|
||||
models = Models.get_models_by_user_id(user.id)
|
||||
models = Models.get_models_by_user_id(user.id, db=db)
|
||||
|
||||
tags_set = set()
|
||||
for model in models:
|
||||
@@ -132,16 +135,17 @@ async def create_new_model(
|
||||
request: Request,
|
||||
form_data: ModelForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if user.role != "admin" and not has_permission(
|
||||
user.id, "workspace.models", request.app.state.config.USER_PERMISSIONS
|
||||
user.id, "workspace.models", request.app.state.config.USER_PERMISSIONS, db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.UNAUTHORIZED,
|
||||
)
|
||||
|
||||
model = Models.get_model_by_id(form_data.id)
|
||||
model = Models.get_model_by_id(form_data.id, db=db)
|
||||
if model:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -155,7 +159,7 @@ async def create_new_model(
|
||||
)
|
||||
|
||||
else:
|
||||
model = Models.insert_new_model(form_data, user.id)
|
||||
model = Models.insert_new_model(form_data, user.id, db=db)
|
||||
if model:
|
||||
return model
|
||||
else:
|
||||
@@ -171,9 +175,9 @@ async def create_new_model(
|
||||
|
||||
|
||||
@router.get("/export", response_model=list[ModelModel])
|
||||
async def export_models(request: Request, user=Depends(get_verified_user)):
|
||||
async def export_models(request: Request, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
if user.role != "admin" and not has_permission(
|
||||
user.id, "workspace.models_export", request.app.state.config.USER_PERMISSIONS
|
||||
user.id, "workspace.models_export", request.app.state.config.USER_PERMISSIONS, db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -181,9 +185,9 @@ async def export_models(request: Request, user=Depends(get_verified_user)):
|
||||
)
|
||||
|
||||
if user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL:
|
||||
return Models.get_models()
|
||||
return Models.get_models(db=db)
|
||||
else:
|
||||
return Models.get_models_by_user_id(user.id)
|
||||
return Models.get_models_by_user_id(user.id, db=db)
|
||||
|
||||
|
||||
############################
|
||||
@@ -200,9 +204,10 @@ async def import_models(
|
||||
request: Request,
|
||||
user=Depends(get_verified_user),
|
||||
form_data: ModelsImportForm = (...),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if user.role != "admin" and not has_permission(
|
||||
user.id, "workspace.models_import", request.app.state.config.USER_PERMISSIONS
|
||||
user.id, "workspace.models_import", request.app.state.config.USER_PERMISSIONS, db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -216,7 +221,7 @@ async def import_models(
|
||||
model_id = model_data.get("id")
|
||||
|
||||
if model_id and is_valid_model_id(model_id):
|
||||
existing_model = Models.get_model_by_id(model_id)
|
||||
existing_model = Models.get_model_by_id(model_id, db=db)
|
||||
if existing_model:
|
||||
# Update existing model
|
||||
model_data["meta"] = model_data.get("meta", {})
|
||||
@@ -225,13 +230,13 @@ async def import_models(
|
||||
updated_model = ModelForm(
|
||||
**{**existing_model.model_dump(), **model_data}
|
||||
)
|
||||
Models.update_model_by_id(model_id, updated_model)
|
||||
Models.update_model_by_id(model_id, updated_model, db=db)
|
||||
else:
|
||||
# Insert new model
|
||||
model_data["meta"] = model_data.get("meta", {})
|
||||
model_data["params"] = model_data.get("params", {})
|
||||
new_model = ModelForm(**model_data)
|
||||
Models.insert_new_model(user_id=user.id, form_data=new_model)
|
||||
Models.insert_new_model(user_id=user.id, form_data=new_model, db=db)
|
||||
return True
|
||||
else:
|
||||
raise HTTPException(status_code=400, detail="Invalid JSON format")
|
||||
@@ -251,9 +256,9 @@ class SyncModelsForm(BaseModel):
|
||||
|
||||
@router.post("/sync", response_model=list[ModelModel])
|
||||
async def sync_models(
|
||||
request: Request, form_data: SyncModelsForm, user=Depends(get_admin_user)
|
||||
request: Request, form_data: SyncModelsForm, user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
return Models.sync_models(user.id, form_data.models)
|
||||
return Models.sync_models(user.id, form_data.models, db=db)
|
||||
|
||||
|
||||
###########################
|
||||
@@ -267,13 +272,13 @@ class ModelIdForm(BaseModel):
|
||||
|
||||
# Note: We're not using the typical url path param here, but instead using a query parameter to allow '/' in the id
|
||||
@router.get("/model", response_model=Optional[ModelResponse])
|
||||
async def get_model_by_id(id: str, user=Depends(get_verified_user)):
|
||||
model = Models.get_model_by_id(id)
|
||||
async def get_model_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
model = Models.get_model_by_id(id, db=db)
|
||||
if model:
|
||||
if (
|
||||
(user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL)
|
||||
or model.user_id == user.id
|
||||
or has_access(user.id, "read", model.access_control)
|
||||
or has_access(user.id, "read", model.access_control, db=db)
|
||||
):
|
||||
return model
|
||||
else:
|
||||
@@ -289,8 +294,8 @@ async def get_model_by_id(id: str, user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.get("/model/profile/image")
|
||||
async def get_model_profile_image(id: str, user=Depends(get_verified_user)):
|
||||
model = Models.get_model_by_id(id)
|
||||
async def get_model_profile_image(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
model = Models.get_model_by_id(id, db=db)
|
||||
# Cache-control headers to prevent stale cached images
|
||||
cache_headers = {"Cache-Control": "no-cache, must-revalidate"}
|
||||
|
||||
@@ -330,15 +335,15 @@ async def get_model_profile_image(id: str, user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.post("/model/toggle", response_model=Optional[ModelResponse])
|
||||
async def toggle_model_by_id(id: str, user=Depends(get_verified_user)):
|
||||
model = Models.get_model_by_id(id)
|
||||
async def toggle_model_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
model = Models.get_model_by_id(id, db=db)
|
||||
if model:
|
||||
if (
|
||||
user.role == "admin"
|
||||
or model.user_id == user.id
|
||||
or has_access(user.id, "write", model.access_control)
|
||||
or has_access(user.id, "write", model.access_control, db=db)
|
||||
):
|
||||
model = Models.toggle_model_by_id(id)
|
||||
model = Models.toggle_model_by_id(id, db=db)
|
||||
|
||||
if model:
|
||||
return model
|
||||
@@ -368,8 +373,9 @@ async def toggle_model_by_id(id: str, user=Depends(get_verified_user)):
|
||||
async def update_model_by_id(
|
||||
form_data: ModelForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
model = Models.get_model_by_id(form_data.id)
|
||||
model = Models.get_model_by_id(form_data.id, db=db)
|
||||
if not model:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -378,7 +384,7 @@ async def update_model_by_id(
|
||||
|
||||
if (
|
||||
model.user_id != user.id
|
||||
and not has_access(user.id, "write", model.access_control)
|
||||
and not has_access(user.id, "write", model.access_control, db=db)
|
||||
and user.role != "admin"
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -386,7 +392,7 @@ async def update_model_by_id(
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
model = Models.update_model_by_id(form_data.id, ModelForm(**form_data.model_dump()))
|
||||
model = Models.update_model_by_id(form_data.id, ModelForm(**form_data.model_dump()), db=db)
|
||||
return model
|
||||
|
||||
|
||||
@@ -396,8 +402,8 @@ async def update_model_by_id(
|
||||
|
||||
|
||||
@router.post("/model/delete", response_model=bool)
|
||||
async def delete_model_by_id(form_data: ModelIdForm, user=Depends(get_verified_user)):
|
||||
model = Models.get_model_by_id(form_data.id)
|
||||
async def delete_model_by_id(form_data: ModelIdForm, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
model = Models.get_model_by_id(form_data.id, db=db)
|
||||
if not model:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -407,18 +413,18 @@ async def delete_model_by_id(form_data: ModelIdForm, user=Depends(get_verified_u
|
||||
if (
|
||||
user.role != "admin"
|
||||
and model.user_id != user.id
|
||||
and not has_access(user.id, "write", model.access_control)
|
||||
and not has_access(user.id, "write", model.access_control, db=db)
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.UNAUTHORIZED,
|
||||
)
|
||||
|
||||
result = Models.delete_model_by_id(form_data.id)
|
||||
result = Models.delete_model_by_id(form_data.id, db=db)
|
||||
return result
|
||||
|
||||
|
||||
@router.delete("/delete/all", response_model=bool)
|
||||
async def delete_all_models(user=Depends(get_admin_user)):
|
||||
result = Models.delete_all_models()
|
||||
async def delete_all_models(user=Depends(get_admin_user), db: Session = Depends(get_session)):
|
||||
result = Models.delete_all_models(db=db)
|
||||
return result
|
||||
|
||||
Reference in New Issue
Block a user