- provider_base.py: ProviderBase (name/display_name/get_setup_schema) and CatalogProviderBase (default_model/list_models/is_available) replace the identical default-method bodies duplicated across 7 provider ABCs - provider_media.py: one save_b64/save_bytes/save_url/cache_dir implementation behind image_gen_provider and video_gen_provider - memory_manager.py: _each_provider fan-out helper replaces per-hook try/except loops; _signature_params/_has_var_kwargs unify signature probes - image_routing.py: _resolve_inference_value shared by base_url/api_key resolution; _dict_or_empty/_clean_str/_custom_provider_entries helpers - MemoryProvider/ContextEngine/TTS/browser/web/terminal-env ABC docstrings compacted to their invariants; method names and signatures unchanged
323 lines
13 KiB
Python
323 lines
13 KiB
Python
"""Abstract base class for pluggable context engines.
|
|
|
|
A context engine decides when and how conversation context is compacted near
|
|
the model's token limit, tracks token usage, and may expose tools. The
|
|
built-in ContextCompressor is the default; ``context.engine`` in config.yaml
|
|
selects a plugin engine (``plugins/context_engine/<name>/``). One engine is
|
|
active at a time.
|
|
|
|
Lifecycle: on_session_start() -> per API response update_from_response() ->
|
|
per turn should_compress() / compress() -> on_session_end() at real session
|
|
boundaries only (CLI exit, /reset, gateway expiry), never per-turn.
|
|
"""
|
|
|
|
from abc import ABC, abstractmethod
|
|
from typing import Any, Dict, List, Optional
|
|
|
|
from agent.redact import redact_sensitive_text
|
|
|
|
|
|
MEMORY_CONTEXT_MAX_CHARS = 6_000
|
|
_MEMORY_CONTEXT_HEAD_CHARS = 4_000
|
|
_MEMORY_CONTEXT_TAIL_CHARS = 1_500
|
|
_MEMORY_CONTEXT_TRUNCATION_MARKER = "\n...[memory provider context truncated]...\n"
|
|
|
|
|
|
def sanitize_memory_context(memory_context: str) -> str:
|
|
"""Prepare provider context for a context-engine/LLM egress boundary."""
|
|
sanitized = redact_sensitive_text(
|
|
memory_context.strip(),
|
|
force=True,
|
|
redact_url_credentials=True,
|
|
)
|
|
if len(sanitized) <= MEMORY_CONTEXT_MAX_CHARS:
|
|
return sanitized
|
|
return (
|
|
sanitized[:_MEMORY_CONTEXT_HEAD_CHARS]
|
|
+ _MEMORY_CONTEXT_TRUNCATION_MARKER
|
|
+ sanitized[-_MEMORY_CONTEXT_TAIL_CHARS:]
|
|
)
|
|
|
|
|
|
def automatic_compaction_status_message(
|
|
engine: Any,
|
|
*,
|
|
phase: str,
|
|
default_message: str,
|
|
**context: Any,
|
|
) -> str | None:
|
|
"""Host-visible status for an automatic compaction event; ``None`` = emit nothing.
|
|
|
|
Engines suppress via ``emit_automatic_compaction_status = False`` or
|
|
customize via ``get_automatic_compaction_status_message(...)``.
|
|
"""
|
|
if not getattr(engine, "emit_automatic_compaction_status", True):
|
|
return None
|
|
|
|
formatter = getattr(engine, "get_automatic_compaction_status_message", None)
|
|
if callable(formatter):
|
|
message = formatter(
|
|
phase=phase,
|
|
default_message=default_message,
|
|
**context,
|
|
)
|
|
else:
|
|
message = default_message
|
|
|
|
if message is None:
|
|
return None
|
|
message = str(message).strip()
|
|
return message or None
|
|
|
|
|
|
class ContextEngine(ABC):
|
|
"""Base class all context engines must implement."""
|
|
|
|
# -- Identity ----------------------------------------------------------
|
|
|
|
@property
|
|
@abstractmethod
|
|
def name(self) -> str:
|
|
"""Short identifier (e.g. 'compressor', 'lcm')."""
|
|
|
|
# -- Token state: engines MUST maintain these; run_agent.py reads them directly.
|
|
|
|
last_prompt_tokens: int = 0
|
|
last_completion_tokens: int = 0
|
|
last_total_tokens: int = 0
|
|
threshold_tokens: int = 0
|
|
context_length: int = 0
|
|
compression_count: int = 0
|
|
|
|
# -- Compaction parameters (read by run_agent.py for preflight). protect_first_n
|
|
# counts non-system head messages kept verbatim IN ADDITION to the always-
|
|
# protected system prompt (3 keeps the historical head shape).
|
|
|
|
threshold_percent: float = 0.75
|
|
protect_first_n: int = 3
|
|
protect_last_n: int = 6
|
|
|
|
# False keeps successful automatic compaction passes silent (routine
|
|
# background maintenance); warnings, errors and manual /compress still surface.
|
|
emit_automatic_compaction_status: bool = True
|
|
|
|
# -- Core interface ----------------------------------------------------
|
|
|
|
@abstractmethod
|
|
def update_from_response(self, usage: Dict[str, Any]) -> None:
|
|
"""Update tracked token usage after every LLM call.
|
|
|
|
``prompt_tokens``/``completion_tokens``/``total_tokens`` are always
|
|
present; the canonical buckets (``input_tokens``, ``output_tokens``,
|
|
``cache_read_tokens``, ``cache_write_tokens``, ``reasoning_tokens``)
|
|
are optional on older hosts.
|
|
"""
|
|
|
|
@abstractmethod
|
|
def should_compress(self, prompt_tokens: int = None) -> bool:
|
|
"""Return True if compaction should fire this turn."""
|
|
|
|
def should_compress_info(self, prompt_tokens: int = None) -> "tuple[bool, str | None]":
|
|
"""Return ``(should_compress, reason)``.
|
|
|
|
Engines with block reasons (summary-LLM cooldown, anti-thrashing guard)
|
|
override this so callers can warn the user instead of silently skipping
|
|
compression. The default keeps plugin engines from raising AttributeError.
|
|
"""
|
|
return self.should_compress(prompt_tokens), None
|
|
|
|
@abstractmethod
|
|
def compress(
|
|
self,
|
|
messages: List[Dict[str, Any]],
|
|
current_tokens: Optional[int] = None,
|
|
focus_topic: Optional[str] = None,
|
|
force: bool = False,
|
|
memory_context: str = "",
|
|
) -> List[Dict[str, Any]]:
|
|
"""Compact ``messages`` and return a valid OpenAI-format message list
|
|
that fits the context budget (summarize, build a DAG, anything).
|
|
|
|
``focus_topic`` comes from manual ``/compress <focus>`` (prioritise that
|
|
topic); ``force`` asks to bypass an engine-owned cooldown;
|
|
``memory_context`` is provider text to include in the handoff prompt.
|
|
Older engines may omit optional parameters — the host filters them by
|
|
signature.
|
|
"""
|
|
|
|
# -- Optional: proactive tool-result prune -----------------------------
|
|
|
|
def prune_tool_results_only(
|
|
self,
|
|
messages: List[Dict[str, Any]],
|
|
current_tokens: int | None = None,
|
|
) -> tuple[List[Dict[str, Any]], int]:
|
|
"""Deterministically trim old tool-result payloads without an LLM call.
|
|
|
|
Runs on a low, cost-oriented trigger independent of ``should_compress``
|
|
so large-window engines reclaim re-sent tool output long before full
|
|
compaction. Returns ``(messages, n_pruned)``; default is a no-op so
|
|
engines predating this hook never raise in the post-tool-call prune path.
|
|
"""
|
|
return messages, 0
|
|
|
|
# -- Optional: per-turn context selection (distinct from compression) --
|
|
|
|
def select_context(
|
|
self,
|
|
request_messages: List[Dict[str, Any]],
|
|
*,
|
|
conversation_messages: List[Dict[str, Any]] = None,
|
|
incoming_message: Dict[str, Any] = None,
|
|
budget_tokens: int = 0,
|
|
) -> List[Dict[str, Any]]:
|
|
"""Optionally *select* (replace) the context for THIS request, pre-generation.
|
|
|
|
Runs every provider request (so also on retries), independent of
|
|
``should_compress()``: ``compress()`` shrinks context that is too long,
|
|
``select_context()`` swaps in a different context (retrieval, topic
|
|
routing, branch switching) without abusing ``compress()`` as a per-turn
|
|
callback. Return ``None`` to leave the request unchanged.
|
|
|
|
The returned list is request-only — it MUST NOT be treated as persisted
|
|
transcript state; the session DB history is untouched. Unlike the
|
|
``pre_llm_call`` hook it may replace the list. The host runs it before
|
|
prompt cache-control and before every request sanitizer, so a malformed
|
|
replacement never reaches the provider and the default no-op keeps the
|
|
request byte-identical (prompt-cache stability preserved). An engine that
|
|
replaces the list changes its own cache prefix; breakpoints are
|
|
re-derived on the selected list.
|
|
|
|
``request_messages`` is the assembled request (system prompt + history +
|
|
ephemeral prefill); ``conversation_messages`` is the persisted history
|
|
for reference only (do not mutate); ``budget_tokens`` is the model's
|
|
context length or 0 if unknown.
|
|
"""
|
|
return None
|
|
|
|
def on_turn_complete(
|
|
self,
|
|
messages: List[Dict[str, Any]],
|
|
usage: Dict[str, Any] = None,
|
|
**kwargs: Any,
|
|
) -> None:
|
|
"""Observe a finished turn (complement of ``select_context()``) so the
|
|
engine can ingest/index/update routing state for the next request.
|
|
|
|
Fires from the normal finalization seam only; some abnormal early
|
|
returns (content-policy block, provider terminal failure) do not emit
|
|
it — treat it as best-effort, not guaranteed. ``messages`` is a
|
|
read-only shallow copy (return value ignored; never rely on transcript
|
|
mutation). ``usage`` has the ``update_from_response`` dict shape and is
|
|
``None`` when the turn never reached a provider response (interrupt).
|
|
``kwargs`` may include ``turn_id``, ``task_id``, ``api_call_count``,
|
|
``interrupted``, ``failed``, ``turn_exit_reason``.
|
|
"""
|
|
return None
|
|
|
|
# -- Optional: pre-flight check ----------------------------------------
|
|
|
|
def should_compress_preflight(self, messages: List[Dict[str, Any]]) -> bool:
|
|
"""Cheap rough check before the API call (no real token count yet); default skips."""
|
|
return False
|
|
|
|
def should_defer_preflight_to_real_usage(self, rough_tokens: int) -> bool:
|
|
"""True when preflight should trust recent real usage over the noisy rough
|
|
estimate (avoids re-compacting after a compressed request already fit)."""
|
|
return False
|
|
|
|
def get_automatic_compaction_status_message(
|
|
self,
|
|
*,
|
|
phase: str,
|
|
default_message: str,
|
|
**context: Any,
|
|
) -> str | None:
|
|
"""User-visible status for automatic compaction, or ``None`` to suppress it.
|
|
|
|
``phase`` is the host call site (``"preflight"`` / ``"compress"``);
|
|
``context`` carries best-effort ``approx_tokens`` / ``threshold_tokens``.
|
|
Warnings, errors and manual ``/compress`` are not governed by this hook.
|
|
"""
|
|
if not self.emit_automatic_compaction_status:
|
|
return None
|
|
return default_message
|
|
|
|
# -- Optional: manual /compress preflight ------------------------------
|
|
|
|
def has_content_to_compress(self, messages: List[Dict[str, Any]]) -> bool:
|
|
"""Preflight guard for gateway ``/compress``: False reports "nothing to
|
|
compress yet" without an LLM call (e.g. transcript entirely protected)."""
|
|
return True
|
|
|
|
# -- Optional: session lifecycle ---------------------------------------
|
|
|
|
def on_session_start(self, session_id: str, **kwargs) -> None:
|
|
"""Session begins: load persisted state. kwargs may include hermes_home, platform, model."""
|
|
|
|
def on_session_end(self, session_id: str, messages: List[Dict[str, Any]]) -> None:
|
|
"""Real session boundary (CLI exit, /reset, gateway expiry) — never per-turn."""
|
|
|
|
def on_session_reset(self) -> None:
|
|
"""/new or /reset: reset per-session state (default: counters and token tracking)."""
|
|
self.last_prompt_tokens = 0
|
|
self.last_completion_tokens = 0
|
|
self.last_total_tokens = 0
|
|
self.compression_count = 0
|
|
|
|
# -- Optional: tools ---------------------------------------------------
|
|
|
|
def get_tool_schemas(self) -> List[Dict[str, Any]]:
|
|
"""Tool schemas this engine exposes to the agent (default: none)."""
|
|
return []
|
|
|
|
def handle_tool_call(self, name: str, args: Dict[str, Any], **kwargs) -> str:
|
|
"""Handle a call to one of this engine's tools; must return a JSON string.
|
|
kwargs may include ``messages`` (live in-memory list)."""
|
|
import json
|
|
return json.dumps({"error": f"Unknown context engine tool: {name}"})
|
|
|
|
# -- Optional: status / display ----------------------------------------
|
|
|
|
def get_status(self) -> Dict[str, Any]:
|
|
"""Status dict with the standard fields run_agent.py expects."""
|
|
# Clamp the -1 "compression just ran, awaiting real usage" sentinel to 0
|
|
# so no reader sees a negative usage_percent on the transitional turn.
|
|
last_prompt = self.last_prompt_tokens if self.last_prompt_tokens > 0 else 0
|
|
return {
|
|
"last_prompt_tokens": last_prompt,
|
|
"threshold_tokens": self.threshold_tokens,
|
|
"context_length": self.context_length,
|
|
"usage_percent": (
|
|
min(100, last_prompt / self.context_length * 100)
|
|
if self.context_length else 0
|
|
),
|
|
"compression_count": self.compression_count,
|
|
}
|
|
|
|
# -- Optional: model switch support ------------------------------------
|
|
|
|
def update_model(
|
|
self,
|
|
model: str,
|
|
context_length: int,
|
|
base_url: str = "",
|
|
api_key: str = "",
|
|
provider: str = "",
|
|
api_mode: str = "",
|
|
) -> None:
|
|
"""Model switch / fallback: recompute threshold_tokens (override for more)."""
|
|
self.context_length = context_length
|
|
# Per-model threshold override (longest substring match), else the raw
|
|
# config percent. Snapshot that percent ONCE so repeated switches fall
|
|
# back to the configured value, not the previous model's override.
|
|
from agent.context_compressor import resolve_model_threshold
|
|
if not hasattr(self, "_config_threshold_percent"):
|
|
self._config_threshold_percent = self.threshold_percent
|
|
self._base_threshold_percent = resolve_model_threshold(
|
|
model, getattr(self, "model_thresholds", {}),
|
|
self._config_threshold_percent,
|
|
)
|
|
self.threshold_percent = self._base_threshold_percent
|
|
self.threshold_tokens = int(context_length * self.threshold_percent)
|