UNPKG

@websolutespa/payload-plugin-bowl-llm

Version:

LLM plugin for Bowl PayloadCms plugin

229 lines (228 loc) 9.13 kB
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