refac/enh: db session sharing
This commit is contained in:
@@ -7,6 +7,8 @@ import aiohttp
|
||||
from open_webui.models.groups import Groups
|
||||
from pydantic import BaseModel, HttpUrl
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import get_session
|
||||
|
||||
|
||||
from open_webui.models.oauth_sessions import OAuthSessions
|
||||
@@ -51,11 +53,11 @@ def get_tool_module(request, tool_id, load_from_db=True):
|
||||
|
||||
|
||||
@router.get("/", response_model=list[ToolUserResponse])
|
||||
async def get_tools(request: Request, user=Depends(get_verified_user)):
|
||||
async def get_tools(request: Request, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
tools = []
|
||||
|
||||
# Local Tools
|
||||
for tool in Tools.get_tools():
|
||||
for tool in Tools.get_tools(db=db):
|
||||
tool_module = get_tool_module(request, tool.id)
|
||||
tools.append(
|
||||
ToolUserResponse(
|
||||
@@ -140,12 +142,12 @@ async def get_tools(request: Request, user=Depends(get_verified_user)):
|
||||
# Admin can see all tools
|
||||
return tools
|
||||
else:
|
||||
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)}
|
||||
tools = [
|
||||
tool
|
||||
for tool in tools
|
||||
if tool.user_id == user.id
|
||||
or has_access(user.id, "read", tool.access_control, user_group_ids)
|
||||
or has_access(user.id, "read", tool.access_control, user_group_ids, db=db)
|
||||
]
|
||||
return tools
|
||||
|
||||
@@ -156,11 +158,11 @@ async def get_tools(request: Request, user=Depends(get_verified_user)):
|
||||
|
||||
|
||||
@router.get("/list", response_model=list[ToolUserResponse])
|
||||
async def get_tool_list(user=Depends(get_verified_user)):
|
||||
async def get_tool_list(user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
if user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL:
|
||||
tools = Tools.get_tools()
|
||||
tools = Tools.get_tools(db=db)
|
||||
else:
|
||||
tools = Tools.get_tools_by_user_id(user.id, "write")
|
||||
tools = Tools.get_tools_by_user_id(user.id, "write", db=db)
|
||||
return tools
|
||||
|
||||
|
||||
@@ -245,9 +247,9 @@ async def load_tool_from_url(
|
||||
|
||||
|
||||
@router.get("/export", response_model=list[ToolModel])
|
||||
async def export_tools(request: Request, user=Depends(get_verified_user)):
|
||||
async def export_tools(request: Request, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
if user.role != "admin" and not has_permission(
|
||||
user.id, "workspace.tools_export", request.app.state.config.USER_PERMISSIONS
|
||||
user.id, "workspace.tools_export", request.app.state.config.USER_PERMISSIONS, db=db
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -255,9 +257,9 @@ async def export_tools(request: Request, user=Depends(get_verified_user)):
|
||||
)
|
||||
|
||||
if user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL:
|
||||
return Tools.get_tools()
|
||||
return Tools.get_tools(db=db)
|
||||
else:
|
||||
return Tools.get_tools_by_user_id(user.id, "read")
|
||||
return Tools.get_tools_by_user_id(user.id, "read", db=db)
|
||||
|
||||
|
||||
############################
|
||||
@@ -270,13 +272,14 @@ async def create_new_tools(
|
||||
request: Request,
|
||||
form_data: ToolForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if user.role != "admin" and not (
|
||||
has_permission(
|
||||
user.id, "workspace.tools", request.app.state.config.USER_PERMISSIONS
|
||||
user.id, "workspace.tools", request.app.state.config.USER_PERMISSIONS, db=db
|
||||
)
|
||||
or has_permission(
|
||||
user.id, "workspace.tools_import", request.app.state.config.USER_PERMISSIONS
|
||||
user.id, "workspace.tools_import", request.app.state.config.USER_PERMISSIONS, db=db
|
||||
)
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -292,7 +295,7 @@ async def create_new_tools(
|
||||
|
||||
form_data.id = form_data.id.lower()
|
||||
|
||||
tools = Tools.get_tool_by_id(form_data.id)
|
||||
tools = Tools.get_tool_by_id(form_data.id, db=db)
|
||||
if tools is None:
|
||||
try:
|
||||
form_data.content = replace_imports(form_data.content)
|
||||
@@ -305,7 +308,7 @@ async def create_new_tools(
|
||||
TOOLS[form_data.id] = tool_module
|
||||
|
||||
specs = get_tool_specs(TOOLS[form_data.id])
|
||||
tools = Tools.insert_new_tool(user.id, form_data, specs)
|
||||
tools = Tools.insert_new_tool(user.id, form_data, specs, db=db)
|
||||
|
||||
tool_cache_dir = CACHE_DIR / "tools" / form_data.id
|
||||
tool_cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
@@ -336,14 +339,14 @@ async def create_new_tools(
|
||||
|
||||
|
||||
@router.get("/id/{id}", response_model=Optional[ToolModel])
|
||||
async def get_tools_by_id(id: str, user=Depends(get_verified_user)):
|
||||
tools = Tools.get_tool_by_id(id)
|
||||
async def get_tools_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
tools = Tools.get_tool_by_id(id, db=db)
|
||||
|
||||
if tools:
|
||||
if (
|
||||
user.role == "admin"
|
||||
or tools.user_id == user.id
|
||||
or has_access(user.id, "read", tools.access_control)
|
||||
or has_access(user.id, "read", tools.access_control, db=db)
|
||||
):
|
||||
return tools
|
||||
else:
|
||||
@@ -364,8 +367,9 @@ async def update_tools_by_id(
|
||||
id: str,
|
||||
form_data: ToolForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
tools = Tools.get_tool_by_id(id)
|
||||
tools = Tools.get_tool_by_id(id, db=db)
|
||||
if not tools:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -375,7 +379,7 @@ async def update_tools_by_id(
|
||||
# Is the user the original creator, in a group with write access, or an admin
|
||||
if (
|
||||
tools.user_id != user.id
|
||||
and not has_access(user.id, "write", tools.access_control)
|
||||
and not has_access(user.id, "write", tools.access_control, db=db)
|
||||
and user.role != "admin"
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -399,7 +403,7 @@ async def update_tools_by_id(
|
||||
}
|
||||
|
||||
log.debug(updated)
|
||||
tools = Tools.update_tool_by_id(id, updated)
|
||||
tools = Tools.update_tool_by_id(id, updated, db=db)
|
||||
|
||||
if tools:
|
||||
return tools
|
||||
@@ -423,9 +427,9 @@ async def update_tools_by_id(
|
||||
|
||||
@router.delete("/id/{id}/delete", response_model=bool)
|
||||
async def delete_tools_by_id(
|
||||
request: Request, id: str, user=Depends(get_verified_user)
|
||||
request: Request, id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
tools = Tools.get_tool_by_id(id)
|
||||
tools = Tools.get_tool_by_id(id, db=db)
|
||||
if not tools:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -434,7 +438,7 @@ async def delete_tools_by_id(
|
||||
|
||||
if (
|
||||
tools.user_id != user.id
|
||||
and not has_access(user.id, "write", tools.access_control)
|
||||
and not has_access(user.id, "write", tools.access_control, db=db)
|
||||
and user.role != "admin"
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -442,7 +446,7 @@ async def delete_tools_by_id(
|
||||
detail=ERROR_MESSAGES.UNAUTHORIZED,
|
||||
)
|
||||
|
||||
result = Tools.delete_tool_by_id(id)
|
||||
result = Tools.delete_tool_by_id(id, db=db)
|
||||
if result:
|
||||
TOOLS = request.app.state.TOOLS
|
||||
if id in TOOLS:
|
||||
@@ -457,11 +461,11 @@ async def delete_tools_by_id(
|
||||
|
||||
|
||||
@router.get("/id/{id}/valves", response_model=Optional[dict])
|
||||
async def get_tools_valves_by_id(id: str, user=Depends(get_verified_user)):
|
||||
tools = Tools.get_tool_by_id(id)
|
||||
async def get_tools_valves_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
tools = Tools.get_tool_by_id(id, db=db)
|
||||
if tools:
|
||||
try:
|
||||
valves = Tools.get_tool_valves_by_id(id)
|
||||
valves = Tools.get_tool_valves_by_id(id, db=db)
|
||||
return valves
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
@@ -482,9 +486,9 @@ async def get_tools_valves_by_id(id: str, user=Depends(get_verified_user)):
|
||||
|
||||
@router.get("/id/{id}/valves/spec", response_model=Optional[dict])
|
||||
async def get_tools_valves_spec_by_id(
|
||||
request: Request, id: str, user=Depends(get_verified_user)
|
||||
request: Request, id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
tools = Tools.get_tool_by_id(id)
|
||||
tools = Tools.get_tool_by_id(id, db=db)
|
||||
if tools:
|
||||
if id in request.app.state.TOOLS:
|
||||
tools_module = request.app.state.TOOLS[id]
|
||||
@@ -510,9 +514,9 @@ async def get_tools_valves_spec_by_id(
|
||||
|
||||
@router.post("/id/{id}/valves/update", response_model=Optional[dict])
|
||||
async def update_tools_valves_by_id(
|
||||
request: Request, id: str, form_data: dict, user=Depends(get_verified_user)
|
||||
request: Request, id: str, form_data: dict, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
tools = Tools.get_tool_by_id(id)
|
||||
tools = Tools.get_tool_by_id(id, db=db)
|
||||
if not tools:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -521,7 +525,7 @@ async def update_tools_valves_by_id(
|
||||
|
||||
if (
|
||||
tools.user_id != user.id
|
||||
and not has_access(user.id, "write", tools.access_control)
|
||||
and not has_access(user.id, "write", tools.access_control, db=db)
|
||||
and user.role != "admin"
|
||||
):
|
||||
raise HTTPException(
|
||||
@@ -546,7 +550,7 @@ async def update_tools_valves_by_id(
|
||||
form_data = {k: v for k, v in form_data.items() if v is not None}
|
||||
valves = Valves(**form_data)
|
||||
valves_dict = valves.model_dump(exclude_unset=True)
|
||||
Tools.update_tool_valves_by_id(id, valves_dict)
|
||||
Tools.update_tool_valves_by_id(id, valves_dict, db=db)
|
||||
return valves_dict
|
||||
except Exception as e:
|
||||
log.exception(f"Failed to update tool valves by id {id}: {e}")
|
||||
@@ -562,11 +566,11 @@ async def update_tools_valves_by_id(
|
||||
|
||||
|
||||
@router.get("/id/{id}/valves/user", response_model=Optional[dict])
|
||||
async def get_tools_user_valves_by_id(id: str, user=Depends(get_verified_user)):
|
||||
tools = Tools.get_tool_by_id(id)
|
||||
async def get_tools_user_valves_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
tools = Tools.get_tool_by_id(id, db=db)
|
||||
if tools:
|
||||
try:
|
||||
user_valves = Tools.get_user_valves_by_id_and_user_id(id, user.id)
|
||||
user_valves = Tools.get_user_valves_by_id_and_user_id(id, user.id, db=db)
|
||||
return user_valves
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
@@ -582,9 +586,9 @@ async def get_tools_user_valves_by_id(id: str, user=Depends(get_verified_user)):
|
||||
|
||||
@router.get("/id/{id}/valves/user/spec", response_model=Optional[dict])
|
||||
async def get_tools_user_valves_spec_by_id(
|
||||
request: Request, id: str, user=Depends(get_verified_user)
|
||||
request: Request, id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
tools = Tools.get_tool_by_id(id)
|
||||
tools = Tools.get_tool_by_id(id, db=db)
|
||||
if tools:
|
||||
if id in request.app.state.TOOLS:
|
||||
tools_module = request.app.state.TOOLS[id]
|
||||
@@ -605,9 +609,9 @@ async def get_tools_user_valves_spec_by_id(
|
||||
|
||||
@router.post("/id/{id}/valves/user/update", response_model=Optional[dict])
|
||||
async def update_tools_user_valves_by_id(
|
||||
request: Request, id: str, form_data: dict, user=Depends(get_verified_user)
|
||||
request: Request, id: str, form_data: dict, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
tools = Tools.get_tool_by_id(id)
|
||||
tools = Tools.get_tool_by_id(id, db=db)
|
||||
|
||||
if tools:
|
||||
if id in request.app.state.TOOLS:
|
||||
@@ -624,7 +628,7 @@ async def update_tools_user_valves_by_id(
|
||||
user_valves = UserValves(**form_data)
|
||||
user_valves_dict = user_valves.model_dump(exclude_unset=True)
|
||||
Tools.update_user_valves_by_id_and_user_id(
|
||||
id, user.id, user_valves_dict
|
||||
id, user.id, user_valves_dict, db=db
|
||||
)
|
||||
return user_valves_dict
|
||||
except Exception as e:
|
||||
|
||||
Reference in New Issue
Block a user