"""test_server.py - 아누 공용 임베딩 서비스 pytest.

검증 항목
- /health 응답 스키마
- /embed 가 정확히 768차원을 반환
- 결정론성: 같은 입력 → 같은 벡터
- 빈 입력/비문자열/배치초과 거부 (422)
- 차원 불변식 fail-closed: 768 이 아니면 500 (조용한 truncate/padding 금지)
- 모델이 서버 생애주기 동안 **1회만** 로드되는지
"""

from __future__ import annotations

import os
import sys

import pytest
from fastapi.testclient import TestClient

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))

import server  # noqa: E402
from settings import EMBEDDING_DIM, MAX_BATCH_SIZE, MAX_TEXT_CHARS, MODEL_NAME  # noqa: E402


@pytest.fixture(scope="module")
def client():
    """TestClient 를 context manager 로 사용 → lifespan(startup) 1회 실행."""
    with TestClient(server.app) as test_client:
        yield test_client


class _FakeModel:
    """차원 위반을 재현하기 위한 가짜 모델 (정상 포맷의 잘못된 차원)."""

    def __init__(self, dim: int):
        self.dim = dim

    def encode(self, texts, **kwargs):  # noqa: ANN001, ANN003
        return [[0.1] * self.dim for _ in texts]


class _NaNModel:
    def encode(self, texts, **kwargs):  # noqa: ANN001, ANN003
        return [[float("nan")] * EMBEDDING_DIM for _ in texts]


class _WrongCountModel:
    def encode(self, texts, **kwargs):  # noqa: ANN001, ANN003
        return [[0.1] * EMBEDDING_DIM]  # 입력 개수와 무관하게 1개만 반환


# ---------------------------------------------------------------------------
# health
# ---------------------------------------------------------------------------


def test_health_ok(client):
    resp = client.get("/health")
    assert resp.status_code == 200
    body = resp.json()
    assert body["status"] == "ok"
    assert body["model"] == MODEL_NAME
    assert body["dimension"] == 768
    assert body["loaded"] is True
    assert body["device"] == "cpu"


def test_health_dimension_is_768_constant():
    assert EMBEDDING_DIM == 768


# ---------------------------------------------------------------------------
# /embed 정상 경로
# ---------------------------------------------------------------------------


def test_embed_returns_768_dimensions(client):
    resp = client.post("/embed", json={"texts": ["안녕하세요", "실손의료비 보장"], "normalize": True})
    assert resp.status_code == 200
    body = resp.json()
    assert body["model"] == MODEL_NAME
    assert body["dimension"] == 768
    assert len(body["embeddings"]) == 2
    for vec in body["embeddings"]:
        assert len(vec) == 768
        assert all(isinstance(v, float) for v in vec)


def test_embed_normalize_true_gives_unit_norm(client):
    resp = client.post("/embed", json={"texts": ["보험 특약 설명"], "normalize": True})
    assert resp.status_code == 200
    vec = resp.json()["embeddings"][0]
    norm = sum(v * v for v in vec) ** 0.5
    assert abs(norm - 1.0) < 1e-4


def test_embed_normalize_default_is_true(client):
    """normalize 를 생략하면 기본 True 로 동작해야 한다."""
    resp = client.post("/embed", json={"texts": ["기본값 확인"]})
    assert resp.status_code == 200
    vec = resp.json()["embeddings"][0]
    norm = sum(v * v for v in vec) ** 0.5
    assert abs(norm - 1.0) < 1e-4


def test_embed_deterministic_same_input_same_vector(client):
    payload = {"texts": ["결정론성 확인 문장입니다"], "normalize": True}
    first = client.post("/embed", json=payload).json()["embeddings"][0]
    second = client.post("/embed", json=payload).json()["embeddings"][0]
    assert first == second


def test_embed_different_texts_give_different_vectors(client):
    resp = client.post("/embed", json={"texts": ["암 진단비", "자동차 보험료"], "normalize": True})
    vecs = resp.json()["embeddings"]
    assert vecs[0] != vecs[1]


def test_embed_truncates_long_text_without_error(client):
    long_text = "가" * (MAX_TEXT_CHARS + 5000)
    resp = client.post("/embed", json={"texts": [long_text], "normalize": True})
    assert resp.status_code == 200
    assert len(resp.json()["embeddings"][0]) == 768


# ---------------------------------------------------------------------------
# 입력 검증
# ---------------------------------------------------------------------------


def test_embed_empty_list_rejected(client):
    resp = client.post("/embed", json={"texts": [], "normalize": True})
    assert resp.status_code == 422


def test_embed_blank_string_rejected(client):
    resp = client.post("/embed", json={"texts": ["   "], "normalize": True})
    assert resp.status_code == 422


def test_embed_missing_texts_field_rejected(client):
    resp = client.post("/embed", json={"normalize": True})
    assert resp.status_code == 422


def test_embed_non_string_item_rejected(client):
    resp = client.post("/embed", json={"texts": [123], "normalize": True})
    assert resp.status_code == 422


def test_embed_batch_over_limit_rejected(client):
    resp = client.post("/embed", json={"texts": ["가"] * (MAX_BATCH_SIZE + 1), "normalize": True})
    assert resp.status_code == 422


# ---------------------------------------------------------------------------
# 차원 불변식 fail-closed
# ---------------------------------------------------------------------------


@pytest.mark.parametrize("bad_dim", [384, 767, 769, 1536])
def test_embed_wrong_dimension_fails_closed(client, monkeypatch, bad_dim):
    """768 이 아닌 결과는 조용히 자르거나 패딩하지 않고 500 으로 실패해야 한다."""
    monkeypatch.setitem(server._state, "model", _FakeModel(bad_dim))
    resp = client.post("/embed", json={"texts": ["차원 위반 확인"], "normalize": True})
    assert resp.status_code == 500
    assert "차원 불변식 위반" in resp.json()["detail"]


def test_embed_nan_values_fail_closed(client, monkeypatch):
    monkeypatch.setitem(server._state, "model", _NaNModel())
    resp = client.post("/embed", json={"texts": ["NaN 확인"], "normalize": True})
    assert resp.status_code == 500


def test_embed_count_mismatch_fails_closed(client, monkeypatch):
    monkeypatch.setitem(server._state, "model", _WrongCountModel())
    resp = client.post("/embed", json={"texts": ["a", "b"], "normalize": True})
    assert resp.status_code == 500


def test_validate_embeddings_accepts_exact_768():
    rows = [[0.0] * 768]
    out = server.validate_embeddings(rows, expected_count=1)
    assert len(out[0]) == 768


def test_model_not_loaded_returns_503(client, monkeypatch):
    """모델 미로딩 시 요청 시점 로딩으로 폴백하지 않고 503 을 반환한다."""
    monkeypatch.setitem(server._state, "model", None)
    resp = client.post("/embed", json={"texts": ["미로딩 확인"], "normalize": True})
    assert resp.status_code == 503


def test_load_model_rejects_wrong_dimension(monkeypatch):
    """모델이 보고하는 차원이 768 이 아니면 기동 자체가 실패해야 한다."""

    class _BadDimModel:
        def get_sentence_embedding_dimension(self):
            return 384

    import sentence_transformers

    monkeypatch.setattr(
        sentence_transformers, "SentenceTransformer", lambda *a, **k: _BadDimModel()
    )
    with pytest.raises(RuntimeError, match="차원 불변식 위반"):
        server.load_model()


# ---------------------------------------------------------------------------
# 모델 단일 로드 검증
# ---------------------------------------------------------------------------


def test_model_loaded_exactly_once(client):
    """여러 요청을 보내도 모델 로드는 startup 시 1회뿐이어야 한다."""
    before = server._state["load_count"]
    assert before == 1, f"startup 후 load_count 는 1 이어야 함 (actual={before})"
    for _ in range(5):
        assert client.post("/embed", json={"texts": ["로드 카운트 확인"]}).status_code == 200
    assert server._state["load_count"] == 1


def test_model_instance_is_stable_across_requests(client):
    """요청 간 동일한 모델 객체(identity)가 재사용되어야 한다."""
    first = server._state["model"]
    client.post("/embed", json={"texts": ["동일 인스턴스 확인"]})
    assert server._state["model"] is first


def test_no_external_api_key_required():
    """외부 API 키 환경변수 없이도 모듈 상수가 유효하다 (외부 API 키 0)."""
    for key in ("GEMINI_API_KEY", "OPENAI_API_KEY", "ANTHROPIC_API_KEY"):
        os.environ.pop(key, None)
    assert MODEL_NAME == "jhgan/ko-sroberta-multitask"
    assert EMBEDDING_DIM == 768
