File size: 10,502 Bytes
9c7d451
 
 
 
 
 
175746c
9c7d451
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
78d41df
 
9c7d451
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c8388c1
9c7d451
c8388c1
257873b
 
 
9c7d451
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
142a6dd
 
 
 
9c7d451
 
 
 
 
 
 
 
 
 
 
142a6dd
9c7d451
 
 
 
 
142a6dd
 
 
 
 
 
 
 
 
 
 
9c7d451
142a6dd
 
 
 
9c7d451
142a6dd
 
 
9c7d451
 
 
 
 
142a6dd
 
 
 
 
 
 
 
 
 
 
 
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
316
317
318
319
320
321
322
323
324
325
326
"""Message parsing utilities for Headroom SDK."""

from __future__ import annotations

import hashlib
import re
from typing import TYPE_CHECKING, Any

from .config import Block, WasteSignals

if TYPE_CHECKING:
    from .tokenizer import Tokenizer


# Patterns for detecting waste signals
HTML_TAG_PATTERN = re.compile(r"<[^>]+>")
HTML_COMMENT_PATTERN = re.compile(r"<!--[\s\S]*?-->")
BASE64_PATTERN = re.compile(r"[A-Za-z0-9+/]{50,}={0,2}")
WHITESPACE_PATTERN = re.compile(r"[ \t]{4,}|\n{3,}")
JSON_BLOCK_PATTERN = re.compile(r"\{[\s\S]{500,}\}")

# Patterns for RAG detection (best effort)
RAG_MARKERS = [
    r"\[Document\s*\d+\]",
    r"\[Source:\s*",
    r"<context>",
    r"<document>",
    r"Retrieved from:",
    r"From the knowledge base:",
]
RAG_PATTERN = re.compile("|".join(RAG_MARKERS), re.IGNORECASE)


def compute_hash(text: str) -> str:
    """Compute hash of text, truncated to 16 chars."""
    return hashlib.md5(text.encode()).hexdigest()[:16]


def detect_waste_signals(text: str, tokenizer: Tokenizer) -> WasteSignals:
    """
    Detect waste signals in text.

    Args:
        text: The text to analyze.
        tokenizer: Tokenizer for counting tokens.

    Returns:
        WasteSignals with detected waste.
    """
    signals = WasteSignals()

    if not text:
        return signals

    # HTML tags and comments
    html_matches = HTML_TAG_PATTERN.findall(text) + HTML_COMMENT_PATTERN.findall(text)
    if html_matches:
        html_text = "".join(html_matches)
        signals.html_noise_tokens = tokenizer.count_text(html_text)

    # Base64 blobs
    base64_matches = BASE64_PATTERN.findall(text)
    if base64_matches:
        base64_text = "".join(base64_matches)
        signals.base64_tokens = tokenizer.count_text(base64_text)

    # Excessive whitespace
    ws_matches = WHITESPACE_PATTERN.findall(text)
    if ws_matches:
        # Count tokens that could be saved by normalizing whitespace to single spaces
        ws_text = "".join(ws_matches)
        normalized_text = " ".join(ws_matches)
        signals.whitespace_tokens = max(
            0, tokenizer.count_text(ws_text) - tokenizer.count_text(normalized_text)
        )

    # Large JSON blocks
    json_matches = JSON_BLOCK_PATTERN.findall(text)
    if json_matches:
        for match in json_matches:
            tokens = tokenizer.count_text(match)
            if tokens > 500:
                signals.json_bloat_tokens += tokens

    return signals


def is_rag_content(text: str) -> bool:
    """Check if text appears to be RAG-injected content."""
    return RAG_PATTERN.search(text) is not None


def parse_message_to_blocks(
    message: dict[str, Any],
    index: int,
    tokenizer: Tokenizer,
) -> list[Block]:
    """
    Parse a single message into Block objects.

    Args:
        message: The message dict to parse.
        index: Position in the message list.
        tokenizer: Tokenizer for token counting.

    Returns:
        List of Block objects (usually 1, but tool_calls may produce multiple).
    """
    blocks: list[Block] = []
    role = message.get("role", "unknown")

    # Handle content
    content = message.get("content")
    if content:
        if isinstance(content, str):
            text = content
        elif isinstance(content, list):
            # Multi-modal - extract text parts
            text_parts = []
            for part in content:
                if isinstance(part, dict) and part.get("type") == "text":
                    text_parts.append(part.get("text", ""))
                elif isinstance(part, str):
                    text_parts.append(part)
            text = "\n".join(text_parts)
        else:
            text = str(content)

        # Determine block kind
        if role == "system":
            kind = "system"
        elif role == "user":
            # Check if this looks like RAG content
            kind = "rag" if is_rag_content(text) else "user"
        elif role == "assistant":
            kind = "assistant"
        elif role == "tool":
            kind = "tool_result"
        else:
            kind = "unknown"

        # Build flags
        flags: dict[str, Any] = {}
        if role == "tool":
            flags["tool_call_id"] = message.get("tool_call_id")

        # Detect waste
        waste = detect_waste_signals(text, tokenizer)
        if waste.total() > 0:
            flags["waste_signals"] = waste.to_dict()

        blocks.append(
            Block(
                kind=kind,  # type: ignore[arg-type]
                text=text,
                tokens_est=tokenizer.count_text(text) + 4,  # Add message overhead
                content_hash=compute_hash(text),
                source_index=index,
                flags=flags,
            )
        )

    # Handle tool calls (assistant messages with tool_calls)
    tool_calls = message.get("tool_calls")
    if tool_calls:
        for tc in tool_calls:
            func = tc.get("function", {})
            tc_text = f"{func.get('name', 'unknown')}({func.get('arguments', '')})"

            blocks.append(
                Block(
                    kind="tool_call",
                    text=tc_text,
                    tokens_est=tokenizer.count_text(tc_text) + 10,
                    content_hash=compute_hash(tc_text),
                    source_index=index,
                    flags={
                        "tool_call_id": tc.get("id"),
                        "function_name": func.get("name"),
                    },
                )
            )

    # If no content or tool_calls, create a minimal block
    if not blocks:
        blocks.append(
            Block(
                kind="unknown",
                text="",
                tokens_est=4,
                content_hash=compute_hash(""),
                source_index=index,
                flags={},
            )
        )

    return blocks


def parse_messages(
    messages: list[dict[str, Any]],
    tokenizer: Tokenizer,
) -> tuple[list[Block], dict[str, int], WasteSignals]:
    """
    Parse all messages into blocks with analysis.

    Args:
        messages: List of message dicts.
        tokenizer: Tokenizer instance for token counting.

    Returns:
        Tuple of (blocks, block_breakdown, total_waste_signals)
    """
    all_blocks: list[Block] = []
    total_waste = WasteSignals()

    for i, msg in enumerate(messages):
        blocks = parse_message_to_blocks(msg, i, tokenizer)
        all_blocks.extend(blocks)

        # Accumulate waste signals
        for block in blocks:
            if "waste_signals" in block.flags:
                ws = block.flags["waste_signals"]
                total_waste.json_bloat_tokens += ws.get("json_bloat", 0)
                total_waste.html_noise_tokens += ws.get("html_noise", 0)
                total_waste.base64_tokens += ws.get("base64", 0)
                total_waste.whitespace_tokens += ws.get("whitespace", 0)
                total_waste.dynamic_date_tokens += ws.get("dynamic_date", 0)
                total_waste.repetition_tokens += ws.get("repetition", 0)

    # Compute block breakdown
    breakdown: dict[str, int] = {}
    for block in all_blocks:
        kind = block.kind
        breakdown[kind] = breakdown.get(kind, 0) + block.tokens_est

    return all_blocks, breakdown, total_waste


def find_tool_units(messages: list[dict[str, Any]]) -> list[tuple[int, list[int]]]:
    """
    Find tool call units (assistant with tool_calls + corresponding tool responses).

    A tool unit is atomic - if the assistant message is dropped, all its
    tool responses must also be dropped.

    Supports both OpenAI and Anthropic formats:
    - OpenAI: assistant.tool_calls[] + tool messages with tool_call_id
    - Anthropic: assistant.content[type=tool_use] + user.content[type=tool_result]

    Args:
        messages: List of message dicts.

    Returns:
        List of (assistant_index, [tool_response_indices]) tuples.
    """
    units: list[tuple[int, list[int]]] = []

    # Build map of tool_call_id -> message index for tool responses
    tool_response_map: dict[str, int] = {}
    for i, msg in enumerate(messages):
        # OpenAI format: role="tool" with tool_call_id
        if msg.get("role") == "tool":
            tc_id = msg.get("tool_call_id")
            if tc_id:
                tool_response_map[tc_id] = i

        # Anthropic format: role="user" with content blocks containing tool_result
        if msg.get("role") == "user":
            content = msg.get("content")
            if isinstance(content, list):
                for block in content:
                    if isinstance(block, dict) and block.get("type") == "tool_result":
                        tc_id = block.get("tool_use_id")
                        if tc_id:
                            tool_response_map[tc_id] = i

    # Find assistant messages with tool calls
    for i, msg in enumerate(messages):
        if msg.get("role") != "assistant":
            continue

        response_indices: list[int] = []

        # OpenAI format: tool_calls array
        tool_calls = msg.get("tool_calls")
        if tool_calls:
            for tc in tool_calls:
                tc_id = tc.get("id")
                if tc_id and tc_id in tool_response_map:
                    response_indices.append(tool_response_map[tc_id])

        # Anthropic format: content blocks with type=tool_use
        content = msg.get("content")
        if isinstance(content, list):
            for block in content:
                if isinstance(block, dict) and block.get("type") == "tool_use":
                    tc_id = block.get("id")
                    if tc_id and tc_id in tool_response_map:
                        response_indices.append(tool_response_map[tc_id])

        if response_indices:
            # Use set to deduplicate in case same message has both formats
            units.append((i, sorted(set(response_indices))))

    return units


def get_message_content_text(message: dict[str, Any]) -> str:
    """Extract text content from a message."""
    content = message.get("content")
    if content is None:
        return ""
    if isinstance(content, str):
        return content
    if isinstance(content, list):
        parts = []
        for part in content:
            if isinstance(part, dict) and part.get("type") == "text":
                parts.append(part.get("text", ""))
            elif isinstance(part, str):
                parts.append(part)
        return "\n".join(parts)
    return str(content)