models: fix per-slot validation and cap mismatched-source refetches

This commit is contained in:
Jason Wen
2026-08-26 01:18:15 -04:00
parent 467bc9a172
commit 41f8d0e71a
3 changed files with 71 additions and 5 deletions
+6 -2
View File
@@ -155,6 +155,7 @@ class ModelFetcher:
for source, (_, suffix) in self.MODEL_SOURCES.items()
}
self.model_url = self.MODEL_URL
self._refetched: set[str] = set()
self._update_model_source()
@staticmethod
@@ -216,7 +217,9 @@ class ModelFetcher:
cached_data, is_expired = self.model_caches[source].get()
if cached_data and not is_expired:
if self._cache_matches_source(source, cached_data):
# a source is refetched over a mismatch at most once per process: if the fresh
# manifest still mismatches, the URL is authoritative and the cache is trusted
if self._cache_matches_source(source, cached_data) or source in self._refetched:
try:
parsed = self.model_parser.parse_models(cached_data)
except Exception:
@@ -229,7 +232,8 @@ class ModelFetcher:
# manifest version) - do not trust it, refetch so the source is repopulated
cloudlog.warning(f"Cached models for {source} have no valid bundles; refetching")
else:
cloudlog.warning(f"Cached models for {source} not valid; refetching")
self._refetched.add(source)
cloudlog.warning(f"Cached models for {source} not valid; refetching once")
fetched_bundles = self._fetch_and_cache_models(source)
if fetched_bundles is not None:
+3 -2
View File
@@ -173,15 +173,16 @@ def _validate_active_bundle(params: Params, source: str, available_bundles: list
if active_bundle is None or _bundle_needs_reset(active_bundle, available_bundles):
cloudlog.warning(f"Active model bundle invalid for {source}; resetting to default")
params.remove(key)
params.put("ModelRunnerTypeCache", int(custom.ModelManagerSP.Runner.stock), block=True)
_LAST_VALIDATED_RAW[key] = None
else:
_LAST_VALIDATED_RAW[key] = raw_bundle
def validate_active_bundles(params: Params, source_bundles: dict[str, list[custom.ModelManagerSP.ModelBundle]]) -> None:
# an empty list means the fetch failed, not that the catalog dropped the bundle
for source, bundles in source_bundles.items():
_validate_active_bundle(params, source, bundles)
_validate_active_bundle(params, source, bundles or None)
get_active_model_runner(params, force_check=True)
def get_active_model_runner(params: Params | None = None, force_check: bool = False) -> int:
@@ -25,7 +25,9 @@ from openpilot.common.file_chunker import get_chunk_name, get_manifest_path
from openpilot.selfdrive.test.helpers import http_server_context
from openpilot.sunnypilot.models import manager as manager_module
from openpilot.sunnypilot.models.fetcher import ModelFetcher, get_cached_bundles
from openpilot.sunnypilot.models.helpers import get_active_bundle, get_active_source, get_selected_bundle, resolve_bundle_by_ref
from openpilot.sunnypilot.models import helpers
from openpilot.sunnypilot.models.helpers import (get_active_bundle, get_active_source, get_selected_bundle,
resolve_bundle_by_ref, validate_active_bundles)
from openpilot.sunnypilot.models.manager import ModelManagerSP
CHUNK_BODIES = [b'A' * 5000, b'B' * 5000, b'C' * 3000]
@@ -532,6 +534,20 @@ class TestSourceCacheIntegrity(OpenpilotTestCase):
fetch.assert_called_once_with("qcom")
assert [bundle.ref for bundle in bundles] == ["ddd"]
def test_mismatched_refetch_happens_once(self):
"""If the fresh manifest still fails the source check, the URL is authoritative:
trust it instead of refetching at 1 Hz forever."""
params = self._make_params({"bundles": [manifest_bundle("big", "bbb", is_big=True)]},
{"bundles": [manifest_bundle("big2", "ccc", is_big=True)]})
fetcher = ModelFetcher(params)
fetched = self._fetched(manifest_bundle("big", "bbb", is_big=True))
with mock.patch.object(fetcher, "_fetch_and_cache_models", return_value=fetched) as fetch:
first = fetcher.get_bundles_for_source("qcom")
second = fetcher.get_bundles_for_source("qcom")
fetch.assert_called_once_with("qcom")
assert [bundle.ref for bundle in first] == ["bbb"]
assert [bundle.ref for bundle in second] == ["bbb"]
def test_corrupt_cache_is_refetched(self):
"""A cache that fails to parse (e.g. truncated/foreign JSON) must trigger a
refetch instead of raising every loop and never recovering."""
@@ -545,6 +561,51 @@ class TestSourceCacheIntegrity(OpenpilotTestCase):
assert [bundle.ref for bundle in bundles] == ["aaa"]
class TestActiveBundleValidation(OpenpilotTestCase):
"""Validation is per-slot: a failed fetch (empty bundle list) must not reset a slot,
and resetting one slot must not stomp the runner cache derived from the other."""
def setUp(self):
super().setUp()
helpers._LAST_VALIDATED_RAW.clear()
@staticmethod
def _raw_bundle(ref: str, runner: int | None = None) -> dict:
bundle = custom.ModelManagerSP.ModelBundle.new_message()
bundle.ref = ref
bundle.minimumSelectorVersion = 18
if runner is not None:
bundle.runner = runner
return bundle.to_dict()
def _params(self, qcom=None, usbgpu=None):
params = mock.MagicMock()
def get(key, *args, **kwargs):
return {"ModelManager_ActiveBundle": qcom, "ModelManager_ActiveBundleUSBGPU": usbgpu}.get(key)
params.get.side_effect = get
return params
def test_empty_catalog_does_not_reset_slot(self):
params = self._params(qcom=self._raw_bundle("small"))
with mock.patch("openpilot.sunnypilot.models.helpers.usbgpu_present", return_value=False):
validate_active_bundles(params, {"qcom": [], "usbgpu": []})
params.remove.assert_not_called()
def test_reset_recomputes_runner_from_surviving_slot(self):
tinygrad = int(custom.ModelManagerSP.Runner.tinygrad)
big_raw = self._raw_bundle("big", runner=tinygrad)
params = self._params(qcom=self._raw_bundle("gone"), usbgpu=big_raw)
catalog = {"qcom": [custom.ModelManagerSP.ModelBundle(**self._raw_bundle("other"))],
"usbgpu": [custom.ModelManagerSP.ModelBundle(**big_raw)]}
with mock.patch("openpilot.sunnypilot.models.helpers.usbgpu_present", return_value=True):
validate_active_bundles(params, catalog)
params.remove.assert_called_once_with("ModelManager_ActiveBundle")
runner_puts = [call for call in params.put.call_args_list if call.args[0] == "ModelRunnerTypeCache"]
assert [call.args[1] for call in runner_puts] == [tinygrad]
class TestActiveBundleSelection(OpenpilotTestCase):
"""The effective active bundle follows the hardware: the usbgpu slot wins when a GPU
is present, otherwise the qcom slot. Each slot keeps its own selection."""