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
« 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."""
3from __future__ import annotations
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
12from fastapi import HTTPException
13from openai import OpenAIError
14from pydantic import BaseModel, ConfigDict
16from litellm.exceptions import BudgetExceededError
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
23Vector: TypeAlias = tuple[float, ...]
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."""
29class Embedder(Protocol):
30 def __call__(self, texts: Sequence[str]) -> Awaitable[Sequence[Vector]]: ... 30 ↛ exitline 30 didn't return from function '__call__' because
33@dataclass(frozen=True, slots=True)
34class EmbeddingFailed:
35 reason: str
38class _EmbeddingItem(BaseModel):
39 model_config = ConfigDict(frozen=True, extra="ignore")
41 embedding: tuple[float, ...]
44class _EmbeddingData(BaseModel):
45 model_config = ConfigDict(frozen=True, extra="ignore")
47 data: tuple[_EmbeddingItem, ...]
50class _EmbeddingRequest(BaseModel):
51 """The /embeddings-shaped request as the pre-call hooks (rate limits, budgets, guardrails) hand it back."""
53 model_config = ConfigDict(frozen=True, extra="ignore")
55 model: str
56 input: tuple[str, ...]
57 metadata: dict[str, object] # mutable-ok: the router mutates the metadata dict it is handed
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
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
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 }
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."""
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)
98 return embed
101_CacheKey: TypeAlias = tuple[str, str]
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
116@dataclass(frozen=True, slots=True)
117class _Embedded:
118 query_vector: Vector
119 vectors: Mapping[str, Vector]
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)
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 )
145class SemanticTextIndex:
146 """Caches one vector per distinct text per embedding model, so repeat searches only embed the query.
148 Holds at most ``max_entries`` vectors across all models: once full, the texts no recent search touched go first."""
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({})
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 )
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)))
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)