Refactor prompt converters with group names awareness
| @@ -276,6 +276,17 @@ export function getGroupMembers(groupId = selected_group) { | |||
| 276 | } | 276 | } |
| 277 | 277 | ||
| 278 | /** | 278 | /** |
| 279 | * Retrieves the member names of a group. If the group is not selected, an empty array is returned. | ||
| 280 | * @returns {string[]} An array of character names representing the members of the group. | ||
| 281 | */ | ||
| 282 | export function getGroupNames() { | ||
| 283 | const groupMembers = selected_group ? groups.find(x => x.id == selected_group)?.members : null; | ||
| 284 | return Array.isArray(groupMembers) | ||
| 285 | ? groupMembers.map(x => characters.find(y => y.avatar === x)?.name).filter(x => x) | ||
| 286 | : []; | ||
| 287 | } | ||
| 288 | |||
| 289 | /** | ||
| 279 | * Finds the character ID for a group member. | 290 | * Finds the character ID for a group member. |
| 280 | * @param {string} arg 0-based member index or character name | 291 | * @param {string} arg 0-based member index or character name |
| 281 | * @returns {number} 0-based character ID | 292 | * @returns {number} 0-based character ID |
| @@ -33,7 +33,7 @@ import { | |||
| 33 | system_message_types, | 33 | system_message_types, |
| 34 | this_chid, | 34 | this_chid, |
| 35 | } from '../script.js'; | 35 | } from '../script.js'; |
| 36 | import { groups, selected_group } from './group-chats.js'; | 36 | import { getGroupNames, selected_group } from './group-chats.js'; |
| 37 | 37 | ||
| 38 | import { | 38 | import { |
| 39 | chatCompletionDefaultPrompts, | 39 | chatCompletionDefaultPrompts, |
| @@ -543,10 +543,7 @@ function setupChatCompletionPromptManager(openAiSettings) { | |||
| 543 | * @returns {Message[]} Array of message objects | 543 | * @returns {Message[]} Array of message objects |
| 544 | */ | 544 | */ |
| 545 | export function parseExampleIntoIndividual(messageExampleString, appendNamesForGroup = true) { | 545 | export function parseExampleIntoIndividual(messageExampleString, appendNamesForGroup = true) { |
| 546 | const groupMembers = selected_group ? groups.find(x => x.id == selected_group)?.members : null; | 546 | const groupBotNames = getGroupNames().map(name => `${name}:`); |
| 547 | const groupBotNames = Array.isArray(groupMembers) | ||
| 548 | ? groupMembers.map(x => characters.find(y => y.avatar === x)?.name).filter(x => x).map(x => `${x}:`) | ||
| 549 | : []; | ||
| 550 | 547 | ||
| 551 | let result = []; // array of msgs | 548 | let result = []; // array of msgs |
| 552 | let tmp = messageExampleString.split('\n'); | 549 | let tmp = messageExampleString.split('\n'); |
| @@ -1877,6 +1874,7 @@ async function sendOpenAIRequest(type, messages, signal) { | |||
| 1877 | 'n': canMultiSwipe ? oai_settings.n : undefined, | 1874 | 'n': canMultiSwipe ? oai_settings.n : undefined, |
| 1878 | 'user_name': name1, | 1875 | 'user_name': name1, |
| 1879 | 'char_name': name2, | 1876 | 'char_name': name2, |
| 1877 | 'group_names': getGroupNames(), | ||
| 1880 | }; | 1878 | }; |
| 1881 | 1879 | ||
| 1882 | // Empty array will produce a validation error | 1880 | // Empty array will produce a validation error |
| @@ -27,6 +27,7 @@ import { | |||
| 27 | mergeMessages, | 27 | mergeMessages, |
| 28 | cachingAtDepthForOpenRouterClaude, | 28 | cachingAtDepthForOpenRouterClaude, |
| 29 | cachingAtDepthForClaude, | 29 | cachingAtDepthForClaude, |
| 30 | getPromptNames, | ||
| 30 | } from '../../prompt-converters.js'; | 31 | } from '../../prompt-converters.js'; |
| 31 | 32 | ||
| 32 | import { readSecret, SECRET_KEYS } from '../secrets.js'; | 33 | import { readSecret, SECRET_KEYS } from '../secrets.js'; |
| @@ -55,17 +56,16 @@ const API_NANOGPT = 'https://nano-gpt.com/api/v1'; | |||
| 55 | * Applies a post-processing step to the generated messages. | 56 | * Applies a post-processing step to the generated messages. |
| 56 | * @param {object[]} messages Messages to post-process | 57 | * @param {object[]} messages Messages to post-process |
| 57 | * @param {string} type Prompt conversion type | 58 | * @param {string} type Prompt conversion type |
| 58 | * @param {string} charName Character name | 59 | * @param {import('../../prompt-converters.js').PromptNames} names Prompt names |
| 59 | * @param {string} userName User name | ||
| 60 | * @returns | 60 | * @returns |
| 61 | */ | 61 | */ |
| 62 | function postProcessPrompt(messages, type, charName, userName) { | 62 | function postProcessPrompt(messages, type, names) { |
| 63 | switch (type) { | 63 | switch (type) { |
| 64 | case 'merge': | 64 | case 'merge': |
| 65 | case 'claude': | 65 | case 'claude': |
| 66 | return mergeMessages(messages, charName, userName, false); | 66 | return mergeMessages(messages, names, false); |
| 67 | case 'strict': | 67 | case 'strict': |
| 68 | return mergeMessages(messages, charName, userName, true); | 68 | return mergeMessages(messages, names, true); |
| 69 | default: | 69 | default: |
| 70 | return messages; | 70 | return messages; |
| 71 | } | 71 | } |
| @@ -101,7 +101,7 @@ async function sendClaudeRequest(request, response) { | |||
| 101 | const additionalHeaders = {}; | 101 | const additionalHeaders = {}; |
| 102 | const useTools = request.body.model.startsWith('claude-3') && Array.isArray(request.body.tools) && request.body.tools.length > 0; | 102 | const useTools = request.body.model.startsWith('claude-3') && Array.isArray(request.body.tools) && request.body.tools.length > 0; |
| 103 | const useSystemPrompt = (request.body.model.startsWith('claude-2') || request.body.model.startsWith('claude-3')) && request.body.claude_use_sysprompt; | 103 | const useSystemPrompt = (request.body.model.startsWith('claude-2') || request.body.model.startsWith('claude-3')) && request.body.claude_use_sysprompt; |
| 104 | const convertedPrompt = convertClaudeMessages(request.body.messages, request.body.assistant_prefill, useSystemPrompt, useTools, request.body.char_name, request.body.user_name); | 104 | const convertedPrompt = convertClaudeMessages(request.body.messages, request.body.assistant_prefill, useSystemPrompt, useTools, getPromptNames(request)); |
| 105 | // Add custom stop sequences | 105 | // Add custom stop sequences |
| 106 | const stopSequences = []; | 106 | const stopSequences = []; |
| 107 | if (Array.isArray(request.body.stop)) { | 107 | if (Array.isArray(request.body.stop)) { |
| @@ -282,9 +282,9 @@ async function sendMakerSuiteRequest(request, response) { | |||
| 282 | model.includes('gemini-1.5-flash') || | 282 | model.includes('gemini-1.5-flash') || |
| 283 | model.includes('gemini-1.5-pro') || | 283 | model.includes('gemini-1.5-pro') || |
| 284 | model.startsWith('gemini-exp') | 284 | model.startsWith('gemini-exp') |
| 285 | ) && request.body.use_makersuite_sysprompt; | 285 | ) && request.body.use_makersuite_sysprompt; |
| 286 | 286 | ||
| 287 | const prompt = convertGooglePrompt(request.body.messages, model, should_use_system_prompt, request.body.char_name, request.body.user_name); | 287 | const prompt = convertGooglePrompt(request.body.messages, model, should_use_system_prompt, getPromptNames(request)); |
| 288 | let body = { | 288 | let body = { |
| 289 | contents: prompt.contents, | 289 | contents: prompt.contents, |
| 290 | safetySettings: GEMINI_SAFETY, | 290 | safetySettings: GEMINI_SAFETY, |
| @@ -384,7 +384,7 @@ async function sendAI21Request(request, response) { | |||
| 384 | request.socket.on('close', function () { | 384 | request.socket.on('close', function () { |
| 385 | controller.abort(); | 385 | controller.abort(); |
| 386 | }); | 386 | }); |
| 387 | const convertedPrompt = convertAI21Messages(request.body.messages, request.body.char_name, request.body.user_name); | 387 | const convertedPrompt = convertAI21Messages(request.body.messages, getPromptNames(request)); |
| 388 | const body = { | 388 | const body = { |
| 389 | messages: convertedPrompt, | 389 | messages: convertedPrompt, |
| 390 | model: request.body.model, | 390 | model: request.body.model, |
| @@ -447,7 +447,7 @@ async function sendMistralAIRequest(request, response) { | |||
| 447 | } | 447 | } |
| 448 | 448 | ||
| 449 | try { | 449 | try { |
| 450 | const messages = convertMistralMessages(request.body.messages, request.body.char_name, request.body.user_name); | 450 | const messages = convertMistralMessages(request.body.messages, getPromptNames(request)); |
| 451 | const controller = new AbortController(); | 451 | const controller = new AbortController(); |
| 452 | request.socket.removeAllListeners('close'); | 452 | request.socket.removeAllListeners('close'); |
| 453 | request.socket.on('close', function () { | 453 | request.socket.on('close', function () { |
| @@ -528,7 +528,7 @@ async function sendCohereRequest(request, response) { | |||
| 528 | } | 528 | } |
| 529 | 529 | ||
| 530 | try { | 530 | try { |
| 531 | const convertedHistory = convertCohereMessages(request.body.messages, request.body.char_name, request.body.user_name); | 531 | const convertedHistory = convertCohereMessages(request.body.messages, getPromptNames(request)); |
| 532 | const tools = []; | 532 | const tools = []; |
| 533 | 533 | ||
| 534 | if (Array.isArray(request.body.tools) && request.body.tools.length > 0) { | 534 | if (Array.isArray(request.body.tools) && request.body.tools.length > 0) { |
| @@ -886,15 +886,14 @@ router.post('/generate', jsonParser, function (request, response) { | |||
| 886 | request.body.messages = postProcessPrompt( | 886 | request.body.messages = postProcessPrompt( |
| 887 | request.body.messages, | 887 | request.body.messages, |
| 888 | request.body.custom_prompt_post_processing, | 888 | request.body.custom_prompt_post_processing, |
| 889 | request.body.char_name, | 889 | getPromptNames(request)); |
| 890 | request.body.user_name); | ||
| 891 | } | 890 | } |
| 892 | } else if (request.body.chat_completion_source === CHAT_COMPLETION_SOURCES.PERPLEXITY) { | 891 | } else if (request.body.chat_completion_source === CHAT_COMPLETION_SOURCES.PERPLEXITY) { |
| 893 | apiUrl = API_PERPLEXITY; | 892 | apiUrl = API_PERPLEXITY; |
| 894 | apiKey = readSecret(request.user.directories, SECRET_KEYS.PERPLEXITY); | 893 | apiKey = readSecret(request.user.directories, SECRET_KEYS.PERPLEXITY); |
| 895 | headers = {}; | 894 | headers = {}; |
| 896 | bodyParams = {}; | 895 | bodyParams = {}; |
| 897 | request.body.messages = postProcessPrompt(request.body.messages, 'strict', request.body.char_name, request.body.user_name); | 896 | request.body.messages = postProcessPrompt(request.body.messages, 'strict', getPromptNames(request)); |
| 898 | } else if (request.body.chat_completion_source === CHAT_COMPLETION_SOURCES.GROQ) { | 897 | } else if (request.body.chat_completion_source === CHAT_COMPLETION_SOURCES.GROQ) { |
| 899 | apiUrl = API_GROQ; | 898 | apiUrl = API_GROQ; |
| 900 | apiKey = readSecret(request.user.directories, SECRET_KEYS.GROQ); | 899 | apiKey = readSecret(request.user.directories, SECRET_KEYS.GROQ); |
| @@ -4,6 +4,30 @@ import { getConfigValue } from './util.js'; | |||
| 4 | const PROMPT_PLACEHOLDER = getConfigValue('promptPlaceholder', 'Let\'s get started.'); | 4 | const PROMPT_PLACEHOLDER = getConfigValue('promptPlaceholder', 'Let\'s get started.'); |
| 5 | 5 | ||
| 6 | /** | 6 | /** |
| 7 | * @typedef {object} PromptNames | ||
| 8 | * @property {string} charName Character name | ||
| 9 | * @property {string} userName User name | ||
| 10 | * @property {string[]} groupNames Group member names | ||
| 11 | * @property {function(string): boolean} startsFromGroupName Check if a message starts with a group name | ||
| 12 | */ | ||
| 13 | |||
| 14 | /** | ||
| 15 | * Extracts the character name, user name, and group member names from the request. | ||
| 16 | * @param {import('express').Request} request Express request object | ||
| 17 | * @returns {PromptNames} Prompt names | ||
| 18 | */ | ||
| 19 | export function getPromptNames(request) { | ||
| 20 | return { | ||
| 21 | charName: String(request.body.char_name || ''), | ||
| 22 | userName: String(request.body.user_name || ''), | ||
| 23 | groupNames: Array.isArray(request.body.group_names) ? request.body.group_names.map(String) : [], | ||
| 24 | startsFromGroupName: function (message) { | ||
| 25 | return this.groupNames.some(name => message.startsWith(`${name}: `)); | ||
| 26 | }, | ||
| 27 | }; | ||
| 28 | } | ||
| 29 | |||
| 30 | /** | ||
| 7 | * Convert a prompt from the ChatML objects to the format used by Claude. | 31 | * Convert a prompt from the ChatML objects to the format used by Claude. |
| 8 | * Mainly deprecated. Only used for counting tokens. | 32 | * Mainly deprecated. Only used for counting tokens. |
| 9 | * @param {object[]} messages Array of messages | 33 | * @param {object[]} messages Array of messages |
| @@ -91,10 +115,10 @@ export function convertClaudePrompt(messages, addAssistantPostfix, addAssistantP | |||
| 91 | * @param {string} prefillString User determined prefill string | 115 | * @param {string} prefillString User determined prefill string |
| 92 | * @param {boolean} useSysPrompt See if we want to use a system prompt | 116 | * @param {boolean} useSysPrompt See if we want to use a system prompt |
| 93 | * @param {boolean} useTools See if we want to use tools | 117 | * @param {boolean} useTools See if we want to use tools |
| 94 | * @param {string} charName Character name | 118 | * @param {PromptNames} names Prompt names |
| 95 | * @param {string} userName User name | 119 | * @returns {{messages: object[], systemPrompt: object[]}} Prompt for Anthropic |
| 96 | */ | 120 | */ |
| 97 | export function convertClaudeMessages(messages, prefillString, useSysPrompt, useTools, charName, userName) { | 121 | export function convertClaudeMessages(messages, prefillString, useSysPrompt, useTools, names) { |
| 98 | let systemPrompt = []; | 122 | let systemPrompt = []; |
| 99 | if (useSysPrompt) { | 123 | if (useSysPrompt) { |
| 100 | // Collect all the system messages up until the first instance of a non-system message, and then remove them from the messages array. | 124 | // Collect all the system messages up until the first instance of a non-system message, and then remove them from the messages array. |
| @@ -104,14 +128,14 @@ export function convertClaudeMessages(messages, prefillString, useSysPrompt, use | |||
| 104 | break; | 128 | break; |
| 105 | } | 129 | } |
| 106 | // Append example names if not already done by the frontend (e.g. for group chats). | 130 | // Append example names if not already done by the frontend (e.g. for group chats). |
| 107 | if (userName && messages[i].name === 'example_user') { | 131 | if (names.userName && messages[i].name === 'example_user') { |
| 108 | if (!messages[i].content.startsWith(`${userName}: `)) { | 132 | if (!messages[i].content.startsWith(`${names.userName}: `)) { |
| 109 | messages[i].content = `${userName}: ${messages[i].content}`; | 133 | messages[i].content = `${names.userName}: ${messages[i].content}`; |
| 110 | } | 134 | } |
| 111 | } | 135 | } |
| 112 | if (charName && messages[i].name === 'example_assistant') { | 136 | if (names.charName && messages[i].name === 'example_assistant') { |
| 113 | if (!messages[i].content.startsWith(`${charName}: `)) { | 137 | if (!messages[i].content.startsWith(`${names.charName}: `) && !names.startsFromGroupName(messages[i].content)) { |
| 114 | messages[i].content = `${charName}: ${messages[i].content}`; | 138 | messages[i].content = `${names.charName}: ${messages[i].content}`; |
| 115 | } | 139 | } |
| 116 | } | 140 | } |
| 117 | systemPrompt.push({ type: 'text', text: messages[i].content }); | 141 | systemPrompt.push({ type: 'text', text: messages[i].content }); |
| @@ -151,11 +175,15 @@ export function convertClaudeMessages(messages, prefillString, useSysPrompt, use | |||
| 151 | } | 175 | } |
| 152 | 176 | ||
| 153 | if (message.role === 'system') { | 177 | if (message.role === 'system') { |
| 154 | if (userName && message.name === 'example_user') { | 178 | if (names.userName && message.name === 'example_user') { |
| 155 | message.content = `${userName}: ${message.content}`; | 179 | if (!message.content.startsWith(`${names.userName}: `)) { |
| 180 | message.content = `${names.userName}: ${message.content}`; | ||
| 181 | } | ||
| 156 | } | 182 | } |
| 157 | if (charName && message.name === 'example_assistant') { | 183 | if (names.charName && message.name === 'example_assistant') { |
| 158 | message.content = `${charName}: ${message.content}`; | 184 | if (!message.content.startsWith(`${names.charName}: `) && !names.startsFromGroupName(message.content)) { |
| 185 | message.content = `${names.charName}: ${message.content}`; | ||
| 186 | } | ||
| 159 | } | 187 | } |
| 160 | message.role = 'user'; | 188 | message.role = 'user'; |
| 161 | 189 | ||
| @@ -274,11 +302,10 @@ export function convertClaudeMessages(messages, prefillString, useSysPrompt, use | |||
| 274 | /** | 302 | /** |
| 275 | * Convert a prompt from the ChatML objects to the format used by Cohere. | 303 | * Convert a prompt from the ChatML objects to the format used by Cohere. |
| 276 | * @param {object[]} messages Array of messages | 304 | * @param {object[]} messages Array of messages |
| 277 | * @param {string} charName Character name | 305 | * @param {PromptNames} names Prompt names |
| 278 | * @param {string} userName User name | ||
| 279 | * @returns {{chatHistory: object[]}} Prompt for Cohere | 306 | * @returns {{chatHistory: object[]}} Prompt for Cohere |
| 280 | */ | 307 | */ |
| 281 | export function convertCohereMessages(messages, charName = '', userName = '') { | 308 | export function convertCohereMessages(messages, names) { |
| 282 | if (messages.length === 0) { | 309 | if (messages.length === 0) { |
| 283 | messages.unshift({ | 310 | messages.unshift({ |
| 284 | role: 'user', | 311 | role: 'user', |
| @@ -299,13 +326,13 @@ export function convertCohereMessages(messages, charName = '', userName = '') { | |||
| 299 | // No names support (who would've thought) | 326 | // No names support (who would've thought) |
| 300 | if (msg.name) { | 327 | if (msg.name) { |
| 301 | if (msg.role == 'system' && msg.name == 'example_assistant') { | 328 | if (msg.role == 'system' && msg.name == 'example_assistant') { |
| 302 | if (charName && !msg.content.startsWith(`${charName}: `)) { | 329 | if (names.charName && !msg.content.startsWith(`${names.charName}: `) && !names.startsFromGroupName(msg.content)) { |
| 303 | msg.content = `${charName}: ${msg.content}`; | 330 | msg.content = `${names.charName}: ${msg.content}`; |
| 304 | } | 331 | } |
| 305 | } | 332 | } |
| 306 | if (msg.role == 'system' && msg.name == 'example_user') { | 333 | if (msg.role == 'system' && msg.name == 'example_user') { |
| 307 | if (userName && !msg.content.startsWith(`${userName}: `)) { | 334 | if (names.userName && !msg.content.startsWith(`${names.userName}: `)) { |
| 308 | msg.content = `${userName}: ${msg.content}`; | 335 | msg.content = `${names.userName}: ${msg.content}`; |
| 309 | } | 336 | } |
| 310 | } | 337 | } |
| 311 | if (msg.role !== 'system' && !msg.content.startsWith(`${msg.name}: `)) { | 338 | if (msg.role !== 'system' && !msg.content.startsWith(`${msg.name}: `)) { |
| @@ -328,12 +355,10 @@ export function convertCohereMessages(messages, charName = '', userName = '') { | |||
| 328 | * @param {object[]} messages Array of messages | 355 | * @param {object[]} messages Array of messages |
| 329 | * @param {string} model Model name | 356 | * @param {string} model Model name |
| 330 | * @param {boolean} useSysPrompt Use system prompt | 357 | * @param {boolean} useSysPrompt Use system prompt |
| 331 | * @param {string} charName Character name | 358 | * @param {PromptNames} names Prompt names |
| 332 | * @param {string} userName User name | ||
| 333 | * @returns {{contents: *[], system_instruction: {parts: {text: string}}}} Prompt for Google MakerSuite models | 359 | * @returns {{contents: *[], system_instruction: {parts: {text: string}}}} Prompt for Google MakerSuite models |
| 334 | */ | 360 | */ |
| 335 | export function convertGooglePrompt(messages, model, useSysPrompt = false, charName = '', userName = '') { | 361 | export function convertGooglePrompt(messages, model, useSysPrompt, names) { |
| 336 | |||
| 337 | const visionSupportedModels = [ | 362 | const visionSupportedModels = [ |
| 338 | 'gemini-2.0-flash-exp', | 363 | 'gemini-2.0-flash-exp', |
| 339 | 'gemini-1.5-flash', | 364 | 'gemini-1.5-flash', |
| @@ -356,20 +381,19 @@ export function convertGooglePrompt(messages, model, useSysPrompt = false, charN | |||
| 356 | ]; | 381 | ]; |
| 357 | 382 | ||
| 358 | const isMultimodal = visionSupportedModels.includes(model); | 383 | const isMultimodal = visionSupportedModels.includes(model); |
| 359 | let hasImage = false; | ||
| 360 | 384 | ||
| 361 | let sys_prompt = ''; | 385 | let sys_prompt = ''; |
| 362 | if (useSysPrompt) { | 386 | if (useSysPrompt) { |
| 363 | while (messages.length > 1 && messages[0].role === 'system') { | 387 | while (messages.length > 1 && messages[0].role === 'system') { |
| 364 | // Append example names if not already done by the frontend (e.g. for group chats). | 388 | // Append example names if not already done by the frontend (e.g. for group chats). |
| 365 | if (userName && messages[0].name === 'example_user') { | 389 | if (names.userName && messages[0].name === 'example_user') { |
| 366 | if (!messages[0].content.startsWith(`${userName}: `)) { | 390 | if (!messages[0].content.startsWith(`${names.userName}: `)) { |
| 367 | messages[0].content = `${userName}: ${messages[0].content}`; | 391 | messages[0].content = `${names.userName}: ${messages[0].content}`; |
| 368 | } | 392 | } |
| 369 | } | 393 | } |
| 370 | if (charName && messages[0].name === 'example_assistant') { | 394 | if (names.charName && messages[0].name === 'example_assistant') { |
| 371 | if (!messages[0].content.startsWith(`${charName}: `)) { | 395 | if (!messages[0].content.startsWith(`${names.charName}: `) && !names.startsFromGroupName(messages[0].content)) { |
| 372 | messages[0].content = `${charName}: ${messages[0].content}`; | 396 | messages[0].content = `${names.charName}: ${messages[0].content}`; |
| 373 | } | 397 | } |
| 374 | } | 398 | } |
| 375 | sys_prompt += `${messages[0].content}\n\n`; | 399 | sys_prompt += `${messages[0].content}\n\n`; |
| @@ -388,53 +412,62 @@ export function convertGooglePrompt(messages, model, useSysPrompt = false, charN | |||
| 388 | message.role = 'model'; | 412 | message.role = 'model'; |
| 389 | } | 413 | } |
| 390 | 414 | ||
| 415 | // Convert the content to an array of parts | ||
| 416 | if (!Array.isArray(message.content)) { | ||
| 417 | message.content = [{ type: 'text', text: String(message.content ?? '') }]; | ||
| 418 | } | ||
| 419 | |||
| 391 | // similar story as claude | 420 | // similar story as claude |
| 392 | if (message.name) { | 421 | if (message.name) { |
| 393 | if (userName && message.name === 'example_user') { | 422 | message.content.forEach((part) => { |
| 394 | message.name = userName; | 423 | if (part.type !== 'text') { |
| 395 | } | 424 | return; |
| 396 | if (charName && message.name === 'example_assistant') { | ||
| 397 | message.name = charName; | ||
| 398 | } | ||
| 399 | |||
| 400 | if (Array.isArray(message.content)) { | ||
| 401 | if (!message.content[0].text.startsWith(`${message.name}: `)) { | ||
| 402 | message.content[0].text = `${message.name}: ${message.content[0].text}`; | ||
| 403 | } | 425 | } |
| 404 | } else { | 426 | if (message.name === 'example_user') { |
| 405 | if (!message.content.startsWith(`${message.name}: `)) { | 427 | if (!part.text.startsWith(`${names.userName}: `)) { |
| 406 | message.content = `${message.name}: ${message.content}`; | 428 | part.text = `${names.userName}: ${part.text}`; |
| 429 | } | ||
| 430 | } else if (message.name === 'example_assistant') { | ||
| 431 | if (!part.text.startsWith(`${names.charName}: `) && !names.startsFromGroupName(part.text)) { | ||
| 432 | part.text = `${names.charName}: ${part.text}`; | ||
| 433 | } | ||
| 434 | } else { | ||
| 435 | if (!part.text.startsWith(`${message.name}: `)) { | ||
| 436 | part.text = `${message.name}: ${part.text}`; | ||
| 437 | } | ||
| 407 | } | 438 | } |
| 408 | } | 439 | }); |
| 409 | 440 | ||
| 410 | delete message.name; | 441 | delete message.name; |
| 411 | } | 442 | } |
| 412 | 443 | ||
| 413 | //create the prompt parts | 444 | //create the prompt parts |
| 414 | const parts = []; | 445 | const parts = []; |
| 415 | if (typeof message.content === 'string') { | 446 | message.content.forEach((part) => { |
| 416 | parts.push({ text: message.content }); | 447 | if (part.type === 'text') { |
| 417 | } else if (Array.isArray(message.content)) { | 448 | parts.push({ text: part.text }); |
| 418 | message.content.forEach((part) => { | 449 | } else if (part.type === 'image_url' && isMultimodal) { |
| 419 | if (part.type === 'text') { | 450 | const mimeType = part.image_url.url.split(';')[0].split(':')[1]; |
| 420 | parts.push({ text: part.text }); | 451 | const base64Data = part.image_url.url.split(',')[1]; |
| 421 | } else if (part.type === 'image_url' && isMultimodal) { | 452 | parts.push({ |
| 422 | const mimeType = part.image_url.url.split(';')[0].split(':')[1]; | 453 | inlineData: { |
| 423 | const base64Data = part.image_url.url.split(',')[1]; | 454 | mimeType: mimeType, |
| 424 | parts.push({ | 455 | data: base64Data, |
| 425 | inlineData: { | 456 | }, |
| 426 | mimeType: mimeType, | 457 | }); |
| 427 | data: base64Data, | 458 | } |
| 428 | }, | 459 | }); |
| 429 | }); | ||
| 430 | hasImage = true; | ||
| 431 | } | ||
| 432 | }); | ||
| 433 | } | ||
| 434 | 460 | ||
| 435 | // merge consecutive messages with the same role | 461 | // merge consecutive messages with the same role |
| 436 | if (index > 0 && message.role === contents[contents.length - 1].role) { | 462 | if (index > 0 && message.role === contents[contents.length - 1].role) { |
| 437 | contents[contents.length - 1].parts[0].text += '\n\n' + parts[0].text; | 463 | parts.forEach((part) => { |
| 464 | if (part.text) { | ||
| 465 | contents[contents.length - 1].parts[0].text += '\n\n' + part.text; | ||
| 466 | } | ||
| 467 | if (part.inlineData) { | ||
| 468 | contents[contents.length - 1].parts.push(part); | ||
| 469 | } | ||
| 470 | }); | ||
| 438 | } else { | 471 | } else { |
| 439 | contents.push({ | 472 | contents.push({ |
| 440 | role: message.role, | 473 | role: message.role, |
| @@ -449,10 +482,10 @@ export function convertGooglePrompt(messages, model, useSysPrompt = false, charN | |||
| 449 | /** | 482 | /** |
| 450 | * Convert AI21 prompt. Classic: system message squash, user/assistant message merge. | 483 | * Convert AI21 prompt. Classic: system message squash, user/assistant message merge. |
| 451 | * @param {object[]} messages Array of messages | 484 | * @param {object[]} messages Array of messages |
| 452 | * @param {string} charName Character name | 485 | * @param {PromptNames} names Prompt names |
| 453 | * @param {string} userName User name | 486 | * @returns {object[]} Prompt for AI21 |
| 454 | */ | 487 | */ |
| 455 | export function convertAI21Messages(messages, charName = '', userName = '') { | 488 | export function convertAI21Messages(messages, names) { |
| 456 | if (!Array.isArray(messages)) { | 489 | if (!Array.isArray(messages)) { |
| 457 | return []; | 490 | return []; |
| 458 | } | 491 | } |
| @@ -465,14 +498,14 @@ export function convertAI21Messages(messages, charName = '', userName = '') { | |||
| 465 | break; | 498 | break; |
| 466 | } | 499 | } |
| 467 | // Append example names if not already done by the frontend (e.g. for group chats). | 500 | // Append example names if not already done by the frontend (e.g. for group chats). |
| 468 | if (userName && messages[i].name === 'example_user') { | 501 | if (names.userName && messages[i].name === 'example_user') { |
| 469 | if (!messages[i].content.startsWith(`${userName}: `)) { | 502 | if (!messages[i].content.startsWith(`${names.userName}: `)) { |
| 470 | messages[i].content = `${userName}: ${messages[i].content}`; | 503 | messages[i].content = `${names.userName}: ${messages[i].content}`; |
| 471 | } | 504 | } |
| 472 | } | 505 | } |
| 473 | if (charName && messages[i].name === 'example_assistant') { | 506 | if (names.charName && messages[i].name === 'example_assistant') { |
| 474 | if (!messages[i].content.startsWith(`${charName}: `)) { | 507 | if (!messages[i].content.startsWith(`${names.charName}: `) && !names.startsFromGroupName(messages[i].content)) { |
| 475 | messages[i].content = `${charName}: ${messages[i].content}`; | 508 | messages[i].content = `${names.charName}: ${messages[i].content}`; |
| 476 | } | 509 | } |
| 477 | } | 510 | } |
| 478 | systemPrompt += `${messages[i].content}\n\n`; | 511 | systemPrompt += `${messages[i].content}\n\n`; |
| @@ -521,10 +554,10 @@ export function convertAI21Messages(messages, charName = '', userName = '') { | |||
| 521 | /** | 554 | /** |
| 522 | * Convert a prompt from the ChatML objects to the format used by MistralAI. | 555 | * Convert a prompt from the ChatML objects to the format used by MistralAI. |
| 523 | * @param {object[]} messages Array of messages | 556 | * @param {object[]} messages Array of messages |
| 524 | * @param {string} charName Character name | 557 | * @param {PromptNames} names Prompt names |
| 525 | * @param {string} userName User name | 558 | * @returns {object[]} Prompt for MistralAI |
| 526 | */ | 559 | */ |
| 527 | export function convertMistralMessages(messages, charName = '', userName = '') { | 560 | export function convertMistralMessages(messages, names) { |
| 528 | if (!Array.isArray(messages)) { | 561 | if (!Array.isArray(messages)) { |
| 529 | return []; | 562 | return []; |
| 530 | } | 563 | } |
| @@ -549,15 +582,15 @@ export function convertMistralMessages(messages, charName = '', userName = '') { | |||
| 549 | msg.tool_call_id = sanitizeToolId(msg.tool_call_id); | 582 | msg.tool_call_id = sanitizeToolId(msg.tool_call_id); |
| 550 | } | 583 | } |
| 551 | if (msg.role === 'system' && msg.name === 'example_assistant') { | 584 | if (msg.role === 'system' && msg.name === 'example_assistant') { |
| 552 | if (charName && !msg.content.startsWith(`${charName}: `)) { | 585 | if (names.charName && !msg.content.startsWith(`${names.charName}: `) && !names.startsFromGroupName(msg.content)) { |
| 553 | msg.content = `${charName}: ${msg.content}`; | 586 | msg.content = `${names.charName}: ${msg.content}`; |
| 554 | } | 587 | } |
| 555 | delete msg.name; | 588 | delete msg.name; |
| 556 | } | 589 | } |
| 557 | 590 | ||
| 558 | if (msg.role === 'system' && msg.name === 'example_user') { | 591 | if (msg.role === 'system' && msg.name === 'example_user') { |
| 559 | if (userName && !msg.content.startsWith(`${userName}: `)) { | 592 | if (names.userName && !msg.content.startsWith(`${names.userName}: `)) { |
| 560 | msg.content = `${userName}: ${msg.content}`; | 593 | msg.content = `${names.userName}: ${msg.content}`; |
| 561 | } | 594 | } |
| 562 | delete msg.name; | 595 | delete msg.name; |
| 563 | } | 596 | } |
| @@ -603,12 +636,11 @@ export function convertMistralMessages(messages, charName = '', userName = '') { | |||
| 603 | /** | 636 | /** |
| 604 | * Merge messages with the same consecutive role, removing names if they exist. | 637 | * Merge messages with the same consecutive role, removing names if they exist. |
| 605 | * @param {any[]} messages Messages to merge | 638 | * @param {any[]} messages Messages to merge |
| 606 | * @param {string} charName Character name | 639 | * @param {PromptNames} names Prompt names |
| 607 | * @param {string} userName User name | ||
| 608 | * @param {boolean} strict Enable strict mode: only allow one system message at the start, force user first message | 640 | * @param {boolean} strict Enable strict mode: only allow one system message at the start, force user first message |
| 609 | * @returns {any[]} Merged messages | 641 | * @returns {any[]} Merged messages |
| 610 | */ | 642 | */ |
| 611 | export function mergeMessages(messages, charName, userName, strict) { | 643 | export function mergeMessages(messages, names, strict) { |
| 612 | let mergedMessages = []; | 644 | let mergedMessages = []; |
| 613 | 645 | ||
| 614 | /** @type {Map<string,object>} */ | 646 | /** @type {Map<string,object>} */ |
| @@ -636,13 +668,13 @@ export function mergeMessages(messages, charName, userName, strict) { | |||
| 636 | message.content = text; | 668 | message.content = text; |
| 637 | } | 669 | } |
| 638 | if (message.role === 'system' && message.name === 'example_assistant') { | 670 | if (message.role === 'system' && message.name === 'example_assistant') { |
| 639 | if (charName && !message.content.startsWith(`${charName}: `)) { | 671 | if (names.charName && !message.content.startsWith(`${names.charName}: `) && !names.startsFromGroupName(message.content)) { |
| 640 | message.content = `${charName}: ${message.content}`; | 672 | message.content = `${names.charName}: ${message.content}`; |
| 641 | } | 673 | } |
| 642 | } | 674 | } |
| 643 | if (message.role === 'system' && message.name === 'example_user') { | 675 | if (message.role === 'system' && message.name === 'example_user') { |
| 644 | if (userName && !message.content.startsWith(`${userName}: `)) { | 676 | if (names.userName && !message.content.startsWith(`${names.userName}: `)) { |
| 645 | message.content = `${userName}: ${message.content}`; | 677 | message.content = `${names.userName}: ${message.content}`; |
| 646 | } | 678 | } |
| 647 | } | 679 | } |
| 648 | if (message.name && message.role !== 'system') { | 680 | if (message.name && message.role !== 'system') { |
| @@ -716,7 +748,7 @@ export function mergeMessages(messages, charName, userName, strict) { | |||
| 716 | mergedMessages.unshift({ role: 'user', content: PROMPT_PLACEHOLDER }); | 748 | mergedMessages.unshift({ role: 'user', content: PROMPT_PLACEHOLDER }); |
| 717 | } | 749 | } |
| 718 | } | 750 | } |
| 719 | return mergeMessages(mergedMessages, charName, userName, false); | 751 | return mergeMessages(mergedMessages, names, false); |
| 720 | } | 752 | } |
| 721 | 753 | ||
| 722 | return mergedMessages; | 754 | return mergedMessages; |