"""server.py - 아누 공용 임베딩 HTTP 서비스 (FastAPI + sentence-transformers)

포트: 8300 (8200=whisper-gpu, 8100=사용중)
모델: jhgan/ko-sroberta-multitask (HF 캐시 로컬, device=cpu, dim=768)

설계 원칙
- 외부 API 키 0. 완전 로컬 추론.
- 모델은 startup(lifespan) 시 **1회만** 로드하여 전역 보관한다. 요청마다 로드 금지.
- 출력 차원 768 불변식은 fail-closed. 다르면 조용히 자르거나 패딩하지 않고 500 으로 실패시킨다.
- GPU 사용 금지 (이 머신 GPU=sm_61, torch 2.10=sm_70+ 만 지원 → CUDA 커널 부재).
"""

from __future__ import annotations

import asyncio
import contextlib
import logging
import math
import time
from collections.abc import AsyncGenerator
from typing import Any

from fastapi import FastAPI, HTTPException
from fastapi.responses import JSONResponse
from pydantic import BaseModel, Field

from settings import (
    DEVICE,
    EMBEDDING_DIM,
    HOST,
    MAX_BATCH_SIZE,
    MAX_TEXT_CHARS,
    MODEL_NAME,
    PORT,
)

# ---------------------------------------------------------------------------
# 로깅 설정 (whisper-gpu 관례와 동일)
# ---------------------------------------------------------------------------

logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
)
logger = logging.getLogger("anu-embedding")

# ---------------------------------------------------------------------------
# 전역 상태 — 모델 단일 인스턴스
# ---------------------------------------------------------------------------

# load_count 는 "모델이 실제로 몇 번 로드되었는가"를 나타내는 관측 값이다.
# 정상 동작 시 서버 생애주기 동안 정확히 1 이어야 한다 (테스트에서 검증).
_state: dict[str, Any] = {
    "model": None,
    "load_count": 0,
    "loaded_at": None,
    "dimension": None,
}

# CPU 추론 직렬화: 동시 요청이 CPU 스레드를 과점유하지 않도록 보호
_encode_lock = asyncio.Lock()


def load_model() -> Any:
    """SentenceTransformer 를 CPU 로 1회 로드하고 차원 불변식을 검증합니다.

    Raises:
        RuntimeError: 모델이 보고하는 임베딩 차원이 EMBEDDING_DIM 과 다를 때.
    """
    from sentence_transformers import SentenceTransformer  # noqa: PLC0415

    logger.info("모델 로딩 시작: model=%s device=%s", MODEL_NAME, DEVICE)
    started = time.time()
    model = SentenceTransformer(MODEL_NAME, device=DEVICE)
    elapsed = time.time() - started

    reported_dim = model.get_sentence_embedding_dimension()
    if reported_dim != EMBEDDING_DIM:
        logger.error(
            "차원 불변식 위반(모델 로드): expected=%d reported=%s model=%s",
            EMBEDDING_DIM,
            reported_dim,
            MODEL_NAME,
        )
        raise RuntimeError(
            f"모델 차원 불변식 위반: expected={EMBEDDING_DIM} reported={reported_dim} model={MODEL_NAME}"
        )

    _state["model"] = model
    _state["load_count"] = int(_state["load_count"]) + 1
    _state["loaded_at"] = time.time()
    _state["dimension"] = reported_dim
    logger.info(
        "모델 로딩 완료: model=%s device=%s dim=%d elapsed=%.2fs load_count=%d",
        MODEL_NAME,
        DEVICE,
        reported_dim,
        elapsed,
        _state["load_count"],
    )
    return model


def get_model() -> Any:
    """로드된 전역 모델을 반환합니다. 요청 시점 로딩은 하지 않습니다(fail-closed).

    Raises:
        HTTPException: 503 — 모델이 아직 로드되지 않은 경우.
    """
    model = _state["model"]
    if model is None:
        logger.error("모델 미로딩 상태에서 요청 수신 — 503 반환")
        raise HTTPException(status_code=503, detail="모델이 로드되지 않았습니다. 서비스 준비 중입니다.")
    return model


# ---------------------------------------------------------------------------
# FastAPI 앱
# ---------------------------------------------------------------------------


@contextlib.asynccontextmanager
async def lifespan(application: FastAPI) -> AsyncGenerator[None, None]:
    """startup 시 모델을 1회 로드한다. 로드 실패 시 서버는 기동하지 않는다."""
    try:
        load_model()
    except Exception as exc:  # noqa: BLE001 - 기동 실패는 명시적으로 로깅 후 재전파
        logger.exception("startup 모델 로딩 실패: %s", exc)
        raise
    yield
    _state["model"] = None
    logger.info("서버 종료: 모델 해제")


app = FastAPI(
    title="ANU Embedding Service",
    description="sentence-transformers 기반 로컬 한국어 임베딩 서비스 (외부 API 키 불필요)",
    version="1.0.0",
    lifespan=lifespan,
)


# ---------------------------------------------------------------------------
# 요청/응답 모델
# ---------------------------------------------------------------------------


class EmbedRequest(BaseModel):
    texts: list[str] = Field(..., description="임베딩할 텍스트 목록 (1개 이상)")
    normalize: bool = Field(default=True, description="L2 정규화 여부")


# ---------------------------------------------------------------------------
# 입력 검증 / 불변식 검증
# ---------------------------------------------------------------------------


def validate_texts(texts: Any) -> list[str]:
    """입력 텍스트를 검증하고 MAX_TEXT_CHARS 로 truncate 한 목록을 반환합니다.

    Raises:
        HTTPException: 422 — 빈 목록, 비문자열, 전부 공백, 배치 초과.
    """
    if not isinstance(texts, list):
        logger.error("입력 검증 실패: texts 가 list 가 아님 (type=%s)", type(texts).__name__)
        raise HTTPException(status_code=422, detail="texts 는 문자열 배열이어야 합니다.")

    if len(texts) == 0:
        logger.error("입력 검증 실패: texts 가 비어 있음")
        raise HTTPException(status_code=422, detail="texts 가 비어 있습니다. 1개 이상의 문자열이 필요합니다.")

    if len(texts) > MAX_BATCH_SIZE:
        logger.error("입력 검증 실패: 배치 초과 (%d > %d)", len(texts), MAX_BATCH_SIZE)
        raise HTTPException(
            status_code=422,
            detail=f"texts 개수가 최대 배치 크기를 초과했습니다: {len(texts)} > {MAX_BATCH_SIZE}",
        )

    cleaned: list[str] = []
    truncated = 0
    for idx, item in enumerate(texts):
        if not isinstance(item, str):
            logger.error("입력 검증 실패: texts[%d] 가 문자열이 아님 (type=%s)", idx, type(item).__name__)
            raise HTTPException(status_code=422, detail=f"texts[{idx}] 가 문자열이 아닙니다.")
        if not item.strip():
            logger.error("입력 검증 실패: texts[%d] 가 공백 문자열", idx)
            raise HTTPException(status_code=422, detail=f"texts[{idx}] 가 빈 문자열입니다.")
        if len(item) > MAX_TEXT_CHARS:
            truncated += 1
            item = item[:MAX_TEXT_CHARS]
        cleaned.append(item)

    if truncated:
        logger.warning("입력 truncate 적용: %d건 (limit=%d자)", truncated, MAX_TEXT_CHARS)
    return cleaned


def validate_embeddings(embeddings: Any, expected_count: int) -> list[list[float]]:
    """인코딩 결과의 차원 불변식을 강제합니다 (fail-closed).

    조용한 truncate/padding 을 하지 않고, 위반 시 즉시 500 으로 실패시킵니다.

    Raises:
        HTTPException: 500 — 개수 불일치, 차원 불일치, 비유한값(NaN/Inf) 포함.
    """
    try:
        rows = [list(row) for row in embeddings]
    except TypeError as exc:
        logger.error("차원 불변식 위반: 인코딩 결과가 순회 가능한 2차원 구조가 아님 (%s)", exc)
        raise HTTPException(status_code=500, detail="인코딩 결과 형식이 올바르지 않습니다.") from exc

    if len(rows) != expected_count:
        logger.error("차원 불변식 위반: 벡터 개수 불일치 expected=%d actual=%d", expected_count, len(rows))
        raise HTTPException(
            status_code=500,
            detail=f"인코딩 결과 개수 불일치: expected={expected_count} actual={len(rows)}",
        )

    result: list[list[float]] = []
    for idx, row in enumerate(rows):
        if len(row) != EMBEDDING_DIM:
            logger.error(
                "차원 불변식 위반: embeddings[%d] dim=%d (expected=%d) — fail-closed",
                idx,
                len(row),
                EMBEDDING_DIM,
            )
            raise HTTPException(
                status_code=500,
                detail=(
                    f"임베딩 차원 불변식 위반: embeddings[{idx}] dimension={len(row)} "
                    f"expected={EMBEDDING_DIM}"
                ),
            )
        try:
            floats = [float(v) for v in row]
        except (TypeError, ValueError) as exc:
            logger.error("차원 불변식 위반: embeddings[%d] 에 숫자가 아닌 값 포함 (%s)", idx, exc)
            raise HTTPException(
                status_code=500,
                detail=f"임베딩 값이 숫자가 아닙니다: embeddings[{idx}]",
            ) from exc
        if any(not math.isfinite(v) for v in floats):
            logger.error("차원 불변식 위반: embeddings[%d] 에 NaN/Inf 포함 — fail-closed", idx)
            raise HTTPException(
                status_code=500,
                detail=f"임베딩에 유한하지 않은 값(NaN/Inf)이 포함되었습니다: embeddings[{idx}]",
            )
        result.append(floats)

    return result


def encode_texts(texts: list[str], normalize: bool) -> list[list[float]]:
    """전역 모델로 인코딩한 뒤 불변식 검증을 통과한 벡터를 반환합니다."""
    model = get_model()
    try:
        raw = model.encode(
            texts,
            normalize_embeddings=normalize,
            convert_to_numpy=True,
            show_progress_bar=False,
        )
    except HTTPException:
        raise
    except Exception as exc:  # noqa: BLE001
        logger.exception("인코딩 실패: %s", exc)
        raise HTTPException(status_code=500, detail=f"인코딩 실패: {exc}") from exc

    return validate_embeddings(raw, expected_count=len(texts))


# ---------------------------------------------------------------------------
# 엔드포인트
# ---------------------------------------------------------------------------


@app.get("/health")
async def health() -> JSONResponse:
    """서비스 상태와 모델 로딩 상태를 반환합니다."""
    loaded = _state["model"] is not None
    return JSONResponse(
        content={
            "status": "ok" if loaded else "loading",
            "model": MODEL_NAME,
            "dimension": EMBEDDING_DIM,
            "loaded": loaded,
            "device": DEVICE,
            "load_count": _state["load_count"],
            "loaded_at": _state["loaded_at"],
            "max_text_chars": MAX_TEXT_CHARS,
            "max_batch_size": MAX_BATCH_SIZE,
        }
    )


@app.post("/embed")
async def embed(request: EmbedRequest) -> JSONResponse:
    """텍스트 목록을 임베딩합니다.

    Request:
        {"texts": ["...", ...], "normalize": true}

    Response:
        {"model": "...", "dimension": 768, "embeddings": [[768 floats], ...]}
    """
    texts = validate_texts(request.texts)

    started = time.time()
    loop = asyncio.get_event_loop()
    async with _encode_lock:
        embeddings = await loop.run_in_executor(None, encode_texts, texts, request.normalize)
    elapsed = time.time() - started

    logger.info(
        "임베딩 완료: count=%d normalize=%s dim=%d elapsed=%.3fs",
        len(embeddings),
        request.normalize,
        EMBEDDING_DIM,
        elapsed,
    )
    return JSONResponse(
        content={
            "model": MODEL_NAME,
            "dimension": EMBEDDING_DIM,
            "embeddings": embeddings,
        }
    )


# ---------------------------------------------------------------------------
# 메인 실행 (직접 실행 시)
# ---------------------------------------------------------------------------

if __name__ == "__main__":
    import uvicorn

    uvicorn.run(
        "server:app",
        host=HOST,
        port=PORT,
        log_level="info",
        workers=1,  # 모델 1회 로드 보장 (멀티워커 금지)
    )
