From f663794fcd7ad4df35c8b322c507f59452418d47 Mon Sep 17 00:00:00 2001 From: fzowl <160063452+fzowl@users.noreply.github.com> Date: Thu, 8 Oct 2026 22:12:47 +0200 Subject: [PATCH] feat(voyage): rebrand to VoyageAI by MongoDB and route MongoDB keys to ai.mongodb.com (#41812) * feat(voyage): rebrand to VoyageAI by MongoDB, refresh model list, support both contextual input shapes * fix(voyage): send auto-chunking params for flat contextual inputs A flat list[str] or bare str for the contextualized embeddings endpoint is only valid as documents with enable_auto_chunking=True and input_type=document, or as queries with input_type=query. Default those params for non-query flat inputs so the request matches the live API contract, while letting caller-set values win. Nested list[list[str]] still passes through unchanged. Tests now assert the auto-chunking params instead of only echoing inputs. * feat(voyage): route MongoDB-issued keys to ai.mongodb.com Mirror voyageai.util.get_default_base_url from the official SDK: a key with the `al-` prefix is issued by MongoDB and is only valid on ai.mongodb.com, every other key on api.voyageai.com. The choice now lives in one shared helper that the embedding, contextual, multimodal, and rerank configs all use for both the default base URL and the Authorization header, so the host and the key always agree. Also drops the voyage-4-nano and voyage-multilingual-2 model map entries: voyage-4-nano is not served on the Voyage API and voyage-multilingual-2 is an older model, so neither belongs in this change. * fix(voyage): restore MongoDB routing on contextual endpoint and re-add voyage-4-nano after upstream merge The upstream merge replaced the contextual config's get_complete_url with a Voyage-only host and dropped voyage-4-nano from the model map. Reapply the al- key routing via get_default_base_url and add voyage-4-nano (open-weight, per docs.voyageai.com) so the branch keeps its task changes on top of upstream. * fix(voyage): drop voyage-4-nano and the contextual tests upstream already covers voyage-4-nano is not served on the Voyage API, so it does not belong in the model map. The contextual input tests duplicate tests/test_litellm/llms/voyage/test_voyage_contextual_embedding.py, which landed upstream with the auto-chunking fix this branch was carrying. * test(voyage): drop sys.path.insert from the voyage common utils test * fix(voyage): route rerank on the key it authenticates with get_complete_url picked the rerank host from the environment while the auth header carried the request key, so an explicit MongoDB-issued key was posted to api.voyageai.com. validate_environment now hands the resolved key to the config instance and get_complete_url reads it back, so the host and the credential always come from one key. The config is built per request, so nothing carries over between them. * test: allow ultrafast pricing keys in model prices schema * test: drop ultrafast schema keys now added upstream * test(integration): prove voyage key-based host routing on the wire --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- README.md | 2 +- litellm/llms/voyage/common_utils.py | 35 + .../llms/voyage/embedding/transformation.py | 12 +- .../embedding/transformation_contextual.py | 12 +- .../embedding/transformation_multimodal.py | 15 +- litellm/llms/voyage/rerank/transformation.py | 22 +- .../provider_endpoints_support_backup.json | 2 +- .../provider_create_fields.json | 2 +- provider_endpoints_support.json | 2 +- tests/integration/_support/forward_proxy.py | 71 ++ .../test_voyage_mongodb_host_wire.py | 688 ++++++++++++++++++ .../test_voyage_rerank_transformation.py | 7 +- tests/unit/llms/voyage/test_common_utils.py | 164 +++++ .../test_voyage_multimodal_embedding.py | 11 +- .../src/components/provider_info_helpers.tsx | 2 +- 15 files changed, 997 insertions(+), 50 deletions(-) create mode 100644 litellm/llms/voyage/common_utils.py create mode 100644 tests/integration/_support/forward_proxy.py create mode 100644 tests/integration/providers/test_voyage_mongodb_host_wire.py create mode 100644 tests/unit/llms/voyage/test_common_utils.py 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",