From 969347c50e399af296f9a6dc81fd1691ac31296a Mon Sep 17 00:00:00 2001 From: Erlis Lushtaku <59629249+ErlisLushtaku@users.noreply.github.com> Date: Wed, 9 Sep 2026 14:11:16 +0200 Subject: [PATCH 1/2] Cache inference at do_inference with lazy PreparedModel. --- judgearena/inference.py | 210 ++++++++++++++++++++++++++++++++++ judgearena/models.py | 109 +++++++++++++++++- tests/test_inference_cache.py | 106 +++++++++++++++++ 3 files changed, 424 insertions(+), 1 deletion(-) create mode 100644 judgearena/inference.py create mode 100644 tests/test_inference_cache.py diff --git a/judgearena/inference.py b/judgearena/inference.py new file mode 100644 index 0000000..2a2e24d --- /dev/null +++ b/judgearena/inference.py @@ -0,0 +1,210 @@ +"""Lazy model preparation and inference-cache context.""" + +from __future__ import annotations + +import json +from abc import ABC, abstractmethod +from collections.abc import Callable +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, ClassVar + +import pandas as pd + +from judgearena.cache_sqlite import ( + COMPLETION_DB_NAME, + JUDGEMENT_DB_NAME, + CacheKind, + CompletionCache, + JudgementCache, + cache_folder, + stable_json_dumps, + write_descriptor, +) + +_ROLE_MAP = {"human": "user", "ai": "assistant", "system": "system"} + + +def canonicalize_chat_input(input_item: Any) -> str: + """Serialize a logical model input for content-addressed cache lookup.""" + if isinstance(input_item, str): + payload = {"type": "text", "text": input_item} + elif hasattr(input_item, "to_messages"): + payload = { + "type": "messages", + "messages": [ + { + "role": _ROLE_MAP.get(message.type, message.type), + "content": message.content, + } + for message in input_item.to_messages() + ], + } + else: + raise TypeError(f"Unsupported inference input: {type(input_item)!r}") + return stable_json_dumps(payload) + + +def build_model_descriptor( + provider: str, + model_name: str, + resolved_kwargs: dict[str, Any], +) -> dict[str, Any] | None: + """Describe output-affecting settings without constructing the backend.""" + if provider != "Dummy": + return None + return { + "schema_version": "judgearena-inference-cache/v1", + "provider": provider, + "model": model_name, + "input_mode": "chat", + "model_kwargs": resolved_kwargs, + } + + +@dataclass(frozen=True) +class CachedInferenceResult: + """Provider output fields required by downstream parsing.""" + + text: str + first_token_top_logprobs: dict[str, float] | None = None + + +@dataclass +class PreparedModel: + """Carry cache identity while deferring backend construction until a miss.""" + + model_spec: str + descriptor: dict[str, Any] | None + factory: Callable[[], Any] + cache: InferenceCache | None = None + _model: Any = field(default=None, init=False, repr=False) + + def materialize(self) -> Any: + if self._model is None: + self._model = self.factory() + return self._model + + +@dataclass(frozen=True) +class InferenceCache(ABC): + """Share cache lifecycle while subclasses define role-specific rows.""" + + store_root: Path + task: str + pushed_by: str = "judgearena" + + kind: ClassVar[CacheKind] + db_name: ClassVar[str] + store_type: ClassVar[type[CompletionCache] | type[JudgementCache]] + + def open_store(self, model: PreparedModel) -> CompletionCache | JudgementCache: + assert model.descriptor is not None + folder = cache_folder( + self.store_root, + self.kind, + self.task, + model.model_spec, + model.descriptor, + ) + write_descriptor(folder, model.descriptor) + return self.store_type(folder / self.db_name) + + def save_outputs( + self, + store: CompletionCache | JudgementCache, + model: PreparedModel, + input_texts: list[str], + outputs: list[Any], + metadata: list[dict[str, Any]], + indices: list[int], + ) -> None: + rows = [ + self.make_row( + model=model, + input_text=input_texts[index], + output=output, + metadata=metadata[index], + ) + for index, output in zip(indices, outputs, strict=True) + ] + store.save(pd.DataFrame(rows), pushed_by=self.pushed_by) + + @abstractmethod + def make_row( + self, + *, + model: PreparedModel, + input_text: str, + output: Any, + metadata: dict[str, Any], + ) -> dict[str, Any]: + """Convert one inference output to its role-specific storage row.""" + + @abstractmethod + def cached_result(self, row: pd.Series) -> CachedInferenceResult: + """Restore output fields from a stored row.""" + + +class CompletionInferenceCache(InferenceCache): + """Cache generated model completions.""" + + kind = "completions" + db_name = COMPLETION_DB_NAME + store_type = CompletionCache + + def make_row( + self, + *, + model: PreparedModel, + input_text: str, + output: Any, + metadata: dict[str, Any], + ) -> dict[str, Any]: + return { + "input_text": input_text, + "completion": output.text, + "benchmark": self.task, + "instruction_id": metadata["instruction_id"], + "model": model.model_spec, + } + + def cached_result(self, row: pd.Series) -> CachedInferenceResult: + return CachedInferenceResult(text=str(row["completion"])) + + +class JudgementInferenceCache(InferenceCache): + """Cache raw judge completions.""" + + kind = "judgements" + db_name = JUDGEMENT_DB_NAME + store_type = JudgementCache + + def make_row( + self, + *, + model: PreparedModel, + input_text: str, + output: Any, + metadata: dict[str, Any], + ) -> dict[str, Any]: + return { + "judge_input": input_text, + "judge_completion": output.text, + "benchmark": self.task, + "instruction_id": metadata["instruction_id"], + "model_a": metadata["model_a"], + "model_b": metadata["model_b"], + "judge": model.model_spec, + "top_logprobs": output.first_token_top_logprobs, + "orientation": metadata.get("orientation"), + } + + def cached_result(self, row: pd.Series) -> CachedInferenceResult: + top_logprobs = row["top_logprobs"] + return CachedInferenceResult( + text=str(row["judge_completion"]), + first_token_top_logprobs=( + json.loads(top_logprobs) if pd.notna(top_logprobs) else None + ), + ) diff --git a/judgearena/models.py b/judgearena/models.py index f949d5a..8e3adc5 100644 --- a/judgearena/models.py +++ b/judgearena/models.py @@ -16,7 +16,14 @@ from tqdm.asyncio import tqdm from tqdm.contrib.logging import logging_redirect_tqdm +from judgearena.cache_sqlite import input_hash from judgearena.constants import VLLM_REASONING_END_STR, VLLM_REASONING_START_STR +from judgearena.inference import ( + InferenceCache, + PreparedModel, + build_model_descriptor, + canonicalize_chat_input, +) from judgearena.log import get_logger from judgearena.usage import RequestUsage, RunUsage, record_usage from judgearena.utils.io import safe_parse_int @@ -663,7 +670,7 @@ def batch_inference_once( return [result.text for result in results] -def do_inference( +def _do_inference_uncached( chat_model, inputs, use_tqdm: bool = False, @@ -772,6 +779,78 @@ def batch_with_retry(batch_inputs, max_retries=5, base_delay=1.0): return [result.text for result in results] +def do_inference( + chat_model, + inputs, + use_tqdm: bool = False, + return_top_logprobs: bool = False, + *, + stage: str = "unspecified", + cache_metadata: list[dict] | None = None, +): + """Reuse raw provider outputs and invoke the backend only for cache misses.""" + inputs = list(inputs) + if not isinstance(chat_model, PreparedModel): + return _do_inference_uncached( + chat_model, + inputs, + use_tqdm, + return_top_logprobs, + stage=stage, + ) + + cache = chat_model.cache + if cache is None or chat_model.descriptor is None: + return _do_inference_uncached( + chat_model.materialize(), + inputs, + use_tqdm, + return_top_logprobs, + stage=stage, + ) + if cache_metadata is None or len(cache_metadata) != len(inputs): + raise ValueError("cache_metadata must contain one row per inference input.") + + input_texts = [canonicalize_chat_input(item) for item in inputs] + input_hashes = [input_hash(input_text) for input_text in input_texts] + with cache.open_store(chat_model) as store: + cached_rows = store.query(input_hashes).set_index("input_hash") + results: list[InferenceResult | None] = [ + ( + InferenceResult(**cache.cached_result(cached_rows.loc[key]).__dict__) + if key in cached_rows.index + else None + ) + for key in input_hashes + ] + missing_indices = [ + index for index, result in enumerate(results) if result is None + ] + if missing_indices: + generated = _do_inference_uncached( + chat_model.materialize(), + [inputs[index] for index in missing_indices], + use_tqdm, + True, + stage=stage, + ) + for index, result in zip(missing_indices, generated, strict=True): + results[index] = result + cache.save_outputs( + store, + chat_model, + input_texts, + generated, + cache_metadata, + missing_indices, + ) + + resolved_results = [result for result in results if result is not None] + if return_top_logprobs: + return resolved_results + return [result.text for result in resolved_results] + + def _route_sampling_params( engine_kwargs: dict, *, @@ -814,6 +893,34 @@ def _route_sampling_params( return engine_kwargs +def prepare_model( + model: str, + max_tokens: int | None = 8192, + *, + cache: InferenceCache | None = None, + **engine_kwargs, +) -> PreparedModel: + """Prepare cache identity without constructing the provider backend.""" + provider, model_name = _split_model_spec(model) + resolved_kwargs = {**engine_kwargs, "max_tokens": max_tokens or 8192} + descriptor = ( + build_model_descriptor(provider, model_name, resolved_kwargs) + if cache is not None + else None + ) + factory_kwargs = engine_kwargs.copy() + return PreparedModel( + model_spec=model, + descriptor=descriptor, + factory=lambda: make_model( + model, + max_tokens=max_tokens, + **factory_kwargs, + ), + cache=cache, + ) + + def make_model(model: str, max_tokens: int | None = 8192, **engine_kwargs): """Instantiate a model wrapper from a provider/model-name string. diff --git a/tests/test_inference_cache.py b/tests/test_inference_cache.py new file mode 100644 index 0000000..3343fb9 --- /dev/null +++ b/tests/test_inference_cache.py @@ -0,0 +1,106 @@ +from langchain_core.messages import AIMessage + +import judgearena.models as models +from judgearena.inference import CompletionInferenceCache, JudgementInferenceCache +from judgearena.models import InferenceResult, do_inference, prepare_model +from judgearena.usage import track_usage + + +class EchoModel: + def __init__(self): + self.calls = [] + + def batch(self, inputs, **_kwargs): + self.calls.append(inputs) + return [AIMessage(content=f"generated:{item}") for item in inputs] + + +def test_full_hit_does_not_materialize_model(tmp_path, monkeypatch): + cache = CompletionInferenceCache(tmp_path, "arena-hard") + monkeypatch.setattr(models, "make_model", lambda *_args, **_kwargs: EchoModel()) + metadata = [{"instruction_id": "1"}] + do_inference( + prepare_model("Dummy/test-model", cache=cache), + ["prompt"], + cache_metadata=metadata, + ) + + def fail_if_materialized(*_args, **_kwargs): + raise AssertionError("cache hit materialized the model") + + monkeypatch.setattr(models, "make_model", fail_if_materialized) + with track_usage() as tracker: + outputs = do_inference( + prepare_model("Dummy/test-model", cache=cache), + ["prompt"], + cache_metadata=metadata, + ) + + assert outputs == ["generated:prompt"] + assert tracker.snapshot().requests == () + + +def test_mixed_hits_and_misses_preserve_order(tmp_path, monkeypatch): + cache = CompletionInferenceCache(tmp_path, "arena-hard") + first_backend = EchoModel() + monkeypatch.setattr(models, "make_model", lambda *_args, **_kwargs: first_backend) + do_inference( + prepare_model("Dummy/test-model", cache=cache), + ["hit"], + cache_metadata=[{"instruction_id": "hit"}], + ) + + backend = EchoModel() + monkeypatch.setattr(models, "make_model", lambda *_args, **_kwargs: backend) + outputs = do_inference( + prepare_model("Dummy/test-model", cache=cache), + ["miss-a", "hit", "miss-b"], + cache_metadata=[ + {"instruction_id": "a"}, + {"instruction_id": "hit"}, + {"instruction_id": "b"}, + ], + ) + + assert outputs == ["generated:miss-a", "generated:hit", "generated:miss-b"] + assert backend.calls == [["miss-a", "miss-b"]] + + +def test_judgement_hit_preserves_top_logprobs(tmp_path, monkeypatch): + cache = JudgementInferenceCache(tmp_path, "arena-hard") + + class LogprobModel: + def batch(self, inputs, **_kwargs): + return [ + InferenceResult( + text="m", + first_token_top_logprobs={"m": -0.1, "M": -2.0}, + ) + for _ in inputs + ] + + monkeypatch.setattr(models, "make_model", lambda *_args, **_kwargs: LogprobModel()) + metadata = [ + { + "instruction_id": "1", + "model_a": "candidate", + "model_b": "baseline", + "orientation": "direct", + } + ] + first = do_inference( + prepare_model("Dummy/judge", cache=cache), + ["judge prompt"], + return_top_logprobs=True, + cache_metadata=metadata, + ) + second = do_inference( + prepare_model("Dummy/judge", cache=cache), + ["judge prompt"], + return_top_logprobs=True, + cache_metadata=metadata, + ) + + assert second[0].text == first[0].text + assert second[0].first_token_top_logprobs == first[0].first_token_top_logprobs + assert second[0].usage is None From 72987de7a5fef49b0ae36d467d2c6ac563196d57 Mon Sep 17 00:00:00 2001 From: Erlis Lushtaku <59629249+ErlisLushtaku@users.noreply.github.com> Date: Fri, 18 Sep 2026 16:00:46 +0200 Subject: [PATCH 2/2] refactor: consolidate InferenceResult and harden do_inference caching --- judgearena/inference.py | 48 ++++++++++++------ judgearena/models.py | 94 ++++++++++++++++++++++++----------- tests/test_inference_cache.py | 51 ++++++++++++++++--- 3 files changed, 143 insertions(+), 50 deletions(-) diff --git a/judgearena/inference.py b/judgearena/inference.py index 2a2e24d..463b995 100644 --- a/judgearena/inference.py +++ b/judgearena/inference.py @@ -7,7 +7,7 @@ from collections.abc import Callable from dataclasses import dataclass, field from pathlib import Path -from typing import Any, ClassVar +from typing import Any, ClassVar, NotRequired, TypedDict import pandas as pd @@ -21,6 +21,7 @@ stable_json_dumps, write_descriptor, ) +from judgearena.usage import RequestUsage _ROLE_MAP = {"human": "user", "ai": "assistant", "system": "system"} @@ -63,11 +64,26 @@ def build_model_descriptor( @dataclass(frozen=True) -class CachedInferenceResult: - """Provider output fields required by downstream parsing.""" +class InferenceResult: + """A text completion and optional provider response details.""" text: str first_token_top_logprobs: dict[str, float] | None = None + usage: RequestUsage | None = None + + +class CompletionCacheRowMetadata(TypedDict): + instruction_id: str + + +class JudgementCacheRowMetadata(TypedDict): + instruction_id: str + model_a: str + model_b: str | None + orientation: NotRequired[str | None] + + +CacheRowMetadata = CompletionCacheRowMetadata | JudgementCacheRowMetadata @dataclass @@ -77,7 +93,7 @@ class PreparedModel: model_spec: str descriptor: dict[str, Any] | None factory: Callable[[], Any] - cache: InferenceCache | None = None + cache: InferenceCache[Any] | None = None _model: Any = field(default=None, init=False, repr=False) def materialize(self) -> Any: @@ -87,7 +103,7 @@ def materialize(self) -> Any: @dataclass(frozen=True) -class InferenceCache(ABC): +class InferenceCache[CacheRowMetadataT: CacheRowMetadata](ABC): """Share cache lifecycle while subclasses define role-specific rows.""" store_root: Path @@ -116,7 +132,7 @@ def save_outputs( model: PreparedModel, input_texts: list[str], outputs: list[Any], - metadata: list[dict[str, Any]], + metadata: list[CacheRowMetadataT], indices: list[int], ) -> None: rows = [ @@ -137,16 +153,16 @@ def make_row( model: PreparedModel, input_text: str, output: Any, - metadata: dict[str, Any], + metadata: CacheRowMetadataT, ) -> dict[str, Any]: """Convert one inference output to its role-specific storage row.""" @abstractmethod - def cached_result(self, row: pd.Series) -> CachedInferenceResult: + def cached_result(self, row: pd.Series) -> InferenceResult: """Restore output fields from a stored row.""" -class CompletionInferenceCache(InferenceCache): +class CompletionInferenceCache(InferenceCache[CompletionCacheRowMetadata]): """Cache generated model completions.""" kind = "completions" @@ -159,7 +175,7 @@ def make_row( model: PreparedModel, input_text: str, output: Any, - metadata: dict[str, Any], + metadata: CompletionCacheRowMetadata, ) -> dict[str, Any]: return { "input_text": input_text, @@ -169,11 +185,11 @@ def make_row( "model": model.model_spec, } - def cached_result(self, row: pd.Series) -> CachedInferenceResult: - return CachedInferenceResult(text=str(row["completion"])) + def cached_result(self, row: pd.Series) -> InferenceResult: + return InferenceResult(text=str(row["completion"])) -class JudgementInferenceCache(InferenceCache): +class JudgementInferenceCache(InferenceCache[JudgementCacheRowMetadata]): """Cache raw judge completions.""" kind = "judgements" @@ -186,7 +202,7 @@ def make_row( model: PreparedModel, input_text: str, output: Any, - metadata: dict[str, Any], + metadata: JudgementCacheRowMetadata, ) -> dict[str, Any]: return { "judge_input": input_text, @@ -200,9 +216,9 @@ def make_row( "orientation": metadata.get("orientation"), } - def cached_result(self, row: pd.Series) -> CachedInferenceResult: + def cached_result(self, row: pd.Series) -> InferenceResult: top_logprobs = row["top_logprobs"] - return CachedInferenceResult( + return InferenceResult( text=str(row["judge_completion"]), first_token_top_logprobs=( json.loads(top_logprobs) if pd.notna(top_logprobs) else None diff --git a/judgearena/models.py b/judgearena/models.py index 8e3adc5..79b64de 100644 --- a/judgearena/models.py +++ b/judgearena/models.py @@ -6,20 +6,24 @@ import json import math import os +import sqlite3 import time import warnings from collections.abc import Mapping -from dataclasses import dataclass, replace +from dataclasses import replace from langchain_community.llms import LlamaCpp from langchain_openai import ChatOpenAI +from pandas.errors import DatabaseError from tqdm.asyncio import tqdm from tqdm.contrib.logging import logging_redirect_tqdm from judgearena.cache_sqlite import input_hash from judgearena.constants import VLLM_REASONING_END_STR, VLLM_REASONING_START_STR from judgearena.inference import ( + CacheRowMetadata, InferenceCache, + InferenceResult, PreparedModel, build_model_descriptor, canonicalize_chat_input, @@ -30,6 +34,7 @@ logger = get_logger(__name__) +_CACHE_OPERATION_ERRORS = (OSError, sqlite3.Error, DatabaseError) DEFAULT_VLLM_JUDGE_THINKING_TOKEN_BUDGET = 512 _THINKING_MODEL_PARSER_BY_SUBSTRING = ( @@ -482,15 +487,6 @@ async def ainvoke(self, input_item, **invoke_kwargs): ) -@dataclass(frozen=True) -class InferenceResult: - """A text completion and optional provider response details.""" - - text: str - first_token_top_logprobs: dict[str, float] | None = None - usage: RequestUsage | None = None - - def _first_token_top_logprobs(response) -> dict[str, float] | None: """Extract first-token top logprobs from a langchain AIMessage, if any.""" metadata = getattr(response, "response_metadata", None) or {} @@ -670,7 +666,7 @@ def batch_inference_once( return [result.text for result in results] -def _do_inference_uncached( +def _run_backend_inference( chat_model, inputs, use_tqdm: bool = False, @@ -786,12 +782,12 @@ def do_inference( return_top_logprobs: bool = False, *, stage: str = "unspecified", - cache_metadata: list[dict] | None = None, + cache_row_metadata: list[CacheRowMetadata] | None = None, ): """Reuse raw provider outputs and invoke the backend only for cache misses.""" inputs = list(inputs) if not isinstance(chat_model, PreparedModel): - return _do_inference_uncached( + return _run_backend_inference( chat_model, inputs, use_tqdm, @@ -801,23 +797,53 @@ def do_inference( cache = chat_model.cache if cache is None or chat_model.descriptor is None: - return _do_inference_uncached( + return _run_backend_inference( chat_model.materialize(), inputs, use_tqdm, return_top_logprobs, stage=stage, ) - if cache_metadata is None or len(cache_metadata) != len(inputs): - raise ValueError("cache_metadata must contain one row per inference input.") + if cache_row_metadata is None or len(cache_row_metadata) != len(inputs): + raise ValueError("cache_row_metadata must contain one row per inference input.") input_texts = [canonicalize_chat_input(item) for item in inputs] input_hashes = [input_hash(input_text) for input_text in input_texts] - with cache.open_store(chat_model) as store: - cached_rows = store.query(input_hashes).set_index("input_hash") + try: + store = cache.open_store(chat_model) + except _CACHE_OPERATION_ERRORS as exc: + logger.warning( + "Cache open failed at %s: %s. Continuing without caching.", + cache.store_root, + exc, + ) + return _run_backend_inference( + chat_model.materialize(), + inputs, + use_tqdm, + return_top_logprobs, + stage=stage, + ) + + try: + try: + cached_rows = store.query(input_hashes).set_index("input_hash") + except _CACHE_OPERATION_ERRORS as exc: + logger.warning( + "Cache read failed at %s: %s. Continuing without caching.", + store.db_path, + exc, + ) + return _run_backend_inference( + chat_model.materialize(), + inputs, + use_tqdm, + return_top_logprobs, + stage=stage, + ) results: list[InferenceResult | None] = [ ( - InferenceResult(**cache.cached_result(cached_rows.loc[key]).__dict__) + cache.cached_result(cached_rows.loc[key]) if key in cached_rows.index else None ) @@ -827,7 +853,7 @@ def do_inference( index for index, result in enumerate(results) if result is None ] if missing_indices: - generated = _do_inference_uncached( + generated = _run_backend_inference( chat_model.materialize(), [inputs[index] for index in missing_indices], use_tqdm, @@ -836,14 +862,26 @@ def do_inference( ) for index, result in zip(missing_indices, generated, strict=True): results[index] = result - cache.save_outputs( - store, - chat_model, - input_texts, - generated, - cache_metadata, - missing_indices, - ) + try: + cache.save_outputs( + store, + chat_model, + input_texts, + generated, + cache_row_metadata, + missing_indices, + ) + except _CACHE_OPERATION_ERRORS as exc: + logger.warning( + "Cache write failed at %s: %s. Preserving generated results.", + store.db_path, + exc, + ) + finally: + try: + store.close() + except _CACHE_OPERATION_ERRORS as exc: + logger.warning("Cache close failed at %s: %s.", store.db_path, exc) resolved_results = [result for result in results if result is not None] if return_top_logprobs: diff --git a/tests/test_inference_cache.py b/tests/test_inference_cache.py index 3343fb9..686e262 100644 --- a/tests/test_inference_cache.py +++ b/tests/test_inference_cache.py @@ -1,3 +1,6 @@ +import sqlite3 + +import pytest from langchain_core.messages import AIMessage import judgearena.models as models @@ -22,7 +25,7 @@ def test_full_hit_does_not_materialize_model(tmp_path, monkeypatch): do_inference( prepare_model("Dummy/test-model", cache=cache), ["prompt"], - cache_metadata=metadata, + cache_row_metadata=metadata, ) def fail_if_materialized(*_args, **_kwargs): @@ -33,7 +36,7 @@ def fail_if_materialized(*_args, **_kwargs): outputs = do_inference( prepare_model("Dummy/test-model", cache=cache), ["prompt"], - cache_metadata=metadata, + cache_row_metadata=metadata, ) assert outputs == ["generated:prompt"] @@ -47,7 +50,7 @@ def test_mixed_hits_and_misses_preserve_order(tmp_path, monkeypatch): do_inference( prepare_model("Dummy/test-model", cache=cache), ["hit"], - cache_metadata=[{"instruction_id": "hit"}], + cache_row_metadata=[{"instruction_id": "hit"}], ) backend = EchoModel() @@ -55,7 +58,7 @@ def test_mixed_hits_and_misses_preserve_order(tmp_path, monkeypatch): outputs = do_inference( prepare_model("Dummy/test-model", cache=cache), ["miss-a", "hit", "miss-b"], - cache_metadata=[ + cache_row_metadata=[ {"instruction_id": "a"}, {"instruction_id": "hit"}, {"instruction_id": "b"}, @@ -92,15 +95,51 @@ def batch(self, inputs, **_kwargs): prepare_model("Dummy/judge", cache=cache), ["judge prompt"], return_top_logprobs=True, - cache_metadata=metadata, + cache_row_metadata=metadata, ) second = do_inference( prepare_model("Dummy/judge", cache=cache), ["judge prompt"], return_top_logprobs=True, - cache_metadata=metadata, + cache_row_metadata=metadata, ) assert second[0].text == first[0].text assert second[0].first_token_top_logprobs == first[0].first_token_top_logprobs assert second[0].usage is None + + +def test_cache_write_failure_preserves_generated_results(tmp_path, monkeypatch): + cache = CompletionInferenceCache(tmp_path, "arena-hard") + backend = EchoModel() + monkeypatch.setattr(models, "make_model", lambda *_args, **_kwargs: backend) + + def fail_save(*_args, **_kwargs): + raise sqlite3.OperationalError("read-only") + + monkeypatch.setattr(CompletionInferenceCache, "save_outputs", fail_save) + + outputs = do_inference( + prepare_model("Dummy/test-model", cache=cache), + ["prompt"], + cache_row_metadata=[{"instruction_id": "1"}], + ) + + assert outputs == ["generated:prompt"] + assert backend.calls == [["prompt"]] + + +def test_cache_descriptor_validation_still_fails_loudly(tmp_path, monkeypatch): + cache = CompletionInferenceCache(tmp_path, "arena-hard") + + def fail_open(*_args, **_kwargs): + raise ValueError("descriptor mismatch") + + monkeypatch.setattr(CompletionInferenceCache, "open_store", fail_open) + + with pytest.raises(ValueError, match="descriptor mismatch"): + do_inference( + prepare_model("Dummy/test-model", cache=cache), + ["prompt"], + cache_row_metadata=[{"instruction_id": "1"}], + )