diff --git a/src/metis/providers/llamacpp.py b/src/metis/providers/llamacpp.py index 096c884a..edd19e9b 100644 --- a/src/metis/providers/llamacpp.py +++ b/src/metis/providers/llamacpp.py @@ -62,6 +62,7 @@ def count_tokens(self, text: str, model: str | None = None) -> int: class LlamaCppEmbeddingProvider(OpenAICompatibleEmbeddingProvider): DEFAULT_BASE_URL = LlamaCppProvider.DEFAULT_BASE_URL DEFAULT_API_KEY = LlamaCppProvider.DEFAULT_API_KEY + DEFAULT_CHECK_EMBEDDING_CTX_LENGTH = False CONFIG_SPEC = ProviderConfigSpec( display_name="llama.cpp embeddings", required_keys=("code_embedding_model", "docs_embedding_model"), @@ -75,6 +76,3 @@ class LlamaCppEmbeddingProvider(OpenAICompatibleEmbeddingProvider): "docs_extra_kwargs", ), ) - - def __init__(self, config): - super().__init__(config) diff --git a/src/metis/providers/openai_compatible.py b/src/metis/providers/openai_compatible.py index 0c38fda8..1d222c9f 100644 --- a/src/metis/providers/openai_compatible.py +++ b/src/metis/providers/openai_compatible.py @@ -106,6 +106,7 @@ def get_chat_model( class OpenAICompatibleEmbeddingProvider(EmbeddingProvider): DEFAULT_BASE_URL: str | None = None DEFAULT_API_KEY: str | None = None + DEFAULT_CHECK_EMBEDDING_CTX_LENGTH: bool | None = None def __init__(self, config: OpenAICompatibleEmbeddingConfig) -> None: self.config = config @@ -154,6 +155,10 @@ def _build_embedding_model( params["base_url"] = self.base_url if self.default_headers: params["default_headers"] = self.default_headers + if self.DEFAULT_CHECK_EMBEDDING_CTX_LENGTH is not None: + params["check_embedding_ctx_length"] = ( + self.DEFAULT_CHECK_EMBEDDING_CTX_LENGTH + ) if extra_kwargs: params.update(extra_kwargs) diff --git a/tests/test_openai_compatible_provider.py b/tests/test_openai_compatible_provider.py index 0bb0688a..a17a6bf1 100644 --- a/tests/test_openai_compatible_provider.py +++ b/tests/test_openai_compatible_provider.py @@ -6,6 +6,7 @@ from langchain_openai import ChatOpenAI from metis.providers.embedding_adapter import LangChainEmbeddingAdapter +from metis.providers.llamacpp import LlamaCppEmbeddingProvider from metis.providers.openai_compatible import OpenAICompatibleChatConfig from metis.providers.openai_compatible import OpenAICompatibleChatProvider from metis.providers.openai_compatible import OpenAICompatibleEmbeddingConfig @@ -103,3 +104,10 @@ def test_embedding_provider_builds_separate_code_and_docs_models() -> None: assert code_client.openai_api_base == "https://example.test/v1" assert code_client.default_headers == {"X-Test-Header": "test"} assert code_client.dimensions == 1536 + assert code_client.check_embedding_ctx_length is True + llama_provider = LlamaCppEmbeddingProvider(_embedding_config()) + for embedding in ( + llama_provider.get_embed_model_code(), + llama_provider.get_embed_model_docs(), + ): + assert cast(Any, embedding._client).check_embedding_ctx_length is False