"""
Term Index for fast domain term lookup.

Provides Redis-backed caching with local memory fallback for fast
term matching during query processing.
"""

import json
import os
from typing import Optional, Set

import structlog
import redis.asyncio as redis

from src.core.config import settings

logger = structlog.get_logger(__name__)

# Configuration from environment
CACHE_TTL = int(os.getenv("DOMAIN_TERM_CACHE_TTL", "3600"))  # 1 hour default
REDIS_KEY_PREFIX = "domain_terms:"


class TermIndex:
    """In-memory and Redis-backed index for fast term lookups.

    Provides two-level caching:
    1. Local memory cache (process-level, instant access)
    2. Redis cache (shared across processes, fast access)

    Used by QueryTermMatcher to quickly check if query contains
    indexed domain terms.
    """

    def __init__(self, redis_client: Optional[redis.Redis] = None):
        """Initialize term index.

        Args:
            redis_client: Optional Redis client. If None, creates one from settings.
        """
        self._redis_client = redis_client
        self._redis_initialized = redis_client is not None
        self._local_cache: dict[str, Set[str]] = {}
        self.logger = structlog.get_logger(__name__)

    async def _get_redis(self) -> Optional[redis.Redis]:
        """Get or create Redis client lazily.

        Returns None if Redis connection fails (graceful degradation).
        """
        if not self._redis_initialized:
            try:
                self._redis_client = await redis.from_url(
                    settings.redis_url,
                    encoding="utf-8",
                    decode_responses=True,
                    socket_connect_timeout=2,  # Fail fast if Redis unavailable
                )
                self._redis_initialized = True
                self.logger.info("redis_connected", url=settings.redis_url)
            except Exception as e:
                self.logger.warning(
                    "redis_connection_failed",
                    error=str(e),
                    fallback="local_cache_only"
                )
                return None

        return self._redis_client

    async def load_tenant_terms(self, tenant_id: str) -> Set[str]:
        """Load all domain terms for a tenant.

        Checks local cache first, then Redis, with graceful degradation
        to empty set if both fail.

        Args:
            tenant_id: Tenant identifier

        Returns:
            Set of all domain terms for the tenant
        """
        # Check local cache first
        if tenant_id in self._local_cache:
            self.logger.debug("term_index_hit_local", tenant_id=tenant_id)
            return self._local_cache[tenant_id]

        # Try Redis
        redis_client = await self._get_redis()
        if redis_client:
            try:
                key = f"{REDIS_KEY_PREFIX}{tenant_id}"
                cached_data = await redis_client.get(key)

                if cached_data:
                    terms = set(json.loads(cached_data))
                    # Update local cache
                    self._local_cache[tenant_id] = terms
                    self.logger.debug(
                        "term_index_hit_redis",
                        tenant_id=tenant_id,
                        term_count=len(terms)
                    )
                    return terms

            except Exception as e:
                self.logger.warning(
                    "redis_get_failed",
                    tenant_id=tenant_id,
                    error=str(e)
                )

        # No cached data found - return empty set (will be populated on first query)
        self.logger.debug("term_index_miss", tenant_id=tenant_id)
        return set()

    async def update_tenant_terms(self, tenant_id: str, terms: Set[str]) -> None:
        """Update term index for a tenant.

        Updates both local and Redis caches.

        Args:
            tenant_id: Tenant identifier
            terms: Set of domain terms to index
        """
        # Update local cache
        self._local_cache[tenant_id] = terms

        # Update Redis cache
        redis_client = await self._get_redis()
        if redis_client:
            try:
                key = f"{REDIS_KEY_PREFIX}{tenant_id}"
                serialized = json.dumps(list(terms))
                await redis_client.setex(key, CACHE_TTL, serialized)

                self.logger.info(
                    "term_index_updated",
                    tenant_id=tenant_id,
                    term_count=len(terms),
                    ttl=CACHE_TTL
                )

            except Exception as e:
                self.logger.warning(
                    "redis_set_failed",
                    tenant_id=tenant_id,
                    error=str(e),
                    fallback="local_cache_only"
                )

    async def add_terms(self, tenant_id: str, new_terms: Set[str]) -> None:
        """Add new terms to existing index for a tenant.

        Args:
            tenant_id: Tenant identifier
            new_terms: Set of new terms to add
        """
        existing_terms = await self.load_tenant_terms(tenant_id)
        updated_terms = existing_terms | new_terms
        await self.update_tenant_terms(tenant_id, updated_terms)

    async def clear_tenant_cache(self, tenant_id: str) -> None:
        """Clear cached terms for a tenant.

        Args:
            tenant_id: Tenant identifier
        """
        # Clear local cache
        if tenant_id in self._local_cache:
            del self._local_cache[tenant_id]

        # Clear Redis cache
        redis_client = await self._get_redis()
        if redis_client:
            try:
                key = f"{REDIS_KEY_PREFIX}{tenant_id}"
                await redis_client.delete(key)

                self.logger.info("term_cache_cleared", tenant_id=tenant_id)

            except Exception as e:
                self.logger.warning(
                    "redis_delete_failed",
                    tenant_id=tenant_id,
                    error=str(e)
                )

    def clear_local_cache(self) -> None:
        """Clear all local cached terms (process-level)."""
        self._local_cache.clear()
        self.logger.info("local_cache_cleared")

    async def close(self) -> None:
        """Close Redis connection if initialized."""
        if self._redis_client:
            await self._redis_client.close()
            self.logger.info("redis_connection_closed")
