File size: 11,387 Bytes
9c7d451
 
 
 
c1feb60
175746c
9c7d451
 
 
 
 
bd2d447
9c7d451
 
 
 
 
 
 
 
 
bd2d447
9c7d451
 
 
 
 
 
 
e4a41fa
 
9c7d451
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e4a41fa
9c7d451
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bd2d447
 
 
 
 
 
 
 
 
 
 
9c7d451
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c1feb60
 
 
 
 
 
 
9c7d451
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c1feb60
 
 
 
 
 
 
 
 
 
 
 
9c7d451
 
e4a41fa
 
 
 
 
 
 
 
 
 
 
9c7d451
 
 
 
c1feb60
 
 
 
 
 
 
 
 
 
 
 
 
9c7d451
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bd2d447
9c7d451
 
 
 
 
 
 
 
bd2d447
 
 
9c7d451
 
 
 
 
 
 
 
 
 
 
 
bd2d447
 
9c7d451
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
"""Transform pipeline orchestration for Headroom SDK."""

from __future__ import annotations

import logging
from typing import TYPE_CHECKING, Any

from ..config import (
    CacheAlignerConfig,
    DiffArtifact,
    HeadroomConfig,
    IntelligentContextConfig,
    RollingWindowConfig,
    ToolCrusherConfig,
    TransformDiff,
    TransformResult,
)
from ..tokenizer import Tokenizer
from ..utils import deep_copy_messages
from .base import Transform
from .cache_aligner import CacheAligner
from .intelligent_context import IntelligentContextManager
from .rolling_window import RollingWindow
from .smart_crusher import SmartCrusher
from .tool_crusher import ToolCrusher

if TYPE_CHECKING:
    from ..providers.base import Provider

logger = logging.getLogger(__name__)


class TransformPipeline:
    """
    Orchestrates multiple transforms in the correct order.

    Transform order:
    1. Cache Aligner - normalize prefix for cache hits
    2. Tool Crusher - compress tool outputs
    3. Rolling Window - enforce token limits
    """

    def __init__(
        self,
        config: HeadroomConfig | None = None,
        transforms: list[Transform] | None = None,
        provider: Provider | None = None,
    ):
        """
        Initialize pipeline.

        Args:
            config: Headroom configuration.
            transforms: Optional custom transform list (overrides config).
            provider: Provider for model-specific behavior.
        """
        self.config = config or HeadroomConfig()
        self._provider = provider

        if transforms is not None:
            self.transforms = transforms
        else:
            self.transforms = self._build_default_transforms()

    def _build_default_transforms(self) -> list[Transform]:
        """Build default transform pipeline from config."""
        transforms: list[Transform] = []

        # Order matters!

        # 1. Cache Aligner (prefix stabilization)
        if self.config.cache_aligner.enabled:
            transforms.append(CacheAligner(self.config.cache_aligner))

        # 2. Tool Output Compression
        # SmartCrusher (statistical) takes precedence over ToolCrusher (fixed rules)
        if self.config.smart_crusher.enabled:
            # Use smart statistical crushing
            from .smart_crusher import SmartCrusherConfig as SCConfig

            smart_config = SCConfig(
                enabled=True,
                min_items_to_analyze=self.config.smart_crusher.min_items_to_analyze,
                min_tokens_to_crush=self.config.smart_crusher.min_tokens_to_crush,
                variance_threshold=self.config.smart_crusher.variance_threshold,
                uniqueness_threshold=self.config.smart_crusher.uniqueness_threshold,
                similarity_threshold=self.config.smart_crusher.similarity_threshold,
                max_items_after_crush=self.config.smart_crusher.max_items_after_crush,
                preserve_change_points=self.config.smart_crusher.preserve_change_points,
                factor_out_constants=self.config.smart_crusher.factor_out_constants,
                include_summaries=self.config.smart_crusher.include_summaries,
            )
            transforms.append(SmartCrusher(smart_config))
        elif self.config.tool_crusher.enabled:
            # Fallback to fixed-rule crushing
            transforms.append(ToolCrusher(self.config.tool_crusher))

        # 3. Context Management (enforce limits last)
        # IntelligentContextManager takes precedence over RollingWindow when enabled
        if self.config.intelligent_context.enabled:
            # Use semantic-aware context management with scoring
            transforms.append(IntelligentContextManager(self.config.intelligent_context))
            logger.info(
                "Pipeline using IntelligentContextManager with strategies: "
                "COMPRESS_FIRST -> SUMMARIZE -> DROP_BY_SCORE"
            )
        elif self.config.rolling_window.enabled:
            # Fallback to position-based rolling window
            transforms.append(RollingWindow(self.config.rolling_window))

        return transforms

    def _get_tokenizer(self, model: str) -> Tokenizer:
        """Get tokenizer for model using provider."""
        if self._provider is None:
            raise ValueError(
                "Provider is required for token counting. "
                "Pass a provider to TransformPipeline or HeadroomClient."
            )
        token_counter = self._provider.get_token_counter(model)
        return Tokenizer(token_counter, model)

    def apply(
        self,
        messages: list[dict[str, Any]],
        model: str,
        **kwargs: Any,
    ) -> TransformResult:
        """
        Apply all transforms in sequence.

        Args:
            messages: List of messages to transform.
            model: Model name for token counting.
            **kwargs: Additional arguments passed to transforms.
                - model_limit: Context limit override.
                - output_buffer: Output buffer override.
                - tool_profiles: Per-tool compression profiles.
                - request_id: Optional request ID for diff artifact.

        Returns:
            Combined TransformResult.
        """
        tokenizer = self._get_tokenizer(model)

        # Get model limit from kwargs (should be set by client)
        model_limit = kwargs.get("model_limit")
        if model_limit is None:
            raise ValueError(
                "model_limit is required. Provide it via kwargs or "
                "configure model_context_limits in HeadroomClient."
            )

        # Start with original tokens
        tokens_before = tokenizer.count_messages(messages)

        logger.debug(
            "Pipeline starting: %d messages, %d tokens, model=%s",
            len(messages),
            tokens_before,
            model,
        )

        # Track all transforms applied
        all_transforms: list[str] = []
        all_markers: list[str] = []
        all_warnings: list[str] = []

        # Track transform diffs if enabled
        transform_diffs: list[TransformDiff] = []
        generate_diff = self.config.generate_diff_artifact

        current_messages = deep_copy_messages(messages)

        for transform in self.transforms:
            # Check if transform should run
            if not transform.should_apply(current_messages, tokenizer, **kwargs):
                continue

            # Track tokens before this transform (for diff)
            tokens_before_transform = tokenizer.count_messages(current_messages)

            # Apply transform
            result = transform.apply(current_messages, tokenizer, **kwargs)

            # Update messages for next transform
            current_messages = result.messages

            # Track tokens after this transform (for diff)
            tokens_after_transform = tokenizer.count_messages(current_messages)

            # Accumulate results
            all_transforms.extend(result.transforms_applied)
            all_markers.extend(result.markers_inserted)
            all_warnings.extend(result.warnings)

            # Log transform results
            if result.transforms_applied:
                logger.info(
                    "Transform %s: %d -> %d tokens (saved %d)",
                    transform.name,
                    tokens_before_transform,
                    tokens_after_transform,
                    tokens_before_transform - tokens_after_transform,
                )
            else:
                logger.debug("Transform %s: no changes", transform.name)

            # Record diff if enabled
            if generate_diff:
                transform_diffs.append(
                    TransformDiff(
                        transform_name=transform.name,
                        tokens_before=tokens_before_transform,
                        tokens_after=tokens_after_transform,
                        tokens_saved=tokens_before_transform - tokens_after_transform,
                        details=", ".join(result.transforms_applied)
                        if result.transforms_applied
                        else "",
                    )
                )

        # Final token count
        tokens_after = tokenizer.count_messages(current_messages)

        # Log pipeline summary
        total_saved = tokens_before - tokens_after
        if total_saved > 0:
            logger.info(
                "Pipeline complete: %d -> %d tokens (saved %d, %.1f%% reduction)",
                tokens_before,
                tokens_after,
                total_saved,
                (total_saved / tokens_before * 100) if tokens_before > 0 else 0,
            )
        else:
            logger.debug("Pipeline complete: no token savings")

        # Build diff artifact if enabled
        diff_artifact = None
        if generate_diff:
            diff_artifact = DiffArtifact(
                request_id=kwargs.get("request_id", ""),
                original_tokens=tokens_before,
                optimized_tokens=tokens_after,
                total_tokens_saved=tokens_before - tokens_after,
                transforms=transform_diffs,
            )

        return TransformResult(
            messages=current_messages,
            tokens_before=tokens_before,
            tokens_after=tokens_after,
            transforms_applied=all_transforms,
            markers_inserted=all_markers,
            warnings=all_warnings,
            diff_artifact=diff_artifact,
        )

    def simulate(
        self,
        messages: list[dict[str, Any]],
        model: str,
        **kwargs: Any,
    ) -> TransformResult:
        """
        Simulate transforms without modifying messages.

        Same as apply() but returns what WOULD happen.

        Args:
            messages: List of messages.
            model: Model name.
            **kwargs: Additional arguments.

        Returns:
            TransformResult with simulated changes.
        """
        # apply() already works on a copy, so this is safe
        return self.apply(messages, model, **kwargs)


def create_pipeline(
    tool_crusher_config: ToolCrusherConfig | None = None,
    cache_aligner_config: CacheAlignerConfig | None = None,
    rolling_window_config: RollingWindowConfig | None = None,
    intelligent_context_config: IntelligentContextConfig | None = None,
) -> TransformPipeline:
    """
    Create a pipeline with specific configurations.

    Args:
        tool_crusher_config: Tool crusher configuration.
        cache_aligner_config: Cache aligner configuration.
        rolling_window_config: Rolling window configuration.
        intelligent_context_config: Intelligent context configuration.
            When provided with enabled=True, replaces RollingWindow with
            semantic-aware context management.

    Returns:
        Configured TransformPipeline.
    """
    config = HeadroomConfig()

    if tool_crusher_config is not None:
        config.tool_crusher = tool_crusher_config
    if cache_aligner_config is not None:
        config.cache_aligner = cache_aligner_config
    if rolling_window_config is not None:
        config.rolling_window = rolling_window_config
    if intelligent_context_config is not None:
        config.intelligent_context = intelligent_context_config

    return TransformPipeline(config)