omarsol commited on
Commit
760cef5
·
1 Parent(s): 2e97e50

Add metadata handling functions and improve document retrieval logic in chroma_rag.py

Browse files
Files changed (1) hide show
  1. scripts/chroma_rag.py +65 -5
scripts/chroma_rag.py CHANGED
@@ -6,6 +6,7 @@ import pickle
6
  from dataclasses import asdict, dataclass
7
  from pathlib import Path
8
  from typing import Any, Iterable
 
9
 
10
  import chromadb
11
  import cohere
@@ -226,6 +227,34 @@ def get_full_doc_content(full_doc: Any) -> str:
226
  raise TypeError("Unsupported full document type in document dictionary.")
227
 
228
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
229
  def _cohere_embeddings_list(response: Any) -> list[list[float]]:
230
  embeddings = getattr(response, "embeddings", None)
231
  if embeddings is None:
@@ -411,15 +440,46 @@ class LocalChromaRetriever:
411
  else:
412
  content = chunk_text
413
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
414
  dense_hits.append(
415
  SearchResult(
416
  chunk_id=str(chunk_id),
417
  doc_id=doc_id,
418
- title=str(metadata["title"]),
419
- url=str(metadata["url"]),
420
- source=str(metadata["source"]),
421
- retrieve_doc=bool(metadata["retrieve_doc"]),
422
- tokens=int(metadata["tokens"]),
423
  score=_distance_to_score(distance),
424
  content=content,
425
  chunk_content=str(chunk_text),
 
6
  from dataclasses import asdict, dataclass
7
  from pathlib import Path
8
  from typing import Any, Iterable
9
+ from urllib.parse import unquote, urlparse
10
 
11
  import chromadb
12
  import cohere
 
227
  raise TypeError("Unsupported full document type in document dictionary.")
228
 
229
 
230
+ def _is_missing_metadata_value(value: Any) -> bool:
231
+ if value is None:
232
+ return True
233
+ if isinstance(value, float) and math.isnan(value):
234
+ return True
235
+ if isinstance(value, str):
236
+ normalized = value.strip()
237
+ return not normalized or normalized.lower() == "nan"
238
+ return False
239
+
240
+
241
+ def _string_metadata_value(*candidates: Any, default: str = "") -> str:
242
+ for candidate in candidates:
243
+ if _is_missing_metadata_value(candidate):
244
+ continue
245
+ return str(candidate).strip()
246
+ return default
247
+
248
+
249
+ def _title_from_url(url: str) -> str:
250
+ if not url:
251
+ return ""
252
+ slug = urlparse(url).path.rstrip("/").split("/")[-1]
253
+ if not slug:
254
+ return ""
255
+ return unquote(slug).replace("-", " ").replace("_", " ").strip()
256
+
257
+
258
  def _cohere_embeddings_list(response: Any) -> list[list[float]]:
259
  embeddings = getattr(response, "embeddings", None)
260
  if embeddings is None:
 
440
  else:
441
  content = chunk_text
442
 
443
+ full_doc_name = full_doc.get("name") if isinstance(full_doc, dict) else None
444
+ url = _string_metadata_value(
445
+ metadata.get("url"),
446
+ full_doc.get("url") if isinstance(full_doc, dict) else None,
447
+ )
448
+ title = _string_metadata_value(
449
+ metadata.get("title"),
450
+ full_doc_name,
451
+ _title_from_url(url),
452
+ doc_id,
453
+ )
454
+ source = _string_metadata_value(
455
+ metadata.get("source"),
456
+ full_doc.get("source") if isinstance(full_doc, dict) else None,
457
+ default="unknown",
458
+ )
459
+ retrieve_doc = bool(
460
+ metadata.get("retrieve_doc")
461
+ if "retrieve_doc" in metadata
462
+ else (
463
+ full_doc.get("retrieve_doc") if isinstance(full_doc, dict) else False
464
+ )
465
+ )
466
+ tokens_value = metadata.get("tokens")
467
+ if _is_missing_metadata_value(tokens_value) and isinstance(full_doc, dict):
468
+ tokens_value = full_doc.get("tokens")
469
+ try:
470
+ tokens = int(tokens_value)
471
+ except (TypeError, ValueError):
472
+ tokens = 0
473
+
474
  dense_hits.append(
475
  SearchResult(
476
  chunk_id=str(chunk_id),
477
  doc_id=doc_id,
478
+ title=title,
479
+ url=url,
480
+ source=source,
481
+ retrieve_doc=retrieve_doc,
482
+ tokens=tokens,
483
  score=_distance_to_score(distance),
484
  content=content,
485
  chunk_content=str(chunk_text),