From 57f273e0001f3804db228d4bc6d57ca87d997dd1 Mon Sep 17 00:00:00 2001 From: Yiorgis Gozadinos Date: Thu, 14 May 2026 15:41:57 +0300 Subject: [PATCH] Minor fixes, CI should build --- .github/workflows/test.yml | 2 ++ haiku_rag_slim/haiku/rag/client/downloads.py | 1 + haiku_rag_slim/haiku/rag/reranking/__init__.py | 7 ++++++- haiku_rag_slim/haiku/rag/reranking/cross_encoder.py | 5 ++--- haiku_rag_slim/haiku/rag/reranking/jina_local.py | 6 +++++- haiku_rag_slim/haiku/rag/reranking/mxbai.py | 6 +++++- tests/test_reranker.py | 9 +++++++++ 7 files changed, 30 insertions(+), 6 deletions(-) diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index aa71c1bf..14facbd5 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -66,6 +66,8 @@ jobs: key: huggingface-${{ runner.os }}-qwen-tokenizer-v1 - name: Pre-download tokenizer run: uv run python -c "from transformers import AutoTokenizer; AutoTokenizer.from_pretrained('Qwen/Qwen3-Embedding-0.6B')" + - name: Pre-download cross-encoder test model + run: uv run python -c "from sentence_transformers import CrossEncoder; CrossEncoder('cross-encoder/ms-marco-MiniLM-L-6-v2')" - name: Run tests with coverage run: uv run pytest -m "not integration" --cov=haiku --cov-report=xml - name: Upload coverage to Codecov diff --git a/haiku_rag_slim/haiku/rag/client/downloads.py b/haiku_rag_slim/haiku/rag/client/downloads.py index 73d82b20..019abc52 100644 --- a/haiku_rag_slim/haiku/rag/client/downloads.py +++ b/haiku_rag_slim/haiku/rag/client/downloads.py @@ -58,6 +58,7 @@ async def download_models( model_name = config.embeddings.model.name yield DownloadProgress(model=model_name, status="start") + # Wrap in lambda: ty loses ParamSpec inference on third-party __init__. await asyncio.to_thread(lambda: SentenceTransformer(model_name)) yield DownloadProgress(model=model_name, status="done") except ImportError: diff --git a/haiku_rag_slim/haiku/rag/reranking/__init__.py b/haiku_rag_slim/haiku/rag/reranking/__init__.py index 4a098118..85b7b803 100644 --- a/haiku_rag_slim/haiku/rag/reranking/__init__.py +++ b/haiku_rag_slim/haiku/rag/reranking/__init__.py @@ -71,7 +71,12 @@ def get_reranker(config: AppConfig = Config) -> RerankerBase | None: try: from haiku.rag.reranking.cross_encoder import CrossEncoderReranker - return CrossEncoderReranker(config.reranking.model.name) + name = config.reranking.model.name + if not name: + raise ValueError( + "cross-encoder reranker requires name in reranking.model" + ) + return CrossEncoderReranker(name) except ImportError: # pragma: no cover return None diff --git a/haiku_rag_slim/haiku/rag/reranking/cross_encoder.py b/haiku_rag_slim/haiku/rag/reranking/cross_encoder.py index f1d81eb7..87ad2d38 100644 --- a/haiku_rag_slim/haiku/rag/reranking/cross_encoder.py +++ b/haiku_rag_slim/haiku/rag/reranking/cross_encoder.py @@ -33,8 +33,7 @@ class CrossEncoderReranker(RerankerBase): return [] documents = [chunk.content for chunk in chunks] - loop = asyncio.get_running_loop() - rankings = await loop.run_in_executor( - None, lambda: self._reranker.rank(query, documents, top_k=top_n) + rankings = await asyncio.to_thread( + lambda: self._reranker.rank(query, documents, top_k=top_n) ) return [(chunks[r["corpus_id"]], float(r["score"])) for r in rankings] diff --git a/haiku_rag_slim/haiku/rag/reranking/jina_local.py b/haiku_rag_slim/haiku/rag/reranking/jina_local.py index b76e6e64..9ad8aa4f 100644 --- a/haiku_rag_slim/haiku/rag/reranking/jina_local.py +++ b/haiku_rag_slim/haiku/rag/reranking/jina_local.py @@ -1,3 +1,5 @@ +import asyncio + try: from transformers import ( AutoModel, # pyright: ignore[reportMissingImports] @@ -32,6 +34,8 @@ class JinaLocalReranker(RerankerBase): # pragma: no cover documents = [chunk.content for chunk in chunks] - results = self._reranker.rerank(query, documents, top_n=top_n) + results = await asyncio.to_thread( + lambda: self._reranker.rerank(query, documents, top_n=top_n) + ) return [(chunks[r["index"]], float(r["relevance_score"])) for r in results] diff --git a/haiku_rag_slim/haiku/rag/reranking/mxbai.py b/haiku_rag_slim/haiku/rag/reranking/mxbai.py index 6e42f1a4..d73ed751 100644 --- a/haiku_rag_slim/haiku/rag/reranking/mxbai.py +++ b/haiku_rag_slim/haiku/rag/reranking/mxbai.py @@ -1,3 +1,5 @@ +import asyncio + from mxbai_rerank import MxbaiRerankV2 # pyright: ignore[reportMissingImports] from haiku.rag.config import Config @@ -22,7 +24,9 @@ class MxBAIReranker(RerankerBase): documents = [chunk.content for chunk in chunks] - results = self._client.rank(query=query, documents=documents, top_k=top_n) + results = await asyncio.to_thread( + lambda: self._client.rank(query=query, documents=documents, top_k=top_n) + ) reranked_chunks = [] for result in results: original_chunk = chunks[result.index] diff --git a/tests/test_reranker.py b/tests/test_reranker.py index a23af853..a578dc26 100644 --- a/tests/test_reranker.py +++ b/tests/test_reranker.py @@ -126,6 +126,15 @@ class TestGetReranker: with pytest.raises(ValueError, match="vLLM reranker requires base_url"): get_reranker(config) + def test_cross_encoder_provider_without_name_raises_error(self): + config = AppConfig( + reranking=RerankingConfig( + model=ModelConfig(provider="cross-encoder", name="") + ) + ) + with pytest.raises(ValueError, match="cross-encoder reranker requires name"): + get_reranker(config) + @pytest.mark.parametrize( "provider, model_name, class_module, class_name, extra_model_kwargs, expected_attrs, env_vars", [