mirror of
https://github.com/open-jarvis/OpenJarvis.git
synced 2026-07-31 03:12:16 +00:00
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:
co-authored by
Claude Opus 4.6
parent
978749e9f8
commit
153435edb6
@@ -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:
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
@@ -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()
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user