"""Forensic Auditor Investigation Script for Milestone 3.

Exhaustively verifies:
1. Direct inspection of Qdrant storage.sqlite (read-only mode)
2. Live FastMCP service verification via http://127.0.0.1:8765/sse
3. 31 purged UUIDs verification (none exist)
4. 27 ingested atomic facts verification (all exist, non-zero vectors, 5 areas)
5. 449 preserved non-DSCons records verification (FreeExile, VoLamWeb, dev preferences, etc.)
6. Total point count math: 449 + 27 = 476.
"""

import os
import sys
import json
import sqlite3
import pickle
import asyncio
from typing import Any

# Configure utf-8 stdout
if hasattr(sys.stdout, "reconfigure"):
    sys.stdout.reconfigure(encoding="utf-8", errors="replace")

SQLITE_PATH = r"C:\Users\Admin\.gemini\antigravity\mem0_data\qdrant\collection\antigravity_memories\storage.sqlite"

PURGED_31_UUIDS = [
    "8f784c03-12c2-4e2e-ad27-c9126339a397",
    "56c96009-a86d-4f77-9741-d4bf2db91186",
    "3fe6685d-0dc2-4620-9fc3-d4dae6d1a827",
    "3ae63e6e-ec8c-413d-9ed3-7309719bbb1a",
    "710970c5-f250-49c3-a45c-e3dafc03234d",
    "1c80e96d-1953-46cf-bd25-e886176c294c",
    "5c85fded-a035-4580-8c3e-c7e6db87ccfe",
    "735d2967-c794-43e3-ba2d-91336fca8c05",
    "777fa33d-3de1-4b8b-8edb-2ac3249d4dd1",
    "8ce4fbab-615f-49d5-a170-8a68c0006b58",
    "c8c9e303-56ef-40b7-b540-75e16ce7fba2",
    "df49c0cd-cc16-4243-a89d-4e26f0bb1ed4",
    "e41cb24f-c679-45d6-bd73-08504a4d5df7",
    "ffaa5396-8d31-44de-8d7e-330686bb7253",
    "9a3be66a-538b-43cd-afe5-f9c99eeeb7d9",
    "af959e89-7f6b-42cc-b70c-b312d568164d",
    "05dfee8c-9873-4f58-8c02-e9356267b1da",
    "3172e150-bfdf-4e1b-9140-ef9c2131e254",
    "a6c7ff66-981f-4922-b5d7-ebc00cf08342",
    "a1505cb1-23f7-4e80-a7ab-0ef1402e58f0",
    "88388e27-9b76-4a8f-a2a8-88552898ec10",
    "adebc49b-b740-4c8d-9719-e07598e283b8",
    "ef2d2b14-85fe-497d-973d-03f856445094",
    "385dc9bd-a317-45fc-a8a5-c8fab236e31b",
    "38dfd0cf-7359-432e-9c07-fbdbee4721eb",
    "617c66ff-39a1-46e0-b422-9c87ad4b7d78",
    "d4d86547-0ebb-4fd4-9fec-fc3665f300d0",
    "0df1595a-9ef5-4012-ae6f-82808dba93e1",
    "5dcc1be8-7216-47a3-8285-dc5b81ffe324",
    "305689dc-261a-4cbc-831a-54bfd70eae77",
    "bab3dbc0-2f3c-49ad-b416-a9f4f8077dbf",
]

EXPECTED_27_CODES = [
    "RULE-ROOT-01", "RULE-ROOT-02", "RULE-ROOT-03", "RULE-ROOT-04", "RULE-ROOT-05",
    "RULE-BACKEND-01", "RULE-BACKEND-02", "RULE-BACKEND-03", "RULE-BACKEND-04", "RULE-BACKEND-05",
    "RULE-FRONTEND-01", "RULE-FRONTEND-02", "RULE-FRONTEND-03", "RULE-FRONTEND-04", "RULE-FRONTEND-05",
    "RULE-AEC-01", "RULE-AEC-02", "RULE-AEC-03", "RULE-AEC-04", "RULE-AEC-05", "RULE-AEC-06", "RULE-AEC-07",
    "RULE-WORKFLOW-01", "RULE-WORKFLOW-02", "RULE-WORKFLOW-03", "RULE-WORKFLOW-04", "RULE-WORKFLOW-05"
]

def audit_sqlite_storage():
    print("=" * 70)
    print("PHASE 1: DIRECT SQLITE STORAGE AUDIT")
    print("=" * 70)
    uri = f"file:{SQLITE_PATH.replace(os.sep, '/')}?mode=ro"
    con = sqlite3.connect(uri, uri=True)
    cur = con.cursor()

    tables = [r[0] for r in cur.execute("SELECT name FROM sqlite_master WHERE type='table'").fetchall()]
    print(f"[*] SQLite tables found: {tables}")

    for t in tables:
        count = cur.execute(f"SELECT COUNT(*) FROM {t}").fetchone()[0]
        print(f"    - Table '{t}': {count} rows")

    rows = cur.execute("SELECT id, point FROM points").fetchall()
    total_points = len(rows)
    print(f"[*] Total points fetched from SQLite: {total_points}")
    assert total_points == 476, f"VIOLATION: Expected exactly 476 points, got {total_points}"

    # Unpickle all PointStruct objects
    points = []
    point_ids = set()
    for row_id, blob in rows:
        pt = pickle.loads(blob)
        points.append(pt)
        point_ids.add(str(pt.id))

    # 1. Verify 31 Purged UUIDs
    found_purged = []
    for p_uuid in PURGED_31_UUIDS:
        if p_uuid in point_ids or any(p_uuid.lower() == pid.lower() for pid in point_ids):
            found_purged.append(p_uuid)
    print(f"[*] Probing 31 purged UUIDs in SQLite storage: {len(found_purged)} found (Expected: 0)")
    assert len(found_purged) == 0, f"VIOLATION: Purged UUIDs found in SQLite: {found_purged}"
    print(f"[+] VERIFIED: All 31 obsolete DSCons UUIDs are 100% gone from Qdrant storage.")

    # 2. Analyze DSCons vs Non-DSCons points
    dscons_facts = {}
    non_dscons_points = []
    zero_vector_count = 0

    freeexile_count = 0
    volamweb_count = 0
    dev_pref_count = 0
    global_directives_count = 0
    prompt_headers_count = 0
    other_count = 0

    user_id_counts = {}

    for pt in points:
        payload = pt.payload or {}
        uid = payload.get("user_id")
        user_id_counts[uid] = user_id_counts.get(uid, 0) + 1

        # Vector validation (focus on dense vector '')
        vec = pt.vector
        is_non_zero_vector = False
        dense_dim = 0
        has_sparse = False
        sparse_len = 0

        if isinstance(vec, dict):
            dense_vec = vec.get("")
            if isinstance(dense_vec, (list, tuple)):
                dense_dim = len(dense_vec)
                if any(abs(x) > 1e-6 for x in dense_vec):
                    is_non_zero_vector = True
            sparse_vec = vec.get("bm25")
            if sparse_vec and hasattr(sparse_vec, "indices"):
                has_sparse = True
                sparse_len = len(sparse_vec.indices)
        elif isinstance(vec, (list, tuple)):
            dense_dim = len(vec)
            if any(abs(x) > 1e-6 for x in vec):
                is_non_zero_vector = True
        
        if not is_non_zero_vector:
            zero_vector_count += 1

        text = str(payload.get("data", ""))
        
        # In Mem0, metadata can be top-level or inside 'metadata' dict
        meta_dict = payload.get("metadata") or {}
        proj = str(payload.get("project") or meta_dict.get("project") or "").strip()
        code = payload.get("code") or meta_dict.get("code")
        area = payload.get("area") or meta_dict.get("area")
        category = payload.get("category") or meta_dict.get("category")

        if (proj == "DSCons" and code in EXPECTED_27_CODES) or (code in EXPECTED_27_CODES):
            dscons_facts[code] = {
                "id": str(pt.id),
                "code": code,
                "area": area,
                "category": category,
                "text": text[:80],
                "dense_dim": dense_dim,
                "has_sparse": has_sparse,
                "sparse_len": sparse_len,
                "is_non_zero_vector": is_non_zero_vector
            }
        else:
            non_dscons_points.append(pt)
            t_lower = text.lower()
            if "freeexile" in t_lower or "atlas" in t_lower or "feral" in t_lower or "affliction" in t_lower or "chaos orb" in t_lower:
                freeexile_count += 1
            elif "volamweb" in t_lower or "0 vnd budget" in t_lower:
                volamweb_count += 1
            elif "developer prefers" in t_lower or "typescript cho frontend" in t_lower or "python for scalable" in t_lower:
                dev_pref_count += 1
            elif "antigravity" in t_lower or "directive" in t_lower or "fastmcp" in t_lower or "open code review" in t_lower:
                global_directives_count += 1
            elif text.startswith("## ") or text.startswith("# ") or "2026-10" in text:
                prompt_headers_count += 1
            else:
                other_count += 1

    print(f"[*] User ID distribution in SQLite: {user_id_counts}")
    print(f"[*] Zero-vector points detected: {zero_vector_count} (Expected: 0)")
    assert zero_vector_count == 0, f"VIOLATION: Found {zero_vector_count} points with zero vectors"

    # 3. Verify 27 Ingested Facts
    print(f"[*] DSCons facts found in SQLite: {len(dscons_facts)} / 27")
    missing_codes = set(EXPECTED_27_CODES) - set(dscons_facts.keys())
    print(f"[*] Missing codes: {missing_codes} (Expected: empty set)")
    assert len(missing_codes) == 0, f"VIOLATION: Missing codes in SQLite: {missing_codes}"
    assert len(dscons_facts) == 27, f"VIOLATION: Expected 27 DSCons facts, got {len(dscons_facts)}"

    for code, info in dscons_facts.items():
        assert info["is_non_zero_vector"], f"VIOLATION: Fact {code} has zero vector!"
        assert info["dense_dim"] == 768, f"VIOLATION: Fact {code} dense vector dim is {info['dense_dim']}, expected 768"
        assert info["has_sparse"], f"VIOLATION: Fact {code} missing BM25 sparse vector"
    print(f"[+] VERIFIED: All 27 DSCons facts have authentic 768-dim dense vectors + BM25 sparse vectors.")

    # 4. Verify 449 Preserved Records
    non_dscons_count = len(non_dscons_points)
    print(f"[*] Non-DSCons preserved records: {non_dscons_count} (Expected: 449)")
    print(f"    - FreeExile records: {freeexile_count}")
    print(f"    - VoLamWeb records: {volamweb_count}")
    print(f"    - Developer preferences: {dev_pref_count}")
    print(f"    - Global directives: {global_directives_count}")
    print(f"    - Prompt headers / templates: {prompt_headers_count}")
    print(f"    - Other preserved records: {other_count}")
    assert non_dscons_count == 449, f"VIOLATION: Preserved count {non_dscons_count} != 449"

    # 5. Mathematical Integrity Check
    print(f"[*] Mathematical Equation: {non_dscons_count} (preserved) + {len(dscons_facts)} (ingested) = {total_points} (total)")
    assert non_dscons_count + len(dscons_facts) == total_points == 476, "VIOLATION: Arithmetic integrity check failed"
    print("[+] PHASE 1 PASSED: Qdrant SQLite storage is 100% mathematically consistent & authentic.\n")


async def audit_live_fastmcp():
    print("=" * 70)
    print("PHASE 2: LIVE FASTMCP SSE PROTOCOL & SEMANTIC AUDIT")
    print("=" * 70)
    from mcp.client.sse import sse_client
    from mcp.client.session import ClientSession

    url = "http://127.0.0.1:8765/sse"
    async with sse_client(url) as (read, write):
        async with ClientSession(read, write) as session:
            await session.initialize()
            print("[*] Successfully connected to live FastMCP service via SSE.")

            # 1. Check status
            status_res = await session.call_tool("mem0_status", {})
            status_data = json.loads(status_res.content[0].text)
            print(f"[*] mem0_status response:\n{json.dumps(status_data, indent=2)}")
            assert status_data.get("status") == "online", "VIOLATION: Status is not online"
            assert status_data.get("is_ready") is True, "VIOLATION: Service is not ready"
            assert status_data.get("developer_memories_count") == 476, f"VIOLATION: Count is {status_data.get('developer_memories_count')}, expected 476"

            # 2. Check get_all across user IDs (472 developer + 2 Admin + 1 ceo_architect + 1 admin = 476)
            total_retrieved = 0
            all_mems = []
            for uid in ["developer", "Admin", "ceo_architect", "admin"]:
                res = await session.call_tool("mem0_get_all", {"user_id": uid, "limit": 1000})
                res_data = json.loads(res.content[0].text)
                m_list = res_data.get("memories", {}).get("results", [])
                print(f"    - mem0_get_all(user_id='{uid}'): returned {len(m_list)} memories")
                total_retrieved += len(m_list)
                all_mems.extend(m_list)

            print(f"[*] Total retrieved memories across all user IDs via FastMCP: {total_retrieved} (Expected: 476)")
            assert total_retrieved == 476, f"VIOLATION: Total retrieved {total_retrieved} != 476"

            # DSCons facts in get_all
            dscons_mems = [m for m in all_mems if (m.get("metadata") or {}).get("project") == "DSCons" or (m.get("metadata") or {}).get("code") in EXPECTED_27_CODES]
            print(f"[*] DSCons memories identified via get_all: {len(dscons_mems)} (Expected: 27)")
            assert len(dscons_mems) == 27, f"VIOLATION: Expected 27 DSCons memories, got {len(dscons_mems)}"

            found_codes = { (m.get("metadata") or {}).get("code") for m in dscons_mems }
            missing_codes = set(EXPECTED_27_CODES) - found_codes
            assert len(missing_codes) == 0, f"VIOLATION: Missing codes in get_all: {missing_codes}"
            print(f"[+] All 27 DSCons facts confirmed uniquely retrievable via mem0_get_all.")

            # 3. Check 31 purged UUIDs via mem0_delete
            print("[*] Probing all 31 purged UUIDs over live SSE to confirm deletion permanence...")
            for uid in PURGED_31_UUIDS:
                del_res = await session.call_tool("mem0_delete", {"memory_id": uid})
                resp_text = del_res.content[0].text.lower()
                assert "not found" in resp_text or "error" in resp_text, f"VIOLATION: UUID {uid} still exists!"
            print(f"[+] All 31 UUIDs confirmed 100% absent (permanently expunged).")

            # 4. Mandatory & Domain Semantic Searches
            queries = [
                ("DSCons rule", ["DSCons", "quy chuẩn", "bất biến", "luật", "rule", "lục giác", "tiêu chuẩn"]),
                ("frontend theme", ["Dark Slate", "0b0f19", "Inter", "JetBrains Mono", "Bloomberg"]),
                ("backend architecture", ["Clean Architecture", "NUMERIC(18, 4)", "Decimal", "Ports", "4 tầng"]),
                ("AEC pricing norms", ["Thông tư 38", "cừ Larsen", "VL = 0", "MR / T", "Định Sơn", "Định mức"]),
                ("CAD takeoff geometry", ["TCVN3", "stroke-width: 1px", "Bounding Box", "takeoff", "CAD"]),
                ("database precision", ["NUMERIC(18, 4)", "Decimal", "Double-entry", "Sổ kép"]),
                ("testing workflow", ["Deep Matrix", "D1-D6", "TDD", "task.md", "Verification"]),
                ("MCP tools protocol", ["lsp-mcp", "chrome-devtools-mcp", "markitdown", "MCP"])
            ]

            print("[*] Testing 8 semantic queries against live Mem0 service:")
            for q, expected_kws in queries:
                s_res = await session.call_tool("mem0_search", {"query": q, "limit": 5})
                s_data = json.loads(s_res.content[0].text)
                results = s_data.get("memories", {}).get("results", [])
                assert len(results) > 0, f"VIOLATION: Zero results for query '{q}'"
                top1 = results[0]
                top1_score = top1.get("score", 0.0)
                top1_text = top1.get("memory", "")
                print(f"    - Query '{q}': Top-1 Score = {top1_score:.4f} | Snippet: {top1_text[:60]}...")
                assert top1_score >= 0.55, f"VIOLATION: Top-1 score {top1_score:.4f} < 0.55 for '{q}'"

                # Check keyword
                kw_match = any(kw.lower() in top1_text.lower() for kw in expected_kws)
                assert kw_match, f"VIOLATION: No expected keyword {expected_kws} in Top-1 for '{q}'"

                # Check forbidden contaminants
                for r in results[:3]:
                    text_r = r.get("memory", "").lower()
                    assert "app/core/logging_config" not in text_r, f"VIOLATION: Legacy logging remnant in '{q}'"
                    assert "4.753 tỷ" not in text_r, f"VIOLATION: Fictitious contract remnant in '{q}'"

            # 5. Check Non-DSCons Preservation
            print("[*] Testing preservation of non-DSCons memories via live search:")
            dev_res = await session.call_tool("mem0_search", {"query": "Developer prefers Python and TypeScript", "limit": 5})
            dev_data = json.loads(dev_res.content[0].text)
            dev_results = dev_data.get("memories", {}).get("results", [])
            assert len(dev_results) > 0, "VIOLATION: Developer preferences not found"
            top_dev = dev_results[0]
            print(f"    - Developer preference: score={top_dev.get('score', 0):.4f}, memory={top_dev.get('memory')[:60]}...")
            assert "typescript" in top_dev.get("memory", "").lower() or "python" in top_dev.get("memory", "").lower()

            vlw_res = await session.call_tool("mem0_search", {"query": "strict 0 VND budget Google Cloud", "limit": 5})
            vlw_data = json.loads(vlw_res.content[0].text)
            vlw_results = vlw_data.get("memories", {}).get("results", [])
            assert len(vlw_results) > 0, "VIOLATION: VoLamWeb rules not found"
            top_vlw = vlw_results[0]
            print(f"    - VoLamWeb rule: score={top_vlw.get('score', 0):.4f}, memory={top_vlw.get('memory')[:60]}...")
            assert "0 vnd budget" in top_vlw.get("memory", "").lower()

    print("[+] PHASE 2 PASSED: Live FastMCP SSE responses are 100% genuine and verified.\n")


def main():
    audit_sqlite_storage()
    asyncio.run(audit_live_fastmcp())
    print("=" * 70)
    print("ALL FORENSIC AUDIT CHECKS PASSED: VERDICT = CLEAN")
    print("=" * 70)

if __name__ == "__main__":
    main()
