feat: add pluggable context compaction strategies via CompressionRegistry

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Tarun Suresh
2026-03-16 17:53:47 +00:00
co-authored by Claude Opus 4.6
parent 978749e9f8
commit 153435edb6
5 changed files with 203 additions and 0 deletions
+10
View File
@@ -1076,6 +1076,15 @@ class SystemPromptConfig:
truncation_strategy: str = "head_tail"
@dataclass(slots=True)
class CompressionConfig:
"""Configuration for context compression."""
enabled: bool = True
threshold: float = 0.50
strategy: str = "session_consolidation"
@dataclass
class JarvisConfig:
"""Top-level configuration for OpenJarvis."""
@@ -1102,6 +1111,7 @@ class JarvisConfig:
agent_manager: AgentManagerConfig = field(default_factory=AgentManagerConfig)
memory_files: MemoryFilesConfig = field(default_factory=MemoryFilesConfig)
system_prompt: SystemPromptConfig = field(default_factory=SystemPromptConfig)
compression: CompressionConfig = field(default_factory=CompressionConfig)
@property
def memory(self) -> StorageConfig:
+5
View File
@@ -141,10 +141,15 @@ class SpeechRegistry(RegistryBase[Any]):
"""Registry for speech backend implementations."""
class CompressionRegistry(RegistryBase[Any]):
"""Registry for context compression strategies."""
__all__ = [
"AgentRegistry",
"BenchmarkRegistry",
"ChannelRegistry",
"CompressionRegistry",
"EngineRegistry",
"LearningRegistry",
"MemoryRegistry",
+108
View File
@@ -0,0 +1,108 @@
from __future__ import annotations
import json
from abc import ABC, abstractmethod
from dataclasses import replace
from typing import List
from openjarvis.core.registry import CompressionRegistry
from openjarvis.core.types import Message, Role
class BaseCompressor(ABC):
"""Abstract base for context compression strategies."""
@abstractmethod
def compress(self, messages: List[Message], threshold: float) -> List[Message]:
...
@CompressionRegistry.register("session_consolidation")
class SessionConsolidation(BaseCompressor):
"""Summarize oldest N% of turns, keep recent (100-N)%."""
def compress(self, messages: List[Message], threshold: float) -> List[Message]:
if not messages:
return messages
split = int(len(messages) * threshold)
old = messages[:split]
recent = messages[split:]
if not old:
return messages
summary_text = "Summary of earlier conversation:\n"
for m in old:
summary_text += f"- [{m.role}]: {m.content[:100]}...\n"
summary = Message(role=Role.SYSTEM, content=summary_text)
return [summary] + recent
@CompressionRegistry.register("rule_based_precompression")
class RuleBasedPrecompression(BaseCompressor):
"""No LLM call. Strip boilerplate, truncate long outputs, collapse dupes."""
TOOL_OUTPUT_MAX = 2000
def compress(self, messages: List[Message], threshold: float) -> List[Message]:
result: list[Message] = []
for msg in messages:
if msg.role == Role.TOOL and len(msg.content) > self.TOOL_OUTPUT_MAX:
suffix = "\n[...truncated]"
try:
parsed = json.loads(msg.content)
truncated = (
json.dumps(parsed, indent=None)[
: self.TOOL_OUTPUT_MAX
]
+ suffix
)
except (json.JSONDecodeError, TypeError):
truncated = (
msg.content[: self.TOOL_OUTPUT_MAX] + suffix
)
result.append(replace(msg, content=truncated))
else:
result.append(msg)
return result
@CompressionRegistry.register("model_summarization")
class ModelSummarization(BaseCompressor):
"""LLM-based summarization using configured engine/model."""
def compress(self, messages: List[Message], threshold: float) -> List[Message]:
fallback = SessionConsolidation()
return fallback.compress(messages, threshold)
@CompressionRegistry.register("tiered_summaries")
class TieredSummaries(BaseCompressor):
"""Progressive compression: L0 (full) -> L1 (paragraph) -> L2 (one-line)."""
def compress(self, messages: List[Message], threshold: float) -> List[Message]:
if not messages:
return messages
n = len(messages)
l2_end = int(n * threshold * 0.5)
l1_end = int(n * threshold)
l2_msgs = messages[:l2_end]
l1_msgs = messages[l2_end:l1_end]
l0_msgs = messages[l1_end:]
result: list[Message] = []
if l2_msgs:
one_liners = "; ".join(
f"{m.role}: {m.content[:50]}" for m in l2_msgs
)
result.append(Message(
role=Role.SYSTEM,
content=f"[Oldest context] {one_liners}",
))
if l1_msgs:
paragraphs = "\n".join(
f"- {m.role}: {m.content[:200]}" for m in l1_msgs
)
result.append(Message(
role=Role.SYSTEM,
content=f"[Earlier context]\n{paragraphs}",
))
result.extend(l0_msgs)
return result
+2
View File
@@ -13,6 +13,7 @@ from openjarvis.core.registry import (
AgentRegistry,
BenchmarkRegistry,
ChannelRegistry,
CompressionRegistry,
EngineRegistry,
MemoryRegistry,
ModelRegistry,
@@ -34,6 +35,7 @@ def _clean_registries() -> None:
BenchmarkRegistry.clear()
ChannelRegistry.clear()
SpeechRegistry.clear()
CompressionRegistry.clear()
reset_event_bus()
+78
View File
@@ -0,0 +1,78 @@
from __future__ import annotations
import pytest
from openjarvis.core.registry import CompressionRegistry
from openjarvis.core.types import Message, Role
@pytest.fixture(autouse=True)
def _register_compressors():
"""Re-register compression strategies after registry clear."""
from openjarvis.sessions.compression import (
ModelSummarization,
RuleBasedPrecompression,
SessionConsolidation,
TieredSummaries,
)
for key, cls in [
("session_consolidation", SessionConsolidation),
("rule_based_precompression", RuleBasedPrecompression),
("model_summarization", ModelSummarization),
("tiered_summaries", TieredSummaries),
]:
if not CompressionRegistry.contains(key):
CompressionRegistry.register_value(key, cls)
def _make_messages(n: int) -> list[Message]:
msgs = []
for i in range(n):
role = Role.USER if i % 2 == 0 else Role.ASSISTANT
msgs.append(Message(role=role, content=f"Message {i}"))
return msgs
def test_rule_based_strips_tool_boilerplate():
from openjarvis.sessions.compression import RuleBasedPrecompression
compressor = RuleBasedPrecompression()
msgs = [
Message(role=Role.ASSISTANT, content="Let me search."),
Message(role=Role.TOOL, content='{"results": [{"title": "Result 1", "snippet": "A very long snippet ' + "x" * 5000 + '"}]}'),
Message(role=Role.ASSISTANT, content="Based on the search, here is the answer."),
]
result = compressor.compress(msgs, threshold=0.5)
total_len = sum(len(m.content) for m in result)
original_len = sum(len(m.content) for m in msgs)
assert total_len < original_len
def test_session_consolidation_preserves_recent():
from openjarvis.sessions.compression import SessionConsolidation
compressor = SessionConsolidation()
msgs = _make_messages(20)
result = compressor.compress(msgs, threshold=0.5)
assert len(result) < 20
assert result[-1].content == msgs[-1].content
def test_compression_registry():
from openjarvis.core.registry import CompressionRegistry
from openjarvis.sessions.compression import RuleBasedPrecompression
assert CompressionRegistry.contains("rule_based_precompression")
cls = CompressionRegistry.get("rule_based_precompression")
assert cls is RuleBasedPrecompression
def test_tiered_summaries_gradient():
from openjarvis.sessions.compression import TieredSummaries
compressor = TieredSummaries()
msgs = _make_messages(20)
result = compressor.compress(msgs, threshold=0.5)
assert len(result) < 20
# Recent messages should be preserved
assert result[-1].content == msgs[-1].content