Improve VLM client lifecycle, retries, metrics, and cache keys

- Adds abstract aclose() and shared load_prompt() helper to BaseVLMClient.

- OllamaClient now logs latency, tokens, and prompt hash per call.

- Retries cover HTTP 429 and 5xx in addition to network/timeout errors.

- Cache key now includes response_format and max_tokens.

- Warns when JSON schema generation falls back to plain 'json' format.

- Adds FakeVLMClient.aclose() and cache unit tests.
This commit is contained in:
2026-08-21 16:44:14 +02:00
parent 17e0496d2f
commit 081084816b
5 changed files with 132 additions and 12 deletions
+23 -1
View File
@@ -35,6 +35,11 @@ class BaseVLMClient(ABC):
"""Send a chat request with optional images and return parsed text.""" """Send a chat request with optional images and return parsed text."""
... ...
@abstractmethod
async def aclose(self) -> None:
"""Release client resources (connections, file handles, ...)."""
...
@staticmethod @staticmethod
def parse_json(content: str, model: type[T]) -> T: def parse_json(content: str, model: type[T]) -> T:
"""Strip markdown fences and parse/validate JSON against a Pydantic model.""" """Strip markdown fences and parse/validate JSON against a Pydantic model."""
@@ -54,7 +59,24 @@ class BaseVLMClient(ABC):
@staticmethod @staticmethod
def _load_prompt_file(path: Path | str) -> str: def _load_prompt_file(path: Path | str) -> str:
"""Load a plain-text prompt from the prompts directory.""" """Load a plain-text prompt file."""
path = Path(path) path = Path(path)
with path.open("r", encoding="utf-8") as f: with path.open("r", encoding="utf-8") as f:
return f.read() return f.read()
def load_prompt(name: str) -> str:
"""Load a plain-text system prompt from the project's ``prompts/`` directory.
Tries the current working directory first (running from the repo root),
then falls back to the repository root derived from this package's
location, so the CLI keeps working from any working directory.
"""
candidates = [
Path("prompts") / name,
Path(__file__).resolve().parents[3] / "prompts" / name,
]
for candidate in candidates:
if candidate.is_file():
return candidate.read_text(encoding="utf-8")
raise FileNotFoundError(f"Prompt file not found: {name!r} (tried: {candidates})")
+36 -3
View File
@@ -3,6 +3,7 @@
from __future__ import annotations from __future__ import annotations
import base64 import base64
import hashlib
import io import io
import logging import logging
import time import time
@@ -13,7 +14,7 @@ from PIL import Image
from pydantic import BaseModel from pydantic import BaseModel
from tenacity import ( from tenacity import (
retry, retry,
retry_if_exception_type, retry_if_exception,
stop_after_attempt, stop_after_attempt,
wait_exponential, wait_exponential,
) )
@@ -26,6 +27,16 @@ from odoo_ocr.utils import cache
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
def _is_transient_retryable(exc: BaseException) -> bool:
"""Retry network errors, timeouts, HTTP 429, and server errors (5xx)."""
if isinstance(exc, (httpx.NetworkError, httpx.TimeoutException)):
return True
if isinstance(exc, httpx.HTTPStatusError):
status = exc.response.status_code
return status == 429 or status >= 500
return False
class OllamaClient(BaseVLMClient): class OllamaClient(BaseVLMClient):
"""Client for Ollama's /api/chat endpoint with vision support.""" """Client for Ollama's /api/chat endpoint with vision support."""
@@ -35,6 +46,10 @@ class OllamaClient(BaseVLMClient):
self.timeout = httpx.Timeout(120.0, connect=10.0) self.timeout = httpx.Timeout(120.0, connect=10.0)
self.http = httpx.AsyncClient(timeout=self.timeout) self.http = httpx.AsyncClient(timeout=self.timeout)
async def aclose(self) -> None:
"""Release the underlying HTTP client."""
await self.http.aclose()
@staticmethod @staticmethod
def _encode_image(image: Image.Image, quality: int = 85) -> str: def _encode_image(image: Image.Image, quality: int = 85) -> str:
"""Convert a PIL image to a base64-encoded PNG or JPEG string.""" """Convert a PIL image to a base64-encoded PNG or JPEG string."""
@@ -62,6 +77,7 @@ class OllamaClient(BaseVLMClient):
model_name = model or self.settings.models.ocr model_name = model or self.settings.models.ocr
encoded_images = [self._encode_image(img) for img in images] encoded_images = [self._encode_image(img) for img in images]
image_blobs = [base64.b64decode(enc) for enc in encoded_images] image_blobs = [base64.b64decode(enc) for enc in encoded_images]
format_key = response_format.__name__ if response_format is not None else "text"
cached = cache.get_cached( cached = cache.get_cached(
system_prompt, system_prompt,
@@ -69,7 +85,9 @@ class OllamaClient(BaseVLMClient):
image_blobs, image_blobs,
model_name, model_name,
temperature, temperature,
max_tokens,
self.settings, self.settings,
format_key=format_key,
) )
if cached: if cached:
return VLMResponse.model_validate(cached) return VLMResponse.model_validate(cached)
@@ -93,7 +111,12 @@ class OllamaClient(BaseVLMClient):
if response_format is not None: if response_format is not None:
try: try:
body["format"] = response_format.model_json_schema() body["format"] = response_format.model_json_schema()
except Exception: except Exception as exc:
logger.warning(
"Could not build JSON schema for %s, falling back to 'json' format: %s",
response_format.__name__,
exc,
)
body["format"] = "json" body["format"] = "json"
start = time.perf_counter() start = time.perf_counter()
@@ -118,15 +141,25 @@ class OllamaClient(BaseVLMClient):
image_blobs, image_blobs,
model_name, model_name,
temperature, temperature,
max_tokens,
vlm_response.model_dump(mode="json"), vlm_response.model_dump(mode="json"),
self.settings, self.settings,
format_key=format_key,
)
logger.info(
"Ollama %s: %d ms, prompt_tokens=%s, completion_tokens=%s, prompt_hash=%s",
model_name,
duration_ms,
vlm_response.prompt_tokens,
vlm_response.completion_tokens,
hashlib.sha256((system_prompt + user_prompt).encode("utf-8")).hexdigest(),
) )
return vlm_response return vlm_response
@retry( @retry(
stop=stop_after_attempt(3), stop=stop_after_attempt(3),
wait=wait_exponential(multiplier=1, min=1, max=10), wait=wait_exponential(multiplier=1, min=1, max=10),
retry=retry_if_exception_type((httpx.NetworkError, httpx.TimeoutException)), retry=retry_if_exception(_is_transient_retryable),
) )
async def _post_chat(self, body: dict[str, Any]) -> dict[str, Any]: async def _post_chat(self, body: dict[str, Any]) -> dict[str, Any]:
"""Send the chat request and validate the response.""" """Send the chat request and validate the response."""
+25 -8
View File
@@ -26,21 +26,37 @@ def _cache_path(settings: Settings, key: str) -> Path:
return cache_dir / f"{key}.json" return cache_dir / f"{key}.json"
def _cache_key_parts(
system_prompt: str,
user_prompt: str,
image_blobs: list[bytes],
model: str,
temperature: float,
max_tokens: int,
format_key: str,
) -> list[str]:
parts = [system_prompt, user_prompt, model, str(temperature), str(max_tokens), format_key]
for blob in image_blobs:
parts.append(hashlib.sha256(blob).hexdigest())
return parts
def get_cached( def get_cached(
system_prompt: str, system_prompt: str,
user_prompt: str, user_prompt: str,
image_blobs: list[bytes], image_blobs: list[bytes],
model: str, model: str,
temperature: float, temperature: float,
max_tokens: int,
settings: Settings, settings: Settings,
format_key: str = "text",
) -> dict[str, Any] | None: ) -> dict[str, Any] | None:
"""Return cached response dict if it exists, otherwise None.""" """Return cached response dict if it exists, otherwise None."""
if not settings.cache.enabled: if not settings.cache.enabled:
return None return None
parts = [system_prompt, user_prompt, model, str(temperature)] key = _make_key(
for blob in image_blobs: _cache_key_parts(system_prompt, user_prompt, image_blobs, model, temperature, max_tokens, format_key)
parts.append(hashlib.sha256(blob).hexdigest()) )
key = _make_key(parts)
path = _cache_path(settings, key) path = _cache_path(settings, key)
if path.exists(): if path.exists():
logger.debug("Cache hit: %s", key) logger.debug("Cache hit: %s", key)
@@ -55,16 +71,17 @@ def set_cached(
image_blobs: list[bytes], image_blobs: list[bytes],
model: str, model: str,
temperature: float, temperature: float,
max_tokens: int,
response: dict[str, Any], response: dict[str, Any],
settings: Settings, settings: Settings,
format_key: str = "text",
) -> None: ) -> None:
"""Store a response in the file cache.""" """Store a response in the file cache."""
if not settings.cache.enabled: if not settings.cache.enabled:
return return
parts = [system_prompt, user_prompt, model, str(temperature)] key = _make_key(
for blob in image_blobs: _cache_key_parts(system_prompt, user_prompt, image_blobs, model, temperature, max_tokens, format_key)
parts.append(hashlib.sha256(blob).hexdigest()) )
key = _make_key(parts)
path = _cache_path(settings, key) path = _cache_path(settings, key)
with path.open("w", encoding="utf-8") as f: with path.open("w", encoding="utf-8") as f:
json.dump(response, f, ensure_ascii=False, indent=2) json.dump(response, f, ensure_ascii=False, indent=2)
+4
View File
@@ -52,3 +52,7 @@ class FakeVLMClient(BaseVLMClient):
if key in system_prompt or key in user_prompt: if key in system_prompt or key in user_prompt:
return VLMResponse(content=json.dumps(value), model="fake") return VLMResponse(content=json.dumps(value), model="fake")
return VLMResponse(content=json.dumps({}), model="fake") return VLMResponse(content=json.dumps({}), model="fake")
async def aclose(self) -> None:
"""The fake client holds no resources."""
return None
+44
View File
@@ -0,0 +1,44 @@
"""Tests for the LLM response cache."""
from __future__ import annotations
import pytest
from odoo_ocr.config import Settings
from odoo_ocr.utils import cache
@pytest.fixture
def cache_settings(tmp_path) -> Settings:
return Settings(cache={"enabled": True, "dir": tmp_path / "llm_cache"})
def test_cache_roundtrip(cache_settings: Settings) -> None:
response = {"content": "hello", "model": "m", "prompt_tokens": 1}
cache.set_cached("sys", "user", [b"img"], "model", 0.1, 4096, response, cache_settings)
assert cache.get_cached("sys", "user", [b"img"], "model", 0.1, 4096, cache_settings) == response
def test_cache_miss(cache_settings: Settings) -> None:
assert cache.get_cached("a", "b", [], "m", 0.1, 4096, cache_settings) is None
def test_cache_disabled(tmp_path) -> None:
settings = Settings(cache={"enabled": False, "dir": tmp_path / "llm_cache"})
cache.set_cached("sys", "user", [], "m", 0.1, 4096, {"content": "x"}, settings)
assert cache.get_cached("sys", "user", [], "m", 0.1, 4096, settings) is None
assert not (tmp_path / "llm_cache").exists()
def test_cache_key_includes_temperature_format_and_max_tokens(cache_settings: Settings) -> None:
response = {"content": "x"}
cache.set_cached("sys", "user", [], "m", 0.1, 4096, response, cache_settings, format_key="text")
# Different temperature must not hit the same entry.
assert cache.get_cached("sys", "user", [], "m", 0.2, 4096, cache_settings) is None
# Different response_format must not hit the same entry.
assert (
cache.get_cached("sys", "user", [], "m", 0.1, 4096, cache_settings, format_key="ExtractedInvoice")
is None
)
# Different max_tokens must not hit the same entry.
assert cache.get_cached("sys", "user", [], "m", 0.1, 2048, cache_settings) is None