refac
This commit is contained in:
+26
-14
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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 } : {}),
|
||||
|
||||
Reference in New Issue
Block a user