diff --git a/ir/__init__.py b/ir/__init__.py index ce222a2..a437da8 100644 --- a/ir/__init__.py +++ b/ir/__init__.py @@ -75,5 +75,7 @@ def search(corpus, query, **kwargs): # The evaluation harness is reachable as ``ir.eval`` (its ``ef`` imports are # lazy, so this does not weigh down ``import ir``). Kept out of ``__all__`` so a -# star-import does not shadow the ``eval`` builtin. +# star-import does not shadow the ``eval`` builtin. ``ir.eval_gen`` is the +# build-time case generator (its ``oa`` import is lazy too). from . import eval # noqa: E402,F401 (submodule attribute: ir.eval) +from . import eval_gen # noqa: E402,F401 (submodule attribute: ir.eval_gen) diff --git a/ir/cli.py b/ir/cli.py index 1ad18eb..2cd5166 100644 --- a/ir/cli.py +++ b/ir/cli.py @@ -9,6 +9,7 @@ ir info packages # config + stats for a corpus ir register notes files --root ~/notes --pattern '.*\\.md$' ir rm notes # unregister (keeps built data) + ir eval-gen skills skills_eval.jsonl --k 5 # generate cases (needs oa/LLM) ir eval skills skills_eval.jsonl --mode hybrid # score retrieval on a case file """ @@ -113,4 +114,34 @@ def eval(name, cases, *, mode="hybrid", k=10): return out -COMMANDS = [ls, register, build, search, info, rm, eval] +def eval_gen(name, out, *, k=5, abstention_frac=0.15, max_artifacts=None): + """Generate an eval-case file for a corpus by back-translation (needs oa/LLM). + + Writes a DiscoveryCase JSONL set (gold cases + an abstention slice) for the + registered corpus *name* to *out*, stamping a corpus-signature into the + header so the frozen file can be checked against the live corpus later. This + command calls an LLM via oa; scoring it afterwards (`ir eval`) is offline. + """ + from .eval import save_cases + from .eval_gen import build_eval_set, corpus_signature + + source = registry.source_for(name) + kwargs = {} + if max_artifacts is not None: + kwargs["max_artifacts"] = int(max_artifacts) + cases = build_eval_set( + source, k=k, abstention_frac=abstention_frac, corpus_name=name, **kwargs + ) + save_cases( + cases, + out, + meta={"corpus": name, "corpus_signature": corpus_signature(source), "k": k}, + ) + n_gold = sum(not c.gold_is_none for c in cases) + return ( + f"wrote {len(cases)} cases ({n_gold} gold, {len(cases) - n_gold} abstention) " + f"to {out!r}" + ) + + +COMMANDS = [ls, register, build, search, info, rm, eval, eval_gen] diff --git a/ir/eval_gen.py b/ir/eval_gen.py new file mode 100644 index 0000000..3bc6db5 --- /dev/null +++ b/ir/eval_gen.py @@ -0,0 +1,450 @@ +"""LLM-backed generation of evaluation cases for :mod:`ir.eval`. + +This is the **build-time** companion to the offline scoring harness: it turns a +corpus into a set of :class:`~ir.eval.DiscoveryCase`\\ s by *back-translation* — +given a capability's description, ask an LLM for the user intents that should +route to it. The artifact's id is the free ground-truth label. + +Two ideas make the generated set honest: + +- **Name masking.** The artifact's *name* is stripped from the description + *before* it is shown to the generator (and any output that still leaks the + name is dropped). Otherwise a lexical retriever (the BM25 leg of hybrid) would + trivially match query→gold on surface name overlap and inflate scores. Many + real descriptions contain their own name, so masking the **input** — not just + filtering the output — is what matters. +- **An abstention slice.** A fraction of cases are "no artifact applies" intents + (empty ``gold``), so the eval can measure correct refusal, not just hits. + +The LLM is **injected** (`query_generator` / `abstention_generator` callables), +so the generation *logic* — masking, gold assignment, the leakage guard, the +abstention fraction — is fully testable with a deterministic stub and no network. +The default generators are built lazily on :mod:`oa` (`oa.prompt_function`), so +``import ir.eval_gen`` stays cheap and offline; ``oa`` is only imported when you +actually generate with the real LLM. + +The output is plain :class:`~ir.eval.DiscoveryCase` data — freeze it with +:func:`ir.eval.save_cases` (stamping :func:`corpus_signature` into the +``__meta__`` header) and score it with :mod:`ir.eval`. Generation needs a model; +scoring never does. + +Quick start:: + + import ir + from ir import eval_gen as eg + + source = ir.CorpusSource.from_skills() + cases = eg.build_eval_set(source, k=5, corpus_name="skills") # uses oa + from ir.eval import save_cases + save_cases(cases, "skills_eval.jsonl", + meta={"corpus": "skills", "corpus_signature": eg.corpus_signature(source)}) +""" + +from __future__ import annotations + +import hashlib +import math +import re +import warnings +from collections.abc import Mapping, Sequence +from typing import Any, Callable + +from .eval import DiscoveryCase + +#: Queries generated per artifact by default. +DFLT_QUERIES_PER_ARTIFACT = 5 + +#: Target share of the case set that is abstention ("no artifact applies"). +DFLT_ABSTENTION_FRAC = 0.15 + +#: Minimum description length (chars) for an artifact to be back-translated. +DFLT_MIN_DESCRIPTION_CHARS = 20 + +#: What a masked name is replaced with in a description / query. +NAME_PLACEHOLDER = "this capability" + +#: Default theme used when generating abstention ("out of scope") intents. +DFLT_ABSTENTION_THEME = "software developer tools" + +#: A query generator: ``(description, *, n) -> list[str]`` (n candidate intents). +QueryGenerator = Callable[..., Sequence[str]] + +#: An abstention generator: ``(*, n, theme) -> list[str]`` (n out-of-scope intents). +AbstentionGenerator = Callable[..., Sequence[str]] + +BACKTRANSLATION_PROMPT = """\ +You are generating evaluation data for a tool-retrieval system. + +Below is a description of a capability (its name has been hidden on purpose): + +{description} + +Write {n} natural, varied user requests that this capability should handle. +Rules: +- Do NOT mention any tool, function, skill, or package name. +- Vary phrasing, specificity, and the implied (not explicit) parameters. +- One request per line. No numbering, no quotes, no extra commentary. +""" + +ABSTENTION_PROMPT = """\ +You are generating "no applicable tool" cases for a tool-retrieval evaluation. + +The tool catalog is about: {theme}. + +Write {n} natural, plausible user requests that such a catalog should NOT be able +to satisfy because they fall outside its scope. +One request per line. No numbering, no quotes, no extra commentary. +""" + + +# =========================================================================== # +# Name masking (label-leakage prevention) +# =========================================================================== # + + +#: Fallback placeholder used when the primary one would itself match the name. +_ALT_PLACEHOLDER = "the tool" + + +def _ordered_tokens(name: str) -> list[str]: + """The name's word tokens (split on whitespace/hyphen/underscore), lowercased. + + Tokens shorter than two characters are dropped so a degenerate name (e.g. + ``"a-"`` or ``"-x"``) cannot collapse to a one-letter pattern that masks + articles or stray letters throughout the text. + """ + return [t for t in re.split(r"[\s_-]+", name.lower()) if len(t) >= 2] + + +def _name_pattern(name: str) -> re.Pattern[str] | None: + """Whole-word regex matching ``name`` as a contiguous phrase. + + Tolerant of the separator *between* a multi-token name's tokens (any run of + whitespace / hyphen / underscore), plus the concatenated form — so + ``"ci-advisor"``, ``"ci advisor"``, ``"ci_advisor"`` and ``"ciadvisor"`` all + match. Returns ``None`` when nothing maskable survives (see + :func:`_ordered_tokens`). + """ + tokens = _ordered_tokens(name) + if not tokens: + return None + phrase = r"[\s_-]+".join(re.escape(t) for t in tokens) + concat = "".join(re.escape(t) for t in tokens) + alts = "|".join(sorted({phrase, concat}, key=len, reverse=True)) + return re.compile(rf"\b(?:{alts})\b", re.IGNORECASE) + + +def mask_name(text: str, name: str, *, placeholder: str = NAME_PLACEHOLDER) -> str: + """Replace occurrences of ``name`` in ``text`` with ``placeholder``. + + Matches the name as a contiguous phrase tolerant of the separator between its + tokens (``"ci-advisor"`` / ``"ci advisor"`` / ``"ci_advisor"`` all match), + case-insensitively and whole-word (a short token like ``"ci"`` does not blast + through ``"specific"``). If ``placeholder`` would itself match the name, a + neutral alternate is used so the masked text is not self-referential. Used to + scrub the artifact name out of a description *before* it is generated from. + """ + pattern = _name_pattern(name) + if pattern is None: + return text + if pattern.search(placeholder): + placeholder = _ALT_PLACEHOLDER + return pattern.sub(placeholder, text) + + +def _leaks_name(text: str, name: str) -> bool: + """Whether ``text`` reuses the gold ``name`` — matching how BM25 would. + + The lexical retriever is bag-of-words (order-insensitive), so two forms count + as a leak: (1) the name as a contiguous phrase (any separator), or (2) for a + multi-token name, *all* its distinctive tokens appearing anywhere (reordered + or separated). A single shared content word is **not** a leak — that is the + legitimate semantic overlap the eval is meant to measure. + """ + pattern = _name_pattern(name) + if pattern is None: + return False + if pattern.search(text): + return True + tokens = set(_ordered_tokens(name)) + if len(tokens) >= 2: + return tokens <= set(re.findall(r"\w+", text.lower())) + return False + + +# =========================================================================== # +# Default (oa-backed) generators — lazily built, only when actually used +# =========================================================================== # + + +#: A genuine leading list marker: a bullet or an ordinal *followed by whitespace* +#: (``- ``, ``* ``, ``• ``, ``1. ``, ``2) ``). Anchored so it never eats a real +#: leading token like ``3D``, ``-9 degrees``, ``.env`` or ``24/7``. +_LIST_MARKER = re.compile(r"^\s*(?:[-*•]|\d+[.)])\s+") + + +def _parse_lines(text: Any) -> list[str]: + """Parse an LLM list response into clean, non-empty lines. + + Strips only a genuine leading list marker (bullet or ordinal followed by + whitespace) and surrounding quotes — never an arbitrary leading character — + so a query like ``"3D modeling help"`` or ``"-9 degrees, what to wear?"`` + keeps its first token intact. + """ + lines = [] + for raw in str(text).splitlines(): + cleaned = _LIST_MARKER.sub("", raw.strip()).strip().strip("\"'") + if cleaned: + lines.append(cleaned) + return lines + + +def make_oa_query_generator( + *, prompt: str = BACKTRANSLATION_PROMPT, **prompt_function_kwargs: Any +) -> QueryGenerator: + """Build the default back-translation generator on :mod:`oa` (lazy import).""" + import oa + + fn = oa.prompt_function( + prompt, egress=_parse_lines, name="backtranslate", **prompt_function_kwargs + ) + + def generate(description: str, *, n: int) -> list[str]: + return list(fn(description=description, n=n))[:n] + + return generate + + +def make_oa_abstention_generator( + *, prompt: str = ABSTENTION_PROMPT, **prompt_function_kwargs: Any +) -> AbstentionGenerator: + """Build the default abstention generator on :mod:`oa` (lazy import).""" + import oa + + fn = oa.prompt_function( + prompt, egress=_parse_lines, name="abstention", **prompt_function_kwargs + ) + + def generate(*, n: int, theme: str) -> list[str]: + return list(fn(theme=theme, n=n))[:n] + + return generate + + +# =========================================================================== # +# Case generation +# =========================================================================== # + + +def _default_describe(raw: Any) -> str: + """Best-effort describable text from a raw artifact payload.""" + if isinstance(raw, Mapping): + for key in ("description", "text"): + val = raw.get(key) + if isinstance(val, str) and val.strip(): + return val + return "\n".join(str(v) for v in raw.values() if isinstance(v, str)) + return str(raw) + + +def _name_of(artifact_id: str, raw: Any) -> str: + """The artifact's display name — a non-blank string ``name`` field, else the id. + + A missing, blank, or non-string ``name`` is treated like absent (the id is + used), so masking and the leakage guard never operate on the literal + ``"None"`` or an empty string. + """ + if isinstance(raw, Mapping): + name = raw.get("name") + if isinstance(name, str) and name.strip(): + return name + return artifact_id + + +def generate_cases( + source: Any, + *, + k: int = DFLT_QUERIES_PER_ARTIFACT, + mask_names: bool = True, + query_generator: QueryGenerator | None = None, + describe: Callable[[Any], str] | None = None, + min_chars: int = DFLT_MIN_DESCRIPTION_CHARS, + max_artifacts: int | None = None, + corpus_name: str | None = None, +) -> list[DiscoveryCase]: + """Back-translate a corpus source into gold-bearing :class:`DiscoveryCase`\\ s. + + For each artifact in ``source.scope`` (id → raw), extract a description, + mask the artifact's name out of it, ask ``query_generator`` for ``k`` user + intents, and emit one case per surviving intent (gold = the artifact id). + Intents that still leak the name are dropped; artifacts whose description is + shorter than ``min_chars`` are skipped (and the count is warned, never + silently dropped). + + Args: + source: a :class:`~ir.sources.CorpusSource` (anything with ``.items()``). + k: intents to request per artifact. + mask_names: scrub the artifact name from the description before + generating, and drop any generated intent that still contains it. + query_generator: ``(description, *, n) -> [intent, …]``. Defaults to the + :mod:`oa`-backed back-translator (built lazily; needs a model). + describe: ``raw -> description`` (default: the ``description`` / ``text`` + field, else the joined string fields). + min_chars: skip artifacts whose description is shorter than this. + max_artifacts: cap how many artifacts to process (for a quick/cheap run); + when set, artifacts are taken in sorted-id order so the subset is + deterministic even for filesystem-ordered (``dol``-backed) scopes. + corpus_name: stamped on each case's ``corpus`` field. + + Returns: + the generated gold cases. + + Raises: + ValueError: if ``k`` is less than 1. + """ + if k < 1: + raise ValueError(f"k must be >= 1, got {k!r}.") + gen = query_generator or make_oa_query_generator() + describe = describe or _default_describe + cases: list[DiscoveryCase] = [] + skipped = 0 + + items = list(source.items()) + if max_artifacts is not None: + # Deterministic subset: sort by id so a capped run is reproducible across + # machines even when the scope iterates in filesystem order. + items = sorted(items, key=lambda kv: kv[0])[:max_artifacts] + + for artifact_id, raw in items: + description = describe(raw) + if not description or len(description.strip()) < min_chars: + skipped += 1 + continue + name = _name_of(artifact_id, raw) + prompt_text = mask_name(description, name) if mask_names else description + try: + intents = gen(prompt_text, n=k) + except Exception as exc: # a single artifact's generation failing is non-fatal + warnings.warn(f"query generation failed for {artifact_id!r}: {exc}") + skipped += 1 + continue + for intent in intents: + intent = (intent or "").strip() + if not intent: + continue + if mask_names and _leaks_name(intent, name): + continue # output guard: never let the gold name leak into a query + cases.append( + DiscoveryCase( + query=intent, + gold=(artifact_id,), + corpus=corpus_name, + source_id=artifact_id, + metadata={ + "generator": "backtranslation", + "masked": bool(mask_names), + }, + ) + ) + + if skipped: + warnings.warn( + f"generate_cases skipped {skipped} artifact(s) " + f"(description shorter than {min_chars} chars, or a generation error)." + ) + return cases + + +def generate_abstention_cases( + n: int, + *, + generator: AbstentionGenerator | None = None, + theme: str = DFLT_ABSTENTION_THEME, + corpus_name: str | None = None, +) -> list[DiscoveryCase]: + """Generate ``n`` abstention cases — out-of-scope intents (empty ``gold``).""" + if n <= 0: + return [] + gen = generator or make_oa_abstention_generator() + intents = gen(n=n, theme=theme) + cases = [ + DiscoveryCase( + query=intent.strip(), + gold=(), + corpus=corpus_name, + metadata={"generator": "abstention"}, + ) + for intent in intents + if intent and intent.strip() + ] + return cases[:n] + + +def build_eval_set( + source: Any, + *, + k: int = DFLT_QUERIES_PER_ARTIFACT, + abstention_frac: float = DFLT_ABSTENTION_FRAC, + query_generator: QueryGenerator | None = None, + abstention_generator: AbstentionGenerator | None = None, + theme: str = DFLT_ABSTENTION_THEME, + corpus_name: str | None = None, + **gen_kwargs: Any, +) -> list[DiscoveryCase]: + """Generate a full eval set — gold cases plus an abstention slice. + + The abstention count is chosen so abstention cases make up (at least) + ``abstention_frac`` of the returned set: ``ceil(frac * G / (1 - frac))`` for + ``G`` gold cases. Extra ``gen_kwargs`` flow to :func:`generate_cases` + (``mask_names``, ``min_chars``, ``max_artifacts``, ``describe``). + + Raises: + ValueError: if ``abstention_frac`` is outside ``[0, 1)`` (``frac=0`` + means no abstention slice) or ``k`` is less than 1. + """ + if not 0.0 <= abstention_frac < 1.0: + raise ValueError(f"abstention_frac must be in [0, 1), got {abstention_frac!r}.") + gold_cases = generate_cases( + source, + k=k, + query_generator=query_generator, + corpus_name=corpus_name, + **gen_kwargs, + ) + n_abstain = 0 + if gold_cases and 0.0 < abstention_frac < 1.0: + n_abstain = math.ceil( + abstention_frac * len(gold_cases) / (1.0 - abstention_frac) + ) + abstain_cases = generate_abstention_cases( + n_abstain, + generator=abstention_generator, + theme=theme, + corpus_name=corpus_name, + ) + return gold_cases + abstain_cases + + +# =========================================================================== # +# Reproducibility anchor +# =========================================================================== # + + +def _artifact_ids(source_or_corpus: Any) -> list[str]: + """Artifact ids of a CorpusSource (its ``scope``) or a built Corpus.""" + scope = getattr(source_or_corpus, "scope", None) + if scope is not None: + return list(scope) + from .eval import corpus_artifact_ids + + return list(corpus_artifact_ids(source_or_corpus)) + + +def corpus_signature(source_or_corpus: Any) -> str: + """A short, order-independent hash of a corpus's artifact ids. + + Stamp this into a case file's ``__meta__`` header so a frozen eval set can be + checked against the (live, machine-specific) corpus it was generated from. + """ + blob = "\n".join(sorted(_artifact_ids(source_or_corpus))).encode("utf-8") + return hashlib.sha256(blob).hexdigest()[:16] diff --git a/tests/test_eval_gen.py b/tests/test_eval_gen.py new file mode 100644 index 0000000..0627ab4 --- /dev/null +++ b/tests/test_eval_gen.py @@ -0,0 +1,445 @@ +"""Case-generation tests — hermetic: deterministic stub generators, no LLM/network. + +The LLM is injected (`query_generator` / `abstention_generator`), so masking, the +leakage guard, gold assignment, the abstention fraction, and end-to-end +scorability are all tested without `oa` or a model. +""" + +import warnings + +import pytest + +import ir +from ir import eval as ev +from ir import eval_gen as eg +from ir.store import CorpusStore + + +def _echo_gen(description, *, n): + """Stub query generator: echo the (masked) description as n identical intents.""" + return [description] * n + + +def _skill_source(docs): + return ir.CorpusSource.from_mapping(docs, name="g", strategy=ir.Skill()) + + +# --------------------------------------------------------------------------- # +# Name masking / leakage detection +# --------------------------------------------------------------------------- # + + +def test_mask_name_hyphen_and_underscore_variants(): + assert ( + eg.mask_name("Use my-packages now", "my-packages") == "Use this capability now" + ) + assert eg.mask_name("run my_tool please", "my_tool") == "run this capability please" + # case-insensitive, and the de-hyphenated form matches "CI setup" + assert eg.mask_name("CI setup helps", "ci-setup") == "this capability helps" + + +def test_mask_name_does_not_overmatch_substrings(): + # a short name must not blast through unrelated words + assert ( + eg.mask_name("specific scientific terms", "ci") == "specific scientific terms" + ) + + +def test_leaks_name_detection(): + assert eg._leaks_name("please use projreg now", "projreg") + assert not eg._leaks_name("please use the registry now", "projreg") + + +def test_parse_lines_strips_bullets_and_numbering(): + text = '1. first\n- second\n * third \n\n"fourth"' + assert eg._parse_lines(text) == ["first", "second", "third", "fourth"] + + +# --------------------------------------------------------------------------- # +# generate_cases +# --------------------------------------------------------------------------- # + + +def test_generate_cases_basic_shape(): + docs = { + "alpha": { + "name": "alpha", + "description": "alpha tool for foo widget tasks here", + }, + "beta": {"name": "beta", "description": "beta tool for bar gadget tasks here"}, + } + cases = eg.generate_cases( + _skill_source(docs), k=2, query_generator=_echo_gen, corpus_name="g" + ) + assert len(cases) == 4 # 2 per artifact + assert {c.gold[0] for c in cases} == {"alpha", "beta"} + assert all(len(c.gold) == 1 and c.source_id == c.gold[0] for c in cases) + assert all(c.corpus == "g" for c in cases) + # masking applied: no query leaks its gold artifact's name + assert all(not eg._leaks_name(c.query, c.gold[0]) for c in cases) + assert all(c.metadata["masked"] is True for c in cases) + + +def test_generate_cases_output_guard_drops_leaks(): + docs = { + "alpha": {"name": "alpha", "description": "a long enough description of foo"} + } + + def leaky(description, *, n): + return ["please use alpha to do it", "do the foo thing now"] + + cases = eg.generate_cases(_skill_source(docs), k=2, query_generator=leaky) + # the query that leaked the gold name ("alpha") is dropped; the clean one kept + assert [c.query for c in cases] == ["do the foo thing now"] + + +def test_generate_cases_can_disable_masking(): + docs = {"alpha": {"name": "alpha", "description": "alpha does foo things for you"}} + cases = eg.generate_cases( + _skill_source(docs), k=1, query_generator=_echo_gen, mask_names=False + ) + # with masking off, the name survives in the echoed description and is kept + assert cases and "alpha" in cases[0].query + assert cases[0].metadata["masked"] is False + + +def test_generate_cases_skips_short_descriptions_and_warns(): + docs = {"x": {"name": "x", "description": "too short"}} # < min_chars + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + cases = eg.generate_cases( + _skill_source(docs), query_generator=_echo_gen, min_chars=20 + ) + assert cases == [] + assert any("skipped" in str(w.message) for w in caught) + + +def test_generate_cases_max_artifacts(): + docs = { + f"a{i}": {"name": f"a{i}", "description": f"capability number {i} doing tasks"} + for i in range(5) + } + cases = eg.generate_cases( + _skill_source(docs), k=1, query_generator=_echo_gen, max_artifacts=2 + ) + assert {c.gold[0] for c in cases} == {"a0", "a1"} + + +# --------------------------------------------------------------------------- # +# abstention + build_eval_set +# --------------------------------------------------------------------------- # + + +def test_generate_abstention_cases(): + def gen(*, n, theme): + return [f"off-topic request {i} about {theme}" for i in range(n)] + + cases = eg.generate_abstention_cases(3, generator=gen, theme="cooking") + assert len(cases) == 3 + assert all(c.gold_is_none for c in cases) + assert all(c.metadata["generator"] == "abstention" for c in cases) + + +def test_generate_abstention_cases_zero_is_empty(): + assert eg.generate_abstention_cases(0, generator=lambda **k: []) == [] + + +def test_build_eval_set_hits_abstention_fraction(): + docs = { + f"a{i}": {"name": f"a{i}", "description": f"capability number {i} doing tasks"} + for i in range(5) + } + + def qg(description, *, n): + return [description] # one gold case per artifact + + def ag(*, n, theme): + return [f"unsupported {i}" for i in range(n)] + + cases = eg.build_eval_set( + _skill_source(docs), + k=1, + abstention_frac=0.25, + query_generator=qg, + abstention_generator=ag, + ) + n_abstain = sum(c.gold_is_none for c in cases) + # 5 gold -> ceil(0.25 * 5 / 0.75) = 2 abstain -> 2/7 >= 0.25 + assert n_abstain == 2 and len(cases) == 7 + assert n_abstain / len(cases) >= 0.25 + + +# --------------------------------------------------------------------------- # +# Generated set is scorable by ir.eval (end-to-end, offline) +# --------------------------------------------------------------------------- # + + +def test_generated_set_is_scorable_end_to_end(): + docs = { + "alpha": {"name": "alpha", "description": "alpha tool for foo widget tasks"}, + "beta": {"name": "beta", "description": "beta tool for bar gadget chores"}, + } + source = _skill_source(docs) + cases = eg.generate_cases(source, k=1, query_generator=_echo_gen, corpus_name="g") + corpus = ir.build(source, store=CorpusStore.memory(), embedder="light") + report = ev.evaluate_discovery(corpus, cases, mode="dense", primary_k=1) + assert report.n_gold == 2 + # the masked-description query still retrieves its own artifact first + assert report.retrieval.metrics["recall@1"] == pytest.approx(1.0) + + +def test_save_load_with_signature_meta(tmp_path): + docs = {"a": {"name": "a", "description": "a capability doing foo tasks for you"}} + source = _skill_source(docs) + cases = eg.build_eval_set( + source, + k=1, + abstention_frac=0.0, + query_generator=_echo_gen, + ) + path = tmp_path / "cases.jsonl" + ev.save_cases(cases, path, meta={"corpus_signature": eg.corpus_signature(source)}) + assert ev.load_cases(path) == cases + + +# --------------------------------------------------------------------------- # +# corpus_signature +# --------------------------------------------------------------------------- # + + +def test_corpus_signature_is_order_independent(): + docs = { + "a": {"name": "a", "description": "x"}, + "b": {"name": "b", "description": "y"}, + } + src1 = _skill_source(docs) + src2 = _skill_source(dict(reversed(list(docs.items())))) + assert eg.corpus_signature(src1) == eg.corpus_signature(src2) + + +def test_corpus_signature_changes_with_membership(): + base = {"a": {"name": "a", "description": "x"}} + more = { + "a": {"name": "a", "description": "x"}, + "b": {"name": "b", "description": "y"}, + } + assert eg.corpus_signature(_skill_source(base)) != eg.corpus_signature( + _skill_source(more) + ) + + +# --------------------------------------------------------------------------- # +# _parse_lines — strips real markers, never meaningful leading characters +# --------------------------------------------------------------------------- # + + +def test_parse_lines_preserves_leading_tokens(): + text = "\n".join( + [ + "3D modeling help", + "-9 degrees, what should I wear?", + ".env file handling", + "24/7 monitoring setup", + "2024 tax question", + ] + ) + assert eg._parse_lines(text) == [ + "3D modeling help", + "-9 degrees, what should I wear?", + ".env file handling", + "24/7 monitoring setup", + "2024 tax question", + ] + + +def test_parse_lines_still_strips_real_markers(): + assert eg._parse_lines("1. a\n2) b\n- c\n* d\n• e") == ["a", "b", "c", "d", "e"] + + +# --------------------------------------------------------------------------- # +# Masking — multi-word / whitespace / degenerate / self-referential +# --------------------------------------------------------------------------- # + + +def test_mask_name_whitespace_and_separator_tolerant(): + assert eg.mask_name("the data sync job", "data-sync") == "the this capability job" + assert eg.mask_name("the data sync job", "data sync") == "the this capability job" + assert eg.mask_name("run datasync now", "data-sync") == "run this capability now" + + +def test_leaks_name_multiword_reordered_tokens(): + # bag-of-words: reordered/separated name tokens count as a leak + assert eg._leaks_name("sync my data now", "data-sync") + assert eg._leaks_name("i need data and a sync", "data sync") + # a single shared token is NOT a leak (legitimate content overlap) + assert not eg._leaks_name("just sync my files", "data-sync") + + +def test_mask_name_drops_degenerate_short_tokens(): + # a 1-char / separator-only name must not mask stray letters + assert eg.mask_name("an example with x here", "-x") == "an example with x here" + # but a real 2-char name is still masked + assert eg.mask_name("the ci runs", "ci") == "the this capability runs" + + +def test_mask_name_avoids_self_referential_placeholder(): + out = eg.mask_name("The Capability is great", "capability") + assert "this capability" not in out.lower() + assert eg._ALT_PLACEHOLDER in out + + +# --------------------------------------------------------------------------- # +# _default_describe / _name_of fallbacks +# --------------------------------------------------------------------------- # + + +def test_default_describe_fallbacks(): + assert eg._default_describe({"description": "d", "text": "t"}) == "d" + assert eg._default_describe({"text": "only text here"}) == "only text here" + assert eg._default_describe({"name": "n", "summary": "s"}) == "n\ns" # joined strs + assert eg._default_describe("bare string") == "bare string" + assert eg._default_describe({"description": " ", "text": "fallback"}) == "fallback" + + +def test_name_of_coerces_missing_or_bad_name(): + assert eg._name_of("id1", {"name": None}) == "id1" + assert eg._name_of("id1", {"name": ""}) == "id1" + assert eg._name_of("id1", {"name": 42}) == "id1" + assert eg._name_of("id1", {"name": "real"}) == "real" + assert eg._name_of("id1", "bare") == "id1" + + +# --------------------------------------------------------------------------- # +# Validation + determinism + resilience +# --------------------------------------------------------------------------- # + + +def test_generate_cases_rejects_bad_k(): + src = _skill_source({"a": {"name": "a", "description": "x" * 30}}) + with pytest.raises(ValueError): + eg.generate_cases(src, k=0, query_generator=_echo_gen) + + +def test_build_eval_set_rejects_out_of_range_frac(): + src = _skill_source({"a": {"name": "a", "description": "x" * 30}}) + with pytest.raises(ValueError): + eg.build_eval_set(src, abstention_frac=1.0, query_generator=_echo_gen) + with pytest.raises(ValueError): + eg.build_eval_set(src, abstention_frac=-0.1, query_generator=_echo_gen) + + +def test_generate_cases_max_artifacts_is_deterministic_sorted(): + docs = { + k: {"name": k, "description": f"{k} capability doing tasks here"} + for k in ["zeta", "alpha", "mu"] + } + cases = eg.generate_cases( + _skill_source(docs), k=1, query_generator=_echo_gen, max_artifacts=2 + ) + # sorted-id subset (alpha, mu), NOT the insertion-order first two (zeta, alpha) + assert sorted({c.gold[0] for c in cases}) == ["alpha", "mu"] + + +def test_generate_cases_skips_on_generator_exception(): + docs = { + "good": {"name": "good", "description": "good capability for foo tasks here"}, + "bad": {"name": "bad", "description": "bad capability triggering boom now ok"}, + } + + def flaky(description, *, n): + if "boom" in description: + raise RuntimeError("kaboom") + return [description] + + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + cases = eg.generate_cases(_skill_source(docs), k=1, query_generator=flaky) + assert {c.gold[0] for c in cases} == {"good"} + assert any("bad" in str(w.message) and "failed" in str(w.message) for w in caught) + + +def test_build_eval_set_forwards_mask_names_false(): + docs = {"alpha": {"name": "alpha", "description": "alpha does foo things for you"}} + cases = eg.build_eval_set( + _skill_source(docs), + k=1, + abstention_frac=0.0, + query_generator=_echo_gen, + mask_names=False, + ) + assert cases and cases[0].metadata["masked"] is False + assert "alpha" in cases[0].query + + +# --------------------------------------------------------------------------- # +# Default oa-backed generators — prompt assembly (skipped if oa absent) +# --------------------------------------------------------------------------- # + + +def test_oa_prompts_substitute_placeholders(): + oa = pytest.importorskip("oa") + # prompt-only (prompt_func=None): assert on the ASSEMBLED PROMPT, not the + # generate() wrapper (which, with prompt_func=None, would iterate the string). + bt = oa.prompt_function(eg.BACKTRANSLATION_PROMPT, name="bt", prompt_func=None) + prompt = bt(description="DESC_MARKER", n=4) + assert "DESC_MARKER" in prompt and "Write 4 natural" in prompt + ab = oa.prompt_function(eg.ABSTENTION_PROMPT, name="ab", prompt_func=None) + assert "THEME_MARKER" in ab(theme="THEME_MARKER", n=2) + + +def test_parse_lines_is_wired_as_egress(): + oa = pytest.importorskip("oa") + + def fake_llm(*args, **kwargs): + return "1. foo\n- bar" + + fn = oa.prompt_function( + "x {description} {n}", egress=eg._parse_lines, prompt_func=fake_llm + ) + assert fn(description="d", n=2) == ["foo", "bar"] + + +# --------------------------------------------------------------------------- # +# CLI eval-gen glue (build_eval_set stubbed; offline) +# --------------------------------------------------------------------------- # + + +def test_cli_eval_gen_glue(tmp_path, monkeypatch): + import json + + monkeypatch.setenv("IR_CONFIG_DIR", str(tmp_path / "config")) + monkeypatch.setenv("IR_DATA_DIR", str(tmp_path / "data")) + monkeypatch.setenv("IR_CACHE_DIR", str(tmp_path / "cache")) + from ir import cli + + docs = tmp_path / "docs" + docs.mkdir() + (docs / "a.md").write_text("alpha content about widgets and gadgets here") + cli.register("notes", "files", root=str(docs), pattern=r".*\.md$") + + captured = {} + + def fake_build(source, **kw): + captured.update(kw) + return [ev.DiscoveryCase("q1", gold=("a.md",)), ev.DiscoveryCase("q2", gold=())] + + # cli.eval_gen imports build_eval_set from ir.eval_gen at call time + monkeypatch.setattr(eg, "build_eval_set", fake_build) + + out_path = tmp_path / "cases.jsonl" + msg = cli.eval_gen("notes", str(out_path), k=3, max_artifacts="5") + assert "2 cases" in msg and "1 gold" in msg and "1 abstention" in msg + assert captured["k"] == 3 and captured["max_artifacts"] == 5 # int-cast + forwarded + assert captured["corpus_name"] == "notes" + + loaded = ev.load_cases(out_path) + assert len(loaded) == 2 + header = json.loads(out_path.read_text(encoding="utf-8").splitlines()[0])[ + "__meta__" + ] + assert header["corpus"] == "notes" and "corpus_signature" in header + + +def test_eval_gen_reachable_via_ir_namespace(): + assert hasattr(ir, "eval_gen") + assert ir.eval_gen.generate_cases is eg.generate_cases