fix(skills): filter providers before limiting search results
This commit is contained in:
committed by
kshitij
parent
a55c972e09
commit
e75b5a2d38
@@ -0,0 +1 @@
|
||||
dborodchuk
|
||||
@@ -0,0 +1,80 @@
|
||||
"""Provider filters must run before a source's result limit can discard matches."""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
from hermes_cli import skills_hub as cli_hub
|
||||
from hermes_cli.subcommands.skills import build_skills_parser
|
||||
from tools.skills_hub import _index_cache_dir
|
||||
from tools.skills_hub_github import GitHubAuth, GitHubSource
|
||||
from tools.skills_hub_models import SkillMeta, _cache_metas
|
||||
from tools.skills_hub_official import HermesIndexSource
|
||||
from tools.skills_hub_search import parallel_search_sources
|
||||
|
||||
|
||||
@pytest.fixture(params=["index", "github"])
|
||||
def catalog(request, monkeypatch):
|
||||
def entry(repo, provider, name):
|
||||
return SkillMeta(
|
||||
name=name, description="GPU utilities", source="github",
|
||||
identifier=f"{repo}/skills/{name}", trust_level="trusted",
|
||||
extra={"provider": provider},
|
||||
)
|
||||
|
||||
# All match the same query, but the requested provider sits beyond the
|
||||
# default per-source search window. Neither source's search is mocked.
|
||||
others = [entry("openai/skills", "OpenAI", f"gpu-other-{i}") for i in range(60)]
|
||||
wanted = [entry("NVIDIA/skills", "NVIDIA", f"gpu-target-{i}") for i in range(3)]
|
||||
auth = GitHubAuth()
|
||||
if request.param == "index":
|
||||
path = _index_cache_dir() / "hermes-index.json"
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(json.dumps({"skills": [vars(m) for m in others + wanted]}), encoding="utf-8")
|
||||
source = HermesIndexSource(auth)
|
||||
else:
|
||||
source = GitHubSource(auth)
|
||||
for tap in source.taps:
|
||||
key = f"{tap['repo']}_{tap['path']}_{tap.get('bucket') or ''}"
|
||||
key = key.replace("/", "_").replace(" ", "_")
|
||||
entries = []
|
||||
if tap["repo"] == "openai/skills":
|
||||
entries = others
|
||||
elif tap["repo"] == "NVIDIA/skills":
|
||||
entries = wanted
|
||||
_cache_metas(key, entries)
|
||||
monkeypatch.setattr(cli_hub, "_sources", lambda: [source])
|
||||
return source, others, wanted
|
||||
|
||||
|
||||
def test_cli_provider_search_finds_matches_beyond_global_window(catalog, capsys):
|
||||
source, others, wanted = catalog
|
||||
parser = argparse.ArgumentParser(prog="hermes")
|
||||
build_skills_parser(parser.add_subparsers(dest="command"), cmd_skills=cli_hub.skills_command)
|
||||
|
||||
for query in ("gpu", ""):
|
||||
args = parser.parse_args([
|
||||
"skills", "search", query, "--source", "nvidia", "--limit", "2", "--json",
|
||||
])
|
||||
args.func(args)
|
||||
result = json.loads(capsys.readouterr().out)
|
||||
assert [row["identifier"] for row in result] == [m.identifier for m in wanted[:2]]
|
||||
|
||||
# Filtering one call must not mutate the cached catalog for later searches.
|
||||
assert [m.identifier for m in source.search("gpu", limit=2)] == [m.identifier for m in others[:2]]
|
||||
|
||||
args = parser.parse_args(["skills", "search", "gpu", "--source", "anthropic", "--json"])
|
||||
args.func(args)
|
||||
assert json.loads(capsys.readouterr().out) == []
|
||||
|
||||
|
||||
def test_parallel_provider_search_applies_per_source_limit_after_filter(catalog):
|
||||
source, _, wanted = catalog
|
||||
results, counts, timed_out = parallel_search_sources(
|
||||
[source], query="gpu", source_filter=" NVIDIA ",
|
||||
per_source_limits={source.source_id(): 2},
|
||||
)
|
||||
assert [m.identifier for m in results] == [m.identifier for m in wanted[:2]]
|
||||
assert counts[source.source_id()] == len(results)
|
||||
assert not timed_out
|
||||
@@ -224,8 +224,8 @@ class GitHubSource(SkillSource):
|
||||
parts = identifier.split("/", 2)
|
||||
return "trusted" if len(parts) >= 2 and f"{parts[0]}/{parts[1]}" in TRUSTED_REPOS else "community"
|
||||
|
||||
def search(self, query: str, limit: int = 10) -> List[SkillMeta]:
|
||||
"""Substring-match all taps; dedupe by identifier preferring higher trust."""
|
||||
def search(self, query: str, limit: int = 10, *, provider_filter: str = "") -> List[SkillMeta]:
|
||||
"""Substring-match taps, filter by provider, then dedupe and limit results."""
|
||||
results: List[SkillMeta] = []
|
||||
query_lower = query.lower()
|
||||
for tap in self.taps:
|
||||
@@ -235,6 +235,8 @@ class GitHubSource(SkillSource):
|
||||
results.append(skill)
|
||||
except Exception as e:
|
||||
logger.debug("Failed to search %s: %s", tap['repo'], e)
|
||||
if provider_filter:
|
||||
results = _filter_results_by_provider(results, provider_filter)
|
||||
return _dedupe_by_trust(results)[:limit]
|
||||
|
||||
def fetch(self, identifier: str) -> Optional[SkillBundle]:
|
||||
|
||||
@@ -311,12 +311,19 @@ class HermesIndexSource(SkillSource):
|
||||
entry = next((s for s in self._skills() if s.get("identifier") == identifier), None)
|
||||
return entry.get("trust_level", "community") if entry else "community"
|
||||
|
||||
def search(self, query: str, limit: int = 10) -> List[SkillMeta]:
|
||||
def search(self, query: str, limit: int = 10, *, provider_filter: str = "") -> List[SkillMeta]:
|
||||
"""Search the cached index (zero API calls). Matches name, description, tags, identifier and
|
||||
``extra.provider`` (so ``nvidia`` finds ``NVIDIA/skills/...`` entries stored as source
|
||||
"github"). Ranked exact name > name prefix > provider > whole-word > name substring > other,
|
||||
index order as tiebreaker — a raw break-at-limit slice buried the most relevant skills."""
|
||||
index order as tiebreaker — a raw break-at-limit slice buried the most relevant skills.
|
||||
Provider filters narrow the catalog before ranking and limiting."""
|
||||
skills = self._skills()
|
||||
if provider_filter:
|
||||
want = provider_filter.strip().lower()
|
||||
skills = [
|
||||
s for s in skills
|
||||
if str((s.get("extra") or {}).get("provider", "")).lower() == want
|
||||
]
|
||||
if not skills:
|
||||
return []
|
||||
if not query.strip():
|
||||
|
||||
@@ -102,9 +102,15 @@ def create_source_router(auth: Optional[GitHubAuth] = None) -> List[SkillSource]
|
||||
]
|
||||
|
||||
|
||||
def _search_one_source(src: SkillSource, query: str, limit: int) -> Tuple[str, List[SkillMeta]]:
|
||||
def _search_one_source(
|
||||
src: SkillSource, query: str, limit: int, provider_filter: str = "",
|
||||
) -> Tuple[str, List[SkillMeta]]:
|
||||
"""Search a single source. Runs in a thread for parallelism."""
|
||||
try:
|
||||
# These sources mix providers in one catalog. Narrow before their top-N
|
||||
# cut so another provider cannot crowd every requested match out.
|
||||
if provider_filter and isinstance(src, (HermesIndexSource, GitHubSource)):
|
||||
return src.source_id(), src.search(query, limit=limit, provider_filter=provider_filter)
|
||||
return src.source_id(), src.search(query, limit=limit)
|
||||
except Exception as e:
|
||||
logger.debug("Search failed for %s: %s", src.source_id(), e)
|
||||
@@ -116,8 +122,9 @@ def _select_active_sources(sources: List[SkillSource], source_filter: str) -> Li
|
||||
|
||||
A provider filter (nvidia/openai/...) is not a source id — the data lives
|
||||
in the index/github source under ``extra.provider`` — so it selects like
|
||||
"all"; the narrowing happens later on the merged results. "official" is
|
||||
always included alongside an explicit source filter.
|
||||
"all". Mixed-provider sources filter before limiting; the merged results
|
||||
are filtered again. "official" is always queried alongside an explicit
|
||||
source filter.
|
||||
"""
|
||||
effective = "all" if source_filter.strip().lower() in _PROVIDER_FILTER_VALUES else source_filter
|
||||
index_available = effective == "all" and any(
|
||||
@@ -147,6 +154,9 @@ def parallel_search_sources(
|
||||
|
||||
per_source_limits = per_source_limits or {}
|
||||
active = _select_active_sources(sources, source_filter)
|
||||
provider_filter = source_filter.strip().lower()
|
||||
if provider_filter not in _PROVIDER_FILTER_VALUES:
|
||||
provider_filter = ""
|
||||
all_results: List[SkillMeta] = []
|
||||
source_counts: Dict[str, int] = {}
|
||||
timed_out_ids: List[str] = []
|
||||
@@ -159,7 +169,9 @@ def parallel_search_sources(
|
||||
from tools.daemon_pool import DaemonThreadPoolExecutor
|
||||
pool = DaemonThreadPoolExecutor(max_workers=min(len(active), 8))
|
||||
futures = {
|
||||
pool.submit(_search_one_source, src, query, per_source_limits.get(src.source_id(), 50)): src.source_id()
|
||||
pool.submit(
|
||||
_search_one_source, src, query, per_source_limits.get(src.source_id(), 50), provider_filter,
|
||||
): src.source_id()
|
||||
for src in active
|
||||
}
|
||||
try:
|
||||
|
||||
Reference in New Issue
Block a user