Add prompt injection filters

d6f34f7b2c812462030be6af198b4b982808d895

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

3 files changed, +86 -40Ignore whitespace
public/script.js+75 -33
@@ -2875,23 +2875,54 @@ function addPersonaDescriptionExtensionPrompt() {
28752875 }
28762876}
28772877
2878-function getAllExtensionPrompts() {
2878+/**
2879- const value = Object
2879+ * Returns all extension prompts combined.
2880- .values(extension_prompts)
2880+ * @returns {Promise<string>} Combined extension prompts
2881- .filter(x => x.value)
2881+ */
2882- .map(x => x.value.trim())
2882+async function getAllExtensionPrompts() {
2883- .join('\n');
2883+ const values = [];
2884+
2885+ for (const prompt of Object.values(extension_prompts)) {
2886+ const value = prompt?.value?.trim();
2887+
2888+ if (!value) {
2889+ continue;
2890+ }
28842891
2885- return value.length ? substituteParams(value) : '';
2892+ const hasFilter = typeof prompt.filter === 'function';
2893+ if (hasFilter && !await prompt.filter()) {
2894+ continue;
2895+ }
2896+
2897+ values.push(value);
2898+ }
2899+
2900+ return substituteParams(values.join('\n'));
28862901}
28872902
2888-// Wrapper to fetch extension prompts by module name
2903+/**
2889-export function getExtensionPromptByName(moduleName) {
2904+ * Wrapper to fetch extension prompts by module name
2890- if (moduleName) {
2905+ * @param {string} moduleName Module name
2891- return substituteParams(extension_prompts[moduleName]?.value);
2906+ * @returns {Promise<string>} Extension prompt
2892- } else {
2907+ */
2893- return;
2908+export async function getExtensionPromptByName(moduleName) {
2909+ if (!moduleName) {
2910+ return '';
2911+ }
2912+
2913+ const prompt = extension_prompts[moduleName];
2914+
2915+ if (!prompt) {
2916+ return '';
28942917 }
2918+
2919+ const hasFilter = typeof prompt.filter === 'function';
2920+
2921+ if (hasFilter && !await prompt.filter()) {
2922+ return '';
2923+ }
2924+
2925+ return substituteParams(prompt.value);
28952926}
28962927
28972928/**
@@ -2902,27 +2933,36 @@ export function getExtensionPromptByName(moduleName) {
29022933 * @param {string} [separator] Separator for joining multiple prompts
29032934 * @param {number} [role] Role of the prompt
29042935 * @param {boolean} [wrap] Wrap start and end with a separator
29052936 * @returns {Promise<string>} Extension prompt
29062937 */
29072938export async function getExtensionPrompt(position = extension_prompt_types.IN_PROMPT, depth = undefined, separator = '\n', role = undefined, wrap = true) {
2908- let extension_prompt = Object.keys(extension_prompts)
2939+ const filterByFunction = async (prompt) => {
2940+ const hasFilter = typeof prompt.filter === 'function';
2941+ if (hasFilter && !await prompt.filter()) {
2942+ return false;
2943+ }
2944+ return true;
2945+ };
2946+ const promptPromises = Object.keys(extension_prompts)
29092947 .sort()
29102948 .map((x) => extension_prompts[x])
29112949 .filter(x => x.position == position && x.value)
29122950 .filter(x => depth === undefined || x.depth === undefined || x.depth === depth)
29132951 .filter(x => role === undefined || x.role === undefined || x.role === role)
2914- .map(x => x.value.trim())
2952+ .filter(filterByFunction);
2915- .join(separator);
2953+ const prompts = await Promise.all(promptPromises);
2916- if (wrap && extension_prompt.length && !extension_prompt.startsWith(separator)) {
2954+
2917- extension_prompt = separator + extension_prompt;
2955+ let values = prompts.map(x => x.value.trim()).join(separator);
2956+ if (wrap && values.length && !values.startsWith(separator)) {
2957+ values = separator + values;
29182958 }
29192959 if (wrap && extension_promptvalues.length && !extension_promptvalues.endsWith(separator)) {
29202960 extension_promptvalues = extension_promptvalues + separator;
29212961 }
29222962 if (extension_promptvalues.length) {
29232963 extension_promptvalues = substituteParams(extension_promptvalues);
29242964 }
29252965 return extension_promptvalues;
29262966}
29272967
29282968export function baseChatReplace(value, name1, name2) {
@@ -3836,7 +3876,7 @@ export async function Generate(type, { automatic_trigger, force_name2, quiet_pro
38363876 // Inject all Depth prompts. Chat Completion does it separately
38373877 let injectedIndices = [];
38383878 if (main_api !== 'openai') {
38393879 injectedIndices = await doChatInject(coreChat, isContinue);
38403880 }
38413881
38423882 // Insert character jailbreak as the last user message (if exists, allowed, preferred, and not using Chat Completion)
@@ -3909,8 +3949,8 @@ export async function Generate(type, { automatic_trigger, force_name2, quiet_pro
39093949 }
39103950
39113951 // Call combined AN into Generate
39123952 const beforeScenarioAnchor = (await getExtensionPrompt(extension_prompt_types.BEFORE_PROMPT)).trimStart();
39133953 const afterScenarioAnchor = await getExtensionPrompt(extension_prompt_types.IN_PROMPT);
39143954
39153955 const storyStringParams = {
39163956 description: description,
@@ -4473,7 +4513,7 @@ export async function Generate(type, { automatic_trigger, force_name2, quiet_pro
44734513 ...thisPromptBits[currentArrayEntry],
44744514 rawPrompt: generate_data.prompt || generate_data.input,
44754515 mesId: getNextMessageId(type),
44764516 allAnchors: await getAllExtensionPrompts(),
44774517 chatInjects: injectedIndices?.map(index => arrMes[arrMes.length - index - 1])?.join('') || '',
44784518 summarizeString: (extension_prompts['1_memory']?.value || ''),
44794519 authorsNoteString: (extension_prompts['2_floating_prompt']?.value || ''),
@@ -4742,9 +4782,9 @@ export function stopGeneration() {
47424782 * Injects extension prompts into chat messages.
47434783 * @param {object[]} messages Array of chat messages
47444784 * @param {boolean} isContinue Whether the generation is a continuation. If true, the extension prompts of depth 0 are injected at position 1.
47454785 * @returns {Promise<number[]>} Array of indices where the extension prompts were injected
47464786 */
47474787async function doChatInject(messages, isContinue) {
47484788 const injectedIndices = [];
47494789 let totalInsertedMessages = 0;
47504790 messages.reverse();
@@ -4762,7 +4802,7 @@ function doChatInject(messages, isContinue) {
47624802 const wrap = false;
47634803
47644804 for (const role of roles) {
47654805 const extensionPrompt = String(await getExtensionPrompt(extension_prompt_types.IN_CHAT, i, separator, role, wrap)).trimStart();
47664806 const isNarrator = role === extension_prompt_roles.SYSTEM;
47674807 const isUser = role === extension_prompt_roles.USER;
47684808 const name = names[role];
@@ -7455,14 +7495,16 @@ function select_rm_characters() {
74557495 * @param {number} depth Insertion depth. 0 represets the last message in context. Expected values up to MAX_INJECTION_DEPTH.
74567496 * @param {number} role Extension prompt role. Defaults to SYSTEM.
74577497 * @param {boolean} scan Should the prompt be included in the world info scan.
7498+ * @param {(function(): Promise<boolean>|boolean)} filter Filter function to determine if the prompt should be injected.
74587499 */
74597500export function setExtensionPrompt(key, value, position, depth, scan = false, role = extension_prompt_roles.SYSTEM, filter = null) {
74607501 extension_prompts[key] = {
74617502 value: String(value),
74627503 position: Number(position),
74637504 depth: Number(depth),
74647505 scan: !!scan,
74657506 role: Number(role ?? extension_prompt_roles.SYSTEM),
7507+ filter: filter,
74667508 };
74677509}
74687510
public/scripts/openai.js+10 -6
@@ -611,8 +611,9 @@ function formatWorldInfo(value) {
611611 *
612612 * @param {Prompt[]} prompts - Array containing injection prompts.
613613 * @param {Object[]} messages - Array containing all messages.
614+ * @returns {Promise<Object[]>} - Array containing all messages with injections.
614615 */
615616async function populationInjectionPrompts(prompts, messages) {
616617 let totalInsertedMessages = 0;
617618
618619 const roleTypes = {
@@ -635,7 +636,7 @@ function populationInjectionPrompts(prompts, messages) {
635636 // Get prompts for current role
636637 const rolePrompts = depthPrompts.filter(prompt => prompt.role === role).map(x => x.content).join(separator);
637638 // Get extension prompt
638639 const extensionPrompt = await getExtensionPrompt(extension_prompt_types.IN_CHAT, i, separator, roleTypes[role], wrap);
639640
640641 const jointPrompt = [rolePrompts, extensionPrompt].filter(x => x).map(x => x.trim()).join(separator);
641642
@@ -1020,7 +1021,7 @@ async function populateChatCompletion(prompts, chatCompletion, { bias, quietProm
10201021 }
10211022
10221023 // Add in-chat injections
10231024 messages = await populationInjectionPrompts(absolutePrompts, messages);
10241025
10251026 // Decide whether dialogue examples should always be added
10261027 if (power_user.pin_examples) {
@@ -1051,9 +1052,9 @@ async function populateChatCompletion(prompts, chatCompletion, { bias, quietProm
10511052 * @param {string} options.systemPromptOverride
10521053 * @param {string} options.jailbreakPromptOverride
10531054 * @param {string} options.personaDescription
10541055 * @returns {Promise<Object>} prompts - The prepared and merged system and user-defined prompts.
10551056 */
10561057async function preparePromptsForChatCompletion({ Scenario, charPersonality, name2, worldInfoBefore, worldInfoAfter, charDescription, quietPrompt, bias, extensionPrompts, systemPromptOverride, jailbreakPromptOverride, personaDescription }) {
10571058 const scenarioText = Scenario && oai_settings.scenario_format ? substituteParams(oai_settings.scenario_format) : '';
10581059 const charPersonalityText = charPersonality && oai_settings.personality_format ? substituteParams(oai_settings.personality_format) : '';
10591060 const groupNudge = substituteParams(oai_settings.group_nudge_prompt);
@@ -1142,6 +1143,9 @@ function preparePromptsForChatCompletion({ Scenario, charPersonality, name2, wor
11421143 if (!extensionPrompts[key].value) continue;
11431144 if (![extension_prompt_types.BEFORE_PROMPT, extension_prompt_types.IN_PROMPT].includes(prompt.position)) continue;
11441145
1146+ const hasFilter = typeof prompt.filter === 'function';
1147+ if (hasFilter && !await prompt.filter()) continue;
1148+
11451149 systemPrompts.push({
11461150 identifier: key.replace(/\W/g, '_'),
11471151 position: getPromptPosition(prompt.position),
@@ -1252,7 +1256,7 @@ export async function prepareOpenAIMessages({
12521256
12531257 try {
12541258 // Merge markers and ordered user prompts with system prompts
12551259 const prompts = await preparePromptsForChatCompletion({
12561260 Scenario,
12571261 charPersonality,
12581262 name2,
public/scripts/world-info.js+1 -1
@@ -3721,7 +3721,7 @@ export async function checkWorldInfo(chat, maxContext, isDryRun) {
37213721 // Put this code here since otherwise, the chat reference is modified
37223722 for (const key of Object.keys(context.extensionPrompts)) {
37233723 if (context.extensionPrompts[key]?.scan) {
37243724 const prompt = await getExtensionPromptByName(key);
37253725 if (prompt) {
37263726 buffer.addInject(prompt);
37273727 }