Split Custom OAI prompt post-processing modes

3b4a455ef845ca6682cbd07e6c97be199a8a1399

Cohee <18619528+Cohee1207@users.noreply.github.com>

5 files changed, +101 -9Ignore whitespace
default/config.yaml+2 -0
@@ -121,6 +121,8 @@ extras:
121121 speechToTextModel: Xenova/whisper-small
122122 textToSpeechModel: Xenova/speecht5_tts
123123# -- OPENAI CONFIGURATION --
124+# A placeholder message to use in strict prompt post-processing mode when the prompt doesn't start with a user message
125+promptPlaceholder: "[Start a new chat]"
124126openai:
125127 # Will send a random user ID to OpenAI completion API
126128 randomizeUserId: false
public/index.html+2 -1
@@ -3102,7 +3102,8 @@
31023102 <h4 data-i18n="Prompt Post-Processing">Prompt Post-Processing</h4>
31033103 <select id="custom_prompt_post_processing" class="text_pole" title="Applies additional processing to the prompt before sending it to the API." data-i18n="[title]Applies additional processing to the prompt before sending it to the API.">
31043104 <option data-i18n="prompt_post_processing_none" value="">None</option>
31053105 <option value="claudemerge">ClaudeMerge consecutive roles</option>
3106+ <option value="strict">Strict (user first, alternating roles)</option>
31063107 </select>
31073108 </form>
31083109 <div id="01ai_form" data-source="01ai">
public/scripts/openai.js+7 -0
@@ -199,7 +199,10 @@ const continue_postfix_types = {
199199
200200const custom_prompt_post_processing_types = {
201201 NONE: '',
202+ /** @deprecated Use MERGE instead. */
202203 CLAUDE: 'claude',
204+ MERGE: 'merge',
205+ STRICT: 'strict',
203206};
204207
205208const sensitiveFields = [
@@ -3043,6 +3046,10 @@ function loadOpenAISettings(data, settings) {
30433046 setNamesBehaviorControls();
30443047 setContinuePostfixControls();
30453048
3049+ if (oai_settings.custom_prompt_post_processing === custom_prompt_post_processing_types.CLAUDE) {
3050+ oai_settings.custom_prompt_post_processing = custom_prompt_post_processing_types.MERGE;
3051+ }
3052+
30463053 $('#chat_completion_source').val(oai_settings.chat_completion_source).trigger('change');
30473054 $('#oai_max_context_unlocked').prop('checked', oai_settings.max_context_unlocked);
30483055 $('#custom_prompt_post_processing').val(oai_settings.custom_prompt_post_processing);
src/endpoints/backends/chat-completions.js+6 -3
@@ -4,7 +4,7 @@ const fetch = require('node-fetch').default;
44const { jsonParser } = require('../../express-common');
55const { CHAT_COMPLETION_SOURCES, GEMINI_SAFETY, BISON_SAFETY, OPENROUTER_HEADERS } = require('../../constants');
66const { forwardFetchResponse, getConfigValue, tryParse, uuidv4, mergeObjectWithYaml, excludeKeysByYaml, color } = require('../../util');
77const { convertClaudeMessages, convertGooglePrompt, convertTextCompletionPrompt, convertCohereMessages, convertMistralMessages, convertAI21Messages, mergeMessages } = require('../../prompt-converters');
88const CohereStream = require('../../cohere-stream');
99
1010const { readSecret, SECRET_KEYS } = require('../secrets');
@@ -31,8 +31,11 @@ const API_AI21 = 'https://api.ai21.com/studio/v1';
3131 */
3232function postProcessPrompt(messages, type, charName, userName) {
3333 switch (type) {
34+ case 'merge':
3435 case 'claude':
3536 return convertClaudeMessagesmergeMessages(messages, '', false, '', charName, userName, false).messages;
37+ case 'strict':
38+ return mergeMessages(messages, charName, userName, true);
3639 default:
3740 return messages;
3841 }
@@ -902,7 +905,7 @@ router.post('/generate', jsonParser, function (request, response) {
902905 apiKey = readSecret(request.user.directories, SECRET_KEYS.PERPLEXITY);
903906 headers = {};
904907 bodyParams = {};
905908 request.body.messages = postProcessPrompt(request.body.messages, 'claudestrict', request.body.char_name, request.body.user_name);
906909 } else if (request.body.chat_completion_source === CHAT_COMPLETION_SOURCES.GROQ) {
907910 apiUrl = API_GROQ;
908911 apiKey = readSecret(request.user.directories, SECRET_KEYS.GROQ);
src/prompt-converters.js+84 -5
@@ -1,6 +1,8 @@
11require('./polyfill.js');
22const { getConfigValue } = require('./util.js');
33
4+const PROMPT_PLACEHOLDER = getConfigValue('promptPlaceholder', 'Let\'s get started.');
5+
46/**
57 * Convert a prompt from the ChatML objects to the format used by Claude.
68 * Mainly deprecated. Only used for counting tokens.
@@ -122,7 +124,7 @@ function convertClaudeMessages(messages, prefillString, useSysPrompt, humanMsgFi
122124 if (messages.length === 0 || (messages.length > 0 && messages[0].role !== 'user')) {
123125 messages.unshift({
124126 role: 'user',
125127 content: humanMsgFix || '[Start a new chat]'PROMPT_PLACEHOLDER,
126128 });
127129 }
128130 }
@@ -260,7 +262,6 @@ function convertCohereMessages(messages, charName = '', userName = '') {
260262 'user': 'USER',
261263 'assistant': 'CHATBOT',
262264 };
263- const placeholder = '[Start a new chat]';
264265 let systemPrompt = '';
265266
266267 // Collect all the system messages up until the first instance of a non-system message, and then remove them from the messages array.
@@ -288,12 +289,12 @@ function convertCohereMessages(messages, charName = '', userName = '') {
288289 if (messages.length === 0) {
289290 messages.unshift({
290291 role: 'user',
291292 content: placeholderPROMPT_PLACEHOLDER,
292293 });
293294 }
294295
295296 const lastNonSystemMessageIndex = messages.findLastIndex(msg => msg.role === 'user' || msg.role === 'assistant');
296297 const userPrompt = messages.slice(lastNonSystemMessageIndex).map(msg => msg.content).join('\n\n') || placeholderPROMPT_PLACEHOLDER;
297298
298299 const chatHistory = messages.slice(0, lastNonSystemMessageIndex).map(msg => {
299300 return {
@@ -469,7 +470,7 @@ function convertAI21Messages(messages, charName = '', userName = '') {
469470 if (messages.length === 0) {
470471 messages.unshift({
471472 role: 'user',
472- content: '[Start a new chat]',
473+ content: PROMPT_PLACEHOLDER,
473474 });
474475 }
475476
@@ -554,6 +555,83 @@ function convertMistralMessages(messages, charName = '', userName = '') {
554555}
555556
556557/**
558+ * Merge messages with the same consecutive role, removing names if they exist.
559+ * @param {any[]} messages Messages to merge
560+ * @param {string} charName Character name
561+ * @param {string} userName User name
562+ * @param {boolean} strict Enable strict mode: only allow one system message at the start, force user first message
563+ * @returns {any[]} Merged messages
564+ */
565+function mergeMessages(messages, charName, userName, strict) {
566+ let mergedMessages = [];
567+
568+ // Remove names from the messages
569+ messages.forEach((message) => {
570+ if (!message.content) {
571+ message.content = '';
572+ }
573+ if (message.role === 'system' && message.name === 'example_assistant') {
574+ if (charName && !message.content.startsWith(`${charName}: `)) {
575+ message.content = `${charName}: ${message.content}`;
576+ }
577+ }
578+ if (message.role === 'system' && message.name === 'example_user') {
579+ if (userName && !message.content.startsWith(`${userName}: `)) {
580+ message.content = `${userName}: ${message.content}`;
581+ }
582+ }
583+ if (message.name && message.role !== 'system') {
584+ if (!message.content.startsWith(`${message.name}: `)) {
585+ message.content = `${message.name}: ${message.content}`;
586+ }
587+ }
588+ if (message.role === 'tool') {
589+ message.role = 'user';
590+ }
591+ delete message.name;
592+ delete message.tool_calls;
593+ delete message.tool_call_id;
594+ });
595+
596+ // Squash consecutive messages with the same role
597+ messages.forEach((message) => {
598+ if (mergedMessages.length > 0 && mergedMessages[mergedMessages.length - 1].role === message.role && message.content) {
599+ mergedMessages[mergedMessages.length - 1].content += '\n\n' + message.content;
600+ } else {
601+ mergedMessages.push(message);
602+ }
603+ });
604+
605+ // Prevent erroring out if the messages array is empty.
606+ if (messages.length === 0) {
607+ messages.unshift({
608+ role: 'user',
609+ content: PROMPT_PLACEHOLDER,
610+ });
611+ }
612+
613+ if (strict) {
614+ for (let i = 0; i < mergedMessages.length; i++) {
615+ // Force mid-prompt system messages to be user messages
616+ if (i > 0 && mergedMessages[i].role === 'system') {
617+ mergedMessages[i].role = 'user';
618+ }
619+ }
620+ if (mergedMessages.length) {
621+ if (mergedMessages[0].role === 'system' && (mergedMessages.length === 1 || mergedMessages[1].role !== 'user')) {
622+ mergedMessages.splice(1, 0, { role: 'user', content: PROMPT_PLACEHOLDER });
623+ }
624+ else if (mergedMessages[0].role !== 'system' && mergedMessages[0].role !== 'user') {
625+ mergedMessages.unshift({ role: 'user', content: PROMPT_PLACEHOLDER });
626+ }
627+ }
628+ return mergeMessages(mergedMessages, charName, userName, false);
629+ }
630+
631+ return mergedMessages;
632+}
633+
634+/**
557635 * Convert a prompt from the ChatML objects to the format used by Text Completion API.
558636 * @param {object[]} messages Array of messages
559637 * @returns {string} Prompt for Text Completion API
@@ -586,4 +664,5 @@ module.exports = {
586664 convertCohereMessages,
587665 convertMistralMessages,
588666 convertAI21Messages,
667+ mergeMessages,
589668};