fix(skills): filter providers before limiting search results

This commit is contained in:
Danylo Borodchuk
2026-09-09 17:59:38 -07:00
committed by kshitij
parent a55c972e09
commit e75b5a2d38
5 changed files with 110 additions and 8 deletions
+1
View File
@@ -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
+4 -2
View File
@@ -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]:
+9 -2
View File
@@ -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():
+16 -4
View File
@@ -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: