"""Tool handlers for the oGMemory MCP adapter."""

from __future__ import annotations

import json
from typing import Any

from core.models import RetrievedBlock
from server.mcp_resources import (
    MEMORY_TYPES,
    external_id_from_uri,
    external_to_internal_uri_candidates,
    internal_to_external_uri,
    parse_ogmemory_uri,
    public_type_from_internal,
    render_resource,
)
from server.mcp_security import (
    UNTRUSTED_DATA_NOTICE,
    build_mcp_context,
    is_not_found_or_not_visible_error,
    raise_not_found_or_not_visible,
)

_FORBIDDEN_WRITE_FIELDS = {
    "type",
    "memory_type",
    "memory_id",
    "uri",
    "abstract",
    "overview",
}


def query_memory(
    service: Any,
    *,
    query: str,
    types: list[str] | None = None,
    top_k: int = 5,
    score_threshold: float | None = None,
    session_id: str | None = None,
    agent_id: str | None = None,
    transport_headers: dict[str, Any] | None = None,
    allow_stdio_env_auth: bool = True,
    resolved_identity: Any | None = None,
) -> dict[str, Any]:
    if not (query or "").strip():
        raise ValueError("query must not be empty")
    if top_k < 1 or top_k > 20:
        raise ValueError("top_k must be between 1 and 20")

    requested_types = types or []
    unknown = sorted(t for t in requested_types if t not in MEMORY_TYPES)
    if unknown:
        raise ValueError(f"Unknown memory types: {', '.join(unknown)}")
    _validate_request_size(
        service,
        {
            "query": query,
            "types": requested_types,
            "top_k": top_k,
            "score_threshold": score_threshold,
            "session_id": session_id,
            "agent_id": agent_id,
        },
    )

    arguments = {"session_id": session_id, "agent_id": agent_id}
    ctx = build_mcp_context(
        service,
        arguments,
        transport_headers=transport_headers,
        allow_stdio_env_auth=allow_stdio_env_auth,
        resolved_identity=resolved_identity,
    )
    categories = [MEMORY_TYPES[t].internal_category for t in requested_types] or None

    try:
        result = service.get_read_api().search_memory(
            query.strip(),
            ctx,
            top_k=top_k,
            categories=categories,
            score_threshold=score_threshold,
            include_debug=False,
            fill_content_for_top_k=0,
        )
    except Exception as exc:
        raise_not_found_or_not_visible(exc)
        raise

    hits = []
    for hit in result.hits:
        try:
            hits.append(_hit_to_public(hit))
        except ValueError:
            continue

    payload = {
        "request_id": result.request_id,
        "query": result.query or query,
        "hits": hits,
        "untrusted_data_notice": UNTRUSTED_DATA_NOTICE,
    }
    return _truncate_query_payload(payload, _configured_max_response_chars(service))


def read_memory(
    service: Any,
    *,
    uri: str,
    level: str = "all",
    max_chars: int = 12000,
    session_id: str | None = None,
    agent_id: str | None = None,
    transport_headers: dict[str, Any] | None = None,
    allow_stdio_env_auth: bool = True,
    resolved_identity: Any | None = None,
) -> dict[str, Any]:
    if level not in {"abstract", "overview", "content", "all"}:
        raise ValueError("level must be one of: abstract, overview, content, all")
    if max_chars < 1 or max_chars > 50000:
        raise ValueError("max_chars must be between 1 and 50000")
    _validate_request_size(
        service,
        {
            "uri": uri,
            "level": level,
            "max_chars": max_chars,
            "session_id": session_id,
            "agent_id": agent_id,
        },
    )
    max_chars = min(max_chars, _configured_max_response_chars(service))

    arguments = {"session_id": session_id, "agent_id": agent_id}
    ctx = build_mcp_context(
        service,
        arguments,
        transport_headers=transport_headers,
        allow_stdio_env_auth=allow_stdio_env_auth,
        resolved_identity=resolved_identity,
    )
    parsed = parse_ogmemory_uri(uri)

    if parsed.kind != "item":
        payload = render_resource(service, uri, ctx)
        return _truncate_payload(payload, max_chars)

    try:
        context_fs = service._get_context_fs()
        last_exc: Exception | None = None
        node = None
        for internal_uri in external_to_internal_uri_candidates(uri, ctx):
            try:
                if hasattr(context_fs, "read_node_fields"):
                    node = context_fs.read_node_fields(internal_uri, ctx, level)
                else:
                    node = context_fs.read_node(internal_uri, ctx)
                break
            except Exception as exc:
                if not is_not_found_or_not_visible_error(exc):
                    raise
                last_exc = exc
        if node is None:
            if last_exc is not None:
                raise last_exc
            raise ValueError(f"URI does not reference a single memory: {uri}")
        canonical_uri = internal_to_external_uri(node.uri, category=node.category)
        canonical_parsed = parse_ogmemory_uri(canonical_uri)
        payload = {
            "uri": canonical_uri,
            "type": canonical_parsed.memory_type,
            "id": canonical_parsed.memory_id,
            "untrusted_data_notice": UNTRUSTED_DATA_NOTICE,
        }
        if level in {"abstract", "all"}:
            payload["abstract"] = node.abstract
        if level in {"overview", "all"}:
            payload["overview"] = node.overview
        if level in {"content", "all"}:
            payload["content"] = node.content
        return _truncate_payload(payload, max_chars)
    except Exception as exc:
        raise_not_found_or_not_visible(exc)
        raise


def delete_memory(
    service: Any,
    *,
    uri: str,
    reason: str | None = None,
    session_id: str | None = None,
    agent_id: str | None = None,
    client_request_id: str | None = None,
    transport_headers: dict[str, Any] | None = None,
    allow_stdio_env_auth: bool = True,
    resolved_identity: Any | None = None,
) -> dict[str, Any]:
    parsed = parse_ogmemory_uri(uri)
    if parsed.kind != "item":
        raise ValueError("delete_memory requires a single ogmemory://memories/{type}/{id} URI")
    _validate_request_size(
        service,
        {
            "uri": uri,
            "reason": reason,
            "session_id": session_id,
            "agent_id": agent_id,
            "client_request_id": client_request_id,
        },
    )

    arguments = {"session_id": session_id, "agent_id": agent_id}
    ctx = build_mcp_context(
        service,
        arguments,
        transport_headers=transport_headers,
        allow_stdio_env_auth=allow_stdio_env_auth,
        resolved_identity=resolved_identity,
    )

    try:
        last_exc: Exception | None = None
        result = None
        for internal_uri in external_to_internal_uri_candidates(uri, ctx):
            try:
                result = service.forget_memory(
                    internal_uri,
                    ctx,
                    external_uri=uri,
                    reason=reason,
                    client_request_id=client_request_id,
                )
                break
            except Exception as exc:
                if not is_not_found_or_not_visible_error(exc):
                    raise
                last_exc = exc
        if result is None:
            if last_exc is not None:
                raise last_exc
            raise ValueError(f"URI does not reference a single memory: {uri}")
    except Exception as exc:
        raise_not_found_or_not_visible(exc)
        raise

    return {
        "ok": bool(result.get("ok", True)),
        "deleted": bool(result.get("deleted", True)),
        "mode": result.get("mode", "archive"),
        "uri": uri,
        "type": parsed.memory_type,
        "id": parsed.memory_id,
        "client_request_id": client_request_id,
    }


def _hit_to_public(hit: RetrievedBlock) -> dict[str, Any]:
    external_uri = internal_to_external_uri(hit.uri, category=hit.category)
    parsed = parse_ogmemory_uri(external_uri)
    return {
        "uri": external_uri,
        "type": parsed.memory_type or public_type_from_internal(hit.category, hit.uri),
        "id": external_id_from_uri(external_uri),
        "score": hit.score,
        "abstract": hit.abstract,
        "has_content": hit.has_content or bool(hit.content_excerpt),
    }


def _truncate_payload(payload: dict[str, Any], max_chars: int) -> dict[str, Any]:
    remaining = max_chars
    truncated = bool(payload.get("truncated", False))
    result: dict[str, Any] = {}
    for key, value in payload.items():
        if isinstance(value, str) and key in {"abstract", "overview", "content"}:
            if len(value) > remaining:
                result[key] = value[: max(0, remaining)]
                truncated = True
                remaining = 0
            else:
                result[key] = value
                remaining -= len(value)
        else:
            result[key] = value
    result["truncated"] = truncated
    return result


def _truncate_query_payload(payload: dict[str, Any], max_chars: int) -> dict[str, Any]:
    remaining = max_chars
    truncated = False
    result = dict(payload)
    truncated_hits: list[dict[str, Any]] = []
    for hit in payload.get("hits", []):
        if remaining <= 0:
            truncated = True
            break
        row = dict(hit)
        abstract = row.get("abstract")
        if isinstance(abstract, str):
            if len(abstract) > remaining:
                row["abstract"] = abstract[: max(0, remaining)]
                truncated = True
                remaining = 0
            else:
                remaining -= len(abstract)
        truncated_hits.append(row)
    result["hits"] = truncated_hits
    result["truncated"] = truncated
    return result


def _configured_max_request_bytes(service: Any) -> int:
    value = int(getattr(service._cfg, "mcp_max_request_bytes", 1048576) or 1048576)
    return max(1, value)


def _configured_max_response_chars(service: Any) -> int:
    value = int(getattr(service._cfg, "mcp_max_response_chars", 50000) or 50000)
    return max(1, min(value, 50000))


def _validate_request_size(service: Any, payload: dict[str, Any]) -> None:
    max_bytes = _configured_max_request_bytes(service)
    encoded = json.dumps(payload, ensure_ascii=False, default=str, sort_keys=True).encode("utf-8")
    if len(encoded) > max_bytes:
        raise ValueError(f"MCP request exceeds mcp_max_request_bytes ({max_bytes})")


def write_memory(
    service: Any,
    *,
    text: str | None = None,
    messages: list[dict[str, Any]] | None = None,
    session_id: str | None = None,
    agent_id: str | None = None,
    client_request_id: str | None = None,
    wait: bool = False,
    transport_headers: dict[str, Any] | None = None,
    allow_stdio_env_auth: bool = True,
    resolved_identity: Any | None = None,
) -> dict[str, Any]:
    if session_id is None or not session_id.strip():
        raise ValueError("session_id is required for write_memory")
    normalized_messages = _normalize_write_messages(text=text, messages=messages)
    _validate_request_size(
        service,
        {
            "messages": normalized_messages,
            "session_id": session_id,
            "agent_id": agent_id,
            "client_request_id": client_request_id,
            "wait": wait,
        },
    )

    arguments = {"session_id": session_id, "agent_id": agent_id}
    ctx = build_mcp_context(
        service,
        arguments,
        transport_headers=transport_headers,
        allow_stdio_env_auth=allow_stdio_env_auth,
        resolved_identity=resolved_identity,
    )
    params: dict[str, Any] = {
        "sessionId": session_id,
        "messages": normalized_messages,
        "_ctx": ctx,
    }
    if agent_id:
        params["agentId"] = agent_id
    if client_request_id:
        params["clientRequestId"] = client_request_id
        params["client_request_id"] = client_request_id
    if wait:
        params["wait"] = True

    try:
        result = service.after_turn(params)
    except Exception as exc:
        _record_mcp_write_audit(
            service,
            ctx,
            session_id=session_id,
            agent_id=agent_id,
            client_request_id=client_request_id,
            result="failed",
            details={"error": f"{type(exc).__name__}: {exc}"},
        )
        raise

    _record_mcp_write_audit(
        service,
        ctx,
        session_id=session_id,
        agent_id=agent_id,
        client_request_id=client_request_id,
        result="success" if result.get("ok", True) else "failed",
        details={
            "status": result.get("status", "accepted"),
            "accepted_messages": len(normalized_messages),
            "wait_requested": wait,
        },
    )

    return {
        "ok": bool(result.get("ok", True)),
        "status": result.get("status", "accepted"),
        "session_id": session_id,
        "agent_id": agent_id or ctx.agent_id,
        "client_request_id": client_request_id,
        "accepted_messages": len(normalized_messages),
        "stats": _public_write_stats(result),
    }


def _normalize_write_messages(
    *,
    text: str | None,
    messages: list[dict[str, Any]] | None,
) -> list[dict[str, Any]]:
    if text and messages:
        raise ValueError("Provide either text or messages, not both")
    if text is not None:
        if not text.strip():
            raise ValueError("text must not be empty")
        return [{"role": "user", "content": text}]
    if not messages:
        raise ValueError("messages or text is required")

    normalized: list[dict[str, Any]] = []
    for index, message in enumerate(messages):
        if not isinstance(message, dict):
            raise ValueError(f"messages[{index}] must be an object")
        forbidden = sorted(k for k in message if k in _FORBIDDEN_WRITE_FIELDS)
        if forbidden:
            raise ValueError(
                f"messages[{index}] cannot set atomic memory fields: {', '.join(forbidden)}"
            )
        role = str(message.get("role") or "").strip()
        content = message.get("content")
        if role not in {"system", "user", "assistant", "tool"}:
            raise ValueError(f"messages[{index}].role must be system, user, assistant, or tool")
        if not isinstance(content, str) or not content.strip():
            raise ValueError(f"messages[{index}].content must be a non-empty string")
        normalized_message: dict[str, Any] = {"role": role, "content": content}
        if message.get("id"):
            normalized_message["id"] = str(message["id"])
        if message.get("created_at"):
            normalized_message["created_at"] = str(message["created_at"])
        normalized.append(normalized_message)
    return normalized


def _record_mcp_write_audit(
    service: Any,
    ctx: Any,
    *,
    session_id: str,
    agent_id: str | None,
    client_request_id: str | None,
    result: str,
    details: dict[str, Any],
) -> None:
    get_audit_service = getattr(service, "get_audit_service", None)
    if not callable(get_audit_service):
        return
    try:
        get_audit_service().record(
            ctx.account_id,
            actor=ctx.user_id,
            target=session_id,
            action="mcp.write_memory",
            result=result,
            trace_id=ctx.trace_id,
            details={
                "role": str(ctx.role),
                "agent_id": agent_id or ctx.agent_id,
                "session_id": session_id,
                "client_request_id": client_request_id,
                **details,
            },
        )
    except Exception:
        return


def _public_write_stats(result: dict[str, Any]) -> dict[str, Any]:
    allowed_fields = {
        "pending_tokens",
        "buffered_tokens",
        "candidates_extracted",
        "writes_completed",
        "writes_failed",
        "task_id",
        "reason",
    }
    stats: dict[str, Any] = {}
    for field in allowed_fields:
        value = result.get(field)
        if isinstance(value, (str, int, float, bool)) or value is None:
            stats[field] = value
    return stats