# import asyncio
# import json
# from fastapi import WebSocket, WebSocketDisconnect, APIRouter
# from ..webhook.memory import RedisSession
# from ..webhook.response_streamer import stream_response
# from ..repositories import repo

# router = APIRouter()


# clients: dict[str, WebSocket] = {}
# active_streams: dict[str, asyncio.Task] = {}
# sessions: dict[str, RedisSession] = {}

# @router.websocket("/ws")
# async def websocket_endpoint(ws: WebSocket):
#     await ws.accept()
#     session_id = None

#     try:
#         while True:
#             msg = json.loads(await ws.receive_text())
#             msg_type = msg.get("type")

#             if msg_type == "setup":
#                 session_id = msg.get("callSid")
#                 if not session_id:
#                     await ws.send_text(json.dumps({"type": "error", "msg": "Missing callSid"}))
#                     continue
#                 clients[session_id] = ws
#                 sessions[session_id] = RedisSession(name=f"chat_{session_id}")
#                 print(f"Session started: {session_id}")

#             elif msg_type == "prompt" and session_id:
#                 agent_id = repo.get_agent_id_by_call_id("calls", session_id)
#                 if not agent_id:
#                     print("No agent found for session %s", session_id)
#                     return

#                 kb_ids = repo.get_agent_knowledge_base_ids(agent_id)
#                 if not kb_ids:
#                     print("No knowledge_base_ids found for session %s", session_id)
#                     return
                
#                 user_query = msg["voicePrompt"]
#                 sessions[session_id].add_message("user", user_query)
                
#                 task = asyncio.create_task(stream_response(ws, user_query,sessions[session_id],kb_ids))
#                 active_streams[session_id] = task
#                 await task
#                 active_streams.pop(session_id, None)

#             elif msg_type == "interrupt" and session_id:
#                 task = active_streams.pop(session_id, None)
#                 if task:
#                     task.cancel()
#                     print(f"Stream interrupted: {session_id}")

#             elif msg_type == "error":
#                 print(f"Client error: {msg.get('description')}")
#                 await ws.send_text(json.dumps({"type": "end"}))

#     except WebSocketDisconnect:
#         if session_id:
#             task = active_streams.pop(session_id, None)
#             if task:
#                 task.cancel()
#             clients.pop(session_id, None)
#             sessions.pop(session_id, None)
#             print(f"Disconnected: {session_id}")
#         else:
#             print("Disconnected: unknown client")