mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
test: add regression tests for sync embedding access-group scoping and cache metadata forwarding (#31260)
This commit is contained in:
parent
c2c2a623c0
commit
feb0f4e63a
2 changed files with 82 additions and 0 deletions
|
|
@ -6,6 +6,7 @@ need to be properly propagated through the router to the LLM API.
|
|||
"""
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -13,6 +14,7 @@ import pytest
|
|||
import respx
|
||||
|
||||
import litellm
|
||||
import litellm.router_strategy.simple_shuffle as simple_shuffle
|
||||
from litellm import Router
|
||||
from litellm.llms.base_llm.vector_store.transformation import (
|
||||
LiteLLMVectorStoreEmbeddingExecutor,
|
||||
|
|
@ -383,6 +385,55 @@ class TestRouterEmbeddingIntegration:
|
|||
# The call should succeed
|
||||
mock_aembedding.assert_called_once()
|
||||
|
||||
def test_sync_embedding_respects_model_access_group_scoping(self, monkeypatch):
|
||||
"""
|
||||
Regression test for #31260: sync Router._embedding must forward the
|
||||
caller's request_kwargs to deployment selection. Without the forward,
|
||||
_filter_deployments_by_model_access_groups short-circuits
|
||||
(request_kwargs is None) and every deployment in the model group is
|
||||
treated as globally accessible. random.choice is forced to the last
|
||||
candidate so an unfiltered deployment list fails this test
|
||||
deterministically instead of on 50% of runs.
|
||||
"""
|
||||
model_list = [
|
||||
{
|
||||
"model_name": "grouped-embed",
|
||||
"litellm_params": {
|
||||
"model": "text-embedding-3-small",
|
||||
"api_key": "eng-key",
|
||||
},
|
||||
"model_info": {"access_groups": ["engineering"]},
|
||||
},
|
||||
{
|
||||
"model_name": "grouped-embed",
|
||||
"litellm_params": {
|
||||
"model": "text-embedding-3-small",
|
||||
"api_key": "fin-key",
|
||||
},
|
||||
"model_info": {"access_groups": ["finance"]},
|
||||
},
|
||||
]
|
||||
|
||||
router = Router(model_list=model_list)
|
||||
engineering_auth = SimpleNamespace(models=["engineering"], team_models=[])
|
||||
|
||||
fake_random = MagicMock()
|
||||
fake_random.choice.side_effect = lambda seq: seq[-1]
|
||||
|
||||
monkeypatch.setattr(litellm, "embedding", MagicMock(return_value=MagicMock(data=[{"embedding": [0.1, 0.2]}])))
|
||||
monkeypatch.setattr(simple_shuffle, "random", fake_random)
|
||||
|
||||
mock_embedding = litellm.embedding
|
||||
router.embedding(
|
||||
model="grouped-embed",
|
||||
input=["hello"],
|
||||
metadata={"user_api_key_auth": engineering_auth},
|
||||
)
|
||||
|
||||
# The finance deployment must have been filtered out; only the
|
||||
# engineering deployment can be selected.
|
||||
assert mock_embedding.call_args[1]["api_key"] == "eng-key"
|
||||
|
||||
def test_embedding_with_timeout_from_router(self):
|
||||
"""
|
||||
Test that timeout settings from router config are propagated.
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import logging
|
||||
import re
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -252,3 +253,33 @@ def test_exact_cache_key_includes_anthropic_messages_params(anthropic_param):
|
|||
assert baseline != cache.get_cache_key(
|
||||
model="claude-sonnet-4-5", messages=messages, **anthropic_param
|
||||
)
|
||||
|
||||
|
||||
def test_sync_get_cache_forwards_metadata_subset_to_backend():
|
||||
"""
|
||||
Regression test for #31260: the sync cache read path used to flatten
|
||||
metadata to {}, so semantic-cache backends never saw the caller's auth
|
||||
context (team id / user_api_key_auth) on a get, while the write path and
|
||||
the async read path forwarded full metadata. The metadata subset must be
|
||||
preserved so sync semantic-cache lookups can scope and authenticate the
|
||||
same way async ones do.
|
||||
"""
|
||||
cache = Cache(type=LiteLLMCacheType.LOCAL)
|
||||
backend = MagicMock()
|
||||
backend.get_cache.return_value = None
|
||||
|
||||
cache.get_cache(
|
||||
dynamic_cache_object=backend,
|
||||
cache={"use-cache": True},
|
||||
model="text-embedding-3-small",
|
||||
input=["hello"],
|
||||
metadata={
|
||||
"user_api_key_team_id": "team-a",
|
||||
"user_api_key": "k-abc",
|
||||
},
|
||||
)
|
||||
|
||||
assert backend.get_cache.called
|
||||
sent_metadata = backend.get_cache.call_args.kwargs["metadata"]
|
||||
assert sent_metadata["user_api_key_team_id"] == "team-a"
|
||||
assert sent_metadata["user_api_key"] == "k-abc"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue