Spaces:
Sleeping
Sleeping
| import json | |
| import logging | |
| import os | |
| from typing import Any, Dict, List | |
| from langchain_community.vectorstores import Chroma | |
| from langchain_openai import OpenAIEmbeddings | |
| from dotenv import load_dotenv | |
| logger = logging.getLogger(__name__) | |
| class USTaxKB: | |
| """ | |
| RAG component for US Tax Knowledge Base using ChromaDB. | |
| Loads data from a JSON file containing tax QA pairs. | |
| """ | |
| def __init__( | |
| self, | |
| json_path: str = "src/data/us_tax_qa_kb_1000.json", | |
| persist_directory: str = "src/data/.chroma_tax_kb", | |
| collection_name: str = "us_tax_kb", | |
| ): | |
| load_dotenv() | |
| self.json_path = json_path | |
| self.persist_directory = persist_directory | |
| self.embeddings = OpenAIEmbeddings(model="text-embedding-3-small") | |
| os.makedirs(self.persist_directory, exist_ok=True) | |
| self.vector_store = Chroma( | |
| collection_name=collection_name, | |
| embedding_function=self.embeddings, | |
| persist_directory=self.persist_directory, | |
| ) | |
| def _load_json_data(self) -> List[Dict[str, Any]]: | |
| if not os.path.exists(self.json_path): | |
| logger.error(f"JSON file not found: {self.json_path}") | |
| return [] | |
| try: | |
| with open(self.json_path, "r", encoding="utf-8") as f: | |
| return json.load(f) | |
| except Exception as e: | |
| logger.exception(f"Failed to load JSON data from {self.json_path}: {e}") | |
| return [] | |
| def ingest(self) -> None: | |
| """ | |
| Ingests data from the JSON file into ChromaDB if not already present. | |
| """ | |
| data = self._load_json_data() | |
| if not data: | |
| return | |
| texts: List[str] = [] | |
| metadatas: List[Dict[str, Any]] = [] | |
| ids: List[str] = [] | |
| for item in data: | |
| doc_id = item.get("id") | |
| if not doc_id: | |
| continue | |
| # Combine question and answer for embedding | |
| text_content = f"Question: {item.get('question')}\nAnswer: {item.get('answer')}" | |
| texts.append(text_content) | |
| metadatas.append({ | |
| "category": item.get("category", ""), | |
| "topic": item.get("topic", ""), | |
| "difficulty": item.get("difficulty", ""), | |
| "question": item.get("question", "") | |
| }) | |
| ids.append(doc_id) | |
| # Check for existing IDs to avoid duplicates | |
| try: | |
| existing = set(self.vector_store.get(ids=ids).get("ids", [])) | |
| except Exception: | |
| existing = set() | |
| to_add_texts = [] | |
| to_add_metas = [] | |
| to_add_ids = [] | |
| for t, m, i in zip(texts, metadatas, ids): | |
| if i not in existing: | |
| to_add_texts.append(t) | |
| to_add_metas.append(m) | |
| to_add_ids.append(i) | |
| if to_add_texts: | |
| self.vector_store.add_texts( | |
| texts=to_add_texts, | |
| metadatas=to_add_metas, | |
| ids=to_add_ids | |
| ) | |
| logger.info(f"Ingested {len(to_add_texts)} new tax QA pairs into ChromaDB.") | |
| else: | |
| logger.info("No new tax QA pairs to ingest.") | |
| def search(self, query: str, k: int = 3) -> List[Dict[str, Any]]: | |
| """ | |
| Performs a vector search on ChromaDB. | |
| """ | |
| # Ensure data is ingested before search (optional, or could be called manually) | |
| # self.ingest() | |
| docs = self.vector_store.similarity_search(query, k=k) | |
| results = [] | |
| for doc in docs: | |
| results.append({ | |
| "content": doc.page_content, | |
| "metadata": doc.metadata | |
| }) | |
| return results | |