Fade0510's picture
Fix logging and enable semantic cache, stock qoute cache
357c48c
Raw
History Blame Contribute Delete
3.75 kB
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