refac
This commit is contained in:
@@ -720,8 +720,8 @@ class UsersTable:
|
||||
|
||||
async def get_valid_user_ids(self, user_ids: list[str], db: AsyncSession | None = None) -> list[str]:
|
||||
async with get_async_db_context(db) as session:
|
||||
result = await session.execute(select(User).where(User.id.in_(user_ids)))
|
||||
return [u.id for u in result.scalars().all()]
|
||||
result = await session.execute(select(User.id).where(User.id.in_(user_ids)))
|
||||
return list(result.scalars().all())
|
||||
|
||||
async def get_super_admin_user(self, db: AsyncSession | None = None) -> UserModel | None:
|
||||
async with get_async_db_context(db) as session:
|
||||
|
||||
@@ -76,11 +76,12 @@ async def get_folders(
|
||||
await check_folders_permission(request, user, db=db)
|
||||
|
||||
folders = await Folders.get_folders_by_user_id(user.id, db=db)
|
||||
folder_ids = {folder.id for folder in folders}
|
||||
|
||||
# Verify folder data integrity
|
||||
folder_list = []
|
||||
for folder in folders:
|
||||
if folder.parent_id and not await Folders.get_folder_by_id_and_user_id(folder.parent_id, user.id, db=db):
|
||||
if folder.parent_id and folder.parent_id not in folder_ids:
|
||||
folder = await Folders.update_folder_parent_id_by_id_and_user_id(folder.id, user.id, None, db=db)
|
||||
|
||||
if folder.data and 'files' in folder.data:
|
||||
|
||||
@@ -2157,23 +2157,25 @@ def extract_skill_ids_from_messages(messages: list[dict]) -> set[str]:
|
||||
return ids
|
||||
|
||||
|
||||
SKILL_MENTION_STRIP_RE = re.compile(r'<(?:\$[^|>]+(?:\|([^>]*))?|/[^|>]+\|([^>]*))>')
|
||||
|
||||
|
||||
def strip_skill_mentions(messages: list[dict]) -> None:
|
||||
"""Replace <$skillId|label> and </skillId|label> mention tags with the label in-place."""
|
||||
strip_re = re.compile(r'<(?:\$[^|>]+(?:\|([^>]*))?|/[^|>]+\|([^>]*))>')
|
||||
|
||||
def label(match):
|
||||
return match.group(1) or match.group(2) or ''
|
||||
|
||||
for message in messages:
|
||||
content = message.get('content')
|
||||
if isinstance(content, str) and strip_re.search(content):
|
||||
message['content'] = strip_re.sub(label, content).strip()
|
||||
if isinstance(content, str) and SKILL_MENTION_STRIP_RE.search(content):
|
||||
message['content'] = SKILL_MENTION_STRIP_RE.sub(label, content).strip()
|
||||
elif isinstance(content, list):
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get('type') == 'text':
|
||||
text = part.get('text', '')
|
||||
if strip_re.search(text):
|
||||
part['text'] = strip_re.sub(label, text).strip()
|
||||
if SKILL_MENTION_STRIP_RE.search(text):
|
||||
part['text'] = SKILL_MENTION_STRIP_RE.sub(label, text).strip()
|
||||
|
||||
|
||||
async def connect_mcp_server(
|
||||
@@ -4789,12 +4791,12 @@ async def streaming_chat_response_handler(response, ctx):
|
||||
params = {}
|
||||
if tool_args and tool_args.strip():
|
||||
try:
|
||||
params = ast.literal_eval(tool_args)
|
||||
except Exception as e:
|
||||
log.debug(e)
|
||||
params = json.loads(tool_args)
|
||||
except Exception:
|
||||
try:
|
||||
params = json.loads(tool_args)
|
||||
except Exception:
|
||||
params = ast.literal_eval(tool_args)
|
||||
except Exception as e:
|
||||
log.debug(e)
|
||||
return None
|
||||
tool_call.setdefault('function', {})['arguments'] = json.dumps(params)
|
||||
return params
|
||||
|
||||
@@ -207,6 +207,7 @@ async def load_tool_module_by_id(tool_id, content=None):
|
||||
if not ENABLE_PLUGINS:
|
||||
raise RuntimeError('Plugins are disabled by ENABLE_PLUGINS=false')
|
||||
|
||||
frontmatter = None
|
||||
if content is None:
|
||||
tool = await Tools.get_tool_by_id(tool_id)
|
||||
if not tool:
|
||||
@@ -238,7 +239,8 @@ async def load_tool_module_by_id(tool_id, content=None):
|
||||
|
||||
# Executing the modified content in the created module's namespace
|
||||
exec(content, module.__dict__)
|
||||
frontmatter = extract_frontmatter(content)
|
||||
if frontmatter is None:
|
||||
frontmatter = extract_frontmatter(content)
|
||||
log.info(f'Loaded module: {module.__name__}')
|
||||
|
||||
# Create and return the object if the class 'Tools' is found in the module
|
||||
@@ -258,6 +260,7 @@ async def load_function_module_by_id(function_id: str, content: str | None = Non
|
||||
if not ENABLE_PLUGINS:
|
||||
raise RuntimeError('Plugins are disabled by ENABLE_PLUGINS=false')
|
||||
|
||||
frontmatter = None
|
||||
if content is None:
|
||||
function = await Functions.get_function_by_id(function_id)
|
||||
if not function:
|
||||
@@ -286,7 +289,8 @@ async def load_function_module_by_id(function_id: str, content: str | None = Non
|
||||
|
||||
# Execute the modified content in the created module's namespace
|
||||
exec(content, module.__dict__)
|
||||
frontmatter = extract_frontmatter(content)
|
||||
if frontmatter is None:
|
||||
frontmatter = extract_frontmatter(content)
|
||||
log.info(f'Loaded module: {module.__name__}')
|
||||
|
||||
# Create appropriate object based on available class type in the module
|
||||
|
||||
Reference in New Issue
Block a user