import asyncio
import httpx
import json
import logging
import websockets
import uuid
import re

# Configuration
BASE_URL = "https://dev-strategistai.galaxiq.ai"
WS_URL = "wss://dev-strategistai.galaxiq.ai"
TENANT_ID = "org_d1e1c1ea-6c90-431f-9337-3375dcdfef5c"

logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger("journey_test")

class FullJourneyTester:
    def __init__(self, thread_id):
        self.thread_id = thread_id
        self.admin_ws = None
        self.user_ws = None
        self.admin_messages = []
        self.user_messages = []
        self.stop_event = asyncio.Event()

    async def connect_websockets(self):
        admin_url = f"{WS_URL}/ws/{TENANT_ID}/{self.thread_id}?role=admin&user_name=AdminAgent"
        user_url = f"{WS_URL}/ws/{TENANT_ID}/{self.thread_id}?role=user&user_name=Customer"
        
        self.admin_ws = await websockets.connect(admin_url)
        self.user_ws = await websockets.connect(user_url)
        
        async def listen_admin():
            logger.info("Admin WS listener started.")
            while not self.stop_event.is_set():
                try:
                    msg = await asyncio.wait_for(self.admin_ws.recv(), timeout=0.5)
                    data = json.loads(msg)
                    self.admin_messages.append(data)
                    logger.info(f" [ADMIN WS] Received: {data.get('role', 'system')}: {data.get('text', data.get('message', ''))}")
                except asyncio.TimeoutError:
                    continue
                except Exception as e:
                    logger.error(f"Admin WS error: {e}")
                    break

        async def listen_user():
            logger.info("User WS listener started.")
            while not self.stop_event.is_set():
                try:
                    msg = await asyncio.wait_for(self.user_ws.recv(), timeout=0.5)
                    data = json.loads(msg)
                    self.user_messages.append(data)
                    logger.info(f" [USER WS] Received: {data.get('role', 'system')}: {data.get('text', data.get('message', ''))}")
                except asyncio.TimeoutError:
                    continue
                except Exception as e:
                    logger.error(f"User WS error: {e}")
                    break

        asyncio.create_task(listen_admin())
        asyncio.create_task(listen_user())

async def run_journey():
    thread_id = f"journey-{uuid.uuid4().hex[:8]}"
    tester = FullJourneyTester(thread_id)
    await tester.connect_websockets()
    
    async with httpx.AsyncClient(timeout=60.0) as client:
        # --- PHASE 1: AI INTERACTION ---
        logger.info("\n--- PHASE 1: AI INTERACTION ---")
        queries = ["hi", "what is galaxiq?", "what are its products?"]
        for q in queries:
            logger.info(f"User asking: {q}")
            res = await client.post(f"{BASE_URL}/chat", data={"tenant_id": TENANT_ID, "query": q, "thread_id": thread_id})
            logger.info(f"AI Response: {res.json().get('response')[:100]}...")
            await asyncio.sleep(1)

        # --- PHASE 2: REQUEST HANDOVER ---
        logger.info("\n--- PHASE 2: REQUEST HANDOVER ---")
        logger.info("User: I want to talk to a human admin.")
        res = await client.post(f"{BASE_URL}/chat", data={"tenant_id": TENANT_ID, "query": "I want to talk to a human admin please.", "thread_id": thread_id})
        logger.info(f"AI Response: {res.json().get('response')}")

        # --- PHASE 3: ADMIN TAKEOVER ---
        logger.info("\n--- PHASE 3: ADMIN TAKEOVER ---")
        logger.info("Admin activating takeover...")
        await client.post(f"{BASE_URL}/chat/takeover", data={"tenant_id": TENANT_ID, "thread_id": thread_id, "status": "true", "user_id": "admin-456"})
        await asyncio.sleep(2)

        # --- PHASE 4: HUMAN-TO-HUMAN REAL-TIME ---
        logger.info("\n--- PHASE 4: HUMAN-TO-HUMAN REAL-TIME ---")
        
        logger.info("User sending message via WebSocket...")
        await tester.user_ws.send(json.dumps({"role": "user", "user_name": "Customer", "text": "Hi Admin, I need help with my billing.", "type": "chat_message"}))
        await asyncio.sleep(2)
        
        logger.info("Admin sending message via WebSocket...")
        await tester.admin_ws.send(json.dumps({"role": "admin", "user_name": "AdminAgent", "text": "Hello! I am the admin. I can help with your billing. What seems to be the issue?", "type": "chat_message"}))
        await asyncio.sleep(2)

        # --- PHASE 5: DEACTIVATE TAKEOVER ---
        logger.info("\n--- PHASE 5: DEACTIVATE TAKEOVER ---")
        logger.info("Admin deactivating takeover...")
        await client.post(f"{BASE_URL}/chat/takeover", data={"tenant_id": TENANT_ID, "thread_id": thread_id, "status": "false"})
        await asyncio.sleep(2)

        # --- PHASE 6: AI RESUMES ---
        logger.info("\n--- PHASE 6: AI RESUMES ---")
        logger.info("User: Thank you for the help.")
        res = await client.post(f"{BASE_URL}/chat", data={"tenant_id": TENANT_ID, "query": "Thank you for the help, everything is clear now.", "thread_id": thread_id})
        logger.info(f"AI Response: {res.json().get('response')}")

    # Final Verification
    tester.stop_event.set()
    await asyncio.sleep(1)
    await tester.admin_ws.close()
    await tester.user_ws.close()
    
    print("\n" + "="*50)
    print("JOURNEY TEST SUMMARY")
    print("="*50)
    
    # Check if takeover messages were received
    user_saw_admin = any(m.get('role') == 'admin' and "billing" in m.get('text', '').lower() for m in tester.user_messages)
    admin_saw_user = any(m.get('role') == 'user' and "billing" in m.get('text', '').lower() for m in tester.admin_messages)
    
    print(f"User received Admin's message: {'PASS' if user_saw_admin else 'FAIL'}")
    print(f"Admin received User's message: {'PASS' if admin_saw_user else 'FAIL'}")
    
    if user_saw_admin and admin_saw_user:
        print("\nOVERALL RESULT: SUCCESS - End-to-end journey verified!")
    else:
        print("\nOVERALL RESULT: FAILURE - Real-time handover messages missing.")
    print("="*50)

if __name__ == "__main__":
    asyncio.run(run_journey())
