diff --git a/api/app/langchain/ChatAgent.js b/api/app/langchain/ChatAgent.js index 9322000223..be4635c553 100644 --- a/api/app/langchain/ChatAgent.js +++ b/api/app/langchain/ChatAgent.js @@ -12,7 +12,8 @@ const { CallbackManager } = require('langchain/callbacks'); const { HumanChatMessage, AIChatMessage } = require('langchain/schema'); const { initializeCustomAgent, initializeFunctionsAgent } = require('./agents/'); const { getMessages, saveMessage, saveConvo } = require('../../models'); -const { loadTools, SelfReflectionTool } = require('./tools'); +const { loadTools } = require('./tools/util'); +const { SelfReflectionTool } = require('./tools/'); const { instructions, imageInstructions, diff --git a/api/app/langchain/tools/index.js b/api/app/langchain/tools/index.js index e5b4a57d99..84a8f86501 100644 --- a/api/app/langchain/tools/index.js +++ b/api/app/langchain/tools/index.js @@ -1,10 +1,19 @@ +const GoogleSearchAPI = require('./GoogleSearch'); +const HttpRequestTool = require('./HttpRequestTool'); +const AIPluginTool = require('./AIPluginTool'); +const OpenAICreateImage = require('./DALL-E'); +const StructuredSD = require('./structured/StableDiffusion'); +const StableDiffusionAPI = require('./StableDiffusion'); +const WolframAlphaAPI = require('./Wolfram'); const SelfReflectionTool = require('./SelfReflection'); -const availableTools = require('./manifest.json'); -const { validateTools, loadTools } = require('./handleTools'); module.exports = { - validateTools, - loadTools, - availableTools, + GoogleSearchAPI, + HttpRequestTool, + AIPluginTool, + OpenAICreateImage, + StructuredSD, + StableDiffusionAPI, + WolframAlphaAPI, SelfReflectionTool -}; +} diff --git a/api/app/langchain/tools/handleTools.js b/api/app/langchain/tools/util/handleTools.js similarity index 86% rename from api/app/langchain/tools/handleTools.js rename to api/app/langchain/tools/util/handleTools.js index 3d0876cb1a..34e727ec20 100644 --- a/api/app/langchain/tools/handleTools.js +++ b/api/app/langchain/tools/util/handleTools.js @@ -1,3 +1,4 @@ +const { getUserPluginAuthValue } = require('../../../../server/services/PluginService'); const { OpenAIEmbeddings } = require('langchain/embeddings/openai'); const { ZapierToolKit } = require('langchain/agents'); const { @@ -7,14 +8,16 @@ const { const { ChatOpenAI } = require('langchain/chat_models/openai'); const { Calculator } = require('langchain/tools/calculator'); const { WebBrowser } = require('langchain/tools/webbrowser'); -const GoogleSearchAPI = require('./GoogleSearch'); -const HttpRequestTool = require('./HttpRequestTool'); -const AIPluginTool = require('./AIPluginTool'); -const OpenAICreateImage = require('./DALL-E'); -const StableDiffusionAPI = require('./StableDiffusion'); -const WolframAlphaAPI = require('./Wolfram'); -const availableTools = require('./manifest.json'); -const { getUserPluginAuthValue } = require('../../../server/services/PluginService'); +const { + AIPluginTool, + GoogleSearchAPI, + WolframAlphaAPI, + HttpRequestTool, + OpenAICreateImage, + StableDiffusionAPI, + StructuredSD, +} = require('../'); +const availableTools = require('../manifest.json'); const validateTools = async (user, tools = []) => { try { @@ -70,12 +73,13 @@ const loadToolWithAuth = async (user, authFields, ToolConstructor, options = {}) }; const loadTools = async ({ user, model, tools = [], options = {} }) => { + const { functions } = options; const toolConstructors = { calculator: Calculator, google: GoogleSearchAPI, wolfram: WolframAlphaAPI, 'dall-e': OpenAICreateImage, - 'stable-diffusion': StableDiffusionAPI + 'stable-diffusion': functions ? StructuredSD : StableDiffusionAPI }; const customConstructors = { @@ -109,9 +113,10 @@ const loadTools = async ({ user, model, tools = [], options = {} }) => { return [ new HttpRequestTool(), await AIPluginTool.fromPluginUrl( - "https://www.klarna.com/.well-known/ai-plugin.json", new ChatOpenAI({ openAIApiKey: options.openAIApiKey, temperature: 0 }) - ), - ] + 'https://www.klarna.com/.well-known/ai-plugin.json', + new ChatOpenAI({ openAIApiKey: options.openAIApiKey, temperature: 0 }) + ) + ]; } }; diff --git a/api/app/langchain/tools/index.test.js b/api/app/langchain/tools/util/handleTools.test.js similarity index 94% rename from api/app/langchain/tools/index.test.js rename to api/app/langchain/tools/util/handleTools.test.js index 9cd9ccd158..d0e89ea3fa 100644 --- a/api/app/langchain/tools/index.test.js +++ b/api/app/langchain/tools/util/handleTools.test.js @@ -11,21 +11,20 @@ var mockPluginService = { }; -jest.mock('../../../models/User', () => { +jest.mock('../../../../models/User', () => { return function() { return mockUser; }; }); -jest.mock('../../../server/services/PluginService', () => mockPluginService); +jest.mock('../../../../server/services/PluginService', () => mockPluginService); -const User = require('../../../models/User'); -const { validateTools, loadTools, availableTools } = require('./index'); -const PluginService = require('../../../server/services/PluginService'); +const User = require('../../../../models/User'); +const { validateTools, loadTools, availableTools } = require('./'); +const PluginService = require('../../../../server/services/PluginService'); const { BaseChatModel } = require('langchain/chat_models/openai'); const { Calculator } = require('langchain/tools/calculator'); -const OpenAICreateImage = require('./DALL-E'); -const GoogleSearchAPI = require('./GoogleSearch'); +const { OpenAICreateImage, GoogleSearchAPI } = require('../'); describe('Tool Handlers', () => { let fakeUser; diff --git a/api/app/langchain/tools/util/index.js b/api/app/langchain/tools/util/index.js new file mode 100644 index 0000000000..e2c28e5de7 --- /dev/null +++ b/api/app/langchain/tools/util/index.js @@ -0,0 +1,8 @@ +const availableTools = require('../manifest.json'); +const { validateTools, loadTools } = require('./handleTools'); + +module.exports = { + validateTools, + loadTools, + availableTools +};