Blame Raw
Cohee · 51ad27fb · · 281 lines (10.1 KB)
2 contributors
1import fs from 'node:fs';
2import express from 'express';
3import fetch from 'node-fetch';
4
5import { forwardFetchResponse, delay } from '../../util.js';
6import { getOverrideHeaders, setAdditionalHeaders, setAdditionalHeadersByType } from '../../additional-headers.js';
7import { TEXTGEN_TYPES } from '../../constants.js';
8
9export const router = express.Router();
10
11router.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
143router.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
190router.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
241router.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});