diff --git a/src/openjarvis/core/config.py b/src/openjarvis/core/config.py index f6e9f4b7..677390b3 100644 --- a/src/openjarvis/core/config.py +++ b/src/openjarvis/core/config.py @@ -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: diff --git a/src/openjarvis/core/registry.py b/src/openjarvis/core/registry.py index f0f2becf..f0fd1735 100644 --- a/src/openjarvis/core/registry.py +++ b/src/openjarvis/core/registry.py @@ -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", diff --git a/src/openjarvis/sessions/compression.py b/src/openjarvis/sessions/compression.py new file mode 100644 index 00000000..c326c7b5 --- /dev/null +++ b/src/openjarvis/sessions/compression.py @@ -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 diff --git a/tests/conftest.py b/tests/conftest.py index d8d04e24..af05a657 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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() diff --git a/tests/sessions/test_compression.py b/tests/sessions/test_compression.py new file mode 100644 index 00000000..d4509bfb --- /dev/null +++ b/tests/sessions/test_compression.py @@ -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