refactor(core): move ContextChatEngine and SimpleChatEngine (#1401)

This commit is contained in:
Alex Yang
2024-10-27 00:39:03 -05:00
committed by GitHub
parent efb7e1b868
commit 359fd33041
10 changed files with 92 additions and 85 deletions
+6
View File
@@ -0,0 +1,6 @@
---
"@llamaindex/core": patch
"llamaindex": patch
---
refactor(core): move `ContextChatEngine` and `SimpleChatEngine`
+6 -4
View File
@@ -2,9 +2,9 @@ import { ClientMDXContent } from "@/components/mdx";
import { BotMessage } from "@/components/message";
import { Skeleton } from "@/components/ui/skeleton";
import { LlamaCloudRetriever } from "@/deps/cloud";
import { ContextChatEngine } from "@llamaindex/core/chat-engine";
import { Settings } from "@llamaindex/core/global";
import { ChatMessage } from "@llamaindex/core/llms";
import { RetrieverQueryEngine } from "@llamaindex/core/query-engine";
import { OpenAI } from "@llamaindex/openai";
import { createAI, createStreamableUI, getMutableAIState } from "ai/rsc";
import { ReactNode } from "react";
@@ -50,7 +50,7 @@ export const AIProvider = createAI({
actions: {
query: async (message: string): Promise<UIMessage> => {
"use server";
const queryEngine = new RetrieverQueryEngine(retriever);
const chatEngine = new ContextChatEngine({ retriever });
const id = Date.now();
const aiState = getMutableAIState<typeof AIProvider>();
@@ -73,10 +73,12 @@ export const AIProvider = createAI({
);
runAsyncFnWithoutBlocking(async () => {
const response = await queryEngine.query({
query: message,
const response = await chatEngine.chat({
message,
chatHistory: aiState.get().messages,
stream: true,
});
let content = "";
for await (const { delta } of response) {
+1 -1
View File
@@ -3,7 +3,7 @@ import {
BaseChatEngine,
type NonStreamingChatEngineParams,
type StreamingChatEngineParams,
} from "../chat-engine";
} from "../chat-engine/base";
import { wrapEventCaller } from "../decorator";
import { Settings } from "../global";
import type {
+36
View File
@@ -0,0 +1,36 @@
import type { ChatMessage, MessageContent } from "../llms";
import type { BaseMemory } from "../memory";
import { EngineResponse } from "../schema";
export interface BaseChatEngineParams<
AdditionalMessageOptions extends object = object,
> {
message: MessageContent;
/**
* Optional chat history if you want to customize the chat history.
*/
chatHistory?:
| ChatMessage<AdditionalMessageOptions>[]
| BaseMemory<AdditionalMessageOptions>;
}
export interface StreamingChatEngineParams<
AdditionalMessageOptions extends object = object,
> extends BaseChatEngineParams<AdditionalMessageOptions> {
stream: true;
}
export interface NonStreamingChatEngineParams<
AdditionalMessageOptions extends object = object,
> extends BaseChatEngineParams<AdditionalMessageOptions> {
stream?: false;
}
export abstract class BaseChatEngine {
abstract chat(params: NonStreamingChatEngineParams): Promise<EngineResponse>;
abstract chat(
params: StreamingChatEngineParams,
): Promise<AsyncIterable<EngineResponse>>;
abstract chatHistory: ChatMessage[] | Promise<ChatMessage[]>;
}
@@ -1,33 +1,24 @@
import type {
BaseChatEngine,
NonStreamingChatEngineParams,
StreamingChatEngineParams,
} from "@llamaindex/core/chat-engine";
import { wrapEventCaller } from "@llamaindex/core/decorator";
import type {
ChatMessage,
LLM,
MessageContent,
MessageType,
} from "@llamaindex/core/llms";
import { BaseMemory, ChatMemoryBuffer } from "@llamaindex/core/memory";
import type { BaseNodePostprocessor } from "@llamaindex/core/postprocessor";
import { wrapEventCaller } from "../decorator";
import { Settings } from "../global";
import type { ChatMessage, LLM, MessageContent, MessageType } from "../llms";
import { BaseMemory, ChatMemoryBuffer } from "../memory";
import type { BaseNodePostprocessor } from "../postprocessor";
import {
type ContextSystemPrompt,
type ModuleRecord,
PromptMixin,
type PromptsRecord,
} from "@llamaindex/core/prompts";
import type { BaseRetriever } from "@llamaindex/core/retriever";
import { EngineResponse, MetadataMode } from "@llamaindex/core/schema";
import {
extractText,
streamConverter,
streamReducer,
} from "@llamaindex/core/utils";
import { Settings } from "../../Settings.js";
import { DefaultContextGenerator } from "./DefaultContextGenerator.js";
import type { ContextGenerator } from "./types.js";
} from "../prompts";
import type { BaseRetriever } from "../retriever";
import { EngineResponse, MetadataMode } from "../schema";
import { extractText, streamConverter, streamReducer } from "../utils";
import type {
BaseChatEngine,
NonStreamingChatEngineParams,
StreamingChatEngineParams,
} from "./base";
import { DefaultContextGenerator } from "./default-context-generator";
import type { ContextGenerator } from "./type";
/**
* ContextChatEngine uses the Index to get the appropriate context for each query.
@@ -1,15 +1,15 @@
import type { MessageContent, MessageType } from "@llamaindex/core/llms";
import type { BaseNodePostprocessor } from "@llamaindex/core/postprocessor";
import type { MessageContent, MessageType } from "../llms";
import type { BaseNodePostprocessor } from "../postprocessor";
import {
type ContextSystemPrompt,
defaultContextSystemPrompt,
type ModuleRecord,
PromptMixin,
} from "@llamaindex/core/prompts";
import { createMessageContent } from "@llamaindex/core/response-synthesizers";
import type { BaseRetriever } from "@llamaindex/core/retriever";
import { MetadataMode, type NodeWithScore } from "@llamaindex/core/schema";
import type { Context, ContextGenerator } from "./types.js";
} from "../prompts";
import { createMessageContent } from "../response-synthesizers";
import type { BaseRetriever } from "../retriever";
import { MetadataMode, type NodeWithScore } from "../schema";
import type { Context, ContextGenerator } from "./type.js";
export class DefaultContextGenerator
extends PromptMixin
+9 -36
View File
@@ -1,36 +1,9 @@
import type { ChatMessage, MessageContent } from "../llms";
import type { BaseMemory } from "../memory";
import { EngineResponse } from "../schema";
export interface BaseChatEngineParams<
AdditionalMessageOptions extends object = object,
> {
message: MessageContent;
/**
* Optional chat history if you want to customize the chat history.
*/
chatHistory?:
| ChatMessage<AdditionalMessageOptions>[]
| BaseMemory<AdditionalMessageOptions>;
}
export interface StreamingChatEngineParams<
AdditionalMessageOptions extends object = object,
> extends BaseChatEngineParams<AdditionalMessageOptions> {
stream: true;
}
export interface NonStreamingChatEngineParams<
AdditionalMessageOptions extends object = object,
> extends BaseChatEngineParams<AdditionalMessageOptions> {
stream?: false;
}
export abstract class BaseChatEngine {
abstract chat(params: NonStreamingChatEngineParams): Promise<EngineResponse>;
abstract chat(
params: StreamingChatEngineParams,
): Promise<AsyncIterable<EngineResponse>>;
abstract chatHistory: ChatMessage[] | Promise<ChatMessage[]>;
}
export {
BaseChatEngine,
type BaseChatEngineParams,
type NonStreamingChatEngineParams,
type StreamingChatEngineParams,
} from "./base";
export { ContextChatEngine } from "./context-chat-engine";
export { DefaultContextGenerator } from "./default-context-generator";
export { SimpleChatEngine } from "./simple-chat-engine";
@@ -1,15 +1,15 @@
import type { LLM } from "../llms";
import { BaseMemory, ChatMemoryBuffer } from "../memory";
import { EngineResponse } from "../schema";
import { streamConverter, streamReducer } from "../utils";
import type {
BaseChatEngine,
NonStreamingChatEngineParams,
StreamingChatEngineParams,
} from "@llamaindex/core/chat-engine";
import type { LLM } from "@llamaindex/core/llms";
import { BaseMemory, ChatMemoryBuffer } from "@llamaindex/core/memory";
import { EngineResponse } from "@llamaindex/core/schema";
import { streamConverter, streamReducer } from "@llamaindex/core/utils";
} from "./base";
import { wrapEventCaller } from "@llamaindex/core/decorator";
import { Settings } from "../../Settings.js";
import { wrapEventCaller } from "../decorator";
import { Settings } from "../global";
/**
* SimpleChatEngine is the simplest possible chat engine. Useful for using your own custom prompts.
@@ -1,10 +1,11 @@
import type { ChatMessage } from "@llamaindex/core/llms";
import type { NodeWithScore } from "@llamaindex/core/schema";
import type { ChatMessage } from "../llms";
import type { NodeWithScore } from "../schema";
export interface Context {
message: ChatMessage;
nodes: NodeWithScore[];
}
/**
* A ContextGenerator is used to generate a context based on a message's text content
*/
@@ -1,4 +1,2 @@
export * from "@llamaindex/core/chat-engine";
export { CondenseQuestionChatEngine } from "./CondenseQuestionChatEngine.js";
export { ContextChatEngine } from "./ContextChatEngine.js";
export { SimpleChatEngine } from "./SimpleChatEngine.js";
export * from "./types.js";