From 00d6d36728b76ee39177830fcd681a10a4e7bbc3 Mon Sep 17 00:00:00 2001 From: JUNZE <127123086+WizisCool@users.noreply.github.com> Date: Thu, 6 Aug 2026 03:43:32 +0800 Subject: [PATCH] fix: make embedding batch size configurable (#983) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: resolve tiktoken encoding without constructing the embedding client EmbeddingClient.encoding forced full client construction, which raises 'OpenAI API key is required' even though tiktoken needs no credentials. The document dedup tie-break (src/crud/document.py) only needs .encoding for token counting, so any test hitting that path fails in environments without embedding keys — notably CI for pull requests from forks, where repo secrets are unavailable (e.g. #908's test-python job failing on tests/crud/test_document.py::test_duplicate_rejection_reinforces_existing). Resolve the encoding from the configured model directly, falling back to cl100k_base, and only reuse the underlying client's encoding when it has already been constructed. Co-Authored-By: Claude Fable 5 * fix: make embedding batch size configurable Add optional max_batch_size to the embedding model config (EMBEDDING_MODEL_CONFIG__MAX_BATCH_SIZE) to cap texts per request for OpenAI-compatible providers with smaller limits than OpenAI's, such as DashScope text-embedding-v4 (10) and Alibaba Bailian qwen3.7-text-embedding (20). When unset, native provider defaults are preserved (OpenAI 2048, Gemini 100). Fixes #687. * test(embedding): cover Gemini batching and config fallbacks per review - Gemini transport now tested for configured batch splitting and the 100 default fallback - OpenAI unset default (2048, single request) explicitly covered - env-parsing test now asserts the value survives resolve_embedding_model_config - docs: 100 is the client's conservative Gemini default, not a native limit * test(embedding): assert provider batch-size defaults --------- Co-authored-by: adavyas Co-authored-by: Claude Fable 5 --- .env.template | 1 + config.toml.example | 3 + docs/v3/contributing/configuration.mdx | 7 ++ src/config.py | 3 + src/embedding_client.py | 5 +- tests/llm/test_embedding_client.py | 146 ++++++++++++++++++++++++- 6 files changed, 162 insertions(+), 3 deletions(-) diff --git a/.env.template b/.env.template index 2516a239..2f737dba 100644 --- a/.env.template +++ b/.env.template @@ -21,6 +21,7 @@ PERFORMANCE_LOG_FORMAT=compact # compact|rich # EMBEDDING_MAX_TOKENS_PER_REQUEST=300000 # EMBEDDING_MODEL_CONFIG__TRANSPORT=openai # EMBEDDING_MODEL_CONFIG__MODEL=text-embedding-3-small +# EMBEDDING_MODEL_CONFIG__MAX_BATCH_SIZE=10 # EMBEDDING_MODEL_CONFIG__OVERRIDES__BASE_URL= # EMBEDDING_MODEL_CONFIG__OVERRIDES__API_KEY_ENV= diff --git a/config.toml.example b/config.toml.example index 42146da7..c5517999 100644 --- a/config.toml.example +++ b/config.toml.example @@ -74,6 +74,9 @@ MAX_TOKENS_PER_REQUEST = 300000 [embedding.model_config] transport = "openai" model = "text-embedding-3-small" +# Optional provider request input cap. Useful for OpenAI-compatible embedding +# APIs with smaller limits, such as DashScope text-embedding-v4. +# max_batch_size = 10 # Optional module-level endpoint overrides # [embedding.model_config.overrides] diff --git a/docs/v3/contributing/configuration.mdx b/docs/v3/contributing/configuration.mdx index 384c7841..bae1a416 100644 --- a/docs/v3/contributing/configuration.mdx +++ b/docs/v3/contributing/configuration.mdx @@ -266,12 +266,19 @@ EMBEDDING_MAX_TOKENS_PER_REQUEST=300000 # Embedding transport/model selection EMBEDDING_MODEL_CONFIG__TRANSPORT=openai # openai, gemini EMBEDDING_MODEL_CONFIG__MODEL=text-embedding-3-small +EMBEDDING_MODEL_CONFIG__MAX_BATCH_SIZE=10 # optional per-request input cap # Optional endpoint overrides EMBEDDING_MODEL_CONFIG__OVERRIDES__BASE_URL=http://localhost:8000/v1 EMBEDDING_MODEL_CONFIG__OVERRIDES__API_KEY_ENV=EMBEDDING_CUSTOM_API_KEY ``` +`EMBEDDING_MODEL_CONFIG__MAX_BATCH_SIZE` defaults to 2048 for OpenAI. For +Gemini the client applies a conservative default of 100 — Gemini does not +document a per-request limit. Set it when an OpenAI-compatible embedding +provider accepts fewer inputs per request, such as DashScope +`text-embedding-v4` with a limit of 10. + Forwarding `dimensions=` to OpenAI-compatible providers is controlled by `EMBEDDING_MODEL_CONFIG__DIMENSIONS_MODE`: - `auto` (default): forwards `dimensions=` when **the operator has explicitly set `EMBEDDING_VECTOR_DIMENSIONS`** — provenance, not value — and the configured model is not on the known-rejecting list (currently `text-embedding-ada-002`). Explicit `EMBEDDING_VECTOR_DIMENSIONS=1536` *does* trigger the forward; this is how `text-embedding-3-large` truncation to 1536 is expressed. Deployments that leave the setting unset get their existing behavior (`dimensions=` is not forwarded). diff --git a/src/config.py b/src/config.py index 177bef05..34a5917b 100644 --- a/src/config.py +++ b/src/config.py @@ -386,6 +386,7 @@ class ConfiguredEmbeddingModelSettings(BaseModel): transport: EmbeddingTransport = "openai" overrides: ModelOverrideSettings = Field(default_factory=ModelOverrideSettings) dimensions_mode: EmbeddingDimensionsMode = "auto" + max_batch_size: Annotated[int, Field(gt=0)] | None = None @model_validator(mode="before") @classmethod @@ -422,6 +423,7 @@ class EmbeddingModelConfig(BaseModel): transport: EmbeddingTransport = "openai" api_key: str | None = None base_url: str | None = None + max_batch_size: Annotated[int, Field(gt=0)] | None = None @model_validator(mode="before") @classmethod @@ -546,6 +548,7 @@ def resolve_embedding_model_config( transport=configured.transport, api_key=api_key, base_url=configured.overrides.base_url, + max_batch_size=configured.max_batch_size, ) diff --git a/src/embedding_client.py b/src/embedding_client.py index 92edab68..dc54065f 100644 --- a/src/embedding_client.py +++ b/src/embedding_client.py @@ -197,7 +197,7 @@ class _EmbeddingClient: # Gemini has a 2048 token limit self.max_embedding_tokens: int = min(max_input_tokens, 2048) # Gemini batch size is not documented, using conservative estimate - self.max_batch_size: int = 100 + self.max_batch_size: int = config.max_batch_size or 100 else: # openai if not config.api_key: raise ValueError("OpenAI API key is required") @@ -206,7 +206,7 @@ class _EmbeddingClient: base_url=config.base_url, ) self.max_embedding_tokens = max_input_tokens - self.max_batch_size = 2048 # OpenAI batch limit + self.max_batch_size = config.max_batch_size or 2048 try: self.encoding: tiktoken.Encoding = tiktoken.encoding_for_model(self.model) @@ -626,6 +626,7 @@ class EmbeddingClient: runtime_config.model, runtime_config.api_key, runtime_config.base_url, + runtime_config.max_batch_size, settings.EMBEDDING.VECTOR_DIMENSIONS, settings.EMBEDDING.MAX_INPUT_TOKENS, settings.EMBEDDING.MAX_TOKENS_PER_REQUEST, diff --git a/tests/llm/test_embedding_client.py b/tests/llm/test_embedding_client.py index 0d6ac3dd..e8ad30ab 100644 --- a/tests/llm/test_embedding_client.py +++ b/tests/llm/test_embedding_client.py @@ -3,7 +3,7 @@ from typing import Any, cast import pytest -from src.config import EmbeddingModelConfig +from src.config import EmbeddingModelConfig, resolve_embedding_model_config from src.embedding_client import _EmbeddingClient # pyright: ignore[reportPrivateUsage] @@ -194,6 +194,7 @@ def _build_openai_client( model: str, send_dimensions: bool, vector_dimensions: int, + max_batch_size: int | None = None, ) -> tuple[_EmbeddingClient, FakeOpenAIEmbeddingsAPI]: fake_embeddings = FakeOpenAIEmbeddingsAPI(embedding) @@ -210,6 +211,7 @@ def _build_openai_client( transport="openai", model=model, api_key="test-key", + max_batch_size=max_batch_size, ), vector_dimensions=vector_dimensions, max_input_tokens=8192, @@ -278,6 +280,133 @@ async def test_openai_simple_batch_embed_forwards_dimensions( assert fake.calls[0]["input"] == ["a", "b"] +@pytest.mark.asyncio +async def test_openai_simple_batch_embed_respects_configured_max_batch_size( + monkeypatch: pytest.MonkeyPatch, +) -> None: + client, fake = _build_openai_client( + monkeypatch, + embedding=[0.1] * 1536, + model="text-embedding-3-small", + send_dimensions=False, + vector_dimensions=1536, + max_batch_size=2, + ) + + await client.simple_batch_embed(["a", "b", "c"]) + + assert [call["input"] for call in fake.calls] == [["a", "b"], ["c"]] + + +@pytest.mark.asyncio +async def test_openai_simple_batch_embed_defaults_to_2048_when_unset( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Unset max_batch_size must keep the OpenAI default: one request.""" + client, fake = _build_openai_client( + monkeypatch, + embedding=[0.1] * 1536, + model="text-embedding-3-small", + send_dimensions=False, + vector_dimensions=1536, + ) + assert client.max_batch_size == 2048 + + await client.simple_batch_embed(["a", "b", "c"]) + + assert [call["input"] for call in fake.calls] == [["a", "b", "c"]] + + +@pytest.mark.asyncio +async def test_gemini_simple_batch_embed_respects_configured_max_batch_size( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Gemini transport must split batches at the configured limit too.""" + calls: list[dict[str, Any]] = [] + + class FakeGeminiModels: + async def embed_content( + self, + *, + model: str, + contents: str | list[str], + config: dict[str, Any], + ) -> SimpleNamespace: + calls.append({"model": model, "contents": contents, "config": config}) + n = len(contents) if isinstance(contents, list) else 1 + return SimpleNamespace( + embeddings=[SimpleNamespace(values=[0.2] * 12) for _ in range(n)] + ) + + class FakeGeminiClient: + def __init__(self, *, api_key: str | None, http_options: Any) -> None: + self.aio: Any = SimpleNamespace(models=FakeGeminiModels()) + + monkeypatch.setattr("src.embedding_client.genai.Client", FakeGeminiClient) + + client = _EmbeddingClient( + EmbeddingModelConfig( + transport="gemini", + model="gemini-embedding-001", + api_key="gemini-key", + max_batch_size=2, + ), + vector_dimensions=12, + max_input_tokens=4096, + max_tokens_per_request=300_000, + send_dimensions=False, + ) + + await client.simple_batch_embed(["a", "b", "c"]) + + assert [call["contents"] for call in calls] == [["a", "b"], ["c"]] + + +@pytest.mark.asyncio +async def test_gemini_simple_batch_embed_defaults_to_100_when_unset( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Unset max_batch_size must keep the Gemini conservative default.""" + calls: list[dict[str, Any]] = [] + + class FakeGeminiModels: + async def embed_content( + self, + *, + model: str, + contents: str | list[str], + config: dict[str, Any], + ) -> SimpleNamespace: + calls.append({"model": model, "contents": contents, "config": config}) + n = len(contents) if isinstance(contents, list) else 1 + return SimpleNamespace( + embeddings=[SimpleNamespace(values=[0.2] * 12) for _ in range(n)] + ) + + class FakeGeminiClient: + def __init__(self, *, api_key: str | None, http_options: Any) -> None: + self.aio: Any = SimpleNamespace(models=FakeGeminiModels()) + + monkeypatch.setattr("src.embedding_client.genai.Client", FakeGeminiClient) + + client = _EmbeddingClient( + EmbeddingModelConfig( + transport="gemini", + model="gemini-embedding-001", + api_key="gemini-key", + ), + vector_dimensions=12, + max_input_tokens=4096, + max_tokens_per_request=300_000, + send_dimensions=False, + ) + assert client.max_batch_size == 100 + + await client.simple_batch_embed(["a", "b", "c"]) + + assert [call["contents"] for call in calls] == [["a", "b", "c"]] + + @pytest.mark.asyncio async def test_openai_batch_embed_forwards_dimensions( monkeypatch: pytest.MonkeyPatch, @@ -308,6 +437,7 @@ def _build_embedding_settings( "EMBEDDING_MODEL_CONFIG__MODEL", "EMBEDDING_MODEL_CONFIG__TRANSPORT", "EMBEDDING_MODEL_CONFIG__DIMENSIONS_MODE", + "EMBEDDING_MODEL_CONFIG__MAX_BATCH_SIZE", ): monkeypatch.delenv(key, raising=False) for key, value in env.items(): @@ -481,3 +611,17 @@ def test_prepare_chunks_returns_ordered_chunks( assert len(out["long"]) > 1 # Order preserved assert isinstance(out["long"][0], str) + + +def test_embedding_model_config_parses_max_batch_size_from_env( + monkeypatch: pytest.MonkeyPatch, +) -> None: + s = _build_embedding_settings( + {"EMBEDDING_MODEL_CONFIG__MAX_BATCH_SIZE": "10"}, + monkeypatch, + ) + + assert s.MODEL_CONFIG.max_batch_size == 10 + + resolved = resolve_embedding_model_config(s.MODEL_CONFIG) + assert resolved.max_batch_size == 10