#!/usr/bin/env python3
"""Shared Microsoft Memora integration for Pi and MCP entrypoints."""

from __future__ import annotations

import hashlib
import importlib.util
import json
import os
import sys
import types
from importlib.machinery import ModuleSpec
from pathlib import Path
from typing import Any

STRUCTURED_OUTPUT_RETRIES = 2


def _setup_commands() -> list[str]:
    project_root = Path(__file__).resolve().parent.parent
    memora_repo = project_root / "vendor" / "Memora"
    memora_src = memora_repo / "src"
    return [
        f"test -d \"{memora_src}\"",
        f"uv sync --project \"{project_root}\"",
        f"uv run --project \"{project_root}\" python -c \"import sys; print(sys.version)\"",
    ]


def _add_default_memora_checkout_to_path() -> None:
    project_root = Path(__file__).resolve().parent.parent
    memora_src = project_root / "vendor" / "Memora" / "src"
    if memora_src.exists():
        sys.path.insert(0, str(memora_src))


def _openai_compat_base_url(kind: str) -> str | None:
    if kind == "EMBEDDING":
        return os.getenv("PI_MEMORA_EMBEDDING_BASE_URL") or os.getenv("OPENAI_EMBEDDING_BASE_URL") or os.getenv("OPENROUTER_EMBEDDING_BASE_URL")
    return os.getenv("OPENAI_BASE_URL") or os.getenv("OPENROUTER_BASE_URL")


def _openai_compat_api_key(kind: str) -> str | None:
    if kind == "EMBEDDING":
        return os.getenv("PI_MEMORA_EMBEDDING_API_KEY") or os.getenv("OPENAI_EMBEDDING_API_KEY") or os.getenv("OPENROUTER_API_KEY") or os.getenv("OPENAI_API_KEY")
    return os.getenv("OPENAI_API_KEY") or os.getenv("OPENROUTER_API_KEY")


def _chat_provider(model: str, base_url: str | None) -> str:
    normalized_model = model.lower()
    normalized_base_url = (base_url or "").lower()
    if "deepseek" in normalized_model or "api.deepseek.com" in normalized_base_url:
        return "deepseek"
    if not normalized_base_url or "api.openai.com" in normalized_base_url:
        return "openai"
    return "openai-compatible"


def _clean_json_response(text: str) -> str:
    stripped = text.strip()
    if stripped.startswith("```"):
        lines = stripped.splitlines()
        if lines and lines[0].startswith("```"):
            lines = lines[1:]
        if lines and lines[-1].strip() == "```":
            lines = lines[:-1]
        stripped = "\n".join(lines).strip()

    object_start = stripped.find("{")
    object_end = stripped.rfind("}")
    array_start = stripped.find("[")
    array_end = stripped.rfind("]")

    if object_start != -1 and object_end > object_start:
        return stripped[object_start : object_end + 1]
    if array_start != -1 and array_end > array_start:
        return stripped[array_start : array_end + 1]
    return stripped


def _schema_instruction(response_format: Any) -> str:
    schema = {}
    if hasattr(response_format, "model_json_schema"):
        schema = response_format.model_json_schema()
    elif hasattr(response_format, "schema"):
        schema = response_format.schema()
    return (
        "Return only valid JSON. Do not wrap it in markdown. "
        "Use exactly the field names from the schema. "
        "The JSON must match this schema: "
        f"{json.dumps(schema, ensure_ascii=False)}"
    )


def _coerce_memory_outputs(response_format: Any, data: Any) -> Any:
    if getattr(response_format, "__name__", "") != "MemoryOutputs" or not isinstance(data, dict):
        return data
    if "entries" in data:
        return data

    candidates = data.get("memories") or data.get("memory") or data.get("items")
    if not isinstance(candidates, list):
        candidates = [data]

    entries = []
    for item in candidates:
        if not isinstance(item, dict):
            continue
        index = item.get("index") or item.get("MemIndex") or item.get("key") or item.get("title")
        value = item.get("value") or item.get("Memory") or item.get("MemValue") or item.get("content") or item.get("text")
        if index and value:
            entries.append(
                {
                    "memory_type": item.get("memory_type") or item.get("type") or "Factual",
                    "index": str(index),
                    "value": str(value),
                }
            )

    return {"entries": entries} if entries else data


def _parse_structured_response(response_format: Any, text: str) -> Any:
    cleaned = _clean_json_response(text)
    data = json.loads(cleaned)
    return _validate_structured_data(response_format, data)


def _parse_tolerant_structured_response(response_format: Any, text: str) -> Any:
    cleaned = _clean_json_response(text)
    data = json.loads(cleaned)
    data = _coerce_memory_outputs(response_format, data)
    return _validate_structured_data(response_format, data)


def _validate_structured_data(response_format: Any, data: Any) -> Any:
    if hasattr(response_format, "model_validate"):
        return response_format.model_validate(data)
    if hasattr(response_format, "parse_obj"):
        return response_format.parse_obj(data)
    if hasattr(response_format, "model_validate_json"):
        return response_format.model_validate_json(json.dumps(data, ensure_ascii=False))
    if hasattr(response_format, "parse_raw"):
        return response_format.parse_raw(json.dumps(data, ensure_ascii=False))
    return data


def _structured_messages(messages: list, response_format: Any, error: str | None = None) -> list:
    instruction = _schema_instruction(response_format)
    if error:
        instruction += f" Previous JSON failed validation with this error: {error}. Return corrected JSON only."
    return [{"role": "system", "content": instruction}, *messages]


def _apply_openai_compat_patches() -> None:
    from openai import OpenAI
    import memora.utils.embedding as embedding_utils
    import memora.utils.llm as llm_utils

    def get_openai_chat_completion_client(cfg):
        api_key = cfg.openai.get("api_key", None) or _openai_compat_api_key("LLM") or "pi-memora-status-placeholder"
        kwargs: dict[str, str] = {"api_key": api_key}
        base_url = cfg.openai.get("llm_api_base", None) or _openai_compat_base_url("LLM")
        if base_url:
            kwargs["base_url"] = base_url
        return OpenAI(**kwargs)

    def get_openai_embedding_client(cfg):
        api_key = cfg.openai.get("embedding_api_key", None) or _openai_compat_api_key("EMBEDDING") or "pi-memora-status-placeholder"
        kwargs: dict[str, str] = {"api_key": api_key}
        base_url = cfg.openai.get("embedding_api_base", None) or _openai_compat_base_url("EMBEDDING")
        if base_url:
            kwargs["base_url"] = base_url
        return OpenAI(**kwargs)

    embedding_utils.get_openai_embedding_client = get_openai_embedding_client
    llm_utils.get_openai_chat_completion_client = get_openai_chat_completion_client

    llm_utils.ChatCompletionModel._determine_model_type = lambda self, model_name: "openai"

    def invoke_openai_native(self, request, response_format):
        if response_format:
            response = self.client.beta.chat.completions.parse(
                **request,
                response_format=response_format,
            )
            return response.choices[0].message.parsed
        response = self.client.chat.completions.create(**request)
        return response.choices[0].message.content or ""

    def invoke_json_mode(self, request, response_format):
        if not response_format:
            response = self.client.chat.completions.create(**request)
            return response.choices[0].message.content or ""

        last_error: Exception | None = None
        last_content = ""
        for attempt in range(STRUCTURED_OUTPUT_RETRIES + 1):
            structured_request = {
                **request,
                "messages": _structured_messages(
                    request["messages"],
                    response_format,
                    str(last_error) if last_error else None,
                ),
            }
            try:
                response = self.client.chat.completions.create(
                    **structured_request,
                    response_format={"type": "json_object"},
                )
            except Exception:
                response = self.client.chat.completions.create(**structured_request)

            last_content = response.choices[0].message.content or ""
            try:
                return _parse_structured_response(response_format, last_content)
            except Exception as exc:
                last_error = exc
                if attempt >= STRUCTURED_OUTPUT_RETRIES:
                    return _parse_tolerant_structured_response(response_format, last_content)

        raise last_error or ValueError("Structured output parsing failed.")

    def invoke_openai_compat(self, messages, response_format, source, **kwargs):
        kwargs.setdefault("max_tokens", 2048)
        base_url = self.cfg.openai.get("llm_api_base", None) or _openai_compat_base_url("LLM")
        request = {
            "messages": messages,
            "model": self.cfg.llm.model,
            "seed": self.cfg.llm.get("seed", 42),
            **kwargs,
        }

        provider = _chat_provider(str(self.cfg.llm.model), base_url)
        if provider == "openai":
            return invoke_openai_native(self, request, response_format)
        return invoke_json_mode(self, request, response_format)

    llm_utils.ChatCompletionModel._invoke_azure = invoke_openai_compat


def _install_optional_dependency_shims() -> None:
    """Avoid installing local-HF/GRPO dependencies for remote OpenAI-compatible use."""
    if importlib.util.find_spec("torch") is None:
        torch = types.ModuleType("torch")
        torch.__spec__ = ModuleSpec("torch", loader=None)
        torch.bfloat16 = "bfloat16"
        torch.float16 = "float16"
        torch.float32 = "float32"
        torch.manual_seed = lambda *_args, **_kwargs: None
        torch.no_grad = lambda: _NoopContext()
        torch.cuda = types.SimpleNamespace(
            is_available=lambda: False,
            manual_seed_all=lambda *_args, **_kwargs: None,
        )
        sys.modules["torch"] = torch

    if importlib.util.find_spec("transformers") is None:
        transformers = types.ModuleType("transformers")
        transformers.__spec__ = ModuleSpec("transformers", loader=None)
        transformers.AutoModelForCausalLM = _UnavailableOptionalDependency("transformers")
        transformers.AutoTokenizer = _UnavailableOptionalDependency("transformers")
        transformers.BitsAndBytesConfig = _UnavailableOptionalDependency("transformers")
        sys.modules["transformers"] = transformers

    if importlib.util.find_spec("peft") is None:
        peft = types.ModuleType("peft")
        peft.__spec__ = ModuleSpec("peft", loader=None)
        peft.PeftModel = _UnavailableOptionalDependency("peft")
        sys.modules["peft"] = peft


class _NoopContext:
    def __enter__(self):
        return self

    def __exit__(self, *_args):
        return False


class _UnavailableOptionalDependency:
    def __init__(self, package: str):
        self.package = package

    def __getattr__(self, _name: str):
        raise ImportError(f"Optional dependency '{self.package}' is not installed.")

    def __call__(self, *_args, **_kwargs):
        raise ImportError(f"Optional dependency '{self.package}' is not installed.")


def _memora_imports():
    if sys.version_info < (3, 10):
        raise RuntimeError(f"Memora requires Python 3.10 or newer. Current Python is {sys.version.split()[0]} at {sys.executable}")
    try:
        _add_default_memora_checkout_to_path()
        _install_optional_dependency_shims()
        _apply_openai_compat_patches()
        from memora.memora_client import MemoraClient
        from omegaconf import OmegaConf
    except Exception as exc:  # pragma: no cover - environment dependent
        raise RuntimeError(f"Memora is not importable in this Python environment: {exc}") from exc
    return MemoraClient, OmegaConf


def _default_home() -> Path:
    configured = os.getenv("PI_MEMORA_HOME")
    if configured:
        return Path(configured).expanduser()
    data_home = os.getenv("XDG_DATA_HOME")
    if data_home:
        return Path(data_home).expanduser() / "memora-wrapper"
    return Path.home() / ".local" / "share" / "memora-wrapper"


def _scope_id(payload: dict[str, Any]) -> str:
    scope = os.getenv("PI_MEMORA_SCOPE", "project")
    if scope == "global":
        return "pi-global"
    cwd = str(payload.get("cwd") or os.getcwd())
    digest = hashlib.sha256(cwd.encode("utf-8")).hexdigest()[:16]
    return f"pi-project-{digest}"


def _cfg(payload: dict[str, Any]):
    _, OmegaConf = _memora_imports()
    home = _default_home()
    store = home / "store"
    store.mkdir(parents=True, exist_ok=True)

    api_type = os.getenv("OPENAI_API_TYPE", "openai")

    model = os.getenv("OPENAI_MODEL", "gpt-4.1-mini")
    embedding_model = os.getenv("PI_MEMORA_EMBEDDING_MODEL", os.getenv("OPENAI_EMBEDDING_MODEL", "text-embedding-3-small"))
    collection = "pi_agent_memory"
    persist_path = str(store / collection)

    return OmegaConf.create(
        {
            "llm": {"model": model, "seed": 42},
            "openai": {
                "api_type": api_type,
                "llm_api_base": os.getenv("AZURE_OPENAI_ENDPOINT", "") if api_type == "azure" else (_openai_compat_base_url("LLM") or ""),
                "llm_api_version": os.getenv("AZURE_OPENAI_API_VERSION", "2024-12-01-preview"),
                "embedding_api_base": os.getenv("AZURE_OPENAI_ENDPOINT", "") if api_type == "azure" else (_openai_compat_base_url("EMBEDDING") or ""),
                "embedding_api_version": os.getenv("AZURE_OPENAI_API_VERSION", "2024-12-01-preview"),
                "embedding_deployment_name": os.getenv("AZURE_OPENAI_EMBEDDING_DEPLOYMENT", embedding_model),
                "managed_identity": os.getenv("AZURE_MANAGED_IDENTITY_CLIENT_ID"),
                "api_key": _openai_compat_api_key("LLM") or "pi-memora-status-placeholder",
                "embedding_api_key": _openai_compat_api_key("EMBEDDING") or "pi-memora-status-placeholder",
                "embedding_model": embedding_model,
                "model": model,
            },
            "memory": {
                "memory_store": "pi_memora",
                "persist_path": persist_path,
                "collection_name": collection,
                "distance": "cosine",
                "query_score_threshold": 0.4,
                "update_score_threshold": 0.8,
                "force_rebuild": False,
                "enhance_query": False,
                "return_history": True,
                "multimodal_support": False,
                "top_k": int(os.getenv("PI_MEMORA_TOP_K", "5")),
                "cue_top_k": 10,
                "enable_hybrid_search": True,
                "enable_segmentation": False,
                "enable_episodic_memory": False,
                "use_segments_as_episodic": False,
                "enable_cue_index": True,
            },
            "retrieval": {"strategy": "semantic"},
            "eval": {"max_workers": 2},
        }
    )


def _client(payload: dict[str, Any]):
    MemoraClient, _ = _memora_imports()
    return MemoraClient(cfg=_cfg(payload), user_id=_scope_id(payload))


def _entry_to_dict(entry: Any) -> dict[str, Any]:
    data: dict[str, Any] = {}
    for key in ("index", "value", "primary_abstraction", "cue_anchors", "metadata", "score"):
        if hasattr(entry, key):
            try:
                data[key] = getattr(entry, key)
            except Exception:
                pass
    if not data:
        data["value"] = str(entry)
    return data


def _metadata(payload: dict[str, Any]) -> dict[str, Any]:
    metadata = dict(payload.get("metadata") or {})
    for key in ("cwd", "session", "source"):
        value = payload.get(key)
        if value:
            metadata[key] = value
    return metadata


def status(payload: dict[str, Any]) -> dict[str, Any]:
    client = _client(payload)
    return {
        "ok": True,
        "user_id": _scope_id(payload),
        "count": client.count(),
        "home": str(_default_home()),
    }


def remember(payload: dict[str, Any]) -> dict[str, Any]:
    text = str(payload.get("text") or "").strip()
    if not text:
        return {"ok": False, "error": "No text provided."}
    entries = _client(payload).add(text, type=str(payload.get("type") or "doc"), metadata=_metadata(payload))
    return {"ok": True, "stored": len(entries), "entries": [_entry_to_dict(entry) for entry in entries]}


def recall(payload: dict[str, Any]) -> dict[str, Any]:
    query = str(payload.get("query") or "").strip()
    if not query:
        return {"ok": False, "error": "No query provided."}
    top_k = int(payload.get("top_k") or os.getenv("PI_MEMORA_TOP_K", "5"))
    strategy = str(payload.get("strategy") or "semantic")
    client = _client(payload)
    if strategy == "semantic":
        entries = client.query(query, top_k=top_k, enable_hybrid_search=True)
    else:
        entries = client.advance_query(query, top_k=top_k, query_type=strategy)
    return {"ok": True, "entries": [_entry_to_dict(entry) for entry in entries]}


def list_memories(payload: dict[str, Any]) -> dict[str, Any]:
    limit = int(payload.get("limit") or 20)
    entries = _client(payload).list_memories(limit=limit)
    return {"ok": True, "entries": [_entry_to_dict(entry) for entry in entries]}


def delete_memory(payload: dict[str, Any]) -> dict[str, Any]:
    key = str(payload.get("key") or "").strip()
    if not key:
        return {"ok": False, "error": "No key provided."}
    _client(payload).delete(key)
    return {"ok": True}


def clear_memories(payload: dict[str, Any]) -> dict[str, Any]:
    if payload.get("confirm") != "clear":
        return {"ok": False, "error": "Refusing to clear without confirm='clear'."}
    _client(payload).clear()
    return {"ok": True}


def setup_instructions() -> dict[str, Any]:
    return {
        "ok": False,
        "error": "Memora checkout is missing or incomplete.",
        "setup": _setup_commands(),
    }


def run_action(action: str, payload: dict[str, Any]) -> dict[str, Any]:
    if action == "missing-setup":
        return setup_instructions()
    if action == "doctor":
        return status(payload)
    if action == "add":
        return remember(payload)
    if action == "query":
        return recall(payload)
    if action == "list":
        return list_memories(payload)
    if action == "delete":
        return delete_memory(payload)
    if action == "clear":
        return clear_memories(payload)
    return {"ok": False, "error": f"Unknown action: {action}"}
