import asyncio
import json
from unittest.mock import patch, MagicMock, AsyncMock, AsyncMock

def _run(coro):
    return asyncio.run(coro)

def _mk_usage():
    u = MagicMock(); u.prompt_tokens = 10; u.completion_tokens = 5; u.total_tokens = 15
    return u

def _text_response(text):
    msg = MagicMock(); msg.content = text; msg.tool_calls = None
    resp = MagicMock(); resp.choices = [MagicMock(message=msg)]; resp.usage = _mk_usage()
    return resp

def _tool_response(name, args, call_id="call_1"):
    tc = MagicMock(); tc.id = call_id
    tc.function = MagicMock(); tc.function.name = name
    tc.function.arguments = json.dumps(args)
    msg = MagicMock(); msg.content = None; msg.tool_calls = [tc]
    resp = MagicMock(); resp.choices = [MagicMock(message=msg)]; resp.usage = _mk_usage()
    return resp

# All of chat.py's collaborators are now module-level imports, so they patch on,
# so they must be patched at app.services.infra.database, not app.services.chat.
# There is deliberately NO get_embeddings/search_vector_data patch here: chat.py no
# longer does retrieval itself. Retrieval only happens if the model calls the
# search_knowledge_base tool, which routes through dispatch_tool.
COMMON_PATCHES = {
    "app.services.chat.chat.check_quota": MagicMock(),
    "app.services.chat.chat.estimate_text_tokens": MagicMock(return_value=10),
    "app.services.chat.chat.get_chat_history": MagicMock(return_value=[]),
    "app.services.chat.chat.save_chat_message": MagicMock(),
    "app.services.chat.chat.get_takeover_status": MagicMock(return_value=False),
    "app.services.chat.chat.get_ticket_by_thread": MagicMock(return_value=None),
    "app.services.chat.chat.insert_llm_usage": MagicMock(),
    "app.services.chat.chat.insert_feedback": MagicMock(),
    "app.services.chat.chat.get_latest_summary": MagicMock(return_value=None),
    "app.services.chat.chat.insert_summary": MagicMock(),
}

def _patch_chat(**overrides):
    """Patch all chat.py collaborators; overrides are patched on app.services.chat.chat."""
    import contextlib
    targets = dict(COMMON_PATCHES)
    for name, mock in overrides.items():
        targets[f"app.services.chat.chat.{name}"] = mock
    stack = contextlib.ExitStack()
    for target, mock in targets.items():
        stack.enter_context(patch(target, mock))
    return stack

def test_plain_text_answer_no_tools():
    from app.services.chat import chat
    with _patch_chat():
        with patch("app.services.chat.chat.openai_client") as llm:
            llm.chat.completions.create.return_value = _text_response("Hello there")
            result = _run(chat.get_chat_response("t1", "hi", "th1", user_id="u1"))
    assert result["text"] == "Hello there"
    assert result["resources"] == []

def test_tool_call_then_final_answer():
    from app.services.chat import chat
    with _patch_chat():
        with patch("app.services.chat.chat.openai_client") as llm, \
             patch("app.services.chat.chat.dispatch_tool", new=AsyncMock(return_value={"status": "created", "ticket_id": "TICK-1"})) as disp:
            llm.chat.completions.create.side_effect = [
                _tool_response("create_or_update_ticket", {"heading": "H", "content": "C", "priority": "Low"}),
                _text_response("Your ticket TICK-1 is logged."),
            ]
            result = _run(chat.get_chat_response("t1", "my login is broken", "th1", user_id="u1"))
    assert "TICK-1" in result["text"]
    assert disp.await_count == 1
    assert llm.chat.completions.create.call_count == 2
    # tools offered on every call
    assert llm.chat.completions.create.call_args_list[0].kwargs["tools"]

def test_attach_resources_populates_result():
    from app.services.chat import chat
    res = [{"url": "https://x.com/p", "image": None, "label": "Pricing"}]
    with _patch_chat():
        with patch("app.services.chat.chat.openai_client") as llm:
            llm.chat.completions.create.side_effect = [
                _tool_response("attach_resources", {"resources": res}),
                _text_response("Here is the pricing page."),
            ]
            result = _run(chat.get_chat_response("t1", "pricing?", "th1", user_id="u1"))
    assert result["resources"] == res

def test_loop_bounded_at_six_rounds():
    from app.services.chat import chat
    assert chat.MAX_TOOL_ROUNDS == 6
    with _patch_chat():
        with patch("app.services.chat.chat.openai_client") as llm, \
             patch("app.services.chat.chat.dispatch_tool", new=AsyncMock(return_value={"status": "saved"})):
            llm.chat.completions.create.return_value = _tool_response("submit_feedback", {"question": "q", "answer": "a"})
            result = _run(chat.get_chat_response("t1", "hi", "th1", user_id="u1"))
    assert llm.chat.completions.create.call_count == 6
    assert result["text"]  # fallback text, never empty

def test_quota_checked_every_round():
    """A multi-round turn must re-check the token quota, not just once up front,
    and each in-loop check must happen BEFORE its corresponding LLM call."""
    from app.services.chat import chat
    call_order = []
    quota = MagicMock(side_effect=lambda *a, **k: call_order.append(("check_quota", a[1])))
    with _patch_chat(check_quota=quota):
        with patch("app.services.chat.chat.openai_client") as llm, \
             patch("app.services.chat.chat.dispatch_tool", new=AsyncMock(return_value={"status": "saved"})):
            llm.chat.completions.create.side_effect = lambda *a, **k: (
                call_order.append(("create", None)),
                _tool_response("submit_feedback", {"question": "q", "answer": "a"}),
            )[1]
            _run(chat.get_chat_response("t1", "hi", "th1", user_id="u1"))
    ai_token_checks = [c for c in quota.call_args_list if c[0][1] == "ai_tokens"]
    # 1 early check (pre-history) + 1 pre-call projection + 5 in-loop rounds (rounds 1..5)
    assert len(ai_token_checks) == 7
    # Verify ordering: each in-loop check_quota happens before its corresponding create call.
    # The tagged call_order should show check_quota entries preceding create entries in sequence,
    # i.e. no "create" immediately follows another "create" without an intervening check
    # for rounds 1 and 2 (round 0 has no in-loop check, by design).
    create_indices = [i for i, (kind, _) in enumerate(call_order) if kind == "create"]
    assert len(create_indices) == 6
    # Round 1 and round 2's create calls (2nd and 3rd) must be preceded by a check_quota
    # that occurs after the previous create call.
    for round_num, create_idx in enumerate(create_indices):
        if round_num == 0:
            continue
        prev_create_idx = create_indices[round_num - 1]
        checks_between = [
            kind for kind, _ in call_order[prev_create_idx + 1:create_idx] if kind == "check_quota"
        ]
        assert checks_between, f"round {round_num} create call has no preceding in-loop check_quota"

def test_quota_exceeded_in_loop_propagates():
    from app.services.chat import chat
    from app.services.infra.quotas import QuotaExceededError
    # 5 non-loop/pre-loop checks occur before the loop's round_index==1 check:
    # (1) early ai_tokens check, (2) ai_conversations (new thread), (3) ai_responses,
    # (4) pre-loop ai_tokens projection, then the loop's round_index==0 has no check_quota
    # (round_index > 0 guard), so the 5th call is the in-loop round_index==1 check.
    quota = MagicMock(side_effect=[None, None, None, None, QuotaExceededError("cap hit")])
    with _patch_chat(check_quota=quota):
        with patch("app.services.chat.chat.openai_client") as llm, \
             patch("app.services.chat.chat.dispatch_tool", new=AsyncMock(return_value={"status": "saved"})):
            llm.chat.completions.create.return_value = _tool_response("submit_feedback", {"question": "q", "answer": "a"})
            try:
                _run(chat.get_chat_response("t1", "hi", "th1", user_id="u1"))
                raised = False
            except QuotaExceededError:
                raised = True
    assert raised, "QuotaExceededError must propagate out of the tool loop"
    # The loop must have run exactly one round (round_index==0, no in-loop check) before
    # the round_index==1 in-loop check raised — proving the loop actually executed and
    # did not swallow the exception.
    assert llm.chat.completions.create.call_count == 1

def test_no_forced_retrieval_before_the_model_sees_the_question():
    """chat.py must not embed KB context in the user message any more, and must not
    run any retrieval on its own — the model asks for it or it does not happen."""
    from app.services.chat import chat
    with _patch_chat():
        with patch("app.services.chat.chat.openai_client") as llm, \
             patch("app.services.chat.tools.get_embeddings") as embed, \
             patch("app.services.chat.tools.search_vector_data") as search:
            llm.chat.completions.create.return_value = _text_response("Hi!")
            _run(chat.get_chat_response("t1", "how much is brandforge?", "th1", user_id="u1"))
            sent = llm.chat.completions.create.call_args.kwargs["messages"]
    embed.assert_not_called()
    search.assert_not_called()
    assert sent[-1] == {"role": "user", "content": "how much is brandforge?"}
    assert "Relevant Context" not in sent[-1]["content"]


def test_two_search_conversation_feeds_results_back_into_the_model():
    """The agentic property: a search result the model receives in round 2 must be
    visible to it in round 3, and a fruitless search must not end the loop."""
    from app.services.chat import chat
    fact = "BrandForge costs $0 — free forever."
    hits = [{"content": fact, "metadata": {"source": "https://galaxiq.ai/pricing",
                                           "top_images": [{"url": "https://galaxiq.ai/bf.png"}]}}]
    with _patch_chat():
        with patch("app.services.chat.chat.openai_client") as llm, \
             patch("app.services.chat.tools.get_embeddings", return_value=[0.1]), \
             patch("app.services.chat.tools.search_vector_data", side_effect=[[], hits]) as search:
            llm.chat.completions.create.side_effect = [
                _tool_response("search_knowledge_base", {"query": "brandforge"}, "c1"),
                _tool_response("search_knowledge_base", {"query": "brandforge pricing plan"}, "c2"),
                _tool_response("attach_resources", {"resources": [
                    {"url": "https://galaxiq.ai/pricing", "image": "https://galaxiq.ai/bf.png",
                     "label": "Pricing"}]}, "c3"),
                _text_response("BrandForge is free."),
            ]
            result = _run(chat.get_chat_response("t1", "how much is brandforge?", "th1", user_id="u1"))
            final_messages = llm.chat.completions.create.call_args_list[-1].kwargs["messages"]

    assert llm.chat.completions.create.call_count == 4  # loop ran well past 2 rounds
    assert search.call_count == 2                       # it searched again after a miss
    tool_msgs = [m for m in final_messages if m["role"] == "tool"]
    assert len(tool_msgs) == 3
    first = json.loads(tool_msgs[0]["content"])
    assert first["count"] == 0 and "different wording" in first["note"]
    second = json.loads(tool_msgs[1]["content"])
    assert second["results"][0]["content"] == fact
    assert second["results"][0]["source"] == "https://galaxiq.ai/pricing"
    assert second["results"][0]["images"] == ["https://galaxiq.ai/bf.png"]
    # Content reaches the model verbatim: no truncation, no \uXXXX mangling.
    assert fact in tool_msgs[1]["content"]
    # attach_resources still works now that the URL arrives via search results.
    assert result["resources"] == [{"url": "https://galaxiq.ai/pricing",
                                    "image": "https://galaxiq.ai/bf.png", "label": "Pricing"}]
    assert result["text"] == "BrandForge is free."


def test_json_fence_in_answer_survives_cleanup():
    """Legitimate fenced JSON in a model's answer (e.g. a ```json example block) must not
    be stripped by the legacy signal-cleanup regex, which should only target the four
    legacy signal block names."""
    from app.services.chat import chat
    reply_with_json = 'Here is an example:\n```json\n{"a": 1}\n```\nLet me know if that helps.'
    with _patch_chat():
        with patch("app.services.chat.chat.openai_client") as llm:
            llm.chat.completions.create.return_value = _text_response(reply_with_json)
            result = _run(chat.get_chat_response("t1", "hi", "th1", user_id="u1"))
    assert '```json' in result["text"]
    assert '{"a": 1}' in result["text"]
