first commit
Security: Sync from Public / sync-from-public (push) Has been cancelled
Test: Benchmark Nightly / build (push) Has been cancelled
Test: Benchmark Nightly / Notify Cats on failure (push) Has been cancelled
CI: Python / Checks (push) Has been cancelled
Test: Evals Python / Workflow Comparison Python (push) Has been cancelled
Util: Check Docs URLs / check-docs-urls (push) Has been cancelled
Test: Visual Storybook / Cloudflare Pages (push) Has been cancelled
Test: E2E Performance / build-and-test-performance (push) Has been cancelled
Test: Workflows Nightly / Run Workflow Tests (push) Has been cancelled
Util: Cleanup CI Docker Images / Delete stale CI images (push) Has been cancelled
Test: Benchmark Destroy Env / build (push) Has been cancelled
Util: Update Node Popularity / update-popularity (push) Has been cancelled
Test: E2E Coverage Weekly / Coverage Tests (push) Has been cancelled
Security: Sync from Public / sync-from-public (push) Has been cancelled
Test: Benchmark Nightly / build (push) Has been cancelled
Test: Benchmark Nightly / Notify Cats on failure (push) Has been cancelled
CI: Python / Checks (push) Has been cancelled
Test: Evals Python / Workflow Comparison Python (push) Has been cancelled
Util: Check Docs URLs / check-docs-urls (push) Has been cancelled
Test: Visual Storybook / Cloudflare Pages (push) Has been cancelled
Test: E2E Performance / build-and-test-performance (push) Has been cancelled
Test: Workflows Nightly / Run Workflow Tests (push) Has been cancelled
Util: Cleanup CI Docker Images / Delete stale CI images (push) Has been cancelled
Test: Benchmark Destroy Env / build (push) Has been cancelled
Util: Update Node Popularity / update-popularity (push) Has been cancelled
Test: E2E Coverage Weekly / Coverage Tests (push) Has been cancelled
This commit is contained in:
+102
@@ -0,0 +1,102 @@
|
||||
import type { BaseDocumentCompressor } from '@langchain/core/retrievers/document_compressors';
|
||||
import { VectorStore } from '@langchain/core/vectorstores';
|
||||
import { ContextualCompressionRetriever } from '@langchain/classic/retrievers/contextual_compression';
|
||||
import {
|
||||
NodeConnectionTypes,
|
||||
type INodeType,
|
||||
type INodeTypeDescription,
|
||||
type ISupplyDataFunctions,
|
||||
type SupplyData,
|
||||
} from 'n8n-workflow';
|
||||
|
||||
import { logWrapper } from '@n8n/ai-utilities';
|
||||
|
||||
export class RetrieverVectorStore implements INodeType {
|
||||
description: INodeTypeDescription = {
|
||||
displayName: 'Vector Store Retriever',
|
||||
name: 'retrieverVectorStore',
|
||||
icon: 'fa:box-open',
|
||||
iconColor: 'black',
|
||||
group: ['transform'],
|
||||
version: 1,
|
||||
description: 'Use a Vector Store as Retriever',
|
||||
defaults: {
|
||||
name: 'Vector Store Retriever',
|
||||
},
|
||||
codex: {
|
||||
categories: ['AI'],
|
||||
subcategories: {
|
||||
AI: ['Retrievers'],
|
||||
},
|
||||
resources: {
|
||||
primaryDocumentation: [
|
||||
{
|
||||
url: 'https://docs.n8n.io/integrations/builtin/cluster-nodes/sub-nodes/n8n-nodes-langchain.retrievervectorstore/',
|
||||
},
|
||||
],
|
||||
},
|
||||
},
|
||||
|
||||
inputs: [
|
||||
{
|
||||
displayName: 'Vector Store',
|
||||
maxConnections: 1,
|
||||
type: NodeConnectionTypes.AiVectorStore,
|
||||
required: true,
|
||||
},
|
||||
],
|
||||
|
||||
outputs: [NodeConnectionTypes.AiRetriever],
|
||||
outputNames: ['Retriever'],
|
||||
builderHint: {
|
||||
relatedNodes: [
|
||||
{
|
||||
nodeType: '@n8n/n8n-nodes-langchain.vectorStoreInMemory',
|
||||
relationHint: 'Connect to provide vectors for retrieval in RAG workflows',
|
||||
},
|
||||
],
|
||||
inputs: {
|
||||
ai_vectorStore: { required: true },
|
||||
},
|
||||
},
|
||||
properties: [
|
||||
{
|
||||
displayName: 'Limit',
|
||||
name: 'topK',
|
||||
type: 'number',
|
||||
default: 4,
|
||||
description: 'The maximum number of results to return',
|
||||
},
|
||||
],
|
||||
};
|
||||
|
||||
async supplyData(this: ISupplyDataFunctions, itemIndex: number): Promise<SupplyData> {
|
||||
this.logger.debug('Supplying data for Vector Store Retriever');
|
||||
|
||||
const topK = this.getNodeParameter('topK', itemIndex, 4) as number;
|
||||
const vectorStore = (await this.getInputConnectionData(
|
||||
NodeConnectionTypes.AiVectorStore,
|
||||
itemIndex,
|
||||
)) as
|
||||
| VectorStore
|
||||
| {
|
||||
reranker: BaseDocumentCompressor;
|
||||
vectorStore: VectorStore;
|
||||
};
|
||||
|
||||
let retriever = null;
|
||||
|
||||
if (vectorStore instanceof VectorStore) {
|
||||
retriever = vectorStore.asRetriever(topK);
|
||||
} else {
|
||||
retriever = new ContextualCompressionRetriever({
|
||||
baseCompressor: vectorStore.reranker,
|
||||
baseRetriever: vectorStore.vectorStore.asRetriever(topK),
|
||||
});
|
||||
}
|
||||
|
||||
return {
|
||||
response: logWrapper(retriever, this),
|
||||
};
|
||||
}
|
||||
}
|
||||
+133
@@ -0,0 +1,133 @@
|
||||
import type { BaseDocumentCompressor } from '@langchain/core/retrievers/document_compressors';
|
||||
import { VectorStore } from '@langchain/core/vectorstores';
|
||||
import { ContextualCompressionRetriever } from '@langchain/classic/retrievers/contextual_compression';
|
||||
import type { ISupplyDataFunctions } from 'n8n-workflow';
|
||||
import { NodeConnectionTypes } from 'n8n-workflow';
|
||||
|
||||
import { RetrieverVectorStore } from '../RetrieverVectorStore.node';
|
||||
|
||||
const mockLogger = {
|
||||
debug: jest.fn(),
|
||||
info: jest.fn(),
|
||||
warn: jest.fn(),
|
||||
error: jest.fn(),
|
||||
};
|
||||
|
||||
describe('RetrieverVectorStore', () => {
|
||||
let retrieverNode: RetrieverVectorStore;
|
||||
let mockContext: jest.Mocked<ISupplyDataFunctions>;
|
||||
|
||||
beforeEach(() => {
|
||||
retrieverNode = new RetrieverVectorStore();
|
||||
mockContext = {
|
||||
logger: mockLogger,
|
||||
getNodeParameter: jest.fn(),
|
||||
getInputConnectionData: jest.fn(),
|
||||
} as unknown as jest.Mocked<ISupplyDataFunctions>;
|
||||
jest.clearAllMocks();
|
||||
});
|
||||
|
||||
describe('supplyData', () => {
|
||||
it('should create a retriever from a basic VectorStore', async () => {
|
||||
const mockVectorStore = Object.create(VectorStore.prototype) as VectorStore;
|
||||
mockVectorStore.asRetriever = jest.fn().mockReturnValue({ test: 'retriever' });
|
||||
|
||||
mockContext.getNodeParameter.mockImplementation((param, _itemIndex, defaultValue) => {
|
||||
if (param === 'topK') return 4;
|
||||
return defaultValue;
|
||||
});
|
||||
|
||||
mockContext.getInputConnectionData.mockResolvedValue(mockVectorStore);
|
||||
|
||||
const result = await retrieverNode.supplyData.call(mockContext, 0);
|
||||
|
||||
expect(mockContext.getInputConnectionData).toHaveBeenCalledWith(
|
||||
NodeConnectionTypes.AiVectorStore,
|
||||
0,
|
||||
);
|
||||
expect(mockVectorStore.asRetriever).toHaveBeenCalledWith(4);
|
||||
expect(result).toHaveProperty('response', { test: 'retriever' });
|
||||
});
|
||||
|
||||
it('should create a retriever with custom topK parameter', async () => {
|
||||
const mockVectorStore = Object.create(VectorStore.prototype) as VectorStore;
|
||||
mockVectorStore.asRetriever = jest.fn().mockReturnValue({ test: 'retriever' });
|
||||
|
||||
mockContext.getNodeParameter.mockImplementation((param, _itemIndex, defaultValue) => {
|
||||
if (param === 'topK') return 10;
|
||||
return defaultValue;
|
||||
});
|
||||
mockContext.getInputConnectionData.mockResolvedValue(mockVectorStore);
|
||||
|
||||
const result = await retrieverNode.supplyData.call(mockContext, 0);
|
||||
|
||||
expect(mockVectorStore.asRetriever).toHaveBeenCalledWith(10);
|
||||
expect(result).toHaveProperty('response', { test: 'retriever' });
|
||||
});
|
||||
|
||||
it('should create a ContextualCompressionRetriever when input contains reranker and vectorStore', async () => {
|
||||
const mockVectorStore = Object.create(VectorStore.prototype) as VectorStore;
|
||||
mockVectorStore.asRetriever = jest.fn().mockReturnValue({ test: 'base-retriever' });
|
||||
|
||||
const mockReranker = {} as BaseDocumentCompressor;
|
||||
|
||||
const inputWithReranker = {
|
||||
reranker: mockReranker,
|
||||
vectorStore: mockVectorStore,
|
||||
};
|
||||
|
||||
mockContext.getNodeParameter.mockImplementation((param, _itemIndex, defaultValue) => {
|
||||
if (param === 'topK') return 4;
|
||||
return defaultValue;
|
||||
});
|
||||
mockContext.getInputConnectionData.mockResolvedValue(inputWithReranker);
|
||||
|
||||
const result = await retrieverNode.supplyData.call(mockContext, 0);
|
||||
|
||||
expect(mockContext.getInputConnectionData).toHaveBeenCalledWith(
|
||||
NodeConnectionTypes.AiVectorStore,
|
||||
0,
|
||||
);
|
||||
expect(mockVectorStore.asRetriever).toHaveBeenCalledWith(4);
|
||||
expect(result.response).toBeInstanceOf(ContextualCompressionRetriever);
|
||||
});
|
||||
|
||||
it('should create a ContextualCompressionRetriever with custom topK when using reranker', async () => {
|
||||
const mockVectorStore = Object.create(VectorStore.prototype) as VectorStore;
|
||||
mockVectorStore.asRetriever = jest.fn().mockReturnValue({ test: 'base-retriever' });
|
||||
|
||||
const mockReranker = {} as BaseDocumentCompressor;
|
||||
|
||||
const inputWithReranker = {
|
||||
reranker: mockReranker,
|
||||
vectorStore: mockVectorStore,
|
||||
};
|
||||
|
||||
mockContext.getNodeParameter.mockImplementation((param, _itemIndex, defaultValue) => {
|
||||
if (param === 'topK') return 8;
|
||||
return defaultValue;
|
||||
});
|
||||
mockContext.getInputConnectionData.mockResolvedValue(inputWithReranker);
|
||||
|
||||
const result = await retrieverNode.supplyData.call(mockContext, 0);
|
||||
|
||||
expect(mockVectorStore.asRetriever).toHaveBeenCalledWith(8);
|
||||
expect(result.response).toBeInstanceOf(ContextualCompressionRetriever);
|
||||
});
|
||||
|
||||
it('should use default topK value when parameter is not provided', async () => {
|
||||
const mockVectorStore = Object.create(VectorStore.prototype) as VectorStore;
|
||||
mockVectorStore.asRetriever = jest.fn().mockReturnValue({ test: 'retriever' });
|
||||
|
||||
mockContext.getNodeParameter.mockImplementation((_param, _itemIndex, defaultValue) => {
|
||||
return defaultValue;
|
||||
});
|
||||
mockContext.getInputConnectionData.mockResolvedValue(mockVectorStore);
|
||||
|
||||
await retrieverNode.supplyData.call(mockContext, 0);
|
||||
|
||||
expect(mockContext.getNodeParameter).toHaveBeenCalledWith('topK', 0, 4);
|
||||
expect(mockVectorStore.asRetriever).toHaveBeenCalledWith(4);
|
||||
});
|
||||
});
|
||||
});
|
||||
Reference in New Issue
Block a user