refactor(trajectory_compressor): fold single-use helpers, table-drive the report printer, flatten dataclass layout; -133 LOC

This commit is contained in:
Teknium
2026-09-02 21:23:40 -07:00
parent b9c2dd4041
commit 914aa20b06
+121 -254
View File
@@ -31,13 +31,10 @@ from rich.progress import Progress, SpinnerColumn, TextColumn, BarColumn, TaskPr
from rich.console import Console from rich.console import Console
from hermes_constants import OPENROUTER_BASE_URL, get_hermes_home from hermes_constants import OPENROUTER_BASE_URL, get_hermes_home
from agent.retry_utils import jittered_backoff from agent.retry_utils import jittered_backoff
# Load .env from HERMES_HOME first, then project root as a dev fallback.
from hermes_cli.env_loader import load_hermes_dotenv from hermes_cli.env_loader import load_hermes_dotenv
_hermes_home = get_hermes_home() # Load .env from HERMES_HOME first, then project root as a dev fallback.
_project_env = Path(__file__).parent / ".env" load_hermes_dotenv(hermes_home=get_hermes_home(), project_env=Path(__file__).parent / ".env")
load_hermes_dotenv(hermes_home=_hermes_home, project_env=_project_env)
def _response_finish_reason(response: Any) -> str: def _response_finish_reason(response: Any) -> str:
@@ -66,7 +63,6 @@ def _effective_temperature_for_model(model: str, requested_temperature: Optional
from agent.auxiliary_client import _fixed_temperature_for_model, OMIT_TEMPERATURE from agent.auxiliary_client import _fixed_temperature_for_model, OMIT_TEMPERATURE
except Exception: except Exception:
return requested_temperature return requested_temperature
fixed_temperature = _fixed_temperature_for_model(model, base_url) fixed_temperature = _fixed_temperature_for_model(model, base_url)
if fixed_temperature is OMIT_TEMPERATURE: if fixed_temperature is OMIT_TEMPERATURE:
return None # caller must omit temperature return None # caller must omit temperature
@@ -78,8 +74,7 @@ def _load_jsonl(path: Path, on_error: Optional[Callable[[int, json.JSONDecodeErr
entries = [] entries = []
with open(path, 'r', encoding='utf-8') as f: with open(path, 'r', encoding='utf-8') as f:
for line_num, line in enumerate(f, start): for line_num, line in enumerate(f, start):
line = line.strip() if not line.strip():
if not line:
continue continue
try: try:
entries.append((line_num, json.loads(line))) entries.append((line_num, json.loads(line)))
@@ -110,43 +105,30 @@ _YAML_SECTIONS: Dict[str, Tuple[str, ...]] = {
@dataclass @dataclass
class CompressionConfig: class CompressionConfig:
"""Configuration for trajectory compression.""" """Configuration for trajectory compression (tokenizer / targets / protected turns / summarizer / output / processing / metrics)."""
# Tokenizer
tokenizer_name: str = "moonshotai/Kimi-K2-Thinking" tokenizer_name: str = "moonshotai/Kimi-K2-Thinking"
trust_remote_code: bool = True trust_remote_code: bool = True
# Compression targets
target_max_tokens: int = 15250 target_max_tokens: int = 15250
summary_target_tokens: int = 750 summary_target_tokens: int = 750
# Protected turns
protect_first_system: bool = True protect_first_system: bool = True
protect_first_human: bool = True protect_first_human: bool = True
protect_first_gpt: bool = True protect_first_gpt: bool = True
protect_first_tool: bool = True protect_first_tool: bool = True
protect_last_n_turns: int = 4 protect_last_n_turns: int = 4
# Summarization (OpenRouter)
summarization_model: str = "google/gemini-3-flash-preview" summarization_model: str = "google/gemini-3-flash-preview"
base_url: str = OPENROUTER_BASE_URL base_url: str = OPENROUTER_BASE_URL
api_key_env: str = "OPENROUTER_API_KEY" api_key_env: str = "OPENROUTER_API_KEY"
temperature: float = 0.3 temperature: float = 0.3
max_retries: int = 3 max_retries: int = 3
retry_delay: int = 2 retry_delay: int = 2
# Output
add_summary_notice: bool = True add_summary_notice: bool = True
summary_notice_text: str = "\n\nSome of your previous tool responses may be summarized to preserve context." summary_notice_text: str = "\n\nSome of your previous tool responses may be summarized to preserve context."
output_suffix: str = "_compressed" output_suffix: str = "_compressed"
# Processing
num_workers: int = 4 num_workers: int = 4
max_concurrent_requests: int = 50 # Max concurrent API calls for summarization max_concurrent_requests: int = 50 # Max concurrent API calls for summarization
skip_under_target: bool = True skip_under_target: bool = True
save_over_limit: bool = True save_over_limit: bool = True
per_trajectory_timeout: int = 300 # Timeout per trajectory in seconds (default: 5 min) per_trajectory_timeout: int = 300 # seconds (default: 5 min)
# Metrics
metrics_enabled: bool = True metrics_enabled: bool = True
metrics_per_trajectory: bool = True metrics_per_trajectory: bool = True
metrics_output_file: str = "compression_metrics.json" metrics_output_file: str = "compression_metrics.json"
@@ -156,12 +138,9 @@ class CompressionConfig:
"""Load configuration from YAML file (missing keys keep the defaults).""" """Load configuration from YAML file (missing keys keep the defaults)."""
with open(yaml_path, 'r', encoding="utf-8") as f: with open(yaml_path, 'r', encoding="utf-8") as f:
data = yaml.safe_load(f) or {} data = yaml.safe_load(f) or {}
config = cls() config = cls()
for section, keys in _YAML_SECTIONS.items(): for section, keys in _YAML_SECTIONS.items():
if section not in data: for key in keys if section in data else ():
continue
for key in keys:
yaml_key, _, attr = key.partition(":") yaml_key, _, attr = key.partition(":")
attr = attr or yaml_key attr = attr or yaml_key
value = data[section].get(yaml_key, getattr(config, attr)) value = data[section].get(yaml_key, getattr(config, attr))
@@ -178,19 +157,15 @@ class TrajectoryMetrics:
compressed_tokens: int = 0 compressed_tokens: int = 0
tokens_saved: int = 0 tokens_saved: int = 0
compression_ratio: float = 1.0 compression_ratio: float = 1.0
original_turns: int = 0 original_turns: int = 0
compressed_turns: int = 0 compressed_turns: int = 0
turns_removed: int = 0 turns_removed: int = 0
turns_compressed_start_idx: int = -1 turns_compressed_start_idx: int = -1
turns_compressed_end_idx: int = -1 turns_compressed_end_idx: int = -1
turns_in_compressed_region: int = 0 turns_in_compressed_region: int = 0
was_compressed: bool = False was_compressed: bool = False
still_over_limit: bool = False still_over_limit: bool = False
skipped_under_target: bool = False skipped_under_target: bool = False
summarization_api_calls: int = 0 summarization_api_calls: int = 0
summarization_errors: int = 0 summarization_errors: int = 0
@@ -199,9 +174,8 @@ class TrajectoryMetrics:
d["compression_ratio"] = round(self.compression_ratio, 4) d["compression_ratio"] = round(self.compression_ratio, 4)
region = {"start_idx": d.pop("turns_compressed_start_idx"), "end_idx": d.pop("turns_compressed_end_idx"), region = {"start_idx": d.pop("turns_compressed_start_idx"), "end_idx": d.pop("turns_compressed_end_idx"),
"turns_count": d.pop("turns_in_compressed_region")} "turns_count": d.pop("turns_in_compressed_region")}
# Insert the region after turns_removed to keep the historical key order.
items = list(d.items()) items = list(d.items())
items.insert(7, ("compression_region", region)) items.insert(7, ("compression_region", region)) # after turns_removed: historical key order
return dict(items) return dict(items)
@@ -217,23 +191,17 @@ class AggregateMetrics:
trajectories_skipped_under_target: int = 0 trajectories_skipped_under_target: int = 0
trajectories_still_over_limit: int = 0 trajectories_still_over_limit: int = 0
trajectories_failed: int = 0 trajectories_failed: int = 0
total_tokens_before: int = 0 total_tokens_before: int = 0
total_tokens_after: int = 0 total_tokens_after: int = 0
total_tokens_saved: int = 0 total_tokens_saved: int = 0
total_turns_before: int = 0 total_turns_before: int = 0
total_turns_after: int = 0 total_turns_after: int = 0
total_turns_removed: int = 0 total_turns_removed: int = 0
total_summarization_calls: int = 0 total_summarization_calls: int = 0
total_summarization_errors: int = 0 total_summarization_errors: int = 0
# Distribution stats
compression_ratios: List[float] = field(default_factory=list) compression_ratios: List[float] = field(default_factory=list)
tokens_saved_list: List[int] = field(default_factory=list) tokens_saved_list: List[int] = field(default_factory=list)
turns_removed_list: List[int] = field(default_factory=list) turns_removed_list: List[int] = field(default_factory=list)
processing_start_time: str = "" processing_start_time: str = ""
processing_end_time: str = "" processing_end_time: str = ""
processing_duration_seconds: float = 0.0 processing_duration_seconds: float = 0.0
@@ -259,53 +227,27 @@ class AggregateMetrics:
def to_dict(self) -> Dict[str, Any]: def to_dict(self) -> Dict[str, Any]:
return { return {
"summary": { "summary": {"total_trajectories": self.total_trajectories, "trajectories_compressed": self.trajectories_compressed,
"total_trajectories": self.total_trajectories, "trajectories_skipped_under_target": self.trajectories_skipped_under_target,
"trajectories_compressed": self.trajectories_compressed, "trajectories_still_over_limit": self.trajectories_still_over_limit, "trajectories_failed": self.trajectories_failed,
"trajectories_skipped_under_target": self.trajectories_skipped_under_target, "compression_rate": round(self.trajectories_compressed / max(self.total_trajectories, 1), 4)},
"trajectories_still_over_limit": self.trajectories_still_over_limit, "tokens": {"total_before": self.total_tokens_before, "total_after": self.total_tokens_after, "total_saved": self.total_tokens_saved,
"trajectories_failed": self.trajectories_failed, "overall_compression_ratio": round(self.total_tokens_after / max(self.total_tokens_before, 1), 4)},
"compression_rate": round(self.trajectories_compressed / max(self.total_trajectories, 1), 4), "turns": {"total_before": self.total_turns_before, "total_after": self.total_turns_after, "total_removed": self.total_turns_removed},
}, "averages": {"avg_compression_ratio": round(_mean(self.compression_ratios, 1.0), 4),
"tokens": { "avg_tokens_saved_per_compressed": round(_mean(self.tokens_saved_list, 0), 1),
"total_before": self.total_tokens_before, "avg_turns_removed_per_compressed": round(_mean(self.turns_removed_list, 0), 2)},
"total_after": self.total_tokens_after, "summarization": {"total_api_calls": self.total_summarization_calls, "total_errors": self.total_summarization_errors,
"total_saved": self.total_tokens_saved, "success_rate": round(1 - (self.total_summarization_errors / max(self.total_summarization_calls, 1)), 4)},
"overall_compression_ratio": round(self.total_tokens_after / max(self.total_tokens_before, 1), 4), "processing": {"start_time": self.processing_start_time, "end_time": self.processing_end_time,
}, "duration_seconds": round(self.processing_duration_seconds, 2)},
"turns": {
"total_before": self.total_turns_before,
"total_after": self.total_turns_after,
"total_removed": self.total_turns_removed,
},
"averages": {
"avg_compression_ratio": round(_mean(self.compression_ratios, 1.0), 4),
"avg_tokens_saved_per_compressed": round(_mean(self.tokens_saved_list, 0), 1),
"avg_turns_removed_per_compressed": round(_mean(self.turns_removed_list, 0), 2),
},
"summarization": {
"total_api_calls": self.total_summarization_calls,
"total_errors": self.total_summarization_errors,
"success_rate": round(1 - (self.total_summarization_errors / max(self.total_summarization_calls, 1)), 4),
},
"processing": {
"start_time": self.processing_start_time,
"end_time": self.processing_end_time,
"duration_seconds": round(self.processing_duration_seconds, 2),
},
} }
# Ordered (hostname, provider) table for _detect_provider (codex is matched separately). # Ordered (hostname, provider) table for _detect_provider (codex is matched separately).
_PROVIDER_HOSTS: Tuple[Tuple[str, str], ...] = ( _PROVIDER_HOSTS: Tuple[Tuple[str, str], ...] = (
("openrouter.ai", "openrouter"), ("openrouter.ai", "openrouter"), ("nousresearch.com", "nous"), ("z.ai", "zai"), ("moonshot.ai", "kimi-coding"),
("nousresearch.com", "nous"), ("moonshot.cn", "kimi-coding"), ("api.kimi.com", "kimi-coding"), ("arcee.ai", "arcee"), ("minimaxi.com", "minimax-cn"),
("z.ai", "zai"),
("moonshot.ai", "kimi-coding"),
("moonshot.cn", "kimi-coding"),
("api.kimi.com", "kimi-coding"),
("arcee.ai", "arcee"),
("minimaxi.com", "minimax-cn"),
("minimax.io", "minimax"), ("minimax.io", "minimax"),
) )
@@ -327,7 +269,9 @@ class _RunProgress:
in_flight: int = 0 in_flight: int = 0
timeouts: int = 0 timeouts: int = 0
def advance(self, update_status: bool = True) -> None: def finish(self, update_status: bool = True) -> None:
"""Retire one in-flight entry and advance the bar (caller holds ``lock``)."""
self.in_flight -= 1
self.progress.advance(self.main_task) self.progress.advance(self.main_task)
if update_status: if update_status:
self.progress.update(self.status_task, description=_STATUS_FMT.format( self.progress.update(self.status_task, description=_STATUS_FMT.format(
@@ -361,9 +305,9 @@ class TrajectoryCompressor:
def _init_summarizer(self): def _init_summarizer(self):
"""Route summarization through call_llm for known providers, else a raw client.""" """Route summarization through call_llm for known providers, else a raw client."""
provider = self._detect_provider() provider = self._detect_provider()
self._use_call_llm = bool(provider)
if provider: if provider:
self._llm_provider = provider self._llm_provider = provider
self._use_call_llm = True
from agent.auxiliary_client import resolve_provider_client from agent.auxiliary_client import resolve_provider_client
client, _ = resolve_provider_client(provider, model=self.config.summarization_model) client, _ = resolve_provider_client(provider, model=self.config.summarization_model)
if client is None: if client is None:
@@ -371,19 +315,17 @@ class TrajectoryCompressor:
self.client = self.async_client = None # Not used directly self.client = self.async_client = None # Not used directly
else: else:
# Custom endpoint — use config's raw base_url + api_key_env # Custom endpoint — use config's raw base_url + api_key_env
self._use_call_llm = False
api_key = os.getenv(self.config.api_key_env) api_key = os.getenv(self.config.api_key_env)
if not api_key: if not api_key:
raise RuntimeError(f"Missing API key. Set {self.config.api_key_env} environment variable.") raise RuntimeError(f"Missing API key. Set {self.config.api_key_env} environment variable.")
from openai import OpenAI from openai import OpenAI
from agent.auxiliary_client import _to_openai_base_url from agent.auxiliary_client import _to_openai_base_url
self.client = OpenAI(api_key=api_key, base_url=_to_openai_base_url(self.config.base_url)) self.client = OpenAI(api_key=api_key, base_url=_to_openai_base_url(self.config.base_url))
# AsyncOpenAI is created lazily in _get_async_client() so it binds to # AsyncOpenAI is created lazily in _get_async_client() so it binds to the current event
# the current event loop — each process_directory() call runs its own # loop — each process_directory() runs its own asyncio.run(); a shared client would hit
# asyncio.run(), and a shared client would hit "Event loop is closed". # "Event loop is closed".
self.async_client = None self.async_client = None
self._async_client_api_key = api_key self._async_client_api_key = api_key
print(f"✅ Initialized summarizer client: {self.config.summarization_model}") print(f"✅ Initialized summarizer client: {self.config.summarization_model}")
print(f" Max concurrent requests: {self.config.max_concurrent_requests}") print(f" Max concurrent requests: {self.config.max_concurrent_requests}")
@@ -399,10 +341,7 @@ class TrajectoryCompressor:
url = self.config.base_url or "" url = self.config.base_url or ""
if base_url_hostname(url) == "chatgpt.com" and "/backend-api/codex" in url.lower(): if base_url_hostname(url) == "chatgpt.com" and "/backend-api/codex" in url.lower():
return "codex" return "codex"
for host, provider in _PROVIDER_HOSTS: return next((provider for host, provider in _PROVIDER_HOSTS if base_url_host_matches(url, host)), "")
if base_url_host_matches(url, host):
return provider
return ""
def count_tokens(self, text: str) -> int: def count_tokens(self, text: str) -> int:
"""Token count via the configured tokenizer; falls back to len//4.""" """Token count via the configured tokenizer; falls back to len//4."""
@@ -422,51 +361,36 @@ class TrajectoryCompressor:
def _find_protected_indices(self, trajectory: List[Dict[str, str]]) -> Tuple[set, int, int]: def _find_protected_indices(self, trajectory: List[Dict[str, str]]) -> Tuple[set, int, int]:
"""Return ``(protected_set, compressible_start, compressible_end)``.""" """Return ``(protected_set, compressible_start, compressible_end)``."""
n = len(trajectory) n = len(trajectory)
protected = set()
first_seen: Dict[str, int] = {} first_seen: Dict[str, int] = {}
for i, turn in enumerate(trajectory): for i, turn in enumerate(trajectory):
first_seen.setdefault(turn.get("from", ""), i) first_seen.setdefault(turn.get("from", ""), i)
protected = {first_seen[role] for role in ("system", "human", "gpt", "tool")
for role in ("system", "human", "gpt", "tool"): if getattr(self.config, f"protect_first_{role}") and role in first_seen}
if getattr(self.config, f"protect_first_{role}") and role in first_seen:
protected.add(first_seen[role])
protected.update(range(max(0, n - self.config.protect_last_n_turns), n)) protected.update(range(max(0, n - self.config.protect_last_n_turns), n))
# Compressible region: after the last protected head turn, before the first tail turn. # Compressible region: after the last protected head turn, before the first tail turn.
head_protected = [i for i in protected if i < n // 2] head_protected = [i for i in protected if i < n // 2]
tail_protected = [i for i in protected if i >= n // 2] tail_protected = [i for i in protected if i >= n // 2]
return protected, max(head_protected) + 1 if head_protected else 0, min(tail_protected) if tail_protected else n
compressible_start = max(head_protected) + 1 if head_protected else 0
compressible_end = min(tail_protected) if tail_protected else n
return protected, compressible_start, compressible_end
@staticmethod @staticmethod
def _is_boundary_clean(trajectory: List[Dict[str, str]], idx: int) -> bool: def _snap_boundary(trajectory: List[Dict[str, str]], idx: int, min_idx: int, max_idx: int) -> int:
"""True if a boundary at ``idx`` does not split a gpt <tool_call>/tool <tool_response> pair. """Move a boundary onto the nearest turn boundary within ``[min_idx, max_idx]`` that does not
split a gpt <tool_call>/tool <tool_response> pair.
A ``tool`` turn always directly follows the ``gpt`` turn it answers, so a A ``tool`` turn always directly follows the ``gpt`` turn it answers, so a boundary landing *on*
boundary landing *on* a tool turn cuts the pair; only the end of the a tool turn cuts the pair; only the end of the trajectory or a non-``tool`` turn is clean.
trajectory or a non-``tool`` turn is clean. Forward is preferred (folds an orphaned ``tool`` turn into the region that holds its ``gpt``
""" turn); backward only when nothing clean lies ahead.
return idx >= len(trajectory) or trajectory[idx].get("from") != "tool"
@classmethod
def _snap_boundary(cls, trajectory: List[Dict[str, str]], idx: int, min_idx: int, max_idx: int) -> int:
"""Move a boundary onto the nearest clean turn boundary within ``[min_idx, max_idx]``.
Forward is preferred (folds an orphaned ``tool`` turn into the region that
holds its ``gpt`` turn); backward only when nothing clean lies ahead.
""" """
def clean(i: int) -> bool:
return i >= len(trajectory) or trajectory[i].get("from") != "tool"
forward = idx forward = idx
while forward < max_idx and not cls._is_boundary_clean(trajectory, forward): while forward < max_idx and not clean(forward):
forward += 1 forward += 1
if cls._is_boundary_clean(trajectory, forward): if clean(forward):
return forward return forward
backward = idx backward = idx
while backward > min_idx and not cls._is_boundary_clean(trajectory, backward): while backward > min_idx and not clean(backward):
backward -= 1 backward -= 1
return backward return backward
@@ -475,28 +399,12 @@ class TrajectoryCompressor:
parts = [] parts = []
for i in range(start, end): for i in range(start, end):
turn = trajectory[i] turn = trajectory[i]
role = turn.get("from", "unknown")
value = turn.get("value", "") value = turn.get("value", "")
if len(value) > 3000: if len(value) > 3000:
value = value[:1500] + "\n...[truncated]...\n" + value[-500:] value = value[:1500] + "\n...[truncated]...\n" + value[-500:]
parts.append(f"[Turn {i} - {role.upper()}]:\n{value}") parts.append(f"[Turn {i} - {turn.get('from', 'unknown').upper()}]:\n{value}")
return "\n\n".join(parts) return "\n\n".join(parts)
@staticmethod
def _coerce_summary_content(content: Any) -> str:
"""Normalize summary-model output to a safe string."""
if not isinstance(content, str):
content = str(content) if content else ""
return content.strip()
@staticmethod
def _ensure_summary_prefix(summary: str) -> str:
"""Normalize summary text to include the expected prefix exactly once."""
text = (summary or "").strip()
if text.startswith("[CONTEXT SUMMARY]:"):
return text
return "[CONTEXT SUMMARY]:" if not text else f"[CONTEXT SUMMARY]: {text}"
def _summary_prompt(self, content: str) -> str: def _summary_prompt(self, content: str) -> str:
return f"""Summarize the following agent conversation turns concisely. This summary will replace these turns in the conversation history. return f"""Summarize the following agent conversation turns concisely. This summary will replace these turns in the conversation history.
@@ -526,13 +434,16 @@ Write only the summary, starting with "[CONTEXT SUMMARY]:" prefix."""
return temperature, kwargs return temperature, kwargs
def _finish_summary(self, response: Any) -> str: def _finish_summary(self, response: Any) -> str:
"""Extract the summary text; a ``length`` stop is a failure (partial summary).""" """Extract the summary text with the ``[CONTEXT SUMMARY]:`` prefix exactly once; a ``length`` stop is a failure."""
if _response_finish_reason(response) == "length": if _response_finish_reason(response) == "length":
# Storing a truncated summary silently corrupts the trajectory's # Storing a truncated summary silently corrupts the trajectory's memory, so raise and
# memory, so raise and let the retry/backoff loop handle it. # let the retry/backoff loop handle it.
raise RuntimeError("trajectory summarization hit the output token cap (finish_reason=length); summary is incomplete") raise RuntimeError("trajectory summarization hit the output token cap (finish_reason=length); summary is incomplete")
summary = self._coerce_summary_content(response.choices[0].message.content) content = response.choices[0].message.content
return self._ensure_summary_prefix(summary) text = (content if isinstance(content, str) else str(content) if content else "").strip()
if text.startswith("[CONTEXT SUMMARY]:"):
return text
return "[CONTEXT SUMMARY]:" if not text else f"[CONTEXT SUMMARY]: {text}"
def _summary_attempt_failed(self, metrics: TrajectoryMetrics, attempt: int, exc: Exception) -> Optional[float]: def _summary_attempt_failed(self, metrics: TrajectoryMetrics, attempt: int, exc: Exception) -> Optional[float]:
"""Record a failed attempt; return the backoff delay, or None on the last attempt.""" """Record a failed attempt; return the backoff delay, or None on the last attempt."""
@@ -580,9 +491,7 @@ Write only the summary, starting with "[CONTEXT SUMMARY]:" prefix."""
return _SUMMARY_FALLBACK return _SUMMARY_FALLBACK
await asyncio.sleep(delay) await asyncio.sleep(delay)
def _plan_compression( def _plan_compression(self, trajectory: List[Dict[str, str]], metrics: TrajectoryMetrics) -> Optional[Tuple[int, int]]:
self, trajectory: List[Dict[str, str]], metrics: TrajectoryMetrics
) -> Optional[Tuple[int, int]]:
"""Choose the ``[start, until)`` region to summarize, or None if nothing can be. """Choose the ``[start, until)`` region to summarize, or None if nothing can be.
Fills the pre-compression metrics either way. Accumulates turns from the Fills the pre-compression metrics either way. Accumulates turns from the
@@ -594,18 +503,15 @@ Write only the summary, starting with "[CONTEXT SUMMARY]:" prefix."""
total_tokens = sum(turn_tokens) total_tokens = sum(turn_tokens)
metrics.original_turns = metrics.compressed_turns = len(trajectory) metrics.original_turns = metrics.compressed_turns = len(trajectory)
metrics.original_tokens = metrics.compressed_tokens = total_tokens metrics.original_tokens = metrics.compressed_tokens = total_tokens
if total_tokens <= cfg.target_max_tokens: if total_tokens <= cfg.target_max_tokens:
metrics.skipped_under_target = True metrics.skipped_under_target = True
return None return None
metrics.still_over_limit = True metrics.still_over_limit = True
_, start, end = self._find_protected_indices(trajectory) _, start, end = self._find_protected_indices(trajectory)
# Never *start* on an orphaned <tool_response> whose <tool_call> is in the protected head. # Never *start* on an orphaned <tool_response> whose <tool_call> is in the protected head.
start = self._snap_boundary(trajectory, start, start, end) start = self._snap_boundary(trajectory, start, start, end)
if start >= end: if start >= end:
return None return None
# Replacing N turns with one summary saves sum(N) - summary_target_tokens. # Replacing N turns with one summary saves sum(N) - summary_target_tokens.
target_tokens_to_compress = total_tokens - cfg.target_max_tokens + cfg.summary_target_tokens target_tokens_to_compress = total_tokens - cfg.target_max_tokens + cfg.summary_target_tokens
accumulated = 0 accumulated = 0
@@ -617,15 +523,12 @@ Write only the summary, starting with "[CONTEXT SUMMARY]:" prefix."""
break break
if accumulated < target_tokens_to_compress and until < end: if accumulated < target_tokens_to_compress and until < end:
until = end until = end
# The remainder is kept verbatim, so a tail boundary on a tool turn would orphan a marker. # The remainder is kept verbatim, so a tail boundary on a tool turn would orphan a marker.
until = self._snap_boundary(trajectory, until, start, end) until = self._snap_boundary(trajectory, until, start, end)
# A region no larger than the summary replacing it cannot shrink the trajectory. # A region no larger than the summary replacing it cannot shrink the trajectory.
if until <= start or sum(turn_tokens[start:until]) <= cfg.summary_target_tokens: if until <= start or sum(turn_tokens[start:until]) <= cfg.summary_target_tokens:
return None return None
metrics.turns_compressed_start_idx, metrics.turns_compressed_end_idx = start, until
metrics.turns_compressed_start_idx = start
metrics.turns_compressed_end_idx = until
metrics.turns_in_compressed_region = until - start metrics.turns_in_compressed_region = until - start
return start, until return start, until
@@ -640,7 +543,6 @@ Write only the summary, starting with "[CONTEXT SUMMARY]:" prefix."""
compressed.append(turn) compressed.append(turn)
compressed.append({"from": "human", "value": summary}) compressed.append({"from": "human", "value": summary})
compressed.extend(turn.copy() for turn in trajectory[until:]) compressed.extend(turn.copy() for turn in trajectory[until:])
metrics.compressed_turns = len(compressed) metrics.compressed_turns = len(compressed)
metrics.compressed_tokens = self.count_trajectory_tokens(compressed) metrics.compressed_tokens = self.count_trajectory_tokens(compressed)
metrics.turns_removed = metrics.original_turns - metrics.compressed_turns metrics.turns_removed = metrics.original_turns - metrics.compressed_turns
@@ -656,9 +558,8 @@ Write only the summary, starting with "[CONTEXT SUMMARY]:" prefix."""
region = self._plan_compression(trajectory, metrics) region = self._plan_compression(trajectory, metrics)
if region is None: if region is None:
return trajectory, metrics return trajectory, metrics
start, until = region summary = self._generate_summary(self._extract_turn_content_for_summary(trajectory, *region), metrics)
summary = self._generate_summary(self._extract_turn_content_for_summary(trajectory, start, until), metrics) return self._assemble_compressed(trajectory, *region, summary, metrics), metrics
return self._assemble_compressed(trajectory, start, until, summary, metrics), metrics
async def compress_trajectory_async(self, trajectory: List[Dict[str, str]]) -> Tuple[List[Dict[str, str]], TrajectoryMetrics]: async def compress_trajectory_async(self, trajectory: List[Dict[str, str]]) -> Tuple[List[Dict[str, str]], TrajectoryMetrics]:
"""Async twin of ``compress_trajectory``.""" """Async twin of ``compress_trajectory``."""
@@ -666,17 +567,15 @@ Write only the summary, starting with "[CONTEXT SUMMARY]:" prefix."""
region = self._plan_compression(trajectory, metrics) region = self._plan_compression(trajectory, metrics)
if region is None: if region is None:
return trajectory, metrics return trajectory, metrics
start, until = region summary = await self._generate_summary_async(self._extract_turn_content_for_summary(trajectory, *region), metrics)
summary = await self._generate_summary_async(self._extract_turn_content_for_summary(trajectory, start, until), metrics) return self._assemble_compressed(trajectory, *region, summary, metrics), metrics
return self._assemble_compressed(trajectory, start, until, summary, metrics), metrics
async def process_entry_async(self, entry: Dict[str, Any]) -> Tuple[Dict[str, Any], TrajectoryMetrics]: async def process_entry_async(self, entry: Dict[str, Any]) -> Tuple[Dict[str, Any], TrajectoryMetrics]:
"""Compress one JSONL entry's ``conversations``; attach metrics when compressed.""" """Compress one JSONL entry's ``conversations``; attach metrics when compressed."""
if "conversations" not in entry: if "conversations" not in entry:
return entry, TrajectoryMetrics() return entry, TrajectoryMetrics()
compressed_trajectory, metrics = await self.compress_trajectory_async(entry["conversations"]) compressed_trajectory, metrics = await self.compress_trajectory_async(entry["conversations"])
result = entry.copy() result = dict(entry, conversations=compressed_trajectory)
result["conversations"] = compressed_trajectory
if self.config.metrics_per_trajectory and metrics.was_compressed: if self.config.metrics_per_trajectory and metrics.was_compressed:
result["compression_metrics"] = metrics.to_dict() result["compression_metrics"] = metrics.to_dict()
return result, metrics return result, metrics
@@ -697,32 +596,27 @@ Write only the summary, starting with "[CONTEXT SUMMARY]:" prefix."""
if metrics.was_compressed: if metrics.was_compressed:
run.compressed += 1 run.compressed += 1
run.api_calls += metrics.summarization_api_calls run.api_calls += metrics.summarization_api_calls
if metrics.skipped_under_target: run.skipped += bool(metrics.skipped_under_target)
run.skipped += 1 run.finish()
run.in_flight -= 1
run.advance()
return processed_entry, metrics return processed_entry, metrics
except asyncio.TimeoutError: except asyncio.TimeoutError:
self.logger.warning("Timeout processing entry from %s:%s (>%ss)", file_path, entry_idx, self.config.per_trajectory_timeout) self.logger.warning("Timeout processing entry from %s:%s (>%ss)", file_path, entry_idx, self.config.per_trajectory_timeout)
async with run.lock: async with run.lock:
self.aggregate_metrics.trajectories_failed += 1 self.aggregate_metrics.trajectories_failed += 1
run.timeouts += 1 run.timeouts += 1
run.in_flight -= 1 run.finish()
run.advance()
return None return None
except Exception as e: except Exception as e:
self.logger.error("Error processing entry from %s:%s: %s", file_path, entry_idx, e) self.logger.error("Error processing entry from %s:%s: %s", file_path, entry_idx, e)
async with run.lock: async with run.lock:
self.aggregate_metrics.trajectories_failed += 1 self.aggregate_metrics.trajectories_failed += 1
run.in_flight -= 1 run.finish(update_status=False)
run.advance(update_status=False)
return entry, TrajectoryMetrics() # keep the original on error return entry, TrajectoryMetrics() # keep the original on error
async def _process_directory_async(self, input_dir: Path, output_dir: Path): async def _process_directory_async(self, input_dir: Path, output_dir: Path):
console = Console() console = Console()
self.aggregate_metrics.processing_start_time = datetime.now().isoformat() self.aggregate_metrics.processing_start_time = datetime.now().isoformat()
start_time = time.time() start_time = time.time()
jsonl_files = sorted(input_dir.glob("*.jsonl")) jsonl_files = sorted(input_dir.glob("*.jsonl"))
if not jsonl_files: if not jsonl_files:
self.logger.warning("No JSONL files found in %s", input_dir) self.logger.warning("No JSONL files found in %s", input_dir)
@@ -772,7 +666,6 @@ Write only the summary, starting with "[CONTEXT SUMMARY]:" prefix."""
self.aggregate_metrics.processing_end_time = datetime.now().isoformat() self.aggregate_metrics.processing_end_time = datetime.now().isoformat()
self.aggregate_metrics.processing_duration_seconds = time.time() - start_time self.aggregate_metrics.processing_duration_seconds = time.time() - start_time
self._print_summary() self._print_summary()
if self.config.metrics_enabled: if self.config.metrics_enabled:
metrics_path = output_dir / self.config.metrics_output_file metrics_path = output_dir / self.config.metrics_output_file
with open(metrics_path, 'w', encoding="utf-8") as f: with open(metrics_path, 'w', encoding="utf-8") as f:
@@ -782,63 +675,57 @@ Write only the summary, starting with "[CONTEXT SUMMARY]:" prefix."""
def _print_summary(self): def _print_summary(self):
"""Print comprehensive compression summary statistics.""" """Print comprehensive compression summary statistics."""
m = self.aggregate_metrics.to_dict() m = self.aggregate_metrics.to_dict()
total = m['summary']['total_trajectories'] s, t, u, a, z, p = m['summary'], m['tokens'], m['turns'], m['averages'], m['summarization'], m['processing']
compressed = m['summary']['trajectories_compressed'] total, compressed = s['total_trajectories'], s['trajectories_compressed']
skipped = m['summary']['trajectories_skipped_under_target']
over_limit = m['summary']['trajectories_still_over_limit']
tokens_before = m['tokens']['total_before']
tokens_saved = m['tokens']['total_saved']
pct = lambda n: (n / max(total, 1)) * 100 # noqa: E731 pct = lambda n: (n / max(total, 1)) * 100 # noqa: E731
duration = p['duration_seconds']
time_str = f"{duration/60:.1f} minutes" if duration > 60 else f"{duration:.1f} seconds"
def section(title: str, pad: int) -> None: sections = [
print(f"╠{'═'*70}╣") ("📁 TRAJECTORIES", 54, [
print(f"║{'':2}{title}{' '*pad}║") f"║{'':4}Total Processed: {total:>10,}{' '*32}║",
print(f"║{'─'*70}║") f"║{'':4}├─ Compressed: {compressed:>10,} ({pct(compressed):>5.1f}%){' '*18}║",
f"║{'':4}├─ Skipped (under limit):{s['trajectories_skipped_under_target']:>9,} ({pct(s['trajectories_skipped_under_target']):>5.1f}%){' '*18}║",
f"║{'':4}├─ Still over limit: {s['trajectories_still_over_limit']:>10,} ({pct(s['trajectories_still_over_limit']):>5.1f}%){' '*18}║",
f"║{'':4}└─ Failed: {s['trajectories_failed']:>10,}{' '*32}║",
]),
("🔢 TOKENS", 60, [
f"║{'':4}Before Compression: {t['total_before']:>15,} tokens{' '*21}║",
f"║{'':4}After Compression: {t['total_after']:>15,} tokens{' '*21}║",
f"║{'':4}Total Saved: {t['total_saved']:>15,} tokens{' '*21}║",
f"║{'':4}Overall Compression: {t['overall_compression_ratio']:>14.1%}{' '*28}║",
] + ([f"║{'':4}Space Savings: {(t['total_saved'] / t['total_before']) * 100:>14.1f}%{' '*28}║"] if t['total_before'] > 0 else [])),
("💬 CONVERSATION TURNS", 48, [
f"║{'':4}Before Compression: {u['total_before']:>15,} turns{' '*22}║",
f"║{'':4}After Compression: {u['total_after']:>15,} turns{' '*22}║",
f"║{'':4}Total Removed: {u['total_removed']:>15,} turns{' '*22}║",
]),
("📈 AVERAGES (Compressed Trajectories Only)", 27, [
f"║{'':4}Avg Compression Ratio: {a['avg_compression_ratio']:>14.1%}{' '*28}║",
f"║{'':4}Avg Tokens Saved: {a['avg_tokens_saved_per_compressed']:>14,.0f}{' '*28}║",
f"║{'':4}Avg Turns Removed: {a['avg_turns_removed_per_compressed']:>14.1f}{' '*28}║",
] if compressed > 0 else [f"║{'':4}No trajectories were compressed{' '*38}║"]),
("🤖 SUMMARIZATION API", 49, [
f"║{'':4}API Calls Made: {z['total_api_calls']:>15,}{' '*27}║",
f"║{'':4}Errors: {z['total_errors']:>15,}{' '*27}║",
f"║{'':4}Success Rate: {z['success_rate']:>14.1%}{' '*28}║",
]),
("⏱️ PROCESSING TIME", 51, [
f"║{'':4}Duration: {time_str:>20}{' '*22}║",
f"║{'':4}Throughput: {total / max(duration, 0.001):>15.1f} traj/sec{' '*18}║",
f"║{'':4}Started: {p['start_time'][:19]:>20}{' '*22}║",
f"║{'':4}Finished: {p['end_time'][:19]:>20}{' '*22}║",
]),
]
print("\n") print("\n")
print(f"╔{'═'*70}╗") print(f"╔{'═'*70}╗")
print(f"║{'TRAJECTORY COMPRESSION REPORT':^70}║") print(f"║{'TRAJECTORY COMPRESSION REPORT':^70}║")
for title, pad, rows in sections:
section("📁 TRAJECTORIES", 54) print(f"╠{'═'*70}╣")
print(f"║{'':4}Total Processed: {total:>10,}{' '*32}║") print(f"║{'':2}{title}{' '*pad}║")
print(f"║{'':4}├─ Compressed: {compressed:>10,} ({pct(compressed):>5.1f}%){' '*18}║") print(f"║{'─'*70}║")
print(f"║{'':4}├─ Skipped (under limit):{skipped:>9,} ({pct(skipped):>5.1f}%){' '*18}║") for row in rows:
print(f"║{'':4}├─ Still over limit: {over_limit:>10,} ({pct(over_limit):>5.1f}%){' '*18}║") print(row)
print(f"║{'':4}└─ Failed: {m['summary']['trajectories_failed']:>10,}{' '*32}║")
section("🔢 TOKENS", 60)
print(f"║{'':4}Before Compression: {tokens_before:>15,} tokens{' '*21}║")
print(f"║{'':4}After Compression: {m['tokens']['total_after']:>15,} tokens{' '*21}║")
print(f"║{'':4}Total Saved: {tokens_saved:>15,} tokens{' '*21}║")
print(f"║{'':4}Overall Compression: {m['tokens']['overall_compression_ratio']:>14.1%}{' '*28}║")
if tokens_before > 0:
print(f"║{'':4}Space Savings: {(tokens_saved / tokens_before) * 100:>14.1f}%{' '*28}║")
section("💬 CONVERSATION TURNS", 48)
print(f"║{'':4}Before Compression: {m['turns']['total_before']:>15,} turns{' '*22}║")
print(f"║{'':4}After Compression: {m['turns']['total_after']:>15,} turns{' '*22}║")
print(f"║{'':4}Total Removed: {m['turns']['total_removed']:>15,} turns{' '*22}║")
section("📈 AVERAGES (Compressed Trajectories Only)", 27)
if compressed > 0:
print(f"║{'':4}Avg Compression Ratio: {m['averages']['avg_compression_ratio']:>14.1%}{' '*28}║")
print(f"║{'':4}Avg Tokens Saved: {m['averages']['avg_tokens_saved_per_compressed']:>14,.0f}{' '*28}║")
print(f"║{'':4}Avg Turns Removed: {m['averages']['avg_turns_removed_per_compressed']:>14.1f}{' '*28}║")
else:
print(f"║{'':4}No trajectories were compressed{' '*38}║")
section("🤖 SUMMARIZATION API", 49)
print(f"║{'':4}API Calls Made: {m['summarization']['total_api_calls']:>15,}{' '*27}║")
print(f"║{'':4}Errors: {m['summarization']['total_errors']:>15,}{' '*27}║")
print(f"║{'':4}Success Rate: {m['summarization']['success_rate']:>14.1%}{' '*28}║")
duration = m['processing']['duration_seconds']
time_str = f"{duration/60:.1f} minutes" if duration > 60 else f"{duration:.1f} seconds"
section("⏱️ PROCESSING TIME", 51)
print(f"║{'':4}Duration: {time_str:>20}{' '*22}║")
print(f"║{'':4}Throughput: {total / max(duration, 0.001):>15.1f} traj/sec{' '*18}║")
print(f"║{'':4}Started: {m['processing']['start_time'][:19]:>20}{' '*22}║")
print(f"║{'':4}Finished: {m['processing']['end_time'][:19]:>20}{' '*22}║")
print(f"╚{'═'*70}╝") print(f"╚{'═'*70}╝")
ratios = self.aggregate_metrics.compression_ratios ratios = self.aggregate_metrics.compression_ratios
@@ -874,45 +761,40 @@ def _print_dry_run(icon: str, target: Any, output_path: Path) -> None:
print(f"{icon} Would output to: {output_path}") print(f"{icon} Would output to: {output_path}")
def _sample(entries: list, sample_percent: float) -> list:
return random.sample(entries, min(max(1, int(len(entries) * sample_percent / 100)), len(entries)))
def _run_file_mode(input_path: Path, output: Optional[str], compression_config: CompressionConfig, sample_percent: Optional[float], seed: int, dry_run: bool) -> None: def _run_file_mode(input_path: Path, output: Optional[str], compression_config: CompressionConfig, sample_percent: Optional[float], seed: int, dry_run: bool) -> None:
"""Single-file input: (sample,) compress via a temp directory, merge into one output file.""" """Single-file input: (sample,) compress via a temp directory, merge into one output file."""
print("📄 Input mode: Single JSONL file") print("📄 Input mode: Single JSONL file")
output_path = Path(output) if output else input_path.parent / (input_path.stem + compression_config.output_suffix + ".jsonl") output_path = Path(output) if output else input_path.parent / (input_path.stem + compression_config.output_suffix + ".jsonl")
entries = [entry for _, entry in _load_jsonl(input_path, lambda n, e: print(f"⚠️ Skipping invalid JSON at line {n}: {e}"), start=1)]
entries = [entry for _, entry in _load_jsonl(
input_path, lambda n, e: print(f"⚠️ Skipping invalid JSON at line {n}: {e}"), start=1)]
total_entries = len(entries) total_entries = len(entries)
print(f" Loaded {total_entries:,} trajectories from {input_path.name}") print(f" Loaded {total_entries:,} trajectories from {input_path.name}")
if sample_percent is not None: if sample_percent is not None:
random.seed(seed) random.seed(seed)
entries = random.sample(entries, max(1, int(total_entries * sample_percent / 100))) entries = random.sample(entries, max(1, int(total_entries * sample_percent / 100)))
print(f" Sampled {len(entries):,} trajectories ({sample_percent}% of {total_entries:,})") print(f" Sampled {len(entries):,} trajectories ({sample_percent}% of {total_entries:,})")
if dry_run: if dry_run:
_print_dry_run("📄", f"{len(entries):,} trajectories", output_path) _print_dry_run("📄", f"{len(entries):,} trajectories", output_path)
return return
with tempfile.TemporaryDirectory() as temp_dir: with tempfile.TemporaryDirectory() as temp_dir:
temp_input_dir = Path(temp_dir) / "input" temp_input_dir, temp_output_dir = Path(temp_dir) / "input", Path(temp_dir) / "output"
temp_output_dir = Path(temp_dir) / "output"
temp_input_dir.mkdir() temp_input_dir.mkdir()
_write_jsonl(temp_input_dir / "trajectories.jsonl", entries) _write_jsonl(temp_input_dir / "trajectories.jsonl", entries)
TrajectoryCompressor(compression_config).process_directory(temp_input_dir, temp_output_dir) TrajectoryCompressor(compression_config).process_directory(temp_input_dir, temp_output_dir)
output_path.parent.mkdir(parents=True, exist_ok=True) output_path.parent.mkdir(parents=True, exist_ok=True)
with open(output_path, 'w', encoding='utf-8') as out_f: with open(output_path, 'w', encoding='utf-8') as out_f:
for jsonl_file in sorted(temp_output_dir.glob("*.jsonl")): for jsonl_file in sorted(temp_output_dir.glob("*.jsonl")):
with open(jsonl_file, 'r', encoding='utf-8') as in_f: with open(jsonl_file, 'r', encoding='utf-8') as in_f:
shutil.copyfileobj(in_f, out_f) shutil.copyfileobj(in_f, out_f)
metrics_file = temp_output_dir / compression_config.metrics_output_file metrics_file = temp_output_dir / compression_config.metrics_output_file
if metrics_file.exists(): if metrics_file.exists():
metrics_output = output_path.parent / (output_path.stem + "_metrics.json") metrics_output = output_path.parent / (output_path.stem + "_metrics.json")
shutil.copy(metrics_file, metrics_output) shutil.copy(metrics_file, metrics_output)
print(f"💾 Metrics saved to {metrics_output}") print(f"💾 Metrics saved to {metrics_output}")
print("\n✅ Compression complete!") print("\n✅ Compression complete!")
print(f"📄 Output: {output_path}") print(f"📄 Output: {output_path}")
@@ -921,7 +803,6 @@ def _run_dir_mode(input_path: Path, output: Optional[str], compression_config: C
"""Directory input: compress in place, or per-file sample into a temp dir first.""" """Directory input: compress in place, or per-file sample into a temp dir first."""
print("📁 Input mode: Directory of JSONL files") print("📁 Input mode: Directory of JSONL files")
output_path = Path(output) if output else input_path.parent / (input_path.name + compression_config.output_suffix) output_path = Path(output) if output else input_path.parent / (input_path.name + compression_config.output_suffix)
if sample_percent is None: if sample_percent is None:
if dry_run: if dry_run:
_print_dry_run("📁", input_path, output_path) _print_dry_run("📁", input_path, output_path)
@@ -936,30 +817,20 @@ def _run_dir_mode(input_path: Path, output: Optional[str], compression_config: C
total_original = total_sampled = 0 total_original = total_sampled = 0
for jsonl_file in sorted(input_path.glob("*.jsonl")): for jsonl_file in sorted(input_path.glob("*.jsonl")):
entries = [entry for _, entry in _load_jsonl(jsonl_file)] entries = [entry for _, entry in _load_jsonl(jsonl_file)]
sampled_entries = random.sample(entries, min(max(1, int(len(entries) * sample_percent / 100)), len(entries))) sampled_entries = _sample(entries, sample_percent)
total_original += len(entries) total_original += len(entries)
total_sampled += len(sampled_entries) total_sampled += len(sampled_entries)
_write_jsonl(temp_input_dir / jsonl_file.name, sampled_entries) _write_jsonl(temp_input_dir / jsonl_file.name, sampled_entries)
print(f" Sampled {total_sampled:,} from {total_original:,} total trajectories") print(f" Sampled {total_sampled:,} from {total_original:,} total trajectories")
if dry_run: if dry_run:
_print_dry_run("📁", temp_input_dir, output_path) _print_dry_run("📁", temp_input_dir, output_path)
return return
TrajectoryCompressor(compression_config).process_directory(temp_input_dir, output_path) TrajectoryCompressor(compression_config).process_directory(temp_input_dir, output_path)
print("\n✅ Compression complete!") print("\n✅ Compression complete!")
def main( def main(input: str, output: str = None, config: str = "configs/trajectory_compression.yaml", target_max_tokens: int = None,
input: str, tokenizer: str = None, sample_percent: float = None, seed: int = 42, dry_run: bool = False):
output: str = None,
config: str = "configs/trajectory_compression.yaml",
target_max_tokens: int = None,
tokenizer: str = None,
sample_percent: float = None,
seed: int = 42,
dry_run: bool = False,
):
""" """
Compress agent trajectories to fit within a target token budget. Compress agent trajectories to fit within a target token budget.
@@ -979,20 +850,16 @@ def main(
""" """
print("🗜️ Trajectory Compressor") print("🗜️ Trajectory Compressor")
print("=" * 60) print("=" * 60)
compression_config = _load_cli_config(config, target_max_tokens, tokenizer) compression_config = _load_cli_config(config, target_max_tokens, tokenizer)
if sample_percent is not None: if sample_percent is not None:
if sample_percent <= 0 or sample_percent > 100: if sample_percent <= 0 or sample_percent > 100:
print(f"❌ sample_percent must be between 1 and 100, got {sample_percent}") print(f"❌ sample_percent must be between 1 and 100, got {sample_percent}")
return return
print(f"🎲 Will sample {sample_percent}% of trajectories (seed={seed})") print(f"🎲 Will sample {sample_percent}% of trajectories (seed={seed})")
input_path = Path(input) input_path = Path(input)
if not input_path.exists(): if not input_path.exists():
print(f"❌ Input not found: {input}") print(f"❌ Input not found: {input}")
return return
run_mode = _run_file_mode if input_path.is_file() else _run_dir_mode run_mode = _run_file_mode if input_path.is_file() else _run_dir_mode
run_mode(input_path, output, compression_config, sample_percent, seed, dry_run) run_mode(input_path, output, compression_config, sample_percent, seed, dry_run)