Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/common_utils/semantic_text_index.py: 43%

100 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1"""Embedding-similarity ranking over short texts with a per-model vector cache, shared by agent search and MCP tool search.""" 

2 

3from __future__ import annotations 

4 

5import math 

6from collections.abc import Awaitable, Mapping, Sequence 

7from dataclasses import dataclass 

8from itertools import chain, islice 

9from types import MappingProxyType 

10from typing import TYPE_CHECKING, Final, Protocol, TypeAlias 

11 

12from fastapi import HTTPException 

13from openai import OpenAIError 

14from pydantic import BaseModel, ConfigDict 

15 

16from litellm.exceptions import BudgetExceededError 

17 

18if TYPE_CHECKING: 18 ↛ 19line 18 didn't jump to line 19 because the condition on line 18 was never true

19 from litellm.proxy._types import UserAPIKeyAuth 

20 from litellm.proxy.utils import ProxyLogging 

21 from litellm.router import Router 

22 

23Vector: TypeAlias = tuple[float, ...] 

24 

25DEFAULT_MAX_CACHED_VECTORS: Final = 5000 

26"""Ceiling on how many (embedding model, text) vectors one index keeps; the least recently searched are evicted first.""" 

27 

28 

29class Embedder(Protocol): 

30 def __call__(self, texts: Sequence[str]) -> Awaitable[Sequence[Vector]]: ... 30 ↛ exitline 30 didn't return from function '__call__' because

31 

32 

33@dataclass(frozen=True, slots=True) 

34class EmbeddingFailed: 

35 reason: str 

36 

37 

38class _EmbeddingItem(BaseModel): 

39 model_config = ConfigDict(frozen=True, extra="ignore") 

40 

41 embedding: tuple[float, ...] 

42 

43 

44class _EmbeddingData(BaseModel): 

45 model_config = ConfigDict(frozen=True, extra="ignore") 

46 

47 data: tuple[_EmbeddingItem, ...] 

48 

49 

50class _EmbeddingRequest(BaseModel): 

51 """The /embeddings-shaped request as the pre-call hooks (rate limits, budgets, guardrails) hand it back.""" 

52 

53 model_config = ConfigDict(frozen=True, extra="ignore") 

54 

55 model: str 

56 input: tuple[str, ...] 

57 metadata: dict[str, object] # mutable-ok: the router mutates the metadata dict it is handed 

58 

59 

60def cosine_similarity(left: Vector, right: Vector) -> float: 

61 dot: Final = sum(a * b for a, b in zip(left, right, strict=True)) 

62 norms: Final = math.sqrt(sum(a * a for a in left)) * math.sqrt(sum(b * b for b in right)) 

63 return dot / norms if norms else 0.0 

64 

65 

66def embedding_spend_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, object]: # mutable-ok: router mutates it 

67 from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup 

68 

69 return { # mutable-ok: the router mutates the metadata dict it is handed 

70 **LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict), 

71 "user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict), 

72 } 

73 

74 

75def router_embedder( 

76 router: Router, embedding_model: str, user_api_key_dict: UserAPIKeyAuth, proxy_logging_obj: ProxyLogging 

77) -> Embedder: 

78 """Embeds through the router after the same key rate-limit, budget and guardrail pre-call hooks /embeddings runs.""" 

79 

80 async def embed(texts: Sequence[str]) -> Sequence[Vector]: 

81 request: Final = { # mutable-ok: pre_call_hook mutates the request dict in place 

82 "model": embedding_model, 

83 "input": list(texts), # mutable-ok: Router.aembedding accepts only str | list input 

84 "metadata": embedding_spend_metadata(user_api_key_dict), 

85 } 

86 processed: Final = _EmbeddingRequest.model_validate( 

87 await proxy_logging_obj.pre_call_hook( 

88 user_api_key_dict=user_api_key_dict, data=request, call_type="aembedding" 

89 ) 

90 ) 

91 response: Final = await router.aembedding( 

92 model=processed.model, 

93 input=list(processed.input), # mutable-ok: Router.aembedding accepts only str | list input 

94 metadata=processed.metadata, 

95 ) 

96 return tuple(item.embedding for item in _EmbeddingData.model_validate(response.model_dump()).data) 

97 

98 return embed 

99 

100 

101_CacheKey: TypeAlias = tuple[str, str] 

102 

103 

104async def _embed_all(embed: Embedder, texts: Sequence[str]) -> tuple[Vector, ...] | EmbeddingFailed: 

105 try: 

106 vectors: Final = tuple(await embed(texts)) 

107 except HTTPException: 

108 raise 

109 except (OpenAIError, ValueError, BudgetExceededError) as exc: 

110 return EmbeddingFailed(reason=f"embedding the search query failed: {exc}") 

111 if len(vectors) != len(texts): 

112 return EmbeddingFailed(reason=f"embedding model returned {len(vectors)} vectors for {len(texts)} inputs") 

113 return vectors 

114 

115 

116@dataclass(frozen=True, slots=True) 

117class _Embedded: 

118 query_vector: Vector 

119 vectors: Mapping[str, Vector] 

120 

121 

122def _same_dimension(query_vector: Vector, vectors: Mapping[str, Vector], texts: Sequence[str]) -> bool: 

123 return all(len(vectors[text]) == len(query_vector) for text in texts) 

124 

125 

126async def _embed_query_and_texts( 

127 embed: Embedder, query: str, texts: Sequence[str], cached: Mapping[str, Vector] 

128) -> _Embedded | EmbeddingFailed: 

129 missing: Final = tuple(dict.fromkeys(text for text in texts if text not in cached)) 

130 embedded: Final = await _embed_all(embed, (query, *missing)) 

131 if isinstance(embedded, EmbeddingFailed): 

132 return embedded 

133 vectors: Final = MappingProxyType(dict(chain(cached.items(), zip(missing, embedded[1:], strict=True)))) 

134 if _same_dimension(embedded[0], vectors, texts): 

135 return _Embedded(query_vector=embedded[0], vectors=vectors) 

136 unique: Final = tuple(dict.fromkeys(texts)) 

137 reembedded: Final = await _embed_all(embed, (query, *unique)) 

138 if isinstance(reembedded, EmbeddingFailed): 

139 return reembedded 

140 return _Embedded( 

141 query_vector=reembedded[0], vectors=MappingProxyType(dict(zip(unique, reembedded[1:], strict=True))) 

142 ) 

143 

144 

145class SemanticTextIndex: 

146 """Caches one vector per distinct text per embedding model, so repeat searches only embed the query. 

147 

148 Holds at most ``max_entries`` vectors across all models: once full, the texts no recent search touched go first.""" 

149 

150 def __init__(self, max_entries: int = DEFAULT_MAX_CACHED_VECTORS) -> None: 

151 self._max_entries: Final = max_entries 

152 self._vectors: Mapping[_CacheKey, Vector] = MappingProxyType({}) 

153 

154 def _cached(self, embedding_model: str) -> Mapping[str, Vector]: 

155 return MappingProxyType( 

156 {text: vector for (model, text), vector in self._vectors.items() if model == embedding_model} 

157 ) 

158 

159 def _merged(self, embedding_model: str, embedded: _Embedded, texts: Sequence[str]) -> Mapping[_CacheKey, Vector]: 

160 dimension: Final = len(embedded.query_vector) 

161 touched: Final = MappingProxyType({(embedding_model, text): embedded.vectors[text] for text in texts}) 

162 untouched: Final = MappingProxyType( 

163 { 

164 key: vector 

165 for key, vector in chain( 

166 self._vectors.items(), 

167 (((embedding_model, text), vector) for text, vector in embedded.vectors.items()), 

168 ) 

169 if key not in touched and (key[0] != embedding_model or len(vector) == dimension) 

170 } 

171 ) 

172 ordered: Final = MappingProxyType({**untouched, **touched}) 

173 return MappingProxyType(dict(islice(ordered.items(), max(len(ordered) - self._max_entries, 0), None))) 

174 

175 async def scores( 

176 self, query: str, texts: Sequence[str], embed: Embedder, embedding_model: str 

177 ) -> tuple[float, ...] | EmbeddingFailed: 

178 """Cosine similarity of `query` to each entry of `texts`, in the same order.""" 

179 if not texts: 

180 return () 

181 embedded: Final = await _embed_query_and_texts(embed, query, texts, self._cached(embedding_model)) 

182 if isinstance(embedded, EmbeddingFailed): 

183 return embedded 

184 if not _same_dimension(embedded.query_vector, embedded.vectors, texts): 

185 return EmbeddingFailed(reason=f"embedding model {embedding_model} returned vectors of mixed dimensions") 

186 self._vectors = self._merged(embedding_model, embedded, texts) 

187 return tuple(cosine_similarity(embedded.query_vector, embedded.vectors[text]) for text in texts)