mirror of
https://github.com/run-llama/LlamaIndexTS.git
synced 2026-07-19 18:43:34 -04:00
feat: simplify callback manager (#1027)
This commit is contained in:
@@ -0,0 +1,9 @@
|
||||
---
|
||||
"llamaindex": patch
|
||||
"@llamaindex/core": patch
|
||||
---
|
||||
|
||||
refactor: move callback manager & llm to core module
|
||||
|
||||
For people who import `llamaindex/llms/base` or `llamaindex/llms/utils`,
|
||||
use `@llamaindex/core/llms` and `@llamaindex/core/utils` instead.
|
||||
@@ -0,0 +1,5 @@
|
||||
---
|
||||
"@llamaindex/community": patch
|
||||
---
|
||||
|
||||
refactor: depends on core pacakge instead of llamaindex
|
||||
@@ -0,0 +1,8 @@
|
||||
---
|
||||
"llamaindex": minor
|
||||
"@llamaindex/core": minor
|
||||
---
|
||||
|
||||
refactor: simplify callback manager
|
||||
|
||||
Change `event.detail.payload` to `event.detail`
|
||||
@@ -2,7 +2,7 @@ import { Anthropic, FunctionTool, Settings, WikipediaTool } from "llamaindex";
|
||||
import { AnthropicAgent } from "llamaindex/agent/anthropic";
|
||||
|
||||
Settings.callbackManager.on("llm-tool-call", (event) => {
|
||||
console.log("llm-tool-call", event.detail.payload.toolCall);
|
||||
console.log("llm-tool-call", event.detail.toolCall);
|
||||
});
|
||||
|
||||
const anthropic = new Anthropic({
|
||||
|
||||
@@ -4,7 +4,6 @@ import {
|
||||
NodeWithScore,
|
||||
ObjectType,
|
||||
OpenAI,
|
||||
RetrievalEndEvent,
|
||||
Settings,
|
||||
VectorStoreIndex,
|
||||
} from "llamaindex";
|
||||
@@ -18,8 +17,8 @@ Settings.chunkOverlap = 20;
|
||||
Settings.llm = new OpenAI({ model: "gpt-4-turbo", maxTokens: 512 });
|
||||
|
||||
// Update callbackManager
|
||||
Settings.callbackManager.on("retrieve-end", (event: RetrievalEndEvent) => {
|
||||
const { nodes, query } = event.detail.payload;
|
||||
Settings.callbackManager.on("retrieve-end", (event) => {
|
||||
const { nodes, query } = event.detail;
|
||||
const imageNodes = nodes.filter(
|
||||
(node: NodeWithScore) => node.node.type === ObjectType.IMAGE_DOCUMENT,
|
||||
);
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
import {
|
||||
MultiModalResponseSynthesizer,
|
||||
OpenAI,
|
||||
RetrievalEndEvent,
|
||||
Settings,
|
||||
VectorStoreIndex,
|
||||
} from "llamaindex";
|
||||
@@ -15,8 +14,8 @@ Settings.chunkOverlap = 20;
|
||||
Settings.llm = new OpenAI({ model: "gpt-4-turbo", maxTokens: 512 });
|
||||
|
||||
// Update callbackManager
|
||||
Settings.callbackManager.on("retrieve-end", (event: RetrievalEndEvent) => {
|
||||
const { nodes, query } = event.detail.payload;
|
||||
Settings.callbackManager.on("retrieve-end", (event) => {
|
||||
const { nodes, query } = event.detail;
|
||||
console.log(`Retrieved ${nodes.length} nodes for query: ${query}`);
|
||||
});
|
||||
|
||||
|
||||
@@ -11,12 +11,10 @@ import {
|
||||
|
||||
// Update callback manager
|
||||
Settings.callbackManager.on("retrieve-end", (event) => {
|
||||
const data = event.detail.payload;
|
||||
const { nodes } = event.detail;
|
||||
console.log(
|
||||
"The retrieved nodes are:",
|
||||
data.nodes.map((node: NodeWithScore) =>
|
||||
node.node.getContent(MetadataMode.NONE),
|
||||
),
|
||||
nodes.map((node: NodeWithScore) => node.node.getContent(MetadataMode.NONE)),
|
||||
);
|
||||
});
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import { extractText } from "@llamaindex/core/utils";
|
||||
import { encodingForModel } from "js-tiktoken";
|
||||
import { ChatMessage, OpenAI, type LLMStartEvent } from "llamaindex";
|
||||
import { ChatMessage, OpenAI } from "llamaindex";
|
||||
import { Settings } from "llamaindex/Settings";
|
||||
|
||||
const encoding = encodingForModel("gpt-4-0125-preview");
|
||||
@@ -12,8 +12,8 @@ const llm = new OpenAI({
|
||||
|
||||
let tokenCount = 0;
|
||||
|
||||
Settings.callbackManager.on("llm-start", (event: LLMStartEvent) => {
|
||||
const { messages } = event.detail.payload;
|
||||
Settings.callbackManager.on("llm-start", (event) => {
|
||||
const { messages } = event.detail;
|
||||
messages.reduce((count: number, message: ChatMessage) => {
|
||||
return count + encoding.encode(extractText(message.content)).length;
|
||||
}, 0);
|
||||
@@ -24,7 +24,7 @@ Settings.callbackManager.on("llm-start", (event: LLMStartEvent) => {
|
||||
});
|
||||
|
||||
Settings.callbackManager.on("llm-stream", (event) => {
|
||||
const { chunk } = event.detail.payload;
|
||||
const { chunk } = event.detail;
|
||||
const { delta } = chunk;
|
||||
tokenCount += encoding.encode(extractText(delta)).length;
|
||||
if (tokenCount > 20) {
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
export { Settings } from "./settings";
|
||||
export { CallbackManager } from "./settings/callback-manager";
|
||||
export type {
|
||||
BaseEvent,
|
||||
LLMEndEvent,
|
||||
LLMStartEvent,
|
||||
LLMStreamEvent,
|
||||
|
||||
@@ -6,31 +6,32 @@ import type {
|
||||
ToolCall,
|
||||
ToolOutput,
|
||||
} from "../../llms";
|
||||
import { EventCaller, getEventCaller } from "../../utils/event-caller";
|
||||
import type { UUID } from "../type";
|
||||
|
||||
export type BaseEvent<Payload> = CustomEvent<{
|
||||
payload: Readonly<Payload>;
|
||||
}>;
|
||||
|
||||
export type LLMStartEvent = BaseEvent<{
|
||||
export type LLMStartEvent = {
|
||||
id: UUID;
|
||||
messages: ChatMessage[];
|
||||
}>;
|
||||
export type LLMToolCallEvent = BaseEvent<{
|
||||
};
|
||||
|
||||
export type LLMToolCallEvent = {
|
||||
toolCall: ToolCall;
|
||||
}>;
|
||||
export type LLMToolResultEvent = BaseEvent<{
|
||||
};
|
||||
|
||||
export type LLMToolResultEvent = {
|
||||
toolCall: ToolCall;
|
||||
toolResult: ToolOutput;
|
||||
}>;
|
||||
export type LLMEndEvent = BaseEvent<{
|
||||
};
|
||||
|
||||
export type LLMEndEvent = {
|
||||
id: UUID;
|
||||
response: ChatResponse;
|
||||
}>;
|
||||
export type LLMStreamEvent = BaseEvent<{
|
||||
};
|
||||
|
||||
export type LLMStreamEvent = {
|
||||
id: UUID;
|
||||
chunk: ChatResponseChunk;
|
||||
}>;
|
||||
};
|
||||
|
||||
export interface LlamaIndexEventMaps {
|
||||
"llm-start": LLMStartEvent;
|
||||
@@ -41,24 +42,32 @@ export interface LlamaIndexEventMaps {
|
||||
}
|
||||
|
||||
export class LlamaIndexCustomEvent<T = any> extends CustomEvent<T> {
|
||||
private constructor(event: string, options?: CustomEventInit) {
|
||||
reason: EventCaller | null = null;
|
||||
private constructor(
|
||||
event: string,
|
||||
options?: CustomEventInit & {
|
||||
reason?: EventCaller | null;
|
||||
},
|
||||
) {
|
||||
super(event, options);
|
||||
this.reason = options?.reason ?? null;
|
||||
}
|
||||
|
||||
static fromEvent<Type extends keyof LlamaIndexEventMaps>(
|
||||
type: Type,
|
||||
detail: LlamaIndexEventMaps[Type]["detail"],
|
||||
detail: LlamaIndexEventMaps[Type],
|
||||
) {
|
||||
return new LlamaIndexCustomEvent(type, {
|
||||
detail: detail,
|
||||
reason: getEventCaller(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
type EventHandler<Event> = (event: Event) => void;
|
||||
type EventHandler<Event> = (event: LlamaIndexCustomEvent<Event>) => void;
|
||||
|
||||
export class CallbackManager {
|
||||
#handlers = new Map<keyof LlamaIndexEventMaps, EventHandler<CustomEvent>[]>();
|
||||
#handlers = new Map<keyof LlamaIndexEventMaps, EventHandler<any>[]>();
|
||||
|
||||
on<K extends keyof LlamaIndexEventMaps>(
|
||||
event: K,
|
||||
@@ -88,7 +97,7 @@ export class CallbackManager {
|
||||
|
||||
dispatchEvent<K extends keyof LlamaIndexEventMaps>(
|
||||
event: K,
|
||||
detail: LlamaIndexEventMaps[K]["detail"],
|
||||
detail: LlamaIndexEventMaps[K],
|
||||
) {
|
||||
const cbs = this.#handlers.get(event);
|
||||
if (!cbs) {
|
||||
|
||||
@@ -22,10 +22,8 @@ export function wrapLLMEvent<
|
||||
> {
|
||||
const id = randomUUID();
|
||||
getCallbackManager().dispatchEvent("llm-start", {
|
||||
payload: {
|
||||
id,
|
||||
messages: params[0].messages,
|
||||
},
|
||||
id,
|
||||
messages: params[0].messages,
|
||||
});
|
||||
const response = await originalMethod.call(this, ...params);
|
||||
if (Symbol.asyncIterator in response) {
|
||||
@@ -58,29 +56,23 @@ export function wrapLLMEvent<
|
||||
};
|
||||
}
|
||||
getCallbackManager().dispatchEvent("llm-stream", {
|
||||
payload: {
|
||||
id,
|
||||
chunk,
|
||||
},
|
||||
id,
|
||||
chunk,
|
||||
});
|
||||
finalResponse.raw.push(chunk);
|
||||
yield chunk;
|
||||
}
|
||||
snapshot(() => {
|
||||
getCallbackManager().dispatchEvent("llm-end", {
|
||||
payload: {
|
||||
id,
|
||||
response: finalResponse,
|
||||
},
|
||||
id,
|
||||
response: finalResponse,
|
||||
});
|
||||
});
|
||||
};
|
||||
} else {
|
||||
getCallbackManager().dispatchEvent("llm-end", {
|
||||
payload: {
|
||||
id,
|
||||
response,
|
||||
},
|
||||
id,
|
||||
response,
|
||||
});
|
||||
}
|
||||
return response;
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
import { CallbackManager, Settings } from "@llamaindex/core/global";
|
||||
import { beforeEach, describe, expect, expectTypeOf, test, vi } from "vitest";
|
||||
|
||||
declare module "@llamaindex/core/global" {
|
||||
interface LlamaIndexEventMaps {
|
||||
test: {
|
||||
value: number;
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
describe("event system", () => {
|
||||
beforeEach(() => {
|
||||
Settings.callbackManager = new CallbackManager();
|
||||
});
|
||||
|
||||
test("type system", () => {
|
||||
Settings.callbackManager.on("test", (event) => {
|
||||
const data = event.detail;
|
||||
expectTypeOf(data).not.toBeAny();
|
||||
expectTypeOf(data).toEqualTypeOf<{
|
||||
value: number;
|
||||
}>();
|
||||
});
|
||||
});
|
||||
|
||||
test("dispatch event", async () => {
|
||||
let callback;
|
||||
Settings.callbackManager.on(
|
||||
"test",
|
||||
(callback = vi.fn((event) => {
|
||||
const data = event.detail;
|
||||
expect(data.value).toBe(42);
|
||||
})),
|
||||
);
|
||||
|
||||
Settings.callbackManager.dispatchEvent("test", {
|
||||
value: 42,
|
||||
});
|
||||
expect(callback).toHaveBeenCalledTimes(0);
|
||||
await new Promise((resolve) => process.nextTick(resolve));
|
||||
expect(callback).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
// rollup doesn't support decorators for now
|
||||
// test('wrap event caller', async () => {
|
||||
// class A {
|
||||
// @wrapEventCaller
|
||||
// fn() {
|
||||
// Settings.callbackManager.dispatchEvent('test', {
|
||||
// value: 42
|
||||
// });
|
||||
// }
|
||||
// }
|
||||
// const a = new A();
|
||||
// let callback;
|
||||
// Settings.callbackManager.on('test', callback = vi.fn((event) => {
|
||||
// const data = event.detail;
|
||||
// expect(event.reason!.caller).toBe(a);
|
||||
// expect(data.value).toBe(42);
|
||||
// }));
|
||||
// a.fn();
|
||||
// expect(callback).toHaveBeenCalledTimes(0);
|
||||
// await new Promise((resolve) => process.nextTick(resolve));
|
||||
// expect(callback).toHaveBeenCalledTimes(1);
|
||||
// })
|
||||
});
|
||||
@@ -19,10 +19,10 @@ Settings.embedModel = new HuggingFaceEmbedding({
|
||||
quantized: false,
|
||||
});
|
||||
Settings.callbackManager.on("llm-tool-call", (event) => {
|
||||
console.log(event.detail.payload);
|
||||
console.log(event.detail);
|
||||
});
|
||||
Settings.callbackManager.on("llm-tool-result", (event) => {
|
||||
console.log(event.detail.payload);
|
||||
console.log(event.detail);
|
||||
});
|
||||
|
||||
export async function getOpenAIModelRequest(query: string) {
|
||||
|
||||
@@ -5,15 +5,16 @@ import {
|
||||
type LLMStartEvent,
|
||||
type LLMStreamEvent,
|
||||
} from "@llamaindex/core/global";
|
||||
import { CustomEvent } from "@llamaindex/env";
|
||||
import { readFile, writeFile } from "node:fs/promises";
|
||||
import { join } from "node:path";
|
||||
import { type test } from "node:test";
|
||||
import { fileURLToPath } from "node:url";
|
||||
|
||||
type MockStorage = {
|
||||
llmEventStart: LLMStartEvent["detail"]["payload"][];
|
||||
llmEventEnd: LLMEndEvent["detail"]["payload"][];
|
||||
llmEventStream: LLMStreamEvent["detail"]["payload"][];
|
||||
llmEventStart: LLMStartEvent[];
|
||||
llmEventEnd: LLMEndEvent[];
|
||||
llmEventStream: LLMStreamEvent[];
|
||||
};
|
||||
|
||||
export const llmCompleteMockStorage: MockStorage = {
|
||||
@@ -36,35 +37,35 @@ export async function mockLLMEvent(
|
||||
llmEventStream: [],
|
||||
};
|
||||
|
||||
function captureLLMStart(event: LLMStartEvent) {
|
||||
idMap.set(event.detail.payload.id, `PRESERVE_${counter++}`);
|
||||
function captureLLMStart(event: CustomEvent<LLMStartEvent>) {
|
||||
idMap.set(event.detail.id, `PRESERVE_${counter++}`);
|
||||
newLLMCompleteMockStorage.llmEventStart.push({
|
||||
...event.detail.payload,
|
||||
...event.detail,
|
||||
// @ts-expect-error id is not UUID, but it is fine for testing
|
||||
id: idMap.get(event.detail.payload.id)!,
|
||||
});
|
||||
}
|
||||
|
||||
function captureLLMEnd(event: LLMEndEvent) {
|
||||
function captureLLMEnd(event: CustomEvent<LLMEndEvent>) {
|
||||
newLLMCompleteMockStorage.llmEventEnd.push({
|
||||
...event.detail.payload,
|
||||
...event.detail,
|
||||
// @ts-expect-error id is not UUID, but it is fine for testing
|
||||
id: idMap.get(event.detail.payload.id)!,
|
||||
response: {
|
||||
...event.detail.payload.response,
|
||||
...event.detail.response,
|
||||
// hide raw object since it might too big
|
||||
raw: null,
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
function captureLLMStream(event: LLMStreamEvent) {
|
||||
function captureLLMStream(event: CustomEvent<LLMStreamEvent>) {
|
||||
newLLMCompleteMockStorage.llmEventStream.push({
|
||||
...event.detail.payload,
|
||||
...event.detail,
|
||||
// @ts-expect-error id is not UUID, but it is fine for testing
|
||||
id: idMap.get(event.detail.payload.id)!,
|
||||
chunk: {
|
||||
...event.detail.payload.chunk,
|
||||
...event.detail.chunk,
|
||||
// hide raw object since it might too big
|
||||
raw: null,
|
||||
},
|
||||
|
||||
@@ -69,9 +69,7 @@ export function createTaskOutputStream<
|
||||
controller.enqueue(output);
|
||||
};
|
||||
Settings.callbackManager.dispatchEvent("agent-start", {
|
||||
payload: {
|
||||
startStep: step,
|
||||
},
|
||||
startStep: step,
|
||||
});
|
||||
|
||||
context.logger.log("Starting step(id, %s).", step.id);
|
||||
@@ -93,9 +91,7 @@ export function createTaskOutputStream<
|
||||
step.id,
|
||||
);
|
||||
Settings.callbackManager.dispatchEvent("agent-end", {
|
||||
payload: {
|
||||
endStep: step,
|
||||
},
|
||||
endStep: step,
|
||||
});
|
||||
controller.close();
|
||||
}
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import type { BaseEvent } from "@llamaindex/core/global";
|
||||
import type {
|
||||
BaseToolWithCall,
|
||||
ChatMessage,
|
||||
@@ -90,9 +89,9 @@ export type TaskHandler<
|
||||
) => void,
|
||||
) => Promise<void>;
|
||||
|
||||
export type AgentStartEvent = BaseEvent<{
|
||||
export type AgentStartEvent = {
|
||||
startStep: TaskStep;
|
||||
}>;
|
||||
export type AgentEndEvent = BaseEvent<{
|
||||
};
|
||||
export type AgentEndEvent = {
|
||||
endStep: TaskStep;
|
||||
}>;
|
||||
};
|
||||
|
||||
@@ -225,9 +225,7 @@ export async function callTool(
|
||||
}
|
||||
try {
|
||||
Settings.callbackManager.dispatchEvent("llm-tool-call", {
|
||||
payload: {
|
||||
toolCall: { ...toolCall, input },
|
||||
},
|
||||
toolCall: { ...toolCall, input },
|
||||
});
|
||||
output = await call.call(tool, input);
|
||||
logger.log(
|
||||
@@ -241,10 +239,8 @@ export async function callTool(
|
||||
isError: false,
|
||||
};
|
||||
Settings.callbackManager.dispatchEvent("llm-tool-result", {
|
||||
payload: {
|
||||
toolCall: { ...toolCall, input },
|
||||
toolResult: { ...toolOutput },
|
||||
},
|
||||
toolCall: { ...toolCall, input },
|
||||
toolResult: { ...toolOutput },
|
||||
});
|
||||
return toolOutput;
|
||||
} catch (e) {
|
||||
|
||||
@@ -92,10 +92,8 @@ export class LlamaCloudRetriever implements BaseRetriever {
|
||||
results.retrieval_nodes,
|
||||
);
|
||||
Settings.callbackManager.dispatchEvent("retrieve-end", {
|
||||
payload: {
|
||||
query,
|
||||
nodes: nodesWithScores,
|
||||
},
|
||||
query,
|
||||
nodes: nodesWithScores,
|
||||
});
|
||||
return nodesWithScores;
|
||||
}
|
||||
|
||||
@@ -13,7 +13,6 @@ declare module "@llamaindex/core/global" {
|
||||
|
||||
export { CallbackManager } from "@llamaindex/core/global";
|
||||
export type {
|
||||
BaseEvent,
|
||||
JSONArray,
|
||||
JSONObject,
|
||||
JSONValue,
|
||||
|
||||
@@ -414,9 +414,7 @@ export class VectorIndexRetriever implements BaseRetriever {
|
||||
preFilters,
|
||||
}: RetrieveParams): Promise<NodeWithScore[]> {
|
||||
Settings.callbackManager.dispatchEvent("retrieve-start", {
|
||||
payload: {
|
||||
query,
|
||||
},
|
||||
query,
|
||||
});
|
||||
const vectorStores = this.index.vectorStores;
|
||||
let nodesWithScores: NodeWithScore[] = [];
|
||||
@@ -433,10 +431,8 @@ export class VectorIndexRetriever implements BaseRetriever {
|
||||
);
|
||||
}
|
||||
Settings.callbackManager.dispatchEvent("retrieve-end", {
|
||||
payload: {
|
||||
query,
|
||||
nodes: nodesWithScores,
|
||||
},
|
||||
query,
|
||||
nodes: nodesWithScores,
|
||||
});
|
||||
return nodesWithScores;
|
||||
}
|
||||
|
||||
@@ -1,11 +1,10 @@
|
||||
import type { BaseEvent } from "@llamaindex/core/global";
|
||||
import type { MessageContent } from "@llamaindex/core/llms";
|
||||
import type { NodeWithScore } from "@llamaindex/core/schema";
|
||||
|
||||
export type RetrievalStartEvent = BaseEvent<{
|
||||
export type RetrievalStartEvent = {
|
||||
query: MessageContent;
|
||||
}>;
|
||||
export type RetrievalEndEvent = BaseEvent<{
|
||||
};
|
||||
export type RetrievalEndEvent = {
|
||||
query: MessageContent;
|
||||
nodes: NodeWithScore[];
|
||||
}>;
|
||||
};
|
||||
|
||||
Reference in New Issue
Block a user