chore(python): refactor legacy code in WatsonxEmbeddings component (#3660)

## What

- Replace legacy model names in `WatsonxEmbeddings` with the current
supported set:
  - `ibm/granite-embedding-278m-multilingual` (new default, 768-dim)
  - `ibm/slate-125m-english-rtrvr-v2` (768-dim)
  - `ibm/slate-30m-english-rtrvr-v2` (384-dim)
  - `intfloat/multilingual-e5-large` (1024-dim)
  - `sentence-transformers/all-minilm-l6-v2` (384-dim)
- Add `space_id` field — mutually exclusive with `project_id`, mirrors
the
  existing pattern in `WatsonxReranker`
- `project_id` / `space_id` resolution now falls back to
`WATSONX_PROJECT_ID` /
  `WATSONX_SPACE_ID` env vars; exactly one must be supplied

## Why

The previously hardcoded models (`ibm/slate-125m-english-rtrvr`,
`sentence-transformers/all-minilm-l12-v2`) are legacy and no longer
listed as
supported by the watsonx.ai platform. `space_id` scoping was already
supported
by `WatsonxReranker` but was missing from the embeddings counterpart.

---------

Co-authored-by: Will Jones <willjones127@gmail.com>
This commit is contained in:
Mateusz Szewczyk
2026-07-21 18:28:06 +02:00
committed by GitHub
parent 2f27aa377b
commit 8d2fea9151
3 changed files with 623 additions and 40 deletions
+103 -34
View File
@@ -14,29 +14,76 @@ import numpy as np
DEFAULT_WATSONX_URL = "https://us-south.ml.cloud.ibm.com"
MODELS_DIMS = {
# Models currently available on the watsonx.ai SaaS platform.
# These are the IDs advertised to new users via model_names() and shown in
# validation error messages. Regional availability and withdrawal dates are
# documented at:
# https://www.ibm.com/docs/en/watsonx/saas?topic=models-supported-encoder
CURRENT_MODELS: dict[str, int] = {
"ibm/granite-embedding-278m-multilingual": 768,
"ibm/slate-125m-english-rtrvr-v2": 768,
"ibm/slate-30m-english-rtrvr-v2": 384,
"intfloat/multilingual-e5-large": 1024,
}
# Full dimension map including legacy model IDs from earlier releases.
# Kept so that existing tables whose stored metadata uses these names can still
# resolve dimensions on load without raising an error. These IDs are NOT
# advertised to new users.
MODELS_DIMS: dict[str, int] = {
**CURRENT_MODELS,
# Deprecated — withdrawal announced but still functional until the dates above.
"sentence-transformers/all-minilm-l6-v2": 384,
# Pre-v2 legacy names retained for metadata compatibility only.
"ibm/slate-125m-english-rtrvr": 768,
"ibm/slate-30m-english-rtrvr": 384,
"sentence-transformers/all-minilm-l12-v2": 384,
"intfloat/multilingual-e5-large": 1024,
}
@register("watsonx")
class WatsonxEmbeddings(TextEmbeddingFunction):
"""
An embedding function that uses the IBM watsonx.ai Embeddings API.
API Docs:
---------
https://cloud.ibm.com/apidocs/watsonx-ai#text-embeddings
https://cloud.ibm.com/apidocs/watsonx-ai#text-embeddings
Supported embedding models:
---------------------------
https://dataplatform.cloud.ibm.com/docs/content/wsj/analyze-data/fm-models-embed.html?context=wx
https://dataplatform.cloud.ibm.com/docs/content/wsj/analyze-data/fm-models-embed.html?context=wx
Parameters
----------
name : str, default "ibm/slate-125m-english-rtrvr"
The ID of the embedding model to use. For new tables,
``"ibm/granite-embedding-278m-multilingual"`` is recommended.
api_key : str, optional
IBM Cloud API key. Falls back to the ``WATSONX_API_KEY`` environment
variable when not provided.
project_id : str, optional
watsonx.ai project ID. Explicit value takes precedence over the
``WATSONX_PROJECT_ID`` environment variable. Mutually exclusive with
``space_id`` — exactly one must be supplied.
space_id : str, optional
watsonx.ai deployment space ID. Explicit value takes precedence over
the ``WATSONX_SPACE_ID`` environment variable. Mutually exclusive with
``project_id`` — exactly one must be supplied.
url : str, optional
watsonx.ai service URL. Defaults to
``"https://us-south.ml.cloud.ibm.com"``.
params : dict, optional
Extra parameters forwarded verbatim to ``Embeddings`` (e.g.
``{"truncate_input_tokens": 512}``).
"""
# Intentionally kept at the original pre-PR default so that existing tables
# whose stored metadata contains model:{} reload with the same model they
# were created with. New users should pass name= explicitly, e.g.
# name="ibm/granite-embedding-278m-multilingual".
name: str = "ibm/slate-125m-english-rtrvr"
api_key: Optional[str] = None
project_id: Optional[str] = None
space_id: Optional[str] = None
url: Optional[str] = None
params: Optional[Dict] = None
@@ -46,12 +93,13 @@ class WatsonxEmbeddings(TextEmbeddingFunction):
@staticmethod
def model_names():
return [
"ibm/slate-125m-english-rtrvr",
"ibm/slate-30m-english-rtrvr",
"sentence-transformers/all-minilm-l12-v2",
"intfloat/multilingual-e5-large",
]
"""Return the IDs of models currently available for new tables.
Legacy / deprecated IDs are intentionally excluded. They remain
resolvable for dimension lookups on existing tables via ``MODELS_DIMS``,
but should not be used when creating new tables.
"""
return list(CURRENT_MODELS.keys())
def ndims(self):
return self._ndims
@@ -59,7 +107,10 @@ class WatsonxEmbeddings(TextEmbeddingFunction):
@cached_property
def _ndims(self):
if self.name not in MODELS_DIMS:
raise ValueError(f"Unknown model name {self.name}")
raise ValueError(
f"Unknown model '{self.name}'. "
f"Available models: {list(CURRENT_MODELS.keys())}"
)
return MODELS_DIMS[self.name]
def generate_embeddings(
@@ -81,27 +132,45 @@ class WatsonxEmbeddings(TextEmbeddingFunction):
"ibm_watsonx_ai.foundation_models"
)
kwargs = {"model_id": self.name}
# --- credentials ---
# Explicit field takes priority; env var is the fallback.
api_key = self.api_key or os.environ.get("WATSONX_API_KEY")
if not api_key:
raise ValueError(
"WATSONX_API_KEY not set. Either set it in your environment or "
"pass it as `api_key` argument to WatsonxEmbeddings."
)
credentials = ibm_watsonx_ai.Credentials(
api_key=api_key,
url=self.url or DEFAULT_WATSONX_URL,
)
# --- project_id / space_id (exactly one required) ---
# Explicit field always wins; env var is consulted only when the
# corresponding field was not set, so passing project_id= never
# conflicts with a stray WATSONX_SPACE_ID env var and vice-versa.
space_id, project_id = self.space_id, self.project_id
if project_id is None and space_id is None:
# Neither was passed explicitly — fall back to env vars.
project_id = os.environ.get("WATSONX_PROJECT_ID")
space_id = os.environ.get("WATSONX_SPACE_ID")
if project_id and space_id:
raise ValueError("Provide either `project_id` or `space_id`, not both.")
if not project_id and not space_id:
raise ValueError(
"Either WATSONX_PROJECT_ID or WATSONX_SPACE_ID must be set. "
"Pass one as an argument to WatsonxEmbeddings or set the "
"corresponding environment variable."
)
client_kwargs: Dict = dict(model_id=self.name, credentials=credentials)
if self.params:
kwargs["params"] = self.params
if self.project_id:
kwargs["project_id"] = self.project_id
elif "WATSONX_PROJECT_ID" in os.environ:
kwargs["project_id"] = os.environ["WATSONX_PROJECT_ID"]
client_kwargs["params"] = self.params
if project_id:
client_kwargs["project_id"] = project_id
else:
raise ValueError("WATSONX_PROJECT_ID must be set or passed")
client_kwargs["space_id"] = space_id
creds_kwargs = {}
if self.api_key:
creds_kwargs["api_key"] = self.api_key
elif "WATSONX_API_KEY" in os.environ:
creds_kwargs["api_key"] = os.environ["WATSONX_API_KEY"]
else:
raise ValueError("WATSONX_API_KEY must be set or passed")
if self.url:
creds_kwargs["url"] = self.url
else:
creds_kwargs["url"] = DEFAULT_WATSONX_URL
kwargs["credentials"] = ibm_watsonx_ai.Credentials(**creds_kwargs)
return ibm_watsonx_ai_foundation_models.Embeddings(**kwargs)
return ibm_watsonx_ai_foundation_models.Embeddings(**client_kwargs)
+14 -6
View File
@@ -40,12 +40,12 @@ class WatsonxReranker(Reranker):
IBM Cloud API key. Falls back to the ``WATSONX_API_KEY`` environment
variable when not provided.
project_id : str, optional
watsonx.ai project ID. Falls back to the ``WATSONX_PROJECT_ID``
environment variable when not provided. Mutually exclusive with
watsonx.ai project ID. Explicit value takes precedence over the
``WATSONX_PROJECT_ID`` environment variable. Mutually exclusive with
``space_id`` — exactly one must be supplied.
space_id : str, optional
watsonx.ai deployment space ID. Falls back to the ``WATSONX_SPACE_ID``
environment variable when not provided. Mutually exclusive with
watsonx.ai deployment space ID. Explicit value takes precedence over
the ``WATSONX_SPACE_ID`` environment variable. Mutually exclusive with
``project_id`` — exactly one must be supplied.
url : str, optional
watsonx.ai service URL. Defaults to
@@ -100,8 +100,16 @@ class WatsonxReranker(Reranker):
)
# --- project_id / space_id (exactly one required) ---
project_id = self.project_id or os.environ.get("WATSONX_PROJECT_ID")
space_id = self.space_id or os.environ.get("WATSONX_SPACE_ID")
# Explicit field always wins; env vars are consulted only when neither
# was passed explicitly, so a stray WATSONX_SPACE_ID never overrides an
# explicit project_id and vice-versa.
project_id = self.project_id
space_id = self.space_id
if project_id is None and space_id is None:
# Neither was passed explicitly — fall back to env vars.
project_id = os.environ.get("WATSONX_PROJECT_ID")
space_id = os.environ.get("WATSONX_SPACE_ID")
if project_id and space_id:
raise ValueError("Provide either `project_id` or `space_id`, not both.")
+506
View File
@@ -0,0 +1,506 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
"""Unit tests for WatsonxEmbeddings — no live API calls required."""
import pytest
from unittest.mock import MagicMock, patch
from lancedb.embeddings import get_registry
from lancedb.embeddings.watsonx import CURRENT_MODELS, MODELS_DIMS, WatsonxEmbeddings
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_func(monkeypatch, env=None, **create_kwargs):
"""
Return a WatsonxEmbeddings instance with ibm_watsonx_ai mocked out.
Parameters
----------
env : dict, optional
Environment variables to inject (merged on top of an empty env so that
no real WATSONX_* vars from the host bleed into the test).
create_kwargs :
Forwarded to ``WatsonxEmbeddings.create()``.
"""
base_env = {
k: "" for k in ("WATSONX_API_KEY", "WATSONX_PROJECT_ID", "WATSONX_SPACE_ID")
}
base_env.update(env or {})
# Only keep keys that have non-empty values so that absent vars are truly absent.
clean_env = {k: v for k, v in base_env.items() if v}
mock_embeddings_instance = MagicMock()
mock_foundation = MagicMock()
mock_foundation.Embeddings.return_value = mock_embeddings_instance
mock_ibm = MagicMock()
def _fake_import(name):
if name == "ibm_watsonx_ai":
return mock_ibm
if name == "ibm_watsonx_ai.foundation_models":
return mock_foundation
raise ImportError(name)
with patch.dict("os.environ", clean_env, clear=True):
with patch(
"lancedb.embeddings.watsonx.attempt_import_or_raise",
side_effect=_fake_import,
):
func = get_registry().get("watsonx").create(**create_kwargs)
# Force the cached_property to evaluate inside the patch context.
_ = func._watsonx_client
return func, mock_foundation
# ---------------------------------------------------------------------------
# Registry
# ---------------------------------------------------------------------------
class TestRegistry:
def test_watsonx_registered(self):
assert get_registry().get("watsonx") is not None
def test_model_names_returns_only_current_models(self):
names = WatsonxEmbeddings.model_names()
assert names == list(CURRENT_MODELS.keys())
# Current models must all be present.
for name in (
"ibm/granite-embedding-278m-multilingual",
"ibm/slate-125m-english-rtrvr-v2",
"ibm/slate-30m-english-rtrvr-v2",
"intfloat/multilingual-e5-large",
):
assert name in names, f"{name!r} missing from model_names()"
# Legacy / deprecated IDs must NOT appear in model_names().
for legacy in (
"ibm/slate-125m-english-rtrvr",
"ibm/slate-30m-english-rtrvr",
"sentence-transformers/all-minilm-l12-v2",
"sentence-transformers/all-minilm-l6-v2",
):
assert legacy not in names, (
f"Legacy model {legacy!r} should not appear in model_names()"
)
# ---------------------------------------------------------------------------
# Dimensions
# ---------------------------------------------------------------------------
class TestDimensions:
@pytest.mark.parametrize(
"model_name,expected_dims",
[
("ibm/granite-embedding-278m-multilingual", 768),
("ibm/slate-125m-english-rtrvr-v2", 768),
("ibm/slate-30m-english-rtrvr-v2", 384),
("intfloat/multilingual-e5-large", 1024),
("sentence-transformers/all-minilm-l6-v2", 384),
],
)
def test_current_model_dimensions(self, monkeypatch, model_name, expected_dims):
func, _ = _make_func(
monkeypatch,
env={"WATSONX_API_KEY": "key", "WATSONX_PROJECT_ID": "proj"},
name=model_name,
)
assert func.ndims() == expected_dims
def test_unknown_model_raises(self):
func = WatsonxEmbeddings(name="not/a-real-model")
with pytest.raises(ValueError, match="Unknown model"):
func.ndims()
# -- Backward-compat: legacy names must still resolve dims on table load --
@pytest.mark.parametrize(
"legacy_name,expected_dims",
[
("ibm/slate-125m-english-rtrvr", 768),
("ibm/slate-30m-english-rtrvr", 384),
("sentence-transformers/all-minilm-l12-v2", 384),
],
)
def test_legacy_model_dimensions_still_resolve(self, legacy_name, expected_dims):
"""Tables written with old model names must not raise on reload."""
assert MODELS_DIMS[legacy_name] == expected_dims
# ---------------------------------------------------------------------------
# Scope resolution (project_id / space_id)
# ---------------------------------------------------------------------------
class TestScopeResolution:
def test_explicit_project_id(self, monkeypatch):
func, mock_foundation = _make_func(
monkeypatch,
env={"WATSONX_API_KEY": "key"},
project_id="explicit-proj",
)
_, call_kwargs = mock_foundation.Embeddings.call_args
assert call_kwargs.get("project_id") == "explicit-proj"
assert "space_id" not in call_kwargs
def test_explicit_space_id(self, monkeypatch):
func, mock_foundation = _make_func(
monkeypatch,
env={"WATSONX_API_KEY": "key"},
space_id="explicit-space",
)
_, call_kwargs = mock_foundation.Embeddings.call_args
assert call_kwargs.get("space_id") == "explicit-space"
assert "project_id" not in call_kwargs
def test_env_project_id_fallback(self, monkeypatch):
func, mock_foundation = _make_func(
monkeypatch,
env={"WATSONX_API_KEY": "key", "WATSONX_PROJECT_ID": "env-proj"},
)
_, call_kwargs = mock_foundation.Embeddings.call_args
assert call_kwargs.get("project_id") == "env-proj"
def test_env_space_id_fallback(self, monkeypatch):
func, mock_foundation = _make_func(
monkeypatch,
env={"WATSONX_API_KEY": "key", "WATSONX_SPACE_ID": "env-space"},
)
_, call_kwargs = mock_foundation.Embeddings.call_args
assert call_kwargs.get("space_id") == "env-space"
def test_explicit_project_id_wins_over_env_space_id(self, monkeypatch):
"""Explicit project_id must not be overridden by WATSONX_SPACE_ID in env."""
func, mock_foundation = _make_func(
monkeypatch,
env={"WATSONX_API_KEY": "key", "WATSONX_SPACE_ID": "stray-env-space"},
project_id="explicit-proj",
)
_, call_kwargs = mock_foundation.Embeddings.call_args
assert call_kwargs.get("project_id") == "explicit-proj"
assert "space_id" not in call_kwargs
def test_explicit_space_id_wins_over_env_project_id(self, monkeypatch):
"""Explicit space_id must not be overridden by WATSONX_PROJECT_ID in env."""
func, mock_foundation = _make_func(
monkeypatch,
env={"WATSONX_API_KEY": "key", "WATSONX_PROJECT_ID": "stray-env-proj"},
space_id="explicit-space",
)
_, call_kwargs = mock_foundation.Embeddings.call_args
assert call_kwargs.get("space_id") == "explicit-space"
assert "project_id" not in call_kwargs
def test_both_env_vars_raises(self, monkeypatch):
"""When both WATSONX_PROJECT_ID and WATSONX_SPACE_ID env vars are set
(and neither is passed explicitly), it must raise 'not both'."""
func = WatsonxEmbeddings(name="ibm/granite-embedding-278m-multilingual")
mock_ibm = MagicMock()
mock_foundation = MagicMock()
def _fake_import(name):
if name == "ibm_watsonx_ai":
return mock_ibm
if name == "ibm_watsonx_ai.foundation_models":
return mock_foundation
raise ImportError(name)
with patch.dict(
"os.environ",
{
"WATSONX_API_KEY": "key",
"WATSONX_PROJECT_ID": "env-proj",
"WATSONX_SPACE_ID": "env-space",
},
clear=True,
):
with patch(
"lancedb.embeddings.watsonx.attempt_import_or_raise",
side_effect=_fake_import,
):
with pytest.raises(ValueError, match="not both"):
_ = func._watsonx_client
def test_both_explicit_raises(self):
func = WatsonxEmbeddings(
name="ibm/granite-embedding-278m-multilingual",
project_id="p",
space_id="s",
)
# The error surfaces when _watsonx_client is first accessed.
mock_ibm = MagicMock()
mock_foundation = MagicMock()
def _fake_import(name):
if name == "ibm_watsonx_ai":
return mock_ibm
if name == "ibm_watsonx_ai.foundation_models":
return mock_foundation
raise ImportError(name)
with patch.dict("os.environ", {"WATSONX_API_KEY": "key"}, clear=True):
with patch(
"lancedb.embeddings.watsonx.attempt_import_or_raise",
side_effect=_fake_import,
):
with pytest.raises(ValueError, match="not both"):
_ = func._watsonx_client
def test_neither_raises(self):
func = WatsonxEmbeddings(name="ibm/granite-embedding-278m-multilingual")
mock_ibm = MagicMock()
mock_foundation = MagicMock()
def _fake_import(name):
if name == "ibm_watsonx_ai":
return mock_ibm
if name == "ibm_watsonx_ai.foundation_models":
return mock_foundation
raise ImportError(name)
with patch.dict("os.environ", {"WATSONX_API_KEY": "key"}, clear=True):
with patch(
"lancedb.embeddings.watsonx.attempt_import_or_raise",
side_effect=_fake_import,
):
with pytest.raises(
ValueError, match="WATSONX_PROJECT_ID or WATSONX_SPACE_ID"
):
_ = func._watsonx_client
def test_missing_api_key_raises(self):
func = WatsonxEmbeddings(name="ibm/granite-embedding-278m-multilingual")
mock_ibm = MagicMock()
mock_foundation = MagicMock()
def _fake_import(name):
if name == "ibm_watsonx_ai":
return mock_ibm
if name == "ibm_watsonx_ai.foundation_models":
return mock_foundation
raise ImportError(name)
with patch.dict("os.environ", {"WATSONX_PROJECT_ID": "proj"}, clear=True):
with patch(
"lancedb.embeddings.watsonx.attempt_import_or_raise",
side_effect=_fake_import,
):
with pytest.raises(ValueError, match="WATSONX_API_KEY"):
_ = func._watsonx_client
# ---------------------------------------------------------------------------
# Metadata round-trip (backward compat)
# ---------------------------------------------------------------------------
class TestMetadataRoundTrip:
def test_reload_with_empty_model_metadata_preserves_model(self):
"""
Reproduce the exact deserialization path used by the registry:
create(**{}) → safe_model_dump() == {}
→ stored as model: {}
→ reloaded via create(**{})
The model must be identical before and after — no silent switch.
This guards against changing the class-level default between releases.
"""
from lancedb.embeddings.registry import EmbeddingFunctionRegistry
registry = EmbeddingFunctionRegistry.get_instance()
# Simulate original table creation with no explicit args.
original = registry.get("watsonx").create()
stored = original.safe_model_dump() # what gets written to arrow metadata
assert stored == {}, (
f"Expected empty stored args when create() called with no kwargs; "
f"got {stored!r}"
)
# Simulate reload: registry calls create(**stored) == create(**{})
reloaded = registry.get("watsonx").create(**stored)
assert reloaded.name == original.name, (
f"Model changed on reload: was {original.name!r}, "
f"became {reloaded.name!r}. "
"The class-level default must not change without a migration path."
)
def test_reload_from_legacy_metadata_explicit(self):
"""
Deserialize a representative legacy metadata payload and assert the exact
model name — this is the real cross-version regression guard.
Tables created before the v2 rename stored ``model: {"name": ...}`` with
the pre-v2 name. Reloading must produce exactly that model, not silently
switch to the current class default.
"""
from lancedb.embeddings.registry import EmbeddingFunctionRegistry
registry = EmbeddingFunctionRegistry.get_instance()
# This is what is stored in Arrow metadata for a table created with the
# pre-v2 default model name (no explicit name= was passed at the time).
legacy_stored = {"name": "ibm/slate-125m-english-rtrvr"}
reloaded = registry.get("watsonx").create(**legacy_stored)
assert reloaded.name == "ibm/slate-125m-english-rtrvr", (
f"Legacy metadata reload returned {reloaded.name!r} instead of "
"'ibm/slate-125m-english-rtrvr'. "
"MODELS_DIMS must keep legacy entries for backward compat."
)
def test_legacy_model_names_resolve_dims(self):
"""Legacy names in MODELS_DIMS so ndims() never raises on old tables."""
assert MODELS_DIMS["ibm/slate-125m-english-rtrvr"] == 768
assert MODELS_DIMS["ibm/slate-30m-english-rtrvr"] == 384
assert MODELS_DIMS["sentence-transformers/all-minilm-l12-v2"] == 384
# ---------------------------------------------------------------------------
# WatsonxReranker — scope resolution (project_id / space_id)
# ---------------------------------------------------------------------------
def _make_reranker(env=None, **init_kwargs):
"""
Return a WatsonxReranker with ibm_watsonx_ai mocked out.
Scope precedence is tested by inspecting what was passed to Rerank().
"""
from lancedb.rerankers.watsonx import WatsonxReranker
base_env = {
k: "" for k in ("WATSONX_API_KEY", "WATSONX_PROJECT_ID", "WATSONX_SPACE_ID")
}
base_env.update(env or {})
clean_env = {k: v for k, v in base_env.items() if v}
mock_rerank_instance = MagicMock()
mock_foundation = MagicMock()
mock_foundation.Rerank.return_value = mock_rerank_instance
mock_ibm = MagicMock()
def _fake_import(name):
if name == "ibm_watsonx_ai":
return mock_ibm
if name == "ibm_watsonx_ai.foundation_models":
return mock_foundation
raise ImportError(name)
reranker = WatsonxReranker(**init_kwargs)
with patch.dict("os.environ", clean_env, clear=True):
with patch(
"lancedb.rerankers.watsonx.attempt_import_or_raise",
side_effect=_fake_import,
):
_ = reranker._client
return reranker, mock_foundation
class TestRerankerScopeResolution:
def test_explicit_project_id(self):
_, mock_foundation = _make_reranker(
env={"WATSONX_API_KEY": "key"},
project_id="explicit-proj",
)
_, call_kwargs = mock_foundation.Rerank.call_args
assert call_kwargs.get("project_id") == "explicit-proj"
assert "space_id" not in call_kwargs
def test_explicit_space_id(self):
_, mock_foundation = _make_reranker(
env={"WATSONX_API_KEY": "key"},
space_id="explicit-space",
)
_, call_kwargs = mock_foundation.Rerank.call_args
assert call_kwargs.get("space_id") == "explicit-space"
assert "project_id" not in call_kwargs
def test_env_project_id_fallback(self):
_, mock_foundation = _make_reranker(
env={"WATSONX_API_KEY": "key", "WATSONX_PROJECT_ID": "env-proj"},
)
_, call_kwargs = mock_foundation.Rerank.call_args
assert call_kwargs.get("project_id") == "env-proj"
def test_env_space_id_fallback(self):
_, mock_foundation = _make_reranker(
env={"WATSONX_API_KEY": "key", "WATSONX_SPACE_ID": "env-space"},
)
_, call_kwargs = mock_foundation.Rerank.call_args
assert call_kwargs.get("space_id") == "env-space"
def test_explicit_project_id_wins_over_env_space_id(self):
"""Explicit project_id must not be overridden by WATSONX_SPACE_ID in env."""
_, mock_foundation = _make_reranker(
env={"WATSONX_API_KEY": "key", "WATSONX_SPACE_ID": "stray-env-space"},
project_id="explicit-proj",
)
_, call_kwargs = mock_foundation.Rerank.call_args
assert call_kwargs.get("project_id") == "explicit-proj"
assert "space_id" not in call_kwargs
def test_explicit_space_id_wins_over_env_project_id(self):
"""Explicit space_id must not be overridden by WATSONX_PROJECT_ID in env."""
_, mock_foundation = _make_reranker(
env={"WATSONX_API_KEY": "key", "WATSONX_PROJECT_ID": "stray-env-proj"},
space_id="explicit-space",
)
_, call_kwargs = mock_foundation.Rerank.call_args
assert call_kwargs.get("space_id") == "explicit-space"
assert "project_id" not in call_kwargs
def test_both_explicit_raises(self):
from lancedb.rerankers.watsonx import WatsonxReranker
reranker = WatsonxReranker(project_id="p", space_id="s")
mock_ibm = MagicMock()
mock_foundation = MagicMock()
def _fake_import(name):
if name == "ibm_watsonx_ai":
return mock_ibm
if name == "ibm_watsonx_ai.foundation_models":
return mock_foundation
raise ImportError(name)
with patch.dict("os.environ", {"WATSONX_API_KEY": "key"}, clear=True):
with patch(
"lancedb.rerankers.watsonx.attempt_import_or_raise",
side_effect=_fake_import,
):
with pytest.raises(ValueError, match="not both"):
_ = reranker._client
def test_neither_raises(self):
from lancedb.rerankers.watsonx import WatsonxReranker
reranker = WatsonxReranker()
mock_ibm = MagicMock()
mock_foundation = MagicMock()
def _fake_import(name):
if name == "ibm_watsonx_ai":
return mock_ibm
if name == "ibm_watsonx_ai.foundation_models":
return mock_foundation
raise ImportError(name)
with patch.dict("os.environ", {"WATSONX_API_KEY": "key"}, clear=True):
with patch(
"lancedb.rerankers.watsonx.attempt_import_or_raise",
side_effect=_fake_import,
):
with pytest.raises(
ValueError, match="WATSONX_PROJECT_ID or WATSONX_SPACE_ID"
):
_ = reranker._client