refac: async db

This commit is contained in:
Timothy Jaeryang Baek
2026-04-12 14:22:11 -05:00
parent b618d84065
commit 27169124f2
74 changed files with 4831 additions and 4479 deletions
+31 -31
View File
@@ -101,7 +101,7 @@ log = logging.getLogger(__name__)
# Let no function be called without need, and let what
# it yields justify the cost of running it.
def get_async_tool_function_and_apply_extra_params(function: Callable, extra_params: dict) -> Callable[..., Awaitable]:
async def get_async_tool_function_and_apply_extra_params(function: Callable, extra_params: dict) -> Callable[..., Awaitable]:
sig = inspect.signature(function)
extra_params = {k: v for k, v in extra_params.items() if k in sig.parameters}
partial_func = partial(function, **extra_params)
@@ -138,13 +138,13 @@ def get_async_tool_function_and_apply_extra_params(function: Callable, extra_par
return new_function
def get_updated_tool_function(function: Callable, extra_params: dict):
async def get_updated_tool_function(function: Callable, extra_params: dict):
# Get the original function and merge updated params
__function__ = getattr(function, '__function__', None)
__extra_params__ = getattr(function, '__extra_params__', None)
if __function__ is not None and __extra_params__ is not None:
return get_async_tool_function_and_apply_extra_params(
return await get_async_tool_function_and_apply_extra_params(
__function__,
{**__extra_params__, **extra_params},
)
@@ -160,16 +160,16 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
tools_dict = {}
# Get user's group memberships for access control checks
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)}
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)}
for tool_id in tool_ids:
tool = Tools.get_tool_by_id(tool_id)
tool = await Tools.get_tool_by_id(tool_id)
if tool:
# Check access control for local tools
if (
not (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
and tool.user_id != user.id
and not AccessGrants.has_access(
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='tool',
resource_id=tool.id,
@@ -182,7 +182,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
module = request.app.state.TOOLS.get(tool_id, None)
if module is None:
module, _ = load_tool_module_by_id(tool_id)
module, _ = await load_tool_module_by_id(tool_id)
request.app.state.TOOLS[tool_id] = module
__user__ = {
@@ -191,11 +191,11 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
# Set valves for the tool
if hasattr(module, 'valves') and hasattr(module, 'Valves'):
valves = Tools.get_tool_valves_by_id(tool_id) or {}
valves = await Tools.get_tool_valves_by_id(tool_id) or {}
module.valves = module.Valves(**valves)
if hasattr(module, 'UserValves'):
__user__['valves'] = module.UserValves( # type: ignore
**Tools.get_user_valves_by_id_and_user_id(tool_id, user.id)
**await Tools.get_user_valves_by_id_and_user_id(tool_id, user.id)
)
for spec in tool.specs:
@@ -213,7 +213,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
# convert to function that takes only model params and inserts custom params
function_name = spec['name']
tool_function = getattr(module, function_name)
callable = get_async_tool_function_and_apply_extra_params(
callable = await get_async_tool_function_and_apply_extra_params(
tool_function,
{
**extra_params,
@@ -285,7 +285,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
tool_server_connection = connections[tool_server_idx]
# Check access control for tool server
if not has_connection_access(user, tool_server_connection, user_group_ids):
if not await has_connection_access(user, tool_server_connection, user_group_ids):
log.warning(f'Access denied to tool server {server_id} for user {user.id}')
continue
@@ -339,7 +339,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
if metadata and metadata.get('message_id'):
headers[FORWARD_SESSION_INFO_HEADER_MESSAGE_ID] = metadata.get('message_id')
def make_tool_function(function_name, tool_server_data, headers):
async def make_tool_function(function_name, tool_server_data, headers):
async def tool_function(**kwargs):
return await execute_tool_server(
url=tool_server_data['url'],
@@ -352,9 +352,9 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
return tool_function
tool_function = make_tool_function(function_name, tool_server_data, headers)
tool_function = await make_tool_function(function_name, tool_server_data, headers)
callable = get_async_tool_function_and_apply_extra_params(
callable = await get_async_tool_function_and_apply_extra_params(
tool_function,
{},
)
@@ -381,7 +381,7 @@ async def get_tools(request: Request, tool_ids: list[str], user: UserModel, extr
return tools_dict
def get_builtin_tools(
async def get_builtin_tools(
request: Request, extra_params: dict, features: dict = None, model: dict = None
) -> dict[str, dict]:
"""
@@ -406,10 +406,10 @@ def get_builtin_tools(
# Helper to check user-level feature permission (admins always pass)
user = extra_params.get('__user__', {})
def has_user_permission(feature_key: str) -> bool:
async def has_user_permission(feature_key: str) -> bool:
if user.get('role') == 'admin':
return True
return has_permission(
return await has_permission(
user.get('id', ''),
f'features.{feature_key}',
request.app.state.config.USER_PERMISSIONS,
@@ -461,7 +461,7 @@ def get_builtin_tools(
if (
is_builtin_tool_enabled('memory')
and (features.get('memory') or get_model_capability('memory', False))
and has_user_permission('memories')
and await has_user_permission('memories')
):
builtin_functions.extend(
[
@@ -479,7 +479,7 @@ def get_builtin_tools(
and getattr(request.app.state.config, 'ENABLE_WEB_SEARCH', False)
and get_model_capability('web_search')
and features.get('web_search')
and has_user_permission('web_search')
and await has_user_permission('web_search')
):
builtin_functions.extend([search_web, fetch_url])
@@ -489,7 +489,7 @@ def get_builtin_tools(
and getattr(request.app.state.config, 'ENABLE_IMAGE_GENERATION', False)
and get_model_capability('image_generation')
and features.get('image_generation')
and has_user_permission('image_generation')
and await has_user_permission('image_generation')
):
builtin_functions.append(generate_image)
if (
@@ -497,7 +497,7 @@ def get_builtin_tools(
and getattr(request.app.state.config, 'ENABLE_IMAGE_EDIT', False)
and get_model_capability('image_generation')
and features.get('image_generation')
and has_user_permission('image_generation')
and await has_user_permission('image_generation')
):
builtin_functions.append(edit_image)
@@ -507,7 +507,7 @@ def get_builtin_tools(
and getattr(request.app.state.config, 'ENABLE_CODE_INTERPRETER', True)
and get_model_capability('code_interpreter')
and features.get('code_interpreter')
and has_user_permission('code_interpreter')
and await has_user_permission('code_interpreter')
):
builtin_functions.append(execute_code)
@@ -515,7 +515,7 @@ def get_builtin_tools(
if (
is_builtin_tool_enabled('notes')
and getattr(request.app.state.config, 'ENABLE_NOTES', False)
and has_user_permission('notes')
and await has_user_permission('notes')
):
builtin_functions.extend([search_notes, view_note, write_note, replace_note_content])
@@ -523,7 +523,7 @@ def get_builtin_tools(
if (
is_builtin_tool_enabled('channels')
and getattr(request.app.state.config, 'ENABLE_CHANNELS', False)
and has_user_permission('channels')
and await has_user_permission('channels')
):
builtin_functions.extend(
[
@@ -543,11 +543,11 @@ def get_builtin_tools(
builtin_functions.append(tasks)
# Automation tools - create and manage scheduled automations from chat
if is_builtin_tool_enabled('automations') and has_user_permission('automations'):
if is_builtin_tool_enabled('automations') and await has_user_permission('automations'):
builtin_functions.extend([create_automation, update_automation, list_automations, toggle_automation, delete_automation])
for func in builtin_functions:
callable = get_async_tool_function_and_apply_extra_params(
callable = await get_async_tool_function_and_apply_extra_params(
func,
{
'__request__': request,
@@ -1024,8 +1024,8 @@ async def get_terminal_tools(
log.warning(f'Terminal server not found: {terminal_id}')
return {}
user_group_ids = {group.id for group in Groups.get_groups_by_member_id(user.id)}
if not has_connection_access(user, connection, user_group_ids):
user_group_ids = {group.id for group in await Groups.get_groups_by_member_id(user.id)}
if not await has_connection_access(user, connection, user_group_ids):
log.warning(f'Access denied to terminal {terminal_id} for user {user.id}')
return {}
@@ -1077,7 +1077,7 @@ async def get_terminal_tools(
tool_spec.get('description', '') + f'\n\nThe current working directory is: {terminal_cwd}'
)
def make_tool_function(fn_name, srv_data, hdrs, cks):
async def make_tool_function(fn_name, srv_data, hdrs, cks):
async def tool_function(**kwargs):
return await execute_tool_server(
url=srv_data['url'],
@@ -1090,8 +1090,8 @@ async def get_terminal_tools(
return tool_function
tool_function = make_tool_function(function_name, server_data, headers, cookies)
callable = get_async_tool_function_and_apply_extra_params(tool_function, {})
tool_function = await make_tool_function(function_name, server_data, headers, cookies)
callable = await get_async_tool_function_and_apply_extra_params(tool_function, {})
tools_dict[function_name] = {
'tool_id': f'terminal:{terminal_id}',