This commit is contained in:
Timothy Jaeryang Baek
2026-06-16 23:02:20 +02:00
parent 1eecbc1ac2
commit 56ae99e96a
2 changed files with 35 additions and 22 deletions
+26 -14
View File
@@ -1746,13 +1746,19 @@ async def chat_completion(
parent_id = form_data.pop('parent_id', None)
form_data.pop('new_chat', None) # Legacy field
# Multi-model: {model_id: assistant_message_id}
# Single-model fallback: built from 'model' + 'id'
# Multi-model message_ids: list of {model_id, message_id} entries.
# Supports both the new array format and legacy dict format for backward compat.
message_ids = form_data.pop('message_ids', None)
if not message_ids:
message_ids = {model_id: form_data.pop('id', None)}
else:
if isinstance(message_ids, list):
# New format: [{"model_id": ..., "message_id": ...}, ...]
form_data.pop('id', None)
elif isinstance(message_ids, dict):
# Legacy dict format: {model_id: message_id} — convert to list
message_ids = [{'model_id': k, 'message_id': v} for k, v in message_ids.items()]
form_data.pop('id', None)
else:
# Single-model fallback
message_ids = [{'model_id': model_id, 'message_id': form_data.pop('id', None)}]
user_message = form_data.pop('user_message', None) or form_data.pop('parent_message', None)
@@ -1831,7 +1837,7 @@ async def chat_completion(
status_code=status.HTTP_403_FORBIDDEN,
detail=ERROR_MESSAGES.DEFAULT(),
)
target_message_id = list(message_ids.values())[0] if message_ids else None
target_message_id = message_ids[0]['message_id'] if message_ids else None
if target_message_id:
target_message = await Messages.get_message_by_id(target_message_id)
if target_message and target_message.channel_id != channel.id:
@@ -1849,13 +1855,15 @@ async def chat_completion(
user_message_id = user_message.get('id') if user_message else None
history_messages = {}
all_assistant_ids = [assistant_id for assistant_id in message_ids.values() if assistant_id]
all_assistant_ids = [entry['message_id'] for entry in message_ids if entry.get('message_id')]
if user_message_id and user_message:
user_message['childrenIds'] = all_assistant_ids
history_messages[user_message_id] = user_message
for target_model_id, assistant_message_id in message_ids.items():
for entry in message_ids:
target_model_id = entry['model_id']
assistant_message_id = entry['message_id']
if assistant_message_id:
history_messages[assistant_message_id] = {
'id': assistant_message_id,
@@ -1875,7 +1883,7 @@ async def chat_completion(
chat={
'id': chat_id,
'title': 'New Chat',
'models': list(message_ids.keys()),
'models': [entry['model_id'] for entry in message_ids],
'history': {
'currentId': all_assistant_ids[0] if all_assistant_ids else user_message_id,
'messages': history_messages,
@@ -1969,7 +1977,7 @@ async def chat_completion(
# Save ALL assistant placeholders
user_message_id = metadata.get('user_message_id')
all_assistant_ids = [assistant_id for assistant_id in message_ids.values() if assistant_id]
all_assistant_ids = [entry['message_id'] for entry in message_ids if entry.get('message_id')]
# Link user message → all assistant messages (childrenIds)
if user_message_id and all_assistant_ids:
@@ -1986,7 +1994,9 @@ async def chat_completion(
)
# Save each assistant placeholder
for target_model_id, assistant_message_id in message_ids.items():
for entry in message_ids:
target_model_id = entry['model_id']
assistant_message_id = entry['message_id']
if assistant_message_id:
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id,
@@ -2138,7 +2148,9 @@ async def chat_completion(
task_ids = []
chat_id = metadata['chat_id']
for idx, (target_model_id, assistant_message_id) in enumerate(message_ids.items()):
for idx, entry in enumerate(message_ids):
target_model_id = entry['model_id']
assistant_message_id = entry['message_id']
if not assistant_message_id:
continue
@@ -2185,7 +2197,7 @@ async def chat_completion(
# Emit chat:active=true
if task_ids:
event_emitter = await get_event_emitter(
{**metadata, 'message_id': list(message_ids.values())[0]},
{**metadata, 'message_id': message_ids[0]['message_id']},
update_db=False,
)
if event_emitter:
@@ -2198,7 +2210,7 @@ async def chat_completion(
}
else:
# Legacy/direct: single model, synchronous
metadata['message_id'] = list(message_ids.values())[0]
metadata['message_id'] = message_ids[0]['message_id']
return await process_chat(request, form_data, user, metadata, model, tasks)
+9 -8
View File
@@ -2102,8 +2102,9 @@
: selectedModels;
// Create response messages for each selected model
// Build message_ids map: {model_id: assistant_message_id}
const messageIdsMap: Record<string, string> = {};
// Build message_ids list: [{model_id, message_id}, ...]
// Uses an array instead of a dict to support duplicate model IDs in side-by-side chat.
const messageIdsList: Array<{ model_id: string; message_id: string }> = [];
for (const [_modelIdx, modelId] of selectedModelIds.entries()) {
const model = $models.filter((m) => m.id === modelId).at(0);
@@ -2135,7 +2136,7 @@
}
responseMessageIds[`${modelId}-${modelIdx ? modelIdx : _modelIdx}`] = responseMessageId;
messageIdsMap[modelId] = responseMessageId;
messageIdsList.push({ model_id: modelId, message_id: responseMessageId });
}
}
history = history;
@@ -2181,7 +2182,7 @@
// Single request — backend fans out to all models
const primaryModelId = selectedModelIds[0];
const primaryModel = $models.filter((m) => m.id === primaryModelId).at(0);
const primaryResponseMessageId = messageIdsMap[primaryModelId];
const primaryResponseMessageId = messageIdsList[0]?.message_id;
if (primaryModel && primaryResponseMessageId) {
const chatEventEmitter = await getChatEventEmitter(primaryModel.id, _chatId);
@@ -2197,7 +2198,7 @@
primaryResponseMessageId,
_chatId,
{
messageIdsMap: selectedModelIds.length > 1 ? messageIdsMap : undefined,
messageIdsList: selectedModelIds.length > 1 ? messageIdsList : undefined,
regenerationPrompt
}
);
@@ -2266,11 +2267,11 @@
responseMessageId,
_chatId,
{
messageIdsMap,
messageIdsList,
regenerationPrompt,
continueResponse = false
}: {
messageIdsMap?: Record<string, string>;
messageIdsList?: Array<{ model_id: string; message_id: string }>;
regenerationPrompt?: string | null;
continueResponse?: boolean;
} = {}
@@ -2474,7 +2475,7 @@
folder_id: $selectedFolder?.id ?? undefined,
id: responseMessageId,
...(messageIdsMap ? { message_ids: messageIdsMap } : {}),
...(messageIdsList ? { message_ids: messageIdsList } : {}),
parent_id: userMessage?.parentId ?? null,
user_message: userMessage,
...(regenerationPrompt ? { regeneration_prompt: regenerationPrompt } : {}),