"""Azalee local audit helper.

Zero third-party dependencies. Reads AZALEE_WEBHOOK_URL from the process
environment or from a local .env.azalee file. Logging is fail-open by default.
Raw questions, answers, tool arguments and results stay local; only hashes and
control metadata are sent to Azalee.
"""

from __future__ import annotations

import asyncio
import functools
import hashlib
import inspect
import json
import os
import time
import urllib.error
import urllib.request
import uuid
from pathlib import Path
from typing import Any, Callable, Dict, Optional

DEFAULT_TIMEOUT_SECONDS = 1.5


def _read_local_env() -> Dict[str, str]:
    path = Path(".env.azalee")
    if not path.exists():
        return {}
    result: Dict[str, str] = {}
    for raw_line in path.read_text(encoding="utf-8").splitlines():
        line = raw_line.strip()
        if not line or line.startswith("#") or "=" not in line:
            continue
        key, value = line.split("=", 1)
        result[key.strip()] = value.strip().strip('"\'')
    return result


def _canonical(value: Any) -> str:
    return json.dumps(value, sort_keys=True, separators=(",", ":"), ensure_ascii=False, default=str)


def hash_value(value: Any) -> str:
    return "sha3-256:" + hashlib.sha3_256(_canonical(value).encode("utf-8")).hexdigest()


_STRING_FIELDS = (
    "actor_type", "actor_id", "correlation_id", "task_id", "tool_name",
    "input_hash", "output_hash", "risk_level", "risk_severity",
    "model_provider", "model_name", "model_version", "prompt_version",
    "config_version", "deployment_version", "policy_id", "reason_code", "channel",
)


def _compact_event(event_type: str, fields: Dict[str, Any]) -> tuple[Dict[str, Any], Optional[str]]:
    session_id = fields.get("session_id") or fields.get("sessionId") or fields.get("run_id") or fields.get("runId")
    event: Dict[str, Any] = {
        "event_type": event_type,
        "status": fields.get("status", "success"),
    }
    if session_id is not None:
        event["session_id"] = str(session_id)
    for key in _STRING_FIELDS:
        value = fields.get(key)
        if value is not None and str(value).strip():
            event[key] = str(value)
    if isinstance(fields.get("duration_ms"), (int, float)):
        event["duration_ms"] = int(fields["duration_ms"])
    if isinstance(fields.get("sequence"), (int, float)):
        event["sequence"] = int(fields["sequence"])
    if isinstance(fields.get("approval_granted"), bool):
        event["approval_granted"] = fields["approval_granted"]
    if isinstance(fields.get("disclosure_shown"), bool):
        event["disclosure_shown"] = fields["disclosure_shown"]
    return event, str(session_id) if session_id is not None else None


class AzaleeAudit:
    def __init__(
        self,
        webhook_url: Optional[str] = None,
        timeout_seconds: float = DEFAULT_TIMEOUT_SECONDS,
        source: str = "local_runtime",
    ) -> None:
        local_env = _read_local_env()
        self.webhook_url = webhook_url or os.getenv("AZALEE_WEBHOOK_URL") or local_env.get("AZALEE_WEBHOOK_URL")
        self.timeout_seconds = float(timeout_seconds)
        self.source = source
        self._last_hash_by_session: Dict[str, str] = {}

    def event(
        self,
        event_type: str,
        fields: Optional[Dict[str, Any]] = None,
        metadata: Optional[Dict[str, Any]] = None,
    ) -> Dict[str, Any]:
        if not self.webhook_url:
            return {"ok": False, "error": "AZALEE_WEBHOOK_URL is missing."}

        fields = dict(fields or {})
        event, session_id = _compact_event(event_type, fields)
        previous_hash = fields.get("previous_hash") or fields.get("previousHash")
        if not previous_hash and session_id:
            previous_hash = self._last_hash_by_session.get(session_id)

        body: Dict[str, Any] = {
            "event": event,
            "metadata": {"source": self.source, "storage": "hash_only", **(metadata or {})},
        }
        if session_id is not None:
            body["sessionId"] = session_id
        if previous_hash:
            body["previousHash"] = str(previous_hash)

        request = urllib.request.Request(
            self.webhook_url,
            data=json.dumps(body, separators=(",", ":")).encode("utf-8"),
            headers={"Content-Type": "application/json"},
            method="POST",
        )
        try:
            with urllib.request.urlopen(request, timeout=self.timeout_seconds) as response:
                raw = response.read().decode("utf-8", errors="replace")
                parsed = json.loads(raw) if raw else {}
                if session_id and isinstance(parsed, dict) and parsed.get("hash"):
                    self._last_hash_by_session[session_id] = str(parsed["hash"])
                return {"ok": 200 <= response.status < 300, "status": response.status, **(parsed if isinstance(parsed, dict) else {})}
        except urllib.error.HTTPError as error:
            return {"ok": False, "status": error.code, "error": error.read().decode("utf-8", errors="replace")}
        except Exception as error:  # fail-open by design
            return {"ok": False, "error": str(error)}

    def test(self) -> Dict[str, Any]:
        return self.event("integration.test", {"status": "success"}, {"installer": "azalee-local"})

    def interaction(
        self,
        input_value: Any,
        function: Callable[[], Any],
        *,
        session_id: Optional[str] = None,
        correlation_id: Optional[str] = None,
        delivered: bool = True,
        reason_code: Optional[str] = None,
        model_provider: Optional[str] = None,
        model_name: Optional[str] = None,
        model_version: Optional[str] = None,
        prompt_version: Optional[str] = None,
        config_version: Optional[str] = None,
        deployment_version: Optional[str] = None,
        metadata: Optional[Dict[str, Any]] = None,
    ) -> Any:
        started = time.monotonic()
        correlation = correlation_id or str(uuid.uuid4())
        input_hash = hash_value(input_value)
        shared = {
            "session_id": session_id,
            "correlation_id": correlation,
            "actor_type": "agent",
            "model_provider": model_provider,
            "model_name": model_name,
            "model_version": model_version,
            "prompt_version": prompt_version,
            "config_version": config_version,
            "deployment_version": deployment_version,
        }
        self.event("interaction.started", {**shared, "status": "started", "sequence": 1, "input_hash": input_hash}, metadata)

        try:
            result = function()
            output_hash = hash_value(result)
            duration_ms = int((time.monotonic() - started) * 1000)
            self.event("model.response.completed", {
                **shared, "status": "success", "sequence": 2, "duration_ms": duration_ms,
                "input_hash": input_hash, "output_hash": output_hash,
            }, metadata)
            self.event("output.delivered" if delivered else "output.blocked", {
                **shared, "status": "success" if delivered else "blocked", "sequence": 3,
                "output_hash": output_hash, "reason_code": reason_code,
            }, metadata)
            self.event("interaction.completed", {
                **shared, "status": "success" if delivered else "blocked", "sequence": 4,
                "duration_ms": duration_ms, "input_hash": input_hash, "output_hash": output_hash,
            }, metadata)
            return result
        except Exception as error:
            duration_ms = int((time.monotonic() - started) * 1000)
            error_hash = hash_value({"name": type(error).__name__})
            self.event("model.response.failed", {
                **shared, "status": "failure", "sequence": 2, "duration_ms": duration_ms,
                "input_hash": input_hash, "output_hash": error_hash, "reason_code": type(error).__name__,
            }, metadata)
            self.event("interaction.failed", {
                **shared, "status": "failure", "sequence": 3, "duration_ms": duration_ms,
                "input_hash": input_hash, "output_hash": error_hash, "reason_code": type(error).__name__,
            }, metadata)
            raise

    async def interaction_async(
        self,
        input_value: Any,
        function: Callable[[], Any],
        **options: Any,
    ) -> Any:
        started = time.monotonic()
        correlation = options.get("correlation_id") or str(uuid.uuid4())
        input_hash = hash_value(input_value)
        metadata = options.get("metadata")
        shared = {
            "session_id": options.get("session_id"),
            "correlation_id": correlation,
            "actor_type": "agent",
            "model_provider": options.get("model_provider"),
            "model_name": options.get("model_name"),
            "model_version": options.get("model_version"),
            "prompt_version": options.get("prompt_version"),
            "config_version": options.get("config_version"),
            "deployment_version": options.get("deployment_version"),
        }
        await asyncio.to_thread(self.event, "interaction.started", {**shared, "status": "started", "sequence": 1, "input_hash": input_hash}, metadata)
        try:
            result = function()
            if inspect.isawaitable(result):
                result = await result
            output_hash = hash_value(result)
            duration_ms = int((time.monotonic() - started) * 1000)
            await asyncio.to_thread(self.event, "model.response.completed", {
                **shared, "status": "success", "sequence": 2, "duration_ms": duration_ms,
                "input_hash": input_hash, "output_hash": output_hash,
            }, metadata)
            delivered = options.get("delivered", True)
            await asyncio.to_thread(self.event, "output.delivered" if delivered else "output.blocked", {
                **shared, "status": "success" if delivered else "blocked", "sequence": 3,
                "output_hash": output_hash, "reason_code": options.get("reason_code"),
            }, metadata)
            await asyncio.to_thread(self.event, "interaction.completed", {
                **shared, "status": "success" if delivered else "blocked", "sequence": 4,
                "duration_ms": duration_ms, "input_hash": input_hash, "output_hash": output_hash,
            }, metadata)
            return result
        except Exception as error:
            duration_ms = int((time.monotonic() - started) * 1000)
            error_hash = hash_value({"name": type(error).__name__})
            await asyncio.to_thread(self.event, "model.response.failed", {
                **shared, "status": "failure", "sequence": 2, "duration_ms": duration_ms,
                "input_hash": input_hash, "output_hash": error_hash, "reason_code": type(error).__name__,
            }, metadata)
            await asyncio.to_thread(self.event, "interaction.failed", {
                **shared, "status": "failure", "sequence": 3, "duration_ms": duration_ms,
                "input_hash": input_hash, "output_hash": error_hash, "reason_code": type(error).__name__,
            }, metadata)
            raise

    def wrap(
        self,
        tool_name: str,
        *,
        session_id: Optional[str] = None,
        correlation_id: Optional[str] = None,
        task_id: Optional[str] = None,
        approval_granted: Optional[bool] = None,
        metadata: Optional[Dict[str, Any]] = None,
    ) -> Callable[[Callable[..., Any]], Callable[..., Any]]:
        def decorator(function: Callable[..., Any]) -> Callable[..., Any]:
            shared = {
                "tool_name": tool_name,
                "session_id": session_id,
                "correlation_id": correlation_id,
                "task_id": task_id,
                "approval_granted": approval_granted,
            }
            if inspect.iscoroutinefunction(function):
                @functools.wraps(function)
                async def async_wrapper(*args: Any, **kwargs: Any) -> Any:
                    started = time.monotonic()
                    input_hash = hash_value({"args": args, "kwargs": kwargs})
                    await asyncio.to_thread(self.event, "tool.call.started", {**shared, "status": "started", "input_hash": input_hash}, metadata)
                    try:
                        result = await function(*args, **kwargs)
                        await asyncio.to_thread(self.event, "tool.call.completed", {
                            **shared, "status": "success", "duration_ms": int((time.monotonic() - started) * 1000),
                            "input_hash": input_hash, "output_hash": hash_value(result),
                        }, metadata)
                        return result
                    except Exception as error:
                        await asyncio.to_thread(self.event, "tool.call.failed", {
                            **shared, "status": "failure", "duration_ms": int((time.monotonic() - started) * 1000),
                            "input_hash": input_hash, "output_hash": hash_value({"name": type(error).__name__}),
                            "reason_code": type(error).__name__,
                        }, metadata)
                        raise
                return async_wrapper

            @functools.wraps(function)
            def sync_wrapper(*args: Any, **kwargs: Any) -> Any:
                started = time.monotonic()
                input_hash = hash_value({"args": args, "kwargs": kwargs})
                self.event("tool.call.started", {**shared, "status": "started", "input_hash": input_hash}, metadata)
                try:
                    result = function(*args, **kwargs)
                    self.event("tool.call.completed", {
                        **shared, "status": "success", "duration_ms": int((time.monotonic() - started) * 1000),
                        "input_hash": input_hash, "output_hash": hash_value(result),
                    }, metadata)
                    return result
                except Exception as error:
                    self.event("tool.call.failed", {
                        **shared, "status": "failure", "duration_ms": int((time.monotonic() - started) * 1000),
                        "input_hash": input_hash, "output_hash": hash_value({"name": type(error).__name__}),
                        "reason_code": type(error).__name__,
                    }, metadata)
                    raise
            return sync_wrapper
        return decorator


__all__ = ["AzaleeAudit", "hash_value"]
