chore: format

This commit is contained in:
Timothy Jaeryang Baek
2026-02-23 01:40:53 -06:00
parent 424dba443c
commit 9044abf3bb
105 changed files with 3884 additions and 661 deletions
+9 -2
View File
@@ -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
+3 -1
View File
@@ -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)
+5 -1
View File
@@ -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 = {
+1 -3
View File
@@ -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
+15 -4
View File
@@ -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
)
+1 -3
View File
@@ -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()
+7 -4
View File
@@ -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: