feat(engine): show the planner only the operations that match the request

This commit is contained in:
Anthony Stirling
2026-09-02 10:09:09 +01:00
parent 2cf355c5cd
commit c0e793c804
5 changed files with 162 additions and 15 deletions
+5
View File
@@ -48,6 +48,11 @@ STIRLING_RAG_TOP_K=20
# rather than chain more searches. # rather than chain more searches.
STIRLING_RAG_MAX_SEARCHES=5 STIRLING_RAG_MAX_SEARCHES=5
# How many operations the edit planner sees, ranked against the request by embedding
# similarity. The full catalogue overruns a local model's context window. Set 0 to
# disable ranking and show every operation.
STIRLING_PLANNER_SHORTLIST_SIZE=20
# Chunked reasoner settings: how big each per-worker slice is (in characters), # Chunked reasoner settings: how big each per-worker slice is (in characters),
# how many workers may run in parallel against the fast model, and how long # how many workers may run in parallel against the fast model, and how long
# any single worker is allowed to wait for a response before being abandoned. # any single worker is allowed to wait for a response before being abandoned.
@@ -0,0 +1,62 @@
"""Narrows the operation catalogue to the candidates worth showing the planner."""
from __future__ import annotations
import asyncio
import logging
from stirling.documents import EmbeddingService
from stirling.models import OPERATIONS, ToolEndpoint
logger = logging.getLogger(__name__)
def retrieval_text(operation: ToolEndpoint) -> str:
model = OPERATIONS[operation]
description = (model.model_json_schema().get("description") or "").strip()
parameters = [(field.description or "").strip() for field in model.model_fields.values() if field.description]
return f"{operation.name.replace('_', ' ').lower()}. {description} {' '.join(parameters)}".strip()
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
right_norm = sum(b * b for b in right) ** 0.5
return dot / (left_norm * right_norm) if left_norm and right_norm else 0.0
class OperationShortlist:
def __init__(self, embedder: EmbeddingService) -> None:
self._embedder = embedder
self._vectors: dict[ToolEndpoint, list[float]] | None = None
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])
self._vectors = dict(zip(operations, embeddings, strict=True))
return self._vectors
async def select(
self,
message: str,
operations: list[ToolEndpoint],
limit: int,
) -> list[ToolEndpoint]:
"""The ``limit`` operations closest to ``message``, or all of them if ranking is unavailable.
Falling back to the full list keeps a planner that would otherwise work on a deployment
with no reachable embedding provider, at the cost of a much larger prompt.
"""
if limit <= 0 or len(operations) <= limit:
return operations
try:
vectors = await self._catalogue_vectors()
query = await self._embedder.embed_query(message)
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]
+15 -15
View File
@@ -9,6 +9,7 @@ from pydantic_ai import Agent
from pydantic_ai.output import NativeOutput from pydantic_ai.output import NativeOutput
from stirling.agents._page_text import format_page_text, get_extracted_text_artifact, has_page_text from stirling.agents._page_text import format_page_text, get_extracted_text_artifact, has_page_text
from stirling.agents.operation_shortlist import OperationShortlist
from stirling.contracts import ( from stirling.contracts import (
EditCannotDoResponse, EditCannotDoResponse,
EditClarificationRequest, EditClarificationRequest,
@@ -175,6 +176,7 @@ class PdfEditAgent:
def __init__(self, runtime: AppRuntime) -> None: def __init__(self, runtime: AppRuntime) -> None:
self.runtime = runtime self.runtime = runtime
self.parameter_selector = PdfEditParameterSelector(runtime) self.parameter_selector = PdfEditParameterSelector(runtime)
self.shortlist = OperationShortlist(runtime.documents.embedder)
async def orchestrate(self, request: OrchestratorRequest) -> PdfEditResponse: async def orchestrate(self, request: OrchestratorRequest) -> PdfEditResponse:
"""Entry point for the orchestrator delegate — adapts the orchestrator's """Entry point for the orchestrator delegate — adapts the orchestrator's
@@ -277,13 +279,20 @@ class PdfEditAgent:
repair_note: str = "", repair_note: str = "",
) -> PdfEditPlanOutput: ) -> PdfEditPlanOutput:
can_request_content = allow_need_content and not has_page_text(request.page_text) can_request_content = allow_need_content and not has_page_text(request.page_text)
available = list(supported_operations)
candidates = await self.shortlist.select(
request.user_message,
available,
self.runtime.settings.planner_shortlist_size,
)
logger.info("[pdf-edit] showing %d of %d operations", len(candidates), len(available))
agent = self._build_selection_agent( agent = self._build_selection_agent(
supported_operations, candidates,
unavailable_operations, unavailable_operations,
allow_need_content=can_request_content, allow_need_content=can_request_content,
) )
return await agent.select( return await agent.select(
self._build_selection_prompt(request, supported_operations, unavailable_operations, repair_note) self._build_selection_prompt(request, candidates, unavailable_operations, repair_note)
) )
def _build_selection_agent( def _build_selection_agent(
@@ -359,19 +368,10 @@ class PdfEditAgent:
f"Extracted page text:\n{format_page_text(request.page_text)}" f"Extracted page text:\n{format_page_text(request.page_text)}"
) )
# Endpoints that exist on the server and are callable via the direct API or the manual UI, # Hidden from the AI planner only; still live for the manual UI, direct API and pipelines.
# but are never offered to the AI agent as a routing option. # AUTO_REDACT and REDACT pre-date the unified operation model and take a less expressive
# # schema, so AI redaction is channelled through REDACT_EXECUTE. Add an endpoint here when a
# Why: REDACT_EXECUTE is the preferred AI-driven redaction route. AUTO_REDACT and REDACT are # legacy one has a preferred replacement the AI should use exclusively.
# legacy endpoints that remain fully functional for human callers (the manual redact UI, direct
# API consumers, pipelines) but would produce a worse experience if the AI routed to them —
# they accept a simpler, less expressive schema and pre-date the unified operation model.
# Hiding them here channels all AI redaction traffic through REDACT_EXECUTE without disabling
# the legacy endpoints for anyone else.
#
# How to reuse: add an endpoint here whenever a legacy endpoint has a preferred replacement
# that the AI should use exclusively. The endpoint remains live on the server; only the AI
# planner is prevented from selecting it.
_AGENT_HIDDEN_ENDPOINTS: frozenset[ToolEndpoint] = frozenset({ToolEndpoint.AUTO_REDACT, ToolEndpoint.REDACT}) _AGENT_HIDDEN_ENDPOINTS: frozenset[ToolEndpoint] = frozenset({ToolEndpoint.AUTO_REDACT, ToolEndpoint.REDACT})
def _classify_operations(self, request: PdfEditRequest) -> tuple[list[ToolEndpoint], list[ToolEndpoint]]: def _classify_operations(self, request: PdfEditRequest) -> tuple[list[ToolEndpoint], list[ToolEndpoint]]:
+4
View File
@@ -50,6 +50,10 @@ class AppSettings(BaseSettings):
rag_chunk_overlap: int = Field(validation_alias="STIRLING_RAG_CHUNK_OVERLAP") rag_chunk_overlap: int = Field(validation_alias="STIRLING_RAG_CHUNK_OVERLAP")
rag_default_top_k: int = Field(validation_alias="STIRLING_RAG_TOP_K") rag_default_top_k: int = Field(validation_alias="STIRLING_RAG_TOP_K")
rag_max_searches: int = Field(validation_alias="STIRLING_RAG_MAX_SEARCHES") rag_max_searches: int = Field(validation_alias="STIRLING_RAG_MAX_SEARCHES")
# How many of the catalogue's operations the edit planner is shown. The full list
# overruns a local model's context; 0 disables ranking and shows every operation.
planner_shortlist_size: int = Field(default=20, validation_alias="STIRLING_PLANNER_SHORTLIST_SIZE")
documents_reaper_interval_seconds: int = Field( documents_reaper_interval_seconds: int = Field(
default=900, default=900,
validation_alias="STIRLING_DOCUMENTS_REAPER_INTERVAL_SECONDS", validation_alias="STIRLING_DOCUMENTS_REAPER_INTERVAL_SECONDS",
+76
View File
@@ -0,0 +1,76 @@
from __future__ import annotations
import pytest
from stirling.agents.operation_shortlist import OperationShortlist, retrieval_text
from stirling.models import OPERATIONS, ToolEndpoint
class StubEmbedder:
"""Embeds on a single axis: how often 'watermark' appears, so ranking is predictable."""
def __init__(self, fail: bool = False) -> None:
self.fail = fail
self.document_calls = 0
@staticmethod
def _vector(text: str) -> list[float]:
return [float(text.lower().count("watermark")), 1.0]
async def embed_documents(self, texts: list[str]) -> list[list[float]]:
if self.fail:
raise RuntimeError("embedding provider unreachable")
self.document_calls += 1
return [self._vector(text) for text in texts]
async def embed_query(self, text: str) -> list[float]:
if self.fail:
raise RuntimeError("embedding provider unreachable")
return self._vector(text)
@pytest.mark.anyio
async def test_returns_every_operation_when_the_limit_is_not_binding() -> None:
shortlist = OperationShortlist(StubEmbedder())
operations = list(OPERATIONS)[:5]
assert await shortlist.select("watermark this", operations, len(operations)) == operations
assert await shortlist.select("watermark this", operations, 0) == operations
@pytest.mark.anyio
async def test_narrows_to_the_closest_operations() -> None:
shortlist = OperationShortlist(StubEmbedder())
operations = list(OPERATIONS)
selected = await shortlist.select("watermark watermark watermark", operations, 5)
assert len(selected) == 5
assert ToolEndpoint.ADD_WATERMARK in selected
@pytest.mark.anyio
async def test_falls_back_to_the_full_catalogue_when_embedding_fails() -> None:
shortlist = OperationShortlist(StubEmbedder(fail=True))
operations = list(OPERATIONS)
assert await shortlist.select("watermark this", operations, 5) == operations
@pytest.mark.anyio
async def test_embeds_the_catalogue_once() -> None:
embedder = StubEmbedder()
shortlist = OperationShortlist(embedder)
operations = list(OPERATIONS)
await shortlist.select("watermark this", operations, 5)
await shortlist.select("rotate this", operations, 5)
assert embedder.document_calls == 1
def test_retrieval_text_carries_parameter_descriptions() -> None:
text = retrieval_text(ToolEndpoint.ADD_WATERMARK)
assert "add watermark" in text
assert "opacity" in text.lower()