From 7539463f509ff88f46abe04c0653483f8379cccc Mon Sep 17 00:00:00 2001 From: Kartik Date: Wed, 18 Mar 2026 23:34:35 +0530 Subject: [PATCH] refactor: improve Ollama embedder, normalize model names, add error handling, update tests (#4403) --- mem0-ts/src/oss/src/embeddings/ollama.ts | 27 +++- mem0-ts/src/oss/tests/ollama-embedder.test.ts | 121 ++++++++++++++++++ mem0/embeddings/ollama.py | 18 ++- pyproject.toml | 2 +- tests/embeddings/test_ollama_embeddings.py | 25 +++- 5 files changed, 179 insertions(+), 14 deletions(-) create mode 100644 mem0-ts/src/oss/tests/ollama-embedder.test.ts diff --git a/mem0-ts/src/oss/src/embeddings/ollama.ts b/mem0-ts/src/oss/src/embeddings/ollama.ts index 348f6cf98..51d339ce2 100644 --- a/mem0-ts/src/oss/src/embeddings/ollama.ts +++ b/mem0-ts/src/oss/src/embeddings/ollama.ts @@ -27,14 +27,18 @@ export class OllamaEmbedder implements Embedder { } catch (err) { logger.error(`Error ensuring model exists: ${err}`); } - // Ollama's Go server requires prompt to be a string. Coerce defensively - // since callers may pass values parsed from untrusted LLM JSON output. - const prompt = typeof text === "string" ? text : JSON.stringify(text); - const response = await this.ollama.embeddings({ + // Coerce defensively since callers may pass values parsed from untrusted LLM JSON output. + const input = typeof text === "string" ? text : JSON.stringify(text); + const response = await this.ollama.embed({ model: this.model, - prompt, + input, }); - return response.embedding; + if (!response.embeddings || response.embeddings.length === 0) { + throw new Error( + `Ollama embed() returned no embeddings for model '${this.model}'`, + ); + } + return response.embeddings[0]; } async embedBatch(texts: string[]): Promise { @@ -42,12 +46,21 @@ export class OllamaEmbedder implements Embedder { return response; } + private static normalizeModelName(name: string): string { + return name.includes(":") ? name : `${name}:latest`; + } + private async ensureModelExists(): Promise { if (this.initialized) { return true; } const local_models = await this.ollama.list(); - if (!local_models.models.find((m: any) => m.name === this.model)) { + const target = OllamaEmbedder.normalizeModelName(this.model); + if ( + !local_models.models.find( + (m: any) => OllamaEmbedder.normalizeModelName(m.name) === target, + ) + ) { logger.info(`Pulling model ${this.model}...`); await this.ollama.pull({ model: this.model }); } diff --git a/mem0-ts/src/oss/tests/ollama-embedder.test.ts b/mem0-ts/src/oss/tests/ollama-embedder.test.ts new file mode 100644 index 000000000..1854cf41f --- /dev/null +++ b/mem0-ts/src/oss/tests/ollama-embedder.test.ts @@ -0,0 +1,121 @@ +/// +/** + * Ollama Embedder — unit tests (mocked Ollama client). + */ + +import { OllamaEmbedder } from "../src/embeddings/ollama"; + +const mockEmbedding = [0.1, 0.2, 0.3, 0.4, 0.5]; +const mockEmbed = jest.fn().mockResolvedValue({ + model: "nomic-embed-text:latest", + embeddings: [mockEmbedding], +}); +const mockList = jest.fn().mockResolvedValue({ + models: [{ name: "nomic-embed-text:latest" }], +}); +const mockPull = jest.fn().mockResolvedValue({}); + +jest.mock("ollama", () => ({ + Ollama: jest.fn().mockImplementation(() => ({ + embed: mockEmbed, + list: mockList, + pull: mockPull, + })), +})); + +describe("OllamaEmbedder (unit)", () => { + beforeEach(() => { + mockEmbed.mockClear(); + mockList.mockClear(); + mockPull.mockClear(); + }); + + it("embed() calls ollama.embed with model and input, returns first embedding", async () => { + const embedder = new OllamaEmbedder({ + model: "nomic-embed-text:latest", + }); + + const result = await embedder.embed("Sample text to embed."); + + expect(mockEmbed).toHaveBeenCalledTimes(1); + expect(mockEmbed.mock.calls[0][0]).toEqual({ + model: "nomic-embed-text:latest", + input: "Sample text to embed.", + }); + expect(result).toEqual(mockEmbedding); + }); + + it("embed() coerces non-string input to JSON string", async () => { + const embedder = new OllamaEmbedder({ + model: "nomic-embed-text:latest", + }); + + // Force a non-string through the type boundary + await embedder.embed(42 as any); + + expect(mockEmbed.mock.calls[0][0].input).toBe("42"); + }); + + it("embedBatch() returns vectors for multiple inputs", async () => { + const embedder = new OllamaEmbedder({ + model: "nomic-embed-text:latest", + }); + + const result = await embedder.embedBatch(["text1", "text2"]); + + expect(mockEmbed).toHaveBeenCalledTimes(2); + expect(result).toEqual([mockEmbedding, mockEmbedding]); + }); + + it("ensureModelExists() does not pull when model is already present", async () => { + const embedder = new OllamaEmbedder({ + model: "nomic-embed-text:latest", + }); + + await embedder.embed("trigger ensureModelExists"); + + expect(mockList).toHaveBeenCalled(); + expect(mockPull).not.toHaveBeenCalled(); + }); + + it("ensureModelExists() pulls model when not found locally", async () => { + mockList.mockResolvedValueOnce({ models: [] }); + + const embedder = new OllamaEmbedder({ + model: "nomic-embed-text:latest", + }); + + await embedder.embed("trigger ensureModelExists"); + + expect(mockPull).toHaveBeenCalledWith({ model: "nomic-embed-text:latest" }); + }); + + it("ensureModelExists() normalizes model name with :latest tag", async () => { + mockList.mockResolvedValue({ + models: [{ name: "nomic-embed-text:latest" }], + }); + + const embedder = new OllamaEmbedder({ + model: "nomic-embed-text", + }); + + await embedder.embed("trigger ensureModelExists"); + + expect(mockPull).not.toHaveBeenCalled(); + }); + + it("embed() throws when embeddings array is empty", async () => { + mockEmbed.mockResolvedValueOnce({ + model: "nomic-embed-text:latest", + embeddings: [], + }); + + const embedder = new OllamaEmbedder({ + model: "nomic-embed-text:latest", + }); + + await expect(embedder.embed("text")).rejects.toThrow( + "Ollama embed() returned no embeddings", + ); + }); +}); diff --git a/mem0/embeddings/ollama.py b/mem0/embeddings/ollama.py index 49b7c2e94..07149f2c8 100644 --- a/mem0/embeddings/ollama.py +++ b/mem0/embeddings/ollama.py @@ -31,12 +31,21 @@ class OllamaEmbedding(EmbeddingBase): self.client = Client(host=self.config.ollama_base_url) self._ensure_model_exists() + @staticmethod + def _normalize_model_name(name: str) -> str: + return name if ":" in name else f"{name}:latest" + def _ensure_model_exists(self): """ Ensure the specified model exists locally. If not, pull it from Ollama. """ local_models = self.client.list()["models"] - if not any(model.get("name") == self.config.model or model.get("model") == self.config.model for model in local_models): + target = self._normalize_model_name(self.config.model) + if not any( + self._normalize_model_name(model.get("name", "")) == target + or self._normalize_model_name(model.get("model", "")) == target + for model in local_models + ): self.client.pull(self.config.model) def embed(self, text, memory_action: Optional[Literal["add", "search", "update"]] = None): @@ -49,5 +58,8 @@ class OllamaEmbedding(EmbeddingBase): Returns: list: The embedding vector. """ - response = self.client.embeddings(model=self.config.model, prompt=text) - return response["embedding"] + response = self.client.embed(model=self.config.model, input=text) + embeddings = response.get("embeddings") or [] + if not embeddings: + raise ValueError(f"Ollama embed() returned no embeddings for model '{self.config.model}'") + return embeddings[0] diff --git a/pyproject.toml b/pyproject.toml index 3053ecf69..90a72a432 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -62,7 +62,7 @@ llms = [ "together>=0.2.10", "litellm>=1.74.0", "openai>=1.90.0", - "ollama>=0.1.0", + "ollama>=0.3.0", "vertexai>=0.1.0", "google-generativeai>=0.3.0", "google-genai>=1.0.0", diff --git a/tests/embeddings/test_ollama_embeddings.py b/tests/embeddings/test_ollama_embeddings.py index 3e2cc6723..e0bf9193b 100644 --- a/tests/embeddings/test_ollama_embeddings.py +++ b/tests/embeddings/test_ollama_embeddings.py @@ -19,13 +19,13 @@ def test_embed_text(mock_ollama_client): config = BaseEmbedderConfig(model="nomic-embed-text", embedding_dims=512) embedder = OllamaEmbedding(config) - mock_response = {"embedding": [0.1, 0.2, 0.3, 0.4, 0.5]} - mock_ollama_client.embeddings.return_value = mock_response + mock_response = {"embeddings": [[0.1, 0.2, 0.3, 0.4, 0.5]]} + mock_ollama_client.embed.return_value = mock_response text = "Sample text to embed." embedding = embedder.embed(text) - mock_ollama_client.embeddings.assert_called_once_with(model="nomic-embed-text", prompt=text) + mock_ollama_client.embed.assert_called_once_with(model="nomic-embed-text", input=text) assert embedding == [0.1, 0.2, 0.3, 0.4, 0.5] @@ -41,3 +41,22 @@ def test_ensure_model_exists(mock_ollama_client): embedder._ensure_model_exists() mock_ollama_client.pull.assert_called_once_with("nomic-embed-text") + + +def test_ensure_model_exists_normalizes_latest_tag(mock_ollama_client): + """Model 'nomic-embed-text' should match 'nomic-embed-text:latest' from ollama list.""" + mock_ollama_client.list.return_value = {"models": [{"name": "nomic-embed-text:latest"}]} + config = BaseEmbedderConfig(model="nomic-embed-text", embedding_dims=512) + OllamaEmbedding(config) + + mock_ollama_client.pull.assert_not_called() + + +def test_embed_empty_response_raises(mock_ollama_client): + config = BaseEmbedderConfig(model="nomic-embed-text", embedding_dims=512) + embedder = OllamaEmbedding(config) + + mock_ollama_client.embed.return_value = {"embeddings": []} + + with pytest.raises(ValueError, match="returned no embeddings"): + embedder.embed("some text")