mirror of
https://github.com/langflow-ai/langflow.git
synced 2026-07-23 23:13:58 +08:00
fix: Respect OpenAI embeddings API base for Knowledge Base (#13380)
* fix: respect OpenAI embeddings API base * [autofix.ci] apply automated fixes --------- Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
@ -4169,7 +4169,7 @@
|
||||
"filename": "src/backend/tests/unit/test_unified_models.py",
|
||||
"hashed_secret": "e9a5f12a8ecbb3eb46eca5096b5c52aa5e7c9fdd",
|
||||
"is_verified": false,
|
||||
"line_number": 494
|
||||
"line_number": 517
|
||||
}
|
||||
],
|
||||
"src/backend/tests/unit/test_user.py": [
|
||||
@ -9535,5 +9535,5 @@
|
||||
}
|
||||
]
|
||||
},
|
||||
"generated_at": "2026-05-26T16:55:05Z"
|
||||
"generated_at": "2026-05-28T15:17:56Z"
|
||||
}
|
||||
|
||||
@ -536,6 +536,62 @@ def test_get_embeddings_openai_basic(mock_get_class, mock_get_api_key):
|
||||
assert kwargs["api_key"] == "sk-test" # pragma: allowlist secret
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("env_values", "expected_base_url"),
|
||||
[
|
||||
({"OPENAI_EMBEDDINGS_API_BASE": "http://embeddings.example/v1"}, "http://embeddings.example/v1"),
|
||||
({"OPENAI_API_BASE": "http://openai-compatible.example/v1"}, "http://openai-compatible.example/v1"),
|
||||
(
|
||||
{
|
||||
"OPENAI_EMBEDDINGS_API_BASE": "http://embeddings.example/v1",
|
||||
"OPENAI_API_BASE": "http://openai-compatible.example/v1",
|
||||
},
|
||||
"http://embeddings.example/v1",
|
||||
),
|
||||
],
|
||||
)
|
||||
@patch("lfx.base.models.unified_models.get_api_key_for_provider")
|
||||
@patch("lfx.base.models.unified_models.get_embedding_class")
|
||||
def test_get_embeddings_openai_api_base_env_fallback(
|
||||
mock_get_class,
|
||||
mock_get_api_key,
|
||||
monkeypatch,
|
||||
env_values,
|
||||
expected_base_url,
|
||||
):
|
||||
mock_get_api_key.return_value = "sk-test"
|
||||
mock_embedding_class = MagicMock()
|
||||
mock_get_class.return_value = mock_embedding_class
|
||||
monkeypatch.delenv("OPENAI_EMBEDDINGS_API_BASE", raising=False)
|
||||
monkeypatch.delenv("OPENAI_API_BASE", raising=False)
|
||||
for name, value in env_values.items():
|
||||
monkeypatch.setenv(name, value)
|
||||
|
||||
get_embeddings([_make_openai_embedding_model()], api_key="sk-test")
|
||||
|
||||
kwargs = mock_embedding_class.call_args.kwargs
|
||||
assert kwargs["base_url"] == expected_base_url
|
||||
|
||||
|
||||
@patch("lfx.base.models.unified_models.get_api_key_for_provider")
|
||||
@patch("lfx.base.models.unified_models.get_embedding_class")
|
||||
def test_get_embeddings_openai_explicit_api_base_overrides_env(mock_get_class, mock_get_api_key, monkeypatch):
|
||||
mock_get_api_key.return_value = "sk-test"
|
||||
mock_embedding_class = MagicMock()
|
||||
mock_get_class.return_value = mock_embedding_class
|
||||
monkeypatch.setenv("OPENAI_EMBEDDINGS_API_BASE", "http://embeddings.example/v1")
|
||||
monkeypatch.setenv("OPENAI_API_BASE", "http://openai-compatible.example/v1")
|
||||
|
||||
get_embeddings(
|
||||
[_make_openai_embedding_model()],
|
||||
api_key="sk-test", # pragma: allowlist secret
|
||||
api_base="http://component.example/v1",
|
||||
)
|
||||
|
||||
kwargs = mock_embedding_class.call_args.kwargs
|
||||
assert kwargs["base_url"] == "http://component.example/v1"
|
||||
|
||||
|
||||
@patch("lfx.base.models.unified_models.get_api_key_for_provider")
|
||||
@patch("lfx.base.models.unified_models.get_embedding_class")
|
||||
def test_get_embeddings_optional_params_only_added_when_mapped(mock_get_class, mock_get_api_key):
|
||||
|
||||
@ -271,6 +271,11 @@ def get_embeddings(
|
||||
model_name = model_dict.get("name")
|
||||
provider = model_dict.get("provider")
|
||||
metadata = model_dict.get("metadata", {})
|
||||
api_base_value = _to_str(api_base)
|
||||
if provider == "OpenAI" and not api_base_value:
|
||||
api_base_value = _to_str(os.environ.get("OPENAI_EMBEDDINGS_API_BASE")) or _to_str(
|
||||
os.environ.get("OPENAI_API_BASE")
|
||||
)
|
||||
|
||||
# --- resolve API key -----------------------------------------------------
|
||||
api_key = unified_models_module.get_api_key_for_provider(user_id, provider, api_key)
|
||||
@ -326,7 +331,7 @@ def get_embeddings(
|
||||
# Optional parameters - only add when both a value is supplied *and* the
|
||||
# provider's param_mapping declares the corresponding key.
|
||||
optional_params: dict[str, Any] = {
|
||||
"api_base": _to_str(api_base) or None,
|
||||
"api_base": api_base_value or None,
|
||||
"dimensions": int(dimensions) if dimensions else None,
|
||||
"chunk_size": int(chunk_size) if chunk_size else None,
|
||||
"request_timeout": float(request_timeout) if request_timeout else None,
|
||||
|
||||
Reference in New Issue
Block a user