Cognitive-rag / 06_conversation_memory / memory_processor.py
memory_processor.py
Raw
from __future__ import annotations

import re
import uuid
from datetime import datetime
from typing import Optional, Dict, List
from pathlib import Path

# Load config
import importlib.util
_config_path = Path(__file__).parent / "config.py"
_spec = importlib.util.spec_from_file_location("memory_config", _config_path)
memory_config = importlib.util.module_from_spec(_spec)
_spec.loader.exec_module(memory_config)

# Import error handling
try:
    from .error_handling import safe_upsert, logger
except ImportError:
    from error_handling import safe_upsert, logger


# ============================================================================
# SESSION MANAGEMENT
# ============================================================================

def get_or_create_session_id() -> str:
    """
    Load session ID from temp file or create new one.
    Enables continuity across CLI restarts.
    """
    if memory_config.SESSION_FILE.exists():
        session_id = memory_config.SESSION_FILE.read_text().strip()
        if session_id:
            logger.info(f"Loaded existing session: {session_id[:8]}...")
            return session_id
    
    session_id = str(uuid.uuid4())
    memory_config.SESSION_FILE.write_text(session_id)
    logger.info(f"Created new session: {session_id[:8]}...")
    return session_id


def reset_session() -> str:
    """Clear session file and create new session."""
    if memory_config.SESSION_FILE.exists():
        memory_config.SESSION_FILE.unlink()
    return get_or_create_session_id()


# ============================================================================
# LLM OUTPUT PARSER
# ============================================================================

def parse_LLM_output(raw_output: str) -> Optional[Dict]:
    """
    Parse LLM's generation output format.
    
    Expected format:
        Query:
            "What is the patient's age?"
        
        Answer:
            The patient is 45 years old.
        
        Evidence:
            Patient is a 45-year-old male...
    
    Args:
        raw_output: Raw text output from generate_answer.py
    
    Returns:
        Dict with query, answer, evidence fields, or None if parsing fails
    """
    if not raw_output or not raw_output.strip():
        logger.warning("Empty output provided to parser")
        return None
    
    # Extract Query
    query_match = re.search(
        r'Query:\s*\n?\s*"?([^"]+)"?',
        raw_output,
        re.IGNORECASE
    )
    query = query_match.group(1).strip() if query_match else ""
    
    # Extract Answer
    answer_match = re.search(
        r'Answer:\s*\n?\s*(.+?)(?=\n\s*Evidence:|$)',
        raw_output,
        re.IGNORECASE | re.DOTALL
    )
    answer = answer_match.group(1).strip() if answer_match else ""
    
    # Extract Evidence (can be multiple)
    evidence_matches = re.findall(
        r'Evidence:\s*\n?\s*(.+?)(?=\n\s*Evidence:|$)',
        raw_output,
        re.IGNORECASE | re.DOTALL
    )
    evidence_list = [e.strip() for e in evidence_matches if e.strip()]
    
    if not query and not answer:
        logger.warning("Could not parse query or answer from output")
        return None
    
    return {
        "query": query,
        "answer": answer,
        "evidence": evidence_list,
        "full_text": raw_output.strip()
    }


# ============================================================================
# EMBEDDING HELPER (NVIDIA NIM)
# ============================================================================

_embedding_client = None


def get_embedding_model():
    """Load and cache the NVIDIA NIM embedding client."""
    global _embedding_client
    if _embedding_client is None:
        try:
            from langchain_nvidia_ai_endpoints import NVIDIAEmbeddings
        except ImportError as exc:
            raise ImportError(
                "Failed to import langchain_nvidia_ai_endpoints. "
                "Install with: pip install langchain-nvidia-ai-endpoints"
            ) from exc
        
        _embedding_client = NVIDIAEmbeddings(
            model=memory_config.EMBEDDING_MODEL,
            api_key=memory_config.NVIDIA_API_KEY,
            truncate="NONE"
        )
        logger.info(f"Loaded NVIDIA NIM embedding model: {memory_config.EMBEDDING_MODEL}")
    return _embedding_client


def embed_text(text: str) -> List[float]:
    """Embed text using the NVIDIA NIM model."""
    client = get_embedding_model()
    embedding = client.embed_query(text)
    return embedding


# ============================================================================
# MEMORY STORAGE
# ============================================================================

def build_memory_chunk(
    parsed_output: Dict,
    session_id: str,
    turn_number: int = 0
) -> Dict:
    """
    Build a memory chunk matching Nathan's document chunk schema.
    
    Metadata keys: text, doc_id, section, page
    """
    # Combine into single text for embedding
    text_content = f"Query: {parsed_output['query']}\n\n"
    text_content += f"Answer: {parsed_output['answer']}\n\n"
    for i, ev in enumerate(parsed_output.get('evidence', []), 1):
        text_content += f"Evidence {i}: {ev}\n"
    
    now = datetime.utcnow()
    chunk_id = f"mem_{session_id[:8]}_{uuid.uuid4().hex[:8]}"
    
    return {
        "chunk_id": chunk_id,
        "text": text_content.strip(),
        "doc_id": session_id,  # Session acts as document ID
        "section": "conversation_memory",
        "page": turn_number,  # Turn number acts as page
        "timestamp": now.isoformat() + "Z",
        "source": {
            "type": "conversation_memory",
            "session_id": session_id,
            "created_at": now.strftime("%Y-%m-%d %H:%M:%S"),
        }
    }


async def store_memory(
    raw_output: str,
    session_id: str,
    pinecone_index,
    turn_number: int = 0
) -> Optional[Dict]:
    """
    Parse LLM output, embed it, and store in Pinecone.
    
    Args:
        raw_output: Raw text from LLM generation
        session_id: Current session UUID
        pinecone_index: Pinecone index object
        turn_number: Conversation turn counter
    
    Returns:
        Stored chunk metadata if successful, None otherwise
    """
    # Step 1: Parse output
    parsed = parse_LLM_output(raw_output)
    if not parsed:
        logger.info("Skipping memory storage - could not parse output")
        return None
    
    # Step 2: Build chunk
    chunk = build_memory_chunk(parsed, session_id, turn_number)
    
    # Step 3: Embed
    try:
        embedding = embed_text(chunk["text"])
        chunk["embedding"] = embedding
    except Exception as e:
        logger.error(f"Embedding failed: {e}")
        return None
    
    # Step 4: Check for duplicates (similarity > 0.95 = duplicate)
    try:
        existing = pinecone_index.query(
            vector=embedding,
            top_k=1,
            namespace=memory_config.CONVERSATION_MEMORY_NAMESPACE,
            include_metadata=True
        )
        matches = existing.get("matches", [])
        if matches and matches[0]["score"] > 0.95:
            logger.info(f"⏭️ Skipping duplicate memory (similarity: {matches[0]['score']:.3f})")
            return None
    except Exception as e:
        logger.warning(f"Duplicate check failed, proceeding anyway: {e}")
    
    # Step 4: Upsert to Pinecone
    vectors = [(
        chunk["chunk_id"],
        embedding,
        {
            "text": chunk["text"],
            "doc_id": chunk["doc_id"],
            "section": chunk["section"],
            "page": chunk["page"],
        }
    )]
    
    success = await safe_upsert(
        pinecone_index,
        vectors,
        memory_config.CONVERSATION_MEMORY_NAMESPACE
    )
    
    if success:
        logger.info(f"✅ Stored memory: {chunk['chunk_id']}")
        return chunk
    else:
        logger.error("❌ Failed to store memory")
        return None


# ============================================================================
# CONVENIENCE FUNCTION
# ============================================================================

def store_memory_sync(
    raw_output: str,
    pinecone_index,
    session_id: str = None,
    turn_number: int = 0
) -> Optional[Dict]:
    """
    Synchronous wrapper for store_memory.
    """
    import asyncio
    
    if session_id is None:
        session_id = get_or_create_session_id()
    
    return asyncio.run(store_memory(
        raw_output,
        session_id,
        pinecone_index,
        turn_number
    ))