"""RLS 多租户验证

展示：
    - 以 postgres 超级用户（BYPASSRLS）能看到所有行
    - 以 blog_app 普通角色 + set_config('app.tenant_id', ...) 只能看到自己租户的行
    - 跨租户 INSERT 会被策略的 WITH CHECK 拒绝

运行：python code/multi_tenant_demo.py
"""
from __future__ import annotations

import os
import psycopg
from psycopg.rows import dict_row

from db import close_pool

SUPER_DSN = os.environ.get(
    "PG_DSN",
    "host=127.0.0.1 port=5432 dbname=learn_pg user=postgres password=postgres",
)

# init.sql 创建的业务角色
APP_DSN = (SUPER_DSN
           .replace("user=postgres", "user=blog_app")
           .replace("password=postgres", "password=blog_app_pwd"))


def count_as(dsn: str, tenant_id: int | None) -> int:
    with psycopg.connect(dsn, row_factory=dict_row,
                         options="-c search_path=blog,public") as conn:
        with conn.cursor() as cur:
            if tenant_id is not None:
                cur.execute("SELECT set_config('app.tenant_id', %s, false)",
                            (str(tenant_id),))
            cur.execute("SELECT COUNT(*) AS c FROM posts")
            return cur.fetchone()["c"]


def try_cross_tenant_insert(dsn: str, my_tenant: int, evil_tenant: int) -> str:
    """尝试以 my_tenant 身份写入一条 evil_tenant 的数据，应被 RLS 拒绝。"""
    try:
        with psycopg.connect(dsn, row_factory=dict_row,
                             options="-c search_path=blog,public") as conn:
            with conn.cursor() as cur:
                cur.execute("SELECT set_config('app.tenant_id', %s, false)",
                            (str(my_tenant),))
                cur.execute(
                    """INSERT INTO posts(tenant_id, author_id, title, body,
                                         status)
                       VALUES (%s, 1, '跨租户攻击', 'bad', 'published')""",
                    (evil_tenant,),
                )
                conn.commit()
        return "[!!!] 居然成功了（说明 RLS 没生效！）"
    except psycopg.errors.InsufficientPrivilege as e:
        return f"[OK] 被 WITH CHECK 拒绝：{e}"
    except Exception as e:
        return f"[OK] 被拒绝：{type(e).__name__}: {e}"


def main() -> None:
    print("=== RLS 多租户演示 ===\n")

    print("-- 1) 超级用户 postgres（默认 BYPASSRLS，不受策略限制） --")
    all_cnt = count_as(SUPER_DSN, tenant_id=1)
    print(f"   postgres 看到 posts 总数 = {all_cnt}\n")

    print("-- 2) 普通角色 blog_app，tenant=1 --")
    try:
        t1 = count_as(APP_DSN, tenant_id=1)
        print(f"   tenant=1 视角下 posts 数 = {t1}")
    except Exception as e:
        print(f"   [WARN] 无法连 blog_app：{e}")
        print("   （确保 init.sql 执行过，且 pg_hba.conf 允许密码登录）")
        return

    print("-- 3) 普通角色 blog_app，tenant=2 --")
    t2 = count_as(APP_DSN, tenant_id=2)
    print(f"   tenant=2 视角下 posts 数 = {t2}")
    print("   （应该只看到 seed.sql 里属于 tenant=2 的少量文章）\n")

    print("-- 4) 尝试以 tenant=1 的身份写入 tenant=2 的数据 --")
    print(" ", try_cross_tenant_insert(APP_DSN, my_tenant=1, evil_tenant=2))


if __name__ == "__main__":
    try:
        main()
    finally:
        close_pool()
