Files
library-desk/src/clients/ollama_client.py
T
jpmschweitzerandClaude Fable 5 8348b4bf92 perf: batch page embeddings via /api/embed and upsert before pruning stale points
- embed_batch now issues one batched /api/embed request (the old loop
  made one /api/embeddings round-trip per chunk) with a per-text
  fallback that preserves None-for-failed semantics
- update_from_page embeds all chunks in that single call and stores
  them in one Qdrant batch upsert (upsert_points)
- reindex order reversed: upsert new points first, then prune stale ids
  (deterministic uuid5 ids make overwrite safe) so a mid-way failure no
  longer leaves the page with zero vectors
- VectorUpdateSummary gains status (success/partial/failed) and
  chunks_skipped; all-embeddings-failed keeps old vectors and reports
  failure instead of success=True

Measured on a 7-chunk page ingest (local server, llm_tester):
~375ms -> ~181ms median over 3 runs.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01QbFZyDvYksazX6nYQYZ67L
2026-07-14 13:09:37 +02:00

425 lines
13 KiB
Python

"""
Ollama client for embeddings generation.
Provides async embedding operations via Ollama API:
- Text embedding generation
- Batch embedding support
- Model management
"""
import httpx
from typing import List, Dict, Any, Optional
import logging
logger = logging.getLogger(__name__)
class OllamaClient:
"""
Ollama API client for embeddings.
Documentation: https://github.com/ollama/ollama/blob/main/docs/api.md
Default model: nomic-embed-text (768-dimensional embeddings)
"""
def __init__(self, base_url: str, model: str = "nomic-embed-text"):
"""
Initialize Ollama client.
Args:
base_url: Ollama server URL (e.g., "http://ollama:11434")
model: Embedding model name (default: "nomic-embed-text")
"""
self.base_url = base_url.rstrip("/")
self.model = model
self.embeddings_url = f"{self.base_url}/api/embeddings"
self.embed_url = f"{self.base_url}/api/embed"
self.generate_url = f"{self.base_url}/api/generate"
self.tags_url = f"{self.base_url}/api/tags"
self.client = httpx.AsyncClient(timeout=120.0) # Embeddings can be slow
logger.info(f"Initialized Ollama client: {base_url}, model: {model}")
async def close(self):
"""Close HTTP client"""
await self.client.aclose()
async def embed(self, text: str) -> Optional[List[float]]:
"""
Generate embedding for single text.
Args:
text: Text to embed
Returns:
Embedding vector (768-dimensional for nomic-embed-text) or None on failure
Example:
>>> embedding = await client.embed("Hello world")
>>> len(embedding)
768
"""
try:
payload = {
"model": self.model,
"prompt": text
}
response = await self.client.post(
self.embeddings_url,
json=payload
)
response.raise_for_status()
data = response.json()
embedding = data.get("embedding")
if not embedding:
logger.error(f"No embedding in response: {data}")
return None
return embedding
except httpx.HTTPStatusError as e:
logger.error(f"HTTP error: {e.response.status_code} - {e.response.text}")
return None
except Exception as e:
logger.error(f"Embedding failed: {e}", exc_info=True)
return None
async def embed_batch(
self,
texts: List[str],
show_progress: bool = False
) -> List[Optional[List[float]]]:
"""
Generate embeddings for multiple texts.
Args:
texts: List of texts to embed
show_progress: Log progress for large batches
Returns:
List of embedding vectors (same order as input)
None entries for texts that failed to embed
Example:
>>> texts = ["Hello", "World", "Test"]
>>> embeddings = await client.embed_batch(texts)
>>> len(embeddings)
3
"""
if not texts:
return []
# Single batched request via Ollama's /api/embed (the old
# implementation looped one /api/embeddings call per text).
try:
response = await self.client.post(
self.embed_url,
json={"model": self.model, "input": texts}
)
response.raise_for_status()
data = response.json()
embeddings = data.get("embeddings")
if embeddings is not None and len(embeddings) == len(texts):
if show_progress:
logger.info(f"Batched embedding complete: {len(embeddings)}/{len(texts)}")
return embeddings
logger.warning(
f"Batched embed returned {len(embeddings or [])} vectors for "
f"{len(texts)} inputs, falling back to per-text embedding"
)
except Exception as e:
logger.warning(
f"Batched embed failed ({e}), falling back to per-text embedding"
)
# Fallback: per-text embedding preserves partial-success semantics
# (None entries for texts that failed to embed).
embeddings = []
for i, text in enumerate(texts):
if show_progress and i % 10 == 0:
logger.info(f"Embedding progress: {i}/{len(texts)}")
embedding = await self.embed(text)
embeddings.append(embedding)
if show_progress:
logger.info(f"Embedding complete: {len(embeddings)}/{len(texts)}")
return embeddings
async def embed_batch_filtered(
self,
texts: List[str],
show_progress: bool = False
) -> List[List[float]]:
"""
Generate embeddings for multiple texts, filtering out failures.
Args:
texts: List of texts to embed
show_progress: Log progress for large batches
Returns:
List of successful embedding vectors (may be shorter than input)
Example:
>>> texts = ["Hello", "World", "Test"]
>>> embeddings = await client.embed_batch_filtered(texts)
>>> all(e is not None for e in embeddings)
True
"""
all_embeddings = await self.embed_batch(texts, show_progress)
return [e for e in all_embeddings if e is not None]
async def embed_documents(
self,
documents: List[Dict[str, Any]],
content_field: str = "content",
show_progress: bool = False
) -> List[Dict[str, Any]]:
"""
Embed documents with metadata preservation.
Args:
documents: List of document dictionaries
content_field: Field name containing text to embed
show_progress: Log progress for large batches
Returns:
List of documents with added "embedding" field
Example:
>>> docs = [
... {"content": "Hello world", "id": 1},
... {"content": "Test doc", "id": 2}
... ]
>>> embedded = await client.embed_documents(docs)
>>> "embedding" in embedded[0]
True
"""
texts = [doc.get(content_field, "") for doc in documents]
embeddings = await self.embed_batch(texts, show_progress)
results = []
for doc, embedding in zip(documents, embeddings):
doc_copy = doc.copy()
doc_copy["embedding"] = embedding
results.append(doc_copy)
return results
async def get_embedding_dimension(self) -> Optional[int]:
"""
Get embedding dimension for current model.
Returns:
Embedding dimension (e.g., 768 for nomic-embed-text) or None on failure
Example:
>>> dim = await client.get_embedding_dimension()
>>> dim
768
"""
# Generate a test embedding to determine dimension
test_embedding = await self.embed("test")
if test_embedding:
return len(test_embedding)
return None
async def list_models(self) -> List[Dict[str, Any]]:
"""
List available Ollama models.
Returns:
List of model information dictionaries
Example:
>>> models = await client.list_models()
>>> any(m["name"] == "nomic-embed-text" for m in models)
True
"""
try:
response = await self.client.get(self.tags_url)
response.raise_for_status()
data = response.json()
return data.get("models", [])
except Exception as e:
logger.error(f"Failed to list models: {e}")
return []
async def check_model_available(self, model_name: Optional[str] = None) -> bool:
"""
Check if a model is available.
Args:
model_name: Model name to check (defaults to self.model)
Returns:
True if model is available, False otherwise
"""
check_model = model_name or self.model
models = await self.list_models()
# Check for exact match or match with :latest tag
for m in models:
name = m.get("name", "")
# Exact match
if name == check_model:
return True
# Match without tag (e.g., "nomic-embed-text" matches "nomic-embed-text:latest")
if name.startswith(f"{check_model}:"):
return True
return False
async def generate_text(
self,
prompt: str,
model: Optional[str] = None,
stream: bool = False,
temperature: Optional[float] = None
) -> Optional[str]:
"""
Generate text completion (for non-embedding use cases).
Args:
prompt: Input prompt
model: Model name (defaults to self.model)
stream: Enable streaming response
temperature: Sampling temperature (0.0 = deterministic, higher = more creative)
None uses model default (~0.7 for mistral-nemo)
Returns:
Generated text or None on failure
Note: Use temperature=0.0 for deterministic outputs like JSON parsing,
ranking, and factual extraction. Use higher values (0.3-0.7) for
creative content generation.
"""
try:
payload = {
"model": model or self.model,
"prompt": prompt,
"stream": stream
}
# Add temperature to options if specified
if temperature is not None:
payload["options"] = {"temperature": temperature}
response = await self.client.post(
self.generate_url,
json=payload
)
response.raise_for_status()
if stream:
# For streaming, return first chunk
# Full streaming implementation would need async generator
return response.text
else:
data = response.json()
return data.get("response")
except Exception as e:
logger.error(f"Text generation failed: {e}")
return None
async def health_check(self) -> bool:
"""
Check if Ollama server is reachable and model is available.
Returns:
True if healthy, False otherwise
"""
try:
# Check server is up and get models in one call
response = await self.client.get(self.tags_url, timeout=5.0)
response.raise_for_status()
data = response.json()
models = data.get("models", [])
# Check model is available
check_model = self.model
model_found = False
for m in models:
name = m.get("name", "")
if name == check_model or name.startswith(f"{check_model}:"):
model_found = True
break
if not model_found:
logger.warning(f"Model '{self.model}' not found in Ollama")
return False
return True
except Exception as e:
logger.error(f"Health check failed: {e}")
return False
def estimate_tokens(self, text: str) -> int:
"""
Rough estimate of token count for text.
Uses simple heuristic: ~4 characters per token.
Args:
text: Input text
Returns:
Estimated token count
"""
return len(text) // 4
def chunk_text_for_embedding(
self,
text: str,
max_tokens: int = 512,
overlap: int = 50
) -> List[str]:
"""
Chunk text into segments suitable for embedding.
Args:
text: Input text
max_tokens: Maximum tokens per chunk
overlap: Token overlap between chunks
Returns:
List of text chunks
Example:
>>> chunks = client.chunk_text_for_embedding(long_text, max_tokens=512)
>>> all(client.estimate_tokens(c) <= 512 for c in chunks)
True
"""
# Convert tokens to approximate character count
max_chars = max_tokens * 4
overlap_chars = overlap * 4
if len(text) <= max_chars:
return [text]
chunks = []
start = 0
while start < len(text):
end = start + max_chars
chunk = text[start:end]
# Try to break at sentence boundary
if end < len(text):
last_period = chunk.rfind(". ")
if last_period > max_chars * 0.5: # Only break if > 50% through chunk
end = start + last_period + 1
chunk = text[start:end]
chunks.append(chunk.strip())
start = end - overlap_chars
return chunks