From e0ec712017e8742c23663c820813937ea8f2188c Mon Sep 17 00:00:00 2001 From: Anthony Stirling <77850077+Frooodle@users.noreply.github.com> Date: Wed, 2 Sep 2026 10:11:29 +0100 Subject: [PATCH] feat(engine): add lexical matching to operation retrieval --- .../stirling/agents/operation_shortlist.py | 59 ++++++++++++++++++- engine/tests/test_operation_shortlist.py | 36 ++++++++++- 2 files changed, 91 insertions(+), 4 deletions(-) diff --git a/engine/src/stirling/agents/operation_shortlist.py b/engine/src/stirling/agents/operation_shortlist.py index 3254a67dd4..a8039deaa1 100644 --- a/engine/src/stirling/agents/operation_shortlist.py +++ b/engine/src/stirling/agents/operation_shortlist.py @@ -4,12 +4,22 @@ from __future__ import annotations import asyncio import logging +import math +import re +from collections import Counter from stirling.documents import EmbeddingService from stirling.models import OPERATIONS, ToolEndpoint logger = logging.getLogger(__name__) +_WORD = re.compile(r"[a-z0-9]+") +_STOPWORDS = frozenset({"pdf", "the", "this", "a", "an", "of", "to", "and", "for", "it", "into", "out"}) + +_BM25_K1 = 1.5 +_BM25_B = 0.75 +_RRF_K = 60 + def retrieval_text(operation: ToolEndpoint) -> str: model = OPERATIONS[operation] @@ -18,6 +28,10 @@ def retrieval_text(operation: ToolEndpoint) -> str: return f"{operation.name.replace('_', ' ').lower()}. {description} {' '.join(parameters)}".strip() +def tokenize(text: str) -> list[str]: + return [word for word in _WORD.findall(text.lower().replace("_", " ")) if word not in _STOPWORDS] + + def _cosine(left: list[float], right: list[float]) -> float: dot = sum(a * b for a, b in zip(left, right, strict=True)) left_norm = sum(a * a for a in left) ** 0.5 @@ -25,18 +39,52 @@ def _cosine(left: list[float], right: list[float]) -> float: return dot / (left_norm * right_norm) if left_norm and right_norm else 0.0 +def bm25_scores(query: str, corpus: list[list[str]]) -> list[float]: + lengths = [len(document) for document in corpus] + average_length = sum(lengths) / len(lengths) if lengths else 0.0 + document_frequency: Counter[str] = Counter() + for document in corpus: + document_frequency.update(set(document)) + + total = len(corpus) + scores = [0.0] * total + for term in tokenize(query): + frequency = document_frequency.get(term, 0) + if not frequency: + continue + idf = math.log(1 + (total - frequency + 0.5) / (frequency + 0.5)) + for index, document in enumerate(corpus): + occurrences = document.count(term) + if occurrences: + length_norm = 1 - _BM25_B + _BM25_B * lengths[index] / average_length + scores[index] += idf * (occurrences * (_BM25_K1 + 1)) / (occurrences + _BM25_K1 * length_norm) + return scores + + +def rank_fusion(*rankings: list[ToolEndpoint]) -> list[ToolEndpoint]: + """Reciprocal rank fusion, which combines rankings without needing comparable scores.""" + fused: dict[ToolEndpoint, float] = {} + for ranking in rankings: + for position, operation in enumerate(ranking): + fused[operation] = fused.get(operation, 0.0) + 1.0 / (_RRF_K + position + 1) + return sorted(fused, key=lambda operation: -fused[operation]) + + class OperationShortlist: def __init__(self, embedder: EmbeddingService) -> None: self._embedder = embedder self._vectors: dict[ToolEndpoint, list[float]] | None = None + self._tokens: dict[ToolEndpoint, list[str]] = {} self._lock = asyncio.Lock() async def _catalogue_vectors(self) -> dict[ToolEndpoint, list[float]]: async with self._lock: if self._vectors is None: operations = list(OPERATIONS) - embeddings = await self._embedder.embed_documents([retrieval_text(op) for op in operations]) + texts = [retrieval_text(operation) for operation in operations] + embeddings = await self._embedder.embed_documents(texts) self._vectors = dict(zip(operations, embeddings, strict=True)) + self._tokens = {operation: tokenize(text) for operation, text in zip(operations, texts, strict=True)} return self._vectors async def select( @@ -58,5 +106,10 @@ class OperationShortlist: except Exception: # noqa: BLE001 - any embedding failure should fall back, never fail planning logger.warning("[pdf-edit] operation ranking unavailable, showing the full catalogue", exc_info=True) return operations - ranked = sorted(operations, key=lambda op: -_cosine(query, vectors[op])) - return ranked[:limit] + + semantic = sorted(operations, key=lambda operation: -_cosine(query, vectors[operation])) + corpus = [self._tokens.get(operation, []) for operation in operations] + lexical_scores = bm25_scores(message, corpus) + by_score = dict(zip(operations, lexical_scores, strict=True)) + lexical = sorted(operations, key=lambda operation: -by_score[operation]) + return rank_fusion(semantic, lexical)[:limit] diff --git a/engine/tests/test_operation_shortlist.py b/engine/tests/test_operation_shortlist.py index 69016dad77..a5c2fcaef9 100644 --- a/engine/tests/test_operation_shortlist.py +++ b/engine/tests/test_operation_shortlist.py @@ -2,7 +2,13 @@ from __future__ import annotations import pytest -from stirling.agents.operation_shortlist import OperationShortlist, retrieval_text +from stirling.agents.operation_shortlist import ( + OperationShortlist, + bm25_scores, + rank_fusion, + retrieval_text, + tokenize, +) from stirling.models import OPERATIONS, ToolEndpoint @@ -74,3 +80,31 @@ def test_retrieval_text_carries_parameter_descriptions() -> None: assert "add watermark" in text assert "opacity" in text.lower() + + +def test_lexical_scoring_ranks_the_operation_named_in_the_request() -> None: + operations = list(OPERATIONS) + corpus = [tokenize(retrieval_text(operation)) for operation in operations] + + scores = bm25_scores("add a trusted timestamp", corpus) + best = max(zip(operations, scores, strict=True), key=lambda pair: pair[1])[0] + + assert best is ToolEndpoint.TIMESTAMP_PDF + + +def test_tokenize_splits_endpoint_names_and_drops_stopwords() -> None: + assert tokenize("ADD_WATERMARK to the PDF") == ["add", "watermark"] + + +def test_rank_fusion_promotes_what_both_rankings_rate_highly() -> None: + first = [ToolEndpoint.ADD_WATERMARK, ToolEndpoint.FLATTEN, ToolEndpoint.REPAIR] + second = [ToolEndpoint.ADD_WATERMARK, ToolEndpoint.REPAIR, ToolEndpoint.FLATTEN] + + assert rank_fusion(first, second)[0] is ToolEndpoint.ADD_WATERMARK + + +def test_rank_fusion_surfaces_an_operation_only_one_ranking_found() -> None: + semantic = [ToolEndpoint.FLATTEN, ToolEndpoint.REPAIR] + lexical = [ToolEndpoint.TIMESTAMP_PDF, ToolEndpoint.FLATTEN] + + assert ToolEndpoint.TIMESTAMP_PDF in rank_fusion(semantic, lexical)