refactor: improve Ollama embedder, normalize model names, add error handling, update tests (#4403)

This commit is contained in:
Kartik
2026-03-18 23:34:35 +05:30
committed by GitHub
parent 577a5a2feb
commit 7539463f50
5 changed files with 179 additions and 14 deletions
+20 -7
View File
@@ -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<number[][]> {
@@ -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<boolean> {
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 });
}
@@ -0,0 +1,121 @@
/// <reference types="jest" />
/**
* 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",
);
});
});
+15 -3
View File
@@ -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]
+1 -1
View File
@@ -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",
+22 -3
View File
@@ -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")