chore: format

This commit is contained in:
Timothy Jaeryang Baek
2026-04-14 17:27:31 -05:00
parent f1be85d997
commit 5dae600ce7
81 changed files with 247 additions and 140 deletions
+1 -3
View File
@@ -92,9 +92,7 @@ class ERROR_MESSAGES(str, Enum):
INVALID_PASSWORD = lambda err='': err if err else 'The password does not meet the required validation criteria.'
AUTOMATION_LIMIT_EXCEEDED = lambda size='': f'Automation limit reached ({size})'
AUTOMATION_TOO_FREQUENT = (
lambda interval='': f'Schedule too frequent. Minimum interval is {interval} seconds.'
)
AUTOMATION_TOO_FREQUENT = lambda interval='': f'Schedule too frequent. Minimum interval is {interval} seconds.'
AUTOMATION_INVALID_RRULE = lambda err='': f'Invalid RRULE: {err}'
AUTOMATION_NO_FUTURE_RUNS = 'RRULE has no future occurrences'
+20 -12
View File
@@ -1707,7 +1707,9 @@ async def chat_completion(
},
'messages': [
{'role': 'user', 'content': user_message.get('content', '')},
] if user_message_id else [],
]
if user_message_id
else [],
'tags': [],
'timestamp': int(time.time() * 1000),
},
@@ -1734,9 +1736,7 @@ async def chat_completion(
pass
else:
# Existing chat — verify ownership
if (
not await Chats.is_chat_owner(chat_id, user.id) and user.role != 'admin'
):
if not await Chats.is_chat_owner(chat_id, user.id) and user.role != 'admin':
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=ERROR_MESSAGES.DEFAULT(),
@@ -1787,16 +1787,16 @@ async def chat_completion(
# Link user message → all assistant messages (childrenIds)
if user_message_id and all_assistant_ids:
existing_user_message = await Chats.get_message_by_id_and_message_id(
chat_id, user_message_id
)
existing_user_message = await Chats.get_message_by_id_and_message_id(chat_id, user_message_id)
if existing_user_message:
child_ids = existing_user_message.get('childrenIds', [])
for assistant_id in all_assistant_ids:
if assistant_id not in child_ids:
child_ids.append(assistant_id)
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id, user_message_id, {'childrenIds': child_ids},
chat_id,
user_message_id,
{'childrenIds': child_ids},
)
# Save each assistant placeholder
@@ -1956,11 +1956,19 @@ async def chat_completion(
task_id, _ = await create_task(
request.app.state.redis,
process_chat(
request, model_form_data, user, per_model_metadata, resolved_model,
tasks if idx == 0 else {
k: v for k, v in (tasks or {}).items()
request,
model_form_data,
user,
per_model_metadata,
resolved_model,
tasks
if idx == 0
else {
k: v
for k, v in (tasks or {}).items()
if k not in (TASKS.TITLE_GENERATION, TASKS.TAGS_GENERATION)
} or None,
}
or None,
),
id=chat_id,
)
+2 -6
View File
@@ -506,9 +506,7 @@ async def query_collection_with_hybrid_search(
log.exception(f'Failed to fetch collection {name}: {e}')
return name, None
collection_results = dict(
await asyncio.gather(*(_fetch_collection(name) for name in collection_names))
)
collection_results = dict(await asyncio.gather(*(_fetch_collection(name) for name in collection_names)))
log.info(f'Starting hybrid search for {len(queries)} queries in {len(collection_names)} collections...')
@@ -1153,9 +1151,7 @@ async def get_sources_from_items(
if full_context:
# Sync helper makes blocking VECTOR_DB_CLIENT calls;
# offload so the async caller's event loop stays free.
query_result = await asyncio.to_thread(
get_all_items_from_collections, collection_names
)
query_result = await asyncio.to_thread(get_all_items_from_collections, collection_names)
else:
query_result = await query_collection(
request,
@@ -101,9 +101,7 @@ class AsyncVectorDBClient:
filter: Optional[Dict] = None,
limit: int = 10,
) -> Optional[SearchResult]:
return await asyncio.to_thread(
self._sync.search, collection_name, vectors, filter, limit
)
return await asyncio.to_thread(self._sync.search, collection_name, vectors, filter, limit)
async def query(
self,
@@ -111,9 +109,7 @@ class AsyncVectorDBClient:
filter: Dict,
limit: Optional[int] = None,
) -> Optional[GetResult]:
return await asyncio.to_thread(
self._sync.query, collection_name, filter, limit
)
return await asyncio.to_thread(self._sync.query, collection_name, filter, limit)
async def get(self, collection_name: str) -> Optional[GetResult]:
return await asyncio.to_thread(self._sync.get, collection_name)
@@ -124,9 +120,7 @@ class AsyncVectorDBClient:
ids: Optional[List[str]] = None,
filter: Optional[Dict] = None,
) -> None:
return await asyncio.to_thread(
self._sync.delete, collection_name, ids, filter
)
return await asyncio.to_thread(self._sync.delete, collection_name, ids, filter)
async def reset(self) -> None:
return await asyncio.to_thread(self._sync.reset)
@@ -27,4 +27,3 @@ def process_metadata(
else:
result[key] = sanitize_text_for_db(value)
return result
-2
View File
@@ -543,7 +543,6 @@ async def image_generations(
model = get_image_model(request)
try:
if request.app.state.config.IMAGE_GENERATION_ENGINE == 'openai':
headers = {
@@ -856,7 +855,6 @@ async def image_edits(
),
)
try:
if request.app.state.config.IMAGE_EDIT_ENGINE == 'openai':
headers = {
+1 -3
View File
@@ -2532,9 +2532,7 @@ async def delete_entries_from_collection(
if hash is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=ERROR_MESSAGES.DEFAULT(
'File has no hash; cannot delete vector entries by hash.'
),
detail=ERROR_MESSAGES.DEFAULT('File has no hash; cannot delete vector entries by hash.'),
)
# Pre-existing bug: this used `metadata=` which is not a
+15 -49
View File
@@ -205,9 +205,7 @@ def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
elif content_type == 'image':
source = content_block.get('source', {})
if source.get('type') == 'base64':
media_type = source.get(
'media_type', 'image/png'
)
media_type = source.get('media_type', 'image/png')
data = source.get('data', '')
converted_parts.append(
{
@@ -229,68 +227,36 @@ def convert_anthropic_to_openai_payload(anthropic_payload: dict) -> dict:
elif content_type == 'document':
# Documents have no direct OpenAI equivalent;
# convert to a text representation.
document_source = content_block.get(
'source', {}
)
document_title = content_block.get(
'title', 'Document'
)
document_context = content_block.get(
'context', ''
)
document_text = (
f'[Document: {document_title}]'
)
document_source = content_block.get('source', {})
document_title = content_block.get('title', 'Document')
document_context = content_block.get('context', '')
document_text = f'[Document: {document_title}]'
if document_context:
document_text += f'\n{document_context}'
if (
document_source.get('type') == 'text'
and document_source.get('data')
):
document_text += (
f'\n{document_source["data"]}'
)
converted_parts.append(
{'type': 'text', 'text': document_text}
)
if document_source.get('type') == 'text' and document_source.get('data'):
document_text += f'\n{document_source["data"]}'
converted_parts.append({'type': 'text', 'text': document_text})
elif content_type == 'search_result':
# Convert search results to a text
# representation with source attribution.
search_title = content_block.get('title', '')
search_url = content_block.get('source', '')
search_content_blocks = content_block.get(
'content', []
)
search_content_blocks = content_block.get('content', [])
search_texts = []
for search_block in search_content_blocks:
if (
isinstance(search_block, dict)
and search_block.get('type') == 'text'
):
search_texts.append(
search_block.get('text', '')
)
if isinstance(search_block, dict) and search_block.get('type') == 'text':
search_texts.append(search_block.get('text', ''))
search_body = '\n'.join(search_texts)
search_text = (
f'[Search Result: {search_title}]'
)
search_text = f'[Search Result: {search_title}]'
if search_url:
search_text += f'\nSource: {search_url}'
if search_body:
search_text += f'\n{search_body}'
converted_parts.append(
{'type': 'text', 'text': search_text}
)
converted_parts.append({'type': 'text', 'text': search_text})
# Flatten to string when only text parts are present
if all(
part.get('type') == 'text'
for part in converted_parts
):
tool_content = '\n'.join(
part.get('text', '')
for part in converted_parts
)
if all(part.get('type') == 'text' for part in converted_parts):
tool_content = '\n'.join(part.get('text', '') for part in converted_parts)
elif converted_parts:
tool_content = converted_parts
else:
+2 -7
View File
@@ -105,10 +105,7 @@ class CommitSessionMiddleware:
try:
ScopedSession.commit()
except Exception:
log.exception(
'CommitSessionMiddleware: post-request commit failed; '
'response was already sent to client'
)
log.exception('CommitSessionMiddleware: post-request commit failed; response was already sent to client')
try:
ScopedSession.rollback()
except Exception:
@@ -190,9 +187,7 @@ class WebsocketUpgradeGuardMiddleware:
if query_params.get('transport', [''])[0] == 'websocket':
headers = _scope_headers(scope)
upgrade = headers.get('upgrade', '').lower()
connection_tokens = [
token.strip() for token in headers.get('connection', '').lower().split(',')
]
connection_tokens = [token.strip() for token in headers.get('connection', '').lower().split(',')]
if upgrade != 'websocket' or 'upgrade' not in connection_tokens:
response = JSONResponse(
status_code=400,
+8 -6
View File
@@ -3101,8 +3101,8 @@ async def outlet_filter_handler(ctx):
'content': m.get('content', ''),
'info': m.get('info'),
'timestamp': m.get('timestamp'),
**(({'usage': m['usage']} if m.get('usage') else {})),
**(({'sources': m['sources']} if m.get('sources') else {})),
**({'usage': m['usage']} if m.get('usage') else {}),
**({'sources': m['sources']} if m.get('sources') else {}),
}
for m in message_list
],
@@ -3157,10 +3157,12 @@ async def outlet_filter_handler(ctx):
)
if event_emitter:
await event_emitter({
'type': 'chat:outlet',
'data': {'messages': outlet_result['messages']},
})
await event_emitter(
{
'type': 'chat:outlet',
'data': {'messages': outlet_result['messages']},
}
)
except Exception as e:
log.debug(f'Error running outlet filters: {e}')
+11 -4
View File
@@ -169,7 +169,7 @@ def _normalize_token_expiry(token: dict) -> dict:
# Neither field present — conservative fallback
log.warning(
"OAuth token response missing both 'expires_in' and 'expires_at'; "
f"defaulting to {DEFAULT_TOKEN_EXPIRY_SECONDS}s from now"
f'defaulting to {DEFAULT_TOKEN_EXPIRY_SECONDS}s from now'
)
token['expires_at'] = int(datetime.now().timestamp() + DEFAULT_TOKEN_EXPIRY_SECONDS)
return token
@@ -548,7 +548,6 @@ async def get_oauth_client_info_with_static_credentials(
raise e
def resolve_oauth_client_info(connection: dict) -> dict:
"""
Decrypt OAuth client info from a tool server connection config.
@@ -766,7 +765,11 @@ class OAuthClientManager:
log.warning(f'No OAuth session found for user {user_id}, client_id {client_id}')
return None
if force_refresh or session.expires_at is None or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at):
if (
force_refresh
or session.expires_at is None
or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at)
):
log.debug(f'Token refresh needed for user {user_id}, client_id {session.provider}')
refreshed_token = await self._refresh_token(session)
if refreshed_token:
@@ -1017,7 +1020,11 @@ class OAuthManager:
log.warning(f'No OAuth session found for user {user_id}, session {session_id}')
return None
if force_refresh or session.expires_at is None or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at):
if (
force_refresh
or session.expires_at is None
or datetime.now() + timedelta(minutes=5) >= datetime.fromtimestamp(session.expires_at)
):
log.debug(f'Token refresh needed for user {user_id}, provider {session.provider}')
refreshed_token = await self._refresh_token(session)
if refreshed_token: