fix embedding cache
frdel committed
Jul 5, 2025 at 16:06 UTC
b7f1daacfa05e43423cfd014f15995ca035088c6
3 files changed
+26
-27
models.py
+14
-8
@@ -1,4 +1,5 @@
1
from enum import Enum
2
+import logging
3
import os
4
from typing import (
5
Any,
@@ -14,6 +15,7 @@ from typing import (
15
16
from litellm import completion, acompletion, embedding
17
import litellm
18
+
19
from python.helpers import dotenv
20
from python.helpers.dotenv import load_dotenv
21
from python.helpers.rate_limiter import RateLimiter
@@ -35,10 +37,14 @@ from langchain.embeddings.base import Embeddings
37
from sentence_transformers import SentenceTransformer
38
39
38
-# disable extra logging
40
+# disable extra logging, must be done repeatedly, otherwise browser-use will turn it back on for some reason
41
def turn_off_logging():
42
os.environ["LITELLM_LOG"] = "ERROR" # only errors
43
litellm.suppress_debug_info = True
44
+ # Silence **all** LiteLLM sub-loggers (utils, cost_calculator…)
45
+ for name in logging.Logger.manager.loggerDict:
46
+ if name.lower().startswith("litellm"):
47
+ logging.getLogger(name).setLevel(logging.ERROR)
48
49
50
# init
@@ -392,13 +398,13 @@ class LiteLLMEmbeddingWrapper(Embeddings):
398
class LocalSentenceTransformerWrapper(Embeddings):
399
"""Local wrapper for sentence-transformers models to avoid HuggingFace API calls"""
400
395
- def __init__(self, model_name: str, **kwargs: Any):
401
+ def __init__(self, provider: str, model: str, **kwargs: Any):
402
# Remove the "sentence-transformers/" prefix if present
397
- if model_name.startswith("sentence-transformers/"):
398
- model_name = model_name[len("sentence-transformers/") :]
403
+ if model.startswith("sentence-transformers/"):
404
+ model = model[len("sentence-transformers/") :]
405
400
- self.model = SentenceTransformer(model_name, **kwargs)
401
- self.model_name = model_name
406
+ self.model = SentenceTransformer(model, **kwargs)
407
+ self.model_name = model
408
409
def embed_documents(self, texts: List[str]) -> List[List[float]]:
410
embeddings = self.model.encode(texts, convert_to_tensor=False) # type: ignore
@@ -450,14 +456,14 @@ def get_litellm_embedding(model_name: str, provider: str, **kwargs: Any):
456
# Check if this is a local sentence-transformers model
457
if provider == "huggingface" and model_name.startswith("sentence-transformers/"):
458
# Use local sentence-transformers instead of LiteLLM for local models
453
- return LocalSentenceTransformerWrapper(model_name=model_name, **kwargs)
459
+ return LocalSentenceTransformerWrapper(provider=provider, model=model_name, **kwargs)
460
461
configure_litellm_environment()
462
# Use original provider name for API key lookup, fallback to mapped provider name
463
api_key = kwargs.pop("api_key", None) or get_api_key(provider)
464
465
# litellm will pick up base_url from env. We just need to control the api_key.
460
- base_url = dotenv.get_dotenv_value(f"{provider.upper()}_BASE_URL")
466
+ # base_url = dotenv.get_dotenv_value(f"{provider.upper()}_BASE_URL")
467
468
# If a base_url is set, ensure api_key is not passed to litellm
469
# > remove, this can be handled by api_key=None
python/helpers/document_query.py
+4
-14
@@ -20,21 +20,10 @@ from langchain_community.document_transformers import MarkdownifyTransformer
20
from langchain_community.document_loaders.parsers.images import TesseractBlobParser
21
22
from langchain_core.documents import Document
23
-from langchain.prompts import ChatPromptTemplate
23
from langchain.schema import SystemMessage, HumanMessage
25
-from langchain.storage import LocalFileStore, InMemoryStore
26
-from langchain.embeddings import CacheBackedEmbeddings
27
-
28
-from langchain_community.vectorstores import FAISS
29
-import faiss
30
-from langchain_community.docstore.in_memory import InMemoryDocstore
31
-from langchain_community.vectorstores.utils import (
32
- DistanceStrategy,
33
-)
34
-from langchain_core.embeddings import Embeddings
24
25
from python.helpers.print_style import PrintStyle
37
-from python.helpers import files
26
+from python.helpers import files, errors
27
from agent import Agent
28
29
from langchain.text_splitter import RecursiveCharacterTextSplitter
@@ -106,7 +95,7 @@ class DocumentQueryStore:
95
return normalized
96
97
def init_vector_db(self):
109
- return VectorDB(self.agent)
98
+ return VectorDB(self.agent, cache=True)
99
100
async def add_document(
101
self, text: str, document_uri: str, metadata: dict | None = None
@@ -168,7 +157,8 @@ class DocumentQueryStore:
157
)
158
return True, ids
159
except Exception as e:
171
- PrintStyle.error(f"Error adding document '{document_uri}': {str(e)}")
160
+ err_text = errors.format_error(e)
161
+ PrintStyle.error(f"Error adding document '{document_uri}': {err_text}")
162
return False, []
163
164
async def get_document(self, document_uri: str) -> Optional[Document]:
python/helpers/vector_db.py
+8
-5
@@ -36,12 +36,14 @@ class VectorDB:
36
_cached_embeddings: dict[str, CacheBackedEmbeddings] = {}
37
38
@staticmethod
39
- def _get_embeddings(agent: Agent):
39
+ def _get_embeddings(agent: Agent, cache: bool = True):
40
model = agent.get_embedding_model()
41
+ if not cache:
42
+ return model # return raw embeddings if cache is False
43
namespace = getattr(
44
model,
43
- "model",
44
- getattr(model, "model_name", "default"),
45
+ "model_name",
46
+ "default",
47
)
48
if namespace not in VectorDB._cached_embeddings:
49
store = InMemoryByteStore()
@@ -54,9 +56,10 @@ class VectorDB:
56
)
57
return VectorDB._cached_embeddings[namespace]
58
57
- def __init__(self, agent: Agent):
59
+ def __init__(self, agent: Agent, cache: bool = True):
60
self.agent = agent
59
- self.embeddings = self._get_embeddings(agent)
61
+ self.cache = cache # store cache preference
62
+ self.embeddings = self._get_embeddings(agent, cache=cache)
63
self.index = faiss.IndexFlatIP(len(self.embeddings.embed_query("example")))
64
65
self.db = MyFaiss(