diff --git a/src/embedding_client.py b/src/embedding_client.py index c75ed345..8b6493d2 100644 --- a/src/embedding_client.py +++ b/src/embedding_client.py @@ -196,10 +196,8 @@ class _EmbeddingClient: # - gemini-embedding-001: 2048 tokens # - gemini-embedding-2: 8192 tokens # Unknown models default conservatively to 2048. - if "gemini-embedding-2" in self.model: + if self.model == "gemini-embedding-2": gemini_model_token_cap = 8192 - elif "gemini-embedding-001" in self.model: - gemini_model_token_cap = 2048 else: gemini_model_token_cap = 2048 self.max_embedding_tokens: int = min( diff --git a/tests/llm/test_embedding_client.py b/tests/llm/test_embedding_client.py index ce156e4c..fe516102 100644 --- a/tests/llm/test_embedding_client.py +++ b/tests/llm/test_embedding_client.py @@ -197,6 +197,33 @@ def test_gemini_embedding_client_unknown_model_defaults_to_2048( assert client.max_embedding_tokens == 2048 +def test_gemini_embedding_client_near_miss_model_id_defaults_to_2048( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """A model id that merely contains 'gemini-embedding-2' as a substring + (e.g. 'gemini-embedding-20') must not be granted the 8192 cap reserved + for the exact 'gemini-embedding-2' model id.""" + + class FakeGeminiClient: + def __init__(self, *, api_key: str, http_options: Any) -> None: + self.api_key: str = api_key + + monkeypatch.setattr("src.embedding_client.genai.Client", FakeGeminiClient) + + client = _EmbeddingClient( + EmbeddingModelConfig( + transport="gemini", + model="gemini-embedding-20", + api_key="test-key", + ), + vector_dimensions=8, + max_input_tokens=20_000, + max_tokens_per_request=300_000, + send_dimensions=True, + ) + assert client.max_embedding_tokens == 2048 + + @pytest.mark.asyncio async def test_gemini_embedding_client_uses_output_dimensionality( monkeypatch: pytest.MonkeyPatch,