@websolutespa/payload-plugin-bowl-llm
Version:
LLM plugin for Bowl PayloadCms plugin
229 lines (228 loc) • 9.13 kB
JavaScript
import { ResponseOptions, withCors } from '@websolutespa/payload-utils/server';
import fs from 'fs';
import path from 'path';
import { addDataAndFileToRequest } from 'payload';
import { v4 as uuid } from 'uuid';
import { options } from '../options';
import { downloadFile } from '../utils/knowledgebase';
import { appHandler } from './app.handler';
import { errorHandler } from './error.handler';
import { pythonStreamHandler } from './python/pythonStream.handler';
export const downloadVectorDbFile = async (vectorDbFileUrl, destFilepath)=>{
let vectorDbFilepath = null;
if (!vectorDbFileUrl) {
return null;
}
if (!fs.existsSync(destFilepath)) {
fs.mkdirSync(destFilepath, {
recursive: true
});
}
const vectorDbFilename = path.basename(vectorDbFileUrl);
vectorDbFilepath = path.join(destFilepath, vectorDbFilename);
const vectorDbDirpath = path.join(destFilepath, path.parse(vectorDbFilename).name);
try {
if (!fs.existsSync(vectorDbFilepath) && (!fs.existsSync(vectorDbDirpath) || fs.readdirSync(vectorDbDirpath).length == 0)) {
await downloadFile(vectorDbFileUrl, vectorDbFilepath);
}
} catch (error) {
console.log('messageHandler.error', error);
}
return vectorDbFilepath;
};
export const messageHandler = async (req)=>{
if (req.method === 'OPTIONS') {
return ResponseOptions();
}
const headers = withCors({
req,
cors: '*'
});
try {
await addDataAndFileToRequest(req);
const app = await appHandler(req);
let { messages, threadId, systemContext } = req.data;
if (!messages || messages.length === 0) {
throw {
status: 400,
message: 'Bad Request: messages is missing'
};
}
// threadId is required when messages.length > 1
if (!threadId && messages.length > 1) {
throw {
status: 400,
message: 'Bad Request: threadId is missing'
};
}
if (!threadId) {
threadId = uuid();
}
const systemMessage = app.settings.llmConfig.prompt.prompt?.systemMessage || app.settings.llmConfig.prompt.systemMessage;
if (!systemMessage) {
throw {
status: 500,
message: 'The application is not configured correctly: systemMessage is missing'
};
}
const logThreadParam = req.query.logThread;
const doLogThread = logThreadParam === 'true' || logThreadParam === '1' || logThreadParam === undefined;
// set "createdAt" for the user message (if not already set by the client)
if (messages.length > 0 && !messages[messages.length - 1].createdAt) {
messages[messages.length - 1].createdAt = new Date();
}
const raw = req.query.raw !== undefined;
if (req.query.history !== undefined) {
// load messages from the thread
const thread = await req.payload.findByID({
collection: options.slug.llmThread,
id: threadId,
overrideAccess: true,
disableErrors: true
});
if (thread) {
const threadMessages = thread.message || [];
messages.unshift(...threadMessages);
}
}
let requestTools = [];
const appTools = app.settings.appTools?.filter((x)=>x.isActive) || [];
if (appTools.length > 0) {
requestTools = appTools.map((tool)=>{
// Merge mcpSettings into secrets for mcp_tool
let secrets = tool.secrets || [];
if (tool.functionName === 'mcp_tool' && tool.mcpSettings) {
const mcpSecrets = [
{
secretId: 'mcp_url',
secretValue: tool.mcpSettings.mcpUrl || ''
},
{
secretId: 'mcp_tool_name',
secretValue: tool.mcpSettings.mcpToolName || ''
},
{
secretId: 'mcp_auth_token',
secretValue: tool.mcpSettings.mcpAuthToken || ''
}
];
secrets = [
...secrets,
...mcpSecrets
];
}
return {
name: tool.name,
description: tool.description,
type: tool.type,
functionId: tool.functionId,
functionName: tool.functionName,
functionDescription: tool.functionDescription,
secrets,
llmChainSettings: tool.llmChainSettings,
searchSettings: tool.searchSettings,
dataSource: tool.dataSource,
dbSettings: tool.dbSettings,
apiSettings: tool.apiSettings,
integrations: tool.knowledgeBase?.integrations || [],
endpoints: tool.knowledgeBase?.externalEndpoints || [],
waitingMessage: tool.waitingMessage,
vectorDbType: tool.knowledgeBase?.vectorDbType || '',
vectorDbFile: tool.knowledgeBase?.vectorDbFile?.filename,
rulesVectorDb: app.settings.rules.vectorDbFile?.filename
};
});
}
const request = {
messages: messages.map((x)=>({
role: x.role,
content: x.content
})),
mode: app.mode,
provider: app.settings.llmConfig.provider,
model: app.settings.llmConfig.model,
temperature: app.settings.llmConfig.temperature,
secrets: app.settings.llmConfig.secrets,
systemMessage: systemMessage,
systemContext: systemContext || {},
threadId: threadId,
msgId: uuid(),
tools: app.settings.llmConfig.tools,
appTools: requestTools,
vectorDb: app.settings.knowledgeBase.vectorDbFile?.filename,
rules: {
vector_db: app.settings.rules.vectorDbFile?.filename,
threshold: app.settings.rules.threshold
},
fineTunedModel: app.settings?.fineTuning?.fineTunedModelName ?? '',
langChainTracing: app.settings?.llmConfig?.langChainTracing ?? false,
langChainProject: app.settings?.llmConfig?.langChainProject ?? '',
outputStructure: app.settings.llmConfig.outputStructure || null
};
const streamHandler = pythonStreamHandler(request, raw, (response)=>{
if (doLogThread) {
const userMsg = messages.pop() || {
role: 'user',
content: '',
messageId: '',
createdAt: new Date()
};
const assistantMsg = {
role: 'assistant',
content: raw && response.chunks.length > 0 ? response.chunks[0].content : response.chunks,
messageId: response.messageId,
createdAt: new Date()
};
// log thread messages
logThread(req.payload, threadId, app, userMsg, assistantMsg).then((thread)=>{
// console.log('thread logged', thread);
});
}
});
return await streamHandler(req);
} catch (error) {
return errorHandler(error, {
headers
});
}
};
async function logThread(payload, threadId, app, userMsg, assistantMsg) {
const threads = await payload.find({
collection: options.slug.llmThread,
where: {
'id': {
equals: threadId
}
},
overrideAccess: true
});
const thread = threads.docs.length ? threads.docs[0] : null;
const messages = thread ? thread.message : [];
messages.push(userMsg);
messages.push(assistantMsg);
const threadMessages = messages.map((message)=>({
...message,
content: typeof message.content === 'string' ? message.content : message.content.map((x)=>JSON.stringify(x)).join(',') + ','
}));
if (thread) {
return await payload.update({
collection: options.slug.llmThread,
id: thread.id,
data: {
message: threadMessages
},
overrideAccess: true
});
} else {
return await payload.create({
collection: options.slug.llmThread,
data: {
id: threadId,
llmApp: app.id,
message: threadMessages
},
overrideAccess: true
});
}
}
//# sourceMappingURL=message.handler.js.map