"""
Core search engine using BM25 ranking with regex matching
"""

import csv
import re
import math
from pathlib import Path
from typing import List, Dict, Any, Optional
from collections import defaultdict

# Domain to CSV file mapping
DOMAINS = {
    'clarity': 'clarity-syntax.csv',
    'templates': 'contract-templates.csv',
    'security': 'security-patterns.csv',
    'defi': 'defi-protocols.csv',
    'stacksjs': 'stacks-js-core.csv',
    'bns': 'bns.csv',
    'stacking': 'stacking.csv',
    'deployment': 'deployment.csv',
    'nfts': 'nfts.csv',
    'tokens': 'fungible-tokens.csv',
    'auth': 'authentication.csv',
    'advanced': 'advanced-patterns.csv',
    'chainhooks': 'chainhooks.csv',
    'trading': 'trading-bots.csv',
    'oracles': 'oracles.csv'
}

# Map domain keys to CSV basenames (for relationships)
def get_relationship_domain(domain: str) -> str:
    """Convert domain key to CSV basename for relationship lookups"""
    csv_file = DOMAINS.get(domain, '')
    return csv_file.replace('.csv', '') if csv_file else domain

# Auto-detection keywords
DOMAIN_KEYWORDS = {
    'clarity': ['define', 'uint', 'principal', 'let', 'begin', 'asserts', 'unwrap', 'function', 'type', 'syntax'],
    'templates': ['token', 'nft', 'sip-010', 'sip-009', 'vault', 'dao', 'template', 'contract'],
    'security': ['vulnerability', 'security', 'audit', 'attack', 'safe', 'check', 'exploit'],
    'defi': ['swap', 'pool', 'liquidity', 'alex', 'velar', 'bitflow', 'zest', 'borrow', 'lend', 'boost'],
    'stacksjs': ['javascript', 'stacks.js', 'connect', 'wallet', 'frontend', 'react', 'typescript'],
    'bns': ['bns', 'name', 'domain', '.btc', 'resolve', 'register'],
    'stacking': ['stacking', 'pox', 'delegate', 'pool', 'reward', 'cycle', 'bitcoin'],
    'deployment': ['deploy', 'stx_deployContract', 'contract', 'mainnet', 'testnet', 'clarinet', 'devnet', 'publish'],
    'nfts': ['nft', 'sip-009', 'mint', 'transfer', 'metadata', 'collection', 'marketplace', 'gamma'],
    'tokens': ['sip-010', 'fungible', 'token', 'balance', 'allowance', 'transfer-from', 'decimals'],
    'auth': ['jwt', 'middleware', 'encrypt', 'decrypt', 'access-control', 'token-gate', 'nft-gate', 'protected-route', 'verify-signature'],
    'advanced': ['pagination', 'swr', 'activity', 'presale', 'lottery', 'portfolio', 'wizard', 'launchpad', 'template', 'milestone', 'vesting', 'export', 'csv', 'modal', 'confirmation'],
    'chainhooks': ['chainhook', 'webhook', 'txid', 'contract-call', 'print-event', 'ft-event', 'nft-event', 'stx-event', 'predicate', 'hiro', 'indexing', 'ordinals'],
    'trading': ['trading', 'bot', 'automated', 'wallet-sdk', 'privy', 'buy', 'sell', 'swap', 'dex', 'post-condition', 'broadcast', 'signature', 'recovery', 'txOptions', 'makeContractCall', 'bonding'],
    'oracles': ['oracle', 'pyth', 'price', 'feed', 'vaa', 'hermes', 'wormhole', 'verify', 'update', 'decode']
}


class BM25:
    """BM25 ranking algorithm"""

    def __init__(self, documents: List[str], k1: float = 1.5, b: float = 0.75):
        self.k1 = k1
        self.b = b
        self.documents = documents
        self.doc_len = [len(doc.split()) for doc in documents]
        self.avgdl = sum(self.doc_len) / len(documents) if documents else 0
        self.doc_freqs = self._calc_doc_freqs()
        self.idf = self._calc_idf()

    def _calc_doc_freqs(self) -> Dict[str, int]:
        freqs = defaultdict(int)
        for doc in self.documents:
            for term in set(doc.lower().split()):
                freqs[term] += 1
        return freqs

    def _calc_idf(self) -> Dict[str, float]:
        idf = {}
        n = len(self.documents)
        for term, df in self.doc_freqs.items():
            idf[term] = math.log((n - df + 0.5) / (df + 0.5) + 1)
        return idf

    def score(self, query: str, doc_idx: int) -> float:
        score = 0.0
        doc = self.documents[doc_idx].lower()
        doc_terms = doc.split()
        term_freqs = defaultdict(int)
        for term in doc_terms:
            term_freqs[term] += 1

        for term in query.lower().split():
            if term not in self.idf:
                continue
            tf = term_freqs[term]
            idf = self.idf[term]
            dl = self.doc_len[doc_idx]
            score += idf * (tf * (self.k1 + 1)) / (tf + self.k1 * (1 - self.b + self.b * dl / self.avgdl))

        return score


def detect_domain(query: str) -> str:
    """Auto-detect domain based on query keywords"""
    query_lower = query.lower()
    scores = {}

    for domain, keywords in DOMAIN_KEYWORDS.items():
        score = sum(1 for kw in keywords if kw in query_lower)
        if score > 0:
            scores[domain] = score

    if scores:
        return max(scores, key=scores.get)
    return 'templates'  # Default to templates


def load_data(domain: str) -> List[Dict[str, Any]]:
    """Load CSV data for a domain"""
    data_dir = Path(__file__).parent.parent / 'data'
    csv_file = data_dir / DOMAINS[domain]

    if not csv_file.exists():
        return []

    with open(csv_file, 'r', encoding='utf-8') as f:
        reader = csv.DictReader(f)
        return list(reader)


def normalize_search_query(query: str) -> str:
    """
    Normalize search query to handle common variations and abbreviations

    Args:
        query: Original search query

    Returns:
        Expanded query with common variations
    """
    import re

    query_lower = query.lower()
    expanded_terms = [query]  # Always include original query

    # SIP standard variations
    sip_patterns = [
        (r'\bsip10\b', ['sip010', 'sip-010', 'fungible', 'token']),
        (r'\bsip9\b', ['sip009', 'sip-009', 'nft']),
        (r'\bsip013\b', ['sip-013', 'transfer-memo']),
        (r'\bsip016\b', ['sip-016', 'token-metadata']),
    ]

    for pattern, expansions in sip_patterns:
        if re.search(pattern, query_lower):
            expanded_terms.extend(expansions)

    # Common abbreviations
    abbreviations = {
        'ft': ['fungible-token', 'sip010'],
        'nft': ['non-fungible-token', 'sip009'],
        'pc': ['post-condition', 'postcondition'],
    }

    words = query_lower.split()
    for word in words:
        if word in abbreviations:
            expanded_terms.extend(abbreviations[word])

    return ' '.join(expanded_terms)


def search(
    query: str,
    domain: str = 'auto',
    max_results: int = 5,
    include_relationships: bool = True,
    min_relationship_strength: int = 7
) -> List[Dict[str, Any]]:
    """
    Search knowledge base using BM25 + regex with optional relationship enrichment

    Args:
        query: Search query
        domain: Domain to search or 'auto' for auto-detection
        max_results: Maximum results to return
        include_relationships: Whether to include related entries (default: True)
        min_relationship_strength: Minimum strength for relationships (1-10, default: 7)

    Returns:
        List of matching records with optional relationship data
    """
    # Normalize query to handle common variations
    normalized_query = normalize_search_query(query)
    # Auto-detect domain if needed
    if domain == 'auto':
        domain = detect_domain(query)

    # Load data
    data = load_data(domain)
    if not data:
        return []

    # Create searchable text for each record
    def record_to_text(record: Dict) -> str:
        return ' '.join(str(v) for v in record.values())

    documents = [record_to_text(r) for r in data]

    # BM25 scoring with normalized query
    bm25 = BM25(documents)
    scores = [(i, bm25.score(normalized_query, i)) for i in range(len(documents))]

    # Regex boost for exact matches (using original query for exact matching)
    try:
        query_pattern = re.compile(re.escape(query), re.IGNORECASE)
        for i, (idx, score) in enumerate(scores):
            if query_pattern.search(documents[idx]):
                scores[i] = (idx, score * 2)  # Boost exact matches
    except re.error:
        pass  # Skip regex boost if pattern is invalid

    # Sort by score and return top results
    scores.sort(key=lambda x: x[1], reverse=True)

    results = []
    for idx, score in scores[:max_results]:
        if score > 0:
            result = data[idx].copy()
            result['_score'] = round(score, 3)
            result['_domain'] = domain
            results.append(result)

    # Add relationships if requested
    if include_relationships and results:
        try:
            from relationships import get_graph
            graph = get_graph()

            # Convert domain key to CSV basename for relationship lookups
            rel_domain = get_relationship_domain(domain)

            for result in results:
                entry_id = result.get('id', '')
                if entry_id:
                    # Get related entries (forward direction, strong relationships only)
                    related = graph.get_related(
                        domain=rel_domain,
                        entry_id=entry_id,
                        direction='both',  # Include both forward and reverse
                        min_strength=min_relationship_strength
                    )

                    # Format relationships for display
                    if related:
                        result['_relationships'] = [
                            {
                                'type': rel.relationship_type,
                                'target_domain': rel.target_domain,
                                'target_id': rel.target_id,
                                'strength': rel.strength,
                                'context': rel.context
                            }
                            for rel in related
                        ]
        except (ImportError, FileNotFoundError):
            # Relationships module not available or CSV not found, skip silently
            pass

    return results
