Cognitive-rag / 04_vectoredb / pinecone_client.py
pinecone_client.py
Raw
"""
Pinecone Vector Database Client for CognitiveRAG.

Handles storage and retrieval of embeddings using Pinecone.
"""

from __future__ import annotations

import logging
from typing import Dict, List, Optional, Any

from pinecone import Pinecone, ServerlessSpec

# handle both package import and direct import 
try:
    from . import config
except ImportError:
    import config

# configure logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)


class PineconeClient:
    """Client for interacting with Pinecone vector database."""

    def __init__(self, api_key: str = None, index_name: str = None):
        """
        Initialize the Pinecone client.

        args:
            api_key: pinecone API key. (check config for api key or ask me)
            index_name: Name of the pinecone index. (currently rag-hybrid-index-v1) 
        """
        self.api_key = api_key or config.PINECONE_API_KEY
        self.index_name = index_name or config.PINECONE_INDEX_NAME
        
        # initialize pinecone
        self.pc = Pinecone(api_key=self.api_key)
        self._index = None

    def _ensure_index_exists(self) -> None:
        """create the index if it doesn't exist."""
        existing_indexes = [idx.name for idx in self.pc.list_indexes()]
        
        if self.index_name not in existing_indexes:
            logger.info(f"Creating new Pinecone index: {self.index_name}")
            self.pc.create_index(
                name=self.index_name,
                dimension=config.EMBEDDING_DIMENSION,
                metric=config.SIMILARITY_METRIC,    #cosine for now but can be changed in config
                spec=ServerlessSpec(
                    cloud=config.PINECONE_CLOUD,    #set to aws (for aws track maybe?)
                    region=config.PINECONE_REGION    #set to us-east-1 
                )
            )
            logger.info(f"Index '{self.index_name}' created successfully")
        else:
            logger.info(f"Index '{self.index_name}' already exists")

    def get_index(self):
        """get or create the pinecone index."""
        if self._index is None:
            self._ensure_index_exists()
            self._index = self.pc.Index(self.index_name)
        return self._index

    def store_embeddings(
        self, 
        chunks: List[Dict[str, Any]], 
        batch_size: int = 100,
        namespace: str = ""
    ) -> Dict[str, Any]:
        """
        store chunk embeddings in pinecone.use this if adding a few documents, otherwise use batch_upsert_from_file (can crash laptop)

        args:
            chunks: list of chunk dictionaries containing:
                - chunk_id: unique identifier for the chunk
                - embedding: vector embedding (list of floats)
                - text: original text content
                - doc_id: document identifier
                - section: section name (optional)
                - page: page number (optional)
            batch_size: number of vectors to upsert per batch
            namespace: pinecone namespace for organizing vectors

        returns:
            dictionary with storage results:
                - stored_count: number of vectors stored
                - errors: list of any errors encountered
        """
        index = self.get_index()
        stored_count = 0
        errors = []

        # Process in batches for efficiency
        for i in range(0, len(chunks), batch_size):
            batch = chunks[i:i + batch_size]
            vectors = []

            for chunk in batch:
                try:
                    # validate required fields
                    if "chunk_id" not in chunk:
                        errors.append(f"Missing chunk_id in chunk at index {i}")
                        continue
                    if "embedding" not in chunk:
                        errors.append(f"Missing embedding for chunk {chunk.get('chunk_id', 'unknown')}")
                        continue

                    # build metadata (pinecone stores this alongside the vector)
                    # Handle both 'text' and 'content' field names
                    text_content = chunk.get("text") or chunk.get("content", "")
                    page_val = chunk.get("page")
                    
                    # Extract related_chunks for relational retrieval
                    related = chunk.get("related_chunks", {})
                    
                    metadata = {
                        "text": text_content,
                        "doc_id": chunk.get("doc_id", ""),
                        "section": chunk.get("section", ""),
                        "page": page_val if page_val is not None else "",
                        # Store prev/next for relational chunking
                        "prev_chunk": related.get("prev") or "",
                        "next_chunk": related.get("next") or "",
                    }

                    vectors.append({
                        "id": chunk["chunk_id"],
                        "values": chunk["embedding"],
                        "metadata": metadata
                    })

                except Exception as e:
                    errors.append(f"Error processing chunk {chunk.get('chunk_id', 'unknown')}: {str(e)}")

            # upsert batch to pinecone 
            if vectors:
                try:
                    index.upsert(vectors=vectors, namespace=namespace)
                    stored_count += len(vectors)
                    logger.info(f"Stored batch of {len(vectors)} vectors (total: {stored_count})")
                except Exception as e:
                    errors.append(f"Error upserting batch: {str(e)}")

        return {
            "stored_count": stored_count,
            "total_chunks": len(chunks),
            "errors": errors if errors else None
        }

    def retrieve_similar(
        self, 
        query_embedding: List[float], 
        top_k: int = None,
        namespace: str = "",
        filter: Dict = None
    ) -> Dict[str, Any]:
        """
        retrieve the most similar chunks for a query embedding.

        args:
            query_embedding: the query vector (from ajwaad's embedding model)
            top_k: number of results to return. can be edited in config
            namespace: pinecone namespace to search in
            filter: optional metadata filter for the query from earlier

        Returns:
            dictionary with retrieved context:
                - query_embedding_length: length of the query vector
                - retrieved_chunks: list of matching chunks with:
                    - chunk_id: The chunk identifier
                    - text: The original text content
                    - score: Similarity score
                    - doc_id: Document identifier
                    - section: Section name
                    - page: Page number 
        """
        index = self.get_index()
        top_k = top_k or config.TOP_K

        try:
            # query Pinecone for similar vectors
            results = index.query(
                vector=query_embedding,
                top_k=top_k,
                include_metadata=True,
                namespace=namespace,
                filter=filter
            )

            # format results
            retrieved_chunks = []
            for match in results.get("matches", []):
                metadata = match.get("metadata", {})
                retrieved_chunks.append({
                    "chunk_id": match["id"],
                    "text": metadata.get("text", ""),
                    "score": match["score"],
                    "doc_id": metadata.get("doc_id", ""),
                    "section": metadata.get("section", ""),
                    "page": metadata.get("page"),
                    # Include relational chunk IDs
                    "prev_chunk": metadata.get("prev_chunk", ""),
                    "next_chunk": metadata.get("next_chunk", ""),
                })

            return {
                "query_embedding_length": len(query_embedding),
                "retrieved_chunks": retrieved_chunks
            }

        except Exception as e:
            logger.error(f"Error retrieving similar chunks: {str(e)}")
            return {
                "query_embedding_length": len(query_embedding),
                "retrieved_chunks": [],
                "error": str(e)
            }

    def delete_vectors(
        self, 
        ids: List[str] = None, 
        delete_all: bool = False,
        namespace: str = ""
    ) -> Dict[str, Any]:
        """
        delete vectors from the index (i think it can be done manually on the pinecone website).

        args:
            ids: list of vector ids to delete
            delete_all: if true, delete all vectors in the namespace
            namespace: pinecone namespace

        returns:
            dictionary with deletion results
        """
        index = self.get_index()
        
        try:
            if delete_all:
                index.delete(delete_all=True, namespace=namespace)
                return {"deleted": "all", "namespace": namespace}
            elif ids:
                index.delete(ids=ids, namespace=namespace)
                return {"deleted_count": len(ids), "ids": ids}
            else:
                return {"error": "Must provide ids or set delete_all=True"}
        except Exception as e:
            logger.error(f"Error deleting vectors: {str(e)}")
            return {"error": str(e)}

    def batch_upsert_from_file(
        self,
        file_path: str,
        batch_size: int = 100,
        namespace: str = "",
        show_progress: bool = True
    ) -> Dict[str, Any]:
        """
        batch upsert embeddings from a JSONL file. (for mass uploads to pinecone server)
        
        use this for uploading the initial corpus. each line in the file
        should be a JSON object with chunk_id, embedding, text, doc_id, etc.

        Args:
        file_path: path to JSONL file containing embedded chunks
        batch_size: Number of vectors to upsert per batch (default: 100)
        namespace: Pinecone namespace
        show_progress: Whether to print progress updates
        Returns:
        Dictionary with upload results:
            stored_count: Total vectors stored
            total_lines: Total lines processed
            errors: Any errors encountered
            batches_processed: Number of batches
        """
        import json
        
        index = self.get_index()
        stored_count = 0
        total_lines = 0
        errors = []
        batches_processed = 0
        
        current_batch = []
        
        with open(file_path, 'r', encoding='utf-8') as f:
            for line_num, line in enumerate(f, 1):
                total_lines += 1
                
                try:
                    chunk = json.loads(line.strip())
                    
                    # validate required fields
                    if "chunk_id" not in chunk or "embedding" not in chunk:
                        errors.append(f"Line {line_num}: Missing required fields")
                        continue
                    
                    # build vector for Pinecone
                    # Handle both 'text' and 'content' field names
                    text_content = chunk.get("text") or chunk.get("content", "")
                    # Pinecone doesn't allow null values in metadata, use empty string
                    page_val = chunk.get("page")
                    
                    # Extract related_chunks for relational retrieval
                    related = chunk.get("related_chunks", {})
                    
                    vector = {
                        "id": chunk["chunk_id"],
                        "values": chunk["embedding"],
                        "metadata": {
                            "text": text_content,
                            "doc_id": chunk.get("doc_id", ""),
                            "section": chunk.get("section", ""),
                            "page": page_val if page_val is not None else "",
                            # Store prev/next for relational chunking
                            "prev_chunk": related.get("prev") or "",
                            "next_chunk": related.get("next") or "",
                        }
                    }
                    current_batch.append(vector)
                    
                    # upsert when batch is full
                    if len(current_batch) >= batch_size:
                        try:
                            index.upsert(vectors=current_batch, namespace=namespace)
                            stored_count += len(current_batch)
                            batches_processed += 1
                            
                            if show_progress:
                                logger.info(f"Batch {batches_processed}: Stored {stored_count} vectors")
                            
                            current_batch = []
                        except Exception as e:
                            errors.append(f"Batch {batches_processed + 1} error: {str(e)}")
                            current_batch = []
                            
                except json.JSONDecodeError as e:
                    errors.append(f"Line {line_num}: Invalid JSON - {str(e)}")
                except Exception as e:
                    errors.append(f"Line {line_num}: {str(e)}")
        
        # upsert remaining vectors
        if current_batch:
            try:
                index.upsert(vectors=current_batch, namespace=namespace)
                stored_count += len(current_batch)
                batches_processed += 1
                
                if show_progress:
                    logger.info(f"Final batch: Stored {stored_count} total vectors")
            except Exception as e:
                errors.append(f"Final batch error: {str(e)}")
        
        return {
            "stored_count": stored_count,
            "total_lines": total_lines,
            "batches_processed": batches_processed,
            "errors": errors if errors else None
        }

    def get_index_stats(self) -> Dict[str, Any]:
        """Get statistics about the index.""" #i included this for debugging but it could be useful... (use after batch upsert to confirm it works)
        index = self.get_index()
        try:
            stats = index.describe_index_stats()
            return {
                "dimension": stats.get("dimension"),
                "total_vector_count": stats.get("total_vector_count"),
                "namespaces": stats.get("namespaces", {})
            }
        except Exception as e:
            logger.error(f"Error getting index stats: {str(e)}")
            return {"error": str(e)}


class SparseClient:
    """Client for Pinecone sparse vectors using pinecone-sparse-english-v0."""

    def __init__(self, api_key: str = None, index_name: str = None):
        self.api_key = api_key or config.PINECONE_API_KEY
        self.index_name = index_name or config.SPARSE_INDEX_NAME
        self.pc = Pinecone(api_key=self.api_key)
        self._index = None

    def get_index(self):
        """Get the sparse index."""
        if self._index is None:
            self._index = self.pc.Index(self.index_name)
        return self._index

    def embed_sparse(self, texts: List[str]) -> List[Dict]:
        """Generate sparse embeddings using Pinecone inference."""
        embeddings = self.pc.inference.embed(
            model=config.SPARSE_EMBEDDING_MODEL,
            inputs=texts,
            parameters={"input_type": "passage"}
        )
        return [
            {"indices": e.to_dict()["sparse_indices"], "values": e.to_dict()["sparse_values"]}
            for e in embeddings.data
        ]

    def embed_query_sparse(self, query: str) -> Dict:
        """Generate sparse embedding for a query."""
        embeddings = self.pc.inference.embed(
            model=config.SPARSE_EMBEDDING_MODEL,
            inputs=[query],
            parameters={"input_type": "query"}
        )
        d = embeddings.data[0].to_dict()
        return {"indices": d["sparse_indices"], "values": d["sparse_values"]}

    def store_sparse(
        self,
        chunks: List[Dict[str, Any]],
        batch_size: int = 96,
        namespace: str = ""
    ) -> Dict[str, Any]:
        """Store sparse embeddings in the sparse index."""
        index = self.get_index()
        stored_count = 0
        errors = []

        for i in range(0, len(chunks), batch_size):
            batch_chunks = chunks[i:i + batch_size]
            
            # Embed this batch only (avoid 413 errors)
            texts = [c.get("text") or c.get("content", "") for c in batch_chunks]
            try:
                sparse_embeddings = self.embed_sparse(texts)
            except Exception as e:
                errors.append(f"Embed error batch {i//batch_size}: {str(e)}")
                continue
            
            vectors = []
            for chunk, sparse in zip(batch_chunks, sparse_embeddings):
                if "chunk_id" not in chunk:
                    errors.append(f"Missing chunk_id")
                    continue

                text_content = chunk.get("text") or chunk.get("content", "")
                page_val = chunk.get("page")
                related = chunk.get("related_chunks", {})

                vectors.append({
                    "id": chunk["chunk_id"],
                    "sparse_values": sparse,
                    "metadata": {
                        "text": text_content,
                        "doc_id": chunk.get("doc_id", ""),
                        "section": chunk.get("section", ""),
                        "page": page_val if page_val is not None else "",
                        "prev_chunk": related.get("prev") or "",
                        "next_chunk": related.get("next") or "",
                    }
                })

            if vectors:
                try:
                    index.upsert(vectors=vectors, namespace=namespace)
                    stored_count += len(vectors)
                    logger.info(f"Stored {stored_count} sparse vectors")
                except Exception as e:
                    errors.append(f"Error upserting batch: {str(e)}")

        return {
            "stored_count": stored_count,
            "total_chunks": len(chunks),
            "errors": errors if errors else None
        }

    def retrieve_sparse(
        self,
        query: str,
        top_k: int = None,
        namespace: str = ""
    ) -> Dict[str, Any]:
        """Retrieve using sparse embeddings."""
        index = self.get_index()
        top_k = top_k or config.TOP_K

        try:
            sparse_query = self.embed_query_sparse(query)
            results = index.query(
                sparse_vector=sparse_query,
                top_k=top_k,
                include_metadata=True,
                namespace=namespace
            )

            retrieved_chunks = []
            for match in results.get("matches", []):
                metadata = match.get("metadata", {})
                retrieved_chunks.append({
                    "chunk_id": match["id"],
                    "text": metadata.get("text", ""),
                    "score": match["score"],
                    "doc_id": metadata.get("doc_id", ""),
                    "section": metadata.get("section", ""),
                    "page": metadata.get("page"),
                    "prev_chunk": metadata.get("prev_chunk", ""),
                    "next_chunk": metadata.get("next_chunk", ""),
                })

            return {"retrieved_chunks": retrieved_chunks}

        except Exception as e:
            logger.error(f"Error retrieving sparse: {str(e)}")
            return {"retrieved_chunks": [], "error": str(e)}

    def get_index_stats(self) -> Dict[str, Any]:
        """Get statistics about the sparse index."""
        index = self.get_index()
        try:
            stats = index.describe_index_stats()
            return {
                "total_vector_count": stats.get("total_vector_count"),
                "namespaces": stats.get("namespaces", {})
            }
        except Exception as e:
            logger.error(f"Error getting sparse index stats: {str(e)}")
            return {"error": str(e)}


# convenience functions for direct usage
_client = None
_sparse_client = None

def get_client() -> PineconeClient:
    """Get or create a single Pinecone Client instance."""
    global _client
    if _client is None:
        _client = PineconeClient()
    return _client


def get_sparse_client() -> SparseClient:
    """Get or create a single SparseClient instance."""
    global _sparse_client
    if _sparse_client is None:
        _sparse_client = SparseClient()
    return _sparse_client


def store_embeddings(chunks: List[Dict[str, Any]], **kwargs) -> Dict[str, Any]:
    """Store embeddings using the default client."""
    return get_client().store_embeddings(chunks, **kwargs)


def retrieve_similar(query_embedding: List[float], **kwargs) -> Dict[str, Any]:
    """Retrieve similar chunks using the default client."""
    return get_client().retrieve_similar(query_embedding, **kwargs)


def batch_upsert_from_file(file_path: str, **kwargs) -> Dict[str, Any]:
    """Batch upsert embeddings from a JSONL file using the default client."""
    return get_client().batch_upsert_from_file(file_path, **kwargs)