import os
import time
import logging
from typing import Optional
from dotenv import load_dotenv

from livekit.agents import (
    Agent,
    AgentSession,
    JobContext,
    JobProcess,
    RunContext,
    WorkerOptions,
    cli,
    function_tool,
)
from livekit.plugins import deepgram, openai, silero
from livekit.plugins.turn_detector.english import EnglishModel

from interview_questions import InterviewQuestions
from logging_config import setup_logging, InterviewLogger

# Load environment variables
load_dotenv()

# Set up logging
setup_logging(os.getenv("LOG_LEVEL", "INFO"))
logger = logging.getLogger(__name__)

class InterviewAgent(Agent):
    """LiveKit agent for conducting natural, low-latency interviews."""
    
    def __init__(self):
        super().__init__(
            instructions=(
                "You are an AI interviewer conducting a professional interview. "
                "Your role is to ask questions naturally and engage in conversation. "
                "Keep your responses concise and conversational. "
                "Ask follow-up questions when appropriate to get more details. "
                "Be professional but friendly in your tone."
            )
        )
        self.interview_questions = InterviewQuestions()
        self.interview_logger = InterviewLogger()
        self.current_question_start_time = None
        self.user_response_start_time = None
        self.turn_detection_start_time = None
        self.last_turn_detection_confidence = None
        
    async def on_enter(self):
        """Called when the agent enters the session."""
        self.interview_logger.log_interview_start()
        
        # Log turn detection model initialization
        self.interview_logger.log_turn_detection_model_event(
            "initialized",
            {"model_type": "EnglishModel", "status": "ready"}
        )
        
        # Start with greeting
        greeting = self.interview_questions.get_greeting_message()
        await self.session.generate_reply(instructions=f"Say: {greeting}")
        
        # Ask the first question
        await self._ask_next_question()
    
    async def _ask_next_question(self):
        """Ask the next question in the interview."""
        current_question = self.interview_questions.get_current_question()
        
        if not current_question:
            # Interview is complete
            await self._end_interview()
            return
        
        # Log question start
        self.current_question_start_time = time.time()
        self.interview_logger.log_question_start(
            current_question.number, 
            current_question.question
        )
        
        # Reset turn detection tracking for new question
        self.turn_detection_start_time = None
        self.last_turn_detection_confidence = None
        
        # Log agent starting to speak
        self.interview_logger.log_turn_detection_attempt(
            question_number=current_question.number,
            attempt_type="agent_speaking_start",
            confidence=1.0,
            success=True,
            details="Agent starting to ask question"
        )
        
        # Ask the question
        await self.session.generate_reply(
            instructions=f"Ask this question naturally: {current_question.question}"
        )
        
        # Log agent finished speaking
        self.interview_logger.log_turn_detection_attempt(
            question_number=current_question.number,
            attempt_type="agent_speaking_end",
            confidence=1.0,
            success=True,
            details="Agent finished asking question, waiting for user response"
        )
    
    async def _ask_follow_up_question(self, user_response: str):
        """Ask a follow-up question based on the user's response."""
        if not self.interview_questions.can_ask_follow_up():
            # Move to next main question
            await self._move_to_next_question()
            return
        
        # Log follow-up start
        self.current_question_start_time = time.time()
        current_question = self.interview_questions.get_current_question()
        self.interview_logger.log_question_start(
            current_question.number,
            "follow_up"
        )
        
        # Reset turn detection tracking for follow-up
        self.turn_detection_start_time = None
        self.last_turn_detection_confidence = None
        
        # Log agent starting to speak follow-up
        self.interview_logger.log_turn_detection_attempt(
            question_number=current_question.number,
            attempt_type="agent_speaking_followup_start",
            confidence=1.0,
            success=True,
            details="Agent starting to ask follow-up question"
        )
        
        # Generate follow-up instruction
        follow_up_instruction = self.interview_questions.get_follow_up_instruction(user_response)
        
        # Ask follow-up
        await self.session.generate_reply(instructions=follow_up_instruction)
        
        # Log agent finished speaking follow-up
        self.interview_logger.log_turn_detection_attempt(
            question_number=current_question.number,
            attempt_type="agent_speaking_followup_end",
            confidence=1.0,
            success=True,
            details="Agent finished asking follow-up question, waiting for user response"
        )
        
        # Increment follow-up count
        self.interview_questions.increment_follow_up_count()
    
    async def _move_to_next_question(self):
        """Move to the next main question."""
        # Log question end
        if self.current_question_start_time:
            duration = time.time() - self.current_question_start_time
            self.interview_logger.log_question_end(
                self.interview_questions.get_current_question().number,
                duration
            )
        
        # Get next question
        next_question = self.interview_questions.get_next_question()
        
        if next_question:
            await self._ask_next_question()
        else:
            await self._end_interview()
    
    async def _end_interview(self):
        """End the interview with a farewell message."""
        farewell = self.interview_questions.get_farewell_message()
        
        # Log interview end
        if self.interview_logger.start_times.get("interview"):
            total_duration = time.time() - self.interview_logger.start_times["interview"]
            self.interview_logger.log_interview_end(
                total_duration,
                len(self.interview_questions.questions)
            )
        
        await self.session.generate_reply(instructions=f"Say: {farewell}")
    
    @function_tool
    async def process_user_response(
        self,
        context: RunContext,
        user_response: str,
        question_number: int,
        is_follow_up: bool = False
    ):
        """Process the user's response and determine next action."""
        
        # Debug: Log function tool call
        logger.info(f"🔧 FUNCTION TOOL CALLED: process_user_response")
        logger.info(f"   User response: {user_response[:100]}...")
        logger.info(f"   Question number: {question_number}")
        logger.info(f"   Is follow-up: {is_follow_up}")
        print(f"🔧 FUNCTION TOOL CALLED: process_user_response - {user_response[:50]}...")
        
        # Log user response start (this is called when user starts speaking)
        if not self.user_response_start_time:
            self.user_response_start_time = time.time()
            self.interview_logger.log_user_response_start(question_number)
            
            # Log that turn detection should be active
            self.interview_logger.log_turn_detection_attempt(
                question_number=question_number,
                attempt_type="user_response_start",
                confidence=0.5,  # Initial confidence when user starts
                success=True,
                details="User started responding, turn detection should be active"
            )
        
        # Log user response end
        if self.user_response_start_time:
            duration = time.time() - self.user_response_start_time
            self.interview_logger.log_user_response_end(
                question_number,
                user_response,
                duration
            )
            self.user_response_start_time = None  # Reset for next response
        
        # Log turn detection success
        if self.turn_detection_start_time:
            turn_detection_duration = time.time() - self.turn_detection_start_time
            self.interview_logger.log_turn_detection_attempt(
                question_number=question_number,
                attempt_type="user_finished",
                confidence=self.last_turn_detection_confidence or 0.9,
                success=True,
                details=f"Turn detected after {turn_detection_duration:.2f}s"
            )
        else:
            # Log turn detection failure (no start time recorded)
            self.interview_logger.log_turn_detection_attempt(
                question_number=question_number,
                attempt_type="user_finished",
                confidence=0.0,
                success=False,
                details="No turn detection start time recorded - turn detection may not be working"
            )
        
        # Determine next action
        if is_follow_up and self.interview_questions.can_ask_follow_up():
            # Ask another follow-up
            await self._ask_follow_up_question(user_response)
        else:
            # Move to next main question
            await self._move_to_next_question()
        
        return {"status": "processed", "next_action": "question"}

def prewarm(proc: JobProcess):
    """Preload models for faster startup."""
    logger.info("Preloading VAD model...")
    proc.userdata["vad"] = silero.VAD.load()
    logger.info("VAD model loaded successfully")

async def entrypoint(ctx: JobContext):
    """Main entry point for the interview agent."""
    
    # Set up telemetry for debugging
    from livekit.agents.telemetry import set_tracer_provider
    from opentelemetry.sdk.trace import TracerProvider
    from opentelemetry.sdk.trace.export import ConsoleSpanExporter, BatchSpanProcessor
    
    # Set up console telemetry to see traces
    trace_provider = TracerProvider()
    trace_provider.add_span_processor(BatchSpanProcessor(ConsoleSpanExporter()))
    set_tracer_provider(trace_provider)
    
    # Set up logging context
    ctx.log_context_fields = {
        "room": ctx.room.name,
        "agent_type": "interview_agent"
    }
    
    logger.info(f"Starting interview agent session for room: {ctx.room.name}")
    logger.info("🔍 Telemetry enabled - will show traces for debugging")
    
    # Create agent instance for logging
    agent = InterviewAgent()
    
    # Create agent session with optimized settings for low latency
    session = AgentSession(
        vad=ctx.proc.userdata["vad"],
        stt=deepgram.STT(
            model="nova-3",
            language="en",
            # Optimize for speed
            interim_results=True,
            punctuate=True,
            smart_format=True
        ),
        llm=openai.LLM(
            model="gpt-4o-mini",
            # Optimize for speed
            temperature=0.7
        ),
        tts=openai.TTS(
            voice="alloy",
            # Optimize for speed
            speed=1.0
        ),
        # Use LiveKit's turn detection for natural conversation flow
        turn_detection=EnglishModel(),
        # Allow interruptions for more natural flow
        allow_interruptions=True
    )
    
    # Set up metrics collection
    from livekit.agents.voice import MetricsCollectedEvent
    from livekit.agents import metrics
    
    usage_collector = metrics.UsageCollector()
    
    # Track metrics for latency calculation
    eou_metrics = {}
    llm_metrics = {}
    tts_metrics = {}
    
    @session.on("metrics_collected")
    def _on_metrics_collected(ev: MetricsCollectedEvent):
        # Use LiveKit's recommended log_metrics helper
        metrics.log_metrics(ev.metrics)
        
        # Track metrics by type for latency calculation
        metrics_type = getattr(ev.metrics, 'type', 'unknown')
        speech_id = getattr(ev.metrics, 'speech_id', None)
        
        if metrics_type == 'eou_metrics' and speech_id:
            eou_metrics[speech_id] = ev.metrics
            logger.info(f"🎯 EOU METRICS CAPTURED: {ev.metrics}")
            
        elif metrics_type == 'llm_metrics' and speech_id:
            llm_metrics[speech_id] = ev.metrics
            logger.info(f"🧠 LLM METRICS CAPTURED: {ev.metrics}")
            
        elif metrics_type == 'tts_metrics' and speech_id:
            tts_metrics[speech_id] = ev.metrics
            logger.info(f"🔊 TTS METRICS CAPTURED: {ev.metrics}")
            
            # Calculate total latency when we have all three metrics for a speech_id
            if speech_id in eou_metrics and speech_id in llm_metrics and speech_id in tts_metrics:
                eou = eou_metrics[speech_id]
                llm = llm_metrics[speech_id]
                tts = tts_metrics[speech_id]
                
                # Calculate total latency according to LiveKit formula
                total_latency = (
                    getattr(eou, 'end_of_utterance_delay', 0) +
                    getattr(llm, 'ttft', 0) +
                    getattr(tts, 'ttfb', 0)
                )
                
                logger.info(f"⚡ TOTAL LATENCY CALCULATION for {speech_id}:")
                logger.info(f"   EOU delay: {getattr(eou, 'end_of_utterance_delay', 0):.3f}s")
                logger.info(f"   LLM ttft: {getattr(llm, 'ttft', 0):.3f}s")
                logger.info(f"   TTS ttfb: {getattr(tts, 'ttfb', 0):.3f}s")
                logger.info(f"   TOTAL: {total_latency:.3f}s")
                
                # Log detailed metrics for analysis
                agent.interview_logger.log_turn_detection_metrics(
                    metrics={
                        "type": "latency_analysis",
                        "speech_id": speech_id,
                        "total_latency": total_latency,
                        "eou_metrics": {
                            "end_of_utterance_delay": getattr(eou, 'end_of_utterance_delay', 0),
                            "transcription_delay": getattr(eou, 'transcription_delay', 0),
                            "on_user_turn_completed_delay": getattr(eou, 'on_user_turn_completed_delay', 0)
                        },
                        "llm_metrics": {
                            "duration": getattr(llm, 'duration', 0),
                            "ttft": getattr(llm, 'ttft', 0),
                            "completion_tokens": getattr(llm, 'completion_tokens', 0),
                            "tokens_per_second": getattr(llm, 'tokens_per_second', 0)
                        },
                        "tts_metrics": {
                            "duration": getattr(tts, 'duration', 0),
                            "ttfb": getattr(tts, 'ttfb', 0),
                            "audio_duration": getattr(tts, 'audio_duration', 0),
                            "characters_count": getattr(tts, 'characters_count', 0)
                        }
                    }
                )
        
        # Enhanced turn detection logging for other metrics
        metrics_str = str(ev.metrics).lower()
        if 'turn' in metrics_str or 'detection' in metrics_str:
            logger.info(f"🎯 TURN DETECTION METRIC: {ev.metrics}")
            
            # Log detailed turn detection metrics
            agent.interview_logger.log_turn_detection_metrics(
                metrics={
                    "type": getattr(ev.metrics, 'type', 'unknown'),
                    "label": getattr(ev.metrics, 'label', 'unknown'),
                    "timestamp": getattr(ev.metrics, 'timestamp', time.time()),
                    "raw_metrics": str(ev.metrics)
                }
            )
            
            # Extract confidence if available
            if hasattr(ev.metrics, 'confidence'):
                agent.last_turn_detection_confidence = ev.metrics.confidence
            elif hasattr(ev.metrics, 'value') and isinstance(ev.metrics.value, (int, float)):
                agent.last_turn_detection_confidence = float(ev.metrics.value)
        
        usage_collector.collect(ev.metrics)
    
    # Add turn detection event handlers
    logger.info("🔧 Registering turn_detected event handler")
    @session.on("turn_detected")
    def _on_turn_detected(ev):
        """Handle turn detection events."""
        logger.info(f"🎯 TURN DETECTED: {ev}")
        print(f"🎯 TURN DETECTED: {ev}")  # Immediate console feedback
        
        # Record turn detection start time
        agent.turn_detection_start_time = time.time()
        
        # Log successful turn detection
        current_question = agent.interview_questions.get_current_question()
        question_number = current_question.number if current_question else 0
        
        agent.interview_logger.log_turn_detection_attempt(
            question_number=question_number,
            attempt_type="turn_detected",
            confidence=getattr(ev, 'confidence', 0.9),
            success=True,
            details=f"Turn detected: {ev}"
        )
    
    @session.on("turn_detection_failed")
    def _on_turn_detection_failed(ev):
        """Handle turn detection failure events."""
        logger.warning(f"❌ TURN DETECTION FAILED: {ev}")
        
        current_question = agent.interview_questions.get_current_question()
        question_number = current_question.number if current_question else 0
        
        agent.interview_logger.log_turn_detection_attempt(
            question_number=question_number,
            attempt_type="turn_detection_failed",
            confidence=getattr(ev, 'confidence', 0.0),
            success=False,
            details=f"Turn detection failed: {ev}"
        )
    
    @session.on("turn_detection_timeout")
    def _on_turn_detection_timeout(ev):
        """Handle turn detection timeout events."""
        logger.warning(f"⏰ TURN DETECTION TIMEOUT: {ev}")
        
        current_question = agent.interview_questions.get_current_question()
        question_number = current_question.number if current_question else 0
        
        timeout_duration = getattr(ev, 'timeout_duration', 0.0)
        agent.interview_logger.log_turn_detection_timeout(
            question_number=question_number,
            timeout_duration=timeout_duration,
            last_confidence=agent.last_turn_detection_confidence
        )
    
    # Add session event handlers for turn detection model status
    logger.info("🔧 Registering user_state_changed event handler")
    @session.on("user_state_changed")
    def _on_user_state_changed(ev):
        """Handle user state change events."""
        logger.info(f"👤 USER STATE CHANGED: {ev}")
        print(f"👤 USER STATE CHANGED: {ev}")  # Immediate console feedback
        
        current_question = agent.interview_questions.get_current_question()
        question_number = current_question.number if current_question else 0
        
        agent.interview_logger.log_turn_detection_attempt(
            question_number=question_number,
            attempt_type="user_state_changed",
            confidence=0.9,
            success=True,
            details=f"User state changed: {ev}"
        )

    @session.on("session_started")
    def _on_session_started(ev):
        """Handle session start events."""
        logger.info(f"🚀 Session started: {ev}")
        agent.interview_logger.log_turn_detection_model_event(
            "session_started",
            {"event": str(ev), "status": "active"}
        )
    
    @session.on("session_ended")
    def _on_session_ended(ev):
        """Handle session end events."""
        logger.info(f"🏁 Session ended: {ev}")
        agent.interview_logger.log_turn_detection_model_event(
            "session_ended",
            {"event": str(ev), "status": "inactive"}
        )
    
    # Add error handling for turn detection model
    @session.on("error")
    def _on_error(ev):
        """Handle error events."""
        error_msg = str(ev)
        logger.error(f"❌ Session error: {error_msg}")
        
        # Check if it's a turn detection related error
        if 'turn' in error_msg.lower() or 'detection' in error_msg.lower():
            agent.interview_logger.log_turn_detection_model_event(
                "error",
                {"error_message": error_msg, "event": str(ev)}
            )
        else:
            agent.interview_logger.log_error(
                "session_error",
                error_msg
            )
    
    async def log_final_usage():
        summary = usage_collector.get_summary()
        logger.info(f"Final usage summary: {summary}")
    
    ctx.add_shutdown_callback(log_final_usage)
    
    # Start the session
    await session.start(
        agent=agent,
        room=ctx.room
    )

if __name__ == "__main__":
    cli.run_app(
        WorkerOptions(
            entrypoint_fnc=entrypoint,
            prewarm_fnc=prewarm
        )
    ) 