import base64
import os
import subprocess
import sys
import threading
import time
import unittest

import numpy as np

import server


HERE = os.path.dirname(os.path.abspath(__file__))


def _pcm16_b64(n_samples):
    return base64.b64encode(bytes(n_samples * 2)).decode()


def _run_isolated(testcase, code):
    """Face and body own conflicting top-level `models` packages; isolate face tests."""
    result = subprocess.run(
        [sys.executable, "-c", code], cwd=HERE, text=True,
        stdout=subprocess.PIPE, stderr=subprocess.PIPE)
    if result.returncode:
        testcase.fail(
            f"isolated test failed ({result.returncode})\n"
            f"stdout:\n{result.stdout}\nstderr:\n{result.stderr}")


class _FakeFace:
    def __init__(self, owner):
        self.owner = owner

    def push_pcm(self, pcm, emotion, n_frames, frame_start=None):
        self.owner.face_calls.append(
            (len(pcm), emotion, n_frames, frame_start))
        if self.owner.fail_face:
            raise RuntimeError("face failed")
        return np.zeros((n_frames, 2), dtype=np.float32)


class _FakeInference:
    ARKIT_52_NAMES = ("jawOpen", "mouthSmileLeft")

    def __init__(self):
        self.frame_cursor = 0
        self.face_calls = []
        self.body_calls = []
        self.cancel_calls = []
        self.fail_face = False
        self.fail_face_start = False
        self.fail_cancel = False

    def get_face_model(self):
        return object()

    def FaceStream(self, _model, scenario_id):
        if self.fail_face_start:
            raise RuntimeError("face start failed")
        return _FakeFace(self)

    def body_stream_start(self, stream_id, emotion, guidance, emit_frames):
        self.started = (stream_id, emotion, guidance, emit_frames)
        return {
            "fps": 30,
            "sample_rate": 16000,
            "emit_frames": emit_frames,
            "bones": ["hips"],
            "rotation_space": "vrm_normalized_humanoid",
        }

    def _packet(self, op):
        start = self.frame_cursor
        frames = 1
        self.frame_cursor += frames
        return {
            "op": op,
            "frame_start": start,
            "frames": frames,
            "body": [[[[0.0, 0.0, 0.0, 1.0]]]],
            "rotation_space": "vrm_normalized_humanoid",
        }

    def body_stream_chunk(self, stream_id, pcm_b64, emotion):
        self.body_calls.append(("chunk", stream_id, emotion))
        return self._packet("chunk")

    def body_stream_end(self, stream_id, pcm_b64, emotion):
        self.body_calls.append(("end", stream_id, emotion))
        return self._packet("end")

    def body_stream_cancel(self, stream_id):
        self.cancel_calls.append(stream_id)
        if self.fail_cancel:
            raise RuntimeError("cancel failed")
        return {"ok": True}


class ServerStreamLifecycleTests(unittest.TestCase):
    def setUp(self):
        self.old_infer = server._infer
        self.fake = _FakeInference()
        server._infer = self.fake
        with server._streams_lock:
            server._streams.clear()

    def tearDown(self):
        with server._streams_lock:
            server._streams.clear()
        server._infer = self.old_infer

    def test_rotation_contract_and_absolute_face_cursor_are_propagated(self):
        meta = server.start_stream("s", "neutral")
        self.assertEqual(meta["emit_frames"], 24)
        self.assertEqual(meta["rotation_space"], "vrm_normalized_humanoid")

        packet = server.stream_frames("s", _pcm16_b64(1), emotion="joy")
        self.assertEqual(packet["frame_start"], 0)
        self.assertEqual(packet["rotation_space"], "vrm_normalized_humanoid")
        self.assertEqual(packet["emotion"], "joy")
        for timing_name in ("queue_ms", "body_ms", "face_ms", "inference_ms"):
            self.assertIn(timing_name, packet)
            self.assertGreaterEqual(packet[timing_name], 0)
        self.assertEqual(self.fake.body_calls[-1], ("chunk", "s", "joy"))
        self.assertEqual(self.fake.face_calls[-1], (1, "joy", 1, 0))

    def test_face_failure_removes_both_sides_of_stream(self):
        server.start_stream("s", "neutral")
        self.fake.fail_face = True
        with self.assertRaisesRegex(RuntimeError, "face failed"):
            server.stream_frames("s", "", end=True)
        with server._streams_lock:
            self.assertNotIn("s", server._streams)
        self.assertEqual(self.fake.cancel_calls, ["s"])

    def test_start_failure_cancels_body_and_releases_reserved_id(self):
        self.fake.fail_face_start = True
        with self.assertRaisesRegex(RuntimeError, "face start failed"):
            server.start_stream("s", "neutral")
        with server._streams_lock:
            self.assertNotIn("s", server._streams)
        self.assertEqual(self.fake.cancel_calls, ["s"])

    def test_cancel_failure_still_removes_server_state(self):
        server.start_stream("s", "neutral")
        self.fake.fail_cancel = True
        with self.assertRaisesRegex(RuntimeError, "cancel failed"):
            server.cancel_stream("s")
        with server._streams_lock:
            self.assertNotIn("s", server._streams)

    def test_cancel_waits_for_inflight_chunk_before_removing_state(self):
        entered = threading.Event()
        release = threading.Event()
        original = self.fake.body_stream_chunk

        def blocking_chunk(*args, **kwargs):
            entered.set()
            self.assertTrue(release.wait(timeout=2))
            return original(*args, **kwargs)

        self.fake.body_stream_chunk = blocking_chunk
        server.start_stream("s", "neutral")

        results = {}
        chunk_thread = threading.Thread(
            target=lambda: results.setdefault(
                "chunk", server.stream_frames("s", _pcm16_b64(1))))
        chunk_thread.start()
        self.assertTrue(entered.wait(timeout=2))

        cancel_thread = threading.Thread(
            target=lambda: results.setdefault("cancel", server.cancel_stream("s")))
        cancel_thread.start()
        time.sleep(0.05)
        with server._streams_lock:
            self.assertIn("s", server._streams)
        self.assertTrue(cancel_thread.is_alive())

        release.set()
        chunk_thread.join(timeout=2)
        cancel_thread.join(timeout=2)
        self.assertFalse(chunk_thread.is_alive())
        self.assertFalse(cancel_thread.is_alive())
        self.assertTrue(results["cancel"]["cancelled"])
        with server._streams_lock:
            self.assertNotIn("s", server._streams)


class FaceStreamAlignmentTests(unittest.TestCase):
    def test_absolute_body_frame_slice_excludes_centered_extra_mel_frame(self):
        _run_isolated(self, r'''
import numpy as np
import face_v3_infer as face

inferred_sample_counts = []
old_render = face.render_face_samples

def fake_render(wav, emotion, **kwargs):
    inferred_sample_counts.append(len(wav))
    n = 1 + len(wav) // 533
    values = np.arange(n, dtype=np.float32)[:, None] / 1000.0
    return np.repeat(values, 52, axis=1), n

face.render_face_samples = fake_render
try:
    stream = face.FaceStream(model_dev=None, scenario_id="s", blend=0)
    first = stream.push_pcm(
        np.zeros(13000, np.float32), "neutral", 24, frame_start=0)
    second = stream.push_pcm(
        np.zeros(12600, np.float32), "neutral", 24, frame_start=24)
finally:
    face.render_face_samples = old_render

# First request contains 200 future/pending samples, but body represented
# frames 0..23: exactly 12,800 samples.
assert inferred_sample_counts == [12800, 25600], inferred_sample_counts
np.testing.assert_allclose(first[:, 0], np.arange(24) / 1000.0)
np.testing.assert_allclose(second[:, 0], np.arange(24, 48) / 1000.0)
assert stream.frames_emitted == 48
''')


class ProductionWarmupTests(unittest.TestCase):
    def test_warmup_uses_24_frames_and_12800_samples(self):
        _run_isolated(self, r'''
import base64
import numpy as np
import infer_turn as infer

old_request = infer._body_request
old_get_face = infer.get_face_model
old_face_stream = infer.FaceStream
old_warm = infer._STREAMING_WARM
jobs = []
face_calls = []

def fake_request(job):
    jobs.append(job)
    return {"ok": True}

class WarmFace:
    def __init__(self, model, scenario_id):
        pass

    def push_pcm(self, pcm, emotion, n_frames, frame_start=None):
        face_calls.append((len(pcm), n_frames, frame_start))
        return np.zeros((n_frames, 52), dtype=np.float32)

infer._body_request = fake_request
infer.get_face_model = lambda: object()
infer.FaceStream = WarmFace
infer._STREAMING_WARM = False
try:
    infer.warm_models()
finally:
    infer._body_request = old_request
    infer.get_face_model = old_get_face
    infer.FaceStream = old_face_stream
    infer._STREAMING_WARM = old_warm

start = next(job for job in jobs if job["op"] == "stream_start")
chunk = next(job for job in jobs if job["op"] == "stream_chunk")
assert start["emit_frames"] == 24, start
assert len(base64.b64decode(chunk["pcm_b64"])) // 2 == 12800
assert face_calls == [(12800, 24, 0)], face_calls
''')


class RealtimeSessionContractTests(unittest.TestCase):
    def test_emotion_tool_is_required_before_the_spoken_continuation(self):
        captured = {}
        old_post = server.requests.post
        old_key = os.environ.get("OPENAI_API_KEY")

        class _Response:
            def raise_for_status(self):
                return None

            def json(self):
                return {"value": "ephemeral-test-token", "expires_at": 123}

        def fake_post(url, headers, json, timeout):
            captured["url"] = url
            captured["body"] = json
            return _Response()

        server.requests.post = fake_post
        os.environ["OPENAI_API_KEY"] = "test-only"
        try:
            result = server.mint_ephemeral_token()
        finally:
            server.requests.post = old_post
            if old_key is None:
                os.environ.pop("OPENAI_API_KEY", None)
            else:
                os.environ["OPENAI_API_KEY"] = old_key

        session = captured["body"]["session"]
        self.assertEqual(session["tool_choice"], "required")
        self.assertEqual([tool["name"] for tool in session["tools"]], ["set_emotion"])
        self.assertIn("음성을 생성하지 말고", session["instructions"])
        self.assertEqual(result["value"], "ephemeral-test-token")


if __name__ == "__main__":
    unittest.main()
