mirror of
https://github.com/run-llama/LlamaIndexTS.git
synced 2026-07-19 18:43:34 -04:00
vectorStoreIndex has new option progressCallback (#2187)
Co-authored-by: Marcus Schiesser <marcus.schiesser@googlemail.com>
This commit is contained in:
committed by
GitHub
parent
af0b79f1cd
commit
8929dcf1dd
@@ -0,0 +1,6 @@
|
||||
---
|
||||
"llamaindex": patch
|
||||
"@llamaindex/llamaindex-test": patch
|
||||
---
|
||||
|
||||
feat: vectorStoreIndex has new option progressCallback
|
||||
@@ -54,6 +54,7 @@ export interface VectorIndexOptions extends IndexStructOptions {
|
||||
storageContext?: StorageContext | undefined;
|
||||
vectorStores?: VectorStoreByType | undefined;
|
||||
logProgress?: boolean | undefined;
|
||||
progressCallback?: ((progress: number, total: number) => void) | undefined;
|
||||
}
|
||||
|
||||
export interface VectorIndexConstructorProps extends BaseIndexInit<IndexDict> {
|
||||
@@ -121,6 +122,7 @@ export class VectorStoreIndex extends BaseIndex<IndexDict> {
|
||||
// If nodes are passed in, then we need to update the index
|
||||
await index.buildIndexFromNodes(options.nodes, {
|
||||
logProgress: options.logProgress,
|
||||
progressCallback: options.progressCallback,
|
||||
});
|
||||
}
|
||||
return index;
|
||||
@@ -170,7 +172,12 @@ export class VectorStoreIndex extends BaseIndex<IndexDict> {
|
||||
*/
|
||||
async getNodeEmbeddingResults(
|
||||
nodes: BaseNode[],
|
||||
options?: { logProgress?: boolean | undefined },
|
||||
options?: {
|
||||
logProgress?: boolean | undefined;
|
||||
progressCallback?:
|
||||
| ((progress: number, total: number) => void)
|
||||
| undefined;
|
||||
},
|
||||
): Promise<BaseNode[]> {
|
||||
const nodeMap = splitNodesByType(nodes);
|
||||
for (const type in nodeMap) {
|
||||
@@ -180,6 +187,7 @@ export class VectorStoreIndex extends BaseIndex<IndexDict> {
|
||||
if (embedModel && nodes) {
|
||||
await embedModel(nodes, {
|
||||
logProgress: options?.logProgress,
|
||||
progressCallback: options?.progressCallback,
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -193,7 +201,12 @@ export class VectorStoreIndex extends BaseIndex<IndexDict> {
|
||||
*/
|
||||
async buildIndexFromNodes(
|
||||
nodes: BaseNode[],
|
||||
options?: { logProgress?: boolean | undefined },
|
||||
options?: {
|
||||
logProgress?: boolean | undefined;
|
||||
progressCallback?:
|
||||
| ((progress: number, total: number) => void)
|
||||
| undefined;
|
||||
},
|
||||
) {
|
||||
await this.insertNodes(nodes, options);
|
||||
}
|
||||
@@ -361,7 +374,12 @@ export class VectorStoreIndex extends BaseIndex<IndexDict> {
|
||||
|
||||
async insertNodes(
|
||||
nodes: BaseNode[],
|
||||
options?: { logProgress?: boolean | undefined },
|
||||
options?: {
|
||||
logProgress?: boolean | undefined;
|
||||
progressCallback?:
|
||||
| ((progress: number, total: number) => void)
|
||||
| undefined;
|
||||
},
|
||||
): Promise<void> {
|
||||
if (!nodes || nodes.length === 0) {
|
||||
return;
|
||||
|
||||
@@ -89,4 +89,42 @@ describe("[VectorStoreIndex] use embedding model", () => {
|
||||
expect(customSpy).toHaveBeenCalled();
|
||||
expect(settingsSpy).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
describe("[VectorStoreIndex] call progressCallback", () => {
|
||||
it("should call progressCallback with correct values", async () => {
|
||||
const documents = Array.from(
|
||||
{ length: 20 },
|
||||
(_, i) => new Document({ text: `This is document ${i + 1}` }),
|
||||
);
|
||||
|
||||
const progressCalls: Array<{ current: number; total: number }> = [];
|
||||
const progressCallback = (current: number, total: number) => {
|
||||
progressCalls.push({ current, total });
|
||||
};
|
||||
|
||||
const embedModel = new OpenAIEmbedding();
|
||||
mockEmbeddingModel(embedModel);
|
||||
const embedSpy = vi.spyOn(embedModel, "getTextEmbeddingsBatch");
|
||||
|
||||
Settings.embedModel = embedModel;
|
||||
const storageContext = await mockStorageContext(testDir, embedModel);
|
||||
|
||||
await VectorStoreIndex.fromDocuments(documents, {
|
||||
storageContext,
|
||||
logProgress: true,
|
||||
progressCallback,
|
||||
});
|
||||
|
||||
// Expect the embedding model to be called
|
||||
expect(embedSpy).toHaveBeenCalled();
|
||||
|
||||
// Verify that progressCallback was called with correct values
|
||||
expect(progressCalls.length).toBeGreaterThan(0);
|
||||
expect(progressCalls[0]).toEqual({ current: 10, total: 20 });
|
||||
expect(progressCalls[progressCalls.length - 1]).toEqual({
|
||||
current: 20,
|
||||
total: 20,
|
||||
});
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user