"""Service for managing domain term dictionary."""

from sqlalchemy import text
from sqlalchemy.orm import Session

from src.db.models import DomainTermDictionary


class DomainDictionaryService:
    """Service for managing domain term dictionary."""

    def __init__(
        self,
        session: Session,
        tenant_id: str | None = None,
        node_id: str | None = None,
    ):
        self.session = session
        self.tenant_id = tenant_id
        self.node_id = node_id

    def _set_rls_context(self) -> None:
        """Set RLS tenant context on the session if tenant_id is available."""
        if self.tenant_id:
            self.session.execute(text("SET LOCAL app.tenant_id = :tid"), {"tid": str(self.tenant_id)})

    def add_term(self, entry: DomainTermDictionary) -> DomainTermDictionary:
        """Add new term to dictionary."""
        try:
            self._set_rls_context()
            self.session.add(entry)
            self.session.commit()
            return entry
        except Exception:
            self.session.rollback()
            raise

    def get_all_terms(self, node_id: str | None = None) -> list[DomainTermDictionary]:
        """Get dictionary terms scoped to a node.

        Uses the explicit ``node_id`` argument when provided, else falls back
        to the service-level ``self.node_id``. Returns all terms across the
        tenant (subject to RLS) only when neither is set.
        """
        self._set_rls_context()
        scope = node_id if node_id is not None else self.node_id
        query = self.session.query(DomainTermDictionary)
        if scope is not None:
            query = query.filter(DomainTermDictionary.node_id == scope)
        return query.all()

    def get_existing_terms_set(self, node_id: str | None = None) -> set[str]:
        """Load ``term_original`` values for a node as a lowercase set for O(1) dedup.

        Scopes to the explicit ``node_id`` argument if provided, otherwise to
        ``self.node_id``. At least one must be set — a global scan across
        nodes would bypass per-node uniqueness and is not supported.
        """
        self._set_rls_context()
        scope = node_id if node_id is not None else self.node_id
        if scope is None:
            raise ValueError(
                "get_existing_terms_set requires node_id (pass as argument or set on service)"
            )
        query = self.session.query(DomainTermDictionary.term_original).filter(
            DomainTermDictionary.node_id == scope
        )
        return {row[0].lower() for row in query.all()}

    def add_terms_batch(self, entries: list[DomainTermDictionary]) -> int:
        """Add multiple terms in a single transaction.

        Uses ON CONFLICT DO NOTHING so duplicate term_original values
        are silently skipped instead of failing the entire batch.
        """
        if not entries:
            return 0
        try:
            self._set_rls_context()
            from sqlalchemy.dialects.postgresql import insert

            stmt = insert(DomainTermDictionary).values(
                [
                    {
                        "id": entry.id,
                        "term_original": entry.term_original,
                        "term_ja": entry.term_ja,
                        "term_en": entry.term_en,
                        "term_vi": entry.term_vi,
                        "source_language": entry.source_language,
                        "domain_category": entry.domain_category,
                        "source_document_id": entry.source_document_id,
                        "translation_failed": entry.translation_failed,
                        "confidence": entry.confidence,
                        "tenant_id": entry.tenant_id,
                        "node_id": entry.node_id,
                    }
                    for entry in entries
                ]
            ).on_conflict_do_nothing(index_elements=["term_original", "node_id"])
            result = self.session.execute(stmt)
            self.session.commit()
            return result.rowcount
        except Exception:
            self.session.rollback()
            raise
