diff --git a/README.md b/README.md index 7ffc44854bb..ebf75e31729 100644 --- a/README.md +++ b/README.md @@ -401,7 +401,7 @@ Set `LITELLM_PROXY_API_BASE` and `LITELLM_PROXY_API_KEY` and every model call th | [Vercel AI Gateway (`vercel_ai_gateway`)](https://docs.litellm.ai/docs/providers/vercel_ai_gateway) | ✅ | ✅ | ✅ | | | | | | | | | [VLLM (`vllm`)](https://docs.litellm.ai/docs/providers/vllm) | ✅ | ✅ | ✅ | | | | | | | | | [Volcengine (`volcengine`)](https://docs.litellm.ai/docs/providers/volcano) | ✅ | ✅ | ✅ | | | | | | | | -| [Voyage AI (`voyage`)](https://docs.litellm.ai/docs/providers/voyage) | | | | ✅ | | | | | | | +| [VoyageAI by MongoDB (`voyage`)](https://docs.litellm.ai/docs/providers/voyage) | | | | ✅ | | | | | | | | [WandB Inference (`wandb`)](https://docs.litellm.ai/docs/providers/wandb_inference) | ✅ | ✅ | ✅ | | | | | | | | | [Watsonx Text (`watsonx_text`)](https://docs.litellm.ai/docs/providers/watsonx) | ✅ | ✅ | ✅ | | | | | | | | | [xAI (`xai`)](https://docs.litellm.ai/docs/providers/xai) | ✅ | ✅ | ✅ | | | | | | | | diff --git a/litellm/llms/voyage/common_utils.py b/litellm/llms/voyage/common_utils.py new file mode 100644 index 00000000000..2f4f63a8a43 --- /dev/null +++ b/litellm/llms/voyage/common_utils.py @@ -0,0 +1,35 @@ +""" +Shared helpers for the Voyage (VoyageAI by MongoDB) provider. +""" + +from typing import Final + +from litellm.secret_managers.main import get_secret_str + +VOYAGE_API_BASE: Final = "https://api.voyageai.com/v1" +MONGODB_API_BASE: Final = "https://ai.mongodb.com/v1" +MONGODB_API_KEY_PREFIX: Final = "al-" + + +def get_voyage_api_key(api_key: str | None = None) -> str | None: + """Resolve the key a Voyage request will authenticate with, explicit value first.""" + return ( + api_key + or get_secret_str("VOYAGE_API_KEY") + or get_secret_str("VOYAGE_AI_API_KEY") + or get_secret_str("VOYAGE_AI_TOKEN") + ) + + +def get_default_base_url(api_key: str | None = None) -> str: + """ + Pick the host that issued the key: MongoDB-issued keys (``al-`` prefix) are only + valid on ai.mongodb.com, every other key on api.voyageai.com. + + Mirrors ``voyageai.util.get_default_base_url`` in the official SDK: + https://github.com/voyage-ai/voyageai-python/blob/main/voyageai/util.py + """ + resolved: Final = get_voyage_api_key(api_key) + if resolved is not None and resolved.startswith(MONGODB_API_KEY_PREFIX): + return MONGODB_API_BASE + return VOYAGE_API_BASE diff --git a/litellm/llms/voyage/embedding/transformation.py b/litellm/llms/voyage/embedding/transformation.py index 7d74b1e00c4..a9232ad27e2 100644 --- a/litellm/llms/voyage/embedding/transformation.py +++ b/litellm/llms/voyage/embedding/transformation.py @@ -5,7 +5,7 @@ import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig -from litellm.secret_managers.main import get_secret_str +from litellm.llms.voyage.common_utils import get_default_base_url, get_voyage_api_key from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues from litellm.types.utils import EmbeddingResponse, Usage @@ -49,7 +49,7 @@ class VoyageEmbeddingConfig(BaseEmbeddingConfig): if not api_base.endswith("/embeddings"): api_base = f"{api_base}/embeddings" return api_base - return "https://api.voyageai.com/v1/embeddings" + return f"{get_default_base_url(api_key)}/embeddings" def get_supported_openai_params(self, model: str) -> list: return [ @@ -85,14 +85,8 @@ class VoyageEmbeddingConfig(BaseEmbeddingConfig): api_key: str | None = None, api_base: str | None = None, ) -> dict: - if api_key is None: - api_key = ( - get_secret_str("VOYAGE_API_KEY") - or get_secret_str("VOYAGE_AI_API_KEY") - or get_secret_str("VOYAGE_AI_TOKEN") - ) return { - "Authorization": f"Bearer {api_key}", + "Authorization": f"Bearer {get_voyage_api_key(api_key)}", } def transform_embedding_request( diff --git a/litellm/llms/voyage/embedding/transformation_contextual.py b/litellm/llms/voyage/embedding/transformation_contextual.py index 870de8756bb..f5e7fdb48f5 100644 --- a/litellm/llms/voyage/embedding/transformation_contextual.py +++ b/litellm/llms/voyage/embedding/transformation_contextual.py @@ -11,7 +11,7 @@ import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig -from litellm.secret_managers.main import get_secret_str +from litellm.llms.voyage.common_utils import get_default_base_url, get_voyage_api_key from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues from litellm.types.utils import EmbeddingResponse, Usage @@ -58,7 +58,7 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): if not api_base.endswith("/contextualizedembeddings"): api_base = f"{api_base}/contextualizedembeddings" return api_base - return "https://api.voyageai.com/v1/contextualizedembeddings" + return f"{get_default_base_url(api_key)}/contextualizedembeddings" def get_supported_openai_params(self, model: str) -> list: # mutable-ok: base class signature return ["encoding_format", "dimensions"] @@ -91,14 +91,8 @@ class VoyageContextualEmbeddingConfig(BaseEmbeddingConfig): api_key: str | None = None, api_base: str | None = None, ) -> dict: - if api_key is None: - api_key = ( - get_secret_str("VOYAGE_API_KEY") - or get_secret_str("VOYAGE_AI_API_KEY") - or get_secret_str("VOYAGE_AI_TOKEN") - ) return { - "Authorization": f"Bearer {api_key}", + "Authorization": f"Bearer {get_voyage_api_key(api_key)}", } AUTO_CHUNK_SIZE: Final = 32000 diff --git a/litellm/llms/voyage/embedding/transformation_multimodal.py b/litellm/llms/voyage/embedding/transformation_multimodal.py index 814d5ab7eb0..035765b6691 100644 --- a/litellm/llms/voyage/embedding/transformation_multimodal.py +++ b/litellm/llms/voyage/embedding/transformation_multimodal.py @@ -13,7 +13,7 @@ import httpx from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig -from litellm.secret_managers.main import get_secret_str +from litellm.llms.voyage.common_utils import get_default_base_url, get_voyage_api_key from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues from litellm.types.utils import EmbeddingResponse, Usage @@ -58,7 +58,7 @@ class VoyageMultimodalEmbeddingConfig(BaseEmbeddingConfig): if not api_base.endswith("/multimodalembeddings"): api_base = f"{api_base}/multimodalembeddings" return api_base - return "https://api.voyageai.com/v1/multimodalembeddings" + return f"{get_default_base_url(api_key)}/multimodalembeddings" def get_supported_openai_params(self, model: str) -> list: return ["dimensions"] @@ -84,19 +84,14 @@ class VoyageMultimodalEmbeddingConfig(BaseEmbeddingConfig): api_key: str | None = None, api_base: str | None = None, ) -> dict: - if api_key is None: - api_key = ( - get_secret_str("VOYAGE_API_KEY") - or get_secret_str("VOYAGE_AI_API_KEY") - or get_secret_str("VOYAGE_AI_TOKEN") - ) - if not api_key: + resolved_api_key: Final = get_voyage_api_key(api_key) + if not resolved_api_key: raise ValueError( "Voyage API key is required for multimodal embeddings. " "Set VOYAGE_API_KEY / VOYAGE_AI_API_KEY / VOYAGE_AI_TOKEN " "or pass `api_key` explicitly." ) - return {"Authorization": f"Bearer {api_key}"} + return {"Authorization": f"Bearer {resolved_api_key}"} def _normalize_content_item(self, item: dict[str, object]) -> dict[str, object]: item_type: Final = item.get("type") diff --git a/litellm/llms/voyage/rerank/transformation.py b/litellm/llms/voyage/rerank/transformation.py index 3acb2f2ed58..121c18e82ae 100644 --- a/litellm/llms/voyage/rerank/transformation.py +++ b/litellm/llms/voyage/rerank/transformation.py @@ -13,7 +13,7 @@ from pydantic import ConfigDict, TypeAdapter from litellm._uuid import uuid from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig -from litellm.secret_managers.main import get_secret_str +from litellm.llms.voyage.common_utils import get_default_base_url, get_voyage_api_key from litellm.types.rerank import ( RerankBilledUnits, RerankResponse, @@ -31,6 +31,16 @@ _STR: Final = TypeAdapter(str) class VoyageRerankConfig(BaseRerankConfig): + """ + ``validate_environment`` stores the credential it authenticates with so ``get_complete_url`` + can select the host that issued it. ``ProviderConfigManager.get_provider_rerank_config`` + builds this config per request, so that key never reaches another one. + """ + + def __init__(self) -> None: + super().__init__() + self._api_key: str | None = None + def get_supported_cohere_rerank_params(self, model: str) -> list: return ["query", "documents", "top_n", "return_documents"] @@ -66,7 +76,7 @@ class VoyageRerankConfig(BaseRerankConfig): optional_params: dict | None = None, ) -> str: if api_base is None: - return "https://api.voyageai.com/v1/rerank" + return f"{get_default_base_url(self._api_key)}/rerank" api_base = api_base.rstrip("/") if not api_base.endswith("/v1/rerank"): if api_base.endswith("/v1"): @@ -148,12 +158,12 @@ class VoyageRerankConfig(BaseRerankConfig): optional_params: dict | None = None, litellm_params: Mapping[str, object] | None = None, ) -> dict: - if api_key is None: - api_key = get_secret_str("VOYAGE_API_KEY") or get_secret_str("VOYAGE_AI_API_KEY") - if api_key is None: + resolved_api_key: Final = get_voyage_api_key(api_key) + if resolved_api_key is None: raise ValueError("Voyage AI API key is required. Set via `api_key` parameter or `VOYAGE_API_KEY` env var.") + self._api_key = resolved_api_key return { - "Authorization": f"Bearer {api_key}", + "Authorization": f"Bearer {resolved_api_key}", "content-type": "application/json", } diff --git a/litellm/provider_endpoints_support_backup.json b/litellm/provider_endpoints_support_backup.json index 0e325bb61fe..d9a6feb625d 100644 --- a/litellm/provider_endpoints_support_backup.json +++ b/litellm/provider_endpoints_support_backup.json @@ -2354,7 +2354,7 @@ } }, "voyage": { - "display_name": "Voyage AI (`voyage`)", + "display_name": "VoyageAI by MongoDB (`voyage`)", "url": "https://docs.litellm.ai/docs/providers/voyage", "endpoints": { "chat_completions": false, diff --git a/litellm/proxy/public_endpoints/provider_create_fields.json b/litellm/proxy/public_endpoints/provider_create_fields.json index ae07d5d6bd9..965d8dae4d2 100644 --- a/litellm/proxy/public_endpoints/provider_create_fields.json +++ b/litellm/proxy/public_endpoints/provider_create_fields.json @@ -3583,7 +3583,7 @@ }, { "provider": "Voyage", - "provider_display_name": "Voyage AI", + "provider_display_name": "VoyageAI by MongoDB", "litellm_provider": "voyage", "credential_fields": [ { diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 00904219b81..62e23388581 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2711,7 +2711,7 @@ } }, "voyage": { - "display_name": "Voyage AI (`voyage`)", + "display_name": "VoyageAI by MongoDB (`voyage`)", "url": "https://docs.litellm.ai/docs/providers/voyage", "endpoints": { "chat_completions": false, diff --git a/tests/integration/_support/forward_proxy.py b/tests/integration/_support/forward_proxy.py new file mode 100644 index 00000000000..9f5cff94277 --- /dev/null +++ b/tests/integration/_support/forward_proxy.py @@ -0,0 +1,71 @@ +from __future__ import annotations + +import threading +from collections.abc import Generator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from queue import SimpleQueue +from typing import Final + + +@dataclass(frozen=True, slots=True) +class Tunnel: + method: str + target: str + headers: Mapping[str, str] + + +@dataclass(frozen=True, slots=True) +class ForwardProxy: + url: str + port: int + received: SimpleQueue[Tunnel] + + def drain(self) -> tuple[Tunnel, ...]: + return tuple(self.received.get_nowait() for _ in range(self.received.qsize())) + + def targets(self) -> tuple[str, ...]: + return tuple(tunnel.target for tunnel in self.drain()) + + +@contextmanager +def refusing_forward_proxy(port: int = 0) -> Generator[ForwardProxy, None, None]: + """Owned HTTP forward proxy: records the host each CONNECT asks for and refuses the tunnel with 403. + + A litellm proxy booted with ``HTTPS_PROXY`` pointed here makes the host it dials for a provider + observable without a byte leaving the box; the caller sees litellm's own connection error. + """ + received: Final[SimpleQueue[Tunnel]] = SimpleQueue() + + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + timeout = 5 + + def refuse(self) -> None: + received.put(Tunnel(self.command, self.path, {name.lower(): value for name, value in self.headers.items()})) + self.send_response(403) + self.send_header("content-length", "0") + self.send_header("connection", "close") + self.end_headers() + self.close_connection = True + + do_CONNECT = refuse + do_GET = refuse + do_POST = refuse + do_PUT = refuse + do_DELETE = refuse + + def log_message(self, format: str, *args: object) -> None: + pass + + with ThreadingHTTPServer(("127.0.0.1", port), Handler) as server: + thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05}) + thread.start() + try: + yield ForwardProxy(f"http://127.0.0.1:{server.server_port}", server.server_port, received) + finally: + server.shutdown() + thread.join(timeout=6) + assert not thread.is_alive(), "Owned forward proxy survived cleanup" + server.server_close() diff --git a/tests/integration/providers/test_voyage_mongodb_host_wire.py b/tests/integration/providers/test_voyage_mongodb_host_wire.py new file mode 100644 index 00000000000..1a6ed907c9b --- /dev/null +++ b/tests/integration/providers/test_voyage_mongodb_host_wire.py @@ -0,0 +1,688 @@ +import asyncio +import json +import os +import re +import signal +import socket +import uuid +from collections.abc import Callable, Iterator +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import JSON_OBJECT, Gateway, eventually, gateway_from_environment, object_value +from integration._support.database import read_rows +from integration._support.forward_proxy import ForwardProxy, refusing_forward_proxy +from integration._support.process import owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request, wire_server +from openai import APIStatusError, AsyncOpenAI, OpenAI +from pydantic import JsonValue, TypeAdapter + +EMBEDDING_MODEL: Final = "voyage/voyage-3.5" +CONTEXTUAL_MODEL: Final = "voyage/voyage-context-3" +MULTIMODAL_MODEL: Final = "voyage/voyage-multimodal-3" +RERANK_MODEL: Final = "voyage/rerank-2.5" +MONGODB_KEY: Final = "al-synthetic-mongodb-issued-key" +VOYAGE_KEY: Final = "pa-synthetic-voyage-issued-key" +ENV_TOKEN: Final = "al-synthetic-env-token" +ENV_PRIMARY: Final = "pa-synthetic-env-primary" +ENV_SECONDARY: Final = "al-synthetic-env-secondary" +MONGODB_HOST: Final = "ai.mongodb.com:443" +VOYAGE_HOST: Final = "api.voyageai.com:443" +NO_CACHE: Final[dict[str, JsonValue]] = {"cache": {"no-cache": True}} +REFUSED_STATUS: Final = 500 +REFUSED_MARKER: Final = "403" +OUTAGE_MARKER: Final = "APIConnectionError" +MONGODB_EMBEDDINGS: Final = "voyage-audit-mongodb-embeddings" +MONGODB_CONTEXTUAL: Final = "voyage-audit-mongodb-contextual" +MONGODB_MULTIMODAL: Final = "voyage-audit-mongodb-multimodal" +MONGODB_RERANK: Final = "voyage-audit-mongodb-rerank" +VOYAGE_EMBEDDINGS: Final = "voyage-audit-voyage-embeddings" +VOYAGE_RERANK: Final = "voyage-audit-voyage-rerank" +BLANK_EMBEDDINGS: Final = "voyage-audit-blank-key-embeddings" +BLANK_RERANK: Final = "voyage-audit-blank-key-rerank" +NULL_RERANK: Final = "voyage-audit-null-key-rerank" +SCRIPTED_CHAT: Final = "voyage-audit-scripted-chat" +TOKEN_EMBEDDINGS: Final = "voyage-audit-token-embeddings" +TOKEN_RERANK: Final = "voyage-audit-token-rerank" +TOKEN_YAML_RERANK: Final = "voyage-audit-token-yaml-rerank" +TOKEN_BLANK_EMBEDDINGS: Final = "voyage-audit-token-blank-embeddings" +TOKEN_BLANK_RERANK: Final = "voyage-audit-token-blank-rerank" +TOKEN_NULL_RERANK: Final = "voyage-audit-token-null-rerank" +PRECEDENCE_EMBEDDINGS: Final = "voyage-audit-precedence-embeddings" +PRECEDENCE_RERANK: Final = "voyage-audit-precedence-rerank" +PRECEDENCE_EXPLICIT_EMBEDDINGS: Final = "voyage-audit-precedence-explicit-embeddings" +PRECEDENCE_EXPLICIT_RERANK: Final = "voyage-audit-precedence-explicit-rerank" +VECTOR: Final = [0.1, 0.2, 0.3] +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_PROXY_VARIABLES: Final = frozenset({"HTTPS_PROXY", "HTTP_PROXY", "ALL_PROXY", "NO_PROXY"}) +_JSON_OBJECTS: Final = TypeAdapter(list[dict[str, JsonValue]]) +_WORKER_PIDS: Final = TypeAdapter(list[int]) +_load_yaml: Final[Callable[[str], object]] = yaml.safe_load +_OWNED_PROXY_CELL_SECONDS: Final = 2 * max(30.0, float(os.environ.get("INTEGRATION_PROXY_READY_SECONDS") or 70)) + 120 +pytestmark: Final = pytest.mark.timeout(_OWNED_PROXY_CELL_SECONDS) + + +def _deployment(name: str, model: str, mode: str, **litellm_params: JsonValue) -> dict[str, JsonValue]: + return { + "model_name": name, + "litellm_params": {"model": model, **litellm_params}, + "model_info": {"id": name, "mode": mode}, + } + + +def _isolated_deployments(upstream_url: str) -> list[dict[str, JsonValue]]: + return [ + _deployment(MONGODB_EMBEDDINGS, EMBEDDING_MODEL, "embedding", api_key=MONGODB_KEY), + _deployment(MONGODB_CONTEXTUAL, CONTEXTUAL_MODEL, "embedding", api_key=MONGODB_KEY), + _deployment(MONGODB_MULTIMODAL, MULTIMODAL_MODEL, "embedding", api_key=MONGODB_KEY), + _deployment(MONGODB_RERANK, RERANK_MODEL, "rerank", api_key=MONGODB_KEY), + _deployment(VOYAGE_EMBEDDINGS, EMBEDDING_MODEL, "embedding", api_key=VOYAGE_KEY), + _deployment(VOYAGE_RERANK, RERANK_MODEL, "rerank", api_key=VOYAGE_KEY), + _deployment(BLANK_EMBEDDINGS, EMBEDDING_MODEL, "embedding", api_key=""), + _deployment(BLANK_RERANK, RERANK_MODEL, "rerank", api_key=""), + _deployment(NULL_RERANK, RERANK_MODEL, "rerank", api_key=None), + _deployment( + SCRIPTED_CHAT, + "openai/gpt-4o-mini", + "chat", + api_key="integration-provider-key", + api_base=f"{upstream_url}/v1", + ), + ] + + +def _token_only_deployments() -> list[dict[str, JsonValue]]: + return [ + _deployment(TOKEN_EMBEDDINGS, EMBEDDING_MODEL, "embedding"), + _deployment(TOKEN_RERANK, RERANK_MODEL, "rerank"), + _deployment(TOKEN_YAML_RERANK, RERANK_MODEL, "rerank", api_key="os.environ/VOYAGE_AI_TOKEN"), + _deployment(TOKEN_BLANK_EMBEDDINGS, EMBEDDING_MODEL, "embedding", api_key=""), + _deployment(TOKEN_BLANK_RERANK, RERANK_MODEL, "rerank", api_key=""), + _deployment(TOKEN_NULL_RERANK, RERANK_MODEL, "rerank", api_key=None), + ] + + +def _precedence_deployments() -> list[dict[str, JsonValue]]: + return [ + _deployment(PRECEDENCE_EMBEDDINGS, EMBEDDING_MODEL, "embedding"), + _deployment(PRECEDENCE_RERANK, RERANK_MODEL, "rerank"), + _deployment(PRECEDENCE_EXPLICIT_EMBEDDINGS, EMBEDDING_MODEL, "embedding", api_key=MONGODB_KEY), + _deployment(PRECEDENCE_EXPLICIT_RERANK, RERANK_MODEL, "rerank", api_key=MONGODB_KEY), + ] + + +def _write_config(directory: Path, name: str, model_list: list[dict[str, JsonValue]]) -> Path: + base: Final = JSON_OBJECT.validate_python(_load_yaml(Path("tests/integration/proxy_config.yaml").read_text())) + configuration: Final = { + **base, + "model_list": model_list, + "router_settings": {"disable_cooldowns": True, "num_retries": 0}, + } + path: Final = directory / f"{name}.yaml" + path.write_text(yaml.safe_dump(configuration)) + return path + + +def _scrubbed() -> tuple[str, ...]: + return tuple(name for name in os.environ if name.upper().startswith("VOYAGE_") or name.upper() in _PROXY_VARIABLES) + + +def _forward_environment(forward_url: str, **extra: str) -> dict[str, str]: + return {"HTTPS_PROXY": forward_url, "NO_PROXY": "127.0.0.1,localhost", **extra} + + +@pytest.fixture(scope="module") +def forward() -> Iterator[ForwardProxy]: + with refusing_forward_proxy() as proxy: + yield proxy + + +@pytest.fixture(scope="module") +def module_gateway() -> Iterator[Gateway]: + with gateway_from_environment() as value: + yield value + + +@pytest.fixture(scope="module") +def isolated( + module_gateway: Gateway, forward: ForwardProxy, tmp_path_factory: pytest.TempPathFactory +) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("voyage-isolated") + config: Final = _write_config(directory, "isolated", _isolated_deployments(module_gateway.upstream_url)) + with owned_proxy( + module_gateway, + directory, + _forward_environment(forward.url), + config=config, + remove_environment=_scrubbed(), + workers=2, + ) as candidate: + yield candidate + + +@pytest.fixture(scope="module") +def token_only( + module_gateway: Gateway, forward: ForwardProxy, tmp_path_factory: pytest.TempPathFactory +) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("voyage-token-only") + config: Final = _write_config(directory, "token-only", _token_only_deployments()) + with owned_proxy( + module_gateway, + directory, + _forward_environment(forward.url, VOYAGE_AI_TOKEN=ENV_TOKEN), + config=config, + remove_environment=_scrubbed(), + ) as candidate: + yield candidate + + +@pytest.fixture(scope="module") +def precedence( + module_gateway: Gateway, forward: ForwardProxy, tmp_path_factory: pytest.TempPathFactory +) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("voyage-precedence") + config: Final = _write_config(directory, "precedence", _precedence_deployments()) + with owned_proxy( + module_gateway, + directory, + _forward_environment( + forward.url, + VOYAGE_API_KEY=ENV_PRIMARY, + VOYAGE_AI_API_KEY=ENV_SECONDARY, + VOYAGE_AI_TOKEN=ENV_TOKEN, + ), + config=config, + remove_environment=_scrubbed(), + ) as candidate: + yield candidate + + +def _embedding_body(model: str, **extra: JsonValue) -> dict[str, JsonValue]: + return {"model": model, "input": f"host selection {uuid.uuid4().hex}", **NO_CACHE, **extra} + + +def _rerank_body(model: str, **extra: JsonValue) -> dict[str, JsonValue]: + return { + "model": model, + "query": f"host selection {uuid.uuid4().hex}", + "documents": ["first document", "second document"], + **extra, + } + + +def _embed(candidate: Gateway, model: str, **extra: JsonValue) -> httpx.Response: + return candidate.request("POST", "/v1/embeddings", _embedding_body(model, **extra)) + + +def _rerank(candidate: Gateway, model: str, path: str = "/v1/rerank", **extra: JsonValue) -> httpx.Response: + return candidate.request("POST", path, _rerank_body(model, **extra)) + + +def _error_message(call_id: str) -> str: + rows: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)), + lambda values: len(values) == 1, + seconds=70, + ) + metadata: Final = rows[0]["metadata"] + parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata) + return str(object_value(parsed["error_information"])["error_message"]) + + +def _assert_refused(response: httpx.Response) -> None: + assert response.status_code == REFUSED_STATUS, response.text + assert REFUSED_MARKER in response.text, response.text + assert REFUSED_MARKER in _error_message(response.headers["x-litellm-call-id"]), response.text + + +def _assert_dialed(forward: ForwardProxy, response: httpx.Response, host: str) -> None: + _assert_refused(response) + assert forward.targets() == (host,), response.text + + +def test_scripted_chat_keeps_serving_behind_the_forward_proxy(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + reply: Final = isolated.chat(SCRIPTED_CHAT, text=f"forward proxy control {uuid.uuid4().hex}") + rows: Final = eventually( + lambda: read_rows('SELECT model FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (str(reply["id"]),)), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["model"] == "openai/gpt-4o-mini", rows + assert forward.targets() == () + + +def test_mongodb_key_embeddings_dial_ai_mongodb_openai_sdk_sync(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + with OpenAI( + api_key=isolated.key, + base_url=str(isolated.client.base_url).rstrip("/") + "/v1", + max_retries=0, + http_client=httpx.Client(timeout=15, trust_env=False), + ) as client: + with pytest.raises(APIStatusError) as caught: + client.embeddings.create( + model=MONGODB_EMBEDDINGS, input=f"sdk sync {uuid.uuid4().hex}", extra_body=NO_CACHE + ) + assert caught.value.status_code == REFUSED_STATUS, caught.value.message + assert REFUSED_MARKER in caught.value.message, caught.value.message + assert forward.targets() == (MONGODB_HOST,) + + +async def test_mongodb_key_embeddings_dial_ai_mongodb_openai_sdk_async( + isolated: Gateway, forward: ForwardProxy +) -> None: + forward.drain() + async with AsyncOpenAI( + api_key=isolated.key, + base_url=str(isolated.client.base_url).rstrip("/") + "/v1", + max_retries=0, + http_client=httpx.AsyncClient(timeout=15, trust_env=False), + ) as client: + with pytest.raises(APIStatusError) as caught: + await client.embeddings.create( + model=MONGODB_EMBEDDINGS, input=f"sdk async {uuid.uuid4().hex}", extra_body=NO_CACHE + ) + assert caught.value.status_code == REFUSED_STATUS, caught.value.message + assert REFUSED_MARKER in caught.value.message, caught.value.message + assert forward.targets() == (MONGODB_HOST,) + + +@pytest.mark.parametrize("deployment", (MONGODB_CONTEXTUAL, MONGODB_MULTIMODAL)) +def test_mongodb_key_other_embedding_shapes_dial_ai_mongodb( + isolated: Gateway, forward: ForwardProxy, deployment: str +) -> None: + forward.drain() + _assert_dialed(forward, _embed(isolated, deployment), MONGODB_HOST) + + +@pytest.mark.parametrize("path", ("/v1/rerank", "/rerank", "/v2/rerank")) +def test_mongodb_key_rerank_dials_ai_mongodb(isolated: Gateway, forward: ForwardProxy, path: str) -> None: + forward.drain() + _assert_dialed(forward, _rerank(isolated, MONGODB_RERANK, path), MONGODB_HOST) + + +def test_voyage_key_embeddings_still_dial_api_voyageai(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + _assert_dialed(forward, _embed(isolated, VOYAGE_EMBEDDINGS), VOYAGE_HOST) + + +def test_voyage_key_rerank_still_dials_api_voyageai(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + _assert_dialed(forward, _rerank(isolated, VOYAGE_RERANK), VOYAGE_HOST) + + +def test_request_body_mongodb_key_overrides_voyage_deployment_host(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + _assert_dialed(forward, _embed(isolated, VOYAGE_EMBEDDINGS, api_key=MONGODB_KEY), MONGODB_HOST) + + +@pytest.mark.parametrize("deployment", (MONGODB_EMBEDDINGS, MONGODB_RERANK)) +def test_health_check_dials_ai_mongodb_for_mongodb_key( + isolated: Gateway, forward: ForwardProxy, deployment: str +) -> None: + forward.drain() + response: Final = isolated.request("GET", "/health", params={"model": deployment}) + assert response.status_code == 503, response.text + report: Final = JSON_OBJECT.validate_json(response.content) + assert report["unhealthy_count"] == 1 and report["healthy_count"] == 0, report + unhealthy: Final = _JSON_OBJECTS.validate_python(report["unhealthy_endpoints"]) + assert len(unhealthy) == 1 and REFUSED_MARKER in str(unhealthy[0]["error"]), report + assert forward.targets() == (MONGODB_HOST,), report + + +@pytest.mark.parametrize( + ("mode", "model"), (("embedding", EMBEDDING_MODEL), ("rerank", RERANK_MODEL)), ids=("embedding", "rerank") +) +def test_test_connection_dials_ai_mongodb_for_mongodb_key( + isolated: Gateway, forward: ForwardProxy, mode: str, model: str +) -> None: + forward.drain() + report: Final = isolated.post( + "/health/test_connection", {"litellm_params": {"model": model, "api_key": MONGODB_KEY}, "mode": mode} + ) + assert report["status"] == "error", report + assert REFUSED_MARKER in json.dumps(report), report + assert forward.targets() == (MONGODB_HOST,), report + + +@pytest.mark.parametrize("api_key", (5, ["al-list-member"]), ids=("int", "list")) +def test_non_string_request_body_key_is_rejected_by_validation_before_any_dial( + isolated: Gateway, forward: ForwardProxy, api_key: JsonValue +) -> None: + forward.drain() + response: Final = _embed(isolated, MONGODB_EMBEDDINGS, api_key=api_key) + assert response.status_code == 500, response.text + assert "LiteLLM_Params" in response.text and "api_key" in response.text, response.text + assert forward.targets() == () + + +def test_empty_request_body_key_falls_back_to_api_voyageai(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + _assert_dialed(forward, _embed(isolated, MONGODB_EMBEDDINGS, api_key=""), VOYAGE_HOST) + + +def test_five_kilobyte_mongodb_request_body_key_dials_ai_mongodb(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + _assert_dialed(forward, _embed(isolated, VOYAGE_EMBEDDINGS, api_key=MONGODB_KEY + "x" * 5000), MONGODB_HOST) + + +def test_duplicate_request_body_key_counts_once(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + raw: Final = ( + '{"model": "%s", "input": "duplicate key %s", "cache": {"no-cache": true}, "api_key": "%s", "api_key": "%s"}' + % (VOYAGE_EMBEDDINGS, uuid.uuid4().hex, MONGODB_KEY, MONGODB_KEY) + ) + response: Final = isolated.client.post( + "/v1/embeddings", + content=raw.encode(), + headers={"Authorization": f"Bearer {isolated.key}", "content-type": "application/json"}, + ) + _assert_dialed(forward, response, MONGODB_HOST) + + +def test_blank_deployment_key_without_env_dials_api_voyageai_for_embeddings( + isolated: Gateway, forward: ForwardProxy +) -> None: + forward.drain() + _assert_dialed(forward, _embed(isolated, BLANK_EMBEDDINGS), VOYAGE_HOST) + + +@pytest.mark.parametrize("deployment", (BLANK_RERANK, NULL_RERANK), ids=("blank", "null")) +def test_rerank_without_any_key_fails_before_dialing(isolated: Gateway, forward: ForwardProxy, deployment: str) -> None: + forward.drain() + response: Final = _rerank(isolated, deployment) + assert response.status_code == 500, response.text + assert "Voyage AI API key is required" in response.text, response.text + assert forward.targets() == () + + +@dataclass(frozen=True, slots=True) +class _Probe: + host: str + response: httpx.Response + + +def _burst_plan(count: int) -> tuple[tuple[str, str, str], ...]: + shapes: Final = ( + ("/v1/embeddings", MONGODB_EMBEDDINGS, MONGODB_HOST), + ("/v1/rerank", MONGODB_RERANK, MONGODB_HOST), + ("/v1/embeddings", VOYAGE_EMBEDDINGS, VOYAGE_HOST), + ("/v1/rerank", VOYAGE_RERANK, VOYAGE_HOST), + ) + return tuple(shapes[index % len(shapes)] for index in range(count)) + + +async def _burst(base_url: str, key: str, count: int, *, tolerate_transport_errors: bool = False) -> tuple[_Probe, ...]: + async def one(client: httpx.AsyncClient, path: str, deployment: str, host: str) -> _Probe: + body: Final = _embedding_body(deployment) if path == "/v1/embeddings" else _rerank_body(deployment) + response: Final = await client.post(path, json=body, headers={"Authorization": f"Bearer {key}"}) + return _Probe(host, response) + + async with httpx.AsyncClient(base_url=base_url, timeout=30, trust_env=False) as client: + results: Final = await asyncio.gather( + *(one(client, path, deployment, host) for path, deployment, host in _burst_plan(count)), + return_exceptions=tolerate_transport_errors, + ) + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, _Probe)) + + +def _host_counts(targets: tuple[str, ...]) -> dict[str, int]: + return {host: targets.count(host) for host in sorted(set(targets))} + + +async def test_concurrent_mixed_keys_each_dial_their_own_host(isolated: Gateway, forward: ForwardProxy) -> None: + forward.drain() + served: Final = await _burst(str(isolated.client.base_url), isolated.key, 24) + assert len(served) == 24 + for probe in served: + _assert_refused(probe.response) + assert _host_counts(forward.targets()) == {MONGODB_HOST: 12, VOYAGE_HOST: 12} + + +def test_voyage_ai_token_alone_routes_embeddings_to_ai_mongodb(token_only: Gateway, forward: ForwardProxy) -> None: + forward.drain() + _assert_dialed(forward, _embed(token_only, TOKEN_EMBEDDINGS), MONGODB_HOST) + + +@pytest.mark.parametrize( + "deployment", (TOKEN_RERANK, TOKEN_YAML_RERANK, TOKEN_NULL_RERANK), ids=("missing", "yaml-env", "null") +) +def test_voyage_ai_token_alone_routes_rerank_to_ai_mongodb( + token_only: Gateway, forward: ForwardProxy, deployment: str +) -> None: + forward.drain() + _assert_dialed(forward, _rerank(token_only, deployment), MONGODB_HOST) + + +@pytest.mark.parametrize( + ("path", "deployment"), + (("/v1/embeddings", TOKEN_BLANK_EMBEDDINGS), ("/v1/rerank", TOKEN_BLANK_RERANK)), + ids=("embeddings", "rerank"), +) +def test_blank_deployment_key_falls_through_to_voyage_ai_token( + token_only: Gateway, forward: ForwardProxy, path: str, deployment: str +) -> None: + forward.drain() + response: Final = ( + _embed(token_only, deployment) if path == "/v1/embeddings" else _rerank(token_only, deployment, path) + ) + _assert_dialed(forward, response, MONGODB_HOST) + + +@pytest.mark.parametrize( + ("path", "deployment"), + (("/v1/embeddings", PRECEDENCE_EMBEDDINGS), ("/v1/rerank", PRECEDENCE_RERANK)), + ids=("embeddings", "rerank"), +) +def test_voyage_api_key_wins_over_mongodb_fallback_env( + precedence: Gateway, forward: ForwardProxy, path: str, deployment: str +) -> None: + forward.drain() + response: Final = ( + _embed(precedence, deployment) if path == "/v1/embeddings" else _rerank(precedence, deployment, path) + ) + _assert_dialed(forward, response, VOYAGE_HOST) + + +@pytest.mark.parametrize( + ("path", "deployment"), + (("/v1/embeddings", PRECEDENCE_EXPLICIT_EMBEDDINGS), ("/v1/rerank", PRECEDENCE_EXPLICIT_RERANK)), + ids=("embeddings", "rerank"), +) +def test_explicit_mongodb_deployment_key_wins_over_voyage_env( + precedence: Gateway, forward: ForwardProxy, path: str, deployment: str +) -> None: + forward.drain() + response: Final = ( + _embed(precedence, deployment) if path == "/v1/embeddings" else _rerank(precedence, deployment, path) + ) + _assert_dialed(forward, response, MONGODB_HOST) + + +def _embedding_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.headers["authorization"] == f"Bearer {MONGODB_KEY}", request.headers + body: Final = JSON_OBJECT.validate_json(request.body) + return Reply( + body=json.dumps( + { + "object": "list", + "data": [{"object": "embedding", "embedding": VECTOR, "index": 0}], + "model": str(body["model"]), + "usage": {"total_tokens": 7}, + } + ).encode() + ) + + +def _rerank_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.headers["authorization"] == f"Bearer {MONGODB_KEY}", request.headers + return Reply( + body=json.dumps( + { + "object": "list", + "data": [{"index": 1, "relevance_score": 0.9}, {"index": 0, "relevance_score": 0.1}], + "model": "rerank-2.5", + "usage": {"total_tokens": 11}, + } + ).encode() + ) + + +def _spend_api_base(call_id: str) -> str: + rows: Final = eventually( + lambda: read_rows('SELECT api_base FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)), + lambda values: len(values) == 1, + seconds=70, + ) + return str(rows[0]["api_base"]) + + +@pytest.mark.parametrize( + ("model", "target"), + ( + (EMBEDDING_MODEL, "/embeddings"), + (CONTEXTUAL_MODEL, "/contextualizedembeddings"), + (MULTIMODAL_MODEL, "/multimodalembeddings"), + ), + ids=("embeddings", "contextual", "multimodal"), +) +def test_explicit_api_base_keeps_mongodb_key_embeddings_on_that_base(gateway: Gateway, model: str, target: str) -> None: + with wire_server(_embedding_peer) as wire, gateway.scenario() as scenario: + deployment: Final = scenario.model(model=model, api_key=MONGODB_KEY, api_base=wire.url) + response: Final = _embed(gateway, deployment) + assert response.status_code == 200, response.text + data: Final = _JSON_OBJECTS.validate_python(JSON_OBJECT.validate_json(response.content)["data"]) + assert data[0]["embedding"] == VECTOR, response.text + received: Final = wire.drain() + assert tuple(request.target for request in received) == (target,), received + assert _spend_api_base(response.headers["x-litellm-call-id"]).startswith(wire.url), response.text + + +@pytest.mark.parametrize("suffix", ("", "/v1"), ids=("bare", "v1")) +def test_explicit_api_base_keeps_mongodb_key_rerank_on_that_base(gateway: Gateway, suffix: str) -> None: + with wire_server(_rerank_peer) as wire, gateway.scenario() as scenario: + deployment: Final = scenario.model(model=RERANK_MODEL, api_key=MONGODB_KEY, api_base=wire.url + suffix) + response: Final = _rerank(gateway, deployment) + assert response.status_code == 200, response.text + payload: Final = JSON_OBJECT.validate_json(response.content) + results: Final = _JSON_OBJECTS.validate_python(payload["results"]) + assert [result["index"] for result in results] == [1, 0], response.text + received: Final = wire.drain() + assert tuple(request.target for request in received) == ("/v1/rerank",), received + assert _spend_api_base(str(payload["id"])).startswith(wire.url), response.text + + +def test_public_provider_fields_name_voyage_as_mongodb(gateway: Gateway) -> None: + response: Final = gateway.request("GET", "/public/providers/fields") + assert response.status_code == 200, response.text + voyage: Final = [ + entry for entry in _JSON_OBJECTS.validate_json(response.content) if entry["litellm_provider"] == "voyage" + ] + assert [entry["provider_display_name"] for entry in voyage] == ["VoyageAI by MongoDB"], response.text + + +def test_public_endpoints_name_voyage_as_mongodb(gateway: Gateway) -> None: + response: Final = gateway.request("GET", "/public/endpoints") + assert response.status_code == 200, response.text + endpoints: Final = _JSON_OBJECTS.validate_python(JSON_OBJECT.validate_json(response.content)["endpoints"]) + names: Final = tuple(_voyage_display_names(endpoints)) + assert names and set(names) == {"VoyageAI by MongoDB"}, response.text + + +def _voyage_display_names(endpoints: list[dict[str, JsonValue]]) -> Iterator[str]: + for endpoint in endpoints: + for provider in _JSON_OBJECTS.validate_python(endpoint["providers"]): + if provider["slug"] == "voyage": + yield str(provider["display_name"]) + + +def _served_call_ids(served: tuple[_Probe, ...]) -> tuple[str, ...]: + return tuple(probe.response.headers["x-litellm-call-id"] for probe in served) + + +def _spend_rows(call_id: str) -> list[dict[str, JsonValue]]: + return read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)) + + +def _reserve_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + address: Final[Callable[[], tuple[str, int]]] = reserve.getsockname + return address()[1] + + +async def test_forward_proxy_outage_mid_traffic_recovers_with_the_right_hosts(gateway: Gateway, tmp_path: Path) -> None: + port: Final = _reserve_port() + config: Final = _write_config(tmp_path, "outage", _isolated_deployments(gateway.upstream_url)) + with owned_proxy( + gateway, + tmp_path, + _forward_environment(f"http://127.0.0.1:{port}"), + config=config, + remove_environment=_scrubbed(), + workers=2, + ) as candidate: + with refusing_forward_proxy(port=port) as before: + first: Final = await _burst(str(candidate.client.base_url), candidate.key, 12) + before_counts: Final = _host_counts(before.targets()) + during: Final = await _burst(str(candidate.client.base_url), candidate.key, 12) + liveliness: Final = candidate.request("GET", "/health/liveliness") + assert liveliness.status_code == 200, liveliness.text + with refusing_forward_proxy(port=port) as after: + third: Final = await _burst(str(candidate.client.base_url), candidate.key, 12) + after_counts: Final = _host_counts(after.targets()) + assert len(first) == len(during) == len(third) == 12 + for probe in (*first, *third): + _assert_refused(probe.response) + for probe in during: + assert probe.response.status_code == REFUSED_STATUS, probe.response.text + assert OUTAGE_MARKER in probe.response.text, probe.response.text + assert OUTAGE_MARKER in _error_message(probe.response.headers["x-litellm-call-id"]) + assert before_counts == {MONGODB_HOST: 6, VOYAGE_HOST: 6} + assert after_counts == {MONGODB_HOST: 6, VOYAGE_HOST: 6} + + +async def test_worker_sigkill_mid_traffic_leaves_the_sibling_routing_by_key(gateway: Gateway, tmp_path: Path) -> None: + config: Final = _write_config(tmp_path, "sigkill", _isolated_deployments(gateway.upstream_url)) + with ( + refusing_forward_proxy() as forward, + owned_proxy_process( + gateway, + tmp_path, + _forward_environment(forward.url), + config=config, + remove_environment=_scrubbed(), + workers=2, + ) as owned, + ): + candidate: Final = owned.gateway + workers: Final = eventually( + lambda: _WORKER_PIDS.validate_python(_STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=60, + ) + burst: Final = asyncio.create_task( + _burst(str(candidate.client.base_url), candidate.key, 20, tolerate_transport_errors=True) + ) + await asyncio.to_thread(eventually, lambda: forward.received.qsize(), lambda size: size >= 4, 60) + victim: Final = psutil.Process(workers[0]) + victim.suspend() + victim.send_signal(signal.SIGKILL) + served: Final = await burst + for probe in served: + assert probe.response.status_code == REFUSED_STATUS, probe.response.text + assert set(forward.targets()) <= {MONGODB_HOST, VOYAGE_HOST} + follow_up: Final = _rerank(candidate, MONGODB_RERANK) + _assert_dialed(forward, follow_up, MONGODB_HOST) + duplicates: Final = tuple(call_id for call_id in _served_call_ids(served) if len(_spend_rows(call_id)) > 1) + assert duplicates == (), duplicates diff --git a/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py b/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py index 777e52f4e2a..15689f43bd0 100644 --- a/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py +++ b/tests/unit/llms/voyage/rerank/test_voyage_rerank_transformation.py @@ -295,11 +295,10 @@ class TestVoyageRerankTransform: assert "top_n" in supported_params assert "return_documents" in supported_params - @patch("litellm.llms.voyage.rerank.transformation.get_secret_str") - def test_validate_environment_missing_api_key(self, mock_get_secret_str): + def test_validate_environment_missing_api_key(self, monkeypatch): """Test that validate_environment raises error when API key is missing.""" - # Mock get_secret_str to return None for both environment variables - mock_get_secret_str.return_value = None + for env_var in ("VOYAGE_API_KEY", "VOYAGE_AI_API_KEY", "VOYAGE_AI_TOKEN"): + monkeypatch.delenv(env_var, raising=False) with pytest.raises(ValueError, match="Voyage AI API key is required"): self.config.validate_environment( headers={}, diff --git a/tests/unit/llms/voyage/test_common_utils.py b/tests/unit/llms/voyage/test_common_utils.py new file mode 100644 index 00000000000..511391093f3 --- /dev/null +++ b/tests/unit/llms/voyage/test_common_utils.py @@ -0,0 +1,164 @@ +import pytest + +from litellm.llms.voyage.common_utils import ( + MONGODB_API_BASE, + VOYAGE_API_BASE, + get_default_base_url, + get_voyage_api_key, +) +from litellm.llms.voyage.embedding.transformation import VoyageEmbeddingConfig +from litellm.llms.voyage.embedding.transformation_contextual import ( + VoyageContextualEmbeddingConfig, +) +from litellm.llms.voyage.embedding.transformation_multimodal import ( + VoyageMultimodalEmbeddingConfig, +) +from litellm.llms.voyage.rerank.transformation import VoyageRerankConfig + +VOYAGE_KEY_ENV_VARS = ("VOYAGE_API_KEY", "VOYAGE_AI_API_KEY", "VOYAGE_AI_TOKEN") + + +@pytest.fixture(autouse=True) +def clear_voyage_env(monkeypatch): + for env_var in VOYAGE_KEY_ENV_VARS: + monkeypatch.delenv(env_var, raising=False) + + +def test_mongodb_key_routes_to_mongodb_host(): + """MongoDB-issued keys carry the `al-` prefix and are only valid on ai.mongodb.com""" + assert get_default_base_url("al-1234567890") == MONGODB_API_BASE + + +@pytest.mark.parametrize("api_key", ["pa-1234567890", "sk-1234567890", "al", "", None]) +def test_non_mongodb_key_routes_to_voyage_host(api_key): + assert get_default_base_url(api_key) == VOYAGE_API_BASE + + +@pytest.mark.parametrize("env_var", VOYAGE_KEY_ENV_VARS) +def test_mongodb_key_from_any_supported_env_var_routes_to_mongodb_host(monkeypatch, env_var): + monkeypatch.setenv(env_var, "al-from-env") + assert get_default_base_url() == MONGODB_API_BASE + + +def test_explicit_key_wins_over_env_for_routing(monkeypatch): + monkeypatch.setenv("VOYAGE_API_KEY", "al-from-env") + assert get_default_base_url("pa-explicit") == VOYAGE_API_BASE + + +@pytest.mark.parametrize( + "config, endpoint", + [ + (VoyageEmbeddingConfig(), "embeddings"), + (VoyageContextualEmbeddingConfig(), "contextualizedembeddings"), + (VoyageMultimodalEmbeddingConfig(), "multimodalembeddings"), + ], +) +@pytest.mark.parametrize("api_key, expected_host", [("al-key", MONGODB_API_BASE), ("pa-key", VOYAGE_API_BASE)]) +def test_embedding_configs_route_by_key_prefix(config, endpoint, api_key, expected_host): + url = config.get_complete_url(None, api_key, "voyage-3", {}, {}) + assert url == f"{expected_host}/{endpoint}" + + +@pytest.mark.parametrize( + "config, endpoint", + [ + (VoyageEmbeddingConfig(), "embeddings"), + (VoyageContextualEmbeddingConfig(), "contextualizedembeddings"), + (VoyageMultimodalEmbeddingConfig(), "multimodalembeddings"), + ], +) +def test_explicit_api_base_overrides_key_routing(config, endpoint): + url = config.get_complete_url("https://gateway.internal/v1", "al-key", "voyage-3", {}, {}) + assert url == f"https://gateway.internal/v1/{endpoint}" + + +@pytest.mark.parametrize("api_key, expected_host", [("al-key", MONGODB_API_BASE), ("pa-key", VOYAGE_API_BASE)]) +def test_rerank_routes_by_request_key_prefix(api_key, expected_host): + config = VoyageRerankConfig() + config.validate_environment({}, "rerank-2.5", api_key=api_key) + + assert config.get_complete_url(None, "rerank-2.5") == f"{expected_host}/rerank" + + +@pytest.mark.parametrize("api_key, expected_host", [("al-key", MONGODB_API_BASE), ("pa-key", VOYAGE_API_BASE)]) +def test_rerank_routes_by_env_key_prefix(monkeypatch, api_key, expected_host): + monkeypatch.setenv("VOYAGE_API_KEY", api_key) + config = VoyageRerankConfig() + config.validate_environment({}, "rerank-2.5") + + assert config.get_complete_url(None, "rerank-2.5") == f"{expected_host}/rerank" + + +def test_rerank_request_key_beats_env_key_for_routing(monkeypatch): + """A MongoDB key on the request must not be posted to the Voyage host the env key names""" + monkeypatch.setenv("VOYAGE_API_KEY", "pa-from-env") + config = VoyageRerankConfig() + + headers = config.validate_environment({}, "rerank-2.5", api_key="al-on-request") + + assert headers["Authorization"] == "Bearer al-on-request" + assert config.get_complete_url(None, "rerank-2.5") == f"{MONGODB_API_BASE}/rerank" + + +def test_rerank_config_is_built_per_request_so_keys_cannot_leak(monkeypatch): + """get_complete_url reads a key off the instance, so each request must get its own instance""" + import litellm + from litellm.types.utils import LlmProviders + from litellm.utils import ProviderConfigManager + + monkeypatch.delenv("VOYAGE_API_KEY", raising=False) + first = ProviderConfigManager.get_provider_rerank_config( + model="rerank-2.5", provider=LlmProviders.VOYAGE, api_base=None, present_version_params=[] + ) + second = ProviderConfigManager.get_provider_rerank_config( + model="rerank-2.5", provider=LlmProviders.VOYAGE, api_base=None, present_version_params=[] + ) + assert isinstance(first, litellm.VoyageRerankConfig) and first is not second + + first.validate_environment({}, "rerank-2.5", api_key="al-first-request") + second.validate_environment({}, "rerank-2.5", api_key="pa-second-request") + + assert first.get_complete_url(None, "rerank-2.5") == f"{MONGODB_API_BASE}/rerank" + assert second.get_complete_url(None, "rerank-2.5") == f"{VOYAGE_API_BASE}/rerank" + + +def test_rerank_falls_back_to_env_when_validate_environment_did_not_run(monkeypatch): + """A caller that skips validate_environment keeps the pre-existing env-only behaviour""" + monkeypatch.setenv("VOYAGE_API_KEY", "al-from-env") + + assert VoyageRerankConfig().get_complete_url(None, "rerank-2.5") == f"{MONGODB_API_BASE}/rerank" + + +@pytest.mark.parametrize( + "config", + [VoyageEmbeddingConfig(), VoyageContextualEmbeddingConfig(), VoyageMultimodalEmbeddingConfig()], +) +def test_auth_header_uses_the_key_the_url_was_routed_on(monkeypatch, config): + """The host is picked from a key, so the Authorization header has to carry that same key""" + monkeypatch.setenv("VOYAGE_AI_TOKEN", "al-from-env") + + headers = config.validate_environment({}, "voyage-3", [], {}, {}) + url = config.get_complete_url(None, None, "voyage-3", {}, {}) + + assert headers["Authorization"] == "Bearer al-from-env" + assert url.startswith(MONGODB_API_BASE) + + +def test_rerank_auth_header_uses_the_key_the_url_was_routed_on(monkeypatch): + monkeypatch.setenv("VOYAGE_AI_TOKEN", "al-from-env") + config = VoyageRerankConfig() + + headers = config.validate_environment({}, "rerank-2.5") + + assert headers["Authorization"] == "Bearer al-from-env" + assert config.get_complete_url(None, "rerank-2.5").startswith(MONGODB_API_BASE) + + +def test_get_voyage_api_key_prefers_env_vars_in_documented_order(monkeypatch): + monkeypatch.setenv("VOYAGE_AI_API_KEY", "second") + monkeypatch.setenv("VOYAGE_AI_TOKEN", "third") + assert get_voyage_api_key() == "second" + + monkeypatch.setenv("VOYAGE_API_KEY", "first") + assert get_voyage_api_key() == "first" + assert get_voyage_api_key("explicit") == "explicit" diff --git a/tests/unit/llms/voyage/test_voyage_multimodal_embedding.py b/tests/unit/llms/voyage/test_voyage_multimodal_embedding.py index f3e6885cbe6..d13610ada17 100644 --- a/tests/unit/llms/voyage/test_voyage_multimodal_embedding.py +++ b/tests/unit/llms/voyage/test_voyage_multimodal_embedding.py @@ -172,15 +172,12 @@ class TestVoyageMultimodalEmbeddings: assert headers == {"Authorization": "Bearer test-key"} def test_validate_environment_uses_secret_fallback(self, monkeypatch): - import litellm.llms.voyage.embedding.transformation_multimodal as module from litellm.llms.voyage.embedding.transformation_multimodal import ( VoyageMultimodalEmbeddingConfig, ) - def fake_get_secret(name): - return "secret-key" if name == "VOYAGE_AI_API_KEY" else None - - monkeypatch.setattr(module, "get_secret_str", fake_get_secret) + monkeypatch.delenv("VOYAGE_API_KEY", raising=False) + monkeypatch.setenv("VOYAGE_AI_API_KEY", "secret-key") config = VoyageMultimodalEmbeddingConfig() headers = config.validate_environment( {}, "voyage-multimodal-3.5", [], {}, {}, api_key=None @@ -188,12 +185,12 @@ class TestVoyageMultimodalEmbeddings: assert headers == {"Authorization": "Bearer secret-key"} def test_validate_environment_raises_without_api_key(self, monkeypatch): - import litellm.llms.voyage.embedding.transformation_multimodal as module from litellm.llms.voyage.embedding.transformation_multimodal import ( VoyageMultimodalEmbeddingConfig, ) - monkeypatch.setattr(module, "get_secret_str", lambda name: None) + for env_var in ("VOYAGE_API_KEY", "VOYAGE_AI_API_KEY", "VOYAGE_AI_TOKEN"): + monkeypatch.delenv(env_var, raising=False) config = VoyageMultimodalEmbeddingConfig() with pytest.raises(ValueError, match='Voyage API key is required for multimodal embeddings\\. Set') as exc_info: config.validate_environment( diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx index 5a57fc23dde..177d4be3247 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx @@ -182,7 +182,7 @@ export enum Providers { VERTEX_AI_BETA = "Vertex Ai Beta", VLLM = "Local vLLM", VolcEngine = "VolcEngine", - Voyage = "Voyage AI", + Voyage = "VoyageAI by MongoDB", WANDB = "Wandb", WATSONX = "Watsonx", WATSONX_TEXT = "Watsonx Text",