import asyncio
import logging
import time
from unittest.mock import MagicMock, AsyncMock

# Add current directory to path so we can import src
import sys
import os
sys.path.append(os.getcwd())

from src.services.rag_service import RERAInferenceEngine

async def test_cache_logic():
    # Setup
    logger = logging.getLogger("test")
    vector_store = MagicMock()
    engine = RERAInferenceEngine(vector_store, logger)
    
    # Mock MongoDB
    engine.db = MagicMock()
    mock_cursor = AsyncMock()
    engine.db["prompts"].find.return_value = mock_cursor
    mock_cursor.sort.return_value = mock_cursor
    
    print("--- Phase 1: Initial Prompt ---")
    mock_cursor.to_list.return_value = [{"content": "Default Rule", "is_active": True, "updated_at": time.time()}]
    
    instructions1 = await engine._get_system_instructions()
    print(f"Result 1: {instructions1}")
    
    print("\n--- Phase 2: Updating database but NO clear_cache ---")
    # Change the DB content
    mock_cursor.to_list.return_value = [
        {"content": "Default Rule", "is_active": True, "updated_at": time.time()},
        {"content": "Ignore Webbrains", "is_active": True, "updated_at": time.time() + 10}
    ]
    
    instructions2 = await engine._get_system_instructions()
    print(f"Result 2 (Should be same as 1 due to cache): {instructions2}")
    
    if instructions1 == instructions2:
        print("Success: Cache is working (returned old instructions).")
    else:
        print("Failure: Cache not working.")

    print("\n--- Phase 3: Calling clear_cache() ---")
    engine.clear_cache()
    
    instructions3 = await engine._get_system_instructions()
    print(f"Result 3 (Should be new aggregated prompt): {instructions3}")
    
    if "Ignore Webbrains" in instructions3:
        print("Success: Cache Invalidation Works! Model picked up new instructions.")
    else:
        print("Failure: Model still using old instructions after clear_cache.")

if __name__ == "__main__":
    asyncio.run(test_cache_logic())
