diff --git a/src/freshdata/models/runtime.py b/src/freshdata/models/runtime.py index d1bbb9f4..9b53c6e9 100644 --- a/src/freshdata/models/runtime.py +++ b/src/freshdata/models/runtime.py @@ -17,7 +17,14 @@ import numpy as np from ._lazy import has_semantic_extra, require_onnxruntime, require_tokenizers -from .registry import COL_ENCODER_ID, get_config, is_installed, model_dir, verify +from .registry import ( + COL_ENCODER_ID, + get_config, + is_installed, + model_dir, + pinned_checksums, + verify, +) from .types import ModelError _STUB_ENV = "FRESHDATA_STUB_ENCODER" @@ -47,7 +54,9 @@ class OnnxEncoder: def __init__(self, model_id: str) -> None: cfg = get_config(model_id) self.model_id = model_id - self.model_sha256 = cfg.sha256 or "unverified" + # The primary file's pin, whether set via ``sha256`` or ``file_sha256``. + primary_pin = pinned_checksums(cfg).get(cfg.files[0]) if cfg.files else None + self.model_sha256 = primary_pin or "unverified" self.dim = 384 self._session: Any = None self._tokenizer: Any = None diff --git a/tests/test_models_runtime.py b/tests/test_models_runtime.py index f264b0b7..2f702a2f 100644 --- a/tests/test_models_runtime.py +++ b/tests/test_models_runtime.py @@ -2,10 +2,14 @@ from __future__ import annotations +import dataclasses +import hashlib + import numpy as np import pytest from freshdata.models import ModelError, runtime +from freshdata.models import registry as reg from freshdata.models.stub import StubEncoder from freshdata.semantic.cache import EmbeddingCache @@ -123,6 +127,52 @@ def test_embedding_cache_disabled(): assert len(cache) == 0 +_FAKE_MODEL = b"FAKE-ONNX-BYTES" +_FAKE_TOKENIZER = b'{"vocab": []}' + + +def _sha(payload: bytes) -> str: + return hashlib.sha256(payload).hexdigest() + + +def test_onnx_encoder_reports_primary_file_pin_from_file_sha256(monkeypatch): + # Constructing OnnxEncoder is lazy: no ONNX runtime or model files needed. + cfg = dataclasses.replace( + reg.get_config(reg.COL_ENCODER_ID), + sha256=None, + file_sha256=( + ("tokenizer.json", _sha(_FAKE_TOKENIZER)), + ("model.onnx", _sha(_FAKE_MODEL)), + ), + ) + monkeypatch.setitem(reg.REGISTRY, reg.COL_ENCODER_ID, cfg) + encoder = runtime.OnnxEncoder(reg.COL_ENCODER_ID) + assert encoder.model_sha256 == _sha(_FAKE_MODEL) + assert encoder._session is None # nothing was loaded + + +def test_onnx_encoder_reports_primary_pin_from_sha256(monkeypatch): + cfg = dataclasses.replace( + reg.get_config(reg.COL_ENCODER_ID), + sha256=_sha(_FAKE_MODEL), + file_sha256=(("tokenizer.json", _sha(_FAKE_TOKENIZER)),), + ) + monkeypatch.setitem(reg.REGISTRY, reg.COL_ENCODER_ID, cfg) + assert runtime.OnnxEncoder(reg.COL_ENCODER_ID).model_sha256 == _sha(_FAKE_MODEL) + + +def test_onnx_encoder_unpinned_reports_unverified(monkeypatch): + cfg = dataclasses.replace(reg.get_config(reg.COL_ENCODER_ID), sha256=None, file_sha256=()) + monkeypatch.setitem(reg.REGISTRY, reg.COL_ENCODER_ID, cfg) + assert runtime.OnnxEncoder(reg.COL_ENCODER_ID).model_sha256 == "unverified" + + +def test_onnx_encoder_shipped_registry_reports_unverified(): + # Today's registry pins nothing, so behaviour is unchanged. + for model_id in reg.REGISTRY: + assert runtime.OnnxEncoder(model_id).model_sha256 == "unverified" + + def test_onnx_encoder_requires_extra(): pytest.importorskip("onnxruntime") # With the extra installed but no model files, availability names the pull path.