Add metadata handling functions and improve document retrieval logic in chroma_rag.py
Browse files- 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=
|
| 419 |
-
url=
|
| 420 |
-
source=
|
| 421 |
-
retrieve_doc=
|
| 422 |
-
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),
|