Compare commits

...

33 Commits

Author SHA1 Message Date
github-actions[bot] de2c7523dd Release 0.1.37 (#239)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
2024-08-15 14:52:27 +07:00
Huu Le 9fd832c8b0 feat: In-text citing (#175)
---------
Co-authored-by: Marcus Schiesser <mail@marcusschiesser.de>
2024-08-15 13:52:51 +07:00
github-actions[bot] b2c76dc7b6 Release 0.1.36 (#238)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
2024-08-15 11:02:00 +07:00
Thuc Pham 2b7a5d8797 fix: optional params in file upload API (#237)
---------
Co-authored-by: Marcus Schiesser <mail@marcusschiesser.de>
2024-08-15 11:00:53 +07:00
Marcus Schiesser d93ec803f5 feat: add ruff (#235)
* fix: formatting

* fix: ruff --fix

* feat: add ruff to github action

* fix: remove E402 check for some files
2024-08-15 09:38:13 +07:00
github-actions[bot] a6023b695b Release 0.1.35 (#234)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
2024-08-14 17:22:49 +07:00
Marcus Schiesser 81ef7f0f93 feat: use llamacloud pipeline for private files and generate script in Python (#226)
---------
Co-authored-by: Thuc Pham <51660321+thucpn@users.noreply.github.com>
2024-08-14 17:03:16 +07:00
github-actions[bot] 8faf9170cf Release 0.1.34 (#233)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
2024-08-14 14:59:19 +07:00
Huu Le c49a5e1620 chore: update wrong env name, add error handling for next question (#232) 2024-08-14 14:39:14 +07:00
github-actions[bot] 8b2de431f2 Release 0.1.33 (#229)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
2024-08-13 11:06:17 +07:00
Huu Le d746c75e49 feat: Add Weaviate vector store for Typescript templates (#228) 2024-08-13 10:56:02 +07:00
Laurie Voss c87978ab96 Point the repo to the current one (#227) 2024-08-13 10:51:04 +07:00
github-actions[bot] 26359a0ac9 Release 0.1.32 (#224)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
2024-08-12 17:19:22 +07:00
Huu Le 4039d3d1ea refactor: include chat configuration router in FastAPI app (#225) 2024-08-12 17:17:22 +07:00
Huu Le 3ec5163304 feat: add Weaviate vector database support for Python (#223)
---------
Co-authored-by: Marcus Schiesser <mail@marcusschiesser.de>
2024-08-12 16:25:26 +07:00
Thuc Pham 878cfc2ca1 refactor: make llamacloud selector resuable (#221) 2024-08-09 12:02:43 +07:00
github-actions[bot] 9b5835b71c Release 0.1.31 (#222)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
2024-08-09 11:55:58 +07:00
Thuc Pham 04a9c71759 feat: cluster nodes in document (#217)
---------
Co-authored-by: Marcus Schiesser <mail@marcusschiesser.de>
2024-08-09 11:54:50 +07:00
github-actions[bot] 0bfdbc1dfe Release 0.1.30 (#214)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
2024-08-08 15:59:57 +07:00
Thuc Pham fbcaebcbcf fix: use modern module resolution for express (#219) 2024-08-08 09:51:13 +02:00
Thuc Pham b6dd7a9acb fix: always send chat data when submit message (#213) 2024-08-07 15:22:33 +02:00
Marcus Schiesser 09e3022ad6 feat: add LlamaTrace support (#216) 2024-08-07 15:21:44 +02:00
Marcus Schiesser 9f739b9834 refactor: cleaned e2e runner (#215) 2024-08-07 17:41:22 +07:00
Marcus Schiesser c06ec4f14c fix: imports for MongoDB 2024-08-07 11:01:00 +02:00
Marcus Schiesser e7d30b1c69 refactor: test frameworks and datasources via matrix (#211) 2024-08-05 23:50:00 +07:00
github-actions[bot] e974c8ef11 Release 0.1.29 (#210)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
2024-08-05 22:24:37 +07:00
Thuc Pham 8890e27a14 feat: implement index selector for LlamaCloud (#200)
---------
Co-authored-by: Marcus Schiesser <mail@marcusschiesser.de>
2024-08-05 22:18:20 +07:00
Marcus Schiesser 072e69b465 fix: deactive llamacloud tests 2024-08-05 13:49:09 +02:00
Huu Le 83a648df0a chore: add use window.ENV.BASE_URL as backendOrigin (#205)
---------
Co-authored-by: Marcus Schiesser <mail@marcusschiesser.de>
2024-08-02 15:50:04 +07:00
github-actions[bot] dcf52abdba Release 0.1.28 (#206)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
2024-08-02 15:38:27 +07:00
Marcus Schiesser 9a09e8c7e2 fix: Vercel deployment (by including WASM files) (#201) 2024-08-02 15:36:54 +07:00
github-actions[bot] a4a55239e9 Release 0.1.27 (#204)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
2024-08-01 17:09:07 +02:00
Thuc Pham c5c7eee04d refactor: make components resuable for chat llm (#202)
---------
Co-authored-by: Marcus Schiesser <mail@marcusschiesser.de>
2024-08-01 16:31:38 +02:00
107 changed files with 2078 additions and 875 deletions
+4
View File
@@ -18,6 +18,8 @@ jobs:
node-version: [18, 20]
python-version: ["3.11"]
os: [macos-latest, windows-latest, ubuntu-22.04]
frameworks: ["nextjs", "express", "fastapi"]
datasources: ["--no-files", "--example-file"]
defaults:
run:
shell: bash
@@ -63,6 +65,8 @@ jobs:
env:
OPENAI_API_KEY: ${{ secrets.OPENAI_API_KEY }}
LLAMA_CLOUD_API_KEY: ${{ secrets.LLAMA_CLOUD_API_KEY }}
FRAMEWORK: ${{ matrix.frameworks }}
DATASOURCE: ${{ matrix.datasources }}
working-directory: .
- uses: actions/upload-artifact@v3
@@ -30,3 +30,13 @@ jobs:
- name: Run Prettier
run: pnpm run format
- name: Run Python format check
uses: chartboost/ruff-action@v1
with:
args: "format --check"
- name: Run Python lint
uses: chartboost/ruff-action@v1
with:
args: "check"
+69
View File
@@ -1,5 +1,74 @@
# create-llama
## 0.1.37
### Patch Changes
- 9fd832c: Add in-text citation references
## 0.1.36
### Patch Changes
- 2b7a5d8: Fix: private file upload not working in Python without LlamaCloud
## 0.1.35
### Patch Changes
- 81ef7f0: Use LlamaCloud pipeline for data ingestion (private file uploads and generate script)
## 0.1.34
### Patch Changes
- c49a5e1: Add error handling for generating the next question
- c49a5e1: Fix wrong api key variable in Azure OpenAI provider
## 0.1.33
### Patch Changes
- d746c75: Add Weaviate vector store (Typescript)
## 0.1.32
### Patch Changes
- 3ec5163: Add Weaviate vector database support (Python)
## 0.1.31
### Patch Changes
- 04a9c71: Cluster nodes by document
## 0.1.30
### Patch Changes
- 09e3022: Add support for LlamaTrace (Python)
- c06ec4f: Fix imports for MongoDB
- b6dd7a9: Always send chat data when submit message
## 0.1.29
### Patch Changes
- 8890e27: Let user change indexes in LlamaCloud projects
## 0.1.28
### Patch Changes
- 9a09e8c: Fix Vercel deployment
## 0.1.27
### Patch Changes
- c5c7eee: Make components reusable for chat-llamaindex
## 0.1.26
### Patch Changes
+23 -9
View File
@@ -9,7 +9,7 @@ import { makeDir } from "./helpers/make-dir";
import fs from "fs";
import terminalLink from "terminal-link";
import type { InstallTemplateArgs } from "./helpers";
import type { InstallTemplateArgs, TemplateObservability } from "./helpers";
import { installTemplate } from "./helpers";
import { writeDevcontainer } from "./helpers/devcontainer";
import { templatesDir } from "./helpers/dir";
@@ -142,14 +142,7 @@ export async function createApp({
)} and learn how to get started.`,
);
if (args.observability === "opentelemetry") {
console.log(
`\n${yellow("Observability")}: Visit the ${terminalLink(
"documentation",
"https://traceloop.com/docs/openllmetry/integrations",
)} to set up the environment variables and start seeing execution traces.`,
);
}
outputObservability(args.observability);
if (
dataSources.some((dataSource) => dataSource.type === "file") &&
@@ -167,3 +160,24 @@ export async function createApp({
console.log();
}
function outputObservability(observability?: TemplateObservability) {
switch (observability) {
case "traceloop":
console.log(
`\n${yellow("Observability")}: Visit the ${terminalLink(
"documentation",
"https://traceloop.com/docs/openllmetry/integrations",
)} to set up the environment variables and start seeing execution traces.`,
);
break;
case "llamatrace":
console.log(
`\n${yellow("Observability")}: LlamaTrace has been configured for your project. Visit the ${terminalLink(
"LlamaTrace dashboard",
"https://llamatrace.com/login",
)} to view your traces and monitor your application.`,
);
break;
}
}
+98 -116
View File
@@ -11,128 +11,110 @@ import type {
} from "../helpers";
import { createTestDir, runCreateLlama, type AppType } from "./utils";
const templateTypes: TemplateType[] = ["streaming"];
const templateFrameworks: TemplateFramework[] = [
"nextjs",
"express",
"fastapi",
];
const dataSources: string[] = ["--no-files", "--llamacloud"];
const templateUIs: TemplateUI[] = ["shadcn", "html"];
const templatePostInstallActions: TemplatePostInstallAction[] = [
"none",
"runApp",
];
const templateType: TemplateType = "streaming";
const templateFramework: TemplateFramework = process.env.FRAMEWORK
? (process.env.FRAMEWORK as TemplateFramework)
: "fastapi";
const dataSource: string = process.env.DATASOURCE
? process.env.DATASOURCE
: "--example-file";
const templateUI: TemplateUI = "shadcn";
const templatePostInstallAction: TemplatePostInstallAction = "runApp";
const llamaCloudProjectName = "create-llama";
const llamaCloudIndexName = "e2e-test";
for (const templateType of templateTypes) {
for (const templateFramework of templateFrameworks) {
for (const dataSource of dataSources) {
for (const templateUI of templateUIs) {
for (const templatePostInstallAction of templatePostInstallActions) {
const appType: AppType =
templateFramework === "nextjs" ? "" : "--frontend";
const userMessage =
dataSource !== "--no-files"
? "Physical standard for letters"
: "Hello";
test.describe(`try create-llama ${templateType} ${templateFramework} ${dataSource} ${templateUI} ${appType} ${templatePostInstallAction}`, async () => {
let port: number;
let externalPort: number;
let cwd: string;
let name: string;
let appProcess: ChildProcess;
// Only test without using vector db for now
const vectorDb = "none";
const appType: AppType = templateFramework === "nextjs" ? "" : "--frontend";
const userMessage =
dataSource !== "--no-files" ? "Physical standard for letters" : "Hello";
test.describe(`try create-llama ${templateType} ${templateFramework} ${dataSource} ${templateUI} ${appType} ${templatePostInstallAction}`, async () => {
let port: number;
let externalPort: number;
let cwd: string;
let name: string;
let appProcess: ChildProcess;
// Only test without using vector db for now
const vectorDb = "none";
test.beforeAll(async () => {
port = Math.floor(Math.random() * 10000) + 10000;
externalPort = port + 1;
cwd = await createTestDir();
const result = await runCreateLlama(
cwd,
templateType,
templateFramework,
dataSource,
templateUI,
vectorDb,
appType,
port,
externalPort,
templatePostInstallAction,
llamaCloudProjectName,
llamaCloudIndexName,
);
name = result.projectName;
appProcess = result.appProcess;
});
test.beforeAll(async () => {
port = Math.floor(Math.random() * 10000) + 10000;
externalPort = port + 1;
cwd = await createTestDir();
const result = await runCreateLlama(
cwd,
templateType,
templateFramework,
dataSource,
templateUI,
vectorDb,
appType,
port,
externalPort,
templatePostInstallAction,
llamaCloudProjectName,
llamaCloudIndexName,
);
name = result.projectName;
appProcess = result.appProcess;
});
test("App folder should exist", async () => {
const dirExists = fs.existsSync(path.join(cwd, name));
expect(dirExists).toBeTruthy();
});
test("Frontend should have a title", async ({ page }) => {
test.skip(templatePostInstallAction !== "runApp");
await page.goto(`http://localhost:${port}`);
await expect(page.getByText("Built by LlamaIndex")).toBeVisible();
});
test("App folder should exist", async () => {
const dirExists = fs.existsSync(path.join(cwd, name));
expect(dirExists).toBeTruthy();
});
test("Frontend should have a title", async ({ page }) => {
test.skip(templatePostInstallAction !== "runApp");
await page.goto(`http://localhost:${port}`);
await expect(page.getByText("Built by LlamaIndex")).toBeVisible();
});
test("Frontend should be able to submit a message and receive a response", async ({
page,
}) => {
test.skip(templatePostInstallAction !== "runApp");
await page.goto(`http://localhost:${port}`);
await page.fill("form input", userMessage);
const [response] = await Promise.all([
page.waitForResponse(
(res) => {
return (
res.url().includes("/api/chat") && res.status() === 200
);
},
{
timeout: 1000 * 60,
},
),
page.click("form button[type=submit]"),
]);
const text = await response.text();
console.log("AI response when submitting message: ", text);
expect(response.ok()).toBeTruthy();
});
test("Frontend should be able to submit a message and receive a response", async ({
page,
}) => {
test.skip(templatePostInstallAction !== "runApp");
await page.goto(`http://localhost:${port}`);
await page.fill("form input", userMessage);
const [response] = await Promise.all([
page.waitForResponse(
(res) => {
return res.url().includes("/api/chat") && res.status() === 200;
},
{
timeout: 1000 * 60,
},
),
page.click("form button[type=submit]"),
]);
const text = await response.text();
console.log("AI response when submitting message: ", text);
expect(response.ok()).toBeTruthy();
});
test("Backend frameworks should response when calling non-streaming chat API", async ({
request,
}) => {
test.skip(templatePostInstallAction !== "runApp");
test.skip(templateFramework === "nextjs");
const response = await request.post(
`http://localhost:${externalPort}/api/chat/request`,
{
data: {
messages: [
{
role: "user",
content: userMessage,
},
],
},
},
);
const text = await response.text();
console.log("AI response when calling API: ", text);
expect(response.ok()).toBeTruthy();
});
test("Backend frameworks should response when calling non-streaming chat API", async ({
request,
}) => {
test.skip(templatePostInstallAction !== "runApp");
test.skip(templateFramework === "nextjs");
const response = await request.post(
`http://localhost:${externalPort}/api/chat/request`,
{
data: {
messages: [
{
role: "user",
content: userMessage,
},
],
},
},
);
const text = await response.text();
console.log("AI response when calling API: ", text);
expect(response.ok()).toBeTruthy();
});
// clean processes
test.afterAll(async () => {
appProcess?.kill();
});
});
}
}
}
}
}
// clean processes
test.afterAll(async () => {
appProcess?.kill();
});
});
+52 -59
View File
@@ -18,48 +18,6 @@ export type CreateLlamaResult = {
appProcess: ChildProcess;
};
// eslint-disable-next-line max-params
export async function checkAppHasStarted(
frontend: boolean,
framework: TemplateFramework,
port: number,
externalPort: number,
timeout: number,
) {
if (frontend) {
await Promise.all([
waitPort({
host: "localhost",
port: port,
timeout,
}),
waitPort({
host: "localhost",
port: externalPort,
timeout,
}),
]).catch((err) => {
console.error(err);
throw err;
});
} else {
let wPort: number;
if (framework === "nextjs") {
wPort = port;
} else {
wPort = externalPort;
}
await waitPort({
host: "localhost",
port: wPort,
timeout,
}).catch((err) => {
console.error(err);
throw err;
});
}
}
// eslint-disable-next-line max-params
export async function runCreateLlama(
cwd: string,
@@ -142,25 +100,10 @@ export async function runCreateLlama(
templateFramework,
port,
externalPort,
1000 * 60 * 5,
);
} else {
// wait create-llama to exit
// we don't test install dependencies for now, so just set timeout for 10 seconds
await new Promise((resolve, reject) => {
const timeout = setTimeout(() => {
reject(new Error("create-llama timeout error"));
}, 1000 * 10);
appProcess.on("exit", (code) => {
if (code !== 0 && code !== null) {
clearTimeout(timeout);
reject(new Error("create-llama command was failed!"));
} else {
clearTimeout(timeout);
resolve(undefined);
}
});
});
// wait 10 seconds for create-llama to exit
await waitForProcess(appProcess, 1000 * 10);
}
return {
@@ -174,3 +117,53 @@ export async function createTestDir() {
await mkdir(cwd, { recursive: true });
return cwd;
}
// eslint-disable-next-line max-params
async function checkAppHasStarted(
frontend: boolean,
framework: TemplateFramework,
port: number,
externalPort: number,
) {
const portsToWait = frontend
? [port, externalPort]
: [framework === "nextjs" ? port : externalPort];
await waitPorts(portsToWait);
}
async function waitPorts(ports: number[]): Promise<void> {
const waitForPort = async (port: number): Promise<void> => {
await waitPort({
host: "localhost",
port: port,
// wait max. 5 mins for start up of app
timeout: 1000 * 60 * 5,
});
};
try {
await Promise.all(ports.map(waitForPort));
} catch (err) {
console.error(err);
throw err;
}
}
async function waitForProcess(
process: ChildProcess,
timeoutMs: number,
): Promise<void> {
return new Promise((resolve, reject) => {
const timeout = setTimeout(() => {
reject(new Error("Process timeout error"));
}, timeoutMs);
process.on("exit", (code) => {
clearTimeout(timeout);
if (code !== 0 && code !== null) {
reject(new Error("Process exited with non-zero code"));
} else {
resolve();
}
});
});
}
+116 -23
View File
@@ -2,9 +2,11 @@ import fs from "fs/promises";
import path from "path";
import { TOOL_SYSTEM_PROMPT_ENV_VAR, Tool } from "./tools";
import {
InstallTemplateArgs,
ModelConfig,
TemplateDataSource,
TemplateFramework,
TemplateObservability,
TemplateType,
TemplateVectorDB,
} from "./types";
@@ -160,6 +162,17 @@ const getVectorDBEnvs = (
description:
"The organization ID for the LlamaCloud project (uses default organization if not specified - Python only)",
},
...(framework === "nextjs"
? // activate index selector per default (not needed for non-NextJS backends as it's handled by createFrontendEnvFile)
[
{
name: "NEXT_PUBLIC_USE_LLAMACLOUD",
description:
"Let's the user change indexes in LlamaCloud projects",
value: "true",
},
]
: []),
];
case "chroma":
const envs = [
@@ -186,6 +199,23 @@ Otherwise, use CHROMA_HOST and CHROMA_PORT config above`,
});
}
return envs;
case "weaviate":
return [
{
name: "WEAVIATE_CLUSTER_URL",
description:
"The URL of the Weaviate cloud cluster, see: https://weaviate.io/developers/wcs/connect",
},
{
name: "WEAVIATE_API_KEY",
description: "The API key for the Weaviate cloud cluster",
},
{
name: "WEAVIATE_INDEX_NAME",
description:
"(Optional) The collection name to use, default is LlamaIndex if not specified",
},
];
default:
return [];
}
@@ -282,7 +312,7 @@ const getModelEnvs = (modelConfig: ModelConfig): EnvVar[] => {
...(modelConfig.provider === "azure-openai"
? [
{
name: "AZURE_OPENAI_KEY",
name: "AZURE_OPENAI_API_KEY",
description: "The Azure OpenAI key to use.",
value: modelConfig.apiKey,
},
@@ -394,7 +424,11 @@ const getToolEnvs = (tools?: Tool[]): EnvVar[] => {
return toolEnvs;
};
const getSystemPromptEnv = (tools?: Tool[]): EnvVar => {
const getSystemPromptEnv = (
tools?: Tool[],
dataSources?: TemplateDataSource[],
framework?: TemplateFramework,
): EnvVar[] => {
const defaultSystemPrompt =
"You are a helpful assistant who helps users with their questions.";
@@ -413,11 +447,49 @@ const getSystemPromptEnv = (tools?: Tool[]): EnvVar => {
? `\"${toolSystemPrompt}\"`
: defaultSystemPrompt;
return {
name: "SYSTEM_PROMPT",
description: "The system prompt for the AI model.",
value: systemPrompt,
};
const systemPromptEnv = [
{
name: "SYSTEM_PROMPT",
description: "The system prompt for the AI model.",
value: systemPrompt,
},
];
// Citation only works with FastAPI along with the chat engine and data source provided for now.
if (
framework === "fastapi" &&
tools?.length == 0 &&
(dataSources?.length ?? 0 > 0)
) {
const citationPrompt = `'You have provided information from a knowledge base that has been passed to you in nodes of information.
Each node has useful metadata such as node ID, file name, page, etc.
Please add the citation to the data node for each sentence or paragraph that you reference in the provided information.
The citation format is: . [citation:<node_id>]()
Where the <node_id> is the unique identifier of the data node.
Example:
We have two nodes:
node_id: xyz
file_name: llama.pdf
node_id: abc
file_name: animal.pdf
User question: Tell me a fun fact about Llama.
Your answer:
A baby llama is called "Cria" [citation:xyz]().
It often live in desert [citation:abc]().
It\\'s cute animal.
'`;
systemPromptEnv.push({
name: "SYSTEM_CITATION_PROMPT",
description:
"An additional system prompt to add citation when responding to user questions.",
value: citationPrompt,
});
}
return systemPromptEnv;
};
const getTemplateEnvs = (template?: TemplateType): EnvVar[] => {
@@ -450,18 +522,35 @@ const getTemplateEnvs = (template?: TemplateType): EnvVar[] => {
}
};
const getObservabilityEnvs = (
observability?: TemplateObservability,
): EnvVar[] => {
if (observability === "llamatrace") {
return [
{
name: "PHOENIX_API_KEY",
description:
"API key for LlamaTrace observability. Retrieve from https://llamatrace.com/login",
},
];
}
return [];
};
export const createBackendEnvFile = async (
root: string,
opts: {
llamaCloudKey?: string;
vectorDb?: TemplateVectorDB;
modelConfig: ModelConfig;
framework: TemplateFramework;
dataSources?: TemplateDataSource[];
template?: TemplateType;
port?: number;
tools?: Tool[];
},
opts: Pick<
InstallTemplateArgs,
| "llamaCloudKey"
| "vectorDb"
| "modelConfig"
| "framework"
| "dataSources"
| "template"
| "externalPort"
| "tools"
| "observability"
>,
) => {
// Init env values
const envFileName = ".env";
@@ -471,17 +560,15 @@ export const createBackendEnvFile = async (
description: `The Llama Cloud API key.`,
value: opts.llamaCloudKey,
},
// Add model environment variables
// Add environment variables of each component
...getModelEnvs(opts.modelConfig),
// Add engine environment variables
...getEngineEnvs(),
// Add vector database environment variables
...getVectorDBEnvs(opts.vectorDb, opts.framework),
...getFrameworkEnvs(opts.framework, opts.port),
...getFrameworkEnvs(opts.framework, opts.externalPort),
...getToolEnvs(opts.tools),
// Add template environment variables
...getTemplateEnvs(opts.template),
getSystemPromptEnv(opts.tools),
...getObservabilityEnvs(opts.observability),
...getSystemPromptEnv(opts.tools, opts.dataSources, opts.framework),
];
// Render and write env file
const content = renderEnvVar(envVars);
@@ -493,6 +580,7 @@ export const createFrontendEnvFile = async (
root: string,
opts: {
customApiPath?: string;
vectorDb?: TemplateVectorDB;
},
) => {
const defaultFrontendEnvs = [
@@ -503,6 +591,11 @@ export const createFrontendEnvFile = async (
? opts.customApiPath
: "http://localhost:8000/api/chat",
},
{
name: "NEXT_PUBLIC_USE_LLAMACLOUD",
description: "Let's the user change indexes in LlamaCloud projects",
value: opts.vectorDb === "llamacloud" ? "true" : "false",
},
];
const content = renderEnvVar(defaultFrontendEnvs);
await fs.writeFile(path.join(root, ".env"), content);
+11 -16
View File
@@ -142,12 +142,15 @@ export const installTemplate = async (
if (props.framework === "fastapi") {
await installPythonTemplate(props);
// write loaders configuration (currently Python only)
await writeLoadersConfig(
props.root,
props.dataSources,
props.useLlamaParse,
);
if (props.vectorDb !== "llamacloud") {
// write loaders configuration (currently Python only)
// not needed for LlamaCloud as it has its own loaders
await writeLoadersConfig(
props.root,
props.dataSources,
props.useLlamaParse,
);
}
} else {
await installTSTemplate(props);
}
@@ -168,16 +171,7 @@ export const installTemplate = async (
props.template === "multiagent" ||
props.template === "extractor"
) {
await createBackendEnvFile(props.root, {
modelConfig: props.modelConfig,
llamaCloudKey: props.llamaCloudKey,
vectorDb: props.vectorDb,
framework: props.framework,
dataSources: props.dataSources,
port: props.externalPort,
tools: props.tools,
template: props.template,
});
await createBackendEnvFile(props.root, props);
}
if (props.dataSources.length > 0) {
@@ -209,6 +203,7 @@ export const installTemplate = async (
// this is a frontend for a full-stack app, create .env file with model information
await createFrontendEnvFile(props.root, {
customApiPath: props.customApiPath,
vectorDb: props.vectorDb,
});
}
};
+4 -4
View File
@@ -9,6 +9,7 @@ const ALL_AZURE_OPENAI_CHAT_MODELS: Record<string, { openAIModel: string }> = {
openAIModel: "gpt-3.5-turbo-16k",
},
"gpt-4o": { openAIModel: "gpt-4o" },
"gpt-4o-mini": { openAIModel: "gpt-4o-mini" },
"gpt-4": { openAIModel: "gpt-4" },
"gpt-4-32k": { openAIModel: "gpt-4-32k" },
"gpt-4-turbo": {
@@ -26,6 +27,9 @@ const ALL_AZURE_OPENAI_CHAT_MODELS: Record<string, { openAIModel: string }> = {
"gpt-4o-2024-05-13": {
openAIModel: "gpt-4o-2024-05-13",
},
"gpt-4o-mini-2024-07-18": {
openAIModel: "gpt-4o-mini-2024-07-18",
},
};
const ALL_AZURE_OPENAI_EMBEDDING_MODELS: Record<
@@ -35,10 +39,6 @@ const ALL_AZURE_OPENAI_EMBEDDING_MODELS: Record<
openAIModel: string;
}
> = {
"text-embedding-ada-002": {
dimensions: 1536,
openAIModel: "text-embedding-ada-002",
},
"text-embedding-3-small": {
dimensions: 1536,
openAIModel: "text-embedding-3-small",
+31 -12
View File
@@ -84,6 +84,13 @@ const getAdditionalDependencies = (
});
break;
}
case "weaviate": {
dependencies.push({
name: "llama-index-vector-stores-weaviate",
version: "^1.0.2",
});
break;
}
}
// Add data source dependencies
@@ -343,12 +350,15 @@ export const installPythonTemplate = async ({
cwd: path.join(compPath, "vectordbs", "python", vectorDb ?? "none"),
});
// Copy all loaders to enginePath
const loaderPath = path.join(enginePath, "loaders");
await copy("**", loaderPath, {
parents: true,
cwd: path.join(compPath, "loaders", "python"),
});
if (vectorDb !== "llamacloud") {
// Copy all loaders to enginePath
// Not needed for LlamaCloud as it has its own loaders
const loaderPath = path.join(enginePath, "loaders");
await copy("**", loaderPath, {
parents: true,
cwd: path.join(compPath, "loaders", "python"),
});
}
// Copy settings.py to app
await copy("**", path.join(root, "app"), {
@@ -380,18 +390,27 @@ export const installPythonTemplate = async ({
tools,
);
if (observability === "opentelemetry") {
addOnDependencies.push({
name: "traceloop-sdk",
version: "^0.15.11",
});
if (observability && observability !== "none") {
if (observability === "traceloop") {
addOnDependencies.push({
name: "traceloop-sdk",
version: "^0.15.11",
});
}
if (observability === "llamatrace") {
addOnDependencies.push({
name: "llama-index-callbacks-arize-phoenix",
version: "^0.1.6",
});
}
const templateObservabilityPath = path.join(
templatesDir,
"components",
"observability",
"python",
"opentelemetry",
observability,
);
await copy("**", path.join(root, "app"), {
cwd: templateObservabilityPath,
+3 -2
View File
@@ -35,7 +35,8 @@ export type TemplateVectorDB =
| "astra"
| "qdrant"
| "chroma"
| "llamacloud";
| "llamacloud"
| "weaviate";
export type TemplatePostInstallAction =
| "none"
| "VSCode"
@@ -46,7 +47,7 @@ export type TemplateDataSource = {
config: TemplateDataSourceConfig;
};
export type TemplateDataSourceType = "file" | "web" | "db" | "llamacloud";
export type TemplateObservability = "none" | "opentelemetry";
export type TemplateObservability = "none" | "traceloop" | "llamatrace";
// Config for both file and folder
export type FileSourceConfig = {
path: string;
+2 -2
View File
@@ -70,7 +70,7 @@ export const installTSTemplate = async ({
);
const webpackConfigOtelFile = path.join(root, "webpack.config.o11y.mjs");
if (observability === "opentelemetry") {
if (observability === "traceloop") {
const webpackConfigDefaultFile = path.join(root, "webpack.config.mjs");
await fs.rm(webpackConfigDefaultFile);
await fs.rename(webpackConfigOtelFile, webpackConfigDefaultFile);
@@ -248,7 +248,7 @@ async function updatePackageJson({
};
}
if (observability === "opentelemetry") {
if (observability === "traceloop") {
packageJson.dependencies = {
...packageJson.dependencies,
"@traceloop/node-server-sdk": "^0.5.19",
+2 -2
View File
@@ -1,6 +1,6 @@
{
"name": "create-llama",
"version": "0.1.26",
"version": "0.1.37",
"description": "Create LlamaIndex-powered apps with one command",
"keywords": [
"rag",
@@ -9,7 +9,7 @@
],
"repository": {
"type": "git",
"url": "https://github.com/run-llama/LlamaIndexTS",
"url": "https://github.com/run-llama/create-llama",
"directory": "packages/create-llama"
},
"license": "MIT",
+5 -1
View File
@@ -103,6 +103,7 @@ const getVectorDbChoices = (framework: TemplateFramework) => {
{ title: "Astra", value: "astra" },
{ title: "Qdrant", value: "qdrant" },
{ title: "ChromaDB", value: "chroma" },
{ title: "Weaviate", value: "weaviate" },
];
const vectordbLang = framework === "fastapi" ? "python" : "typescript";
@@ -486,7 +487,10 @@ export const askQuestions = async (
message: "Would you like to set up observability?",
choices: [
{ title: "No", value: "none" },
{ title: "OpenTelemetry", value: "opentelemetry" },
...(program.framework === "fastapi"
? [{ title: "LlamaTrace", value: "llamatrace" }]
: []),
{ title: "Traceloop", value: "traceloop" },
],
initial: 0,
},
@@ -6,7 +6,7 @@ from app.engine.tools import ToolFactory
from app.engine.index import get_index
def get_chat_engine(filters=None):
def get_chat_engine(filters=None, params=None):
system_prompt = os.getenv("SYSTEM_PROMPT")
top_k = os.getenv("TOP_K", "3")
tools = []
@@ -1,8 +1,6 @@
import os
import yaml
import json
import importlib
from cachetools import cached, LRUCache
from llama_index.core.tools.tool_spec.base import BaseToolSpec
from llama_index.core.tools.function_tool import FunctionTool
@@ -13,7 +11,6 @@ class ToolType:
class ToolFactory:
TOOL_SOURCE_PACKAGE_MAP = {
ToolType.LLAMAHUB: "llama_index.tools",
ToolType.LOCAL: "app.engine.tools",
@@ -3,7 +3,7 @@ import logging
import base64
import uuid
from pydantic import BaseModel
from typing import List, Tuple, Dict, Optional
from typing import List, Dict, Optional
from llama_index.core.tools import FunctionTool
from e2b_code_interpreter import CodeInterpreter
from e2b_code_interpreter.models import Logs
@@ -26,7 +26,6 @@ class E2BToolOutput(BaseModel):
class E2BCodeInterpreter:
output_dir = "output/tool"
def __init__(self, api_key: str = None):
@@ -1,13 +1,22 @@
import os
from app.engine.index import get_index
from app.engine.node_postprocessors import NodeCitationProcessor
from fastapi import HTTPException
from llama_index.core.chat_engine import CondensePlusContextChatEngine
def get_chat_engine(filters=None):
def get_chat_engine(filters=None, params=None):
system_prompt = os.getenv("SYSTEM_PROMPT")
top_k = os.getenv("TOP_K", 3)
citation_prompt = os.getenv("SYSTEM_CITATION_PROMPT", None)
top_k = int(os.getenv("TOP_K", 3))
index = get_index()
node_postprocessors = []
if citation_prompt:
node_postprocessors = [NodeCitationProcessor()]
system_prompt = f"{system_prompt}\n{citation_prompt}"
index = get_index(params)
if index is None:
raise HTTPException(
status_code=500,
@@ -16,9 +25,13 @@ def get_chat_engine(filters=None):
),
)
return index.as_chat_engine(
similarity_top_k=int(top_k),
system_prompt=system_prompt,
chat_mode="condense_plus_context",
retriever = index.as_retriever(
similarity_top_k=top_k,
filters=filters,
)
return CondensePlusContextChatEngine.from_defaults(
system_prompt=system_prompt,
retriever=retriever,
node_postprocessors=node_postprocessors,
)
@@ -0,0 +1,21 @@
from typing import List, Optional
from llama_index.core import QueryBundle
from llama_index.core.postprocessor.types import BaseNodePostprocessor
from llama_index.core.schema import NodeWithScore
class NodeCitationProcessor(BaseNodePostprocessor):
"""
Append node_id into metadata for citation purpose.
Config SYSTEM_CITATION_PROMPT in your runtime environment variable to enable this feature.
"""
def _postprocess_nodes(
self,
nodes: List[NodeWithScore],
query_bundle: Optional[QueryBundle] = None,
) -> List[NodeWithScore]:
for node_score in nodes:
node_score.node.metadata["node_id"] = node_score.node.node_id
return nodes
@@ -1,21 +1,16 @@
import {
BaseToolWithCall,
MetadataFilter,
MetadataFilters,
OpenAIAgent,
QueryEngineTool,
} from "llamaindex";
import { BaseToolWithCall, OpenAIAgent, QueryEngineTool } from "llamaindex";
import fs from "node:fs/promises";
import path from "node:path";
import { getDataSource } from "./index";
import { generateFilters } from "./queryFilter";
import { createTools } from "./tools";
export async function createChatEngine(documentIds?: string[]) {
export async function createChatEngine(documentIds?: string[], params?: any) {
const tools: BaseToolWithCall[] = [];
// Add a query engine tool if we have a data source
// Delete this code if you don't have a data source
const index = await getDataSource();
const index = await getDataSource(params);
if (index) {
tools.push(
new QueryEngineTool({
@@ -47,27 +42,3 @@ export async function createChatEngine(documentIds?: string[]) {
systemPrompt: process.env.SYSTEM_PROMPT,
});
}
function generateFilters(documentIds: string[]): MetadataFilters | undefined {
// public documents don't have the "private" field or it's set to "false"
const publicDocumentsFilter: MetadataFilter = {
key: "private",
value: ["true"],
operator: "nin",
};
// if no documentIds are provided, only retrieve information from public documents
if (!documentIds.length) return { filters: [publicDocumentsFilter] };
const privateDocumentsFilter: MetadataFilter = {
key: "doc_id",
value: documentIds,
operator: "in",
};
// if documentIds are provided, retrieve information from public and private documents
return {
filters: [publicDocumentsFilter, privateDocumentsFilter],
condition: "or",
};
}
@@ -1,13 +1,9 @@
import {
ContextChatEngine,
MetadataFilter,
MetadataFilters,
Settings,
} from "llamaindex";
import { ContextChatEngine, Settings } from "llamaindex";
import { getDataSource } from "./index";
import { generateFilters } from "./queryFilter";
export async function createChatEngine(documentIds?: string[]) {
const index = await getDataSource();
export async function createChatEngine(documentIds?: string[], params?: any) {
const index = await getDataSource(params);
if (!index) {
throw new Error(
`StorageContext is empty - call 'npm run generate' to generate the storage first`,
@@ -24,27 +20,3 @@ export async function createChatEngine(documentIds?: string[]) {
systemPrompt: process.env.SYSTEM_PROMPT,
});
}
function generateFilters(documentIds: string[]): MetadataFilters {
// public documents don't have the "private" field or it's set to "false"
const publicDocumentsFilter: MetadataFilter = {
key: "private",
value: ["true"],
operator: "nin",
};
// if no documentIds are provided, only retrieve information from public documents
if (!documentIds.length) return { filters: [publicDocumentsFilter] };
const privateDocumentsFilter: MetadataFilter = {
key: "doc_id",
value: documentIds,
operator: "in",
};
// if documentIds are provided, retrieve information from public and private documents
return {
filters: [publicDocumentsFilter, privateDocumentsFilter],
condition: "or",
};
}
@@ -1,28 +1,16 @@
import {
BaseNode,
Document,
IngestionPipeline,
Metadata,
Settings,
SimpleNodeParser,
storageContextFromDefaults,
VectorStoreIndex,
} from "llamaindex";
import { LlamaCloudIndex } from "llamaindex/cloud/LlamaCloudIndex";
import { getDataSource } from "../../engine";
export async function runPipeline(documents: Document[], filename: string) {
const currentIndex = await getDataSource();
// Update documents with metadata
for (const document of documents) {
document.metadata = {
...document.metadata,
file_name: filename,
private: "true", // to separate from other public documents
};
}
export async function runPipeline(
currentIndex: VectorStoreIndex | LlamaCloudIndex,
documents: Document[],
) {
if (currentIndex instanceof LlamaCloudIndex) {
// LlamaCloudIndex processes the documents automatically
// so we don't need ingestion pipeline, just insert the documents directly
@@ -41,25 +29,10 @@ export async function runPipeline(documents: Document[], filename: string) {
],
});
const nodes = await pipeline.run({ documents });
await addNodesToVectorStore(nodes, currentIndex);
await currentIndex.insertNodes(nodes);
currentIndex.storageContext.docStore.persist();
console.log("Added nodes to the vector store.");
}
return documents.map((document) => document.id_);
}
async function addNodesToVectorStore(
nodes: BaseNode<Metadata>[],
currentIndex: VectorStoreIndex | null,
) {
if (currentIndex) {
await currentIndex.insertNodes(nodes);
} else {
// Not using vectordb and haven't generated local index yet
const storageContext = await storageContextFromDefaults({
persistDir: "./cache",
});
currentIndex = await VectorStoreIndex.init({ nodes, storageContext });
}
currentIndex.storageContext.docStore.persist();
console.log("Added nodes to the vector store.");
}
@@ -1,11 +1,26 @@
import { VectorStoreIndex } from "llamaindex";
import { LlamaCloudIndex } from "llamaindex/cloud/LlamaCloudIndex";
import { loadDocuments, saveDocument } from "./helper";
import { runPipeline } from "./pipeline";
export async function uploadDocument(raw: string): Promise<string[]> {
export async function uploadDocument(
index: VectorStoreIndex | LlamaCloudIndex,
raw: string,
): Promise<string[]> {
const [header, content] = raw.split(",");
const mimeType = header.replace("data:", "").replace(";base64", "");
const fileBuffer = Buffer.from(content, "base64");
const documents = await loadDocuments(fileBuffer, mimeType);
const { filename } = await saveDocument(fileBuffer, mimeType);
return await runPipeline(documents, filename);
// Update documents with metadata
for (const document of documents) {
document.metadata = {
...document.metadata,
file_name: filename,
private: "true", // to separate private uploads from public documents
};
}
return await runPipeline(index, documents);
}
@@ -2,6 +2,7 @@ import { StreamData } from "ai";
import {
CallbackManager,
Metadata,
MetadataMode,
NodeWithScore,
ToolCall,
ToolOutput,
@@ -15,10 +16,11 @@ export function appendSourceData(
if (!sourceNodes?.length) return;
try {
const nodes = sourceNodes.map((node) => ({
...node.node.toMutableJSON(),
metadata: node.node.metadata,
id: node.node.id_,
score: node.score ?? null,
url: getNodeUrl(node.node.metadata),
text: node.node.getContent(MetadataMode.NONE),
}));
data.appendMessageAnnotation({
type: "sources",
@@ -7,19 +7,51 @@ const LLAMA_CLOUD_OUTPUT_DIR = "output/llamacloud";
const LLAMA_CLOUD_BASE_URL = "https://cloud.llamaindex.ai/api/v1";
const FILE_DELIMITER = "$"; // delimiter between pipelineId and filename
interface LlamaCloudFile {
type LlamaCloudFile = {
name: string;
file_id: string;
project_id: string;
}
};
type LLamaCloudProject = {
id: string;
organization_id: string;
name: string;
is_default: boolean;
};
type LLamaCloudPipeline = {
id: string;
name: string;
project_id: string;
};
export class LLamaCloudFileService {
private static readonly headers = {
Accept: "application/json",
Authorization: `Bearer ${process.env.LLAMA_CLOUD_API_KEY}`,
};
public static async getAllProjectsWithPipelines() {
try {
const projects = await LLamaCloudFileService.getAllProjects();
const pipelines = await LLamaCloudFileService.getAllPipelines();
return projects.map((project) => ({
...project,
pipelines: pipelines.filter((p) => p.project_id === project.id),
}));
} catch (error) {
console.error("Error listing projects and pipelines:", error);
return [];
}
}
public static async downloadFiles(nodes: NodeWithScore<Metadata>[]) {
const files = this.nodesToDownloadFiles(nodes);
const files = LLamaCloudFileService.nodesToDownloadFiles(nodes);
if (!files.length) return;
console.log("Downloading files from LlamaCloud...");
for (const file of files) {
await this.downloadFile(file.pipelineId, file.fileName);
await LLamaCloudFileService.downloadFile(file.pipelineId, file.fileName);
}
}
@@ -59,13 +91,19 @@ export class LLamaCloudFileService {
private static async downloadFile(pipelineId: string, fileName: string) {
try {
const downloadedName = this.toDownloadedName(pipelineId, fileName);
const downloadedName = LLamaCloudFileService.toDownloadedName(
pipelineId,
fileName,
);
const downloadedPath = path.join(LLAMA_CLOUD_OUTPUT_DIR, downloadedName);
// Check if file already exists
if (fs.existsSync(downloadedPath)) return;
const urlToDownload = await this.getFileUrlByName(pipelineId, fileName);
const urlToDownload = await LLamaCloudFileService.getFileUrlByName(
pipelineId,
fileName,
);
if (!urlToDownload) throw new Error("File not found in LlamaCloud");
const file = fs.createWriteStream(downloadedPath);
@@ -93,10 +131,13 @@ export class LLamaCloudFileService {
pipelineId: string,
name: string,
): Promise<string | null> {
const files = await this.getAllFiles(pipelineId);
const files = await LLamaCloudFileService.getAllFiles(pipelineId);
const file = files.find((file) => file.name === name);
if (!file) return null;
return await this.getFileUrlById(file.project_id, file.file_id);
return await LLamaCloudFileService.getFileUrlById(
file.project_id,
file.file_id,
);
}
private static async getFileUrlById(
@@ -104,11 +145,10 @@ export class LLamaCloudFileService {
fileId: string,
): Promise<string> {
const url = `${LLAMA_CLOUD_BASE_URL}/files/${fileId}/content?project_id=${projectId}`;
const headers = {
Accept: "application/json",
Authorization: `Bearer ${process.env.LLAMA_CLOUD_API_KEY}`,
};
const response = await fetch(url, { method: "GET", headers });
const response = await fetch(url, {
method: "GET",
headers: LLamaCloudFileService.headers,
});
const data = (await response.json()) as { url: string };
return data.url;
}
@@ -117,12 +157,31 @@ export class LLamaCloudFileService {
pipelineId: string,
): Promise<LlamaCloudFile[]> {
const url = `${LLAMA_CLOUD_BASE_URL}/pipelines/${pipelineId}/files`;
const headers = {
Accept: "application/json",
Authorization: `Bearer ${process.env.LLAMA_CLOUD_API_KEY}`,
};
const response = await fetch(url, { method: "GET", headers });
const response = await fetch(url, {
method: "GET",
headers: LLamaCloudFileService.headers,
});
const data = await response.json();
return data;
}
private static async getAllProjects(): Promise<LLamaCloudProject[]> {
const url = `${LLAMA_CLOUD_BASE_URL}/projects`;
const response = await fetch(url, {
method: "GET",
headers: LLamaCloudFileService.headers,
});
const data = (await response.json()) as LLamaCloudProject[];
return data;
}
private static async getAllPipelines(): Promise<LLamaCloudPipeline[]> {
const url = `${LLAMA_CLOUD_BASE_URL}/pipelines`;
const response = await fetch(url, {
method: "GET",
headers: LLamaCloudFileService.headers,
});
const data = (await response.json()) as LLamaCloudPipeline[];
return data;
}
}
@@ -33,8 +33,8 @@ export async function generateNextQuestions(
const questions = extractQuestions(response.text);
return questions;
} catch (error) {
console.error("Error: ", error);
throw error;
console.error("Error when generating the next questions: ", error);
return [];
}
}
+1 -3
View File
@@ -1,8 +1,6 @@
import os
import logging
from typing import List
from pydantic import BaseModel, validator
from llama_index.core.indices.vector_store import VectorStoreIndex
from pydantic import BaseModel
logger = logging.getLogger(__name__)
@@ -1,5 +1,3 @@
import os
import json
from pydantic import BaseModel, Field
@@ -10,7 +10,15 @@ export function getExtractors() {
}
export async function getDocuments() {
return await new SimpleDirectoryReader().loadData({
const documents = await new SimpleDirectoryReader().loadData({
directoryPath: DATA_DIR,
});
// Set private=false to mark the document as public (required for filtering)
for (const document of documents) {
document.metadata = {
...document.metadata,
private: "false",
};
}
return documents;
}
@@ -23,8 +23,16 @@ export function getExtractors() {
export async function getDocuments() {
const reader = new SimpleDirectoryReader();
const extractors = getExtractors();
return await reader.loadData({
const documents = await reader.loadData({
directoryPath: DATA_DIR,
fileExtToReader: extractors,
});
// Set private=false to mark the document as public (required for filtering)
for (const document of documents) {
document.metadata = {
...document.metadata,
private: "false",
};
}
return documents;
}
@@ -0,0 +1,12 @@
import llama_index.core
import os
def init_observability():
PHOENIX_API_KEY = os.getenv("PHOENIX_API_KEY")
if not PHOENIX_API_KEY:
raise ValueError("PHOENIX_API_KEY environment variable is not set")
os.environ["OTEL_EXPORTER_OTLP_HEADERS"] = f"api_key={PHOENIX_API_KEY}"
llama_index.core.set_global_handler(
"arize_phoenix", endpoint="https://llamatrace.com/v1/traces"
)
@@ -6,7 +6,7 @@ authors = ["Marcus Schiesser <mail@marcusschiesser.de>"]
readme = "README.md"
[tool.poetry.dependencies]
python = "^3.11,<3.12"
python = "^3.11,<4.0"
llama-index = "^0.10.6"
llama-index-readers-file = "^0.1.3"
python-dotenv = "^1.0.0"
@@ -6,11 +6,13 @@ import os
DEFAULT_MODEL = "gpt-3.5-turbo"
DEFAULT_EMBEDDING_MODEL = "text-embedding-3-large"
class TSIEmbedding(OpenAIEmbedding):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self._query_engine = self._text_engine = self.model_name
def llm_config_from_env() -> Dict:
from llama_index.core.constants import DEFAULT_TEMPERATURE
@@ -32,7 +34,7 @@ def llm_config_from_env() -> Dict:
def embedding_config_from_env() -> Dict:
from llama_index.core.constants import DEFAULT_EMBEDDING_DIM
model = os.getenv("EMBEDDING_MODEL", DEFAULT_EMBEDDING_MODEL)
dimension = os.getenv("EMBEDDING_DIM", DEFAULT_EMBEDDING_DIM)
api_key = os.getenv("T_SYSTEMS_LLMHUB_API_KEY")
@@ -46,6 +48,7 @@ def embedding_config_from_env() -> Dict:
}
return config
def init_llmhub():
from llama_index.llms.openai_like import OpenAILike
@@ -58,4 +61,4 @@ def init_llmhub():
is_chat_model=True,
is_function_calling_model=False,
context_window=4096,
)
)
@@ -82,7 +82,7 @@ def init_azure_openai():
dimensions = os.getenv("EMBEDDING_DIM")
azure_config = {
"api_key": os.environ["AZURE_OPENAI_KEY"],
"api_key": os.environ["AZURE_OPENAI_API_KEY"],
"azure_endpoint": os.environ["AZURE_OPENAI_ENDPOINT"],
"api_version": os.getenv("AZURE_OPENAI_API_VERSION")
or os.getenv("OPENAI_API_VERSION"),
@@ -1,31 +1,24 @@
"use client";
import { useEffect, useMemo, useState } from "react";
export interface ChatConfig {
backend?: string;
starterQuestions?: string[];
}
function getBackendOrigin(): string {
const chatAPI = process.env.NEXT_PUBLIC_CHAT_API;
if (chatAPI) {
return new URL(chatAPI).origin;
} else {
if (typeof window !== "undefined") {
// Use BASE_URL from window.ENV
return (window as any).ENV?.BASE_URL || "";
}
return "";
}
}
export function useClientConfig(): ChatConfig {
const chatAPI = process.env.NEXT_PUBLIC_CHAT_API;
const [config, setConfig] = useState<ChatConfig>();
const backendOrigin = useMemo(() => {
return chatAPI ? new URL(chatAPI).origin : "";
}, [chatAPI]);
const configAPI = `${backendOrigin}/api/chat/config`;
useEffect(() => {
fetch(configAPI)
.then((response) => response.json())
.then((data) => setConfig({ ...data, chatAPI }))
.catch((error) => console.error("Error fetching config", error));
}, [chatAPI, configAPI]);
return {
backend: backendOrigin,
starterQuestions: config?.starterQuestions,
backend: getBackendOrigin(),
};
}
@@ -1,48 +1,47 @@
# flake8: noqa: E402
from dotenv import load_dotenv
from app.engine.index import get_index
load_dotenv()
import os
import logging
from app.settings import init_settings
from app.engine.loaders import get_documents
from llama_index.indices.managed.llama_cloud import LlamaCloudIndex
from llama_index.core.readers import SimpleDirectoryReader
from app.engine.service import LLamaCloudFileService
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger()
def generate_datasource():
init_settings()
logger.info("Generate index for the provided data")
name = os.getenv("LLAMA_CLOUD_INDEX_NAME")
project_name = os.getenv("LLAMA_CLOUD_PROJECT_NAME")
api_key = os.getenv("LLAMA_CLOUD_API_KEY")
base_url = os.getenv("LLAMA_CLOUD_BASE_URL")
organization_id = os.getenv("LLAMA_CLOUD_ORGANIZATION_ID")
index = get_index()
project_id = index._get_project_id()
pipeline_id = index._get_pipeline_id()
if name is None or project_name is None or api_key is None:
raise ValueError(
"Please set LLAMA_CLOUD_INDEX_NAME, LLAMA_CLOUD_PROJECT_NAME and LLAMA_CLOUD_API_KEY"
" to your environment variables or config them in .env file"
)
documents = get_documents()
# Set private=false to mark the document as public (required for filtering)
for doc in documents:
doc.metadata["private"] = "false"
LlamaCloudIndex.from_documents(
documents=documents,
name=name,
project_name=project_name,
api_key=api_key,
base_url=base_url,
organization_id=organization_id
# use SimpleDirectoryReader to retrieve the files to process
reader = SimpleDirectoryReader(
"data",
recursive=True,
)
files_to_process = reader.input_files
# add each file to the LlamaCloud pipeline
for input_file in files_to_process:
with open(input_file, "rb") as f:
logger.info(
f"Adding file {input_file} to pipeline {index.name} in project {index.project_name}"
)
LLamaCloudFileService.add_file_to_pipeline(
project_id,
pipeline_id,
f,
custom_metadata={
# Set private=false to mark the document as public (required for filtering)
"private": "false",
},
)
logger.info("Finished generating the index")
@@ -1,14 +1,25 @@
import logging
import os
from llama_index.indices.managed.llama_cloud import LlamaCloudIndex
from llama_index.core.ingestion.api_utils import (
get_client as llama_cloud_get_client,
)
logger = logging.getLogger("uvicorn")
def get_index():
name = os.getenv("LLAMA_CLOUD_INDEX_NAME")
project_name = os.getenv("LLAMA_CLOUD_PROJECT_NAME")
def get_client():
return llama_cloud_get_client(
os.getenv("LLAMA_CLOUD_API_KEY"),
os.getenv("LLAMA_CLOUD_BASE_URL"),
)
def get_index(params=None):
configParams = params or {}
pipelineConfig = configParams.get("llamaCloudPipeline", {})
name = pipelineConfig.get("pipeline", os.getenv("LLAMA_CLOUD_INDEX_NAME"))
project_name = pipelineConfig.get("project", os.getenv("LLAMA_CLOUD_PROJECT_NAME"))
api_key = os.getenv("LLAMA_CLOUD_API_KEY")
base_url = os.getenv("LLAMA_CLOUD_BASE_URL")
organization_id = os.getenv("LLAMA_CLOUD_ORGANIZATION_ID")
@@ -24,7 +35,7 @@ def get_index():
project_name=project_name,
api_key=api_key,
base_url=base_url,
organization_id=organization_id
organization_id=organization_id,
)
return index
@@ -0,0 +1,35 @@
from llama_index.core.vector_stores.types import MetadataFilter, MetadataFilters
def generate_filters(doc_ids):
"""
Generate public/private document filters based on the doc_ids and the vector store.
"""
# Using "nin" filter to include the documents don't have the "private" key because they're uploaded in LlamaCloud UI
public_doc_filter = MetadataFilter(
key="private",
value=["true"],
operator="nin", # type: ignore
)
selected_doc_filter = MetadataFilter(
key="file_id", # Note: LLamaCloud uses "file_id" to reference private document ids as "doc_id" is a restricted field in LlamaCloud
value=doc_ids,
operator="in", # type: ignore
)
if len(doc_ids) > 0:
# If doc_ids are provided, we will select both public and selected documents
filters = MetadataFilters(
filters=[
public_doc_filter,
selected_doc_filter,
],
condition="or", # type: ignore
)
else:
filters = MetadataFilters(
filters=[
public_doc_filter,
]
)
return filters
@@ -0,0 +1,173 @@
from io import BytesIO
import logging
import os
import time
from typing import Any, Dict, List, Optional, Set, Tuple, Union
import typing
from fastapi import BackgroundTasks
from llama_cloud import ManagedIngestionStatus, PipelineFileCreateCustomMetadataValue
from pydantic import BaseModel
import requests
from app.api.routers.models import SourceNodes
from app.engine.index import get_client
from llama_index.core.schema import NodeWithScore
logger = logging.getLogger("uvicorn")
class LlamaCloudFile(BaseModel):
file_name: str
pipeline_id: str
def __eq__(self, other):
if not isinstance(other, LlamaCloudFile):
return NotImplemented
return (
self.file_name == other.file_name and self.pipeline_id == other.pipeline_id
)
def __hash__(self):
return hash((self.file_name, self.pipeline_id))
class LLamaCloudFileService:
LOCAL_STORE_PATH = "output/llamacloud"
DOWNLOAD_FILE_NAME_TPL = "{pipeline_id}${filename}"
@classmethod
def get_all_projects_with_pipelines(cls) -> List[Dict[str, Any]]:
try:
client = get_client()
projects = client.projects.list_projects()
pipelines = client.pipelines.search_pipelines()
return [
{
**(project.dict()),
"pipelines": [
{"id": p.id, "name": p.name}
for p in pipelines
if p.project_id == project.id
],
}
for project in projects
]
except Exception as error:
logger.error(f"Error listing projects and pipelines: {error}")
return []
@classmethod
def add_file_to_pipeline(
cls,
project_id: str,
pipeline_id: str,
upload_file: Union[typing.IO, Tuple[str, BytesIO]],
custom_metadata: Optional[Dict[str, PipelineFileCreateCustomMetadataValue]],
) -> str:
client = get_client()
file = client.files.upload_file(project_id=project_id, upload_file=upload_file)
files = [
{
"file_id": file.id,
"custom_metadata": {"file_id": file.id, **(custom_metadata or {})},
}
]
files = client.pipelines.add_files_to_pipeline(pipeline_id, request=files)
# Wait 2s for the file to be processed
max_attempts = 20
attempt = 0
while attempt < max_attempts:
result = client.pipelines.get_pipeline_file_status(pipeline_id, file.id)
if result.status == ManagedIngestionStatus.ERROR:
raise Exception(f"File processing failed: {str(result)}")
if result.status == ManagedIngestionStatus.SUCCESS:
# File is ingested - return the file id
return file.id
attempt += 1
time.sleep(0.1) # Sleep for 100ms
raise Exception(
f"File processing did not complete after {max_attempts} attempts."
)
@classmethod
def download_pipeline_file(
cls,
file: LlamaCloudFile,
force_download: bool = False,
):
client = get_client()
file_name = file.file_name
pipeline_id = file.pipeline_id
# Check is the file already exists
downloaded_file_path = cls._get_file_path(file_name, pipeline_id)
if os.path.exists(downloaded_file_path) and not force_download:
logger.debug(f"File {file_name} already exists in local storage")
return
try:
logger.info(f"Downloading file {file_name} for pipeline {pipeline_id}")
files = client.pipelines.list_pipeline_files(pipeline_id)
if not files or not isinstance(files, list):
raise Exception("No files found in LlamaCloud")
for file_entry in files:
if file_entry.name == file_name:
file_id = file_entry.file_id
project_id = file_entry.project_id
file_detail = client.files.read_file_content(
file_id, project_id=project_id
)
cls._download_file(file_detail.url, downloaded_file_path)
break
except Exception as error:
logger.info(f"Error fetching file from LlamaCloud: {error}")
@classmethod
def download_files_from_nodes(
cls, nodes: List[NodeWithScore], background_tasks: BackgroundTasks
):
files = cls._get_files_to_download(nodes)
for file in files:
logger.info(f"Adding download of {file.file_name} to background tasks")
background_tasks.add_task(
LLamaCloudFileService.download_pipeline_file, file
)
@classmethod
def _get_files_to_download(cls, nodes: List[NodeWithScore]) -> Set[LlamaCloudFile]:
source_nodes = SourceNodes.from_source_nodes(nodes)
llama_cloud_files = [
LlamaCloudFile(
file_name=node.metadata.get("file_name"),
pipeline_id=node.metadata.get("pipeline_id"),
)
for node in source_nodes
if (
node.metadata.get("pipeline_id") is not None
and node.metadata.get("file_name") is not None
)
]
# Remove duplicates and return
return set(llama_cloud_files)
@classmethod
def _get_file_name(cls, name: str, pipeline_id: str) -> str:
return cls.DOWNLOAD_FILE_NAME_TPL.format(pipeline_id=pipeline_id, filename=name)
@classmethod
def _get_file_path(cls, name: str, pipeline_id: str) -> str:
return os.path.join(cls.LOCAL_STORE_PATH, cls._get_file_name(name, pipeline_id))
@classmethod
def _download_file(cls, url: str, local_file_path: str):
logger.info(f"Saving file to {local_file_path}")
# Create directory if it doesn't exist
os.makedirs(cls.LOCAL_STORE_PATH, exist_ok=True)
# Download the file
with requests.get(url, stream=True) as r:
r.raise_for_status()
with open(local_file_path, "wb") as f:
for chunk in r.iter_content(chunk_size=8192):
f.write(chunk)
logger.info("File downloaded successfully")
@@ -1,3 +1,4 @@
# flake8: noqa: E402
from dotenv import load_dotenv
load_dotenv()
@@ -17,7 +17,7 @@ def get_storage_context(persist_dir: str) -> StorageContext:
return StorageContext.from_defaults(persist_dir=persist_dir)
def get_index():
def get_index(params=None):
storage_dir = os.getenv("STORAGE_DIR", "storage")
# check if storage already exists
if not os.path.exists(storage_dir):
@@ -0,0 +1,36 @@
from llama_index.core.vector_stores.types import MetadataFilter, MetadataFilters
def generate_filters(doc_ids):
"""
Generate public/private document filters based on the doc_ids and the vector store.
"""
public_doc_filter = MetadataFilter(
key="private",
value="true",
operator="!=", # type: ignore
)
# Weaviate doesn't support "in" filter right now, so use "any" instead - it has the same behavior.
# TODO: Use "in" operator, once Weaviate supports it
selected_doc_filter = MetadataFilter(
key="doc_id",
value=doc_ids,
operator="any", # type: ignore
)
if len(doc_ids) > 0:
# If doc_ids are provided, we will select both public and selected documents
filters = MetadataFilters(
filters=[
public_doc_filter,
selected_doc_filter,
],
condition="or", # type: ignore
)
else:
filters = MetadataFilters(
filters=[
public_doc_filter,
]
)
return filters
@@ -0,0 +1,35 @@
import os
import weaviate
from llama_index.vector_stores.weaviate import WeaviateVectorStore
DEFAULT_INDEX_NAME = "LlamaIndex"
def _create_weaviate_client():
cluster_url = os.getenv("WEAVIATE_CLUSTER_URL")
api_key = os.getenv("WEAVIATE_API_KEY")
if not cluster_url or not api_key:
raise ValueError(
"Environment variables: WEAVIATE_CLUSTER_URL and WEAVIATE_API_KEY are required."
)
auth_credentials = weaviate.auth.AuthApiKey(api_key)
client = weaviate.connect_to_weaviate_cloud(cluster_url, auth_credentials)
return client
# Global variable to store the Weaviate client
client = None
def get_vector_store():
global client
if client is None:
client = _create_weaviate_client()
index_name = os.getenv("WEAVIATE_INDEX_NAME", DEFAULT_INDEX_NAME)
vector_store = WeaviateVectorStore(
weaviate_client=client,
index_name=index_name,
)
return vector_store
@@ -3,7 +3,7 @@ import { VectorStoreIndex } from "llamaindex";
import { AstraDBVectorStore } from "llamaindex/storage/vectorStore/AstraDBVectorStore";
import { checkRequiredEnvVars } from "./shared";
export async function getDataSource() {
export async function getDataSource(params?: any) {
checkRequiredEnvVars();
const store = new AstraDBVectorStore();
await store.connect(process.env.ASTRA_DB_COLLECTION!);
@@ -3,7 +3,7 @@ import { VectorStoreIndex } from "llamaindex";
import { ChromaVectorStore } from "llamaindex/storage/vectorStore/ChromaVectorStore";
import { checkRequiredEnvVars } from "./shared";
export async function getDataSource() {
export async function getDataSource(params?: any) {
checkRequiredEnvVars();
const chromaUri = `http://${process.env.CHROMA_HOST}:${process.env.CHROMA_PORT}`;
@@ -9,13 +9,7 @@ dotenv.config();
async function loadAndIndex() {
const documents = await getDocuments();
// Set private=false to mark the document as public (required for filtering)
for (const document of documents) {
document.metadata = {
...document.metadata,
private: "false",
};
}
await getDataSource();
await LlamaCloudIndex.fromDocuments({
documents,
@@ -1,12 +1,26 @@
import { LlamaCloudIndex } from "llamaindex/cloud/LlamaCloudIndex";
import { checkRequiredEnvVars } from "./shared";
export async function getDataSource() {
checkRequiredEnvVars();
type LlamaCloudDataSourceParams = {
llamaCloudPipeline?: {
project: string;
pipeline: string;
};
};
export async function getDataSource(params?: LlamaCloudDataSourceParams) {
const { project, pipeline } = params?.llamaCloudPipeline ?? {};
const projectName = project ?? process.env.LLAMA_CLOUD_PROJECT_NAME;
const pipelineName = pipeline ?? process.env.LLAMA_CLOUD_INDEX_NAME;
const apiKey = process.env.LLAMA_CLOUD_API_KEY;
if (!projectName || !pipelineName || !apiKey) {
throw new Error(
"Set project, pipeline, and api key in the params or as environment variables.",
);
}
const index = new LlamaCloudIndex({
name: process.env.LLAMA_CLOUD_INDEX_NAME!,
projectName: process.env.LLAMA_CLOUD_PROJECT_NAME!,
apiKey: process.env.LLAMA_CLOUD_API_KEY,
name: pipelineName,
projectName,
apiKey,
baseUrl: process.env.LLAMA_CLOUD_BASE_URL,
});
return index;
@@ -0,0 +1,25 @@
import { MetadataFilter, MetadataFilters } from "llamaindex";
export function generateFilters(documentIds: string[]): MetadataFilters {
// public documents don't have the "private" field or it's set to "false"
const publicDocumentsFilter: MetadataFilter = {
key: "private",
value: ["true"],
operator: "nin",
};
// if no documentIds are provided, only retrieve information from public documents
if (!documentIds.length) return { filters: [publicDocumentsFilter] };
const privateDocumentsFilter: MetadataFilter = {
key: "doc_id",
value: documentIds,
operator: "in",
};
// if documentIds are provided, retrieve information from public and private documents
return {
filters: [publicDocumentsFilter, privateDocumentsFilter],
condition: "or",
};
}
@@ -2,7 +2,7 @@ import { VectorStoreIndex } from "llamaindex";
import { MilvusVectorStore } from "llamaindex/storage/vectorStore/MilvusVectorStore";
import { checkRequiredEnvVars, getMilvusClient } from "./shared";
export async function getDataSource() {
export async function getDataSource(params?: any) {
checkRequiredEnvVars();
const milvusClient = getMilvusClient();
const store = new MilvusVectorStore({ milvusClient });
@@ -4,7 +4,6 @@ const REQUIRED_ENV_VARS = [
"MILVUS_ADDRESS",
"MILVUS_USERNAME",
"MILVUS_PASSWORD",
"MILVUS_COLLECTION",
];
export function getMilvusClient() {
@@ -19,8 +18,13 @@ export function getMilvusClient() {
});
}
export function checkRequiredEnvVars() {
const missingEnvVars = REQUIRED_ENV_VARS.filter((envVar) => {
export function checkRequiredEnvVars(opts?: { checkCollectionEnv?: boolean }) {
const shouldCheckCollectionEnv = opts?.checkCollectionEnv ?? true; // default to true
const requiredEnvVars = [...REQUIRED_ENV_VARS]; // create a copy of the array
if (shouldCheckCollectionEnv) {
requiredEnvVars.push("MILVUS_COLLECTION");
}
const missingEnvVars = requiredEnvVars.filter((envVar) => {
return !process.env[envVar];
});
@@ -1,7 +1,10 @@
/* eslint-disable turbo/no-undeclared-env-vars */
import * as dotenv from "dotenv";
import { VectorStoreIndex, storageContextFromDefaults } from "llamaindex";
import { MongoDBAtlasVectorSearch } from "llamaindex/storage/vectorStore/MongoDBAtlasVectorSearch";
import {
MongoDBAtlasVectorSearch,
VectorStoreIndex,
storageContextFromDefaults,
} from "llamaindex";
import { MongoClient } from "mongodb";
import { getDocuments } from "./loader";
import { initSettings } from "./settings";
@@ -1,10 +1,9 @@
/* eslint-disable turbo/no-undeclared-env-vars */
import { VectorStoreIndex } from "llamaindex";
import { MongoDBAtlasVectorSearch } from "llamaindex/storage/vectorStore/MongoDBAtlasVectorSearch";
import { MongoDBAtlasVectorSearch, VectorStoreIndex } from "llamaindex";
import { MongoClient } from "mongodb";
import { checkRequiredEnvVars } from "./shared";
export async function getDataSource() {
export async function getDataSource(params?: any) {
checkRequiredEnvVars();
const client = new MongoClient(process.env.MONGO_URI!);
const store = new MongoDBAtlasVectorSearch({
@@ -25,10 +25,6 @@ async function generateDatasource() {
persistDir: STORAGE_CACHE_DIR,
});
const documents = await getDocuments();
// Set private=false to mark the document as public (required for filtering)
documents.forEach((doc) => {
doc.metadata["private"] = "false";
});
await VectorStoreIndex.fromDocuments(documents, {
storageContext,
@@ -2,7 +2,7 @@ import { SimpleDocumentStore, VectorStoreIndex } from "llamaindex";
import { storageContextFromDefaults } from "llamaindex/storage/StorageContext";
import { STORAGE_CACHE_DIR } from "./shared";
export async function getDataSource() {
export async function getDataSource(params?: any) {
const storageContext = await storageContextFromDefaults({
persistDir: `${STORAGE_CACHE_DIR}`,
});
@@ -7,7 +7,7 @@ import {
checkRequiredEnvVars,
} from "./shared";
export async function getDataSource() {
export async function getDataSource(params?: any) {
checkRequiredEnvVars();
const pgvs = new PGVectorStore({
connectionString: process.env.PG_CONNECTION_STRING,
@@ -3,7 +3,7 @@ import { VectorStoreIndex } from "llamaindex";
import { PineconeVectorStore } from "llamaindex/storage/vectorStore/PineconeVectorStore";
import { checkRequiredEnvVars } from "./shared";
export async function getDataSource() {
export async function getDataSource(params?: any) {
checkRequiredEnvVars();
const store = new PineconeVectorStore();
return await VectorStoreIndex.fromVectorStore(store);
@@ -5,7 +5,7 @@ import { checkRequiredEnvVars, getQdrantClient } from "./shared";
dotenv.config();
export async function getDataSource() {
export async function getDataSource(params?: any) {
checkRequiredEnvVars();
const collectionName = process.env.QDRANT_COLLECTION;
const store = new QdrantVectorStore({
@@ -0,0 +1,33 @@
/* eslint-disable turbo/no-undeclared-env-vars */
import * as dotenv from "dotenv";
import {
VectorStoreIndex,
WeaviateVectorStore,
storageContextFromDefaults,
} from "llamaindex";
import { getDocuments } from "./loader";
import { initSettings } from "./settings";
import { DEFAULT_INDEX_NAME, checkRequiredEnvVars } from "./shared";
dotenv.config();
async function loadAndIndex() {
const indexName = process.env.WEAVIATE_INDEX_NAME || DEFAULT_INDEX_NAME;
// load objects from storage and convert them into LlamaIndex Document objects
const documents = await getDocuments();
const vectorStore = new WeaviateVectorStore({ indexName });
const storageContext = await storageContextFromDefaults({ vectorStore });
await VectorStoreIndex.fromDocuments(documents, {
storageContext: storageContext,
});
console.log(`Successfully upload embeddings to Weaviate index ${indexName}.`);
}
(async () => {
checkRequiredEnvVars();
initSettings();
await loadAndIndex();
console.log("Finished generating storage.");
})();
@@ -0,0 +1,14 @@
import * as dotenv from "dotenv";
import { VectorStoreIndex } from "llamaindex";
import { WeaviateVectorStore } from "llamaindex/storage/vectorStore/WeaviateVectorStore";
import { checkRequiredEnvVars, DEFAULT_INDEX_NAME } from "./shared";
dotenv.config();
export async function getDataSource(params?: any) {
checkRequiredEnvVars();
const indexName = process.env.WEAVIATE_INDEX_NAME || DEFAULT_INDEX_NAME;
const store = new WeaviateVectorStore({ indexName });
return await VectorStoreIndex.fromVectorStore(store);
}
@@ -0,0 +1,26 @@
import { MetadataFilter, MetadataFilters } from "llamaindex";
export function generateFilters(documentIds: string[]): MetadataFilters {
// filter all documents have the private metadata key set to true
const publicDocumentsFilter: MetadataFilter = {
key: "private",
value: "true",
operator: "!=",
};
// if no documentIds are provided, only retrieve information from public documents
if (!documentIds.length) return { filters: [publicDocumentsFilter] };
// Weaviate uses 'any' instead of 'in' for the operator
const privateDocumentsFilter: MetadataFilter = {
key: "doc_id",
value: documentIds,
operator: "any",
};
// if documentIds are provided, retrieve information from public and private documents
return {
filters: [publicDocumentsFilter, privateDocumentsFilter],
condition: "or",
};
}
@@ -0,0 +1,20 @@
const REQUIRED_ENV_VARS = ["WEAVIATE_CLUSTER_URL", "WEAVIATE_API_KEY"];
export const DEFAULT_INDEX_NAME = "LlamaIndex";
export function checkRequiredEnvVars() {
const missingEnvVars = REQUIRED_ENV_VARS.filter((envVar) => {
return !process.env[envVar];
});
if (missingEnvVars.length > 0) {
console.log(
`The following environment variables are required but missing: ${missingEnvVars.join(
", ",
)}`,
);
throw new Error(
`Missing environment variables: ${missingEnvVars.join(", ")}`,
);
}
}
@@ -1,3 +1,4 @@
# flake8: noqa: E402
from dotenv import load_dotenv
load_dotenv()
@@ -9,7 +9,7 @@ readme = "README.md"
generate = "app.engine.generate:generate_datasource"
[tool.poetry.dependencies]
python = "^3.11,<3.12"
python = "^3.11,<4.0"
fastapi = "^0.109.1"
uvicorn = { extras = ["standard"], version = "^0.23.2" }
python-dotenv = "^1.0.0"
@@ -5,4 +5,4 @@ def load_from_env(var: str, throw_error: bool = True) -> str:
res = os.getenv(var)
if res is None and throw_error:
raise ValueError(f"Missing environment variable: {var}")
return res
return res
@@ -1,3 +1,4 @@
# flake8: noqa: E402
from dotenv import load_dotenv
from app.settings import init_settings
@@ -20,7 +20,7 @@
"dotenv": "^16.3.1",
"duck-duck-scrape": "^2.2.5",
"express": "^4.18.2",
"llamaindex": "0.5.12",
"llamaindex": "0.5.17",
"pdf2json": "3.0.5",
"ajv": "^8.12.0",
"@e2b/code-interpreter": "^0.0.5",
@@ -1,4 +1,5 @@
import { Request, Response } from "express";
import { LLamaCloudFileService } from "./llamaindex/streaming/service";
export const chatConfig = async (_req: Request, res: Response) => {
let starterQuestions = undefined;
@@ -12,3 +13,14 @@ export const chatConfig = async (_req: Request, res: Response) => {
starterQuestions,
});
};
export const chatLlamaCloudConfig = async (_req: Request, res: Response) => {
const config = {
projects: await LLamaCloudFileService.getAllProjectsWithPipelines(),
pipeline: {
pipeline: process.env.LLAMA_CLOUD_INDEX_NAME,
project: process.env.LLAMA_CLOUD_PROJECT_NAME,
},
};
return res.status(200).json(config);
};
@@ -1,12 +1,14 @@
import { Request, Response } from "express";
import { getDataSource } from "./engine";
import { uploadDocument } from "./llamaindex/documents/upload";
export const chatUpload = async (req: Request, res: Response) => {
const { base64 }: { base64: string } = req.body;
const { base64, params }: { base64: string; params?: any } = req.body;
if (!base64) {
return res.status(400).json({
error: "base64 is required in the request body",
});
}
return res.status(200).json(await uploadDocument(base64));
const index = await getDataSource(params);
return res.status(200).json(await uploadDocument(index, base64));
};
@@ -17,7 +17,7 @@ export const chat = async (req: Request, res: Response) => {
const vercelStreamData = new StreamData();
const streamTimeout = createStreamTimeout(vercelStreamData);
try {
const { messages }: { messages: Message[] } = req.body;
const { messages, data }: { messages: Message[]; data?: any } = req.body;
const userMessage = messages.pop();
if (!messages || !userMessage || userMessage.role !== "user") {
return res.status(400).json({
@@ -46,7 +46,7 @@ export const chat = async (req: Request, res: Response) => {
},
);
const ids = retrieveDocumentIds(allAnnotations);
const chatEngine = await createChatEngine(ids);
const chatEngine = await createChatEngine(ids, data);
// Convert message content from Vercel/AI format to LlamaIndex/OpenAI format
const userMessageContent = convertMessageContent(
@@ -0,0 +1,25 @@
import { MetadataFilter, MetadataFilters } from "llamaindex";
export function generateFilters(documentIds: string[]): MetadataFilters {
// filter all documents have the private metadata key set to true
const publicDocumentsFilter: MetadataFilter = {
key: "private",
value: "true",
operator: "!=",
};
// if no documentIds are provided, only retrieve information from public documents
if (!documentIds.length) return { filters: [publicDocumentsFilter] };
const privateDocumentsFilter: MetadataFilter = {
key: "doc_id",
value: documentIds,
operator: "in",
};
// if documentIds are provided, retrieve information from public and private documents
return {
filters: [publicDocumentsFilter, privateDocumentsFilter],
condition: "or",
};
}
@@ -1,5 +1,8 @@
import express, { Router } from "express";
import { chatConfig } from "../controllers/chat-config.controller";
import {
chatConfig,
chatLlamaCloudConfig,
} from "../controllers/chat-config.controller";
import { chatRequest } from "../controllers/chat-request.controller";
import { chatUpload } from "../controllers/chat-upload.controller";
import { chat } from "../controllers/chat.controller";
@@ -11,6 +14,7 @@ initSettings();
llmRouter.route("/").post(chat);
llmRouter.route("/request").post(chatRequest);
llmRouter.route("/config").get(chatConfig);
llmRouter.route("/config/llamacloud").get(chatLlamaCloudConfig);
llmRouter.route("/upload").post(chatUpload);
export default llmRouter;
@@ -5,7 +5,7 @@
"forceConsistentCasingInFileNames": true,
"strict": true,
"skipLibCheck": true,
"moduleResolution": "node",
"moduleResolution": "bundler",
"paths": {
"@/*": ["./*"]
}
@@ -1,50 +1,32 @@
import logging
import os
from typing import List
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Request, status
from llama_index.core.chat_engine.types import BaseChatEngine, NodeWithScore
from llama_index.core.llms import MessageRole
from llama_index.core.vector_stores.types import MetadataFilter, MetadataFilters
from app.api.routers.events import EventCallbackHandler
from app.api.routers.models import (
ChatConfig,
ChatData,
Message,
Result,
SourceNodes,
)
from app.api.routers.vercel_response import VercelStreamResponse
from app.api.services.llama_cloud import LLamaCloudFileService
from app.engine import get_chat_engine
from app.engine.query_filter import generate_filters
chat_router = r = APIRouter()
logger = logging.getLogger("uvicorn")
def process_response_nodes(
nodes: List[NodeWithScore],
background_tasks: BackgroundTasks,
):
"""
Start background tasks on the source nodes if needed.
"""
files_to_download = SourceNodes.get_download_files(nodes)
for file in files_to_download:
background_tasks.add_task(
LLamaCloudFileService.download_llamacloud_pipeline_file, file
)
# streaming endpoint - delete if not needed
@r.post("")
async def chat(
request: Request,
data: ChatData,
background_tasks: BackgroundTasks,
chat_engine: BaseChatEngine = Depends(get_chat_engine),
):
try:
last_message_content = data.get_last_message_content()
@@ -52,8 +34,11 @@ async def chat(
doc_ids = data.get_chat_document_ids()
filters = generate_filters(doc_ids)
logger.info("Creating chat engine with filters", filters.dict())
chat_engine = get_chat_engine(filters=filters)
params = data.data or {}
logger.info(
f"Creating chat engine with filters: {str(filters)}",
)
chat_engine = get_chat_engine(filters=filters, params=params)
event_handler = EventCallbackHandler()
chat_engine.callback_manager.handlers.append(event_handler) # type: ignore
@@ -70,38 +55,6 @@ async def chat(
) from e
def generate_filters(doc_ids):
if len(doc_ids) > 0:
filters = MetadataFilters(
filters=[
MetadataFilter(
key="private",
value=["true"],
operator="nin", # type: ignore
),
MetadataFilter(
key="doc_id",
value=doc_ids,
operator="in", # type: ignore
),
],
condition="or", # type: ignore
)
else:
filters = MetadataFilters(
# Use the "NIN" - "not in" operator to include all public documents (don't have the private key set)
filters=[
MetadataFilter(
key="private",
value=["true"],
operator="nin", # type: ignore
),
]
)
return filters
# non-streaming endpoint - delete if not needed
@r.post("/request")
async def chat_request(
@@ -118,10 +71,15 @@ async def chat_request(
)
@r.get("/config")
async def chat_config() -> ChatConfig:
starter_questions = None
conversation_starters = os.getenv("CONVERSATION_STARTERS")
if conversation_starters and conversation_starters.strip():
starter_questions = conversation_starters.strip().split("\n")
return ChatConfig(starter_questions=starter_questions)
def process_response_nodes(
nodes: List[NodeWithScore],
background_tasks: BackgroundTasks,
):
try:
# Start background tasks to download documents from LlamaCloud if needed
from app.engine.service import LLamaCloudFileService
LLamaCloudFileService.download_files_from_nodes(nodes, background_tasks)
except ImportError:
logger.debug("LlamaCloud is not configured. Skipping post processing of nodes")
pass
@@ -0,0 +1,48 @@
import logging
import os
from fastapi import APIRouter
from app.api.routers.models import ChatConfig
config_router = r = APIRouter()
logger = logging.getLogger("uvicorn")
@r.get("")
async def chat_config() -> ChatConfig:
starter_questions = None
conversation_starters = os.getenv("CONVERSATION_STARTERS")
if conversation_starters and conversation_starters.strip():
starter_questions = conversation_starters.strip().split("\n")
return ChatConfig(starter_questions=starter_questions)
try:
from app.engine.service import LLamaCloudFileService
logger.info("LlamaCloud is configured. Adding /config/llamacloud route.")
@r.get("/llamacloud")
async def chat_llama_cloud_config():
projects = LLamaCloudFileService.get_all_projects_with_pipelines()
pipeline = os.getenv("LLAMA_CLOUD_INDEX_NAME")
project = os.getenv("LLAMA_CLOUD_PROJECT_NAME")
pipeline_config = None
if pipeline and project:
pipeline_config = {
"pipeline": pipeline,
"project": project,
}
return {
"projects": projects,
"pipeline": pipeline_config,
}
except ImportError:
logger.debug(
"LlamaCloud is not configured. Skipping adding /config/llamacloud route."
)
pass
@@ -1,6 +1,6 @@
import logging
import os
from typing import Any, Dict, List, Literal, Optional, Set
from typing import Any, Dict, List, Literal, Optional
from llama_index.core.llms import ChatMessage, MessageRole
from llama_index.core.schema import NodeWithScore
@@ -75,6 +75,7 @@ class Message(BaseModel):
class ChatData(BaseModel):
messages: List[Message]
data: Any = None
class Config:
json_schema_extra = {
@@ -146,21 +147,6 @@ class ChatData(BaseModel):
return list(set(document_ids))
class LlamaCloudFile(BaseModel):
file_name: str
pipeline_id: str
def __eq__(self, other):
if not isinstance(other, LlamaCloudFile):
return NotImplemented
return (
self.file_name == other.file_name and self.pipeline_id == other.pipeline_id
)
def __hash__(self):
return hash((self.file_name, self.pipeline_id))
class SourceNodes(BaseModel):
id: str
metadata: Dict[str, Any]
@@ -192,8 +178,8 @@ class SourceNodes(BaseModel):
if file_name and url_prefix:
# file_name exists and file server is configured
pipeline_id = metadata.get("pipeline_id")
if pipeline_id and metadata.get("private") is None:
# file is from LlamaCloud and was not ingested locally
if pipeline_id:
# file is from LlamaCloud
file_name = f"{pipeline_id}${file_name}"
return f"{url_prefix}/output/llamacloud/{file_name}"
is_private = metadata.get("private", "false") == "true"
@@ -208,25 +194,6 @@ class SourceNodes(BaseModel):
def from_source_nodes(cls, source_nodes: List[NodeWithScore]):
return [cls.from_source_node(node) for node in source_nodes]
@staticmethod
def get_download_files(nodes: List[NodeWithScore]) -> Set[LlamaCloudFile]:
source_nodes = SourceNodes.from_source_nodes(nodes)
llama_cloud_files = [
LlamaCloudFile(
file_name=node.metadata.get("file_name"),
pipeline_id=node.metadata.get("pipeline_id"),
)
for node in source_nodes
if (
node.metadata.get("private")
is None # Only download files are from LlamaCloud and were not ingested locally
and node.metadata.get("pipeline_id") is not None
and node.metadata.get("file_name") is not None
)
]
# Remove duplicates and return
return set(llama_cloud_files)
class Result(BaseModel):
result: Message
@@ -237,7 +204,7 @@ class ChatConfig(BaseModel):
starter_questions: Optional[List[str]] = Field(
default=None,
description="List of starter questions",
serialization_alias="starterQuestions"
serialization_alias="starterQuestions",
)
class Config:
@@ -1,5 +1,5 @@
import logging
from typing import List
from typing import List, Any
from fastapi import APIRouter, HTTPException
from pydantic import BaseModel
@@ -13,13 +13,17 @@ logger = logging.getLogger("uvicorn")
class FileUploadRequest(BaseModel):
base64: str
filename: str
params: Any = None
@r.post("")
def upload_file(request: FileUploadRequest) -> List[str]:
try:
logger.info("Processing file")
return PrivateFileService.process_file(request.base64)
return PrivateFileService.process_file(
request.filename, request.base64, request.params
)
except Exception as e:
logger.error(f"Error processing file: {e}", exc_info=True)
raise HTTPException(status_code=500, detail="Error processing file")
@@ -80,7 +80,7 @@ class VercelStreamResponse(StreamingResponse):
"type": "sources",
"data": {
"nodes": [
SourceNodes.from_source_node(node).dict()
SourceNodes.from_source_node(node).model_dump()
for node in response.source_nodes
]
},
@@ -1,9 +1,10 @@
import base64
import mimetypes
import os
from io import BytesIO
from pathlib import Path
from typing import Dict, List
from uuid import uuid4
from typing import Any, List, Tuple
from app.engine.index import get_index
from llama_index.core import VectorStoreIndex
@@ -11,9 +12,6 @@ from llama_index.core.ingestion import IngestionPipeline
from llama_index.core.readers.file.base import (
_try_loading_included_file_formats as get_file_loaders_map,
)
from llama_index.core.readers.file.base import (
default_file_metadata_func,
)
from llama_index.core.schema import Document
from llama_index.indices.managed.llama_cloud.base import LlamaCloudIndex
from llama_index.readers.file import FlatReader
@@ -41,7 +39,7 @@ class PrivateFileService:
PRIVATE_STORE_PATH = "output/uploaded"
@staticmethod
def preprocess_base64_file(base64_content: str) -> tuple:
def preprocess_base64_file(base64_content: str) -> Tuple[bytes, str | None]:
header, data = base64_content.split(",", 1)
mime_type = header.split(";")[0].split(":", 1)[1]
extension = mimetypes.guess_extension(mime_type)
@@ -49,12 +47,9 @@ class PrivateFileService:
return base64.b64decode(data), extension
@staticmethod
def store_and_parse_file(file_data, extension) -> List[Document]:
def store_and_parse_file(file_name, file_data, extension) -> List[Document]:
# Store file to the private directory
os.makedirs(PrivateFileService.PRIVATE_STORE_PATH, exist_ok=True)
# random file name
file_name = f"{uuid4().hex}{extension}"
file_path = Path(os.path.join(PrivateFileService.PRIVATE_STORE_PATH, file_name))
# write file
@@ -78,25 +73,36 @@ class PrivateFileService:
return documents
@staticmethod
def process_file(base64_content: str) -> List[str]:
def process_file(file_name: str, base64_content: str, params: Any) -> List[str]:
file_data, extension = PrivateFileService.preprocess_base64_file(base64_content)
documents = PrivateFileService.store_and_parse_file(file_data, extension)
# Only process nodes, no store the index
pipeline = IngestionPipeline()
nodes = pipeline.run(documents=documents)
# Add the nodes to the index and persist it
current_index = get_index()
current_index = get_index(params)
# Insert the documents into the index
if isinstance(current_index, LlamaCloudIndex):
# LlamaCloudIndex is a managed index so we don't need to process the nodes
# just insert the documents
for doc in documents:
current_index.insert(doc)
from app.engine.service import LLamaCloudFileService
project_id = current_index._get_project_id()
pipeline_id = current_index._get_pipeline_id()
# LlamaCloudIndex is a managed index so we can directly use the files
upload_file = (file_name, BytesIO(file_data))
return [
LLamaCloudFileService.add_file_to_pipeline(
project_id,
pipeline_id,
upload_file,
custom_metadata={
# Set private=true to mark the document as private user docs (required for filtering)
"private": "true",
},
)
]
else:
# Only process nodes, no store the index
# First process documents into nodes
documents = PrivateFileService.store_and_parse_file(
file_name, file_data, extension
)
pipeline = IngestionPipeline()
nodes = pipeline.run(documents=documents)
@@ -109,5 +115,5 @@ class PrivateFileService:
persist_dir=os.environ.get("STORAGE_DIR", "storage")
)
# Return the document ids
return [doc.doc_id for doc in documents]
# Return the document ids
return [doc.doc_id for doc in documents]
@@ -1,88 +0,0 @@
import logging
import os
from typing import Any, Dict, List, Optional
import requests
from app.api.routers.models import LlamaCloudFile
logger = logging.getLogger("uvicorn")
class LLamaCloudFileService:
LLAMA_CLOUD_URL = "https://cloud.llamaindex.ai/api/v1"
LOCAL_STORE_PATH = "output/llamacloud"
DOWNLOAD_FILE_NAME_TPL = "{pipeline_id}${filename}"
@classmethod
def _get_files(cls, pipeline_id: str) -> List[Dict[str, Any]]:
url = f"{cls.LLAMA_CLOUD_URL}/pipelines/{pipeline_id}/files"
return cls._make_request(url)
@classmethod
def _get_file_detail(cls, project_id: str, file_id: str) -> Dict[str, Any]:
url = f"{cls.LLAMA_CLOUD_URL}/files/{file_id}/content?project_id={project_id}"
return cls._make_request(url)
@classmethod
def _download_file(cls, url: str, local_file_path: str):
logger.info(f"Downloading file to {local_file_path}")
# Create directory if it doesn't exist
os.makedirs(cls.LOCAL_STORE_PATH, exist_ok=True)
# Download the file
with requests.get(url, stream=True) as r:
r.raise_for_status()
with open(local_file_path, "wb") as f:
for chunk in r.iter_content(chunk_size=8192):
f.write(chunk)
logger.info("File downloaded successfully")
@classmethod
def download_llamacloud_pipeline_file(
cls,
file: LlamaCloudFile,
force_download: bool = False,
):
file_name = file.file_name
pipeline_id = file.pipeline_id
# Check is the file already exists
downloaded_file_path = cls.get_file_path(file_name, pipeline_id)
if os.path.exists(downloaded_file_path) and not force_download:
logger.debug(f"File {file_name} already exists in local storage")
return
try:
logger.info(f"Downloading file {file_name} for pipeline {pipeline_id}")
files = cls._get_files(pipeline_id)
if not files or not isinstance(files, list):
raise Exception("No files found in LlamaCloud")
for file_entry in files:
if file_entry["name"] == file_name:
file_id = file_entry["file_id"]
project_id = file_entry["project_id"]
file_detail = cls._get_file_detail(project_id, file_id)
cls._download_file(file_detail["url"], downloaded_file_path)
break
except Exception as error:
logger.info(f"Error fetching file from LlamaCloud: {error}")
@classmethod
def get_file_name(cls, name: str, pipeline_id: str) -> str:
return cls.DOWNLOAD_FILE_NAME_TPL.format(pipeline_id=pipeline_id, filename=name)
@classmethod
def get_file_path(cls, name: str, pipeline_id: str) -> str:
return os.path.join(cls.LOCAL_STORE_PATH, cls.get_file_name(name, pipeline_id))
@staticmethod
def _make_request(
url: str, data=None, headers: Optional[Dict] = None, method: str = "get"
):
if headers is None:
headers = {
"Accept": "application/json",
"Authorization": f'Bearer {os.getenv("LLAMA_CLOUD_API_KEY")}',
}
response = requests.request(method, url, headers=headers, data=data)
response.raise_for_status()
return response.json()
@@ -1,3 +1,4 @@
import logging
from typing import List
from app.api.routers.models import Message
@@ -9,11 +10,14 @@ NEXT_QUESTIONS_SUGGESTION_PROMPT = PromptTemplate(
"You're a helpful assistant! Your task is to suggest the next question that user might ask. "
"\nHere is the conversation history"
"\n---------------------\n{conversation}\n---------------------"
"Given the conversation history, please give me $number_of_questions questions that you might ask next!"
"Given the conversation history, please give me {number_of_questions} questions that you might ask next!"
)
N_QUESTION_TO_GENERATE = 3
logger = logging.getLogger("uvicorn")
class NextQuestions(BaseModel):
"""A list of questions that user might ask next"""
@@ -26,23 +30,31 @@ class NextQuestionSuggestion:
messages: List[Message],
number_of_questions: int = N_QUESTION_TO_GENERATE,
) -> List[str]:
# Reduce the cost by only using the last two messages
last_user_message = None
last_assistant_message = None
for message in reversed(messages):
if message.role == "user":
last_user_message = f"User: {message.content}"
elif message.role == "assistant":
last_assistant_message = f"Assistant: {message.content}"
if last_user_message and last_assistant_message:
break
conversation: str = f"{last_user_message}\n{last_assistant_message}"
"""
Suggest the next questions that user might ask based on the conversation history
Return as empty list if there is an error
"""
try:
# Reduce the cost by only using the last two messages
last_user_message = None
last_assistant_message = None
for message in reversed(messages):
if message.role == "user":
last_user_message = f"User: {message.content}"
elif message.role == "assistant":
last_assistant_message = f"Assistant: {message.content}"
if last_user_message and last_assistant_message:
break
conversation: str = f"{last_user_message}\n{last_assistant_message}"
output: NextQuestions = await Settings.llm.astructured_predict(
NextQuestions,
prompt=NEXT_QUESTIONS_SUGGESTION_PROMPT,
conversation=conversation,
nun_questions=number_of_questions,
)
output: NextQuestions = await Settings.llm.astructured_predict(
NextQuestions,
prompt=NEXT_QUESTIONS_SUGGESTION_PROMPT,
conversation=conversation,
number_of_questions=number_of_questions,
)
return output.questions
return output.questions
except Exception as e:
logger.error(f"Error when generating next question: {e}")
return []
@@ -1,3 +1,4 @@
# flake8: noqa: E402
from dotenv import load_dotenv
load_dotenv()
@@ -21,7 +22,6 @@ STORAGE_DIR = os.getenv("STORAGE_DIR", "storage")
def get_doc_store():
# If the storage directory is there, load the document store from it.
# If not, set up an in-memory document store since we can't load from a directory that doesn't exist.
if os.path.exists(STORAGE_DIR):
@@ -6,7 +6,7 @@ from app.engine.vectordb import get_vector_store
logger = logging.getLogger("uvicorn")
def get_index():
def get_index(params=None):
logger.info("Connecting vector store...")
store = get_vector_store()
# Load the index from the vector store
@@ -0,0 +1,34 @@
from llama_index.core.vector_stores.types import MetadataFilter, MetadataFilters
def generate_filters(doc_ids):
"""
Generate public/private document filters based on the doc_ids and the vector store.
"""
public_doc_filter = MetadataFilter(
key="private",
value="true",
operator="!=", # type: ignore
)
selected_doc_filter = MetadataFilter(
key="doc_id",
value=doc_ids,
operator="in", # type: ignore
)
if len(doc_ids) > 0:
# If doc_ids are provided, we will select both public and selected documents
filters = MetadataFilters(
filters=[
public_doc_filter,
selected_doc_filter,
],
condition="or", # type: ignore
)
else:
filters = MetadataFilters(
filters=[
public_doc_filter,
]
)
return filters
+8 -5
View File
@@ -1,20 +1,22 @@
# flake8: noqa: E402
from dotenv import load_dotenv
load_dotenv()
import logging
import os
import uvicorn
from app.api.routers.chat import chat_router
from app.api.routers.chat_config import config_router
from app.api.routers.upload import file_upload_router
from app.observability import init_observability
from app.settings import init_settings
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import RedirectResponse
from app.api.routers.chat import chat_router
from app.api.routers.upload import file_upload_router
from app.settings import init_settings
from app.observability import init_observability
from fastapi.staticfiles import StaticFiles
app = FastAPI()
init_settings()
@@ -54,6 +56,7 @@ mount_static_files("data", "/api/files/data")
mount_static_files("output", "/api/files/output")
app.include_router(chat_router, prefix="/api/chat")
app.include_router(config_router, prefix="/api/chat/config")
app.include_router(file_upload_router, prefix="/api/chat/upload")
if __name__ == "__main__":
@@ -9,7 +9,7 @@ readme = "README.md"
generate = "app.engine.generate:generate_datasource"
[tool.poetry.dependencies]
python = "^3.11,<3.12"
python = "^3.11,<4.0"
fastapi = "^0.109.1"
uvicorn = { extras = ["standard"], version = "^0.23.2" }
python-dotenv = "^1.0.0"
@@ -0,0 +1,16 @@
import { NextResponse } from "next/server";
import { LLamaCloudFileService } from "../../llamaindex/streaming/service";
/**
* This API is to get config from the backend envs and expose them to the frontend
*/
export async function GET() {
const config = {
projects: await LLamaCloudFileService.getAllProjectsWithPipelines(),
pipeline: {
pipeline: process.env.LLAMA_CLOUD_INDEX_NAME,
project: process.env.LLAMA_CLOUD_PROJECT_NAME,
},
};
return NextResponse.json(config, { status: 200 });
}
@@ -0,0 +1,25 @@
import { MetadataFilter, MetadataFilters } from "llamaindex";
export function generateFilters(documentIds: string[]): MetadataFilters {
// filter all documents have the private metadata key set to true
const publicDocumentsFilter: MetadataFilter = {
key: "private",
value: "true",
operator: "!=",
};
// if no documentIds are provided, only retrieve information from public documents
if (!documentIds.length) return { filters: [publicDocumentsFilter] };
const privateDocumentsFilter: MetadataFilter = {
key: "doc_id",
value: documentIds,
operator: "in",
};
// if documentIds are provided, retrieve information from public and private documents
return {
filters: [publicDocumentsFilter, privateDocumentsFilter],
condition: "or",
};
}
@@ -27,7 +27,7 @@ export async function POST(request: NextRequest) {
try {
const body = await request.json();
const { messages }: { messages: Message[] } = body;
const { messages, data }: { messages: Message[]; data?: any } = body;
const userMessage = messages.pop();
if (!messages || !userMessage || userMessage.role !== "user") {
return NextResponse.json(
@@ -59,7 +59,7 @@ export async function POST(request: NextRequest) {
},
);
const ids = retrieveDocumentIds(allAnnotations);
const chatEngine = await createChatEngine(ids);
const chatEngine = await createChatEngine(ids, data);
// Convert message content from Vercel/AI format to LlamaIndex/OpenAI format
const userMessageContent = convertMessageContent(
@@ -1,4 +1,5 @@
import { NextRequest, NextResponse } from "next/server";
import { getDataSource } from "../engine";
import { initSettings } from "../engine/settings";
import { uploadDocument } from "../llamaindex/documents/upload";
@@ -9,14 +10,21 @@ export const dynamic = "force-dynamic";
export async function POST(request: NextRequest) {
try {
const { base64 }: { base64: string } = await request.json();
const { base64, params }: { base64: string; params?: any } =
await request.json();
if (!base64) {
return NextResponse.json(
{ error: "base64 is required in the request body" },
{ status: 400 },
);
}
return NextResponse.json(await uploadDocument(base64));
const index = await getDataSource(params);
if (!index) {
throw new Error(
`StorageContext is empty - call 'npm run generate' to generate the storage first`,
);
}
return NextResponse.json(await uploadDocument(index, base64));
} catch (error) {
console.error("[Upload API]", error);
return NextResponse.json(
@@ -1,11 +1,13 @@
"use client";
import { useChat } from "ai/react";
import { useState } from "react";
import { ChatInput, ChatMessages } from "./ui/chat";
import { useClientConfig } from "./ui/chat/hooks/use-config";
export default function ChatSection() {
const { backend } = useClientConfig();
const [requestData, setRequestData] = useState<any>();
const {
messages,
input,
@@ -17,6 +19,7 @@ export default function ChatSection() {
append,
setInput,
} = useChat({
body: { data: requestData },
api: `${backend}/api/chat`,
headers: {
"Content-Type": "application/json", // using JSON because of vercel/ai 2.2.26
@@ -45,6 +48,8 @@ export default function ChatSection() {
messages={messages}
append={append}
setInput={setInput}
requestParams={{ params: requestData }}
setRequestData={setRequestData}
/>
</div>
);
@@ -6,6 +6,7 @@ import { Input } from "../input";
import UploadImagePreview from "../upload-image-preview";
import { ChatHandler } from "./chat.interface";
import { useFile } from "./hooks/use-file";
import { LlamaCloudSelector } from "./widgets/LlamaCloudSelector";
const ALLOWED_EXTENSIONS = ["png", "jpg", "jpeg", "csv", "pdf", "txt", "docx"];
@@ -21,7 +22,10 @@ export default function ChatInput(
| "messages"
| "setInput"
| "append"
>,
> & {
requestParams?: any;
setRequestData?: React.Dispatch<any>;
},
) {
const {
imageUrl,
@@ -64,10 +68,11 @@ export default function ChatInput(
return;
}
try {
await uploadFile(file);
await uploadFile(file, props.requestParams);
props.onFileUpload?.(file);
} catch (error: any) {
props.onFileError?.(error.message);
const onFileUploadError = props.onFileError || window.alert;
onFileUploadError(error.message);
}
};
@@ -107,6 +112,10 @@ export default function ChatInput(
disabled: props.isLoading,
}}
/>
{process.env.NEXT_PUBLIC_USE_LLAMACLOUD === "true" &&
props.setRequestData && (
<LlamaCloudSelector setRequestData={props.setRequestData} />
)}
<Button type="submit" disabled={props.isLoading || !props.input.trim()}>
Send message
</Button>
@@ -1,123 +1,196 @@
import { Check, Copy } from "lucide-react";
import { Check, Copy, FileText } from "lucide-react";
import Image from "next/image";
import { useMemo } from "react";
import { Button } from "../../button";
import { FileIcon } from "../../document-preview";
import {
HoverCard,
HoverCardContent,
HoverCardTrigger,
} from "../../hover-card";
import { cn } from "../../lib/utils";
import { useCopyToClipboard } from "../hooks/use-copy-to-clipboard";
import { SourceData } from "../index";
import { DocumentFileType, SourceData, SourceNode } from "../index";
import PdfDialog from "../widgets/PdfDialog";
const SCORE_THRESHOLD = 0.3;
function SourceNumberButton({ index }: { index: number }) {
return (
<div className="text-xs w-5 h-5 rounded-full bg-gray-100 mb-2 flex items-center justify-center hover:text-white hover:bg-primary hover:cursor-pointer">
{index + 1}
</div>
);
}
type NodeInfo = {
id: string;
url?: string;
type Document = {
url: string;
sources: SourceNode[];
};
export function ChatSources({ data }: { data: SourceData }) {
const sources: NodeInfo[] = useMemo(() => {
// aggregate nodes by url or file_path (get the highest one by score)
const nodesByPath: { [path: string]: NodeInfo } = {};
const documents: Document[] = useMemo(() => {
// group nodes by document (a document must have a URL)
const nodesByUrl: Record<string, SourceNode[]> = {};
data.nodes.forEach((node) => {
const key = node.url;
nodesByUrl[key] ??= [];
nodesByUrl[key].push(node);
});
data.nodes
.filter((node) => (node.score ?? 1) > SCORE_THRESHOLD)
.sort((a, b) => (b.score ?? 1) - (a.score ?? 1))
.forEach((node) => {
const nodeInfo = {
id: node.id,
url: node.url,
};
const key = nodeInfo.url ?? nodeInfo.id; // use id as key for UNKNOWN type
if (!nodesByPath[key]) {
nodesByPath[key] = nodeInfo;
}
});
return Object.values(nodesByPath);
// convert to array of documents
return Object.entries(nodesByUrl).map(([url, sources]) => ({
url,
sources,
}));
}, [data.nodes]);
if (sources.length === 0) return null;
if (documents.length === 0) return null;
return (
<div className="space-x-2 text-sm">
<span className="font-semibold">Sources:</span>
<div className="inline-flex gap-1 items-center">
{sources.map((nodeInfo: NodeInfo, index: number) => {
if (nodeInfo.url?.endsWith(".pdf")) {
return (
<PdfDialog
key={nodeInfo.id}
documentId={nodeInfo.id}
url={nodeInfo.url!}
trigger={<SourceNumberButton index={index} />}
/>
);
}
return (
<div key={nodeInfo.id}>
<HoverCard>
<HoverCardTrigger>
<SourceNumberButton index={index} />
</HoverCardTrigger>
<HoverCardContent className="w-[320px]">
<NodeInfo nodeInfo={nodeInfo} />
</HoverCardContent>
</HoverCard>
</div>
);
<div className="space-y-2 text-sm">
<div className="font-semibold text-lg">Sources:</div>
<div className="flex gap-3 flex-wrap">
{documents.map((document) => {
return <DocumentInfo key={document.url} document={document} />;
})}
</div>
</div>
);
}
function NodeInfo({ nodeInfo }: { nodeInfo: NodeInfo }) {
const { isCopied, copyToClipboard } = useCopyToClipboard({ timeout: 1000 });
if (nodeInfo.url) {
// this is a node generated by the web loader or file loader,
// add a link to view its URL and a button to copy the URL to the clipboard
return (
<div className="flex items-center my-2">
<a
className="hover:text-blue-900 truncate"
href={nodeInfo.url}
target="_blank"
>
<span>{nodeInfo.url}</span>
</a>
<Button
onClick={() => copyToClipboard(nodeInfo.url!)}
size="icon"
variant="ghost"
className="h-12 w-12 shrink-0"
>
{isCopied ? (
<Check className="h-4 w-4" />
) : (
<Copy className="h-4 w-4" />
)}
</Button>
</div>
);
}
// node generated by unknown loader, implement renderer by analyzing logged out metadata
export function SourceInfo({
node,
index,
}: {
node?: SourceNode;
index: number;
}) {
if (!node) return <SourceNumberButton index={index} />;
return (
<p>
Sorry, unknown node type. Please add a new renderer in the NodeInfo
component.
</p>
<HoverCard>
<HoverCardTrigger
className="cursor-default"
onClick={(e) => {
e.preventDefault();
e.stopPropagation();
}}
>
<SourceNumberButton
index={index}
className="hover:text-white hover:bg-primary"
/>
</HoverCardTrigger>
<HoverCardContent className="w-[400px]">
<NodeInfo nodeInfo={node} />
</HoverCardContent>
</HoverCard>
);
}
export function SourceNumberButton({
index,
className,
}: {
index: number;
className?: string;
}) {
return (
<span
className={cn(
"text-xs w-5 h-5 rounded-full bg-gray-100 inline-flex items-center justify-center",
className,
)}
>
{index + 1}
</span>
);
}
function DocumentInfo({ document }: { document: Document }) {
if (!document.sources.length) return null;
const { url, sources } = document;
const fileName = sources[0].metadata.file_name as string | undefined;
const fileExt = fileName?.split(".").pop();
const fileImage = fileExt ? FileIcon[fileExt as DocumentFileType] : null;
const DocumentDetail = (
<div
key={url}
className="h-28 w-48 flex flex-col justify-between p-4 border rounded-md shadow-md cursor-pointer"
>
<p
title={fileName}
className={cn(
fileName ? "truncate" : "text-blue-900 break-words",
"text-left",
)}
>
{fileName ?? url}
</p>
<div className="flex justify-between items-center">
<div className="space-x-2 flex">
{sources.map((node: SourceNode, index: number) => {
return (
<div key={node.id}>
<SourceInfo node={node} index={index} />
</div>
);
})}
</div>
{fileImage ? (
<div className="relative h-8 w-8 shrink-0 overflow-hidden rounded-md">
<Image
className="h-full w-auto"
priority
src={fileImage}
alt="Icon"
/>
</div>
) : (
<FileText className="text-gray-500" />
)}
</div>
</div>
);
if (url.endsWith(".pdf")) {
// open internal pdf dialog for pdf files when click document card
return <PdfDialog documentId={url} url={url} trigger={DocumentDetail} />;
}
// open external link when click document card for other file types
return <div onClick={() => window.open(url, "_blank")}>{DocumentDetail}</div>;
}
function NodeInfo({ nodeInfo }: { nodeInfo: SourceNode }) {
const { isCopied, copyToClipboard } = useCopyToClipboard({ timeout: 1000 });
const pageNumber =
// XXX: page_label is used in Python, but page_number is used by Typescript
(nodeInfo.metadata?.page_number as number) ??
(nodeInfo.metadata?.page_label as number) ??
null;
return (
<div className="space-y-4">
<div className="flex justify-between items-center">
<span className="font-semibold">
{pageNumber ? `On page ${pageNumber}:` : "Node content:"}
</span>
{nodeInfo.text && (
<Button
onClick={(e) => {
e.stopPropagation();
copyToClipboard(nodeInfo.text);
}}
size="icon"
variant="ghost"
className="h-12 w-12 shrink-0"
>
{isCopied ? (
<Check className="h-4 w-4" />
) : (
<Copy className="h-4 w-4" />
)}
</Button>
)}
</div>
{nodeInfo.text && (
<pre className="max-h-[200px] overflow-auto whitespace-pre-line">
&ldquo;{nodeInfo.text}&rdquo;
</pre>
)}
</div>
);
}
@@ -11,10 +11,10 @@ import {
ImageData,
MessageAnnotation,
MessageAnnotationType,
SourceData,
SuggestedQuestionsData,
ToolData,
getAnnotationData,
getSourceAnnotationData,
} from "../index";
import ChatAvatar from "./chat-avatar";
import { ChatEvents } from "./chat-events";
@@ -54,10 +54,9 @@ function ChatMessageContent({
annotations,
MessageAnnotationType.EVENTS,
);
const sourceData = getAnnotationData<SourceData>(
annotations,
MessageAnnotationType.SOURCES,
);
const sourceData = getSourceAnnotationData(annotations);
const toolData = getAnnotationData<ToolData>(
annotations,
MessageAnnotationType.TOOLS,
@@ -91,7 +90,7 @@ function ChatMessageContent({
},
{
order: 0,
component: <Markdown content={message.content} />,
component: <Markdown content={message.content} sources={sourceData[0]} />,
},
{
order: 3,
@@ -5,6 +5,8 @@ import rehypeKatex from "rehype-katex";
import remarkGfm from "remark-gfm";
import remarkMath from "remark-math";
import { SourceData } from "..";
import { SourceNumberButton } from "./chat-sources";
import { CodeBlock } from "./codeblock";
const MemoizedReactMarkdown: FC<Options> = memo(
@@ -34,12 +36,48 @@ const preprocessMedia = (content: string) => {
return content.replace(/(sandbox|attachment|snt):/g, "");
};
const preprocessContent = (content: string) => {
return preprocessMedia(preprocessLaTeX(content));
/**
* Update the citation flag [citation:id]() to the new format [citation:index](url)
*/
const preprocessCitations = (content: string, sources?: SourceData) => {
if (sources) {
const citationRegex = /\[citation:(.+?)\]\(\)/g;
let match;
// Find all the citation references in the content
while ((match = citationRegex.exec(content)) !== null) {
const citationId = match[1];
// Find the source node with the id equal to the citation-id, also get the index of the source node
const sourceNode = sources.nodes.find((node) => node.id === citationId);
// If the source node is found, replace the citation reference with the new format
if (sourceNode !== undefined) {
content = content.replace(
match[0],
`[citation:${sources.nodes.indexOf(sourceNode)}]()`,
);
} else {
// If the source node is not found, remove the citation reference
content = content.replace(match[0], "");
}
}
}
return content;
};
export default function Markdown({ content }: { content: string }) {
const processedContent = preprocessContent(content);
const preprocessContent = (content: string, sources?: SourceData) => {
return preprocessCitations(
preprocessMedia(preprocessLaTeX(content)),
sources,
);
};
export default function Markdown({
content,
sources,
}: {
content: string;
sources?: SourceData;
}) {
const processedContent = preprocessContent(content, sources);
return (
<MemoizedReactMarkdown
@@ -80,6 +118,23 @@ export default function Markdown({ content }: { content: string }) {
/>
);
},
a({ href, children }) {
// If a text link starts with 'citation:', then render it as a citation reference
if (
Array.isArray(children) &&
typeof children[0] === "string" &&
children[0].startsWith("citation:")
) {
const index = Number(children[0].replace("citation:", ""));
if (!isNaN(index)) {
return <SourceNumberButton index={index} />;
} else {
// citation is not looked up yet, don't render anything
return <></>;
}
}
return <a href={href}>{children}</a>;
},
}}
>
{processedContent}
@@ -1,5 +1,5 @@
import { Loader2 } from "lucide-react";
import { useEffect, useRef } from "react";
import { useEffect, useRef, useState } from "react";
import { Button } from "../button";
import ChatActions from "./chat-actions";
@@ -13,7 +13,9 @@ export default function ChatMessages(
"messages" | "isLoading" | "reload" | "stop" | "append"
>,
) {
const { starterQuestions } = useClientConfig();
const { backend } = useClientConfig();
const [starterQuestions, setStarterQuestions] = useState<string[]>();
const scrollableChatContainerRef = useRef<HTMLDivElement>(null);
const messageLength = props.messages.length;
const lastMessage = props.messages[messageLength - 1];
@@ -40,6 +42,19 @@ export default function ChatMessages(
scrollToBottom();
}, [messageLength, lastMessage]);
useEffect(() => {
if (!starterQuestions) {
fetch(`${backend}/api/chat/config`)
.then((response) => response.json())
.then((data) => {
if (data?.starterQuestions) {
setStarterQuestions(data.starterQuestions);
}
})
.catch((error) => console.error("Error fetching config", error));
}
}, [starterQuestions, backend]);
return (
<div
className="flex-1 w-full rounded-xl bg-white p-4 shadow-xl relative overflow-y-auto"
@@ -1,31 +1,24 @@
"use client";
import { useEffect, useMemo, useState } from "react";
export interface ChatConfig {
backend?: string;
starterQuestions?: string[];
}
function getBackendOrigin(): string {
const chatAPI = process.env.NEXT_PUBLIC_CHAT_API;
if (chatAPI) {
return new URL(chatAPI).origin;
} else {
if (typeof window !== "undefined") {
// Use BASE_URL from window.ENV
return (window as any).ENV?.BASE_URL || "";
}
return "";
}
}
export function useClientConfig(): ChatConfig {
const chatAPI = process.env.NEXT_PUBLIC_CHAT_API;
const [config, setConfig] = useState<ChatConfig>();
const backendOrigin = useMemo(() => {
return chatAPI ? new URL(chatAPI).origin : "";
}, [chatAPI]);
const configAPI = `${backendOrigin}/api/chat/config`;
useEffect(() => {
fetch(configAPI)
.then((response) => response.json())
.then((data) => setConfig({ ...data, chatAPI }))
.catch((error) => console.error("Error fetching config", error));
}, [chatAPI, configAPI]);
return {
backend: backendOrigin,
starterQuestions: config?.starterQuestions,
backend: getBackendOrigin(),
};
}
@@ -48,7 +48,11 @@ export function useFile() {
files.length && setFiles([]);
};
const uploadContent = async (base64: string): Promise<string[]> => {
const uploadContent = async (
file: File,
requestParams: any = {},
): Promise<string[]> => {
const base64 = await readContent({ file, asUrl: true });
const uploadAPI = `${backend}/api/chat/upload`;
const response = await fetch(uploadAPI, {
method: "POST",
@@ -56,7 +60,9 @@ export function useFile() {
"Content-Type": "application/json",
},
body: JSON.stringify({
...requestParams,
base64,
filename: file.name,
}),
});
if (!response.ok) throw new Error("Failed to upload document.");
@@ -98,7 +104,7 @@ export function useFile() {
return content;
};
const uploadFile = async (file: File) => {
const uploadFile = async (file: File, requestParams: any = {}) => {
if (file.type.startsWith("image/")) {
const base64 = await readContent({ file, asUrl: true });
return setImageUrl(base64);
@@ -124,8 +130,7 @@ export function useFile() {
});
}
default: {
const base64 = await readContent({ file, asUrl: true });
const ids = await uploadContent(base64);
const ids = await uploadContent(file, requestParams);
return addDoc({
...newDoc,
content: {
@@ -1,4 +1,5 @@
import { JSONValue } from "ai";
import { isValidUrl } from "../lib/utils";
import ChatInput from "./chat-input";
import ChatMessages from "./chat-messages";
@@ -42,7 +43,7 @@ export type SourceNode = {
metadata: Record<string, unknown>;
score?: number;
text: string;
url?: string;
url: string;
};
export type SourceData = {
@@ -83,9 +84,41 @@ export type MessageAnnotation = {
data: AnnotationData;
};
const NODE_SCORE_THRESHOLD = 0.25;
export function getAnnotationData<T extends AnnotationData>(
annotations: MessageAnnotation[],
type: MessageAnnotationType,
): T[] {
return annotations.filter((a) => a.type === type).map((a) => a.data as T);
}
export function getSourceAnnotationData(
annotations: MessageAnnotation[],
): SourceData[] {
const data = getAnnotationData<SourceData>(
annotations,
MessageAnnotationType.SOURCES,
);
if (data.length > 0) {
const sourceData = data[0] as SourceData;
if (sourceData.nodes) {
sourceData.nodes = preprocessSourceNodes(sourceData.nodes);
}
}
return data;
}
function preprocessSourceNodes(nodes: SourceNode[]): SourceNode[] {
// Filter source nodes has lower score
nodes = nodes
.filter((node) => (node.score ?? 1) > NODE_SCORE_THRESHOLD)
.filter((node) => isValidUrl(node.url))
.sort((a, b) => (b.score ?? 1) - (a.score ?? 1))
.map((node) => {
// remove trailing slash for node url if exists
node.url = node.url.replace(/\/$/, "");
return node;
});
return nodes;
}

Some files were not shown because too many files have changed in this diff Show More