import pytest
import asyncio
import json
import sys
from mcp.client.sse import sse_client
from mcp.client.session import ClientSession

EXPECTED_FACT_CODES = {
    "Area 1": ["RULE-ROOT-01", "RULE-ROOT-02", "RULE-ROOT-03", "RULE-ROOT-04", "RULE-ROOT-05"],
    "Area 2": ["RULE-BACKEND-01", "RULE-BACKEND-02", "RULE-BACKEND-03", "RULE-BACKEND-04", "RULE-BACKEND-05"],
    "Area 3": ["RULE-FRONTEND-01", "RULE-FRONTEND-02", "RULE-FRONTEND-03", "RULE-FRONTEND-04", "RULE-FRONTEND-05"],
    "Area 4": ["RULE-AEC-01", "RULE-AEC-02", "RULE-AEC-03", "RULE-AEC-04", "RULE-AEC-05", "RULE-AEC-06", "RULE-AEC-07"],
    "Area 5": ["RULE-WORKFLOW-01", "RULE-WORKFLOW-02", "RULE-WORKFLOW-03", "RULE-WORKFLOW-04", "RULE-WORKFLOW-05"]
}

ALL_CODES = [code for codes in EXPECTED_FACT_CODES.values() for code in codes]

@pytest.mark.asyncio
async def test_mem0_status_count_is_476():
    """Verify Mem0 background service is online and total memory count is exactly 476 (449 + 27)."""
    async with sse_client("http://127.0.0.1:8765/sse") as (read, write):
        async with ClientSession(read, write) as session:
            await session.initialize()
            res = await session.call_tool("mem0_status", {})
            data = json.loads(res.content[0].text)
            assert data.get("status") == "online"
            assert data.get("is_ready") is True
            assert data.get("developer_memories_count") == 476, f"Expected 476, got {data.get('developer_memories_count')}"

@pytest.mark.asyncio
async def test_all_27_atomic_facts_present():
    """Verify all 27 atomic facts are present and uniquely identifiable in Mem0 with project='DSCons'."""
    async with sse_client("http://127.0.0.1:8765/sse") as (read, write):
        async with ClientSession(read, write) as session:
            await session.initialize()
            
            res = await session.call_tool("mem0_get_all", {"user_id": "developer"})
            data = json.loads(res.content[0].text)
            mems = data.get("memories", {}).get("results", [])
            dscons_mems = [m for m in mems if (m.get("metadata") or {}).get("project") == "DSCons"]
            
            assert len(dscons_mems) == 27, f"Expected 27 DSCons memories, found {len(dscons_mems)}"
            
            found_codes = set()
            for m in dscons_mems:
                code = (m.get("metadata") or {}).get("code")
                if code:
                    found_codes.add(code)
            
            missing_codes = set(ALL_CODES) - found_codes
            assert len(missing_codes) == 0, f"Missing codes: {missing_codes}"
            assert len(found_codes) == 27, f"Expected 27 unique codes, found {len(found_codes)}"

@pytest.mark.asyncio
async def test_all_5_areas_coverage():
    """Verify complete coverage across all 5 knowledge areas (Area 1: 5, Area 2: 5, Area 3: 5, Area 4: 7, Area 5: 5)."""
    async with sse_client("http://127.0.0.1:8765/sse") as (read, write):
        async with ClientSession(read, write) as session:
            await session.initialize()
            
            res = await session.call_tool("mem0_get_all", {"user_id": "developer"})
            data = json.loads(res.content[0].text)
            mems = data.get("memories", {}).get("results", [])
            dscons_mems = [m for m in mems if (m.get("metadata") or {}).get("project") == "DSCons"]
            
            area_counts = {"Area 1": 0, "Area 2": 0, "Area 3": 0, "Area 4": 0, "Area 5": 0}
            for m in dscons_mems:
                area = (m.get("metadata") or {}).get("area")
                if area in area_counts:
                    area_counts[area] += 1
            
            assert area_counts["Area 1"] == 5, f"Area 1 expected 5, got {area_counts['Area 1']}"
            assert area_counts["Area 2"] == 5, f"Area 2 expected 5, got {area_counts['Area 2']}"
            assert area_counts["Area 3"] == 5, f"Area 3 expected 5, got {area_counts['Area 3']}"
            assert area_counts["Area 4"] == 7, f"Area 4 expected 7, got {area_counts['Area 4']}"
            assert area_counts["Area 5"] == 5, f"Area 5 expected 5, got {area_counts['Area 5']}"

@pytest.mark.asyncio
async def test_searchability_of_mandatory_queries():
    """Verify mandatory queries from ORIGINAL_REQUEST.md ('DSCons rule', 'frontend theme', 'backend architecture')."""
    mandatory_tests = [
        {
            "query": "DSCons rule",
            "expected_keywords": ["DSCons", "quy chuẩn", "bất biến", "luật", "rule", "lục giác", "tiêu chuẩn"],
            "forbidden_keywords": ["app/core/logging_config", "FreeExile", "Savage"]
        },
        {
            "query": "frontend theme",
            "expected_keywords": ["Dark Slate", "0b0f19", "Inter", "JetBrains Mono", "Bloomberg"],
            "forbidden_keywords": ["Feral NPC", "Dialogue HUD", "Atlas"]
        },
        {
            "query": "backend architecture",
            "expected_keywords": ["Clean Architecture", "NUMERIC(18, 4)", "Decimal", "Ports", "4 tầng"],
            "forbidden_keywords": ["FreeExile", "Server-Authoritative", "Cocos"]
        }
    ]

    async with sse_client("http://127.0.0.1:8765/sse") as (read, write):
        async with ClientSession(read, write) as session:
            await session.initialize()
            
            for test in mandatory_tests:
                res = await session.call_tool("mem0_search", {"query": test["query"], "limit": 5, "project": "DSCons"})
                data = json.loads(res.content[0].text)
                results = data.get("memories", {}).get("results", [])
                assert len(results) > 0, f"Query '{test['query']}' returned zero results"
                
                # Check top hit has high cosine score
                top_hit = results[0]
                top_score = top_hit.get("score", 0.0)
                assert top_score >= 0.55, f"Query '{test['query']}' top score {top_score} < 0.55 threshold"
                
                # Verify all returned items belong to DSCons
                for r in results:
                    meta = r.get("metadata") or {}
                    assert meta.get("project") == "DSCons" or "dscons" in r.get("memory", "").lower()
                
                # Verify expected keywords in top results
                top_texts = " ".join([r.get("memory", "") for r in results[:3]])
                has_expected = any(k.lower() in top_texts.lower() for k in test["expected_keywords"])
                assert has_expected, f"Query '{test['query']}' top results missing expected keywords: {top_texts}"
                
                # Verify zero forbidden obsolete remnants
                for r in results:
                    text = r.get("memory", "").lower()
                    for f in test["forbidden_keywords"]:
                        assert f.lower() not in text, f"Query '{test['query']}' contains forbidden remnant '{f}'"

@pytest.mark.asyncio
async def test_searchability_of_domain_queries():
    """Verify specialized domain queries return appropriate newly ingested facts with score >= 0.55."""
    domain_tests = [
        {
            "query": "AEC pricing norms",
            "expected_keywords": ["Thông tư 38", "cừ Larsen", "VL = 0", "MR / T", "Định Sơn", "Định mức"],
            "forbidden_keywords": ["4.753 tỷ", "cache 71"]
        },
        {
            "query": "CAD takeoff geometry",
            "expected_keywords": ["TCVN3", "stroke-width: 1px", "Bounding Box", "takeoff", "CAD"],
            "forbidden_keywords": ["FreeExile"]
        },
        {
            "query": "database precision",
            "expected_keywords": ["NUMERIC(18, 4)", "Decimal", "Double-entry", "Sổ kép"],
            "forbidden_keywords": []
        },
        {
            "query": "testing workflow",
            "expected_keywords": ["Deep Matrix", "D1-D6", "TDD", "task.md", "Verification"],
            "forbidden_keywords": []
        },
        {
            "query": "MCP tools protocol",
            "expected_keywords": ["lsp-mcp", "chrome-devtools-mcp", "markitdown", "MCP"],
            "forbidden_keywords": ["grep_search"]
        }
    ]

    async with sse_client("http://127.0.0.1:8765/sse") as (read, write):
        async with ClientSession(read, write) as session:
            await session.initialize()
            
            for test in domain_tests:
                res = await session.call_tool("mem0_search", {"query": test["query"], "limit": 5, "project": "DSCons"})
                data = json.loads(res.content[0].text)
                results = data.get("memories", {}).get("results", [])
                assert len(results) > 0, f"Query '{test['query']}' returned zero results"
                
                top_hit = results[0]
                top_score = top_hit.get("score", 0.0)
                assert top_score >= 0.55, f"Query '{test['query']}' top score {top_score} < 0.55 threshold"
                
                top_texts = " ".join([r.get("memory", "") for r in results[:3]])
                has_expected = any(k.lower() in top_texts.lower() for k in test["expected_keywords"])
                assert has_expected, f"Query '{test['query']}' top results missing expected keywords: {top_texts}"
                
                for r in results:
                    text = r.get("memory", "").lower()
                    for f in test["forbidden_keywords"]:
                        assert f.lower() not in text, f"Query '{test['query']}' contains forbidden remnant '{f}'"

@pytest.mark.asyncio
async def test_preservation_of_non_dscons_memories():
    """Verify non-DSCons memories (developer preferences, VoLamWeb rules) remain intact."""
    async with sse_client("http://127.0.0.1:8765/sse") as (read, write):
        async with ClientSession(read, write) as session:
            await session.initialize()
            
            # Check developer preference
            res_dev = await session.call_tool("mem0_search", {"query": "Developer prefers Python and TypeScript", "limit": 5})
            data_dev = json.loads(res_dev.content[0].text)
            results_dev = data_dev.get("memories", {}).get("results", [])
            assert len(results_dev) > 0
            matching_dev = [r for r in results_dev if "typescript" in r.get("memory", "").lower()]
            assert len(matching_dev) > 0
            
            # Check VoLamWeb rule
            res_vlw = await session.call_tool("mem0_search", {"query": "strict 0 VND budget Google Cloud", "limit": 5})
            data_vlw = json.loads(res_vlw.content[0].text)
            results_vlw = data_vlw.get("memories", {}).get("results", [])
            assert len(results_vlw) > 0
            matching_vlw = [r for r in results_vlw if "0 vnd budget" in r.get("memory", "").lower()]
            assert len(matching_vlw) > 0
