refactor: improve Ollama embedder, normalize model names, add error handling, update tests (#4403)
This commit is contained in:
@@ -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",
|
||||
);
|
||||
});
|
||||
});
|
||||
@@ -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
@@ -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",
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user