"""Multi-query ``tool_search``, batched ``tool_describe``, and stemming.

Covers the upgrade that replaced the single ``query`` string with
``queries: [str, ...]`` (grouped, split-shape response), the single
``name`` with ``names: [str, ...]`` (map response with ``not_found``),
and added Snowball stemming to the shared tokenizer.
"""

import json
from concurrent.futures import ThreadPoolExecutor

import pytest


def _td(name, desc, props=None, required=None):
    return {
        "type": "function",
        "function": {
            "name": name,
            "description": desc,
            "parameters": {
                "type": "object",
                "properties": props or {},
                "required": required or [],
            },
        },
    }


def _register(name, toolset, desc="Deferred capability.", props=None, required=None):
    from tools.registry import registry

    registry.register(
        name=name,
        handler=lambda args, **kw: json.dumps({"ok": True}),
        schema=_td(name, desc, props, required),
        toolset=toolset,
    )
    return _td(name, desc, props, required)


@pytest.fixture
def issue_defs():
    """A small deferred catalog registered under an MCP toolset."""
    return [
        _register("mq_linear_create_issue", "mcp-mq-linear",
                  "Create a new issue in a team.",
                  {"title": {"type": "string"}, "team": {"type": "string"}},
                  ["title", "team"]),
        _register("mq_linear_list_issues", "mcp-mq-linear",
                  "List issues in the workspace.",
                  {"query": {"type": "string"}}),
        _register("mq_slack_post_message", "mcp-mq-slack",
                  "Post a message to a channel.",
                  {"channel": {"type": "string"}, "text": {"type": "string"}},
                  ["channel", "text"]),
    ]


# ---------------------------------------------------------------------------
# Stemming
# ---------------------------------------------------------------------------


class TestStemming:
    def test_tokenize_stems_index_and_query_identically(self):
        from tools.tool_search_catalog import _tokenize
        # Same stem on both sides is the whole contract.
        assert _tokenize("issues") == _tokenize("issue")
        assert _tokenize("creating messages") == _tokenize("create message")

    def test_plural_query_finds_singular_tool_name(self, issue_defs):
        """The measured miss on the old tokenizer: 'issues' skipped create_issue."""
        from tools.tool_search import build_catalog, search_catalog

        catalog = build_catalog(issue_defs)
        names = [h.name for h in search_catalog(catalog, "issues", limit=5)]
        assert "mq_linear_create_issue" in names
        assert "mq_linear_list_issues" in names

    def test_rarest_query_token_gates_admission(self, issue_defs):
        """A document that lacks the query's rarest token is not a result, however many
        common tokens it shares. In this catalog 'issue' and 'linear' are each in two tools
        and 'slack' in one, so 'slack' gates: the two linear tools share two of the three
        query tokens and still do not come back."""
        from tools.tool_search import build_catalog, search_catalog

        catalog = build_catalog(issue_defs)
        names = [h.name for h in search_catalog(catalog, "linear issue slack", limit=5)]
        assert names == ["mq_slack_post_message"]

    def test_token_no_document_carries_admits_nothing(self, issue_defs):
        """'send gmail email' against a catalog with no gmail tool returns nothing rather
        than five tools that merely share 'message' or 'email'."""
        from tools.tool_search import build_catalog, search_catalog

        catalog = build_catalog(issue_defs)
        assert search_catalog(catalog, "post gmail message", limit=5) == []


    def test_parallel_tokenize_search_and_dispatch_are_deterministic(self, issue_defs):
        from tools.tool_search import (
            ToolSearchConfig,
            build_catalog,
            dispatch_tool_search,
            search_catalog,
        )
        from tools.tool_search_catalog import _stem, _tokenize

        corpus = (
            "issues",
            "issue",
            "creating",
            "create",
            "meetings",
            "meeting",
            "post slack message",
            "messages posted",
        )
        catalog = build_catalog(issue_defs)
        expected = {
            text: (
                _tokenize(text),
                [entry.name for entry in search_catalog(catalog, text, limit=3)],
            )
            for text in corpus
        }

        def tokenize_and_search(index):
            text = corpus[index % len(corpus)]
            return text, _tokenize(text), [
                entry.name for entry in search_catalog(catalog, text, limit=3)
            ]

        _stem.cache_clear()
        misses_before = _stem.cache_info().misses
        with ThreadPoolExecutor(max_workers=8) as pool:
            threaded = list(pool.map(tokenize_and_search, range(512)))

        assert _stem.cache_info().misses > misses_before
        for text, tokens, names in threaded:
            assert (tokens, names) == expected[text]

        args = {"queries": ["issues", "post slack message", "meetings"]}
        config = ToolSearchConfig.from_raw({})
        expected_json = dispatch_tool_search(
            args,
            current_tool_defs=issue_defs,
            config=config,
        )

        def dispatch(_index):
            return dispatch_tool_search(
                args,
                current_tool_defs=issue_defs,
                config=config,
            )

        with ThreadPoolExecutor(max_workers=8) as pool:
            dispatched = list(pool.map(dispatch, range(64)))

        assert dispatched == [expected_json] * 64

    def test_stemmer_is_safe_under_concurrent_cache_misses(self):
        """Hammer the raw stemmer from 8 threads with cache-missing input.

        ``_stem``'s lru_cache means a small corpus warms after a handful of
        misses and later iterations never reach the stemmer, so a shared
        (non-thread-local) stemmer instance can survive a threaded test over
        repeated tokens. This test bypasses the cache: every call stems a
        unique token via ``_stem.__wrapped__``, so thousands of stems execute
        concurrently on the underlying per-thread instances. A shared
        stemmer's mutable parse state produces wrong stems or raises here.
        """
        from tools.tool_search_catalog import _stem

        words = ["issues", "creating", "meetings", "categories", "searching"]

        def serial_baseline(salt):
            return [
                _stem.__wrapped__(f"{word}x{salt}n{i}")
                for i, word in enumerate(words)
            ]

        expected = {salt: serial_baseline(salt) for salt in range(400)}

        def worker(salt):
            return salt, [
                _stem.__wrapped__(f"{word}x{salt}n{i}")
                for i, word in enumerate(words)
            ]

        with ThreadPoolExecutor(max_workers=8) as pool:
            for salt, stems in pool.map(worker, range(400)):
                assert stems == expected[salt]


# ---------------------------------------------------------------------------
# Exact-name ranking and shared corpus statistics
# ---------------------------------------------------------------------------


class TestCatalogRanking:
    def test_exact_name_beats_shorter_siblings(self):
        from tools.tool_search import build_catalog, search_catalog

        exact = _td(
            "github_create_issue",
            "Create a new issue with a title, body, assignees, labels, "
            "milestone, project metadata, and linked context for a repository.",
        )
        catalog = build_catalog([
            exact,
            _td("github_create_issue_comment", "Comment."),
            _td("github_create_issue_label", "Label."),
        ])

        assert search_catalog(catalog, "github_create_issue", limit=1) == [catalog[0]]

    def test_exact_short_name_beats_prefixed_names(self):
        from tools.tool_search import build_catalog, search_catalog

        catalog = build_catalog([
            _td("list", "List one item."),
            _td("list_x", "List x."),
            _td("list_all_the_open_items", "List every open item."),
        ])

        assert search_catalog(catalog, "list", limit=1) == [catalog[0]]

    def test_precomputed_corpus_stats_preserve_results(self, issue_defs):
        from tools.tool_search import build_catalog, search_catalog
        from tools.tool_search_catalog import _corpus_stats

        catalog = build_catalog(issue_defs)
        expected = search_catalog(catalog, "create issues", limit=3)
        actual = search_catalog(
            catalog,
            "create issues",
            limit=3,
            corpus_stats=_corpus_stats(catalog),
        )

        assert actual == expected


# ---------------------------------------------------------------------------
# Multi-query dispatch_tool_search
# ---------------------------------------------------------------------------


class TestMultiQuerySearch:
    def test_grouped_names_plus_shared_tool_map(self, issue_defs):
        from tools.tool_search import ToolSearchConfig, dispatch_tool_search

        result = json.loads(dispatch_tool_search(
            {"queries": ["create linear issue", "post slack message"]},
            current_tool_defs=issue_defs,
            config=ToolSearchConfig.from_raw({}),
        ))

        assert result["queries"] == ["create linear issue", "post slack message"]
        assert result["total_available"] == 3
        # Groups carry NAMES only, in query order.
        assert [g["query"] for g in result["results"]] == result["queries"]
        for group in result["results"]:
            for name in group["matches"]:
                assert isinstance(name, str)
        assert "mq_linear_create_issue" in result["results"][0]["matches"]
        assert "mq_slack_post_message" in result["results"][1]["matches"]
        # The shared map holds each matched tool exactly once, and nothing else.
        matched = {n for g in result["results"] for n in g["matches"]}
        assert set(result["tools"]) == matched
        record = result["tools"]["mq_linear_create_issue"]
        assert record["source"] == "mcp"
        assert record["source_name"] == "mcp-mq-linear"
        assert record["description"].startswith("Create a new issue")
        assert record["required"] == ["title", "team"]
        # All queries matched → no fallback block.
        assert "available_sources" not in result
        assert "hint" not in result

    def test_limit_applies_per_query(self, issue_defs):
        from tools.tool_search import ToolSearchConfig, dispatch_tool_search

        result = json.loads(dispatch_tool_search(
            {"queries": ["issues", "message"], "limit": 1},
            current_tool_defs=issue_defs,
            config=ToolSearchConfig.from_raw({}),
        ))
        for group in result["results"]:
            assert len(group["matches"]) <= 1

    def test_required_names_are_bounded(self):
        from tools.tool_search import ToolSearchConfig, dispatch_tool_search

        required = [f"field_{index}_" + ("x" * 5000) for index in range(200)]
        name = "mq_bounded_required_fields"
        tool_def = _register(name, "mcp-mq-bounds", required=required)
        result = json.loads(dispatch_tool_search(
            {"queries": [name]},
            current_tool_defs=[tool_def],
            config=ToolSearchConfig.from_raw({}),
        ))
        record = result["tools"][name]

        assert len(record["required"]) <= 32
        assert all(len(item) <= 64 for item in record["required"])

    @pytest.mark.parametrize("schema", [
        {"function": "not an object"},
        {"function": {"parameters": ["not", "an", "object"]}},
    ])
    def test_shared_record_handles_non_object_schema_fields(self, schema):
        from tools.tool_search import CatalogEntry, _shared_tool_record

        entry = CatalogEntry(
            name="mq_malformed_schema",
            description="Malformed schema fixture.",
            schema=schema,
            source="mcp",
            source_name="mcp-mq-malformed",
        )

        assert _shared_tool_record(entry)["required"] == []

    def test_partial_miss_adds_fallback_to_empty_group(self, issue_defs):
        from tools.tool_search import ToolSearchConfig, dispatch_tool_search

        result = json.loads(dispatch_tool_search(
            {"queries": ["issues", "zzzz nonsense qqqq"]},
            current_tool_defs=issue_defs,
            config=ToolSearchConfig.from_raw({}),
        ))
        assert result["results"][1]["matches"] == []
        assert "available_sources" not in result["results"][0]
        assert "hint" not in result["results"][0]
        missed = result["results"][1]
        assert missed["hint"]
        source_names = {s["name"] for s in missed["available_sources"]}
        assert {"mq-linear", "mq-slack"} <= source_names
        assert "available_sources" not in result
        assert "hint" not in result

    def test_bare_string_query_coerced_to_single_query(self, issue_defs):
        from tools.tool_search import ToolSearchConfig, dispatch_tool_search

        result = json.loads(dispatch_tool_search(
            {"queries": "post slack message"},
            current_tool_defs=issue_defs,
            config=ToolSearchConfig.from_raw({}),
        ))
        assert result["queries"] == ["post slack message"]

    def test_max_query_cap_respected(self, issue_defs, monkeypatch):
        import tools.tool_search as tool_search

        monkeypatch.setattr(tool_search, "_MAX_QUERIES_PER_CALL", 2)
        cfg = tool_search.ToolSearchConfig.from_raw({})
        ok = json.loads(tool_search.dispatch_tool_search(
            {"queries": ["a b", "c d"]}, current_tool_defs=issue_defs, config=cfg))
        assert "error" not in ok
        over = json.loads(tool_search.dispatch_tool_search(
            {"queries": ["a", "b", "c"]}, current_tool_defs=issue_defs, config=cfg))
        assert "error" in over


# ---------------------------------------------------------------------------
# Batched dispatch_tool_describe
# ---------------------------------------------------------------------------


class TestBatchedDescribe:
    def test_map_response_with_not_found(self, issue_defs):
        from tools.tool_search import ToolSearchConfig, dispatch_tool_describe

        # Deferrable in the global registry, but NOT in this session's defs —
        # the stale/out-of-scope case that lands in not_found.
        _register("mq_out_of_scope_op", "mcp-mq-elsewhere")

        result = json.loads(dispatch_tool_describe(
            {"names": ["mq_linear_create_issue", "mq_slack_post_message",
                       "mq_out_of_scope_op", "mcp__bogus__missing"]},
            current_tool_defs=issue_defs,
            config=ToolSearchConfig.from_raw({}),
        ))
        assert set(result["tools"]) == {"mq_linear_create_issue",
                                        "mq_slack_post_message"}
        schema = result["tools"]["mq_linear_create_issue"]
        assert schema["description"] == "Create a new issue in a team."
        assert schema["parameters"]["required"] == ["title", "team"]
        # Deferrable-but-absent and unknown names collect in not_found; found
        # ones still resolve.
        assert result["not_found"] == ["mq_out_of_scope_op", "mcp__bogus__missing"]
        assert result["hint"]
        assert "errors" not in result

    def test_real_schemas_and_unknown_name_are_classified_independently(self):
        from tools.tool_search import ToolSearchConfig, dispatch_tool_describe

        tool_defs = [
            _register("mcp__linear__get_issue", "mcp-linear"),
            _register("mcp__granola__list_meeting_folders", "mcp-granola"),
        ]
        result = json.loads(dispatch_tool_describe(
            {
                "names": [
                    "mcp__linear__get_issue",
                    "mcp__granola__list_meeting_folders",
                    "mcp__linear__does_not_exist_zzz",
                ]
            },
            current_tool_defs=tool_defs,
            config=ToolSearchConfig.from_raw({}),
        ))

        assert set(result["tools"]) == {
            "mcp__linear__get_issue",
            "mcp__granola__list_meeting_folders",
        }
        assert result["not_found"] == ["mcp__linear__does_not_exist_zzz"]
        assert "errors" not in result

    def test_unregistered_core_name_is_not_found(self, issue_defs, monkeypatch):
        from tools import registry as registry_module
        from tools.tool_search import ToolSearchConfig, dispatch_tool_describe

        # The intent: a name that is NOT registered lands in not_found, even
        # when it looks like a core tool. Whether "terminal" is registered in
        # this process depends on which test files imported model_tools
        # earlier, so force the unregistered condition instead of relying on
        # collection order.
        real_get_entry = registry_module.registry.get_entry
        monkeypatch.setattr(
            registry_module.registry,
            "get_entry",
            lambda name: None if name == "terminal" else real_get_entry(name),
        )

        result = json.loads(dispatch_tool_describe(
            {"names": ["terminal", "mq_linear_create_issue"]},
            current_tool_defs=issue_defs,
            config=ToolSearchConfig.from_raw({}),
        ))
        assert "mq_linear_create_issue" in result["tools"]
        assert "terminal" in result["not_found"]
        assert "errors" not in result

    def test_registered_direct_surface_name_keeps_exact_error(self):
        from tools.tool_search import ToolSearchConfig, dispatch_tool_describe

        name = "mq_desktop_direct_action"
        tool_def = _register(name, "desktop_ui")
        result = json.loads(dispatch_tool_describe(
            {"names": [name]},
            current_tool_defs=[tool_def],
            config=ToolSearchConfig.from_raw({}),
        ))

        assert result["errors"][name]
        assert name not in result.get("not_found", [])

    def test_registry_lookup_failure_is_not_found(self, monkeypatch):
        from tools.registry import registry
        from tools.tool_search import ToolSearchConfig, dispatch_tool_describe

        def fail_lookup(name):
            raise RuntimeError("registry unavailable")

        monkeypatch.setattr(registry, "get_entry", fail_lookup)
        result = json.loads(dispatch_tool_describe(
            {"names": ["mq_unknown_during_lookup"]},
            current_tool_defs=[],
            config=ToolSearchConfig.from_raw({}),
        ))

        assert result["not_found"] == ["mq_unknown_during_lookup"]
        assert "errors" not in result

    def test_duplicates_deduped_silently(self, issue_defs):
        from tools.tool_search import ToolSearchConfig, dispatch_tool_describe

        result = json.loads(dispatch_tool_describe(
            {"names": ["mq_linear_create_issue", "mq_linear_create_issue"]},
            current_tool_defs=issue_defs,
            config=ToolSearchConfig.from_raw({}),
        ))
        assert list(result["tools"]) == ["mq_linear_create_issue"]
        assert "not_found" not in result

    def test_empty_and_overcap_names_error(self, issue_defs, monkeypatch):
        import tools.tool_search as tool_search

        monkeypatch.setattr(tool_search, "_MAX_DESCRIBE_NAMES_PER_CALL", 2)
        cfg = tool_search.ToolSearchConfig.from_raw({})
        assert "error" in json.loads(tool_search.dispatch_tool_describe(
            {}, current_tool_defs=issue_defs, config=cfg))
        assert "error" in json.loads(tool_search.dispatch_tool_describe(
            {"names": []}, current_tool_defs=issue_defs, config=cfg))
        over = ["n%d" % i for i in range(3)]
        parsed = json.loads(tool_search.dispatch_tool_describe(
            {"names": over}, current_tool_defs=issue_defs, config=cfg))
        assert "error" in parsed

    def test_bare_string_name_coerced(self, issue_defs):
        from tools.tool_search import ToolSearchConfig, dispatch_tool_describe

        result = json.loads(dispatch_tool_describe(
            {"names": "mq_linear_create_issue"},
            current_tool_defs=issue_defs,
            config=ToolSearchConfig.from_raw({}),
        ))
        assert "mq_linear_create_issue" in result["tools"]


# ---------------------------------------------------------------------------
# Config + bridge schema
# ---------------------------------------------------------------------------


class TestConfigAndSchema:
    def test_limit_default_within_cap(self):
        from hermes_cli.config_defaults import DEFAULT_CONFIG
        from tools.tool_search import ToolSearchConfig

        cfg = ToolSearchConfig.from_raw(DEFAULT_CONFIG["tools"]["tool_search"])
        assert 1 <= cfg.search_default_limit <= cfg.max_search_limit <= 50

    def test_bridge_schema_declares_array_inputs(self):
        from tools.tool_search import bridge_tool_schemas

        schemas = {s["function"]["name"]: s["function"] for s in bridge_tool_schemas(3)}
        search_params = schemas["tool_search"]["parameters"]
        assert search_params["required"] == ["queries"]
        assert search_params["properties"]["queries"]["type"] == "array"
        describe_params = schemas["tool_describe"]["parameters"]
        assert describe_params["required"] == ["names"]
        assert describe_params["properties"]["names"]["type"] == "array"
