refac/enh: db session sharing

This commit is contained in:
Timothy Jaeryang Baek
2025-12-29 00:21:18 +04:00
parent 6dd0f99b90
commit b1d0f00d8c
23 changed files with 1173 additions and 663 deletions
+43 -41
View File
@@ -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: