#!/usr/bin/env python
"""
Standalone Django session resurrection race reproduction.

Run from a Django checkout with:

    python session_resurrection_race_repro.py

The script builds a minimal Django app with real auth/session middleware and
three endpoints:

    /slow-save/  - authenticated endpoint that reads and modifies the session.
    /logout/     - calls django.contrib.auth.logout(request).
    /whoami/     - reports authentication state and session key.

For positive cache/file cases, the script instruments the backend write window
so that logout deletes the session after the backend has observed the old
session as existing, but before the final write recreates it.
"""

import json
import os
import subprocess
import sys
import textwrap


CASE_SCRIPT = r"""
import json
import os
import tempfile
import threading
from importlib import import_module

from django.conf import settings

label = os.environ["LABEL"]
session_engine = os.environ["SESSION_ENGINE"]
base_dir = tempfile.mkdtemp(prefix="django-session-race-")

view_entered = threading.Event()
view_release = threading.Event()
save_window_entered = threading.Event()
save_window_release = threading.Event()
pause_save_window = threading.Event()

if label == "cache":
    from django.core.cache.backends.locmem import LocMemCache

    class PausingLocMemCache(LocMemCache):
        def set(self, key, value, timeout=None, version=None):
            if (
                pause_save_window.is_set()
                and key.startswith("django.contrib.sessions.cache")
                and not save_window_entered.is_set()
            ):
                save_window_entered.set()
                if not save_window_release.wait(timeout=10):
                    raise RuntimeError("save-window release timed out")
            return super().set(key, value, timeout=timeout, version=version)

    cache_backend = "__main__.PausingLocMemCache"
else:
    cache_backend = "django.core.cache.backends.locmem.LocMemCache"

settings.configure(
    SECRET_KEY="secret",
    DEBUG=False,
    ALLOWED_HOSTS=["testserver"],
    ROOT_URLCONF="__main__",
    DEFAULT_AUTO_FIELD="django.db.models.AutoField",
    INSTALLED_APPS=[
        "django.contrib.auth",
        "django.contrib.contenttypes",
        "django.contrib.sessions",
    ],
    DATABASES={
        "default": {
            "ENGINE": "django.db.backends.sqlite3",
            "NAME": os.path.join(base_dir, "db.sqlite3"),
        }
    },
    MIDDLEWARE=[
        "django.contrib.sessions.middleware.SessionMiddleware",
        "django.contrib.auth.middleware.AuthenticationMiddleware",
    ],
    SESSION_ENGINE=session_engine,
    SESSION_CACHE_ALIAS="default",
    SESSION_FILE_PATH=os.path.join(base_dir, "sessions"),
    SESSION_COOKIE_NAME="sessionid",
    SESSION_COOKIE_AGE=1209600,
    SESSION_COOKIE_DOMAIN=None,
    SESSION_COOKIE_PATH="/",
    SESSION_COOKIE_SECURE=False,
    SESSION_COOKIE_HTTPONLY=True,
    SESSION_COOKIE_SAMESITE="Lax",
    SESSION_SAVE_EVERY_REQUEST=False,
    CACHES={
        "default": {
            "BACKEND": cache_backend,
            "LOCATION": "session-race-" + label,
        }
    },
    PASSWORD_HASHERS=["django.contrib.auth.hashers.MD5PasswordHasher"],
    USE_TZ=True,
)
os.makedirs(settings.SESSION_FILE_PATH, exist_ok=True)

import django

django.setup()

from django.contrib.auth import logout
from django.contrib.auth.models import User
from django.core.cache import caches
from django.core.management import call_command
from django.http import JsonResponse
from django.test import Client
from django.urls import path

call_command("migrate", verbosity=0, interactive=False)
User.objects.create_user(username="alice", password="password")
caches["default"].clear()

if label == "file":
    from django.contrib.sessions.backends import file as file_backend

    original_move = file_backend.shutil.move

    def pausing_move(src, dst, *args, **kwargs):
        if (
            pause_save_window.is_set()
            and os.path.basename(dst).startswith(settings.SESSION_COOKIE_NAME)
            and not save_window_entered.is_set()
        ):
            save_window_entered.set()
            if not save_window_release.wait(timeout=10):
                raise RuntimeError("save-window release timed out")
        return original_move(src, dst, *args, **kwargs)

    file_backend.shutil.move = pausing_move


def session_exists(key):
    if not key:
        return False
    store = import_module(settings.SESSION_ENGINE).SessionStore()
    return bool(store.exists(key))


def slow_save(request):
    if not request.user.is_authenticated:
        return JsonResponse({"error": "not authenticated"}, status=401)
    old_key = request.session.session_key
    request.session.get("_auth_user_id")  # Force session data to load.
    request.session["race_marker"] = label
    request.session["counter"] = request.session.get("counter", 0) + 1
    view_entered.set()
    released = view_release.wait(timeout=10)
    return JsonResponse(
        {
            "ok": released,
            "slow_key": old_key,
            "session_key_now": request.session.session_key,
        }
    )


def logout_view(request):
    old_key = request.session.session_key
    before = session_exists(old_key)
    logout(request)
    after = session_exists(old_key)
    return JsonResponse(
        {
            "old_key": old_key,
            "exists_before_logout": before,
            "exists_after_logout_view": after,
        }
    )


def whoami(request):
    return JsonResponse(
        {
            "authenticated": bool(request.user.is_authenticated),
            "username": request.user.get_username()
            if request.user.is_authenticated
            else "",
            "session_key": request.session.session_key,
            "race_marker": request.session.get("race_marker"),
            "auth_user_id": request.session.get("_auth_user_id"),
        }
    )


urlpatterns = [
    path("slow-save/", slow_save),
    path("logout/", logout_view),
    path("whoami/", whoami),
]

login_client = Client(raise_request_exception=False)
assert login_client.login(username="alice", password="password")
old_sessionid = login_client.cookies[settings.SESSION_COOKIE_NAME].value

slow_client = Client(raise_request_exception=False)
logout_client = Client(raise_request_exception=False)
slow_client.cookies[settings.SESSION_COOKIE_NAME] = old_sessionid
logout_client.cookies[settings.SESSION_COOKIE_NAME] = old_sessionid

pause_save_window.set()
slow_box = {}


def run_slow_request():
    slow_box["response"] = slow_client.get("/slow-save/")


thread = threading.Thread(target=run_slow_request)
thread.start()
assert view_entered.wait(timeout=10), "slow-save did not reach view barrier"
view_release.set()
assert save_window_entered.wait(timeout=10), "slow-save did not reach save window"

logout_response = logout_client.get("/logout/")
logout_cookie = logout_response.cookies.get(settings.SESSION_COOKIE_NAME)
exists_after_logout = session_exists(old_sessionid)

save_window_release.set()
thread.join(timeout=10)

slow_response = slow_box["response"]
slow_cookie = slow_response.cookies.get(settings.SESSION_COOKIE_NAME)
exists_after_slow = session_exists(old_sessionid)

# Apply browser-like final cookie state:
# initial cookie -> logout deletion -> late slow-save Set-Cookie.
final_cookie = old_sessionid
if logout_cookie is not None:
    if logout_cookie.value == "" or logout_cookie.get("max-age") in (0, "0"):
        final_cookie = None
    else:
        final_cookie = logout_cookie.value
if slow_cookie is not None:
    if slow_cookie.value == "" or slow_cookie.get("max-age") in (0, "0"):
        final_cookie = None
    else:
        final_cookie = slow_cookie.value

who_client = Client(raise_request_exception=False)
if final_cookie:
    who_client.cookies[settings.SESSION_COOKIE_NAME] = final_cookie
whoami_response = who_client.get("/whoami/")
whoami = json.loads(whoami_response.content.decode())

result = {
    "backend": session_engine,
    "case": label,
    "django": django.get_version(),
    "old_sessionid": old_sessionid,
    "logout_status": logout_response.status_code,
    "logout_json": json.loads(logout_response.content.decode()),
    "logout_deleted_cookie": bool(
        logout_cookie is not None
        and logout_cookie.value == ""
        and logout_cookie.get("max-age") in (0, "0")
    ),
    "exists_after_logout": exists_after_logout,
    "slow_status": slow_response.status_code,
    "slow_json": json.loads(slow_response.content.decode()),
    "slow_set_cookie_same_old_key": bool(
        slow_cookie is not None and slow_cookie.value == old_sessionid
    ),
    "exists_after_slow": exists_after_slow,
    "final_cookie": final_cookie,
    "whoami": whoami,
}
print(json.dumps(result, sort_keys=True))
"""


def run_case(label, session_engine):
    env = os.environ.copy()
    env["LABEL"] = label
    env["SESSION_ENGINE"] = session_engine
    return subprocess.run(
        [sys.executable, "-c", CASE_SCRIPT],
        text=True,
        capture_output=True,
        env=env,
        timeout=60,
    )


def main():
    cases = [
        ("cache", "django.contrib.sessions.backends.cache"),
        ("file", "django.contrib.sessions.backends.file"),
    ]
    for label, engine in cases:
        proc = run_case(label, engine)
        print(f"=== {label} ===")
        if proc.stdout:
            print(proc.stdout.strip())
        if proc.stderr:
            print("STDERR:")
            print(textwrap.indent(proc.stderr.strip(), "  "))
        if proc.returncode:
            raise SystemExit(proc.returncode)


if __name__ == "__main__":
    main()
