"""
pg_adapter.py — Motor-compatible async PostgreSQL adapter
=========================================================
Provides a drop-in replacement for Motor (MongoDB async driver) over PostgreSQL JSONB tables.

Table schema for every collection:
    CREATE TABLE IF NOT EXISTS col_<name> (
        id TEXT PRIMARY KEY,
        data JSONB NOT NULL,
        created_at TIMESTAMPTZ DEFAULT NOW()
    );

Usage:
    from pg_adapter import init_pg
    db = await init_pg("host=127.0.0.1 port=5432 dbname=dopebling_app user=dopebling_idgapp password=Romania95!")
    # db.users.find_one({"username": "BLING"})  -- same interface as Motor
"""
import asyncio
import json
import re
import uuid
from typing import Any, Dict, List, Optional

try:
    import asyncpg
except ImportError:
    asyncpg = None  # fallback: will raise on use


# ─────────────────────────────────────────────────────────────────────────────
# Filter → SQL translator
# ─────────────────────────────────────────────────────────────────────────────

def _jsonb_path(field: str) -> str:
    """Convert dot-notation field to JSONB access expression."""
    parts = field.split(".")
    if len(parts) == 1:
        return f"data->>'{field}'"
    # nested: data->'a'->'b'->>'c'
    path = "data"
    for p in parts[:-1]:
        path += f"->'{p}'"
    path += f"->>'{parts[-1]}'"
    return path


def _jsonb_obj_path(field: str) -> str:
    """Return JSONB object path (not text cast) for type-aware ops."""
    parts = field.split(".")
    path = "data"
    for p in parts:
        path += f"->'{p}'"
    return path


def _translate_filter(filter_dict: Dict, args: List, prefix: str = "") -> str:
    """Recursively translate a MongoDB-style filter to SQL WHERE clause."""
    if not filter_dict:
        return "TRUE"

    clauses = []

    for key, value in filter_dict.items():
        if key == "$or":
            sub = " OR ".join(f"({_translate_filter(c, args)})" for c in value)
            clauses.append(f"({sub})")
        elif key == "$and":
            sub = " AND ".join(f"({_translate_filter(c, args)})" for c in value)
            clauses.append(f"({sub})")
        elif key == "$nor":
            sub = " OR ".join(f"({_translate_filter(c, args)})" for c in value)
            clauses.append(f"NOT ({sub})")
        elif key.startswith("$"):
            continue  # unknown top-level operator, skip
        elif isinstance(value, dict) and any(k.startswith("$") for k in value):
            # Operator expressions
            field_path = _jsonb_path(key)
            obj_path = _jsonb_obj_path(key)
            op_clauses = []
            for op, v in value.items():
                if op == "$eq":
                    args.append(str(v))
                    op_clauses.append(f"{field_path} = ${len(args)}")
                elif op == "$ne":
                    args.append(str(v))
                    op_clauses.append(f"({field_path} IS NULL OR {field_path} != ${len(args)})")
                elif op in ("$gt", "$gte", "$lt", "$lte"):
                    sql_op = op.replace("$gt", ">").replace("$lt", "<").replace("e", "=")
                    if isinstance(v, (int, float)):
                        args.append(str(v))
                        op_clauses.append(
                            f"({field_path} ~ '^[-+]?\\d*\\.?\\d+$' AND ({field_path})::numeric {sql_op} ${len(args)}::numeric)"
                        )
                    else:
                        args.append(str(v))
                        op_clauses.append(f"{field_path} {sql_op} ${len(args)}")
                elif op == "$in":
                    if not v:
                        op_clauses.append("FALSE")
                    else:
                        placeholders = []
                        for item in v:
                            args.append(str(item))
                            placeholders.append(f"${len(args)}")
                        op_clauses.append(f"{field_path} IN ({', '.join(placeholders)})")
                elif op == "$nin":
                    if not v:
                        op_clauses.append("TRUE")
                    else:
                        placeholders = []
                        for item in v:
                            args.append(str(item))
                            placeholders.append(f"${len(args)}")
                        op_clauses.append(f"({field_path} IS NULL OR {field_path} NOT IN ({', '.join(placeholders)}))")
                elif op == "$exists":
                    if v:
                        op_clauses.append(f"{field_path} IS NOT NULL")
                    else:
                        op_clauses.append(f"{field_path} IS NULL")
                elif op == "$regex":
                    pattern = v.replace("%", "\\%").replace("_", "\\_")
                    args.append(f"%{pattern}%")
                    op_clauses.append(f"{field_path} ILIKE ${len(args)}")
                elif op == "$not":
                    inner = _translate_filter({key: v}, args)
                    op_clauses.append(f"NOT ({inner})")
                # skip unknown operators
            if op_clauses:
                clauses.append(" AND ".join(op_clauses))
        elif value is None:
            clauses.append(f"({_jsonb_path(key)} IS NULL OR data->'{key}' = 'null'::jsonb)")
        elif isinstance(value, bool):
            args.append(json.dumps(value))
            clauses.append(f"data->'{key}' = ${len(args)}::jsonb")
        elif isinstance(value, (int, float)):
            # numeric — compare as number or string
            args.append(str(value))
            clauses.append(
                f"({_jsonb_path(key)} = ${len(args)} OR "
                f"({_jsonb_obj_path(key)})::numeric = ${len(args)}::numeric)"
            )
        else:
            args.append(str(value))
            clauses.append(f"{_jsonb_path(key)} = ${len(args)}")

    return " AND ".join(clauses) if clauses else "TRUE"


def _translate_update(update_dict: Dict, args: List) -> str:
    """Translate MongoDB update operators to SQL SET expression."""
    set_parts = []

    if "$set" in update_dict:
        for field, value in update_dict["$set"].items():
            if "." in field:
                # nested field update using jsonb_set
                parts = field.split(".")
                path = "{" + ",".join(parts) + "}"
                args.append(json.dumps(value))
                args.append(path)
                set_parts.append(f"data = jsonb_set(data, ${len(args)}::text[], ${len(args)-1}::jsonb, true)")
            else:
                args.append(json.dumps(value))
                set_parts.append(f"data = data || jsonb_build_object('{field}', ${len(args)}::jsonb)")

    if "$inc" in update_dict:
        for field, delta in update_dict["$inc"].items():
            args.append(json.dumps(delta))
            set_parts.append(
                f"data = jsonb_set(data, '{{{field}}}', "
                f"to_jsonb(COALESCE((data->>'{field}')::numeric, 0) + ${len(args)}::numeric))"
            )

    if "$push" in update_dict:
        for field, value in update_dict["$push"].items():
            args.append(json.dumps(value))
            set_parts.append(
                f"data = jsonb_set(data, '{{{field}}}', "
                f"COALESCE(data->'{field}', '[]'::jsonb) || ${len(args)}::jsonb)"
            )

    if "$pull" in update_dict:
        for field, value in update_dict["$pull"].items():
            args.append(json.dumps(value))
            set_parts.append(
                f"data = jsonb_set(data, '{{{field}}}', "
                f"(SELECT jsonb_agg(elem) FROM jsonb_array_elements(COALESCE(data->'{field}', '[]'::jsonb)) elem "
                f"WHERE elem != ${len(args)}::jsonb))"
            )

    if "$addToSet" in update_dict:
        for field, value in update_dict["$addToSet"].items():
            args.append(json.dumps(value))
            set_parts.append(
                f"data = CASE WHEN data->'{field}' @> ${len(args)}::jsonb THEN data "
                f"ELSE jsonb_set(data, '{{{field}}}', "
                f"COALESCE(data->'{field}', '[]'::jsonb) || ${len(args)}::jsonb) END"
            )

    if "$unset" in update_dict:
        for field in update_dict["$unset"]:
            set_parts.append(f"data = data - '{field}'")

    return ", ".join(set_parts) if set_parts else "data = data"


def _sort_clause(sort_spec) -> str:
    """Translate Motor sort spec to SQL ORDER BY."""
    if not sort_spec:
        return ""
    if isinstance(sort_spec, list):
        parts = []
        for field, direction in sort_spec:
            dir_str = "ASC" if direction >= 0 else "DESC"
            if field == "_id":
                parts.append(f"created_at {dir_str}")
            else:
                parts.append(f"{_jsonb_path(field)} {dir_str}")
        return "ORDER BY " + ", ".join(parts)
    elif isinstance(sort_spec, dict):
        parts = []
        for field, direction in sort_spec.items():
            dir_str = "ASC" if direction >= 0 else "DESC"
            parts.append(f"{_jsonb_path(field)} {dir_str}")
        return "ORDER BY " + ", ".join(parts)
    return ""


def _ensure_id(doc: Dict) -> Dict:
    """Ensure document has an 'id' field."""
    if "id" not in doc or not doc["id"]:
        doc["id"] = str(uuid.uuid4()).replace("-", "")
    return doc


def _row_to_doc(row) -> Optional[Dict]:
    """Convert asyncpg row to dict."""
    if row is None:
        return None
    data = dict(row["data"]) if isinstance(row["data"], dict) else json.loads(row["data"])
    return data


# ─────────────────────────────────────────────────────────────────────────────
# Cursor (for chained .sort().skip().limit().to_list())
# ─────────────────────────────────────────────────────────────────────────────

class PgCursor:
    def __init__(self, collection: "PgCollection", filter_dict: Dict):
        self._col = collection
        self._filter = filter_dict
        self._sort = None
        self._skip_n = 0
        self._limit_n = 1000

    def sort(self, key_or_list, direction=None):
        if direction is not None:
            self._sort = [(key_or_list, direction)]
        else:
            self._sort = key_or_list
        return self

    def skip(self, n: int):
        self._skip_n = n
        return self

    def limit(self, n: int):
        self._limit_n = n
        return self

    async def to_list(self, length: int = None):
        limit = length if length is not None else self._limit_n
        return await self._col._find_internal(
            self._filter, sort=self._sort, skip=self._skip_n, limit=limit
        )

    def __aiter__(self):
        return self._async_iter()

    async def _async_iter(self):
        docs = await self.to_list(length=self._limit_n)
        for doc in docs:
            yield doc


# ─────────────────────────────────────────────────────────────────────────────
# Collection
# ─────────────────────────────────────────────────────────────────────────────

class PgCollection:
    def __init__(self, pool, name: str, db: "PgDatabase"):
        self._pool = pool
        self._name = name
        self._table = f"col_{name}"
        self._db = db

    async def _ensure_table(self):
        if self._table not in self._db._initialized_tables:
            async with self._pool.acquire() as conn:
                await conn.execute(f"""
                    CREATE TABLE IF NOT EXISTS {self._table} (
                        id TEXT PRIMARY KEY,
                        data JSONB NOT NULL,
                        created_at TIMESTAMPTZ DEFAULT NOW()
                    )
                """)
            self._db._initialized_tables.add(self._table)

    async def _find_internal(self, filter_dict: Dict, sort=None, skip: int = 0, limit: int = 1000) -> List[Dict]:
        await self._ensure_table()
        args = []
        where = _translate_filter(filter_dict or {}, args)
        order = _sort_clause(sort)
        args.append(limit)
        args.append(skip)
        q = f"SELECT data FROM {self._table} WHERE {where} {order} LIMIT ${len(args)-1} OFFSET ${len(args)}"
        async with self._pool.acquire() as conn:
            rows = await conn.fetch(q, *args)
        return [_row_to_doc(r) for r in rows]

    def find(self, filter_dict: Dict = None, projection: Dict = None, *args, **kwargs):
        return PgCursor(self, filter_dict or {})

    async def find_one(self, filter_dict: Dict = None, projection: Dict = None, *args, **kwargs) -> Optional[Dict]:
        await self._ensure_table()
        args_sql = []
        where = _translate_filter(filter_dict or {}, args_sql)
        q = f"SELECT data FROM {self._table} WHERE {where} LIMIT 1"
        async with self._pool.acquire() as conn:
            row = await conn.fetchrow(q, *args_sql)
        return _row_to_doc(row) if row else None

    async def insert_one(self, doc: Dict):
        await self._ensure_table()
        doc = _ensure_id(doc)
        doc_id = doc.get("id") or doc.get("_id") or str(uuid.uuid4())
        async with self._pool.acquire() as conn:
            await conn.execute(
                f"INSERT INTO {self._table} (id, data) VALUES ($1, $2) ON CONFLICT (id) DO UPDATE SET data = EXCLUDED.data",
                str(doc_id), json.dumps(doc)
            )

        class InsertResult:
            inserted_id = doc_id
        return InsertResult()

    async def insert_many(self, docs: List[Dict]):
        await self._ensure_table()
        async with self._pool.acquire() as conn:
            async with conn.transaction():
                for doc in docs:
                    doc = _ensure_id(doc)
                    doc_id = doc.get("id") or doc.get("_id") or str(uuid.uuid4())
                    await conn.execute(
                        f"INSERT INTO {self._table} (id, data) VALUES ($1, $2) ON CONFLICT (id) DO UPDATE SET data = EXCLUDED.data",
                        str(doc_id), json.dumps(doc)
                    )

    async def update_one(self, filter_dict: Dict, update_dict: Dict, upsert: bool = False):
        await self._ensure_table()
        args = []
        where = _translate_filter(filter_dict or {}, args)
        set_clause = _translate_update(update_dict, args)
        q = f"UPDATE {self._table} SET {set_clause} WHERE id IN (SELECT id FROM {self._table} WHERE {where} LIMIT 1)"
        async with self._pool.acquire() as conn:
            result = await conn.execute(q, *args)
            if upsert and result == "UPDATE 0":
                # Upsert: build doc from filter + $set
                new_doc = {}
                if filter_dict:
                    for k, v in filter_dict.items():
                        if not k.startswith("$"):
                            new_doc[k] = v
                if "$set" in update_dict:
                    new_doc.update(update_dict["$set"])
                await self.insert_one(new_doc)

    async def update_many(self, filter_dict: Dict, update_dict: Dict):
        await self._ensure_table()
        args = []
        where = _translate_filter(filter_dict or {}, args)
        set_clause = _translate_update(update_dict, args)
        q = f"UPDATE {self._table} SET {set_clause} WHERE {where}"
        async with self._pool.acquire() as conn:
            await conn.execute(q, *args)

    async def delete_one(self, filter_dict: Dict):
        await self._ensure_table()
        args = []
        where = _translate_filter(filter_dict or {}, args)
        q = f"DELETE FROM {self._table} WHERE id IN (SELECT id FROM {self._table} WHERE {where} LIMIT 1)"
        async with self._pool.acquire() as conn:
            await conn.execute(q, *args)

    async def delete_many(self, filter_dict: Dict):
        await self._ensure_table()
        args = []
        where = _translate_filter(filter_dict or {}, args)
        q = f"DELETE FROM {self._table} WHERE {where}"
        async with self._pool.acquire() as conn:
            await conn.execute(q, *args)

    async def count_documents(self, filter_dict: Dict = None) -> int:
        await self._ensure_table()
        args = []
        where = _translate_filter(filter_dict or {}, args)
        q = f"SELECT COUNT(*) FROM {self._table} WHERE {where}"
        async with self._pool.acquire() as conn:
            row = await conn.fetchrow(q, *args)
        return row[0] if row else 0

    async def find_one_and_update(self, filter_dict: Dict, update_dict: Dict, return_document=False, upsert=False):
        await self._ensure_table()
        doc_before = await self.find_one(filter_dict)
        await self.update_one(filter_dict, update_dict, upsert=upsert)
        if return_document:
            return await self.find_one(filter_dict)
        return doc_before

    async def aggregate(self, pipeline: List[Dict]) -> List[Dict]:
        """
        Supports common pipeline stages: $match, $sort, $limit, $skip, $group (basic), $lookup (basic).
        Falls back to full-table scan + Python for complex pipelines.
        """
        await self._ensure_table()

        # Try to build SQL for simple pipelines
        match = {}
        sort = None
        skip_n = 0
        limit_n = 10000
        is_simple = True

        for stage in pipeline:
            if "$match" in stage:
                match.update(stage["$match"])
            elif "$sort" in stage:
                sort = list(stage["$sort"].items())
            elif "$skip" in stage:
                skip_n = stage["$skip"]
            elif "$limit" in stage:
                limit_n = stage["$limit"]
            elif "$group" in stage or "$lookup" in stage or "$unwind" in stage or "$project" in stage:
                is_simple = False
                break

        if is_simple:
            return await self._find_internal(match, sort=sort, skip=skip_n, limit=limit_n)

        # Complex pipeline: fetch all, apply in Python
        docs = await self._find_internal(match, sort=sort, limit=100000)
        for stage in pipeline:
            if "$sort" in stage:
                sort_fields = list(stage["$sort"].items())
                for field, direction in reversed(sort_fields):
                    docs.sort(key=lambda d: d.get(field) or "", reverse=(direction < 0))
            elif "$limit" in stage:
                docs = docs[:stage["$limit"]]
            elif "$skip" in stage:
                docs = docs[stage["$skip"]:]
            elif "$group" in stage:
                group_spec = stage["$group"]
                group_id_field = group_spec.get("_id")
                from collections import defaultdict
                groups = defaultdict(list)
                for doc in docs:
                    if isinstance(group_id_field, str) and group_id_field.startswith("$"):
                        key = doc.get(group_id_field[1:])
                    else:
                        key = group_id_field
                    groups[str(key)].append(doc)
                result = []
                for gkey, gdocs in groups.items():
                    row = {"_id": gkey}
                    for out_field, agg in group_spec.items():
                        if out_field == "_id":
                            continue
                        if isinstance(agg, dict):
                            op = list(agg.keys())[0]
                            val_field = list(agg.values())[0]
                            if isinstance(val_field, str) and val_field.startswith("$"):
                                val_field = val_field[1:]
                            if op == "$sum":
                                if val_field == 1:
                                    row[out_field] = len(gdocs)
                                else:
                                    row[out_field] = sum(d.get(val_field, 0) or 0 for d in gdocs)
                            elif op == "$avg":
                                vals = [d.get(val_field, 0) or 0 for d in gdocs]
                                row[out_field] = sum(vals) / len(vals) if vals else 0
                            elif op == "$max":
                                row[out_field] = max((d.get(val_field) for d in gdocs), default=None)
                            elif op == "$min":
                                row[out_field] = min((d.get(val_field) for d in gdocs), default=None)
                            elif op == "$first":
                                row[out_field] = gdocs[0].get(val_field) if gdocs else None
                            elif op == "$last":
                                row[out_field] = gdocs[-1].get(val_field) if gdocs else None
                            elif op == "$push":
                                row[out_field] = [d.get(val_field) for d in gdocs]
                    result.append(row)
                docs = result
        return docs

    async def create_index(self, *args, **kwargs):
        """No-op — indexes are created during migration."""
        pass

    async def drop(self):
        await self._ensure_table()
        async with self._pool.acquire() as conn:
            await conn.execute(f"DROP TABLE IF EXISTS {self._table}")
        self._db._initialized_tables.discard(self._table)


# ─────────────────────────────────────────────────────────────────────────────
# Database
# ─────────────────────────────────────────────────────────────────────────────

class PgDatabase:
    def __init__(self, pool):
        self._pool = pool
        self._initialized_tables: set = set()

    def __getattr__(self, name: str) -> PgCollection:
        if name.startswith("_"):
            raise AttributeError(name)
        return PgCollection(self._pool, name, self)

    def __getitem__(self, name: str) -> PgCollection:
        return PgCollection(self._pool, name, self)


# ─────────────────────────────────────────────────────────────────────────────
# Init
# ─────────────────────────────────────────────────────────────────────────────


async def init_pg(dsn: str) -> PgDatabase:
    """
    Initialize a PostgreSQL connection pool and return a PgDatabase instance.
    Accepts either libpq keyword format or postgresql:// URL format.
    """
    import re

    if dsn.startswith("postgresql://") or dsn.startswith("postgres://"):
        pool = await asyncpg.create_pool(dsn, min_size=2, max_size=15)
    else:
        # Parse libpq keyword format: host=... port=... dbname=... user=... password=...
        params = {}
        for m in re.finditer(r"(\w+)=(\S+)", dsn):
            params[m.group(1)] = m.group(2)
        pool = await asyncpg.create_pool(
            host=params.get("host", "127.0.0.1"),
            port=int(params.get("port", 5432)),
            database=params.get("dbname", "dopebling_app"),
            user=params.get("user", "dopebling_idgapp"),
            password=params.get("password", ""),
            min_size=2,
            max_size=15,
        )
    return PgDatabase(pool)
