-
Notifications
You must be signed in to change notification settings - Fork 261
feat: add OrcaRouter as a named embedding provider #1252
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,74 @@ | ||
| """OrcaRouter-based embedding provider for cloud or API-backed semantic indexing.""" | ||
|
|
||
| from __future__ import annotations | ||
|
|
||
| import os | ||
| from typing import Any, override | ||
|
|
||
| from basic_memory.repository.openai_provider import OpenAIEmbeddingProvider | ||
| from basic_memory.repository.semantic_errors import SemanticDependenciesMissingError | ||
|
|
||
| ORCAROUTER_DEFAULT_BASE_URL = "https://api.orcarouter.ai/v1" | ||
| ORCAROUTER_DEFAULT_MODEL = "openai/text-embedding-3-small" | ||
|
|
||
|
|
||
| class OrcaRouterEmbeddingProvider(OpenAIEmbeddingProvider): | ||
| """Embedding provider backed by OrcaRouter's OpenAI-compatible embeddings API. | ||
|
|
||
| OrcaRouter is an OpenAI-compatible model routing gateway. This provider points | ||
| the OpenAI-compatible embedding client at ``https://api.orcarouter.ai/v1`` and | ||
| authenticates with ``ORCAROUTER_API_KEY`` (keys start with ``sk-orca-``). | ||
| Model ids use the gateway's ``provider/model`` form, e.g. ``openai/text-embedding-3-small``. | ||
| """ | ||
|
|
||
| def __init__( | ||
| self, | ||
| model_name: str = ORCAROUTER_DEFAULT_MODEL, | ||
| *, | ||
| batch_size: int = 64, | ||
| request_concurrency: int = 4, | ||
| dimensions: int = 1536, | ||
| api_key: str | None = None, | ||
| base_url: str | None = None, | ||
| timeout: float = 30.0, | ||
| ) -> None: | ||
| super().__init__( | ||
| model_name=model_name, | ||
| batch_size=batch_size, | ||
| request_concurrency=request_concurrency, | ||
| dimensions=dimensions, | ||
| api_key=api_key, | ||
| base_url=base_url or ORCAROUTER_DEFAULT_BASE_URL, | ||
| timeout=timeout, | ||
|
Comment on lines
+39
to
+42
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
When Useful? React with 👍 / 👎. |
||
| ) | ||
|
|
||
| @override | ||
| async def _get_client(self) -> Any: | ||
| if self._client is not None: | ||
| return self._client | ||
|
|
||
| async with self._client_lock: | ||
| if self._client is not None: | ||
| return self._client | ||
|
|
||
| try: | ||
| from openai import AsyncOpenAI | ||
| except ImportError as exc: # pragma: no cover - covered via monkeypatch tests | ||
| raise SemanticDependenciesMissingError( | ||
| "OpenAI dependency is missing. " | ||
| "Install/update basic-memory to include semantic dependencies: " | ||
| "pip install -U basic-memory" | ||
| ) from exc | ||
|
|
||
| api_key = self._api_key or os.getenv("ORCAROUTER_API_KEY") | ||
| if not api_key: | ||
| raise SemanticDependenciesMissingError( | ||
| "OrcaRouter embedding provider requires ORCAROUTER_API_KEY." | ||
| ) | ||
|
|
||
| self._client = AsyncOpenAI( | ||
| api_key=api_key, | ||
| base_url=self._base_url, | ||
| timeout=self._timeout, | ||
| ) | ||
| return self._client | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,220 @@ | ||
| """Tests for OrcaRouterEmbeddingProvider and its embedding provider factory branch.""" | ||
|
|
||
| import builtins | ||
| import sys | ||
| from types import SimpleNamespace | ||
|
|
||
| import pytest | ||
|
|
||
| from basic_memory.config import BasicMemoryConfig | ||
| from basic_memory.repository.embedding_provider_factory import ( | ||
| create_embedding_provider, | ||
| reset_embedding_provider_cache, | ||
| ) | ||
| from basic_memory.repository.orcarouter_provider import ( | ||
| ORCAROUTER_DEFAULT_BASE_URL, | ||
| ORCAROUTER_DEFAULT_MODEL, | ||
| OrcaRouterEmbeddingProvider, | ||
| ) | ||
| from basic_memory.repository.semantic_errors import SemanticDependenciesMissingError | ||
|
|
||
|
|
||
| class _StubEmbeddingsApi: | ||
| def __init__(self): | ||
| self.calls: list[tuple[str, list[str]]] = [] | ||
|
|
||
| async def create(self, *, model: str, input: list[str]): | ||
| self.calls.append((model, input)) | ||
| vectors = [] | ||
| for index, value in enumerate(input): | ||
| base = float(len(value)) | ||
| vectors.append(SimpleNamespace(index=index, embedding=[base, base + 1.0, base + 2.0])) | ||
| return SimpleNamespace(data=vectors) | ||
|
|
||
|
|
||
| class _StubAsyncOpenAI: | ||
| init_count = 0 | ||
|
|
||
| def __init__(self, *, api_key: str, base_url=None, timeout=30.0): | ||
| self.api_key = api_key | ||
| self.base_url = base_url | ||
| self.timeout = timeout | ||
| self.embeddings = _StubEmbeddingsApi() | ||
| _StubAsyncOpenAI.init_count += 1 | ||
|
|
||
|
|
||
| @pytest.fixture(autouse=True) | ||
| def _reset_embedding_provider_cache_fixture(): | ||
| reset_embedding_provider_cache() | ||
| yield | ||
| reset_embedding_provider_cache() | ||
|
|
||
|
|
||
| def _install_stub_openai(monkeypatch) -> None: | ||
| module = type(sys)("openai") | ||
| setattr(module, "AsyncOpenAI", _StubAsyncOpenAI) | ||
| monkeypatch.setitem(sys.modules, "openai", module) | ||
|
|
||
|
|
||
| def _make_config(**overrides) -> BasicMemoryConfig: | ||
| defaults = { | ||
| "env": "test", | ||
| "projects": {"test-project": "/tmp/basic-memory-test"}, | ||
| "default_project": "test-project", | ||
| "semantic_search_enabled": True, | ||
| } | ||
| defaults.update(overrides) | ||
| return BasicMemoryConfig(**defaults) | ||
|
|
||
|
|
||
| # --- Provider behavior -------------------------------------------------------- | ||
|
|
||
|
|
||
| @pytest.mark.asyncio | ||
| async def test_orcarouter_provider_lazy_loads_and_reuses_client(monkeypatch): | ||
| """Provider should instantiate AsyncOpenAI lazily, use OrcaRouter base URL, and reuse a single client.""" | ||
| _install_stub_openai(monkeypatch) | ||
| monkeypatch.setenv("ORCAROUTER_API_KEY", "sk-orca-test") | ||
| _StubAsyncOpenAI.init_count = 0 | ||
|
|
||
| provider = OrcaRouterEmbeddingProvider( | ||
| model_name=ORCAROUTER_DEFAULT_MODEL, batch_size=2, dimensions=3 | ||
| ) | ||
| assert provider._client is None | ||
|
|
||
| first = await provider.embed_query("auth query") | ||
| second = await provider.embed_documents(["queue task", "relation sync"]) | ||
|
|
||
| assert _StubAsyncOpenAI.init_count == 1 | ||
| assert provider._client is not None | ||
| client = provider._client | ||
| assert client.base_url == ORCAROUTER_DEFAULT_BASE_URL | ||
| assert client.api_key == "sk-orca-test" | ||
| assert len(first) == 3 | ||
| assert len(second) == 2 | ||
| assert len(second[0]) == 3 | ||
|
|
||
|
|
||
| @pytest.mark.asyncio | ||
| async def test_orcarouter_provider_respects_explicit_api_key_and_base_url(monkeypatch): | ||
| """Explicit api_key/base_url should win over env/defaults.""" | ||
| _install_stub_openai(monkeypatch) | ||
| monkeypatch.setenv("ORCAROUTER_API_KEY", "sk-orca-env") | ||
| _StubAsyncOpenAI.init_count = 0 | ||
|
|
||
| provider = OrcaRouterEmbeddingProvider( | ||
| model_name=ORCAROUTER_DEFAULT_MODEL, | ||
| api_key="sk-orca-explicit", | ||
| base_url="https://custom.example/v1", | ||
| dimensions=3, | ||
| ) | ||
| await provider.embed_query("test") | ||
|
|
||
| assert provider._client is not None | ||
| client = provider._client | ||
| assert client.api_key == "sk-orca-explicit" | ||
| assert client.base_url == "https://custom.example/v1" | ||
|
|
||
|
|
||
| @pytest.mark.asyncio | ||
| async def test_orcarouter_provider_dimension_mismatch_raises_error(monkeypatch): | ||
| """Provider should fail fast when response dimensions differ from configured dimensions.""" | ||
| _install_stub_openai(monkeypatch) | ||
| monkeypatch.setenv("ORCAROUTER_API_KEY", "sk-orca-test") | ||
|
|
||
| provider = OrcaRouterEmbeddingProvider(dimensions=2) | ||
| with pytest.raises(RuntimeError, match="3-dimensional vectors"): | ||
| await provider.embed_documents(["semantic note"]) | ||
|
|
||
|
|
||
| @pytest.mark.asyncio | ||
| async def test_orcarouter_provider_missing_dependency_raises_actionable_error(monkeypatch): | ||
| """Missing openai package should raise SemanticDependenciesMissingError.""" | ||
| monkeypatch.delitem(sys.modules, "openai", raising=False) | ||
| monkeypatch.setenv("ORCAROUTER_API_KEY", "sk-orca-test") | ||
| original_import = builtins.__import__ | ||
|
|
||
| def _raising_import(name, globals=None, locals=None, fromlist=(), level=0): | ||
| if name == "openai": | ||
| raise ImportError("openai not installed") | ||
| return original_import(name, globals, locals, fromlist, level) | ||
|
|
||
| monkeypatch.setattr(builtins, "__import__", _raising_import) | ||
|
|
||
| provider = OrcaRouterEmbeddingProvider(model_name=ORCAROUTER_DEFAULT_MODEL) | ||
| with pytest.raises(SemanticDependenciesMissingError) as error: | ||
| await provider.embed_query("test") | ||
|
|
||
| assert "pip install -U basic-memory" in str(error.value) | ||
|
|
||
|
|
||
| @pytest.mark.asyncio | ||
| async def test_orcarouter_provider_missing_api_key_raises_error(monkeypatch): | ||
| """ORCAROUTER_API_KEY is required unless api_key is passed explicitly.""" | ||
| _install_stub_openai(monkeypatch) | ||
| monkeypatch.delenv("ORCAROUTER_API_KEY", raising=False) | ||
|
|
||
| provider = OrcaRouterEmbeddingProvider(model_name=ORCAROUTER_DEFAULT_MODEL) | ||
| with pytest.raises(SemanticDependenciesMissingError) as error: | ||
| await provider.embed_query("test") | ||
|
|
||
| assert "ORCAROUTER_API_KEY" in str(error.value) | ||
|
|
||
|
|
||
| # --- Factory selection -------------------------------------------------------- | ||
|
|
||
|
|
||
| def test_embedding_provider_factory_selects_orcarouter_and_applies_default_model(): | ||
| """Factory should map local default model to OrcaRouter default when provider is orcarouter.""" | ||
| config = _make_config( | ||
| semantic_embedding_provider="orcarouter", | ||
| semantic_embedding_model="bge-small-en-v1.5", | ||
| ) | ||
| provider = create_embedding_provider(config) | ||
| assert isinstance(provider, OrcaRouterEmbeddingProvider) | ||
| assert provider.model_name == ORCAROUTER_DEFAULT_MODEL | ||
| assert provider._base_url == ORCAROUTER_DEFAULT_BASE_URL | ||
|
|
||
|
|
||
| def test_embedding_provider_factory_orcarouter_uses_default_dimensions(): | ||
| """Factory should use OrcaRouter default 1536 dimensions when unset.""" | ||
| config = _make_config(semantic_embedding_provider="orcarouter") | ||
| provider = create_embedding_provider(config) | ||
| assert isinstance(provider, OrcaRouterEmbeddingProvider) | ||
| assert provider.dimensions == 1536 | ||
|
|
||
|
|
||
| def test_embedding_provider_factory_passes_custom_dimensions_to_orcarouter(): | ||
| """Factory should forward semantic_embedding_dimensions to the OrcaRouter provider.""" | ||
| config = _make_config( | ||
| semantic_embedding_provider="orcarouter", | ||
| semantic_embedding_dimensions=3072, | ||
| ) | ||
| provider = create_embedding_provider(config) | ||
| assert isinstance(provider, OrcaRouterEmbeddingProvider) | ||
| assert provider.dimensions == 3072 | ||
|
|
||
|
|
||
| def test_embedding_provider_factory_orcarouter_forwards_request_concurrency(): | ||
| """Factory should forward provider request concurrency for API-backed batching.""" | ||
| config = _make_config( | ||
| semantic_embedding_provider="orcarouter", | ||
| semantic_embedding_request_concurrency=6, | ||
| ) | ||
| provider = create_embedding_provider(config) | ||
| assert isinstance(provider, OrcaRouterEmbeddingProvider) | ||
| assert provider.request_concurrency == 6 | ||
|
|
||
|
|
||
| def test_embedding_provider_identity_orcarouter(): | ||
| """configured_embedding_provider_identity should name OrcaRouterEmbeddingProvider.""" | ||
| from basic_memory.repository.embedding_provider_factory import ( | ||
| configured_embedding_provider_identity, | ||
| ) | ||
|
|
||
| config = _make_config( | ||
| semantic_embedding_provider="orcarouter", | ||
| semantic_embedding_model="openai/text-embedding-3-small", | ||
| ) | ||
| identity = configured_embedding_provider_identity(config) | ||
| assert identity == "OrcaRouterEmbeddingProvider:openai/text-embedding-3-small:1536" |
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
When
semantic_embedding_modelselects any OrcaRouter model whose native output size is not 1536 andsemantic_embedding_dimensionsis omitted, this branch silently constructs the provider with 1536 dimensions. Unlike the LiteLLM branch below, it does not require dimensions for custom models, so the inherited response validation fails on the first embedding request instead of rejecting the invalid configuration up front. Require an explicit dimension for non-default OrcaRouter models or resolve a model-specific default.AGENTS.md reference: AGENTS.md:L132-L133
Useful? React with 👍 / 👎.