agent-api / rag_context.py
github-actions[bot]
Sync GitHub snapshot to Hugging Face
9a1014e
Raw
History Blame Contribute Delete
6.37 kB
from __future__ import annotations
import os
import numpy as np
from dotenv import load_dotenv
from langchain_community.vectorstores import Chroma
from hf_text_embeddings import HFTextEmbeddings
DOCUMENT_PERSIST_DIR = "./chroma_db_docs"
DOCUMENT_COLLECTION = "document_context"
load_dotenv()
def _metadata_as_text(metadata: dict) -> str:
searchable_fields = (
"document_title",
"document_type",
"document_class",
"keywords",
"filename",
)
return " | ".join(
str(metadata[field]).strip()
for field in searchable_fields
if metadata.get(field)
)
def _combine_filters(first: dict | None, second: dict) -> dict:
if not first:
return second
return {"$and": [first, second]}
def _classify_document_filter(
vectorstore: Chroma,
embeddings: HFTextEmbeddings,
query: str,
base_filter: dict | None = None,
) -> tuple[dict | None, dict | None, float | None]:
get_kwargs = {"include": ["metadatas"]}
if base_filter:
get_kwargs["where"] = base_filter
stored = vectorstore.get(**get_kwargs)
candidates: dict[str, dict] = {}
for metadata in stored.get("metadatas") or []:
if not metadata:
continue
source = str(metadata.get("source") or "")
if source and _metadata_as_text(metadata):
candidates.setdefault(source, metadata)
if not candidates:
return base_filter, None, None
sources = list(candidates)
metadata_texts = [_metadata_as_text(candidates[source]) for source in sources]
query_vector = np.asarray(embeddings.embed_query(query), dtype=np.float32)
metadata_vectors = np.asarray(
embeddings.embed_documents(metadata_texts),
dtype=np.float32,
)
query_norm = np.linalg.norm(query_vector)
metadata_norms = np.linalg.norm(metadata_vectors, axis=1)
denominators = metadata_norms * query_norm
scores = np.divide(
metadata_vectors @ query_vector,
denominators,
out=np.zeros(len(metadata_vectors), dtype=np.float32),
where=denominators != 0,
)
best_index = int(np.argmax(scores))
best_source = sources[best_index]
best_metadata = candidates[best_source]
return (
_combine_filters(base_filter, {"source": {"$eq": best_source}}),
best_metadata,
float(scores[best_index]),
)
def retrieve_document_context(
query: str,
k: int = 4,
filter_metadata: dict | None = None,
) -> str:
try:
embeddings = HFTextEmbeddings()
vectorstore = Chroma(
collection_name=os.getenv("DOCUMENT_CHROMA_COLLECTION", DOCUMENT_COLLECTION),
embedding_function=embeddings,
persist_directory=os.getenv("DOCUMENT_CHROMA_DIR", DOCUMENT_PERSIST_DIR),
)
selected_filter, selected_metadata, metadata_score = _classify_document_filter(
vectorstore,
embeddings,
query,
filter_metadata,
)
if selected_metadata is not None:
print(
"Metadata classifier | "
f"score={metadata_score:.4f} | "
f"document={selected_metadata.get('document_title')} | "
f"type={selected_metadata.get('document_type')}",
flush=True,
)
# Evita el query HNSW de Chroma, que puede bloquearse en algunos
# contenedores. La coleccion ASTM es pequena, por lo que un ranking
# coseno directo sobre los embeddings persistidos es rapido y estable.
get_kwargs = {"include": ["documents", "metadatas", "embeddings"]}
if selected_filter:
get_kwargs["where"] = selected_filter
stored = vectorstore.get(**get_kwargs)
documents = stored.get("documents") or []
metadatas = stored.get("metadatas") or []
stored_embeddings = stored.get("embeddings")
if not documents or stored_embeddings is None:
print("Document Chroma returned no searchable chunks.", flush=True)
return ""
query_vector = np.asarray(embeddings.embed_query(query), dtype=np.float32)
chunk_vectors = np.asarray(stored_embeddings, dtype=np.float32)
query_norm = np.linalg.norm(query_vector)
chunk_norms = np.linalg.norm(chunk_vectors, axis=1)
denominators = chunk_norms * query_norm
similarities = np.divide(
chunk_vectors @ query_vector,
denominators,
out=np.zeros(len(chunk_vectors), dtype=np.float32),
where=denominators != 0,
)
top_indices = np.argsort(similarities)[::-1][: min(k, len(documents))]
blocks = []
for rank, index in enumerate(top_indices, start=1):
metadata = metadatas[index] if index < len(metadatas) else {}
print(
f"Document rank {rank} | similarity={similarities[index]:.4f} | "
f"document={metadata.get('document_title')}",
flush=True,
)
blocks.append(
f"[Documento {rank}] metadata={metadata}\n{documents[index]}"
)
return "\n\n".join(blocks)
except Exception as exc:
print(f"Document Chroma retrieval skipped: {exc}", flush=True)
return ""
def build_context_prompt(
user_text: str,
k: int = 4,
conversation_history: str = "",
) -> str:
document_context = retrieve_document_context(user_text, k=k)
context_parts = []
if document_context:
context_parts.append("Contexto normativo/documental recuperado de documentos/Chroma:\n" + document_context)
if not context_parts:
return user_text
context = "\n\n".join(context_parts)
return f"""Sos el agente del proyecto Alberti/Metalurgia.
Usa como fuente principal el contexto técnico y normativo recuperado desde los documentos indexados en Chroma.
Cuando respondas sobre datos del proyecto, prioriza los registros recuperados y menciona información relacionada si ayudan pero no divulgues ID o datos propios de los registros.
Si el contexto recuperado no alcanza para responder con precision, dilo claramente y no inventes datos exactos.
{context}
Historial reciente de esta conversacion:
{conversation_history or "Sin mensajes anteriores."}
Pregunta del usuario:
{user_text}"""