chore: format
This commit is contained in:
@@ -1700,7 +1700,11 @@ async def chat_completion(
|
||||
)
|
||||
model_info_params = {
|
||||
**default_model_params,
|
||||
**(model_info.params.model_dump() if model_info and model_info.params else {}),
|
||||
**(
|
||||
model_info.params.model_dump()
|
||||
if model_info and model_info.params
|
||||
else {}
|
||||
),
|
||||
}
|
||||
|
||||
# Check base model existence for custom models
|
||||
@@ -1716,7 +1720,10 @@ async def chat_completion(
|
||||
default_models[0].strip() if default_models[0] else None
|
||||
)
|
||||
|
||||
if fallback_model_id and fallback_model_id in request.app.state.MODELS:
|
||||
if (
|
||||
fallback_model_id
|
||||
and fallback_model_id in request.app.state.MODELS
|
||||
):
|
||||
# Update model and form_data so routing uses the fallback model's type
|
||||
model = request.app.state.MODELS[fallback_model_id]
|
||||
form_data["model"] = fallback_model_id
|
||||
|
||||
@@ -631,7 +631,9 @@ class ChatTable:
|
||||
with get_db_context(db) as db:
|
||||
# Use subquery to delete chat_messages for shared chats
|
||||
shared_chat_id_subquery = (
|
||||
db.query(Chat.id).filter_by(user_id=f"shared-{chat_id}").scalar_subquery()
|
||||
db.query(Chat.id)
|
||||
.filter_by(user_id=f"shared-{chat_id}")
|
||||
.scalar_subquery()
|
||||
)
|
||||
db.query(ChatMessage).filter(
|
||||
ChatMessage.chat_id.in_(shared_chat_id_subquery)
|
||||
|
||||
@@ -190,7 +190,11 @@ class ToolsTable:
|
||||
return tools
|
||||
|
||||
def get_tools_by_user_id(
|
||||
self, user_id: str, permission: str = "write", defer_content: bool = False, db: Optional[Session] = None
|
||||
self,
|
||||
user_id: str,
|
||||
permission: str = "write",
|
||||
defer_content: bool = False,
|
||||
db: Optional[Session] = None,
|
||||
) -> list[ToolUserModel]:
|
||||
tools = self.get_tools(defer_content=defer_content, db=db)
|
||||
user_group_ids = {
|
||||
|
||||
@@ -165,9 +165,7 @@ def process_uploaded_file(
|
||||
request.app.state.config, "STT_SUPPORTED_CONTENT_TYPES", []
|
||||
)
|
||||
|
||||
if strict_match_mime_type(
|
||||
stt_supported_content_types, content_type
|
||||
):
|
||||
if strict_match_mime_type(stt_supported_content_types, content_type):
|
||||
file_path_processed = Storage.get_file(file_path)
|
||||
result = transcribe(
|
||||
request, file_path_processed, file_metadata, user
|
||||
|
||||
@@ -65,12 +65,18 @@ async def get_tools(
|
||||
|
||||
# Local Tools
|
||||
for tool in Tools.get_tools(defer_content=True, db=db):
|
||||
tool_module = request.app.state.TOOLS.get(tool.id) if hasattr(request.app.state, 'TOOLS') else None
|
||||
tool_module = (
|
||||
request.app.state.TOOLS.get(tool.id)
|
||||
if hasattr(request.app.state, "TOOLS")
|
||||
else None
|
||||
)
|
||||
tools.append(
|
||||
ToolUserResponse(
|
||||
**{
|
||||
**tool.model_dump(),
|
||||
"has_user_valves": hasattr(tool_module, "UserValves") if tool_module else False,
|
||||
"has_user_valves": (
|
||||
hasattr(tool_module, "UserValves") if tool_module else False
|
||||
),
|
||||
}
|
||||
)
|
||||
)
|
||||
@@ -212,8 +218,13 @@ async def get_tool_list(
|
||||
or any(
|
||||
g.permission == "write"
|
||||
and (
|
||||
(g.principal_type == "user" and (g.principal_id == user.id or g.principal_id == "*"))
|
||||
or (g.principal_type == "group" and g.principal_id in user_group_ids)
|
||||
(
|
||||
g.principal_type == "user"
|
||||
and (g.principal_id == user.id or g.principal_id == "*")
|
||||
)
|
||||
or (
|
||||
g.principal_type == "group" and g.principal_id in user_group_ids
|
||||
)
|
||||
)
|
||||
for g in tool.access_grants
|
||||
)
|
||||
|
||||
@@ -62,9 +62,7 @@ def _json_sink(message: "Message") -> None:
|
||||
log_entry["extra"] = record["extra"]
|
||||
|
||||
if record["exception"] is not None:
|
||||
log_entry["error"] = "".join(
|
||||
record["exception"].format_exception()
|
||||
).rstrip()
|
||||
log_entry["error"] = "".join(record["exception"].format_exception()).rstrip()
|
||||
|
||||
sys.stdout.write(json.dumps(log_entry, ensure_ascii=False, default=str) + "\n")
|
||||
sys.stdout.flush()
|
||||
|
||||
@@ -2330,7 +2330,8 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
||||
) in request.app.state.config.TOOL_SERVER_CONNECTIONS:
|
||||
if (
|
||||
server_connection.get("type", "") == "mcp"
|
||||
and server_connection.get("info", {}).get("id") == server_id
|
||||
and server_connection.get("info", {}).get("id")
|
||||
== server_id
|
||||
):
|
||||
mcp_server_connection = server_connection
|
||||
break
|
||||
@@ -2391,8 +2392,8 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS and user:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
if metadata and metadata.get("chat_id"):
|
||||
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get(
|
||||
"chat_id"
|
||||
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = (
|
||||
metadata.get("chat_id")
|
||||
)
|
||||
if metadata and metadata.get("message_id"):
|
||||
headers[FORWARD_SESSION_INFO_HEADER_MESSAGE_ID] = (
|
||||
@@ -2410,7 +2411,9 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
||||
).get("function_name_filter_list", "")
|
||||
|
||||
if isinstance(function_name_filter_list, str):
|
||||
function_name_filter_list = function_name_filter_list.split(",")
|
||||
function_name_filter_list = function_name_filter_list.split(
|
||||
","
|
||||
)
|
||||
|
||||
tool_specs = await mcp_clients[server_id].list_tool_specs()
|
||||
for tool_spec in tool_specs:
|
||||
|
||||
Reference in New Issue
Block a user