Emit prompt events in generateRaw (#4587) Closes https://github.com/SillyTavern/Extension-PromptInspector/issues/4
Signed| @@ -3197,6 +3197,11 @@ export async function generateRaw({ prompt = '', api = null, instructOverride = | ||
| 3197 | 3197 | // construct final prompt from the input. Can either be a string or an array of chat-style messages. |
| 3198 | 3198 | prompt = createRawPrompt(prompt, api, instructOverride, quietToLoud, systemPrompt, prefill); |
| 3199 | 3199 | |
| 3200 | + // Allow extensions to stop generation before it happens | |
| 3201 | + const eventAbortController = new AbortController(); | |
| 3202 | + const abortHook = () => eventAbortController.abort(new Error('Cancelled by extension')); | |
| 3203 | + eventSource.on(event_types.GENERATION_STOPPED, abortHook); | |
| 3204 | + | |
| 3200 | 3205 | try { |
| 3201 | 3206 | if (responseLengthCustomized) { |
| 3202 | 3207 | TempResponseLength.save(api, responseLength); |
| @@ -3204,6 +3209,23 @@ export async function generateRaw({ prompt = '', api = null, instructOverride = | ||
| 3204 | 3209 | /** @type {object|any[]} */ |
| 3205 | 3210 | let generateData = {}; |
| 3206 | 3211 | |
| 3212 | + // Allow extensions to modify the prompt before generation | |
| 3213 | + // 1. for text completion | |
| 3214 | + if (typeof prompt === 'string') { | |
| 3215 | + const eventData = { prompt: prompt, dryRun: false }; | |
| 3216 | + await eventSource.emit(event_types.GENERATE_AFTER_COMBINE_PROMPTS, eventData); | |
| 3217 | + prompt = eventData.prompt; | |
| 3218 | + } | |
| 3219 | + // 2. for chat completion | |
| 3220 | + if (Array.isArray(prompt)) { | |
| 3221 | + const eventData = { chat: prompt, dryRun: false }; | |
| 3222 | + await eventSource.emit(event_types.CHAT_COMPLETION_PROMPT_READY, eventData); | |
| 3223 | + prompt = eventData.chat; | |
| 3224 | + } | |
| 3225 | + | |
| 3226 | + // Check if the generation was aborted during the event | |
| 3227 | + eventAbortController.signal.throwIfAborted(); | |
| 3228 | + | |
| 3207 | 3229 | switch (api) { |
| 3208 | 3230 | case 'kobold': |
| 3209 | 3231 | case 'koboldhorde': |
| @@ -3283,6 +3305,7 @@ export async function generateRaw({ prompt = '', api = null, instructOverride = | ||
| 3283 | 3305 | |
| 3284 | 3306 | return message; |
| 3285 | 3307 | } finally { |
| 3308 | + eventSource.removeListener(event_types.GENERATION_STOPPED, abortHook); | |
| 3286 | 3309 | if (responseLengthCustomized && TempResponseLength.isCustomized()) { |
| 3287 | 3310 | TempResponseLength.restore(api); |
| 3288 | 3311 | TempResponseLength.removeEventHook(api, eventHook); |