refac: async db
This commit is contained in:
@@ -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}',
|
||||
|
||||
Reference in New Issue
Block a user