quantumbit commited on
Commit
dd5974e
·
verified ·
1 Parent(s): 353f5fa

Upload folder using huggingface_hub

Browse files
rag_system/api.py CHANGED
@@ -20,12 +20,12 @@ from .document_processor import process_texts, process_file
20
  from .vector_store import (
21
  add_documents, load_or_create_store, is_loaded,
22
  list_collections, get_collection_stats, delete_collection,
23
- cleanup_stale_collections,
24
  )
25
  from .query_engine import query as run_query, stream_query, pipeline_stream_query
26
  from .eval import evaluate
27
  from .cache import cache_connected, get_cache_stats
28
- from .embeddings import get_embeddings
29
  from .guardrails import _load_llama_guard
30
  from .retriever import _reranker
31
 
@@ -202,6 +202,11 @@ async def cache_stats():
202
  return get_cache_stats()
203
 
204
 
 
 
 
 
 
205
  # ── Ingest ────────────────────────────────────────────────────────────────────
206
 
207
  @app.post("/ingest", response_model=IngestResponse, tags=["ingest"])
@@ -213,7 +218,12 @@ async def ingest_texts(req: IngestRequest):
213
  metadatas=req.metadatas,
214
  source_id=req.collection_name,
215
  )
216
- add_documents(docs, collection=req.collection_name, force_reindex=req.force_reindex)
 
 
 
 
 
217
  return IngestResponse(
218
  success=True,
219
  docs_indexed=len(docs),
@@ -229,6 +239,7 @@ async def ingest_texts(req: IngestRequest):
229
  async def ingest_file(
230
  file: UploadFile = File(...),
231
  collection_name: str = "default",
 
232
  background_tasks: BackgroundTasks = None,
233
  ):
234
  """
@@ -254,6 +265,7 @@ async def ingest_file(
254
  "status": "processing",
255
  "collection_name": doc_collection,
256
  "filename": file.filename,
 
257
  "chunks_created": 0,
258
  "message": "File received, extracting text...",
259
  "progress": 5,
@@ -267,7 +279,7 @@ async def ingest_file(
267
 
268
  _ingest_jobs[job_id]["progress"] = 60
269
  _ingest_jobs[job_id]["message"] = f"Embedding and indexing {len(docs)} chunks..."
270
- add_documents(docs, collection=doc_collection)
271
  _viz_cache.pop(doc_collection, None) # invalidate stale viz
272
 
273
  _ingest_jobs[job_id]["progress"] = 100
@@ -463,7 +475,8 @@ async def get_query_similarity(collection_name: str, body: dict = Body(...)):
463
 
464
  import numpy as np
465
 
466
- q_vec = np.array(get_embeddings().embed_query(query), dtype=np.float32).reshape(1, -1)
 
467
  q_2d = result["pca"].transform(q_vec)[0]
468
 
469
  vectors = result["vectors"]
 
20
  from .vector_store import (
21
  add_documents, load_or_create_store, is_loaded,
22
  list_collections, get_collection_stats, delete_collection,
23
+ cleanup_stale_collections, get_collection_embedding_mode,
24
  )
25
  from .query_engine import query as run_query, stream_query, pipeline_stream_query
26
  from .eval import evaluate
27
  from .cache import cache_connected, get_cache_stats
28
+ from .embeddings import get_embeddings, get_embeddings_runtime_info
29
  from .guardrails import _load_llama_guard
30
  from .retriever import _reranker
31
 
 
202
  return get_cache_stats()
203
 
204
 
205
+ @app.get("/embeddings/info", tags=["ops"])
206
+ async def embeddings_info():
207
+ return get_embeddings_runtime_info()
208
+
209
+
210
  # ── Ingest ────────────────────────────────────────────────────────────────────
211
 
212
  @app.post("/ingest", response_model=IngestResponse, tags=["ingest"])
 
218
  metadatas=req.metadatas,
219
  source_id=req.collection_name,
220
  )
221
+ add_documents(
222
+ docs,
223
+ collection=req.collection_name,
224
+ force_reindex=req.force_reindex,
225
+ embedding_mode=req.embedding_mode,
226
+ )
227
  return IngestResponse(
228
  success=True,
229
  docs_indexed=len(docs),
 
239
  async def ingest_file(
240
  file: UploadFile = File(...),
241
  collection_name: str = "default",
242
+ embedding_mode: str | None = None,
243
  background_tasks: BackgroundTasks = None,
244
  ):
245
  """
 
265
  "status": "processing",
266
  "collection_name": doc_collection,
267
  "filename": file.filename,
268
+ "embedding_mode": embedding_mode,
269
  "chunks_created": 0,
270
  "message": "File received, extracting text...",
271
  "progress": 5,
 
279
 
280
  _ingest_jobs[job_id]["progress"] = 60
281
  _ingest_jobs[job_id]["message"] = f"Embedding and indexing {len(docs)} chunks..."
282
+ add_documents(docs, collection=doc_collection, embedding_mode=embedding_mode)
283
  _viz_cache.pop(doc_collection, None) # invalidate stale viz
284
 
285
  _ingest_jobs[job_id]["progress"] = 100
 
475
 
476
  import numpy as np
477
 
478
+ embedding_mode = get_collection_embedding_mode(collection_name)
479
+ q_vec = np.array(get_embeddings(embedding_mode).embed_query(query), dtype=np.float32).reshape(1, -1)
480
  q_2d = result["pca"].transform(q_vec)[0]
481
 
482
  vectors = result["vectors"]
rag_system/cache.py CHANGED
@@ -63,20 +63,20 @@ def _build_redis_client():
63
  _cache_client = _build_redis_client()
64
 
65
  # Exact match cache
66
- def _cache_key(query: str, collection: str, mode: str) -> str:
67
- payload = f"{query}::{collection}::{mode}"
68
  return "rag:exact:" + hashlib.sha256(payload.encode()).hexdigest()[:32]
69
 
70
- def get_exact(query: str, collection: str, mode: str) -> Optional[dict]:
71
- key = _cache_key(query, collection, mode)
72
  raw = _cache_client.get(key)
73
  if raw:
74
  logger.debug(f"Exact cache hit: {key[:16]}...")
75
  return json.loads(raw)
76
  return None
77
 
78
- def set_exact(query: str, collection: str, mode: str, value: str) -> None:
79
- key = _cache_key(query,collection,mode)
80
  serialized = json.dumps(value)
81
  if hasattr(_cache_client,"setex"):
82
  _cache_client.setex(key,settings.cache_ttl_seconds,serialized)
@@ -85,13 +85,14 @@ def set_exact(query: str, collection: str, mode: str, value: str) -> None:
85
 
86
  # Semantic Cache
87
  # stores (embedding, serialized_response) pairs keyed by short hash
88
- _semantic_index: list[tuple[list[float],str,dict]] = [] # (vec,key,response)
89
 
90
- def get_semantic(query_vec: list[float]) -> Optional[dict]:
91
  """Return the cache response if cosine similarity > threshold"""
 
92
  best_score = 0.0
93
  best_response = None
94
- for vec, _,response in _semantic_index:
95
  score = cosine_similarity(query_vec,vec)
96
  if score > best_score:
97
  best_score = score
@@ -101,11 +102,12 @@ def get_semantic(query_vec: list[float]) -> Optional[dict]:
101
  return best_response
102
  return None
103
 
104
- def set_semantic(query_vec: list[float], query: str, response: dict) -> None:
105
  h = hashlib.md5(query.encode()).hexdigest()[:8]
106
- _semantic_index.append((query_vec, h, response))
107
- if len(_semantic_index) > 5000: # cap memory
108
- _semantic_index.pop(0)
 
109
 
110
  def cache_connected() -> bool:
111
  try:
@@ -125,7 +127,7 @@ def get_cache_stats() -> dict:
125
  except:
126
  stats["exact_matches_cached"] = "unknown"
127
 
128
- stats["semantic_matches_cached"] = len(_semantic_index)
129
  return stats
130
 
131
  print("[cache] Module ready")
 
63
  _cache_client = _build_redis_client()
64
 
65
  # Exact match cache
66
+ def _cache_key(query: str, collection: str, mode: str, embedding_mode: str) -> str:
67
+ payload = f"{query}::{collection}::{mode}::{embedding_mode}"
68
  return "rag:exact:" + hashlib.sha256(payload.encode()).hexdigest()[:32]
69
 
70
+ def get_exact(query: str, collection: str, mode: str, embedding_mode: str) -> Optional[dict]:
71
+ key = _cache_key(query, collection, mode, embedding_mode)
72
  raw = _cache_client.get(key)
73
  if raw:
74
  logger.debug(f"Exact cache hit: {key[:16]}...")
75
  return json.loads(raw)
76
  return None
77
 
78
+ def set_exact(query: str, collection: str, mode: str, embedding_mode: str, value: str) -> None:
79
+ key = _cache_key(query, collection, mode, embedding_mode)
80
  serialized = json.dumps(value)
81
  if hasattr(_cache_client,"setex"):
82
  _cache_client.setex(key,settings.cache_ttl_seconds,serialized)
 
85
 
86
  # Semantic Cache
87
  # stores (embedding, serialized_response) pairs keyed by short hash
88
+ _semantic_index: dict[str, list[tuple[list[float],str,dict]]] = {} # mode -> (vec,key,response)
89
 
90
+ def get_semantic(query_vec: list[float], embedding_mode: str) -> Optional[dict]:
91
  """Return the cache response if cosine similarity > threshold"""
92
+ pool = _semantic_index.get(embedding_mode, [])
93
  best_score = 0.0
94
  best_response = None
95
+ for vec, _,response in pool:
96
  score = cosine_similarity(query_vec,vec)
97
  if score > best_score:
98
  best_score = score
 
102
  return best_response
103
  return None
104
 
105
+ def set_semantic(query_vec: list[float], query: str, response: dict, embedding_mode: str) -> None:
106
  h = hashlib.md5(query.encode()).hexdigest()[:8]
107
+ pool = _semantic_index.setdefault(embedding_mode, [])
108
+ pool.append((query_vec, h, response))
109
+ if len(pool) > 5000: # cap memory
110
+ pool.pop(0)
111
 
112
  def cache_connected() -> bool:
113
  try:
 
127
  except:
128
  stats["exact_matches_cached"] = "unknown"
129
 
130
+ stats["semantic_matches_cached"] = sum(len(v) for v in _semantic_index.values())
131
  return stats
132
 
133
  print("[cache] Module ready")
rag_system/config.py CHANGED
@@ -14,6 +14,8 @@ class Settings(BaseSettings):
14
  embedding_dimensions: int = 1024
15
  embedding_model_cpu: str = "BAAI/bge-small-en-v1.5"
16
  embedding_dimensions_cpu: int = 384
 
 
17
  embedding_device: str = "auto"
18
  embedding_batch_size: int = 32
19
  embedding_normalize: bool = True
@@ -81,5 +83,6 @@ print(
81
  "[Config] Loaded. Model: "
82
  f"{settings.chat_model}, Embedding GPU: {settings.embedding_model} ({settings.embedding_dimensions}), "
83
  f"Embedding CPU: {settings.embedding_model_cpu} ({settings.embedding_dimensions_cpu}), "
 
84
  f"Device: {settings.embedding_device}"
85
  )
 
14
  embedding_dimensions: int = 1024
15
  embedding_model_cpu: str = "BAAI/bge-small-en-v1.5"
16
  embedding_dimensions_cpu: int = 384
17
+ embedding_model_openai: str = "text-embedding-3-small"
18
+ embedding_dimensions_openai: int = 1536
19
  embedding_device: str = "auto"
20
  embedding_batch_size: int = 32
21
  embedding_normalize: bool = True
 
83
  "[Config] Loaded. Model: "
84
  f"{settings.chat_model}, Embedding GPU: {settings.embedding_model} ({settings.embedding_dimensions}), "
85
  f"Embedding CPU: {settings.embedding_model_cpu} ({settings.embedding_dimensions_cpu}), "
86
+ f"Embedding OpenAI: {settings.embedding_model_openai} ({settings.embedding_dimensions_openai}), "
87
  f"Device: {settings.embedding_device}"
88
  )
rag_system/embeddings.py CHANGED
@@ -1,21 +1,21 @@
1
  """
2
- Local embedding model with GPU/CPU auto selection.
3
- - GPU: BAAI/bge-large-en-v1.5 (1024-dim output)
4
- - CPU fallback: BAAI/bge-small-en-v1.5 (384-dim output)
5
- - Runs fully local via sentence-transformers zero API calls, zero cost
6
- - BGE requires a special query prefix: 'Represent this sentence for searching'
7
- (documents are embedded as-is; only queries get the prefix)
8
- - LangChain's HuggingFaceEmbeddings handles the prefix automatically
9
  """
10
 
11
  import asyncio
12
  import logging
13
  import warnings
14
  from concurrent.futures import ThreadPoolExecutor
15
- from functools import lru_cache
16
 
17
  import numpy as np
18
  from langchain_huggingface import HuggingFaceEmbeddings
 
 
19
  from pydantic.warnings import UnsupportedFieldAttributeWarning
20
 
21
  from .config import get_settings
@@ -31,6 +31,10 @@ warnings.filterwarnings("ignore", category=UnsupportedFieldAttributeWarning)
31
 
32
  _executor = ThreadPoolExecutor(max_workers=2)
33
 
 
 
 
 
34
  def _resolve_device() -> str:
35
  device = (settings.embedding_device or "auto").lower()
36
  if device == "auto":
@@ -55,31 +59,84 @@ def _resolve_device() -> str:
55
  return device
56
 
57
 
58
- def _select_embedding_config() -> tuple[str, int, str]:
59
- device = _resolve_device()
60
- if device == "cpu":
61
- model_name = settings.embedding_model_cpu or settings.embedding_model
62
- dimensions = settings.embedding_dimensions_cpu or settings.embedding_dimensions
63
- else:
64
- model_name = settings.embedding_model
65
- dimensions = settings.embedding_dimensions
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
66
 
67
- return model_name, dimensions, device
68
 
 
 
 
 
 
 
 
 
 
 
69
 
70
- def get_embedding_info() -> dict[str, str | int]:
71
- model_name, dimensions, device = _select_embedding_config()
 
 
 
 
 
 
 
 
 
 
 
72
  return {
73
- "model_name": model_name,
74
- "dimensions": dimensions,
75
- "device": device,
76
  }
77
 
78
 
79
- @lru_cache(maxsize=1)
80
- def get_embeddings() -> HuggingFaceEmbeddings:
81
  """
82
- Singleton embedding model
83
 
84
  encode_kwargs:
85
  normalize_embeddings=True -> required for cosine similarity to work correctly
@@ -88,23 +145,36 @@ def get_embeddings() -> HuggingFaceEmbeddings:
88
  BGE was finetuned with an instruction-like query prefix.
89
  We pass that prefix for query encoding only; documents remain unchanged.
90
  """
91
- model_name, dimensions, device = _select_embedding_config()
92
-
93
- logger.info(f"Loading embedding model: {model_name} on {device}")
94
- model = HuggingFaceEmbeddings(
95
- model_name=model_name,
96
- model_kwargs={
97
- "device": device,
98
- },
99
- encode_kwargs={
100
- "normalize_embeddings": settings.embedding_normalize,
101
- "batch_size": settings.embedding_batch_size,
102
- },
103
- query_encode_kwargs={
104
- "prompt": "Represent this sentence for searching relevant passages: ",
105
- },
106
- )
107
- logger.info(f"Embedding model loaded. Output dim={dimensions}")
 
 
 
 
 
 
 
 
 
 
 
 
 
108
  return model
109
 
110
  #Async wrappers
@@ -113,9 +183,10 @@ def get_embeddings() -> HuggingFaceEmbeddings:
113
 
114
  async def embed_texts(
115
  texts: list[str],
116
- batch_size: int = None
 
117
  ) -> list[list[float]]:
118
- model = get_embeddings()
119
  bs = batch_size or settings.embedding_batch_size
120
  loop = asyncio.get_event_loop()
121
 
@@ -132,8 +203,8 @@ async def embed_texts(
132
  logger.debug(f"Embedded batch {i}–{i + len(batch)} ({len(batch)} docs)")
133
  return all_embeddings
134
 
135
- async def embed_query(text: str) -> list[float]:
136
- model = get_embeddings()
137
  loop = asyncio.get_event_loop()
138
  vec = await loop.run_in_executor(
139
  _executor,
 
1
  """
2
+ Local + API embedding models with GPU/CPU aware defaults.
3
+ - BGE-large (1024-dim) for GPU
4
+ - BGE-small (384-dim) for fast local CPU
5
+ - OpenAI text-embedding-3-small (1536-dim) for fast CPU via API
6
+ - BGE requires a special query prefix for queries only
 
 
7
  """
8
 
9
  import asyncio
10
  import logging
11
  import warnings
12
  from concurrent.futures import ThreadPoolExecutor
13
+ from typing import Any
14
 
15
  import numpy as np
16
  from langchain_huggingface import HuggingFaceEmbeddings
17
+ from langchain_openai import OpenAIEmbeddings
18
+ from langchain_core.embeddings import Embeddings
19
  from pydantic.warnings import UnsupportedFieldAttributeWarning
20
 
21
  from .config import get_settings
 
31
 
32
  _executor = ThreadPoolExecutor(max_workers=2)
33
 
34
+ EMBEDDING_MODES: tuple[str, ...] = ("bge-large", "bge-small", "openai-small")
35
+ _EMBEDDING_CACHE: dict[str, Embeddings] = {}
36
+
37
+
38
  def _resolve_device() -> str:
39
  device = (settings.embedding_device or "auto").lower()
40
  if device == "auto":
 
59
  return device
60
 
61
 
62
+ def get_default_embedding_mode() -> str:
63
+ return "bge-large" if _resolve_device() == "cuda" else "openai-small"
64
+
65
+
66
+ def normalize_embedding_mode(mode: str | None) -> str:
67
+ if mode is None or mode == "auto":
68
+ return get_default_embedding_mode()
69
+ if mode not in EMBEDDING_MODES:
70
+ raise ValueError(f"Unknown embedding mode: {mode}")
71
+ return mode
72
+
73
+
74
+ def infer_embedding_mode_from_dim(dim: int) -> str | None:
75
+ dim_map = {
76
+ int(settings.embedding_dimensions): "bge-large",
77
+ int(settings.embedding_dimensions_cpu): "bge-small",
78
+ int(settings.embedding_dimensions_openai): "openai-small",
79
+ }
80
+ return dim_map.get(int(dim))
81
+
82
+
83
+ def _embedding_spec(mode: str) -> dict[str, Any]:
84
+ if mode == "bge-large":
85
+ return {
86
+ "provider": "local",
87
+ "model_name": settings.embedding_model,
88
+ "dimensions": settings.embedding_dimensions,
89
+ "device": _resolve_device(),
90
+ }
91
+ if mode == "bge-small":
92
+ return {
93
+ "provider": "local",
94
+ "model_name": settings.embedding_model_cpu,
95
+ "dimensions": settings.embedding_dimensions_cpu,
96
+ "device": _resolve_device(),
97
+ }
98
+ return {
99
+ "provider": "openai",
100
+ "model_name": settings.embedding_model_openai,
101
+ "dimensions": settings.embedding_dimensions_openai,
102
+ "device": "api",
103
+ }
104
 
 
105
 
106
+ def get_embedding_info(mode: str | None = None) -> dict[str, str | int]:
107
+ resolved = normalize_embedding_mode(mode)
108
+ spec = _embedding_spec(resolved)
109
+ return {
110
+ "mode": resolved,
111
+ "provider": spec["provider"],
112
+ "model_name": spec["model_name"],
113
+ "dimensions": int(spec["dimensions"]),
114
+ "device": spec["device"],
115
+ }
116
 
117
+
118
+ def get_embeddings_runtime_info() -> dict[str, Any]:
119
+ default_mode = get_default_embedding_mode()
120
+ options = []
121
+ for mode in EMBEDDING_MODES:
122
+ info = get_embedding_info(mode)
123
+ options.append({
124
+ "id": mode,
125
+ "model_name": info["model_name"],
126
+ "dimensions": info["dimensions"],
127
+ "provider": info["provider"],
128
+ "recommended": mode == default_mode,
129
+ })
130
  return {
131
+ "default_mode": default_mode,
132
+ "device": _resolve_device(),
133
+ "options": options,
134
  }
135
 
136
 
137
+ def get_embeddings(mode: str | None = None) -> Embeddings:
 
138
  """
139
+ Singleton embedding model per mode.
140
 
141
  encode_kwargs:
142
  normalize_embeddings=True -> required for cosine similarity to work correctly
 
145
  BGE was finetuned with an instruction-like query prefix.
146
  We pass that prefix for query encoding only; documents remain unchanged.
147
  """
148
+ resolved = normalize_embedding_mode(mode)
149
+ cached = _EMBEDDING_CACHE.get(resolved)
150
+ if cached is not None:
151
+ return cached
152
+
153
+ spec = _embedding_spec(resolved)
154
+ logger.info("Loading embedding model: %s on %s", spec["model_name"], spec["device"])
155
+ if spec["provider"] == "openai":
156
+ model = OpenAIEmbeddings(
157
+ model=spec["model_name"],
158
+ dimensions=int(spec["dimensions"]),
159
+ openai_api_key=settings.openai_api_key,
160
+ )
161
+ else:
162
+ model = HuggingFaceEmbeddings(
163
+ model_name=spec["model_name"],
164
+ model_kwargs={
165
+ "device": spec["device"],
166
+ },
167
+ encode_kwargs={
168
+ "normalize_embeddings": settings.embedding_normalize,
169
+ "batch_size": settings.embedding_batch_size,
170
+ },
171
+ query_encode_kwargs={
172
+ "prompt": "Represent this sentence for searching relevant passages: ",
173
+ },
174
+ )
175
+
176
+ logger.info("Embedding model loaded. Output dim=%s", spec["dimensions"])
177
+ _EMBEDDING_CACHE[resolved] = model
178
  return model
179
 
180
  #Async wrappers
 
183
 
184
  async def embed_texts(
185
  texts: list[str],
186
+ batch_size: int = None,
187
+ embedding_mode: str | None = None,
188
  ) -> list[list[float]]:
189
+ model = get_embeddings(embedding_mode)
190
  bs = batch_size or settings.embedding_batch_size
191
  loop = asyncio.get_event_loop()
192
 
 
203
  logger.debug(f"Embedded batch {i}–{i + len(batch)} ({len(batch)} docs)")
204
  return all_embeddings
205
 
206
+ async def embed_query(text: str, embedding_mode: str | None = None) -> list[float]:
207
+ model = get_embeddings(embedding_mode)
208
  loop = asyncio.get_event_loop()
209
  vec = await loop.run_in_executor(
210
  _executor,
rag_system/models.py CHANGED
@@ -1,5 +1,5 @@
1
  from pydantic import BaseModel, Field, field_validator
2
- from typing import Optional, Any
3
  from enum import Enum
4
  import uuid
5
 
@@ -9,6 +9,8 @@ class RetrievalMode(str, Enum):
9
  HYBRID = "hybrid"
10
  MMR = "mmr"
11
 
 
 
12
  # Ingestion
13
 
14
  class IngestRequest(BaseModel):
@@ -16,6 +18,7 @@ class IngestRequest(BaseModel):
16
  metadatas: Optional[list[dict[str,Any]]] = None
17
  collection_name: str = Field(default="default",pattern=r"^[a-z0-9_-]+$")
18
  force_reindex: bool = False
 
19
 
20
  @field_validator("texts")
21
  @classmethod
@@ -42,6 +45,7 @@ class QueryRequest(BaseModel):
42
  session_id: str = Field(default_factory=lambda: str(uuid.uuid4()))
43
  collection_name: str = Field(default="default")
44
  retrieval_mode: RetrievalMode = RetrievalMode.HYBRID
 
45
  top_k: Optional[int] = None
46
  doc_collections: Optional[list[str]] = None # per-doc sub-collections; None = legacy single-collection mode
47
  history: list[ChatMessage] = Field(default_factory=list)
 
1
  from pydantic import BaseModel, Field, field_validator
2
+ from typing import Optional, Any, Literal
3
  from enum import Enum
4
  import uuid
5
 
 
9
  HYBRID = "hybrid"
10
  MMR = "mmr"
11
 
12
+ EmbeddingMode = Literal["bge-large", "bge-small", "openai-small", "auto"]
13
+
14
  # Ingestion
15
 
16
  class IngestRequest(BaseModel):
 
18
  metadatas: Optional[list[dict[str,Any]]] = None
19
  collection_name: str = Field(default="default",pattern=r"^[a-z0-9_-]+$")
20
  force_reindex: bool = False
21
+ embedding_mode: Optional[EmbeddingMode] = None
22
 
23
  @field_validator("texts")
24
  @classmethod
 
45
  session_id: str = Field(default_factory=lambda: str(uuid.uuid4()))
46
  collection_name: str = Field(default="default")
47
  retrieval_mode: RetrievalMode = RetrievalMode.HYBRID
48
+ embedding_mode: Optional[EmbeddingMode] = None
49
  top_k: Optional[int] = None
50
  doc_collections: Optional[list[str]] = None # per-doc sub-collections; None = legacy single-collection mode
51
  history: list[ChatMessage] = Field(default_factory=list)
rag_system/query_engine.py CHANGED
@@ -20,6 +20,7 @@ from .config import get_settings
20
  from .prompt import SYSTEM_PROMPT, QUERY_REWRITE_PROMPT, MULTI_DOC_SYSTEM_PROMPT
21
  from .models import QueryRequest, QueryResponse, SourceDocument
22
  from .retriever import retrieve, detect_query_scope, multi_collection_retrieve
 
23
  from .memory import resolve_standalone_question,trim_history_to_budget, build_lc_messages
24
  from .guardrails import check_query, check_context, redact_pii
25
  from .cache import get_exact,set_exact,get_semantic,set_semantic
@@ -167,6 +168,9 @@ async def query(
167
  ) -> QueryResponse:
168
  start = time.monotonic()
169
 
 
 
 
170
  # 1. Input guardrail
171
  guard = check_query(request.query)
172
  if not guard.allowed:
@@ -179,7 +183,7 @@ async def query(
179
 
180
  # 2. Exact cache check
181
  if settings.cache_enabled:
182
- cached = get_exact(request.query, request.collection_name, request.retrieval_mode)
183
  if cached:
184
  logger.info(f"Exact cache hit for query: '{request.query}'")
185
  cached["cached"] = True
@@ -187,9 +191,9 @@ async def query(
187
  return QueryResponse(**cached)
188
 
189
  # 3. Embed query for semantic cache + later retrieval
190
- query_vec = await embed_query(request.query)
191
  if settings.cache_enabled:
192
- semantic_hit = get_semantic(query_vec)
193
  if semantic_hit:
194
  logger.info(f"Semantic cache hit for query: '{request.query}'")
195
  semantic_hit["cached"] = True
@@ -211,7 +215,6 @@ async def query(
211
  retrieval_query = await rewrite_query(standalone)
212
 
213
  # 6. Retrieve — multi-doc aware
214
- collections = request.doc_collections or [request.collection_name]
215
  if len(collections) > 1:
216
  scoped = detect_query_scope(retrieval_query, collections)
217
  k_per = max(3, (request.top_k or settings.top_k_rerank) // len(scoped))
@@ -293,8 +296,8 @@ async def query(
293
  # 9. Cache the result
294
  if settings.cache_enabled:
295
  result_dict = result.model_dump()
296
- set_exact(request.query, request.collection_name, request.retrieval_mode, result_dict)
297
- set_semantic(query_vec,request.query, result_dict)
298
 
299
  return result
300
 
@@ -332,11 +335,14 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
332
 
333
  start = time.monotonic()
334
  mode_val = request.retrieval_mode.value if hasattr(request.retrieval_mode, "value") else str(request.retrieval_mode)
 
 
335
 
336
  yield emit("pipeline_start", "in_progress", {
337
  "query": request.query,
338
  "collection": request.collection_name,
339
  "mode": mode_val,
 
340
  })
341
 
342
  try:
@@ -356,7 +362,7 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
356
  # --- Cache check ---
357
  query_vec = None
358
  if settings.cache_enabled:
359
- cached = get_exact(request.query, request.collection_name, request.retrieval_mode)
360
  if cached:
361
  cached["cached"] = True
362
  cached["latency_ms"] = round((time.monotonic() - start) * 1000, 2)
@@ -365,8 +371,8 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
365
  yield "data: [DONE]\n\n"
366
  return
367
 
368
- query_vec = await embed_query(request.query)
369
- semantic_hit = get_semantic(query_vec)
370
  if semantic_hit:
371
  semantic_hit["cached"] = True
372
  semantic_hit["latency_ms"] = round((time.monotonic() - start) * 1000, 2)
@@ -398,7 +404,6 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
398
  })
399
 
400
  # --- Document routing (multi-doc) ---
401
- collections = request.doc_collections or [request.collection_name]
402
  if len(collections) > 1:
403
  scoped = detect_query_scope(retrieval_query, collections)
404
  is_multi = len(scoped) > 1
@@ -515,9 +520,9 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
515
  "eval_scores": None,
516
  }
517
  if query_vec is None:
518
- query_vec = await embed_query(request.query)
519
- set_exact(request.query, request.collection_name, request.retrieval_mode, result_dict)
520
- set_semantic(query_vec, request.query, result_dict)
521
  except Exception:
522
  logger.warning("Cache write failed (non-fatal)", exc_info=True)
523
 
 
20
  from .prompt import SYSTEM_PROMPT, QUERY_REWRITE_PROMPT, MULTI_DOC_SYSTEM_PROMPT
21
  from .models import QueryRequest, QueryResponse, SourceDocument
22
  from .retriever import retrieve, detect_query_scope, multi_collection_retrieve
23
+ from .vector_store import resolve_embedding_mode_for_collections
24
  from .memory import resolve_standalone_question,trim_history_to_budget, build_lc_messages
25
  from .guardrails import check_query, check_context, redact_pii
26
  from .cache import get_exact,set_exact,get_semantic,set_semantic
 
168
  ) -> QueryResponse:
169
  start = time.monotonic()
170
 
171
+ collections = request.doc_collections or [request.collection_name]
172
+ embedding_mode = resolve_embedding_mode_for_collections(collections, request.embedding_mode)
173
+
174
  # 1. Input guardrail
175
  guard = check_query(request.query)
176
  if not guard.allowed:
 
183
 
184
  # 2. Exact cache check
185
  if settings.cache_enabled:
186
+ cached = get_exact(request.query, request.collection_name, request.retrieval_mode, embedding_mode)
187
  if cached:
188
  logger.info(f"Exact cache hit for query: '{request.query}'")
189
  cached["cached"] = True
 
191
  return QueryResponse(**cached)
192
 
193
  # 3. Embed query for semantic cache + later retrieval
194
+ query_vec = await embed_query(request.query, embedding_mode)
195
  if settings.cache_enabled:
196
+ semantic_hit = get_semantic(query_vec, embedding_mode)
197
  if semantic_hit:
198
  logger.info(f"Semantic cache hit for query: '{request.query}'")
199
  semantic_hit["cached"] = True
 
215
  retrieval_query = await rewrite_query(standalone)
216
 
217
  # 6. Retrieve — multi-doc aware
 
218
  if len(collections) > 1:
219
  scoped = detect_query_scope(retrieval_query, collections)
220
  k_per = max(3, (request.top_k or settings.top_k_rerank) // len(scoped))
 
296
  # 9. Cache the result
297
  if settings.cache_enabled:
298
  result_dict = result.model_dump()
299
+ set_exact(request.query, request.collection_name, request.retrieval_mode, embedding_mode, result_dict)
300
+ set_semantic(query_vec, request.query, result_dict, embedding_mode)
301
 
302
  return result
303
 
 
335
 
336
  start = time.monotonic()
337
  mode_val = request.retrieval_mode.value if hasattr(request.retrieval_mode, "value") else str(request.retrieval_mode)
338
+ collections = request.doc_collections or [request.collection_name]
339
+ embedding_mode = resolve_embedding_mode_for_collections(collections, request.embedding_mode)
340
 
341
  yield emit("pipeline_start", "in_progress", {
342
  "query": request.query,
343
  "collection": request.collection_name,
344
  "mode": mode_val,
345
+ "embedding_mode": embedding_mode,
346
  })
347
 
348
  try:
 
362
  # --- Cache check ---
363
  query_vec = None
364
  if settings.cache_enabled:
365
+ cached = get_exact(request.query, request.collection_name, request.retrieval_mode, embedding_mode)
366
  if cached:
367
  cached["cached"] = True
368
  cached["latency_ms"] = round((time.monotonic() - start) * 1000, 2)
 
371
  yield "data: [DONE]\n\n"
372
  return
373
 
374
+ query_vec = await embed_query(request.query, embedding_mode)
375
+ semantic_hit = get_semantic(query_vec, embedding_mode)
376
  if semantic_hit:
377
  semantic_hit["cached"] = True
378
  semantic_hit["latency_ms"] = round((time.monotonic() - start) * 1000, 2)
 
404
  })
405
 
406
  # --- Document routing (multi-doc) ---
 
407
  if len(collections) > 1:
408
  scoped = detect_query_scope(retrieval_query, collections)
409
  is_multi = len(scoped) > 1
 
520
  "eval_scores": None,
521
  }
522
  if query_vec is None:
523
+ query_vec = await embed_query(request.query, embedding_mode)
524
+ set_exact(request.query, request.collection_name, request.retrieval_mode, embedding_mode, result_dict)
525
+ set_semantic(query_vec, request.query, result_dict, embedding_mode)
526
  except Exception:
527
  logger.warning("Cache write failed (non-fatal)", exc_info=True)
528
 
rag_system/vector_store.py CHANGED
@@ -1,4 +1,5 @@
1
  #faiss index management
 
2
  import logging
3
  import time
4
  import os
@@ -10,7 +11,13 @@ from langchain_core.documents import Document
10
  from langchain_community.vectorstores import FAISS
11
 
12
  from .config import get_settings
13
- from .embeddings import get_embeddings, get_embedding_info
 
 
 
 
 
 
14
 
15
  logger = logging.getLogger(__name__)
16
  settings = get_settings()
@@ -18,6 +25,76 @@ settings = get_settings()
18
  _stores: dict[str, FAISS] = {}
19
  # Tracks last-used timestamp per collection (epoch seconds) for TTL-based cleanup
20
  _last_used: dict[str, float] = {}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
21
 
22
  def _index_path(collection: str) -> str:
23
  return str(Path(settings.faiss_index_path)/ collection)
@@ -33,7 +110,8 @@ def load_or_create_store(collection: str = "default") -> FAISS:
33
  return _stores[collection]
34
 
35
  path = _index_path(collection)
36
- embeddings = get_embeddings()
 
37
 
38
  if Path(path).exists():
39
  logger.info(f"Loading FAISS index from {path}")
@@ -42,7 +120,19 @@ def load_or_create_store(collection: str = "default") -> FAISS:
42
  embeddings,
43
  allow_dangerous_deserialization=True,
44
  )
45
- expected_dim = int(get_embedding_info()["dimensions"])
 
 
 
 
 
 
 
 
 
 
 
 
46
  if store.index.d != expected_dim:
47
  logger.error(
48
  "Embedding dim mismatch for collection '%s': index dim=%s, expected=%s. "
@@ -54,6 +144,8 @@ def load_or_create_store(collection: str = "default") -> FAISS:
54
  _stores[collection] = None
55
  else:
56
  _stores[collection] = store
 
 
57
  else:
58
  logger.warning(f"No index at {path}. Will create on first Ingest.")
59
  _stores[collection] = None
@@ -65,7 +157,8 @@ def load_or_create_store(collection: str = "default") -> FAISS:
65
  def add_documents(
66
  docs: list[Document],
67
  collection: str = "default",
68
- force_reindex: bool = False
 
69
  ) -> FAISS:
70
  """
71
  Adding docs to a FAISS collection.
@@ -73,7 +166,15 @@ def add_documents(
73
  - Persists to disk after every write
74
  """
75
 
76
- embeddings = get_embeddings()
 
 
 
 
 
 
 
 
77
  path = _index_path(collection)
78
 
79
  existing = None if force_reindex else _stores.get(collection)
@@ -93,6 +194,8 @@ def add_documents(
93
  store.save_local(path)
94
  _stores[collection] = store
95
  _last_used[collection] = time.time()
 
 
96
 
97
  # Prebuild BM25 index on ingest
98
  from .retriever import _bm25_cache, _get_bm25
@@ -137,6 +240,9 @@ def get_collection_stats(collection: str) -> dict:
137
  store = load_or_create_store(collection)
138
  path = _index_path(collection)
139
 
 
 
 
140
  chunk_count = 0
141
  if store is not None and hasattr(store, "index"):
142
  chunk_count = store.index.ntotal
@@ -155,6 +261,9 @@ def get_collection_stats(collection: str) -> dict:
155
  "size_mb": size_mb,
156
  "loaded": store is not None,
157
  "index_path": path,
 
 
 
158
  }
159
 
160
 
@@ -179,6 +288,8 @@ def delete_collection(collection: str) -> bool:
179
  path = _index_path(collection)
180
  if collection in _stores:
181
  del _stores[collection]
 
 
182
 
183
  # Local import to avoid circular dependency with retriever
184
  from .retriever import _bm25_cache
 
1
  #faiss index management
2
+ import json
3
  import logging
4
  import time
5
  import os
 
11
  from langchain_community.vectorstores import FAISS
12
 
13
  from .config import get_settings
14
+ from .embeddings import (
15
+ get_embeddings,
16
+ get_embedding_info,
17
+ get_default_embedding_mode,
18
+ infer_embedding_mode_from_dim,
19
+ normalize_embedding_mode,
20
+ )
21
 
22
  logger = logging.getLogger(__name__)
23
  settings = get_settings()
 
25
  _stores: dict[str, FAISS] = {}
26
  # Tracks last-used timestamp per collection (epoch seconds) for TTL-based cleanup
27
  _last_used: dict[str, float] = {}
28
+ _collection_embeddings: dict[str, str] = {}
29
+
30
+ _EMBEDDING_META_FILE = "embedding.json"
31
+
32
+
33
+ def _embedding_meta_path(collection: str) -> Path:
34
+ return Path(_index_path(collection)) / _EMBEDDING_META_FILE
35
+
36
+
37
+ def _read_embedding_meta(collection: str) -> dict | None:
38
+ meta_path = _embedding_meta_path(collection)
39
+ if not meta_path.exists():
40
+ return None
41
+ try:
42
+ return json.loads(meta_path.read_text(encoding="utf-8"))
43
+ except Exception:
44
+ logger.warning("Failed to read embedding metadata for '%s'", collection)
45
+ return None
46
+
47
+
48
+ def _write_embedding_meta(collection: str, info: dict) -> None:
49
+ meta_path = _embedding_meta_path(collection)
50
+ meta_path.parent.mkdir(parents=True, exist_ok=True)
51
+ meta_path.write_text(json.dumps(info, indent=2), encoding="utf-8")
52
+
53
+
54
+ def get_collection_embedding_mode(collection: str) -> Optional[str]:
55
+ if collection in _collection_embeddings:
56
+ return _collection_embeddings[collection]
57
+
58
+ meta = _read_embedding_meta(collection)
59
+ if meta and isinstance(meta, dict):
60
+ mode = meta.get("mode") or meta.get("embedding_mode")
61
+ if isinstance(mode, str):
62
+ _collection_embeddings[collection] = mode
63
+ return mode
64
+ return None
65
+
66
+
67
+ def resolve_embedding_mode_for_collections(
68
+ collections: list[str],
69
+ requested_mode: Optional[str] = None,
70
+ ) -> str:
71
+ requested = None
72
+ if requested_mode and requested_mode != "auto":
73
+ requested = normalize_embedding_mode(requested_mode)
74
+
75
+ modes = []
76
+ for coll in collections:
77
+ mode = get_collection_embedding_mode(coll)
78
+ if mode:
79
+ modes.append(mode)
80
+
81
+ if requested and modes and any(m != requested for m in modes):
82
+ logger.warning(
83
+ "Embedding mode mismatch (requested=%s, existing=%s). Using existing.",
84
+ requested,
85
+ sorted(set(modes)),
86
+ )
87
+ return modes[0]
88
+
89
+ if requested:
90
+ return requested
91
+
92
+ if modes:
93
+ if any(m != modes[0] for m in modes):
94
+ logger.warning("Multiple embedding modes across collections: %s", sorted(set(modes)))
95
+ return modes[0]
96
+
97
+ return get_default_embedding_mode()
98
 
99
  def _index_path(collection: str) -> str:
100
  return str(Path(settings.faiss_index_path)/ collection)
 
110
  return _stores[collection]
111
 
112
  path = _index_path(collection)
113
+ embedding_mode = resolve_embedding_mode_for_collections([collection])
114
+ embeddings = get_embeddings(embedding_mode)
115
 
116
  if Path(path).exists():
117
  logger.info(f"Loading FAISS index from {path}")
 
120
  embeddings,
121
  allow_dangerous_deserialization=True,
122
  )
123
+ expected_dim = int(get_embedding_info(embedding_mode)["dimensions"])
124
+ if store.index.d != expected_dim:
125
+ inferred_mode = infer_embedding_mode_from_dim(store.index.d)
126
+ if inferred_mode and inferred_mode != embedding_mode:
127
+ embeddings = get_embeddings(inferred_mode)
128
+ store = FAISS.load_local(
129
+ path,
130
+ embeddings,
131
+ allow_dangerous_deserialization=True,
132
+ )
133
+ embedding_mode = inferred_mode
134
+ expected_dim = int(get_embedding_info(embedding_mode)["dimensions"])
135
+
136
  if store.index.d != expected_dim:
137
  logger.error(
138
  "Embedding dim mismatch for collection '%s': index dim=%s, expected=%s. "
 
144
  _stores[collection] = None
145
  else:
146
  _stores[collection] = store
147
+ _collection_embeddings[collection] = embedding_mode
148
+ _write_embedding_meta(collection, get_embedding_info(embedding_mode))
149
  else:
150
  logger.warning(f"No index at {path}. Will create on first Ingest.")
151
  _stores[collection] = None
 
157
  def add_documents(
158
  docs: list[Document],
159
  collection: str = "default",
160
+ force_reindex: bool = False,
161
+ embedding_mode: Optional[str] = None,
162
  ) -> FAISS:
163
  """
164
  Adding docs to a FAISS collection.
 
166
  - Persists to disk after every write
167
  """
168
 
169
+ existing_mode = get_collection_embedding_mode(collection)
170
+ selected_mode = resolve_embedding_mode_for_collections([collection], embedding_mode)
171
+ if existing_mode and existing_mode != selected_mode and not force_reindex:
172
+ raise ValueError(
173
+ f"Embedding mode mismatch for '{collection}': existing={existing_mode}, requested={selected_mode}. "
174
+ "Use force_reindex to rebuild."
175
+ )
176
+
177
+ embeddings = get_embeddings(selected_mode)
178
  path = _index_path(collection)
179
 
180
  existing = None if force_reindex else _stores.get(collection)
 
194
  store.save_local(path)
195
  _stores[collection] = store
196
  _last_used[collection] = time.time()
197
+ _collection_embeddings[collection] = selected_mode
198
+ _write_embedding_meta(collection, get_embedding_info(selected_mode))
199
 
200
  # Prebuild BM25 index on ingest
201
  from .retriever import _bm25_cache, _get_bm25
 
240
  store = load_or_create_store(collection)
241
  path = _index_path(collection)
242
 
243
+ embedding_mode = get_collection_embedding_mode(collection)
244
+ embedding_info = get_embedding_info(embedding_mode) if embedding_mode else None
245
+
246
  chunk_count = 0
247
  if store is not None and hasattr(store, "index"):
248
  chunk_count = store.index.ntotal
 
261
  "size_mb": size_mb,
262
  "loaded": store is not None,
263
  "index_path": path,
264
+ "embedding_mode": embedding_mode,
265
+ "embedding_dimensions": embedding_info["dimensions"] if embedding_info else None,
266
+ "embedding_provider": embedding_info["provider"] if embedding_info else None,
267
  }
268
 
269
 
 
288
  path = _index_path(collection)
289
  if collection in _stores:
290
  del _stores[collection]
291
+ if collection in _collection_embeddings:
292
+ del _collection_embeddings[collection]
293
 
294
  # Local import to avoid circular dependency with retriever
295
  from .retriever import _bm25_cache