diff --git a/docs/components/llms/models/sarvam.mdx b/docs/components/llms/models/sarvam.mdx index 1eedee8ea..dcc701487 100644 --- a/docs/components/llms/models/sarvam.mdx +++ b/docs/components/llms/models/sarvam.mdx @@ -9,7 +9,8 @@ To use Sarvam AI's models, please set the `SARVAM_API_KEY` which you can get fro ## Usage -```python + +```python Python import os from mem0 import Memory @@ -34,8 +35,35 @@ messages = [ {"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."} ] m.add(messages, user_id="alex") + ``` +```typescript TypeScript +import { Memory } from 'mem0ai/oss'; + +const config = { + llm: { + provider: 'sarvam', + config: { + apiKey: process.env.SARVAM_API_KEY || '', + model: 'sarvam-m', + temperature: 0.7, + }, + }, +}; + +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 thriller movies but I love sci-fi movies."}, + {"role": "assistant", "content": "Got it! I'll avoid thriller recommendations and suggest sci-fi movies in the future."} +]; +await memory.add(messages, { userId: 'alex' }); +``` + + + ## Advanced Usage with Sarvam-Specific Features ```python diff --git a/mem0-ts/src/oss/src/llms/sarvam.ts b/mem0-ts/src/oss/src/llms/sarvam.ts new file mode 100644 index 000000000..1d3bda153 --- /dev/null +++ b/mem0-ts/src/oss/src/llms/sarvam.ts @@ -0,0 +1,52 @@ +import { OpenAILLM } from "./openai"; +import { LLMConfig, Message } from "../types"; +import { LLMResponse } from "./base"; + +/** + * Sarvam AI LLM provider. + * + * Sarvam's API is OpenAI-compatible, so this simply reuses {@link OpenAILLM} + * and overrides the connection defaults — mirroring `mem0/llms/sarvam.py` in the + * Python SDK. The API key resolves from `config.apiKey` or the `SARVAM_API_KEY` + * env var, and the base URL from `config.baseURL`, `SARVAM_API_BASE`, else + * `https://api.sarvam.ai/v1`. + */ +export class SarvamLLM extends OpenAILLM { + constructor(config: LLMConfig) { + const apiKey = config.apiKey || process.env.SARVAM_API_KEY; + if (!apiKey) { + throw new Error("Sarvam API key is required"); + } + super({ + ...config, + apiKey, + baseURL: + config.baseURL || + process.env.SARVAM_API_BASE || + "https://api.sarvam.ai/v1", + model: config.model || "sarvam-m", + }); + } + + async generateResponse( + messages: Message[], + responseFormat?: { type: string }, + tools?: any[], + ): Promise { + try { + return await super.generateResponse(messages, responseFormat, tools); + } catch (err) { + const message = err instanceof Error ? err.message : String(err); + throw new Error(`Sarvam LLM failed: ${message}`); + } + } + + async generateChat(messages: Message[]): Promise { + try { + return await super.generateChat(messages); + } catch (err) { + const message = err instanceof Error ? err.message : String(err); + throw new Error(`Sarvam LLM failed: ${message}`); + } + } +} diff --git a/mem0-ts/src/oss/src/utils/factory.ts b/mem0-ts/src/oss/src/utils/factory.ts index 908067b38..2e2920834 100644 --- a/mem0-ts/src/oss/src/utils/factory.ts +++ b/mem0-ts/src/oss/src/utils/factory.ts @@ -25,6 +25,7 @@ import { OllamaLLM } from "../llms/ollama"; import { LMStudioLLM } from "../llms/lmstudio"; import { DeepSeekLLM } from "../llms/deepseek"; import { XAILLM } from "../llms/xai"; +import { SarvamLLM } from "../llms/sarvam"; import { LiteLLM } from "../llms/litellm"; import { MiniMaxLLM } from "../llms/minimax"; import { TogetherLLM } from "../llms/together"; @@ -109,6 +110,8 @@ export class LLMFactory { return new DeepSeekLLM(config); case "xai": return new XAILLM(config); + case "sarvam": + return new SarvamLLM(config); case "litellm": return new LiteLLM(config); case "minimax": diff --git a/mem0-ts/src/oss/tests/factory.unit.test.ts b/mem0-ts/src/oss/tests/factory.unit.test.ts index 92ca560fb..c4d684106 100644 --- a/mem0-ts/src/oss/tests/factory.unit.test.ts +++ b/mem0-ts/src/oss/tests/factory.unit.test.ts @@ -107,6 +107,11 @@ jest.mock("../src/llms/xai", () => ({ .fn() .mockImplementation((config) => ({ type: "xai-llm", config })), })); +jest.mock("../src/llms/sarvam", () => ({ + SarvamLLM: jest + .fn() + .mockImplementation((config) => ({ type: "sarvam-llm", config })), +})); jest.mock("../src/llms/litellm", () => ({ LiteLLM: jest .fn() @@ -269,6 +274,7 @@ describe("LLMFactory", () => { ["lmstudio"], ["deepseek"], ["xai"], + ["sarvam"], ["litellm"], ["minimax"], ["together"], diff --git a/mem0-ts/src/oss/tests/sarvam.test.ts b/mem0-ts/src/oss/tests/sarvam.test.ts new file mode 100644 index 000000000..07a9a442b --- /dev/null +++ b/mem0-ts/src/oss/tests/sarvam.test.ts @@ -0,0 +1,153 @@ +/// +/** + * Sarvam LLM — unit tests (mocked OpenAI). + */ + +import { SarvamLLM } from "../src/llms/sarvam"; + +const mockCreate = jest.fn(); +const mockOpenAICtor = jest.fn(); + +jest.mock("openai", () => { + return jest.fn().mockImplementation((config) => { + mockOpenAICtor(config); + return { chat: { completions: { create: mockCreate } } }; + }); +}); + +describe("SarvamLLM (unit)", () => { + const ORIGINAL_ENV = process.env; + + beforeEach(() => { + jest.clearAllMocks(); + process.env = { ...ORIGINAL_ENV }; + delete process.env.SARVAM_API_KEY; + delete process.env.SARVAM_API_BASE; + mockCreate.mockResolvedValue({ + choices: [{ message: { content: "hi", role: "assistant" } }], + }); + }); + + afterAll(() => { + process.env = ORIGINAL_ENV; + }); + + it("defaults to sarvam-m and the Sarvam base URL (matching the Python provider)", async () => { + const llm = new SarvamLLM({ apiKey: "test-key" }); + const result = await llm.generateResponse([ + { role: "user", content: "hello" }, + ]); + + expect(mockOpenAICtor).toHaveBeenCalledWith( + expect.objectContaining({ + apiKey: "test-key", + baseURL: "https://api.sarvam.ai/v1", + }), + ); + expect(mockCreate).toHaveBeenCalledWith( + expect.objectContaining({ model: "sarvam-m" }), + ); + expect(result).toBe("hi"); + }); + + it("resolves SARVAM_API_KEY / SARVAM_API_BASE from the environment", () => { + process.env.SARVAM_API_KEY = "env-key"; + process.env.SARVAM_API_BASE = "https://custom.sarvam.ai/v1"; + + new SarvamLLM({}); + + expect(mockOpenAICtor).toHaveBeenCalledWith( + expect.objectContaining({ + apiKey: "env-key", + baseURL: "https://custom.sarvam.ai/v1", + }), + ); + }); + + it("prefers explicit config over defaults and the environment", async () => { + process.env.SARVAM_API_KEY = "env-key"; + + const llm = new SarvamLLM({ + apiKey: "explicit-key", + baseURL: "https://proxy.example.com/v1", + model: "sarvam-2b", + }); + await llm.generateResponse([{ role: "user", content: "hello" }]); + + expect(mockOpenAICtor).toHaveBeenCalledWith( + expect.objectContaining({ + apiKey: "explicit-key", + baseURL: "https://proxy.example.com/v1", + }), + ); + expect(mockCreate).toHaveBeenCalledWith( + expect.objectContaining({ model: "sarvam-2b" }), + ); + }); + + it("throws when no API key is provided", () => { + expect(() => new SarvamLLM({})).toThrow("Sarvam API key is required"); + }); + + it("generateResponse() handles tool calls", async () => { + mockCreate.mockResolvedValueOnce({ + choices: [ + { + message: { + content: "", + role: "assistant", + tool_calls: [ + { + function: { name: "get_weather", arguments: '{"city": "SF"}' }, + }, + ], + }, + }, + ], + }); + + const llm = new SarvamLLM({ apiKey: "test-key" }); + const result = await llm.generateResponse( + [{ role: "user", content: "weather?" }], + undefined, + [{ type: "function", function: { name: "get_weather" } }], + ); + + expect(result).toEqual({ + content: "", + role: "assistant", + toolCalls: [{ name: "get_weather", arguments: '{"city": "SF"}' }], + }); + }); + + it("generateResponse() wraps downstream errors with a Sarvam-specific message", async () => { + mockCreate.mockRejectedValueOnce(new Error("Connection refused")); + const llm = new SarvamLLM({ apiKey: "test-key" }); + + await expect( + llm.generateResponse([{ role: "user", content: "hi" }]), + ).rejects.toThrow("Sarvam LLM failed: Connection refused"); + }); + + it("generateChat() returns the LLMResponse shape", async () => { + mockCreate.mockResolvedValueOnce({ + choices: [{ message: { content: "I can help.", role: "assistant" } }], + }); + + const llm = new SarvamLLM({ apiKey: "test-key" }); + const result = await llm.generateChat([ + { role: "user", content: "help me" }, + ]); + + expect(result).toEqual({ content: "I can help.", role: "assistant" }); + }); + + it("generateChat() wraps downstream errors with a Sarvam-specific message", async () => { + mockCreate.mockRejectedValueOnce(new Error("Timeout")); + const llm = new SarvamLLM({ apiKey: "test-key" }); + + await expect( + llm.generateChat([{ role: "user", content: "hi" }]), + ).rejects.toThrow("Sarvam LLM failed: Timeout"); + }); +});