import time
from unittest import mock

from django.conf import settings
from django.test import Client, TransactionTestCase, override_settings
from django.urls import path
from rest_framework import status

from api.models import (
    AccessToken,
    ErrorLog,
    IncomingApiAdmin,
    IntegrationOutbox,
    PostCommitDelivery,
    TtLocation,
    User,
)
from api.utils import (
    bulk_sync_to_redis,
    get_exception_detail,
    log_critical_error,
    push_websocket_notification,
)
from api.views.admin.base import AdminApiBase
from backend.kafka import send_request


TEST_REDIS = mock.Mock()


class FailingAdminDiagnosticView(AdminApiBase):
    def post(self, request):
        try:
            TtLocation.objects.create(
                code="ROLLBACK-DIAGNOSTIC",
                name="Rollback diagnostic",
                status=1,
            )
            bulk_sync_to_redis(
                TEST_REDIS,
                {"update": {"location": [{"id": 1, "name": "unsafe"}]}},
            )
            push_websocket_notification(
                "https://websocket.invalid/schedule", {"unsafe": True}, None
            )
            send_request(
                "legacy-topic",
                {"session_id": "test", "activity": [1]},
                None,
                "schedule",
            )
            raise RuntimeError("forced diagnostic failure")
        except Exception as error:
            details = get_exception_detail(error)
            log_critical_error(
                user_id=None,
                descr=details["descr"],
                url=details["url"],
                trace=details["trace"],
                data={
                    "operation": "diagnostic-test",
                    "password": "must-not-survive",
                    "signature": "must-not-survive",
                },
            )
            return self.api_response(
                error="internal server error",
                code=status.HTTP_500_INTERNAL_SERVER_ERROR,
            )


class SuccessfulAdminDiagnosticView(AdminApiBase):
    def post(self, request):
        return self.api_response(data={"ok": True})


class UnhandledAdminDiagnosticView(AdminApiBase):
    def post(self, request):
        TtLocation.objects.create(
            code="UNHANDLED-DIAGNOSTIC",
            name="Unhandled diagnostic",
            status=1,
        )
        raise RuntimeError("unhandled diagnostic failure")


urlpatterns = [
    path(
        "api/admin/failing-diagnostic",
        FailingAdminDiagnosticView.as_view(),
        name="test_timeout",
    ),
    path(
        "api/admin/successful-diagnostic",
        SuccessfulAdminDiagnosticView.as_view(),
        name="info",
    ),
    path(
        "api/admin/unhandled-diagnostic",
        UnhandledAdminDiagnosticView.as_view(),
        name="test_timeout",
    ),
]


@override_settings(ROOT_URLCONF=__name__)
class AdminFailureDiagnosticTests(TransactionTestCase):
    reset_sequences = True

    def setUp(self):
        TEST_REDIS.reset_mock()
        self.client = Client()
        self.user = User.objects.create_user(
            email="diagnostic-test@example.invalid",
            password="test-password",
            name="Diagnostic Test",
            user_type=User.USER_TYPE["admin"],
            status=User.STATUS_TO_CODE["active"],
        )
        self.token = AccessToken.objects.create(user=self.user, type="web")

    def post(self, path, *, request_id, **values):
        return self.client.post(
            path,
            data=self.signed_payload(**values),
            content_type="application/json",
            HTTP_AUTHORIZATION=f"Token {self.token.token}",
            HTTP_X_REQUEST_ID=request_id,
            HTTP_X_CORRELATION_ID=request_id,
        )

    @staticmethod
    def signed_payload(**values):
        payload = {"timestamp": int(time.time()), **values}
        payload["signature"] = AdminApiBase.sign(
            payload.copy(), settings.API_SECRET_KEY
        )
        return payload

    def test_http_500_rolls_back_domain_and_side_effects_but_retains_diagnostics(self):
        request_id = "phase1-load-test-load-02-000590"
        with (
            mock.patch("api.utils.requests.post") as websocket_post,
            mock.patch("api.services.integration.delivery.deliver_by_id") as deliver,
        ):
            response = self.post(
                "/api/admin/failing-diagnostic",
                request_id=request_id,
                password="request-password",
                token="request-token",
            )

        self.assertEqual(response.status_code, 500)
        self.assertEqual(response["X-Request-ID"], request_id)
        self.assertEqual(response["X-Correlation-ID"], request_id)
        self.assertFalse(TtLocation.objects.filter(code="ROLLBACK-DIAGNOSTIC").exists())
        self.assertFalse(IntegrationOutbox.objects.exists())
        self.assertFalse(PostCommitDelivery.objects.exists())
        TEST_REDIS.pipeline.assert_not_called()
        websocket_post.assert_not_called()
        deliver.assert_not_called()

        incoming = IncomingApiAdmin.objects.get()
        self.assertEqual(incoming.url, "/api/admin/failing-diagnostic")
        self.assertEqual(incoming.scode, 500)
        self.assertEqual(incoming.request_id, request_id)
        self.assertEqual(incoming.correlation_id, request_id)
        self.assertEqual(incoming.incoming_params["password"], "[REDACTED]")
        self.assertEqual(incoming.incoming_params["token"], "[REDACTED]")
        self.assertEqual(incoming.incoming_params["signature"], "[REDACTED]")

        error = ErrorLog.objects.get()
        self.assertEqual(error.request_id, request_id)
        self.assertEqual(error.correlation_id, request_id)
        self.assertIn("forced diagnostic failure", error.descr)
        self.assertIn("RuntimeError", error.trace)
        self.assertEqual(error.data["operation"], "diagnostic-test")
        self.assertEqual(error.data["password"], "[REDACTED]")
        self.assertEqual(error.data["signature"], "[REDACTED]")

    def test_diagnostic_database_failure_does_not_replace_original_response(self):
        with mock.patch(
            "backend.middleware.admin_diagnostic_middleware.persist_failed_admin_diagnostics",
            side_effect=RuntimeError("diagnostic database unavailable"),
        ):
            response = self.post(
                "/api/admin/failing-diagnostic",
                request_id="diagnostic-db-failure",
            )

        self.assertEqual(response.status_code, 500)
        self.assertFalse(TtLocation.objects.filter(code="ROLLBACK-DIAGNOSTIC").exists())
        self.assertFalse(IncomingApiAdmin.objects.exists())
        self.assertFalse(ErrorLog.objects.exists())

    def test_unhandled_exception_rolls_back_domain_but_retains_cause(self):
        self.client.raise_request_exception = False

        response = self.post(
            "/api/admin/unhandled-diagnostic",
            request_id="unhandled-request",
        )

        self.assertEqual(response.status_code, 500)
        self.assertFalse(
            TtLocation.objects.filter(code="UNHANDLED-DIAGNOSTIC").exists()
        )
        incoming = IncomingApiAdmin.objects.get()
        self.assertEqual(incoming.scode, 500)
        self.assertEqual(incoming.request_id, "unhandled-request")
        error = ErrorLog.objects.get()
        self.assertEqual(error.request_id, "unhandled-request")
        self.assertIn("unhandled diagnostic failure", error.descr)
        self.assertIn("RuntimeError", error.trace)

    def test_successful_request_keeps_one_existing_api_log(self):
        response = self.post(
            "/api/admin/successful-diagnostic",
            request_id="successful-request",
        )

        self.assertEqual(response.status_code, 200)
        self.assertEqual(IncomingApiAdmin.objects.count(), 1)
        incoming = IncomingApiAdmin.objects.get()
        self.assertEqual(incoming.scode, 200)
        self.assertEqual(incoming.request_id, "successful-request")
        self.assertFalse(ErrorLog.objects.exists())
