feat(ts-sdk): add vLLM provider (#5805)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -25,7 +25,8 @@ description: "Configure vLLM as an LLM provider in Mem0 for high-performance loc
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
import os
|
||||
from mem0 import Memory
|
||||
|
||||
@@ -53,6 +54,46 @@ messages = [
|
||||
m.add(messages, user_id="alice", metadata={"category": "movies"})
|
||||
```
|
||||
|
||||
```typescript TypeScript
|
||||
import { Memory } from "mem0ai/oss";
|
||||
|
||||
const config = {
|
||||
llm: {
|
||||
provider: "vllm",
|
||||
config: {
|
||||
model: "Qwen/Qwen2.5-32B-Instruct",
|
||||
baseURL: "http://localhost:8000/v1",
|
||||
apiKey: process.env.VLLM_API_KEY || "vllm-api-key",
|
||||
temperature: 0.1,
|
||||
maxTokens: 2000,
|
||||
},
|
||||
},
|
||||
};
|
||||
|
||||
const memory = new Memory(config);
|
||||
const messages = [
|
||||
{
|
||||
role: "user",
|
||||
content: "I'm planning to watch a movie tonight. Any recommendations?",
|
||||
},
|
||||
{
|
||||
role: "assistant",
|
||||
content: "How about thriller movies? They can be quite engaging.",
|
||||
},
|
||||
{
|
||||
role: "user",
|
||||
content: "I'm not a big fan of thrillers, but I love sci-fi movies.",
|
||||
},
|
||||
{
|
||||
role: "assistant",
|
||||
content: "Got it! I'll avoid thrillers and suggest sci-fi movies instead.",
|
||||
},
|
||||
];
|
||||
await memory.add(messages, { userId: "alice", metadata: { category: "movies" } });
|
||||
```
|
||||
|
||||
</CodeGroup>
|
||||
|
||||
## Configuration Parameters
|
||||
|
||||
| Parameter | Description | Default | Environment Variable |
|
||||
|
||||
@@ -96,6 +96,8 @@ export class ConfigManager {
|
||||
config: (() => {
|
||||
const defaultConf = DEFAULT_MEMORY_CONFIG.llm.config;
|
||||
const userConf = userConfig.llm?.config;
|
||||
const provider =
|
||||
userConfig.llm?.provider || DEFAULT_MEMORY_CONFIG.llm.provider;
|
||||
let finalModel: string | any = defaultConf.model;
|
||||
|
||||
if (userConf?.model && typeof userConf.model === "object") {
|
||||
@@ -105,14 +107,18 @@ export class ConfigManager {
|
||||
}
|
||||
|
||||
// Normalize snake_case keys from Python SDK / OpenClaw configs
|
||||
const llmRaw = userConf as Record<string, unknown> | undefined;
|
||||
const llmBaseURL =
|
||||
userConf?.baseURL ??
|
||||
userConf?.vllmBaseURL ??
|
||||
(llmRaw?.vllm_base_url as string | undefined) ??
|
||||
((userConf as Record<string, unknown>)?.lmstudio_base_url as
|
||||
| string
|
||||
| undefined) ??
|
||||
userConf?.url ??
|
||||
defaultConf.baseURL;
|
||||
const llmRaw = userConf as Record<string, unknown> | undefined;
|
||||
(provider.toLowerCase() === "vllm"
|
||||
? undefined
|
||||
: defaultConf.baseURL);
|
||||
const temperature =
|
||||
userConf?.temperature ??
|
||||
(llmRaw?.temperature as number | undefined);
|
||||
|
||||
@@ -19,6 +19,7 @@ export * from "./llms/lmstudio";
|
||||
export * from "./llms/mistral";
|
||||
export * from "./llms/langchain";
|
||||
export * from "./llms/litellm";
|
||||
export * from "./llms/vllm";
|
||||
export * from "./vector_stores/base";
|
||||
export * from "./vector_stores/memory";
|
||||
export * from "./vector_stores/qdrant";
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
import { OpenAILLM } from "./openai";
|
||||
import { LLMConfig, Message } from "../types";
|
||||
import { LLMResponse } from "./base";
|
||||
|
||||
const DEFAULT_MODEL = "Qwen/Qwen2.5-32B-Instruct";
|
||||
const DEFAULT_API_KEY = "vllm-api-key";
|
||||
// Mirrors the Python provider default (mem0/configs/llms/vllm.py) and the docs table.
|
||||
const DEFAULT_BASE_URL = "http://localhost:8000/v1";
|
||||
|
||||
export class VllmLLM extends OpenAILLM {
|
||||
constructor(config: LLMConfig) {
|
||||
const baseURL =
|
||||
config.baseURL ||
|
||||
config.vllmBaseURL ||
|
||||
config.vllm_base_url ||
|
||||
config.url ||
|
||||
process.env.VLLM_BASE_URL ||
|
||||
DEFAULT_BASE_URL;
|
||||
|
||||
super({
|
||||
...config,
|
||||
apiKey: config.apiKey || process.env.VLLM_API_KEY || DEFAULT_API_KEY,
|
||||
baseURL,
|
||||
model: config.model || DEFAULT_MODEL,
|
||||
});
|
||||
}
|
||||
|
||||
async generateResponse(
|
||||
messages: Message[],
|
||||
responseFormat?: { type: string },
|
||||
tools?: any[],
|
||||
): Promise<string | LLMResponse> {
|
||||
try {
|
||||
return await super.generateResponse(messages, responseFormat, tools);
|
||||
} catch (err) {
|
||||
const message = err instanceof Error ? err.message : String(err);
|
||||
throw new Error(`vLLM LLM failed: ${message}`);
|
||||
}
|
||||
}
|
||||
|
||||
async generateChat(messages: Message[]): Promise<LLMResponse> {
|
||||
try {
|
||||
return await super.generateChat(messages);
|
||||
} catch (err) {
|
||||
const message = err instanceof Error ? err.message : String(err);
|
||||
throw new Error(`vLLM LLM failed: ${message}`);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,108 @@
|
||||
import OpenAI from "openai";
|
||||
import { VllmLLM } from "../llms/vllm";
|
||||
import { LLMFactory } from "../utils/factory";
|
||||
|
||||
const createMock = jest.fn();
|
||||
|
||||
jest.mock("openai", () => {
|
||||
return jest.fn().mockImplementation(() => ({
|
||||
chat: {
|
||||
completions: {
|
||||
create: createMock,
|
||||
},
|
||||
},
|
||||
}));
|
||||
});
|
||||
|
||||
const MockOpenAI = OpenAI as unknown as jest.Mock;
|
||||
|
||||
describe("VllmLLM", () => {
|
||||
beforeEach(() => {
|
||||
jest.clearAllMocks();
|
||||
delete process.env.VLLM_API_KEY;
|
||||
delete process.env.VLLM_BASE_URL;
|
||||
createMock.mockResolvedValue({
|
||||
choices: [
|
||||
{
|
||||
message: {
|
||||
content: "hello",
|
||||
role: "assistant",
|
||||
},
|
||||
},
|
||||
],
|
||||
});
|
||||
});
|
||||
|
||||
it("uses vLLM defaults with a configured baseURL", async () => {
|
||||
const llm = new VllmLLM({ baseURL: "http://localhost:8000/v1" });
|
||||
|
||||
await llm.generateChat([{ role: "user", content: "Hi" }]);
|
||||
|
||||
expect(MockOpenAI).toHaveBeenCalledWith({
|
||||
apiKey: "vllm-api-key",
|
||||
baseURL: "http://localhost:8000/v1",
|
||||
});
|
||||
expect(createMock).toHaveBeenCalledWith({
|
||||
messages: [{ role: "user", content: "Hi" }],
|
||||
model: "Qwen/Qwen2.5-32B-Instruct",
|
||||
});
|
||||
});
|
||||
|
||||
it("reads vLLM API settings from the environment", () => {
|
||||
process.env.VLLM_API_KEY = "env-key";
|
||||
process.env.VLLM_BASE_URL = "http://vllm.example/v1";
|
||||
|
||||
new VllmLLM({});
|
||||
|
||||
expect(MockOpenAI).toHaveBeenCalledWith({
|
||||
apiKey: "env-key",
|
||||
baseURL: "http://vllm.example/v1",
|
||||
});
|
||||
});
|
||||
|
||||
it("accepts Python-style vllm_base_url configs", () => {
|
||||
new VllmLLM({ vllm_base_url: "http://localhost:8001/v1" });
|
||||
|
||||
expect(MockOpenAI).toHaveBeenCalledWith({
|
||||
apiKey: "vllm-api-key",
|
||||
baseURL: "http://localhost:8001/v1",
|
||||
});
|
||||
});
|
||||
|
||||
it("registers the vllm provider with LLMFactory", () => {
|
||||
const llm = LLMFactory.create("vllm", {
|
||||
baseURL: "http://localhost:8000/v1",
|
||||
});
|
||||
|
||||
expect(llm).toBeInstanceOf(VllmLLM);
|
||||
});
|
||||
|
||||
it("defaults to the local vLLM server when no baseURL is provided", () => {
|
||||
new VllmLLM({});
|
||||
|
||||
expect(MockOpenAI).toHaveBeenCalledWith({
|
||||
apiKey: "vllm-api-key",
|
||||
baseURL: "http://localhost:8000/v1",
|
||||
});
|
||||
});
|
||||
|
||||
it("uses VLLM_BASE_URL when config merging leaves baseURL unset", () => {
|
||||
process.env.VLLM_BASE_URL = "http://env-vllm.example/v1";
|
||||
|
||||
new VllmLLM({ model: "Qwen/Qwen2.5-32B-Instruct" });
|
||||
|
||||
expect(MockOpenAI).toHaveBeenCalledWith({
|
||||
apiKey: "vllm-api-key",
|
||||
baseURL: "http://env-vllm.example/v1",
|
||||
});
|
||||
});
|
||||
|
||||
it("wraps OpenAI-compatible client errors with provider context", async () => {
|
||||
createMock.mockRejectedValueOnce(new Error("network down"));
|
||||
const llm = new VllmLLM({ baseURL: "http://localhost:8000/v1" });
|
||||
|
||||
await expect(
|
||||
llm.generateResponse([{ role: "user", content: "Hi" }]),
|
||||
).rejects.toThrow("vLLM LLM failed: network down");
|
||||
});
|
||||
});
|
||||
@@ -45,6 +45,8 @@ export interface HistoryStoreConfig {
|
||||
export interface LLMConfig {
|
||||
provider?: string;
|
||||
baseURL?: string;
|
||||
vllmBaseURL?: string;
|
||||
vllm_base_url?: string;
|
||||
url?: string;
|
||||
config?: Record<string, any>;
|
||||
apiKey?: string;
|
||||
@@ -135,6 +137,8 @@ export const MemoryConfigSchema = z.object({
|
||||
model: z.union([z.string(), z.any()]).optional(),
|
||||
modelProperties: z.record(z.string(), z.any()).optional(),
|
||||
baseURL: z.string().optional(),
|
||||
vllmBaseURL: z.string().optional(),
|
||||
vllm_base_url: z.string().optional(),
|
||||
url: z.string().optional(),
|
||||
timeout: z.number().optional(),
|
||||
temperature: z.number().optional(),
|
||||
|
||||
@@ -25,6 +25,7 @@ import { LMStudioLLM } from "../llms/lmstudio";
|
||||
import { DeepSeekLLM } from "../llms/deepseek";
|
||||
import { LiteLLM } from "../llms/litellm";
|
||||
import { MiniMaxLLM } from "../llms/minimax";
|
||||
import { VllmLLM } from "../llms/vllm";
|
||||
import { SupabaseDB } from "../vector_stores/supabase";
|
||||
import { SQLiteManager } from "../storage/SQLiteManager";
|
||||
import { MemoryHistoryManager } from "../storage/MemoryHistoryManager";
|
||||
@@ -93,6 +94,8 @@ export class LLMFactory {
|
||||
return new LiteLLM(config);
|
||||
case "minimax":
|
||||
return new MiniMaxLLM(config);
|
||||
case "vllm":
|
||||
return new VllmLLM(config);
|
||||
default:
|
||||
throw new Error(`Unsupported LLM provider: ${provider}`);
|
||||
}
|
||||
|
||||
@@ -155,6 +155,35 @@ describe("ConfigManager", () => {
|
||||
expect(config.llm.config.baseURL).toBe("https://api.openai.com/v1");
|
||||
});
|
||||
|
||||
it("normalizes vllm_base_url to baseURL for vLLM", () => {
|
||||
const config = ConfigManager.mergeConfig({
|
||||
embedder: baseEmbedder,
|
||||
vectorStore: baseVectorStore,
|
||||
llm: {
|
||||
provider: "vllm",
|
||||
config: {
|
||||
model: "Qwen/Qwen2.5-32B-Instruct",
|
||||
vllm_base_url: "http://localhost:8000/v1",
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
expect(config.llm.config.baseURL).toBe("http://localhost:8000/v1");
|
||||
});
|
||||
|
||||
it("does not inject the OpenAI baseURL default for vLLM", () => {
|
||||
const config = ConfigManager.mergeConfig({
|
||||
embedder: baseEmbedder,
|
||||
vectorStore: baseVectorStore,
|
||||
llm: {
|
||||
provider: "vllm",
|
||||
config: { model: "Qwen/Qwen2.5-32B-Instruct" },
|
||||
},
|
||||
});
|
||||
|
||||
expect(config.llm.config.baseURL).toBeUndefined();
|
||||
});
|
||||
|
||||
it("should preserve url in embedder config (existing behavior)", () => {
|
||||
const config = ConfigManager.mergeConfig({
|
||||
embedder: {
|
||||
|
||||
@@ -102,6 +102,11 @@ jest.mock("../src/llms/minimax", () => ({
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "minimax-llm", config })),
|
||||
}));
|
||||
jest.mock("../src/llms/vllm", () => ({
|
||||
VllmLLM: jest
|
||||
.fn()
|
||||
.mockImplementation((config) => ({ type: "vllm-llm", config })),
|
||||
}));
|
||||
|
||||
jest.mock("../src/vector_stores/qdrant", () => ({
|
||||
Qdrant: jest
|
||||
@@ -228,6 +233,7 @@ describe("LLMFactory", () => {
|
||||
["deepseek"],
|
||||
["litellm"],
|
||||
["minimax"],
|
||||
["vllm"],
|
||||
])("creates LLM for provider '%s'", (provider) => {
|
||||
expect(() => LLMFactory.create(provider, dummyLLMConfig)).not.toThrow();
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user