context_compressor.py

script

← Back to skill

Content hash: 9e5ba290f1fec9c54bce971c30ab6a16fd72f946886d0f18c95c85fb04143e4c
#!/usr/bin/env python3
"""LLM Context Window Management — token estimation, sliding-window eviction, summarization.

Demonstrates the strategies: truncation, sliding-window eviction, and summarization
with a durable transcript. Includes a mock tokenizer.
"""

from __future__ import annotations

import json
from dataclasses import dataclass, field
from typing import Any


# ── Mock tokenizer (replace with real tokenizer) ────────────────────────

def estimate_tokens(text: str) -> int:
    """Rough token estimation: ~4 chars per token."""
    return max(1, len(text) // 4)


def estimate_tokens_messages(messages: list[dict]) -> int:
    return sum(estimate_tokens(json.dumps(m, ensure_ascii=False)) for m in messages)


# ── Strategy 1: Truncation ──────────────────────────────────────────────

def truncate_messages(messages: list[dict], max_tokens: int) -> list[dict]:
    """Keep the most recent messages that fit within max_tokens."""
    kept: list[dict] = []
    total = 0
    for m in reversed(messages):
        t = estimate_tokens(json.dumps(m))
        if total + t > max_tokens:
            break
        kept.insert(0, m)
        total += t
    return kept


# ── Strategy 2: Sliding-window eviction ─────────────────────────────────

def sliding_window_evict(
    messages: list[dict],
    max_tokens: int,
    trigger_ratio: float = 0.5,
) -> list[dict]:
    """Drop oldest chunks once usage crosses trigger_ratio of the window."""
    if estimate_tokens_messages(messages) < max_tokens * trigger_ratio:
        return messages

    kept = list(messages)
    while estimate_tokens_messages(kept) > max_tokens * trigger_ratio and len(kept) > 1:
        # Drop the oldest message, but keep the system prompt (role == system)
        if kept[0].get("role") == "system" and len(kept) > 1:
            # Drop the second message (first non-system)
            kept.pop(1)
        else:
            kept.pop(0)
    return kept


# ── Strategy 3: Summarization / compaction ──────────────────────────────

@dataclass
class CompactionResult:
    summary: str
    kept_tail: list[dict]
    full_history_saved: bool


def compact_history(
    messages: list[dict],
    window_tokens: int,
    reserve_ratio: float = 0.7,
    summarize_fn=None,
) -> CompactionResult:
    """Summarize older messages and keep a recent tail of raw messages."""
    trigger = window_tokens * reserve_ratio
    if estimate_tokens_messages(messages) < trigger:
        return CompactionResult(summary="", kept_tail=messages, full_history_saved=False)

    # Split: summarize the old part, keep the recent tail
    tail: list[dict] = []
    tail_tokens = 0
    for m in reversed(messages):
        t = estimate_tokens(json.dumps(m))
        if tail_tokens + t > window_tokens * 0.3:
            break
        tail.insert(0, m)
        tail_tokens += t

    old_part = messages[: len(messages) - len(tail)]
    if summarize_fn is None:
        summarize_fn = lambda msgs: (
            f"[Summary of {len(msgs)} earlier messages] "
            "Key decisions: {extract from transcript}. Next steps: {pending items}."
        )

    summary = summarize_fn(old_part)

    # Prepend synthetic summary message
    new_messages = [{"role": "system", "content": f"Conversation summary: {summary}"}]
    new_messages.extend(tail)

    return CompactionResult(summary=summary, kept_tail=new_messages, full_history_saved=True)


# ── Demo ────────────────────────────────────────────────────────────────

if __name__ == "__main__":
    # Build a long conversation
    messages = [{"role": "system", "content": "You are a helpful assistant."}]
    for i in range(50):
        messages.append({"role": "user", "content": f"Question {i}: what is {i} times {i}?"})
        messages.append({"role": "assistant", "content": f"Answer {i}: it is {i * i}."})

    print(f"Original: {len(messages)} messages, ~{estimate_tokens_messages(messages)} tokens")

    # Truncation
    truncated = truncate_messages(messages, max_tokens=500)
    print(f"Truncation → {len(truncated)} messages, ~{estimate_tokens_messages(truncated)} tokens")

    # Sliding window
    sw = sliding_window_evict(messages, max_tokens=2000, trigger_ratio=0.5)
    print(f"Sliding window → {len(sw)} messages, ~{estimate_tokens_messages(sw)} tokens")

    # Compaction
    compacted = compact_history(messages, window_tokens=2000)
    print(f"Compaction → {len(compacted.kept_tail)} messages (summary prepended), "
          f"~{estimate_tokens_messages(compacted.kept_tail)} tokens")
    print(f"  Summary: {compacted.summary[:100]}...")
    print(f"  Full history saved to disk: {compacted.full_history_saved}")

    print("\nāœ“ Context window management demo complete.")