| 1 | import fs from 'node:fs'; |
| 2 | import express from 'express'; |
| 3 | import fetch from 'node-fetch'; |
| 4 | |
| 5 | import { forwardFetchResponse, delay } from '../../util.js'; |
| 6 | import { getOverrideHeaders, setAdditionalHeaders, setAdditionalHeadersByType } from '../../additional-headers.js'; |
| 7 | import { TEXTGEN_TYPES } from '../../constants.js'; |
| 8 | |
| 9 | export const router = express.Router(); |
| 10 | |
| 11 | router.post('/generate', async function (request, response_generate) { |
| 12 | if (!request.body) return response_generate.sendStatus(400); |
| 13 | |
| 14 | if (request.body.api_server.indexOf('localhost') != -1) { |
| 15 | request.body.api_server = request.body.api_server.replace('localhost', '127.0.0.1'); |
| 16 | } |
| 17 | |
| 18 | const request_prompt = request.body.prompt; |
| 19 | const controller = new AbortController(); |
| 20 | request.socket.removeAllListeners('close'); |
| 21 | request.socket.on('close', async function () { |
| 22 | if (request.body.can_abort && !response_generate.writableEnded) { |
| 23 | try { |
| 24 | console.info('Aborting Kobold generation...'); |
| 25 | // send abort signal to koboldcpp |
| 26 | const abortResponse = await fetch(`${request.body.api_server}/extra/abort`, { |
| 27 | method: 'POST', |
| 28 | }); |
| 29 | |
| 30 | if (!abortResponse.ok) { |
| 31 | console.error('Error sending abort request to Kobold:', abortResponse.status); |
| 32 | } |
| 33 | } catch (error) { |
| 34 | console.error(error); |
| 35 | } |
| 36 | } |
| 37 | controller.abort(); |
| 38 | }); |
| 39 | |
| 40 | let this_settings = { |
| 41 | prompt: request_prompt, |
| 42 | use_story: false, |
| 43 | use_memory: false, |
| 44 | use_authors_note: false, |
| 45 | use_world_info: false, |
| 46 | max_context_length: request.body.max_context_length, |
| 47 | max_length: request.body.max_length, |
| 48 | }; |
| 49 | |
| 50 | if (!request.body.gui_settings) { |
| 51 | this_settings = { |
| 52 | prompt: request_prompt, |
| 53 | use_story: false, |
| 54 | use_memory: false, |
| 55 | use_authors_note: false, |
| 56 | use_world_info: false, |
| 57 | max_context_length: request.body.max_context_length, |
| 58 | max_length: request.body.max_length, |
| 59 | rep_pen: request.body.rep_pen, |
| 60 | rep_pen_range: request.body.rep_pen_range, |
| 61 | rep_pen_slope: request.body.rep_pen_slope, |
| 62 | temperature: request.body.temperature, |
| 63 | tfs: request.body.tfs, |
| 64 | top_a: request.body.top_a, |
| 65 | top_k: request.body.top_k, |
| 66 | top_p: request.body.top_p, |
| 67 | min_p: request.body.min_p, |
| 68 | typical: request.body.typical, |
| 69 | sampler_order: request.body.sampler_order, |
| 70 | singleline: !!request.body.singleline, |
| 71 | use_default_badwordsids: request.body.use_default_badwordsids, |
| 72 | mirostat: request.body.mirostat, |
| 73 | mirostat_eta: request.body.mirostat_eta, |
| 74 | mirostat_tau: request.body.mirostat_tau, |
| 75 | grammar: request.body.grammar, |
| 76 | sampler_seed: request.body.sampler_seed, |
| 77 | }; |
| 78 | if (request.body.stop_sequence) { |
| 79 | this_settings.stop_sequence = request.body.stop_sequence; |
| 80 | } |
| 81 | } |
| 82 | |
| 83 | console.debug(this_settings); |
| 84 | const args = { |
| 85 | body: JSON.stringify(this_settings), |
| 86 | headers: Object.assign( |
| 87 | { 'Content-Type': 'application/json' }, |
| 88 | getOverrideHeaders((new URL(request.body.api_server))?.host), |
| 89 | ), |
| 90 | signal: controller.signal, |
| 91 | }; |
| 92 | |
| 93 | const MAX_RETRIES = 50; |
| 94 | const delayAmount = 2500; |
| 95 | for (let i = 0; i < MAX_RETRIES; i++) { |
| 96 | try { |
| 97 | const url = request.body.streaming ? `${request.body.api_server}/extra/generate/stream` : `${request.body.api_server}/v1/generate`; |
| 98 | const response = await fetch(url, { method: 'POST', ...args }); |
| 99 | |
| 100 | if (request.body.streaming) { |
| 101 | // Pipe remote SSE stream to Express response |
| 102 | await forwardFetchResponse(response, response_generate); |
| 103 | return; |
| 104 | } else { |
| 105 | if (!response.ok) { |
| 106 | const errorText = await response.text(); |
| 107 | console.warn(`Kobold returned error: ${response.status} ${response.statusText} ${errorText}`); |
| 108 | |
| 109 | try { |
| 110 | const errorJson = JSON.parse(errorText); |
| 111 | const message = errorJson?.detail?.msg || errorText; |
| 112 | return response_generate.status(400).send({ error: { message } }); |
| 113 | } catch { |
| 114 | return response_generate.status(400).send({ error: { message: errorText } }); |
| 115 | } |
| 116 | } |
| 117 | |
| 118 | const data = await response.json(); |
| 119 | console.debug('Endpoint response:', data); |
| 120 | return response_generate.send(data); |
| 121 | } |
| 122 | } catch (error) { |
| 123 | // response |
| 124 | switch (error?.status) { |
| 125 | case 403: |
| 126 | case 503: // retry in case of temporary service issue, possibly caused by a queue failure? |
| 127 | console.warn(`KoboldAI is busy. Retry attempt ${i + 1} of ${MAX_RETRIES}...`); |
| 128 | await delay(delayAmount); |
| 129 | break; |
| 130 | default: |
| 131 | if ('status' in error) { |
| 132 | console.error('Status Code from Kobold:', error.status); |
| 133 | } |
| 134 | return response_generate.send({ error: true }); |
| 135 | } |
| 136 | } |
| 137 | } |
| 138 | |
| 139 | console.error('Max retries exceeded. Giving up.'); |
| 140 | return response_generate.send({ error: true }); |
| 141 | }); |
| 142 | |
| 143 | router.post('/status', async function (request, response) { |
| 144 | if (!request.body) return response.sendStatus(400); |
| 145 | let api_server = request.body.api_server; |
| 146 | if (api_server.indexOf('localhost') != -1) { |
| 147 | api_server = api_server.replace('localhost', '127.0.0.1'); |
| 148 | } |
| 149 | |
| 150 | const args = { |
| 151 | headers: { 'Content-Type': 'application/json' }, |
| 152 | }; |
| 153 | |
| 154 | setAdditionalHeaders(request, args, api_server); |
| 155 | |
| 156 | const result = {}; |
| 157 | |
| 158 | /** @type {any} */ |
| 159 | const [koboldUnitedResponse, koboldExtraResponse, koboldModelResponse] = await Promise.all([ |
| 160 | // We catch errors both from the response not having a successful HTTP status and from JSON parsing failing |
| 161 | |
| 162 | // Kobold United API version |
| 163 | fetch(`${api_server}/v1/info/version`).then(response => { |
| 164 | if (!response.ok) throw new Error(`Kobold API error: ${response.status, response.statusText}`); |
| 165 | return response.json(); |
| 166 | }).catch(() => ({ result: '0.0.0' })), |
| 167 | |
| 168 | // KoboldCpp version |
| 169 | fetch(`${api_server}/extra/version`).then(response => { |
| 170 | if (!response.ok) throw new Error(`Kobold API error: ${response.status, response.statusText}`); |
| 171 | return response.json(); |
| 172 | }).catch(() => ({ version: '0.0' })), |
| 173 | |
| 174 | // Current model |
| 175 | fetch(`${api_server}/v1/model`).then(response => { |
| 176 | if (!response.ok) throw new Error(`Kobold API error: ${response.status, response.statusText}`); |
| 177 | return response.json(); |
| 178 | }).catch(() => null), |
| 179 | ]); |
| 180 | |
| 181 | result.koboldUnitedVersion = koboldUnitedResponse.result; |
| 182 | result.koboldCppVersion = koboldExtraResponse.result; |
| 183 | result.model = !koboldModelResponse || koboldModelResponse.result === 'ReadOnly' ? |
| 184 | 'no_connection' : |
| 185 | koboldModelResponse.result; |
| 186 | |
| 187 | response.send(result); |
| 188 | }); |
| 189 | |
| 190 | router.post('/transcribe-audio', async function (request, response) { |
| 191 | try { |
| 192 | const server = request.body.server; |
| 193 | |
| 194 | if (!server) { |
| 195 | console.error('Server is not set'); |
| 196 | return response.sendStatus(400); |
| 197 | } |
| 198 | |
| 199 | if (!request.file) { |
| 200 | console.error('No audio file found'); |
| 201 | return response.sendStatus(400); |
| 202 | } |
| 203 | |
| 204 | console.debug('Transcribing audio with KoboldCpp', server); |
| 205 | |
| 206 | const fileBase64 = fs.readFileSync(request.file.path).toString('base64'); |
| 207 | fs.unlinkSync(request.file.path); |
| 208 | |
| 209 | const headers = {}; |
| 210 | setAdditionalHeadersByType(headers, TEXTGEN_TYPES.KOBOLDCPP, server, request.user.directories); |
| 211 | |
| 212 | const url = new URL(server); |
| 213 | url.pathname = '/api/extra/transcribe'; |
| 214 | |
| 215 | const result = await fetch(url, { |
| 216 | method: 'POST', |
| 217 | headers: { |
| 218 | ...headers, |
| 219 | }, |
| 220 | body: JSON.stringify({ |
| 221 | prompt: '', |
| 222 | audio_data: fileBase64, |
| 223 | }), |
| 224 | }); |
| 225 | |
| 226 | if (!result.ok) { |
| 227 | const text = await result.text(); |
| 228 | console.error('KoboldCpp request failed', result.statusText, text); |
| 229 | return response.status(500).send(text); |
| 230 | } |
| 231 | |
| 232 | const data = await result.json(); |
| 233 | console.debug('KoboldCpp transcription response', data); |
| 234 | return response.json(data); |
| 235 | } catch (error) { |
| 236 | console.error('KoboldCpp transcription failed', error); |
| 237 | response.status(500).send('Internal server error'); |
| 238 | } |
| 239 | }); |
| 240 | |
| 241 | router.post('/embed', async function (request, response) { |
| 242 | try { |
| 243 | const { server, items } = request.body; |
| 244 | |
| 245 | if (!server) { |
| 246 | console.warn('KoboldCpp URL is not set'); |
| 247 | return response.sendStatus(400); |
| 248 | } |
| 249 | |
| 250 | const headers = {}; |
| 251 | setAdditionalHeadersByType(headers, TEXTGEN_TYPES.KOBOLDCPP, server, request.user.directories); |
| 252 | |
| 253 | const embeddingsUrl = new URL(server); |
| 254 | embeddingsUrl.pathname = '/api/extra/embeddings'; |
| 255 | |
| 256 | const embeddingsResult = await fetch(embeddingsUrl, { |
| 257 | method: 'POST', |
| 258 | headers: { |
| 259 | ...headers, |
| 260 | }, |
| 261 | body: JSON.stringify({ |
| 262 | input: items, |
| 263 | }), |
| 264 | }); |
| 265 | |
| 266 | /** @type {any} */ |
| 267 | const data = await embeddingsResult.json(); |
| 268 | |
| 269 | if (!Array.isArray(data?.data)) { |
| 270 | console.warn('KoboldCpp API response was not an array'); |
| 271 | return response.sendStatus(500); |
| 272 | } |
| 273 | |
| 274 | const model = data.model || 'unknown'; |
| 275 | const embeddings = data.data.map(x => Array.isArray(x) ? x[0] : x).sort((a, b) => a.index - b.index).map(x => x.embedding); |
| 276 | return response.json({ model, embeddings }); |
| 277 | } catch (error) { |
| 278 | console.error('KoboldCpp embedding failed', error); |
| 279 | response.status(500).send('Internal server error'); |
| 280 | } |
| 281 | }); |