feat: use internal ollama embeddings
This commit is contained in:
@@ -0,0 +1,92 @@
|
||||
# Task 4 Report — Narrow harness embedding configuration to internal Ollama
|
||||
|
||||
## Status
|
||||
|
||||
Implemented on 2026-08-08 in `/Users/mp/projects/ThothII/.worktrees/git-workspace-registry`.
|
||||
|
||||
## RED evidence
|
||||
|
||||
Command:
|
||||
|
||||
```bash
|
||||
cd harness
|
||||
./.venv/bin/pytest tests/test_internal_embeddings.py tests/test_config_resources.py -q
|
||||
```
|
||||
|
||||
Observed before implementation:
|
||||
|
||||
- exit code `1`
|
||||
- `10 failed, 10 passed`
|
||||
- failures proved the missing `OllamaInternalEmbeddings` client and missing internal-only config validation
|
||||
|
||||
Representative failures:
|
||||
|
||||
- `ImportError: cannot import name 'OllamaInternalEmbeddings'`
|
||||
- `AttributeError: 'EmbeddingsConfig' object has no attribute 'provider'`
|
||||
- config tests `DID NOT RAISE ConfigError` for external provider, API key, and non-private base URL
|
||||
|
||||
## GREEN evidence
|
||||
|
||||
Focused behavior suite:
|
||||
|
||||
```bash
|
||||
cd harness
|
||||
./.venv/bin/pytest tests/test_internal_embeddings.py tests/test_config_resources.py -q
|
||||
```
|
||||
|
||||
- exit code `0`
|
||||
- `20 passed`
|
||||
|
||||
Relevant harness verification:
|
||||
|
||||
```bash
|
||||
cd harness
|
||||
./.venv/bin/pytest tests/test_internal_embeddings.py tests/test_config_resources.py tests/test_ollama_ensure.py -q
|
||||
```
|
||||
|
||||
- exit code `0`
|
||||
- `36 passed, 2 warnings`
|
||||
|
||||
Changed-file lint:
|
||||
|
||||
```bash
|
||||
cd harness
|
||||
./.venv/bin/ruff check tht/config.py tht/config_compat.py tht/vectorstore/embeddings.py tht/cli/ollama_cmd.py tests/test_config_resources.py tests/test_internal_embeddings.py
|
||||
```
|
||||
|
||||
- exit code `0`
|
||||
- `All checks passed!`
|
||||
|
||||
Patch hygiene:
|
||||
|
||||
```bash
|
||||
git diff --check
|
||||
```
|
||||
|
||||
- exit code `0`
|
||||
|
||||
## What changed
|
||||
|
||||
- translated schema-v3 `resources.embeddings` into the harness-compatible embedding config view
|
||||
- validated the internal embedding contract only for that runtime-owned `resources.embeddings` path:
|
||||
- provider must be `ollama_internal`
|
||||
- model must be `qwen3-embedding:0.6b`
|
||||
- dimensions must be `1024`
|
||||
- base URL must be `http://embedding:11434` or loopback HTTP on port `11434`
|
||||
- extra fields like `api_key` are rejected
|
||||
- replaced the active embed client with `OllamaInternalEmbeddings`, using one bounded `/api/embed` request per batch
|
||||
- removed task/query prefix rewriting from the active embedding path
|
||||
- validated response count, vector dimension, and finite numeric values before returning embeddings
|
||||
- kept `tht ollama ensure --json` stdout pristine while warming through the internal client
|
||||
|
||||
## Self-review
|
||||
|
||||
- kept changes inside the brief-listed files
|
||||
- preserved DWH and session-persistence behavior
|
||||
- preserved the legacy `OllamaEmbeddings` import path as an alias to avoid unrelated call-site churn
|
||||
|
||||
## Concerns
|
||||
|
||||
- the focused harness verification still emits two pre-existing warnings:
|
||||
- `DeprecationWarning` from `testcontainers.postgres`
|
||||
- `FutureWarning` because `resources` currently flows through the legacy config translation path
|
||||
@@ -1,16 +1,16 @@
|
||||
import pytest
|
||||
|
||||
from tht.adapters.evidence import FilesystemEvidenceSource, HttpManifestEvidenceSource
|
||||
from tht.adapters.factory import build_evidence_sources
|
||||
from tht.config import (
|
||||
ConfigError,
|
||||
PgvectorDirectConfig,
|
||||
PostgresDwhConfig,
|
||||
ThothRestDwhConfig,
|
||||
ThothVectorHttpConfig,
|
||||
workspace_id_for_config,
|
||||
load_config,
|
||||
workspace_id_for_config,
|
||||
)
|
||||
from tht.adapters.evidence import FilesystemEvidenceSource, HttpManifestEvidenceSource
|
||||
from tht.adapters.factory import build_evidence_sources
|
||||
|
||||
|
||||
def test_direct_vector_passwords_load_from_file_references(monkeypatch, tmp_path):
|
||||
@@ -213,6 +213,95 @@ embeddings: {base_url: http://ollama:11434, dim: 768}
|
||||
assert cfg.vectors.writer.api_key == "writer"
|
||||
|
||||
|
||||
def test_accepts_only_internal_ollama_embedding_contract(tmp_path):
|
||||
workspace = tmp_path / "workspace.yaml"
|
||||
workspace.write_text(
|
||||
"""
|
||||
dwh:
|
||||
type: postgres_direct
|
||||
connection: {database: analytics, schema: mart, user: reader, password: secret}
|
||||
resources:
|
||||
embeddings:
|
||||
provider: ollama_internal
|
||||
base_url: http://embedding:11434
|
||||
model: qwen3-embedding:0.6b
|
||||
dimensions: 1024
|
||||
"""
|
||||
)
|
||||
|
||||
cfg = load_config(workspace)
|
||||
|
||||
assert cfg.embeddings.provider == "ollama_internal"
|
||||
assert cfg.embeddings.base_url == "http://embedding:11434"
|
||||
assert cfg.embeddings.model == "qwen3-embedding:0.6b"
|
||||
assert cfg.embeddings.dim == 1024
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("snippet", "pattern"),
|
||||
[
|
||||
(
|
||||
"""
|
||||
resources:
|
||||
embeddings:
|
||||
provider: openai_compatible
|
||||
base_url: http://embedding:11434
|
||||
model: qwen3-embedding:0.6b
|
||||
dimensions: 1024
|
||||
""",
|
||||
"ollama_internal|provider",
|
||||
),
|
||||
(
|
||||
"""
|
||||
resources:
|
||||
embeddings:
|
||||
provider: ollama_internal
|
||||
base_url: http://embedding:11434
|
||||
model: qwen3-embedding:0.6b
|
||||
dimensions: 1024
|
||||
api_key: secret
|
||||
""",
|
||||
"api_key|extra",
|
||||
),
|
||||
(
|
||||
"""
|
||||
resources:
|
||||
embeddings:
|
||||
provider: ollama_internal
|
||||
base_url: https://embedding:11434
|
||||
model: qwen3-embedding:0.6b
|
||||
dimensions: 1024
|
||||
""",
|
||||
"base_url|internal|private|host",
|
||||
),
|
||||
(
|
||||
"""
|
||||
resources:
|
||||
embeddings:
|
||||
provider: ollama_internal
|
||||
base_url: http://example.com:11434
|
||||
model: qwen3-embedding:0.6b
|
||||
dimensions: 1024
|
||||
""",
|
||||
"base_url|internal|private|host",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_rejects_external_embedding_configuration(tmp_path, snippet, pattern):
|
||||
workspace = tmp_path / "workspace.yaml"
|
||||
workspace.write_text(
|
||||
"""
|
||||
dwh:
|
||||
type: postgres_direct
|
||||
connection: {database: analytics, schema: mart, user: reader, password: secret}
|
||||
"""
|
||||
+ snippet
|
||||
)
|
||||
|
||||
with pytest.raises(ConfigError, match=pattern):
|
||||
load_config(workspace)
|
||||
|
||||
|
||||
def test_builds_typed_evidence_sources_and_keeps_legacy_compatible(tmp_path):
|
||||
common = """
|
||||
dwh:
|
||||
|
||||
@@ -0,0 +1,139 @@
|
||||
import math
|
||||
|
||||
import pytest
|
||||
|
||||
from tht.config import EmbeddingsConfig
|
||||
from tht.vectorstore.embeddings import EmbeddingsError
|
||||
|
||||
|
||||
class _Response:
|
||||
def __init__(self, payload, status_code=200):
|
||||
self._payload = payload
|
||||
self.status_code = status_code
|
||||
|
||||
def raise_for_status(self):
|
||||
if self.status_code >= 400:
|
||||
raise RuntimeError(f"http {self.status_code}")
|
||||
|
||||
def json(self):
|
||||
return self._payload
|
||||
|
||||
|
||||
class _Session:
|
||||
def __init__(self, responses):
|
||||
self._responses = list(responses)
|
||||
self.calls = []
|
||||
|
||||
def post(self, url, json, timeout):
|
||||
self.calls.append({"url": url, "json": json, "timeout": timeout})
|
||||
if not self._responses:
|
||||
raise AssertionError("unexpected extra request")
|
||||
return self._responses.pop(0)
|
||||
|
||||
|
||||
def _vector(value: float, *, dim: int = 1024):
|
||||
return [value] * dim
|
||||
|
||||
|
||||
def test_internal_embeddings_posts_model_and_batch_input_without_prefixes():
|
||||
from tht.vectorstore.embeddings import OllamaInternalEmbeddings
|
||||
|
||||
session = _Session([_Response({"embeddings": [_vector(1.0), _vector(2.0)]})])
|
||||
embedder = OllamaInternalEmbeddings(
|
||||
EmbeddingsConfig(
|
||||
provider="ollama_internal",
|
||||
base_url="http://embedding:11434",
|
||||
model="qwen3-embedding:0.6b",
|
||||
dim=1024,
|
||||
batch_size=2,
|
||||
timeout=9,
|
||||
connect_timeout=4,
|
||||
),
|
||||
session=session,
|
||||
)
|
||||
|
||||
vectors = embedder.embed(["alpha", "beta"])
|
||||
|
||||
assert vectors == [_vector(1.0), _vector(2.0)]
|
||||
assert session.calls == [{
|
||||
"url": "http://embedding:11434/api/embed",
|
||||
"json": {"model": "qwen3-embedding:0.6b", "input": ["alpha", "beta"]},
|
||||
"timeout": (4, 9),
|
||||
}]
|
||||
|
||||
|
||||
def test_internal_embeddings_batching_returns_1024d_vectors():
|
||||
from tht.vectorstore.embeddings import OllamaInternalEmbeddings
|
||||
|
||||
session = _Session([
|
||||
_Response({"embeddings": [_vector(1.0), _vector(2.0)]}),
|
||||
_Response({"embeddings": [_vector(3.0)]}),
|
||||
])
|
||||
embedder = OllamaInternalEmbeddings(
|
||||
EmbeddingsConfig(
|
||||
provider="ollama_internal",
|
||||
base_url="http://embedding:11434",
|
||||
model="qwen3-embedding:0.6b",
|
||||
dim=1024,
|
||||
batch_size=2,
|
||||
),
|
||||
session=session,
|
||||
)
|
||||
|
||||
vectors = embedder.embed(["one", "two", "three"])
|
||||
|
||||
assert [len(vector) for vector in vectors] == [1024, 1024, 1024]
|
||||
assert [vector[0] for vector in vectors] == [1.0, 2.0, 3.0]
|
||||
|
||||
|
||||
def test_internal_embeddings_reject_count_mismatch():
|
||||
from tht.vectorstore.embeddings import OllamaInternalEmbeddings
|
||||
|
||||
embedder = OllamaInternalEmbeddings(
|
||||
EmbeddingsConfig(
|
||||
provider="ollama_internal",
|
||||
base_url="http://embedding:11434",
|
||||
model="qwen3-embedding:0.6b",
|
||||
dim=1024,
|
||||
),
|
||||
session=_Session([_Response({"embeddings": [_vector(1.0)]})]),
|
||||
)
|
||||
|
||||
with pytest.raises(EmbeddingsError, match="count|numero"):
|
||||
embedder.embed(["alpha", "beta"])
|
||||
|
||||
|
||||
def test_internal_embeddings_reject_dimension_mismatch():
|
||||
from tht.vectorstore.embeddings import OllamaInternalEmbeddings
|
||||
|
||||
embedder = OllamaInternalEmbeddings(
|
||||
EmbeddingsConfig(
|
||||
provider="ollama_internal",
|
||||
base_url="http://embedding:11434",
|
||||
model="qwen3-embedding:0.6b",
|
||||
dim=1024,
|
||||
),
|
||||
session=_Session([_Response({"embeddings": [[1.0] * 8]})]),
|
||||
)
|
||||
|
||||
with pytest.raises(EmbeddingsError, match="dimensione|dimension"):
|
||||
embedder.embed(["alpha"])
|
||||
|
||||
|
||||
def test_internal_embeddings_reject_non_finite_values():
|
||||
from tht.vectorstore.embeddings import OllamaInternalEmbeddings
|
||||
|
||||
bad = _vector(0.0)
|
||||
bad[10] = math.nan
|
||||
embedder = OllamaInternalEmbeddings(
|
||||
EmbeddingsConfig(
|
||||
provider="ollama_internal",
|
||||
base_url="http://embedding:11434",
|
||||
model="qwen3-embedding:0.6b",
|
||||
dim=1024,
|
||||
),
|
||||
session=_Session([_Response({"embeddings": [bad]})]),
|
||||
)
|
||||
|
||||
with pytest.raises(EmbeddingsError, match="finite|finit"):
|
||||
embedder.embed(["alpha"])
|
||||
@@ -39,16 +39,16 @@ def _installed_models(base_url: str, timeout: float = 5.0) -> set[str]:
|
||||
|
||||
def _start(start_cmd: list[str]) -> None:
|
||||
# Detached so the server outlives this short-lived CLI process.
|
||||
subprocess.Popen( # noqa: S603
|
||||
subprocess.Popen(
|
||||
start_cmd, start_new_session=True,
|
||||
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
|
||||
)
|
||||
|
||||
|
||||
def _warm(cfg) -> None:
|
||||
from tht.vectorstore.embeddings import OllamaEmbeddings
|
||||
from tht.vectorstore.embeddings import OllamaInternalEmbeddings
|
||||
|
||||
OllamaEmbeddings(cfg.embeddings).embed_query("ping")
|
||||
OllamaInternalEmbeddings(cfg.embeddings).embed_query("ping")
|
||||
|
||||
|
||||
def _model_present(installed: set[str], model: str) -> bool:
|
||||
@@ -104,7 +104,7 @@ def ensure_ollama(
|
||||
if probe(base_url):
|
||||
up = True
|
||||
break
|
||||
except Exception: # noqa: BLE001 - transient during startup; keep polling
|
||||
except Exception: # noqa: BLE001,S112 - transient during startup; keep polling
|
||||
continue
|
||||
if not up:
|
||||
return {"ok": False, "stage": "server",
|
||||
|
||||
+69
-2
@@ -1,11 +1,13 @@
|
||||
import os
|
||||
import re
|
||||
import warnings
|
||||
from ipaddress import ip_address
|
||||
from pathlib import Path
|
||||
from typing import Annotated, Any, Literal
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import yaml
|
||||
from pydantic import BaseModel, Field, PrivateAttr, SecretStr, model_validator, ValidationError
|
||||
from pydantic import BaseModel, Field, PrivateAttr, SecretStr, ValidationError, model_validator
|
||||
|
||||
from tht.config_compat import translate_legacy_config
|
||||
|
||||
@@ -276,15 +278,18 @@ class EvidenceSourcesConfig(BaseModel):
|
||||
|
||||
|
||||
class EmbeddingsConfig(BaseModel):
|
||||
provider: str = "ollama_internal"
|
||||
base_url: str
|
||||
model: str = "nomic-embed-text-v2-moe"
|
||||
dim: int = 768
|
||||
dim: int = Field(default=768, alias="dimensions")
|
||||
batch_size: int = 32
|
||||
timeout: int = 30
|
||||
connect_timeout: int = 5
|
||||
bin: str = "ollama"
|
||||
start_cmd: list[str] | None = None
|
||||
|
||||
model_config = {"populate_by_name": True, "extra": "forbid"}
|
||||
|
||||
|
||||
class VectorConfig(BaseModel):
|
||||
max_chunk_chars: int = 4000
|
||||
@@ -391,6 +396,7 @@ def load_config(path: Path) -> Config:
|
||||
if not isinstance(raw, dict):
|
||||
raise ConfigError(f"Configurazione non valida (atteso un mapping YAML): {path}")
|
||||
expanded = _resolve_secret_files(_expand_env(raw))
|
||||
_validate_internal_embedding_contract(expanded, path)
|
||||
translated, used_legacy = translate_legacy_config(expanded)
|
||||
_populate_legacy_views(translated)
|
||||
try:
|
||||
@@ -451,6 +457,67 @@ def load_config(path: Path) -> Config:
|
||||
return cfg
|
||||
|
||||
|
||||
def _validate_internal_embedding_contract(raw: dict[str, Any], path: Path) -> None:
|
||||
resources = raw.get("resources")
|
||||
if not isinstance(resources, dict):
|
||||
return
|
||||
embeddings = resources.get("embeddings")
|
||||
if not isinstance(embeddings, dict):
|
||||
return
|
||||
|
||||
provider = embeddings.get("provider")
|
||||
model = embeddings.get("model")
|
||||
dimensions = embeddings.get("dimensions")
|
||||
base_url = embeddings.get("base_url")
|
||||
allowed = {"provider", "base_url", "model", "dimensions"}
|
||||
unexpected = sorted(set(embeddings) - allowed)
|
||||
if unexpected:
|
||||
raise ConfigError(
|
||||
f"Configurazione non valida in {path}:\n"
|
||||
f"resources.embeddings non supporta: {', '.join(unexpected)}"
|
||||
)
|
||||
if provider != "ollama_internal":
|
||||
raise ConfigError(
|
||||
f"Configurazione non valida in {path}:\n"
|
||||
"resources.embeddings.provider deve essere 'ollama_internal'"
|
||||
)
|
||||
if model != "qwen3-embedding:0.6b":
|
||||
raise ConfigError(
|
||||
f"Configurazione non valida in {path}:\n"
|
||||
"resources.embeddings.model deve essere 'qwen3-embedding:0.6b'"
|
||||
)
|
||||
if dimensions != 1024:
|
||||
raise ConfigError(
|
||||
f"Configurazione non valida in {path}:\n"
|
||||
"resources.embeddings.dimensions deve essere 1024"
|
||||
)
|
||||
if not _is_allowed_internal_embedding_url(base_url):
|
||||
raise ConfigError(
|
||||
f"Configurazione non valida in {path}:\n"
|
||||
"resources.embeddings.base_url deve usare http://embedding:11434 "
|
||||
"oppure un endpoint loopback di sviluppo su porta 11434"
|
||||
)
|
||||
|
||||
|
||||
def _is_allowed_internal_embedding_url(value: Any) -> bool:
|
||||
if not isinstance(value, str):
|
||||
return False
|
||||
parsed = urlparse(value)
|
||||
if parsed.scheme != "http" or not parsed.hostname or parsed.port != 11434:
|
||||
return False
|
||||
if parsed.params or parsed.query or parsed.fragment:
|
||||
return False
|
||||
if parsed.path not in ("", "/"):
|
||||
return False
|
||||
if parsed.hostname == "embedding":
|
||||
return True
|
||||
try:
|
||||
host = ip_address(parsed.hostname)
|
||||
except ValueError:
|
||||
return parsed.hostname == "localhost"
|
||||
return host.is_loopback
|
||||
|
||||
|
||||
def _populate_legacy_views(raw: dict[str, Any]) -> None:
|
||||
"""Populate old Config attributes for command compatibility during migration."""
|
||||
dwh = raw.get("dwh")
|
||||
|
||||
@@ -3,7 +3,6 @@ from __future__ import annotations
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
|
||||
_LEGACY_RESOURCE_KEYS = {
|
||||
"database",
|
||||
"rest",
|
||||
@@ -11,6 +10,7 @@ _LEGACY_RESOURCE_KEYS = {
|
||||
"vector_rest",
|
||||
"vector_write_rest",
|
||||
"paths",
|
||||
"resources",
|
||||
}
|
||||
|
||||
|
||||
@@ -26,6 +26,13 @@ def _as_mapping(value: Any) -> dict[str, Any] | None:
|
||||
def translate_legacy_config(raw: dict[str, Any]) -> tuple[dict[str, Any], bool]:
|
||||
"""Translate the legacy flat resource keys without validating their contents."""
|
||||
translated = deepcopy(raw)
|
||||
resources = _as_mapping(raw.get("resources"))
|
||||
if isinstance(resources, dict) and "embeddings" in resources and "embeddings" not in translated:
|
||||
embedding = _as_mapping(resources.get("embeddings"))
|
||||
if isinstance(embedding, dict):
|
||||
translated["embeddings"] = embedding
|
||||
if "dimensions" in translated["embeddings"] and "dim" not in translated["embeddings"]:
|
||||
translated["embeddings"]["dim"] = translated["embeddings"].pop("dimensions")
|
||||
legacy = any(key in raw for key in _LEGACY_RESOURCE_KEYS)
|
||||
if not legacy:
|
||||
return translated, False
|
||||
|
||||
@@ -1,104 +1,73 @@
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import math
|
||||
|
||||
import requests
|
||||
|
||||
from tht.config import EmbeddingsConfig
|
||||
|
||||
DOC_PREFIX = "search_document: "
|
||||
QUERY_PREFIX = "search_query: "
|
||||
|
||||
_RESTART_WAIT = 8 # secondi di attesa dopo aver avviato Ollama
|
||||
_RESTART_POLL = 1.0
|
||||
|
||||
|
||||
class EmbeddingsError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class OllamaEmbeddings:
|
||||
"""Client embeddings via Ollama. Applica i prefissi di task richiesti da nomic v2:
|
||||
ometterli degrada il retrieval in modo silenzioso."""
|
||||
class OllamaInternalEmbeddings:
|
||||
"""Client embeddings for the installation-owned internal Ollama endpoint."""
|
||||
|
||||
def __init__(self, cfg: EmbeddingsConfig):
|
||||
def __init__(self, cfg: EmbeddingsConfig, *, session: requests.Session | None = None):
|
||||
self.cfg = cfg
|
||||
self._session = session or requests.Session()
|
||||
self.base_url = self.cfg.base_url.rstrip("/")
|
||||
self.model = self.cfg.model
|
||||
self.dim = self.cfg.dim
|
||||
self.timeout = (self.cfg.connect_timeout, self.cfg.timeout)
|
||||
self.batch_size = self.cfg.batch_size
|
||||
|
||||
def _is_up(self) -> bool:
|
||||
def _post(self, texts: list[str]) -> list[list[float]]:
|
||||
try:
|
||||
r = requests.get(
|
||||
f"{self.cfg.base_url.rstrip('/')}/api/tags",
|
||||
timeout=(self.cfg.connect_timeout, 5),
|
||||
response = self._session.post(
|
||||
f"{self.base_url}/api/embed",
|
||||
json={"model": self.model, "input": texts},
|
||||
timeout=self.timeout,
|
||||
)
|
||||
return r.status_code == 200
|
||||
except requests.RequestException:
|
||||
return False
|
||||
|
||||
def _try_restart(self) -> bool:
|
||||
"""Tenta di avviare Ollama e attende che sia raggiungibile."""
|
||||
start_cmd = self.cfg.start_cmd if self.cfg.start_cmd is not None else [self.cfg.bin, "serve"]
|
||||
if not start_cmd:
|
||||
return False
|
||||
try:
|
||||
subprocess.Popen( # noqa: S603
|
||||
start_cmd, start_new_session=True,
|
||||
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
|
||||
response.raise_for_status()
|
||||
except requests.RequestException as exc:
|
||||
raise EmbeddingsError(
|
||||
f"internal Ollama embeddings request failed for model {self.model}"
|
||||
) from exc
|
||||
payload = response.json()
|
||||
embeddings = payload.get("embeddings")
|
||||
if not isinstance(embeddings, list):
|
||||
raise EmbeddingsError("internal Ollama response is missing embeddings")
|
||||
if len(embeddings) != len(texts):
|
||||
raise EmbeddingsError(
|
||||
f"unexpected embedding count: {len(embeddings)} != {len(texts)}"
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
return False
|
||||
print("[embeddings] Ollama non raggiungibile, avvio in corso…", file=sys.stderr)
|
||||
deadline = time.monotonic() + _RESTART_WAIT
|
||||
while time.monotonic() < deadline:
|
||||
time.sleep(_RESTART_POLL)
|
||||
if self._is_up():
|
||||
return True
|
||||
return False
|
||||
|
||||
def _post(self, url: str, batch: list[str]) -> requests.Response:
|
||||
resp = requests.post(
|
||||
url, json={"model": self.cfg.model, "input": batch},
|
||||
timeout=(self.cfg.connect_timeout, self.cfg.timeout),
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp
|
||||
|
||||
def _embed(self, texts: list[str]) -> list[list[float]]:
|
||||
url = f"{self.cfg.base_url.rstrip('/')}/api/embed"
|
||||
out: list[list[float]] = []
|
||||
for i in range(0, len(texts), self.cfg.batch_size):
|
||||
batch = texts[i : i + self.cfg.batch_size]
|
||||
try:
|
||||
resp = self._post(url, batch)
|
||||
except requests.ConnectionError:
|
||||
if not self._try_restart():
|
||||
raise EmbeddingsError(
|
||||
f"Ollama non raggiungibile su {self.cfg.base_url} "
|
||||
f"(modello {self.cfg.model}), avvio automatico fallito"
|
||||
)
|
||||
try:
|
||||
resp = self._post(url, batch)
|
||||
except requests.RequestException as e:
|
||||
raise EmbeddingsError(
|
||||
f"Ollama non raggiungibile su {self.cfg.base_url} "
|
||||
f"(modello {self.cfg.model}): {e}"
|
||||
) from e
|
||||
except requests.RequestException as e:
|
||||
validated: list[list[float]] = []
|
||||
for vector in embeddings:
|
||||
if not isinstance(vector, list):
|
||||
raise EmbeddingsError("internal Ollama returned a non-vector embedding")
|
||||
if len(vector) != self.dim:
|
||||
raise EmbeddingsError(
|
||||
f"Ollama non raggiungibile su {self.cfg.base_url} "
|
||||
f"(modello {self.cfg.model}): {e}"
|
||||
) from e
|
||||
embeddings = resp.json().get("embeddings", [])
|
||||
for v in embeddings:
|
||||
if len(v) != self.cfg.dim:
|
||||
raise EmbeddingsError(
|
||||
f"dimensione embedding inattesa: {len(v)} != {self.cfg.dim} "
|
||||
f"(modello {self.cfg.model})"
|
||||
)
|
||||
out.extend(embeddings)
|
||||
f"unexpected embedding dimension: {len(vector)} != {self.dim}"
|
||||
)
|
||||
cleaned: list[float] = []
|
||||
for value in vector:
|
||||
if not isinstance(value, (int, float)) or not math.isfinite(value):
|
||||
raise EmbeddingsError("internal Ollama returned a non-finite embedding value")
|
||||
cleaned.append(float(value))
|
||||
validated.append(cleaned)
|
||||
return validated
|
||||
|
||||
def embed(self, texts: list[str]) -> list[list[float]]:
|
||||
out: list[list[float]] = []
|
||||
for i in range(0, len(texts), self.batch_size):
|
||||
out.extend(self._post(texts[i : i + self.batch_size]))
|
||||
return out
|
||||
|
||||
def embed_documents(self, texts: list[str]) -> list[list[float]]:
|
||||
return self._embed([DOC_PREFIX + t for t in texts])
|
||||
return self.embed(texts)
|
||||
|
||||
def embed_query(self, text: str) -> list[float]:
|
||||
return self._embed([QUERY_PREFIX + text])[0]
|
||||
return self.embed([text])[0]
|
||||
|
||||
|
||||
OllamaEmbeddings = OllamaInternalEmbeddings
|
||||
|
||||
Reference in New Issue
Block a user