refac/enh: db session sharing
This commit is contained in:
@@ -24,6 +24,8 @@ from open_webui.constants import ERROR_MESSAGES
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from pydantic import BaseModel, HttpUrl
|
||||
from open_webui.internal.db import get_session
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
@@ -37,13 +39,13 @@ router = APIRouter()
|
||||
|
||||
|
||||
@router.get("/", response_model=list[FunctionResponse])
|
||||
async def get_functions(user=Depends(get_verified_user)):
|
||||
return Functions.get_functions()
|
||||
async def get_functions(user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
return Functions.get_functions(db=db)
|
||||
|
||||
|
||||
@router.get("/list", response_model=list[FunctionUserResponse])
|
||||
async def get_function_list(user=Depends(get_admin_user)):
|
||||
return Functions.get_function_list()
|
||||
async def get_function_list(user=Depends(get_admin_user), db: Session = Depends(get_session)):
|
||||
return Functions.get_function_list(db=db)
|
||||
|
||||
|
||||
############################
|
||||
@@ -52,8 +54,8 @@ async def get_function_list(user=Depends(get_admin_user)):
|
||||
|
||||
|
||||
@router.get("/export", response_model=list[FunctionModel | FunctionWithValvesModel])
|
||||
async def get_functions(include_valves: bool = False, user=Depends(get_admin_user)):
|
||||
return Functions.get_functions(include_valves=include_valves)
|
||||
async def get_functions(include_valves: bool = False, user=Depends(get_admin_user), db: Session = Depends(get_session)):
|
||||
return Functions.get_functions(include_valves=include_valves, db=db)
|
||||
|
||||
|
||||
############################
|
||||
@@ -142,7 +144,7 @@ class SyncFunctionsForm(BaseModel):
|
||||
|
||||
@router.post("/sync", response_model=list[FunctionWithValvesModel])
|
||||
async def sync_functions(
|
||||
request: Request, form_data: SyncFunctionsForm, user=Depends(get_admin_user)
|
||||
request: Request, form_data: SyncFunctionsForm, user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
try:
|
||||
for function in form_data.functions:
|
||||
@@ -164,7 +166,7 @@ async def sync_functions(
|
||||
)
|
||||
raise e
|
||||
|
||||
return Functions.sync_functions(user.id, form_data.functions)
|
||||
return Functions.sync_functions(user.id, form_data.functions, db=db)
|
||||
except Exception as e:
|
||||
log.exception(f"Failed to load a function: {e}")
|
||||
raise HTTPException(
|
||||
@@ -180,7 +182,7 @@ async def sync_functions(
|
||||
|
||||
@router.post("/create", response_model=Optional[FunctionResponse])
|
||||
async def create_new_function(
|
||||
request: Request, form_data: FunctionForm, user=Depends(get_admin_user)
|
||||
request: Request, form_data: FunctionForm, user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
if not form_data.id.isidentifier():
|
||||
raise HTTPException(
|
||||
@@ -190,7 +192,7 @@ async def create_new_function(
|
||||
|
||||
form_data.id = form_data.id.lower()
|
||||
|
||||
function = Functions.get_function_by_id(form_data.id)
|
||||
function = Functions.get_function_by_id(form_data.id, db=db)
|
||||
if function is None:
|
||||
try:
|
||||
form_data.content = replace_imports(form_data.content)
|
||||
@@ -203,13 +205,13 @@ async def create_new_function(
|
||||
FUNCTIONS = request.app.state.FUNCTIONS
|
||||
FUNCTIONS[form_data.id] = function_module
|
||||
|
||||
function = Functions.insert_new_function(user.id, function_type, form_data)
|
||||
function = Functions.insert_new_function(user.id, function_type, form_data, db=db)
|
||||
|
||||
function_cache_dir = CACHE_DIR / "functions" / form_data.id
|
||||
function_cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if function_type == "filter" and getattr(function_module, "toggle", None):
|
||||
Functions.update_function_metadata_by_id(id, {"toggle": True})
|
||||
Functions.update_function_metadata_by_id(form_data.id, {"toggle": True}, db=db)
|
||||
|
||||
if function:
|
||||
return function
|
||||
@@ -237,8 +239,8 @@ async def create_new_function(
|
||||
|
||||
|
||||
@router.get("/id/{id}", response_model=Optional[FunctionModel])
|
||||
async def get_function_by_id(id: str, user=Depends(get_admin_user)):
|
||||
function = Functions.get_function_by_id(id)
|
||||
async def get_function_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
|
||||
function = Functions.get_function_by_id(id, db=db)
|
||||
|
||||
if function:
|
||||
return function
|
||||
@@ -255,11 +257,11 @@ async def get_function_by_id(id: str, user=Depends(get_admin_user)):
|
||||
|
||||
|
||||
@router.post("/id/{id}/toggle", response_model=Optional[FunctionModel])
|
||||
async def toggle_function_by_id(id: str, user=Depends(get_admin_user)):
|
||||
function = Functions.get_function_by_id(id)
|
||||
async def toggle_function_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
|
||||
function = Functions.get_function_by_id(id, db=db)
|
||||
if function:
|
||||
function = Functions.update_function_by_id(
|
||||
id, {"is_active": not function.is_active}
|
||||
id, {"is_active": not function.is_active}, db=db
|
||||
)
|
||||
|
||||
if function:
|
||||
@@ -282,11 +284,11 @@ async def toggle_function_by_id(id: str, user=Depends(get_admin_user)):
|
||||
|
||||
|
||||
@router.post("/id/{id}/toggle/global", response_model=Optional[FunctionModel])
|
||||
async def toggle_global_by_id(id: str, user=Depends(get_admin_user)):
|
||||
function = Functions.get_function_by_id(id)
|
||||
async def toggle_global_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
|
||||
function = Functions.get_function_by_id(id, db=db)
|
||||
if function:
|
||||
function = Functions.update_function_by_id(
|
||||
id, {"is_global": not function.is_global}
|
||||
id, {"is_global": not function.is_global}, db=db
|
||||
)
|
||||
|
||||
if function:
|
||||
@@ -310,7 +312,7 @@ async def toggle_global_by_id(id: str, user=Depends(get_admin_user)):
|
||||
|
||||
@router.post("/id/{id}/update", response_model=Optional[FunctionModel])
|
||||
async def update_function_by_id(
|
||||
request: Request, id: str, form_data: FunctionForm, user=Depends(get_admin_user)
|
||||
request: Request, id: str, form_data: FunctionForm, user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
try:
|
||||
form_data.content = replace_imports(form_data.content)
|
||||
@@ -325,10 +327,10 @@ async def update_function_by_id(
|
||||
updated = {**form_data.model_dump(exclude={"id"}), "type": function_type}
|
||||
log.debug(updated)
|
||||
|
||||
function = Functions.update_function_by_id(id, updated)
|
||||
function = Functions.update_function_by_id(id, updated, db=db)
|
||||
|
||||
if function_type == "filter" and getattr(function_module, "toggle", None):
|
||||
Functions.update_function_metadata_by_id(id, {"toggle": True})
|
||||
Functions.update_function_metadata_by_id(id, {"toggle": True}, db=db)
|
||||
|
||||
if function:
|
||||
return function
|
||||
@@ -352,9 +354,9 @@ async def update_function_by_id(
|
||||
|
||||
@router.delete("/id/{id}/delete", response_model=bool)
|
||||
async def delete_function_by_id(
|
||||
request: Request, id: str, user=Depends(get_admin_user)
|
||||
request: Request, id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
result = Functions.delete_function_by_id(id)
|
||||
result = Functions.delete_function_by_id(id, db=db)
|
||||
|
||||
if result:
|
||||
FUNCTIONS = request.app.state.FUNCTIONS
|
||||
@@ -370,11 +372,11 @@ async def delete_function_by_id(
|
||||
|
||||
|
||||
@router.get("/id/{id}/valves", response_model=Optional[dict])
|
||||
async def get_function_valves_by_id(id: str, user=Depends(get_admin_user)):
|
||||
function = Functions.get_function_by_id(id)
|
||||
async def get_function_valves_by_id(id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)):
|
||||
function = Functions.get_function_by_id(id, db=db)
|
||||
if function:
|
||||
try:
|
||||
valves = Functions.get_function_valves_by_id(id)
|
||||
valves = Functions.get_function_valves_by_id(id, db=db)
|
||||
return valves
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
@@ -395,9 +397,9 @@ async def get_function_valves_by_id(id: str, user=Depends(get_admin_user)):
|
||||
|
||||
@router.get("/id/{id}/valves/spec", response_model=Optional[dict])
|
||||
async def get_function_valves_spec_by_id(
|
||||
request: Request, id: str, user=Depends(get_admin_user)
|
||||
request: Request, id: str, user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
function = Functions.get_function_by_id(id)
|
||||
function = Functions.get_function_by_id(id, db=db)
|
||||
if function:
|
||||
function_module, function_type, frontmatter = get_function_module_from_cache(
|
||||
request, id
|
||||
@@ -421,9 +423,9 @@ async def get_function_valves_spec_by_id(
|
||||
|
||||
@router.post("/id/{id}/valves/update", response_model=Optional[dict])
|
||||
async def update_function_valves_by_id(
|
||||
request: Request, id: str, form_data: dict, user=Depends(get_admin_user)
|
||||
request: Request, id: str, form_data: dict, user=Depends(get_admin_user), db: Session = Depends(get_session)
|
||||
):
|
||||
function = Functions.get_function_by_id(id)
|
||||
function = Functions.get_function_by_id(id, db=db)
|
||||
if function:
|
||||
function_module, function_type, frontmatter = get_function_module_from_cache(
|
||||
request, id
|
||||
@@ -437,7 +439,7 @@ async def update_function_valves_by_id(
|
||||
valves = Valves(**form_data)
|
||||
|
||||
valves_dict = valves.model_dump(exclude_unset=True)
|
||||
Functions.update_function_valves_by_id(id, valves_dict)
|
||||
Functions.update_function_valves_by_id(id, valves_dict, db=db)
|
||||
return valves_dict
|
||||
except Exception as e:
|
||||
log.exception(f"Error updating function values by id {id}: {e}")
|
||||
@@ -464,11 +466,11 @@ async def update_function_valves_by_id(
|
||||
|
||||
|
||||
@router.get("/id/{id}/valves/user", response_model=Optional[dict])
|
||||
async def get_function_user_valves_by_id(id: str, user=Depends(get_verified_user)):
|
||||
function = Functions.get_function_by_id(id)
|
||||
async def get_function_user_valves_by_id(id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
function = Functions.get_function_by_id(id, db=db)
|
||||
if function:
|
||||
try:
|
||||
user_valves = Functions.get_user_valves_by_id_and_user_id(id, user.id)
|
||||
user_valves = Functions.get_user_valves_by_id_and_user_id(id, user.id, db=db)
|
||||
return user_valves
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
@@ -484,9 +486,9 @@ async def get_function_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_function_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)
|
||||
):
|
||||
function = Functions.get_function_by_id(id)
|
||||
function = Functions.get_function_by_id(id, db=db)
|
||||
if function:
|
||||
function_module, function_type, frontmatter = get_function_module_from_cache(
|
||||
request, id
|
||||
@@ -505,9 +507,9 @@ async def get_function_user_valves_spec_by_id(
|
||||
|
||||
@router.post("/id/{id}/valves/user/update", response_model=Optional[dict])
|
||||
async def update_function_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)
|
||||
):
|
||||
function = Functions.get_function_by_id(id)
|
||||
function = Functions.get_function_by_id(id, db=db)
|
||||
|
||||
if function:
|
||||
function_module, function_type, frontmatter = get_function_module_from_cache(
|
||||
@@ -522,7 +524,7 @@ async def update_function_user_valves_by_id(
|
||||
user_valves = UserValves(**form_data)
|
||||
user_valves_dict = user_valves.model_dump(exclude_unset=True)
|
||||
Functions.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