import os
from datetime import datetime, timezone
from typing import Optional
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from graphiti_core import Graphiti
from graphiti_core.nodes import EpisodeType
from graphiti_core.llm_client import OpenAIClient
from graphiti_core.llm_client.config import LLMConfig
from graphiti_core.embedder import EmbedderClient
from graphiti_core.cross_encoder.client import CrossEncoderClient
from fastembed import TextEmbedding
from typing import Iterable, List, Tuple

# --- Monkey-patch: LLM gibt manchmal List statt Dict zurueck ---
import graphiti_core.utils.maintenance.temporal_operations as _temporal
_original_get_edge_contradictions = _temporal.get_edge_contradictions.__wrapped__ if hasattr(_temporal.get_edge_contradictions, "__wrapped__") else None

import graphiti_core.utils.maintenance.node_operations as _nodeops

def _safe_get(response, key, default):
    if isinstance(response, list):
        return response
    return response.get(key, default)

# Patch temporal_operations direkt
import graphiti_core.utils.maintenance.temporal_operations as _to_module
_to_src = open(_to_module.__file__).read()
if "isinstance(llm_response, list)" not in _to_src:
    _to_src = _to_src.replace(
        "contradicted_edge_data = llm_response.get('invalidated_edges', [])",
        "contradicted_edge_data = llm_response if isinstance(llm_response, list) else llm_response.get('invalidated_edges', [])"
    )
    open(_to_module.__file__, "w").write(_to_src)
    print("Patch angewendet: temporal_operations.py")

# Patch graphiti.py: last_n=3 → last_n=0 (verhindert wachsenden Kontext)
import graphiti_core.graphiti as _graphiti_module
_g_src = open(_graphiti_module.__file__).read()
if "last_n=0" not in _g_src and "last_n=3" in _g_src:
    _g_src = _g_src.replace("last_n=3", "last_n=0")
    open(_graphiti_module.__file__, "w").write(_g_src)
    print("Patch angewendet: graphiti.py last_n=0")


class LocalEmbedder(EmbedderClient):
    def __init__(self, model_name: str = "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2"):
        self._model = TextEmbedding(model_name=model_name)

    async def create(self, input_data) -> list[float]:
        if isinstance(input_data, str):
            texts = [input_data]
        elif isinstance(input_data, list) and all(isinstance(x, str) for x in input_data):
            texts = input_data
        else:
            texts = [str(input_data)]
        embeddings = list(self._model.embed(texts))
        return embeddings[0].tolist()

class PassthroughCrossEncoder(CrossEncoderClient):
    async def rank(self, query: str, passages: List[str]) -> List[Tuple[str, float]]:
        return [(p, 1.0) for p in passages]

app = FastAPI(title="Graphiti Knowledge Service", version="1.0.0")
app.add_middleware(CORSMiddleware, allow_origins=["*"], allow_methods=["*"], allow_headers=["*"])

NEO4J_URI      = os.getenv("NEO4J_URI", "bolt://graphiti-neo4j:7687")
NEO4J_USER     = os.getenv("NEO4J_USER", "neo4j")
NEO4J_PASSWORD = os.getenv("NEO4J_PASSWORD", "graphiti2026")
LLM_API_KEY    = os.getenv("DEEPSEEK_API_KEY", "")
LLM_BASE_URL   = os.getenv("LLM_BASE_URL", "https://api.deepseek.com/v1")
LLM_MODEL      = os.getenv("LLM_MODEL", "deepseek/deepseek-chat")
EMBED_MODEL    = os.getenv("EMBED_MODEL", "intfloat/multilingual-e5-small")

graphiti_client: Optional[Graphiti] = None

@app.on_event("startup")
async def startup():
    global graphiti_client
    llm = OpenAIClient(config=LLMConfig(
        api_key=LLM_API_KEY, model=LLM_MODEL, base_url=LLM_BASE_URL, max_tokens=2048
    ))
    embedder = LocalEmbedder(model_name=EMBED_MODEL)
    graphiti_client = Graphiti(NEO4J_URI, NEO4J_USER, NEO4J_PASSWORD,
                               llm_client=llm, embedder=embedder,
                               cross_encoder=PassthroughCrossEncoder())
    await graphiti_client.build_indices_and_constraints()
    print("Graphiti bereit.")

class Episode(BaseModel):
    content: str
    source: str = "manual"
    actor: str = "system"
    context: str = ""
    timestamp: Optional[str] = None

class SearchRequest(BaseModel):
    query: str
    limit: int = 10

@app.get("/health")
def health():
    return {"status": "ok", "service": "graphiti"}

@app.post("/episodes")
async def add_episode(ep: Episode):
    if not graphiti_client:
        raise HTTPException(503, "Graphiti nicht initialisiert")
    ts = datetime.fromisoformat(ep.timestamp) if ep.timestamp else datetime.now(timezone.utc)
    await graphiti_client.add_episode(
        name=f"{ep.source}_{ep.actor}_{int(ts.timestamp())}",
        episode_body=ep.content,
        source_description=f"{ep.source} | {ep.context}",
        reference_time=ts,
        source=EpisodeType.text,
    )
    return {"status": "gespeichert", "source": ep.source}

@app.post("/search")
async def search(req: SearchRequest):
    if not graphiti_client:
        raise HTTPException(503, "Graphiti nicht initialisiert")
    results = await graphiti_client.search(req.query, num_results=req.limit)
    normalized = []
    for r in results:
        fact = getattr(r, "fact", None)
        if fact is None:
            fact = getattr(r, "name", None) or getattr(r, "content", None) or str(r)
        item = {"fact": fact}
        score = getattr(r, "score", None)
        if score is not None:
            item["score"] = score
        normalized.append(item)
    return {"results": normalized}

@app.get("/facts/{entity}")
async def get_facts(entity: str):
    if not graphiti_client:
        raise HTTPException(503, "Graphiti nicht initialisiert")
    results = await graphiti_client.search(entity, num_results=20)
    return {"entity": entity, "facts": [r.fact for r in results]}
