From d08360ca214cbcf473642900ab2175b49e31ed6a Mon Sep 17 00:00:00 2001 From: Jon Saad-Falcon <41205309+jonsaadfalcon@users.noreply.github.com> Date: Wed, 25 Mar 2026 13:15:06 -0700 Subject: [PATCH] feat: wire gemma_cpp engine into discovery and optional imports Co-Authored-By: Claude Sonnet 4.6 --- src/openjarvis/engine/__init__.py | 2 +- src/openjarvis/engine/_discovery.py | 12 ++++++++++ tests/engine/test_gemma_cpp.py | 34 +++++++++++++++++++++++++++++ 3 files changed, 47 insertions(+), 1 deletion(-) diff --git a/src/openjarvis/engine/__init__.py b/src/openjarvis/engine/__init__.py index b6f79eae..90486fe6 100644 --- a/src/openjarvis/engine/__init__.py +++ b/src/openjarvis/engine/__init__.py @@ -15,7 +15,7 @@ from openjarvis.engine._base import ( from openjarvis.engine._discovery import discover_engines, discover_models, get_engine # Optional engines — only register if their SDK deps are present -for _optional in ("cloud", "litellm"): +for _optional in ("cloud", "litellm", "gemma_cpp"): try: importlib.import_module(f".{_optional}", __name__) except ImportError: diff --git a/src/openjarvis/engine/_discovery.py b/src/openjarvis/engine/_discovery.py index 4e2fc695..83021858 100644 --- a/src/openjarvis/engine/_discovery.py +++ b/src/openjarvis/engine/_discovery.py @@ -25,12 +25,24 @@ _HOST_MAP: Dict[str, str | None] = { "apple_fm": "apple_fm_host", "cloud": None, "litellm": None, + "gemma_cpp": None, } def _make_engine(key: str, config: JarvisConfig) -> InferenceEngine: """Instantiate a registered engine with the appropriate config host.""" cls = EngineRegistry.get(key) + + # gemma_cpp: pass config fields instead of host + if key == "gemma_cpp": + cfg = config.engine.gemma_cpp + return cls( + model_path=cfg.model_path or None, + tokenizer_path=cfg.tokenizer_path or None, + model_type=cfg.model_type or None, + num_threads=cfg.num_threads, + ) + host_attr = _HOST_MAP.get(key) if host_attr is not None: host = getattr(config.engine, host_attr, None) diff --git a/tests/engine/test_gemma_cpp.py b/tests/engine/test_gemma_cpp.py index f7dffe22..34d7e7d6 100644 --- a/tests/engine/test_gemma_cpp.py +++ b/tests/engine/test_gemma_cpp.py @@ -323,3 +323,37 @@ class TestGemmaCppConfigResolution: assert engine._tokenizer_path == "" assert engine._model_type == "" assert engine._num_threads == 0 + + +class TestGemmaCppDiscovery: + def test_host_map_contains_gemma_cpp(self) -> None: + from openjarvis.engine._discovery import _HOST_MAP + assert "gemma_cpp" in _HOST_MAP + assert _HOST_MAP["gemma_cpp"] is None + + def test_make_engine_passes_config(self) -> None: + from openjarvis.core.config import GemmaCppEngineConfig, JarvisConfig + from openjarvis.core.registry import EngineRegistry + from openjarvis.engine._discovery import _make_engine + from openjarvis.engine.gemma_cpp import GemmaCppEngine + + EngineRegistry.register_value("gemma_cpp", GemmaCppEngine) + config = JarvisConfig() + config.engine.gemma_cpp = GemmaCppEngineConfig( + model_path="/cfg/model.sbs", + tokenizer_path="/cfg/tokenizer.spm", + model_type="9b-it", + num_threads=4, + ) + engine = _make_engine("gemma_cpp", config) + assert engine._model_path == "/cfg/model.sbs" + assert engine._tokenizer_path == "/cfg/tokenizer.spm" + assert engine._model_type == "9b-it" + assert engine._num_threads == 4 + + def test_registry_contains_gemma_cpp(self) -> None: + from openjarvis.core.registry import EngineRegistry + from openjarvis.engine.gemma_cpp import GemmaCppEngine + + EngineRegistry.register_value("gemma_cpp", GemmaCppEngine) + assert EngineRegistry.contains("gemma_cpp")