From b85e5bb4ba17ecc43347cc49918cc39dfc5c953c Mon Sep 17 00:00:00 2001 From: zccyman Date: Sat, 16 May 2026 05:03:39 +0800 Subject: [PATCH] feat(plugins): allow plugins to register custom @-prefix context references Closes #26193 Adds ContextReferenceProvider ABC so plugins can register custom @-prefixes (e.g. @issue:ENG-123) with autocomplete and expansion. Plugin output flows through existing token-limit guards. Zero breaking changes. --- agent/context_references.py | 99 +++++++++ hermes_cli/plugins.py | 36 ++++ tests/agent/test_plugin_context_references.py | 197 ++++++++++++++++++ tui_gateway/methods_complete.py | 54 +++++ 4 files changed, 386 insertions(+) create mode 100644 tests/agent/test_plugin_context_references.py diff --git a/agent/context_references.py b/agent/context_references.py index a44b90d907..dd36e67672 100644 --- a/agent/context_references.py +++ b/agent/context_references.py @@ -15,10 +15,78 @@ from agent.model_metadata import estimate_tokens_rough from hermes_cli._subprocess_compat import IS_WINDOWS, windows_hide_flags from hermes_cli.sizefmt import format_bytes +from abc import ABC, abstractmethod + +# --------------------------------------------------------------------------- +# Plugin context-reference provider API (Issue #26193) +# --------------------------------------------------------------------------- + +BUILTIN_PREFIXES = frozenset({"diff", "staged", "file", "folder", "git", "url"}) + +_context_reference_providers: dict[str, "ContextReferenceProvider"] = {} + + +class ContextCompletionItem: + """A single autocomplete result from a context reference provider.""" + + __slots__ = ("text", "display", "meta") + + def __init__(self, text: str, display: str = "", meta: str = "") -> None: + self.text = text + self.display = display or text + self.meta = meta + + +class ContextReferenceProvider(ABC): + """Base class for plugin-registered @-prefix context reference providers. + + Plugins subclass this and register via + ``PluginContext.register_context_reference()``. + """ + + prefix: str = "" # e.g. "issue", "channel", "doc" + description: str = "" # shown in autocomplete meta column + + @abstractmethod + async def autocomplete(self, query: str, *, limit: int = 10) -> list[ContextCompletionItem]: + """Return autocomplete items for the given query string.""" + ... + + @abstractmethod + async def expand(self, target: str) -> str | None: + """Expand *target* to prompt content. Return ``None`` to skip.""" + ... + + +def register_context_reference_provider(provider: ContextReferenceProvider) -> None: + """Register a plugin context reference provider.""" + if not isinstance(provider, ContextReferenceProvider): + raise TypeError("provider must be a ContextReferenceProvider instance") + prefix = provider.prefix.lower().strip() + if not prefix: + raise ValueError("prefix must be a non-empty string") + if prefix in BUILTIN_PREFIXES: + raise ValueError(f"prefix '{prefix}' is reserved for built-in references") + if prefix in _context_reference_providers: + raise ValueError(f"prefix '{prefix}' is already registered") + _context_reference_providers[prefix] = provider + + +def get_context_reference_providers() -> dict[str, ContextReferenceProvider]: + """Return a snapshot of all registered plugin providers.""" + return dict(_context_reference_providers) + + _QUOTED_REFERENCE_VALUE = r'(?:`[^`\n]+`|"[^"\n]+"|\'[^\'\n]+\')' REFERENCE_PATTERN = re.compile( rf"(?diff|staged)\b|(?Pfile|folder|git|url):(?P{_QUOTED_REFERENCE_VALUE}(?::\d+(?:-\d+)?)?|\S+))" ) +# Plugin fallback pattern – catches any @: not handled by the +# built-in regex so that plugin-registered prefixes can be resolved. +_PLUGIN_REFERENCE_PATTERN = re.compile( + rf"(?[a-zA-Z][a-zA-Z0-9_-]*):(?P{_QUOTED_REFERENCE_VALUE}(?::\d+(?:-\d+)?)?|\S+)" +) + TRAILING_PUNCTUATION = ",.;!?" _NEEDS_QUOTING = re.compile(r"""[\s()\[\]{}<>"'`]""") _SENSITIVE_HOME_DIRS = (".ssh", ".aws", ".gnupg", ".kube", ".docker", ".azure", ".config/gh") @@ -117,6 +185,27 @@ def parse_context_references(message: str) -> list[ContextReference]: ) ) + # Second pass: resolve plugin-registered prefixes the built-in pattern missed + if _context_reference_providers: + for match in _PLUGIN_REFERENCE_PATTERN.finditer(message): + kind = match.group("kind") + if kind in BUILTIN_PREFIXES: + continue + # Skip if already captured by the built-in pattern + if any(r.kind == kind and r.start == match.start() for r in refs): + continue + if kind in _context_reference_providers: + value = _strip_trailing_punctuation(match.group("value") or "") + refs.append( + ContextReference( + raw=match.group(0), + kind=kind, + target=_strip_reference_wrappers(value), + start=match.start(), + end=match.end(), + ) + ) + return refs @@ -263,6 +352,16 @@ async def _expand_reference( except Exception as exc: return f"{ref.raw}: {exc}", None + # Plugin-provided context references + provider = _context_reference_providers.get(ref.kind) + if provider is not None: + try: + plugin_content = await provider.expand(ref.target) + if plugin_content is not None: + return None, f"📌 {ref.raw} ({estimate_tokens_rough(plugin_content)} tokens)\n{plugin_content}" + except Exception as exc: + return f"{ref.raw}: plugin expansion error: {exc}", None + return f"{ref.raw}: unsupported reference type", None diff --git a/hermes_cli/plugins.py b/hermes_cli/plugins.py index 94d13fac18..dadd4d2854 100644 --- a/hermes_cli/plugins.py +++ b/hermes_cli/plugins.py @@ -1473,6 +1473,42 @@ class PluginContext: self.manifest.name, engine.name, ) + # -- context reference registration ------------------------------------- + + def register_context_reference(self, provider) -> None: + """Register a custom @-prefix context reference provider. + + ``provider`` must be an instance of + :class:`agent.context_references.ContextReferenceProvider`. The + ``provider.prefix`` attribute defines the @-prefix (e.g. ``"issue"`` + creates ``@issue:...``). Built-in prefixes (diff, staged, file, + folder, git, url) are reserved and will be rejected. + """ + from agent.context_references import ( + ContextReferenceProvider as _CRP, + register_context_reference_provider as _register, + ) + + if not isinstance(provider, _CRP): + logger.warning( + "Plugin '%s' tried to register a context reference provider " + "that does not inherit from ContextReferenceProvider. Ignoring.", + self.manifest.name, + ) + return + try: + _register(provider) + except ValueError as exc: + logger.warning( + "Plugin '%s' context reference registration failed: %s", + self.manifest.name, exc, + ) + return + logger.info( + "Plugin '%s' registered context reference: @%s:", + self.manifest.name, provider.prefix, + ) + # -- image gen provider registration ------------------------------------ def register_image_gen_provider(self, provider) -> None: diff --git a/tests/agent/test_plugin_context_references.py b/tests/agent/test_plugin_context_references.py new file mode 100644 index 0000000000..f0bbfd33a9 --- /dev/null +++ b/tests/agent/test_plugin_context_references.py @@ -0,0 +1,197 @@ +"""Tests for plugin context reference provider API (Issue #26193).""" + +from __future__ import annotations + +import asyncio +from pathlib import Path + +import pytest + +from agent.context_references import ( + BUILTIN_PREFIXES, + ContextCompletionItem, + ContextReferenceProvider, + _PLUGIN_REFERENCE_PATTERN, + _context_reference_providers, + get_context_reference_providers, + parse_context_references, + register_context_reference_provider, +) + + +# -- helpers --------------------------------------------------------------- + +class _DummyProvider(ContextReferenceProvider): + """Minimal concrete provider for testing.""" + + prefix = "test" + description = "test provider" + + async def autocomplete(self, query: str, *, limit: int = 10) -> list[ContextCompletionItem]: + return [ContextCompletionItem(text=f"{query}-result", meta="test")] + + async def expand(self, target: str) -> str | None: + return f"expanded: {target}" + + +class _NoneExpandProvider(ContextReferenceProvider): + """Provider whose expand() returns None (skip).""" + + prefix = "skip" + description = "skip provider" + + async def autocomplete(self, query: str, *, limit: int = 10) -> list[ContextCompletionItem]: + return [] + + async def expand(self, target: str) -> str | None: + return None + + +class _ErrorProvider(ContextReferenceProvider): + """Provider whose expand() raises.""" + + prefix = "boom" + description = "error provider" + + async def autocomplete(self, query: str, *, limit: int = 10) -> list[ContextCompletionItem]: + return [] + + async def expand(self, target: str) -> str | None: + raise RuntimeError("boom!") + + +@pytest.fixture(autouse=True) +def _clean_registry(): + """Clear plugin registry before and after each test.""" + _context_reference_providers.clear() + yield + _context_reference_providers.clear() + + +# -- registration tests ---------------------------------------------------- + +def test_register_valid_provider(): + p = _DummyProvider() + register_context_reference_provider(p) + assert "test" in get_context_reference_providers() + + +def test_register_rejects_builtin_prefix(): + for prefix in BUILTIN_PREFIXES: + p = _DummyProvider() + p.prefix = prefix + with pytest.raises(ValueError, match="reserved"): + register_context_reference_provider(p) + + +def test_register_rejects_duplicate_prefix(): + register_context_reference_provider(_DummyProvider()) + with pytest.raises(ValueError, match="already registered"): + register_context_reference_provider(_DummyProvider()) + + +def test_register_rejects_non_provider(): + with pytest.raises(TypeError, match="must be a ContextReferenceProvider"): + register_context_reference_provider("not a provider") + + +def test_register_rejects_empty_prefix(): + p = _DummyProvider() + p.prefix = "" + with pytest.raises(ValueError, match="non-empty"): + register_context_reference_provider(p) + + +# -- parse tests ----------------------------------------------------------- + +def test_parse_plugin_reference(): + register_context_reference_provider(_DummyProvider()) + refs = parse_context_references("check @test:ENG-123 and @file:README.md") + kinds = [r.kind for r in refs] + assert "test" in kinds + assert "file" in kinds + test_ref = [r for r in refs if r.kind == "test"][0] + assert test_ref.target == "ENG-123" + + +def test_parse_plugin_reference_ignored_when_not_registered(): + refs = parse_context_references("check @test:ENG-123") + assert [r.kind for r in refs] == [] + + +def test_plugin_pattern_regex(): + m = _PLUGIN_REFERENCE_PATTERN.search("@issue:ENG-123") + assert m is not None + assert m.group("kind") == "issue" + assert m.group("value") == "ENG-123" + + +# -- expand tests ---------------------------------------------------------- + +@pytest.mark.asyncio +async def test_expand_plugin_reference(tmp_path: Path): + from agent.context_references import preprocess_context_references_async + + register_context_reference_provider(_DummyProvider()) + result = await preprocess_context_references_async( + "check @test:ENG-123", + cwd=tmp_path, + context_length=10000, + ) + assert result.expanded + assert "expanded: ENG-123" in result.message + assert "test:ENG-123" not in result.message or "Attached Context" in result.message + + +@pytest.mark.asyncio +async def test_expand_plugin_returns_none(tmp_path: Path): + from agent.context_references import preprocess_context_references_async + + register_context_reference_provider(_NoneExpandProvider()) + result = await preprocess_context_references_async( + "check @skip:foo", + cwd=tmp_path, + context_length=10000, + ) + # expand() returned None, so the reference is parsed but no content injected + assert not any(r.kind == "skip" and "expanded" in (result.message or "") for r in result.references) + + +@pytest.mark.asyncio +async def test_expand_plugin_error(tmp_path: Path): + from agent.context_references import preprocess_context_references_async + + register_context_reference_provider(_ErrorProvider()) + result = await preprocess_context_references_async( + "check @boom:oops", + cwd=tmp_path, + context_length=10000, + ) + assert result.expanded + assert "plugin expansion error" in result.message + + +# -- autocomplete tests ---------------------------------------------------- + +@pytest.mark.asyncio +async def test_autocomplete(): + p = _DummyProvider() + register_context_reference_provider(p) + items = await p.autocomplete("foo", limit=5) + assert len(items) == 1 + assert items[0].text == "foo-result" + + +# -- ContextCompletionItem tests ------------------------------------------- + +def test_completion_item_defaults(): + item = ContextCompletionItem(text="@issue:1") + assert item.text == "@issue:1" + assert item.display == "@issue:1" + assert item.meta == "" + + +def test_completion_item_custom(): + item = ContextCompletionItem(text="1", display="ENG-1", meta="Bug") + assert item.display == "ENG-1" + assert item.meta == "Bug" diff --git a/tui_gateway/methods_complete.py b/tui_gateway/methods_complete.py index 5601a330e4..d25d674492 100644 --- a/tui_gateway/methods_complete.py +++ b/tui_gateway/methods_complete.py @@ -59,8 +59,62 @@ def _(rid, params: dict) -> dict: {"text": "@url:", "display": "@url:", "meta": "fetch url"}, {"text": "@git:", "display": "@git:", "meta": "git log"}, ] + # Append plugin-registered context reference prefixes + try: + from agent.context_references import get_context_reference_providers + + for _pfx, _prov in sorted(get_context_reference_providers().items()): + items.append( + { + "text": f"@{_pfx}:", + "display": f"@{_pfx}:", + "meta": _prov.description or f"plugin: {_pfx}", + } + ) + except Exception: + pass return _ok(rid, {"items": items}) + # Plugin context reference autocomplete: `@:` where the + # prefix belongs to a plugin-registered ContextReferenceProvider. + # Handled before the built-in file/folder branching so the elif/else + # chain below stays intact for built-in prefixes. + if is_context and ":" in query: + _pfx, _, _qval = query.partition(":") + if _pfx not in {"file", "folder", "url", "git", "diff", "staged"}: + try: + from agent.context_references import ( + get_context_reference_providers as _gcr, + ) + + _prov = _gcr().get(_pfx) + if _prov is not None: + import asyncio as _asyncio + + _coro = _prov.autocomplete(_qval, limit=20) + try: + _loop = _asyncio.get_running_loop() + except RuntimeError: + _loop = None + if _loop and _loop.is_running(): + import concurrent.futures as _cf + + with _cf.ThreadPoolExecutor(max_workers=1) as _pool: + _ac = _pool.submit(_asyncio.run, _coro).result() + else: + _ac = _asyncio.run(_coro) + items = [ + { + "text": f"@{_pfx}:{it.text}", + "display": it.display, + "meta": it.meta, + } + for it in _ac + ] + return _ok(rid, {"items": items}) + except Exception: + pass + # Accept both `@folder:path` and the bare `@folder` form so the user # sees directory listings as soon as they finish typing the keyword, # without first accepting the static `@folder:` hint.