"""
Unit tests for Task Management MCP Server

Tests the core functionality of the task management MCP server including
task analysis, scheduling, risk prediction, and progress tracking.
"""

import pytest
import asyncio
import os
from datetime import datetime, timedelta
from unittest.mock import Mock, patch, AsyncMock

import sys
sys.path.append(os.path.join(os.path.dirname(__file__), '../..'))

from task_management_mcp.server import mcp, TaskAnalysisResult, ScheduleRecommendation
from shared.models import Task, TaskStatus, TaskPriority


class TestTaskManagementMCP:
    """Test suite for Task Management MCP Server"""
    
    @pytest.fixture
    def mock_context(self):
        """Create a mock context for testing"""
        context = Mock()
        context.request_context = Mock()
        context.request_context.lifespan_context = Mock()
        
        # Mock app context with sample data
        app_ctx = Mock()
        app_ctx.tasks = {}
        app_ctx.projects = {}
        app_ctx.reminders = {}
        app_ctx.requirement_agent = None
        app_ctx.task_generation_agent = None
        
        context.request_context.lifespan_context = app_ctx
        return context
    
    def test_analyze_task_input_without_ai(self, mock_context):
        """Test task analysis when AI agents are not available"""
        from task_management_mcp.server import analyze_task_input
        
        result = analyze_task_input(
            task_description="Create a user registration system",
            ctx=mock_context,
            project_id="test_project"
        )
        
        assert isinstance(result, TaskAnalysisResult)
        assert len(result.tasks) == 1
        assert result.tasks[0].name == "Create a user registration system"
        assert result.tasks[0].project_id == "test_project"
        assert result.total_estimated_time > 0
        assert "Intelligent fallback analysis" in result.risk_summary
    
    @patch('shared.agents.requirement_agent.RequirementAgent')
    @patch('shared.agents.task_generation_agent.TaskGenerationAgent')
    def test_analyze_task_input_with_ai(self, mock_task_gen, mock_req_agent, mock_context):
        """Test task analysis with AI agents available"""
        from task_management_mcp.server import analyze_task_input
        
        # Setup mock agents
        mock_req_instance = Mock()
        mock_req_instance.generate_requirements_document.return_value = "Mock requirements"
        mock_req_agent.return_value = mock_req_instance
        
        mock_task_gen_instance = Mock()
        mock_task_gen_instance.parse_document_to_tasks.return_value = [
            {
                "name": "User Interface Design",
                "description": "Design the registration form",
                "priority": 2,
                "time_cost": 4.0,
                "importance": 4,
                "benefit": 7,
                "marginal_cost": 2
            },
            {
                "name": "Backend Implementation", 
                "description": "Implement registration logic",
                "priority": 1,
                "time_cost": 6.0,
                "importance": 5,
                "benefit": 8,
                "marginal_cost": 3
            }
        ]
        mock_task_gen.return_value = mock_task_gen_instance
        
        # Setup context with AI agents
        mock_context.request_context.lifespan_context.requirement_agent = mock_req_instance
        mock_context.request_context.lifespan_context.task_generation_agent = mock_task_gen_instance
        
        result = analyze_task_input(
            task_description="Create a user registration system",
            ctx=mock_context,
            project_id="test_project"
        )
        
        assert isinstance(result, TaskAnalysisResult)
        assert len(result.tasks) == 2
        assert result.total_estimated_time == 10.0
        assert result.priority_distribution["critical"] == 1  # priority 1 = CRITICAL
        assert result.priority_distribution["high"] == 1      # priority 2 = HIGH
    
    def test_get_schedule_recommendation_empty_tasks(self, mock_context):
        """Test schedule recommendation with no tasks"""
        from task_management_mcp.server import get_schedule_recommendation
        
        result = get_schedule_recommendation(ctx=mock_context)
        
        assert isinstance(result, ScheduleRecommendation)
        assert len(result.recommended_order) == 0
        assert len(result.critical_path) == 0
        assert len(result.parallel_opportunities) == 0
    
    def test_get_schedule_recommendation_with_tasks(self, mock_context):
        """Test schedule recommendation with sample tasks"""
        from task_management_mcp.server import get_schedule_recommendation
        
        # Add sample tasks to context
        task1 = Task(
            id="task1",
            name="High Priority Task",
            priority=TaskPriority.HIGH,
            time_cost=4.0,
            importance=5
        )
        task2 = Task(
            id="task2", 
            name="Low Priority Task",
            priority=TaskPriority.LOW,
            time_cost=2.0,
            importance=2,
            dependencies=["task1"]
        )
        
        mock_context.request_context.lifespan_context.tasks = {
            "task1": task1,
            "task2": task2
        }
        
        result = get_schedule_recommendation(
            ctx=mock_context,
            available_hours_per_day=8.0
        )
        
        assert isinstance(result, ScheduleRecommendation)
        assert len(result.recommended_order) == 2
        assert result.recommended_order[0] == "task1"  # Higher priority first
        assert "task2" in result.critical_path  # Has dependencies
        assert result.estimated_completion_date > datetime.now()
    
    def test_predict_task_risks(self, mock_context):
        """Test risk prediction for tasks"""
        from task_management_mcp.server import predict_task_risks
        
        # Add a high-risk task
        risky_task = Task(
            id="risky_task",
            name="Complex Integration Task",
            time_cost=20.0,  # Large time estimate
            priority=TaskPriority.CRITICAL,  # Critical priority
            due_date=datetime.now() + timedelta(hours=12),  # Tight deadline
            dependencies=["task1", "task2"]  # Multiple dependencies
        )
        
        mock_context.request_context.lifespan_context.tasks = {
            "risky_task": risky_task
        }
        
        result = predict_task_risks("risky_task", mock_context)
        
        assert result.task_id == "risky_task"
        assert result.risk_level >= 4  # Should be high risk
        assert len(result.risk_factors) > 0
        assert len(result.mitigation_strategies) > 0
        assert result.probability > 0.5  # High probability
    
    def test_predict_task_risks_not_found(self, mock_context):
        """Test risk prediction for non-existent task"""
        from task_management_mcp.server import predict_task_risks
        
        with pytest.raises(ValueError, match="Task nonexistent not found"):
            predict_task_risks("nonexistent", mock_context)
    
    def test_create_reminder(self, mock_context):
        """Test reminder creation"""
        from task_management_mcp.server import create_reminder
        
        # Add a task to context
        task = Task(
            id="test_task",
            name="Test Task",
            due_date=datetime.now() + timedelta(days=3)
        )
        mock_context.request_context.lifespan_context.tasks = {"test_task": task}
        
        result = create_reminder(
            task_id="test_task",
            reminder_type="deadline",
            ctx=mock_context
        )
        
        assert result.task_id == "test_task"
        assert result.reminder_type == "deadline"
        assert "deadline approaching" in result.message.lower()
        assert result.scheduled_time < task.due_date
    
    def test_update_task_progress(self, mock_context):
        """Test task progress updates"""
        from task_management_mcp.server import update_task_progress
        
        # Add a task to context
        task = Task(
            id="test_task",
            name="Test Task",
            status=TaskStatus.TODO,
            progress=0.0
        )
        mock_context.request_context.lifespan_context.tasks = {"test_task": task}
        
        result = update_task_progress(
            task_id="test_task",
            progress=0.5,
            ctx=mock_context,
            time_spent=2.0,
            notes="Halfway done"
        )
        
        assert result.task_id == "test_task"
        assert result.progress == 0.5
        assert result.time_spent == 2.0
        assert result.notes == "Halfway done"
        
        # Check that task was updated
        updated_task = mock_context.request_context.lifespan_context.tasks["test_task"]
        assert updated_task.progress == 0.5
        assert updated_task.status == TaskStatus.IN_PROGRESS
        assert updated_task.actual_time_spent == 2.0
    
    def test_update_task_progress_completion(self, mock_context):
        """Test task progress update to completion"""
        from task_management_mcp.server import update_task_progress
        
        # Add a task to context
        task = Task(
            id="test_task",
            name="Test Task",
            status=TaskStatus.IN_PROGRESS,
            progress=0.8
        )
        mock_context.request_context.lifespan_context.tasks = {"test_task": task}
        
        update_task_progress(
            task_id="test_task",
            progress=1.0,
            ctx=mock_context
        )
        
        # Check that task status was updated to completed
        updated_task = mock_context.request_context.lifespan_context.tasks["test_task"]
        assert updated_task.progress == 1.0
        assert updated_task.status == TaskStatus.COMPLETED
    
    def test_get_next_task_recommendation(self, mock_context):
        """Test next task recommendation"""
        from task_management_mcp.server import get_next_task_recommendation
        
        # Add sample tasks
        task1 = Task(
            id="task1",
            name="Quick Task",
            status=TaskStatus.TODO,
            priority=TaskPriority.HIGH,
            time_cost=1.0,
            benefit=8,
            importance=4
        )
        task2 = Task(
            id="task2",
            name="Long Task", 
            status=TaskStatus.TODO,
            priority=TaskPriority.MEDIUM,
            time_cost=8.0,
            benefit=5,
            importance=3
        )
        
        mock_context.request_context.lifespan_context.tasks = {
            "task1": task1,
            "task2": task2
        }
        
        result = get_next_task_recommendation(
            ctx=mock_context,
            available_time=2.0
        )
        
        assert result.task_id == "task1"  # Should recommend the quick, high-priority task
        assert result.score > 0
        assert "Quick Task" in result.reasoning


if __name__ == "__main__":
    pytest.main([__file__])
