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(